大模型微调数据工程:从数据采集到质量评估的完整流程
2026/7/3大约 19 分钟
大模型微调数据工程:从数据采集到质量评估的完整流程
"Garbage In, Garbage Out。微调效果 80% 取决于数据质量,20% 取决于训练方法。"
引言:为什么数据工程是微调的核心?
很多人把精力花在调学习率、选 LoRA rank、试不同优化器上,但忽略了最重要的事情:数据质量。
实践中,数据质量对微调效果的影响远超算法选择:
┌─────────────────────────────────────────────────┐
│ 微调效果影响因素分析 │
│ │
│ 数据质量 ████████████████████░░ 80% │
│ 数据配比 ████████░░░░░░░░░░░░░░ 15% │
│ 训练方法 ████░░░░░░░░░░░░░░░░░░ 3% │
│ 超参调优 ██░░░░░░░░░░░░░░░░░░░░ 2% │
│ │
│ 结论:把时间花在数据上! │
└─────────────────────────────────────────────────┘一、微调数据来源
1.1 三大数据来源
┌──────────────────────────────────────────────────────┐
│ 微调数据来源分类 │
├──────────────┬───────────────────────────────────────┤
│ 公开数据 │ Alpaca, ShareGPT, FLAN, OpenOrca │
│ │ 优点:量大、免费 │
│ │ 缺点:质量参差不齐、可能含噪声 │
├──────────────┼───────────────────────────────────────┤
│ 合成数据 │ 用 GPT-4/Claude 生成训练数据 │
│ │ 优点:质量可控、可定制 │
│ │ 缺点:成本高、可能有模型偏差 │
├──────────────┼───────────────────────────────────────┤
│ 业务数据 │ 真实用户日志、客服记录、专家标注 │
│ │ 优点:最贴近实际场景 │
│ │ 缺点:需要清洗脱敏、量少 │
└──────────────┴───────────────────────────────────────┘1.2 公开数据集
"""
常用公开微调数据集加载
"""
from datasets import load_dataset
# 1. Alpaca 格式数据集
alpaca_ds = load_dataset("tatsu-lab/alpaca", split="train")
# 字段: instruction, input, output
# 2. ShareGPT 格式数据集
sharegpt_ds = load_dataset("RyokoAI/ShareGPT52K", split="train")
# 字段: conversations (多轮对话)
# 3. FLAN 集合(多任务)
flan_ds = load_dataset("Open-Orca/FLAN", split="train")
# 4. OpenOrca(GPT-4 蒸馏数据)
orca_ds = load_dataset("Open-Orca/OpenOrca", split="train")
# 5. UltraChat(多轮对话)
ultra_ds = load_dataset("stingning/ultrachat", split="train")
print(f"Alpaca: {len(alpaca_ds)} 条")
print(f"ShareGPT: {len(sharegpt_ds)} 条")
print(f"FLAN: {len(flan_ds)} 条")
print(f"OpenOrca: {len(orca_ds)} 条")
print(f"UltraChat: {len(ultra_ds)} 条")1.3 合成数据生成
"""
使用大模型生成合成训练数据
关键:质量控制 + 多样性
"""
import json
import random
from openai import OpenAI
client = OpenAI()
# ── 合成数据生成器 ──
SYNTHESIS_PROMPT = """你是一个训练数据生成专家。请生成高质量的指令-回答对。
主题领域:{domain}
难度级别:{difficulty}
输出语言:{language}
要求:
1. 指令要具体、多样化,不要模板化
2. 回答要准确、详细、有结构
3. 避免有害内容
4. 每条数据独立完整
请生成 {n} 条数据,JSON 数组格式:
[
{{
"instruction": "具体的指令",
"input": "可选的输入内容(可为空)",
"output": "高质量的回答"
}}
]
领域范围参考:{domain_hints}
"""
DOMAIN_CONFIGS = {
"编程": {
"hints": "Python/JavaScript/Go、数据结构、算法、系统设计、Debug、代码审查、最佳实践",
"difficulty_levels": ["初级", "中级", "高级", "专家"]
},
"写作": {
"hints": "邮件、报告、文案、小说、技术文档、营销文案",
"difficulty_levels": ["简单", "中等", "复杂"]
},
"推理": {
"hints": "逻辑推理、数学题、脑筋急转弯、案例分析、决策",
"difficulty_levels": ["简单", "中等", "困难"]
},
}
def generate_synthetic_data(
domain: str,
n: int = 10,
difficulty: str = "中等",
language: str = "中文"
):
"""生成合成数据"""
config = DOMAIN_CONFIGS.get(domain, DOMAIN_CONFIGS["编程"])
response = client.chat.completions.create(
model="gpt-4o",
messages=[{
"role": "user",
"content": SYNTHESIS_PROMPT.format(
domain=domain,
difficulty=difficulty,
language=language,
n=n,
domain_hints=config["hints"]
)
}],
response_format={"type": "json_object"},
temperature=0.9, # 高温度增加多样性
)
data = json.loads(response.choices[0].message.content)
return data if isinstance(data, list) else data.get("data", [])
# ── 质量过滤 ──
def quality_filter(data_item: dict) -> bool:
"""过滤低质量合成数据"""
instruction = data_item.get("instruction", "")
output = data_item.get("output", "")
# 规则 1: 长度检查
if len(instruction) < 5 or len(output) < 20:
return False
# 规则 2: 拒绝模板化输出
template_phrases = ["作为一个AI", "我是一个语言模型", "我无法"]
if any(phrase in output for phrase in template_phrases):
return False
# 规则 3: 指令不能太相似(简化版)
if instruction.count("?") > 3: # 过多问号
return False
# 规则 4: 回答不能只是重复指令
if output.strip() == instruction.strip():
return False
return True
# ── 去重 ──
def deduplicate(data_list: list, key: str = "instruction") -> list:
"""基于指令文本去重"""
seen = set()
result = []
for item in data_list:
text = item.get(key, "").strip().lower()
if text not in seen:
seen.add(text)
result.append(item)
return result
# ── 批量生成 ──
def batch_generate(domains: list, samples_per_domain: int = 100):
"""批量生成多领域数据"""
all_data = []
for domain in domains:
config = DOMAIN_CONFIGS.get(domain, DOMAIN_CONFIGS["编程"])
for difficulty in config["difficulty_levels"]:
batch = generate_synthetic_data(domain, samples_per_domain, difficulty)
# 质量过滤
batch = [d for d in batch if quality_filter(d)]
all_data.extend(batch)
print(f"{domain}/{difficulty}: 生成 {len(batch)} 条")
# 全局去重
all_data = deduplicate(all_data)
print(f"\n总计: {len(all_data)} 条(去重后)")
return all_data1.4 业务数据采集
"""
从业务日志中提取训练数据
"""
import re
from datetime import datetime
class BusinessDataExtractor:
"""从业务日志提取训练数据"""
def __init__(self):
self.sensitive_patterns = [
(r'\b\d{3}-\d{3}-\d{4}\b', '[PHONE]'), # 电话号码
(r'\b[\w.+-]+@[\w-]+\.[\w.-]+\b', '[EMAIL]'), # 邮箱
(r'\b\d{15,18}\b', '[ID_CARD]'), # 身份证
(r'\b\d{16,19}\b', '[BANK_CARD]'), # 银行卡
(r'https?://\S+', '[URL]'), # URL
]
def desensitize(self, text: str) -> str:
"""脱敏处理"""
for pattern, replacement in self.sensitive_patterns:
text = re.sub(pattern, replacement, text)
return text
def extract_from_chat_log(self, log_entry: dict) -> dict:
"""从客服对话日志提取训练数据"""
# 原始日志格式:
# {"user_id": "xxx", "messages": [{"role": "user", "content": "..."}, ...]}
messages = log_entry.get("messages", [])
if len(messages) < 2:
return None
# 提取第一轮问答
user_msg = None
agent_msg = None
for msg in messages:
if msg["role"] == "user" and not user_msg:
user_msg = msg["content"]
elif msg["role"] == "assistant" and user_msg and not agent_msg:
agent_msg = msg["content"]
break
if not user_msg or not agent_msg:
return None
# 脱敏
user_msg = self.desensitize(user_msg)
agent_msg = self.desensitize(agent_msg)
# 质量检查
if len(user_msg) < 5 or len(agent_msg) < 20:
return None
return {
"instruction": user_msg,
"input": "",
"output": agent_msg,
"source": "business_log",
"timestamp": datetime.now().isoformat()
}
def extract_from_qa_pairs(self, qa_logs: list) -> list:
"""从 FAQ/知识库中提取"""
training_data = []
for log in qa_logs:
item = self.extract_from_chat_log(log)
if item:
training_data.append(item)
return training_data二、数据清洗
2.1 清洗管线
┌──────────────────────────────────────────────────┐
│ 数据清洗管线 │
│ │
│ 原始数据 │
│ │ │
│ ▼ │
│ ① 格式校验 ──── 去除格式错误的数据 │
│ │ │
│ ▼ │
│ ② 去重 ──────── 精确去重 + 模糊去重 │
│ │ │
│ ▼ │
│ ③ 长度过滤 ──── 过滤过短/过长的数据 │
│ │ │
│ ▼ │
│ ④ 质量过滤 ──── 规则 + 模型双重过滤 │
│ │ │
│ ▼ │
│ ⑤ 脱敏处理 ──── 去除个人信息 │
│ │ │
│ ▼ │
│ ⑥毒性检测 ───── 去除有害内容 │
│ │ │
│ ▼ │
│ ⑦ 语言检测 ──── 确保语言一致性 │
│ │ │
│ ▼ │
│ 清洗后数据 │
└──────────────────────────────────────────────────┘2.2 完整清洗实现
"""
数据清洗完整实现
"""
import re
import hashlib
from typing import List, Dict
from collections import Counter
class DataCleaner:
def __init__(self):
self.stats = {
"total": 0,
"format_error": 0,
"duplicated": 0,
"too_short": 0,
"too_long": 0,
"low_quality": 0,
"sensitive": 0,
"toxic": 0,
"passed": 0,
}
def clean(self, data: List[Dict]) -> List[Dict]:
"""完整清洗流程"""
self.stats["total"] = len(data)
result = []
# 1. 格式校验
data = self._validate_format(data)
# 2. 精确去重
data = self._exact_dedup(data)
# 3. 模糊去重
data = self._fuzzy_dedup(data)
# 4. 长度过滤
data = self._length_filter(data)
# 5. 质量过滤
data = self._quality_filter(data)
# 6. 脱敏
data = self._desensitize(data)
# 7. 毒性检测
data = self._toxicity_filter(data)
self.stats["passed"] = len(data)
result = data
return result
def _validate_format(self, data: List[Dict]) -> List[Dict]:
"""格式校验"""
result = []
for item in data:
if not isinstance(item, dict):
self.stats["format_error"] += 1
continue
if "instruction" not in item or "output" not in item:
self.stats["format_error"] += 1
continue
if not isinstance(item["instruction"], str) or not isinstance(item["output"], str):
self.stats["format_error"] += 1
continue
result.append(item)
return result
def _exact_dedup(self, data: List[Dict]) -> List[Dict]:
"""精确去重(基于 hash)"""
seen = set()
result = []
for item in data:
# 对 instruction 做hash
h = hashlib.md5(item["instruction"].strip().encode()).hexdigest()
if h in seen:
self.stats["duplicated"] += 1
continue
seen.add(h)
result.append(item)
return result
def _fuzzy_dedup(self, data: List[Dict], threshold: float = 0.85) -> List[Dict]:
"""模糊去重(基于 Jaccard 相似度)"""
def jaccard_similarity(s1: str, s2: str) -> float:
words1 = set(s1.lower().split())
words2 = set(s2.lower().split())
if not words1 or not words2:
return 0
intersection = words1 & words2
union = words1 | words2
return len(intersection) / len(union)
# 简化版:分桶后桶内比较
result = []
for i, item in enumerate(data):
is_dup = False
for j in range(max(0, len(result) - 100), len(result)): # 只和最近 100 条比较
if jaccard_similarity(item["instruction"], result[j]["instruction"]) > threshold:
is_dup = True
self.stats["duplicated"] += 1
break
if not is_dup:
result.append(item)
return result
def _length_filter(self, data: List[Dict]) -> List[Dict]:
"""长度过滤"""
result = []
for item in data:
inst_len = len(item["instruction"])
out_len = len(item["output"])
# 太短
if inst_len < 5 or out_len < 10:
self.stats["too_short"] += 1
continue
# 太长
if inst_len > 4000 or out_len > 8000:
self.stats["too_long"] += 1
continue
result.append(item)
return result
def _quality_filter(self, data: List[Dict]) -> List[Dict]:
"""质量过滤"""
# 低质量模式
low_quality_patterns = [
r"^(好的|是的|对的|嗯|好的呢)$", # 过于简单的回答
r"我不知道", # 无效回答
r"请稍等", # 客服套话
r"^(.{1,5})\1{5,}", # 重复字符
]
# 指令中的低质量模式
bad_instruction_patterns = [
r"^test", # 测试数据
r"^\d+$", # 纯数字
r"^[a-z]{1,3}$", # 过短英文
]
result = []
for item in data:
output = item["output"].strip()
instruction = item["instruction"].strip()
# 检查低质量回答
if any(re.match(p, output, re.IGNORECASE) for p in low_quality_patterns):
self.stats["low_quality"] += 1
continue
# 检查低质量指令
if any(re.match(p, instruction, re.IGNORECASE) for p in bad_instruction_patterns):
self.stats["low_quality"] += 1
continue
# 检查重复内容
if output.count(instruction) > 0 and len(instruction) > 50:
self.stats["low_quality"] += 1
continue
result.append(item)
return result
def _desensitize(self, data: List[Dict]) -> List[Dict]:
"""脱敏处理"""
sensitive_patterns = [
(r'\b1[3-9]\d{9}\b', '[PHONE]'), # 手机号
(r'\b[\w.+-]+@[\w-]+\.[\w.-]+\b', '[EMAIL]'), # 邮箱
(r'\b\d{15,18}[Xx]?\b', '[ID_CARD]'), # 身份证
(r'\b\d{16,19}\b', '[BANK_CARD]'), # 银行卡号
(r'密码[是为::]\s*\S+', '密码: [REDACTED]'), # 密码
]
result = []
for item in data:
for key in ["instruction", "input", "output"]:
if key in item:
for pattern, replacement in sensitive_patterns:
item[key] = re.sub(pattern, replacement, item[key])
result.append(item)
return result
def _toxicity_filter(self, data: List[Dict]) -> List[Dict]:
"""毒性检测(简化版)"""
toxic_keywords = [
"炸弹", "毒品", "自杀", "杀人", "恐怖袭击",
"racial slur equivalents in Chinese",
]
result = []
for item in data:
text = item.get("instruction", "") + item.get("output", "")
if any(kw in text for kw in toxic_keywords):
self.stats["toxic"] += 1
continue
result.append(item)
return result
def print_stats(self):
"""打印清洗统计"""
print("\n=== 数据清洗统计 ===")
for key, value in self.stats.items():
if key == "total":
print(f" 总计: {value}")
elif key == "passed":
print(f" 通过: {value} ({value/self.stats['total']*100:.1f}%)")
else:
print(f" {key}: {value}")三、SFT 数据格式
3.1 常见格式对比
┌──────────────────────────────────────────────────────────────┐
│ SFT 数据格式对比 │
├──────────────┬───────────────────────────────────────────────┤
│ Alpaca │ instruction + input + output │
│ │ 适合单轮指令任务 │
├──────────────┼───────────────────────────────────────────────┤
│ ShareGPT │ conversations: [{role, content}, ...] │
│ │ 适合多轮对话 │
├──────────────┼───────────────────────────────────────────────┤
│ ChatML │ <|im_start|>system\n...\n<|im_end|> │
│ │ <|im_start|>user\n...\n<|im_end|> │
│ │ Qwen 系列使用 │
├──────────────┼───────────────────────────────────────────────┤
│ Llama3 │ <|begin_of_text|><|start_header_id|> │
│ │ system<|end_header_id|>... │
│ │ Llama 3 专用 │
└──────────────┴───────────────────────────────────────────────┘3.2 格式转换
"""
不同数据格式之间的转换
"""
from typing import List, Dict
# ── Alpaca 格式 ──
alpaca_example = {
"instruction": "解释什么是闭包",
"input": "",
"output": "闭包是指一个函数能够访问其外部作用域中的变量..."
}
# ── ShareGPT 格式 ──
sharegpt_example = {
"conversations": [
{"from": "human", "value": "解释什么是闭包"},
{"from": "gpt", "value": "闭包是指一个函数能够访问其外部作用域中的变量..."},
{"from": "human", "value": "给一个 JavaScript 例子"},
{"from": "gpt", "value": "function outer() {\n let x = 10;\n return function inner() {\n return x;\n };\n}"}
]
}
# ── 格式转换器 ──
class FormatConverter:
@staticmethod
def alpaca_to_sharegpt(item: Dict) -> Dict:
"""Alpaca → ShareGPT"""
conversations = []
if item.get("input"):
conversations.append({
"from": "human",
"value": f"{item['instruction']}\n\n{item['input']}"
})
else:
conversations.append({
"from": "human",
"value": item["instruction"]
})
conversations.append({
"from": "gpt",
"value": item["output"]
})
return {"conversations": conversations}
@staticmethod
def sharegpt_to_alpaca(item: Dict) -> List[Dict]:
"""ShareGPT → Alpaca(拆分为多条单轮数据)"""
convs = item["conversations"]
results = []
current_user_msg = ""
for conv in convs:
if conv["from"] == "human":
current_user_msg = conv["value"]
elif conv["from"] == "gpt" and current_user_msg:
results.append({
"instruction": current_user_msg,
"input": "",
"output": conv["value"]
})
current_user_msg = ""
return results
@staticmethod
def alpaca_to_chatml(item: Dict, system_prompt: str = "") -> str:
"""Alpaca → ChatML 格式文本"""
text = ""
if system_prompt:
text += f"<|im_start|>system\n{system_prompt}<|im_end|>\n"
user_msg = item["instruction"]
if item.get("input"):
user_msg += f"\n\n{item['input']}"
text += f"<|im_start|>user\n{user_msg}<|im_end|>\n"
text += f"<|im_start|>assistant\n{item['output']}<|im_end|>"
return text
@staticmethod
def alpaca_to_llama3(item: Dict, system_prompt: str = "") -> str:
"""Alpaca → Llama3 格式"""
text = "<|begin_of_text|>"
if system_prompt:
text += f"<|start_header_id|>system<|end_header_id|>\n\n{system_prompt}<|eot_id|>"
text += f"<|start_header_id|>user<|end_header_id|>\n\n{item['instruction']}"
if item.get("input"):
text += f"\n\n{item['input']}"
text += f"<|eot_id|><|start_header_id|>assistant<|end_header_id|>\n\n{item['output']}<|eot_id|>"
return text
# 使用示例
converter = FormatConverter()
# Alpaca → ShareGPT
sg = converter.alpaca_to_sharegpt(alpaca_example)
# Alpaca → ChatML
chatml = converter.alpaca_to_chatml(alpaca_example, system_prompt="你是一个编程助手")
# Alpaca → Llama3
llama3 = converter.alpaca_to_llama3(alpaca_example, system_prompt="你是一个编程助手")
print("=== ChatML ===")
print(chatml)
print("\n=== Llama3 ===")
print(llama3)3.3 构建训练数据集
"""
构建完整的 SFT 训练数据集
"""
import json
import random
from torch.utils.data import Dataset
class SFTDataset(Dataset):
"""SFT 训练数据集"""
def __init__(self, data_path: str, tokenizer, max_length: int = 2048,
format_type: str = "chatml", system_prompt: str = ""):
self.tokenizer = tokenizer
self.max_length = max_length
self.format_type = format_type
self.system_prompt = system_prompt
self.data = self._load_data(data_path)
def _load_data(self, path: str) -> list:
"""加载数据(支持 json/jsonl)"""
data = []
if path.endswith(".json"):
with open(path, "r", encoding="utf-8") as f:
data = json.load(f)
elif path.endswith(".jsonl"):
with open(path, "r", encoding="utf-8") as f:
for line in f:
data.append(json.loads(line))
return data
def _format_to_text(self, item: dict) -> str:
"""格式化为模型输入文本"""
if self.format_type == "chatml":
return FormatConverter.alpaca_to_chatml(item, self.system_prompt)
elif self.format_type == "llama3":
return FormatConverter.alpaca_to_llama3(item, self.system_prompt)
else:
raise ValueError(f"未知格式: {self.format_type}")
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
item = self.data[idx]
text = self._format_to_text(item)
# Tokenize
encodings = self.tokenizer(
text,
truncation=True,
max_length=self.max_length,
padding="max_length",
return_tensors="pt"
)
# 构建标签:只对 assistant 部分计算 loss
input_ids = encodings["input_ids"][0]
labels = input_ids.clone()
# 找到 assistant 部分的起始位置
if self.format_type == "chatml":
assistant_start = text.find("<|im_start|>assistant\n")
if assistant_start != -1:
# 找到对应的 token 位置
prefix = text[:assistant_start + len("<|im_start|>assistant\n")]
prefix_ids = self.tokenizer(prefix, truncation=True, max_length=self.max_length)["input_ids"]
# 将非 assistant 部分的 label 设为 -100
labels[:len(prefix_ids)] = -100
elif self.format_type == "llama3":
assistant_start = text.find("<|start_header_id|>assistant<|end_header_id|>")
if assistant_start != -1:
prefix = text[:assistant_start]
prefix_ids = self.tokenizer(prefix, truncation=True, max_length=self.max_length)["input_ids"]
labels[:len(prefix_ids)] = -100
return {
"input_ids": input_ids,
"attention_mask": encodings["attention_mask"][0],
"labels": labels
}四、数据质量评估
4.1 评估维度
┌──────────────────────────────────────────────────┐
│ 数据质量评估维度 │
│ │
│ ┌──────────┐ ┌──────────┐ ┌──────────┐ │
│ │ 多样性 │ │ 难度 │ │ 正确性 │ │
│ │ Diversity │ │ Difficulty│ │ Correctness│ │
│ ├──────────┤ ├──────────┤ ├──────────┤ │
│ │ 指令多样 │ │ 难度分布 │ │ 事实正确 │ │
│ │ 领域覆盖 │ │ 推理深度 │ │ 逻辑正确 │ │
│ │ 表达多样 │ │ 知识要求 │ │ 格式正确 │ │
│ └──────────┘ └──────────┘ └──────────┘ │
│ │
│ ┌──────────┐ ┌──────────┐ ┌──────────┐ │
│ │ 一致性 │ │ 完整性 │ │ 平衡性 │ │
│ │Consistency│ │Completeness│ │ Balance │ │
│ ├──────────┤ ├──────────┤ ├──────────┤ │
│ │ 格式一致 │ │ 指令完整 │ │ 领域平衡 │ │
│ │ 风格一致 │ │ 回答完整 │ │ 难度平衡 │ │
│ │ 语气一致 │ │ 上下文完整│ │ 长度平衡 │ │
│ └──────────┘ └──────────┘ └──────────┘ │
└──────────────────────────────────────────────────┘4.2 多样性评估
"""
数据多样性评估
"""
import numpy as np
from collections import Counter
from sklearn.feature_extraction.text import TfidfVectorizer
from sklearn.metrics.pairwise import cosine_similarity
class DiversityEvaluator:
def __init__(self, data: list):
self.data = data
self.instructions = [d["instruction"] for d in data]
def lexical_diversity(self) -> float:
"""词汇多样性(Type-Token Ratio)"""
all_words = []
for inst in self.instructions:
all_words.extend(inst.lower().split())
if not all_words:
return 0
return len(set(all_words)) / len(all_words)
def instruction_length_diversity(self) -> dict:
"""指令长度分布"""
lengths = [len(inst) for inst in self.instructions]
return {
"mean": np.mean(lengths),
"std": np.std(lengths),
"min": min(lengths),
"max": max(lengths),
"median": np.median(lengths),
"p25": np.percentile(lengths, 25),
"p75": np.percentile(lengths, 75),
}
def semantic_diversity(self) -> float:
"""语义多样性(基于 TF-IDF 余弦相似度)"""
if len(self.instructions) < 2:
return 0
vectorizer = TfidfVectorizer(max_features=1000)
tfidf_matrix = vectorizer.fit_transform(self.instructions)
# 计算两两相似度的平均值
sim_matrix = cosine_similarity(tfidf_matrix)
# 取上三角(不含对角线)
n = len(self.instructions)
upper_tri = sim_matrix[np.triu_indices(n, k=1)]
avg_sim = np.mean(upper_tri)
# 多样性 = 1 - 平均相似度
return 1 - avg_sim
def topic_diversity(self, n_topics: int = 10) -> dict:
"""主题分布"""
from sklearn.decomposition import LatentDirichletAllocation
vectorizer = TfidfVectorizer(max_features=500, stop_words='english')
tfidf = vectorizer.fit_transform(self.instructions)
lda = LatentDirichletAllocation(n_components=n_topics, random_state=42)
topics = lda.fit_transform(tfidf)
topic_counts = Counter(topics.argmax(axis=1))
return {
f"topic_{i}": topic_counts.get(i, 0) / len(self.instructions)
for i in range(n_topics)
}
def evaluate(self) -> dict:
"""综合评估"""
return {
"lexical_diversity": self.lexical_diversity(),
"length_distribution": self.instruction_length_diversity(),
"semantic_diversity": self.semantic_diversity(),
"topic_distribution": self.topic_diversity(),
}4.3 难度评估
"""
数据难度评估
"""
from openai import OpenAI
client = OpenAI()
class DifficultyEvaluator:
def __init__(self):
self.client = client
def estimate_difficulty(self, instruction: str, output: str) -> str:
"""用 LLM 评估难度"""
response = self.client.chat.completions.create(
model="gpt-4o-mini",
messages=[{
"role": "user",
"content": f"""评估以下指令-回答对的难度级别。
指令: {instruction}
回答: {output[:500]}
难度级别:
- easy: 简单事实查询或基础操作
- medium: 需要一定推理或多步骤
- hard: 需要深入分析、复杂推理或专业知识
- expert: 需要领域专家级别的知识或创造性解决
只输出一个词: easy/medium/hard/expert"""
}],
temperature=0.0,
)
return response.choices[0].message.content.strip()
def batch_evaluate(self, data: list, sample_size: int = 100) -> dict:
"""批量评估难度分布"""
import random
sample = random.sample(data, min(sample_size, len(data)))
difficulty_counts = {"easy": 0, "medium": 0, "hard": 0, "expert": 0}
for item in sample:
level = self.estimate_difficulty(item["instruction"], item["output"])
if level in difficulty_counts:
difficulty_counts[level] += 1
# 归一化
total = sum(difficulty_counts.values())
return {k: v/total for k, v in difficulty_counts.items()} if total > 0 else difficulty_counts
def rule_based_difficulty(self, instruction: str, output: str) -> str:
"""基于规则的难度评估"""
inst_len = len(instruction)
out_len = len(output)
# 简单启发式规则
if inst_len < 20 and out_len < 100:
return "easy"
elif inst_len < 100 and out_len < 500:
return "medium"
elif inst_len < 500 and out_len < 2000:
return "hard"
else:
return "expert"4.4 正确性评估
"""
数据正确性评估
"""
class CorrectnessEvaluator:
def __init__(self, llm_client):
self.client = llm_client
def llm_judge(self, instruction: str, output: str) -> dict:
"""用 LLM 作为裁判评估正确性"""
response = self.client.chat.completions.create(
model="gpt-4o",
messages=[{
"role": "user",
"content": f"""请评估以下回答的质量,给出 0-10 分的评分。
指令: {instruction}
回答: {output}
评分维度:
1. 准确性 (0-3分): 信息是否正确
2. 完整性 (0-3分): 是否完整回答了指令
3. 清晰度 (0-2分): 表达是否清晰
4. 有用性 (0-2分): 对用户是否有帮助
输出 JSON:
{{"accuracy": X, "completeness": X, "clarity": X, "usefulness": X, "total": X, "reason": "简要说明"}}"""
}],
response_format={"type": "json_object"},
temperature=0.0,
)
import json
return json.loads(response.choices[0].message.content)
def consistency_check(self, data: list, sample_size: int = 50) -> float:
"""一致性检查:同一指令多次生成是否一致"""
import random
sample = random.sample(data, min(sample_size, len(data)))
consistent = 0
for item in sample:
# 用 LLM 重新回答,检查一致性
result = self.llm_judge(item["instruction"], item["output"])
if result.get("total", 0) >= 7:
consistent += 1
return consistent / len(sample)
def format_check(self, data: list) -> dict:
"""格式正确性检查"""
issues = {
"missing_fields": 0,
"empty_output": 0,
"encoding_issues": 0,
"truncated": 0,
}
for item in data:
if not item.get("instruction") or not item.get("output"):
issues["missing_fields"] += 1
continue
if not item["output"].strip():
issues["empty_output"] += 1
continue
# 检查编码问题
if "\\u" in item["output"] or "\\x" in item["output"]:
issues["encoding_issues"] += 1
# 检查截断
if item["output"].rstrip().endswith(("...", "…", "[截断]", "[truncated]")):
issues["truncated"] += 1
return issues五、数据配比策略
5.1 配比的重要性
┌──────────────────────────────────────────────────────┐
│ 数据配比对微调效果的影响 │
│ │
│ 场景:你有一个代码助手,需要微调 │
│ │
│ 方案 A (单一数据): │
│ ████████████████████████████ 100% 代码数据 │
│ → 代码能力强,但对话能力差,不会拒绝有害请求 │
│ │
│ 方案 B (合理配比): │
│ ████████████████░░░░░░░░░ 60% 代码数据 │
│ ██████░░░░░░░░░░░░░░░░░░░ 20% 通用对话 │
│ ████░░░░░░░░░░░░░░░░░░░░░ 15% 安全对齐 │
│ ██░░░░░░░░░░░░░░░░░░░░░░░ 5% 拒绝样本 │
│ → 代码能力强,对话流畅,能拒绝有害请求 │
│ │
│ 方案 C (过度稀释): │
│ ██████░░░░░░░░░░░░░░░░░░░ 30% 代码数据 │
│ ██████░░░░░░░░░░░░░░░░░░░ 30% 通用对话 │
│ ██████░░░░░░░░░░░░░░░░░░░ 30% 写作数据 │
│ ██░░░░░░░░░░░░░░░░░░░░░░░ 10% 其他 │
│ → 各方面都一般,没有突出能力 │
└──────────────────────────────────────────────────────┘5.2 配比策略实现
"""
数据配比管理器
"""
import random
from collections import defaultdict
class DataMixer:
"""数据配比管理"""
def __init__(self):
self.datasets = {} # {name: [data]}
def add_dataset(self, name: str, data: list, weight: float = 1.0):
"""添加数据集"""
self.datasets[name] = {
"data": data,
"weight": weight,
"count": len(data)
}
def mix(self, total_size: int = None, strategy: str = "weighted") -> list:
"""
混合数据
strategy:
- "weighted": 按权重比例混合
- "balanced": 每个数据集等量
- "oversample": 少量数据过采样
- "curriculum": 按难度排序(简单→难)
"""
if strategy == "weighted":
return self._weighted_mix(total_size)
elif strategy == "balanced":
return self._balanced_mix(total_size)
elif strategy == "oversample":
return self._oversample_mix(total_size)
elif strategy == "curriculum":
return self._curriculum_mix()
else:
raise ValueError(f"未知策略: {strategy}")
def _weighted_mix(self, total_size: int) -> list:
"""按权重比例混合"""
if total_size is None:
total_size = sum(d["count"] for d in self.datasets.values())
total_weight = sum(d["weight"] for d in self.datasets.values())
result = []
for name, ds in self.datasets.items():
target_count = int(total_size * (ds["weight"] / total_weight))
if target_count > ds["count"]:
# 数据不够,全部使用
result.extend(ds["data"])
else:
# 随机采样
result.extend(random.sample(ds["data"], target_count))
random.shuffle(result)
return result
def _balanced_mix(self, total_size: int) -> list:
"""均衡混合(每个数据集等量)"""
n_datasets = len(self.datasets)
per_dataset = (total_size or min(d["count"] for d in self.datasets.values())) // n_datasets
result = []
for name, ds in self.datasets.items():
sample_size = min(per_dataset, ds["count"])
result.extend(random.sample(ds["data"], sample_size))
random.shuffle(result)
return result
def _oversample_mix(self, total_size: int) -> list:
"""过采样少量数据"""
max_count = max(d["count"] for d in self.datasets.values())
result = []
for name, ds in self.datasets.items():
if ds["count"] < max_count:
# 过采样:重复数据
oversampled = ds["data"] * (max_count // ds["count"])
remaining = max_count - len(oversampled)
oversampled.extend(random.sample(ds["data"], remaining))
result.extend(oversampled)
else:
result.extend(random.sample(ds["data"], max_count))
random.shuffle(result)
return result
def _curriculum_mix(self) -> list:
"""课程学习:简单→难排序"""
all_data = []
for name, ds in self.datasets.items():
for item in ds["data"]:
# 简单难度评估
difficulty = len(item.get("output", ""))
all_data.append((difficulty, item))
# 按难度排序
all_data.sort(key=lambda x: x[0])
return [item for _, item in all_data]
def print_distribution(self):
"""打印数据分布"""
print("\n=== 数据分布 ===")
total = sum(d["count"] for d in self.datasets.values())
for name, ds in self.datasets.items():
print(f" {name}: {ds['count']} ({ds['count']/total*100:.1f}%) [weight={ds['weight']}]")
print(f" 总计: {total}")
# 使用示例
mixer = DataMixer()
# 添加不同来源的数据
mixer.add_dataset("code", code_data, weight=0.5) # 代码数据 50%
mixer.add_dataset("general", general_data, weight=0.2) # 通用对话 20%
mixer.add_dataset("safety", safety_data, weight=0.15) # 安全对齐 15%
mixer.add_dataset("writing", writing_data, weight=0.1) # 写作 10%
mixer.add_dataset("reject", reject_data, weight=0.05) # 拒绝样本 5%
mixer.print_distribution()
# 按权重混合
mixed_data = mixer.mix(total_size=10000, strategy="weighted")
print(f"\n混合后数据量: {len(mixed_data)}")5.3 不同任务的推荐配比
# 推荐配比模板
RECOMMENDED_RATIOS = {
"代码助手": {
"代码": 0.55,
"通用对话": 0.15,
"数学推理": 0.10,
"安全对齐": 0.10,
"拒绝样本": 0.05,
"格式化输出": 0.05,
},
"客服机器人": {
"客服对话": 0.40,
"通用对话": 0.20,
"产品知识": 0.20,
"安全对齐": 0.10,
"情绪安抚": 0.10,
},
"写作助手": {
"写作": 0.40,
"通用对话": 0.25,
"知识问答": 0.15,
"安全对齐": 0.10,
"格式化输出": 0.10,
},
"通用助手": {
"通用对话": 0.30,
"代码": 0.15,
"数学推理": 0.10,
"写作": 0.15,
"知识问答": 0.15,
"安全对齐": 0.10,
"拒绝样本": 0.05,
},
}六、面试要点
Q1:微调数据多少条合适?
回答:取决于任务复杂度和模型大小。一般经验值:
- 7B 模型简单任务:5000-10000 条
- 7B 模型复杂任务:50000-100000 条
- 70B 模型:100000+ 条
- 少于 1000 条通常不够,容易过拟合
- 多于 500000 条边际收益递减
Q2:如何判断微调数据是否足够?
回答:
- 训练 loss 持续下降但验证 loss 开始上升 → 可能过拟合,数据不够
- 模型在测试集上表现差,但在训练集上很好 → 数据不够或分布不匹配
- 生成答案多样性降低 → 数据量不足导致模式坍缩
- 不同 epoch 的输出质量变化不大 → 数据可能足够
Q3:合成数据和真实数据哪个更好?
回答:真实数据更贴近实际场景,但量少且有噪音。合成数据质量可控但可能有模型偏差。最佳实践是混合使用:真实数据为主(60-70%),合成数据补充(30-40%)。关键是要确保合成数据的多样性。
Q4:如何评估微调数据的质量?
回答:从五个维度评估:
- 多样性:指令的语义多样性、长度分布、领域覆盖
- 正确性:LLM-as-Judge 评分、人工抽检
- 难度分布:easy/medium/hard/expert 的合理分布
- 一致性:格式一致、风格一致、质量一致
- 安全性:无有害内容、已脱敏
Q5:数据清洗中最重要的一步是什么?
回答:去重。重复数据会导致模型对某些模式过拟合,严重影响泛化能力。去重包括精确去重(hash)和模糊去重(语义相似度)。其次是质量过滤,去掉低质量回答。
七、避坑指南
坑 1:过度依赖合成数据
合成数据会让模型学到生成它的模型的风格和偏见。如果用 GPT-4 生成全部数据,微调后的模型会"GPT-4化"。建议混合真实数据和合成数据。
坑 2:忽略数据顺序
# ❌ 错误:按原始顺序训练(同类数据扎堆)
data = load_data() # 可能按类别排序
# ✅ 正确:打乱顺序
import random
random.shuffle(data)
# ✅ 更好:课程学习策略(简单→难)
data.sort(key=lambda x: len(x.get("output", "")))坑 3:验证集污染训练集
# ❌ 错误:从同一来源切分训练集和验证集
all_data = load_dataset("alpaca")
train_data = all_data[:8000]
val_data = all_data[8000:10000] # 可能与训练集有语义重复
# ✅ 正确:确保验证集与训练集无语义重复
# 用模糊去重确保训练集和验证集不重叠坑 4:忘记处理多语言
如果数据包含中英文混合,确保 tokenizer 支持中文。某些英文 tokenizer 对中文的 tokenize 效果很差,导致训练效率低下。
坑 5:指令和输出不匹配
# ❌ 问题:指令问的是 A,输出回答的是 B
{"instruction": "解释递归", "output": "迭代是一种..."} # 指令和输出不匹配
# ✅ 检查方法:用 LLM 验证匹配度
def check_alignment(instruction, output):
"""检查指令和输出是否匹配"""
response = client.chat.completions.create(
model="gpt-4o-mini",
messages=[{
"role": "user",
"content": f"以下回答是否正确回应了指令?只回答 是/否\n\n指令: {instruction}\n回答: {output[:300]}"
}],
temperature=0.0,
)
return "是" in response.choices[0].message.content总结
微调数据工程是一个系统工程,涵盖:
- 数据采集:公开数据打底 + 合成数据补充 + 业务数据增值
- 数据清洗:格式校验 → 去重 → 长度过滤 → 质量过滤 → 脱敏 → 毒性检测
- 格式标准化:统一为 Alpaca/ShareGPT 格式,再转换为模型特定格式
- 质量评估:多样性 + 难度 + 正确性 + 一致性 + 平衡性
- 配比策略:根据任务类型合理配比不同来源数据
一句话总结:数据决定上限,算法逼近上限。在微调中,花 80% 的时间在数据上,20% 在训练调参上,才是正确的优先级。