| 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 |
|
|