RELIEF: Reinforcement Learning Empowered Graph Feature Prompt Tuning---论文总结
一、研究动机与核心问题
1.1 背景:从预训练-微调到预训练-提示
近年来吗,图神经网络在图表示学习中取得了巨大成功,为了应对少样本和分布外场景的挑战,“预训练-微调”范式被广泛应用。然而,该范式存在两个主要问题:
(1)目标不一致:预训练任务(如边预测、属性掩码)与下游任务(如图分类)的目标存在差异,导致微调后性能次优。
(2)灾难性遗忘:在有限的标注数据上微调整个模型,容易导致模型遗忘在预训练中学到的知识,损害其泛化能力。
受自然语言处理领域的启发,“预训练-提示”范式被引入图学习,其核心思想是:冻结预训练的GNN模型参数,通过设计“提示”来调整输入数据,从而将下游任务“伪装”成预训练任务,激活模型已有知识。
1.2 现有图提示方法的局限与我们的洞察
现有的图提示方法可分为两类:
(1)依赖训练策略的方法:
他们将所有的下游任务统一转换成边预测问题,也就是说,无论你的实际任务是节点分类、图分类还是其他类型的任务,这些方法都会尝试通过预测图中的边来间接解决问题。
局限性:
通用性差:这种方法的适用范围非常有限,因为并不是所有的任务都可以转化为边预测形式。如果下游任务与边预测不匹配(例如,如果你的任务需要识别特定类型的节点而不是边),那么模型的性能会大幅下降。
灵活性不足:由于这种转换是固定的,它限制了模型适应不同任务的能力。这意味着,对于那些本质上不是关于边的问题,模型可能无法利用其全部潜力。
(2)与预训练策略无关的方法:
此类方法以基于特征的提示为代表,如GPF和GPF-plus。他们的基本思想是在每个节点的特征向量上添加一个可学习的提示向量,以此引导模型完成特定任务。
局限性:
效率问题:有必要给所有的节点都添加提示吗?这是该方法面临的一个核心问题。虽然这种方法具有较好的通用性,因为他不依赖于特定的预训练策略,但是对每个节点都添加提示可能会导致不必要的计算开销,并且在实践中可能并非必要。
潜在的负面影响:如果提示幅度过大,可能会淹没原始输入特征空间,使得预训练好的GNN认为输入的图与其预训练时见到的数据分布有很大差异,从而损害迁移性能。换句话说,过多的提示可能会让模型忽略原本重要的信息,影响最终的表现。
1.3 核心洞察
作者从NLP领域获得灵感,提出即使是强大的预训练模型也只需要少量的“条件信号”就能激发期望的行为。应用到图学习中,意味着不需要对所有节点都添加提示。相反,只需在少数关键节点上添加轻量级的提示就足够了这可以更有效的利用预训练模型的能力,同时避免因过度修改输入而导致的负面效果。这样不仅提高了效率,还增强了模型处理不同任务的能力。
1.4 核心问题
基于上述洞察,我们的目标转变为:如何为原始图“策略性地”引入“必要且轻量”的特征提示。这引出了一个组合优化问题,我们需要一个策略来决定:
1. 给哪些节点添加提示?(离散决策)
2. 给这些节点添加什么内容的提示?(连续决策)
2、方法设计:RELIEF框架详解
我们提出了RELIEF方法,其核心思想是将提示的添加过程建模为一个序列决策问题,并使用RL来优化这个策略。(relief把加提示看作一个多轮游戏,每一轮做两件事:(1)选哪个节点(2)加什么提示,把这个新图(已加提示)输入与训练好的GNN,看下游任务(比如分类)的损失有没有下降,如果损失下降,给智能体一个正奖励(说明这步做的好)如果损失上升了,给负奖励(说明这步做的不好)。重复N次这个过程,就完成了一轮提示构建。)
2.1 问题形式转换为马尔可夫决策过程
状态:在时间步t,状态s_t定义为上一个被提示后的图输入到冻结的预训练GNN后得到的节点表示。这确保了智能体能够感知到与GNN模型紧密的相关的图状态。
动作:这是一个混合动作空间。
离散动作a:从n个节点中选择一个节点v_a。
连续动作z:为一个D维的实值向量,作为要添加到节点的特征提示p(a,z)
状态转移:当智能体执行动作(a,z)后,提示矩阵p更新(在对应节点v_a的位置加上p(a,z)),从而得到新的带提示特征x*_t=x+p。将x*_t输入预训练GNN,得到新的节点表示,即下一状态s_t+1.
- 初始特征是
X(比如每个节点是一个 64 维向量)。 - 智能体维护一个“提示矩阵”
P,初始全为 0。 - 当它决定给节点 42 加提示
z,就把P[42] += z。 - 新输入特征变成
X* = X + P。 - 把
X*喂给冻结的预训练 GNN(参数 θ 不更新!),得到新嵌入 → 这就是下一个状态s_{t+1}。
奖励函数:我们使用损失函数的下降值作为即时奖励。
r_t=l_t-1 - l_t:这个设计是目标导向的,如果当前步骤添加的提示降低了任务损失,则获得正奖励,反之获得负奖励。累积奖励反映了从开始到结束的总性能提升。
整个流程:
sₜ → 智能体选 (节点a, 提示z) → 更新 P → X* = X + P → GNN(X*) → s_{t+1}
2.2策略网络架构
我们采用H-PPO算法,其策略网络n_w包含:
1、一个输入Actor:输入状态s,输出一个概率分布,表示选择每个节点的概率。
2、一个连续Actor:输入状态s和选定的离散动作a,输出一个高斯分布的参数,从中采样得到一个连续动作z,为确保提示“轻量”,我们将z的每维度都裁剪到[-z_max,z_max].
3、一个Critic:输入状态s,输出一个标量V(s),评估当前状态的价值。
关键点:策略网路的状态编码器直接复制了预训练GNN并保持冻结。只有后面的MLP层棵学习。
这与RLHF中在冻结的LLM上加可学习层的思路一致。
策略网络有三个大脑:
离散Actor:负责选人(选哪个节点)
连续Actor:负责选内容(给选中的节点加什么样的提示向量)
Critic:负责做评价(当前这个图的状态好不好,值不值得继续优化)
三个大脑是怎样合作的:
1、离散Actor-选人:
- 输入:当前图的状态s,这个状态s实际上就是被预训练GNN处理后的所有节点的表示,你可以把它理解为GNN对当前图的一个内部看法或总结。
- 输出:一个概率列表。比如图中有5个节点,他会输出
[0.1, 0.7, 0.05, 0.1, 0.05]。这个列表的意思是:“我有70%的把握认为,给第2个节点加提示效果最好!” - 作用:解决“在哪里加提示”的问题。这是一个离散选择问题(从有限的几个节点里选一个)。
2、连续Actor-定内容:
- 输入:同样是图的状态s,以及离散Actor刚刚选中的那个节点a。
- 输出:一个具体的提示向量z。这个向量和原始节点特征维度一样,比如原始特征是128维,那z也是128维的实数向量。
- 如何生成:他先预测一个高斯分布的均值,然后从这个分布中随机采样出最终的z。这样做是为了在训练时引入一定的探索性(试试不同的提示),而不是每次都用完全一样的。
- 关键约束:为了保证提示是轻量级的(不会把原始特征搞得太面目全非),他会对z的每个数值进行裁剪,确保他们都在[-z_max,z_max]这个范围内,z_max是一个超参数,比如设为0.5,那么任何提示值都不会超过0.5或低于-0.5。
- 作用:解决了加什么样的提示的问题。这是一个连续值问题。
3、Critic-做评价
- 输入:当前图的状态(同样是GNN处理后的节点表示)
- 输出:一个标量数字V(s)。这个数字代表了评论家对当前状态的整体价值评估,比如,V(s)=10,可能意味着“这个图现在的样子很有潜力,再优化一下分数能很高”;而v(s)=-5可能意味着“完蛋了,加的提示把事情搞砸了”
- 作用:为两个Actor提供反馈信号。Actor门根据Critic的评价来调整自己的策略--如果Critic给高分,就多做类似的事,如果给低分就少做。
2.3 整体训练框架
RELIEF有两个可训练模块:策略网络和投影头,我们交替训练的策略:
1、策略网络训练阶段:
- 对于一个图,初始化提示矩阵维零
- 进行n步(节点数)添加提示,每一步,智能体根据当前策略,选择一个节点并生成提示向量,添加到图中。
- 计算每一步的即时奖励,收集n条转移经验
(s, a, z, r, s‘)。 - 用一个批次的图经验,使用PPO的目标函数,来更新离散和连续Actor,用MESE损失更新Critic。
具体训练步骤:
步骤 1: 计算优势估计
优势函数 A(s, a) 衡量了在状态 s 下采取动作 a 相比平均情况有多好。它是策略更新的核心信号。
-
使用Critic网络计算状态价值:
-
对于经验中的每一个状态
s,使用当前的Critic网络计算其价值估计V(s)。 -
对于下一个状态
s‘,同样计算V(s’)。
-
-
计算TD误差:
-
对于每一步,时间差分误差为:
δ_t = r_t + γ * V(s‘_t) - V(s_t) -
其中
γ是折扣因子。
-
-
使用GAE计算优势估计:
-
RELIEF使用广义优势估计来得到更平滑、方差更低的优势值。
-
A_t = Σ_{l=0}^{∞} (γλ)^l * δ_{t+l} -
其中
λ是GAE参数(通常0.95)。在实践中,这是一个从后向前的迭代计算。
-
步骤 2: 计算回报
为了更新Critic,我们需要一个目标值。
-
回报:
R_t = A_t + V(s_t) -
这个
R_t将作为更新Critic时的目标。
步骤 3: 计算概率比
这是PPO的核心,用于衡量新策略与旧策略的差异。
-
记录旧策略的概率:
-
在收集经验时,我们已经记录了在状态
s_t下,选择动作a_t和z_t的旧策略的概率。 -
对于离散动作:
π_d_old(a_t | s_t) -
对于连续动作:
π_c_old(z_t | s_t, a_t)(基于旧策略的高斯分布概率密度)
-
-
计算新策略的概率:
-
使用当前的策略网络(即我们准备更新的网络),对同一批经验数据
(s_t, a_t, z_t)重新计算概率。 -
对于离散动作:
π_d_new(a_t | s_t) -
对于连续动作:
π_c_new(z_t | s_t, a_t)
-
-
计算概率比:
-
离散动作概率比:
r_d(t) = π_d_new(a_t | s_t) / π_d_old(a_t | s_t) -
连续动作概率比:
r_c(t) = π_c_new(z_t | s_t, a_t) / π_c_old(z_t | s_t, a_t)
-
步骤 4: 计算PPO替代目标损失并更新Actor
PPO通过裁剪概率比来防止新策略与旧策略偏离过远。
-
计算未裁剪的目标:
-
离散Actor:
L_d_unclipped(t) = r_d(t) * A_t -
连续Actor:
L_c_unclipped(t) = r_c(t) * A_t
-
-
计算裁剪后的目标:
-
离散Actor:
L_d_clipped(t) = clip(r_d(t), 1-ε, 1+ε) * A_t -
连续Actor:
L_c_clipped(t) = clip(r_c(t), 1-ε, 1+ε) * A_t -
其中
ε是裁剪参数(如0.2),clip函数将概率比限制在[1-ε, 1+ε]范围内。
-
-
取最小值作为最终目标:
-
PPO最终取未裁剪和裁剪目标中的最小值,这是一个悲观的选择,确保我们不会因为某个动作的偶然高收益而进行危险的巨大更新。
-
离散Actor最终损失:
L_d_PPO(t) = min(L_d_unclipped(t), L_d_clipped(t)) -
连续Actor最终损失:
L_c_PPO(t) = min(L_c_unclipped(t), L_c_clipped(t))
-
-
加入熵正则项:
-
为了鼓励探索,在损失中加入策略的熵。
-
L_d_total = - (L_d_PPO + β * H(π_d_new))(目标是最大化,所以用负号) -
L_c_total = - (L_c_PPO + β * H(π_c_new)) -
其中
β是熵系数,H是熵。
-
-
梯度上升更新Actor:
-
对
L_d_total和L_c_total分别求导,并使用优化器(如Adam)更新离散Actor和连续Actor的网络参数。
-
步骤 5: 更新Critic
Critic的目标是更好地估计状态价值。
-
计算MSE损失:
-
L_critic = MSE(V(s_t), R_t) = (V(s_t) - R_t)^2 -
这个损失衡量了当前Critic的预测
V(s_t)与目标回报R_t之间的差距。
-
-
梯度下降更新Critic:
-
对
L_critic求导,并使用优化器更新Critic网络的参数。
-
2、投影头训练阶段:
- 使用训练好的策略(此时改为确定性策略,即离散选择最高概率节点,连续选均值),对训练集中的每个图进行完整的n步提示添加,得到最终的提示图G*
- 用这些G*和其标签,通过最小化损失函数来更新投影头,使其与新的图表示协调。
以上两个阶段构成一个训练周期,交替进行。
2.4 策略泛化技术
在少样本场景下,RL智能体容易对有限的训练环境过拟合。我们集成了 LEEP 方法进行策略泛化。
-
我们不再学习一个策略,而是学习
l个子策略{π_{d,1}, ..., π_{d,l}}和{π_{c,1}, ..., π_{c,l}}。 -
每个子策略在从总训练集 bootstrap 采样的不同子集上进行训练。
-
在更新每个子策略时,除了最大化PPO目标,还增加了一个正则化项,以最小化该子策略与联合策略
π_J之间的KL散度。-
离散联合策略
π_{d,J}取所有子策略对每个动作概率的最大值,然后归一化。 -
连续联合策略
π_{c,J}取所有子策略输出均值的平均值。
-
-
这鼓励了子策略的多样性,同时让它们共同趋近于一个更鲁棒的联合策略,有效防止过拟合。
2.5 量化提示影响的指标
为了衡量我们的方法是否实现了“必要且轻量”,我们设计了两个指标:
-
提示覆盖率:在整个提示添加过程中,至少被提示过一次的节点比例。PCR越低,说明提示越“精准”。
-
平均提示幅度:所有有效提示向量的L1范数之和,除以节点数和特征维度。APM越低,说明提示越“轻量”。
三、 实验验证
我们在图级和节点级的少样本任务上进行了广泛的实验。
3.1 少试图分类
-
设置:在8个分子性质预测数据集上,采用4种不同的预训练策略,在50样本场景下进行实验。
-
基线:对比了微调、多种特征提示方法 以及依赖预训练的All in One。
-
结果:
-
RELIEF在 32个任务中的28个 上取得了最佳性能,平均ROC-AUC超过微调 1.64%,超过第二名 1.04%。
-
RELIEF是唯一一个在所有任务上都稳定超越微调的方法。
-
我们的PCR和APM指标证实,RELIEF添加的提示覆盖节点最少、幅度最小,整体影响最低,真正做到了“必要且轻量”。
-
3.2 数据效率
-
实验:我们测试了每个方法需要多少比例的训练数据才能达到全量数据微调的性能。
-
结果:RELIEF所需数据量最少,而其他基线方法甚至在用上全部数据后,在某些任务上仍无法超越全量微调。这证明了RL范式通过逐步决策和评估,在极少数据下也能高效学习。
3.3 少样本节点分类
-
适配:通过为每个目标节点生成其k跳子图,将节点分类任务转化为子图分类任务。
-
结果:RELIEF在多个数据集和预训练策略下,在准确率和Macro F1分数上均达到第一或第二,平均超过微调约2%,证明了其通用性和有效性。
3.4 消融研究与分析
-
消融实验:分别将离散Actor替换为随机选择,将连续Actor替换为随机生成。结果表明,两者缺一不可,但离散Actor(节点选择策略)的作用更为关键。仅训练投影头的“线性探测”性能最差,证明了提示的必要性。
-
参数分析:RELIEF对超参数表现出良好的鲁棒性。
-
为什么RELIEF有效?
-
强大:RL善于解决组合优化问题,通过目标导向的奖励,能找到最优提示策略。
-
必要:智能体学会根据状态模式选择关键节点,错误的节点选择会导致累积奖励下降,从而被策略更新所修正。
-
轻量:过大尺度的提示会导致GNN产生不熟悉的表示,与投影头不协调,产生负奖励,从而迫使智能体生成小幅度的提示。
-
四、 总结与贡献
总结:本论文受大模型中提示边际效应启发,首次将强化学习引入图特征提示 tuning,提出了RELIEF方法。它通过序列决策,策略性地为图注入必要且轻量的特征提示,从而在少样本场景下显著提升了预训练GNN的下游任务性能。
我们的核心贡献:
-
新视角:提出了通过添加“必要且轻量”的特征提示来增强预训练GNN性能的新思路,并设计了PCR和APM指标进行量化。
-
新方法:首次将特征提示的添加过程形式化为序列决策问题,并提出了RELIEF这一基于混合动作空间RL的解决方案,集成了策略泛化技术以保证稳定性。
-
强验证:在图和节点级别的多种任务和预训练策略下,通过大量实验证明了RELIEF在分类性能和数据效率上均优于微调及其他先进的提示方法。
更多推荐

所有评论(0)