Revert "Implement shared event-trajectory reasoning backbone"
This reverts commit 06f29c0f0a.
This commit is contained in:
@@ -27,12 +27,7 @@ from tqdm.auto import tqdm
|
||||
|
||||
from dataset import AllFutureHealthDataset, all_future_collate_fn
|
||||
from losses import build_loss
|
||||
from models import (
|
||||
EVENT_TRAJECTORY_ARCHITECTURE,
|
||||
MODEL_SIZE_NAMES,
|
||||
DeepHealth,
|
||||
resolve_model_size,
|
||||
)
|
||||
from models import TRAJ_MIXER_ARCHITECTURE, DeepHealth
|
||||
from targets import CHECKUP_IDX, PAD_IDX
|
||||
from train_util import (
|
||||
configure_torch_for_training,
|
||||
@@ -83,13 +78,10 @@ def parse_args() -> argparse.Namespace:
|
||||
parser.add_argument("--min_future_events", type=int, default=1)
|
||||
parser.add_argument("--validation_query_seed", type=int, default=None)
|
||||
|
||||
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_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("--n_bins", type=int, default=16)
|
||||
parser.add_argument("--extra_pool_reduce", type=str, default="mean",
|
||||
choices=["mean", "sum"])
|
||||
@@ -154,8 +146,10 @@ def move_batch_to_device(batch: Dict[str, torch.Tensor], device: torch.device) -
|
||||
def build_model(args: argparse.Namespace, dataset: AllFutureHealthDataset) -> DeepHealth:
|
||||
return DeepHealth(
|
||||
vocab_size=dataset.vocab_size,
|
||||
model_size=args.model_size,
|
||||
n_reasoning_rounds=args.n_reasoning_rounds,
|
||||
n_embd=args.n_embd,
|
||||
n_head=args.n_head,
|
||||
n_hist_layer=args.n_hist_layer,
|
||||
n_tab_layer=args.n_tab_layer,
|
||||
n_types=dataset.n_types,
|
||||
n_cont_types=dataset.n_cont_types,
|
||||
n_categories=dataset.n_categories,
|
||||
@@ -300,17 +294,12 @@ 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": "AllFutureHealthDataset",
|
||||
"collate_fn": "all_future_collate_fn",
|
||||
"model_class": "DeepHealth",
|
||||
"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_architecture": TRAJ_MIXER_ARCHITECTURE,
|
||||
"model_target_mode": "all_future",
|
||||
"target_mode": "all_future",
|
||||
"dist_mode": args.dist_mode,
|
||||
@@ -348,26 +337,12 @@ def main() -> None:
|
||||
configure_torch_for_training(device)
|
||||
|
||||
run_dir, run_name = create_unique_run_dir(
|
||||
lambda timestamp: (
|
||||
f"{args.model_size}_r{args.n_reasoning_rounds}_"
|
||||
f"{args.time_mode}_{args.dist_mode}_"
|
||||
f"all_future_pure_disease_{timestamp}"
|
||||
)
|
||||
lambda timestamp: f"{args.time_mode}_{args.dist_mode}_all_future_pure_disease_{timestamp}"
|
||||
)
|
||||
logger = setup_logging(run_dir)
|
||||
|
||||
logger.info(f"Starting all-future 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("Loading all-future datasets...")
|
||||
|
||||
Reference in New Issue
Block a user