1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115
| import torch import torch.nn as nn
class LoRALinear(nn.Module): def __init__(self, in_features, out_features, r): super(LoRALinear, self).__init__() self.in_features = in_features self.out_features = out_features self.r = r
self.weight = nn.Parameter(torch.randn(in_features, out_features)) self.weight.requires_grad = False self.bias = nn.Parameter(torch.zeros(out_features)) self.bias.requires_grad = False
self.B = nn.Parameter(torch.empty(out_features, r)) self.A = nn.Parameter(torch.zeros(r, in_features)) nn.init.normal_(self.A, mean=0.0, std=0.02)
def forward(self, x): original_output = torch.nn.functional.linear(x, self.weight, self.bias) delta_W = torch.matmul(self.B, self.A) lora_output = torch.nn.functional.linear(x, delta_W) return original_output + lora_output
class LoRAAttention(nn.Module): def __init__(self, dim, rank): super(LoRAAttention, self).__init__() self.dim = dim self.rank = rank
self.original_q_W = nn.Linear(dim, dim) self.original_k_W = nn.Linear(dim, dim) self.original_v_W = nn.Linear(dim, dim) self.original_o_W = nn.Linear(dim, dim)
for p in self.original_q_W.parameters(): p.requires_grad = False for p in self.original_k_W.parameters(): p.requires_grad = False for p in self.original_v_W.parameters(): p.requires_grad = False
self.q_B = nn.Parameter(torch.zeros(dim, rank)) self.q_A = nn.Parameter(torch.empty(rank, dim)) nn.init.normal_(self.q_A, mean=0.0, std=0.02)
self.k_B = nn.Parameter(torch.zeros(dim, rank)) self.k_A = nn.Parameter(torch.empty(rank, dim)) nn.init.normal_(self.k_A, mean=0.0, std=0.02)
self.v_B = nn.Parameter(torch.zeros(dim, rank)) self.v_A = nn.Parameter(torch.empty(rank, dim)) nn.init.normal_(self.v_A, mean=0.0, std=0.02)
def forward(self, q, k, v): Q = self.original_q_W(q) K = self.original_k_W(k) V = self.original_v_W(v)
delta_q = torch.matmul(q, self.q_B) delta_q = torch.matmul(delta_q, self.q_A)
delta_k = torch.matmul(k, self.k_B) delta_k = torch.matmul(delta_k, self.k_A)
delta_v = torch.matmul(v, self.v_B) delta_v = torch.matmul(delta_v, self.v_A)
Q = Q + delta_q K = K + delta_k V = V + delta_v
attn = torch.matmul(Q, K.transpose(-2, -1)) / (self.embed_dim ** 0.5) attn = torch.nn.functional.softmax(attn, dim=-1) out = torch.matmul(attn, V) out = self.o_W(out) return out
def count_params(module: torch.nn.Module): trainable = 0 total = 0 for p in module.parameters(): total += p.numel() if p.requires_grad: trainable += p.numel() return trainable, total
if __name__ == "__main__": b, l, d = 1, 512, 768 r = 8 lora = LoRALinear(in_features=d, out_features=d, r = r).to(device='cuda') trainable, total = count_params(lora) print(f'trainable : {trainable:,}') print(f'total : {total:,}')
lora = LoRAAttention(d, r).to(device='cuda') trainable, total = count_params(lora) print(f'trainable : {trainable:,}') print(f'total : {total:,}')
|