迁移学习这几年在计算机视觉里的地位,基本等同于“省力杠杆”——同样一个图像分类任务,从零训练 ResNet 可能要好几天才能收敛,加载 ImageNet 预训练权重之后,往往几个 epoch 就能达到可用效果。这篇文章不聊概念堆砌,直接讲清楚迁移学习在视觉任务里的落地方案:怎么选预训练模型、怎么改分类头、怎么冻结特征层、怎么在少量数据集上做微调,以及最终如何把训练好的模型封装成 API 接口和批量推理流程。如果你想做图像分类、目标检测或者特征提取,又不想从零开始烧算力,这篇文章可以直接收藏。
迁移学习最核心的价值有三点:第一,预训练模型已经学到了通用的边缘、纹理、形状等底层特征,这些特征在不同视觉任务之间是可复用的;第二,微调阶段的训练时间通常是从零训练的十分之一甚至更短,对硬件要求也更友好;第三,即使只有几百张甚至几十张标注图片,也能训练出可用的视觉模型。本文会用 PyTorch 和 Torchvision 作为主技术栈,走一遍从加载预训练权重到模型微调、效果验证、服务部署的完整流程。
文章里涉及的具体操作包括:环境准备与依赖安装、预训练模型加载与结构查看、分类头替换、特征层冻结、全参数微调、数据增强策略、损失函数与优化器配置、模型保存与推理测试、FastAPI 封装接口、批量推理脚本编写,以及资源占用观察和常见问题排查。整套流程跑通之后,你手里就有一套可以横向迁移到其他视觉任务的工程模板。
1. 核心能力速览
迁移学习在计算机视觉中的应用,本质上是一套方法论加工程组合,不是某个单一的开源工具。为了让读者快速判断自己是否需要用迁移学习,下面把它拆成能力维度来看:
| 能力项 | 说明 |
|---|---|
| 核心作用 | 利用预训练模型权重加速视觉任务训练,降低数据量和算力要求 |
| 适用任务 | 图像分类、目标检测、图像分割、特征提取、相似度检索、OCR 预处理 |
| 常用预训练模型 | ResNet、VGG、EfficientNet、MobileNet、ViT、ConvNeXt 等 |
| 主流微调方式 | 冻结特征层只训练分类头、全参数微调、线性探测、渐进解冻 |
| 显存需求 | 根据模型和 batch size 变化,通常 4G 到 24G 不等 |
| 硬件要求 | 支持 CPU 推理和训练,但 GPU 能显著提升效率 |
| 开发框架 | PyTorch、Torchvision、TensorFlow、Keras |
| 数据要求 | 最少每类几十张图片即可启动微调,效果随数据量递增 |
| 接口能力 | 训练完成后可导出为 ONNX、TorchScript,或用 FastAPI 封装 |
| 批量任务 | 支持文件夹级批量推理、批量特征提取、批量预测结果导出 |
| 适合场景 | 小样本分类、工业质检、医学影像辅助分析、安防识别、风格迁移 |
这套能力组合决定了迁移学习在视觉项目里几乎是“默认起手式”。无论是做算法验证还是上线部署,直接拿预训练权重做初始化,都能省掉大量无效训练时间。
2. 适用场景与使用边界
迁移学习不是万能药,它在很多场景下表现优秀,但也有明确的适用边界。搞清楚这些边界,能避免在错误的方向上浪费时间。
2.1 适合什么场景
- 标注数据有限。工业场景里,有缺陷的样本往往很少,有些类别的图片可能只有几十张。这种情况下从零训练容易过拟合,迁移学习可以借用预训练模型已经学到的通用视觉特征,在小数据上依然得到可接受的效果。
- 任务与 ImageNet 或其他大型数据集分布接近。通用图像分类、物体识别、场景理解等任务,底层特征高度相似,迁移效果明显。
- 需要快速出原型。算法验证阶段,从预训练模型出发微调,往往几十分钟内就能看到一个 baseline 效果,方便快速评估可行性。
- 算力资源有限。加载预训练权重后,用较小的学习率微调,通常只需要单张中端 GPU 甚至 CPU 就能完成。
2.2 不适合什么场景
- 图像分布与预训练数据差异极大。比如医学 CT 影像、卫星遥感雷达图、工业 X 光图像,这类图像跟自然图像在底层特征上差异很大,直接复用 ImageNet 权重未必有明显收益,甚至可能引入噪声。
- 需要极高精度的专业领域任务。对于病灶分割、特定零部件检测等任务,单纯做分类头微调往往不够,需要额外的领域预训练或自监督学习。
- 数据集本身足够大。如果每类有上万张图片,数据分布又比较独特,从零训练或重新预训练的效果可能更好,迁移学习的优势会被削弱。
2.3 使用边界与合规要求
使用预训练模型时需要注意模型的开源协议,不同的模型有不同的商用许可限制。涉及人脸、车辆、医疗影像等敏感数据时,必须确保数据的来源合法、脱敏合规,并对模型输出进行人工复核。模型训练和部署应当限定在授权测试环境和生产环境范围内,不能未经授权采集或使用他人图像数据。
3. 环境准备与前置条件
迁移学习在计算机视觉中的工程实践,推荐使用 Python 3.8 及以上版本,搭配 PyTorch 2.x 和 Torchvision。GPU 不是硬性要求,CPU 也能跑,但训练速度和显存占用表现会有差异。
3.1 硬件环境参考
- GPU:建议 NVIDIA 显卡,显存 4G 以上,显存越高可以支持更大的 batch size 和输入分辨率。
- CPU:支持训练和推理,但速度会慢很多,适合调试和小批量验证。
- 内存:16G 以上比较稳妥。
- 磁盘空间:预训练模型权重文件加数据集,预留 20G 以上空间比较合适。
3.2 软件环境清单
| 依赖项 | 推荐版本 | 用途 |
|---|---|---|
| Python | 3.8 - 3.11 | 开发语言 |
| PyTorch | 2.0 及以上 | 深度学习框架 |
| Torchvision | 0.15 及以上 | 预训练模型和数据集工具 |
| CUDA | 11.8 或 12.1 | GPU 加速 |
| fastapi | 最新稳定版 | 接口服务封装 |
| uvicorn | 最新稳定版 | API 服务启动 |
| Pillow | 最新稳定版 | 图像读取处理 |
| tensorboard | 最新稳定版 | 训练过程可视化 |
3.3 环境安装命令
# 创建 Python 虚拟环境 python -m venv venv source venv/bin/activate # Windows 下使用 venv\Scripts\activate # 安装 PyTorch GPU 版本,请根据实际 CUDA 版本调整命令 pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 # 安装常用依赖 pip install fastapi uvicorn pillow tensorboard安装完成后,可以用下面的代码验证 GPU 是否可用:
import torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0) if torch.cuda.is_available() else "CPU mode")这里建议在开始训练前先确认 PyTorch 能正常调用 GPU。如果输出cuda.is_available()为False,需要检查显卡驱动和 CUDA 版本的匹配关系,或者在 CPU 模式下运行。
4. 安装部署与模型加载
这一部分从代码层面演示迁移学习的核心操作:加载预训练模型、替换分类头、冻结特征层。全部代码基于 PyTorch 和 Torchvision。
4.1 加载预训练模型
Torchvision 提供了很多常用的视觉模型和预训练权重,使用weights参数可以直接下载并加载。下面以 ResNet18 为例:
import torchvision.models as models # 加载 ResNet18 预训练权重 weights = models.ResNet18_Weights.DEFAULT model = models.resnet18(weights=weights) # 查看模型结构 print(model)第一次运行时,PyTorch 会自动下载权重文件到本机缓存目录。下载完成后,后续加载会直接使用缓存,不再重复下载。权重文件默认存放在用户目录下的.cache/torch/hub/checkpoints文件夹,Windows 下一般在C:\Users\<用户名>\.cache\torch\hub\checkpoints。
除了 ResNet18,Torchvision 还支持多种常用模型:
| 模型 | 参数量 | 特点 | 适用场景 |
|---|---|---|---|
| ResNet18/34/50 | 11M - 25M | 精度和速度平衡 | 通用分类、特征提取 |
| MobileNetV3 | 4M - 6M | 轻量级 | 移动端、嵌入式设备 |
| EfficientNet-B0 | 5M | 高精度高效率 | 资源受限场景 |
| ViT-B/16 | 86M | Transformer 架构 | 大规模数据微调 |
| ConvNeXt | 28M - 350M | 现代 CNN | 高精度任务 |
| Swin Transformer | 28M - 88M | 层级注意力 | 检测、分割 |
模型选择的核心原则是:先用轻量级模型跑通流程,再根据精度需求逐步换更大的模型。不要一上来就用最大模型,否则显存和训练时间都会失控。
4.2 替换分类头
预训练模型默认输出是 1000 类,对应 ImageNet 的分类任务。如果我们的任务是二分类或者自定义类别数,需要把最后一层全连接层替换掉。
import torch.nn as nn import torchvision.models as models weights = models.ResNet18_Weights.DEFAULT model = models.resnet18(weights=weights) # 获取模型的特征提取层输出维度 num_features = model.fc.in_features print(f"特征维度: {num_features}") # 替换全连接层,输出自定义类别数 num_classes = 10 # 这里用 10 分类做示例 model.fc = nn.Linear(num_features, num_classes) print(model.fc)替换分类头是整个迁移学习流程中最重要的步骤之一。无论原来的模型结构多复杂,只需要修改这一层,就能让预训练特征复用到新的分类任务上。
4.3 冻结特征层
冻结特征层的意思是让预训练提取特征的参数在反向传播时不更新,只训练新加的分类头。这种方式适合数据量很少的情况,可以显著降低过拟合风险,同时减少训练时间。
# 先冻结所有层 for param in model.parameters(): param.requires_grad = False # 只让分类头参与训练 for param in model.fc.parameters(): param.requires_grad = True上面这种方式就是经典的“冻结特征层 + 训练分类头”,在迁移学习中被称为线性探测的一种变体。好处是训练速度快、不容易过拟合,缺点是由于底层特征不做适配,精度上限可能稍低一些。
如果想做全参数微调,也就是让所有层都参与训练,那么不使用上面的冻结逻辑,直接定义优化器时传入所有参数即可。全参数微调适合数据量比较充足的情况,效果通常更好。
5. 功能测试与效果验证
迁移学习最终要落在真实数据集上才能验证效果。这一节用 CIFAR-10 数据集演示完整的微调流程,包括数据加载、训练、验证和推理。
5.1 准备数据集
为了方便演示,直接使用 Torchvision 内置的 CIFAR-10 数据集。如果是自定义数据,只需要用ImageFolder方式加载目录数据即可。
import torch import torchvision import torchvision.transforms as transforms from torch.utils.data import DataLoader # 数据预处理:训练集和验证集使用不同的增强策略 transform_train = transforms.Compose([ transforms.RandomResizedCrop(224), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) transform_val = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # 加载 CIFAR-10 train_dataset = torchvision.datasets.CIFAR10( root='./data', train=True, download=True, transform=transform_train ) val_dataset = torchvision.datasets.CIFAR10( root='./data', train=False, download=True, transform=transform_val ) train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=2) val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=2)这里的核心点在于Normalize层使用的 mean 和 std 是 ImageNet 数据集的统计值。因为预训练模型是在 ImageNet 上训练的,数据标准化应该使用相同的统计值,否则预训练特征分布会被破坏。
5.2 训练分类头
下面演示冻结特征层,只训练分类头的方式。这种方式在少量数据上非常稳定。
import torch.nn as nn import torch.optim as optim import torchvision.models as models from tqdm import tqdm device = torch.device("cuda" if torch.cuda.is_available() else "cpu") # 加载预训练模型 weights = models.ResNet18_Weights.DEFAULT model = models.resnet18(weights=weights) num_features = model.fc.in_features model.fc = nn.Linear(num_features, 10) model = model.to(device) # 冻结特征层 for param in model.parameters(): param.requires_grad = False for param in model.fc.parameters(): param.requires_grad = True # 损失函数和优化器 criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.fc.parameters(), lr=0.001) # 训练分类头 epochs = 5 for epoch in range(epochs): model.train() running_loss = 0.0 correct = 0 total = 0 pbar = tqdm(train_loader, desc=f"Epoch {epoch + 1}/{epochs}") for images, labels in pbar: images = images.to(device) labels = labels.to(device) optimizer.zero_grad() outputs = model(images) loss = criterion(outputs, labels) loss.backward() optimizer.step() running_loss += loss.item() _, predicted = torch.max(outputs, 1) total += labels.size(0) correct += (predicted == labels).sum().item() pbar.set_postfix(loss=loss.item(), acc=100.0 * correct / total) epoch_acc = 100.0 * correct / total epoch_loss = running_loss / len(train_loader) print(f"Epoch {epoch + 1}: loss={epoch_loss:.4f}, acc={epoch_acc:.2f}%")这段代码直接在训练过程中打印 loss 和 accuracy,可以在终端直观看到模型收敛情况。由于只训练分类头,参数数量很少,训练速度非常快。在 CIFAR-10 上,即使只训练 5 个 epoch,准确率通常也能达到比较可观的水平。
5.3 评估验证集
训练完成后,需要在验证集上评估模型效果,避免过拟合判断出错。
# 验证模型 model.eval() correct = 0 total = 0 with torch.no_grad(): for images, labels in val_loader: images = images.to(device) labels = labels.to(device) outputs = model(images) _, predicted = torch.max(outputs, 1) total += labels.size(0) correct += (predicted == labels).sum().item() val_acc = 100.0 * correct / total print(f"验证集准确率: {val_acc:.2f}%")5.4 自定义数据集加载
如果使用自己的目录数据集,目录结构需要满足下面的格式:
data/ train/ class1/ 1.jpg 2.jpg class2/ 1.jpg 2.jpg val/ class1/ 1.jpg class2/ 1.jpg对应的加载方式:
from torchvision.datasets import ImageFolder train_dataset = ImageFolder(root='data/train', transform=transform_train) val_dataset = ImageFolder(root='data/val', transform=transform_val) print(train_dataset.classes) print(train_dataset.class_to_idx)ImageFolder会按照子目录名称自动生成类别标签,使用起来非常方便。分类头的输出类别数需要根据len(train_dataset.classes)动态设置。
5.5 推理测试
训练完成后,对单张图片做推理测试的完整代码:
from PIL import Image import torchvision.transforms as transforms # 图片预处理 def preprocess_image(image_path): transform = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) image = Image.open(image_path).convert("RGB") return transform(image).unsqueeze(0) # 推理 model.eval() image_tensor = preprocess_image("test.jpg").to(device) with torch.no_grad(): outputs = model(image_tensor) _, predicted = torch.max(outputs, 1) print(f"预测类别索引: {predicted.item()}")如果想知道具体的类别名称,可以用val_dataset.classes[predicted.item()]映射回原始类别名。
6. 接口 API 与批量任务
模型训练完成之后,实际工程中很少直接在训练脚本里做推理,更多是把模型封装成接口服务或者批量推理脚本。这一部分演示两种方式:FastAPI 封装和批量推理。
6.1 FastAPI 封装模型接口
封装成 HTTP 接口之后,其他系统可以通过请求直接调用模型,便于集成到业务系统里。
import io import torch import torchvision.models as models from fastapi import FastAPI, UploadFile, File from PIL import Image import torchvision.transforms as transforms app = FastAPI() # 定义设备 device = torch.device("cuda" if torch.cuda.is_available() else "cpu") # 加载训练好的模型 model = models.resnet18(weights=models.ResNet18_Weights.DEFAULT) num_features = model.fc.in_features model.fc = torch.nn.Linear(num_features, 10) model.load_state_dict(torch.load("resnet18_cifar10.pth")) model = model.to(device) model.eval() # 类别名称 classes = ["airplane", "automobile", "bird", "cat", "deer", "dog", "frog", "horse", "ship", "truck"] # 图像预处理 transform = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) @app.post("/predict") async def predict(file: UploadFile = File(...)): image = Image.open(io.BytesIO(await file.read())).convert("RGB") image_tensor = transform(image).unsqueeze(0).to(device) with torch.no_grad(): outputs = model(image_tensor) probabilities = torch.nn.functional.softmax(outputs[0], dim=0) confidence, predicted = torch.max(probabilities, 0) return { "class_id": predicted.item(), "class_name": classes[predicted.item()], "confidence": round(confidence.item(), 4) } # 启动: uvicorn main:app --host 0.0.0.0 --port 8000业务系统可以直接用requests调用这个接口:
import requests url = "http://127.0.0.1:8000/predict" files = {"file": open("test.jpg", "rb")} response = requests.post(url, files=files, timeout=30) print(response.json())返回结果示例:
{ "class_id": 3, "class_name": "cat", "confidence": 0.9234 }接口封装的核心思路很简单:加载模型、预处理输入、推理、返回结果。生产环境还需要考虑并发限制、超时设置和日志记录,这里给出的是最小可运行版本。
6.2 批量推理脚本
批量推理是视觉任务中最高频的需求之一。比如给一个文件夹里几千张图片做分类,需要写一个脚本自动遍历图片、执行推理、保存结果。
import os import csv import torch import torchvision.models as models from PIL import Image import torchvision.transforms as transforms device = torch.device("cuda" if torch.cuda.is_available() else "cpu") # 加载模型 model = models.resnet18(weights=models.ResNet18_Weights.DEFAULT) num_features = model.fc.in_features model.fc = torch.nn.Linear(num_features, 10) model.load_state_dict(torch.load("resnet18_cifar10.pth")) model = model.to(device) model.eval() transform = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) def predict_image(model, image_path, device): image = Image.open(image_path).convert("RGB") image_tensor = transform(image).unsqueeze(0).to(device) with torch.no_grad(): outputs = model(image_tensor) _, predicted = torch.max(outputs, 1) return predicted.item() # 批量处理 input_dir = "./images" output_file = "./predictions.csv" results = [] for filename in os.listdir(input_dir): if filename.lower().endswith((".jpg", ".jpeg", ".png")): image_path = os.path.join(input_dir, filename) try: class_id = predict_image(model, image_path, device) results.append([filename, class_id]) print(f"{filename} -> class {class_id}") except Exception as e: print(f"处理 {filename} 失败: {e}") # 保存结果 with open(output_file, "w", newline="") as f: writer = csv.writer(f) writer.writerow(["filename", "class_id"]) writer.writerows(results) print(f"结果已保存到 {output_file}")批量推理的关键在于容错处理。单张图片损坏、格式异常或者网络中断,都不应该导致整个批处理任务崩溃。上面代码中加入了try-except结构,单张图片失败时记录日志并继续处理下一张。
6.3 批量任务设计建议
如果图片数量很大,比如几万张,需要考虑以下优化策略:
- 使用 DataLoader 替代单张循环,利用
batch_size提升吞吐。 - 将推理结果分批写入 CSV 或数据库,而不是全部攒在内存里。
- 多进程并行处理不同文件夹或分片。
- 推理前先做图片大小压缩和去重,减少无效计算。
- 设置失败重试机制,对超时或异常任务重新处理。
7. 资源占用与性能观察
迁移学习虽然在训练效率上远胜从零训练,但依然需要关注资源占用。这一部分给出观察方法和优化思路。
7.1 显存占用观察
在训练过程中,可以用 NVIDIA 提供的命令实时查看显存占用:
nvidia-smi如果希望更精确地监控 PyTorch 的显存分配情况,可以在代码里加入以下逻辑:
import torch # 查看当前 PyTorch 缓存占用的显存 print(f"allocated: {torch.cuda.memory_allocated() / 1024 ** 2:.2f} MB") print(f"cached: {torch.cuda.memory_cached() / 1024 ** 2:.2f} MB")显存占用的主要影响因素包括:
- 输入图片的分辨率,分辨率越高,中间特征图占用的显存越大。
- batch size 越大,显存占用线性增长。
- 模型的层数和宽度,ResNet50 比 ResNet18 显存占用高很多。
- 是否使用混合精度训练,使用 AMP 可以降低显存占用。
7.2 CPU 与 GPU 训练差异
CPU 可以完成迁移学习的完整流程,但速度会慢很多。以 ResNet18 在 CIFAR-10 上训练为例,GPU 训练一个 epoch 可能只需要几十秒,CPU 可能要好几分钟甚至十几分钟,具体取决于 CPU 核心数和内存带宽。
如果 CPU 资源有限,建议降低输入分辨率、减小 batch size、简化数据增强,或者直接使用 MobileNet 这类轻量模型。
7.3 降低显存占用的常用方法
- 降低分辨率。从 224 降到 160 或 128,显存占用会大幅下降。
- 减小 batch size。一次处理 8 张和一次处理 32 张,显存占用差异非常明显。
- 使用混合精度训练。PyTorch 的
torch.cuda.amp可以将部分计算转为 FP16,显存占用和训练速度都有改善。 - 使用梯度累积。在小 batch size 下模拟较大 batch size 的训练效果。
- 选择更轻量的模型。MobileNetV3 明显比 ResNet50 省显存。
7.4 混精度训练示例
from torch.cuda.amp import autocast, GradScaler scaler = GradScaler() for epoch in range(epochs): model.train() for images, labels in train_loader: images = images.to(device) labels = labels.to(device) optimizer.zero_grad() with autocast(): outputs = model(images) loss = criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()混合精度训练在 NVIDIA 显卡上通常能带来明显加速和显存节省,尤其是 Turing 架构及以上的 GPU 效果更明显。
8. 常见问题与排查方法
迁移学习在视觉任务中踩坑的地方不少,下面整理了一份高频率问题清单。
| 问题现象 | 可能原因 | 排查方式 | 解决方案 |
|---|---|---|---|
| 模型加载时报错 | 预训练权重版本与模型结构不匹配 | 检查 torchvision 版本和 weights 参数 | 升级 torchvision 或改用不带 weights 加载后手动加载权重 |
| 训练时显存溢出 | batch size 太大或分辨率过高 | 用 nvidia-smi 查看显存占用 | 减小 batch size、降低分辨率、使用梯度累积 |
| CPU 训练速度极慢 | 数据加载线程数不足或模型太大 | 查看 CPU 利用率和内存占用 | 增加 num_workers,使用轻量模型 |
| 验证集准确率远低于训练集 | 过拟合 | 对比训练集和验证集准确率 | 增加数据增强、降低模型复杂度、增加 dropout |
| 图像分类结果全是某一类 | 类别不均衡或学习率过大 | 打印每个类别的预测分布 | 使用类别权重、降低学习率、检查数据标签 |
| 模型输出 NaN 丢失 | 学习率过高或数据有异常 | 打印 loss 曲线 | 降低学习率、检查数据标准化、清除异常数据 |
| 接口服务请求超时 | 推理时间过长或并发过高 | 检查单张推理耗时和并发连接数 | 使用异步推理、增加 batch 推理、横向扩展服务 |
| 换数据集后模型效果很差 | 新数据分布与预训练数据差异大 | 检查数据标准化和类别分布 | 使用更小的学习率、全参数微调、或做领域预训练 |
| 冻结层后效果不如预期 | 分类头能力有限 | 检查底层特征是否适配 | 改为全参数微调或渐进解冻 |
| 预训练权重下载失败 | 网络问题或镜像站不可达 | 检查网络连通性和缓存目录 | 手动下载权重文件放到缓存目录 |
对于多分类任务中类别不均衡的问题,可以在损失函数中传入权重:
class_counts = [5000, 500, 50] # 每个类别的样本数量 total_samples = sum(class_counts) class_weights = [total_samples / count for count in class_counts] class_weights = torch.tensor(class_weights, device=device) criterion = nn.CrossEntropyLoss(weight=class_weights)这种方式会提高小类别样本在损失函数中的权重,从而缓解模型偏向大类别的问题。
9. 最佳实践与使用建议
迁移学习的工程化落地,建议从一开始就建立起一套规范的流程,避免后期返工。
第一,第一次跑通流程时使用小数据集、小模型和较小分辨率。比如先拿每类 100 张图片、ResNet18、128 分辨率跑通全流程,确认代码没有问题后再扩展到完整数据和更大模型。这样能大幅缩短调试时间,也让资源占用保持在可控范围。
第二,数据目录、模型权重和训练日志要分目录管理。推荐的目录结构如下:
project/ data/ train/ val/ test/ models/ checkpoints/ logs/ scripts/ train.py predict.py api.py output/ predictions.csv第三,批量任务必须加日志和失败重试机制。实际项目中,图片文件损坏、格式异常、网络中断等情况经常发生。代码里要有完整的try-except结构和日志记录,单张图片失败不能导致整个任务终止。
第四,接口服务要限制访问范围。如果是内网使用,可以绑定内网 IP 启动服务;如果需要暴露到公网,必须在前面加鉴权层,防止接口被滥用。
# 只监听本机地址,不暴露到外部网络 uvicorn main:app --host 127.0.0.1 --port 8000第五,涉及人脸、车辆、医疗影像等数据时,必须确保数据来源合法,已完成脱敏处理,并确认符合相关法规要求。模型输出只能作为辅助判断,不能直接替代人的决策。
第六,每次训练都要保存模型结构和超参数配置。建议使用配置字典或 YAML 文件记录 batch size、学习率、优化器、数据增强方式等信息,方便后续复现实验。
# config.yaml 示例 model: name: resnet18 pretrained: true num_classes: 10 training: batch_size: 32 epochs: 10 learning_rate: 0.001 freeze_backbone: true mixed_precision: true data: train_dir: data/train val_dir: data/val image_size: 224第七,模型发布前要做效果复核。可以准备一份独立的测试集,验证模型在未见数据上的表现。如果模型对某些特定场景表现不佳,需要针对性地补充数据或者调整模型结构。
10. 总结与下一步
迁移学习在计算机视觉里的实用价值非常直接:用预训练模型省时间、省数据、省算力,同时还能得到不错的精度。这篇文章从环境准备、模型加载、分类头替换、特征冻结、全参数微调、自定义数据集训练、API 封装、批量推理、资源占用和问题排查,走完了一整套视觉任务的迁移学习流程。
最先应该验证的功能是预训练模型加载和分类头替换,这两步跑通之后,后续的整套流程都能顺畅推进。最容易踩的坑是数据标准化不一致、分类头输出类别数不匹配、显存溢出这三个点,遇到问题优先检查这三项。
后续可以继续扩展的方向包括:引入更强的基础模型如 ConvNeXt 或 ViT、使用更大的输入分辨率、加入数据增强策略如 MixUp 和 CutMix、在目标检测和语义分割任务中复现同样的迁移学习思路,以及把模型导出为 ONNX 格式部署到生产环境。迁移学习不仅是一个方法,更是一套可以复用的工程模板,值得花时间搭好。