From dcfb33fda52551897ebb83ba37242de6e1fbf545 Mon Sep 17 00:00:00 2001 From: Joy Zhang Date: Wed, 9 Sep 2026 14:51:44 -0700 Subject: [PATCH] Describe a backend by what it can do, not by its name (#151) Summary: The generation path decided device behaviour with a single Jinja conditional: ``` {% if device_string == "xpu" %} if not hasattr(torch, 'xpu') or not torch.xpu.is_available(): {% else %} if not torch.cuda.is_available(): {% endif %} ``` An if/else over two known names is not an abstraction. It silently routes every third backend down the CUDA path -- generating CUDA availability checks under another backend's name -- and a backend defined outside this repository cannot add an arm to it. `PlatformConfig` now carries the five capabilities a generated harness actually needs: `availability_check`, `device_setup`, `synchronize_call`, `test_prelude` and `default_num_workers`. The template reads them; it no longer knows any device name. `register_platform()` is the seam that lets a backend defined elsewhere become selectable without editing anything here -- an explicit call rather than an import-time decorator, so registration order stays something the caller controls. Three smaller corrections fall out of the same change: - `DEFAULT_PLATFORM` was declared and then ignored in favour of hardcoded `"cuda"` literals. It is now the value actually used. - `TritonKernelAgent` passed the *raw* `target_platform` to `PromptManager` while passing the *normalized* one to `WorkerManager`, so two places independently re-derived the same default. The resolved config is now used for both, and it is resolved before the worker count, because the backend supplies that default. - The two `device='cuda'` literals in the mock test-generation fallback now follow the selected backend. A `fake` backend is registered alongside `cuda` and `xpu`: no accelerator, a check that asserts nothing because there is nothing to assert, an empty `synchronize_call`, and one worker. It is named for what it is, and its guidance block says outright that nothing it produces is a performance claim -- the same reasoning as the existing `noop` implementations in `triton_kernel_agent.platform`. Empty `synchronize_call` is a real answer, not a gap to be filled with the CUDA call. No Meta-internal import enters the generic tree; a test asserts that. Differential Revision: D118187263 --- README.md | 13 + tests/test_platform_capabilities.py | 386 ++++++++++++++++++ tests/test_platform_config.py | 20 +- triton_kernel_agent/agent.py | 40 +- triton_kernel_agent/platform_config.py | 155 ++++++- triton_kernel_agent/prompt_manager.py | 3 + .../templates/test_generation.j2 | 11 +- 7 files changed, 604 insertions(+), 24 deletions(-) create mode 100644 tests/test_platform_capabilities.py diff --git a/README.md b/README.md index 5e2eaa79..ddcc8d86 100644 --- a/README.md +++ b/README.md @@ -226,6 +226,19 @@ KernelAgent supports multiple GPU platforms for Triton kernel execution: |----------|---------------|------|--------| | NVIDIA CUDA | `cuda` | `--target-platform cuda` (default) | Fully supported | | Intel XPU | `xpu` | `--target-platform xpu` | Supported | +| Fake (no accelerator) | `cpu` | `--target-platform fake` | Dry runs and CI without hardware | + +A backend is described by the capabilities it declares — availability check, +device setup, synchronization, test prelude and default worker count — rather +than by its name, so nothing in the generic templates branches on a device +string. Register one with +`triton_kernel_agent.platform_config.register_platform`, and register its +optimization components with +`triton_kernel_agent.platform.registry.registry.register`. + +The `fake` backend exists to exercise the pipeline where there is no +accelerator. It reports a placeholder time so the pipeline can complete, and +nothing it produces is a performance claim. ### Intel XPU Notes diff --git a/tests/test_platform_capabilities.py b/tests/test_platform_capabilities.py new file mode 100644 index 00000000..944d99e2 --- /dev/null +++ b/tests/test_platform_capabilities.py @@ -0,0 +1,386 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Tests for backend capabilities in the generation path. + +Host-only: no accelerator, no torch, no network, no model call. Where a test +needs ``torch``, it injects a stub, which is what lets the emitted harness be +*executed* here rather than only inspected. + +The snapshot tests compare against the template as it was before capabilities +were introduced, rendered through the same Jinja environment. That is stronger +than a hand-written golden string: it cannot drift from the historical +behaviour it claims to preserve, and it does not depend on anybody reasoning +correctly about ``trim_blocks``. +""" + +from __future__ import annotations + +import sys +import tempfile +import types +import unittest +from pathlib import Path + +from jinja2 import Environment, FileSystemLoader + +from triton_kernel_agent.platform_config import ( + DEFAULT_NUM_WORKERS, + DEFAULT_PLATFORM, + get_platform, + get_platform_choices, + PLATFORMS, + PlatformConfig, + register_platform, + UnknownPlatformError, +) +from triton_kernel_agent.prompt_manager import PromptManager + +# The device block exactly as it stood before backend capabilities existed. +# Rendering this and the current template must produce identical text for every +# backend that existed then. +_LEGACY_DEVICE_BLOCK = """\ + # Device setup + device = "{{ device_string }}" +{% if device_string == "xpu" %} + if not hasattr(torch, 'xpu') or not torch.xpu.is_available(): + raise RuntimeError("Intel XPU not available. Install PyTorch with Intel GPU support.") +{% else %} + if not torch.cuda.is_available(): + raise RuntimeError("CUDA not available") +{% endif %} + + # Create test data +""" + +# The same region of the current template, kept in step with test_generation.j2. +_CURRENT_DEVICE_BLOCK = """\ + # Device setup + device = "{{ device_string }}" +{% if device_setup %}{{ device_setup | indent(8, true) }} +{% endif %} +{{ availability_check | indent(8, true) }} + + # Create test data +""" + +_PRE_CAPABILITY_PLATFORMS = ("cuda", "xpu") + + +def _env() -> Environment: + """A Jinja environment configured exactly as PromptManager configures its own.""" + return Environment(trim_blocks=True, lstrip_blocks=True) + + +def _render_block(source: str, platform: PlatformConfig) -> str: + """Render one template fragment for one backend.""" + return _env().from_string(source).render( + device_string=platform.device_string, + availability_check=platform.availability_check, + device_setup=platform.device_setup, + test_prelude=platform.test_prelude, + ) + + +def _stub_torch(*, cuda_available: bool = False, xpu_available: bool = False): + """A torch stand-in with just enough surface for an availability check.""" + torch = types.ModuleType("torch") + cuda = types.SimpleNamespace( + is_available=lambda: cuda_available, synchronize=lambda: None + ) + torch.cuda = cuda + if xpu_available: + torch.xpu = types.SimpleNamespace( + is_available=lambda: True, synchronize=lambda: None + ) + return torch + + +def _exec_check(platform: PlatformConfig, torch_stub) -> None: + """Execute a backend's availability check against a stubbed torch. + + Raises whatever the emitted code raises, which is the point: this proves the + check is valid Python *and* that it actually refuses an unavailable device. + """ + namespace = {"torch": torch_stub} + exec(compile(platform.availability_check, "", "exec"), namespace) + + +class SnapshotTest(unittest.TestCase): + """Existing backends must render byte-identically.""" + + def test_the_device_block_is_unchanged_for_every_pre_capability_backend( + self, + ) -> None: + for name in _PRE_CAPABILITY_PLATFORMS: + platform = get_platform(name) + legacy = _render_block(_LEGACY_DEVICE_BLOCK, platform) + current = _render_block(_CURRENT_DEVICE_BLOCK, platform) + self.assertEqual(current, legacy, f"{name} device block drifted") + + def test_the_snapshot_would_catch_a_drift(self) -> None: + # A snapshot that cannot fail proves nothing. A backend whose check + # differs from the legacy if/else must produce different text. + drifted = PlatformConfig( + name="drifted", + device_string="cuda", + guidance_block="", + kernel_guidance="", + availability_check="pass", + ) + self.assertNotEqual( + _render_block(_CURRENT_DEVICE_BLOCK, drifted), + _render_block(_LEGACY_DEVICE_BLOCK, drifted), + ) + + def test_the_full_prompt_still_contains_the_backends_check(self) -> None: + # Guards against the template fragment above drifting out of step with + # the real file: this renders test_generation.j2 itself. + for name in _PRE_CAPABILITY_PLATFORMS: + platform = get_platform(name) + prompt = PromptManager(target_platform=platform).render_test_generation_prompt( + "add two tensors" + ) + self.assertIn(platform.availability_check.splitlines()[0], prompt) + self.assertIn(f'device = "{platform.device_string}"', prompt) + + def test_the_template_no_longer_branches_on_a_device_name(self) -> None: + # The defect being removed: an if/else over two known names routes every + # third backend down the CUDA path without saying so. + source = (_templates_dir() / "test_generation.j2").read_text(encoding="utf-8") + self.assertNotIn('device_string == "xpu"', source) + self.assertNotIn("torch.cuda.is_available()", source) + + +def _templates_dir() -> Path: + """Locate the bundled templates directory.""" + import triton_kernel_agent + + return Path(triton_kernel_agent.__file__).parent / "templates" + + +class CapabilityTest(unittest.TestCase): + def test_every_registered_backend_declares_an_availability_check(self) -> None: + # A backend with no check emits an empty statement block, which is a + # syntax error inside the generated `try:`. + for name, platform in PLATFORMS.items(): + self.assertTrue(platform.availability_check.strip(), name) + + def test_every_availability_check_is_valid_python(self) -> None: + for name, platform in PLATFORMS.items(): + with self.subTest(platform=name): + compile(platform.availability_check, "", "exec") + + def test_an_accelerator_check_refuses_an_unavailable_device(self) -> None: + # Proves the emitted check is not vacuous. + with self.assertRaises(RuntimeError): + _exec_check(get_platform("cuda"), _stub_torch(cuda_available=False)) + with self.assertRaises(RuntimeError): + _exec_check(get_platform("xpu"), _stub_torch(cuda_available=True)) + + def test_an_accelerator_check_passes_when_the_device_is_there(self) -> None: + _exec_check(get_platform("cuda"), _stub_torch(cuda_available=True)) + _exec_check(get_platform("xpu"), _stub_torch(xpu_available=True)) + + def test_synchronization_is_declared_per_backend(self) -> None: + self.assertEqual(get_platform("cuda").synchronize_call, "torch.cuda.synchronize()") + self.assertEqual(get_platform("xpu").synchronize_call, "torch.xpu.synchronize()") + # Empty is a real answer for a backend with nothing to wait for, not a + # gap to be filled with the CUDA call. + self.assertEqual(get_platform("fake").synchronize_call, "") + + def test_concurrency_is_declared_per_backend(self) -> None: + for name in _PRE_CAPABILITY_PLATFORMS: + self.assertEqual( + get_platform(name).default_num_workers, DEFAULT_NUM_WORKERS, name + ) + self.assertEqual(get_platform("fake").default_num_workers, 1) + + +class FakeBackendTest(unittest.TestCase): + """The fake backend must drive the pipeline with no accelerator present.""" + + def test_it_is_registered_and_selectable(self) -> None: + self.assertIn("fake", get_platform_choices()) + self.assertEqual(get_platform("fake").device_string, "cpu") + + def test_its_check_executes_with_no_device_and_no_torch_calls(self) -> None: + # The stub reports nothing available; the fake backend must still pass. + _exec_check(get_platform("fake"), _stub_torch(cuda_available=False)) + + def test_it_generates_a_harness_that_runs_on_the_host(self) -> None: + # Take the emitted device setup and availability check, assemble the + # harness the template describes, and actually execute it. + platform = get_platform("fake") + rendered = _render_block(_CURRENT_DEVICE_BLOCK, platform) + # The template renders into a `def test_kernel(): try:` body, so the + # emitted block sits at eight spaces. Reproduce that exactly. + # Jinja drops the template's final newline, so rejoin explicitly or the + # last rendered line swallows whatever is appended to it. + harness = ( + "def test_kernel():\n" + " try:\n" + + rendered.rstrip("\n") + + "\n return device\n" + " except Exception:\n" + " raise\n" + ) + namespace = {"torch": _stub_torch()} + exec(compile(harness, "", "exec"), namespace) + self.assertEqual(namespace["test_kernel"](), "cpu") + + def test_its_guidance_says_its_numbers_are_not_measurements(self) -> None: + # Naming matters here: a backend that fabricates timings must not be + # mistaken for one that measures them. + guidance = get_platform("fake").guidance_block + self.assertIn("no accelerator", guidance) + self.assertIn("is a performance claim", guidance) + self.assertIn("Nothing generated for it", guidance) + + def test_the_full_prompt_renders_for_it(self) -> None: + prompt = PromptManager( + target_platform=get_platform("fake") + ).render_test_generation_prompt("add two tensors") + self.assertIn('device = "cpu"', prompt) + self.assertNotIn("torch.cuda.is_available()", prompt) + + +class RegistrationTest(unittest.TestCase): + def setUp(self) -> None: + self._saved = dict(PLATFORMS) + + def tearDown(self) -> None: + PLATFORMS.clear() + PLATFORMS.update(self._saved) + + def _config(self, name: str = "outside") -> PlatformConfig: + return PlatformConfig( + name=name, + device_string=name, + guidance_block="", + kernel_guidance="", + availability_check="pass", + ) + + def test_a_backend_defined_elsewhere_becomes_selectable(self) -> None: + # The seam that lets a backend outside this repository be used without + # editing anything in it. + register_platform(self._config()) + self.assertIn("outside", get_platform_choices()) + self.assertIs(get_platform("outside"), PLATFORMS["outside"]) + + def test_registering_over_an_existing_backend_needs_saying_so(self) -> None: + with self.assertRaises(ValueError) as ctx: + register_platform(self._config("cuda")) + self.assertIn("already registered", str(ctx.exception)) + register_platform(self._config("cuda"), replace=True) + self.assertEqual(get_platform("cuda").device_string, "cuda") + + def test_a_nameless_backend_is_refused(self) -> None: + with self.assertRaises(ValueError): + register_platform(self._config(" ")) + + def test_an_unknown_backend_raises_rather_than_defaulting_to_cuda(self) -> None: + # Returning CUDA for an unrecognized name would generate CUDA code under + # another backend's name. + with self.assertRaises(UnknownPlatformError) as ctx: + get_platform("nonexistent") + self.assertIn("Available:", str(ctx.exception)) + # Still a ValueError, so existing callers keep working. + self.assertIsInstance(ctx.exception, ValueError) + + def test_the_default_platform_constant_is_the_one_actually_used(self) -> None: + # DEFAULT_PLATFORM used to be declared and then ignored in favour of + # hardcoded "cuda" literals. + self.assertIn(DEFAULT_PLATFORM, PLATFORMS) + import triton_kernel_agent.agent as agent_module + + source = Path(agent_module.__file__).read_text(encoding="utf-8") + self.assertNotIn('get_platform("cuda")', source) + + +class NoInternalImportsTest(unittest.TestCase): + def test_the_generic_tree_imports_nothing_meta_internal(self) -> None: + # The exported package must stay installable outside fbsource. + import triton_kernel_agent + + root = Path(triton_kernel_agent.__file__).parent + offenders = [] + for path in sorted(root.rglob("*.py")): + text = path.read_text(encoding="utf-8") + for marker in ("kernelagent.fb", "from fb.", "import fb.", "libfb"): + if marker in text: + offenders.append(f"{path.name}: {marker}") + self.assertEqual(offenders, []) + + def test_platform_config_needs_only_the_standard_library(self) -> None: + # It is imported by every entry point, including ones with no jinja2. + import triton_kernel_agent.platform_config as module + + source = Path(module.__file__).read_text(encoding="utf-8") + for banned in ("import torch", "import jinja2", "import numpy"): + self.assertNotIn(banned, source) + + +class ExecutedHarnessIsolationTest(unittest.TestCase): + def test_executing_a_harness_does_not_import_real_torch(self) -> None: + # Guards the host-only claim: if these tests ever start importing torch + # for real, they stop being runnable on a machine without it. + before = "torch" in sys.modules + with tempfile.TemporaryDirectory(): + _exec_check(get_platform("fake"), _stub_torch()) + self.assertEqual("torch" in sys.modules, before) + + +class RegistryConsistencyTest(unittest.TestCase): + """Mirrors the invariants in the pytest-only test_platform_config.py. + + That file cannot be collected by python_unittest, so the facts it asserts + are re-asserted here where they are actually run. + """ + + def test_every_config_name_matches_its_registry_key(self) -> None: + for key, config in PLATFORMS.items(): + self.assertEqual(config.name, key) + + def test_every_registered_backend_is_reachable_and_named(self) -> None: + for name in get_platform_choices(): + config = get_platform(name) + self.assertEqual(config.name, name) + # Not an allow-list of device strings: the registry is extensible. + self.assertTrue(config.device_string) + + def test_choices_match_the_registry_and_are_sorted(self) -> None: + choices = get_platform_choices() + self.assertEqual(set(choices), set(PLATFORMS)) + self.assertEqual(choices, sorted(choices)) + + def test_every_backend_has_the_full_capability_surface(self) -> None: + for name in get_platform_choices(): + config = get_platform(name) + with self.subTest(platform=name): + self.assertIsInstance(config.guidance_block, str) + self.assertIsInstance(config.kernel_guidance, str) + self.assertIsInstance(config.cuda_hacks_to_strip, tuple) + self.assertIsInstance(config.availability_check, str) + self.assertIsInstance(config.device_setup, str) + self.assertIsInstance(config.synchronize_call, str) + self.assertIsInstance(config.test_prelude, str) + self.assertIsInstance(config.default_num_workers, int) + self.assertGreater(config.default_num_workers, 0) + + def test_a_blank_or_padded_name_is_refused(self) -> None: + for bad in ("", " ", " cuda "): + with self.assertRaises(ValueError): + get_platform(bad) diff --git a/tests/test_platform_config.py b/tests/test_platform_config.py index f795589d..b4e12969 100644 --- a/tests/test_platform_config.py +++ b/tests/test_platform_config.py @@ -151,7 +151,11 @@ def test_all_platforms_accessible(self): for name in get_platform_choices(): config = get_platform(name) assert config.name == name - assert config.device_string in ["cuda", "xpu"] + # Not an allow-list of known device strings: the registry is + # extensible, and a backend registered from outside this repository + # would fail a hardcoded list without being wrong. + assert isinstance(config.device_string, str) + assert config.device_string class TestEdgeCases: @@ -178,7 +182,7 @@ def test_platform_with_extra_whitespace_raises(self): get_platform(" cuda ") -@pytest.mark.parametrize("platform_name", ["cuda", "xpu"]) +@pytest.mark.parametrize("platform_name", get_platform_choices()) def test_all_platforms_have_consistent_structure(platform_name): """All platforms should have consistent field types.""" config = get_platform(platform_name) @@ -187,11 +191,21 @@ def test_all_platforms_have_consistent_structure(platform_name): assert isinstance(config.guidance_block, str) assert isinstance(config.kernel_guidance, str) assert isinstance(config.cuda_hacks_to_strip, tuple) + assert isinstance(config.availability_check, str) + assert isinstance(config.device_setup, str) + assert isinstance(config.synchronize_call, str) + assert isinstance(config.test_prelude, str) + assert isinstance(config.default_num_workers, int) @pytest.mark.parametrize("platform_name", ["cuda", "xpu"]) def test_platform_name_equals_device_string(platform_name): - """Platform name should equal device string for simplicity.""" + """Accelerator backends name themselves after their device. + + Deliberately scoped to the accelerator backends rather than the whole + registry: the `fake` backend is named for what it is and runs on `cpu`, so + the two differ there on purpose. + """ config = get_platform(platform_name) assert config.name == config.device_string diff --git a/triton_kernel_agent/agent.py b/triton_kernel_agent/agent.py index 84bd5d54..4d907600 100644 --- a/triton_kernel_agent/agent.py +++ b/triton_kernel_agent/agent.py @@ -26,7 +26,11 @@ from .manager import WorkerManager from .prompt_manager import PromptManager from utils.providers import BaseProvider, get_model_provider -from triton_kernel_agent.platform_config import PlatformConfig, get_platform +from triton_kernel_agent.platform_config import ( + DEFAULT_PLATFORM, + get_platform, + PlatformConfig, +) from triton_kernel_agent.worker_util import format_test_code_for_llm @@ -60,8 +64,18 @@ def __init__( # Load environment variables load_dotenv() - # Load configuration from environment - self.num_workers = num_workers or int(os.getenv("NUM_KERNEL_SEEDS", "4")) + # Resolve the backend first: it supplies the default worker count, so it + # has to exist before the concurrency decision is made. + self._platform_config = ( + target_platform if target_platform else get_platform(DEFAULT_PLATFORM) + ) + + # Load configuration from environment. Precedence is caller, then + # environment, then the backend's own default -- a backend with no + # accelerator behind it has no reason to fan out four ways. + self.num_workers = num_workers or int( + os.getenv("NUM_KERNEL_SEEDS", str(self._platform_config.default_num_workers)) + ) self.max_rounds = max_rounds or int(os.getenv("MAX_REFINEMENT_ROUNDS", "10")) self.model_name = model_name or os.getenv( "OPENAI_MODEL", "claude-sonnet-4-20250514" @@ -87,18 +101,16 @@ def __init__( self.log_dir = Path.cwd() / "triton_kernel_logs" self.log_dir.mkdir(exist_ok=True, parents=True) - # Normalize to PlatformConfig - self._platform_config = ( - target_platform if target_platform else get_platform("cuda") - ) self.no_cusolver = no_cusolver self.test_timeout_s = test_timeout_s # Setup main logger self._setup_logging() - # Initialize prompt manager - self.prompt_manager = PromptManager(target_platform=target_platform) + # Initialize prompt manager with the resolved config, not the raw + # argument: passing None here made PromptManager re-derive the default + # independently, so two places decided the same thing. + self.prompt_manager = PromptManager(target_platform=self._platform_config) # Initialize worker manager self.manager = WorkerManager( @@ -281,7 +293,7 @@ def test_kernel(): # Adapted from provided test code try: # Create test data (standardized format) - test_input = torch.randn(1024, device='cuda') + test_input = torch.randn(1024, device='__DEVICE__') # Call kernel_function as a normal Python function result = kernel_function(test_input) @@ -302,6 +314,9 @@ def test_kernel(): success = test_kernel() sys.exit(0 if success else 1) ''' + test_code = test_code.replace( + "__DEVICE__", self._platform_config.device_string + ) else: test_code = '''""" Test for kernel implementation. @@ -315,7 +330,7 @@ def test_kernel(): # Mock test - replace with actual test logic try: # Create test data - test_input = torch.randn(1024, device='cuda') + test_input = torch.randn(1024, device='__DEVICE__') # Call kernel_function as a normal Python function # (kernel launch logic is handled inside kernel.py) @@ -332,6 +347,9 @@ def test_kernel(): success = test_kernel() sys.exit(0 if success else 1) ''' + test_code = test_code.replace( + "__DEVICE__", self._platform_config.device_string + ) return test_code def _generate_kernel_seeds( diff --git a/triton_kernel_agent/platform_config.py b/triton_kernel_agent/platform_config.py index f69c6b5e..4a33e6b0 100644 --- a/triton_kernel_agent/platform_config.py +++ b/triton_kernel_agent/platform_config.py @@ -15,28 +15,83 @@ """ Platform configuration registry for multi-backend support. +A backend is described by what it *can do*, not by its name. Templates and +callers read capabilities off a :class:`PlatformConfig`; nothing branches on +``if device == "xpu"``. That matters because an ``if/else`` on two known names +is not an abstraction --- it silently routes every third backend down the CUDA +path --- and because a backend that lives outside this repository cannot add an +arm to an ``if``, but it can register a config. + +The five capabilities a generated harness needs are: + +``availability_check`` + Code that raises if the device cannot be used. Emitted verbatim into the + generated test. +``device_setup`` + Any extra setup the backend needs after the device string is bound. +``synchronize_call`` + How to wait for the device, or empty when there is nothing to wait for. +``test_prelude`` + Imports or statements a generated test needs before anything else. +``default_num_workers`` + How many generation workers this backend can usefully run at once. + Usage: from triton_kernel_agent.platform_config import get_platform, get_platform_choices platform = get_platform("xpu") print(platform.device_string) # "xpu" print(platform.guidance_block) # Intel XPU-specific guidance + +Registering a backend from outside this module: + from triton_kernel_agent.platform_config import PlatformConfig, register_platform + + register_platform(PlatformConfig(name="mybackend", device_string="mybackend", ...)) """ from dataclasses import dataclass, field DEFAULT_PLATFORM = "cuda" +# How many generation workers a backend runs by default. Kept as the historical +# value so registering the capability changes no existing behaviour. +DEFAULT_NUM_WORKERS = 4 + @dataclass(frozen=True) class PlatformConfig: - """Configuration for a specific hardware platform/backend.""" + """Configuration for a specific hardware platform/backend. + + Attributes: + name: Registry key, and the value a CLI accepts. + device_string: What ``torch`` calls this device. + guidance_block: Platform requirements injected into the test prompt. + kernel_guidance: Platform optimization notes injected into the kernel + prompt. + cuda_hacks_to_strip: Literal snippets to remove from model output that + tried to force a CUDA path. + availability_check: Code raising if the device is unusable. Rendered + verbatim into the generated test, so it must be valid Python at zero + indentation; the template indents it. + device_setup: Extra setup emitted after the device string is bound. + Empty for backends that need none. + synchronize_call: Expression that waits for the device, or empty when + the backend is synchronous. Empty is a real answer, not a gap. + test_prelude: Statements a generated test needs before anything else. + default_num_workers: Generation workers to run concurrently when the + caller and the environment do not say. + """ name: str device_string: str guidance_block: str kernel_guidance: str cuda_hacks_to_strip: tuple = field(default_factory=tuple) + availability_check: str = "" + device_setup: str = "" + synchronize_call: str = "" + test_prelude: str = "" + default_num_workers: int = DEFAULT_NUM_WORKERS # Platform-specific constants @@ -76,6 +131,31 @@ class PlatformConfig: "XPUDriver.is_available = classmethod(lambda cls: False)", ) +# Availability checks. These are emitted verbatim into the generated test and +# are byte-for-byte what the `{% if device_string == "xpu" %}` branch used to +# render, so moving them out of the template changes no generated output. +_CUDA_AVAILABILITY = """\ +if not torch.cuda.is_available(): + raise RuntimeError("CUDA not available")""" + +_XPU_AVAILABILITY = """\ +if not hasattr(torch, 'xpu') or not torch.xpu.is_available(): + raise RuntimeError("Intel XPU not available. Install PyTorch with Intel GPU support.")""" + +# The fake backend asserts nothing, because there is nothing to assert: the +# check has to remain a statement so the emitted block is never empty. +_FAKE_AVAILABILITY = """\ +# The fake backend has no device; nothing to check. +pass""" + +_FAKE_GUIDANCE = """\ +**PLATFORM: FAKE BACKEND (no accelerator).** +- This backend exists to exercise the generation pipeline without hardware. +- Nothing generated for it is a performance claim, and no timing it reports is + meaningful. +- Allocate on device='cpu' and do not call any accelerator API.""" + + # Platform registry PLATFORMS: dict[str, PlatformConfig] = { "cuda": PlatformConfig( @@ -84,6 +164,8 @@ class PlatformConfig: guidance_block="", kernel_guidance="", cuda_hacks_to_strip=(), + availability_check=_CUDA_AVAILABILITY, + synchronize_call="torch.cuda.synchronize()", ), "xpu": PlatformConfig( name="xpu", @@ -91,15 +173,82 @@ class PlatformConfig: guidance_block=_XPU_GUIDANCE, kernel_guidance=_XPU_KERNEL_GUIDANCE, cuda_hacks_to_strip=_XPU_CUDA_HACKS, + availability_check=_XPU_AVAILABILITY, + synchronize_call="torch.xpu.synchronize()", + ), + # A backend with no accelerator behind it, for exercising the pipeline on a + # host with no device. Mirrors the `noop` implementations in + # `triton_kernel_agent.platform`, and is named for what it is so nothing + # downstream mistakes its output for a measurement. + "fake": PlatformConfig( + name="fake", + device_string="cpu", + guidance_block=_FAKE_GUIDANCE, + kernel_guidance="", + cuda_hacks_to_strip=(), + availability_check=_FAKE_AVAILABILITY, + # Nothing to synchronize. Empty is the answer, not a missing value. + synchronize_call="", + default_num_workers=1, ), } +class UnknownPlatformError(ValueError): + """Raised for a platform name that is not registered. + + A subclass rather than a bare ``ValueError`` so a caller can distinguish + "that backend does not exist" from any other bad argument, while remaining + backward compatible with code that catches ``ValueError``. + """ + + +def register_platform(config: PlatformConfig, *, replace: bool = False) -> None: + """Register a backend. + + This is the seam that lets a backend defined outside this repository become + selectable without editing anything here. It is an explicit call rather than + an import-time decorator so that registration order is something a caller + controls and can reason about. + + Args: + config: The backend to add. + replace: Whether to overwrite an existing registration. + + Raises: + ValueError: If the name is blank, or is already registered and + ``replace`` is not set. Silently overwriting would let one import + change another backend's behaviour invisibly. + """ + if not config.name.strip(): + raise ValueError("a platform must have a name") + if config.name in PLATFORMS and not replace: + raise ValueError( + f"platform {config.name!r} is already registered; pass replace=True " + "to override it deliberately" + ) + PLATFORMS[config.name] = config + + def get_platform(name: str) -> PlatformConfig: - """Get platform configuration by name.""" + """Get platform configuration by name. + + Args: + name: A registered platform name. + + Returns: + The registered configuration. + + Raises: + UnknownPlatformError: If the name is not registered. There is no default + fallback: silently returning CUDA for an unrecognized backend would + generate CUDA code under another backend's name. + """ if name not in PLATFORMS: available = ", ".join(sorted(PLATFORMS.keys())) - raise ValueError(f"Unknown platform '{name}'. Available: {available}") + raise UnknownPlatformError( + f"Unknown platform '{name}'. Available: {available}" + ) return PLATFORMS[name] diff --git a/triton_kernel_agent/prompt_manager.py b/triton_kernel_agent/prompt_manager.py index 208d6cd0..e6dc1324 100644 --- a/triton_kernel_agent/prompt_manager.py +++ b/triton_kernel_agent/prompt_manager.py @@ -161,6 +161,9 @@ def render_test_generation_prompt( problem_description=problem_description, provided_test_code=provided_test_code, device_string=self.target_platform.device_string, + availability_check=self.target_platform.availability_check, + device_setup=self.target_platform.device_setup, + test_prelude=self.target_platform.test_prelude, ) def render_kernel_generation_prompt( diff --git a/triton_kernel_agent/templates/test_generation.j2 b/triton_kernel_agent/templates/test_generation.j2 index 61c9b2e2..7338d7e0 100644 --- a/triton_kernel_agent/templates/test_generation.j2 +++ b/triton_kernel_agent/templates/test_generation.j2 @@ -104,7 +104,8 @@ The test file structure should be: ```python import torch # other imports as needed (but NOT triton - that's handled in kernel.py) - +{% if test_prelude %}{{ test_prelude }} +{% endif %} # Comment for summary of original problem description def test_kernel(): """Test the kernel implementation.""" @@ -117,13 +118,9 @@ def test_kernel(): # Device setup device = "{{ device_string }}" -{% if device_string == "xpu" %} - if not hasattr(torch, 'xpu') or not torch.xpu.is_available(): - raise RuntimeError("Intel XPU not available. Install PyTorch with Intel GPU support.") -{% else %} - if not torch.cuda.is_available(): - raise RuntimeError("CUDA not available") +{% if device_setup %}{{ device_setup | indent(8, true) }} {% endif %} +{{ availability_check | indent(8, true) }} # Create test data using EXACT specifications from problem description # If problem specifies shape/dtype, use those exact values