【实战】使用 Unsloth 框架微调 Qwen3-4B 医疗大模型
在当今的大模型(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 行),更重要的是让消费级显卡也能跑得动大模型微调任务。
关键点回顾:
- 4-bit 量化:降低显存门槛。
- LoRA:只训练 <1% 参数,高效且效果好。
- 数据清洗:真实世界的数据往往是杂乱的,多种编码检测和列名匹配是必要的。
- 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
更多推荐

所有评论(0)