第48课:TensorFlow|TF模型线上部署入门【本地服务封装、接口快速开发】

文章目录

    • [1. 课前导读](#1. 课前导读)
      • [1.1 本节课学习目标](#1.1 本节课学习目标)
      • [1.2 知识重难点](#1.2 知识重难点)
      • [1.3 学习前置条件](#1.3 学习前置条件)
      • [1.4 学完可掌握能力](#1.4 学完可掌握能力)
      • [1.5 行业应用场景](#1.5 行业应用场景)
    • [2. 核心理论精讲](#2. 核心理论精讲)
      • [2.1 模型服务的架构模式](#2.1 模型服务的架构模式)
      • [2.2 模型导出与优化](#2.2 模型导出与优化)
      • [2.3 服务框架选型](#2.3 服务框架选型)
      • [2.4 批处理与合并请求](#2.4 批处理与合并请求)
      • [2.5 容器化部署](#2.5 容器化部署)
      • [2.6 监控与日志](#2.6 监控与日志)
    • [3. 环境搭建与工具配置](#3. 环境搭建与工具配置)
      • [3.1 安装依赖](#3.1 安装依赖)
      • [3.2 准备预训练模型](#3.2 准备预训练模型)
    • [4. 代码实战教学](#4. 代码实战教学)
      • [4.1 使用FastAPI创建同步推理API](#4.1 使用FastAPI创建同步推理API)
      • [4.2 异步批处理队列(高级)](#4.2 异步批处理队列(高级))
      • [4.3 使用TensorFlow Serving部署](#4.3 使用TensorFlow Serving部署)
      • [4.4 使用Docker封装FastAPI服务](#4.4 使用Docker封装FastAPI服务)
      • [4.5 性能测试(使用locust)](#4.5 性能测试(使用locust))
    • [5. 案例实操演练](#5. 案例实操演练)
      • [5.1 模型准备](#5.1 模型准备)
      • [5.2 FastAPI服务增强](#5.2 FastAPI服务增强)
      • [5.3 部署到云服务(示例:AWS Lambda 不适合大模型,建议ECS)](#5.3 部署到云服务(示例:AWS Lambda 不适合大模型,建议ECS))
      • [5.4 压力测试与结果分析](#5.4 压力测试与结果分析)
    • [6. 常见坑点与排错总结](#6. 常见坑点与排错总结)
      • [6.1 模型加载与签名](#6.1 模型加载与签名)
      • [6.2 性能问题](#6.2 性能问题)
      • [6.3 并发与稳定性](#6.3 并发与稳定性)
      • [6.4 TensorFlow Serving坑点](#6.4 TensorFlow Serving坑点)
    • [7. 知识点总结 + 课后作业](#7. 知识点总结 + 课后作业)
      • [7.1 核心知识点梳理](#7.1 核心知识点梳理)
      • [7.2 基础作业](#7.2 基础作业)
      • [7.3 进阶实操作业](#7.3 进阶实操作业)
      • [7.4 思考拓展题](#7.4 思考拓展题)
  • [🔗《TensorFlow2.x: 深度学习入门到高阶实战教程》系列课程导航](#🔗《TensorFlow2.x: 深度学习入门到高阶实战教程》系列课程导航)

1. 课前导读

1.1 本节课学习目标

  • 理解线上模型服务的需求与挑战(延迟、吞吐量、扩展性)。
  • 掌握将Keras模型导出为优化后的SavedModel,支持批量推理。
  • 学会使用FastAPI构建轻量级高性能API服务,包括同步和异步端点。
  • 实现批处理与请求合并(动态batching),提高GPU利用率。
  • 了解TensorFlow Serving的安装与基本使用,实现模型版本管理。
  • 掌握使用Docker打包模型服务,实现环境一致性。
  • 通过实战案例(图像分类服务)完成从模型导出到服务启动的全过程。

1.2 知识重难点

类别 内容
重点 SavedModel导出与签名定义;FastAPI路由与请求处理;批处理推理优化;TensorFlow Serving的部署与调用
难点 动态批处理的请求队列管理;异步处理与WebSocket;TensorFlow Serving的gRPC客户端编写;Docker镜像优化(减小体积)
易混淆点 REST vs gRPC;同步API vs 异步API;模型签名中的输入输出名称;TensorFlow Serving的模型版本策略

1.3 学习前置条件

  • 已完成第20课模型保存与加载,熟悉SavedModel。
  • 能够训练简单的CNN模型。
  • 了解Python Web框架基础和Docker概念(可选)。

1.4 学完可掌握能力

  • 独立将模型封装为REST API,提供稳定推理服务。
  • 使用FastAPI实现高并发接口。
  • 使用TensorFlow Serving进行企业级模型部署。
  • 容器化模型服务,便于迁移和扩展。

1.5 行业应用场景

  • 在线图像识别:手机APP调用云端API。
  • 推荐系统:实时获取用户推荐。
  • 金融反欺诈:低延迟评分服务。
  • IoT边缘计算:在边缘节点部署模型服务。

2. 核心理论精讲

2.1 模型服务的架构模式

模型服务通常以HTTP/REST或gRPC协议对外提供API。客户端发送请求(图像、文本等),服务端进行预处理、模型推理、后处理,返回结果。架构要点:

  • 无状态:每个请求独立,便于水平扩展。
  • 高并发:利用异步IO、线程池、动态批处理。
  • 可观测:日志、监控、链路追踪。

2.2 模型导出与优化

训练后的模型(Keras H5或SavedModel)需要导出为优化格式。推荐使用SavedModel,它包含图结构和权重,支持跨语言。优化手段:

  • 静态图转换@tf.function 提前编译。
  • 量化:INT8量化减小体积和加速(需评估精度损失)。
  • 移除训练相关操作:如Dropout、BatchNormalization的统计。

导出代码:

python 复制代码
tf.saved_model.save(model, 'model/1', signatures=...)  # 版本号子目录

2.3 服务框架选型

框架 特点 适用场景
Flask 轻量、简单 低并发、原型验证
FastAPI 异步、高性能、自动文档 高并发、生产级API
gRPC 二进制协议、低延迟、多语言 内部微服务调用
TensorFlow Serving 专为TF模型设计,版本管理、批处理 大规模生产部署

本课重点使用FastAPI(因其易用性和性能)和TensorFlow Serving。

2.4 批处理与合并请求

深度学习推理在批量处理时效率更高(GPU并行)。若请求是随机到达的单个样本,服务端可将多个请求合并为一个批次,称为动态批处理。实现方式:

  • 使用队列收集请求,等待固定时间或积累到一定数量。
  • 调用模型批量推理,将结果返回给各个请求。
  • TensorFlow Serving内置了动态批处理功能(--enable_batching)。

2.5 容器化部署

Docker封装模型服务,确保环境一致性和快速部署。Dockerfile示例:

dockerfile 复制代码
FROM tensorflow/serving:2.13.0
COPY ./model /models/my_model
ENV MODEL_NAME=my_model

然后使用docker run启动。

2.6 监控与日志

生产服务需记录:

  • 请求数量、延迟分布、错误率。
  • 模型版本、硬件资源使用。
    使用Prometheus + Grafana或云服务。

3. 环境搭建与工具配置

3.1 安装依赖

bash 复制代码
conda activate tf213
pip install fastapi uvicorn[standard] python-multipart pillow tensorflow-serving-api

安装Docker(可选,用于TensorFlow Serving容器)。

3.2 准备预训练模型

我们使用一个简单的MNIST CNN模型,训练并导出为SavedModel。

python 复制代码
import tensorflow as tf
from tensorflow import keras
from tensorflow.keras import layers

# 训练简单CNN
(x_train, y_train), (x_test, y_test) = keras.datasets.mnist.load_data()
x_train = x_train.reshape(-1,28,28,1).astype('float32') / 255.0
x_test = x_test.reshape(-1,28,28,1).astype('float32') / 255.0
model = keras.Sequential([
    layers.Conv2D(32, 3, activation='relu', input_shape=(28,28,1)),
    layers.MaxPooling2D(),
    layers.Conv2D(64, 3, activation='relu'),
    layers.Flatten(),
    layers.Dense(10, activation='softmax')
])
model.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy'])
model.fit(x_train, y_train, epochs=3, batch_size=128, validation_split=0.1)

# 导出为SavedModel,定义签名
@tf.function(input_signature=[tf.TensorSpec(shape=[None,28,28,1], dtype=tf.float32, name='input')])
def serve_fn(x):
    return {'output': model(x)}
model.save('mnist_model/1', signatures={'serving_default': serve_fn})

验证SavedModel:

bash 复制代码
saved_model_cli show --dir mnist_model/1 --all

4. 代码实战教学

4.1 使用FastAPI创建同步推理API

python 复制代码
# app.py
import tensorflow as tf
import numpy as np
from fastapi import FastAPI, File, UploadFile, HTTPException
from PIL import Image
import io

app = FastAPI(title="MNIST Classifier API")

# 加载模型
model = tf.saved_model.load('mnist_model/1')
infer = model.signatures['serving_default']

def preprocess_image(image_bytes):
    img = Image.open(io.BytesIO(image_bytes)).convert('L')
    img = img.resize((28,28))
    img_array = np.array(img, dtype=np.float32) / 255.0
    img_array = img_array.reshape(1,28,28,1)
    return img_array

@app.post("/predict")
async def predict(file: UploadFile = File(...)):
    try:
        contents = await file.read()
        input_tensor = preprocess_image(contents)
        output = infer(tf.constant(input_tensor))
        predictions = output['output'].numpy()[0]
        digit = int(np.argmax(predictions))
        confidence = float(np.max(predictions))
        return {"digit": digit, "confidence": confidence}
    except Exception as e:
        raise HTTPException(status_code=400, detail=str(e))

# 健康检查
@app.get("/health")
async def health():
    return {"status": "ok"}

if __name__ == "__main__":
    import uvicorn
    uvicorn.run(app, host="0.0.0.0", port=8000)

启动服务:

bash 复制代码
python app.py

测试请求:

bash 复制代码
curl -X POST -F "file=@test.png" http://localhost:8000/predict

4.2 异步批处理队列(高级)

为了充分利用GPU,可以实现请求队列,合并多个请求为一个批次。下面是一个简单的批处理服务器(使用asyncio队列)。

python 复制代码
import asyncio
import threading
import queue
import time
from fastapi import FastAPI, BackgroundTasks
from pydantic import BaseModel

app_batch = FastAPI()
request_queue = queue.Queue()
response_dict = {}

def batch_processor():
    """后台线程处理批次"""
    while True:
        time.sleep(0.1)  # 等待积累请求
        batch_requests = []
        while not request_queue.empty() and len(batch_requests) < 32:
            req = request_queue.get()
            batch_requests.append(req)
        if batch_requests:
            # 合并图像
            images = np.array([req['image'] for req in batch_requests])
            preds = infer(tf.constant(images))['output'].numpy()
            for req, pred in zip(batch_requests, preds):
                response_dict[req['id']] = {"digit": int(np.argmax(pred)), "confidence": float(np.max(pred))}

threading.Thread(target=batch_processor, daemon=True).start()

@app_batch.post("/predict_batch")
async def predict_batch(file: UploadFile):
    import uuid
    req_id = str(uuid.uuid4())
    contents = await file.read()
    input_tensor = preprocess_image(contents)[0]  # 去掉batch维
    request_queue.put({'id': req_id, 'image': input_tensor})
    # 轮询等待结果
    timeout = 5
    start = time.time()
    while req_id not in response_dict:
        if time.time() - start > timeout:
            raise HTTPException(status_code=504, detail="Timeout")
        await asyncio.sleep(0.01)
    result = response_dict.pop(req_id)
    return result

4.3 使用TensorFlow Serving部署

启动TensorFlow Serving Docker容器:

bash 复制代码
docker run -t --rm -p 8501:8501 -v "$(pwd)/mnist_model:/models/mnist" -e MODEL_NAME=mnist tensorflow/serving:2.13.0

然后使用REST API调用:

bash 复制代码
curl -d '{"instances": [[[[0.0]*28]*28]]}' -X POST http://localhost:8501/v1/models/mnist:predict

或者gRPC客户端:

python 复制代码
import grpc
import tensorflow as tf
from tensorflow_serving.apis import predict_pb2, prediction_service_pb2_grpc
import numpy as np

channel = grpc.insecure_channel('localhost:8500')
stub = prediction_service_pb2_grpc.PredictionServiceStub(channel)
request = predict_pb2.PredictRequest()
request.model_spec.name = 'mnist'
request.model_spec.signature_name = 'serving_default'
input_tensor = tf.make_tensor_proto(np.random.rand(1,28,28,1).astype(np.float32))
request.inputs['input'].CopyFrom(input_tensor)
response = stub.Predict(request, timeout=10.0)
print(response.outputs['output'])

4.4 使用Docker封装FastAPI服务

编写Dockerfile:

dockerfile 复制代码
FROM python:3.9-slim
WORKDIR /app
COPY requirements.txt .
RUN pip install --no-cache-dir -r requirements.txt
COPY . .
CMD ["uvicorn", "app:app", "--host", "0.0.0.0", "--port", "8000"]

requirements.txt:

复制代码
tensorflow==2.13.0
fastapi==0.100.0
uvicorn[standard]==0.23.0
pillow==9.5.0
python-multipart==0.0.6

构建并运行:

bash 复制代码
docker build -t mnist-api .
docker run -p 8000:8000 mnist-api

4.5 性能测试(使用locust)

安装locust:pip install locust编写locustfile.py

python 复制代码
from locust import HttpUser, task, between
import random

class MNISTUser(HttpUser):
    wait_time = between(1, 2)
    @task
    def predict(self):
        with open('sample.png', 'rb') as f:
            self.client.post("/predict", files={"file": f})

启动:locust -f locustfile.py --host=http://localhost:8000

5. 案例实操演练

案例:部署一个商品分类模型(ResNet50微调)并提供API

5.1 模型准备

假设我们已经微调了ResNet50,输出类别为5种商品。导出为SavedModel,输入尺寸224×224×3。

python 复制代码
# 导出签名
@tf.function(input_signature=[tf.TensorSpec(shape=[None,224,224,3], dtype=tf.float32)])
def predict_fn(x):
    return {'probs': model(x)}
tf.saved_model.save(model, 'product_model/1', signatures={'serving_default': predict_fn})

5.2 FastAPI服务增强

添加多类别返回、日志记录、请求限流(使用慢速)。

python 复制代码
from slowapi import Limiter, _rate_limit_exceeded_handler
from slowapi.util import get_remote_address

limiter = Limiter(key_func=get_remote_address)
app.state.limiter = limiter
app.add_exception_handler(429, _rate_limit_exceeded_handler)

@app.post("/predict")
@limiter.limit("10/minute")
async def predict(request: Request, file: UploadFile = File(...)):
    # ... 同前

5.3 部署到云服务(示例:AWS Lambda 不适合大模型,建议ECS)

本课不展开云部署,但可提及。

5.4 压力测试与结果分析

使用wrk或locust,记录QPS和延迟。

6. 常见坑点与排错总结

6.1 模型加载与签名

  • 坑1:加载SavedModel后直接调用模型失败,需使用签名。

    • 解决model = tf.saved_model.load(path); infer = model.signatures['serving_default']
  • 坑2:输入张量名称与签名不匹配,调用时出现KeyError。

    • 解决 :查看签名输入名称:print(infer.structured_input_signature)

6.2 性能问题

  • 坑3:单个请求推理慢,未使用批处理。

    • 解决:实现动态批处理或使用TensorFlow Serving内置批处理。
  • 坑4:GPU利用率低,服务成为CPU瓶颈。

    • 解决:确保预处理(图像解码、缩放)在CPU上高效,模型推理在GPU上;使用异步并发。

6.3 并发与稳定性

  • 坑5:高并发下模型加载多次(每个worker进程加载一份),内存溢出。

    • 解决 :使用全局单例模型,gunicorn使用--preload
  • 坑6:请求体过大导致内存爆炸。

    • 解决:限制上传文件大小,使用流式读取。

6.4 TensorFlow Serving坑点

  • 坑7:模型版本路径必须包含数字子目录,否则无法识别。
  • 坑8:gRPC客户端与REST端口混淆(REST:8501, gRPC:8500)。

7. 知识点总结 + 课后作业

7.1 核心知识点梳理

  • 模型导出:SavedModel + 签名定义。
  • 服务框架:FastAPI(高性能、异步)、TensorFlow Serving(专业模型服务)。
  • 性能优化:动态批处理、模型量化、并发控制。
  • 容器化:Docker打包服务,便于部署。
  • 监控:日志、健康检查、压力测试。

7.2 基础作业

  1. 使用FastAPI部署MNIST模型,并实现一个简单的前端页面调用API(可用HTML表单)。
  2. 使用TensorFlow Serving部署同一模型,通过REST API测试。
  3. 用locust对API进行并发测试,记录不同并发下的平均延迟。

7.3 进阶实操作业

任务:实现支持文本分类的服务(含预处理)

  • 训练或使用预训练文本分类模型(如IMDb情感分析)。
  • 导出模型(包含文本向量化层)。
  • 构建FastAPI服务,接收原始文本字符串,返回情感类别和概率。
  • 实现请求限流和批量处理(可选)。
  • 使用Docker部署,并编写docker-compose.yml管理服务和数据库(可选)。

7.4 思考拓展题

  1. 在分布式环境中,如何保证模型服务的高可用和负载均衡?

  2. 如果模型推理延迟很高(如1秒),如何设计API以提升用户体验(例如异步任务队列)?

  3. TensorFlow Serving与FastAPI + TF模型直接调用相比,有哪些优势和不足?


下一课预告:项目性能调优全方案------我们将学习训练和推理的性能调优技巧,包括数据加载优化、混合精度、XLA编译等,榨干硬件性能。


🔗《TensorFlow2.x: 深度学习入门到高阶实战教程》系列课程导航

去订阅

第一部分:基础入门(1-10 课)

第二部分:神经网络核心(11-25 课)

第三部分:进阶网络与框架高阶(26-40 课)

第四部分:企业实战与项目落地(41-50 课)
🌟 感谢您耐心阅读到这里!

💡 如果本文对您有所启发欢迎:

👍 点赞📌 收藏 📤 分享给更多需要的伙伴。

🗣️ 期待在评论区看到您的想法, 共同进步。

🔔 关注我,持续获取更多干货内容~

🤗 我们下篇文章见~

相关推荐
Xxtaoaooo1 小时前
AI 视频生成能不能真正跑通?MoneyPrinterTurbo 本地制作、远程访问与授权验证
人工智能·音视频·cpolar·ai视频·自动化短视频
Bruce_Liuxiaowei1 小时前
录屏片段无损合并:从 ffmpeg concat 原理到 Python 自动化脚本
python·ffmpeg·自动化
海浪仙人掌1 小时前
流动比率有哪些分析陷阱?流动比率怎么避开这些陷阱?
大数据·数据库·人工智能
宝桥南山2 小时前
DeepSeek - 尝试安装和使用一下DeepSeek Harness
人工智能·ai·aigc·ai编程
小海豚儿2 小时前
不是所有 Workflow,都值得升级成 Graph
人工智能·ai编程
X54先生(人文科技)2 小时前
Yuri演唱会碳硅协同思维推演备忘录
人工智能·深度学习·知识图谱·零知识证明
网易云信2 小时前
更强、更快、更普惠!网易智企帝王蟹率先接入 DeepSeek V4.1 Flash
人工智能·agent·deepseek
renzao_ai2 小时前
本地 35B 大模型部署实战:Ollama 跑 Ornith-35B 全流程
python·llama·免费ai大模型