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

如何使用PyTorch view保留张量维度的分组结构?

如何在PyTorch张量重塑时保留指定维度的分组?

核心思路是:先把包含目标分组大小的维度移动到最后一位,再将前面所有维度展平,就能把所有目标大小的分组整合到一起。

针对两种具体场景的实现

场景1:张量形状为(4,5,2,3)

这里5是第2个维度(索引从0开始为1),先通过permute调整维度顺序,把5所在的维度移到最后,再展平前面的维度:

import torch
x = torch.randn(4,5,2,3)
# 调整维度顺序:原维度0、2、3保留,把维度1(对应5)移到最后
x_reshaped = x.permute(0, 2, 3, 1).view(-1, 5)
# 最终形状为 (4*2*3, 5) = (24, 5)

场景2:张量形状为(4,2,5,3)

这里5是第3个维度(索引从0开始为2),同样调整维度顺序后展平:

x = torch.randn(4,2,5,3)
# 调整维度顺序:原维度0、1、3保留,把维度2(对应5)移到最后
x_reshaped = x.permute(0, 1, 3, 2).view(-1, 5)
# 最终形状为 (4*2*3, 5) = (24, 5)

通用化实现(适配任意维度位置)

如果不想手动写维度顺序,可以写个简单函数自动处理:

def gather_target_groups(x, target_group_size):
    # 找到目标分组大小对应的维度索引
    target_dim = x.shape.index(target_group_size)
    # 构造新的维度顺序:除目标维度外,其他维度保持原顺序,最后追加目标维度
    new_dim_order = [i for i in range(len(x.shape)) if i != target_dim] + [target_dim]
    # 调整维度并展平
    return x.permute(new_dim_order).view(-1, target_group_size)

# 测试场景1
x1 = torch.randn(4,5,2,3)
print(gather_target_groups(x1, 5).shape)  # 输出 torch.Size([24, 5])

# 测试场景2
x2 = torch.randn(4,2,5,3)
print(gather_target_groups(x2, 5).shape)  # 输出 torch.Size([24, 5])

内容的提问来源于stack exchange,提问作者JobHunter69

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 03:45:01