You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch中如何将random_split生成的Subset转为张量/矩阵?

解决PyTorch Subset无法用列名提取特征标签的问题

问题原因

torch.utils.data.random_split()返回的Subset是PyTorch的数据集包装类,它仅支持整数/切片索引(比如train[0]取第一个样本),不支持pandas DataFrame的列名字符串索引,所以直接用train[['Sepal.Length', ...]]会触发TypeError。

解决方案

根据你原始iris数据集的类型,选择对应方法:

方法一:如果原始iris是pandas DataFrame(推荐)

直接拆分索引,再通过索引提取DataFrame子集,保留DataFrame的列操作能力:

import torch
import pandas as pd

# 拆分索引(range(len(iris))生成所有行的索引)
train_idx, test_idx = torch.utils.data.random_split(
    range(len(iris)), 
    [112, 38], 
    generator=torch.Generator().manual_seed(42)
)

# 通过iloc提取子集
train_df = iris.iloc[train_idx.indices]
test_df = iris.iloc[test_idx.indices]

# 正常提取特征和标签
train_X = train_df[['Sepal.Length', 'Sepal.Width', 'Petal.Length', 'Petal.Width']]
train_y = train_df.Species

test_X = test_df[['Sepal.Length', 'Sepal.Width', 'Petal.Length', 'Petal.Width']]
test_y = test_df.Species

方法二:如果原始iris是PyTorch Dataset

遍历Subset中的每个样本,拼接成张量:

import torch

# 提取训练集特征和标签(假设每个样本是(特征张量, 标签)格式)
train_X = torch.stack([sample[0] for sample in train])
train_y = torch.tensor([sample[1] for sample in train])

# 提取测试集特征和标签
test_X = torch.stack([sample[0] for sample in test])
test_y = torch.tensor([sample[1] for sample in test])

如果你的自定义Dataset返回的是DataFrame行对象,也可以在遍历的时候直接取列:

train_X = pd.DataFrame([sample[['Sepal.Length', 'Sepal.Width', 'Petal.Length', 'Petal.Width']] for sample in train])
train_y = pd.Series([sample['Species'] for sample in train])

内容的提问来源于stack exchange,提问作者Emrecan Serin

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.14 03:50:48