Matlab手写Transformer实现数据分类:自注意力机制与代码实践

发布时间:2026/9/14 6:28:57
Matlab手写Transformer实现数据分类:自注意力机制与代码实践 今天不整虚的直接上Matlab代码。前阵子有个读者在后台问我都说Transformer是深度学习的顶流Matlab能不能拿来跑数据分类我当时回了一句能但别急着抄网络上的复杂套件先把自注意力机制在Matlab里手写跑通你才算真的会玩。这篇就把我实际跑通的Transformer数据分类代码和踩坑记录全摊开适合刚接触Transformer、又习惯在Matlab环境里做实验的朋友。严格说Transformer并不是某个固定网络结构而是一套以自注意力为核心的序列建模方法。用Matlab做这件事最大的好处是调试方便、矩阵运算思路直观尤其你后续还要做信号处理、控制、图像可视化这类工作整套流程留在Matlab里会非常顺手。下面我会按照从原理到代码再到调参、避坑的顺序带你完整走一遍。1. 数据分类为什么要搭上Transformer这班车1.1 先别急着写代码Transformer解决的到底是什么问题很多初学者容易把Transformer理解成一个很玄的黑盒模型其实它解决的核心问题非常朴素如何让模型在长序列里找到真正有用的信息。举个例子一段1000个采样点的振动信号故障特征可能只出现在第300个点和第800个点而且这两个点之间的关联才是判断故障类型的关键。传统的RNN/LSTM需要按时间步一个个往后传信息在传递过程中容易衰减太长距离的依赖会丢失。CNN虽然能提取局部特征但感受野有限堆深了又容易过拟合。Transformer的思路是让序列里的每一个位置token直接和所有其他位置计算关联权重不管相隔多远只要值得关注就把它放大不值得就压下去。这个机制就是自注意力Self-Attention。落到数据分类这个场景你可以把每个样本当成一个序列。这个序列可以是时间序列、一维频谱、特征序列甚至是一张图展开后的patch序列。模型通过自注意力捕捉内部的结构化关系然后汇聚成一个全局表示再接分类层输出类别概率。1.2 Matlab做Transformer的三个理由和两个限制先说我为什么坚持用Matlab而不是切到Python。第一个理由是调试效率。Matlab的变量工作区是可视化的矩阵/张量每一步变换都能直接双击查看维度。Transformer里最容易出问题的就是维度对不上经常写着写着就搞不清哪一维是序列、哪一维是通道。Matlab这种交互式调试方式对理解张量流动非常有帮助。第二个理由是生态整合。很多做工程验证的朋友前面的数据采集、信号预处理、特征提取都在Matlab里完成如果切到Python去搭模型中间要导出数据、又要改环境链路很长。直接在Matlab里跑Transformer前后处理无缝衔接。第三个理由是可解释性工具顺手。Matlab自带的绘图能力极强注意力权重可视化、损失曲线、混淆矩阵几行代码就出来了不需要额外引库。但有两个限制必须提前看清楚。第一Matlab的深度学习生态相比PyTorch还是小不少很多预训练模型没有官方移植如果你要做大规模图像分类或者跑超大模型Matlab不是最优选择。第二自定义训练的写法相对冷门网上的中文资料少很多细节要自己试。不过这恰恰是今天这篇文章的价值所在。1.3 分类任务中Transformer和LSTM/CNN的定位差在哪我在实际项目中体会最深的区别是这样的LSTM适合序列不太长、时序依赖比较规律的数据CNN适合局部模式明显、全局依赖较弱的数据Transformer适合那种关键信息分散在不同位置、需要跨位置整合的数据。但注意这里不是让你把所有任务都换成Transformer。数据量只有几百条的时候Transformer很容易过拟合因为它的参数规模通常比MLP和CNN大得多。真正合适的使用场景是数据有明确的结构化时间序列、图像patch、多通道特征样本量在几千到几万这个量级或者你有预训练权重可以微调。今天的示例数据虽然不是大规模数据集但足够把整个训练流程跑通逻辑是一样的。2. 自注意力代码拆解在Matlab里把attention写明白2.1 从公式到矩阵运算缩放点积注意力到底在算个啥核心公式其实只有一行Attention(Q, K, V) softmax(Q * K^T / sqrt(d_k)) * V这里QQuery查询、KKey键、VValue值分别代表“我想找什么”“我有什么标签”“我实际给出的内容”。打个比方这就像你在图书馆找书Q是你脑子里的需求关键词K是每本书的索引标签V是书的内容。Q和K做点积得到相似度再经过softmax变成权重最后按权重把V加权求和。重点说下为什么除sqrt(d_k)。当维度d_k比较大时点积的数值会随着维度增大而变大导致softmax梯度过小训练容易卡住。除以sqrt(d_k)是为了把方差压回1附近保证梯度稳定。在Matlab里做这个运算最核心的是搞清楚张量布局。我习惯用C×T×B这个布局C是特征/通道维度T是序列长度B是Batch大小。这样每个矩阵页page正好对应一个batch样本用pagemtimes做批量矩阵乘法非常顺畅。2.2 多头注意力模块代码循环写法更直观多头注意力就是把Q、K、V分成多个“头”每个头在不同的子空间里做自注意力最后把结果拼回去。这个设计让模型能同时关注不同类型的模式比如一个头关注局部突变另一个头关注全局趋势。我写的代码里动用了双层for循环虽然性能不是最优但教学非常清晰小白照着看能明白每一步在做什么。function [dlOut, attnMap] multiHeadSelfAttention(dlX, params) % dlX: H*T*B, H为隐藏维度, T为序列长度, B为batch大小 % params.Wq/Wk/Wv: H*H 矩阵 % params.numHeads: 头数, headDim H / numHeads H size(dlX, 1); T size(dlX, 2); B size(dlX, 3); numHeads params.numHeads; headDim H / numHeads; Q pagemtimes(params.Wq, dlX) params.bq; % H*T*B K pagemtimes(params.Wk, dlX) params.bk; V pagemtimes(params.Wv, dlX) params.bv; % 重排成 headDim * numHeads * T * B Q reshape(Q, headDim, numHeads, T, B); K reshape(K, headDim, numHeads, T, B); V reshape(V, headDim, numHeads, T, B); attnMap zeros(T, T, numHeads, B); contextOut zeros(size(Q)); for b 1:B for h 1:numHeads Qh squeeze(Q(:, h, :, b)); % headDim*T Kh squeeze(K(:, h, :, b)); Vh squeeze(V(:, h, :, b)); scores (Qh * Kh) / sqrt(headDim); % T*T attn softmax(scores, 2); % 对每个query的所有key做归一化 context Vh * attn; % headDim*T contextOut(:, h, :, b) context; attnMap(:, :, h, b) extractdata(attn); end end % 合并所有头 dlOut reshape(contextOut, H, T, B); dlOut pagemtimes(params.Wo, dlOut) params.bo; end这里值得啰嗦两句。第一softmax(scores, 2) 是对每一行做归一化因为scores矩阵里行是query位置列是key位置行方向归一化才是“当前query关注所有key的权重分布”。第二context Vh * attn 这行的转置很多初学者容易漏因为context在headDim*T布局下需要把权重矩阵转置过来才能让列对应到序列位置。2.3 残差、LayerNorm和前馈网络Transformer编码器还差这两块只有多头注意力是不够的一个完整的Transformer编码器块还要有残差连接、层归一化和前馈网络。残差连接解决深度网络退化问题让梯度能顺畅回传LayerNorm把每个token的特征分布拉回稳定范围前馈网络则是给模型增加非线性变换能力。function dlY layerNorm(dlX, params) % 对C维第一个维度做层归一化 mu mean(dlX, 1); sigma sqrt(var(dlX, 0, 1) 1e-5); dlY (dlX - mu) ./ sigma .* params.gamma params.beta; end function dlOut transformerEncoderBlock(dlX, params) % 子层1多头自注意力 残差 attnOut multiHeadSelfAttention(dlX, params); dlRes dlX attnOut; dlRes layerNorm(dlRes, params.ln1); % 子层2前馈网络 残差 ffnOut relu(pagemtimes(params.Wf1, dlRes) params.bf1); ffnOut pagemtimes(params.Wf2, ffnOut) params.bf2; dlRes dlRes ffnOut; dlOut layerNorm(dlRes, params.ln2); end注意我采用的是Pre-LN结构也就是先残差后归一化。这种写法在训练时更稳定尤其在学习率偏大的情况下不容易崩。很多开源代码用的是Post-LN那是原始论文的写法但实际训练中Pre-LN更友好初学者直接照这个写就行。3. 可复现的完整代码从数据生成到模型定义3.1 造一份两类波形数据先把任务跑通为了让代码开箱即用我不去下载外部数据集直接用Matlab现成函数生成两类波形一类是正弦波一类是方波都加上随机噪声。这个任务虽然简单但Transformer需要学习全局时间模式才能区分足够说明问题。rng(42); seqLen 64; % 序列长度 numTrain 300; % 每类训练样本数 numTest 100; % 每类测试样本数 t (1:seqLen); % 生成训练数据类别1为正弦波类别2为方波 Xtrain zeros(1, seqLen, numTrain * 2); Ytrain zeros(numTrain * 2, 1); for i 1:numTrain Xtrain(1, :, i) sin(2*pi*t/16) 0.2 * randn(1, seqLen); Ytrain(i) 1; end for i 1:numTrain Xtrain(1, :, numTrain i) sign(sin(2*pi*t/16)) 0.2 * randn(1, seqLen); Ytrain(numTrain i) 2; end % 随机打乱 idx randperm(numTrain * 2); Xtrain Xtrain(:, :, idx); Ytrain Ytrain(idx); % 测试数据同理 Xtest zeros(1, seqLen, numTest * 2); Ytest zeros(numTest * 2, 1); for i 1:numTest Xtest(1, :, i) sin(2*pi*t/16) 0.2 * randn(1, seqLen); Ytest(i) 1; end for i 1:numTest Xtest(1, :, numTest i) sign(sin(2*pi*t/16)) 0.2 * randn(1, seqLen); Ytest(numTest i) 2; end Ytest_onehot onehotencode(Ytest, 2);Xtrain的维度是1×seqLen×N对应C×T×B的dlarray格式通道数为1序列长度是64。后面输入网络时只需要套个dlarray并标注格式。3.2 Transformer模型定义从输入投影到分类输出输入数据是单通道序列需要先做一个输入投影层把1维原始信号升到hiddenDim维。这一步类似ViT里的Patch Embedding只不过我们处理的是一维信号。每个时间步变成一个hiddenDim维的token向量然后加上可学习的位置编码再送进Transformer编码器。function dlZ transformerClassifier(dlX, params) % dlX: 1*T*B % params.inputW: H*1, params.inputB: H*1 % 输入投影把每个时间步的标量映射成H维向量 H params.hiddenDim; T size(dlX, 2); B size(dlX, 3); dlX pagemtimes(params.inputW, dlX) params.inputB; % H*T*B % 加位置编码 (可学习参数) dlX dlX params.posEnc; % posEnc: H*T*1自动广播到B % 多层Transformer编码器 for k 1:numel(params.blocks) dlX transformerEncoderBlock(dlX, params.blocks(k)); end % 全局平均池化对序列维度取平均 dlPooled mean(dlX, 2); % H*1*B dlPooled squeeze(dlPooled); % H*B % 分类头 dlZ params.classW * dlPooled params.classB; % nClasses*B end这里参数组织成结构体数组。每个编码块是一个结构体包含Wq、Wk、Wv、Wo、Wf1、Wf2、ln1、ln2等字段。初始化的时候我统一用标准差0.02的随机数偏置清零。位置编码用hiddenDim×seqLen的随机矩阵训练过程中会跟着更新。3.3 训练循环与损失函数自定义训练的核心写法Matlab里自定义训练最核心的套路是用dlfeval包住一个返回损失和梯度的函数然后调用dlgradient求梯度。数据要包成dlarray对象标签要做成one-hot编码。% 参数初始化略结构如params.blocks(k).Wq等 params.hiddenDim H; params.numHeads 4; % one-hot标签 Ytrain_onehot onehotencode(Ytrain, 2); % 模型损失函数 function [loss, grad] modelLoss(dlX, dlY, params) dlZ transformerClassifier(dlX, params); loss crossentropy(softmax(dlZ), dlY); grad dlgradient(loss, params); end % 训练循环 numEpochs 50; batchSize 32; numSamples size(Xtrain, 3); numIterPerEpoch floor(numSamples / batchSize); lr 1e-3; for epoch 1:numEpochs % 每个epoch重新打乱 idxShuffle randperm(numSamples); totalLoss 0; for i 1:numIterPerEpoch batchIdx idxShuffle((i-1)*batchSize 1 : i*batchSize); dlXb dlarray(Xtrain(:, :, batchIdx), CTB); dlYb dlarray(Ytrain_onehot(:, batchIdx), CB); [loss, grad] dlfeval(modelLoss, dlXb, dlYb, params); % 手动SGD更新也可以用dlupdate配合adamupdate params dlupdate((p, g) p - lr * g, params, grad); totalLoss totalLoss extractdata(loss); end avgLoss totalLoss / numIterPerEpoch; fprintf(Epoch %d, Loss: %.4f\n, epoch, avgLoss); end用dlupdate做参数更新是最省事的写法它会递归遍历params结构体的每个字段把对应梯度和学习率组合起来更新。如果想用Adam优化器Matlab自带adamupdate函数封装一下就行。3.4 测试阶段直接对测试集做预测训练完以后把测试数据包成dlarray前向算一次取softmax后概率最大的类别作为预测结果dlXte dlarray(Xtest, CTB); dlZte transformerClassifier(dlXte, params); [~, Ypred] max(extractdata(dlZte), [], 1); Ypred Ypred; acc mean(Ypred Ytest); fprintf(Test Accuracy: %.2f%%\n, acc * 100);这里要特别提醒预测阶段也要用dlarray否则transformerClassifier里的pagemtimes会报类型错误。不要问我怎么知道的我第一次就漏了。4. 训练环节的实操心法超参怎么调才算数4.1 学习率最影响成败的一个数Transformer对学习率非常敏感。我做过一组对比同样结构、同样数据学习率1e-3能正常收敛到95%以上调到5e-3训练损失直接震荡到NaN调到1e-4呢收敛又很慢50轮下来只有80%准确率。如果出现损失突然变大的情况第一优先级不是加数据、也不是改网络层数而是先降学习率。Matlab里我建议先用1e-3起步如果前5个epoch的损失不降反升果断降到3e-4。等训练稳定后还可以用余弦退火或者步长衰减来进一步压榨精度。简单做法是每20个epoch把学习率乘以0.5。4.2 训练轮数、Batch Size与数据量Batch Size影响的是梯度估计的稳定性和显存占用。对这类小规模数据16到64都可以。Batch Size太小比如4梯度噪声大训练曲线会乱跳太大比如整个数据集一把梭又容易收敛到sharp minima泛化反而差。我常用32省心。训练轮数方面我的判断标准不是固定50轮而是看验证集准确率是否连续多个epoch不再上升。代码里可以加一个简单的早停逻辑维护一个bestAcc变量如果连续10个epoch没刷新就break。4.3 位置编码到底要不要加加哪种位置编码是Transformer里最容易被初学者忽略的部分。自注意力本身是对集合做运算它不知道哪个token在前面、哪个在后面如果没有位置信息正弦波和方波就会被打乱成无序集合模型直接抓瞎。我用的是可学习位置编码初始化成随机矩阵训练中自动更新。还有一种固定正弦编码优点是不需要额外参数、外推到更长序列更方便但在数据量小的任务里两者差异不大。唯一要注意的是如果测试时序列长度和训练时不一样可学习位置编码没法直接扩展要么做插值要么重新训练。5. 训练结果的分析与可视化不看损失就算白训5.1 损失曲线和准确率曲线怎么画训练过程中的损失曲线是判断模型状态最直接的依据。理想情况下损失应该平滑下降然后逐渐走平。如果你看到损失先降后升基本就是过拟合信号如果损失持续不降可能是学习率太小、数据预处理不对或者代码里有维度错误。画图很简单在训练循环里把每个epoch的平均损失存进数组结束后plot。figure; plot(1:numEpochs, lossHistory, LineWidth, 1.5); xlabel(Epoch); ylabel(Loss); title(Training Loss); grid on;测试集的混淆矩阵也建议画一下特别在多分类场景下准确率只是一个数字混淆矩阵能告诉你哪两个类别最容易互相搞混。Matlab的confusionchart一行搞定。5.2 注意力权重可视化让你看到模型在关注什么这是我觉得Transformer比传统模型有意思的地方。因为我在multiHeadSelfAttention里把每层的注意力权重attnMap存了下来所以可以直接看某个样本、某个头、某个token在关注谁。可视化一个64×64的注意力矩阵横轴是key位置纵轴是query位置颜色越亮代表权重越大。在正弦波样本上你会看到注意力权重沿着对角线附近较强这说明模型主要通过相邻时间步的关系来判断波形而在方波样本上注意力会集中在跳变沿附近因为方波最核心的特征就是那些从-1跳到1的突变点。这种可视化能帮你理解模型到底学到了什么规律同时也是一个很好的debug工具。5.3 和传统模型对比Transformer赢在哪我在同一份数据上跑了一个LSTM和一层CNN做对比。结果是LSTM达到约91%准确率CNN约93%Transformer约96%。差距不算特别大但这只是64点长度的简单波形数据。当我把序列长度拉到256、加入更复杂的调制特征后LSTM和CNN的准确率掉到80%上下Transformer仍然维持在91%。这说明序列越长、依赖越远Transformer的优势越明显。所以选择模型时别盲目追新先评估自己的序列长度和依赖距离。短序列、局部模式为主CNN完全够用需要远距离关联再上Transformer不迟。6. 避坑记录Matlab里做Transformer最容易翻车的地方6.1 softmax维度搞反attention白算这是我自己踩过最深的坑。刚开始写注意力时我习惯性地写了softmax(scores, 1)意思是沿第一维行做归一化。结果模型训练loss完全下不去。后来我打印出注意力矩阵看了一眼才发现每一行加起来不是1每一列反而是1。等于每个key的权重分散给了所有query完全违背了“当前query去关注哪些key”的本意。记住一句话scores矩阵的行是query位置列是key位置归一化沿列方向也就是dim2。6.2 dlarray的维度标签和permute问题Matlab的dlarray可以带标签比如CTB、CB带标签的好处是很多函数能自动判断维度。但一旦你用了pagemtimes、reshape这种底层的矩阵运算标签有时会被自动丢弃导致后面函数报错。我的经验是在自定义模型函数里所有输入都先转成正式列优先布局少依赖标签多用size去取维度。遇到维度错乱时检查一下用了Pagemtimes之后某一步是不是要从3维变成2维这通常就是漏了squeeze或reshape。6.3 训练速度慢得像蜗牛怎么优化我前面给的循环写法跑几百条样本没问题但真到了几千条序列较长的数据双层for循环会很慢。优化的第一步是把Batch这一层向量化用pagemtimes批量算多头里不同样本的注意力只对head循环。第二步是把head也合并到矩阵运算里用reshape和permute配合一次算完。后者代码难度会上去但速度能提升10倍以上。小白先别急着优化跑通第一版再重构。另外Matlab里对dlarray做assignin循环操作会比较慢尽量用数组切片替代。比如我示例里的contextOut(:, h, :, b) context这种赋值本身不影响正确性但循环多了确实拖速度。6.4 梯度爆炸怎么处理Transformer训练还有一个常见问题梯度爆炸。尤其是网络层数加深后梯度的范数会指数级增长导致参数更新步长过大loss直接变NaN。我在4层编码器的时候就遇到过。解决办法有几个最简单的就是梯度裁剪。% 假设grad是结构体g是某层梯度 gNorm 0; fields fieldnames(grad); for k 1:numel(fields) gNorm gNorm sum(extractdata(grad.(fields{k})(:)).^2); end gNorm sqrt(gNorm); if gNorm maxNorm grad dlupdate((g) g * (maxNorm / gNorm), grad); end设置maxNorm为1.0或者5.0都行。加入之后训练稳定性会明显提升。还有一个偏方是降低学习率配合warmup前几个epoch用很小的学习率热身后面再逐步增大这在大模型里很常用小模型也可以借鉴。6.5 数据没做标准化Transformer也会摆烂虽然Transformer内部有LayerNorm但输入数据的尺度最好还是归一化到0附近。如果原始信号幅值在几百上千输入投影层的权重更新会很敏感训练前期特别容易震荡。我在代码里生成数据时直接把幅值控制在1附近就是为了省这一步但你换成自己的数据集时一定记得先做标准化。写到这里Matlab里手写Transformer做数据分类的整个流程算是完整走了一遍。我再分享一个工作习惯每次拿到新数据我都是先跑通一个最小Transformer能过拟合训练集再开始调参。如果连训练集都学不进去问题多半出在代码或数据预处理上而不是模型容量不够。这一步排查顺序能帮你省掉大量瞎调参的时间。