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