refactor(models): move model selection logic to chat_model and embedding modules
This commit is contained in:
parent
342c541645
commit
33dbe78914
@ -1,65 +1,4 @@
|
|||||||
import os
|
from src.models.chat_model import select_model, get_custom_model
|
||||||
import traceback
|
from src.models.embedding import select_embedding_model
|
||||||
|
|
||||||
from src import config
|
__all__ = ['select_model', 'select_embedding_model', 'get_custom_model']
|
||||||
from src.models.chat_model import OpenAIBase
|
|
||||||
from src.models.embedding import OllamaEmbedding, OtherEmbedding
|
|
||||||
from src.utils.logging_config import logger
|
|
||||||
|
|
||||||
|
|
||||||
def select_model(model_provider, model_name=None):
|
|
||||||
"""根据模型提供者选择模型"""
|
|
||||||
assert model_provider is not None, "Model provider not specified"
|
|
||||||
model_info = config.model_names.get(model_provider, {})
|
|
||||||
model_name = model_name or model_info.get("default", "")
|
|
||||||
|
|
||||||
logger.info(f"Selecting model from `{model_provider}` with `{model_name}`")
|
|
||||||
|
|
||||||
if model_provider == "openai":
|
|
||||||
from src.models.chat_model import OpenModel
|
|
||||||
|
|
||||||
return OpenModel(model_name)
|
|
||||||
|
|
||||||
if model_provider == "custom":
|
|
||||||
model_info = get_custom_model(model_name)
|
|
||||||
|
|
||||||
from src.models.chat_model import CustomModel
|
|
||||||
|
|
||||||
return CustomModel(model_info)
|
|
||||||
|
|
||||||
# 其他模型,默认使用OpenAIBase
|
|
||||||
try:
|
|
||||||
model = OpenAIBase(
|
|
||||||
api_key=os.getenv(model_info["env"][0]),
|
|
||||||
base_url=model_info["base_url"],
|
|
||||||
model_name=model_name,
|
|
||||||
)
|
|
||||||
return model
|
|
||||||
except Exception as e:
|
|
||||||
raise ValueError(f"Model provider {model_provider} load failed, {e} \n {traceback.format_exc()}")
|
|
||||||
|
|
||||||
|
|
||||||
def select_embedding_model(model_id):
|
|
||||||
provider, model_name = model_id.split("/", 1) if model_id else ("", "")
|
|
||||||
support_embed_models = config.embed_model_names.keys()
|
|
||||||
assert model_id in support_embed_models, f"Unsupported embed model: {model_id}, only support {support_embed_models}"
|
|
||||||
logger.debug(f"Loading embedding model {model_id}")
|
|
||||||
if provider == "local":
|
|
||||||
raise ValueError("Local embedding model is not supported, please use other embedding models")
|
|
||||||
|
|
||||||
elif provider == "ollama":
|
|
||||||
model = OllamaEmbedding(**config.embed_model_names[model_id])
|
|
||||||
|
|
||||||
else:
|
|
||||||
model = OtherEmbedding(**config.embed_model_names[model_id])
|
|
||||||
|
|
||||||
return model
|
|
||||||
|
|
||||||
|
|
||||||
def get_custom_model(model_id):
|
|
||||||
"""return model_info"""
|
|
||||||
assert config.custom_models is not None, "custom_models is not set"
|
|
||||||
modle_info = next((x for x in config.custom_models if x["custom_id"] == model_id), None)
|
|
||||||
if modle_info is None:
|
|
||||||
raise ValueError(f"Model {model_id} not found in custom models")
|
|
||||||
return modle_info
|
|
||||||
|
|||||||
@ -1,7 +1,9 @@
|
|||||||
import os
|
import os
|
||||||
|
import traceback
|
||||||
|
|
||||||
from openai import OpenAI
|
from openai import OpenAI
|
||||||
|
|
||||||
|
from src import config
|
||||||
from src.utils import get_docker_safe_url, logger
|
from src.utils import get_docker_safe_url, logger
|
||||||
|
|
||||||
|
|
||||||
@ -83,5 +85,41 @@ class GeneralResponse:
|
|||||||
self.is_full = False
|
self.is_full = False
|
||||||
|
|
||||||
|
|
||||||
|
def select_model(model_provider, model_name=None):
|
||||||
|
"""根据模型提供者选择模型"""
|
||||||
|
assert model_provider is not None, "Model provider not specified"
|
||||||
|
model_info = config.model_names.get(model_provider, {})
|
||||||
|
model_name = model_name or model_info.get("default", "")
|
||||||
|
|
||||||
|
logger.info(f"Selecting model from `{model_provider}` with `{model_name}`")
|
||||||
|
|
||||||
|
if model_provider == "openai":
|
||||||
|
return OpenModel(model_name)
|
||||||
|
|
||||||
|
if model_provider == "custom":
|
||||||
|
model_info = get_custom_model(model_name)
|
||||||
|
return CustomModel(model_info)
|
||||||
|
|
||||||
|
# 其他模型,默认使用OpenAIBase
|
||||||
|
try:
|
||||||
|
model = OpenAIBase(
|
||||||
|
api_key=os.getenv(model_info["env"][0]),
|
||||||
|
base_url=model_info["base_url"],
|
||||||
|
model_name=model_name,
|
||||||
|
)
|
||||||
|
return model
|
||||||
|
except Exception as e:
|
||||||
|
raise ValueError(f"Model provider {model_provider} load failed, {e} \n {traceback.format_exc()}")
|
||||||
|
|
||||||
|
|
||||||
|
def get_custom_model(model_id):
|
||||||
|
"""return model_info"""
|
||||||
|
assert config.custom_models is not None, "custom_models is not set"
|
||||||
|
modle_info = next((x for x in config.custom_models if x["custom_id"] == model_id), None)
|
||||||
|
if modle_info is None:
|
||||||
|
raise ValueError(f"Model {model_id} not found in custom models")
|
||||||
|
return modle_info
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
pass
|
pass
|
||||||
|
|||||||
@ -6,6 +6,7 @@ from abc import ABC, abstractmethod
|
|||||||
import httpx
|
import httpx
|
||||||
import requests
|
import requests
|
||||||
|
|
||||||
|
from src import config
|
||||||
from src.utils import get_docker_safe_url, hashstr, logger
|
from src.utils import get_docker_safe_url, hashstr, logger
|
||||||
|
|
||||||
|
|
||||||
@ -174,3 +175,20 @@ class OtherEmbedding(BaseEmbeddingModel):
|
|||||||
except (httpx.RequestError, json.JSONDecodeError) as e:
|
except (httpx.RequestError, json.JSONDecodeError) as e:
|
||||||
logger.error(f"Other Embedding async request failed: {e}, {payload}")
|
logger.error(f"Other Embedding async request failed: {e}, {payload}")
|
||||||
raise ValueError(f"Other Embedding async request failed: {e}")
|
raise ValueError(f"Other Embedding async request failed: {e}")
|
||||||
|
|
||||||
|
|
||||||
|
def select_embedding_model(model_id):
|
||||||
|
provider, model_name = model_id.split("/", 1) if model_id else ("", "")
|
||||||
|
support_embed_models = config.embed_model_names.keys()
|
||||||
|
assert model_id in support_embed_models, f"Unsupported embed model: {model_id}, only support {support_embed_models}"
|
||||||
|
logger.debug(f"Loading embedding model {model_id}")
|
||||||
|
if provider == "local":
|
||||||
|
raise ValueError("Local embedding model is not supported, please use other embedding models")
|
||||||
|
|
||||||
|
elif provider == "ollama":
|
||||||
|
model = OllamaEmbedding(**config.embed_model_names[model_id])
|
||||||
|
|
||||||
|
else:
|
||||||
|
model = OtherEmbedding(**config.embed_model_names[model_id])
|
||||||
|
|
||||||
|
return model
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user