Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
oracleai.py195 linesDownload Raw Back to embeddings
1# Authors:2#   Harichandan Roy (hroy)3#   David Jiang (ddjiang)4#5# -----------------------------------------------------------------------------6# oracleai.py7# -----------------------------------------------------------------------------8 9from __future__ import annotations10 11import json12import logging13import traceback14from typing import TYPE_CHECKING, Any, Dict, List, Optional15 16from langchain_core.embeddings import Embeddings17from pydantic import BaseModel, ConfigDict18 19if TYPE_CHECKING:20    from oracledb import Connection21 22logger = logging.getLogger(__name__)23 24"""OracleEmbeddings class"""25 26 27class OracleEmbeddings(BaseModel, Embeddings):28    """Get Embeddings"""29 30    """Oracle Connection"""31    conn: Any = None32    """Embedding Parameters"""33    params: Dict[str, Any]34    """Proxy"""35    proxy: Optional[str] = None36 37    def __init__(self, **kwargs: Any):38        super().__init__(**kwargs)39 40    model_config = ConfigDict(41        extra="forbid",42    )43 44    """45    1 - user needs to have create procedure, 46        create mining model, create any directory privilege.47    2 - grant create procedure, create mining model, 48        create any directory to <user>;49    """50 51    @staticmethod52    def load_onnx_model(53        conn: Connection, dir: str, onnx_file: str, model_name: str54    ) -> None:55        """Load an ONNX model to Oracle Database.56        Args:57            conn: Oracle Connection,58            dir: Oracle Directory,59            onnx_file: ONNX file name,60            model_name: Name of the model.61        """62 63        try:64            if conn is None or dir is None or onnx_file is None or model_name is None:65                raise Exception("Invalid input")66 67            cursor = conn.cursor()68            cursor.execute(69                """70                begin71                    dbms_data_mining.drop_model(model_name => :model, force => true);72                    SYS.DBMS_VECTOR.load_onnx_model(:path, :filename, :model, 73                        json('{"function" : "embedding", 74                            "embeddingOutput" : "embedding", 75                            "input": {"input": ["DATA"]}}'));76                end;""",77                path=dir,78                filename=onnx_file,79                model=model_name,80            )81 82            cursor.close()83 84        except Exception as ex:85            logger.info(f"An exception occurred :: {ex}")86            traceback.print_exc()87            cursor.close()88            raise89 90    def embed_documents(self, texts: List[str]) -> List[List[float]]:91        """Compute doc embeddings using an OracleEmbeddings.92        Args:93            texts: The list of texts to embed.94        Returns:95            List of embeddings, one for each input text.96        """97 98        try:99            import oracledb100        except ImportError as e:101            raise ImportError(102                "Unable to import oracledb, please install with "103                "`pip install -U oracledb`."104            ) from e105 106        if texts is None:107            return None108 109        embeddings: List[List[float]] = []110        try:111            # returns strings or bytes instead of a locator112            oracledb.defaults.fetch_lobs = False113            cursor = self.conn.cursor()114 115            if self.proxy:116                cursor.execute(117                    "begin utl_http.set_proxy(:proxy); end;", proxy=self.proxy118                )119 120            chunks = []121            for i, text in enumerate(texts, start=1):122                chunk = {"chunk_id": i, "chunk_data": text}123                chunks.append(json.dumps(chunk))124 125            vector_array_type = self.conn.gettype("SYS.VECTOR_ARRAY_T")126            inputs = vector_array_type.newobject(chunks)127            cursor.execute(128                "select t.* "129                + "from dbms_vector_chain.utl_to_embeddings(:content, "130                + "json(:params)) t",131                content=inputs,132                params=json.dumps(self.params),133            )134 135            for row in cursor:136                if row is None:137                    embeddings.append([])138                else:139                    rdata = json.loads(row[0])140                    # dereference string as array141                    vec = json.loads(rdata["embed_vector"])142                    embeddings.append(vec)143 144            cursor.close()145            return embeddings146        except Exception as ex:147            logger.info(f"An exception occurred :: {ex}")148            traceback.print_exc()149            cursor.close()150            raise151 152    def embed_query(self, text: str) -> List[float]:153        """Compute query embedding using an OracleEmbeddings.154        Args:155            text: The text to embed.156        Returns:157            Embedding for the text.158        """159        return self.embed_documents([text])[0]160 161 162# uncomment the following code block to run the test163 164"""165# A sample unit test.166 167import oracledb168# get the Oracle connection 169conn = oracledb.connect(170    user="<user>",171    password="<password>",172    dsn="<hostname>/<service_name>",173)174print("Oracle connection is established...")175 176# params 177embedder_params = {"provider": "database", "model": "demo_model"}178proxy = ""179 180# instance181embedder = OracleEmbeddings(conn=conn, params=embedder_params, proxy=proxy)182 183docs = ["hello world!", "hi everyone!", "greetings!"]184embeds = embedder.embed_documents(docs)185print(f"Total Embeddings: {len(embeds)}")186print(f"Embedding generated by OracleEmbeddings: {embeds[0]}\n")187 188embed = embedder.embed_query("Hello World!")189print(f"Embedding generated by OracleEmbeddings: {embed}")190 191conn.close()192print("Connection is closed.")193 194"""195 
codekingpro/portable-devtools · Team Ai