Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
oci_data_science.py1037 linesDownload Raw Back to chat_models
1# Copyright (c) 2024, Oracle and/or its affiliates.2 3"""Chat model for OCI data science model deployment endpoint."""4 5import importlib6import json7import logging8from operator import itemgetter9from typing import (10    Any,11    AsyncIterator,12    Callable,13    Dict,14    Iterator,15    List,16    Literal,17    Optional,18    Sequence,19    Type,20    Union,21)22 23from langchain_core.callbacks import (24    AsyncCallbackManagerForLLMRun,25    CallbackManagerForLLMRun,26)27from langchain_core.language_models import LanguageModelInput28from langchain_core.language_models.chat_models import (29    BaseChatModel,30    agenerate_from_stream,31    generate_from_stream,32)33from langchain_core.messages import (34    AIMessage,35    AIMessageChunk,36    BaseMessage,37    BaseMessageChunk,38)39from langchain_core.output_parsers import (40    JsonOutputParser,41    PydanticOutputParser,42)43from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult44from langchain_core.runnables import Runnable, RunnableMap, RunnablePassthrough45from langchain_core.tools import BaseTool46from langchain_core.utils.function_calling import convert_to_openai_tool47from pydantic import BaseModel, Field, model_validator48 49from langchain_community.llms.oci_data_science_model_deployment_endpoint import (50    DEFAULT_MODEL_NAME,51    BaseOCIModelDeployment,52)53 54logger = logging.getLogger(__name__)55DEFAULT_INFERENCE_ENDPOINT_CHAT = "/v1/chat/completions"56 57 58def _is_pydantic_class(obj: Any) -> bool:59    return isinstance(obj, type) and issubclass(obj, BaseModel)60 61 62class ChatOCIModelDeployment(BaseChatModel, BaseOCIModelDeployment):63    """OCI Data Science Model Deployment chat model integration.64 65    Prerequisite66        The OCI Model Deployment plugins are installable only on67        python version 3.9 and above. If you're working inside the notebook,68        try installing the python 3.10 based conda pack and running the69        following setup.70 71 72    Setup:73        Install ``oracle-ads`` and ``langchain-openai``.74 75        .. code-block:: bash76 77            pip install -U oracle-ads langchain-openai78 79        Use `ads.set_auth()` to configure authentication.80        For example, to use OCI resource_principal for authentication:81 82        .. code-block:: python83 84            import ads85            ads.set_auth("resource_principal")86 87        For more details on authentication, see:88        https://accelerated-data-science.readthedocs.io/en/latest/user_guide/cli/authentication.html89 90        Make sure to have the required policies to access the OCI Data91        Science Model Deployment endpoint. See:92        https://docs.oracle.com/en-us/iaas/data-science/using/model-dep-policies-auth.htm93 94 95    Key init args - completion params:96        endpoint: str97            The OCI model deployment endpoint.98        temperature: float99            Sampling temperature.100        max_tokens: Optional[int]101            Max number of tokens to generate.102 103    Key init args — client params:104        auth: dict105            ADS auth dictionary for OCI authentication.106        default_headers: Optional[Dict]107            The headers to be added to the Model Deployment request.108 109    Instantiate:110        .. code-block:: python111 112            from langchain_community.chat_models import ChatOCIModelDeployment113 114            chat = ChatOCIModelDeployment(115                endpoint="https://modeldeployment.<region>.oci.customer-oci.com/<ocid>/predict",116                model="odsc-llm", # this is the default model name if deployed with AQUA117                streaming=True,118                max_retries=3,119                model_kwargs={120                    "max_token": 512,121                    "temperature": 0.2,122                    # other model parameters ...123                },124                default_headers={125                    "route": "/v1/chat/completions",126                    # other request headers ...127                },128            )129 130    Invocation:131        .. code-block:: python132 133            messages = [134                ("system", "Translate the user sentence to French."),135                ("human", "Hello World!"),136            ]137            chat.invoke(messages)138 139        .. code-block:: python140 141            AIMessage(142                content='Bonjour le monde!',143                response_metadata={144                    'token_usage': {145                        'prompt_tokens': 40,146                        'total_tokens': 50,147                        'completion_tokens': 10148                    },149                    'model_name': 'odsc-llm',150                    'system_fingerprint': '',151                    'finish_reason': 'stop'152                },153                id='run-cbed62da-e1b3-4abd-9df3-ec89d69ca012-0'154            )155 156    Streaming:157        .. code-block:: python158 159            for chunk in chat.stream(messages):160                print(chunk)161 162        .. code-block:: python163 164            content='' id='run-02c6-c43f-42de'165            content='\n' id='run-02c6-c43f-42de'166            content='B' id='run-02c6-c43f-42de'167            content='on' id='run-02c6-c43f-42de'168            content='j' id='run-02c6-c43f-42de'169            content='our' id='run-02c6-c43f-42de'170            content=' le' id='run-02c6-c43f-42de'171            content=' monde' id='run-02c6-c43f-42de'172            content='!' id='run-02c6-c43f-42de'173            content='' response_metadata={'finish_reason': 'stop'} id='run-02c6-c43f-42de'174 175    Async:176        .. code-block:: python177 178            await chat.ainvoke(messages)179 180            # stream:181            # async for chunk in (await chat.astream(messages))182 183        .. code-block:: python184 185            AIMessage(186                content='Bonjour le monde!',187                response_metadata={'finish_reason': 'stop'},188                id='run-8657a105-96b7-4bb6-b98e-b69ca420e5d1-0'189            )190 191    Structured output:192        .. code-block:: python193 194            from typing import Optional195            from pydantic import BaseModel, Field196 197            class Joke(BaseModel):198                setup: str = Field(description="The setup of the joke")199                punchline: str = Field(description="The punchline to the joke")200 201            structured_llm = chat.with_structured_output(Joke, method="json_mode")202            structured_llm.invoke(203                "Tell me a joke about cats, "204                "respond in JSON with `setup` and `punchline` keys"205            )206 207        .. code-block:: python208 209            Joke(210                setup='Why did the cat get stuck in the tree?',211                punchline='Because it was chasing its tail!'212            )213 214        See ``ChatOCIModelDeployment.with_structured_output()`` for more.215 216    Customized Usage:217        You can inherit from base class and overwrite the `_process_response`,218        `_process_stream_response`, `_construct_json_body` for customized usage.219 220        .. code-block:: python221 222            class MyChatModel(ChatOCIModelDeployment):223                def _process_stream_response(self, response_json: dict) -> ChatGenerationChunk:224                    print("My customized streaming result handler.")225                    return GenerationChunk(...)226 227                def _process_response(self, response_json:dict) -> ChatResult:228                    print("My customized output handler.")229                    return ChatResult(...)230 231                def _construct_json_body(self, messages: list, params: dict) -> dict:232                    print("My customized payload handler.")233                    return {234                        "messages": messages,235                        **params,236                    }237 238            chat = MyChatModel(239                endpoint=f"https://modeldeployment.<region>.oci.customer-oci.com/{ocid}/predict",240                model="odsc-llm",241            }242 243            chat.invoke("tell me a joke")244 245    Response metadata246        .. code-block:: python247 248            ai_msg = chat.invoke(messages)249            ai_msg.response_metadata250 251        .. code-block:: python252 253            {254                'token_usage': {255                    'prompt_tokens': 40,256                    'total_tokens': 50,257                    'completion_tokens': 10258                },259                'model_name': 'odsc-llm',260                'system_fingerprint': '',261                'finish_reason': 'stop'262            }263 264    """  # noqa: E501265 266    model_kwargs: Dict[str, Any] = Field(default_factory=dict)267    """Keyword arguments to pass to the model."""268 269    model: str = DEFAULT_MODEL_NAME270    """The name of the model."""271 272    stop: Optional[List[str]] = None273    """Stop words to use when generating. Model output is cut off274    at the first occurrence of any of these substrings."""275 276    @model_validator(mode="before")277    @classmethod278    def validate_openai(cls, values: Any) -> Any:279        """Checks if langchain_openai is installed."""280        if not importlib.util.find_spec("langchain_openai"):281            raise ImportError(282                "Could not import langchain_openai package. "283                "Please install it with `pip install langchain_openai`."284            )285        return values286 287    @property288    def _llm_type(self) -> str:289        """Return type of llm."""290        return "oci_model_depolyment_chat_endpoint"291 292    @property293    def _identifying_params(self) -> Dict[str, Any]:294        """Get the identifying parameters."""295        _model_kwargs = self.model_kwargs or {}296        return {297            **{"endpoint": self.endpoint, "model_kwargs": _model_kwargs},298            **self._default_params,299        }300 301    @property302    def _default_params(self) -> Dict[str, Any]:303        """Get the default parameters."""304        return {305            "model": self.model,306            "stop": self.stop,307            "stream": self.streaming,308        }309 310    def _headers(311        self, is_async: Optional[bool] = False, body: Optional[dict] = None312    ) -> Dict:313        """Construct and return the headers for a request.314 315        Args:316            is_async (bool, optional): Indicates if the request is asynchronous.317                Defaults to `False`.318            body (optional): The request body to be included in the headers if319                the request is asynchronous.320 321        Returns:322            `dict` containing the appropriate headers for the request.323        """324        return {325            "route": DEFAULT_INFERENCE_ENDPOINT_CHAT,326            **super()._headers(is_async=is_async, body=body),327        }328 329    def _generate(330        self,331        messages: List[BaseMessage],332        stop: Optional[List[str]] = None,333        run_manager: Optional[CallbackManagerForLLMRun] = None,334        **kwargs: Any,335    ) -> ChatResult:336        """Call out to an OCI Model Deployment Online endpoint.337 338        Args:339            messages:  The messages in the conversation with the chat model.340            stop: Optional list of stop words to use when generating.341 342        Returns:343            LangChain ChatResult344 345        Raises:346            RuntimeError:347                Raise when invoking endpoint fails.348 349        Example:350 351            .. code-block:: python352 353                messages = [354                    (355                        "system",356                        "You are a helpful assistant that translates English to French. Translate the user sentence.",357                    ),358                    ("human", "Hello World!"),359                ]360 361                response = chat.invoke(messages)362        """  # noqa: E501363        if self.streaming:364            stream_iter = self._stream(365                messages, stop=stop, run_manager=run_manager, **kwargs366            )367            return generate_from_stream(stream_iter)368 369        requests_kwargs = kwargs.pop("requests_kwargs", {})370        params = self._invocation_params(stop, **kwargs)371        body = self._construct_json_body(messages, params)372        res = self.completion_with_retry(373            data=body, run_manager=run_manager, **requests_kwargs374        )375        return self._process_response(res.json())376 377    def _stream(378        self,379        messages: List[BaseMessage],380        stop: Optional[List[str]] = None,381        run_manager: Optional[CallbackManagerForLLMRun] = None,382        **kwargs: Any,383    ) -> Iterator[ChatGenerationChunk]:384        """Stream OCI Data Science Model Deployment endpoint on given messages.385 386        Args:387            messages (List[BaseMessage]):388                The messagaes to pass into the model.389            stop (List[str], Optional):390                List of stop words to use when generating.391            kwargs:392                requests_kwargs:393                    Additional ``**kwargs`` to pass to requests.post394 395        Returns:396            An iterator of ChatGenerationChunk.397 398        Raises:399            RuntimeError:400                Raise when invoking endpoint fails.401 402        Example:403 404            .. code-block:: python405 406                messages = [407                    (408                        "system",409                        "You are a helpful assistant that translates English to French. Translate the user sentence.",410                    ),411                    ("human", "Hello World!"),412                ]413 414                chunk_iter = chat.stream(messages)415 416        """  # noqa: E501417        requests_kwargs = kwargs.pop("requests_kwargs", {})418        self.streaming = True419        params = self._invocation_params(stop, **kwargs)420        body = self._construct_json_body(messages, params)  # request json body421 422        response = self.completion_with_retry(423            data=body, run_manager=run_manager, stream=True, **requests_kwargs424        )425        default_chunk_class = AIMessageChunk426        for line in self._parse_stream(response.iter_lines()):427            chunk = self._handle_sse_line(line, default_chunk_class)428            if run_manager:429                run_manager.on_llm_new_token(chunk.text, chunk=chunk)430            yield chunk431 432    async def _agenerate(433        self,434        messages: List[BaseMessage],435        stop: Optional[List[str]] = None,436        run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,437        **kwargs: Any,438    ) -> ChatResult:439        """Asynchronously call out to OCI Data Science Model Deployment440        endpoint on given messages.441 442        Args:443            messages (List[BaseMessage]):444                The messagaes to pass into the model.445            stop (List[str], Optional):446                List of stop words to use when generating.447            kwargs:448                requests_kwargs:449                    Additional ``**kwargs`` to pass to requests.post450 451        Returns:452            LangChain ChatResult.453 454        Raises:455            ValueError:456                Raise when invoking endpoint fails.457 458        Example:459 460            .. code-block:: python461 462                messages = [463                    (464                        "system",465                        "You are a helpful assistant that translates English to French. Translate the user sentence.",466                    ),467                    ("human", "I love programming."),468                ]469 470                resp = await chat.ainvoke(messages)471 472        """  # noqa: E501473        if self.streaming:474            stream_iter = self._astream(475                messages, stop=stop, run_manager=run_manager, **kwargs476            )477            return await agenerate_from_stream(stream_iter)478 479        requests_kwargs = kwargs.pop("requests_kwargs", {})480        params = self._invocation_params(stop, **kwargs)481        body = self._construct_json_body(messages, params)482        response = await self.acompletion_with_retry(483            data=body,484            run_manager=run_manager,485            **requests_kwargs,486        )487        return self._process_response(response)488 489    async def _astream(490        self,491        messages: List[BaseMessage],492        stop: Optional[List[str]] = None,493        run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,494        **kwargs: Any,495    ) -> AsyncIterator[ChatGenerationChunk]:496        """Asynchronously streaming OCI Data Science Model Deployment497        endpoint on given messages.498 499        Args:500            messages (List[BaseMessage]):501                The messagaes to pass into the model.502            stop (List[str], Optional):503                List of stop words to use when generating.504            kwargs:505                requests_kwargs:506                    Additional ``**kwargs`` to pass to requests.post507 508        Returns:509            An Asynciterator of ChatGenerationChunk.510 511        Raises:512            ValueError:513                Raise when invoking endpoint fails.514 515        Example:516 517            .. code-block:: python518 519                messages = [520                    (521                        "system",522                        "You are a helpful assistant that translates English to French. Translate the user sentence.",523                    ),524                    ("human", "I love programming."),525                ]526 527                chunk_iter = await chat.astream(messages)528 529        """  # noqa: E501530        requests_kwargs = kwargs.pop("requests_kwargs", {})531        self.streaming = True532        params = self._invocation_params(stop, **kwargs)533        body = self._construct_json_body(messages, params)  # request json body534 535        default_chunk_class = AIMessageChunk536        async for line in await self.acompletion_with_retry(537            data=body, run_manager=run_manager, stream=True, **requests_kwargs538        ):539            chunk = self._handle_sse_line(line, default_chunk_class)540            if run_manager:541                await run_manager.on_llm_new_token(chunk.text, chunk=chunk)542            yield chunk543 544    def with_structured_output(545        self,546        schema: Optional[Union[Dict, Type[BaseModel]]] = None,547        *,548        method: Literal["json_mode"] = "json_mode",549        include_raw: bool = False,550        **kwargs: Any,551    ) -> Runnable[LanguageModelInput, Union[Dict, BaseModel]]:552        """Model wrapper that returns outputs formatted to match the given schema.553 554        Args:555            schema: The output schema as a dict or a Pydantic class. If a Pydantic class556                then the model output will be an object of that class. If a dict then557                the model output will be a dict. With a Pydantic class the returned558                attributes will be validated, whereas with a dict they will not be. If559                `method` is "function_calling" and `schema` is a dict, then the dict560                must match the OpenAI function-calling spec.561            method: The method for steering model generation, currently only support562                for "json_mode". If "json_mode" then JSON mode will be used. Note that563                if using "json_mode" then you must include instructions for formatting564                the output into the desired schema into the model call.565            include_raw: If False then only the parsed structured output is returned. If566                an error occurs during model output parsing it will be raised. If True567                then both the raw model response (a BaseMessage) and the parsed model568                response will be returned. If an error occurs during output parsing it569                will be caught and returned as well. The final output is always a dict570                with keys "raw", "parsed", and "parsing_error".571 572        Returns:573            A Runnable that takes any ChatModel input and returns as output:574 575                If include_raw is True then a dict with keys:576                    raw: BaseMessage577                    parsed: Optional[_DictOrPydantic]578                    parsing_error: Optional[BaseException]579 580                If include_raw is False then just _DictOrPydantic is returned,581                where _DictOrPydantic depends on the schema:582 583                If schema is a Pydantic class then _DictOrPydantic is the Pydantic584                    class.585 586                If schema is a dict then _DictOrPydantic is a dict.587 588        """  # noqa: E501589        if kwargs:590            raise ValueError(f"Received unsupported arguments {kwargs}")591        is_pydantic_schema = _is_pydantic_class(schema)592        if method == "json_mode":593            llm = self.bind(response_format={"type": "json_object"})594            output_parser = (595                PydanticOutputParser(pydantic_object=schema)  # type: ignore[arg-type]596                if is_pydantic_schema597                else JsonOutputParser()598            )599        else:600            raise ValueError(601                f"Unrecognized method argument. Expected `json_mode`."602                f"Received: `{method}`."603            )604 605        if include_raw:606            parser_assign = RunnablePassthrough.assign(607                parsed=itemgetter("raw") | output_parser, parsing_error=lambda _: None608            )609            parser_none = RunnablePassthrough.assign(parsed=lambda _: None)610            parser_with_fallback = parser_assign.with_fallbacks(611                [parser_none], exception_key="parsing_error"612            )613            return RunnableMap(raw=llm) | parser_with_fallback614        else:615            return llm | output_parser616 617    def _invocation_params(self, stop: Optional[List[str]], **kwargs: Any) -> dict:618        """Combines the invocation parameters with default parameters."""619        params = self._default_params620        _model_kwargs = self.model_kwargs or {}621        params["stop"] = stop or params.get("stop", [])622        return {**params, **_model_kwargs, **kwargs}623 624    def _handle_sse_line(625        self, line: str, default_chunk_cls: Type[BaseMessageChunk] = AIMessageChunk626    ) -> ChatGenerationChunk:627        """Handle a single Server-Sent Events (SSE) line and process it into628        a chat generation chunk.629 630        Args:631            line (str): A single line from the SSE stream in string format.632            default_chunk_cls (AIMessageChunk): The default class for message633                chunks to be used during the processing of the stream response.634 635        Returns:636            ChatGenerationChunk: The processed chat generation chunk. If an error637                occurs, an empty `ChatGenerationChunk` is returned.638        """639        try:640            obj = json.loads(line)641            return self._process_stream_response(obj, default_chunk_cls)642        except Exception as e:643            logger.debug(f"Error occurs when processing line={line}: {str(e)}")644            return ChatGenerationChunk(message=AIMessageChunk(content=""))645 646    def _construct_json_body(self, messages: list, params: dict) -> dict:647        """Constructs the request body as a dictionary (JSON).648 649        Args:650            messages (list): A list of message objects to be included in the651                request body.652            params (dict): A dictionary of additional parameters to be included653                in the request body.654 655        Returns:656            dict: A dictionary representing the JSON request body, including657                converted messages and additional parameters.658 659        """660        from langchain_openai.chat_models.base import _convert_message_to_dict661 662        return {663            "messages": [_convert_message_to_dict(m) for m in messages],664            **params,665        }666 667    def _process_stream_response(668        self,669        response_json: dict,670        default_chunk_cls: Type[BaseMessageChunk] = AIMessageChunk,671    ) -> ChatGenerationChunk:672        """Formats streaming response in OpenAI spec.673 674        Args:675            response_json (dict): The JSON response from the streaming endpoint.676            default_chunk_cls (type, optional): The default class to use for677                creating message chunks. Defaults to `AIMessageChunk`.678 679        Returns:680            ChatGenerationChunk: An object containing the processed message681                chunk and any relevant generation information such as finish682                reason and usage.683 684        Raises:685            ValueError: If the response JSON is not well-formed or does not686                contain the expected structure.687        """688        from langchain_openai.chat_models.base import _convert_delta_to_message_chunk689 690        try:691            choice = response_json["choices"][0]692            if not isinstance(choice, dict):693                raise TypeError("Endpoint response is not well formed.")694        except (KeyError, IndexError, TypeError) as e:695            raise ValueError(696                "Error while formatting response payload for chat model of type"697            ) from e698 699        chunk = _convert_delta_to_message_chunk(choice["delta"], default_chunk_cls)700        default_chunk_cls = chunk.__class__701        finish_reason = choice.get("finish_reason")702        usage = choice.get("usage")703        gen_info = {}704        if finish_reason is not None:705            gen_info.update({"finish_reason": finish_reason})706        if usage is not None:707            gen_info.update({"usage": usage})708 709        return ChatGenerationChunk(710            message=chunk, generation_info=gen_info if gen_info else None711        )712 713    def _process_response(self, response_json: dict) -> ChatResult:714        """Formats response in OpenAI spec.715 716        Args:717            response_json (dict): The JSON response from the chat model endpoint.718 719        Returns:720            ChatResult: An object containing the list of `ChatGeneration` objects721            and additional LLM output information.722 723        Raises:724            ValueError: If the response JSON is not well-formed or does not725            contain the expected structure.726 727        """728        from langchain_openai.chat_models.base import _convert_dict_to_message729 730        generations = []731        try:732            choices = response_json["choices"]733            if not isinstance(choices, list):734                raise TypeError("Endpoint response is not well formed.")735        except (KeyError, TypeError) as e:736            raise ValueError(737                "Error while formatting response payload for chat model of type"738            ) from e739 740        for choice in choices:741            message = _convert_dict_to_message(choice["message"])742            generation_info = {"finish_reason": choice.get("finish_reason")}743            if "logprobs" in choice:744                generation_info["logprobs"] = choice["logprobs"]745 746            gen = ChatGeneration(747                message=message,748                generation_info=generation_info,749            )750            generations.append(gen)751 752        token_usage = response_json.get("usage", {})753        llm_output = {754            "token_usage": token_usage,755            "model_name": self.model,756            "system_fingerprint": response_json.get("system_fingerprint", ""),757        }758        return ChatResult(generations=generations, llm_output=llm_output)759 760    def bind_tools(761        self,762        tools: Sequence[Union[Dict[str, Any], Type[BaseModel], Callable, BaseTool]],763        **kwargs: Any,764    ) -> Runnable[LanguageModelInput, AIMessage]:765        formatted_tools = [convert_to_openai_tool(tool) for tool in tools]766        return super().bind(tools=formatted_tools, **kwargs)767 768 769class ChatOCIModelDeploymentVLLM(ChatOCIModelDeployment):770    """OCI large language chat models deployed with vLLM.771 772    To use, you must provide the model HTTP endpoint from your deployed773    model, e.g. https://modeldeployment.us-ashburn-1.oci.customer-oci.com/<ocid>/predict.774 775    To authenticate, `oracle-ads` has been used to automatically load776    credentials: https://accelerated-data-science.readthedocs.io/en/latest/user_guide/cli/authentication.html777 778    Make sure to have the required policies to access the OCI Data779    Science Model Deployment endpoint. See:780    https://docs.oracle.com/en-us/iaas/data-science/using/model-dep-policies-auth.htm#model_dep_policies_auth__predict-endpoint781 782    Example:783 784        .. code-block:: python785 786            from langchain_community.chat_models import ChatOCIModelDeploymentVLLM787 788            chat = ChatOCIModelDeploymentVLLM(789                endpoint="https://modeldeployment.us-ashburn-1.oci.customer-oci.com/<ocid>/predict",790                frequency_penalty=0.1,791                max_tokens=512,792                temperature=0.2,793                top_p=1.0,794                # other model parameters...795            )796 797    """  # noqa: E501798 799    frequency_penalty: float = 0.0800    """Penalizes repeated tokens according to frequency. Between 0 and 1."""801 802    logit_bias: Optional[Dict[str, float]] = None803    """Adjust the probability of specific tokens being generated."""804 805    max_tokens: Optional[int] = 256806    """The maximum number of tokens to generate in the completion."""807 808    n: int = 1809    """Number of output sequences to return for the given prompt."""810 811    presence_penalty: float = 0.0812    """Penalizes repeated tokens. Between 0 and 1."""813 814    temperature: float = 0.2815    """What sampling temperature to use."""816 817    top_p: float = 1.0818    """Total probability mass of tokens to consider at each step."""819 820    best_of: Optional[int] = None821    """Generates best_of completions server-side and returns the "best"822    (the one with the highest log probability per token).823    """824 825    use_beam_search: Optional[bool] = False826    """Whether to use beam search instead of sampling."""827 828    top_k: Optional[int] = -1829    """Number of most likely tokens to consider at each step."""830 831    min_p: Optional[float] = 0.0832    """Float that represents the minimum probability for a token to be considered.833    Must be in [0,1]. 0 to disable this."""834 835    repetition_penalty: Optional[float] = 1.0836    """Float that penalizes new tokens based on their frequency in the837    generated text. Values > 1 encourage the model to use new tokens."""838 839    length_penalty: Optional[float] = 1.0840    """Float that penalizes sequences based on their length. Used only841    when `use_beam_search` is True."""842 843    early_stopping: Optional[bool] = False844    """Controls the stopping condition for beam search. It accepts the845    following values: `True`, where the generation stops as soon as there846    are `best_of` complete candidates; `False`, where a heuristic is applied847    to the generation stops when it is very unlikely to find better candidates;848    `never`, where the beam search procedure only stops where there cannot be849    better candidates (canonical beam search algorithm)."""850 851    ignore_eos: Optional[bool] = False852    """Whether to ignore the EOS token and continue generating tokens after853    the EOS token is generated."""854 855    min_tokens: Optional[int] = 0856    """Minimum number of tokens to generate per output sequence before857    EOS or stop_token_ids can be generated"""858 859    stop_token_ids: Optional[List[int]] = None860    """List of tokens that stop the generation when they are generated.861    The returned output will contain the stop tokens unless the stop tokens862    are special tokens."""863 864    skip_special_tokens: Optional[bool] = True865    """Whether to skip special tokens in the output. Defaults to True."""866 867    spaces_between_special_tokens: Optional[bool] = True868    """Whether to add spaces between special tokens in the output.869    Defaults to True."""870 871    tool_choice: Optional[str] = None872    """Whether to use tool calling.873    Defaults to None, tool calling is disabled.874    Tool calling requires model support and the vLLM to be configured875    with `--tool-call-parser`.876    Set this to `auto` for the model to make tool calls automatically.877    Set this to `required` to force the model to always call one or more tools.878    """879 880    chat_template: Optional[str] = None881    """Use customized chat template.882    Defaults to None. The chat template from the tokenizer will be used.883    """884 885    @property886    def _llm_type(self) -> str:887        """Return type of llm."""888        return "oci_model_depolyment_chat_endpoint_vllm"889 890    @property891    def _default_params(self) -> Dict[str, Any]:892        """Get the default parameters."""893        params = {894            "model": self.model,895            "stop": self.stop,896            "stream": self.streaming,897        }898        for attr_name in self._get_model_params():899            try:900                value = getattr(self, attr_name)901                if value is not None:902                    params.update({attr_name: value})903            except Exception:904                pass905 906        return params907 908    def _get_model_params(self) -> List[str]:909        """Gets the name of model parameters."""910        return [911            "best_of",912            "early_stopping",913            "frequency_penalty",914            "ignore_eos",915            "length_penalty",916            "logit_bias",917            "logprobs",918            "max_tokens",919            "min_p",920            "min_tokens",921            "n",922            "presence_penalty",923            "repetition_penalty",924            "skip_special_tokens",925            "spaces_between_special_tokens",926            "stop_token_ids",927            "temperature",928            "top_k",929            "top_p",930            "use_beam_search",931            "tool_choice",932            "chat_template",933        ]934 935 936class ChatOCIModelDeploymentTGI(ChatOCIModelDeployment):937    """OCI large language chat models deployed with Text Generation Inference.938 939    To use, you must provide the model HTTP endpoint from your deployed940    model, e.g. https://modeldeployment.us-ashburn-1.oci.customer-oci.com/<ocid>/predict.941 942    To authenticate, `oracle-ads` has been used to automatically load943    credentials: https://accelerated-data-science.readthedocs.io/en/latest/user_guide/cli/authentication.html944 945    Make sure to have the required policies to access the OCI Data946    Science Model Deployment endpoint. See:947    https://docs.oracle.com/en-us/iaas/data-science/using/model-dep-policies-auth.htm#model_dep_policies_auth__predict-endpoint948 949    Example:950 951        .. code-block:: python952 953            from langchain_community.chat_models import ChatOCIModelDeploymentTGI954 955            chat = ChatOCIModelDeploymentTGI(956                endpoint="https://modeldeployment.us-ashburn-1.oci.customer-oci.com/<ocid>/predict",957                max_token=512,958                temperature=0.2,959                frequency_penalty=0.1,960                seed=42,961                # other model parameters...962            )963 964    """  # noqa: E501965 966    frequency_penalty: Optional[float] = None967    """Penalizes repeated tokens according to frequency. Between 0 and 1."""968 969    logit_bias: Optional[Dict[str, float]] = None970    """Adjust the probability of specific tokens being generated."""971 972    logprobs: Optional[bool] = None973    """Whether to return log probabilities of the output tokens or not."""974 975    max_tokens: int = 256976    """The maximum number of tokens to generate in the completion."""977 978    n: int = 1979    """Number of output sequences to return for the given prompt."""980 981    presence_penalty: Optional[float] = None982    """Penalizes repeated tokens. Between 0 and 1."""983 984    seed: Optional[int] = None985    """To sample deterministically,"""986 987    temperature: float = 0.2988    """What sampling temperature to use."""989 990    top_p: Optional[float] = None991    """Total probability mass of tokens to consider at each step."""992 993    top_logprobs: Optional[int] = None994    """An integer between 0 and 5 specifying the number of most995    likely tokens to return at each token position, each with an996    associated log probability. logprobs must be set to true if997    this parameter is used."""998 999    @property1000    def _llm_type(self) -> str:1001        """Return type of llm."""1002        return "oci_model_depolyment_chat_endpoint_tgi"1003 1004    @property1005    def _default_params(self) -> Dict[str, Any]:1006        """Get the default parameters."""1007        params = {1008            "model": self.model,1009            "stop": self.stop,1010            "stream": self.streaming,1011        }1012        for attr_name in self._get_model_params():1013            try:1014                value = getattr(self, attr_name)1015                if value is not None:1016                    params.update({attr_name: value})1017            except Exception:1018                pass1019 1020        return params1021 1022    def _get_model_params(self) -> List[str]:1023        """Gets the name of model parameters."""1024        return [1025            "frequency_penalty",1026            "logit_bias",1027            "logprobs",1028            "max_tokens",1029            "n",1030            "presence_penalty",1031            "seed",1032            "temperature",1033            "top_k",1034            "top_p",1035            "top_logprobs",1036        ]1037