用 32B 当老师蒸馏 7B:一条 NL2SQL Agent 的数据工程流水线
用 32B 当老师蒸馏 7B:一条 NL2SQL Agent 的数据工程流水线
一、问题:7B 的天花板到底在哪
BotForge QA 是一个 NL2SQL/Agent 系统:运营人员用自然语言问数据问题 → 系统检索 RAG schema 文档 → 模型生成 SQL → 真执行 MySQL → 失败则把错误回传模型自修正(最多 3 轮)→ 生成自然语言回复。
Qwen3-32B-FP8 零样本(无任何微调)在固定验收集上 49/49、ad-hoc 全修,效果最好,但要常驻 4×RTX 5090(TP=4),太贵。目标:把这份能力压到 Qwen2.5-7B(LoRA 微调)上,省下 4 张卡。
经过几轮迭代,最终结论一句话:
瓶颈是训练数据的「模式覆盖面」,不是参数量。
证据:用 32B 打标 + SQL 真执行过滤产出 1348 条蒸馏数据(含 151 条错误自修正 retry 样本)+ 历史 547 条手写 = 1895 条微调 7B → 固定验收 49/49,全新 held-out 10 题 9/10(90%);而旧手写数据训出的 7B 同类 ad-hoc 仅 17~50%。差距根因是数据覆盖面窄,不是 7B 容量不足。
于是 32B 的角色从"生产推理引擎"变成"数据生产机 + 效果上限参照 + 兜底":7B 要继续提升,靠让 32B 产出更多覆盖新模式的样本,而不是调 7B 的超参。
二、技术路线总览
迭代闭环
线上 7B 失败 case
→ 归因:先查 RAG 是否命中正确文档(不要先改模型)
├─ 属知识缺失 → 补 RAG 文档,即时生效,无需训练
└─ 属推理模式缺失 → 32B 打标产出新样本 + SQL 真执行过滤
→ 增量微调 7B
→ 双轨验证:49 项验收(防回归) + held-out(验泛化)关键设计是把"补知识"和"补能力"分流:前者改 RAG 文档、成本近零;后者才走蒸馏+训练。早期吃过亏——把两类问题都当"训练数据不够"来解,投在了最低 ROI 的层。
三个口径(容易混,先说清)
| 口径 | 定义 | 用途 |
|---|---|---|
| 49 项验收通过率 | 固定题库 49 题,程序化断言 | 回归防护,每次改动后必跑 |
| ad-hoc 通过率 | 人工临时提问 | 探索性发现新失败模式 |
| held-out 通过率 | 不在训练集、程序化验证的新题 | 泛化能力的唯一可信口径 |
只看验收集会系统性高估泛化:7B 曾在固定验收集 49/49,同类 ad-hoc 仅 17~50%。所以 NL2SQL/Agent 微调评估必须「固定验收集 + held-out」双轨。
三、核心代码实现细节
整条流水线是一个可分段运行的脚本 gen_v33_distill.py,6 个 stage 通过 sys.argv[1] 指定,每个 stage 都支持断点续跑(增量写、已完成的跳过):
seeds → paraphrase → label → filter → retry → assemble1. label:复刻线上 Agent 的多轮循环当老师
蒸馏的核心是让 32B 跑一遍和线上完全一样的 Agent 流程,把完整 messages 轨迹留作训练样本。label_question() 精确复刻 app_v30.do_chat() 的核心循环(只去掉安全拦截,因为种子已过滤为 query intent):
def label_question(question):
from service.app_v30 import (build_messages, call_llm, parse_tool_call, run_sql,
_assistant_msg, _tool_response, _GLOBAL_RE, LLMError, MAX_ROUNDS)
messages = build_messages(question, []) # 注入 SYSTEM_PROMPT + RAG schema + 实时在线数/日期
last_sql, last_rows = None, None
seen_sqls = set()
for round_i in range(MAX_ROUNDS):
msg = call_llm(messages) # 调 32B(localhost:30000, temp=0)
tc = parse_tool_call(msg) # 三层 fallback: OpenAI tool_calls → XML → 正则
if not tc: # 不再产出 tool_call = 回复完成
final_reply = msg.get("content", "") or ""
if last_sql is None:
return {"question": question, "success": False, "reason": "no_sql_generated"}
return {"question": question, "success": True,
"final_sql": last_sql, "final_rows": last_rows,
"final_reply": final_reply,
"messages": messages + [{"role": "assistant", "content": final_reply}],
"rounds": round_i + 1}
raw_sql = tc["arguments"].get("sql", "")
if _GLOBAL_RE.search(question): # 全局性问题剥掉 user= 过滤,否则查不到全局
raw_sql = re.sub(r"\s+WHERE\s+user\s*=\s*'[^']*'(?=GROUP|ORDER|HAVING|LIMIT|$)", ' ', raw_sql, flags=re.I)
if raw_sql in seen_sqls: # 🔴 dedup_loop 检测:又生成相同 SQL → 死循环
return {"question": question, "success": False, "reason": "dedup_loop"}
seen_sqls.add(raw_sql)
res = run_sql(raw_sql) # 🔴 SQL 真执行(MySQL 双库路由)
messages.append(_assistant_msg(msg))
if not res["success"]: # 执行失败 → 错误回传让 32B 自修正
messages.append(_tool_response(json.dumps(
{"error": res["error"], "hint": "请检查表名和列名是否正确,修正SQL后重试"})))
continue
last_sql, last_rows = raw_sql, res["rows"]
messages.append(_tool_response(json.dumps(res["rows"][:50], default=str)))
return {"question": question, "success": False, "reason": "max_rounds_exhausted"}两个细节决定成败:
- dedup_loop 检测:
seen_sqls记录每轮生成的 SQL,一旦重复直接判失败。retry 时生成完全相同 SQL 是 7B 的残留失败模式,从 50~83% 压到 10%,靠的就是这道检测把病态轨迹挡在训练集外。 - SQL 真执行是免费 Ground Truth:不需要人工标注"这个 SQL 对不对"——能执行成功且返回非空,就是 32B 在这个问题上的一次成功示范。失败案例不进训练集,只在 retry 阶段被改造成"纠错样本"。
12 worker 并发跑(ThreadPoolExecutor),写文件加锁,断点续跑靠启动时扫已完成的 question 集合。
2. filter:SQL 执行成功 = 免费过滤 + 启发式存疑
_EXPECT_NONEMPTY_RE = re.compile(r"(哪些|哪个|谁|多少|几个|排名|前\d+|top)")
def gen_filter():
labeled = [json.loads(l) for l in open(f"{OUT_DIR}/labeled.jsonl")]
ok = [r for r in labeled if r.get("success")]
suspect, clean = [], []
for r in ok:
rows = r.get("final_rows") or []
if not rows and _EXPECT_NONEMPTY_RE.search(r["question"]):
suspect.append(r) # 执行成功但空结果、且问题明显期待非空 → 存疑,不丢弃
else:
clean.append(r)
# → filtered_ok.jsonl(进训练) + filtered_suspect.jsonl(留待 triage)关键:suspect 不丢弃。执行成功+空结果有两种可能——数据本就无匹配行(真空),或 SQL 取数错了(比如"在线几个号"查了 accounts.login_status,而该列当前全为空,应该查 session_logs MAX(id) 取最新状态)。一删了之会丢掉真问题。下一轮我们对这 70 条 suspect 做 triage 分类:格式坑(user_key vs user 不兼容)补 RAG 可即时修,其余待 32B 重打标。
3. retry:错误注入 + 让 32B 自己修正
最有意思的设计。对一条已知正确的 SQL,用 mutator 注入一种真实出现过的错误模式,构造第 1 轮错误 → 把执行报错喂回 32B → 让它自己改对 → 产出"错误→自修正"的两轮轨迹:
# 复用之前确认过的真实 bug 模式(列名混淆/括号失衡/表名单复数)
_COLUMN_CONFUSION = {
"user_key": "user", "created_at": "create_time", "start_time": "begin_time",
"executed_at": "exec_time", "task_name": "taskname", "account_id": "accountid",
"status": "state", "success_count": "success_cnt", "total_count": "total_cnt",
}
def _mutate_wrong_column(sql):
for correct, wrong in _COLUMN_CONFUSION.items():
if correct in sql:
return sql.replace(correct, wrong, 1), f"column: {correct}->{wrong}"
return None, None
def _mutate_paren_imbalance(sql):
m = re.search(r"DATE\(([\w.]+)\)", sql)
if m:
pos = m.end()
return sql[:pos] + ")" + sql[pos:], "extra_paren_after_DATE()"
idx = sql.rfind(")")
if idx > 0:
return sql[:idx] + sql[idx+1:], "missing_closing_paren"
return None, None
def _mutate_wrong_table(sql):
m = re.search(r"\bFROM\s+(\w+\.)?(\w+)", sql, re.I)
if m:
table = m.group(2)
wrong = table.rstrip("s") if table.endswith("s") else table + "s"
if wrong != table:
return re.sub(rf"\b{table}\b", wrong, sql, count=1), f"table: {table}->{wrong}"
return None, None
MUTATORS = [_mutate_wrong_column, _mutate_paren_imbalance, _mutate_wrong_table]然后构造两轮轨迹:第 1 轮塞入坏 SQL + 执行错误,第 2 轮让 32B 看着错误自己改:
# 第1轮: question → bad tool_call → error tool_response
bad_res = run_sql(bad_sql)
if bad_res["success"]:
continue # mutation 没产生真实错误,跳过(必须是真能报错的坏样本)
messages = build_messages(q, [])
messages.append(bad_tc_msg) # 塞坏 SQL
messages.append(_tool_response(json.dumps({"error": bad_res["error"], "hint": "..."})))
# 第2轮:32B 看着错误自己改(不是硬塞 good_sql)
msg2 = call_llm(messages)
fixed_sql = parse_tool_call(msg2)["arguments"].get("sql", "")
fixed_res = run_sql(fixed_sql)
if not fixed_res["success"]:
continue # 32B 也没改对,不用这条
# → 保留 "错误→自修正→成功" 完整轨迹教训:曾经试过加 11 条手写 error-retry 样本,不泛化——遇到没见过的错误类型仍重复错 SQL。错误恢复能力靠机制(线上重试循环)而非靠样例灌注。retry 样本的价值是让 7B 学会"看到报错就改"的范式,而不是背某一种具体错误的修法。
4. assemble:千分位清洗 + 去重 + 静态 SYSTEM_PROMPT
三条产出汇合,做三件清洗:
_THOUSAND_SEP_RE = re.compile(r"(?<=\d),(?=\d{3}(?!\d))")
def _strip_thousand_separators(text):
"""移除数字里的千分位逗号(如 886,469 -> 886469)。
已确认: 7B 推理若训练数据含千分位逗号,会系统性在数字末尾多加 0(放大 10 倍),3B 无此现象。"""
prev = None
while prev != text: # 循环处理连续多个逗号(如 1,234,567)
prev = text
text = _THOUSAND_SEP_RE.sub("", text)
return text
def _build_sft_sample(system_prompt, tools, full_messages):
msgs = [{"role": "system", "content": system_prompt}]
for m in full_messages:
if m["role"] == "assistant" and m.get("content"):
m = {**m, "content": _strip_thousand_separators(m["content"])}
msgs.append({"role": m["role"], "content": m.get("content", "")})
return {"messages": msgs, "tools": tools}
def gen_assemble():
...
def add_from_file(path, msg_key="messages", is_retry=False):
for r in ...:
if not is_retry: # retry 即使(question,SQL)重复也保留——
key = (r["question"], r.get("final_sql") or r.get("fixed_sql"))
if key in seen_qsql: continue # 训练价值不同:一条是直接答对,一条是纠错轨迹
seen_qsql.add(key)
# raw_msgs[0] 是 build_messages 生成的 system,含实时动态注入(current_date/在线计数)
# 训练数据不该固化瞬时值 → 统一替换成静态 SYSTEM_PROMPT(推理时 build_messages 会重新生成 ctx)
sample = _build_sft_sample(SYSTEM_PROMPT, TOOLS, raw_msgs[1:])
samples.append(sample)千分位这个 bug 很隐蔽:7B 看到 886,469 训练样本,会学到"数字结尾多加个 0",于是把 886469 输出成 8864690。3B 没这毛病。解法不是改模型,是训练数据一律禁千分位逗号。
一个反直觉的订正:用宽正则
\d,\d{3}复查残留会误报 1 条——实际是坐标串<a href=5921116,1200.7328...>里的字段分隔逗号(5921116,1200),不是千分位,assemble 的正则(?<=\d),(?=\d{3}(?!\d))正确地没动它(后面是 4 位数,剥离会破坏坐标)。测量工具的宽匹配会伪装成数据缺陷——审计用的正则必须和生产清洗逻辑一致。
5. seeds + held-out:schema-grounded 种子与独立性
gen_v34_seeds.py 解决"问题多样性不够"。早期每表只有 3~5 个浅问,复杂模式(跨库 JOIN、日期范围、HAVING、子查询、排名、三跳)几乎没覆盖。改成基于真实 schema 结构化生成:
def gen_cross_db_join(joins): # 遍历 join_paths.jsonl 每条 → 生成跨表问
... # 含 gameagentforge↔fantasy 的 user_key vs user 不兼容坑
def gen_date_range(): ... # 近7天/今昨对比/CURDATE/DATE_SUB/按周聚合
def gen_having(): ... # 成功率最高/失败最多/超时>120s/重试>3
def gen_subquery_maxid(): ... # MAX(id) 取每用户最新状态(patterns.jsonl 模板衍生)
def gen_three_hop(): ... # task_ledger↔task_pool↔accounts 三跳统计historical 问从 qa_conversations.jsonl(intent=query)混入——真实分布最有代表性。去重后历史问仅 165 条,这本身就印证了"分布多样性有限,模式覆盖是瓶颈"。
held-out 必须独立于蒸馏数据(否则测不出"32B 的错误被 7B 继承")。早期发现 ragas_dataset(1261)是 RAG grounding 评测(与 seeds 同源、重叠可接受),不是泛化 held-out。真正 held-out 靠 stratified carve:
random.seed(42)
by_pat = defaultdict(list)
for s in seeds: by_pat[s["pattern"]].append(s)
held_out = []
for pat, lst in by_pat.items(): # 每个 pattern 分层抽 ~15%,stratified 覆盖
random.shuffle(lst)
k = max(2, min(25, len(lst) // 7))
held_out.extend(lst[:k])
held_norms = set(norm(s["question"]) for s in held_out)
train_seeds = [s for s in seeds if norm(s["question"]) not in held_norms] # 构造保证不进训练校验:carved held-out 与 v34 seeds 重叠 = 0(构造保证);与现有 v33 训练重叠 = 5(对当前 7B 污染,71 条干净;对下一轮 v34-7B 全部 76 条干净)。held-out 从 n=10 扩到 76 且分层均衡。
四、方法论教训
- 高分先怀疑口径,不是先庆祝。7B 固定验收集 49/49、同类 ad-hoc 仅 17~50%——差点得出"7B 已达标"的错误结论。
- 效果崩塌的第一嫌疑人是"数据/检索链路静默降级",不是模型。
.237→.34迁移漏拷 GTE embedding 权重(205MB 只拷了 config),RAG 静默退化为 BM25-only,系统照常返回答案只是答错。一整阶段的失败被误归因为"SFT 数据不够"。 - 错误恢复靠机制,不靠样例。加手写 retry 样本不泛化;线上重试循环 + 蒸馏"纠错轨迹"才解决。
- 换更新的基座不等于更好。同数据训 Qwen3-8B 劣于 Qwen2.5-7B(think 标签泄露 + 功能回归)。
- 测量工具的缺陷会伪装成被测对象的缺陷(同上千分位坐标串案例)。
五、效果与下一轮
| 口径 | 值 |
|---|---|
| 49 项验收(32B 零样本) | 49/49 |
| 49 项验收(1895 条蒸馏后 7B) | 49/49 |
| held-out 新题(蒸馏后 7B) | 9/10 = 90%(n=10,样本偏小) |
| held-out 新题(旧手写 7B) | 17~50% |
下一轮(已铺好路):GPU 授权后拉起 32B(TP=4)→ paraphrase + label 466 条 schema-grounded 新种子 → filter → retry → assemble → merge 产出 sft_v34_full → LoRA(r=16/alpha=32/lr=1e-4)微调 Qwen2.5-7B → 双轨 eval(49 验收 + 76 条 clean held-out + 7B-vs-32B gap)。
7B 与 32B 的差距 = 蒸馏数据覆盖面还差多少,这个差值本身就是最有用的进度指标。瓶颈从来不是参数量,而是训练数据的覆盖完整性。