欢迎光临
我们一直在努力

手写 AI 合成数据生成系统:从零构建高质量训练数据流水线

引言

2025 年以来,大模型训练数据的"墙"正在逼近——网络上可用的高质量文本数据几乎被耗尽。OpenAI、Google 等头部 AI 公司的论文中反复出现一个关键词:合成数据(Synthetic Data)。从 GPT-4 的推理能力提升到 Llama 3 的训练,合成数据已经成为突破数据瓶颈的核心手段。

但合成数据不是简单的"让 AI 自己编数据"。一个高质量的合成数据生成系统,需要解决数据多样性、质量控制和领域覆盖三大核心挑战。

本文将从零实现一个完整的 AI 合成数据生成系统,涵盖文本生成、结构化数据合成、质量过滤与数据增强四大模块。所有代码可直接运行,无需框架依赖。


一、合成数据的技术背景

1.1 为什么需要合成数据?

传统数据采集面临三个难以逾越的瓶颈:

  • 成本问题:人工标注 100 万条高质量指令数据,需要数百人月的工作量和数十万成本
  • 隐私问题:真实用户数据涉及隐私合规(GDPR、个保法),难以大规模使用
  • 长尾问题:高价值的长尾场景(医疗推理、代码审计、数学证明)真实数据极度稀缺

合成数据的核心价值在于:用模型的知识蒸馏出针对特定场景的高质量数据,再反哺模型训练。

1.2 合成数据的路线图

一个完整的合成数据系统包含四个层次:

原始种子数据 → 数据扩展引擎 → 质量控制 → 数据增强 → 高质量训练数据
↓ ↓ ↓ ↓
Seed Generation Filtering Augmentation

本文将从这四个层次逐一实现。


二、系统架构设计

2.1 整体设计

我们先设计系统的顶层架构。合成数据生成系统由以下核心组件组成:

# config.py – 系统配置
from dataclasses import dataclass, field
from typing import List, Optional

@dataclass
class SynthesizerConfig:
"""合成数据系统配置"""
# LLM 配置
api_base: str = "https://api.deepseek.com/v1" # 中性 API 地址示例
api_key: str = "your-api-key"
model: str = "deepseek-chat"

# 生成参数
temperature: float = 0.8
max_tokens: int = 2048
batch_size: int = 5

# 质量过滤配置
min_length: int = 50
max_length: int = 4096
diversity_threshold: float = 0.3
quality_score_threshold: float = 0.7

# 数据增强配置
enable_augmentation: bool = True
augmentation_methods: List[str] = field(
default_factory=lambda: ["paraphrase", "back_translate", "noise_injection"]
)

2.2 核心数据模型

定义贯穿系统的数据模型:

# models.py – 数据模型
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional
import json
import hashlib
from datetime import datetime

@dataclass
class DataPoint:
"""单条数据记录"""
id: str
content: str
metadata: Dict[str, Any] = field(default_factory=dict)
quality_score: float = 0.0
source: str = "synthetic"
created_at: str = field(default_factory=lambda: datetime.now().isoformat())

def compute_hash(self) -> str:
"""计算内容哈希,用于去重"""
return hashlib.md5(self.content.encode()).hexdigest()

@dataclass
class DataBatch:
"""数据批次"""
datapoints: List[DataPoint] = field(default_factory=list)
batch_metadata: Dict[str, Any] = field(default_factory=dict)

def add(self, content: str, **metadata):
dp = DataPoint(
id=f"dp_{len(self.datapoints)}_{datetime.now().timestamp()}",
content=content,
metadata=metadata
)
self.datapoints.append(dp)
return dp

def to_jsonl(self, filepath: str):
"""导出为 JSONL 格式"""
with open(filepath, 'w', encoding='utf-8') as f:
for dp in self.datapoints:
f.write(json.dumps({
'id': dp.id,
'content': dp.content,
'metadata': dp.metadata,
'quality_score': dp.quality_score,
'source': dp.source
}, ensure_ascii=False) + '\\n')
print(f"✅ 导出 {len(self.datapoints)} 条数据到 {filepath}")

@classmethod
def from_jsonl(cls, filepath: str) -> 'DataBatch':
batch = cls()
with open(filepath, 'r', encoding='utf-8') as f:
for line in f:
data = json.loads(line.strip())
dp = DataPoint(**data)
batch.datapoints.append(dp)
return batch

def stats(self) -> Dict:
"""数据统计"""
if not self.datapoints:
return {"count": 0}
scores = [dp.quality_score for dp in self.datapoints]
lengths = [len(dp.content) for dp in self.datapoints]
return {
"count": len(self.datapoints),
"avg_quality": sum(scores) / len(scores),
"avg_length": sum(lengths) / len(lengths),
"min_length": min(lengths),
"max_length": max(lengths),
"unique_sources": len(set(dp.source for dp in self.datapoints))
}


