Skip to content

第 04 章:强化学习与条件引导生成:从 PPO 到 GFlowNets

“传统强化学习(RL)的目标是‘竭尽全力找到唯一的冠军’;
而生成流网络(GFlowNets)与条件扩散模型的哲学是‘按照每个候选的才华概率,百花齐放地采样整个生态’。”


1. 引言:离散马尔可夫决策过程的呼唤

在第 02 章与第 03 章中,我们深入研讨了基于变分自编码器(VAE)的连续潜在流形优化。这种“先连续化、再最优化、最后离散还原”的间接范式在许多场景下极为优雅。

然而,在面对具有严苛离散组合规则的材料化学设计时,直接在离散动作空间中构建生成智能体往往展现出更为直截了当的威力:

  • 我们可以像化学家在实验室中一样,一步一步组装分子:“先放一个吡啶环,在第 4 位连一个酰胺键,再接一个氟代苯”
  • 目标性能可以完全是一个不可微的黑盒函数(例如湿法合成产率、分子对接打分、经验毒性过滤器);
  • 我们不仅需要最优解,更需要找到几百个结构迥异、机理不同的多样化骨架,以抵御后期实验失败的高昂风险。

为了达成这一壮举,计算科学界先后发展出了基于强化学习的离散策略优化Bengio 团队开创的生成流网络(GFlowNets)以及连续扩散模型的条件得分引导(Score Guidance)。本章将系统解构这一前沿武器库。


2. 经典强化学习(RL)在微观分子生成中的建模与困境

2.1 马尔可夫决策过程(MDP)的形式化构建

我们将分子或晶体碎片的逐原子组装,建模为一个有限时域的马尔可夫决策过程(Markov Decision Process, MDP),由五元组 (S,A,P,R,γ)(\mathcal{S}, \mathcal{A}, \mathcal{P}, \mathcal{R}, \gamma) 唯一定义:

  • 状态空间 S\mathcal{S}:状态 sts_t 代表当前已部分构建完成的不完整分子拓扑图(s0=s_0 = \emptyset 为空图起点);
  • 动作空间 A\mathcal{A}:动作 atA(st)a_t \in \mathcal{A}(s_t) 代表合法的化学拼接操作(例如:添加特定原子、闭环形成芳香环、添加官能团,或发出终止信号 [Stop]);
  • 转移概率 P\mathcal{P}:确定性图更新规则 st+1=AddStep(st,at)s_{t+1} = \text{AddStep}(s_t, a_t)
  • 奖励函数 R(sT)\mathcal{R}(s_T):仅在轨迹终止到达完整结构 x=sTx = s_T 时,调用外部分子性质预测器或物理模型赋予最终标量奖励 R(x)0R(x) \ge 0(中间过程奖励通常为 0)。

2.2 策略梯度(Policy Gradient)与 PPO

在强化学习框架下,智能体策略由参数化网络 πθ(as)\pi_\theta(a \mid s) 表示。目标是最大化整条生成轨迹的期望累积奖励:

J(θ)=Eτπθ[R(xτ)]\mathcal{J}(\theta) = \mathbb{E}_{\tau \sim \pi_\theta}\left[ R(x_\tau) \right]

根据策略梯度定理(Policy Gradient Theorem),损失函数的解析梯度为:

θJ(θ)=Eτπθ[t=0T1θlogπθ(atst)A^(st,at)]\nabla_\theta \mathcal{J}(\theta) = \mathbb{E}_{\tau \sim \pi_\theta}\left[ \sum_{t=0}^{T-1} \nabla_\theta \log \pi_\theta(a_t \mid s_t) \cdot \hat{A}(s_t, a_t) \right]

其中 A^(st,at)\hat{A}(s_t, a_t) 为优势函数(Advantage Function)。为了解决策略更新步长过大导致的训练崩溃,**近端策略优化(PPO)**引入了重要性采样与裁剪目标函数(Clipped Surrogate Objective):

LCLIP(θ)=E^t[min(rt(θ)A^t,clip(rt(θ),1ϵ,1+ϵ)A^t)],rt(θ)=πθ(atst)πθold(atst)\mathcal{L}^{\text{CLIP}}(\theta) = \hat{\mathbb{E}}_t \left[ \min\left( r_t(\theta)\hat{A}_t, \, \text{clip}(r_t(\theta), 1-\epsilon, 1+\epsilon)\hat{A}_t \right) \right], \quad r_t(\theta) = \frac{\pi_\theta(a_t \mid s_t)}{\pi_{\theta_{\text{old}}}(a_t \mid s_t)}


