工业实战:用ResNet50和Flask快速搭建一个Web版图像分类API(附Docker部署脚本)
工业实战:用ResNet50和Flask快速搭建Web版图像分类API
在初创团队的日常开发中,快速将AI模型转化为可用的服务是提升效率的关键。假设你手头已经有一个训练好的ResNet50模型,如何让它从实验室走向生产线?本文将带你用Flask构建一个轻量级Web API,并通过Docker实现一键部署,让图像分类能力随时待命。
1. 环境准备与模型加载
首先确保你的开发环境已安装Python 3.7+。建议使用conda创建独立环境:
conda create -n resnet_api python=3.8
conda activate resnet_api
安装核心依赖库:
pip install torch torchvision flask pillow
加载预训练模型时,需要注意PyTorch的模型保存方式。假设你的模型权重保存为resnet50_custom.pth,加载代码应该这样写:
import torch
import torchvision.models as models
model = models.resnet50(pretrained=False)
num_classes = 10 # 根据你的分类任务调整
model.fc = torch.nn.Linear(model.fc.in_features, num_classes)
model.load_state_dict(torch.load('resnet50_custom.pth'))
model.eval()
常见问题排查:
- 如果遇到维度不匹配错误,检查
num_classes是否与训练时一致 - 在CPU环境下运行时添加
map_location='cpu'参数 - 模型文件较大时(>100MB),建议使用
.pt格式替代.pth
2. Flask API开发实战
Flask的轻量级特性使其成为快速开发API的理想选择。我们先构建基础路由:
from flask import Flask, request, jsonify
from PIL import Image
import io
app = Flask(__name__)
@app.route('/predict', methods=['POST'])
def predict():
if 'file' not in request.files:
return jsonify({'error': 'No file uploaded'}), 400
file = request.files['file']
image = Image.open(io.BytesIO(file.read()))
# 图像预处理和预测逻辑将放在这里
return jsonify({'class_id': 0, 'confidence': 0.95})
if __name__ == '__main__':
app.run(host='0.0.0.0', port=5000)
图像预处理需要与训练时保持一致,通常包括:
- 调整大小至224×224像素
- 转换为RGB格式
- 归一化处理(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
- 转换为Tensor并添加batch维度
完整预处理函数示例:
from torchvision import transforms
preprocess = transforms.Compose([
transforms.Resize(256),
transforms.CenterCrop(224),
transforms.ToTensor(),
transforms.Normalize(
mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225]
)
])
def transform_image(image):
return preprocess(image).unsqueeze(0)
3. 性能优化技巧
直接在生产环境运行上述代码会遇到性能瓶颈,以下是几个关键优化点:
批处理支持: 修改API以支持多图同时预测,显著提升吞吐量:
@app.route('/batch_predict', methods=['POST'])
def batch_predict():
files = request.files.getlist('files')
batch = torch.stack([transform_image(Image.open(io.BytesIO(f.read())))
for f in files])
with torch.no_grad():
outputs = model(batch)
# 后续处理...
GPU加速: 如果服务器配有NVIDIA显卡,添加以下代码启用CUDA:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model.to(device)
# 在预测时移动数据到GPU
inputs = inputs.to(device)
异步处理: 对于高延迟请求,可以使用Celery实现异步队列:
from celery import Celery
celery = Celery('tasks', broker='redis://localhost:6379/0')
@celery.task
def async_predict(image_data):
# 预测逻辑
return result
4. Docker化部署方案
将服务容器化是保证环境一致性的最佳实践。创建Dockerfile:
FROM python:3.8-slim
WORKDIR /app
COPY requirements.txt .
RUN pip install --no-cache-dir -r requirements.txt
COPY . .
EXPOSE 5000
CMD ["gunicorn", "--bind", "0.0.0.0:5000", "--workers", "4", "app:app"]
对应的docker-compose.yml配置:
version: '3'
services:
web:
build: .
ports:
- "5000:5000"
volumes:
- ./models:/app/models
environment:
- FLASK_ENV=production
redis:
image: "redis:alpine"
部署时建议使用Gunicorn作为WSGI服务器:
# 构建镜像
docker-compose build
# 启动服务
docker-compose up -d
# 查看日志
docker-compose logs -f
5. 监控与扩展
生产环境还需要考虑监控和自动扩展。Prometheus+Granfana是经典组合:
from prometheus_client import start_http_server, Counter
REQUEST_COUNT = Counter('request_count', 'API request count')
@app.before_request
def before_request():
REQUEST_COUNT.inc()
对于流量波动大的场景,可以配置Kubernetes的HPA(Horizontal Pod Autoscaler):
apiVersion: autoscaling/v2beta2
kind: HorizontalPodAutoscaler
metadata:
name: resnet-api
spec:
scaleTargetRef:
apiVersion: apps/v1
kind: Deployment
name: resnet-api
minReplicas: 2
maxReplicas: 10
metrics:
- type: Resource
resource:
name: cpu
target:
type: Utilization
averageUtilization: 70
6. 安全防护措施
公开API必须考虑安全性,以下是基本防护:
请求限流: 使用Flask-Limiter防止暴力请求:
from flask_limiter import Limiter
from flask_limiter.util import get_remote_address
limiter = Limiter(
app,
key_func=get_remote_address,
default_limits=["200 per day", "50 per hour"]
)
输入验证: 添加文件类型和大小检查:
ALLOWED_EXTENSIONS = {'png', 'jpg', 'jpeg'}
MAX_CONTENT_LENGTH = 8 * 1024 * 1024 # 8MB
app.config['MAX_CONTENT_LENGTH'] = MAX_CONTENT_LENGTH
def allowed_file(filename):
return '.' in filename and \
filename.rsplit('.', 1)[1].lower() in ALLOWED_EXTENSIONS
API密钥认证: 简单的基于令牌的认证:
from functools import wraps
def require_api_key(view_function):
@wraps(view_function)
def decorated_function(*args, **kwargs):
if request.headers.get('X-API-KEY') != os.getenv('API_KEY'):
return jsonify({'error': 'Unauthorized'}), 403
return view_function(*args, **kwargs)
return decorated_function
7. 客户端集成示例
最后提供一个完整的Python客户端示例,方便其他系统集成:
import requests
def classify_image(image_path, api_url, api_key=None):
with open(image_path, 'rb') as f:
files = {'file': f}
headers = {'X-API-KEY': api_key} if api_key else {}
response = requests.post(api_url, files=files, headers=headers)
return response.json()
# 使用示例
result = classify_image('test.jpg', 'http://localhost:5000/predict')
print(result)
对于Web前端,可以使用简单的HTML表单:
<form action="/predict" method="post" enctype="multipart/form-data">
<input type="file" name="file" accept="image/*">
<button type="submit">分类</button>
</form>
<div id="result"></div>
<script>
document.querySelector('form').addEventListener('submit', async (e) => {
e.preventDefault();
const formData = new FormData(e.target);
const response = await fetch('/predict', {
method: 'POST',
body: formData
});
const result = await response.json();
document.getElementById('result').innerText =
`类别: ${result.class_id}, 置信度: ${result.confidence}`;
});
</script>
更多推荐
所有评论(0)