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