From c105c26b2f5271a86ebd0d70c5fd80132c1fd017 Mon Sep 17 00:00:00 2001 From: Marco Vinciguerra Date: Mon, 12 Aug 2024 18:32:41 +0200 Subject: [PATCH] Update abstract_graph.py --- scrapegraphai/graphs/abstract_graph.py | 122 +++++++++++++------------ 1 file changed, 64 insertions(+), 58 deletions(-) diff --git a/scrapegraphai/graphs/abstract_graph.py b/scrapegraphai/graphs/abstract_graph.py index 7a0c4d04..6d1d4afe 100644 --- a/scrapegraphai/graphs/abstract_graph.py +++ b/scrapegraphai/graphs/abstract_graph.py @@ -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. - """ + """ \ No newline at end of file