From b9d03d9a29646985ab6a2e6d03db3da4a604e54d Mon Sep 17 00:00:00 2001 From: Julian Ng-Thow-Hing Date: Tue, 28 Jul 2026 16:58:47 -0700 Subject: [PATCH] Update [ghstack-poisoned] --- backends/webgpu/scripts/gen_wgsl_headers.py | 270 +++++++++++--- backends/webgpu/test/test_wgsl_codegen.py | 368 ++++++++++++++++++++ 2 files changed, 582 insertions(+), 56 deletions(-) diff --git a/backends/webgpu/scripts/gen_wgsl_headers.py b/backends/webgpu/scripts/gen_wgsl_headers.py index dd8cc4cc635..1e35b205888 100644 --- a/backends/webgpu/scripts/gen_wgsl_headers.py +++ b/backends/webgpu/scripts/gen_wgsl_headers.py @@ -26,11 +26,14 @@ import copy import hashlib import io +import os import re +import stat import sys +import tempfile from itertools import product from pathlib import Path -from typing import Any, Dict, List, NamedTuple, Optional, Set +from typing import Any, Dict, List, NamedTuple, Optional, Set, Tuple import yaml from yaml.constructor import ConstructorError @@ -523,23 +526,51 @@ def registry_path() -> Path: return BACKEND_ROOT / "runtime/WebGPUShaderRegistry.cpp" -def registry_entries() -> List[RegistryEntry]: - """Return one registry entry for every concrete generated shader.""" - entries = [] +def _registry_entry(header: Path) -> RegistryEntry: + suffix = "_wgsl.h" + if not header.name.endswith(suffix): + raise ValueError(f"unexpected generated header name: {header.name}") + name = header.name[: -len(suffix)] + return RegistryEntry( + name=name, + include=header.relative_to(BACKEND_ROOT).as_posix(), + symbol=symbol_base(name), + ) + + +def _collect_header_outputs() -> Tuple[Dict[Path, str], List[RegistryEntry]]: + """Render every concrete header once and reject global collisions.""" + outputs: Dict[Path, str] = {} + entries: List[RegistryEntry] = [] + registry_names: Set[str] = set() + registry_symbols: Set[str] = set() for wgsl in discover(): - for header, _ in headers_for_shader(wgsl): - suffix = "_wgsl.h" - if not header.name.endswith(suffix): - raise ValueError(f"unexpected generated header name: {header.name}") - name = header.name[: -len(suffix)] - entries.append( - RegistryEntry( - name=name, - include=header.relative_to(BACKEND_ROOT).as_posix(), - symbol=symbol_base(name), + try: + rendered_headers = list(headers_for_shader(wgsl)) + except Exception as error: + raise ValueError(f"{wgsl.relative_to(BACKEND_ROOT)}: {error}") from error + for header, rendered in rendered_headers: + if header in outputs: + raise ValueError( + "duplicate generated header path: " + f"{header.relative_to(BACKEND_ROOT)}" ) - ) - return sorted(entries) + entry = _registry_entry(header) + if entry.name in registry_names: + raise ValueError(f"duplicate shader registry name: {entry.name}") + if entry.symbol in registry_symbols: + raise ValueError(f"duplicate shader registry symbol: {entry.symbol}") + outputs[header] = rendered + entries.append(entry) + registry_names.add(entry.name) + registry_symbols.add(entry.symbol) + return outputs, sorted(entries) + + +def registry_entries() -> List[RegistryEntry]: + """Return one registry entry for every concrete generated shader.""" + _, entries = _collect_header_outputs() + return entries def render_registry(entries: List[RegistryEntry]) -> str: @@ -630,7 +661,143 @@ def headers_for_shader(wgsl): yield header, render_header(stem, text, stem) -def _report_drift(missing, stale) -> None: +def collect_outputs() -> Tuple[Dict[Path, bytes], List[Path]]: + """Render the complete output tree and report unexpected old headers.""" + header_outputs, entries = _collect_header_outputs() + outputs = { + path: rendered.encode("utf-8") for path, rendered in header_outputs.items() + } + registry = registry_path() + if registry in outputs: + raise ValueError(f"duplicate generated output path: {registry}") + outputs[registry] = render_registry(entries).encode("utf-8") + + expected_headers = set(header_outputs) + actual_headers = set((BACKEND_ROOT / "runtime/ops").glob("**/*_wgsl.h")) + return outputs, sorted(actual_headers - expected_headers) + + +class _OriginalOutput(NamedTuple): + existed: bool + contents: bytes + mode: int + + +def _stage_bytes(destination: Path, contents: bytes, mode: int) -> Path: + """Write one same-directory candidate without changing its destination.""" + fd, name = tempfile.mkstemp( + prefix=f".{destination.name}.wgsl-gen-", + suffix=".tmp", + dir=destination.parent, + ) + temporary = Path(name) + try: + with os.fdopen(fd, "wb") as output: + output.write(contents) + temporary.chmod(mode) + except BaseException: + try: + temporary.unlink(missing_ok=True) + except OSError: + pass + raise + return temporary + + +def _cleanup_temporaries(temporaries) -> List[str]: + errors = [] + for temporary in temporaries: + try: + temporary.unlink(missing_ok=True) + except OSError as error: + errors.append(f"cannot remove temporary {temporary}: {error}") + return errors + + +def _stage_outputs( + outputs: Dict[Path, bytes], changed: List[Path] +) -> Tuple[Dict[Path, _OriginalOutput], Dict[Path, Path], List[str]]: + originals: Dict[Path, _OriginalOutput] = {} + staged: Dict[Path, Path] = {} + try: + for destination in sorted(changed): + if destination.exists(): + original = _OriginalOutput( + existed=True, + contents=destination.read_bytes(), + mode=stat.S_IMODE(destination.stat().st_mode), + ) + else: + original = _OriginalOutput(False, b"", 0o644) + originals[destination] = original + staged[destination] = _stage_bytes( + destination, outputs[destination], original.mode + ) + except BaseException as error: + cleanup_errors = _cleanup_temporaries(staged.values()) + if isinstance(error, OSError): + errors = [f"cannot stage generated output: {error}"] + cleanup_errors + return originals, staged, errors + raise + return originals, staged, [] + + +def _rollback_outputs( + originals: Dict[Path, _OriginalOutput], + replaced: List[Path], + staged: Dict[Path, Path], +) -> List[str]: + errors = [] + for destination in reversed(replaced): + original = originals[destination] + restore_temporary: Optional[Path] = None + try: + if original.existed: + restore_temporary = _stage_bytes( + destination, original.contents, original.mode + ) + os.replace(restore_temporary, destination) + else: + destination.unlink(missing_ok=True) + except OSError as error: + errors.append(f"cannot roll back {destination}: {error}") + finally: + if restore_temporary is not None: + errors.extend(_cleanup_temporaries([restore_temporary])) + errors.extend(_cleanup_temporaries(staged.values())) + return errors + + +def _publish_outputs(outputs: Dict[Path, bytes], changed: List[Path]) -> List[str]: + """Stage and publish changed outputs, rolling back reported failures.""" + originals, staged, stage_errors = _stage_outputs(outputs, changed) + if stage_errors: + return stage_errors + + replaced: List[Path] = [] + try: + for destination in sorted(changed): + try: + os.replace(staged[destination], destination) + except OSError: + raise + except BaseException: + replaced.append(destination) + raise + else: + replaced.append(destination) + except OSError as commit_error: + return [f"cannot publish generated output: {commit_error}"] + _rollback_outputs( + originals, replaced, staged + ) + except BaseException: + _rollback_outputs(originals, replaced, staged) + raise + + return _cleanup_temporaries(staged.values()) + + +def _report_drift(missing, stale, orphans) -> None: """Print the --check report for missing/stale committed headers.""" if missing: print("Missing embedded WGSL headers (run scripts/gen_wgsl_headers.py):") @@ -640,16 +807,10 @@ def _report_drift(missing, stale) -> None: print("Stale embedded WGSL headers (run scripts/gen_wgsl_headers.py):") for h in stale: print(f" {h.relative_to(BACKEND_ROOT)}") - - -def _sync_generated_output(output, want, check, missing, stale) -> None: - """Write one generated file, or record its --check drift.""" - if output.exists() and output.read_text() == want: - return - if check: - (missing if not output.exists() else stale).append(output) - else: - output.write_text(want) + if orphans: + print("Orphan embedded WGSL headers (remove or restore their sources):") + for h in orphans: + print(f" {h.relative_to(BACKEND_ROOT)}") def main(argv=None) -> int: @@ -661,38 +822,35 @@ def main(argv=None) -> int: ) args = parser.parse_args(argv) - stale = [] - missing = [] - errors = [] - for wgsl in discover(): - try: - rendered = list(headers_for_shader(wgsl)) - # A malformed spec raises yaml.YAMLError (incl. UniqueKeyLoader's - # ConstructorError) / ValueError / KeyError from parse_template_spec, and - # a malformed template raises AssertionError from preprocess; catch them - # all so a bad shader is a clean --check report, not a traceback. - except (ValueError, KeyError, AssertionError, yaml.YAMLError) as e: - errors.append(f"{wgsl.relative_to(BACKEND_ROOT)}: {e}") - continue - for header, want in rendered: - # Full-content compare (not just the sha) catches generator-logic drift too. - _sync_generated_output(header, want, args.check, missing, stale) + try: + outputs, orphans = collect_outputs() + missing = [] + stale = [] + for output, want in sorted(outputs.items()): + if not output.exists(): + missing.append(output) + elif output.read_bytes() != want: + stale.append(output) + except Exception as error: + print("Cannot generate WGSL outputs:") + print(f" {error}") + return 1 - if not errors: - try: - registry = render_registry(registry_entries()) - output = registry_path() - _sync_generated_output(output, registry, args.check, missing, stale) - except ValueError as e: - errors.append(f"shader registry: {e}") + if orphans: + _report_drift([], [], orphans) + return 1 + + if args.check: + if stale or missing: + _report_drift(missing, stale, []) + return 1 + return 0 + errors = _publish_outputs(outputs, missing + stale) if errors: - print("Cannot generate header (malformed shader):") - for e in errors: - print(f" {e}") - return 1 - if args.check and (stale or missing): - _report_drift(missing, stale) + print("Cannot publish WGSL outputs:") + for error in errors: + print(f" {error}") return 1 return 0 diff --git a/backends/webgpu/test/test_wgsl_codegen.py b/backends/webgpu/test/test_wgsl_codegen.py index 704278482c9..4a8be392351 100644 --- a/backends/webgpu/test/test_wgsl_codegen.py +++ b/backends/webgpu/test/test_wgsl_codegen.py @@ -13,10 +13,13 @@ import hashlib import importlib.util import io +import os import re +import stat import tempfile import unittest from pathlib import Path +from unittest import mock import yaml @@ -211,6 +214,25 @@ def test_committed_headers_match_generator(self) -> None: got, want, f"{header.name} stale; run scripts/gen_wgsl_headers.py" ) + def test_generated_output_manifest_digest(self) -> None: + outputs = sorted( + [ + *(g.BACKEND_ROOT / "runtime/ops").glob("**/*_wgsl.h"), + g.registry_path(), + ] + ) + digest = hashlib.sha256() + for output in outputs: + digest.update(output.relative_to(g.BACKEND_ROOT).as_posix().encode()) + digest.update(b"\0") + digest.update(output.read_bytes()) + digest.update(b"\0") + self.assertEqual(len(outputs), 134) + self.assertEqual( + digest.hexdigest(), + "19a0baf9345bec02fe2091a0e6320b81d966724bedc33feb0779c6a824925972", + ) + def test_rope_hf_reconstructs_full_2d_grid_stride(self) -> None: shader = ( g.BACKEND_ROOT / "runtime" / "ops" / "rope" / "rotary_embedding_hf.wgsl" @@ -351,6 +373,352 @@ def test_render_header_3d_emits_xyz(self) -> None: self.assertIn("inline constexpr uint32_t kFooWorkgroupSizeZ = 2;", h) +class WgslGenerationTransactionTest(unittest.TestCase): + _VALID_SHADER = "@compute @workgroup_size(1)\nfn main() {}\n" + + def setUp(self) -> None: + self._tmp = tempfile.TemporaryDirectory() + self.root = Path(self._tmp.name) + (self.root / "runtime/ops").mkdir(parents=True) + self._original_root = g.BACKEND_ROOT + g.BACKEND_ROOT = self.root + + def tearDown(self) -> None: + g.BACKEND_ROOT = self._original_root + self._tmp.cleanup() + + def _write_shader( + self, directory: str, stem: str, text: str = _VALID_SHADER + ) -> Path: + op_dir = self.root / "runtime/ops" / directory + op_dir.mkdir(parents=True, exist_ok=True) + shader = op_dir / f"{stem}.wgsl" + shader.write_text(text) + return shader + + def _write_template( + self, directory: str, stem: str, text: str, names: list[str] + ) -> Path: + shader = self._write_shader(directory, stem, text) + spec = { + stem: { + "parameter_names_with_default_values": {}, + "shader_variants": [{"NAME": name} for name in names], + } + } + shader.with_suffix(".yaml").write_text(yaml.safe_dump(spec)) + return shader + + def _snapshot(self): + return { + path.relative_to(self.root).as_posix(): ( + path.read_bytes(), + stat.S_IMODE(path.stat().st_mode), + ) + for path in sorted(self.root.rglob("*")) + if path.is_file() + } + + def _run(self, *args: str): + output = io.StringIO() + with contextlib.redirect_stdout(output): + result = g.main(list(args)) + return result, output.getvalue() + + def _assert_no_temps(self) -> None: + self.assertEqual(list(self.root.rglob("*.tmp")), []) + + @staticmethod + def _fail_nth(real_fn, n: int): + calls = 0 + + def wrapped(*args, **kwargs): + nonlocal calls + calls += 1 + if calls == n: + raise OSError(f"injected failure on call {n}") + return real_fn(*args, **kwargs) + + return wrapped + + def test_late_malformed_shader_leaves_tree_unchanged(self) -> None: + good = self._write_shader("a", "good") + good.with_name("good_wgsl.h").write_text("stale\n") + self._write_shader("z", "bad", "${MISSING\n") + before = self._snapshot() + + result, _ = self._run() + + self.assertEqual(result, 1) + self.assertEqual(self._snapshot(), before) + + def test_duplicate_registry_name_leaves_tree_unchanged(self) -> None: + self._write_shader("a", "shared") + self._write_shader("b", "shared") + before = self._snapshot() + + result, _ = self._run() + + self.assertEqual(result, 1) + self.assertEqual(self._snapshot(), before) + + def test_duplicate_registry_symbol_leaves_tree_unchanged(self) -> None: + self._write_shader("a", "foo_bar") + self._write_shader("b", "foo__bar") + before = self._snapshot() + + result, _ = self._run() + + self.assertEqual(result, 1) + self.assertEqual(self._snapshot(), before) + + def test_duplicate_output_path_leaves_tree_unchanged(self) -> None: + self._write_template("op", "op", self._VALID_SHADER, ["duplicate", "duplicate"]) + before = self._snapshot() + + result, _ = self._run() + + self.assertEqual(result, 1) + self.assertEqual(self._snapshot(), before) + + def test_second_stage_failure_leaves_tree_unchanged(self) -> None: + self._write_shader("a", "first") + self._write_shader("b", "second") + before = self._snapshot() + real_mkstemp = tempfile.mkstemp + + with mock.patch( + "tempfile.mkstemp", side_effect=self._fail_nth(real_mkstemp, 2) + ): + result, _ = self._run() + + self.assertEqual(result, 1) + self.assertEqual(self._snapshot(), before) + self._assert_no_temps() + + def test_staging_interrupt_leaves_tree_unchanged(self) -> None: + self._write_shader("a", "first") + self._write_shader("b", "second") + before = self._snapshot() + real_chmod = Path.chmod + + def interrupt_second(path, mode, **kwargs): + interrupt_second.calls += 1 + if interrupt_second.calls == 2: + raise KeyboardInterrupt("injected staging interruption") + return real_chmod(path, mode, **kwargs) + + interrupt_second.calls = 0 + with mock.patch.object( + Path, "chmod", autospec=True, side_effect=interrupt_second + ): + with self.assertRaises(KeyboardInterrupt): + self._run() + + self.assertEqual(self._snapshot(), before) + self._assert_no_temps() + + def test_replace_failure_restores_existing_destination(self) -> None: + self._write_shader("op", "op") + registry = g.registry_path() + registry.write_text("old registry\n") + registry.chmod(0o600) + before = self._snapshot() + real_replace = os.replace + + with mock.patch("os.replace", side_effect=self._fail_nth(real_replace, 2)): + result, _ = self._run() + + self.assertEqual(result, 1) + self.assertEqual(self._snapshot(), before) + self._assert_no_temps() + + def test_replace_failure_removes_new_destination(self) -> None: + self._write_shader("op", "op") + before = self._snapshot() + real_replace = os.replace + + with mock.patch("os.replace", side_effect=self._fail_nth(real_replace, 2)): + result, _ = self._run() + + self.assertEqual(result, 1) + self.assertEqual(self._snapshot(), before) + self._assert_no_temps() + + def test_multiple_rollback_errors_do_not_stop_later_restores(self) -> None: + headers = [] + for directory in ("a", "b", "c"): + shader = self._write_shader(directory, directory) + header = shader.with_name(f"{directory}_wgsl.h") + header.write_text(f"old {directory}\n") + headers.append(header) + registry = g.registry_path() + registry.write_text("old registry\n") + real_replace = os.replace + calls = [] + + def fail_commit_and_two_rollbacks(source, destination): + calls.append(Path(destination)) + if len(calls) in (4, 5, 6): + raise OSError(f"injected failure on replace {len(calls)}") + return real_replace(source, destination) + + with mock.patch("os.replace", side_effect=fail_commit_and_two_rollbacks): + result, output = self._run() + + self.assertEqual(result, 1) + self.assertEqual(len(calls), 7) + self.assertEqual(calls[-3:], [headers[1], headers[0], registry]) + self.assertIn(f"cannot roll back {headers[1]}", output) + self.assertIn(f"cannot roll back {headers[0]}", output) + self.assertEqual(registry.read_text(), "old registry\n") + self.assertNotEqual(headers[0].read_text(), "old a\n") + self.assertNotEqual(headers[1].read_text(), "old b\n") + self.assertEqual(headers[2].read_text(), "old c\n") + self._assert_no_temps() + + def test_success_preserves_existing_mode_and_creates_0644(self) -> None: + shader = self._write_shader("op", "op") + registry = g.registry_path() + registry.write_text("old registry\n") + registry.chmod(0o600) + + result, _ = self._run() + + self.assertEqual(result, 0) + self.assertEqual(stat.S_IMODE(registry.stat().st_mode), 0o600) + self.assertEqual( + stat.S_IMODE(shader.with_name("op_wgsl.h").stat().st_mode), 0o644 + ) + + def test_orphans_are_sorted_reported_and_never_deleted(self) -> None: + self._write_shader("new", "new") + orphan_z = self.root / "runtime/ops/z/old_z_wgsl.h" + orphan_a = self.root / "runtime/ops/a/old_a_wgsl.h" + orphan_z.parent.mkdir(parents=True) + orphan_a.parent.mkdir(parents=True) + orphan_z.write_text("// @generated\n") + orphan_a.write_text("// @generated\n") + before = self._snapshot() + + check_result, check_output = self._run("--check") + normal_result, normal_output = self._run() + + self.assertEqual(check_result, 1) + self.assertEqual(normal_result, 1) + for output in (check_output, normal_output): + self.assertIn("Orphan", output) + self.assertLess(output.index("old_a_wgsl.h"), output.index("old_z_wgsl.h")) + self.assertEqual(self._snapshot(), before) + + def test_check_fails_read_only_when_outputs_are_only_missing(self) -> None: + self._write_shader("op", "op") + before = self._snapshot() + + result, output = self._run("--check") + + self.assertEqual(result, 1) + self.assertIn("Missing embedded WGSL headers", output) + self.assertEqual(self._snapshot(), before) + + def test_check_catches_template_syntax_error_without_writing(self) -> None: + self._write_template( + "a", "syntax", "$if :\n " + self._VALID_SHADER, ["syntax"] + ) + before = self._snapshot() + + result, output = self._run("--check") + + self.assertEqual(result, 1) + self.assertIn("runtime/ops/a/syntax.wgsl", output) + self.assertEqual(self._snapshot(), before) + + def test_check_catches_template_name_error_without_writing(self) -> None: + self._write_template( + "op", "name", "$if MISSING:\n " + self._VALID_SHADER, ["name"] + ) + before = self._snapshot() + + result, output = self._run("--check") + + self.assertEqual(result, 1) + self.assertIn("runtime/ops/op/name.wgsl", output) + self.assertEqual(self._snapshot(), before) + + def test_interrupted_commit_is_detected_and_repaired(self) -> None: + self._write_shader("op", "op") + before = self._snapshot() + real_replace = os.replace + + def interrupt_second(source, destination): + interrupt_second.calls += 1 + if interrupt_second.calls == 2: + raise KeyboardInterrupt("injected interruption") + return real_replace(source, destination) + + interrupt_second.calls = 0 + with mock.patch("os.replace", side_effect=interrupt_second): + with self.assertRaises(KeyboardInterrupt): + self._run() + + self.assertEqual(self._snapshot(), before) + self._assert_no_temps() + check_result, check_output = self._run("--check") + self.assertEqual(check_result, 1) + self.assertNotIn("Orphan", check_output) + + normal_result, _ = self._run() + self.assertEqual(normal_result, 0) + final_check_result, _ = self._run("--check") + self.assertEqual(final_check_result, 0) + self._assert_no_temps() + + def test_interrupt_after_replace_restores_tree(self) -> None: + self._write_shader("op", "op") + before = self._snapshot() + real_replace = os.replace + + def interrupt_after_second(source, destination): + interrupt_after_second.calls += 1 + result = real_replace(source, destination) + if interrupt_after_second.calls == 2: + raise KeyboardInterrupt("injected post-replace interruption") + return result + + interrupt_after_second.calls = 0 + with mock.patch("os.replace", side_effect=interrupt_after_second): + with self.assertRaises(KeyboardInterrupt): + self._run() + + self.assertEqual(self._snapshot(), before) + self._assert_no_temps() + + def test_generation_renders_once_and_second_run_does_no_io(self) -> None: + shaders = [ + self._write_shader("a", "first"), + self._write_shader("b", "second"), + ] + render_counts = {shader: 0 for shader in shaders} + real_headers_for_shader = g.headers_for_shader + + def counted(shader): + render_counts[shader] += 1 + return real_headers_for_shader(shader) + + with mock.patch.object(g, "headers_for_shader", side_effect=counted): + first_result, _ = self._run() + self.assertEqual(first_result, 0) + self.assertEqual(render_counts, {shader: 1 for shader in shaders}) + + with mock.patch( + "tempfile.mkstemp", wraps=tempfile.mkstemp + ) as mkstemp, mock.patch("os.replace", wraps=os.replace) as replace: + second_result, _ = self._run() + self.assertEqual(second_result, 0) + mkstemp.assert_not_called() + replace.assert_not_called() + + class WgslTemplateEngineTest(unittest.TestCase): """Coverage for the $-block template engine + DTYPE/VEC variant matrix."""