bert_base/services/tax/app.py

451 lines
14 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
from flask import Flask, request, jsonify
import torch
from transformers import BertTokenizer, BertForSequenceClassification
import joblib
import re
from functools import lru_cache
from concurrent.futures import ThreadPoolExecutor
from typing import List, Dict
import threading
import logging
import atexit
from datetime import datetime
import time
# ==================== 从环境变量读取配置 ====================
SERVICE_NAME = os.environ.get('SERVICE_NAME', 'tax_classifier')
SERVICE_PORT = int(os.environ.get('SERVICE_PORT', 5002))
MODEL_ID = os.environ.get('MODEL_ID', 'circles1/tax_rate')
MODEL_DIR = os.environ.get('MODEL_DIR', '/app/services/tax_classifier/model')
# 兼容旧版路径如果MODEL_DIR未设置尝试从SERVICE_NAME推断
if not os.environ.get('MODEL_DIR'):
# 从SERVICE_NAME推断模型目录
service_folder = SERVICE_NAME.replace('_classifier', '')
MODEL_DIR = f'/app/services/{service_folder}/model'
# 模型文件路径
LABEL_ENCODER_PATH = os.path.join(MODEL_DIR, 'label_encoder_roberta_large.pkl')
# 如果上面路径不存在,尝试在模型根目录查找
if not os.path.exists(LABEL_ENCODER_PATH):
LABEL_ENCODER_PATH = os.path.join(MODEL_DIR, 'label_encoder.pkl')
# ==================== 配置参数 ====================
MAX_LENGTH = 512
DEVICE = os.environ.get('DEVICE', 'cuda' if torch.cuda.is_available() else 'cpu')
# 如果设备设置为cuda但不可用回退到cpu
if DEVICE == 'cuda' and not torch.cuda.is_available():
DEVICE = 'cpu'
print(f"⚠️ CUDA not available, falling back to {DEVICE}")
BATCH_SIZE = int(os.environ.get('BATCH_SIZE', 32))
MAX_WORKERS = int(os.environ.get('MAX_WORKERS', 4))
CACHE_SIZE = 2000
TOKEN_CACHE_SIZE = 1000
# ==================== 日志配置 ====================
logging.basicConfig(
level=logging.INFO,
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s'
)
logger = logging.getLogger(__name__)
# 打印启动信息
logger.info("=" * 50)
logger.info(f"Starting service: {SERVICE_NAME}")
logger.info(f"Service port: {SERVICE_PORT}")
logger.info(f"Model ID: {MODEL_ID}")
logger.info(f"Model directory: {MODEL_DIR}")
logger.info(f"Device: {DEVICE}")
logger.info(f"Batch size: {BATCH_SIZE}")
logger.info(f"Max workers: {MAX_WORKERS}")
logger.info("=" * 50)
app = Flask(__name__)
app.config['JSON_AS_ASCII'] = False
# 全局变量锁
model_lock = threading.Lock()
class Predictor:
"""预测器类,用于管理模型和分词器的生命周期"""
_instance = None
def __new__(cls):
if cls._instance is None:
with model_lock:
if cls._instance is None:
cls._instance = super().__new__(cls)
cls._instance._initialized = False
return cls._instance
def __init__(self):
if self._initialized:
return
global LABEL_ENCODER_PATH # 关键修复
logger.info(f"Initializing Predictor on device: {DEVICE}")
start_time = time.time()
try:
# 检查模型目录是否存在
if not os.path.exists(MODEL_DIR):
raise FileNotFoundError(f"Model directory not found: {MODEL_DIR}")
# 检查label_encoder是否存在
if not os.path.exists(LABEL_ENCODER_PATH):
# 尝试在目录中查找label_encoder文件
import glob
encoder_files = glob.glob(os.path.join(MODEL_DIR, 'label_encoder*.pkl'))
if encoder_files:
actual_encoder_path = encoder_files[0]
logger.info(f"Found label_encoder at: {actual_encoder_path}")
LABEL_ENCODER_PATH = actual_encoder_path
else:
raise FileNotFoundError(f"Label encoder not found in {MODEL_DIR}")
# 加载模型和分词器
logger.info(f"Loading tokenizer and model from {MODEL_DIR}...")
self.tokenizer = BertTokenizer.from_pretrained(MODEL_DIR)
self.model = BertForSequenceClassification.from_pretrained(MODEL_DIR).to(DEVICE)
self.model.eval()
logger.info(f"Model loaded successfully")
# 加载标签映射器
logger.info(f"Loading label encoder from {LABEL_ENCODER_PATH}...")
self.label_encoder = joblib.load(LABEL_ENCODER_PATH)
self.num_labels = len(self.label_encoder.classes_)
logger.info(f"Loaded {self.num_labels} labels")
# 创建线程池
self.executor = ThreadPoolExecutor(max_workers=MAX_WORKERS)
logger.info(f"ThreadPoolExecutor created with {MAX_WORKERS} workers")
# 注册关闭钩子
atexit.register(self._shutdown)
# 模型预热
self._warmup()
self._initialized = True
elapsed_time = time.time() - start_time
logger.info(f"Predictor initialized successfully in {elapsed_time:.2f}s")
except Exception as e:
logger.error(f"Failed to initialize Predictor: {e}")
raise
def _shutdown(self):
"""清理资源"""
logger.info("Shutting down executor...")
self.executor.shutdown(wait=True)
logger.info("Executor shutdown complete")
def _warmup(self):
"""模型预热确保CUDA初始化"""
logger.info("Warming up model...")
try:
test_text = "测试文本"
self.predict_single(test_text)
logger.info("Model warmup completed")
except Exception as e:
logger.warning(f"Model warmup failed: {e}")
@staticmethod
@lru_cache(maxsize=CACHE_SIZE)
def clean_text(text: str) -> str:
"""文本清洗函数,带缓存"""
if not isinstance(text, str):
return ""
# 只保留中文字符
cleaned_text = re.sub(r'[^\u4e00-\u9fa5]', '', text)
return cleaned_text.strip()
@lru_cache(maxsize=TOKEN_CACHE_SIZE)
def tokenize_text(self, text: str):
"""缓存tokenization结果"""
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 {"type": "", "tax": "", "confidence": 0.0, "error": "Invalid input"}
try:
# 清洗文本
cleaned_text = self.clean_text(text)
if not cleaned_text:
logger.warning(f"Empty text after cleaning: {text[:50]}...")
return {"type": "", "tax": "", "confidence": 0.0}
# Tokenization使用缓存
inputs = self.tokenize_text(cleaned_text)
if inputs is None:
return {"type": "", "tax": "", "confidence": 0.0}
# 移动到设备
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_idx = top_idx.item()
tax_label = self.label_encoder.inverse_transform([label_idx])[0]
confidence = top_prob.item()
# 尝试解析标签格式
tax_arr = tax_label.split("_")
if len(tax_arr) >= 2:
result = {
"type": tax_arr[0],
"tax": tax_arr[1],
"confidence": round(confidence, 4)
}
else:
result = {
"type": tax_label,
"tax": "",
"confidence": round(confidence, 4)
}
return result
except torch.cuda.OutOfMemoryError as e:
logger.error(f"CUDA OOM error: {e}")
torch.cuda.empty_cache()
return {"type": "", "tax": "", "confidence": 0.0, "error": "GPU memory exhausted"}
except Exception as e:
logger.error(f"Prediction error for text '{text[:50]}...': {e}")
return {"type": "", "tax": "", "confidence": 0.0, "error": str(e)}
def batch_predict(self, texts: List[str], batch_size: int = BATCH_SIZE) -> List[Dict]:
"""批量预测(并发处理)"""
if not texts:
return []
logger.info(f"Processing {len(texts)} texts with batch_size={batch_size}")
start_time = time.time()
results = []
total_batches = (len(texts) + batch_size - 1) // batch_size
for batch_idx in range(total_batches):
batch_start = batch_idx * batch_size
batch_end = min(batch_start + batch_size, len(texts))
batch_texts = texts[batch_start:batch_end]
# 提交批次任务
futures = [
self.executor.submit(self.predict_single, text)
for text in batch_texts
]
# 收集结果
batch_results = [future.result() for future in futures]
results.extend(batch_results)
if batch_idx % 10 == 0 and batch_idx > 0:
logger.info(f"Processed {batch_end}/{len(texts)} texts")
elapsed_time = time.time() - start_time
logger.info(f"Batch prediction completed in {elapsed_time:.2f}s for {len(texts)} texts")
return results
def get_stats(self) -> Dict:
"""获取预测器统计信息"""
return {
"service_name": SERVICE_NAME,
"model_id": MODEL_ID,
"model_dir": MODEL_DIR,
"device": DEVICE,
"num_labels": self.num_labels,
"max_length": MAX_LENGTH,
"cache_size": {
"text_clean": CACHE_SIZE,
"tokenization": TOKEN_CACHE_SIZE
},
"thread_pool": {
"max_workers": MAX_WORKERS,
"batch_size": BATCH_SIZE
}
}
# 初始化预测器(单例)
try:
predictor = Predictor()
logger.info("Predictor instance created successfully")
except Exception as e:
logger.error(f"Failed to create predictor instance: {e}")
raise
@app.route('/predicts', methods=['POST'])
def predicts():
"""预测接口"""
try:
# 解析请求
data = request.get_json()
if not data:
return jsonify({
"status": "error",
"error": "Invalid request, JSON body required"
}), 400
if 'product_names' not in data:
return jsonify({
"status": "error",
"error": "Invalid request, 'product_names' is required"
}), 400
product_names = data['product_names']
if not isinstance(product_names, list):
product_names = [product_names]
# 限制最大处理数量
max_items = data.get('max_items', 1000)
if len(product_names) > max_items:
logger.warning(f"Too many items: {len(product_names)} > {max_items}, truncating")
product_names = product_names[:max_items]
# 去重(可选)
if data.get('deduplicate', False):
product_names = list(dict.fromkeys(product_names))
logger.info(f"After deduplication: {len(product_names)} unique texts")
# 批量预测
results = predictor.batch_predict(product_names)
# 构建响应
response = {
"status": "success",
"predictions": results,
"metadata": {
"total": len(results),
"timestamp": datetime.now().isoformat(),
"service": SERVICE_NAME,
"device": DEVICE
}
}
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('/predict', methods=['POST'])
def predict():
"""单条预测接口(简化版)"""
try:
data = request.get_json()
if not data or 'product_name' not in data:
return jsonify({
"status": "error",
"error": "Invalid request, 'product_name' is required"
}), 400
product_name = data['product_name']
result = predictor.predict_single(product_name)
return jsonify({
"status": "success",
"prediction": result
})
except Exception as e:
logger.error(f"Single prediction endpoint error: {e}", exc_info=True)
return jsonify({
"status": "error",
"error": str(e)
}), 500
@app.route('/health', methods=['GET'])
def health_check():
"""健康检查接口"""
try:
stats = predictor.get_stats()
return jsonify({
"status": "healthy",
"timestamp": datetime.now().isoformat(),
"stats": stats
})
except Exception as e:
return jsonify({
"status": "unhealthy",
"error": str(e)
}), 500
@app.route('/stats', methods=['GET'])
def get_stats():
"""获取预测器统计信息"""
try:
stats = predictor.get_stats()
return jsonify({
"status": "success",
"stats": 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 Flask application: {SERVICE_NAME}")
logger.info(f"Model path: {MODEL_DIR}")
logger.info(f"Device: {DEVICE}")
logger.info(f"Port: {SERVICE_PORT}")
logger.info(f"Batch size: {BATCH_SIZE}")
logger.info(f"Max workers: {MAX_WORKERS}")
app.run(
host='0.0.0.0',
port=SERVICE_PORT,
threaded=True,
debug=False
)
else:
# 兼容 WSGI 标准(如 Gunicorn
application = app
logger.info(f"Application loaded for WSGI server: {SERVICE_NAME}")