Merge branch 'dev'

This commit is contained in:
kijai
2025-04-25 16:20:59 +03:00
6 changed files with 2195 additions and 107 deletions
File diff suppressed because it is too large Load Diff
+2 -2
View File
@@ -7,8 +7,8 @@ def fp8_linear_forward(cls, original_dtype, input):
weight_dtype = cls.weight.dtype
if weight_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
if len(input.shape) == 3:
target_dtype = torch.float8_e5m2 if weight_dtype == torch.float8_e4m3fn else torch.float8_e4m3fn
inn = input.reshape(-1, input.shape[2]).to(target_dtype)
#target_dtype = torch.float8_e5m2 if weight_dtype == torch.float8_e4m3fn else torch.float8_e4m3fn
inn = input.reshape(-1, input.shape[2]).to(weight_dtype)
w = cls.weight.t()
scale = torch.ones((1), device=input.device, dtype=torch.float32)
+160 -30
View File
@@ -465,7 +465,7 @@ class WanVideoModelLoader:
"model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' -folder",}),
"base_precision": (["fp32", "bf16", "fp16", "fp16_fast"], {"default": "bf16"}),
"quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e4m3fn_fast', 'fp8_e5m2', 'torchao_fp8dq', "torchao_fp8dqrow", "torchao_int8dq", "torchao_fp6", "torchao_int4", "torchao_int8"], {"default": 'disabled', "tooltip": "optional quantization method"}),
"quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e4m3fn_fast', 'fp8_e5m2', 'fp8_e4m3fn_fast_no_ffn', 'torchao_fp8dq', "torchao_fp8dqrow", "torchao_int8dq", "torchao_fp6", "torchao_int4", "torchao_int8"], {"default": 'disabled', "tooltip": "optional quantization method"}),
"load_device": (["main_device", "offload_device"], {"default": "main_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}),
},
"optional": {
@@ -621,6 +621,8 @@ class WanVideoModelLoader:
"vace_layers": vace_layers,
"vace_in_dim": vace_in_dim,
"inject_sample_info": True if "fps_embedding.weight" in sd else False,
"add_ref_conv": True if "ref_conv.weight" in sd else False,
"in_dim_ref_conv": sd["ref_conv.weight"].shape[1] if "ref_conv.weight" in sd else None,
}
with init_empty_weights():
@@ -637,7 +639,7 @@ class WanVideoModelLoader:
block.cam_encoder.bias.data.zero_()
block.projector.weight = nn.Parameter(torch.eye(dim))
block.projector.bias = nn.Parameter(torch.zeros(dim))
comfy_model = WanVideoModel(
WanVideoModelConfig(base_dtype),
model_type=comfy.model_base.ModelType.FLOW,
@@ -646,13 +648,13 @@ class WanVideoModelLoader:
if not "torchao" in quantization:
if quantization == "fp8_e4m3fn" or quantization == "fp8_e4m3fn_fast" or quantization == "fp8_scaled":
if "fp8_e4m3fn" in quantization:
dtype = torch.float8_e4m3fn
elif quantization == "fp8_e5m2":
dtype = torch.float8_e5m2
else:
dtype = base_dtype
params_to_keep = {"norm", "head", "bias", "time_in", "vector_in", "patch_embedding", "time_", "img_emb", "modulation"}
params_to_keep = {"norm", "head", "bias", "time_in", "vector_in", "patch_embedding", "time_", "img_emb", "modulation", "text_embedding"}
#if lora is not None:
# transformer_load_device = device
if not lora_low_mem_load:
@@ -663,7 +665,7 @@ 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])
comfy_model.diffusion_model = transformer
@@ -708,7 +710,7 @@ class WanVideoModelLoader:
transformer.patch_embedding.kernel_size,
transformer.patch_embedding.stride,
transformer.patch_embedding.padding,
).to(device=device, dtype=torch.bfloat16)
).to(device=device, dtype=torch.float32)
new_in.weight.zero_()
new_in.bias.zero_()
@@ -728,13 +730,16 @@ class WanVideoModelLoader:
#patcher.load(device, full_load=True)
patcher.model.is_patched = True
del sd
if quantization == "fp8_e4m3fn_fast":
if "fast" in quantization:
from .fp8_optimization import convert_fp8_linear
#params_to_keep.update({"ffn"})
if quantization == "fp8_e4m3fn_fast_no_ffn":
params_to_keep.update({"ffn"})
print(params_to_keep)
convert_fp8_linear(patcher.model.diffusion_model, base_dtype, params_to_keep=params_to_keep)
convert_fp8_linear(patcher.model.diffusion_model, base_dtype, params_to_keep=params_to_keep, sd=sd)
del sd
if vram_management_args is not None:
from .diffsynth.vram_management import enable_vram_management, AutoWrappedModule, AutoWrappedLinear
@@ -864,7 +869,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 +1755,68 @@ 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_latent_1": ("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"}),
"phantom_start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the phantom model"}),
"phantom_end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the phantom model"}),
},
"optional": {
"phantom_latent_2": ("LATENT", {"tooltip": "reference latents for the phantom model"}),
"phantom_latent_3": ("LATENT", {"tooltip": "reference latents for the phantom model"}),
"phantom_latent_4": ("LATENT", {"tooltip": "reference latents for the phantom model"}),
"vace_embeds": ("WANVIDIMAGE_EMBEDS", {"tooltip": "VACE embeds"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", )
RETURN_NAMES = ("image_embeds",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, num_frames, phantom_cfg_scale, phantom_start_percent, phantom_end_percent, phantom_latent_1, phantom_latent_2=None, phantom_latent_3=None, phantom_latent_4=None, vace_embeds=None):
vae_stride = (4, 8, 8)
samples = phantom_latent_1["samples"].squeeze(0)
if phantom_latent_2 is not None:
samples = torch.cat([samples, phantom_latent_2["samples"].squeeze(0)], dim=1)
if phantom_latent_3 is not None:
samples = torch.cat([samples, phantom_latent_3["samples"].squeeze(0)], dim=1)
if phantom_latent_4 is not None:
samples = torch.cat([samples, phantom_latent_4["samples"].squeeze(0)], dim=1)
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,
"phantom_start_percent": phantom_start_percent,
"phantom_end_percent": phantom_end_percent,
}
if vace_embeds is not None:
vace_input = {
"vace_context": vace_embeds["vace_context"],
"vace_scale": vace_embeds["vace_scale"],
"has_ref": vace_embeds["has_ref"],
"vace_start_percent": vace_embeds["vace_start_percent"],
"vace_end_percent": vace_embeds["vace_end_percent"],
"vace_seq_len": vace_embeds["vace_seq_len"],
"additional_vace_inputs": vace_embeds["additional_vace_inputs"],
}
embeds.update(vace_input)
return (embeds,)
class WanVideoControlEmbeds:
@classmethod
def INPUT_TYPES(s):
@@ -1758,6 +1825,9 @@ class WanVideoControlEmbeds:
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the control signal"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the control signal"}),
},
"optional": {
"fun_ref_image": ("LATENT", {"tooltip": "Reference latent for the Fun 1.1 -model"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", )
@@ -1765,7 +1835,7 @@ class WanVideoControlEmbeds:
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, latents, start_percent, end_percent):
def process(self, latents, start_percent, end_percent, fun_ref_image=None):
samples = latents["samples"].squeeze(0)
C, T, H, W = samples.shape
@@ -1780,7 +1850,8 @@ class WanVideoControlEmbeds:
"control_embeds": {
"control_images": samples,
"start_percent": start_percent,
"end_percent": end_percent
"end_percent": end_percent,
"fun_ref_image": fun_ref_image["samples"][:,:, 0] if fun_ref_image is not None else None,
}
}
@@ -2213,7 +2284,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()
@@ -2269,7 +2340,9 @@ 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
fun_or_fl2v_model, has_ref, drop_last, = False, False, False
phantom_latents = None
fun_ref_image = None
image_cond = image_embeds.get("image_embeds", None)
@@ -2292,7 +2365,11 @@ class WanVideoSampler:
image_cond = image_embeds.get("image_embeds", None)
print("image_cond", image_cond.shape)
clip_fea = image_embeds.get("clip_context", None)
if clip_fea is not None:
clip_fea = clip_fea.to(dtype)
clip_fea_neg = image_embeds.get("negative_clip_context", None)
if clip_fea_neg is not None:
clip_fea_neg = clip_fea_neg.to(dtype)
control_embeds = image_embeds.get("control_embeds", None)
if control_embeds is not None:
@@ -2371,7 +2448,7 @@ class WanVideoSampler:
raise ValueError("Control signal only works with Fun-Control model")
image_cond = torch.zeros_like(control_latents).to(device) #fun control
clip_fea = None
fun_ref_image = control_embeds.get("fun_ref_image", None)
control_start_percent = control_embeds.get("start_percent", 0.0)
control_end_percent = control_embeds.get("end_percent", 1.0)
else:
@@ -2382,6 +2459,13 @@ 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)
phantom_start_percent = image_embeds.get("phantom_start_percent", 0.0)
phantom_end_percent = image_embeds.get("phantom_end_percent", 1.0)
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 +2661,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 +2678,9 @@ class WanVideoSampler:
transformer.slg_blocks = None
self.teacache_state = [None, None]
if phantom_latents is not None:
log.info(f"Phantom latents shape: {phantom_latents.shape}")
self.teacache_state = [None, None, None]
self.teacache_state_source = [None, None]
self.teacache_states_context = []
@@ -2601,6 +2689,8 @@ class WanVideoSampler:
source_image_embeds = flowedit_args.get("source_image_embeds", image_embeds)
source_image_cond = source_image_embeds.get("image_embeds", None)
source_clip_fea = source_image_embeds.get("clip_fea", clip_fea)
if source_image_cond is not None:
source_image_cond = source_image_cond.to(dtype)
skip_steps = flowedit_args["skip_steps"]
drift_steps = flowedit_args["drift_steps"]
source_cfg = flowedit_args["source_cfg"]
@@ -2649,7 +2739,8 @@ 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):
z = z.to(dtype)
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
@@ -2664,9 +2755,14 @@ class WanVideoSampler:
else:
if (control_start_percent <= current_step_percentage <= control_end_percent) or \
(control_end_percent > 0 and idx == 0 and current_step_percentage >= control_start_percent):
image_cond_input = torch.cat([control_latents, image_cond])
image_cond_input = torch.cat([control_latents.to(z), image_cond.to(z)])
else:
image_cond_input = torch.cat([torch.zeros_like(image_cond), image_cond])
image_cond_input = torch.cat([torch.zeros_like(image_cond, dtype=dtype), image_cond.to(z)])
if fun_ref_image is not None:
fun_ref_input = fun_ref_image.to(z)
else:
fun_ref_input = torch.zeros_like(z, dtype=z.dtype)[:, 0].unsqueeze(1)
#fun_ref_input = None
if control_lora:
if not control_start_percent <= current_step_percentage <= control_end_percent:
@@ -2676,17 +2772,30 @@ class WanVideoSampler:
patcher.unpatch_model(device)
patcher.model.is_patched = False
else:
image_cond_input = control_latents.to(device)
image_cond_input = control_latents.to(z)
if not patcher.model.is_patched:
log.info("Loading LoRA...")
patcher = apply_lora(patcher, device, device, low_mem_load=False)
patcher.model.is_patched = True
else:
image_cond_input = image_cond
image_cond_input = image_cond.to(z) if image_cond is not None else None
if recammaster is not None:
z = torch.cat([z, recam_latents.to(z)], dim=1)
use_phantom = False
if phantom_latents is not None:
if (phantom_start_percent <= current_step_percentage <= phantom_end_percent) or \
(phantom_end_percent > 0 and idx == 0 and current_step_percentage >= phantom_start_percent):
z_pos = torch.cat([z[:,:-phantom_latents.shape[1]], phantom_latents.to(z)], dim=1)
z_phantom_img = torch.cat([z[:,:-phantom_latents.shape[1]], phantom_latents.to(z)], dim=1)
z_neg = torch.cat([z[:,:-phantom_latents.shape[1]], torch.zeros_like(phantom_latents).to(z)], dim=1)
use_phantom = True
if len(teacache_state) != 3:
teacache_state.append(None)
if not use_phantom:
z_pos = z_neg = z
base_params = {
'seq_len': seq_len,
'device': device,
@@ -2694,9 +2803,9 @@ class WanVideoSampler:
't': timestep,
'current_step': idx,
'control_lora_enabled': control_lora_enabled,
'vace_data': vace_data,
'camera_embed': camera_embed,
'unianim_data': unianim_data,
'fun_ref': fun_ref_input if fun_ref_image is not None else None,
}
batch_size = 1
@@ -2707,9 +2816,10 @@ 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,
vace_data=vace_data,
**base_params
)
noise_pred_cond = noise_pred_cond[0].to(intermediate_device)
@@ -2724,13 +2834,28 @@ 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,
vace_data=vace_data,
**base_params
)
noise_pred_uncond = noise_pred_uncond[0].to(intermediate_device)
#phantom
if use_phantom:
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,
vace_data=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
@@ -2827,11 +2952,11 @@ class WanVideoSampler:
latent_model_input = torch.cat([latent_model_input[:, shift_idx:]] + [latent_model_input[:, :shift_idx]], dim=1)
#enhance-a-video
if feta_args is not None:
if feta_start_percent <= current_step_percentage <= feta_end_percent:
enable_enhance()
else:
disable_enhance()
if feta_args is not None and feta_start_percent <= current_step_percentage <= feta_end_percent:
enable_enhance()
else:
disable_enhance()
#flow-edit
if flowedit_args is not None:
sigma = t / 1000.0
@@ -3103,6 +3228,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 +3490,7 @@ NODE_CLASS_MAPPINGS = {
"WanVideoVACEEncode": WanVideoVACEEncode,
"WanVideoVACEStartToEndFrame": WanVideoVACEStartToEndFrame,
"WanVideoVACEModelSelect": WanVideoVACEModelSelect,
"WanVideoPhantomEmbeds": WanVideoPhantomEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoSampler": "WanVideo Sampler",
@@ -3397,4 +3526,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",
}
+7 -3
View File
@@ -12,6 +12,8 @@ from ..wanvideo.utils.scheduling_flow_match_lcm import FlowMatchLCMScheduler
from ..nodes import optimized_scale
from einops import rearrange
from ..enhance_a_video.globals import disable_enhance
import comfy.model_management as mm
from comfy.utils import load_torch_file, ProgressBar, common_upscale
from comfy.clip_vision import clip_preprocess, ClipVisionModel
@@ -139,7 +141,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()
@@ -305,6 +307,7 @@ class WanVideoDiffusionForcingSampler:
"end_percent": unianimate_poses["end_percent"]
}
disable_enhance() #not sure if this can work, disabling for now to avoid errors if it's enabled by another sampler
freqs = None
transformer.rope_embedder.k = None
@@ -371,6 +374,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 +414,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 +529,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)
k = k.transpose(1, 2)
v = v.transpose(1, 2)
out = sageattn_func(
q, k, v, attn_mask=attn_mask, is_causal=causal, dropout_p=dropout_p)
+117 -69
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
@@ -734,6 +748,8 @@ class WanModel(ModelMixin, ConfigMixin):
vace_layers=None,
vace_in_dim=None,
inject_sample_info=False,
add_ref_conv=False,
in_dim_ref_conv=16,
):
r"""
Initialize the diffusion model backbone.
@@ -885,9 +901,15 @@ class WanModel(ModelMixin, ConfigMixin):
if model_type == 'i2v' or model_type == 'fl2v':
self.img_emb = MLPProj(1280, dim, fl_pos_emb=model_type == 'fl2v')
#skyreels v2
if inject_sample_info:
self.fps_embedding = nn.Embedding(2, dim)
self.fps_projection = nn.Sequential(nn.Linear(dim, dim), nn.SiLU(), nn.Linear(dim, dim * 6))
#fun 1.1
if add_ref_conv:
self.ref_conv = nn.Conv2d(in_dim_ref_conv, dim, kernel_size=patch_size[1:], stride=patch_size[1:])
else:
self.ref_conv = None
def block_swap(self, blocks_to_swap, offload_txt_emb=False, offload_img_emb=False, vace_blocks_to_swap=None):
log.info(f"Swapping {blocks_to_swap + 1} transformer blocks")
@@ -944,7 +966,7 @@ class WanModel(ModelMixin, ConfigMixin):
kwargs
):
# embeddings
c = [self.vace_patch_embedding(u.unsqueeze(0)) for u in vace_context]
c = [self.vace_patch_embedding(u.unsqueeze(0).float()).to(x.dtype) for u in vace_context]
c = [u.flatten(2).transpose(1, 2) for u in c]
c = torch.cat([
torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))],
@@ -991,6 +1013,7 @@ class WanModel(ModelMixin, ConfigMixin):
camera_embed=None,
unianim_data=None,
fps_embeds=None,
fun_ref = None
):
r"""
Forward pass through the diffusion model
@@ -1032,13 +1055,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
]
@@ -1046,6 +1069,14 @@ class WanModel(ModelMixin, ConfigMixin):
[torch.tensor(u.shape[2:], dtype=torch.long) for u in x])
x = [u.flatten(2).transpose(1, 2) for u in x]
if self.ref_conv is not None and fun_ref is not None:
fun_ref = self.ref_conv(fun_ref).flatten(2).transpose(1, 2)
grid_sizes = torch.stack([torch.tensor([u[0] + 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
seq_len += fun_ref.size(1)
F += 1
x = [torch.concat([_fun_ref.unsqueeze(0), u], dim=1) for _fun_ref, u in zip(fun_ref, x)]
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long)
assert seq_lens.max() <= seq_len
x = torch.cat([
@@ -1069,39 +1100,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).to(e0.dtype)
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 +1147,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)
@@ -1143,10 +1178,13 @@ class WanModel(ModelMixin, ConfigMixin):
if self.teacache_use_coefficients:
rescale_func = np.poly1d(self.teacache_coefficients[self.teacache_mode])
temb = e if self.teacache_mode == 'e' else e0
accumulated_rel_l1_distance += rescale_func(((temb-previous_modulated_input).abs().mean() / previous_modulated_input.abs().mean()).cpu().item())
accumulated_rel_l1_distance += rescale_func((
(temb.to(device) - previous_modulated_input).abs().mean() / previous_modulated_input.abs().mean()
).cpu().item())
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 +1193,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 +1214,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,
@@ -1198,11 +1237,11 @@ class WanModel(ModelMixin, ConfigMixin):
if (data["start"] <= current_step_percentage <= data["end"]) or \
(data["end"] > 0 and current_step == 0 and current_step_percentage >= data["start"]):
vace_hints = self.forward_vace(x.to(torch.float32), data["context"], data["seq_len"], kwargs)
vace_hints = self.forward_vace(x, data["context"], data["seq_len"], kwargs)
vace_hint_list.append(vace_hints)
vace_scale_list.append(data["scale"])
else:
vace_hints = self.forward_vace(x.to(torch.float32), vace_data, seq_len, kwargs)
vace_hints = self.forward_vace(x, vace_data, seq_len, kwargs)
vace_hint_list.append(vace_hints)
vace_scale_list.append(1.0)
@@ -1216,7 +1255,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 +1263,16 @@ 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)
if self.ref_conv is not None and fun_ref is not None:
full_ref_length = fun_ref.size(1)
x = x[:, full_ref_length:]
grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
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 +1305,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
@@ -1296,7 +1340,7 @@ class TeaCacheState:
del self.states[pred_id]
def clear_all(self):
self.states.clear()
self.states = {}
self._next_pred_id = 0
def relative_l1_distance(last_tensor, current_tensor):
@@ -1304,3 +1348,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"