Remove legacy event and mixed distribution paths

This commit is contained in:
2026-08-01 14:23:18 +08:00
parent dfb22adf2d
commit de6f9b75b9
22 changed files with 370 additions and 463 deletions

View File

@@ -10,10 +10,10 @@ from torch.nn.utils.rnn import pad_sequence
from torch.utils.data import Dataset
from targets import (
CHECKUP_IDX,
DAYS_PER_YEAR,
NO_EVENT_IDX,
PAD_IDX,
RESERVED_IDX,
build_next_token_targets,
)
@@ -188,12 +188,12 @@ def load_label_vocab(
) -> Tuple[Dict[str, int], Dict[int, str]]:
label_id_to_code: Dict[int, str] = {
PAD_IDX: "<PAD>",
CHECKUP_IDX: "<CHECKUP>",
RESERVED_IDX: "<RESERVED>",
}
if include_no_event:
label_id_to_code[NO_EVENT_IDX] = "<NO_EVENT>"
offset = NO_EVENT_IDX + 1 if include_no_event else CHECKUP_IDX + 1
offset = NO_EVENT_IDX + 1 if include_no_event else RESERVED_IDX + 1
label_code_to_id: Dict[str, int] = {}
with open(labels_file, encoding="utf-8") as f:
for i, line in enumerate(f):
@@ -416,14 +416,12 @@ class _ExpoBaseDataset(Dataset):
times_days_raw = rows[:, 1].astype(np.float32)
labels_raw = rows[:, 2].astype(np.int64)
# CHECKUP is the assessment landmark for selected extra-info tokens.
# An explicitly empty selection represents a disease-only history,
# so retaining CHECKUP in that case would introduce an empty
# landmark token that is not part of the disease sequence.
if not self.extra_info_types:
keep = labels_raw != CHECKUP_IDX
times_days_raw = times_days_raw[keep]
labels_raw = labels_raw[keep]
# Label 1 was emitted as a CHECKUP event by older prepared files.
# It is now an unused reserved slot and must never enter either the
# next-token or all-future disease sequence.
keep = labels_raw != RESERVED_IDX
times_days_raw = times_days_raw[keep]
labels_raw = labels_raw[keep]
if len(labels_raw) == 0:
yield eid, times_days_raw, labels_raw
@@ -615,7 +613,7 @@ class AllFutureHealthDataset(_ExpoBaseDataset):
labels = patient["labels"]
real_event_mask = ~np.isin(
labels,
np.array([PAD_IDX, CHECKUP_IDX, NO_EVENT_IDX], dtype=np.int64),
np.array([PAD_IDX, RESERVED_IDX, NO_EVENT_IDX], dtype=np.int64),
)
n_hist = int((times <= t_query).sum())
n_future = int(((times > t_query) & real_event_mask).sum())
@@ -634,7 +632,7 @@ class AllFutureHealthDataset(_ExpoBaseDataset):
labels = np.asarray(patient["labels"], dtype=np.int64)
real_event_mask = ~np.isin(
labels,
np.array([PAD_IDX, CHECKUP_IDX, NO_EVENT_IDX], dtype=np.int64),
np.array([PAD_IDX, RESERVED_IDX, NO_EVENT_IDX], dtype=np.int64),
)
real_times = np.sort(times[real_event_mask].astype(np.float32, copy=False))
n_real_events = int(real_times.size)