From 88f50a8fd992cadf4ed77cd085a56a39a4912949 Mon Sep 17 00:00:00 2001 From: Alexander Measure Date: Mon, 14 Jul 2025 14:10:56 -0400 Subject: [PATCH] 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. --- utils.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/utils.py b/utils.py index ea4bc50..b7f76e6 100644 --- a/utils.py +++ b/utils.py @@ -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 \ No newline at end of file + return x_filtered