Files
2026-04-14 10:28:22 +08:00

257 lines
8.2 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
Text2SQL Few-shot 集成示例
展示如何在现有系统中集成经验数据集few-shot功能
"""
import sys
from pathlib import Path
_root = Path(__file__).resolve().parent.parent
if str(_root / "backend") not in sys.path:
sys.path.insert(0, str(_root / "backend"))
from typing import Optional, List
def enhance_orchestrator_with_fewshot(orchestrator_class):
"""
装饰器:为Text2SQLOrchestrator添加few-shot功能
用法:
from agents.orchestrator import Text2SQLOrchestrator
EnhancedOrchestrator = enhance_orchestrator_with_fewshot(Text2SQLOrchestrator)
orch = EnhancedOrchestrator(..., fewshot_enabled=True, fewshot_top_k=3)
"""
class FewShotEnhancedOrchestrator(orchestrator_class):
"""增强版Orchestrator,支持few-shot"""
def __init__(
self,
*args,
fewshot_enabled: bool = True,
fewshot_samples_path: Optional[str] = None,
fewshot_top_k: int = 3,
fewshot_min_rating: int = 7,
**kwargs
):
super().__init__(*args, **kwargs)
self.fewshot_enabled = fewshot_enabled
self.fewshot_top_k = fewshot_top_k
self.fewshot_min_rating = fewshot_min_rating
self.fewshot_selector = None
if fewshot_enabled:
from utils.fewshot_selector import FewShotSelector
path = fewshot_samples_path or "data/experiences/all_samples.jsonl"
self.fewshot_selector = FewShotSelector(path)
print(f"✅ Few-shot已启用: top_k={fewshot_top_k}, min_rating={fewshot_min_rating}")
def _generate_sql_with_fewshot(
self,
question: str,
schema_str: str,
dialect: str = "tsql"
) -> str:
"""带few-shot的SQL生成"""
from llm.deepseek_client import DeepSeekConfig
# 获取few-shot示例
examples_prompt = ""
if self.fewshot_selector:
examples = self.fewshot_selector.select(
question=question,
top_k=self.fewshot_top_k,
min_rating=self.fewshot_min_rating,
exclude_qids=[] # 可排除当前问题(如果已存在)
)
if examples:
examples_prompt = "\n\n".join([
f"示例 {i+1}:\n问题:{ex.question_zh}\nSQL:\n{ex.sql}"
for i, ex in enumerate(examples)
])
examples_prompt = f"参考以下相似示例:\n\n{examples_prompt}\n\n"
# 构造增强的Prompt
from config.prompts import SQL_GENERATOR_SYSTEM, SQL_GENERATOR_USER
user_content = f"""{examples_prompt}
当前问题:{question}
Schema信息:
{schema_str}
数据库方言:{dialect}
请参考上述示例的SQL编写风格,生成**有用 SQL**:
- 与示例保持相同的版式(大写关键字、4空格缩进、PascalCase别名)
- 根据当前Schema选择正确的表和字段
- 输出纯SQL,不要其他内容"""
messages = [
{"role": "system", "content": SQL_GENERATOR_SYSTEM},
{"role": "user", "content": user_content}
]
# 调用LLM
config = DeepSeekConfig(
api_key=self.deepseek.config.api_key,
base_url=self.deepseek.config.base_url,
model_name=self.deepseek.config.model_name,
temperature=0.3,
max_tokens=4096
)
from llm.deepseek_client import DeepSeekClient
client = DeepSeekClient(config)
response = client.chat(messages)
sql = response.content.strip()
# 清理代码块
if "```sql" in sql:
sql = sql[sql.find("```sql")+6:sql.find("```", sql.find("```sql")+6)].strip()
elif "```" in sql:
sql = sql[sql.find("```")+3:sql.find("```", sql.find("```")+3)].strip()
from utils.sql_parser import normalize_sql_for_dialect
return normalize_sql_for_dialect(sql, dialect)
# 重写 _generate_sql 方法
def _generate_sql(
self,
question: str,
schema_str: str,
dialect: str = "tsql"
) -> str:
if self.fewshot_enabled and self.fewshot_selector:
return self._generate_sql_with_fewshot(question, schema_str, dialect)
else:
return super()._generate_sql(question, schema_str, dialect)
return FewShotEnhancedOrchestrator
# ========== 集成步骤 ==========
def integrate_fewshot():
"""
集成few-shot功能的完整步骤
"""
print("=" * 60)
print("Few-shot 集成指南")
print("=" * 60)
steps = """
**步骤1:确保经验数据集已生成**
```bash
cd text2sql_agent_camel
python scripts/parse_examples.py
```
生成: data/experiences/all_samples.jsonl
**步骤2:修改 backend/agents/orchestrator.py**
在文件开头添加:
```python
from utils.fewshot_selector import FewShotSelector
```
修改 Text2SQLOrchestrator.__init__:
```python
def __init__(self, ..., fewshot_enabled=True, fewshot_top_k=3, ...):
# ... 原有代码 ...
self.fewshot_enabled = fewshot_enabled
self.fewshot_top_k = fewshot_top_k
self.fewshot_selector = None
if fewshot_enabled:
self.fewshot_selector = FewShotSelector(
os.getenv("FEWSHOT_DATA_PATH", "data/experiences/all_samples.jsonl")
)
logger.info(f"Few-shot已启用: top_k={fewshot_top_k}")
```
修改 _generate_sql 方法(约第251行):
```python
def _generate_sql(self, question: str, schema_str: str, dialect: str = "tsql") -> str:
# 添加few-shot增强
if self.fewshot_enabled and self.fewshot_selector:
examples = self.fewshot_selector.select(
question=question,
top_k=self.fewshot_top_k,
min_rating=7,
exclude_qids=[] # 可选:排除当前问题对应的QID
)
if examples:
examples_prompt = "\\n\\n".join([
f"示例({ex.qid}):\\n问题:{ex.question_zh}\\nSQL:\\n{ex.sql}"
for ex in examples
])
# 将examples_prompt注入到用户消息中
# (需要调整 deepseek_client 调用逻辑)
# ... 原有生成逻辑 ...
```
**步骤3:配置环境变量(可选)**
```bash
# .env 文件添加
FEWSHOT_ENABLED=true
FEWSHOT_TOP_K=3
FEWSHOT_MIN_RATING=7
FEWSHOT_DATA_PATH=./data/experiences/all_samples.jsonl
```
**步骤4:测试效果**
```bash
# 对比测试
python backend/main.py "查询2024年1月的销售额" --verbose
# 观察日志中的 "Few-shot已选择: Q31, Q6, Q44"
```
**步骤5:A/B测试评估**
```python
# 创建对比脚本
from agents.orchestrator import Text2SQLOrchestrator
# 无few-shot
orch_baseline = Text2SQLOrchestrator(..., fewshot_enabled=False)
result1 = orch_baseline.generate(question)
# 有few-shot
orch_fewshot = Text2SQLOrchestrator(..., fewshot_enabled=True, fewshot_top_k=3)
result2 = orch_fewshot.generate(question)
print(f"Baseline: {result1.sql[:100]}")
print(f"Few-shot: {result2.sql[:100]}")
```
"""
print(steps)
def quick_test_fewshot():
"""快速测试few-shot选择器"""
from utils.fewshot_selector import FewShotSelector
selector = FewShotSelector("data/experiences/all_samples.jsonl")
test_questions = [
"查询2024年1月的销售额",
"统计所有活跃账户数量",
"按对手方列出未结算交易",
"查询今天的现金余额",
]
print("\nFew-shot选择测试:\n")
for q in test_questions:
print(f"问题: {q}")
examples = selector.select(q, top_k=2, min_rating=8)
for ex in examples:
print(f" [{ex.qid}] {ex.question_zh[:40]}... (rating={ex.rating})")
print()
if __name__ == "__main__":
integrate_fewshot()
# quick_test_fewshot()