tiny-search-reader-60M / searchreader.py
ahkamboh's picture
tiny-search-reader-60M v3.1a: model, GGUF phone files, prompt helper, eval results
e563482 verified
Raw History Blame Contribute Delete
8.82 kB
"""Ask tiny-search-reader-60M a question about a list of search results.
build_prompt(tokenizer, question, snippets) -> prompt token ids that fit the 512-token context
answer_transformers(question, snippets) -> answer string, using the Hugging Face files in this folder
answer_gguf(question, snippets) -> answer string, using a GGUF file through llama-cpp-python
A snippet is a dict {"title": "...", "text": "..."} (a plain string is read as text with no title). Results go in
the order the search engine returned them. The answer is a few words, or "not found" when the results do not
answer the question.
Needs: `tokenizers` for build_prompt; `torch` + `transformers` for answer_transformers; `tokenizers` +
`llama-cpp-python` for answer_gguf. Run `python searchreader.py` for a demo.
"""
import atexit
import os
import re
HERE = os.path.dirname(os.path.abspath(__file__))
TOKENIZER_JSON = os.path.join(HERE, "tokenizer.json")
DEFAULT_GGUF = os.path.join(HERE, "gguf", "tiny-search-reader-60M-Q8_0.gguf")
CTX = 512 # the model's context length
MAX_NEW = 24 # answer tokens; the exams were graded with 24
NOT_FOUND = "not found"
QUESTION, SNIPPET, ANSWER, EOT, PAD = "<|question|>", "<|snippet|>", "<|answer|>", "<|endoftext|>", "<|pad|>"
EOT_ID, PAD_ID = 0, 4
# ---- prompt building (a port of searchfmt.clean / render_one / Prompt.build as used at test time) ----
def clean(t):
return re.sub(r"\s+", " ", (t or "").replace("\u00a0", " ")).strip()
def render_one(i, s):
if isinstance(s, str):
s = {"text": s}
title, text = clean(s.get("title")), clean(s.get("text"))
return f"[{i}] {title}: {text}" if title else f"[{i}] {text}"
class _Tok:
"""The same three calls on a `tokenizers.Tokenizer` or on a transformers tokenizer (AutoTokenizer)."""
def __init__(self, tok):
self.tok = tok
self.raw = hasattr(tok, "token_to_id") # tokenizers.Tokenizer
def encode(self, text):
if self.raw:
return self.tok.encode(text).ids
return self.tok.encode(text, add_special_tokens=False, verbose=False) # no "longer than 512" warning
def token_to_id(self, t):
return self.tok.token_to_id(t) if self.raw else self.tok.convert_tokens_to_ids(t)
def decode(self, ids):
if self.raw:
return self.tok.decode(ids)
return self.tok.decode(ids, skip_special_tokens=True)
def _build(tok, question, snippets, ctx, reserve):
q, s, a = (tok.token_to_id(t) for t in (QUESTION, SNIPPET, ANSWER))
head = [q] + tok.encode(" " + clean(question)) + [s]
room = ctx - len(head) - 1 - reserve
parts = [tok.encode(" " + render_one(i + 1, sn)) for i, sn in enumerate(snippets)]
keep = list(range(len(parts)))
# Too long: drop results from the end, one at a time, but always keep the first.
while sum(len(parts[i]) for i in keep) > room and len(keep) > 1:
keep.pop()
ids = [t for i in keep for t in parts[i]]
# Still too long (one huge first result): cut it at the token limit.
if len(ids) > room:
ids = ids[:room]
return head + ids + [a]
def build_prompt(tokenizer, question, snippets, ctx=CTX, reserve=MAX_NEW):
"""Token ids of "<|question|> {question}<|snippet|> [1] {title}: {text} [2] ...<|answer|>".
Results keep their order. If they do not fit in ctx - reserve tokens, results are dropped from the end
until they do. The first one always stays, and if it alone is still too long it is cut short.
A question too long to leave any room is cut to its first ctx // 2 tokens. This is exactly what the model
was graded on (finetune_search.Exam.prompt); at inference no gold answer is known, so the training-time
rule in searchfmt.Prompt.build that protects snippets holding the answer never applies.
tokenizer: a `tokenizers.Tokenizer` (Tokenizer.from_file("tokenizer.json")) or a transformers tokenizer."""
tok = _Tok(tokenizer)
ids = _build(tok, question, snippets, ctx, reserve)
if len(ids) > ctx - reserve:
q = tok.encode(" " + clean(question))
ids = _build(tok, tok.decode(q[:ctx // 2]), snippets, ctx, reserve)
return ids
# ---- answering ----
_cache = {}
@atexit.register
def close():
"""Free cached llama.cpp models. Runs at exit too: letting Python tear them down late can abort on macOS."""
for key in [k for k in _cache if k[0] == "gguf"]:
_cache.pop(key)[0].close()
def _hf(model_dir):
if ("hf", model_dir) not in _cache:
from transformers import AutoModelForCausalLM, AutoTokenizer
tok = AutoTokenizer.from_pretrained(model_dir)
model = AutoModelForCausalLM.from_pretrained(model_dir).eval() # fp32, as stored
_cache["hf", model_dir] = model, tok
return _cache["hf", model_dir]
def answer_transformers(question, snippets, model_dir=HERE, max_new_tokens=MAX_NEW):
"""Greedy answer with transformers (fp32 on CPU unless you move the model)."""
import torch
model, tok = _hf(model_dir)
ids = build_prompt(tok, question, snippets, reserve=max_new_tokens)
x = torch.tensor([ids], device=model.device)
with torch.no_grad():
out = model.generate(x, attention_mask=torch.ones_like(x), max_new_tokens=max_new_tokens, do_sample=False,
eos_token_id=EOT_ID, pad_token_id=PAD_ID)
new = out[0, len(ids):].tolist()
if EOT_ID in new:
new = new[:new.index(EOT_ID)]
return tok.decode(new, skip_special_tokens=True).strip()
def _gguf(gguf_path, tokenizer_path, n_threads):
key = ("gguf", gguf_path, tokenizer_path, n_threads)
if key not in _cache:
from llama_cpp import Llama
from tokenizers import Tokenizer
llm = Llama(model_path=gguf_path, n_ctx=CTX, n_batch=CTX, n_threads=n_threads, n_gpu_layers=0, verbose=False)
_cache[key] = llm, Tokenizer.from_file(tokenizer_path)
return _cache[key]
def answer_gguf(question, snippets, gguf_path=DEFAULT_GGUF, tokenizer_path=TOKENIZER_JSON, n_threads=4,
max_new_tokens=MAX_NEW):
"""Greedy answer with llama-cpp-python on CPU.
The prompt goes in as token ids from our own tokenizer, so the special tokens are single ids and no BOS is
added. Decoding stops at <|endoftext|> (id 0) or after max_new_tokens."""
llm, tok = _gguf(gguf_path, tokenizer_path, n_threads)
ids = build_prompt(tok, question, snippets, reserve=max_new_tokens)
llm.reset() # evaluate every prompt from scratch
out = []
for t in llm.generate(ids, temp=0.0, top_k=1, top_p=1.0, min_p=0.0, repeat_penalty=1.0, reset=True):
if t == EOT_ID:
break
out.append(t)
if len(out) == max_new_tokens:
break
return tok.decode(out).strip()
# ---- demo (made-up results) ----
DEMO = [
("When does the Harbor Street library open on Saturday?", [
{"title": "Harbor Street Library - Hours and location",
"text": "The Harbor Street Library is open Monday to Friday from 8 am to 8 pm. On Saturday the library "
"opens at 10 am and closes at 4 pm. It is closed on Sunday."},
{"title": "Harbor Street Library events",
"text": "Story time for children runs every Wednesday morning in the reading room."},
{"title": "City libraries | Parking",
"text": "Free parking is available behind the Harbor Street Library for up to two hours."},
]),
("Who won the 2031 Harbor Street Chess Open?", [
{"title": "Harbor Street Chess Open 2025 results",
"text": "Mara Lindqvist won the 2025 Harbor Street Chess Open with 7 points from 9 games."},
{"title": "Harbor Street Chess Open 2026",
"text": "The 2026 Harbor Street Chess Open was won by Tomas Varga after a tie-break against Mara Lindqvist."},
{"title": "About the Harbor Street Chess Open",
"text": "The open has been played every June at the Harbor Street Library since 2019."},
]),
]
if __name__ == "__main__":
import argparse
ap = argparse.ArgumentParser(description="Demo: answer two made-up questions from made-up search results.")
ap.add_argument("--backend", choices=["transformers", "gguf", "both"], default="both")
ap.add_argument("--gguf", default=DEFAULT_GGUF, help="GGUF file for the llama-cpp-python backend")
a = ap.parse_args()
for question, snippets in DEMO:
print("Q:", question)
if a.backend in ("transformers", "both"):
print(" transformers:", answer_transformers(question, snippets))
if a.backend in ("gguf", "both"):
print(f" gguf ({os.path.basename(a.gguf)}):", answer_gguf(question, snippets, gguf_path=a.gguf))