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