首页
学习
活动
专区
圈层
工具
发布
社区首页 >专栏 >开源模型层出不穷,一线工程师该怎样做好模型调优与业务落地

开源模型层出不穷,一线工程师该怎样做好模型调优与业务落地

原创
作者头像
大盘鸡拌面
发布2026-08-23 20:09:52
发布2026-08-23 20:09:52
550
举报

说个扎心的事实:2024 年到现在,光 Qwen 一个系列就发布了不下二十个版本,DeepSeek、GLM、Yi、Baichuan 轮番上阵,每周都有"最强开源模型"的帽子换了主人。你在 HuggingFace 上刷到一篇新论文,激动得想换模型,结果一上业务测试集,效果跟上个版本半斤八两。 问题不是模型不够好,是你没有一套自己的调优和落地方法论。这篇文章就是讲这个的。

我不是来给你安利哪个模型的——今天最强的模型下周可能就被超越。我想聊的是:不管模型怎么换,你该怎么搭一套能持续评估、持续调优、安全上线的体系。


一、先搞清楚一件事:调优不是玄学,是工程

很多人理解的"模型调优"就是改改 Prompt、换换 temperature,看几条 case 感觉"好像好了一点"就算完事。这跟算命差不太多。

真正的调优是工程化的:有评估集、有指标、有基线、有对比、有版本管理。 你改了一个参数,得能说出"准确率从 72% 提升到 78%,P95 延迟增加了 200ms",而不是"感觉更顺了"。

先画个全局图,看看调优在落地流程中的位置:

这套流程看着多,但每一环都有存在的道理。下面我用一个完整的业务场景把它走一遍。


二、业务场景:企业 Text-to-SQL 助手

为什么选这个场景

数据团队每天收到提数需求——"帮我查一下上个月华东区客单价下降超过 20% 的品类"。分析师手写 SQL,快的十分钟、慢的半小时。如果是业务同学自己提需求走工单流程,来回沟通至少半天。

Text-to-SQL 是一个特别适合用大模型的场景:输入是自然语言,输出是结构化的 SQL,而且结果可验证——跑一下 SQL 对不对就知道。不像"帮我写个方案"这种开放性任务,Text-to-SQL 有明确的对错标准。

但也是个特别容易翻车的场景:表名记错了、关联条件写反了、聚合维度漏了——任何一个错都导致结果完全不对。

评估集先行
代码语言:javascript
复制
import json
import sqlite3
import time
from dataclasses import dataclass, field
from typing import List, Dict, Tuple
import requests

@dataclass
class SQLTestCase:
    """Text-to-SQL 评估用例"""
    id: str
    question: str           # 自然语言问题
    db_schema: str          # 数据库 schema 描述
    expected_sql: str       # 标准答案 SQL
    difficulty: str        # easy / medium / hard
    category: str           # 单表查询/多表关联/子查询/窗口函数等
    expected_result_check: str = ""  # 预期结果特征描述


# 评估集:覆盖不同难度和类型
SQL_TEST_CASES = [
    # === Easy: 单表基础查询 ===
    SQLTestCase(
        id="sql-001",
        question="查询订单表中状态为已完成的订单数量",
        db_schema="orders(id, user_id, status, amount, created_at) status值: pending/paid/shipped/completed/cancelled",
        expected_sql="SELECT COUNT(*) FROM orders WHERE status = 'completed'",
        difficulty="easy",
        category="单表聚合"
    ),
    SQLTestCase(
        id="sql-002",
        question="最近7天每天的新增用户数",
        db_schema="users(id, name, phone, created_at, status)",
        expected_sql="""SELECT DATE(created_at) as dt, COUNT(*) as new_users 
                        FROM users 
                        WHERE created_at >= DATE('now', '-7 days') 
                        GROUP BY DATE(created_at) 
                        ORDER BY dt DESC""",
        difficulty="easy",
        category="时间范围聚合"
    ),

    # === Medium: 多表关联 ===
    SQLTestCase(
        id="sql-003",
        question="查询每个用户的订单总金额,按金额降序排列,只看前10名",
        db_schema="""users(id, name, phone) 
                    orders(id, user_id, amount, status, created_at) 
                    关联: orders.user_id = users.id""",
        expected_sql="""SELECT u.name, SUM(o.amount) as total_amount 
                        FROM users u 
                        JOIN orders o ON u.id = o.user_id 
                        WHERE o.status = 'completed'
                        GROUP BY u.id, u.name 
                        ORDER BY total_amount DESC 
                        LIMIT 10""",
        difficulty="medium",
        category="多表关联+聚合+排序"
    ),
    SQLTestCase(
        id="sql-004",
        question="上个月客单价(总金额/订单数)低于前一个月的地区有哪些",
        db_schema="""orders(id, user_id, amount, status, created_at)
                    users(id, name, region, phone)
                    关联: orders.user_id = users.id
                    status: 只有 completed 算有效订单""",
        expected_sql="""-- 这个查询需要同比计算,比较两个月的数据
                        WITH last_month AS (
                            SELECT u.region, SUM(o.amount)/COUNT(*) as avg_order
                            FROM orders o JOIN users u ON o.user_id = u.id
                            WHERE o.status = 'completed' 
                              AND o.created_at >= DATE('now','start of month','-1 month')
                              AND o.created_at < DATE('now','start of month')
                            GROUP BY u.region
                        ),
                        prev_month AS (
                            SELECT u.region, SUM(o.amount)/COUNT(*) as avg_order
                            FROM orders o JOIN users u ON o.user_id = u.id
                            WHERE o.status = 'completed'
                              AND o.created_at >= DATE('now','start of month','-2 month')
                              AND o.created_at < DATE('now','start of month','-1 month')
                            GROUP BY u.region
                        )
                        SELECT l.region 
                        FROM last_month l JOIN prev_month p ON l.region = p.region
                        WHERE l.avg_order < p.avg_order""",
        difficulty="medium",
        category="多表关联+同比计算"
    ),

    # === Hard: 复杂业务逻辑 ===
    SQLTestCase(
        id="sql-005",
        question="找出连续3天都有下单行为的用户",
        db_schema="""orders(id, user_id, amount, status, created_at)
                    status: completed 才算有效""",
        expected_sql="""-- 连续N天问题,经典窗口函数场景
                        WITH daily_orders AS (
                            SELECT DISTINCT user_id, DATE(created_at) as dt
                            FROM orders WHERE status = 'completed'
                        ),
                        grouped AS (
                            SELECT user_id, dt,
                                   DATE(dt, '-' || (
                                       ROW_NUMBER() OVER(PARTITION BY user_id ORDER BY dt) - 1
                                   ) || ' days') as grp
                            FROM daily_orders
                        )
                        SELECT DISTINCT user_id 
                        FROM grouped 
                        GROUP BY user_id, grp 
                        HAVING COUNT(*) >= 3""",
        difficulty="hard",
        category="窗口函数+连续行为"
    ),
    SQLTestCase(
        id="sql-006",
        question="每个品类销量前三的商品,如果并列也算",
        db_schema="""products(id, name, category, price) 
                    orders(id, product_id, quantity, amount, status, created_at)
                    关联: orders.product_id = products.id
                    status: completed""",
        expected_sql="""WITH ranked AS (
                            SELECT p.category, p.name, 
                                   SUM(o.quantity) as total_qty,
                                   RANK() OVER(PARTITION BY p.category 
                                               ORDER BY SUM(o.quantity) DESC) as rnk
                            FROM products p 
                            JOIN orders o ON p.id = o.product_id
                            WHERE o.status = 'completed'
                            GROUP BY p.category, p.name
                        )
                        SELECT category, name, total_qty 
                        FROM ranked WHERE rnk <= 3""",
        difficulty="hard",
        category="窗口函数+分组TopN"
    ),
]


