offloading fixes

This commit is contained in:
kijai
2025-08-19 19:11:33 +03:00
parent 39908c9aea
commit 4aac86a828
2 changed files with 66 additions and 69 deletions
+65 -65
View File
@@ -2218,38 +2218,38 @@ class WanVideoSampler:
gc.collect()
#blockswap init
if block_swap_args is not None and not transformer.patched_linear:
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False)
for name, param in transformer.named_parameters():
if "block" not in name:
param.data = param.data.to(device)
if "control_adapter" in name:
param.data = param.data.to(device)
elif block_swap_args["offload_txt_emb"] and "txt_emb" in name:
param.data = param.data.to(offload_device)
elif block_swap_args["offload_img_emb"] and "img_emb" in name:
param.data = param.data.to(offload_device)
if not transformer.patched_linear:
if block_swap_args is not None:
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False)
for name, param in transformer.named_parameters():
if "block" not in name:
param.data = param.data.to(device)
if "control_adapter" in name:
param.data = param.data.to(device)
elif block_swap_args["offload_txt_emb"] and "txt_emb" in name:
param.data = param.data.to(offload_device)
elif block_swap_args["offload_img_emb"] and "img_emb" in name:
param.data = param.data.to(offload_device)
transformer.block_swap(
block_swap_args["blocks_to_swap"] - 1 ,
block_swap_args["offload_txt_emb"],
block_swap_args["offload_img_emb"],
vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None),
prefetch_blocks = block_swap_args.get("prefetch_blocks", 0),
block_swap_debug = block_swap_args.get("block_swap_debug", False),
)
elif model["auto_cpu_offload"]:
for module in transformer.modules():
if hasattr(module, "offload"):
module.offload()
if hasattr(module, "onload"):
module.onload()
for block in transformer.blocks:
block.modulation = torch.nn.Parameter(block.modulation.to(device))
transformer.head.modulation = torch.nn.Parameter(transformer.head.modulation.to(device))
#elif model["manual_offloading"]:
# transformer.to(device)
transformer.block_swap(
block_swap_args["blocks_to_swap"] - 1 ,
block_swap_args["offload_txt_emb"],
block_swap_args["offload_img_emb"],
vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None),
prefetch_blocks = block_swap_args.get("prefetch_blocks", 0),
block_swap_debug = block_swap_args.get("block_swap_debug", False),
)
elif model["auto_cpu_offload"]:
for module in transformer.modules():
if hasattr(module, "offload"):
module.offload()
if hasattr(module, "onload"):
module.onload()
for block in transformer.blocks:
block.modulation = torch.nn.Parameter(block.modulation.to(device))
transformer.head.modulation = torch.nn.Parameter(transformer.head.modulation.to(device))
else:
transformer.to(device)
# Initialize Cache if enabled
transformer.enable_teacache = transformer.enable_magcache = transformer.enable_easycache = False
@@ -2657,7 +2657,7 @@ class WanVideoSampler:
except Exception as e:
log.error(f"Error during model prediction: {e}")
if force_offload:
if model["manual_offloading"]:
if not model["auto_cpu_offload"]:
offload_transformer(transformer)
raise e
@@ -3256,36 +3256,36 @@ class WanVideoSampler:
if offload:
#blockswap init
if transformer_options is not None:
block_swap_args = transformer_options.get("block_swap_args", None)
if not transformer.patched_linear:
if block_swap_args is not None:
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False)
for name, param in transformer.named_parameters():
if "block" not in name:
param.data = param.data.to(device)
if "control_adapter" in name:
param.data = param.data.to(device)
elif block_swap_args["offload_txt_emb"] and "txt_emb" in name:
param.data = param.data.to(offload_device)
elif block_swap_args["offload_img_emb"] and "img_emb" in name:
param.data = param.data.to(offload_device)
if block_swap_args is not None:
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False)
for name, param in transformer.named_parameters():
if "block" not in name:
param.data = param.data.to(device)
if "control_adapter" in name:
param.data = param.data.to(device)
elif block_swap_args["offload_txt_emb"] and "txt_emb" in name:
param.data = param.data.to(offload_device)
elif block_swap_args["offload_img_emb"] and "img_emb" in name:
param.data = param.data.to(offload_device)
transformer.block_swap(
block_swap_args["blocks_to_swap"] - 1 ,
block_swap_args["offload_txt_emb"],
block_swap_args["offload_img_emb"],
vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None),
)
elif model["auto_cpu_offload"]:
for module in transformer.modules():
if hasattr(module, "offload"):
module.offload()
if hasattr(module, "onload"):
module.onload()
elif model["manual_offloading"]:
transformer.to(device)
transformer.block_swap(
block_swap_args["blocks_to_swap"] - 1 ,
block_swap_args["offload_txt_emb"],
block_swap_args["offload_img_emb"],
vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None),
)
elif model["auto_cpu_offload"]:
for module in transformer.modules():
if hasattr(module, "offload"):
module.offload()
if hasattr(module, "onload"):
module.onload()
for block in transformer.blocks:
block.modulation = torch.nn.Parameter(block.modulation.to(device))
transformer.head.modulation = torch.nn.Parameter(transformer.head.modulation.to(device))
else:
transformer.to(device)
comfy_pbar = ProgressBar(len(timesteps)-1)
for i in tqdm(range(len(timesteps)-1)):
@@ -3337,7 +3337,7 @@ class WanVideoSampler:
comfy_pbar.update(1)
if offload:
transformer.to(offload_device)
offload_transformer(transformer)
vae.to(device)
videos = vae.decode(x0.unsqueeze(0).to(vae.dtype), device=device, tiled=tiled_vae)
vae.to(offload_device)
@@ -3410,7 +3410,7 @@ class WanVideoSampler:
del noise, latent
if force_offload:
if model["manual_offloading"]:
if not model["auto_cpu_offload"]:
transformer.to(offload_device)
mm.soft_empty_cache()
gc.collect()
@@ -3508,7 +3508,7 @@ class WanVideoSampler:
except Exception as e:
log.error(f"Error during sampling: {e}")
if force_offload:
if model["manual_offloading"]:
if not model["auto_cpu_offload"]:
offload_transformer(transformer)
raise e
@@ -3516,7 +3516,7 @@ class WanVideoSampler:
cache_report(transformer, cache_args)
if force_offload:
if model["manual_offloading"]:
if not model["auto_cpu_offload"]:
offload_transformer(transformer)
try:
+1 -4
View File
@@ -948,7 +948,7 @@ class WanVideoModelLoader:
mm.unload_all_models()
mm.cleanup_models()
mm.soft_empty_cache()
manual_offloading = True
if "sage" in attention_mode:
try:
from sageattention import sageattn
@@ -964,8 +964,6 @@ class WanVideoModelLoader:
if merge_loras is True:
raise ValueError("GGUF models do not support LoRA merging, please disable merge_loras in the LoRA select node.")
manual_offloading = True
transformer_load_device = device if load_device == "main_device" else offload_device
base_dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp16_fast": torch.float16, "fp32": torch.float32}[base_precision]
@@ -1366,7 +1364,6 @@ class WanVideoModelLoader:
patcher.model["weight_dtype"] = weight_dtype
patcher.model["base_path"] = model_path
patcher.model["model_name"] = model
patcher.model["manual_offloading"] = manual_offloading
patcher.model["quantization"] = quantization
patcher.model["auto_cpu_offload"] = True if vram_management_args is not None else False
patcher.model["control_lora"] = control_lora