diff --git a/wanvideo.py b/wanvideo.py index 8f4f327..39491f3 100644 --- a/wanvideo.py +++ b/wanvideo.py @@ -2,6 +2,7 @@ import logging import torch import inspect import copy +import importlib import folder_paths import comfy.model_management as mm from nodes import NODE_CLASS_MAPPINGS @@ -27,6 +28,38 @@ scheduler_list = [ rope_functions = ["default", "comfy", "comfy_chunked"] + +def _module_has_meta_tensors(module: torch.nn.Module) -> bool: + return any(p.is_meta for p in module.parameters()) or any(b.is_meta for b in module.buffers()) + + +def _patch_wanvideo_vram_management_meta_safe(loader_module): + layers_module = importlib.import_module(".diffsynth.vram_management.layers", package=loader_module.__package__) + + if getattr(layers_module, "_multigpu_meta_safe_patch_applied", False): + return + + def _auto_wrapped_module_init(self, module: torch.nn.Module, offload_dtype, offload_device, onload_dtype, onload_device, computation_dtype, computation_device): + torch.nn.Module.__init__(self) + + if _module_has_meta_tensors(module): + module = module.to_empty(device=offload_device) + module = module.to(dtype=offload_dtype) + else: + module = module.to(dtype=offload_dtype, device=offload_device) + + self.module = module + self.offload_dtype = offload_dtype + self.offload_device = offload_device + self.onload_dtype = onload_dtype + self.onload_device = onload_device + self.computation_dtype = computation_dtype + self.computation_device = computation_device + self.state = 0 + + layers_module.AutoWrappedModule.__init__ = _auto_wrapped_module_init + layers_module._multigpu_meta_safe_patch_applied = True + class WanVideoModelLoader: @classmethod def INPUT_TYPES(s): @@ -80,6 +113,9 @@ class WanVideoModelLoader: if kwargs.get("extra_model") is None and vace_model is not None: kwargs["extra_model"] = vace_model + if kwargs.get("vram_management_args") is not None: + _patch_wanvideo_vram_management_meta_safe(loader_module) + set_current_device(compute_device) compute_device_to_be_patched = mm.get_torch_device()