三、种子数据管理

种子数据是合成数据的"基因"。我们需要一个种子管理器来维护初始样本。

# seed_manager.py – 种子数据管理器
from typing import List, Dict, Optional
import json
import random

class SeedManager:
"""种子数据管理器"""

def __init__(self):
self.seeds: Dict[str, List[str]] = {
"qa": self._load_default_qa_seeds(),
"code": self._load_default_code_seeds(),
"reasoning": self._load_default_reasoning_seeds(),
"creative": self._load_default_creative_seeds(),
}

def _load_default_qa_seeds(self) -> List[str]:
"""加载默认问答种子"""
return [
"解释什么是数据库索引以及它是如何提高查询性能的。",
"描述 RESTful API 的设计原则和最佳实践。",
"什么是死锁?如何避免和解决死锁问题?",
"解释面向对象编程中的继承和多态。",
"描述 HTTP 和 HTTPS 的区别及其工作原理。",
]

def _load_default_code_seeds(self) -> List[str]:
return [
"实现一个二分查找算法。",
"编写一个函数来反转链表。",
"实现 LRU 缓存淘汰策略。",
"使用 Python 实现快速排序算法。",
"编写一个 SQL 查询来统计每个类别的销售总额。",
]

def _load_default_reasoning_seeds(self) -> List[str]:
return [
"如果所有 A 都是 B,所有 B 都是 C,那么以下哪个结论是正确的?",
"一个水池有一个进水口和一个出水口,单独开进水口需要 6 小时注满,单独开出水管需要 9 小时排空,两个同时开需要多少小时?",
]

def _load_default_creative_seeds(self) -> List[str]:
return [
"写一篇关于未来城市的短篇科幻故事。",
"以\\"最后一片叶子\\"为主题创作一首诗。",
]

def add_seed(self, category: str, seed: str):
"""添加种子"""
if category not in self.seeds:
self.seeds[category] = []
self.seeds[category].append(seed)

def get_seeds(self, category: Optional[str] = None,
n: int = -1, shuffle: bool = True) -> List[str]:
"""获取种子"""
if category:
seeds = self.seeds.get(category, [])
else:
seeds = [s for cat in self.seeds.values() for s in cat]

if shuffle:
random.shuffle(seeds)

if n > 0:
seeds = seeds[:n]

return seeds

def add_domain_seeds(self, domain: str, seeds: List[str]):
"""添加领域专有种子"""
self.seeds[domain] = seeds
print(f"✅ 添加领域 [{domain}]: {len(seeds)} 条种子")


四、数据扩展引擎

这是整个系统的核心——从种子数据出发,生成多样化的合成数据。

4.1 基础生成器

# generators/base_generator.py
from typing import List, Dict, Optional, Callable
import json
import time

class BaseGenerator:
"""基础生成器"""

def __init__(self, config, llm_client):
self.config = config
self.client = llm_client
self.generation_models: Dict[str, Dict] = {}

def register_model(self, name: str,
system_prompt: str,
template_fn: Callable):
"""注册生成模板"""
self.generation_models[name] = {
"system_prompt": system_prompt,
"template_fn": template_fn,
}

def generate(self, model_name: str, seed: str,
**kwargs) -> Optional[str]:
"""执行单次生成"""
if model_name not in self.generation_models:
raise ValueError(f"未知生成模型: {model_name}")

model = self.generation_models[model_name]
prompt = model["template_fn"](seed, **kwargs)

try:
response = self.client.chat(
model=self.config.model,
messages=[
{"role": "system", "content": model["system_prompt"]},
{"role": "user", "content": prompt}
],
temperature=self.config.temperature,
max_tokens=self.config.max_tokens,
)
return response
except Exception as e:
print(f"❌ 生成失败 [{model_name}]: {e}")
return None

def batch_generate(self, model_name: str, seeds: List[str],
**kwargs) -> List[str]:
"""批量生成"""
results = []
for i in range(0, len(seeds), self.config.batch_size):
batch = seeds[i:i + self.config.batch_size]
batch_results = []
for seed in batch:
result = self.generate(model_name, seed, **kwargs)
if result:
batch_results.append(result)
time.sleep(0.5) # 限流保护
results.extend(batch_results)
print(f"📊 批次完成: {i + len(batch)}/{len(seeds)}")
return results

4.2 指令数据生成器

专门用于生成高质量指令-回复对:

# generators/instruction_generator.py
from .base_generator import BaseGenerator

class InstructionGenerator(BaseGenerator):
"""指令数据生成器"""

def __init__(self, config, llm_client):
super().__init__(config, llm_client)
self._register_templates()

