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

如何导出固定Batch Size的PyTorch模型至ONNX?

解决ONNX导出时固定Batch Size的问题

方法一:强制固定Batch Size为1

你的代码已经用unsqueeze(0)生成了batch size=1的输入张量,只需在导出时不启用动态维度,ONNX就会完全按照输入的形状固定模型的batch维度,后续只能用batch size=1的输入运行模型。

修改导出代码:

torch.onnx.export(
    bert, 
    example, 
    "model.onnx", 
    export_params=True, 
    opset_version=10, 
    do_constant_folding=True, 
    input_names=['text'], 
    output_names=['output']
    # 移除dynamic_axes参数,默认所有维度固定
)

这样导出的模型会把输入的第0维(batch维)固定为1,彻底消除警告。

方法二:修改模型支持动态Batch Size(同时消除警告)

如果后续需要用不同batch size运行模型,可以修改BERTGRUSentiment模型,将GRU的初始隐藏状态h0作为输入参数,而非在模型内部初始化:

import torch.nn as nn

class BERTGRUSentiment(nn.Module):
    def __init__(self, bert, hidden_dim, output_dim, n_layers, dropout):
        super().__init__()
        self.bert = bert
        self.gru = nn.GRU(
            bert.config.hidden_size, 
            hidden_dim, 
            num_layers=n_layers, 
            bidirectional=True, 
            batch_first=True, 
            dropout=dropout if n_layers > 1 else 0
        )
        self.fc = nn.Linear(hidden_dim * 2, output_dim)
        self.dropout = nn.Dropout(dropout)

    # 将h0作为输入参数传入
    def forward(self, text, h0=None):
        with torch.no_grad():
            embedded = self.bert(text)[0]
        
        batch_size = text.shape[0]
        # 当h0未传入时,根据当前batch size初始化
        if h0 is None:
            h0 = torch.zeros(
                self.gru.num_layers * 2, 
                batch_size, 
                self.gru.hidden_size
            ).to(text.device)
        
        output, hidden = self.gru(embedded, h0)
        hidden = self.dropout(torch.cat((hidden[-2,:,:], hidden[-1,:,:]), dim=1))
        return self.fc(hidden)

导出时需要同时传入输入文本和初始隐藏状态:

# 生成对应batch size的h0
h0 = torch.zeros(N_LAYERS * 2, example.shape[0], HIDDEN_DIM).to(device)

torch.onnx.export(
    bert, 
    (example, h0),  # 传入两个输入
    "model.onnx", 
    export_params=True, 
    opset_version=10, 
    do_constant_folding=True, 
    input_names=['text', 'h0'], 
    output_names=['output'],
    # 允许batch维度动态变化
    dynamic_axes={
        'text': {0: 'batch_size'}, 
        'h0': {1: 'batch_size'}, 
        'output': {0: 'batch_size'}
    }
)

总结

  • 仅需固定batch size时,用batch=1的输入导出,不设置dynamic_axes即可。
  • 需要动态batch size时,修改模型将GRU初始状态作为输入,导出时传入对应参数并配置动态维度。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 15:37:32