8.4 KiB
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_group 由 n_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:
intra_norm只服务于组内 Mixer;cross_norm只服务于跨组 Mixer;- 第一阶段 residual 的输出是第二阶段的输入;
- 两个 residual 都在
TrajMixer内部完成; - 不再保留 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}) 的正态分布; - 两个 LayerNorm:PyTorch 默认 affine 初始化;
- 两个 residual stage 的 Dropout 均沿用
mlp_dropout。
两个 output projection 的小方差初始化使两阶段在训练初期都接近恒等 residual update。
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_v3、model_parameter_count 和 trainable_parameter_count 写入 train_config.json,并在训练日志中显式打印参数量。
本分支的评估和导出入口只接受 traj_mixer_v3 checkpoint,并检查两阶段 Norm、组内 projection、Group Alignment 和跨组 projection 参数是否齐全。traj_mixer_v2 及更早 checkpoint 不向后兼容,直接拒绝加载。