From 5774a9489803d8255d5841d410eb4fb1de7b7759 Mon Sep 17 00:00:00 2001 From: namastex888 Date: Sat, 15 Jun 2024 02:32:17 +0000 Subject: [PATCH] upd: llm models --- .../langflow/components/models/OllamaModel.py | 346 +++++++++--------- .../components/models/VertexAiModel.py | 164 ++++----- 2 files changed, 228 insertions(+), 282 deletions(-) diff --git a/src/backend/base/langflow/components/models/OllamaModel.py b/src/backend/base/langflow/components/models/OllamaModel.py index 8e999f83d..bcca0b107 100644 --- a/src/backend/base/langflow/components/models/OllamaModel.py +++ b/src/backend/base/langflow/components/models/OllamaModel.py @@ -1,249 +1,229 @@ -from typing import Any, List, Optional - -import httpx -from langchain_community.chat_models.ollama import ChatOllama +from typing import Any, Dict, List, Optional +from langchain_community.chat_models import ChatOllama from langflow.base.constants import STREAM_INFO_TEXT from langflow.base.models.model import LCModelComponent from langflow.field_typing import BaseLanguageModel, Text -from langflow.template.field.base import Input, Output - +from langflow.inputs import BoolInput, DictInput, DropdownInput, FloatInput, IntInput, StrInput +from langflow.template import Output class ChatOllamaComponent(LCModelComponent): display_name = "Ollama" description = "Generate text using Ollama Local LLMs." icon = "Ollama" - def update_build_config(self, build_config: dict, field_value: Any, field_name: str | None = None): - if field_name == "mirostat": - if field_value == "Disabled": - build_config["mirostat_eta"]["advanced"] = True - build_config["mirostat_tau"]["advanced"] = True - build_config["mirostat_eta"]["value"] = None - build_config["mirostat_tau"]["value"] = None - - else: - build_config["mirostat_eta"]["advanced"] = False - build_config["mirostat_tau"]["advanced"] = False - - if field_value == "Mirostat 2.0": - build_config["mirostat_eta"]["value"] = 0.2 - build_config["mirostat_tau"]["value"] = 10 - else: - build_config["mirostat_eta"]["value"] = 0.1 - build_config["mirostat_tau"]["value"] = 5 - - if field_name == "model": - base_url_dict = build_config.get("base_url", {}) - base_url_load_from_db = base_url_dict.get("load_from_db", False) - base_url_value = base_url_dict.get("value") - if base_url_load_from_db: - base_url_value = self.variables(base_url_value) - elif not base_url_value: - base_url_value = "http://localhost:11434" - build_config["model"]["options"] = self.get_model(base_url_value + "/api/tags") - - if field_name == "keep_alive_flag": - if field_value == "Keep": - build_config["keep_alive"]["value"] = "-1" - build_config["keep_alive"]["advanced"] = True - elif field_value == "Immediately": - build_config["keep_alive"]["value"] = "0" - build_config["keep_alive"]["advanced"] = True - else: - build_config["keep_alive"]["advanced"] = False - - return build_config - - def get_model(self, url: str) -> List[str]: - try: - with httpx.Client() as client: - response = client.get(url) - response.raise_for_status() - data = response.json() - - model_names = [model["name"] for model in data.get("models", [])] - return model_names - except Exception as e: - raise ValueError("Could not retrieve models") from e - inputs = [ - Input( + StrInput( name="base_url", - field_type=Optional[str], display_name="Base URL", info="Endpoint of the Ollama API. Defaults to 'http://localhost:11434' if not specified.", - value="http://localhost:11434", + advanced=True, ), - Input( + StrInput( name="model", - field_type=str, display_name="Model Name", - options=[], # This should be dynamically loaded if possible + value="llama2", info="Refer to https://ollama.ai/library for more models.", - real_time_refresh=True, - refresh_button=True, ), - Input( + FloatInput( + name="temperature", + display_name="Temperature", + value=0.8, + info="Controls the creativity of model responses.", + ), + StrInput( + name="format", + display_name="Format", + info="Specify the format of the output (e.g., json).", + advanced=True, + ), + DictInput( + name="metadata", + display_name="Metadata", + info="Metadata to add to the run trace.", + advanced=True, + ), + DropdownInput( name="mirostat", - field_type=str, display_name="Mirostat", options=["Disabled", "Mirostat", "Mirostat 2.0"], info="Enable/disable Mirostat sampling for controlling perplexity.", - advanced=False, - real_time_refresh=True, - refresh_button=True, value="Disabled", + advanced=True, ), - Input( + FloatInput( name="mirostat_eta", - field_type=Optional[float], display_name="Mirostat Eta", - info="Learning rate for Mirostat algorithm.", + info="Learning rate for Mirostat algorithm. (Default: 0.1)", advanced=True, - real_time_refresh=True, - value=None, # Default can vary based on mirostat status ), - Input( + FloatInput( name="mirostat_tau", - field_type=Optional[float], display_name="Mirostat Tau", - info="Controls the balance between coherence and diversity of the output.", + info="Controls the balance between coherence and diversity of the output. (Default: 5.0)", advanced=True, - real_time_refresh=True, - value=None, # Default can vary based on mirostat status ), - Input( - name="temperature", - field_type=float, - display_name="Temperature", - info="Controls the creativity of model responses.", - value=0.8, + IntInput( + name="num_ctx", + display_name="Context Window Size", + info="Size of the context window for generating tokens. (Default: 2048)", + advanced=True, ), - Input(name="input_value", type=str, display_name="Input", input_types=["Text", "Data", "Prompt"]), - Input(name="stream", type=bool, display_name="Stream", info=STREAM_INFO_TEXT, value=False), - Input( + IntInput( + name="num_gpu", + display_name="Number of GPUs", + info="Number of GPUs to use for computation. (Default: 1 on macOS, 0 to disable)", + advanced=True, + ), + IntInput( + name="num_thread", + display_name="Number of Threads", + info="Number of threads to use during computation. (Default: detected for optimal performance)", + advanced=True, + ), + IntInput( + name="repeat_last_n", + display_name="Repeat Last N", + info="How far back the model looks to prevent repetition. (Default: 64, 0 = disabled, -1 = num_ctx)", + advanced=True, + ), + FloatInput( + name="repeat_penalty", + display_name="Repeat Penalty", + info="Penalty for repetitions in generated text. (Default: 1.1)", + advanced=True, + ), + FloatInput( + name="tfs_z", + display_name="TFS Z", + info="Tail free sampling value. (Default: 1)", + advanced=True, + ), + IntInput( + name="timeout", + display_name="Timeout", + info="Timeout for the request stream.", + advanced=True, + ), + IntInput( + name="top_k", + display_name="Top K", + info="Limits token selection to top K. (Default: 40)", + advanced=True, + ), + FloatInput( + name="top_p", + display_name="Top P", + info="Works together with top-k. (Default: 0.9)", + advanced=True, + ), + BoolInput( + name="verbose", + display_name="Verbose", + info="Whether to print out response text.", + ), + StrInput( + name="tags", + display_name="Tags", + info="Comma-separated list of tags to add to the run trace.", + advanced=True, + ), + StrInput( + name="stop", + display_name="Stop Tokens", + info="Comma-separated list of tokens to signal the model to stop generating text.", + advanced=True, + ), + StrInput( + name="system", + display_name="System", + info="System to use for generating text.", + advanced=True, + ), + StrInput( + name="template", + display_name="Template", + info="Template to use for generating text.", + advanced=True, + ), + StrInput( + name="input_value", + display_name="Input", + input_types=["Text", "Data", "Prompt"], + ), + BoolInput( + name="stream", + display_name="Stream", + info=STREAM_INFO_TEXT, + ), + StrInput( name="system_message", - field_type=Optional[str], display_name="System Message", info="System message to pass to the model.", advanced=True, - value=None, - ), - Input( - name="headers", - field_type=dict, - display_name="Headers", - info="Additional headers to send with the request.", - advanced=True, - ), - Input( - name="keep_alive_flag", - field_type=str, - display_params=["Keep", "Immediately", "Minute", "Hour", "sec"], - display_name="Unload interval", - info="Controls how the model unload interval is managed.", - real_time_refresh=True, - refresh_button=True, - ), - Input( - name="keep_alive", - field_type=int, - display_name="Interval", - info="How long the model will stay loaded into memory.", - value=None, ), ] outputs = [ Output(display_name="Text", name="text_output", method="text_response"), - Output(display_name="Language Model", name="model_output", method="model_response"), + Output(display_name="Language Model", name="model_output", method="build_model"), ] def text_response(self) -> Text: input_value = self.input_value stream = self.stream system_message = self.system_message - output = self.model_response() + output = self.build_model() result = self.get_chat_result(output, stream, input_value, system_message) self.status = result return result - def model_response(self) -> BaseLanguageModel: - base_url = self.base_url or "http://localhost:11434" - model = self.model - mirostat = self.mirostat or "Disabled" - mirostat_eta = self.mirostat_eta - mirostat_tau = self.mirostat_tau - repeat_last_n = self.repeat_last_n - verbose = self.verbose - keep_alive = self.keep_alive - keep_alive_flag = self.keep_alive_flag or "Keep" - num_ctx = self.num_ctx - num_gpu = self.num_gpu - _format = self.format - metadata = self.metadata - num_thread = self.num_thread - repeat_penalty = self.repeat_penalty - stop = self.stop - system = self.system - tags = self.tags - temperature = self.temperature - template = self.template - tfs_z = self.tfs_z - timeout = self.timeout - top_k = self.top_k - top_p = self.top_p - headers = self.headers + def build_model(self) -> BaseLanguageModel: + # Mapping mirostat settings to their corresponding values + mirostat_options = {"Mirostat": 1, "Mirostat 2.0": 2} - if keep_alive_flag == "Minute": - keep_alive_instance = f"{keep_alive}m" - elif keep_alive_flag == "Hour": - keep_alive_instance = f"{keep_alive}h" - elif keep_alive_flag == "sec": - keep_alive_instance = f"{keep_alive}s" - elif keep_alive_flag == "Keep": - keep_alive_instance = "-1" - elif keep_alive_flag == "Immediately": - keep_alive_instance = "0" + # Default to 0 for 'Disabled' + mirostat_value = mirostat_options.get(self.mirostat, 0) # type: ignore + + # Set mirostat_eta and mirostat_tau to None if mirostat is disabled + if mirostat_value == 0: + mirostat_eta = None + mirostat_tau = None else: - keep_alive_instance = "Invalid option" - - mirostat_instance = 0 - if mirostat == "disable": - mirostat_instance = 0 + mirostat_eta = self.mirostat_eta + mirostat_tau = self.mirostat_tau + # Mapping system settings to their corresponding values llm_params = { - "base_url": base_url, - "model": model, - "mirostat": mirostat_instance, - "keep_alive": keep_alive_instance, - "format": _format, - "metadata": metadata, - "tags": tags, + "base_url": self.base_url, + "model": self.model, + "mirostat": mirostat_value, + "format": self.format, + "metadata": self.metadata, + "tags": self.tags.split(",") if self.tags else None, "mirostat_eta": mirostat_eta, "mirostat_tau": mirostat_tau, - "num_ctx": num_ctx, - "num_gpu": num_gpu, - "num_thread": num_thread, - "repeat_last_n": repeat_last_n, - "repeat_penalty": repeat_penalty, - "temperature": temperature, - "stop": stop, - "system": system, - "template": template, - "tfs_z": tfs_z, - "timeout": timeout, - "top_k": top_k, - "top_p": top_p, - "verbose": verbose, - "headers": headers, + "num_ctx": self.num_ctx or None, + "num_gpu": self.num_gpu or None, + "num_thread": self.num_thread or None, + "repeat_last_n": self.repeat_last_n or None, + "repeat_penalty": self.repeat_penalty or None, + "temperature": self.temperature or None, + "stop": self.stop.split(",") if self.stop else None, + "system": self.system, + "template": self.template, + "tfs_z": self.tfs_z or None, + "timeout": self.timeout or None, + "top_k": self.top_k or None, + "top_p": self.top_p or None, + "verbose": self.verbose, } + # Remove parameters with None values llm_params = {k: v for k, v in llm_params.items() if v is not None} try: - output = ChatOllama(**llm_params) + output = ChatOllama(**llm_params) # type: ignore except Exception as e: raise ValueError("Could not initialize Ollama LLM.") from e - return output + return output \ No newline at end of file diff --git a/src/backend/base/langflow/components/models/VertexAiModel.py b/src/backend/base/langflow/components/models/VertexAiModel.py index aed218c08..c494fb07a 100644 --- a/src/backend/base/langflow/components/models/VertexAiModel.py +++ b/src/backend/base/langflow/components/models/VertexAiModel.py @@ -1,112 +1,77 @@ -from typing import Optional - +from langchain_google_vertexai import ChatVertexAI from langflow.base.constants import STREAM_INFO_TEXT from langflow.base.models.model import LCModelComponent -from langflow.field_typing import Text - +from langflow.field_typing import BaseLanguageModel, Text +from langflow.inputs import BoolInput, FloatInput, IntInput, StrInput +from langflow.template import Output class ChatVertexAIComponent(LCModelComponent): display_name = "Vertex AI" description = "Generate text using Vertex AI LLMs." icon = "VertexAI" - field_order = [ - "credentials", - "project", - "examples", - "location", - "max_output_tokens", - "model_name", - "temperature", - "top_k", - "top_p", - "verbose", - "input_value", - "system_message", - "stream", + inputs = [ + StrInput(name="input_value", display_name="Input", input_types=["Text", "Data", "Prompt"]), + StrInput( + name="credentials", + display_name="Credentials", + info="Path to the JSON file containing the credentials.", + file_types=[".json"], + advanced=True, + ), + StrInput(name="project", display_name="Project", info="The project ID."), + StrInput( + name="examples", + display_name="Examples", + info="Examples to pass to the model.", + multiline=True, + advanced=True, + ), + StrInput(name="location", display_name="Location", value="us-central1", advanced=True), + IntInput( + name="max_output_tokens", + display_name="Max Output Tokens", + value=128, + advanced=True, + ), + StrInput(name="model_name", display_name="Model Name", value="gemini-1.5-pro"), + FloatInput(name="temperature", display_name="Temperature", value=0.0), + IntInput(name="top_k", display_name="Top K", value=40, advanced=True), + FloatInput(name="top_p", display_name="Top P", value=0.95, advanced=True), + BoolInput(name="verbose", display_name="Verbose", value=False, advanced=True), + BoolInput(name="stream", display_name="Stream", info=STREAM_INFO_TEXT, advanced=True), + StrInput( + name="system_message", + display_name="System Message", + info="System message to pass to the model.", + advanced=True, + ), + ] + outputs = [ + Output(display_name="Text", name="text_output", method="text_response"), + Output(display_name="Language Model", name="model_output", method="build_model"), ] - def build_config(self): - return { - "credentials": { - "display_name": "Credentials", - "field_type": "file", - "file_types": [".json"], - "file_path": None, - }, - "examples": { - "display_name": "Examples", - "multiline": True, - }, - "location": { - "display_name": "Location", - "value": "us-central1", - }, - "max_output_tokens": { - "display_name": "Max Output Tokens", - "value": 128, - "advanced": True, - }, - "model_name": { - "display_name": "Model Name", - "value": "chat-bison", - }, - "project": { - "display_name": "Project", - }, - "temperature": { - "display_name": "Temperature", - "value": 0.0, - }, - "top_k": { - "display_name": "Top K", - "value": 40, - "advanced": True, - }, - "top_p": { - "display_name": "Top P", - "value": 0.95, - "advanced": True, - }, - "verbose": { - "display_name": "Verbose", - "value": False, - "advanced": True, - }, - "input_value": {"display_name": "Input", "input_types": ["Text", "Data", "Prompt"]}, - "stream": { - "display_name": "Stream", - "info": STREAM_INFO_TEXT, - "advanced": True, - }, - "system_message": { - "display_name": "System Message", - "info": "System message to pass to the model.", - "advanced": True, - }, - } + def text_response(self) -> Text: + input_value = self.input_value + stream = self.stream + system_message = self.system_message + output = self.build_model() + result = self.get_chat_result(output, stream, input_value, system_message) + self.status = result + return result + + def build_model(self) -> BaseLanguageModel: + credentials = self.credentials + location = self.location + max_output_tokens = self.max_output_tokens + model_name = self.model_name + project = self.project + temperature = self.temperature + top_k = self.top_k + top_p = self.top_p + verbose = self.verbose - def build( - self, - input_value: Text, - credentials: Optional[str], - project: str, - location: str = "us-central1", - max_output_tokens: int = 128, - model_name: str = "chat-bison", - temperature: float = 0.0, - top_k: int = 40, - top_p: float = 0.95, - verbose: bool = False, - stream: bool = False, - system_message: Optional[str] = None, - ) -> Text: - try: - from langchain_google_vertexai import ChatVertexAI # type: ignore - except ImportError: - raise ImportError( - "To use the ChatVertexAI model, you need to install the langchain-google-vertexai package." - ) output = ChatVertexAI( credentials=credentials, location=location, @@ -119,4 +84,5 @@ class ChatVertexAIComponent(LCModelComponent): verbose=verbose, ) - return self.get_chat_result(output, stream, input_value, system_message) + return output + \ No newline at end of file