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

如何无for循环将shape为[2,2,1,2]的PyTorch张量扩展为[2,4,1,2]

PyTorch张量维度转换:无循环实现[2,2,1,2]转[2,4,1,2]

需求说明

现有shape为torch.Size([2, 2, 1, 2])的张量a:

>>> print(a, a.shape)
tensor([[[[0.2955, 0.8836]],

         [[0.7607, 0.6657]]],


        [[[0.6779, 0.5109]],

         [[0.0785, 0.6564]]]]) torch.Size([2, 2, 1, 2])

需要将其转换为shape为torch.Size([2, 4, 1, 2])的张量b,其中第二维度的每个元素重复两次:

>>> print(b, b.shape)
tensor([[[[0.2955, 0.8836]],

         [[0.2955, 0.8836]],

         [[0.7607, 0.6657]],

         [[0.7607, 0.6657]]],


        [[[0.6779, 0.5109]],

         [[0.6779, 0.5109]],

         [[0.0785, 0.6564]],

         [[0.0785, 0.6564]]]]) torch.Size([2, 4, 1, 2])

无循环实现方法

方法1:使用torch.repeat_interleave(推荐)

torch.repeat_interleave可以直接在指定维度上对每个元素重复指定次数,是最直观的解决方案:

import torch

# 构造示例张量
a = torch.tensor([[[[0.2955, 0.8836]],
                   [[0.7607, 0.6657]]],
                  [[[0.6779, 0.5109]],
                   [[0.0785, 0.6564]]]])

# 转换张量
b = torch.repeat_interleave(a, repeats=2, dim=1)

# 验证结果
print(b.shape)  # 输出: torch.Size([2, 4, 1, 2])
print(b)
  • 参数说明:repeats=2表示每个元素重复2次,dim=1指定在第二维度(索引从0开始)执行重复操作。

方法2:结合unsqueeze、expand与flatten

如果需要用expand实现,可以先插入临时维度再展平:

# 转换张量
b = a.unsqueeze(2).expand(-1, -1, 2, -1, -1).flatten(1, 2)

# 验证结果
print(b.shape)  # 输出: torch.Size([2, 4, 1, 2])
  • 步骤解析:
    1. unsqueeze(2):在第二维度后插入一个新维度,shape变为[2, 2, 1, 1, 2]
    2. expand(-1, -1, 2, -1, -1):将新插入的维度扩展为2,shape变为[2, 2, 2, 1, 2]
    3. flatten(1, 2):将第1、第2维度展平合并,得到目标shape[2, 4, 1, 2]

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.22 22:27:19