ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

Matlab中实现Transformer数据分类:从数据准备到模型训练

Matlab中实现Transformer数据分类:从数据准备到模型训练 老铁们今天来点硬核的。平时聊到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或者zscoreX 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 训练不收敛或者直接NaNNaN 问题八成出在梯度爆炸上。Transformer的自注意力层对学习率比较敏感这时候先把学习率降到0.0001试试。如果还是NaN检查数据里是不是有缺失值或者无穷大值。我用any(isnan(X), all)排查过数据源里一个NaN就能让整个训练崩掉。没有NaN但损失不降第一反应看标签是不是从1开始连续编号。classificationLayer要求标签必须是categorical类型数字标签直接喂进去有时会出问题。5.3 CPU训练太慢不是所有人都有NVIDIA GPUMatlab在纯CPU上跑Transformer确实煎熬。我的建议是先把seqLen截短到模型能忍受的范围比如50或者100别一上来就处理500个时间步。MiniBatchSize调小到8减少内存压力。用ExecutionEnvironment, auto让Matlab自动选CPU还是GPU。如果你的数据确实很长先考虑降采样或者用滑动窗口切成短片段。窗口长度在保证覆盖特征跨度的前提下越短越好。5.4 数据量太小Transformer是数据饥渴型模型没有足够样本很容易过拟合。如果训练集只有几百条样本我一般会做两件事减小模型降低numHeadDimensions和hiddenSize让模型没那么多参数去死记硬背。加dropout层在fc1后面加dropoutLayer(0.5)正则化效果立竿见影。如果你的数据允许也可以做简单的数据增强。比如时间序列上做小幅平移、加噪声、随机缩放这些操作在Matlab里用几行代码就能实现能有效缓解过拟合。最后再分享一个我自己的小习惯不管任务看起来多简单我都会先写一个极简版本跑通全流程——用50个样本、10个epoch、最小模型确认数据流没问题后再把规模慢慢加上去。这样排查问题的成本最低也不容易一开始就在某个隐蔽的维度错误上卡死。希望这套Matlab Transformer分类流程能帮你少走点弯路。
返回列表