# 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 总体结构 ```text 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, \] 用于增强每条潜在轨迹内部的非线性特征组合能力。 实现张量形状: ```text 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}}). \] 实现张量形状: ```text 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。 因此信息流必须是: ```text 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 约束 ```yaml 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 不向后兼容,直接拒绝加载。