ARTICLE DETAIL

资讯详情

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

Watch Your Step 图注意力嵌入模型:论文复现、环境搭建与源码级参数解析

Watch Your Step 图注意力嵌入模型:论文复现、环境搭建与源码级参数解析 人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载导读本文围绕 Google Research 开源仓库中graph_embedding/watch_your_step目录下的官方实现完整讲解 NIPS 2018 论文《Watch Your Step: Learning Node Embeddings via Graph Attention》的复现流程包括环境搭建、数据集准备、命令行运行、全部可调参数及其底层实现原理。读完本文你将掌握如何在本地一键训练带上下文注意力的图节点嵌入模型能够读懂其基于转移矩阵幂与可学习注意力权重的核心算法并学会解读训练过程中输出的 AUC 指标与学到的上下文分布。项目定位与论文背景Watch Your Step是 Google Research 仓库中图表示学习graph embedding方向的经典实现对应论文《Watch Your Step: Learning Node Embeddings via Graph Attention》Sami Abu-El-Haija、Bryan Perozzi、Rami Al-Rfou、Alex AlemiNIPS 2018。该目录位于 graph_embedding/watch_your_step主要包含以下文件README.md官方使用说明本文的主体依据graph_attention_learning.py完整训练/评估实现约 460 行requirements.txt依赖清单run.sh一键冒烟测试脚本。核心思想传统基于随机游走的图嵌入方法如 DeepWalk依赖一个固定大小的上下文窗口context window来定义节点共现窗口大小是人工设定的超参数。Watch Your Step 的核心创新在于将上下文距离分布本身变成一组可学习的注意力权重context distribution / attention让模型自己决定在随机游走中应该关注多远的邻居从而自动选择正确的上下文。实现上这组权重以 softmax 归一化的方式作用在归一化邻接矩阵转移矩阵的各次幂上构成一个参数化的共现矩阵期望E[D; q]。环境准备与依赖安装按照官方 README建议先创建一个全新的 Python 3 虚拟环境再从仓库根目录google-research/安装依赖# 从 google-research/ 根目录执行 virtualenv -p python3 . source ./bin/activate pip install -r graph_embedding/watch_your_step/requirements.txt依赖清单requirements.txt的关键条目如下依赖版本用途tensorflow1.12.1深度学习框架模型训练主体absl-py0.6.1命令行 flags 解析app、flags、loggingnumpy1.15.4邻接矩阵、转移矩阵幂的数值计算与磁盘缓存scikit-learn0.20.0链接预测评估指标roc_auc_scorescipy/h5py/protobuf等固定版本TensorFlow 1.x 运行依赖需要注意的兼容性事实源码在 graph_attention_learning.py 中通过import tensorflow.compat.v1 as tf与from tensorflow.contrib import slim as contrib_slim同时兼容 TF1 与 TF2 环境但由于tensorflow.contrib自 TensorFlow 2.x 起已被移除实际运行时建议使用1.12.1的 TensorFlow 1.x 环境以获得完整功能。数据集准备官方 README 指定使用论文作者在 CIKM17 工作中使用的图数据集Abu-El-Haija et al., CIKM17通过如下命令下载并解压# 从 google-research/ 根目录执行 curl http://sami.haija.org/graph/datasets.tgz datasets.tgz tar zxvf datasets.tgz export DATA_DIRdatasets解压后得到的datasets目录应包含多个数据集子目录如wiki-vote。结合 graph_attention_learning.py 中的dataset_dir说明每个数据集目录内需要预置以下.npy/.pkl文件文件作用index.pkl节点索引字典含index键用于确定节点总数NUM_NODES见 GetNumNodestrain.txt.npy训练正样本边对每行[源节点, 目标节点]test.txt.npy测试正样本边对train.neg.txt.npy训练负样本边对test.neg.txt.npy测试负样本边对若图是有向图则还需test.directed.neg.txt.npy其中是否有向的判定逻辑位于 IsDirected只要数据目录中存在test.directed.neg.txt.npy文件即视为有向图。运行过程中程序还会自动生成并缓存两个中间产物由训练边构造的邻接矩阵a.npy见 GetOrMakeAdjacencyMatrix以及转移矩阵的各次幂t_i.npy见IterPowerTransitionPairs二次运行时可显著加速。运行模型依赖与数据就绪后从仓库根目录执行# 从 google-research/ 根目录执行 python -m graph_embedding.watch_your_step.graph_attention_learning --dataset_dir ${DATA_DIR}/wiki-vote官方 README 说明上述命令只需几分钟即可复现论文中的主要结果。若希望保存输出请追加--output_dir参数输出文件将包含训练/测试指标、节点嵌入embeddings以及学到的上下文分布context distributions。仓库还提供了 run.sh 一键脚本它在环境准备与数据下载之外使用快速测试参数运行python -m graph_embedding.watch_your_step.graph_attention_learning \ --dataset_dir ${DATA_DIR}/wiki-vote \ --transition_powers 2 \ --max_number_of_steps 10注意 run.sh 内的注释明确提醒--transition_powers 2 --max_number_of_steps 10并非该数据集上的推荐配置只是为了保证开源冒烟测试能在短时间内跑完。复现论文效果请使用默认参数。命令行参数详解源码级所有命令行参数均在 graph_attention_learning.py 顶部通过 absl flags 定义汇总如下参数类型默认值说明--dataset_dirstring必填无默认数据集所在目录须包含{train,test}.txt.npy与{train,test}.neg.txt.npy等文件未指定时直接报错退出flags.mark_flag_as_required--output_dirstringNone若设置则将评估指标、训练好的参数写入该目录目录不存在时会自动创建main--max_number_of_stepsint100最大梯度更新步数--learning_ratefloat0.2学习率README 中注释为 PercentDelta learning rate实际作用在 PercentDelta 梯度缩放机制中--dint4嵌入维度embedding dimensions对应左/右两个嵌入字典的维度--transition_powersint5归一化邻接转移矩阵的最高幂次即注意力分布覆盖的最大上下文步数--context_regularizerfloat0.1作用于上下文分布参数向量q的正则化系数--objectivestringnlgl目标函数二选一rmse均方误差或nlglNegative Log Graph Likelihood负对数图似然--share_embeddingsboolFalse若开启左右两个嵌入字典共享同一组参数其中--objective的选择直接影响损失构造。在 CreateObjective 中nlgl对模型打分g施加 sigmoid 后计算二分类交叉熵正样本部分使用参数化的期望矩阵E[D; q]作为权重负样本部分使用真实邻接矩阵的补集1 - true_adjacency这与论文中positive part / negative part的负对数似然设计一致rmse直接计算(g - target_matrix)^2的均值。算法原理与源码解析参数化的共现期望E[D; q]模型的核心数学结构在 GetParametrizedExpectation 中实现其定义式为E[D; q] P_0 · (Q_1·T Q_2·T² Q_3·T³ ...)其中T为归一化邻接矩阵转移矩阵由a.T / (degree 1e-7)计算得到Q softmax(q)是可学习的上下文分布q是长度为transition_powers的可训练向量源码中初始化为全 1见 graph_attention_learning.py各次幂T^k通过GetPowerTransitionPairs惰性计算并在磁盘缓存为t_i.npy首项T由邻接矩阵直接推导、不缓存IterPowerTransitionPairs该期望再乘以度向量与GetNumNodes() * 80的缩放系数作为训练目标graph_attention_learning.py。训练结束后q的 softmax 结果就是模型学到的上下文注意力分布它会告诉研究者模型认为随机游走中第几跳的邻居对共现贡献最大——这正是论文标题 Watch Your Step注意脚下的步子的含义所在。双嵌入字典与打分函数模型为每个节点维护左、右两个嵌入向量source 侧与 target 侧由 CreateEmbeddingDictionary 以U(-0.1, 0.1)均匀分布初始化并附带1e-6的 L2 正则。当--share_embeddings开启时右侧字典直接复用左侧main。打分函数g net_l · net_rᵀCreateGFn即左右嵌入的内积矩阵用于链接预测评分。PercentDelta 优化技巧值得注意的是源码没有直接使用普通的 SGD而是通过 CreateGradMultipliers 构造了一个PercentDelta百分比增量机制根据当前全局步数线性地把目标增量从 1 衰减到 0.01见 GetPD 中的推导注释并据此缩放每个变量的梯度使 SGD 的更新步长变为按参数当前幅值的百分比调整。这也是 README 中把learning_rate注释为 PercentDelta learning rate 的原因。总损失由目标损失、上下文分布正则mean(q²) * context_regularizer与嵌入 L2 正则叠加得到main。每 4 步评估一次 AUC训练循环每 4 步执行一次评估graph_attention_learning.py通过sklearn.metrics.roc_auc_score在训练/测试正负样本对上计算 AUCRunEval同时记录该时刻的qmults与 softmax 归一化后的分布normed_mults。若连续 100 步未刷新最佳训练 AUC训练提前终止graph_attention_learning.py。输出文件与结果解读当指定--output_dir时程序会生成三类文件文件名以ds.数据集名.e.维度.o.目标函数为前缀见 Description前缀.json评估指标序列包含train auc、test auc、total loss、objective loss、mults原始注意力参数、normed_multssoftmax 后的上下文分布以及best train auc、test auc at best train、i at best train等最佳值字段main前缀.best.pkl在最佳训练 AUC 时刻保存的全部可训练变量含全局步数前缀.last.pkl训练结束时最后一轮的全部变量。其中.pkl文件按(变量名, 值)配对列表序列化graph_attention_learning.py便于后续加载嵌入向量做下游任务。评估日志格式为步数 test/train auc测试AUC/训练AUC obj.loss目标损失 total.loss总损失。相关实现与延伸阅读该论文提出的思想在仓库内其他模块中也有对应实现可作为交叉验证graph_embedding/huge/model.py 中的negative_log_graph_likelihood实现了同一论文定义的 Negative Log Graph Likelihood 损失正/负样本对分别计算并聚合graph_embedding/huge/io.py 中的ComputeExpectedEdgeScore实现了论文第 3 节与 3.1 节的共现矩阵期望E[D]D_{v,u}表示在上下文距离c ~ U[1, C]内节点v与u的共现次数。若想了解更广泛的图嵌入工作可继续浏览仓库的 graph_embedding 目录下的ddgk、dmon、monet、persona、slaq等子模块。引用论文若你的研究工作中使用了 Watch Your Step官方 README 建议引用以下论文Abu-El-Haija, S., Perozzi, B., Al-Rfou, R., and Alemi, A. (2018). Watch Your Step: Learning Node Embeddings via Graph Attention. InNeural Information Processing Systems.BibTeXinproceedings{abu2018watch, author{Abu-El-Haija, Sami and Perozzi, Bryan and Al-Rfou, Rami and Alemi, Alex}, title{Watch Your Step: Learning Node Embeddings via Graph Attention}, booktitle{Neural Information Processing Systems}, year{2018}, }使用注意事项TensorFlow 版本实现基于 TensorFlow 1.x API 与tensorflow.contrib请按 requirements.txt 使用tensorflow1.12.1的 1.x 环境若在 TF2 中运行tensorflow.contrib相关代码contrib_slim可能不可用。数据集格式dataset_dir下的.npy边文件格式必须与源码读取逻辑一致train_edges[:, 0]/train_edges[:, 1]分别为源、目标节点并配套index.pkl。快速测试与复现参数run.sh 中的--transition_powers 2 --max_number_of_steps 10仅供快速验证流程复现论文结果应使用默认参数transition_powers5、max_number_of_steps100。有向图处理有向图需要额外提供test.directed.neg.txt.npy无向图会对称回填邻接矩阵GetOrMakeAdjacencyMatrix。赞分享人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载相关推荐为什么选择react-image-magnify5大优势碾压传统图片放大插件为什么选择react image magnify5大优势碾压传统图片放大插件 react image magnify是一款专为购物网站设计的响应式图片放大组件AutoGPT Forge llamafile 本地大模型接入指南环境搭建、配置参数与源码原理AutoGPT Forge llamafile 本地大模型接入指南环境搭建、配置参数与源码原理 本文基于 AutoGPT 仓库中 Forge 引擎的 llam人工智能AI Agent自主智能体Agent 工作流工作流自动化后端前端掌握transformers.js模型注意力机制自注意力与交叉注意力完整指南掌握transformers.js模型注意力机制自注意力与交叉注意力完整指南 transformers.js是一个在浏览器中直接运行Transformer人工智能大模型本地部署上一篇Tiptap水平分割线5分钟掌握文档结构分隔的终极指南下一篇3个关键问题为什么ESP32是智能硬件开发的首选平台创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表