Skip to content

LoRA with PiSSA init and merge

ModuleL6.6 · build · Python · Pass 5 · 3 to 4 h, plus your graded tests (rung R5)
You buildpython/tinyllm/obj/lora.py: LoRALinear, inject_lora, merge_lora, lora_state_dict, load_lora_state_dict, trainable_fraction; and your own oracle tests in python/tests/l6-6-lora/
Contractcourse/contracts/py/tinyllm/obj/lora.pyi
Testscourse/tests/L6.6/test_lora.py (what they check: section 4); the PiSSA oracle is LAPACK’s SVD; your tests are graded by mutation, threshold 0.80 with every required pitfall fault killed
NeedsL0.1 Tensor · L0.2 ops · L0.4 Linear, Dropout, Module · M03.5 SVD and low_rank (or --ref-deps)
Used byL6.5 lora_classifier (the finetune classify --lora r=8 step of MS-L6) · later: L12.1 SFT, sq.multi-lora
MilestoneMS-L6 (trainable_frac < 0.05 for the LoRA fine-tune)
Optional depthHu et al., “LoRA: Low-Rank Adaptation of Large Language Models” (2021), sections 4 and 7; Meng, Wang, and Zhang, “PiSSA: Principal Singular Values and Singular Vectors Adaptation” (2024), sections 3 and 4; Aghajanyan et al., “Intrinsic Dimensionality Explains the Effectiveness of Language Model Fine-Tuning” (2020)
  • A LoRA layer computes Wx+b+s BAxW x + b + s\,B A x with the base WW frozen; at default init B=0B = 0, so the adapted model equals the base bit for bit (test_default_init_equals_base_bitwise).
  • The scale is s=α/rs = \alpha / r, not α\alpha: doubling the rank at a fixed α\alpha halves each direction’s step (test_scaling_is_alpha_over_r).
  • Only AA and BB get gradients; the frozen weight stays in the checkpoint but out of the optimizer (test_only_adapter_parameters_get_gradients, test_inject_freezes_everything_else).
  • PiSSA starts the adapter at the top-rr singular part of WW and freezes the residual, so step 0 still computes the base function (test_pissa_reproduces_the_base, test_pissa_adapter_is_the_top_singular_part).
  • Merging folds sBAs B A into WW and puts the plain Linear back: same outputs, same keys, no serving cost (test_merge_equals_unmerged, test_merge_restores_the_base_keys).
Terminal window
ol start L6.6 # stubs lora.py; prints your test path and rung (R5)
ol tests L6.6 # the course tests
# write your oracle tests in python/tests/l6-6-lora/ (section 4 lists what to cover), then:
ol check L6.6 # course tests and the mutation grade of your tests
ol mutate L6.6 # the full grade, cached by your test files' hash
ol diff L6.6 # after passing: your code against the reference

You can now train a BERT encoder (L6.2) and want it to classify sentences (L6.5). Fine-tuning every weight works on a laptop for a 50k-parameter model, but the habit does not scale: the AdamW state (M10.3) is two extra copies of every weight, and each task you fine-tune stores a full copy of the model. At the 135M parameters of SmolLM2, which your engine serves later, that is half a gigabyte of optimizer state and of checkpoint per task. The milestone of this part asks for the opposite: a classifier fine-tune where less than 5% of the parameters train (trainable_frac < 0.05). LoRA is the tool, and it is built from a piece you already own: the low-rank approximation of M03.5.

SymbolMeaningType / shape
WWthe base Linear’s weight, frozenfloat32[out, in]
bbthe base Linear’s bias, frozenfloat32[out]
xxone input rowfloat32[in]
rrthe adapter rank, 1≤r≤min⁡(in,out)1 \le r \le \min(\text{in}, \text{out})int
AAlora_A.weight, the down projectionfloat32[r, in]
BBlora_B.weight, the up projectionfloat32[out, r]
α\alphaalpha, the adapter’s scale knobfloat >0> 0
s=α/rs = \alpha / rscalingfloat
ΔW=sBA\Delta W = s B Athe update the adapter adds, rank at most rrfloat32[out, in]
U,S,V⊤U, S, V^\topthe SVD of WW (M03.5): W=U diag(S) V⊤W = U\,\mathrm{diag}(S)\,V^\topas svd returns
WrW_rthe best rank-rr approximation of WW (Eckart-Young)[out, in]

A fine-tune changes WW into W+ΔWW + \Delta W. Hu et al. observed that the useful ΔW\Delta W of a fine-tune has low “intrinsic” rank, so they write it as a product ΔW=sBA\Delta W = s B A with rr much smaller than either dimension. The layer computes

