etomoscow/mff_lora / code /src /mfflora /train /lora_runner.py
etomoscow's picture
download
raw
2.87 kB
"""LoRA training runner with custom init injection.
We use ``peft.LoraConfig(init_lora_weights=False)`` and then *manually* set
``lora_A.default.weight`` and ``lora_B.default.weight`` from a user-supplied
init function. After injection we read back the actual weights to verify peft
did not silently rewrite them.
PEFT convention: for a target Linear of shape ``[n, m]``,
- ``lora_A.default.weight`` has shape ``[r, m]`` (i.e. our ``V``).
- ``lora_B.default.weight`` has shape ``[n, r]`` (i.e. our ``U``).
So ``ΔW (peft) = lora_B @ lora_A = U V`` matches our convention.
"""
from __future__ import annotations
import torch
from torch import nn
def inject_lora_init(
peft_model: nn.Module,
inits: dict[str, dict[str, torch.Tensor | None]],
*,
verify: bool = True,
) -> None:
"""Apply per-target ``{U, V, residual}`` dicts to a peft-wrapped model.
Builds the module dict once to avoid O(n²) traversals for large models.
"""
# Pre-build module map once — avoids repeated O(n_modules) traversal
# when injecting 96+ target layers in Llama-3.1-8B.
all_modules = dict(peft_model.named_modules())
for target_dotted, payload in inits.items():
U = payload["U"]
V = payload["V"]
residual = payload.get("residual")
if U is None or V is None:
raise ValueError(f"Init for {target_dotted} missing U or V")
if target_dotted not in all_modules:
raise KeyError(f"Module {target_dotted!r} not found in peft model")
target = all_modules[target_dotted]
if not hasattr(target, "lora_A") or not hasattr(target, "lora_B"):
raise AttributeError(f"{target_dotted!r} is not a peft LoRA wrapper")
lora_A = target.lora_A["default"]
lora_B = target.lora_B["default"]
with torch.no_grad():
lora_A.weight.copy_(V.to(lora_A.weight.dtype).to(lora_A.weight.device))
lora_B.weight.copy_(U.to(lora_B.weight.dtype).to(lora_B.weight.device))
if residual is not None:
base = target.base_layer.weight
base.copy_(residual.to(base.dtype).to(base.device))
if verify:
for target_dotted, payload in inits.items():
target = all_modules[target_dotted]
lora_A_w = target.lora_A["default"].weight.detach().cpu().to(torch.float32)
lora_B_w = target.lora_B["default"].weight.detach().cpu().to(torch.float32)
V_cpu = payload["V"].cpu().to(torch.float32)
U_cpu = payload["U"].cpu().to(torch.float32)
if not torch.allclose(lora_A_w, V_cpu, atol=1e-4):
raise RuntimeError(f"Verify failed: lora_A mismatch on {target_dotted}")
if not torch.allclose(lora_B_w, U_cpu, atol=1e-4):
raise RuntimeError(f"Verify failed: lora_B mismatch on {target_dotted}")

Xet Storage Details

Size:
2.87 kB
·
Xet hash:
a18e2874b831660f09617a6dfebe4197e51e2821a8bbf98345b3f6e9fbdaaae8

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.