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

如何从PyTorch DataLoader中提取数据与标签?

从ImageFolder数据集/DataLoader提取数据与标签

首先纠正你代码里的一个小错误:test_loader初始化时误用了training_dataset,应该改为test_dataset:

test_loader = torch.utils.data.DataLoader(test_dataset, batch_size=2, shuffle=True)

下面分两种场景给出提取方法:

1. 直接从ImageFolder数据集对象提取(更高效)

ImageFolder没有像CIFAR10FL那样直接暴露data和target属性,但可以通过以下方式获取完整的特征和标签:

方法1:遍历数据集收集

# 提取训练数据和标签
X_train = []
y_train = []
for img, label in training_dataset:
    X_train.append(img)
    y_train.append(label)
# 转换为张量格式
X_train = torch.stack(X_train)
y_train = torch.tensor(y_train)

# 测试集同理操作
X_test = []
y_test = []
for img, label in test_dataset:
    X_test.append(img)
    y_test.append(label)
X_test = torch.stack(X_test)
y_test = torch.tensor(y_test)

方法2:利用ImageFolder的内部属性

ImageFolder的samples属性存储了所有(图片路径, 标签)对,你可以结合已定义的transform手动加载处理图片:

from PIL import Image

X_train = []
y_train = []
for img_path, label in training_dataset.samples:
    img = Image.open(img_path).convert('RGB')
    # 应用预设的图像变换
    img_tensor = transform(img)
    X_train.append(img_tensor)
    y_train.append(label)
X_train = torch.stack(X_train)
y_train = torch.tensor(y_train)

2. 从DataLoader中提取数据与标签

DataLoader是批量迭代器,需要遍历所有批次并拼接数据:

# 提取训练集
X_train = []
y_train = []
for batch_imgs, batch_labels in train_set:
    X_train.append(batch_imgs)
    y_train.append(batch_labels)
# 拼接所有批次的数据
X_train = torch.cat(X_train, dim=0)
y_train = torch.cat(y_train, dim=0)

# 测试集同理操作
X_test = []
y_test = []
for batch_imgs, batch_labels in test_loader:
    X_test.append(batch_imgs)
    y_test.append(batch_labels)
X_test = torch.cat(X_test, dim=0)
y_test = torch.cat(y_test, dim=0)

注意:如果数据集规模很大,直接提取完整的X和Y会占用大量内存,这种情况下建议保留DataLoader的批量迭代方式进行训练。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 23:58:24