当前位置:首页 > 技术 > 正文内容

LoRA微调:原理、实现与调参实践

访客 技术 2026年10月9日 2

一、LoRA的核心思路

LoRA(Low-Rank Adaptation)并非全新的训练算法,而是一种高效的模型适配技术。其基本思路是:冻结预训练模型的全部参数,仅在模型特定层旁插入少量可训练的低秩矩阵。训练时只更新这些小型矩阵,从而大幅降低计算和存储开销。

类比来说,全参数微调就像重写整本书去添加新知识,而LoRA则是在书中夹几页"速查卡",只修改卡片内容即可快速适应新任务。

方案策略
全量微调更新模型全部参数,计算量和存储需求巨大
LoRA冻结原参数,插入低秩矩阵,仅优化极少量参数

二、低秩分解工作原理

2.1 原始权重矩阵

假设大模型某层权重矩阵W为5×4,共有20个参数。

全量微调需更新全部20个参数,微调后权重变为W + ΔW。

2.2 低秩分解ΔW

LoRA将增量矩阵ΔW分解为两个小矩阵的乘积:

  • 矩阵A:尺寸5×2,10个参数
  • 矩阵B:尺寸2×4,8个参数

总参数量为18,比原来少2个(本例因矩阵小减幅有限,实际大型矩阵效果显著)。

训练对象从20个参数降为18个,且计算图更紧凑。

2.3 为什么低秩可行

权重更新ΔW通常存在大量冗余,其有效信息可用较低秩表达。矩阵的秩表示线性独立方向的数量,若秩远小于维度,则可用少数方向近似描述变化。

例如,一个5×5矩阵中若某行是其他行的线性组合,则该行不提供新信息,秩降低。

2.4 实际节省比例

对于512×512的权重矩阵(262,144参数):

  • 全微调:更新262,144个参数
  • LoRA(r=8):A矩阵512×8=4,096 + B矩阵8×512=4,096,总计8,192个参数,节省约97%参数

三、资源估算

可使用显存计算工具(如llamafactory的GPU内存估算页面)预估不同模型和参数下的显存需求。根据模型大小、LoRA秩、批量大小等参数,工具会给出训练所需显存及公式说明。

注意:实际消耗受框架实现、梯度累积等因素影响。

四、关键参数调优

4.1 LoRA主要配置项

参数含义
r低秩矩阵的秩,控制可训练参数量
lora_alpha缩放因子,通常设置为2r或根据实验调整
lora_dropout正则化,防过拟合,常用0.05~0.1
target_modules应用LoRA的模块列表(如q_proj, v_proj)

4.2 调参经验

  • 起始秩选择:大多数任务从r=8或r=16开始,根据评估结果增减。
  • 数据集规模建议:样本<5k时r=8即可;>50k时可尝试r=32或r=64。
  • 复杂任务策略:对推理密集型任务,可增大r至32~64,启用rsLoRA稳定训练,并扩展target_modules(如加入FFN层)。
  • 领域差异处理:若LoRA训练效果不佳且任务与预训练数据领域差异大,可提升秩(如从8提至64)。
  • 基座模型选择:参数更多、量化更低的模型通常优于参数量小但精度高的模型。例如33B-fp4优于13B-fp16。
  • 数据清洗要点:将PDF、Word、HTML等格式统一转为纯文本,过滤噪声和乱码;过滤长度小于100字符的文档。

五、开源项目参考

  • ChatGLM-Tuning(GitHub):包含LoRA微调脚本,适合对话模型。
  • baichuan_sft_lora(GitHub):基于QLoRA微调百川模型,展示高效微调流程。

六、完整微调代码实战

以下示例基于Transformers、PEFT和Datasets库,演示从数据加载、模型准备、训练到推理合并的完整流程。

import json
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, TrainingArguments, Trainer, DataCollatorForLanguageModeling
from peft import LoraConfig, get_peft_model, PeftModel
from datasets import Dataset
from modelscope import snapshot_download

def prepare_dataset(path, tokenizer=None):
    """加载JSON格式数据集,转为Dataset并进行tokenize"""
    with open(path, 'r', encoding='utf-8') as f:
        raw_data = json.load(f)
    
    formatted = []
    for item in raw_data:
        text = f"Human: {item['instruction']}\nAssistant: {item['output']}"
        formatted.append({'text': text})
    
    dataset = Dataset.from_list(formatted)
    if tokenizer:
        def tokenize_fn(examples):
            return tokenizer(
                examples['text'],
                truncation=True,
                max_length=512,
                padding=False
            )
        dataset = dataset.map(tokenize_fn, batched=True, remove_columns=dataset.column_names)
    return dataset

