from __future__ import annotations import argparse import json import math import os import sys import time from pathlib import Path from typing import Any import numpy as np import pandas as pd import torch SCRIPT_DIR = Path(__file__).resolve().parent V3P5_SCRIPT_DIR = SCRIPT_DIR.parents[1] / "v3p5_static" / "scripts" if str(SCRIPT_DIR) not in sys.path: sys.path.insert(0, str(SCRIPT_DIR)) if str(V3P5_SCRIPT_DIR) not in sys.path: sys.path.insert(0, str(V3P5_SCRIPT_DIR)) from loss_v4p4 import sctm_v4p4_loss, sctm_v4p4_next_visit_metrics from model_v4p4 import SCTMv4p4, SCTMv4p4Config from smoke_v3p4_architecture import build_field_value_mask, load_json, load_service_prior from config_v4p4_train import ( CURRICULUM_PHASES, CURRICULUM_PHASE_TO_LEGACY_STAGE, LEGACY_STAGE_TO_CURRICULUM_PHASE, LOSS_WEIGHTS_V4P4_WORLD_MODEL_20260523, ROLLOUT_SCHEDULE_DEFAULTS_V4P4_WORLD_MODEL_20260523, ROLLOUT_WEIGHTS_V4P4_WORLD_MODEL_20260523, curriculum_phase_summary as config_curriculum_phase_summary, ) DEFAULT_TENSOR_DIR = Path("/Users/yang/Downloads/精卫_SCTM_数据交付包/60_pretraining/sctm_v3p1/tensor_shards_v3p5_static_collapsed_seq128_stride64") DEFAULT_STAGE0_DIR = Path("/Users/yang/Downloads/精卫_SCTM_数据交付包/60_pretraining/sctm_v3p1/stage0") DEFAULT_OUT_DIR = Path("/Users/yang/Downloads/精卫_SCTM_数据交付包/60_pretraining/sctm_v3p1/v4p4_world_model_train") TENSOR_KEYS = [ "cat_value_ids", "missing_ids", "numeric_values", "numeric_mask", "ordinal_cbe", "ordinal_mask", "drug_name_ids", "drug_class_ids", "drug_mask", "service_state", "terminal_label", "event_labels", "delta_t_next_log", "time_since_start_days", "visit_year", "visit_indices", "valid_mask", "static_value_ids", ] INTEGER_TENSOR_KEYS = { "cat_value_ids", "missing_ids", "drug_name_ids", "drug_class_ids", "service_state", "terminal_label", "visit_year", "visit_indices", "static_value_ids", } FLOAT_TENSOR_KEYS = {"numeric_values", "ordinal_cbe", "event_labels", "delta_t_next_log", "time_since_start_days"} BOOL_TENSOR_KEYS = {"numeric_mask", "ordinal_mask", "drug_mask", "valid_mask"} def resolve_split_path(meta: dict[str, Any], tensor_dir: Path, split: str) -> Path: seq_len = int(meta["max_seq_len"]) local = tensor_dir / f"v3p1_{split}_seq{seq_len}_tensor_shard.npz" if local.exists(): return local original = Path(meta["split_paths"][split]) if original.exists(): return original raise FileNotFoundError(f"Could not resolve {split} tensor shard in {tensor_dir} or {original}") def load_npz_to_memory(path: Path) -> dict[str, np.ndarray]: with np.load(path, allow_pickle=True) as z: keep = set(TENSOR_KEYS) | {"patient_ids", "window_starts"} return {k: z[k] for k in z.files if k in keep} def inspect_npz_schema(path: Path) -> dict[str, Any]: with np.load(path, allow_pickle=True) as z: return {k: {"shape": list(z[k].shape), "dtype": str(z[k].dtype)} for k in z.files if k in set(TENSOR_KEYS)} def tensor_alignment_audit(meta: dict[str, Any], paths: dict[str, Path]) -> dict[str, Any]: expected = { "cat_value_ids": [int(meta["max_seq_len"]), len(meta["cat_cols"])], "missing_ids": [int(meta["max_seq_len"]), len(meta["cat_cols"])], "numeric_values": [int(meta["max_seq_len"]), len(meta["num_cols"])], "numeric_mask": [int(meta["max_seq_len"]), len(meta["num_cols"])], "ordinal_cbe": [int(meta["max_seq_len"]), len(meta["ord_cols"]), int(meta["cbe_dim"])], "ordinal_mask": [int(meta["max_seq_len"]), len(meta["ord_cols"])], "drug_name_ids": [int(meta["max_seq_len"]), int(meta["max_drug_atoms"])], "drug_class_ids": [int(meta["max_seq_len"]), int(meta["max_drug_atoms"])], "drug_mask": [int(meta["max_seq_len"]), int(meta["max_drug_atoms"])], "service_state": [int(meta["max_seq_len"])], "terminal_label": [int(meta["max_seq_len"])], "event_labels": [int(meta["max_seq_len"]), len(meta["event_cols"])], "delta_t_next_log": [int(meta["max_seq_len"])], "time_since_start_days": [int(meta["max_seq_len"])], "visit_year": [int(meta["max_seq_len"])], "visit_indices": [int(meta["max_seq_len"])], "valid_mask": [int(meta["max_seq_len"])], "static_value_ids": [len(meta.get("static_cols", []))], } split_audits: dict[str, Any] = {} failures: list[str] = [] for split, path in paths.items(): schema = inspect_npz_schema(path) missing = [key for key in TENSOR_KEYS if key not in schema] shape_errors: dict[str, Any] = {} for key, tail in expected.items(): if key not in schema: continue actual = schema[key]["shape"] if actual[1:] != tail: shape_errors[key] = {"actual": actual, "expected_tail": ["n_windows", *tail]} if missing or shape_errors: failures.append(split) split_audits[split] = { "path": str(path), "n_windows": int(schema["valid_mask"]["shape"][0]) if "valid_mask" in schema else None, "missing_required_keys": missing, "shape_errors": shape_errors, } return { "status": "passed" if not failures else "failed", "failed_splits": failures, "meta_dims": { "max_seq_len": int(meta["max_seq_len"]), "n_cat_fields": len(meta["cat_cols"]), "n_numeric_fields": len(meta["num_cols"]), "n_ordinal_fields": len(meta["ord_cols"]), "n_events": len(meta["event_cols"]), "n_static_fields": len(meta.get("static_cols", [])), "cat_vocab_size": int(meta["cat_vocab_size"]), "drug_vocab_size": int(meta["drug_vocab_size"]), "static_vocab_size": int(meta.get("static_vocab_size", 0)), }, "splits": split_audits, } def train_value_alignment_audit(arrays: dict[str, np.ndarray], meta: dict[str, Any], config: SCTMv4p4Config) -> dict[str, Any]: valid = arrays["valid_mask"].astype(bool) service = arrays["service_state"] terminal = arrays["terminal_label"] event_labels = arrays["event_labels"] delta_days = np.expm1(arrays["delta_t_next_log"]) valid_next = valid[:, :-1] & valid[:, 1:] time_gap = arrays["time_since_start_days"][:, 1:] - arrays["time_since_start_days"][:, :-1] dt_diff = np.abs(time_gap[valid_next] - delta_days[:, :-1][valid_next]) if valid_next.any() else np.array([0.0]) service_ok = bool((service[valid] < config.n_service_states).all()) if valid.any() else True terminal_ok = bool(((terminal[valid] >= 0) & (terminal[valid] <= 4)).all()) if valid.any() else True event_binary = bool(((event_labels == 0.0) | (event_labels == 1.0)).all()) cat_ok = bool(arrays["cat_value_ids"].max() < config.cat_vocab_size) missing_ok = bool(arrays["missing_ids"].max() < config.n_missing) static_ok = bool((not config.n_static_fields) or arrays["static_value_ids"].max() < config.static_vocab_size) delta_ok = bool(float(dt_diff.max()) <= 1.0e-2) status = "passed" if all([service_ok, terminal_ok, event_binary, cat_ok, missing_ok, static_ok, delta_ok]) else "failed" return { "status": status, "valid_slots": int(valid.sum()), "valid_next_positions": int(valid_next.sum()), "service_state_min": int(service[valid].min()) if valid.any() else None, "service_state_max": int(service[valid].max()) if valid.any() else None, "terminal_label_min": int(terminal[valid].min()) if valid.any() else None, "terminal_label_max": int(terminal[valid].max()) if valid.any() else None, "event_labels_last_dim": int(event_labels.shape[-1]), "event_cols": list(meta["event_cols"]), "event_positive_by_col": event_labels[valid].sum(axis=0).astype(int).tolist() if valid.any() else [], "cat_value_min": int(arrays["cat_value_ids"].min()), "cat_value_max": int(arrays["cat_value_ids"].max()), "cat_vocab_size": int(config.cat_vocab_size), "missing_id_min": int(arrays["missing_ids"].min()), "missing_id_max": int(arrays["missing_ids"].max()), "raw_missing_map": dict(meta["missing_map"]), "static_value_min": int(arrays["static_value_ids"].min()) if config.n_static_fields else None, "static_value_max": int(arrays["static_value_ids"].max()) if config.n_static_fields else None, "static_vocab_size": int(config.static_vocab_size), "delta_days_min_valid_next": float(delta_days[:, :-1][valid_next].min()) if valid_next.any() else None, "delta_days_max_valid_next": float(delta_days[:, :-1][valid_next].max()) if valid_next.any() else None, "delta_matches_time_gap_abs_diff_max": float(dt_diff.max()), "delta_matches_time_gap_abs_diff_mean": float(dt_diff.mean()), "service_state_within_config": service_ok, "terminal_label_within_expected_0_to_4": terminal_ok, "event_labels_binary": event_binary, "cat_value_within_vocab": cat_ok, "missing_id_within_v4_vocab": missing_ok, "static_value_within_vocab": static_ok, "delta_matches_time_gap_within_tolerance": delta_ok, } def batch_from_indices(arrays: dict[str, np.ndarray], idx: np.ndarray, device: torch.device) -> dict[str, torch.Tensor]: out: dict[str, torch.Tensor] = {} for key in TENSOR_KEYS: if key not in arrays: continue arr = arrays[key][idx] if key in INTEGER_TENSOR_KEYS: out[key] = torch.as_tensor(arr, dtype=torch.long, device=device) elif key in BOOL_TENSOR_KEYS: out[key] = torch.as_tensor(arr, dtype=torch.bool, device=device) elif key in FLOAT_TENSOR_KEYS: out[key] = torch.as_tensor(arr, dtype=torch.float32, device=device) else: raise KeyError(f"untyped tensor key in batch_from_indices: {key}") return out class TrainIndexSampler: def __init__(self, rng: np.random.Generator, n: int, sample_with_replacement: bool) -> None: self.rng = rng self.n = int(n) self.sample_with_replacement = bool(sample_with_replacement) self.perm = self.rng.permutation(self.n) if self.n > 0 and not self.sample_with_replacement else np.empty(0, dtype=np.int64) self.ptr = 0 def next(self, batch_size: int) -> np.ndarray: size = min(int(batch_size), self.n) if size <= 0: return np.empty(0, dtype=np.int64) if self.sample_with_replacement: return self.rng.integers(0, self.n, size=size) if self.ptr + size <= self.n: idx = self.perm[self.ptr : self.ptr + size] self.ptr += size return idx tail = self.perm[self.ptr :] self.perm = self.rng.permutation(self.n) take = size - int(tail.shape[0]) head = self.perm[:take] self.ptr = take return np.concatenate([tail, head]) def autocast_context(device: torch.device, precision: str): enabled = device.type == "cuda" and precision in {"bf16", "fp16"} dtype = torch.bfloat16 if precision == "bf16" else torch.float16 return torch.autocast(device_type=device.type, dtype=dtype, enabled=enabled) def cosine_lr(step: int, *, max_steps: int, warmup_steps: int, lr: float, min_lr: float) -> float: if step <= warmup_steps: return lr * step / max(1, warmup_steps) progress = (step - warmup_steps) / max(1, max_steps - warmup_steps) coeff = 0.5 * (1.0 + math.cos(math.pi * min(1.0, progress))) return min_lr + coeff * (lr - min_lr) def build_conditional_specs(meta: dict[str, Any], vocab: dict[str, int]) -> tuple[dict[str, Any], ...]: cat_cols = list(meta["cat_cols"]) def field_idx(name: str) -> int: return cat_cols.index(name) if name in cat_cols else -1 specs = [] tables = [ ("是否转诊", ["是"], ["转诊类型", "转诊原因", "转诊至机构"]), ("治疗方式", ["住院"], ["本次入院形式", "本次住院名称", "本次入院日期", "末次出院日期"]), ] for parent, active_values, children in tables: pidx = field_idx(parent) child_indices = [field_idx(c) for c in children if field_idx(c) >= 0] active_ids = [int(vocab[f"{parent}={v}"]) for v in active_values if f"{parent}={v}" in vocab] if pidx >= 0 and child_indices and active_ids: specs.append({"parent": parent, "parent_idx": pidx, "active_value_ids": active_ids, "child_indices": child_indices}) return tuple(specs) def build_color_map(meta: dict[str, Any], vocab: dict[str, int]) -> tuple[int, tuple[tuple[int, int], ...]]: cat_cols = list(meta["cat_cols"]) color_field = "新规评估信号(颜色)" if color_field not in cat_cols: return -1, tuple() order = {"绿色": 0, "绿": 0, "蓝色": 1, "蓝": 1, "黄色": 2, "黄": 2, "橙色": 3, "橙": 3, "红色": 4, "红": 4} pairs: list[tuple[int, int]] = [] prefix = f"{color_field}=" for token, value_id in vocab.items(): if not token.startswith(prefix): continue value = token.removeprefix(prefix) if value in order: pairs.append((int(value_id), int(order[value]))) return cat_cols.index(color_field), tuple(sorted(pairs)) def metadata_config(meta: dict[str, Any], vocab: dict[str, int], args: argparse.Namespace) -> SCTMv4p4Config: color_idx, color_map = build_color_map(meta, vocab) return SCTMv4p4Config( cat_vocab_size=int(meta["cat_vocab_size"]), drug_vocab_size=int(meta["drug_vocab_size"]), n_cat_fields=len(meta["cat_cols"]), n_numeric_fields=len(meta["num_cols"]), n_ordinal_fields=len(meta["ord_cols"]), cbe_dim=int(meta["cbe_dim"]), max_drug_atoms=int(meta["max_drug_atoms"]), n_events=len(meta["event_cols"]), static_vocab_size=int(meta.get("static_vocab_size", 0)), n_static_fields=len(meta.get("static_cols", [])), d_model=args.d_model, n_heads=args.n_heads, n_layers=args.n_layers, dropout=args.dropout, n_missing=args.n_missing, year_min=int(meta.get("year_min", 2010)), year_max=int(meta.get("year_max", 2026)), conditional_specs=build_conditional_specs(meta, vocab), color_cat_field_index=color_idx, color_value_to_class=color_map, ) def try_train_microbatch( model: SCTMv4p4, config: SCTMv4p4Config, arrays: dict[str, np.ndarray], batch_size: int, device: torch.device, args: argparse.Namespace, ) -> bool: idx = np.arange(min(batch_size, arrays["valid_mask"].shape[0])) try: model.train() batch = batch_from_indices(arrays, idx, device) rollout_steps, teacher_forcing_rate, gumbel_temperature = rollout_schedule(args, 1) with autocast_context(device, args.precision): out = model( batch, rollout_steps=rollout_steps, teacher_forcing_rate=teacher_forcing_rate, gumbel_temperature=gumbel_temperature, compute_pwe_diagnostics=False, ) loss, _ = sctm_v4p4_loss( out, batch, config, field_loss_weight=args.field_loss_weight, numeric_loss_weight=args.numeric_loss_weight, ordinal_loss_weight=args.ordinal_loss_weight, missingness_loss_weight=args.missingness_loss_weight, pre_event_risk_weight=args.pre_event_risk_weight, event_risk_weight=args.event_risk_weight, post_event_risk_weight=args.post_event_risk_weight, event_generation_weight=args.event_generation_weight, ontology_event_weight=args.ontology_event_weight, history_aux_weight=args.history_aux_weight, contact_aux_weight=args.contact_aux_weight, color_loss_weight=args.color_loss_weight, pwe_pre_nll_weight=args.pwe_pre_nll_weight, pwe_contact_nll_weight=args.pwe_contact_nll_weight, pwe_post_nll_weight=args.pwe_post_nll_weight, pwe_pre_distill_weight=args.pwe_pre_distill_weight, pwe_contact_distill_weight=args.pwe_contact_distill_weight, rollout_weights=(args.rollout_weight_1, args.rollout_weight_2, args.rollout_weight_3), return_metrics=False, ) loss.backward() model.zero_grad(set_to_none=True) if device.type == "cuda": torch.cuda.empty_cache() return True except RuntimeError as exc: if "out of memory" in str(exc).lower(): model.zero_grad(set_to_none=True) if device.type == "cuda": torch.cuda.empty_cache() return False raise def select_batch_size( model: SCTMv4p4, config: SCTMv4p4Config, arrays: dict[str, np.ndarray], args: argparse.Namespace, device: torch.device, ) -> int: if not args.auto_batch: return args.batch_size candidates = [] value = args.batch_size while value >= args.min_batch_size: candidates.append(value) value = int(value * 0.75) candidates.append(args.min_batch_size) for batch_size in sorted(set(candidates), reverse=True): if try_train_microbatch(model, config, arrays, batch_size, device, args): return batch_size raise RuntimeError(f"No valid batch size found down to {args.min_batch_size}") def adapt_v3p5_tensor_for_v4p4(key: str, src: torch.Tensor, dst: torch.Tensor, config: SCTMv4p4Config) -> torch.Tensor | None: """Copy v3p5 weights into v4p4 tensors whose leading semantics expanded. Generic prefix-copying is intentionally avoided for flattened heads because rows encode `(field, class)` pairs. Only known v3p5->v4p4 expansions are adapted here. """ if key == "tab_encoder.axis_slot_emb" and src.ndim == 4 and dst.ndim == 4 and src.shape[-1] == dst.shape[-1]: adapted = torch.zeros_like(dst) n_slots = min(src.shape[2], dst.shape[2]) adapted[:, :, :n_slots, :] = src[:, :, :n_slots, :].to(dtype=dst.dtype, device=dst.device) return adapted if key == "tab_encoder.missing_emb.weight" and src.ndim == 2 and dst.ndim == 2 and src.shape[1] == dst.shape[1]: adapted = torch.zeros_like(dst) n_rows = min(src.shape[0], dst.shape[0]) adapted[:n_rows, :] = src[:n_rows, :].to(dtype=dst.dtype, device=dst.device) return adapted if key == "missingness_head.weight" and src.ndim == 2 and dst.ndim == 2 and src.shape[1] == dst.shape[1]: n_fields = int(config.n_cat_fields) if src.shape[0] % n_fields == 0 and dst.shape[0] == n_fields * int(config.n_missing): src_missing = src.shape[0] // n_fields n_missing = min(src_missing, int(config.n_missing)) adapted = torch.zeros_like(dst) src_view = src.reshape(n_fields, src_missing, src.shape[1]) dst_view = adapted.reshape(n_fields, int(config.n_missing), dst.shape[1]) dst_view[:, :n_missing, :] = src_view[:, :n_missing, :].to(dtype=dst.dtype, device=dst.device) return adapted if key == "missingness_head.bias" and src.ndim == 1 and dst.ndim == 1: n_fields = int(config.n_cat_fields) if src.shape[0] % n_fields == 0 and dst.shape[0] == n_fields * int(config.n_missing): src_missing = src.shape[0] // n_fields n_missing = min(src_missing, int(config.n_missing)) adapted = torch.zeros_like(dst) src_view = src.reshape(n_fields, src_missing) dst_view = adapted.reshape(n_fields, int(config.n_missing)) dst_view[:, :n_missing] = src_view[:, :n_missing].to(dtype=dst.dtype, device=dst.device) return adapted return None def load_compatible_state_dict(model: torch.nn.Module, checkpoint: Path, device: torch.device) -> dict[str, Any]: payload = torch.load(checkpoint, map_location="cpu", weights_only=False) payload.pop("optimizer_state_dict", None) state = payload["model_state_dict"] if any(k.startswith("_orig_mod.") for k in state): state = {k.removeprefix("_orig_mod."): v for k, v in state.items()} own = model.state_dict() compatible: dict[str, torch.Tensor] = {} skipped: dict[str, list[int]] = {} adapted: dict[str, dict[str, list[int]]] = {} base_model = model._orig_mod if hasattr(model, "_orig_mod") else model config = base_model.config for key, value in state.items(): if key not in own: continue target = own[key] if tuple(value.shape) == tuple(target.shape): compatible[key] = value continue adapted_value = adapt_v3p5_tensor_for_v4p4(key, value, target, config) if adapted_value is None: skipped[key] = list(value.shape) continue compatible[key] = adapted_value adapted[key] = {"source_shape": list(value.shape), "target_shape": list(target.shape)} for suffix in ("weight", "bias"): src_key = f"ontology_event_head.{suffix}" dst_key = f"event_generation_head.{suffix}" if dst_key not in compatible and src_key in state and dst_key in own and tuple(state[src_key].shape) == tuple(own[dst_key].shape): compatible[dst_key] = state[src_key].to(dtype=own[dst_key].dtype, device=own[dst_key].device) adapted[dst_key] = {"source_key": src_key, "source_shape": list(state[src_key].shape), "target_shape": list(own[dst_key].shape)} incompatible = model.load_state_dict(compatible, strict=False) return { "checkpoint": str(checkpoint), "source_step": int(payload.get("step", -1)), "loaded_keys": len(compatible), "shape_adapted": adapted, "missing_keys": list(incompatible.missing_keys), "unexpected_keys": list(incompatible.unexpected_keys), "shape_skipped": skipped, } def unwrap_compiled_model(model: torch.nn.Module) -> torch.nn.Module: return model._orig_mod if hasattr(model, "_orig_mod") else model def strip_compile_prefix(state: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: if any(key.startswith("_orig_mod.") for key in state): return {key.removeprefix("_orig_mod."): value for key, value in state.items()} return state V4P4_ALLOWED_RESUME_MISSING_PREFIXES = ( "deployable_risk_post_residual_head.", ) def load_resume_state_compatible(model: torch.nn.Module, state: dict[str, torch.Tensor]) -> dict[str, Any]: incompatible = unwrap_compiled_model(model).load_state_dict(state, strict=False) bad_missing = [key for key in incompatible.missing_keys if not key.startswith(V4P4_ALLOWED_RESUME_MISSING_PREFIXES)] if bad_missing or incompatible.unexpected_keys: raise RuntimeError( "resume checkpoint is not compatible with current v4p4 model: " f"bad_missing={bad_missing}, unexpected={list(incompatible.unexpected_keys)}" ) return { "missing_keys": list(incompatible.missing_keys), "unexpected_keys": list(incompatible.unexpected_keys), "allowed_missing_prefixes": list(V4P4_ALLOWED_RESUME_MISSING_PREFIXES), } BACKBONE_PREFIXES = ("tab_encoder", "traj_decoder", "static_encoder", "static_injection_scale") def normalize_curriculum_args(args: argparse.Namespace) -> argparse.Namespace: """Keep legacy --stage reproducible while exposing the unified-model curriculum.""" requested = getattr(args, "curriculum_phase", "") or "" if requested: if requested not in CURRICULUM_PHASE_TO_LEGACY_STAGE: raise ValueError(f"unknown curriculum phase: {requested}") legacy = CURRICULUM_PHASE_TO_LEGACY_STAGE[requested] if getattr(args, "stage", legacy) != legacy: raise ValueError(f"--stage {args.stage!r} conflicts with --curriculum-phase {requested!r}") args.stage = legacy args.curriculum_phase = requested else: args.curriculum_phase = LEGACY_STAGE_TO_CURRICULUM_PHASE.get(args.stage, args.stage) return args def curriculum_phase_summary(args: argparse.Namespace) -> dict[str, Any]: phase = getattr(args, "curriculum_phase", LEGACY_STAGE_TO_CURRICULUM_PHASE.get(args.stage, args.stage)) return config_curriculum_phase_summary(phase, legacy_stage=args.stage) def apply_stage_trainability(model: torch.nn.Module, args: argparse.Namespace) -> dict[str, Any]: for param in model.parameters(): param.requires_grad = True frozen: list[str] = [] phase = getattr(args, "curriculum_phase", LEGACY_STAGE_TO_CURRICULUM_PHASE.get(args.stage, args.stage)) if phase == "adapter_warmup" and args.freeze_backbone_stage1: for name, param in model.named_parameters(): if name == "static_injection_scale" or name.startswith(BACKBONE_PREFIXES[:3]): param.requires_grad = False frozen.append(name) trainable_params = int(sum(p.numel() for p in model.parameters() if p.requires_grad)) frozen_params = int(sum(p.numel() for p in model.parameters() if not p.requires_grad)) if trainable_params == 0: raise RuntimeError("stage trainability configuration froze all parameters") return { "stage": args.stage, "curriculum_phase": phase, "freeze_backbone_stage1": bool(args.freeze_backbone_stage1), "trainable_params": trainable_params, "frozen_params": frozen_params, "frozen_param_names_sample": frozen[:20], "n_frozen_param_names": len(frozen), } def curriculum_schedule(args: argparse.Namespace, step: int) -> tuple[int, float, float]: phase = getattr(args, "curriculum_phase", LEGACY_STAGE_TO_CURRICULUM_PHASE.get(args.stage, args.stage)) if phase == "joint_one_step": return 1, 1.0, args.gumbel_temperature_start return 1, 1.0, args.gumbel_temperature_start def rollout_schedule(args: argparse.Namespace, step: int) -> tuple[int, float, float]: return curriculum_schedule(args, step) @torch.no_grad() def evaluate( model: SCTMv4p4, config: SCTMv4p4Config, arrays: dict[str, np.ndarray], batch_size: int, max_batches: int, device: torch.device, precision: str, loss_kwargs: dict[str, Any] | None = None, rollout_steps: int = 1, teacher_forcing_rate: float = 1.0, gumbel_temperature: float = 1.0, ) -> dict[str, float]: model.eval() weighted: dict[str, float] = {} weights: dict[str, float] = {} batches = 0 n = arrays["valid_mask"].shape[0] for start in range(0, n, batch_size): if max_batches > 0 and batches >= max_batches: break idx = np.arange(start, min(start + batch_size, n)) batch = batch_from_indices(arrays, idx, device) with autocast_context(device, precision): out = model( batch, rollout_steps=rollout_steps, teacher_forcing_rate=teacher_forcing_rate, gumbel_temperature=gumbel_temperature, compute_pwe_diagnostics=False, ) _, loss_metrics = sctm_v4p4_loss(out, batch, config, **(loss_kwargs or {}), return_metrics=True) metric = sctm_v4p4_next_visit_metrics(out, batch, config) w = float(loss_metrics.get("valid_next_positions", 1.0)) for row in (loss_metrics, {f"metric_{k}": v for k, v in metric.items()}): for key, value in row.items(): if isinstance(value, (int, float)) and np.isfinite(value): if key == "valid_next_positions": continue weighted[key] = weighted.get(key, 0.0) + float(value) * w weights[key] = weights.get(key, 0.0) + w weighted["valid_next_positions"] = weighted.get("valid_next_positions", 0.0) + w batches += 1 out = {key: weighted[key] / max(1.0e-12, weights[key]) for key in weights} out["valid_next_positions"] = weighted.get("valid_next_positions", 0.0) out["eval_batches"] = float(batches) return out def json_ready(value: Any) -> Any: if isinstance(value, Path): return str(value) if isinstance(value, dict): return {str(k): json_ready(v) for k, v in value.items()} if isinstance(value, (list, tuple)): return [json_ready(v) for v in value] if isinstance(value, np.generic): return value.item() return value def save_checkpoint( path: Path, *, model: SCTMv4p4, optimizer: torch.optim.Optimizer, step: int, config: SCTMv4p4Config, tensor_metadata: dict[str, Any], args: argparse.Namespace, train_log_tail: dict[str, Any], ) -> None: path.parent.mkdir(parents=True, exist_ok=True) torch.save( { "step": step, "model_state_dict": unwrap_compiled_model(model).state_dict(), "optimizer_state_dict": optimizer.state_dict(), "config": config.__dict__, "tensor_metadata": tensor_metadata, "args": json_ready(vars(args)), "train_log_tail": train_log_tail, }, path, ) def maybe_upload(args: argparse.Namespace, folder: Path, stage: str) -> None: if not args.hub_repo_id: return token = os.environ.get("HF_TOKEN") if not token: print(json.dumps({"stage": "hub_upload_skipped", "reason": "HF_TOKEN missing", "upload_stage": stage}), flush=True) return from huggingface_hub import HfApi api = HfApi(token=token) api.upload_folder( folder_path=str(folder), path_in_repo=args.hub_path, repo_id=args.hub_repo_id, repo_type=args.hub_repo_type, token=token, commit_message=f"Upload SCTM-v4p4 {args.run_label} at {stage}", ) def run(args: argparse.Namespace) -> dict[str, Any]: args = normalize_curriculum_args(args) args.out_dir.mkdir(parents=True, exist_ok=True) torch.manual_seed(args.seed) np.random.seed(args.seed) if args.matmul_precision: torch.set_float32_matmul_precision(args.matmul_precision) device = torch.device(args.device) meta = load_json(args.tensor_dir / "tensor_metadata.json") vocab = load_json(args.tensor_dir / "cat_value_vocab.json") paths = {split: resolve_split_path(meta, args.tensor_dir, split) for split in ["train", "val", "test"]} shape_audit = tensor_alignment_audit(meta, paths) print(json.dumps({"stage": "tensor_alignment_audit", **shape_audit}, ensure_ascii=False), flush=True) if shape_audit["status"] != "passed": raise RuntimeError(f"tensor schema does not align with v4p4 config contract: {shape_audit}") print(json.dumps({"stage": "load_npz_to_memory", "paths": {k: str(v) for k, v in paths.items()}}, ensure_ascii=False), flush=True) train_arrays = load_npz_to_memory(paths["train"]) val_arrays = load_npz_to_memory(paths["val"]) missing = [key for key in TENSOR_KEYS if key not in train_arrays] if missing: raise KeyError(f"train tensor shard missing required v4p4 keys: {missing}") field_value_mask, field_mask_audit = build_field_value_mask(meta, vocab, {k: train_arrays[k] for k in ["cat_value_ids", "missing_ids"]}) service_prior, prior_audit = load_service_prior(args.stage0_dir, int(meta.get("n_service_states", 8)), args.service_prior_file) config = metadata_config(meta, vocab, args) value_audit = train_value_alignment_audit(train_arrays, meta, config) print(json.dumps({"stage": "train_value_alignment_audit", **value_audit}, ensure_ascii=False), flush=True) if value_audit["status"] != "passed": raise RuntimeError(f"tensor values do not align with v4p4 config contract: {value_audit}") model = SCTMv4p4(config, field_value_mask=field_value_mask, service_prior_bias=service_prior).to(device) init_audit: dict[str, Any] = {} if args.init_checkpoint: init_audit = load_compatible_state_dict(model, args.init_checkpoint, device) print(json.dumps({"stage": "init_checkpoint_loaded", **init_audit}, ensure_ascii=False), flush=True) trainability_audit = apply_stage_trainability(model, args) param_count = int(sum(p.numel() for p in model.parameters())) selected_batch_size = select_batch_size(model, config, train_arrays, args, device) compile_enabled = bool(args.compile and hasattr(torch, "compile")) if compile_enabled: model = torch.compile(model) # type: ignore[assignment] print( json.dumps( { "stage": "model_ready", "param_count": param_count, "selected_batch_size": selected_batch_size, "requested_batch_size": args.batch_size, "device": str(device), "run_stage": args.stage, "cuda_device": torch.cuda.get_device_name(0) if device.type == "cuda" else None, "compile_enabled": compile_enabled, "sample_with_replacement": bool(args.sample_with_replacement), }, ensure_ascii=False, ), flush=True, ) use_fused = bool(args.fused_adamw and device.type == "cuda") optimizer = torch.optim.AdamW( [p for p in model.parameters() if p.requires_grad], lr=args.learning_rate, betas=(args.beta1, args.beta2), weight_decay=args.weight_decay, fused=use_fused, ) resume_step = 0 resume_audit: dict[str, Any] = {} if args.resume_checkpoint: payload = torch.load(args.resume_checkpoint, map_location=device, weights_only=False) state = strip_compile_prefix(payload["model_state_dict"]) resume_state_audit = load_resume_state_compatible(model, state) resume_audit = { "checkpoint": str(args.resume_checkpoint), "source_step": int(payload.get("step", 0)), "model_state_loaded": True, "model_state_audit": resume_state_audit, "optimizer_state_loaded": False, "optimizer_state_skip_reason": "", } if "optimizer_state_dict" in payload: try: optimizer.load_state_dict(payload["optimizer_state_dict"]) resume_audit["optimizer_state_loaded"] = True except ValueError as exc: resume_audit["optimizer_state_skip_reason"] = str(exc) resume_step = int(payload.get("step", 0)) print(json.dumps({"stage": "resume_checkpoint_loaded", **resume_audit}, ensure_ascii=False), flush=True) run_meta = { "run_label": args.run_label, "model_family": "sctm_v4p4_world_model", "stage": args.stage, "curriculum_phase": args.curriculum_phase, "curriculum_phase_summary": curriculum_phase_summary(args), "resolved_train_config": { "preset": "v4p4_world_model_20260523", "loss_weights": {key: getattr(args, key) for key in LOSS_WEIGHTS_V4P4_WORLD_MODEL_20260523}, "rollout_weights": { "rollout_weight_1": args.rollout_weight_1, "rollout_weight_2": args.rollout_weight_2, "rollout_weight_3": args.rollout_weight_3, }, "rollout_schedule": { "rollout_steps": args.rollout_steps, "teacher_forcing_rate": args.teacher_forcing_rate, "teacher_forcing_start": args.teacher_forcing_start, "teacher_forcing_end": args.teacher_forcing_end, "gumbel_temperature_start": args.gumbel_temperature_start, "gumbel_temperature_end": args.gumbel_temperature_end, }, }, "param_count": param_count, "config": json_ready(config.__dict__), "tensor_dir": str(args.tensor_dir), "tensor_paths": {k: str(v) for k, v in paths.items()}, "tensor_alignment_audit": shape_audit, "train_value_alignment_audit": value_audit, "field_value_mask_audit": field_mask_audit, "service_prior_audit": prior_audit, "init_audit": init_audit, "trainability_audit": trainability_audit, "selected_batch_size": selected_batch_size, "requested_batch_size": args.batch_size, "device": str(device), "precision": args.precision, "compile_enabled": compile_enabled, "sample_with_replacement": bool(args.sample_with_replacement), "finite_check_every": int(args.finite_check_every), "resume_checkpoint": str(args.resume_checkpoint) if args.resume_checkpoint else "", "resume_step": resume_step, "resume_audit": resume_audit, } (args.out_dir / "run_meta.json").write_text(json.dumps(run_meta, ensure_ascii=False, indent=2), encoding="utf-8") rng = np.random.default_rng(args.seed) n_train = train_arrays["valid_mask"].shape[0] train_sampler = TrainIndexSampler(rng, n_train, args.sample_with_replacement) start_time = time.time() log_rows: list[dict[str, Any]] = [] latest_eval: dict[str, float] = {} best_metric_value: float | None = None best_step = resume_step evals_since_best = 0 early_stopped = False stop_reason = "max_steps" loss_kwargs = { key: getattr(args, key) for key in LOSS_WEIGHTS_V4P4_WORLD_MODEL_20260523 } | { "rollout_weights": (args.rollout_weight_1, args.rollout_weight_2, args.rollout_weight_3), } for step in range(resume_step + 1, args.max_steps + 1): need_step_metrics = ( step == 1 or step % args.log_every == 0 or step % args.eval_every == 0 or step % args.save_every == 0 or step % args.hub_upload_every == 0 or step == args.max_steps ) need_finite_check = args.finite_check_every > 0 and ( step == 1 or step % args.finite_check_every == 0 or step == args.max_steps ) lr = cosine_lr(step, max_steps=args.max_steps, warmup_steps=args.warmup_steps, lr=args.learning_rate, min_lr=args.min_learning_rate) for group in optimizer.param_groups: group["lr"] = lr step_start = time.time() if device.type == "cuda" and need_step_metrics: torch.cuda.reset_peak_memory_stats(device) model.train() optimizer.zero_grad(set_to_none=True) rollout_steps, teacher_forcing_rate, gumbel_temperature = rollout_schedule(args, step) accum_losses: list[dict[str, float]] = [] for _ in range(args.grad_accum_steps): idx = train_sampler.next(selected_batch_size) batch = batch_from_indices(train_arrays, idx, device) with autocast_context(device, args.precision): out = model( batch, rollout_steps=rollout_steps, teacher_forcing_rate=teacher_forcing_rate, gumbel_temperature=gumbel_temperature, compute_pwe_diagnostics=need_step_metrics, ) loss, loss_metrics = sctm_v4p4_loss( out, batch, config, **loss_kwargs, return_metrics=need_step_metrics, ) loss = loss / args.grad_accum_steps if need_finite_check and not torch.isfinite(loss): raise RuntimeError(f"non-finite loss at step {step}: {float(loss.detach().cpu())}") loss.backward() if need_step_metrics: accum_losses.append(loss_metrics) grad_norm_tensor = torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip) optimizer.step() if need_step_metrics and device.type == "cuda": torch.cuda.synchronize(device) step_sec = max(1.0e-12, time.time() - step_start) effective_windows = int(selected_batch_size * args.grad_accum_steps) seq_len = int(train_arrays["valid_mask"].shape[1]) base_model = model._orig_mod if hasattr(model, "_orig_mod") else model row: dict[str, Any] = { "step": step, "curriculum_phase": args.curriculum_phase, "legacy_stage": args.stage, "lr": lr, "grad_norm": float(grad_norm_tensor.detach().cpu()) if need_step_metrics else float("nan"), "elapsed_sec": round(time.time() - start_time, 3), "param_count": param_count, "step_sec": round(step_sec, 4), "selected_batch_size": selected_batch_size, "effective_windows": effective_windows, "visit_slots_per_sec": float(effective_windows * seq_len / step_sec), "rollout_steps": rollout_steps, "teacher_forcing_rate": teacher_forcing_rate, "gumbel_temperature": gumbel_temperature, "contact_injection_scale": ( float(base_model.contact_injector.scale.detach().float().cpu().item()) if need_step_metrics and hasattr(base_model, "contact_injector") else float("nan") ), } if device.type == "cuda" and need_step_metrics: row["cuda_max_memory_allocated_gb"] = float(torch.cuda.max_memory_allocated(device) / 1024**3) row["cuda_memory_reserved_gb"] = float(torch.cuda.memory_reserved(device) / 1024**3) for key in sorted(set().union(*(x.keys() for x in accum_losses))): vals = [x[key] for x in accum_losses if key in x and isinstance(x[key], (int, float))] if vals: row[f"train_{key}"] = float(np.mean(vals)) if step == 1 or step % args.eval_every == 0 or step == args.max_steps: val_metrics = evaluate( model, config, val_arrays, args.eval_batch_size, args.eval_batches, device, args.precision, loss_kwargs=loss_kwargs, rollout_steps=rollout_steps, teacher_forcing_rate=teacher_forcing_rate, gumbel_temperature=gumbel_temperature, ) latest_eval = {f"val_{k}": v for k, v in val_metrics.items()} row.update(latest_eval) metric_value = row.get(args.early_stopping_metric) if args.early_stopping_patience > 0 and isinstance(metric_value, (int, float)) and math.isfinite(float(metric_value)): metric_value = float(metric_value) improved = best_metric_value is None or ( metric_value < best_metric_value - args.early_stopping_min_delta if args.early_stopping_mode == "min" else metric_value > best_metric_value + args.early_stopping_min_delta ) if improved: best_metric_value = metric_value best_step = step evals_since_best = 0 row["early_stopping_best_step"] = best_step row["early_stopping_best_metric"] = best_metric_value save_checkpoint(args.out_dir / "checkpoints" / "best.pt", model=model, optimizer=optimizer, step=step, config=config, tensor_metadata=meta, args=args, train_log_tail=row) else: evals_since_best += 1 row["early_stopping_evals_since_best"] = evals_since_best row["early_stopping_best_step"] = best_step row["early_stopping_best_metric"] = best_metric_value log_rows.append(row) if step == 1 or step % args.log_every == 0 or step % args.eval_every == 0 or step == args.max_steps: print(json.dumps(row, ensure_ascii=False), flush=True) should_stop = ( args.early_stopping_patience > 0 and step >= args.early_stopping_start_step and evals_since_best >= args.early_stopping_patience ) should_save = step % args.save_every == 0 or step == args.max_steps or should_stop should_upload = step % args.hub_upload_every == 0 or step == args.max_steps or should_stop if should_save or should_upload: pd.DataFrame(log_rows).to_csv(args.out_dir / "train_log.csv", index=False) save_checkpoint(args.out_dir / "checkpoints" / "latest.pt", model=model, optimizer=optimizer, step=step, config=config, tensor_metadata=meta, args=args, train_log_tail=row) if args.keep_step_checkpoints and should_save: save_checkpoint(args.out_dir / "checkpoints" / f"step_{step:06d}.pt", model=model, optimizer=optimizer, step=step, config=config, tensor_metadata=meta, args=args, train_log_tail=row) if should_upload: maybe_upload(args, args.out_dir, f"step_{step}") if should_stop: early_stopped = True stop_reason = f"early_stopping_patience_{args.early_stopping_patience}" print( json.dumps( { "stage": "early_stopping_triggered", "step": step, "metric": args.early_stopping_metric, "best_step": best_step, "best_metric": best_metric_value, "evals_since_best": evals_since_best, }, ensure_ascii=False, ), flush=True, ) break summary = { "status": "completed", "run_label": args.run_label, "stop_reason": stop_reason, "early_stopped": early_stopped, "best_step": best_step, "best_metric": best_metric_value, "early_stopping_metric": args.early_stopping_metric if args.early_stopping_patience > 0 else "", "selected_batch_size": selected_batch_size, "max_steps": args.max_steps, "actual_steps": int(log_rows[-1]["step"]) if log_rows else resume_step, "latest_eval": latest_eval, "elapsed_sec": time.time() - start_time, "out_dir": str(args.out_dir), "checkpoint": str(args.out_dir / "checkpoints" / "latest.pt"), } (args.out_dir / "train_result.json").write_text(json.dumps(summary, ensure_ascii=False, indent=2), encoding="utf-8") pd.DataFrame(log_rows).to_csv(args.out_dir / "train_log.csv", index=False) maybe_upload(args, args.out_dir, "final") return summary def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description="Train the SCTM-v4.4 closed-loop world-model heads on v3p5-compatible tensors.") parser.add_argument("--tensor-dir", type=Path, default=DEFAULT_TENSOR_DIR) parser.add_argument("--stage0-dir", type=Path, default=DEFAULT_STAGE0_DIR) parser.add_argument("--out-dir", type=Path, default=DEFAULT_OUT_DIR) parser.add_argument("--service-prior-file", type=str, default="service_state_transitions_train.json") parser.add_argument("--run-label", type=str, default="sctm_v4p4_world_model") parser.add_argument("--stage", choices=["stage1", "stage2a"], default="stage2a") parser.add_argument( "--curriculum-phase", choices=CURRICULUM_PHASES, default="", help="Unified-model curriculum alias for --stage; kept separate so old run scripts remain reproducible.", ) parser.add_argument("--device", type=str, default="cuda" if torch.cuda.is_available() else "cpu") parser.add_argument("--precision", choices=["fp32", "bf16", "fp16"], default="bf16") parser.add_argument("--seed", type=int, default=20260521) parser.add_argument("--d-model", type=int, default=896) parser.add_argument("--n-heads", type=int, default=14) parser.add_argument("--n-layers", type=int, default=10) parser.add_argument("--dropout", type=float, default=0.05) parser.add_argument("--n-missing", type=int, default=5) parser.add_argument("--batch-size", type=int, default=128) parser.add_argument("--min-batch-size", type=int, default=32) parser.add_argument("--auto-batch", action="store_true") parser.add_argument("--eval-batch-size", type=int, default=128) parser.add_argument("--grad-accum-steps", type=int, default=2) parser.add_argument("--max-steps", type=int, default=1000) parser.add_argument("--warmup-steps", type=int, default=100) parser.add_argument("--learning-rate", type=float, default=1.5e-4) parser.add_argument("--min-learning-rate", type=float, default=1.0e-5) parser.add_argument("--weight-decay", type=float, default=0.05) parser.add_argument("--beta1", type=float, default=0.9) parser.add_argument("--beta2", type=float, default=0.95) parser.add_argument("--grad-clip", type=float, default=1.0) parser.add_argument("--log-every", type=int, default=25) parser.add_argument("--eval-every", type=int, default=500) parser.add_argument("--eval-batches", type=int, default=5) parser.add_argument("--save-every", type=int, default=1000) parser.add_argument("--hub-upload-every", type=int, default=1000) parser.add_argument("--field-loss-weight", type=float, default=LOSS_WEIGHTS_V4P4_WORLD_MODEL_20260523["field_loss_weight"]) parser.add_argument("--numeric-loss-weight", type=float, default=LOSS_WEIGHTS_V4P4_WORLD_MODEL_20260523["numeric_loss_weight"]) parser.add_argument("--ordinal-loss-weight", type=float, default=LOSS_WEIGHTS_V4P4_WORLD_MODEL_20260523["ordinal_loss_weight"]) parser.add_argument("--missingness-loss-weight", type=float, default=LOSS_WEIGHTS_V4P4_WORLD_MODEL_20260523["missingness_loss_weight"]) parser.add_argument("--pre-event-risk-weight", type=float, default=LOSS_WEIGHTS_V4P4_WORLD_MODEL_20260523["pre_event_risk_weight"]) parser.add_argument("--event-risk-weight", type=float, default=LOSS_WEIGHTS_V4P4_WORLD_MODEL_20260523["event_risk_weight"]) parser.add_argument("--post-event-risk-weight", type=float, default=LOSS_WEIGHTS_V4P4_WORLD_MODEL_20260523["post_event_risk_weight"]) parser.add_argument("--event-generation-weight", type=float, default=LOSS_WEIGHTS_V4P4_WORLD_MODEL_20260523["event_generation_weight"]) parser.add_argument("--ontology-event-weight", type=float, default=LOSS_WEIGHTS_V4P4_WORLD_MODEL_20260523["ontology_event_weight"]) parser.add_argument("--history-aux-weight", type=float, default=LOSS_WEIGHTS_V4P4_WORLD_MODEL_20260523["history_aux_weight"]) parser.add_argument("--contact-aux-weight", type=float, default=LOSS_WEIGHTS_V4P4_WORLD_MODEL_20260523["contact_aux_weight"]) parser.add_argument("--color-loss-weight", type=float, default=LOSS_WEIGHTS_V4P4_WORLD_MODEL_20260523["color_loss_weight"]) parser.add_argument("--pwe-pre-nll-weight", type=float, default=LOSS_WEIGHTS_V4P4_WORLD_MODEL_20260523["pwe_pre_nll_weight"]) parser.add_argument("--pwe-contact-nll-weight", type=float, default=LOSS_WEIGHTS_V4P4_WORLD_MODEL_20260523["pwe_contact_nll_weight"]) parser.add_argument("--pwe-post-nll-weight", type=float, default=LOSS_WEIGHTS_V4P4_WORLD_MODEL_20260523["pwe_post_nll_weight"]) parser.add_argument("--pwe-pre-distill-weight", type=float, default=LOSS_WEIGHTS_V4P4_WORLD_MODEL_20260523["pwe_pre_distill_weight"]) parser.add_argument("--pwe-contact-distill-weight", type=float, default=LOSS_WEIGHTS_V4P4_WORLD_MODEL_20260523["pwe_contact_distill_weight"]) parser.add_argument("--rollout-steps", type=int, default=ROLLOUT_SCHEDULE_DEFAULTS_V4P4_WORLD_MODEL_20260523["rollout_steps"]) parser.add_argument("--rollout-weight-1", type=float, default=ROLLOUT_WEIGHTS_V4P4_WORLD_MODEL_20260523["rollout_weight_1"]) parser.add_argument("--rollout-weight-2", type=float, default=ROLLOUT_WEIGHTS_V4P4_WORLD_MODEL_20260523["rollout_weight_2"]) parser.add_argument("--rollout-weight-3", type=float, default=ROLLOUT_WEIGHTS_V4P4_WORLD_MODEL_20260523["rollout_weight_3"]) parser.add_argument("--teacher-forcing-rate", type=float, default=ROLLOUT_SCHEDULE_DEFAULTS_V4P4_WORLD_MODEL_20260523["teacher_forcing_rate"]) parser.add_argument("--teacher-forcing-start", type=float, default=ROLLOUT_SCHEDULE_DEFAULTS_V4P4_WORLD_MODEL_20260523["teacher_forcing_start"]) parser.add_argument("--teacher-forcing-end", type=float, default=ROLLOUT_SCHEDULE_DEFAULTS_V4P4_WORLD_MODEL_20260523["teacher_forcing_end"]) parser.add_argument("--gumbel-temperature-start", type=float, default=ROLLOUT_SCHEDULE_DEFAULTS_V4P4_WORLD_MODEL_20260523["gumbel_temperature_start"]) parser.add_argument("--gumbel-temperature-end", type=float, default=ROLLOUT_SCHEDULE_DEFAULTS_V4P4_WORLD_MODEL_20260523["gumbel_temperature_end"]) parser.add_argument("--compile", action="store_true") parser.add_argument("--sample-with-replacement", action="store_true") parser.add_argument("--finite-check-every", type=int, default=25) parser.add_argument("--no-freeze-backbone-stage1", dest="freeze_backbone_stage1", action="store_false") parser.set_defaults(freeze_backbone_stage1=True) parser.add_argument("--fused-adamw", action="store_true") parser.add_argument("--matmul-precision", type=str, default="high") parser.add_argument("--init-checkpoint", type=Path, default=None) parser.add_argument("--resume-checkpoint", type=Path, default=None) parser.add_argument("--keep-step-checkpoints", action="store_true") parser.add_argument("--early-stopping-patience", type=int, default=0) parser.add_argument("--early-stopping-metric", type=str, default="val_loss_total") parser.add_argument("--early-stopping-mode", choices=["min", "max"], default="min") parser.add_argument("--early-stopping-min-delta", type=float, default=0.0) parser.add_argument("--early-stopping-start-step", type=int, default=0) parser.add_argument("--hub-repo-id", type=str, default="") parser.add_argument("--hub-repo-type", type=str, default="model") parser.add_argument("--hub-path", type=str, default="sctm_v4p4_world_model") parsed = parser.parse_args() if parsed.curriculum_phase and "--stage" not in sys.argv[1:]: parsed.stage = CURRICULUM_PHASE_TO_LEGACY_STAGE[parsed.curriculum_phase] return normalize_curriculum_args(parsed) if __name__ == "__main__": result = run(parse_args()) print(json.dumps(result, ensure_ascii=False), flush=True)