"""Pure-Python cross-correlation and a bounded translation-equivariance audit. Only points whose translated receptive fields remain inside the declared valid domain are compared. Boundary exclusion is reported, not hidden. """ from __future__ import annotations from dataclasses import asdict, dataclass import hashlib import json import math import re from typing import Any, Iterable MAX_SIDE = 256 MAX_CASES = 256 MAX_OPERATIONS = 10_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()}" def _matrix( name: str, value: object, *, height: int, width: int, maximum_absolute_value: float, minimum_numeric_resolution: float, ) -> tuple[tuple[float, ...], ...]: if type(value) is not tuple or len(value) != height: raise TypeError(f"{name} must be an immutable tuple with declared height") rows: list[tuple[float, ...]] = [] for row in value: if type(row) is not tuple or len(row) != width: raise TypeError(f"{name} rows must be immutable tuples with declared width") checked_row: list[float] = [] for item in row: checked = _finite(name, item) if abs(checked) > maximum_absolute_value: raise ValueError(f"{name} exceeds the contracted magnitude") if checked != 0.0 and abs(checked) < minimum_numeric_resolution: raise ValueError(f"{name} contains a sub-resolution value") checked_row.append(checked) rows.append(tuple(checked_row)) return tuple(rows) @dataclass(frozen=True) class EquivarianceContract: contract_version: str operation: str layout: str padding: str stride_y: int stride_x: int dilation_y: int dilation_x: int translation_y: int translation_x: int input_height: int input_width: int kernel_height: int kernel_width: int kernel_id: str kernel_version: str kernel_values: tuple[tuple[float, ...], ...] maximum_cases: int maximum_operations: int minimum_numeric_resolution: float maximum_absolute_input: float maximum_absolute_output: float absolute_tolerance: float numeric_convention: str data_owner: str kernel_owner: str audit_owner: str def __post_init__(self) -> None: for field in ("contract_version", "kernel_id", "kernel_version"): _identifier(field, getattr(self, field)) for field in ("data_owner", "kernel_owner", "audit_owner"): _owner(field, getattr(self, field)) if self.operation != "cross-correlation-2d": raise ValueError("unsupported operation") if self.layout != "row-major-hw": raise ValueError("unsupported layout") if self.padding != "valid": raise ValueError("only explicit valid padding is supported") for field in ("input_height", "input_width", "kernel_height", "kernel_width"): _integer(field, getattr(self, field), minimum=1, maximum=MAX_SIDE) for field in ("stride_y", "stride_x", "dilation_y", "dilation_x"): _integer(field, getattr(self, field), minimum=1, maximum=MAX_SIDE) for field in ("translation_y", "translation_x"): _integer(field, getattr(self, field), minimum=-MAX_SIDE, maximum=MAX_SIDE) if self.translation_y == 0 and self.translation_x == 0: raise ValueError("translation must move at least one axis") effective_height = self.dilation_y * (self.kernel_height - 1) + 1 effective_width = self.dilation_x * (self.kernel_width - 1) + 1 if effective_height > self.input_height or effective_width > self.input_width: raise ValueError("effective kernel does not fit the declared input") _integer("maximum_cases", self.maximum_cases, minimum=1, maximum=MAX_CASES) _integer( "maximum_operations", self.maximum_operations, minimum=1, maximum=MAX_OPERATIONS, ) resolution = _finite( "minimum_numeric_resolution", self.minimum_numeric_resolution ) input_limit = _finite("maximum_absolute_input", self.maximum_absolute_input) output_limit = _finite( "maximum_absolute_output", self.maximum_absolute_output ) tolerance = _finite("absolute_tolerance", self.absolute_tolerance) if not MIN_BINARY64_RESOLUTION <= resolution <= MAX_BINARY64_RESOLUTION: raise ValueError("minimum_numeric_resolution is outside the safe range") if input_limit < resolution or input_limit > 1.0 / resolution: raise ValueError("maximum_absolute_input is outside the safe range") if output_limit < input_limit or output_limit > 1.0 / resolution: raise ValueError("maximum_absolute_output is outside the safe range") if tolerance < 0.0 or tolerance > resolution: raise ValueError("absolute_tolerance exceeds numeric resolution") if self.numeric_convention != "binary64-fsum-valid-domain-v1": raise ValueError("unsupported numeric_convention") checked_kernel = _matrix( "kernel_values", self.kernel_values, height=self.kernel_height, width=self.kernel_width, maximum_absolute_value=input_limit, minimum_numeric_resolution=resolution, ) object.__setattr__(self, "kernel_values", checked_kernel) @property def output_shape(self) -> tuple[int, int]: effective_height = self.dilation_y * (self.kernel_height - 1) + 1 effective_width = self.dilation_x * (self.kernel_width - 1) + 1 return ( (self.input_height - effective_height) // self.stride_y + 1, (self.input_width - effective_width) // self.stride_x + 1, ) @property def content_id(self) -> str: return _content_id("equivariance-contract", asdict(self)) @property def kernel_content_id(self) -> str: return _content_id( "convolution-kernel", { "kernel_id": self.kernel_id, "kernel_version": self.kernel_version, "kernel_height": self.kernel_height, "kernel_width": self.kernel_width, "kernel_values": self.kernel_values, "numeric_convention": self.numeric_convention, }, ) @dataclass(frozen=True) class ConvolutionCase: case_id: str input_id: str input_version: str kernel_id: str kernel_version: str kernel_content_id: str input_values: tuple[tuple[float, ...], ...] kernel_values: tuple[tuple[float, ...], ...] def __post_init__(self) -> None: for field in ( "case_id", "input_id", "input_version", "kernel_id", "kernel_version", "kernel_content_id", ): _identifier(field, getattr(self, field)) @classmethod def capture( cls, *, contract: EquivarianceContract, case_id: str, input_id: str, input_version: str, kernel_id: str, kernel_version: str, input_values: object, kernel_values: object, ) -> "ConvolutionCase": if type(contract) is not EquivarianceContract: raise TypeError("contract must be a concrete EquivarianceContract") contract.__post_init__() checked_kernel_id = _identifier("kernel_id", kernel_id) checked_kernel_version = _identifier("kernel_version", kernel_version) if (checked_kernel_id, checked_kernel_version) != ( contract.kernel_id, contract.kernel_version, ): raise ValueError("kernel identity does not match the contract") checked_kernel = _matrix( "kernel_values", kernel_values, height=contract.kernel_height, width=contract.kernel_width, maximum_absolute_value=contract.maximum_absolute_input, minimum_numeric_resolution=contract.minimum_numeric_resolution, ) if checked_kernel != contract.kernel_values: raise ValueError("kernel content does not match the contracted kernel") return cls( case_id=_identifier("case_id", case_id), input_id=_identifier("input_id", input_id), input_version=_identifier("input_version", input_version), kernel_id=checked_kernel_id, kernel_version=checked_kernel_version, kernel_content_id=contract.kernel_content_id, input_values=_matrix( "input_values", input_values, height=contract.input_height, width=contract.input_width, maximum_absolute_value=contract.maximum_absolute_input, minimum_numeric_resolution=contract.minimum_numeric_resolution, ), kernel_values=checked_kernel, ) @dataclass(frozen=True) class EquivarianceSuite: evidence_id: str evidence_content_id: str contract_content_id: str cases: tuple[ConvolutionCase, ...] 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.cases) is not tuple: raise TypeError("cases must be an immutable tuple") if not self.cases or len(self.cases) > MAX_CASES: raise ValueError("cases are empty or exceed the global safety limit") if any(type(item) is not ConvolutionCase for item in self.cases): raise TypeError("cases must contain concrete ConvolutionCase values") @classmethod def capture( cls, *, evidence_id: str, contract: EquivarianceContract, cases: Iterable[ConvolutionCase], ) -> "EquivarianceSuite": if type(contract) is not EquivarianceContract: raise TypeError("contract must be a concrete EquivarianceContract") contract.__post_init__() frozen = tuple(cases) if not frozen or len(frozen) > contract.maximum_cases: raise ValueError("cases are empty or exceed contract maximum_cases") if any(type(item) is not ConvolutionCase for item in frozen): raise TypeError("cases must contain concrete ConvolutionCase values") selected_id = _identifier("evidence_id", evidence_id) return cls( evidence_id=selected_id, evidence_content_id=_suite_content_id(contract.content_id, selected_id, frozen), contract_content_id=contract.content_id, cases=frozen, ) @dataclass(frozen=True) class CaseResult: case_id: str compared_points: int excluded_boundary_points: int maximum_absolute_error: float status: str @dataclass(frozen=True) class EquivarianceReport: contract_content_id: str kernel_content_id: str evidence_id: str evidence_content_id: str output_shape: tuple[int, int] status: str assumption: str cases: tuple[CaseResult, ...] def _suite_content_id( contract_content_id: str, evidence_id: str, cases: tuple[ConvolutionCase, ...], ) -> str: return _content_id( "equivariance-evidence", { "contract_content_id": contract_content_id, "evidence_id": evidence_id, "cases": tuple(asdict(item) for item in cases), }, ) def cross_correlate_2d( contract: EquivarianceContract, input_values: tuple[tuple[float, ...], ...], kernel_values: tuple[tuple[float, ...], ...], ) -> tuple[tuple[float, ...], ...]: if type(contract) is not EquivarianceContract: raise TypeError("contract must be a concrete EquivarianceContract") contract.__post_init__() checked_input = _matrix( "input_values", input_values, height=contract.input_height, width=contract.input_width, maximum_absolute_value=contract.maximum_absolute_input, minimum_numeric_resolution=contract.minimum_numeric_resolution, ) checked_kernel = _matrix( "kernel_values", kernel_values, height=contract.kernel_height, width=contract.kernel_width, maximum_absolute_value=contract.maximum_absolute_input, minimum_numeric_resolution=contract.minimum_numeric_resolution, ) if checked_kernel != contract.kernel_values: raise ValueError("kernel content does not match the contracted kernel") out_height, out_width = contract.output_shape operations = out_height * out_width * contract.kernel_height * contract.kernel_width if operations > contract.maximum_operations: raise ValueError("cross-correlation exceeds maximum_operations") output: list[tuple[float, ...]] = [] for out_y in range(out_height): row: list[float] = [] for out_x in range(out_width): terms: list[float] = [] for kernel_y in range(contract.kernel_height): for kernel_x in range(contract.kernel_width): input_y = out_y * contract.stride_y + kernel_y * contract.dilation_y input_x = out_x * contract.stride_x + kernel_x * contract.dilation_x term = checked_input[input_y][input_x] * checked_kernel[kernel_y][kernel_x] if not math.isfinite(term): raise OverflowError("cross-correlation product overflowed") terms.append(term) value = math.fsum(terms) if not math.isfinite(value) or abs(value) > contract.maximum_absolute_output: raise OverflowError("cross-correlation output exceeds the numeric contract") row.append(value) output.append(tuple(row)) return tuple(output) def _translate( matrix: tuple[tuple[float, ...], ...], translation_y: int, translation_x: int ) -> tuple[tuple[float, ...], ...]: height, width = len(matrix), len(matrix[0]) return tuple( tuple( matrix[y - translation_y][x - translation_x] if 0 <= y - translation_y < height and 0 <= x - translation_x < width else 0.0 for x in range(width) ) for y in range(height) ) def audit_translation_equivariance( contract: EquivarianceContract, suite: EquivarianceSuite ) -> EquivarianceReport: if type(contract) is not EquivarianceContract: raise TypeError("contract must be a concrete EquivarianceContract") if type(suite) is not EquivarianceSuite: raise TypeError("suite must be a concrete EquivarianceSuite") contract.__post_init__() suite.__post_init__() if suite.contract_content_id != contract.content_id: raise ValueError("suite belongs to a different contract") expected_content_id = _suite_content_id( contract.content_id, suite.evidence_id, suite.cases ) if suite.evidence_content_id != expected_content_id: raise ValueError("suite content identity does not match its cases") if len(suite.cases) > contract.maximum_cases: raise ValueError("cases exceed contract maximum_cases") case_ids: set[str] = set() input_ids: set[str] = set() per_pass_operations = ( contract.output_shape[0] * contract.output_shape[1] * contract.kernel_height * contract.kernel_width ) if 2 * per_pass_operations * len(suite.cases) > contract.maximum_operations: raise ValueError("equivariance audit exceeds maximum_operations") aligned = ( contract.translation_y % contract.stride_y == 0 and contract.translation_x % contract.stride_x == 0 ) results: list[CaseResult] = [] for case in suite.cases: case.__post_init__() if case.case_id in case_ids: raise ValueError("duplicate case identity") if case.input_id in input_ids: raise ValueError("duplicate input identity") case_ids.add(case.case_id) input_ids.add(case.input_id) if (case.kernel_id, case.kernel_version, case.kernel_content_id) != ( contract.kernel_id, contract.kernel_version, contract.kernel_content_id, ): raise ValueError("kernel identity or content ID does not match the contract") input_values = _matrix( "input_values", case.input_values, height=contract.input_height, width=contract.input_width, maximum_absolute_value=contract.maximum_absolute_input, minimum_numeric_resolution=contract.minimum_numeric_resolution, ) kernel_values = _matrix( "kernel_values", case.kernel_values, height=contract.kernel_height, width=contract.kernel_width, maximum_absolute_value=contract.maximum_absolute_input, minimum_numeric_resolution=contract.minimum_numeric_resolution, ) if kernel_values != contract.kernel_values: raise ValueError("kernel content does not match the contracted kernel") if not aligned: results.append( CaseResult(case.case_id, 0, 0, 0.0, "UNSUPPORTED_ASSUMPTIONS") ) continue original = cross_correlate_2d(contract, input_values, kernel_values) shifted = cross_correlate_2d( contract, _translate(input_values, contract.translation_y, contract.translation_x), kernel_values, ) output_shift_y = contract.translation_y // contract.stride_y output_shift_x = contract.translation_x // contract.stride_x errors: list[float] = [] total_points = contract.output_shape[0] * contract.output_shape[1] for y in range(contract.output_shape[0]): for x in range(contract.output_shape[1]): shifted_y, shifted_x = y + output_shift_y, x + output_shift_x if ( 0 <= shifted_y < contract.output_shape[0] and 0 <= shifted_x < contract.output_shape[1] ): errors.append(abs(original[y][x] - shifted[shifted_y][shifted_x])) if not errors: results.append( CaseResult(case.case_id, 0, total_points, 0.0, "UNSUPPORTED_DOMAIN") ) continue maximum_error = max(errors) status = "EQUIVARIANT" if maximum_error <= contract.absolute_tolerance else "VIOLATION" results.append( CaseResult( case_id=case.case_id, compared_points=len(errors), excluded_boundary_points=total_points - len(errors), maximum_absolute_error=maximum_error, status=status, ) ) statuses = {item.status for item in results} if "VIOLATION" in statuses: status = "VIOLATION" elif statuses != {"EQUIVARIANT"}: status = "UNSUPPORTED" else: status = "EQUIVARIANT" return EquivarianceReport( contract_content_id=contract.content_id, kernel_content_id=contract.kernel_content_id, evidence_id=suite.evidence_id, evidence_content_id=suite.evidence_content_id, output_shape=contract.output_shape, status=status, assumption="WEIGHT_SHARING_VALID_DOMAIN_ONLY", cases=tuple(results), ) def _example() -> None: kernel_values = ((1.0, 0.0), (0.0, -1.0)) contract = EquivarianceContract( contract_version="equivariance-v1", operation="cross-correlation-2d", layout="row-major-hw", padding="valid", stride_y=1, stride_x=1, dilation_y=1, dilation_x=1, translation_y=1, translation_x=1, input_height=5, input_width=5, kernel_height=2, kernel_width=2, kernel_id="edge-kernel", kernel_version="kernel-v1", kernel_values=kernel_values, maximum_cases=4, maximum_operations=10_000, minimum_numeric_resolution=1e-12, maximum_absolute_input=1_000.0, maximum_absolute_output=100_000.0, absolute_tolerance=1e-12, numeric_convention="binary64-fsum-valid-domain-v1", data_owner="team:data", kernel_owner="team:model", audit_owner="team:verification", ) case = ConvolutionCase.capture( contract=contract, case_id="grid-case-001", input_id="grid-input-001", input_version="input-v1", kernel_id=contract.kernel_id, kernel_version=contract.kernel_version, input_values=tuple( tuple(float(y * contract.input_width + x + 1) for x in range(contract.input_width)) for y in range(contract.input_height) ), kernel_values=kernel_values, ) suite = EquivarianceSuite.capture( evidence_id="equivariance-suite-001", contract=contract, cases=(case,) ) report = audit_translation_equivariance(contract, suite) result = report.cases[0] print("example=illustrative_only") print(f"contract_version={contract.contract_version}") print(f"evidence={report.evidence_id}") print(f"evidence_content_id={report.evidence_content_id}") print(f"kernel_content_id={report.kernel_content_id}") print(f"output_shape={report.output_shape[0]}x{report.output_shape[1]}") print(f"compared={result.compared_points}") print(f"excluded_boundary={result.excluded_boundary_points}") print(f"maximum_error={result.maximum_absolute_error:.3f}") print(f"status={report.status}") print(f"assumption={report.assumption}") if __name__ == "__main__": _example()