如何无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])
- 步骤解析:
unsqueeze(2):在第二维度后插入一个新维度,shape变为[2, 2, 1, 1, 2]expand(-1, -1, 2, -1, -1):将新插入的维度扩展为2,shape变为[2, 2, 2, 1, 2]flatten(1, 2):将第1、第2维度展平合并,得到目标shape[2, 4, 1, 2]
内容的提问来源于stack exchange,提问作者Vishak Raj
相关产品推荐
相关产品推荐

