Revert "Implement shared event-trajectory reasoning backbone"

This reverts commit 06f29c0f0a.
This commit is contained in:
2026-07-23 16:05:00 +08:00
parent 22faee7c51
commit 85352dae0f
11 changed files with 753 additions and 1392 deletions

View File

@@ -31,8 +31,8 @@ from dataset import HealthDataset
from eval_data import load_sequence_eval_dataset
from models import (
DeepHealth,
validate_event_trajectory_config,
validate_event_trajectory_state_dict,
validate_traj_mixer_config,
validate_traj_mixer_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_event_trajectory_config(cfg)
validate_traj_mixer_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,
model_size=str(cfg_get(args, cfg, "model_size", "nano")),
n_reasoning_rounds=int(
cfg_get(args, cfg, "n_reasoning_rounds", 12)
),
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)),
n_types=dataset.n_types,
n_cont_types=dataset.n_cont_types,
n_categories=dataset.n_categories,
@@ -209,12 +209,7 @@ 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_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,
)
validate_traj_mixer_state_dict(state_dict)
model.load_state_dict(state_dict, strict=True)