y=Wx+b+s B(A x).y = W x + b + s\,B (A\,x).

The order of the product matters for cost: AxA x is rr numbers, so the adapter costs r(in+out)r(\text{in} + \text{out}) multiply-adds per row instead of in⋅out\text{in} \cdot \text{out}. The parameters are AA and BB, r(in+out)r(\text{in} + \text{out}) numbers per layer. For a 768×768768 \times 768 projection at r=8r = 8 that is 12 288 trainable numbers instead of 589 824: about 2%.

The scale s=α/rs = \alpha / r is a convention with a reason. With AA initialized at a fixed size, the update BAxB A x sums rr terms, so its size grows with rr; dividing by rr keeps the effective learning rate of the adapter roughly independent of the rank, and you tune α\alpha once.

LoRA’s dropout (dropout) acts on the adapter’s input only: y=Wx+b+sBA dropout(x)y = W x + b + s B A\,\mathrm{dropout}(x). The frozen path must see exactly the base model’s input.

Default. AA is drawn like any Linear weight (L0.4) and B=0B = 0. Then ΔW=0\Delta W = 0 exactly: the adapted model starts as the pretrained one, and sBAxs B A x is a vector of exact zeros, so yy is bit for bit the base output. The first gradient step moves BB only (∂y/∂A\partial y / \partial A contains BB, which is zero), then both.

PiSSA. Meng et al. start the adapter at the most important part of WW instead. With the SVD W=U diag(S) V⊤W = U\,\mathrm{diag}(S)\,V^\top and singular values in decreasing order, Eckart-Young (M03.5) says the best rank-rr approximation is Wr=U:,:r diag(S:r) V:,:r⊤W_r = U_{:, :r}\,\mathrm{diag}(S_{:r})\,V_{:, :r}^\top. low_rank(W, r) returns it as balanced factors P=U:,:rS:rP = U_{:, :r}\sqrt{S_{:r}} and Q=S:r V:,:r⊤Q = \sqrt{S_{:r}}\,V_{:, :r}^\top, with PQ=WrP Q = W_r. PiSSA sets

B=P/s,A=Q/s,W←W−Wr,B = P / \sqrt{s}, \qquad A = Q / \sqrt{s}, \qquad W \leftarrow W - W_r,

so sBA=Wrs B A = W_r and the frozen weight keeps only the residual. At step 0 the layer computes (W−Wr)x+Wrx=Wx(W - W_r)x + W_r x = W x: the base function again (to float32 rounding, since it is now a sum of two products). What changed is which directions train: the adapter now owns the top singular directions of WW, the ones a fine-tune most often needs to move, and PiSSA converges faster than zero-initialized LoRA in the paper’s experiments. Dividing by s\sqrt{s} is what makes sBAs B A come out as WrW_r and not sWrs W_r.

Freezing is requires_grad = False on a registered parameter. L0.4 registers a Tensor when it is assigned with requires_grad set, and it stays registered afterwards, so a frozen weight is still in named_parameters and state_dict (the checkpoint keeps it) but produces no gradient: an op (L0.1’s from_op) records a vjp only when some input requires grad. inject_lora freezes the whole model first, then wraps the targeted Linears; after it, the adapters are the only trainable parameters, and an optimizer built from p for p in model.parameters() if p.requires_grad holds AA and BB only. trainable_fraction counts elements: trainable over all registered, frozen included.

A LoRALinear reuses the base Linear’s own Tensors under the base’s names (weight, bias) and adds lora_A.weight and lora_B.weight. That keeps every base checkpoint key valid on the adapted model.

To serve, fold the update into the weight: W′=W+sBAW' = W + s B A, computed in float64 and stored in float32, then put the original Linear object back in its parent. The merged model has the base’s keys in the base’s order, no adapter left, and the same outputs as the adapted model up to float32 rounding (the two compute W′xW' x and Wx+sB(Ax)W x + s B (A x), which round differently). A merged model never runs the adapter again: forgetting to remove it counts the update twice.

Adapters are shared without the base, under PEFT’s names in adapter_model.safetensors: base_model.model.<module name>.lora_A.weight and .lora_B.weight. lora_state_dict writes those keys; load_lora_state_dict reads them back and refuses a mismatched file before copying anything.

Forward and merge. W=(1234)W = \begin{pmatrix} 1 & 2 \\ 3 & 4 \end{pmatrix}, b=(0.5,−0.5)b = (0.5, -0.5), r=1r = 1, α=2\alpha = 2, so s=2s = 2. Take A=(1,−1)A = (1, -1) and B=(0.5,1)⊤B = (0.5, 1)^\top, and x=(2,1)x = (2, 1).

