在当今的大模型(LLM)应用开发中,**微调(Fine-tuning)**是将通用模型转化为垂直领域专家的关键步骤。

本文将带你深入一段基于 Unsloth 框架的代码,该框架以其极高的训练速度和极低的显存占用而闻名。我们将演示如何在单张显卡(如 Tesla T4 或 RTX 30/40系列)上,利用 LoRA 技术对 Qwen3-4B(注:此处以用户代码为准,实际目前主流为 Qwen2.5)进行医疗数据的微调。

📊 整体运行流程图

在开始代码之前,让我们通过一张流程图来理解整个微调的生命周期:

在这里插入图片描述


📝 代码详解与详细注释

以下是完整的代码解析。为了方便阅读,我将代码拆分为功能模块,并添加了详细的“逐行级”注释。

1. 环境初始化与模型加载

这一步利用 Unsloth 的 FastLanguageModel 快速加载 4-bit 量化的基础模型,这是节省显存的关键。

# 导入必要的库
# Unsloth 是一个针对 LLM 训练优化的库,速度比 huggingface 快 2-5 倍,显存减少 80%
from unsloth import FastLanguageModel
import torch

# === 模型参数配置 ===
max_seq_length = 2048  # 设置最大上下文长度。Unsloth 支持 RoPE 自动缩放,不仅限于 2048
dtype = None  # 数据精度设置。None = 自动检测 (T4用Float16, Ampere架构如A100/3090用Bfloat16)
load_in_4bit = True  # 核心开关:启用 4bit 量化加载,大幅降低显存需求 (4B模型仅需约3-4G显存)

# === 加载预训练模型 ===
# path: 可以是 HuggingFace Hub ID 或本地路径
model, tokenizer = FastLanguageModel.from_pretrained(
    model_name="/root/autodl-tmp/models/Qwen/Qwen3-4B", # 指定本地模型路径
    max_seq_length=max_seq_length,
    dtype=dtype,
    load_in_4bit=load_in_4bit,
)

2. 配置 LoRA 适配器 (PEFT)

我们不直接训练整个模型(全量微调),而是使用 LoRA (Low-Rank Adaptation) 技术,冻结主模型,只训练极少量的附加参数。

# === LoRA 参数配置 ===
model = FastLanguageModel.get_peft_model(
    model,
    r=16,  # LoRA 的秩 (Rank)。数值越大,可学习参数越多,但也更易过拟合。常用 8, 16, 32
    target_modules=["q_proj", "k_proj", "v_proj", "o_proj",
                    "gate_proj", "up_proj", "down_proj"],  # 指定由于哪些层应用 LoRA。全量指定效果通常最好
    lora_alpha=16,  # LoRA 的缩放系数,通常设置为 r 的 1 倍或 2 倍
    lora_dropout=0,  # 为了优化训练速度,建议设为 0
    bias="none",    # 是否训练偏置项,none 是最节省显存的设置
    use_gradient_checkpointing="unsloth",  # 显存优化黑科技,使用 unsloth 特有的长序列优化
    random_state=3407,  # 这里的 3407 是一个著名的"幸运"随机种子 (arxiv:2109.08203)
    use_rslora=False,  # Rank Stabilized LoRA,一种改进的 LoRA 变体
    loftq_config=None, # LoftQ 初始化配置,通常不用
)

3. 数据准备与清洗

这一部分是实际工程中最繁琐的环节。代码包含了一个鲁棒的 CSV 读取器,能够处理医疗数据中常见的编码问题(GBK/UTF-8)和列名不统一问题。

import os
import pandas as pd
from datasets import Dataset

# === 定义提示词模板 ===
# 这是一个标准的 Alpaca 风格模板,让模型学会"指令-输入-输出"的格式
medical_prompt = """你是一个专业的医疗助手。请根据患者的问题提供专业、准确的回答。

### 问题:
{}

### 回答:
{}"""

# 获取模型特定的结束符 (EOS),这对于模型知道何时停止生成至关重要
EOS_TOKEN = tokenizer.eos_token

def read_csv_with_encoding(file_path):
    """
    工具函数:解决中文 CSV 文件常见的编码报错问题。
    依次尝试 gbk, utf-8 等常见编码。
    """
    encodings = ['gbk', 'gb2312', 'gb18030', 'utf-8']
    for encoding in encodings:
        try:
            return pd.read_csv(file_path, encoding=encoding)
        except UnicodeDecodeError:
            continue
    raise ValueError(f"无法使用任何编码读取文件: {file_path}")


