Fix UniAnimate and MultiTalk loading
This commit is contained in:
+1
-8
@@ -22,14 +22,8 @@ class MultiTalkModelLoader:
|
||||
|
||||
def loadmodel(self, model, base_precision=None):
|
||||
from .multitalk import AudioProjModel
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
model_path = folder_paths.get_full_path_or_raise("diffusion_models", model)
|
||||
if model_path.endswith(".gguf"):
|
||||
from diffusers.models.model_loading_utils import load_gguf_checkpoint
|
||||
sd = load_gguf_checkpoint(model_path)
|
||||
else:
|
||||
sd = load_torch_file(model_path, device=offload_device, safe_load=True)
|
||||
|
||||
audio_window=5
|
||||
intermediate_dim=512
|
||||
@@ -50,8 +44,7 @@ class MultiTalkModelLoader:
|
||||
|
||||
multitalk = {
|
||||
"proj_model": multitalk_proj_model,
|
||||
"sd": sd,
|
||||
"is_gguf": model_path.endswith(".gguf"),
|
||||
"model_path": model_path,
|
||||
"model_type": "InfiniteTalk" if "infinite" in model.lower() else "MultiTalk",
|
||||
}
|
||||
|
||||
|
||||
@@ -1701,20 +1701,21 @@ class WanVideoSampler:
|
||||
is_5b = transformer.out_dim == 48
|
||||
vae_upscale_factor = 16 if is_5b else 8
|
||||
|
||||
# Load weights
|
||||
if transformer.patched_linear 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)
|
||||
|
||||
if gguf_reader is not None:
|
||||
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)
|
||||
set_lora_params_gguf(transformer, patcher.patches)
|
||||
transformer.patched_linear = True
|
||||
elif len(patcher.patches) != 0 and transformer.patched_linear:
|
||||
elif len(patcher.patches) != 0 and transformer.patched_linear: #handle patched linear layers (unmerged loras, fp8 scaled)
|
||||
log.info(f"Using {len(patcher.patches)} LoRA weight patches for WanVideo model")
|
||||
if not merge_loras and fp8_matmul:
|
||||
raise NotImplementedError("FP8 matmul with unmerged LoRAs is not supported")
|
||||
set_lora_params(transformer, patcher.patches)
|
||||
else:
|
||||
remove_lora_from_module(transformer)
|
||||
remove_lora_from_module(transformer) #clear possible unmerged lora weights
|
||||
|
||||
transformer.lora_scheduling_enabled = transformer_options.get("lora_scheduling_enabled", False)
|
||||
|
||||
@@ -2386,16 +2387,18 @@ class WanVideoSampler:
|
||||
import copy
|
||||
sample_scheduler_flipped = copy.deepcopy(sample_scheduler)
|
||||
|
||||
#rope
|
||||
# Rotary positional embeddings (RoPE)
|
||||
|
||||
# RoPE base freq scaling as used with CineScale
|
||||
ntk_alphas = [1.0, 1.0, 1.0]
|
||||
if isinstance(rope_function, dict):
|
||||
ntk_alphas = rope_function["ntk_scale_f"], rope_function["ntk_scale_h"], rope_function["ntk_scale_w"]
|
||||
rope_function = rope_function["rope_function"]
|
||||
|
||||
|
||||
freqs = None
|
||||
transformer.rope_embedder.k = None
|
||||
transformer.rope_embedder.num_frames = None
|
||||
if "default" in rope_function or bidirectional_sampling:
|
||||
if "default" in rope_function or bidirectional_sampling: # original RoPE
|
||||
d = transformer.dim // transformer.num_heads
|
||||
freqs = torch.cat([
|
||||
rope_params(1024, d - 4 * (d // 6), L_test=latent_video_length, k=riflex_freq_index),
|
||||
@@ -2403,7 +2406,7 @@ class WanVideoSampler:
|
||||
rope_params(1024, 2 * (d // 6))
|
||||
],
|
||||
dim=1)
|
||||
elif "comfy" in rope_function:
|
||||
elif "comfy" in rope_function: # comfy's rope
|
||||
transformer.rope_embedder.k = riflex_freq_index
|
||||
transformer.rope_embedder.num_frames = latent_video_length
|
||||
|
||||
@@ -2565,37 +2568,37 @@ class WanVideoSampler:
|
||||
|
||||
|
||||
base_params = {
|
||||
'seq_len': seq_len,
|
||||
'device': device,
|
||||
'freqs': freqs,
|
||||
't': timestep,
|
||||
'current_step': idx,
|
||||
'last_step': len(timesteps) - 1 == idx,
|
||||
'control_lora_enabled': control_lora_enabled,
|
||||
'enhance_enabled': enhance_enabled,
|
||||
'camera_embed': camera_embed,
|
||||
'unianim_data': unianim_data,
|
||||
'fun_ref': fun_ref_input if fun_ref_image is not None else None,
|
||||
'fun_camera': control_camera_input if control_camera_latents is not None else None,
|
||||
'audio_proj': audio_proj if fantasytalking_embeds is not None else None,
|
||||
'audio_scale': audio_scale,
|
||||
"pcd_data": pcd_data_input,
|
||||
"controlnet": controlnet,
|
||||
"add_cond": add_cond_input,
|
||||
"nag_params": text_embeds.get("nag_params", {}),
|
||||
"nag_context": text_embeds.get("nag_prompt_embeds", None),
|
||||
"multitalk_audio": multitalk_audio_input if multitalk_audio_embedding is not None else None,
|
||||
"ref_target_masks": ref_target_masks if multitalk_audio_embedding is not None else None,
|
||||
"inner_t": [shot_len] if shot_len else None,
|
||||
"standin_input": standin_input,
|
||||
"fantasy_portrait_input": fantasy_portrait_input,
|
||||
"phantom_ref": phantom_ref,
|
||||
"reverse_time": reverse_time,
|
||||
"ntk_alphas": ntk_alphas,
|
||||
"mtv_motion_tokens": mtv_motion_tokens if mtv_input is not None else None,
|
||||
"mtv_motion_rotary_emb": mtv_motion_rotary_emb if mtv_input is not None else None,
|
||||
"mtv_strength": mtv_strength[idx] if mtv_input is not None else 1.0,
|
||||
"mtv_freqs": mtv_freqs if mtv_input is not None else None,
|
||||
'seq_len': seq_len, # sequence length
|
||||
'device': device, # main device
|
||||
'freqs': freqs, # rope freqs
|
||||
't': timestep, # current timestep
|
||||
'current_step': idx, # current step
|
||||
'last_step': len(timesteps) - 1 == idx, # is last step
|
||||
'control_lora_enabled': control_lora_enabled, # control lora toggle for patch embed selection
|
||||
'enhance_enabled': enhance_enabled, # enhance-a-video toggle
|
||||
'camera_embed': camera_embed, # recammaster embedding
|
||||
'unianim_data': unianim_data, # unianimate input
|
||||
'fun_ref': fun_ref_input if fun_ref_image is not None else None, # Fun model reference latent
|
||||
'fun_camera': control_camera_input if control_camera_latents is not None else None, # Fun model camera embed
|
||||
'audio_proj': audio_proj if fantasytalking_embeds is not None else None, # FantasyTalking audio projection
|
||||
'audio_scale': audio_scale, # FantasyTalking audio scale
|
||||
"pcd_data": pcd_data_input, # Uni3C input
|
||||
"controlnet": controlnet, # TheDenk's controlnet input
|
||||
"add_cond": add_cond_input, # additional conditioning input
|
||||
"nag_params": text_embeds.get("nag_params", {}), # normalized attention guidance
|
||||
"nag_context": text_embeds.get("nag_prompt_embeds", None), # normalized attention guidance context
|
||||
"multitalk_audio": multitalk_audio_input if multitalk_audio_embedding is not None else None, # Multi/InfiniteTalk audio input
|
||||
"ref_target_masks": ref_target_masks if multitalk_audio_embedding is not None else None, # Multi/InfiniteTalk reference target masks
|
||||
"inner_t": [shot_len] if shot_len else None, # inner timestep for EchoShot
|
||||
"standin_input": standin_input, # Stand-in reference input
|
||||
"fantasy_portrait_input": fantasy_portrait_input, # Fantasy portrait input
|
||||
"phantom_ref": phantom_ref, # Phantom reference input
|
||||
"reverse_time": reverse_time, # Reverse RoPE toggle
|
||||
"ntk_alphas": ntk_alphas, # RoPE freq scaling values
|
||||
"mtv_motion_tokens": mtv_motion_tokens if mtv_input is not None else None, # MTV-Crafter motion tokens
|
||||
"mtv_motion_rotary_emb": mtv_motion_rotary_emb if mtv_input is not None else None, # MTV-Crafter RoPE
|
||||
"mtv_strength": mtv_strength[idx] if mtv_input is not None else 1.0, # MTV-Crafter scaling
|
||||
"mtv_freqs": mtv_freqs if mtv_input is not None else None, # MTV-Crafter extra RoPE freqs
|
||||
}
|
||||
|
||||
batch_size = 1
|
||||
|
||||
+90
-189
@@ -741,6 +741,13 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
log.info("Using GGUF to load and assign model weights to device...")
|
||||
|
||||
# Prepare sd from GGUF readers
|
||||
|
||||
# UniAnimate embedding weight workaround
|
||||
unianimate_sd = {}
|
||||
for key in sd.keys():
|
||||
if "dwpose_embedding" in key or "randomref_embedding_pose" in key:
|
||||
unianimate_sd[key] = sd[key]
|
||||
|
||||
sd = {}
|
||||
all_tensors = []
|
||||
for r in reader:
|
||||
@@ -769,6 +776,8 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
is_gguf_quant = tensor.tensor_type not in [GGMLQuantizationType.F32, GGMLQuantizationType.F16]
|
||||
weights = torch.from_numpy(tensor.data.copy()).to(load_device)
|
||||
sd[tensor.name] = GGUFParameter(weights, quant_type=tensor.tensor_type) if is_gguf_quant else weights
|
||||
sd.update(unianimate_sd)
|
||||
del unianimate_sd
|
||||
|
||||
if not getattr(transformer, "gguf_patched", False):
|
||||
transformer = _replace_with_gguf_linear(
|
||||
@@ -795,7 +804,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
except Exception:
|
||||
block_idx = None
|
||||
|
||||
if "loras" in name or "dwpose" in name or "randomref" in name:
|
||||
if "loras" in name:
|
||||
continue
|
||||
|
||||
# GGUF: skip GGUFParameter params
|
||||
@@ -820,7 +829,6 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
elif vace_block_idx is not None:
|
||||
if vace_block_idx >= len(transformer.vace_blocks) - block_swap_args.get("vace_blocks_to_swap", 0):
|
||||
load_device = offload_device
|
||||
|
||||
# Set tensor to device
|
||||
set_module_tensor_to_device(transformer, name, device=load_device, dtype=dtype_to_use, value=sd[name.replace("_orig_mod.", "")])
|
||||
cnt += 1
|
||||
@@ -869,6 +877,7 @@ def patch_stand_in_lora(transformer, lora_sd, transformer_load_device, base_dtyp
|
||||
param.data.copy_(lora_sd["diffusion_model." + name].to(param.device, dtype=param.dtype))
|
||||
|
||||
def add_lora_weights(patcher, lora, base_dtype, merge_loras=False):
|
||||
unianimate_sd = None
|
||||
#spacepxl's control LoRA patch
|
||||
for l in lora:
|
||||
log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}")
|
||||
@@ -885,7 +894,7 @@ def add_lora_weights(patcher, lora, base_dtype, merge_loras=False):
|
||||
if "dwpose_embedding.0.weight" in lora_sd: #unianimate
|
||||
from .unianimate.nodes import update_transformer
|
||||
log.info("Unianimate LoRA detected, patching model...")
|
||||
patcher.model.diffusion_model = update_transformer(patcher.model.diffusion_model, lora_sd)
|
||||
patcher.model.diffusion_model, unianimate_sd = update_transformer(patcher.model.diffusion_model, lora_sd)
|
||||
|
||||
lora_sd = standardize_lora_key_format(lora_sd)
|
||||
|
||||
@@ -906,7 +915,7 @@ def add_lora_weights(patcher, lora, base_dtype, merge_loras=False):
|
||||
patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0)
|
||||
|
||||
del lora_sd
|
||||
return patcher, control_lora
|
||||
return patcher, control_lora, unianimate_sd
|
||||
|
||||
#region Model loading
|
||||
class WanVideoModelLoader:
|
||||
@@ -1033,15 +1042,18 @@ class WanVideoModelLoader:
|
||||
if "vace_blocks.0.after_proj.weight" in sd and not "patch_embedding.weight" in sd:
|
||||
raise ValueError("You are attempting to load a VACE module as a WanVideo model, instead you should use the vace_model input and matching T2V base model")
|
||||
|
||||
# currently this can be VAE or MTV-Crafter weights
|
||||
if extra_model is not None:
|
||||
if gguf:
|
||||
if not extra_model["path"].endswith(".gguf"):
|
||||
raise ValueError("With GGUF main model the extra model must also be a GGUF quantized, if the main model already has extra included, you can disconnect the extra module loader")
|
||||
extra_sd, extra_reader = load_gguf(extra_model["path"])
|
||||
gguf_reader.append(extra_reader)
|
||||
del extra_reader
|
||||
else:
|
||||
extra_sd = load_torch_file(extra_model["path"], device=transformer_load_device, safe_load=True)
|
||||
sd.update(extra_sd)
|
||||
del extra_sd
|
||||
|
||||
first_key = next(iter(sd))
|
||||
if first_key.startswith("model.diffusion_model."):
|
||||
@@ -1206,40 +1218,54 @@ class WanVideoModelLoader:
|
||||
log.info("FantasyPortrait model detected, patching model...")
|
||||
context_dim = fantasyportrait_model["sd"]["ip_adapter.blocks.0.cross_attn.ip_adapter_single_stream_k_proj.weight"].shape[1]
|
||||
|
||||
for block in transformer.blocks:
|
||||
block.cross_attn.ip_adapter_single_stream_k_proj = nn.Linear(context_dim, dim, bias=False)
|
||||
block.cross_attn.ip_adapter_single_stream_v_proj = nn.Linear(context_dim, dim, bias=False)
|
||||
with init_empty_weights():
|
||||
for block in transformer.blocks:
|
||||
block.cross_attn.ip_adapter_single_stream_k_proj = nn.Linear(context_dim, dim, bias=False)
|
||||
block.cross_attn.ip_adapter_single_stream_v_proj = nn.Linear(context_dim, dim, bias=False)
|
||||
ip_adapter_sd = {}
|
||||
for k, v in fantasyportrait_model["sd"].items():
|
||||
if k.startswith("ip_adapter."):
|
||||
ip_adapter_sd[k.replace("ip_adapter.", "")] = v
|
||||
sd.update(ip_adapter_sd)
|
||||
del ip_adapter_sd
|
||||
|
||||
if multitalk_model is not None:
|
||||
if multitalk_model["is_gguf"] and not gguf:
|
||||
raise ValueError("Multitalk/InfiniteTalk model is a GGUF model, main model also has to be a GGUF model.")
|
||||
multitalk_model_type = multitalk_model.get("model_type", "MultiTalk")
|
||||
log.info(f"{multitalk_model_type} detected, patching model...")
|
||||
|
||||
multitalk_model_path = multitalk_model["model_path"]
|
||||
if multitalk_model_path.endswith(".gguf") and not gguf:
|
||||
raise ValueError("Multitalk/InfiniteTalk model is a GGUF model, main model also has to be a GGUF model.")
|
||||
|
||||
# init audio module
|
||||
from .multitalk.multitalk import SingleStreamMultiAttention
|
||||
from .wanvideo.modules.model import WanLayerNorm
|
||||
norm_input_visual = True #dunno what this is
|
||||
|
||||
for block in transformer.blocks:
|
||||
block.audio_cross_attn = SingleStreamMultiAttention(
|
||||
dim=dim,
|
||||
encoder_hidden_states_dim=768,
|
||||
num_heads=num_heads,
|
||||
with init_empty_weights():
|
||||
block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True)
|
||||
block.audio_cross_attn = SingleStreamMultiAttention(
|
||||
dim=dim,
|
||||
encoder_hidden_states_dim=768,
|
||||
num_heads=num_heads,
|
||||
qkv_bias=True,
|
||||
class_range=24,
|
||||
class_interval=4,
|
||||
attention_mode=attention_mode,
|
||||
)
|
||||
block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True) if norm_input_visual else nn.Identity()
|
||||
log.info(f"{multitalk_model_type} detected, patching model...")
|
||||
transformer.audio_proj = multitalk_model["proj_model"]
|
||||
transformer.multitalk_model_type = multitalk_model_type
|
||||
sd.update(multitalk_model["sd"])
|
||||
|
||||
|
||||
extra_model_path = multitalk_model["model_path"]
|
||||
if gguf:
|
||||
extra_sd, extra_reader = load_gguf(extra_model_path)
|
||||
gguf_reader.append(extra_reader)
|
||||
del extra_reader
|
||||
else:
|
||||
extra_sd = load_torch_file(extra_model_path, device=transformer_load_device, safe_load=True)
|
||||
sd.update(extra_sd)
|
||||
del extra_sd
|
||||
|
||||
# Additional cond latents
|
||||
if "add_conv_in.weight" in sd:
|
||||
def zero_module(module):
|
||||
@@ -1266,178 +1292,54 @@ class WanVideoModelLoader:
|
||||
patcher.model.is_patched = False
|
||||
|
||||
scale_weights = {}
|
||||
if "fp8" in quantization:
|
||||
for k, v in sd.items():
|
||||
if k.endswith(".scale_weight"):
|
||||
scale_weights[k] = v.to(base_dtype)
|
||||
|
||||
if "fp8_e4m3fn" in quantization:
|
||||
weight_dtype = torch.float8_e4m3fn
|
||||
elif "fp8_e5m2" in quantization:
|
||||
weight_dtype = torch.float8_e5m2
|
||||
else:
|
||||
weight_dtype = base_dtype
|
||||
|
||||
params_to_keep = {"norm", "bias", "time_in", "patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter", "add", "ref_conv", "audio_proj"}
|
||||
|
||||
control_lora = False
|
||||
|
||||
if not merge_loras and control_lora:
|
||||
log.warning("Control-LoRA patching is only supported with merge_loras=True")
|
||||
|
||||
if lora is not None:
|
||||
patcher, control_lora, unianimate_sd = add_lora_weights(patcher, lora, base_dtype, merge_loras=merge_loras)
|
||||
if unianimate_sd is not None:
|
||||
log.info("Merging UniAnimate weights to the model...")
|
||||
sd.update(unianimate_sd)
|
||||
del unianimate_sd
|
||||
|
||||
if not gguf:
|
||||
if "fp8" in quantization:
|
||||
for k, v in sd.items():
|
||||
if k.endswith(".scale_weight"):
|
||||
scale_weights[k] = v
|
||||
if not merge_loras:
|
||||
if merge_loras and lora is not None:
|
||||
if not lora_low_mem_load:
|
||||
load_weights(transformer, sd, weight_dtype, base_dtype, transformer_load_device)
|
||||
|
||||
if control_lora:
|
||||
patch_control_lora(patcher.model.diffusion_model, device)
|
||||
patcher.model.is_patched = True
|
||||
|
||||
log.info("Merging LoRA to the model...")
|
||||
patcher = apply_lora(
|
||||
patcher, device, transformer_load_device, params_to_keep=params_to_keep, dtype=weight_dtype, base_dtype=base_dtype, state_dict=sd,
|
||||
low_mem_load=lora_low_mem_load, control_lora=control_lora, scale_weights=scale_weights,)
|
||||
if not control_lora:
|
||||
scale_weights.clear()
|
||||
patcher.patches.clear()
|
||||
transformer.patched_linear = False
|
||||
else:
|
||||
from .custom_linear import _replace_linear
|
||||
transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights)
|
||||
|
||||
if "fp8_e4m3fn" in quantization:
|
||||
dtype = torch.float8_e4m3fn
|
||||
elif "fp8_e5m2" in quantization:
|
||||
dtype = torch.float8_e5m2
|
||||
else:
|
||||
dtype = base_dtype
|
||||
params_to_keep = {"norm", "bias", "time_in", "patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter", "add", "ref_conv", "audio_proj"}
|
||||
if not lora_low_mem_load:
|
||||
log.info("Using accelerate to load and assign model weights to device...")
|
||||
param_count = sum(1 for _ in transformer.named_parameters())
|
||||
pbar = ProgressBar(param_count)
|
||||
cnt = 0
|
||||
for name, param in tqdm(transformer.named_parameters(),
|
||||
desc=f"Loading transformer parameters to {transformer_load_device}",
|
||||
total=param_count,
|
||||
leave=True):
|
||||
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype
|
||||
dtype_to_use = dtype if sd[name].dtype == dtype else dtype_to_use
|
||||
if "modulation" in name or "norm" in name or "bias" in name:
|
||||
dtype_to_use = base_dtype
|
||||
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])
|
||||
cnt += 1
|
||||
if cnt % 100 == 0:
|
||||
pbar.update(100)
|
||||
transformer.patched_linear = True
|
||||
|
||||
#for name, param in transformer.named_parameters():
|
||||
# print(name, param.dtype, param.device, param.shape)
|
||||
pbar.update_absolute(param_count)
|
||||
pbar.update_absolute(0)
|
||||
|
||||
comfy_model.diffusion_model = transformer
|
||||
comfy_model.load_device = transformer_load_device
|
||||
|
||||
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
|
||||
patcher.model.is_patched = False
|
||||
|
||||
control_lora = False
|
||||
if lora is not None:
|
||||
for l in lora:
|
||||
log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}")
|
||||
lora_path = l["path"]
|
||||
lora_strength = l["strength"]
|
||||
if isinstance(lora_strength, list):
|
||||
if merge_loras:
|
||||
raise ValueError("LoRA strength should be a single value when merge_loras=True")
|
||||
transformer.lora_scheduling_enabled = True
|
||||
if lora_strength == 0:
|
||||
log.warning(f"LoRA {lora_path} has strength 0, skipping...")
|
||||
continue
|
||||
lora_sd = load_torch_file(lora_path, safe_load=True)
|
||||
if "dwpose_embedding.0.weight" in lora_sd: #unianimate
|
||||
from .unianimate.nodes import update_transformer
|
||||
log.info("Unianimate LoRA detected, patching model...")
|
||||
transformer = update_transformer(transformer, lora_sd)
|
||||
|
||||
lora_sd = standardize_lora_key_format(lora_sd)
|
||||
|
||||
if l["blocks"]:
|
||||
lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"], l.get("layer_filter", []))
|
||||
|
||||
# Filter out any LoRA keys containing 'img' if the base model state_dict has no 'img' keys
|
||||
if not any('img' in k for k in sd.keys()):
|
||||
lora_sd = {k: v for k, v in lora_sd.items() if 'img' not in k}
|
||||
|
||||
#spacepxl's control LoRA patch
|
||||
# for key in lora_sd.keys():
|
||||
# print(key)
|
||||
|
||||
if "diffusion_model.patch_embedding.lora_A.weight" in lora_sd:
|
||||
log.info("Control-LoRA detected, patching model...")
|
||||
if not merge_loras:
|
||||
log.warning("Control-LoRA patching is only supported with merge_loras=True, setting it to True")
|
||||
merge_loras = True
|
||||
control_lora = True
|
||||
|
||||
in_cls = transformer.patch_embedding.__class__ # nn.Conv3d
|
||||
old_in_dim = transformer.in_dim # 16
|
||||
new_in_dim = lora_sd["diffusion_model.patch_embedding.lora_A.weight"].shape[1]
|
||||
assert new_in_dim == 32
|
||||
|
||||
new_in = in_cls(
|
||||
new_in_dim,
|
||||
transformer.patch_embedding.out_channels,
|
||||
transformer.patch_embedding.kernel_size,
|
||||
transformer.patch_embedding.stride,
|
||||
transformer.patch_embedding.padding,
|
||||
).to(device=device, dtype=torch.float32)
|
||||
|
||||
new_in.weight.zero_()
|
||||
new_in.bias.zero_()
|
||||
|
||||
new_in.weight[:, :old_in_dim].copy_(transformer.patch_embedding.weight)
|
||||
new_in.bias.copy_(transformer.patch_embedding.bias)
|
||||
|
||||
transformer.patch_embedding = new_in
|
||||
transformer.expanded_patch_embedding = new_in
|
||||
|
||||
if "diffusion_model.blocks.0.self_attn.q_loras.down.weight" in lora_sd:
|
||||
log.info("Stand-In LoRA detected")
|
||||
for block in transformer.blocks:
|
||||
block.self_attn.q_loras = LoRALinearLayer(dim, dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength)
|
||||
block.self_attn.k_loras = LoRALinearLayer(dim, dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength)
|
||||
block.self_attn.v_loras = LoRALinearLayer(dim, dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength)
|
||||
for lora in [block.self_attn.q_loras, block.self_attn.k_loras, block.self_attn.v_loras]:
|
||||
for param in lora.parameters():
|
||||
param.requires_grad = False
|
||||
for name, param in transformer.named_parameters():
|
||||
if "lora" in name:
|
||||
param.data.copy_(lora_sd["diffusion_model." + name].to(param.device, dtype=param.dtype))
|
||||
else:
|
||||
patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0)
|
||||
|
||||
del lora_sd
|
||||
|
||||
if not gguf and merge_loras:
|
||||
log.info("Patching LoRA to the model...")
|
||||
patcher = apply_lora(
|
||||
patcher, device, transformer_load_device,
|
||||
params_to_keep=params_to_keep, dtype=dtype, base_dtype=base_dtype, state_dict=sd,
|
||||
low_mem_load=lora_low_mem_load, control_lora=control_lora, scale_weights=scale_weights)
|
||||
scale_weights.clear()
|
||||
patcher.patches.clear()
|
||||
|
||||
if gguf:
|
||||
#from diffusers.quantizers.gguf.utils import _replace_with_gguf_linear, GGUFParameter
|
||||
from .gguf.gguf import _replace_with_gguf_linear, GGUFParameter
|
||||
log.info("Using GGUF to load and assign model weights to device...")
|
||||
param_count = sum(1 for _ in transformer.named_parameters())
|
||||
|
||||
out_features = sd["blocks.0.self_attn.k.weight"].shape[1]
|
||||
|
||||
patcher.model.diffusion_model = _replace_with_gguf_linear(patcher.model.diffusion_model, base_dtype, sd, patches=patcher.patches)
|
||||
pbar = ProgressBar(param_count)
|
||||
cnt = 0
|
||||
for name, param in tqdm(patcher.model.diffusion_model.named_parameters(),
|
||||
desc=f"Loading transformer parameters to {transformer_load_device}",
|
||||
total=param_count,
|
||||
leave=True):
|
||||
if "loras" in name:
|
||||
continue
|
||||
#print(name, param.dtype, param.device, param.shape)
|
||||
if isinstance(param, GGUFParameter):
|
||||
dtype_to_use = torch.uint8
|
||||
elif "patch_embedding" in name:
|
||||
dtype_to_use = torch.float32
|
||||
else:
|
||||
dtype_to_use = base_dtype
|
||||
set_module_tensor_to_device(patcher.model.diffusion_model, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name])
|
||||
cnt += 1
|
||||
if cnt % 100 == 0:
|
||||
pbar.update(100)
|
||||
|
||||
#for name, param in transformer.named_parameters():
|
||||
# print(name, param.dtype, param.device, param.shape)
|
||||
#patcher.load(device, full_load=True)
|
||||
pbar.update_absolute(param_count)
|
||||
|
||||
patcher.model.is_patched = True
|
||||
|
||||
patch_linear = (True if "scaled" in quantization or (lora is not None and not merge_loras) else False)
|
||||
|
||||
if "fast" in quantization:
|
||||
if lora is not None and not merge_loras:
|
||||
raise NotImplementedError("fp8_fast is not supported with unmerged LoRAs")
|
||||
@@ -1577,8 +1479,6 @@ class WanVideoVAELoader:
|
||||
|
||||
def loadmodel(self, model_name, precision, compile_args=None):
|
||||
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
|
||||
#with open(os.path.join(script_directory, 'configs', 'hy_vae_config.json')) as f:
|
||||
# vae_config = json.load(f)
|
||||
model_path = folder_paths.get_full_path("vae", model_name)
|
||||
vae_sd = load_torch_file(model_path, safe_load=True)
|
||||
|
||||
@@ -1592,8 +1492,9 @@ class WanVideoVAELoader:
|
||||
vae = WanVideoVAE38(dtype=dtype)
|
||||
|
||||
vae.load_state_dict(vae_sd)
|
||||
del vae_sd
|
||||
vae.eval()
|
||||
vae.to(device = offload_device, dtype = dtype)
|
||||
vae.to(device=offload_device, dtype=dtype)
|
||||
if compile_args is not None:
|
||||
vae.model.decoder = torch.compile(vae.model.decoder, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||
|
||||
|
||||
Reference in New Issue
Block a user