Merge branch 'main' into vap
This commit is contained in:
+42
-28
@@ -3,15 +3,20 @@ import torch.nn as nn
|
||||
from accelerate import init_empty_weights
|
||||
|
||||
#based on https://github.com/huggingface/diffusers/blob/main/src/diffusers/quantizers/gguf/utils.py
|
||||
def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, scale_weights=None):
|
||||
def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, scale_weights=None, compile_args=None):
|
||||
|
||||
has_children = list(model.children())
|
||||
if not has_children:
|
||||
return
|
||||
|
||||
allow_compile = False
|
||||
|
||||
for name, module in model.named_children():
|
||||
if compile_args is not None:
|
||||
allow_compile = compile_args.get("allow_unmerged_lora_compile", False)
|
||||
module_prefix = prefix + name + "."
|
||||
module_prefix = module_prefix.replace("_orig_mod.", "")
|
||||
_replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights)
|
||||
_replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights, compile_args)
|
||||
|
||||
if isinstance(module, nn.Linear) and "loras" not in module_prefix:
|
||||
in_features = state_dict[module_prefix + "weight"].shape[1]
|
||||
@@ -25,7 +30,8 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s
|
||||
out_features,
|
||||
module.bias is not None,
|
||||
compute_dtype=compute_dtype,
|
||||
scale_weight=scale_weights.get(scale_key) if scale_weights else None
|
||||
scale_weight=scale_weights.get(scale_key) if scale_weights else None,
|
||||
allow_compile=allow_compile
|
||||
)
|
||||
model._modules[name].source_cls = type(module)
|
||||
model._modules[name].requires_grad_(False)
|
||||
@@ -77,7 +83,8 @@ class CustomLinear(nn.Linear):
|
||||
bias=False,
|
||||
compute_dtype=None,
|
||||
device=None,
|
||||
scale_weight=None
|
||||
scale_weight=None,
|
||||
allow_compile=False
|
||||
) -> None:
|
||||
super().__init__(in_features, out_features, bias, device)
|
||||
self.compute_dtype = compute_dtype
|
||||
@@ -85,38 +92,28 @@ class CustomLinear(nn.Linear):
|
||||
self.step = 0
|
||||
self.scale_weight = scale_weight
|
||||
self.lora_strengths = []
|
||||
self.allow_compile = allow_compile
|
||||
|
||||
if not allow_compile:
|
||||
self._get_weight_with_lora = torch.compiler.disable()(self._get_weight_with_lora)
|
||||
|
||||
def set_lora_diffs(self, lora_diffs, device=torch.device("cpu")):
|
||||
self.lora_diffs = []
|
||||
for i, diff in enumerate(lora_diffs):
|
||||
if isinstance(diff, tuple):
|
||||
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device))
|
||||
self.register_buffer(f"lora_diff_{i}_1", diff[1].to(device))
|
||||
if len(diff) > 1:
|
||||
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype))
|
||||
self.register_buffer(f"lora_diff_{i}_1", diff[1].to(device, self.compute_dtype))
|
||||
setattr(self, f"lora_diff_{i}_2", diff[2])
|
||||
self.lora_diffs.append((f"lora_diff_{i}_0", f"lora_diff_{i}_1", f"lora_diff_{i}_2"))
|
||||
else:
|
||||
self.register_buffer(f"lora_diff_{i}", diff.to(device))
|
||||
self.lora_diffs.append(f"lora_diff_{i}")
|
||||
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype))
|
||||
self.lora_diffs.append(f"lora_diff_{i}_0")
|
||||
|
||||
def forward(self, input):
|
||||
if self.bias is not None:
|
||||
bias = self.bias.to(input)
|
||||
else:
|
||||
bias = None
|
||||
weight = self.weight.to(input)
|
||||
|
||||
if self.scale_weight is not None:
|
||||
if weight.numel() < input.numel():
|
||||
weight = weight * self.scale_weight
|
||||
else:
|
||||
input = input * self.scale_weight
|
||||
|
||||
if hasattr(self, f"lora_diff_0_0"):
|
||||
weight = self.apply_lora(weight).to(self.compute_dtype)
|
||||
|
||||
return torch.nn.functional.linear(input, weight, bias)
|
||||
|
||||
def apply_lora(self, weight):
|
||||
def _get_weight_with_lora(self, weight):
|
||||
"""Apply LoRA outside compiled region"""
|
||||
if not hasattr(self, "lora_diff_0_0"):
|
||||
return weight
|
||||
|
||||
for lora_diff_names, lora_strength in zip(self.lora_diffs, self.lora_strengths):
|
||||
if isinstance(lora_strength, list):
|
||||
lora_strength = lora_strength[self.step]
|
||||
@@ -139,6 +136,23 @@ class CustomLinear(nn.Linear):
|
||||
lora_diff = getattr(self, lora_diff_names)
|
||||
weight = weight.add(lora_diff, alpha=lora_strength)
|
||||
return weight
|
||||
|
||||
def forward(self, input):
|
||||
if self.bias is not None:
|
||||
bias = self.bias.to(input)
|
||||
else:
|
||||
bias = None
|
||||
weight = self.weight.to(input)
|
||||
|
||||
if self.scale_weight is not None:
|
||||
if weight.numel() < input.numel():
|
||||
weight = weight * self.scale_weight
|
||||
else:
|
||||
input = input * self.scale_weight
|
||||
|
||||
weight = self._get_weight_with_lora(weight)
|
||||
|
||||
return torch.nn.functional.linear(input, weight, bias)
|
||||
|
||||
def remove_lora_from_module(module):
|
||||
for name, submodule in module.named_modules():
|
||||
|
||||
+24
-11
@@ -19,7 +19,7 @@ def load_gguf(model_path):
|
||||
return parsed_parameters, reader
|
||||
|
||||
#based on https://github.com/huggingface/diffusers/blob/main/src/diffusers/quantizers/gguf/utils.py
|
||||
def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modules_to_not_convert=[], patches=None):
|
||||
def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modules_to_not_convert=[], patches=None, compile_args=None):
|
||||
def _should_convert_to_gguf(state_dict, prefix):
|
||||
weight_key = prefix + "weight"
|
||||
return weight_key in state_dict and isinstance(state_dict[weight_key], GGUFParameter)
|
||||
@@ -27,10 +27,14 @@ def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modul
|
||||
has_children = list(model.children())
|
||||
if not has_children:
|
||||
return
|
||||
|
||||
allow_compile = False
|
||||
|
||||
for name, module in model.named_children():
|
||||
if compile_args is not None:
|
||||
allow_compile = compile_args.get("allow_unmerged_lora_compile", False)
|
||||
module_prefix = prefix + name + "."
|
||||
_replace_with_gguf_linear(module, compute_dtype, state_dict, module_prefix, modules_to_not_convert, patches)
|
||||
_replace_with_gguf_linear(module, compute_dtype, state_dict, module_prefix, modules_to_not_convert, patches, compile_args)
|
||||
|
||||
if (
|
||||
isinstance(module, nn.Linear)
|
||||
@@ -46,7 +50,8 @@ def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modul
|
||||
in_features,
|
||||
out_features,
|
||||
module.bias is not None,
|
||||
compute_dtype=compute_dtype
|
||||
compute_dtype=compute_dtype,
|
||||
allow_compile=allow_compile
|
||||
)
|
||||
|
||||
model._modules[name].source_cls = type(module)
|
||||
@@ -95,35 +100,43 @@ class GGUFLinear(nn.Linear):
|
||||
bias=False,
|
||||
compute_dtype=None,
|
||||
device=None,
|
||||
allow_compile=False
|
||||
) -> None:
|
||||
super().__init__(in_features, out_features, bias, device)
|
||||
self.compute_dtype = compute_dtype
|
||||
self.lora_diffs = []
|
||||
self.lora_strengths = []
|
||||
self.step = 0
|
||||
self.allow_compile = allow_compile
|
||||
|
||||
if not allow_compile:
|
||||
self._get_weight_with_lora = torch.compiler.disable()(self._get_weight_with_lora)
|
||||
|
||||
def forward(self, inputs):
|
||||
weight = dequantize_gguf_tensor(self.weight).to(self.compute_dtype)
|
||||
bias = self.bias.to(self.compute_dtype) if self.bias is not None else None
|
||||
|
||||
if hasattr(self, f"lora_diff_0_0"):
|
||||
weight = self.apply_lora(weight).to(self.compute_dtype)
|
||||
weight = self._get_weight_with_lora(weight)#.to(self.compute_dtype)
|
||||
|
||||
return torch.nn.functional.linear(inputs, weight, bias)
|
||||
|
||||
def set_lora_diffs(self, lora_diffs, device=torch.device("cpu")):
|
||||
self.lora_diffs = []
|
||||
for i, diff in enumerate(lora_diffs):
|
||||
if isinstance(diff, tuple):
|
||||
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device))
|
||||
self.register_buffer(f"lora_diff_{i}_1", diff[1].to(device))
|
||||
if len(diff) > 1:
|
||||
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype))
|
||||
self.register_buffer(f"lora_diff_{i}_1", diff[1].to(device, self.compute_dtype))
|
||||
setattr(self, f"lora_diff_{i}_2", diff[2])
|
||||
self.lora_diffs.append((f"lora_diff_{i}_0", f"lora_diff_{i}_1", f"lora_diff_{i}_2"))
|
||||
else:
|
||||
self.register_buffer(f"lora_diff_{i}", diff.to(device))
|
||||
self.lora_diffs.append(f"lora_diff_{i}")
|
||||
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype))
|
||||
self.lora_diffs.append(f"lora_diff_{i}_0")
|
||||
|
||||
def apply_lora(self, weight):
|
||||
def _get_weight_with_lora(self, weight):
|
||||
"""Apply LoRA outside compiled region"""
|
||||
if not hasattr(self, "lora_diff_0_0"):
|
||||
return weight
|
||||
|
||||
for lora_diff_names, lora_strength in zip(self.lora_diffs, self.lora_strengths):
|
||||
if isinstance(lora_strength, list):
|
||||
lora_strength = lora_strength[self.step]
|
||||
|
||||
@@ -198,7 +198,7 @@ class AudioProjModel(nn.Module):
|
||||
context_tokens = self.proj3(audio_embeds_c).reshape(batch_size_c*N_t, self.context_tokens, self.output_dim)
|
||||
|
||||
# normalization and reshape
|
||||
context_tokens = self.norm(context_tokens)
|
||||
context_tokens = self.norm(context_tokens.to(self.norm.weight.dtype)).to(context_tokens.dtype)
|
||||
context_tokens = rearrange(context_tokens, "(bz f) m c -> bz f m c", f=video_length)
|
||||
|
||||
return context_tokens
|
||||
|
||||
@@ -330,6 +330,7 @@ class WanVideoTorchCompileSettings:
|
||||
"optional": {
|
||||
"dynamo_recompile_limit": ("INT", {"default": 128, "min": 0, "max": 1024, "step": 1, "tooltip": "torch._dynamo.config.recompile_limit"}),
|
||||
"force_parameter_static_shapes": ("BOOLEAN", {"default": False, "tooltip": "torch._dynamo.config.force_parameter_static_shapes"}),
|
||||
"allow_unmerged_lora_compile": ("BOOLEAN", {"default": False, "tooltip": "Allow LoRA application to be compiled with torch.compile to avoid graph breaks, causes issues with some LoRAs, mostly dynamic ones"}),
|
||||
},
|
||||
}
|
||||
RETURN_TYPES = ("WANCOMPILEARGS",)
|
||||
@@ -338,7 +339,8 @@ class WanVideoTorchCompileSettings:
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
DESCRIPTION = "torch.compile settings, when connected to the model loader, torch.compile of the selected layers is attempted. Requires Triton and torch > 2.7.0 is recommended"
|
||||
|
||||
def set_args(self, backend, fullgraph, mode, dynamic, dynamo_cache_size_limit, compile_transformer_blocks_only, dynamo_recompile_limit=128, force_parameter_static_shapes=True):
|
||||
def set_args(self, backend, fullgraph, mode, dynamic, dynamo_cache_size_limit, compile_transformer_blocks_only, dynamo_recompile_limit=128,
|
||||
force_parameter_static_shapes=True, allow_unmerged_lora_compile=False):
|
||||
|
||||
compile_args = {
|
||||
"backend": backend,
|
||||
@@ -349,6 +351,7 @@ class WanVideoTorchCompileSettings:
|
||||
"dynamo_recompile_limit": dynamo_recompile_limit,
|
||||
"compile_transformer_blocks_only": compile_transformer_blocks_only,
|
||||
"force_parameter_static_shapes": force_parameter_static_shapes,
|
||||
"allow_unmerged_lora_compile": allow_unmerged_lora_compile,
|
||||
}
|
||||
|
||||
return (compile_args, )
|
||||
@@ -782,7 +785,7 @@ def rename_fuser_block(name):
|
||||
return new_name
|
||||
|
||||
def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
transformer_load_device=None, block_swap_args=None, gguf=False, reader=None, patcher=None):
|
||||
transformer_load_device=None, block_swap_args=None, gguf=False, reader=None, patcher=None, compile_args=None):
|
||||
params_to_keep = {"time_in", "patch_embedding", "time_", "modulation", "text_embedding",
|
||||
"adapter", "add", "ref_conv", "casual_audio_encoder", "cond_encoder", "frame_packer", "audio_proj_glob", "face_encoder", "fuser_block"}
|
||||
param_count = sum(1 for _ in transformer.named_parameters())
|
||||
@@ -837,7 +840,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
|
||||
if not getattr(transformer, "gguf_patched", False):
|
||||
transformer = _replace_with_gguf_linear(
|
||||
transformer, base_dtype, sd, patches=patcher.patches
|
||||
transformer, base_dtype, sd, patches=patcher.patches, compile_args=compile_args
|
||||
)
|
||||
transformer.gguf_patched = True
|
||||
else:
|
||||
@@ -1535,7 +1538,7 @@ class WanVideoModelLoader:
|
||||
transformer.patched_linear = False
|
||||
sd = None
|
||||
elif "scaled" in quantization or lora is not None:
|
||||
transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights)
|
||||
transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights, compile_args=compile_args)
|
||||
transformer.patched_linear = True
|
||||
|
||||
if "fast" in quantization:
|
||||
|
||||
+11
-4
@@ -189,14 +189,21 @@ class WanVideoSampler:
|
||||
vae_upscale_factor = 16 if is_5b else 8
|
||||
|
||||
# Load weights
|
||||
if not transformer.patched_linear and patcher.model["sd"] is not None and len(patcher.patches) != 0:
|
||||
transformer = _replace_linear(transformer, dtype, patcher.model["sd"])
|
||||
if transformer.audio_model is not None:
|
||||
for block in transformer.blocks:
|
||||
if hasattr(block, 'audio_block'):
|
||||
block.audio_block = None
|
||||
|
||||
if not transformer.patched_linear and patcher.model["sd"] is not None and len(patcher.patches) != 0 and gguf_reader is None:
|
||||
transformer = _replace_linear(transformer, dtype, patcher.model["sd"], model["compile_args"])
|
||||
transformer.patched_linear = True
|
||||
if patcher.model["sd"] is not None and gguf_reader is None:
|
||||
load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, base_dtype=dtype, transformer_load_device=device, block_swap_args=block_swap_args)
|
||||
load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, base_dtype=dtype, transformer_load_device=device,
|
||||
block_swap_args=block_swap_args, compile_args=model["compile_args"])
|
||||
|
||||
if gguf_reader is not None: #handle GGUF
|
||||
load_weights(transformer, patcher.model["sd"], base_dtype=dtype, transformer_load_device=device, patcher=patcher, gguf=True, reader=gguf_reader, block_swap_args=block_swap_args)
|
||||
load_weights(transformer, patcher.model["sd"], base_dtype=dtype, transformer_load_device=device, patcher=patcher, gguf=True,
|
||||
reader=gguf_reader, block_swap_args=block_swap_args, compile_args=model["compile_args"])
|
||||
set_lora_params_gguf(transformer, patcher.patches)
|
||||
transformer.patched_linear = True
|
||||
elif len(patcher.patches) != 0: #handle patched linear layers (unmerged loras, fp8 scaled)
|
||||
|
||||
@@ -74,6 +74,13 @@ HuMo: https://github.com/Phantom-video/HuMo
|
||||
|
||||
WanAnimate: https://github.com/Wan-Video/Wan2.2/tree/main/wan/modules/animate
|
||||
|
||||
Lynx: https://github.com/bytedance/lynx
|
||||
|
||||
|
||||
Not exactly Wan model, but close enough to work with the code base:
|
||||
|
||||
LongCat-Video: https://meituan-longcat.github.io/LongCat-Video/
|
||||
|
||||
|
||||
Examples:
|
||||
---
|
||||
|
||||
+1
-1
@@ -172,7 +172,7 @@ class WanVideoDiffusionForcingSampler:
|
||||
|
||||
# Load weights
|
||||
if not transformer.patched_linear and patcher.model["sd"] is not None and len(patcher.patches) != 0:
|
||||
transformer = _replace_linear(transformer, dtype, patcher.model["sd"])
|
||||
transformer = _replace_linear(transformer, dtype, patcher.model["sd"], model["compile_args"])
|
||||
transformer.patched_linear = True
|
||||
if patcher.model["sd"] is not None and gguf_reader is None:
|
||||
load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, base_dtype=dtype, transformer_load_device=device, block_swap_args=block_swap_args)
|
||||
|
||||
+52
-54
@@ -29,14 +29,14 @@ from comfy import model_management as mm
|
||||
__all__ = ['WanModel']
|
||||
|
||||
class AdaLayerNorm(nn.Module):
|
||||
def __init__(self, embedding_dim, output_dim=None, norm_elementwise_affine=False, norm_eps=1e-5, dtype=None, device=None, operations=None):
|
||||
def __init__(self, embedding_dim, output_dim=None, norm_elementwise_affine=False, norm_eps=1e-5):
|
||||
super().__init__()
|
||||
|
||||
output_dim = output_dim or embedding_dim * 2
|
||||
|
||||
self.silu = nn.SiLU()
|
||||
self.linear = operations.Linear(embedding_dim, output_dim, dtype=dtype, device=device)
|
||||
self.norm = operations.LayerNorm(output_dim // 2, norm_eps, norm_elementwise_affine, dtype=dtype, device=device)
|
||||
self.linear = nn.Linear(embedding_dim, output_dim)
|
||||
self.norm = nn.LayerNorm(output_dim // 2, norm_eps, norm_elementwise_affine)
|
||||
|
||||
def forward(self, x, temb):
|
||||
temb = self.linear(self.silu(temb))
|
||||
@@ -672,7 +672,7 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
|
||||
if is_longcat:
|
||||
if num_cond_latents is not None and num_cond_latents > 0:
|
||||
num_cond_latents_thw = num_cond_latents * (s // grid_sizes[0][0])
|
||||
num_cond_latents_thw = num_cond_latents * (s // num_latent_frames)
|
||||
x = x[:, num_cond_latents_thw:]
|
||||
q = self.norm_q(self.q(x).view(b, -1, n, d).to(self.norm_q.weight.dtype)).to(x.dtype)
|
||||
else:
|
||||
@@ -1042,7 +1042,7 @@ class WanAttentionBlock(nn.Module):
|
||||
return torch.addcmul(shift_msa, norm_x, 1 + scale_msa)
|
||||
|
||||
def ffn_chunked(self, x, shift_mlp, scale_mlp, num_chunks=4):
|
||||
modulated_input = torch.addcmul(shift_mlp, self.norm2(x), 1 + scale_mlp)
|
||||
modulated_input = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp).to(x.dtype)
|
||||
|
||||
result = torch.empty_like(x)
|
||||
seq_len = modulated_input.shape[1]
|
||||
@@ -1241,7 +1241,7 @@ class WanAttentionBlock(nn.Module):
|
||||
full_v = torch.cat([v, v_ip], dim=1)
|
||||
y = self.self_attn.forward(q, full_k, full_v, seq_lens)
|
||||
elif is_longcat and num_cond_latents is not None and num_cond_latents > 0:
|
||||
num_cond_latents_thw = num_cond_latents * (N // grid_sizes[0][0])
|
||||
num_cond_latents_thw = num_cond_latents * (N // num_latent_frames)
|
||||
# process the condition tokens
|
||||
x_cond = self.self_attn.forward(
|
||||
q[:, :num_cond_latents_thw].contiguous(),
|
||||
@@ -1338,6 +1338,7 @@ class WanAttentionBlock(nn.Module):
|
||||
x_mot_ref = x_mot_ref.to(input_dtype)
|
||||
# MultiTalk
|
||||
if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock):
|
||||
x_audio = self.audio_cross_attn(self.norm_x(x.to(self.norm_x.weight.dtype)).to(input_dtype), encoder_hidden_states=multitalk_audio_embedding,
|
||||
x_audio = self.audio_cross_attn(self.norm_x(x.to(self.norm_x.weight.dtype)).to(input_dtype), encoder_hidden_states=multitalk_audio_embedding,
|
||||
shape=grid_sizes[0], x_ref_attn_map=x_ref_attn_map, human_num=human_num)
|
||||
x = x.add(x_audio, alpha=audio_scale)
|
||||
@@ -1352,39 +1353,39 @@ class WanAttentionBlock(nn.Module):
|
||||
x = self.audio_cross_attn_wrapper(x, humo_audio_input, grid_sizes, humo_audio_scale)
|
||||
|
||||
|
||||
# ffn
|
||||
if self.rope_func == "comfy_chunked":
|
||||
x_ffn = self.ffn_chunked(x, shift_mlp, scale_mlp)
|
||||
else:
|
||||
if zero_timestep:
|
||||
norm2_x = self.norm2(x.to(self.norm2.weight.dtype)).to(input_dtype)
|
||||
parts = []
|
||||
for i in range(2):
|
||||
parts.append(norm2_x[:, self.seg_idx[i]:self.seg_idx[i + 1]] *
|
||||
(1 + scale_mlp[:, i:i + 1]) + shift_mlp[:, i:i + 1])
|
||||
norm2_x = torch.cat(parts, dim=1)
|
||||
x_ffn = self.ffn(norm2_x)
|
||||
else:
|
||||
if not is_longcat:
|
||||
mod_x = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp)
|
||||
else:
|
||||
mod_x = torch.addcmul(shift_mlp, self.norm2(x.view(B, -1, N//T, C).float()), 1 + scale_mlp).view(B, -1, C)
|
||||
x_ffn = self.ffn(mod_x.to(input_dtype))
|
||||
del shift_mlp, scale_mlp
|
||||
|
||||
# gate_mlp
|
||||
# ffn
|
||||
if self.rope_func == "comfy_chunked":
|
||||
x_ffn = self.ffn_chunked(x, shift_mlp, scale_mlp)
|
||||
else:
|
||||
if zero_timestep:
|
||||
z = []
|
||||
norm2_x = self.norm2(x)
|
||||
parts = []
|
||||
for i in range(2):
|
||||
z.append(x_ffn[:, self.seg_idx[i]:self.seg_idx[i + 1]] * gate_mlp[:, i:i + 1])
|
||||
x_ffn = torch.cat(z, dim=1)
|
||||
x = x.add(x_ffn)
|
||||
parts.append(norm2_x[:, self.seg_idx[i]:self.seg_idx[i + 1]] *
|
||||
(1 + scale_mlp[:, i:i + 1]) + shift_mlp[:, i:i + 1])
|
||||
norm2_x = torch.cat(parts, dim=1)
|
||||
x_ffn = self.ffn(norm2_x)
|
||||
else:
|
||||
if not is_longcat:
|
||||
x = x.addcmul(x_ffn.to(gate_mlp.dtype), gate_mlp).to(input_dtype)
|
||||
mod_x = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp)
|
||||
else:
|
||||
x = x + (gate_mlp * x_ffn.view(B, -1, N//T, C).float()).to(input_dtype).view(B, -1, C)
|
||||
del gate_mlp
|
||||
mod_x = torch.addcmul(shift_mlp, self.norm2(x.view(B, -1, N//T, C).float()), 1 + scale_mlp).view(B, -1, C)
|
||||
x_ffn = self.ffn(mod_x.to(input_dtype))
|
||||
del shift_mlp, scale_mlp
|
||||
|
||||
# gate_mlp
|
||||
if zero_timestep:
|
||||
z = []
|
||||
for i in range(2):
|
||||
z.append(x_ffn[:, self.seg_idx[i]:self.seg_idx[i + 1]] * gate_mlp[:, i:i + 1])
|
||||
x_ffn = torch.cat(z, dim=1)
|
||||
x = x.add(x_ffn)
|
||||
else:
|
||||
if not is_longcat:
|
||||
x = x.addcmul(x_ffn.to(gate_mlp.dtype), gate_mlp).to(input_dtype)
|
||||
else:
|
||||
x = x + (gate_mlp * x_ffn.view(B, -1, N//T, C).float()).to(input_dtype).view(B, -1, C)
|
||||
del gate_mlp
|
||||
|
||||
if x_ip is not None: #stand-in
|
||||
x_ip = x_ip.addcmul(y_ip, gate_msa_ip)
|
||||
@@ -1699,7 +1700,7 @@ class AudioInjector_WAN(nn.Module):
|
||||
if enable_adain:
|
||||
self.injector_adain_layers = nn.ModuleList([
|
||||
AdaLayerNorm(
|
||||
output_dim=dim * 2, embedding_dim=adain_dim, chunk_dim=1)
|
||||
output_dim=dim * 2, embedding_dim=adain_dim)
|
||||
for _ in range(audio_injector_id)
|
||||
])
|
||||
if need_adain_ont:
|
||||
@@ -2662,7 +2663,17 @@ class WanModel(torch.nn.Module):
|
||||
|
||||
e = e.to(self.offload_device, non_blocking=self.use_non_blocking)
|
||||
|
||||
|
||||
# clip vision embedding
|
||||
clip_embed = None
|
||||
if clip_fea is not None and hasattr(self, "img_emb"):
|
||||
clip_fea = clip_fea.to(self.main_device)
|
||||
if self.offload_img_emb:
|
||||
self.img_emb.to(self.main_device)
|
||||
clip_embed = self.img_emb(clip_fea) # bs x 257 x dim
|
||||
#context = torch.concat([context_clip, context], dim=1)
|
||||
if self.offload_img_emb:
|
||||
self.img_emb.to(self.offload_device, non_blocking=self.use_non_blocking)
|
||||
|
||||
#context (text embedding)
|
||||
if hasattr(self, "text_embedding") and context != []:
|
||||
text_embed_dtype = self.text_embedding[0].weight.dtype
|
||||
@@ -2707,24 +2718,13 @@ class WanModel(torch.nn.Module):
|
||||
|
||||
if self.offload_txt_emb:
|
||||
self.text_embedding.to(self.offload_device, non_blocking=self.use_non_blocking)
|
||||
|
||||
seq_chunks = max(context.shape[0], clip_embed.shape[0] if clip_embed is not None else 0)
|
||||
chunked_self_attention = seq_chunks > 1 and current_step in self.video_attention_split_steps
|
||||
else:
|
||||
context = None
|
||||
|
||||
clip_embed = clip_embed_mot_ref = None
|
||||
if clip_fea is not None and hasattr(self, "img_emb"):
|
||||
clip_fea = clip_fea.to(self.main_device)
|
||||
if self.offload_img_emb:
|
||||
self.img_emb.to(self.main_device)
|
||||
clip_embed = self.img_emb(clip_fea) # bs x 257 x dim
|
||||
if self.offload_img_emb:
|
||||
self.img_emb.to(self.offload_device, non_blocking=self.use_non_blocking)
|
||||
if mot_ref_clip_embeds is not None:
|
||||
mot_ref_clip_embeds = mot_ref_clip_embeds.to(self.main_device)
|
||||
if self.offload_img_emb:
|
||||
self.img_emb.to(self.main_device)
|
||||
clip_embed_mot_ref = self.img_emb_mot_ref(mot_ref_clip_embeds) # bs x 257 x dim
|
||||
if self.offload_img_emb:
|
||||
self.img_emb.to(self.offload_device, non_blocking=self.use_non_blocking)
|
||||
chunked_self_attention = False
|
||||
seq_chunks = 0
|
||||
|
||||
# MultiTalk
|
||||
if multitalk_audio is not None:
|
||||
@@ -2902,8 +2902,6 @@ class WanModel(torch.nn.Module):
|
||||
dwpose_emb = rearrange(unianim_data['dwpose'], 'b c f h w -> b (f h w) c').contiguous()
|
||||
x.add_(dwpose_emb, alpha=unianim_data['strength'])
|
||||
|
||||
seq_chunks = max(context.shape[0], clip_embed.shape[0] if clip_embed is not None else 0)
|
||||
chunked_self_attention = seq_chunks > 1 and current_step in self.video_attention_split_steps
|
||||
# arguments
|
||||
kwargs = dict(
|
||||
e=e0,
|
||||
|
||||
@@ -60,7 +60,7 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
|
||||
if flowedit_args: #seems to work better
|
||||
timesteps, _ = retrieve_timesteps(sample_scheduler, device=device, sigmas=get_sampling_sigmas(steps, shift))
|
||||
else:
|
||||
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas)
|
||||
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas[:-1].tolist() if sigmas is not None else None)
|
||||
elif 'dpm' in scheduler:
|
||||
if 'sde' in scheduler:
|
||||
algorithm_type = "sde-dpmsolver++"
|
||||
|
||||
Reference in New Issue
Block a user