| 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 |
| from build_v5_action_tensors import ( |
| 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 |
| from export_v4p4_visit_outputs import load_model as load_v4p4_model |
| from loss_v4p3 import piecewise_exp_competing_risk_nll, pwe_target_from_terminal, v4_missing_targets |
| from model_v4p4 import pwe_closed_form_cif |
| from model_v5 import SCTMv5, SCTMv5Config |
| from smoke_v3p4_architecture import build_field_value_mask, load_service_prior |
| from train_v4p4_cloud import autocast_context, batch_from_indices, load_npz_to_memory |
|
|
| try: |
| from sklearn.metrics import average_precision_score, roc_auc_score |
| except Exception: |
| 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) |
| |
| |
| 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() |
|
|