Files
DeepHealth/TrajMixer_设计方案.md

8.8 KiB
Raw Blame History

TrajMixer Block 最终设计方案

状态:Frozen implementation baseline

版本:v2.0 / traj_mixer_v3

固化日期:2026-07-24

本文档是 TrajMixer 后续实现与实验的结构基线。本版本将原先仅含跨轨迹交互的 TrajMixer 扩展为“组内混合 + 跨组混合”的两阶段结构。

1. 目标

在保持原始 Delphi Transformer Attention 结构不变的前提下,用轻量、可并行的两阶段 TrajMixer 替换 FFN。

保持不变的组件包括:

  • 原始 causal mask
  • 原始 TimeRoPE / Relative Time Attention Bias
  • 原始 Multi-Head Attention包括 (W_Q/W_K/W_V/W_O)
  • 原始序列建模和训练目标。

TrajMixer 不沿序列维度混合,也不引入时间递归。

2. Block 总体结构

PreNorm Causal Multi-Head Attention
→ Attention Residual
→ reshape [B, L, n_group, d_group]
→ Intra-Group PreNorm
→ Per-Group SwiGLU: d_group → 4d_group → d_group
→ Intra-Group Residual
→ Cross-Group PreNorm
→ Group-wise Feature Alignment
→ Cross-Group SwiGLU: n_group → 4n_group → n_group
→ Cross-Group Residual
→ reshape [B, L, d]

Attention 阶段保持原样:

[ U=X^{(l)}+\operatorname{CausalMHA} \left(\operatorname{LN}_{\mathrm{attn}}(X^{(l)})\right). ]

随后:

[ H^{(0)} =\operatorname{reshape}(U) \in\mathbb{R}^{B\times L\times G\times D}, ]

其中 (G=n_{\mathrm{group}})(D=d_{\mathrm{group}})。

两阶段 TrajMixer 为:

[ H^{(1)} =H^{(0)} +\operatorname{Dropout} \left(\operatorname{IntraMixer} \left(\operatorname{LN}_{D}(H^{(0)})\right)\right), ]

[ H^{(2)} =H^{(1)} +\operatorname{Dropout} \left(\operatorname{CrossMixer} \left(\operatorname{LN}_{G}(H^{(1)})\right)\right), ]

[ X^{(l+1)}=\operatorname{reshape}(H^{(2)}) \in\mathbb{R}^{B\times L\times d}. ]

TrajMixer.forward() 返回的是已经完成两次 residual update 的完整状态,而不是单个 residual delta。因此 GPTBlock 在 Attention residual 后直接返回 TrajMixer(U),不得再写成 U + TrajMixer(U)

3. Latent Trajectory Group 定义

Attention 输出经过 (W_O) 后仍是标准 residual representation

[ U\in\mathbb{R}^{B\times L\times d}. ]

固定:

[ G:=n_{\mathrm{group}}=n_{\mathrm{head}}, \qquad D:=d_{\mathrm{group}}=\frac{d}{G}, \qquad d=GD. ]

默认配置:

[ d=120,\qquad G=10,\qquad D=12. ]

reshape 后:

[ H^{(0)}\in\mathbb{R}^{B\times L\times G\times D}. ]

n_groupn_head 决定,但 residual groups 只是 residual space 的连续分区,不等同于 Attention heads。

4. 第一阶段:组内 SwiGLU Mixer

第一阶段对每个 group 独立进行特征变换。不同 group 使用各自的投影参数,不发生 group 间信息交换。

先对每个 ((b,t,g)) 的 (D) 维向量独立执行 LayerNorm

[ \widetilde H^{(0)} =\operatorname{LN}_{D}(H^{(0)}). ]

归一化统计量在每个 group 内独立计算为保持轻量LayerNorm 的 affine 参数在各 group 间共享。

对于 (g=1,\ldots,G),定义:

[ W_{g,\mathrm{intra}}^{(g)}, W_{v,\mathrm{intra}}^{(g)} \in\mathbb{R}^{D\times 4D}, ]

[ W_{o,\mathrm{intra}}^{(g)} \in\mathbb{R}^{4D\times D}. ]

计算:

[ P_g =\operatorname{SiLU} \left(\widetilde H^{(0)}g W{g,\mathrm{intra}}^{(g)}\right) \odot \left(\widetilde H^{(0)}g W{v,\mathrm{intra}}^{(g)}\right), ]

[ \Delta_{\mathrm{intra},g} =P_gW_{o,\mathrm{intra}}^{(g)}, ]

[ H^{(1)} =H^{(0)} +\operatorname{Dropout}(\Delta_{\mathrm{intra}}). ]

该阶段完成:

[ D\rightarrow4D\rightarrow D, ]

用于增强每条潜在轨迹内部的非线性特征组合能力。

实现张量形状:

intra_norm:        LayerNorm(d_group)
intra_gate_proj:   [n_group, d_group, 4 * d_group]
intra_value_proj:  [n_group, d_group, 4 * d_group]
intra_output_proj: [n_group, 4 * d_group, d_group]

三个 projection 均不带 bias。

5. 第二阶段:跨组 TrajMixer

第二阶段沿 group 维度进行交互。对于每个内部坐标 (r),独立执行:

[ G\rightarrow4G\rightarrow G. ]

首先将 (H^{(1)}) 的最后两个维度交换,并在 group 维度执行 LayerNorm

[ \widetilde H^{(1)}{b,t,:,r} =\operatorname{LN}{G} \left(H^{(1)}_{b,t,:,r}\right). ]

归一化统计量对每个内部坐标 (r) 独立计算LayerNorm 的 affine 参数在各内部坐标间共享。

