Refactor model loading
This commit is contained in:
+1
-1
@@ -12,7 +12,7 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s
|
||||
module_prefix = prefix + name + "."
|
||||
_replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights)
|
||||
|
||||
if isinstance(module, nn.Linear):
|
||||
if isinstance(module, nn.Linear) and "loras" not in module_prefix:
|
||||
in_features = state_dict[module_prefix + "weight"].shape[1]
|
||||
out_features = state_dict[module_prefix + "weight"].shape[0]
|
||||
if scale_weights is not None:
|
||||
|
||||
+14
-2
@@ -1,13 +1,25 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import numpy as np
|
||||
from diffusers.quantizers.gguf.utils import GGUFParameter, dequantize_gguf_tensor
|
||||
import gguf
|
||||
from diffusers.utils import is_accelerate_available
|
||||
from contextlib import nullcontext
|
||||
|
||||
from ..utils import log
|
||||
if is_accelerate_available():
|
||||
import accelerate
|
||||
from accelerate import init_empty_weights
|
||||
|
||||
def load_gguf(model_path):
|
||||
from gguf import GGUFReader
|
||||
reader = GGUFReader(model_path)
|
||||
parsed_parameters = {}
|
||||
for tensor in reader.tensors:
|
||||
# if the tensor is a torch supported dtype do not use GGUFParameter
|
||||
is_gguf_quant = tensor.tensor_type not in [gguf.GGMLQuantizationType.F32, gguf.GGMLQuantizationType.F16]
|
||||
meta_tensor = torch.empty(tensor.data.shape, dtype=torch.from_numpy(np.empty(0, dtype=tensor.data.dtype)).dtype, device='meta')
|
||||
parsed_parameters[tensor.name] = GGUFParameter(meta_tensor, quant_type=tensor.tensor_type) if is_gguf_quant else meta_tensor
|
||||
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 _should_convert_to_gguf(state_dict, prefix):
|
||||
|
||||
@@ -42,6 +42,17 @@ def offload_transformer(transformer):
|
||||
transformer.magcache_state.clear_all()
|
||||
transformer.easycache_state.clear_all()
|
||||
transformer.to(offload_device)
|
||||
# for name, param in transformer.named_parameters():
|
||||
# module = transformer
|
||||
# subnames = name.split('.')
|
||||
# for subname in subnames[:-1]:
|
||||
# module = getattr(module, subname)
|
||||
# attr_name = subnames[-1]
|
||||
# if param.data.is_floating_point():
|
||||
# meta_param = torch.nn.Parameter(torch.empty_like(param.data, device='meta'), requires_grad=False)
|
||||
# setattr(module, attr_name, meta_param)
|
||||
# else:
|
||||
# pass
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
|
||||
@@ -348,8 +359,11 @@ class WanVideoTextEncode:
|
||||
raise ValueError("No cached text embeds found for prompts, please provide a T5 encoder.")
|
||||
|
||||
if model_to_offload is not None and device == "gpu":
|
||||
log.info(f"Moving video model to {offload_device}")
|
||||
model_to_offload.model.to(offload_device)
|
||||
try:
|
||||
log.info(f"Moving video model to {offload_device}")
|
||||
model_to_offload.model.to(offload_device)
|
||||
except:
|
||||
pass
|
||||
|
||||
encoder = t5["model"]
|
||||
dtype = t5["dtype"]
|
||||
@@ -1080,7 +1094,7 @@ class WanVideoPhantomEmbeds:
|
||||
|
||||
log.info(f"Phantom latents shape: {samples.shape}")
|
||||
|
||||
target_shape = (16, (num_frames - 1) // VAE_STRIDE[0] + 1 + T,
|
||||
target_shape = (16, (num_frames - 1) // VAE_STRIDE[0] + 1,
|
||||
H * 8 // VAE_STRIDE[1],
|
||||
W * 8 // VAE_STRIDE[2])
|
||||
|
||||
@@ -1583,20 +1597,37 @@ class WanVideoSampler:
|
||||
model = model.model
|
||||
transformer = model.diffusion_model
|
||||
|
||||
dtype = model["dtype"]
|
||||
dtype = model["base_dtype"]
|
||||
weight_dtype = model["weight_dtype"]
|
||||
fp8_matmul = model["fp8_matmul"]
|
||||
gguf = model["gguf"]
|
||||
gguf_reader = model["gguf_reader"]
|
||||
control_lora = model["control_lora"]
|
||||
|
||||
transformer_options = patcher.model_options.get("transformer_options", None)
|
||||
merge_loras = transformer_options["merge_loras"]
|
||||
|
||||
block_swap_args = transformer_options.get("block_swap_args", None)
|
||||
if block_swap_args is not None:
|
||||
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False)
|
||||
transformer.blocks_to_swap = block_swap_args.get("blocks_to_swap", 0)
|
||||
transformer.vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", 0)
|
||||
transformer.prefetch_blocks = block_swap_args.get("prefetch_blocks", 0)
|
||||
transformer.block_swap_debug = block_swap_args.get("block_swap_debug", False)
|
||||
transformer.offload_img_emb = block_swap_args.get("offload_img_emb", False)
|
||||
transformer.offload_txt_emb = block_swap_args.get("offload_txt_emb", False)
|
||||
|
||||
is_5b = transformer.out_dim == 48
|
||||
vae_upscale_factor = 16 if is_5b else 8
|
||||
|
||||
patch_linear = transformer_options.get("patch_linear", False)
|
||||
from .nodes_model_loading import load_weights, load_weights_gguf
|
||||
weights_assigned = False
|
||||
if not merge_loras and gguf_reader is None:
|
||||
load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, dtype, device, block_swap_args=block_swap_args)
|
||||
weights_assigned = True
|
||||
|
||||
if gguf:
|
||||
if gguf_reader is not None:
|
||||
load_weights_gguf(transformer, gguf_reader, patcher.model["sd"], dtype, device, patcher)
|
||||
set_lora_params_gguf(transformer, patcher.patches)
|
||||
elif len(patcher.patches) != 0 and patch_linear:
|
||||
log.info(f"Using {len(patcher.patches)} LoRA weight patches for WanVideo model")
|
||||
@@ -1867,8 +1898,6 @@ class WanVideoSampler:
|
||||
phantom_cfg_scale = [phantom_cfg_scale] * (steps +1)
|
||||
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]
|
||||
|
||||
@@ -2153,11 +2182,8 @@ class WanVideoSampler:
|
||||
mm.unload_all_models()
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
|
||||
if transformer_options is not None:
|
||||
block_swap_args = transformer_options.get("block_swap_args", None)
|
||||
|
||||
if block_swap_args is not None:
|
||||
if block_swap_args is not None and not weights_assigned:
|
||||
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False)
|
||||
for name, param in transformer.named_parameters():
|
||||
if "block" not in name:
|
||||
@@ -2187,8 +2213,8 @@ class WanVideoSampler:
|
||||
block.modulation = torch.nn.Parameter(block.modulation.to(device))
|
||||
transformer.head.modulation = torch.nn.Parameter(transformer.head.modulation.to(device))
|
||||
|
||||
elif model["manual_offloading"]:
|
||||
transformer.to(device)
|
||||
#elif model["manual_offloading"]:
|
||||
# transformer.to(device)
|
||||
|
||||
# Initialize Cache if enabled
|
||||
transformer.enable_teacache = transformer.enable_magcache = transformer.enable_easycache = False
|
||||
@@ -2338,18 +2364,14 @@ class WanVideoSampler:
|
||||
z = torch.cat([z, recam_latents.to(z)], dim=1)
|
||||
|
||||
use_phantom = False
|
||||
phantom_ref = None
|
||||
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)
|
||||
phantom_ref = phantom_latents.to(z)
|
||||
use_phantom = True
|
||||
if cache_state is not None and len(cache_state) != 3:
|
||||
cache_state.append(None)
|
||||
if not use_phantom:
|
||||
z_pos = z_neg = z
|
||||
|
||||
if controlnet_latents is not None:
|
||||
if (controlnet_start <= current_step_percentage < controlnet_end):
|
||||
@@ -2375,9 +2397,9 @@ class WanVideoSampler:
|
||||
|
||||
if minimax_latents is not None:
|
||||
if context_window is not None:
|
||||
z_pos = z_neg = torch.cat([z, minimax_latents[:, context_window], minimax_mask_latents[:, context_window]], dim=0)
|
||||
z = torch.cat([z, minimax_latents[:, context_window], minimax_mask_latents[:, context_window]], dim=0)
|
||||
else:
|
||||
z_pos = z_neg = torch.cat([z, minimax_latents, minimax_mask_latents], dim=0)
|
||||
z = torch.cat([z, minimax_latents, minimax_mask_latents], dim=0)
|
||||
|
||||
if not multitalk_sampling and multitalk_audio_embedding is not None:
|
||||
audio_embedding = multitalk_audio_embedding
|
||||
@@ -2442,6 +2464,7 @@ class WanVideoSampler:
|
||||
"inner_t": [shot_len] if shot_len else None,
|
||||
"standin_input": standin_input,
|
||||
"fantasy_portrait_input": fantasy_portrait_input,
|
||||
"phantom_ref": phantom_ref
|
||||
}
|
||||
|
||||
batch_size = 1
|
||||
@@ -2456,7 +2479,7 @@ class WanVideoSampler:
|
||||
if not batched_cfg:
|
||||
#cond
|
||||
noise_pred_cond, cache_state_cond = transformer(
|
||||
[z_pos], context=positive_embeds, y=[image_cond_input] if image_cond_input is not None else None,
|
||||
[z], 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=cache_state[0] if cache_state else None,
|
||||
vace_data=vace_data, attn_cond=attn_cond,
|
||||
@@ -2477,7 +2500,7 @@ class WanVideoSampler:
|
||||
if not math.isclose(audio_cfg_scale[idx], 1.0):
|
||||
base_params['audio_proj'] = None
|
||||
noise_pred_uncond, cache_state_uncond = transformer(
|
||||
[z_neg], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea,
|
||||
[z], 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=cache_state[1] if cache_state else None,
|
||||
@@ -2488,7 +2511,7 @@ class WanVideoSampler:
|
||||
#phantom
|
||||
if use_phantom and not math.isclose(phantom_cfg_scale[idx], 1.0):
|
||||
noise_pred_phantom, cache_state_phantom = transformer(
|
||||
[z_phantom_img], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea,
|
||||
[z], 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=cache_state[2] if cache_state else None,
|
||||
@@ -2506,7 +2529,7 @@ class WanVideoSampler:
|
||||
cache_state.append(None)
|
||||
base_params['audio_proj'] = None
|
||||
noise_pred_no_audio, cache_state_audio = transformer(
|
||||
[z_pos], context=positive_embeds, y=[image_cond_input] if image_cond_input is not None else None,
|
||||
[z], 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=cache_state[2] if cache_state else None,
|
||||
vace_data=vace_data,
|
||||
@@ -2525,7 +2548,7 @@ class WanVideoSampler:
|
||||
cache_state.append(None)
|
||||
base_params['multitalk_audio'] = torch.zeros_like(multitalk_audio_input)[-1:]
|
||||
noise_pred_no_audio, cache_state_audio = transformer(
|
||||
[z_pos], context=negative_embeds, y=[image_cond_input] if image_cond_input is not None else None,
|
||||
[z], context=negative_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=cache_state[2] if cache_state else None,
|
||||
vace_data=vace_data,
|
||||
@@ -3283,8 +3306,8 @@ class WanVideoSampler:
|
||||
if callback is not None:
|
||||
if recammaster is not None:
|
||||
callback_latent = (latent_model_input[:, :orig_noise_len].to(device) - noise_pred[:, :orig_noise_len].to(device) * t.to(device) / 1000).detach()
|
||||
elif phantom_latents is not None:
|
||||
callback_latent = (latent_model_input[:,:-phantom_latents.shape[1]].to(device) - noise_pred[:,:-phantom_latents.shape[1]].to(device) * t.to(device) / 1000).detach()
|
||||
#elif phantom_latents is not None:
|
||||
# callback_latent = (latent_model_input[:,:-phantom_latents.shape[1]].to(device) - noise_pred[:,:-phantom_latents.shape[1]].to(device) * t.to(device) / 1000).detach()
|
||||
else:
|
||||
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device) / 1000).detach()
|
||||
callback(idx, callback_latent.permute(1,0,2,3), None, len(timesteps))
|
||||
@@ -3303,8 +3326,8 @@ class WanVideoSampler:
|
||||
offload_transformer(transformer)
|
||||
raise e
|
||||
|
||||
if phantom_latents is not None:
|
||||
latent = latent[:,:-phantom_latents.shape[1]]
|
||||
#if phantom_latents is not None:
|
||||
# latent = latent[:,:-phantom_latents.shape[1]]
|
||||
|
||||
if cache_args is not None:
|
||||
cache_report(transformer, cache_args)
|
||||
|
||||
+231
-175
@@ -708,6 +708,178 @@ class WanVideoSetLoRAs:
|
||||
|
||||
return (patcher,)
|
||||
|
||||
def load_weights(transformer, sd, weight_dtype, base_dtype, transformer_load_device, block_swap_args=None):
|
||||
print("block_swap_args", block_swap_args)
|
||||
params_to_keep = {"norm", "bias", "time_in", "patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter", "add", "ref_conv"}
|
||||
|
||||
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):
|
||||
block_idx = vace_block_idx = None
|
||||
if "vace_blocks." in name:
|
||||
try:
|
||||
vace_block_idx = int(name.split("vace_blocks.")[1].split(".")[0])
|
||||
except Exception:
|
||||
vace_block_idx = None
|
||||
elif "blocks." in name:
|
||||
try:
|
||||
block_idx = int(name.split("blocks.")[1].split(".")[0])
|
||||
except Exception:
|
||||
block_idx = None
|
||||
print("block_idx:", block_idx)
|
||||
print("vace_block_idx:", vace_block_idx)
|
||||
if "loras" in name:
|
||||
continue
|
||||
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else weight_dtype
|
||||
dtype_to_use = weight_dtype if sd[name].dtype == weight_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
|
||||
|
||||
load_device = device
|
||||
if block_swap_args is not None:
|
||||
if block_idx is not None:
|
||||
if block_idx >= len(transformer.blocks) - block_swap_args["blocks_to_swap"]:
|
||||
load_device = offload_device
|
||||
elif vace_block_idx is not None:
|
||||
if vace_block_idx >= len(transformer.vace_blocks) - block_swap_args["vace_blocks_to_swap"]:
|
||||
load_device = offload_device
|
||||
|
||||
print("load_device:", load_device)
|
||||
|
||||
set_module_tensor_to_device(transformer, name, device=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)
|
||||
pbar.update_absolute(param_count)
|
||||
pbar.update_absolute(0)
|
||||
|
||||
def load_weights_gguf(transformer, reader, sd, base_dtype, transformer_load_device, patcher):
|
||||
from .gguf.gguf import _replace_with_gguf_linear, GGUFParameter
|
||||
import gguf
|
||||
log.info("Using GGUF to load and assign model weights to device...")
|
||||
param_count = sum(1 for _ in transformer.named_parameters())
|
||||
|
||||
sd = {}
|
||||
for tensor in reader.tensors:
|
||||
# if the tensor is a torch supported dtype do not use GGUFParameter
|
||||
is_gguf_quant = tensor.tensor_type not in [gguf.GGMLQuantizationType.F32, gguf.GGMLQuantizationType.F16]
|
||||
weights = torch.from_numpy(tensor.data.copy()).to(transformer_load_device)
|
||||
sd[tensor.name] = GGUFParameter(weights, quant_type=tensor.tensor_type) if is_gguf_quant else weights
|
||||
|
||||
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)
|
||||
pbar.update_absolute(0)
|
||||
|
||||
def patch_control_lora(transformer, device):
|
||||
log.info("Control-LoRA detected, patching model...")
|
||||
|
||||
in_cls = transformer.patch_embedding.__class__ # nn.Conv3d
|
||||
old_in_dim = transformer.in_dim # 16
|
||||
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
|
||||
|
||||
def patch_stand_in_lora(transformer, lora_sd, transformer_load_device, base_dtype, lora_strength):
|
||||
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(transformer.dim, transformer.dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength)
|
||||
block.self_attn.k_loras = LoRALinearLayer(transformer.dim, transformer.dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength)
|
||||
block.self_attn.v_loras = LoRALinearLayer(transformer.dim, transformer.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))
|
||||
|
||||
def add_lora_weights(patcher, lora, base_dtype, merge_loras=False):
|
||||
#spacepxl's control LoRA patch
|
||||
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}
|
||||
control_lora=False
|
||||
if "diffusion_model.patch_embedding.lora_A.weight" in lora_sd:
|
||||
control_lora = True
|
||||
#stand-in LoRA patch
|
||||
if "diffusion_model.blocks.0.self_attn.q_loras.down.weight" in lora_sd:
|
||||
patch_stand_in_lora(patcher.model.diffusion_model, lora_sd, device, base_dtype, lora_strength)
|
||||
# normal LoRA patch
|
||||
else:
|
||||
patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0)
|
||||
|
||||
del lora_sd
|
||||
return patcher, control_lora
|
||||
|
||||
#region Model loading
|
||||
class WanVideoModelLoader:
|
||||
@classmethod
|
||||
@@ -797,13 +969,16 @@ class WanVideoModelLoader:
|
||||
|
||||
|
||||
model_path = folder_paths.get_full_path_or_raise("diffusion_models", model)
|
||||
|
||||
|
||||
gguf_reader = None
|
||||
if not gguf:
|
||||
sd = load_torch_file(model_path, device=transformer_load_device, safe_load=True)
|
||||
else:
|
||||
from diffusers.models.model_loading_utils import load_gguf_checkpoint
|
||||
sd = load_gguf_checkpoint(model_path)
|
||||
|
||||
#from diffusers.models.model_loading_utils import load_gguf_checkpoint
|
||||
#sd = load_gguf_checkpoint(model_path)
|
||||
from .gguf.gguf import load_gguf
|
||||
sd, gguf_reader = load_gguf(model_path)
|
||||
|
||||
if quantization == "disabled":
|
||||
for k, v in sd.items():
|
||||
if isinstance(v, torch.Tensor):
|
||||
@@ -1055,178 +1230,56 @@ class WanVideoModelLoader:
|
||||
model_type=comfy.model_base.ModelType.FLOW,
|
||||
device=device,
|
||||
)
|
||||
scale_weights = {}
|
||||
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:
|
||||
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"}
|
||||
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)
|
||||
|
||||
#for name, param in transformer.named_parameters():
|
||||
# print(name, param.dtype, param.device, param.shape)
|
||||
pbar.update_absolute(param_count)
|
||||
|
||||
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]
|
||||
scale_weights = {}
|
||||
if "fp8" in quantization:
|
||||
for k, v in sd.items():
|
||||
if k.endswith(".scale_weight"):
|
||||
scale_weights[k] = v
|
||||
|
||||
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
|
||||
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"}
|
||||
|
||||
control_lora = False
|
||||
patch_linear = (True if "scaled" in quantization or (lora is not None and not merge_loras) else 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 = add_lora_weights(patcher, lora, base_dtype, merge_loras=merge_loras)
|
||||
|
||||
if not gguf:
|
||||
if merge_loras and not patch_linear:
|
||||
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()
|
||||
else:
|
||||
from .custom_linear import _replace_linear
|
||||
transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights)
|
||||
|
||||
if "fast" in quantization:
|
||||
if lora is not None and not merge_loras:
|
||||
raise NotImplementedError("fp8_fast is not supported with unmerged LoRAs")
|
||||
@@ -1234,7 +1287,7 @@ class WanVideoModelLoader:
|
||||
convert_fp8_linear(transformer, base_dtype, params_to_keep, scale_weight_keys=scale_weights)
|
||||
patch_linear = False
|
||||
|
||||
del sd
|
||||
#del sd
|
||||
|
||||
if multitalk_model is not None:
|
||||
transformer.audio_proj = multitalk_model["proj_model"]
|
||||
@@ -1263,18 +1316,18 @@ class WanVideoModelLoader:
|
||||
WanRMSNorm: AutoWrappedModule,
|
||||
},
|
||||
module_config = dict(
|
||||
offload_dtype=dtype,
|
||||
offload_dtype=weight_dtype,
|
||||
offload_device=offload_device,
|
||||
onload_dtype=dtype,
|
||||
onload_dtype=weight_dtype,
|
||||
onload_device=device,
|
||||
computation_dtype=base_dtype,
|
||||
computation_device=device,
|
||||
),
|
||||
max_num_param=params_to_keep,
|
||||
overflow_module_config = dict(
|
||||
offload_dtype=dtype,
|
||||
offload_dtype=weight_dtype,
|
||||
offload_device=offload_device,
|
||||
onload_dtype=dtype,
|
||||
onload_dtype=weight_dtype,
|
||||
onload_device=offload_device,
|
||||
computation_dtype=base_dtype,
|
||||
computation_device=device,
|
||||
@@ -1282,13 +1335,14 @@ class WanVideoModelLoader:
|
||||
compile_args = compile_args,
|
||||
)
|
||||
|
||||
if load_device == "offload_device" and patcher.model.diffusion_model.device != offload_device:
|
||||
if merge_loras:
|
||||
log.info(f"Moving diffusion model from {patcher.model.diffusion_model.device} to {offload_device}")
|
||||
patcher.model.diffusion_model.to(offload_device)
|
||||
gc.collect()
|
||||
mm.soft_empty_cache()
|
||||
|
||||
patcher.model["dtype"] = base_dtype
|
||||
patcher.model["base_dtype"] = base_dtype
|
||||
patcher.model["weight_dtype"] = weight_dtype
|
||||
patcher.model["base_path"] = model_path
|
||||
patcher.model["model_name"] = model
|
||||
patcher.model["manual_offloading"] = manual_offloading
|
||||
@@ -1296,9 +1350,11 @@ class WanVideoModelLoader:
|
||||
patcher.model["auto_cpu_offload"] = True if vram_management_args is not None else False
|
||||
patcher.model["control_lora"] = control_lora
|
||||
patcher.model["compile_args"] = compile_args
|
||||
patcher.model["gguf"] = gguf
|
||||
patcher.model["gguf_reader"] = gguf_reader
|
||||
patcher.model["fp8_matmul"] = "fast" in quantization
|
||||
patcher.model["scale_weights"] = scale_weights
|
||||
patcher.model["sd"] = sd
|
||||
patcher.model["lora"] = lora
|
||||
|
||||
if 'transformer_options' not in patcher.model_options:
|
||||
patcher.model_options['transformer_options'] = {}
|
||||
|
||||
@@ -173,7 +173,7 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d
|
||||
to_load.append((n, m, params))
|
||||
|
||||
to_load.sort(reverse=True)
|
||||
pbar = ProgressBar(len(to_load))
|
||||
#pbar = ProgressBar(len(to_load))
|
||||
for x in tqdm(to_load, desc="Loading model and applying LoRA weights:", leave=True):
|
||||
name = x[0]
|
||||
m = x[1]
|
||||
@@ -207,7 +207,7 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d
|
||||
except:
|
||||
continue
|
||||
m.comfy_patched_weights = True
|
||||
pbar.update(1)
|
||||
#pbar.update(1)
|
||||
|
||||
model.current_weight_patches_uuid = model.patches_uuid
|
||||
if low_mem_load:
|
||||
|
||||
+40
-15
@@ -1377,7 +1377,10 @@ class WanModel(torch.nn.Module):
|
||||
return block_mask
|
||||
|
||||
def block_swap(self, blocks_to_swap, offload_txt_emb=False, offload_img_emb=False, vace_blocks_to_swap=None, prefetch_blocks=0, block_swap_debug=False):
|
||||
log.info(f"Swapping {blocks_to_swap + 1} transformer blocks")
|
||||
# Clamp blocks_to_swap to valid range
|
||||
blocks_to_swap = max(0, min(blocks_to_swap, len(self.blocks)))
|
||||
|
||||
log.info(f"Swapping {blocks_to_swap} transformer blocks")
|
||||
self.blocks_to_swap = blocks_to_swap
|
||||
self.prefetch_blocks = prefetch_blocks
|
||||
self.block_swap_debug = block_swap_debug
|
||||
@@ -1387,11 +1390,14 @@ class WanModel(torch.nn.Module):
|
||||
|
||||
total_offload_memory = 0
|
||||
total_main_memory = 0
|
||||
|
||||
# Calculate the index where swapping starts
|
||||
swap_start_idx = len(self.blocks) - blocks_to_swap
|
||||
|
||||
for b, block in tqdm(enumerate(self.blocks), total=len(self.blocks), desc="Initializing block swap"):
|
||||
block_memory = get_module_memory_mb(block)
|
||||
|
||||
if b > self.blocks_to_swap:
|
||||
if b < swap_start_idx:
|
||||
block.to(self.main_device)
|
||||
total_main_memory += block_memory
|
||||
else:
|
||||
@@ -1402,12 +1408,17 @@ class WanModel(torch.nn.Module):
|
||||
vace_blocks_to_swap = 1
|
||||
|
||||
if vace_blocks_to_swap > 0 and self.vace_layers is not None:
|
||||
# Clamp vace_blocks_to_swap to valid range
|
||||
vace_blocks_to_swap = max(0, min(vace_blocks_to_swap, len(self.vace_blocks)))
|
||||
self.vace_blocks_to_swap = vace_blocks_to_swap
|
||||
|
||||
# Calculate the index where VACE swapping starts
|
||||
vace_swap_start_idx = len(self.vace_blocks) - vace_blocks_to_swap
|
||||
|
||||
for b, block in tqdm(enumerate(self.vace_blocks), total=len(self.vace_blocks), desc="Initializing vace block swap"):
|
||||
block_memory = get_module_memory_mb(block)
|
||||
|
||||
if b > self.vace_blocks_to_swap:
|
||||
if b < vace_swap_start_idx:
|
||||
block.to(self.main_device)
|
||||
total_main_memory += block_memory
|
||||
else:
|
||||
@@ -1447,9 +1458,10 @@ class WanModel(torch.nn.Module):
|
||||
|
||||
hints = []
|
||||
current_c = c
|
||||
vace_swap_start_idx = len(self.vace_blocks) - self.vace_blocks_to_swap if self.vace_blocks_to_swap > 0 else len(self.vace_blocks)
|
||||
|
||||
for b, block in enumerate(self.vace_blocks):
|
||||
if b <= self.vace_blocks_to_swap and self.vace_blocks_to_swap >= 0:
|
||||
if b >= vace_swap_start_idx and self.vace_blocks_to_swap > 0:
|
||||
block.to(self.main_device)
|
||||
|
||||
if b == 0:
|
||||
@@ -1462,13 +1474,13 @@ class WanModel(torch.nn.Module):
|
||||
# Store skip connection
|
||||
c_skip = block.after_proj(c_processed)
|
||||
hints.append(c_skip.to(
|
||||
self.offload_device if self.vace_blocks_to_swap != -1 else self.main_device,
|
||||
self.offload_device if self.vace_blocks_to_swap > 0 else self.main_device,
|
||||
non_blocking=self.use_non_blocking
|
||||
))
|
||||
|
||||
current_c = c_processed
|
||||
|
||||
if b <= self.vace_blocks_to_swap and self.vace_blocks_to_swap >= 0:
|
||||
if b >= vace_swap_start_idx and self.vace_blocks_to_swap > 0:
|
||||
block.to(self.offload_device, non_blocking=self.use_non_blocking)
|
||||
|
||||
return hints
|
||||
@@ -1509,7 +1521,8 @@ class WanModel(torch.nn.Module):
|
||||
ref_target_masks=None,
|
||||
inner_t=None,
|
||||
standin_input=None,
|
||||
fantasy_portrait_input=None
|
||||
fantasy_portrait_input=None,
|
||||
phantom_ref=None
|
||||
):
|
||||
r"""
|
||||
Forward pass through the diffusion model
|
||||
@@ -1622,6 +1635,15 @@ class WanModel(torch.nn.Module):
|
||||
F += 1
|
||||
x = [torch.concat([_fun_ref.unsqueeze(0), u], dim=1) for _fun_ref, u in zip(fun_ref, x)]
|
||||
|
||||
if phantom_ref is not None:
|
||||
phantom_ref_frames = phantom_ref.size(1)
|
||||
phantom_ref = self.original_patch_embedding(phantom_ref.unsqueeze(0).to(torch.float32)).flatten(2).transpose(1, 2).to(x[0].dtype)
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] + phantom_ref_frames, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
phantom_ref_seq_len = phantom_ref.size(1)
|
||||
seq_len += phantom_ref_seq_len
|
||||
F += phantom_ref_frames
|
||||
x = [torch.concat([u, phantom_ref.unsqueeze(0)], dim=1) for phantom_ref, u in zip(phantom_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([
|
||||
@@ -2005,20 +2027,21 @@ class WanModel(torch.nn.Module):
|
||||
# Asynchronous block offloading with CUDA streams and events
|
||||
cuda_stream = mm.get_offload_stream(device)
|
||||
events = [torch.cuda.Event() for _ in self.blocks]
|
||||
swap_start_idx = len(self.blocks) - self.blocks_to_swap if self.blocks_to_swap > 0 else len(self.blocks)
|
||||
|
||||
for b, block in enumerate(self.blocks):
|
||||
# Prefetch blocks if enabled
|
||||
if self.prefetch_blocks > 0:
|
||||
for prefetch_offset in range(1, self.prefetch_blocks + 1):
|
||||
prefetch_idx = b + prefetch_offset
|
||||
if prefetch_idx < len(self.blocks) and self.blocks_to_swap >= 0 and prefetch_idx <= self.blocks_to_swap:
|
||||
if prefetch_idx < len(self.blocks) and self.blocks_to_swap > 0 and prefetch_idx >= swap_start_idx:
|
||||
with torch.cuda.stream(cuda_stream):
|
||||
self.blocks[prefetch_idx].to(self.main_device, non_blocking=self.use_non_blocking)
|
||||
events[prefetch_idx].record(cuda_stream)
|
||||
if self.block_swap_debug:
|
||||
transfer_start = time.perf_counter()
|
||||
# Wait for block to be ready
|
||||
if b <= self.blocks_to_swap and self.blocks_to_swap >= 0:
|
||||
if b >= swap_start_idx and self.blocks_to_swap > 0:
|
||||
if self.prefetch_blocks > 0:
|
||||
if not events[b].query():
|
||||
events[b].synchronize()
|
||||
@@ -2037,7 +2060,7 @@ class WanModel(torch.nn.Module):
|
||||
compute_end = time.perf_counter()
|
||||
compute_time = compute_end - compute_start
|
||||
to_cpu_transfer_start = time.perf_counter()
|
||||
if b <= self.blocks_to_swap and self.blocks_to_swap >= 0:
|
||||
if b >= swap_start_idx and self.blocks_to_swap > 0:
|
||||
block.to(self.offload_device, non_blocking=self.use_non_blocking)
|
||||
if self.block_swap_debug:
|
||||
to_cpu_transfer_end = time.perf_counter()
|
||||
@@ -2051,9 +2074,6 @@ class WanModel(torch.nn.Module):
|
||||
if (controlnet is not None) and (b % controlnet["controlnet_stride"] == 0) and (b // controlnet["controlnet_stride"] < len(controlnet["controlnet_states"])):
|
||||
x[:, :x_len] += controlnet["controlnet_states"][b // controlnet["controlnet_stride"]].to(x) * controlnet["controlnet_weight"]
|
||||
|
||||
if b <= self.blocks_to_swap and self.blocks_to_swap >= 0:
|
||||
block.to(self.offload_device, non_blocking=self.use_non_blocking)
|
||||
|
||||
if self.enable_teacache and (self.teacache_start_step <= current_step <= self.teacache_end_step) and pred_id is not None:
|
||||
self.teacache_state.update(
|
||||
pred_id,
|
||||
@@ -2076,9 +2096,14 @@ class WanModel(torch.nn.Module):
|
||||
)
|
||||
|
||||
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:]
|
||||
fun_ref_length = fun_ref.size(1)
|
||||
x = x[:, fun_ref_length:]
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
|
||||
if phantom_ref is not None:
|
||||
phantom_ref_length = phantom_ref.size(1)
|
||||
x = x[:, :-phantom_ref_length]
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] - phantom_ref_frames, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
|
||||
if attn_cond is not None:
|
||||
x = x[:, :x_len]
|
||||
|
||||
@@ -1037,7 +1037,7 @@ class VideoVAE_(nn.Module):
|
||||
feat_cache=self._enc_feat_map,
|
||||
feat_idx=self._enc_conv_idx)
|
||||
out = torch.cat([out, out_], 2)
|
||||
pbar.update(1)
|
||||
pbar.update(iter_)
|
||||
mu = self.conv1(out).chunk(2, dim=1)[0]
|
||||
|
||||
mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu)
|
||||
|
||||
Reference in New Issue
Block a user