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