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,有问题时可直接看到主进程异常 |
完整的 CSV 图片数据集
跳转到“完整的 CSV 图片数据集”保存为 dataset_demo.py 后直接运行。使用 torchvision 0.23.0、PyTorch 2.8.0 CPU;--workers 2 可以另测多进程读取。Dataset 类位于模块顶层,程序入口由 if __name__ == "__main__" 保护,便于 Windows 的 spawn 模式导入。
import argparseimport csvfrom pathlib import Pathfrom tempfile import TemporaryDirectoryimport torchfrom torch.utils.data import Dataset, DataLoaderfrom torchvision.io import decode_image, ImageReadMode, write_pngfrom 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 torchfrom 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 参考