Files
DeepHealth/TrajMixer_设计方案.md

386 lines
7.5 KiB
Markdown
Raw Permalink 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**
>
> 版本:**v3.0 / traj_mixer_v5**
>
> 固化日期:**2026-07-24**
本文档是当前 TrajMixer 的实现与实验基线。本版本采用单 PreNorm、单外层 residual、静态门控组内融合和跨 group SwiGLU。
## 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
→ Full-width TrajMixer PreNorm
→ reshape [B, L, n_group, d_group]
→ Per-Group SwiGLU: d_group → 4d_group → d_group
→ Static Gated Fusion
→ Cross-Group SwiGLU: n_group → 4n_group → n_group
→ reshape [B, L, n_embd]
→ Dropout
→ One TrajMixer Residual
```
Attention 阶段:
\[
X
=X^{(l)}
+\operatorname{CausalMHA}
\left(\operatorname{LN}_{\mathrm{attn}}(X^{(l)})\right).
\]
TrajMixer 阶段:
\[
N=\operatorname{LN}_{d}(X),
\]
\[
G=\operatorname{reshape}(N)
\in\mathbb{R}^{B\times L\times n_{\mathrm{group}}\times d_{\mathrm{group}}},
\]
\[
P=\operatorname{IntraMixer}(G),
\]
\[
U=G+\sigma(\Theta)\odot P,
\]
\[
\Delta=\operatorname{reshape}
\left(\operatorname{CrossMixer}(U)\right),
\]
\[
X^{(l+1)}=X+\operatorname{Dropout}(\Delta).
\]
整个 TrajMixer 只有最后一次 `X + update` 是 residual。`U=G+\sigma(\Theta)\odot P` 是 update 分支内部的静态门控特征融合,不是相对于主 residual stream 的独立 residual stage。
## 3. Group 定义
Attention 输出经过 \(W_O\) 后仍是标准 residual representation
\[
X\in\mathbb{R}^{B\times L\times d}.
\]
定义:
\[
n_{\mathrm{group}}:=n_{\mathrm{head}},
\qquad
d_{\mathrm{group}}=\frac{d}{n_{\mathrm{group}}},
\qquad
d=n_{\mathrm{group}}d_{\mathrm{group}}.
\]
默认:
\[
d=120,\qquad
n_{\mathrm{group}}=10,\qquad
d_{\mathrm{group}}=12.
\]
这些 group 是 residual space 的连续分区,不等同于 Attention heads二者只共享数量。
## 4. 唯一的 Full-Width PreNorm
TrajMixer 只使用一个:
```text
norm: LayerNorm(n_embd)
```
LayerNorm 作用于完整 \(d\) 维 residual representation然后才 reshape
\[
G=\operatorname{reshape}
\left(\operatorname{LN}_{d}(X)\right).
\]
本版本明确删除:
```text
intra_norm
cross_norm
group_align
```
不得在组内或跨组阶段再增加额外 LayerNorm。
## 5. 组内 SwiGLU
每个 group 使用独立参数,对其 \(d_{\mathrm{group}}\) 维内部特征执行:
\[
d_{\mathrm{group}}
\rightarrow
4d_{\mathrm{group}}
\rightarrow
d_{\mathrm{group}}.
\]
对 group \(g\)
\[
W_{g,\mathrm{intra}}^{(g)},
W_{v,\mathrm{intra}}^{(g)}
\in
\mathbb{R}^{d_{\mathrm{group}}\times4d_{\mathrm{group}}},
\]
\[
W_{o,\mathrm{intra}}^{(g)}
\in
\mathbb{R}^{4d_{\mathrm{group}}\times d_{\mathrm{group}}}.
\]
计算:
\[
H_g
=
\operatorname{SiLU}
\left(G_gW_{g,\mathrm{intra}}^{(g)}\right)
\odot
\left(G_gW_{v,\mathrm{intra}}^{(g)}\right),
\]
\[
P_g=H_gW_{o,\mathrm{intra}}^{(g)}.
\]
实现形状:
```text
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。
## 6. 静态门控融合
定义可学习 gate logits
\[
\Theta\in
\mathbb{R}^{n_{\mathrm{group}}\times d_{\mathrm{group}}}.
\]
实际门值为:
\[
\Gamma=\sigma(\Theta).
\]
初始化:
\[
\Theta_{g,r}
=\operatorname{logit}(0.1)
=\log\frac{0.1}{0.9}
\approx-2.1972,
\]
因此:
\[
\Gamma_{g,r}\approx0.1.
\]
融合:
\[
U=G+\Gamma\odot P.
\]
\(\Gamma\) 对 batch 和序列位置共享,但每个 group、每个内部坐标拥有独立可学习值。
## 7. 跨 Group SwiGLU
对于每个内部坐标 \(r\),独立沿 group 维度执行:
\[
n_{\mathrm{group}}
\rightarrow
4n_{\mathrm{group}}
\rightarrow
n_{\mathrm{group}}.
\]
定义:
\[
A_g^{(r)},A_v^{(r)}
\in
\mathbb{R}^{n_{\mathrm{group}}\times4n_{\mathrm{group}}},
\]
\[
A_o^{(r)}
\in
\mathbb{R}^{4n_{\mathrm{group}}\times n_{\mathrm{group}}}.
\]
计算:
\[
Q_{:,r}
=
\operatorname{SiLU}\left(U_{:,r}A_g^{(r)}\right)
\odot
\left(U_{:,r}A_v^{(r)}\right),
\]
\[
\Delta_{:,r}=Q_{:,r}A_o^{(r)}.
\]
实现形状:
```text
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。不同内部坐标拥有独立的跨 group 参数,且不沿序列维度交互。
## 8. 唯一的外层 Residual
跨 group 输出 reshape 回:
\[
\Delta\in\mathbb{R}^{B\times L\times d}.
\]
最终:
\[
\operatorname{TrajMixer}(X)
=X+\operatorname{Dropout}(\Delta).
\]
固定约束:
- 组内阶段后不执行独立 residual
- 跨组阶段后不执行独立 residual
- `GPTBlock` 不再额外执行 `X + TrajMixer(X)`
- 整个 TrajMixer 只有一次主 residual。
## 9. 初始化
固定初始化:
- `intra_gate_proj/intra_value_proj`:每个 group 独立 Xavier uniform
- `intra_output_proj`:每个 group 独立 Xavier uniform
- `intra_gate_logits`:初始化为 \(\operatorname{logit}(0.1)\)
- 跨组 `gate_proj/value_proj`:每个内部坐标独立 Xavier uniform
- 最终跨组 `output_proj`:均值 0、标准差 \(10^{-3}\) 的正态分布;
- Full-width LayerNormPyTorch 默认 affine 初始化;
- Dropout沿用 `mlp_dropout`
组内输出使用正常 Xavier 初始化以保证其具有完整表达能力;静态门控将其初始贡献限制在约 0.1。最终跨 group 输出投影保持小值初始化,使整个 TrajMixer residual update 在训练初期接近零。
Relative Time Attention Bias 初始化固定为:
- `rbf_proj.weight`:零初始化;
- `time_bias_scale`:初始化为 \(1.0\)
- 初始 RBF attention bias 严格为零;
- `rbf_proj.weight` 从第一个优化步骤即可获得梯度。
## 10. 参数量
默认 \(d=120\)、\(n_{\mathrm{group}}=10\)、\(d_{\mathrm{group}}=12\)。
Full-width LayerNorm
\[
2d=240.
\]
组内 projections
\[
3n_{\mathrm{group}}d_{\mathrm{group}}
\left(4d_{\mathrm{group}}\right)
=17{,}280.
\]
静态门控:
\[
n_{\mathrm{group}}d_{\mathrm{group}}
=120.
\]
跨 group projections
\[
3d_{\mathrm{group}}n_{\mathrm{group}}
\left(4n_{\mathrm{group}}\right)
=14{,}400.
\]
每层 TrajMixer 合计:
\[
240+17{,}280+120+14{,}400
=\boxed{32{,}040}.
\]
默认 relative-time、12 层、`vocab_size=1256`、无额外信息类型时,完整模型参数量为:
\[
\boxed{1{,}232{,}428}.
\]
## 11. 固定配置与 checkpoint 约束
```yaml
model_architecture: traj_mixer_v5
d_model: 120
n_head: 10
n_group_rule: n_head
d_group_rule: d_model / n_group
traj_mixer_norm: layer_norm_over_n_embd
intra_hidden_rule: 4 * d_group
intra_gate_shape: [n_group, d_group]
intra_gate_initial_sigmoid: 0.1
cross_hidden_rule: 4 * n_group
group_alignment: false
intra_residual: false
cross_residual: false
traj_mixer_outer_residual: true
projection_bias: false
intra_output_init: xavier_uniform
cross_output_init_std: 0.001
```
训练时必须将 `model_architecture: traj_mixer_v5``model_parameter_count``trainable_parameter_count` 写入 `train_config.json`,并在日志中打印参数量。
评估和导出入口只接受 `traj_mixer_v5` checkpoint并检查 Full-width LayerNorm、组内 projections、静态门控和跨 group projections 是否齐全。`traj_mixer_v4` 及更早 checkpoint 不向后兼容,直接拒绝加载。