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))