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

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