Matlab中实现Transformer数据分类:从数据准备到模型训练
2026/9/7 22:52:55 网站建设 项目流程

老铁们,今天来点硬核的。平时聊到Transformer,默认都是Python加PyTorch,再不济也是TensorFlow,用Matlab做Transformer总觉得差点意思。但实际上,只要你的Matlab版本不太老,并且装了Deep Learning Toolbox,用Transformer做数据分类完全能落地。今天我把这套直接能跑的思路和代码整理出来,从数据收拾、网络搭建、训练评估到踩坑排查,一条龙讲清楚。不管你是搞工业信号分类、传感器多通道时序识别,还是金融特征序列分类,都能套这套流程。文章里给到的代码我尽可能写得能照抄,个别层名字以你本机帮助文档为准,但整体思路绝对通用。

1. 为什么用Transformer做数据分类?Matlab到底行不行

1.1 先搞清楚Transformer在分类场景里的优势

很多人一听到Transformer就想到大语言模型,其实把它用在数据分类上,逻辑很简单:它靠自注意力机制(Self-Attention)直接建模输入序列中任意两个位置之间的关系。这意味着,如果某个关键特征出现在序列的前面,而决定类别的信号在序列末尾,Transformer依然能轻松把这两端关联起来。

这一点和LSTM不同。LSTM是逐步传递隐藏状态,长距离信息会一路衰减;CNN则受限卷积核大小,需要堆很多层才能扩大感受野。Transformer一上来就把整个序列摊开,所有位置两两互动,在特征跨度大的数据上优势非常明显。实际做分类时,我会把输入看成“多通道的时间序列”或者“多特征的有序样本”,让网络自己决定该重点看哪个区域、哪几个特征组合在一起有判别力。

但话说回来,如果你的数据很短、特征维度也不高,用个随机森林或者简单的全连接网络可能更快更稳。Transformer不是银弹,它适合的是序列长度中等以上、特征之间确实存在长程依赖的场景。

1.2 Matlab做Transformer的三种路线,别选错

Matlab没有一个叫“Transformer”的终极一键层,但你完全可以用官方提供的基础层把它拼出来。我实际用下来,主要有三条路线:

  • 路线一:用 Deep Learning Toolbox 里的selfAttentionLayer自己搭 layerGraph。这是 R2021a 之后引入的层,本质上就是把多头自注意力打包好,是今天文章的主角。
  • 路线二:如果你的版本比较新(R2023a 之后),可能已经有transformerLayer,这更接近标准Transformer block,直接当普通层用就行。
  • 路线三:在 Matlab 里通过 Python 交互调用训练好的 PyTorch 模型。这条路适合你不想用Matlab重写模型的情况,但今天不展开,因为要搞定Python环境,就失去了用Matlab图省事的意义。

我的建议是:先查一下自己的版本有没有selfAttentionLayer,在命令行敲help selfAttentionLayer,能看到帮助说明就用这条路线,兼容性最好,可控性也最强。我用ver('deep')确认过,很多同学明明装了工具箱,却一直没发现这些好东西。

2. 第一步不是写网络,而是把数据收拾成能训练的样子

2.1 网络到底希望吃到什么形状的数据

很多新手一上来就纠结网络结构,结果在数据格式上卡了两小时。这里先统一标准:Matlab里训练序列分类网络,输入一般是N×1的 cell 数组,每个 cell 存的是一个numFeatures × sequenceLength的 double 矩阵。换句话说,特征放行,时间步放列。

举个例子,如果你的样本是8个传感器通道、每个样本采集50个时间点的数据,那每个样本就是一个8×50的矩阵,最终XTrain样本数×1的 cell。标签用categorical,比如{"正常";"故障";"异常"}转成分类向量。

这个格式和sequenceInputLayer的要求是严格对应的,后面的自注意力层也延续了这个排布习惯。搞清楚这一点,后面维度报错会少一大半。

2.2 CSV导入、归一化、划分训练验证集

假设数据放在CSV里,每行是一个样本,前若干列是特征,最后一列是标签。第一步建议用readmatrix快速读进来,而不是readtable,因为纯数值矩阵后面处理更方便。

data = readmatrix("your_data.csv"); X = data(:, 1:end-1); Y = data(:, end);

这里有个坑:如果你的数据是每条样本一个固定长度的时间序列,CSV可能把时间步也展开成了列,那每个样本的维度就是特征数 × 时间步数。读进来之后,需要先reshape成 cell 数组:

numFeatures = 8; seqLen = 50; XTrain = cell(size(X, 1), 1); for i = 1:size(X, 1) XTrain{i} = reshape(X(i, :), numFeatures, seqLen); end YTrain = categorical(Y);

归一化别偷懒。Transformer对输入尺度比较敏感,直接喂原始数据容易让注意力权重被个别大值特征带偏。我用的是mapminmax或者zscore

X = zscore(X, 0, 'all');

注意要在训练集上计算均值方差,再应用到验证集和测试集,避免数据泄漏。划分数据集用cvpartition做分层抽样更稳,类别不平衡时尤其重要:

