Skip to content

Commit 714fb84

Browse files
committed
fix:修复模型不能下载问题
1 parent a7dccc4 commit 714fb84

3 files changed

Lines changed: 316 additions & 3 deletions

File tree

scripts/check_models.py

Lines changed: 20 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -83,9 +83,6 @@ async def download_models():
8383
# 显示友好的下载前提醒
8484
show_download_warning()
8585

86-
# 设置环境变量确保显示进度
87-
os.environ['HF_HUB_DISABLE_PROGRESS_BARS'] = 'false'
88-
8986
# 尝试导入模型管理器
9087
try:
9188
from any2markdown_mcp.models.model_manager import ModelManager
@@ -96,6 +93,26 @@ async def download_models():
9693
print_error(" pip install -r requirements.txt")
9794
return False
9895

96+
# 🔧 在创建模型管理器之前,先确保所有环境变量都正确设置
97+
print_progress("🔧 设置模型缓存环境变量...")
98+
99+
# 设置模型缓存相关的环境变量
100+
env_mappings = {
101+
'MODEL_CACHE_DIR': settings.model_cache_dir,
102+
'HF_HOME': settings.hf_home,
103+
'HF_HUB_CACHE': settings.hf_hub_cache,
104+
'HF_ASSETS_CACHE': settings.hf_assets_cache,
105+
'TORCH_HOME': settings.torch_home,
106+
'TRANSFORMERS_CACHE': settings.transformers_cache,
107+
'HF_HUB_ENABLE_HF_TRANSFER': str(settings.hf_hub_enable_hf_transfer).lower(),
108+
'HF_HUB_DISABLE_PROGRESS_BARS': 'false', # 确保显示进度
109+
'HF_HUB_DISABLE_TELEMETRY': str(settings.hf_hub_disable_telemetry).lower(),
110+
}
111+
112+
for env_var, value in env_mappings.items():
113+
os.environ[env_var] = value
114+
print_info(f" ✅ 设置 {env_var} = {value}")
115+
99116
print_progress("📦 创建模型管理器实例...")
100117

