514 lines
17 KiB
Python
514 lines
17 KiB
Python
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()}") |