(arXiv 2026)GLiBRL :可学习基函数的深度贝叶斯元强化学习 ----待补充

文章目录

导读

论文标题:Generalised Linear Models in Deep Bayesian RL with Learnable Basis Functions

本文 由澳大利亚国立大学团队发表于 arXiv(2025 年 12 月首发,2026 年 5 月更新),是深度贝叶斯强化学习(Deep BRL) 领域的突破性工作,在元强化学习(Meta-RL)基准上取得了 SOTA 性能,同时提供了严谨的理论保证。

背景动机

经典 BRL 方法 假设转移和奖励模型的函数形式已知(如线性、二次型),只能适配结构简单的任务,泛化性极差,且计算复杂度随特征维度呈四次方增长,无法扩展到高维连续控制场景。

现有深度 BRL 方法(如 PEARL、VariBAD)用神经网络学习模型形式,但引入了无法解决的结构性问题:

  • 依赖变分推断(VI) :神经网络让任务参数与数据高度非线性耦合,精确边际似然不可解,只能优化证据下界(ELBO),带来高方差蒙特卡洛估计、摊销间隙(amortization gap)、后验崩溃(posterior collapse) 三大顽疾,最终学到的任务表示模糊、区分度差,直接损害策略性能。
  • 排列变异性 :多数方法用 RNN/Transformer 编码历史样本,样本顺序会影响输出,无法与样本高效的离策略(off-policy)算法(如 SAC) 兼容,回放缓冲区的乱序样本会导致性能暴跌。
  • 黑盒表示:任务嵌入是神经网络输出的黑盒,缺乏理论解释性,无法保证表示差异与任务差异的对应关系。

GLiBRL 核心突破是: 用 可学习的非线性基函数 + 参数空间广义线性模型 的架构,在保留深度模型表达能力的同时,实现完全可解的精确贝叶斯推断。

任何连续的非线性函数,都可以表示为 非线性基函数映射 + 线性组合 的形式。比如支持向量机(SVM)把低维数据映射到高维特征空间,在高维空间做线性分类;多项式回归 把原始特征 x 映射为 x , x 2 , x 3 , ...   x, x\^2, x\^3, \\dots x,x2,x3,...,再做线性回归;神经网络 每一层都是 非线性激活 + 线性变换 的堆叠。

