Implement shared event-trajectory reasoning backbone
This commit is contained in:
@@ -42,8 +42,8 @@ from dataset import HealthDataset
|
||||
from eval_data import load_sequence_eval_dataset, sequence_eval_collate_fn
|
||||
from models import (
|
||||
DeepHealth,
|
||||
validate_traj_mixer_config,
|
||||
validate_traj_mixer_state_dict,
|
||||
validate_event_trajectory_config,
|
||||
validate_event_trajectory_state_dict,
|
||||
)
|
||||
from readouts import build_readout
|
||||
from targets import PAD_IDX, CHECKUP_IDX, NO_EVENT_IDX
|
||||
@@ -313,7 +313,7 @@ def split_indices(n: int, train_ratio: float, val_ratio: float, test_ratio: floa
|
||||
|
||||
|
||||
def build_model_from_dataset(args: argparse.Namespace, cfg: Dict[str, Any], dataset: HealthDataset) -> DeepHealth:
|
||||
validate_traj_mixer_config(cfg)
|
||||
validate_event_trajectory_config(cfg)
|
||||
model_target_mode = str(cfg_get(
|
||||
args, cfg, "model_target_mode", "next_token")).lower()
|
||||
if model_target_mode not in {"next_token", "all_future"}:
|
||||
@@ -322,10 +322,10 @@ def build_model_from_dataset(args: argparse.Namespace, cfg: Dict[str, Any], data
|
||||
)
|
||||
return DeepHealth(
|
||||
vocab_size=dataset.vocab_size,
|
||||
n_embd=int(cfg_get(args, cfg, "n_embd", 120)),
|
||||
n_head=int(cfg_get(args, cfg, "n_head", 10)),
|
||||
n_hist_layer=int(cfg_get(args, cfg, "n_hist_layer", 12)),
|
||||
n_tab_layer=int(cfg_get(args, cfg, "n_tab_layer", 4)),
|
||||
model_size=str(cfg_get(args, cfg, "model_size", "nano")),
|
||||
n_reasoning_rounds=int(
|
||||
cfg_get(args, cfg, "n_reasoning_rounds", 12)
|
||||
),
|
||||
n_types=dataset.n_types,
|
||||
n_cont_types=dataset.n_cont_types,
|
||||
n_categories=dataset.n_categories,
|
||||
@@ -391,7 +391,12 @@ def load_model_state(
|
||||
state = state_dict if state_dict is not None else load_checkpoint_state_dict(
|
||||
checkpoint_path, map_location=device)
|
||||
|
||||
validate_traj_mixer_state_dict(state)
|
||||
validate_event_trajectory_state_dict(
|
||||
state,
|
||||
expected_d_model=model.d_model,
|
||||
expected_n_trajectory=model.n_trajectory,
|
||||
expected_n_reasoning_rounds=model.n_reasoning_rounds,
|
||||
)
|
||||
model.load_state_dict(state, strict=True)
|
||||
|
||||
|
||||
@@ -528,7 +533,7 @@ def infer_readout_hidden(
|
||||
hidden = torch.zeros(
|
||||
batch_size,
|
||||
seq_len,
|
||||
model.n_embd,
|
||||
model.d_model,
|
||||
device=event_seq.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user