2.3 传统强化学习在物质设计中的阿喀琉斯之踵:模式坍塌(Mode Collapse)

尽管 PPO 和 REINFORCE 在优化性质指标上极其高效,但在真实的新药与新材料研发中,实验学家很快发现了其致命缺陷:

  1. 追求极值的狂热(Mode Collapse): 强化学习的数学目标是 maxθE[R(x)]\max_\theta \mathbb{E}[R(x)]。在经过充分训练后,智能体会不可避免地发现某一个单一的高分母核,并随后以接近 100% 的概率反复生成这一种骨架的微小变体,彻底丧失探索新骨架的动力;
  2. 高维多峰奖励地形的盲区: 在真实的化学空间中,具有极高结合活性或优良带隙的结构往往分散在几个完全不同的化学大类中(Multimodal Landscape)。传统 RL 算法极难跨越不同类目之间的深邃低分荒漠,容易被锁定在第一个发现的局部极小峰上。

3. 生成流网络(GFlowNets):采样多样性物质的数学革命

为了从根本上治愈强化学习的模式坍塌,图灵奖得主 Yoshua Bengio 团队于 2021 年提出了震惊 AI 界的理论突破——生成流网络(Generative Flow Networks, GFlowNets)

3.1 核心范式转向:从“赢家通吃”到“按概率多样化采样”

GFlowNet 放弃了传统的“极大化期望奖励”,而是将生成目标设定为:训练一个采样策略,使其生成任意完整分子 xx 的边际概率 P(x)P(x),严格正比于该分子的非负奖励值 R(x)R(x)

P(x)=R(x)xXR(x)=R(x)ZP(x) = \frac{R(x)}{\sum_{x' \in \mathcal{X}} R(x')} = \frac{R(x)}{\mathcal{Z}}

  • 如果某种构型的奖励极高(R=100R=100),它被采样的概率就很大;
  • 如果另一种构型的奖励适中(R=20R=20),它依然有相当的概率被采样到;
  • 如果某种结构完全不合逻辑(R=0R=0),它被采样的概率严格归零!

这意味着,GFlowNet 天生能够在多个截然不同的化学骨架高峰之间按比例同时分配采样资源,完美满足了真实研发中对“高活性且结构多样化”的极端渴求。

图 4-1:生成流网络(GFlowNets)在有向无环图上的分子离散序列构建与流平衡机制示意图

图 4-2:强化学习模式坍塌(单一峰收敛)与 GFlowNet 按奖励概率多峰采样(多样性骨架探索)对比实验


3.2 网络流理论与流守恒定律(Flow Conservation)

