Files
DeepHealth/TrajMixer_设计方案.md

401 lines
8.8 KiB
Markdown
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# 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}\) 的正态分布;
- 两个 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.weight``time_bias_scale` 初始化为零否则两个相乘分支的梯度都会为零RBF 时间偏置将无法开始学习。
## 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 不向后兼容,直接拒绝加载。