Files
DeepHealth/TrajMixer_设计方案.md

7.5 KiB
Raw Blame History

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 结构

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 只使用一个:

norm: LayerNorm(n_embd)

LayerNorm 作用于完整 (d) 维 residual representation然后才 reshape

[ G=\operatorname{reshape} \left(\operatorname{LN}_{d}(X)\right). ]

本版本明确删除:

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)}. ]

实现形状:

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)}. ]

实现形状:

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 约束

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_v5model_parameter_counttrainable_parameter_count 写入 train_config.json,并在日志中打印参数量。

评估和导出入口只接受 traj_mixer_v5 checkpoint并检查 Full-width LayerNorm、组内 projections、静态门控和跨 group projections 是否齐全。traj_mixer_v4 及更早 checkpoint 不向后兼容,直接拒绝加载。