#!/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()