Support Phantom, refactor model dtypes, reduce DF model memory use

This commit is contained in:
kijai
2025-04-22 21:13:27 +03:00
parent a623f87dca
commit 6fadcbd957
4 changed files with 171 additions and 77 deletions
+77 -7
View File
@@ -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
View File
@@ -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"],
+3 -3
View File
@@ -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
View File
@@ -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"