diff --git a/README.md b/README.md index 3ad28e4..e7afd7b 100644 --- a/README.md +++ b/README.md @@ -6,19 +6,32 @@ Timestep Embedding Aware Cache (TeaCache) is a training-free caching approach th TeaCache has now been integrated into ComfyUI and is compatible with the ComfyUI native nodes. ComfyUI-TeaCache is easy to use, simply connect the TeaCache node with the ComfyUI native nodes for seamless usage. ## Updates +- Jan 10 2025: ComfyUI-TeaCache supports LTX-Video: + - It can achieve a 1.4x lossless speedup and a 1.7x speedup without much visual quality degradation. + - Support Text to Video and Image to Video! - Jan 9 2025: ComfyUI-TeaCache supports HunyuanVideo: - - It can achieve a 1.6x lossless speedup and a 2x speedup without much visual quality degradation, which are consistent with the original [TeaCache4HunyuanVideo](https://github.com/ali-vilab/TeaCache/tree/main/TeaCache4HunyuanVideo). + - It can achieve a 1.6x lossless speedup and a 2x speedup without much visual quality degradation. - Jan 8 2025: ComfyUI-TeaCache supports FLUX: - - It can achieve a 1.4x lossless speedup and a 2x speedup without much visual quality degradation, which are consistent with the original [TeaCache4FLUX](https://github.com/ali-vilab/TeaCache/tree/main/TeaCache4FLUX). + - It can achieve a 1.4x lossless speedup and a 2x speedup without much visual quality degradation. - Support FLUX LoRA! - Support FLUX ControlNet! ## Installation +Installation via ComfyUI-Manager is preferred. Simply search for ComfyUI-TeaCache in the list of nodes and click install. +### Manual installation 1. Go to comfyUI custom_nodes folder, `ComfyUI/custom_nodes/` 2. git clone https://github.com/welltop-cn/ComfyUI-TeaCache.git +## Recommended settings +The following table gives the recommended rel_l1_thresh ​for different models: + +| | FLUX | HunyuanVideo | LTX-Video | +|:---------------------:|:----------------------------:|:---------------------:|:---------------------:| +| rel_l1_thresh | 0.4 | 0.15 | 0.06 | +| speedup | ~2x | ~2x | ~1.7x | + ## Usage -The demo workflow is placed in examples folder. +The demo workflows are placed in examples folder. ## Demo -

FLUX

@@ -27,6 +40,9 @@ https://github.com/user-attachments/assets/e977cf34-f7d0-4b25-a2e3-10fd62ebfe30 -

HunyuanVideo

https://github.com/user-attachments/assets/4d8e9f12-2c54-40c5-a992-c2cecbde019a +-

LTX-Video

+https://github.com/user-attachments/assets/19e63dd8-ecdf-418c-8ec2-b9b9dcf9a655 + ## Result comparison -

FLUX

![](./assets/compare_flux.png) @@ -34,9 +50,8 @@ https://github.com/user-attachments/assets/4d8e9f12-2c54-40c5-a992-c2cecbde019a -

HunyuanVideo

https://github.com/user-attachments/assets/b3aca64d-c2ae-440c-a362-f3a7b6c633e0 -## Roadmap - -- [ ] Support LTX-Video +-

LTX-Video

