Revert "Implement shared event-trajectory reasoning backbone"
This reverts commit 06f29c0f0a.
This commit is contained in:
@@ -1,386 +0,0 @@
|
||||
# Event–Trajectory Shared Reasoning Backbone
|
||||
|
||||
> 状态:**Frozen implementation baseline**
|
||||
> 架构标识:`event_trajectory_shared_v1`
|
||||
> 固化日期:**2026-07-23**
|
||||
|
||||
## 1. 核心定义
|
||||
|
||||
使用一个共享的 Attention–TrajMixer 推理核心,对固定 Event Memory 进行多轮读取,并持续更新 Trajectory State。
|
||||
|
||||
模型只实例化:
|
||||
|
||||
```python
|
||||
self.reasoning_core = SharedEventTrajectoryCore(...)
|
||||
```
|
||||
|
||||
禁止为不同推理轮创建独立 Transformer blocks。参数只保存一套,计算上顺序运行多轮。
|
||||
|
||||
模型规模固定为四档:
|
||||
|
||||
| model_size | d_model | n_trajectory | trajectory_dim | traj_hidden |
|
||||
|---|---:|---:|---:|---:|
|
||||
| nano | 256 | 8 | 32 | 32 |
|
||||
| small | 512 | 8 | 64 | 32 |
|
||||
| medium | 768 | 12 | 64 | 48 |
|
||||
| huge | 1024 | 16 | 64 | 64 |
|
||||
|
||||
默认使用 `model_size=nano`。`n_reasoning_rounds` 是独立参数,默认值为12,
|
||||
不属于模型规模预设;任意模型规模均可单独指定推理轮数。
|
||||
|
||||
必须满足:
|
||||
|
||||
\[
|
||||
d_{\mathrm{model}}
|
||||
=n_{\mathrm{trajectory}}d_{\mathrm{trajectory}}.
|
||||
\]
|
||||
|
||||
## 2. 固定 Event Memory
|
||||
|
||||
疾病事件与可选协变量首先组成事件序列:
|
||||
|
||||
\[
|
||||
X_E\in\mathbb{R}^{B\times L\times d_{\mathrm{model}}}.
|
||||
\]
|
||||
|
||||
事件特征由以下信息相加:
|
||||
|
||||
```text
|
||||
disease / covariate embedding
|
||||
+ age/time encoding
|
||||
+ sex context
|
||||
```
|
||||
|
||||
然后只编码一次:
|
||||
|
||||
\[
|
||||
E=\operatorname{EventNorm}
|
||||
\left(\operatorname{EventProjection}(X_E)\right).
|
||||
\]
|
||||
|
||||
进入 reasoning loop 后,\(E\) 的数值保持不变,但不执行 `detach`,梯度仍可回传到事件编码器。
|
||||
|
||||
Key 和 Value 同样每次 forward 只投影一次:
|
||||
|
||||
```python
|
||||
event_key_value = reasoning_core.project_event_memory(E)
|
||||
```
|
||||
|
||||
12 轮共享并复用该结果。
|
||||
|
||||
## 3. Trajectory State
|
||||
|
||||
每个查询维护 `n_trajectory` 个显式 trajectory slots;nano 默认使用8个:
|
||||
|
||||
\[
|
||||
S\in\mathbb{R}^{B\times Q\times8\times32}.
|
||||
\]
|
||||
|
||||
其中:
|
||||
|
||||
- all-future:\(Q=1\);
|
||||
- next-token:\(Q=L\),所有查询位置并行计算。
|
||||
|
||||
定义可学习原型:
|
||||
|
||||
\[
|
||||
P\in\mathbb{R}^{8\times32}.
|
||||
\]
|
||||
|
||||
查询上下文经过投影并 reshape:
|
||||
|
||||
\[
|
||||
C_Q
|
||||
=\operatorname{QueryProjection}(\text{query features})
|
||||
\in\mathbb{R}^{B\times Q\times8\times32},
|
||||
\]
|
||||
|
||||
\[
|
||||
S^{(0)}=P+C_Q.
|
||||
\]
|
||||
|
||||
all-future 的 query features 包含可学习 query token、查询年龄和性别;next-token 的 query features 使用当前位置的事件、时间和性别表示,以保留 token-level 预测语义。
|
||||
|
||||
## 4. 共享 Trajectory-to-Event Attention
|
||||
|
||||
Trajectory State 作为 Query,固定 Event Memory 作为 Key 和 Value:
|
||||
|
||||
\[
|
||||
Q^{(r)}
|
||||
=W_Q\operatorname{LN}_{A}(S^{(r)}),
|
||||
\]
|
||||
|
||||
\[
|
||||
K=W_KE,\qquad V=W_VE.
|
||||
\]
|
||||
|
||||
形状为:
|
||||
|
||||
```text
|
||||
Q: [B, query, trajectory, trajectory_dim]
|
||||
K: [B, trajectory, event, trajectory_dim]
|
||||
V: [B, trajectory, event, trajectory_dim]
|
||||
```
|
||||
|
||||
Attention:
|
||||
|
||||
\[
|
||||
\operatorname{score}_{b,q,h,l}
|
||||
=
|
||||
\frac{
|
||||
\left\langle Q_{b,q,h,:},K_{b,h,l,:}\right\rangle
|
||||
}{
|
||||
\sqrt{d_{\mathrm{trajectory}}}
|
||||
}.
|
||||
\]
|
||||
|
||||
每个 trajectory slot 独立读取整段 Event Memory。Attention 不包含跨 trajectory 的完整输出投影;trajectory 之间的交互只由后续 TrajMixer 完成。
|
||||
|
||||
外部 `padding_mask` 的语义固定为 `True = valid`。
|
||||
|
||||
all-future 的内部 mask 必须满足:
|
||||
|
||||
\[
|
||||
\operatorname{valid}_{b,q,l}
|
||||
=
|
||||
\operatorname{eventValid}_{b,l}
|
||||
\land
|
||||
(t_l\le t_q).
|
||||
\]
|
||||
|
||||
next-token 还必须对相同时间戳加入位置因果约束:
|
||||
|
||||
\[
|
||||
\operatorname{valid}_{b,q,l}
|
||||
=
|
||||
\operatorname{eventValid}_{b,l}
|
||||
\land
|
||||
\left[
|
||||
(t_l<t_q)
|
||||
\lor
|
||||
\left((t_l=t_q)\land(l\le q)\right)
|
||||
\right].
|
||||
\]
|
||||
|
||||
这可以阻止前一个 token 直接读到同时间的后续目标 token;它只改变并行
|
||||
Attention 的可见性矩阵,不沿疾病时间轴递归。
|
||||
|
||||
两种 mask 均不能读取未来事件或协变量。
|
||||
|
||||
全 masked query 的 Attention readout 必须显式返回零,不能产生 NaN。
|
||||
|
||||
## 5. 时间信息
|
||||
|
||||
所有模式都在 Event Memory 和 query context 中加入 age/time encoding。
|
||||
|
||||
当 `time_mode=relative` 时,shared cross-attention 额外使用:
|
||||
|
||||
- query-time 对 event-time 的 Cross-TimeRoPE;
|
||||
- query–event 时间差的 Gaussian RBF bias。
|
||||
|
||||
对应缓存形状为:
|
||||
|
||||
\[
|
||||
\text{RBF cache}\in\mathbb{R}^{B\times Q\times L\times n_{\mathrm{rbf}}}.
|
||||
\]
|
||||
|
||||
## 6. 共享 TrajMixer
|
||||
|
||||
TrajMixer 输入:
|
||||
|
||||
\[
|
||||
U\in\mathbb{R}^{B\times Q\times H\times D_h}.
|
||||
\]
|
||||
|
||||
它只沿 trajectory 轴交互,不使用普通全维度 FFN,也不包含旧版 Group Alignment。
|
||||
|
||||
对每个内部坐标 \(r\):
|
||||
|
||||
\[
|
||||
G_r=U_rW_g^{(r)},\qquad
|
||||
V_r=U_rW_v^{(r)},
|
||||
\]
|
||||
|
||||
\[
|
||||
M_r
|
||||
=
|
||||
\operatorname{SiLU}(G_r)\odot V_r,
|
||||
\]
|
||||
|
||||
\[
|
||||
Y_r=M_rW_o^{(r)}.
|
||||
\]
|
||||
|
||||
参数形状:
|
||||
|
||||
```text
|
||||
W_g: [trajectory_dim, n_trajectory, traj_hidden]
|
||||
W_v: [trajectory_dim, n_trajectory, traj_hidden]
|
||||
W_o: [trajectory_dim, traj_hidden, n_trajectory]
|
||||
```
|
||||
|
||||
默认:
|
||||
|
||||
```text
|
||||
n_trajectory -> 4 * n_trajectory -> n_trajectory
|
||||
```
|
||||
|
||||
其中:
|
||||
|
||||
\[
|
||||
d_{\mathrm{trajHidden}}=4n_{\mathrm{trajectory}}.
|
||||
\]
|
||||
|
||||
## 7. 单轮共享核心
|
||||
|
||||
单轮计算:
|
||||
|
||||
\[
|
||||
R^{(r)}
|
||||
=
|
||||
A_\theta\left(
|
||||
\operatorname{LN}_A(S^{(r)}),E
|
||||
\right),
|
||||
\]
|
||||
|
||||
\[
|
||||
U^{(r)}
|
||||
=
|
||||
S^{(r)}
|
||||
+\alpha_A\operatorname{Dropout}(R^{(r)}),
|
||||
\]
|
||||
|
||||
\[
|
||||
S^{(r+1)}
|
||||
=
|
||||
U^{(r)}
|
||||
+\alpha_M\operatorname{Dropout}
|
||||
\left(
|
||||
M_\phi(\operatorname{LN}_M(U^{(r)}))
|
||||
\right).
|
||||
\]
|
||||
|
||||
残差统一使用加法。
|
||||
|
||||
## 8. 多轮参数共享
|
||||
|
||||
同一个核心重复运行:
|
||||
|
||||
```python
|
||||
for _ in range(n_reasoning_rounds):
|
||||
S = self.reasoning_core(...)
|
||||
```
|
||||
|
||||
所有轮次共享:
|
||||
|
||||
```text
|
||||
q_proj / k_proj / v_proj
|
||||
relative-time projection
|
||||
TrajMixer parameters
|
||||
LayerNorm parameters
|
||||
attn_scale / mixer_scale
|
||||
```
|
||||
|
||||
因此推理轮数不改变模型参数量:
|
||||
|
||||
\[
|
||||
A^{(1)}=\cdots=A^{(12)}=A_\theta,
|
||||
\]
|
||||
|
||||
\[
|
||||
M^{(1)}=\cdots=M^{(12)}=M_\phi.
|
||||
\]
|
||||
|
||||
但每轮 state 不同,因此 Query 与 Attention weights 也不同。
|
||||
|
||||
## 9. 稳定性设计
|
||||
|
||||
共享残差缩放初始化为:
|
||||
|
||||
\[
|
||||
\alpha_A=\alpha_M
|
||||
=
|
||||
\frac{1}{\sqrt{n_{\mathrm{reasoningRounds}}}}.
|
||||
\]
|
||||
|
||||
两个标量可学习,并由全部轮次共享。
|
||||
|
||||
第一版不加入:
|
||||
|
||||
```text
|
||||
round-specific parameters
|
||||
round embedding
|
||||
每轮独立 LayerNorm
|
||||
每轮独立 residual scale
|
||||
GRU 或其他时间递归
|
||||
```
|
||||
|
||||
## 10. 输出接口
|
||||
|
||||
推理结束后按固定顺序 flatten trajectory slots:
|
||||
|
||||
\[
|
||||
H
|
||||
=
|
||||
\operatorname{FinalNorm}
|
||||
\left(
|
||||
\operatorname{Flatten}(S^{(R)})
|
||||
\right).
|
||||
\]
|
||||
|
||||
- all-future 输出:`[B, d_model]`;
|
||||
- next-token 输出:`[B, L, d_model]`;
|
||||
- next-token 的 risk-head weight tying 保持不变;
|
||||
- Weibull 与 mixed heads 继续使用同一最终 hidden。
|
||||
|
||||
next-token 的 query 位置全部并行,只有 reasoning rounds 顺序执行,因此不存在沿疾病时间轴的状态递归。
|
||||
|
||||
## 11. 配置与 checkpoint 约束
|
||||
|
||||
训练配置必须写入:
|
||||
|
||||
```yaml
|
||||
model_architecture: event_trajectory_shared_v1
|
||||
model_size: nano
|
||||
d_model: 256
|
||||
n_trajectory: 8
|
||||
trajectory_dim: 32
|
||||
traj_hidden: 32
|
||||
n_reasoning_rounds: 12
|
||||
model_parameter_count: <runtime count>
|
||||
trainable_parameter_count: <runtime count>
|
||||
```
|
||||
|
||||
评估和导出入口必须同时验证:
|
||||
|
||||
1. `model_architecture` 完全匹配;
|
||||
2. `model_size` 属于 `nano / small / medium / huge`;
|
||||
3. `d_model`、`n_trajectory`、`trajectory_dim` 和 `traj_hidden`
|
||||
与对应规模预设完全匹配;
|
||||
4. checkpoint 包含一套且仅一套 `reasoning_core` 关键参数;
|
||||
5. checkpoint 内持久化的 `d_model`、`n_trajectory` 和
|
||||
`n_reasoning_rounds` 架构指纹与训练配置完全一致;
|
||||
6. 不接受旧 `traj_mixer_v2` checkpoint。
|
||||
|
||||
其中 `n_reasoning_rounds` 必须进入 checkpoint 架构指纹,因为改变轮数
|
||||
不会改变参数 shape,不能仅依赖 `load_state_dict(strict=True)` 检出错配。
|
||||
|
||||
## 12. 信息流
|
||||
|
||||
```text
|
||||
E ─────────────┬──────────────┬──────────────┬──────────────┐
|
||||
│ │ │ │
|
||||
▼ ▼ ▼ ▼
|
||||
S0 -> Shared Core -> S1 -> Shared Core -> S2 -> ... -> Shared Core -> S12
|
||||
同一套参数 同一套参数 同一套参数
|
||||
```
|
||||
|
||||
整体定义:
|
||||
|
||||
\[
|
||||
\boxed{
|
||||
\text{一个共享 Event–Trajectory 推理核心}
|
||||
\times
|
||||
\text{多轮状态依赖推理}
|
||||
}
|
||||
\]
|
||||
333
TrajMixer_设计方案.md
Normal file
333
TrajMixer_设计方案.md
Normal file
@@ -0,0 +1,333 @@
|
||||
# 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}}}.
|
||||
\]
|
||||
|
||||
隐藏宽度不再独立配置,固定为:
|
||||
|
||||
\[
|
||||
h_{\mathrm{group}}=4n_{\mathrm{head}}
|
||||
=4n_{\mathrm{group}}.
|
||||
\]
|
||||
|
||||
当前 \(n_{\mathrm{head}}=10\),因此 \(h_{\mathrm{group}}=40\)。
|
||||
|
||||
对固定的 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 数量;
|
||||
- hidden group 表示固定扩展为 group 数量的 4 倍。
|
||||
|
||||
所有 \(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}}.
|
||||
\]
|
||||
|
||||
固定 \(h_{\mathrm{group}}=4n_{\mathrm{group}}=40\) 时,Mixer 每层权重参数量为:
|
||||
|
||||
\[
|
||||
3d_{\mathrm{group}}n_{\mathrm{group}}h_{\mathrm{group}}
|
||||
=3\times12\times10\times40
|
||||
=14{,}400.
|
||||
\]
|
||||
|
||||
加上 Group Feature Alignment 后,TrajMixer residual branch 每层共有:
|
||||
|
||||
\[
|
||||
14{,}400+1{,}440=15{,}840
|
||||
\]
|
||||
|
||||
个主要权重参数。作为对照,原始 \(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
|
||||
model_architecture: traj_mixer_v2
|
||||
d_model: 120
|
||||
n_head: 10 # 同时决定 residual group 数量
|
||||
d_group: 12
|
||||
hidden_group_rule: 4 * n_head # 不单独配置
|
||||
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
|
||||
```
|
||||
|
||||
训练时必须将 `model_architecture: traj_mixer_v2`、`model_parameter_count` 和 `trainable_parameter_count` 写入 `train_config.json`,并在训练日志中显式打印总参数量与可训练参数量。本分支的评估和导出入口只接受带有该架构标识、且 checkpoint 中包含 TrajMixer 参数张量的模型;其他版本或分支生成的模型应直接拒绝加载。
|
||||
|
||||
必须满足:
|
||||
|
||||
\[
|
||||
d=n_{\mathrm{group}}d_{\mathrm{group}}.
|
||||
\]
|
||||
|
||||
后续实现、单元测试、参数量核验和主实验均以以上配置为默认基线。
|
||||
390
backbones.py
390
backbones.py
@@ -29,14 +29,6 @@ class TimeRoPE(nn.Module):
|
||||
x2 = x[..., 1::2]
|
||||
return torch.stack((-x2, x1), dim=-1).flatten(-2)
|
||||
|
||||
@staticmethod
|
||||
def apply_single_from_cache(
|
||||
x: torch.Tensor,
|
||||
rope_cache: tuple[torch.Tensor, torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
cos, sin = rope_cache
|
||||
return x * cos + TimeRoPE._rotate_half(x) * sin
|
||||
|
||||
@staticmethod
|
||||
def apply_from_cache(
|
||||
q: torch.Tensor,
|
||||
@@ -93,280 +85,232 @@ class GaussianRBFTimeBasis(nn.Module):
|
||||
)
|
||||
return rbf_acts
|
||||
|
||||
def precompute_cross_cache(
|
||||
self,
|
||||
query_tau: torch.Tensor,
|
||||
key_tau: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Return RBF activations for query-time minus event-time."""
|
||||
diff = query_tau.float().unsqueeze(2) - key_tau.float().unsqueeze(1)
|
||||
widths = self.log_widths.exp()
|
||||
return torch.exp(
|
||||
-0.5
|
||||
* (
|
||||
(diff.unsqueeze(-1) - self.centers)
|
||||
/ widths
|
||||
).square()
|
||||
)
|
||||
|
||||
|
||||
class TrajectoryCrossAttention(nn.Module):
|
||||
"""Shared trajectory-slot queries reading a fixed event memory."""
|
||||
|
||||
class TemporalAttention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
d_model: int,
|
||||
n_trajectory: int,
|
||||
n_embd: int,
|
||||
n_head: int,
|
||||
n_rbf_bases: int = 16,
|
||||
use_time_rope: bool = False,
|
||||
use_rbf_bias: bool = False,
|
||||
dropout: float = 0.0,
|
||||
use_time_rope: bool = True,
|
||||
use_rbf_bias: bool = True,
|
||||
):
|
||||
super().__init__()
|
||||
if d_model <= 0 or n_trajectory <= 0:
|
||||
raise ValueError("d_model and n_trajectory must be positive")
|
||||
if d_model % n_trajectory != 0:
|
||||
raise ValueError(
|
||||
"d_model must be divisible by n_trajectory, got "
|
||||
f"{d_model} and {n_trajectory}"
|
||||
)
|
||||
self.d_model = d_model
|
||||
self.n_trajectory = n_trajectory
|
||||
self.trajectory_dim = d_model // n_trajectory
|
||||
self.scale = self.trajectory_dim ** -0.5
|
||||
assert n_embd % n_head == 0, "n_embd must be divisible by n_head"
|
||||
self.n_head = n_head
|
||||
self.d_head = n_embd // n_head
|
||||
self.scale = 1.0 / math.sqrt(self.d_head)
|
||||
self.use_time_rope = use_time_rope
|
||||
self.use_rbf_bias = use_rbf_bias
|
||||
|
||||
# q_proj acts on each slot independently and is shared across slots.
|
||||
self.q_proj = nn.Linear(
|
||||
self.trajectory_dim,
|
||||
self.trajectory_dim,
|
||||
bias=False,
|
||||
)
|
||||
self.k_proj = nn.Linear(d_model, d_model, bias=False)
|
||||
self.v_proj = nn.Linear(d_model, d_model, bias=False)
|
||||
if use_rbf_bias:
|
||||
self.rbf_proj = nn.Linear(
|
||||
n_rbf_bases,
|
||||
n_trajectory,
|
||||
bias=False,
|
||||
)
|
||||
self.time_bias_scale = nn.Parameter(torch.tensor(0.0))
|
||||
else:
|
||||
self.rbf_proj = None
|
||||
self.register_parameter("time_bias_scale", None)
|
||||
# QKV projection (fused for efficiency)
|
||||
self.qkv = nn.Linear(n_embd, 3 * n_embd, bias=False)
|
||||
# Output projection
|
||||
self.out_proj = nn.Linear(n_embd, n_embd, bias=False)
|
||||
|
||||
# Layer-specific projection from shared RBF basis activations to per-head attention bias.
|
||||
self.rbf_proj = nn.Linear(n_rbf_bases, n_head, bias=False)
|
||||
self.time_bias_scale = nn.Parameter(torch.tensor(0.0))
|
||||
|
||||
self.resid_drop = nn.Dropout(dropout)
|
||||
self.reset_parameters()
|
||||
|
||||
def reset_parameters(self) -> None:
|
||||
nn.init.normal_(self.q_proj.weight, mean=0.0, std=0.02)
|
||||
nn.init.normal_(self.k_proj.weight, mean=0.0, std=0.02)
|
||||
nn.init.normal_(self.v_proj.weight, mean=0.0, std=0.02)
|
||||
if self.rbf_proj is not None:
|
||||
# The scalar gate starts at zero, so the relative-time bias still
|
||||
# starts disabled. A nonzero projection is necessary for the gate
|
||||
# itself to receive a gradient on the first optimization step.
|
||||
nn.init.xavier_uniform_(self.rbf_proj.weight)
|
||||
|
||||
def project_event_memory(
|
||||
self,
|
||||
event_memory: torch.Tensor,
|
||||
event_rope_cache: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Project K/V once for reuse by every reasoning round."""
|
||||
batch_size, memory_len, _ = event_memory.shape
|
||||
key = self.k_proj(event_memory).reshape(
|
||||
batch_size,
|
||||
memory_len,
|
||||
self.n_trajectory,
|
||||
self.trajectory_dim,
|
||||
).transpose(1, 2)
|
||||
value = self.v_proj(event_memory).reshape(
|
||||
batch_size,
|
||||
memory_len,
|
||||
self.n_trajectory,
|
||||
self.trajectory_dim,
|
||||
).transpose(1, 2)
|
||||
if self.use_time_rope:
|
||||
if event_rope_cache is None:
|
||||
raise ValueError(
|
||||
"event_rope_cache is required when TimeRoPE is enabled"
|
||||
)
|
||||
key = TimeRoPE.apply_single_from_cache(key, event_rope_cache)
|
||||
return key, value
|
||||
"""Match the previous version's GPT-style weight initialization."""
|
||||
nn.init.normal_(self.qkv.weight, mean=0.0, std=0.02)
|
||||
nn.init.normal_(self.out_proj.weight, mean=0.0, std=0.02)
|
||||
nn.init.zeros_(self.rbf_proj.weight)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
trajectory_state: torch.Tensor,
|
||||
event_key_value: tuple[torch.Tensor, torch.Tensor],
|
||||
event_invalid_mask: torch.Tensor,
|
||||
query_rope_cache: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
x: torch.Tensor,
|
||||
rope_cache: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
rbf_cache: torch.Tensor | None = None,
|
||||
attn_mask: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Read memory for states shaped ``(B, Q, H, Dh)``."""
|
||||
if trajectory_state.ndim != 4:
|
||||
raise ValueError(
|
||||
"trajectory_state must have shape (B, Q, H, Dh), got "
|
||||
f"{tuple(trajectory_state.shape)}"
|
||||
)
|
||||
batch_size, n_query, n_trajectory, trajectory_dim = (
|
||||
trajectory_state.shape
|
||||
)
|
||||
if (n_trajectory, trajectory_dim) != (
|
||||
self.n_trajectory,
|
||||
self.trajectory_dim,
|
||||
):
|
||||
raise ValueError(
|
||||
"Unexpected trajectory shape: "
|
||||
f"{(n_trajectory, trajectory_dim)}"
|
||||
)
|
||||
key, value = event_key_value
|
||||
memory_len = key.size(2)
|
||||
if event_invalid_mask.shape != (batch_size, n_query, memory_len):
|
||||
raise ValueError(
|
||||
"event_invalid_mask must have shape "
|
||||
f"{(batch_size, n_query, memory_len)}, got "
|
||||
f"{tuple(event_invalid_mask.shape)}"
|
||||
)
|
||||
|
||||
query = self.q_proj(trajectory_state).transpose(1, 2)
|
||||
if self.use_time_rope:
|
||||
if query_rope_cache is None:
|
||||
raise ValueError(
|
||||
"query_rope_cache is required when TimeRoPE is enabled"
|
||||
)
|
||||
query = TimeRoPE.apply_single_from_cache(query, query_rope_cache)
|
||||
query = query.transpose(1, 2)
|
||||
|
||||
scores = torch.einsum("bqhd,bhld->bqhl", query, key) * self.scale
|
||||
assert rope_cache is not None, "rope_cache must be provided when use_time_rope is True"
|
||||
if self.use_rbf_bias:
|
||||
if rbf_cache is None or self.rbf_proj is None:
|
||||
raise ValueError(
|
||||
"rbf_cache is required when relative time bias is enabled"
|
||||
)
|
||||
time_bias = self.rbf_proj(rbf_cache).permute(0, 1, 3, 2)
|
||||
scores = scores + self.time_bias_scale.tanh() * time_bias
|
||||
assert rbf_cache is not None, "rbf_cache must be provided when use_rbf_bias is True"
|
||||
|
||||
mask = event_invalid_mask.unsqueeze(2)
|
||||
min_value = torch.finfo(scores.dtype).min
|
||||
masked_scores = scores.masked_fill(mask, min_value)
|
||||
weights = torch.softmax(masked_scores.float(), dim=-1).to(scores.dtype)
|
||||
weights = weights.masked_fill(mask, 0.0)
|
||||
denominator = weights.sum(dim=-1, keepdim=True)
|
||||
weights = weights / denominator.clamp_min(
|
||||
torch.finfo(weights.dtype).eps
|
||||
B, L, _ = x.shape
|
||||
H, D = self.n_head, self.d_head
|
||||
|
||||
# --- QKV ----------------------------------------------------------
|
||||
qkv = self.qkv(x).reshape(B, L, 3, H, D).permute(2, 0, 3, 1, 4)
|
||||
q, k, v = qkv.unbind(0) # each (B, H, L, D)
|
||||
|
||||
# --- Apply RoPE (from shared cache) --------------------------------
|
||||
if self.use_time_rope:
|
||||
q, k = TimeRoPE.apply_from_cache(q, k, rope_cache)
|
||||
|
||||
# Build additive attention bias mask: time bias + causal/padding mask.
|
||||
time_bias = None
|
||||
if self.use_rbf_bias:
|
||||
time_bias = self.rbf_proj(rbf_cache).permute(
|
||||
0, 3, 1, 2) # (B, H, L, L)
|
||||
time_bias = self.time_bias_scale.tanh() * time_bias
|
||||
|
||||
if time_bias is not None and attn_mask is not None:
|
||||
attn_bias = time_bias + attn_mask.to(time_bias.dtype)
|
||||
elif time_bias is not None:
|
||||
attn_bias = time_bias
|
||||
elif attn_mask is not None:
|
||||
attn_bias = attn_mask
|
||||
else:
|
||||
attn_bias = None
|
||||
|
||||
out = F.scaled_dot_product_attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
attn_mask=attn_bias,
|
||||
dropout_p=0.0,
|
||||
is_causal=False,
|
||||
scale=self.scale,
|
||||
)
|
||||
return torch.einsum("bqhl,bhld->bqhd", weights, value)
|
||||
|
||||
# --- Aggregate & project out --------------------------------------
|
||||
out = out.transpose(1, 2).reshape(B, L, H * D)
|
||||
return self.resid_drop(self.out_proj(out))
|
||||
|
||||
|
||||
class SharedTrajectoryMixer(nn.Module):
|
||||
"""SwiGLU interaction along the trajectory axis only."""
|
||||
class TrajMixer(nn.Module):
|
||||
"""Lightweight gated interaction across latent residual-space groups.
|
||||
|
||||
def __init__(self, n_trajectory: int, trajectory_dim: int):
|
||||
The groups are contiguous partitions of the post-``W_O`` residual
|
||||
representation. They are deliberately not treated as attention heads.
|
||||
All operations are position-wise, so the sequence dimension remains fully
|
||||
parallel and no temporal information can leak between positions here.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
n_embd: int,
|
||||
n_head: int = 10,
|
||||
dropout: float = 0.0,
|
||||
):
|
||||
super().__init__()
|
||||
if n_trajectory <= 0 or trajectory_dim <= 0:
|
||||
raise ValueError("trajectory dimensions must be positive")
|
||||
self.n_trajectory = n_trajectory
|
||||
self.trajectory_dim = trajectory_dim
|
||||
self.traj_hidden = 4 * n_trajectory
|
||||
if n_embd <= 0:
|
||||
raise ValueError(f"n_embd must be > 0, got {n_embd}")
|
||||
if n_head <= 0:
|
||||
raise ValueError(f"n_head must be > 0, got {n_head}")
|
||||
if n_embd % n_head != 0:
|
||||
raise ValueError(
|
||||
f"n_embd must be divisible by n_head, got {n_embd} and {n_head}"
|
||||
)
|
||||
self.n_embd = n_embd
|
||||
# The residual-group count is tied to n_head, but the resulting groups
|
||||
# are still residual-space partitions rather than attention heads.
|
||||
self.n_group = n_head
|
||||
self.d_group = n_embd // n_head
|
||||
self.hidden_group = 4 * n_head
|
||||
|
||||
# Per-group feature alignment: [group, input feature, output feature].
|
||||
self.group_align = nn.Parameter(
|
||||
torch.empty(self.n_group, self.d_group, self.d_group)
|
||||
)
|
||||
|
||||
# Per-feature cross-group projections. The feature index is kept
|
||||
# independent, exactly as specified by the TrajMixer baseline.
|
||||
self.gate_proj = nn.Parameter(
|
||||
torch.empty(trajectory_dim, n_trajectory, self.traj_hidden)
|
||||
torch.empty(self.d_group, self.n_group, self.hidden_group)
|
||||
)
|
||||
self.value_proj = nn.Parameter(
|
||||
torch.empty(trajectory_dim, n_trajectory, self.traj_hidden)
|
||||
torch.empty(self.d_group, self.n_group, self.hidden_group)
|
||||
)
|
||||
self.output_proj = nn.Parameter(
|
||||
torch.empty(trajectory_dim, self.traj_hidden, n_trajectory)
|
||||
torch.empty(self.d_group, self.hidden_group, self.n_group)
|
||||
)
|
||||
self.drop = nn.Dropout(dropout)
|
||||
self.reset_parameters()
|
||||
|
||||
def reset_parameters(self) -> None:
|
||||
for feature_idx in range(self.trajectory_dim):
|
||||
with torch.no_grad():
|
||||
identity = torch.eye(
|
||||
self.d_group,
|
||||
dtype=self.group_align.dtype,
|
||||
device=self.group_align.device,
|
||||
)
|
||||
self.group_align.copy_(identity.unsqueeze(0).expand_as(self.group_align))
|
||||
|
||||
# Initialise each feature-specific matrix independently so Xavier's
|
||||
# fan-in/fan-out calculation sees a two-dimensional matrix.
|
||||
for feature_idx in range(self.d_group):
|
||||
nn.init.xavier_uniform_(self.gate_proj[feature_idx])
|
||||
nn.init.xavier_uniform_(self.value_proj[feature_idx])
|
||||
nn.init.normal_(self.output_proj, mean=0.0, std=1e-3)
|
||||
|
||||
def forward(self, state: torch.Tensor) -> torch.Tensor:
|
||||
if state.shape[-2:] != (self.n_trajectory, self.trajectory_dim):
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Map ``(B, L, n_embd)`` to an equally shaped residual update."""
|
||||
if x.ndim != 3:
|
||||
raise ValueError(f"TrajMixer expects a 3D tensor, got shape {tuple(x.shape)}")
|
||||
if x.size(-1) != self.n_embd:
|
||||
raise ValueError(
|
||||
"Expected trailing trajectory shape "
|
||||
f"{(self.n_trajectory, self.trajectory_dim)}, got "
|
||||
f"{tuple(state.shape[-2:])}"
|
||||
f"Expected hidden size {self.n_embd}, got {x.size(-1)}"
|
||||
)
|
||||
gate = torch.einsum("...hr,rhk->...kr", state, self.gate_proj)
|
||||
value = torch.einsum("...hr,rhk->...kr", state, self.value_proj)
|
||||
|
||||
batch_size, seq_len, _ = x.shape
|
||||
grouped = x.reshape(
|
||||
batch_size, seq_len, self.n_group, self.d_group
|
||||
)
|
||||
aligned = torch.einsum(
|
||||
"blgd,gde->blge", grouped, self.group_align
|
||||
)
|
||||
|
||||
gate = torch.einsum(
|
||||
"blgr,rgh->blhr", aligned, self.gate_proj
|
||||
)
|
||||
value = torch.einsum(
|
||||
"blgr,rgh->blhr", aligned, self.value_proj
|
||||
)
|
||||
hidden = F.silu(gate) * value
|
||||
return torch.einsum("...kr,rkh->...hr", hidden, self.output_proj)
|
||||
mixed = torch.einsum(
|
||||
"blhr,rhg->blgr", hidden, self.output_proj
|
||||
)
|
||||
return self.drop(mixed.reshape(batch_size, seq_len, self.n_embd))
|
||||
|
||||
|
||||
class SharedEventTrajectoryCore(nn.Module):
|
||||
"""One parameter-shared reasoning core reused across all rounds."""
|
||||
|
||||
class GPTBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
d_model: int,
|
||||
n_trajectory: int,
|
||||
n_reasoning_rounds: int,
|
||||
dropout: float = 0.0,
|
||||
n_rbf_bases: int = 16,
|
||||
n_embd: int,
|
||||
n_head: int,
|
||||
|
||||
attn_dropout: float = 0.0,
|
||||
mlp_dropout: float = 0.0,
|
||||
use_time_rope: bool = False,
|
||||
use_rbf_bias: bool = False,
|
||||
n_rbf_bases: int = 16,
|
||||
):
|
||||
super().__init__()
|
||||
if n_reasoning_rounds <= 0:
|
||||
raise ValueError("n_reasoning_rounds must be positive")
|
||||
if d_model <= 0 or n_trajectory <= 0:
|
||||
raise ValueError("d_model and n_trajectory must be positive")
|
||||
if d_model % n_trajectory != 0:
|
||||
raise ValueError("d_model must equal n_trajectory * trajectory_dim")
|
||||
trajectory_dim = d_model // n_trajectory
|
||||
self.norm_attn = nn.LayerNorm(trajectory_dim)
|
||||
self.cross_attention = TrajectoryCrossAttention(
|
||||
d_model=d_model,
|
||||
n_trajectory=n_trajectory,
|
||||
self.attn = TemporalAttention(
|
||||
n_embd=n_embd,
|
||||
n_head=n_head,
|
||||
n_rbf_bases=n_rbf_bases,
|
||||
dropout=attn_dropout,
|
||||
use_time_rope=use_time_rope,
|
||||
use_rbf_bias=use_rbf_bias,
|
||||
)
|
||||
self.norm_mixer = nn.LayerNorm(trajectory_dim)
|
||||
self.traj_mixer = SharedTrajectoryMixer(
|
||||
n_trajectory=n_trajectory,
|
||||
trajectory_dim=trajectory_dim,
|
||||
)
|
||||
initial_scale = 1.0 / math.sqrt(n_reasoning_rounds)
|
||||
self.attn_scale = nn.Parameter(torch.tensor(initial_scale))
|
||||
self.mixer_scale = nn.Parameter(torch.tensor(initial_scale))
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
|
||||
def project_event_memory(
|
||||
self,
|
||||
event_memory: torch.Tensor,
|
||||
event_rope_cache: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
return self.cross_attention.project_event_memory(
|
||||
event_memory,
|
||||
event_rope_cache=event_rope_cache,
|
||||
self.mlp = TrajMixer(
|
||||
n_embd=n_embd,
|
||||
n_head=n_head,
|
||||
dropout=mlp_dropout,
|
||||
)
|
||||
self.ln1 = nn.LayerNorm(n_embd)
|
||||
self.ln2 = nn.LayerNorm(n_embd)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
trajectory_state: torch.Tensor,
|
||||
event_key_value: tuple[torch.Tensor, torch.Tensor],
|
||||
event_invalid_mask: torch.Tensor,
|
||||
query_rope_cache: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
x: torch.Tensor,
|
||||
rope_cache: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
rbf_cache: torch.Tensor | None = None,
|
||||
attn_mask: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
readout = self.cross_attention(
|
||||
trajectory_state=self.norm_attn(trajectory_state),
|
||||
event_key_value=event_key_value,
|
||||
event_invalid_mask=event_invalid_mask,
|
||||
query_rope_cache=query_rope_cache,
|
||||
rbf_cache=rbf_cache,
|
||||
)
|
||||
updated = (
|
||||
trajectory_state
|
||||
+ self.attn_scale * self.dropout(readout)
|
||||
)
|
||||
mixed = self.traj_mixer(self.norm_mixer(updated))
|
||||
return updated + self.mixer_scale * self.dropout(mixed)
|
||||
x = x + self.attn(self.ln1(x), rope_cache, rbf_cache, attn_mask)
|
||||
x = x + self.mlp(self.ln2(x))
|
||||
return x
|
||||
|
||||
|
||||
class TokenAutoDiscretization(nn.Module):
|
||||
|
||||
@@ -42,8 +42,8 @@ from dataset import HealthDataset
|
||||
from eval_data import load_sequence_eval_dataset, sequence_eval_collate_fn
|
||||
from models import (
|
||||
DeepHealth,
|
||||
validate_event_trajectory_config,
|
||||
validate_event_trajectory_state_dict,
|
||||
validate_traj_mixer_config,
|
||||
validate_traj_mixer_state_dict,
|
||||
)
|
||||
from readouts import build_readout
|
||||
from targets import PAD_IDX, CHECKUP_IDX, NO_EVENT_IDX
|
||||
@@ -313,7 +313,7 @@ def split_indices(n: int, train_ratio: float, val_ratio: float, test_ratio: floa
|
||||
|
||||
|
||||
def build_model_from_dataset(args: argparse.Namespace, cfg: Dict[str, Any], dataset: HealthDataset) -> DeepHealth:
|
||||
validate_event_trajectory_config(cfg)
|
||||
validate_traj_mixer_config(cfg)
|
||||
model_target_mode = str(cfg_get(
|
||||
args, cfg, "model_target_mode", "next_token")).lower()
|
||||
if model_target_mode not in {"next_token", "all_future"}:
|
||||
@@ -322,10 +322,10 @@ def build_model_from_dataset(args: argparse.Namespace, cfg: Dict[str, Any], data
|
||||
)
|
||||
return DeepHealth(
|
||||
vocab_size=dataset.vocab_size,
|
||||
model_size=str(cfg_get(args, cfg, "model_size", "nano")),
|
||||
n_reasoning_rounds=int(
|
||||
cfg_get(args, cfg, "n_reasoning_rounds", 12)
|
||||
),
|
||||
n_embd=int(cfg_get(args, cfg, "n_embd", 120)),
|
||||
n_head=int(cfg_get(args, cfg, "n_head", 10)),
|
||||
n_hist_layer=int(cfg_get(args, cfg, "n_hist_layer", 12)),
|
||||
n_tab_layer=int(cfg_get(args, cfg, "n_tab_layer", 4)),
|
||||
n_types=dataset.n_types,
|
||||
n_cont_types=dataset.n_cont_types,
|
||||
n_categories=dataset.n_categories,
|
||||
@@ -391,12 +391,7 @@ def load_model_state(
|
||||
state = state_dict if state_dict is not None else load_checkpoint_state_dict(
|
||||
checkpoint_path, map_location=device)
|
||||
|
||||
validate_event_trajectory_state_dict(
|
||||
state,
|
||||
expected_d_model=model.d_model,
|
||||
expected_n_trajectory=model.n_trajectory,
|
||||
expected_n_reasoning_rounds=model.n_reasoning_rounds,
|
||||
)
|
||||
validate_traj_mixer_state_dict(state)
|
||||
model.load_state_dict(state, strict=True)
|
||||
|
||||
|
||||
@@ -533,7 +528,7 @@ def infer_readout_hidden(
|
||||
hidden = torch.zeros(
|
||||
batch_size,
|
||||
seq_len,
|
||||
model.d_model,
|
||||
model.n_embd,
|
||||
device=event_seq.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
|
||||
@@ -31,8 +31,8 @@ from dataset import HealthDataset
|
||||
from eval_data import load_sequence_eval_dataset
|
||||
from models import (
|
||||
DeepHealth,
|
||||
validate_event_trajectory_config,
|
||||
validate_event_trajectory_state_dict,
|
||||
validate_traj_mixer_config,
|
||||
validate_traj_mixer_state_dict,
|
||||
)
|
||||
from readouts import build_readout
|
||||
from targets import CHECKUP_IDX, NO_EVENT_IDX, PAD_IDX
|
||||
@@ -182,7 +182,7 @@ def resolve_dist_mode_for_checkpoint(cfg_dist_mode: str, state_dict: Dict[str, A
|
||||
|
||||
|
||||
def build_model_from_dataset(args: argparse.Namespace, cfg: Dict[str, Any], dataset: HealthDataset) -> DeepHealth:
|
||||
validate_event_trajectory_config(cfg)
|
||||
validate_traj_mixer_config(cfg)
|
||||
model_target_mode = str(cfg_get(
|
||||
args, cfg, "model_target_mode", "next_token")).lower()
|
||||
if model_target_mode not in {"next_token", "all_future"}:
|
||||
@@ -191,10 +191,10 @@ def build_model_from_dataset(args: argparse.Namespace, cfg: Dict[str, Any], data
|
||||
)
|
||||
return DeepHealth(
|
||||
vocab_size=dataset.vocab_size,
|
||||
model_size=str(cfg_get(args, cfg, "model_size", "nano")),
|
||||
n_reasoning_rounds=int(
|
||||
cfg_get(args, cfg, "n_reasoning_rounds", 12)
|
||||
),
|
||||
n_embd=int(cfg_get(args, cfg, "n_embd", 120)),
|
||||
n_head=int(cfg_get(args, cfg, "n_head", 10)),
|
||||
n_hist_layer=int(cfg_get(args, cfg, "n_hist_layer", 12)),
|
||||
n_tab_layer=int(cfg_get(args, cfg, "n_tab_layer", 4)),
|
||||
n_types=dataset.n_types,
|
||||
n_cont_types=dataset.n_cont_types,
|
||||
n_categories=dataset.n_categories,
|
||||
@@ -209,12 +209,7 @@ def build_model_from_dataset(args: argparse.Namespace, cfg: Dict[str, Any], data
|
||||
|
||||
|
||||
def load_model_state(model: torch.nn.Module, state_dict: Dict[str, Any]) -> None:
|
||||
validate_event_trajectory_state_dict(
|
||||
state_dict,
|
||||
expected_d_model=model.d_model,
|
||||
expected_n_trajectory=model.n_trajectory,
|
||||
expected_n_reasoning_rounds=model.n_reasoning_rounds,
|
||||
)
|
||||
validate_traj_mixer_state_dict(state_dict)
|
||||
model.load_state_dict(state_dict, strict=True)
|
||||
|
||||
|
||||
|
||||
@@ -205,7 +205,7 @@ def main() -> None:
|
||||
|
||||
n_rows = len(landmark_dataset)
|
||||
vocab_size = int(dataset.vocab_size)
|
||||
hidden_dim = int(model.d_model)
|
||||
hidden_dim = int(getattr(model, "n_embd", cfg_get(args, cfg_model, "n_embd", 120)))
|
||||
logits_dtype = numpy_float_dtype(args.logits_dtype)
|
||||
hidden_dtype = numpy_float_dtype(args.hidden_dtype)
|
||||
|
||||
|
||||
465
models.py
465
models.py
@@ -7,197 +7,39 @@ import torch.nn.functional as F
|
||||
|
||||
from backbones import (
|
||||
AgeSinusoidalEncoding,
|
||||
GPTBlock,
|
||||
GaussianRBFTimeBasis,
|
||||
SharedEventTrajectoryCore,
|
||||
TimeRoPE,
|
||||
TokenAutoDiscretization,
|
||||
)
|
||||
from targets import PAD_IDX
|
||||
|
||||
|
||||
EVENT_TRAJECTORY_ARCHITECTURE = "event_trajectory_shared_v1"
|
||||
TRAJ_MIXER_ARCHITECTURE = "traj_mixer_v2"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class EventTrajectoryModelSize:
|
||||
d_model: int
|
||||
n_trajectory: int
|
||||
|
||||
@property
|
||||
def trajectory_dim(self) -> int:
|
||||
return self.d_model // self.n_trajectory
|
||||
|
||||
@property
|
||||
def traj_hidden(self) -> int:
|
||||
return 4 * self.n_trajectory
|
||||
|
||||
|
||||
MODEL_SIZE_PRESETS = {
|
||||
"nano": EventTrajectoryModelSize(d_model=256, n_trajectory=8),
|
||||
"small": EventTrajectoryModelSize(d_model=512, n_trajectory=8),
|
||||
"medium": EventTrajectoryModelSize(d_model=768, n_trajectory=12),
|
||||
"huge": EventTrajectoryModelSize(d_model=1024, n_trajectory=16),
|
||||
}
|
||||
MODEL_SIZE_NAMES = tuple(MODEL_SIZE_PRESETS)
|
||||
|
||||
|
||||
def resolve_model_size(model_size: str) -> EventTrajectoryModelSize:
|
||||
if not isinstance(model_size, str):
|
||||
raise ValueError(
|
||||
f"model_size must be a string, got {type(model_size).__name__}"
|
||||
)
|
||||
normalized = model_size.strip().lower()
|
||||
try:
|
||||
return MODEL_SIZE_PRESETS[normalized]
|
||||
except KeyError as exc:
|
||||
choices = ", ".join(MODEL_SIZE_NAMES)
|
||||
raise ValueError(
|
||||
f"Unknown model_size {model_size!r}; expected one of: {choices}"
|
||||
) from exc
|
||||
|
||||
|
||||
def _required_config_int(
|
||||
config: Mapping[str, object],
|
||||
key: str,
|
||||
) -> int:
|
||||
raw_value = config.get(key)
|
||||
if isinstance(raw_value, bool):
|
||||
raise ValueError(f"Config field {key!r} must be an integer")
|
||||
try:
|
||||
value = int(raw_value)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ValueError(
|
||||
f"Config field {key!r} must be present and integer-valued; "
|
||||
f"got {raw_value!r}"
|
||||
) from exc
|
||||
if isinstance(raw_value, float) and not raw_value.is_integer():
|
||||
raise ValueError(f"Config field {key!r} must be an integer")
|
||||
return value
|
||||
|
||||
|
||||
def validate_event_trajectory_config(config: Mapping[str, object]) -> None:
|
||||
def validate_traj_mixer_config(config: Mapping[str, object]) -> None:
|
||||
actual = config.get("model_architecture")
|
||||
if actual != EVENT_TRAJECTORY_ARCHITECTURE:
|
||||
if actual != TRAJ_MIXER_ARCHITECTURE:
|
||||
raise ValueError(
|
||||
"This branch only accepts models trained with the shared "
|
||||
"event-trajectory architecture marker "
|
||||
f"{EVENT_TRAJECTORY_ARCHITECTURE!r}; got {actual!r}."
|
||||
)
|
||||
raw_model_size = config.get("model_size")
|
||||
if not isinstance(raw_model_size, str):
|
||||
raise ValueError(
|
||||
"Config field 'model_size' must be one of: "
|
||||
+ ", ".join(MODEL_SIZE_NAMES)
|
||||
)
|
||||
model_size = raw_model_size.strip().lower()
|
||||
preset = resolve_model_size(model_size)
|
||||
d_model = _required_config_int(config, "d_model")
|
||||
n_trajectory = _required_config_int(config, "n_trajectory")
|
||||
n_reasoning_rounds = _required_config_int(
|
||||
config,
|
||||
"n_reasoning_rounds",
|
||||
)
|
||||
trajectory_dim = _required_config_int(config, "trajectory_dim")
|
||||
traj_hidden = _required_config_int(config, "traj_hidden")
|
||||
if n_reasoning_rounds <= 0:
|
||||
raise ValueError(
|
||||
"n_reasoning_rounds must be positive"
|
||||
)
|
||||
expected_values = {
|
||||
"d_model": preset.d_model,
|
||||
"n_trajectory": preset.n_trajectory,
|
||||
"trajectory_dim": preset.trajectory_dim,
|
||||
"traj_hidden": preset.traj_hidden,
|
||||
}
|
||||
actual_values = {
|
||||
"d_model": d_model,
|
||||
"n_trajectory": n_trajectory,
|
||||
"trajectory_dim": trajectory_dim,
|
||||
"traj_hidden": traj_hidden,
|
||||
}
|
||||
mismatches = [
|
||||
f"{key}: expected {expected}, got {actual_values[key]}"
|
||||
for key, expected in expected_values.items()
|
||||
if actual_values[key] != expected
|
||||
]
|
||||
if mismatches:
|
||||
raise ValueError(
|
||||
f"Config does not match model_size={model_size!r}: "
|
||||
+ "; ".join(mismatches)
|
||||
"This branch only accepts models trained with the TrajMixer "
|
||||
f"architecture marker {TRAJ_MIXER_ARCHITECTURE!r}; got {actual!r}."
|
||||
)
|
||||
|
||||
|
||||
def _checkpoint_scalar_int(
|
||||
state_dict: Mapping[str, object],
|
||||
key: str,
|
||||
) -> int:
|
||||
value = state_dict[key]
|
||||
if not isinstance(value, torch.Tensor) or value.numel() != 1:
|
||||
raise ValueError(
|
||||
f"Checkpoint architecture field {key!r} must be a scalar tensor"
|
||||
)
|
||||
return int(value.detach().cpu().item())
|
||||
|
||||
|
||||
def validate_event_trajectory_state_dict(
|
||||
state_dict: Mapping[str, object],
|
||||
*,
|
||||
expected_d_model: int | None = None,
|
||||
expected_n_trajectory: int | None = None,
|
||||
expected_n_reasoning_rounds: int | None = None,
|
||||
) -> None:
|
||||
def validate_traj_mixer_state_dict(state_dict: Mapping[str, object]) -> None:
|
||||
required_keys = {
|
||||
"architecture_d_model",
|
||||
"architecture_n_trajectory",
|
||||
"architecture_n_reasoning_rounds",
|
||||
"event_projection.weight",
|
||||
"trajectory_prototypes",
|
||||
"query_projection.weight",
|
||||
"reasoning_core.cross_attention.q_proj.weight",
|
||||
"reasoning_core.cross_attention.k_proj.weight",
|
||||
"reasoning_core.cross_attention.v_proj.weight",
|
||||
"reasoning_core.traj_mixer.gate_proj",
|
||||
"reasoning_core.traj_mixer.value_proj",
|
||||
"reasoning_core.traj_mixer.output_proj",
|
||||
"reasoning_core.attn_scale",
|
||||
"reasoning_core.mixer_scale",
|
||||
"blocks.0.mlp.group_align",
|
||||
"blocks.0.mlp.gate_proj",
|
||||
"blocks.0.mlp.value_proj",
|
||||
"blocks.0.mlp.output_proj",
|
||||
}
|
||||
missing = sorted(required_keys.difference(state_dict))
|
||||
if missing:
|
||||
raise ValueError(
|
||||
"Checkpoint is not a shared event-trajectory checkpoint; "
|
||||
"missing required "
|
||||
"Checkpoint is not a TrajMixer checkpoint; missing required "
|
||||
f"parameters: {', '.join(missing)}"
|
||||
)
|
||||
checkpoint_values = {
|
||||
"d_model": _checkpoint_scalar_int(
|
||||
state_dict,
|
||||
"architecture_d_model",
|
||||
),
|
||||
"n_trajectory": _checkpoint_scalar_int(
|
||||
state_dict,
|
||||
"architecture_n_trajectory",
|
||||
),
|
||||
"n_reasoning_rounds": _checkpoint_scalar_int(
|
||||
state_dict,
|
||||
"architecture_n_reasoning_rounds",
|
||||
),
|
||||
}
|
||||
expected_values = {
|
||||
"d_model": expected_d_model,
|
||||
"n_trajectory": expected_n_trajectory,
|
||||
"n_reasoning_rounds": expected_n_reasoning_rounds,
|
||||
}
|
||||
mismatches = [
|
||||
f"{name}: checkpoint={checkpoint_values[name]}, expected={expected}"
|
||||
for name, expected in expected_values.items()
|
||||
if expected is not None and checkpoint_values[name] != expected
|
||||
]
|
||||
if mismatches:
|
||||
raise ValueError(
|
||||
"Checkpoint architecture does not match the constructed model: "
|
||||
+ "; ".join(mismatches)
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -331,8 +173,10 @@ class DeepHealth(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
vocab_size: int,
|
||||
model_size: str,
|
||||
n_reasoning_rounds: int,
|
||||
n_embd: int,
|
||||
n_head: int,
|
||||
n_hist_layer: int,
|
||||
n_tab_layer: int,
|
||||
n_types: int,
|
||||
n_cont_types: int,
|
||||
n_categories: int,
|
||||
@@ -357,21 +201,11 @@ class DeepHealth(nn.Module):
|
||||
"dist_mode must be either 'exponential', 'weibull' or 'mixed'")
|
||||
if extra_pool_reduce not in {"mean", "sum"}:
|
||||
raise ValueError("extra_pool_reduce must be either 'mean' or 'sum'")
|
||||
if n_reasoning_rounds <= 0:
|
||||
raise ValueError(
|
||||
"n_reasoning_rounds must be positive, got "
|
||||
f"{n_reasoning_rounds}"
|
||||
)
|
||||
size_config = resolve_model_size(model_size)
|
||||
normalized_model_size = model_size.strip().lower()
|
||||
d_model = size_config.d_model
|
||||
n_trajectory = size_config.n_trajectory
|
||||
|
||||
self.token_embedding = nn.Embedding(vocab_size, d_model, padding_idx=0)
|
||||
self.token_embedding = nn.Embedding(vocab_size, n_embd, padding_idx=0)
|
||||
self.gender_embedding = nn.Embedding(
|
||||
2, d_model) # Assuming binary gender
|
||||
2, n_embd) # Assuming binary gender
|
||||
self.tokenizer = OtherInfoTokenizer(
|
||||
n_embd=d_model,
|
||||
n_embd=n_embd,
|
||||
n_types=n_types,
|
||||
n_cont_types=n_cont_types,
|
||||
n_categories=n_categories,
|
||||
@@ -383,101 +217,70 @@ class DeepHealth(nn.Module):
|
||||
self.time_mode = time_mode
|
||||
self.dist_mode = dist_mode
|
||||
self.extra_pool_reduce = extra_pool_reduce
|
||||
self.model_size = normalized_model_size
|
||||
self.d_model = d_model
|
||||
self.n_trajectory = n_trajectory
|
||||
self.trajectory_dim = d_model // n_trajectory
|
||||
self.traj_hidden = 4 * n_trajectory
|
||||
self.n_reasoning_rounds = n_reasoning_rounds
|
||||
self.n_embd = n_embd
|
||||
self.vocab_size = vocab_size
|
||||
self.register_buffer(
|
||||
"architecture_d_model",
|
||||
torch.tensor(d_model, dtype=torch.int64),
|
||||
)
|
||||
self.register_buffer(
|
||||
"architecture_n_trajectory",
|
||||
torch.tensor(n_trajectory, dtype=torch.int64),
|
||||
)
|
||||
self.register_buffer(
|
||||
"architecture_n_reasoning_rounds",
|
||||
torch.tensor(n_reasoning_rounds, dtype=torch.int64),
|
||||
)
|
||||
nn.init.normal_(self.token_embedding.weight, mean=0.0, std=0.02)
|
||||
nn.init.zeros_(self.token_embedding.weight[0])
|
||||
nn.init.normal_(self.gender_embedding.weight, mean=0.0, std=0.02)
|
||||
if dist_mode == "weibull":
|
||||
self.rho_head = nn.Linear(d_model, vocab_size)
|
||||
self.rho_head = nn.Linear(n_embd, vocab_size)
|
||||
nn.init.zeros_(self.rho_head.weight)
|
||||
nn.init.constant_(self.rho_head.bias, 0.5413)
|
||||
|
||||
if dist_mode == "mixed":
|
||||
self.death_idx = vocab_size - 1
|
||||
self.rho_death_head = nn.Linear(d_model, 1)
|
||||
self.rho_death_head = nn.Linear(n_embd, 1)
|
||||
nn.init.zeros_(self.rho_death_head.weight)
|
||||
nn.init.constant_(self.rho_death_head.bias, 0.5413)
|
||||
|
||||
# Event and query time are encoded once before shared reasoning. In
|
||||
# relative mode, cross-attention additionally uses TimeRoPE and RBF.
|
||||
self.age_encoding = AgeSinusoidalEncoding(d_model)
|
||||
self.event_projection = nn.Linear(d_model, d_model, bias=False)
|
||||
self.event_norm = nn.LayerNorm(d_model)
|
||||
self.query_projection = nn.Linear(d_model, d_model, bias=False)
|
||||
self.trajectory_prototypes = nn.Parameter(
|
||||
torch.empty(n_trajectory, self.trajectory_dim)
|
||||
)
|
||||
self.query_token = nn.Parameter(torch.empty(d_model))
|
||||
nn.init.normal_(self.event_projection.weight, mean=0.0, std=0.02)
|
||||
nn.init.normal_(self.query_projection.weight, mean=0.0, std=0.02)
|
||||
nn.init.normal_(self.trajectory_prototypes, mean=0.0, std=0.02)
|
||||
nn.init.normal_(self.query_token, mean=0.0, std=0.02)
|
||||
|
||||
use_relative_time = time_mode == "relative"
|
||||
self.reasoning_core = SharedEventTrajectoryCore(
|
||||
d_model=d_model,
|
||||
n_trajectory=n_trajectory,
|
||||
n_reasoning_rounds=n_reasoning_rounds,
|
||||
dropout=dropout,
|
||||
n_rbf_bases=16,
|
||||
use_time_rope=use_relative_time,
|
||||
use_rbf_bias=use_relative_time,
|
||||
)
|
||||
if use_relative_time:
|
||||
self.rope = TimeRoPE(self.trajectory_dim)
|
||||
self.rbf = GaussianRBFTimeBasis(
|
||||
n_bases=16,
|
||||
max_time_diff=40.0,
|
||||
)
|
||||
else:
|
||||
if time_mode == "absolute":
|
||||
self.age_encoding = AgeSinusoidalEncoding(n_embd)
|
||||
self.blocks = nn.ModuleList([
|
||||
GPTBlock(
|
||||
n_embd=n_embd,
|
||||
n_head=n_head,
|
||||
use_time_rope=False,
|
||||
use_rbf_bias=False,
|
||||
mlp_dropout=dropout,
|
||||
) for _ in range(n_hist_layer)
|
||||
])
|
||||
self.rope = None
|
||||
self.rbf = None
|
||||
elif time_mode == "relative":
|
||||
self.age_encoding = None
|
||||
self.blocks = nn.ModuleList([
|
||||
GPTBlock(
|
||||
n_embd=n_embd,
|
||||
n_head=n_head,
|
||||
use_time_rope=True,
|
||||
use_rbf_bias=True,
|
||||
mlp_dropout=dropout,
|
||||
) for _ in range(n_hist_layer)
|
||||
])
|
||||
self.rope = TimeRoPE(n_embd // n_head)
|
||||
self.rbf = GaussianRBFTimeBasis(n_bases=16, max_time_diff=40.0)
|
||||
|
||||
self.final_ln = nn.LayerNorm(d_model)
|
||||
self.risk_head = nn.Linear(d_model, vocab_size, bias=False)
|
||||
self.final_ln = nn.LayerNorm(n_embd)
|
||||
self.risk_head = nn.Linear(n_embd, vocab_size, bias=False)
|
||||
if target_mode == "next_token":
|
||||
self.risk_head.weight = self.token_embedding.weight
|
||||
self.query_token = nn.Parameter(torch.zeros(n_embd))
|
||||
nn.init.normal_(self.query_token, mean=0.0, std=0.02)
|
||||
|
||||
def _make_event_invalid_mask(
|
||||
def _make_history_attn_mask(
|
||||
self,
|
||||
event_valid_mask: torch.Tensor,
|
||||
event_time: torch.Tensor,
|
||||
query_time: torch.Tensor,
|
||||
query_position: torch.Tensor | None = None,
|
||||
padding_mask: torch.Tensor,
|
||||
time_seq: torch.Tensor,
|
||||
dtype: torch.dtype,
|
||||
) -> torch.Tensor:
|
||||
valid_key = event_valid_mask[:, None, :]
|
||||
key_time = event_time[:, None, :]
|
||||
query_time = query_time[:, :, None]
|
||||
if query_position is None:
|
||||
visible_by_time = key_time <= query_time
|
||||
else:
|
||||
key_position = torch.arange(
|
||||
event_time.size(1),
|
||||
device=event_time.device,
|
||||
).view(1, 1, -1)
|
||||
visible_by_time = (key_time < query_time) | (
|
||||
(key_time == query_time)
|
||||
& (key_position <= query_position[:, :, None])
|
||||
)
|
||||
return ~(valid_key & visible_by_time)
|
||||
valid_key = padding_mask[:, None, :] # (B, 1, L)
|
||||
visible_by_time = time_seq[:, None, :] <= time_seq[:, :, None]
|
||||
valid = valid_key & visible_by_time
|
||||
return torch.zeros(
|
||||
valid.shape,
|
||||
device=valid.device,
|
||||
dtype=dtype,
|
||||
).masked_fill(~valid, -1e4)[:, None, :, :]
|
||||
|
||||
def _pool_other_by_time(
|
||||
self,
|
||||
@@ -581,8 +384,8 @@ class DeepHealth(nn.Module):
|
||||
padding_mask = padding_mask.to(device=event_seq.device, dtype=torch.bool)
|
||||
|
||||
event_len = event_seq.size(1)
|
||||
event_features = self.token_embedding(event_seq)
|
||||
event_time = time_seq
|
||||
h_disease = self.token_embedding(event_seq)
|
||||
t_disease = time_seq
|
||||
|
||||
if other_time.shape != other_type.shape:
|
||||
raise ValueError(
|
||||
@@ -590,120 +393,64 @@ class DeepHealth(nn.Module):
|
||||
f"{tuple(other_time.shape)} vs {tuple(other_type.shape)}"
|
||||
)
|
||||
other_time = other_time.to(device=event_seq.device, dtype=time_seq.dtype)
|
||||
other_features, other_mask = self.tokenizer(
|
||||
h_other, other_mask = self.tokenizer(
|
||||
other_type=other_type,
|
||||
other_value=other_value,
|
||||
other_value_kind=other_value_kind,
|
||||
)
|
||||
other_features = other_features.to(device=event_seq.device)
|
||||
h_other = h_other.to(device=event_seq.device)
|
||||
other_mask = other_mask.to(device=event_seq.device, dtype=torch.bool)
|
||||
|
||||
event_features = torch.cat([event_features, other_features], dim=1)
|
||||
event_time = torch.cat([event_time, other_time], dim=1)
|
||||
event_valid_mask = torch.cat([padding_mask, other_mask], dim=1)
|
||||
|
||||
batch_size = event_seq.size(0)
|
||||
sex_context = self.gender_embedding(sex)[:, None, :]
|
||||
event_features = (
|
||||
event_features
|
||||
+ sex_context
|
||||
+ self.age_encoding(event_time)
|
||||
)
|
||||
event_features = event_features * event_valid_mask.unsqueeze(-1).to(
|
||||
event_features.dtype
|
||||
)
|
||||
event_memory = self.event_norm(
|
||||
self.event_projection(event_features)
|
||||
)
|
||||
event_memory = event_memory * event_valid_mask.unsqueeze(-1).to(
|
||||
event_memory.dtype
|
||||
)
|
||||
h_disease = torch.cat([h_disease, h_other], dim=1)
|
||||
t_disease = torch.cat([t_disease, other_time], dim=1)
|
||||
padding_mask = torch.cat([padding_mask, other_mask], dim=1)
|
||||
h_disease = h_disease * padding_mask.unsqueeze(-1).to(h_disease.dtype)
|
||||
|
||||
if mode == "all_future":
|
||||
query_time = t_query[:, None]
|
||||
query_position = None
|
||||
query_features = (
|
||||
self.query_token.view(1, 1, -1)
|
||||
+ sex_context
|
||||
+ self.age_encoding(query_time)
|
||||
)
|
||||
query_valid_mask = torch.ones(
|
||||
batch_size = event_seq.size(0)
|
||||
query = self.query_token.view(1, 1, -1).expand(batch_size, 1, -1)
|
||||
h_disease = torch.cat([h_disease, query], dim=1)
|
||||
t_disease = torch.cat([t_disease, t_query[:, None]], dim=1)
|
||||
query_mask = torch.ones(
|
||||
batch_size,
|
||||
1,
|
||||
dtype=torch.bool,
|
||||
device=event_seq.device,
|
||||
)
|
||||
else:
|
||||
# Each event position is an independent parallel query. Including
|
||||
# its event feature preserves token-level next-step semantics.
|
||||
# Equal-time memory is additionally position-causal so a token
|
||||
# cannot read a later token that may be its Delphi2M target.
|
||||
query_time = event_time
|
||||
query_position = torch.arange(
|
||||
event_time.size(1),
|
||||
device=event_time.device,
|
||||
).view(1, -1).expand(batch_size, -1)
|
||||
query_features = event_features
|
||||
query_valid_mask = event_valid_mask
|
||||
padding_mask = torch.cat([padding_mask, query_mask], dim=1)
|
||||
|
||||
n_query = query_time.size(1)
|
||||
query_context = self.query_projection(query_features).reshape(
|
||||
batch_size,
|
||||
n_query,
|
||||
self.n_trajectory,
|
||||
self.trajectory_dim,
|
||||
)
|
||||
trajectory_state = (
|
||||
self.trajectory_prototypes.view(
|
||||
1,
|
||||
1,
|
||||
self.n_trajectory,
|
||||
self.trajectory_dim,
|
||||
)
|
||||
+ query_context
|
||||
)
|
||||
event_invalid_mask = self._make_event_invalid_mask(
|
||||
event_valid_mask=event_valid_mask,
|
||||
event_time=event_time,
|
||||
query_time=query_time,
|
||||
query_position=query_position,
|
||||
)
|
||||
sex_emb = self.gender_embedding(sex)[:, None, :]
|
||||
h_disease = h_disease + sex_emb
|
||||
h_disease = h_disease * padding_mask.unsqueeze(-1).to(h_disease.dtype)
|
||||
|
||||
event_rope_cache = None
|
||||
query_rope_cache = None
|
||||
rope_cache = None
|
||||
rbf_cache = None
|
||||
if self.time_mode == "relative":
|
||||
if self.rope is None or self.rbf is None:
|
||||
raise RuntimeError("Relative-time modules are not initialized")
|
||||
event_rope_cache = self.rope.precompute_cache(event_time)
|
||||
query_rope_cache = self.rope.precompute_cache(query_time)
|
||||
rbf_cache = self.rbf.precompute_cross_cache(
|
||||
query_time,
|
||||
event_time,
|
||||
)
|
||||
if self.time_mode == "absolute":
|
||||
h_disease = h_disease + self.age_encoding(t_disease)
|
||||
h_disease = h_disease * padding_mask.unsqueeze(-1).to(h_disease.dtype)
|
||||
elif self.time_mode == "relative":
|
||||
rope_cache = self.rope.precompute_cache(t_disease)
|
||||
rbf_cache = self.rbf.precompute_cache(t_disease)
|
||||
|
||||
event_key_value = self.reasoning_core.project_event_memory(
|
||||
event_memory,
|
||||
event_rope_cache=event_rope_cache,
|
||||
attn_mask = self._make_history_attn_mask(
|
||||
padding_mask=padding_mask,
|
||||
time_seq=t_disease,
|
||||
dtype=h_disease.dtype,
|
||||
)
|
||||
for _ in range(self.n_reasoning_rounds):
|
||||
trajectory_state = self.reasoning_core(
|
||||
trajectory_state=trajectory_state,
|
||||
event_key_value=event_key_value,
|
||||
event_invalid_mask=event_invalid_mask,
|
||||
query_rope_cache=query_rope_cache,
|
||||
for block in self.blocks:
|
||||
h_disease = block(
|
||||
h_disease,
|
||||
rope_cache=rope_cache,
|
||||
rbf_cache=rbf_cache,
|
||||
attn_mask=attn_mask,
|
||||
)
|
||||
h_disease = h_disease * padding_mask.unsqueeze(-1).to(h_disease.dtype)
|
||||
|
||||
hidden_sequence = self.final_ln(
|
||||
trajectory_state.reshape(batch_size, n_query, self.d_model)
|
||||
)
|
||||
hidden_sequence = hidden_sequence * query_valid_mask.unsqueeze(-1).to(
|
||||
hidden_sequence.dtype
|
||||
)
|
||||
h_disease = self.final_ln(h_disease)
|
||||
h_disease = h_disease * padding_mask.unsqueeze(-1).to(h_disease.dtype)
|
||||
|
||||
if mode == "all_future":
|
||||
hidden = hidden_sequence[:, 0, :]
|
||||
hidden = h_disease[:, -1, :]
|
||||
if return_output:
|
||||
return DeepHealthOutput(
|
||||
hidden=hidden,
|
||||
@@ -718,13 +465,13 @@ class DeepHealth(nn.Module):
|
||||
)
|
||||
return hidden
|
||||
if return_output:
|
||||
h_event = hidden_sequence[:, :event_len, :]
|
||||
t_event = event_time[:, :event_len]
|
||||
event_mask = event_valid_mask[:, :event_len]
|
||||
h_event = h_disease[:, :event_len, :]
|
||||
t_event = t_disease[:, :event_len]
|
||||
event_mask = padding_mask[:, :event_len]
|
||||
h_extra, t_extra, extra_mask = self._pool_other_by_time(
|
||||
h_other=hidden_sequence[:, event_len:, :],
|
||||
other_time=event_time[:, event_len:],
|
||||
other_mask=event_valid_mask[:, event_len:],
|
||||
h_other=h_disease[:, event_len:, :],
|
||||
other_time=t_disease[:, event_len:],
|
||||
other_mask=padding_mask[:, event_len:],
|
||||
)
|
||||
return DeepHealthOutput(
|
||||
hidden=torch.cat([h_event, h_extra], dim=1),
|
||||
@@ -732,7 +479,7 @@ class DeepHealth(nn.Module):
|
||||
padding_mask=torch.cat([event_mask, extra_mask], dim=1),
|
||||
event_len=event_len,
|
||||
)
|
||||
return hidden_sequence[:, :event_len, :]
|
||||
return h_disease[:, :event_len, :]
|
||||
|
||||
def forward_next_token(self, **kwargs) -> torch.Tensor:
|
||||
return self._forward_shared(mode="next_token", **kwargs)
|
||||
|
||||
@@ -1,325 +0,0 @@
|
||||
import math
|
||||
import unittest
|
||||
|
||||
import torch
|
||||
|
||||
from backbones import (
|
||||
SharedEventTrajectoryCore,
|
||||
SharedTrajectoryMixer,
|
||||
TrajectoryCrossAttention,
|
||||
)
|
||||
from models import (
|
||||
EVENT_TRAJECTORY_ARCHITECTURE,
|
||||
MODEL_SIZE_PRESETS,
|
||||
DeepHealth,
|
||||
resolve_model_size,
|
||||
validate_event_trajectory_config,
|
||||
validate_event_trajectory_state_dict,
|
||||
)
|
||||
from train_util import get_model_parameter_counts
|
||||
|
||||
|
||||
def build_test_model(
|
||||
*,
|
||||
target_mode: str = "next_token",
|
||||
time_mode: str = "absolute",
|
||||
n_reasoning_rounds: int = 3,
|
||||
) -> DeepHealth:
|
||||
return DeepHealth(
|
||||
vocab_size=32,
|
||||
model_size="nano",
|
||||
n_reasoning_rounds=n_reasoning_rounds,
|
||||
n_types=2,
|
||||
n_cont_types=0,
|
||||
n_categories=2,
|
||||
cont_type_ids=[],
|
||||
target_mode=target_mode,
|
||||
time_mode=time_mode,
|
||||
)
|
||||
|
||||
|
||||
def model_inputs() -> dict[str, torch.Tensor]:
|
||||
return {
|
||||
"event_seq": torch.tensor([[1, 2, 3, 4], [5, 6, 0, 0]]),
|
||||
"time_seq": torch.tensor(
|
||||
[[1.0, 2.0, 3.0, 4.0], [1.0, 2.0, 0.0, 0.0]]
|
||||
),
|
||||
"sex": torch.tensor([0, 1]),
|
||||
"padding_mask": torch.tensor(
|
||||
[[True, True, True, True], [True, True, False, False]]
|
||||
),
|
||||
"other_type": torch.zeros(2, 1, dtype=torch.long),
|
||||
"other_value": torch.zeros(2, 1),
|
||||
"other_value_kind": torch.zeros(2, 1, dtype=torch.long),
|
||||
"other_time": torch.zeros(2, 1),
|
||||
}
|
||||
|
||||
|
||||
class EventTrajectoryBackboneTest(unittest.TestCase):
|
||||
def test_model_size_presets(self) -> None:
|
||||
expected = {
|
||||
"nano": (256, 8, 32, 32),
|
||||
"small": (512, 8, 64, 32),
|
||||
"medium": (768, 12, 64, 48),
|
||||
"huge": (1024, 16, 64, 64),
|
||||
}
|
||||
self.assertEqual(set(MODEL_SIZE_PRESETS), set(expected))
|
||||
for name, values in expected.items():
|
||||
preset = resolve_model_size(name)
|
||||
self.assertEqual(
|
||||
(
|
||||
preset.d_model,
|
||||
preset.n_trajectory,
|
||||
preset.trajectory_dim,
|
||||
preset.traj_hidden,
|
||||
),
|
||||
values,
|
||||
)
|
||||
|
||||
def test_default_mixer_shapes_and_parameter_count(self) -> None:
|
||||
mixer = SharedTrajectoryMixer(
|
||||
n_trajectory=8,
|
||||
trajectory_dim=32,
|
||||
)
|
||||
state = torch.randn(2, 5, 8, 32)
|
||||
self.assertEqual(mixer(state).shape, state.shape)
|
||||
self.assertEqual(mixer.traj_hidden, 32)
|
||||
self.assertEqual(tuple(mixer.gate_proj.shape), (32, 8, 32))
|
||||
self.assertEqual(tuple(mixer.value_proj.shape), (32, 8, 32))
|
||||
self.assertEqual(tuple(mixer.output_proj.shape), (32, 32, 8))
|
||||
self.assertEqual(
|
||||
get_model_parameter_counts(mixer),
|
||||
{
|
||||
"model_parameter_count": 24_576,
|
||||
"trainable_parameter_count": 24_576,
|
||||
},
|
||||
)
|
||||
|
||||
def test_reasoning_rounds_share_one_core_parameter_set(self) -> None:
|
||||
core_one = SharedEventTrajectoryCore(
|
||||
d_model=256,
|
||||
n_trajectory=8,
|
||||
n_reasoning_rounds=1,
|
||||
)
|
||||
core_twelve = SharedEventTrajectoryCore(
|
||||
d_model=256,
|
||||
n_trajectory=8,
|
||||
n_reasoning_rounds=12,
|
||||
)
|
||||
self.assertEqual(
|
||||
sum(p.numel() for p in core_one.parameters()),
|
||||
sum(p.numel() for p in core_twelve.parameters()),
|
||||
)
|
||||
self.assertAlmostEqual(core_one.attn_scale.item(), 1.0)
|
||||
self.assertAlmostEqual(
|
||||
core_twelve.attn_scale.item(),
|
||||
1.0 / math.sqrt(12),
|
||||
places=6,
|
||||
)
|
||||
self.assertAlmostEqual(
|
||||
core_twelve.mixer_scale.item(),
|
||||
1.0 / math.sqrt(12),
|
||||
places=6,
|
||||
)
|
||||
|
||||
def test_all_masked_attention_is_finite_and_zero(self) -> None:
|
||||
attention = TrajectoryCrossAttention(
|
||||
d_model=32,
|
||||
n_trajectory=4,
|
||||
)
|
||||
memory = torch.randn(2, 3, 32)
|
||||
key_value = attention.project_event_memory(memory)
|
||||
state = torch.randn(2, 2, 4, 8)
|
||||
invalid_mask = torch.ones(2, 2, 3, dtype=torch.bool)
|
||||
output = attention(
|
||||
trajectory_state=state,
|
||||
event_key_value=key_value,
|
||||
event_invalid_mask=invalid_mask,
|
||||
)
|
||||
self.assertTrue(torch.isfinite(output).all())
|
||||
torch.testing.assert_close(output, torch.zeros_like(output))
|
||||
|
||||
def test_next_token_future_events_do_not_change_earlier_query(self) -> None:
|
||||
torch.manual_seed(0)
|
||||
model = build_test_model(n_reasoning_rounds=2)
|
||||
model.eval()
|
||||
inputs = model_inputs()
|
||||
original = model(**inputs)
|
||||
changed_inputs = dict(inputs)
|
||||
changed_inputs["event_seq"] = inputs["event_seq"].clone()
|
||||
changed_inputs["event_seq"][0, 3] = 9
|
||||
changed = model(**changed_inputs)
|
||||
torch.testing.assert_close(original[0, 1], changed[0, 1])
|
||||
|
||||
def test_next_token_later_equal_time_event_is_not_visible(self) -> None:
|
||||
torch.manual_seed(0)
|
||||
model = build_test_model(n_reasoning_rounds=2)
|
||||
model.eval()
|
||||
inputs = model_inputs()
|
||||
inputs["time_seq"] = inputs["time_seq"].clone()
|
||||
inputs["time_seq"][0] = torch.tensor([1.0, 1.0, 2.0, 3.0])
|
||||
original = model(**inputs)
|
||||
changed_inputs = dict(inputs)
|
||||
changed_inputs["event_seq"] = inputs["event_seq"].clone()
|
||||
changed_inputs["event_seq"][0, 1] = 9
|
||||
changed = model(**changed_inputs)
|
||||
torch.testing.assert_close(original[0, 0], changed[0, 0])
|
||||
|
||||
def test_padding_content_does_not_change_valid_queries(self) -> None:
|
||||
torch.manual_seed(0)
|
||||
model = build_test_model(
|
||||
time_mode="relative",
|
||||
n_reasoning_rounds=2,
|
||||
)
|
||||
model.eval()
|
||||
inputs = model_inputs()
|
||||
original = model(**inputs)
|
||||
changed_inputs = dict(inputs)
|
||||
changed_inputs["event_seq"] = inputs["event_seq"].clone()
|
||||
changed_inputs["time_seq"] = inputs["time_seq"].clone()
|
||||
changed_inputs["event_seq"][1, 2:] = torch.tensor([9, 10])
|
||||
changed_inputs["time_seq"][1, 2:] = torch.tensor([30.0, 40.0])
|
||||
changed = model(**changed_inputs)
|
||||
torch.testing.assert_close(original[1, :2], changed[1, :2])
|
||||
|
||||
def test_next_token_and_all_future_output_contracts(self) -> None:
|
||||
inputs = model_inputs()
|
||||
next_model = build_test_model(target_mode="next_token")
|
||||
next_hidden = next_model(**inputs)
|
||||
self.assertEqual(tuple(next_hidden.shape), (2, 4, 256))
|
||||
next_output = next_model(**inputs, return_output=True)
|
||||
self.assertEqual(tuple(next_output.hidden.shape), (2, 4, 256))
|
||||
self.assertEqual(tuple(next_output.padding_mask.shape), (2, 4))
|
||||
|
||||
future_model = build_test_model(target_mode="all_future")
|
||||
future_hidden = future_model(
|
||||
**inputs,
|
||||
t_query=torch.tensor([5.0, 3.0]),
|
||||
)
|
||||
self.assertEqual(tuple(future_hidden.shape), (2, 256))
|
||||
|
||||
def test_model_contains_one_shared_core_and_no_block_stack(self) -> None:
|
||||
model = build_test_model(n_reasoning_rounds=12)
|
||||
self.assertFalse(hasattr(model, "blocks"))
|
||||
reasoning_keys = [
|
||||
key
|
||||
for key in model.state_dict()
|
||||
if key.startswith("reasoning_core.")
|
||||
]
|
||||
self.assertTrue(reasoning_keys)
|
||||
self.assertFalse(any("blocks." in key for key in model.state_dict()))
|
||||
self.assertFalse(any("out_proj" in key for key in reasoning_keys))
|
||||
self.assertFalse(any("group_align" in key for key in reasoning_keys))
|
||||
|
||||
def test_event_key_and_value_are_projected_once_per_forward(self) -> None:
|
||||
model = build_test_model(n_reasoning_rounds=12)
|
||||
call_counts = {"key": 0, "value": 0}
|
||||
|
||||
def count_key(*_args) -> None:
|
||||
call_counts["key"] += 1
|
||||
|
||||
def count_value(*_args) -> None:
|
||||
call_counts["value"] += 1
|
||||
|
||||
key_handle = (
|
||||
model.reasoning_core.cross_attention.k_proj
|
||||
.register_forward_hook(count_key)
|
||||
)
|
||||
value_handle = (
|
||||
model.reasoning_core.cross_attention.v_proj
|
||||
.register_forward_hook(count_value)
|
||||
)
|
||||
try:
|
||||
model(**model_inputs())
|
||||
finally:
|
||||
key_handle.remove()
|
||||
value_handle.remove()
|
||||
self.assertEqual(call_counts, {"key": 1, "value": 1})
|
||||
|
||||
def test_relative_time_forward_and_backward_are_finite(self) -> None:
|
||||
torch.manual_seed(0)
|
||||
model = build_test_model(
|
||||
target_mode="all_future",
|
||||
time_mode="relative",
|
||||
n_reasoning_rounds=2,
|
||||
)
|
||||
hidden = model(
|
||||
**model_inputs(),
|
||||
t_query=torch.tensor([5.0, 3.0]),
|
||||
)
|
||||
(hidden * torch.randn_like(hidden)).sum().backward()
|
||||
self.assertTrue(torch.isfinite(hidden).all())
|
||||
self.assertIsNotNone(model.event_projection.weight.grad)
|
||||
self.assertTrue(torch.isfinite(model.event_projection.weight.grad).all())
|
||||
time_scale = model.reasoning_core.cross_attention.time_bias_scale
|
||||
self.assertIsNotNone(time_scale)
|
||||
self.assertIsNotNone(time_scale.grad)
|
||||
self.assertGreater(abs(float(time_scale.grad)), 0.0)
|
||||
|
||||
def test_architecture_marker_and_checkpoint_are_required(self) -> None:
|
||||
validate_event_trajectory_config(
|
||||
{
|
||||
"model_architecture": EVENT_TRAJECTORY_ARCHITECTURE,
|
||||
"model_size": "nano",
|
||||
"d_model": 256,
|
||||
"n_trajectory": 8,
|
||||
"trajectory_dim": 32,
|
||||
"traj_hidden": 32,
|
||||
"n_reasoning_rounds": 3,
|
||||
}
|
||||
)
|
||||
with self.assertRaisesRegex(ValueError, "only accepts models trained"):
|
||||
validate_event_trajectory_config(
|
||||
{"model_architecture": "traj_mixer_v2"}
|
||||
)
|
||||
with self.assertRaisesRegex(ValueError, "trajectory_dim"):
|
||||
validate_event_trajectory_config(
|
||||
{
|
||||
"model_architecture": EVENT_TRAJECTORY_ARCHITECTURE,
|
||||
"model_size": "nano",
|
||||
"d_model": 256,
|
||||
"n_trajectory": 8,
|
||||
"trajectory_dim": 16,
|
||||
"traj_hidden": 32,
|
||||
"n_reasoning_rounds": 3,
|
||||
}
|
||||
)
|
||||
|
||||
model = build_test_model()
|
||||
state_dict = model.state_dict()
|
||||
validate_event_trajectory_state_dict(
|
||||
state_dict,
|
||||
expected_d_model=256,
|
||||
expected_n_trajectory=8,
|
||||
expected_n_reasoning_rounds=3,
|
||||
)
|
||||
with self.assertRaisesRegex(
|
||||
ValueError,
|
||||
"Checkpoint architecture does not match",
|
||||
):
|
||||
validate_event_trajectory_state_dict(
|
||||
state_dict,
|
||||
expected_n_reasoning_rounds=12,
|
||||
)
|
||||
state_dict.pop("reasoning_core.attn_scale")
|
||||
with self.assertRaisesRegex(
|
||||
ValueError,
|
||||
"not a shared event-trajectory checkpoint",
|
||||
):
|
||||
validate_event_trajectory_state_dict(state_dict)
|
||||
|
||||
def test_unknown_model_size_is_rejected(self) -> None:
|
||||
with self.assertRaisesRegex(ValueError, "Unknown model_size"):
|
||||
DeepHealth(
|
||||
vocab_size=32,
|
||||
model_size="giant",
|
||||
n_reasoning_rounds=2,
|
||||
n_types=2,
|
||||
n_cont_types=0,
|
||||
n_categories=2,
|
||||
cont_type_ids=[],
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
107
test_traj_mixer.py
Normal file
107
test_traj_mixer.py
Normal file
@@ -0,0 +1,107 @@
|
||||
import unittest
|
||||
|
||||
import torch
|
||||
|
||||
from backbones import GPTBlock, TrajMixer
|
||||
from models import (
|
||||
TRAJ_MIXER_ARCHITECTURE,
|
||||
validate_traj_mixer_config,
|
||||
validate_traj_mixer_state_dict,
|
||||
)
|
||||
from train_util import get_model_parameter_counts
|
||||
|
||||
|
||||
class TrajMixerTest(unittest.TestCase):
|
||||
def test_default_shape_parameters_and_initialization(self) -> None:
|
||||
mixer = TrajMixer(
|
||||
n_embd=120,
|
||||
n_head=10,
|
||||
dropout=0.0,
|
||||
)
|
||||
|
||||
x = torch.randn(2, 7, 120)
|
||||
self.assertEqual(mixer(x).shape, x.shape)
|
||||
self.assertEqual(sum(p.numel() for p in mixer.parameters()), 15_840)
|
||||
|
||||
expected = torch.eye(12).expand(10, 12, 12)
|
||||
torch.testing.assert_close(mixer.group_align.detach(), expected)
|
||||
self.assertEqual(mixer.hidden_group, 40)
|
||||
self.assertEqual(tuple(mixer.gate_proj.shape), (12, 10, 40))
|
||||
self.assertEqual(tuple(mixer.value_proj.shape), (12, 10, 40))
|
||||
self.assertEqual(tuple(mixer.output_proj.shape), (12, 40, 10))
|
||||
|
||||
def test_mixer_does_not_mix_sequence_positions(self) -> None:
|
||||
torch.manual_seed(0)
|
||||
mixer = TrajMixer(120, n_head=10, dropout=0.0)
|
||||
mixer.eval()
|
||||
x = torch.randn(2, 5, 120)
|
||||
changed = x.clone()
|
||||
changed[:, 3, :] += torch.randn_like(changed[:, 3, :])
|
||||
|
||||
original_out = mixer(x)
|
||||
changed_out = mixer(changed)
|
||||
unchanged_positions = torch.tensor([0, 1, 2, 4])
|
||||
torch.testing.assert_close(
|
||||
original_out.index_select(1, unchanged_positions),
|
||||
changed_out.index_select(1, unchanged_positions),
|
||||
)
|
||||
|
||||
def test_gradients_reach_all_projection_families(self) -> None:
|
||||
torch.manual_seed(1)
|
||||
mixer = TrajMixer(120, n_head=10, dropout=0.0)
|
||||
x = torch.randn(2, 4, 120, requires_grad=True)
|
||||
|
||||
mixer(x).square().mean().backward()
|
||||
|
||||
self.assertIsNotNone(x.grad)
|
||||
for name, parameter in mixer.named_parameters():
|
||||
self.assertIsNotNone(parameter.grad, name)
|
||||
self.assertTrue(torch.isfinite(parameter.grad).all(), name)
|
||||
|
||||
def test_gpt_block_defaults_to_traj_mixer_and_standard_layer_norm(self) -> None:
|
||||
block = GPTBlock(n_embd=120, n_head=10)
|
||||
self.assertIsInstance(block.mlp, TrajMixer)
|
||||
self.assertIsInstance(block.ln2, torch.nn.LayerNorm)
|
||||
self.assertEqual(tuple(block.ln2.normalized_shape), (120,))
|
||||
|
||||
x = torch.randn(2, 6, 120)
|
||||
self.assertEqual(block(x).shape, x.shape)
|
||||
|
||||
def test_architecture_marker_is_required(self) -> None:
|
||||
validate_traj_mixer_config(
|
||||
{"model_architecture": TRAJ_MIXER_ARCHITECTURE}
|
||||
)
|
||||
with self.assertRaisesRegex(ValueError, "only accepts models trained"):
|
||||
validate_traj_mixer_config({})
|
||||
with self.assertRaisesRegex(ValueError, "only accepts models trained"):
|
||||
validate_traj_mixer_config({"model_architecture": "delphi_swiglu"})
|
||||
|
||||
def test_checkpoint_must_contain_traj_mixer_parameters(self) -> None:
|
||||
block = GPTBlock(n_embd=120, n_head=10)
|
||||
state_dict = {
|
||||
f"blocks.0.{key}": value
|
||||
for key, value in block.state_dict().items()
|
||||
}
|
||||
validate_traj_mixer_state_dict(state_dict)
|
||||
|
||||
state_dict.pop("blocks.0.mlp.group_align")
|
||||
with self.assertRaisesRegex(ValueError, "not a TrajMixer checkpoint"):
|
||||
validate_traj_mixer_state_dict(state_dict)
|
||||
|
||||
def test_invalid_group_partition_is_rejected(self) -> None:
|
||||
with self.assertRaisesRegex(ValueError, "divisible"):
|
||||
TrajMixer(n_embd=121, n_head=10)
|
||||
|
||||
def test_parameter_counts_match_traj_mixer_parameters(self) -> None:
|
||||
mixer = TrajMixer(n_embd=120, n_head=10)
|
||||
self.assertEqual(
|
||||
get_model_parameter_counts(mixer),
|
||||
{
|
||||
"model_parameter_count": 15_840,
|
||||
"trainable_parameter_count": 15_840,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -27,12 +27,7 @@ from tqdm.auto import tqdm
|
||||
|
||||
from dataset import AllFutureHealthDataset, all_future_collate_fn
|
||||
from losses import build_loss
|
||||
from models import (
|
||||
EVENT_TRAJECTORY_ARCHITECTURE,
|
||||
MODEL_SIZE_NAMES,
|
||||
DeepHealth,
|
||||
resolve_model_size,
|
||||
)
|
||||
from models import TRAJ_MIXER_ARCHITECTURE, DeepHealth
|
||||
from targets import CHECKUP_IDX, PAD_IDX
|
||||
from train_util import (
|
||||
configure_torch_for_training,
|
||||
@@ -83,13 +78,10 @@ def parse_args() -> argparse.Namespace:
|
||||
parser.add_argument("--min_future_events", type=int, default=1)
|
||||
parser.add_argument("--validation_query_seed", type=int, default=None)
|
||||
|
||||
parser.add_argument(
|
||||
"--model_size",
|
||||
type=str,
|
||||
default="nano",
|
||||
choices=MODEL_SIZE_NAMES,
|
||||
)
|
||||
parser.add_argument("--n_reasoning_rounds", type=int, default=12)
|
||||
parser.add_argument("--n_embd", type=int, default=120)
|
||||
parser.add_argument("--n_head", type=int, default=10)
|
||||
parser.add_argument("--n_hist_layer", type=int, default=12)
|
||||
parser.add_argument("--n_tab_layer", type=int, default=4)
|
||||
parser.add_argument("--n_bins", type=int, default=16)
|
||||
parser.add_argument("--extra_pool_reduce", type=str, default="mean",
|
||||
choices=["mean", "sum"])
|
||||
@@ -154,8 +146,10 @@ def move_batch_to_device(batch: Dict[str, torch.Tensor], device: torch.device) -
|
||||
def build_model(args: argparse.Namespace, dataset: AllFutureHealthDataset) -> DeepHealth:
|
||||
return DeepHealth(
|
||||
vocab_size=dataset.vocab_size,
|
||||
model_size=args.model_size,
|
||||
n_reasoning_rounds=args.n_reasoning_rounds,
|
||||
n_embd=args.n_embd,
|
||||
n_head=args.n_head,
|
||||
n_hist_layer=args.n_hist_layer,
|
||||
n_tab_layer=args.n_tab_layer,
|
||||
n_types=dataset.n_types,
|
||||
n_cont_types=dataset.n_cont_types,
|
||||
n_categories=dataset.n_categories,
|
||||
@@ -300,17 +294,12 @@ def build_metadata(
|
||||
val_subset,
|
||||
test_subset,
|
||||
) -> Dict[str, Any]:
|
||||
size_config = resolve_model_size(args.model_size)
|
||||
return {
|
||||
"run_name": run_name,
|
||||
"dataset_class": "AllFutureHealthDataset",
|
||||
"collate_fn": "all_future_collate_fn",
|
||||
"model_class": "DeepHealth",
|
||||
"model_architecture": EVENT_TRAJECTORY_ARCHITECTURE,
|
||||
"d_model": size_config.d_model,
|
||||
"n_trajectory": size_config.n_trajectory,
|
||||
"trajectory_dim": size_config.trajectory_dim,
|
||||
"traj_hidden": size_config.traj_hidden,
|
||||
"model_architecture": TRAJ_MIXER_ARCHITECTURE,
|
||||
"model_target_mode": "all_future",
|
||||
"target_mode": "all_future",
|
||||
"dist_mode": args.dist_mode,
|
||||
@@ -348,26 +337,12 @@ def main() -> None:
|
||||
configure_torch_for_training(device)
|
||||
|
||||
run_dir, run_name = create_unique_run_dir(
|
||||
lambda timestamp: (
|
||||
f"{args.model_size}_r{args.n_reasoning_rounds}_"
|
||||
f"{args.time_mode}_{args.dist_mode}_"
|
||||
f"all_future_pure_disease_{timestamp}"
|
||||
)
|
||||
lambda timestamp: f"{args.time_mode}_{args.dist_mode}_all_future_pure_disease_{timestamp}"
|
||||
)
|
||||
logger = setup_logging(run_dir)
|
||||
|
||||
logger.info(f"Starting all-future training run: {run_name}")
|
||||
logger.info(f"Device: {device}")
|
||||
size_config = resolve_model_size(args.model_size)
|
||||
logger.info(
|
||||
"Model size: "
|
||||
f"{args.model_size} "
|
||||
f"(d_model={size_config.d_model}, "
|
||||
f"n_trajectory={size_config.n_trajectory}, "
|
||||
f"trajectory_dim={size_config.trajectory_dim}, "
|
||||
f"traj_hidden={size_config.traj_hidden}); "
|
||||
f"reasoning_rounds={args.n_reasoning_rounds}"
|
||||
)
|
||||
logger.info(f"extra_info_types: {format_extra_info_types(args.extra_info_types)}")
|
||||
|
||||
logger.info("Loading all-future datasets...")
|
||||
|
||||
@@ -24,13 +24,7 @@ from tqdm.auto import tqdm
|
||||
|
||||
from dataset import HealthDataset, collate_fn
|
||||
from losses import build_loss
|
||||
from models import (
|
||||
EVENT_TRAJECTORY_ARCHITECTURE,
|
||||
MODEL_SIZE_NAMES,
|
||||
DeepHealth,
|
||||
DeepHealthOutput,
|
||||
resolve_model_size,
|
||||
)
|
||||
from models import TRAJ_MIXER_ARCHITECTURE, DeepHealth, DeepHealthOutput
|
||||
from readouts import build_readout
|
||||
from targets import CHECKUP_IDX, NO_EVENT_IDX, PAD_IDX
|
||||
from train_util import (
|
||||
@@ -80,13 +74,10 @@ def parse_args() -> argparse.Namespace:
|
||||
parser.add_argument("--val_eid_file", type=str, default="ukb_val_eid.csv")
|
||||
parser.add_argument("--test_eid_file", type=str, default="ukb_test_eid.csv")
|
||||
|
||||
parser.add_argument(
|
||||
"--model_size",
|
||||
type=str,
|
||||
default="nano",
|
||||
choices=MODEL_SIZE_NAMES,
|
||||
)
|
||||
parser.add_argument("--n_reasoning_rounds", type=int, default=12)
|
||||
parser.add_argument("--n_embd", type=int, default=120)
|
||||
parser.add_argument("--n_head", type=int, default=10)
|
||||
parser.add_argument("--n_hist_layer", type=int, default=12)
|
||||
parser.add_argument("--n_tab_layer", type=int, default=4)
|
||||
parser.add_argument("--n_bins", type=int, default=16)
|
||||
parser.add_argument("--extra_pool_reduce", type=str, default="mean",
|
||||
choices=["mean", "sum"])
|
||||
@@ -160,8 +151,10 @@ def move_batch_to_device(batch: Dict[str, torch.Tensor], device: torch.device) -
|
||||
def build_model(args: argparse.Namespace, dataset: HealthDataset) -> DeepHealth:
|
||||
return DeepHealth(
|
||||
vocab_size=dataset.vocab_size,
|
||||
model_size=args.model_size,
|
||||
n_reasoning_rounds=args.n_reasoning_rounds,
|
||||
n_embd=args.n_embd,
|
||||
n_head=args.n_head,
|
||||
n_hist_layer=args.n_hist_layer,
|
||||
n_tab_layer=args.n_tab_layer,
|
||||
n_types=dataset.n_types,
|
||||
n_cont_types=dataset.n_cont_types,
|
||||
n_categories=dataset.n_categories,
|
||||
@@ -487,17 +480,12 @@ def build_metadata(
|
||||
val_subset,
|
||||
test_subset,
|
||||
) -> Dict[str, Any]:
|
||||
size_config = resolve_model_size(args.model_size)
|
||||
return {
|
||||
"run_name": run_name,
|
||||
"dataset_class": "NextStepHealthDataset",
|
||||
"collate_fn": "next_step_collate_fn",
|
||||
"model_class": "DeepHealth",
|
||||
"model_architecture": EVENT_TRAJECTORY_ARCHITECTURE,
|
||||
"d_model": size_config.d_model,
|
||||
"n_trajectory": size_config.n_trajectory,
|
||||
"trajectory_dim": size_config.trajectory_dim,
|
||||
"traj_hidden": size_config.traj_hidden,
|
||||
"model_architecture": TRAJ_MIXER_ARCHITECTURE,
|
||||
"model_target_mode": "next_token",
|
||||
"target_mode": args.target_mode,
|
||||
"dist_mode": "exponential",
|
||||
@@ -533,9 +521,7 @@ def main() -> None:
|
||||
|
||||
run_dir, run_name = create_unique_run_dir(
|
||||
lambda timestamp: (
|
||||
f"{args.model_size}_r{args.n_reasoning_rounds}_"
|
||||
f"{args.time_mode}_exponential_"
|
||||
f"next_token_{args.target_mode}_"
|
||||
f"{args.time_mode}_exponential_next_token_{args.target_mode}_"
|
||||
f"gap_{args.no_event_interval_years:g}y_{timestamp}"
|
||||
)
|
||||
)
|
||||
@@ -543,16 +529,6 @@ def main() -> None:
|
||||
|
||||
logger.info(f"Starting next-step training run: {run_name}")
|
||||
logger.info(f"Device: {device}")
|
||||
size_config = resolve_model_size(args.model_size)
|
||||
logger.info(
|
||||
"Model size: "
|
||||
f"{args.model_size} "
|
||||
f"(d_model={size_config.d_model}, "
|
||||
f"n_trajectory={size_config.n_trajectory}, "
|
||||
f"trajectory_dim={size_config.trajectory_dim}, "
|
||||
f"traj_hidden={size_config.traj_hidden}); "
|
||||
f"reasoning_rounds={args.n_reasoning_rounds}"
|
||||
)
|
||||
logger.info(f"extra_info_types: {format_extra_info_types(args.extra_info_types)}")
|
||||
logger.info(f"readout={args.readout_name}, target_mode={args.target_mode}")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user