From bd679f47d637702a6e38735280a4cce20e1b5277 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Sat, 23 Nov 2024 22:25:46 +0800 Subject: [PATCH] =?UTF-8?q?=E6=9C=AC=E5=9C=B0=E6=A8=A1=E5=9E=8B=E5=8F=82?= =?UTF-8?q?=E6=95=B0=E4=BF=AE=E6=94=B9=EF=BC=88=E6=9C=AA=E6=B5=8B=E8=AF=95?= =?UTF-8?q?=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/models/embedding.py | 14 +- src/plugins/oneke.py | 4 +- web/src/components/TableConfigComponent.vue | 202 ++++++++++++++++++++ web/src/views/SettingView.vue | 14 +- 4 files changed, 217 insertions(+), 17 deletions(-) create mode 100644 web/src/components/TableConfigComponent.vue diff --git a/src/models/embedding.py b/src/models/embedding.py index 7b576abe..604d1c1d 100644 --- a/src/models/embedding.py +++ b/src/models/embedding.py @@ -15,11 +15,7 @@ GLOBAL_EMBED_STATE = {} class EmbeddingModel(FlagModel): def __init__(self, model_info, config, **kwargs): self.info = model_info - model_name_or_path = handle_local_model( - paths=config.model_local_paths, - model_name=model_info["name"], - default_path=model_info.get("default_path", None)) - + model_name_or_path = config.model_local_paths.get(model_info["name"], model_info.get("default_path", None)) logger.info(f"Loading embedding model {model_info['name']} from {model_name_or_path}") super().__init__(model_name_or_path, @@ -34,11 +30,7 @@ class Reranker(FlagReranker): assert config.reranker in RERANKER_LIST.keys(), f"Unsupported Reranker: {config.reranker}, only support {RERANKER_LIST.keys()}" - model_name_or_path = handle_local_model( - paths=config.model_local_paths, - model_name=config.reranker, - default_path=RERANKER_LIST[config.reranker]) - + model_name_or_path = config.model_local_paths.get(config.reranker, default_path=RERANKER_LIST[config.reranker]) logger.info(f"Loading Reranker model {config.reranker} from {model_name_or_path}") super().__init__(model_name_or_path, use_fp16=True, **kwargs) @@ -113,6 +105,4 @@ def get_embedding_model(config): def handle_local_model(paths, model_name, default_path): model_path = paths.get(model_name, default_path) - if os.getenv("MODEL_ROOT_DIR") and not os.path.isabs(model_path): - model_path = os.path.join(os.getenv("MODEL_ROOT_DIR"), model_path) return model_path \ No newline at end of file diff --git a/src/plugins/oneke.py b/src/plugins/oneke.py index aea4c5c8..3ca3d60b 100644 --- a/src/plugins/oneke.py +++ b/src/plugins/oneke.py @@ -16,8 +16,6 @@ logger = setup_logger("OneKE") dotenv.load_dotenv() -MODEL_NAME_OR_PATH = os.path.join(os.getenv('MODEL_ROOT_DIR', './'), 'OneKE') - instruction_mapper = { 'NERzh': "你是专门进行实体抽取的专家。请从input中抽取出符合schema定义的实体,不存在的实体类型返回空列表。请按照JSON字符串的格式回答。", 'REzh': "你是专门进行关系抽取的专家。请从input中抽取出符合schema定义的关系三元组。请按照JSON字符串的格式回答。", @@ -42,7 +40,7 @@ class OneKE: def __init__(self, config=None): self.config = config - model_name_or_path = config.model_local_paths.get('oneke', "zjunlp/OneKE") + model_name_or_path = config.model_local_paths.get('zjunlp/OneKE', "zjunlp/OneKE") logger.info(f"Loading KGC model OneKE from {model_name_or_path}") model_config = AutoConfig.from_pretrained(model_name_or_path, trust_remote_code=True) diff --git a/web/src/components/TableConfigComponent.vue b/web/src/components/TableConfigComponent.vue new file mode 100644 index 00000000..7b372d4c --- /dev/null +++ b/web/src/components/TableConfigComponent.vue @@ -0,0 +1,202 @@ + + + + + + diff --git a/web/src/views/SettingView.vue b/web/src/views/SettingView.vue index ca0fe1ab..a314595a 100644 --- a/web/src/views/SettingView.vue +++ b/web/src/views/SettingView.vue @@ -185,7 +185,11 @@
-

暂无配置

+

本地模型配置

+
@@ -206,6 +210,7 @@ import { InfoCircleOutlined, } from '@ant-design/icons-vue'; import HeaderComponent from '@/components/HeaderComponent.vue'; +import TableConfigComponent from '@/components/TableConfigComponent.vue'; import { notification, Button } from 'ant-design-vue'; const configStore = useConfigStore() @@ -246,6 +251,10 @@ const generateRandomHash = (length) => { return hash; } +const handleModelLocalPathsUpdate = (config) => { + handleChange('model_local_paths', config) +} + const handleChange = (key, e) => { if (key == 'enable_knowledge_graph' && e && !configStore.config.enable_knowledge_base) { message.error('启动知识图谱必须请先启用知识库功能') @@ -264,7 +273,8 @@ const handleChange = (key, e) => { || key == 'model_provider' || key == 'model_name' || key == 'embed_model' - || key == 'reranker') { + || key == 'reranker' + || key == 'model_local_paths') { if (!isNeedRestart.value) { isNeedRestart.value = true notification.info({