如何将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
相关产品推荐
相关产品推荐

