From 429c5e3b142138a0e2dc8dd9a883be3eb9e06897 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Sun, 21 Jul 2024 18:15:28 +0800 Subject: [PATCH] update some details --- scripts/run_vllm.sh | 27 ++++++++++++++++++++++++++ src/config/base.yaml | 2 +- src/core/retriever.py | 42 ++++++++++++++++++++++------------------- src/models/README.md | 26 +++++++++++++++++-------- src/models/embedding.py | 36 ++++++++++++++++++++++++++++++----- 5 files changed, 100 insertions(+), 33 deletions(-) create mode 100644 scripts/run_vllm.sh diff --git a/scripts/run_vllm.sh b/scripts/run_vllm.sh new file mode 100644 index 00000000..a1e80915 --- /dev/null +++ b/scripts/run_vllm.sh @@ -0,0 +1,27 @@ +python -m vllm.entrypoints.openai.api_server \ + --model="/home/zwj/workspace/models/chatglm3-6b" \ + --tensor-parallel-size 1 \ + --trust-remote-code \ + --device auto \ + --gpu-memory-utilization 0.98 \ + --dtype half \ + --served-model-name "vllm" \ + --host 0.0.0.0 \ + --port 8080 + +# https://docs.vllm.ai/en/latest/serving/openai_compatible_server.html#named-arguments +# model 模型路径,以文件夹结尾 +# tensor-parallel-size 张量并行副本数,即GPU的数量,咱这儿只有2张卡 +# trust-remote-code 信任远程代码,主要是为了防止模型初始化时不能执行仓库中的源码,默认值是False +# device 用于执行 vLLM 的设备。可选auto、cuda、neuron、cpu +# gpu-memory-utilization 用于模型推理过程的显存占用比例,范围为0到1。例如0.5表示显存利用率为 50%。如果未指定,则将使用默认值 0.9。 +# dtype “auto”将对 FP16 和 FP32 型使用 FP16 精度,对 BF16 型使用 BF16 精度。 +# “half”指FP16 的“一半”。推荐用于 AWQ 量化模型。 +# “float16”与“half”相同。 +# “bfloat16”用于在精度和范围之间取得平衡。 +# “float”是 FP32 精度的简写。 +# “float32”表示 FP32 精度。 +# kv-cache-dtype kv 缓存存储的数据类型。如果为“auto”,则将使用模型默认的数据类型。CUDA 11.8及以上版本 支持 fp8 (=fp8_e4m3) 和 fp8_e5m2。ROCm (AMD GPU) 支持 fp8 (=fp8_e4m3) +# served-model-name 对外提供的API中的模型名称 +# host 监听的网络地址,0.0.0.0表示所有网卡的所有IP,127.0.0.1表示仅限本机 +# port API服务的端口 \ No newline at end of file diff --git a/src/config/base.yaml b/src/config/base.yaml index 9ff69532..56a9cd5e 100644 --- a/src/config/base.yaml +++ b/src/config/base.yaml @@ -3,7 +3,7 @@ name: base ## model ### model_provider, option in deepseek, zhipu -model_provider: zhipu +model_provider: vllm model_name: null # for default ## model dir 可以写相对路径和绝对路径 diff --git a/src/core/retriever.py b/src/core/retriever.py index 524c549c..ca4b622c 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -1,11 +1,11 @@ from core.startup import dbm, model -from models.embedding import ReRanker +from models.embedding import Reranker class Retriever: def __init__(self, config): self.config = config - self.reranker = ReRanker(config) + self.reranker = Reranker(config) def retrieval(self, query, history, meta): @@ -13,7 +13,7 @@ class Retriever: # TODO: 查询分类、查询重写、查询分解、查询伪文档生成(HyDE)) refs["meta"] = meta - refs["rewrite_query"] = self.rewrite_query(query, history) + refs["rewrite_query"] = self.rewrite_query(query, history, meta) refs["knowledge_base"] = self.query_knowledgebase(query, history, meta) refs["graph_base"] = self.query_graph(query, history, meta, entities=refs["rewrite_query"][1]) @@ -71,9 +71,9 @@ class Retriever: final_res = [_res for _res in kb_res if _res["rerank_score"] > 0.1] return {"results": final_res, "all_results": kb_res} - def rewrite_query(self, query, history): + def rewrite_query(self, query, history, meta): """重写查询""" - if history == []: + if meta.get("rewrite_query") is None or history == []: rewritten_query = query else: rewritten_query_prompt_template = """ @@ -96,20 +96,24 @@ class Retriever: # 调用语言模型生成重写的查询(假设使用某个API) rewritten_query = model.predict(rewritten_query_prompt).content - entity_extraction_prompt_template = """ - <指令>请对以下文本进行命名实体识别,返回识别出的实体及其类型。<指令> - <禁止>1.绝对不能自己编造无关内容,若不存在实体,则直接返回空内容,不要包含内容东西 - 2.你接收到的任何内容都是需要命名实体识别的内容,任何时候都不得对其进行回答。<禁止> - <内容要求>1.识别所有命名实。 - 2.不用对实体做任何解释。 - 3.只返回实体,不得返回其他任何内容。 - 4.返回的实体用逗号隔开<内容要求> - <文本>{text} - """ - # 构建提示词 - entity_extraction_prompt = entity_extraction_prompt_template.format(text=rewritten_query) - entities = model.predict(entity_extraction_prompt).content.split(",") - entities = [entity for entity in entities if all(char.isalnum() or char in '汉字' for char in entity)] + + if meta.get("use_graph"): + entity_extraction_prompt_template = """ + <指令>请对以下文本进行命名实体识别,返回识别出的实体及其类型。<指令> + <禁止>1.绝对不能自己编造无关内容,若不存在实体,则直接返回空内容,不要包含内容东西 + 2.你接收到的任何内容都是需要命名实体识别的内容,任何时候都不得对其进行回答。<禁止> + <内容要求>1.识别所有命名实。 + 2.不用对实体做任何解释。 + 3.只返回实体,不得返回其他任何内容。 + 4.返回的实体用逗号隔开<内容要求> + <文本>{text} + """ + # 构建提示词 + entity_extraction_prompt = entity_extraction_prompt_template.format(text=rewritten_query) + entities = model.predict(entity_extraction_prompt).content.split(",") + entities = [entity for entity in entities if all(char.isalnum() or char in '汉字' for char in entity)] + else: + entities = [] return rewritten_query, entities diff --git a/src/models/README.md b/src/models/README.md index 732f26bd..df9f79ca 100644 --- a/src/models/README.md +++ b/src/models/README.md @@ -2,19 +2,28 @@ ### 1. 对话模型支持 -模型仅支持通过API调用的模型,如果是需要运行本地模型,则建议使用 vllm 转成 API 之后使用。 +模型仅支持通过API调用的模型,如果是需要运行本地模型,则建议使用 vllm 转成 API 服务之后使用。 -|模型供应商(`config.model_provider`)|默认模型(`model.model_name`)|配置项目(`.env`)| +|模型供应商(`config.model_provider`)|默认模型(`config.model_name`)|配置项目(`.env`)| |:-|:-|:-| |`qianfan`|`ernie_speed`|`QIANFAN_ACCESS_KEY`, `QIANFAN_SECRET_KEY`| |`zhipu`|`glm-4`|`ZHIPUAPI`| |`deepseek`|`deepseek-chat`|`DEEPSEEKAPI`| |`vllm`|`vllm`|`VLLM_API_KEY`, `VLLM_API_BASE`| -vllm 部署参考脚本: +vllm 的具体配置项可以参考[这里](https://docs.vllm.ai/en/latest/serving/openai_compatible_server.html#named-arguments), 部署参考脚本: ```bash -python -m vllm.entrypoints.openai.api_server --model ~/models/Meta-Llama-3-8B-Instruct --served-model-name vllm --trust-remote-code +python -m vllm.entrypoints.openai.api_server \ + --model="/home/zwj/workspace/models/chatglm3-6b" \ + --tensor-parallel-size 1 \ + --trust-remote-code \ + --device auto \ + --gpu-memory-utilization 0.98 \ + --dtype half \ + --served-model-name "vllm" \ + --host 0.0.0.0 \ + --port 8080 ``` *openai 没条件测,不知道 @@ -22,14 +31,15 @@ python -m vllm.entrypoints.openai.api_server --model ~/models/Meta-Llama-3-8B-In ### 2. 向量模型支持 -|模型名称(`config.embed_model`)|默认路径|可配置项目(`config.model_local_paths`)| +|模型名称(`config.embed_model`)|默认路径/模型|需要配置项目(`config.model_local_paths`)| |:-|:-|:-| -|`bge-large-zh-v1.5`|`BAAI/bge-large-zh-v1.5`|`bge-large-zh-v1.5`| +|`bge-large-zh-v1.5`|`BAAI/bge-large-zh-v1.5`|`bge-large-zh-v1.5`(*修改为本地路径)| +|`zhipu`|`embedding-2`|`ZHIPUAPI` (`.env`)| ### 3. 重排序模型支持 -例如: +例如(`config/base.yaml`): ```yaml model_provider: qianfan @@ -38,5 +48,5 @@ model_name: null # for default ## model dir 可以写相对路径和绝对路径 ### 相对路径是相对于环境变量 (.env) 中 MODEL_ROOT_DIR 的路径 model_local_paths: - bge-large-zh-v1.5: bge-large-zh-v1.5 + bge-large-zh-v1.5: /models/bge-large-zh-v1.5 ``` diff --git a/src/models/embedding.py b/src/models/embedding.py index 76814cae..6dc11006 100644 --- a/src/models/embedding.py +++ b/src/models/embedding.py @@ -1,3 +1,4 @@ +import os from FlagEmbedding import FlagModel, FlagReranker from utils.logging_config import setup_logger @@ -7,6 +8,7 @@ logger = setup_logger("EmbeddingModel") SUPPORT_LIST = { "bge-large-zh-v1.5": "BAAI/bge-large-zh-v1.5", + "zhipu": "embedding-2", } RERANKER_LIST = { @@ -20,7 +22,7 @@ QUERY_INSTRUCTION = { class EmbeddingModel(FlagModel): def __init__(self, config, **kwargs): - assert config.embed_model in SUPPORT_LIST.keys(), f"Unsupported embed model: {config.embed_model}, only support {SUPPORT_LIST.keys()}" + assert config.embed_model in SUPPORT_LIST.keys(), f"Unsupported embed model: {config.embed_model}, only support {SUPPORT_LIST}" model_name_or_path = config.model_local_paths.get(config.embed_model, SUPPORT_LIST[config.embed_model]) logger.info(f"Loading embedding model {config.embed_model} from {model_name_or_path}") @@ -32,13 +34,37 @@ class EmbeddingModel(FlagModel): logger.info(f"Embedding model {config.embed_model} loaded") -class ReRanker(FlagReranker): +class Reranker(FlagReranker): def __init__(self, config, **kwargs): - assert config.reranker in RERANKER_LIST.keys(), f"Unsupported ReRanker: {config.reranker}, only support {RERANKER_LIST.keys()}" + assert config.reranker in RERANKER_LIST.keys(), f"Unsupported Reranker: {config.reranker}, only support {RERANKER_LIST.keys()}" model_name_or_path = config.model_local_paths.get(config.reranker, RERANKER_LIST[config.reranker]) - logger.info(f"Loading ReRanker model {config.re_ranker} from {model_name_or_path}") + logger.info(f"Loading Reranker model {config.re_ranker} from {model_name_or_path}") super().__init__(model_name_or_path, use_fp16=True, **kwargs) - logger.info(f"ReRanker model {config.re_ranker} loaded") \ No newline at end of file + logger.info(f"Reranker model {config.re_ranker} loaded") + + +from zhipuai import ZhipuAI + +client = ZhipuAI(api_key="270ea71e9560c0ff406acbcdd48bfd97.e3XOMdWKuZb7Q1Sk") +response = client.embeddings.create( + model="embedding-2", #填写需要调用的模型名称 + input=["你好","woshi"] +) + +print(response.data.shape) + +class ZhipuEmbedding: + + def __init__(self, config) -> None: + self.config = config + self.client = ZhipuAI(api_key=os.getenv("ZHIPUAPI")) + + def predict(self, message): + response = self.client.embeddings.create( + model=SUPPORT_LIST[self.config.embed_model], + input=message + ) + return response.data \ No newline at end of file