sctm2-r5-code-20260611 / tests /test_tensor_bridge_artifacts.py
agoniii97's picture
Normalize datetime precision for HF P1 tensor build
ae6d94c verified
Raw
History Blame Contribute Delete
2.51 kB
from __future__ import annotations
import json
import os
from pathlib import Path
import numpy as np
import pandas as pd
import pytest
pytestmark = pytest.mark.integration
def test_tensor_bridge_smoke_artifacts_contract() -> None:
out = os.environ.get("SCTM_TENSOR_BRIDGE_OUTPUT")
if not out:
return
root = Path(out)
tensor_dir = root / "tensor"
stage0_dir = root / "stage0"
manifest = json.loads((root / "tensor_builder_manifest.json").read_text(encoding="utf-8"))
split_counts = pd.read_csv(root / "tensor_split_counts.csv")
meta = json.loads((tensor_dir / "tensor_metadata.json").read_text(encoding="utf-8"))
vocab = json.loads((tensor_dir / "cat_value_vocab.json").read_text(encoding="utf-8"))
assert "变动情况" not in "".join(meta["cat_cols"])
assert "失访原因" not in "".join(meta["cat_cols"])
assert not any("变动情况=" in key or "失访原因=" in key for key in vocab)
for field in [
"治疗方式",
"服药方式",
"服药时间",
"是否免费服药",
"是否转诊",
"转诊类型",
"随访类型",
"本次访视方式",
"本次访视对象",
"是否应急处置",
"健康体检",
"专科医生意见-治疗",
"基础管理分级",
]:
assert field in meta["cat_cols"]
assert "公安滋事总数" in meta["num_cols"]
assert "公安肇事总数" in meta["num_cols"]
assert "has_static" in meta["static_cols"]
assert manifest["n_windowed_patients"] == manifest["n_patients_used"]
assert {"n_singleton_patients", "n_transition_eligible_windows", "n_windowed_patients"}.issubset(split_counts.columns)
assert int(split_counts["n_windowed_patients"].sum()) == manifest["n_patients_used"]
assert int(split_counts["n_transition_eligible_windows"].sum()) > 0
assert (stage0_dir / "ordinal_direction_table.json").exists()
assert (stage0_dir / "service_state_transitions_train.json").exists()
with np.load(Path(meta["split_paths"]["val"]), allow_pickle=True) as z:
assert z["cat_value_ids"].shape[1:] == (meta["max_seq_len"], len(meta["cat_cols"]))
assert z["numeric_values"].shape[1:] == (meta["max_seq_len"], len(meta["num_cols"]))
assert z["ordinal_cbe"].shape[1:] == (meta["max_seq_len"], len(meta["ord_cols"]), meta["cbe_dim"])
assert z["static_value_ids"].shape[1] == len(meta["static_cols"])
assert z["valid_mask"].sum() > 0