如何使用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
相关产品推荐
相关产品推荐