def _register_templates(self):
"""注册各种生成模板"""

# 1. 基础问答生成
self.register_model(
"qa_generation",
system_prompt="你是一个AI训练数据工程师。根据给定的主题,生成一个高质量的技术问答对。"
"要求答案准确、详细、有条理,包含具体示例。",
template_fn=lambda seed, **kw:
f"根据以下主题生成一个技术问答对(instruction + response):\\n"
f"主题:{seed}\\n"
f"格式要求:以 '## Q:' 和 '## A:' 标记问题和答案。\\n"
f"答案至少 200 字,包含代码示例或具体场景。"
)

# 2. 思维链推理生成
self.register_model(
"cot_generation",
system_prompt="你是一个推理数据生成专家。生成包含完整思维过程的问题-推理-答案三元组。"
"推理步骤要细致,每步都要有明确的逻辑依据。",
template_fn=lambda seed, **kw:
f"基于以下种子问题,生成一个包含思维链(Chain of Thought)的推理样本:\\n"
f"种子:{seed}\\n\\n"
f"请包含:\\n"
f"1. 完整问题描述\\n"
f"2. 逐步推理过程(用 Step 1, Step 2… 标记)\\n"
f"3. 最终答案"
)

# 3. 多轮对话生成
self.register_model(
"multi_turn_generation",
system_prompt="你是一个对话数据生成专家。根据给定主题,生成一个包含 3-5 轮的多轮对话。"
"对话要自然流畅,展现上下文理解和对话连贯性。",
template_fn=lambda seed, **kw:
f"基于以下主题生成一个多轮技术对话(3-5 轮):\\n"
f"主题:{seed}\\n\\n"
f"要求:\\n"
f"- 每轮对话包含 User 和 Assistant 的交互\\n"
f"- 体现对历史对话的理解和引用\\n"
f"- 难度逐步递进,从基础到深入"
)

4.3 代码数据生成器

# generators/code_generator.py
from .base_generator import BaseGenerator

class CodeGenerator(BaseGenerator):
"""代码数据生成器"""

def __init__(self, config, llm_client):
super().__init__(config, llm_client)
self._register_templates()

def _register_templates(self):
# 代码实现 – 含注释的完整代码
self.register_model(
"code_implementation",
system_prompt="你是一个代码教学专家。生成完整、可运行、包含详细中文注释的代码样本。"
"代码必须风格规范,包含错误处理和边界检查。",
template_fn=lambda seed, **kw:
f"请实现以下算法/功能,要求代码完整可用:\\n"
f"需求:{seed}\\n\\n"
f"格式要求:\\n"
f"1. 用 markdown 代码块包裹\\n"
f"2. 包含详细中文注释\\n"
f"3. 包含使用示例和测试用例\\n"
f"4. 包含时间复杂度说明"
)

# Bug修复样本生成
self.register_model(
"bug_fix_generation",
system_prompt="你是一个代码审查专家。生成一个有 Bug 的代码段及其对应的修复版本。",
template_fn=lambda seed, **kw:
f"针对以下功能需求,生成一个含有 Bug 的代码段和对应的修复方案:\\n"
f"功能:{seed}\\n\\n"
f"输出格式:\\n"
f"## Bug 代码\\n[有问题的代码]\\n"
f"## Bug 描述\\n[Bug 分析]\\n"
f"## 修复代码\\n[修复后的代码]\\n"
f"## 修复说明\\n[解释为什么这样修复]"
)

4.4 LLM 客户端封装

需要一个简单的 LLM 客户端来与 API 交互:

# llm_client.py – LLM API 客户端
import requests
import json
from typing import List, Dict, Optional

class LLMClient:
"""LLM API 客户端(通用接口)"""

def __init__(self, config):
self.api_base = config.api_base
self.api_key = config.api_key
self.session = requests.Session()
self.session.headers.update({
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json"
})

def chat(self, model: str, messages: List[Dict],
temperature: float = 0.8,
max_tokens: int = 2048) -> Optional[str]:
"""调用 LLM API"""
payload = {
"model": model,
"messages": messages,
"temperature": temperature,
"max_tokens": max_tokens,
}

try:
resp = self.session.post(
f"{self.api_base}/chat/completions",
json=payload,
timeout=60
)
resp.raise_for_status()
data = resp.json()
return data['choices'][0]['message']['content']
except Exception as e:
print(f"❌ API 调用失败: {e}")
return None


五、质量控制模块

合成数据的质量直接决定训练效果。我们需要一个多维度质量评估系统。

# quality_control.py – 质量控制
from typing import List, Dict, Tuple
import re
from collections import Counter

class QualityController:
"""质量控制器"""

