refactor: isolate Delphi2M next-token pipeline
This commit is contained in:
32
README.md
32
README.md
@@ -59,8 +59,8 @@ python prepare_data.py
|
||||
`dataset.py` 提供两个 dataset:
|
||||
|
||||
- `NextStepHealthDataset`
|
||||
- 用于 next-token / next-time-point 监督
|
||||
- 对应 `Delphi2MLoss` 和 `UniqueTimeSetExponentialLoss`
|
||||
- 仅用于 absolute-time Delphi2M next-token 复现
|
||||
- 对应 `Delphi2MLoss`
|
||||
|
||||
- `AllFutureHealthDataset`
|
||||
- 用于 query-conditioned all-future 监督
|
||||
@@ -173,7 +173,7 @@ extra-info 不再通过独立的 `BaselineEncoder` 或 `CrossAttention` 注入
|
||||
如果需要拿到完整 next-token 输出,可使用结构化返回:
|
||||
|
||||
```python
|
||||
out = model(..., target_mode="next_token", return_output=True)
|
||||
out = model(..., return_output=True)
|
||||
out.hidden # disease tokens + pooled extra-info readout tokens
|
||||
out.time_seq # 与 hidden 对齐的时间
|
||||
out.padding_mask # 与 hidden 对齐的有效位置
|
||||
@@ -187,8 +187,7 @@ model = DeepHealth(
|
||||
vocab_size=dataset.vocab_size,
|
||||
n_embd=120,
|
||||
n_head=10,
|
||||
n_hist_layer=12,
|
||||
n_tab_layer=4, # 兼容旧配置;当前不再创建独立 tabular transformer
|
||||
n_layer=12,
|
||||
n_types=dataset.n_types,
|
||||
n_cont_types=dataset.n_cont_types,
|
||||
n_categories=dataset.n_categories,
|
||||
@@ -204,14 +203,13 @@ model = DeepHealth(
|
||||
next-token 监督:
|
||||
|
||||
- `Delphi2MLoss`
|
||||
- `UniqueTimeSetExponentialLoss`
|
||||
|
||||
next-token 训练中,模型会请求 `return_output=True`,因此 loss 的预测位置包括:
|
||||
|
||||
- 原 disease token readout 位置
|
||||
- 同一时间点 extra hidden 池化后的 pooled extra-info readout token
|
||||
|
||||
pooled extra-info readout token 的监督目标在训练时动态构造:对 pooled extra-info readout token 的时间 `t`,寻找该患者 `t` 之后的下一个 disease 事件时间;`delphi2m` 使用第一个未来事件作为 next-token target,`uts` 使用下一唯一时间点上的事件集合做 multi-hot target。若该 extra-info 时间点之后没有未来 disease target,则该位置不参与 loss。
|
||||
pooled extra-info readout token 的监督目标在训练时动态构造:对 pooled extra-info readout token 的时间 `t`,寻找该患者 `t` 之后的第一个 disease 事件作为 next-token target。若该时间点之后没有未来 disease target,则该位置不参与 loss。
|
||||
|
||||
all-future / query-conditioned 监督:
|
||||
|
||||
@@ -221,17 +219,15 @@ all-future / query-conditioned 监督:
|
||||
|
||||
all-future 训练只读出 `t_query` 对应的 query hidden。展开的 extra-info tokens 作为主序列上下文输入,但不会被单独读出,也不会被纳入 loss 监督。
|
||||
|
||||
`UniqueTimeSetExponentialLoss` 的 observed term 固定使用 sum reduction,不再暴露旧的 `observed_reduction` 参数。
|
||||
|
||||
## 训练
|
||||
|
||||
当前提供两类训练入口:
|
||||
|
||||
- `train_next_step.py`
|
||||
- 使用 `NextStepHealthDataset`
|
||||
- `--target_mode delphi2m` 默认搭配 `Delphi2MLoss` + `token` readout
|
||||
- `--target_mode uts` 默认搭配 `UniqueTimeSetExponentialLoss` + `same_time_group_end` readout
|
||||
- 当前 next-token 训练只支持 exponential time loss
|
||||
- 仅用于 Delphi2M 复现,固定 `time_mode=absolute`、`target_mode=delphi2m`
|
||||
- 固定使用 `Delphi2MLoss` + `token` readout
|
||||
- next-token 训练只支持 exponential time loss
|
||||
- 展开的 extra-info tokens 进入主序列;读出端 pooled extra-info tokens 会加入 prediction/loss 监督
|
||||
- `train_all_future.py`
|
||||
- 使用 `AllFutureHealthDataset`
|
||||
@@ -243,8 +239,7 @@ all-future 训练只读出 `t_query` 对应的 query hidden。展开的 extra-in
|
||||
|
||||
| 训练模式 | 时间模式 | 分布/监督 | 默认 loss/readout |
|
||||
| --- | --- | --- | --- |
|
||||
| `next_token` | `relative`, `absolute` | `target_mode=delphi2m`, `dist_mode=exponential` | `Delphi2MLoss` + `token` |
|
||||
| `next_token` | `relative`, `absolute` | `target_mode=uts`, `dist_mode=exponential` | `UniqueTimeSetExponentialLoss` + `same_time_group_end` |
|
||||
| `next_token` | `absolute` | `target_mode=delphi2m`, `dist_mode=exponential` | `Delphi2MLoss` + `token` |
|
||||
| `all_future` | `relative`, `absolute` | `dist_mode=exponential` | `ExponentialLoss`,无 readout |
|
||||
| `all_future` | `relative`, `absolute` | `dist_mode=weibull` | `WeibullLoss`,无 readout |
|
||||
| `all_future` | `relative`, `absolute` | `dist_mode=mixed` | `MixedLoss`,无 readout |
|
||||
@@ -255,11 +250,9 @@ all-future 训练只读出 `t_query` 对应的 query hidden。展开的 extra-in
|
||||
python train_next_step.py \
|
||||
--data_prefix ukb \
|
||||
--labels_file labels.csv \
|
||||
--target_mode uts \
|
||||
--n_embd 120 \
|
||||
--n_head 10 \
|
||||
--n_hist_layer 12 \
|
||||
--n_tab_layer 4 \
|
||||
--n_layer 12 \
|
||||
--extra_pool_reduce mean
|
||||
```
|
||||
|
||||
@@ -424,11 +417,6 @@ python evaluate_auc_v2.py \
|
||||
- `losses.py`
|
||||
- next-token 和 all-future losses
|
||||
|
||||
- `readouts.py`
|
||||
- token readout
|
||||
- same-time group readout
|
||||
- last-valid readout
|
||||
|
||||
- `evaluate_auc.py`
|
||||
- next-step/token-level 疾病 AUC 评估
|
||||
- 使用 prediction offset、sex、age bracket 分层
|
||||
|
||||
Reference in New Issue
Block a user