stepvalue
frozen path Wx+bW x + b(1⋅2+2⋅1,3⋅2+4⋅1)+b=(4,10)+(0.5,−0.5)=(4.5,9.5)(1 \cdot 2 + 2 \cdot 1, 3 \cdot 2 + 4 \cdot 1) + b = (4, 10) + (0.5, -0.5) = (4.5, 9.5)
AxA x2−1=12 - 1 = 1
B(Ax)B (A x)(0.5,1)(0.5, 1)
times s=2s = 2(1,2)(1, 2)
yy(5.5,11.5)(5.5, 11.5)
ΔW=sBA\Delta W = s B A2(0.5−0.51−1)=(1−12−2)2 \begin{pmatrix} 0.5 & -0.5 \\ 1 & -1 \end{pmatrix} = \begin{pmatrix} 1 & -1 \\ 2 & -2 \end{pmatrix}
merged W′W'(2152)\begin{pmatrix} 2 & 1 \\ 5 & 2 \end{pmatrix}, and W′x+b=(5,12)+b=(5.5,11.5)W' x + b = (5, 12) + b = (5.5, 11.5)

PiSSA split. W=diag(3,1)W = \mathrm{diag}(3, 1), r=1r = 1, α=4\alpha = 4, so s=4s = 4. The SVD is U=V=IU = V = I, S=(3,1)S = (3, 1), so W1=diag(3,0)W_1 = \mathrm{diag}(3, 0) and P=(3,0)⊤P = (\sqrt 3, 0)^\top, Q=(3,0)Q = (\sqrt 3, 0). Divide each by s=2\sqrt s = 2: B=(3/2,0)⊤B = (\sqrt 3 / 2, 0)^\top, A=(3/2,0)A = (\sqrt 3 / 2, 0), and sBA=4⋅34diag(1,0)=diag(3,0)s B A = 4 \cdot \tfrac{3}{4} \mathrm{diag}(1, 0) = \mathrm{diag}(3, 0). The frozen residual is diag(0,1)\mathrm{diag}(0, 1). (The pair (−B,−A)(-B, -A) is the same split.)

These are test_hand_example_forward_and_merge and test_hand_example_pissa_split.

class LoRALinear(Module):
def __init__(self, base: Linear, r: int, alpha: float, dropout: float = 0.0,
init: Literal["default", "pissa"] = "default", rng=None) -> None: ...
def delta_weight(self) -> NDArray: ... # s * B @ A, float32 [out, in]
def forward(self, x: Tensor) -> Tensor: ...
def inject_lora(model, target: Callable[[str, Module], bool], r: int, alpha: float,
dropout: float = 0.0, init="default", rng=None) -> list[str]: ...
def merge_lora(model) -> list[str]: ...
def lora_state_dict(model) -> dict[str, NDArray]: ... # PEFT key names
def load_lora_state_dict(model, sd) -> None: ...
def trainable_fraction(model) -> float: ...
TestKINDChecksWhy it matters downstream
test_hand_example_forward_and_mergeunitsection 3: y=(5.5,11.5)y = (5.5, 11.5) before and after merging, W′=[[2,1],[5,2]]W' = [[2, 1], [5, 2]]you and the test agree on the formula and on ss
test_hand_example_pissa_splitunitsection 3: BB, AA, and the residual diag(0,1)\mathrm{diag}(0, 1)the PiSSA split, sign aside
test_default_init_equals_base_bitwisedifferentialadapted outputs equal the base’s bit for bit; B=0B = 0, A≠0A \ne 0fine-tuning starts from the pretrained function
test_scaling_is_alpha_over_runits=α/rs = \alpha / r and ΔW=sBA\Delta W = s B A for four (r,α)(r, \alpha)one α\alpha works across ranks
test_only_adapter_parameters_get_gradientspropertyflags and gradients: frozen weight and bias get nonethe optimizer state is the adapter’s only
test_inject_freezes_everything_elsepropertyembeddings, norms, untargeted Linears frozen; a target matching nothing raisesLoRA trains adapters, nothing else
test_adapter_gradcheckgradcheck∂/∂A\partial / \partial A and ∂/∂B\partial / \partial B against the frozen central differencesthe gradients L6.5 trains with
test_pissa_reproduces_the_basedifferentialPiSSA outputs equal the base’s within float32step 0 is still the pretrained model
test_pissa_adapter_is_the_top_singular_partgoldensBA=Wrs B A = W_r from LAPACK; BB‘s columns are ±Sj/s uj\pm\sqrt{S_j / s}\,u_jPiSSA’s point: train the principal directions
test_merge_equals_unmergeddifferentialmerged plain Linears give the adapted outputs within 1e-5, default and PiSSAserving pays nothing for the adapter
test_merge_restores_the_base_keyspropertysame keys, same order as before injection; with B=0B = 0 the same valuesevery loader reads a merged checkpoint
test_peft_key_namesgoldenbase_model.model.q_proj.lora_A.weight and friends, shapes [r,in][r, \text{in}] and [out,r][\text{out}, r]adapters move between tools
test_adapter_roundtrippropertyan adapter file restores the adapter; wrong keys fail before copyingresuming and sharing adapters
test_dropout_only_on_the_adapter_pathpropertywith B=0B = 0 dropout changes nothing; in eval mode it is the identitythe frozen path sees the true input
test_trainable_fractionunit16 trainable elements out of the totalthe MS-L6 trainable_frac bar
test_validationboundaryrank bounds, α≤0\alpha \le 0, unknown init, a non-Linear basecaller bugs fail loudly