def __init__(self, config):
self.config = config
self.filters: List[callable] = []
self._register_default_filters()

def _register_default_filters(self):
"""注册默认过滤规则"""
self.filters.extend([
self._length_filter,
self._language_filter,
self._repetition_filter,
self._toxicity_filter,
])

def _length_filter(self, dp) -> Tuple[bool, str]:
"""长度过滤"""
length = len(dp.content)
if length < self.config.min_length:
return False, f"长度不足: {length} < {self.config.min_length}"
if length > self.config.max_length:
return False, f"长度超标: {length} > {self.config.max_length}"
return True, ""

def _language_filter(self, dp) -> Tuple[bool, str]:
"""语言质量检查"""
content = dp.content

# 检查是否包含完整句子
sentences = re.split(r'[。!?\\n.!?]', content)
valid_sentences = [s.strip() for s in sentences if len(s.strip()) > 10]

if len(valid_sentences) < 2:
return False, "有效句子不足"

# 检查中文占比
chinese_chars = len(re.findall(r'[\\u4e00-\\u9fff]', content))
total_chars = len(content.strip())
if total_chars > 0 and chinese_chars / total_chars < 0.3:
return False, "中文占比过低"

return True, ""

def _repetition_filter(self, dp) -> Tuple[bool, str]:
"""重复检测"""
content = dp.content

# N-gram 重复检测
def ngram_repetition_rate(text: str, n: int = 3) -> float:
words = list(text)
if len(words) < n:
return 0.0
ngrams = [tuple(words[i:i+n]) for i in range(len(words)-n+1)]
if not ngrams:
return 0.0
unique = len(set(ngrams))
total = len(ngrams)
return 1.0 – unique / total

rep_rate = ngram_repetition_rate(content)
if rep_rate > 0.5:
return False, f"重复率过高: {rep_rate:.2%}"

return True, ""

def _toxicity_filter(self, dp) -> Tuple[bool, str]:
"""基础毒性检测"""
toxic_patterns = [
r'(?:fuck|shit|damn|asshole|草泥马|他妈)',
]
for pattern in toxic_patterns:
if re.search(pattern, dp.content, re.IGNORECASE):
return False, "包含不当内容"
return True, ""

def register_filter(self, filter_fn: callable):
"""注册自定义过滤器"""
self.filters.append(filter_fn)

def evaluate_quality(self, dp) -> float:
"""质量评分(0-1)"""
score = 1.0

# 长度分值
length = len(dp.content)
if length < self.config.min_length:
score *= 0.5
elif length < self.config.min_length * 2:
score *= 0.8

# 多样性分值
unique_chars = len(set(dp.content))
char_ratio = unique_chars / max(len(dp.content), 1)
if char_ratio < 0.3:
score *= 0.6

# 信息密度
code_blocks = len(re.findall(r'```', dp.content)) // 2
if code_blocks > 0:
score = min(1.0, score + 0.1)

dp.quality_score = round(score, 2)
return dp.quality_score

def filter_and_score(self, batch: List) -> List:
"""完整过滤+评分流水线"""
results = []
stats = {"total": len(batch), "passed": 0, "failed": 0, "reasons": Counter()}

for dp in batch:
all_passed = True

# 运行所有过滤器
for filter_fn in self.filters:
passed, reason = filter_fn(dp)
if not passed:
all_passed = False
stats["reasons"][reason] += 1
break

if not all_passed:
stats["failed"] += 1
continue

# 质量评分
self.evaluate_quality(dp)

# 评分阈值过滤
if dp.quality_score >= self.config.quality_score_threshold:
results.append(dp)
stats["passed"] += 1
else:
stats["failed"] += 1
stats["reasons"][f"评分不足: {dp.quality_score}"] += 1

print(f"📊 质控结果: 通过 {stats['passed']}/{stats['total']} | "
f"失败 {stats['failed']}/{stats['total']}")
return results


六、数据增强模块

数据增强可以进一步提升数据的多样性。

# augmentation.py – 数据增强
from typing import List, Optional
import random
import re

class DataAugmentor:
"""数据增强器"""

def __init__(self, config, llm_client=None):
self.config = config
self.client = llm_client
self.methods = {
"paraphrase": self._paraphrase,
"noise_injection": self._noise_injection,
"entity_replacement": self._entity_replacement,
}

def _paraphrase(self, text: str) -> str:
"""同义改写"""
# 简单同义词替换
synonyms = {
"实现": ["实现", "构建", "开发", "创建"],
"系统": ["系统", "平台", "框架", "工具"],
"使用": ["使用", "采用", "利用", "应用"],
"方法": ["方法", "方式", "手段", "方案"],
"数据": ["数据", "信息", "资料", "内容"],
"模型": ["模型", "算法", "架构", "方案"],
"性能": ["性能", "效率", "表现", "效果"],
"问题": ["问题", "挑战", "难点", "需求"],
"分析": ["分析", "研究", "探讨", "解析"],
"优化": ["优化", "改善", "改进", "提升"],
}

