jasonfan's picture
2026-03-19
3b2d368 verified
Raw
History Blame Contribute Delete
1.32 kB
import math
import torch
import torch.nn as nn
class LoRALinear(nn.Linear):
"""
Linear layer with LoRA - inherits from nn.Linear to keep parameter names unchanged
This allows loading pretrained weights without name conversion
"""
def __init__(self, in_features, out_features, bias=True,
lora_rank=8, lora_alpha=16, lora_dropout=0.0):
super().__init__(in_features, out_features, bias=bias)
self.lora_rank = lora_rank
self.lora_alpha = lora_alpha
self.scaling = lora_alpha / lora_rank
# LoRA parameters
self.lora_A = nn.Parameter(torch.zeros(in_features, lora_rank))
self.lora_B = nn.Parameter(torch.zeros(lora_rank, out_features))
self.lora_dropout = nn.Dropout(lora_dropout) if lora_dropout > 0 else nn.Identity()
# Initialize LoRA weights
nn.init.kaiming_uniform_(self.lora_A, a=math.sqrt(5))
nn.init.zeros_(self.lora_B)
def forward(self, x):
# Original linear transformation
result = super().forward(x)
# Add LoRA adaptation
lora_result = (x @ self.lora_A) @ self.lora_B
lora_result = self.lora_dropout(lora_result)
result = result + lora_result * self.scaling
return result