mirror of
https://github.com/VinciGit00/Scrapegraph-ai.git
synced 2026-07-20 21:10:05 +08:00
Update abstract_graph.py
This commit is contained in:
parent
8cece1d31b
commit
c105c26b2f
@ -146,78 +146,84 @@ class AbstractGraph(ABC):
|
||||
with warnings.catch_warnings():
|
||||
warnings.simplefilter("ignore")
|
||||
return init_chat_model(**llm_params)
|
||||
|
||||
known_models = ["azure", "fireworks", "gemini", "claude", "vertexai", "hugging_face", "groq", "gpt-", "ollama", "claude-3-", "bedrock", "mistral", "ernie", "oneapi", "nvidia"]
|
||||
|
||||
if "fireworks" in llm_params["model"]:
|
||||
model_name = "/".join(llm_params["model"].split("/")[1:])
|
||||
token_key = llm_params["model"].split("/")[-1]
|
||||
return handle_model(model_name, "fireworks", token_key)
|
||||
if llm_params["model"] not in known_models:
|
||||
raise ValueError(f"Model '{llm_params['model']}' is not supported")
|
||||
|
||||
elif "gemini" in llm_params["model"]:
|
||||
model_name = llm_params["model"].split("/")[-1]
|
||||
return handle_model(model_name, "google_genai", model_name)
|
||||
try:
|
||||
if "fireworks" in llm_params["model"]:
|
||||
model_name = "/".join(llm_params["model"].split("/")[1:])
|
||||
token_key = llm_params["model"].split("/")[-1]
|
||||
return handle_model(model_name, "fireworks", token_key)
|
||||
|
||||
elif llm_params["model"].startswith("claude"):
|
||||
model_name = llm_params["model"].split("/")[-1]
|
||||
return handle_model(model_name, "anthropic", model_name)
|
||||
elif "gemini" in llm_params["model"]:
|
||||
model_name = llm_params["model"].split("/")[-1]
|
||||
return handle_model(model_name, "google_genai", model_name)
|
||||
|
||||
elif llm_params["model"].startswith("vertexai"):
|
||||
return handle_model(llm_params["model"], "google_vertexai", llm_params["model"])
|
||||
elif "gpt-" in llm_params["model"]:
|
||||
return handle_model(llm_params["model"], "openai", llm_params["model"])
|
||||
elif llm_params["model"].startswith("claude"):
|
||||
model_name = llm_params["model"].split("/")[-1]
|
||||
return handle_model(model_name, "anthropic", model_name)
|
||||
|
||||
elif "ollama" in llm_params["model"]:
|
||||
model_name = llm_params["model"].split("ollama/")[-1]
|
||||
token_key = model_name if "model_tokens" not in llm_params else llm_params["model_tokens"]
|
||||
return handle_model(model_name, "ollama", token_key)
|
||||
elif llm_params["model"].startswith("vertexai"):
|
||||
return handle_model(llm_params["model"], "google_vertexai", llm_params["model"])
|
||||
|
||||
elif "claude-3-" in llm_params["model"]:
|
||||
return handle_model(llm_params["model"], "anthropic", "claude3")
|
||||
elif "gpt-" in llm_params["model"]:
|
||||
return handle_model(llm_params["model"], "openai", llm_params["model"])
|
||||
|
||||
elif llm_params["model"].startswith("mistral"):
|
||||
model_name = llm_params["model"].split("/")[-1]
|
||||
return handle_model(model_name, "mistralai", model_name)
|
||||
elif "ollama" in llm_params["model"]:
|
||||
model_name = llm_params["model"].split("ollama/")[-1]
|
||||
token_key = model_name if "model_tokens" not in llm_params else llm_params["model_tokens"]
|
||||
return handle_model(model_name, "ollama", token_key)
|
||||
|
||||
# Instantiate the language model based on the model name (models that do not use the common interface)
|
||||
elif "deepseek" in llm_params["model"]:
|
||||
try:
|
||||
self.model_token = models_tokens["deepseek"][llm_params["model"]]
|
||||
except KeyError:
|
||||
print("model not found, using default token size (8192)")
|
||||
self.model_token = 8192
|
||||
return DeepSeek(llm_params)
|
||||
elif "claude-3-" in llm_params["model"]:
|
||||
return handle_model(llm_params["model"], "anthropic", "claude3")
|
||||
|
||||
elif "ernie" in llm_params["model"]:
|
||||
try:
|
||||
self.model_token = models_tokens["ernie"][llm_params["model"]]
|
||||
except KeyError:
|
||||
print("model not found, using default token size (8192)")
|
||||
self.model_token = 8192
|
||||
return ErnieBotChat(llm_params)
|
||||
elif llm_params["model"].startswith("mistral"):
|
||||
model_name = llm_params["model"].split("/")[-1]
|
||||
return handle_model(model_name, "mistralai", model_name)
|
||||
|
||||
elif "oneapi" in llm_params["model"]:
|
||||
# Instantiate the language model based on the model name (models that do not use the common interface)
|
||||
elif "deepseek" in llm_params["model"]:
|
||||
try:
|
||||
self.model_token = models_tokens["deepseek"][llm_params["model"]]
|
||||
except KeyError:
|
||||
print("model not found, using default token size (8192)")
|
||||
self.model_token = 8192
|
||||
return DeepSeek(llm_params)
|
||||
|
||||
# take the model after the last dash
|
||||
llm_params["model"] = llm_params["model"].split("/")[-1]
|
||||
try:
|
||||
self.model_token = models_tokens["oneapi"][llm_params["model"]]
|
||||
except KeyError as exc:
|
||||
raise KeyError("Model not supported") from exc
|
||||
return OneApi(llm_params)
|
||||
elif "ernie" in llm_params["model"]:
|
||||
try:
|
||||
self.model_token = models_tokens["ernie"][llm_params["model"]]
|
||||
except KeyError:
|
||||
print("model not found, using default token size (8192)")
|
||||
self.model_token = 8192
|
||||
return ErnieBotChat(llm_params)
|
||||
|
||||
elif "nvidia" in llm_params["model"]:
|
||||
elif "oneapi" in llm_params["model"]:
|
||||
# take the model after the last dash
|
||||
llm_params["model"] = llm_params["model"].split("/")[-1]
|
||||
try:
|
||||
self.model_token = models_tokens["oneapi"][llm_params["model"]]
|
||||
except KeyError:
|
||||
raise KeyError("Model not supported")
|
||||
return OneApi(llm_params)
|
||||
|
||||
try:
|
||||
self.model_token = models_tokens["nvidia"][llm_params["model"].split("/")[-1]]
|
||||
llm_params["model"] = "/".join(llm_params["model"].split("/")[1:])
|
||||
except KeyError as exc:
|
||||
raise KeyError("Model not supported") from exc
|
||||
return ChatNVIDIA(llm_params)
|
||||
else:
|
||||
model_name = llm_params["model"].split("/")[-1]
|
||||
return handle_model(model_name, llm_params["model"], model_name)
|
||||
elif "nvidia" in llm_params["model"]:
|
||||
try:
|
||||
self.model_token = models_tokens["nvidia"][llm_params["model"].split("/")[-1]]
|
||||
llm_params["model"] = "/".join(llm_params["model"].split("/")[1:])
|
||||
except KeyError:
|
||||
raise KeyError("Model not supported")
|
||||
return ChatNVIDIA(llm_params)
|
||||
|
||||
raise ValueError("Model provided by the configuration not supported")
|
||||
else:
|
||||
model_name = llm_params["model"].split("/")[-1]
|
||||
return handle_model(model_name, llm_params["model"], model_name)
|
||||
|
||||
except KeyError as e:
|
||||
print(f"Model not supported: {e}")
|
||||
|
||||
def get_state(self, key=None) -> dict:
|
||||
""" ""
|
||||
@ -264,4 +270,4 @@ class AbstractGraph(ABC):
|
||||
def run(self) -> str:
|
||||
"""
|
||||
Abstract method to execute the graph and return the result.
|
||||
"""
|
||||
"""
|
||||
Loading…
Reference in New Issue
Block a user