LoRA Linear
2026/6/6大约 1 分钟
LoRA Linear
题目描述
实现 LoRA(Low-Rank Adaptation)线性层的前向传播。给定输入矩阵 ()、基础权重矩阵 ()、LoRA 降维投影矩阵 ()和 LoRA 升维投影矩阵 (),计算:
其中 。所有张量为 float32。
实现要求
- 实现
solve函数,签名保持不变。 - 不允许使用外部库。
- 将结果写入
output。
示例
Input: x = [[1,0,-1,2],[0,1,1,-1]], W = I_3x4 (前3行), A = [[1,0,0,0],[0,1,0,0]], B = [[1,0,0],[0,1,0]]
lora_scale = 0.5
x @ W^T = [[1,0,-1],[0,1,1]]
x @ A^T = [[1,0],[0,1]]
α * (x@A^T) @ B^T = 0.5 * [[1,0,0],[0,1,0]]
Output: [[1.5, 0, -1], [0, 1.5, 1]]约束条件
- ,。
- ,。
- 性能测试在 下进行。
解题思路
LoRA 将全秩权重更新分解为两个低秩矩阵的乘积。计算流程:(GEMM)+ (两次小 GEMM 或一次 batched GEMM)。由于 ,LoRA 分支的计算量远小于主分支,但增加了一次额外的矩阵乘。可以将 视为合并权重来一次性完成计算。
代码实现
CUDA
#include <cuda_runtime.h>
__global__ void lora_kernel(const float* x, const float* W, const float* A, const float* B,
float* out, int batch, int din, int dout, int rank, float alpha) {
int b=blockIdx.x, o=threadIdx.x;
if(b<batch && o<dout) {
float base=0,lora=0;
for(int i=0;i<din;i++) base+=x[b*din+i]*W[o*din+i];
float* tmp=(float*)malloc(rank*sizeof(float));
for(int r=0;r<rank;r++){tmp[r]=0;for(int i=0;i<din;i++)tmp[r]+=x[b*din+i]*A[r*din+i];}
for(int r=0;r<rank;r++) lora+=tmp[r]*B[o*rank+r];
out[b*dout+o]=base+alpha*lora;
free(tmp);
}
}
extern "C" void solve(const float* x, const float* W, const float* A, const float* B,
float* out, int batch, int din, int dout, int rank, float alpha) {
lora_kernel<<<batch,dout>>>(x,W,A,B,out,batch,din,dout,rank,alpha);
cudaDeviceSynchronize();
}Triton
import triton, triton.language as tl
@triton.jit
def lora_kernel(x_ptr,W_ptr,A_ptr,B_ptr,out_ptr, batch,din,dout,rank,alpha, BLOCK:tl.constexpr):
b=tl.program_id(0)*BLOCK+tl.arange(0,BLOCK); mask=b<batch
din_range=tl.arange(0,din); r_range=tl.arange(0,rank)
xb=tl.load(x_ptr+b[:,None]*din+din_range[None,:],mask=mask[:,None])
base=xb@tl.load(W_ptr) # [batch, dout]
xa=xb@tl.trans(tl.load(A_ptr+r_range[:,None]*din+din_range[None,:])) # [batch, rank]
lora=xa@tl.load(B_ptr) # [batch, dout]
tl.store(out_ptr+b[:,None]*dout+tl.arange(0,dout)[None,:],base+alpha*lora,mask=mask[:,None])