对1D张量执行Pooling操作调用MaxPool1d报错如何解决?
1D张量池化操作报错解决方法
报错原因
PyTorch的nn.MaxPool1d要求输入张量维度为2或3维,对应[通道数, 序列长度]或[批量大小, 通道数, 序列长度]的格式,你传入的是形状为(768,)的1维张量,不符合输入维度要求,因此触发报错。
解决方法
你只需要给1维张量补充对应维度即可正常执行池化,修改后的可运行代码如下:
import numpy as np import torch import torch.nn as nn A = np.random.rand(768) # 定义池化层,核大小4步长4,长度缩短为原来的1/4 m = nn.MaxPool1d(4, 4) A_tensor = torch.from_numpy(A).float() # 补充通道维度,形状从(768,)变为(1, 768),符合2维输入要求 A_tensor = A_tensor.unsqueeze(0) output = m(A_tensor) # 去掉冗余的通道维度,得到缩短后的1维张量,形状为(192,) output_1d = output.squeeze()
批量场景适配
如果需要同时处理多条1D序列,可直接整理为[批量大小, 1, 序列长度]的3维格式传入池化层,输出结果批量维度会自动保留。
内容的提问来源于stack exchange,提问作者albusdemens
相关产品推荐
相关产品推荐

