ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

Matlab实现GCN图卷积分类:邻接矩阵与特征聚合实战

Matlab实现GCN图卷积分类:邻接矩阵与特征聚合实战 简介这份压缩包为基于图卷积神经网络GCN的节点分类任务提供了一整套可直接运行的 Matlab 实现适用于计算机、电子信息工程、数学等专业学生在课程设计、期末大作业及毕业设计中快速上手图数据分类研究。包内包含十个文件涵盖 Matlab 脚本.m、交互式脚本.mlx、示例数据.mat、邻接矩阵转换与 Glorot 初始化等工具函数以及分类结果可视化图片可帮助使用者从数据准备、模型训练到结果展示完成闭环验证。压缩包大小约 4.96MB已有 174 人学习浏览。代码采用参数化编程并附详细注释支持 Matlab 2014/2019a/2024a 多个版本用户可直接运行 qm7 示例数据进行 GCN 节点分类实验也可修改参数适配其他图结构数据。整体资源结构清晰、案例完整对初学者理解 GCN 原理和开展图数据分类实践具有较高参考价值。1. 这个“GCN分类”方案能解决什么问题给数据分类贴上“图结构”的标签做数据分类的人迟早会遇到一类数据样本和样本之间是有连线关系的。城市交通站点、传感器阵列、论文引用网络特征挂在节点上样本之间又存在明确的邻接关系。传统BP神经网络假设样本独立卷积神经网络又要求网格状规则输入这两类网络在这类非欧几里得数据上都不顺手。图卷积神经网络GCN把邻接矩阵和节点特征一起喂进网络用邻居信息反复更新节点表示成了图数据分类的主流做法。这个标题给的是一套完整、可以直接跑的方案GCN分类核心代码、Matlab实现、数据分类验证流程整体打包成rar压缩包。Matlab里没有开箱即用的官方图卷积层所以这套代码里的构图、归一化、前向传播、反向传播基本都是手写的能看能改。适合两类人第一类是学校里有Matlab授权、论文里需要图数据实验结果的师生第二类是公司里已有Matlab数据管线不想为了一个分类任务引入整套Python环境的人。下面按“原理 → 核心实现 → 数据构图 → 避坑 → 验证”的顺序展开先用最小例子跑通再替换成自己的数据。2. 图卷积的前向计算从邻接矩阵到两层GCN分类器的数学过程2.1 为什么是“图卷积”而不是“图全连接”邻居聚合的数学表达图卷积这个名字听起来吓人剥开看核心运算只有一个让每个节点沿着边把自己的特征和邻居的特征加权求和再经过一个非线性激活函数。这一层做完节点表示里就带上了邻居的信息再做一层就带上了“邻居的邻居”的信息。两层GCN能覆盖到二跳邻居对大多数分类任务已经够用。用矩阵语言描述图卷积层的前向传播写成H^(l1) ReLU(Ã H^(l) W^(l))其中 H^(l) 是第 l 层的节点特征矩阵行数是节点数 n列数是该层特征维度W^(l) 是该层的可学习参数矩阵最关键的是 Ã它是由邻接矩阵 A 加自环、归一化得到的传播算子。如果不加自环每个节点在聚合时只会看到邻居自己上一层的特征反而被丢掉了所以一般先构造 A I。归一化则是为了防止“社交达人”节点——也就是度数特别高的节点——在聚合时把特征尺度拉得过大导致训练不稳定。为什么用对称归一化而不用普通平均归一化普通做法是 D⁻¹A每个邻居的贡献按中心节点度数平均相当于只从目标节点视角做一次归一化。对称归一化写成 D⁻¹ᐟ²(AI)D⁻¹ᐟ²同时除以两端节点度数平方根的乘积。这样一来低度数节点之间的边不会被过度放大高度数节点也不会被自身度数过度稀释。这是图卷积里最常见的选择也是标题这套代码里用的形式。2.2 两层GCN做分类从输入特征到Softmax输出的完整张量流动把分类任务套进去两层GCN的完整前向过程是H ReLU(A_hat X W1)Z softmax(A_hat H W2)这里 X 是 n×d 的原始特征矩阵W1 是 d×h 的第一层参数W2 是 h×c 的第二层参数c 是类别数。第一层把原始特征映射到隐藏空间第二层把隐藏表示映射到类别空间最后 softmax 输出每个节点的类别概率。损失函数用交叉熵只对训练集掩码范围内的节点计算。用手推一个小例子看尺度假设图有 4 个节点每个节点 2 维特征X 是 4×2W1 取 2×3则 A_hat X 得到 4×2再乘 W1 得到 4×3 的隐藏表示再经过第二层 A_hat H 得到 4×3乘 W2 得到 4×c。每一步矩阵乘法都在做同一件事把邻居的特征沿边“搬运”到中心节点上。如果某个节点的两个邻居特征分别是 [1,0] 和 [0,1]聚合后这个节点就拿到了 [1,1] 的组合信息两层之后这个组合信息被映射到类别空间分类依据就不再只是节点自身特征而是它在一张图里的局部结构。2.3 和传统模型放在一起比BP、CNN、GCN各自的前提假设模型数据假设特征聚合方式典型适用场景BP神经网络样本独立同分布无聚合逐样本计算表格数据、拟合曲线CNN规则网格结构固定大小卷积核滑窗图像、语音、时序GCN任意图结构邻接矩阵驱动的邻居聚合社交网络、引文网络、传感器网络BP神经网络做拟合曲线很顺手但把图数据拉平成向量喂进去会丢掉结构信息。CNN要求输入是规则网格图像可以社交网络不行。GCN把结构信息编码进邻接矩阵理论上讲只要你能写出邻接矩阵它就能学。这也是标题里“数据分类”四个字的落脚点——如果你的数据天然有连线关系或者可以通过相似度构造出连线关系GCN就比前两者多了一个维度可以利用。2.4 层数不是越深越好图卷积的过平滑与模型容量GCN常见配置是两层最多三层。原因是一个叫“过平滑”的现象随着层数加深每个节点的表示反复与邻居聚合趋同于整个连通分量的平均信息节点之间的差异性被磨平分类精度反而断崖式下跌。这和CNN可以堆几十层完全不同。所以标题这套代码里如果只看到两层GCN不要觉得“浅”。在图上两层覆盖二跳邻居已经能捕捉到大部分局部结构信息。非要加深的话需要配合残差连接、跳跃连接或者归一化技巧那是进阶话题先在两层上把流程跑通更重要。参数规模上两层GCN的参数只有 W1 和 W2 两个矩阵比同容量的BP和CNN小得多从这个角度看图分类任务真正吃数据的部分是邻接矩阵的构造质量而不是模型参数数量。3. 用Matlab从零实现GCN前向传播、反向传播与训练主循环3.1 先把邻接矩阵变成归一化传播算子A_hat的Matlab写法拿到邻接矩阵 A 之后第一步不是直接喂给网络而是做加自环和对称归一化。这也是整个实现里最容易写错的地方。Matlab里推荐全程用稀疏矩阵节点数量上千时两者速度能差一个数量级。% A: 稀疏邻接矩阵 (n x n)非零值表示节点之间有边 n size(A, 1); A2 A speye(n); % 加自环speye 生成稀疏单位矩阵 d sum(A2, 2); % 计算每个节点的度数得到 n x 1 列向量 D_inv_sqrt spdiags(1 ./ sqrt(d), 0, n, n); % 度数的 -1/2 次方对角矩阵 A_hat D_inv_sqrt * A2 * D_inv_sqrt; % 对称归一化传播算子逻辑说明A2 在原始邻接矩阵对角线补 1让节点聚合时把自己的特征也保留一份。sum(A2, 2) 按行求和得到每个节点的度数。spdiags 把一个列向量放到稀疏对角阵上1./sqrt(d) 是逐元素的倒数开方。最后三步连乘得到完整的 A_hat。这里用 * 而不是 .* 因为我们要做的是矩阵乘法不是逐元素乘法。如果 d 里某个元素为 0说明存在孤立节点1./sqrt(0) 会给出 Inf后面就会出问题这种情况建议在构图阶段就把孤立节点过滤掉。3.2 两层GCN前向传播一个函数搞定前向传播是整个模型最直观的部分把它封装成独立函数后面反反复复调试会轻松很多。function [Z, cache] gcn_forward(X, A_hat, W1, W2) % X: n x d 节点特征矩阵 % A_hat: n x n 归一化邻接矩阵稀疏 % W1: d x h 第一层参数 % W2: h x c 第二层参数 % Z: n x c 每个节点的类别概率 H1 A_hat * X * W1; % 第一层线性变换n x h A1 max(H1, 0); % ReLU 激活n x h H2 A_hat * A1 * W2; % 第二层线性变换n x c Z softmax(H2); % 概率输出n x c cache struct(H1, H1, A1, A1, H2, H2); % 缓存中间量给反向传播用 end逻辑说明A_hat 首先作用于特征矩阵 X完成邻居特征聚合然后乘 W1 做线性映射ReLU 激活后得到第一层隐藏表示 A1。第二层重复同样的聚合和线性映射最后用 softmax 转成概率。cache 里缓存了 H1、A1、H2 三个中间量反向传播时要靠它们计算梯度省得再算一遍前向。softmax 在 Matlab 里没有内置函数手写时要注意数值稳定性。直接对 H2 做 exp 经常会溢出因为矩阵里可能出现几十甚至上百的值exp(100) 直接就是 Inf。常规做法是每行减去该行的最大值再取指数function Z softmax(X) X X - max(X, [], 2); % 每行减去最大值防止 exp 溢出 eX exp(X); Z eX ./ sum(eX, 2); end参数说明max(X, [], 2) 表示按行取最大值没写第三个参数的话默认是按列。减去最大值不会改变 softmax 的结果因为分子分母同时缩放了同一个倍数但指数运算的输入被压到了非正数范围exp 的输出最大只有 1彻底避开溢出。3.3 反向传播与交叉熵损失一组能直接拿去用的梯度公式手写反向传播是这套代码里技术含量最高的部分。好在两层GCN结构简单梯度公式可以一步步推导出来。注意这里的 dZ 是 softmax 和交叉熵联合后的梯度不用分开算两步。% Y_onehot: n x c 独热标签矩阵 % mask: n x 1 逻辑向量true 表示该节点属于训练集 m sum(mask); % 训练集样本数 loss -sum(sum(Y_onehot .* log(Z 1e-8), 2) .* mask) / m; % 交叉熵 % Softmax 交叉熵的联合梯度 dZ (Z - Y_onehot) .* mask / m; % n x c % 第二层参数梯度链式法则 dW2 A1 * A_hat * dZ; % h x c dA1 A_hat * dZ * W2; % n x h % ReLU 反向负半轴梯度为 0 dH1 dA1 .* (H1 0); % n x h % 第一层参数梯度 dW1 X * A_hat * dH1; % d x h逻辑说明loss 只对 mask 选中的训练节点计算未标记节点的误差被 mask 乘 0分母除以 m保证 batch 大小变化时损失尺度一致。dZ 是 softmax 输出与独热标签的差这是 softmax交叉熵组合的已知结论可以直接用。dW2 的维度是 h×c和 W2 完全一致dW1 的维度是 d×h和 W1 完全一致写完后务必检查一遍维度对不上说明某个环节转置错了。这里有个细节mask 必须是 n×1 的列向量。Matlab R2016b 之后的版本支持隐式扩展所以 (Z - Y_onehot) .* mask 会把列向量自动扩展到每一列。如果 mask 是 1×n 的行向量维度匹配不上会直接报错或者更隐蔽地给出错误结果。我一般习惯在数据准备阶段统一用 mask mask(:) 强制转成列向量。3.4 训练主循环与参数初始化rng、学习率、断点输出训练脚本的主体是一个标准的梯度下降循环。这里的关键是固定随机种子否则每次运行结果都不一样无法对比实验。rng(42); % 固定随机种子保证结果可复现 W1 0.1 * randn(d, 64); % 第一层参数d 是特征维度 W2 0.1 * randn(64, c); % 第二层参数c 是类别数 lr 0.01; % 学习率 for epoch 1:300 [Z, cache] gcn_forward(X, A_hat, W1, W2); % 计算交叉熵损失代码见 3.3 节 m sum(mask); loss -sum(sum(Y_onehot .* log(Z 1e-8), 2) .* mask) / m; % 计算梯度代码见 3.3 节省略中间变量 dZ (Z - Y_onehot) .* mask / m; dW2 cache.A1 * A_hat * dZ; dA1 A_hat * dZ * W2; dH1 dA1 .* (cache.H1 0); dW1 X * A_hat * dH1; % 梯度下降更新 W1 W1 - lr * dW1; W2 W2 - lr * dW2; if mod(epoch, 20) 0 fprintf(epoch %d, loss %.4f\n, epoch, loss); end end参数说明初始化的幅值 0.1 是一个经验值太大会导致早期输出集中在 softmax 饱和区梯度极小太小又会让训练变慢。隐藏维度 64 对大多数中小规模分类任务足够节点多、特征复杂时可以调到 128 或 256。学习率 0.01 在两层小模型上比较稳如果发现 loss 震荡先降到 0.005。fprintf 每 20 个 epoch 打一次日志既能看到收敛趋势又不会刷屏。从输出里应该看到 loss 平稳下降如果出现 NaN 或剧烈震荡直接跳到第 5 章排查。4. 数据分类的构图落地邻接矩阵构造、标签对齐与训练集划分4.1 你的数据有没有“图”两类数据的构图路线GCN 分类的第一步不是调参而是想清楚数据以什么形式进入模型。现实中碰到的情况大概分两类。第一类是数据天然带图结构论文引用网络里论文是节点、引用关系是边社交网络里用户是节点、关注关系是边交通路网里路口是节点、道路是边。这类数据只需要把边列表转成邻接矩阵再补上节点特征矩阵就可以直接进入第 3 章的前向流程。第二类是普通表格数据每一行是一个样本每列是一个特征没有显式的边。想把 GCN 用起来就要先按特征相似度构图。这也是“数据分类”这个标题下最容易被忽略的一步——很多人把表格数据直接塞进 GCN结果不如 XGBoost原因不是模型不行而是图本身没构造好。GCN 在图上做的是特征平滑如果边连接的是相似样本平滑能够抑制噪声特征如果边连接的是语义相反的样本平滑反而把特征搞混模型自然学不好。4.2 从特征矩阵到邻接矩阵KNN构图与距离阈值构图的Matlab实现给普通表格数据构图常见做法有两种KNN 图和距离阈值图。KNN 图保证每个节点恰好有 k 个邻居图的连通性相对稳定距离阈值图只连接距离小于阈值的节点对稀疏程度由阈值控制。先对特征做标准化再算距离否则量纲大的特征会主导相似度计算。% X: n x d 原始特征矩阵先标准化 X_std (X - mean(X, 1)) ./ std(X, 1); % 手动 zscore按列标准化 % 方法一KNN 构图 k 5; % 邻居数 D pdist2(X_std, X_std, euclidean); % n x n 距离矩阵 [~, idx] sort(D, 2); % 每行按距离升序排列 A zeros(n, n); for ii 1:n A(ii, idx(ii, 2:k1)) 1; % 自己到自己的距离为0跳过第1列 end A max(A, A); % 对称化你有我我就有你逻辑说明pdist2 计算两两欧氏距离返回稠密矩阵节点数超过五千时这一步内存压力很大可改用 knnsearch 分块查询。sort 之后每行第 1 列是自己所以取第 2 到 k1 列作为邻居。对称化这一步不能省因为 KNN 关系不是天然对称的——A 是 B 的最近邻不代表 B 是 A 的最近邻而 GCN 的邻接矩阵应该是对称的否则传播方向就有偏置。% 方法二高斯核权重 距离阈值截断 sigma 1.0; % 高斯核带宽 W exp(-D.^2 / (2 * sigma^2)); % 距离越近权重越大 W(D 1.5) 0; % 超过阈值的边直接断开 A (W W) / 2; % 对称化并取平均权重参数说明sigma 控制权重随距离衰减的快慢一般取所有样本对距离的标准差或者用交叉验证调。阈值 1.5 表示只保留距离小于 1.5 倍单位距离的边实际使用时可以改成 prctile(D(:), 10) 这类按分位数截断的方式保证图的边数符合预期。高斯核构图比 KNN 更精细保留了边的权重信息但多了一个带宽参数要调。小数据集上我习惯先用 KNN 图跑通流程再对比换成高斯核效果有没有提升。4.3 标签对齐与mask划分训练集、验证集、测试集按节点索引切构图完成之后紧接着是一个让不少人翻车的环节标签对齐。邻接矩阵的每一行对应一个节点特征矩阵的每一行也对应一个节点标签必须和这个顺序完全一致。如果中间对数据做过排序、去重、拼接很容易出现特征行序和标签行序错开的情况GCN 带着错位的标签训练还能正常收敛只是精度上不去很难察觉。rng(2024); n size(X, 1); idx randperm(n); % 随机打乱节点索引 n_train round(n * 0.6); % 60% 训练 n_val round(n * 0.2); % 20% 验证 train_mask false(n, 1); val_mask false(n, 1); test_mask false(n, 1); train_mask(idx(1:n_train)) true; val_mask(idx(n_train1:n_trainn_val)) true; test_mask(idx(n_trainn_val1:end)) true;逻辑说明mask 是 n×1 的逻辑向量true 位置对应的节点属于该集合。这是图节点分类的标准做法——训练、验证、测试都在同一张图上只是用 mask 区分哪些节点参与计算。和传统机器学习随机切分样本不同图上的邻居关系是共享的测试集节点仍然会出现在训练节点的聚合范围里这是 GCN 的设定不用纠结。关键点是 idx 用 randperm 生成后train_mask、val_mask、test_mask 全部基于同一个 idx 划分三者互不重叠。如果分开用三次 randperm就会有一些节点既在训练集又在测试集精度虚高论文里容易被审稿人追问。4.4 特征标准化为什么 zscore 比 min-max 归一化更稳特征尺度对 GCN 的影响比普通神经网络更明显。原因是特征矩阵 X 在第一层要乘 A_hat等于每个节点的特征被邻居特征加权平均。如果某个特征列的范围是 0 到 1000另一列是 0 到 1前者的贡献在聚合后被放大了一千倍模型的实际输入被这一列主导其他特征形同虚设。zscore 按列减去均值、除以标准差让每个特征列的尺度都在 1 附近。min-max 归一化也能把特征压到 0-1但它对离群点敏感一个极端值会把其他所有样本的值压缩到很窄的区间。GCN 的聚合操作天然会放大异常值的影响所以 zscore 是更稳妥的选择。上面代码里用 (X - mean(X, 1)) ./ std(X, 1) 手动实现是因为 zscore 函数需要 Statistics and Machine Learning Toolbox手动写几行不依赖任何工具箱换机器跑也不会报错。5. GCN分类在Matlab中的避坑记录5个高频问题与排查方法5.1 训练第一个epoch就出现NaN现象fprintf 输出的 loss 第一轮就是 NaN或者前几轮正常突然变成 NaN 后再也没恢复。原因最常见的是 softmax 溢出。H2 矩阵里出现较大数值时exp 直接上溢成 InfInf/Inf 得到 NaN。另一个常见来源是特征矩阵 X 里本身含有 NaN 或 Inf进入 A_hat * X 之后被邻居聚合放大。解决先检查 X 里有没有 NaN用 sum(isnan(X), all) 看一眼。排除数据问题后确认 softmax 里有没有做减去行最大值的操作这步能挡住绝大多数数值溢出。如果还不行把初始化幅值从 0.1 降到 0.01降低第一轮 H2 的数值范围。最后检查构图孤立节点会让 d 中出现 01/sqrt(0) 直接产生 Inf。5.2 训练loss在下降测试精度却一直上不去现象训练集准确率能到 90% 以上测试集始终在 50% 附近徘徊甚至随着训练震荡。原因九成情况是标签顺序和节点顺序没对齐。pdist2 算距离时用的是 X 的行顺序如果你之前对表格做过 sortrows而标签向量没有同步排序错位样本在图上会学到完全错误的邻居关系。解决从数据加载开始就维护一个统一的节点索引特征、标签用同一个索引顺序。检查邻接矩阵 A 和标签 Y 是不是按同一个 id 排序的。另一个可能原因是构图时 k 选得太小图上出现大量孤立的小连通分量信息传不远模型退化成只看自身特征。把 k 从 5 调到 10 或 15观察测试精度是否回升。5.3 节点到几千个以后训练慢到无法忍受现象一千个节点以内跑得还算流畅节点数到五千或者一万一次前向要等几十秒内存占用飙升。原因用了 pdist2 生成稠密距离矩阵n10000 时距离矩阵就是 10000×10000占 800MB 内存构图阶段直接把机器拖垮。邻接矩阵 A 本身应该是极度稀疏的但中间变量 W exp(-D.^2) 会把所有节点对都算一遍等于把稀疏图强制变成稠密图。解决构图阶段用 knnsearch 代替 pdist2knnsearch 内部用 KD 树只需要返回每个节点的前 k 个邻居不会生成完整的距离矩阵。% 需要 Statistics and Machine Learning Toolbox [nbr_idx, nbr_dist] knnsearch(X_std, X_std, K, k1); A sparse(n, n); for ii 1:n A(ii, nbr_idx(ii, 2:end)) 1; end A max(A, A); A_hat sparse(A_hat); % 确保后续参与运算的是稀疏矩阵逻辑说明knnsearch 返回每个节点的 k 个最近邻索引循环里只填稀疏矩阵的非零位置内存占用从 O(n²) 降为 O(n·k)。注意最后再用 sparse() 显式转换一次防止前面某些操作把稀疏矩阵转成稠密的。如果没有统计工具箱可以按块计算距离例如每 500 个节点一批算完只保留每行的前 k 个最小值效果一样但代码略长。5.4 同一份数据每次跑出来的精度都不一样现象什么都没改只是重新运行脚本测试精度差了 5 到 10 个百分点。原因随机种子没有固定。参数初始化用 randn训练集划分用 randperm这两处每次运行都会生成不同的随机数。本身这不是 bug但如果想对比不同参数的效果变量被随机性污染结论不可信。解决在脚本最开头加一行 rng(42)固定全局随机种子。训练集划分那里也显式调用一次 rng。保存模型参数时把 mask 和种子一起存下来方便下次加载复现。rng(42); save(gcn_run1.mat, W1, W2, mask, A_hat); % 后续加载load(gcn_run1.mat)参数说明种子值 42 只是个习惯也可以换成任意整数。关键是两次实验之间保持同一套划分和初始化。需要对比不同 k 值时每次把参数存成一个独立文件最后汇总比较测试精度。5.5 打开别人的.m文件中文注释变成乱码现象下载的代码在 Matlab 2023 之后的版本里打开中文注释显示成乱码但代码本身能跑。原因文件保存时的编码格式和当前 Matlab 版本默认的编码格式不一致。老版本脚本默认 GBKMatlab 2023 之后默认 UTF-8编译器按 UTF-8 去解码 GBK 的字节流中文自然乱掉。解决在 Matlab 主页找到“预设项”进入“编辑器 → 语言”把 MATLAB 文件的默认编码改成与源文件一致或者直接用文本编辑器把文件转存成 UTF-8 编码。更省心的做法是代码注释尽量用英文。给不会读英文的人一个折中方案——关键注释用拼音缩写配合变量名的语义也能把代码逻辑看懂。不要在这个问题上花太多时间它不影响任何数值结果。6. 训练完成之后怎么验证数值梯度检查与三层信号确认6.1 三个信号判断模型是在学习还是在背答案训练跑完先别急着把测试精度写进报告。我一般会看三个信号训练 loss 是不是平稳下降、验证集精度是否随训练轮次同步上升、测试集错误样本是不是集中在少数几个类。如果 loss 下降但验证精度纹丝不动大概率是过拟合给 GCN 加 dropout 或者减小隐藏维度如果错误样本集中在某个类上去看这个类的样本数是不是特别少图上的该类节点是不是高度稀疏。把 loss 曲线画出来是最直观的验证figure; plot(1:epochs, loss_history, LineWidth, 1.5); xlabel(epoch); ylabel(loss); title(训练loss曲线); grid on;曲线应该是一条单调下降、尾部趋于平坦的线。如果出现明显的周期性震荡学习率偏大如果下降极其缓慢学习率偏小。这个图比任何精度数字都能说明问题。6.2 数值梯度检查验证反向传播没写错的唯一可靠方法手写反向传播最怕的不是公式推导错而是某个转置或者下标写错代码照样能跑梯度方向大致正确但不精确模型训练半天停在次优位置。数值梯度检查是唯一的后悔药——用差分近似计算梯度和手写解析梯度比较误差在 1e-6 量级说明反向传播正确。epsilon 1e-6; num_grad zeros(size(W1)); for ii 1:numel(W1) W1_p W1; W1_m W1; W1_p(ii) W1_p(ii) epsilon; W1_m(ii) W1_m(ii) - epsilon; % compute_loss 需要把前向传播和交叉熵封装成一个函数 loss_p compute_loss(W1_p, W2, X, A_hat, Y_onehot, mask); loss_m compute_loss(W1_m, W2, X, A_hat, Y_onehot, mask); num_grad(ii) (loss_p - loss_m) / (2 * epsilon); end rel_err norm(num_grad - dW1(:)) / (norm(num_grad) norm(dW1(:))); fprintf(相对误差: %.2e\n, rel_err);逻辑说明中心差分公式比前向差分精度高一个量级是检查梯度时的标准做法。相对误差小于 1e-6 说明解析梯度正确1e-4 以内还能用超过 1e-2 基本可以断定反向传播有 bug。检查用小模型跑比如 20 个节点、隐藏维度 5全量参数也就一百多个几秒就能算完。跑通这套验证流程之后这个 GCN 分类方案才算真正做完。下一步往哪个方向走取决于你的数据规模图很大就研究邻接矩阵的 mini-batch 采样图很小就尝试换一种构图方式或者加一层 dropout。我自己的习惯是每换一个数据集先跑一遍数值梯度检查再对照三个训练信号判断模型状态——这套流程能挡住大多数暗坑。希望帮到你。本文还有配套的精品资源点击获取
返回列表