基于深度学习的语义分割:从理论到实践(MATLAB 版详解)

本文基于 MathWorks 官方示例,系统讲解语义分割的核心概念、DeepLab v3+ 网络结构、CamVid 数据集处理、训练与评估全流程。适合希望快速上手语义分割的开发者与研究者。


📌 目录


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 分割,返回对应的 imageDatastorepixelLabelDatastore

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);
  • 修改 classeslabelIDs 映射;
  • 根据数据规模调整训练轮次和 batch size;
  • 若数据较少,可冻结骨干网络部分层,仅训练分类头。

9. 总结与参考资料

9.1 关键 takeaways

  • 语义分割是像素级分类任务,DeepLab v3+ 是目前主流框架之一。
  • 类别不平衡是真实场景的常态,需通过加权损失或重采样处理。
  • 数据增强、预训练权重、合理超参数是提升性能的三大法宝。
  • 评估指标应结合全局 IoU 和类别 IoU 综合判断。

9.2 相关工具箱与函数

  • Computer Vision Toolboxsemanticseg, labeloverlay, pixelLabelDatastore
  • Deep Learning Toolboxtrainnet, trainingOptions, dlnetwork
  • Parallel Computing Toolbox:GPU 加速支持

9.3 参考文献

  1. Chen, L.-C., et al. “Encoder-Decoder with Atrous Separable Convolution for Semantic Image Segmentation.” ECCV 2018.
  2. Brostow, G. J., et al. “Semantic object classes in video: A high-definition ground truth database.” Pattern Recognition Letters, 2009.
  3. 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

🚀 运行方式

  1. 将以下代码完整复制到一个新建的 MATLAB 脚本文件(例如 SemanticSegmentationDemo.m)中。
  2. 直接运行脚本。
  3. 首次运行会下载预训练模型(约 58 MB)和 CamVid 数据集(图像 557 MB + 标签 16 MB),请确保网络通畅且有足够磁盘空间。
  4. 若希望从头训练,将脚本开头的 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 等指标,并可视化分割结果。
辅助函数 包含类别映射、颜色表、数据划分、增强和损失函数等,保持脚本自包含。

⚠️ 注意事项

  1. GPU 内存:训练时 MiniBatchSize=4,若显存不足可改为 1 或减小输入尺寸(imageSize)。
  2. 下载失败:若 websave 因网络问题失败,可手动下载数据集并修改 outputFolder 指向本地路径。
  3. 训练耗时:在 RTX 3090 Ti 上约 50 分钟,CPU 训练极慢,不建议。
  4. 结果复现:代码设置了随机种子(rng(0)),确保划分和增强的可重复性。

📚 扩展阅读


如有问题,欢迎在评论区留言讨论! 😊

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