upd: llm models

This commit is contained in:
namastex888 2024-06-15 02:32:17 +00:00
commit 5774a94898
2 changed files with 228 additions and 282 deletions

View file

@ -1,248 +1,228 @@
from typing import Any, List, Optional from typing import Any, Dict, List, Optional
import httpx
from langchain_community.chat_models.ollama import ChatOllama
from langchain_community.chat_models import ChatOllama
from langflow.base.constants import STREAM_INFO_TEXT from langflow.base.constants import STREAM_INFO_TEXT
from langflow.base.models.model import LCModelComponent from langflow.base.models.model import LCModelComponent
from langflow.field_typing import BaseLanguageModel, Text 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): class ChatOllamaComponent(LCModelComponent):
display_name = "Ollama" display_name = "Ollama"
description = "Generate text using Ollama Local LLMs." description = "Generate text using Ollama Local LLMs."
icon = "Ollama" 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 = [ inputs = [
Input( StrInput(
name="base_url", name="base_url",
field_type=Optional[str],
display_name="Base URL", display_name="Base URL",
info="Endpoint of the Ollama API. Defaults to 'http://localhost:11434' if not specified.", info="Endpoint of the Ollama API. Defaults to 'http://localhost:11434' if not specified.",
value="http://localhost:11434", advanced=True,
), ),
Input( StrInput(
name="model", name="model",
field_type=str,
display_name="Model Name", display_name="Model Name",
options=[], # This should be dynamically loaded if possible value="llama2",
info="Refer to https://ollama.ai/library for more models.", 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", name="mirostat",
field_type=str,
display_name="Mirostat", display_name="Mirostat",
options=["Disabled", "Mirostat", "Mirostat 2.0"], options=["Disabled", "Mirostat", "Mirostat 2.0"],
info="Enable/disable Mirostat sampling for controlling perplexity.", info="Enable/disable Mirostat sampling for controlling perplexity.",
advanced=False,
real_time_refresh=True,
refresh_button=True,
value="Disabled", value="Disabled",
advanced=True,
), ),
Input( FloatInput(
name="mirostat_eta", name="mirostat_eta",
field_type=Optional[float],
display_name="Mirostat Eta", display_name="Mirostat Eta",
info="Learning rate for Mirostat algorithm.", info="Learning rate for Mirostat algorithm. (Default: 0.1)",
advanced=True, advanced=True,
real_time_refresh=True,
value=None, # Default can vary based on mirostat status
), ),
Input( FloatInput(
name="mirostat_tau", name="mirostat_tau",
field_type=Optional[float],
display_name="Mirostat Tau", 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, advanced=True,
real_time_refresh=True,
value=None, # Default can vary based on mirostat status
), ),
Input( IntInput(
name="temperature", name="num_ctx",
field_type=float, display_name="Context Window Size",
display_name="Temperature", info="Size of the context window for generating tokens. (Default: 2048)",
info="Controls the creativity of model responses.", advanced=True,
value=0.8,
), ),
Input(name="input_value", type=str, display_name="Input", input_types=["Text", "Data", "Prompt"]), IntInput(
Input(name="stream", type=bool, display_name="Stream", info=STREAM_INFO_TEXT, value=False), name="num_gpu",
Input( 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", name="system_message",
field_type=Optional[str],
display_name="System Message", display_name="System Message",
info="System message to pass to the model.", info="System message to pass to the model.",
advanced=True, 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 = [ outputs = [
Output(display_name="Text", name="text_output", method="text_response"), 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: def text_response(self) -> Text:
input_value = self.input_value input_value = self.input_value
stream = self.stream stream = self.stream
system_message = self.system_message system_message = self.system_message
output = self.model_response() output = self.build_model()
result = self.get_chat_result(output, stream, input_value, system_message) result = self.get_chat_result(output, stream, input_value, system_message)
self.status = result self.status = result
return result return result
def model_response(self) -> BaseLanguageModel: def build_model(self) -> BaseLanguageModel:
base_url = self.base_url or "http://localhost:11434" # Mapping mirostat settings to their corresponding values
model = self.model mirostat_options = {"Mirostat": 1, "Mirostat 2.0": 2}
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
if keep_alive_flag == "Minute": # Default to 0 for 'Disabled'
keep_alive_instance = f"{keep_alive}m" mirostat_value = mirostat_options.get(self.mirostat, 0) # type: ignore
elif keep_alive_flag == "Hour":
keep_alive_instance = f"{keep_alive}h" # Set mirostat_eta and mirostat_tau to None if mirostat is disabled
elif keep_alive_flag == "sec": if mirostat_value == 0:
keep_alive_instance = f"{keep_alive}s" mirostat_eta = None
elif keep_alive_flag == "Keep": mirostat_tau = None
keep_alive_instance = "-1"
elif keep_alive_flag == "Immediately":
keep_alive_instance = "0"
else: else:
keep_alive_instance = "Invalid option" mirostat_eta = self.mirostat_eta
mirostat_tau = self.mirostat_tau
mirostat_instance = 0
if mirostat == "disable":
mirostat_instance = 0
# Mapping system settings to their corresponding values
llm_params = { llm_params = {
"base_url": base_url, "base_url": self.base_url,
"model": model, "model": self.model,
"mirostat": mirostat_instance, "mirostat": mirostat_value,
"keep_alive": keep_alive_instance, "format": self.format,
"format": _format, "metadata": self.metadata,
"metadata": metadata, "tags": self.tags.split(",") if self.tags else None,
"tags": tags,
"mirostat_eta": mirostat_eta, "mirostat_eta": mirostat_eta,
"mirostat_tau": mirostat_tau, "mirostat_tau": mirostat_tau,
"num_ctx": num_ctx, "num_ctx": self.num_ctx or None,
"num_gpu": num_gpu, "num_gpu": self.num_gpu or None,
"num_thread": num_thread, "num_thread": self.num_thread or None,
"repeat_last_n": repeat_last_n, "repeat_last_n": self.repeat_last_n or None,
"repeat_penalty": repeat_penalty, "repeat_penalty": self.repeat_penalty or None,
"temperature": temperature, "temperature": self.temperature or None,
"stop": stop, "stop": self.stop.split(",") if self.stop else None,
"system": system, "system": self.system,
"template": template, "template": self.template,
"tfs_z": tfs_z, "tfs_z": self.tfs_z or None,
"timeout": timeout, "timeout": self.timeout or None,
"top_k": top_k, "top_k": self.top_k or None,
"top_p": top_p, "top_p": self.top_p or None,
"verbose": verbose, "verbose": self.verbose,
"headers": headers,
} }
# Remove parameters with None values
llm_params = {k: v for k, v in llm_params.items() if v is not None} llm_params = {k: v for k, v in llm_params.items() if v is not None}
try: try:
output = ChatOllama(**llm_params) output = ChatOllama(**llm_params) # type: ignore
except Exception as e: except Exception as e:
raise ValueError("Could not initialize Ollama LLM.") from e raise ValueError("Could not initialize Ollama LLM.") from e

View file

@ -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.constants import STREAM_INFO_TEXT
from langflow.base.models.model import LCModelComponent 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): class ChatVertexAIComponent(LCModelComponent):
display_name = "Vertex AI" display_name = "Vertex AI"
description = "Generate text using Vertex AI LLMs." description = "Generate text using Vertex AI LLMs."
icon = "VertexAI" icon = "VertexAI"
field_order = [ inputs = [
"credentials", StrInput(name="input_value", display_name="Input", input_types=["Text", "Data", "Prompt"]),
"project", StrInput(
"examples", name="credentials",
"location", display_name="Credentials",
"max_output_tokens", info="Path to the JSON file containing the credentials.",
"model_name", file_types=[".json"],
"temperature", advanced=True,
"top_k", ),
"top_p", StrInput(name="project", display_name="Project", info="The project ID."),
"verbose", StrInput(
"input_value", name="examples",
"system_message", display_name="Examples",
"stream", 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): def text_response(self) -> Text:
return { input_value = self.input_value
"credentials": { stream = self.stream
"display_name": "Credentials", system_message = self.system_message
"field_type": "file", output = self.build_model()
"file_types": [".json"], result = self.get_chat_result(output, stream, input_value, system_message)
"file_path": None, self.status = result
}, return result
"examples": {
"display_name": "Examples", def build_model(self) -> BaseLanguageModel:
"multiline": True, credentials = self.credentials
}, location = self.location
"location": { max_output_tokens = self.max_output_tokens
"display_name": "Location", model_name = self.model_name
"value": "us-central1", project = self.project
}, temperature = self.temperature
"max_output_tokens": { top_k = self.top_k
"display_name": "Max Output Tokens", top_p = self.top_p
"value": 128, verbose = self.verbose
"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 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( output = ChatVertexAI(
credentials=credentials, credentials=credentials,
location=location, location=location,
@ -119,4 +84,5 @@ class ChatVertexAIComponent(LCModelComponent):
verbose=verbose, verbose=verbose,
) )
return self.get_chat_result(output, stream, input_value, system_message) return output