2026-07-22 11:52:44 +08:00
|
|
|
|
# TrajMixer Block 最终设计方案
|
|
|
|
|
|
|
|
|
|
|
|
> 状态:**Frozen implementation baseline**
|
|
|
|
|
|
>
|
|
|
|
|
|
> 版本:**v1.0**
|
|
|
|
|
|
>
|
|
|
|
|
|
> 固化日期:**2026-07-22**
|
|
|
|
|
|
|
|
|
|
|
|
本文档是 TrajMixer 后续实现与实验的唯一结构基线。除显式标记为消融项的配置外,所有实现均应遵循本文档;若结构发生变化,应先更新版本和实验记录。
|
|
|
|
|
|
|
|
|
|
|
|
## 1. 目标
|
|
|
|
|
|
|
|
|
|
|
|
在保持原始 Delphi Transformer Attention 结构不变的前提下,用轻量、可并行的轨迹交互模块替换 FFN。
|
|
|
|
|
|
|
|
|
|
|
|
保持不变的组件包括:
|
|
|
|
|
|
|
|
|
|
|
|
- 原始 causal mask;
|
|
|
|
|
|
- 原始 TimeRoPE / Relative Time Attention Bias;
|
|
|
|
|
|
- 原始 Multi-Head Attention,包括 \(W_Q/W_K/W_V/W_O\);
|
|
|
|
|
|
- 原始序列建模与训练目标。
|
|
|
|
|
|
|
|
|
|
|
|
TrajMixer 不修改 Attention,只替换每个 Transformer block 中的 FFN residual branch。
|
|
|
|
|
|
|
|
|
|
|
|
## 2. Block 总体结构
|
|
|
|
|
|
|
|
|
|
|
|
概念结构:
|
|
|
|
|
|
|
|
|
|
|
|
```text
|
|
|
|
|
|
PreNorm Causal Multi-Head Attention
|
|
|
|
|
|
→ Residual
|
|
|
|
|
|
→ Standard Mixer PreNorm
|
|
|
|
|
|
→ Group-wise Feature Alignment
|
|
|
|
|
|
→ SwiGLU Cross-Group Mixer
|
|
|
|
|
|
→ Residual
|
|
|
|
|
|
```
|
|
|
|
|
|
|
|
|
|
|
|
完整计算为:
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
U = X^{(l)} + \operatorname{Dropout}\!\left(
|
|
|
|
|
|
\operatorname{CausalMHA}\left(
|
|
|
|
|
|
\operatorname{LN}_{\mathrm{attn}}(X^{(l)}),
|
|
|
|
|
|
\text{time information}
|
|
|
|
|
|
\right)\right),
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
N = \operatorname{LN}_{\mathrm{mixer}}(U),
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
\Delta = \operatorname{TrajMixer}(N),
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
X^{(l+1)} = U + \operatorname{Dropout}(\Delta).
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
首版中的 \(\operatorname{LN}_{\mathrm{mixer}}\) 是作用于完整 \(d=120\) 维 residual representation 的标准 LayerNorm。
|
|
|
|
|
|
|
|
|
|
|
|
## 3. Latent Trajectory Group 定义
|
|
|
|
|
|
|
|
|
|
|
|
Attention 输出经过 \(W_O\) 后仍是标准 residual representation:
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
N\in\mathbb{R}^{B\times L\times d},\qquad d=120.
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
将 hidden dimension 划分为与 Attention head 数量相同的 group 数量:
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
n_{\mathrm{group}}:=n_{\mathrm{head}}=10,
|
|
|
|
|
|
\qquad d_{\mathrm{group}}=\frac{d}{n_{\mathrm{head}}}=12,
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
`n_group` 不再是独立超参数,代码统一使用 `n_head` 确定 residual group 数量。二者只共享数量;这些 residual groups 在语义和张量来源上仍不等同于原始 Attention heads。
|
|
|
|
|
|
|
|
|
|
|
|
并 reshape 为:
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
N_{\mathrm{group}}in
|
|
|
|
|
|
\mathbb{R}^{B\times L\times n_{\mathrm{group}}\times d_{\mathrm{group}}}.
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
这些 group 是 residual space 中的 **latent trajectory groups**,不等同于原始 Attention heads。本文中的 group、trajectory group 均指这一 residual-channel partition。
|
|
|
|
|
|
|
|
|
|
|
|
## 4. Group-wise Feature Alignment
|
|
|
|
|
|
|
|
|
|
|
|
为缓解不同 group 内部坐标不对齐的问题,每个 group 使用独立的小矩阵:
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
B_i\in\mathbb{R}^{d_{\mathrm{group}}\times d_{\mathrm{group}}},
|
|
|
|
|
|
\qquad i=1,\ldots,n_{\mathrm{group}}.
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
对每个 group 内的特征进行可学习对齐:
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
Z_{b,t,i,:}=N_{\mathrm{group},b,t,i,:}B_i.
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
因此:
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
Z\in
|
|
|
|
|
|
\mathbb{R}^{B\times L\times n_{\mathrm{group}}\times d_{\mathrm{group}}}.
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
首版实现约定:
|
|
|
|
|
|
|
|
|
|
|
|
- \(B_i\) 不带 bias;
|
|
|
|
|
|
- \(B_i\) 使用单位矩阵初始化;
|
|
|
|
|
|
- Alignment 只作用于 Mixer residual branch,不改变 Attention residual stream;
|
|
|
|
|
|
- 首版不增加逆变换或额外的 group 内输出投影。
|
|
|
|
|
|
|
|
|
|
|
|
Alignment 每层权重参数量为:
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
n_{\mathrm{group}}d_{\mathrm{group}}^2
|
|
|
|
|
|
=10\times12^2
|
|
|
|
|
|
=1{,}440.
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
## 5. SwiGLU Cross-Group Mixer
|
|
|
|
|
|
|
|
|
|
|
|
Mixer 只沿 group 维度交互,不沿序列维度交互,因此不会引入时间递归或未来信息泄漏。
|
|
|
|
|
|
|
|
|
|
|
|
对于每个 group 内特征维度:
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
r=1,\ldots,d_{\mathrm{group}},
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
定义:
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
A_g^{(r)},A_v^{(r)}
|
|
|
|
|
|
\in\mathbb{R}^{n_{\mathrm{group}}\times h_{\mathrm{group}}},
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
A_o^{(r)}
|
|
|
|
|
|
\in\mathbb{R}^{h_{\mathrm{group}}\times n_{\mathrm{group}}}.
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-22 16:14:52 +08:00
|
|
|
|
隐藏宽度不再独立配置,固定为:
|
2026-07-22 11:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
\[
|
2026-07-22 16:14:52 +08:00
|
|
|
|
h_{\mathrm{group}}=4n_{\mathrm{head}}
|
|
|
|
|
|
=4n_{\mathrm{group}}.
|
2026-07-22 11:52:44 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-22 16:14:52 +08:00
|
|
|
|
当前 \(n_{\mathrm{head}}=10\),因此 \(h_{\mathrm{group}}=40\)。
|
|
|
|
|
|
|
2026-07-22 11:52:44 +08:00
|
|
|
|
对固定的 batch、时间位置和内部特征维度 \(r\),将:
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
Z_{b,t,:,r}\in\mathbb{R}^{n_{\mathrm{group}}}
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
视为 row vector,计算:
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
G_{b,t,:,r}=Z_{b,t,:,r}A_g^{(r)},
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
V_{b,t,:,r}=Z_{b,t,:,r}A_v^{(r)},
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
M_{b,t,:,r}=\operatorname{SiLU}(G_{b,t,:,r})\odot V_{b,t,:,r},
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
Y_{b,t,:,r}=M_{b,t,:,r}A_o^{(r)}.
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
其中:
|
|
|
|
|
|
|
|
|
|
|
|
- gate 分支控制信息写入;
|
|
|
|
|
|
- value 分支提供交互内容;
|
|
|
|
|
|
- output matrix 将隐藏 group 表示投影回原始 group 数量;
|
2026-07-22 16:14:52 +08:00
|
|
|
|
- hidden group 表示固定扩展为 group 数量的 4 倍。
|
2026-07-22 11:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
所有 \(r\) 的输出组合为:
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
Y\in
|
|
|
|
|
|
\mathbb{R}^{B\times L\times n_{\mathrm{group}}\times d_{\mathrm{group}}},
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
再 reshape 为:
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
\Delta\in\mathbb{R}^{B\times L\times d}.
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
## 6. 参数张量与无歧义索引
|
|
|
|
|
|
|
|
|
|
|
|
建议的实现存储形状为:
|
|
|
|
|
|
|
|
|
|
|
|
```text
|
|
|
|
|
|
group_align: [n_group, d_group, d_group]
|
|
|
|
|
|
gate_proj: [d_group, n_group, hidden_group]
|
|
|
|
|
|
value_proj: [d_group, n_group, hidden_group]
|
|
|
|
|
|
output_proj: [d_group, hidden_group, n_group]
|
|
|
|
|
|
```
|
|
|
|
|
|
|
|
|
|
|
|
对应的索引公式为:
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
G_{b,t,q,r}
|
|
|
|
|
|
=\sum_i Z_{b,t,i,r}\,A_{g,r,i,q},
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
V_{b,t,q,r}
|
|
|
|
|
|
=\sum_i Z_{b,t,i,r}\,A_{v,r,i,q},
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
Y_{b,t,i,r}
|
|
|
|
|
|
=\sum_q
|
|
|
|
|
|
\left[\operatorname{SiLU}(G_{b,t,q,r})V_{b,t,q,r}\right]
|
|
|
|
|
|
A_{o,r,q,i}.
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
首版的三个 Mixer projection 均不带 bias。
|
|
|
|
|
|
|
|
|
|
|
|
## 7. Mixer Hidden Width 与参数量
|
|
|
|
|
|
|
|
|
|
|
|
Mixer 的 group 维变换为:
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
n_{\mathrm{group}}
|
|
|
|
|
|
\rightarrow
|
|
|
|
|
|
h_{\mathrm{group}}
|
|
|
|
|
|
\rightarrow
|
|
|
|
|
|
n_{\mathrm{group}}.
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-22 16:14:52 +08:00
|
|
|
|
固定 \(h_{\mathrm{group}}=4n_{\mathrm{group}}=40\) 时,Mixer 每层权重参数量为:
|
2026-07-22 11:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
3d_{\mathrm{group}}n_{\mathrm{group}}h_{\mathrm{group}}
|
2026-07-22 16:14:52 +08:00
|
|
|
|
=3\times12\times10\times40
|
|
|
|
|
|
=14{,}400.
|
2026-07-22 11:52:44 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
加上 Group Feature Alignment 后,TrajMixer residual branch 每层共有:
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
2026-07-22 16:14:52 +08:00
|
|
|
|
14{,}400+1{,}440=15{,}840
|
2026-07-22 11:52:44 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
个主要权重参数。作为对照,原始 \(120\rightarrow480\rightarrow120\) FFN 每层约有 115,800 个参数。
|
|
|
|
|
|
|
|
|
|
|
|
参数对照口径说明:上面的 115,800 对应结构方案中的标准两层 FFN。当前代码库在 TrajMixer 替换前实际使用的是隐藏宽度 300 的全维度 SwiGLU(gate/value/output 三个线性层),每层共有 108,720 个参数(含 bias)。代码实验和 checkpoint 参数量比较必须以 108,720 作为历史实现基线,不能与概念方案中的标准 FFN 参数量混用。
|
|
|
|
|
|
|
|
|
|
|
|
## 8. LayerNorm 基线与消融
|
|
|
|
|
|
|
|
|
|
|
|
为保持与原始 Transformer 的可比性,首版固定使用:
|
|
|
|
|
|
|
|
|
|
|
|
```text
|
|
|
|
|
|
原始 FFN baseline:FFN + 标准 LayerNorm
|
|
|
|
|
|
TrajMixer baseline:Mixer + 标准 LayerNorm
|
|
|
|
|
|
```
|
|
|
|
|
|
|
|
|
|
|
|
以下配置不属于首版主实验,只作为独立消融:
|
|
|
|
|
|
|
|
|
|
|
|
```text
|
|
|
|
|
|
Mixer + Group-wise LayerNorm
|
|
|
|
|
|
```
|
|
|
|
|
|
|
|
|
|
|
|
不得将 Group-wise LayerNorm 的结果直接作为“仅替换 FFN”的对照结果。
|
|
|
|
|
|
|
|
|
|
|
|
## 9. 初始化
|
|
|
|
|
|
|
|
|
|
|
|
首版初始化约定:
|
|
|
|
|
|
|
|
|
|
|
|
- Group Alignment \(B_i\):单位矩阵初始化;
|
|
|
|
|
|
- \(A_g/A_v\):Xavier uniform 初始化;
|
|
|
|
|
|
- \(A_o\):均值为 0、标准差为 \(10^{-3}\) 的正态初始化;
|
|
|
|
|
|
- Dropout 概率沿用原始 FFN residual branch 的配置。
|
|
|
|
|
|
|
|
|
|
|
|
小初始化的 \(A_o\) 使新增分支在训练初期接近恒等残差更新,同时允许模型逐步学习轨迹交互。
|
|
|
|
|
|
|
|
|
|
|
|
## 10. 核心设计思想
|
|
|
|
|
|
|
|
|
|
|
|
**Attention**:负责从历史疾病序列中选择并整合相关信息。
|
|
|
|
|
|
|
|
|
|
|
|
**Group Feature Alignment**:负责学习不同 latent trajectory groups 的内部特征对齐。
|
|
|
|
|
|
|
|
|
|
|
|
**Cross-Group Mixer**:负责不同潜在疾病轨迹之间的非线性门控交互。
|
|
|
|
|
|
|
|
|
|
|
|
整个模块保持:
|
|
|
|
|
|
|
|
|
|
|
|
- 无时间递归;
|
|
|
|
|
|
- 序列维度完全并行;
|
|
|
|
|
|
- 参数量远低于原始 FFN;
|
|
|
|
|
|
- 保留 Transformer 的因果历史建模能力;
|
|
|
|
|
|
- 不把 residual groups 误解释为原始 Attention heads。
|
|
|
|
|
|
|
|
|
|
|
|
## 11. 首版固定配置
|
|
|
|
|
|
|
|
|
|
|
|
```yaml
|
2026-07-22 16:14:52 +08:00
|
|
|
|
model_architecture: traj_mixer_v2
|
2026-07-22 11:52:44 +08:00
|
|
|
|
d_model: 120
|
|
|
|
|
|
n_head: 10 # 同时决定 residual group 数量
|
|
|
|
|
|
d_group: 12
|
2026-07-22 16:14:52 +08:00
|
|
|
|
hidden_group_rule: 4 * n_head # 不单独配置
|
2026-07-22 11:52:44 +08:00
|
|
|
|
attention: unchanged
|
|
|
|
|
|
attention_output_projection: unchanged
|
|
|
|
|
|
mixer_norm: standard_layer_norm
|
|
|
|
|
|
group_alignment: per_group_12x12
|
|
|
|
|
|
group_alignment_bias: false
|
|
|
|
|
|
group_alignment_init: identity
|
|
|
|
|
|
mixer_bias: false
|
|
|
|
|
|
gate_value_init: xavier_uniform
|
|
|
|
|
|
output_init_std: 0.001
|
|
|
|
|
|
group_wise_layer_norm: false
|
|
|
|
|
|
```
|
|
|
|
|
|
|
2026-07-22 16:14:52 +08:00
|
|
|
|
训练时必须将 `model_architecture: traj_mixer_v2`、`model_parameter_count` 和 `trainable_parameter_count` 写入 `train_config.json`,并在训练日志中显式打印总参数量与可训练参数量。本分支的评估和导出入口只接受带有该架构标识、且 checkpoint 中包含 TrajMixer 参数张量的模型;其他版本或分支生成的模型应直接拒绝加载。
|
2026-07-22 11:52:44 +08:00
|
|
|
|
|
|
|
|
|
|
必须满足:
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
d=n_{\mathrm{group}}d_{\mathrm{group}}.
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
后续实现、单元测试、参数量核验和主实验均以以上配置为默认基线。
|