Update unload model comfyui for saving VRAM. (#340)

This commit is contained in:
Bubbliiiing
2025-09-28 14:45:22 +08:00
committed by GitHub
parent 7851e14319
commit 6183d41df8
11 changed files with 224 additions and 73 deletions
+9 -1
View File
@@ -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":
+25
View File
@@ -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"):
+26 -9
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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 -9
View File
@@ -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,)
+1 -1
View File
@@ -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
+7 -1
View File
@@ -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")
+26
View File
@@ -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