def build_lora_model(base_path, rank=16, device="auto", precision=torch.float32):
    """加载基座模型并插入LoRA适配层"""
    tokenizer = AutoTokenizer.from_pretrained(base_path)
    tokenizer.pad_token = tokenizer.eos_token

    base_model = AutoModelForCausalLM.from_pretrained(
        base_path,
        torch_dtype=precision,
        device_map=device
    )
    base_model.config.pad_token_id = base_model.config.eos_token_id

    # 冻结原模型
    for param in base_model.parameters():
        param.requires_grad = False

    lora_cfg = LoraConfig(
        r=rank,
        lora_alpha=32,
        lora_dropout=0.1,
        bias="none",
        task_type="CAUSAL_LM",
        target_modules=["q_proj", "k_proj", "v_proj"]
    )
    peft_model = get_peft_model(base_model, lora_cfg)
    peft_model.print_trainable_parameters()
    return peft_model, tokenizer

def train_lora(base_path, rank=16, device="auto", precision=torch.float32, dataset_path="./data.json", output_dir="./lora_weights"):
    """执行LoRA微调训练"""
    model, tokenizer = build_lora_model(base_path, rank, device, precision)
    train_dataset = prepare_dataset(dataset_path, tokenizer)

    collator = DataCollatorForLanguageModeling(tokenizer, mlm=False)

    args = TrainingArguments(
        output_dir=output_dir,
        num_train_epochs=3,
        per_device_train_batch_size=1,
        gradient_accumulation_steps=4,
        learning_rate=1e-4,
        logging_steps=10,
        remove_unused_columns=False,
        push_to_hub=False,
        logging_dir="./logs"
    )

    trainer = Trainer(
        model=model,
        args=args,
        train_dataset=train_dataset,
        data_collator=collator
    )
    trainer.train()
    model.save_pretrained(output_dir)
    tokenizer.save_pretrained(output_dir)

def inference_with_lora(base_path, lora_dir="./lora_weights", precision=torch.float32):
    """加载原始模型和LoRA权重进行推理"""
    model = AutoModelForCausalLM.from_pretrained(base_path, torch_dtype=precision)
    tokenizer = AutoTokenizer.from_pretrained(base_path)
    model = PeftModel.from_pretrained(model, lora_dir)
    model = model.merge_and_unload()
    tokenizer.pad_token = tokenizer.eos_token
    model.config.pad_token_id = model.config.eos_token_id

    prompts = [
        "只剩一个心脏了还能活吗?",
        "如何学习一门编程语言?",
        "你是谁?能帮我解决什么问题?"
    ]
    for q in prompts:
        full = f"Human: {q}\nAssistant:"
        inputs = tokenizer(full, return_tensors="pt")
        outputs = model.generate(**inputs, max_length=100, temperature=0.7)
        reply = tokenizer.decode(outputs[0], skip_special_tokens=True)
        print(f"Q: {q}\nA: {reply}\n")

def merge_lora(base_path, lora_dir="./lora_weights", precision=torch.float32, save_path="./merged_model"):
    """合并LoRA权重到基座模型并保存完整模型"""
    model = AutoModelForCausalLM.from_pretrained(base_path, torch_dtype=precision)
    tokenizer = AutoTokenizer.from_pretrained(base_path)
    model = PeftModel.from_pretrained(model, lora_dir)
    model = model.merge_and_unload()
    tokenizer.pad_token = tokenizer.eos_token
    model.config.pad_token_id = model.config.eos_token_id
    model.save_pretrained(save_path, safe_serialization=True)
    tokenizer.save_pretrained(save_path)
    print(f"合并后模型已保存至 {save_path}")

def download_model(model_id="LLM-Research/Llama-3.2-1B", cache="./model_cache"):
    """从ModelScope下载模型"""
    snapshot_download(model_id=model_id, cache_dir=cache)

