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