Rung R5 asks for oracles: expected values from an independent computation. For LoRA the oracles are numpy itself: recompute xW⊤+b+s xA⊤B⊤x W^\top + b + s\,x A^\top B^\top, take LAPACK’s np.linalg.svd for PiSSA, and check gradients with your own finite differences in float64. Cover the formula with r>1r > 1, the bitwise start, which parameters train and get gradients, PiSSA against LAPACK, merging (outputs, removed adapters, unchanged keys), PEFT names, dropout on the adapter path only, and the rank bound. Import only the contract (tinyllm.obj.lora and the L0.4 layers). ol check L6.6 requires a mutation score of at least 0.80 with every required pitfall fault killed.

PitfallSymptomCaught by
1. scaling by α\alpha instead of α/r\alpha / ra learning rate that works at r=4r = 4 diverges at r=16r = 16test_scaling_is_alpha_over_r, test_adapter_gradcheck (mutant s01)
2. a random BB at default initthe “fine-tune” starts from a damaged model; step-0 loss above the base’stest_default_init_equals_base_bitwise (mutant s02)
3. the wrapped weight left trainable, or only the wrapped layers frozenthe whole model trains: optimizer memory and checkpoints as large as a full fine-tunetest_only_adapter_parameters_get_gradients (mutant s03), test_inject_freezes_everything_else (mutant s09)
4. PiSSA without the residual, or without dividing by s\sqrt sstep 0 computes W+WrW + W_r or W−Wr+sWrW - W_r + s W_r: not the base modeltest_pissa_reproduces_the_base, test_pissa_adapter_is_the_top_singular_part (mutants s04, s05)
5. merging without ss, or leaving the adapter in placemerged outputs differ, or the update is counted twicetest_merge_equals_unmerged (mutants s06, s07)
6. adapter keys without PEFT’s prefix; dropout on the frozen pathadapters that no other tool loads; a noisy base model in trainingtest_peft_key_names (mutant s08), test_dropout_only_on_the_adapter_path (mutant s10)
DirectionModuleHow it uses this
BackL0.1freezing is requires_grad; from_op records no vjp for frozen inputs
BackL0.2the adapter path is matmul and transpose of the op library
BackL0.4LoRALinear wraps a Linear, reuses its Tensors, and builds lora_A and lora_B as Linears
BackM03.5low_rank gives PiSSA’s balanced factors
ForwardL6.5lora_classifier adapts a classifier’s attention queries and values and keeps the new head trainable
ForwardL12.1supervised fine-tuning of the chat model trains LoRA adapters (optional, Pass 10)
Forwardsq.multi-lorathe engine serves many adapters over one base, unmerged

If you skip this module, ol check L6.5 stops with needs L6.6: build it, or pass --ref-deps.

Your pieceProduction equivalentWhat it addsWhere to look
LoRALinear, inject_loraHugging Face PEFT LoraConfig, get_peft_modeladapters for Conv1D and embeddings, several named adapters per layer, modules_to_save, rank patterns per modulepeft/tuners/lora/layer.py, peft/tuners/lora/model.py
init="pissa"PEFT init_lora_weights="pissa" and "pissa_niter_4"a fast randomized SVD for large weights; also OLoRA, LoftQ, EVA initializationspeft/tuners/lora/layer.py (pissa_init)
merge_loramerge_and_unload(), add_weighted_adaptermerging several adapters with weights (TIES, DARE)peft/tuners/lora/model.py
unmerged adapters at serving timevLLM and S-LoRA multi-LoRA servingthousands of adapters over one base, batched with custom kernels (Punica SGMV)vLLM vllm/lora/; Sheng et al., “S-LoRA” (2023)
low-rank trainingQLoRA, DoRAadapters over a 4-bit base; a magnitude and direction split of WWDettmers et al. 2023; Liu et al. 2024