Files
ai-g3sb-backman2.0/scripts/parse_examples.py
T
2026-04-10 16:52:07 +08:00

440 lines
16 KiB
Python
Raw 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
"""
将 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()