first commit

This commit is contained in:
陈辅元
2026-04-10 16:52:07 +08:00
commit 84fe545640
87 changed files with 20842 additions and 0 deletions
+249
View File
@@ -0,0 +1,249 @@
"""
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()
+439
View File
@@ -0,0 +1,439 @@
#!/usr/bin/env python3
"""
将 Example_text2sql.md 解析为结构化数据集
输出格式:
1. JSONL格式(每行一个样本)
2. 按难度/类别分类
3. 包含Schema、SQL、执行结果、评分
"""
import re
import json
import yaml
from pathlib import Path
from typing import List, Dict, Optional, Any
from dataclasses import dataclass, asdict
from datetime import datetime
@dataclass
class Text2SQLSample:
"""单条Text2SQL样本"""
qid: str # 问题ID,如 "Q1"
question_zh: str # 中文问题
sql: str # SQL语句
schema_info: Dict[str, Any] # Schema信息(表、字段)
explanation: str = "" # SQL理由
question_en: Optional[str] = None # 英文问题(如有)
rating: Optional[int] = None # 准确度评分 1-10
execution_result: Optional[List[Dict]] = None # 执行结果(表格)
result_note: str = "" # 结果说明(如"当前查询无结果")
tags: List[str] = None # 标签(如["聚合", "JOIN", "日期筛选"])
difficulty: str = "medium" # 难度:easy/medium/hard
def __post_init__(self):
if self.tags is None:
self.tags = []
def to_dict(self):
return asdict(self)
def to_fewshot_example(self) -> str:
"""转换为few-shot示例字符串"""
return f"""问题:{self.question_zh}
Schema:{json.dumps(self.schema_info, ensure_ascii=False)[:500]}
SQL:
{self.sql}
"""
class ExampleText2SQLParser:
"""解析 Example_text2sql.md 文件"""
def __init__(self, md_path: str):
self.md_path = Path(md_path)
self.content = self.md_path.read_text(encoding='utf-8')
self.samples: List[Text2SQLSample] = []
def parse(self) -> List[Text2SQLSample]:
"""解析整个MD文件"""
# 按问题分割(## Qxx. 标题)
q_blocks = re.split(r'\n## Q(\d+)\. ', self.content)
# q_blocks[0] 是文件头(目录、固定日期说明等),跳过
for i in range(1, len(q_blocks)):
qid_num = q_blocks[i].split('\n')[0].strip() # 从分割中获取问题号
# 实际qid应为 "Q{num}"
# 重新用正则匹配完整的Q标题
pass
# 更稳健的方法:直接用正则查找所有Q标题块
pattern = r'(## Q(\d+)\. .*?)(?=\n## Q\d+\. |\Z)'
matches = re.findall(pattern, self.content, re.DOTALL)
print(f"找到 {len(matches)} 个问答对")
for full_match, qid in matches:
sample = self._parse_single_block(qid, full_match)
if sample:
self.samples.append(sample)
return self.samples
def _parse_single_block(self, qid: str, block: str) -> Optional[Text2SQLSample]:
"""解析单个问答块"""
try:
lines = block.strip().split('\n')
# 提取问题(支持中英文双语)
# 格式:### 问题 后跟英文段落,空行,然后中文段落
question_zh = ""
question_en = ""
found_question_header = False
paragraph_lines = []
paragraphs = []
for line in lines:
if line.startswith('### 问题'):
found_question_header = True
continue
if not found_question_header:
continue
if line.startswith('###'):
# 遇到下一节,停止
break
if line.strip() == '':
# 空行:段落结束
if paragraph_lines:
paragraphs.append(' '.join(paragraph_lines).strip())
paragraph_lines = []
else:
paragraph_lines.append(line.strip())
# 最后一个段落
if paragraph_lines:
paragraphs.append(' '.join(paragraph_lines).strip())
# 第一段为英文,第二段为中文
if len(paragraphs) >= 1:
question_en = paragraphs[0]
if len(paragraphs) >= 2:
question_zh = paragraphs[1]
# 提取SQL
sql_match = re.search(r'```sql\n(.*?)\n```', block, re.DOTALL)
if not sql_match:
return None
sql = sql_match.group(1).strip()
# 提取SQL理由(### SQL理由 后面的内容)
reason_match = re.search(r'### SQL理由\n\n(.*?)(?=\n### |\Z)', block, re.DOTALL)
explanation = reason_match.group(1).strip() if reason_match else ""
# 提取评分(### SQL准确度评分)
rating_match = re.search(r'\*\*评分:\s*(\d+)/10\*\*', block)
rating = int(rating_match.group(1)) if rating_match else None
# 提取结果说明(### SQL执行结果 后面)
result_match = re.search(r'### SQL执行结果\n\n(.*?)(?=\n### |\Z)', block, re.DOTALL)
result_note = result_match.group(1).strip() if result_match else ""
# 推断标签
tags = self._infer_tags(question_zh, sql)
# 推断难度
difficulty = self._infer_difficulty(sql, tags)
# 构建Schema信息(从SQL中提取的表和字段)
schema_info = self._extract_schema_from_sql(sql)
return Text2SQLSample(
qid=f"Q{qid}",
question_zh=question_zh.strip(),
question_en=question_en.strip() if question_en else None,
sql=sql,
schema_info=schema_info,
explanation=explanation,
rating=rating,
result_note=result_note,
tags=tags,
difficulty=difficulty
)
except Exception as e:
print(f"解析 {qid} 失败: {e}")
return None
def _infer_tags(self, question: str, sql: str) -> List[str]:
"""推断问题标签"""
tags = []
q_lower = question.lower()
sql_upper = sql.upper()
# 按业务域
if any(k in q_lower for k in ['对手方', 'broker', '经纪商']):
tags.append('broker')
if any(k in q_lower for k in ['账户', 'account', '客户']):
tags.append('account')
if any(k in q_lower for k in ['现金', 'cash', '余额']):
tags.append('cash')
if any(k in q_lower for k in ['持仓', 'holding', 'position']):
tags.append('holding')
if any(k in q_lower for k in ['交易', 'trade', 'commission']):
tags.append('trade')
if any(k in q_lower for k in ['IPO', '认购', 'subscription']):
tags.append('ipo')
if any(k in q_lower for k in ['公司行动', 'corporate action', '权益']):
tags.append('corporate_action')
if any(k in q_lower for k in ['汇率', 'exchange rate']):
tags.append('exchange_rate')
if any(k in q_lower for k in ['保证金', 'margin']):
tags.append('margin')
if any(k in q_lower for k in ['利息', 'interest']):
tags.append('interest')
if any(k in q_lower for k in ['报表', 'statement']):
tags.append('statement')
# 按SQL特征
if 'GROUP BY' in sql_upper:
tags.append('aggregation')
if 'JOIN' in sql_upper:
tags.append('join')
if 'WHERE' in sql_upper:
tags.append('filter')
if re.search(r'DATE.*BETWEEN|>=.*AND', sql_upper) or 'BETWEEN' in sql_upper:
tags.append('date_range')
if re.search(r'\bCOUNT\b|\bSUM\b|\bAVG\b|\bMAX\b|\bMIN\b', sql_upper):
tags.append('aggregation')
if 'ORDER BY' in sql_upper:
tags.append('ordering')
if 'TOP' in sql_upper or 'LIMIT' in sql_upper:
tags.append('limit')
if 'LEFT JOIN' in sql_upper:
tags.append('left_join')
if 'UNION' in sql_upper:
tags.append('union')
if re.search(r'DATEADD|DATEDIFF|CAST.*GETDATE', sql_upper):
tags.append('date_function')
return list(set(tags))
def _infer_difficulty(self, sql: str, tags: List[str]) -> str:
"""推断难度"""
sql_upper = sql.upper()
# 简单规则
if 'UNION' in sql_upper or 'EXCEPT' in sql_upper or 'INTERSECT' in sql_upper:
return 'hard'
if sql.count('JOIN') >= 3:
return 'hard'
if 'GROUP BY' in sql_upper and 'HAVING' in sql_upper:
return 'hard'
if re.search(r'SUBQUERY|EXISTS.*SELECT', sql_upper, re.IGNORECASE):
return 'hard'
if len(tags) >= 5:
return 'hard'
if sql.count('JOIN') == 2 or 'GROUP BY' in sql_upper:
return 'medium'
return 'easy'
def _extract_schema_from_sql(self, sql: str) -> Dict[str, Any]:
"""从SQL中提取表结构信息"""
import sqlglot
from sqlglot import exp
schema = {
"tables": {},
"joins": []
}
try:
parsed = sqlglot.parse_one(sql, dialect='tsql')
# 提取所有表
for table in parsed.find_all(exp.Table):
table_name = table.name
alias = table.alias_or_name
schema["tables"][table_name] = {
"alias": alias,
"columns": []
}
# 提取所有列引用
for col in parsed.find_all(exp.Column):
table_name = col.table
col_name = col.name
if table_name in schema["tables"]:
if col_name not in schema["tables"][table_name]["columns"]:
schema["tables"][table_name]["columns"].append(col_name)
# 提取JOIN关系
for join in parsed.find_all(exp.Join):
schema["joins"].append({
"type": join.args.get('side', 'INNER'),
"on": str(join.args.get('on'))
})
except Exception as e:
print(f"SQL解析失败: {e}")
return schema
def save_jsonl(self, output_path: Path):
"""保存为JSONL格式"""
with output_path.open('w', encoding='utf-8') as f:
for sample in self.samples:
f.write(json.dumps(sample.to_dict(), ensure_ascii=False) + '\n')
print(f"[OK] 保存 {len(self.samples)} 条样本到 {output_path}")
def save_fewshot_prompt(self, output_path: Path, limit_per_tag: int = 2):
"""生成few-shot提示文件"""
examples_by_tag = {}
for sample in self.samples:
for tag in sample.tags:
if tag not in examples_by_tag:
examples_by_tag[tag] = []
if len(examples_by_tag[tag]) < limit_per_tag:
examples_by_tag[tag].append(sample)
with output_path.open('w', encoding='utf-8') as f:
f.write("# Few-shot Examples for Text2SQL\n\n")
f.write("以下是从经验数据集中精选的示例,用于辅助SQL生成。\n\n")
for tag, samples in sorted(examples_by_tag.items()):
f.write(f"## 标签: {tag}\n\n")
for i, sample in enumerate(samples, 1):
f.write(f"### 示例 {i}\n")
f.write(f"**问题**:{sample.question_zh}\n\n")
f.write(f"**SQL**:\n```sql\n{sample.sql}\n```\n\n")
f.write(f"**说明**:{sample.explanation[:200]}...\n\n")
f.write("\n")
print(f"[OK] Few-shot示例已保存到 {output_path}")
def save_by_category(self, output_dir: Path):
"""按标签分类保存"""
output_dir.mkdir(parents=True, exist_ok=True)
by_tag = {}
for sample in self.samples:
for tag in sample.tags:
if tag not in by_tag:
by_tag[tag] = []
by_tag[tag].append(sample)
for tag, samples in by_tag.items():
tag_file = output_dir / f"{tag}.jsonl"
with tag_file.open('w', encoding='utf-8') as f:
for sample in samples:
f.write(json.dumps(sample.to_dict(), ensure_ascii=False) + '\n')
print(f"[OK] 按标签分类保存到 {output_dir}({len(by_tag)}个类别)")
def generate_stats(self) -> Dict:
"""生成数据集统计"""
stats = {
"total_samples": len(self.samples),
"by_difficulty": {},
"by_tag": {},
"by_rating": {},
"avg_rating": 0,
}
# 难度分布
for diff in ['easy', 'medium', 'hard']:
stats["by_difficulty"][diff] = sum(
1 for s in self.samples if s.difficulty == diff
)
# 标签分布
tag_counts = {}
for sample in self.samples:
for tag in sample.tags:
tag_counts[tag] = tag_counts.get(tag, 0) + 1
stats["by_tag"] = tag_counts
# 评分分布
ratings = [s.rating for s in self.samples if s.rating is not None]
if ratings:
stats["avg_rating"] = sum(ratings) / len(ratings)
for r in range(1, 11):
stats["by_rating"][r] = ratings.count(r)
return stats
def main():
md_path = Path(__file__).parent.parent / "data" / "Example" / "Example_text2sql.md"
output_dir = Path(__file__).parent.parent / "data" / "experiences"
print("=" * 60)
print("Text2SQL 经验数据集构建工具")
print("=" * 60)
parser = ExampleText2SQLParser(md_path)
samples = parser.parse()
print(f"\n[OK] 解析完成,共 {len(samples)} 个样本")
# 生成统计
stats = parser.generate_stats()
print(f"\n[INFO] 数据集统计:")
print(f" 总样本数: {stats['total_samples']}")
print(f" 难度分布: {stats['by_difficulty']}")
print(f" 平均评分: {stats['avg_rating']:.1f}/10")
print(f"\n 标签分布(Top 10):")
sorted_tags = sorted(stats["by_tag"].items(), key=lambda x: -x[1])[:10]
for tag, count in sorted_tags:
print(f" {tag}: {count}")
# 保存输出
output_dir.mkdir(parents=True, exist_ok=True)
# 1. JSONL格式(所有样本)
parser.save_jsonl(output_dir / "all_samples.jsonl")
# 2. 按难度分类
for diff in ['easy', 'medium', 'hard']:
diff_samples = [s for s in samples if s.difficulty == diff]
if diff_samples:
with (output_dir / f"difficulty_{diff}.jsonl").open('w', encoding='utf-8') as f:
for s in diff_samples:
f.write(json.dumps(s.to_dict(), ensure_ascii=False) + '\n')
# 3. Few-shot提示文件
parser.save_fewshot_prompt(output_dir / "fewshot_examples.md")
# 4. 按标签分类
parser.save_by_category(output_dir / "by_tag")
# 5. 高评分样本(rating >= 8)
high_rating = [s for s in samples if s.rating and s.rating >= 8]
with (output_dir / "high_rating_samples.jsonl").open('w', encoding='utf-8') as f:
for s in high_rating:
f.write(json.dumps(s.to_dict(), ensure_ascii=False) + '\n')
print(f"\n[OK] 高评分样本(rating>=8): {len(high_rating)} 条")
# 6. 生成README
with (output_dir / "README.md").open('w', encoding='utf-8') as f:
f.write("# Text2SQL 经验数据集\n\n")
f.write(f"**来源**: `Example_text2sql.md`\n")
f.write(f"**样本数**: {len(samples)}\n")
f.write(f"**生成时间**: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n\n")
f.write("## 文件说明\n\n")
f.write("- `all_samples.jsonl` - 所有样本(JSONL格式)\n")
f.write("- `difficulty_*.jsonl` - 按难度分类\n")
f.write("- `by_tag/` - 按标签分类(每个标签一个文件)\n")
f.write("- `fewshot_examples.md` - Few-shot示例(Markdown格式)\n")
f.write("- `high_rating_samples.jsonl` - 高评分样本(rating≥8)\n\n")
f.write("## 使用建议\n\n")
f.write("1. **Few-shot prompting**: 使用 `fewshot_examples.md` 中的示例\n")
f.write("2. **训练数据**: 使用 `all_samples.jsonl` 进行模型微调\n")
f.write("3. **评估测试**: 使用 `high_rating_samples.jsonl` 作为测试集\n")
f.write("4. **领域适应**: 按标签选择相关示例(如broker相关查询)\n")
print(f"\n[OK] 所有文件已保存到: {output_dir}")
if __name__ == "__main__":
main()
+237
View File
@@ -0,0 +1,237 @@
#!/usr/bin/env python3
"""
自动集成few-shot功能到 Text2SQL 系统
用法:
python scripts/patch_fewshot.py --enable --top-k 3 --min-rating 7
"""
import argparse
import re
from pathlib import Path
def patch_orchestrator(
enable: bool = True,
top_k: int = 3,
min_rating: int = 7,
data_path: str = "data/experiences/all_samples.jsonl"
):
"""
修改 agents/orchestrator.py 添加few-shot支持
"""
orch_path = Path(__file__).parent.parent / "agents" / "orchestrator.py"
if not orch_path.exists():
print(f"❌ 文件不存在: {orch_path}")
return False
content = orch_path.read_text(encoding='utf-8')
# 1. 在文件头部添加import
import_line = "from utils.fewshot_selector import FewShotSelector"
if import_line not in content:
# 找到合适的插入点(在其他import之后)
last_import = content.rfind("import ")
if last_import != -1:
insert_pos = content.find("\n", last_import) + 1
content = content[:insert_pos] + import_line + "\n" + content[insert_pos:]
print("✅ 添加 import: from utils.fewshot_selector import FewShotSelector")
# 2. 修改 __init__ 方法
init_pattern = r'def __init__\([\s\S]*?\):'
init_match = re.search(init_pattern, content)
if init_match:
init_text = content[init_match.start():init_match.end()]
# 检查是否已添加few-shot参数
if "fewshot_enabled" not in init_text:
# 在参数列表末尾添加(在max_retry之后)
new_params = """,
# Few-shot 配置
fewshot_enabled: bool = True,
fewshot_samples_path: Optional[str] = None,
fewshot_top_k: int = 3,
fewshot_min_rating: int = 7,"""
# 找到max_retry参数后插入
max_retry_pos = init_text.find("max_retry:")
if max_retry_pos != -1:
comma_pos = init_text.find(",", max_retry_pos) + 1
new_init = (
init_text[:comma_pos] +
new_params +
init_text[comma_pos:]
)
content = content[:init_match.start()] + new_init + content[init_match.end():]
print("✅ 修改 __init__: 添加few-shot参数")
# 在__init__方法体内添加few-shot初始化代码
# 找到"self.max_retry = max_retry"这一行
max_retry_assign = "self.max_retry = max_retry"
if max_retry_assign in content[init_match.start():init_match.end()]:
init_body_start = init_match.end()
init_body = content[init_body_start:init_body_start + 2000] # 读取前2000字符
if "self.fewshot_selector" not in init_body:
fewshot_init = '''
# Few-shot 初始化
self.fewshot_enabled = fewshot_enabled
self.fewshot_top_k = fewshot_top_k
self.fewshot_min_rating = fewshot_min_rating
self.fewshot_selector = None
if self.fewshot_enabled:
try:
from utils.fewshot_selector import FewShotSelector
path = fewshot_samples_path or os.getenv(
"FEWSHOT_DATA_PATH",
"./data/experiences/all_samples.jsonl"
)
self.fewshot_selector = FewShotSelector(path)
logger.info(
f"Few-shot已启用: top_k={fewshot_top_k}, "
f"min_rating={fewshot_min_rating}"
)
except Exception as e:
logger.warning(f"Few-shot加载失败: {e},将使用标准生成")
self.fewshot_enabled = False
'''
# 插入到self.max_retry赋值之后
insert_pos = content[init_match.start():].find(max_retry_assign)
if insert_pos != -1:
absolute_pos = init_match.start() + insert_pos + len(max_retry_assign) + 1
content = content[:absolute_pos] + fewshot_init + content[absolute_pos:]
print("✅ 修改 __init__: 添加few-shot初始化")
# 3. 修改 _generate_sql 方法
generate_pattern = r'def _generate_sql\([\s\S]*?\) -> str:'
generate_match = re.search(generate_pattern, content)
if generate_match:
generate_text = content[generate_match.start():generate_match.end()]
if "if self.fewshot_enabled" not in generate_text:
# 在方法开头注入few-shot逻辑
fewshot_logic = ''' # Few-shot 增强
if self.fewshot_enabled and self.fewshot_selector:
try:
examples = self.fewshot_selector.select(
question=question,
top_k=self.fewshot_top_k,
min_rating=self.fewshot_min_rating
)
if examples:
examples_prompt = "\\n".join([
f"示例 {i+1}:\\n问题:{ex.question_zh}\\nSQL:\\n{ex.sql}"
for i, ex in enumerate(examples)
])
examples_prompt = (
f"参考以下相似问题的SQL示例:\\n\\n{examples_prompt}\\n\\n"
f"请学习上述示例的SQL编写风格和模式,为当前问题生成SQL。"
)
# 注入到schema_str前或user_content中
# 这里选择在schema_str前添加
schema_str = f"{examples_prompt}Schema信息:\\n{schema_str}"
logger.debug(f"已注入 {len(examples)} 个few-shot示例")
except Exception as e:
logger.warning(f"Few-shot检索失败: {e}")
'''
# 插入到方法开头("dialect_label = dialect" 之前)
dialect_pos = content[generate_match.start():].find("dialect_label = dialect")
if dialect_pos != -1:
absolute_pos = generate_match.start() + dialect_pos
content = content[:absolute_pos] + fewshot_logic + content[absolute_pos:]
print("✅ 修改 _generate_sql: 注入few-shot逻辑")
# 保存修改
backup_path = orch_path.with_suffix('.py.bak')
orch_path.write_text(content, encoding='utf-8')
print(f"✅ 已修改 {orch_path}")
print(f" 备份文件: {backup_path}")
return True
def create_env_config(top_k: int, min_rating: int):
"""更新.env文件添加few-shot配置"""
env_path = Path(__file__).parent.parent / ".env"
if not env_path.exists():
print(f"⚠️ .env文件不存在: {env_path}")
return
content = env_path.read_text(encoding='utf-8')
additions = f"""
# Few-shot 配置
FEWSHOT_ENABLED=true
FEWSHOT_TOP_K={top_k}
FEWSHOT_MIN_RATING={min_rating}
FEWSHOT_DATA_PATH=./data/experiences/all_samples.jsonl
"""
# 检查是否已存在
if "FEWSHOT_ENABLED" not in content:
with env_path.open('a', encoding='utf-8') as f:
f.write(additions)
print(f"✅ 已添加Few-shot配置到 .env")
else:
print(f"ℹ️ .env 中已存在FEWSHOT配置,跳过")
def main():
parser = argparse.ArgumentParser(description="集成Few-shot功能到Text2SQL系统")
parser.add_argument("--enable", action="store_true", help="启用few-shot")
parser.add_argument("--top-k", type=int, default=3, help="每次使用的示例数量")
parser.add_argument("--min-rating", type=int, default=7, help="示例最低评分")
parser.add_argument("--env-only", action="store_true", help="仅修改.env")
args = parser.parse_args()
if not args.enable and not args.env_only:
print("请使用 --enable 启用few-shot,或 --env-only 仅修改环境变量")
return
print("=" * 60)
print("Text2SQL Few-shot 集成工具")
print("=" * 60)
if args.env_only:
create_env_config(args.top_k, args.min_rating)
return
if args.enable:
print(f"\n配置: top_k={args.top_k}, min_rating={args.min_rating}")
# 1. 修改orchestrator.py
success = patch_orchestrator(
enable=True,
top_k=args.top_k,
min_rating=args.min_rating
)
if success:
# 2. 更新.env
create_env_config(args.top_k, args.min_rating)
print("\n" + "=" * 60)
print("✅ Few-shot集成完成!")
print("=" * 60)
print(f"\n下一步:")
print(f"1. 确保经验数据集已生成:")
print(f" python scripts/parse_examples.py")
print(f"\n2. 测试系统:")
print(f" python main.py \"查询2024年1月的销售额\" --verbose")
print(f"\n3. 查看日志中的few-shot选择:")
print(f" Few-shot已选择: Q31, Q6, Q44")
print(f"\n4. 调整参数:")
print(f" 编辑 .env 修改 FEWSHOT_TOP_K 和 FEWSHOT_MIN_RATING")
else:
print("❌ 集成失败")
if __name__ == "__main__":
main()