Crocos / app /rag.py
Beibars Ybraiakhyn
Upload Gradio Space
75f4c14
Raw
History Blame Contribute Delete
3.41 kB
"""
RAG pipeline: load knowledge base -> chunk -> embed (bge-m3) -> Chroma -> retrieve.
Embeddings are computed explicitly via the LLM backend (Ollama/vLLM) and passed
to Chroma, so the retrieval path is fully under our control and uses the same
multilingual model for documents and queries.
"""
import os
import glob
import chromadb
from .config import settings
from .embeddings import embed_texts
def get_collection(reset: bool = False):
db = chromadb.PersistentClient(path=settings.chroma_dir)
if reset:
try:
db.delete_collection(settings.collection)
except Exception:
pass
# cosine distance pairs well with the multilingual embedders we use
return db.get_or_create_collection(settings.collection, metadata={"hnsw:space": "cosine"})
def chunk_text(text: str, size: int, overlap: int) -> list[str]:
"""Pack paragraphs into ~`size`-char chunks, then add `overlap` between them."""
paras = [p.strip() for p in text.split("\n\n") if p.strip()]
chunks, cur = [], ""
for p in paras:
if len(cur) + len(p) + 2 <= size:
cur = (cur + "\n\n" + p).strip()
else:
if cur:
chunks.append(cur)
if len(p) <= size:
cur = p
else: # hard-split an oversized paragraph
for i in range(0, len(p), size - overlap):
chunks.append(p[i:i + size])
cur = ""
if cur:
chunks.append(cur)
if overlap > 0 and len(chunks) > 1:
stitched = [chunks[0]]
for i in range(1, len(chunks)):
tail = chunks[i - 1][-overlap:]
stitched.append((tail + "\n" + chunks[i]).strip())
chunks = stitched
return chunks
def ingest(kb_dir: str = "knowledge_base", reset: bool = False) -> int:
"""(Re)build the vector store from every *.md file in the knowledge base.
Pass reset=True to drop the collection first (needed when the embedding
model — and therefore the vector dimension — changes).
"""
col = get_collection(reset=reset)
ids, docs, metas = [], [], []
for path in sorted(glob.glob(os.path.join(kb_dir, "*.md"))):
lang = "kk" if path.endswith("_kk.md") else "ru"
with open(path, encoding="utf-8") as f:
text = f.read()
for i, chunk in enumerate(chunk_text(text, settings.chunk_size, settings.chunk_overlap)):
ids.append(f"{os.path.basename(path)}::{i}")
docs.append(chunk)
metas.append({"source": os.path.basename(path), "lang": lang})
if not docs:
return 0
embeddings = embed_texts(docs)
col.upsert(ids=ids, documents=docs, embeddings=embeddings, metadatas=metas)
return len(docs)
def retrieve(query: str, top_k: int | None = None, lang: str | None = None) -> list[dict]:
col = get_collection()
q_emb = embed_texts([query])[0]
where = {"lang": lang} if lang else None
res = col.query(
query_embeddings=[q_emb],
n_results=top_k or settings.top_k,
where=where,
)
docs = res["documents"][0]
metas = res["metadatas"][0]
dists = res["distances"][0]
# cosine distance -> similarity score in [0, 1]
return [
{"text": d, "source": m["source"], "score": max(0.0, 1.0 - dist)}
for d, m, dist in zip(docs, metas, dists)
]