工业实战:用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)

图像预处理需要与训练时保持一致,通常包括:

  1. 调整大小至224×224像素
  2. 转换为RGB格式
  3. 归一化处理(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
  4. 转换为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>

更多推荐