Implement shared event-trajectory reasoning backbone

This commit is contained in:
2026-07-23 14:23:34 +08:00
parent 68a6a3df88
commit 06f29c0f0a
11 changed files with 1405 additions and 766 deletions

View File

@@ -24,7 +24,13 @@ from tqdm.auto import tqdm
from dataset import HealthDataset, collate_fn
from losses import build_loss
from models import TRAJ_MIXER_ARCHITECTURE, DeepHealth, DeepHealthOutput
from models import (
EVENT_TRAJECTORY_ARCHITECTURE,
MODEL_SIZE_NAMES,
DeepHealth,
DeepHealthOutput,
resolve_model_size,
)
from readouts import build_readout
from targets import CHECKUP_IDX, NO_EVENT_IDX, PAD_IDX
from train_util import (
@@ -74,10 +80,13 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--val_eid_file", type=str, default="ukb_val_eid.csv")
parser.add_argument("--test_eid_file", type=str, default="ukb_test_eid.csv")
parser.add_argument("--n_embd", type=int, default=120)
parser.add_argument("--n_head", type=int, default=10)
parser.add_argument("--n_hist_layer", type=int, default=12)
parser.add_argument("--n_tab_layer", type=int, default=4)
parser.add_argument(
"--model_size",
type=str,
default="nano",
choices=MODEL_SIZE_NAMES,
)
parser.add_argument("--n_reasoning_rounds", type=int, default=12)
parser.add_argument("--n_bins", type=int, default=16)
parser.add_argument("--extra_pool_reduce", type=str, default="mean",
choices=["mean", "sum"])
@@ -151,10 +160,8 @@ def move_batch_to_device(batch: Dict[str, torch.Tensor], device: torch.device) -
def build_model(args: argparse.Namespace, dataset: HealthDataset) -> DeepHealth:
return DeepHealth(
vocab_size=dataset.vocab_size,
n_embd=args.n_embd,
n_head=args.n_head,
n_hist_layer=args.n_hist_layer,
n_tab_layer=args.n_tab_layer,
model_size=args.model_size,
n_reasoning_rounds=args.n_reasoning_rounds,
n_types=dataset.n_types,
n_cont_types=dataset.n_cont_types,
n_categories=dataset.n_categories,
@@ -480,12 +487,17 @@ def build_metadata(
val_subset,
test_subset,
) -> Dict[str, Any]:
size_config = resolve_model_size(args.model_size)
return {
"run_name": run_name,
"dataset_class": "NextStepHealthDataset",
"collate_fn": "next_step_collate_fn",
"model_class": "DeepHealth",
"model_architecture": TRAJ_MIXER_ARCHITECTURE,
"model_architecture": EVENT_TRAJECTORY_ARCHITECTURE,
"d_model": size_config.d_model,
"n_trajectory": size_config.n_trajectory,
"trajectory_dim": size_config.trajectory_dim,
"traj_hidden": size_config.traj_hidden,
"model_target_mode": "next_token",
"target_mode": args.target_mode,
"dist_mode": "exponential",
@@ -521,7 +533,9 @@ def main() -> None:
run_dir, run_name = create_unique_run_dir(
lambda timestamp: (
f"{args.time_mode}_exponential_next_token_{args.target_mode}_"
f"{args.model_size}_r{args.n_reasoning_rounds}_"
f"{args.time_mode}_exponential_"
f"next_token_{args.target_mode}_"
f"gap_{args.no_event_interval_years:g}y_{timestamp}"
)
)
@@ -529,6 +543,16 @@ def main() -> None:
logger.info(f"Starting next-step training run: {run_name}")
logger.info(f"Device: {device}")
size_config = resolve_model_size(args.model_size)
logger.info(
"Model size: "
f"{args.model_size} "
f"(d_model={size_config.d_model}, "
f"n_trajectory={size_config.n_trajectory}, "
f"trajectory_dim={size_config.trajectory_dim}, "
f"traj_hidden={size_config.traj_hidden}); "
f"reasoning_rounds={args.n_reasoning_rounds}"
)
logger.info(f"extra_info_types: {format_extra_info_types(args.extra_info_types)}")
logger.info(f"readout={args.readout_name}, target_mode={args.target_mode}")