| 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 |
| from sklearn.ensemble import HistGradientBoostingClassifier |
| from sklearn.isotonic import IsotonicRegression |
| from sklearn.linear_model import LogisticRegression |
| from sklearn.metrics import average_precision_score, roc_auc_score |
| from sklearn.pipeline import make_pipeline |
| from sklearn.preprocessing import StandardScaler |
|
|
|
|
| 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 |
| from build_v5_action_tensors import ( |
| build_action_arrays, |
| build_target_trial_labels, |
| medication_flags_from_stage0, |
| resolve_split_path, |
| validate_medication_flags, |
| ) |
| from run_service_process_target_trial_v5 import pre_action_features |
| from scan_target_trial_support_v5 import era_tokens_np, future_event_within, load_or_build_action_arrays, unique_landmark_filter |
| from target_trial_estimators_v5 import effective_sample_size |
|
|
| try: |
| from lightgbm import LGBMClassifier |
| except Exception: |
| LGBMClassifier = None |
|
|
|
|
| CAUSE_NAMES = {0: "next_contact", 1: "death", 2: "disengagement"} |
| PWE_BIN_EDGES_DAYS = (7.0, 14.0, 30.0, 60.0, 90.0, 365.0, math.inf) |
| RISK_ENDPOINTS = { |
| "death": {"kind": "pwe", "cause_id": 1}, |
| "disengagement": {"kind": "pwe", "cause_id": 2}, |
| "primary_referral_relapse": {"kind": "event", "event_index": 0}, |
| "service_escalation": {"kind": "event", "event_index": 3}, |
| "clinical_deterioration": {"kind": "event", "event_index": 4}, |
| "high_acuity_state": {"kind": "event", "event_index": 5}, |
| } |
| 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], |
| "service_escalation": [0.03, 0.05, 0.10, 0.15, 0.20], |
| "clinical_deterioration": [0.30, 0.40, 0.50, 0.55, 0.60], |
| "high_acuity_state": [0.10, 0.20, 0.30, 0.40, 0.50], |
| "default": [0.01, 0.03, 0.05, 0.10, 0.15], |
| } |
|
|
|
|
| def load_npz_to_memory(path: Path) -> dict[str, np.ndarray]: |
| with np.load(path, allow_pickle=True) as z: |
| return {key: z[key] for key in z.files} |
|
|
|
|
| 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 service_era_key(service: np.ndarray, years: np.ndarray) -> np.ndarray: |
| return service.astype(np.int64) * 10 + era_tokens_np(years).astype(np.int64) |
|
|
|
|
| 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, sample_weight: np.ndarray | None = None) -> float: |
| x = np.asarray(x, dtype=np.float64).reshape(-1) |
| if not x.size: |
| return math.nan |
| if sample_weight is None: |
| return float(np.mean(x)) |
| w = np.asarray(sample_weight, dtype=np.float64).reshape(-1) |
| ok = np.isfinite(x) & np.isfinite(w) & (w > 0) |
| if not ok.any(): |
| return math.nan |
| return float(np.sum(w[ok] * x[ok]) / np.sum(w[ok])) |
|
|
|
|
| def calibration_intercept_slope(y: np.ndarray, p: np.ndarray, sample_weight: np.ndarray | None = None) -> tuple[float, float]: |
| y = np.asarray(y, dtype=np.int8).reshape(-1) |
| p = clip_prob(p) |
| if y.size == 0 or np.unique(y).size < 2: |
| return math.nan, math.nan |
| try: |
| x = np.log(p / (1.0 - p)).reshape(-1, 1) |
| clf = LogisticRegression(C=1.0e6, solver="lbfgs", max_iter=500) |
| clf.fit(x, y, sample_weight=sample_weight) |
| return float(clf.intercept_[0]), float(clf.coef_[0, 0]) |
| except Exception: |
| return math.nan, 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 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 w is not None: |
| ok = np.isfinite(w) & (w > 0) |
| y = y[ok] |
| p = p[ok] |
| w = w[ok] |
| 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, |
| } |
| auc = float(roc_auc_score(y, p, sample_weight=w)) if np.unique(y).size >= 2 else math.nan |
| ap = float(average_precision_score(y, p, sample_weight=w)) if np.unique(y).size >= 2 else math.nan |
| order = np.argsort(p) |
| ece = 0.0 |
| ece_den = float(y.size) if w is None else float(w.sum()) |
| for idx in np.array_split(order, min(10, max(1, y.size))): |
| if idx.size: |
| bin_w = None if w is None else w[idx] |
| bin_weight = float(idx.size) if w is None else float(bin_w.sum()) |
| ece += abs(weighted_mean(p[idx], bin_w) - weighted_mean(y[idx], bin_w)) * bin_weight / ece_den |
| intercept, slope = calibration_intercept_slope(y, p, w) if include_calibration_model else (math.nan, math.nan) |
| return { |
| "auc": auc, |
| "average_precision": ap, |
| "brier": weighted_mean((p - y) ** 2, w), |
| "ece": float(ece), |
| "ici": integrated_calibration_index(y, p, w) if include_calibration_model else math.nan, |
| "mean_predicted": weighted_mean(p, w), |
| "observed_rate": weighted_mean(y, w), |
| "calibration_intercept": intercept, |
| "calibration_slope": slope, |
| } |
|
|
|
|
| def net_benefit(y: np.ndarray, p: np.ndarray, threshold: float, sample_weight: np.ndarray | None = None) -> float: |
| y = np.asarray(y, dtype=np.int8).reshape(-1) |
| p = clip_prob(p) |
| 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) |
| ok = np.isfinite(w) & (w > 0) |
| y = y[ok] |
| p = p[ok] |
| w = w[ok] |
| if y.size == 0 or float(w.sum()) <= 0.0: |
| return math.nan |
| pred_pos = p >= float(threshold) |
| den = float(w.sum()) |
| 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 calibration_curve_rows( |
| *, |
| split: str, |
| endpoint: str, |
| horizon_days: float, |
| model: str, |
| 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, |
| "endpoint": endpoint, |
| "horizon_days": float(horizon_days), |
| "model": model, |
| "metric_family": "calibration_curve", |
| "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 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: dict[str, list[float]] = 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_net_benefit_ci( |
| y: np.ndarray, |
| p: np.ndarray, |
| patient_ids: np.ndarray, |
| threshold: float, |
| 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) |
| vals = [] |
| for _ in range(n_bootstrap): |
| counts = rng.multinomial(clusters.size, np.full(clusters.size, 1.0 / clusters.size)) |
| vals.append(net_benefit(y, p, threshold, sample_weight=counts[inv].astype(np.float64))) |
| arr = np.asarray(vals, dtype=np.float64) |
| return { |
| "net_benefit_ci_low": float(np.nanquantile(arr, 0.025)) if arr.size else math.nan, |
| "net_benefit_ci_high": float(np.nanquantile(arr, 0.975)) if arr.size else math.nan, |
| } |
|
|
|
|
| def active_next_mask(arrays: dict[str, np.ndarray]) -> np.ndarray: |
| return arrays["valid_mask"][:, :-1].astype(bool) & arrays["valid_mask"][:, 1:].astype(bool) |
|
|
|
|
| def pwe_targets(arrays: dict[str, np.ndarray]) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]: |
| terminal = arrays["terminal_label"][:, 1:].astype(np.int64) |
| service = arrays["service_state"][:, 1:].astype(np.int64) |
| valid = active_next_mask(arrays) |
| cause = np.zeros_like(terminal, dtype=np.int64) |
| cause[(terminal == 1) | (service == 7)] = 1 |
| cause[(terminal == 2) | (terminal == 3) | (service == 6)] = 2 |
| censored = (terminal == 4) | (service == 5) |
| days = np.expm1(arrays["delta_t_next_log"][:, :-1].astype(np.float64)).clip(min=1.0e-6) |
| return cause, censored.astype(bool), valid, days |
|
|
|
|
| def build_conditional_specs(meta: dict[str, Any], vocab: dict[str, Any]) -> tuple[dict[str, Any], ...]: |
| cat_cols = list(meta.get("cat_cols", [])) |
|
|
| def field_idx(name: str) -> int: |
| return cat_cols.index(name) if name in cat_cols else -1 |
|
|
| specs = [] |
| tables = [ |
| ("是否转诊", ["是"], ["转诊类型", "转诊原因", "转诊至机构"]), |
| ("治疗方式", ["住院"], ["本次入院形式", "本次住院名称", "本次入院日期", "末次出院日期"]), |
| ] |
| for parent, active_values, children in tables: |
| parent_idx = field_idx(parent) |
| child_indices = [field_idx(child) for child in children if field_idx(child) >= 0] |
| active_ids = [int(vocab[f"{parent}={value}"]) for value in active_values if f"{parent}={value}" in vocab] |
| if parent_idx >= 0 and child_indices and active_ids: |
| specs.append({"parent_idx": parent_idx, "active_value_ids": active_ids, "child_indices": child_indices}) |
| return tuple(specs) |
|
|
|
|
| def conditional_status_np(cat_value_ids: np.ndarray, missing_ids: np.ndarray, specs: tuple[dict[str, Any], ...]) -> np.ndarray: |
| status = np.zeros_like(cat_value_ids, dtype=np.int64) |
| if not specs: |
| return status |
| observed = missing_ids.astype(np.int64) == 0 |
| unknown = 1 |
| for spec in specs: |
| parent_idx = int(spec["parent_idx"]) |
| child_indices = [int(x) for x in spec["child_indices"]] |
| active_ids = np.asarray([int(x) for x in spec["active_value_ids"]], dtype=cat_value_ids.dtype) |
| if parent_idx < 0 or parent_idx >= cat_value_ids.shape[-1] or not child_indices or active_ids.size == 0: |
| continue |
| parent_val = cat_value_ids[:, :, parent_idx] |
| parent_observed = observed[:, :, parent_idx] & (parent_val != unknown) |
| parent_active = np.isin(parent_val, active_ids) & parent_observed |
| parent_unknown = ~parent_observed |
| for child_idx in child_indices: |
| if child_idx < 0 or child_idx >= cat_value_ids.shape[-1]: |
| continue |
| child_val = cat_value_ids[:, :, child_idx] |
| child_observed = observed[:, :, child_idx] & (child_val != unknown) |
| code = np.zeros_like(parent_val, dtype=np.int64) |
| code[parent_unknown] = 2 |
| code[(~parent_unknown) & (~parent_active) & (~child_observed)] = 1 |
| code[(~parent_unknown) & (~parent_active) & child_observed] = 3 |
| status[:, :, child_idx] = code |
| return status |
|
|
|
|
| def mapped_missing_targets( |
| arrays: dict[str, np.ndarray], |
| meta: dict[str, Any], |
| vocab: dict[str, Any], |
| *, |
| start: int, |
| stop: int | None, |
| ) -> np.ndarray: |
| """Match the v4/v5 loss missingness target: 5 classes incl. structural/no-clinical.""" |
|
|
| raw = arrays["missing_ids"][:, start:stop, :].astype(np.int64) |
| mapped = np.full_like(raw, 3, dtype=np.int64) |
| mapped[raw == 0] = 0 |
| mapped[raw == 1] = 1 |
| mapped[raw == 2] = 3 |
| specs = build_conditional_specs(meta, vocab) |
| if specs: |
| cat = arrays["cat_value_ids"][:, start:stop, :].astype(np.int64) |
| conditional = conditional_status_np(cat, raw, specs) |
| mapped[conditional == 1] = 2 |
| service = arrays["service_state"][:, start:stop].astype(np.int64) |
| terminal = arrays["terminal_label"][:, start:stop].astype(np.int64) |
| no_clinical = (terminal > 0) | (service >= 5) |
| mapped = np.where(no_clinical[:, :, None], 4, mapped) |
| return mapped.clip(0, 4) |
|
|
|
|
| def exposure_by_bin(days: np.ndarray) -> np.ndarray: |
| upper = np.asarray(PWE_BIN_EDGES_DAYS[:-1], dtype=np.float64) |
| lower = np.asarray((0.0, 7.0, 14.0, 30.0, 60.0, 90.0), dtype=np.float64) |
| finite = (np.minimum(days[..., None], upper) - lower).clip(min=0.0) |
| tail = (days - 365.0).clip(min=0.0)[..., None] |
| return np.concatenate([finite, tail], axis=-1) |
|
|
|
|
| def pwe_bin_index(days: np.ndarray) -> np.ndarray: |
| upper = np.asarray(PWE_BIN_EDGES_DAYS[:-1], dtype=np.float64) |
| return np.searchsorted(upper, days, side="left").clip(max=len(PWE_BIN_EDGES_DAYS) - 1) |
|
|
|
|
| def fit_semi_markov(train: dict[str, np.ndarray], smoothing: float) -> dict[str, Any]: |
| cause, censored, valid, days = pwe_targets(train) |
| cur_state = train["service_state"][:, :-1].astype(np.int64) |
| next_state = train["service_state"][:, 1:].astype(np.int64) |
| key = service_era_key(cur_state, train["visit_year"][:, :-1]) |
| n_states = int(max(train["service_state"].max(), 7) + 1) |
| n_bins = len(PWE_BIN_EDGES_DAYS) |
| trans_counts: dict[int, np.ndarray] = {} |
| exposure: dict[int, np.ndarray] = {} |
| events: dict[int, np.ndarray] = {} |
| for k in np.unique(key[valid]): |
| m = valid & (key == k) |
| counts = np.bincount(next_state[m].clip(0, n_states - 1), minlength=n_states).astype(np.float64) + smoothing |
| trans_counts[int(k)] = counts / counts.sum() |
| ex = exposure_by_bin(days[m]).sum(axis=0) |
| ev = np.zeros((3, n_bins), dtype=np.float64) |
| bin_idx = pwe_bin_index(days[m]) |
| c = cause[m] |
| z = censored[m] |
| for cause_id in range(3): |
| np.add.at(ev[cause_id], bin_idx[(~z) & (c == cause_id)], 1.0) |
| exposure[int(k)] = ex + smoothing |
| events[int(k)] = ev + smoothing |
| global_counts = np.bincount(next_state[valid].clip(0, n_states - 1), minlength=n_states).astype(np.float64) + smoothing |
| global_trans = global_counts / global_counts.sum() |
| global_ex = exposure_by_bin(days[valid]).sum(axis=0) + smoothing |
| global_ev = np.zeros((3, n_bins), dtype=np.float64) + smoothing |
| bin_idx = pwe_bin_index(days[valid]) |
| for cause_id in range(3): |
| np.add.at(global_ev[cause_id], bin_idx[(~censored[valid]) & (cause[valid] == cause_id)], 1.0) |
| return { |
| "trans": trans_counts, |
| "exposure": exposure, |
| "events": events, |
| "global_trans": global_trans, |
| "global_exposure": global_ex, |
| "global_events": global_ev, |
| "n_states": n_states, |
| "n_bins": n_bins, |
| "smoothing": smoothing, |
| } |
|
|
|
|
| def semi_markov_metrics(model: dict[str, Any], arrays: dict[str, np.ndarray], split: str) -> list[dict[str, Any]]: |
| cause, censored, valid, days = pwe_targets(arrays) |
| cur_state = arrays["service_state"][:, :-1].astype(np.int64) |
| next_state = arrays["service_state"][:, 1:].astype(np.int64).clip(0, model["n_states"] - 1) |
| key = service_era_key(cur_state, arrays["visit_year"][:, :-1]) |
| rows = [] |
| key_flat = key[valid].astype(np.int64) |
| y = next_state[valid] |
| pred = np.empty(y.shape, dtype=np.int64) |
| ce_sum = 0.0 |
| for k in np.unique(key_flat): |
| idx = np.flatnonzero(key_flat == int(k)) |
| probs_k = model["trans"].get(int(k), model["global_trans"]) |
| pred[idx] = int(np.argmax(probs_k)) |
| ce_sum += float(-np.log(np.clip(probs_k[y[idx]], 1.0e-12, 1.0)).sum()) |
| ce = ce_sum / y.size if y.size else math.nan |
| conf = np.zeros((model["n_states"], model["n_states"]), dtype=np.int64) |
| if y.size: |
| np.add.at(conf, (y, pred), 1) |
| f1 = [] |
| for cls in range(model["n_states"]): |
| if conf[cls].sum() <= 0: |
| continue |
| tp = conf[cls, cls] |
| fp = conf[:, cls].sum() - tp |
| fn = conf[cls].sum() - tp |
| f1.append(2 * tp / max(1, 2 * tp + fp + fn)) |
| rows.append( |
| { |
| "split": split, |
| "model": "Semi-Markov Care-Process", |
| "metric_family": "service_transition", |
| "n": int(y.size), |
| "cross_entropy": float(ce), |
| "accuracy": float((pred == y).mean()) if y.size else math.nan, |
| "macro_f1": float(np.mean(f1)) if f1 else math.nan, |
| } |
| ) |
| cause_flat = cause[valid].astype(np.int64) |
| cens_flat = censored[valid].astype(bool) |
| days_flat = days[valid].astype(np.float64) |
| nll_sum = 0.0 |
| nll_n = int(days_flat.size) |
| bin_flat = pwe_bin_index(days_flat) |
| for k in np.unique(key_flat): |
| idx = np.flatnonzero(key_flat == int(k)) |
| rates = model["events"].get(int(k), model["global_events"]) / model["exposure"].get(int(k), model["global_exposure"])[None, :] |
| total_rate = rates.sum(axis=0) |
| ex = exposure_by_bin(days_flat[idx]) |
| ll = -np.sum(ex * total_rate[None, :], axis=1) |
| event_idx = idx[~cens_flat[idx]] |
| if event_idx.size: |
| rel = np.flatnonzero(~cens_flat[idx]) |
| ll[rel] += np.log(np.clip(rates[cause_flat[event_idx], bin_flat[event_idx]], 1.0e-12, None)) |
| nll_sum += float((-ll).sum()) |
| rows.append( |
| { |
| "split": split, |
| "model": "Semi-Markov Care-Process", |
| "metric_family": "pwe_time_to_next_process", |
| "n": nll_n, |
| "pwe_nll": nll_sum / nll_n if nll_n else math.nan, |
| } |
| ) |
| return rows |
|
|
|
|
| def fit_empirical_grammar(train: dict[str, np.ndarray], meta: dict[str, Any], vocab_map: dict[str, Any], smoothing: float) -> dict[str, Any]: |
| valid = active_next_mask(train) |
| next_contact = valid & (train["service_state"][:, 1:] < 5) |
| key = service_era_key(train["service_state"][:, :-1], train["visit_year"][:, :-1]) |
| n_missing = 5 |
| missing_target = mapped_missing_targets(train, meta, vocab_map, start=1, stop=None) |
| vocab = int(meta.get("cat_vocab_size", int(train["cat_value_ids"].max()) + 1)) |
| cbe_dim = int(meta.get("cbe_dim", 5)) |
| out: dict[str, Any] = {"missing": {}, "cat": {}, "numeric": {}, "ordinal": {}, "global": {}, "target_contract": "v4_v5_mapped_5class_missingness"} |
| for k in np.unique(key[valid]): |
| m = valid & (key == k) |
| miss = missing_target[m].reshape(-1).astype(int) |
| out["missing"][int(k)] = (np.bincount(miss.clip(0, n_missing - 1), minlength=n_missing) + smoothing) |
| out["missing"][int(k)] /= out["missing"][int(k)].sum() |
| miss_global = missing_target[valid].reshape(-1).astype(int) |
| out["global"]["missing"] = (np.bincount(miss_global.clip(0, n_missing - 1), minlength=n_missing) + smoothing) |
| out["global"]["missing"] /= out["global"]["missing"].sum() |
| for field_idx in range(train["cat_value_ids"].shape[2]): |
| out["cat"][field_idx] = {} |
| observed = next_contact & (missing_target[:, :, field_idx] == 0) |
| vals_global = train["cat_value_ids"][:, 1:, field_idx][observed].astype(int) |
| g = np.bincount(vals_global.clip(0, vocab - 1), minlength=vocab).astype(np.float64) + smoothing |
| out["global"][f"cat_{field_idx}"] = g / g.sum() |
| for k in np.unique(key[observed]): |
| vals = train["cat_value_ids"][:, 1:, field_idx][observed & (key == k)].astype(int) |
| counts = np.bincount(vals.clip(0, vocab - 1), minlength=vocab).astype(np.float64) + smoothing |
| out["cat"][field_idx][int(k)] = counts / counts.sum() |
| for field_idx in range(train["numeric_values"].shape[2]): |
| out["numeric"][field_idx] = {} |
| observed = next_contact & train["numeric_mask"][:, 1:, field_idx] |
| vals = train["numeric_values"][:, 1:, field_idx][observed].astype(float) |
| out["global"][f"num_{field_idx}"] = float(np.nanmean(vals)) if vals.size else 0.0 |
| for k in np.unique(key[observed]): |
| sub = train["numeric_values"][:, 1:, field_idx][observed & (key == k)].astype(float) |
| out["numeric"][field_idx][int(k)] = float(np.nanmean(sub)) if sub.size else out["global"][f"num_{field_idx}"] |
| for field_idx in range(train["ordinal_cbe"].shape[2]): |
| out["ordinal"][field_idx] = {} |
| observed = next_contact & train["ordinal_mask"][:, 1:, field_idx] |
| levels = train["ordinal_cbe"][:, 1:, field_idx, :][observed].sum(axis=1).astype(int) |
| counts = np.bincount(levels.clip(0, cbe_dim), minlength=cbe_dim + 1).astype(np.float64) + smoothing |
| out["global"][f"ord_{field_idx}"] = int(np.argmax(counts)) |
| for k in np.unique(key[observed]): |
| sub = train["ordinal_cbe"][:, 1:, field_idx, :][observed & (key == k)].sum(axis=1).astype(int) |
| c = np.bincount(sub.clip(0, cbe_dim), minlength=cbe_dim + 1).astype(np.float64) + smoothing |
| out["ordinal"][field_idx][int(k)] = int(np.argmax(c)) |
| return out |
|
|
|
|
| def empirical_grammar_metrics(model: dict[str, Any], arrays: dict[str, np.ndarray], split: str, meta: dict[str, Any], vocab: dict[str, Any]) -> list[dict[str, Any]]: |
| valid = active_next_mask(arrays) |
| next_contact = valid & (arrays["service_state"][:, 1:] < 5) |
| key = service_era_key(arrays["service_state"][:, :-1], arrays["visit_year"][:, :-1]) |
| rows = [] |
| missing_target = mapped_missing_targets(arrays, meta, vocab, start=1, stop=None) |
| miss_y = missing_target[valid].reshape(-1).astype(int) |
| miss_keys = np.repeat(key[valid], arrays["missing_ids"].shape[2]) |
| miss_correct = 0 |
| miss_brier_sum = 0.0 |
| for k in np.unique(miss_keys): |
| idx = np.flatnonzero(miss_keys == int(k)) |
| probs = model["missing"].get(int(k), model["global"]["missing"]) |
| yk = miss_y[idx].clip(0, probs.shape[0] - 1) |
| miss_correct += int((int(np.argmax(probs)) == yk).sum()) |
| miss_brier_sum += float((np.sum(probs * probs) + 1.0) * yk.size - 2.0 * np.sum(probs[yk])) |
| rows.append( |
| { |
| "split": split, |
| "model": "Empirical state-year grammar", |
| "family": "missingness", |
| "n": int(miss_y.size), |
| "accuracy": float(miss_correct / miss_y.size) if miss_y.size else math.nan, |
| "multiclass_brier": float(miss_brier_sum / miss_y.size) if miss_y.size else math.nan, |
| "target_contract": model.get("target_contract", "v4_v5_mapped_5class_missingness"), |
| } |
| ) |
| cat_ce = [] |
| cat_top1 = [] |
| cat_top3 = [] |
| cat_n = 0 |
| for field_idx in range(arrays["cat_value_ids"].shape[2]): |
| observed = next_contact & (missing_target[:, :, field_idx] == 0) |
| y = arrays["cat_value_ids"][:, 1:, field_idx][observed].astype(int) |
| if y.size: |
| cat_n += int(y.size) |
| k_obs = key[observed].astype(np.int64) |
| ce_sum = 0.0 |
| top1_sum = 0 |
| top3_sum = 0 |
| for k in np.unique(k_obs): |
| idx = np.flatnonzero(k_obs == int(k)) |
| probs = model["cat"][field_idx].get(int(k), model["global"][f"cat_{field_idx}"]) |
| yk = y[idx].clip(0, probs.shape[0] - 1) |
| ce_sum += float(-np.log(np.clip(probs[yk], 1.0e-12, 1.0)).sum()) |
| top1 = int(np.argmax(probs)) |
| top1_sum += int((yk == top1).sum()) |
| k_top = min(3, probs.shape[0]) |
| top3 = np.argpartition(-probs, kth=k_top - 1)[:k_top] |
| top3_sum += int(np.isin(yk, top3).sum()) |
| cat_ce.append(ce_sum / y.size) |
| cat_top1.append(top1_sum / y.size) |
| cat_top3.append(top3_sum / y.size) |
| rows.append( |
| { |
| "split": split, |
| "model": "Empirical state-year grammar", |
| "family": "categorical", |
| "n": int(cat_n), |
| "mean_cross_entropy": float(np.nanmean(cat_ce)) if cat_ce else math.nan, |
| "mean_top1_accuracy": float(np.nanmean(cat_top1)) if cat_top1 else math.nan, |
| "mean_top3_accuracy": float(np.nanmean(cat_top3)) if cat_top3 else math.nan, |
| } |
| ) |
| abs_err = [] |
| sq_err = [] |
| for field_idx in range(arrays["numeric_values"].shape[2]): |
| observed = next_contact & arrays["numeric_mask"][:, 1:, field_idx] |
| y = arrays["numeric_values"][:, 1:, field_idx][observed].astype(float) |
| pred = np.asarray([model["numeric"][field_idx].get(int(k), model["global"][f"num_{field_idx}"]) for k in key[observed]], dtype=float) if y.size else np.asarray([]) |
| abs_err.append(np.abs(pred - y)) |
| sq_err.append((pred - y) ** 2) |
| ae = np.concatenate(abs_err) if abs_err else np.asarray([]) |
| se = np.concatenate(sq_err) if sq_err else np.asarray([]) |
| rows.append( |
| { |
| "split": split, |
| "model": "Empirical state-year grammar", |
| "family": "numeric", |
| "n": int(ae.size), |
| "mae": float(ae.mean()) if ae.size else math.nan, |
| "rmse": float(math.sqrt(float(se.mean()))) if se.size else math.nan, |
| } |
| ) |
| ord_abs = [] |
| for field_idx in range(arrays["ordinal_cbe"].shape[2]): |
| observed = next_contact & arrays["ordinal_mask"][:, 1:, field_idx] |
| y = arrays["ordinal_cbe"][:, 1:, field_idx, :][observed].sum(axis=1).astype(int) |
| pred = np.asarray([model["ordinal"][field_idx].get(int(k), model["global"][f"ord_{field_idx}"]) for k in key[observed]], dtype=int) if y.size else np.asarray([]) |
| ord_abs.append(np.abs(pred - y)) |
| oa = np.concatenate(ord_abs) if ord_abs else np.asarray([]) |
| rows.append({"split": split, "model": "Empirical state-year grammar", "family": "ordinal", "n": int(oa.size), "mae": float(oa.mean()) if oa.size else math.nan}) |
| return rows |
|
|
|
|
| def horizon_label(arrays: dict[str, np.ndarray], endpoint: str, horizon: float) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]: |
| spec = RISK_ENDPOINTS[endpoint] |
| valid = active_next_mask(arrays) |
| row_idx, pos_idx = np.nonzero(valid) |
| if endpoint in ("death", "disengagement"): |
| cause, censored, _, days = pwe_targets(arrays) |
| cause_id = int(spec["cause_id"]) |
| y_full = ((~censored) & (cause == cause_id) & (days <= float(horizon))).astype(np.int8) |
| observed = valid & ~(censored & (days <= float(horizon))) |
| row_idx, pos_idx = np.nonzero(observed) |
| y = y_full[row_idx, pos_idx] |
| else: |
| y_full = future_event_within(arrays, int(spec["event_index"]), float(horizon)).astype(np.int8) |
| valid_full = arrays["valid_mask"].astype(bool) |
| times = arrays["time_since_start_days"].astype(float) |
| seq_len = valid_full.shape[1] |
| current_len = seq_len - 1 |
| any_future = np.zeros((valid_full.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] |
| any_future[:, :cur_len] |= valid_full[:, offset : offset + cur_len] & (dt > 0.0) & (dt <= float(horizon)) |
| last_time = np.where(valid_full, times, -np.inf).max(axis=1, keepdims=True) |
| followup = last_time - times[:, :-1] |
| observed = valid & any_future & ((y_full[:, :-1].astype(bool)) | (followup >= float(horizon))) |
| row_idx, pos_idx = np.nonzero(observed) |
| y = y_full[:, :-1][row_idx, pos_idx] |
| if row_idx.size and "patient_ids" in arrays: |
| keep = unique_landmark_filter(arrays, row_idx, pos_idx) |
| row_idx = row_idx[keep] |
| pos_idx = pos_idx[keep] |
| y = y[keep] |
| return row_idx, pos_idx, y.astype(np.int8), arrays["patient_ids"][row_idx].astype(str) if "patient_ids" in arrays else row_idx.astype(str) |
|
|
|
|
| def fit_classifier(X: np.ndarray, y: np.ndarray, model_kind: str, seed: int, max_train: int, require_lightgbm: bool) -> tuple[Any, str]: |
| rng = np.random.default_rng(seed) |
| if max_train > 0 and X.shape[0] > max_train: |
| idx = rng.choice(X.shape[0], size=max_train, replace=False) |
| X_fit = X[idx] |
| y_fit = y[idx] |
| else: |
| X_fit, y_fit = X, y |
| if np.unique(y_fit).size < 2: |
| rate = float(y_fit.mean()) if y_fit.size else 0.0 |
| return rate, "constant_rate" |
| if model_kind == "clinical_landmark": |
| model = make_pipeline( |
| StandardScaler(with_mean=True, with_std=True), |
| LogisticRegression(max_iter=1000, C=0.5, solver="lbfgs"), |
| ) |
| model.fit(X_fit, y_fit) |
| return model, "Clinical Landmark logistic (L2)" |
| if LGBMClassifier is not None: |
| model = LGBMClassifier( |
| n_estimators=400, |
| learning_rate=0.03, |
| num_leaves=31, |
| subsample=0.85, |
| colsample_bytree=0.85, |
| reg_lambda=1.0, |
| random_state=seed, |
| n_jobs=-1, |
| verbose=-1, |
| ) |
| model.fit(X_fit, y_fit) |
| return model, "LightGBM Landmark" |
| if require_lightgbm: |
| raise RuntimeError("LightGBM backend is required but lightgbm could not be imported.") |
| model = HistGradientBoostingClassifier( |
| learning_rate=0.05, |
| max_iter=160, |
| max_leaf_nodes=31, |
| min_samples_leaf=80, |
| l2_regularization=0.01, |
| early_stopping=True, |
| random_state=seed, |
| ) |
| model.fit(X_fit, y_fit) |
| return model, "HistGradientBoosting fallback for LightGBM" |
|
|
|
|
| def predict_classifier(model: Any, X: np.ndarray) -> np.ndarray: |
| if isinstance(model, float): |
| return np.full(X.shape[0], model, dtype=np.float64) |
| return model.predict_proba(X)[:, 1].astype(np.float64) |
|
|
|
|
| def risk_baseline_rows( |
| train: dict[str, np.ndarray], |
| eval_arrays: dict[str, np.ndarray], |
| split: str, |
| endpoints: list[str], |
| horizons: list[float], |
| seed: int, |
| max_train: int, |
| bootstrap_clusters: int, |
| require_lightgbm: bool, |
| model_cache: dict[tuple[str, float, str], tuple[Any, str]] | None = None, |
| ) -> list[dict[str, Any]]: |
| rows = [] |
| for endpoint in endpoints: |
| for horizon in horizons: |
| print(json.dumps({"stage": "risk_horizon_labels", "split": split, "endpoint": endpoint, "horizon_days": float(horizon)}), flush=True) |
| tr_row, tr_pos, y_train, _ = horizon_label(train, endpoint, horizon) |
| ev_row, ev_pos, y_eval, patient_ids = horizon_label(eval_arrays, endpoint, horizon) |
| if y_train.size == 0 or y_eval.size == 0: |
| continue |
| X_train, _ = pre_action_features(train, tr_row, tr_pos) |
| X_eval, _ = pre_action_features(eval_arrays, ev_row, ev_pos) |
| for kind in ("clinical_landmark", "lightgbm_landmark"): |
| print( |
| json.dumps( |
| { |
| "stage": "risk_model_fit", |
| "split": split, |
| "endpoint": endpoint, |
| "horizon_days": float(horizon), |
| "model_kind": kind, |
| "train_n": int(y_train.size), |
| "eval_n": int(y_eval.size), |
| "max_train": int(max_train), |
| } |
| ), |
| flush=True, |
| ) |
| cache_key = (endpoint, float(horizon), kind) |
| if model_cache is not None and cache_key in model_cache: |
| model, model_label = model_cache[cache_key] |
| else: |
| model, model_label = fit_classifier(X_train, y_train, kind, seed + int(horizon) + len(endpoint), max_train, require_lightgbm) |
| if model_cache is not None: |
| model_cache[cache_key] = (model, model_label) |
| print( |
| json.dumps( |
| { |
| "stage": "risk_model_predict", |
| "split": split, |
| "endpoint": endpoint, |
| "horizon_days": float(horizon), |
| "model": model_label, |
| } |
| ), |
| flush=True, |
| ) |
| p = predict_classifier(model, X_eval) |
| metrics = weighted_binary_metrics(y_eval, p) |
| boot = bootstrap_metric_ci( |
| y_eval, |
| p, |
| patient_ids, |
| bootstrap_clusters, |
| seed + int(horizon) * 13 + len(endpoint) * 101 + (0 if kind == "clinical_landmark" else 1), |
| ) |
| base = { |
| "split": split, |
| "endpoint": endpoint, |
| "horizon_days": float(horizon), |
| "model": model_label, |
| "n": int(y_eval.size), |
| "events": int(y_eval.sum()), |
| **metrics, |
| **boot, |
| } |
| rows.append(base) |
| rows.extend( |
| calibration_curve_rows( |
| split=split, |
| endpoint=endpoint, |
| horizon_days=float(horizon), |
| model=model_label, |
| y=y_eval, |
| p=p, |
| ) |
| ) |
| for threshold in RISK_THRESHOLD_GRID.get(endpoint, RISK_THRESHOLD_GRID["default"]): |
| dca_boot = bootstrap_net_benefit_ci( |
| y_eval, |
| p, |
| patient_ids, |
| float(threshold), |
| bootstrap_clusters, |
| seed + int(horizon) * 17 + int(round(float(threshold) * 10000)) + (0 if kind == "clinical_landmark" else 1), |
| ) |
| rows.append( |
| { |
| "split": split, |
| "endpoint": endpoint, |
| "horizon_days": float(horizon), |
| "model": model_label, |
| "metric_family": "decision_curve", |
| "threshold": float(threshold), |
| "net_benefit": net_benefit(y_eval, p, float(threshold)), |
| "bootstrap_n": int(bootstrap_clusters), |
| "n": int(y_eval.size), |
| "events": int(y_eval.sum()), |
| **dca_boot, |
| } |
| ) |
| return rows |
|
|
|
|
| def propensity_baseline_rows( |
| train: dict[str, np.ndarray], |
| eval_arrays: dict[str, np.ndarray], |
| train_actions: dict[str, np.ndarray], |
| eval_actions: dict[str, np.ndarray], |
| ontology: dict[str, Any], |
| split: str, |
| seed: int, |
| max_train: int, |
| require_lightgbm: bool, |
| model_cache: dict[int, Any] | None = None, |
| ) -> list[dict[str, Any]]: |
| rows = [] |
| slot_names = [slot["slot_name"] for slot in ontology["slots"]] |
| local = ontology.get("local_label_to_action_value_id_by_slot", {}) |
| valid_train = train["valid_mask"].astype(bool) |
| valid_eval = eval_arrays["valid_mask"].astype(bool) |
| for slot_id, slot_name in enumerate(slot_names): |
| present_id = local.get(slot_name, {}).get("present") |
| if present_id is None: |
| positive_label = next((label for label in ("是", "有", "面访", "住院", "持续", "社区转医院") if label in local.get(slot_name, {})), None) |
| if positive_label is None: |
| continue |
| present_id = int(local[slot_name][positive_label]) |
| tr_mask = valid_train & train_actions["action_mask"][:, :, slot_id].astype(bool) |
| ev_mask = valid_eval & eval_actions["action_mask"][:, :, slot_id].astype(bool) |
| tr_row, tr_pos = np.nonzero(tr_mask) |
| ev_row, ev_pos = np.nonzero(ev_mask) |
| if tr_row.size == 0 or ev_row.size == 0: |
| continue |
| y_train = (train_actions["action_value_ids"][:, :, slot_id][tr_row, tr_pos] == int(present_id)).astype(np.int8) |
| y_eval = (eval_actions["action_value_ids"][:, :, slot_id][ev_row, ev_pos] == int(present_id)).astype(np.int8) |
| X_train, _ = pre_action_features(train, tr_row, tr_pos) |
| X_eval, _ = pre_action_features(eval_arrays, ev_row, ev_pos) |
| for kind, label in (("empirical_state_era", "Empirical service-state x era propensity"), ("lightgbm_landmark", "LightGBM propensity")): |
| if kind == "empirical_state_era": |
| key_train = service_era_key(train["service_state"][tr_row, tr_pos], train["visit_year"][tr_row, tr_pos]) |
| key_eval = service_era_key(eval_arrays["service_state"][ev_row, ev_pos], eval_arrays["visit_year"][ev_row, ev_pos]) |
| rates = {} |
| for k in np.unique(key_train): |
| m = key_train == k |
| rates[int(k)] = float((y_train[m].sum() + 1.0) / (m.sum() + 2.0)) |
| fallback = float((y_train.sum() + 1.0) / (y_train.size + 2.0)) |
| p = np.asarray([rates.get(int(k), fallback) for k in key_eval], dtype=np.float64) |
| else: |
| if model_cache is not None and slot_id in model_cache: |
| model = model_cache[slot_id] |
| else: |
| model, _ = fit_classifier(X_train, y_train, "lightgbm_landmark", seed + slot_id, max_train, require_lightgbm) |
| if model_cache is not None: |
| model_cache[slot_id] = model |
| p = predict_classifier(model, X_eval) |
| p = np.clip(p, 1.0e-4, 1.0 - 1.0e-4) |
| w = y_eval / p + (1 - y_eval) / (1.0 - p) |
| rows.append( |
| { |
| "split": split, |
| "model": label, |
| "slot_id": int(slot_id), |
| "slot_name": slot_name, |
| "n": int(y_eval.size), |
| "present_rate": float(y_eval.mean()) if y_eval.size else math.nan, |
| "mean_propensity": float(p.mean()) if p.size else math.nan, |
| "p01": float(np.quantile(p, 0.01)) if p.size else math.nan, |
| "p05": float(np.quantile(p, 0.05)) if p.size else math.nan, |
| "p50": float(np.quantile(p, 0.50)) if p.size else math.nan, |
| "p95": float(np.quantile(p, 0.95)) if p.size else math.nan, |
| "p99": float(np.quantile(p, 0.99)) if p.size else math.nan, |
| "extreme_propensity_rate": float(((p < 0.01) | (p > 0.99)).mean()) if p.size else math.nan, |
| "effective_sample_size_binary_ipw": effective_sample_size(w), |
| **weighted_binary_metrics(y_eval, p), |
| } |
| ) |
| return rows |
|
|
|
|
| def load_actions( |
| split: str, |
| arrays: dict[str, np.ndarray], |
| action_dir: Path | None, |
| tensor_metadata: dict[str, Any], |
| ordinal_direction: dict[str, Any], |
| cat_value_vocab: dict[str, int], |
| ontology: dict[str, Any], |
| stage0_dir: Path, |
| max_windows: int, |
| ) -> dict[str, np.ndarray]: |
| actions, _, _ = load_or_build_action_arrays( |
| split=split, |
| arrays=arrays, |
| action_dir=action_dir, |
| tensor_metadata=tensor_metadata, |
| ordinal_direction=ordinal_direction, |
| cat_value_vocab=cat_value_vocab, |
| ontology=ontology, |
| grace_days=(30.0, 60.0, 90.0), |
| stage0_dir=stage0_dir, |
| ) |
| for key, value in actions.items(): |
| if max_windows > 0 and value.shape[0] > max_windows: |
| actions[key] = value[:max_windows] |
| return actions |
|
|
|
|
| def main() -> None: |
| parser = argparse.ArgumentParser(description="Run matched non-neural baselines for SCTM-v5 top-journal tables.") |
| 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("--out-dir", type=Path, required=True) |
| parser.add_argument("--splits", default="val,test") |
| parser.add_argument("--horizons", default="90,365") |
| parser.add_argument("--risk-endpoints", default="death,disengagement,primary_referral_relapse,service_escalation,clinical_deterioration,high_acuity_state") |
| parser.add_argument("--max-windows", type=int, default=0) |
| parser.add_argument("--max-risk-train", type=int, default=500000) |
| parser.add_argument("--bootstrap-clusters", type=int, default=200) |
| parser.add_argument("--seed", type=int, default=20260526) |
| parser.add_argument("--require-lightgbm", action="store_true") |
| args = parser.parse_args() |
|
|
| args.out_dir.mkdir(parents=True, exist_ok=True) |
| 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) |
| train = load_npz_to_memory(resolve_split_path(args.tensor_dir, "train")) |
| if args.max_windows > 0: |
| train = {k: v[: args.max_windows] if getattr(v, "shape", (0,))[0] == train["valid_mask"].shape[0] else v for k, v in train.items()} |
| train_actions = load_actions("train", train, args.action_dir, meta, ordinal_direction, vocab, ontology, args.stage0_dir, args.max_windows) |
|
|
| print(json.dumps({"stage": "fit_semi_markov", "train_windows": int(train["valid_mask"].shape[0])}), flush=True) |
| semi = fit_semi_markov(train, smoothing=1.0) |
| print(json.dumps({"stage": "fit_empirical_grammar"}), flush=True) |
| grammar = fit_empirical_grammar(train, meta, vocab, smoothing=1.0) |
| print(json.dumps({"stage": "fit_baseline_models_start", "splits": args.splits}), flush=True) |
| horizons = [float(x) for x in args.horizons.split(",") if x.strip()] |
| endpoints = [x.strip() for x in args.risk_endpoints.split(",") if x.strip()] |
|
|
| semi_rows = [] |
| grammar_rows = [] |
| risk_rows = [] |
| propensity_rows = [] |
| risk_model_cache: dict[tuple[str, float, str], tuple[Any, str]] = {} |
| propensity_model_cache: dict[int, Any] = {} |
| for split in [s.strip() for s in args.splits.split(",") if s.strip()]: |
| print(json.dumps({"stage": "evaluate_split_start", "split": split}), flush=True) |
| 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()} |
| actions = load_actions(split, arrays, args.action_dir, meta, ordinal_direction, vocab, ontology, args.stage0_dir, args.max_windows) |
| print(json.dumps({"stage": "semi_markov_metrics", "split": split}), flush=True) |
| semi_rows.extend(semi_markov_metrics(semi, arrays, split)) |
| print(json.dumps({"stage": "empirical_grammar_metrics", "split": split}), flush=True) |
| grammar_rows.extend(empirical_grammar_metrics(grammar, arrays, split, meta, vocab)) |
| print(json.dumps({"stage": "risk_baseline_metrics", "split": split}), flush=True) |
| risk_rows.extend( |
| risk_baseline_rows( |
| train, |
| arrays, |
| split, |
| endpoints, |
| horizons, |
| args.seed, |
| args.max_risk_train, |
| args.bootstrap_clusters, |
| args.require_lightgbm, |
| risk_model_cache, |
| ) |
| ) |
| print(json.dumps({"stage": "propensity_baseline_metrics", "split": split}), flush=True) |
| propensity_rows.extend(propensity_baseline_rows(train, arrays, train_actions, actions, ontology, split, args.seed, args.max_risk_train, args.require_lightgbm, propensity_model_cache)) |
| print(json.dumps({"stage": "evaluate_split_done", "split": split}), flush=True) |
|
|
| paths = { |
| "semi_markov": args.out_dir / "semi_markov_care_process_metrics.csv", |
| "empirical_grammar": args.out_dir / "empirical_state_year_grammar_metrics.csv", |
| "risk": args.out_dir / "clinical_lightgbm_risk_baseline_metrics.csv", |
| "propensity": args.out_dir / "empirical_lightgbm_propensity_baseline_metrics.csv", |
| "behrt_status": args.out_dir / "behrt_style_baseline_status.json", |
| } |
| pd.DataFrame(semi_rows).to_csv(paths["semi_markov"], index=False) |
| pd.DataFrame(grammar_rows).to_csv(paths["empirical_grammar"], index=False) |
| pd.DataFrame(risk_rows).to_csv(paths["risk"], index=False) |
| pd.DataFrame(propensity_rows).to_csv(paths["propensity"], index=False) |
| behrt_status = { |
| "status": "not_run_in_this_script", |
| "reason": "Exact BEHRT-style baseline requires separate event-token MLM pretraining and landmark fine-tuning; this non-neural matched baseline suite intentionally does not substitute a tabular transformer for that baseline.", |
| "required_spec": { |
| "architecture": "6-layer Transformer encoder, 12 heads, hidden size 768", |
| "pretraining": "15% masked language modeling over m1_event_token_v2 tokens", |
| "finetuning": "landmark risk heads for death, disengagement, event0 primary referral relapse, service escalation, clinical deterioration, and high acuity", |
| }, |
| } |
| paths["behrt_status"].write_text(json.dumps(behrt_status, ensure_ascii=False, indent=2), encoding="utf-8") |
| summary = { |
| "status": "completed", |
| "script": "evaluate_v5_matched_baselines.py", |
| "splits": args.splits, |
| "horizons": horizons, |
| "risk_endpoints": endpoints, |
| "max_risk_train": int(args.max_risk_train), |
| "bootstrap_clusters": int(args.bootstrap_clusters), |
| "require_lightgbm": bool(args.require_lightgbm), |
| "lightgbm_backend": "lightgbm" if LGBMClassifier is not None else "unavailable", |
| "outputs": {k: str(v) for k, v in paths.items()}, |
| "notes": ( |
| "LightGBM rows use the lightgbm backend; use --require-lightgbm for fail-fast publication runs. " |
| "Clinical Landmark uses unweighted L2 logistic regression so predict_proba remains on the natural prevalence scale. " |
| "Event-horizon endpoints use the same teacher-forced observed-future-contact evaluable mask as evaluate_v5_topjournal.py; " |
| "PWE endpoints use the same next-transition CIF target convention. " |
| "Supported event endpoints include primary_referral_relapse, service_escalation, clinical_deterioration, and high_acuity_state." |
| ), |
| } |
| (args.out_dir / "matched_baseline_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() |
|
|