remove torchao quants, fix long unianimate crash
This commit is contained in:
@@ -664,142 +664,150 @@ class WanVideoModelLoader:
|
||||
model_type=comfy.model_base.ModelType.FLOW,
|
||||
device=device,
|
||||
)
|
||||
|
||||
|
||||
if quantization == "disabled":
|
||||
for k, v in sd.items():
|
||||
if isinstance(v, torch.Tensor):
|
||||
if v.dtype == torch.float8_e4m3fn:
|
||||
quantization = "fp8_e4m3fn"
|
||||
break
|
||||
elif v.dtype == torch.float8_e5m2:
|
||||
quantization = "fp8_e5m2"
|
||||
break
|
||||
|
||||
if not "torchao" in quantization:
|
||||
if "fp8_e4m3fn" in quantization:
|
||||
dtype = torch.float8_e4m3fn
|
||||
elif quantization == "fp8_e5m2":
|
||||
dtype = torch.float8_e5m2
|
||||
else:
|
||||
dtype = base_dtype
|
||||
params_to_keep = {"norm", "head", "bias", "time_in", "vector_in", "patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter"}
|
||||
#if lora is not None:
|
||||
# transformer_load_device = device
|
||||
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())
|
||||
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
|
||||
if "patch_embedding" in name:
|
||||
dtype_to_use = torch.float32
|
||||
set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name])
|
||||
comfy_model.diffusion_model = transformer
|
||||
comfy_model.load_device = transformer_load_device
|
||||
|
||||
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
|
||||
patcher.model.is_patched = False
|
||||
if "fp8_e4m3fn" in quantization:
|
||||
dtype = torch.float8_e4m3fn
|
||||
elif quantization == "fp8_e5m2":
|
||||
dtype = torch.float8_e5m2
|
||||
else:
|
||||
dtype = base_dtype
|
||||
params_to_keep = {"norm", "head", "bias", "time_in", "vector_in", "patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter"}
|
||||
#if lora is not None:
|
||||
# transformer_load_device = device
|
||||
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())
|
||||
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
|
||||
if "patch_embedding" in name:
|
||||
dtype_to_use = torch.float32
|
||||
set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name])
|
||||
comfy_model.diffusion_model = transformer
|
||||
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"]
|
||||
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)
|
||||
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"]
|
||||
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"])
|
||||
lora_sd = standardize_lora_key_format(lora_sd)
|
||||
if l["blocks"]:
|
||||
lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"])
|
||||
|
||||
#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...")
|
||||
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
|
||||
transformer.register_to_config(in_dim=new_in_dim)
|
||||
|
||||
patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0)
|
||||
|
||||
del lora_sd
|
||||
#spacepxl's control LoRA patch
|
||||
# for key in lora_sd.keys():
|
||||
# print(key)
|
||||
|
||||
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)
|
||||
#patcher.load(device, full_load=True)
|
||||
patcher.model.is_patched = True
|
||||
if "diffusion_model.patch_embedding.lora_A.weight" in lora_sd:
|
||||
log.info("Control-LoRA detected, patching model...")
|
||||
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
|
||||
transformer.register_to_config(in_dim=new_in_dim)
|
||||
|
||||
patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0)
|
||||
|
||||
del lora_sd
|
||||
|
||||
|
||||
if "fast" in quantization:
|
||||
from .fp8_optimization import convert_fp8_linear
|
||||
if quantization == "fp8_e4m3fn_fast_no_ffn":
|
||||
params_to_keep.update({"ffn"})
|
||||
print(params_to_keep)
|
||||
convert_fp8_linear(patcher.model.diffusion_model, base_dtype, params_to_keep=params_to_keep)
|
||||
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)
|
||||
#patcher.load(device, full_load=True)
|
||||
patcher.model.is_patched = True
|
||||
|
||||
del sd
|
||||
|
||||
|
||||
if "fast" in quantization:
|
||||
from .fp8_optimization import convert_fp8_linear
|
||||
if quantization == "fp8_e4m3fn_fast_no_ffn":
|
||||
params_to_keep.update({"ffn"})
|
||||
print(params_to_keep)
|
||||
convert_fp8_linear(patcher.model.diffusion_model, base_dtype, params_to_keep=params_to_keep)
|
||||
|
||||
if vram_management_args is not None:
|
||||
from .diffsynth.vram_management import enable_vram_management, AutoWrappedModule, AutoWrappedLinear
|
||||
from .wanvideo.modules.model import WanLayerNorm, WanRMSNorm
|
||||
del sd
|
||||
|
||||
total_params_in_model = sum(p.numel() for p in patcher.model.diffusion_model.parameters())
|
||||
log.info(f"Total number of parameters in the loaded model: {total_params_in_model}")
|
||||
if vram_management_args is not None:
|
||||
from .diffsynth.vram_management import enable_vram_management, AutoWrappedModule, AutoWrappedLinear
|
||||
from .wanvideo.modules.model import WanLayerNorm, WanRMSNorm
|
||||
|
||||
offload_percent = vram_management_args["offload_percent"]
|
||||
offload_params = int(total_params_in_model * offload_percent)
|
||||
params_to_keep = total_params_in_model - offload_params
|
||||
log.info(f"Selected params to offload: {offload_params}")
|
||||
|
||||
enable_vram_management(
|
||||
patcher.model.diffusion_model,
|
||||
module_map = {
|
||||
torch.nn.Linear: AutoWrappedLinear,
|
||||
torch.nn.Conv3d: AutoWrappedModule,
|
||||
torch.nn.LayerNorm: AutoWrappedModule,
|
||||
WanLayerNorm: AutoWrappedModule,
|
||||
WanRMSNorm: AutoWrappedModule,
|
||||
},
|
||||
module_config = dict(
|
||||
offload_dtype=dtype,
|
||||
offload_device=offload_device,
|
||||
onload_dtype=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_device=offload_device,
|
||||
onload_dtype=dtype,
|
||||
onload_device=offload_device,
|
||||
computation_dtype=base_dtype,
|
||||
computation_device=device,
|
||||
),
|
||||
compile_args = compile_args,
|
||||
)
|
||||
total_params_in_model = sum(p.numel() for p in patcher.model.diffusion_model.parameters())
|
||||
log.info(f"Total number of parameters in the loaded model: {total_params_in_model}")
|
||||
|
||||
offload_percent = vram_management_args["offload_percent"]
|
||||
offload_params = int(total_params_in_model * offload_percent)
|
||||
params_to_keep = total_params_in_model - offload_params
|
||||
log.info(f"Selected params to offload: {offload_params}")
|
||||
|
||||
enable_vram_management(
|
||||
patcher.model.diffusion_model,
|
||||
module_map = {
|
||||
torch.nn.Linear: AutoWrappedLinear,
|
||||
torch.nn.Conv3d: AutoWrappedModule,
|
||||
torch.nn.LayerNorm: AutoWrappedModule,
|
||||
WanLayerNorm: AutoWrappedModule,
|
||||
WanRMSNorm: AutoWrappedModule,
|
||||
},
|
||||
module_config = dict(
|
||||
offload_dtype=dtype,
|
||||
offload_device=offload_device,
|
||||
onload_dtype=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_device=offload_device,
|
||||
onload_dtype=dtype,
|
||||
onload_device=offload_device,
|
||||
computation_dtype=base_dtype,
|
||||
computation_device=device,
|
||||
),
|
||||
compile_args = compile_args,
|
||||
)
|
||||
|
||||
#compile
|
||||
if compile_args is not None and vram_management_args is None:
|
||||
@@ -824,66 +832,6 @@ class WanVideoModelLoader:
|
||||
gc.collect()
|
||||
mm.soft_empty_cache()
|
||||
|
||||
elif "torchao" in quantization:
|
||||
try:
|
||||
from torchao.quantization import (
|
||||
quantize_,
|
||||
fpx_weight_only,
|
||||
float8_dynamic_activation_float8_weight,
|
||||
int8_dynamic_activation_int8_weight,
|
||||
int8_weight_only,
|
||||
int4_weight_only
|
||||
)
|
||||
except:
|
||||
raise ImportError("torchao is not installed")
|
||||
|
||||
# def filter_fn(module: nn.Module, fqn: str) -> bool:
|
||||
# target_submodules = {'attn1', 'ff'} # avoid norm layers, 1.5 at least won't work with quantized norm1 #todo: test other models
|
||||
# if any(sub in fqn for sub in target_submodules):
|
||||
# return isinstance(module, nn.Linear)
|
||||
# return False
|
||||
|
||||
if "fp6" in quantization:
|
||||
quant_func = fpx_weight_only(3, 2)
|
||||
elif "int4" in quantization:
|
||||
quant_func = int4_weight_only()
|
||||
elif "int8" in quantization:
|
||||
quant_func = int8_weight_only()
|
||||
elif "fp8dq" in quantization:
|
||||
quant_func = float8_dynamic_activation_float8_weight()
|
||||
elif 'fp8dqrow' in quantization:
|
||||
from torchao.quantization.quant_api import PerRow
|
||||
quant_func = float8_dynamic_activation_float8_weight(granularity=PerRow())
|
||||
elif 'int8dq' in quantization:
|
||||
quant_func = int8_dynamic_activation_int8_weight()
|
||||
|
||||
log.info(f"Quantizing model with {quant_func}")
|
||||
comfy_model.diffusion_model = transformer
|
||||
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
|
||||
|
||||
for i, block in enumerate(patcher.model.diffusion_model.blocks):
|
||||
log.info(f"Quantizing block {i}")
|
||||
for name, _ in block.named_parameters(prefix=f"blocks.{i}"):
|
||||
#print(f"Parameter name: {name}")
|
||||
set_module_tensor_to_device(patcher.model.diffusion_model, name, device=transformer_load_device, dtype=base_dtype, value=sd[name])
|
||||
if compile_args is not None:
|
||||
patcher.model.diffusion_model.blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||
quantize_(block, quant_func)
|
||||
print(block)
|
||||
#block.to(offload_device)
|
||||
for name, param in patcher.model.diffusion_model.named_parameters():
|
||||
if "blocks" not in name:
|
||||
set_module_tensor_to_device(patcher.model.diffusion_model, name, device=transformer_load_device, dtype=base_dtype, value=sd[name])
|
||||
|
||||
manual_offloading = False # to disable manual .to(device) calls
|
||||
log.info(f"Quantized transformer blocks to {quantization}")
|
||||
for name, param in patcher.model.diffusion_model.named_parameters():
|
||||
print(name, param.dtype)
|
||||
#param.data = param.data.to(self.vae_dtype).to(device)
|
||||
|
||||
del sd
|
||||
mm.soft_empty_cache()
|
||||
|
||||
patcher.model["dtype"] = base_dtype
|
||||
patcher.model["base_path"] = model_path
|
||||
patcher.model["model_name"] = model
|
||||
@@ -2553,12 +2501,10 @@ class WanVideoSampler:
|
||||
latent_video_length = noise.shape[1]
|
||||
|
||||
if unianimate_poses is not None:
|
||||
transformer.dwpose_embedding.to(device)
|
||||
transformer.randomref_embedding_pose.to(device)
|
||||
dwpose_data = unianimate_poses["pose"]
|
||||
dwpose_data = transformer.dwpose_embedding(
|
||||
(torch.cat([dwpose_data[:,:,:1].repeat(1,1,3,1,1), dwpose_data], dim=2)
|
||||
).to(device)).to(model["dtype"])
|
||||
transformer.dwpose_embedding.to(device, model["dtype"])
|
||||
dwpose_data = unianimate_poses["pose"].to(device, model["dtype"])
|
||||
dwpose_data = torch.cat([dwpose_data[:,:,:1].repeat(1,1,3,1,1), dwpose_data], dim=2)
|
||||
dwpose_data = transformer.dwpose_embedding(dwpose_data)
|
||||
log.info(f"UniAnimate pose embed shape: {dwpose_data.shape}")
|
||||
if dwpose_data.shape[2] > latent_video_length:
|
||||
log.warning(f"UniAnimate pose embed length {dwpose_data.shape[2]} is longer than the video length {latent_video_length}, truncating")
|
||||
@@ -2572,6 +2518,7 @@ class WanVideoSampler:
|
||||
|
||||
random_ref_dwpose_data = None
|
||||
if image_cond is not None:
|
||||
transformer.randomref_embedding_pose.to(device)
|
||||
random_ref_dwpose = unianimate_poses.get("ref", None)
|
||||
if random_ref_dwpose is not None:
|
||||
random_ref_dwpose_data = transformer.randomref_embedding_pose(
|
||||
|
||||
Reference in New Issue
Block a user