250 lines
8.1 KiB
Python
250 lines
8.1 KiB
Python
"""
|
||
Text2SQL Few-shot 集成示例
|
||
|
||
展示如何在现有系统中集成经验数据集few-shot功能
|
||
"""
|
||
|
||
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:修改 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 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()
|