sctm2-r5-code-20260611 / tests /test_step5_preflight_code.py
agoniii97's picture
Normalize datetime precision for HF P1 tensor build
ae6d94c verified
Raw
History Blame Contribute Delete
13.2 kB
import json
from pathlib import Path
import numpy as np
import pandas as pd
import pytest
pytest.importorskip("pyarrow", reason="pyarrow missing; install requirements-step5.txt to run Step5 package tests")
from sctm_next_round.step5_baseline_evaluation import (
FeatureScaleSpec,
SplitData,
censoring_km_survival_at_times,
fast_ipcw_auc_brier_at_horizons,
field_indices_for_metrics,
fast_harrell_uno_concordance,
load_feature_scale_spec,
locf_field_values,
predict_xgboost_aft,
relapse_metric_verification_checks,
relapse_labels_from_arrays,
survival_integrated_brier_rows,
write_feature_scaling_audit,
)
def _split_for_relapse(cat: np.ndarray, missing: np.ndarray, terminal: np.ndarray) -> SplitData:
npz = {
"cat_value_ids": cat,
"missing_ids": missing,
"terminal_label": terminal,
}
row_df = pd.DataFrame(
{
"row_id": ["r0", "r1", "r2"],
"patient_id": ["p0", "p1", "p2"],
"window_index": [0, 1, 2],
"position": [0, 0, 0],
"current_visit_index": [0, 0, 0],
"split": ["val", "val", "val"],
}
)
return SplitData(
split="val",
npz=npz,
row_df=row_df,
X=np.zeros((3, 1), dtype=np.float32),
row_key_hash="dummy",
patient_ids=np.array(["p0", "p1", "p2"]),
window_index=np.array([0, 1, 2]),
position=np.array([0, 0, 0]),
time_to_death=np.zeros(3, dtype=np.float32),
event_death=np.zeros(3, dtype=bool),
time_to_disengagement=np.zeros(3, dtype=np.float32),
event_disengagement=np.zeros(3, dtype=bool),
censor_time=np.zeros(3, dtype=np.float32),
y_by_endpoint_horizon={},
)
def test_load_feature_scale_spec_uses_global_vocab_denominators(tmp_path: Path) -> None:
(tmp_path / "tensor_metadata.json").write_text(json.dumps({"cat_vocab_size": 7, "static_vocab_size": 11}), encoding="utf-8")
(tmp_path / "cat_value_vocab.json").write_text(json.dumps({"a": 0, "b": 1}), encoding="utf-8")
(tmp_path / "static_value_vocab.json").write_text(json.dumps({"s": 0}), encoding="utf-8")
spec = load_feature_scale_spec(tmp_path)
assert spec.cat_id_denominator == 7.0
assert spec.static_id_denominator == 11.0
assert spec.policy == "fixed_global_vocab_denominator_id_coded_tabular_baseline"
def test_feature_scaling_audit_requires_ids_below_global_denominator(tmp_path: Path) -> None:
spec = FeatureScaleSpec(
cat_id_denominator=7.0,
static_id_denominator=11.0,
cat_vocab_size=7,
static_vocab_size=11,
tensor_metadata_sha256="m",
cat_value_vocab_sha256="c",
static_value_vocab_sha256="s",
)
data = SplitData(
split="train",
npz={
"cat_value_ids": np.array([[[6]]], dtype=np.int64),
"static_value_ids": np.array([[10]], dtype=np.int64),
},
row_df=pd.DataFrame(),
X=np.zeros((1, 1), dtype=np.float32),
row_key_hash="dummy",
patient_ids=np.array(["p"]),
window_index=np.array([0]),
position=np.array([0]),
time_to_death=np.zeros(1, dtype=np.float32),
event_death=np.zeros(1, dtype=bool),
time_to_disengagement=np.zeros(1, dtype=np.float32),
event_disengagement=np.zeros(1, dtype=bool),
censor_time=np.zeros(1, dtype=np.float32),
y_by_endpoint_horizon={},
)
report = write_feature_scaling_audit(tmp_path, spec, [data])
assert report["status"] == "pass"
audit = pd.read_csv(tmp_path / "baseline_feature_scaling_audit.csv")
assert set(audit["status"]) == {"pass"}
assert audit.loc[audit["split"] == "train", "cat_id_scale_denominator"].iloc[0] == 7.0
def test_relapse_evaluable_matches_step4_active_next_visit_rule() -> None:
# field 0 = 是否转诊, field 1 = 转诊类型; positions 0=current, 1=next.
cat = np.zeros((3, 2, 2), dtype=np.int64)
missing = np.full((3, 2, 2), 3, dtype=np.int64)
terminal = np.zeros((3, 2), dtype=np.int64)
spec = {
"referral_field_index": 0,
"type_field_index": 1,
"referral_yes_value_id": 33,
"community_to_hospital_value_id": 35,
"label_definition": "next_visit(是否转诊=是 AND 转诊类型=社区转医院)",
}
cat[:, 1, 0] = np.array([32, 33, 33]) # no, yes, yes
cat[:, 1, 1] = np.array([0, 35, 35])
missing[:, 1, 0] = 0
missing[:, 1, 1] = np.array([3, 0, 0])
terminal[2, 1] = 3
y, evaluable, _ = relapse_labels_from_arrays(_split_for_relapse(cat, missing, terminal), spec, next_visit=True)
assert y.tolist() == [0, 1, 1]
assert evaluable.tolist() == [True, True, False]
def test_survival_integrated_brier_rows_are_metric_rows() -> None:
rows = survival_integrated_brier_rows(
[
{
"model_name": "m",
"endpoint": "death",
"split": "test",
"ibs_raw": 0.12,
"rowset_hash": "h",
"raw_path": "raw.parquet",
"raw_sha256": "abc",
},
{"model_name": "relapse", "endpoint": "primary_referral_relapse_next_visit", "split": "test"},
]
)
assert len(rows) == 1
assert rows[0]["integrated_brier_score"] == 0.12
assert rows[0]["calibration_status"] == "raw"
def test_xgboost_aft_prediction_uses_fixed_train_scaling() -> None:
pytest.importorskip("xgboost", reason="xgboost missing; install requirements-step5.txt to run AFT tests")
class FakeBooster:
def predict(self, dmat):
n = int(dmat.num_row())
if n == 1:
return np.asarray([10.0], dtype=np.float32)
return np.asarray([10.0, 20.0, 30.0], dtype=np.float32)[:n]
fitted = (
"xgboost_aft_survival",
FakeBooster(),
np.asarray([0.1, 0.2, 0.3, 0.4], dtype=np.float32),
0.0,
2.0,
10.0,
)
one = predict_xgboost_aft(fitted, np.zeros((1, 3), dtype=np.float32))[0]
batch = predict_xgboost_aft(fitted, np.zeros((3, 3), dtype=np.float32))[0]
np.testing.assert_allclose(one, batch, rtol=0.0, atol=0.0)
def test_true_locf_uses_last_observed_value_when_current_missing() -> None:
cat = np.zeros((2, 3, 1), dtype=np.int64)
missing = np.full((2, 3, 1), 3, dtype=np.int64)
cat[0, 0, 0] = 2
cat[0, 1, 0] = 0
missing[0, 0, 0] = 0
missing[0, 1, 0] = 3
cat[1, 0, 0] = 0
missing[1, 0, 0] = 3
data = SplitData(
split="val",
npz={
"cat_value_ids": cat,
"missing_ids": missing,
"valid_mask": np.ones((2, 3), dtype=bool),
},
row_df=pd.DataFrame(
{
"row_id": ["r0", "r1"],
"patient_id": ["p0", "p1"],
"window_index": [0, 1],
"position": [1, 0],
"current_visit_index": [1, 0],
"split": ["val", "val"],
}
),
X=np.zeros((2, 1), dtype=np.float32),
row_key_hash="dummy",
patient_ids=np.array(["p0", "p1"]),
window_index=np.array([0, 1]),
position=np.array([1, 0]),
time_to_death=np.zeros(2, dtype=np.float32),
event_death=np.zeros(2, dtype=bool),
time_to_disengagement=np.zeros(2, dtype=np.float32),
event_disengagement=np.zeros(2, dtype=bool),
censor_time=np.zeros(2, dtype=np.float32),
y_by_endpoint_horizon={},
)
locf, available = locf_field_values(data, field_family="categorical", field_index=0, fallback_value=7)
assert locf.tolist() == [2, 7]
assert available.tolist() == [True, False]
assert int(cat[0, 1, 0]) == 0
def test_next_visit_field_indices_use_tensor_metadata_not_metric_order(tmp_path: Path) -> None:
(tmp_path / "tensor_metadata.json").write_text(
json.dumps(
{
"cat_cols": ["a", "b", "c"],
"ord_cols": ["o0", "o1"],
"num_cols": ["n0", "n1", "n2"],
},
ensure_ascii=False,
),
encoding="utf-8",
)
metrics = pd.DataFrame(
[
{"field_family": "categorical", "field": "c"},
{"field_family": "categorical", "field": "a"},
{"field_family": "ordinal", "field": "o1"},
{"field_family": "numeric", "field": "n2"},
{"field_family": "numeric", "field": "n0"},
]
)
indices = field_indices_for_metrics(metrics, tmp_path, tmp_path)
assert indices["categorical"] == [("c", 2), ("a", 0)]
assert indices["ordinal"] == [("o1", 1)]
assert indices["numeric"] == [("n2", 2), ("n0", 0)]
audit = pd.read_csv(tmp_path / "next_visit_field_order_audit.csv")
assert set(audit["status"]) == {"pass"}
assert int((audit["metric_order_equals_tensor_index"].astype(str).str.lower() == "false").sum()) > 0
def test_fast_concordance_matches_pairwise_definition() -> None:
from sksurv.metrics import concordance_index_censored, concordance_index_ipcw
from sksurv.util import Surv
train_time = np.asarray([1.0, 2.0, 2.0, 4.0, 5.0, 8.0], dtype=np.float64)
train_event = np.asarray([True, False, True, False, True, False], dtype=bool)
test_time = np.asarray([1.0, 2.0, 2.0, 3.0, 5.0, 7.0], dtype=np.float64)
test_event = np.asarray([True, True, False, False, True, False], dtype=bool)
score = np.asarray([0.9, 0.4, 0.4, 0.6, 0.7, 0.1], dtype=np.float64)
harrell, uno = fast_harrell_uno_concordance(test_time, test_event, score, train_time, train_event)
y_train = Surv.from_arrays(train_event, train_time)
y_test = Surv.from_arrays(test_event, test_time)
assert harrell == pytest.approx(concordance_index_censored(test_event, test_time, score)[0])
assert uno == pytest.approx(concordance_index_ipcw(y_train, y_test, score)[0])
def test_fast_ipcw_auc_brier_matches_bruteforce_definition() -> None:
from sksurv.metrics import brier_score, cumulative_dynamic_auc
from sksurv.nonparametric import CensoringDistributionEstimator
from sksurv.util import Surv
train_time = np.asarray([1.0, 2.0, 2.0, 4.0, 5.0, 8.0], dtype=np.float64)
train_event = np.asarray([True, False, True, False, True, False], dtype=bool)
test_time = np.asarray([1.0, 2.0, 2.0, 3.0, 5.0, 7.0], dtype=np.float64)
test_event = np.asarray([True, True, False, False, True, False], dtype=bool)
horizons = np.asarray([2.0, 4.0], dtype=np.float64)
risks = np.asarray(
[
[0.8, 0.9],
[0.6, 0.7],
[0.6, 0.5],
[0.3, 0.4],
[0.2, 0.5],
[0.1, 0.2],
],
dtype=np.float64,
)
aucs, briers, ibs = fast_ipcw_auc_brier_at_horizons(train_time, train_event, test_time, test_event, risks, horizons)
y_train = Surv.from_arrays(train_event, train_time)
y_test = Surv.from_arrays(test_event, test_time)
expected_auc, _ = cumulative_dynamic_auc(y_train, y_test, risks, horizons)
_, expected_briers = brier_score(y_train, y_test, 1.0 - risks, horizons)
expected_g = CensoringDistributionEstimator().fit(y_train).predict_proba(test_time)
actual_g = censoring_km_survival_at_times(train_time, train_event, test_time)
np.testing.assert_allclose(actual_g, expected_g)
for j, h in enumerate(horizons):
assert aucs[float(h)] == pytest.approx(float(expected_auc[j]))
assert briers[float(h)] == pytest.approx(float(expected_briers[j]))
expected_ibs = np.trapezoid(np.asarray(expected_briers), horizons) / (horizons[-1] - horizons[0])
assert ibs == pytest.approx(expected_ibs)
def test_relapse_metric_verifier_requires_raw_and_calibrated_rows() -> None:
rows = []
for model in ["sctm_relapse_next_visit", "clinical_current_anchor_relapse", "gbm_relapse_next_visit"]:
for split in ["val", "test"]:
for stratum in ["all", "changed", "unchanged"]:
for status in ["raw", "calibrated_validation_only"]:
rows.append(
{
"model_name": model,
"endpoint": "primary_referral_relapse_next_visit",
"split": split,
"metric_prefix": stratum,
"calibration_status": status,
"n": 10,
"n_events": 2,
"rowset_hash": "full",
"evaluable_rowset_hash": "eval",
"static_spec_hash": "static",
}
)
checks = relapse_metric_verification_checks(pd.DataFrame(rows))
assert {c["status"] for c in checks} == {"pass"}
broken = pd.DataFrame(rows).query("not (model_name == 'sctm_relapse_next_visit' and split == 'test' and metric_prefix == 'all' and calibration_status == 'calibrated_validation_only')")
broken_checks = relapse_metric_verification_checks(broken)
assert any(c["status"] == "fail" for c in broken_checks)