跳转到内容
新建笔记

PyTorch 数据集:CSV 图片、DataLoader 与批次边界

Dataset 负责一个样本,DataLoader 负责取样与组批

跳转到“Dataset 负责一个样本,DataLoader 负责取样与组批”

以图片分类为例,一个样本由图像和类别组成。对于实现 __len__、__getitem__ 的 map-style Dataset,DataLoader 再根据采样器选择索引、读取样本、调用 collate_fn 组成批次。数据集不必一次把所有图片读进内存。

原笔记使用两列 CSV、图片目录以及 transform / target_transform,但代码缺少 Dataset 导入、可用样本和批次边界检查。下面保留这一结构,补成可独立运行的本地实验:生成 5 张 RGB PNG,写入标签 CSV,读取后统一尺寸,再检查训练与测试加载器的输出。示例不下载数据,也不使用这 5 张合成图片评价分类效果。

边界本例约定
标签文件UTF-8 CSV,表头严格为 image,label
图片位置同一目录第一层的文件名,不接受子目录或路径分隔符
类别整数 0 <= label < num_classes
解码显式 RGB;本例仅接受 uint8 图像
输入变换统一为 float32 的 3×16×16,范围 0–1
标签变换可选;调用者负责维持后续任务要求的类型与范围
默认执行num_workers=0,有问题时可直接看到主进程异常

保存为 dataset_demo.py 后直接运行。使用 torchvision 0.23.0、PyTorch 2.8.0 CPU;--workers 2 可以另测多进程读取。Dataset 类位于模块顶层,程序入口由 if __name__ == "__main__" 保护,便于 Windows 的 spawn 模式导入。

import argparse
import csv
from pathlib import Path
from tempfile import TemporaryDirectory
import torch
from torch.utils.data import Dataset, DataLoader
from torchvision.io import decode_image, ImageReadMode, write_png
from torchvision.transforms import v2
class CustomImageDataset(Dataset):
def __init__(self, annotations_file, img_dir, num_classes,
transform=None, target_transform=None):
if type(num_classes) is not int or num_classes <= 0:
raise ValueError("num_classes must be a positive integer")
self.img_dir = Path(img_dir)
self.transform = transform
self.target_transform = target_transform
self.samples = []
with open(annotations_file, "r", encoding="utf-8", newline="") as stream:
reader = csv.DictReader(stream)
if reader.fieldnames != ["image", "label"]:
raise ValueError("expected CSV header: image,label")
for line, row in enumerate(reader, start=2):
if set(row) != {"image", "label"}:
raise ValueError(f"bad column count on CSV line {line}")
filename = row["image"]
if (not filename or filename in {".", ".."}
or "/" in filename or "\\" in filename
or ":" in filename):
raise ValueError(f"expected a plain filename on line {line}")
try:
label = int(row["label"])
except (TypeError, ValueError) as error:
raise ValueError(f"invalid integer label on line {line}") from error
if not 0 <= label < num_classes:
raise ValueError(f"label out of range on line {line}")
path = self.img_dir / filename
if not path.is_file():
raise FileNotFoundError(path)
self.samples.append((path, label))
if not self.samples:
raise ValueError("dataset is empty")
def __len__(self):
return len(self.samples)
def __getitem__(self, index):
path, label = self.samples[index]
image = decode_image(str(path), mode=ImageReadMode.RGB)
if image.dtype != torch.uint8:
raise ValueError("this dataset expects 8-bit images")
if self.transform is not None:
image = self.transform(image)
if self.target_transform is not None:
label = self.target_transform(label)
return image, label
def main(argv=None):
parser = argparse.ArgumentParser()
parser.add_argument("--workers", type=int, choices=[0, 2], default=0)
args = parser.parse_args(argv)
with TemporaryDirectory() as directory:
root = Path(directory)
rows = []
for index in range(5):
image = torch.full((3, 8 + index, 10 + index),
index * 40, dtype=torch.uint8)
image[0] = 255
filename = f"sample_{index}.png"
write_png(image, str(root / filename))
rows.append((filename, index % 2))
with (root / "labels.csv").open("w", encoding="utf-8", newline="") as stream:
writer = csv.writer(stream)
writer.writerow(["image", "label"])
writer.writerows(rows)
transform = v2.Compose([
v2.ToImage(),
v2.Resize((16, 16), antialias=True),
v2.ToDtype(torch.float32, scale=True),
])
dataset = CustomImageDataset(root / "labels.csv", root, 2,
transform=transform)
generator = torch.Generator().manual_seed(12)
train_loader = DataLoader(dataset, batch_size=2, shuffle=True,
generator=generator,
num_workers=args.workers)
test_loader = DataLoader(dataset, batch_size=2, shuffle=False,
num_workers=args.workers)
batches = list(train_loader)
assert [x.shape[0] for x, y in batches] == [2, 2, 1]
for features, labels in batches:
assert features.shape[1:] == (3, 16, 16)
assert features.dtype == torch.float32
assert 0 <= features.min().item() <= features.max().item() <= 1
assert labels.dtype == torch.int64
assert labels.shape == (features.shape[0],)
ordered_labels = torch.cat([y for x, y in test_loader])
assert ordered_labels.tolist() == [0, 1, 0, 1, 0]
assert sum(len(y) for x, y in batches) == len(dataset)
image, label = dataset[0]
display_rgb = image.permute(1, 2, 0).numpy()
assert display_rgb.shape == (16, 16, 3)
preview = (image * 255).round().clamp(0, 255).to(torch.uint8)
write_png(preview, str(root / "preview.png"))
assert (root / "preview.png").is_file()
print("batch sizes:", [len(y) for x, y in batches])
print("first sample label:", label)
print("workers:", args.workers)
if __name__ == "__main__":
main()