def load_medical_data(data_dir):
    """
    核心数据加载函数:遍历指定目录下的所有科室数据
    """
    data = []
    # 映射目录名到中文科室名
    departments = {
        'IM_内科': '内科',
        'Surgical_外科': '外科',
        'Pediatric_儿科': '儿科',
        'Oncology_肿瘤科': '肿瘤科',
        'OAGD_妇产科': '妇产科',
        'Andriatria_男科': '男科'
    }

    # 遍历科室目录
    for dept_dir, dept_name in departments.items():
        dept_path = os.path.join(data_dir, dept_dir)
        if not os.path.exists(dept_path):
            print(f"目录不存在: {dept_path}")
            continue

        print(f"\n处理{dept_name}数据...")
        csv_files = [f for f in os.listdir(dept_path) if f.endswith('.csv')]

        for csv_file in csv_files:
            file_path = os.path.join(dept_path, csv_file)
            print(f"正在处理文件: {csv_file}")

            try:
                df = read_csv_with_encoding(file_path)
                
                # 逐行处理
                for _, row in df.iterrows():
                    try:
                        # === 模糊匹配列名 ===
                        # 数据集来源可能不同,列名可能是 question/ask/问题 等
                        question = None
                        answer = None

                        if 'question' in row: question = str(row['question']).strip()
                        elif '问题' in row: question = str(row['问题']).strip()
                        elif 'ask' in row: question = str(row['ask']).strip()

                        if 'answer' in row: answer = str(row['answer']).strip()
                        elif '回答' in row: answer = str(row['回答']).strip()
                        elif 'response' in row: answer = str(row['response']).strip()

                        # 数据校验:跳过空数据或过长的数据(过长可能导致 OOM 或包含噪音)
                        if not question or not answer: continue
                        if len(question) > 200 or len(answer) > 200: continue

                        # 构建训练样本
                        data.append({
                            "instruction": "请回答以下医疗相关问题",
                            "input": question,
                            "output": answer
                        })

                    except Exception as e:
                        continue # 跳过坏行

            except Exception as e:
                print(f"处理文件出错: {e}")
                continue

    print(f"\n成功处理 {len(data)} 条数据")
    # 将 Python 列表转换为 HuggingFace Dataset 对象
    return Dataset.from_list(data)


def formatting_prompts_func(examples):
    """
    数据预处理函数:将 instruction, input, output 拼装成最终的训练文本。
    并加上 EOS_TOKEN,否则模型可能会无限生成。
    """
    instructions = examples["instruction"]
    inputs = examples["input"]
    outputs = examples["output"]
    texts = []
    for instruction, input, output in zip(instructions, inputs, outputs):
        # 格式化文本 + 添加结束符
        text = medical_prompt.format(input, output) + EOS_TOKEN
        texts.append(text)
    return {"text": texts}


# 执行数据加载流水线
dataset = load_medical_data("Data_数据")
dataset = dataset.map(formatting_prompts_func, batched=True)

4. 训练器配置 (SFTTrainer)

这里使用 SFTTrainer (Supervised Fine-tuning Trainer)。这是 transformers 库的高级封装,专门用于指令微调。

from trl import SFTTrainer
from transformers import TrainingArguments
from unsloth import is_bfloat16_supported

# === 训练参数详解 ===
training_args = TrainingArguments(
    per_device_train_batch_size=2,  # 显存如果只有 16G,建议设为 2 或 4
    gradient_accumulation_steps=4,  # 梯度累积。实际 batch_size = 2 * 4 = 8。用于模拟大 Batch 训练
    warmup_steps=5,      # 预热步数,训练开始时缓慢增加学习率,防止梯度爆炸
    max_steps=-1,        # 设置为 -1 表示按 epoch 训练,不按步数限制
    num_train_epochs=3,  # 遍历数据集 3 次
    learning_rate=2e-4,  # LoRA 常用学习率。全量微调通常用 1e-5
    fp16=not is_bfloat16_supported(),  # 旧显卡(T4)用 FP16
    bf16=is_bfloat16_supported(),      # 新显卡(A100/3090)用 BF16,更稳定
    logging_steps=1,     # 每一步都打印日志,方便观察 Loss
    optim="adamw_8bit",  # 使用 8-bit AdamW 优化器,进一步节省显存
    weight_decay=0.01,   # 权重衰减,防止过拟合
    lr_scheduler_type="linear", # 学习率随时间线性下降
    seed=3407,
    output_dir="outputs", # 模型检查点保存路径
    report_to="none",     # 不上传到 WandB 等平台
)

