基于 Qwen/Llama 的领域大模型微调实战:从训练到部署的全流程记录

本文完整记录了一次真实的大模型领域微调实践:从环境搭建、数据处理、SFT 微调、量化压缩到推理部署的全链路,覆盖 LoRA/QLoRA、DeepSpeed、vLLM 等核心技术栈。


目录


一、背景与动机

通用大模型(如 GPT-4、Qwen、Llama 系列)在广泛任务上表现出色,但在垂直领域往往存在以下问题:

  • 领域知识浅层化:对行业术语、业务流程理解不到位;
  • 输出风格不匹配:通用模型的回答方式不符合企业内部规范;
  • 数据隐私要求:不能将敏感数据发送到公有云 API。

因此,在特定领域的高质量语料上对开源基座模型进行监督微调(Supervised Fine-Tuning, SFT),是当前企业落地大模型的主流路径。

本文基于 Qwen2.5-7B 和 Llama-3.1-8B 两个模型家族,以"法律咨询"领域作为示例场景,完整走通从数据到部署的全流程。所有方法可平替至医疗、金融、教育等其他垂直领域。


二、硬件与软件环境

2.1 训练环境

组件配置
GPU4 × NVIDIA A100 80GB SXM
CPU2 × AMD EPYC 7763 (128 核)
内存512 GB DDR4
存储4 TB NVMe SSD
OSUbuntu 22.04 LTS
CUDA12.1
Python3.10

单卡 A100 40G / RTX 4090 24G 也可以跑,后文会说明对应参数的调整方式。

2.2 关键依赖

# 创建虚拟环境
conda create -n llm-finetune python=3.10 -y
conda activate llm-finetune

# PyTorch (CUDA 12.1)
pip install torch==2.4.0 torchvision==0.19.0 torchaudio==2.4.0 --index-url https://download.pytorch.org/whl/cu121

# 训练框架
pip install transformers==4.44.0
pip install accelerate==0.33.0
pip install peft==0.12.0
pip install trl==0.9.6
pip install deepspeed==0.15.1

# 数据处理
pip install datasets==2.20.0
pip install sentencepiece==0.2.0
pip install tiktoken==0.7.0

# 量化与推理
pip install bitsandbytes==0.43.3
pip install auto-gptq==0.7.1
pip install vllm==0.5.4

# 其他工具
pip install wandb==0.17.5        # 训练监控
pip install flash-attn --no-build-isolation  # Flash Attention 2

注意flash-attn 编译耗时较长(约 10-20 分钟),建议提前安装并验证:

python -c "import flash_attn; print(flash_attn.__version__)"

三、数据准备与清洗

数据是微调的灵魂。下面是一套可复用的数据处理管线。

3.1 数据格式

采用 LLaMA-Factory / Alpaca 兼容的 JSON 格式,每条样本包含三轮对话的系统提示、用户输入和助手输出:

[
  {
    "conversations": [
      {
        "role": "system",
        "content": "你是一名专业法律顾问,擅长解答民商事法律问题。回答应严谨、引用法条、避免主观臆断。"
      },
      {
        "role": "user",
        "content": "劳动合同到期不续签,公司需要支付经济补偿吗?"
      },
      {
        "role": "assistant",
        "content": "根据《中华人民共和国劳动合同法》第四十六条第五项,除用人单位维持或者提高劳动合同约定条件续订劳动合同,劳动者不同意续订的情形外,劳动合同期满终止固定期限劳动合同的,用人单位应当向劳动者支付经济补偿。\n\n经济补偿的计算标准依据第四十七条:按劳动者在本单位工作的年限,每满一年支付一个月工资。六个月以上不满一年的,按一年计算;不满六个月的,支付半个月工资。\n\n参考案例:(2023)京0105民初12345号判决中,法院支持了劳动者因合同到期不续签而主张的经济补偿请求。"
      }
    ]
  }
]

3.2 数据清洗管线

import json
import re
from pathlib import Path
from typing import List, Dict, Any


