From 81a08d0e59922a54797e9a400c73d26b5be7db6d Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 16 Mar 2025 01:48:45 +0200 Subject: [PATCH] Make non_blocking offloading optional for low RAM usecases --- latent_preview.py | 1 - nodes.py | 11 ++++++++--- wanvideo/modules/model.py | 11 +++++++---- 3 files changed, 15 insertions(+), 8 deletions(-) diff --git a/latent_preview.py b/latent_preview.py index 81225f6..fb94e54 100644 --- a/latent_preview.py +++ b/latent_preview.py @@ -73,7 +73,6 @@ def get_previewer(device, latent_format): method = args.preview_method if method != LatentPreviewMethod.NoPreviews: # TODO previewer methods - taesd_decoder_path = None if method == LatentPreviewMethod.Auto: method = LatentPreviewMethod.Latent2RGB diff --git a/nodes.py b/nodes.py index 326f653..2388a84 100644 --- a/nodes.py +++ b/nodes.py @@ -46,6 +46,9 @@ class WanVideoBlockSwap: "offload_img_emb": ("BOOLEAN", {"default": False, "tooltip": "Offload img_emb to offload_device"}), "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"}), + }, } RETURN_TYPES = ("BLOCKSWAPARGS",) RETURN_NAMES = ("block_swap_args",) @@ -1454,19 +1457,21 @@ class WanVideoSampler: #blockswap init if model["block_swap_args"] is not None: + transformer.use_non_blocking = model["block_swap_args"].get("use_non_blocking", True) for name, param in transformer.named_parameters(): if "block" not in name: - param.data = param.data.to(device) + param.data = param.data.to(device, non_blocking=transformer.use_non_blocking) elif model["block_swap_args"]["offload_txt_emb"] and "txt_emb" in name: - param.data = param.data.to(offload_device) + param.data = param.data.to(offload_device, non_blocking=transformer.use_non_blocking) elif model["block_swap_args"]["offload_img_emb"] and "img_emb" in name: - param.data = param.data.to(offload_device) + param.data = param.data.to(offload_device, non_blocking=transformer.use_non_blocking) transformer.block_swap( model["block_swap_args"]["blocks_to_swap"] - 1 , model["block_swap_args"]["offload_txt_emb"], model["block_swap_args"]["offload_img_emb"], ) + elif model["auto_cpu_offload"]: for module in transformer.modules(): if hasattr(module, "offload"): diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index bb82ddd..c191612 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -601,6 +601,8 @@ class WanModel(ModelMixin, ConfigMixin): self.slg_start_percent = 0.0 self.slg_end_percent = 1.0 + self.use_non_blocking = True + # embeddings self.patch_embedding = nn.Conv3d( in_dim, dim, kernel_size=patch_size, stride=patch_size) @@ -663,7 +665,7 @@ class WanModel(ModelMixin, ConfigMixin): block.to(self.main_device) total_main_memory += block_memory else: - block.to(self.offload_device) + block.to(self.offload_device, non_blocking=self.use_non_blocking) total_offload_memory += block_memory mm.soft_empty_cache() @@ -675,6 +677,7 @@ class WanModel(ModelMixin, ConfigMixin): log.info(f"Transformer blocks on {self.offload_device}: {total_offload_memory:.2f}MB") log.info(f"Transformer blocks on {self.main_device}: {total_main_memory:.2f}MB") log.info(f"Total memory used by transformer blocks: {(total_offload_memory + total_main_memory):.2f}MB") + log.info(f"Non-blocking memory transfer: {self.use_non_blocking}") log.info("----------------------") def forward( @@ -783,7 +786,7 @@ class WanModel(ModelMixin, ConfigMixin): for u in context ])) if self.offload_txt_emb: - self.text_embedding.to(self.offload_device, non_blocking=True) + self.text_embedding.to(self.offload_device, non_blocking=self.use_non_blocking) if clip_fea is not None: if self.offload_img_emb: @@ -791,7 +794,7 @@ class WanModel(ModelMixin, ConfigMixin): context_clip = self.img_emb(clip_fea) # bs x 257 x dim context = torch.concat([context_clip, context], dim=1) if self.offload_img_emb: - self.img_emb.to(self.offload_device, non_blocking=True) + self.img_emb.to(self.offload_device, non_blocking=self.use_non_blocking) should_calc = True accumulated_rel_l1_distance = torch.tensor(0.0, dtype=torch.float32, device=device) @@ -855,7 +858,7 @@ class WanModel(ModelMixin, ConfigMixin): block.to(self.main_device) x = block(x, **kwargs) if b <= self.blocks_to_swap and self.blocks_to_swap >= 0: - block.to(self.offload_device, non_blocking=True) + block.to(self.offload_device, non_blocking=self.use_non_blocking) if self.enable_teacache and pred_id is not None: self.teacache_state.update(