Update unload model comfyui for saving VRAM. (#340)
This commit is contained in:
@@ -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":
|
||||
|
||||
@@ -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"):
|
||||
|
||||
@@ -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,)
|
||||
|
||||
+26
-5
@@ -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)
|
||||
|
||||
+20
-13
@@ -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,)
|
||||
|
||||
|
||||
+32
-15
@@ -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,)
|
||||
+26
-19
@@ -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,)
|
||||
@@ -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,)
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
)
|
||||
)
|
||||
|
||||
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")
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user