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