class TextToSQLEvaluator:
    """Text-to-SQL 模型评估器"""

    def __init__(self, model_url: str, model_name: str):
        self.url = model_url
        self.model = model_name
        # 测试数据库(内存 SQLite,用于验证 SQL 执行不报错)
        self.test_db = self._init_test_db()

    def _init_test_db(self):
        """初始化测试数据库,造一些假数据"""
        conn = sqlite3.connect(":memory:")
        conn.executescript("""
            CREATE TABLE users(id INT, name TEXT, phone TEXT, 
                              region TEXT, created_at TEXT, status TEXT);
            CREATE TABLE orders(id INT, user_id INT, product_id INT,
                               amount REAL, quantity INT, status TEXT, created_at TEXT);
            CREATE TABLE products(id INT, name TEXT, category TEXT, price REAL);
            
            INSERT INTO users VALUES 
                (1,'张三','138xxxx','华东','2024-01-15','active'),
                (2,'李四','139xxxx','华南','2024-02-20','active'),
                (3,'王五','137xxxx','华东','2024-03-10','active');
            
            INSERT INTO products VALUES
                (101,'商品A','电子',99.9),(102,'商品B','食品',29.9),
                (103,'商品C','电子',199.0),(104,'商品D','服装',59.0);
            
            INSERT INTO orders VALUES
                (1,1,101,99.9,1,'completed','2024-06-01'),
                (2,1,102,59.8,2,'completed','2024-06-02'),
                (3,2,103,199.0,1,'completed','2024-06-01'),
                (4,3,104,59.0,1,'pending','2024-06-03');
        """)
        return conn

    def generate_sql(self, question: str, schema: str) -> str:
        """让模型生成 SQL"""
        prompt = f"""你是一个 SQL 专家,请根据问题生成 SQLite SQL。

数据库表结构:
{schema}

问题:{question}

要求:
1. 只输出 SQL 语句,不要解释
2. 使用 SQLite 语法
3. 复杂查询使用 CTE (WITH 子句) 提高可读性
"""

        response = requests.post(
            f"{self.url}/v1/chat/completions",
            json={
                "model": self.model,
                "messages": [{"role": "user", "content": prompt}],
                "temperature": 0.0,  # SQL 生成用 0 温度,要确定性
                "max_tokens": 800,
            },
            timeout=60
        )
        raw = response.json()["choices"][0]["message"]["content"]
        # 清理 markdown 包裹
        sql = raw.strip()
        if sql.startswith("```"):
            sql = sql.split("\n", 1)[1] if "\n" in sql else sql[3:]
            if sql.endswith("```"):
                sql = sql[:-3]
        return sql.strip()

    def evaluate(self, test_cases: List[SQLTestCase]) -> Dict:
        """完整评估"""
        results = []
        for tc in test_cases:
            start = time.time()
            generated_sql = self.generate_sql(tc.question, tc.db_schema)
            latency = time.time() - start

            # 评估维度 1: SQL 语法正确性(能否执行不报错)
            syntax_ok = self._check_syntax(generated_sql)

            # 评估维度 2: 执行结果匹配(在测试库上跑,看结果是否正确)
            result_match = False
            if syntax_ok:
                result_match = self._check_result(generated_sql, tc.expected_sql)

            # 评估维度 3: SQL 结构相似度(简单字符串相似度)
            similarity = self._sql_similarity(generated_sql, tc.expected_sql)

            results.append({
                "case_id": tc.id,
                "difficulty": tc.difficulty,
                "category": tc.category,
                "syntax_ok": syntax_ok,
                "result_match": result_match,
                "similarity": round(similarity, 2),
                "latency_ms": round(latency * 1000),
                "generated_sql": generated_sql[:200],
            })

        # 汇总统计
        total = len(results)
        syntax_pass = sum(r["syntax_ok"] for r in results) / total
        result_pass = sum(r["result_match"] for r in results) / total
        avg_latency = sum(r["latency_ms"] for r in results) / total

        # 按难度和类型分组
        by_diff = {}
        for r in results:
            d = r["difficulty"]
            if d not in by_diff:
                by_diff[d] = {"total": 0, "syntax": 0, "result": 0}
            by_diff[d]["total"] += 1
            by_diff[d]["syntax"] += 1 if r["syntax_ok"] else 0
            by_diff[d]["result"] += 1 if r["result_match"] else 0

        report = f"""
============ Text-to-SQL 评估报告 ============
模型: {self.model}
测试用例: {total} 条

【核心指标】
语法正确率: {syntax_pass:.1%}
结果匹配率: {result_pass:.1%}  ← 这个最重要
平均延迟: {avg_latency:.0f}ms

【按难度分组】
"""
        for diff in ["easy", "medium", "hard"]:
            if diff in by_diff:
                s = by_diff[diff]
                report += (f"  {diff}: 语法 {s['syntax']}/{s['total']}, "
                          f"结果 {s['result']}/{s['total']}\n")

        # 失败用例
        failures = [r for r in results if not r["result_match"]]
        if failures:
            report += f"\n【失败用例 ({len(failures)})】\n"
            for f in failures:
                status = "语法错误" if not f["syntax_ok"] else "结果不匹配"
                report += f"  [{f['case_id']}] {f['difficulty']}/{f['category']} - {status}\n"

        return {"report": report, "details": results}

    def _check_syntax(self, sql: str) -> bool:
        """检查 SQL 能否在测试库上执行不报错"""
        try:
            self.test_db.execute(f"EXPLAIN {sql}")
            return True
        except Exception:
            return False

    def _check_result(self, generated: str, expected: str) -> bool:
        """对比生成 SQL 和标准 SQL 的执行结果"""
        try:
            cur1 = self.test_db.execute(generated)
            cur2 = self.test_db.execute(expected)
            r1 = sorted([str(row) for row in cur1.fetchall()])
            r2 = sorted([str(row) for row in cur2.fetchall()])
            return r1 == r2
        except Exception:
            return False

    def _sql_similarity(self, s1: str, s2: str) -> float:
        """简单的 token 级 Jaccard 相似度"""
        t1 = set(s1.lower().split())
        t2 = set(s2.lower().split())
        if not t1 and not t2:
            return 1.0
        return len(t1 & t2) / len(t1 | t2) if t1 | t2 else 0


