First_agent_template / eval /generate_local_options_eval.py
mathidot's picture
build option trading agent modules
8f1601b
Raw
History Blame Contribute Delete
5.4 kB
from __future__ import annotations
import argparse
import json
import random
import re
from pathlib import Path
from typing import Any
from tools.query_knowledge import RAW_DIR, iter_source_files, load_source_file
KEY_TERMS = [
"volatility smile",
"implied volatility",
"local volatility",
"stochastic volatility",
"Black-Scholes",
"delta",
"gamma",
"vega",
"theta",
"rho",
"skew",
"straddle",
"correlation",
"at-the-money",
"forward",
"risk-neutral",
]
PROJECT_ROOT = Path(__file__).resolve().parents[1]
OUTPUT_PATH = PROJECT_ROOT / "eval" / "local_options_eval.jsonl"
def normalize_space(text: str) -> str:
return re.sub(r"\s+", " ", text).strip()
def extract_keywords(text: str, max_keywords: int = 4) -> list[str]:
lowered = text.lower()
keywords = [term for term in KEY_TERMS if term.lower() in lowered]
equation_ids = re.findall(r"\(\d+\.\d+[a-z]?\)", text)
formulas = re.findall(r"[A-Za-z𝜎𝜇𝜌𝜃𝛴][A-Za-z0-9𝜎𝜇𝜌𝜃𝛴_{}^]*\s*=", text)
keywords.extend(equation_ids[:2])
keywords.extend(item.strip() for item in formulas[:2])
if not keywords:
candidates = [
word
for word in re.findall(r"[A-Za-z][A-Za-z-]{4,}", text)
if word.lower() not in {"there", "where", "which", "would", "could", "should", "chapter"}
]
keywords.extend(candidates[:max_keywords])
deduped = []
banned = {"id=", "FORMULA", "value ="}
for keyword in keywords:
if keyword and keyword not in banned and keyword not in deduped:
deduped.append(keyword)
return deduped[:max_keywords]
def is_sane_section(section: str | None) -> bool:
if not section:
return False
section = section.strip()
if not 6 <= len(section) <= 90:
return False
if section.count(",") >= 2:
return False
digit_count = sum(char.isdigit() for char in section)
letter_count = sum(char.isalpha() for char in section)
if digit_count > max(2, letter_count // 3):
return False
if re.search(r"\b(figure|table|printed|united states|amount unit price|call price|under)$", section, re.I):
return False
if "figure" in section.lower() or "table" in section.lower():
return False
if re.search(r"\b(figure|table|printed|united states|amount unit price|call price)\b", section, re.I):
return False
words = section.split()
if len(words) > 12:
return False
return True
def make_case(document: Any, index: int) -> dict[str, Any] | None:
metadata = document.metadata
text = normalize_space(document.text)
if len(text) < 80:
return None
page = metadata.get("page_number")
if isinstance(page, int) and (page < 25 or page > 500):
return None
section = metadata.get("section_path") or metadata.get("section_title")
content_type = metadata.get("content_type", "text")
formula_id = metadata.get("formula_id")
keywords = extract_keywords(text)
if not keywords and not section:
return None
if content_type == "formula" or formula_id:
question = f"What formula or equation is described on page {page}?"
answer_type = "formula"
elif is_sane_section(section):
question = f"What does the section {section} discuss?"
answer_type = "section"
keywords.append(section.split(">")[-1].strip())
else:
if not keywords:
return None
term = keywords[0]
if term.lower() in {"formula", "id=", "value ="}:
return None
question = f"Where does the options reference discuss {term}?"
answer_type = "concept"
expected_pages = [page] if page is not None else []
return {
"id": f"auto_options_{index:03d}",
"question": question,
"expected_pages": expected_pages,
"expected_keywords": keywords[:5],
"answer_type": answer_type,
}
def generate_cases(count: int, seed: int) -> list[dict[str, Any]]:
documents = []
for source_file in iter_source_files(RAW_DIR):
documents.extend(load_source_file(source_file))
random.Random(seed).shuffle(documents)
cases = []
seen_questions = set()
for document in documents:
case = make_case(document, len(cases) + 1)
if not case:
continue
if case["question"] in seen_questions:
continue
seen_questions.add(case["question"])
cases.append(case)
if len(cases) >= count:
break
if len(cases) < count:
raise RuntimeError(f"Only generated {len(cases)} cases; requested {count}.")
return cases
def main() -> None:
parser = argparse.ArgumentParser(description="Generate local options RAG eval cases.")
parser.add_argument("--count", type=int, default=40)
parser.add_argument("--seed", type=int, default=20260525)
parser.add_argument("--output", type=Path, default=OUTPUT_PATH)
args = parser.parse_args()
cases = generate_cases(args.count, args.seed)
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(
"\n".join(json.dumps(case, ensure_ascii=False) for case in cases) + "\n",
encoding="utf-8",
)
print(f"Wrote {len(cases)} cases to {args.output}")
if __name__ == "__main__":
main()