WIP
This commit is contained in:
+129
-254
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user