Spaces:
Sleeping
Sleeping
File size: 8,174 Bytes
b2931f4 | 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 203 204 205 206 207 208 209 210 211 212 | """BM25 lexical retrieval over the same chunks indexed in Qdrant.
Why this exists: dense embeddings (Cohere v3) lose exact-match signal on
years, tickers, GAAP terminology, and dollar amounts β exactly the tokens
that matter most in financial documents. BM25 catches these.
This module is *not* a full retriever. It returns (chunk_id, score) pairs.
The fusion + hydration into full RetrievedChunk objects happens in
retrieval.hybrid (Decision 10).
Pipeline:
build: data/processed/*.jsonl βββΆ tokenize βββΆ BM25Okapi
β
βΌ
pickled to data/bm25_index.pkl
load: pickle load β ready to search
"""
from __future__ import annotations
import pickle
import re
from dataclasses import dataclass
from pathlib import Path
import numpy as np
from rank_bm25 import BM25Okapi
from finrag.ingestion.parse import PROCESSED_DIR, Chunk
# parse.py β ingestion/ β finrag/ β src/ β backend/ β ROOT
REPO_ROOT = Path(__file__).resolve().parents[4]
INDEX_PATH = REPO_ROOT / "data" / "bm25_index.pkl"
# Token regex: any run of alphanumeric chars including underscores. Drops
# punctuation, splits on whitespace + symbols. Lowercased before splitting.
# Critical: this exact function is also called on queries β the same vocab
# must be used on both sides or no terms will match.
_TOKEN_RE = re.compile(r"\w+")
def tokenize(text: str) -> list[str]:
"""Lowercase + simple word tokenization.
The same function runs on chunk text at index time and on user queries
at search time. Don't tweak one side without the other.
"""
return _TOKEN_RE.findall(text.lower())
# ββ On-disk format ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
@dataclass
class _BM25Bundle:
"""What we pickle. Separated so we can version the schema later.
`chunks` is a list of (id, ticker, fiscal_year, chunk_type) tuples β the
minimal payload we need for post-search filtering and lookups. The full
chunk text/metadata lives in Qdrant; storing it twice would double disk
use and risk drift between stores.
"""
bm25: BM25Okapi
# Parallel arrays β index `i` in `bm25` corresponds to chunks[i].
# We use a tuple-list rather than a dict because BM25Okapi indexes by
# position, not by chunk_id.
chunk_ids: list[str]
tickers: list[str]
fiscal_years: list[int]
chunk_types: list[str]
# Pin __module__ so pickle records the dotted path "finrag.retrieval.lexical"
# instead of "__main__" when this file is run via `python -m`. Without this,
# the pickle is only loadable from the same entrypoint that built it.
_BM25Bundle.__module__ = "finrag.retrieval.lexical"
# ββ Build βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
def _load_all_chunks(processed_dir: Path) -> list[Chunk]:
chunks: list[Chunk] = []
for jsonl in sorted(processed_dir.glob("*.jsonl")):
for line in jsonl.read_text(encoding="utf-8").splitlines():
if line.strip():
chunks.append(Chunk.model_validate_json(line))
return chunks
def build_index(processed_dir: Path = PROCESSED_DIR) -> _BM25Bundle:
"""Build a BM25 index from all chunks in processed_dir and persist it."""
chunks = _load_all_chunks(processed_dir)
if not chunks:
raise RuntimeError(f"No chunks found in {processed_dir}")
print(f"Tokenizing {len(chunks)} chunksβ¦")
tokenized = [tokenize(c.text) for c in chunks]
print("Building BM25Okapi indexβ¦")
# k1=1.5, b=0.75 are BM25's standard defaults. The rank_bm25 library
# exposes these as kwargs; leave them at defaults unless we have a
# specific reason β these are well-calibrated for English text and
# any tuning we'd do should be eval-driven, not guess-driven.
bm25 = BM25Okapi(tokenized)
bundle = _BM25Bundle(
bm25=bm25,
chunk_ids=[c.chunk_id for c in chunks],
tickers=[c.ticker for c in chunks],
fiscal_years=[c.fiscal_year for c in chunks],
chunk_types=[c.chunk_type for c in chunks],
)
INDEX_PATH.parent.mkdir(parents=True, exist_ok=True)
with INDEX_PATH.open("wb") as f:
pickle.dump(bundle, f, protocol=pickle.HIGHEST_PROTOCOL)
print(f"Wrote {INDEX_PATH} ({INDEX_PATH.stat().st_size / 1024:.0f} KB)")
return bundle
# ββ Load + search ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
_cached_bundle: _BM25Bundle | None = None
def load_index() -> _BM25Bundle:
"""Load the pickled index from disk, caching in-process.
Module-level cache (not lru_cache) because the underlying BM25Okapi
object is heavyweight (~30 MB at our scale) β we want exactly one in
memory regardless of how many callers ask for it.
"""
global _cached_bundle
if _cached_bundle is None:
if not INDEX_PATH.exists():
raise FileNotFoundError(
f"BM25 index not found at {INDEX_PATH}. "
"Run `uv run python -m finrag.retrieval.lexical` to build it."
)
with INDEX_PATH.open("rb") as f:
_cached_bundle = pickle.load(f)
return _cached_bundle
def search(
query: str,
top_k: int = 50,
ticker: str | None = None,
fiscal_year: int | None = None,
chunk_type: str | None = None,
) -> list[tuple[str, float]]:
"""Return ranked (chunk_id, score) tuples for a query.
Filtering is post-ranking: we ask BM25 for top-N (where N > top_k to
leave headroom after filtering), then drop chunks that don't match.
This is fine at our scale (~4k chunks); at 1M+ you'd want a filter-
aware index structure or pre-shard by ticker.
"""
bundle = load_index()
tokens = tokenize(query)
if not tokens:
return []
# get_scores returns one score per indexed document, in index order
scores = bundle.bm25.get_scores(tokens)
# Build candidate list β over-fetch to allow for filter attrition.
# 4x is a heuristic; if filters are tight (e.g. one ticker Γ one year),
# we may want more β but unbounded over-fetch defeats the purpose.
candidate_count = top_k * 4 if (ticker or fiscal_year or chunk_type) else top_k
candidate_count = min(candidate_count, len(scores))
# argpartition is O(n) vs argsort's O(n log n) β meaningful at scale.
# We get the top-K unordered, then sort just those K.
top_indices = np.argpartition(-scores, candidate_count - 1)[:candidate_count]
# Sort the candidates by descending score
top_indices = top_indices[np.argsort(-scores[top_indices])]
results: list[tuple[str, float]] = []
for i in top_indices:
if ticker and bundle.tickers[i] != ticker:
continue
if fiscal_year and bundle.fiscal_years[i] != fiscal_year:
continue
if chunk_type and bundle.chunk_types[i] != chunk_type:
continue
results.append((bundle.chunk_ids[i], float(scores[i])))
if len(results) >= top_k:
break
return results
# ββ CLI βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
def main() -> None:
build_index()
# Sanity check: run a couple of test queries
print("\nSanity-check queries:")
for q in [
"services revenue 2023",
"iPhone net sales",
"SG&A expense",
"Tesla R&D",
]:
results = search(q, top_k=3)
print(f"\n Q: {q!r}")
for chunk_id, score in results:
print(f" {chunk_id} score={score:.3f}")
if __name__ == "__main__":
main()
|