diff --git a/pyproject.toml b/pyproject.toml index 139e019..105ba34 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -18,16 +18,16 @@ dependencies = [ "python-frontmatter>=1.0.0", "PyPDF2>=3.0.0", "python-docx>=1.0.0", - "numpy>=1.24.0", "tqdm>=4.66.0", - "rapidfuzz>=3.6.0", - "sentence-transformers>=2.5.0", "orjson>=3.9.0", - "click-spinner>=0.1.10", "sqlite-fts5", ] [project.optional-dependencies] +embeddings = [ + "numpy>=1.24.0", + "sentence-transformers>=2.5.0", +] openai = ["openai>=1.0.0"] ocr = [ "pytesseract>=0.3.10", diff --git a/raglite_sqlite/api.py b/raglite_sqlite/api.py index 02d6064..dc63239 100644 --- a/raglite_sqlite/api.py +++ b/raglite_sqlite/api.py @@ -7,6 +7,7 @@ from .chunking import chunk_blocks from .db import Database from .embeddings.base import EmbeddingBackend +from .embeddings.hash_backend import HashingBackend from .parsers.csv import CSVParser from .parsers.docx import DocxParser from .parsers.html import HTMLParser @@ -37,9 +38,7 @@ def _default_backend() -> EmbeddingBackend: - from .embeddings.sentence_transformers_backend import SentenceTransformersBackend - - return SentenceTransformersBackend() + return HashingBackend() class RagLite: diff --git a/raglite_sqlite/cli.py b/raglite_sqlite/cli.py index 7d10086..033825e 100644 --- a/raglite_sqlite/cli.py +++ b/raglite_sqlite/cli.py @@ -17,10 +17,24 @@ def get_rag(db: Path) -> RagLite: return RagLite(str(db)) -def get_backend(model: Optional[str]): - from .embeddings.sentence_transformers_backend import SentenceTransformersBackend - - return SentenceTransformersBackend(model_name=model or "sentence-transformers/all-MiniLM-L6-v2") +def get_backend(backend: str, model: Optional[str]): + backend_name = (backend or "hash").lower() + if backend_name == "hash": + from .embeddings.hash_backend import HashingBackend + + return HashingBackend() + if backend_name in {"sentence-transformers", "st"}: + try: + from .embeddings.sentence_transformers_backend import SentenceTransformersBackend + except ImportError as exc: # pragma: no cover - optional dependency missing + raise typer.BadParameter( + "Sentence Transformers backend requires raglite-sqlite[embeddings]" + ) from exc + + return SentenceTransformersBackend( + model_name=model or "sentence-transformers/all-MiniLM-L6-v2" + ) + raise typer.BadParameter(f"Unknown embedding backend '{backend}'") @app.command() @@ -35,6 +49,7 @@ def index( path: Path = typer.Argument(..., exists=True, file_okay=True, dir_okay=True), db: Path = typer.Option(..., help="Database path"), tags: Optional[str] = typer.Option(None, help="Comma-separated tags"), + backend: str = typer.Option("hash", help="Embedding backend to use (hash or sentence-transformers)"), model: Optional[str] = typer.Option(None, help="Embedding model name"), chunk_size: int = typer.Option(512, help="Chunk size in tokens"), overlap: int = typer.Option(64, help="Chunk overlap"), @@ -42,14 +57,14 @@ def index( recursive: bool = typer.Option(True, help="Recurse into directories"), skip_unchanged: bool = typer.Option(True, help="Skip unchanged files"), ) -> None: - backend = get_backend(model) + backend_impl = get_backend(backend, model) rag = get_rag(db) result = rag.index( [str(path)], tags=tags, chunk_size_tokens=chunk_size, chunk_overlap_tokens=overlap, - embedding_backend=backend, + embedding_backend=backend_impl, model_name=model, glob=glob, recurse=recursive, diff --git a/raglite_sqlite/db.py b/raglite_sqlite/db.py index fdcc687..ec9686a 100644 --- a/raglite_sqlite/db.py +++ b/raglite_sqlite/db.py @@ -3,7 +3,7 @@ import sqlite3 from array import array from pathlib import Path -from typing import Iterable, List, Optional, Sequence +from typing import Iterable, Iterator, List, Optional, Sequence from .utils import dumps_json, ensure_directory, loads_json @@ -186,6 +186,19 @@ def get_all_vectors(self, model_name: str | None = None) -> tuple[list[list[floa chunk_ids.append(row["chunk_id"]) return vectors, chunk_ids + def iter_vectors(self, model_name: str | None = None) -> Iterator[tuple[str, list[float]]]: + if model_name: + cur = self.conn.execute( + "SELECT chunk_id, dim, dtype, vec FROM vectors WHERE model_name = ?", + (model_name,), + ) + else: + cur = self.conn.execute("SELECT chunk_id, dim, dtype, vec FROM vectors") + for row in cur: + data = array("f") + data.frombytes(row["vec"]) + yield row["chunk_id"], list(data) + def get_embedding_cache(self, content_sha: str, model_name: str) -> Optional[list[float]]: cur = self.conn.execute( "SELECT dim, dtype, vec FROM cache_embeddings WHERE content_sha = ? AND model_name = ?", diff --git a/raglite_sqlite/embeddings/__init__.py b/raglite_sqlite/embeddings/__init__.py index 3cef675..90db5b3 100644 --- a/raglite_sqlite/embeddings/__init__.py +++ b/raglite_sqlite/embeddings/__init__.py @@ -1,5 +1,6 @@ """Embedding backends for RagLite.""" +from .hash_backend import HashingBackend from .sentence_transformers_backend import SentenceTransformersBackend -__all__ = ["SentenceTransformersBackend"] +__all__ = ["HashingBackend", "SentenceTransformersBackend"] diff --git a/raglite_sqlite/embeddings/hash_backend.py b/raglite_sqlite/embeddings/hash_backend.py new file mode 100644 index 0000000..8915786 --- /dev/null +++ b/raglite_sqlite/embeddings/hash_backend.py @@ -0,0 +1,38 @@ +from __future__ import annotations + +import hashlib +import math +from collections import Counter +from typing import Sequence + +from ..utils import normalize_text +from .base import EmbeddingBackend + + +class HashingBackend(EmbeddingBackend): + """Lightweight hashing-based embedding backend.""" + + def __init__(self, dim: int = 384) -> None: + self.dim = dim + self.model_name = f"hashing-{dim}" + + def embed_texts(self, texts: Sequence[str], model_name: str | None = None) -> Sequence[Sequence[float]]: + return [self._embed(text) for text in texts] + + def _embed(self, text: str) -> list[float]: + tokens = [token for token in normalize_text(text).split() if token] + if not tokens: + return [0.0] * self.dim + counts = Counter(tokens) + vector = [0.0] * self.dim + for token, weight in counts.items(): + index = self._bucket(token) + vector[index] += float(weight) + norm = math.sqrt(sum(value * value for value in vector)) + if norm == 0.0: + return vector + return [value / norm for value in vector] + + def _bucket(self, token: str) -> int: + digest = hashlib.blake2b(token.encode("utf-8"), digest_size=16).digest() + return int.from_bytes(digest, "big") % self.dim diff --git a/raglite_sqlite/search.py b/raglite_sqlite/search.py index add8eff..74281f2 100644 --- a/raglite_sqlite/search.py +++ b/raglite_sqlite/search.py @@ -1,9 +1,10 @@ from __future__ import annotations +import heapq import math -from typing import Dict, Iterable, List, Tuple +from typing import Dict, Iterable, List, Sequence, Tuple -from .db import Database, cosine_search +from .db import Database from .embeddings.base import EmbeddingBackend from .typing import SearchResult from .utils import normalize_text @@ -46,13 +47,33 @@ def vector_search( model_name: str | None, k: int, ) -> list[tuple[str, float]]: - matrix, chunk_ids = db.get_all_vectors(model_name=model_name) - if not matrix: - return [] query_vectors = backend.embed_texts([query], model_name=model_name) + if not query_vectors: + return [] query_vec = list(query_vectors[0]) - results = cosine_search(matrix, query_vec, top_k=min(len(chunk_ids), max(k * 5, 10))) - return [(chunk_ids[idx], score) for idx, score in results] + + def norm(vec: Sequence[float]) -> float: + return math.sqrt(sum(value * value for value in vec)) + + query_norm = norm(query_vec) + if math.isclose(query_norm, 0.0): + return [] + + limit = max(k * 5, 10) + heap: list[tuple[float, str]] = [] + for chunk_id, vector in db.iter_vectors(model_name=model_name): + row_norm = norm(vector) + if math.isclose(row_norm, 0.0): + score = 0.0 + else: + score = sum(a * b for a, b in zip(query_vec, vector)) / (row_norm * query_norm) + if len(heap) < limit: + heapq.heappush(heap, (score, chunk_id)) + continue + if score > heap[0][0]: + heapq.heapreplace(heap, (score, chunk_id)) + heap.sort(reverse=True) + return [(chunk_id, score) for score, chunk_id in heap] def hybrid_fuse( diff --git a/raglite_sqlite/server.py b/raglite_sqlite/server.py index 661aa01..4e5b80f 100644 --- a/raglite_sqlite/server.py +++ b/raglite_sqlite/server.py @@ -12,9 +12,9 @@ def _default_backend() -> EmbeddingBackend: - from .embeddings.sentence_transformers_backend import SentenceTransformersBackend + from .embeddings.hash_backend import HashingBackend - return SentenceTransformersBackend() + return HashingBackend() class QueryPayload(BaseModel): diff --git a/raglite_sqlite/utils.py b/raglite_sqlite/utils.py index 0dd02fd..c7aaee4 100644 --- a/raglite_sqlite/utils.py +++ b/raglite_sqlite/utils.py @@ -80,18 +80,20 @@ def now_ts() -> int: def iter_files(paths: Sequence[str], recurse: bool = True, glob: str | None = None) -> list[Path]: - candidates: list[Path] = [] + candidates: dict[Path, Path] = {} for input_path in paths: path = Path(input_path) if path.is_dir(): - pattern = glob or "**/*" if recurse else "*" - for child in path.glob(pattern): + if glob: + iterator = path.rglob(glob) if recurse else path.glob(glob) + else: + iterator = path.rglob("*") if recurse else path.glob("*") + for child in iterator: if child.is_file(): - candidates.append(child) + candidates.setdefault(child.resolve(), child) elif path.is_file(): - candidates.append(path) - unique = {p.resolve(): p for p in candidates} - return list(unique.values()) + candidates.setdefault(path.resolve(), path) + return list(candidates.values()) def dumps_json(data: object) -> str: diff --git a/tests/test_cli.py b/tests/test_cli.py index 104086a..792a80d 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -23,7 +23,7 @@ def test_cli_flow(tmp_path: Path, monkeypatch) -> None: # monkeypatch backend to avoid heavy model dummy = DummyBackend(dim=16) - monkeypatch.setattr("raglite_sqlite.cli.get_backend", lambda model: dummy) + monkeypatch.setattr("raglite_sqlite.cli.get_backend", lambda backend, model: dummy) data_dir = Path(__file__).parent / "data" result = runner.invoke(