You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何将TensorFlow鸟类识别项目封装为可POST请求的接口

实现鸟类识别模型的POST接口封装

一、环境准备

先安装所需依赖库:

pip install tensorflow tensorflow-hub flask pillow

若选择更高效的FastAPI框架,额外安装:

pip install fastapi uvicorn

二、Flask版本实现(适合新手)

以下代码完成模型加载、图片预处理、POST接口定义全流程:

import tensorflow as tf
import tensorflow_hub as hub
from flask import Flask, request, jsonify
from PIL import Image
import numpy as np

app = Flask(__name__)

# 加载TF Hub鸟类识别模型
model = hub.load("https://tfhub.dev/google/aiy/vision/classifier/birds_V1/1")
# 加载标签文件(需提前准备,每行对应一个鸟类物种名称,顺序与模型输出索引匹配)
with open("birds_labels.txt", "r") as f:
    labels = [line.strip() for line in f.readlines()]

def preprocess_image(image):
    # 转换为模型要求的224x224尺寸
    image = image.resize((224, 224))
    # 归一化像素值至0-1区间
    img_array = np.array(image) / 255.0
    # 添加batch维度(模型输入格式要求:[batch_size, height, width, channels])
    return np.expand_dims(img_array, axis=0)

@app.route('/predict-bird', methods=['POST'])
def predict_bird():
    if 'file' not in request.files:
        return jsonify({"error": "未上传图片文件"}), 400
    
    file = request.files['file']
    if file.filename == '':
        return jsonify({"error": "未选择图片文件"}), 400
    
    try:
        # 读取图片并转换为RGB格式
        image = Image.open(file.stream).convert('RGB')
        # 预处理图片
        processed_img = preprocess_image(image)
        # 模型预测
        predictions = model(processed_img)
        # 获取置信度最高的物种索引
        top_idx = np.argmax(predictions[0])
        # 提取对应物种名称和置信度
        bird_name = labels[top_idx]
        confidence = float(predictions[0][top_idx])
        
        return jsonify({
            "bird_species": bird_name,
            "confidence": round(confidence, 4)
        })
    except Exception as e:
        return jsonify({"error": f"处理失败: {str(e)}"}), 500

if __name__ == '__main__':
    app.run(host='0.0.0.0', port=5000, debug=False)

三、FastAPI版本实现(高效现代框架)

若追求更高性能和自动API文档,可使用FastAPI:

import tensorflow as tf
import tensorflow_hub as hub
from fastapi import FastAPI, File, UploadFile
from PIL import Image
import numpy as np
from fastapi.responses import JSONResponse

app = FastAPI()

model = hub.load("https://tfhub.dev/google/aiy/vision/classifier/birds_V1/1")
with open("birds_labels.txt", "r") as f:
    labels = [line.strip() for line in f.readlines()]

def preprocess_image(image):
    image = image.resize((224, 224))
    img_array = np.array(image) / 255.0
    return np.expand_dims(img_array, axis=0)

@app.post("/predict-bird")
async def predict_bird(file: UploadFile = File(...)):
    try:
        image = Image.open(file.file).convert('RGB')
        processed_img = preprocess_image(image)
        predictions = model(processed_img)
        top_idx = np.argmax(predictions[0])
        bird_name = labels[top_idx]
        confidence = float(predictions[0][top_idx])
        
        return {
            "bird_species": bird_name,
            "confidence": round(confidence, 4)
        }
    except Exception as e:
        return JSONResponse(status_code=500, content={"error": f"处理失败: {str(e)}"})

运行命令:

uvicorn main:app --host 0.0.0.0 --port 5000 --reload

四、接口测试

使用curl命令快速测试:

curl -X POST -F "file=@your_bird_photo.jpg" http://localhost:5000/predict-bird

也可通过Postman等工具,选择POST请求,在form-data中设置key为file,上传本地鸟类图片后发送请求,即可获取识别结果。

关键注意事项

  • 标签文件birds_labels.txt的顺序必须与模型输出的类别索引严格对应,否则会返回错误的物种名称。
  • 模型首次加载会自动下载权重文件,耗时取决于网络速度。
  • 生产环境部署时,建议关闭调试模式,添加图片格式(仅允许JPG/PNG)、大小限制的校验逻辑,避免恶意请求。

内容的提问来源于stack exchange,提问作者Hazza

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.13 06:15:30