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