深入理解PyTorch transforms:原理、实践与性能优化 📅 发布时间:2026/9/8 11:29:40 👁 浏览次数: 1. 为什么PyTorch要把数据预处理做成 transforms 这套体系先聊点实际的。我第一次接触 PyTorch 的 transforms 时最大的困惑不是它怎么用而是“为什么非要搞出这么一套东西”。毕竟 TensorFlow 那边有 tf.imageKeras 有 ImageDataGeneratorOpenCV 也能做图像处理PyTorch 为什么要单独设计一个 torchvision.transforms这个问题的答案其实藏在使用场景里。PyTorch 的训练流程通常是Dataset 负责取数据DataLoader 负责按 batch 取数据模型负责前向计算。而 transforms 恰好卡在 Dataset 和 DataLoader 之间——它负责把“原始数据”变成“模型能吃的数据”。这里说的“原始数据”可以是图片文件、PIL Image、numpy 数组甚至是一段文本、一个音频信号而“模型能吃的数据”则是固定 shape、固定数值范围、固定 dtype 的 tensor。如果没有 transforms你得在每个 Dataset 的getitem里手动写一大堆预处理逻辑读图、转 RGB、缩放、转 tensor、归一化……而且每个项目都得重写一遍。有了 transforms这套流程被抽象成了可组合、可复用的函数式模块。你要做的只是把预处理步骤“拼”起来像搭积木一样。它的核心价值有几个组合性强Compose 把多个变换串成流水线想加一个增强步骤、去掉一个步骤改一行代码就行。延迟执行transforms 的变换发生在取数据时而不是预处理整个数据集。这意味着你可以对同一份数据做不同的随机增强每个 epoch 看到的都是“新”数据。与 Dataset 解耦Dataset 只负责“拿数据”transforms 只负责“改数据”职责清晰调试方便。生态成熟torchvision 内置了常用的几十种变换加上第三方库比如 albumentations、Kornia基本覆盖了所有常见需求。我见过很多初学者直接跳过 transforms在 Dataset 里用 OpenCV 写完所有预处理结果代码又长又难维护换个数据集基本要重写。等你用熟 transforms 再回头看会发现它的设计是真的省心。顺便说一句PyTorch 从 1.x 到 2.xtransforms 的 API 基本保持稳定你不用担心版本升级导致大量代码重写。网上有些教程用的是老版 API比如 FiveCrop、RandomSizedCrop 这种已经改名或调整的变换用新版本时注意看下文档就行。2. 最常用的几个 transforms 背后的原理与踩坑点transforms 里有一批“常驻选手”几乎每个图像项目都会用到。很多人是别人怎么用自己就怎么用从没想过这些变换到底做了什么、为什么要做。这部分我把它们的原理和容易踩的坑说清楚。2.1 ToTensor 到底做了什么from torchvision import transforms transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ])这是最常见的写法但很多人对 ToTensor 的认知停留在“把 PIL Image 变成 tensor”。其实它做了三件事数据格式转换把 PIL Image 或 numpy 数组HWC 布局转成 tensorCHW 布局。数值范围缩放把 0~255 的像素值除以 255缩放到 0.0~1.0。dtype 转换转为 torch.float32。这里有个容易忽略的细节如果输入是 numpy 数组ToTensor 会先检查其 dtype。如果输入已经是 uint8除以 255 后是 float但如果你传入的是 float32 的 numpy 数组且值域是 0~255ToTensor 不会帮你缩放它只会执行除 255。为什么因为 ToTensor 的源码里有一个判断只有输入为 uint8、int8、int16 等整数类型时才会除以 255。如果你的 numpy 图像是 float32值域 0~255直接 ToTensor 得到的是 0~255 的 tensor后面 Normalize 就会出问题。我自己就踩过这个坑。用 cv2.imread 读图时默认是 uint8没问题但有一次我先做了些 float 运算图像变成了 float32直接进 ToTensor结果模型训练 loss 异常排查了半天才发现数值范围不对。经验如果自己做预处理确保传给 ToTensor 的图像是 uint8 类型或者手动除以 255。2.2 Resize 的插值方式选择Resize 看起来没什么技术含量但插值方式的选择在特定场景下能带来明显差异。torchvision 的 Resize 默认用双线性插值bilinear这对大多数自然图像是合适的。但如果你做的是医学图像分割、卫星图像分割这类对边缘细节敏感的任务改用最近邻插值可以避免边缘模糊而做图像超分辨率时经常用 bicubic 插值。# 最近邻插值适合分割任务 transforms.Resize((512, 512), interpolationtransforms.InterpolationMode.NEAREST) # 双三次插值适合超分任务 transforms.Resize((256, 256), interpolationtransforms.InterpolationMode.BICUBIC)还有一个细节Resize 接受两种输入方式一是(h, w)元组二是单个 int。单 int 时表示短边缩放到该值长边按比例缩放保持宽高比。但注意这里的“保持宽高比”并不是说输出一定是(h, w)那种固定尺寸而是短边等于你给的值长边等比缩放后取整。后续如果还需要固定尺寸得再配合 CenterCrop 使用。2.3 Normalize 的参数为什么是 0.485、0.456、0.406这三个数在很多代码里直接复制粘贴但知道它们来历的人不多。它们其实是 ImageNet 数据集的 RGB 三个通道的均值mean和标准差std。使用 ImageNet 预训练模型时输入数据必须用这些数做标准化否则预训练权重就“不认识”你的输入分布。有人问我我自己训练一个新模型也需要用这三个数吗如果从零训练这三个数不是必须的你可以统计自己数据集的均值和标准差用它替换。但如果你要用 ImageNet 预训练模型做迁移学习这三个数必须保留。还有个折中做法即使从零训练只要数据量够大用 ImageNet 的均值和标准差也不会有什么坏处很多研究就是这么干的省去统计的麻烦。Normalize 的公式是(x - mean) / std经过这步处理后原来的 0~1 范围会被映射到大约 -2~2 的范围不同通道的分布被拉齐到接近标准正态分布。这对模型训练的稳定性是有帮助的尤其是使用 BN 层时可以让初始激活不那么极端。注意一点Normalize 必须在 ToTensor 之后调用因为 ToTensor 先把像素缩放到 0~1然后用 ImageNet 均值做标准化才能得到正确的数值范围。如果顺序反了等于对 0~255 的数据直接减去 0.485结果完全不对。2.4 RandomResizedCrop数据增强的“万金油”如果说让我只保留一个数据增强变换我会选 RandomResizedCrop。它在训练阶段做了一件很关键的事随机裁剪图像的不同区域然后缩放到固定尺寸。这模拟了不同尺度、不同位置的观察视角也是一种尺度不变性的近似。transforms.RandomResizedCrop(size(224, 224), scale(0.08, 1.0), ratio(0.75, 1.333))参数含义size输出尺寸。scale裁剪面积占原图面积的比例范围。默认是 0.08~1.0即最小可以裁剪原图 8% 的区域。ratio裁剪区域的宽高比范围默认 0.75~1.33允许一定程度的横竖变化。实际使用时如果目标对象比较小可以把 scale 的下限调高一点避免裁剪到太多背景如果目标对象的宽高比很固定比如车牌ratio 的范围可以缩小。推理阶段要用 CenterCrop因为推理是确定性的不需要随机性。很多人把这个切换忘了结果训练效果很好推理结果波动大那就是因为在推理时还在用随机裁剪。3. Compose 的执行机制与自定义 transforms 的注意事项3.1 Compose 是按顺序执行的函数管道Compose 的实现非常直白把所有 transforms 按传入顺序存到一个 list 里调用时依次对数据执行。它的核心逻辑就相当于def compose(transforms_list, data): for t in transforms_list: data t(data) return data虽然代码简单但“顺序”永远是理解 transforms 的关键。数据在每个变换中不断改变类型和数值范围后面的变换必须能接受前面的输出格式。一个常见场景如果你要在 transforms 里做 OpenCV 的操作比如光流计算因为 OpenCV 处理的是 BGR 格式的 numpy 数组你就必须把 ToTensor 之前的阶段用 OpenCV 处理或者把 ToTensor 转换后的 tensor 再转回 numpy。这个顺序绕不过去。还有一点Compose 里的变换是同步执行的全部在主进程中完成。如果你的预处理特别耗时比如大尺寸图像的 RandomResizedCropDataLoader 的num_workers参数可以帮你并行但 transforms 本身是 CPU 密集操作大量 worker 会带来 CPU 竞争。后面我会专门说这个。3.2 自定义 transforms 的两种写法很多时候内置 transforms 不够用比如你要做一个“随机调整色温”的增强或者项目特有的归一化方式。自定义 transforms 有两种常见写法。第一种继承nn.Module推荐:import torch from torch import nn class RandomColorTemperature(nn.Module): def __init__(self, delta_range(-20, 20)): super().__init__() self.delta_range delta_range def forward(self, img): if not torch.is_tensor(img): raise TypeError(fExpected tensor, got {type(img)}) delta torch.randint(self.delta_range[0], self.delta_range[1], (1,)).item() img[0] (img[0] delta / 255).clamp(0, 1) img[2] (img[2] - delta / 255).clamp(0, 1) return img第二种定义成普通函数。torchvision 的 transforms 本身既支持可调用对象也支持函数。建议继承 nn.Module 是因为它在torch.compile和模型保存加载时更友好。自定义 transforms 有几个容易被忽略的坑类型判断要明确你的变换是接收 tensor 还是 PIL Image。一个管道里有多个变换前面的变换决定了输出类型自定义变换必须与之匹配。如果类型不对直接 raise TypeError别让错误一路传到训练循环里才爆出来否则排查成本太高。默认参数的变换应该是“恒等变换”如果你的自定义变换带随机参数注意设置合理的默认值否则测试阶段很难保持确定性。支持 batch 和不支持 batch上面的代码逐张处理单张图。如果你在 DataLoader 的collate_fn之前调用 transform那它就是逐张处理的。如果你要在 GPU 上做 batch 级别的增强就得自己用 einsum 或卷积来向量化实现torchvision.transforms.v2 里提供了批量变换支持但那是另一个话题。3.3 与第三方增强库的集成方式torchvision 内置的变换对自然图像足够用但在某个细分领域就不一定了。比如做目标检测、实例分割时图像变换的同时bbox 也得跟着变换torchvision 的通用变换做不到——你需要同步变换标注信息。这个需求刚好是 albumentations 的强项。albumentations 的用法和 torchvision.transforms 很像只是它同步处理 image、mask、bboxes、keypoints 这些参数。集成方式有两种# 方式一在 Dataset 里用 albumentations import albumentations as A from albumentations.pytorch import ToTensorV2 def get_train_transforms(): return A.Compose([ A.RandomResizedCrop(height224, width224, scale(0.08, 1.0)), A.HorizontalFlip(p0.5), A.ColorJitter(brightness0.2, contrast0.2, saturation0.2, hue0.1), A.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ToTensorV2(), ])然后在 Dataset 的getitem里def __getitem__(self, idx): image cv2.imread(self.image_paths[idx]) image cv2.cvtColor(image, cv2.COLOR_BGR2RGB) mask cv2.imread(self.mask_paths[idx], cv2.IMREAD_GRAYSCALE) transformed self.transform(imageimage, maskmask) return transformed[image], transformed[mask]这种方式的好处是bboxes、masks 的变换由库自动处理不用你手动写坐标变换省心很多。我是这样选择的项目需要同步变换图像和标注分割/检测时直接用 albumentations然后把自定义的变换写成 torchvision.transforms 格式因为 torchvision 生态里的预训练模型默认假设你用的是它的 transforms 流程兼容性更好。两边混用要注意数值范围一致——albumentations 的 Normalize 需要你在参数里指定是否将输入从 0~255 转为 0~1如果直接传 PIL Image 进去类型也可能是问题务必先看文档。4. 训练阶段与推理阶段使用 transforms 的关键差异这是我看到被坑最多的一块比 API 用法本身更值得单独拿出来讲。4.1 训练阶段随机增强要多但别过度训练阶段的 transforms 设计目标只有一个通过随机变换增加数据多样性使模型学到更鲁棒的特征。常规训练 transform 长这样train_transforms transforms.Compose([ transforms.RandomResizedCrop(size(224, 224), scale(0.08, 1.0)), transforms.RandomHorizontalFlip(p0.5), transforms.RandomRotation(degrees15), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.1, hue0.05), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ])这里面每个变换的随机性都要有度。翻转概率 0.5 是约定俗成的旋转角度别看 30 度那会让图像变形很严重一般 10~15 度ColorJitter 的幅度也得根据任务来比如做车牌识别颜色抖动太大反而让模型学会“忽略颜色”对某些任务不是好事。一个常用的经验法增强强度应该保证“人类还能认得出这是什么”。如果增强后连你自己都看不出来原图里是什么物体模型也很难从中学到有效特征。还要注意增强会变相增加训练时间因为它们发生在每个 batch 加载时。如果你的数据集特别大随机增强的耗时可能超过 GPU 计算耗时这时候要优化增强逻辑而不是盲目加更多变换。4.2 推理阶段确定性优先随机全关推理阶段的 transform 就没那么多花样了核心是按序执行固定变换val_transforms transforms.Compose([ transforms.Resize((256, 256)), transforms.CenterCrop((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ])这里用 Resize 再 CenterCrop 的组合是 ImageNet 时代的标准做法。Resize 到 256CenterCrop 到 224实际上是在模仿训练时的“裁剪”行为但保证每次结果是确定的。有个常见疑问训练用 RandomResizedCrop那推理应该用 ResizeCrop 还是直接 Resize 到目标尺寸我测过不少模型直接 Resize 到目标尺寸和 ResizeCrop 相比在大多数分类任务上差距不大但在细分类任务比如区分不同鸟的品种上ResizeCrop 会更好一点因为它保留的中央区域信息更完整。如果你懒直接 Resize 到目标尺寸也可以但前提是训练时也尽量别只依赖 RandomResizedCrop而是加一些固定 Resize 的对比实验。4.3 TTATest Time Augmentation——推理阶段主动加增强TTA 是个有意思的思路推理时对同一张图做多次变换通常包括水平翻转、多尺度缩放得到多个预测结果然后对预测取平均。这相当于用“计算量换精度”。def tta_inference(model, image, device): model.eval() with torch.no_grad(): # 原图 pred1 model(image.unsqueeze(0).to(device)) # 水平翻转 flipped torch.flip(image, dims[2]) pred2 model(flipped.unsqueeze(0).to(device)) pred2 torch.flip(pred2, dims[2]) # 多尺度 scale_small transforms.Resize((196, 196))(image) scale_large transforms.Resize((256, 256))(image) pred3 model(transforms.CenterCrop((224, 224))(scale_small).unsqueeze(0).to(device)) pred4 model(transforms.CenterCrop((224, 224))(scale_large).unsqueeze(0).to(device)) pred (pred1 pred2 pred3 pred4) / 4 return predTTA 对分类任务通常能稳定提升 0.5~1 个点对检测和分割也有帮助但代价是推理时间翻倍。工业部署时一般不用 TTA除非精度指标就差这零点几个点。5. DataLoader 并行与 transforms 性能优化很多人在小数据集上跑通了代码一换大规模真实数据就发现训练循环里数据加载成了瓶颈。这里分享几个我在实际项目中用过的性能优化手段。5.1 num_workers 与持久化 workerDataLoader 的num_workers控制的是“有几个子进程负责执行 Dataset 和 transforms”。设置为 0 表示主进程加载数据加载和计算混在一起容易让 GPU 频繁空转设置为大于 0 时数据加载在子进程并行执行主进程拿到的是已经处理好的 batch。有人以为 num_workers 越大越好其实不是。worker 数量过大会导致大量进程切换开销反而变慢。经验上num_workers 设置为 CPU 核心数的一半到四分之三比较合适。比如服务器有 16 核设 8~12 个 worker具体用多少建议做一个简单的梯度测试import torch, time from torch.utils.data import DataLoader for nw in [0, 2, 4, 8, 12]: loader DataLoader(dataset, batch_size32, num_workersnw, pin_memoryTrue) start time.time() for batch in loader: pass print(fnum_workers{nw}, 耗时{time.time() - start:.2f}s)另外PyTorch 1.13 的 DataLoader 支持persistent_workersTrue意思是每个 epoch 结束时不销毁 worker 进程复用它们继续干活。如果你的数据集不大、每个 epoch 很短这个参数能省下反复创建进程的开销明显提升训练速度。5.2 pin_memory 与非阻塞传输pin_memoryTrue的含义是每个 batch 的数据在取回主内存时被锁定在固定的物理内存地址上这样从 CPU 到 GPU 的拷贝更快。配合tensor.to(device, non_blockingTrue)可以实现异步传输让 CPU 准备下一批数据的同时GPU 正在计算当前批次。for images, labels in train_loader: images images.to(device, non_blockingTrue) labels labels.to(device, non_blockingTrue)这个组合本身不会减少处理时间但能减少等待时间。如果你的机器支持强烈建议开启。5.3 瓶颈定位transforms 太慢还是 IO 太慢如果数据加载明显拖慢训练第一步是定位瓶颈是磁盘读图慢还是 transforms 计算慢方法是分阶段计时直接从硬盘读取图片不经过任何 transform看看耗时在 Dataset 里读图后立刻做 transforms看看额外耗时多少。如果磁盘 IO 是主要瓶颈建议先把图片缩放到较小尺寸再存盘或者用 TAR、TFRecord、WebDataset 这类紧凑格式减少随机 IO 次数。如果 transforms 本身就是重头那就要从算法层面优化比如用 tensor 操作代替逐像素 for 循环、减少不必要的拷贝和转换。我印象很深的一次一个医学图像项目每张图是 512x512 的 uint16 灰度图transforms 里有个“窗宽窗位调整”的操作我用 numpy 逐像素做了个 for 循环结果一个 worker 每秒只能处理 3 张图训练完全跑不动。后来改成向量化操作秒开 30 张差距十倍以上。凡是能用 numpy/torch 向量化的绝不写 for 循环。6. 实际项目中的 transforms 调试三板斧代码写多了你会发现transforms 相关的问题往往特别隐蔽报错不报错先不说模型跑起来 loss 不降、指标上不去很多时候就是预处理哪一步出了问题。分享几个高效的调试手段。6.1 可视化 transform 输出最简单的调试方式把 transform 后的结果画出来看。但要注意直接plt.imshow(tensor)会报错或显示异常因为 tensor 是 CHW 布局且经过了 Normalize很多像素是负数。正确做法是import matplotlib.pyplot as plt import torchvision.transforms.functional as TF def visualize_transformed(img_tensor, title, denormTrue): if denorm: inv_normalize transforms.Normalize( mean[-0.485/0.229, -0.456/0.224, -0.406/0.225], std[1/0.229, 1/0.224, 1/0.225] ) img_tensor inv_normalize(img_tensor) img img_tensor.cpu().detach() img TF.to_pil_image(img.clamp(0, 1)) plt.imshow(img) plt.title(title) plt.axis(off) plt.show() # 批量可视化增强效果 train_ds torchvision.datasets.ImageFolder(./data/train, transformtrain_transforms) for i in range(6): img, _ train_ds[i] visualize_transformed(img, titlefsample {i})多采几个样对比看看增强后的图像分布是否合理。如果有的图被裁到几乎没有有效内容说明 scale 参数下限设得太低了如果颜色失真严重说明 ColorJitter 的幅度过猛。6.2 检查数值范围与 dtype有些 bug 发生在数值范围上手动抽样打印几个统计量就能快速定位def inspect_tensor(tensor, nametensor): print(f{name}: shape{tuple(tensor.shape)}, dtype{tensor.dtype}, fmin{tensor.min():.4f}, max{tensor.max():.4f}, mean{tensor.mean():.4f})正常的输入数据经过 ToTensorNormalize 后各通道均值应该接近 0min/max 在 -2~2 附近。如果 min/max 还在 0~1 之间说明 Normalize 没有生效顺序错了如果 min/max 高达几十上百说明 ToTensor 没有把像素缩放到 0~1可能是 float32 输入的问题。6.3 固定随机种子复现问题transforms 里的随机增强默认是全局随机的如果你要复现某个具体的 bug 场景可以固定随机种子import random import numpy as np import torch def seed_everything(seed42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) torch.backends.cudnn.deterministic True torch.backends.cudnn.benchmark False seed_everything(42)这样每次运行数据增强的顺序和结果都是一致的方便对比问题是否有复现性。但注意训练时不要开 deterministic因为它会牺牲一部分性能只在调试阶段用一下。7. 版本演进与未来方向从 torchvision.transforms 到 v2PyTorch 生态一直在演进transforms 也一样。很多人知道 torchvision.transforms但不太了解 torchvision.transforms.v2 这个新版 API。我简单说说它解决了什么问题。v2 的核心理念是“batch-first”所有变换都是批量友好的输入可以是 (B, C, H, W) 的 batch tensor而不再是单张图。这意味着你可以把整个 batch 的增强放到 GPU 上执行完全摆脱 CPU 预处理的瓶颈。它还统一了 mask、bounding boxes、video 等数据的变换逻辑和 albumentations 的思路是一致的。from torchvision.transforms import v2 transform v2.Compose([ v2.ToImage(), v2.RandomResizedCrop(size(224, 224)), v2.RandomHorizontalFlip(p0.5), v2.ToDtype(torch.float32, scaleTrue), v2.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ])v2 里多了ToImage把 PIL Image 或 tensor 统一转成 v2 的内部图像格式、ToDtype显式转换 dtype 并可选缩放这些更底层、更可控的变换。如果你在开新项目建议直接学 v2它代表了 torchvision 的未来方向老 API 目前只是兼容保留。不过目前很多开源代码还是老 API如果你看别人的代码有transforms.Compose不用急着改成 v2。能用、能跑、能复现就行等踩到性能瓶颈了再迁移也不迟。还有一点PyTorch 2.x 时代把 transforms 放进torch.compile是可能的。如果你用 v2 的变换PyTorch 可以把它编译进计算图里做优化这在以往是做不到的。但这属于进阶玩法现阶段知道有这回事就行不需要一上来就用。关于环境版本顺便说一句我在本地项目里习惯用pip install torch torchvision同时安装并确保两者的主版本号匹配。torchvision 和 torch 是强耦合的版本不匹配时经常会出现 “cannot import name xxx from torchvision” 之类的诡异报错。遇到这种问题先检查是不是 torch 和 torchvision 版本没对上用pip list | grep torch看一眼版本号再排查其他原因。8. 一组可直接上手的完整流程模板说了这么多最后给一份我实际项目里反复用的模板从数据读取到训练循环一步到位。你可以直接复制改改参数就能跑。import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import Dataset, DataLoader import torchvision from torchvision import transforms from PIL import Image import os # ---------- 1. 定义训练/验证 transform ---------- train_transforms transforms.Compose([ transforms.RandomResizedCrop(size(224, 224), scale(0.08, 1.0)), transforms.RandomHorizontalFlip(p0.5), transforms.RandomRotation(degrees15), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.1, hue0.05), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) val_transforms transforms.Compose([ transforms.Resize((256, 256)), transforms.CenterCrop((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) # ---------- 2. 自定义 Dataset ---------- class ImageFolderDataset(Dataset): def __init__(self, root_dir, transformNone): self.file_list [] self.labels [] self.classes sorted(os.listdir(root_dir)) self.class_to_idx {cls: i for i, cls in enumerate(self.classes)} self.transform transform for cls in self.classes: cls_dir os.path.join(root_dir, cls) if not os.path.isdir(cls_dir): continue for fname in os.listdir(cls_dir): if fname.lower().endswith((.jpg, .jpeg, .png)): self.file_list.append(os.path.join(cls_dir, fname)) self.labels.append(self.class_to_idx[cls]) def __len__(self): return len(self.file_list) def __getitem__(self, idx): img Image.open(self.file_list[idx]).convert(RGB) label self.labels[idx] if self.transform is not None: img self.transform(img) return img, label # ---------- 3. 构建 DataLoader ---------- train_ds ImageFolderDataset(./data/train, transformtrain_transforms) val_ds ImageFolderDataset(./data/val, transformval_transforms) train_loader DataLoader(train_ds, batch_size64, shuffleTrue, num_workers4, pin_memoryTrue, persistent_workersTrue) val_loader DataLoader(val_ds, batch_size64, shuffleFalse, num_workers4, pin_memoryTrue, persistent_workersTrue) # ---------- 4. 模型、损失、优化器 ---------- model torchvision.models.resnet18(pretrainedTrue) model.fc nn.Linear(model.fc.in_features, len(train_ds.classes)) device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device) criterion nn.CrossEntropyLoss() optimizer optim.SGD(model.parameters(), lr0.01, momentum0.9, weight_decay1e-4) scheduler optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max20) # ---------- 5. 训练循环 ---------- def train_one_epoch(model, loader, criterion, optimizer, device): model.train() total_loss, total_correct, total_num 0.0, 0, 0 for images, labels in loader: images images.to(device, non_blockingTrue) labels labels.to(device, non_blockingTrue) outputs model(images) loss criterion(outputs, labels) optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() * images.size(0) total_correct (outputs.argmax(dim1) labels).sum().item() total_num images.size(0) return total_loss / total_num, total_correct / total_num def evaluate(model, loader, criterion, device): model.eval() total_loss, total_correct, total_num 0.0, 0, 0 with torch.no_grad(): for images, labels in loader: images images.to(device, non_blockingTrue) labels labels.to(device, non_blockingTrue) outputs model(images) loss criterion(outputs, labels) total_loss loss.item() * images.size(0) total_correct (outputs.argmax(dim1) labels).sum().item() total_num images.size(0) return total_loss / total_num, total_correct / total_num for epoch in range(20): train_loss, train_acc train_one_epoch(model, train_loader, criterion, optimizer, device) val_loss, val_acc evaluate(model, val_loader, criterion, device) scheduler.step() print(fEpoch {epoch1:03d} | train_loss{train_loss:.4f} | train_acc{train_acc:.4f} | fval_loss{val_loss:.4f} | val_acc{val_acc:.4f})这个模板里用到了前面说的所有关键点训练/验证 transform 分开设计、pin_memory non_blocking persistent_workers 优化数据传输、ImageNet 预训练模型配 ImageNet normalize 参数。你可以直接拿它作为项目的起点再根据实际情况调整。我自己做项目时的习惯是先把这套流水线跑通确认模型能拟合几张小批量数据判断代码有没有 bug再加大数据量训练、调参。transforms 这块代码基本不动因为模板已经稳定了偶尔换数据集时只改变换的参数就行。希望这份经验对你有用。