默认运行应得到批次大小 [2, 2, 1]。临时目录离开作用域后自动清理,包括 preview.png;若要留下预览,可在自己的项目中选择明确的输出目录。这里训练和测试加载器故意读取同一份合成数据,目的只是对比顺序与批次行为。真实实验必须先独立划分训练、验证、测试数据,防止数据泄漏。

transform is not None 表达“提供了一个变换”。不要用对象真假性代替存在性:用户定义的可调用对象可以同时具有假值语义。文件存在检查只提前发现明显问题;文件可能随后变化,解码也可能失败,因此不能声称在构造时已验证所有像素内容。

shuffle、最后一批与多进程

跳转到“shuffle、最后一批与多进程”
参数或操作含义注意点
shuffle=True每次建立遍历使用随机采样顺序不是先顺序读完全部批次才洗牌;小数据集的两次随机排列也可能相同
batch_size=64期望每批最多 64 个样本默认 drop_last=False,末批可能不足 64
drop_last=True丢弃不能组成完整批次的尾部小于一个批次的数据集可能产生零批次
shuffle=False默认顺序采样验证/测试通常更易追踪预测与原样本的对应关系
num_workers=0在主进程执行取样便于调试,并非性能一定最优
num_workers>0使用工作进程加载样本内存、启动成本、磁盘和变换开销都会影响收益
next(iter(loader))新建一次迭代并取首批多次这样写不是推进同一个迭代器;需连续取样时保存迭代器

当指定 sampler 或 batch_sampler 时,要遵守它们与 shuffle、batch_size 等参数的互斥规则。这个讨论针对普通 map-style 数据集;IterableDataset 的数据划分、多工作进程重复与 drop_last 边界另有约定,不能直接照搬。DataLoader 与多进程官方说明

Windows 多进程运行时,避免把 Dataset、collate_fn 定义在局部函数中,避免给工作进程传无法序列化的 lambda。若数据增强使用 Python random 或 NumPy 的随机数,还应在工作进程初始化时显式管理相应种子;只设置 torch.manual_seed 不代表整个多库、多设备流程完全可复现。可复现性说明

不同尺寸的样本不能默认堆叠

跳转到“不同尺寸的样本不能默认堆叠”

默认组批会尝试沿新批次维堆叠 Tensor。上例通过 Resize 统一形状;检测任务若需要保留原尺寸,可以让 collate_fn 返回列表,把调整尺寸与目标处理留给模型约定。

import torch
from torch.utils.data import DataLoader
def keep_as_lists(batch):
images, targets = zip(*batch)
return list(images), list(targets)
samples = [(torch.zeros(3, 8, 10), {"label": 0}),
(torch.ones(3, 12, 7), {"label": 1})]
loader = DataLoader(samples, batch_size=2, collate_fn=keep_as_lists)
images, targets = next(iter(loader))
assert [tuple(x.shape) for x in images] == [(3, 8, 10), (3, 12, 7)]
assert [y["label"] for y in targets] == [0, 1]
print("preserved two differently sized images")

这只是演示自定义组批,不是某个检测器要求的完整标签格式。RGB 可视化要转换为 H,W,C;灰度图再按单通道处理,原记录中的无条件 squeeze() 加 cmap="gray" 不能适用于所有彩色批次。若输入已经标准化,还需要先按同一均值、标准差反变换。视觉变换与图像布局、decode_image 参考