Spaces:
Running
Running
| """Build a ChromaDB corpus of NASA passages for the playground. | |
| Downloads `govt.jsonl.zip` from IBM mt-rag-benchmark, filters to NASA domains, | |
| merges consecutive same-title passages to double chunk size (~halving embedding | |
| count), embeds with `ibm-granite/granite-embedding-small-english-r2`. | |
| Persists to `/data/nasa_chroma` when a storage bucket is mounted (survives | |
| restarts). Falls back to in-memory when `/data` is unavailable. | |
| """ | |
| import io | |
| import json | |
| import os | |
| import time | |
| import zipfile | |
| import chromadb | |
| import httpx | |
| import torch | |
| from chromadb import Documents, EmbeddingFunction, Embeddings | |
| from tqdm.auto import tqdm | |
| from transformers import AutoModel, AutoTokenizer | |
| EMBEDDING_MODEL_ID = "ibm-granite/granite-embedding-small-english-r2" | |
| GOVT_JSONL_URL = "https://github.com/IBM/mt-rag-benchmark/raw/main/corpora/passage_level/govt.jsonl.zip" | |
| PERSIST_PATH = "/data/nasa_chroma" | |
| NASA_DOMAINS = { | |
| "www.nasa.gov", | |
| "www.jpl.nasa.gov", | |
| "earthobservatory.nasa.gov", | |
| "europa.nasa.gov", | |
| } | |
| class GraniteEmbeddingFunction(EmbeddingFunction): | |
| """ChromaDB EmbeddingFunction backed by ibm-granite/granite-embedding-*-r2.""" | |
| def __init__(self, model_id=EMBEDDING_MODEL_ID, batch_size=64, device="cpu"): | |
| self._device = device | |
| self._batch = batch_size | |
| self._tokenizer = AutoTokenizer.from_pretrained(model_id) | |
| self._model = AutoModel.from_pretrained(model_id).to(device).eval() | |
| print(f"Granite embedding model ready on {self._device} ({model_id})") | |
| def __call__(self, input: Documents) -> Embeddings: | |
| all_embs = [] | |
| for i in range(0, len(input), self._batch): | |
| batch = list(input[i : i + self._batch]) | |
| enc = self._tokenizer( | |
| batch, return_tensors="pt", truncation=True, max_length=512, padding=True | |
| ) | |
| enc = {k: v.to(self._device) for k, v in enc.items()} | |
| with torch.no_grad(): | |
| out = self._model(**enc) | |
| mask = enc["attention_mask"].unsqueeze(-1).float() | |
| emb = (out.last_hidden_state * mask).sum(1) / mask.sum(1).clamp(min=1e-9) | |
| all_embs.extend(emb.cpu().float().tolist()) | |
| return all_embs | |
| def _extract_domain(url): | |
| parts = url.split("/") | |
| return parts[2] if len(parts) > 2 else "" | |
| def _filter_and_merge(raw_docs): | |
| """Filter to NASA domains and merge consecutive same-title passages.""" | |
| nasa_docs = [] | |
| for doc in raw_docs: | |
| url = doc.get("url", "") | |
| if _extract_domain(url) in NASA_DOMAINS: | |
| text = doc.get("text", "").strip() | |
| if text: | |
| nasa_docs.append(doc) | |
| nasa_docs.sort(key=lambda d: (d.get("title", ""), d.get("_id", d.get("id", "")))) | |
| merged = [] | |
| i = 0 | |
| while i < len(nasa_docs): | |
| title = nasa_docs[i].get("title", "") | |
| text = nasa_docs[i].get("text", "").strip() | |
| url = nasa_docs[i].get("url", "") | |
| doc_id = nasa_docs[i].get("_id", nasa_docs[i].get("id", str(i))) | |
| if i + 1 < len(nasa_docs) and nasa_docs[i + 1].get("title", "") == title: | |
| text = text + "\n\n" + nasa_docs[i + 1].get("text", "").strip() | |
| i += 2 | |
| else: | |
| i += 1 | |
| merged.append({"id": doc_id, "text": text, "title": title, "url": url}) | |
| return merged | |
| def _persist_path(): | |
| """Return the best available persist path, or None for in-memory.""" | |
| if os.path.isdir("/data"): | |
| return PERSIST_PATH | |
| app_dir = os.path.dirname(os.path.abspath(__file__)) | |
| return os.path.join(app_dir, "nasa_chroma") | |
| def build_nasa_chroma(embedding_model_id=EMBEDDING_MODEL_ID): | |
| """Return a ready-to-query Chroma collection for NASA passages. | |
| Persists to disk (/data on HF Spaces, ./nasa_chroma locally) and loads | |
| instantly on subsequent runs. First run embeds ~1,965 passages. | |
| """ | |
| granite_ef = GraniteEmbeddingFunction(model_id=embedding_model_id) | |
| persist_path = _persist_path() | |
| client = chromadb.PersistentClient(path=persist_path) | |
| collection = client.get_or_create_collection( | |
| name="nasa", | |
| embedding_function=granite_ef, | |
| metadata={"hnsw:space": "cosine"}, | |
| ) | |
| if collection.count() > 0: | |
| print(f"Loaded persisted DB from {persist_path} ({collection.count():,} docs).") | |
| return collection | |
| print(f"Building DB — will persist to {persist_path}.") | |
| print(f"Downloading {GOVT_JSONL_URL} …") | |
| t0 = time.time() | |
| with httpx.Client(follow_redirects=True, timeout=120.0) as c: | |
| resp = c.get(GOVT_JSONL_URL) | |
| resp.raise_for_status() | |
| with zipfile.ZipFile(io.BytesIO(resp.content)) as zf: | |
| inner = next(n for n in zf.namelist() if n.endswith(".jsonl")) | |
| data = zf.open(inner).read().decode() | |
| print(f"Downloaded in {time.time() - t0:.1f}s.") | |
| raw_docs = [json.loads(line) for line in data.strip().split("\n") if line.strip()] | |
| merged = _filter_and_merge(raw_docs) | |
| print(f"Filtered {len(raw_docs):,} → {len(merged):,} NASA passages (merged pairs).") | |
| ids = [d["id"] for d in merged] | |
| texts = [d["text"] for d in merged] | |
| metas = [{"title": d["title"], "url": d["url"]} for d in merged] | |
| t1 = time.time() | |
| batch = 500 | |
| for i in tqdm(range(0, len(ids), batch), unit="batch", desc="indexing"): | |
| collection.upsert( | |
| ids=ids[i : i + batch], | |
| documents=texts[i : i + batch], | |
| metadatas=metas[i : i + batch], | |
| ) | |
| print(f"Done. {collection.count():,} docs indexed in {time.time() - t1:.1f}s.") | |
| return collection | |