基于深度学习的语义分割:从理论到实践(MATLAB 版详解)
基于深度学习的语义分割:从理论到实践(MATLAB 版详解)
本文基于 MathWorks 官方示例,系统讲解语义分割的核心概念、DeepLab v3+ 网络结构、CamVid 数据集处理、训练与评估全流程。适合希望快速上手语义分割的开发者与研究者。
📌 目录
- 1. 语义分割是什么
- 2. 为什么选择 DeepLab v3+ 与 ResNet-18
- 3. 数据集:CamVid 详解与预处理
- 4. 网络构建与类平衡处理
- 5. 训练策略与超参数调优
- 6. 模型评估与指标解读
- 7. 代码逐段剖析(含重点函数)
- 8. 常见问题与优化建议
- 9. 总结与参考资料
1. 语义分割是什么
语义分割(Semantic Segmentation) 是计算机视觉中的核心任务,目标是对图像中的每个像素分配一个类别标签,从而产生“像素级”的分类结果。与目标检测(边界框)和图像分类(整体标签)不同,语义分割能够精细区分物体轮廓、遮挡关系,广泛应用于:
- 自动驾驶:区分道路、车辆、行人、交通标志等;
- 医学影像:分割肿瘤、器官、细胞等;
- 遥感分析:识别土地覆盖类型、建筑物等。
本示例使用 DeepLab v3+ 网络,它是一种基于空洞卷积(Atrous Convolution)的编码器-解码器结构,能够在保持高分辨率的同时扩大感受野,获得精细的分割边界。
2. 为什么选择 DeepLab v3+ 与 ResNet-18
2.1 DeepLab v3+ 核心亮点
- 编码器-解码器架构:编码器使用空洞卷积提取多尺度特征,解码器逐步恢复空间分辨率,弥补下采样造成的位置信息损失。
- Atrous Spatial Pyramid Pooling (ASPP):通过不同膨胀率的并行空洞卷积,捕捉多个尺度的上下文信息。
- 深度可分离卷积:在 ASPP 和 decoder 中采用,大幅减少参数量和计算量。
2.2 ResNet-18 作为骨干网络
- 轻量级残差网络,18 层深度,适合资源受限场景(如嵌入式设备)。
- 在 ImageNet 上预训练,提供了良好的特征提取能力,迁移学习效果好。
- 也可替换为 ResNet-50、MobileNet v2 等,依据应用场景权衡精度与速度。
3. 数据集:CamVid 详解与预处理
3.1 CamVid 数据集简介
- 来自剑桥大学,包含 701 张 街道驾驶场景图像(白天、黄昏)。
- 原始标注 32 个语义类别,本示例将其合并为 11 个常用类别(见表 1)。
| 合并后类别 | 原始包含类别(部分) |
|---|---|
| Sky | Sky |
| Building | Bridge, Building, Wall, Tunnel, Archway |
| Pole | Column_Pole, TrafficCone |
| Road | Road, LaneMkgsDriv, LaneMkgsNonDriv |
| Pavement | Sidewalk, ParkingBlock, RoadShoulder |
| Tree | Tree, VegetationMisc |
| SignSymbol | SignSymbol, Misc_Text, TrafficLight |
| Fence | Fence |
| Car | Car, SUVPickupTruck, Truck_Bus, Train, OtherMoving |
| Pedestrian | Pedestrian, Child, CartLuggagePram, Animal |
| Bicyclist | Bicyclist, MotorcycleScooter |
📂 数据下载:图像压缩包约 557 MB,标签压缩包约 16 MB。示例中通过
websave自动下载到临时目录。
3.2 数据加载与统计分析
使用 imageDatastore 加载原始图像,使用 pixelLabelDatastore 加载像素标签图像。标签图像以 RGB 颜色编码,通过 camvidPixelLabelIDs 函数映射到 11 个类别 ID。
类别分布分析:调用 countEachLabel 统计每个类别的像素数,并绘制频率直方图。结果显示 道路、天空、建筑 占主导,而 行人、骑车人 像素极少——这是典型的类别不平衡问题,需在训练时通过加权损失函数缓解。
4. 网络构建与类平衡处理
4.1 创建 DeepLab v3+ 网络
network = deeplabv3plus(imageSize, numClasses, "resnet18");
imageSize = [720, 960, 3](高度×宽度×通道)numClasses = 11- 该函数自动加载预训练的 ResNet-18 作为编码器,并构建完整的 DeepLab v3+ 结构。
4.2 类别加权(Class Weighting)
为了对抗类别不平衡,采用 中值频率平衡(Median Frequency Balancing):
imageFreq = tbl.PixelCount ./ tbl.ImagePixelCount; % 每个类别的图像级频率
classWeights = median(imageFreq) ./ imageFreq; % 权重 = 中位数 / 各自频率
- 频率较低的类别(如 Pedestrian)获得较大权重,迫使网络更加关注它们。
- 在损失函数中,每个像素的损失乘以其类别的权重。

