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

@@ -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)