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