diff --git a/code/benchmarks/__init__.py b/code/benchmarks/__init__.py new file mode 100644 index 0000000..81f0abb --- /dev/null +++ b/code/benchmarks/__init__.py @@ -0,0 +1,4 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Benchmark dataset integrations.""" diff --git a/code/benchmarks/google_qec.py b/code/benchmarks/google_qec.py new file mode 100644 index 0000000..0fda551 --- /dev/null +++ b/code/benchmarks/google_qec.py @@ -0,0 +1,298 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Google Quantum AI QEC benchmark dataset integration. + +The source dataset is the Zenodo record for "Quantum error correction below the +surface code threshold". This module deliberately treats the data as an +external benchmark archive: the files are multi-GB zip archives with their own +README files and are not committed to this repository. +""" + +from __future__ import annotations + +import hashlib +import json +import shutil +import urllib.request +import zipfile +from dataclasses import asdict, dataclass +from pathlib import Path +from typing import Iterable, Sequence + + +GOOGLE_QEC_RECORD_ID = 13273331 +GOOGLE_QEC_RECORD_URL = f"https://zenodo.org/api/records/{GOOGLE_QEC_RECORD_ID}" +GOOGLE_QEC_RECORD_HTML = f"https://zenodo.org/records/{GOOGLE_QEC_RECORD_ID}" + +DEFAULT_BENCHMARK_KEY = "google_105Q_surface_code_d3_d5_d7.zip" + + +@dataclass(frozen=True) +class GoogleQECFile: + key: str + size_bytes: int + md5: str + url: str + code_family: str + distances: tuple[int, ...] + + +@dataclass(frozen=True) +class GoogleQECManifest: + record_id: int + title: str + license_id: str + record_url: str + files: tuple[GoogleQECFile, ...] + + def by_key(self) -> dict[str, GoogleQECFile]: + return {entry.key: entry for entry in self.files} + + +@dataclass(frozen=True) +class DownloadItem: + entry: GoogleQECFile + path: Path + exists: bool + + +@dataclass(frozen=True) +class DownloadPlan: + root: Path + items: tuple[DownloadItem, ...] + required_bytes: int + + +@dataclass(frozen=True) +class GoogleQECIndex: + root: Path + manifest_path: Path | None + archives: dict[str, Path] + extracted_dirs: dict[str, Path] + + +def _infer_code_family(key: str) -> str: + if "surface_code" in key: + return "surface" + if "repetition_code" in key: + return "repetition" + return "unknown" + + +def _infer_distances(key: str) -> tuple[int, ...]: + stem = key.removesuffix(".zip") + values = [] + for part in stem.split("_"): + if len(part) > 1 and part[0] == "d" and part[1:].isdigit(): + values.append(int(part[1:])) + return tuple(values) + + +def parse_zenodo_record(record: dict) -> GoogleQECManifest: + """Parse the Zenodo API response into a stable local manifest.""" + + files = [] + for file_info in record.get("files", []): + checksum = str(file_info.get("checksum", "")) + if not checksum.startswith("md5:"): + raise ValueError(f"Unsupported checksum for {file_info.get('key')!r}: {checksum!r}") + key = str(file_info["key"]) + files.append( + GoogleQECFile( + key=key, + size_bytes=int(file_info["size"]), + md5=checksum.split(":", 1)[1], + url=str(file_info["links"]["self"]), + code_family=_infer_code_family(key), + distances=_infer_distances(key), + ) + ) + + metadata = record.get("metadata", {}) + license_info = metadata.get("license") or {} + return GoogleQECManifest( + record_id=int(record["id"]), + title=str(metadata.get("title", record.get("title", ""))), + license_id=str(license_info.get("id", "")), + record_url=str(record.get("links", {}).get("self_html", GOOGLE_QEC_RECORD_HTML)), + files=tuple(sorted(files, key=lambda entry: entry.size_bytes)), + ) + + +def fetch_zenodo_manifest(url: str = GOOGLE_QEC_RECORD_URL, timeout: float = 60.0) -> GoogleQECManifest: + """Fetch and parse the official Zenodo record.""" + + with urllib.request.urlopen(url, timeout=timeout) as response: + payload = json.loads(response.read().decode("utf-8")) + return parse_zenodo_record(payload) + + +def build_download_plan( + manifest: GoogleQECManifest, + root: Path, + keys: Sequence[str] | None = None, +) -> DownloadPlan: + """Build a concrete download plan without performing network or disk writes.""" + + selected_keys = tuple(keys) if keys else (DEFAULT_BENCHMARK_KEY,) + by_key = manifest.by_key() + missing = [key for key in selected_keys if key not in by_key] + if missing: + raise KeyError(f"Unknown Google QEC benchmark file(s): {missing}") + + root = Path(root) + items = [] + required = 0 + for key in selected_keys: + entry = by_key[key] + path = root / entry.key + exists = path.exists() + items.append(DownloadItem(entry=entry, path=path, exists=exists)) + if not exists: + required += entry.size_bytes + return DownloadPlan(root=root, items=tuple(items), required_bytes=required) + + +def ensure_sufficient_space(path: Path, required_bytes: int, margin: float = 1.10) -> None: + """Raise before starting a large download if the filesystem is too full.""" + + if required_bytes <= 0: + return + usage = shutil.disk_usage(path) + needed = int(required_bytes * float(margin)) + if usage.free < needed: + raise RuntimeError( + f"Not enough free space under {path}: need at least {needed:,} bytes " + f"including margin, found {usage.free:,} bytes" + ) + + +def _md5_file(path: Path, chunk_size: int = 16 * 1024 * 1024) -> str: + digest = hashlib.md5() + with path.open("rb") as f: + while True: + chunk = f.read(chunk_size) + if not chunk: + break + digest.update(chunk) + return digest.hexdigest() + + +def verify_archive(path: Path, entry: GoogleQECFile) -> None: + if path.stat().st_size != entry.size_bytes: + raise RuntimeError( + f"Size mismatch for {path}: expected {entry.size_bytes}, got {path.stat().st_size}" + ) + got = _md5_file(path) + if got != entry.md5: + raise RuntimeError(f"MD5 mismatch for {path}: expected {entry.md5}, got {got}") + + +def build_download_request(entry: GoogleQECFile, resume_from: int = 0) -> urllib.request.Request: + """Build a request for a benchmark archive, optionally using HTTP Range.""" + + headers = {} + if int(resume_from) > 0: + headers["Range"] = f"bytes={int(resume_from)}-" + return urllib.request.Request(entry.url, headers=headers) + + +def download_entry(entry: GoogleQECFile, path: Path, force: bool = False) -> Path: + """Download one benchmark archive and verify size + md5.""" + + path.parent.mkdir(parents=True, exist_ok=True) + if path.exists() and not force: + verify_archive(path, entry) + return path + + tmp_path = path.with_suffix(path.suffix + ".part") + if force and tmp_path.exists(): + tmp_path.unlink() + + resume_from = tmp_path.stat().st_size if tmp_path.exists() else 0 + if resume_from >= entry.size_bytes: + tmp_path.replace(path) + verify_archive(path, entry) + return path + + request = build_download_request(entry, resume_from=resume_from) + with urllib.request.urlopen(request, timeout=60.0) as response: + status = getattr(response, "status", None) or response.getcode() + mode = "ab" if resume_from > 0 and status == 206 else "wb" + if mode == "wb": + resume_from = 0 + with tmp_path.open(mode) as out: + while True: + chunk = response.read(16 * 1024 * 1024) + if not chunk: + break + out.write(chunk) + tmp_path.replace(path) + verify_archive(path, entry) + return path + + +def extract_archive(path: Path, output_dir: Path | None = None) -> Path: + """Extract a downloaded benchmark zip next to the archive by default.""" + + target = output_dir or path.with_suffix("") + target.mkdir(parents=True, exist_ok=True) + with zipfile.ZipFile(path) as zf: + zf.extractall(target) + return target + + +class GoogleQECBenchmarkStore: + """Local project store for Google QEC benchmark archives.""" + + def __init__(self, root: Path | str = "benchmarks/google_qec"): + self.root = Path(root) + + @property + def manifest_path(self) -> Path: + return self.root / "manifest.json" + + def write_manifest(self, manifest: GoogleQECManifest) -> Path: + self.root.mkdir(parents=True, exist_ok=True) + payload = asdict(manifest) + self.manifest_path.write_text(json.dumps(payload, indent=2, sort_keys=True) + "\n") + return self.manifest_path + + def index(self) -> GoogleQECIndex: + archives = {path.name: path for path in sorted(self.root.glob("*.zip"))} + extracted_dirs = { + path.name: path + for path in sorted(self.root.iterdir()) if path.is_dir() and path.name != "__pycache__" + } if self.root.exists() else {} + manifest_path = self.manifest_path if self.manifest_path.exists() else None + return GoogleQECIndex( + root=self.root, + manifest_path=manifest_path, + archives=archives, + extracted_dirs=extracted_dirs, + ) + + def download( + self, + manifest: GoogleQECManifest, + keys: Sequence[str] | None = None, + *, + force: bool = False, + extract: bool = False, + check_space: bool = True, + ) -> DownloadPlan: + self.root.mkdir(parents=True, exist_ok=True) + plan = build_download_plan(manifest, self.root, keys) + if check_space: + ensure_sufficient_space(self.root, plan.required_bytes) + self.write_manifest(manifest) + for item in plan.items: + archive_path = download_entry(item.entry, item.path, force=force) + if extract: + extract_archive(archive_path) + return plan + + +def benchmark_keys(files: Iterable[GoogleQECFile]) -> list[str]: + return [entry.key for entry in sorted(files, key=lambda entry: (entry.code_family, entry.size_bytes))] diff --git a/code/examples/infer_ood.py b/code/examples/infer_ood.py new file mode 100644 index 0000000..7d255e5 --- /dev/null +++ b/code/examples/infer_ood.py @@ -0,0 +1,123 @@ +#!/usr/bin/env python3 +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Run released pre-decoders on the fixed training-axis OOD grid.""" + +from __future__ import annotations + +import argparse +import sys +from pathlib import Path + +CODE_ROOT = Path(__file__).resolve().parents[1] +if str(CODE_ROOT) not in sys.path: + sys.path.insert(0, str(CODE_ROOT)) + +from scripts.experiments.unknown_noise.generate_unknown_axismix_grid_u1p2_5p0_configs import ( # noqa: E402 + write_axismix_grid_configs, +) +from scripts.qadapt_example_utils import ( # noqa: E402 + InferenceJob, + add_common_inference_args, + build_paired_command, + parse_gpus, + run_jobs, +) + + +PAPER_DISTANCES = (7, 9) +PAPER_MULTIPLIERS = (1.2, 1.5, 2.0, 2.5, 3.0) + + +def parse_distances(value: str) -> list[int]: + result = [int(item.strip()) for item in value.split(",") if item.strip()] + if not result or result != sorted(set(result)): + raise argparse.ArgumentTypeError( + "distances must be a non-empty, increasing comma-separated list" + ) + return result + + +def parse_multipliers(value: str) -> list[float]: + result = [float(item.strip()) for item in value.split(",") if item.strip()] + if not result or result != sorted(set(result)) or any(item <= 0 for item in result): + raise argparse.ArgumentTypeError( + "multipliers must be a non-empty, increasing comma-separated list " + "of positive numbers" + ) + return result + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--distances", + type=parse_distances, + default=list(PAPER_DISTANCES), + help="Comma-separated distances; defaults to the paper's d=7,9 grid.", + ) + parser.add_argument("--n-rounds", type=int, default=9) + parser.add_argument( + "--multipliers", + type=parse_multipliers, + default=list(PAPER_MULTIPLIERS), + help="Comma-separated OOD multipliers; defaults to the paper's 1.2--3.0 grid.", + ) + parser.add_argument( + "--generated-config-dir", + type=Path, + default=Path("outputs/generated_configs/ood"), + ) + parser.add_argument( + "--manifest", + type=Path, + default=Path("outputs/generated_configs/ood/manifest.json"), + ) + add_common_inference_args( + parser, + default_output_dir=Path("outputs/examples/released_models/ood"), + ) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + _, manifest = write_axismix_grid_configs( + base_config="conf/examples/qadapt/config_qadapt_t0_base.yaml", + output_dir=args.generated_config_dir, + manifest=args.manifest, + grid_multipliers=args.multipliers, + ) + jobs = [] + for distance in args.distances: + for environment in manifest["environments"]: + config_file = args.generated_config_dir / environment["config_filename"] + label = ( + f"d{distance}_{environment['env_key']}_" + f"{environment['multiplier_key']}" + ) + output_path = args.output_dir / f"d{distance}" / f"{label}.json" + jobs.append( + InferenceJob( + label=label, + command=build_paired_command( + args, + config_file=config_file, + output_path=output_path, + distance=distance, + n_rounds=args.n_rounds, + ), + output_path=output_path, + ) + ) + run_jobs( + jobs, + gpus=parse_gpus(args.gpus), + parallelism=args.parallelism, + resume=args.resume, + dry_run=args.dry_run, + ) + + +if __name__ == "__main__": + main() diff --git a/code/examples/infer_t0_t4.py b/code/examples/infer_t0_t4.py new file mode 100644 index 0000000..d3fd27b --- /dev/null +++ b/code/examples/infer_t0_t4.py @@ -0,0 +1,109 @@ +#!/usr/bin/env python3 +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Run released pre-decoders on the five T0-T4 simulated noise tasks.""" + +from __future__ import annotations + +import argparse +import sys +from pathlib import Path + +CODE_ROOT = Path(__file__).resolve().parents[1] +if str(CODE_ROOT) not in sys.path: + sys.path.insert(0, str(CODE_ROOT)) + +from scripts.qadapt_example_utils import ( # noqa: E402 + InferenceJob, + TASK_CONFIGS, + add_common_inference_args, + build_paired_command, + parse_gpus, + run_jobs, +) + + +TASK_BY_ID = { + f"T{index}": (task_key, config_name) + for index, (task_key, config_name) in enumerate(TASK_CONFIGS) +} + + +def parse_distances(value: str) -> list[int]: + result = [int(item.strip()) for item in value.split(",") if item.strip()] + if not result or result != sorted(set(result)): + raise argparse.ArgumentTypeError( + "distances must be a non-empty, increasing comma-separated list" + ) + return result + + +def parse_tasks(value: str) -> list[str]: + result = [item.strip().upper() for item in value.split(",") if item.strip()] + if not result or len(result) != len(set(result)): + raise argparse.ArgumentTypeError( + "tasks must be a non-empty comma-separated subset of T0,T1,T2,T3,T4" + ) + unknown = [item for item in result if item not in TASK_BY_ID] + if unknown: + raise argparse.ArgumentTypeError(f"unknown task(s): {','.join(unknown)}") + return result + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--distances", + type=parse_distances, + default=[9], + help=( + "Comma-separated distances. Use 7,9 with --tasks T0 for the " + "paper's mapped-noise geometry; the default is release coverage at d=9." + ), + ) + parser.add_argument( + "--tasks", + type=parse_tasks, + default=list(TASK_BY_ID), + help="Comma-separated task subset; defaults to T0,T1,T2,T3,T4.", + ) + parser.add_argument("--n-rounds", type=int, default=9) + add_common_inference_args( + parser, + default_output_dir=Path("outputs/examples/released_models/t0_t4"), + ) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + jobs = [] + for distance in args.distances: + for task_id in args.tasks: + task_key, config_name = TASK_BY_ID[task_id] + label = f"d{distance}_{task_key}" + output_path = args.output_dir / f"d{distance}" / f"{task_key}.json" + jobs.append( + InferenceJob( + label=label, + command=build_paired_command( + args, + config_name=config_name, + output_path=output_path, + distance=distance, + n_rounds=args.n_rounds, + ), + output_path=output_path, + ) + ) + run_jobs( + jobs, + gpus=parse_gpus(args.gpus), + parallelism=args.parallelism, + resume=args.resume, + dry_run=args.dry_run, + ) + + +if __name__ == "__main__": + main() diff --git a/code/examples/infer_willow.py b/code/examples/infer_willow.py new file mode 100644 index 0000000..fc0e7f7 --- /dev/null +++ b/code/examples/infer_willow.py @@ -0,0 +1,115 @@ +#!/usr/bin/env python3 +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Reproduce the paper's d=5/d=7, ten-round Google Willow evaluation.""" + +from __future__ import annotations + +import argparse +import os +import shlex +import sys +from pathlib import Path + +CODE_ROOT = Path(__file__).resolve().parents[1] +if str(CODE_ROOT) not in sys.path: + sys.path.insert(0, str(CODE_ROOT)) + +from scripts.qadapt_example_utils import ( # noqa: E402 + add_common_inference_args, + checkpoint_specs, + parse_gpus, +) + + +def parse_args(argv: list[str] | None = None) -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--benchmark-root", + type=Path, + default=Path("benchmarks/google_qec/google_105Q_surface_code_d3_d5_d7"), + ) + parser.add_argument( + "--distances", + nargs="+", + type=int, + default=[5, 7], + help="Paper default: d=5 and d=7.", + ) + parser.add_argument( + "--rounds", + nargs="+", + type=int, + default=[10], + help="Paper default: ten syndrome-extraction rounds.", + ) + add_common_inference_args( + parser, + default_output_dir=Path("outputs/examples/released_models/willow"), + default_num_samples=0, + ) + return parser.parse_args(argv) + + +def main(argv: list[str] | None = None) -> int: + args = parse_args(argv) + output_path = args.output_dir / "results.json" + if args.resume and output_path.is_file(): + print(f"[resume] output exists: {output_path}") + return 0 + + selected_gpus = parse_gpus(args.gpus) + bases = ["X", "Z"] if args.basis == "both" else [args.basis] + specs = checkpoint_specs(args) + command_preview = [ + str(args.python), + "-m", + "scripts.providers.google_qec_decoder_benchmark", + "--benchmark-root", + str(args.benchmark_root), + "--distances", + *(str(value) for value in args.distances), + "--rounds", + *(str(value) for value in args.rounds), + "--bases", + *bases, + "--models", + *(spec.name for spec in specs), + "--max-shots", + str(args.num_samples), + "--batch-size", + str(args.batch_size), + "--latency-shots", + str(args.latency_num_samples), + "--output", + str(output_path), + ] + if args.dry_run: + print( + f"[dry-run] gpu={selected_gpus[0]} seed={args.seed} " + + shlex.join(command_preview) + ) + for spec in specs: + print( + f"[dry-run] model {spec.name}: " + f"model_id={spec.model_id} checkpoint={spec.checkpoint}" + ) + return 0 + + os.environ["CUDA_VISIBLE_DEVICES"] = selected_gpus[0] + from scripts.providers import google_qec_decoder_benchmark as benchmark + + benchmark.DEFAULT_MODELS = { + spec.name: benchmark.BenchmarkModel( + spec.name, + spec.model_id, + spec.checkpoint, + ) + for spec in specs + } + benchmark.DEFAULT_BENCHMARK_ROOT = args.benchmark_root + return benchmark.main(command_preview[3:]) + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/code/model/checkpoint_loader.py b/code/model/checkpoint_loader.py new file mode 100644 index 0000000..2387be2 --- /dev/null +++ b/code/model/checkpoint_loader.py @@ -0,0 +1,48 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Load one explicitly identified pre-decoder checkpoint.""" + +from __future__ import annotations + +from pathlib import Path +from typing import Any + +import torch + + +def load_model_checkpoint( + cfg: Any, + *, + checkpoint: Path, + model_id: int, + distributed: Any, +) -> torch.nn.Module: + """Load a ``.pt`` or ``.safetensors`` checkpoint for one public model ID.""" + + path = Path(checkpoint).expanduser().resolve() + if not path.is_file(): + raise FileNotFoundError(f"Checkpoint not found: {path}") + + if path.suffix.lower() != ".safetensors": + from workflows.run import _load_model + + cfg.model_checkpoint_file = str(path) + return _load_model(cfg, distributed) + + from export.safetensors_utils import load_safetensors + + model, metadata = load_safetensors( + str(path), + model_id=None, + device=str(distributed.device), + ) + embedded_model_id = metadata.get("model_id") + if embedded_model_id is not None and str(embedded_model_id) != str(model_id): + raise ValueError( + f"SafeTensors model_id mismatch for {path}: " + f"CLI requested {model_id}, file metadata contains {embedded_model_id}" + ) + cfg.enable_fp16 = metadata.get("quant_format") == "fp16" + cfg.model_checkpoint_file = str(path) + return model diff --git a/code/model/factory.py b/code/model/factory.py index cd436a7..dabdd81 100644 --- a/code/model/factory.py +++ b/code/model/factory.py @@ -1,5 +1,6 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 +# Modified in 2026 for the QAdapt Hugging Face release: added HTNet dispatch. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. @@ -38,6 +39,9 @@ class ModelFactory: from model.predecoder import PreDecoderModelMemory_v1 model = PreDecoderModelMemory_v1(cfg) return model + elif cfg.model.version == "htnet": + from model.qadapt import HTnet + return HTnet(cfg) elif cfg.model.version == "predecoder_memory_v2": from model.predecoder import PreDecoderModelMemory_v2 model = PreDecoderModelMemory_v2(cfg) diff --git a/code/model/qadapt.py b/code/model/qadapt.py new file mode 100644 index 0000000..d2893c7 --- /dev/null +++ b/code/model/qadapt.py @@ -0,0 +1,252 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""HTNet architecture used by the QAdapt surface-code pre-decoder.""" + +from __future__ import annotations + +import torch +from torch import nn + + +def _activation(name: str) -> nn.Module: + if name == "relu": + return nn.ReLU() + if name == "gelu": + return nn.GELU(approximate="tanh") + if name == "leakyrelu": + return nn.LeakyReLU() + raise ValueError(f"Unsupported activation: {name}") + + +class AdaptiveBranchFusion3D(nn.Module): + """Input-adaptive fusion of spatial, temporal, and joint branches.""" + + def __init__(self, channels: int, reduction: int, activation_name: str): + super().__init__() + self.num_branches = 3 + hidden_channels = max(1, channels // (reduction + 2)) + self.pool = nn.AdaptiveAvgPool3d(1) + self.weight_net = nn.Sequential( + nn.Conv3d(channels * self.num_branches, hidden_channels, kernel_size=1), + _activation(activation_name), + nn.Conv3d(hidden_channels, channels * self.num_branches, kernel_size=1), + ) + nn.init.zeros_(self.weight_net[-1].weight) + nn.init.zeros_(self.weight_net[-1].bias) + + def forward( + self, + spatial: torch.Tensor, + temporal: torch.Tensor, + joint: torch.Tensor, + ) -> torch.Tensor: + batch_size, channels = spatial.shape[:2] + pooled = torch.cat( + [self.pool(spatial), self.pool(temporal), self.pool(joint)], + dim=1, + ) + weights = self.weight_net(pooled).view( + batch_size, + self.num_branches, + channels, + 1, + 1, + 1, + ) + weights = torch.softmax(weights, dim=1) + fused = ( + weights[:, 0] * spatial + + weights[:, 1] * temporal + + weights[:, 2] * joint + ) + return fused * self.num_branches + + +class AxisChannelGate3D(nn.Module): + """Joint channel, temporal-axis, and spatial-axis gating.""" + + def __init__(self, channels: int, reduction: int, activation_name: str): + super().__init__() + hidden_channels = max(1, channels // reduction) + self.channel_net = nn.Sequential( + nn.AdaptiveAvgPool3d(1), + nn.Conv3d(channels, hidden_channels, kernel_size=1), + _activation(activation_name), + nn.Conv3d(hidden_channels, channels, kernel_size=1), + ) + self.temporal_conv = nn.Conv3d( + 1, + 1, + kernel_size=(3, 1, 1), + padding=(1, 0, 0), + ) + self.spatial_conv = nn.Conv3d( + 1, + 1, + kernel_size=(1, 3, 3), + padding=(0, 1, 1), + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + channel_logits = self.channel_net(x) + temporal_logits = self.temporal_conv( + x.mean(dim=(1, 3, 4), keepdim=True) + ) + spatial_logits = self.spatial_conv(x.mean(dim=(1, 2), keepdim=True)) + return x * torch.sigmoid( + channel_logits + temporal_logits + spatial_logits + ) + + +class STFusionBlockV2(nn.Module): + """One HTNet block with separable space/time and grouped joint evidence.""" + + def __init__( + self, + channels: int, + expand_channels: int, + joint_groups: int, + norm_groups: int, + se_reduction: int, + dropout_p: float, + activation_name: str, + ): + super().__init__() + if expand_channels % joint_groups != 0: + raise ValueError( + "expand_channels must be divisible by joint_groups: " + f"{expand_channels} vs {joint_groups}" + ) + if channels % norm_groups != 0 or expand_channels % norm_groups != 0: + raise ValueError( + "channels and expand_channels must be divisible by norm_groups" + ) + + self.pre = nn.Sequential( + nn.GroupNorm(num_groups=norm_groups, num_channels=channels), + nn.Conv3d(channels, expand_channels, kernel_size=1), + _activation(activation_name), + ) + self.spatial = nn.Conv3d( + expand_channels, + expand_channels, + kernel_size=(1, 3, 3), + padding=(0, 1, 1), + groups=expand_channels, + ) + self.temporal = nn.Conv3d( + expand_channels, + expand_channels, + kernel_size=(3, 1, 1), + padding=(1, 0, 0), + groups=expand_channels, + ) + self.joint = nn.Sequential( + nn.GroupNorm( + num_groups=norm_groups, + num_channels=expand_channels, + ), + nn.Conv3d( + expand_channels, + expand_channels, + kernel_size=3, + padding=1, + groups=joint_groups, + ), + ) + self.branch_fusion = AdaptiveBranchFusion3D( + expand_channels, + se_reduction, + activation_name, + ) + self.branch_mixer = nn.Sequential( + nn.Conv3d( + expand_channels, + expand_channels, + kernel_size=1, + groups=joint_groups, + ), + _activation(activation_name), + ) + self.project = nn.Sequential( + nn.Conv3d(expand_channels, channels, kernel_size=1), + _activation(activation_name), + ) + self.gate = AxisChannelGate3D( + channels, + se_reduction, + activation_name, + ) + self.dropout = nn.Dropout3d(p=dropout_p) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + residual = x + y = self.pre(x) + y = self.branch_fusion( + self.spatial(y), + self.temporal(y), + self.joint(y), + ) + y = self.branch_mixer(y) + y = self.project(y) + y = self.gate(y) + return residual + self.dropout(y) + + +class HTnet(nn.Module): + """QAdapt HTNet model with an effective receptive field of nine.""" + + def __init__(self, cfg): + super().__init__() + self.distance = cfg.distance + self.n_rounds = cfg.n_rounds + self.dropout_p = cfg.model.dropout_p + + input_channels = int(cfg.model.input_channels) + out_channels = int(cfg.model.out_channels) + channels = int(cfg.model.channels) + expand_channels = int(cfg.model.expand_channels) + num_blocks = int(cfg.model.num_blocks) + joint_groups = int(cfg.model.joint_groups) + norm_groups = int(cfg.model.norm_groups) + se_reduction = int(cfg.model.se_reduction) + activation_name = str(cfg.model.activation) + + self.stem = nn.Sequential( + nn.Conv3d(input_channels, channels, kernel_size=3, padding=1), + nn.GroupNorm(num_groups=norm_groups, num_channels=channels), + _activation(activation_name), + ) + self.blocks = nn.Sequential( + *[ + STFusionBlockV2( + channels=channels, + expand_channels=expand_channels, + joint_groups=joint_groups, + norm_groups=norm_groups, + se_reduction=se_reduction, + dropout_p=self.dropout_p, + activation_name=activation_name, + ) + for _ in range(num_blocks) + ] + ) + self.head_norm = nn.GroupNorm( + num_groups=norm_groups, + num_channels=channels, + ) + self.head_hidden = nn.Conv3d( + channels + input_channels, + channels, + kernel_size=1, + ) + self.head_activation = _activation(activation_name) + self.head_out = nn.Conv3d(channels, out_channels, kernel_size=1) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + y = self.blocks(self.stem(x)) + y = self.head_norm(y) + y = torch.cat([y, x], dim=1) + y = self.head_activation(self.head_hidden(y)) + return self.head_out(y) diff --git a/code/model/registry.py b/code/model/registry.py index a17bf9f..96e31af 100644 --- a/code/model/registry.py +++ b/code/model/registry.py @@ -1,5 +1,6 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 +# Modified in 2026 for the QAdapt Hugging Face release: added model ID 111. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. @@ -52,6 +53,12 @@ class PublicModelSpec: kernel_size: List[int] receptive_field: int model_version: str = "predecoder_memory_v1" + channels: Optional[int] = None + expand_channels: Optional[int] = None + num_blocks: Optional[int] = None + joint_groups: Optional[int] = None + norm_groups: Optional[int] = None + se_reduction: Optional[int] = None # Non-convolutional models (e.g. the cascade/bottleneck model "B") are not # described by num_filters/kernel_size. For those, `model_overrides` carries # the full `model.*` block that should be written into the merged config. @@ -86,6 +93,21 @@ _MODEL_SPECS: Dict[Union[int, str], PublicModelSpec] = { kernel_size=[3, 3, 3, 3], receptive_field=compute_receptive_field([3, 3, 3, 3]), ), + # QAdapt: three HTNet blocks with an effective receptive field of nine. + 111: + PublicModelSpec( + model_id=111, + num_filters=[112, 112, 112, 112, 4], + kernel_size=[3, 3, 3, 3], + receptive_field=compute_receptive_field([3, 3, 3, 3]), + model_version="htnet", + channels=112, + expand_channels=168, + num_blocks=3, + joint_groups=6, + norm_groups=8, + se_reduction=4, + ), # Model 2: 4 conv layers, k=3, wider 2: PublicModelSpec( @@ -152,13 +174,17 @@ def _normalize_model_id(model_id: Union[int, str]) -> Union[int, str]: def get_model_spec(model_id: Union[int, str]) -> PublicModelSpec: - """Return the public model spec for a given model_id (1..5 or "B").""" + """Return a public model spec, including QAdapt model_id 111.""" try: key = _normalize_model_id(model_id) except Exception as e: - raise ValueError(f"model_id must be one of [1..5] or 'B', got: {model_id!r}") from e + raise ValueError( + f"model_id must be one of [1..5], 111, or 'B', got: {model_id!r}" + ) from e if key == 0: raise ValueError("model_id=0 is not supported in the public release") if key not in _MODEL_SPECS: - raise ValueError(f"model_id must be one of [1..5] or 'B', got: {model_id!r}") + raise ValueError( + f"model_id must be one of [1..5], 111, or 'B', got: {model_id!r}" + ) return _MODEL_SPECS[key] diff --git a/code/scripts/__init__.py b/code/scripts/__init__.py new file mode 100644 index 0000000..ec0797d --- /dev/null +++ b/code/scripts/__init__.py @@ -0,0 +1,4 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Developer and experiment command modules.""" diff --git a/code/scripts/config_paths.py b/code/scripts/config_paths.py new file mode 100644 index 0000000..18aef86 --- /dev/null +++ b/code/scripts/config_paths.py @@ -0,0 +1,72 @@ +#!/usr/bin/env python3 +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Shared helpers for Hydra config names stored below ``conf/``.""" + +from __future__ import annotations + +from pathlib import Path +from typing import Any, Mapping + + +CODE_ROOT = Path(__file__).resolve().parents[1] +REPO_ROOT = CODE_ROOT.parent +CONF_ROOT = REPO_ROOT / "conf" + + +def rel(path: str | Path) -> Path: + path = Path(path) + return path if path.is_absolute() else REPO_ROOT / path + + +def config_path(config_name: str | Path) -> Path: + """Return the YAML path for a Hydra config name below ``conf/``. + + Configs are grouped in nested preset and experiment directories. For callers + that still pass a historical basename, return its unique recursive match. + """ + raw = str(config_name) + if raw.endswith(".yaml"): + raw = raw[:-5] + direct = CONF_ROOT / f"{raw}.yaml" + if direct.exists() or "/" in raw or "\\" in raw: + return direct + matches = sorted(CONF_ROOT.rglob(f"{raw}.yaml")) + if len(matches) == 1: + return matches[0] + return direct + + +def config_name_from_path(path: str | Path) -> str: + """Return the Hydra config name for a YAML path when it is below a ``conf/`` dir.""" + path = rel(path) + try: + relative = path.relative_to(CONF_ROOT) + except ValueError: + parts = path.parts + if "conf" not in parts: + return path.stem + conf_index = len(parts) - 1 - list(reversed(parts)).index("conf") + relative = Path(*parts[conf_index + 1 :]) + return relative.with_suffix("").as_posix() + + +def config_basename(config_name: str | Path) -> str: + """Return the final component of a Hydra config name.""" + raw = str(config_name) + if raw.endswith(".yaml"): + raw = raw[:-5] + return Path(raw).name + + +def config_lookup_with_basename( + environments: list[Mapping[str, Any]] | tuple[Mapping[str, Any], ...], +) -> dict[str, dict[str, Any]]: + """Map both full config names and historical basenames to manifest rows.""" + lookup: dict[str, dict[str, Any]] = {} + for env in environments: + item = dict(env) + full = str(item["config_name"]) + lookup[full] = item + lookup.setdefault(config_basename(full), item) + return lookup diff --git a/code/scripts/download_google_qec_benchmark.py b/code/scripts/download_google_qec_benchmark.py new file mode 100644 index 0000000..a02de79 --- /dev/null +++ b/code/scripts/download_google_qec_benchmark.py @@ -0,0 +1,89 @@ +#!/usr/bin/env python3 +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Download Google Quantum AI QEC benchmark archives from Zenodo.""" + +from __future__ import annotations + +import argparse +from pathlib import Path + +from benchmarks.google_qec import ( + DEFAULT_BENCHMARK_KEY, + GoogleQECBenchmarkStore, + benchmark_keys, + build_download_plan, + fetch_zenodo_manifest, +) + + +def _parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--output-dir", + type=Path, + default=Path("benchmarks/google_qec"), + help="Directory for manifest and downloaded zip archives.", + ) + parser.add_argument( + "--file", + action="append", + dest="files", + help=( + "Zenodo file key to download. May be repeated. " + f"Default: {DEFAULT_BENCHMARK_KEY}" + ), + ) + parser.add_argument("--all", action="store_true", help="Download all Google QEC archives.") + parser.add_argument("--list", action="store_true", help="List available archives and exit.") + parser.add_argument("--manifest-only", action="store_true", help="Only write manifest.json.") + parser.add_argument("--extract", action="store_true", help="Extract downloaded zip archives.") + parser.add_argument("--force", action="store_true", help="Re-download archives that already exist.") + parser.add_argument("--skip-space-check", action="store_true", help="Skip free-space guard.") + return parser.parse_args() + + +def main() -> int: + args = _parse_args() + manifest = fetch_zenodo_manifest() + store = GoogleQECBenchmarkStore(args.output_dir) + + if args.list: + for entry in manifest.files: + gib = entry.size_bytes / (1024**3) + distances = ",".join(str(d) for d in entry.distances) or "unknown" + print(f"{entry.key}\t{gib:.2f} GiB\t{entry.code_family}\td={distances}") + return 0 + + if args.all: + keys = benchmark_keys(manifest.files) + else: + keys = tuple(args.files) if args.files else (DEFAULT_BENCHMARK_KEY,) + + store.write_manifest(manifest) + plan = build_download_plan(manifest, args.output_dir, keys) + print(f"Google QEC Zenodo record: {manifest.record_url}") + print(f"Output directory: {args.output_dir}") + for item in plan.items: + status = "exists" if item.exists else "download" + gib = item.entry.size_bytes / (1024**3) + print(f" [{status}] {item.entry.key} ({gib:.2f} GiB, md5={item.entry.md5})") + + if args.manifest_only: + print(f"Wrote manifest: {store.manifest_path}") + return 0 + + store.download( + manifest, + keys, + force=args.force, + extract=args.extract, + check_space=not args.skip_space_check, + ) + print("Download complete.") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/code/scripts/experiments/__init__.py b/code/scripts/experiments/__init__.py new file mode 100644 index 0000000..45fad83 --- /dev/null +++ b/code/scripts/experiments/__init__.py @@ -0,0 +1,4 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Experiment orchestration modules.""" diff --git a/code/scripts/experiments/unknown_noise/__init__.py b/code/scripts/experiments/unknown_noise/__init__.py new file mode 100644 index 0000000..d5bbf76 --- /dev/null +++ b/code/scripts/experiments/unknown_noise/__init__.py @@ -0,0 +1,4 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Unknown-noise experiment configuration and comparison modules.""" diff --git a/code/scripts/experiments/unknown_noise/generate_unknown_axismix_grid_u1p2_5p0_configs.py b/code/scripts/experiments/unknown_noise/generate_unknown_axismix_grid_u1p2_5p0_configs.py new file mode 100644 index 0000000..fb3d72d --- /dev/null +++ b/code/scripts/experiments/unknown_noise/generate_unknown_axismix_grid_u1p2_5p0_configs.py @@ -0,0 +1,348 @@ +#!/usr/bin/env python3 +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Generate fixed multiplier-grid training-axis mixed OOD noise configs.""" + +from __future__ import annotations + +import argparse +import json +import sys +from itertools import combinations +from pathlib import Path +from typing import Any, Mapping, Sequence + +from omegaconf import OmegaConf + +CODE_ROOT = Path(__file__).resolve().parents[3] +REPO_ROOT = CODE_ROOT.parent +if str(CODE_ROOT) not in sys.path: + sys.path.insert(0, str(CODE_ROOT)) + +from qec.noise_model import NoiseModel # noqa: E402 +from scripts.config_paths import config_name_from_path # noqa: E402 + + +DEFAULT_BASE_CONFIG = "conf/examples/qadapt/config_qadapt_t0_base.yaml" +DESIGN_LABEL = "training-axis fixed multiplier grid OOD stress test" +DEFAULT_PREFIX = "config_unknown_axismix_grid_u1p2_5p0" +DEFAULT_OUTPUT_DIR = "outputs/generated_configs/ood" +DEFAULT_MANIFEST = "outputs/generated_configs/ood/manifest.json" +AXIS_ORDER = ("meas_all", "cnot_all", "idle_all", "z_bias") +GRID_MULTIPLIERS = (1.2, 1.5, 2.0, 2.5, 3.0, 3.5, 4.0, 4.5, 5.0) + +CNOT_KEYS = ( + "p_cnot_IX", + "p_cnot_IY", + "p_cnot_IZ", + "p_cnot_XI", + "p_cnot_XX", + "p_cnot_XY", + "p_cnot_XZ", + "p_cnot_YI", + "p_cnot_YX", + "p_cnot_YY", + "p_cnot_YZ", + "p_cnot_ZI", + "p_cnot_ZX", + "p_cnot_ZY", + "p_cnot_ZZ", +) + +AXES: dict[str, tuple[str, ...]] = { + "meas_all": ("p_meas_X", "p_meas_Z"), + "cnot_all": CNOT_KEYS, + "idle_all": ( + "p_idle_cnot_X", + "p_idle_cnot_Y", + "p_idle_cnot_Z", + "p_idle_spam_X", + "p_idle_spam_Y", + "p_idle_spam_Z", + ), + "z_bias": ( + "p_prep_X", + "p_meas_X", + "p_idle_cnot_Z", + "p_idle_spam_Z", + "p_cnot_IZ", + "p_cnot_XZ", + "p_cnot_YZ", + "p_cnot_ZI", + "p_cnot_ZX", + "p_cnot_ZY", + "p_cnot_ZZ", + ), +} + + +def rel(path: str | Path) -> Path: + path = Path(path) + return path if path.is_absolute() else REPO_ROOT / path + + +def _plain_mapping(value: Any) -> dict[str, float]: + raw = OmegaConf.to_container(value, resolve=True) if hasattr(value, "items") else value + if raw is None: + raise ValueError("base config does not contain data.noise_model") + return {str(key): float(item) for key, item in dict(raw).items()} + + +def load_base_noise_model(base_config: str | Path) -> dict[str, float]: + cfg = OmegaConf.load(rel(base_config)) + noise_cfg = getattr(getattr(cfg, "data", None), "noise_model", None) + noise = _plain_mapping(noise_cfg) + return NoiseModel.from_config_dict(noise).to_config_dict() + + +def multiplier_key(multiplier: float) -> str: + return f"m{float(multiplier):.1f}".replace(".", "p") + + +def _axis_signature(active_axes: Sequence[str]) -> str: + return "+".join(active_axes) + + +def _default_env_specs() -> list[dict[str, Any]]: + specs = [] + for size in (2, 3, 4): + for active_axes in combinations(AXIS_ORDER, size): + env_index = len(specs) + specs.append( + { + "env_index": env_index, + "env_key": f"e{env_index:02d}", + "active_axes": tuple(active_axes), + "axis_signature": _axis_signature(active_axes), + "combination_size": size, + "contains_z_bias": "z_bias" in active_axes, + "contains_cnot_z_bias": "cnot_all" in active_axes and "z_bias" in active_axes, + "purpose": f"{size}-axis fixed multiplier grid composite", + } + ) + return specs + + +DEFAULT_ENV_SPECS: list[dict[str, Any]] = _default_env_specs() + + +def _normalize_spec(raw_spec: Mapping[str, Any]) -> dict[str, Any]: + env_index = int(raw_spec["env_index"]) + active_axes = tuple(str(axis) for axis in raw_spec["active_axes"]) + if not 2 <= len(active_axes) <= 4: + raise ValueError(f"grid env must activate 2, 3, or 4 axes, got {active_axes}") + unknown = [axis for axis in active_axes if axis not in AXIS_ORDER] + if unknown: + raise ValueError(f"unknown grid axes: {unknown}") + if len(set(active_axes)) != len(active_axes): + raise ValueError(f"duplicate active axes: {active_axes}") + return { + "env_index": env_index, + "env_key": str(raw_spec.get("env_key", f"e{env_index:02d}")), + "active_axes": active_axes, + "axis_signature": str(raw_spec.get("axis_signature", _axis_signature(active_axes))), + "combination_size": len(active_axes), + "contains_z_bias": "z_bias" in active_axes, + "contains_cnot_z_bias": "cnot_all" in active_axes and "z_bias" in active_axes, + "purpose": str(raw_spec.get("purpose", f"{len(active_axes)}-axis fixed multiplier grid composite")), + } + + +def _parameter_multipliers( + base_noise: Mapping[str, float], + active_axes: Sequence[str], + multiplier: float, +) -> dict[str, float]: + multipliers = {key: 1.0 for key in base_noise} + for axis_name in active_axes: + if axis_name not in AXES: + raise ValueError(f"unknown training noise axis: {axis_name}") + for key in AXES[axis_name]: + if key not in base_noise: + raise ValueError(f"axis {axis_name} references missing noise parameter {key}") + multipliers[key] = max(multipliers[key], float(multiplier)) + return multipliers + + +def _probability_totals(noise: Mapping[str, float]) -> dict[str, float]: + return { + "cnot_total": sum(value for key, value in noise.items() if key.startswith("p_cnot_")), + "idle_cnot_total": sum(value for key, value in noise.items() if key.startswith("p_idle_cnot_")), + "idle_spam_total": sum(value for key, value in noise.items() if key.startswith("p_idle_spam_")), + } + + +def generate_axismix_grid_noise_models( + base_noise: Mapping[str, float], + env_specs: Sequence[Mapping[str, Any]] = DEFAULT_ENV_SPECS, + *, + grid_multipliers: Sequence[float] = GRID_MULTIPLIERS, +) -> list[dict[str, Any]]: + if not grid_multipliers: + raise ValueError("grid_multipliers must not be empty") + base = NoiseModel.from_config_dict(dict(base_noise)).to_config_dict() + generated = [] + for raw_spec in env_specs: + spec = _normalize_spec(raw_spec) + for multiplier_index, multiplier in enumerate(grid_multipliers): + multiplier = float(multiplier) + if multiplier < 0: + raise ValueError(f"multiplier must be non-negative, got {multiplier}") + param_multipliers = _parameter_multipliers(base, spec["active_axes"], multiplier) + axis_multipliers = { + axis: (multiplier if axis in spec["active_axes"] else 1.0) + for axis in AXIS_ORDER + } + noise = { + key: float(base_value) * float(param_multipliers[key]) + for key, base_value in base.items() + } + validated = NoiseModel.from_config_dict(noise) + noise = validated.to_config_dict() + generated.append( + { + **spec, + "multiplier_index": multiplier_index, + "multiplier": multiplier, + "multiplier_key": multiplier_key(multiplier), + "axis_multipliers": axis_multipliers, + "parameter_multipliers": param_multipliers, + "noise_model": {key: float(value) for key, value in noise.items()}, + "probability_totals": _probability_totals(noise), + "noise_model_sha256": validated.sha256(), + } + ) + return generated + + +def _render_config(base_cfg: Any, noise_model: Mapping[str, float], *, header: str) -> str: + cfg = OmegaConf.create(OmegaConf.to_container(base_cfg, resolve=True)) + cfg.data.noise_model = dict(noise_model) + return header + OmegaConf.to_yaml(cfg, resolve=True) + + +def _config_name(prefix: str, env_index: int, multiplier: float) -> str: + return f"{prefix}_e{int(env_index):02d}_{multiplier_key(multiplier)}" + + +def write_axismix_grid_configs( + *, + base_config: str | Path = DEFAULT_BASE_CONFIG, + output_dir: str | Path = DEFAULT_OUTPUT_DIR, + prefix: str = DEFAULT_PREFIX, + manifest: str | Path = DEFAULT_MANIFEST, + env_specs: Sequence[Mapping[str, Any]] = DEFAULT_ENV_SPECS, + grid_multipliers: Sequence[float] = GRID_MULTIPLIERS, +) -> tuple[list[Path], dict[str, Any]]: + base_path = rel(base_config) + if not base_path.exists(): + raise FileNotFoundError(base_path) + out_dir = rel(output_dir) + out_dir.mkdir(parents=True, exist_ok=True) + + base_cfg = OmegaConf.load(base_path) + base_noise = load_base_noise_model(base_path) + generated = generate_axismix_grid_noise_models( + base_noise, + env_specs, + grid_multipliers=grid_multipliers, + ) + + paths = [] + environments = [] + for item in generated: + config_name = _config_name(prefix, int(item["env_index"]), float(item["multiplier"])) + filename = f"{config_name}.yaml" + path = out_dir / filename + axis_json = json.dumps(item["axis_multipliers"], sort_keys=True) + header = ( + "# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.\n" + "# SPDX-License-Identifier: Apache-2.0\n" + "\n" + "# Auto-generated training-axis fixed multiplier grid OOD noise environment.\n" + f"# design: {DESIGN_LABEL}\n" + f"# base_config: {base_path.name}\n" + f"# env_key: {item['env_key']}\n" + f"# active_axes: {item['axis_signature']}\n" + f"# multiplier: {float(item['multiplier']):.6g}\n" + f"# axis_multipliers: {axis_json}\n" + f"# noise_model_sha256: {item['noise_model_sha256']}\n\n" + ) + path.write_text( + _render_config(base_cfg, item["noise_model"], header=header), + encoding="utf-8", + ) + paths.append(path) + environments.append( + { + "env_index": int(item["env_index"]), + "env_key": item["env_key"], + "multiplier_index": int(item["multiplier_index"]), + "multiplier_key": item["multiplier_key"], + "multiplier": float(item["multiplier"]), + "config_name": config_name_from_path(path), + "config_filename": filename, + "active_axes": list(item["active_axes"]), + "axis_signature": item["axis_signature"], + "axis_multipliers": item["axis_multipliers"], + "combination_size": int(item["combination_size"]), + "contains_z_bias": bool(item["contains_z_bias"]), + "contains_cnot_z_bias": bool(item["contains_cnot_z_bias"]), + "parameter_multipliers": item["parameter_multipliers"], + "probability_totals": item["probability_totals"], + "noise_model_sha256": item["noise_model_sha256"], + } + ) + + env_count = len({int(item["env_index"]) for item in generated}) + manifest_payload = { + "design": DESIGN_LABEL, + "base_config": str(base_config), + "prefix": prefix, + "axis_order": list(AXIS_ORDER), + "grid_multipliers": [float(value) for value in grid_multipliers], + "num_envs": env_count, + "num_configs": len(generated), + "axes": {name: list(keys) for name, keys in AXES.items()}, + "environments": environments, + } + manifest_path = rel(manifest) + manifest_path.parent.mkdir(parents=True, exist_ok=True) + manifest_path.write_text( + json.dumps(manifest_payload, indent=2, sort_keys=True), + encoding="utf-8", + ) + manifest_payload["manifest_path"] = str(manifest_path) + return paths, manifest_payload + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--base-config", default=DEFAULT_BASE_CONFIG) + parser.add_argument("--output-dir", default=DEFAULT_OUTPUT_DIR) + parser.add_argument("--prefix", default=DEFAULT_PREFIX) + parser.add_argument("--manifest", default=DEFAULT_MANIFEST) + parser.add_argument( + "--grid-multipliers", + default=",".join(str(value) for value in GRID_MULTIPLIERS), + help="Comma-separated multiplier grid.", + ) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + grid = [float(item.strip()) for item in args.grid_multipliers.split(",") if item.strip()] + paths, manifest = write_axismix_grid_configs( + base_config=args.base_config, + output_dir=args.output_dir, + prefix=args.prefix, + manifest=args.manifest, + grid_multipliers=grid, + ) + print(f"[write] {manifest['manifest_path']}") + print(f"[write] {len(paths)} configs") + + +if __name__ == "__main__": + main() diff --git a/code/scripts/paired_inference_compare.py b/code/scripts/paired_inference_compare.py new file mode 100644 index 0000000..4d8373b --- /dev/null +++ b/code/scripts/paired_inference_compare.py @@ -0,0 +1,942 @@ +#!/usr/bin/env python3 +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Paired inference comparison on one shared inference dataset. + +This script compares pure PyMatching with one or more predecoder models on the +same samples for each measurement basis. Samples are generated by Stim unless +``--stim-samples-dir`` points to external ``.dets`` artifacts. It is +intentionally separate from the Hydra workflow so the standard train/inference +entry points stay unchanged. +""" + +from __future__ import annotations + +import argparse +import csv +import json +import math +import os +import random +import sys +import time +from dataclasses import dataclass +from itertools import combinations +from pathlib import Path +from types import SimpleNamespace +from typing import Any + +import numpy as np +import pymatching +import torch +from omegaconf import OmegaConf +from torch.utils.data import DataLoader + +CODE_ROOT = Path(__file__).resolve().parents[1] +REPO_ROOT = CODE_ROOT.parent +if str(CODE_ROOT) not in sys.path: + sys.path.insert(0, str(CODE_ROOT)) + +from scripts.config_paths import config_path # noqa: E402 +from data.factory import DatapipeFactory # noqa: E402 +from evaluation.logical_error_rate import ( # noqa: E402 + PreDecoderMemoryEvalModule, + _build_stab_maps, +) +from training.utils import dict_to_device # noqa: E402 +from workflows.config_validator import ( # noqa: E402 + apply_public_defaults_and_model, + validate_public_config, +) +from model.checkpoint_loader import load_model_checkpoint # noqa: E402 + + +@dataclass(frozen=True) +class ModelSpec: + name: str + model_id: int + checkpoint: Path + + +@dataclass(frozen=True) +class ComparisonSpec: + candidate: str + baseline: str + + +@dataclass(frozen=True) +class FactorialContrastSpec: + name: str + cell_11: str + cell_10: str + cell_01: str + cell_00: str + + +@dataclass +class SyndromeDensityAccumulator: + """Stream shot-level syndrome-density moments without storing every sample.""" + + shots: int = 0 + syndrome_ones: int = 0 + syndrome_elements: int = 0 + shot_density_sum: float = 0.0 + shot_density_sum_squares: float = 0.0 + + def update(self, syndromes: np.ndarray) -> None: + values = np.asarray(syndromes, dtype=np.uint8) + if values.ndim == 1: + values = values.reshape(1, -1) + if values.ndim != 2 or values.shape[1] == 0: + raise ValueError( + "syndromes must be a non-empty-width 2D array, " + f"got shape={values.shape}" + ) + ones_per_shot = np.count_nonzero(values, axis=1).astype(np.float64) + densities = ones_per_shot / float(values.shape[1]) + self.shots += int(values.shape[0]) + self.syndrome_ones += int(ones_per_shot.sum()) + self.syndrome_elements += int(values.size) + self.shot_density_sum += float(densities.sum()) + self.shot_density_sum_squares += float(np.square(densities).sum()) + + def statistics(self, prefix: str) -> dict[str, float | int]: + if not prefix: + raise ValueError("density prefix must not be empty") + center = ( + float(self.syndrome_ones / self.syndrome_elements) + if self.syndrome_elements + else float("nan") + ) + if self.shots > 1: + numerator = self.shot_density_sum_squares - ( + self.shot_density_sum * self.shot_density_sum / self.shots + ) + variance = max(0.0, numerator / (self.shots - 1)) + standard_error = float(np.sqrt(variance / self.shots)) + else: + standard_error = 0.0 if self.shots == 1 else float("nan") + margin = 1.96 * standard_error + return { + f"{prefix}_density_shots": self.shots, + f"{prefix}_syndrome_ones": self.syndrome_ones, + f"{prefix}_syndrome_elements": self.syndrome_elements, + f"{prefix}_density_shot_sum": self.shot_density_sum, + f"{prefix}_density_shot_sum_squares": self.shot_density_sum_squares, + f"{prefix}_syndrome_density": center, + f"{prefix}_density_standard_error": standard_error, + f"{prefix}_density_ci95_low": max(0.0, center - margin), + f"{prefix}_density_ci95_high": min(1.0, center + margin), + } + + +def combine_density_statistics( + rows: list[dict[str, Any]], + prefix: str, +) -> dict[str, float | int]: + """Combine density sufficient statistics using detector-element weighting.""" + + accumulator = SyndromeDensityAccumulator() + for row in rows: + accumulator.shots += int(row.get(f"{prefix}_density_shots", 0)) + accumulator.syndrome_ones += int(row.get(f"{prefix}_syndrome_ones", 0)) + accumulator.syndrome_elements += int( + row.get(f"{prefix}_syndrome_elements", 0) + ) + accumulator.shot_density_sum += float( + row.get(f"{prefix}_density_shot_sum", 0.0) + ) + accumulator.shot_density_sum_squares += float( + row.get(f"{prefix}_density_shot_sum_squares", 0.0) + ) + return accumulator.statistics(prefix) + + +def density_reduction_statistics( + input_density: float, + residual_density: float, +) -> dict[str, float]: + input_value = float(input_density) + residual_value = float(residual_density) + delta = residual_value - input_value + if input_value > 0 and math.isfinite(input_value): + reduction_fraction = (input_value - residual_value) / input_value + else: + reduction_fraction = float("nan") + if residual_value > 0 and math.isfinite(residual_value): + reduction_factor = input_value / residual_value + elif input_value > 0 and residual_value == 0: + reduction_factor = float("inf") + else: + reduction_factor = float("nan") + return { + "density_delta": delta, + "density_reduction_fraction": reduction_fraction, + "density_reduction_factor": reduction_factor, + } + + +def model_density_statistics( + input_accumulator: SyndromeDensityAccumulator, + residual_accumulator: SyndromeDensityAccumulator, +) -> dict[str, float | int]: + input_stats = input_accumulator.statistics("input") + residual_stats = residual_accumulator.statistics("residual") + return { + **input_stats, + **residual_stats, + **density_reduction_statistics( + float(input_stats["input_syndrome_density"]), + float(residual_stats["residual_syndrome_density"]), + ), + } + + +def parse_model_spec(value: str) -> ModelSpec: + parts = value.split(":", 2) + if len(parts) != 3: + raise argparse.ArgumentTypeError( + "--model must be formatted as name:model_id:/path/to/checkpoint" + ) + name, model_id_raw, checkpoint_raw = parts + if not name: + raise argparse.ArgumentTypeError("model name must not be empty") + try: + model_id = int(model_id_raw) + except ValueError as exc: + raise argparse.ArgumentTypeError(f"invalid model_id: {model_id_raw}") from exc + checkpoint = Path(checkpoint_raw).expanduser() + if not checkpoint.is_absolute(): + checkpoint = REPO_ROOT / checkpoint + return ModelSpec(name=name, model_id=model_id, checkpoint=checkpoint) + + +def parse_comparison_spec(value: str) -> ComparisonSpec: + parts = value.split(":", 1) + if len(parts) != 2 or not all(part.strip() for part in parts): + raise argparse.ArgumentTypeError( + "--paired-comparison must be formatted as candidate:baseline" + ) + candidate, baseline = (part.strip() for part in parts) + if candidate == baseline: + raise argparse.ArgumentTypeError("candidate and baseline must be different methods") + return ComparisonSpec(candidate=candidate, baseline=baseline) + + +def parse_factorial_contrast_spec(value: str) -> FactorialContrastSpec: + parts = [part.strip() for part in value.split(":")] + if len(parts) != 5 or not all(parts): + raise argparse.ArgumentTypeError( + "--factorial-contrast must be formatted as " + "name:cell_11:cell_10:cell_01:cell_00" + ) + name, cell_11, cell_10, cell_01, cell_00 = parts + if len({cell_11, cell_10, cell_01, cell_00}) != 4: + raise argparse.ArgumentTypeError("factorial contrast cells must be four distinct methods") + return FactorialContrastSpec(name, cell_11, cell_10, cell_01, cell_00) + + +def factorial_contrast_statistics( + cell_11_errors: np.ndarray, + cell_10_errors: np.ndarray, + cell_01_errors: np.ndarray, + cell_00_errors: np.ndarray, +) -> dict[str, float | int]: + masks = [ + np.asarray(errors, dtype=np.bool_).reshape(-1) + for errors in (cell_11_errors, cell_10_errors, cell_01_errors, cell_00_errors) + ] + shapes = {mask.shape for mask in masks} + if len(shapes) != 1: + raise ValueError(f"factorial contrast masks must have one shape: {sorted(shapes)}") + samples = int(masks[0].size) + if samples == 0: + raise ValueError("factorial contrast masks must not be empty") + contrast = ( + masks[0].astype(np.int8) + - masks[1].astype(np.int8) + - masks[2].astype(np.int8) + + masks[3].astype(np.int8) + ) + interaction = float(contrast.mean()) + standard_error = ( + float(contrast.std(ddof=1) / np.sqrt(samples)) if samples > 1 else 0.0 + ) + margin = 1.96 * standard_error + result: dict[str, float | int] = { + "samples": samples, + "interaction_ler": interaction, + "standard_error": standard_error, + "ci95_low": max(-2.0, interaction - margin), + "ci95_high": min(2.0, interaction + margin), + } + result.update( + { + f"contrast_count_{value:+d}": int(np.count_nonzero(contrast == value)) + for value in range(-2, 3) + } + ) + return result + + +def paired_error_statistics( + candidate_errors: np.ndarray, + baseline_errors: np.ndarray, +) -> dict[str, float | int]: + candidate = np.asarray(candidate_errors, dtype=np.bool_).reshape(-1) + baseline = np.asarray(baseline_errors, dtype=np.bool_).reshape(-1) + if candidate.shape != baseline.shape: + raise ValueError( + f"paired error masks must have the same shape: {candidate.shape} != {baseline.shape}" + ) + samples = int(candidate.size) + if samples == 0: + raise ValueError("paired error masks must not be empty") + + candidate_only = int(np.count_nonzero(candidate & ~baseline)) + baseline_only = int(np.count_nonzero(~candidate & baseline)) + both = int(np.count_nonzero(candidate & baseline)) + neither = samples - candidate_only - baseline_only - both + differences = candidate.astype(np.int8) - baseline.astype(np.int8) + delta = float(differences.mean()) + standard_error = ( + float(differences.std(ddof=1) / np.sqrt(samples)) if samples > 1 else 0.0 + ) + margin = 1.96 * standard_error + return { + "samples": samples, + "candidate_only_errors": candidate_only, + "baseline_only_errors": baseline_only, + "both_errors": both, + "neither_errors": neither, + "delta_ler": delta, + "standard_error": standard_error, + "ci95_low": max(-1.0, delta - margin), + "ci95_high": min(1.0, delta + margin), + } + + +def paired_error_comparison( + method_a: str, + errors_a: np.ndarray, + method_b: str, + errors_b: np.ndarray, + *, + basis: str, +) -> dict[str, Any]: + """Summarize two shot-aligned logical-error masks.""" + stats = paired_error_statistics(errors_a, errors_b) + return { + "basis": basis, + "method_a": method_a, + "method_b": method_b, + "samples": stats["samples"], + "both_error": stats["both_errors"], + "a_only_error": stats["candidate_only_errors"], + "b_only_error": stats["baseline_only_errors"], + "neither_error": stats["neither_errors"], + "ler_delta_a_minus_b": stats["delta_ler"], + "paired_standard_error": stats["standard_error"], + "ler_delta_ci95_normal": [stats["ci95_low"], stats["ci95_high"]], + } + + +def parse_args(argv: list[str] | None = None) -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Compare PyMatching and replaceable predecoder models on identical samples." + ) + config_group = parser.add_mutually_exclusive_group() + config_group.add_argument( + "--config-name", default="examples/qadapt/config_qadapt_t0_base" + ) + config_group.add_argument( + "--config-file", type=Path, help="Explicit YAML path, including generated OOD configs." + ) + parser.add_argument("--distance", type=int, default=9) + parser.add_argument("--n-rounds", type=int, default=9) + parser.add_argument("--num-samples", type=int, default=262144) + parser.add_argument("--latency-num-samples", type=int, default=10000) + parser.add_argument("--batch-size", type=int, default=2048) + parser.add_argument("--num-workers", type=int, default=0) + parser.add_argument("--seed", type=int, default=12345) + parser.add_argument("--device", default=None) + parser.add_argument( + "--basis", + choices=("both", "X", "Z"), + default="both", + help="Measurement basis to evaluate.", + ) + parser.add_argument( + "--stim-samples-dir", + default=None, + help=( + "Optional directory containing samples_X.dets/metadata_X.json and/or " + "samples_Z.dets/metadata_Z.json. When omitted, Stim generates samples." + ), + ) + parser.add_argument( + "--model", + action="append", + type=parse_model_spec, + required=True, + help=( + "Repeatable model spec: name:model_id:/path/to/checkpoint " + "(.pt or .safetensors)." + ), + ) + parser.add_argument( + "--paired-comparison", + action="append", + type=parse_comparison_spec, + default=[], + help="Repeatable paired comparison: candidate:baseline.", + ) + parser.add_argument( + "--factorial-contrast", + action="append", + type=parse_factorial_contrast_spec, + default=[], + help="Repeatable contrast: name:cell_11:cell_10:cell_01:cell_00.", + ) + parser.add_argument( + "--output", + default="outputs/examples/released_models/paired_inference.json", + help="JSON output path. A CSV summary is written next to it.", + ) + parser.add_argument( + "--residual-output-dir", + default=None, + help=( + "Optional directory for full residual detector tensors. One uint8 " + "PyTorch tensor is written per basis and model." + ), + ) + return parser.parse_args(argv) + + +def set_all_seeds(seed: int) -> None: + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + if torch.cuda.is_available(): + torch.cuda.manual_seed_all(seed) + + +def resolve_stim_samples_dir(args: argparse.Namespace) -> Path | None: + value = getattr(args, "stim_samples_dir", None) or os.environ.get( + "PREDECODER_STIM_SAMPLES_DIR" + ) + if not value: + return None + path = Path(value).expanduser() + return path if path.is_absolute() else REPO_ROOT / path + + +def build_cfg(args: argparse.Namespace, model: ModelSpec, basis: str) -> Any: + explicit_path = getattr(args, "config_file", None) + cfg_path = ( + Path(explicit_path).expanduser() + if explicit_path is not None + else config_path(args.config_name) + ) + cfg = OmegaConf.load(cfg_path) + cfg.model_id = model.model_id + cfg.distance = args.distance + cfg.n_rounds = args.n_rounds + cfg.workflow.task = "inference" + + spec = validate_public_config(cfg) + cfg = apply_public_defaults_and_model(cfg, spec) + cfg.model_checkpoint_file = str(model.checkpoint) + cfg.test.meas_basis_test = basis + cfg.test.num_samples = int(args.num_samples) + cfg.test.latency_num_samples = int(args.latency_num_samples) + cfg.test.batch_size = int(args.batch_size) + cfg.test.dataloader_num_workers = int(args.num_workers) + stim_samples_dir = resolve_stim_samples_dir(args) + if stim_samples_dir: + cfg.test.stim_samples_dir = str(stim_samples_dir) + return cfg + + +def make_dataset(cfg: Any, seed: int): + py_state = random.getstate() + np_state = np.random.get_state() + torch_state = torch.get_rng_state() + cuda_state = torch.cuda.get_rng_state_all() if torch.cuda.is_available() else None + try: + set_all_seeds(seed) + return DatapipeFactory.create_datapipe_inference(cfg) + finally: + random.setstate(py_state) + np.random.set_state(np_state) + torch.set_rng_state(torch_state) + if cuda_state is not None: + torch.cuda.set_rng_state_all(cuda_state) + + +def time_single_shot(matcher: pymatching.Matching, syndromes: np.ndarray, n_rounds: int) -> float: + n_rounds = max(int(n_rounds), 1) + if syndromes.size == 0: + return float("nan") + if torch.cuda.is_available(): + torch.cuda.synchronize() + warmup_n = min(50, len(syndromes)) + for i in range(warmup_n): + matcher.decode(np.asarray(syndromes[i], dtype=np.uint8)) + + times = [] + for row in syndromes: + start = time.perf_counter() + matcher.decode(np.asarray(row, dtype=np.uint8)) + times.append(time.perf_counter() - start) + return float(np.mean(times) / n_rounds * 1e6) + + +def build_matcher(dataset) -> tuple[pymatching.Matching, int]: + circuit = dataset.circ.stim_circuit + det_model = circuit.detector_error_model(decompose_errors=True, approximate_disjoint_errors=True) + return pymatching.Matching.from_detector_error_model(det_model), int(circuit.num_observables) + + +def evaluate_pymatching( + matcher: pymatching.Matching, + dets_and_obs: np.ndarray, + num_obs: int, + latency_samples: int, + n_rounds: int, +) -> tuple[dict[str, float | int], np.ndarray]: + dets = np.ascontiguousarray(dets_and_obs[:, :-num_obs], dtype=np.uint8) + obs = np.ascontiguousarray(dets_and_obs[:, -num_obs:], dtype=np.uint8) + pred = matcher.decode_batch(dets).reshape(obs.shape) + error_mask = np.asarray(pred != obs, dtype=np.bool_).reshape(obs.shape[0], -1).any(axis=1) + errors = int(error_mask.sum()) + total = int(obs.shape[0]) + latency_rows = dets[: min(latency_samples, len(dets))] + input_density = SyndromeDensityAccumulator() + input_density.update(dets) + return { + "logical_errors": errors, + "samples": total, + "ler": float(errors / total) if total else float("nan"), + "latency_us_per_round": time_single_shot(matcher, latency_rows, n_rounds), + **input_density.statistics("input"), + }, error_mask + + +def evaluate_model( + model: torch.nn.Module, + cfg: Any, + dataset, + matcher: pymatching.Matching, + num_obs: int, + device: torch.device, + latency_samples: int, + n_rounds: int, + residual_tensor_path: Path | None = None, +) -> tuple[dict[str, Any], np.ndarray]: + maps = _build_stab_maps(int(cfg.distance), getattr(cfg, "rotation", "XV")) + module = PreDecoderMemoryEvalModule(model, cfg, maps, device).to(device) + module.eval() + loader = DataLoader( + dataset, + batch_size=int(cfg.test.batch_size), + shuffle=False, + num_workers=int(cfg.test.dataloader_num_workers), + pin_memory=(device.type == "cuda"), + ) + + logical_errors = 0 + total = 0 + residual_chunks: list[np.ndarray] = [] + saved_residual_chunks: list[np.ndarray] = [] + error_chunks: list[np.ndarray] = [] + residual_count = 0 + input_density = SyndromeDensityAccumulator() + residual_density = SyndromeDensityAccumulator() + + with torch.no_grad(): + for batch in loader: + batch = dict_to_device(batch, device) + dets_and_obs = batch["dets_and_obs"] + dets_only = dets_and_obs[:, :-num_obs] + gt_obs = dets_and_obs[:, -num_obs:].to(torch.int64).cpu() + + output = module(dets_only) + pre_l = output[:, 0].to(torch.int64).cpu() + residual = output[:, 1:].to(torch.uint8).cpu().numpy() + input_density.update(dets_only.to(torch.uint8).cpu().numpy()) + residual_density.update(residual) + if residual_tensor_path is not None: + saved_residual_chunks.append( + np.ascontiguousarray(residual, dtype=np.uint8) + ) + pred_obs = torch.from_numpy(matcher.decode_batch(residual)).reshape(gt_obs.shape) + final_l = (pre_l.reshape(gt_obs.shape) + pred_obs).remainder(2) + + error_mask = (final_l != gt_obs).reshape(gt_obs.shape[0], -1).any(dim=1) + logical_errors += int(error_mask.sum().item()) + total += int(gt_obs.shape[0]) + error_chunks.append(error_mask.numpy()) + + if residual_count < latency_samples: + take = min(latency_samples - residual_count, residual.shape[0]) + residual_chunks.append(np.ascontiguousarray(residual[:take], dtype=np.uint8)) + residual_count += take + + residual_rows = ( + np.concatenate(residual_chunks, axis=0) if residual_chunks else np.empty((0, 0), dtype=np.uint8) + ) + all_errors = np.concatenate(error_chunks) if error_chunks else np.empty(0, dtype=np.bool_) + result: dict[str, Any] = { + "logical_errors": logical_errors, + "samples": total, + "ler": float(logical_errors / total) if total else float("nan"), + "latency_us_per_round": time_single_shot(matcher, residual_rows, n_rounds), + **model_density_statistics(input_density, residual_density), + } + if residual_tensor_path is not None: + residual_tensor_path.parent.mkdir(parents=True, exist_ok=True) + saved_residual = ( + np.concatenate(saved_residual_chunks, axis=0) + if saved_residual_chunks + else np.empty((0, 0), dtype=np.uint8) + ) + torch.save(torch.from_numpy(saved_residual), residual_tensor_path) + result.update( + residual_tensor_path=str(residual_tensor_path), + residual_tensor_rows=int(saved_residual.shape[0]), + residual_tensor_detectors=int(saved_residual.shape[1]), + residual_tensor_dtype="torch.uint8", + ) + return result, all_errors + + +def build_paired_comparison_rows( + error_masks_by_basis: dict[str, dict[str, np.ndarray]], + comparisons: list[ComparisonSpec], +) -> list[dict[str, Any]]: + rows: list[dict[str, Any]] = [] + basis_order = [basis for basis in ("X", "Z") if basis in error_masks_by_basis] + for comparison in comparisons: + candidate_chunks = [] + baseline_chunks = [] + for basis in basis_order: + masks = error_masks_by_basis[basis] + missing = { + method + for method in (comparison.candidate, comparison.baseline) + if method not in masks + } + if missing: + raise KeyError(f"paired comparison methods missing for {basis}: {sorted(missing)}") + candidate = masks[comparison.candidate] + baseline = masks[comparison.baseline] + rows.append( + { + "basis": basis, + "candidate": comparison.candidate, + "baseline": comparison.baseline, + **paired_error_statistics(candidate, baseline), + } + ) + candidate_chunks.append(candidate) + baseline_chunks.append(baseline) + if len(basis_order) > 1: + rows.append( + { + "basis": "both", + "candidate": comparison.candidate, + "baseline": comparison.baseline, + **paired_error_statistics( + np.concatenate(candidate_chunks), + np.concatenate(baseline_chunks), + ), + } + ) + return rows + + +def build_factorial_contrast_rows( + error_masks_by_basis: dict[str, dict[str, np.ndarray]], + contrasts: list[FactorialContrastSpec], +) -> list[dict[str, Any]]: + rows: list[dict[str, Any]] = [] + basis_order = [basis for basis in ("X", "Z") if basis in error_masks_by_basis] + for contrast in contrasts: + chunks = {field: [] for field in ("cell_11", "cell_10", "cell_01", "cell_00")} + for basis in basis_order: + masks = error_masks_by_basis[basis] + methods = { + field: getattr(contrast, field) + for field in ("cell_11", "cell_10", "cell_01", "cell_00") + } + missing = set(methods.values()) - set(masks) + if missing: + raise KeyError(f"factorial contrast methods missing for {basis}: {sorted(missing)}") + stats = factorial_contrast_statistics(*(masks[methods[field]] for field in chunks)) + rows.append( + { + "basis": basis, + "name": contrast.name, + **methods, + **stats, + } + ) + for field, method in methods.items(): + chunks[field].append(masks[method]) + if len(basis_order) > 1: + rows.append( + { + "basis": "both", + "name": contrast.name, + "cell_11": contrast.cell_11, + "cell_10": contrast.cell_10, + "cell_01": contrast.cell_01, + "cell_00": contrast.cell_00, + **factorial_contrast_statistics( + *(np.concatenate(chunks[field]) for field in chunks) + ), + } + ) + return rows + + +def mean_metric(rows: list[dict[str, Any]], name: str) -> float: + values = [float(row[name]) for row in rows if row.get(name) is not None] + return float(np.mean(values)) if values else float("nan") + + +def main() -> None: + args = parse_args() + stim_samples_dir = resolve_stim_samples_dir(args) + if stim_samples_dir is not None: + # DatapipeFactory historically gives the environment variable priority. + # Synchronize it so an explicit CLI path cannot be silently shadowed. + os.environ["PREDECODER_STIM_SAMPLES_DIR"] = str(stim_samples_dir) + output_path = Path(args.output) + if not output_path.is_absolute(): + output_path = REPO_ROOT / output_path + output_path.parent.mkdir(parents=True, exist_ok=True) + residual_output_dir = ( + Path(args.residual_output_dir) if args.residual_output_dir else None + ) + if residual_output_dir is not None and not residual_output_dir.is_absolute(): + residual_output_dir = REPO_ROOT / residual_output_dir + + for spec in args.model: + if not spec.checkpoint.exists(): + raise FileNotFoundError(f"Checkpoint not found for {spec.name}: {spec.checkpoint}") + available_methods = {"pymatching", *(spec.name for spec in args.model)} + for comparison in args.paired_comparison: + missing = {comparison.candidate, comparison.baseline} - available_methods + if missing: + raise ValueError(f"Unknown paired comparison methods: {sorted(missing)}") + for contrast in args.factorial_contrast: + missing = { + contrast.cell_11, + contrast.cell_10, + contrast.cell_01, + contrast.cell_00, + } - available_methods + if missing: + raise ValueError(f"Unknown factorial contrast methods: {sorted(missing)}") + + device = torch.device(args.device or ("cuda:0" if torch.cuda.is_available() else "cpu")) + dist = SimpleNamespace(rank=0, world_size=1, device=device) + bases = ["X", "Z"] if args.basis == "both" else [args.basis] + + model_cfgs = {spec.name: build_cfg(args, spec, basis=bases[0]) for spec in args.model} + models = {} + for spec in args.model: + print(f"[load] {spec.name}: model_id={spec.model_id}, checkpoint={spec.checkpoint}") + model = load_model_checkpoint( + model_cfgs[spec.name], + checkpoint=spec.checkpoint, + model_id=spec.model_id, + distributed=dist, + ) + model.eval() + models[spec.name] = model + + rows: list[dict[str, Any]] = [] + error_masks_by_basis: dict[str, dict[str, np.ndarray]] = {} + paired_comparisons: list[dict[str, Any]] = [] + sample_metadata: dict[str, Any] = {} + for basis_index, basis in enumerate(bases): + dataset_cfg = build_cfg(args, args.model[0], basis=basis) + dataset_seed = int(args.seed) + basis_index + print(f"[data] basis={basis}, seed={dataset_seed}, samples={args.num_samples}") + dataset = make_dataset(dataset_cfg, dataset_seed) + if hasattr(dataset, "metadata"): + sample_metadata[basis] = dict(dataset.metadata) + matcher, num_obs = build_matcher(dataset) + dets_and_obs = np.asarray(dataset.dets_and_obs, dtype=np.uint8) + + baseline, baseline_errors = evaluate_pymatching( + matcher, + dets_and_obs, + num_obs, + int(args.latency_num_samples), + int(args.n_rounds), + ) + error_masks_by_basis[basis] = {"pymatching": baseline_errors} + baseline_row = { + "basis": basis, + "method": "pymatching", + "model_id": "", + "checkpoint": "", + **baseline, + "speedup_vs_pymatching": 1.0, + } + rows.append(baseline_row) + basis_errors = {"pymatching": baseline_errors} + print( + f"[result] {basis} pymatching ler={baseline['ler']:.6f}, " + f"latency={baseline['latency_us_per_round']:.3f} us/round" + ) + + for spec in args.model: + cfg = build_cfg(args, spec, basis=basis) + residual_tensor_path = ( + residual_output_dir / f"{basis}_{spec.name}_residual_detectors.pt" + if residual_output_dir is not None + else None + ) + result, model_errors = evaluate_model( + models[spec.name], + cfg, + dataset, + matcher, + num_obs, + device, + int(args.latency_num_samples), + int(args.n_rounds), + residual_tensor_path, + ) + error_masks_by_basis[basis][spec.name] = model_errors + speedup = float(baseline["latency_us_per_round"]) / float(result["latency_us_per_round"]) + row = { + "basis": basis, + "method": spec.name, + "model_id": spec.model_id, + "checkpoint": str(spec.checkpoint), + **result, + "speedup_vs_pymatching": speedup, + } + rows.append(row) + basis_errors[spec.name] = model_errors + print( + f"[result] {basis} {spec.name} ler={result['ler']:.6f}, " + f"latency={result['latency_us_per_round']:.3f} us/round, speedup={speedup:.3f}x" + ) + for method_a, method_b in combinations(basis_errors, 2): + paired_comparisons.append( + paired_error_comparison( + method_a, + basis_errors[method_a], + method_b, + basis_errors[method_b], + basis=basis, + ) + ) + + if args.paired_comparison: + paired_comparisons = build_paired_comparison_rows( + error_masks_by_basis, + args.paired_comparison, + ) + factorial_contrasts = build_factorial_contrast_rows( + error_masks_by_basis, + args.factorial_contrast, + ) + methods = sorted({row["method"] for row in rows}) + summary = [] + for method in methods: + method_rows = [row for row in rows if row["method"] == method] + summary_row: dict[str, Any] = { + "method": method, + "ler_avg": mean_metric(method_rows, "ler"), + "latency_us_per_round_avg": mean_metric(method_rows, "latency_us_per_round"), + "speedup_vs_pymatching_avg": mean_metric(method_rows, "speedup_vs_pymatching"), + **combine_density_statistics(method_rows, "input"), + } + if any(row.get("residual_syndrome_elements") for row in method_rows): + residual_stats = combine_density_statistics(method_rows, "residual") + summary_row.update(residual_stats) + summary_row.update( + density_reduction_statistics( + float(summary_row["input_syndrome_density"]), + float(residual_stats["residual_syndrome_density"]), + ) + ) + summary.append(summary_row) + + payload = { + "config_name": args.config_name, + "distance": args.distance, + "n_rounds": args.n_rounds, + "num_samples": args.num_samples, + "latency_num_samples": args.latency_num_samples, + "seed": args.seed, + "device": str(device), + "sample_source": "stim_files" if stim_samples_dir else "generated", + "stim_samples_dir": str(stim_samples_dir) if stim_samples_dir else None, + "sample_metadata": sample_metadata, + "rows": rows, + "summary": summary, + "paired_comparisons": paired_comparisons, + "factorial_contrasts": factorial_contrasts, + } + output_path.write_text(json.dumps(payload, indent=2), encoding="utf-8") + + csv_path = output_path.with_suffix(".csv") + fieldnames = [ + "basis", + "method", + "model_id", + "logical_errors", + "samples", + "ler", + "latency_us_per_round", + "speedup_vs_pymatching", + "input_density_shots", + "input_syndrome_ones", + "input_syndrome_elements", + "input_density_shot_sum", + "input_density_shot_sum_squares", + "input_syndrome_density", + "input_density_standard_error", + "input_density_ci95_low", + "input_density_ci95_high", + "residual_density_shots", + "residual_syndrome_ones", + "residual_syndrome_elements", + "residual_density_shot_sum", + "residual_density_shot_sum_squares", + "residual_syndrome_density", + "residual_density_standard_error", + "residual_density_ci95_low", + "residual_density_ci95_high", + "density_delta", + "density_reduction_fraction", + "density_reduction_factor", + "residual_tensor_path", + "residual_tensor_rows", + "residual_tensor_detectors", + "residual_tensor_dtype", + "checkpoint", + ] + with csv_path.open("w", newline="", encoding="utf-8") as f: + writer = csv.DictWriter(f, fieldnames=fieldnames) + writer.writeheader() + for row in rows: + writer.writerow({field: row.get(field, "") for field in fieldnames}) + + print(f"[write] {output_path}") + print(f"[write] {csv_path}") + + +if __name__ == "__main__": + main() diff --git a/code/scripts/providers/__init__.py b/code/scripts/providers/__init__.py new file mode 100644 index 0000000..f89f0e8 --- /dev/null +++ b/code/scripts/providers/__init__.py @@ -0,0 +1,4 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""External benchmark and circuit-data command modules.""" diff --git a/code/scripts/providers/google_qec_decoder_benchmark.py b/code/scripts/providers/google_qec_decoder_benchmark.py new file mode 100644 index 0000000..2a48c01 --- /dev/null +++ b/code/scripts/providers/google_qec_decoder_benchmark.py @@ -0,0 +1,1390 @@ +#!/usr/bin/env python3 +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Benchmark PyMatching and released pre-decoders on Google Willow QEC data.""" + +from __future__ import annotations + +import argparse +import csv +import json +import math +import sys +import time +from dataclasses import asdict, dataclass +from datetime import datetime, timezone +from pathlib import Path +from types import SimpleNamespace +from typing import Any, Iterable, Mapping, Sequence + +import numpy as np +import pymatching +import stim +import torch +from omegaconf import OmegaConf + +CODE_ROOT = Path(__file__).resolve().parents[2] +REPO_ROOT = CODE_ROOT.parent +if str(CODE_ROOT) not in sys.path: + sys.path.insert(0, str(CODE_ROOT)) + +from evaluation.logical_error_rate import ( # noqa: E402 + PreDecoderMemoryEvalModule, + _build_stab_maps, +) +from qec.surface_code.memory_circuit import SurfaceCode # noqa: E402 +from scripts.config_paths import config_path # noqa: E402 +from scripts.paired_inference_compare import ( # noqa: E402 + SyndromeDensityAccumulator, + model_density_statistics, +) +from workflows.config_validator import ( # noqa: E402 + apply_public_defaults_and_model, + validate_public_config, +) +from model.checkpoint_loader import load_model_checkpoint # noqa: E402 + + +DEFAULT_BENCHMARK_ROOT = ( + REPO_ROOT / "benchmarks/google_qec/google_105Q_surface_code_d3_d5_d7" +) + + +@dataclass(frozen=True) +class BenchmarkModel: + name: str + model_id: int + checkpoint: Path + + +# The public wrapper injects explicitly named checkpoint paths before parsing. +# Keep the backend free of internal training-output defaults. +DEFAULT_MODELS: dict[str, BenchmarkModel] = {} + + +def maybe_compile_model( + model: torch.nn.Module, + *, + enabled: bool, + mode: str = "default", +) -> torch.nn.Module: + """Optionally compile one cached model with dynamic detector dimensions.""" + + return torch.compile(model, mode=mode, dynamic=True) if enabled else model + + +@dataclass(frozen=True) +class GoogleQECCase: + path: Path + patch: str + distance: int + basis: str + rounds: int + shots: int + + +REQUIRED_CASE_FILES = ( + "circuit_ideal.stim", + "circuit_noisy_si1000.stim", + "detection_events.b8", + "obs_flips_actual.b8", +) + + +def discover_cases( + root: Path, + *, + distances: set[int] | None = None, + rounds: set[int] | None = None, + bases: set[str] | None = None, + patches: set[str] | None = None, +) -> list[GoogleQECCase]: + """Discover complete Google benchmark cases selected by metadata.""" + + root = Path(root) + cases = [] + for metadata_path in root.glob("d*_at_q*/[XZ]/r*/metadata.json"): + metadata = json.loads(metadata_path.read_text()) + case_dir = metadata_path.parent + patch = case_dir.parents[1].name + distance = int(metadata["distance"]) + basis = str(metadata["basis"]).upper() + n_rounds = int(metadata["rounds"]) + if distances is not None and distance not in distances: + continue + if rounds is not None and n_rounds not in rounds: + continue + if bases is not None and basis not in bases: + continue + if patches is not None and patch not in patches: + continue + missing = [name for name in REQUIRED_CASE_FILES if not (case_dir / name).is_file()] + if missing: + raise FileNotFoundError(f"Incomplete Google QEC case {case_dir}: missing {missing}") + cases.append( + GoogleQECCase( + path=case_dir, + patch=patch, + distance=distance, + basis=basis, + rounds=n_rounds, + shots=int(metadata["shots"]), + ) + ) + return sorted(cases, key=lambda case: (case.distance, case.patch, case.basis, case.rounds)) + + +def _google_to_xv_coordinate( + coordinate: Sequence[float], + *, + min_difference: int, + min_sum: int, +) -> tuple[int, int]: + if len(coordinate) < 2: + raise ValueError(f"Google coordinate must contain x and y, got {coordinate!r}") + x = float(coordinate[0]) + y = float(coordinate[1]) + if not x.is_integer() or not y.is_integer(): + raise ValueError(f"Google coordinate must be integral, got {coordinate!r}") + x_int = int(x) + y_int = int(y) + return ( + x_int - y_int - int(min_difference) + 1, + x_int + y_int - int(min_sum) + 1, + ) + + +def build_detector_permutation( + circuit: stim.Circuit, + metadata: Mapping[str, Any], +) -> np.ndarray: + """Return indices that map Google detector columns to the model's XV order. + + Google emits each bulk round in physical measurement-qubit order. The + predecoder consumes initial-boundary, X-block, Z-block, ..., final-boundary + order, with stabilizers indexed by the repository's XV patch convention. + """ + + distance = int(metadata["distance"]) + rounds = int(metadata["rounds"]) + basis = str(metadata["basis"]).upper() + if basis not in {"X", "Z"}: + raise ValueError(f"basis must be X or Z, got {basis!r}") + if distance < 3 or distance % 2 == 0: + raise ValueError(f"distance must be an odd integer >= 3, got {distance}") + if rounds < 1: + raise ValueError(f"rounds must be positive, got {rounds}") + + half = (distance * distance - 1) // 2 + expected_detectors = 2 * rounds * half + if int(circuit.num_detectors) != expected_detectors: + raise ValueError( + "detector count mismatch: " + f"circuit has {circuit.num_detectors}, expected {expected_detectors} " + f"for d={distance}, rounds={rounds}" + ) + + data_coordinates = [tuple(item) for item in metadata["data_qubit_coords"]] + if len(data_coordinates) != distance * distance: + raise ValueError( + f"data coordinate count mismatch: {len(data_coordinates)} != {distance * distance}" + ) + min_difference = min(int(x) - int(y) for x, y in data_coordinates) + min_sum = min(int(x) + int(y) for x, y in data_coordinates) + transformed_data = { + _google_to_xv_coordinate( + coordinate, + min_difference=min_difference, + min_sum=min_sum, + ) + for coordinate in data_coordinates + } + odd_coordinates = range(1, 2 * distance, 2) + expected_data = {(x, y) for x in odd_coordinates for y in odd_coordinates} + if transformed_data != expected_data: + raise ValueError("Google data-qubit coordinates do not form the expected rotated patch") + + code = SurfaceCode(distance, first_bulk_syndrome_type="X", rotated_type="V") + x_indices = { + tuple(map(int, code.xcheck_qubits_dict[int(qubit)]["coord"])): index + for index, qubit in enumerate(code.xcheck_qubits) + } + z_indices = { + tuple(map(int, code.zcheck_qubits_dict[int(qubit)]["coord"])): index + for index, qubit in enumerate(code.zcheck_qubits) + } + detector_coordinates = circuit.get_detector_coordinates() + if len(detector_coordinates) != expected_detectors: + raise ValueError( + "detector coordinate count mismatch: " + f"{len(detector_coordinates)} != {expected_detectors}" + ) + + canonical_to_source = np.full(expected_detectors, -1, dtype=np.int64) + boundary_start = expected_detectors - half + for source_index in range(expected_detectors): + raw_coordinate = detector_coordinates[source_index] + if len(raw_coordinate) < 3: + raise ValueError(f"detector {source_index} has no spatial/time coordinate") + # Initial and bulk detectors end in their stabilizer coordinate. Google + # final-boundary detectors list data coordinates first and the previous + # ancilla/stabilizer coordinate last, so the last coordinate triple is + # the uniform choice for every phase. + model_coordinate = _google_to_xv_coordinate( + raw_coordinate[-3:-1], + min_difference=min_difference, + min_sum=min_sum, + ) + if model_coordinate in x_indices: + stabilizer_type = "X" + stabilizer_index = x_indices[model_coordinate] + elif model_coordinate in z_indices: + stabilizer_type = "Z" + stabilizer_index = z_indices[model_coordinate] + else: + raise ValueError( + f"detector {source_index} coordinate {raw_coordinate!r} maps to " + f"unknown XV stabilizer {model_coordinate}" + ) + + if source_index < half: + if stabilizer_type != basis: + raise ValueError( + f"initial detector {source_index} is {stabilizer_type}, expected {basis}" + ) + canonical_index = stabilizer_index + elif source_index >= boundary_start: + if stabilizer_type != basis: + raise ValueError( + f"boundary detector {source_index} is {stabilizer_type}, expected {basis}" + ) + canonical_index = boundary_start + stabilizer_index + else: + bulk_offset = source_index - half + bulk_round = bulk_offset // (2 * half) + type_offset = 0 if stabilizer_type == "X" else half + canonical_index = half + bulk_round * 2 * half + type_offset + stabilizer_index + + if canonical_to_source[canonical_index] != -1: + raise ValueError( + f"duplicate detector mapping for canonical index {canonical_index}" + ) + canonical_to_source[canonical_index] = source_index + + if np.any(canonical_to_source < 0): + missing = np.flatnonzero(canonical_to_source < 0).tolist() + raise ValueError(f"incomplete detector mapping; missing canonical indices {missing}") + return canonical_to_source + + +def google_to_canonical(data: np.ndarray, canonical_to_source: np.ndarray) -> np.ndarray: + rows = np.asarray(data) + permutation = np.asarray(canonical_to_source, dtype=np.int64) + if rows.ndim != 2 or rows.shape[1] != permutation.size: + raise ValueError( + f"Google detector shape {rows.shape} is incompatible with permutation " + f"width {permutation.size}" + ) + return np.ascontiguousarray(rows[:, permutation]) + + +def canonical_to_google(data: np.ndarray, canonical_to_source: np.ndarray) -> np.ndarray: + rows = np.asarray(data) + permutation = np.asarray(canonical_to_source, dtype=np.int64) + if rows.ndim != 2 or rows.shape[1] != permutation.size: + raise ValueError( + f"canonical detector shape {rows.shape} is incompatible with permutation " + f"width {permutation.size}" + ) + restored = np.empty_like(rows) + restored[:, permutation] = rows + return np.ascontiguousarray(restored) + + +def verify_bulk_data_fault_equivalence( + circuit: stim.Circuit, + metadata: Mapping[str, Any], +) -> dict[str, Any]: + """Compare all inter-cycle physical X/Y/Z faults with CSS signatures.""" + + distance = int(metadata["distance"]) + basis = str(metadata["basis"]).upper() + if basis not in {"X", "Z"}: + raise ValueError(f"basis must be X or Z, got {basis!r}") + + data_coordinates = [tuple(map(int, item)) for item in metadata["data_qubit_coords"]] + if len(data_coordinates) != distance * distance: + raise ValueError( + f"data coordinate count mismatch: {len(data_coordinates)} != {distance * distance}" + ) + qubit_coordinates = { + int(qubit): tuple(map(int, coordinate)) + for qubit, coordinate in circuit.get_final_qubit_coordinates().items() + } + coordinate_to_qubit = {coordinate: qubit for qubit, coordinate in qubit_coordinates.items()} + missing_qubits = [coordinate for coordinate in data_coordinates if coordinate not in coordinate_to_qubit] + if missing_qubits: + raise ValueError(f"data coordinates missing from circuit: {missing_qubits}") + data_qubits = {coordinate_to_qubit[coordinate] for coordinate in data_coordinates} + + cycle_boundaries = [] + for instruction_index in range(len(circuit)): + instruction = circuit[instruction_index] + if instruction.name != "Y": + continue + targets = { + int(target.value) + for target in instruction.targets_copy() + if target.is_qubit_target + } + if targets == data_qubits: + cycle_boundaries.append(instruction_index) + expected_boundaries = int(metadata["rounds"]) - 1 + if len(cycle_boundaries) != expected_boundaries: + raise ValueError( + "inter-cycle boundary count mismatch: " + f"{len(cycle_boundaries)} != {expected_boundaries}" + ) + + permutation = build_detector_permutation(circuit, metadata) + maps = _build_stab_maps(distance, "XV") + hx = maps["Hx_i32"].to(torch.uint8).cpu().numpy() + hz = maps["Hz_i32"].to(torch.uint8).cpu().numpy() + half = (distance * distance - 1) // 2 + min_difference = min(x - y for x, y in data_coordinates) + min_sum = min(x + y for x, y in data_coordinates) + mismatches = [] + error_names = {"X": "X_ERROR", "Y": "Y_ERROR", "Z": "Z_ERROR"} + + for pair_index, boundary_index in enumerate(cycle_boundaries): + insertion_index = boundary_index + 1 + pair_start = half + pair_index * 2 * half + for coordinate in data_coordinates: + qubit = coordinate_to_qubit[coordinate] + model_x, model_y = _google_to_xv_coordinate( + coordinate, + min_difference=min_difference, + min_sum=min_sum, + ) + row = (model_x - 1) // 2 + column = (model_y - 1) // 2 + data_index = row * distance + column + has_local_hadamard = (row + column) % 2 == 1 + + for physical_pauli, error_name in error_names.items(): + if physical_pauli == "Y": + css_components = {"x", "z"} + elif physical_pauli == "X": + css_components = {"z" if has_local_hadamard else "x"} + else: + css_components = {"x" if has_local_hadamard else "z"} + + faulty = circuit[:insertion_index] + faulty.append(error_name, [qubit], 1.0) + faulty += circuit[insertion_index:] + google_detectors, observables = faulty.compile_detector_sampler().sample( + shots=1, + separate_observables=True, + ) + actual_detectors = google_to_canonical( + np.asarray(google_detectors, dtype=np.uint8), + permutation, + )[0] + actual_observable = int(np.asarray(observables, dtype=np.uint8)[0, 0]) + + expected_detectors = np.zeros(int(circuit.num_detectors), dtype=np.uint8) + if "z" in css_components: + expected_detectors[pair_start : pair_start + half] ^= hx[:, data_index] + if "x" in css_components: + expected_detectors[pair_start + half : pair_start + 2 * half] ^= hz[:, data_index] + expected_observable = int( + (basis == "X" and "z" in css_components and row == 0) + or (basis == "Z" and "x" in css_components and column == 0) + ) + if not np.array_equal(actual_detectors, expected_detectors) or ( + actual_observable != expected_observable + ): + mismatches.append( + { + "bulk_pair_index": pair_index, + "coordinate": list(coordinate), + "qubit": qubit, + "physical_pauli": physical_pauli, + "local_hadamard": has_local_hadamard, + "css_components": sorted(css_components), + "actual_detector_indices": np.flatnonzero(actual_detectors).tolist(), + "expected_detector_indices": np.flatnonzero(expected_detectors).tolist(), + "actual_observable": actual_observable, + "expected_observable": expected_observable, + } + ) + + return { + "distance": distance, + "basis": basis, + "bulk_pair_indices": list(range(len(cycle_boundaries))), + "faults_checked": 3 * len(data_coordinates) * len(cycle_boundaries), + "mismatches": mismatches, + } + + + +def verify_final_data_fault_equivalence( + circuit: stim.Circuit, + metadata: Mapping[str, Any], +) -> dict[str, Any]: + """Compare Google final-measurement fault signatures with CSS-frame signatures. + + An X immediately before the final data-qubit measurement flips exactly one + physical measurement result. For every data qubit this checks that the + resulting Google detector/observable signature, after canonicalization, + equals the CSS parity-check column and logical-string parity used by the + predecoder. + """ + + distance = int(metadata["distance"]) + basis = str(metadata["basis"]).upper() + if basis not in {"X", "Z"}: + raise ValueError(f"basis must be X or Z, got {basis!r}") + + data_coordinates = [tuple(map(int, item)) for item in metadata["data_qubit_coords"]] + if len(data_coordinates) != distance * distance: + raise ValueError( + f"data coordinate count mismatch: {len(data_coordinates)} != {distance * distance}" + ) + qubit_coordinates = { + int(qubit): tuple(map(int, coordinate)) + for qubit, coordinate in circuit.get_final_qubit_coordinates().items() + } + coordinate_to_qubit = {coordinate: qubit for qubit, coordinate in qubit_coordinates.items()} + missing_qubits = [coordinate for coordinate in data_coordinates if coordinate not in coordinate_to_qubit] + if missing_qubits: + raise ValueError(f"data coordinates missing from circuit: {missing_qubits}") + data_qubits = {coordinate_to_qubit[coordinate] for coordinate in data_coordinates} + + final_measurement_index = None + for instruction_index in range(len(circuit) - 1, -1, -1): + instruction = circuit[instruction_index] + if instruction.name not in {"M", "MX", "MY"}: + continue + measured_qubits = { + int(target.value) + for target in instruction.targets_copy() + if target.is_qubit_target + } + if measured_qubits == data_qubits: + final_measurement_index = instruction_index + break + if final_measurement_index is None: + raise ValueError("could not find the final all-data-qubit measurement") + + permutation = build_detector_permutation(circuit, metadata) + maps = _build_stab_maps(distance, "XV") + parity_matrix = ( + maps["Hx_i32"] if basis == "X" else maps["Hz_i32"] + ).to(torch.uint8).cpu().numpy() + half = (distance * distance - 1) // 2 + boundary_start = int(circuit.num_detectors) - half + min_difference = min(x - y for x, y in data_coordinates) + min_sum = min(x + y for x, y in data_coordinates) + mismatches = [] + + for coordinate in data_coordinates: + qubit = coordinate_to_qubit[coordinate] + model_x, model_y = _google_to_xv_coordinate( + coordinate, + min_difference=min_difference, + min_sum=min_sum, + ) + row = (model_x - 1) // 2 + column = (model_y - 1) // 2 + data_index = row * distance + column + + faulty = circuit[:final_measurement_index] + faulty.append("X_ERROR", [qubit], 1.0) + faulty += circuit[final_measurement_index:] + google_detectors, observables = faulty.compile_detector_sampler().sample( + shots=1, + separate_observables=True, + ) + actual_detectors = google_to_canonical( + np.asarray(google_detectors, dtype=np.uint8), + permutation, + )[0] + actual_observable = int(np.asarray(observables, dtype=np.uint8)[0, 0]) + + expected_detectors = np.zeros(int(circuit.num_detectors), dtype=np.uint8) + expected_detectors[boundary_start:] = parity_matrix[:, data_index] % 2 + expected_observable = int(row == 0) if basis == "X" else int(column == 0) + if not np.array_equal(actual_detectors, expected_detectors) or ( + actual_observable != expected_observable + ): + mismatches.append( + { + "coordinate": list(coordinate), + "qubit": qubit, + "model_data_index": data_index, + "actual_detector_indices": np.flatnonzero(actual_detectors).tolist(), + "expected_detector_indices": np.flatnonzero(expected_detectors).tolist(), + "actual_observable": actual_observable, + "expected_observable": expected_observable, + } + ) + + return { + "distance": distance, + "basis": basis, + "faults_checked": len(data_coordinates), + "mismatches": mismatches, + } + + +def wilson_interval(errors: int, shots: int, z: float = 1.96) -> tuple[float, float]: + if shots <= 0: + return float("nan"), float("nan") + p = float(errors) / float(shots) + denominator = 1.0 + z * z / shots + center = (p + z * z / (2.0 * shots)) / denominator + half_width = ( + z + * math.sqrt((p * (1.0 - p) + z * z / (4.0 * shots)) / shots) + / denominator + ) + return max(0.0, center - half_width), min(1.0, center + half_width) + + +def paired_error_counts( + candidate_errors: np.ndarray, + baseline_errors: np.ndarray, +) -> dict[str, int | float]: + candidate = np.asarray(candidate_errors, dtype=np.bool_).reshape(-1) + baseline = np.asarray(baseline_errors, dtype=np.bool_).reshape(-1) + if candidate.shape != baseline.shape: + raise ValueError( + f"paired error shape mismatch: {candidate.shape} != {baseline.shape}" + ) + candidate_only = int(np.count_nonzero(candidate & ~baseline)) + baseline_only = int(np.count_nonzero(~candidate & baseline)) + both = int(np.count_nonzero(candidate & baseline)) + neither = int(candidate.size - candidate_only - baseline_only - both) + result = _paired_statistics_from_counts( + samples=int(candidate.size), + candidate_only=candidate_only, + baseline_only=baseline_only, + both=both, + neither=neither, + ) + # Kept for backward compatibility with existing candidate-vs-PyMatching rows. + result["delta_ler_vs_pymatching"] = result["delta_ler"] + return result + + +def _paired_statistics_from_counts( + *, + samples: int, + candidate_only: int, + baseline_only: int, + both: int, + neither: int, +) -> dict[str, int | float]: + if samples < 0 or min(candidate_only, baseline_only, both, neither) < 0: + raise ValueError("paired counts must be non-negative") + if candidate_only + baseline_only + both + neither != samples: + raise ValueError("paired outcome counts must sum to samples") + delta_errors = candidate_only - baseline_only + delta_ler = float(delta_errors / samples) if samples else float("nan") + if samples > 1: + difference_square_sum = candidate_only + baseline_only + variance = max( + 0.0, + (difference_square_sum - samples * delta_ler * delta_ler) + / (samples - 1), + ) + standard_error = math.sqrt(variance / samples) + else: + standard_error = 0.0 if samples == 1 else float("nan") + margin = 1.96 * standard_error + return { + "samples": samples, + "candidate_only_errors": candidate_only, + "baseline_only_errors": baseline_only, + "both_errors": both, + "neither_errors": neither, + "delta_logical_errors": delta_errors, + "delta_ler": delta_ler, + "standard_error": standard_error, + "ci95_low": max(-1.0, delta_ler - margin), + "ci95_high": min(1.0, delta_ler + margin), + } + + +MODEL_PAIRWISE_PRIORITY = ( + "qadapt", + "ising-fast", + "ising_fast_t0_e100", +) + + +def build_model_pairwise_rows( + error_masks: Mapping[str, np.ndarray], + case_fields: Mapping[str, Any], +) -> list[dict[str, Any]]: + """Build pairwise rows when more than one neural model is selected.""" + + known = [name for name in MODEL_PAIRWISE_PRIORITY if name in error_masks] + extras = sorted(set(error_masks) - set(known) - {"pymatching"}) + methods = known + extras + rows: list[dict[str, Any]] = [] + for candidate_index, candidate in enumerate(methods): + for baseline in methods[candidate_index + 1 :]: + rows.append( + { + **dict(case_fields), + "candidate": candidate, + "baseline": baseline, + **paired_error_counts( + error_masks[candidate], + error_masks[baseline], + ), + } + ) + rows[-1].pop("delta_ler_vs_pymatching", None) + return rows + + +def aggregate_paired_rows( + rows: Iterable[Mapping[str, Any]], +) -> list[dict[str, Any]]: + """Pool case-level paired outcomes without treating cases as independent CIs.""" + + totals: dict[tuple[str, str], dict[str, Any]] = {} + for row in rows: + key = (str(row["candidate"]), str(row["baseline"])) + entry = totals.setdefault( + key, + { + "candidate": key[0], + "baseline": key[1], + "cases": 0, + "samples": 0, + "candidate_only_errors": 0, + "baseline_only_errors": 0, + "both_errors": 0, + "neither_errors": 0, + }, + ) + entry["cases"] += 1 + for field in ( + "samples", + "candidate_only_errors", + "baseline_only_errors", + "both_errors", + "neither_errors", + ): + entry[field] += int(row[field]) + + results = [] + for entry in totals.values(): + stats = _paired_statistics_from_counts( + samples=int(entry["samples"]), + candidate_only=int(entry["candidate_only_errors"]), + baseline_only=int(entry["baseline_only_errors"]), + both=int(entry["both_errors"]), + neither=int(entry["neither_errors"]), + ) + results.append( + { + "candidate": entry["candidate"], + "baseline": entry["baseline"], + "cases": entry["cases"], + **stats, + } + ) + return sorted(results, key=lambda row: (row["candidate"], row["baseline"])) + + +def aggregate_rows(rows: Iterable[Mapping[str, Any]]) -> dict[str, dict[str, Any]]: + totals: dict[str, dict[str, Any]] = {} + for row in rows: + if row.get("status", "ok") != "ok": + continue + method = str(row["method"]) + entry = totals.setdefault( + method, + {"method": method, "cases": 0, "shots": 0, "logical_errors": 0}, + ) + entry["cases"] += 1 + entry["shots"] += int(row["shots"]) + entry["logical_errors"] += int(row["logical_errors"]) + for entry in totals.values(): + shots = int(entry["shots"]) + errors = int(entry["logical_errors"]) + low, high = wilson_interval(errors, shots) + entry.update( + ler=float(errors / shots) if shots else float("nan"), + ci95_low=low, + ci95_high=high, + ) + return totals + + +def _read_b8( + path: Path, + *, + num_detectors: int, + num_observables: int, +) -> np.ndarray: + data = stim.read_shot_data_file( + path=str(path), + format="b8", + num_detectors=int(num_detectors), + num_observables=int(num_observables), + ) + return np.asarray(data, dtype=np.uint8) + + +def load_case_data( + case: GoogleQECCase, + *, + max_shots: int = 0, +) -> tuple[stim.Circuit, stim.Circuit, dict[str, Any], np.ndarray, np.ndarray]: + metadata = json.loads((case.path / "metadata.json").read_text()) + ideal = stim.Circuit.from_file(case.path / "circuit_ideal.stim") + noisy = stim.Circuit.from_file(case.path / "circuit_noisy_si1000.stim") + if ideal.num_detectors != noisy.num_detectors: + raise ValueError(f"ideal/noisy detector mismatch in {case.path}") + if ideal.num_observables != noisy.num_observables: + raise ValueError(f"ideal/noisy observable mismatch in {case.path}") + detectors = _read_b8( + case.path / "detection_events.b8", + num_detectors=int(ideal.num_detectors), + num_observables=0, + ) + observables = _read_b8( + case.path / "obs_flips_actual.b8", + num_detectors=0, + num_observables=int(ideal.num_observables), + ) + if detectors.shape[0] != observables.shape[0]: + raise ValueError( + f"detector/observable shot mismatch in {case.path}: " + f"{detectors.shape[0]} != {observables.shape[0]}" + ) + if detectors.shape[0] != int(metadata["shots"]): + raise ValueError( + f"metadata shot mismatch in {case.path}: " + f"{detectors.shape[0]} != {metadata['shots']}" + ) + limit = int(max_shots) + if limit > 0: + detectors = detectors[:limit] + observables = observables[:limit] + return ideal, noisy, metadata, detectors, observables + + +def build_matcher(noisy_circuit: stim.Circuit) -> pymatching.Matching: + dem = noisy_circuit.detector_error_model(decompose_errors=True) + return pymatching.Matching.from_detector_error_model(dem) + + +def _decode_batch(matcher: pymatching.Matching, detectors: np.ndarray) -> np.ndarray: + predictions = np.asarray( + matcher.decode_batch(np.ascontiguousarray(detectors, dtype=np.uint8)), + dtype=np.uint8, + ) + if predictions.ndim == 1: + predictions = predictions.reshape(-1, 1) + return predictions + + +def time_single_shot( + matcher: pymatching.Matching, + detectors: np.ndarray, + *, + rounds: int, +) -> float: + rows = np.asarray(detectors, dtype=np.uint8) + if len(rows) == 0: + return float("nan") + for row in rows[: min(20, len(rows))]: + matcher.decode(row) + timings = [] + for row in rows: + start = time.perf_counter() + matcher.decode(row) + timings.append(time.perf_counter() - start) + return float(np.mean(timings) * 1e6 / max(1, int(rounds))) + + +def _error_metrics(predictions: np.ndarray, observables: np.ndarray) -> tuple[dict[str, Any], np.ndarray]: + predicted = np.asarray(predictions, dtype=np.uint8) + actual = np.asarray(observables, dtype=np.uint8) + if predicted.shape != actual.shape: + raise ValueError(f"prediction/observable shape mismatch: {predicted.shape} != {actual.shape}") + error_mask = np.any(predicted != actual, axis=1) + errors = int(error_mask.sum()) + shots = int(len(error_mask)) + low, high = wilson_interval(errors, shots) + return ( + { + "logical_errors": errors, + "shots": shots, + "ler": float(errors / shots) if shots else float("nan"), + "ci95_low": low, + "ci95_high": high, + }, + error_mask, + ) + + +def evaluate_pymatching( + matcher: pymatching.Matching, + detectors: np.ndarray, + observables: np.ndarray, + *, + rounds: int, + latency_shots: int, +) -> tuple[dict[str, Any], np.ndarray]: + start = time.perf_counter() + predictions = _decode_batch(matcher, detectors) + batch_seconds = time.perf_counter() - start + metrics, error_mask = _error_metrics(predictions, observables) + latency_rows = detectors[: min(int(latency_shots), len(detectors))] + input_density = SyndromeDensityAccumulator() + input_density.update(detectors) + metrics.update( + { + "method": "pymatching", + "decoder": "uncorrelated_pymatching_si1000_prior", + "batch_decode_us_per_shot": float(batch_seconds * 1e6 / max(1, len(detectors))), + "pymatching_latency_us_per_round": time_single_shot( + matcher, + latency_rows, + rounds=rounds, + ), + **input_density.statistics("input"), + } + ) + return metrics, error_mask + + +def build_model_cfg( + spec: BenchmarkModel, + case: GoogleQECCase, + *, + config_name: str, + batch_size: int, + latency_shots: int, +) -> Any: + cfg = OmegaConf.load(config_path(config_name)) + cfg.model_id = int(spec.model_id) + cfg.distance = int(case.distance) + cfg.n_rounds = int(case.rounds) + cfg.workflow.task = "inference" + public_spec = validate_public_config(cfg) + cfg = apply_public_defaults_and_model(cfg, public_spec) + cfg.model_checkpoint_file = str(spec.checkpoint) + cfg.test.meas_basis_test = str(case.basis) + cfg.test.num_samples = int(case.shots) + cfg.test.latency_num_samples = int(latency_shots) + cfg.test.batch_size = int(batch_size) + cfg.test.dataloader_num_workers = 0 + return cfg + + +def evaluate_predecoder( + model: torch.nn.Module, + cfg: Any, + matcher: pymatching.Matching, + google_detectors: np.ndarray, + canonical_detectors: np.ndarray, + observables: np.ndarray, + canonical_to_source: np.ndarray, + *, + device: torch.device, + rounds: int, + batch_size: int, + latency_shots: int, +) -> tuple[dict[str, Any], np.ndarray]: + maps = _build_stab_maps(int(cfg.distance), str(cfg.data.code_rotation)) + module = PreDecoderMemoryEvalModule(model, cfg, maps, device).to(device).eval() + predictions = [] + residual_google_rows = [] + model_seconds = 0.0 + residual_matching_seconds = 0.0 + + input_density = SyndromeDensityAccumulator() + residual_density = SyndromeDensityAccumulator() + input_density.update(google_detectors) + def synchronize() -> None: + if device.type == "cuda": + torch.cuda.synchronize(device) + + with torch.inference_mode(): + for start_index in range(0, len(canonical_detectors), int(batch_size)): + canonical_batch = canonical_detectors[ + start_index : start_index + int(batch_size) + ] + tensor = torch.from_numpy(canonical_batch).to( + device=device, + dtype=torch.uint8, + ) + synchronize() + started = time.perf_counter() + output = module(tensor) + synchronize() + model_seconds += time.perf_counter() - started + + pre_logical = output[:, :1].to(torch.uint8).cpu().numpy() + canonical_residual = output[:, 1:].to(torch.uint8).cpu().numpy() + google_residual = canonical_to_google( + canonical_residual, + canonical_to_source, + ) + started = time.perf_counter() + residual_prediction = _decode_batch(matcher, google_residual) + residual_density.update(google_residual) + residual_matching_seconds += time.perf_counter() - started + predictions.append((pre_logical + residual_prediction) % 2) + residual_google_rows.append(google_residual) + + final_predictions = np.concatenate(predictions, axis=0) + residual_google = np.concatenate(residual_google_rows, axis=0) + metrics, error_mask = _error_metrics(final_predictions, observables) + latency_rows = residual_google[: min(int(latency_shots), len(residual_google))] + residual_latency = time_single_shot(matcher, latency_rows, rounds=rounds) + density_statistics = model_density_statistics(input_density, residual_density) + shots = max(1, len(google_detectors)) + metrics.update( + { + "model_latency_us_per_shot": float(model_seconds * 1e6 / shots), + "residual_pymatching_batch_us_per_shot": float( + residual_matching_seconds * 1e6 / shots + ), + "end_to_end_batch_us_per_shot": float( + (model_seconds + residual_matching_seconds) * 1e6 / shots + ), + "pymatching_latency_us_per_round": residual_latency, + **density_statistics, + "syndrome_reduction": float(density_statistics["density_reduction_fraction"]), + } + ) + return metrics, error_mask + + +def _case_fields(case: GoogleQECCase) -> dict[str, Any]: + return { + "patch": case.patch, + "distance": case.distance, + "basis": case.basis, + "rounds": case.rounds, + } + + +def run_benchmark(args: argparse.Namespace) -> dict[str, Any]: + root = Path(args.benchmark_root).resolve() + selected_models = [DEFAULT_MODELS[name] for name in args.models] + missing_checkpoints = [ + str(spec.checkpoint) for spec in selected_models if not spec.checkpoint.is_file() + ] + if missing_checkpoints: + raise FileNotFoundError(f"Missing model checkpoint(s): {missing_checkpoints}") + cases = discover_cases( + root, + distances=set(args.distances), + rounds=set(args.rounds), + bases={basis.upper() for basis in args.bases}, + patches=set(args.patches) if args.patches else None, + ) + if not cases: + raise RuntimeError("No Google QEC benchmark cases match the selected filters") + if args.list_cases: + for case in cases: + print(case.path.relative_to(root)) + return {"cases": [str(case.path.relative_to(root)) for case in cases]} + + device = torch.device( + args.device or ("cuda:0" if torch.cuda.is_available() else "cpu") + ) + print(f"[google-qec] device={device} cases={len(cases)}") + model_cache: dict[str, torch.nn.Module] = {} + rows: list[dict[str, Any]] = [] + paired_comparisons: list[dict[str, Any]] = [] + + for case_index, case in enumerate(cases, start=1): + print( + f"[google-qec] case {case_index}/{len(cases)} " + f"{case.patch}/{case.basis}/r{case.rounds}" + ) + ideal, noisy, metadata, detectors, observables = load_case_data( + case, + max_shots=int(args.max_shots), + ) + matcher = build_matcher(noisy) + permutation = build_detector_permutation(ideal, metadata) + canonical_detectors = google_to_canonical(detectors, permutation) + baseline, baseline_errors = evaluate_pymatching( + matcher, + detectors, + observables, + rounds=case.rounds, + latency_shots=int(args.latency_shots), + ) + baseline.update(_case_fields(case), status="ok") + rows.append(baseline) + print( + f" pymatching: LER={baseline['ler']:.6g} " + f"({baseline['logical_errors']}/{baseline['shots']})" + ) + + if case.rounds < 2: + for spec in selected_models: + rows.append( + { + **_case_fields(case), + "method": spec.name, + "status": "unsupported", + "reason": "predecoder requires rounds >= 2", + "shots": int(len(detectors)), + } + ) + print(" neural predecoders skipped: rounds=1 is unsupported") + continue + + model_error_masks: dict[str, np.ndarray] = {} + for spec in selected_models: + cfg = build_model_cfg( + spec, + case, + config_name=args.config_name, + batch_size=int(args.batch_size), + latency_shots=int(args.latency_shots), + ) + if spec.name not in model_cache: + distributed = SimpleNamespace(rank=0, device=device) + loaded_model = load_model_checkpoint( + cfg, + checkpoint=spec.checkpoint, + model_id=spec.model_id, + distributed=distributed, + ).to(device).eval() + model_cache[spec.name] = maybe_compile_model( + loaded_model, + enabled=bool(args.torch_compile), + mode=str(args.torch_compile_mode), + ) + if args.torch_compile: + print(f" {spec.name}: torch.compile mode={args.torch_compile_mode}") + metrics, error_mask = evaluate_predecoder( + model_cache[spec.name], + cfg, + matcher, + detectors, + canonical_detectors, + observables, + permutation, + device=device, + rounds=case.rounds, + batch_size=int(args.batch_size), + latency_shots=int(args.latency_shots), + ) + model_error_masks[spec.name] = error_mask + metrics.update( + _case_fields(case), + method=spec.name, + checkpoint=str(spec.checkpoint), + status="ok", + ) + paired_vs_pymatching = paired_error_counts(error_mask, baseline_errors) + for field in ( + "candidate_only_errors", + "baseline_only_errors", + "both_errors", + "neither_errors", + "delta_logical_errors", + "delta_ler_vs_pymatching", + ): + metrics[field] = paired_vs_pymatching[field] + metrics.update( + paired_samples_vs_pymatching=paired_vs_pymatching["samples"], + paired_standard_error_vs_pymatching=paired_vs_pymatching["standard_error"], + paired_ci95_low_vs_pymatching=paired_vs_pymatching["ci95_low"], + paired_ci95_high_vs_pymatching=paired_vs_pymatching["ci95_high"], + ) + baseline_latency = float(baseline["pymatching_latency_us_per_round"]) + residual_latency = float(metrics["pymatching_latency_us_per_round"]) + metrics["pymatching_speedup"] = ( + baseline_latency / residual_latency + if residual_latency > 0 and math.isfinite(residual_latency) + else float("nan") + ) + rows.append(metrics) + print( + f" {spec.name}: LER={metrics['ler']:.6g} " + f"delta={metrics['delta_ler_vs_pymatching']:+.6g} " + f"syndrome_reduction={metrics['syndrome_reduction']:.3f}" + ) + + paired_comparisons.extend( + build_model_pairwise_rows(model_error_masks, _case_fields(case)) + ) + payload = { + "schema_version": 2, + "generated_at": datetime.now(timezone.utc).isoformat(), + "benchmark_root": str(root), + "decoder_prior": "Google circuit_noisy_si1000.stim DEM", + "detector_mapping": "Google physical order <-> repository XV canonical order", + "device": str(device), + "filters": { + "distances": list(args.distances), + "rounds": list(args.rounds), + "bases": list(args.bases), + "patches": list(args.patches or []), + "max_shots": int(args.max_shots), + "batch_size": int(args.batch_size), + "latency_shots": int(args.latency_shots), + "torch_compile": bool(args.torch_compile), + "torch_compile_mode": str(args.torch_compile_mode), + }, + "models": { + spec.name: { + "model_id": spec.model_id, + "checkpoint": str(spec.checkpoint), + } + for spec in selected_models + }, + "rows": rows, + "aggregate": aggregate_rows(rows), + "paired_comparisons": paired_comparisons, + "paired_aggregate": aggregate_paired_rows(paired_comparisons), + } + return payload + + + +def merge_benchmark_payloads( + payloads: Sequence[Mapping[str, Any]], +) -> dict[str, Any]: + """Merge disjoint benchmark shards and recompute all pooled statistics.""" + + if not payloads: + raise ValueError("at least one benchmark payload is required") + reference = payloads[0] + for index, payload in enumerate(payloads): + if int(payload.get("schema_version", 0)) != 2: + raise ValueError(f"benchmark shard {index} is not schema_version=2") + for field in ( + "benchmark_root", + "decoder_prior", + "detector_mapping", + "models", + ): + if payload.get(field) != reference.get(field): + raise ValueError(f"benchmark shard {index} disagrees on {field}") + + rows = [dict(row) for payload in payloads for row in payload.get("rows", [])] + paired = [ + dict(row) + for payload in payloads + for row in payload.get("paired_comparisons", []) + ] + row_keys = [ + ( + str(row.get("patch")), + int(row.get("distance", 0)), + str(row.get("basis")), + int(row.get("rounds", 0)), + str(row.get("method")), + ) + for row in rows + ] + if len(row_keys) != len(set(row_keys)): + raise ValueError("benchmark shards contain duplicate case/method rows") + paired_keys = [ + ( + str(row.get("patch")), + int(row.get("distance", 0)), + str(row.get("basis")), + int(row.get("rounds", 0)), + str(row.get("candidate")), + str(row.get("baseline")), + ) + for row in paired + ] + if len(paired_keys) != len(set(paired_keys)): + raise ValueError("benchmark shards contain duplicate paired comparisons") + + rows.sort( + key=lambda row: ( + int(row.get("distance", 0)), + str(row.get("patch")), + str(row.get("basis")), + int(row.get("rounds", 0)), + str(row.get("method")), + ) + ) + paired.sort( + key=lambda row: ( + int(row.get("distance", 0)), + str(row.get("patch")), + str(row.get("basis")), + int(row.get("rounds", 0)), + str(row.get("candidate")), + str(row.get("baseline")), + ) + ) + max_shots = { + int(payload.get("filters", {}).get("max_shots", 0)) for payload in payloads + } + if len(max_shots) != 1: + raise ValueError("benchmark shards disagree on max_shots") + execution_filters = {} + for field in ( + "batch_size", + "latency_shots", + "torch_compile", + "torch_compile_mode", + ): + values = {payload.get("filters", {}).get(field) for payload in payloads} + if len(values) != 1: + raise ValueError(f"benchmark shards disagree on {field}") + execution_filters[field] = values.pop() + return { + "schema_version": 2, + "generated_at": datetime.now(timezone.utc).isoformat(), + "benchmark_root": reference["benchmark_root"], + "decoder_prior": reference["decoder_prior"], + "detector_mapping": reference["detector_mapping"], + "device": "merged_shards", + "filters": { + "distances": sorted({int(row["distance"]) for row in rows}), + "rounds": sorted({int(row["rounds"]) for row in rows}), + "bases": sorted({str(row["basis"]) for row in rows}), + "patches": sorted({str(row["patch"]) for row in rows}), + "max_shots": max_shots.pop(), + **execution_filters, + }, + "models": reference["models"], + "rows": rows, + "aggregate": aggregate_rows(rows), + "paired_comparisons": paired, + "paired_aggregate": aggregate_paired_rows(paired), + } + +def write_results(payload: Mapping[str, Any], output_path: Path) -> tuple[Path, Path]: + output_path = Path(output_path) + output_path.parent.mkdir(parents=True, exist_ok=True) + output_path.write_text(json.dumps(payload, indent=2, sort_keys=True) + "\n") + csv_path = output_path.with_suffix(".csv") + rows = list(payload.get("rows", [])) + fieldnames = sorted({str(key) for row in rows for key in row}) + with csv_path.open("w", newline="") as stream: + writer = csv.DictWriter(stream, fieldnames=fieldnames) + writer.writeheader() + writer.writerows(rows) + paired_rows = list(payload.get("paired_comparisons", [])) + paired_csv_path = output_path.with_name( + f"{output_path.stem}_paired.csv" + ) + paired_fields = sorted({str(key) for row in paired_rows for key in row}) + with paired_csv_path.open("w", newline="") as stream: + writer = csv.DictWriter(stream, fieldnames=paired_fields) + if paired_fields: + writer.writeheader() + writer.writerows(paired_rows) + return output_path, csv_path + + +def parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace: + parser = argparse.ArgumentParser( + description=( + "Evaluate PyMatching and released pre-decoders on Google Willow " + "QEC hardware samples." + ) + ) + parser.add_argument("--benchmark-root", type=Path, default=DEFAULT_BENCHMARK_ROOT) + parser.add_argument("--distances", nargs="+", type=int, default=[3, 5, 7]) + parser.add_argument( + "--rounds", + nargs="+", + type=int, + default=[13], + help="Google cycle counts. The default r13 is the calibration slice.", + ) + parser.add_argument("--bases", nargs="+", choices=("X", "Z"), default=["X", "Z"]) + parser.add_argument( + "--patches", + nargs="+", + default=None, + help="Optional exact patch directory names, for example d7_at_q6_7.", + ) + parser.add_argument( + "--models", + nargs="+", + choices=tuple(DEFAULT_MODELS), + default=list(DEFAULT_MODELS), + ) + parser.add_argument("--config-name", default="examples/qadapt/config_qadapt_t0_base") + parser.add_argument("--max-shots", type=int, default=0, help="0 uses all shots.") + parser.add_argument("--batch-size", type=int, default=512) + parser.add_argument("--latency-shots", type=int, default=512) + parser.add_argument("--device", default=None) + parser.add_argument( + "--torch-compile", + action="store_true", + help="Compile each neural model once with dynamic input shapes.", + ) + parser.add_argument( + "--torch-compile-mode", + choices=( + "default", + "reduce-overhead", + "max-autotune", + "max-autotune-no-cudagraphs", + ), + default="default", + ) + parser.add_argument("--output", type=Path, default=None) + parser.add_argument( + "--merge-inputs", + nargs="+", + type=Path, + default=None, + help="Merge disjoint schema-v2 benchmark JSON shards instead of running inference.", + ) + parser.add_argument("--list-cases", action="store_true") + args = parser.parse_args(argv) + if args.max_shots < 0: + parser.error("--max-shots must be >= 0") + if args.batch_size <= 0: + parser.error("--batch-size must be positive") + if args.latency_shots <= 0: + parser.error("--latency-shots must be positive") + if args.output is None: + args.output = Path(args.benchmark_root) / "ising_decoder_results/results.json" + if args.merge_inputs and args.list_cases: + parser.error("--merge-inputs cannot be combined with --list-cases") + return args + + +def main(argv: Sequence[str] | None = None) -> int: + args = parse_args(argv) + if args.merge_inputs: + payload = merge_benchmark_payloads( + [json.loads(Path(path).read_text(encoding="utf-8")) for path in args.merge_inputs] + ) + payload["merged_inputs"] = [str(Path(path).resolve()) for path in args.merge_inputs] + print(f"[google-qec] merged {len(args.merge_inputs)} shards") + else: + payload = run_benchmark(args) + if args.list_cases: + return 0 + json_path, csv_path = write_results(payload, args.output) + print(f"[google-qec] JSON: {json_path}") + print(f"[google-qec] CSV: {csv_path}") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/code/scripts/qadapt_example_utils.py b/code/scripts/qadapt_example_utils.py new file mode 100644 index 0000000..daf6316 --- /dev/null +++ b/code/scripts/qadapt_example_utils.py @@ -0,0 +1,229 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Shared command construction and execution for the public QAdapt examples.""" + +from __future__ import annotations + +import argparse +import os +import shlex +import subprocess +import sys +from concurrent.futures import ThreadPoolExecutor, as_completed +from dataclasses import dataclass +from pathlib import Path +from typing import Sequence + + +REPO_ROOT = Path(__file__).resolve().parents[2] +PAIRED_INFERENCE_SCRIPT = REPO_ROOT / "code" / "scripts" / "paired_inference_compare.py" +TASK_CONFIGS = ( + ("t0_base", "examples/qadapt/config_qadapt_t0_base"), + ("t1_meas_1p5", "examples/qadapt/config_qadapt_t1_meas_1p5"), + ("t2_cnot_1p5", "examples/qadapt/config_qadapt_t2_cnot_1p5"), + ("t3_idle_1p5", "examples/qadapt/config_qadapt_t3_idle_1p5"), + ("t4_z_bias_1p5", "examples/qadapt/config_qadapt_t4_z_bias_1p5"), +) + + +@dataclass(frozen=True) +class ModelArgument: + name: str + model_id: int + checkpoint: Path + + +@dataclass(frozen=True) +class InferenceJob: + label: str + command: tuple[str, ...] + output_path: Path + + +def parse_model_argument(value: str) -> ModelArgument: + parts = value.split(":", 2) + if len(parts) != 3: + raise argparse.ArgumentTypeError( + "--model must be formatted as name:model_id:/path/to/checkpoint" + ) + name, model_id_raw, checkpoint_raw = (part.strip() for part in parts) + if not name or not checkpoint_raw: + raise argparse.ArgumentTypeError("model name and checkpoint must not be empty") + try: + model_id = int(model_id_raw) + except ValueError as exc: + raise argparse.ArgumentTypeError( + f"invalid model_id: {model_id_raw}" + ) from exc + checkpoint = Path(checkpoint_raw).expanduser() + if not checkpoint.is_absolute(): + checkpoint = REPO_ROOT / checkpoint + return ModelArgument(name=name, model_id=model_id, checkpoint=checkpoint) + + +def _default_gpus() -> str: + visible = os.environ.get("CUDA_VISIBLE_DEVICES", "").strip() + return visible or "0" + + +def add_common_inference_args( + parser: argparse.ArgumentParser, + *, + default_output_dir: Path, + default_num_samples: int = 262144, +) -> None: + parser.add_argument( + "--model", + action="append", + type=parse_model_argument, + required=True, + help=( + "Repeat for each released model: name:model_id:/path/to/checkpoint. " + "Both .pt and .safetensors are supported." + ), + ) + parser.add_argument("--num-samples", type=int, default=default_num_samples) + parser.add_argument("--latency-num-samples", type=int, default=10000) + parser.add_argument("--batch-size", type=int, default=2048) + parser.add_argument("--num-workers", type=int, default=0) + parser.add_argument("--basis", choices=("both", "X", "Z"), default="both") + parser.add_argument("--seed", type=int, default=12345) + parser.add_argument("--gpus", default=_default_gpus()) + parser.add_argument("--parallelism", type=int, default=1) + parser.add_argument( + "--python", + default=os.environ.get("PREDECODER_PYTHON", sys.executable), + ) + parser.add_argument("--output-dir", type=Path, default=default_output_dir) + parser.add_argument("--resume", action="store_true") + parser.add_argument("--dry-run", action="store_true") + + +def checkpoint_specs(args: argparse.Namespace) -> tuple[ModelArgument, ...]: + specs = tuple(args.model) + names = [spec.name for spec in specs] + if len(names) != len(set(names)): + raise ValueError(f"model names must be unique: {names}") + return specs + + +def parse_gpus(value: str | Sequence[str]) -> list[str]: + raw = value.split(",") if isinstance(value, str) else value + result = [str(item).strip() for item in raw if str(item).strip()] + if not result: + raise ValueError("at least one GPU must be selected") + return result + + +def build_paired_command( + args: argparse.Namespace, + *, + output_path: Path, + distance: int, + n_rounds: int, + config_name: str | None = None, + config_file: Path | None = None, +) -> tuple[str, ...]: + if (config_name is None) == (config_file is None): + raise ValueError("provide exactly one of config_name or config_file") + command = [ + str(args.python), + "-u", + str(PAIRED_INFERENCE_SCRIPT), + ] + if config_name is not None: + command.extend(("--config-name", config_name)) + else: + command.extend(("--config-file", str(Path(config_file)))) + command.extend( + ( + "--distance", + str(distance), + "--n-rounds", + str(n_rounds), + "--num-samples", + str(args.num_samples), + "--latency-num-samples", + str(args.latency_num_samples), + "--batch-size", + str(args.batch_size), + "--num-workers", + str(args.num_workers), + "--seed", + str(args.seed), + "--basis", + str(args.basis), + "--device", + "cuda:0", + "--output", + str(output_path), + ) + ) + for spec in checkpoint_specs(args): + command.extend( + ("--model", f"{spec.name}:{spec.model_id}:{spec.checkpoint}") + ) + return tuple(command) + + +def _run_one(job: InferenceJob, gpu: str) -> tuple[InferenceJob, int, Path]: + job.output_path.parent.mkdir(parents=True, exist_ok=True) + log_path = job.output_path.with_suffix(".log") + env = dict(os.environ) + env["CUDA_VISIBLE_DEVICES"] = gpu + with log_path.open("w", encoding="utf-8") as stream: + completed = subprocess.run( + job.command, + cwd=REPO_ROOT, + env=env, + stdout=stream, + stderr=subprocess.STDOUT, + check=False, + ) + return job, int(completed.returncode), log_path + + +def run_jobs( + jobs: Sequence[InferenceJob], + *, + gpus: Sequence[str], + parallelism: int, + resume: bool, + dry_run: bool, +) -> None: + selected_gpus = parse_gpus(gpus) + workers = max(1, min(int(parallelism), len(selected_gpus))) + pending = [ + job for job in jobs + if not (resume and job.output_path.is_file()) + ] + skipped = len(jobs) - len(pending) + if skipped: + print(f"[resume] skipped {skipped} existing outputs") + if dry_run: + for index, job in enumerate(pending): + gpu = selected_gpus[index % workers] + print( + f"[dry-run] gpu={gpu} label={job.label} " + + shlex.join(job.command) + ) + return + failures = [] + with ThreadPoolExecutor(max_workers=workers) as executor: + futures = { + executor.submit(_run_one, job, selected_gpus[index % workers]): job + for index, job in enumerate(pending) + } + for future in as_completed(futures): + job, returncode, log_path = future.result() + if returncode: + failures.append((job, returncode, log_path)) + print(f"[fail] {job.label} log={log_path}") + else: + print(f"[done] {job.label} output={job.output_path}") + if failures: + details = "\n".join( + f" - {job.label}: exit={returncode}, log={log_path}" + for job, returncode, log_path in failures + ) + raise RuntimeError(f"Released-model inference jobs failed:\n{details}") diff --git a/code/workflows/config_validator.py b/code/workflows/config_validator.py index a5ac02f..32eb256 100644 --- a/code/workflows/config_validator.py +++ b/code/workflows/config_validator.py @@ -1,5 +1,6 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 +# Modified in 2026 for the QAdapt Hugging Face release: added HTNet defaults. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. @@ -45,6 +46,7 @@ _INTERNAL_ROTATION_TO_PUBLIC = {v: k for k, v in _PUBLIC_ROTATION_TO_INTERNAL.it _PUBLIC_MODEL_ID_TO_LR = { 1: 3e-4, + 111: 3e-4, 2: 2e-4, 3: 1e-4, 4: 2e-4, @@ -557,6 +559,18 @@ def apply_public_defaults_and_model(cfg: DictConfig, model_spec: PublicModelSpec merged.model.version = model_spec.model_version merged.model.num_filters = list(model_spec.num_filters) merged.model.kernel_size = list(model_spec.kernel_size) + if model_spec.channels is not None: + merged.model.channels = int(model_spec.channels) + if model_spec.expand_channels is not None: + merged.model.expand_channels = int(model_spec.expand_channels) + if model_spec.num_blocks is not None: + merged.model.num_blocks = int(model_spec.num_blocks) + if model_spec.joint_groups is not None: + merged.model.joint_groups = int(model_spec.joint_groups) + if model_spec.norm_groups is not None: + merged.model.norm_groups = int(model_spec.norm_groups) + if model_spec.se_reduction is not None: + merged.model.se_reduction = int(model_spec.se_reduction) _apply_code_specific_defaults(merged, code, model_spec) diff --git a/conf/examples/qadapt/config_qadapt_t0_base.yaml b/conf/examples/qadapt/config_qadapt_t0_base.yaml new file mode 100644 index 0000000..d3631a2 --- /dev/null +++ b/conf/examples/qadapt/config_qadapt_t0_base.yaml @@ -0,0 +1,40 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +# QAdapt T0 inference environment. + +model_id: 111 +distance: 9 +n_rounds: 9 + +workflow: + task: inference + +data: + code_rotation: O1 + noise_model: + p_prep_X: 0.0010000 + p_prep_Z: 0.0010000 + p_meas_X: 0.0100000 + p_meas_Z: 0.0100000 + p_idle_cnot_X: 0.0003330 + p_idle_cnot_Y: 0.0003330 + p_idle_cnot_Z: 0.0003330 + p_idle_spam_X: 0.0006670 + p_idle_spam_Y: 0.0006670 + p_idle_spam_Z: 0.0006670 + p_cnot_IX: 0.0006670 + p_cnot_IY: 0.0006670 + p_cnot_IZ: 0.0006670 + p_cnot_XI: 0.0006670 + p_cnot_XX: 0.0006670 + p_cnot_XY: 0.0006670 + p_cnot_XZ: 0.0006670 + p_cnot_YI: 0.0006670 + p_cnot_YX: 0.0006670 + p_cnot_YY: 0.0006670 + p_cnot_YZ: 0.0006670 + p_cnot_ZI: 0.0006670 + p_cnot_ZX: 0.0006670 + p_cnot_ZY: 0.0006670 + p_cnot_ZZ: 0.0006670 diff --git a/conf/examples/qadapt/config_qadapt_t1_meas_1p5.yaml b/conf/examples/qadapt/config_qadapt_t1_meas_1p5.yaml new file mode 100644 index 0000000..46dd700 --- /dev/null +++ b/conf/examples/qadapt/config_qadapt_t1_meas_1p5.yaml @@ -0,0 +1,40 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +# Shared QAdapt T1 measurement-noise task. + +model_id: 111 +distance: 9 +n_rounds: 9 + +workflow: + task: inference + +data: + code_rotation: O1 + noise_model: + p_prep_X: 0.0010000 + p_prep_Z: 0.0010000 + p_meas_X: 0.0150000 + p_meas_Z: 0.0150000 + p_idle_cnot_X: 0.0003330 + p_idle_cnot_Y: 0.0003330 + p_idle_cnot_Z: 0.0003330 + p_idle_spam_X: 0.0006670 + p_idle_spam_Y: 0.0006670 + p_idle_spam_Z: 0.0006670 + p_cnot_IX: 0.0006670 + p_cnot_IY: 0.0006670 + p_cnot_IZ: 0.0006670 + p_cnot_XI: 0.0006670 + p_cnot_XX: 0.0006670 + p_cnot_XY: 0.0006670 + p_cnot_XZ: 0.0006670 + p_cnot_YI: 0.0006670 + p_cnot_YX: 0.0006670 + p_cnot_YY: 0.0006670 + p_cnot_YZ: 0.0006670 + p_cnot_ZI: 0.0006670 + p_cnot_ZX: 0.0006670 + p_cnot_ZY: 0.0006670 + p_cnot_ZZ: 0.0006670 diff --git a/conf/examples/qadapt/config_qadapt_t2_cnot_1p5.yaml b/conf/examples/qadapt/config_qadapt_t2_cnot_1p5.yaml new file mode 100644 index 0000000..1eda622 --- /dev/null +++ b/conf/examples/qadapt/config_qadapt_t2_cnot_1p5.yaml @@ -0,0 +1,40 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +# Shared QAdapt T2 CNOT-noise task. + +model_id: 111 +distance: 9 +n_rounds: 9 + +workflow: + task: inference + +data: + code_rotation: O1 + noise_model: + p_prep_X: 0.0010000 + p_prep_Z: 0.0010000 + p_meas_X: 0.0100000 + p_meas_Z: 0.0100000 + p_idle_cnot_X: 0.0003330 + p_idle_cnot_Y: 0.0003330 + p_idle_cnot_Z: 0.0003330 + p_idle_spam_X: 0.0006670 + p_idle_spam_Y: 0.0006670 + p_idle_spam_Z: 0.0006670 + p_cnot_IX: 0.0010005 + p_cnot_IY: 0.0010005 + p_cnot_IZ: 0.0010005 + p_cnot_XI: 0.0010005 + p_cnot_XX: 0.0010005 + p_cnot_XY: 0.0010005 + p_cnot_XZ: 0.0010005 + p_cnot_YI: 0.0010005 + p_cnot_YX: 0.0010005 + p_cnot_YY: 0.0010005 + p_cnot_YZ: 0.0010005 + p_cnot_ZI: 0.0010005 + p_cnot_ZX: 0.0010005 + p_cnot_ZY: 0.0010005 + p_cnot_ZZ: 0.0010005 diff --git a/conf/examples/qadapt/config_qadapt_t3_idle_1p5.yaml b/conf/examples/qadapt/config_qadapt_t3_idle_1p5.yaml new file mode 100644 index 0000000..26acbdd --- /dev/null +++ b/conf/examples/qadapt/config_qadapt_t3_idle_1p5.yaml @@ -0,0 +1,40 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +# Shared QAdapt T3 idle-noise task. + +model_id: 111 +distance: 9 +n_rounds: 9 + +workflow: + task: inference + +data: + code_rotation: O1 + noise_model: + p_prep_X: 0.0010000 + p_prep_Z: 0.0010000 + p_meas_X: 0.0100000 + p_meas_Z: 0.0100000 + p_idle_cnot_X: 0.0004995 + p_idle_cnot_Y: 0.0004995 + p_idle_cnot_Z: 0.0004995 + p_idle_spam_X: 0.0010005 + p_idle_spam_Y: 0.0010005 + p_idle_spam_Z: 0.0010005 + p_cnot_IX: 0.0006670 + p_cnot_IY: 0.0006670 + p_cnot_IZ: 0.0006670 + p_cnot_XI: 0.0006670 + p_cnot_XX: 0.0006670 + p_cnot_XY: 0.0006670 + p_cnot_XZ: 0.0006670 + p_cnot_YI: 0.0006670 + p_cnot_YX: 0.0006670 + p_cnot_YY: 0.0006670 + p_cnot_YZ: 0.0006670 + p_cnot_ZI: 0.0006670 + p_cnot_ZX: 0.0006670 + p_cnot_ZY: 0.0006670 + p_cnot_ZZ: 0.0006670 diff --git a/conf/examples/qadapt/config_qadapt_t4_z_bias_1p5.yaml b/conf/examples/qadapt/config_qadapt_t4_z_bias_1p5.yaml new file mode 100644 index 0000000..1a147e6 --- /dev/null +++ b/conf/examples/qadapt/config_qadapt_t4_z_bias_1p5.yaml @@ -0,0 +1,40 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +# Shared QAdapt T4 Z-biased-noise task. + +model_id: 111 +distance: 9 +n_rounds: 9 + +workflow: + task: inference + +data: + code_rotation: O1 + noise_model: + p_prep_X: 0.0015000 + p_prep_Z: 0.0010000 + p_meas_X: 0.0150000 + p_meas_Z: 0.0100000 + p_idle_cnot_X: 0.0003330 + p_idle_cnot_Y: 0.0003330 + p_idle_cnot_Z: 0.0004995 + p_idle_spam_X: 0.0006670 + p_idle_spam_Y: 0.0006670 + p_idle_spam_Z: 0.0010005 + p_cnot_IX: 0.0006670 + p_cnot_IY: 0.0006670 + p_cnot_IZ: 0.0010005 + p_cnot_XI: 0.0006670 + p_cnot_XX: 0.0006670 + p_cnot_XY: 0.0006670 + p_cnot_XZ: 0.0010005 + p_cnot_YI: 0.0006670 + p_cnot_YX: 0.0006670 + p_cnot_YY: 0.0006670 + p_cnot_YZ: 0.0010005 + p_cnot_ZI: 0.0010005 + p_cnot_ZX: 0.0010005 + p_cnot_ZY: 0.0010005 + p_cnot_ZZ: 0.0010005