+https://github.com/user-attachments/assets/8fce9b48-2243-46f1-b411-80e4a53f6f7d ## Acknowledgments Thanks to TeaCache repo owner [ali-vilab/TeaCache: Timestep Embedding Tells: It's Time to Cache for Video Diffusion Model](https://github.com/ali-vilab/TeaCache) diff --git a/assets/teacache_ltx_video.png b/assets/teacache_ltx_video.png new file mode 100644 index 0000000..553d3e6 Binary files /dev/null and b/assets/teacache_ltx_video.png differ diff --git a/examples/teacache_ltx_video.json b/examples/teacache_ltx_video.json new file mode 100644 index 0000000..5ee7d68 --- /dev/null +++ b/examples/teacache_ltx_video.json @@ -0,0 +1,713 @@ +{ + "last_node_id": 88, + "last_link_id": 185, + "nodes": [ + { + "id": 71, + "type": "LTXVScheduler", + "pos": [ + 856, + 531 + ], + "size": [ + 315, + 154 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [ + { + "name": "latent", + "type": "LATENT", + "link": 168, + "shape": 7 + } + ], + "outputs": [ + { + "name": "SIGMAS", + "type": "SIGMAS", + "links": [ + 182 + ], + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "LTXVScheduler" + }, + "widgets_values": [ + 30, + 2.05, + 0.95, + true, + 0.1 + ] + }, + { + "id": 6, + "type": "CLIPTextEncode", + "pos": [ + 420, + 190 + ], + "size": [ + 422.84503173828125, + 164.31304931640625 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "clip", + "type": "CLIP", + "link": 74 + } + ], + "outputs": [ + { + "name": "CONDITIONING", + "type": "CONDITIONING", + "links": [ + 169 + ], + "slot_index": 0 + } + ], + "title": "CLIP Text Encode (Positive Prompt)", + "properties": { + "Node name for S&R": "CLIPTextEncode" + }, + "widgets_values": [ + "A woman with long brown hair and light skin smiles at another woman with long blonde hair. The woman with brown hair wears a black jacket and has a small, barely noticeable mole on her right cheek. The camera angle is a close-up, focused on the woman with brown hair's face. The lighting is warm and natural, likely from the setting sun, casting a soft glow on the scene. The scene appears to be real-life footage.", + true + ], + "color": "#232", + "bgcolor": "#353" + }, + { + "id": 7, + "type": "CLIPTextEncode", + "pos": [ + 420, + 390 + ], + "size": [ + 425.27801513671875, + 180.6060791015625 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "clip", + "type": "CLIP", + "link": 75 + } + ], + "outputs": [ + { + "name": "CONDITIONING", + "type": "CONDITIONING", + "links": [ + 170 + ], + "slot_index": 0 + } + ], + "title": "CLIP Text Encode (Negative Prompt)", + "properties": { + "Node name for S&R": "CLIPTextEncode" + }, + "widgets_values": [ + "low quality, worst quality, deformed, distorted, disfigured, motion smear, motion artifacts, fused fingers, bad anatomy, weird hand, ugly", + true + ], + "color": "#322", + "bgcolor": "#533" + }, + { + "id": 73, + "type": "KSamplerSelect", + "pos": [ + 860, + 420 + ], + "size": [ + 315, + 58 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "SAMPLER", + "type": "SAMPLER", + "links": [ + 172 + ] + } + ], + "properties": { + "Node name for S&R": "KSamplerSelect" + }, + "widgets_values": [ + "euler" + ] + }, + { + "id": 76, + "type": "Note", + "pos": [ + 40, + 350 + ], + "size": [ + 360, + 200 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": {}, + "widgets_values": [ + "This model needs long descriptive prompts, if the prompt is too short the quality will suffer greatly." + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 70, + "type": "EmptyLTXVLatentVideo", + "pos": [ + 860, + 240 + ], + "size": [ + 315, + 130 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "LATENT", + "type": "LATENT", + "links": [ + 168, + 175 + ], + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "EmptyLTXVLatentVideo" + }, + "widgets_values": [ + 768, + 768, + 97, + 1 + ] + }, + { + "id": 69, + "type": "LTXVConditioning", + "pos": [ + 920, + 60 + ], + "size": [ + 223.8660125732422, + 78 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "positive", + "type": "CONDITIONING", + "link": 169 + }, + { + "name": "negative", + "type": "CONDITIONING", + "link": 170 + } + ], + "outputs": [ + { + "name": "positive", + "type": "CONDITIONING", + "links": [ + 166 + ], + "slot_index": 0 + }, + { + "name": "negative", + "type": "CONDITIONING", + "links": [ + 167 + ], + "slot_index": 1 + } + ], + "properties": { + "Node name for S&R": "LTXVConditioning" + }, + "widgets_values": [ + 25 + ] + }, + { + "id": 38, + "type": "CLIPLoader", + "pos": [ + 60, + 190 + ], + "size": [ + 315, + 82 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "CLIP", + "type": "CLIP", + "links": [ + 74, + 75 + ], + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "CLIPLoader" + }, + "widgets_values": [ + "t5xxl_fp16.safetensors", + "ltxv" + ] + }, + { + "id": 8, + "type": "VAEDecode", + "pos": [ + 1600, + 30 + ], + "size": [ + 210, + 46 + ], + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [ + { + "name": "samples", + "type": "LATENT", + "link": 171 + }, + { + "name": "vae", + "type": "VAE", + "link": 87 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 185 + ], + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "VAEDecode" + }, + "widgets_values": [] + }, + { + "id": 86, + "type": "VHS_VideoCombine", + "pos": [ + 1890.7164306640625, + 30.7105770111084 + ], + "size": [ + 312.7515869140625, + 616.7515869140625 + ], + "flags": {}, + "order": 12, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 185 + }, + { + "name": "audio", + "type": "AUDIO", + "link": null, + "shape": 7 + }, + { + "name": "meta_batch", + "type": "VHS_BatchManager", + "link": null, + "shape": 7 + }, + { + "name": "vae", + "type": "VAE", + "link": null, + "shape": 7 + } + ], + "outputs": [ + { + "name": "Filenames", + "type": "VHS_FILENAMES", + "links": null + } + ], + "properties": { + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 24, + "loop_count": 0, + "filename_prefix": "ltxv", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 19, + "save_metadata": true, + "pingpong": false, + "save_output": true, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "ltxv_00032.mp4", + "subfolder": "", + "type": "output", + "format": "video/h264-mp4", + "frame_rate": 24 + }, + "muted": false + } + } + }, + { + "id": 72, + "type": "SamplerCustom", + "pos": [ + 1206.866943359375, + 26.604873657226562 + ], + "size": [ + 355.20001220703125, + 230 + ], + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "MODEL", + "link": 184 + }, + { + "name": "positive", + "type": "CONDITIONING", + "link": 166 + }, + { + "name": "negative", + "type": "CONDITIONING", + "link": 167 + }, + { + "name": "sampler", + "type": "SAMPLER", + "link": 172 + }, + { + "name": "sigmas", + "type": "SIGMAS", + "link": 182 + }, + { + "name": "latent_image", + "type": "LATENT", + "link": 175 + } + ], + "outputs": [ + { + "name": "output", + "type": "LATENT", + "links": [ + 171 + ], + "slot_index": 0 + }, + { + "name": "denoised_output", + "type": "LATENT", + "links": null + } + ], + "properties": { + "Node name for S&R": "SamplerCustom" + }, + "widgets_values": [ + true, + 11905454606274, + "fixed", + 3 + ] + }, + { + "id": 85, + "type": "TeaCacheForVidGen", + "pos": [ + 864.0098266601562, + -156.047119140625 + ], + "size": [ + 315, + 130 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "MODEL", + "link": 183 + } + ], + "outputs": [ + { + "name": "MODEL", + "type": "MODEL", + "links": [ + 184 + ], + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "TeaCacheForVidGen" + }, + "widgets_values": [ + true, + "ltxv", + 0.06, + 30 + ] + }, + { + "id": 44, + "type": "CheckpointLoaderSimple", + "pos": [ + 520.5762329101562, + 17.9000244140625 + ], + "size": [ + 315, + 98 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "MODEL", + "type": "MODEL", + "links": [ + 183 + ], + "slot_index": 0 + }, + { + "name": "CLIP", + "type": "CLIP", + "links": null + }, + { + "name": "VAE", + "type": "VAE", + "links": [ + 87 + ], + "slot_index": 2 + } + ], + "properties": { + "Node name for S&R": "CheckpointLoaderSimple" + }, + "widgets_values": [ + "ltx-video-2b-v0.9.1.safetensors" + ] + } + ], + "links": [ + [ + 74, + 38, + 0, + 6, + 0, + "CLIP" + ], + [ + 75, + 38, + 0, + 7, + 0, + "CLIP" + ], + [ + 87, + 44, + 2, + 8, + 1, + "VAE" + ], + [ + 166, + 69, + 0, + 72, + 1, + "CONDITIONING" + ], + [ + 167, + 69, + 1, + 72, + 2, + "CONDITIONING" + ], + [ + 168, + 70, + 0, + 71, + 0, + "LATENT" + ], + [ + 169, + 6, + 0, + 69, + 0, + "CONDITIONING" + ], + [ + 170, + 7, + 0, + 69, + 1, + "CONDITIONING" + ], + [ + 171, + 72, + 0, + 8, + 0, + "LATENT" + ], + [ + 172, + 73, + 0, + 72, + 3, + "SAMPLER" + ], + [ + 175, + 70, + 0, + 72, + 5, + "LATENT" + ], + [ + 182, + 71, + 0, + 72, + 4, + "SIGMAS" + ], + [ + 183, + 44, + 0, + 85, + 0, + "MODEL" + ], + [ + 184, + 85, + 0, + 72, + 0, + "MODEL" + ], + [ + 185, + 8, + 0, + 86, + 0, + "IMAGE" + ] + ], + "groups": [], + "config": {}, + "extra": { + "ds": { + "scale": 0.7513148009015777, + "offset": [ + -12.24569670436394, + 293.53734923538843 + ] + } + }, + "version": 0.4 +} \ No newline at end of file diff --git a/nodes.py b/nodes.py index 5fb4787..e2f7288 100644 --- a/nodes.py +++ b/nodes.py @@ -1,12 +1,16 @@ +import math import torch import numpy as np -from comfy.ldm.flux.model import Flux -from comfy.ldm.hunyuan_video.model import HunyuanVideo -from comfy.ldm.flux.layers import timestep_embedding - from torch import Tensor +from comfy.ldm.flux.model import Flux +from comfy.ldm.flux.layers import timestep_embedding +from comfy.ldm.hunyuan_video.model import HunyuanVideo +from comfy.ldm.lightricks.model import LTXVModel, precompute_freqs_cis +from comfy.ldm.common_dit import rms_norm + + def teacache_flux_forward( self, img: Tensor, @@ -266,6 +270,162 @@ def teacache_hunyuanvideo_forward( img = img.reshape(initial_shape) return img +def teacache_ltxvmodel_forward( + self, + x, + timestep, + context, + attention_mask, + frame_rate=25, + guiding_latent=None, + guiding_latent_noise_scale=0, + transformer_options={}, + **kwargs + ): + patches_replace = transformer_options.get("patches_replace", {}) + + indices_grid = self.patchifier.get_grid( + orig_num_frames=x.shape[2], + orig_height=x.shape[3], + orig_width=x.shape[4], + batch_size=x.shape[0], + scale_grid=((1 / frame_rate) * 8, 32, 32), + device=x.device, + ) + + if guiding_latent is not None: + ts = torch.ones([x.shape[0], 1, x.shape[2], x.shape[3], x.shape[4]], device=x.device, dtype=x.dtype) + input_ts = timestep.view([timestep.shape[0]] + [1] * (x.ndim - 1)) + ts *= input_ts + ts[:, :, 0] = guiding_latent_noise_scale * (input_ts[:, :, 0] ** 2) + timestep = self.patchifier.patchify(ts) + input_x = x.clone() + x[:, :, 0] = guiding_latent[:, :, 0] + if guiding_latent_noise_scale > 0: + if self.generator is None: + self.generator = torch.Generator(device=x.device).manual_seed(42) + elif self.generator.device != x.device: + self.generator = torch.Generator(device=x.device).set_state(self.generator.get_state()) + + noise_shape = [guiding_latent.shape[0], guiding_latent.shape[1], 1, guiding_latent.shape[3], guiding_latent.shape[4]] + scale = guiding_latent_noise_scale * (input_ts ** 2) + guiding_noise = scale * torch.randn(size=noise_shape, device=x.device, generator=self.generator) + + x[:, :, 0] = guiding_noise[:, :, 0] + x[:, :, 0] * (1.0 - scale[:, :, 0]) + + + orig_shape = list(x.shape) + + x = self.patchifier.patchify(x) + + x = self.patchify_proj(x) + timestep = timestep * 1000.0 + + attention_mask = 1.0 - attention_mask.to(x.dtype).reshape((attention_mask.shape[0], 1, -1, attention_mask.shape[-1])) + attention_mask = attention_mask.masked_fill(attention_mask.to(torch.bool), float("-inf")) # not sure about this + # attention_mask = (context != 0).any(dim=2).to(dtype=x.dtype) + + pe = precompute_freqs_cis(indices_grid, dim=self.inner_dim, out_dtype=x.dtype) + + batch_size = x.shape[0] + timestep, embedded_timestep = self.adaln_single( + timestep.flatten(), + {"resolution": None, "aspect_ratio": None}, + batch_size=batch_size, + hidden_dtype=x.dtype, + ) + # Second dimension is 1 or number of tokens (if timestep_per_token) + timestep = timestep.view(batch_size, -1, timestep.shape[-1]) + embedded_timestep = embedded_timestep.view( + batch_size, -1, embedded_timestep.shape[-1] + ) + + # 2. Blocks + if self.caption_projection is not None: + batch_size = x.shape[0] + context = self.caption_projection(context) + context = context.view( + batch_size, -1, x.shape[-1] + ) + + blocks_replace = patches_replace.get("dit", {}) + + # enable teacache + inp = x.clone() + timestep_ = timestep.clone() + num_ada_params = self.transformer_blocks[0].scale_shift_table.shape[0] + ada_values = self.transformer_blocks[0].scale_shift_table[None, None] + timestep_.reshape(batch_size, timestep_.size(1), num_ada_params, -1) + shift_msa, scale_msa, _, _, _, _ = ada_values.unbind(dim=2) + modulated_inp = rms_norm(inp) + modulated_inp = modulated_inp * (1 + scale_msa) + shift_msa + + if self.cnt == 0 or self.cnt == self.steps - 1: + should_calc = True + self.accumulated_rel_l1_distance = 0 + else: + coefficients = [2.14700694e+01, -1.28016453e+01, 2.31279151e+00, 7.92487521e-01, 9.69274326e-03] + rescale_func = np.poly1d(coefficients) + self.accumulated_rel_l1_distance += rescale_func(((modulated_inp-self.previous_modulated_input).abs().mean() / self.previous_modulated_input.abs().mean()).cpu().item()) + if self.accumulated_rel_l1_distance < self.rel_l1_thresh: + should_calc = False + else: + should_calc = True + self.accumulated_rel_l1_distance = 0 + + self.previous_modulated_input = modulated_inp + self.cnt += 1 + + if self.cnt == self.steps: + self.cnt = 0 + + if not should_calc: + x += self.previous_residual + else: + ori_x = x.clone() + for i, block in enumerate(self.transformer_blocks): + if ("double_block", i) in blocks_replace: + def block_wrap(args): + out = {} + out["img"] = block(args["img"], context=args["txt"], attention_mask=args["attention_mask"], timestep=args["vec"], pe=args["pe"]) + return out + + out = blocks_replace[("double_block", i)]({"img": x, "txt": context, "attention_mask": attention_mask, "vec": timestep, "pe": pe}, {"original_block": block_wrap}) + x = out["img"] + else: + x = block( + x, + context=context, + attention_mask=attention_mask, + timestep=timestep, + pe=pe + ) + + # 3. Output + scale_shift_values = ( + self.scale_shift_table[None, None].to(device=x.device, dtype=x.dtype) + embedded_timestep[:, :, None] + ) + shift, scale = scale_shift_values[:, :, 0], scale_shift_values[:, :, 1] + x = self.norm_out(x) + # Modulation + x = x * (1 + scale) + shift + self.previous_residual = x - ori_x + + x = self.proj_out(x) + + x = self.patchifier.unpatchify( + latents=x, + output_height=orig_shape[3], + output_width=orig_shape[4], + output_num_frames=orig_shape[2], + out_channels=orig_shape[1] // math.prod(self.patchifier.patch_size), + ) + + if guiding_latent is not None: + x[:, :, 0] = (input_x[:, :, 0] - guiding_latent[:, :, 0]) / input_ts[:, :, 0] + + # print("res", x) + return x + class TeaCacheForImgGen: @classmethod def INPUT_TYPES(s): @@ -286,10 +446,10 @@ class TeaCacheForImgGen: def apply_teacache(self, model, enable_teacache: bool, model_type: str, rel_l1_thresh: float, steps: int): if enable_teacache: + model.model.diffusion_model.__class__.cnt = 0 + model.model.diffusion_model.__class__.rel_l1_thresh = rel_l1_thresh + model.model.diffusion_model.__class__.steps = steps if model_type == "flux": - model.model.diffusion_model.__class__.cnt = 0 - model.model.diffusion_model.__class__.rel_l1_thresh = rel_l1_thresh - model.model.diffusion_model.__class__.steps = steps model.model.diffusion_model.forward_orig = teacache_flux_forward.__get__( model.model.diffusion_model, model.model.diffusion_model.__class__ @@ -314,7 +474,7 @@ class TeaCacheForVidGen: "required": { "model": ("MODEL", {"tooltip": "The video diffusion model the TeaCache will be applied to."}), "enable_teacache": ("BOOLEAN", {"default": True, "tooltip": "Enable teacache will speed up inference but may lose visual quality."}), - "model_type": (["hunyuan_video"],), + "model_type": (["hunyuan_video", "ltxv"],), "rel_l1_thresh": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "How strongly to cache the output of diffusion model. This value must be non-negative."}), "steps": ("INT", {"default": 25, "min": 1, "max": 10000, "step": 1}), } @@ -327,14 +487,19 @@ class TeaCacheForVidGen: def apply_teacache(self, model, enable_teacache: bool, model_type: str, rel_l1_thresh: float, steps: int): if enable_teacache: + model.model.diffusion_model.__class__.cnt = 0 + model.model.diffusion_model.__class__.rel_l1_thresh = rel_l1_thresh + model.model.diffusion_model.__class__.steps = steps if model_type == "hunyuan_video": - model.model.diffusion_model.__class__.cnt = 0 - model.model.diffusion_model.__class__.rel_l1_thresh = rel_l1_thresh - model.model.diffusion_model.__class__.steps = steps model.model.diffusion_model.forward_orig = teacache_hunyuanvideo_forward.__get__( model.model.diffusion_model, model.model.diffusion_model.__class__ ) + elif model_type == "ltxv": + model.model.diffusion_model.forward = teacache_ltxvmodel_forward.__get__( + model.model.diffusion_model, + model.model.diffusion_model.__class__ + ) else: raise ValueError(f"Unknown type {model_type}") else: @@ -343,6 +508,11 @@ class TeaCacheForVidGen: model.model.diffusion_model, model.model.diffusion_model.__class__ ) + elif model_type == "ltxv": + model.model.diffusion_model.forward = LTXVModel.forward.__get__( + model.model.diffusion_model, + model.model.diffusion_model.__class__ + ) else: raise ValueError(f"Unknown type {model_type}")