5.1 Group-wise Feature Alignment

沿用现有的可学习 group 特征对齐矩阵:

[ B_g\in\mathbb{R}^{D\times D}, \qquad g=1,\ldots,G, ]

[ Z_{b,t,g,:} =\widetilde H^{(1)}_{b,t,g,:}B_g. ]

(B_g) 不带 bias并使用单位矩阵初始化。

5.2 Cross-Group SwiGLU

对每个内部坐标 (r=1,\ldots,D),定义:

[ A_g^{(r)},A_v^{(r)} \in\mathbb{R}^{G\times4G}, \qquad A_o^{(r)} \in\mathbb{R}^{4G\times G}. ]

计算:

[ Q_{b,t,:,r} =\operatorname{SiLU} \left(Z_{b,t,:,r}A_g^{(r)}\right) \odot \left(Z_{b,t,:,r}A_v^{(r)}\right), ]

[ \Delta_{\mathrm{cross},b,t,:,r} =Q_{b,t,:,r}A_o^{(r)}, ]

[ H^{(2)} =H^{(1)} +\operatorname{Dropout}(\Delta_{\mathrm{cross}}). ]

实现张量形状:

cross_norm:  LayerNorm(n_group)
group_align: [n_group, d_group, d_group]
gate_proj:   [d_group, n_group, 4 * n_group]
value_proj:  [d_group, n_group, 4 * n_group]
output_proj: [d_group, 4 * n_group, n_group]

三个 projection 均不带 bias。

6. PreNorm 与残差约束

本版本固定使用两个独立的 PreNorm residual stage

  1. intra_norm 只服务于组内 Mixer
  2. cross_norm 只服务于跨组 Mixer
  3. 第一阶段 residual 的输出是第二阶段的输入;
  4. 两个 residual 都在 TrajMixer 内部完成;
  5. 不再保留 block 外部的全维度 ln2 或额外 Mixer residual。

因此信息流必须是:

U
→ U + IntraMixer(IntraNorm(U))
→ H1 + CrossMixer(CrossNorm(H1))
→ output

7. 参数量

默认 (d=120,G=10,D=12)。

7.1 组内阶段

投影权重:

[ 3G D(4D) =12GD^2 =17{,}280. ]

LayerNorm(D)

[ 2D=24. ]

7.2 跨组阶段

跨组投影权重:

[ 3D G(4G) =12DG^2 =14{,}400. ]

Group Feature Alignment

[ GD^2 =1{,}440. ]

LayerNorm(G)

[ 2G=20. ]

7.3 每个 TrajMixer 合计

[ 17{,}280+24+14{,}400+1{,}440+20 =\boxed{33{,}164}. ]

相对于 traj_mixer_v2 的跨组单阶段结构 (15{,}840),每层增加 (17{,}324) 个参数。作为历史实现对照,代码库原全维度 SwiGLU FFN 每层为 (108{,}720) 个参数。

8. 初始化

固定初始化约定:

  • 组内 intra_gate_proj/intra_value_proj:每个 group 独立 Xavier uniform
  • 组内 intra_output_proj:均值 0、标准差 (10^{-3}) 的正态分布;
  • Group Alignment单位矩阵
  • 跨组 gate_proj/value_proj:每个内部坐标独立 Xavier uniform
  • 跨组 output_proj:均值 0、标准差 (10^{-3}) 的正态分布;
  • 两个 LayerNormPyTorch 默认 affine 初始化;
  • 两个 residual stage 的 Dropout 均沿用 mlp_dropout

两个 output projection 的小方差初始化使两阶段在训练初期都接近恒等 residual update。

Relative Time Attention Bias 的初始化固定为:

  • rbf_proj.weight:零初始化;
  • time_bias_scale:初始化为 (1.0)
  • 初始 RBF attention bias 仍严格为零;
  • rbf_proj.weight 从第一个优化步骤即可获得梯度。

不得同时将 rbf_proj.weighttime_bias_scale 初始化为零否则两个相乘分支的梯度都会为零RBF 时间偏置将无法开始学习。

9. 信息流与语义

Attention:从历史疾病事件中选择和整合相关信息。

Intra-Group Mixer:学习每条潜在轨迹内部的非线性特征组合。

Group Feature Alignment:对齐不同潜在轨迹的内部坐标。

Cross-Group Mixer:学习不同潜在轨迹在相同内部坐标上的门控交互。

整个模块保持:

  • 无时间递归;
  • 不混合序列位置;
  • 序列维度完全并行;
  • 保留原始因果 Attention
  • residual groups 不等同于 Attention heads。

10. 固定配置与 checkpoint 约束

model_architecture: traj_mixer_v3
d_model: 120
n_head: 10
n_group_rule: n_head
d_group_rule: d_model / n_group
intra_hidden_rule: 4 * d_group
cross_hidden_rule: 4 * n_group
attention: unchanged
intra_norm: layer_norm_over_d_group
cross_norm: layer_norm_over_n_group
group_alignment: per_group_d_group_x_d_group
projection_bias: false
gate_value_init: xavier_uniform
output_init_std: 0.001

必须满足:

[ d=n_{\mathrm{group}}d_{\mathrm{group}}. ]

训练时必须将 model_architecture: traj_mixer_v3model_parameter_counttrainable_parameter_count 写入 train_config.json,并在训练日志中显式打印参数量。

本分支的评估和导出入口只接受 traj_mixer_v3 checkpoint并检查两阶段 Norm、组内 projection、Group Alignment 和跨组 projection 参数是否齐全。traj_mixer_v2 及更早 checkpoint 不向后兼容,直接拒绝加载。