Files
2026-04-14 10:28:22 +08:00

238 lines
8.9 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/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"
):
"""
修改 backend/agents/orchestrator.py 添加few-shot支持
"""
orch_path = Path(__file__).parent.parent / "backend" / "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 backend/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()