如何从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
相关产品推荐
相关产品推荐

