From f1bfbc530a96593788c8825a195706aa04e8a4b7 Mon Sep 17 00:00:00 2001 From: maxiaosong1124 Date: Wed, 12 Aug 2026 17:35:00 +0800 Subject: [PATCH] feat(ws1): land C3 forward config-invariance harness (#269) Add the shared forward accuracy/invariance API, C2 config matrix, backend provenance fail-closed checks, selected-logprob smoke, GPU gate CLI, CPU tests, and closeout evidence for WS1 C3. --- .github/workflows/ci.yml | 1 + docs/design/ws1-c3-269-closeout-evidence.md | 64 ++ rl_engine/kernels/gtest/__init__.py | 18 + rl_engine/kernels/gtest/forward_invariance.py | 726 ++++++++++++++++++ scripts/check_forward_invariance.py | 262 +++++++ tests/test_forward_invariance.py | 462 +++++++++++ 6 files changed, 1533 insertions(+) create mode 100644 docs/design/ws1-c3-269-closeout-evidence.md create mode 100644 rl_engine/kernels/gtest/forward_invariance.py create mode 100644 scripts/check_forward_invariance.py create mode 100644 tests/test_forward_invariance.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 92cd0433..5d7225a7 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -67,6 +67,7 @@ jobs: run: | python -m pytest rl_engine/tests/test_dispatch.py -v PYTEST_DISABLE_PLUGIN_AUTOLOAD=1 python -m pytest tests/test_attention_correctness.py -q -rs + python -m pytest tests/test_forward_invariance.py -q - name: Run Attention Ground-Truth Tests (CPU-safe) run: | diff --git a/docs/design/ws1-c3-269-closeout-evidence.md b/docs/design/ws1-c3-269-closeout-evidence.md new file mode 100644 index 00000000..a876476c --- /dev/null +++ b/docs/design/ws1-c3-269-closeout-evidence.md @@ -0,0 +1,64 @@ +# WS1 C3 (#269) closeout evidence + +**Parent:** #266 · **Depends on:** #267 / #268 · **Scope:** shared forward harness only + +## Acceptance map + +| #269 criterion | Evidence | +| --- | --- | +| Accuracy and invariance separate | `ForwardInvarianceReport.accuracy_reports` and `invariance_reports` | +| Batch/chunk bitwise after logical unpadding | C1 `forward_invariance` resolver plus exact C2 logical-key validation | +| C2 transforms | `build_config_matrix`: fixed 2×2 matrix, permutation, packing, left/right padding | +| Diagnostics | tensor name, config pair, max/mean absolute error, max relative error | +| Backend provenance | profile, requested/actual backend, candidate/kernel id, device, CC, dtype, seed, fallback reason | +| Silent/cross-profile fallback | missing or mismatched provenance fails; CLI rejects candidate/profile mismatch | +| Selected-logprob smoke | C1 `max_abs_dlogp`, `approx_kl0`, and `clipfrac0` verdict | +| CUDA and Triton same schema | one API/CLI/report schema; both profile contracts are parametrically tested | +| No private thresholds | all tensor and aggregate thresholds resolve through the C1 contract | + +CPU-safe contract regression: + +```bash +python -m pytest -q \ + tests/test_tolerance_contract.py \ + tests/test_ws1_workload.py \ + tests/test_forward_invariance.py \ + tests/test_op_checks.py +``` + +Required-profile runtime examples (must run on CUDA hardware and must not be skipped): + +```bash +python scripts/check_forward_invariance.py \ + --op logp --candidate cuda \ + --backend-profile cuda_bf16 --json + +python scripts/check_forward_invariance.py \ + --op batch_invariant_logp --candidate triton \ + --backend-profile triton_cuda_bf16 --json +``` + +The CLI exits red when CUDA is unavailable, a candidate is absent, the C2 node is +`missing_required`, the compute capability cannot run a declared SM90 candidate, provenance +does not match the profile, or any accuracy/invariance/logprob verdict fails. + +## Runtime verification + +Verified on NVIDIA GeForce RTX 3060 Laptop GPU (`sm86`) with PyTorch 2.8.0+cu128: + +| Gate | Result | +| --- | --- | +| Full pytest suite | `1524 passed, 121 skipped` | +| Full pre-commit | trailing whitespace, EOF, YAML, large-file, black, isort, flake8 passed | +| `cuda_bf16` / generic CUDA logp C3 matrix | passed; all invariance max-abs errors `0.0` | +| `triton_cuda_bf16` / Triton batch-invariant-logp C3 matrix | passed; all invariance max-abs errors `0.0` | +| CUDA operator accuracy check | passed, max absolute error `0.0287590` | +| Triton operator accuracy check | passed, max absolute error `9.536743e-07` | + +The CUDA profile uses the manifest-declared generic CUDA logp candidate on SM86. No SM90 +candidate or fallback path is claimed on this device. + +## Parent boundary + +This closes only C3. It supplies the report and canonicalization contract that C10 must reuse. +It does not claim the full-model, backward, KV-cache, or CI EXIT requirements of #266. diff --git a/rl_engine/kernels/gtest/__init__.py b/rl_engine/kernels/gtest/__init__.py index a12db99e..0fa103c9 100644 --- a/rl_engine/kernels/gtest/__init__.py +++ b/rl_engine/kernels/gtest/__init__.py @@ -1,6 +1,16 @@ # SPDX-License-Identifier: Apache-2.0 # Copyright (c) 2026 RL-Kernel Contributors +from .forward_invariance import ( + AccuracyReport, + ConfigSpec, + ForwardInvarianceReport, + InvarianceReport, + LogprobSmokeResult, + TensorComparisonDetail, + assert_forward_batch_invariant, + build_config_matrix, +) from .op_checks import CandidateSpec, OperatorCase, run_operator_suite from .tolerance import ( BackendProvenance, @@ -15,8 +25,16 @@ ) __all__ = [ + "AccuracyReport", "CandidateSpec", + "ConfigSpec", + "ForwardInvarianceReport", + "InvarianceReport", + "LogprobSmokeResult", "OperatorCase", + "TensorComparisonDetail", + "assert_forward_batch_invariant", + "build_config_matrix", "run_operator_suite", "BackendProvenance", "ContractError", diff --git a/rl_engine/kernels/gtest/forward_invariance.py b/rl_engine/kernels/gtest/forward_invariance.py new file mode 100644 index 00000000..f837334d --- /dev/null +++ b/rl_engine/kernels/gtest/forward_invariance.py @@ -0,0 +1,726 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""WS1 C3 (#269): Forward config-invariance and backend provenance harness. + +Provides a shared forward accuracy/invariance API so downstream gates (C8, C10) +do not invent private thresholds, canonicalize wrong tokens, or compare outputs +from silently-fallback backends. + +All thresholds come from the C1 tolerance contract. Logical identity and config +transforms come from the C2 canonical workload. +""" + +from __future__ import annotations + +from collections.abc import Callable, Mapping, Sequence +from dataclasses import asdict, dataclass +from typing import Any + +import torch + +from rl_engine.kernels.gtest.tolerance import ( + BackendProvenance, + ContractResolveError, + LogprobAggregateVerdict, +) +from rl_engine.kernels.gtest.tolerance import _dtype_name as _normalize_dtype_name +from rl_engine.kernels.gtest.tolerance import ( + compute_logprob_aggregates, + default_clip_interval, + judge_logprob_aggregates, + load_contract, + resolve_comparison_roles, + resolve_tolerance, + validate_backend_provenance, +) +from rl_engine.testing.ws1_workload import ( + LogicalBatch, + PaddedBatch, + PhysicalLayout, + WS1Manifest, + apply_chunking, + apply_packing, + apply_padding, + batch_permutation_from_manifest, + build_logical_batch, + chunk_plan_from_manifest, + load_manifest, + permute_batch, + restore_logical_order, + restore_logical_order_from_padded, +) + + +@dataclass(frozen=True) +class ConfigSpec: + """One workload configuration (batch/chunk/padding/packing variant).""" + + config_id: str + transform_kind: str + logical_batch: LogicalBatch + physical_layout: PhysicalLayout | PaddedBatch + is_canonical: bool = False + + +@dataclass(frozen=True) +class TensorComparisonDetail: + """Per-tensor comparison result with full diagnostics.""" + + tensor_name: str + config_pair: tuple[str, str] + shape: tuple[int, ...] + dtype: str + max_abs_error: float + mean_abs_error: float + max_rel_error: float + atol: float + rtol: float + passed: bool + judgment: str + comparison_lhs_role: str + comparison_rhs_role: str + + def to_dict(self) -> dict[str, Any]: + return asdict(self) + + +@dataclass(frozen=True) +class AccuracyReport: + """Forward accuracy: bf16_candidate vs fp32_reference.""" + + config_id: str + op_class: str + dtype: str + backend_profile: str + details: tuple[TensorComparisonDetail, ...] + passed: bool + backend_provenance: BackendProvenance | None = None + + def to_dict(self) -> dict[str, Any]: + data = asdict(self) + if self.backend_provenance is not None: + data["backend_provenance"] = self.backend_provenance.to_dict() + return data + + +@dataclass(frozen=True) +class InvarianceReport: + """Forward invariance: transformed vs canonical (bitwise atol=0 rtol=0).""" + + canonical_config_id: str + transformed_config_id: str + transform_kind: str + op_class: str + dtype: str + backend_profile: str + details: tuple[TensorComparisonDetail, ...] + passed: bool + + def to_dict(self) -> dict[str, Any]: + return asdict(self) + + +@dataclass(frozen=True) +class LogprobSmokeResult: + """Selected-logprob aggregate smoke on fixed workload.""" + + config_id: str + backend_profile: str + verdict: LogprobAggregateVerdict + passed: bool + + def to_dict(self) -> dict[str, Any]: + return { + "config_id": self.config_id, + "backend_profile": self.backend_profile, + "verdict": self.verdict.to_dict(), + "passed": self.passed, + } + + +@dataclass(frozen=True) +class ForwardInvarianceReport: + """Suite-level report combining accuracy, invariance, and logprob smoke.""" + + op_name: str + backend_profile: str + accuracy_reports: tuple[AccuracyReport, ...] + invariance_reports: tuple[InvarianceReport, ...] + logprob_smoke: LogprobSmokeResult | None + backend_provenance: BackendProvenance | None + candidate_id: str + device: str + compute_capability: str | None + seed: int + fallback_reason: str | None + passed: bool + provenance_valid: bool + metadata_valid: bool + + def to_dict(self) -> dict[str, Any]: + return { + "op_name": self.op_name, + "backend_profile": self.backend_profile, + "accuracy_reports": [r.to_dict() for r in self.accuracy_reports], + "invariance_reports": [r.to_dict() for r in self.invariance_reports], + "logprob_smoke": (self.logprob_smoke.to_dict() if self.logprob_smoke else None), + "backend_provenance": ( + self.backend_provenance.to_dict() if self.backend_provenance else None + ), + "candidate_id": self.candidate_id, + "device": self.device, + "compute_capability": self.compute_capability, + "seed": self.seed, + "fallback_reason": self.fallback_reason, + "passed": self.passed, + "provenance_valid": self.provenance_valid, + "metadata_valid": self.metadata_valid, + } + + +def build_config_matrix( + manifest: WS1Manifest | None = None, +) -> list[ConfigSpec]: + """Build the C2 primary 2x2 matrix + permutation + padding + packing configs.""" + + m = manifest if manifest is not None else load_manifest() + chunk_plan = chunk_plan_from_manifest(m) + batch_bn = build_logical_batch(m) + configs: list[ConfigSpec] = [] + + packed_bn = apply_packing(batch_bn) + configs.append( + ConfigSpec( + config_id="BN/full", + transform_kind="canonical", + logical_batch=batch_bn, + physical_layout=packed_bn, + is_canonical=True, + ) + ) + + chunked_bn = apply_chunking(batch_bn, chunk_size=chunk_plan.chunk_size) + configs.append( + ConfigSpec( + config_id="BN/chunked", + transform_kind="chunk", + logical_batch=batch_bn, + physical_layout=chunked_bn, + ) + ) + + for sample in batch_bn.samples: + single_batch = LogicalBatch( + workload_id=batch_bn.workload_id, + seed=batch_bn.seed, + samples=(sample,), + cell_id="B1-singleton_aggregate/full", + ) + packed_single = apply_packing(single_batch) + configs.append( + ConfigSpec( + config_id=f"B1-singleton_aggregate/full/{sample.sample_id}", + transform_kind="batch_size", + logical_batch=single_batch, + physical_layout=packed_single, + ) + ) + + for sample in batch_bn.samples: + single_batch = LogicalBatch( + workload_id=batch_bn.workload_id, + seed=batch_bn.seed, + samples=(sample,), + cell_id="B1-singleton_aggregate/chunked", + ) + chunked_single = apply_chunking(single_batch, chunk_size=chunk_plan.chunk_size) + configs.append( + ConfigSpec( + config_id=f"B1-singleton_aggregate/chunked/{sample.sample_id}", + transform_kind="chunk", + logical_batch=single_batch, + physical_layout=chunked_single, + ) + ) + + perm = batch_permutation_from_manifest(m) + permuted = permute_batch(batch_bn, perm) + packed_perm = apply_packing(permuted) + configs.append( + ConfigSpec( + config_id="BN/permuted", + transform_kind="permutation", + logical_batch=permuted, + physical_layout=packed_perm, + ) + ) + + padded_right = apply_padding(batch_bn, pad_side="right", manifest=m) + configs.append( + ConfigSpec( + config_id="BN/padded_right", + transform_kind="padding", + logical_batch=batch_bn, + physical_layout=padded_right, + is_canonical=False, + ) + ) + + padded_left = apply_padding(batch_bn, pad_side="left", manifest=m) + configs.append( + ConfigSpec( + config_id="BN/padded_left", + transform_kind="padding", + logical_batch=batch_bn, + physical_layout=padded_left, + is_canonical=False, + ) + ) + + return configs + + +def _compare_logical_tensors( + canonical: torch.Tensor, + transformed: torch.Tensor, + *, + judgment: str, + contract: Mapping[str, Any], + op_class: str, + dtype: str | torch.dtype, + backend_profile: str | None = None, + tensor_name: str = "output", + config_pair: tuple[str, str] = ("canonical", "transformed"), +) -> TensorComparisonDetail: + """Compare two tensors aligned to the same logical token order.""" + + spec = resolve_tolerance( + contract, + judgment=judgment, + op_class=op_class, + dtype=dtype, + backend_profile=backend_profile, + ) + atol, rtol = spec.atol, spec.rtol + roles = resolve_comparison_roles(contract, judgment) + + canonical_fp32 = canonical.float() + transformed_fp32 = transformed.float() + + if canonical_fp32.shape != transformed_fp32.shape: + return TensorComparisonDetail( + tensor_name=tensor_name, + config_pair=config_pair, + shape=tuple(transformed_fp32.shape), + dtype=_normalize_dtype_name(transformed.dtype), + max_abs_error=float("inf"), + mean_abs_error=float("inf"), + max_rel_error=float("inf"), + atol=atol, + rtol=rtol, + passed=False, + judgment=judgment, + comparison_lhs_role=roles.comparison_lhs_role, + comparison_rhs_role=roles.comparison_rhs_role, + ) + + abs_error = (canonical_fp32 - transformed_fp32).abs() + if abs_error.numel() == 0: + max_abs = 0.0 + mean_abs = 0.0 + max_rel = 0.0 + else: + max_abs = float(abs_error.max().item()) + mean_abs = float(abs_error.mean().item()) + rel_error = abs_error / canonical_fp32.abs().clamp_min(1e-12) + max_rel = float(rel_error.max().item()) + + passed = bool(torch.allclose(transformed_fp32, canonical_fp32, atol=atol, rtol=rtol)) + + return TensorComparisonDetail( + tensor_name=tensor_name, + config_pair=config_pair, + shape=tuple(canonical_fp32.shape), + dtype=_normalize_dtype_name(canonical.dtype), + max_abs_error=max_abs, + mean_abs_error=mean_abs, + max_rel_error=max_rel, + atol=atol, + rtol=rtol, + passed=passed, + judgment=judgment, + comparison_lhs_role=roles.comparison_lhs_role, + comparison_rhs_role=roles.comparison_rhs_role, + ) + + +def _validate_provenance( + contract: Mapping[str, Any], + provenance: BackendProvenance | None, + backend_profile: str, +) -> bool: + """Validate backend provenance; return False if silent/cross-profile fallback.""" + + if provenance is None: + return False + try: + validate_backend_provenance(contract, provenance) + except ContractResolveError: + return False + if provenance.backend_profile != backend_profile: + return False + return True + + +def _collect_logical_outputs( + op: Callable[..., Any] | Any, + config: ConfigSpec, + *, + op_kwargs: Mapping[str, Any] | None = None, +) -> dict[tuple[str, int], torch.Tensor]: + """Run op on a config and restore outputs to logical (sample_id, position) order.""" + + kwargs = dict(op_kwargs) if op_kwargs else {} + if hasattr(op, "forward") and callable(op.forward): + raw_output = op.forward(config=config, **kwargs) + else: + raw_output = op(config=config, **kwargs) + + if isinstance(raw_output, dict): + return raw_output + + if isinstance(raw_output, torch.Tensor): + if isinstance(config.physical_layout, PaddedBatch): + if raw_output.shape != ( + len(config.physical_layout.restore_map), + config.physical_layout.padded_len, + ): + raise ValueError( + f"padded output shape {tuple(raw_output.shape)} does not match " + f"({len(config.physical_layout.restore_map)}, " + f"{config.physical_layout.padded_len})" + ) + return restore_logical_order_from_padded(config.physical_layout, list(raw_output)) + flat = raw_output.reshape(-1) + return restore_logical_order(config.physical_layout, list(flat)) + + raise TypeError(f"op must return dict or Tensor, got {type(raw_output)!r}") + + +def _align_and_compare_invariance( + canonical_map: dict[tuple[str, int], Any], + transformed_map: dict[tuple[str, int], Any], + *, + contract: Mapping[str, Any], + op_class: str, + dtype: str | torch.dtype, + backend_profile: str | None, + canonical_id: str, + transformed_id: str, + tensor_name: str = "output", + expected_keys: set[tuple[str, int]], +) -> TensorComparisonDetail: + """Align two logical output maps and compare for bitwise invariance.""" + + canonical_keys = set(canonical_map) + transformed_keys = set(transformed_map) + if ( + not expected_keys + or not expected_keys.issubset(canonical_keys) + or not expected_keys.issubset(transformed_keys) + ): + return TensorComparisonDetail( + tensor_name=tensor_name, + config_pair=(canonical_id, transformed_id), + shape=(0,), + dtype=( + _normalize_dtype_name(dtype) + if isinstance(dtype, str) + else _normalize_dtype_name(dtype) + ), + max_abs_error=float("inf"), + mean_abs_error=float("inf"), + max_rel_error=float("inf"), + atol=0.0, + rtol=0.0, + passed=False, + judgment="forward_invariance", + comparison_lhs_role="transformed_config", + comparison_rhs_role="canonical_config", + ) + + shared_keys = sorted(expected_keys) + canonical_vals = torch.stack([torch.as_tensor(canonical_map[k]) for k in shared_keys]) + transformed_vals = torch.stack([torch.as_tensor(transformed_map[k]) for k in shared_keys]) + + return _compare_logical_tensors( + canonical_vals, + transformed_vals, + judgment="forward_invariance", + contract=contract, + op_class=op_class, + dtype=dtype, + backend_profile=backend_profile, + tensor_name=tensor_name, + config_pair=(canonical_id, transformed_id), + ) + + +def assert_forward_batch_invariant( + op: Callable[..., Any] | Any, + configs: Sequence[ConfigSpec] | None = None, + contract: Mapping[str, Any] | None = None, + *, + manifest: WS1Manifest | None = None, + backend_profile: str, + provenance: BackendProvenance | None = None, + gold_fn: Callable[..., Any] | None = None, + op_class: str = "logprob", + dtype: torch.dtype = torch.bfloat16, + op_name: str = "operator", + op_kwargs: Mapping[str, Any] | None = None, + include_logprob_smoke: bool = True, + active_only: bool = True, + candidate_id: str = "unspecified", + device: str = "unspecified", + compute_capability: str | None = None, + fallback_reason: str | None = None, +) -> ForwardInvarianceReport: + """Run forward config-invariance and accuracy checks. + + This is the sole C3 API. C10 must reuse this harness/report schema. + + Args: + op: Operator callable. Must accept (config=ConfigSpec, **op_kwargs) and + return either a dict[(sample_id, position) -> Tensor] or a flat Tensor. + configs: Config matrix; built from C2 manifest if None. + contract: C1 tolerance contract; loaded from default path if None. + manifest: C2 workload manifest; loaded from default path if None. + backend_profile: Required profile id (cuda_bf16 or triton_cuda_bf16). + provenance: Runtime-observed backend provenance. Missing provenance fails closed. + gold_fn: FP32 reference callable for accuracy checks. + op_class: Operator class for tolerance resolution. + dtype: Execution dtype. + op_name: Name for reporting. + op_kwargs: Extra kwargs passed to op. + include_logprob_smoke: Whether to run logprob aggregate smoke. + active_only: Only compare active (non-prompt) tokens for invariance. + + Returns: + ForwardInvarianceReport with accuracy, invariance, and logprob sub-reports. + """ + + loaded_contract = dict(contract or load_contract()) + m = manifest if manifest is not None else load_manifest() + config_list = list(configs) if configs is not None else build_config_matrix(m) + if not config_list: + raise ValueError("configs must contain at least one configuration") + if gold_fn is None: + raise ValueError("gold_fn is required for forward accuracy") + if include_logprob_smoke and op_class != "logprob": + raise ValueError("selected-logprob smoke requires op_class='logprob'") + + provenance_valid = _validate_provenance(loaded_contract, provenance, backend_profile) + if not provenance_valid and fallback_reason is None: + fallback_reason = "missing or contract-invalid backend provenance" + metadata_valid = ( + candidate_id != "unspecified" + and device != "unspecified" + and compute_capability is not None + and fallback_reason is None + ) + + canonical_config = next((c for c in config_list if c.is_canonical), config_list[0]) + canonical_outputs = _collect_logical_outputs(op, canonical_config, op_kwargs=op_kwargs) + + def expected_keys(config: ConfigSpec) -> set[tuple[str, int]]: + return set(config.logical_batch.logical_keys(active_only=active_only)) + + def validate_keys( + outputs: Mapping[tuple[str, int], Any], config: ConfigSpec, label: str + ) -> None: + required = expected_keys(config) + allowed = set(config.logical_batch.logical_keys(active_only=False)) + actual = set(outputs) + if not required.issubset(actual) or not actual.issubset(allowed): + raise ValueError( + f"{label} output keys for {config.config_id!r} do not match the " + "C2 logical identity" + ) + + canonical_keys = expected_keys(canonical_config) + validate_keys(canonical_outputs, canonical_config, "canonical") + + invariance_reports: list[InvarianceReport] = [] + for config in config_list: + if config.is_canonical: + continue + transformed_outputs = _collect_logical_outputs(op, config, op_kwargs=op_kwargs) + validate_keys(transformed_outputs, config, "transformed") + detail = _align_and_compare_invariance( + canonical_outputs, + transformed_outputs, + contract=loaded_contract, + op_class=op_class, + dtype=dtype, + backend_profile=backend_profile, + canonical_id=canonical_config.config_id, + transformed_id=config.config_id, + expected_keys=expected_keys(config), + ) + invariance_reports.append( + InvarianceReport( + canonical_config_id=canonical_config.config_id, + transformed_config_id=config.config_id, + transform_kind=config.transform_kind, + op_class=op_class, + dtype=_normalize_dtype_name(dtype), + backend_profile=backend_profile, + details=(detail,), + passed=detail.passed, + ) + ) + + accuracy_reports: list[AccuracyReport] = [] + for config in config_list: + candidate_outputs = ( + canonical_outputs + if config.is_canonical + else _collect_logical_outputs(op, config, op_kwargs=op_kwargs) + ) + gold_outputs = _collect_logical_outputs(gold_fn, config, op_kwargs=op_kwargs) + keys = expected_keys(config) + validate_keys(candidate_outputs, config, "candidate accuracy") + validate_keys(gold_outputs, config, "reference accuracy") + ordered_keys = sorted(keys) + candidate_vals = torch.stack([torch.as_tensor(candidate_outputs[k]) for k in ordered_keys]) + gold_vals = torch.stack([torch.as_tensor(gold_outputs[k]) for k in ordered_keys]) + acc_detail = _compare_logical_tensors( + gold_vals, + candidate_vals, + judgment="forward_accuracy", + contract=loaded_contract, + op_class=op_class, + dtype=dtype, + backend_profile=backend_profile, + tensor_name="selected_logprob" if op_class == "logprob" else "output", + config_pair=(config.config_id, "fp32_reference"), + ) + accuracy_reports.append( + AccuracyReport( + config_id=config.config_id, + op_class=op_class, + dtype=_normalize_dtype_name(dtype), + backend_profile=backend_profile, + details=(acc_detail,), + passed=acc_detail.passed, + backend_provenance=provenance, + ) + ) + + logprob_smoke: LogprobSmokeResult | None = None + if include_logprob_smoke: + logprob_smoke = _run_logprob_smoke( + canonical_outputs, + gold_fn, + canonical_config, + loaded_contract, + m, + backend_profile=backend_profile, + op_kwargs=op_kwargs, + active_keys=canonical_keys, + ) + + all_invariance_passed = all(r.passed for r in invariance_reports) + all_accuracy_passed = all(r.passed for r in accuracy_reports) + smoke_passed = logprob_smoke.passed if logprob_smoke is not None else True + + overall_passed = ( + all_invariance_passed + and all_accuracy_passed + and smoke_passed + and provenance_valid + and metadata_valid + ) + + return ForwardInvarianceReport( + op_name=op_name, + backend_profile=backend_profile, + accuracy_reports=tuple(accuracy_reports), + invariance_reports=tuple(invariance_reports), + logprob_smoke=logprob_smoke, + backend_provenance=provenance, + candidate_id=candidate_id, + device=device, + compute_capability=compute_capability, + seed=m.seed, + fallback_reason=fallback_reason, + passed=overall_passed, + provenance_valid=provenance_valid, + metadata_valid=metadata_valid, + ) + + +def _run_logprob_smoke( + candidate_outputs: dict[tuple[str, int], Any], + gold_fn: Callable[..., Any] | Any, + config: ConfigSpec, + contract: Mapping[str, Any], + manifest: WS1Manifest, + *, + backend_profile: str, + op_kwargs: Mapping[str, Any] | None = None, + active_keys: set[tuple[str, int]] | None = None, +) -> LogprobSmokeResult: + """Run selected-logprob aggregate smoke check.""" + + gold_outputs = _collect_logical_outputs(gold_fn, config, op_kwargs=op_kwargs) + if active_keys is not None: + shared = sorted(k for k in candidate_outputs if k in gold_outputs and k in active_keys) + else: + shared = sorted(k for k in candidate_outputs if k in gold_outputs) + + if not shared: + raise ContractResolveError("no shared active tokens for logprob smoke") + + lhs_logp = torch.stack([torch.as_tensor(candidate_outputs[k]).float() for k in shared]) + rhs_logp = torch.stack([torch.as_tensor(gold_outputs[k]).float() for k in shared]) + active_mask = torch.ones(len(shared), dtype=torch.bool) + + clip_interval = default_clip_interval(contract) + roles = resolve_comparison_roles(contract, "forward_accuracy") + + aggregates = compute_logprob_aggregates( + lhs_logp, + rhs_logp, + active_mask, + contract=contract, + report_kind="forward_accuracy", + clip_interval=clip_interval, + comparison_lhs_role=roles.comparison_lhs_role, + comparison_rhs_role=roles.comparison_rhs_role, + ) + verdict = judge_logprob_aggregates( + aggregates, + contract, + execution_dtype="bfloat16", + ) + return LogprobSmokeResult( + config_id=config.config_id, + backend_profile=backend_profile, + verdict=verdict, + passed=verdict.passed, + ) + + +__all__ = [ + "AccuracyReport", + "ConfigSpec", + "ForwardInvarianceReport", + "InvarianceReport", + "LogprobSmokeResult", + "TensorComparisonDetail", + "assert_forward_batch_invariant", + "build_config_matrix", +] diff --git a/scripts/check_forward_invariance.py b/scripts/check_forward_invariance.py new file mode 100644 index 00000000..d4f8dd9e --- /dev/null +++ b/scripts/check_forward_invariance.py @@ -0,0 +1,262 @@ +#!/usr/bin/env python +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Run the WS1 C3 selected-logprob forward invariance gate on a real GPU.""" + +from __future__ import annotations + +import argparse +import json +import pathlib +import sys +from typing import Any + +import torch + +REPO_ROOT = pathlib.Path(__file__).resolve().parents[1] +if str(REPO_ROOT) not in sys.path: + sys.path.insert(0, str(REPO_ROOT)) + +from rl_engine.kernels.gtest import ( # noqa: E402 + BackendProvenance, + ConfigSpec, + assert_forward_batch_invariant, + load_contract, +) +from rl_engine.kernels.gtest.operator_specs import OP_SPECS, _load_object # noqa: E402 +from rl_engine.kernels.gtest.tolerance import resolve_dtype_policy # noqa: E402 +from rl_engine.testing.ws1_workload import PaddedBatch, load_manifest # noqa: E402 + + +def _object_path(value: Any) -> str: + cls = value.__class__ + return f"{cls.__module__}.{cls.__qualname__}" + + +def _profile_node(manifest: Any, profile: str, op_name: str) -> dict[str, Any]: + node_name = "logprob" if op_name == "logp" else op_name + nodes = manifest.backend_profiles[profile]["required_nodes"] + node = next((dict(item) for item in nodes if item["node"] == node_name), None) + if node is None: + raise RuntimeError(f"profile {profile!r} does not declare node {node_name!r}") + if node["status"] != "declared": + raise RuntimeError( + f"profile {profile!r} node {node_name!r} is {node['status']!r}; " + "missing required candidates are red, not fallback or N/A" + ) + return node + + +def _candidate_family(candidate: str) -> str: + if candidate.startswith("cuda"): + return "cuda" + if candidate == "triton": + return "triton" + return candidate + + +def _validate_candidate_selection( + *, manifest: Any, profile: str, op_name: str, candidate: str +) -> dict[str, Any]: + node = _profile_node(manifest, profile, op_name) + expected_family = manifest.backend_profiles[profile]["backend_family"] + actual_family = _candidate_family(candidate) + if actual_family != expected_family: + raise RuntimeError( + f"candidate {candidate!r} belongs to {actual_family!r}, but profile " + f"{profile!r} requires {expected_family!r}" + ) + if candidate != node["expected_backend_id"]: + raise RuntimeError( + f"candidate {candidate!r} does not match the C2 declaration " + f"{node['expected_backend_id']!r} for {profile}/{node['node']}" + ) + return node + + +def _physical_rows( + config: ConfigSpec, +) -> tuple[list[tuple[str, int] | None], tuple[int, ...]]: + layout = config.physical_layout + if isinstance(layout, PaddedBatch): + keys = [key for row in layout.restore_map for key in row] + return keys, (len(layout.restore_map), layout.padded_len) + return list(layout.restore_map), (len(layout.restore_map),) + + +def _make_inputs( + config: ConfigSpec, + *, + device: torch.device, + dtype: torch.dtype, + vocab_size: int, +) -> tuple[torch.Tensor, torch.Tensor]: + """Create row-local deterministic logits from C2 logical identity.""" + + keys, leading_shape = _physical_rows(config) + vocab_axis = torch.arange(vocab_size, device=device, dtype=torch.int64) + rows: list[torch.Tensor] = [] + targets: list[int] = [] + token_by_key = { + (token.sample_id, token.token_position): token.token_id + for sample in config.logical_batch.samples + for token in sample.tokens() + } + for key in keys: + if key is None: + position, token_id = 0, 0 + else: + position = key[1] + token_id = token_by_key[key] + # Integer construction makes each logical row independent of batching, + # chunking, permutation, padding, and RNG consumption order. + values = ((vocab_axis + token_id * 17 + position * 13) % 257) - 128 + rows.append((values.to(torch.float32) / 1024.0).to(dtype)) + targets.append(token_id % vocab_size) + logits = torch.stack(rows).reshape(leading_shape + (vocab_size,)) + target_tensor = torch.tensor(targets, device=device, dtype=torch.long).reshape(leading_shape) + return logits, target_tensor + + +def _make_runner( + operator: Any, + *, + device: torch.device, + dtype: torch.dtype, + vocab_size: int, + reference: bool, +): + def run(config: ConfigSpec, **_: Any) -> torch.Tensor: + logits, targets = _make_inputs(config, device=device, dtype=dtype, vocab_size=vocab_size) + if reference: + logits = logits.float() + return operator(logits, targets) + + return run + + +def _summarize(report: Any) -> None: + print( + f"op={report.op_name} profile={report.backend_profile} " + f"candidate={report.candidate_id} passed={report.passed}" + ) + print( + f" device={report.device} cc={report.compute_capability} seed={report.seed} " + f"provenance_valid={report.provenance_valid}" + ) + for acc in report.accuracy_reports: + detail = acc.details[0] + print( + f" accuracy config={acc.config_id} max_abs={detail.max_abs_error:.8e} " + f"max_rel={detail.max_rel_error:.8e} passed={acc.passed}" + ) + for inv in report.invariance_reports: + detail = inv.details[0] + print( + f" invariance pair={detail.config_pair} transform={inv.transform_kind} " + f"max_abs={detail.max_abs_error:.8e} passed={inv.passed}" + ) + if report.logprob_smoke is not None: + print(f" selected_logprob_smoke passed={report.logprob_smoke.passed}") + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="WS1 C3 forward invariance GPU gate") + parser.add_argument("--op", choices=("logp", "batch_invariant_logp"), default="logp") + parser.add_argument( + "--candidate", required=True, help="Manifest-declared CUDA/Triton candidate" + ) + parser.add_argument( + "--backend-profile", + choices=("cuda_bf16", "triton_cuda_bf16"), + required=True, + ) + parser.add_argument("--device", default="cuda") + parser.add_argument("--vocab", type=int, default=151936) + parser.add_argument("--json", action="store_true") + return parser.parse_args() + + +def main() -> None: + args = parse_args() + device = torch.device(args.device) + if device.type != "cuda" or not torch.cuda.is_available(): + raise SystemExit("ERROR: C3 required-profile evidence requires an available CUDA device") + if args.vocab <= 240: + raise SystemExit("ERROR: --vocab must cover every fixed C2 workload token id") + + contract = load_contract() + manifest = load_manifest() + node = _validate_candidate_selection( + manifest=manifest, + profile=args.backend_profile, + op_name=args.op, + candidate=args.candidate, + ) + spec = OP_SPECS[args.op] + if args.candidate not in spec.candidate_paths: + raise SystemExit(f"ERROR: operator {args.op!r} has no candidate {args.candidate!r}") + + candidate_op = _load_object(spec.candidate_paths[args.candidate])() + gold_op = _load_object(spec.gold_path)() + gold_method = getattr(gold_op, spec.gold_method) + policy = resolve_dtype_policy(contract) + family = _candidate_family(args.candidate) + torch.backends.cuda.matmul.allow_tf32 = False + torch.backends.cudnn.allow_tf32 = False + cc_tuple = torch.cuda.get_device_capability(device) + cc = f"sm{cc_tuple[0]}{cc_tuple[1]}" + if args.candidate == "cuda-sm90" and cc_tuple[0] != 9: + raise SystemExit( + "ERROR: cuda-sm90 candidate requested on non-SM90 hardware; fallback forbidden" + ) + + provenance = BackendProvenance( + backend_profile=args.backend_profile, + requested_backend=manifest.backend_profiles[args.backend_profile]["backend_family"], + actual_backend=family, + execution_dtype=policy.execution_dtype, + accumulation_dtype=policy.accumulation_dtype, + output_dtype=policy.output_dtype_default, + reference_dtype=policy.reference_dtype, + candidate_tf32_enabled=torch.backends.cuda.matmul.allow_tf32, + reference_tf32_enabled=torch.backends.cuda.matmul.allow_tf32, + ) + report = assert_forward_batch_invariant( + _make_runner( + candidate_op, + device=device, + dtype=torch.bfloat16, + vocab_size=args.vocab, + reference=False, + ), + contract=contract, + manifest=manifest, + backend_profile=args.backend_profile, + provenance=provenance, + gold_fn=_make_runner( + gold_method, + device=device, + dtype=torch.bfloat16, + vocab_size=args.vocab, + reference=True, + ), + op_class="logprob", + dtype=torch.bfloat16, + op_name=args.op, + candidate_id=f"{_object_path(candidate_op)}::{node['expected_kernel_config_id']}", + device=f"{device}:{torch.cuda.get_device_name(device)}", + compute_capability=cc, + ) + + if args.json: + print(json.dumps(report.to_dict(), indent=2, default=str)) + else: + _summarize(report) + if not report.passed: + raise SystemExit(1) + + +if __name__ == "__main__": + main() diff --git a/tests/test_forward_invariance.py b/tests/test_forward_invariance.py new file mode 100644 index 00000000..d3f89ed8 --- /dev/null +++ b/tests/test_forward_invariance.py @@ -0,0 +1,462 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Unit tests for WS1 C3 forward config-invariance harness.""" + +from __future__ import annotations + +from typing import Any + +import pytest +import torch + +from rl_engine.kernels.gtest.forward_invariance import ( + ConfigSpec, + ForwardInvarianceReport, + TensorComparisonDetail, + _validate_provenance, +) +from rl_engine.kernels.gtest.forward_invariance import ( + assert_forward_batch_invariant as _assert_forward_batch_invariant, +) +from rl_engine.kernels.gtest.forward_invariance import build_config_matrix +from rl_engine.kernels.gtest.tolerance import BackendProvenance, load_contract, resolve_tolerance +from rl_engine.testing.ws1_workload import LogicalBatch, LogicalSample, PaddedBatch, load_manifest + + +def assert_forward_batch_invariant(*args: Any, **kwargs: Any) -> ForwardInvarianceReport: + """Supply explicit synthetic runtime metadata for CPU-safe harness tests.""" + + kwargs.setdefault("candidate_id", "synthetic-test-candidate") + kwargs.setdefault("device", "cpu:test-double") + kwargs.setdefault("compute_capability", "synthetic") + return _assert_forward_batch_invariant(*args, **kwargs) + + +@pytest.fixture() +def contract() -> dict[str, Any]: + return load_contract() + + +@pytest.fixture() +def manifest(): + return load_manifest() + + +@pytest.fixture() +def simple_batch() -> LogicalBatch: + samples = ( + LogicalSample(sample_id="s0", token_ids=(1, 2, 3, 4), prompt_len=2, seq_len=4), + LogicalSample(sample_id="s1", token_ids=(5, 6, 7, 8), prompt_len=1, seq_len=4), + ) + return LogicalBatch(workload_id="test", seed=42, samples=samples) + + +def _make_identity_op(value: float = 1.0): + """Op that returns identical outputs regardless of config (batch-invariant).""" + + def op(config: ConfigSpec, **kwargs: Any) -> dict[tuple[str, int], torch.Tensor]: + result: dict[tuple[str, int], torch.Tensor] = {} + for sample in config.logical_batch.samples: + for tok in sample.active_tokens(): + result[(tok.sample_id, tok.token_position)] = torch.tensor( + value, dtype=torch.bfloat16 + ) + return result + + return op + + +def _make_drifting_op(drift: float = 0.1): + """Op that adds drift per sample to break invariance.""" + + def op(config: ConfigSpec, **kwargs: Any) -> dict[tuple[str, int], torch.Tensor]: + result: dict[tuple[str, int], torch.Tensor] = {} + for idx, sample in enumerate(config.logical_batch.samples): + for tok in sample.active_tokens(): + result[(tok.sample_id, tok.token_position)] = torch.tensor( + 1.0 + idx * drift, dtype=torch.bfloat16 + ) + return result + + return op + + +def _make_provenance( + backend_profile: str = "cuda_bf16", + requested: str = "cuda", + actual: str = "cuda", +) -> BackendProvenance: + return BackendProvenance( + backend_profile=backend_profile, + requested_backend=requested, + actual_backend=actual, + execution_dtype="bfloat16", + accumulation_dtype="float32", + output_dtype="bfloat16", + reference_dtype="float32", + candidate_tf32_enabled=False, + reference_tf32_enabled=False, + ) + + +class TestReportStructure: + def test_accuracy_and_invariance_reported_separately(self, contract, manifest): + op = _make_identity_op() + report = assert_forward_batch_invariant( + op, + contract=contract, + manifest=manifest, + backend_profile="cuda_bf16", + provenance=_make_provenance(), + gold_fn=_make_identity_op(1.0), + op_class="logprob", + dtype=torch.bfloat16, + op_name="test_op", + include_logprob_smoke=False, + ) + assert isinstance(report, ForwardInvarianceReport) + assert hasattr(report, "accuracy_reports") + assert hasattr(report, "invariance_reports") + assert isinstance(report.accuracy_reports, tuple) + assert isinstance(report.invariance_reports, tuple) + assert len(report.invariance_reports) > 0 + assert len(report.accuracy_reports) == len(build_config_matrix(manifest)) + + def test_report_contains_required_runtime_metadata(self, contract, manifest): + report = assert_forward_batch_invariant( + _make_identity_op(), + contract=contract, + manifest=manifest, + backend_profile="cuda_bf16", + provenance=_make_provenance(), + gold_fn=_make_identity_op(), + op_class="logprob", + include_logprob_smoke=False, + candidate_id="cuda-test-kernel", + device="cuda:0:test-device", + compute_capability="sm90", + ) + payload = report.to_dict() + assert payload["candidate_id"] == "cuda-test-kernel" + assert payload["device"] == "cuda:0:test-device" + assert payload["compute_capability"] == "sm90" + assert payload["seed"] == manifest.seed + assert payload["fallback_reason"] is None + + def test_missing_runtime_metadata_fails_closed(self, contract, manifest): + report = _assert_forward_batch_invariant( + _make_identity_op(), + contract=contract, + manifest=manifest, + backend_profile="cuda_bf16", + provenance=_make_provenance(), + gold_fn=_make_identity_op(), + include_logprob_smoke=False, + ) + assert report.provenance_valid + assert not report.metadata_valid + assert not report.passed + + def test_report_contains_max_abs_rel_tensor_name(self, contract, manifest): + op = _make_identity_op() + report = assert_forward_batch_invariant( + op, + contract=contract, + manifest=manifest, + backend_profile="cuda_bf16", + provenance=_make_provenance(), + gold_fn=_make_identity_op(), + op_class="logprob", + dtype=torch.bfloat16, + op_name="test_op", + include_logprob_smoke=False, + ) + for inv in report.invariance_reports: + for detail in inv.details: + assert isinstance(detail, TensorComparisonDetail) + assert detail.tensor_name is not None + assert detail.max_abs_error is not None + assert detail.max_rel_error is not None + assert detail.config_pair is not None + assert len(detail.config_pair) == 2 + + +class TestInvariance: + def test_invariance_bitwise_zero_tolerance(self, contract, manifest): + op = _make_identity_op() + report = assert_forward_batch_invariant( + op, + contract=contract, + manifest=manifest, + backend_profile="cuda_bf16", + provenance=_make_provenance(), + gold_fn=_make_identity_op(), + op_class="logprob", + dtype=torch.bfloat16, + op_name="test_op", + include_logprob_smoke=False, + ) + for inv in report.invariance_reports: + for detail in inv.details: + assert detail.judgment == "forward_invariance" + assert detail.atol == 0.0 + assert detail.rtol == 0.0 + + def test_identity_op_passes_invariance(self, contract, manifest): + op = _make_identity_op() + report = assert_forward_batch_invariant( + op, + contract=contract, + manifest=manifest, + backend_profile="cuda_bf16", + provenance=_make_provenance(), + gold_fn=_make_identity_op(), + op_class="logprob", + dtype=torch.bfloat16, + op_name="test_op", + include_logprob_smoke=False, + ) + for inv in report.invariance_reports: + assert inv.passed, f"invariance failed for {inv.transformed_config_id}" + assert report.passed + + def test_logical_unpadding_before_compare(self, contract, manifest): + op = _make_identity_op() + report = assert_forward_batch_invariant( + op, + contract=contract, + manifest=manifest, + backend_profile="cuda_bf16", + provenance=_make_provenance(), + gold_fn=_make_identity_op(), + op_class="logprob", + dtype=torch.bfloat16, + op_name="test_op", + include_logprob_smoke=False, + active_only=True, + ) + for inv in report.invariance_reports: + assert inv.passed + + def test_padding_configs_use_c2_padded_layout(self, manifest): + padded = [c for c in build_config_matrix(manifest) if c.transform_kind == "padding"] + assert {c.physical_layout.pad_side for c in padded} == {"left", "right"} + assert all(isinstance(c.physical_layout, PaddedBatch) for c in padded) + + def test_missing_active_token_hard_fails(self, contract, manifest): + def incomplete(config: ConfigSpec, **kwargs: Any): + result = _make_identity_op()(config, **kwargs) + result.pop(next(iter(result))) + return result + + with pytest.raises(ValueError, match="C2 logical identity"): + assert_forward_batch_invariant( + incomplete, + contract=contract, + manifest=manifest, + backend_profile="cuda_bf16", + provenance=_make_provenance(), + gold_fn=_make_identity_op(), + include_logprob_smoke=False, + ) + + def test_padded_tensor_is_logically_unpadded(self, contract, manifest): + def physical_identity(config: ConfigSpec, **kwargs: Any): + layout = config.physical_layout + if isinstance(layout, PaddedBatch): + return torch.ones( + (len(layout.restore_map), layout.padded_len), dtype=torch.bfloat16 + ) + return torch.ones(len(layout.restore_map), dtype=torch.bfloat16) + + report = assert_forward_batch_invariant( + physical_identity, + contract=contract, + manifest=manifest, + backend_profile="cuda_bf16", + provenance=_make_provenance(), + gold_fn=physical_identity, + include_logprob_smoke=False, + ) + padding_reports = [r for r in report.invariance_reports if r.transform_kind == "padding"] + assert len(padding_reports) == 2 + assert all(r.passed for r in padding_reports) + + +class TestAccuracy: + def test_missing_reference_is_rejected(self, contract, manifest): + with pytest.raises(ValueError, match="gold_fn is required"): + assert_forward_batch_invariant( + _make_identity_op(), + contract=contract, + manifest=manifest, + backend_profile="cuda_bf16", + provenance=_make_provenance(), + gold_fn=None, + include_logprob_smoke=False, + ) + + def test_accuracy_uses_c1_tolerances(self, contract, manifest): + op = _make_identity_op(1.0) + gold = _make_identity_op(1.0) + report = assert_forward_batch_invariant( + op, + contract=contract, + manifest=manifest, + backend_profile="cuda_bf16", + provenance=_make_provenance(), + gold_fn=gold, + op_class="logprob", + dtype=torch.bfloat16, + op_name="test_op", + include_logprob_smoke=False, + ) + for acc in report.accuracy_reports: + for detail in acc.details: + assert detail.judgment == "forward_accuracy" + spec = resolve_tolerance( + contract, + judgment="forward_accuracy", + op_class="logprob", + dtype=torch.bfloat16, + backend_profile="cuda_bf16", + ) + assert detail.atol == spec.atol + assert detail.rtol == spec.rtol + + def test_no_private_thresholds(self, contract, manifest): + op = _make_identity_op(1.0) + gold = _make_identity_op(1.0) + report = assert_forward_batch_invariant( + op, + contract=contract, + manifest=manifest, + backend_profile="cuda_bf16", + provenance=_make_provenance(), + gold_fn=gold, + op_class="logprob", + dtype=torch.bfloat16, + op_name="test_op", + include_logprob_smoke=False, + ) + for acc in report.accuracy_reports: + for detail in acc.details: + spec = resolve_tolerance( + contract, + judgment=detail.judgment, + op_class=acc.op_class, + dtype=torch.bfloat16, + backend_profile=acc.backend_profile, + ) + assert detail.atol == spec.atol + assert detail.rtol == spec.rtol + + +class TestBackendProvenance: + def test_valid_provenance_passes(self, contract): + provenance = _make_provenance("cuda_bf16", "cuda", "cuda") + assert _validate_provenance(contract, provenance, "cuda_bf16") is True + + def test_silent_fallback_rejected(self, contract): + provenance = _make_provenance("cuda_bf16", "cuda", "triton") + assert _validate_provenance(contract, provenance, "cuda_bf16") is False + + def test_cross_profile_fallback_rejected(self, contract): + provenance = _make_provenance("triton_cuda_bf16", "triton", "triton") + assert _validate_provenance(contract, provenance, "cuda_bf16") is False + + def test_none_provenance_fails_closed(self, contract): + assert _validate_provenance(contract, None, "cuda_bf16") is False + + @pytest.mark.parametrize( + ("profile", "family"), + [("cuda_bf16", "cuda"), ("triton_cuda_bf16", "triton")], + ) + def test_required_profiles_share_report_schema(self, contract, manifest, profile, family): + provenance = _make_provenance(profile, family, family) + report = assert_forward_batch_invariant( + _make_identity_op(), + contract=contract, + manifest=manifest, + backend_profile=profile, + provenance=provenance, + gold_fn=_make_identity_op(), + include_logprob_smoke=False, + ) + assert report.passed + assert set(report.to_dict()) == set( + ForwardInvarianceReport( + op_name="x", + backend_profile=profile, + accuracy_reports=(), + invariance_reports=(), + logprob_smoke=None, + backend_provenance=provenance, + candidate_id="x", + device="x", + compute_capability=None, + seed=manifest.seed, + fallback_reason=None, + passed=True, + provenance_valid=True, + metadata_valid=True, + ).to_dict() + ) + + def test_provenance_failure_fails_report(self, contract, manifest): + op = _make_identity_op() + bad_provenance = _make_provenance("cuda_bf16", "cuda", "triton") + report = assert_forward_batch_invariant( + op, + contract=contract, + manifest=manifest, + backend_profile="cuda_bf16", + provenance=bad_provenance, + gold_fn=_make_identity_op(), + op_class="logprob", + dtype=torch.bfloat16, + op_name="test_op", + include_logprob_smoke=False, + ) + assert report.provenance_valid is False + assert report.passed is False + + +class TestConfigMatrix: + def test_config_matrix_covers_c2_cells(self, manifest): + configs = build_config_matrix(manifest) + config_ids = [c.config_id for c in configs] + assert any("BN/full" in cid for cid in config_ids) + assert any("BN/chunked" in cid for cid in config_ids) + assert any("B1-singleton_aggregate/full" in cid for cid in config_ids) + assert any("B1-singleton_aggregate/chunked" in cid for cid in config_ids) + assert any("permuted" in cid for cid in config_ids) + assert any("padded_right" in cid for cid in config_ids) + assert any("padded_left" in cid for cid in config_ids) + + def test_canonical_config_exists(self, manifest): + configs = build_config_matrix(manifest) + canonical = [c for c in configs if c.is_canonical] + assert len(canonical) == 1 + assert canonical[0].config_id == "BN/full" + + +class TestLogprobSmoke: + def test_logprob_smoke_passes_for_identical(self, contract, manifest): + op = _make_identity_op(0.0) + gold = _make_identity_op(0.0) + report = assert_forward_batch_invariant( + op, + contract=contract, + manifest=manifest, + backend_profile="cuda_bf16", + provenance=_make_provenance(), + gold_fn=gold, + op_class="logprob", + dtype=torch.bfloat16, + op_name="test_op", + include_logprob_smoke=True, + ) + assert report.logprob_smoke is not None + assert report.logprob_smoke.passed