From 6fadcbd957ba2990582fe6d494872be653956620 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 22 Apr 2025 21:13:27 +0300 Subject: [PATCH] Support Phantom, refactor model dtypes, reduce DF model memory use --- nodes.py | 84 +++++++++++++++++-- skyreels/nodes.py | 7 +- wanvideo/modules/attention.py | 6 +- wanvideo/modules/model.py | 151 ++++++++++++++++++++-------------- 4 files changed, 171 insertions(+), 77 deletions(-) diff --git a/nodes.py b/nodes.py index 47de9e9..112f541 100644 --- a/nodes.py +++ b/nodes.py @@ -663,9 +663,11 @@ class WanVideoModelLoader: total=param_count, leave=True): dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype - if "modulation" in name or "time_" in name: + if "patch_embedding" in name: dtype_to_use = torch.float32 set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name]) + for name, param in transformer.named_parameters(): + print(name, param.device, param.dtype) comfy_model.diffusion_model = transformer comfy_model.load_device = transformer_load_device @@ -864,7 +866,7 @@ class WanVideoModelLoader: patcher.model["base_path"] = model_path patcher.model["model_name"] = model patcher.model["manual_offloading"] = manual_offloading - patcher.model["quantization"] = "disabled" + patcher.model["quantization"] = quantization patcher.model["auto_cpu_offload"] = True if vram_management_args is not None else False patcher.model["control_lora"] = control_lora @@ -1750,6 +1752,40 @@ class WanVideoEmptyEmbeds: return (embeds,) +# region phantom +class WanVideoPhantomEmbeds: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}), + "phantom_latents": ("LATENT", {"tooltip": "reference latents for the phantom model"}), + "phantom_cfg_scale": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "CFG scale for the extra phantom cond pass"}), + }, + } + + RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", ) + RETURN_NAMES = ("image_embeds",) + FUNCTION = "process" + CATEGORY = "WanVideoWrapper" + + def process(self, num_frames, phantom_latents, phantom_cfg_scale): + vae_stride = (4, 8, 8) + samples = phantom_latents["samples"].squeeze(0) + C, T, H, W = samples.shape + + target_shape = (16, (num_frames - 1) // vae_stride[0] + 1 + T, + H * 8 // vae_stride[1], + W * 8 // vae_stride[2]) + + embeds = { + "target_shape": target_shape, + "num_frames": num_frames, + "phantom_latents": samples, + "phantom_cfg_scale": phantom_cfg_scale, + } + + return (embeds,) + class WanVideoControlEmbeds: @classmethod def INPUT_TYPES(s): @@ -2213,7 +2249,7 @@ class WanVideoSampler: patcher = model model = model.model transformer = model.diffusion_model - + dtype = model["dtype"] control_lora = model["control_lora"] device = mm.get_torch_device() @@ -2270,6 +2306,7 @@ class WanVideoSampler: control_latents, clip_fea, clip_fea_neg, end_image, recammaster, camera_embed, unianim_data = None, None, None, None, None, None, None vace_data, vace_context, vace_scale = None, None, None fun_or_fl2v_model, has_ref, drop_last = False, False, False + phantom_latents = None image_cond = image_embeds.get("image_embeds", None) @@ -2382,6 +2419,11 @@ class WanVideoSampler: masked_video_latents_input = torch.zeros_like(noise) image_cond = torch.cat([mask_latents, masked_video_latents_input], dim=0).to(device) + phantom_latents = image_embeds.get("phantom_latents", None) + phantom_cfg_scale = image_embeds.get("phantom_cfg_scale", None) + if phantom_latents is not None: + phantom_latents = phantom_latents.to(device) + latent_video_length = noise.shape[1] if unianimate_poses is not None: @@ -2577,6 +2619,7 @@ class WanVideoSampler: transformer.rel_l1_thresh = teacache_args["rel_l1_thresh"] transformer.teacache_start_step = teacache_args["start_step"] transformer.teacache_cache_device = teacache_args["cache_device"] + log.info(f"TeaCache: Using cache device: {transformer.teacache_state.cache_device}") transformer.teacache_end_step = len(timesteps)-1 if teacache_args["end_step"] == -1 else teacache_args["end_step"] transformer.teacache_use_coefficients = teacache_args["use_coefficients"] transformer.teacache_mode = teacache_args["mode"] @@ -2593,6 +2636,8 @@ class WanVideoSampler: transformer.slg_blocks = None self.teacache_state = [None, None] + if phantom_latents is not None: + self.teacache_state = [None, None, None] self.teacache_state_source = [None, None] self.teacache_states_context = [] @@ -2649,7 +2694,7 @@ class WanVideoSampler: #region model pred def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, control_latents=None, vace_data=None, unianim_data=None, teacache_state=None): - with torch.autocast(device_type=mm.get_autocast_device(device), dtype=model["dtype"], enabled=True): + with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype, enabled=("fp8" in model["quantization"])): if use_cfg_zero_star and (idx <= zero_star_steps) and use_zero_init: return latent_model_input*0, None @@ -2684,9 +2729,16 @@ class WanVideoSampler: else: image_cond_input = image_cond + z = z.to(dtype) + z_pos = z_neg = z + if recammaster is not None: z = torch.cat([z, recam_latents.to(z)], dim=1) - + if phantom_latents is not None: + z_pos = torch.cat([z_pos[:,:-phantom_latents.shape[1]], phantom_latents.to(z)], dim=1) + z_phantom_img = torch.cat([z_pos[:,:-phantom_latents.shape[1]], phantom_latents.to(z)], dim=1) + z_neg = torch.cat([z_pos[:,:-phantom_latents.shape[1]], torch.zeros_like(phantom_latents).to(z)], dim=1) + base_params = { 'seq_len': seq_len, 'device': device, @@ -2707,7 +2759,7 @@ class WanVideoSampler: if not batched_cfg: #cond noise_pred_cond, teacache_state_cond = transformer( - [z], context=positive_embeds, y=[image_cond_input] if image_cond_input is not None else None, + [z_pos], context=positive_embeds, y=[image_cond_input] if image_cond_input is not None else None, clip_fea=clip_fea, is_uncond=False, current_step_percentage=current_step_percentage, pred_id=teacache_state[0] if teacache_state else None, **base_params @@ -2724,13 +2776,26 @@ class WanVideoSampler: return noise_pred_cond, [teacache_state_cond] #uncond noise_pred_uncond, teacache_state_uncond = transformer( - [z], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea, + [z_neg], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea, y=[image_cond_input] if image_cond_input is not None else None, is_uncond=True, current_step_percentage=current_step_percentage, pred_id=teacache_state[1] if teacache_state else None, **base_params ) noise_pred_uncond = noise_pred_uncond[0].to(intermediate_device) + #phantom + if phantom_latents is not None: + noise_pred_phantom, teacache_state_phantom = transformer( + [z_phantom_img], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea, + y=[image_cond_input] if image_cond_input is not None else None, + is_uncond=True, current_step_percentage=current_step_percentage, + pred_id=teacache_state[2] if teacache_state else None, + **base_params + ) + noise_pred_phantom = noise_pred_phantom[0].to(intermediate_device) + + noise_pred = noise_pred_uncond + phantom_cfg_scale * (noise_pred_phantom - noise_pred_uncond) + cfg_scale * (noise_pred_cond - noise_pred_phantom) + return noise_pred, [teacache_state_cond, teacache_state_uncond, teacache_state_phantom] #batched else: teacache_state_uncond = None @@ -3103,6 +3168,9 @@ class WanVideoSampler: callback(idx, callback_latent, None, steps) else: pbar.update(1) + + if phantom_latents is not None: + x0 = x0[:,:-phantom_latents.shape[1]] if teacache_args is not None: states = transformer.teacache_state.states @@ -3362,6 +3430,7 @@ NODE_CLASS_MAPPINGS = { "WanVideoVACEEncode": WanVideoVACEEncode, "WanVideoVACEStartToEndFrame": WanVideoVACEStartToEndFrame, "WanVideoVACEModelSelect": WanVideoVACEModelSelect, + "WanVideoPhantomEmbeds": WanVideoPhantomEmbeds, } NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoSampler": "WanVideo Sampler", @@ -3397,4 +3466,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoVACEEncode": "WanVideo VACE Encode", "WanVideoVACEStartToEndFrame": "WanVideo VACE Start To End Frame", "WanVideoVACEModelSelect": "WanVideo VACE Model Select", + "WanVideoPhantomEmbeds": "WanVideo Phantom Embeds", } diff --git a/skyreels/nodes.py b/skyreels/nodes.py index 8a301f2..d90afa2 100644 --- a/skyreels/nodes.py +++ b/skyreels/nodes.py @@ -139,7 +139,7 @@ class WanVideoDiffusionForcingSampler: patcher = model model = model.model transformer = model.diffusion_model - + dtype = model["dtype"] device = mm.get_torch_device() offload_device = mm.unet_offload_device() @@ -371,6 +371,7 @@ class WanVideoDiffusionForcingSampler: transformer.rel_l1_thresh = teacache_args["rel_l1_thresh"] transformer.teacache_start_step = teacache_args["start_step"] transformer.teacache_cache_device = teacache_args["cache_device"] + log.info(f"TeaCache: Using cache device: {transformer.teacache_state.cache_device}") transformer.teacache_end_step = len(init_timesteps)-1 if teacache_args["end_step"] == -1 else teacache_args["end_step"] transformer.teacache_use_coefficients = teacache_args["use_coefficients"] transformer.teacache_mode = teacache_args["mode"] @@ -410,7 +411,7 @@ class WanVideoDiffusionForcingSampler: #region model pred def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, vace_data=None, unianim_data=None, teacache_state=None): - with torch.autocast(device_type=mm.get_autocast_device(device), dtype=model["dtype"], enabled=True): + with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype, enabled=("fp8" in model["quantization"])): if use_cfg_zero_star and (idx <= zero_star_steps) and use_zero_init: return latent_model_input*0, None @@ -525,7 +526,7 @@ class WanVideoDiffusionForcingSampler: #print("timestep", timestep) noise_pred, self.teacache_state = predict_with_cfg( - latent_model_input, + latent_model_input.to(dtype), cfg[i], text_embeds["prompt_embeds"], text_embeds["negative_prompt_embeds"], diff --git a/wanvideo/modules/attention.py b/wanvideo/modules/attention.py index 96fcea5..12a6cd4 100644 --- a/wanvideo/modules/attention.py +++ b/wanvideo/modules/attention.py @@ -196,9 +196,9 @@ def attention( elif attention_mode == 'sageattn': attn_mask = None - q = q.transpose(1, 2).to(dtype) - k = k.transpose(1, 2).to(dtype) - v = v.transpose(1, 2).to(dtype) + q = q.transpose(1, 2)#.to(dtype) + k = k.transpose(1, 2)#.to(dtype) + v = v.transpose(1, 2)#.to(dtype) out = sageattn_func( q, k, v, attn_mask=attn_mask, is_causal=causal, dropout_p=dropout_p) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 478ab25..034aafd 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -134,10 +134,10 @@ class WanRMSNorm(nn.Module): Args: x(Tensor): Shape [B, L, C] """ - return self._norm(x.float()).type_as(x) * self.weight + return self._norm(x)* self.weight def _norm(self, x): - return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps) + return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps).to(x.dtype) class WanLayerNorm(nn.LayerNorm): @@ -150,7 +150,7 @@ class WanLayerNorm(nn.LayerNorm): Args: x(Tensor): Shape [B, L, C] """ - return super().forward(x.float()).type_as(x) + return super().forward(x) class WanSelfAttention(nn.Module): @@ -442,6 +442,20 @@ class WanAttentionBlock(nn.Module): # modulation self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5) + @torch.compiler.disable() + def get_mod(self, e): + if e.dim() == 3: + modulation = self.modulation # 1, 6, dim + e = (modulation.to(e.device) + e).chunk(6, dim=1) + elif e.dim() == 4: + modulation = self.modulation.unsqueeze(2) # 1, 6, 1, dim + e = (modulation.to(e.device) + e).chunk(6, dim=1) + e = [ei.squeeze(1) for ei in e] + return e + + def modulate(self, x, e): + return x * (1 + e[1]) + e[0] + def forward( self, x, @@ -467,16 +481,9 @@ class WanAttentionBlock(nn.Module): freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2] """ #e = (self.modulation.to(e.device) + e).chunk(6, dim=1) - - if e.dim() == 3: - modulation = self.modulation # 1, 6, dim - e = (modulation.to(e.device) + e).chunk(6, dim=1) - elif e.dim() == 4: - modulation = self.modulation.unsqueeze(2) # 1, 6, 1, dim - e = (modulation.to(e.device) + e).chunk(6, dim=1) - e = [ei.squeeze(1) for ei in e] + e = self.get_mod(e) - input_x = self.norm1(x) * (1 + e[1]) + e[0] + input_x = self.modulate(self.norm1(x), e) if camera_embed is not None: # encode ReCamMaster camera @@ -506,20 +513,23 @@ class WanAttentionBlock(nn.Module): if camera_embed is not None: y = self.projector(y) - x = x.to(torch.float32) + (y.to(torch.float32) * e[2].to(torch.float32)) + del input_x + + x = x + (y * e[2]) + del y # cross-attention & ffn function if (context.shape[0] > 1 or (clip_embed is not None and clip_embed.shape[0] > 1)) and x.shape[0] == 1: x = self.split_cross_attn_ffn(x, context, context_lens, e, clip_embed=clip_embed, grid_sizes=grid_sizes) else: x = self.cross_attn_ffn(x, context, context_lens, e, clip_embed=clip_embed, grid_sizes=grid_sizes) - + del e return x - + @torch.compiler.disable() def cross_attn_ffn(self, x, context, context_lens, e, clip_embed=None, grid_sizes=None): x = x + self.cross_attn(self.norm3(x), context, context_lens, clip_embed=clip_embed) - y = self.ffn(self.norm2(x).float() * (1 + e[4]) + e[3]) - x = x.to(torch.float32) + (y.to(torch.float32) * e[5]) + y = self.ffn(self.norm2(x) * (1 + e[4]) + e[3]) + x = x + (y * e[5]) return x @torch.compiler.disable() @@ -574,9 +584,9 @@ class WanAttentionBlock(nn.Module): # Continue with FFN x = x + x_combined - y = self.ffn(self.norm2(x).float() * (1 + e[4]) + e[3]) - x = x.to(torch.float32) + (y.to(torch.float32) * e[5].to(torch.float32)) - return x + y = self.ffn(self.norm2(x) * (1 + e[4]) + e[3]) + x = x + (y * e[5]) + return x class VaceWanAttentionBlock(WanAttentionBlock): def __init__( @@ -659,6 +669,16 @@ class Head(nn.Module): # modulation self.modulation = nn.Parameter(torch.randn(1, 2, dim) / dim**0.5) + def get_mod(self, e): + if e.dim() == 2: + modulation = self.modulation.to(e.device) # 1, 2, dim + e = (modulation + e.unsqueeze(1)).chunk(2, dim=1) + elif e.dim() == 3: + modulation = self.modulation.to(e.device).unsqueeze(2) # 1, 2, seq, dim + e = (modulation + e.unsqueeze(1)).chunk(2, dim=1) + e = [ei.squeeze(1) for ei in e] + return e + def forward(self, x, e): r""" Args: @@ -670,13 +690,7 @@ class Head(nn.Module): # normed = self.norm(x) # x = self.head(normed * (1 + e[1]) + e[0]) - if e.dim() == 2: - modulation = self.modulation.to(e.device) # 1, 2, dim - e = (modulation + e.unsqueeze(1)).chunk(2, dim=1) - elif e.dim() == 3: - modulation = self.modulation.to(e.device).unsqueeze(2) # 1, 2, seq, dim - e = (modulation + e.unsqueeze(1)).chunk(2, dim=1) - e = [ei.squeeze(1) for ei in e] + e = self.get_mod(e) x = self.head(self.norm(x) * (1 + e[1]) + e[0]) return x @@ -1032,13 +1046,13 @@ class WanModel(ModelMixin, ConfigMixin): if control_lora_enabled: self.expanded_patch_embedding.to(device) x = [ - self.expanded_patch_embedding(u.unsqueeze(0)) + self.expanded_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in x ] else: self.original_patch_embedding.to(self.main_device) x = [ - self.original_patch_embedding(u.unsqueeze(0)) + self.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in x ] @@ -1069,39 +1083,43 @@ class WanModel(ModelMixin, ConfigMixin): rope_func = "default" # time embeddings - with torch.autocast(device_type='cuda', dtype=torch.float32): - # e = self.time_embedding( - # sinusoidal_embedding_1d(self.freq_dim, t).float()) - # e0 = self.time_projection(e).unflatten(1, (6, self.dim)) - # assert e.dtype == torch.float32 and e0.dtype == torch.float32 - if t.dim() == 2: - b, f = t.shape - _flag_df = True - else: - _flag_df = False + + # e = self.time_embedding( + # sinusoidal_embedding_1d(self.freq_dim, t).float()) + # e0 = self.time_projection(e).unflatten(1, (6, self.dim)) + # assert e.dtype == torch.float32 and e0.dtype == torch.float32 + if t.dim() == 2: + b, f = t.shape + _flag_df = True + else: + _flag_df = False - e = self.time_embedding( - sinusoidal_embedding_1d(self.freq_dim, t.flatten()).to(self.patch_embedding.weight.dtype) - ) # b, dim - e0 = self.time_projection(e).unflatten(1, (6, self.dim)) # b, 6, dim + e = self.time_embedding( + sinusoidal_embedding_1d(self.freq_dim, t.flatten()).to(x.dtype) + ) # b, dim + e0 = self.time_projection(e).unflatten(1, (6, self.dim)) # b, 6, dim - if fps_embeds is not None: - fps_embeds = torch.tensor(fps_embeds, dtype=torch.long, device=device) - - fps_emb = self.fps_embedding(fps_embeds).float() - if _flag_df: - e0 = e0 + self.fps_projection(fps_emb).unflatten(1, (6, self.dim)).repeat(t.shape[1], 1, 1) - else: - e0 = e0 + self.fps_projection(fps_emb).unflatten(1, (6, self.dim)) + if fps_embeds is not None: + fps_embeds = torch.tensor(fps_embeds, dtype=torch.long, device=device) + fps_emb = self.fps_embedding(fps_embeds).float() if _flag_df: - e = e.view(b, f, 1, 1, self.dim) - e0 = e0.view(b, f, 1, 1, 6, self.dim) - e = e.repeat(1, 1, grid_sizes[0][1], grid_sizes[0][2], 1).flatten(1, 3) - e0 = e0.repeat(1, 1, grid_sizes[0][1], grid_sizes[0][2], 1, 1).flatten(1, 3) - e0 = e0.transpose(1, 2).contiguous() + e0 = e0 + self.fps_projection(fps_emb).unflatten(1, (6, self.dim)).repeat(t.shape[1], 1, 1) + else: + e0 = e0 + self.fps_projection(fps_emb).unflatten(1, (6, self.dim)) - assert e.dtype == torch.float32 and e0.dtype == torch.float32 + if _flag_df: + e = e.view(b, f, 1, 1, self.dim).expand(b, f, grid_sizes[0][1], grid_sizes[0][2], self.dim) + e0 = e0.view(b, f, 1, 1, 6, self.dim).expand(b, f, grid_sizes[0][1], grid_sizes[0][2], 6, self.dim) + + e = e.flatten(1, 3) + e0 = e0.flatten(1, 3) + + e0 = e0.transpose(1, 2) + if not e0.is_contiguous(): + e0 = e0.contiguous() + + e = e.to(self.offload_device, non_blocking=self.use_non_blocking) # context context_lens = None @@ -1112,7 +1130,7 @@ class WanModel(ModelMixin, ConfigMixin): torch.cat( [u, u.new_zeros(self.text_len - u.size(0), u.size(1))]) for u in context - ])) + ]).to(x.dtype)) if self.offload_txt_emb: self.text_embedding.to(self.offload_device, non_blocking=self.use_non_blocking) @@ -1147,6 +1165,7 @@ class WanModel(ModelMixin, ConfigMixin): else: temb_relative_l1 = relative_l1_distance(previous_modulated_input, e0) accumulated_rel_l1_distance = accumulated_rel_l1_distance.to(e0.device) + temb_relative_l1 + del temb #print("accumulated_rel_l1_distance", accumulated_rel_l1_distance) @@ -1155,8 +1174,10 @@ class WanModel(ModelMixin, ConfigMixin): else: should_calc = True accumulated_rel_l1_distance = torch.tensor(0.0, dtype=torch.float32, device=device) + accumulated_rel_l1_distance = accumulated_rel_l1_distance.to(self.teacache_cache_device, non_blocking=self.use_non_blocking) previous_modulated_input = e.clone() if (self.teacache_use_coefficients and self.teacache_mode == 'e') else e0.clone() + previous_modulated_input = previous_modulated_input.to(self.teacache_cache_device, non_blocking=self.use_non_blocking) if not should_calc: x = x.to(previous_residual.dtype) + previous_residual.to(x.device) #log.info(f"TeaCache: Skipping uncond step {current_step+1}") @@ -1174,7 +1195,6 @@ class WanModel(ModelMixin, ConfigMixin): if unianim_data['start_percent'] <= current_step_percentage <= unianim_data['end_percent']: dwpose_emb = unianim_data['dwpose'] x += dwpose_emb * unianim_data['strength'] - # arguments kwargs = dict( e=e0, @@ -1216,7 +1236,7 @@ class WanModel(ModelMixin, ConfigMixin): continue if b <= self.blocks_to_swap and self.blocks_to_swap >= 0: block.to(self.main_device) - x = block(x.to(torch.float32), **kwargs) + x = block(x, **kwargs) if b <= self.blocks_to_swap and self.blocks_to_swap >= 0: block.to(self.offload_device, non_blocking=self.use_non_blocking) @@ -1224,10 +1244,10 @@ class WanModel(ModelMixin, ConfigMixin): self.teacache_state.update( pred_id, previous_residual=(x.to(original_x.device) - original_x), - accumulated_rel_l1_distance=accumulated_rel_l1_distance.to(self.teacache_cache_device, non_blocking=self.use_non_blocking), - previous_modulated_input=previous_modulated_input.to(self.teacache_cache_device, non_blocking=self.use_non_blocking) + accumulated_rel_l1_distance=accumulated_rel_l1_distance, + previous_modulated_input=previous_modulated_input ) - x = self.head(x, e) + x = self.head(x, e.to(x.device)) x = self.unpatchify(x, grid_sizes) # type: ignore[arg-type] x = [u.float() for u in x] return (x, pred_id) if pred_id is not None else (x, None) @@ -1260,7 +1280,6 @@ class WanModel(ModelMixin, ConfigMixin): class TeaCacheState: def __init__(self, cache_device='cpu'): self.cache_device = cache_device - log.info(f"TeaCache: Using cache device: {self.cache_device}") self.states = {} self._next_pred_id = 0 @@ -1304,3 +1323,7 @@ def relative_l1_distance(last_tensor, current_tensor): norm = torch.abs(last_tensor).mean() relative_l1_distance = l1_distance / norm return relative_l1_distance.to(torch.float32).to(current_tensor.device) + +def get_tensor_memory(tensor): + memory_bytes = tensor.element_size() * tensor.nelement() + return f"{memory_bytes / (1024 * 1024):.2f} MB" \ No newline at end of file