# ===== 基线评估 =====
evaluator = TextToSQLEvaluator(
    model_url="http://localhost:8000",
    model_name="qwen2.5-14b-instruct"
)

result = evaluator.evaluate(SQL_TEST_CASES)
print(result["report"])

# ===== 基线结果示例 =====
# ============ Text-to-SQL 评估报告 ============
# 模型: qwen2.5-14b-instruct
# 测试用例: 6 条
#
# 【核心指标】
# 语法正确率: 83.3%   (5/6)
# 结果匹配率: 50.0%   (3/6)
# 平均延迟: 650ms
#
# 【按难度分组】
#   easy:   语法 2/2, 结果 2/2
#   medium: 语法 2/2, 结果 1/2
#   hard:   语法 1/2, 结果 0/2
#
# 【失败用例 (3)】
#   [sql-004] medium/多表关联+同比计算 - 结果不匹配
#   [sql-005] hard/窗口函数+连续行为 - 语法错误
#   [sql-006] hard/窗口函数+分组TopN - 结果不匹配

看到结果了吧?easy 全过,medium 差强人意,hard 直接拉胯。没有评估集,你根本不知道问题出在哪。有了它,你知道接下来该干什么:攻克窗口函数和复杂关联。


三、调优策略:按问题类型对症下药

评估结果告诉你"哪里不行",但"怎么改"需要选对策略。我的经验是:不同类型的错误用不同的手段,别一上来就微调

策略 1:Few-shot + 结构化输出(解决语法和格式问题)

基线评估发现 sql-005 的语法直接错了——模型不熟悉 SQLite 的日期计算语法。加几个示例比改 Prompt 措辞管用得多:

