用 PyTorch 搭建一个可复用的 CNN 图像分类训练闭环

用 PyTorch 搭建一个可复用的 CNN 图像分类训练闭环
很多开发者第一次学习 CNN 时容易停留在“卷积层提特征、池化层降维、全连接层分类”的概念层面真正写训练代码时却会遇到一组更工程化的问题数据目录怎么组织、输入尺寸如何统一、训练和验证如何拆分、模型参数怎样保存、推理脚本如何复用训练时的预处理逻辑。本文不追求刷榜精度也不编造某个数据集上的测试结果而是搭建一个可复用的最小训练闭环。你可以先用自己的小型图片数据跑通流程再替换模型结构、增强策略或部署方式。典型目录如下cnn-demo/ config.py model.py train.py predict.py data/ train/ cat/ dog/ val/ cat/ dog/ checkpoints/这里使用ImageFolder约定每个类别一个子目录目录名就是类别名。真实项目中建议将训练集和验证集提前固定下来避免每次随机划分导致结果不可复现。CNN 的核心原理CNN 的优势来自局部连接和参数共享。普通全连接层会让每个输入像素都连接到每个输出神经元参数量随图片尺寸快速膨胀卷积层只在局部窗口内计算并让同一个卷积核在整张图上滑动因此能用较少参数捕捉边缘、纹理、局部形状等视觉模式。一个基础图像分类 CNN 通常包含四类组件卷积层提取局部特征例如边缘、颜色块、纹理组合。激活函数引入非线性常用ReLU。池化层降低空间尺寸减少计算量并提高一定的位置鲁棒性。分类头将高维特征映射为类别 logits再交给损失函数计算误差。需要注意训练时模型输出通常不是概率而是 logits。使用nn.CrossEntropyLoss时不需要在模型末尾手动加Softmax因为该损失函数内部会处理对数概率计算。推理阶段如果要展示置信度再对 logits 做softmax即可。环境与配置先安装依赖。具体版本应以你的项目环境为准如果使用 GPU还需要安装与你 CUDA 环境匹配的 PyTorch 构建包。pipinstalltorch torchvision pillow把可变参数集中到config.py便于后续调整frompathlibimportPath ROOTPath(__file__).resolve().parent DATA_DIRROOT/dataTRAIN_DIRDATA_DIR/trainVAL_DIRDATA_DIR/valCKPT_DIRROOT/checkpointsCKPT_PATHCKPT_DIR/cnn_best.ptIMAGE_SIZE128BATCH_SIZE32EPOCHS10LR1e-3NUM_WORKERS2如果你的项目需要访问私有对象存储或远程服务不要把密钥写进代码应从环境变量读取例如importos access_keyos.environ.get(APP_ACCESS_KEY)ifnotaccess_key:raiseRuntimeError(APP_ACCESS_KEY is required)本文示例本身不需要任何密钥。定义模型下面是一个小型 CNN适合用来验证训练链路。它不是面向生产精度优化的结构但层次清晰便于理解和修改。importtorchfromtorchimportnnclassSmallCNN(nn.Module):def__init__(self,num_classes:int):super().__init__()self.featuresnn.Sequential(nn.Conv2d(3,32,kernel_size3,padding1),nn.BatchNorm2d(32),nn.ReLU(inplaceTrue),nn.MaxPool2d(2),nn.Conv2d(32,64,kernel_size3,padding1),nn.BatchNorm2d(64),nn.ReLU(inplaceTrue),nn.MaxPool2d(2),nn.Conv2d(64,128,kernel_size3,padding1),nn.BatchNorm2d(128),nn.ReLU(inplaceTrue),nn.AdaptiveAvgPool2d((1,1)),)self.classifiernn.Linear(128,num_classes)defforward(self,x:torch.Tensor)-torch.Tensor:xself.features(x)xtorch.flatten(x,1)returnself.classifier(x)这里使用AdaptiveAvgPool2d((1, 1))可以让分类头不依赖固定的中间特征图尺寸。只要输入图片经过预处理后尺寸一致模型结构就更容易维护。训练与验证流程训练脚本要完成五件事加载数据、构建模型、定义损失和优化器、循环训练、保存验证集表现最好的权重。importtorchfromtorchimportnnfromtorch.utils.dataimportDataLoaderfromtorchvisionimportdatasets,transformsfromconfigimportTRAIN_DIR,VAL_DIR,CKPT_DIR,CKPT_PATH,IMAGE_SIZE,BATCH_SIZE,EPOCHS,LR,NUM_WORKERSfrommodelimportSmallCNNdefbuild_loaders():train_tftransforms.Compose([transforms.Resize((IMAGE_SIZE,IMAGE_SIZE)),transforms.RandomHorizontalFlip(),transforms.ToTensor(),transforms.Normalize(mean[0.485,0.456,0.406],std[0.229,0.224,0.225]),])val_tftransforms.Compose([transforms.Resize((IMAGE_SIZE,IMAGE_SIZE)),transforms.ToTensor(),transforms.Normalize(mean[0.485,0.456,0.406],std[0.229,0.224,0.225]),])train_setdatasets.ImageFolder(TRAIN_DIR,transformtrain_tf)val_setdatasets.ImageFolder(VAL_DIR,transformval_tf)train_loaderDataLoader(train_set,batch_sizeBATCH_SIZE,shuffleTrue,num_workersNUM_WORKERS)val_loaderDataLoader(val_set,batch_sizeBATCH_SIZE,shuffleFalse,num_workersNUM_WORKERS)returntrain_loader,val_loader,train_set.classesdefevaluate(model,loader,criterion,device):model.eval()total_loss,correct,total0.0,0,0withtorch.no_grad():forimages,labelsinloader:images,labelsimages.to(device),labels.to(device)logitsmodel(images)losscriterion(logits,labels)total_lossloss.item()*images.size(0)predslogits.argmax(dim1)correct(predslabels).sum().item()totallabels.size(0)returntotal_loss/total,correct/totaldefmain():devicetorch.device(cudaiftorch.cuda.is_available()elsecpu)train_loader,val_loader,classesbuild_loaders()modelSmallCNN(num_classeslen(classes)).to(device)criterionnn.CrossEntropyLoss()optimizertorch.optim.Adam(model.parameters(),lrLR)CKPT_DIR.mkdir(parentsTrue,exist_okTrue)best_acc0.0forepochinrange(1,EPOCHS1):model.train()running_loss0.0forimages,labelsintrain_loader:images,labelsimages.to(device),labels.to(device)optimizer.zero_grad()logitsmodel(images)losscriterion(logits,labels)loss.backward()optimizer.step()running_lossloss.item()*images.size(0)train_lossrunning_loss/len(train_loader.dataset)val_loss,val_accevaluate(model,val_loader,criterion,device)print(fepoch{epoch}train_loss{train_loss:.4f}val_loss{val_loss:.4f}val_acc{val_acc:.4f})ifval_accbest_acc:best_accval_acc torch.save({model:model.state_dict(),classes:classes},CKPT_PATH)if__name____main__:main()执行训练python train.py如果你的机器没有 GPU代码会自动使用 CPU只是训练速度可能较慢。示例中的准确率输出只能反映当前数据、划分、增强方式和训练轮数不能作为通用性能结论。推理脚本推理阶段必须复用验证阶段的尺寸调整和归一化逻辑否则训练和推理的数据分布会不一致。importsysimporttorchfromPILimportImagefromtorchvisionimporttransformsfromconfigimportCKPT_PATH,IMAGE_SIZEfrommodelimportSmallCNNdefmain(image_path:str):devicetorch.device(cudaiftorch.cuda.is_available()elsecpu)checkpointtorch.load(CKPT_PATH,map_locationdevice)classescheckpoint[classes]modelSmallCNN(num_classeslen(classes)).to(device)model.load_state_dict(checkpoint[model])model.eval()tftransforms.Compose([transforms.Resize((IMAGE_SIZE,IMAGE_SIZE)),transforms.ToTensor(),transforms.Normalize(mean[0.485,0.456,0.406],std[0.229,0.224,0.225]),])imageImage.open(image_path).convert(RGB)tensortf(image).unsqueeze(0).to(device)withtorch.no_grad():logitsmodel(tensor)probstorch.softmax(logits,dim1)[0]idxint(probs.argmax().item())print({class:classes[idx],confidence:float(probs[idx].item())})if__name____main__:iflen(sys.argv)!2:raiseSystemExit(usage: python predict.py path/to/image.jpg)main(sys.argv[1])执行python predict.py ./sample.jpg可执行改造建议跑通最小闭环后可以按优先级逐步改造先检查数据质量类别目录是否正确、是否存在损坏图片、训练集和验证集是否混入重复样本。再调整输入尺寸和 batch size显存不足时优先降低 batch size而不是盲目删模型层。引入更强的数据增强例如随机裁剪、颜色扰动但验证集不要使用随机增强。替换骨干网络可以用torchvision.models中的预训练模型做迁移学习但要确认输入归一化和分类头修改正确。增加日志与配置管理生产项目建议记录参数、代码版本、数据版本和模型文件路径。这些改造的前提是先有稳定的训练、验证、保存和推理闭环。没有闭环时直接堆复杂模型往往只会增加排查难度。常见问题1. 为什么训练集准确率升高验证集不升反降常见原因是过拟合、训练验证分布不一致、数据量过小或验证集标注质量差。可以先减少模型容量、增加数据增强、固定划分方式并人工抽查错误样本。2. 为什么CrossEntropyLoss前不要加Softmax因为CrossEntropyLoss期望输入 logits并在内部组合了对数 softmax 与负对数似然损失。提前加Softmax可能带来数值稳定性和梯度表达问题。3. 为什么推理结果类别对不上ImageFolder会按类别目录名生成类别索引。保存模型时应同时保存classes推理时读取同一份类别列表避免手写类别顺序导致错位。4. 小数据集是否适合从零训练 CNN可以用于学习流程但未必适合获得稳定泛化能力。真实业务中如果数据量有限通常优先考虑迁移学习、冻结部分骨干层和更严格的数据清洗。5. 多进程 DataLoader 在 Windows 上报错怎么办确保训练入口放在if __name__ __main__:下如果仍不稳定可以先把NUM_WORKERS改为0验证主流程。总结一个可维护的 CNN 项目不只是模型结构本身还包括数据约定、预处理一致性、训练验证拆分、权重保存和推理复用。本文给出的 PyTorch 示例刻意保持简单目标是让训练闭环清晰可运行。后续无论替换为 ResNet、MobileNet还是加入更复杂的增强和部署逻辑都应保留这条主线输入可追踪训练可复现验证可解释推理与训练保持同一套数据处理规则。