words = list(text)
result = text
for word, syns in synonyms.items():
if word in result and random.random() < 0.3:
result = result.replace(word, random.choice(syns), 1)

return result if result != text else text

def _noise_injection(self, text: str) -> str:
"""噪声注入"""
# 仅在代码段外注入噪声
parts = re.split(r'(```.*?```)', text, flags=re.DOTALL)

for i, part in enumerate(parts):
if not part.startswith('```'): # 非代码部分
# 轻微拼写错误(概率很低)
if random.random() < 0.05:
chars = list(part)
if len(chars) > 10:
idx = random.randint(0, len(chars) – 1)
# 不修改中文字符和数字
if chars[idx].isalpha():
chars[idx] = random.choice('abcdefghijklmnopqrstuvwxyz')
parts[i] = ''.join(chars)

return ''.join(parts)

def _entity_replacement(self, text: str) -> str:
"""实体替换"""
replacements = {
"Python": ["Python", "Java", "Go", "Rust", "TypeScript"],
"JavaScript": ["JavaScript", "TypeScript", "Python", "Ruby"],
"MySQL": ["MySQL", "PostgreSQL", "MongoDB", "Redis"],
"Linux": ["Linux", "macOS", "Windows"],
"GPU": ["GPU", "TPU", "NPU"],
"Transformer": ["Transformer", "BERT", "GPT", "LLaMA"],
}

result = text
for entity, options in replacements.items():
if entity in result and random.random() < 0.2:
result = result.replace(entity, random.choice(options), 1)

return result

def augment(self, text: str, methods: Optional[List[str]] = None) -> List[str]:
"""执行增强,返回多个变体"""
if methods is None:
methods = list(self.methods.keys())

augmented = [text] # 保留原始版本

for method in methods:
if method in self.methods:
try:
result = self.methods[method](text)
if result and result != text:
augmented.append(result)
except Exception as e:
print(f"⚠️ 增强失败 [{method}]: {e}")

return augmented

def augment_batch(self, batch: List, methods: Optional[List[str]] = None) -> List:
"""批量增强"""
augmented_batch = []
for dp in batch:
variants = self.augment(dp.content, methods)
for i, variant in enumerate(variants):
new_dp = DataPoint(
id=f"{dp.id}_aug_{i}",
content=variant,
metadata={**dp.metadata, "augmented": i > 0, "original_id": dp.id},
quality_score=dp.quality_score,
source=f"{dp.source}_augmented"
)
augmented_batch.append(new_dp)

print(f"🔧 数据增强: {len(batch)} → {len(augmented_batch)} 条")
return augmented_batch


七、数据去重模块

合成数据很容易产生大量重复内容,去重是保证数据质量的关键环节。

# deduplication.py – 数据去重
from typing import List, Set, Dict
import hashlib
from collections import defaultdict

class Deduplicator:
"""数据去重器"""

def __init__(self):
self.content_hashes: Set[str] = set()
self.near_duplicates: Dict[str, List[str]] = defaultdict(list)

def exact_dedup(self, batch: List) -> List:
"""精确去重(基于 MD5 哈希)"""
unique = []
for dp in batch:
h = hashlib.md5(dp.content.encode()).hexdigest()
if h not in self.content_hashes:
self.content_hashes.add(h)
unique.append(dp)

removed = len(batch) – len(unique)
if removed > 0:
print(f"🗑️ 精确去重: 移除 {removed} 条重复")
return unique

def near_dedup(self, batch: List, threshold: float = 0.85) -> List:
"""近似去重(基于 Jaccard 相似度)"""
def shingles(text: str, k: int = 3) -> Set[str]:
"""生成 k-shingles"""
words = text.split()
if len(words) < k:
return {text}
return {' '.join(words[i:i+k]) for i in range(len(words)-k+1)}

def jaccard(a: Set, b: Set) -> float:
intersection = len(a & b)
union = len(a | b)
return intersection / union if union > 0 else 0.0

selected = []
for dp in batch:
dp_shingles = shingles(dp.content)
is_duplicate = False

for selected_dp in selected:
similarity = jaccard(dp_shingles, shingles(selected_dp.content))
if similarity >= threshold:
is_duplicate = True
break

if not is_duplicate:
selected.append(dp)

removed = len(batch) – len(selected)
if removed > 0:
print(f"🗑️ 近似去重: 移除 {removed} 条(阈值 {threshold})")
return selected

