sctm2-r5-code-20260611 / tests /test_sctm_test_plugin.py
agoniii97's picture
Normalize datetime precision for HF P1 tensor build
ae6d94c verified
Raw
History Blame Contribute Delete
3.22 kB
from __future__ import annotations
import inspect
import torch
from sctm_next_round import canonical_outputs, strict_temporal_stress, train_formal_sctm
from sctm_next_round.formal_model import SCTMTestFormalModel, ZeroInitSCTMTestAdapter
from sctm_next_round.sctm_model_plugins import MODEL_PLUGIN_REGISTRY, build_sctm_model, normalize_model_plugin_id
from sctm_next_round import sctm_test_plugin_smoke
def test_sctm_test_registry_is_a_narrow_plugin_boundary() -> None:
assert MODEL_PLUGIN_REGISTRY["canonical"].__name__ == "SCTMFormalModel"
assert MODEL_PLUGIN_REGISTRY["SCTM-test"] is SCTMTestFormalModel
assert normalize_model_plugin_id(None) == "canonical"
assert normalize_model_plugin_id("SCTM-test") == "SCTM-test"
source = inspect.getsource(build_sctm_model)
assert "MODEL_PLUGIN_REGISTRY" in source
assert "return cls(" in source
def test_sctm_test_adapter_is_exact_identity_at_initialization() -> None:
torch.manual_seed(20260608)
adapter = ZeroInitSCTMTestAdapter(d_model=12, dropout=0.0)
z = torch.randn(3, 5, 12)
out = adapter(z)
assert torch.allclose(out, z, atol=0.0, rtol=0.0)
assert any(name.startswith("net.") for name, _ in adapter.named_parameters())
def test_sctm_test_forward_only_replaces_direct_readout_branch() -> None:
source = inspect.getsource(SCTMTestFormalModel.forward)
assert "super().forward" in source
assert "sctm_test_direct_readout_adapter(z_joint)" in source
assert "self.direct_horizon_risk_head(z_test)" in source
assert "self.direct_relapse_head(z_test)" in source
assert "direct_readout_source" in source
assert "diagnostic_plugin_compatibility_only" in source
assert "teacher" not in source.lower()
assert "next_active" not in source
assert "next_delta" not in source
def test_step3_step4_and_stress_entrypoints_use_model_plugin_factory() -> None:
train_source = inspect.getsource(train_formal_sctm.run)
assert "build_sctm_model" in train_source
assert "plugin_id=args.model_plugin" in train_source
assert '"model_plugin_id": args.model_plugin' in train_source
assert "--model-plugin" in inspect.getsource(train_formal_sctm.build_parser)
step4_source = inspect.getsource(canonical_outputs.instantiate_model)
assert "build_sctm_model" in step4_source
assert 'plugin_id=context.get("model_plugin_id")' in step4_source
assert "strict=True" in step4_source
assert "--model-plugin" in inspect.getsource(canonical_outputs.parse_args)
strict_source = inspect.getsource(strict_temporal_stress.load_model_context)
assert "build_sctm_model" in strict_source
assert "plugin_id=args.model_plugin" in strict_source
assert "--model-plugin" in inspect.getsource(strict_temporal_stress.main)
def test_sctm_test_smoke_records_step5_light_claim_boundary() -> None:
smoke_source = inspect.getsource(sctm_test_plugin_smoke)
assert "step5_light_sctm_test_fixture" in smoke_source
assert "sctm_rows_identity_compatible" in smoke_source
assert "baseline_predictions_reused" in smoke_source
assert "pipeline_contract_audit" in smoke_source
assert "not a retrained SCTM-test performance result" in smoke_source