代码语言:javascript
复制
class FewShotTuner:
    """Few-shot 调优器:根据错误类型自动选择示例"""

    # 示例库:按查询类型分类
    EXAMPLE_BANK = {
        "窗口函数": [
            {
                "question": "每个用户消费金额最高的那笔订单",
                "schema": "orders(id, user_id, amount, status, created_at)",
                "sql": """SELECT * FROM (
                    SELECT o.*, 
                           ROW_NUMBER() OVER(PARTITION BY user_id ORDER BY amount DESC) as rn
                    FROM orders o WHERE status = 'completed'
                ) WHERE rn = 1"""
            },
            {
                "question": "每个地区销售额排名第一的商品",
                "schema": "products(id, name, category) orders(id, product_id, amount, region, status)",
                "sql": """WITH ranked AS (
                    SELECT region, product_id, SUM(amount) as total,
                           RANK() OVER(PARTITION BY region ORDER BY SUM(amount) DESC) as rnk
                    FROM orders WHERE status = 'completed'
                    GROUP BY region, product_id
                ) SELECT * FROM ranked WHERE rnk = 1"""
            }
        ],
        "日期计算": [
            {
                "question": "上个月的总销售额",
                "schema": "orders(id, amount, status, created_at)",
                "sql": """SELECT SUM(amount) FROM orders 
                         WHERE status = 'completed' 
                         AND created_at >= DATE('now','start of month','-1 month')
                         AND created_at < DATE('now','start of month')"""
            },
            {
                "question": "最近30天每天的下单量趋势",
                "schema": "orders(id, status, created_at)",
                "sql": """SELECT DATE(created_at) as dt, COUNT(*) as cnt 
                         FROM orders 
                         WHERE created_at >= DATE('now','-30 days')
                         GROUP BY DATE(created_at) ORDER BY dt"""
            }
        ],
        "连续行为": [
            {
                "question": "连续2天访问的用户",
                "schema": "user_visits(id, user_id, visit_date)",
                "sql": """WITH grouped AS (
                    SELECT user_id, visit_date,
                           DATE(visit_date, '-' || (
                               ROW_NUMBER() OVER(PARTITION BY user_id ORDER BY visit_date) - 1
                           ) || ' days') as grp
                    FROM (SELECT DISTINCT user_id, visit_date FROM user_visits)
                )
                SELECT DISTINCT user_id FROM grouped 
                GROUP BY user_id, grp HAVING COUNT(*) >= 2"""
            }
        ]
    }

    def build_few_shot_prompt(self, question: str, schema: str, 
                               error_type: str) -> str:
        """根据错误类型选择最相关的示例"""
        
        # 简单的关键词匹配选示例
        relevant_examples = []
        for category, examples in self.EXAMPLE_BANK.items():
            if self._is_relevant(question, category, error_type):
                relevant_examples.extend(examples)

        # 没匹配到就用通用示例
        if not relevant_examples:
            relevant_examples = [
                self.EXAMPLE_BANK["日期计算"][0],
                self.EXAMPLE_BANK["窗口函数"][0]
            ]

        # 构造 Few-shot Prompt
        examples_text = ""
        for ex in relevant_examples[:3]:  # 最多 3 个示例,控制 token 数
            examples_text += (
                f"\n--- 示例 ---\n"
                f"问题: {ex['question']}\n"
                f"表结构: {ex['schema']}\n"
                f"SQL: {ex['sql']}\n"
            )

        prompt = f"""你是一个 SQL 专家。请参考以下示例,根据问题生成 SQLite SQL。

{examples_text}

--- 请回答以下问题 ---
表结构: {schema}
问题: {question}

要求:
1. 只输出 SQL,不要解释
2. 使用 SQLite 语法
3. 参考 示例 的写法风格
"""
        return prompt

    def _is_relevant(self, question: str, category: str, error_type: str) -> bool:
        """判断示例类别是否与当前问题相关"""
        keywords_map = {
            "窗口函数": ["排名", "最高", "前N", "每个", "Top"],
            "日期计算": ["上个月", "最近", "同比", "环比", "趋势", "每天"],
            "连续行为": ["连续", "累计", "每天"],
        }
        keywords = keywords_map.get(category, [])
        return any(kw in question for kw in keywords)


# ===== 调优后重新评估 =====
tuner = FewShotTuner()

# 对失败用例用 Few-shot 重新生成
for tc_id in ["sql-005", "sql-006"]:
    tc = next(t for t in SQL_TEST_CASES if t.id == tc_id)
    
    # 用 Few-shot prompt
    prompt = tuner.build_few_shot_prompt(
        tc.question, tc.db_schema, 
        error_type="窗口函数" if "窗口" in tc.category else "复杂查询"
    )
    
    response = requests.post(
        "http://localhost:8000/v1/chat/completions",
        json={
            "model": "qwen2.5-14b-instruct",
            "messages": [{"role": "user", "content": prompt}],
            "temperature": 0.0,
            "max_tokens": 800,
        },
        timeout=60
    )
    new_sql = response.json()["choices"][0]["message"]["content"]
    print(f"\n{tc_id} 调优后 SQL:")
    print(new_sql[:300])

# 调优后评估结果:
# 语法正确率: 83.3% → 100%    (6/6)
# 结果匹配率: 50.0% → 83.3%   (5/6)
# 平均延迟: 650ms → 980ms     (Few-shot 更长,延迟增加)

Few-shot 加上之后,语法全过了,结果匹配率从 50% 涨到 83%。但 sql-004(同比计算)还是不行——这种需要理解业务逻辑"上月 vs 上上月"的,示例管不了,得让模型先理思路。

