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

如何在Django中加载.tar格式PyTorch权重实现推理并展示预测结果

1. PyTorch相关文件存放位置

在Django项目根目录下新建ml文件夹,专门存放所有机器学习相关内容,内部结构参考如下:

你的Django项目根目录/
├── 已有的Django应用目录/
├── ml/
│   ├── __init__.py  # 让Python将该目录识别为可导入的包
│   ├── model_def.py # 存放PyTorch模型类定义,代码必须和训练时的模型结构完全一致
│   ├── weights/     # 权重存放目录,建议加入.gitignore避免大文件上传到代码仓库
│   │   └── your_model_weights.tar
│   └── utils.py     # 存放预测函数、数据预处理、后处理相关工具代码
2. 核心推理实现逻辑

2.1 服务启动时预加载模型和权重

不要在每次用户请求时重复加载模型,会严重拖慢响应速度,建议在Django服务启动时一次性完成模型加载,操作如下:
在ml/__init__.py中写入加载逻辑:

import torch
import os
from .model_def import YourModelClass # 替换为你自己定义的模型类

# 全局变量存储加载好的模型和运行设备
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = None

def load_model():
    global model
    # 初始化模型结构,参数要和训练时的初始化参数完全一致
    model = YourModelClass()
    # 拼接权重文件绝对路径,避免相对路径报错
    weight_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'weights/your_model_weights.tar')
    checkpoint = torch.load(weight_path, map_location=device)
    # 加载权重,key根据你自己保存tar时的设置调整
    model.load_state_dict(checkpoint['model_state_dict'])
    model.to(device)
    model.eval()

然后在Django项目同名目录下的wsgi.py(如果用异步部署就修改asgi.py)中添加代码,让服务启动时自动调用加载函数:

# 在原有代码基础上添加以下两行
import ml
ml.load_model()

2.2 封装预测接口

在ml/utils.py中整合你的预测逻辑,补充数据预处理步骤,适配用户输入:

import torch
from . import model, device
# 导入你自己用的数据处理、构造Dataloader的相关依赖

def predict(raw_input_data):
    # 第一步:将用户传入的原始数据处理为模型需要的Dataloader格式
    # 此处替换为你自己的预处理逻辑,生成solute_graphs、solvent_graphs等输入
    loader = your_custom_data_process(raw_input_data)
    
    # 第二步:调用原有预测逻辑
    total_preds = torch.Tensor()
    with torch.no_grad():
        for solute_graphs, solvent_graphs, solute_lens, solvent_lens in loader:
            outputs, i_map = model(
                [solute_graphs.to(device), solvent_graphs.to(device), torch.tensor(solute_lens).to(device),
                torch.tensor(solvent_lens).to(device)]
            )   
            total_preds = torch.cat((total_preds, outputs.cpu()), 0)
    # 转成普通列表方便Django返回给前端
    return total_preds.numpy().flatten().tolist()

2.3 Django视图调用预测逻辑

在你需要提供预测能力的Django应用的views.py中写接口,接收用户输入并返回预测结果:

from django.http import JsonResponse
from ml.utils import predict

def predict_view(request):
    if request.method == 'POST':
        # 接收用户传参,根据前端传参格式调整,文件上传请用request.FILES获取
        raw_input = request.POST.get('input_data')
        pred_result = predict(raw_input)
        return JsonResponse({'code': 200, 'result': pred_result})
    return JsonResponse({'code': 400, 'msg': '仅支持POST请求'})

最后将该视图路由添加到对应应用的urls.py中即可正常访问。

3. 注意事项
  • 生产环境部署时建议用Gunicorn等WSGI服务器,设置worker数量为1,避免多个worker重复加载模型占用过多内存
  • 输入数据的预处理逻辑必须和训练时的预处理逻辑完全一致,否则会出现预测结果偏差
  • 并发请求量高的场景,建议单独将推理部分拆为独立的FastAPI服务,Django直接调用该服务,避免PyTorch GIL问题影响性能

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 03:39:03