Spaces:
Sleeping
Sleeping
File size: 7,183 Bytes
005e9fd | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 | """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))
|