# === 初始化训练器 ===
trainer = SFTTrainer(
    model=model,
    tokenizer=tokenizer,
    train_dataset=dataset,
    dataset_text_field="text", # 指定数据集中哪一列是训练文本
    max_seq_length=max_seq_length,
    dataset_num_proc=2,    # 数据预处理的进程数
    packing=False,         # 是否将多个短序列打包成一个长序列(设为 True 可加速,但需注意截断)
    args=training_args,
)

5. 执行训练与监控

# === 打印训练前显存状态 ===
gpu_stats = torch.cuda.get_device_properties(0)
start_gpu_memory = round(torch.cuda.max_memory_reserved() / 1024 / 1024 / 1024, 3)
max_memory = round(gpu_stats.total_memory / 1024 / 1024 / 1024, 3)
print(f"GPU: {gpu_stats.name}. 总显存: {max_memory} GB. 已占用: {start_gpu_memory} GB.")

# === 开始训练 ===
trainer_stats = trainer.train()

# === 打印训练统计信息 ===
# 计算 LoRA 训练额外占用的显存
used_memory = round(torch.cuda.max_memory_reserved() / 1024 / 1024 / 1024, 3)
used_memory_for_lora = round(used_memory - start_gpu_memory, 3)
print(f"训练耗时: {trainer_stats.metrics['train_runtime']} 秒")
print(f"LoRA 训练占用显存: {used_memory_for_lora} GB")

6. 模型推理 (Inference)

训练完成后,我们需要验证效果。这里使用 TextStreamer 实现打字机式的流式输出。

# === 定义推理函数 ===
def generate_medical_response(question):
    """
    输入问题,生成回答
    """
    # 切换到推理模式,Unsloth 会优化推理速度 (2x)
    FastLanguageModel.for_inference(model)
    
    # 构建输入并转为 Tensor
    inputs = tokenizer(
        [medical_prompt.format(question, "")],
        return_tensors="pt"
    ).to("cuda")

    # 使用流式输出器,看到字一个个蹦出来
    from transformers import TextStreamer
    text_streamer = TextStreamer(tokenizer)
    
    # 生成回答
    _ = model.generate(
        **inputs,
        streamer=text_streamer,
        max_new_tokens=256,    # 最大生成长度
        temperature=0.7,       # 随机性控制 (0.7 比较平衡)
        top_p=0.9,
        repetition_penalty=1.1 # 惩罚重复内容
    )

# === 测试 ===
test_questions = ["我最近总是感觉头晕,应该怎么办?", "感冒发烧应该吃什么药?"]
for q in test_questions:
    print(f"\n问题:{q}")
    generate_medical_response(q)

7. 模型保存与加载

最后,我们将训练好的 LoRA 权重保存下来,并演示如何重新加载。注意:我们保存的只是几百 MB 的 Adapter,而不是整个 4GB 的模型。

# === 保存 LoRA 权重 ===
# 保存到本地文件夹 'lora_model_medical'
model.save_pretrained("lora_model_medical")
tokenizer.save_pretrained("lora_model_medical")

# === 加载方式 1: 直接加载刚保存的 LoRA ===
if True:
    from unsloth import FastLanguageModel
    model, tokenizer = FastLanguageModel.from_pretrained(
        model_name="lora_model_medical", # 直接指向保存的目录
        max_seq_length=max_seq_length,
        dtype=dtype,
        load_in_4bit=load_in_4bit,
    )
    FastLanguageModel.for_inference(model)
    
    print("加载 LoRA 模型测试:")
    generate_medical_response("我最近总是感觉头晕,应该怎么办?")

# === 加载方式 2: 基础模型 + Adapter (更通用的方式) ===
if True:
    model, tokenizer = FastLanguageModel.from_pretrained(
        model_name="/root/autodl-tmp/models/Qwen/Qwen3-4B", # 基础模型路径
        adapter_name="lora_model_medical",                  # 挂载 LoRA 权重
        max_seq_length=max_seq_length,
        dtype=dtype,
        load_in_4bit=load_in_4bit,
    )
    # ... 后续推理代码相同

💡 总结

通过上述代码,我们完成了一个完整的医疗垂直领域大模型微调流程。Unsloth 框架不仅让代码变得简洁(核心逻辑不到 50 行),更重要的是让消费级显卡也能跑得动大模型微调任务。

关键点回顾:

  1. 4-bit 量化:降低显存门槛。
  2. LoRA:只训练 <1% 参数,高效且效果好。
  3. 数据清洗:真实世界的数据往往是杂乱的,多种编码检测和列名匹配是必要的。
  4. Prompt 模板:格式化输入是让模型听懂指令的关键。

完整代码:

#!/usr/bin/env python
# coding: utf-8

