Make non_blocking offloading optional for low RAM usecases

This commit is contained in:
kijai
2025-03-16 01:48:45 +02:00
parent f28d3f4225
commit 81a08d0e59
3 changed files with 15 additions and 8 deletions
-1
View File
@@ -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
+8 -3
View File
@@ -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"):
+7 -4
View File
@@ -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(