if __name__ == "__main__":
    MODEL_ID = "LLM-Research/Llama-3.2-1B"
    MODEL_CACHE = "./model_cache/LLM-Research/Llama-3___2-1B"
    DATA_PATH = "./data.json"
    RANK = 16
    DEVICE = "mps"   # 根据环境选择 "cuda" / "mps" / "cpu"
    DTYPE = torch.float32

    # 1. 下载基座模型(首次运行)
    # download_model(MODEL_ID)

    # 2. 开始训练
    train_lora(MODEL_CACHE, rank=RANK, device=DEVICE, precision=DTYPE, dataset_path=DATA_PATH)

    # 3. 测试推理
    # inference_with_lora(MODEL_CACHE, lora_dir="./lora_weights", precision=DTYPE)

    # 4. 合并保存完整模型
    # merge_lora(MODEL_CACHE, lora_dir="./lora_weights", precision=DTYPE, save_path="./merged_model")

执行步骤

  1. 下载基座模型:取消注释 download_model() 并运行。
  2. 训练:直接运行 train_lora()(已默认执行)。
  3. 测试推理:取消注释 inference_with_lora() 查看效果。
  4. 合并模型:取消注释 merge_lora() 生成完整权重,可用于部署或上传平台。

合并后的模型保存在 ./merged_model 下,包含 model.safetensors 文件。

部署到Ollama

将合并模型转换为GGUF格式后可导入Ollama使用。例如:

python convert_hf_to_gguf.py /path/to/merged_model --outtype q8_0 --outfile model.gguf

支持精度选项:f32(全精)、f16(半精)、bf16、q8_0(8位量化)、tq1_0/tq2_0(1/2位量化,体积极小且质量损失大)。

完整项目代码可参考:Gitee仓库。

返回列表

上一篇:使用 Docker Compose 快速搭建常用基础服务

没有最新的文章了...

相关文章

Linux crontab 详解

1) crontab 是什么cron 是 Linux 的定时任务守护进程;crontab 是用来编辑/查看“按时间周期执行命令”的表(cron table)。常见两类:用户 crontab:每个用户一份(crontab -e 编辑)系统级 crontab / cron.d:可指定执行用户(/etc/crontab、/etc/cron.d/*)2) crontab 时间...

富文本里可以允许的 HTML 属性

一、所有标签默认允许的安全属性(极少)class        (可选)id           (通常建议禁用)title️ 注意:id 容易被滥用做锚点注入,很多系统直接禁用class 允许的话最好只允许固定前缀(如 editor-*)二、a 标签允许属性<a href="" t...

Mac 安装 Node.js 指南

方法一:通过官网安装包(最简单,适合初学者)如果你只是想快速安装并开始使用,这是最直接的方法。访问 Node.js 官网。页面会显示两个版本:LTS (Recommended For Most Users):长期支持版,最稳定。建议选这个。Current:最新特性版,包含最新功能但可能不够稳定。下载 .pkg 安装包并运行。按照安装向导点击“下一步”即可完成。方法二:使用 Homebrew 安装(...

Dom\HTML_NO_DEFAULT_NS 的副作用:自动加闭合标签

在使用Dom\HTMLDocument时,Dom\HTML_NO_DEFAULT_NS 将禁止在解析过程中设置元素的命名空间, 此设置是为了与DOMDocument向后兼容而存在的。当使用它时,已知的一个副作用就是:自动加闭合标签例如 </img> 为什么会这样?当你使用:Dom\HTML_NO_DEFAULT_NS文档会变成 无命名空间模式,此时内部更接近 XML...

Laravel 事件和监听器创建

在 Laravel 中,使用 Artisan 命令创建 Events(事件) 和 Listeners(监听器) 是非常高效的。你可以通过以下几种方式来实现:1. 手动创建单个 Event如果你只想创建一个事件类,可以使用 make:event 命令:Bashphp artisan make:event UserRegistered执行后,文件将生成在 app/Even...

自定义域名解析神器 dnsmasq

什么是 dnsmasq?dnsmasq 是一个轻量级、功能强大的网络服务工具,专为小型和中等规模网络设计。它是一个综合的网络基础设施解决方案[1]。dnsmasq 能做什么?功能说明应用场景DNS 转发与缓存将 DNS 查询转发到上游服务器(ISP、Google DNS 等),并在本地缓存结果加快 DNS 查询速度,减少外部 DNS 流量本地 DNS解析本地网络设备的主机名,无需编辑&n...

发表评论

访客

◎欢迎参与讨论,请在这里发表您的看法和观点。