(engine_name: str, **kwargs)
| 29 | f"The engine provided is not multimodal. Please provide a multimodal engine, one of the following: {__MULTIMODAL_ENGINES__}") |
| 30 | |
| 31 | def get_engine(engine_name: str, **kwargs) -> EngineLM: |
| 32 | if engine_name in __ENGINE_NAME_SHORTCUTS__: |
| 33 | engine_name = __ENGINE_NAME_SHORTCUTS__[engine_name] |
| 34 | |
| 35 | if "seed" in kwargs and "gpt-4" not in engine_name and "gpt-3.5" not in engine_name and "gpt-35" not in engine_name: |
| 36 | raise ValueError(f"Seed is currently supported only for OpenAI engines, not {engine_name}") |
| 37 | |
| 38 | if "cache" in kwargs and "experimental" not in engine_name: |
| 39 | raise ValueError(f"Cache is currently supported only for LiteLLM engines, not {engine_name}") |
| 40 | |
| 41 | # check if engine_name starts with "experimental:" |
| 42 | if engine_name.startswith("experimental:"): |
| 43 | engine_name = engine_name.split("experimental:")[1] |
| 44 | return LiteLLMEngine(model_string=engine_name, **kwargs) |
| 45 | if engine_name.startswith("azure"): |
| 46 | from .openai import AzureChatOpenAI |
| 47 | # remove engine_name "azure-" prefix |
| 48 | engine_name = engine_name[6:] |
| 49 | return AzureChatOpenAI(model_string=engine_name, **kwargs) |
| 50 | elif (("gpt-4" in engine_name) or ("gpt-3.5" in engine_name)): |
| 51 | from .openai import ChatOpenAI |
| 52 | return ChatOpenAI(model_string=engine_name, is_multimodal=_check_if_multimodal(engine_name), **kwargs) |
| 53 | elif "claude" in engine_name: |
| 54 | from .anthropic import ChatAnthropic |
| 55 | return ChatAnthropic(model_string=engine_name, is_multimodal=_check_if_multimodal(engine_name), **kwargs) |
| 56 | elif "gemini" in engine_name: |
| 57 | from .gemini import ChatGemini |
| 58 | return ChatGemini(model_string=engine_name, **kwargs) |
| 59 | elif "together" in engine_name: |
| 60 | from .together import ChatTogether |
| 61 | engine_name = engine_name.replace("together-", "") |
| 62 | return ChatTogether(model_string=engine_name, **kwargs) |
| 63 | elif engine_name in ["command-r-plus", "command-r", "command", "command-light"]: |
| 64 | from .cohere import ChatCohere |
| 65 | return ChatCohere(model_string=engine_name, **kwargs) |
| 66 | elif engine_name.startswith("ollama"): |
| 67 | from .openai import ChatOpenAI, OLLAMA_BASE_URL |
| 68 | model_string = engine_name.replace("ollama-", "") |
| 69 | return ChatOpenAI( |
| 70 | model_string=model_string, |
| 71 | base_url=OLLAMA_BASE_URL, |
| 72 | **kwargs |
| 73 | ) |
| 74 | elif "vllm" in engine_name: |
| 75 | from .vllm import ChatVLLM |
| 76 | engine_name = engine_name.replace("vllm-", "") |
| 77 | return ChatVLLM(model_string=engine_name, **kwargs) |
| 78 | elif "groq" in engine_name: |
| 79 | from .groq import ChatGroq |
| 80 | engine_name = engine_name.replace("groq-", "") |
| 81 | return ChatGroq(model_string=engine_name, **kwargs) |
| 82 | else: |
| 83 | raise ValueError(f"Engine {engine_name} not supported") |
no test coverage detected