| """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.