mathidot111/First_agent_template
0
1from __future__ import annotations2 3import argparse4import json5import random6import re7from pathlib import Path8from typing import Any9 10from tools.query_knowledge import RAW_DIR, iter_source_files, load_source_file11 12 13KEY_TERMS = [14 "volatility smile",15 "implied volatility",16 "local volatility",17 "stochastic volatility",18 "Black-Scholes",19 "delta",20 "gamma",21 "vega",22 "theta",23 "rho",24 "skew",25 "straddle",26 "correlation",27 "at-the-money",28 "forward",29 "risk-neutral",30]31 32 33PROJECT_ROOT = Path(__file__).resolve().parents[1]34OUTPUT_PATH = PROJECT_ROOT / "eval" / "local_options_eval.jsonl"35 36 37def normalize_space(text: str) -> str:38 return re.sub(r"\s+", " ", text).strip()39 40 41def extract_keywords(text: str, max_keywords: int = 4) -> list[str]:42 lowered = text.lower()43 keywords = [term for term in KEY_TERMS if term.lower() in lowered]44 equation_ids = re.findall(r"\(\d+\.\d+[a-z]?\)", text)45 formulas = re.findall(r"[A-Za-z๐๐๐๐๐ด][A-Za-z0-9๐๐๐๐๐ด_{}^]*\s*=", text)46 keywords.extend(equation_ids[:2])47 keywords.extend(item.strip() for item in formulas[:2])48 49 if not keywords:50 candidates = [51 word52 for word in re.findall(r"[A-Za-z][A-Za-z-]{4,}", text)53 if word.lower() not in {"there", "where", "which", "would", "could", "should", "chapter"}54 ]55 keywords.extend(candidates[:max_keywords])56 57 deduped = []58 banned = {"id=", "FORMULA", "value ="}59 for keyword in keywords:60 if keyword and keyword not in banned and keyword not in deduped:61 deduped.append(keyword)62 return deduped[:max_keywords]63 64 65def is_sane_section(section: str | None) -> bool:66 if not section:67 return False68 section = section.strip()69 if not 6 <= len(section) <= 90:70 return False71 if section.count(",") >= 2:72 return False73 digit_count = sum(char.isdigit() for char in section)74 letter_count = sum(char.isalpha() for char in section)75 if digit_count > max(2, letter_count // 3):76 return False77 if re.search(r"\b(figure|table|printed|united states|amount unit price|call price|under)$", section, re.I):78 return False79 if "figure" in section.lower() or "table" in section.lower():80 return False81 if re.search(r"\b(figure|table|printed|united states|amount unit price|call price)\b", section, re.I):82 return False83 words = section.split()84 if len(words) > 12:85 return False86 return True87 88 89def make_case(document: Any, index: int) -> dict[str, Any] | None:90 metadata = document.metadata91 text = normalize_space(document.text)92 if len(text) < 80:93 return None94 95 page = metadata.get("page_number")96 if isinstance(page, int) and (page < 25 or page > 500):97 return None98 section = metadata.get("section_path") or metadata.get("section_title")99 content_type = metadata.get("content_type", "text")100 formula_id = metadata.get("formula_id")101 keywords = extract_keywords(text)102 if not keywords and not section:103 return None104 105 if content_type == "formula" or formula_id:106 question = f"What formula or equation is described on page {page}?"107 answer_type = "formula"108 elif is_sane_section(section):109 question = f"What does the section {section} discuss?"110 answer_type = "section"111 keywords.append(section.split(">")[-1].strip())112 else:113 if not keywords:114 return None115 term = keywords[0]116 if term.lower() in {"formula", "id=", "value ="}:117 return None118 question = f"Where does the options reference discuss {term}?"119 answer_type = "concept"120 121 expected_pages = [page] if page is not None else []122 return {123 "id": f"auto_options_{index:03d}",124 "question": question,125 "expected_pages": expected_pages,126 "expected_keywords": keywords[:5],127 "answer_type": answer_type,128 }129 130 131def generate_cases(count: int, seed: int) -> list[dict[str, Any]]:132 documents = []133 for source_file in iter_source_files(RAW_DIR):134 documents.extend(load_source_file(source_file))135 136 random.Random(seed).shuffle(documents)137 cases = []138 seen_questions = set()139 for document in documents:140 case = make_case(document, len(cases) + 1)141 if not case:142 continue143 if case["question"] in seen_questions:144 continue145 seen_questions.add(case["question"])146 cases.append(case)147 if len(cases) >= count:148 break149 150 if len(cases) < count:151 raise RuntimeError(f"Only generated {len(cases)} cases; requested {count}.")152 return cases153 154 155def main() -> None:156 parser = argparse.ArgumentParser(description="Generate local options RAG eval cases.")157 parser.add_argument("--count", type=int, default=40)158 parser.add_argument("--seed", type=int, default=20260525)159 parser.add_argument("--output", type=Path, default=OUTPUT_PATH)160 args = parser.parse_args()161 162 cases = generate_cases(args.count, args.seed)163 args.output.parent.mkdir(parents=True, exist_ok=True)164 args.output.write_text(165 "\n".join(json.dumps(case, ensure_ascii=False) for case in cases) + "\n",166 encoding="utf-8",167 )168 print(f"Wrote {len(cases)} cases to {args.output}")169 170 171if __name__ == "__main__":172 main()173 