agoniii97's picture
Normalize datetime precision for HF P1 tensor build
ae6d94c verified
Raw
History Blame Contribute Delete
48.9 kB
from __future__ import annotations
import argparse
import json
import math
import sys
from collections import defaultdict
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 build_action_ontology, load_json # noqa: E402
from build_v5_action_tensors import ( # noqa: E402
build_action_arrays,
medication_flags_from_stage0,
resolve_split_path,
validate_medication_flags,
)
from evaluate_v5_full import ( # noqa: E402
OBSERVED_ACTION,
V5_BASE,
binary_auc_ap,
config_from_payload,
json_ready,
load_v5_model,
multiclass_confusion,
quadratic_weighted_kappa,
safe_mean,
tensorize_actions,
validate_action_arrays,
)
from loss_v4p3 import ( # noqa: E402
exposure_per_pwe_bin,
pwe_bin_indices,
pwe_target_from_terminal,
v4_missing_targets,
)
from model_v4p4 import pwe_closed_form_cif # noqa: E402
from scan_target_trial_support_v5 import era_tokens_np, future_event_within, unique_landmark_filter # noqa: E402
from train_v4p4_cloud import autocast_context, batch_from_indices, load_npz_to_memory # noqa: E402
try:
from scipy.spatial.distance import jensenshannon
from scipy.stats import wasserstein_distance
except Exception: # pragma: no cover
jensenshannon = None
wasserstein_distance = None
try:
from sklearn.isotonic import IsotonicRegression
from sklearn.linear_model import LogisticRegression
from sklearn.metrics import average_precision_score, roc_auc_score
except Exception: # pragma: no cover
IsotonicRegression = None
LogisticRegression = None
average_precision_score = None
roc_auc_score = None
CAUSE_NAMES = {0: "next_contact", 1: "death", 2: "disengagement"}
EVENT_HORIZON_ENDPOINTS = {
0: "primary_referral_relapse",
3: "service_escalation",
4: "clinical_deterioration",
5: "high_acuity_state",
}
RISK_THRESHOLD_GRID = {
"death": [0.005, 0.01, 0.02, 0.03, 0.05],
"primary_referral_relapse": [0.01, 0.03, 0.05, 0.10, 0.15],
"disengagement": [0.01, 0.02, 0.05, 0.08, 0.10],
"clinical_deterioration": [0.30, 0.40, 0.50, 0.55, 0.60],
"default": [0.01, 0.03, 0.05, 0.10, 0.15],
}
PRIMARY_BOOTSTRAP_ENDPOINTS = {"death", "disengagement", "next_contact", "primary_referral_relapse", "clinical_deterioration"}
PRIMARY_BOOTSTRAP_HORIZONS = {90.0, 365.0}
def clip_prob(p: np.ndarray, eps: float = 1.0e-7) -> np.ndarray:
return np.clip(np.asarray(p, dtype=np.float64).reshape(-1), eps, 1.0 - eps)
def weighted_mean(x: np.ndarray, w: np.ndarray | None = None) -> float:
x = np.asarray(x, dtype=np.float64).reshape(-1)
if w is None:
return float(np.mean(x)) if x.size else math.nan
w = np.asarray(w, dtype=np.float64).reshape(-1)
ok = np.isfinite(x) & np.isfinite(w) & (w > 0)
den = float(w[ok].sum())
return float(np.sum(x[ok] * w[ok]) / den) if den > 0 else math.nan
def integrated_calibration_index(y: np.ndarray, p: np.ndarray, sample_weight: np.ndarray | None = None) -> float:
y = np.asarray(y, dtype=np.int8).reshape(-1)
p = clip_prob(p)
w = None if sample_weight is None else np.asarray(sample_weight, dtype=np.float64).reshape(-1)
if w is not None:
ok = np.isfinite(w) & (w > 0)
y = y[ok]
p = p[ok]
w = w[ok]
if IsotonicRegression is None or y.size == 0 or np.unique(y).size < 2:
return math.nan
try:
model = IsotonicRegression(y_min=0.0, y_max=1.0, out_of_bounds="clip")
calibrated = model.fit_transform(p, y, sample_weight=w)
return weighted_mean(np.abs(calibrated - p), w)
except Exception:
return math.nan
def weighted_binary_metrics(
y: np.ndarray,
p: np.ndarray,
sample_weight: np.ndarray | None = None,
*,
include_calibration_model: bool = True,
) -> dict[str, float]:
y = np.asarray(y, dtype=np.int8).reshape(-1)
p = clip_prob(p)
w = None if sample_weight is None else np.asarray(sample_weight, dtype=np.float64).reshape(-1)
if y.size == 0:
return {
"auc": math.nan,
"average_precision": math.nan,
"brier": math.nan,
"ece": math.nan,
"ici": math.nan,
"mean_predicted": math.nan,
"observed_rate": math.nan,
"calibration_intercept": math.nan,
"calibration_slope": math.nan,
}
if w is not None:
ok = np.isfinite(w) & (w > 0)
y = y[ok]
p = p[ok]
w = w[ok]
if y.size == 0:
return weighted_binary_metrics(y, p, None)
auc = math.nan
ap = math.nan
if np.unique(y).size >= 2:
try:
auc = float(roc_auc_score(y, p, sample_weight=w)) if roc_auc_score is not None else math.nan
except Exception:
auc = math.nan
try:
ap = float(average_precision_score(y, p, sample_weight=w)) if average_precision_score is not None else math.nan
except Exception:
ap = math.nan
brier = weighted_mean((p - y) ** 2, w)
ici = integrated_calibration_index(y, p, w) if include_calibration_model else math.nan
order = np.argsort(p)
bins = np.array_split(order, min(10, max(1, y.size)))
ece_parts = []
ece_weights = []
for idx in bins:
if idx.size == 0:
continue
ww = None if w is None else w[idx]
ece_parts.append(abs(weighted_mean(p[idx], ww) - weighted_mean(y[idx], ww)))
ece_weights.append(float(idx.size) if w is None else float(w[idx].sum()))
den = float(sum(ece_weights))
ece = float(sum(v * wt for v, wt in zip(ece_parts, ece_weights)) / den) if den > 0 else math.nan
if include_calibration_model:
intercept, slope = calibration_intercept_slope(y, p, w)
else:
intercept, slope = math.nan, math.nan
return {
"auc": auc,
"average_precision": ap,
"brier": brier,
"ece": ece,
"ici": ici,
"mean_predicted": weighted_mean(p, w),
"observed_rate": weighted_mean(y, w),
"calibration_intercept": intercept,
"calibration_slope": slope,
}
def calibration_intercept_slope(y: np.ndarray, p: np.ndarray, sample_weight: np.ndarray | None = None) -> tuple[float, float]:
if LogisticRegression is None or y.size == 0 or np.unique(y).size < 2:
return math.nan, math.nan
x = np.log(clip_prob(p) / (1.0 - clip_prob(p))).reshape(-1, 1)
try:
clf = LogisticRegression(C=1.0e6, solver="lbfgs", max_iter=500)
clf.fit(x, y.astype(int), sample_weight=sample_weight)
return float(clf.intercept_[0]), float(clf.coef_[0, 0])
except Exception:
return math.nan, math.nan
def calibration_curve_rows(
*,
split: str,
model: str,
pipeline: str,
endpoint: str,
horizon_days: float,
y: np.ndarray,
p: np.ndarray,
n_bins: int = 10,
) -> list[dict[str, Any]]:
y = np.asarray(y, dtype=np.int8).reshape(-1)
p = clip_prob(p)
if y.size == 0:
return []
rows = []
for bin_idx, idx in enumerate(np.array_split(np.argsort(p), min(n_bins, max(1, y.size))), start=1):
if idx.size == 0:
continue
rows.append(
{
"split": split,
"model": model,
"pipeline": pipeline,
"endpoint": endpoint,
"horizon_days": float(horizon_days),
"calibration_bin": int(bin_idx),
"n": int(idx.size),
"events": int(y[idx].sum()),
"mean_predicted": float(p[idx].mean()),
"observed_rate": float(y[idx].mean()),
"predicted_min": float(p[idx].min()),
"predicted_max": float(p[idx].max()),
}
)
return rows
def decision_curve_rows(
*,
split: str,
model: str,
pipeline: str,
endpoint: str,
horizon_days: float,
y: np.ndarray,
p: np.ndarray,
patient_ids: np.ndarray,
thresholds: list[float],
bootstrap_clusters: int,
seed: int,
) -> list[dict[str, Any]]:
y = np.asarray(y, dtype=np.int8).reshape(-1)
p = clip_prob(p)
base_rows = []
boot = bootstrap_decision_curves(y, p, patient_ids, thresholds, bootstrap_clusters, seed)
for threshold in thresholds:
nb = net_benefit(y, p, threshold)
row = {
"split": split,
"model": model,
"pipeline": pipeline,
"endpoint": endpoint,
"horizon_days": float(horizon_days),
"threshold": float(threshold),
"net_benefit": nb,
"bootstrap_n": int(bootstrap_clusters),
}
row.update(boot.get(float(threshold), {}))
base_rows.append(row)
return base_rows
def net_benefit(y: np.ndarray, p: np.ndarray, threshold: float, sample_weight: np.ndarray | None = None) -> float:
if y.size == 0 or threshold <= 0.0 or threshold >= 1.0:
return math.nan
w = np.ones(y.size, dtype=np.float64) if sample_weight is None else np.asarray(sample_weight, dtype=np.float64).reshape(-1)
pred_pos = p >= float(threshold)
den = float(w.sum())
if den <= 0:
return math.nan
tp = float(np.sum(w * pred_pos * (y == 1))) / den
fp = float(np.sum(w * pred_pos * (y == 0))) / den
return float(tp - fp * threshold / (1.0 - threshold))
def bootstrap_decision_curves(
y: np.ndarray,
p: np.ndarray,
patient_ids: np.ndarray,
thresholds: list[float],
n_bootstrap: int,
seed: int,
) -> dict[float, dict[str, float]]:
if n_bootstrap <= 0 or y.size == 0:
return {}
rng = np.random.default_rng(seed)
clusters, inv = np.unique(patient_ids.astype(str), return_inverse=True)
values = {float(t): [] for t in thresholds}
for _ in range(n_bootstrap):
counts = rng.multinomial(clusters.size, np.full(clusters.size, 1.0 / clusters.size))
w = counts[inv].astype(np.float64)
if w.sum() <= 0:
continue
for t in thresholds:
values[float(t)].append(net_benefit(y, p, float(t), sample_weight=w))
return {
t: {
"net_benefit_ci_low": float(np.nanquantile(vals, 0.025)) if vals else math.nan,
"net_benefit_ci_high": float(np.nanquantile(vals, 0.975)) if vals else math.nan,
}
for t, vals in values.items()
}
def bootstrap_metric_ci(
y: np.ndarray,
p: np.ndarray,
patient_ids: np.ndarray,
n_bootstrap: int,
seed: int,
) -> dict[str, float]:
if n_bootstrap <= 0 or y.size == 0:
return {}
rng = np.random.default_rng(seed)
clusters, inv = np.unique(patient_ids.astype(str), return_inverse=True)
tracked = defaultdict(list)
for _ in range(n_bootstrap):
counts = rng.multinomial(clusters.size, np.full(clusters.size, 1.0 / clusters.size))
w = counts[inv].astype(np.float64)
metrics = weighted_binary_metrics(y, p, sample_weight=w, include_calibration_model=False)
for key, value in metrics.items():
if math.isfinite(float(value)):
tracked[key].append(float(value))
out: dict[str, float] = {"bootstrap_n": int(n_bootstrap), "bootstrap_clusters": int(clusters.size)}
for key, vals in tracked.items():
arr = np.asarray(vals, dtype=np.float64)
out[f"{key}_ci_low"] = float(np.nanquantile(arr, 0.025)) if arr.size else math.nan
out[f"{key}_ci_high"] = float(np.nanquantile(arr, 0.975)) if arr.size else math.nan
return out
def bootstrap_reps_for_metric(endpoint: str, horizon_days: float, requested: int) -> int:
if endpoint in PRIMARY_BOOTSTRAP_ENDPOINTS and float(horizon_days) in PRIMARY_BOOTSTRAP_HORIZONS:
return int(requested)
return 0
def patient_matrix(arrays: dict[str, np.ndarray], idx: np.ndarray, seq_len: int) -> np.ndarray:
if "_patient_code" in arrays:
patient_ids = arrays["_patient_code"][idx].astype(np.int64)
elif "patient_ids" in arrays:
patient_ids = pd.factorize(arrays["patient_ids"].astype(str), sort=True)[0][idx].astype(np.int64)
else:
patient_ids = idx.astype(np.int64)
return np.repeat(patient_ids[:, None], seq_len, axis=1)
def prediction_export_slug(split: str, model: str, pipeline: str, endpoint: str, horizon_days: float) -> str:
horizon_text = str(int(horizon_days)) if float(horizon_days).is_integer() else str(horizon_days).replace(".", "p")
parts = [split, model, pipeline, endpoint, f"{horizon_text}d"]
return "__".join(str(part).replace("/", "_").replace(" ", "_") for part in parts)
def write_prediction_export(
export_dir: Path,
key: tuple[str, str, str, str, float],
parts: dict[str, list[np.ndarray]],
*,
y: np.ndarray,
p: np.ndarray,
pid: np.ndarray,
) -> dict[str, Any]:
split, model_name, pipeline, endpoint, horizon = key
export_dir.mkdir(parents=True, exist_ok=True)
payload: dict[str, Any] = {
"y": y.astype(np.int8),
"p": p.astype(np.float32),
"patient_code": pid.astype(np.int64, copy=False),
"row_index": np.concatenate(parts["row"]).astype(np.int64) if parts.get("row") else np.asarray([], dtype=np.int64),
"position": np.concatenate(parts["pos"]).astype(np.int16) if parts.get("pos") else np.asarray([], dtype=np.int16),
"visit_year": np.concatenate(parts["year"]).astype(np.int16) if parts.get("year") else np.asarray([], dtype=np.int16),
"time_since_start_days": np.concatenate(parts["time"]).astype(np.float32) if parts.get("time") else np.asarray([], dtype=np.float32),
}
out_path = export_dir / f"{prediction_export_slug(split, model_name, pipeline, endpoint, horizon)}.npz"
np.savez_compressed(out_path, **payload)
meta = {
"split": split,
"model": model_name,
"pipeline": pipeline,
"endpoint": endpoint,
"horizon_days": float(horizon),
"n": int(y.size),
"events": int(y.sum()) if y.size else 0,
"path": str(out_path),
}
return meta
def aggregate_grammar_rows(rows: list[dict[str, Any]]) -> pd.DataFrame:
"""Aggregate batch-level grammar rows without turning n into a batch mean."""
df = pd.DataFrame(rows)
if df.empty:
return df
out: list[dict[str, Any]] = []
metric_cols = ["accuracy", "macro_f1", "mae", "rmse", "quadratic_weighted_kappa"]
for keys, group in df.groupby(["split", "model", "family"], dropna=False):
n = pd.to_numeric(group["n"], errors="coerce").fillna(0.0).to_numpy(dtype=np.float64)
den = float(n.sum())
row: dict[str, Any] = {"split": keys[0], "model": keys[1], "family": keys[2], "n": den}
for col in metric_cols:
if col not in group.columns:
continue
values = pd.to_numeric(group[col], errors="coerce").to_numpy(dtype=np.float64)
ok = np.isfinite(values) & (n > 0)
if not ok.any():
continue
if col == "rmse":
row[col] = float(math.sqrt(np.sum(n[ok] * values[ok] ** 2) / np.sum(n[ok])))
else:
row[col] = float(np.sum(n[ok] * values[ok]) / np.sum(n[ok]))
row["aggregation_note"] = "weighted_by_batch_n; exact full-eval grammar tables remain preferred for confusion-matrix metrics"
out.append(row)
return pd.DataFrame(out)
def pwe_nll_elements(log_lambda_kj: torch.Tensor, batch: dict[str, torch.Tensor]) -> tuple[np.ndarray, np.ndarray]:
cause, censored, valid, _ = pwe_target_from_terminal(batch)
event_time_days = torch.expm1(batch["delta_t_next_log"][:, :-1]).clamp(min=1.0e-6)
lambda_kj = F.softplus(log_lambda_kj[:, :-1, :, :]) + 1.0e-8
lambda_j = lambda_kj.sum(dim=2)
exposure = exposure_per_pwe_bin(event_time_days)
cum_hazard = lambda_j * exposure
bin_target = pwe_bin_indices(event_time_days).clamp(max=log_lambda_kj.shape[-1] - 1)
log_survival_total = -cum_hazard.sum(dim=-1)
lambda_kj_star = lambda_kj.gather(3, bin_target[:, :, None, None].expand(-1, -1, lambda_kj.shape[2], 1)).squeeze(-1)
lambda_k_star = lambda_kj_star.gather(2, cause.clamp(min=0).unsqueeze(-1)).squeeze(-1)
event_loglik = log_survival_total + torch.log(lambda_k_star.clamp(min=1.0e-8))
loglik = torch.where(censored, log_survival_total, event_loglik)
return (-loglik).detach().cpu().numpy(), valid.detach().cpu().numpy().astype(bool)
def event_horizon_from_teacher_forced(
event_prob: np.ndarray,
event_labels: np.ndarray,
valid: np.ndarray,
times: np.ndarray,
event_index: int,
horizon_days: float,
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
seq_len = valid.shape[1]
current_len = seq_len - 1
no_event = np.ones((valid.shape[0], current_len), dtype=np.float64)
y = np.zeros((valid.shape[0], current_len), dtype=bool)
any_future = np.zeros((valid.shape[0], current_len), dtype=bool)
for offset in range(1, seq_len):
cur_len = min(current_len, seq_len - offset)
if cur_len <= 0:
break
dt = times[:, offset : offset + cur_len] - times[:, :cur_len]
fut_ok = valid[:, offset : offset + cur_len] & (dt > 0.0) & (dt <= float(horizon_days))
p_step = event_prob[:, offset - 1 : offset - 1 + cur_len, int(event_index)]
no_event[:, :cur_len] *= np.where(fut_ok, 1.0 - p_step, 1.0)
y[:, :cur_len] |= fut_ok & event_labels[:, offset : offset + cur_len, int(event_index)].astype(bool)
any_future[:, :cur_len] |= fut_ok
last_time = np.where(valid, times, -np.inf).max(axis=1, keepdims=True)
followup = last_time - times[:, :current_len]
evaluable = valid[:, :current_len] & any_future & (y | (followup >= float(horizon_days)))
return y.astype(np.int8), np.clip(1.0 - no_event, 0.0, 1.0), evaluable
def subgroup_values(arrays: dict[str, np.ndarray], row_idx: np.ndarray, pos_idx: np.ndarray, meta: dict[str, Any]) -> dict[str, np.ndarray]:
out: dict[str, np.ndarray] = {}
years = arrays["visit_year"][row_idx, pos_idx].astype(int)
out["era_pre2019"] = np.where(years <= 2018, "pre_2019", "post_2019")
out["calendar_year"] = years.astype(str)
if "static_value_ids" in arrays:
static_cols = list(meta.get("static_cols", []))
static = arrays["static_value_ids"][row_idx]
for field in ("sex_token", "primary_diagnosis_token", "birth_cohort_token"):
if field in static_cols:
out[field] = np.asarray([f"{field}={int(x)}" for x in static[:, static_cols.index(field)]])
miss = arrays["missing_ids"][row_idx, pos_idx].astype(np.int64)
missing_burden = (miss != 0).mean(axis=1)
out["missingness_burden"] = np.where(missing_burden <= np.nanmedian(missing_burden), "low", "high")
return out
def split_audit_rows(tensor_dir: Path, splits: list[str]) -> list[dict[str, Any]]:
rows = []
seen: dict[str, set[str]] = {}
for split in splits:
arrays = load_npz_to_memory(resolve_split_path(tensor_dir, split))
valid = arrays["valid_mask"].astype(bool)
pids = arrays["patient_ids"].astype(str) if "patient_ids" in arrays else np.arange(valid.shape[0]).astype(str)
seen[split] = set(pids.tolist())
years = arrays["visit_year"][valid]
terminal = arrays["terminal_label"][valid].astype(int)
events = arrays["event_labels"][valid].astype(float)
rows.append(
{
"split": split,
"n_windows": int(valid.shape[0]),
"n_patients": int(len(seen[split])),
"n_valid_visits": int(valid.sum()),
"visit_year_min": int(np.nanmin(years)) if years.size else None,
"visit_year_max": int(np.nanmax(years)) if years.size else None,
"terminal_death_visits": int((terminal == 1).sum()),
"terminal_loss_migration_visits": int(((terminal == 2) | (terminal == 3)).sum()),
"event0_primary_referral_relapse": int(events[:, 0].sum()) if events.size else 0,
"event6_hospitalization_state_indicator": int(events[:, 6].sum()) if events.size else 0,
}
)
del arrays
for i, a in enumerate(splits):
for b in splits[i + 1 :]:
rows.append({"split": f"{a}__vs__{b}", "patient_overlap": int(len(seen[a] & seen[b]))})
return rows
def repeated_information_boundary_audit(
model: torch.nn.Module,
arrays: dict[str, np.ndarray],
action_arrays: dict[str, np.ndarray],
device: torch.device,
n_checks: int,
seed: int,
) -> dict[str, Any]:
rng = np.random.default_rng(seed)
n = min(int(arrays["valid_mask"].shape[0]), 4096)
checks = []
if n <= 0:
return {"checked": 0, "passed": False, "reason": "empty split"}
candidates = rng.choice(n, size=min(n_checks, n), replace=False)
for row in candidates:
valid = arrays["valid_mask"][row].astype(bool)
pos_candidates = np.flatnonzero(valid[:-1] & valid[1:])
if pos_candidates.size == 0:
continue
pos = int(rng.choice(pos_candidates))
batch = batch_from_indices(arrays, np.asarray([row]), device)
batch.update(tensorize_actions(action_arrays, np.asarray([row]), device))
with torch.inference_mode():
out0 = model(batch, rollout_steps=1, compute_pwe_diagnostics=False)
future = {k: v.clone() if torch.is_tensor(v) else v for k, v in batch.items()}
if pos + 1 < future["action_value_ids"].shape[1]:
future["action_value_ids"][0, pos + 1, 0] = (future["action_value_ids"][0, pos + 1, 0] + 1) % model.config.action_value_vocab_size
outf = model(future, rollout_steps=1, compute_pwe_diagnostics=False)
current = {k: v.clone() if torch.is_tensor(v) else v for k, v in batch.items()}
current["action_value_ids"][0, pos, 0] = (current["action_value_ids"][0, pos, 0] + 1) % model.config.action_value_vocab_size
outc = model(current, rollout_steps=1, compute_pwe_diagnostics=False)
future_delta = float((out0["pwe_log_lambda_action"][0, pos] - outf["pwe_log_lambda_action"][0, pos]).abs().max().cpu())
behavior_delta = float((out0["behavior_policy_logits"][0, pos] - outc["behavior_policy_logits"][0, pos]).abs().max().cpu())
action_context_delta = float((out0["action_context"][0, pos] - outc["action_context"][0, pos]).abs().max().cpu())
checks.append(
{
"row": int(row),
"position": int(pos),
"future_action_current_output_max_delta": future_delta,
"behavior_policy_current_action_max_delta": behavior_delta,
"action_context_current_action_max_delta": action_context_delta,
"passed": bool(future_delta < 1.0e-6 and behavior_delta < 1.0e-6 and action_context_delta > 1.0e-8),
}
)
return {
"checked": len(checks),
"passed": bool(checks and all(row["passed"] for row in checks)),
"max_future_action_current_output_delta": max((row["future_action_current_output_max_delta"] for row in checks), default=math.nan),
"max_behavior_policy_current_action_delta": max((row["behavior_policy_current_action_max_delta"] for row in checks), default=math.nan),
"min_action_context_current_action_delta": min((row["action_context_current_action_max_delta"] for row in checks), default=math.nan),
"checks": checks,
}
def jsd(p: np.ndarray, q: np.ndarray) -> float:
p = np.asarray(p, dtype=np.float64)
q = np.asarray(q, dtype=np.float64)
p = p / max(float(p.sum()), 1.0e-12)
q = q / max(float(q.sum()), 1.0e-12)
if jensenshannon is not None:
return float(jensenshannon(p, q, base=2.0) ** 2)
m = 0.5 * (p + q)
kl_pm = np.sum(np.where(p > 0, p * np.log2(p / np.clip(m, 1.0e-12, None)), 0.0))
kl_qm = np.sum(np.where(q > 0, q * np.log2(q / np.clip(m, 1.0e-12, None)), 0.0))
return float(0.5 * (kl_pm + kl_qm))
def rollout_fidelity_rows(
model: torch.nn.Module,
arrays: dict[str, np.ndarray],
action_arrays: dict[str, np.ndarray],
split: str,
device: torch.device,
batch_size: int,
rollout_steps: int,
rollout_max_windows: int,
seed: int,
) -> list[dict[str, Any]]:
torch.manual_seed(seed)
n = min(int(arrays["valid_mask"].shape[0]), int(rollout_max_windows)) if rollout_max_windows > 0 else int(arrays["valid_mask"].shape[0])
rows = []
service_obs = np.zeros(model.config.n_service_states, dtype=np.float64)
service_gen = np.zeros(model.config.n_service_states, dtype=np.float64)
event_obs = []
event_gen = []
dt_obs = []
dt_gen = []
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))
rollout = model.ancestral_rollout(batch, start_pos=0, steps=rollout_steps, deterministic=False, rao_blackwell_rare=True)
gen = rollout["generated_batch"]
obs_valid = batch["valid_mask"][:, 1 : rollout_steps + 1].detach().cpu().numpy().astype(bool)
gen_valid = gen["valid_mask"][:, 1 : rollout_steps + 1].detach().cpu().numpy().astype(bool)
obs_state = batch["service_state"][:, 1 : rollout_steps + 1].detach().cpu().numpy().astype(int)
gen_state = gen["service_state"][:, 1 : rollout_steps + 1].detach().cpu().numpy().astype(int)
if obs_valid.any():
service_obs += np.bincount(obs_state[obs_valid].clip(0, model.config.n_service_states - 1), minlength=model.config.n_service_states)
if gen_valid.any():
service_gen += np.bincount(gen_state[gen_valid].clip(0, model.config.n_service_states - 1), minlength=model.config.n_service_states)
if "event_labels" in gen:
event_obs.append(batch["event_labels"][:, 1 : rollout_steps + 1, :].detach().cpu().numpy()[obs_valid])
event_gen.append(gen["event_labels"][:, 1 : rollout_steps + 1, :].detach().cpu().numpy()[gen_valid])
obs_time = batch["time_since_start_days"][:, : rollout_steps + 1].detach().cpu().numpy()
gen_time = gen["time_since_start_days"][:, : rollout_steps + 1].detach().cpu().numpy()
dt_obs.append(np.diff(obs_time, axis=1)[obs_valid])
dt_gen.append(np.diff(gen_time, axis=1)[gen_valid])
ev_obs = np.concatenate(event_obs, axis=0) if event_obs else np.empty((0, model.config.n_events))
ev_gen = np.concatenate(event_gen, axis=0) if event_gen else np.empty((0, model.config.n_events))
dto = np.concatenate(dt_obs) if dt_obs else np.asarray([], dtype=np.float64)
dtg = np.concatenate(dt_gen) if dt_gen else np.asarray([], dtype=np.float64)
wdist = float(wasserstein_distance(dto, dtg)) if wasserstein_distance is not None and dto.size and dtg.size else math.nan
rows.append(
{
"split": split,
"model": OBSERVED_ACTION,
"rollout_steps": int(rollout_steps),
"generated_positions": int(service_gen.sum()),
"observed_positions": int(service_obs.sum()),
"state_occupancy_jsd": jsd(service_obs, service_gen),
"intervisit_time_wasserstein_days": wdist,
}
)
if ev_obs.size and ev_gen.size:
obs_rate = ev_obs.mean(axis=0)
gen_rate = ev_gen.mean(axis=0)
for event_idx in range(model.config.n_events):
rows.append(
{
"split": split,
"model": OBSERVED_ACTION,
"rollout_steps": int(rollout_steps),
"event_index": int(event_idx),
"observed_event_rate": float(obs_rate[event_idx]),
"generated_event_rate": float(gen_rate[event_idx]),
"absolute_event_rate_error": float(abs(obs_rate[event_idx] - gen_rate[event_idx])),
}
)
return rows
def load_actions_for_split(
*,
split: str,
arrays: dict[str, np.ndarray],
action_dir: Path | None,
meta: dict[str, Any],
ordinal_direction: dict[str, Any],
vocab: dict[str, int],
ontology: dict[str, Any],
stage0_dir: Path,
max_windows: int,
) -> dict[str, np.ndarray]:
if action_dir is not None and (action_dir / f"v5_{split}_action_tensors.npz").exists():
with np.load(action_dir / f"v5_{split}_action_tensors.npz", allow_pickle=False) as z:
out = {k: (z[k][:max_windows] if max_windows > 0 and z[k].ndim > 0 else z[k]) for k in z.files}
validate_action_arrays(out, arrays, ontology, source=action_dir / f"v5_{split}_action_tensors.npz")
validate_medication_flags(out, medication_flags_from_stage0(arrays, stage0_dir), source=action_dir / f"v5_{split}_action_tensors.npz")
return out
medication_flags = medication_flags_from_stage0(arrays, stage0_dir)
out = build_action_arrays(arrays, meta, ordinal_direction, vocab, ontology, medication_flags=medication_flags)
validate_action_arrays(out, arrays, ontology, source="built_in_memory")
validate_medication_flags(out, medication_flags, source="built_in_memory")
return out
def main() -> None:
parser = argparse.ArgumentParser(description="Top-journal supplemental SCTM-v5 evaluation layer.")
parser.add_argument("--v5-checkpoint", type=Path, required=True)
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("--split-audit-splits", default="train,val,test")
parser.add_argument("--batch-size", type=int, default=256)
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("--bootstrap-clusters", type=int, default=200)
parser.add_argument("--subgroup-bootstrap-clusters", type=int, default=0)
parser.add_argument("--rollout-steps", type=int, default=10)
parser.add_argument("--rollout-max-windows", type=int, default=4096)
parser.add_argument("--boundary-checks", type=int, default=32)
parser.add_argument("--seed", type=int, default=20260526)
parser.add_argument(
"--export-predictions",
action="store_true",
help="Write row-level y/p/patient/row/position exports for post-hoc utility, recalibration, patient-weighting, and paired-CI analyses.",
)
args = parser.parse_args()
torch.manual_seed(args.seed)
np.random.seed(args.seed)
args.out_dir.mkdir(parents=True, exist_ok=True)
audit_dir = args.out_dir / "00_audit"
risk_dir = args.out_dir / "01_risk_calibration_dca"
rollout_dir = args.out_dir / "02_rollout_fidelity"
subgroup_dir = args.out_dir / "03_subgroup_temporal"
for subdir in (audit_dir, risk_dir, rollout_dir, subgroup_dir):
subdir.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)
v5_model, payload, model_audit = load_v5_model(args.v5_checkpoint, args.tensor_dir, args.stage0_dir, args.service_prior_file, device)
cfg = v5_model.config
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()]
split_audit = split_audit_rows(args.tensor_dir, [x.strip() for x in args.split_audit_splits.split(",") if x.strip()])
pd.DataFrame(split_audit).to_csv(audit_dir / "patient_split_audit.csv", index=False)
pred_store: dict[tuple[str, str, str, str, float], dict[str, list[np.ndarray]]] = defaultdict(
lambda: {"y": [], "p": [], "pid": [], "row": [], "pos": [], "year": [], "time": []}
)
subgroup_store: dict[tuple[str, str, str, str, float, str, str], dict[str, list[np.ndarray]]] = defaultdict(lambda: {"y": [], "p": []})
process_rows = []
grammar_rows = []
rollout_rows = []
boundary_payloads = []
for split_idx, split in enumerate(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] if getattr(v, "shape", (0,))[0] == arrays["valid_mask"].shape[0] else v for k, v in arrays.items()}
if "patient_ids" in arrays:
arrays["_patient_code"] = pd.factorize(arrays["patient_ids"].astype(str), sort=True)[0].astype(np.int64)
actions = load_actions_for_split(
split=split,
arrays=arrays,
action_dir=args.action_dir,
meta=meta,
ordinal_direction=ordinal_direction,
vocab=vocab,
ontology=ontology,
stage0_dir=args.stage0_dir,
max_windows=args.max_windows,
)
boundary_payloads.append(
{
"split": split,
**repeated_information_boundary_audit(v5_model, arrays, actions, device, args.boundary_checks, args.seed + split_idx),
}
)
n = int(arrays["valid_mask"].shape[0])
for start in range(0, n, args.batch_size):
idx = np.arange(start, min(start + args.batch_size, n))
batch = batch_from_indices(arrays, idx, device)
batch.update(tensorize_actions(actions, idx, device))
with torch.inference_mode(), autocast_context(device, args.precision):
out = v5_model(batch, rollout_steps=1, compute_pwe_diagnostics=False)
pid_next = patient_matrix(arrays, idx, batch["valid_mask"].shape[1] - 1)
current_len = batch["valid_mask"].shape[1] - 1
row_grid = np.repeat(idx[:, None], current_len, axis=1)
pos_grid = np.repeat(np.arange(current_len, dtype=np.int16)[None, :], idx.size, axis=0)
year_grid = batch["visit_year"][:, :-1].detach().cpu().numpy().astype(np.int16)
time_grid = batch["time_since_start_days"][:, :-1].detach().cpu().numpy().astype(np.float32)
valid_np = (batch["valid_mask"][:, :-1] & batch["valid_mask"][:, 1:]).detach().cpu().numpy().astype(bool)
cause, censored, _, _ = pwe_target_from_terminal(batch)
cause_np = cause.detach().cpu().numpy()
cens_np = censored.detach().cpu().numpy().astype(bool)
days_np = torch.expm1(batch["delta_t_next_log"][:, :-1]).detach().cpu().numpy()
horizons_t = torch.tensor(horizons, device=device, dtype=torch.float32)
for model_name, log_key in ((OBSERVED_ACTION, "pwe_log_lambda_action"), (V5_BASE, "pwe_log_lambda_post")):
nll, nll_valid = pwe_nll_elements(out[log_key], batch)
process_rows.append(
{
"split": split,
"model": model_name,
"metric": "pwe_nll_unweighted",
"n": int(nll_valid.sum()),
"estimate": float(nll[nll_valid].mean()) if nll_valid.any() else math.nan,
}
)
pwe = pwe_closed_form_cif(out[log_key][:, :-1, :, :], horizons_t)["cif"].detach().cpu().float().numpy()
for h_idx, horizon in enumerate(horizons):
for cause_id, endpoint in CAUSE_NAMES.items():
evaluable = valid_np & ~(cens_np & (days_np <= float(horizon)))
target = ((~cens_np) & (cause_np == cause_id) & (days_np <= float(horizon))).astype(np.int8)
key = (split, model_name, "pwe_cif", endpoint, float(horizon))
pred_store[key]["y"].append(target[evaluable])
pred_store[key]["p"].append(pwe[:, :, h_idx, cause_id][evaluable])
pred_store[key]["pid"].append(pid_next[evaluable])
pred_store[key]["row"].append(row_grid[evaluable])
pred_store[key]["pos"].append(pos_grid[evaluable])
pred_store[key]["year"].append(year_grid[evaluable])
pred_store[key]["time"].append(time_grid[evaluable])
event_target = batch["event_labels"].detach().cpu().numpy().astype(np.int8)
valid_full = batch["valid_mask"].detach().cpu().numpy().astype(bool)
times = batch["time_since_start_days"].detach().cpu().numpy().astype(np.float64)
for model_name, event_key in ((OBSERVED_ACTION, "event_generation_logits_action"), (V5_BASE, "event_generation_logits")):
event_prob = torch.sigmoid(out[event_key].float()).detach().cpu().numpy()
for event_idx, endpoint in EVENT_HORIZON_ENDPOINTS.items():
for horizon in horizons:
y_h, p_h, evaluable = event_horizon_from_teacher_forced(event_prob, event_target, valid_full, times, event_idx, horizon)
key = (split, model_name, "event_teacher_forced_horizon", endpoint, float(horizon))
pred_store[key]["y"].append(y_h[evaluable])
pred_store[key]["p"].append(p_h[evaluable])
pred_store[key]["pid"].append(pid_next[evaluable])
pred_store[key]["row"].append(row_grid[evaluable])
pred_store[key]["pos"].append(pos_grid[evaluable])
pred_store[key]["year"].append(year_grid[evaluable])
pred_store[key]["time"].append(time_grid[evaluable])
if horizon in (90.0, 365.0) and model_name == OBSERVED_ACTION:
row_idx, pos_idx = np.nonzero(evaluable)
subgroups = subgroup_values(arrays, idx[row_idx], pos_idx, meta)
for axis, values in subgroups.items():
for value in np.unique(values):
mask = values == value
skey = (split, model_name, "event_teacher_forced_horizon", endpoint, float(horizon), axis, str(value))
subgroup_store[skey]["y"].append(y_h[evaluable][mask])
subgroup_store[skey]["p"].append(p_h[evaluable][mask])
miss_target = v4_missing_targets(out, batch, cfg)
next_contact = (valid_np & (batch["service_state"][:, 1:].detach().cpu().numpy() < cfg.n_active_states))
miss_probs = torch.softmax(out["missingness_logits_action"][:, :-1, :, :].float(), dim=-1)
miss_pred = miss_probs.argmax(dim=-1)
miss_mask = torch.as_tensor(valid_np, device=device)[:, :, None].expand_as(miss_target)
if miss_mask.any():
y = miss_target[miss_mask].detach().cpu().numpy().astype(np.int64)
pred = miss_pred[miss_mask].detach().cpu().numpy().astype(np.int64)
conf = multiclass_confusion(y, pred, cfg.n_missing)
grammar_rows.append(
{
"split": split,
"model": OBSERVED_ACTION,
"family": "missingness",
"n": int(y.size),
"accuracy": float((pred == y).mean()),
"macro_f1": float(np.nanmean([2 * conf[i, i] / max(1, conf[i, :].sum() + conf[:, i].sum()) for i in range(conf.shape[0]) if conf[i, :].sum() > 0])),
}
)
ord_logits = out["ordinal_cum_logits_action"][:, :-1, :, :].float()
ord_pred = torch.sigmoid(ord_logits).ge(0.5).long().sum(dim=-1).detach().cpu().numpy()
ord_target = batch["ordinal_cbe"][:, 1:, :, :].ge(0.5).long().sum(dim=-1).detach().cpu().numpy()
ord_mask = next_contact[:, :, None] & batch["ordinal_mask"][:, 1:, :].detach().cpu().numpy().astype(bool)
if ord_mask.any():
oy = ord_target[ord_mask]
op = ord_pred[ord_mask]
grammar_rows.append(
{
"split": split,
"model": OBSERVED_ACTION,
"family": "ordinal",
"n": int(oy.size),
"mae": float(np.abs(op - oy).mean()),
"rmse": float(math.sqrt(float(((op - oy) ** 2).mean()))),
"quadratic_weighted_kappa": quadratic_weighted_kappa(oy, op, cfg.cbe_dim + 1),
}
)
if device.type == "cuda":
torch.cuda.empty_cache()
rollout_rows.extend(
rollout_fidelity_rows(
v5_model,
arrays,
actions,
split,
device,
args.batch_size,
args.rollout_steps,
args.rollout_max_windows,
args.seed + 100 + split_idx,
)
)
(audit_dir / "information_boundary_repeated_audit.json").write_text(json.dumps(json_ready(boundary_payloads), ensure_ascii=False, indent=2), encoding="utf-8")
process_df = pd.DataFrame(process_rows)
if not process_df.empty:
process_df = (
process_df.assign(weighted_value=process_df["estimate"] * process_df["n"])
.groupby(["split", "model", "metric"], dropna=False)
.agg(n=("n", "sum"), weighted_value=("weighted_value", "sum"))
.reset_index()
)
process_df["estimate"] = process_df["weighted_value"] / process_df["n"].replace(0, np.nan)
process_df = process_df.drop(columns=["weighted_value"])
process_df.to_csv(risk_dir / "pwe_process_nll_unweighted.csv", index=False)
aggregate_grammar_rows(grammar_rows).to_csv(risk_dir / "grammar_topjournal_spotcheck.csv", index=False)
pd.DataFrame(rollout_rows).to_csv(rollout_dir / "rollout_distributional_fidelity.csv", index=False)
metric_rows = []
dca_rows = []
calibration_rows = []
prediction_export_rows = []
prediction_export_dir = args.out_dir / "04_prediction_exports"
for i, (key, parts) in enumerate(pred_store.items()):
split, model_name, pipeline, endpoint, horizon = key
y = np.concatenate(parts["y"]).astype(np.int8) if parts["y"] else np.asarray([], dtype=np.int8)
p = np.concatenate(parts["p"]).astype(np.float64) if parts["p"] else np.asarray([], dtype=np.float64)
pid = np.concatenate(parts["pid"]).reshape(-1) if parts["pid"] else np.asarray([], dtype=np.int64)
metrics = weighted_binary_metrics(y, p)
if args.export_predictions:
prediction_export_rows.append(write_prediction_export(prediction_export_dir, key, parts, y=y, p=p, pid=pid))
metric_bootstrap = bootstrap_reps_for_metric(endpoint, horizon, args.bootstrap_clusters) if model_name == OBSERVED_ACTION else 0
ci = bootstrap_metric_ci(y, p, pid, metric_bootstrap, args.seed + i)
row = {
"split": split,
"model": model_name,
"pipeline": pipeline,
"endpoint": endpoint,
"horizon_days": horizon,
"n": int(y.size),
"events": int(y.sum()) if y.size else 0,
**metrics,
**ci,
}
metric_rows.append(row)
calibration_rows.extend(
calibration_curve_rows(
split=split,
model=model_name,
pipeline=pipeline,
endpoint=endpoint,
horizon_days=horizon,
y=y,
p=p,
)
)
thresholds = RISK_THRESHOLD_GRID.get(endpoint, RISK_THRESHOLD_GRID["default"])
dca_bootstrap = metric_bootstrap if model_name == OBSERVED_ACTION else 0
dca_rows.extend(
decision_curve_rows(
split=split,
model=model_name,
pipeline=pipeline,
endpoint=endpoint,
horizon_days=horizon,
y=y,
p=p,
patient_ids=pid,
thresholds=thresholds,
bootstrap_clusters=dca_bootstrap,
seed=args.seed + 1000 + i,
)
)
pd.DataFrame(metric_rows).to_csv(risk_dir / "risk_calibration_metrics_with_cluster_ci.csv", index=False)
pd.DataFrame(calibration_rows).to_csv(risk_dir / "risk_calibration_curve_deciles.csv", index=False)
pd.DataFrame(dca_rows).to_csv(risk_dir / "decision_curve_net_benefit.csv", index=False)
if args.export_predictions:
pd.DataFrame(prediction_export_rows).to_csv(prediction_export_dir / "prediction_export_manifest.csv", index=False)
subgroup_rows = []
for key, parts in subgroup_store.items():
split, model_name, pipeline, endpoint, horizon, axis, level = key
y = np.concatenate(parts["y"]).astype(np.int8) if parts["y"] else np.asarray([], dtype=np.int8)
p = np.concatenate(parts["p"]).astype(np.float64) if parts["p"] else np.asarray([], dtype=np.float64)
if y.size < 30 or np.unique(y).size < 2:
continue
subgroup_rows.append(
{
"split": split,
"model": model_name,
"pipeline": pipeline,
"endpoint": endpoint,
"horizon_days": horizon,
"subgroup_axis": axis,
"subgroup_level": level,
"n": int(y.size),
"events": int(y.sum()),
**weighted_binary_metrics(y, p),
}
)
pd.DataFrame(subgroup_rows).to_csv(subgroup_dir / "subgroup_temporal_risk_metrics.csv", index=False)
summary = {
"status": "completed",
"script": "evaluate_v5_topjournal.py",
"v5_checkpoint": str(args.v5_checkpoint),
"checkpoint_step": payload.get("best_step", payload.get("step")),
"device": str(device),
"precision": args.precision,
"splits": split_names,
"horizons": horizons,
"bootstrap_clusters": int(args.bootstrap_clusters),
"outputs": {
"patient_split_audit": str(audit_dir / "patient_split_audit.csv"),
"information_boundary_repeated_audit": str(audit_dir / "information_boundary_repeated_audit.json"),
"risk_calibration_metrics": str(risk_dir / "risk_calibration_metrics_with_cluster_ci.csv"),
"risk_calibration_curve_deciles": str(risk_dir / "risk_calibration_curve_deciles.csv"),
"decision_curve_net_benefit": str(risk_dir / "decision_curve_net_benefit.csv"),
"rollout_distributional_fidelity": str(rollout_dir / "rollout_distributional_fidelity.csv"),
"subgroup_temporal_risk_metrics": str(subgroup_dir / "subgroup_temporal_risk_metrics.csv"),
"prediction_export_manifest": str(prediction_export_dir / "prediction_export_manifest.csv") if args.export_predictions else None,
},
"model_audit": model_audit,
"claim_boundary": (
"PWE CIF is observed-action-conditioned post-contact/current-care-contact prognosis, "
"not a pure pre-contact deployable risk score. Event-label horizon aggregation is "
"teacher-forced over observed future contact contexts and must be labelled accordingly."
),
}
(args.out_dir / "topjournal_v5_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)}, ensure_ascii=False), flush=True)
if __name__ == "__main__":
main()