def context_dedup(self, batch: List, context_window: int = 100) -> List:
"""上下文去重(保留多样性)"""
# 按种子来源分组,每组只保留质量最高的
from collections import defaultdict

groups = defaultdict(list)
for dp in batch:
seed_key = dp.metadata.get('seed', 'unknown')[:50]
groups[seed_key].append(dp)

result = []
for key, group in groups.items():
# 按质量评分排序,保留评分最高的
group.sort(key=lambda x: x.quality_score, reverse=True)
# 保留 top 3 或全部(如果少于 3)
result.extend(group[:min(3, len(group))])

removed = len(batch) – len(result)
if removed > 0:
print(f"🗑️ 上下文去重: 移除 {removed} 条")
return result


八、编排引擎——组装完整流水线

现在,我们把所有模块组装成一个端到端的合成数据生成流水线。

# pipeline.py – 完整流水线
from typing import List, Optional, Dict
import time
import os

class SyntheticDataPipeline:
"""合成数据生成流水线"""

def __init__(self, config):
self.config = config
self.client = LLMClient(config)
self.seed_manager = SeedManager()
self.generator = BaseGenerator(config, self.client)
self.instruction_gen = InstructionGenerator(config, self.client)
self.code_gen = CodeGenerator(config, self.client)
self.quality_controller = QualityController(config)
self.augmentor = DataAugmentor(config, self.client)
self.deduplicator = Deduplicator()

# 注册生成模板
self._register_generation_strategies()

def _register_generation_strategies(self):
"""注册生成策略"""
self.strategies = {
"qa": {
"generator": self.instruction_gen,
"model": "qa_generation",
},
"code": {
"generator": self.code_gen,
"model": "code_implementation",
},
"reasoning": {
"generator": self.instruction_gen,
"model": "cot_generation",
},
}

def run(self, strategy: str = "qa",
seed_category: Optional[str] = None,
num_seeds: int = 5,
output_dir: str = "./synthetic_data") -> DataBatch:
"""执行完整流水线"""

print(f"🚀 开始合成数据生成 [{strategy}]")
start_time = time.time()

# Step 1: 获取种子
seeds = self.seed_manager.get_seeds(
category=seed_category, n=num_seeds
)
print(f"📦 种子: {len(seeds)} 条")

# Step 2: 数据生成
raw_results = DataBatch()
strategy_config = self.strategies.get(strategy)
if not strategy_config:
raise ValueError(f"未知策略: {strategy}")

generator = strategy_config["generator"]
model_name = strategy_config["model"]

for seed in seeds:
result = generator.generate(model_name, seed)
if result:
raw_results.add(
content=result,
seed=seed,
strategy=strategy,
generation_model=model_name
)
time.sleep(0.3)

print(f"📝 原始生成: {len(raw_results.datapoints)} 条")

# Step 3: 质量过滤
qc_results = self.quality_controller.filter_and_score(
raw_results.datapoints
)
qc_batch = DataBatch(datapoints=qc_results)

# Step 4: 数据增强
if self.config.enable_augmentation:
augmented = self.augmentor.augment_batch(
qc_results,
methods=self.config.augmentation_methods
)
else:
augmented = qc_results
aug_batch = DataBatch(datapoints=augmented)

# Step 5: 去重
deduped = self.deduplicator.exact_dedup(aug_batch.datapoints)
deduped = self.deduplicator.context_dedup(deduped)
final_batch = DataBatch(datapoints=deduped)

# Step 6: 导出
os.makedirs(output_dir, exist_ok=True)
timestamp = time.strftime("%Y%m%d_%H%M%S")
output_path = f"{output_dir}/{strategy}_{timestamp}.jsonl"
final_batch.to_jsonl(output_path)

elapsed = time.time() – start_time
print(f"\\n✅ 流水线完成!")
print(f"⏱️ 耗时: {elapsed:.2f}s")
print(f"📊 最终产出: {len(final_batch.datapoints)} 条高质量数据")
print(f"📂 保存路径: {output_path}")

return final_batch

def stats(self, batch: DataBatch):
"""输出统计数据"""
stats = batch.stats()
print(f"""
📊 合成数据统计
━━━━━━━━━━━━━━━━━━
总条数: {stats['count']}
平均质量分: {stats['avg_quality']:.2f}
平均长度: {stats['avg_length']:.0f} 字符
最短: {stats['min_length']} 字符
最长: {stats['max_length']} 字符
来源多样性: {stats['unique_sources']}
━━━━━━━━━━━━━━━━━━
""")


九、高级技巧与生产化实践

9.1 领域自适应生成

不同领域的数据格式差异巨大。我们可以实现领域适配器:

# domain_adapter.py
class DomainAdapter:
"""领域适配器"""

