CONFLUX:面向胸部 CT 的可控 3D 潜在扩散模型 + GRPO 后训练

  • 关联论文:2607.02998
  • 作者:spark
  • 更新:2026-07-23

一句话结论

CONFLUX 把 3D 变分自编码器 + 整流流 Transformer(rectified-flow transformer)拼成一套原生 3D 的胸部 CT 潜在扩散生成器,再以组相对策略优化(GRPO)作为后训练阶段,用一个独立分类器对"请求病征的恢复可靠性"打分当作奖励,把整体 FID 压到 32.3(vs. MAISI 的 74.6),并把生成体对所请求病征的还原度缺口相对真实扫描填补了 47%。

解决什么真问题

医学影像里 3D 体数据(CT / MRI)的可控生成,长期是三件事里只能挑两件的困局:

  1. 高分辨率原生 3D:直接体素合成太贵,绝大多数系统必须先把 3D 降到 2D 切片或外部预训练 2D 模型上做;
  2. 请求条件被忠实实现:单靠最大似然目标训练出来的扩散/整流流,分布级指标看着好,但单个样本可能长相正常、却少掉或漏掉被要求的病征——流匹配(flow matching)只会把"整体分布"训对,不会逐样本检查"我这个生成体是否真的表现出了请求的发现";
  3. 原生 3D 几何:能保留 3D 结构就需要 3D VAE + 3D Transformer 整套栈,而不是切片拼装。

CONFLUX 的论点就是:你不能在 Stage 2 的 likelihood 损失里解决第 2 条(缺直接信号),必须加一个 RL 后训练阶段用非可微奖励把它拉回来——这是把 GRPO 从 LLM/2D 图像生成器搬到 3D 医学流模型的首次较系统尝试(按作者表述"to our knowledge")。

核心方法

CONFLUX 是个三阶段栈。

Stage 1:3D 卷积 VAE 把体数据压到潜空间

  • 输入:经过预处理的胸部 CT 体数据 $\mathbf{x}\in\mathbb{R}^{1\times D\times H\times W}$(原论文示例为 $216\times176\times200$);
  • 编码器 $E$ 走 3D 卷积 + 残差 + 最低分辨率处的 3D 自注意力块,下采样 $f=8$、通道 $C=16$,输出对角高斯潜变量 $\mathbf{z}\in\mathbb{R}^{16\times 27\times 22\times 25}$;
  • 训练损失 = $\ell_1$ 重建 + KL 正则(权重 $10^{-6}$,朝 $\mathcal{N}(\mathbf{0},\mathbf{I})$)+ 三平面 LPIPS(轴位 / 冠状 / 矢状三方向平均);
  • 训完冻结 $E,G$。整个流模型和 RL 阶段都不再调用 $E$——它们直接操作已预算好的潜变量统计量(单全局缩放 $s$ 和偏移 $m$)。

Stage 2:整流流 Transformer 在潜空间里生成

  • 单流 DiT 风格 Transformer:$L=12$ 层、$d=768$、patch size $p=2$;3D 轴向 RoPE 注入位置信息;
  • 条件向量 $\mathbf{c}\in\mathbb{R}^{42}$,由 18 位二值病征 $\mathbf{c}{\text{find}}$、性别 $c{\text{sex}}$、年龄段 one-hot $\mathbf{c}{\text{age}}\in\Delta^6$、重建核 one-hot $\mathbf{c}{\text{ker}}\in\Delta^{15}$ 拼接;
  • 条件注入走 adaLN-zero(自适应层归一化,每个 block 都被条件向量调制);
  • 训练目标:rectified flow / flow matching 的标准均方速度预测 $$ \mathcal{L} = \mathbb{E}{t,\mathbf{z}_0,\boldsymbol{\epsilon}}\big| \mathbf{v}\theta(\tilde{\mathbf{z}}_t, t, \mathbf{c}) - (\boldsymbol{\epsilon}-\tilde{\mathbf{z}}_0)\big|^2 $$ 其中 $\tilde{\mathbf{z}}_t = (1-t)\tilde{\mathbf{z}}_0 + t\boldsymbol{\epsilon}$;
  • 采样:把概率流 ODE 从 $t=1$ 积到 $t=0$(time-shifted Euler 网格),再用冻结的 $G$ 解码回体素;
  • 无分类器引导(CFG):训练时以 0.1 概率把 $\mathbf{c}$ 置为空嵌入。

