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:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user