做了一个关键的结构拆分:

  1. 用神经网络 ( ϕ T , ϕ R (\phi_T, \phi_R (ϕT,ϕR)作为可学习基函数 ,把原始状态、动作等高维输入映射到低维特征空间( C T , C R C_T, C_R CT,CR),捕捉环境动力学的非线性;

  2. 在特征空间中,任务参数与输出(下一状态、奖励)之间严格保持线性关系。

把所有非线性 "外包" 给基函数,让参数空间保持线性高斯结构,从而用共轭先验得到精确闭式后验。表达能力没有损失

这种设计看似限制了线性,实则通过非线性基函数保留了完整的表达能力,同时换来了三大核心收益:

  • 贝叶斯推断完全闭式可解,彻底抛弃变分推断和 ELBO;
  • 精确贝叶斯更新天然具备排列不变性,无缝兼容 on-policy(PPO)和 off-policy(SAC)算法;
  • 任务表示具备可解释的理论结构,其距离与样本核相似度严格等价。

方法框架

概率建模:共轭先验与似然

GLiBRL 对转移和奖励分别独立建模,采用Normal-Wishart 先验 + 矩阵正态似然的共轭组合,保证后验与先验同分布,存在闭式解。

(1)先验分布

任务参数分为转移参数 θ T = { T μ , T σ } \theta_T = \{T_\mu, T_\sigma\} θT={Tμ,Tσ} 和奖励参数 θ R = { R μ , R σ } \theta_R = \{R_\mu, R_\sigma\} θR={Rμ,Rσ},其中:

  • T μ / R μ T_\mu / R_\mu Tμ/Rμ 是线性映射的均值参数;
  • T σ / R σ T_\sigma / R_\sigma Tσ/Rσ 是模型噪声(协方差矩阵)。

先验采用 Normal-Wishart 分布: 均值参数 T μ / R μ T_\mu / R_\mu Tμ/Rμ 服从矩阵正态分布 M N \mathcal{MN} MN;精度矩阵(噪声的逆)服从 Wishart 分布 W \mathcal{W} W。

这一设计的关键是:对模型噪声也做贝叶斯推断,而不是像前人工作(如 ALPaCA)假设噪声已知,进一步降低了未见任务的预测误差。

(2)似然函数

下一个状态 S ′ S' S′ 和奖励 r 服从矩阵正态分布,其均值是 基函数特征 × 任务参数 的线性组合:

S ′ ∼ M N ( C T T μ , I N , T σ ) , r ∼ M N ( C R R μ , I N , R σ ) S' \sim \mathcal{MN}(C_T T_\mu, I_N, T_\sigma), \quad r \sim \mathcal{MN}(C_R R_\mu, I_N, R_\sigma) S′∼MN(CTTμ,IN,Tσ),r∼MN(CRRμ,IN,Rσ)

其中 C T = ϕ T ( S , A ) C_T = \phi_T(S,A) CT=ϕT(S,A)、 C R = ϕ R ( S , A , S ′ ) C_R = \phi_R(S,A,S') CR=ϕR(S,A,S′) 是基函数网络输出的特征矩阵。

精确后验推断与在线高效更新

由于先验与似然共轭,给定一批样本后,后验仍然是 Normal-Wishart 分布,所有参数都有闭式更新公式(式 15)。

在线更新效率优化 :逐样本在线更新时,直接计算矩阵求逆的复杂度是 O ( D T 3 ) O(D_T^3) O(DT3)。GLiBRL 利用矩阵求逆引理(Sherman-Morrison-Woodbury 公式) 做了化简:

  • 单样本更新时,求逆操作退化为标量倒数;
  • 配合缓存逆矩阵,最终在线更新复杂度降至 O ( max ⁡ ( D S 2 , D T 2 ) ) O(\max(D_S^2, D_T^2)) O(max(DS2,DT2))(转移)和 O ( D R 2 ) O(D_R^2) O(DR2)(奖励),完全满足实时交互的效率要求。

无 ELBO 的模型学习目标

因为推断完全可解,GLiBRL 可以直接计算精确的边际对数似然,彻底不需要 ELBO,从根源上避免了变分推断的所有问题。

模型的损失函数定义为: L m o d e l = − log ⁡ p ( C ) + λ T ∥ C T ∥ F 2 + λ R ∥ C R ∥ F 2 \mathcal{L}_{model} = -\log p(\mathcal{C}) + \lambda_T \|C_T\|_F^2 + \lambda_R \|C_R\|_F^2 Lmodel=−logp(C)+λT∥CT∥F2+λR∥CR∥F2

即负边际对数似然加 Frobenius 范数正则项,防止基函数过拟合,直接用梯度下降优化基函数网络 ϕ T , ϕ R \phi_T, \phi_R ϕT,ϕR 即可。

任务表示与策略融合

策略网络的输入是状态 + 任务信念表示。GLiBRL 用后验分布的确定性均值构造信念表示,并做了归一化处理:

f T , R ( M T ′ , M R ′ ) = tril ( M T ′ M T ′ T ) T ∥ tril ( M T ′ M T ′ T ) ∥ 2 M R ′ T ∥ M R ′ ∥ 2 T f_{T,R}(M_T', M_R') = \left \\frac{\\text{tril}(M_T' M_T'\^T)\^T}{\\\|\\text{tril}(M_T' M_T'\^T)\\\|_2} \\quad \\frac{M_R'\^T}{\\\|M_R'\\\|_2} \\right^T fT,R(MT′,MR′)=∥tril(MT′MT′T)∥2tril(MT′MT′T)T∥MR′∥2MR′TT

其中 tril ( ⋅ ) \text{tril}(\cdot) tril(⋅) 提取矩阵下三角并展平,降低表示维度。

理论贡献:表示距离与核相似度的等价性 ------------------ 首次为在线深度 BRL 建立了任务表示的结构性理论保证。

在温和假设下,GLiBRL 任务表示之间的 L 2 L_2 L2 距离,与任务样本在学习到的特征空间中的核化经验相似度存在严格的闭式等价关系:

  1. 未归一化表示 :其 L 2 L_2 L2 距离平方 = 带符号测度的最大均值差异(sMMD);
  2. 归一化表示 :其 L 2 L_2 L2 距离平方 = 4 − 2 × ( 转移sCMS + 奖励sCMS ) 4 - 2\times(\text{转移sCMS} + \text{奖励sCMS}) 4−2×(转移sCMS+奖励sCMS),其中 sCMS 是带符号测度的余弦均值相似度。

理论意义

  • 打破了深度元 RL 任务表示的黑盒性质:表示不是任意嵌入,其差异直接对应样本层面的特征差异;
  • 保证了「任务差异大 → 表示差异大」的一致性,为策略学习提供了可靠的输入;
  • 从理论上解释了为什么 GLiBRL 的任务区分能力远强于变分方法。

实验分析

  • 基准环境
    • MuJoCo 运动任务:HalfCheetahDir、AntDir、HalfCheetahVel;
    • MetaWorld 操作任务:ML10、ML45(元 RL 最具挑战性的基准之一)。
  • 对比基线:覆盖三类元 RL 方法 ------PPG 类(MAML)、黑盒类(RL²、AMAGO-v2、TrMRL、ECET)、任务推断类(PEARL、VariBAD、SDVT)。
  • 评估指标:零样本测试回报(MuJoCo)、零样本测试成功率(MetaWorld)。

主实验结果

  1. MuJoCo 运动任务

    • 仅用 100 万步训练(仅为其他基线的 10%),性能超过所有 PPO 基线最高 1.5 倍;
    • 相比同是 off-policy 的 SOTA 方法 PEARL,回报提升最高达1.8 倍,验证了精确推断相比变分近似的巨大优势。
  2. MetaWorld 操作任务

    • ML10:与 MAML、RL² 持平,显著超越其他所有深度 BRL 方法;
    • ML45(任务更多、难度更高):全面超越所有基线,成功率提升最高达1.1 倍,体现了更强的任务识别与泛化能力。

总结分析

GLiBRL 天然学习转移和奖励模型,非常适合结合基于模型的规划,但从高维 Wishart 分布采样是计算瓶颈,有待优化; 不确定性的利用 :目前策略仅输入后验均值,尚未利用协方差中的不确定性信息,但直接输入协方差会导致训练不稳定,是未来的探索方向; 更复杂的观测空间:当前实验以向量状态为主,向图像等高维观测扩展是潜在的应用方向。

额外补充

  1. 贝叶斯强化学习(BRL)

贝叶斯强化学习是元强化学习的一个子类,核心思想是:

假设环境的转移函数、奖励函数由未知的任务参数 ( θ T (\theta_T (θT 控制转移, θ R \theta_R θR 控制奖励)决定;

智能体在交互中不断对任务参数做贝叶斯推断,形成 信念(belief),并基于信念决策; 天然具备不确定性量化、快速适配新任务、平衡探索 - 利用的优势。

这类问题可以形式化为贝叶斯自适应 MDP(BAMDP):将信念融入状态空间形成 超状态 ,最终目标是学习一个以 原始状态 + 任务信念 为输入的策略。

  1. 为什么深度 BRL 必须依赖变分推断?

贝叶斯强化学习的核心目标是:给定一个任务的交互样本集合 C = { c t } t = 1 N \mathcal{C} = \{c_t\}{t=1}^N C={ct}t=1N(其中 c t = ( s t , a t , s t + 1 , r t + 1 ) c_t=(s_t,a_t,s{t+1},r_{t+1}) ct=(st,at,st+1,rt+1)),最大化数据的边际对数似然

log ⁡ p ϕ ( C ) = log ⁡ ∫ p ( θ ) ⋅ p ϕ ( C ∣ θ )   d θ \log p_{\phi}(\mathcal{C}) = \log \int p(\theta) \cdot p_{\phi}(\mathcal{C} \mid \theta) \, d\theta logpϕ(C)=log∫p(θ)⋅pϕ(C∣θ)dθ 其中:

  • θ = { θ T , θ R } \theta = \{\theta_T, \theta_R\} θ={θT,θR} 是未知的任务参数(分别控制转移动力学和奖励函数);
  • p ( θ ) p(\theta) p(θ) 是任务参数的先验分布;
  • p ϕ ( C ∣ θ ) p_{\phi}(\mathcal{C} \mid \theta) pϕ(C∣θ) 是似然函数,由参数为 ϕ \phi ϕ 的神经网络建模(比如用神经网络输出高斯分布的均值和方差)。

不可解的根源 :在深度 BRL 方法中,似然 p ϕ ( s ′ , r ∣ s , a , θ ) p_{\phi}(s',r \mid s,a,\theta) pϕ(s′,r∣s,a,θ) 的均值由神经网络输出,神经网络的非线性激活(ReLU、Tanh 等)让 θ \theta θ 与原始数据 ( s , a , s ′ , r ) (s,a,s',r) (s,a,s′,r) 形成高度非线性的耦合关系。这导致上述积分没有解析闭式解,无法直接对 ϕ \phi ϕ 求导最大化边际似然。

因此,PEARL、VariBAD 等方法只能引入变分推断 :用一个由推理网络参数化的近似后验 q ψ ( θ ∣ C ) q_\psi(\theta \mid \mathcal{C}) qψ(θ∣C) 去拟合真实后验 p ( θ ∣ C ) p(\theta \mid \mathcal{C}) p(θ∣C),转而优化证据下界(ELBO)

log ⁡ p ϕ ( C ) ≥ E q ψ ( θ ∣ C ) log ⁡ p ϕ ( C ∣ θ ) − D K L ( q ψ ( θ ∣ C ) ∥ p ( θ ) ) ⏟ ELBO \log p_{\phi}(\mathcal{C}) \geq \underbrace{\mathbb{E}{q\psi(\theta \mid \mathcal{C})}\left \\log p_{\\phi}(\\mathcal{C} \\mid \\theta) \\right - D_{KL}\left( q_\psi(\theta \mid \mathcal{C}) \parallel p(\theta) \right)}_{\text{ELBO}} logpϕ(C)≥ELBO Eqψ(θ∣C)logpϕ(C∣θ)−DKL(qψ(θ∣C)∥p(θ)) ELBO 是真实边际似然的一个下界,最大化 ELBO 可以间接最大化边际似然,但这一近似会引入三类难以避免的问题。

  • 高方差蒙特卡洛估计

ELBO 中的期望项 E q log ⁡ p ( C ∣ θ ) \mathbb{E}_{q}\\log p(\\mathcal{C} \\mid \\theta) Eqlogp(C∣θ) 通常依然没有解析解,必须通过蒙特卡洛采样近似计算:

E q ψ ( θ ∣ C ) log ⁡ p ϕ ( C ∣ θ ) ≈ 1 K ∑ k = 1 K log ⁡ p ϕ ( C ∣ θ k ) , θ k ∼ q ψ ( θ ∣ C ) \mathbb{E}{q\psi(\theta \mid \mathcal{C})}\left \\log p_{\\phi}(\\mathcal{C} \\mid \\theta) \\right \approx \frac{1}{K} \sum_{k=1}^K \log p_{\phi}(\mathcal{C} \mid \theta_k), \quad \theta_k \sim q_\psi(\theta \mid \mathcal{C}) Eqψ(θ∣C)logpϕ(C∣θ)≈K1k=1∑Klogpϕ(C∣θk),θk∼qψ(θ∣C) 其中 K 是采样数量。

为了计算效率,绝大多数深度 BRL 方法采用单样本估计 ( K = 1 K=1 K=1),这会导致梯度估计的方差极高。

  • 摊销间隙(Amortization Gap)

深度元学习为了效率,采用摊销推理 机制:训练一个共享的推理网络(如 RNN 编码器),输入任意任务的样本集合,直接输出近似后验的参数(均值、方差)。所有任务复用同一个推理网络,无需为每个新任务单独优化后验。

但共享网络的表达容量是有限的,不可能对所有任务都达到最优的后验近似。摊销间隙 就是 单个任务单独优化得到的最优 ELBO摊销推理网络输出的 ELBO 之间的固有差距:

Amortization Gap = max ⁡ q   ELBO ( q ) − ELBO ( q ψ ( ⋅ ∣ C ) ) \text{Amortization Gap} = \max_{q} \, \text{ELBO}(q) - \text{ELBO}(q_\psi(\cdot \mid \mathcal{C})) Amortization Gap=maxqELBO(q)−ELBO(qψ(⋅∣C))

共享的推理网络为了兼顾所有任务,只能学到一个 平均水平 的后验近似;对每个具体任务而言,近似后验都存在系统性偏差,无法精准刻画任务的真实动力学参数;

  • 后验崩溃(Posterior Collapse)

后验崩溃是变分类架构最严重的退化现象:近似后验完全退化为先验分布 ,即: D K L ( q ψ ( θ ∣ C ) ∥ p ( θ ) ) ≈ 0 D_{KL}\left( q_\psi(\theta \mid \mathcal{C}) \parallel p(\theta) \right) \approx 0 DKL(qψ(θ∣C)∥p(θ))≈0

此时任务参数 θ \theta θ 完全不携带样本 C \mathcal{C} C 的信息,相当于模型「忽略」了交互数据,所有任务的表示都收敛到同一个先验分布,彻底失去区分度。

生成模型(转移 / 奖励网络 p ϕ p_\phi pϕ)的表达能力过强,即使不依赖任务参数 θ \theta θ,也能很好地重构样本数据。此时 ELBO 的优化会倾向于让 KL 惩罚项趋近于 0,而重构项几乎不受影响,最终导致 θ \theta θ 被模型「废弃」。

  1. 排列变异性:Off-Policy 算法的兼容性

对于一个以集合为输入的函数 f ( ⋅ ) f(\cdot) f(⋅),若对样本集合 C = { c 1 , c 2 , ... , c N } \mathcal{C} = \{c_1, c_2, \dots, c_N\} C={c1,c2,...,cN} 的任意排列 σ \sigma σ,都满足:

f ( { c σ ( 1 ) , c σ ( 2 ) , ... , c σ ( N ) } ) = f ( { c 1 , c 2 , ... , c N } ) f\left( \{c_{\sigma(1)}, c_{\sigma(2)}, \dots, c_{\sigma(N)}\} \right) = f\left( \{c_1, c_2, \dots, c_N\} \right) f({cσ(1),cσ(2),...,cσ(N)})=f({c1,c2,...,cN})

则称 f 是排列不变的。核心性质:输出只和样本的内容有关,和样本的输入顺序完全无关。

  • RNN/LSTM 系列:按时间步顺序逐个输入样本,每一步的隐藏状态严格依赖上一步的输出。样本顺序改变,隐藏状态的计算路径就完全不同,最终输出的任务表示也不同。
  • Transformer 系列 :自注意力运算本身是排列不变的,但为了利用时序信息,标准实现都会加入位置编码,这就引入了顺序依赖。打乱样本顺序会改变位置编码与样本的对应关系,最终输出发生变化。

Off-Policy 强化学习的核心机制是回放缓冲区(Replay Buffer)

交互过程中收集的所有样本都会存入缓冲区;训练时,从缓冲区中随机采样一批样本用于更新网络,采样顺序是乱序的,和真实交互的时序顺序无关。

如果任务推理网络是排列可变的,就会出现致命问题:

同一组样本,只要输入顺序不同,输出的任务信念 b 就不同。

这会导致策略 / 价值网络的输入(状态 + 任务信念)极不稳定:同样的状态和样本集合,每次训练得到的信念都不一样,Q 值估计震荡,策略更新方向混乱,最终性能暴跌。

相关推荐
喜欢吃豆1 小时前
GEO 到底怎么做?从 Prompt 研究到 AI Citation 的完整落地方法
人工智能·prompt
yingyuecom1 小时前
映悦AI × Joverse全球发布:四大重磅更新,AI创作进入工业化时代
人工智能·gpt·chatgpt·prompt·aigc
sunshine22 girl1 小时前
Angular7,9,学习笔记一 创建项目,基本语法
笔记·学习
会周易的程序员1 小时前
软件接入大模型实现 Agent —— 从原理到 C++ 落地完全指南
c++·人工智能·物联网·架构·agent·工业协议·mcp
小雪崩1 小时前
嵌入式学习 day30:minilog
linux·c语言·学习
躺柒1 小时前
读数据可视化13空间标量场(上)
人工智能·深度学习·信息可视化·数据可视化·空间·大数据分析
小王2041 小时前
Day 28:目标检测入门 — 两阶段 vs 单阶段
人工智能·目标检测·计算机视觉
CIO_Alliance1 小时前
AI深度系列(3)| 从RNN到LSTM:序列数据处理的技术逻辑与企业AI化转型启示
人工智能·rnn·深度学习·神经网络·lstm·企业cio联盟·企业级ai化转型
像风一样自由20201 小时前
11.PostgreSQ、-MySQL与MongoDB-AI应用如何选择数据库
数据库·人工智能·mysql·mongodb·大模型·rag·智能体