bert_base/services/unified/app.py

506 lines
17 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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:
"""根据模型类型格式化输出"""
result = {"confidence": confidence}
if self.model_type == 'tax':
# tax 模型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:
# 通用模型:直接输出 label
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()}")