cv = cvpartition(YTrain, 'HoldOut', 0.2); idxTrain = training(cv); idxTest = test(cv); XTrain = XTrain(idxTrain); YTrain = YTrain(idxTrain); XTest = XTrain(idxTest); YTest = YTrain(idxTest);

数据量小的时候,一定要把验证集也保留好,别全扔进训练。后面调参全靠它判断过拟合。

3. 核心代码:用Layer Graph搭一个Transformer分类网络

3.1 按标准Transformer思路拆解网络

标准的Transformer block包含:位置编码、多头自注意力、残差连接、层归一化、前馈网络。我们做分类任务时,不用把Decoder那部分搬过来,只需要Encoder的特征提取能力,最后接一个分类头。

在Matlab里,我搭网络时喜欢拆成四段:

  • 第一段:输入层 + 位置编码。输入层用sequenceInputLayer,位置编码如果嫌麻烦可以先不加,但如果序列顺序本身对分类有影响,建议加上。
  • 第二段:多头自注意力层。用selfAttentionLayer(numHeads, numHeadDimensions),它会把输入序列转成Query、Key、Value,然后算注意力权重。
  • 第三段:归一化 + 池化。加一个layerNormalizationLayer稳定训练,然后用全局平均池化把序列方向压缩成一个向量,给后面的全连接层用。
  • 第四段:分类头。全连接层 + ReLU + 全连接层 + Softmax + 分类层。

这个结构不算严格意义上带残差的完整Transformer,但抓住了最核心的自注意力机制,代码能跑,效果也可控,特别适合入门。

3.2 完整可跑的搭建代码

先解决位置编码的问题。Matlab没有现成的位置编码层,我们可以自己写一个简单的。新建一个PositionalEncodingLayer.m,把下面这段保存下来:

classdef PositionalEncodingLayer < nnet.layer.Layer properties Pe end methods function layer = PositionalEncodingLayer(maxSeqLen, dModel) layer.Name = "pos"; layer.Description = "sinusoidal positional encoding"; pos = (0:maxSeqLen-1)'; divTerm = 1 ./ (10000 .^ ((0:2:dModel-2) / dModel)); pe = zeros(maxSeqLen, dModel); pe(:,1:2:end) = sin(pos * divTerm); pe(:,2:2:end) = cos(pos * divTerm); layer.Pe = pe'; end function Z = predict(layer, X) % X: [dModel, S, B] [~, S, ~] = size(X); Z = X + layer.Pe(:, 1:S); end end end

这个自定义层做的事就是在输入特征上叠加一个和时间步位置有关的固定向量。dModel 建议是偶数,这样正余弦能均匀分配。

接下来是主网络搭建:

numFeatures = 8; numClasses = 3; numHeads = 4; numHeadDimensions = 16; hiddenSize = 64; seqLen = 50; posLayer = PositionalEncodingLayer(seqLen, numFeatures); layers = [ sequenceInputLayer(numFeatures, "Name", "in") posLayer selfAttentionLayer(numHeads, numHeadDimensions, "Name", "attn") layerNormalizationLayer("Name", "ln1") globalAveragePooling1dLayer("Name", "gap") fullyConnectedLayer(hiddenSize, "Name", "fc1") reluLayer("Name", "relu1") fullyConnectedLayer(numClasses, "Name", "fc2") softmaxLayer("Name", "softmax") classificationLayer("Name", "out") ]; lgraph = layerGraph(layers);

如果你的Matlab版本没有globalAveragePooling1dLayer,可以用一个自定义层替代,核心代码就一行:

Z = mean(X, 2);

方法是在自定义层的predict函数里取时序维度的均值,输出就变成特征数×1×batch,后面接全连接层完全没问题。

3.3 训练选项配置:先把训练跑通再说调参

网络搭好了,先不要追求最优效果,最重要的是让它能在几分钟内跑通。训练选项我用adam优化器,初始学习率给0.001,这个值在多数小数据集上不会太激进,也不会慢到让人失去耐心。

options = trainingOptions("adam", ... "InitialLearnRate", 0.001, ... "MaxEpochs", 50, ... "MiniBatchSize", 16, ... "ValidationData", {XTest, YTest}, ... "ValidationFrequency", 20, ... "Plots", "training-progress", ... "Verbose", false); net = trainNetwork(XTrain, YTrain, lgraph, options);

这里提醒一下,MiniBatchSize不要太贪。Transformer的自注意力计算量和序列长度的平方成正比,序列长度50还好,如果到了几百,显存或者内存会迅速吃紧。我一般从16开始,跑通了再往上加。

如果你的数据量很小,把MaxEpochs降到20,观察验证集准确率的变化,避免过拟合。训练过程会弹出一个实时曲线图,看到损失下降就说明网络在正常学习。

4. 结果分析:训练完怎么看指标、怎么调参数

4.1 看训练曲线:损失、准确率、验证集表现

训练跑完以后,先别急着看准确率,先看损失曲线。我见过很多同学一看到训练准确率99%就开心得不行,结果测试集上一塌糊涂,这就是典型的过拟合。

