import os import sys import warnings from flask import Flask, request, jsonify import torch from transformers import BertTokenizer, BertForSequenceClassification import joblib import re from functools import lru_cache from typing import List, Dict import threading import logging import atexit from datetime import datetime import time # 忽略 scikit-learn 版本警告 warnings.filterwarnings("ignore", category=UserWarning, module="sklearn") # ==================== 从环境变量读取配置 ==================== SERVICE_PORT = int(os.environ.get('SERVICE_PORT', 5003)) MAX_WORKERS = int(os.environ.get('MAX_WORKERS', 2)) CACHE_SIZE = 2000 TOKEN_CACHE_SIZE = 1000 # ==================== 设备配置 ==================== MAX_LENGTH = 512 DEVICE = os.environ.get('DEVICE', 'cuda' if torch.cuda.is_available() else 'cpu') if DEVICE == 'cuda' and not torch.cuda.is_available(): DEVICE = 'cpu' print(f"⚠️ CUDA not available, falling back to {DEVICE}") # ==================== 日志配置 ==================== logging.basicConfig( level=logging.INFO, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s' ) logger = logging.getLogger(__name__) # 如果使用 GPU,限制并发线程数为 1 以避免并发 GPU 推理导致 OOM if isinstance(DEVICE, str) and DEVICE.startswith('cuda') and MAX_WORKERS > 1: logger.warning("CUDA in use — limiting MAX_WORKERS to 1 to avoid concurrent GPU inference") MAX_WORKERS = 1 app = Flask(__name__) app.config['JSON_AS_ASCII'] = False # 全局变量锁 model_lock = threading.Lock() def discover_models(): """ 从环境变量中发现所有模型配置 约定:每个模型需要以下环境变量: - {PREFIX}_MODEL_ID: 模型ID - {PREFIX}_MODEL_DIR: 模型目录 - {PREFIX}_SERVICE_NAME: 服务名称(用于日志和标识) - {PREFIX}_TYPE: 模型类型(tax/address/other),用于决定输出格式 """ models = {} for key, value in os.environ.items(): if key.endswith('_MODEL_ID'): prefix = key[:-9] # 去掉 _MODEL_ID service_name = os.getenv(f'{prefix}_SERVICE_NAME', f'{prefix.lower()}_classifier') model_type = os.getenv(f'{prefix}_TYPE', 'general') models[prefix.lower()] = { 'prefix': prefix, 'model_id': value, 'model_dir': os.getenv(f'{prefix}_MODEL_DIR', f'/app/services/{prefix.lower()}/model'), 'service_name': service_name, 'type': model_type, # tax, address, general } return models class BasePredictor: """基础预测器类 - 每个模型实例""" def __init__(self, config): self.config = config self.model_type = config['type'] self.model_dir = config['model_dir'] self.service_name = config['service_name'] self.model_id = config['model_id'] self._initialized = False self._init_model() def _find_label_encoder(self): """自动查找 label_encoder 文件""" # 尝试常见的文件名 possible_names = [ 'label_encoder.pkl', 'label_encoder_roberta_large.pkl', 'label_encoder_roberta.pkl', 'label_encoder_bert.pkl' ] for name in possible_names: path = os.path.join(self.model_dir, name) if os.path.exists(path): return path # 尝试通配符查找 import glob encoder_files = glob.glob(os.path.join(self.model_dir, 'label_encoder*.pkl')) if encoder_files: return encoder_files[0] raise FileNotFoundError(f"No label_encoder found in {self.model_dir}") def _init_model(self): """初始化模型""" logger.info(f"Loading model '{self.service_name}' from {self.model_dir}...") try: if not os.path.exists(self.model_dir): raise FileNotFoundError(f"Model directory not found: {self.model_dir}") # 查找 label_encoder label_encoder_path = self._find_label_encoder() logger.info(f"Found label_encoder: {label_encoder_path}") # 加载模型和分词器 self.tokenizer = BertTokenizer.from_pretrained(self.model_dir) self.model = BertForSequenceClassification.from_pretrained(self.model_dir).to(DEVICE) self.model.eval() # 加载标签映射器 self.label_encoder = joblib.load(label_encoder_path) self.label_map = {i: label for i, label in enumerate(self.label_encoder.classes_)} self.num_labels = len(self.label_map) self._initialized = True logger.info(f"✅ Model '{self.service_name}' loaded, {self.num_labels} labels") except Exception as e: logger.error(f"❌ Failed to load model '{self.service_name}': {e}") raise @lru_cache(maxsize=CACHE_SIZE) def clean_text(self, text: str) -> str: """文本清洗函数,带缓存""" if not isinstance(text, str): return "" cleaned_text = re.sub(r'[^\u4e00-\u9fa5]', '', text) return cleaned_text.strip() def tokenize_text(self, text: str): """Tokenize input text (no large-cache to avoid memory growth). Returns tokenizer output on CPU. """ if not text: return None return self.tokenizer( text, return_tensors="pt", truncation=True, padding=True, max_length=MAX_LENGTH ) def predict_single(self, text: str) -> Dict: """单个文本预测""" if not text or not isinstance(text, str): return {"error": "Invalid input"} try: cleaned_text = self.clean_text(text) if not cleaned_text: return {"error": "Empty text after cleaning"} inputs = self.tokenize_text(cleaned_text) if inputs is None: return {"error": "Tokenization failed"} inputs = {k: v.to(DEVICE) for k, v in inputs.items()} with torch.no_grad(): outputs = self.model(**inputs) probs = torch.softmax(outputs.logits, dim=1).cpu() top_prob, top_idx = torch.topk(probs, k=1) label = self.label_map[top_idx.item()] confidence = round(top_prob.item(), 4) # 根据模型类型格式化输出 return self._format_output(label, confidence) except torch.cuda.OutOfMemoryError as e: logger.error(f"CUDA OOM in {self.service_name}: {e}") torch.cuda.empty_cache() return {"error": "GPU memory exhausted"} except Exception as e: logger.error(f"Prediction error in {self.service_name}: {e}") return {"error": str(e)} def _format_output(self, label: str, confidence: float) -> Dict: """根据模型类型格式化输出""" # 确保所有类型都是 Python 原生类型 confidence = float(confidence) # 确保 label 是字符串 if not isinstance(label, str): if hasattr(label, 'item'): label = str(label.item()) else: label = str(label) result = {"confidence": confidence} if self.model_type == 'tax': parts = label.split('_') if len(parts) >= 2: result['type'] = parts[0] result['tax'] = '_'.join(parts[1:]) else: result['type'] = label result['tax'] = '' elif self.model_type == 'address': result['address'] = label elif self.model_type == 'brand': result['brand'] = label else: result['label'] = label return result def get_stats(self) -> Dict: """获取模型统计信息""" return { "service_name": self.service_name, "model_id": self.model_id, "model_type": self.model_type, "model_dir": self.model_dir, "device": DEVICE, "num_labels": self.num_labels, "labels": list(self.label_map.values())[:10] # 只显示前10个 } class ModelManager: """模型管理器 - 管理所有模型实例""" def __init__(self): self.predictors = {} self._load_all_models() def _load_all_models(self): """加载所有模型""" logger.info("=" * 50) logger.info("Initializing Model Manager...") logger.info(f"Device: {DEVICE}") logger.info("=" * 50) models_config = discover_models() if not models_config: logger.warning("⚠️ No models configured in .env") return logger.info(f"Found {len(models_config)} model(s):") for name, config in models_config.items(): logger.info(f" - {name}: {config['model_id']} (type: {config['type']})") logger.info("-" * 50) for name, config in models_config.items(): try: self.predictors[name] = BasePredictor(config) logger.info(f"✅ {name} loaded successfully") except Exception as e: logger.error(f"❌ Failed to load {name}: {e}") logger.info("=" * 50) logger.info(f"Loaded {len(self.predictors)}/{len(models_config)} models") logger.info(f"Available: {list(self.predictors.keys())}") logger.info("=" * 50) def get_predictor(self, name: str): """获取指定名称的预测器。 尝试多种匹配:精确 key,大小写无关,service_name 或 model_id 匹配。 返回匹配的 Predictor 或 None。 """ if not name: return None # 直接按 key 查找(优先) if name in self.predictors: logger.debug(f"Predictor matched by key: {name}") return self.predictors[name] # 尝试大小写不敏感的 key lower_name = name.lower() for k in self.predictors.keys(): if k.lower() == lower_name: logger.debug(f"Predictor matched by case-insensitive key: {k} for request '{name}'") return self.predictors[k] # 尝试按 predictor 的 service_name 或 model_id 匹配 for k, predictor in self.predictors.items(): try: if getattr(predictor, 'service_name', '').lower() == lower_name: logger.debug(f"Predictor matched by service_name: {k} -> {predictor.service_name}") return predictor if getattr(predictor, 'model_id', '').lower() == lower_name: logger.debug(f"Predictor matched by model_id: {k} -> {predictor.model_id}") return predictor except Exception: continue # 不再做模糊/子串匹配以避免错误映射;只做严格或基于 service_name/model_id 的匹配 return None def list_models(self) -> List[str]: """列出所有可用模型""" return list(self.predictors.keys()) def get_stats(self) -> Dict: """获取所有模型的统计信息""" return { name: predictor.get_stats() for name, predictor in self.predictors.items() } # ==================== 初始化模型管理器 ==================== model_manager = ModelManager() @app.route('/predict', methods=['POST']) def predict(): """统一预测接口""" try: # 检查 Content-Type if not request.is_json: return jsonify({ "status": "error", "error": "Content-Type must be application/json" }), 400 # 解析请求 data = request.get_json(silent=True) if not data: return jsonify({ "status": "error", "error": "Invalid JSON body" }), 400 # 检查必填字段 if 'model' not in data: return jsonify({ "status": "error", "error": f"Missing required field: 'model'. Available: {model_manager.list_models()}" }), 400 if 'text' not in data: return jsonify({ "status": "error", "error": "Missing required field: 'text'" }), 400 model_name = data['model'] text = data['text'] # 验证模型是否存在 predictor = model_manager.get_predictor(model_name) if predictor is None: return jsonify({ "status": "error", "error": f"Invalid model: {model_name}. Available: {model_manager.list_models()}" }), 400 # 验证文本 if not isinstance(text, str): return jsonify({ "status": "error", "error": "Invalid 'text' field, must be a string" }), 400 if not text.strip(): return jsonify({ "status": "error", "error": "Empty text" }), 400 # 记录并返回所使用的 predictor 信息(仅当请求中包含 debug=True 时会把详细信息返给客户端) predictor_info = { "key": None, "service_name": getattr(predictor, 'service_name', None), "model_id": getattr(predictor, 'model_id', None), "model_dir": getattr(predictor, 'model_dir', None) } # 尝试找出 predictor 的注册 key for k, p in model_manager.predictors.items(): if p is predictor: predictor_info['key'] = k break # 进行预测 result = predictor.predict_single(text) response = { "status": "success", "model": model_name, "prediction": result, "metadata": { "timestamp": datetime.now().isoformat(), "device": DEVICE } } # 如果请求开启 debug,则在响应中包含 predictor 信息 if data.get('debug', False): response['predictor'] = predictor_info logger.info(f"Request for model='{model_name}' routed to predictor key='{predictor_info.get('key')}', service_name='{predictor_info.get('service_name')}', model_id='{predictor_info.get('model_id')}'") return jsonify(response) except Exception as e: logger.error(f"Prediction endpoint error: {e}", exc_info=True) return jsonify({ "status": "error", "error": f"Internal server error: {str(e)}" }), 500 @app.route('/models', methods=['GET']) def list_models(): """列出所有可用的模型""" return jsonify({ "status": "success", "models": model_manager.list_models(), "details": model_manager.get_stats() }) @app.route('/debug_models', methods=['GET']) def debug_models(): """返回 model_manager 的映射信息,便于诊断模型与 key 的对应关系(仅内部/运维使用)。""" try: data = {} for k, p in model_manager.predictors.items(): data[k] = { 'service_name': getattr(p, 'service_name', None), 'model_id': getattr(p, 'model_id', None), 'model_dir': getattr(p, 'model_dir', None), } return jsonify({'status': 'success', 'mappings': data}) except Exception as e: logger.error(f"debug_models error: {e}") return jsonify({'status': 'error', 'error': str(e)}), 500 @app.route('/health', methods=['GET']) def health_check(): """健康检查接口""" try: return jsonify({ "status": "healthy", "timestamp": datetime.now().isoformat(), "device": DEVICE, "models": model_manager.list_models(), "loaded": len(model_manager.predictors) }) except Exception as e: return jsonify({ "status": "unhealthy", "error": str(e) }), 500 @app.route('/stats', methods=['GET']) def get_stats(): """获取统计信息""" try: return jsonify({ "status": "success", "device": DEVICE, "models": model_manager.get_stats() }) except Exception as e: return jsonify({ "status": "error", "error": str(e) }), 500 @app.errorhandler(404) def not_found(error): return jsonify({ "status": "error", "error": "Endpoint not found" }), 404 @app.errorhandler(500) def internal_error(error): return jsonify({ "status": "error", "error": "Internal server error" }), 500 if __name__ == '__main__': logger.info(f"Starting unified Flask application on port {SERVICE_PORT}") logger.info(f"Loaded models: {model_manager.list_models()}") logger.info(f"Device: {DEVICE}") app.run( host='0.0.0.0', port=SERVICE_PORT, threaded=True, debug=True, ) else: application = app logger.info(f"Application loaded for WSGI server with models: {model_manager.list_models()}")