from __future__ import annotations import argparse import json import math import sys from collections import defaultdict from dataclasses import asdict from pathlib import Path from typing import Any import numpy as np import pandas as pd import torch import torch.nn.functional as F SCRIPT_DIR = Path(__file__).resolve().parent ROOT_DIR = SCRIPT_DIR.parents[1] V4P4_SCRIPT_DIR = ROOT_DIR / "v4p4_world_model" / "scripts" V3P5_SCRIPT_DIR = ROOT_DIR / "v3p5_static" / "scripts" for path in (SCRIPT_DIR, V4P4_SCRIPT_DIR, V3P5_SCRIPT_DIR): if str(path) not in sys.path: sys.path.insert(0, str(path)) from action_ontology_v5 import action_timing_audit_rows, build_action_ontology, load_json # noqa: E402 from build_v5_action_tensors import ( # noqa: E402 build_action_arrays, build_target_trial_labels, medication_flags_from_stage0, resolve_split_path, validate_medication_flags, ) from config_v5_train import TARGET_TRIAL_DEFAULTS # noqa: E402 from export_v4p4_visit_outputs import load_model as load_v4p4_model # noqa: E402 from loss_v4p3 import piecewise_exp_competing_risk_nll, pwe_target_from_terminal, v4_missing_targets # noqa: E402 from model_v4p4 import pwe_closed_form_cif # noqa: E402 from model_v5 import SCTMv5, SCTMv5Config # noqa: E402 from smoke_v3p4_architecture import build_field_value_mask, load_service_prior # noqa: E402 from train_v4p4_cloud import autocast_context, batch_from_indices, load_npz_to_memory # noqa: E402 try: from sklearn.metrics import average_precision_score, roc_auc_score except Exception: # pragma: no cover average_precision_score = None roc_auc_score = None CAUSE_NAMES = {0: "next_contact", 1: "death", 2: "disengagement"} OBSERVED_ACTION = "v5_action" V5_BASE = "v5_base_post" V4_BASE = "v4_post" 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() if isinstance(value, float) and not math.isfinite(value): return None return value def safe_mean(values: list[float]) -> float: arr = np.asarray([x for x in values if math.isfinite(float(x))], dtype=np.float64) return float(arr.mean()) if arr.size else math.nan def binary_auc_ap(y: np.ndarray, p: np.ndarray) -> tuple[float, float]: y = y.astype(np.int64).reshape(-1) p = np.clip(p.astype(np.float64).reshape(-1), 1.0e-7, 1.0 - 1.0e-7) if y.size == 0 or np.unique(y).size < 2: return math.nan, math.nan auc = float(roc_auc_score(y, p)) if roc_auc_score is not None else math.nan ap = float(average_precision_score(y, p)) if average_precision_score is not None else math.nan return auc, ap def slot_allowed_value_mask(ontology: dict[str, Any], slot_name: str, vocab_size: int, device: torch.device) -> torch.Tensor: allowed = ontology.get("allowed_action_value_ids_by_slot", {}).get(slot_name) mask = torch.zeros(vocab_size, dtype=torch.bool, device=device) if not allowed: mask.fill_(True) return mask ids = torch.as_tensor([int(x) for x in allowed if 0 <= int(x) < vocab_size], dtype=torch.long, device=device) if ids.numel() == 0: mask.fill_(True) return mask mask[ids] = True return mask def binary_calibration(y: np.ndarray, p: np.ndarray, n_bins: int = 10) -> dict[str, float]: y = y.astype(np.float64).reshape(-1) p = np.clip(p.astype(np.float64).reshape(-1), 1.0e-7, 1.0 - 1.0e-7) if y.size == 0: return {"brier": math.nan, "ece": math.nan, "ici": math.nan, "mean_predicted": math.nan, "observed_rate": math.nan} order = np.argsort(p) bins = np.array_split(order, min(n_bins, max(1, y.size))) weighted_abs_calibration_error = 0.0 for idx in bins: if idx.size == 0: continue weighted_abs_calibration_error += abs(float(p[idx].mean()) - float(y[idx].mean())) * float(idx.size) / float(y.size) # Legacy full-eval tables retain the column name `ici`; here it is an equal-count # binned calibration proxy. The manuscript-grade path uses isotonic/smooth ICI. return { "brier": float(np.mean((p - y) ** 2)), "ece": float(weighted_abs_calibration_error), "ici": float(weighted_abs_calibration_error), "mean_predicted": float(p.mean()), "observed_rate": float(y.mean()), } def multiclass_confusion(y: np.ndarray, pred: np.ndarray, n_classes: int) -> np.ndarray: conf = np.zeros((n_classes, n_classes), dtype=np.int64) y = y.reshape(-1).astype(np.int64) pred = pred.reshape(-1).astype(np.int64) mask = (y >= 0) & (y < n_classes) & (pred >= 0) & (pred < n_classes) if mask.any(): np.add.at(conf, (y[mask], pred[mask]), 1) return conf def macro_f1(conf: np.ndarray) -> float: vals = [] for cls in range(conf.shape[0]): support = float(conf[cls, :].sum()) if support <= 0: continue tp = float(conf[cls, cls]) fp = float(conf[:, cls].sum() - tp) fn = float(conf[cls, :].sum() - tp) precision = tp / (tp + fp) if tp + fp > 0 else 0.0 recall = tp / (tp + fn) if tp + fn > 0 else 0.0 vals.append(2.0 * precision * recall / (precision + recall) if precision + recall > 0 else 0.0) return float(np.mean(vals)) if vals else math.nan def multiclass_ece(target: np.ndarray, probs: np.ndarray, n_bins: int = 10) -> tuple[float, float]: if target.size == 0: return math.nan, math.nan pred = probs.argmax(axis=-1) conf = probs.max(axis=-1) correct = pred == target order = np.argsort(conf) bins = np.array_split(order, min(n_bins, max(1, target.size))) ece = 0.0 for idx in bins: if idx.size: ece += abs(float(conf[idx].mean()) - float(correct[idx].mean())) * float(idx.size) / float(target.size) ici = float(np.mean(np.abs(conf - correct.astype(np.float64)))) return float(ece), ici def quadratic_weighted_kappa(y_true: np.ndarray, y_pred: np.ndarray, n_levels: int) -> float: if y_true.size == 0: return math.nan y_true = np.clip(y_true.astype(int), 0, n_levels - 1) y_pred = np.clip(y_pred.astype(int), 0, n_levels - 1) observed = np.zeros((n_levels, n_levels), dtype=np.float64) for t, p in zip(y_true, y_pred): observed[t, p] += 1.0 hist_t = observed.sum(axis=1) hist_p = observed.sum(axis=0) expected = np.outer(hist_t, hist_p) / max(1.0, observed.sum()) weights = np.zeros_like(observed) denom = float((n_levels - 1) ** 2) if n_levels > 1 else 1.0 for i in range(n_levels): for j in range(n_levels): weights[i, j] = ((i - j) ** 2) / denom obs = float((weights * observed).sum()) exp = float((weights * expected).sum()) return float(1.0 - obs / exp) if exp > 0 else math.nan def tensorize_actions(action_arrays: dict[str, np.ndarray], idx: np.ndarray, device: torch.device) -> dict[str, torch.Tensor]: return { "action_type_ids": torch.as_tensor(action_arrays["action_type_ids"][idx], dtype=torch.long, device=device), "action_value_ids": torch.as_tensor(action_arrays["action_value_ids"][idx], dtype=torch.long, device=device), "action_mask": torch.as_tensor(action_arrays["action_mask"][idx], dtype=torch.bool, device=device), "action_available_at": torch.as_tensor(action_arrays["action_available_at"][idx], dtype=torch.long, device=device), } def validate_action_arrays(action_arrays: dict[str, np.ndarray], arrays: dict[str, np.ndarray], ontology: dict[str, Any], *, source: Path | str) -> None: required = ("action_type_ids", "action_value_ids", "action_mask", "action_available_at") missing = [key for key in required if key not in action_arrays] if missing: raise ValueError(f"Action tensor source {source} is missing required arrays: {missing}") n_windows, seq_len = arrays["valid_mask"].shape n_slots = len(ontology["slots"]) expected = (n_windows, seq_len, n_slots) for key in required: shape = tuple(action_arrays[key].shape) if shape != expected: raise ValueError(f"Action tensor source {source} has {key} shape {shape}, expected {expected}") if not np.array_equal(action_arrays["action_mask"].astype(bool), action_arrays["action_mask"]): raise ValueError(f"Action tensor source {source} has non-boolean-compatible action_mask values") def config_from_payload(payload: dict[str, Any]) -> SCTMv5Config: cfg = dict(payload["config"]) cfg["conditional_specs"] = tuple(cfg.get("conditional_specs", ())) cfg["color_value_to_class"] = tuple(tuple(x) for x in cfg.get("color_value_to_class", ())) return SCTMv5Config(**cfg) def load_v5_model( checkpoint: Path, tensor_dir: Path, stage0_dir: Path, service_prior_file: str, device: torch.device, ) -> tuple[SCTMv5, dict[str, Any], dict[str, Any]]: payload = torch.load(checkpoint, map_location="cpu", weights_only=False) meta = load_json(tensor_dir / "tensor_metadata.json") vocab = load_json(tensor_dir / "cat_value_vocab.json") config = config_from_payload(payload) train_path = resolve_split_path(tensor_dir, "train") with np.load(train_path, allow_pickle=True) as z: mask_arrays = {"cat_value_ids": z["cat_value_ids"], "missing_ids": z["missing_ids"]} field_value_mask, field_mask_audit = build_field_value_mask(meta, vocab, mask_arrays) prior, prior_audit = load_service_prior(stage0_dir, int(config.n_service_states), service_prior_file) model = SCTMv5(config, field_value_mask=field_value_mask, service_prior_bias=prior).to(device) state = payload.get("model", payload.get("model_state_dict")) if state is None: raise ValueError("v5 checkpoint lacks model/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()} incompatible = model.load_state_dict(state, strict=True) model.eval() audit = { "checkpoint_path": str(checkpoint), "checkpoint_keys": list(payload.keys()), "state_missing_keys": list(incompatible.missing_keys), "state_unexpected_keys": list(incompatible.unexpected_keys), "field_value_mask_audit": field_mask_audit, "service_prior_audit": prior_audit, } return model, payload, audit def source_audit(model: SCTMv5, arrays: dict[str, np.ndarray], action_arrays: dict[str, np.ndarray], device: torch.device) -> dict[str, Any]: idx = np.arange(min(2, arrays["valid_mask"].shape[0])) batch = batch_from_indices(arrays, idx, device) batch.update(tensorize_actions(action_arrays, idx, device)) model.eval() with torch.inference_mode(): out = model(batch, rollout_steps=1, compute_pwe_diagnostics=False) pos = None for t in range(1, batch["valid_mask"].shape[1] - 1): if bool(batch["valid_mask"][0, t]) and bool(batch["valid_mask"][0, t + 1]): pos = t break if pos is None: return {"checked": False, "passed": False, "reason": "no valid adjacent visits"} future_action = {k: v.clone() if torch.is_tensor(v) else v for k, v in batch.items()} future_action["action_value_ids"][0, pos + 1, 0] = (future_action["action_value_ids"][0, pos + 1, 0] + 1) % model.config.action_value_vocab_size out_future_action = model(future_action, rollout_steps=1, compute_pwe_diagnostics=False) current_action = {k: v.clone() if torch.is_tensor(v) else v for k, v in batch.items()} current_action["action_value_ids"][0, pos, 0] = (current_action["action_value_ids"][0, pos, 0] + 1) % model.config.action_value_vocab_size out_current_action = model(current_action, rollout_steps=1, compute_pwe_diagnostics=False) future_deltas = { "pwe_log_lambda_action": float((out["pwe_log_lambda_action"][0, pos] - out_future_action["pwe_log_lambda_action"][0, pos]).abs().max().cpu()), "active_state_logits_action": float((out["active_state_logits_action"][0, pos] - out_future_action["active_state_logits_action"][0, pos]).abs().max().cpu()), "event_generation_logits_action": float((out["event_generation_logits_action"][0, pos] - out_future_action["event_generation_logits_action"][0, pos]).abs().max().cpu()), } behavior_delta = float((out["behavior_policy_logits"][0, pos] - out_current_action["behavior_policy_logits"][0, pos]).abs().max().cpu()) action_context_delta = float((out["action_context"][0, pos] - out_current_action["action_context"][0, pos]).abs().max().cpu()) return { "checked": True, "position": int(pos), "future_action_current_output_deltas": future_deltas, "behavior_policy_current_action_delta": behavior_delta, "action_context_current_action_delta": action_context_delta, "formal_risk_export_keys": [ "pwe_log_lambda_action", "active_state_logits_action", "missingness_logits_action", "field_logits_action", "numeric_mu_action", "ordinal_cum_logits_action", "event_generation_logits_action", "behavior_policy_logits", ], "excluded_from_formal_risk_export": ["ontology_event_logits", "event_logits", "diagnostic_ontology_event_logits"], "passed": bool(max(future_deltas.values()) < 1.0e-6 and behavior_delta < 1.0e-6 and action_context_delta > 1.0e-8), } def evaluate_outputs( *, model_name: str, out: dict[str, torch.Tensor], batch: dict[str, torch.Tensor], cfg: Any, log_lambda_key: str, active_key: str, missing_key: str, field_key: str, numeric_key: str, ordinal_key: str, event_key: str, split: str, horizons: list[float], ) -> dict[str, list[dict[str, Any]]]: rows: dict[str, list[dict[str, Any]]] = defaultdict(list) cause, censored, valid_next, terminal = pwe_target_from_terminal(batch) event_days = torch.expm1(batch["delta_t_next_log"][:, :-1]).clamp(min=1.0e-6) next_contact = valid_next & terminal.eq(0) & (~censored) & batch["service_state"][:, 1:].lt(cfg.n_active_states) valid_count = int(valid_next.sum().detach().cpu()) next_contact_count = int(next_contact.sum().detach().cpu()) pwe_nll = piecewise_exp_competing_risk_nll(out[log_lambda_key][:, :-1, :, :], event_days, cause, censored, valid_next) rows["process_nll"].append( { "split": split, "model": model_name, "metric": "pwe_nll", "value": float(pwe_nll.detach().cpu()), "weight": valid_count, } ) active_target = batch["service_state"][:, 1:].clamp(min=0, max=cfg.n_active_states - 1) if next_contact.any(): active_logits = out[active_key][:, :-1, :][next_contact].float() active_y = active_target[next_contact].long() active_pred = active_logits.argmax(dim=-1) conf = multiclass_confusion(active_y.detach().cpu().numpy(), active_pred.detach().cpu().numpy(), cfg.n_active_states) rows["grammar_active_state"].append( { "split": split, "model": model_name, "n": int(active_y.numel()), "cross_entropy": float(F.cross_entropy(active_logits, active_y).detach().cpu()), "accuracy": float((active_pred == active_y).float().mean().detach().cpu()), "macro_f1": macro_f1(conf), } ) missing_target = v4_missing_targets(out, batch, cfg) miss_logits = out[missing_key][:, :-1, :, :].float() miss_probs = torch.softmax(miss_logits, dim=-1) miss_pred = miss_probs.argmax(dim=-1) miss_mask = valid_next[:, :, None].expand_as(missing_target) if miss_mask.any(): y = missing_target[miss_mask].detach().cpu().numpy().astype(np.int64) p = miss_probs[miss_mask].detach().cpu().numpy() pred = miss_pred[miss_mask].detach().cpu().numpy().astype(np.int64) conf = multiclass_confusion(y, pred, cfg.n_missing) onehot = np.eye(cfg.n_missing, dtype=np.float32)[y] ece, ici = multiclass_ece(y, p) rows["obs_missingness"].append( { "split": split, "model": model_name, "n": int(y.size), "accuracy": float((pred == y).mean()), "macro_f1": macro_f1(conf), "multiclass_brier": float(np.mean(np.sum((p - onehot) ** 2, axis=-1))), "confidence_ece": ece, "confidence_ici": ici, } ) field_logits = out[field_key][:, :-1, :, :].float() field_mask = next_contact[:, :, None] & missing_target.eq(cfg.missing_observed_id) field_losses = [] field_top1 = [] field_top3 = [] for field_idx in range(cfg.n_cat_fields): fm = field_mask[:, :, field_idx] if not fm.any(): continue logits = field_logits[:, :, field_idx, :][fm] target = batch["cat_value_ids"][:, 1:, field_idx][fm].long() top3 = logits.topk(k=min(3, logits.shape[-1]), dim=-1).indices field_losses.append(float(F.cross_entropy(logits, target).detach().cpu())) field_top1.append(float((logits.argmax(dim=-1) == target).float().mean().detach().cpu())) field_top3.append(float(top3.eq(target[:, None]).any(dim=-1).float().mean().detach().cpu())) rows["grammar_cat_fields"].append( { "split": split, "model": model_name, "n_next_contacts": next_contact_count, "mean_cross_entropy": safe_mean(field_losses), "mean_top1_accuracy": safe_mean(field_top1), "mean_top3_accuracy": safe_mean(field_top3), } ) num_pred = out[numeric_key][:, :-1, :] num_abs = [] num_sq = [] for num_idx in range(cfg.n_numeric_fields): nm = next_contact & batch["numeric_mask"][:, 1:, num_idx] if nm.any(): diff = (num_pred[:, :, num_idx][nm] - batch["numeric_values"][:, 1:, num_idx][nm]).detach().cpu().float().numpy() num_abs.append(np.abs(diff)) num_sq.append(diff * diff) abs_cat = np.concatenate(num_abs) if num_abs else np.array([], dtype=np.float32) sq_cat = np.concatenate(num_sq) if num_sq else np.array([], dtype=np.float32) rows["grammar_numeric"].append( { "split": split, "model": model_name, "n": int(abs_cat.size), "mae": float(abs_cat.mean()) if abs_cat.size else math.nan, "rmse": float(math.sqrt(float(sq_cat.mean()))) if sq_cat.size else math.nan, } ) ord_logits = out[ordinal_key][:, :-1, :, :].float() ord_pred = torch.sigmoid(ord_logits).ge(0.5).long().sum(dim=-1) ord_target = batch["ordinal_cbe"][:, 1:, :, :].ge(0.5).long().sum(dim=-1) ord_y = [] ord_p = [] for ord_idx in range(cfg.n_ordinal_fields): om = next_contact & batch["ordinal_mask"][:, 1:, ord_idx] if om.any(): ord_y.append(ord_target[:, :, ord_idx][om].detach().cpu().numpy()) ord_p.append(ord_pred[:, :, ord_idx][om].detach().cpu().numpy()) if ord_y: y = np.concatenate(ord_y) p = np.concatenate(ord_p) rows["grammar_ordinal"].append( { "split": split, "model": model_name, "n": int(y.size), "mae": float(np.abs(p - y).mean()), "rmse": float(math.sqrt(float(((p - y) ** 2).mean()))), "quadratic_weighted_kappa": quadratic_weighted_kappa(y, p, cfg.cbe_dim + 1), } ) event_prob = torch.sigmoid(out[event_key][:, :-1, :].float()).detach().cpu().numpy() event_target = batch["event_labels"][:, 1:, :].detach().cpu().numpy().astype(np.int64) next_np = next_contact.detach().cpu().numpy().astype(bool) for event_idx in range(cfg.n_events): y = event_target[:, :, event_idx][next_np] p = event_prob[:, :, event_idx][next_np] auc, ap = binary_auc_ap(y, p) cal = binary_calibration(y, p) rows["grammar_events"].append( { "split": split, "model": model_name, "event_index": int(event_idx), "n": int(y.size), "events": int(y.sum()) if y.size else 0, "y_values": y.astype(np.int8), "p_values": p.astype(np.float32), "auc": auc, "average_precision": ap, **cal, } ) horizons_t = torch.tensor(horizons, device=batch["valid_mask"].device, dtype=torch.float32) pwe = pwe_closed_form_cif(out[log_lambda_key][:, :-1, :, :], horizons_t) cif = pwe["cif"].detach().cpu().float().numpy() cause_np = cause.detach().cpu().numpy() cens_np = censored.detach().cpu().numpy().astype(bool) valid_np = valid_next.detach().cpu().numpy().astype(bool) days_np = event_days.detach().cpu().numpy() for h_idx, horizon in enumerate(horizons): for cause_id, cause_name in CAUSE_NAMES.items(): pred = cif[:, :, h_idx, cause_id] evaluable = valid_np & ~(cens_np & (days_np <= float(horizon))) target = ((~cens_np) & (cause_np == cause_id) & (days_np <= float(horizon))).astype(np.int64) y = target[evaluable] p = pred[evaluable] auc, ap = binary_auc_ap(y, p) cal = binary_calibration(y, p) rows["pwe_horizon_risk"].append( { "split": split, "model": model_name, "cause_id": int(cause_id), "cause_name": cause_name, "horizon_days": float(horizon), "n": int(y.size), "events": int(y.sum()) if y.size else 0, "y_values": y.astype(np.int8), "p_values": p.astype(np.float32), "auc": auc, "average_precision": ap, **cal, } ) return rows def behavior_policy_rows( *, out: dict[str, torch.Tensor], batch: dict[str, torch.Tensor], action_arrays: dict[str, np.ndarray], idx: np.ndarray, ontology: dict[str, Any], split: str, ) -> dict[str, list[dict[str, Any]]]: rows: dict[str, list[dict[str, Any]]] = defaultdict(list) logits = out["behavior_policy_logits"].float() target = batch["action_value_ids"].long().clamp(min=0, max=logits.shape[-1] - 1) mask = batch["action_mask"].bool() & batch["valid_mask"][:, :, None].bool() slot_names = [slot["slot_name"] for slot in ontology["slots"]] value_vocab_inv = {int(v): k for k, v in ontology["action_value_vocab"].items()} for slot_id, slot_name in enumerate(slot_names): m = mask[:, :, slot_id] if not m.any(): continue allowed_mask = slot_allowed_value_mask(ontology, slot_name, logits.shape[-1], logits.device) y_t = target[:, :, slot_id][m] invalid_target_count = int((~allowed_mask[y_t]).sum().detach().cpu()) if invalid_target_count: allowed_mask = allowed_mask.clone() allowed_mask[y_t.unique()] = True logits_slot = logits[:, :, slot_id, :].masked_fill(~allowed_mask.view(1, 1, -1), -1.0e9) probs_slot = torch.softmax(logits_slot, dim=-1) logits_t = logits_slot[m] probs_t = probs_slot[m] pred_t = logits_t.argmax(dim=-1) y = y_t.detach().cpu().numpy() pred = pred_t.detach().cpu().numpy() p = probs_t.detach().cpu().numpy() conf = multiclass_confusion(y, pred, logits.shape[-1]) ece, ici = multiclass_ece(y, p) rows["policy_discrimination"].append( { "split": split, "slot_id": int(slot_id), "slot_name": slot_name, "n": int(y.size), "cross_entropy": float(F.cross_entropy(logits_t, y_t).detach().cpu()), "accuracy": float((pred_t == y_t).float().mean().detach().cpu()), "macro_f1": macro_f1(conf), "confidence_ece": ece, "confidence_ici": ici, "invalid_target_count": invalid_target_count, "invalid_target_rate": float(invalid_target_count / max(1, y.size)), } ) values, counts = np.unique(y, return_counts=True) for value_id, count in zip(values.tolist(), counts.tolist()): rows["policy_prevalence"].append( { "split": split, "slot_id": int(slot_id), "slot_name": slot_name, "value_id": int(value_id), "value_label": value_vocab_inv.get(int(value_id), str(value_id)), "n": int(count), "rate_among_available": float(count / max(1, y.size)), } ) local = ontology.get("local_label_to_action_value_id_by_slot", {}).get(slot_name, {}) present_id = local.get("present") if present_id is not None: yy = (y == int(present_id)).astype(np.int64) pp = p[:, int(present_id)] cal = binary_calibration(yy, pp) weights = yy / np.clip(pp, 1.0e-3, 1.0) + (1 - yy) / np.clip(1.0 - pp, 1.0e-3, 1.0) ess = float((weights.sum() ** 2) / np.sum(weights**2)) if weights.size else math.nan rows["policy_overlap"].append( { "split": split, "slot_id": int(slot_id), "slot_name": slot_name, "present_value_id": int(present_id), "n": int(yy.size), "present_rate": float(yy.mean()) if yy.size else math.nan, "mean_propensity": float(pp.mean()) if pp.size else math.nan, "p01": float(np.quantile(pp, 0.01)) if pp.size else math.nan, "p05": float(np.quantile(pp, 0.05)) if pp.size else math.nan, "p50": float(np.quantile(pp, 0.50)) if pp.size else math.nan, "p95": float(np.quantile(pp, 0.95)) if pp.size else math.nan, "p99": float(np.quantile(pp, 0.99)) if pp.size else math.nan, "extreme_propensity_rate": float(((pp < 0.01) | (pp > 0.99)).mean()) if pp.size else math.nan, "effective_sample_size_binary_ipw": ess, **cal, } ) service = batch["service_state"].detach().cpu().numpy() years = batch["visit_year"].detach().cpu().numpy() action_values = batch["action_value_ids"].detach().cpu().numpy() action_mask = mask.detach().cpu().numpy() for slot_id, slot_name in enumerate(slot_names): m = action_mask[:, :, slot_id] if not m.any(): continue vals = action_values[:, :, slot_id][m] svc = service[m] era = np.clip((years[m] - 2018) // 3, 0, 2) for era_id in sorted(np.unique(era).tolist()): sub = vals[era == era_id] rows["policy_by_era"].append({"split": split, "slot_id": int(slot_id), "slot_name": slot_name, "era_id": int(era_id), "n": int(sub.size), "n_unique": int(np.unique(sub).size)}) for state in sorted(np.unique(svc).tolist()): sub = vals[svc == state] rows["policy_by_state"].append({"split": split, "slot_id": int(slot_id), "slot_name": slot_name, "service_state": int(state), "n": int(sub.size), "n_unique": int(np.unique(sub).size)}) return rows @torch.inference_mode() def evaluate_split( *, split: str, arrays: dict[str, np.ndarray], action_arrays: dict[str, np.ndarray], ontology: dict[str, Any], v5_model: SCTMv5, v4_model: torch.nn.Module | None, device: torch.device, precision: str, batch_size: int, max_windows: int, horizons: list[float], rollout_steps: int, rollout_max_windows: int, ) -> dict[str, list[dict[str, Any]]]: all_rows: dict[str, list[dict[str, Any]]] = defaultdict(list) n = int(arrays["valid_mask"].shape[0]) if max_windows > 0: n = min(n, int(max_windows)) for start in range(0, n, batch_size): idx = np.arange(start, min(start + batch_size, n)) batch = batch_from_indices(arrays, idx, device) batch.update(tensorize_actions(action_arrays, idx, device)) with autocast_context(device, precision): out_v5 = v5_model(batch, rollout_steps=1, compute_pwe_diagnostics=False) if v4_model is not None: out_v4 = v4_model(batch, rollout_steps=1, compute_pwe_diagnostics=False) else: out_v4 = None for name, frame_rows in evaluate_outputs( model_name=OBSERVED_ACTION, out=out_v5, batch=batch, cfg=v5_model.config, log_lambda_key="pwe_log_lambda_action", active_key="active_state_logits_action", missing_key="missingness_logits_action", field_key="field_logits_action", numeric_key="numeric_mu_action", ordinal_key="ordinal_cum_logits_action", event_key="event_generation_logits_action", split=split, horizons=horizons, ).items(): all_rows[name].extend(frame_rows) for name, frame_rows in evaluate_outputs( model_name=V5_BASE, out=out_v5, batch=batch, cfg=v5_model.config, log_lambda_key="pwe_log_lambda_post", active_key="active_state_logits", missing_key="missingness_logits", field_key="field_logits", numeric_key="numeric_mu", ordinal_key="ordinal_cum_logits", event_key="event_generation_logits", split=split, horizons=horizons, ).items(): all_rows[name].extend(frame_rows) if out_v4 is not None: for name, frame_rows in evaluate_outputs( model_name=V4_BASE, out=out_v4, batch=batch, cfg=v4_model.config, log_lambda_key="pwe_log_lambda_post", active_key="active_state_logits", missing_key="missingness_logits", field_key="field_logits", numeric_key="numeric_mu", ordinal_key="ordinal_cum_logits", event_key="event_generation_logits", split=split, horizons=horizons, ).items(): all_rows[name].extend(frame_rows) for name, rows in behavior_policy_rows(out=out_v5, batch=batch, action_arrays=action_arrays, idx=idx, ontology=ontology, split=split).items(): all_rows[name].extend(rows) if rollout_steps > 0: rn = min(n, rollout_max_windows) if rollout_max_windows > 0 else n legality_rows = [] for model_name, model in [(OBSERVED_ACTION, v5_model), (V4_BASE, v4_model)]: if model is None: continue counts = defaultdict(int) for start in range(0, rn, batch_size): idx = np.arange(start, min(start + batch_size, rn)) batch = batch_from_indices(arrays, idx, device) batch.update(tensorize_actions(action_arrays, idx, device)) rollout = model.ancestral_rollout(batch, start_pos=0, steps=rollout_steps, deterministic=False, rao_blackwell_rare=True) generated = rollout["generated_batch"] gen_valid = generated["valid_mask"].detach().cpu().numpy().astype(bool) gen_service = generated["service_state"].detach().cpu().numpy().astype(np.int64) gen_missing = generated["missing_ids"].detach().cpu().numpy().astype(np.int64) cfg = model.config for pos in range(1, min(gen_valid.shape[1], rollout_steps + 1)): valid_pos = gen_valid[:, pos] counts["generated_positions"] += int(valid_pos.sum()) if not valid_pos.any(): continue svc = gen_service[:, pos] miss = gen_missing[:, pos, :] active = valid_pos & (svc < cfg.n_active_states) terminal = valid_pos & (svc >= cfg.n_active_states) counts["invalid_service_id"] += int(((svc[valid_pos] < 0) | (svc[valid_pos] >= cfg.n_service_states)).sum()) counts["active_contact_no_clinical_missingness"] += int((miss[active] == cfg.missing_no_clinical_id).sum()) if active.any() else 0 counts["active_contact_visit_missing"] += int((miss[active] == cfg.missing_visit_missing_id).sum()) if active.any() else 0 counts["terminal_non_no_clinical_missingness"] += int((miss[terminal] != cfg.missing_no_clinical_id).sum()) if terminal.any() else 0 den = max(1, counts["generated_positions"]) for key, value in counts.items(): if key == "generated_positions": continue legality_rows.append( { "split": split, "model": model_name, "rollout_steps": int(rollout_steps), "generated_positions": int(counts["generated_positions"]), "violation": key, "count": int(value), "rate_per_generated_position": float(value / den), } ) all_rows["rollout_legality"].extend(legality_rows) return all_rows def aggregate_frame(rows: list[dict[str, Any]], group_cols: list[str]) -> pd.DataFrame: if not rows: return pd.DataFrame() df = pd.DataFrame(rows) if {"y_values", "p_values"}.issubset(df.columns): out_rows = [] for keys, group in df.groupby(group_cols, dropna=False): if not isinstance(keys, tuple): keys = (keys,) y_parts = [np.asarray(v).reshape(-1) for v in group["y_values"].tolist() if np.asarray(v).size] p_parts = [np.asarray(v).reshape(-1) for v in group["p_values"].tolist() if np.asarray(v).size] y = np.concatenate(y_parts) if y_parts else np.array([], dtype=np.int8) p = np.concatenate(p_parts) if p_parts else np.array([], dtype=np.float32) auc, ap = binary_auc_ap(y, p) cal = binary_calibration(y, p) row = {col: key for col, key in zip(group_cols, keys)} row.update({"n": int(y.size), "events": int(y.sum()) if y.size else 0, "auc": auc, "average_precision": ap, **cal}) out_rows.append(row) return pd.DataFrame(out_rows) numeric_cols = [c for c in df.columns if c not in group_cols and pd.api.types.is_numeric_dtype(df[c])] if "weight" in df.columns and "value" in df.columns: def wavg(g: pd.DataFrame) -> pd.Series: w = g["weight"].fillna(0).to_numpy(dtype=float) v = g["value"].to_numpy(dtype=float) den = float(w.sum()) return pd.Series({"value": float(np.sum(v * w) / den) if den > 0 else float(np.nan), "weight": den}) return df.groupby(group_cols, dropna=False).apply(wavg, include_groups=False).reset_index() if "n" in df.columns: weighted_cols = [c for c in numeric_cols if c not in {"n", "events"}] count_cols = [c for c in ("n", "events", "n_next_contacts", "generated_positions", "count") if c in df.columns] def weighted_by_n(g: pd.DataFrame) -> pd.Series: w = g["n"].fillna(0).to_numpy(dtype=float) den = float(w.sum()) out: dict[str, float] = {} for col in weighted_cols: vals = g[col].to_numpy(dtype=float) ok = np.isfinite(vals) & np.isfinite(w) & (w > 0) out[col] = float(np.sum(vals[ok] * w[ok]) / np.sum(w[ok])) if ok.any() and np.sum(w[ok]) > 0 else math.nan for col in count_cols: out[col] = float(g[col].fillna(0).sum()) return pd.Series(out) return df.groupby(group_cols, dropna=False).apply(weighted_by_n, include_groups=False).reset_index() agg = {c: "mean" for c in numeric_cols} for count_col in ("n", "events", "n_next_contacts"): if count_col in agg: agg[count_col] = "sum" return df.groupby(group_cols, dropna=False).agg(agg).reset_index() def compare_models(df: pd.DataFrame, index_cols: list[str], metric_cols: list[str], baseline: str = V4_BASE, candidate: str = OBSERVED_ACTION) -> pd.DataFrame: if df.empty or "model" not in df.columns: return pd.DataFrame() base = df[df["model"].eq(baseline)] cand = df[df["model"].eq(candidate)] if base.empty or cand.empty: return pd.DataFrame() merged = cand.merge(base, on=index_cols, suffixes=("_v5", "_v4")) rows = [] lower_better = {"value", "cross_entropy", "mean_cross_entropy", "brier", "ece", "ici", "confidence_ece", "confidence_ici", "multiclass_brier", "mae", "rmse"} for _, row in merged.iterrows(): out = {col: row[col] for col in index_cols} for metric in metric_cols: v5 = row.get(f"{metric}_v5") v4 = row.get(f"{metric}_v4") if pd.isna(v5) or pd.isna(v4): continue out[f"{metric}_v5"] = float(v5) out[f"{metric}_v4"] = float(v4) out[f"{metric}_raw_delta_v5_minus_v4"] = float(v5 - v4) out[f"{metric}_signed_improvement"] = float(v4 - v5) if metric in lower_better else float(v5 - v4) rows.append(out) return pd.DataFrame(rows) def target_trial_diagnostics( arrays_by_split: dict[str, dict[str, np.ndarray]], action_by_split: dict[str, dict[str, np.ndarray]], out_dir: Path, stage0_dir: Path | None = None, ) -> dict[str, Any]: rows = [] for split, arrays in arrays_by_split.items(): medication_flags = medication_flags_from_stage0(arrays, stage0_dir) if stage0_dir is not None else None labels = build_target_trial_labels(arrays, (30.0, 60.0, 90.0), medication_flags=medication_flags, stage0_dir=stage0_dir) eligible = labels["target_trial_eligible"].astype(bool) for grace in (30, 60, 90): treated = labels[f"lai_initiation_within_{grace}d"].astype(bool) rate = float(treated[eligible].mean()) if eligible.any() else math.nan ok = bool(math.isfinite(rate) and 0.005 <= rate <= 0.995) rows.append( { "split": split, "strategy": "lai_initiation_vs_continued_oral", "grace_days": float(grace), "n_eligible": int(eligible.sum()), "n_treated": int((treated & eligible).sum()), "treatment_rate": rate, "liberal_positivity_ok": ok, } ) pos = pd.DataFrame(rows) causal_dir = out_dir / "05_target_trial" causal_dir.mkdir(parents=True, exist_ok=True) pos.to_csv(causal_dir / "causal_positivity.csv", index=False) not_authorized = pd.DataFrame( [ { "estimator": estimator, "status": "not_run", "reason": "Target-trial effect estimation is blocked until positivity/overlap and pre-action confounding balance pass. This run reports diagnostics only.", } for estimator in ("ipw", "gformula", "aipw_tmle") ] ) for name in ("ipw_effect_estimates.csv", "gformula_effect_estimates.csv", "aipw_tmle_effect_estimates.csv", "risk_curves_by_strategy.csv"): not_authorized.to_csv(causal_dir / name, index=False) summary = { "status": "diagnostics_only", "safe_claim": "target-trial-augmented sensitivity analysis only after positivity and balance gates pass", "positivity_failures": pos.loc[~pos["liberal_positivity_ok"]].to_dict(orient="records"), } (causal_dir / "target_trial_go_no_go.json").write_text(json.dumps(json_ready(summary), ensure_ascii=False, indent=2), encoding="utf-8") return summary def write_protocol_files(out_dir: Path, ontology: dict[str, Any], tensor_metadata: dict[str, Any], cat_value_vocab: dict[str, int]) -> None: protocol_dir = out_dir / "00_protocol" protocol_dir.mkdir(parents=True, exist_ok=True) pd.DataFrame(ontology["slots"]).to_csv(protocol_dir / "action_ontology.csv", index=False) pd.DataFrame(action_timing_audit_rows(ontology, tensor_metadata, cat_value_vocab)).to_csv(protocol_dir / "action_timing_audit.csv", index=False) target_protocol = { "name": TARGET_TRIAL_DEFAULTS["name"], "eligibility": "oral-treated active nonterminal patient-windows without prior/current LAI", "time_zero": "current visit/action landmark after contact-state construction", "strategies": ["LAI initiation within grace period", "continued oral/no LAI within grace period"], "grace_days_primary": TARGET_TRIAL_DEFAULTS["primary_grace_days"], "followup_days_primary": TARGET_TRIAL_DEFAULTS["primary_followup_days"], "claim_boundary": "diagnostic and sensitivity-analysis protocol; no automatic individualized counterfactual or clinical recommendation claim", } (protocol_dir / "target_trial_protocol.json").write_text(json.dumps(target_protocol, ensure_ascii=False, indent=2), encoding="utf-8") estimand = { "primary_estimand": TARGET_TRIAL_DEFAULTS["primary_estimand"], "allowed_language": ["observed-action-conditioned dynamics", "target-trial-augmented sensitivity analysis"], "prohibited_language": ["individualized counterfactual futures", "clinical treatment policy recommendation"], } (protocol_dir / "estimand_sheet.json").write_text(json.dumps(estimand, ensure_ascii=False, indent=2), encoding="utf-8") def main() -> None: parser = argparse.ArgumentParser(description="Run SCTM-v5 full observed-action evaluation with v4 comparison and target-trial boundary diagnostics.") parser.add_argument("--v5-checkpoint", type=Path, required=True) parser.add_argument("--v4-checkpoint", type=Path, default=None) parser.add_argument("--tensor-dir", type=Path, required=True) parser.add_argument("--stage0-dir", type=Path, required=True) parser.add_argument("--action-dir", type=Path, default=None) parser.add_argument("--service-prior-file", default="service_state_transitions_train.json") parser.add_argument("--out-dir", type=Path, required=True) parser.add_argument("--splits", default="val,test") parser.add_argument("--batch-size", type=int, default=128) parser.add_argument("--device", default="cuda") parser.add_argument("--precision", choices=["bf16", "fp16", "fp32"], default="bf16") parser.add_argument("--max-windows", type=int, default=0) parser.add_argument("--horizons", default="30,90,180,365") parser.add_argument("--rollout-steps", type=int, default=10) parser.add_argument("--rollout-max-windows", type=int, default=2048) parser.add_argument("--seed", type=int, default=20260525) args = parser.parse_args() torch.manual_seed(args.seed) np.random.seed(args.seed) args.out_dir.mkdir(parents=True, exist_ok=True) device = torch.device(args.device if args.device == "cpu" or torch.cuda.is_available() else "cpu") meta = load_json(args.tensor_dir / "tensor_metadata.json") vocab = load_json(args.tensor_dir / "cat_value_vocab.json") ordinal_direction = load_json(args.stage0_dir / "ordinal_direction_table.json") ontology = build_action_ontology(meta, vocab) write_protocol_files(args.out_dir, ontology, meta, vocab) v5_model, v5_payload, v5_audit = load_v5_model(args.v5_checkpoint, args.tensor_dir, args.stage0_dir, args.service_prior_file, device) v4_model = None v4_audit = None if args.v4_checkpoint is not None: v4_model, v4_payload, v4_audit = load_v4p4_model(args.v4_checkpoint, args.tensor_dir, args.stage0_dir, args.service_prior_file, device) split_names = [x.strip() for x in args.splits.split(",") if x.strip()] horizons = [float(x) for x in args.horizons.split(",") if x.strip()] all_outputs: dict[str, list[pd.DataFrame]] = defaultdict(list) arrays_by_split: dict[str, dict[str, np.ndarray]] = {} action_by_split: dict[str, dict[str, np.ndarray]] = {} source_audit_written = False for split in split_names: arrays = load_npz_to_memory(resolve_split_path(args.tensor_dir, split)) if args.max_windows > 0: arrays = {k: v[: args.max_windows] for k, v in arrays.items()} if args.action_dir is not None and (args.action_dir / f"v5_{split}_action_tensors.npz").exists(): with np.load(args.action_dir / f"v5_{split}_action_tensors.npz", allow_pickle=False) as z: action_arrays = { k: (z[k][: args.max_windows] if args.max_windows > 0 and z[k].ndim > 0 else z[k]) for k in z.files } validate_action_arrays(action_arrays, arrays, ontology, source=args.action_dir / f"v5_{split}_action_tensors.npz") validate_medication_flags(action_arrays, medication_flags_from_stage0(arrays, args.stage0_dir), source=args.action_dir / f"v5_{split}_action_tensors.npz") else: medication_flags = medication_flags_from_stage0(arrays, args.stage0_dir) action_arrays = build_action_arrays(arrays, meta, ordinal_direction, vocab, ontology, medication_flags=medication_flags) validate_action_arrays(action_arrays, arrays, ontology, source="built_in_memory") validate_medication_flags(action_arrays, medication_flags, source="built_in_memory") arrays_by_split[split] = arrays action_by_split[split] = action_arrays if not source_audit_written: audit = source_audit(v5_model, arrays, action_arrays, device) audit_dir = args.out_dir / "01_source_audits" audit_dir.mkdir(parents=True, exist_ok=True) (audit_dir / "action_antileakage_smoke.json").write_text(json.dumps(json_ready(audit), ensure_ascii=False, indent=2), encoding="utf-8") (audit_dir / "information_boundary_audit.json").write_text(json.dumps(json_ready(audit), ensure_ascii=False, indent=2), encoding="utf-8") (audit_dir / "export_source_audit.json").write_text(json.dumps(json_ready({"status": "passed" if audit.get("passed") else "failed", **audit}), ensure_ascii=False, indent=2), encoding="utf-8") source_audit_written = True rows = evaluate_split( split=split, arrays=arrays, action_arrays=action_arrays, ontology=ontology, v5_model=v5_model, v4_model=v4_model, device=device, precision=args.precision, batch_size=args.batch_size, max_windows=args.max_windows, horizons=horizons, rollout_steps=args.rollout_steps, rollout_max_windows=args.rollout_max_windows, ) groups = { "process_nll": ["split", "model", "metric"], "grammar_active_state": ["split", "model"], "obs_missingness": ["split", "model"], "grammar_cat_fields": ["split", "model"], "grammar_numeric": ["split", "model"], "grammar_ordinal": ["split", "model"], "grammar_events": ["split", "model", "event_index"], "pwe_horizon_risk": ["split", "model", "cause_id", "cause_name", "horizon_days"], "policy_discrimination": ["split", "slot_id", "slot_name"], "policy_prevalence": ["split", "slot_id", "slot_name", "value_id", "value_label"], "policy_overlap": ["split", "slot_id", "slot_name"], "policy_by_era": ["split", "slot_id", "slot_name", "era_id"], "policy_by_state": ["split", "slot_id", "slot_name", "service_state"], "rollout_legality": ["split", "model", "rollout_steps", "violation"], } for name, raw_rows in rows.items(): all_outputs[name].append(aggregate_frame(raw_rows, groups.get(name, ["split"]))) if device.type == "cuda": torch.cuda.empty_cache() frames = {name: pd.concat(parts, ignore_index=True) if parts else pd.DataFrame() for name, parts in all_outputs.items()} out_paths: dict[str, str] = {} routing = { "process_nll": "04_action_gain", "grammar_active_state": "02_v4_comparison", "obs_missingness": "02_v4_comparison", "grammar_cat_fields": "02_v4_comparison", "grammar_numeric": "02_v4_comparison", "grammar_ordinal": "02_v4_comparison", "grammar_events": "02_v4_comparison", "pwe_horizon_risk": "02_v4_comparison", "policy_discrimination": "03_action_policy", "policy_prevalence": "03_action_policy", "policy_overlap": "03_action_policy", "policy_by_era": "03_action_policy", "policy_by_state": "03_action_policy", "rollout_legality": "02_v4_comparison", } filenames = { "process_nll": "action_gain_next_process_nll.csv", "grammar_active_state": "grammar_active_state.csv", "obs_missingness": "obs_missingness_macro_f1.csv", "grammar_cat_fields": "grammar_cat_field_accuracy.csv", "grammar_numeric": "grammar_numeric.csv", "grammar_ordinal": "grammar_ordinal.csv", "grammar_events": "grammar_event_metrics.csv", "pwe_horizon_risk": "pwe_horizon_risk_metrics.csv", "policy_discrimination": "policy_propensity_discrimination.csv", "policy_prevalence": "policy_action_prevalence.csv", "policy_overlap": "policy_overlap_histograms.csv", "policy_by_era": "policy_action_prevalence_by_era.csv", "policy_by_state": "policy_action_by_state.csv", "rollout_legality": "rollout_legality.csv", } for name, frame in frames.items(): subdir = args.out_dir / routing.get(name, "") subdir.mkdir(parents=True, exist_ok=True) path = subdir / filenames.get(name, f"{name}.csv") frame.to_csv(path, index=False) out_paths[name] = str(path) comparison_dir = args.out_dir / "02_v4_comparison" comparison_dir.mkdir(parents=True, exist_ok=True) compare_specs = [ ("pwe_horizon_risk", ["split", "cause_id", "cause_name", "horizon_days"], ["auc", "average_precision", "brier", "ece", "ici"], "risk_v5_minus_v4.csv"), ("process_nll", ["split", "metric"], ["value"], "pwe_v5_minus_v4.csv"), ("obs_missingness", ["split"], ["accuracy", "macro_f1", "multiclass_brier", "confidence_ece"], "observation_v5_minus_v4.csv"), ("grammar_active_state", ["split"], ["accuracy", "macro_f1", "cross_entropy"], "grammar_v5_minus_v4.csv"), ("rollout_legality", ["split", "rollout_steps", "violation"], ["rate_per_generated_position"], "rollout_v5_minus_v4.csv"), ] for frame_name, index_cols, metric_cols, filename in compare_specs: comp = compare_models(frames.get(frame_name, pd.DataFrame()), index_cols, metric_cols) comp.to_csv(comparison_dir / filename, index=False) out_paths[filename.removesuffix(".csv")] = str(comparison_dir / filename) action_gain_dir = args.out_dir / "04_action_gain" action_gain_dir.mkdir(parents=True, exist_ok=True) for frame_name, index_cols, metric_cols, filename in [ ("pwe_horizon_risk", ["split", "cause_id", "cause_name", "horizon_days"], ["auc", "average_precision", "brier", "ece"], "action_gain_pwe_horizon_risk.csv"), ("grammar_events", ["split", "event_index"], ["auc", "average_precision", "brier"], "action_gain_event_metrics.csv"), ("obs_missingness", ["split"], ["accuracy", "macro_f1", "multiclass_brier", "confidence_ece"], "action_gain_observation.csv"), ]: comp = compare_models(frames.get(frame_name, pd.DataFrame()), index_cols, metric_cols, baseline=V5_BASE, candidate=OBSERVED_ACTION) comp.to_csv(action_gain_dir / filename, index=False) out_paths[filename.removesuffix(".csv")] = str(action_gain_dir / filename) pd.DataFrame( [ {"ablation": "all_actions", "status": "evaluated", "comparison": "v5_action"}, {"ablation": "no_action", "status": "evaluated", "comparison": "v5_base_post"}, {"ablation": "group_masked", "status": "not_run", "reason": "requires retracing action tensors with masked blocks; not mixed into primary evaluation"}, ] ).to_csv(action_gain_dir / "action_group_ablation.csv", index=False) policy_dir = args.out_dir / "03_action_policy" policy_dir.mkdir(parents=True, exist_ok=True) frames.get("policy_overlap", pd.DataFrame()).to_csv(policy_dir / "policy_effective_sample_size.csv", index=False) frames.get("policy_overlap", pd.DataFrame()).to_csv(policy_dir / "policy_weight_diagnostics.csv", index=False) pd.DataFrame([{"status": "not_run", "reason": "SMD balance requires a finalized target-trial confounder table; positivity diagnostics are reported separately."}]).to_csv(policy_dir / "policy_covariate_balance.csv", index=False) frames.get("policy_discrimination", pd.DataFrame()).to_csv(policy_dir / "policy_propensity_calibration.csv", index=False) target_summary = target_trial_diagnostics(arrays_by_split, action_by_split, args.out_dir, stage0_dir=args.stage0_dir) sensitivity_dir = args.out_dir / "06_sensitivity" sensitivity_dir.mkdir(parents=True, exist_ok=True) pd.DataFrame(target_summary.get("positivity_failures", [])).to_csv(sensitivity_dir / "sensitivity_off_support.csv", index=False) pd.DataFrame([{"status": "not_run", "reason": "Strategy rollout sensitivity is blocked until target-trial positivity and support gates pass."}]).to_csv(sensitivity_dir / "sensitivity_v5_effects.csv", index=False) pd.DataFrame([{"status": "not_run", "reason": "Uncertainty intervals require a post-evaluation strategy rollout design."}]).to_csv(sensitivity_dir / "sensitivity_rollout_uncertainty.csv", index=False) summary = { "status": "completed", "v5_checkpoint": str(args.v5_checkpoint), "v5_checkpoint_best_step": v5_payload.get("best_step", v5_payload.get("step")), "v4_checkpoint": str(args.v4_checkpoint) if args.v4_checkpoint else None, "device": str(device), "precision": args.precision, "splits": split_names, "max_windows": int(args.max_windows), "horizons": horizons, "rollout_steps": int(args.rollout_steps), "rollout_max_windows": int(args.rollout_max_windows), "claim_boundary": "Observed-action world-model metrics are primary; target-trial outputs are diagnostics/sensitivity only unless positivity and balance gates pass.", "paths": out_paths, "v5_audit": v5_audit, "v4_audit": v4_audit, "target_trial_summary": target_summary, } (args.out_dir / "v5_full_evaluation_summary.json").write_text(json.dumps(json_ready(summary), ensure_ascii=False, indent=2), encoding="utf-8") print(json.dumps({"status": "completed", "out_dir": str(args.out_dir), "summary": str(args.out_dir / "v5_full_evaluation_summary.json")}, ensure_ascii=False), flush=True) if __name__ == "__main__": main()