agoniii97's picture
Normalize datetime precision for HF P1 tensor build
ae6d94c verified
Raw
History Blame Contribute Delete
54.9 kB
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)