"""LoRA Linear：初始输出与基础层一致。"""
import math
import torch
from torch import nn


class LoRALinear(nn.Module):
    def __init__(self, base, rank=8, alpha=16):
        super().__init__(); self.base=base; self.scale=alpha/rank
        for p in base.parameters(): p.requires_grad=False
        self.A=nn.Parameter(torch.empty(rank,base.in_features))
        self.B=nn.Parameter(torch.zeros(base.out_features,rank))
        nn.init.kaiming_uniform_(self.A,a=math.sqrt(5))
    def forward(self,x): return self.base(x)+self.scale*((x@self.A.T)@self.B.T)


if __name__ == "__main__":
    base=nn.Linear(16,12); lora=LoRALinear(base)
    x=torch.randn(3,16)
    assert torch.allclose(base(x),lora(x))
    print("trainable parameters:",sum(p.numel() for p in lora.parameters() if p.requires_grad))

