From 6183d41df8fbeb44f24b6c617efb8ef0ab690052 Mon Sep 17 00:00:00 2001 From: Bubbliiiing <47347516+bubbliiiing@users.noreply.github.com> Date: Sun, 28 Sep 2025 14:45:22 +0800 Subject: [PATCH] Update unload model comfyui for saving VRAM. (#340) --- comfyui/cogvideox_fun/nodes.py | 10 +++++- comfyui/comfyui_nodes.py | 25 +++++++++++++++ comfyui/qwenimage/nodes.py | 35 +++++++++++++++------ comfyui/wan2_1/nodes.py | 31 +++++++++++++++--- comfyui/wan2_1_fun/nodes.py | 33 +++++++++++-------- comfyui/wan2_2/nodes.py | 47 +++++++++++++++++++--------- comfyui/wan2_2_fun/nodes.py | 45 +++++++++++++++----------- comfyui/wan2_2_vace_fun/nodes.py | 35 +++++++++++++++------ videox_fun/utils/__init__.py | 2 +- videox_fun/utils/fp8_optimization.py | 8 ++++- videox_fun/utils/utils.py | 26 +++++++++++++++ 11 files changed, 224 insertions(+), 73 deletions(-) diff --git a/comfyui/cogvideox_fun/nodes.py b/comfyui/cogvideox_fun/nodes.py index 22a4ac9..a6c97b0 100755 --- a/comfyui/cogvideox_fun/nodes.py +++ b/comfyui/cogvideox_fun/nodes.py @@ -28,7 +28,7 @@ from ...videox_fun.utils.lora_utils import merge_lora, unmerge_lora from ...videox_fun.utils.utils import (get_image_to_video_latent, get_video_to_video_latent, save_videos_grid) -from ...videox_fun.utils.fp8_optimization import convert_weight_dtype_wrapper +from ...videox_fun.utils.fp8_optimization import convert_weight_dtype_wrapper, undo_convert_weight_dtype_wrapper from ..comfyui_utils import (eas_cache_dir, script_directory, search_model_in_possible_folders, to_pil) @@ -90,6 +90,10 @@ class LoadCogVideoXFunModel: offload_device = mm.unet_offload_device() weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] + mm.unload_all_models() + mm.cleanup_models() + mm.soft_empty_cache() + # Init processbar pbar = ProgressBar(5) @@ -158,6 +162,10 @@ class LoadCogVideoXFunModel: transformer=transformer, scheduler=scheduler, ) + + pipeline.remove_all_hooks() + undo_convert_weight_dtype_wrapper(transformer) + if GPU_memory_mode == "sequential_cpu_offload": pipeline.enable_sequential_cpu_offload() elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": diff --git a/comfyui/comfyui_nodes.py b/comfyui/comfyui_nodes.py index 49716d2..7dde5c2 100755 --- a/comfyui/comfyui_nodes.py +++ b/comfyui/comfyui_nodes.py @@ -88,6 +88,31 @@ class FunCompile: def compile(self, cache_size_limit, funmodels): torch._dynamo.config.cache_size_limit = cache_size_limit + + if funmodels["pipeline"].transformer.device == torch.device(type="meta"): + if hasattr(funmodels["pipeline"].transformer, "blocks"): + for i, block in enumerate(funmodels["pipeline"].transformer.blocks): + if hasattr(block, "_orig_mod"): + block = block._orig_mod + + if hasattr(funmodels["pipeline"], "transformer_2") and funmodels["pipeline"].transformer_2 is not None: + for i, block in enumerate(funmodels["pipeline"].transformer_2.blocks): + if hasattr(block, "_orig_mod"): + block = block._orig_mod + + elif hasattr(funmodels["pipeline"].transformer, "transformer_blocks"): + for i, block in enumerate(funmodels["pipeline"].transformer.transformer_blocks): + if hasattr(block, "_orig_mod"): + block = block._orig_mod + + if hasattr(funmodels["pipeline"], "transformer_2") and funmodels["pipeline"].transformer_2 is not None: + for i, block in enumerate(funmodels["pipeline"].transformer_2.transformer_blocks): + if hasattr(block, "_orig_mod"): + block = block._orig_mod + + print("Sequential cpu offload can not work with compile. Continue") + return (funmodels,) + if hasattr(funmodels["pipeline"].transformer, "blocks"): for i, block in enumerate(funmodels["pipeline"].transformer.blocks): if hasattr(block, "_orig_mod"): diff --git a/comfyui/qwenimage/nodes.py b/comfyui/qwenimage/nodes.py index 30d3364..1085f32 100644 --- a/comfyui/qwenimage/nodes.py +++ b/comfyui/qwenimage/nodes.py @@ -27,10 +27,10 @@ from ...videox_fun.pipeline import QwenImagePipeline, QwenImageEditPipeline from ...videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler from ...videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler from ...videox_fun.utils.fp8_optimization import ( - convert_model_weight_to_float8, convert_weight_dtype_wrapper, + convert_model_weight_to_float8, convert_weight_dtype_wrapper, undo_convert_weight_dtype_wrapper, replace_parameters_by_name) from ...videox_fun.utils.lora_utils import merge_lora, unmerge_lora -from ...videox_fun.utils.utils import filter_kwargs, get_image +from ...videox_fun.utils.utils import filter_kwargs, get_image, get_autocast_dtype from ..comfyui_utils import (eas_cache_dir, script_directory, to_pil, search_model_in_possible_folders, search_sub_dir_in_possible_folders) @@ -93,6 +93,11 @@ class LoadQwenImageTransformerModel: offload_device = mm.unet_offload_device() weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16}[precision] + mm.unload_all_models() + mm.cleanup_models() + mm.soft_empty_cache() + transformer = None + model_path = folder_paths.get_full_path("diffusion_models", model_name) transformer_state_dict = load_torch_file(model_path, safe_load=True) @@ -471,7 +476,7 @@ class CombineQwenImagePipeline: def loadmodel(self, model_name, GPU_memory_mode, transformer, vae, text_encoder, tokenizer, processor=None, transformer_2=None): # Get pipeline - weight_dtype = transformer.dtype + weight_dtype = transformer.dtype if transformer.dtype not in [torch.float32, torch.float8_e4m3fn, torch.float8_e5m2] else get_autocast_dtype() device = mm.get_torch_device() offload_device = mm.unet_offload_device() @@ -498,6 +503,11 @@ class CombineQwenImagePipeline: else: raise ValueError("Not supported now.") + pipeline.remove_all_hooks() + undo_convert_weight_dtype_wrapper(transformer) + pipeline.to(device=offload_device) + transformer = transformer.to(weight_dtype) + if GPU_memory_mode == "sequential_cpu_offload": pipeline.enable_sequential_cpu_offload(device=device) elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": @@ -563,6 +573,10 @@ class LoadQwenImageModel: offload_device = mm.unet_offload_device() weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] + mm.unload_all_models() + mm.cleanup_models() + mm.soft_empty_cache() + # Init processbar pbar = ProgressBar(5) @@ -642,6 +656,9 @@ class LoadQwenImageModel: else: raise ValueError("Not supported now.") + pipeline.remove_all_hooks() + undo_convert_weight_dtype_wrapper(transformer) + if GPU_memory_mode == "sequential_cpu_offload": pipeline.enable_sequential_cpu_offload(device=device) elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": @@ -808,7 +825,7 @@ class QwenImageT2VSampler: lora_path_before = copy.deepcopy(lora_path_now) pipeline.transformer.load_state_dict(transformer_cpu_cache) for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) else: print('Merge Lora') @@ -820,7 +837,7 @@ class QwenImageT2VSampler: gc.collect() for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) sample = pipeline( prompt, @@ -837,7 +854,7 @@ class QwenImageT2VSampler: if not funmodels.get("lora_cache", False): print('Unmerge Lora') for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) return (image,) class QwenImageEditSampler: @@ -955,7 +972,7 @@ class QwenImageEditSampler: lora_path_before = copy.deepcopy(lora_path_now) pipeline.transformer.load_state_dict(transformer_cpu_cache) for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) else: print('Merge Lora') @@ -967,7 +984,7 @@ class QwenImageEditSampler: gc.collect() for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) image = [to_pil(image) for image in image] image = get_image(image[0]) if image is not None else image @@ -988,5 +1005,5 @@ class QwenImageEditSampler: if not funmodels.get("lora_cache", False): print('Unmerge Lora') for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) return (image,) diff --git a/comfyui/wan2_1/nodes.py b/comfyui/wan2_1/nodes.py index 32c5945..a01d850 100755 --- a/comfyui/wan2_1/nodes.py +++ b/comfyui/wan2_1/nodes.py @@ -29,13 +29,13 @@ from ...videox_fun.ui.controller import all_cheduler_dict from ...videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler from ...videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler from ...videox_fun.utils.fp8_optimization import ( - convert_model_weight_to_float8, convert_weight_dtype_wrapper, + convert_model_weight_to_float8, convert_weight_dtype_wrapper, undo_convert_weight_dtype_wrapper, replace_parameters_by_name) from ...videox_fun.utils.lora_utils import merge_lora, unmerge_lora from ...videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent, get_video_to_video_latent, - save_videos_grid) + save_videos_grid, get_autocast_dtype) from ..comfyui_utils import (eas_cache_dir, script_directory, search_sub_dir_in_possible_folders, search_model_in_possible_folders, to_pil) @@ -89,7 +89,16 @@ class LoadWanTransformerModel: # Init weight_dtype and device device = mm.get_torch_device() offload_device = mm.unet_offload_device() - weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16}[precision] + weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] + + mm.unload_all_models() + mm.cleanup_models() + mm.soft_empty_cache() + + mm.unload_all_models() + mm.cleanup_models() + mm.soft_empty_cache() + transformer = None model_path = folder_paths.get_full_path("diffusion_models", model_name) transformer_state_dict = load_torch_file(model_path, safe_load=True) @@ -376,7 +385,7 @@ class CombineWanPipeline: def loadmodel(self, model_name, GPU_memory_mode, model_type, transformer, vae, text_encoder, tokenizer, clip_encoder=None): # Get pipeline - weight_dtype = transformer.dtype + weight_dtype = transformer.dtype if transformer.dtype not in [torch.float32, torch.float8_e4m3fn, torch.float8_e5m2] else get_autocast_dtype() device = mm.get_torch_device() offload_device = mm.unet_offload_device() @@ -408,6 +417,11 @@ class CombineWanPipeline: clip_image_encoder=clip_encoder ) + pipeline.remove_all_hooks() + undo_convert_weight_dtype_wrapper(transformer) + pipeline.to(device=offload_device) + transformer = transformer.to(weight_dtype) + if GPU_memory_mode == "sequential_cpu_offload": replace_parameters_by_name(transformer, ["modulation",], device=device) transformer.freqs = transformer.freqs.to(device=device) @@ -483,7 +497,11 @@ class LoadWanModel: # Init weight_dtype and device device = mm.get_torch_device() offload_device = mm.unet_offload_device() - weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16}[precision] + weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] + + mm.unload_all_models() + mm.cleanup_models() + mm.soft_empty_cache() # Init processbar pbar = ProgressBar(5) @@ -567,6 +585,9 @@ class LoadWanModel: else: raise ValueError(f"Model type {model_type} not supported") + pipeline.remove_all_hooks() + undo_convert_weight_dtype_wrapper(transformer) + if GPU_memory_mode == "sequential_cpu_offload": replace_parameters_by_name(transformer, ["modulation",], device=device) transformer.freqs = transformer.freqs.to(device=device) diff --git a/comfyui/wan2_1_fun/nodes.py b/comfyui/wan2_1_fun/nodes.py index 91bb64f..e5704d9 100755 --- a/comfyui/wan2_1_fun/nodes.py +++ b/comfyui/wan2_1_fun/nodes.py @@ -27,7 +27,7 @@ from ...videox_fun.pipeline import (WanFunControlPipeline, WanFunInpaintPipeline, WanFunPipeline) from ...videox_fun.ui.controller import all_cheduler_dict from ...videox_fun.utils.fp8_optimization import ( - convert_model_weight_to_float8, convert_weight_dtype_wrapper, + convert_model_weight_to_float8, convert_weight_dtype_wrapper, undo_convert_weight_dtype_wrapper, replace_parameters_by_name) from ...videox_fun.utils.lora_utils import merge_lora, unmerge_lora from ...videox_fun.utils.utils import (filter_kwargs, get_image_latent, @@ -105,6 +105,10 @@ class LoadWanFunModel: offload_device = mm.unet_offload_device() weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] + mm.unload_all_models() + mm.cleanup_models() + mm.soft_empty_cache() + # Init processbar pbar = ProgressBar(5) @@ -193,9 +197,12 @@ class LoadWanFunModel: clip_image_encoder=clip_image_encoder ) + pipeline.remove_all_hooks() + undo_convert_weight_dtype_wrapper(transformer) + if GPU_memory_mode == "sequential_cpu_offload": - replace_parameters_by_name(transformer, ["modulation",], device="cuda") - transformer.freqs = transformer.freqs.to(device="cuda") + replace_parameters_by_name(transformer, ["modulation",], device=device) + transformer.freqs = transformer.freqs.to(device=device) pipeline.enable_sequential_cpu_offload() elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",]) @@ -208,7 +215,7 @@ class LoadWanFunModel: convert_weight_dtype_wrapper(transformer, weight_dtype) pipeline.to(device=device) else: - pipeline.to("cuda") + pipeline.to(device) funmodels = { 'pipeline': pipeline, @@ -380,7 +387,7 @@ class WanFunT2VSampler: lora_path_before = copy.deepcopy(lora_path_now) pipeline.transformer.load_state_dict(transformer_cpu_cache) for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) else: # Clear lora when switch from lora_cache=True to lora_cache=False. if len(transformer_cpu_cache) != 0: @@ -390,7 +397,7 @@ class WanFunT2VSampler: gc.collect() print('Merge Lora') for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) if pipeline.transformer.config.in_channels != pipeline.vae.config.latent_channels: input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=(height, width)) @@ -426,7 +433,7 @@ class WanFunT2VSampler: if not funmodels.get("lora_cache", False): print('Unmerge Lora') for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) return (videos,) @@ -567,7 +574,7 @@ class WanFunInpaintSampler: lora_path_before = copy.deepcopy(lora_path_now) pipeline.transformer.load_state_dict(transformer_cpu_cache) for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) else: # Clear lora when switch from lora_cache=True to lora_cache=False. if len(transformer_cpu_cache) != 0: @@ -578,7 +585,7 @@ class WanFunInpaintSampler: gc.collect() print('Merge Lora') for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) sample = pipeline( prompt, @@ -600,7 +607,7 @@ class WanFunInpaintSampler: if not funmodels.get("lora_cache", False): print('Unmerge Lora') for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) return (videos,) @@ -794,7 +801,7 @@ class WanFunV2VSampler: lora_path_before = copy.deepcopy(lora_path_now) pipeline.transformer.load_state_dict(transformer_cpu_cache) for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) else: # Clear lora when switch from lora_cache=True to lora_cache=False. if len(transformer_cpu_cache) != 0: @@ -804,7 +811,7 @@ class WanFunV2VSampler: gc.collect() print('Merge Lora') for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) if model_type == "Inpaint": sample = pipeline( @@ -846,6 +853,6 @@ class WanFunV2VSampler: if not funmodels.get("lora_cache", False): print('Unmerge Lora') for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) return (videos,) diff --git a/comfyui/wan2_2/nodes.py b/comfyui/wan2_2/nodes.py index 8d86fdb..b8153cc 100755 --- a/comfyui/wan2_2/nodes.py +++ b/comfyui/wan2_2/nodes.py @@ -30,13 +30,13 @@ from ...videox_fun.pipeline import (Wan2_2FunControlPipeline, Wan2_2Pipeline, Wan2_2TI2VPipeline) from ...videox_fun.ui.controller import all_cheduler_dict from ...videox_fun.utils.fp8_optimization import ( - convert_model_weight_to_float8, convert_weight_dtype_wrapper, + convert_model_weight_to_float8, convert_weight_dtype_wrapper, undo_convert_weight_dtype_wrapper, replace_parameters_by_name) from ...videox_fun.utils.lora_utils import merge_lora, unmerge_lora from ...videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent, get_video_to_video_latent, - save_videos_grid) + save_videos_grid, get_autocast_dtype) from ..wan2_1.nodes import get_wan_scheduler from ..comfyui_utils import (eas_cache_dir, script_directory, search_model_in_possible_folders, to_pil) @@ -73,6 +73,11 @@ class LoadWan2_2TransformerModel: offload_device = mm.unet_offload_device() weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16}[precision] + mm.unload_all_models() + mm.cleanup_models() + mm.soft_empty_cache() + transformer = None + model_path = folder_paths.get_full_path("diffusion_models", model_name) transformer_state_dict = load_torch_file(model_path, safe_load=True) @@ -186,7 +191,7 @@ class CombineWan2_2Pipeline: def loadmodel(self, model_name, GPU_memory_mode, model_type, transformer, vae, text_encoder, tokenizer, clip_encoder=None, transformer_2=None): # Get pipeline - weight_dtype = transformer.dtype + weight_dtype = transformer.dtype if transformer.dtype not in [torch.float32, torch.float8_e4m3fn, torch.float8_e5m2] else get_autocast_dtype() device = mm.get_torch_device() offload_device = mm.unet_offload_device() @@ -230,6 +235,11 @@ class CombineWan2_2Pipeline: scheduler=None, ) + pipeline.remove_all_hooks() + undo_convert_weight_dtype_wrapper(transformer) + pipeline.to(device=offload_device) + transformer = transformer.to(weight_dtype) + if GPU_memory_mode == "sequential_cpu_offload": replace_parameters_by_name(transformer, ["modulation",], device=device) transformer.freqs = transformer.freqs.to(device=device) @@ -308,6 +318,10 @@ class LoadWan2_2Model: offload_device = mm.unet_offload_device() weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] + mm.unload_all_models() + mm.cleanup_models() + mm.soft_empty_cache() + # Init processbar pbar = ProgressBar(5) @@ -408,6 +422,9 @@ class LoadWan2_2Model: else: raise ValueError(f"Model type {model_type} not supported") + pipeline.remove_all_hooks() + undo_convert_weight_dtype_wrapper(transformer) + if GPU_memory_mode == "sequential_cpu_offload": replace_parameters_by_name(transformer, ["modulation",], device=device) transformer.freqs = transformer.freqs.to(device=device) @@ -619,7 +636,7 @@ class Wan2_2T2VSampler: lora_path_before = copy.deepcopy(lora_path_now) pipeline.transformer.load_state_dict(transformer_cpu_cache) for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) if pipeline.transformer_2 is not None: # Save the original weights to cpu @@ -635,7 +652,7 @@ class Wan2_2T2VSampler: lora_high_path_before = copy.deepcopy(lora_high_path_now) pipeline.transformer_2.load_state_dict(transformer_cpu_cache) for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2") else: print('Merge Lora') # Clear lora when switch from lora_cache=True to lora_cache=False. @@ -646,7 +663,7 @@ class Wan2_2T2VSampler: gc.collect() for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) # Clear lora when switch from lora_cache=True to lora_cache=False. if pipeline.transformer_2 is not None: @@ -657,7 +674,7 @@ class Wan2_2T2VSampler: gc.collect() for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2") sample = pipeline( prompt, @@ -676,10 +693,10 @@ class Wan2_2T2VSampler: if not funmodels.get("lora_cache", False): print('Unmerge Lora') for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) if pipeline.transformer_2 is not None: for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2") return (videos,) @@ -832,7 +849,7 @@ class Wan2_2I2VSampler: lora_path_before = copy.deepcopy(lora_path_now) pipeline.transformer.load_state_dict(transformer_cpu_cache) for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) if pipeline.transformer_2 is not None: # Save the original weights to cpu @@ -848,7 +865,7 @@ class Wan2_2I2VSampler: lora_high_path_before = copy.deepcopy(lora_high_path_now) pipeline.transformer_2.load_state_dict(transformer_cpu_cache) for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2") else: print('Merge Lora') # Clear lora when switch from lora_cache=True to lora_cache=False. @@ -859,7 +876,7 @@ class Wan2_2I2VSampler: gc.collect() for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) # Clear lora when switch from lora_cache=True to lora_cache=False. if pipeline.transformer_2 is not None: @@ -870,7 +887,7 @@ class Wan2_2I2VSampler: gc.collect() for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2") sample = pipeline( prompt, @@ -892,8 +909,8 @@ class Wan2_2I2VSampler: if not funmodels.get("lora_cache", False): print('Unmerge Lora') for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) if pipeline.transformer_2 is not None: for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2") return (videos,) \ No newline at end of file diff --git a/comfyui/wan2_2_fun/nodes.py b/comfyui/wan2_2_fun/nodes.py index a862705..02ab6fd 100755 --- a/comfyui/wan2_2_fun/nodes.py +++ b/comfyui/wan2_2_fun/nodes.py @@ -29,7 +29,7 @@ from ...videox_fun.pipeline import (Wan2_2FunControlPipeline, Wan2_2FunPipeline) from ...videox_fun.ui.controller import all_cheduler_dict from ...videox_fun.utils.fp8_optimization import ( - convert_model_weight_to_float8, convert_weight_dtype_wrapper, + convert_model_weight_to_float8, convert_weight_dtype_wrapper, undo_convert_weight_dtype_wrapper, replace_parameters_by_name) from ...videox_fun.utils.lora_utils import merge_lora, unmerge_lora from ...videox_fun.utils.utils import (filter_kwargs, get_image_latent, @@ -107,6 +107,10 @@ class LoadWan2_2FunModel: offload_device = mm.unet_offload_device() weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] + mm.unload_all_models() + mm.cleanup_models() + mm.soft_empty_cache() + # Init processbar pbar = ProgressBar(5) @@ -203,6 +207,9 @@ class LoadWan2_2FunModel: scheduler=scheduler, ) + pipeline.remove_all_hooks() + undo_convert_weight_dtype_wrapper(transformer) + if GPU_memory_mode == "sequential_cpu_offload": replace_parameters_by_name(transformer, ["modulation",], device=device) transformer.freqs = transformer.freqs.to(device=device) @@ -415,7 +422,7 @@ class Wan2_2FunT2VSampler: lora_path_before = copy.deepcopy(lora_path_now) pipeline.transformer.load_state_dict(transformer_cpu_cache) for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) if pipeline.transformer_2 is not None: # Save the original weights to cpu @@ -431,7 +438,7 @@ class Wan2_2FunT2VSampler: lora_high_path_before = copy.deepcopy(lora_high_path_now) pipeline.transformer_2.load_state_dict(transformer_cpu_cache) for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2") else: print('Merge Lora') # Clear lora when switch from lora_cache=True to lora_cache=False. @@ -442,7 +449,7 @@ class Wan2_2FunT2VSampler: gc.collect() for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) # Clear lora when switch from lora_cache=True to lora_cache=False. if pipeline.transformer_2 is not None: @@ -453,7 +460,7 @@ class Wan2_2FunT2VSampler: gc.collect() for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2") sample = pipeline( prompt, @@ -475,10 +482,10 @@ class Wan2_2FunT2VSampler: if not funmodels.get("lora_cache", False): print('Unmerge Lora') for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) if pipeline.transformer_2 is not None: for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2") return (videos,) @@ -633,7 +640,7 @@ class Wan2_2FunInpaintSampler: lora_path_before = copy.deepcopy(lora_path_now) pipeline.transformer.load_state_dict(transformer_cpu_cache) for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) if pipeline.transformer_2 is not None: # Save the original weights to cpu @@ -649,7 +656,7 @@ class Wan2_2FunInpaintSampler: lora_high_path_before = copy.deepcopy(lora_high_path_now) pipeline.transformer_2.load_state_dict(transformer_cpu_cache) for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2") else: print('Merge Lora') # Clear lora when switch from lora_cache=True to lora_cache=False. @@ -660,7 +667,7 @@ class Wan2_2FunInpaintSampler: gc.collect() for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) # Clear lora when switch from lora_cache=True to lora_cache=False. if pipeline.transformer_2 is not None: @@ -671,7 +678,7 @@ class Wan2_2FunInpaintSampler: gc.collect() for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2") sample = pipeline( prompt, @@ -693,10 +700,10 @@ class Wan2_2FunInpaintSampler: if not funmodels.get("lora_cache", False): print('Unmerge Lora') for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) if pipeline.transformer_2 is not None: for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2") return (videos,) @@ -906,7 +913,7 @@ class Wan2_2FunV2VSampler: lora_path_before = copy.deepcopy(lora_path_now) pipeline.transformer.load_state_dict(transformer_cpu_cache) for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) if pipeline.transformer_2 is not None: # Save the original weights to cpu @@ -922,7 +929,7 @@ class Wan2_2FunV2VSampler: lora_high_path_before = copy.deepcopy(lora_high_path_now) pipeline.transformer_2.load_state_dict(transformer_cpu_cache) for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2") else: print('Merge Lora') # Clear lora when switch from lora_cache=True to lora_cache=False. @@ -933,7 +940,7 @@ class Wan2_2FunV2VSampler: gc.collect() for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) # Clear lora when switch from lora_cache=True to lora_cache=False. if pipeline.transformer_2 is not None: @@ -944,7 +951,7 @@ class Wan2_2FunV2VSampler: gc.collect() for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2") if model_type == "Inpaint": sample = pipeline( @@ -987,8 +994,8 @@ class Wan2_2FunV2VSampler: if not funmodels.get("lora_cache", False): print('Unmerge Lora') for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) if pipeline.transformer_2 is not None: for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2") return (videos,) \ No newline at end of file diff --git a/comfyui/wan2_2_vace_fun/nodes.py b/comfyui/wan2_2_vace_fun/nodes.py index 72a1538..6e7af93 100644 --- a/comfyui/wan2_2_vace_fun/nodes.py +++ b/comfyui/wan2_2_vace_fun/nodes.py @@ -26,12 +26,12 @@ from ...videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, from ...videox_fun.models.cache_utils import get_teacache_coefficients from ...videox_fun.pipeline import Wan2_2VaceFunPipeline from ...videox_fun.utils.fp8_optimization import ( - convert_model_weight_to_float8, convert_weight_dtype_wrapper, + convert_model_weight_to_float8, convert_weight_dtype_wrapper, undo_convert_weight_dtype_wrapper, replace_parameters_by_name) from ...videox_fun.utils.lora_utils import merge_lora, unmerge_lora from ...videox_fun.utils.utils import (filter_kwargs, get_image_latent, get_image_to_video_latent, - get_video_to_video_latent) + get_video_to_video_latent, get_autocast_dtype) from ..comfyui_utils import (script_directory, search_model_in_possible_folders, to_pil) from ..wan2_1.nodes import get_wan_scheduler @@ -68,6 +68,11 @@ class LoadVaceWanTransformer3DModel: offload_device = mm.unet_offload_device() weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16}[precision] + mm.unload_all_models() + mm.cleanup_models() + mm.soft_empty_cache() + transformer = None + model_path = folder_paths.get_full_path("diffusion_models", model_name) transformer_state_dict = load_torch_file(model_path, safe_load=True) @@ -167,7 +172,7 @@ class CombineWan2_2VaceFunPipeline: CATEGORY = "CogVideoXFUNWrapper" def loadmodel(self, model_name, GPU_memory_mode, transformer, vae, text_encoder, tokenizer, clip_encoder=None, transformer_2=None, model_type="Control"): - weight_dtype = transformer.dtype + weight_dtype = transformer.dtype if transformer.dtype not in [torch.float32, torch.float8_e4m3fn, torch.float8_e5m2] else get_autocast_dtype() device = mm.get_torch_device() offload_device = mm.unet_offload_device() @@ -181,6 +186,11 @@ class CombineWan2_2VaceFunPipeline: scheduler=None, ) + pipeline.remove_all_hooks() + undo_convert_weight_dtype_wrapper(transformer) + pipeline.to(device=offload_device) + transformer = transformer.to(weight_dtype) + if GPU_memory_mode == "sequential_cpu_offload": replace_parameters_by_name(transformer, ["modulation",], device=device) transformer.freqs = transformer.freqs.to(device=device) @@ -255,6 +265,10 @@ class LoadWan2_2VaceFunModel: offload_device = mm.unet_offload_device() weight_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] + mm.unload_all_models() + mm.cleanup_models() + mm.soft_empty_cache() + # Init processbar pbar = ProgressBar(5) @@ -331,6 +345,9 @@ class LoadWan2_2VaceFunModel: scheduler=scheduler, ) + pipeline.remove_all_hooks() + undo_convert_weight_dtype_wrapper(transformer) + if GPU_memory_mode == "sequential_cpu_offload": replace_parameters_by_name(transformer, ["modulation",], device=device) transformer.freqs = transformer.freqs.to(device=device) @@ -549,7 +566,7 @@ class Wan2_2VaceFunSampler: lora_path_before = copy.deepcopy(lora_path_now) pipeline.transformer.load_state_dict(transformer_cpu_cache) for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) if pipeline.transformer_2 is not None: # Save the original weights to cpu @@ -565,7 +582,7 @@ class Wan2_2VaceFunSampler: lora_high_path_before = copy.deepcopy(lora_high_path_now) pipeline.transformer_2.load_state_dict(transformer_cpu_cache) for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2") else: print('Merge Lora') # Clear lora when switch from lora_cache=True to lora_cache=False. @@ -576,7 +593,7 @@ class Wan2_2VaceFunSampler: gc.collect() for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) # Clear lora when switch from lora_cache=True to lora_cache=False. if pipeline.transformer_2 is not None: @@ -587,7 +604,7 @@ class Wan2_2VaceFunSampler: gc.collect() for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + pipeline = merge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2") sample = pipeline( prompt, @@ -611,8 +628,8 @@ class Wan2_2VaceFunSampler: if not funmodels.get("lora_cache", False): print('Unmerge Lora') for _lora_path, _lora_weight in zip(funmodels.get("loras", []), funmodels.get("strength_model", [])): - pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype) + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype) if pipeline.transformer_2 is not None: for _lora_path, _lora_weight in zip(funmodels.get("loras_high", []), funmodels.get("strength_model", [])): - pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device="cuda", dtype=weight_dtype, sub_transformer_name="transformer_2") + pipeline = unmerge_lora(pipeline, _lora_path, _lora_weight, device=device, dtype=weight_dtype, sub_transformer_name="transformer_2") return (videos,) \ No newline at end of file diff --git a/videox_fun/utils/__init__.py b/videox_fun/utils/__init__.py index cb5b7ec..009df37 100755 --- a/videox_fun/utils/__init__.py +++ b/videox_fun/utils/__init__.py @@ -7,7 +7,7 @@ from .fp8_optimization import (autocast_model_forward, convert_weight_dtype_wrapper, replace_parameters_by_name) from .lora_utils import merge_lora, unmerge_lora -from .utils import (filter_kwargs, get_image_latent, get_image_to_video_latent, +from .utils import (filter_kwargs, get_image_latent, get_image_to_video_latent, get_autocast_dtype, get_video_to_video_latent, save_videos_grid) from .cfg_optimization import cfg_skip from .discrete_sampler import DiscreteSampling diff --git a/videox_fun/utils/fp8_optimization.py b/videox_fun/utils/fp8_optimization.py index 0c55cbf..b8b0038 100755 --- a/videox_fun/utils/fp8_optimization.py +++ b/videox_fun/utils/fp8_optimization.py @@ -55,4 +55,10 @@ def convert_weight_dtype_wrapper(module, origin_dtype): module, "forward", lambda *inputs, m=module, **kwargs: autocast_model_forward(m, origin_dtype, *inputs, **kwargs) - ) \ No newline at end of file + ) + +def undo_convert_weight_dtype_wrapper(module): + for name, module in module.named_modules(): + if hasattr(module, "original_forward") and module.weight is not None: + setattr(module, "forward", module.original_forward) + delattr(module, "original_forward") \ No newline at end of file diff --git a/videox_fun/utils/utils.py b/videox_fun/utils/utils.py index 8e00716..98453ad 100755 --- a/videox_fun/utils/utils.py +++ b/videox_fun/utils/utils.py @@ -419,3 +419,29 @@ def _write_to_excel(model_name, time_sum): df.iloc[row_idx, col_idx] = time_sum df.to_excel(file_path, index=False, header=False, sheet_name="Sheet1") + +def get_autocast_dtype(): + try: + if not torch.cuda.is_available(): + print("CUDA not available, using float16 by default.") + return torch.float16 + + device = torch.cuda.current_device() + prop = torch.cuda.get_device_properties(device) + + print(f"GPU: {prop.name}, Compute Capability: {prop.major}.{prop.minor}") + + if prop.major >= 8: + if torch.cuda.is_bf16_supported(): + print("Using bfloat16.") + return torch.bfloat16 + else: + print("Compute capability >= 8.0 but bfloat16 not supported, falling back to float16.") + return torch.float16 + else: + print("GPU does not support bfloat16 natively, using float16.") + return torch.float16 + + except Exception as e: + print(f"Error detecting GPU capability: {e}, falling back to float16.") + return torch.float16 \ No newline at end of file