refactor: isolate Delphi2M next-token pipeline
This commit is contained in:
57
dataset.py
57
dataset.py
@@ -14,7 +14,7 @@ from targets import (
|
||||
DAYS_PER_YEAR,
|
||||
NO_EVENT_IDX,
|
||||
PAD_IDX,
|
||||
build_all_targets,
|
||||
build_next_token_targets,
|
||||
)
|
||||
|
||||
|
||||
@@ -86,13 +86,11 @@ class _ExpoBaseDataset(Dataset):
|
||||
data_prefix: str = "ukb",
|
||||
labels_file: str = "labels.csv",
|
||||
no_event_interval_years: float = 5.0,
|
||||
include_no_event_in_uts_target: bool = False,
|
||||
extra_info_types: Iterable[int] | None = None,
|
||||
) -> None:
|
||||
self.data_prefix = data_prefix
|
||||
self.labels_file = labels_file
|
||||
self.no_event_interval_years = float(no_event_interval_years)
|
||||
self.include_no_event_in_uts_target = bool(include_no_event_in_uts_target)
|
||||
self.requested_extra_info_types = (
|
||||
None
|
||||
if extra_info_types is None
|
||||
@@ -138,11 +136,6 @@ class _ExpoBaseDataset(Dataset):
|
||||
max_id_in_data += 1
|
||||
self.vocab_size = max(max_id_in_vocab, max_id_in_data) + 1
|
||||
|
||||
if not self.include_no_event_in_uts_target:
|
||||
self.ignored_uts_target_ids = {PAD_IDX, CHECKUP_IDX, NO_EVENT_IDX}
|
||||
else:
|
||||
self.ignored_uts_target_ids = {PAD_IDX, CHECKUP_IDX}
|
||||
|
||||
def _prepare_sex(self, basic_table: pd.DataFrame, unique_eids: np.ndarray) -> None:
|
||||
sex_values = pd.to_numeric(basic_table["sex"], errors="coerce").to_numpy()
|
||||
if np.isnan(sex_values).any():
|
||||
@@ -289,12 +282,7 @@ class _ExpoBaseDataset(Dataset):
|
||||
|
||||
class NextStepHealthDataset(_ExpoBaseDataset):
|
||||
"""
|
||||
Dataset for next-token and next-time-point losses with unified other-info
|
||||
tokens.
|
||||
|
||||
Returned targets cover both:
|
||||
- Delphi2MLoss: target_event_seq, target_time_seq
|
||||
- UniqueTimeSetExponentialLoss: readout_mask, target_dt_unique, target_multi_hot
|
||||
Delphi2M next-token dataset with unified other-info tokens.
|
||||
"""
|
||||
|
||||
CACHE_VERSION = 3
|
||||
@@ -304,14 +292,12 @@ class NextStepHealthDataset(_ExpoBaseDataset):
|
||||
data_prefix: str = "ukb",
|
||||
labels_file: str = "labels.csv",
|
||||
no_event_interval_years: float = 5.0,
|
||||
include_no_event_in_uts_target: bool = False,
|
||||
extra_info_types: Iterable[int] | None = None,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
data_prefix=data_prefix,
|
||||
labels_file=labels_file,
|
||||
no_event_interval_years=no_event_interval_years,
|
||||
include_no_event_in_uts_target=include_no_event_in_uts_target,
|
||||
extra_info_types=extra_info_types,
|
||||
)
|
||||
|
||||
@@ -326,23 +312,18 @@ class NextStepHealthDataset(_ExpoBaseDataset):
|
||||
if features is None:
|
||||
continue
|
||||
|
||||
target_pack = build_all_targets(
|
||||
targets = build_next_token_targets(
|
||||
labels=labels,
|
||||
times_days=times_days,
|
||||
vocab_size=self.vocab_size,
|
||||
ignored_uts_target_ids=self.ignored_uts_target_ids,
|
||||
require_sorted=True,
|
||||
)
|
||||
|
||||
self.samples.append({
|
||||
"eid": eid,
|
||||
"event_seq": target_pack.next_token.input_events,
|
||||
"time_seq": target_pack.next_token.input_times_years,
|
||||
"target_event_seq": target_pack.next_token.target_events,
|
||||
"target_time_seq": target_pack.next_token.target_times_years,
|
||||
"readout_mask": target_pack.unique_time_set.readout_mask,
|
||||
"target_dt_unique": target_pack.unique_time_set.target_dt_unique,
|
||||
"target_multi_hot": target_pack.unique_time_set.target_multi_hot,
|
||||
"event_seq": targets.input_events,
|
||||
"time_seq": targets.input_times_years,
|
||||
"target_event_seq": targets.target_events,
|
||||
"target_time_seq": targets.target_times_years,
|
||||
**features,
|
||||
})
|
||||
|
||||
@@ -361,9 +342,6 @@ class NextStepHealthDataset(_ExpoBaseDataset):
|
||||
"other_time": torch.from_numpy(s["other_time"]).float(),
|
||||
"target_event_seq": torch.from_numpy(s["target_event_seq"]).long(),
|
||||
"target_time_seq": torch.from_numpy(s["target_time_seq"]).float(),
|
||||
"readout_mask": torch.from_numpy(s["readout_mask"]).bool(),
|
||||
"target_dt_unique": torch.from_numpy(s["target_dt_unique"]).float(),
|
||||
"target_multi_hot": torch.from_numpy(s["target_multi_hot"]).bool(),
|
||||
}
|
||||
|
||||
|
||||
@@ -386,7 +364,6 @@ class AllFutureHealthDataset(_ExpoBaseDataset):
|
||||
labels_file: str = "labels.csv",
|
||||
split: Literal["train", "valid", "test"] = "train",
|
||||
no_event_interval_years: float = 5.0,
|
||||
include_no_event_in_uts_target: bool = False,
|
||||
min_history_events: int = 1,
|
||||
min_future_events: int = 1,
|
||||
validation_query_seed: int = 42,
|
||||
@@ -399,7 +376,6 @@ class AllFutureHealthDataset(_ExpoBaseDataset):
|
||||
data_prefix=data_prefix,
|
||||
labels_file=labels_file,
|
||||
no_event_interval_years=no_event_interval_years,
|
||||
include_no_event_in_uts_target=include_no_event_in_uts_target,
|
||||
extra_info_types=extra_info_types,
|
||||
)
|
||||
|
||||
@@ -605,31 +581,12 @@ def next_step_collate_fn(batch: List[Dict]) -> Dict:
|
||||
batch_first=True,
|
||||
padding_value=0.0,
|
||||
)
|
||||
readout_mask = pad_sequence(
|
||||
[s["readout_mask"] for s in batch],
|
||||
batch_first=True,
|
||||
padding_value=False,
|
||||
)
|
||||
target_dt_unique = pad_sequence(
|
||||
[s["target_dt_unique"] for s in batch],
|
||||
batch_first=True,
|
||||
padding_value=0.0,
|
||||
)
|
||||
target_multi_hot = pad_sequence(
|
||||
[s["target_multi_hot"] for s in batch],
|
||||
batch_first=True,
|
||||
padding_value=False,
|
||||
)
|
||||
|
||||
out = {
|
||||
"event_seq": event_seq,
|
||||
"time_seq": time_seq,
|
||||
"padding_mask": event_seq > PAD_IDX,
|
||||
"target_event_seq": target_event_seq,
|
||||
"target_time_seq": target_time_seq,
|
||||
"readout_mask": readout_mask,
|
||||
"target_dt_unique": target_dt_unique,
|
||||
"target_multi_hot": target_multi_hot,
|
||||
}
|
||||
out.update(_collate_common_static(batch))
|
||||
return out
|
||||
|
||||
Reference in New Issue
Block a user