Add disease history ablation modes

This commit is contained in:
2026-07-29 13:53:15 +08:00
parent 31f129a7dc
commit c622ec50f7
10 changed files with 1178 additions and 21 deletions

View File

@@ -25,7 +25,12 @@ from torch.optim import AdamW
from torch.utils.data import DataLoader, RandomSampler
from tqdm.auto import tqdm
from dataset import AllFutureHealthDataset, all_future_collate_fn
from dataset import (
DISEASE_HISTORY_MODES,
DISEASE_HISTORY_MODE_TIMED,
AllFutureHealthDataset,
all_future_collate_fn,
)
from losses import build_loss
from model_architectures import (
DEFAULT_MODEL_ARCHITECTURE,
@@ -74,6 +79,16 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--runs_root", type=str, default="runs")
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--extra_info_types_file", type=str, default=None)
parser.add_argument(
"--disease_history_mode",
type=str,
default=DISEASE_HISTORY_MODE_TIMED,
choices=DISEASE_HISTORY_MODES,
help=(
"timed=real disease times; ordered=chronological disease order with "
"ordinal positions; set=unordered disease set with no disease time"
),
)
parser.add_argument("--train_ratio", type=float, default=0.7)
parser.add_argument("--val_ratio", type=float, default=0.15)
@@ -134,6 +149,27 @@ def parse_args() -> argparse.Namespace:
if args.extra_info_types_file is not None
else None
)
if args.disease_history_mode != DISEASE_HISTORY_MODE_TIMED:
expected = {
"model_architecture": "traj_mixer_v5",
"time_mode": "relative",
"dist_mode": "weibull",
}
mismatches = [
f"{name}={getattr(args, name)!r} (expected {value!r})"
for name, value in expected.items()
if getattr(args, name) != value
]
if args.extra_info_types != []:
mismatches.append(
"extra_info_types must be [] via extra_info_types_none.txt"
)
if mismatches:
raise ValueError(
f"disease_history_mode={args.disease_history_mode!r} is reserved "
"for the no-extra TrajMixer + all_future + relative + Weibull "
"ablation; " + "; ".join(mismatches)
)
return args
@@ -296,6 +332,7 @@ def build_metadata(
"model_target_mode": "all_future",
"target_mode": "all_future",
"dist_mode": args.dist_mode,
"disease_history_mode": args.disease_history_mode,
"all_future_min_history_events": int(args.min_history_events),
"all_future_min_future_events": int(args.min_future_events),
"all_future_validation_query_seed": int(args.validation_query_seed),
@@ -330,7 +367,15 @@ 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: (
(
""
if args.disease_history_mode == DISEASE_HISTORY_MODE_TIMED
else f"{args.disease_history_mode}_"
)
+ f"{args.time_mode}_{args.dist_mode}_"
f"all_future_pure_disease_{timestamp}"
),
runs_root=Path(args.runs_root) / args.model_architecture,
)
logger = setup_logging(run_dir)
@@ -338,6 +383,7 @@ def main() -> None:
logger.info(f"Starting all-future training run: {run_name}")
logger.info(f"Device: {device}")
logger.info(f"Model architecture: {args.model_architecture}")
logger.info(f"Disease history mode: {args.disease_history_mode}")
logger.info(f"extra_info_types: {format_extra_info_types(args.extra_info_types)}")
logger.info("Loading all-future datasets...")
@@ -349,6 +395,7 @@ def main() -> None:
min_future_events=args.min_future_events,
validation_query_seed=args.validation_query_seed,
extra_info_types=args.extra_info_types,
disease_history_mode=args.disease_history_mode,
)
val_dataset = AllFutureHealthDataset(
data_prefix=args.data_prefix,
@@ -358,6 +405,7 @@ def main() -> None:
min_future_events=args.min_future_events,
validation_query_seed=args.validation_query_seed,
extra_info_types=args.extra_info_types,
disease_history_mode=args.disease_history_mode,
)
test_dataset = AllFutureHealthDataset(
data_prefix=args.data_prefix,
@@ -367,6 +415,7 @@ def main() -> None:
min_future_events=args.min_future_events,
validation_query_seed=args.validation_query_seed,
extra_info_types=args.extra_info_types,
disease_history_mode=args.disease_history_mode,
)
if args.train_eid_file and args.val_eid_file and args.test_eid_file:
train_subset, val_subset, test_subset = split_all_future_datasets_by_eid_files(