策略 2:Chain-of-Thought(解决复杂逻辑问题)
代码语言:javascript
复制
class CoTTuner:
    """思维链调优:让模型先分析再生成"""

    def generate_with_cot(self, question: str, schema: str) -> str:
        prompt = f"""你是 SQL 专家,请按以下步骤生成 SQL。

表结构:
{schema}

问题:{question}

请按以下格式输出:
1. 分析:这个问题需要查询什么数据,涉及哪些表,需要什么关联条件
2. 思路:先做什么计算,再做什么计算,是否需要子查询或 CTE
3. 注意:有没有边界情况需要处理(空值、去重、时间范围)
4. SQL:最终的 SQLite SQL 语句

示例格式:
1. 分析:需要查询用户表和订单表关联,按地区分组计算客单价
2. 思路:先过滤有效订单(status=completed),再按地区分组求 sum(amount)/count(*),最后两个月做对比
3. 注意:时间范围要精确到月初月末,除零保护
4. SQL:WITH ... SELECT ...
"""

        response = requests.post(
            "http://localhost:8000/v1/chat/completions",
            json={
                "model": self.model,
                "messages": [{"role": "user", "content": prompt}],
                "temperature": 0.0,
                "max_tokens": 1500,  # CoT 需要更多 token
            },
            timeout=90
        )
        full_output = response.json()["choices"][0]["message"]["content"]
        
        # 提取 SQL 部分
        if "4. SQL:" in full_output or "4. SQL:" in full_output:
            sql_part = full_output.split("4. SQL")[1].strip(":: \n")
            if sql_part.startswith("```"):
                sql_part = sql_part.split("\n", 1)[1] if "\n" in sql_part else sql_part[3:]
                if sql_part.endswith("```"):
                    sql_part = sql_part[:-3]
            return sql_part.strip()
        return full_output  # 提取失败就返回原文


# CoT 调优后的效果:
# sql-004 (同比计算): 模型先分析了"需要两个月对比",思路对了,SQL 结果匹配 ✓
# sql-005 (连续行为): 模型分析"需要窗口函数分组",SQL 语法正确,结果匹配 ✓
# 结果匹配率: 83.3% → 100% (6/6)
# 平均延迟: 980ms → 1450ms (CoT 多生成推理内容,延迟显著增加)

CoT 一上,准确率到 100% 了。但延迟从 650ms 涨到 1450ms——这就是调优的核心矛盾:准确率和延迟的 trade-off


四、工程落地:让调优成果稳定跑在生产线上

调优搞定了,但生产环境跟测试环境完全两回事。用户的问题千奇百怪,schema 可能比你的测试集复杂十倍,并发来了你怎么扛?

生产架构:分层处理 + 安全兜底
代码语言:javascript
复制
import re
import hashlib
from dataclasses import dataclass
from typing import Optional
import requests
import json
import time

@dataclass
class QueryRequest:
    """用户提数请求"""
    user_id: str
    question: str
    tables: list          # 用户有权限访问的表
    session_id: str
    timestamp: float = None


@dataclass
class QueryResult:
    """提数结果"""
    sql: str
    sql_explanation: str       # SQL 的自然语言解释
    confidence: float
    needs_review: bool         # 是否需要人工确认
    cache_hit: bool = False
    latency_ms: int = 0


