如何高效将含嵌套列表的Pandas DataFrame转为PyTorch TensorDataset?
解决方法
针对这种包含嵌套列表列的DataFrame,高效转换为TensorDataset的关键是先将嵌套列表列转换为统一的二维numpy数组,避免触发警告或报错,具体步骤如下:
处理嵌套列表列(列A):
由于列A中每个元素都是长度相同的列表,直接用numpy.array()将其转换为二维numpy数组,再转为PyTorch张量,这会比从列表转张量高效得多,也不会触发警告。处理标量列(列B):
直接提取其numpy数组后转为张量即可。组合为TensorDataset:
将处理好的特征张量和标签张量传入TensorDataset。
完整代码
import pandas as pd import numpy as np import torch from torch.utils.data import TensorDataset # 原始DataFrame df = pd.DataFrame({'A': [[1, 2, 3], [1, 2, 3], [1, 2, 3]], 'B': [0, 1, 0]}) # 高效转换 A_tensor = torch.from_numpy(np.array(df['A'].values)) B_tensor = torch.from_numpy(df['B'].values) # 创建TensorDataset dataset = TensorDataset(A_tensor, B_tensor)
为什么之前的方法会有问题?
- 第一种方法用
df['A'].values.tolist()转列表再创建张量,因为df['A'].values是object类型的numpy数组(每个元素是独立的小数组),转列表后再创建张量会逐个处理元素,效率极低,因此触发警告。 - 第二种方法直接用
df.to_numpy()得到的是object类型的二维数组(混合了列表和标量),PyTorch无法直接将这种类型的数组转为张量,所以报错。
内容的提问来源于stack exchange,提问作者Broxy
相关产品推荐
相关产品推荐

