diff --git a/app/create_app.py b/app/create_app.py index 7ed9bae..6972ec1 100644 --- a/app/create_app.py +++ b/app/create_app.py @@ -10,6 +10,7 @@ from starlette.staticfiles import StaticFiles from app.middlewares.app_middlewares import add_app_middlewares from app.utils.db_utils import check_database_connection from app.utils.milvus_utils import milvus_service +from app.utils.nltk_utils import check_nltk from app.utils.postgres_checkpointer import check_postgres_connection, close_postgres_connection from app.utils.redis_utils import redis_utils @@ -23,6 +24,7 @@ def create_app(): asyncio.create_task(check_postgres_connection()), asyncio.create_task(milvus_service.check_milvus_connection()), asyncio.create_task(redis_utils.check_redis_connection()), + asyncio.create_task(check_nltk()), ) async_engine = async_results[0] yield diff --git a/app/utils/milvus_utils.py b/app/utils/milvus_utils.py index af9d8b3..62b4adc 100644 --- a/app/utils/milvus_utils.py +++ b/app/utils/milvus_utils.py @@ -72,10 +72,6 @@ class MilvusService: embed_model=self.embeddings, ) - # 预先加载NLTK语料库以避免多线程环境中的竞争条件 - # 否则在多线程环境下,parser.get_nodes_from_documents 可能会出现报错信息:'WordListCorpusReader' object has no attribute '_LazyCorpusLoader__args' - load_nltk() - from llama_index.core.node_parser import SimpleNodeParser # parser = HierarchicalNodeParser.from_defaults(chunk_sizes=[2048, 512, 128]) parser = SimpleNodeParser() diff --git a/app/utils/nltk_utils.py b/app/utils/nltk_utils.py index 12f4ffd..a22b306 100644 --- a/app/utils/nltk_utils.py +++ b/app/utils/nltk_utils.py @@ -1,14 +1,35 @@ # 预先加载NLTK语料库以避免多线程环境中的竞争条件 +import asyncio + import nltk - -load_flag = False +from nltk.data import find +# 判断nltk是否已经存在 +def is_nltk_package_downloaded(package_name): + try: + find(f'tokenizers/{package_name}') + return True + except LookupError: + return False + + +# 下载nltk所需要文件 def load_nltk(): - global load_flag - if not load_flag: + if not is_nltk_package_downloaded("punkt"): + print("ℹ️ 下载nltk") nltk.download('punkt') nltk.download('averaged_perceptron_tagger') nltk.download('maxent_ne_chunker') nltk.download('words') - load_flag = True + print("✅ nltk下载成功") + else: + print("✅ nltk已经存在,无需下载") + + +# 检查nltk是否已经下载,用于llama-index做文本分割 +async def check_nltk(): + print("nltk_path:", nltk.data.path) + # 预先加载NLTK语料库以避免多线程环境中的竞争条件 + # 否则在多线程环境下,parser.get_nodes_from_documents 可能会出现报错信息:'WordListCorpusReader' object has no attribute '_LazyCorpusLoader__args' + await asyncio.wait_for(asyncio.to_thread(load_nltk), timeout=10)