如何在Django项目中引入.h5格式TensorFlow预训练模型?
在Django中加载TensorFlow .h5预训练模型的实现方法
核心思路
避免在每次请求时重复加载模型(会严重影响性能),建议在Django启动时全局加载模型,之后在视图中复用已加载的模型实例处理请求。
具体实现步骤
1. 放置模型文件
在你的Django应用(比如myapp)下创建models目录,将.h5预训练模型文件放入该目录,路径示例:myapp/models/ai_model.h5。
2. 全局加载模型
在views.py中通过全局函数管理模型实例,确保只加载一次:
import tensorflow as tf from django.conf import settings import os import numpy as np # 全局变量存储模型实例 _loaded_model = None def get_pretrained_model(): global _loaded_model if _loaded_model is None: # 拼接模型文件的绝对路径 model_path = os.path.join(settings.BASE_DIR, 'myapp', 'models', 'ai_model.h5') # 加载.h5模型 _loaded_model = tf.keras.models.load_model(model_path) return _loaded_model
注:Django开发服务器自动重载时可能会重复加载模型,生产环境(如Gunicorn/uWSGI)下无此问题,若需规避开发环境重复加载,可结合
sys.modules判断是否为首次加载。
3. 在视图中使用模型处理请求
编写视图函数接收前端请求,调用模型完成预测:
from django.http import JsonResponse def ai_predict(request): if request.method != 'POST': return JsonResponse({'error': '仅支持POST请求'}, status=405) # 解析前端传来的输入数据(示例为表单格式,可根据实际调整为JSON) try: input_data = request.POST.get('input_features') # 数据预处理:转换为模型所需的输入格式(需与训练时的预处理逻辑一致) processed_data = np.array([float(val) for val in input_data.split(',')]).reshape(1, -1) except ValueError: return JsonResponse({'error': '输入数据格式错误'}, status=400) # 获取已加载的模型 model = get_pretrained_model() # 执行预测 prediction_result = model.predict(processed_data).tolist() # 返回预测结果 return JsonResponse({'prediction': prediction_result})
4. 配置URL路由
在应用的urls.py中添加路由映射:
from django.urls import path from . import views urlpatterns = [ path('api/predict/', views.ai_predict, name='ai_predict'), ]
注意事项
- 依赖安装:确保环境中已安装
tensorflow和h5py,执行pip install tensorflow h5py即可。 - 权限设置:确保服务器进程拥有模型文件的读取权限。
- 数据一致性:前端输入数据的预处理逻辑必须与模型训练时完全一致(如归一化、维度匹配),否则预测结果会失真。
内容的提问来源于stack exchange,提问作者Gyuyeon Cho
相关产品推荐
相关产品推荐

