feat: 启动时就自动检查下载nltk
This commit is contained in:
@@ -10,6 +10,7 @@ from starlette.staticfiles import StaticFiles
|
|||||||
from app.middlewares.app_middlewares import add_app_middlewares
|
from app.middlewares.app_middlewares import add_app_middlewares
|
||||||
from app.utils.db_utils import check_database_connection
|
from app.utils.db_utils import check_database_connection
|
||||||
from app.utils.milvus_utils import milvus_service
|
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.postgres_checkpointer import check_postgres_connection, close_postgres_connection
|
||||||
from app.utils.redis_utils import redis_utils
|
from app.utils.redis_utils import redis_utils
|
||||||
|
|
||||||
@@ -23,6 +24,7 @@ def create_app():
|
|||||||
asyncio.create_task(check_postgres_connection()),
|
asyncio.create_task(check_postgres_connection()),
|
||||||
asyncio.create_task(milvus_service.check_milvus_connection()),
|
asyncio.create_task(milvus_service.check_milvus_connection()),
|
||||||
asyncio.create_task(redis_utils.check_redis_connection()),
|
asyncio.create_task(redis_utils.check_redis_connection()),
|
||||||
|
asyncio.create_task(check_nltk()),
|
||||||
)
|
)
|
||||||
async_engine = async_results[0]
|
async_engine = async_results[0]
|
||||||
yield
|
yield
|
||||||
|
|||||||
@@ -72,10 +72,6 @@ class MilvusService:
|
|||||||
embed_model=self.embeddings,
|
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
|
from llama_index.core.node_parser import SimpleNodeParser
|
||||||
# parser = HierarchicalNodeParser.from_defaults(chunk_sizes=[2048, 512, 128])
|
# parser = HierarchicalNodeParser.from_defaults(chunk_sizes=[2048, 512, 128])
|
||||||
parser = SimpleNodeParser()
|
parser = SimpleNodeParser()
|
||||||
|
|||||||
+26
-5
@@ -1,14 +1,35 @@
|
|||||||
# 预先加载NLTK语料库以避免多线程环境中的竞争条件
|
# 预先加载NLTK语料库以避免多线程环境中的竞争条件
|
||||||
|
import asyncio
|
||||||
|
|
||||||
import nltk
|
import nltk
|
||||||
|
from nltk.data import find
|
||||||
load_flag = False
|
|
||||||
|
|
||||||
|
|
||||||
|
# 判断nltk是否已经存在
|
||||||
|
def is_nltk_package_downloaded(package_name):
|
||||||
|
try:
|
||||||
|
find(f'tokenizers/{package_name}')
|
||||||
|
return True
|
||||||
|
except LookupError:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
# 下载nltk所需要文件
|
||||||
def load_nltk():
|
def load_nltk():
|
||||||
global load_flag
|
if not is_nltk_package_downloaded("punkt"):
|
||||||
if not load_flag:
|
print("ℹ️ 下载nltk")
|
||||||
nltk.download('punkt')
|
nltk.download('punkt')
|
||||||
nltk.download('averaged_perceptron_tagger')
|
nltk.download('averaged_perceptron_tagger')
|
||||||
nltk.download('maxent_ne_chunker')
|
nltk.download('maxent_ne_chunker')
|
||||||
nltk.download('words')
|
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)
|
||||||
|
|||||||
Reference in New Issue
Block a user