2026-07-23 16:05:00 +08:00
|
|
|
|
# TrajMixer Block 最终设计方案
|
|
|
|
|
|
|
|
|
|
|
|
> 状态:**Frozen implementation baseline**
|
|
|
|
|
|
>
|
2026-07-24 10:51:41 +08:00
|
|
|
|
> 版本:**v2.0 / traj_mixer_v3**
|
2026-07-23 16:05:00 +08:00
|
|
|
|
>
|
2026-07-24 10:51:41 +08:00
|
|
|
|
> 固化日期:**2026-07-24**
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
本文档是 TrajMixer 后续实现与实验的结构基线。本版本将原先仅含跨轨迹交互的 TrajMixer 扩展为“组内混合 + 跨组混合”的两阶段结构。
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
|
|
|
|
|
## 1. 目标
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
在保持原始 Delphi Transformer Attention 结构不变的前提下,用轻量、可并行的两阶段 TrajMixer 替换 FFN。
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
|
|
|
|
|
保持不变的组件包括:
|
|
|
|
|
|
|
|
|
|
|
|
- 原始 causal mask;
|
|
|
|
|
|
- 原始 TimeRoPE / Relative Time Attention Bias;
|
|
|
|
|
|
- 原始 Multi-Head Attention,包括 \(W_Q/W_K/W_V/W_O\);
|
2026-07-24 10:51:41 +08:00
|
|
|
|
- 原始序列建模和训练目标。
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
TrajMixer 不沿序列维度混合,也不引入时间递归。
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
|
|
|
|
|
## 2. Block 总体结构
|
|
|
|
|
|
|
|
|
|
|
|
```text
|
|
|
|
|
|
PreNorm Causal Multi-Head Attention
|
2026-07-24 10:51:41 +08:00
|
|
|
|
→ 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
|
2026-07-23 16:05:00 +08:00
|
|
|
|
→ Group-wise Feature Alignment
|
2026-07-24 10:51:41 +08:00
|
|
|
|
→ Cross-Group SwiGLU: n_group → 4n_group → n_group
|
|
|
|
|
|
→ Cross-Group Residual
|
|
|
|
|
|
→ reshape [B, L, d]
|
2026-07-23 16:05:00 +08:00
|
|
|
|
```
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
Attention 阶段保持原样:
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
|
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
U=X^{(l)}+\operatorname{CausalMHA}
|
|
|
|
|
|
\left(\operatorname{LN}_{\mathrm{attn}}(X^{(l)})\right).
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
随后:
|
|
|
|
|
|
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
H^{(0)}
|
|
|
|
|
|
=\operatorname{reshape}(U)
|
|
|
|
|
|
\in\mathbb{R}^{B\times L\times G\times D},
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
其中 \(G=n_{\mathrm{group}}\),\(D=d_{\mathrm{group}}\)。
|
|
|
|
|
|
|
|
|
|
|
|
两阶段 TrajMixer 为:
|
|
|
|
|
|
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
H^{(1)}
|
|
|
|
|
|
=H^{(0)}
|
|
|
|
|
|
+\operatorname{Dropout}
|
|
|
|
|
|
\left(\operatorname{IntraMixer}
|
|
|
|
|
|
\left(\operatorname{LN}_{D}(H^{(0)})\right)\right),
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
H^{(2)}
|
|
|
|
|
|
=H^{(1)}
|
|
|
|
|
|
+\operatorname{Dropout}
|
|
|
|
|
|
\left(\operatorname{CrossMixer}
|
|
|
|
|
|
\left(\operatorname{LN}_{G}(H^{(1)})\right)\right),
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
\[
|
|
|
|
|
|
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)`。
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
|
|
|
|
|
## 3. Latent Trajectory Group 定义
|
|
|
|
|
|
|
|
|
|
|
|
Attention 输出经过 \(W_O\) 后仍是标准 residual representation:
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
U\in\mathbb{R}^{B\times L\times d}.
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
固定:
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
|
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
G:=n_{\mathrm{group}}=n_{\mathrm{head}},
|
|
|
|
|
|
\qquad
|
|
|
|
|
|
D:=d_{\mathrm{group}}=\frac{d}{G},
|
|
|
|
|
|
\qquad
|
|
|
|
|
|
d=GD.
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
默认配置:
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
|
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
d=120,\qquad G=10,\qquad D=12.
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
reshape 后:
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
|
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
H^{(0)}\in\mathbb{R}^{B\times L\times G\times D}.
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
`n_group` 由 `n_head` 决定,但 residual groups 只是 residual space 的连续分区,不等同于 Attention heads。
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
## 4. 第一阶段:组内 SwiGLU Mixer
|
|
|
|
|
|
|
|
|
|
|
|
第一阶段对每个 group 独立进行特征变换。不同 group 使用各自的投影参数,不发生 group 间信息交换。
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
先对每个 \((b,t,g)\) 的 \(D\) 维向量独立执行 LayerNorm:
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
|
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
\widetilde H^{(0)}
|
|
|
|
|
|
=\operatorname{LN}_{D}(H^{(0)}).
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
归一化统计量在每个 group 内独立计算;为保持轻量,LayerNorm 的 affine 参数在各 group 间共享。
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
对于 \(g=1,\ldots,G\),定义:
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
|
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
W_{g,\mathrm{intra}}^{(g)},
|
|
|
|
|
|
W_{v,\mathrm{intra}}^{(g)}
|
|
|
|
|
|
\in\mathbb{R}^{D\times 4D},
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
\[
|
|
|
|
|
|
W_{o,\mathrm{intra}}^{(g)}
|
|
|
|
|
|
\in\mathbb{R}^{4D\times D}.
|
|
|
|
|
|
\]
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
计算:
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
|
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
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),
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
\Delta_{\mathrm{intra},g}
|
|
|
|
|
|
=P_gW_{o,\mathrm{intra}}^{(g)},
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
H^{(1)}
|
|
|
|
|
|
=H^{(0)}
|
|
|
|
|
|
+\operatorname{Dropout}(\Delta_{\mathrm{intra}}).
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
该阶段完成:
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
|
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
D\rightarrow4D\rightarrow D,
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
用于增强每条潜在轨迹内部的非线性特征组合能力。
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
实现张量形状:
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
```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]
|
|
|
|
|
|
```
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
三个 projection 均不带 bias。
|
|
|
|
|
|
|
|
|
|
|
|
## 5. 第二阶段:跨组 TrajMixer
|
|
|
|
|
|
|
|
|
|
|
|
第二阶段沿 group 维度进行交互。对于每个内部坐标 \(r\),独立执行:
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
|
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
G\rightarrow4G\rightarrow G.
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
首先将 \(H^{(1)}\) 的最后两个维度交换,并在 group 维度执行 LayerNorm:
|
|
|
|
|
|
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
\widetilde H^{(1)}_{b,t,:,r}
|
|
|
|
|
|
=\operatorname{LN}_{G}
|
|
|
|
|
|
\left(H^{(1)}_{b,t,:,r}\right).
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
归一化统计量对每个内部坐标 \(r\) 独立计算;LayerNorm 的 affine 参数在各内部坐标间共享。
|
|
|
|
|
|
|
|
|
|
|
|
### 5.1 Group-wise Feature Alignment
|
|
|
|
|
|
|
|
|
|
|
|
沿用现有的可学习 group 特征对齐矩阵:
|
|
|
|
|
|
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
B_g\in\mathbb{R}^{D\times D},
|
|
|
|
|
|
\qquad g=1,\ldots,G,
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
Z_{b,t,g,:}
|
|
|
|
|
|
=\widetilde H^{(1)}_{b,t,g,:}B_g.
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
\(B_g\) 不带 bias,并使用单位矩阵初始化。
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
### 5.2 Cross-Group SwiGLU
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
对每个内部坐标 \(r=1,\ldots,D\),定义:
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
|
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
A_g^{(r)},A_v^{(r)}
|
|
|
|
|
|
\in\mathbb{R}^{G\times4G},
|
|
|
|
|
|
\qquad
|
|
|
|
|
|
A_o^{(r)}
|
|
|
|
|
|
\in\mathbb{R}^{4G\times G}.
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
计算:
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
|
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
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),
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
\[
|
|
|
|
|
|
\Delta_{\mathrm{cross},b,t,:,r}
|
|
|
|
|
|
=Q_{b,t,:,r}A_o^{(r)},
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
H^{(2)}
|
|
|
|
|
|
=H^{(1)}
|
|
|
|
|
|
+\operatorname{Dropout}(\Delta_{\mathrm{cross}}).
|
|
|
|
|
|
\]
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
实现张量形状:
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
|
|
|
|
|
```text
|
2026-07-24 10:51:41 +08:00
|
|
|
|
cross_norm: LayerNorm(n_group)
|
2026-07-23 16:05:00 +08:00
|
|
|
|
group_align: [n_group, d_group, d_group]
|
2026-07-24 10:51:41 +08:00
|
|
|
|
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]
|
2026-07-23 16:05:00 +08:00
|
|
|
|
```
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
三个 projection 均不带 bias。
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
## 6. PreNorm 与残差约束
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
本版本固定使用两个独立的 PreNorm residual stage:
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
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. 参数量
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
默认 \(d=120,G=10,D=12\)。
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
### 7.1 组内阶段
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
投影权重:
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
|
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
3G D(4D)
|
|
|
|
|
|
=12GD^2
|
|
|
|
|
|
=17{,}280.
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
`LayerNorm(D)`:
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
|
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
2D=24.
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
### 7.2 跨组阶段
|
|
|
|
|
|
|
|
|
|
|
|
跨组投影权重:
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
|
|
|
|
|
\[
|
2026-07-24 10:51:41 +08:00
|
|
|
|
3D G(4G)
|
|
|
|
|
|
=12DG^2
|
|
|
|
|
|
=14{,}400.
|
2026-07-23 16:05:00 +08:00
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
Group Feature Alignment:
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
\[
|
|
|
|
|
|
GD^2
|
|
|
|
|
|
=1{,}440.
|
|
|
|
|
|
\]
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
`LayerNorm(G)`:
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
\[
|
|
|
|
|
|
2G=20.
|
|
|
|
|
|
\]
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
### 7.3 每个 TrajMixer 合计
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
\[
|
|
|
|
|
|
17{,}280+24+14{,}400+1{,}440+20
|
|
|
|
|
|
=\boxed{33{,}164}.
|
|
|
|
|
|
\]
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
相对于 `traj_mixer_v2` 的跨组单阶段结构 \(15{,}840\),每层增加 \(17{,}324\) 个参数。作为历史实现对照,代码库原全维度 SwiGLU FFN 每层为 \(108{,}720\) 个参数。
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
## 8. 初始化
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
固定初始化约定:
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
- 组内 `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`。
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
两个 output projection 的小方差初始化使两阶段在训练初期都接近恒等 residual update。
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 14:39:15 +08:00
|
|
|
|
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 时间偏置将无法开始学习。
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
## 9. 信息流与语义
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
**Attention**:从历史疾病事件中选择和整合相关信息。
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
**Intra-Group Mixer**:学习每条潜在轨迹内部的非线性特征组合。
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
**Group Feature Alignment**:对齐不同潜在轨迹的内部坐标。
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
**Cross-Group Mixer**:学习不同潜在轨迹在相同内部坐标上的门控交互。
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
|
|
|
|
|
整个模块保持:
|
|
|
|
|
|
|
|
|
|
|
|
- 无时间递归;
|
2026-07-24 10:51:41 +08:00
|
|
|
|
- 不混合序列位置;
|
2026-07-23 16:05:00 +08:00
|
|
|
|
- 序列维度完全并行;
|
2026-07-24 10:51:41 +08:00
|
|
|
|
- 保留原始因果 Attention;
|
|
|
|
|
|
- residual groups 不等同于 Attention heads。
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
## 10. 固定配置与 checkpoint 约束
|
2026-07-23 16:05:00 +08:00
|
|
|
|
|
|
|
|
|
|
```yaml
|
2026-07-24 10:51:41 +08:00
|
|
|
|
model_architecture: traj_mixer_v3
|
2026-07-23 16:05:00 +08:00
|
|
|
|
d_model: 120
|
2026-07-24 10:51:41 +08:00
|
|
|
|
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
|
2026-07-23 16:05:00 +08:00
|
|
|
|
attention: unchanged
|
2026-07-24 10:51:41 +08:00
|
|
|
|
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
|
2026-07-23 16:05:00 +08:00
|
|
|
|
gate_value_init: xavier_uniform
|
|
|
|
|
|
output_init_std: 0.001
|
|
|
|
|
|
```
|
|
|
|
|
|
|
|
|
|
|
|
必须满足:
|
|
|
|
|
|
|
|
|
|
|
|
\[
|
|
|
|
|
|
d=n_{\mathrm{group}}d_{\mathrm{group}}.
|
|
|
|
|
|
\]
|
|
|
|
|
|
|
2026-07-24 10:51:41 +08:00
|
|
|
|
训练时必须将 `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 不向后兼容,直接拒绝加载。
|