前言关于GAT参考知乎https://zhuanlan.zhihu.com/p/81350196

GAT与GATEAU层的区别

要解释这个问题就是主要解释这张图:
在这里插入图片描述

什么是GAT

注意力权重计算公式

公式定义

e i j = a ( [ W h i ∥ W h j ] ) , j ∈ N i (1) e_{ij} = a\left( [W h_i \parallel W h_j] \right), \quad j \in \mathcal{N}_i \tag{1} eij=a([WhiWhj]),jNi(1)
α i j = exp ⁡ ( LeakyReLU ( e i j ) ) ∑ k ∈ N i exp ⁡ ( LeakyReLU ( e i k ) ) (2) \alpha_{ij} = \frac{\exp\left(\text{LeakyReLU}(e_{ij})\right)}{\sum_{k \in \mathcal{N}_i} \exp\left(\text{LeakyReLU}(e_{ik})\right)} \tag{2} αij=kNiexp(LeakyReLU(eik))exp(LeakyReLU(eij))(2)

参数说明

  • 输入
  • 特征增强:通过共享参数 W 的线性映射对顶点特征增维。
    • e i j e_{ij} eij: 节点对 ( i , j ) (i,j) (i,j) 的原始注意力得分。
    • N i \mathcal{N}_i Ni: 节点 i i i 的邻居集合。
  • 输出
    • α i j \alpha_{ij} αij: 归一化后的注意力权重,满足 ∑ j ∈ N i α i j = 1 \sum_{j \in \mathcal{N}_i} \alpha_{ij} = 1 jNiαij=1

应用场景

  • 图神经网络(如GAT、GATEAU)中的特征聚合。
  • 社交网络分析、推荐系统等需要动态权重的场景。

GATEAU层(Graph Attention neTwork with Edge features from Attention weight Updates)​

核心改进

GATEAU是对传统GAT的扩展,​同时处理节点特征(h)和边特征(g)​,通过引入边特征增强棋类游戏移动逻辑建模能力。


公式解析

边特征更新(公式4)

g i j ′ = W u h i + W e g i j + W v h j (3) g_{ij}' = W_u h_i + W_e g_{ij} + W_v h_j \tag{3} gij=Wuhi+Wegij+Wvhj(3)

  • 输入
    • 源节点特征 $ h_i $
    • 目标节点特征 $ h_j $
    • 原始边特征 $ g_{ij} $
  • 参数
    • $ W_u, W_v $: 节点到边的映射权重
    • $ W_e $: 边特征自更新权重

注意力权重计算(公式5)

α i j = softmax j ( LeakyReLU ( a T g i ⋅ ′ ) ) (4) \alpha_{ij} = \text{softmax}_j \left( \text{LeakyReLU}(a^T g_{i\cdot}') \right) \tag{4} αij=softmaxj(LeakyReLU(aTgi))(4)

  • 功能:通过更新后的边特征 $ g_{ij}’ $ 计算节点间注意力权重

节点特征更新(公式6)

h i ′ = W 0 h i + ∑ j ∈ N i α i j ( W h h j + W g g i j ) (5) h_i' = W_0 h_i + \sum_{j \in \mathcal{N}_i} \alpha_{ij} (W_h h_j + W_g g_{ij}) \tag{5} hi=W0hi+jNiαij(Whhj+Wggij)(5)

  • 设计
    • 残差连接 $ W_0 h_i $
    • 聚合邻居节点与边特征 $ W_h h_j + W_g g_{ij} $

2. 节点(h)与边(g)的表示

节点(h)

  • 定义:棋盘格子作为节点
  • 特征​(表1):
    • 棋子类型(12维软独热编码)
    • 历史位置(过去7步状态)
    • 玩家信息(王车易位权等)
    • 游戏状态(移动次数等)

边(g)

  • 定义:合法移动作为边
  • 特征​(表2):
    • 移动合法性(1维)
    • 方向信息(左右/上下步数)
    • 升变类型(兵升后等)
    • 特殊移动规则(王、兵等)

3. ResGATEAU(残差GATEAU层)

结构(公式9)

ResGATEAU ( h , g ) = ( h , g ) + GATEAU ( BNR ( GATEAU ( BNR ( h , g ) ) ) ) ( ) \text{ResGATEAU}(h, g) = (h, g) + \text{GATEAU}(\text{BNR}(\text{GATEAU}(\text{BNR}(h, g)))) \tag{} ResGATEAU(h,g)=(h,g)+GATEAU(BNR(GATEAU(BNR(h,g))))()

  • BNR层:BatchNorm + ReLU
  • 功能:残差连接缓解梯度消失

4. 价值头(Value Head)

注意力池化(公式7)

α i p = softmax i ( LeakyReLU ( a T h i ) ) \alpha_i^p = \text{softmax}_i(\text{LeakyReLU}(a^T h_i)) αip=softmaxi(LeakyReLU(aThi))
H = ∑ i α i p h i (7) H = \sum_i \alpha_i^p h_i \tag{7} H=iαiphi(7)

  • 输出:标量值 $ v \in [-1, 1] $ 表示胜率

5. 策略头(Policy Head)

实现步骤

  1. 边特征映射:用 $ g’ $ 生成边logits
    logit ( g i j ′ ) = W p g i j ′ \text{logit}(g_{ij}') = W_p g_{ij}' logit(gij)=Wpgij
  2. Softmax归一化:生成移动概率分布
    π i j = exp ⁡ ( logit ( g i j ′ ) ) ∑ k ∈ M i exp ⁡ ( logit ( g i k ′ ) ) \pi_{ij} = \frac{\exp(\text{logit}(g_{ij}'))}{\sum_{k \in \mathcal{M}_i} \exp(\text{logit}(g_{ik}'))} πij=kMiexp(logit(gik))exp(logit(gij))
  • 优势:直接基于边特征计算策略,天然支持多棋盘尺寸

关键总结

  • GATEAU:联合更新节点和边特征(公式4-6)
  • ResGATEAU:残差结构提升训练稳定性(公式9)
  • 注意力机制:全局特征聚合(公式7)与策略生成(Policy Head)
  • 实验效果:AlphaGateau在参数较少时,训练速度更快且泛化更强(5x5→8x8棋盘迁移)
Logo

开源鸿蒙跨平台开发社区汇聚开发者与厂商,共建“一次开发,多端部署”的开源生态,致力于降低跨端开发门槛,推动万物智联创新。

更多推荐