Update nodes.py

This commit is contained in:
kijai
2025-09-19 18:40:49 +03:00
parent ca48988742
commit de4fcd6fb3
+3 -3
View File
@@ -4361,11 +4361,10 @@ class WanVideoSampler:
vae.model.clear_cache()
vae.to(offload_device)
mm.soft_empty_cache()
gc.collect()
temporal_ref_latents = torch.cat([msk, temporal_ref_latents], dim=0) # 4+C T H W
image_cond_in = torch.cat([ref_latent.to(device), temporal_ref_latents], dim=1) # 4+C T+trefs H W
del temporal_ref_latents, msk, bg_image_slice
noise = torch.randn(16, latent_window_size + 1, lat_h, lat_w, dtype=torch.float32, device=torch.device("cpu"), generator=seed_g).to(device)
seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * noise.shape[1])
@@ -4379,6 +4378,7 @@ class WanVideoSampler:
pad_len = latent_window_size - pose_input_slice.shape[2]
pad = torch.zeros(pose_input_slice.shape[0], pose_input_slice.shape[1], pad_len, pose_input_slice.shape[3], pose_input_slice.shape[4], device=pose_input_slice.device, dtype=pose_input_slice.dtype)
pose_input_slice = torch.cat([pose_input_slice, pad], dim=2)
del pad
pose_input_slice = pose_input_slice.to(device, dtype)
if samples is not None:
@@ -4496,7 +4496,6 @@ class WanVideoSampler:
videos = vae.decode(latent[:, 1:].unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu()
del latent
vae.model.clear_cache()
vae.to(offload_device)
sampling_pbar.close()
@@ -4540,6 +4539,7 @@ class WanVideoSampler:
gen_video_samples = torch.zeros(3, 1, 64, 64) # dummy output
if force_offload:
vae.to(offload_device)
if not model["auto_cpu_offload"]:
offload_transformer(transformer)
try: