Support Phantom, refactor model dtypes, reduce DF model memory use
This commit is contained in:
@@ -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",
|
||||
}
|
||||
|
||||
+4
-3
@@ -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"],
|
||||
|
||||
@@ -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)
|
||||
|
||||
+87
-64
@@ -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"
|
||||
Reference in New Issue
Block a user