class DataCleaner:
    """数据清洗器:去重、长度过滤、格式校验"""

    def __init__(
        self,
        min_input_len: int = 10,
        max_input_len: int = 4096,
        min_output_len: int = 50,
        max_output_len: int = 4096,
        dedup_threshold: float = 0.85,
    ):
        self.min_input_len = min_input_len
        self.max_input_len = max_input_len
        self.min_output_len = min_output_len
        self.max_output_len = max_output_len
        self.dedup_threshold = dedup_threshold

    def clean(self, raw_path: str, output_path: str) -> Dict[str, int]:
        raw_data = self._load(raw_path)
        stats = {"total": len(raw_data)}

        # Step 1: 格式校验
        raw_data = [d for d in raw_data if self._validate_format(d)]
        stats["format_valid"] = len(raw_data)

        # Step 2: 长度过滤
        raw_data = [d for d in raw_data if self._validate_length(d)]
        stats["length_valid"] = len(raw_data)

        # Step 3: 特殊字符清洗
        raw_data = [self._clean_special_chars(d) for d in raw_data]
        stats["chars_cleaned"] = len(raw_data)

        # Step 4: 语义去重(基于 MinHash LSH)
        raw_data = self._deduplicate(raw_data)
        stats["after_dedup"] = len(raw_data)

        self._save(raw_data, output_path)
        return stats

    def _validate_format(self, item: Dict) -> bool:
        if "conversations" not in item:
            return False
        conv = item["conversations"]
        if len(conv) < 2:
            return False
        # 必须有 user 和 assistant 角色
        roles = {turn["role"] for turn in conv}
        return "user" in roles and "assistant" in roles

    def _validate_length(self, item: Dict) -> bool:
        conv = item["conversations"]
        user_content = " ".join(
            t["content"] for t in conv if t["role"] == "user"
        )
        assistant_content = " ".join(
            t["content"] for t in conv if t["role"] == "assistant"
        )
        return (
            self.min_input_len <= len(user_content) <= self.max_input_len
            and self.min_output_len <= len(assistant_content) <= self.max_output_len
        )

    def _clean_special_chars(self, item: Dict) -> Dict:
        """去除不可见字符、统一换行符"""
        for turn in item["conversations"]:
            turn["content"] = re.sub(r"[\x00-\x08\x0b\x0c\x0e-\x1f]", "", turn["content"])
            turn["content"] = turn["content"].replace("\r\n", "\n").strip()
        return item

    def _deduplicate(self, data: List[Dict]) -> List[Dict]:
        """基于 n-gram Jaccard 相似度的近似去重"""
        seen_hashes = set()
        result = []
        for item in data:
            sig = self._signature(item)
            is_dup = False
            for sh in seen_hashes:
                if self._jaccard(sig, sh) > self.dedup_threshold:
                    is_dup = True
                    break
            if not is_dup:
                seen_hashes.add(sig)
                result.append(item)
        return result

    def _signature(self, item: Dict) -> set:
        text = " ".join(t["content"] for t in item["conversations"])
        tokens = text[:500]  # 只用前 500 字符做签名
        return set(tokens[i : i + 8] for i in range(0, len(tokens) - 7, 4))

    def _jaccard(self, a: set, b: set) -> float:
        if not a or not b:
            return 0.0
        return len(a & b) / len(a | b)

    def _load(self, path: str) -> List[Dict]:
        with open(path, "r", encoding="utf-8") as f:
            return json.load(f)

    def _save(self, data: List[Dict], path: str):
        with open(path, "w", encoding="utf-8") as f:
            json.dump(data, f, ensure_ascii=False, indent=2)

3.3 数据集拆分

from sklearn.model_selection import train_test_split

with open("data/clean/legal_sft.json", "r", encoding="utf-8") as f:
    data = json.load(f)

train, temp = train_test_split(data, test_size=0.1, random_state=42)
val, test = train_test_split(temp, test_size=0.5, random_state=42)

for name, subset in [("train", train), ("val", val), ("test", test)]:
    with open(f"data/split/{name}.json", "w", encoding="utf-8") as f:
        json.dump(subset, f, ensure_ascii=False, indent=2)
    print(f"{name}: {len(subset)} samples")

典型规模参考:

用途样本量说明
训练集5,000 - 50,000多轮对话、单轮问答混合
验证集500 - 5,000与训练集同分布
测试集500 - 2,000部分 OOD 样本用于泛化性评估

四、基座模型选型

当前开源社区的两大主流选择:

维度Qwen2.5-7B-InstructLlama-3.1-8B-Instruct
参数量7B8B
词表大小152,064128,256
上下文长度128K (原生)128K (RoPE 扩展)
中文能力⭐⭐⭐⭐⭐⭐⭐⭐
多语言中 / 英 / 29 语种英 / 德 / 法 / 意 / 葡 / 西 / 泰 (中文较弱)
许可协议Apache 2.0Llama 3.1 Community
推荐场景中文为主 / 多语言场景英文为主 / 国际化场景

选型建议

  • 中文场景 → Qwen2.5-7B-Instruct,原生中文能力强,微调成本低;
  • 英文 / 代码场景 → Llama-3.1-8B-Instruct,社区生态丰富;
  • 资源受限(单卡 24G)→ 两者均可通过 QLoRA (4-bit) 完成微调;
  • 需要更大容量 → Qwen2.5-14B/32B/72B 或 Llama-3.1-70B。

五、微调方案设计

5.1 三种主流方案对比

方案显存需求 (7B)训练速度效果适用场景
Full Fine-Tuning~112 GB最优数据量大、资源充足
LoRA (r=64, α=128)~48 GB接近全参最常用,性价比高
QLoRA (4-bit + LoRA)~16 GB中等略低于 LoRA单卡 24G 可用

5.2 本方案选择:QLoRA + DeepSpeed ZeRO-2

  • QLoRA:基座 4-bit 量化 + LoRA 低秩适配,大幅降低显存;
  • DeepSpeed ZeRO-2:分片优化器状态与梯度,支持多卡并行;
  • Flash Attention 2:加速注意力计算;
  • Target Modules:所有线性层 (q_proj, k_proj, v_proj, o_proj, gate_proj, up_proj, down_proj)

5.3 核心训练脚本

# train_sft.py
import os
import torch
from datasets import load_dataset
from transformers import (
    AutoTokenizer,
    AutoModelForCausalLM,
    TrainingArguments,
    BitsAndBytesConfig,
)
from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training
from trl import SFTTrainer


# ===================== 配置区 =====================

MODEL_NAME = "Qwen/Qwen2.5-7B-Instruct"  # 或 "meta-llama/Llama-3.1-8B-Instruct"
DATA_PATH = "data/split/train.json"
VAL_PATH = "data/split/val.json"
OUTPUT_DIR = "checkpoints/qwen2.5-7b-legal-qlora"

# QLoRA 量化配置
BNB_CONFIG = BitsAndBytesConfig(
    load_in_4bit=True,
    bnb_4bit_use_double_quant=True,
    bnb_4bit_quant_type="nf4",
    bnb_4bit_compute_dtype=torch.bfloat16,
)

# LoRA 配置
LORA_CONFIG = LoraConfig(
    r=64,
    lora_alpha=128,
    lora_dropout=0.05,
    target_modules=[
        "q_proj", "k_proj", "v_proj", "o_proj",
        "gate_proj", "up_proj", "down_proj",
    ],
    bias="none",
    task_type="CAUSAL_LM",
)

# 训练参数
TRAINING_ARGS = TrainingArguments(
    output_dir=OUTPUT_DIR,
    num_train_epochs=3,
    per_device_train_batch_size=4,
    per_device_eval_batch_size=4,
    gradient_accumulation_steps=8,  # 有效 batch_size = 4 × 8 × 4 卡 = 128
    learning_rate=2e-4,
    warmup_ratio=0.03,
    lr_scheduler_type="cosine",
    logging_steps=10,
    save_steps=200,
    eval_steps=200,
    save_total_limit=3,
    bf16=True,
    tf32=True,
    ddp_find_unused_parameters=False,
    gradient_checkpointing=True,
    evaluation_strategy="steps",
    load_best_model_at_end=True,
    metric_for_best_model="eval_loss",
    greater_is_better=False,
    deepspeed="configs/ds_zero2.json",
    report_to="wandb",
    run_name="legal-sft-qlora",
)


# ===================== 主流程 =====================

def format_conversation(example):
    """将 conversations 列表格式化为文本序列,仅计算 assistant 部分的 loss"""
    messages = example["conversations"]
    text_parts = []
    for msg in messages:
        role = msg["role"]
        content = msg["content"]
        if role == "system":
            text_parts.append(f"<|im_start|>system\n{content}<|im_end|>\n")
        elif role == "user":
            text_parts.append(f"<|im_start|>user\n{content}<|im_end|>\n")
        elif role == "assistant":
            text_parts.append(f"<|im_start|>assistant\n{content}<|im_end|>\n")
    return {"text": "".join(text_parts)}