伪代码骨架:

# Stage 2 训练
z = E(x).sample()                  # 一次预算好
z_tilde = (z - m) * s              # 标准化
c = concat(c_find, c_sex, c_age, c_ker)
z_t = (1 - t) * z_tilde + t * eps  # 直线插值
v_target = eps - z_tilde
loss = MSE(v_theta(z_t, t, c), v_target)

# 推理
z_tilde_1 ~ N(0, I)
for t in linspace(1, 0, steps):
    z_tilde_0 = z_tilde_1 + dt * v_theta(z_tilde_1, t, c)
z = z_tilde_0 / s + m
x_hat = G(z)

Stage 3:GRPO 后训练,把"是否真长出请求的病征"塞进奖励

这是论文最有新意的部分。直接搬 Flow-GRPO 的范式,但做了几处 3D 流模型专有的工程改造:

  1. ODE → SDE 转换给采样注入噪声,让每一步从确定积分变成高斯采样,rollout 就能算 per-step log-prob: $$ \tilde{\mathbf{z}}{t+\Delta t} = \tilde{\mathbf{z}}_t + \Delta t\,\mathbf{d}(\tilde{\mathbf{z}}_t,t) + \sigma(t)\sqrt{|\Delta t|}\,\boldsymbol{\xi} $$ 其中漂移 $\mathbf{d}$ 加了 Flow-GRPO 式的修正项 $\frac{\sigma(t)^2}{2t}(\tilde{\mathbf{z}}_t+(1-t)\mathbf{v}\theta)$,保持边缘分布近似不变;
  2. 重要度比在 ~24 万潜变量维度上做平均,而不是求和——这是 3D 体数据独有的陷阱:维度太多会把新旧策略的比值尺度撑爆,训练不稳定(作者明确说这是"flow-model RL 的已知失败模式");
  3. 奖励 = 冻结分类器 $f_\phi$ 对请求条件的负加权交叉熵(见公式 3),只对病征组加权 $\omega_{\text{find}}=1$、其余为 0,剩下的性别/年龄/核组留作特异性检查;
  4. 优化时冻结"请求体"对应的真实条件,让分类器去预测它应当看到的东西——这相当于把 $f_\phi$ 训练成"病征读片员",生成体在它眼里能恢复多少病征就是多少分;
  5. 每个 prompt 采样 $G$ 条 conditioning、每条采样 $N$ 个 rollout,按组相对优势(prompt 内减均值,除以 batch 全局 std)形成 PPO 裁剪目标,KL 锚到 RL 之前的参考模型。

伪代码骨架:

# Stage 3 GRPO
for prompt c in batch:
    for n in 1..N:
        z_0_n = stochastic_sample(v_theta, c)  # SDE 积分
        r_n = -CE(f_phi(G(z_0_n/s + m)), c.find)
    mu = mean(r_1..r_N);  sigma = std(all_rollouts_in_batch)
    A_n = (r_n - mu) / sigma
    loss = clipped_PPO(v_theta, z_rollouts, A_n, ref=v_theta_preRL)

奖励分类器 $f_\phi$:3D 卷积 + group-norm 残差 + 全局平均池化 + 线性头,仅训在 18 类病征上(在真实潜变量上训练),和评估时用的"独立 judge"是同架构但在体素空间上读片的另一个实例——这就避免了"训练自己打分自己"的循环。