domain_configs = {
"medical": {
"templates": [
"解释 {term} 的病理机制和临床表现。",
"鉴别诊断:{symptom_a} vs {symptom_b}",
"描述 {treatment} 的治疗方案和适应症。",
],
"quality_rules": {
"min_domain_terms": 3,
"require_references": True,
}
},
"legal": {
"templates": [
"分析 {law} 条文的立法目的和适用范围。",
"案例评析:{case} 中的法律争议焦点。",
"对比 {regulation_a} 和 {regulation_b} 的异同。",
],
"quality_rules": {
"min_domain_terms": 5,
"require_citations": True,
}
},
"finance": {
"templates": [
"分析 {indicator} 对市场的影响机制。",
"风险评估框架:{scenario} 场景下的风控方案。",
"量化策略实现:{strategy_name} 策略的回测分析。",
]
}
}

9.2 迭代优化循环

单次生成往往不够完美,需要迭代优化:

# iterative_refinement.py
class IterativeRefiner:
"""迭代优化器"""

def refine(self, batch: DataBatch,
max_iterations: int = 3) -> DataBatch:
"""迭代优化数据质量"""
current_batch = batch

for iteration in range(max_iterations):
print(f"🔄 迭代优化: 第 {iteration + 1} 轮")

# 质量评分
low_quality = [
dp for dp in current_batch.datapoints
if dp.quality_score < 0.6
]

if not low_quality:
print(f"✅ 所有数据质量达标,提前终止")
break

# 重新生成低质量数据
for dp in low_quality:
new_dp = self._regenerate(dp)
if new_dp and new_dp.quality_score > dp.quality_score:
# 替换为优化版本
idx = current_batch.datapoints.index(dp)
current_batch.datapoints[idx] = new_dp

print(f"📊 本轮优化: {len(low_quality)} 条低质量数据")

return current_batch

9.3 数据平衡策略

生成数据时需要注意类别平衡:

# balancing.py
from collections import Counter
import random

class DataBalancer:
"""数据平衡器"""

def balance_by_category(self, batch: DataBatch,
max_per_category: int = 100) -> DataBatch:
"""按类别平衡数据"""
categories = Counter()
for dp in batch.datapoints:
cat = dp.metadata.get('category', 'unknown')
categories[cat] += 1

# 对超大类采样
balanced = []
cat_counts = Counter()
for dp in batch.datapoints:
cat = dp.metadata.get('category', 'unknown')
if cat_counts[cat] < max_per_category:
balanced.append(dp)
cat_counts[cat] += 1

result = DataBatch(datapoints=balanced)
print(f"⚖️ 数据平衡: {len(batch.datapoints)} → {len(balanced)} 条")
return result


十、完整使用示例

10.1 基本用法

# example_usage.py
from config import SynthesizerConfig
from pipeline import SyntheticDataPipeline

# 初始化配置
config = SynthesizerConfig(
api_key="your-api-key",
temperature=0.8,
batch_size=3,
enable_augmentation=True,
)

# 初始化流水线
pipeline = SyntheticDataPipeline(config)

# 方法 1:使用内置种子生成问答数据
qa_batch = pipeline.run(
strategy="qa",
seed_category="qa",
num_seeds=3,
output_dir="./synthetic_data"
)

# 输出统计
pipeline.stats(qa_batch)

# 方法 2:自定义领域种子
pipeline.seed_manager.add_domain_seeds("nlp", [
"解释 Transformer 中的多头注意力机制。",
"什么是 BPE 分词算法?它的优缺点是什么?",
"如何评估文本摘要模型的质量?",
])

nlp_batch = pipeline.run(
strategy="qa",
seed_category="nlp",
num_seeds=3,
)

# 方法 3:生成代码数据
code_batch = pipeline.run(
strategy="code",
seed_category="code",
num_seeds=3,
)

# 方法 4:批量导出
nlp_batch.to_jsonl("./synthetic_data/nlp_training_data.jsonl")

10.2 自定义种子 + 手动生成

# 自定义种子
custom_seeds = [
"解释向量数据库中的 HNSW 索引原理。",
"对比 RAG 和 Fine-tuning 两种知识注入方式的优劣。",
"描述 LoRA 微调的训练流程和参数更新机制。",
]

for seed in custom_seeds:
pipeline.seed_manager.add_seed("custom", seed)

custom_batch = pipeline.run(
strategy="qa",
seed_category="custom",
num_seeds=3,
output_dir="./synthetic_data"
)

10.3 与训练框架集成

合成数据的最终目的是用于训练。以下是与主流训练框架的集成示例:

# integration.py
class TrainingDataExporter:
"""训练数据导出器"""