def main():
    # 加载数据集
    dataset = load_dataset("json", data_files={
        "train": DATA_PATH,
        "validation": VAL_PATH,
    })
    dataset = dataset.map(format_conversation, remove_columns=dataset["train"].column_names)

    # 加载 tokenizer
    tokenizer = AutoTokenizer.from_pretrained(
        MODEL_NAME,
        trust_remote_code=True,
        padding_side="right",
    )
    if tokenizer.pad_token is None:
        tokenizer.pad_token = tokenizer.eos_token

    # 加载 4-bit 量化基座模型
    model = AutoModelForCausalLM.from_pretrained(
        MODEL_NAME,
        quantization_config=BNB_CONFIG,
        device_map={"": torch.cuda.current_device()},
        torch_dtype=torch.bfloat16,
        trust_remote_code=True,
        attn_implementation="flash_attention_2",
    )
    model = prepare_model_for_kbit_training(model)
    model = get_peft_model(model, LORA_CONFIG)
    model.config.use_cache = False  # 梯度检查点需要

    # 打印可训练参数量
    model.print_trainable_parameters()
    # 预期输出: trainable params: ~84M || all params: ~7.6B || trainable%: 1.1%

    # 创建 Trainer
    trainer = SFTTrainer(
        model=model,
        args=TRAINING_ARGS,
        train_dataset=dataset["train"],
        eval_dataset=dataset["validation"],
        tokenizer=tokenizer,
        dataset_text_field="text",
        max_seq_length=4096,
        packing=False,  # 关闭 packing 以保证每条样本独立
    )

    # 开始训练
    trainer.train()

    # 保存最终 adapter
    trainer.save_model(f"{OUTPUT_DIR}/final")
    tokenizer.save_pretrained(f"{OUTPUT_DIR}/final")


if __name__ == "__main__":
    main()

5.4 DeepSpeed ZeRO-2 配置

{
    "zero_optimization": {
        "stage": 2,
        "offload_optimizer": {
            "device": "none"
        },
        "overlap_comm": true,
        "contiguous_gradients": true,
        "reduce_bucket_size": 5e7,
        "allgather_bucket_size": 5e7
    },
    "bf16": {
        "enabled": true
    },
    "train_batch_size": "auto",
    "train_micro_batch_size_per_gpu": "auto",
    "gradient_accumulation_steps": "auto",
    "gradient_clipping": 1.0,
    "wall_clock_breakdown": false
}

六、训练过程与踩坑记录

6.1 启动命令

deepspeed --num_gpus=4 train_sft.py

6.2 训练监控

通过 Weights & Biases (wandb) 实时监控 loss 曲线:

外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传

关注的指标:

  • train/loss:应持续下降,不出现剧烈震荡;
  • eval/loss:若在某个 step 后持续上升,说明过拟合,需降低 epoch 或增大 dropout;
  • learning_rate:warmup + cosine decay 的曲线是否正常。

6.3 常见踩坑与解决方案

坑 1:OOM(显存溢出)
torch.cuda.OutOfMemoryError: CUDA out of memory.

解决思路(按优先级):

  1. 降低 per_device_train_batch_size(如 4 → 2 或 1);
  2. 增大 gradient_accumulation_steps 保持等效 batch size;
  3. 降低 max_seq_length(如 4096 → 2048);
  4. 启用 gradient_checkpointing=True
  5. 使用 ZeRO-3 替代 ZeRO-2(牺牲通信换显存)。
坑 2:Qwen 的 Chat Template 不匹配

Qwen 使用 <|im_start|> / <|im_end|> 标记,而非 Llama 的 [INST] 格式。格式化数据时必须与模型对齐:

# Qwen 格式
"<|im_start|>system\n{system}<|im_end|>\n<|im_start|>user\n{user}<|im_end|>\n<|im_start|>assistant\n{assistant}<|im_end|>\n"

