Compare commits
3 Commits
6f7b5be405
...
codex/unif
| Author | SHA1 | Date | |
|---|---|---|---|
| b13db5e407 | |||
| 4526191fe1 | |||
| 3af823f2e1 |
@@ -1,389 +0,0 @@
|
|||||||
# Event–Trajectory Shared Reasoning Backbone
|
|
||||||
|
|
||||||
> 状态:**Frozen implementation baseline**
|
|
||||||
> 架构标识:`event_trajectory_shared_v2`
|
|
||||||
> 固化日期:**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 | 120 | 6 | 20 | 24 |
|
|
||||||
| tiny | 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 默认使用6个:
|
|
||||||
|
|
||||||
\[
|
|
||||||
S\in\mathbb{R}^{B\times Q\times n_{\mathrm{trajectory}}\times
|
|
||||||
d_{\mathrm{trajectory}}}.
|
|
||||||
\]
|
|
||||||
|
|
||||||
其中:
|
|
||||||
|
|
||||||
- all-future:\(Q=1\);
|
|
||||||
- next-token:\(Q=L\),所有查询位置并行计算。
|
|
||||||
|
|
||||||
定义可学习原型:
|
|
||||||
|
|
||||||
\[
|
|
||||||
P\in\mathbb{R}^{n_{\mathrm{trajectory}}\times d_{\mathrm{trajectory}}}.
|
|
||||||
\]
|
|
||||||
|
|
||||||
查询上下文经过投影并 reshape:
|
|
||||||
|
|
||||||
\[
|
|
||||||
C_Q
|
|
||||||
=\operatorname{QueryProjection}(\text{query features})
|
|
||||||
\in\mathbb{R}^{B\times Q\times n_{\mathrm{trajectory}}\times
|
|
||||||
d_{\mathrm{trajectory}}},
|
|
||||||
\]
|
|
||||||
|
|
||||||
\[
|
|
||||||
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_v2
|
|
||||||
model_size: nano
|
|
||||||
d_model: 120
|
|
||||||
n_trajectory: 6
|
|
||||||
trajectory_dim: 20
|
|
||||||
traj_hidden: 24
|
|
||||||
n_reasoning_rounds: 12
|
|
||||||
model_parameter_count: <runtime count>
|
|
||||||
trainable_parameter_count: <runtime count>
|
|
||||||
```
|
|
||||||
|
|
||||||
评估和导出入口必须同时验证:
|
|
||||||
|
|
||||||
1. `model_architecture` 完全匹配;
|
|
||||||
2. `model_size` 属于 `nano / tiny / 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{多轮状态依赖推理}
|
|
||||||
}
|
|
||||||
\]
|
|
||||||
64
MODEL_ARCHITECTURES.md
Normal file
64
MODEL_ARCHITECTURES.md
Normal file
@@ -0,0 +1,64 @@
|
|||||||
|
# Model architectures
|
||||||
|
|
||||||
|
DeepHealth uses one codebase for both supported history-block architectures.
|
||||||
|
Select the architecture explicitly when starting a training run:
|
||||||
|
|
||||||
|
| `model_architecture` | History block | Checkpoint fingerprint |
|
||||||
|
| --- | --- | --- |
|
||||||
|
| `transformer_ffn_v1` | Temporal attention + SwiGLU FFN | `blocks.*.mlp.w1/w2/w3` and `blocks.*.ln2` |
|
||||||
|
| `traj_mixer_v5` | Temporal attention + TrajMixer | `blocks.*.mlp.intra_*`, `gate_proj`, and `output_proj` |
|
||||||
|
|
||||||
|
`transformer_ffn_v1` is the CLI default; pass `traj_mixer_v5` explicitly for
|
||||||
|
TrajMixer runs.
|
||||||
|
|
||||||
|
## Training
|
||||||
|
|
||||||
|
Next-step example:
|
||||||
|
|
||||||
|
```powershell
|
||||||
|
python train_next_step.py --model_architecture traj_mixer_v5 --n_layer 12
|
||||||
|
```
|
||||||
|
|
||||||
|
All-future example:
|
||||||
|
|
||||||
|
```powershell
|
||||||
|
python train_all_future.py --model_architecture transformer_ffn_v1 --n_layer 12
|
||||||
|
```
|
||||||
|
|
||||||
|
New runs are separated by architecture:
|
||||||
|
|
||||||
|
```text
|
||||||
|
runs/
|
||||||
|
transformer_ffn_v1/
|
||||||
|
<run_name>/
|
||||||
|
traj_mixer_v5/
|
||||||
|
<run_name>/
|
||||||
|
```
|
||||||
|
|
||||||
|
Use `--runs_root` to place this structure under a different root. Existing run
|
||||||
|
directories are not moved or renamed.
|
||||||
|
|
||||||
|
Each generated `train_config.json` records `model_architecture`, total parameter
|
||||||
|
count, and trainable parameter count.
|
||||||
|
|
||||||
|
Both training entry points use the single `--n_layer` option to set the number
|
||||||
|
of history backbone blocks. The same value is passed to `DeepHealth.n_layer`
|
||||||
|
and saved as `n_layer` in `train_config.json`; it must be at least 1.
|
||||||
|
|
||||||
|
## Architecture validation
|
||||||
|
|
||||||
|
Evaluation resolves the architecture before constructing the model and always
|
||||||
|
loads weights with `strict=True`.
|
||||||
|
|
||||||
|
- Every config must include an explicit `model_architecture` marker.
|
||||||
|
- Checkpoint fingerprints are used to validate that the selected architecture
|
||||||
|
matches the stored weights.
|
||||||
|
- A config marker that conflicts with the checkpoint fingerprint raises an
|
||||||
|
error instead of silently choosing one architecture.
|
||||||
|
- Unsupported historical TrajMixer markers such as `traj_mixer_v2`,
|
||||||
|
`traj_mixer_v3`, and `traj_mixer_v4` are rejected.
|
||||||
|
- Checkpoints and configs created before architecture markers were introduced
|
||||||
|
are intentionally unsupported.
|
||||||
|
|
||||||
|
Project code should use the architecture factory rather than instantiate a
|
||||||
|
history block directly.
|
||||||
544
backbones.py
544
backbones.py
@@ -4,6 +4,12 @@ import torch
|
|||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
from model_architectures import (
|
||||||
|
TRAJ_MIXER_ARCHITECTURE,
|
||||||
|
TRANSFORMER_FFN_ARCHITECTURE,
|
||||||
|
resolve_model_architecture,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TimeRoPE(nn.Module):
|
class TimeRoPE(nn.Module):
|
||||||
def __init__(self, dim: int, base: float = 10000.0):
|
def __init__(self, dim: int, base: float = 10000.0):
|
||||||
@@ -29,14 +35,6 @@ class TimeRoPE(nn.Module):
|
|||||||
x2 = x[..., 1::2]
|
x2 = x[..., 1::2]
|
||||||
return torch.stack((-x2, x1), dim=-1).flatten(-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
|
@staticmethod
|
||||||
def apply_from_cache(
|
def apply_from_cache(
|
||||||
q: torch.Tensor,
|
q: torch.Tensor,
|
||||||
@@ -93,280 +91,364 @@ class GaussianRBFTimeBasis(nn.Module):
|
|||||||
)
|
)
|
||||||
return rbf_acts
|
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__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
d_model: int,
|
n_embd: int,
|
||||||
n_trajectory: int,
|
n_head: int,
|
||||||
n_rbf_bases: int = 16,
|
n_rbf_bases: int = 16,
|
||||||
use_time_rope: bool = False,
|
dropout: float = 0.0,
|
||||||
use_rbf_bias: bool = False,
|
use_time_rope: bool = True,
|
||||||
|
use_rbf_bias: bool = True,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
if d_model <= 0 or n_trajectory <= 0:
|
assert n_embd % n_head == 0, "n_embd must be divisible by n_head"
|
||||||
raise ValueError("d_model and n_trajectory must be positive")
|
self.n_head = n_head
|
||||||
if d_model % n_trajectory != 0:
|
self.d_head = n_embd // n_head
|
||||||
raise ValueError(
|
self.scale = 1.0 / math.sqrt(self.d_head)
|
||||||
"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
|
|
||||||
self.use_time_rope = use_time_rope
|
self.use_time_rope = use_time_rope
|
||||||
self.use_rbf_bias = use_rbf_bias
|
self.use_rbf_bias = use_rbf_bias
|
||||||
|
|
||||||
# q_proj acts on each slot independently and is shared across slots.
|
# QKV projection (fused for efficiency)
|
||||||
self.q_proj = nn.Linear(
|
self.qkv = nn.Linear(n_embd, 3 * n_embd, bias=False)
|
||||||
self.trajectory_dim,
|
# Output projection
|
||||||
self.trajectory_dim,
|
self.out_proj = nn.Linear(n_embd, n_embd, bias=False)
|
||||||
bias=False,
|
|
||||||
)
|
# Layer-specific projection from shared RBF basis activations to per-head attention bias.
|
||||||
self.k_proj = nn.Linear(d_model, d_model, bias=False)
|
self.rbf_proj = nn.Linear(n_rbf_bases, n_head, bias=False)
|
||||||
self.v_proj = nn.Linear(d_model, d_model, bias=False)
|
# Keep the initial RBF attention bias exactly zero through the
|
||||||
if use_rbf_bias:
|
# zero-initialized projection, while leaving that projection with a
|
||||||
self.rbf_proj = nn.Linear(
|
# live gradient from the first optimization step.
|
||||||
n_rbf_bases,
|
self.time_bias_scale = nn.Parameter(torch.tensor(1.0))
|
||||||
n_trajectory,
|
|
||||||
bias=False,
|
self.resid_drop = nn.Dropout(dropout)
|
||||||
)
|
|
||||||
self.time_bias_scale = nn.Parameter(torch.tensor(0.0))
|
|
||||||
else:
|
|
||||||
self.rbf_proj = None
|
|
||||||
self.register_parameter("time_bias_scale", None)
|
|
||||||
self.reset_parameters()
|
self.reset_parameters()
|
||||||
|
|
||||||
def reset_parameters(self) -> None:
|
def reset_parameters(self) -> None:
|
||||||
nn.init.normal_(self.q_proj.weight, mean=0.0, std=0.02)
|
"""Match the previous version's GPT-style weight initialization."""
|
||||||
nn.init.normal_(self.k_proj.weight, mean=0.0, std=0.02)
|
nn.init.normal_(self.qkv.weight, mean=0.0, std=0.02)
|
||||||
nn.init.normal_(self.v_proj.weight, mean=0.0, std=0.02)
|
nn.init.normal_(self.out_proj.weight, mean=0.0, std=0.02)
|
||||||
if self.rbf_proj is not None:
|
nn.init.zeros_(self.rbf_proj.weight)
|
||||||
# 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
|
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
trajectory_state: torch.Tensor,
|
x: torch.Tensor,
|
||||||
event_key_value: tuple[torch.Tensor, torch.Tensor],
|
rope_cache: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||||
event_invalid_mask: torch.Tensor,
|
|
||||||
query_rope_cache: tuple[torch.Tensor, torch.Tensor] | None = None,
|
|
||||||
rbf_cache: torch.Tensor | None = None,
|
rbf_cache: torch.Tensor | None = None,
|
||||||
|
attn_mask: torch.Tensor | None = None,
|
||||||
) -> torch.Tensor:
|
) -> 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 self.use_time_rope:
|
||||||
if query_rope_cache is None:
|
assert rope_cache is not None, "rope_cache must be provided when use_time_rope is True"
|
||||||
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
|
|
||||||
if self.use_rbf_bias:
|
if self.use_rbf_bias:
|
||||||
if rbf_cache is None or self.rbf_proj is None:
|
assert rbf_cache is not None, "rbf_cache must be provided when use_rbf_bias is True"
|
||||||
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
|
|
||||||
|
|
||||||
mask = event_invalid_mask.unsqueeze(2)
|
B, L, _ = x.shape
|
||||||
min_value = torch.finfo(scores.dtype).min
|
H, D = self.n_head, self.d_head
|
||||||
masked_scores = scores.masked_fill(mask, min_value)
|
|
||||||
weights = torch.softmax(masked_scores.float(), dim=-1).to(scores.dtype)
|
# --- QKV ----------------------------------------------------------
|
||||||
weights = weights.masked_fill(mask, 0.0)
|
qkv = self.qkv(x).reshape(B, L, 3, H, D).permute(2, 0, 3, 1, 4)
|
||||||
denominator = weights.sum(dim=-1, keepdim=True)
|
q, k, v = qkv.unbind(0) # each (B, H, L, D)
|
||||||
weights = weights / denominator.clamp_min(
|
|
||||||
torch.finfo(weights.dtype).eps
|
# --- 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):
|
class SwiGLU(nn.Module):
|
||||||
"""SwiGLU interaction along the trajectory axis only."""
|
def __init__(
|
||||||
|
self,
|
||||||
def __init__(self, n_trajectory: int, trajectory_dim: int):
|
n_embd: int,
|
||||||
|
hidden_dim: int | None = None,
|
||||||
|
dropout: float = 0.0,
|
||||||
|
bias: bool = True,
|
||||||
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
if n_trajectory <= 0 or trajectory_dim <= 0:
|
hidden_dim = hidden_dim if hidden_dim is not None else int(
|
||||||
raise ValueError("trajectory dimensions must be positive")
|
n_embd * 2.5)
|
||||||
self.n_trajectory = n_trajectory
|
|
||||||
self.trajectory_dim = trajectory_dim
|
self.w1 = nn.Linear(n_embd, hidden_dim, bias=bias) # gate path
|
||||||
self.traj_hidden = 4 * n_trajectory
|
self.w2 = nn.Linear(n_embd, hidden_dim, bias=bias) # value path
|
||||||
self.gate_proj = nn.Parameter(
|
# output projection
|
||||||
torch.empty(trajectory_dim, n_trajectory, self.traj_hidden)
|
self.w3 = nn.Linear(hidden_dim, n_embd, bias=bias)
|
||||||
)
|
self.drop = nn.Dropout(dropout)
|
||||||
self.value_proj = nn.Parameter(
|
|
||||||
torch.empty(trajectory_dim, n_trajectory, self.traj_hidden)
|
|
||||||
)
|
|
||||||
self.output_proj = nn.Parameter(
|
|
||||||
torch.empty(trajectory_dim, self.traj_hidden, n_trajectory)
|
|
||||||
)
|
|
||||||
self.reset_parameters()
|
self.reset_parameters()
|
||||||
|
|
||||||
def reset_parameters(self) -> None:
|
def reset_parameters(self) -> None:
|
||||||
for feature_idx in range(self.trajectory_dim):
|
"""GPT-style parameter initialization for MLP paths."""
|
||||||
|
nn.init.normal_(self.w1.weight, mean=0.0, std=0.02)
|
||||||
|
nn.init.normal_(self.w2.weight, mean=0.0, std=0.02)
|
||||||
|
nn.init.normal_(self.w3.weight, mean=0.0, std=0.02)
|
||||||
|
if self.w1.bias is not None:
|
||||||
|
nn.init.zeros_(self.w1.bias)
|
||||||
|
nn.init.zeros_(self.w2.bias)
|
||||||
|
nn.init.zeros_(self.w3.bias)
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""``(B, L, n_embd) -> (B, L, n_embd)``."""
|
||||||
|
return self.drop(self.w3(F.silu(self.w1(x)) * self.w2(x)))
|
||||||
|
|
||||||
|
|
||||||
|
class TrajMixer(nn.Module):
|
||||||
|
"""PreNorm gated mixing within and across latent trajectory groups.
|
||||||
|
|
||||||
|
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_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
|
||||||
|
self.n_group = n_head
|
||||||
|
self.d_group = n_embd // n_head
|
||||||
|
self.intra_hidden = 4 * self.d_group
|
||||||
|
self.hidden_group = 4 * n_head
|
||||||
|
|
||||||
|
self.norm = nn.LayerNorm(self.n_embd)
|
||||||
|
|
||||||
|
self.intra_gate_proj = nn.Parameter(
|
||||||
|
torch.empty(self.n_group, self.d_group, self.intra_hidden)
|
||||||
|
)
|
||||||
|
self.intra_value_proj = nn.Parameter(
|
||||||
|
torch.empty(self.n_group, self.d_group, self.intra_hidden)
|
||||||
|
)
|
||||||
|
self.intra_output_proj = nn.Parameter(
|
||||||
|
torch.empty(self.n_group, self.intra_hidden, self.d_group)
|
||||||
|
)
|
||||||
|
self.intra_gate_logits = nn.Parameter(
|
||||||
|
torch.empty(self.n_group, self.d_group)
|
||||||
|
)
|
||||||
|
|
||||||
|
self.gate_proj = nn.Parameter(
|
||||||
|
torch.empty(self.d_group, self.n_group, self.hidden_group)
|
||||||
|
)
|
||||||
|
self.value_proj = nn.Parameter(
|
||||||
|
torch.empty(self.d_group, self.n_group, self.hidden_group)
|
||||||
|
)
|
||||||
|
self.output_proj = nn.Parameter(
|
||||||
|
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 group_idx in range(self.n_group):
|
||||||
|
nn.init.xavier_uniform_(self.intra_gate_proj[group_idx])
|
||||||
|
nn.init.xavier_uniform_(self.intra_value_proj[group_idx])
|
||||||
|
nn.init.xavier_uniform_(self.intra_output_proj[group_idx])
|
||||||
|
nn.init.constant_(
|
||||||
|
self.intra_gate_logits,
|
||||||
|
math.log(0.1 / 0.9),
|
||||||
|
)
|
||||||
|
|
||||||
|
for feature_idx in range(self.d_group):
|
||||||
nn.init.xavier_uniform_(self.gate_proj[feature_idx])
|
nn.init.xavier_uniform_(self.gate_proj[feature_idx])
|
||||||
nn.init.xavier_uniform_(self.value_proj[feature_idx])
|
nn.init.xavier_uniform_(self.value_proj[feature_idx])
|
||||||
nn.init.normal_(self.output_proj, mean=0.0, std=1e-3)
|
nn.init.normal_(self.output_proj, mean=0.0, std=1e-3)
|
||||||
|
|
||||||
def forward(self, state: torch.Tensor) -> torch.Tensor:
|
def _intra_mix(self, grouped: torch.Tensor) -> torch.Tensor:
|
||||||
if state.shape[-2:] != (self.n_trajectory, self.trajectory_dim):
|
"""Mix features independently inside each residual-space group."""
|
||||||
raise ValueError(
|
intra_gate = torch.einsum(
|
||||||
"Expected trailing trajectory shape "
|
"blgd,gdh->blgh", grouped, self.intra_gate_proj
|
||||||
f"{(self.n_trajectory, self.trajectory_dim)}, got "
|
)
|
||||||
f"{tuple(state.shape[-2:])}"
|
intra_value = torch.einsum(
|
||||||
)
|
"blgd,gdh->blgh", grouped, self.intra_value_proj
|
||||||
gate = torch.einsum("...hr,rhk->...kr", state, self.gate_proj)
|
)
|
||||||
value = torch.einsum("...hr,rhk->...kr", state, self.value_proj)
|
intra_hidden = F.silu(intra_gate) * intra_value
|
||||||
|
return torch.einsum(
|
||||||
|
"blgh,ghd->blgd", intra_hidden, self.intra_output_proj
|
||||||
|
)
|
||||||
|
|
||||||
|
def _cross_mix(self, grouped: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""Mix groups independently for each within-group coordinate."""
|
||||||
|
gate = torch.einsum(
|
||||||
|
"blgr,rgh->blhr", grouped, self.gate_proj
|
||||||
|
)
|
||||||
|
value = torch.einsum(
|
||||||
|
"blgr,rgh->blhr", grouped, self.value_proj
|
||||||
|
)
|
||||||
hidden = F.silu(gate) * value
|
hidden = F.silu(gate) * value
|
||||||
return torch.einsum("...kr,rkh->...hr", hidden, self.output_proj)
|
return torch.einsum(
|
||||||
|
"blhr,rhg->blgr", hidden, self.output_proj
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""Apply one full-width PreNorm and one outer 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(
|
||||||
|
f"Expected hidden size {self.n_embd}, got {x.size(-1)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
batch_size, seq_len, _ = x.shape
|
||||||
|
grouped = self.norm(x).reshape(
|
||||||
|
batch_size, seq_len, self.n_group, self.d_group
|
||||||
|
)
|
||||||
|
|
||||||
|
intra_output = self._intra_mix(grouped)
|
||||||
|
intra_gate = torch.sigmoid(self.intra_gate_logits).view(
|
||||||
|
1, 1, self.n_group, self.d_group
|
||||||
|
)
|
||||||
|
mixed_input = grouped + intra_gate * intra_output
|
||||||
|
|
||||||
|
update = self._cross_mix(mixed_input).reshape(
|
||||||
|
batch_size, seq_len, self.n_embd
|
||||||
|
)
|
||||||
|
return x + self.drop(update)
|
||||||
|
|
||||||
|
|
||||||
class SharedEventTrajectoryCore(nn.Module):
|
class TransformerFFNBlock(nn.Module):
|
||||||
"""One parameter-shared reasoning core reused across all rounds."""
|
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
d_model: int,
|
n_embd: int,
|
||||||
n_trajectory: int,
|
n_head: int,
|
||||||
n_reasoning_rounds: int,
|
|
||||||
dropout: float = 0.0,
|
attn_dropout: float = 0.0,
|
||||||
n_rbf_bases: int = 16,
|
mlp_dropout: float = 0.0,
|
||||||
use_time_rope: bool = False,
|
use_time_rope: bool = False,
|
||||||
use_rbf_bias: bool = False,
|
use_rbf_bias: bool = False,
|
||||||
|
n_rbf_bases: int = 16,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
if n_reasoning_rounds <= 0:
|
self.attn = TemporalAttention(
|
||||||
raise ValueError("n_reasoning_rounds must be positive")
|
n_embd=n_embd,
|
||||||
if d_model <= 0 or n_trajectory <= 0:
|
n_head=n_head,
|
||||||
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,
|
|
||||||
n_rbf_bases=n_rbf_bases,
|
n_rbf_bases=n_rbf_bases,
|
||||||
|
dropout=attn_dropout,
|
||||||
use_time_rope=use_time_rope,
|
use_time_rope=use_time_rope,
|
||||||
use_rbf_bias=use_rbf_bias,
|
use_rbf_bias=use_rbf_bias,
|
||||||
)
|
)
|
||||||
self.norm_mixer = nn.LayerNorm(trajectory_dim)
|
self.mlp = SwiGLU(n_embd=n_embd, dropout=mlp_dropout)
|
||||||
self.traj_mixer = SharedTrajectoryMixer(
|
self.ln1 = nn.LayerNorm(n_embd)
|
||||||
n_trajectory=n_trajectory,
|
self.ln2 = nn.LayerNorm(n_embd)
|
||||||
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,
|
|
||||||
)
|
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
trajectory_state: torch.Tensor,
|
x: torch.Tensor,
|
||||||
event_key_value: tuple[torch.Tensor, torch.Tensor],
|
rope_cache: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||||
event_invalid_mask: torch.Tensor,
|
|
||||||
query_rope_cache: tuple[torch.Tensor, torch.Tensor] | None = None,
|
|
||||||
rbf_cache: torch.Tensor | None = None,
|
rbf_cache: torch.Tensor | None = None,
|
||||||
|
attn_mask: torch.Tensor | None = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
readout = self.cross_attention(
|
x = x + self.attn(self.ln1(x), rope_cache, rbf_cache, attn_mask)
|
||||||
trajectory_state=self.norm_attn(trajectory_state),
|
x = x + self.mlp(self.ln2(x))
|
||||||
event_key_value=event_key_value,
|
return x
|
||||||
event_invalid_mask=event_invalid_mask,
|
|
||||||
query_rope_cache=query_rope_cache,
|
|
||||||
rbf_cache=rbf_cache,
|
class TrajMixerBlock(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
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__()
|
||||||
|
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,
|
||||||
)
|
)
|
||||||
updated = (
|
self.mlp = TrajMixer(
|
||||||
trajectory_state
|
n_embd=n_embd,
|
||||||
+ self.attn_scale * self.dropout(readout)
|
n_head=n_head,
|
||||||
|
dropout=mlp_dropout,
|
||||||
)
|
)
|
||||||
mixed = self.traj_mixer(self.norm_mixer(updated))
|
self.ln1 = nn.LayerNorm(n_embd)
|
||||||
return updated + self.mixer_scale * self.dropout(mixed)
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
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:
|
||||||
|
x = x + self.attn(self.ln1(x), rope_cache, rbf_cache, attn_mask)
|
||||||
|
return self.mlp(x)
|
||||||
|
|
||||||
|
|
||||||
|
def build_backbone_block(
|
||||||
|
model_architecture: str,
|
||||||
|
*,
|
||||||
|
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,
|
||||||
|
) -> nn.Module:
|
||||||
|
"""Build one history block for a supported model architecture."""
|
||||||
|
architecture = resolve_model_architecture(model_architecture)
|
||||||
|
block_class: type[nn.Module]
|
||||||
|
if architecture == TRANSFORMER_FFN_ARCHITECTURE:
|
||||||
|
block_class = TransformerFFNBlock
|
||||||
|
elif architecture == TRAJ_MIXER_ARCHITECTURE:
|
||||||
|
block_class = TrajMixerBlock
|
||||||
|
else: # pragma: no cover - guarded by resolve_model_architecture.
|
||||||
|
raise ValueError(f"Unsupported model architecture: {architecture!r}")
|
||||||
|
return block_class(
|
||||||
|
n_embd=n_embd,
|
||||||
|
n_head=n_head,
|
||||||
|
attn_dropout=attn_dropout,
|
||||||
|
mlp_dropout=mlp_dropout,
|
||||||
|
use_time_rope=use_time_rope,
|
||||||
|
use_rbf_bias=use_rbf_bias,
|
||||||
|
n_rbf_bases=n_rbf_bases,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TokenAutoDiscretization(nn.Module):
|
class TokenAutoDiscretization(nn.Module):
|
||||||
|
|||||||
183
delphi2m_auc_report.py
Normal file
183
delphi2m_auc_report.py
Normal file
@@ -0,0 +1,183 @@
|
|||||||
|
"""Build Delphi2M-style sex-specific AUC reports.
|
||||||
|
|
||||||
|
The Delphi2M evaluation code uses 0.1 years for the no-gap evaluation. The
|
||||||
|
published report displays that point as 0 months, while retaining the actual
|
||||||
|
0.1-year evaluation period in this project's report output.
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Dict, Optional
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import pandas as pd
|
||||||
|
|
||||||
|
|
||||||
|
DEFAULT_DELPHI2M_PERIODS_YEARS = (0.1, 1.0, 5.0, 10.0)
|
||||||
|
|
||||||
|
_CHAPTER_SHORT_NAMES = {
|
||||||
|
"I": "I. Infectious Diseases",
|
||||||
|
"II": "II. Neoplasms",
|
||||||
|
"III": "III. Blood & Immune Disorders",
|
||||||
|
"IV": "IV. Metabolic Diseases",
|
||||||
|
"V": "V. Mental Disorders",
|
||||||
|
"VI": "VI. Nervous System Diseases",
|
||||||
|
"VII": "VII. Eye Diseases",
|
||||||
|
"VIII": "VIII. Ear Diseases",
|
||||||
|
"IX": "IX. Circulatory Diseases",
|
||||||
|
"X": "X. Respiratory Diseases",
|
||||||
|
"XI": "XI. Digestive Diseases",
|
||||||
|
"XII": "XII. Skin Diseases",
|
||||||
|
"XIII": "XIII. Musculoskeletal Diseases",
|
||||||
|
"XIV": "XIV. Genitourinary Diseases",
|
||||||
|
"XV": "XV. Pregnancy & Childbirth",
|
||||||
|
"XVI": "XVI. Perinatal Conditions",
|
||||||
|
"XVII": "XVII. Congenital Abnormalities",
|
||||||
|
"XVIII": "XVIII. Symptoms & Signs",
|
||||||
|
"XIX": "XIX. Injury & Poisoning",
|
||||||
|
"XX": "XX. External Causes",
|
||||||
|
"XXI": "XXI. Health Services",
|
||||||
|
"XXII": "XXII. Special Purposes",
|
||||||
|
"Death": "Death",
|
||||||
|
"Unmapped": "Unmapped",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _is_no_gap(period_years: float) -> bool:
|
||||||
|
return bool(np.isclose(float(period_years), 0.1, rtol=0.0, atol=1e-8))
|
||||||
|
|
||||||
|
|
||||||
|
def _canonical_period_years(period_years: float) -> float:
|
||||||
|
value = float(period_years)
|
||||||
|
for canonical in DEFAULT_DELPHI2M_PERIODS_YEARS:
|
||||||
|
if np.isclose(value, canonical, rtol=0.0, atol=1e-6):
|
||||||
|
return float(canonical)
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _gap_months(period_years: float) -> int:
|
||||||
|
if _is_no_gap(period_years):
|
||||||
|
return 0
|
||||||
|
return int(round(float(period_years) * 12.0))
|
||||||
|
|
||||||
|
|
||||||
|
def _gap_label(period_years: float) -> str:
|
||||||
|
if _is_no_gap(period_years):
|
||||||
|
return "No gap"
|
||||||
|
value = float(period_years)
|
||||||
|
value_text = f"{value:g}"
|
||||||
|
unit = "year" if np.isclose(value, 1.0) else "years"
|
||||||
|
return f"{value_text} {unit}"
|
||||||
|
|
||||||
|
|
||||||
|
def _load_chapter_by_code(
|
||||||
|
chapter_mapping_path: Optional[str | Path] = None,
|
||||||
|
) -> Dict[str, str]:
|
||||||
|
if chapter_mapping_path is None:
|
||||||
|
chapter_mapping_path = Path(__file__).with_name(
|
||||||
|
"icd10_chapter_organ_mapping.csv"
|
||||||
|
)
|
||||||
|
path = Path(chapter_mapping_path)
|
||||||
|
if not path.exists():
|
||||||
|
return {}
|
||||||
|
|
||||||
|
mapping = pd.read_csv(
|
||||||
|
path,
|
||||||
|
usecols=["code", "icd10_chapter"],
|
||||||
|
dtype={"code": str, "icd10_chapter": str},
|
||||||
|
)
|
||||||
|
mapping["code"] = mapping["code"].str.strip()
|
||||||
|
mapping["chapter"] = (
|
||||||
|
mapping["icd10_chapter"]
|
||||||
|
.str.strip()
|
||||||
|
.map(_CHAPTER_SHORT_NAMES)
|
||||||
|
.fillna("Unmapped")
|
||||||
|
)
|
||||||
|
return dict(zip(mapping["code"], mapping["chapter"]))
|
||||||
|
|
||||||
|
|
||||||
|
def build_delphi2m_auc_report(
|
||||||
|
df_unpooled: pd.DataFrame,
|
||||||
|
*,
|
||||||
|
period_col: str,
|
||||||
|
chapter_mapping_path: Optional[str | Path] = None,
|
||||||
|
) -> pd.DataFrame:
|
||||||
|
"""Aggregate age strata by sex and return a Delphi2M-style AUC report.
|
||||||
|
|
||||||
|
Required input columns are ``token``, ``label_code``, ``sex``,
|
||||||
|
``auc_delong``, and the supplied ``period_col`` (``offset`` or
|
||||||
|
``horizon``). The output begins with the five columns used by Delphi2M
|
||||||
|
Fig. 2e and then records the actual evaluation period and ICD-10 code.
|
||||||
|
"""
|
||||||
|
required = {"token", "label_code", "sex", "auc_delong", period_col}
|
||||||
|
missing = sorted(required - set(df_unpooled.columns))
|
||||||
|
if missing:
|
||||||
|
raise ValueError(
|
||||||
|
"Cannot build Delphi2M AUC report; missing columns: "
|
||||||
|
+ ", ".join(missing)
|
||||||
|
)
|
||||||
|
|
||||||
|
source = df_unpooled.loc[
|
||||||
|
:,
|
||||||
|
["token", "label_code", "sex", "auc_delong", period_col],
|
||||||
|
].copy()
|
||||||
|
source["sex"] = source["sex"].astype(str).str.strip().str.lower()
|
||||||
|
source = source[source["sex"].isin(["female", "male"])]
|
||||||
|
source["auc_delong"] = pd.to_numeric(
|
||||||
|
source["auc_delong"], errors="coerce"
|
||||||
|
)
|
||||||
|
source[period_col] = pd.to_numeric(source[period_col], errors="coerce")
|
||||||
|
source = source.dropna(subset=[period_col, "auc_delong"])
|
||||||
|
source[period_col] = source[period_col].map(_canonical_period_years)
|
||||||
|
|
||||||
|
if source.empty:
|
||||||
|
raise ValueError("Cannot build Delphi2M AUC report from empty AUC data.")
|
||||||
|
|
||||||
|
grouped = (
|
||||||
|
source.groupby(
|
||||||
|
["token", "label_code", period_col, "sex"],
|
||||||
|
dropna=False,
|
||||||
|
as_index=False,
|
||||||
|
)
|
||||||
|
.agg(auc=("auc_delong", "mean"))
|
||||||
|
)
|
||||||
|
report = (
|
||||||
|
grouped.pivot(
|
||||||
|
index=["token", "label_code", period_col],
|
||||||
|
columns="sex",
|
||||||
|
values="auc",
|
||||||
|
)
|
||||||
|
.reset_index()
|
||||||
|
.rename_axis(columns=None)
|
||||||
|
.rename(columns={"female": "Female", "male": "Male"})
|
||||||
|
)
|
||||||
|
for col in ["Female", "Male"]:
|
||||||
|
if col not in report.columns:
|
||||||
|
report[col] = np.nan
|
||||||
|
|
||||||
|
chapter_by_code = _load_chapter_by_code(chapter_mapping_path)
|
||||||
|
report["chapter"] = (
|
||||||
|
report["label_code"].astype(str).map(chapter_by_code).fillna("Unmapped")
|
||||||
|
)
|
||||||
|
report["Gap, months"] = report[period_col].map(_gap_months).astype("Int64")
|
||||||
|
report["Gap label"] = report[period_col].map(_gap_label)
|
||||||
|
report["icd10"] = pd.to_numeric(report["token"], errors="coerce").astype(
|
||||||
|
"Int64"
|
||||||
|
)
|
||||||
|
|
||||||
|
report = report.sort_values(
|
||||||
|
["icd10", period_col], kind="stable", ignore_index=True
|
||||||
|
)
|
||||||
|
return report.loc[
|
||||||
|
:,
|
||||||
|
[
|
||||||
|
"Gap, months",
|
||||||
|
"chapter",
|
||||||
|
"icd10",
|
||||||
|
"Female",
|
||||||
|
"Male",
|
||||||
|
period_col,
|
||||||
|
"Gap label",
|
||||||
|
"label_code",
|
||||||
|
],
|
||||||
|
]
|
||||||
@@ -7,7 +7,8 @@ This script follows the logic of the Delphi evaluation script supplied by the us
|
|||||||
at least `offset` years before the target time;
|
at least `offset` years before the target time;
|
||||||
3. run model inference by disease chunks to avoid materializing all logits;
|
3. run model inference by disease chunks to avoid materializing all logits;
|
||||||
4. compute AUC separately by sex and age bracket;
|
4. compute AUC separately by sex and age bracket;
|
||||||
5. aggregate age brackets with DeLong variance.
|
5. average age-bracket AUCs within each sex and write a Delphi2M-style
|
||||||
|
Female/Male report.
|
||||||
|
|
||||||
Efficiency notes:
|
Efficiency notes:
|
||||||
- transformer/readout inference is executed once and cached;
|
- transformer/readout inference is executed once and cached;
|
||||||
@@ -39,12 +40,13 @@ from torch.utils.data import DataLoader, Subset
|
|||||||
from tqdm.auto import tqdm
|
from tqdm.auto import tqdm
|
||||||
|
|
||||||
from dataset import HealthDataset
|
from dataset import HealthDataset
|
||||||
from eval_data import load_sequence_eval_dataset, sequence_eval_collate_fn
|
from delphi2m_auc_report import (
|
||||||
from models import (
|
DEFAULT_DELPHI2M_PERIODS_YEARS,
|
||||||
DeepHealth,
|
build_delphi2m_auc_report,
|
||||||
validate_event_trajectory_config,
|
|
||||||
validate_event_trajectory_state_dict,
|
|
||||||
)
|
)
|
||||||
|
from eval_data import load_sequence_eval_dataset, sequence_eval_collate_fn
|
||||||
|
from model_architectures import resolve_model_architecture
|
||||||
|
from models import DeepHealth
|
||||||
from readouts import build_readout
|
from readouts import build_readout
|
||||||
from targets import PAD_IDX, CHECKUP_IDX, NO_EVENT_IDX
|
from targets import PAD_IDX, CHECKUP_IDX, NO_EVENT_IDX
|
||||||
|
|
||||||
@@ -312,20 +314,24 @@ def split_indices(n: int, train_ratio: float, val_ratio: float, test_ratio: floa
|
|||||||
return idx[:n_train], idx[n_train:n_train + n_val], idx[n_train + n_val:]
|
return idx[:n_train], idx[n_train:n_train + n_val], idx[n_train + n_val:]
|
||||||
|
|
||||||
|
|
||||||
def build_model_from_dataset(args: argparse.Namespace, cfg: Dict[str, Any], dataset: HealthDataset) -> DeepHealth:
|
def build_model_from_dataset(
|
||||||
validate_event_trajectory_config(cfg)
|
args: argparse.Namespace,
|
||||||
|
cfg: Dict[str, Any],
|
||||||
|
dataset: HealthDataset,
|
||||||
|
state_dict: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> DeepHealth:
|
||||||
model_target_mode = str(cfg_get(
|
model_target_mode = str(cfg_get(
|
||||||
args, cfg, "model_target_mode", "next_token")).lower()
|
args, cfg, "model_target_mode", "next_token")).lower()
|
||||||
if model_target_mode not in {"next_token", "all_future"}:
|
if model_target_mode not in {"next_token", "all_future"}:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"model_target_mode must be next_token or all_future, got {model_target_mode!r}"
|
f"model_target_mode must be next_token or all_future, got {model_target_mode!r}"
|
||||||
)
|
)
|
||||||
|
model_architecture = resolve_model_architecture(cfg, state_dict)
|
||||||
return DeepHealth(
|
return DeepHealth(
|
||||||
vocab_size=dataset.vocab_size,
|
vocab_size=dataset.vocab_size,
|
||||||
model_size=str(cfg_get(args, cfg, "model_size", "nano")),
|
n_embd=int(cfg_get(args, cfg, "n_embd", 120)),
|
||||||
n_reasoning_rounds=int(
|
n_head=int(cfg_get(args, cfg, "n_head", 10)),
|
||||||
cfg_get(args, cfg, "n_reasoning_rounds", 12)
|
n_layer=int(cfg["n_layer"]),
|
||||||
),
|
|
||||||
n_types=dataset.n_types,
|
n_types=dataset.n_types,
|
||||||
n_cont_types=dataset.n_cont_types,
|
n_cont_types=dataset.n_cont_types,
|
||||||
n_categories=dataset.n_categories,
|
n_categories=dataset.n_categories,
|
||||||
@@ -336,6 +342,7 @@ def build_model_from_dataset(args: argparse.Namespace, cfg: Dict[str, Any], data
|
|||||||
time_mode=str(cfg_get(args, cfg, "time_mode", "relative")),
|
time_mode=str(cfg_get(args, cfg, "time_mode", "relative")),
|
||||||
dist_mode=str(cfg_get(args, cfg, "dist_mode", "exponential")),
|
dist_mode=str(cfg_get(args, cfg, "dist_mode", "exponential")),
|
||||||
dropout=float(cfg_get(args, cfg, "dropout", 0.0)),
|
dropout=float(cfg_get(args, cfg, "dropout", 0.0)),
|
||||||
|
model_architecture=model_architecture,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -383,7 +390,7 @@ def resolve_dist_mode_for_checkpoint(cfg_dist_mode: str, state_dict: Dict[str, A
|
|||||||
|
|
||||||
|
|
||||||
def load_model_state(
|
def load_model_state(
|
||||||
model: torch.nn.Module,
|
model: DeepHealth,
|
||||||
checkpoint_path: str,
|
checkpoint_path: str,
|
||||||
device: torch.device,
|
device: torch.device,
|
||||||
state_dict: Optional[Dict[str, Any]] = None,
|
state_dict: Optional[Dict[str, Any]] = None,
|
||||||
@@ -391,12 +398,7 @@ def load_model_state(
|
|||||||
state = state_dict if state_dict is not None else load_checkpoint_state_dict(
|
state = state_dict if state_dict is not None else load_checkpoint_state_dict(
|
||||||
checkpoint_path, map_location=device)
|
checkpoint_path, map_location=device)
|
||||||
|
|
||||||
validate_event_trajectory_state_dict(
|
resolve_model_architecture(model.model_architecture, state)
|
||||||
state,
|
|
||||||
expected_d_model=model.d_model,
|
|
||||||
expected_n_trajectory=model.n_trajectory,
|
|
||||||
expected_n_reasoning_rounds=model.n_reasoning_rounds,
|
|
||||||
)
|
|
||||||
model.load_state_dict(state, strict=True)
|
model.load_state_dict(state, strict=True)
|
||||||
|
|
||||||
|
|
||||||
@@ -533,7 +535,7 @@ def infer_readout_hidden(
|
|||||||
hidden = torch.zeros(
|
hidden = torch.zeros(
|
||||||
batch_size,
|
batch_size,
|
||||||
seq_len,
|
seq_len,
|
||||||
model.d_model,
|
model.n_embd,
|
||||||
device=event_seq.device,
|
device=event_seq.device,
|
||||||
dtype=torch.float32,
|
dtype=torch.float32,
|
||||||
)
|
)
|
||||||
@@ -1169,30 +1171,23 @@ def evaluate_auc_pipeline(
|
|||||||
df_auc_unpooled["label_code"] = df_auc_unpooled["token"].map(
|
df_auc_unpooled["label_code"] = df_auc_unpooled["token"].map(
|
||||||
dataset.label_id_to_code)
|
dataset.label_id_to_code)
|
||||||
|
|
||||||
print("Using DeLong method to calculate AUC confidence intervals.")
|
print(
|
||||||
grouped = df_auc_unpooled.groupby(
|
"Building Delphi2M-style report: mean AUC across age strata, "
|
||||||
["token", "label_code", "offset"], dropna=False, as_index=False)
|
"reported separately for Female and Male."
|
||||||
df_auc = grouped.agg(
|
|
||||||
auc=("auc_delong", "mean"),
|
|
||||||
n_strata=("auc_delong", "size"),
|
|
||||||
n_diseased=("n_diseased", "sum"),
|
|
||||||
n_healthy=("n_healthy", "sum"),
|
|
||||||
auc_variance_sum=("auc_variance_delong", "sum"),
|
|
||||||
)
|
)
|
||||||
df_auc["auc_variance_delong"] = (
|
df_report = build_delphi2m_auc_report(
|
||||||
df_auc["auc_variance_sum"]
|
df_auc_unpooled,
|
||||||
/ (df_auc["n_strata"].clip(lower=1).astype(np.float64) ** 2)
|
period_col="offset",
|
||||||
)
|
)
|
||||||
df_auc = df_auc.drop(columns=["auc_variance_sum"])
|
|
||||||
|
|
||||||
if output_path is not None:
|
if output_path is not None:
|
||||||
out_dir = Path(output_path)
|
out_dir = Path(output_path)
|
||||||
out_dir.mkdir(parents=True, exist_ok=True)
|
out_dir.mkdir(parents=True, exist_ok=True)
|
||||||
df_auc.to_csv(out_dir / "df_both.csv", index=False)
|
report_path = out_dir / "df_auc_delphi2m_report.csv"
|
||||||
df_auc_unpooled.to_csv(
|
df_report.to_csv(report_path, index=False)
|
||||||
out_dir / "df_auc_unpooled.csv", index=False)
|
print(f"Saved Delphi2M-style AUC report: {report_path}")
|
||||||
|
|
||||||
return df_auc_unpooled, df_auc
|
return df_auc_unpooled, df_report
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -1242,8 +1237,18 @@ def make_auc_offsets(args: argparse.Namespace, cfg: Dict[str, Any]) -> List[floa
|
|||||||
if explicit_offsets is not None:
|
if explicit_offsets is not None:
|
||||||
base_offsets = explicit_offsets
|
base_offsets = explicit_offsets
|
||||||
else:
|
else:
|
||||||
next_token_offset = float(cfg_get(args, cfg, "offset", 0.1))
|
next_token_offset = float(
|
||||||
base_offsets = [next_token_offset, 1.0, 5.0, 10.0]
|
cfg_get(
|
||||||
|
args,
|
||||||
|
cfg,
|
||||||
|
"offset",
|
||||||
|
DEFAULT_DELPHI2M_PERIODS_YEARS[0],
|
||||||
|
)
|
||||||
|
)
|
||||||
|
base_offsets = [
|
||||||
|
next_token_offset,
|
||||||
|
*DEFAULT_DELPHI2M_PERIODS_YEARS[1:],
|
||||||
|
]
|
||||||
|
|
||||||
offsets: List[float] = []
|
offsets: List[float] = []
|
||||||
seen = set()
|
seen = set()
|
||||||
@@ -1291,9 +1296,9 @@ def main() -> None:
|
|||||||
parser.add_argument("--filter_min_total", type=int, default=None,
|
parser.add_argument("--filter_min_total", type=int, default=None,
|
||||||
help="Minimum metadata count for disease selection; default 0.")
|
help="Minimum metadata count for disease selection; default 0.")
|
||||||
parser.add_argument("--offset", type=float, default=None,
|
parser.add_argument("--offset", type=float, default=None,
|
||||||
help="Next-token prediction offset in years; preserved and evaluated alongside 1, 5, and 10 years by default.")
|
help="Next-token prediction offset in years; 0.1 is Delphi2M no gap and is evaluated alongside 1, 5, and 10 years by default.")
|
||||||
parser.add_argument("--offsets", type=str, default=None,
|
parser.add_argument("--offsets", type=str, default=None,
|
||||||
help="Comma-separated prediction offsets in years. Overrides the default set of offset,1,5,10.")
|
help="Comma-separated prediction offsets in years. Overrides the default set of 0.1,1,5,10.")
|
||||||
parser.add_argument("--age_start", type=float, default=None)
|
parser.add_argument("--age_start", type=float, default=None)
|
||||||
parser.add_argument("--age_stop", type=float, default=None)
|
parser.add_argument("--age_stop", type=float, default=None)
|
||||||
parser.add_argument("--age_step", type=float, default=None)
|
parser.add_argument("--age_step", type=float, default=None)
|
||||||
@@ -1374,14 +1379,19 @@ def main() -> None:
|
|||||||
cfg = dict(cfg)
|
cfg = dict(cfg)
|
||||||
cfg["dist_mode"] = dist_mode
|
cfg["dist_mode"] = dist_mode
|
||||||
cfg["model_target_mode"] = model_target_mode
|
cfg["model_target_mode"] = model_target_mode
|
||||||
|
model_architecture = resolve_model_architecture(cfg, state_dict)
|
||||||
|
cfg["model_architecture"] = model_architecture
|
||||||
print(f"Resolved dist_mode for evaluation: {dist_mode}")
|
print(f"Resolved dist_mode for evaluation: {dist_mode}")
|
||||||
|
print(f"Resolved model architecture: {model_architecture}")
|
||||||
print(f"Model target mode for AUC: {model_target_mode}")
|
print(f"Model target mode for AUC: {model_target_mode}")
|
||||||
print(
|
print(
|
||||||
"AUC score semantics: evaluate_auc.py uses disease-specific eta/logit scores; "
|
"AUC score semantics: evaluate_auc.py uses disease-specific eta/logit scores; "
|
||||||
"dist_mode affects model loading but is not converted to horizon-specific risk probability."
|
"dist_mode affects model loading but is not converted to horizon-specific risk probability."
|
||||||
)
|
)
|
||||||
|
|
||||||
model = build_model_from_dataset(args, cfg, dataset).to(device)
|
model = build_model_from_dataset(
|
||||||
|
args, cfg, dataset, state_dict=state_dict
|
||||||
|
).to(device)
|
||||||
load_model_state(model, str(model_ckpt_path),
|
load_model_state(model, str(model_ckpt_path),
|
||||||
device, state_dict=state_dict)
|
device, state_dict=state_dict)
|
||||||
model.eval()
|
model.eval()
|
||||||
|
|||||||
@@ -3,6 +3,9 @@
|
|||||||
This script supports DeepHealth fixed-horizon risk scores for exponential,
|
This script supports DeepHealth fixed-horizon risk scores for exponential,
|
||||||
Weibull, and mixed all-future distributions.
|
Weibull, and mixed all-future distributions.
|
||||||
|
|
||||||
|
The default horizons are 0.1, 1, 5, and 10 years. As in Delphi2M, 0.1 years
|
||||||
|
is reported as the no-gap evaluation.
|
||||||
|
|
||||||
Landmark querying depends on the model target mode saved in train_config.json:
|
Landmark querying depends on the model target mode saved in train_config.json:
|
||||||
- next_token: insert a <NO_EVENT> token at landmark age and read it out;
|
- next_token: insert a <NO_EVENT> token at landmark age and read it out;
|
||||||
- all_future: pass landmark age directly as t_query.
|
- all_future: pass landmark age directly as t_query.
|
||||||
@@ -28,12 +31,13 @@ from torch.utils.data import DataLoader, Dataset
|
|||||||
from tqdm.auto import tqdm
|
from tqdm.auto import tqdm
|
||||||
|
|
||||||
from dataset import HealthDataset
|
from dataset import HealthDataset
|
||||||
from eval_data import load_sequence_eval_dataset
|
from delphi2m_auc_report import (
|
||||||
from models import (
|
DEFAULT_DELPHI2M_PERIODS_YEARS,
|
||||||
DeepHealth,
|
build_delphi2m_auc_report,
|
||||||
validate_event_trajectory_config,
|
|
||||||
validate_event_trajectory_state_dict,
|
|
||||||
)
|
)
|
||||||
|
from eval_data import load_sequence_eval_dataset
|
||||||
|
from model_architectures import resolve_model_architecture
|
||||||
|
from models import DeepHealth
|
||||||
from readouts import build_readout
|
from readouts import build_readout
|
||||||
from targets import CHECKUP_IDX, NO_EVENT_IDX, PAD_IDX
|
from targets import CHECKUP_IDX, NO_EVENT_IDX, PAD_IDX
|
||||||
|
|
||||||
@@ -181,20 +185,24 @@ def resolve_dist_mode_for_checkpoint(cfg_dist_mode: str, state_dict: Dict[str, A
|
|||||||
return mode if mode in {"exponential", "weibull", "mixed"} else "exponential"
|
return mode if mode in {"exponential", "weibull", "mixed"} else "exponential"
|
||||||
|
|
||||||
|
|
||||||
def build_model_from_dataset(args: argparse.Namespace, cfg: Dict[str, Any], dataset: HealthDataset) -> DeepHealth:
|
def build_model_from_dataset(
|
||||||
validate_event_trajectory_config(cfg)
|
args: argparse.Namespace,
|
||||||
|
cfg: Dict[str, Any],
|
||||||
|
dataset: HealthDataset,
|
||||||
|
state_dict: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> DeepHealth:
|
||||||
model_target_mode = str(cfg_get(
|
model_target_mode = str(cfg_get(
|
||||||
args, cfg, "model_target_mode", "next_token")).lower()
|
args, cfg, "model_target_mode", "next_token")).lower()
|
||||||
if model_target_mode not in {"next_token", "all_future"}:
|
if model_target_mode not in {"next_token", "all_future"}:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"model_target_mode must be next_token or all_future, got {model_target_mode!r}"
|
f"model_target_mode must be next_token or all_future, got {model_target_mode!r}"
|
||||||
)
|
)
|
||||||
|
model_architecture = resolve_model_architecture(cfg, state_dict)
|
||||||
return DeepHealth(
|
return DeepHealth(
|
||||||
vocab_size=dataset.vocab_size,
|
vocab_size=dataset.vocab_size,
|
||||||
model_size=str(cfg_get(args, cfg, "model_size", "nano")),
|
n_embd=int(cfg_get(args, cfg, "n_embd", 120)),
|
||||||
n_reasoning_rounds=int(
|
n_head=int(cfg_get(args, cfg, "n_head", 10)),
|
||||||
cfg_get(args, cfg, "n_reasoning_rounds", 12)
|
n_layer=int(cfg["n_layer"]),
|
||||||
),
|
|
||||||
n_types=dataset.n_types,
|
n_types=dataset.n_types,
|
||||||
n_cont_types=dataset.n_cont_types,
|
n_cont_types=dataset.n_cont_types,
|
||||||
n_categories=dataset.n_categories,
|
n_categories=dataset.n_categories,
|
||||||
@@ -205,16 +213,12 @@ def build_model_from_dataset(args: argparse.Namespace, cfg: Dict[str, Any], data
|
|||||||
time_mode=str(cfg_get(args, cfg, "time_mode", "relative")),
|
time_mode=str(cfg_get(args, cfg, "time_mode", "relative")),
|
||||||
dist_mode=str(cfg_get(args, cfg, "dist_mode", "exponential")),
|
dist_mode=str(cfg_get(args, cfg, "dist_mode", "exponential")),
|
||||||
dropout=float(cfg_get(args, cfg, "dropout", 0.0)),
|
dropout=float(cfg_get(args, cfg, "dropout", 0.0)),
|
||||||
|
model_architecture=model_architecture,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def load_model_state(model: torch.nn.Module, state_dict: Dict[str, Any]) -> None:
|
def load_model_state(model: DeepHealth, state_dict: Dict[str, Any]) -> None:
|
||||||
validate_event_trajectory_state_dict(
|
resolve_model_architecture(model.model_architecture, state_dict)
|
||||||
state_dict,
|
|
||||||
expected_d_model=model.d_model,
|
|
||||||
expected_n_trajectory=model.n_trajectory,
|
|
||||||
expected_n_reasoning_rounds=model.n_reasoning_rounds,
|
|
||||||
)
|
|
||||||
model.load_state_dict(state_dict, strict=True)
|
model.load_state_dict(state_dict, strict=True)
|
||||||
|
|
||||||
|
|
||||||
@@ -335,44 +339,6 @@ def _first_existing_column(df: pd.DataFrame, candidates: Sequence[str]) -> Optio
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
def build_metadata_for_merge(dataset: HealthDataset, labels_meta: Optional[pd.DataFrame]) -> pd.DataFrame:
|
|
||||||
base_rows = []
|
|
||||||
for token, code in dataset.label_id_to_code.items():
|
|
||||||
token = int(token)
|
|
||||||
code_text = str(code)
|
|
||||||
if token in SPECIAL_TOKENS or code_text.startswith("<"):
|
|
||||||
continue
|
|
||||||
base_rows.append({"token": token, "label_code": code_text})
|
|
||||||
base = pd.DataFrame(base_rows)
|
|
||||||
if labels_meta is None or labels_meta.empty:
|
|
||||||
return base
|
|
||||||
|
|
||||||
meta = labels_meta.copy()
|
|
||||||
code_col = _first_existing_column(
|
|
||||||
meta, ["Name", "code", "ICD10", "icd10", "label", "token", "disease_code"])
|
|
||||||
if code_col is not None:
|
|
||||||
meta["_label_code"] = meta[code_col].astype(
|
|
||||||
str).map(lambda s: s.split()[0].strip())
|
|
||||||
merged = base.merge(meta, left_on="label_code",
|
|
||||||
right_on="_label_code", how="left")
|
|
||||||
return merged.drop(columns=["_label_code"], errors="ignore")
|
|
||||||
|
|
||||||
if "index" in meta.columns:
|
|
||||||
idx = pd.to_numeric(meta["index"], errors="coerce")
|
|
||||||
has_no_event = (
|
|
||||||
NO_EVENT_IDX in dataset.label_id_to_code
|
|
||||||
and dataset.label_id_to_code.get(NO_EVENT_IDX) == "<NO_EVENT>"
|
|
||||||
)
|
|
||||||
if has_no_event:
|
|
||||||
idx = idx.where(idx < NO_EVENT_IDX, idx + 1)
|
|
||||||
meta["_index_int"] = idx.astype("Int64")
|
|
||||||
merged = base.merge(meta, left_on="token",
|
|
||||||
right_on="_index_int", how="left")
|
|
||||||
return merged.drop(columns=["_index_int"], errors="ignore")
|
|
||||||
|
|
||||||
return base
|
|
||||||
|
|
||||||
|
|
||||||
def _metadata_count_map(dataset: HealthDataset, labels_meta: Optional[pd.DataFrame]) -> Dict[int, float]:
|
def _metadata_count_map(dataset: HealthDataset, labels_meta: Optional[pd.DataFrame]) -> Dict[int, float]:
|
||||||
if labels_meta is None or labels_meta.empty or "count" not in labels_meta.columns:
|
if labels_meta is None or labels_meta.empty or "count" not in labels_meta.columns:
|
||||||
return {}
|
return {}
|
||||||
@@ -1112,7 +1078,6 @@ def evaluate_landmark_auc(
|
|||||||
loader: DataLoader,
|
loader: DataLoader,
|
||||||
landmark_dataset: LandmarkDataset,
|
landmark_dataset: LandmarkDataset,
|
||||||
output_path: Path,
|
output_path: Path,
|
||||||
labels_meta: Optional[pd.DataFrame],
|
|
||||||
disease_ids: Sequence[int],
|
disease_ids: Sequence[int],
|
||||||
disease_chunk_size: int,
|
disease_chunk_size: int,
|
||||||
score_mode: str,
|
score_mode: str,
|
||||||
@@ -1129,7 +1094,6 @@ def evaluate_landmark_auc(
|
|||||||
use_amp: bool,
|
use_amp: bool,
|
||||||
hidden_cache_dtype: str,
|
hidden_cache_dtype: str,
|
||||||
logit_batch_size: int,
|
logit_batch_size: int,
|
||||||
meta_info: Dict[str, Any],
|
|
||||||
) -> Tuple[pd.DataFrame, pd.DataFrame]:
|
) -> Tuple[pd.DataFrame, pd.DataFrame]:
|
||||||
model.eval().to(device)
|
model.eval().to(device)
|
||||||
|
|
||||||
@@ -1246,54 +1210,21 @@ def evaluate_landmark_auc(
|
|||||||
df_unpooled["label_code"] = df_unpooled["token"].map(
|
df_unpooled["label_code"] = df_unpooled["token"].map(
|
||||||
landmark_dataset.dataset.label_id_to_code)
|
landmark_dataset.dataset.label_id_to_code)
|
||||||
|
|
||||||
for k, v in meta_info.items():
|
print(
|
||||||
df_unpooled[k] = v
|
"Building Delphi2M-style report: mean AUC across landmark-age "
|
||||||
|
"strata, reported separately for Female and Male."
|
||||||
meta_table = build_metadata_for_merge(landmark_dataset.dataset, labels_meta)
|
|
||||||
df_unpooled = df_unpooled.merge(
|
|
||||||
meta_table, on=["token", "label_code"], how="left")
|
|
||||||
|
|
||||||
grouped = df_unpooled.groupby(
|
|
||||||
["token", "label_code", "horizon"], dropna=False, as_index=False)
|
|
||||||
df_merged = grouped.agg(
|
|
||||||
auc=("auc_delong", "mean"),
|
|
||||||
n_strata=("auc_delong", "size"),
|
|
||||||
n_diseased=("n_diseased", "sum"),
|
|
||||||
n_healthy=("n_healthy", "sum"),
|
|
||||||
auc_variance_sum=("auc_variance_delong", "sum"),
|
|
||||||
)
|
)
|
||||||
df_merged["auc_variance_delong"] = (
|
df_report = build_delphi2m_auc_report(
|
||||||
df_merged["auc_variance_sum"]
|
df_unpooled,
|
||||||
/ (df_merged["n_strata"].clip(lower=1).astype(np.float64) ** 2)
|
period_col="horizon",
|
||||||
)
|
)
|
||||||
df_merged = df_merged.drop(columns=["auc_variance_sum"])
|
|
||||||
|
|
||||||
keep_meta = [
|
|
||||||
c for c in [
|
|
||||||
"model_ckpt_path",
|
|
||||||
"config_path",
|
|
||||||
"target_mode",
|
|
||||||
"model_target_mode",
|
|
||||||
"dist_mode",
|
|
||||||
"time_mode",
|
|
||||||
"attn_mask_mode",
|
|
||||||
"readout_name",
|
|
||||||
"landmark_query_mode",
|
|
||||||
"landmark_token_mode",
|
|
||||||
"score_mode",
|
|
||||||
"eval_split",
|
|
||||||
]
|
|
||||||
if c in df_unpooled.columns
|
|
||||||
]
|
|
||||||
for col in keep_meta:
|
|
||||||
df_merged[col] = meta_info[col]
|
|
||||||
|
|
||||||
output_path.mkdir(parents=True, exist_ok=True)
|
output_path.mkdir(parents=True, exist_ok=True)
|
||||||
df_unpooled.to_csv(
|
report_path = output_path / "df_auc_landmark_delphi2m_report.csv"
|
||||||
output_path / "df_auc_landmark_unpooled.csv", index=False)
|
df_report.to_csv(report_path, index=False)
|
||||||
df_merged.to_csv(output_path / "df_auc_landmark.csv", index=False)
|
print(f"Saved Delphi2M-style landmark AUC report: {report_path}")
|
||||||
|
|
||||||
return df_unpooled, df_merged
|
return df_unpooled, df_report
|
||||||
|
|
||||||
|
|
||||||
def main() -> None:
|
def main() -> None:
|
||||||
@@ -1319,7 +1250,12 @@ def main() -> None:
|
|||||||
parser.add_argument("--landmark_start", type=float, default=None)
|
parser.add_argument("--landmark_start", type=float, default=None)
|
||||||
parser.add_argument("--landmark_stop", type=float, default=None)
|
parser.add_argument("--landmark_stop", type=float, default=None)
|
||||||
parser.add_argument("--landmark_step", type=float, default=None)
|
parser.add_argument("--landmark_step", type=float, default=None)
|
||||||
parser.add_argument("--horizons", type=str, default=None)
|
parser.add_argument(
|
||||||
|
"--horizons",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
help="Comma-separated horizons in years; defaults to 0.1,1,5,10, where 0.1 is Delphi2M no gap.",
|
||||||
|
)
|
||||||
|
|
||||||
parser.add_argument("--min_cases", type=int, default=None)
|
parser.add_argument("--min_cases", type=int, default=None)
|
||||||
parser.add_argument("--min_history_events", type=int, default=None)
|
parser.add_argument("--min_history_events", type=int, default=None)
|
||||||
@@ -1439,8 +1375,9 @@ def main() -> None:
|
|||||||
"Landmark ages are empty. Check landmark_start/landmark_stop/landmark_step.")
|
"Landmark ages are empty. Check landmark_start/landmark_stop/landmark_step.")
|
||||||
|
|
||||||
horizons = np.asarray(
|
horizons = np.asarray(
|
||||||
parse_float_list(cfg_get(args, cfg, "horizons", "1,5,10")) or [
|
parse_float_list(
|
||||||
1.0, 5.0, 10.0],
|
cfg_get(args, cfg, "horizons", "0.1,1,5,10")
|
||||||
|
) or list(DEFAULT_DELPHI2M_PERIODS_YEARS),
|
||||||
dtype=np.float32,
|
dtype=np.float32,
|
||||||
)
|
)
|
||||||
if horizons.size == 0:
|
if horizons.size == 0:
|
||||||
@@ -1463,12 +1400,17 @@ def main() -> None:
|
|||||||
|
|
||||||
cfg_model = dict(cfg)
|
cfg_model = dict(cfg)
|
||||||
cfg_model["dist_mode"] = dist_mode
|
cfg_model["dist_mode"] = dist_mode
|
||||||
|
model_architecture = resolve_model_architecture(cfg_model, state_dict)
|
||||||
|
cfg_model["model_architecture"] = model_architecture
|
||||||
|
print(f"Resolved model architecture: {model_architecture}")
|
||||||
|
|
||||||
device = resolve_eval_device(args.device)
|
device = resolve_eval_device(args.device)
|
||||||
if device.type == "cuda":
|
if device.type == "cuda":
|
||||||
torch.backends.cudnn.benchmark = True
|
torch.backends.cudnn.benchmark = True
|
||||||
|
|
||||||
model = build_model_from_dataset(args, cfg_model, dataset).to(device)
|
model = build_model_from_dataset(
|
||||||
|
args, cfg_model, dataset, state_dict=state_dict
|
||||||
|
).to(device)
|
||||||
|
|
||||||
if (
|
if (
|
||||||
model_target_mode == "next_token"
|
model_target_mode == "next_token"
|
||||||
@@ -1531,8 +1473,6 @@ def main() -> None:
|
|||||||
if model_target_mode == "next_token"
|
if model_target_mode == "next_token"
|
||||||
else "direct_t_query"
|
else "direct_t_query"
|
||||||
)
|
)
|
||||||
score_mode_out = f"{landmark_query_mode}_{score_mode}"
|
|
||||||
|
|
||||||
num_workers_auc = int(
|
num_workers_auc = int(
|
||||||
cfg_get(args, cfg, "num_workers_auc", max(1, (os.cpu_count() or 2) - 1)))
|
cfg_get(args, cfg, "num_workers_auc", max(1, (os.cpu_count() or 2) - 1)))
|
||||||
auc_task_chunk_size = int(cfg_get(args, cfg, "auc_task_chunk_size", 0))
|
auc_task_chunk_size = int(cfg_get(args, cfg, "auc_task_chunk_size", 0))
|
||||||
@@ -1564,27 +1504,11 @@ def main() -> None:
|
|||||||
print(f"AUC workers: {num_workers_auc}")
|
print(f"AUC workers: {num_workers_auc}")
|
||||||
print(f"Output path: {output_path}")
|
print(f"Output path: {output_path}")
|
||||||
|
|
||||||
meta_info = {
|
|
||||||
"score_mode": score_mode_out,
|
|
||||||
"eval_split": eval_split,
|
|
||||||
"model_ckpt_path": str(model_ckpt_path),
|
|
||||||
"config_path": str(config_path),
|
|
||||||
"target_mode": str(target_mode),
|
|
||||||
"model_target_mode": str(model_target_mode),
|
|
||||||
"dist_mode": str(dist_mode),
|
|
||||||
"time_mode": str(time_mode),
|
|
||||||
"attn_mask_mode": str(attn_mask_mode),
|
|
||||||
"readout_name": str(readout_name),
|
|
||||||
"landmark_query_mode": landmark_query_mode,
|
|
||||||
"landmark_token_mode": "no_event" if model_target_mode == "next_token" else "none",
|
|
||||||
}
|
|
||||||
|
|
||||||
evaluate_landmark_auc(
|
evaluate_landmark_auc(
|
||||||
model=model,
|
model=model,
|
||||||
loader=loader,
|
loader=loader,
|
||||||
landmark_dataset=landmark_dataset,
|
landmark_dataset=landmark_dataset,
|
||||||
output_path=output_path,
|
output_path=output_path,
|
||||||
labels_meta=labels_meta,
|
|
||||||
disease_ids=disease_ids,
|
disease_ids=disease_ids,
|
||||||
disease_chunk_size=disease_chunk_size,
|
disease_chunk_size=disease_chunk_size,
|
||||||
score_mode=score_mode,
|
score_mode=score_mode,
|
||||||
@@ -1601,7 +1525,6 @@ def main() -> None:
|
|||||||
use_amp=use_amp,
|
use_amp=use_amp,
|
||||||
hidden_cache_dtype=hidden_cache_dtype,
|
hidden_cache_dtype=hidden_cache_dtype,
|
||||||
logit_batch_size=logit_batch_size,
|
logit_batch_size=logit_batch_size,
|
||||||
meta_info=meta_info,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -649,7 +649,9 @@ def main() -> None:
|
|||||||
cfg_model = dict(cfg)
|
cfg_model = dict(cfg)
|
||||||
cfg_model["dist_mode"] = dist_mode
|
cfg_model["dist_mode"] = dist_mode
|
||||||
device = resolve_eval_device(args.device)
|
device = resolve_eval_device(args.device)
|
||||||
model = build_model_from_dataset(args, cfg_model, dataset).to(device)
|
model = build_model_from_dataset(
|
||||||
|
args, cfg_model, dataset, state_dict=state_dict
|
||||||
|
).to(device)
|
||||||
load_model_state(model, state_dict)
|
load_model_state(model, state_dict)
|
||||||
model.eval()
|
model.eval()
|
||||||
|
|
||||||
|
|||||||
@@ -758,7 +758,9 @@ def main() -> None:
|
|||||||
cfg_model = dict(cfg)
|
cfg_model = dict(cfg)
|
||||||
cfg_model["dist_mode"] = dist_mode
|
cfg_model["dist_mode"] = dist_mode
|
||||||
device = resolve_eval_device(args.device)
|
device = resolve_eval_device(args.device)
|
||||||
model = build_model_from_dataset(args, cfg_model, dataset).to(device)
|
model = build_model_from_dataset(
|
||||||
|
args, cfg_model, dataset, state_dict=state_dict
|
||||||
|
).to(device)
|
||||||
load_model_state(model, state_dict)
|
load_model_state(model, state_dict)
|
||||||
model.eval()
|
model.eval()
|
||||||
|
|
||||||
|
|||||||
@@ -553,7 +553,9 @@ def main() -> None:
|
|||||||
device = resolve_eval_device(args.device)
|
device = resolve_eval_device(args.device)
|
||||||
selected_token_mask = np.zeros(int(dataset.vocab_size), dtype=bool)
|
selected_token_mask = np.zeros(int(dataset.vocab_size), dtype=bool)
|
||||||
selected_token_mask[np.asarray(scanned_disease_tokens, dtype=np.int64)] = True
|
selected_token_mask[np.asarray(scanned_disease_tokens, dtype=np.int64)] = True
|
||||||
model = build_model_from_dataset(args, cfg_model, dataset).to(device)
|
model = build_model_from_dataset(
|
||||||
|
args, cfg_model, dataset, state_dict=state_dict
|
||||||
|
).to(device)
|
||||||
load_model_state(model, state_dict)
|
load_model_state(model, state_dict)
|
||||||
model.eval()
|
model.eval()
|
||||||
|
|
||||||
|
|||||||
@@ -180,7 +180,9 @@ def main() -> None:
|
|||||||
cfg_model = dict(cfg)
|
cfg_model = dict(cfg)
|
||||||
cfg_model["dist_mode"] = dist_mode
|
cfg_model["dist_mode"] = dist_mode
|
||||||
device = resolve_eval_device(args.device)
|
device = resolve_eval_device(args.device)
|
||||||
model = build_model_from_dataset(args, cfg_model, dataset).to(device)
|
model = build_model_from_dataset(
|
||||||
|
args, cfg_model, dataset, state_dict=state_dict
|
||||||
|
).to(device)
|
||||||
load_model_state(model, state_dict)
|
load_model_state(model, state_dict)
|
||||||
model.eval()
|
model.eval()
|
||||||
|
|
||||||
@@ -205,7 +207,7 @@ def main() -> None:
|
|||||||
|
|
||||||
n_rows = len(landmark_dataset)
|
n_rows = len(landmark_dataset)
|
||||||
vocab_size = int(dataset.vocab_size)
|
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)
|
logits_dtype = numpy_float_dtype(args.logits_dtype)
|
||||||
hidden_dtype = numpy_float_dtype(args.hidden_dtype)
|
hidden_dtype = numpy_float_dtype(args.hidden_dtype)
|
||||||
|
|
||||||
|
|||||||
@@ -381,7 +381,9 @@ def main() -> None:
|
|||||||
cfg_model = dict(cfg)
|
cfg_model = dict(cfg)
|
||||||
cfg_model["dist_mode"] = dist_mode
|
cfg_model["dist_mode"] = dist_mode
|
||||||
device = resolve_eval_device(args.device)
|
device = resolve_eval_device(args.device)
|
||||||
model = build_model_from_dataset(args, cfg_model, dataset).to(device)
|
model = build_model_from_dataset(
|
||||||
|
args, cfg_model, dataset, state_dict=state_dict
|
||||||
|
).to(device)
|
||||||
load_model_state(model, state_dict)
|
load_model_state(model, state_dict)
|
||||||
model.eval()
|
model.eval()
|
||||||
|
|
||||||
|
|||||||
131
model_architectures.py
Normal file
131
model_architectures.py
Normal file
@@ -0,0 +1,131 @@
|
|||||||
|
"""Model-architecture identifiers and checkpoint validation helpers."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import re
|
||||||
|
from collections.abc import Mapping
|
||||||
|
|
||||||
|
|
||||||
|
TRANSFORMER_FFN_ARCHITECTURE = "transformer_ffn_v1"
|
||||||
|
TRAJ_MIXER_ARCHITECTURE = "traj_mixer_v5"
|
||||||
|
DEFAULT_MODEL_ARCHITECTURE = TRANSFORMER_FFN_ARCHITECTURE
|
||||||
|
SUPPORTED_MODEL_ARCHITECTURES = (
|
||||||
|
TRANSFORMER_FFN_ARCHITECTURE,
|
||||||
|
TRAJ_MIXER_ARCHITECTURE,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
_FFN_STATE_KEY = re.compile(
|
||||||
|
r"(?:^|\.)blocks\.\d+\.mlp\.w[123]\.(?:weight|bias)$"
|
||||||
|
)
|
||||||
|
_TRAJ_MIXER_STATE_KEY = re.compile(
|
||||||
|
r"(?:^|\.)blocks\.\d+\.mlp\.(?:"
|
||||||
|
r"norm\.(?:weight|bias)|"
|
||||||
|
r"intra_gate_proj|"
|
||||||
|
r"intra_value_proj|"
|
||||||
|
r"intra_output_proj|"
|
||||||
|
r"intra_gate_logits|"
|
||||||
|
r"gate_proj|"
|
||||||
|
r"value_proj|"
|
||||||
|
r"output_proj"
|
||||||
|
r")$"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_model_architecture(model_architecture: object) -> str:
|
||||||
|
if not isinstance(model_architecture, str):
|
||||||
|
raise ValueError(
|
||||||
|
"model_architecture must be one of "
|
||||||
|
f"{SUPPORTED_MODEL_ARCHITECTURES}, got {model_architecture!r}"
|
||||||
|
)
|
||||||
|
if model_architecture not in SUPPORTED_MODEL_ARCHITECTURES:
|
||||||
|
raise ValueError(
|
||||||
|
f"Unsupported model_architecture={model_architecture!r}; "
|
||||||
|
f"expected one of {SUPPORTED_MODEL_ARCHITECTURES}."
|
||||||
|
)
|
||||||
|
return model_architecture
|
||||||
|
|
||||||
|
|
||||||
|
def detect_model_architecture_from_state_dict(
|
||||||
|
state_dict: Mapping[str, object],
|
||||||
|
) -> str:
|
||||||
|
"""Infer the architecture from block parameter names.
|
||||||
|
|
||||||
|
Detection deliberately accepts any ``blocks.<index>`` prefix rather than
|
||||||
|
assuming that block zero is present.
|
||||||
|
"""
|
||||||
|
|
||||||
|
if not isinstance(state_dict, Mapping):
|
||||||
|
raise TypeError(
|
||||||
|
"state_dict must be a mapping, got "
|
||||||
|
f"{type(state_dict).__name__}"
|
||||||
|
)
|
||||||
|
|
||||||
|
has_ffn = False
|
||||||
|
has_traj_mixer = False
|
||||||
|
for raw_key in state_dict:
|
||||||
|
key = str(raw_key)
|
||||||
|
has_ffn = has_ffn or _FFN_STATE_KEY.search(key) is not None
|
||||||
|
has_traj_mixer = (
|
||||||
|
has_traj_mixer
|
||||||
|
or _TRAJ_MIXER_STATE_KEY.search(key) is not None
|
||||||
|
)
|
||||||
|
if has_ffn and has_traj_mixer:
|
||||||
|
raise ValueError(
|
||||||
|
"Checkpoint contains both Transformer FFN and TrajMixer "
|
||||||
|
"block parameters; its model architecture is ambiguous."
|
||||||
|
)
|
||||||
|
|
||||||
|
if has_ffn:
|
||||||
|
return TRANSFORMER_FFN_ARCHITECTURE
|
||||||
|
if has_traj_mixer:
|
||||||
|
return TRAJ_MIXER_ARCHITECTURE
|
||||||
|
raise ValueError(
|
||||||
|
"Could not detect model architecture from checkpoint parameters. "
|
||||||
|
"Expected a blocks.<index>.mlp FFN or TrajMixer parameter."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_model_architecture(
|
||||||
|
config_or_marker: Mapping[str, object] | str | None = None,
|
||||||
|
state_dict: Mapping[str, object] | None = None,
|
||||||
|
) -> str:
|
||||||
|
"""Resolve and cross-check a configured and checkpoint architecture.
|
||||||
|
|
||||||
|
Every saved run must provide an explicit architecture marker. Checkpoint
|
||||||
|
parameter names are used only to verify that the marker describes the
|
||||||
|
weights being loaded.
|
||||||
|
"""
|
||||||
|
|
||||||
|
if isinstance(config_or_marker, Mapping):
|
||||||
|
configured = config_or_marker.get("model_architecture")
|
||||||
|
elif isinstance(config_or_marker, str) or config_or_marker is None:
|
||||||
|
configured = config_or_marker
|
||||||
|
else:
|
||||||
|
raise TypeError(
|
||||||
|
"config_or_marker must be a config mapping, string, or None, got "
|
||||||
|
f"{type(config_or_marker).__name__}"
|
||||||
|
)
|
||||||
|
|
||||||
|
resolved_config = (
|
||||||
|
_validate_model_architecture(configured)
|
||||||
|
if configured is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
detected = (
|
||||||
|
detect_model_architecture_from_state_dict(state_dict)
|
||||||
|
if state_dict is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
|
||||||
|
if resolved_config is None:
|
||||||
|
raise ValueError(
|
||||||
|
"model_architecture is required; expected one of "
|
||||||
|
f"{SUPPORTED_MODEL_ARCHITECTURES}."
|
||||||
|
)
|
||||||
|
if detected is not None and resolved_config != detected:
|
||||||
|
raise ValueError(
|
||||||
|
"Configured model architecture conflicts with checkpoint: "
|
||||||
|
f"config={resolved_config!r}, checkpoint={detected!r}."
|
||||||
|
)
|
||||||
|
return resolved_config
|
||||||
480
models.py
480
models.py
@@ -1,4 +1,3 @@
|
|||||||
from collections.abc import Mapping
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -8,199 +7,14 @@ import torch.nn.functional as F
|
|||||||
from backbones import (
|
from backbones import (
|
||||||
AgeSinusoidalEncoding,
|
AgeSinusoidalEncoding,
|
||||||
GaussianRBFTimeBasis,
|
GaussianRBFTimeBasis,
|
||||||
SharedEventTrajectoryCore,
|
|
||||||
TimeRoPE,
|
TimeRoPE,
|
||||||
TokenAutoDiscretization,
|
TokenAutoDiscretization,
|
||||||
|
build_backbone_block,
|
||||||
)
|
)
|
||||||
|
from model_architectures import resolve_model_architecture
|
||||||
from targets import PAD_IDX
|
from targets import PAD_IDX
|
||||||
|
|
||||||
|
|
||||||
EVENT_TRAJECTORY_ARCHITECTURE = "event_trajectory_shared_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=120, n_trajectory=6),
|
|
||||||
"tiny": 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:
|
|
||||||
actual = config.get("model_architecture")
|
|
||||||
if actual != EVENT_TRAJECTORY_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)
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
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:
|
|
||||||
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",
|
|
||||||
}
|
|
||||||
missing = sorted(required_keys.difference(state_dict))
|
|
||||||
if missing:
|
|
||||||
raise ValueError(
|
|
||||||
"Checkpoint is not a shared event-trajectory 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
|
@dataclass
|
||||||
class DeepHealthOutput:
|
class DeepHealthOutput:
|
||||||
hidden: torch.Tensor
|
hidden: torch.Tensor
|
||||||
@@ -332,8 +146,9 @@ class DeepHealth(nn.Module):
|
|||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
vocab_size: int,
|
vocab_size: int,
|
||||||
model_size: str,
|
n_embd: int,
|
||||||
n_reasoning_rounds: int,
|
n_head: int,
|
||||||
|
n_layer: int,
|
||||||
n_types: int,
|
n_types: int,
|
||||||
n_cont_types: int,
|
n_cont_types: int,
|
||||||
n_categories: int,
|
n_categories: int,
|
||||||
@@ -345,6 +160,7 @@ class DeepHealth(nn.Module):
|
|||||||
dist_mode: str = "exponential", # "exponential", "weibull" or "mixed"
|
dist_mode: str = "exponential", # "exponential", "weibull" or "mixed"
|
||||||
extra_pool_reduce: str = "mean",
|
extra_pool_reduce: str = "mean",
|
||||||
dropout: float = 0.0,
|
dropout: float = 0.0,
|
||||||
|
model_architecture: str | None = None,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
if target_mode not in ["next_token", "all_future"]:
|
if target_mode not in ["next_token", "all_future"]:
|
||||||
@@ -358,21 +174,14 @@ class DeepHealth(nn.Module):
|
|||||||
"dist_mode must be either 'exponential', 'weibull' or 'mixed'")
|
"dist_mode must be either 'exponential', 'weibull' or 'mixed'")
|
||||||
if extra_pool_reduce not in {"mean", "sum"}:
|
if extra_pool_reduce not in {"mean", "sum"}:
|
||||||
raise ValueError("extra_pool_reduce must be either 'mean' or 'sum'")
|
raise ValueError("extra_pool_reduce must be either 'mean' or 'sum'")
|
||||||
if n_reasoning_rounds <= 0:
|
if n_layer < 1:
|
||||||
raise ValueError(
|
raise ValueError(f"n_layer must be >= 1, got {n_layer}")
|
||||||
"n_reasoning_rounds must be positive, got "
|
model_architecture = resolve_model_architecture(model_architecture)
|
||||||
f"{n_reasoning_rounds}"
|
self.token_embedding = nn.Embedding(vocab_size, n_embd, padding_idx=0)
|
||||||
)
|
|
||||||
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.gender_embedding = nn.Embedding(
|
self.gender_embedding = nn.Embedding(
|
||||||
2, d_model) # Assuming binary gender
|
2, n_embd) # Assuming binary gender
|
||||||
self.tokenizer = OtherInfoTokenizer(
|
self.tokenizer = OtherInfoTokenizer(
|
||||||
n_embd=d_model,
|
n_embd=n_embd,
|
||||||
n_types=n_types,
|
n_types=n_types,
|
||||||
n_cont_types=n_cont_types,
|
n_cont_types=n_cont_types,
|
||||||
n_categories=n_categories,
|
n_categories=n_categories,
|
||||||
@@ -384,101 +193,74 @@ class DeepHealth(nn.Module):
|
|||||||
self.time_mode = time_mode
|
self.time_mode = time_mode
|
||||||
self.dist_mode = dist_mode
|
self.dist_mode = dist_mode
|
||||||
self.extra_pool_reduce = extra_pool_reduce
|
self.extra_pool_reduce = extra_pool_reduce
|
||||||
self.model_size = normalized_model_size
|
self.model_architecture = model_architecture
|
||||||
self.d_model = d_model
|
self.n_layer = n_layer
|
||||||
self.n_trajectory = n_trajectory
|
self.n_embd = n_embd
|
||||||
self.trajectory_dim = d_model // n_trajectory
|
|
||||||
self.traj_hidden = 4 * n_trajectory
|
|
||||||
self.n_reasoning_rounds = n_reasoning_rounds
|
|
||||||
self.vocab_size = vocab_size
|
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.normal_(self.token_embedding.weight, mean=0.0, std=0.02)
|
||||||
nn.init.zeros_(self.token_embedding.weight[0])
|
nn.init.zeros_(self.token_embedding.weight[0])
|
||||||
nn.init.normal_(self.gender_embedding.weight, mean=0.0, std=0.02)
|
nn.init.normal_(self.gender_embedding.weight, mean=0.0, std=0.02)
|
||||||
if dist_mode == "weibull":
|
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.zeros_(self.rho_head.weight)
|
||||||
nn.init.constant_(self.rho_head.bias, 0.5413)
|
nn.init.constant_(self.rho_head.bias, 0.5413)
|
||||||
|
|
||||||
if dist_mode == "mixed":
|
if dist_mode == "mixed":
|
||||||
self.death_idx = vocab_size - 1
|
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.zeros_(self.rho_death_head.weight)
|
||||||
nn.init.constant_(self.rho_death_head.bias, 0.5413)
|
nn.init.constant_(self.rho_death_head.bias, 0.5413)
|
||||||
|
|
||||||
# Event and query time are encoded once before shared reasoning. In
|
if time_mode == "absolute":
|
||||||
# relative mode, cross-attention additionally uses TimeRoPE and RBF.
|
self.age_encoding = AgeSinusoidalEncoding(n_embd)
|
||||||
self.age_encoding = AgeSinusoidalEncoding(d_model)
|
self.blocks = nn.ModuleList([
|
||||||
self.event_projection = nn.Linear(d_model, d_model, bias=False)
|
build_backbone_block(
|
||||||
self.event_norm = nn.LayerNorm(d_model)
|
model_architecture,
|
||||||
self.query_projection = nn.Linear(d_model, d_model, bias=False)
|
n_embd=n_embd,
|
||||||
self.trajectory_prototypes = nn.Parameter(
|
n_head=n_head,
|
||||||
torch.empty(n_trajectory, self.trajectory_dim)
|
use_time_rope=False,
|
||||||
)
|
use_rbf_bias=False,
|
||||||
self.query_token = nn.Parameter(torch.empty(d_model))
|
mlp_dropout=dropout,
|
||||||
nn.init.normal_(self.event_projection.weight, mean=0.0, std=0.02)
|
) for _ in range(n_layer)
|
||||||
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:
|
|
||||||
self.rope = None
|
self.rope = None
|
||||||
self.rbf = None
|
self.rbf = None
|
||||||
|
elif time_mode == "relative":
|
||||||
|
self.age_encoding = None
|
||||||
|
self.blocks = nn.ModuleList([
|
||||||
|
build_backbone_block(
|
||||||
|
model_architecture,
|
||||||
|
n_embd=n_embd,
|
||||||
|
n_head=n_head,
|
||||||
|
use_time_rope=True,
|
||||||
|
use_rbf_bias=True,
|
||||||
|
mlp_dropout=dropout,
|
||||||
|
) for _ in range(n_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.final_ln = nn.LayerNorm(n_embd)
|
||||||
self.risk_head = nn.Linear(d_model, vocab_size, bias=False)
|
self.risk_head = nn.Linear(n_embd, vocab_size, bias=False)
|
||||||
if target_mode == "next_token":
|
if target_mode == "next_token":
|
||||||
self.risk_head.weight = self.token_embedding.weight
|
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,
|
self,
|
||||||
event_valid_mask: torch.Tensor,
|
padding_mask: torch.Tensor,
|
||||||
event_time: torch.Tensor,
|
time_seq: torch.Tensor,
|
||||||
query_time: torch.Tensor,
|
dtype: torch.dtype,
|
||||||
query_position: torch.Tensor | None = None,
|
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
valid_key = event_valid_mask[:, None, :]
|
valid_key = padding_mask[:, None, :] # (B, 1, L)
|
||||||
key_time = event_time[:, None, :]
|
visible_by_time = time_seq[:, None, :] <= time_seq[:, :, None]
|
||||||
query_time = query_time[:, :, None]
|
valid = valid_key & visible_by_time
|
||||||
if query_position is None:
|
return torch.zeros(
|
||||||
visible_by_time = key_time <= query_time
|
valid.shape,
|
||||||
else:
|
device=valid.device,
|
||||||
key_position = torch.arange(
|
dtype=dtype,
|
||||||
event_time.size(1),
|
).masked_fill(~valid, -1e4)[:, None, :, :]
|
||||||
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)
|
|
||||||
|
|
||||||
def _pool_other_by_time(
|
def _pool_other_by_time(
|
||||||
self,
|
self,
|
||||||
@@ -582,8 +364,8 @@ class DeepHealth(nn.Module):
|
|||||||
padding_mask = padding_mask.to(device=event_seq.device, dtype=torch.bool)
|
padding_mask = padding_mask.to(device=event_seq.device, dtype=torch.bool)
|
||||||
|
|
||||||
event_len = event_seq.size(1)
|
event_len = event_seq.size(1)
|
||||||
event_features = self.token_embedding(event_seq)
|
h_disease = self.token_embedding(event_seq)
|
||||||
event_time = time_seq
|
t_disease = time_seq
|
||||||
|
|
||||||
if other_time.shape != other_type.shape:
|
if other_time.shape != other_type.shape:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -591,120 +373,64 @@ class DeepHealth(nn.Module):
|
|||||||
f"{tuple(other_time.shape)} vs {tuple(other_type.shape)}"
|
f"{tuple(other_time.shape)} vs {tuple(other_type.shape)}"
|
||||||
)
|
)
|
||||||
other_time = other_time.to(device=event_seq.device, dtype=time_seq.dtype)
|
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_type=other_type,
|
||||||
other_value=other_value,
|
other_value=other_value,
|
||||||
other_value_kind=other_value_kind,
|
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)
|
other_mask = other_mask.to(device=event_seq.device, dtype=torch.bool)
|
||||||
|
|
||||||
event_features = torch.cat([event_features, other_features], dim=1)
|
h_disease = torch.cat([h_disease, h_other], dim=1)
|
||||||
event_time = torch.cat([event_time, other_time], dim=1)
|
t_disease = torch.cat([t_disease, other_time], dim=1)
|
||||||
event_valid_mask = torch.cat([padding_mask, other_mask], dim=1)
|
padding_mask = torch.cat([padding_mask, other_mask], dim=1)
|
||||||
|
h_disease = h_disease * padding_mask.unsqueeze(-1).to(h_disease.dtype)
|
||||||
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
|
|
||||||
)
|
|
||||||
|
|
||||||
if mode == "all_future":
|
if mode == "all_future":
|
||||||
query_time = t_query[:, None]
|
batch_size = event_seq.size(0)
|
||||||
query_position = None
|
query = self.query_token.view(1, 1, -1).expand(batch_size, 1, -1)
|
||||||
query_features = (
|
h_disease = torch.cat([h_disease, query], dim=1)
|
||||||
self.query_token.view(1, 1, -1)
|
t_disease = torch.cat([t_disease, t_query[:, None]], dim=1)
|
||||||
+ sex_context
|
query_mask = torch.ones(
|
||||||
+ self.age_encoding(query_time)
|
|
||||||
)
|
|
||||||
query_valid_mask = torch.ones(
|
|
||||||
batch_size,
|
batch_size,
|
||||||
1,
|
1,
|
||||||
dtype=torch.bool,
|
dtype=torch.bool,
|
||||||
device=event_seq.device,
|
device=event_seq.device,
|
||||||
)
|
)
|
||||||
else:
|
padding_mask = torch.cat([padding_mask, query_mask], dim=1)
|
||||||
# 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
|
|
||||||
|
|
||||||
n_query = query_time.size(1)
|
sex_emb = self.gender_embedding(sex)[:, None, :]
|
||||||
query_context = self.query_projection(query_features).reshape(
|
h_disease = h_disease + sex_emb
|
||||||
batch_size,
|
h_disease = h_disease * padding_mask.unsqueeze(-1).to(h_disease.dtype)
|
||||||
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,
|
|
||||||
)
|
|
||||||
|
|
||||||
event_rope_cache = None
|
rope_cache = None
|
||||||
query_rope_cache = None
|
|
||||||
rbf_cache = None
|
rbf_cache = None
|
||||||
if self.time_mode == "relative":
|
if self.time_mode == "absolute":
|
||||||
if self.rope is None or self.rbf is None:
|
h_disease = h_disease + self.age_encoding(t_disease)
|
||||||
raise RuntimeError("Relative-time modules are not initialized")
|
h_disease = h_disease * padding_mask.unsqueeze(-1).to(h_disease.dtype)
|
||||||
event_rope_cache = self.rope.precompute_cache(event_time)
|
elif self.time_mode == "relative":
|
||||||
query_rope_cache = self.rope.precompute_cache(query_time)
|
rope_cache = self.rope.precompute_cache(t_disease)
|
||||||
rbf_cache = self.rbf.precompute_cross_cache(
|
rbf_cache = self.rbf.precompute_cache(t_disease)
|
||||||
query_time,
|
|
||||||
event_time,
|
|
||||||
)
|
|
||||||
|
|
||||||
event_key_value = self.reasoning_core.project_event_memory(
|
attn_mask = self._make_history_attn_mask(
|
||||||
event_memory,
|
padding_mask=padding_mask,
|
||||||
event_rope_cache=event_rope_cache,
|
time_seq=t_disease,
|
||||||
|
dtype=h_disease.dtype,
|
||||||
)
|
)
|
||||||
for _ in range(self.n_reasoning_rounds):
|
for block in self.blocks:
|
||||||
trajectory_state = self.reasoning_core(
|
h_disease = block(
|
||||||
trajectory_state=trajectory_state,
|
h_disease,
|
||||||
event_key_value=event_key_value,
|
rope_cache=rope_cache,
|
||||||
event_invalid_mask=event_invalid_mask,
|
|
||||||
query_rope_cache=query_rope_cache,
|
|
||||||
rbf_cache=rbf_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(
|
h_disease = self.final_ln(h_disease)
|
||||||
trajectory_state.reshape(batch_size, n_query, self.d_model)
|
h_disease = h_disease * padding_mask.unsqueeze(-1).to(h_disease.dtype)
|
||||||
)
|
|
||||||
hidden_sequence = hidden_sequence * query_valid_mask.unsqueeze(-1).to(
|
|
||||||
hidden_sequence.dtype
|
|
||||||
)
|
|
||||||
|
|
||||||
if mode == "all_future":
|
if mode == "all_future":
|
||||||
hidden = hidden_sequence[:, 0, :]
|
hidden = h_disease[:, -1, :]
|
||||||
if return_output:
|
if return_output:
|
||||||
return DeepHealthOutput(
|
return DeepHealthOutput(
|
||||||
hidden=hidden,
|
hidden=hidden,
|
||||||
@@ -719,13 +445,13 @@ class DeepHealth(nn.Module):
|
|||||||
)
|
)
|
||||||
return hidden
|
return hidden
|
||||||
if return_output:
|
if return_output:
|
||||||
h_event = hidden_sequence[:, :event_len, :]
|
h_event = h_disease[:, :event_len, :]
|
||||||
t_event = event_time[:, :event_len]
|
t_event = t_disease[:, :event_len]
|
||||||
event_mask = event_valid_mask[:, :event_len]
|
event_mask = padding_mask[:, :event_len]
|
||||||
h_extra, t_extra, extra_mask = self._pool_other_by_time(
|
h_extra, t_extra, extra_mask = self._pool_other_by_time(
|
||||||
h_other=hidden_sequence[:, event_len:, :],
|
h_other=h_disease[:, event_len:, :],
|
||||||
other_time=event_time[:, event_len:],
|
other_time=t_disease[:, event_len:],
|
||||||
other_mask=event_valid_mask[:, event_len:],
|
other_mask=padding_mask[:, event_len:],
|
||||||
)
|
)
|
||||||
return DeepHealthOutput(
|
return DeepHealthOutput(
|
||||||
hidden=torch.cat([h_event, h_extra], dim=1),
|
hidden=torch.cat([h_event, h_extra], dim=1),
|
||||||
@@ -733,7 +459,7 @@ class DeepHealth(nn.Module):
|
|||||||
padding_mask=torch.cat([event_mask, extra_mask], dim=1),
|
padding_mask=torch.cat([event_mask, extra_mask], dim=1),
|
||||||
event_len=event_len,
|
event_len=event_len,
|
||||||
)
|
)
|
||||||
return hidden_sequence[:, :event_len, :]
|
return h_disease[:, :event_len, :]
|
||||||
|
|
||||||
def forward_next_token(self, **kwargs) -> torch.Tensor:
|
def forward_next_token(self, **kwargs) -> torch.Tensor:
|
||||||
return self._forward_shared(mode="next_token", **kwargs)
|
return self._forward_shared(mode="next_token", **kwargs)
|
||||||
|
|||||||
@@ -1,10 +1,11 @@
|
|||||||
#!/usr/bin/env bash
|
#!/usr/bin/env bash
|
||||||
set -euo pipefail
|
set -euo pipefail
|
||||||
|
|
||||||
# Run all non-wrapper evaluation scripts for every completed experiment under
|
# Run all non-wrapper evaluation scripts for every completed current-format
|
||||||
# runs/. The script is written for Linux servers with bash 4.2.
|
# experiment under runs/. The script is written for Linux servers with bash 4.2.
|
||||||
|
|
||||||
cd "$(dirname "${BASH_SOURCE[0]}")"
|
cd "$(dirname "${BASH_SOURCE[0]}")"
|
||||||
|
shopt -s globstar nullglob
|
||||||
|
|
||||||
PYTHON_BIN="${PYTHON_BIN:-python}"
|
PYTHON_BIN="${PYTHON_BIN:-python}"
|
||||||
DEVICE="${DEVICE:-cuda}"
|
DEVICE="${DEVICE:-cuda}"
|
||||||
@@ -81,6 +82,20 @@ run_dir_result_if_missing() {
|
|||||||
run_command "$@"
|
run_command "$@"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
run_file_result_if_missing() {
|
||||||
|
local label="$1"
|
||||||
|
local result_dir="$2"
|
||||||
|
local required="$3"
|
||||||
|
shift 3
|
||||||
|
|
||||||
|
if [[ -s "${result_dir}/${required}" ]]; then
|
||||||
|
echo " skip ${label}: found ${result_dir}/${required}"
|
||||||
|
return 0
|
||||||
|
fi
|
||||||
|
|
||||||
|
run_command "$@"
|
||||||
|
}
|
||||||
|
|
||||||
run_has_extra_info() {
|
run_has_extra_info() {
|
||||||
"${PYTHON_BIN}" - "$1" <<'PY'
|
"${PYTHON_BIN}" - "$1" <<'PY'
|
||||||
import json
|
import json
|
||||||
@@ -115,8 +130,30 @@ raise SystemExit(0 if mode == "all_future" else 1)
|
|||||||
PY
|
PY
|
||||||
}
|
}
|
||||||
|
|
||||||
for run_path in runs/*; do
|
run_has_current_model_config() {
|
||||||
[[ -d "${run_path}" ]] || continue
|
"${PYTHON_BIN}" - "$1" <<'PY'
|
||||||
|
import json
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
cfg_path = Path(sys.argv[1]) / "train_config.json"
|
||||||
|
try:
|
||||||
|
cfg = json.loads(cfg_path.read_text(encoding="utf-8"))
|
||||||
|
n_layer = int(cfg.get("n_layer", 0))
|
||||||
|
except Exception:
|
||||||
|
raise SystemExit(1)
|
||||||
|
|
||||||
|
supported = {"transformer_ffn_v1", "traj_mixer_v5"}
|
||||||
|
raise SystemExit(
|
||||||
|
0
|
||||||
|
if cfg.get("model_architecture") in supported and n_layer >= 1
|
||||||
|
else 1
|
||||||
|
)
|
||||||
|
PY
|
||||||
|
}
|
||||||
|
|
||||||
|
for config_path in runs/**/train_config.json; do
|
||||||
|
run_path="${config_path%/train_config.json}"
|
||||||
|
|
||||||
echo "==> ${run_path}"
|
echo "==> ${run_path}"
|
||||||
if [[ ! -f "${run_path}/train_config.json" ]]; then
|
if [[ ! -f "${run_path}/train_config.json" ]]; then
|
||||||
@@ -127,6 +164,10 @@ for run_path in runs/*; do
|
|||||||
echo " skip run: missing best_model.pt"
|
echo " skip run: missing best_model.pt"
|
||||||
continue
|
continue
|
||||||
fi
|
fi
|
||||||
|
if ! run_has_current_model_config "${run_path}"; then
|
||||||
|
echo " skip run: config lacks current model_architecture/n_layer fields"
|
||||||
|
continue
|
||||||
|
fi
|
||||||
|
|
||||||
common=()
|
common=()
|
||||||
while IFS= read -r arg; do common+=("${arg}"); done < <(common_args_with_device "${run_path}")
|
while IFS= read -r arg; do common+=("${arg}"); done < <(common_args_with_device "${run_path}")
|
||||||
@@ -137,18 +178,16 @@ for run_path in runs/*; do
|
|||||||
cpu_reduce_extra=()
|
cpu_reduce_extra=()
|
||||||
while IFS= read -r arg; do cpu_reduce_extra+=("${arg}"); done < <(cpu_reduce_args)
|
while IFS= read -r arg; do cpu_reduce_extra+=("${arg}"); done < <(cpu_reduce_args)
|
||||||
|
|
||||||
run_dir_result_if_missing \
|
run_file_result_if_missing \
|
||||||
"evaluate_auc.py" \
|
"evaluate_auc.py" \
|
||||||
"${run_path}" \
|
"${run_path}" \
|
||||||
"df_both.csv" \
|
"df_auc_delphi2m_report.csv" \
|
||||||
"df_auc_unpooled.csv" \
|
|
||||||
"${PYTHON_BIN}" evaluate_auc.py "${common[@]}" "${auc_extra[@]}"
|
"${PYTHON_BIN}" evaluate_auc.py "${common[@]}" "${auc_extra[@]}"
|
||||||
|
|
||||||
run_dir_result_if_missing \
|
run_file_result_if_missing \
|
||||||
"evaluate_auc_v2.py" \
|
"evaluate_auc_v2.py" \
|
||||||
"${run_path}" \
|
"${run_path}" \
|
||||||
"df_auc_landmark.csv" \
|
"df_auc_landmark_delphi2m_report.csv" \
|
||||||
"df_auc_landmark_unpooled.csv" \
|
|
||||||
"${PYTHON_BIN}" evaluate_auc_v2.py "${common[@]}" "${auc_extra[@]}"
|
"${PYTHON_BIN}" evaluate_auc_v2.py "${common[@]}" "${auc_extra[@]}"
|
||||||
|
|
||||||
if ! run_is_all_future "${run_path}"; then
|
if ! run_is_all_future "${run_path}"; then
|
||||||
|
|||||||
@@ -10,7 +10,8 @@ set -euo pipefail
|
|||||||
# all_future + relative time + mixed death/risk head
|
# all_future + relative time + mixed death/risk head
|
||||||
#
|
#
|
||||||
# This script only launches those missing training jobs. It intentionally does
|
# This script only launches those missing training jobs. It intentionally does
|
||||||
# not call evaluate_*.py and does not add extra random seeds.
|
# not call evaluate_*.py and does not add extra random seeds. Set
|
||||||
|
# MODEL_ARCHITECTURE=traj_mixer_v5 to run the TrajMixer variant.
|
||||||
|
|
||||||
cd "$(dirname "${BASH_SOURCE[0]}")"
|
cd "$(dirname "${BASH_SOURCE[0]}")"
|
||||||
|
|
||||||
@@ -18,6 +19,8 @@ PYTHON_BIN="${PYTHON_BIN:-python}"
|
|||||||
DEVICE="${DEVICE:-cuda}"
|
DEVICE="${DEVICE:-cuda}"
|
||||||
NUM_WORKERS="${NUM_WORKERS:-4}"
|
NUM_WORKERS="${NUM_WORKERS:-4}"
|
||||||
PROGRESS_INTERVAL="${PROGRESS_INTERVAL:-20}"
|
PROGRESS_INTERVAL="${PROGRESS_INTERVAL:-20}"
|
||||||
|
MODEL_ARCHITECTURE="${MODEL_ARCHITECTURE:-transformer_ffn_v1}"
|
||||||
|
N_LAYER="${N_LAYER:-12}"
|
||||||
|
|
||||||
TIME_MODE="relative"
|
TIME_MODE="relative"
|
||||||
DIST_MODE="mixed"
|
DIST_MODE="mixed"
|
||||||
@@ -36,8 +39,8 @@ COMMON_ARGS=(
|
|||||||
--min_future_events 1
|
--min_future_events 1
|
||||||
--n_embd 120
|
--n_embd 120
|
||||||
--n_head 10
|
--n_head 10
|
||||||
--n_hist_layer 12
|
--n_layer "${N_LAYER}"
|
||||||
--n_tab_layer 4
|
--model_architecture "${MODEL_ARCHITECTURE}"
|
||||||
--n_bins 16
|
--n_bins 16
|
||||||
--extra_pool_reduce mean
|
--extra_pool_reduce mean
|
||||||
--dropout 0.0
|
--dropout 0.0
|
||||||
@@ -57,15 +60,23 @@ COMMON_ARGS=(
|
|||||||
|
|
||||||
already_trained() {
|
already_trained() {
|
||||||
local extra_file="$1"
|
local extra_file="$1"
|
||||||
"${PYTHON_BIN}" - "$TIME_MODE" "$DIST_MODE" "$extra_file" "$SEED" "$VALIDATION_QUERY_SEED" <<'PY'
|
"${PYTHON_BIN}" - "$TIME_MODE" "$DIST_MODE" "$extra_file" "$SEED" "$VALIDATION_QUERY_SEED" "$MODEL_ARCHITECTURE" "$N_LAYER" <<'PY'
|
||||||
import json
|
import json
|
||||||
import sys
|
import sys
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
time_mode, dist_mode, extra_file, seed, validation_query_seed = sys.argv[1:6]
|
(
|
||||||
|
time_mode,
|
||||||
|
dist_mode,
|
||||||
|
extra_file,
|
||||||
|
seed,
|
||||||
|
validation_query_seed,
|
||||||
|
model_architecture,
|
||||||
|
n_layer,
|
||||||
|
) = sys.argv[1:8]
|
||||||
extra_name = Path(extra_file).name
|
extra_name = Path(extra_file).name
|
||||||
|
|
||||||
for config_path in Path("runs").glob("*/train_config.json"):
|
for config_path in Path("runs").rglob("train_config.json"):
|
||||||
try:
|
try:
|
||||||
cfg = json.loads(config_path.read_text(encoding="utf-8"))
|
cfg = json.loads(config_path.read_text(encoding="utf-8"))
|
||||||
except Exception:
|
except Exception:
|
||||||
@@ -78,6 +89,8 @@ for config_path in Path("runs").glob("*/train_config.json"):
|
|||||||
|
|
||||||
if (
|
if (
|
||||||
cfg.get("model_target_mode") == "all_future"
|
cfg.get("model_target_mode") == "all_future"
|
||||||
|
and cfg.get("model_architecture") == model_architecture
|
||||||
|
and int(cfg.get("n_layer", -1)) == int(n_layer)
|
||||||
and cfg.get("time_mode") == time_mode
|
and cfg.get("time_mode") == time_mode
|
||||||
and cfg.get("dist_mode") == dist_mode
|
and cfg.get("dist_mode") == dist_mode
|
||||||
and Path(str(cfg.get("extra_info_types_file", ""))).name == extra_name
|
and Path(str(cfg.get("extra_info_types_file", ""))).name == extra_name
|
||||||
@@ -100,7 +113,7 @@ train_if_missing() {
|
|||||||
return 2
|
return 2
|
||||||
fi
|
fi
|
||||||
|
|
||||||
echo "==> Checking ${label}: ${TIME_MODE} ${DIST_MODE} all_future with ${extra_file}"
|
echo "==> Checking ${label}: ${MODEL_ARCHITECTURE} n_layer=${N_LAYER} ${TIME_MODE} ${DIST_MODE} all_future with ${extra_file}"
|
||||||
if existing_run="$(already_trained "$extra_file")"; then
|
if existing_run="$(already_trained "$extra_file")"; then
|
||||||
echo " skip: already trained at ${existing_run}"
|
echo " skip: already trained at ${existing_run}"
|
||||||
return 0
|
return 0
|
||||||
|
|||||||
@@ -1,330 +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": (120, 6, 20, 24),
|
|
||||||
"tiny": (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, 120))
|
|
||||||
next_output = next_model(**inputs, return_output=True)
|
|
||||||
self.assertEqual(tuple(next_output.hidden.shape), (2, 4, 120))
|
|
||||||
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, 120))
|
|
||||||
|
|
||||||
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": 120,
|
|
||||||
"n_trajectory": 6,
|
|
||||||
"trajectory_dim": 20,
|
|
||||||
"traj_hidden": 24,
|
|
||||||
"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, "only accepts models trained"):
|
|
||||||
validate_event_trajectory_config(
|
|
||||||
{"model_architecture": "event_trajectory_shared_v1"}
|
|
||||||
)
|
|
||||||
with self.assertRaisesRegex(ValueError, "trajectory_dim"):
|
|
||||||
validate_event_trajectory_config(
|
|
||||||
{
|
|
||||||
"model_architecture": EVENT_TRAJECTORY_ARCHITECTURE,
|
|
||||||
"model_size": "nano",
|
|
||||||
"d_model": 120,
|
|
||||||
"n_trajectory": 6,
|
|
||||||
"trajectory_dim": 10,
|
|
||||||
"traj_hidden": 24,
|
|
||||||
"n_reasoning_rounds": 3,
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
model = build_test_model()
|
|
||||||
state_dict = model.state_dict()
|
|
||||||
validate_event_trajectory_state_dict(
|
|
||||||
state_dict,
|
|
||||||
expected_d_model=120,
|
|
||||||
expected_n_trajectory=6,
|
|
||||||
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()
|
|
||||||
247
test_model_architectures.py
Normal file
247
test_model_architectures.py
Normal file
@@ -0,0 +1,247 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from backbones import (
|
||||||
|
SwiGLU,
|
||||||
|
TrajMixer,
|
||||||
|
TrajMixerBlock,
|
||||||
|
TransformerFFNBlock,
|
||||||
|
build_backbone_block,
|
||||||
|
)
|
||||||
|
from model_architectures import (
|
||||||
|
TRAJ_MIXER_ARCHITECTURE,
|
||||||
|
TRANSFORMER_FFN_ARCHITECTURE,
|
||||||
|
detect_model_architecture_from_state_dict,
|
||||||
|
resolve_model_architecture,
|
||||||
|
)
|
||||||
|
from models import DeepHealth
|
||||||
|
|
||||||
|
|
||||||
|
def _build_block(model_architecture: str):
|
||||||
|
return build_backbone_block(
|
||||||
|
model_architecture,
|
||||||
|
n_embd=12,
|
||||||
|
n_head=3,
|
||||||
|
use_time_rope=False,
|
||||||
|
use_rbf_bias=False,
|
||||||
|
mlp_dropout=0.0,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _as_model_state_dict(block: torch.nn.Module) -> dict[str, torch.Tensor]:
|
||||||
|
return {
|
||||||
|
f"blocks.0.{name}": value.detach().clone()
|
||||||
|
for name, value in block.state_dict().items()
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _build_model(
|
||||||
|
model_architecture: str | None,
|
||||||
|
*,
|
||||||
|
n_layer: int = 1,
|
||||||
|
) -> DeepHealth:
|
||||||
|
return DeepHealth(
|
||||||
|
vocab_size=8,
|
||||||
|
n_embd=12,
|
||||||
|
n_head=3,
|
||||||
|
n_layer=n_layer,
|
||||||
|
n_types=2,
|
||||||
|
n_cont_types=0,
|
||||||
|
n_categories=2,
|
||||||
|
cont_type_ids=[],
|
||||||
|
time_mode="absolute",
|
||||||
|
model_architecture=model_architecture,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ModelArchitectureFactoryTest(unittest.TestCase):
|
||||||
|
def test_factory_builds_both_architectures_with_expected_topology(self) -> None:
|
||||||
|
ffn_block = _build_block(TRANSFORMER_FFN_ARCHITECTURE)
|
||||||
|
self.assertIsInstance(ffn_block, TransformerFFNBlock)
|
||||||
|
self.assertIsInstance(ffn_block.mlp, SwiGLU)
|
||||||
|
self.assertTrue(hasattr(ffn_block, "ln1"))
|
||||||
|
self.assertTrue(hasattr(ffn_block, "ln2"))
|
||||||
|
|
||||||
|
traj_block = _build_block(TRAJ_MIXER_ARCHITECTURE)
|
||||||
|
self.assertIsInstance(traj_block, TrajMixerBlock)
|
||||||
|
self.assertIsInstance(traj_block.mlp, TrajMixer)
|
||||||
|
self.assertTrue(hasattr(traj_block, "ln1"))
|
||||||
|
self.assertFalse(hasattr(traj_block, "ln2"))
|
||||||
|
|
||||||
|
def test_both_architectures_forward_and_backward(self) -> None:
|
||||||
|
for architecture in (
|
||||||
|
TRANSFORMER_FFN_ARCHITECTURE,
|
||||||
|
TRAJ_MIXER_ARCHITECTURE,
|
||||||
|
):
|
||||||
|
with self.subTest(architecture=architecture):
|
||||||
|
torch.manual_seed(0)
|
||||||
|
block = _build_block(architecture)
|
||||||
|
x = torch.randn(2, 5, 12, requires_grad=True)
|
||||||
|
|
||||||
|
output = block(x)
|
||||||
|
self.assertEqual(output.shape, x.shape)
|
||||||
|
output.square().mean().backward()
|
||||||
|
|
||||||
|
self.assertIsNotNone(x.grad)
|
||||||
|
self.assertTrue(torch.isfinite(x.grad).all())
|
||||||
|
self.assertGreater(x.grad.abs().sum().item(), 0.0)
|
||||||
|
self.assertIsNotNone(block.attn.qkv.weight.grad)
|
||||||
|
self.assertGreater(
|
||||||
|
block.attn.qkv.weight.grad.abs().sum().item(),
|
||||||
|
0.0,
|
||||||
|
)
|
||||||
|
|
||||||
|
if architecture == TRANSFORMER_FFN_ARCHITECTURE:
|
||||||
|
branch_parameters = (
|
||||||
|
block.mlp.w1.weight,
|
||||||
|
block.mlp.w2.weight,
|
||||||
|
block.mlp.w3.weight,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
branch_parameters = (
|
||||||
|
block.mlp.intra_gate_proj,
|
||||||
|
block.mlp.intra_value_proj,
|
||||||
|
block.mlp.output_proj,
|
||||||
|
)
|
||||||
|
for parameter in branch_parameters:
|
||||||
|
self.assertIsNotNone(parameter.grad)
|
||||||
|
self.assertTrue(torch.isfinite(parameter.grad).all())
|
||||||
|
self.assertGreater(parameter.grad.abs().sum().item(), 0.0)
|
||||||
|
|
||||||
|
def test_unknown_architecture_is_rejected(self) -> None:
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
_build_block("unknown_architecture")
|
||||||
|
with self.assertRaisesRegex(ValueError, "model_architecture is required"):
|
||||||
|
_build_model(None)
|
||||||
|
|
||||||
|
def test_deephealth_rejects_fewer_than_one_layer(self) -> None:
|
||||||
|
for n_layer in (0, -1):
|
||||||
|
with self.subTest(n_layer=n_layer):
|
||||||
|
with self.assertRaisesRegex(ValueError, "n_layer must be >= 1"):
|
||||||
|
_build_model(
|
||||||
|
TRANSFORMER_FFN_ARCHITECTURE,
|
||||||
|
n_layer=n_layer,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_deephealth_uses_factory_and_strictly_reloads_both_models(self) -> None:
|
||||||
|
for architecture, block_class in (
|
||||||
|
(TRANSFORMER_FFN_ARCHITECTURE, TransformerFFNBlock),
|
||||||
|
(TRAJ_MIXER_ARCHITECTURE, TrajMixerBlock),
|
||||||
|
):
|
||||||
|
with self.subTest(architecture=architecture):
|
||||||
|
model = _build_model(architecture)
|
||||||
|
self.assertEqual(model.model_architecture, architecture)
|
||||||
|
self.assertIsInstance(model.blocks[0], block_class)
|
||||||
|
self.assertEqual(
|
||||||
|
detect_model_architecture_from_state_dict(
|
||||||
|
model.state_dict()
|
||||||
|
),
|
||||||
|
architecture,
|
||||||
|
)
|
||||||
|
|
||||||
|
reloaded = _build_model(architecture)
|
||||||
|
incompatible = reloaded.load_state_dict(
|
||||||
|
model.state_dict(),
|
||||||
|
strict=True,
|
||||||
|
)
|
||||||
|
self.assertEqual(incompatible.missing_keys, [])
|
||||||
|
self.assertEqual(incompatible.unexpected_keys, [])
|
||||||
|
|
||||||
|
|
||||||
|
class ModelArchitectureResolutionTest(unittest.TestCase):
|
||||||
|
def setUp(self) -> None:
|
||||||
|
self.ffn_block = _build_block(TRANSFORMER_FFN_ARCHITECTURE)
|
||||||
|
self.traj_block = _build_block(TRAJ_MIXER_ARCHITECTURE)
|
||||||
|
self.ffn_state = _as_model_state_dict(self.ffn_block)
|
||||||
|
self.traj_state = _as_model_state_dict(self.traj_block)
|
||||||
|
|
||||||
|
def test_state_dict_detection_recognizes_both_architectures(self) -> None:
|
||||||
|
self.assertEqual(
|
||||||
|
detect_model_architecture_from_state_dict(self.ffn_state),
|
||||||
|
TRANSFORMER_FFN_ARCHITECTURE,
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
detect_model_architecture_from_state_dict(self.traj_state),
|
||||||
|
TRAJ_MIXER_ARCHITECTURE,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_explicit_markers_resolve_when_checkpoint_matches(self) -> None:
|
||||||
|
for architecture, state_dict in (
|
||||||
|
(TRANSFORMER_FFN_ARCHITECTURE, self.ffn_state),
|
||||||
|
(TRAJ_MIXER_ARCHITECTURE, self.traj_state),
|
||||||
|
):
|
||||||
|
with self.subTest(architecture=architecture):
|
||||||
|
self.assertEqual(
|
||||||
|
resolve_model_architecture(
|
||||||
|
{"model_architecture": architecture},
|
||||||
|
state_dict,
|
||||||
|
),
|
||||||
|
architecture,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_architecture_marker_is_required_for_checkpoint_loading(self) -> None:
|
||||||
|
with self.assertRaisesRegex(ValueError, "model_architecture is required"):
|
||||||
|
resolve_model_architecture({}, self.ffn_state)
|
||||||
|
with self.assertRaisesRegex(ValueError, "model_architecture is required"):
|
||||||
|
resolve_model_architecture(None, self.traj_state)
|
||||||
|
|
||||||
|
def test_explicit_marker_conflicting_with_state_dict_is_rejected(self) -> None:
|
||||||
|
conflicts = (
|
||||||
|
(TRANSFORMER_FFN_ARCHITECTURE, self.traj_state),
|
||||||
|
(TRAJ_MIXER_ARCHITECTURE, self.ffn_state),
|
||||||
|
)
|
||||||
|
for architecture, state_dict in conflicts:
|
||||||
|
with self.subTest(architecture=architecture):
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
resolve_model_architecture(
|
||||||
|
{"model_architecture": architecture},
|
||||||
|
state_dict,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_unknown_marker_and_ambiguous_state_dict_are_rejected(self) -> None:
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
resolve_model_architecture(
|
||||||
|
{"model_architecture": "traj_mixer_v4"}
|
||||||
|
)
|
||||||
|
|
||||||
|
ambiguous_state = dict(self.ffn_state)
|
||||||
|
ambiguous_state.update(self.traj_state)
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
detect_model_architecture_from_state_dict(ambiguous_state)
|
||||||
|
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
detect_model_architecture_from_state_dict(
|
||||||
|
{"token_embedding.weight": torch.empty(2, 2)}
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_ffn_block_schema_is_stable_and_strictly_loadable(self) -> None:
|
||||||
|
expected_keys = {
|
||||||
|
"attn.time_bias_scale",
|
||||||
|
"attn.qkv.weight",
|
||||||
|
"attn.out_proj.weight",
|
||||||
|
"attn.rbf_proj.weight",
|
||||||
|
"mlp.w1.weight",
|
||||||
|
"mlp.w1.bias",
|
||||||
|
"mlp.w2.weight",
|
||||||
|
"mlp.w2.bias",
|
||||||
|
"mlp.w3.weight",
|
||||||
|
"mlp.w3.bias",
|
||||||
|
"ln1.weight",
|
||||||
|
"ln1.bias",
|
||||||
|
"ln2.weight",
|
||||||
|
"ln2.bias",
|
||||||
|
}
|
||||||
|
state = self.ffn_block.state_dict()
|
||||||
|
self.assertSetEqual(set(state), expected_keys)
|
||||||
|
self.assertEqual(tuple(state["mlp.w1.weight"].shape), (30, 12))
|
||||||
|
self.assertEqual(tuple(state["mlp.w3.weight"].shape), (12, 30))
|
||||||
|
|
||||||
|
reloaded = _build_block(TRANSFORMER_FFN_ARCHITECTURE)
|
||||||
|
incompatible = reloaded.load_state_dict(state, strict=True)
|
||||||
|
self.assertEqual(incompatible.missing_keys, [])
|
||||||
|
self.assertEqual(incompatible.unexpected_keys, [])
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
46
test_temporal_attention.py
Normal file
46
test_temporal_attention.py
Normal file
@@ -0,0 +1,46 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from backbones import TemporalAttention
|
||||||
|
|
||||||
|
|
||||||
|
class TemporalAttentionTest(unittest.TestCase):
|
||||||
|
def test_zero_rbf_bias_has_live_projection_gradient(self) -> None:
|
||||||
|
torch.manual_seed(0)
|
||||||
|
attention = TemporalAttention(
|
||||||
|
n_embd=12,
|
||||||
|
n_head=3,
|
||||||
|
use_time_rope=False,
|
||||||
|
use_rbf_bias=True,
|
||||||
|
)
|
||||||
|
features = torch.randn(2, 4, 4, 16)
|
||||||
|
target = torch.randn(2, 4, 4, 3)
|
||||||
|
|
||||||
|
initial_bias = (
|
||||||
|
attention.time_bias_scale.tanh()
|
||||||
|
* attention.rbf_proj(features)
|
||||||
|
)
|
||||||
|
torch.testing.assert_close(initial_bias, torch.zeros_like(initial_bias))
|
||||||
|
|
||||||
|
(initial_bias * target).sum().backward()
|
||||||
|
projection_grad = attention.rbf_proj.weight.grad
|
||||||
|
self.assertIsNotNone(projection_grad)
|
||||||
|
self.assertGreater(projection_grad.abs().sum().item(), 0.0)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
attention.rbf_proj.weight.add_(projection_grad, alpha=-1e-3)
|
||||||
|
attention.zero_grad(set_to_none=True)
|
||||||
|
updated_bias = (
|
||||||
|
attention.time_bias_scale.tanh()
|
||||||
|
* attention.rbf_proj(features)
|
||||||
|
)
|
||||||
|
(updated_bias * target).sum().backward()
|
||||||
|
|
||||||
|
scale_grad = attention.time_bias_scale.grad
|
||||||
|
self.assertIsNotNone(scale_grad)
|
||||||
|
self.assertGreater(scale_grad.abs().item(), 0.0)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
166
test_traj_mixer.py
Normal file
166
test_traj_mixer.py
Normal file
@@ -0,0 +1,166 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from backbones import TrajMixer
|
||||||
|
|
||||||
|
|
||||||
|
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()), 32_040)
|
||||||
|
self.assertEqual(tuple(mixer.norm.normalized_shape), (120,))
|
||||||
|
self.assertEqual(tuple(mixer.intra_gate_logits.shape), (10, 12))
|
||||||
|
torch.testing.assert_close(
|
||||||
|
torch.sigmoid(mixer.intra_gate_logits.detach()),
|
||||||
|
torch.full((10, 12), 0.1),
|
||||||
|
)
|
||||||
|
self.assertEqual(mixer.intra_hidden, 48)
|
||||||
|
self.assertEqual(
|
||||||
|
tuple(mixer.intra_gate_proj.shape),
|
||||||
|
(10, 12, 48),
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
tuple(mixer.intra_value_proj.shape),
|
||||||
|
(10, 12, 48),
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
tuple(mixer.intra_output_proj.shape),
|
||||||
|
(10, 48, 12),
|
||||||
|
)
|
||||||
|
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_zero_final_output_projection_makes_mixer_identity(self) -> None:
|
||||||
|
torch.manual_seed(0)
|
||||||
|
mixer = TrajMixer(12, n_head=3, dropout=0.0)
|
||||||
|
with torch.no_grad():
|
||||||
|
mixer.output_proj.zero_()
|
||||||
|
x = torch.randn(2, 5, 12)
|
||||||
|
torch.testing.assert_close(mixer(x), x)
|
||||||
|
|
||||||
|
def test_forward_matches_single_outer_residual_formula(self) -> None:
|
||||||
|
torch.manual_seed(0)
|
||||||
|
mixer = TrajMixer(12, n_head=3, dropout=0.0)
|
||||||
|
mixer.eval()
|
||||||
|
x = torch.randn(2, 5, 12)
|
||||||
|
|
||||||
|
grouped = mixer.norm(x).reshape(2, 5, 3, 4)
|
||||||
|
intra_output = mixer._intra_mix(grouped)
|
||||||
|
static_gate = torch.sigmoid(mixer.intra_gate_logits).view(
|
||||||
|
1, 1, 3, 4
|
||||||
|
)
|
||||||
|
mixed_input = grouped + static_gate * intra_output
|
||||||
|
update = mixer._cross_mix(mixed_input).reshape(2, 5, 12)
|
||||||
|
|
||||||
|
torch.testing.assert_close(mixer(x), x + update)
|
||||||
|
|
||||||
|
def test_intra_stage_is_independent_across_groups(self) -> None:
|
||||||
|
torch.manual_seed(0)
|
||||||
|
mixer = TrajMixer(12, n_head=3, dropout=0.0)
|
||||||
|
mixer.eval()
|
||||||
|
|
||||||
|
grouped = torch.randn(2, 4, 3, 4)
|
||||||
|
changed = grouped.clone()
|
||||||
|
changed[:, :, 1, :] += torch.randn_like(changed[:, :, 1, :])
|
||||||
|
|
||||||
|
original_out = mixer._intra_mix(grouped)
|
||||||
|
changed_out = mixer._intra_mix(changed)
|
||||||
|
unchanged_groups = torch.tensor([0, 2])
|
||||||
|
torch.testing.assert_close(
|
||||||
|
original_out.index_select(2, unchanged_groups),
|
||||||
|
changed_out.index_select(2, unchanged_groups),
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_cross_stage_mixes_groups_without_mixing_coordinates(self) -> None:
|
||||||
|
mixer = TrajMixer(6, n_head=3, dropout=0.0)
|
||||||
|
mixer.eval()
|
||||||
|
with torch.no_grad():
|
||||||
|
mixer.gate_proj.zero_()
|
||||||
|
mixer.value_proj.zero_()
|
||||||
|
mixer.output_proj.zero_()
|
||||||
|
|
||||||
|
# Coordinate 0 reads group 0 through hidden unit 0 and writes it
|
||||||
|
# into group 1. Coordinate 1 must remain independent.
|
||||||
|
mixer.gate_proj[0, 0, 0] = 1.0
|
||||||
|
mixer.value_proj[0, 0, 0] = 1.0
|
||||||
|
mixer.output_proj[0, 0, 1] = 1.0
|
||||||
|
|
||||||
|
grouped = torch.tensor(
|
||||||
|
[[[
|
||||||
|
[-1.0, 4.0],
|
||||||
|
[0.0, 5.0],
|
||||||
|
[1.0, 6.0],
|
||||||
|
]]]
|
||||||
|
)
|
||||||
|
changed = grouped.clone()
|
||||||
|
changed[0, 0, 0, 0] = 2.0
|
||||||
|
|
||||||
|
original_out = mixer._cross_mix(grouped)
|
||||||
|
changed_out = mixer._cross_mix(changed)
|
||||||
|
|
||||||
|
self.assertNotEqual(
|
||||||
|
original_out[0, 0, 1, 0].item(),
|
||||||
|
changed_out[0, 0, 1, 0].item(),
|
||||||
|
)
|
||||||
|
torch.testing.assert_close(
|
||||||
|
original_out[..., 1],
|
||||||
|
changed_out[..., 1],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_mixer_does_not_mix_sequence_positions(self) -> None:
|
||||||
|
torch.manual_seed(0)
|
||||||
|
mixer = TrajMixer(12, n_head=3, dropout=0.0)
|
||||||
|
mixer.eval()
|
||||||
|
x = torch.randn(2, 5, 12)
|
||||||
|
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_every_projection_family(self) -> None:
|
||||||
|
torch.manual_seed(1)
|
||||||
|
mixer = TrajMixer(12, n_head=3, dropout=0.0)
|
||||||
|
x = torch.randn(2, 4, 12, requires_grad=True)
|
||||||
|
|
||||||
|
mixer(x).square().mean().backward()
|
||||||
|
|
||||||
|
self.assertIsNotNone(x.grad)
|
||||||
|
self.assertTrue(torch.isfinite(x.grad).all())
|
||||||
|
self.assertGreater(x.grad.abs().sum().item(), 0.0)
|
||||||
|
projection_names = (
|
||||||
|
"intra_gate_proj",
|
||||||
|
"intra_value_proj",
|
||||||
|
"intra_output_proj",
|
||||||
|
"gate_proj",
|
||||||
|
"value_proj",
|
||||||
|
"output_proj",
|
||||||
|
)
|
||||||
|
for name in projection_names:
|
||||||
|
parameter = getattr(mixer, name)
|
||||||
|
self.assertIsNotNone(parameter.grad, name)
|
||||||
|
self.assertTrue(torch.isfinite(parameter.grad).all(), name)
|
||||||
|
self.assertGreater(parameter.grad.abs().sum().item(), 0.0, name)
|
||||||
|
|
||||||
|
def test_invalid_group_partition_is_rejected(self) -> None:
|
||||||
|
with self.assertRaisesRegex(ValueError, "divisible"):
|
||||||
|
TrajMixer(n_embd=121, n_head=10)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -27,12 +27,11 @@ from tqdm.auto import tqdm
|
|||||||
|
|
||||||
from dataset import AllFutureHealthDataset, all_future_collate_fn
|
from dataset import AllFutureHealthDataset, all_future_collate_fn
|
||||||
from losses import build_loss
|
from losses import build_loss
|
||||||
from models import (
|
from model_architectures import (
|
||||||
EVENT_TRAJECTORY_ARCHITECTURE,
|
DEFAULT_MODEL_ARCHITECTURE,
|
||||||
MODEL_SIZE_NAMES,
|
SUPPORTED_MODEL_ARCHITECTURES,
|
||||||
DeepHealth,
|
|
||||||
resolve_model_size,
|
|
||||||
)
|
)
|
||||||
|
from models import DeepHealth
|
||||||
from targets import CHECKUP_IDX, PAD_IDX
|
from targets import CHECKUP_IDX, PAD_IDX
|
||||||
from train_util import (
|
from train_util import (
|
||||||
configure_torch_for_training,
|
configure_torch_for_training,
|
||||||
@@ -70,6 +69,7 @@ def parse_args() -> argparse.Namespace:
|
|||||||
|
|
||||||
parser.add_argument("--data_prefix", type=str, default="ukb")
|
parser.add_argument("--data_prefix", type=str, default="ukb")
|
||||||
parser.add_argument("--labels_file", type=str, default="labels.csv")
|
parser.add_argument("--labels_file", type=str, default="labels.csv")
|
||||||
|
parser.add_argument("--runs_root", type=str, default="runs")
|
||||||
parser.add_argument("--seed", type=int, default=42)
|
parser.add_argument("--seed", type=int, default=42)
|
||||||
parser.add_argument("--extra_info_types_file", type=str, default=None)
|
parser.add_argument("--extra_info_types_file", type=str, default=None)
|
||||||
|
|
||||||
@@ -83,13 +83,9 @@ def parse_args() -> argparse.Namespace:
|
|||||||
parser.add_argument("--min_future_events", type=int, default=1)
|
parser.add_argument("--min_future_events", type=int, default=1)
|
||||||
parser.add_argument("--validation_query_seed", type=int, default=None)
|
parser.add_argument("--validation_query_seed", type=int, default=None)
|
||||||
|
|
||||||
parser.add_argument(
|
parser.add_argument("--n_embd", type=int, default=120)
|
||||||
"--model_size",
|
parser.add_argument("--n_head", type=int, default=10)
|
||||||
type=str,
|
parser.add_argument("--n_layer", type=int, default=12)
|
||||||
default="nano",
|
|
||||||
choices=MODEL_SIZE_NAMES,
|
|
||||||
)
|
|
||||||
parser.add_argument("--n_reasoning_rounds", type=int, default=12)
|
|
||||||
parser.add_argument("--n_bins", type=int, default=16)
|
parser.add_argument("--n_bins", type=int, default=16)
|
||||||
parser.add_argument("--extra_pool_reduce", type=str, default="mean",
|
parser.add_argument("--extra_pool_reduce", type=str, default="mean",
|
||||||
choices=["mean", "sum"])
|
choices=["mean", "sum"])
|
||||||
@@ -98,6 +94,12 @@ def parse_args() -> argparse.Namespace:
|
|||||||
parser.add_argument("--dist_mode", type=str, default="exponential",
|
parser.add_argument("--dist_mode", type=str, default="exponential",
|
||||||
choices=["exponential", "weibull", "mixed"])
|
choices=["exponential", "weibull", "mixed"])
|
||||||
parser.add_argument("--dropout", type=float, default=0.0)
|
parser.add_argument("--dropout", type=float, default=0.0)
|
||||||
|
parser.add_argument(
|
||||||
|
"--model_architecture",
|
||||||
|
type=str,
|
||||||
|
default=DEFAULT_MODEL_ARCHITECTURE,
|
||||||
|
choices=SUPPORTED_MODEL_ARCHITECTURES,
|
||||||
|
)
|
||||||
|
|
||||||
parser.add_argument("--batch_size", type=int, default=128)
|
parser.add_argument("--batch_size", type=int, default=128)
|
||||||
parser.add_argument("--base_lr", type=float, default=3e-4)
|
parser.add_argument("--base_lr", type=float, default=3e-4)
|
||||||
@@ -154,8 +156,9 @@ def move_batch_to_device(batch: Dict[str, torch.Tensor], device: torch.device) -
|
|||||||
def build_model(args: argparse.Namespace, dataset: AllFutureHealthDataset) -> DeepHealth:
|
def build_model(args: argparse.Namespace, dataset: AllFutureHealthDataset) -> DeepHealth:
|
||||||
return DeepHealth(
|
return DeepHealth(
|
||||||
vocab_size=dataset.vocab_size,
|
vocab_size=dataset.vocab_size,
|
||||||
model_size=args.model_size,
|
n_embd=args.n_embd,
|
||||||
n_reasoning_rounds=args.n_reasoning_rounds,
|
n_head=args.n_head,
|
||||||
|
n_layer=args.n_layer,
|
||||||
n_types=dataset.n_types,
|
n_types=dataset.n_types,
|
||||||
n_cont_types=dataset.n_cont_types,
|
n_cont_types=dataset.n_cont_types,
|
||||||
n_categories=dataset.n_categories,
|
n_categories=dataset.n_categories,
|
||||||
@@ -166,6 +169,7 @@ def build_model(args: argparse.Namespace, dataset: AllFutureHealthDataset) -> De
|
|||||||
time_mode=args.time_mode,
|
time_mode=args.time_mode,
|
||||||
dist_mode=args.dist_mode,
|
dist_mode=args.dist_mode,
|
||||||
dropout=args.dropout,
|
dropout=args.dropout,
|
||||||
|
model_architecture=args.model_architecture,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -300,17 +304,12 @@ def build_metadata(
|
|||||||
val_subset,
|
val_subset,
|
||||||
test_subset,
|
test_subset,
|
||||||
) -> Dict[str, Any]:
|
) -> Dict[str, Any]:
|
||||||
size_config = resolve_model_size(args.model_size)
|
|
||||||
return {
|
return {
|
||||||
"run_name": run_name,
|
"run_name": run_name,
|
||||||
"dataset_class": "AllFutureHealthDataset",
|
"dataset_class": "AllFutureHealthDataset",
|
||||||
"collate_fn": "all_future_collate_fn",
|
"collate_fn": "all_future_collate_fn",
|
||||||
"model_class": "DeepHealth",
|
"model_class": "DeepHealth",
|
||||||
"model_architecture": EVENT_TRAJECTORY_ARCHITECTURE,
|
"model_architecture": args.model_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_target_mode": "all_future",
|
"model_target_mode": "all_future",
|
||||||
"target_mode": "all_future",
|
"target_mode": "all_future",
|
||||||
"dist_mode": args.dist_mode,
|
"dist_mode": args.dist_mode,
|
||||||
@@ -348,26 +347,14 @@ def main() -> None:
|
|||||||
configure_torch_for_training(device)
|
configure_torch_for_training(device)
|
||||||
|
|
||||||
run_dir, run_name = create_unique_run_dir(
|
run_dir, run_name = create_unique_run_dir(
|
||||||
lambda timestamp: (
|
lambda timestamp: f"{args.time_mode}_{args.dist_mode}_all_future_pure_disease_{timestamp}",
|
||||||
f"{args.model_size}_r{args.n_reasoning_rounds}_"
|
runs_root=Path(args.runs_root) / args.model_architecture,
|
||||||
f"{args.time_mode}_{args.dist_mode}_"
|
|
||||||
f"all_future_pure_disease_{timestamp}"
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
logger = setup_logging(run_dir)
|
logger = setup_logging(run_dir)
|
||||||
|
|
||||||
logger.info(f"Starting all-future training run: {run_name}")
|
logger.info(f"Starting all-future training run: {run_name}")
|
||||||
logger.info(f"Device: {device}")
|
logger.info(f"Device: {device}")
|
||||||
size_config = resolve_model_size(args.model_size)
|
logger.info(f"Model architecture: {args.model_architecture}")
|
||||||
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"extra_info_types: {format_extra_info_types(args.extra_info_types)}")
|
||||||
|
|
||||||
logger.info("Loading all-future datasets...")
|
logger.info("Loading all-future datasets...")
|
||||||
|
|||||||
@@ -24,13 +24,11 @@ from tqdm.auto import tqdm
|
|||||||
|
|
||||||
from dataset import HealthDataset, collate_fn
|
from dataset import HealthDataset, collate_fn
|
||||||
from losses import build_loss
|
from losses import build_loss
|
||||||
from models import (
|
from model_architectures import (
|
||||||
EVENT_TRAJECTORY_ARCHITECTURE,
|
DEFAULT_MODEL_ARCHITECTURE,
|
||||||
MODEL_SIZE_NAMES,
|
SUPPORTED_MODEL_ARCHITECTURES,
|
||||||
DeepHealth,
|
|
||||||
DeepHealthOutput,
|
|
||||||
resolve_model_size,
|
|
||||||
)
|
)
|
||||||
|
from models import DeepHealth, DeepHealthOutput
|
||||||
from readouts import build_readout
|
from readouts import build_readout
|
||||||
from targets import CHECKUP_IDX, NO_EVENT_IDX, PAD_IDX
|
from targets import CHECKUP_IDX, NO_EVENT_IDX, PAD_IDX
|
||||||
from train_util import (
|
from train_util import (
|
||||||
@@ -68,6 +66,7 @@ def parse_args() -> argparse.Namespace:
|
|||||||
|
|
||||||
parser.add_argument("--data_prefix", type=str, default="ukb")
|
parser.add_argument("--data_prefix", type=str, default="ukb")
|
||||||
parser.add_argument("--labels_file", type=str, default="labels.csv")
|
parser.add_argument("--labels_file", type=str, default="labels.csv")
|
||||||
|
parser.add_argument("--runs_root", type=str, default="runs")
|
||||||
parser.add_argument("--seed", type=int, default=42)
|
parser.add_argument("--seed", type=int, default=42)
|
||||||
parser.add_argument("--extra_info_types_file", type=str, default=None)
|
parser.add_argument("--extra_info_types_file", type=str, default=None)
|
||||||
parser.add_argument("--no_event_interval_years", type=float, default=5.0)
|
parser.add_argument("--no_event_interval_years", type=float, default=5.0)
|
||||||
@@ -80,19 +79,21 @@ def parse_args() -> argparse.Namespace:
|
|||||||
parser.add_argument("--val_eid_file", type=str, default="ukb_val_eid.csv")
|
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("--test_eid_file", type=str, default="ukb_test_eid.csv")
|
||||||
|
|
||||||
parser.add_argument(
|
parser.add_argument("--n_embd", type=int, default=120)
|
||||||
"--model_size",
|
parser.add_argument("--n_head", type=int, default=10)
|
||||||
type=str,
|
parser.add_argument("--n_layer", type=int, default=12)
|
||||||
default="nano",
|
|
||||||
choices=MODEL_SIZE_NAMES,
|
|
||||||
)
|
|
||||||
parser.add_argument("--n_reasoning_rounds", type=int, default=12)
|
|
||||||
parser.add_argument("--n_bins", type=int, default=16)
|
parser.add_argument("--n_bins", type=int, default=16)
|
||||||
parser.add_argument("--extra_pool_reduce", type=str, default="mean",
|
parser.add_argument("--extra_pool_reduce", type=str, default="mean",
|
||||||
choices=["mean", "sum"])
|
choices=["mean", "sum"])
|
||||||
parser.add_argument("--time_mode", type=str, default="relative",
|
parser.add_argument("--time_mode", type=str, default="relative",
|
||||||
choices=["relative", "absolute"])
|
choices=["relative", "absolute"])
|
||||||
parser.add_argument("--dropout", type=float, default=0.0)
|
parser.add_argument("--dropout", type=float, default=0.0)
|
||||||
|
parser.add_argument(
|
||||||
|
"--model_architecture",
|
||||||
|
type=str,
|
||||||
|
default=DEFAULT_MODEL_ARCHITECTURE,
|
||||||
|
choices=SUPPORTED_MODEL_ARCHITECTURES,
|
||||||
|
)
|
||||||
|
|
||||||
parser.add_argument("--target_mode", type=str, default="uts",
|
parser.add_argument("--target_mode", type=str, default="uts",
|
||||||
choices=["delphi2m", "uts"])
|
choices=["delphi2m", "uts"])
|
||||||
@@ -160,8 +161,9 @@ def move_batch_to_device(batch: Dict[str, torch.Tensor], device: torch.device) -
|
|||||||
def build_model(args: argparse.Namespace, dataset: HealthDataset) -> DeepHealth:
|
def build_model(args: argparse.Namespace, dataset: HealthDataset) -> DeepHealth:
|
||||||
return DeepHealth(
|
return DeepHealth(
|
||||||
vocab_size=dataset.vocab_size,
|
vocab_size=dataset.vocab_size,
|
||||||
model_size=args.model_size,
|
n_embd=args.n_embd,
|
||||||
n_reasoning_rounds=args.n_reasoning_rounds,
|
n_head=args.n_head,
|
||||||
|
n_layer=args.n_layer,
|
||||||
n_types=dataset.n_types,
|
n_types=dataset.n_types,
|
||||||
n_cont_types=dataset.n_cont_types,
|
n_cont_types=dataset.n_cont_types,
|
||||||
n_categories=dataset.n_categories,
|
n_categories=dataset.n_categories,
|
||||||
@@ -172,6 +174,7 @@ def build_model(args: argparse.Namespace, dataset: HealthDataset) -> DeepHealth:
|
|||||||
time_mode=args.time_mode,
|
time_mode=args.time_mode,
|
||||||
dist_mode="exponential",
|
dist_mode="exponential",
|
||||||
dropout=args.dropout,
|
dropout=args.dropout,
|
||||||
|
model_architecture=args.model_architecture,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -487,17 +490,12 @@ def build_metadata(
|
|||||||
val_subset,
|
val_subset,
|
||||||
test_subset,
|
test_subset,
|
||||||
) -> Dict[str, Any]:
|
) -> Dict[str, Any]:
|
||||||
size_config = resolve_model_size(args.model_size)
|
|
||||||
return {
|
return {
|
||||||
"run_name": run_name,
|
"run_name": run_name,
|
||||||
"dataset_class": "NextStepHealthDataset",
|
"dataset_class": "NextStepHealthDataset",
|
||||||
"collate_fn": "next_step_collate_fn",
|
"collate_fn": "next_step_collate_fn",
|
||||||
"model_class": "DeepHealth",
|
"model_class": "DeepHealth",
|
||||||
"model_architecture": EVENT_TRAJECTORY_ARCHITECTURE,
|
"model_architecture": args.model_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_target_mode": "next_token",
|
"model_target_mode": "next_token",
|
||||||
"target_mode": args.target_mode,
|
"target_mode": args.target_mode,
|
||||||
"dist_mode": "exponential",
|
"dist_mode": "exponential",
|
||||||
@@ -533,26 +531,16 @@ def main() -> None:
|
|||||||
|
|
||||||
run_dir, run_name = create_unique_run_dir(
|
run_dir, run_name = create_unique_run_dir(
|
||||||
lambda timestamp: (
|
lambda timestamp: (
|
||||||
f"{args.model_size}_r{args.n_reasoning_rounds}_"
|
f"{args.time_mode}_exponential_next_token_{args.target_mode}_"
|
||||||
f"{args.time_mode}_exponential_"
|
|
||||||
f"next_token_{args.target_mode}_"
|
|
||||||
f"gap_{args.no_event_interval_years:g}y_{timestamp}"
|
f"gap_{args.no_event_interval_years:g}y_{timestamp}"
|
||||||
)
|
),
|
||||||
|
runs_root=Path(args.runs_root) / args.model_architecture,
|
||||||
)
|
)
|
||||||
logger = setup_logging(run_dir)
|
logger = setup_logging(run_dir)
|
||||||
|
|
||||||
logger.info(f"Starting next-step training run: {run_name}")
|
logger.info(f"Starting next-step training run: {run_name}")
|
||||||
logger.info(f"Device: {device}")
|
logger.info(f"Device: {device}")
|
||||||
size_config = resolve_model_size(args.model_size)
|
logger.info(f"Model architecture: {args.model_architecture}")
|
||||||
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"extra_info_types: {format_extra_info_types(args.extra_info_types)}")
|
||||||
logger.info(f"readout={args.readout_name}, target_mode={args.target_mode}")
|
logger.info(f"readout={args.readout_name}, target_mode={args.target_mode}")
|
||||||
|
|
||||||
|
|||||||
@@ -301,7 +301,7 @@ def build_optimizer(args: Any, model: DeepHealth) -> AdamW:
|
|||||||
|
|
||||||
|
|
||||||
def get_model_parameter_counts(model: torch.nn.Module) -> Dict[str, int]:
|
def get_model_parameter_counts(model: torch.nn.Module) -> Dict[str, int]:
|
||||||
"""Return stable parameter-count fields for logs and train_config.json."""
|
"""Return stable total and trainable parameter counts."""
|
||||||
return {
|
return {
|
||||||
"model_parameter_count": sum(
|
"model_parameter_count": sum(
|
||||||
parameter.numel() for parameter in model.parameters()
|
parameter.numel() for parameter in model.parameters()
|
||||||
|
|||||||
Reference in New Issue
Block a user