101118
# 创建模型管理器实例
Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,7 @@
1+
"""
2+
Models module for Any2Markdown MCP Server
3+
"""
4+
5+
from .model_manager import ModelManager
6+
7+
__all__ = ['ModelManager']
Lines changed: 289 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,289 @@
1+
"""
2+
模型管理器模块
3+
4+
负责加载和管理marker-pdf模型以及其他必要的机器学习模型
5+
支持GPU/CPU自动检测和模型缓存
6+
"""
7+
8+
import asyncio
9+
import os
10+
import sys
11+
from pathlib import Path
12+
from typing import Optional, Dict, Any
13+
import warnings
14+
import time
15+
16+
import structlog
17+
import torch
18+
from huggingface_hub import snapshot_download
19+
20+
# 抑制一些不必要的警告
21+
warnings.filterwarnings("ignore", category=FutureWarning)
22+
warnings.filterwarnings("ignore", category=UserWarning)
23+
24+
logger = structlog.get_logger(__name__)
25+
26+
27+
class ModelManager:
28+
"""模型管理器,负责加载和管理所有机器学习模型"""
29+
30+
def __init__(self, config: Dict[str, Any]):
31+
self.config = config
32+
self.device = self._detect_device()
33+
self.models: Dict[str, Any] = {}
34+
self.model_cache_dir = Path(config.get("model_cache_dir", "~/.cache/marker")).expanduser()
35+
self.model_cache_dir.mkdir(parents=True, exist_ok=True)
36+
37+
# 初始化标志
38+
self._initialized = False
39+
self._initialization_lock = asyncio.Lock()
40+
41+
# 设置模型下载相关的环境变量(确保进度显示)
42+
self._setup_model_env_vars()
43+
44+
logger.info("ModelManager initialized",
45+
device=self.device,
46+
cache_dir=str(self.model_cache_dir),
47+
progress_bars_enabled=not self._is_progress_disabled())
48+
49+
def _setup_model_env_vars(self):
50+
"""设置模型下载相关的环境变量"""
51+
# 确保显示下载进度条
52+
if self.config.get("hf_hub_disable_progress_bars", False):
53+
logger.warning("⚠️ Progress bars are disabled in config. Model download progress will not be visible!")
54+
os.environ["HF_HUB_DISABLE_PROGRESS_BARS"] = "true"
55+
else:
56+
logger.info("✅ Model download progress bars are enabled")
57+
os.environ["HF_HUB_DISABLE_PROGRESS_BARS"] = "false"
58+
59+
# 设置其他环境变量
60+
env_vars = {
61+
"MODEL_CACHE_DIR": self.config.get("model_cache_dir", "~/.cache/marker"),
62+
"HF_HOME": self.config.get("hf_home", "~/.cache/huggingface"),
63+
"HF_HUB_CACHE": self.config.get("hf_hub_cache", "~/.cache/huggingface/hub"),
64+
"HF_ASSETS_CACHE": self.config.get("hf_assets_cache", "~/.cache/huggingface/assets"),
65+
"TORCH_HOME": self.config.get("torch_home", "~/.cache/torch"),
66+
"TRANSFORMERS_CACHE": self.config.get("transformers_cache", "~/.cache/transformers"),
67+
"HF_HUB_ENABLE_HF_TRANSFER": str(self.config.get("hf_hub_enable_hf_transfer", False)).lower(),
68+
"HF_HUB_DISABLE_TELEMETRY": str(self.config.get("hf_hub_disable_telemetry", True)).lower(),
69+
}
70+
71+
for key, value in env_vars.items():
72+
# 展开用户目录路径
73+
if value.startswith("~"):
74+
expanded_value = str(Path(value).expanduser())
75+
else:
76+
expanded_value = value
77+
78+
# 设置环境变量(即使已存在也要更新,确保使用配置中的值)
79+
os.environ[key] = expanded_value
80+
logger.debug("Set environment variable", var=key, value=expanded_value)
81+
82+
# 确保缓存目录存在
83+
for key in ["HF_HOME", "HF_HUB_CACHE", "HF_ASSETS_CACHE", "TORCH_HOME", "TRANSFORMERS_CACHE", "MODEL_CACHE_DIR"]:
84+
cache_dir = Path(os.environ[key])
85+
cache_dir.mkdir(parents=True, exist_ok=True)
86+
logger.debug("Ensured cache directory exists", dir=str(cache_dir))
87+
88+
def _is_progress_disabled(self) -> bool:
89+
"""检查是否禁用了进度条"""
90+
return os.environ.get("HF_HUB_DISABLE_PROGRESS_BARS", "false").lower() == "true"
91+
92+
def _detect_device(self) -> str:
93+
"""自动检测最佳设备"""
94+
device_config = self.config.get("device", "auto").lower()
95+
96+
if device_config == "cpu":
97+
return "cpu"
98+
elif device_config == "cuda" and torch.cuda.is_available():
99+
return "cuda"
100+
elif device_config == "mps" and torch.backends.mps.is_available():
101+
return "mps"
102+
elif device_config == "auto":
103+
# 自动检测最佳设备
104+
if torch.cuda.is_available():
105+
device = "cuda"
106+
gpu_count = torch.cuda.device_count()
107+
gpu_name = torch.cuda.get_device_name(0)
108+
logger.info("CUDA available", gpu_count=gpu_count, gpu_name=gpu_name)
109+
return device
110+
elif torch.backends.mps.is_available():
111+
logger.info("MPS (Apple Silicon) available")
112+
return "mps"
113+
else:
114+
logger.info("Using CPU (no GPU acceleration available)")
115+
return "cpu"
116+
else:
117+
logger.warning("Invalid device config, falling back to CPU",
118+
device_config=device_config)
119+
return "cpu"
120+
121+
async def initialize(self) -> None:
122+
"""异步初始化所有模型"""
123+
if self._initialized:
124+
return
125+
126+
async with self._initialization_lock:
127+
if self._initialized: # 双重检查
128+
return
129+
130+
logger.info("🚀 Starting model initialization...")
131+
logger.info("📥 This may take some time for first run as models need to be downloaded...")
132+
133+
start_time = time.time()
134+
135+
try:
136+
# 初始化marker模型
137+
await self._initialize_marker_model()
138+
139+
# 可以在这里添加其他模型的初始化
140+
# await self._initialize_other_models()
141+
142+
self._initialized = True
143+
144+
elapsed_time = time.time() - start_time
145+
logger.info("✅ All models initialized successfully",
146+
initialization_time=f"{elapsed_time:.2f}s")
147+
148+
except Exception as e:
149+
elapsed_time = time.time() - start_time
150+
logger.error("❌ Failed to initialize models",
151+
error=str(e),
152+
elapsed_time=f"{elapsed_time:.2f}s")
153+
raise
154+
155+
async def _initialize_marker_model(self) -> None:
156+
"""初始化marker-pdf模型"""
157+
try:
158+
logger.info("📦 Loading marker-pdf models...")
159+
logger.info("ℹ️ If this is the first run, marker will download ~3-5GB of AI models")
160+
logger.info("📊 Download progress will be shown below:")
161+
162+
# 导入marker模块(延迟导入以避免启动时的依赖问题)
163+
try:
164+
from marker.scripts.convert import process_single_pdf, create_model_dict
165+
logger.info("✅ Marker modules imported successfully")
166+
except ImportError as e:
167+
logger.error("❌ Failed to import marker modules", error=str(e))
168+
logger.error("💡 Please ensure marker-pdf is installed: pip install marker-pdf")
169+
raise
170+
171+
# 在executor中运行模型加载(避免阻塞事件循环)
172+
loop = asyncio.get_event_loop()
173+
174+
logger.info("🔄 Loading AI models (this may take several minutes)...")
175+
models = await loop.run_in_executor(
176+
None,
177+
self._load_marker_models
178+
)
179+
180+
self.models["marker"] = {
181+
"models": models,
182+
"convert_func": process_single_pdf
183+
}
184+
185+
logger.info("✅ Marker models loaded successfully")
186+
187+
except ImportError as e:
188+
logger.error("❌ Failed to import marker modules", error=str(e))
189+
raise
190+
except Exception as e:
191+
logger.error("❌ Failed to load marker models", error=str(e))
192+
raise
193+
194+
def _load_marker_models(self) -> Any:
195+
"""在线程池中加载marker模型"""
196+
try:
197+
from marker.scripts.convert import create_model_dict
198+
199+
logger.info("🔧 Setting up model cache directories...")
200+
201+
# 设置环境变量
202+
os.environ["TORCH_HOME"] = str(self.model_cache_dir)
203+
204+
logger.info("🤖 Creating marker model dictionary...")
205+
logger.info("📡 Models will be downloaded to:", cache_dir=str(self.model_cache_dir))
206+
207+
# 创建模型字典
208+
device = None if self.device == "auto" else self.device
209+
logger.info("🚀 Initializing models...", target_device=device or "auto")
210+
211+
models = create_model_dict(device=device)
212+
213+
logger.info("✅ Marker models created successfully")
214+
return models
215+
216+
except Exception as e:
217+
logger.error("❌ Error in _load_marker_models", error=str(e))
218+
import traceback
219+
logger.error("🔍 Full traceback:", traceback=traceback.format_exc())
220+
raise
221+
222+
async def get_marker_models(self) -> Dict[str, Any]:
223+
"""获取marker模型"""
224+
await self.initialize()
225+
return self.models["marker"]["models"]
226+
227+
async def convert_pdf_with_marker(self, pdf_path: str, **kwargs) -> tuple:
228+
"""使用marker转换PDF"""
229+
await self.initialize()
230+
231+
try:
232+
convert_func = self.models["marker"]["convert_func"]
233+
models = self.models["marker"]["models"]
234+
235+
# 在executor中运行转换(避免阻塞事件循环)
236+
loop = asyncio.get_event_loop()
237+
result = await loop.run_in_executor(
238+
None,
239+
lambda: convert_func(pdf_path, models, **kwargs)
240+
)
241+
242+
return result
243+
244+
except Exception as e:
245+
logger.error("PDF conversion failed", error=str(e), pdf_path=pdf_path)
246+
raise
247+
248+
def get_device_info(self) -> Dict[str, Any]:
249+
"""获取设备信息"""
250+
info = {
251+
"device": self.device,
252+
"torch_version": torch.__version__,
253+
}
254+
255+
if self.device == "cuda":
256+
info.update({
257+
"cuda_available": torch.cuda.is_available(),
258+
"cuda_version": torch.version.cuda,
259+
"gpu_count": torch.cuda.device_count(),
260+
"gpu_names": [torch.cuda.get_device_name(i) for i in range(torch.cuda.device_count())],
261+
"current_device": torch.cuda.current_device(),
262+
})
263+
elif self.device == "mps":
264+
info.update({
265+
"mps_available": torch.backends.mps.is_available(),
266+
})
267+
268+
return info
269+
270+
async def cleanup(self) -> None:
271+
"""清理资源"""
272+
logger.info("Cleaning up model manager...")
273+
274+
# 清理GPU缓存
275+
if self.device == "cuda" and torch.cuda.is_available():
276+
torch.cuda.empty_cache()
277+
278+
# 清理模型引用
279+
self.models.clear()
280+
self._initialized = False
281+
282+
logger.info("Model manager cleanup completed")
283+
284+
def __del__(self):
285+
"""析构函数"""
286+
if hasattr(self, 'models') and self.models:
287+
# 注意:在析构函数中不能使用async
288+
if self.device == "cuda" and torch.cuda.is_available():
289+
torch.cuda.empty_cache()

0 commit comments

Comments
 (0)