diff --git a/nodes.py b/nodes.py index 263c2ab..a33b382 100644 --- a/nodes.py +++ b/nodes.py @@ -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(