Remove legacy event and mixed distribution paths
This commit is contained in:
24
dataset.py
24
dataset.py
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user