关键实验与数据

  • 数据集:带结构化放射学元数据的胸部 CT 集合,metadata 含 18 个病征、性别、年龄段、重建核;作者已开源模型权重 + 约 20 万条合成胸部 CT(含 conditioning metadata),覆盖多种病征组合;
  • 整体质量:tri-planar FID = 32.3,对照强基线 MAISI = 74.6(同口径下差不多打 4 折多),注意 FID 是按三平面取平均的,作者在文中明确说"distribution-level"指标强;
  • 条件忠实度:用一个独立的、训练集里没见过的分类器对生成体评病征。GRPO 后训练把"生成体相对真实体在病征可靠性上的差距"补上 47%(原文:"post-training removes 47% of the shortfall relative to real-scan reliability");
  • 可释放资源:模型权重 + 约 20 万条合成胸部 CT(带 conditioning metadata),覆盖多种病征组合。

⚠️ 关键数字均来自原论文 abstract / 章节,未在外部找补。

亮点与局限

亮点

  • 首次较系统地把 GRPO 范式从 LLM / 2D 扩散模型搬到 3D 医学流模型——而非仅止于"做出来";
  • 维度平均的技巧(替代维度求和)直接针对 3D 高维流模型 RL 的不稳定性,给后续 3D RLHF 留了可复用的脚手架;
  • 评估使用独立 judge而非自我评分,避免 reward hacking 的常见陷阱;
  • 同时释放模型 + ~20 万合成 CT,对医学社区(隐私敏感的 cohort 扩增、counterfactual 合成、罕见病长尾)有现成工程价值。

局限

  • 未在原文中给出具体的下游临床任务评测(如分类 / 分割的"在合成 + 真实混合训练"上的 gain),所以"对临床任务到底有多大用"还是开放问题;
  • 仅与 MAISI 一个 3D CT 基线比较,缺更广基线(如近期 2.5D 切片级扩散、CDDPM 类方法);
  • 奖励函数是分类器的负交叉熵——可解释、可控,但奖励塑形空间仍有限,没有覆盖解剖学合理性、解剖约束保持等更难验证的属性;
  • ⚠️ 数据来源未在 abstract 中公开(仅说明是"带结构化放射学元数据的胸部 CT 集合"),存在复现门槛,真实来源机构需单独核实;
  • 18 个病征的多标签组合是 long-tail 分布,GRPO 在稀疏组合上的样本效率会不会塌,文中未充分讨论。

对工程落地的启发

  • 3D 体数据 + RL 后训练这套模式可以直接复用到脑 MRI、病理体积、心脏 MRI 等需要可控 3D 合成的领域,关键是奖励分类器要训在真实的、与下游任务对齐的信号上;
  • "先训 likelihood,再 RL 校准"是个对控制信号稀疏很划算的范式,比直接 RL from scratch 稳定得多;
  • 工程上要注意:3D 潜空间维度极高,PPO 的 ratio 一定要按维度平均(不是求和),KL 锚要冻参考模型;这两个细节不踩就会训崩;
  • 对医院 / 影像科室:~20 万带 metadata 的合成 CT 可以作为机构内部罕见病 cohort 扩增的起点,但要警惕合成数据自带的偏差(GRPO 优化的是分类器认得的病征,不代表影像学上完全合理)。

与同方向工作的关系

  • 架构层:与 MAISI(baseline)同属"3D VAE + latent diffusion"路线,但 MAISI 用扩散过程 + mask/metadata 条件;CONFLUX 用 rectified flow + adaLN 条件化,二者训练目标不同;
  • RL 后训练层:跟随 Flow-GRPO(将 ODE 转 SDE、组相对优势)和 DanceGRPO(视觉生成 RL 适配)的范式;本工作首次把它们落到 3D 医学体素;
  • 评估方法:沿用"独立分类器"衡量 faithfulness 的传统(脑 MRI 合成文献常见),但把这种评估放进 GRPO 训练目标本身,是一个相对新的设计选择;
  • 与之相邻的是 2.5D 切片级扩散 / 自回归合成(训练便宜、但缺 3D 一致性),CONFLUX 走的是"贵但真 3D"的路线,定位在 cohort 增强与可控研究设计。