class ProductionTextToSQLService:
    """
    生产级 Text-to-SQL 服务
    设计原则:
    1. 简单问题走缓存/小模型,复杂问题走大模型——成本控制
    2. 所有生成的 SQL 必须经过安全校验——防止删库/全表扫描
    3. 低置信度走人工确认——安全兜底
    """

    # SQL 安全规则
    FORBIDDEN_KEYWORDS = [
        "DROP", "DELETE", "UPDATE", "INSERT", "ALTER", 
        "CREATE", "TRUNCATE", "ATTACH", "DETACH"
    ]
    
    # 全表扫描检测:没有 WHERE 的 SELECT
    FULL_TABLE_SCAN_THRESHOLD = 100000  # 预估行数超过此值必须有 WHERE

    def __init__(self, config: dict):
        # 分层模型配置
        self.fast_model = config["fast_model"]    # 7B,处理简单查询
        self.full_model = config["full_model"]    # 14B+CoT,处理复杂查询
        self.fast_url = config["fast_url"]
        self.full_url = config["full_url"]
        
        # 缓存
        self.cache = {}  # 生产环境用 Redis
        self.cache_ttl = config.get("cache_ttl", 3600)
        
        # Schema 管理器
        self.schema_manager = config["schema_manager"]

    def process(self, request: QueryRequest) -> QueryResult:
        """主处理流程"""
        start = time.time()
        request.timestamp = start

        # 1. 缓存检查
        cache_key = self._make_cache_key(request)
        if cache_key in self.cache:
            cached = self.cache[cache_key]
            if time.time() - cached["timestamp"] < self.cache_ttl:
                return QueryResult(
                    sql=cached["sql"],
                    sql_explanation=cached["explanation"],
                    confidence=cached["confidence"],
                    needs_review=False,
                    cache_hit=True,
                    latency_ms=round((time.time() - start) * 1000)
                )

        # 2. 获取相关 Schema(不是全部表,只取相关的)
        relevant_schema = self.schema_manager.get_relevant_tables(
            request.question, request.tables
        )

        # 3. 复杂度评估:决定走哪个模型
        complexity = self._estimate_complexity(request.question, relevant_schema)

        if complexity == "simple":
            result = self._generate_with_fast_model(
                request.question, relevant_schema)
        else:
            result = self._generate_with_full_model(
                request.question, relevant_schema)

        # 4. 安全校验
        security_check = self._security_check(result["sql"], relevant_schema)
        if not security_check["passed"]:
            return QueryResult(
                sql="",
                sql_explanation=f"安全校验未通过:{security_check['reason']}",
                confidence=0,
                needs_review=True,
                latency_ms=round((time.time() - start) * 1000)
            )

        # 5. 置信度评估
        needs_review = result["confidence"] < 0.75

        # 6. 缓存写入
        self.cache[cache_key] = {
            "sql": result["sql"],
            "explanation": result["explanation"],
            "confidence": result["confidence"],
            "timestamp": time.time()
        }

        return QueryResult(
            sql=result["sql"],
            sql_explanation=result["explanation"],
            confidence=result["confidence"],
            needs_review=needs_review,
            cache_hit=False,
            latency_ms=round((time.time() - start) * 1000)
        )

    def _estimate_complexity(self, question: str, schema: str) -> str:
        """评估查询复杂度,决定模型路由"""
        # 复杂特征关键词
        complex_keywords = [
            "连续", "同比", "环比", "排名", "前N", "Top",
            "占比", "累计", "中位数", "分位数",
            "每个", "分别", "对比", "趋势"
        ]
        
        # 多表关联特征
        multi_table = schema.count("CREATE TABLE") > 1 or "JOIN" in question.upper()
        
        # 判断逻辑
        complex_score = sum(1 for kw in complex_keywords if kw in question)
        
        if complex_score >= 2 or (multi_table and complex_score >= 1):
            return "complex"
        return "simple"

    def _generate_with_fast_model(self, question, schema):
        """简单查询:用 7B 快速生成"""
        prompt = f"""根据问题生成 SQLite SQL。

表结构:{schema}

问题:{question}

只输出 SQL,不要解释。"""
        
        response = requests.post(
            f"{self.fast_url}/v1/chat/completions",
            json={
                "model": self.fast_model,
                "messages": [{"role": "user", "content": prompt}],
                "temperature": 0.0,
                "max_tokens": 500,
            },
            timeout=15
        )
        sql = self._extract_sql(response.json()["choices"][0]["message"]["content"])
        
        return {
            "sql": sql,
            "explanation": self._explain_sql(sql, question),
            "confidence": 0.85  # 简单查询给高置信度
        }

    def _generate_with_full_model(self, question, schema):
        """复杂查询:用 14B + CoT"""
        prompt = f"""你是 SQL 专家,请按以下步骤分析并生成 SQL。

表结构:{schema}

问题:{question}

输出格式:
1. 分析:需要什么数据,涉及哪些表,关联条件
2. 思路:计算步骤,是否需要 CTE/子查询
3. 注意:边界情况处理
4. SQL:SQLite 语句
5. 置信度:0.0-1.0(你对这个 SQL 正确性的信心)"""

        response = requests.post(
            f"{self.full_url}/v1/chat/completions",
            json={
                "model": self.full_model,
                "messages": [{"role": "user", "content": prompt}],
                "temperature": 0.0,
                "max_tokens": 1500,
            },
            timeout=90
        )
        output = response.json()["choices"][0]["message"]["content"]
        
        sql = self._extract_from_cot(output)
        confidence = self._extract_confidence(output)
        
        return {
            "sql": sql,
            "explanation": self._extract_analysis(output),
            "confidence": confidence
        }

    def _security_check(self, sql: str, schema: str) -> dict:
        """SQL 安全校验"""
        sql_upper = sql.upper()
        
        # 检查禁止的关键词
        for kw in self.FORBIDDEN_KEYWORDS:
            if kw in sql_upper:
                return {"passed": False, 
                        "reason": f"检测到禁止操作: {kw}"}
        
        # 检查是否有 WHERE(全表扫描防护)
        if "SELECT" in sql_upper and "WHERE" not in sql_upper:
            if "COUNT(*)" not in sql_upper and "LIMIT" not in sql_upper:
                return {"passed": False, 
                        "reason": "缺少 WHERE 条件,可能全表扫描"}
        
        # 检查 LIMIT(防止返回过多数据)
        if "SELECT" in sql_upper and "LIMIT" not in sql_upper:
            # 自动补充 LIMIT
            sql = sql.rstrip(";") + " LIMIT 1000;"
        
        return {"passed": True, "sql": sql}

    def _make_cache_key(self, request: QueryRequest) -> str:
        """生成缓存 key:问题 + 可用表 的 hash"""
        key_str = f"{request.question}|{'|'.join(sorted(request.tables))}"
        return hashlib.md5(key_str.encode()).hexdigest()

    def _extract_sql(self, text: str) -> str:
        """从模型输出中提取 SQL"""
        text = text.strip()
        if text.startswith("```"):
            lines = text.split("\n")
            text = "\n".join(lines[1:-1]) if len(lines) > 2 else text[3:]
        return text.strip("`").strip()

    def _extract_from_cot(self, output: str) -> str:
        """从 CoT 输出中提取 SQL"""
        for marker in ["4. SQL:", "4. SQL:", "SQL:", "SQL:"]:
            if marker in output:
                sql_part = output.split(marker)[1].strip()
                for next_marker in ["5.", "置信度"]:
                    if next_marker in sql_part:
                        sql_part = sql_part.split(next_marker)[0].strip()
                return self._extract_sql(sql_part)
        return self._extract_sql(output)

    def _extract_confidence(self, output: str) -> float:
        """提取置信度"""
        import re
        match = re.search(r'置信度[::]\s*([0-9.]+)', output)
        if match:
            return float(match.group(1))
        return 0.5  # 默认中等置信度

    def _extract_analysis(self, output: str) -> str:
        """提取分析部分作为解释"""
        for marker in ["1. 分析:", "1. 分析:", "分析:", "分析:"]:
            if marker in output:
                analysis = output.split(marker)[1].split("2.")[0].strip()
                return analysis
        return ""

    def _explain_sql(self, sql: str, question: str) -> str:
        """生成 SQL 的自然语言解释"""
        return f"此 SQL 用于回答:{question}"


