diff --git a/tests/end_to_end/test_profiler.py b/tests/end_to_end/test_profiler.py index 04726f169..e7b09dbab 100644 --- a/tests/end_to_end/test_profiler.py +++ b/tests/end_to_end/test_profiler.py @@ -1,3 +1,5 @@ +import importlib + import torch import pytest @@ -421,3 +423,40 @@ def test_load_store_skip_enabled(_isolate_profiler_cfg): assert torch.allclose( y, torch.zeros_like(y) ), f"Skipped execution should leave output unchanged, got {y[:10]}" + + +# ======== Buffer Load Check ======== +def test_buffer_load_check_without_active_driver( + _isolate_profiler_cfg, monkeypatch, capsys +): + """The default buffer load check must not fail launches on driverless hosts.""" + # ``triton.runtime.driver`` resolves to the DriverConfig object, not the module. + triton_driver = importlib.import_module("triton.runtime.driver") + + def _no_driver(): + raise RuntimeError("0 active drivers ([]). There should only be one.") + + # Simulate a CPU-only host even when a GPU driver is available. + monkeypatch.setattr(triton_driver, "_create_driver", _no_driver) + monkeypatch.setattr(triton_driver.driver, "_default", None) + monkeypatch.setattr(triton_driver.driver, "_active", None) + + N = 128 + BLOCK_SIZE = 32 + x = torch.ones(N, dtype=torch.float32) * 3.0 + y = torch.zeros(N, dtype=torch.float32) + + cfg.profiler_enable_load_store_skipping = False + cfg.profiler_enable_block_sampling = False + cfg.profiler_disable_buffer_load_check = False + + profiler = Profiler() + traced_kernel = tilelens.trace(profiler)(simple_kernel) + grid = (triton.cdiv(N, BLOCK_SIZE),) + traced_kernel[grid](x, y, N, BLOCK_SIZE) + + assert torch.allclose(y, torch.ones_like(y) * 6.0) + assert profiler.buffer_load_check_skipped_no_driver is True + assert profiler.disable_buffer_load_check is True + assert profiler.potential_buffer_load_issue_found is False + assert "Skipped: no active Triton driver" in capsys.readouterr().out diff --git a/tilelens/clients/profiler/profiler.py b/tilelens/clients/profiler/profiler.py index e95723549..b55a83557 100644 --- a/tilelens/clients/profiler/profiler.py +++ b/tilelens/clients/profiler/profiler.py @@ -10,11 +10,22 @@ extract_user_frames, extract_complete_statement_from_line, ) +from triton.runtime.driver import driver from triton.runtime.interpreter import _get_np_dtype, TensorHandle import numpy as np from dataclasses import dataclass, replace +def _has_active_driver() -> bool: + """Return whether Triton has a backend driver that can compile kernels.""" + try: + driver.active + except RuntimeError: + # e.g. CPU-only hosts: "0 active drivers ([]). There should only be one." + return False + return True + + @dataclass(frozen=False) class LoopInfo: length: int | None = None @@ -81,6 +92,7 @@ def __init__( # Case 4: Buffer Load Check self.has_buffer_load = False self.disable_buffer_load_check = cfg.profiler_disable_buffer_load_check + self.buffer_load_check_skipped_no_driver = False self.potential_buffer_load_issue_found = False # Block sampling @@ -104,7 +116,16 @@ def post_run_callback(self, fn: Callable) -> bool: def pre_warmup_callback(self, jit_fn, *args, **kwargs) -> bool: # Skip warmup if buffer load check is disabled - return not self.disable_buffer_load_check + if self.disable_buffer_load_check: + return False + # The buffer load check inspects compiled ASM; without a backend driver + # Triton cannot compile the kernel, so drop the check instead of failing + # the launch. + if not _has_active_driver(): + self.disable_buffer_load_check = True + self.buffer_load_check_skipped_no_driver = True + return False + return True def post_warmup_callback(self, jit_fn, ret) -> None: if not ret: @@ -506,7 +527,10 @@ def finalize(self) -> list: print("\n" + "=" * 60 + "\n") - if not self.disable_buffer_load_check: + if ( + not self.disable_buffer_load_check + or self.buffer_load_check_skipped_no_driver + ): print("\n" + "=" * 60) print( "-" * 10 @@ -516,7 +540,12 @@ def finalize(self) -> list: + "-" * 11 ) print("=" * 60) - if self.potential_buffer_load_issue_found: + if self.buffer_load_check_skipped_no_driver: + print( + "Skipped: no active Triton driver to compile the kernel " + "(e.g. CPU-only host)." + ) + elif self.potential_buffer_load_issue_found: print("\n>>>>>> Warning: Potential Buffer Load Issue Detected! <<<<<<") print( "\nSome memory access offsets are within 32-bit range, "