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()
|