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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 14:27:05