File size: 6,417 Bytes
56fa10c
 
 
 
 
 
6e92226
56fa10c
 
6e92226
56fa10c
 
 
 
 
 
 
6e92226
 
 
3c9dd57
 
6e92226
56fa10c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6e92226
 
 
 
 
 
 
 
56fa10c
6e92226
56fa10c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6e92226
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
56fa10c
 
 
 
 
 
 
 
 
 
 
 
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
"""ChromaDB collection builder with section-aware chunking."""
from __future__ import annotations

import json
from pathlib import Path

import torch
import chromadb
from chromadb.utils.embedding_functions import SentenceTransformerEmbeddingFunction
from rich.progress import BarColumn, MofNCompleteColumn, Progress, TextColumn, TimeElapsedColumn

from config import CHROMA_COLLECTION, CHROMA_DIR, PAPERS_PATH
from logging_config import get_logger
from models import ALSPaper

_logger = get_logger("rag.indexer")

# Use MPS on Apple Silicon, CUDA on NVIDIA, otherwise CPU.
_DEVICE = "mps" if torch.backends.mps.is_available() else ("cuda" if torch.cuda.is_available() else "cpu")

# BioLORD-2023-C: anchored to UMLS/SNOMED CT/MeSH ontologies — natively understands
# biomedical synonyms (TARDBP = TDP-43, SOD1 = superoxide dismutase) and clinical phrasing.
_EMBED_FN = SentenceTransformerEmbeddingFunction(model_name="FremyCompany/BioLORD-2023-C", device=_DEVICE)


def _chunk_paper(paper: ALSPaper) -> list[dict]:
    """
    Split a paper into indexable chunks.
    - Full text available: one chunk per section (split on [Section Title] markers)
    - Abstract only: single chunk = title + mesh terms + abstract
    """
    base_meta = {
        "pmid": paper.pmid,
        "title": paper.title,
        "year": paper.year,
        "doi": paper.doi,
        "citation_count": paper.citation_count,
        # ChromaDB metadata must be scalar — serialize lists as comma-separated strings
        "entity_names": ",".join(paper.entity_names),
        "mesh_terms": ",".join(paper.mesh_terms[:10]),  # cap to avoid huge metadata
        "has_full_text": int(bool(paper.full_text)),  # bool not supported → int
    }

    if paper.full_text:
        sections: list[tuple[str, str]] = []
        current_title = "Abstract"
        current_lines: list[str] = [paper.abstract]

        for line in paper.full_text.split("\n"):
            stripped = line.strip()
            if stripped.startswith("[") and stripped.endswith("]") and len(stripped) < 80:
                if current_lines:
                    body = "\n".join(current_lines).strip()
                    if body:
                        sections.append((current_title, body))
                current_title = stripped[1:-1]
                current_lines = []
            else:
                current_lines.append(line)

        if current_lines:
            body = "\n".join(current_lines).strip()
            if body:
                sections.append((current_title, body))

        # Prioritise high-value sections; cap at 6 total to keep index lean.
        # With 500 papers the cross-encoder only sees 20 candidates anyway —
        # 50+ chunks per paper adds noise without improving recall.
        _PRIORITY = {"abstract", "introduction", "results", "discussion", "conclusion", "methods"}
        priority = [s for s in sections if s[0].lower() in _PRIORITY]
        others = [s for s in sections if s[0].lower() not in _PRIORITY]
        selected = (priority + others)[:6]

        chunks = []
        for i, (section_title, section_text) in enumerate(selected):
            doc = f"{paper.title}\n[{section_title}]\n{section_text}"
            chunks.append({
                "id": f"{paper.pmid}_s{i}",
                "document": doc,
                "metadata": {**base_meta, "section": section_title, "chunk_index": i},
            })
        return chunks if chunks else [_abstract_chunk(paper, base_meta)]

    return [_abstract_chunk(paper, base_meta)]


def _abstract_chunk(paper: ALSPaper, base_meta: dict) -> dict:
    doc = f"{paper.title}\n{' '.join(paper.mesh_terms)}\n{paper.abstract}"
    return {
        "id": paper.pmid,
        "document": doc,
        "metadata": {**base_meta, "section": "abstract", "chunk_index": 0},
    }


def build_collection(
    papers_path: Path = PAPERS_PATH,
    chroma_dir: Path = CHROMA_DIR,
    collection_name: str = CHROMA_COLLECTION,
    reset: bool = False,
) -> chromadb.Collection:
    """Build ChromaDB collection from papers.jsonl. Idempotent — skips already-indexed chunks."""
    chroma_dir.mkdir(parents=True, exist_ok=True)
    client = chromadb.PersistentClient(path=str(chroma_dir))

    if reset:
        try:
            client.delete_collection(collection_name)
            _logger.info(f"Deleted collection: {collection_name}")
        except Exception:
            pass

    collection = client.get_or_create_collection(
        name=collection_name,
        embedding_function=_EMBED_FN,
        metadata={"hnsw:space": "cosine"},
    )

    papers: list[ALSPaper] = []
    with open(papers_path, encoding="utf-8") as f:
        for line in f:
            line = line.strip()
            if line:
                papers.append(ALSPaper.from_dict(json.loads(line)))

    _logger.info(f"Loaded {len(papers)} papers")

    all_chunks = []
    for paper in papers:
        all_chunks.extend(_chunk_paper(paper))

    # Skip already-indexed chunks (safe to re-run)
    existing_ids = set(collection.get(include=[])["ids"])
    new_chunks = [c for c in all_chunks if c["id"] not in existing_ids]

    if not new_chunks:
        _logger.info("All chunks already indexed")
        return collection

    _logger.info(f"Indexing {len(new_chunks)} new chunks from {len(papers)} papers")

    batch_size = 100
    batches = [new_chunks[i : i + batch_size] for i in range(0, len(new_chunks), batch_size)]

    with Progress(
        TextColumn("[cyan]Embedding chunks[/cyan]"),
        BarColumn(),
        MofNCompleteColumn(),
        TimeElapsedColumn(),
    ) as progress:
        task = progress.add_task("", total=len(new_chunks))
        for batch in batches:
            collection.add(
                ids=[c["id"] for c in batch],
                documents=[c["document"] for c in batch],
                metadatas=[c["metadata"] for c in batch],
            )
            progress.advance(task, len(batch))

    _logger.info(f"Collection '{collection_name}': {collection.count()} total chunks")
    return collection


def load_collection(
    chroma_dir: Path = CHROMA_DIR,
    collection_name: str = CHROMA_COLLECTION,
) -> chromadb.Collection:
    """Load an existing collection at query time (fast, no re-embedding)."""
    client = chromadb.PersistentClient(path=str(chroma_dir))
    return client.get_collection(name=collection_name, embedding_function=_EMBED_FN)