PyTorch实现Ordinal Pooling神经网络函数维度不匹配报错解决问询
Ordinal Pooling 实现维度不匹配问题解决方案
报错原因
RuntimeError: The size of tensor a (4) must match the size of tensor b (16) at non-singleton dimension 2
上述报错的核心是权重张量和TopK输出张量的维度不匹配:
- 你当前的输入
x为卷积输出的标准四维张量,形状为[batch_size, 通道数, 高度, 宽度] - 在
dim=1维度取Top4后,top.values的形状为[batch_size, 4, 高度, 宽度] - 你重复后的权重形状仅为
[batch_size, 4],缺少高度、宽度两个空间维度,相乘时广播机制无法适配,触发维度冲突。
额外补充:你原代码中权重变量名拼写为wights(缺少字母e),如果未修正也会触发命名错误,下述代码已统一修正为weights。
修复代码
场景1:逐空间位置做序数池化(保留空间维度)
如果需要对每个空间位置的通道维度做Top4加权,保留输出的空间维度,用如下实现:
import torch def Ordinal_Pooling_NN(x): # 权重设置为和输入x同设备,避免CPU/GPU运行报错 weights = torch.tensor([0.6, 0.25, 0.10, 0.05], device=x.device) top = torch.topk(x, 4, dim = 1) # 扩展权重维度适配广播规则,无需重复空间维度,性能更优 weights = weights[None, :, None, None] result = torch.sum(weights * top.values, dim = 1) return result
场景2:全局序数池化(输出一维向量,无空间维度)
如果需要对单样本所有特征元素做全局Top4加权,输出为长度等于batch_size的一维向量,用如下实现:
import torch def Ordinal_Pooling_NN(x): weights = torch.tensor([0.6, 0.25, 0.10, 0.05], device=x.device) # 先把单样本的所有特征展平为一维 x_flat = x.flatten(start_dim=1) top = torch.topk(x_flat, 4, dim = 1) weights = weights[None, :].repeat(x.shape[0], 1) result = torch.sum(weights * top.values, dim = 1) return result
内容的提问来源于stack exchange,提问作者Driss AL
相关产品推荐
相关产品推荐

