如何用for循环转换PyTorch张量列表中的张量类型?
如何用for循环实现PyTorch张量从DoubleTensor到FloatTensor的类型转换?
你尝试将张量从DoubleTensor转换为FloatTensor,但循环内显示转换成功,原张量却未改变。问题出在float()方法的特性以及循环中的赋值逻辑。
你的代码
import itertools import torch # 假设这里已定义xq_train_tensor等DoubleTensor类型的张量 train_tensorset = [xq_train_tensor,y_train_tensor] val_tensorset = [xq_val_tensor,y_val_tensor] test_tensorset = [xq_test_tensor,y_test_tensor] tensor_list = [train_tensorset,val_tensorset,test_tensorset] flat_tensor_list = list(itertools.chain.from_iterable(tensor_list)) print(f"Num tensors:{len(flat_tensor_list)}") for i, tensor in enumerate(flat_tensor_list): tensor = flat_tensor_list[i].float() print(f"{i}: {tensor.type()}") print(f"xq_train_tensor: {xq_train_tensor.type()}")
代码输出
Num tensors:6 0: torch.FloatTensor 1: torch.FloatTensor 2: torch.FloatTensor 3: torch.FloatTensor 4: torch.FloatTensor 5: torch.FloatTensor xq_train_tensor: torch.DoubleTensor
问题原因
float()方法会创建一个新的FloatTensor对象返回,而不是原地修改原张量。循环中的tensor = flat_tensor_list[i].float()只是让局部变量tensor指向了新张量,但原列表flat_tensor_list中的元素以及原始张量(如xq_train_tensor)并没有被更新。
解决方案
方案1:使用原地转换方法(推荐)
PyTorch提供了后缀带下划线的原地操作方法float_(),会直接修改原张量的类型,无需创建新对象:
for i, tensor in enumerate(flat_tensor_list): flat_tensor_list[i].float_() # 原地转换类型,直接修改原张量 print(f"{i}: {flat_tensor_list[i].type()}") # 验证原张量类型 print(f"xq_train_tensor: {xq_train_tensor.type()}")
运行后原张量xq_train_tensor的类型会变为torch.FloatTensor。
方案2:更新列表元素并同步原变量
如果不想原地修改,可将新生成的FloatTensor赋值回列表,再重新拆分并更新原变量:
# 循环生成新张量并更新列表 for i in range(len(flat_tensor_list)): flat_tensor_list[i] = flat_tensor_list[i].float() print(f"{i}: {flat_tensor_list[i].type()}") # 重新拆分回原数据集列表 train_tensorset = flat_tensor_list[0:2] val_tensorset = flat_tensor_list[2:4] test_tensorset = flat_tensor_list[4:6] # 更新原张量变量 xq_train_tensor, y_train_tensor = train_tensorset xq_val_tensor, y_val_tensor = val_tensorset xq_test_tensor, y_test_tensor = test_tensorset # 验证类型 print(f"xq_train_tensor: {xq_train_tensor.type()}")
内容的提问来源于stack exchange,提问作者D Danne
相关产品推荐
相关产品推荐

