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