238 lines
8.9 KiB
Python
238 lines
8.9 KiB
Python
#!/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()
|