如何在不分离、不转Numpy的情况下将Tensor标量转为Tensor列表
PyTorch张量标量转列表形式张量问题解答
问题1:将张量标量转换为张量列表且不执行detach操作
这里的“张量列表”分为两种场景,对应实现方式如下:
- 若要生成**单元素形状的张量(如shape为(1,))**且保留计算图:
使用unsqueeze扩展维度,直接调用标量张量的方法:
或使用全局函数:# x 是标量张量 tensor_list = x.unsqueeze(0)
两种方式均不会触发tensor_list = torch.unsqueeze(x, 0)detach,完整保留梯度追踪能力。 - 若要生成包含标量张量的Python列表:
直接将标量张量放入列表即可:
此操作不会对原张量做任何修改或detach。tensor_list = [x]
问题2:基于long类型标量张量生成单元素张量(无需转Numpy)
对于x = tensor.long()类型的标量,要得到类似tensor[(x.detached().numpy())]的单元素张量,无需依赖Numpy,直接用维度扩展方法即可:
# 生成shape为(1,)的long类型张量,保留计算图 result_tensor = x.unsqueeze(0)
如果不需要保留梯度,也可以用reshape实现相同形状转换:
result_tensor = x.reshape(1)
两种方式生成的张量与目标效果完全一致,全程不涉及Numpy转换,效率更高。
内容的提问来源于stack exchange,提问作者Tommy Yu
相关产品推荐
相关产品推荐

