Files

257 lines
8.2 KiB
Python
Raw Permalink Normal View History

2026-04-10 16:52:07 +08:00
"""
Text2SQL Few-shot 集成示例
展示如何在现有系统中集成经验数据集few-shot功能
"""
2026-04-14 10:28:22 +08:00
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"))
2026-04-10 16:52:07 +08:00
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
2026-04-14 10:28:22 +08:00
**步骤2:修改 backend/agents/orchestrator.py**
2026-04-10 16:52:07 +08:00
在文件开头添加:
```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
# 对比测试
2026-04-14 10:28:22 +08:00
python backend/main.py "查询2024年1月的销售额" --verbose
2026-04-10 16:52:07 +08:00
# 观察日志中的 "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()