修复 model_provider_lite 的 bug

This commit is contained in:
Wenjie Zhang 2025-05-06 22:48:16 +08:00
parent a20143b09c
commit 0dd1d61204
3 changed files with 4 additions and 21 deletions

View File

@ -129,21 +129,6 @@ async def call(query: str = Body(...), meta: dict = Body(None)):
return {"response": response.content}
@chat.post("/call_lite")
async def call_lite(query: str = Body(...), meta: dict = Body(None)):
meta = meta or {}
async def predict_async(query):
loop = asyncio.get_event_loop()
model_provider = meta.get("model_provider", config.model_provider_lite)
model_name = meta.get("model_name", config.model_name_lite)
model = select_model(model_provider=model_provider, model_name=model_name)
return await loop.run_in_executor(executor, model.predict, query)
response = await predict_async(query)
logger.debug({"query": query, "response": response.content})
return {"response": response.content}
@chat.get("/agent")
async def get_agent():
agents = [agent.get_info() for agent in agent_manager.agents.values()]

View File

@ -51,9 +51,7 @@ class Config(SimpleConfig):
## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径
## 如果需要自定义本地模型路径,则在 src/.env 中配置 MODEL_DIR
self.add_item("model_provider", default="siliconflow", des="模型提供商", choices=list(self.model_names.keys()))
self.add_item("model_provider_lite", default="siliconflow", des="模型提供商(用于轻量任务)", choices=list(self.model_names.keys()))
self.add_item("model_name", default="Qwen/Qwen2.5-7B-Instruct", des="模型名称")
self.add_item("model_name_lite", default="Qwen/Qwen2.5-7B-Instruct", des="模型名称(用于轻量任务)")
self.add_item("embed_model", default="siliconflow/BAAI/bge-m3", des="Embedding 模型", choices=list(self.embed_model_names.keys()))
self.add_item("reranker", default="siliconflow/BAAI/bge-reranker-v2-m3", des="Re-Ranker 模型", choices=list(self.reranker_names.keys()))

View File

@ -136,8 +136,8 @@ class Retriever:
def rewrite_query(self, query, history, refs):
"""重写查询"""
model_provider = config.model_provider_lite
model_name = config.model_name_lite
model_provider = config.model_provider
model_name = config.model_name
model = select_model(model_provider=model_provider, model_name=model_name)
if refs["meta"].get("mode") == "search": # 如果是搜索模式,就使用 meta 的配置,否则就使用全局的配置
rewrite_query_span = refs["meta"].get("use_rewrite_query", "off")
@ -162,8 +162,8 @@ class Retriever:
def reco_entities(self, query, history, refs):
"""识别句子中的实体"""
query = refs.get("rewritten_query", query)
model_provider = config.model_provider_lite
model_name = config.model_name_lite
model_provider = config.model_provider
model_name = config.model_name
model = select_model(model_provider=model_provider, model_name=model_name)
entities = []