PyTorch打包序列输入模型绘制TensorBoard图时参数报错求助
TensorBoard绘制含打包序列输入的RNN模型图报错问题
问题描述
我定义了一个以打包序列(packed sequence)为输入的神经网络,其forward函数实现如下:
def forward(self, x): y, _ = self.rnn(x) output ,lengths = torch.nn.utils.rnn.pad_packed_sequence(y, batch_first = True) out = [output[e, i-1,:].unsqueeze(0) for e, i in enumerate(lengths)] out = torch.cat(out, dim = 0) return out
随后尝试使用TensorBoard绘制该网络的模型图,代码如下:
tb = torch.utils.tensorboard.SummaryWriter() tb.add_graph(model, input) # input is a packed sequence tensor
运行后出现错误:
TypeError: MyRNN.forward() takes 2 positional arguments but 5 were given
想请教该错误的原因,以及是否可以使用打包序列完成模型图的绘制?
错误原因
add_graph在解析模型计算图时,会对传入的输入对象执行自动解包操作。PackedSequence内部包含tensor、batch_sizes、sorted_indices、unsorted_indices四个核心属性,加上模型实例本身(self),总共会向forward传递5个参数,但你的forward方法只定义了self和x两个参数,因此触发参数不匹配的报错。
解决方案与可行性说明
可以使用打包序列完成模型图绘制,只需要调整输入的传递方式:
- 将
PackedSequence用元组(或列表)包裹后传入add_graph,这样TensorBoard会将整个包裹容器作为单个参数传递,不会解包内部属性:tb = torch.utils.tensorboard.SummaryWriter() tb.add_graph(model, (input,)) # 用元组打包输入,避免自动解包
这种方式下,forward方法的x参数会正确接收完整的PackedSequence对象,模型图绘制可以正常执行。
内容的提问来源于stack exchange,提问作者Tadgh Kelly
相关产品推荐
相关产品推荐