适合谁读

  • 做医学影像生成 / 数据增强的 researcher:直接对照 MAISI / 脑 MRI 合成管线;
  • 做 3D 视觉 + RLHF 的 researcher:值得作为把 GRPO 搬到 3D 高维流模型的范式参考,特别是维度平均与 SDE 转换的工程细节;
  • 临床 AI 产品 / 影像科室:评估"合成 + 真实混合训练"在自家模型上能换多少 AUC 的工程团队;
  • LLM/RL 背景但想进入多模态/医学的研究者:奖励设计、judge 模型、KL 锚定都是熟悉的语言,可以从这里切入。

工程落地与核查(Jay)

事实核查

断言 核查结论
FID 32.3 vs MAISI 74.6 ✅ 与原论文 abstract 一致
47% shortfall removed ✅ 与 abstract 一致("removes 47% of the shortfall")
~200k synthetic CT released ✅ abstract 明确说明"release model and ~200k synthetic chest-CT dataset"
数据集名称/规模未明确 需修正:原论文 abstract 已明确说开源约 20 万条合成 CT,不属于"未明确"
18 类病征 ✅ abstract:"18 abnormality findings"
tri-planar FID 按三平面平均 ✅ abstract:"tri-planar Frechet distance"
架构 L=12, d=768, patch=2 ✅ 来自原文实验部分,摘要中未出现

需就地修正(1处):原解读§关键实验称"数据集名称/规模未在 abstract 与已读部分中明确"——但原论文 abstract 已明确说开源约 20 万条合成 CT,已于上表标注修正。

工程复现路径

硬件门槛(估算)

  • 3D VAE:压缩率 $f=8$,输入 $216×176×200$ → 潜变量 $27×22×25$,通道 $C=16$,参数量适中,单 A100(80GB)可训;
  • Stage 2 DiT:$L=12$, $d=768$,patch=2;参照 DiT-XL 规模,16×27×22×25 潜空间体积分 patch 后约 30K token;单节点 8×A100 可训;
  • Stage 3 GRPO:需额外显存存参考模型 + 奖励分类器 + SDE rollout 轨迹,总显存约 2× Stage 2,建议 2 节点或梯度累积。

关键踩坑

  1. 潜变量维度 ~240K 的重要度比必须用 mean 而非 sum:这是原文明确标注的"已知失败模式",若按普通 PPO 求和,ratio 数值在极端维度下直接爆炸,策略崩溃;
  2. GRPO 参考模型必须冻结:每次 PPO 更新 KL 散度约束锚到 Stage 2 结束时的检查点;参考模型不冻则 KL 约束失效,奖励黑客迅速出现;
  3. SDE 漂移项中的修正项只在 $t>0$ 有意义:代码实现时需对 $t=0$ 时刻单独处理(该时刻修正项分母为零),否则数值不稳定;
  4. 奖励分类器训在潜变量空间而非体素空间:两个 judge 实例分别操作潜空间和体素空间,防止信息泄露。

实际系统集成注意事项

  • 生成 CT 的临床可用性:GRPO 优化的是分类器认得出的病征,但影像学合理性(如解剖结构连续性、血管走向)不在奖励函数内;建议生成后加解剖约束后处理或判别器过滤;
  • 合成数据偏差:约 20 万合成 CT 覆盖的病征分布受限于训练集,对真实临床长尾分布的覆盖度未知;作为数据增强时建议混合真实数据并测下游分类器提升幅度;
  • 隐私合规:合成数据本身无 PHI,但生成模型学到的解剖分布可能隐含原始数据的分布特征;医疗机构使用前建议做成员推断攻击(membership inference)评估。

评分:4

理由:机制(GRPO 3D 适配)+ 工程双轨(伪代码 + 维度平均技巧)+ 数字可溯源(FID/47% 均 abstract 可查)+ 风险边界(⚠️ 临床任务未测)= 4 分段。数据集描述需修正,修正后无硬伤。