@staticmethod
def to_sharegpt_format(batch: DataBatch, output_path: str):
"""导出为 ShareGPT 格式(用于 LLaMA-Factory 等框架)"""
conversations = []

for dp in batch.datapoints:
conv = {
"conversations": [
{
"from": "human",
"value": dp.metadata.get('instruction', '')
},
{
"from": "gpt",
"value": dp.content
}
]
}
conversations.append(conv)

import json
with open(output_path, 'w', encoding='utf-8') as f:
json.dump(conversations, f, ensure_ascii=False, indent=2)

print(f"✅ 导出 {len(conversations)} 条 ShareGPT 格式数据")

@staticmethod
def to_openai_format(batch: DataBatch, output_path: str):
"""导出为 OpenAI Fine-tuning 格式"""
examples = []

for dp in batch.datapoints:
example = {
"messages": [
{"role": "user", "content": dp.metadata.get('instruction', '')},
{"role": "assistant", "content": dp.content}
]
}
examples.append(example)

with open(output_path, 'w', encoding='utf-8') as f:
for ex in examples:
f.write(json.dumps(ex, ensure_ascii=False) + '\\n')

print(f"✅ 导出 {len(examples)} 条 OpenAI 格式数据")


性能考量与最佳实践

1. 成本控制

调用 LLM API 生成合成数据是有成本的。以下优化策略可以显著降低成本:

  • 批处理调用:每次调用生成多条候选数据
  • 复用中间结果:缓存质量评估结果,避免重复生成
  • 渐进式生成:先小批量测试质量,再大规模生产

2. 质量保障

  • 黄金测试集:保留 100-500 条人工验证的高质量数据作为评估基准
  • A/B 验证:对比合成数据训练的模型 vs 基线模型的性能差异
  • 人工抽样:每 1000 条抽样 50 条人工审核

3. 常见陷阱

❌ 循环自引用:A 生成 B,B 生成 A,导致模型退化(模型崩溃/MAD)
❌ 模式固化:所有数据格式高度雷同,降低模型泛化能力
❌ 偏差放大:生成数据会放大 LLM 自身的偏见和错误
✅ 解决方案:混合 10-30% 真实数据,引入多模型交叉生成

4. 模型崩溃(Model Collapse)的避免

模型崩溃是合成数据领域最大的隐患——用模型生成的数据训练新模型,新模型的质量会逐代下降。应对策略:

# anti_collapse.py
class AntiCollapseStrategy:
"""防止模型崩溃的策略"""

@staticmethod
def mixed_training_data(real_data: DataBatch,
synthetic_data: DataBatch,
ratio: float = 0.3) -> DataBatch:
"""混合真实与合成数据"""
import random

real_count = int(len(synthetic_data.datapoints) * ratio / (1 – ratio))
real_count = min(real_count, len(real_data.datapoints))

selected_real = random.sample(real_data.datapoints, real_count)
mixed = selected_real + synthetic_data.datapoints
random.shuffle(mixed)

return DataBatch(datapoints=mixed)


总结

本文从零实现了一个完整的 AI 合成数据生成系统,涵盖:

  • 种子管理:维护高质量的初始样本库
  • 数据扩展引擎:从种子出发多样化生成
  • 质量控制:多维度过滤和评分
  • 数据增强:提升数据多样性
  • 去重机制:保证数据唯一性
  • 编排流水线:端到端自动化生成
  • 合成数据不是"伪数据",而是一种高效的数据工程手段。正确使用合成数据,可以显著提升模型在特定场景下的表现,同时大幅降低数据采集成本。

    关键要点: – 质量 > 数量:1000 条高质量合成数据优于 10 万条低质量数据 – 混合策略:合成数据应与真实数据混合使用 – 迭代优化:合成数据的质量需要不断评估和迭代改进 – 警惕模型崩溃:始终保留真实数据锚点


    📚 延伸阅读

    如果你对 DeepSeek 的实战用法感兴趣,推荐阅读我的另一篇文章:

    👉 DeepSeek 实战指南:提示词工程、API 集成与效率提升全攻略

    这篇文章系统地拆解了 DeepSeek 的提示词工程技巧、API 封装方法以及日常效率提升场景,全文代码可直接运行,适合已经上手 DeepSeek 但希望更高效使用的开发者。


    本文是"手写 AI 系统"系列文章之一。该系列从零实现 AI 系统中的关键组件,涵盖 RAG、Agent、Function Calling、MCP 等核心技术,帮助你深入理解底层原理,构建属于自己的 AI 工具。

    赞(0)
    未经允许不得转载:171主机测评 » 手写 AI 合成数据生成系统:从零构建高质量训练数据流水线
    分享到: 更多 (0)

    评论 抢沙发

    • 昵称 (必填)
    • 邮箱 (必填)
    • 网址