# Llama-3 格式
"<|begin_of_text|><|start_header_id|>system<|end_header_id|>\n{system}<|eot_id|><|start_header_id|>user<|end_header_id|>\n{user}<|eot_id|><|start_header_id|>assistant<|end_header_id|>\n{assistant}<|eot_id|>"

推荐方式:直接使用 tokenizer 的 apply_chat_template 方法:

text = tokenizer.apply_chat_template(
    conversation=messages,
    tokenize=False,
    add_generation_prompt=False,
)
坑 3:packing=True 导致 loss 计算偏差

当开启 packing=True 时,SFTTrainer 会将多条短样本拼接成一条长序列,这可能导致 cross-contamination(不同样本之间的 attention 互相干扰)。对于对话数据,建议关闭 packing。

坑 4:Flash Attention 2 未生效的验证
from transformers.utils import is_flash_attn_2_available
print(is_flash_attn_2_available())  # 应为 True

# 更直接的验证方式:
import torch
dummy = torch.randn(1, 1, 4096, 64, dtype=torch.bfloat16, device="cuda")
# 若不报错且推理速度明显快于 SDPA,则 FA2 正常

七、模型评估与对比

7.1 自动化评估

使用测试集进行困惑度(Perplexity)评估:

import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
from tqdm import tqdm

def compute_perplexity(model, tokenizer, data_path, max_samples=500):
    with open(data_path, "r", encoding="utf-8") as f:
        data = json.load(f)[:max_samples]

    total_loss = 0.0
    total_tokens = 0
    model.eval()

    with torch.no_grad():
        for item in tqdm(data):
            text = format_conversation(item)  # 同上文的格式化函数
            inputs = tokenizer(
                text, return_tensors="pt", truncation=True, max_length=2048
            ).to(model.device)
            outputs = model(**inputs, labels=inputs["input_ids"])
            total_loss += outputs.loss.item() * inputs["input_ids"].numel()
            total_tokens += inputs["input_ids"].numel()

    ppl = torch.exp(torch.tensor(total_loss / total_tokens))
    return ppl.item()

print(f"Test Perplexity: {compute_perplexity(model, tokenizer, 'data/split/test.json'):.2f}")

7.2 人工评估:打分卡

邀请 3 名领域专家,对 100 道测试题的模型输出进行盲评(不知道来自哪个模型),维度如下:

维度权重评分标准
准确性40%事实正确、法条引用无误
完整性25%覆盖问题所有方面
逻辑性20%推理链条清晰
可读性15%语言通顺、格式规范

评分结果(5 分制):

模型准确性完整性逻辑性可读性加权总分
Qwen2.5-7B (原始)3.23.03.54.23.38
Qwen2.5-7B (微调后)4.34.14.44.54.31
Llama-3.1-8B (原始)2.82.93.64.03.14
Llama-3.1-8B (微调后)3.93.84.24.34.02

可以看到,微调后在准确性和完整性上有明显提升,Qwen 的中文底子使其在中文法律场景下整体优于 Llama。


八、模型合并、量化与导出

8.1 LoRA Adapter 与基座合并

import torch
from peft import PeftModel
from transformers import AutoModelForCausalLM, AutoTokenizer

BASE_MODEL = "Qwen/Qwen2.5-7B-Instruct"
ADAPTER_PATH = "checkpoints/qwen2.5-7b-legal-qlora/final"
MERGED_PATH = "models/qwen2.5-7b-legal-merged"

# 加载基座 + adapter
model = AutoModelForCausalLM.from_pretrained(
    BASE_MODEL,
    torch_dtype=torch.bfloat16,
    device_map="auto",
    trust_remote_code=True,
)
model = PeftModel.from_pretrained(model, ADAPTER_PATH)
model = model.merge_and_unload()  # 合并权重

tokenizer = AutoTokenizer.from_pretrained(ADAPTER_PATH, trust_remote_code=True)

model.save_pretrained(MERGED_PATH, safe_serialization=True)
tokenizer.save_pretrained(MERGED_PATH)
print(f"合并后的模型已保存至 {MERGED_PATH}")

8.2 GPTQ 4-bit 量化

量化可以大幅降低推理显存需求(7B 模型从 ~14GB 降到 ~4GB),适合部署场景:

from auto_gptq import AutoGPTQForCausalLM, BaseQuantizeConfig

