Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
39 changes: 39 additions & 0 deletions tests/end_to_end/test_profiler.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
import importlib

import torch
import pytest

Expand Down Expand Up @@ -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
35 changes: 32 additions & 3 deletions tilelens/clients/profiler/profiler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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:
Expand Down Expand Up @@ -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
Expand All @@ -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, "
Expand Down
Loading