diff --git a/wanvideo.py b/wanvideo.py index 58b485a..5362f4e 100644 --- a/wanvideo.py +++ b/wanvideo.py @@ -17,6 +17,123 @@ import importlib.util logger = logging.getLogger("MultiGPU") + +class WanVideoModelLoader: + @classmethod + def INPUT_TYPES(s): + devices = get_device_list() + default_device = devices[1] if len(devices) > 1 else devices[0] + return { + "required": { + "model": (folder_paths.get_filename_list("unet_gguf") + folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' -folder",}), + + "base_precision": (["fp32", "bf16", "fp16", "fp16_fast"], {"default": "bf16"}), + "quantization": (["disabled", "fp8_e4m3fn", "fp8_e4m3fn_fast", "fp8_e4m3fn_scaled", "fp8_e4m3fn_scaled_fast", "fp8_e5m2", "fp8_e5m2_fast", "fp8_e5m2_scaled", "fp8_e5m2_scaled_fast"], {"default": "disabled", + "tooltip": "Optional quantization method, 'disabled' acts as autoselect based by weights. Scaled modes only work with matching weights, _fast modes (fp8 matmul) require CUDA compute capability >= 8.9 (NVIDIA 4000 series and up), e4m3fn generally can not be torch.compiled on compute capability < 8.9 (3000 series and under)"}), + "load_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}), + "compute_device": (devices, {"default": default_device}), + }, + "optional": { + "attention_mode": ([ + "sdpa", + "flash_attn_2", + "flash_attn_3", + "sageattn", + "sageattn_3", + "radial_sage_attention", + ], {"default": "sdpa"}), + "compile_args": ("WANCOMPILEARGS", ), + "block_swap_args": ("BLOCKSWAPARGS", ), + "lora": ("WANVIDLORA", {"default": None}), + "vram_management_args": ("VRAM_MANAGEMENTARGS", {"default": None, "tooltip": "Alternative offloading method from DiffSynth-Studio, more aggressive in reducing memory use than block swapping, but can be slower"}), + "extra_model": ("VACEPATH", {"default": None, "tooltip": "Extra model to add to the main model, ie. VACE or MTV Crafter"}), + "fantasytalking_model": ("FANTASYTALKINGMODEL", {"default": None, "tooltip": "FantasyTalking model https://github.com/Fantasy-AMAP"}), + "multitalk_model": ("MULTITALKMODEL", {"default": None, "tooltip": "Multitalk model"}), + "fantasyportrait_model": ("FANTASYPORTRAITMODEL", {"default": None, "tooltip": "FantasyPortrait model"}), + "rms_norm_function": (["default", "pytorch"], {"default": "default", "tooltip": "RMSNorm function to use, 'pytorch' is the new native torch RMSNorm, which is faster (when not using torch.compile mostly) but changes results slightly. 'default' is the original WanRMSNorm"}), + } + } + + RETURN_TYPES = ("WANVIDEOMODEL", "MULTIGPUDEVICE",) + RETURN_NAMES = ("model", "compute_device",) + FUNCTION = "loadmodel" + CATEGORY = "multigpu/WanVideoWrapper" + + def loadmodel(self, model, base_precision, compute_device, quantization, load_device, + compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None, + vram_management_args=None, extra_model=None, vace_model=None, + fantasytalking_model=None, multitalk_model=None, fantasyportrait_model=None, + rms_norm_function="default"): + from . import set_current_device + + set_current_device(compute_device) + compute_device_to_be_patched = mm.get_torch_device() + + original_loader = NODE_CLASS_MAPPINGS["WanVideoModelLoader"]() + loader_module = inspect.getmodule(original_loader) + + original_module_device = loader_module.device + + loader_module.device = compute_device_to_be_patched + + result = original_loader.loadmodel(model, base_precision, load_device, quantization, compile_args, attention_mode, block_swap_args, lora, vram_management_args, extra_model=extra_model, + vace_model=vace_model, fantasytalking_model=fantasytalking_model, multitalk_model=multitalk_model, fantasyportrait_model=fantasyportrait_model, rms_norm_function=rms_norm_function,) + + loader_module.device = original_module_device + + patcher = result[0] + + return (patcher, compute_device) + +class WanVideoTextEncode: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "positive_prompt": ("STRING", {"default": "", "multiline": True} ), + "negative_prompt": ("STRING", {"default": "", "multiline": True} ), + }, + "optional": { + "t5": ("WANTEXTENCODER",), + "load_device": ("MULTIGPUDEVICE",), + "force_offload": ("BOOLEAN", {"default": True}), + "model_to_offload": ("WANVIDEOMODEL", {"tooltip": "Model to move to offload_device before encoding"}), + "use_disk_cache": ("BOOLEAN", {"default": False, "tooltip": "Cache the text embeddings to disk for faster re-use, under the custom_nodes/ComfyUI-WanVideoWrapper/text_embed_cache directory"}), + } + } + + RETURN_TYPES = ("WANVIDEOTEXTEMBEDS", ) + RETURN_NAMES = ("text_embeds",) + FUNCTION = "process" + CATEGORY = "multigpu/WanVideoWrapper" + DESCRIPTION = "Encodes text prompts into text embeddings. For rudimentary prompt travel you can input multiple prompts separated by '|', they will be equally spread over the video length" + + def process(self, positive_prompt, negative_prompt, t5=None, load_device=None,force_offload=True, model_to_offload=None, use_disk_cache=False): + from . import set_current_device + + set_current_device(load_device) + + if load_device == "cpu": + device = "cpu" + else: + device = "gpu" + + if t5 is not None: + text_encoder = t5[0] + else: + text_encoder = None + + logger.info(f"[MultiGPU WanVideoWrapper][WanVideoTextEncodeMulitiGPU] current_device set to: {load_device}") + logger.info(f"[MultiGPU WanVideoWrapper][WanVideoTextEncodeMulitiGPU] device set to: {device}") + + original_encoder = NODE_CLASS_MAPPINGS["WanVideoTextEncode"]() + prompt_embeds_dict = original_encoder.process(positive_prompt, negative_prompt, text_encoder, force_offload, model_to_offload, use_disk_cache, device) + return (prompt_embeds_dict) + + def parse_prompt_weights(self, prompt): + """Extract text and weights from prompts with (text:weight) format""" + original_parser = NODE_CLASS_MAPPINGS["WanVideoTextEncode"]() + return original_parser.parse_prompt_weights(prompt) + class LoadWanVideoT5TextEncoder: @classmethod def INPUT_TYPES(s): @@ -58,7 +175,6 @@ class LoadWanVideoT5TextEncoder: original_loader = NODE_CLASS_MAPPINGS["LoadWanVideoT5TextEncoder"]() text_encoder = original_loader.loadmodel(model_name, precision, load_device, quantization) - # Return both the text encoder AND the selected device return text_encoder, device @@ -110,57 +226,6 @@ class WanVideoTextEncodeCached: return prompt_embeds_dict, negative_text_embeds, positive_prompt_out - -class WanVideoTextEncode: - @classmethod - def INPUT_TYPES(s): - return {"required": { - "positive_prompt": ("STRING", {"default": "", "multiline": True} ), - "negative_prompt": ("STRING", {"default": "", "multiline": True} ), - }, - "optional": { - "t5": ("WANTEXTENCODER",), - "load_device": ("MULTIGPUDEVICE",), - "force_offload": ("BOOLEAN", {"default": True}), - "model_to_offload": ("WANVIDEOMODEL", {"tooltip": "Model to move to offload_device before encoding"}), - "use_disk_cache": ("BOOLEAN", {"default": False, "tooltip": "Cache the text embeddings to disk for faster re-use, under the custom_nodes/ComfyUI-WanVideoWrapper/text_embed_cache directory"}), - } - } - - RETURN_TYPES = ("WANVIDEOTEXTEMBEDS", ) - RETURN_NAMES = ("text_embeds",) - FUNCTION = "process" - CATEGORY = "multigpu/WanVideoWrapper" - DESCRIPTION = "Encodes text prompts into text embeddings. For rudimentary prompt travel you can input multiple prompts separated by '|', they will be equally spread over the video length" - - def process(self, positive_prompt, negative_prompt, t5=None, load_device=None,force_offload=True, model_to_offload=None, use_disk_cache=False): - from . import set_current_device - - if load_device is not None: - set_current_device(load_device) - - if load_device == "cpu": - device = "cpu" - else: - device = "gpu" - - if t5 is not None: - text_encoder = t5[0] - else: - text_encoder = None - - logger.info(f"[MultiGPU WanVideoWrapper][WanVideoTextEncodeMulitiGPU] current_device set to: {load_device}") - logger.info(f"[MultiGPU WanVideoWrapper][WanVideoTextEncodeMulitiGPU] device set to: {device}") - - original_encoder = NODE_CLASS_MAPPINGS["WanVideoTextEncode"]() - prompt_embeds_dict = original_encoder.process(positive_prompt, negative_prompt, text_encoder, force_offload, model_to_offload, use_disk_cache, device) - return (prompt_embeds_dict) - - def parse_prompt_weights(self, prompt): - """Extract text and weights from prompts with (text:weight) format""" - original_parser = NODE_CLASS_MAPPINGS["WanVideoTextEncode"]() - return original_parser.parse_prompt_weights(prompt) - class WanVideoTextEncodeSingle: @classmethod def INPUT_TYPES(s): @@ -235,7 +300,7 @@ class WanVideoVAELoader: if load_device is not None: set_current_device(load_device) - logger.info(f"[MultiGPU WanVideoWrapper][WanVideoVAELoader] load_device set to: {load_device}") + logger.info(f"[MultiGPU WanVideoWrapper][WanVideoVAELoaderMultiGPU] load_device set to: {load_device}") original_loader = NODE_CLASS_MAPPINGS["WanVideoVAELoader"]() vae_model = original_loader.loadmodel(model_name, precision, compile_args) @@ -345,8 +410,7 @@ class WanVideoImageToVideoEncode: temporal_mask=None, extra_latents=None, clip_embeds=None, tiled_vae=False, add_cond_latents=None, vae=None, load_device=None): from . import set_current_device - if load_device is not None: - set_current_device(load_device) + set_current_device(load_device) logger.info(f"[MultiGPU WanVideoWrapper][WanVideoImageToVideoEncodeMultiGPU] load device: {load_device}") @@ -537,214 +601,25 @@ class WanVideoDecode: def decode(self, vae, load_device, samples, enable_vae_tiling, tile_x, tile_y, tile_stride_x, tile_stride_y, normalization="default"): from . import set_current_device - if load_device is not None: - set_current_device(load_device) + set_current_device(load_device) + compute_device_to_be_patched = mm.get_torch_device() - logger.info(f"[MultiGPU WanVideoWrapper][WanVideoImageToVideoEncodeMultiGPU] load device: {load_device}") + logger.info(f"[MultiGPU WanVideoWrapper][WanVideoDecodeMultiGPU] load device: {load_device}") - device = mm.get_torch_device() - PATCH_SIZE = (1, 2, 2) - offload_device = mm.unet_offload_device() - - logger.info(f"[MultiGPU WanVideoWrapper][WanVideoImageToVideoEncodeMultiGPU] torch device: {device}") - - if vae is not None: - vae = vae[0] - - mm.soft_empty_cache() - video = samples.get("video", None) - if video is not None: - video.clamp_(-1.0, 1.0) - video.add_(1.0).div_(2.0) - return video.cpu().float(), - latents = samples["samples"] - end_image = samples.get("end_image", None) - has_ref = samples.get("has_ref", False) - drop_last = samples.get("drop_last", False) - is_looped = samples.get("looped", False) - - vae.to(device) - - latents = latents.to(device = device, dtype = vae.dtype) - - mm.soft_empty_cache() - - if has_ref: - latents = latents[:, :, 1:] - if drop_last: - latents = latents[:, :, :-1] - - if type(vae).__name__ == "TAEHV": - images = vae.decode_video(latents.permute(0, 2, 1, 3, 4))[0].permute(1, 0, 2, 3) - images = torch.clamp(images, 0.0, 1.0) - images = images.permute(1, 2, 3, 0).cpu().float() - return (images,) - else: - if end_image is not None: - enable_vae_tiling = False - images = vae.decode(latents, device=device, end_=(end_image is not None), tiled=enable_vae_tiling, tile_size=(tile_x//8, tile_y//8), tile_stride=(tile_stride_x//8, tile_stride_y//8))[0] - - - images = images.cpu().float() - - if normalization == "minmax": - images.sub_(images.min()).div_(images.max() - images.min()) - else: - images.clamp_(-1.0, 1.0) - images.add_(1.0).div_(2.0) - - if is_looped: - temp_latents = torch.cat([latents[:, :, -3:]] + [latents[:, :, :2]], dim=2) - temp_images = vae.decode(temp_latents, device=device, end_=(end_image is not None), tiled=enable_vae_tiling, tile_size=(tile_x//vae.upsampling_factor, tile_y//vae.upsampling_factor), tile_stride=(tile_stride_x//vae.upsampling_factor, tile_stride_y//vae.upsampling_factor))[0] - temp_images = temp_images.cpu().float() - temp_images = (temp_images - temp_images.min()) / (temp_images.max() - temp_images.min()) - images = torch.cat([temp_images[:, 9:].to(images), images[:, 5:]], dim=1) - - if end_image is not None: - images = images[:, 0:-1] - - - vae.to(offload_device) - mm.soft_empty_cache() - - images.clamp_(0.0, 1.0) - - return (images.permute(1, 2, 3, 0),) -class WanVideoModelLoader: - @classmethod - def INPUT_TYPES(s): - devices = get_device_list() - default_device = devices[1] if len(devices) > 1 else devices[0] - # Get the original node's input types to stay up-to-date - original_types = NODE_CLASS_MAPPINGS["WanVideoModelLoader"].INPUT_TYPES() - - # Update with our custom device selection - original_types["required"]["compute_device"] = (devices, {"default": default_device}) - - return original_types - - RETURN_TYPES = ("WANVIDEOMODEL", "MULTIGPUDEVICE",) - RETURN_NAMES = ("model", "compute_device",) - FUNCTION = "loadmodel" - CATEGORY = "multigpu/WanVideoWrapper" - - def loadmodel(self, model, base_precision, compute_device, quantization, load_device, - compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None, - vram_management_args=None, extra_model=None, vace_model=None, - fantasytalking_model=None, multitalk_model=None, fantasyportrait_model=None, - rms_norm_function="default"): - from . import set_current_device - logger.info( - f"[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] User selected device: {compute_device}" - ) - - selected_device = torch.device(compute_device) - set_current_device(selected_device) - - normalized_block_swap = None - swap_device_override = None - if block_swap_args is not None: - normalized_block_swap = dict(block_swap_args) - swap_selection = normalized_block_swap.pop("swap_device", None) - if swap_selection is None: - swap_selection = normalized_block_swap.get("resolved_swap_device") - if swap_selection is None: - swap_selection = "cpu" - try: - swap_device_override = torch.device(str(swap_selection)) - except (TypeError, ValueError): - logger.warning( - "[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] Invalid swap_device '%s', falling back to CPU", - swap_selection, - ) - swap_device_override = torch.device("cpu") - normalized_block_swap["resolved_swap_device"] = str(swap_device_override) - - original_loader = NODE_CLASS_MAPPINGS["WanVideoModelLoader"]() + original_loader = NODE_CLASS_MAPPINGS["WanVideoDecode"]() loader_module = inspect.getmodule(original_loader) - if not loader_module: - logger.error( - "[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] Could not resolve loader module; invoking original implementation without patches." - ) - result = original_loader.loadmodel( - model, - base_precision, - load_device, - quantization, - compile_args, - attention_mode, - normalized_block_swap if normalized_block_swap is not None else block_swap_args, - lora, - vram_management_args, - extra_model=extra_model, - vace_model=vace_model, - fantasytalking_model=fantasytalking_model, - multitalk_model=multitalk_model, - fantasyportrait_model=fantasyportrait_model, - rms_norm_function=rms_norm_function, - ) - return (result[0], compute_device) + original_module_device = loader_module.device - logger.debug( - f"[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] Patching '{loader_module.__name__}'" - ) - original_module_device = getattr(loader_module, "device", None) - had_offload_attr = hasattr(loader_module, "offload_device") - original_module_offload = getattr(loader_module, "offload_device", None) + loader_module.device = compute_device_to_be_patched - setattr(loader_module, "device", selected_device) - if swap_device_override is not None: - setattr(loader_module, "offload_device", swap_device_override) - elif compute_device == "cpu": - setattr(loader_module, "offload_device", selected_device) + result = original_loader.decode(vae[0], samples, enable_vae_tiling, tile_x, tile_y, tile_stride_x, tile_stride_y, normalization) - try: - result = original_loader.loadmodel( - model, - base_precision, - load_device, - quantization, - compile_args, - attention_mode, - normalized_block_swap if normalized_block_swap is not None else block_swap_args, - lora, - vram_management_args, - extra_model=extra_model, - vace_model=vace_model, - fantasytalking_model=fantasytalking_model, - multitalk_model=multitalk_model, - fantasyportrait_model=fantasyportrait_model, - rms_norm_function=rms_norm_function, - ) - finally: - if original_module_device is not None: - setattr(loader_module, "device", original_module_device) - if had_offload_attr: - setattr(loader_module, "offload_device", original_module_offload) - else: - try: - delattr(loader_module, "offload_device") - except AttributeError: - pass + loader_module.device = original_module_device - patcher = result[0] - if normalized_block_swap is not None: - try: - transformer_options = patcher.model_options.setdefault("transformer_options", {}) - transformer_options["block_swap_args"] = normalized_block_swap - except AttributeError: - logger.warning( - "[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] Unable to propagate normalized block swap settings" - ) + decode = result[0] - logger.info( - "[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] WanVideo model loaded on %s with swap_device=%s", - selected_device, - str(swap_device_override) if swap_device_override is not None else "default", - ) - - return (patcher, compute_device) + return (decode,) class WanVideoSampler: