Update apply lora low_mem_load

Avoid creating 2 copies of WAN in VRAM when applying a LoRA by setting inplace_update=True on model.patch_weight_to_device.
This commit is contained in:
Alexander Measure
2025-07-14 14:10:56 -04:00
committed by GitHub
parent 17d48e3e45
commit 88f50a8fd9
+5 -2
View File
@@ -73,7 +73,10 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d
set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[key])
except:
continue
model.patch_weight_to_device("{}.{}".format(name, param), device_to=device_to)
if low_mem_load:
model.patch_weight_to_device("{}.{}".format(name, param), device_to=device_to, inplace_update=True)
else:
model.patch_weight_to_device("{}.{}".format(name, param), device_to=device_to)
if low_mem_load:
try:
set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=model.model.diffusion_model.state_dict()[key])
@@ -232,4 +235,4 @@ def fourier_filter(x, scale_low=1.0, scale_high=1.5, freq_cutoff=20):
# 5) Restore original dtype
x_filtered = x_filtered.to(dtype)
return x_filtered
return x_filtered