5. 训练策略与超参数调优
5.1 数据划分与增强
- 训练集:60%(421 张)
- 验证集:20%(140 张)
- 测试集:20%(140 张)
数据增强(仅应用于训练集):
- 随机水平翻转(左右镜像)
- 随机平移(X/Y 方向 ±10 像素)
- 使用
transform结合augmentImageAndLabel函数,确保图像与标签同步变换。
5.2 训练选项(trainingOptions)
| 参数 | 值 | 说明 |
|---|---|---|
| 优化器 | SGDM | 带动量的随机梯度下降 |
| 初始学习率 | 1e-2 | 较高起步,加速收敛 |
| 学习率调度 | piecewise | 每 6 个 epoch 乘以 0.1 |
| 动量 | 0.9 | 常用值 |
| L2 正则化 | 0.005 | 防止过拟合 |
| MaxEpochs | 18 | 总训练轮次 |
| MiniBatchSize | 4 | 根据 GPU 内存调整(可降至 1) |
| 验证数据 | dsVal |
每 epoch 结束后评估 |
| ValidationPatience | 4 | 验证损失连续 4 轮不下降则停止训练 |
| CheckpointPath | tempdir | 每 epoch 保存中间模型,便于恢复 |
5.3 损失函数
使用交叉熵损失,并增加掩码处理(忽略无标签像素)和类别权重:
function loss = modelLoss(Y,T,classWeights)
weights = dlarray(classWeights,"C");
mask = ~isnan(T); % 标签为 NaN 的像素不参与损失
T(isnan(T)) = 0;
loss = crossentropy(Y,T,weights,Mask=mask,NormalizationFactor="mask-included");
end
6. 模型评估与指标解读
6.1 单张图像测试
使用 semanticseg 对测试集单张图像进行预测,并利用 labeloverlay 可视化叠加结果。对比预测与真实标签(imshowpair 显示差异),可见道路、天空等大目标分割良好,行人、自行车等小目标误差较大。
6.2 交并比(IoU)
IoU(Intersection over Union) 是语义分割最常用的指标:
IoU = (预测 ∩ 真实) / (预测 ∪ 真实)
单张图像各类别 IoU 示例:
- Road: 0.953
- Sky: 0.936
- Tree: 0.926
- Pedestrian: 0.267 (明显偏低)
6.3 整体评估
使用 evaluateSemanticSegmentation 对整个测试集(140 张)计算指标,输出数据集级指标:
| 指标 | 值 | 含义 |
|---|---|---|
| Global Accuracy | 0.9075 | 正确分类的像素占比 |
| Mean Accuracy | 0.8883 | 各类别准确率的平均值 |
| Mean IoU | 0.6957 | 各类别 IoU 的平均值 |
| Weighted IoU | 0.8490 | 按像素数加权的 IoU |
| Mean BF Score | 0.7431 | 边界 F1 分数(关注边缘质量) |
类别级 IoU 显示:Bicyclist (0.677)、Pedestrian (0.505) 仍较低,说明需要更多该类样本或更高级的特征提取。
7. 代码逐段剖析(含重点函数)
7.1 下载预训练模型(直接应用)
若只想快速推理,可直接下载已训练好的 DeepLab v3+(基于 CamVid):
pretrainedURL = "https://ssd.mathworks.com/.../deeplabv3plusResnet18CamVid_v2.zip";
unzip(...);
load('deeplabv3plusResnet18CamVid_v2.mat');
C = semanticseg(I, net);
7.2 数据集分区函数 partitionCamVidData
- 使用
randperm打乱图像索引; - 按 60/20/20 分割,返回对应的
imageDatastore和pixelLabelDatastore。
7.3 数据增强函数 augmentImageAndLabel
- 构造
affine2d变换矩阵,包含反射和平移; - 使用
affineOutputView确保输出图像尺寸固定; - 对图像和标签应用相同
imwarp变换。
7.4 颜色映射和可视化
camvidColorMap返回 11×3 的 RGB 矩阵(归一化到 [0,1]);pixelLabelColorbar生成带有类别名的彩条,便于查看图例。
7.5 训练主流程(trainnet)
[net, info] = trainnet(dsTrain, network, @(Y,T) modelLoss(Y,T,classWeights), options);
dsTrain是经过增强的组合数据存储(图像 + 标签);- 自定义损失函数句柄传递类别权重;
info返回训练历史(损失、准确率等)。
8. 常见问题与优化建议
8.1 GPU 内存不足(OOM)
- 解决方案:
- 减小
MiniBatchSize(如从 4 降至 1); - 降低输入图像尺寸(如 512×512),同时调整
imageSize并重新缩放训练数据; - 使用
'ExecutionEnvironment', 'cpu'但训练速度会大幅下降。
- 减小
8.2 训练不收敛或震荡
- 检查学习率是否过大(可尝试 1e-3);
- 检查类别权重是否合理(极端权重可能导致梯度爆炸);
- 增加
ValidationPatience或关闭早停。
8.3 小目标分割效果差
- 增加数据增强(如随机裁剪、旋转、缩放);
- 尝试更深的骨干网络(ResNet-50)或更复杂的解码器;
- 使用 边界损失 或 focal loss 来强化难例。
8.4 如何迁移到自己的数据集
- 替换数据存储(
imageDatastore+pixelLabelDatastore); - 修改
classes和labelIDs映射; - 根据数据规模调整训练轮次和 batch size;
- 若数据较少,可冻结骨干网络部分层,仅训练分类头。
9. 总结与参考资料
9.1 关键 takeaways
- 语义分割是像素级分类任务,DeepLab v3+ 是目前主流框架之一。
- 类别不平衡是真实场景的常态,需通过加权损失或重采样处理。
- 数据增强、预训练权重、合理超参数是提升性能的三大法宝。
- 评估指标应结合全局 IoU 和类别 IoU 综合判断。
9.2 相关工具箱与函数
- Computer Vision Toolbox:
semanticseg,labeloverlay,pixelLabelDatastore - Deep Learning Toolbox:
trainnet,trainingOptions,dlnetwork - Parallel Computing Toolbox:GPU 加速支持
9.3 参考文献
- Chen, L.-C., et al. “Encoder-Decoder with Atrous Separable Convolution for Semantic Image Segmentation.” ECCV 2018.
- Brostow, G. J., et al. “Semantic object classes in video: A high-definition ground truth database.” Pattern Recognition Letters, 2009.
- MathWorks 官方文档:Semantic Segmentation Using Deep Learning
✍️ 本文是对 MathWorks 官方示例的深度解读,所有代码均可在 MATLAB R2024a+ 环境中运行。如需完整项目文件,可访问 MATLAB File Exchange 搜索相关示例。
完整代码:基于 DeepLab v3+ 的语义分割(MATLAB)
以下提供一份完整、可直接运行的 MATLAB 脚本,包含:
- 自动下载预训练 DeepLab v3+ 网络(基于 CamVid)并进行单张图像分割演示;
- 可选下载 CamVid 数据集并从头训练(需 GPU,默认关闭);
- 训练后评估模型性能(IoU、全局准确率等);
- 所有必要的辅助函数(数据划分、增强、颜色映射、损失函数等)。