GFlowNet 的数学本质是有向无环图(Directed Acyclic Graph, DAG)上的物理水流网络。 设状态转移构成一个庞大的图网络:

  • 定义源点(Source)为初始空状态 s0s_0
  • 定义汇点(Sinks)为所有终止的完整合法分子状态 xXx \in \mathcal{X}
  • 引入边流量 F(ss)0F(s \to s') \ge 0

依据流体力学基尔霍夫电流定律与流守恒公理:流经任意中间状态 ss 的总入流,必须严格等于从该状态流出的总出流

sparentParents(s)F(sparents)=schildChildren(s)F(sschild)\sum_{s_{\text{parent}} \in \text{Parents}(s)} F(s_{\text{parent}} \to s) = \sum_{s_{\text{child}} \in \text{Children}(s)} F(s \to s_{\text{child}})

在汇点 xx 处,流出的最终流量必须严格等于物质的奖励值:F(x)=R(x)F(x) = R(x);在源点 s0s_0 处,注入的总流量恰好等于整个化学空间的配分函数 Z=F(s0)=xR(x)\mathcal{Z} = F(s_0) = \sum_{x} R(x)


3.3 轨迹平衡损失(Trajectory Balance Loss, TB)

Bengio 等人推导证明,流守恒条件可以极其优雅地转化为参数化网络在单条完整轨迹 τ=(s0s1sT=x)\tau = (s_0 \to s_1 \to \dots \to s_T = x) 上的代数平衡关系。

定义:

  • 前向策略(Forward Policy)PF(st+1st;θ)P_F(s_{t+1} \mid s_t; \theta),代表从前向后构建分子的选择概率;
  • 后向策略(Backward Policy)PB(stst+1;θ)P_B(s_t \mid s_{t+1}; \theta),代表从完整结构逐步拆解的逆向概率;
  • 参数化总流量常数ZθZ_\theta(作为一个可学习的标量网络参数)。

对于网络流守恒的任意轨迹,以下代数恒等式必须严格成立:

Zθt=0T1PF(st+1st;θ)=R(x)t=0T1PB(stst+1;θ)Z_\theta \prod_{t=0}^{T-1} P_F(s_{t+1} \mid s_t; \theta) = R(x) \prod_{t=0}^{T-1} P_B(s_t \mid s_{t+1}; \theta)

由此,我们可以直接写出参数优化的轨迹平衡损失函数(Trajectory Balance Objective)

LTB(τ;θ)=(logZθt=0T1PF(st+1st;θ)R(x)t=0T1PB(stst+1;θ))2\mathcal{L}_{\text{TB}}(\tau; \theta) = \left( \log \frac{Z_\theta \prod_{t=0}^{T-1} P_F(s_{t+1} \mid s_t; \theta)}{R(x) \prod_{t=0}^{T-1} P_B(s_t \mid s_{t+1}; \theta)} \right)^2

通过对采样的轨迹极小化 LTB\mathcal{L}_{\text{TB}},模型便在无需显式计算全空间超大配分函数 ZZ 的情况下,自动迫使前向策略 PFP_F 收敛于精确正比于奖励的无偏采样器!


4. 连续几何扩散模型的条件引导前沿

在处理具有连续三维欧氏坐标与晶格矩阵的复杂固态晶体材料时,连续扩散模型(Score-based Diffusion)通过引入**条件引导(Conditioning Guidance)**实现了物理维度的可控生成。

4.1 分类器引导(Classifier Guidance)

根据贝叶斯定理,条件得分函数(Score Function)可精确展开为无条件先验得分与性质似然梯度的线性叠加:

xtlogpt(xty)=xtlogpt(xt)无条件生成得分网络 sθ(xt,t)+γxtlogpt(yxt)外部性质分类器/回归器提供的物理梯度\nabla_{\mathbf{x}_t} \log p_t(\mathbf{x}_t \mid \mathbf{y}) = \underbrace{\nabla_{\mathbf{x}_t} \log p_t(\mathbf{x}_t)}_{\text{无条件生成得分网络 } s_\theta(\mathbf{x}_t, t)} + \underbrace{\gamma \cdot \nabla_{\mathbf{x}_t} \log p_t(\mathbf{y} \mid \mathbf{x}_t)}_{\text{外部性质分类器/回归器提供的物理梯度}}

在逆向时间扩散去噪的每一步中,外部预训练的性质预测器(如带隙回归器或晶体稳定性分类器)对当前带噪结构 xt\mathbf{x}_t 计算梯度,像一双无形的手,在降噪过程中持续将原子坐标与晶胞矢量拽向目标性能区域。

4.2 无分类器引导(Classifier-Free Guidance, CFG)

为了避免训练带噪分类器的高昂成本,CFG 在同一个网络中以一定概率随机丢弃条件输入 y\mathbf{y}(用空标记 \emptyset 代替),在去噪时将外推方向定义为:

ϵ~θ(xt,t,y)=ϵθ(xt,t,)+w(ϵθ(xt,t,y)ϵθ(xt,t,))\tilde{\epsilon}_\theta(\mathbf{x}_t, t, \mathbf{y}) = \epsilon_\theta(\mathbf{x}_t, t, \emptyset) + w \cdot \left( \epsilon_\theta(\mathbf{x}_t, t, \mathbf{y}) - \epsilon_\theta(\mathbf{x}_t, t, \emptyset) \right)

通过设置引导权重 w>1w > 1,模型可以显著提高生成样本对目标物理性能的对齐度(Alignment Fidelity)。


5. 工业级代码实战:从零构建生成流网络(GFlowNet)多峰采样器

下面我们使用 PyTorch 搭建一个最小可运行的**轨迹平衡生成流网络(GFlowNet-TB)**原型,在一个具有多个孤立高分峰谷的离散组合图空间中,验证其如何精准按照奖励比例同时捕获全部高价值模式:

python
"""
文件名: gflownet_trajectory_balance.py
功能: 实现基于轨迹平衡损失 (Trajectory Balance Loss) 的生成流网络原型,
      验证其在离散多峰奖励地形下的无偏多样性采样能力。
"""

import math
import torch
import torch.nn as nn
import torch.nn.functional as F


class DiscreteActionGFlowNet(nn.Module):
    """
    轻量级生成流网络 (GFlowNet-TB) 智能体
    状态表示为一维二进制构型向量 (模拟离散化学片段的添加)
    """
    def __init__(self, state_dim: int = 8, num_actions: int = 8):
        super().__init__()
        self.state_dim = state_dim
        self.num_actions = num_actions  # 每个动作代表激活某一特定化学位点

        # 1. 前向策略网络 P_F(a | s)
        self.forward_policy = nn.Sequential(
            nn.Linear(state_dim, 64),
            nn.SiLU(),
            nn.Linear(64, 64),
            nn.SiLU(),
            nn.Linear(64, num_actions + 1)  # 包含停止动作 [Stop Action]
        )

        # 2. 可学习的全局配分对数常数 log(Z)
        self.log_Z = nn.Parameter(torch.tensor(0.0))

    def get_forward_logits(self, state: torch.Tensor) -> torch.Tensor:
        return self.forward_policy(state)


def multi_modal_chemical_reward(state: torch.Tensor) -> torch.Tensor:
    """
    构造多峰物理化学奖励函数 (模拟化学空间中存在的两个不同高分母核峰)
    Mode A: [1, 1, 0, 0, 0, 0, 0, 0] -> 极高活性
    Mode B: [0, 0, 0, 0, 1, 1, 1, 1] -> 极高活性
    其他散乱组合具有低奖励或惩罚
    """
    mode_A = torch.tensor([1.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0], device=state.device)
    mode_B = torch.tensor([0.0, 0.0, 0.0, 0.0, 1.0, 1.0, 1.0, 1.0], device=state.device)

    dist_A = torch.norm(state - mode_A, p=1, dim=-1)
    dist_B = torch.norm(state - mode_B, p=1, dim=-1)

    reward_A = 100.0 * torch.exp(-dist_A)
    reward_B = 80.0 * torch.exp(-dist_B)

    # 复合奖励且加上底噪保护
    total_reward = torch.maximum(reward_A, reward_B) + 1e-4
    return total_reward


class GFlowNetTrainer:
    """
    GFlowNet 轨迹采样与轨迹平衡 (Trajectory Balance) 训练循环
    """
    def __init__(self, model: DiscreteActionGFlowNet, lr: float = 1e-3, lr_z: float = 1e-2):
        self.model = model
        self.optimizer = torch.optim.Adam([
            {'params': model.forward_policy.parameters(), 'lr': lr},
            {'params': [model.log_Z], 'lr': lr_z}
        ])

    def sample_trajectories(self, batch_size: int = 32):
        """采样生成完整轨迹并记录前向对数概率总和"""
        device = next(self.model.parameters()).device
        state = torch.zeros(batch_size, self.model.state_dim, device=device)

        sum_log_pf = torch.zeros(batch_size, device=device)
        active_mask = torch.ones(batch_size, dtype=torch.bool, device=device)

        max_steps = self.model.state_dim + 1

        for step in range(max_steps):
            if not active_mask.any():
                break

            logits = self.model.get_forward_logits(state)
            # 掩蔽已经添加过的动作位点,防止重复添加
            masked_logits = logits.clone()
            masked_logits[:, :self.model.state_dim] = torch.where(
                state == 1.0,
                torch.full_like(masked_logits[:, :self.model.state_dim], -1e9),
                masked_logits[:, :self.model.state_dim]
            )

            dist = torch.distributions.Categorical(logits=masked_logits)
            actions = dist.sample()
            log_prob = dist.log_prob(actions)

            # 仅对仍在活跃的样本累计对数概率
            sum_log_pf += torch.where(active_mask, log_prob, torch.zeros_like(sum_log_pf))

            # 判定停止动作 (最后一个动作为终止)
            stop_action_idx = self.model.num_actions
            is_stop = (actions == stop_action_idx)
            active_mask = active_mask & (~is_stop)

            # 更新状态 (非停止动作激活相应位点)
            for b in range(batch_size):
                if active_mask[b] and actions[b] < self.model.state_dim:
                    state[b, actions[b]] = 1.0

        rewards = multi_modal_chemical_reward(state)
        return state, rewards, sum_log_pf

    def train_step(self, batch_size: int = 64) -> dict:
        """执行单步轨迹平衡 (Trajectory Balance) 优化"""
        self.optimizer.zero_grad()
        states, rewards, sum_log_pf = self.sample_trajectories(batch_size)

        # 假设均匀后向转移概率 P_B (统一基线)
        sum_log_pb = 0.0

        # Trajectory Balance Loss: (log Z + sum log P_F - log R - sum log P_B)^2
        loss = (self.model.log_Z + sum_log_pf - torch.log(rewards) - sum_log_pb).pow(2).mean()

        loss.backward()
        torch.nn.utils.clip_grad_norm_(self.model.parameters(), max_norm=10.0)
        self.optimizer.step()

        return {
            "tb_loss": loss.item(),
            "mean_reward": rewards.mean().item(),
            "estimated_log_Z": self.model.log_Z.item()
        }


# =====================================================================
# 单元验证模块: 验证 GFlowNet 多峰捕获与多样性生成能力
# =====================================================================
if __name__ == "__main__":
    torch.manual_seed(42)
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

    gfn = DiscreteActionGFlowNet(state_dim=8, num_actions=8).to(device)
    trainer = GFlowNetTrainer(gfn, lr=2e-3, lr_z=5e-2)

    print(">>> [STEP 1] 执行 GFlowNet 轨迹平衡训练 (迭代 100 步)...")
    for epoch in range(1, 101):
        metrics = trainer.train_step(batch_size=64)
        if epoch % 25 == 0:
            print(f"    Epoch {epoch:03d} | TB Loss: {metrics['tb_loss']:.4f} | 平均奖励: {metrics['mean_reward']:.2f} | 估算 log(Z): {metrics['estimated_log_Z']:.2f}")

    print(">>> [STEP 2] 检验训练后模型的采样多样性 (是否能同时采样到 Mode A 与 Mode B)...")
    with torch.no_grad():
        final_states, final_rewards, _ = trainer.sample_trajectories(batch_size=200)

        # 统计命中 Mode A 与 Mode B 的数量
        mode_A_hits = 0
        mode_B_hits = 0
        for s in final_states:
            if s[0] == 1 and s[1] == 1 and s[2:].sum() == 0:
                mode_A_hits += 1
            elif s[:4].sum() == 0 and s[4:].sum() == 4:
                mode_B_hits += 1

        print(f"    200 个生成样本中:")
        print(f"    - Mode A 命中数: {mode_A_hits} 个")
        print(f"    - Mode B 命中数: {mode_B_hits} 个")
        print(f"    - 其他探索性样本: {200 - mode_A_hits - mode_B_hits} 个")

        assert mode_A_hits > 0 and mode_B_hits > 0, "GFlowNet 发生了模式坍塌,未能同时覆盖双峰!"
        print(">>> [SUCCESS] GFlowNet 成功跨越不同骨架荒漠,实现多峰多样性无偏采样!")

5. 本章小结

本章深入探讨了直接在离散动作空间中开展目标导向逆向设计的最前沿架构:

  1. 反思传统强化学习局限:剖析了策略梯度与 PPO 在分子图生成中因极度追求单一最高分而必然发生的“模式坍塌”病理;
  2. 生成流网络(GFlowNets)的理论优雅:从网络流守恒公理推导了轨迹平衡(TB)损失函数,展现了其如何以概率正比于奖励(P(x)R(x)P(x) \propto R(x))从容捕获多个结构迥异的高价值化学骨架;
  3. 连续扩散的条件引导技术:厘清了分类器引导与无分类器引导(CFG)在结合物理力场与目标带隙约束时的连续去噪机制。

在掌握了这些算法利器后,我们即将在下一章迎来真正的“硬核物理战场”——第 05 章:无机晶体生成实战(从 CDVAE 到 MatterGen 与 GNoME),直面三维无限周期性晶胞的生成奇迹!

《原子智能》· 纸质出版预备版 · PolyAI Team 著