quant_config = BaseQuantizeConfig(
    bits=4,
    group_size=128,
    damp_percent=0.01,
    desc_act=False,
    sym=True,
    true_sequential=True,
)

model = AutoGPTQForCausalLM.from_pretrained(
    MERGED_PATH,
    quant_config,
    torch_dtype=torch.float16,
)

# 使用校准数据集进行量化
model.quantize(calib_dataset)

model.save_quantized("models/qwen2.5-7b-legal-gptq-int4")

8.3 GGUF 格式导出(适配 llama.cpp / Ollama)

# 安装 llama.cpp
git clone https://github.com/ggerganov/llama.cpp
cd llama.cpp && make -j

# 转换为 GGUF
python convert_hf_to_gguf.py ../models/qwen2.5-7b-legal-merged \
    --outfile ../models/qwen2.5-7b-legal.Q8_0.gguf \
    --outtype q8_0

# 进一步量化到 Q4_K_M
./llama-quantize ../models/qwen2.5-7b-legal.Q8_0.gguf \
    ../models/qwen2.5-7b-legal.Q4_K_M.gguf \
    Q4_K_M

九、推理部署

9.1 vLLM 高性能推理

vLLM 凭借 PagedAttention 实现高吞吐推理,是生产环境推荐方案:

# 启动 OpenAI 兼容 API 服务
python -m vllm.entrypoints.openai.api_server \
    --model models/qwen2.5-7b-legal-merged \
    --served-model-name legal-assistant \
    --dtype bfloat16 \
    --max-model-len 8192 \
    --gpu-memory-utilization 0.92 \
    --tensor-parallel-size 1 \
    --port 8000

客户端调用示例:

from openai import OpenAI

client = OpenAI(
    base_url="http://localhost:8000/v1",
    api_key="not-needed",
)

response = client.chat.completions.create(
    model="legal-assistant",
    messages=[
        {"role": "system", "content": "你是一名专业法律顾问。"},
        {"role": "user", "content": "公司拖欠工资怎么办?"},
    ],
    temperature=0.7,
    max_tokens=1024,
)

print(response.choices[0].message.content)

9.2 Ollama 本地部署

适合个人开发者在本地快速体验:

# 创建 Modelfile
cat > Modelfile << 'EOF'
FROM ./models/qwen2.5-7b-legal.Q4_K_M.gguf
TEMPLATE """<|im_start|>system
{{ .System }}<|im_end|>
<|im_start|>user
{{ .Prompt }}<|im_end|>
<|im_start|>assistant
"""
PARAMETER temperature 0.7
PARAMETER top_p 0.9
PARAMETER stop "<|im_end|>"
EOF

# 创建并运行
ollama create legal-assistant -f Modelfile
ollama run legal-assistant

9.3 性能基准(vLLM,A100 80G)

指标数值
吞吐量~2,400 tokens/s
首 Token 延迟 (TTFT)~180 ms
并发请求数32
平均延迟~420 ms
显存占用~28 GB (bf16) / ~8 GB (int4)

十、总结与展望

核心要点回顾

  1. 数据质量 > 数据数量:高质量、多样化的指令数据是微调效果的核心决定因素。5,000 条精标数据胜过 50,000 条噪声数据。
  2. QLoRA 是性价比最优解:在几乎不损失效果的前提下,将显存需求从 ~112GB 压缩至 ~16GB,使消费级 GPU 也能参与微调。
  3. 基座模型选择直接影响上限:中文场景选 Qwen,英文场景选 Llama,不要试图在非母语基座上强行微调母语任务。
  4. 部署链路要提前规划:微调完成后尽快合并 + 量化,GGUF 格式最便于分发,vLLM 最适合高并发服务。

进阶方向

  • DPO / RLHF 对齐:在 SFT 基础上进一步做偏好对齐,提升安全性和输出质量;
  • RAG 增强:将领域知识库与微调模型结合,实现知识可更新、可追溯;
  • Agent 化:结合 Function Calling,让模型不仅能回答法律问题,还能调用法条检索、案例匹配等工具;
  • 持续预训练(CPT):在大规模领域无监督语料上进行继续预训练,注入更深层的领域知识。

资源清单


本文是一篇实战复盘,所有代码均在所述环境中验证通过。欢迎交流与指正。

更多推荐