"""
Unsloth Qwen 医疗大模型微调脚本
功能:加载 Qwen 模型,读取医疗 CSV 数据,进行 LoRA 微调,保存并测试。
环境需求:Unsloth, PyTorch, Transformers, TRL, Pandas, Datasets
"""

import os
import torch
import pandas as pd
from datasets import Dataset
from unsloth import FastLanguageModel, is_bfloat16_supported
from trl import SFTTrainer
from transformers import TrainingArguments, TextStreamer

# ===========================
# 1. 全局配置与参数
# ===========================

# 模型路径 (请根据实际情况修改,支持本地路径或 HuggingFace ID)
# 注意:如果您的环境没有 Qwen3,可改为 "unsloth/Qwen2.5-7B-Instruct-bnb-4bit"
MODEL_NAME = "/root/autodl-tmp/models/Qwen/Qwen3-4B"

# 数据集目录
DATA_DIR = "Data_数据"

# 模型保存名称
OUTPUT_MODEL_NAME = "lora_model_medical"

# 训练参数
MAX_SEQ_LENGTH = 2048   # 最大序列长度
LOAD_IN_4BIT = True     # 开启 4bit 量化 (省显存关键)
DTYPE = None            # None 为自动检测 (T4: Float16, Ampere: Bfloat16)
SEED = 3407             # 随机种子

# 提示词模板
MEDICAL_PROMPT = """你是一个专业的医疗助手。请根据患者的问题提供专业、准确的回答。

### 问题:
{}

### 回答:
{}"""

# ===========================
# 2. 数据处理工具函数
# ===========================

def read_csv_with_encoding(file_path):
    """尝试使用不同的编码读取 CSV 文件,解决中文乱码问题"""
    encodings = ['gbk', 'gb2312', 'gb18030', 'utf-8']
    for encoding in encodings:
        try:
            return pd.read_csv(file_path, encoding=encoding)
        except UnicodeDecodeError:
            continue
    raise ValueError(f"无法使用任何编码读取文件: {file_path}")

def load_medical_data(data_dir):
    """
    遍历指定目录下的科室文件夹,加载并清洗 CSV 数据
    """
    data = []
    # 科室目录映射
    departments = {
        'IM_内科': '内科',
        'Surgical_外科': '外科',
        'Pediatric_儿科': '儿科',
        'Oncology_肿瘤科': '肿瘤科',
        'OAGD_妇产科': '妇产科',
        'Andriatria_男科': '男科'
    }

    print(f"开始加载数据,路径: {data_dir}")

    # 遍历所有科室目录
    for dept_dir, dept_name in departments.items():
        dept_path = os.path.join(data_dir, dept_dir)
        if not os.path.exists(dept_path):
            print(f"警告: 目录不存在 {dept_path},跳过。")
            continue

        print(f"正在处理 {dept_name} 数据...")
        csv_files = [f for f in os.listdir(dept_path) if f.endswith('.csv')]

        for csv_file in csv_files:
            file_path = os.path.join(dept_path, csv_file)
            try:
                df = read_csv_with_encoding(file_path)
                
                # 逐行处理
                for _, row in df.iterrows():
                    try:
                        # 尝试不同的列名匹配
                        question = None
                        answer = None

                        # 匹配问题列
                        if 'question' in row: question = str(row['question']).strip()
                        elif '问题' in row: question = str(row['问题']).strip()
                        elif 'ask' in row: question = str(row['ask']).strip()

                        # 匹配回答列
                        if 'answer' in row: answer = str(row['answer']).strip()
                        elif '回答' in row: answer = str(row['回答']).strip()
                        elif 'response' in row: answer = str(row['response']).strip()

                        # 数据校验与过滤
                        if not question or not answer: continue
                        if len(question) > 200 or len(answer) > 200: continue

                        data.append({
                            "instruction": "请回答以下医疗相关问题",
                            "input": question,
                            "output": answer
                        })

                    except Exception:
                        continue # 跳过单行错误

            except Exception as e:
                print(f"读取文件 {csv_file} 失败: {e}")
                continue

    if not data:
        raise ValueError("未加载到任何有效数据,请检查数据路径!")

    print(f"数据加载完成,共 {len(data)} 条样本。")
    return Dataset.from_list(data)

def format_prompts(examples, tokenizer):
    """将数据格式化为训练所需的 Prompt 格式"""
    EOS_TOKEN = tokenizer.eos_token
    instructions = examples["instruction"]
    inputs = examples["i

相关资源:
百度网盘:https://pan.baidu.com/s/14TeT6bC8Wp93o6IL-UeAWw?pwd=rqxx
在这里插入图片描述

更多推荐