Spaces:
Running on Zero
Running on Zero
Download engine/retrieval.py from AngeloUNIMI/document_exam_trainer: direct link, hf CLI and curl.
- Browser
- Download file 5.99 kB
-
https://huggingface.co/spaces/AngeloUNIMI/document_exam_trainer/resolve/main/engine/retrieval.py
- Command line
-
hf download hf://spaces/AngeloUNIMI/document_exam_trainer/engine/retrieval.py
-
curl -L -o retrieval.py https://huggingface.co/spaces/AngeloUNIMI/document_exam_trainer/resolve/main/engine/retrieval.py
5.99 kB
| from __future__ import annotations | |
| import hashlib | |
| import re | |
| import secrets | |
| from typing import Callable | |
| import numpy as np | |
| from .schemas import Chunk, ParsedCorpus, Topic | |
| from .workers import CPUWorker | |
| class Embedder: | |
| def __init__(self, worker: CPUWorker, demo: bool = False): | |
| self.worker, self.demo = worker, demo | |
| def encode(self, texts: list[str], progress: Callable[[str], None] | None = None) -> np.ndarray: | |
| if not texts: | |
| return np.zeros((0, 0), dtype="float32") | |
| if self.demo: | |
| # Explicit demo mode ONLY. Never used as a silent semantic fallback. | |
| out = np.zeros((len(texts), 256), dtype="float32") | |
| for i, text in enumerate(texts): | |
| for token in re.findall(r"\w+", text.casefold()): | |
| out[i, int.from_bytes(hashlib.sha256(token.encode()).digest()[:4], "big") % 256] += 1 | |
| return out / np.maximum(np.linalg.norm(out, axis=1, keepdims=True), 1e-9) | |
| vectors = [] | |
| for start in range(0, len(texts), 16): | |
| if progress: | |
| progress(f"Embedding passages {start+1}-{min(start+16, len(texts))} of {len(texts)} on CPU...") | |
| result = self.worker.request("embed", {"texts": texts[start:start+16]}, timeout=480, progress=progress) | |
| vectors.extend(result["vectors"]) | |
| out = np.asarray(vectors, dtype="float32") | |
| if out.ndim != 2 or len(out) != len(texts) or not np.isfinite(out).all(): | |
| raise RuntimeError("Embedding model returned invalid vectors.") | |
| return out / np.maximum(np.linalg.norm(out, axis=1, keepdims=True), 1e-9) | |
| class CorpusIndex: | |
| """Per-session vectors; FAISS when installed, mathematically equivalent NumPy otherwise.""" | |
| def __init__(self, parsed: ParsedCorpus, vectors: np.ndarray): | |
| if len(vectors) != len(parsed.chunks): | |
| raise ValueError("Chunk/vector count mismatch.") | |
| self.parsed = parsed | |
| self.chunks = parsed.chunks | |
| self.by_id = {c.id: c for c in self.chunks} | |
| self.topics = {t.id: t for t in parsed.topics} | |
| self.vectors = np.ascontiguousarray(vectors, dtype="float32") | |
| self.faiss = None | |
| try: | |
| import faiss | |
| self.faiss = faiss.IndexFlatIP(self.vectors.shape[1]) | |
| self.faiss.add(self.vectors) | |
| except ImportError: | |
| # The vectors are still real semantic embeddings; only the exact | |
| # similarity-search implementation differs for small corpora. | |
| pass | |
| def options(self) -> list[tuple[str, str]]: | |
| return [("Any primary topic (random selection)", "all")] + [(t.label, t.id) for t in self.parsed.topics] | |
| def hits(self, vector: np.ndarray, *, role: str, ids: set[str] | None = None, | |
| query: str = "", top_k: int = 6) -> list[Chunk]: | |
| vector = np.asarray(vector, dtype="float32").reshape(-1) | |
| if vector.shape[0] != self.vectors.shape[1]: | |
| raise ValueError("Embedding dimension mismatch. Rebuild the corpus.") | |
| # Filter scope BEFORE ranking; a globally similar chapter cannot starve | |
| # a selected narrow topic of results. | |
| candidates = [i for i, c in enumerate(self.chunks) if c.role == role and (ids is None or c.id in ids)] | |
| if not candidates: | |
| return [] | |
| if self.faiss is not None: | |
| scores, rows = self.faiss.search(vector.reshape(1, -1), len(self.chunks)) | |
| score_map = dict(zip(map(int, rows[0]), map(float, scores[0]))) | |
| dense = np.asarray([score_map[i] for i in candidates]) | |
| else: | |
| dense = self.vectors[candidates] @ vector | |
| dense_order = np.argsort(-dense) | |
| tokens = set(re.findall(r"\w+", query.casefold())) | |
| lexical = [len(tokens & set(re.findall(r"\w+", self.chunks[i].text.casefold()))) for i in candidates] | |
| lex_order = np.argsort(-np.asarray(lexical)) | |
| fused = np.zeros(len(candidates)) | |
| for rank, pos in enumerate(dense_order): | |
| fused[pos] += 1 / (60 + rank) | |
| if any(lexical): | |
| for rank, pos in enumerate(lex_order): | |
| fused[pos] += 0.35 / (60 + rank) | |
| return [self.chunks[candidates[int(p)]] for p in np.argsort(-fused)[:top_k]] | |
| def question_context(self, topic_id: str, focus: str, embedder: Embedder) -> tuple[Topic, list[Chunk]]: | |
| if topic_id == "all": | |
| # Select a coherent topic, rather than the first N pages of the corpus. | |
| subtopics = [t for t in self.parsed.topics if "-T" in t.id] | |
| topic = secrets.choice(subtopics or self.parsed.topics) | |
| else: | |
| if topic_id not in self.topics: | |
| raise ValueError("Choose a topic from the current corpus.") | |
| topic = self.topics[topic_id] | |
| candidates = [self.by_id[sid] for sid in topic.source_ids] | |
| if focus.strip(): | |
| vector = embedder.encode([focus])[0] | |
| selected = self.hits(vector, role="primary", ids=set(topic.source_ids), query=focus, top_k=7) | |
| # Neighbour passages remain within the selected topic. | |
| nearby = {c.ordinal + d for c in selected for d in (-1, 0, 1)} | |
| candidates = [c for c in candidates if c.ordinal in nearby][:12] | |
| elif len(candidates) > 10: | |
| start = secrets.randbelow(len(candidates) - 7) | |
| candidates = candidates[start:start+8] | |
| return topic, candidates | |
| def pack_context(chunks: list[Chunk], max_chars: int) -> tuple[str, list[Chunk]]: | |
| blocks, admitted, length = [], [], 0 | |
| seen: set[str] = set() | |
| for c in chunks: | |
| if c.id in seen: | |
| continue | |
| block = f"[SOURCE {c.id}; {c.role}; {c.filename}; {c.location}; section: {c.heading}]\n{c.text}" | |
| if length + len(block) > max_chars: | |
| continue | |
| blocks.append(block) | |
| admitted.append(c) | |
| seen.add(c.id) | |
| length += len(block) + 2 | |
| return "\n\n".join(blocks), admitted | |