Default non_blocking to False
too many RAM issues
This commit is contained in:
@@ -1816,16 +1816,16 @@ class WanVideoSampler:
|
||||
block_swap_args = transformer_options.get("block_swap_args", None)
|
||||
|
||||
if block_swap_args is not None:
|
||||
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", True)
|
||||
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, non_blocking=transformer.use_non_blocking)
|
||||
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, non_blocking=transformer.use_non_blocking)
|
||||
param.data = param.data.to(offload_device)
|
||||
|
||||
transformer.block_swap(
|
||||
block_swap_args["blocks_to_swap"] - 1 ,
|
||||
@@ -2695,16 +2695,16 @@ class WanVideoSampler:
|
||||
block_swap_args = transformer_options.get("block_swap_args", None)
|
||||
|
||||
if block_swap_args is not None:
|
||||
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", True)
|
||||
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, non_blocking=transformer.use_non_blocking)
|
||||
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, non_blocking=transformer.use_non_blocking)
|
||||
param.data = param.data.to(offload_device)
|
||||
|
||||
transformer.block_swap(
|
||||
block_swap_args["blocks_to_swap"] - 1 ,
|
||||
@@ -2871,18 +2871,18 @@ class WanVideoSampler:
|
||||
|
||||
if callback is not None:
|
||||
if recammaster is not None:
|
||||
callback_latent = (latent_model_input[:, :orig_noise_len].to(device) - noise_pred[:, :orig_noise_len].to(device) * t.to(device) / 1000).detach().permute(1,0,2,3)
|
||||
callback_latent = (latent_model_input[:, :orig_noise_len].to(device) - noise_pred[:, :orig_noise_len].to(device) * t.to(device) / 1000).detach()
|
||||
elif phantom_latents is not None:
|
||||
callback_latent = (latent_model_input[:,:-phantom_latents.shape[1]].to(device) - noise_pred[:,:-phantom_latents.shape[1]].to(device) * t.to(device) / 1000).detach().permute(1,0,2,3)
|
||||
callback_latent = (latent_model_input[:,:-phantom_latents.shape[1]].to(device) - noise_pred[:,:-phantom_latents.shape[1]].to(device) * t.to(device) / 1000).detach()
|
||||
else:
|
||||
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device) / 1000).detach().permute(1,0,2,3)
|
||||
callback(idx, callback_latent, None, len(timesteps))
|
||||
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device) / 1000).detach()
|
||||
callback(idx, callback_latent.permute(1,0,2,3), None, len(timesteps))
|
||||
else:
|
||||
pbar.update(1)
|
||||
else:
|
||||
if callback is not None:
|
||||
callback_latent = (zt_tgt.to(device) - vt_tgt.to(device) * t.to(device) / 1000).detach().permute(1,0,2,3)
|
||||
callback(idx, callback_latent, None, len(timesteps))
|
||||
callback_latent = (zt_tgt.to(device) - vt_tgt.to(device) * t.to(device) / 1000).detach()
|
||||
callback(idx, callback_latent.permute(1,0,2,3), None, len(timesteps))
|
||||
else:
|
||||
pbar.update(1)
|
||||
|
||||
@@ -2905,7 +2905,6 @@ class WanVideoSampler:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
print("samples out stats: mean", latent.mean().item(), "std", latent.std().item(), "min", latent.min().item(), "max", latent.max().item())
|
||||
return ({
|
||||
"samples": latent.unsqueeze(0).cpu(),
|
||||
"looped": is_looped,
|
||||
@@ -2914,7 +2913,7 @@ class WanVideoSampler:
|
||||
"drop_last": drop_last,
|
||||
"generator_state": seed_g.get_state(),
|
||||
},{
|
||||
"samples": (latent_model_input.cpu() - noise_pred.cpu() * t.cpu() / 1000).unsqueeze(0),
|
||||
"samples": callback_latent.unsqueeze(0).cpu(),
|
||||
})
|
||||
|
||||
#region VideoDecode
|
||||
|
||||
@@ -273,7 +273,7 @@ class WanVideoBlockSwap:
|
||||
"offload_txt_emb": ("BOOLEAN", {"default": False, "tooltip": "Offload time_emb to offload_device"}),
|
||||
},
|
||||
"optional": {
|
||||
"use_non_blocking": ("BOOLEAN", {"default": True, "tooltip": "Use non-blocking memory transfer for offloading, reserves more RAM but is faster"}),
|
||||
"use_non_blocking": ("BOOLEAN", {"default": False, "tooltip": "Use non-blocking memory transfer for offloading, reserves more RAM but is faster"}),
|
||||
"vace_blocks_to_swap": ("INT", {"default": 0, "min": 0, "max": 15, "step": 1, "tooltip": "Number of VACE blocks to swap, the VACE model has 15 blocks"}),
|
||||
},
|
||||
}
|
||||
|
||||
+3
-3
@@ -365,14 +365,14 @@ class WanVideoDiffusionForcingSampler:
|
||||
block_swap_args = transformer_options.get("block_swap_args", None)
|
||||
|
||||
if block_swap_args is not None:
|
||||
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", True)
|
||||
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)
|
||||
elif block_swap_args["offload_txt_emb"] and "txt_emb" in name:
|
||||
param.data = param.data.to(offload_device, non_blocking=transformer.use_non_blocking)
|
||||
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, non_blocking=transformer.use_non_blocking)
|
||||
param.data = param.data.to(offload_device)
|
||||
|
||||
transformer.block_swap(
|
||||
block_swap_args["blocks_to_swap"] - 1 ,
|
||||
|
||||
@@ -1,10 +1,7 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from einops import repeat, rearrange
|
||||
from ...enhance_a_video.enhance import get_feta_scores
|
||||
|
||||
@@ -1096,7 +1093,7 @@ class WanModel(torch.nn.Module):
|
||||
self.slg_start_percent = 0.0
|
||||
self.slg_end_percent = 1.0
|
||||
|
||||
self.use_non_blocking = True
|
||||
self.use_non_blocking = False
|
||||
|
||||
self.video_attention_split_steps = []
|
||||
|
||||
|
||||
@@ -109,5 +109,4 @@ def get_scheduler(scheduler, steps, shift, device, transformer_dim, flowedit_arg
|
||||
sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength, sigmas=sigmas[:-1].tolist() if sigmas is not None else None)
|
||||
if timesteps is None:
|
||||
timesteps = sample_scheduler.timesteps
|
||||
log.info(f"timesteps: {timesteps}")
|
||||
return sample_scheduler, timesteps
|
||||
Reference in New Issue
Block a user