如何将单个torch.Tensor转换为包含多个tensor元素的列表
PyTorch张量拆分为单元素张量列表的实现方法
假设你提到的P1是一维张量(形状为(60,)),拆分后每个列表元素为单个值的张量,可通过以下几种常用方式实现:
- 方法1:使用
torch.unbind()(最高效,官方推荐)unbind会按指定维度直接将张量拆分为张量元组,转列表即可得到目标结果:
import torch # 示例构造60长度的一维张量 P1 = torch.randn(60) # 按第0维拆分后转列表 tensor_list = list(torch.unbind(P1)) # 验证结果 print(len(tensor_list)) # 输出 60 print(isinstance(tensor_list[0], torch.Tensor)) # 输出 True
如果是更高维度的张量,比如形状为(60, 3)需要拆分为60个形状为(3,)的张量,指定维度参数即可:list(torch.unbind(P1, dim=0))
- 方法2:列表推导式遍历
如果需要对拆分后的单个张量做自定义处理,这种写法更灵活:
tensor_list = [P1[i] for i in range(P1.size(0))]
该方法输出结果和unbind完全一致。
- 方法3:使用
torch.split()拆分
指定每个拆分块的大小为1即可,输出为张量元组,直接转列表即可:
# 拆分后每个张量形状为(1,),如果需要标量张量可额外加squeeze操作 tensor_list = list(torch.split(P1, split_size_or_sections=1)) # 若要得到标量张量写法如下 tensor_list = [t.squeeze(0) for t in torch.split(P1, 1)]
注意:如果原始张量存储在GPU上,拆分得到的所有张量也会保留在GPU显存中,无需额外做设备迁移操作。
内容的提问来源于stack exchange,提问作者Saran Zeb
相关产品推荐
相关产品推荐

