PyTorch替换列表中Tensor为不同形状Tensor时遇RuntimeError求助
问题:扩展Tensor列表中每个元素的通道数
我有一个包含多个单通道Tensor的列表t_list:
import torch t_list = [ torch.tensor([[1], [1], [1]]), torch.tensor([[1], [1], [1]]), torch.tensor([[1], [1], [1]]) ]
需要将其转换为每个元素为3通道(原1列数据+2列0)的Tensor列表:
[ torch.tensor([[1,0,0], [1,0,0], [1,0,0]]), torch.tensor([[1,0,0], [1,0,0], [1,0,0]]), torch.tensor([[1,0,0], [1,0,0], [1,0,0]]) ]
尝试以下代码后报错:
import torch z= torch.zeros(1,2) for i, item in enumerate(t_list): for ii, item2 in enumerate(item): unsqueezed = torch.unsqueeze(item2,0) cat1 = torch.cat((unsqueezed,z),-1) squeezed = torch.squeeze(cat1,0) t[i][ii] = squeezed
错误信息:
RuntimeError: expand(torch.FloatTensor{[5]}, size=[]): the number of sizes provided (0) must be greater or equal to the number of dimensions in the tensor (1)
解决方案
代码报错的核心原因:
- 尝试给形状为
[3,1]的原Tensor赋值1维张量(squeezed是[3]形状),维度不匹配引发错误; - 嵌套逐元素循环冗余,PyTorch支持批量操作,无需逐元素处理。
推荐两种简洁实现方式:
方法1:批量拼接零张量
import torch t_list = [ torch.tensor([[1], [1], [1]]), torch.tensor([[1], [1], [1]]), torch.tensor([[1], [1], [1]]) ] # 生成与列表中Tensor行数一致的零张量,形状为[3,2] zero_tensor = torch.zeros_like(t_list[0]).repeat(1, 2) # 遍历列表拼接原张量与零张量 result = [torch.cat((tensor, zero_tensor), dim=1) for tensor in t_list]
方法2:使用填充函数补零
import torch.nn.functional as F t_list = [ torch.tensor([[1], [1], [1]]), torch.tensor([[1], [1], [1]]), torch.tensor([[1], [1], [1]]) ] # 在最后一个维度(列方向)填充2个0 result = [F.pad(tensor, (0, 2), mode='constant', value=0) for tensor in t_list]
以上两种方法均可高效得到目标结果,规避了原代码的维度错误与冗余循环问题。
内容的提问来源于stack exchange,提问作者M.Z
相关产品推荐
相关产品推荐

