multimodalart's picture
multimodalart HF Staff
Upload folder using huggingface_hub
267f80e verified
Raw
History Blame Contribute Delete
1.93 kB
import torch
import torch.nn as nn
from safetensors.torch import safe_open
def build_lora_names(key, lora_down_key, lora_up_key, is_native_weight):
base = "diffusion_model." if is_native_weight else ""
lora_down = base + key.replace(".weight", lora_down_key)
lora_up = base + key.replace(".weight", lora_up_key)
lora_alpha = base + key.replace(".weight", ".alpha")
return lora_down, lora_up, lora_alpha
def load_and_merge_lora_weight(
model: nn.Module,
lora_state_dict: dict,
lora_down_key: str = ".lora_down.weight",
lora_up_key: str = ".lora_up.weight",
):
is_native_weight = any("diffusion_model." in key for key in lora_state_dict)
for key, value in model.named_parameters():
lora_down_name, lora_up_name, lora_alpha_name = build_lora_names(
key, lora_down_key, lora_up_key, is_native_weight
)
if lora_down_name in lora_state_dict:
lora_down = lora_state_dict[lora_down_name]
lora_up = lora_state_dict[lora_up_name]
lora_alpha = float(lora_state_dict[lora_alpha_name])
rank = lora_down.shape[0]
scaling_factor = lora_alpha / rank
assert lora_up.dtype == torch.float32
assert lora_down.dtype == torch.float32
delta_W = scaling_factor * torch.matmul(lora_up, lora_down).to(value.device)
value.data = (value.data + delta_W).type_as(value.data)
return model
def load_and_merge_lora_weight_from_safetensors(
model: nn.Module,
lora_weight_path: str,
lora_down_key: str = ".lora_down.weight",
lora_up_key: str = ".lora_up.weight",
):
lora_state_dict = {}
with safe_open(lora_weight_path, framework="pt", device="cpu") as f:
for key in f.keys():
lora_state_dict[key] = f.get_tensor(key)
model = load_and_merge_lora_weight(model, lora_state_dict, lora_down_key, lora_up_key)
return model