Implement shared event-trajectory reasoning backbone
This commit is contained in:
@@ -31,8 +31,8 @@ from dataset import HealthDataset
|
||||
from eval_data import load_sequence_eval_dataset
|
||||
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 CHECKUP_IDX, NO_EVENT_IDX, PAD_IDX
|
||||
@@ -182,7 +182,7 @@ def resolve_dist_mode_for_checkpoint(cfg_dist_mode: str, state_dict: Dict[str, A
|
||||
|
||||
|
||||
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"}:
|
||||
@@ -191,10 +191,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,
|
||||
@@ -209,7 +209,12 @@ def build_model_from_dataset(args: argparse.Namespace, cfg: Dict[str, Any], data
|
||||
|
||||
|
||||
def load_model_state(model: torch.nn.Module, state_dict: Dict[str, Any]) -> None:
|
||||
validate_traj_mixer_state_dict(state_dict)
|
||||
validate_event_trajectory_state_dict(
|
||||
state_dict,
|
||||
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_dict, strict=True)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user