Spaces:
Sleeping
Sleeping
| """Generate the retrieval evaluation dataset. | |
| Samples chunks and asks an LLM for questions each one answers. The chunk a | |
| question came from is that question's expected answer. | |
| uv run python -m eval.build_dataset | |
| uv run python -m eval.build_dataset --sample 50 | |
| """ | |
| import argparse | |
| import json | |
| import os | |
| import random | |
| import re | |
| import sys | |
| from concurrent.futures import ThreadPoolExecutor | |
| from dotenv import load_dotenv | |
| from llama_index.core.base.llms.types import ChatMessage, MessageRole | |
| from rag.config import CHUNKS_PATH, EVAL_DATASET_PATH, PROVIDERS | |
| from rag.providers import make_llm | |
| from rag.types import Chunk, EvalRow | |
| SAMPLE_SIZE = 200 | |
| SEED = 20260804 | |
| CONCURRENCY = 16 | |
| MIN_CHARS = 400 | |
| CANDIDATES = 4 | |
| WORD = re.compile(r"[a-z_][a-z0-9_]+") | |
| # Real Stack Overflow titles used as few-shot prompt examples | |
| # | |
| # stackoverflow.com/questions/29483365 multiline string literal | |
| # stackoverflow.com/questions/24990520 integer to string | |
| # stackoverflow.com/questions/39204908 release / debug builds with cfg | |
| # stackoverflow.com/questions/57756927 main.rs and lib.rs | |
| # stackoverflow.com/questions/26469715 asserting a panic in a test | |
| # stackoverflow.com/questions/31192956 reading and writing files | |
| EXAMPLES = ( | |
| "What is the syntax for a multiline string literal?", | |
| "How do I convert from an integer to a string?", | |
| "How to check release / debug builds using cfg in Rust?", | |
| "Rust modules confusion when there is main.rs and lib.rs", | |
| "How do I write a Rust unit test that ensures that a panic has occurred?", | |
| "What's the de-facto way of reading and writing files in Rust 1.x?", | |
| ) | |
| PROMPT = """Below is an excerpt from the official Rust documentation. | |
| --------------------- | |
| {context} | |
| --------------------- | |
| Write {count} different questions that this excerpt answers, as a Rust | |
| programmer would type them into a search box before finding this page. | |
| Match the register of these real questions: | |
| {examples} | |
| Rules: | |
| - Under twelve words each. | |
| - One thing per question. No compound questions. | |
| - Reach for the words a programmer would use before reading this page, not the | |
| excerpt's own phrasing. | |
| - Do not mention "the excerpt", "the text", or "the documentation". | |
| Reply with one question per line and nothing else.""" | |
| JUDGE = """Below is an excerpt from the official Rust documentation, followed by | |
| questions someone might search for. | |
| --------------------- | |
| {context} | |
| --------------------- | |
| {questions} | |
| For each question, reply KEEP or DROP. | |
| KEEP if this excerpt is where a reader searching that question should land: it | |
| answers them directly, and more specifically than the rest of the documentation | |
| would. | |
| DROP if the excerpt only touches the subject in passing, or if the question is | |
| broad enough that several other pages would answer it just as well. | |
| Reply with one verdict per line, in order, as "1. KEEP" or "1. DROP", and | |
| nothing else.""" | |
| def lexical_overlap(question: str, text: str) -> float: | |
| asked = set(WORD.findall(question.lower())) | |
| if not asked: | |
| return 1.0 | |
| return len(asked.intersection(WORD.findall(text.lower()))) / len(asked) | |
| def sample_chunks(sample_size: int) -> list[Chunk]: | |
| with CHUNKS_PATH.open(encoding="utf-8") as handle: | |
| chunks = [json.loads(line) for line in handle] | |
| usable = [chunk for chunk in chunks if len(chunk["text"]) >= MIN_CHARS] | |
| print(f"{len(chunks):,} chunks, {len(usable):,} long enough to sample from") | |
| random.seed(SEED) | |
| return random.sample(usable, min(sample_size, len(usable))) | |
| def main(argv: list[str]) -> int: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--sample", type=int, default=SAMPLE_SIZE) | |
| parser.add_argument("--provider", default="openai", choices=sorted(PROVIDERS)) | |
| args = parser.parse_args(argv[1:]) | |
| load_dotenv(".env") | |
| spec = PROVIDERS[args.provider] | |
| api_key = os.environ.get(spec.env_var, "").strip() | |
| if not api_key: | |
| print(f"Set {spec.env_var} to generate the dataset.", file=sys.stderr) | |
| return 2 | |
| if not CHUNKS_PATH.exists(): | |
| print(f"No {CHUNKS_PATH}. Run: uv run python -m ingest.parse_books", file=sys.stderr) | |
| return 2 | |
| llm = make_llm(args.provider, api_key, spec.default_model) | |
| chunks = sample_chunks(args.sample) | |
| by_id = {chunk["id"]: chunk for chunk in chunks} | |
| prompt = PROMPT.replace("{examples}", "\n".join(f" {q}" for q in EXAMPLES)) | |
| def say(content: str, chunk_id: str) -> str | None: | |
| try: | |
| return str(llm.chat([ChatMessage(role=MessageRole.USER, content=content)]).message.content) | |
| except Exception as error: | |
| print(f" {chunk_id}: {error}", file=sys.stderr) | |
| return None | |
| def numbered(reply: str) -> list[str]: | |
| """Strip whatever list markers the model chose to use.""" | |
| return [ | |
| re.sub(r"^\s*(?:[-*\d.)]+\s*)+", "", line).strip() | |
| for line in reply.splitlines() | |
| if line.strip() | |
| ] | |
| def ask(chunk: Chunk) -> str | None: | |
| reply = say(prompt.format(context=chunk["text"], count=CANDIDATES), chunk["id"]) | |
| if reply is None: | |
| return None | |
| candidates = [q for q in numbered(reply) if q.endswith("?") or len(q.split()) >= 4] | |
| if not candidates: | |
| return None | |
| listing = "\n".join(f"{n}. {q}" for n, q in enumerate(candidates, 1)) | |
| verdicts = say(JUDGE.format(context=chunk["text"], questions=listing), chunk["id"]) | |
| if verdicts is None: | |
| return None | |
| kept = [ | |
| question | |
| for question, verdict in zip(candidates, numbered(verdicts)) | |
| if verdict.upper().startswith("KEEP") | |
| ] | |
| if not kept: | |
| return None | |
| return min(kept, key=lambda q: lexical_overlap(q, chunk["text"])) | |
| rows: list[EvalRow] = [] | |
| with ThreadPoolExecutor(max_workers=CONCURRENCY) as pool: | |
| for position, (chunk, question) in enumerate(zip(chunks, pool.map(ask, chunks)), 1): | |
| if question: | |
| rows.append( | |
| { | |
| "question": question, | |
| "chunk_id": chunk["id"], | |
| "book": chunk["metadata"]["book"], | |
| } | |
| ) | |
| if position % 25 == 0 or position == len(chunks): | |
| print(f" {position}/{len(chunks)}", flush=True) | |
| failed = len(chunks) - len(rows) | |
| if failed: | |
| print(f"{failed} chunk(s) produced no question and were dropped", file=sys.stderr) | |
| overlaps = sorted( | |
| lexical_overlap(row["question"], by_id[row["chunk_id"]]["text"]) for row in rows | |
| ) | |
| words = sorted(len(row["question"].split()) for row in rows) | |
| print( | |
| f"\nmedian overlap with the source passage {overlaps[len(overlaps) // 2]:.2f}" | |
| f", median length {words[len(words) // 2]} words" | |
| ) | |
| with EVAL_DATASET_PATH.open("w", encoding="utf-8") as handle: | |
| json.dump(rows, handle, indent=2, ensure_ascii=False) | |
| print(f"\nwrote {len(rows)} question/chunk pairs to {EVAL_DATASET_PATH}") | |
| return 0 | |
| if __name__ == "__main__": | |
| sys.exit(main(sys.argv)) | |