# ===== 生产部署示例 =====
service = ProductionTextToSQLService(config={
    "fast_model": "qwen2.5-7b-instruct",
    "full_model": "qwen2.5-14b-instruct",
    "fast_url": "http://localhost:8001",
    "full_url": "http://localhost:8002",
    "cache_ttl": 3600,
    "schema_manager": type("SM", (), {
        "get_relevant_tables": lambda self, q, tables: 
            "users(id, name, region, created_at)\norders(id, user_id, amount, status, created_at)"
    })()
})

# 测试简单查询(走 fast 模型)
req1 = QueryRequest(
    user_id="analyst_01",
    question="查询已完成订单的数量",
    tables=["orders"],
    session_id="sess-001"
)
result1 = service.process(req1)
print(f"简单查询: {result1.sql}")
print(f"  延迟: {result1.latency_ms}ms, 置信度: {result1.confidence}")
print(f"  需要人工确认: {result1.needs_review}")

# 测试复杂查询(走 full 模型 + CoT)
req2 = QueryRequest(
    user_id="analyst_01",
    question="每个地区上个月和上上个月的客单价对比,找出下降的地区",
    tables=["users", "orders"],
    session_id="sess-002"
)
result2 = service.process(req2)
print(f"\n复杂查询: {result2.sql[:200]}")
print(f"  延迟: {result2.latency_ms}ms, 置信度: {result2.confidence}")
print(f"  需要人工确认: {result2.needs_review}")
print(f"  解释: {result2.sql_explanation}")

# 输出示例:
# 简单查询: SELECT COUNT(*) FROM orders WHERE status = 'completed' LIMIT 1000;
#   延迟: 320ms, 置信度: 0.85
#   需要人工确认: False
#
# 复杂查询: WITH last_month AS (...), prev_month AS (...) SELECT ... 
#   延迟: 1380ms, 置信度: 0.82
#   需要人工确认: False
#   解释: 需要查询用户表和订单表关联,按地区分组计算两个月客单价并对比

这套生产架构的完整时序:


五、模型更迭:新模型出来了,要不要换?

这是最实际的问题。Qwen 从 2.0 到 2.5 到 3.0,DeepSeek 从 V2 到 V3,每次更新都说自己更强。你到底要不要换?

我的原则很简单:用你的业务评估集跑一遍,达标了才换。 别信论文里的数字。

代码语言:javascript
复制
class ModelMigrationGate:
    """
    模型迁移检查站:新模型必须通过业务评估才能替换旧模型
    流程:新模型评估 → 与旧模型对比 → A/B 测试 → 灰度 → 全量切换
    """

    def __init__(self, eval_cases, old_model_config, new_model_config):
        self.cases = eval_cases
        self.old_config = old_model_config
        self.new_config = new_model_config

    def run_comparison(self) -> dict:
        """对比新旧模型"""
        print("=" * 60)
        print("模型迁移评估")
        print(f"旧模型: {self.old_config['model_name']}")
        print(f"新模型: {self.new_config['model_name']}")
        print("=" * 60)

        # 用同一套评估集分别跑
        old_evaluator = TextToSQLEvaluator(
            self.old_config["url"], self.old_config["model_name"])
        new_evaluator = TextToSQLEvaluator(
            self.new_config["url"], self.new_config["model_name"])

        old_result = old_evaluator.evaluate(self.cases)
        new_result = new_evaluator.evaluate(self.cases)

        old_details = old_result["details"]
        new_details = new_result["details"]

        # 逐条对比
        comparison = []
        for old, new in zip(old_details, new_details):
            comparison.append({
                "case_id": old["case_id"],
                "difficulty": old["difficulty"],
                "old_result_match": old["result_match"],
                "new_result_match": new["result_match"],
                "old_latency": old["latency_ms"],
                "new_latency": new["latency_ms"],
                "improved": new["result_match"] and not old["result_match"],
                "regressed": old["result_match"] and not new["result_match"],
            })

        # 汇总
        old_pass = sum(c["old_result_match"] for c in comparison)
        new_pass = sum(c["new_result_match"] for c in comparison)
        improved = sum(c["improved"] for c in comparison)
        regressed = sum(c["regressed"] for c in comparison)
        old_latency = sum(c["old_latency"] for c in comparison) / len(comparison)
        new_latency = sum(c["new_latency"] for c in comparison) / len(comparison)

        # 迁移决策
        should_migrate = (new_pass > old_pass and 
                         regressed == 0 and
                         new_latency <= old_latency * 1.5)

        report = f"""
============ 模型迁移对比报告 ============

【核心指标对比】
                    旧模型          新模型
结果匹配率:        {old_pass}/{len(comparison)} ({old_pass/len(comparison):.0%})    {new_pass}/{len(comparison)} ({new_pass/len(comparison):.0%})
平均延迟:          {old_latency:.0f}ms          {new_latency:.0f}ms

【变化分析】
提升的用例: {improved} 个
退步的用例: {regressed} 个

【迁移决策】
{'✅ 建议迁移:准确率提升且无退步' if should_migrate else '❌ 暂缓迁移'}
"""
        if not should_migrate:
            if new_pass <= old_pass:
                report += "原因:准确率未提升\n"
            if regressed > 0:
                report += f"原因:有 {regressed} 个用例退步\n"
            if new_latency > old_latency * 1.5:
                report += f"原因:延迟增加超过 50%({old_latency:.0f}→{new_latency:.0f}ms)\n"

        report += "\n【逐条对比】\n"
        for c in comparison:
            symbol = "→" if c["old_result_match"] == c["new_result_match"] else \
                     "↑" if c["improved"] else "↓"
            report += (f"  {c['case_id']} [{c['difficulty']}] "
                      f"{c['old_result_match']} {symbol} {c['new_result_match']}  "
                      f"延迟: {c['old_latency']}→{c['new_latency']}ms\n")

        return {"report": report, "should_migrate": should_migrate,
                "comparison": comparison}


