Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
sqlitevec.py242 linesDownload Raw Back to vectorstores
1from __future__ import annotations2 3import json4import logging5import struct6import warnings7from typing import (8    TYPE_CHECKING,9    Any,10    Iterable,11    List,12    Optional,13    Tuple,14    Type,15)16 17from langchain_core.documents import Document18from langchain_core.embeddings import Embeddings19from langchain_core.vectorstores import VectorStore20 21if TYPE_CHECKING:22    import sqlite323 24logger = logging.getLogger(__name__)25 26 27def serialize_f32(vector: List[float]) -> bytes:28    """Serializes a list of floats into a compact "raw bytes" format29 30    Source: https://github.com/asg017/sqlite-vec/blob/21c5a14fc71c83f135f5b00c84115139fd12c492/examples/simple-python/demo.py#L8-L1031    """32    return struct.pack("%sf" % len(vector), *vector)33 34 35class SQLiteVec(VectorStore):36    """SQLite with Vec extension as a vector database.37 38    To use, you should have the ``sqlite-vec`` python package installed.39    Example:40        .. code-block:: python41            from langchain_community.vectorstores import SQLiteVec42            from langchain_community.embeddings.openai import OpenAIEmbeddings43            ...44    """45 46    def __init__(47        self,48        table: str,49        connection: Optional[sqlite3.Connection],50        embedding: Embeddings,51        db_file: str = "vec.db",52    ):53        """Initialize with sqlite client with vss extension."""54        try:55            import sqlite_vec  # noqa  # pylint: disable=unused-import56        except ImportError:57            raise ImportError(58                "Could not import sqlite-vec python package. "59                "Please install it with `pip install sqlite-vec`."60            )61 62        if not connection:63            connection = self.create_connection(db_file)64 65        if not isinstance(embedding, Embeddings):66            warnings.warn("embeddings input must be Embeddings object.")67 68        self._connection = connection69        self._table = table70        self._embedding = embedding71 72        self.create_table_if_not_exists()73 74    def create_table_if_not_exists(self) -> None:75        self._connection.execute(76            f"""77            CREATE TABLE IF NOT EXISTS {self._table}78            (79                rowid INTEGER PRIMARY KEY AUTOINCREMENT,80                text TEXT,81                metadata BLOB,82                text_embedding BLOB83            )84            ;85            """86        )87        self._connection.execute(88            f"""89            CREATE VIRTUAL TABLE IF NOT EXISTS {self._table}_vec USING vec0(90                rowid INTEGER PRIMARY KEY,91                text_embedding float[{self.get_dimensionality()}]92            )93            ;94            """95        )96        self._connection.execute(97            f"""98                CREATE TRIGGER IF NOT EXISTS {self._table}_embed_text 99                AFTER INSERT ON {self._table}100                BEGIN101                    INSERT INTO {self._table}_vec(rowid, text_embedding)102                    VALUES (new.rowid, new.text_embedding) 103                    ;104                END;105            """106        )107        self._connection.commit()108 109    def add_texts(110        self,111        texts: Iterable[str],112        metadatas: Optional[List[dict]] = None,113        **kwargs: Any,114    ) -> List[str]:115        """Add more texts to the vectorstore index.116        Args:117            texts: Iterable of strings to add to the vectorstore.118            metadatas: Optional list of metadatas associated with the texts.119            kwargs: vectorstore specific parameters120        """121        max_id = self._connection.execute(122            f"SELECT max(rowid) as rowid FROM {self._table}"123        ).fetchone()["rowid"]124        if max_id is None:  # no text added yet125            max_id = 0126 127        embeds = self._embedding.embed_documents(list(texts))128        if not metadatas:129            metadatas = [{} for _ in texts]130        data_input = [131            (text, json.dumps(metadata), serialize_f32(embed))132            for text, metadata, embed in zip(texts, metadatas, embeds)133        ]134        self._connection.executemany(135            f"INSERT INTO {self._table}(text, metadata, text_embedding) VALUES (?,?,?)",136            data_input,137        )138        self._connection.commit()139        # pulling every ids we just inserted140        results = self._connection.execute(141            f"SELECT rowid FROM {self._table} WHERE rowid > {max_id}"142        )143        return [row["rowid"] for row in results]144 145    def similarity_search_with_score_by_vector(146        self, embedding: List[float], k: int = 4, **kwargs: Any147    ) -> List[Tuple[Document, float]]:148        sql_query = f"""149            SELECT 150                text,151                metadata,152                distance153            FROM {self._table} AS e154            INNER JOIN {self._table}_vec AS v on v.rowid = e.rowid  155            WHERE156                v.text_embedding MATCH ?157                AND k = ?158            ORDER BY distance159        """160        cursor = self._connection.cursor()161        cursor.execute(162            sql_query,163            [serialize_f32(embedding), k],164        )165        results = cursor.fetchall()166 167        documents = []168        for row in results:169            metadata = json.loads(row["metadata"]) or {}170            doc = Document(page_content=row["text"], metadata=metadata)171            documents.append((doc, row["distance"]))172 173        return documents174 175    def similarity_search(176        self, query: str, k: int = 4, **kwargs: Any177    ) -> List[Document]:178        """Return docs most similar to query."""179        embedding = self._embedding.embed_query(query)180        documents = self.similarity_search_with_score_by_vector(181            embedding=embedding, k=k182        )183        return [doc for doc, _ in documents]184 185    def similarity_search_with_score(186        self, query: str, k: int = 4, **kwargs: Any187    ) -> List[Tuple[Document, float]]:188        """Return docs most similar to query."""189        embedding = self._embedding.embed_query(query)190        documents = self.similarity_search_with_score_by_vector(191            embedding=embedding, k=k192        )193        return documents194 195    def similarity_search_by_vector(196        self, embedding: List[float], k: int = 4, **kwargs: Any197    ) -> List[Document]:198        documents = self.similarity_search_with_score_by_vector(199            embedding=embedding, k=k200        )201        return [doc for doc, _ in documents]202 203    @classmethod204    def from_texts(205        cls: Type[SQLiteVec],206        texts: List[str],207        embedding: Embeddings,208        metadatas: Optional[List[dict]] = None,209        table: str = "langchain",210        db_file: str = "vec.db",211        **kwargs: Any,212    ) -> SQLiteVec:213        """Return VectorStore initialized from texts and embeddings."""214        connection = cls.create_connection(db_file)215        vec = cls(216            table=table, connection=connection, db_file=db_file, embedding=embedding217        )218        vec.add_texts(texts=texts, metadatas=metadatas)219        return vec220 221    @staticmethod222    def create_connection(db_file: str) -> sqlite3.Connection:223        import sqlite3224 225        import sqlite_vec226 227        connection = sqlite3.connect(db_file)228        connection.row_factory = sqlite3.Row229        connection.enable_load_extension(True)230        sqlite_vec.load(connection)231        connection.enable_load_extension(False)232        return connection233 234    def get_dimensionality(self) -> int:235        """236        Function that does a dummy embedding to figure out how many dimensions237        this embedding function returns. Needed for the virtual table DDL.238        """239        dummy_text = "This is a dummy text"240        dummy_embedding = self._embedding.embed_query(dummy_text)241        return len(dummy_embedding)242 
codekingpro/portable-devtools · Team Ai