你要关注两个信号:

  • 训练损失下降但验证损失开始回升,说明模型开始死记硬背训练集了。这时候减小模型规模或加正则化。
  • 训练损失和验证损失都在高位震荡,说明学习率可能太大,或者模型结构有问题,需要回退到更简单的配置。

在Matlab的训练进度图里,上面的子图是准确率,下面的是损失。两条曲线都平滑下降,才是健康的状态。

4.2 混淆矩阵与分类指标

训练完成后,用classify对测试集做预测,然后直接画混淆矩阵:

YPred = classify(net, XTest); figure; confusionchart(YTest, YPred);

这张图能让你一眼看出模型在哪两个类别之间容易混淆。如果某个类别的召回率特别低,通常是这个类的样本量太少或者特征重叠度高。这时候可以回看归一化过程是不是在全体数据上做了,导致验证集信息泄漏。

如果想计算更细的指标,可以手写几行:

acc = mean(YPred == YTest); % 每类的precision、recall、F1可以用confusionmat自己算 C = confusionmat(YTest, YPred); precision = diag(C) ./ sum(C, 2); recall = diag(C) ./ sum(C, 1)'; f1 = 2 * precision .* recall ./ (precision + recall);

注意sum(C,1)的维度,用confusionmat之前最好先summary(YTest)确认类别顺序。

4.3 调参优先级和我的经验值

很多同学第一次上手就被一堆超参数搞晕:注意力头数、头维度、学习率、batch size、层数到底先调哪个?我的经验是,按这个优先级来:

学习率大于一切。先把学习率调到能让损失稳定下降,再谈结构。如果损失震荡太厉害,降到0.0003;如果收敛太慢,试试0.003,但注意配合更大的batch size。

然后是batch size。它对训练的稳定性和显存占用影响很大。小数据集上,16到32通常够用。接着才是注意力头数和维度。头数我一般在4到8之间选,头维度16到32之间选。这两个参数影响的是模型表达能力的上限,但对最终结果的影响往往没有学习率那么立竿见影。

超参数我的常用范围调参方向
InitialLearnRate0.0003~0.003损失震荡就调小,收敛慢就调大
MiniBatchSize16~64显存不足时调小,模型收敛不稳时调大
NumHeads4~8序列长、特征维度高时可适当增大
NumHeadDimensions16~64过大容易过拟合
MaxEpochs20~100看验证集停止提升就提前停

如果验证集提升不明显,我还会检查位置编码是否加对了。之前有个项目,序列的前后顺序其实是关键信息,我漏了位置编码,结果模型准确率卡在60%上不去,加上之后直接到85%。这个坑我印象太深了。

5. 容易翻车的几个地方:问题排查速查表

5.1 维度错误

这是新手最常遇到的报错,比如“Layer 'attn': Invalid input size”或者“Expected input to have 3 dimensions”。绝大多数情况下,问题出在输入数据的cell格式上。记住前面说的:每个cell必须是numFeatures × seqLen的矩阵,而且numFeatures一定要和sequenceInputLayer第一个参数完全一致。

另一个容易踩的是位置编码层的维度。PositionalEncodingLayer里的dModel必须等于numFeatures,否则在自定义层的predict里做加法时矩阵尺寸对不上。

5.2 训练不收敛或者直接NaN

NaN 问题八成出在梯度爆炸上。Transformer的自注意力层对学习率比较敏感,这时候先把学习率降到0.0001试试。如果还是NaN,检查数据里是不是有缺失值或者无穷大值。我用any(isnan(X), 'all')排查过,数据源里一个NaN就能让整个训练崩掉。

没有NaN但损失不降,第一反应看标签是不是从1开始连续编号。classificationLayer要求标签必须是categorical类型,数字标签直接喂进去有时会出问题。

5.3 CPU训练太慢

不是所有人都有NVIDIA GPU,Matlab在纯CPU上跑Transformer确实煎熬。我的建议是:

  • 先把seqLen截短到模型能忍受的范围,比如50或者100,别一上来就处理500个时间步。
  • MiniBatchSize调小到8,减少内存压力。
  • "ExecutionEnvironment", "auto",让Matlab自动选CPU还是GPU。

如果你的数据确实很长,先考虑降采样或者用滑动窗口切成短片段。窗口长度在保证覆盖特征跨度的前提下越短越好。

5.4 数据量太小

Transformer是数据饥渴型模型,没有足够样本很容易过拟合。如果训练集只有几百条样本,我一般会做两件事:

  • 减小模型:降低numHeadDimensionshiddenSize,让模型没那么多参数去死记硬背。
  • 加dropout层:在fc1后面加dropoutLayer(0.5),正则化效果立竿见影。

如果你的数据允许,也可以做简单的数据增强。比如时间序列上做小幅平移、加噪声、随机缩放,这些操作在Matlab里用几行代码就能实现,能有效缓解过拟合。

最后再分享一个我自己的小习惯:不管任务看起来多简单,我都会先写一个极简版本跑通全流程——用50个样本、10个epoch、最小模型,确认数据流没问题后,再把规模慢慢加上去。这样排查问题的成本最低,也不容易一开始就在某个隐蔽的维度错误上卡死。希望这套Matlab Transformer分类流程能帮你少走点弯路。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询