# ===== 新模型评估 =====
gate = ModelMigrationGate(
    eval_cases=SQL_TEST_CASES,
    old_model_config={"url": "http://localhost:8000", 
                      "model_name": "qwen2.5-14b-instruct"},
    new_model_config={"url": "http://localhost:8001", 
                      "model_name": "qwen3-14b-instruct"}  # 假设新出了 Qwen3
)

result = gate.run_comparison()
print(result["report"])

# 示例输出:
# ============ 模型迁移对比报告 ============
# 
# 【核心指标对比】
#                     旧模型          新模型
# 结果匹配率:        5/6 (83%)      6/6 (100%)
# 平均延迟:          980ms          750ms
# 
# 【变化分析】
# 提升的用例: 1 个
# 退步的用例: 0 个
# 
# 【迁移决策】
# ✅ 建议迁移:准确率提升且无退步
# 
# 【逐条对比】
#   sql-001 [easy]    True → True   延迟: 320→280ms
#   sql-002 [easy]    True → True   延迟: 350→310ms
#   sql-003 [medium] True → True   延迟: 680→520ms
#   sql-004 [medium] False ↑ True  延迟: 950→720ms  ← 新模型搞定了同比计算
#   sql-005 [hard]   True → True   延迟: 1450→980ms
#   sql-006 [hard]   True → True   延迟: 1130→790ms

六、踩过的坑,你可以绕着走

坑 1:没有评估集就开调。 这是最基本的错误。没有尺子你怎么量身高?评估集不需要大,20-50 条覆盖核心场景就够,但必须有。

坑 2:用通用 benchmark 判断业务效果。 Spider(Text-to-SQL 标准 benchmark)上跑分高的模型,到你自己的表结构上可能拉胯——因为你的表名是 ​​t_order_detail_2024_h1​​ 这种不规范命名。业务评估集必须用你自己的真实数据。

坑 3:CoT 万能论。 CoT 确实能提升复杂查询准确率,但延迟翻倍是实打实的。简单查询不需要 CoT,分层路由很重要。我们的方案是简单查询走 7B 无 CoT(320ms),复杂查询走 14B + CoT(1380ms),总体平均延迟控制在 500ms 以内。

坑 4:不校验直接执行 SQL。 模型生成的 SQL 你敢直接丢数据库跑?万一它生成了 ​​DELETE FROM orders​​ 怎么办?安全校验不是可选项,是必须项。我们加了禁止关键词检查、全表扫描检查、自动 LIMIT 补充三道防线。

坑 5:缓存命中率低。 一开始我们的缓存 key 用的是完整问题文本,结果"查询已完成订单数"和"查一下已完成订单有多少"就缓存不到一起。后来改成语义缓存(先用小模型做问题归一化,再查缓存),命中率从 8% 涨到 35%。

坑 6:置信度阈值设得不合理。 一开始设 0.9 才自动执行,结果 60% 的查询都走人工确认,用户体验很差。后来降到 0.75,自动执行率提到 85%,错误率只增加了 2 个百分点。这个阈值得根据你的业务容错度调,没有标准答案。

坑 7:不跟踪线上效果。 上线后不监控 = 裸奔。我们做了三个看板:准确率趋势(每周抽样人工标注 100 条算准确率)、延迟分位数(P50/P95/P99)、用户反馈率(点"结果不对"的比例)。准确率连续 3 天低于基线 5 个百分点就自动告警。


开源模型层出不穷是好事——竞争推动进步。但对一线工程师来说,每出一个新模型就要重新评估一遍,这个负担不轻。所以搭一套可复用的评估和调优体系,比追新模型重要得多

这篇文章的核心思路就三句话:

  1. 评估先行——没有评估集,一切调优都是玄学
  2. 按需调优——不同类型的错误用不同策略,别一上来就微调
  3. 工程兜底——模型可能出错,但你的系统不能跟着错。安全校验、置信度路由、人工兜底、缓存优化——这些工程手段才是生产系统稳定运行的基石。

模型会变,方法论不会。把评估集、调优策略库、安全校验规则、A/B 测试框架这些基建搭好了,不管明天出的是 Qwen3 还是 DeepSeek V4,你都能从容地跑一遍评估、做个决策、安全地迁移。

这才是"做好模型调优与业务落地"的真正含义——不是追着模型跑,是让模型在你的体系里跑。

原创声明:本文系作者授权腾讯云开发者社区发表,未经许可,不得转载。

如有侵权,请联系 cloudcommunity@tencent.com 删除。

目录
  • 一、先搞清楚一件事:调优不是玄学,是工程
  • 二、业务场景:企业 Text-to-SQL 助手
    • 为什么选这个场景
    • 评估集先行
  • 三、调优策略:按问题类型对症下药
    • 策略 1:Few-shot + 结构化输出(解决语法和格式问题)
    • 策略 2:Chain-of-Thought(解决复杂逻辑问题)
  • 四、工程落地:让调优成果稳定跑在生产线上
    • 生产架构:分层处理 + 安全兜底
  • 五、模型更迭:新模型出来了,要不要换?
  • 六、踩过的坑,你可以绕着走
问题归档专栏文章快讯文章归档关键词归档开发者手册归档开发者手册 Section 归档