"""Triage the earliest evidenced non-finite neural-training boundary. The ordered telemetry supports containment and replay. It deliberately reports no root cause: a non-finite symptom can have multiple upstream mechanisms. """ from __future__ import annotations from dataclasses import asdict, dataclass import hashlib import json import math import re from typing import Any, Iterable STAGES = ("input", "forward", "loss", "backward", "optimizer") MAX_ELEMENTS_PER_STAGE = 10_000_000_000 MAX_TOTAL_ELEMENTS = 20_000_000_000 MIN_BINARY64_RESOLUTION = 1e-15 MAX_BINARY64_RESOLUTION = 1e-3 _IDENTIFIER = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._:/@-]{0,127}$") def _identifier(name: str, value: object) -> str: if type(value) is not str or not _IDENTIFIER.fullmatch(value): raise ValueError(f"{name} must be a stable identifier") return value def _owner(name: str, value: object) -> str: checked = _identifier(name, value) if not checked.startswith("team:"): raise ValueError(f"{name} must identify an accountable team owner") return checked def _finite(name: str, value: object) -> float: if isinstance(value, bool) or not isinstance(value, (int, float)): raise ValueError(f"{name} must be a real number") checked = float(value) if not math.isfinite(checked): raise ValueError(f"{name} must be finite") return checked def _integer(name: str, value: object, *, minimum: int, maximum: int) -> int: if type(value) is not int or not minimum <= value <= maximum: raise ValueError(f"{name} must be an integer in [{minimum}, {maximum}]") return value def _canonical(value: Any) -> Any: if isinstance(value, dict): return {key: _canonical(item) for key, item in sorted(value.items())} if isinstance(value, tuple): return [_canonical(item) for item in value] return value def _content_id(prefix: str, payload: dict[str, Any]) -> str: encoded = json.dumps( _canonical(payload), sort_keys=True, separators=(",", ":"), allow_nan=False ).encode("utf-8") return f"{prefix}@sha256:{hashlib.sha256(encoded).hexdigest()}" @dataclass(frozen=True) class TrainingIncidentContract: contract_version: str run_id: str step_id: int model_version: str dataset_version: str batch_id: str objective_version: str optimizer_version: str precision: str loss_scale_version: str loss_scale: float minimum_numeric_resolution: float maximum_absolute_finite: float maximum_elements_per_stage: int maximum_total_elements: int numeric_convention: str data_owner: str model_owner: str objective_owner: str optimizer_owner: str incident_owner: str def __post_init__(self) -> None: for field in ( "contract_version", "run_id", "model_version", "dataset_version", "batch_id", "objective_version", "optimizer_version", "precision", "loss_scale_version", ): _identifier(field, getattr(self, field)) for field in ( "data_owner", "model_owner", "objective_owner", "optimizer_owner", "incident_owner", ): _owner(field, getattr(self, field)) _integer("step_id", self.step_id, minimum=0, maximum=1_000_000_000_000) per_stage = _integer( "maximum_elements_per_stage", self.maximum_elements_per_stage, minimum=1, maximum=MAX_ELEMENTS_PER_STAGE, ) total = _integer( "maximum_total_elements", self.maximum_total_elements, minimum=1, maximum=MAX_TOTAL_ELEMENTS, ) if total < per_stage or total > len(STAGES) * per_stage: raise ValueError("maximum_total_elements is inconsistent with stage bound") resolution = _finite( "minimum_numeric_resolution", self.minimum_numeric_resolution ) finite_limit = _finite( "maximum_absolute_finite", self.maximum_absolute_finite ) loss_scale = _finite("loss_scale", self.loss_scale) if not MIN_BINARY64_RESOLUTION <= resolution <= MAX_BINARY64_RESOLUTION: raise ValueError("minimum_numeric_resolution is outside the safe range") if finite_limit < resolution or finite_limit > 1.0 / resolution: raise ValueError("maximum_absolute_finite is outside the safe range") if loss_scale < resolution or loss_scale > 1.0 / resolution: raise ValueError("loss_scale is outside the safe range") if self.numeric_convention != "binary64-summary-mixed-precision-v1": raise ValueError("unsupported numeric_convention") @property def content_id(self) -> str: return _content_id("training-incident-contract", asdict(self)) def expected_owner(self, stage: str) -> str: return { "input": self.data_owner, "forward": self.model_owner, "loss": self.objective_owner, "backward": self.model_owner, "optimizer": self.optimizer_owner, }[stage] @dataclass(frozen=True) class StageTelemetry: telemetry_id: str tensor_set_id: str stage: str stage_order: int state: str checked_elements: int finite_elements: int nonfinite_elements: int maximum_absolute_finite: float run_id: str step_id: int model_version: str dataset_version: str batch_id: str objective_version: str optimizer_version: str precision: str loss_scale_version: str loss_scale: float evidence_source_id: str owner: str def __post_init__(self) -> None: for field in ( "telemetry_id", "tensor_set_id", "run_id", "model_version", "dataset_version", "batch_id", "objective_version", "optimizer_version", "precision", "loss_scale_version", "evidence_source_id", ): _identifier(field, getattr(self, field)) _owner("owner", self.owner) if self.stage not in STAGES: raise ValueError("stage is not in the ordered training path") if self.state not in ("OBSERVED", "NOT_RUN"): raise ValueError("state must be OBSERVED or NOT_RUN") _integer("stage_order", self.stage_order, minimum=0, maximum=len(STAGES) - 1) _integer("step_id", self.step_id, minimum=0, maximum=1_000_000_000_000) for field in ("checked_elements", "finite_elements", "nonfinite_elements"): _integer( field, getattr(self, field), minimum=0, maximum=MAX_ELEMENTS_PER_STAGE, ) maximum = _finite("maximum_absolute_finite", self.maximum_absolute_finite) if maximum < 0.0: raise ValueError("maximum_absolute_finite must be non-negative") loss_scale = _finite("loss_scale", self.loss_scale) if loss_scale <= 0.0: raise ValueError("loss_scale must be positive") if self.state == "OBSERVED": if self.checked_elements == 0: raise ValueError("OBSERVED telemetry must check at least one element") if self.finite_elements + self.nonfinite_elements != self.checked_elements: raise ValueError("element counts do not reconcile") elif any( value != 0 for value in ( self.checked_elements, self.finite_elements, self.nonfinite_elements, self.maximum_absolute_finite, ) ): raise ValueError("NOT_RUN telemetry cannot contain observed values") @classmethod def capture(cls, **values: object) -> "StageTelemetry": return cls( telemetry_id=_identifier("telemetry_id", values["telemetry_id"]), tensor_set_id=_identifier("tensor_set_id", values["tensor_set_id"]), stage=_identifier("stage", values["stage"]), stage_order=_integer( "stage_order", values["stage_order"], minimum=0, maximum=len(STAGES) - 1 ), state=_identifier("state", values["state"]), checked_elements=_integer( "checked_elements", values["checked_elements"], minimum=0, maximum=MAX_ELEMENTS_PER_STAGE, ), finite_elements=_integer( "finite_elements", values["finite_elements"], minimum=0, maximum=MAX_ELEMENTS_PER_STAGE, ), nonfinite_elements=_integer( "nonfinite_elements", values["nonfinite_elements"], minimum=0, maximum=MAX_ELEMENTS_PER_STAGE, ), maximum_absolute_finite=_finite( "maximum_absolute_finite", values["maximum_absolute_finite"] ), run_id=_identifier("run_id", values["run_id"]), step_id=_integer( "step_id", values["step_id"], minimum=0, maximum=1_000_000_000_000 ), model_version=_identifier("model_version", values["model_version"]), dataset_version=_identifier("dataset_version", values["dataset_version"]), batch_id=_identifier("batch_id", values["batch_id"]), objective_version=_identifier( "objective_version", values["objective_version"] ), optimizer_version=_identifier( "optimizer_version", values["optimizer_version"] ), precision=_identifier("precision", values["precision"]), loss_scale_version=_identifier( "loss_scale_version", values["loss_scale_version"] ), loss_scale=_finite("loss_scale", values["loss_scale"]), evidence_source_id=_identifier( "evidence_source_id", values["evidence_source_id"] ), owner=_owner("owner", values["owner"]), ) @dataclass(frozen=True) class TrainingRunEvidence: evidence_id: str evidence_content_id: str contract_content_id: str stages: tuple[StageTelemetry, ...] def __post_init__(self) -> None: _identifier("evidence_id", self.evidence_id) _identifier("evidence_content_id", self.evidence_content_id) _identifier("contract_content_id", self.contract_content_id) if type(self.stages) is not tuple: raise TypeError("stages must be an immutable tuple") if len(self.stages) != len(STAGES): raise ValueError("evidence must contain every ordered training stage") if any(type(item) is not StageTelemetry for item in self.stages): raise TypeError("stages must contain concrete StageTelemetry values") @classmethod def capture( cls, *, evidence_id: str, contract: TrainingIncidentContract, stages: Iterable[StageTelemetry], ) -> "TrainingRunEvidence": if type(contract) is not TrainingIncidentContract: raise TypeError("contract must be a concrete TrainingIncidentContract") contract.__post_init__() frozen = tuple(stages) selected_id = _identifier("evidence_id", evidence_id) return cls( evidence_id=selected_id, evidence_content_id=_evidence_content_id( contract.content_id, selected_id, frozen ), contract_content_id=contract.content_id, stages=frozen, ) @dataclass(frozen=True) class IncidentReport: contract_content_id: str evidence_id: str evidence_content_id: str decision: str earliest_nonfinite_stage: str | None action: str root_cause_claim: str def _evidence_content_id( contract_content_id: str, evidence_id: str, stages: tuple[StageTelemetry, ...], ) -> str: return _content_id( "training-incident-evidence", { "contract_content_id": contract_content_id, "evidence_id": evidence_id, "stages": tuple(asdict(item) for item in stages), }, ) def triage_training_incident( contract: TrainingIncidentContract, evidence: TrainingRunEvidence ) -> IncidentReport: if type(contract) is not TrainingIncidentContract: raise TypeError("contract must be a concrete TrainingIncidentContract") if type(evidence) is not TrainingRunEvidence: raise TypeError("evidence must be concrete TrainingRunEvidence") contract.__post_init__() evidence.__post_init__() if evidence.contract_content_id != contract.content_id: raise ValueError("evidence belongs to a different contract") expected_content_id = _evidence_content_id( contract.content_id, evidence.evidence_id, evidence.stages ) if evidence.evidence_content_id != expected_content_id: raise ValueError("evidence content identity does not match its stages") telemetry_ids: set[str] = set() tensor_set_ids: set[str] = set() earliest: str | None = None saw_not_run = False total_elements = 0 scope = ( contract.run_id, contract.step_id, contract.model_version, contract.dataset_version, contract.batch_id, contract.objective_version, contract.optimizer_version, contract.precision, contract.loss_scale_version, contract.loss_scale, ) for order, telemetry in enumerate(evidence.stages): telemetry.__post_init__() if telemetry.stage != STAGES[order] or telemetry.stage_order != order: raise ValueError("training stages are missing or out of order") if telemetry.telemetry_id in telemetry_ids: raise ValueError("duplicate telemetry identity") if telemetry.tensor_set_id in tensor_set_ids: raise ValueError("duplicate tensor-set identity") telemetry_ids.add(telemetry.telemetry_id) tensor_set_ids.add(telemetry.tensor_set_id) observed_scope = ( telemetry.run_id, telemetry.step_id, telemetry.model_version, telemetry.dataset_version, telemetry.batch_id, telemetry.objective_version, telemetry.optimizer_version, telemetry.precision, telemetry.loss_scale_version, telemetry.loss_scale, ) if observed_scope != scope: raise ValueError("stage scope does not match the incident contract") if telemetry.owner != contract.expected_owner(telemetry.stage): raise ValueError("stage owner does not match the incident contract") if telemetry.checked_elements > contract.maximum_elements_per_stage: raise ValueError("stage exceeds maximum_elements_per_stage") total_elements += telemetry.checked_elements maximum = _finite( "maximum_absolute_finite", telemetry.maximum_absolute_finite ) if maximum > contract.maximum_absolute_finite: raise ValueError("finite telemetry exceeds maximum_absolute_finite") if maximum != 0.0 and maximum < contract.minimum_numeric_resolution: raise ValueError("finite telemetry is below numeric resolution") if telemetry.state == "NOT_RUN": saw_not_run = True if earliest is None: raise ValueError("NOT_RUN stage has no prior non-finite boundary") continue if saw_not_run: raise ValueError("an observed stage cannot follow NOT_RUN telemetry") if telemetry.nonfinite_elements > 0 and earliest is None: earliest = telemetry.stage if total_elements > contract.maximum_total_elements: raise ValueError("evidence exceeds maximum_total_elements") actions = { "input": "hold-and-inspect-input-data-contract", "forward": "hold-and-replay-forward-activations", "loss": "hold-and-inspect-objective-reduction", "backward": "hold-and-inspect-gradient-path-and-loss-scaling", "optimizer": "hold-and-inspect-optimizer-state-and-update", } return IncidentReport( contract_content_id=contract.content_id, evidence_id=evidence.evidence_id, evidence_content_id=evidence.evidence_content_id, decision="HOLD" if earliest is not None else "CONTINUE_BOUNDED_DIAGNOSTICS", earliest_nonfinite_stage=earliest, action=actions[earliest] if earliest is not None else "no-nonfinite-evidenced-check-other-signals", root_cause_claim="NONE_SYMPTOM_LOCATION_IS_NOT_CAUSALITY", ) def _example() -> None: contract = TrainingIncidentContract( contract_version="training-incident-v1", run_id="run-2026-08-25-001", step_id=1842, model_version="model-v12", dataset_version="data-v8", batch_id="batch-1842", objective_version="next-token-v3", optimizer_version="adamw-v5", precision="bf16-with-fp32-reductions", loss_scale_version="dynamic-scale-v2", loss_scale=1024.0, minimum_numeric_resolution=1e-12, maximum_absolute_finite=1e9, maximum_elements_per_stage=1_000_000, maximum_total_elements=5_000_000, numeric_convention="binary64-summary-mixed-precision-v1", data_owner="team:data", model_owner="team:model", objective_owner="team:objective", optimizer_owner="team:optimizer", incident_owner="team:oncall", ) stages: list[StageTelemetry] = [] for order, stage in enumerate(STAGES): not_run = stage == "optimizer" nonfinite = 2 if stage == "backward" else 0 checked = 0 if not_run else 4096 stages.append( StageTelemetry.capture( telemetry_id=f"telemetry-{stage}", tensor_set_id=f"tensor-set-{stage}", stage=stage, stage_order=order, state="NOT_RUN" if not_run else "OBSERVED", checked_elements=checked, finite_elements=checked - nonfinite, nonfinite_elements=nonfinite, maximum_absolute_finite=0.0 if not_run else float(order + 1), run_id=contract.run_id, step_id=contract.step_id, model_version=contract.model_version, dataset_version=contract.dataset_version, batch_id=contract.batch_id, objective_version=contract.objective_version, optimizer_version=contract.optimizer_version, precision=contract.precision, loss_scale_version=contract.loss_scale_version, loss_scale=contract.loss_scale, evidence_source_id=f"probe-{stage}-v1", owner=contract.expected_owner(stage), ) ) evidence = TrainingRunEvidence.capture( evidence_id="incident-evidence-001", contract=contract, stages=stages ) report = triage_training_incident(contract, evidence) print("example=illustrative_only") print(f"contract_version={contract.contract_version}") print(f"run={contract.run_id}") print(f"step={contract.step_id}") print(f"evidence_content_id={report.evidence_content_id}") print(f"earliest_nonfinite={report.earliest_nonfinite_stage}") print(f"decision={report.decision}") print(f"action={report.action}") print(f"root_cause_claim={report.root_cause_claim}") if __name__ == "__main__": _example()