如何导出固定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
相关产品推荐
相关产品推荐

