简介:卷积神经网络(CNN)是计算机视觉领域的核心技术,其通过卷积和池化操作,能够高效提取图像的局部与层次化特征,解决了传统全连接网络处理图像时参数爆炸的问题。这一原理奠定了现代图像识别的基础,其技术价值在于能够端到端地学习从原始像素到高级语义的映射。在工程实践中,迁移学习利用在大规模数据集(如ImageNet)上预训练的模型,通过微调快速适应特定任务,极大降低了数据需求和训练成本。本文以动物图像分类为具体应用场景,详细阐述了如何使用PyTorch框架,结合数据增强、模型微调等技巧,从数据准备、模型搭建、训练优化到可视化部署,完整实现一个高效的分类系统,为入门者提供了清晰的深度学习项目实践路径。
1. 项目概述:从零到一构建一个动物分类器
最近在整理硬盘,翻出来一个老项目——“基于深度学习的动物图像分类.zip”。这让我想起了几年前刚开始接触计算机视觉时,那种既兴奋又迷茫的状态。当时就想,能不能让电脑像我们一样,看一眼图片就知道里面是猫还是狗,是老虎还是狮子?这个压缩包里的代码,就是我当年磕磕绊绊实现这个想法的完整记录。今天,我打算把这个项目重新梳理一遍,把里面的坑、走过的弯路以及最终跑通的喜悦,完整地分享出来。无论你是刚入门深度学习的新手,还是想找一个完整的端到端项目练手,这篇内容都能给你提供一个清晰的路线图。
这个项目的核心目标非常明确:训练一个模型,让它能自动识别并分类图片中的动物。听起来像是ImageNet竞赛的简化版,对吧?但它麻雀虽小,五脏俱全。从数据的收集与清洗、模型的选择与搭建、到训练策略的调整、性能的评估与优化,乃至最后封装成一个简单的应用,整个流程覆盖了深度学习项目实践中的绝大多数关键环节。我们不会使用任何现成的、封装过度的API(比如某些云服务的一键训练),而是从最基础的PyTorch或TensorFlow开始,亲手搭建每一个部件。这样做的目的,是让你真正理解数据是如何流动的,梯度是如何计算的,模型又是如何“学会”区分不同特征的。整个过程,就像教一个孩子认识世界,你需要准备足够多且清晰的“教材”(数据),设计合理的“教学方法”(模型与损失函数),并耐心地“纠正错误”(优化与调参)。
2. 核心思路与方案选型:为什么是CNN?
当我们拿到“动物图像分类”这个任务时,第一个要回答的问题是:用什么模型?在深度学习的武器库里,卷积神经网络(CNN)几乎是处理图像问题的“标准答案”。但为什么是它?这得从图像数据的本质说起。
一张图片在计算机眼里,就是一个巨大的数字矩阵(如果是彩色图,就是三个这样的矩阵堆叠在一起)。传统的全连接神经网络(就是那种一层层神经元全部相连的网络)来处理图像,参数数量会爆炸。想象一下,一张224x224的彩色图片,输入层就有2242243=150,528个神经元。如果下一层也有10万个神经元,那么这一层的参数就超过150亿!这根本无法训练。CNN的巧妙之处在于它引入了“卷积”和“池化”操作。
卷积可以理解为拿着一个小滤镜(卷积核)在图片上滑动。这个滤镜专门负责检测某种局部特征,比如边缘、纹理、颜色块。通过多个不同的卷积核,网络就能自动学习到从简单到复杂的各种特征。池化(通常是最大池化)则像一个“信息浓缩”的过程,它把一个小区域内的特征值(比如2x2)压缩成一个(取最大值),这样做的目的是降低数据维度,增强特征的不变性(比如动物稍微移动一点位置,最大的特征可能还在),同时也能减少计算量。这种“局部连接”和“权值共享”的特性,使得CNN特别擅长捕捉图像的空间层次信息,从边缘到纹理,再到局部器官(如眼睛、耳朵),最后到整个物体。
注意:对于刚入门的朋友,可能会纠结于选择PyTorch还是TensorFlow。我的建议是,如果你是纯粹的新手,从PyTorch开始会更容易上手。它的设计更“Pythonic”,动态图机制让调试像写普通Python代码一样直观。TensorFlow的静态图虽然在大规模部署上仍有优势,但其2.x版本也拥抱了动态图(Eager Execution),两者差距在缩小。这个项目我们将以PyTorch为主线进行讲解,但核心思想是通用的。
确定了CNN这个大方向后,我们面临几个具体的方案选择:
- 从零开始训练(Scratch):自己设计网络结构(比如几个卷积层、几个全连接层),然后用我们的动物数据集从头训练。优点是完全可控,理解深刻。缺点是需要大量的数据和时间,对于小数据集极易过拟合。
- 迁移学习(Transfer Learning):使用在大型数据集(如ImageNet)上预训练好的成熟模型(如ResNet, VGG, EfficientNet),只替换其最后的分类头,然后用我们的动物数据对其最后几层或全部层进行微调(Fine-tuning)。这是本项目最推荐、也是实际中最常用的方法。它相当于让模型站在巨人的肩膀上,利用已学到的通用图像特征,快速适应我们的特定任务,在数据量有限的情况下也能取得非常好的效果。
我们的项目方案很明确:采用迁移学习策略,使用预训练的ResNet-18作为基础模型,在其后接一个适配我们动物类别数量的全连接层,对整个网络进行微调。ResNet-18结构相对简单,训练速度快,且在ImageNet上表现优异,其学习到的特征足以迁移到动物分类任务上。
3. 数据准备:项目的基石与第一个大坑
都说数据和特征决定了机器学习的上限,模型和算法只是逼近这个上限。在动物分类项目里,数据准备往往是耗时最长、也最容易出问题的环节。我们的zip包里,通常应该包含一个data/目录,里面是已经分好类的图片,比如train/dog/,train/cat/,val/dog/等等。但如果你的压缩包里没有,或者你想用自己的数据,那就得从头开始。
3.1 数据收集与爬取
最直接的数据来源是公开数据集,比如:
- Kaggle: 搜索“cats and dogs”、“animals-10”等,有很多高质量的标注数据集。
- ImageNet: 你可以下载其中与动物相关的子集(但过程较复杂)。
- 谷歌开源图像数据集(Open Images)等。
如果公开数据集不满足要求(比如你需要特定种类的动物),可能需要自己爬取。可以使用Python的requests和BeautifulSoup库,或者更高效的scrapy框架,从搜索引擎或图片网站抓取。这里有一个非常重要的注意事项:务必遵守网站的robots.txt协议,尊重版权,并且控制爬取速度和频率,避免对目标服务器造成压力。
3.2 数据清洗与整理
爬取或下载的图片往往是“脏数据”,这是第一个大坑。清洗步骤必不可少:
- 去重:使用哈希(如MD5)或感知哈希(pHash)找出并删除完全相同的或高度相似的图片。
- 过滤:删除损坏的、无法打开的图片(用PIL或OpenCV读取时捕获异常)。删除分辨率过低的图片(如小于64x64)。
- 人工审核(关键!):这是最耗时但无法省略的一步。你需要快速浏览图片,剔除明显错误的图片(比如标签是“狗”,图片却是汽车),以及质量极差的图片(极度模糊、主体不完整)。可以写一个简单的脚本,将图片以网格形式展示,方便快速标记和删除。
清洗后,按照以下目录结构进行整理,这是PyTorchImageFolder类所期望的格式,能极大简化后续数据加载工作:
animal_dataset/ ├── train/ │ ├── dog/ │ │ ├── dog001.jpg │ │ ├── dog002.jpg │ │ └── ... │ ├── cat/ │ │ ├── cat001.jpg │ │ └── ... │ └── tiger/ ├── val/ │ ├── dog/ │ ├── cat/ │ └── tiger/ └── test/ (可选,也可用val代替) ├── dog/ ├── cat/ └── tiger/train用于训练,val用于在训练过程中监控模型表现、防止过拟合,test用于最终评估。通常按7:2:1或8:1:1的比例随机划分。
3.3 数据增强(Data Augmentation)
我们的数据量通常不会像ImageNet那样庞大。为了提升模型的泛化能力,防止过拟合,数据增强是神器。它通过对训练图片进行一系列随机变换,来“创造”出新的训练样本。PyTorch的torchvision.transforms提供了非常方便的工具。
一个典型的训练数据增强流水线如下:
from torchvision import transforms train_transform = transforms.Compose([ transforms.RandomResizedCrop(224), # 随机裁剪并缩放到224x224 transforms.RandomHorizontalFlip(p=0.5), # 随机水平翻转,概率50% transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1), # 随机颜色抖动 transforms.RandomRotation(degrees=15), # 随机旋转±15度 transforms.ToTensor(), # 转换为Tensor,并归一化到[0,1] transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) # 标准化 ])而验证和测试集通常只进行确定性变换(裁剪、缩放、归一化),不进行随机增强,以保证评估的一致性。
val_transform = transforms.Compose([ transforms.Resize(256), # 将短边缩放到256 transforms.CenterCrop(224), # 中心裁剪到224x224 transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ])实操心得:
Normalize使用的均值和标准差是ImageNet数据集上的统计值。因为我们使用在ImageNet上预训练的模型,所以输入数据必须采用相同的归一化参数,这样才能保证模型之前学到的特征分布是匹配的。这是一个非常容易忽略但至关重要的细节。
4. 模型搭建与迁移学习实战
数据管道准备好了,接下来就是模型部分。我们将使用torchvision.models中提供的预训练模型。
4.1 加载预训练模型
import torch import torchvision.models as models import torch.nn as nn # 检查是否有可用的GPU device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') print(f'Using device: {device}') # 加载预训练的resnet18模型 model = models.resnet18(pretrained=True) # 冻结模型的所有参数(在初始阶段,可选) # for param in model.parameters(): # param.requires_grad = Falsepretrained=True会自动下载预训练的权重。此时,model最后的全连接层(fc)是针对ImageNet的1000个类别的。
4.2 修改分类头(Classifier Head)
我们的动物类别数(假设是10类)与1000不同,因此需要替换最后的全连接层。
# 获取原始fc层的输入特征数 num_ftrs = model.fc.in_features # 假设我们的动物有10个类别 num_classes = 10 # 替换fc层为一个新的Sequential模块,可以加入Dropout防止过拟合 model.fc = nn.Sequential( nn.Dropout(p=0.5), # 丢弃概率0.5 nn.Linear(num_ftrs, num_classes) ) # 将模型移动到GPU(如果可用) model = model.to(device)这里我添加了一个Dropout层,它在训练时会随机“关闭”一部分神经元,是一种有效的正则化手段。对于小数据集上的微调,这通常是个好主意。
4.3 设置差分学习率(Differential Learning Rates)
这是微调时的另一个关键技巧。模型的前面几层学习到的是通用特征(如边缘、纹理),这些特征对于动物分类也很有用,我们不想让它们改变太多。而越靠近输出的层(尤其是我们新加的fc层),其任务特异性越强,需要更快地学习。
因此,我们可以为模型的不同部分设置不同的学习率。
# 将模型参数分为三组: # 1. 新添加的fc层的参数 # 2. ResNet最后两个基本块(layer4)的参数 # 3. 其他层的参数 fc_params = list(map(id, model.fc.parameters())) # 获取新fc层参数的id base_params = filter(lambda p: id(p) not in fc_params, model.parameters()) # 定义优化器,为不同参数组设置不同的学习率 optimizer = torch.optim.Adam([ {'params': base_params, 'lr': 1e-4}, # 基础层,学习率较小 {'params': model.fc.parameters(), 'lr': 1e-3} # 新层,学习率较大 ], weight_decay=1e-4) # 加入权重衰减(L2正则化)Adam优化器结合了动量和自适应学习率,是当前非常流行的选择。weight_decay参数用于控制模型复杂度,防止过拟合。
5. 训练循环与核心技巧
训练是模型“学习”的过程。我们需要定义一个循环,反复执行“前向传播 -> 计算损失 -> 反向传播 -> 更新参数”这个过程。
5.1 定义损失函数与训练/验证步骤
对于多分类问题,交叉熵损失(CrossEntropyLoss)是标准选择。
criterion = nn.CrossEntropyLoss() def train_one_epoch(model, dataloader, criterion, optimizer, device): model.train() # 设置为训练模式(启用Dropout等) running_loss = 0.0 correct = 0 total = 0 for batch_idx, (inputs, labels) in enumerate(dataloader): inputs, labels = inputs.to(device), labels.to(device) # 清零梯度 optimizer.zero_grad() # 前向传播 outputs = model(inputs) loss = criterion(outputs, labels) # 反向传播与优化 loss.backward() optimizer.step() # 统计 running_loss += loss.item() _, predicted = outputs.max(1) total += labels.size(0) correct += predicted.eq(labels).sum().item() epoch_loss = running_loss / len(dataloader) epoch_acc = 100. * correct / total return epoch_loss, epoch_acc def validate(model, dataloader, criterion, device): model.eval() # 设置为评估模式(关闭Dropout等) running_loss = 0.0 correct = 0 total = 0 with torch.no_grad(): # 关闭梯度计算,节省内存和计算 for inputs, labels in dataloader: inputs, labels = inputs.to(device), labels.to(device) outputs = model(inputs) loss = criterion(outputs, labels) running_loss += loss.item() _, predicted = outputs.max(1) total += labels.size(0) correct += predicted.eq(labels).sum().item() epoch_loss = running_loss / len(dataloader) epoch_acc = 100. * correct / total return epoch_loss, epoch_acc5.2 主训练循环与模型保存
我们需要循环多个轮次(Epoch),并在每个Epoch后验证模型,保存最好的那个。
num_epochs = 30 best_val_acc = 0.0 train_losses, val_losses = [], [] train_accs, val_accs = [], [] for epoch in range(num_epochs): print(f'\nEpoch {epoch+1}/{num_epochs}') print('-' * 50) # 训练 train_loss, train_acc = train_one_epoch(model, train_loader, criterion, optimizer, device) train_losses.append(train_loss) train_accs.append(train_acc) print(f'Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.2f}%') # 验证 val_loss, val_acc = validate(model, val_loader, criterion, device) val_losses.append(val_loss) val_accs.append(val_acc) print(f'Val Loss: {val_loss:.4f} | Val Acc: {val_acc:.2f}%') # 保存最佳模型 if val_acc > best_val_acc: best_val_acc = val_acc torch.save({ 'epoch': epoch, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'val_acc': val_acc, }, 'best_animal_classifier.pth') print(f'>>> Model saved with Val Acc: {val_acc:.2f}%')5.3 学习率调度(Learning Rate Scheduling)
固定学习率可能不是最优的。我们可以在训练过程中动态调整它,例如当验证集准确率不再提升时,降低学习率。
# 使用ReduceLROnPlateau调度器 scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='max', factor=0.5, patience=5, verbose=True) # 在每个epoch的验证步骤后调用 scheduler.step(val_acc)这个调度器会监控val_acc,如果连续patience个epoch没有提升,就将学习率乘以factor。
6. 可视化、调试与性能分析
训练不是黑盒。我们需要工具来洞察模型的行为。
6.1 使用TensorBoard可视化
TensorBoard是TensorFlow的可视化工具包,但PyTorch通过torch.utils.tensorboard可以无缝使用。
from torch.utils.tensorboard import SummaryWriter writer = SummaryWriter('runs/animal_experiment_1') # 在训练循环中记录标量 writer.add_scalar('Loss/train', train_loss, epoch) writer.add_scalar('Accuracy/train', train_acc, epoch) writer.add_scalar('Loss/val', val_loss, epoch) writer.add_scalar('Accuracy/val', val_acc, epoch) # 还可以记录模型图、直方图等 # writer.add_graph(model, inputs.to(device))训练后,在终端运行tensorboard --logdir=runs,然后在浏览器打开提示的地址,就能看到漂亮的损失和准确率曲线,方便我们判断模型是否过拟合/欠拟合。
6.2 绘制混淆矩阵(Confusion Matrix)
混淆矩阵能详细展示模型在哪些类别上容易混淆。这是分析模型弱点、指导数据收集或后处理的关键工具。
from sklearn.metrics import confusion_matrix import seaborn as sns import matplotlib.pyplot as plt def plot_confusion_matrix(model, dataloader, class_names, device): model.eval() all_preds = [] all_labels = [] with torch.no_grad(): for inputs, labels in dataloader: inputs = inputs.to(device) outputs = model(inputs) _, preds = torch.max(outputs, 1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.numpy()) cm = confusion_matrix(all_labels, all_preds) plt.figure(figsize=(10, 8)) sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=class_names, yticklabels=class_names) plt.ylabel('True Label') plt.xlabel('Predicted Label') plt.title('Confusion Matrix') plt.tight_layout() plt.show() # 使用测试集或验证集绘制 plot_confusion_matrix(model, test_loader, ['cat', 'dog', 'tiger', ...], device)如果发现“豹”和“猎豹”经常分错,那很可能是训练数据中这两类图片的特征不够区分度,或者数据量不足。
6.3 可视化卷积核与特征图
为了理解CNN到底学到了什么,我们可以可视化第一层的卷积核,或者查看中间层对某张图片的激活(特征图)。
# 获取第一层卷积的权重 first_layer_weights = model.conv1.weight.data.cpu().numpy() # first_layer_weights的形状是 [out_channels, in_channels, kernel_h, kernel_w] # 可以将其可视化为多个小滤镜图像这有助于直观感受底层特征检测器(如边缘检测器)的样子。
7. 模型优化与部署前准备
训练出一个可用的模型只是第一步,要让其真正“好用”,还需要进行优化和封装。
7.1 模型剪枝与量化(可选,用于部署)
如果考虑将模型部署到移动端或嵌入式设备,需要对模型进行“瘦身”。
- 剪枝:移除网络中不重要的连接(权重接近0的),减少参数数量。
- 量化:将模型权重和激活从32位浮点数(FP32)转换为8位整数(INT8),大幅减少模型体积和推理时的计算量,通常只会带来极小的精度损失。 PyTorch提供了
torch.quantization和torch.nn.utils.prune工具包来实现这些功能。不过,对于初版项目,可以暂不进行,先以保证精度为主。
7.2 创建简单的推理脚本
我们需要一个脚本,能够加载训练好的模型,并对单张或批量图片进行预测。
import torch from PIL import Image from torchvision import transforms class AnimalClassifier: def __init__(self, model_path, class_names, device='cpu'): self.device = torch.device(device) self.class_names = class_names # 加载模型结构 self.model = models.resnet18(pretrained=False) num_ftrs = self.model.fc.in_features self.model.fc = nn.Linear(num_ftrs, len(class_names)) # 加载训练好的权重 checkpoint = torch.load(model_path, map_location=self.device) self.model.load_state_dict(checkpoint['model_state_dict']) self.model.to(self.device) self.model.eval() # 定义与训练时验证集相同的转换 self.transform = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) def predict(self, image_path): """预测单张图片""" image = Image.open(image_path).convert('RGB') image_tensor = self.transform(image).unsqueeze(0) # 增加batch维度 image_tensor = image_tensor.to(self.device) with torch.no_grad(): outputs = self.model(image_tensor) probabilities = torch.nn.functional.softmax(outputs, dim=1) confidence, predicted_idx = torch.max(probabilities, 1) predicted_class = self.class_names[predicted_idx.item()] confidence = confidence.item() return predicted_class, confidence # 使用示例 classifier = AnimalClassifier('best_animal_classifier.pth', ['cat', 'dog', 'elephant', ...], device='cuda:0') pred_class, conf = classifier.predict('my_pet.jpg') print(f'Predicted: {pred_class} with confidence {conf:.2%}')7.3 使用Gradio或Streamlit构建Web Demo
为了让没有编程背景的人也能体验你的模型,可以快速搭建一个Web界面。Gradio尤其适合快速原型开发。
# pip install gradio import gradio as gr classifier = AnimalClassifier(...) def predict_image(image): # image是gradio上传的PIL Image对象 # 需要先保存到一个临时路径,或用其直接转换 import tempfile with tempfile.NamedTemporaryFile(suffix='.jpg', delete=False) as tmp: image.save(tmp.name) pred, conf = classifier.predict(tmp.name) return f"{pred} ({conf:.1%})" # 创建界面 iface = gr.Interface( fn=predict_image, inputs=gr.Image(type="pil"), outputs="text", title="动物图像分类器", description="上传一张动物图片,模型会预测它是什么动物。" ) iface.launch(share=True) # share=True会生成一个临时公网链接运行这段代码,就会在本地启动一个Web服务,并提供一个链接,任何人都可以通过浏览器上传图片查看分类结果。
8. 常见问题、踩坑记录与排查指南
在实际操作中,你几乎一定会遇到下面这些问题。我把它们和解决方案整理出来,希望能帮你节省大量时间。
8.1 训练问题
| 问题现象 | 可能原因 | 排查与解决思路 |
|---|---|---|
| Loss居高不下,准确率随机(~1/类别数) | 学习率设置不当(通常太大) | 尝试大幅降低学习率(如从1e-3降到1e-5),使用学习率查找器(LR Finder)工具。检查数据标签是否正确。 |
| 训练Loss下降,但验证Loss上升(过拟合) | 模型复杂度过高,训练数据不足或噪声大,训练轮次太多。 | 1. 增强数据增强的强度。2. 增加Dropout比率。3. 加强L2正则化(增大weight_decay)。4. 使用更简单的模型(如ResNet-18换成ResNet-18)。5. 早停(Early Stopping)。 |
| 训练Loss和验证Loss都不动(欠拟合) | 模型能力不足,学习率太小,特征提取层(预训练部分)被冻结且未解冻。 | 1. 使用更复杂的模型。2. 增大学习率。3. 解冻更多预训练层进行微调。4. 检查数据预处理是否正确(特别是归一化参数)。 |
| GPU内存溢出(CUDA out of memory) | 批次大小(Batch Size)太大,模型太大。 | 1. 减小batch_size。2. 使用梯度累积(Gradient Accumulation):每N个小批次累加梯度后再更新一次权重,等效于增大batch size。3. 使用混合精度训练(AMP),减少显存占用并加速。 |
| 训练速度很慢 | 数据加载是瓶颈,没有使用GPU,CPU模式运行。 | 1. 使用DataLoader的num_workers参数(通常设为CPU核心数)进行多进程数据加载。2. 使用pin_memory=True加速GPU数据传输。3. 确保模型和.to(device)在GPU上。 |
8.2 数据与预处理问题
- 问题:验证准确率远低于训练准确率,且差距随训练持续扩大。
- 排查:首先检查数据泄露!确保训练集和验证集是严格分离的,没有重复的图片。检查数据增强是否错误地应用到了验证集(验证集应该只做确定性变换)。
- 问题:模型对所有样本都预测为同一个类别。
- 排查:检查数据类别是否极度不平衡。如果90%的图片都是狗,模型可能会学会永远预测“狗”来获得高准确率。需要采用过采样、欠采样或为损失函数添加类别权重(
nn.CrossEntropyLoss(weight=class_weights))。
- 排查:检查数据类别是否极度不平衡。如果90%的图片都是狗,模型可能会学会永远预测“狗”来获得高准确率。需要采用过采样、欠采样或为损失函数添加类别权重(
- 问题:归一化后图片看起来很奇怪(全黑或全白)。
- 排查:确认
ToTensor()是否在Normalize()之前。ToTensor()会将像素值从[0,255]转换到[0.0,1.0]。如果顺序反了,对[0,255]的整数应用ImageNet的均值和标准差,结果会超出正常范围。
- 排查:确认
8.3 环境与依赖问题
- “CUDA driver version is insufficient”:升级你的NVIDIA显卡驱动。
- “No module named ‘torch’:确保在正确的Python环境下用pip或conda安装了PyTorch,并且安装的是支持CUDA的版本(如果需要GPU)。去PyTorch官网根据你的系统配置生成安装命令是最稳妥的。
- 训练时出现NaN损失:可能是学习率太大导致梯度爆炸,尝试降低学习率,或使用梯度裁剪(
torch.nn.utils.clip_grad_norm_)。
8.4 模型保存与加载问题
- 保存时最好保存整个checkpoint(如我们示例中的字典),而不仅仅是
model.state_dict()。这样恢复训练时,可以连同优化器状态、epoch数一起加载,方便断点续训。 - 在不同设备(如从GPU训练切换到CPU推理)上加载模型时,需要使用
map_location参数:torch.load(‘model.pth’, map_location=torch.device(‘cpu’))。
这个“基于深度学习的动物图像分类”项目,虽然标题简单,但完整走一遍,你会对深度学习项目的全生命周期有一个扎实的把握。从数据工程的琐碎,到模型调参的耐心,再到问题排查的抓狂,最后到模型跑通、准确率提升时的成就感,这些都是书本和课程难以给予的实战经验。最关键的是,你拥有了一个可以随时运行、展示甚至扩展的完整项目,这才是你简历和知识库里最实在的东西。
本文还有配套的精品资源,点击获取