AI工程 8 分钟阅读

用 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 → assemble

1. 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 且分层均衡。

四、方法论教训

  1. 高分先怀疑口径,不是先庆祝。7B 固定验收集 49/49、同类 ad-hoc 仅 17~50%——差点得出"7B 已达标"的错误结论。
  2. 效果崩塌的第一嫌疑人是"数据/检索链路静默降级",不是模型.237→.34 迁移漏拷 GTE embedding 权重(205MB 只拷了 config),RAG 静默退化为 BM25-only,系统照常返回答案只是答错。一整阶段的失败被误归因为"SFT 数据不够"。
  3. 错误恢复靠机制,不靠样例。加手写 retry 样本不泛化;线上重试循环 + 蒸馏"纠错轨迹"才解决。
  4. 换更新的基座不等于更好。同数据训 Qwen3-8B 劣于 Qwen2.5-7B(think 标签泄露 + 功能回归)。
  5. 测量工具的缺陷会伪装成被测对象的缺陷(同上千分位坐标串案例)。

五、效果与下一轮

口径
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 的差距 = 蒸馏数据覆盖面还差多少,这个差值本身就是最有用的进度指标。瓶颈从来不是参数量,而是训练数据的覆盖完整性。