update TeaCache

This commit is contained in:
hkunzhe
2025-02-10 19:52:29 +08:00
parent ad83593912
commit 10478729cd
6 changed files with 41 additions and 16 deletions
+16 -4
View File
@@ -121,6 +121,22 @@ class TeaCache():
self.previous_residual = None
def get_teacache_coefficients(model_name):
# The coefficients for EasyAnimateV5-7b-zh-InP should be:
# [-3.64204720e+03, 1.43764725e+03, -1.93045263e+02, 1.09596499e+01, -1.70663507e-01]
if "v5.1-7b" in model_name.lower():
# The coefficient was obtained by sampling videos from T2V CompBench using EasyAnimateV5.1-7b-zh-InP.
# This coefficient can be applied to both the EasyAnimateV5.1-7b-zh and EasyAnimateV5.1-7b-Control.
return [1.07862322, -4.19362456, 3.06725828, 0.33161686, 0.02374758]
elif "v5.1-12b" in model_name.lower():
# The coefficient was obtained by sampling videos from T2V CompBench using EasyAnimateV5.1-12b-zh-InP.
# This coefficient can be applied to both the EasyAnimateV5.1-12b-zh and EasyAnimateV5.1-12b-Control.
return [-10.47857366, 8.33844143, -0.78477557, 0.68798618, 0.0136149]
else:
print(f"The model {model_name} is not supported by TeaCache.")
return None
class Transformer3DModel(ModelMixin, ConfigMixin):
"""
A 3D Transformer model for image-like data.
@@ -1472,10 +1488,6 @@ class EasyAnimateTransformer3DModel(ModelMixin, ConfigMixin):
rel_l1_thresh: float,
coefficients: list[float] = [-10.47857366, 8.33844143, -0.78477557, 0.68798618, 0.0136149]
):
# The coefficient was obtained by sampling videos from T2V CompBench using EasyAnimateV5.1-12b-zh-InP.
# This coefficient can be applied to both the EasyAnimateV5.1-12b-zh and EasyAnimateV5.1-12b-Control.
# The coefficients for EasyAnimateV5.1-7b-zh-InP should be:
# [-3.64204720e+03, 1.43764725e+03, -1.93045263e+02, 1.09596499e+01, -1.70663507e-01]
self.teacache = TeaCache(coefficients, num_steps, rel_l1_thresh=rel_l1_thresh)
def _set_gradient_checkpointing(self, module, value=False):
+9 -4
View File
@@ -30,6 +30,7 @@ from transformers import (BertModel, BertTokenizer, CLIPImageProcessor,
from ..data.bucket_sampler import ASPECT_RATIO_512, get_closest_ratio
from ..models import (name_to_autoencoder_magvit,
name_to_transformer3d)
from ..models.transformer3d import get_teacache_coefficients
from ..pipeline.pipeline_easyanimate import \
EasyAnimatePipeline
from ..pipeline.pipeline_easyanimate_control import \
@@ -467,8 +468,10 @@ class EasyAnimateController:
# lora part
self.pipeline = merge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider)
if self.edition == "v5.1" and self.enable_teacache:
self.pipeline.transformer.enable_teacache(sample_step_slider, self.teacache_threshold)
coefficients = get_teacache_coefficients(self.base_model_path)
if coefficients is not None and self.enable_teacache:
print(f"Enable TeaCache with threshold: {self.teacache_threshold}.")
self.pipeline.transformer.enable_teacache(sample_step_slider, self.teacache_threshold, coefficients=coefficients)
try:
if self.model_type == "Inpaint":
@@ -1020,6 +1023,7 @@ class EasyAnimateController_Modelscope:
# Config and model path
self.model_type = model_type
self.edition = edition
self.model_name = model_name
self.enable_teacache = enable_teacache
self.teacache_threshold = teacache_threshold
self.weight_dtype = weight_dtype
@@ -1272,9 +1276,10 @@ class EasyAnimateController_Modelscope:
# lora part
self.pipeline = merge_lora(self.pipeline, self.lora_model_path, multiplier=lora_alpha_slider)
if self.edition == "v5.1" and self.enable_teacache:
coefficients = get_teacache_coefficients(self.model_name)
if coefficients is not None and self.enable_teacache:
print(f"Enable TeaCache with threshold: {self.teacache_threshold}.")
self.pipeline.transformer.enable_teacache(sample_step_slider, self.teacache_threshold)
self.pipeline.transformer.enable_teacache(sample_step_slider, self.teacache_threshold, coefficients=coefficients)
try:
if self.model_type == "Inpaint":
+4 -2
View File
@@ -14,6 +14,7 @@ from transformers import (BertModel, BertTokenizer, CLIPImageProcessor,
from easyanimate.models import (name_to_autoencoder_magvit,
name_to_transformer3d)
from easyanimate.models.transformer3d import get_teacache_coefficients
from easyanimate.pipeline.pipeline_easyanimate_inpaint import \
EasyAnimateInpaintPipeline
from easyanimate.utils.fp8_optimization import convert_weight_dtype_wrapper
@@ -256,9 +257,10 @@ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
else:
pipeline.enable_model_cpu_offload()
if "v5.1" in config_path and enable_teacache:
coefficients = get_teacache_coefficients(model_name)
if coefficients is not None and enable_teacache:
print(f"Enable TeaCache with threshold: {teacache_threshold}.")
pipeline.transformer.enable_teacache(num_inference_steps, teacache_threshold)
pipeline.transformer.enable_teacache(num_inference_steps, teacache_threshold, coefficients=coefficients)
generator = torch.Generator(device="cuda").manual_seed(seed)
+4 -2
View File
@@ -14,6 +14,7 @@ from transformers import (BertModel, BertTokenizer,
from easyanimate.models import (name_to_autoencoder_magvit,
name_to_transformer3d)
from easyanimate.models.transformer3d import get_teacache_coefficients
from easyanimate.pipeline.pipeline_easyanimate import \
EasyAnimatePipeline
from easyanimate.pipeline.pipeline_easyanimate_inpaint import \
@@ -264,9 +265,10 @@ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
else:
pipeline.enable_model_cpu_offload()
if "v5.1" in config_path and enable_teacache:
coefficients = get_teacache_coefficients(model_name)
if coefficients is not None and enable_teacache:
print(f"Enable TeaCache with threshold: {teacache_threshold}.")
pipeline.transformer.enable_teacache(num_inference_steps, teacache_threshold)
pipeline.transformer.enable_teacache(num_inference_steps, teacache_threshold, coefficients=coefficients)
generator = torch.Generator(device="cuda").manual_seed(seed)
+4 -2
View File
@@ -14,6 +14,7 @@ from transformers import (BertModel, BertTokenizer, CLIPImageProcessor,
from easyanimate.models import (name_to_autoencoder_magvit,
name_to_transformer3d)
from easyanimate.models.transformer3d import get_teacache_coefficients
from easyanimate.pipeline.pipeline_easyanimate_inpaint import \
EasyAnimateInpaintPipeline
from easyanimate.utils.fp8_optimization import convert_weight_dtype_wrapper
@@ -251,9 +252,10 @@ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
else:
pipeline.enable_model_cpu_offload()
if "v5.1" in config_path and enable_teacache:
coefficients = get_teacache_coefficients(model_name)
if coefficients is not None and enable_teacache:
print(f"Enable TeaCache with threshold: {teacache_threshold}.")
pipeline.transformer.enable_teacache(num_inference_steps, teacache_threshold)
pipeline.transformer.enable_teacache(num_inference_steps, teacache_threshold, coefficients=coefficients)
generator = torch.Generator(device="cuda").manual_seed(seed)
+4 -2
View File
@@ -15,6 +15,7 @@ from transformers import (BertModel, BertTokenizer,
from easyanimate.data.dataset_image_video import process_pose_file
from easyanimate.models import (name_to_autoencoder_magvit,
name_to_transformer3d)
from easyanimate.models.transformer3d import get_teacache_coefficients
from easyanimate.pipeline.pipeline_easyanimate_control import \
EasyAnimateControlPipeline
from easyanimate.utils.lora_utils import merge_lora, unmerge_lora
@@ -236,9 +237,10 @@ elif GPU_memory_mode == "model_cpu_offload_and_qfloat8":
else:
pipeline.enable_model_cpu_offload()
if "v5.1" in config_path and enable_teacache:
coefficients = get_teacache_coefficients(model_name)
if coefficients is not None and enable_teacache:
print(f"Enable TeaCache with threshold: {teacache_threshold}.")
pipeline.transformer.enable_teacache(num_inference_steps, teacache_threshold)
pipeline.transformer.enable_teacache(num_inference_steps, teacache_threshold, coefficients=coefficients)
generator = torch.Generator(device="cuda").manual_seed(seed)