feat: chunked bge embeddings + lancedb sutta index
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
10
scripts/build_index.py
Normal file
10
scripts/build_index.py
Normal file
@@ -0,0 +1,10 @@
|
||||
#!/usr/bin/env python
|
||||
"""Build the sutta embeddings index."""
|
||||
from pathlib import Path
|
||||
from buddhagpt.index import build_index
|
||||
|
||||
if __name__ == "__main__":
|
||||
jsonl = Path("corpus/suttas.jsonl")
|
||||
db_path = Path("data/lancedb")
|
||||
n = build_index(jsonl, db_path)
|
||||
print(f"Built index with {n} chunks at {db_path}")
|
||||
40
src/buddhagpt/index.py
Normal file
40
src/buddhagpt/index.py
Normal file
@@ -0,0 +1,40 @@
|
||||
import json
|
||||
from pathlib import Path
|
||||
import lancedb
|
||||
from sentence_transformers import SentenceTransformer
|
||||
|
||||
EMBED_MODEL = "BAAI/bge-small-en-v1.5"
|
||||
|
||||
|
||||
def chunk_text(text: str, max_words: int = 300, overlap: int = 50) -> list[str]:
|
||||
words = text.split()
|
||||
chunks, start = [], 0
|
||||
while start < len(words):
|
||||
chunks.append(" ".join(words[start:start + max_words]))
|
||||
if start + max_words >= len(words):
|
||||
break
|
||||
start += max_words - overlap
|
||||
return chunks
|
||||
|
||||
|
||||
def build_index(jsonl: Path, db_path: Path) -> int:
|
||||
model = SentenceTransformer(EMBED_MODEL, device="mps")
|
||||
rows = []
|
||||
for line in jsonl.read_text().splitlines():
|
||||
s = json.loads(line)
|
||||
for i, chunk in enumerate(chunk_text(s["text"])):
|
||||
rows.append({"uid": s["uid"], "title": s["title"], "chunk": chunk, "chunk_i": i})
|
||||
vecs = model.encode([r["chunk"] for r in rows], batch_size=64, show_progress_bar=True)
|
||||
for r, v in zip(rows, vecs):
|
||||
r["vector"] = v.tolist()
|
||||
db = lancedb.connect(db_path)
|
||||
db.create_table("suttas", rows, mode="overwrite")
|
||||
return len(rows)
|
||||
|
||||
|
||||
def search(db_path: Path, query: str, k: int = 4) -> list[dict]:
|
||||
model = SentenceTransformer(EMBED_MODEL, device="mps")
|
||||
tbl = lancedb.connect(db_path).open_table("suttas")
|
||||
q = model.encode([query])[0].tolist()
|
||||
hits = tbl.search(q).limit(k).to_list()
|
||||
return [{"uid": h["uid"], "title": h["title"], "chunk": h["chunk"], "score": h["_distance"]} for h in hits]
|
||||
14
tests/test_index.py
Normal file
14
tests/test_index.py
Normal file
@@ -0,0 +1,14 @@
|
||||
from buddhagpt.index import chunk_text
|
||||
|
||||
|
||||
def test_chunk_text_respects_max_words():
|
||||
text = " ".join(f"w{i}" for i in range(700))
|
||||
chunks = chunk_text(text, max_words=300, overlap=50)
|
||||
assert all(len(c.split()) <= 300 for c in chunks)
|
||||
assert len(chunks) == 3
|
||||
|
||||
|
||||
def test_chunks_overlap():
|
||||
text = " ".join(f"w{i}" for i in range(400))
|
||||
a, b = chunk_text(text, max_words=300, overlap=50)
|
||||
assert a.split()[-50:] == b.split()[:50]
|
||||
Reference in New Issue
Block a user