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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 17:47:26