📋 环境要求
- MATLAB R2020b 及以上(推荐 R2024a)
- 工具箱:
- Computer Vision Toolbox
- Deep Learning Toolbox
- Parallel Computing Toolbox(推荐 GPU 加速)
- 支持包(会自动提示下载):
- Deep Learning Toolbox Model for ResNet-18 Network
🚀 运行方式
- 将以下代码完整复制到一个新建的 MATLAB 脚本文件(例如
SemanticSegmentationDemo.m)中。 - 直接运行脚本。
- 首次运行会下载预训练模型(约 58 MB)和 CamVid 数据集(图像 557 MB + 标签 16 MB),请确保网络通畅且有足够磁盘空间。
- 若希望从头训练,将脚本开头的
doTraining = false;改为true(需 GPU 且显存 ≥ 4 GB,训练耗时约 50 分钟)。
📝 完整代码
%% 基于深度学习的语义分割 —— DeepLab v3+ 完整示例
% 该脚本演示了:
% 1) 使用预训练网络对图像进行语义分割;
% 2) 下载 CamVid 数据集,进行数据增强、类别加权,并训练 DeepLab v3+;
% 3) 在测试集上评估模型性能(IoU、准确率等)。
%
% 作者: MathWorks 官方示例 (改编)
% 最后更新: 2026-07-19
clear; clc; close all;
%% ========== 1. 参数设置 ==========
% 是否执行训练(若为 false,则只运行预训练模型推理演示)
doTraining = false; % 改为 true 以启动训练
% 数据集存储路径(默认使用临时目录,可自定义)
outputFolder = fullfile(tempdir, 'CamVid');
pretrainedFolder = fullfile(tempdir, 'pretrainedNetwork');
% 网络输入尺寸(与预训练模型一致)
imageSize = [720 960 3];
numClasses = 11; % 合并后的类别数
%% ========== 2. 下载并加载预训练 DeepLab v3+ 网络 ==========
% 预训练模型 URL(基于 CamVid 训练)
pretrainedURL = "https://ssd.mathworks.com/supportfiles/vision/data/deeplabv3plusResnet18CamVid_v2.zip";
pretrainedNetworkZip = fullfile(pretrainedFolder, 'deeplabv3plusResnet18CamVid_v2.zip');
if ~exist(pretrainedNetworkZip, 'file')
mkdir(pretrainedFolder);
disp('正在下载预训练网络 (58 MB) ...');
websave(pretrainedNetworkZip, pretrainedURL);
end
unzip(pretrainedNetworkZip, pretrainedFolder);
% 加载网络
load(fullfile(pretrainedFolder, 'deeplabv3plusResnet18CamVid_v2.mat'), 'net');
classes = getClassNames(); % 获取 11 个类别名称
disp('预训练网络加载完成!');
%% ========== 3. 使用预训练网络进行单张图像分割 ==========
% 读取测试图像(示例图像来自工具箱)
I = imread('parkinglot_left.png');
% 调整至网络输入尺寸
I = imresize(I, imageSize(1:2));
% 执行语义分割
C = semanticseg(I, net);
% 可视化
cmap = camvidColorMap();
B = labeloverlay(I, C, 'Colormap', cmap, 'Transparency', 0.4);
figure('Name', '预训练网络分割结果');
imshow(B);
pixelLabelColorbar(cmap, classes);
title('预训练 DeepLab v3+ 分割效果');
% 若仅演示推理,至此可以结束,但若继续训练则进入下一步。
if ~doTraining
disp('预训练网络演示结束。若需训练,请将 doTraining 设为 true。');
return;
end
%% ========== 4. 准备训练数据:下载 CamVid 数据集 ==========
disp('开始下载 CamVid 数据集 ...');
imageURL = "http://web4.cs.ucl.ac.uk/staff/g.brostow/MotionSegRecData/files/701_StillsRaw_full.zip";
labelURL = "http://web4.cs.ucl.ac.uk/staff/g.brostow/MotionSegRecData/data/LabeledApproved_full.zip";
labelsZip = fullfile(outputFolder, 'labels.zip');
imagesZip = fullfile(outputFolder, 'images.zip');
if ~exist(labelsZip, 'file') || ~exist(imagesZip, 'file')
mkdir(outputFolder);
disp('正在下载标签 (16 MB) ...');
websave(labelsZip, labelURL);
unzip(labelsZip, fullfile(outputFolder, 'labels'));
disp('正在下载图像 (557 MB) ...');
websave(imagesZip, imageURL);
unzip(imagesZip, fullfile(outputFolder, 'images'));
end
%% ========== 5. 加载图像和像素标签数据 ==========
imgDir = fullfile(outputFolder, 'images', '701_StillsRaw_full');
imds = imageDatastore(imgDir);
% 获取类别名称和对应的标签 ID
labelIDs = camvidPixelLabelIDs(); % 将原始 32 类映射到 11 类
labelDir = fullfile(outputFolder, 'labels');
pxds = pixelLabelDatastore(labelDir, classes, labelIDs);
% 显示一张标注样例
I_sample = readimage(imds, 559);
C_sample = readimage(pxds, 559);
B_sample = labeloverlay(I_sample, C_sample, 'Colormap', cmap);
figure('Name', 'CamVid 标注样例');
imshow(B_sample);
pixelLabelColorbar(cmap, classes);
title('CamVid 数据集样例(像素级标注)');
%% ========== 6. 分析类别分布并计算权重 ==========
tbl = countEachLabel(pxds);
imageFreq = tbl.PixelCount ./ tbl.ImagePixelCount;
classWeights = median(imageFreq) ./ imageFreq; % 中值频率平衡
% 绘制频率直方图
figure('Name', '类别频率分布');
frequency = tbl.PixelCount / sum(tbl.PixelCount);
bar(1:numClasses, frequency);
xticks(1:numClasses);
xticklabels(tbl.Name);
xtickangle(45);
ylabel('像素频率');
title('各类别像素占比(严重不平衡)');
%% ========== 7. 划分训练、验证、测试集 (60/20/20) ==========
[imdsTrain, imdsVal, imdsTest, pxdsTrain, pxdsVal, pxdsTest] = partitionCamVidData(imds, pxds);
fprintf('训练集: %d 张, 验证集: %d 张, 测试集: %d 张\n', ...
numel(imdsTrain.Files), numel(imdsVal.Files), numel(imdsTest.Files));
% 组合验证数据
dsVal = combine(imdsVal, pxdsVal);
%% ========== 8. 数据增强(仅训练集) ==========
% 随机水平翻转和 ±10 像素平移
xTrans = [-10 10];
yTrans = [-10 10];
dsTrain = combine(imdsTrain, pxdsTrain);
dsTrain = transform(dsTrain, @(data) augmentImageAndLabel(data, xTrans, yTrans));
%% ========== 9. 创建 DeepLab v3+ 网络 ==========
% 基于 ResNet-18 预训练权重
network = deeplabv3plus(imageSize, numClasses, 'resnet18');
disp('网络构建完成!');
%% ========== 10. 设置训练选项 ==========
options = trainingOptions('sgdm', ...
'LearnRateSchedule', 'piecewise', ...
'LearnRateDropPeriod', 6, ...
'LearnRateDropFactor', 0.1, ...
'Momentum', 0.9, ...
'InitialLearnRate', 1e-2, ...
'L2Regularization', 0.005, ...
'ValidationData', dsVal, ...
'MaxEpochs', 18, ...
'MiniBatchSize', 4, ...
'Shuffle', 'every-epoch', ...
'CheckpointPath', tempdir, ...
'VerboseFrequency', 10, ...
'ValidationPatience', 4, ...
'Plots', 'training-progress'); % 显示训练曲线
%% ========== 11. 开始训练 ==========
disp('开始训练 ... 这可能需要较长时间,请耐心等待。');
[net, info] = trainnet(dsTrain, network, ...
@(Y,T) modelLoss(Y,T,classWeights), options);
disp('训练完成!');
%% ========== 12. 在测试集上评估 ==========
% 对整个测试集进行分割
disp('正在测试集上进行推理 ...');
pxdsResults = semanticseg(imdsTest, net, ...
'Classes', classes, ...
'MiniBatchSize', 4, ...
'WriteLocation', tempdir, ...
'Verbose', false);
% 计算评估指标
metrics = evaluateSemanticSegmentation(pxdsResults, pxdsTest, 'Verbose', false);
% 显示数据集级指标
disp('数据集级指标:');
disp(metrics.DataSetMetrics);
% 显示各类别指标
disp('各类别 IoU 和准确率:');
disp(metrics.ClassMetrics);
% 可视化测试集中一张图像的分割结果
I_test = readimage(imdsTest, 35);
C_test = semanticseg(I_test, net, 'Classes', classes);
B_test = labeloverlay(I_test, C_test, 'Colormap', cmap, 'Transparency', 0.4);
figure('Name', '测试集分割示例');
imshow(B_test);
pixelLabelColorbar(cmap, classes);
title('训练后模型在测试图像上的分割效果');
%% ========== 辅助函数定义(全部放在脚本末尾) ==========
% 类别名称
function classes = getClassNames()
classes = [
"Sky"
"Building"
"Pole"
"Road"
"Pavement"
"Tree"
"SignSymbol"
"Fence"
"Car"
"Pedestrian"
"Bicyclist"
];
end
% CamVid 原始 RGB 到 11 类标签 ID 的映射
function labelIDs = camvidPixelLabelIDs()
labelIDs = { ...
% Sky
[128 128 128]; ...
% Building
[000 128 064; 128 000 000; 064 192 000; 064 000 064; 192 000 128]; ...
% Pole
[192 192 128; 000 000 064]; ...
% Road
[128 064 128; 128 000 192; 192 000 064]; ...
% Pavement
[000 000 192; 064 192 128; 128 128 192]; ...
% Tree
[128 128 000; 192 192 000]; ...
% SignSymbol
[192 128 128; 128 128 064; 000 064 064]; ...
% Fence
[064 064 128]; ...
% Car
[064 000 128; 064 128 192; 192 128 192; 192 064 128; 128 064 064]; ...
% Pedestrian
[064 064 000; 192 128 064; 064 000 192; 064 128 064]; ...
% Bicyclist
[000 128 192; 192 000 192]; ...
};
end
% CamVid 颜色映射 (11 类)
function cmap = camvidColorMap()
cmap = [
128 128 128 % Sky
128 0 0 % Building
192 192 192 % Pole
128 64 128 % Road
60 40 222 % Pavement
128 128 0 % Tree
192 128 128 % SignSymbol
64 64 128 % Fence
64 0 128 % Car
64 64 0 % Pedestrian
0 128 192 % Bicyclist
] / 255;
end
% 颜色条显示
function pixelLabelColorbar(cmap, classNames)
colormap(gca, cmap);
c = colorbar(gca);
c.TickLabels = classNames;
numClasses = size(cmap, 1);
c.Ticks = 1/(numClasses*2):1/numClasses:1;
c.TickLength = 0;
end
% 数据集划分
function [imdsTrain, imdsVal, imdsTest, pxdsTrain, pxdsVal, pxdsTest] = partitionCamVidData(imds, pxds)
rng(0);
numFiles = numel(imds.Files);
shuffledIndices = randperm(numFiles);
numTrain = round(0.60 * numFiles);
numVal = round(0.20 * numFiles);
trainingIdx = shuffledIndices(1:numTrain);
valIdx = shuffledIndices(numTrain+1:numTrain+numVal);
testIdx = shuffledIndices(numTrain+numVal+1:end);
imdsTrain = subset(imds, trainingIdx);
imdsVal = subset(imds, valIdx);
imdsTest = subset(imds, testIdx);
pxdsTrain = subset(pxds, trainingIdx);
pxdsVal = subset(pxds, valIdx);
pxdsTest = subset(pxds, testIdx);
end
% 数据增强:随机翻转和平移
function data = augmentImageAndLabel(data, xTrans, yTrans)
for i = 1:size(data,1)
tform = randomAffine2d('XReflection', true, ...
'XTranslation', xTrans, 'YTranslation', yTrans);
rout = affineOutputView(size(data{i,1}), tform, 'BoundsStyle', 'centerOutput');
data{i,1} = imwarp(data{i,1}, tform, 'OutputView', rout);
data{i,2} = imwarp(data{i,2}, tform, 'OutputView', rout);
end
end
% 自定义损失函数(类别加权交叉熵)
function loss = modelLoss(Y, T, classWeights)
weights = dlarray(classWeights, 'C');
mask = ~isnan(T);
T(isnan(T)) = 0;
loss = crossentropy(Y, T, weights, 'Mask', mask, 'NormalizationFactor', 'mask-included');
end
🔧 代码说明
| 段落 | 功能 |
|---|---|
| 1–3 | 参数设置,下载预训练模型并展示推理结果。 |
| 4–8 | 下载 CamVid 数据集,加载图像和标签,分析类别分布,计算类别权重,划分数据集,并应用数据增强。 |
| 9–11 | 创建 DeepLab v3+ 网络,配置训练选项,执行训练(需 doTraining=true)。 |
| 12 | 在测试集上评估模型,输出 IoU 等指标,并可视化分割结果。 |
| 辅助函数 | 包含类别映射、颜色表、数据划分、增强和损失函数等,保持脚本自包含。 |
⚠️ 注意事项
- GPU 内存:训练时
MiniBatchSize=4,若显存不足可改为1或减小输入尺寸(imageSize)。 - 下载失败:若
websave因网络问题失败,可手动下载数据集并修改outputFolder指向本地路径。 - 训练耗时:在 RTX 3090 Ti 上约 50 分钟,CPU 训练极慢,不建议。
- 结果复现:代码设置了随机种子(
rng(0)),确保划分和增强的可重复性。
📚 扩展阅读
如有问题,欢迎在评论区留言讨论! 😊
更多推荐

所有评论(0)