【新手入门】全网最全的 LoRA 代码讲解

前言

文章性质:简陋的学习记录 📖

主要内容:本文简单记录了作者在阅读 LoRA 代码时的笔记。

冷知识+1:小伙伴们不经意的 点赞 👍🏻 与 收藏 ✨ 可以让作者更有创作动力! 

目录

一、LoRALayer

二、LoRALinear

三、apply_lora_to_linear_layers

四、get_submodule

五、get_lora_params


一、LoRALayer

这里定义了一个名为 LoRALayer 的类, 继承自 nn.Moudle,这是 PyTorch 中所有神经网络模块的基类。

__init__ 方法是初始化函数,接收输入维度 in_dim、输出维度 out_dim,以及可选的秩 rank 和缩放因子 alpha 参数。

super().__init__() 调用了父类的初始化方法,再将 alpha 和 rank 存储为实例变量。

接着定义两个可训练的参数矩阵 lora_A 和 lora_B,lora_A 的形状为 (rank, in_dim),lora_B 的形状为 (out_dim, rank)。

这两个矩阵初始化为全零,并被包装为 nn.Parameter,使其在训练过程中被优化器更新。

lora_A 用 Kaiming 均匀分布初始化,适用于具有 ReLU 激活函数的层,a=math.sqrt(5) 是与 Kaiming 初始化相关的参数。 

lora_B 被初始化为全零,以确保其初始时对原始模型的影响最小。

class LoRALayer(nn.Module):
    def __init__(self, in_dim, out_dim, rank=4, alpha=1.0):
        super().__init__()
        self.alpha = alpha
        self.rank = rank
        
        self.lora_A = nn.Parameter(torch.zeros((rank, in_dim)))
        self.lora_B = nn.Parameter(torch.zeros((out_dim, rank)))
        
        nn.init.kaiming_uniform_(self.lora_A, a=math.sqrt(5))
        nn.init.zeros_(self.lora_B)

forward 是前向传播方法 ,输入 x 是数据张量。​首先,获取输入张量的形状并存储在 orig_shape 中。

检查输入张量 x 的维度是否大于 2:

如果大于 2,则说明 x 可能是图像或更高维度的数据,需将其展平为二维张量,形状为 (batch_size*序列长度, features)。

否则说明 x 本身就是二维张量,即 (batch_size, features),无需变换,直接将其赋值给 x_2d。

接着计算 LoRA 的权重文件,先使 lora_B lora_A 进行矩阵乘法,得到形状为 (out_dim, in_dim) 的矩阵。

然后乘以缩放因子 alpha / rank,以控制更新的幅度。再将输入数据 x_2d 与 LoRA 权重矩阵的转置相乘,得到输出 output。

输出 output 是新张量,其形状通常为 (batch_size*序列长度, lora_weight.shape[0])。

如果输入数据最初是多维的,则将输出 output 重塑回原始的多维形状,以匹配输入的形状。

    def forward(self, x):
        orig_shape = x.shape
        
        if len(orig_shape) > 2:
            x_2d = x.reshape(-1, orig_shape[-1])
        else:
            x_2d = x
            
        lora_weight = (self.lora_B @ self.lora_A) * (self.alpha / self.rank)
        output = x_2d @ lora_weight.T
        
        i
评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包

打赏作者

作者正在煮茶

你的鼓励将是我创作的最大动力

¥1 ¥2 ¥4 ¥6 ¥10 ¥20
扫码支付:¥1
获取中
扫码支付

您的余额不足,请更换扫码支付或充值

打赏作者

实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值