Page 252 - 《软件学报》2026年第6期
P. 252

蔡瑞初 等: 隐变量因果模型视角下的策略梯度方差优化                                                      2571


                 果编码网络或因果解码网络无关, 最大化             ELBO  调整为:

                            max(ELBO) = min(−ELBO)
                                                                                            
                                           T                        T ∑
                                          ∑                                                 
                                               (               )                            
                                                                                            
                                     = min    D KL q φ (h t |s t ,a t , s t+1 ) ∥ p(h t ) −  E h t ∼q φ (h t |s t ,a t ,s t+1 ) log p ψ (s t+1 |s t ,a t ,h t )  
                                                                                            
                                           t=0                     t=0
                                       T ∑
                                            [   (               )                         ]
                                     =   min D KL q φ (h t |s t ,a t , s t+1 ) ∥ p(h t ) −E h t ∼q φ (h t |s t ,a t ,s t+1 ) log p ψ (s t+1 |s t ,a t ,h t )  (13)
                                       t=0
                    最大化   ELBO  分成各个时间点内单独控制每一步损失, 则因果模型的损失为:

                                            (               )
                                   L causal = D KL q φ (h t |s t ,a t , s t+1 ) ∥ p(h t ) −E h t ∼q φ (h t |s t ,a t ,s t+1 ) log p ψ (s t+1 |s t ,a t ,h t )  (14)
                 其中, 等式右侧第     1  项是近似后验分布      (由因果编码网络推断) 和假设先验分布            (一般假设为高斯分布) 的        KL  散
                 度. 最小化   KL  散度项可以促使因果编码网络学习的隐变量分布尽可能接近先验分布, 同时起到正则化优化问
                 题的作用, 保持参数推断的不确定性. 等式右侧第               2  项是在近似后验分布下对隐藏变量           h t  进行采样, 并计算因果
                 编码网络中    p ψ (s t+1 |s t ,a t ,h t ) 的对数概率负期望值. 最小化对数概率负期望项, 可以使因果编码网络和因果解码网
                 络协同学习, 在给定隐变量的情况下, 使重构数据尽可能接近真实的可观测状态信息, 保持模型对数据的解释能
                 力. 通过平衡这两部分, 保证了变分推断和因果模型拟合的有效性, 进而确保算法可推断因果价值函数所需要的
                 隐变量.
                    假设隐变量的先验分布服从各向同性多元高斯分布, 即                   p(h t ) = N (h t ;0,I). 同时, 也假设真实后验分布也满足
                                                                               (  (j)  (j)  (j)  )
                 对角协方差矩阵的多元高斯分布. 通过向因果编码网络输入采样批次样本对应的                           s t ,a t , s   信息, 获得对应情况
                                                                                      t+1
                                         (j)         (j)               J  的采样批次样本, 因果编码网络的输出通
                 的后验分布的参数, 如均值为         µ t  、标准差为  σ t  . 相应地, 基于数量为
                 过重参数化技巧整理为近似后验分布:

                                                               (           )
                                               (           )         [   ] 2
                                                                 (j)
                                                                    (j)
                                                   (j)
                                                 (j)
                                                      (j)
                                              q φ h t |s t ,a t , s (j)  = N h t ;µ t , σ t (j)  I   (15)
                                                         t+1
                    在公式   (15) 中, 重参数化技巧将隐变量       h t ( j)  的采样过程重新表述为一个确定性函数和一个独立的噪声变量的
                          (j)
                      ( j)
                               (j)
                 函数:  h t = µ t +σ t ⊙ε t (j) . 其中,  ε t (j)   是从标准正态分布  N (0,1) 中采样的噪声变量. 故因果模型的损失中, 近似后验
                      (  (j)  ( j)  (j)  (j)  )
                 分布  q φ h t |s t ,a t , s   和假设先验分布  p(h t ) 的  KL  散度项整理为:
                                t+1
                                    (  (          )     )   1  J ∑ (  { [  ] 2  }  [  ] 2 [  ] 2 )
                                              (j)
                                        (j)
                                           (j)
                                 D KL q φ h t |s t ,a t , s (j)  ∥ p(h t ) = −  1+log σ t (j)  − µ t (j)  − σ t (j)  (16)
                                                t+1
                                                            2
                                                             j=1
                    而对于因果解码网络, 直接优化对数概率负期望值在计算上比较复杂. 在已知可观测状态信息维度为                                d 时, 借
                 助真实后验分布      (已假设为高斯分布) 的概率密度函数, 可以如公式              (17) 将对数概率负期望项简化为重构损失项:

                                                                           (                  ) 
                                                                   1            
            
 
                                                                             1 
            
 2  
                        −E h t ∼q φ (h t |s t ,a t ,s t+1 ) log p ψ (s t+1 |s t ,a t ,h t ) = −E h t ∼q φ (h t |s t ,a t ,s t+1 ) log   √  exp −  
s t+1 − F ψ (s t ,a t ,h t )
  
                                                                    d        2σ 2              
                                                                  (2π) σ 2d
                                                            [                                 ]
                                                             d               1 
 
           
 
 2
                                                 = E h t ∼q φ (h t |s t ,a t ,s t+1 )  log(2π)+dlogσ+  
s t+1 − F ψ (s t ,a t ,h t )
                                                             2              2σ 2
                                                   1  J ∑
 
 (j)  (  (j) 
 2
                                                                      )
                                                                  (j)
                                                               (j)

                                                 ∝     
s t+1  − F ψ s t ,a t ,h t 
 
               (17)
                                                   J
                                                     j=1
                    学习隐变量因果模型的过程, 如算法            1  所示.
                 算法  1. 隐变量因果模型学习算法.
                           β, 采样批次样本量      J, 迭代次数                              F ψ , 因果模型的  KL  散度项权重
                 输入: 缓存池                            I, 因果编码网络    F φ , 因果解码网络
                 w KL , 因果模型的重构损失项权重      w recon ;
                 输出: 隐变量因果模型更新后的因果编码网络参数                φ, 因果解码网络参数      ψ, 隐变量信息   h.
   247   248   249   250   251   252   253   254   255   256   257