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
6 changes: 5 additions & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -238,7 +238,11 @@ def kernel(x_ptr, n, BLOCK: tl.constexpr):
`is_hip()`) answer for that target too, and `TRITON_OVERRIDE_ARCH` does not
apply: the target is the one named. The TTIR can differ between targets (e.g.
tensor descriptors, target-dependent branches), and a verdict holds for the
target it was checked for.
target it was checked for. To check a launch for several targets, stack one
sanitizer per target (`@tilelens.trace(Sanitizer(compile=True,
target="cuda:90"))` over `@tilelens.trace(Sanitizer(compile=True,
target="cuda:80"))`): each keeps its own verdict (`last_verdict`), and the
launch's records hold them in trace order.
- Targets also differ in what compiles at all: `fp8e4nv` (`torch.float8_e4m3fn`,
`tl.float8e4nv`) needs `cuda:89` or later, `num_ctas > 1` and 16-bit tensor
descriptor atomic min/max need `cuda:90`. A kernel or autotune config that fails
Expand Down
20 changes: 20 additions & 0 deletions tests/end_to_end/test_compiled_sanitizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -1577,6 +1577,26 @@ def test_target_dependent_branches_are_the_ir_targets(make, target, status):
assert _check_branches(make, target)[0] == status


def test_sanitizers_for_two_targets_share_one_trace():
"""A second compiled sanitizer, for another target, is not dropped from
the trace: one launch is checked for each target, each with its own
verdict."""
sm80 = Sanitizer(compile=True, abort_on_error=False, target="cuda:80")
sm90 = Sanitizer(compile=True, abort_on_error=False, target="cuda:90")
traced = tilelens.trace(sm90)(tilelens.trace(sm80)(_make_unmasked_from_sm89()))
assert traced.client_manager.ir_clients() == [sm80, sm90]

traced[(8,)](torch.zeros(64), 64, BLOCK=16)

assert (sm80.last_status, sm80.records) == ("ok", [])
assert sm90.last_status == "violations" and sm90.records
assert trace_module.launches[-1].records == [
sm80.last_verdict,
*sm90.records,
sm90.last_verdict,
]


class _Machine:
"""A stand-in for Triton's active driver on a machine with a GPU of
``target``; it must never be asked during IR mode's compile."""
Expand Down
24 changes: 13 additions & 11 deletions tests/end_to_end/test_core.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,19 +16,18 @@ def test_trace_decorator_add_clients():
Test goal:
1. Apply @trace("sanitizer") and @trace("profiler") to add the Sanitizer and Profiler clients.
2. Apply @trace("tracer") to append a Tracer client.
3. Apply @trace(("sanitizer",)) with a duplicate Sanitizer, which should be
ignored by the de-duplication logic.
3. Apply @trace("sanitizer") over a Sanitizer instance: the name asks for a
default Sanitizer, which the one already in the trace serves.

The final Trace object should contain exactly one instance each of
Sanitizer, Profiler, and Tracer (total = 3 clients).
Sanitizer, Profiler, and Tracer (total = 3 clients). A second Sanitizer
instance, whose settings would be lost, is refused.
"""

@tilelens.trace("sanitizer")
@tilelens.trace("profiler")
@tilelens.trace("tracer")
@tilelens.trace(
Sanitizer(abort_on_error=True)
) # Duplicate Sanitizer (should be ignored)
@tilelens.trace(Sanitizer(abort_on_error=False))
@triton.jit
def my_kernel(x_ptr, y_ptr, out_ptr, BLOCK_SIZE: tl.constexpr):
pid = tl.program_id(0)
Expand All @@ -41,11 +40,14 @@ def my_kernel(x_ptr, y_ptr, out_ptr, BLOCK_SIZE: tl.constexpr):
assert isinstance(my_kernel, TritonTrace)

# Verify client de-duplication and addition logic
clients = my_kernel.client_manager.clients
assert len(clients) == 3
assert sum(c == "sanitizer" for c in clients) == 1
assert sum(c == "profiler" for c in clients) == 1
assert sum(c == "tracer" for c in clients) == 1
names = [c.NAME for c in my_kernel.client_manager.clients]
assert sorted(names) == ["profiler", "sanitizer", "tracer"]
# The instance's own settings were kept, not a default's.
assert my_kernel.client_manager.get_client("sanitizer").abort_on_error is False

with pytest.raises(ValueError, match="interpreting client named 'sanitizer'"):
tilelens.trace(Sanitizer(abort_on_error=True))(my_kernel)
assert len(my_kernel.client_manager.clients) == 3


def test_trace_decorator_supports_gluon_frontend():
Expand Down
2 changes: 1 addition & 1 deletion tests/end_to_end/test_race_detector.py
Original file line number Diff line number Diff line change
Expand Up @@ -125,7 +125,7 @@ def test_string_dispatch_and_manager_lookup(_isolate_race_detector_cfg):
x = torch.zeros(8, dtype=torch.float32)
traced[(1,)](x, BLOCK=8)

rd = traced.client_manager.clients["race_detector"]
rd = traced.client_manager.get_client("race_detector")
assert isinstance(rd, SymbolicRaceDetector)
assert len(rd.records) == 2

Expand Down
2 changes: 1 addition & 1 deletion tests/unit/sanitizer_compiled/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -285,7 +285,7 @@ def test_declarations_and_composition():
# The eager sanitizer can share its trace (D4b); a client that needs the
# real launch cannot (D4a).
manager = ClientManager([san, SymbolicSanitizer(abort_on_error=False)])
assert set(manager.clients) == {"compiled_sanitizer", "sanitizer"}
assert [c.NAME for c in manager.clients] == ["compiled_sanitizer", "sanitizer"]
with pytest.raises(RuntimeError, match="disagree on whether the real kernel"):
ClientManager([san, _RunIR()])

Expand Down
80 changes: 63 additions & 17 deletions tests/unit/test_ir_lifecycle.py
Original file line number Diff line number Diff line change
Expand Up @@ -465,22 +465,24 @@ def test_add_clients_rejects_skip_run_conflict_before_inserting():
manager.add_clients([_IndifferentIRClient(), _RunIRClient()])

# Nothing from the rejected batch was inserted.
assert list(manager.clients) == ["ir_skip", "eager"]
assert [c.NAME for c in manager.clients] == ["ir_skip", "eager"]

with pytest.raises(RuntimeError, match="LAUNCH='skip'"):
ClientManager([_RunIRClient(), _SkipIRClient()])


def test_launch_conflict_check_sees_the_resulting_ir_clients():
# A same-NAME client replaces the one it would otherwise conflict with.
def test_launch_conflict_check_sees_every_ir_client():
# A same-NAME IR client joins the first rather than replacing it, so
# their conflict is seen.
class _RunInSkipSlot(_IRClient):
NAME = "ir_skip"
LAUNCH = "run"

manager = ClientManager([_SkipIRClient()])
manager.add_clients([_RunInSkipSlot()])
assert [type(c) for c in manager.clients.values()] == [_RunInSkipSlot]
assert manager.launch_policy() == "run"
first = _SkipIRClient()
manager = ClientManager([first])
with pytest.raises(RuntimeError, match="cannot share one trace"):
manager.add_clients([_RunInSkipSlot()])
assert manager.clients == [first]

# An interpreting client's LAUNCH takes no part in the vote.
class _EagerSkip(_EagerClient):
Expand All @@ -492,13 +494,34 @@ class _EagerSkip(_EagerClient):
assert manager.launch_policy() == "run"


def test_add_clients_keeps_duplicate_rule_and_accepts_indifferent():
first = _SkipIRClient()
manager = ClientManager([first, _IndifferentIRClient(), _EagerClient()])
manager.add_clients([_SkipIRClient()])
def test_add_clients_keeps_every_ir_client_instance():
# Adding a client already in the trace changes nothing; another
# instance of its class is kept, not dropped (e.g. one per target).
first, second = _SkipIRClient(), _SkipIRClient()
indifferent, eager = _IndifferentIRClient(), _EagerClient()
manager = ClientManager([first, indifferent, eager])
manager.add_clients([first, second, second])

assert list(manager.clients) == ["ir_skip", "ir_indifferent", "eager"]
assert manager.clients["ir_skip"] is first
assert manager.clients == [first, indifferent, eager, second]
assert manager.ir_clients() == [first, indifferent, second]
assert manager.get_client("ir_skip") is first


def test_add_clients_refuses_a_second_interpreting_client_of_one_name():
# One interpreted run serves one client per NAME: another one, of the
# same class or not, raises rather than being dropped, and nothing from
# its batch is inserted.
class _SameName(_SiblingEagerClient):
NAME = "eager"

first = _EagerClient()
manager = ClientManager([first])
for duplicate in (_EagerClient(), _SameName()):
with pytest.raises(ValueError, match="interpreting client named 'eager'"):
manager.add_clients([_IndifferentIRClient(), duplicate])
assert manager.clients == [first]
manager.add_clients([first])
assert manager.clients == [first]


def test_add_clients_rejects_unknown_launch_value():
Expand All @@ -516,7 +539,7 @@ def test_trace_decorator_rejects_conflicting_launch_preferences():
with pytest.raises(RuntimeError, match="cannot share one trace"):
tilelens.trace(_RunIRClient())(traced)

assert list(traced.client_manager.clients) == ["ir_skip"]
assert [c.NAME for c in traced.client_manager.clients] == ["ir_skip"]


def test_client_partition_and_launch_policy():
Expand Down Expand Up @@ -989,8 +1012,8 @@ def test_ir_capture_warmup_call_never_launches():
jit_fn.run(torch.zeros(4), 4, grid=None, warmup=True)

assert log == ["compile", "before", "after"]
assert manager.clients["ir_run"].events[0].launched is False
assert manager.clients["ir_run"].events[0].resolved_grid is None
assert manager.get_client("ir_run").events[0].launched is False
assert manager.get_client("ir_run").events[0].resolved_grid is None


def test_ir_capture_restores_on_error_and_does_not_double_wrap():
Expand Down Expand Up @@ -1286,6 +1309,29 @@ def test_each_target_compiles_once_and_reaches_only_its_clients(_fake_host_compi
assert [e.target for e in hip.events] == [cuda89]


def test_instances_of_one_ir_client_class_keep_their_own_targets(
_fake_host_compile,
):
"""Stacked traces of one IR client class, one instance per target, keep
both instances: each target compiles once and reaches only its own."""
from triton.backends.compiler import GPUTarget

cuda80, cuda90 = GPUTarget("cuda", 80, 32), GPUTarget("cuda", 90, 32)
sm80, sm90 = _SkipIRClient(), _SkipIRClient()
sm80.ir_target, sm90.ir_target = "cuda:80", "cuda:90"
traced = tilelens.trace(sm90)(tilelens.trace(sm80)(_make_plain_kernel()))
manager = traced.client_manager
assert manager.clients == [sm80, sm90]
jit_fn = _FakeJit(compile_error=None)

with manager.ir_capture(jit_fn, compile_only=True):
jit_fn.run(torch.zeros(4), 4, grid=(1,), warmup=True)

assert [t for _, t, _ in _fake_host_compile] == [cuda80, cuda90]
assert [e.target for e in sm80.events] == [cuda80]
assert [e.target for e in sm90.events] == [cuda90]


def test_the_configured_target_is_the_default(monkeypatch, _fake_host_compile):
from triton.backends.compiler import GPUTarget

Expand Down Expand Up @@ -2283,7 +2329,7 @@ def begin_launch(self, call):
# Only IR clients, one of them skipping: nothing is interpreted.
assert traced[(2,)](torch.zeros(4)) is None
assert log == ["begin", "finalize"]
(call,) = traced.client_manager.clients["ir_skip"].launch_calls
(call,) = traced.client_manager.get_client("ir_skip").launch_calls
assert call.jit_fn is None and call.capture is False and call.grid == (2,)

log = []
Expand Down
2 changes: 1 addition & 1 deletion tests/unit/test_wrapper.py
Original file line number Diff line number Diff line change
Expand Up @@ -295,7 +295,7 @@ def kernel(x_ptr):
pass


print(sorted(kernel.client_manager.clients), sys.argv[1:])
print(sorted(c.NAME for c in kernel.client_manager.clients), sys.argv[1:])
"""


Expand Down
68 changes: 46 additions & 22 deletions tilelens/core/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -143,6 +143,9 @@ class CompileGroup:


class Client(ABC):
# Names the client's records and its ClientManager.get_client lookup. A
# trace holds one interpreting client per NAME, since the interpreted
# run serves one; IR clients may repeat a NAME (see ir_target).
NAME: ClassVar[str]
# Whether the client consumes the interpreted run (op/loop callbacks,
# pre/post_run votes, arg/grid callbacks). IR clients set this to False
Expand All @@ -163,8 +166,9 @@ class Client(ABC):
# GPUTarget or a spec such as "cuda:90" or "hip:gfx942" (see
# tilelens.core.host_compile.parse_ir_target); None for the configured
# default (tilelens.config.ir_target: TILELENS_IR_TARGET, else
# "cuda:89"). A client may set it per instance. Core compiles once per
# distinct target and gives each client only its own target's events.
# "cuda:89"). A client may set it per instance, so one trace can hold an
# instance of a client class per target. Core compiles once per distinct
# target and gives each client only its own target's events.
ir_target: Any = None

def __init__(self) -> None:
Expand Down Expand Up @@ -486,7 +490,8 @@ def gate(*args, **kwargs):

class ClientManager:
def __init__(self, clients: list[Client] | None = None):
self.clients: dict[str, Client] = {}
# In trace order: the order clients are called and finalized in.
self.clients: list[Client] = []
if clients:
self.add_clients(clients)
self.launch = Launch()
Expand Down Expand Up @@ -528,23 +533,42 @@ def _lock_context(self):
return nullcontext()

def get_client(self, name: str) -> Client | None:
return self.clients.get(name)
"""The first client in trace order whose NAME is ``name``."""
return next((c for c in self.clients if c.NAME == name), None)

def add_clients(self, new_clients_list: list[Client]) -> None:
"""Append each client, in order; one already in the trace (the same
object) is skipped. Every other client is kept or the call raises:
nothing is dropped silently."""
# Validate the whole resulting set before inserting anything, so a
# rejected composition leaves the manager unchanged.
additions: dict[str, Client] = {}
resulting = list(self.clients)
for new_client in new_clients_list:
duplicate = any(
isinstance(existing_client, new_client.__class__)
for existing_client in (*self.clients.values(), *additions.values())
)
if not duplicate:
additions[new_client.NAME] = new_client
# A same-NAME addition replaces the existing client, so check the set
# that will result, not the one before replacement.
self._check_launch_preferences(list({**self.clients, **additions}.values()))
self.clients.update(additions)
if any(new_client is client for client in resulting):
continue
if new_client.NEEDS_INTERPRETER:
taken = next(
(
c
for c in resulting
if c.NEEDS_INTERPRETER and c.NAME == new_client.NAME
),
None,
)
if taken is not None:
raise ValueError(
"this trace already has an interpreting client named "
f"{new_client.NAME!r} ({type(taken).__name__}); one "
"interpreted run serves one client per name, so "
f"another ({type(new_client).__name__}) cannot share "
"the trace. Trace the kernel twice instead, e.g. "
"tilelens.trace(a)(kernel) and tilelens.trace(b)(kernel), "
"and launch each; stacked trace decorators merge into "
"one trace."
)
resulting.append(new_client)
self._check_launch_preferences(resulting)
self.clients = resulting

@staticmethod
def _check_launch_preferences(clients: list[Client]) -> None:
Expand All @@ -569,10 +593,10 @@ def _check_launch_preferences(clients: list[Client]) -> None:
)

def interpreting_clients(self) -> list[Client]:
return [c for c in self.clients.values() if c.NEEDS_INTERPRETER]
return [c for c in self.clients if c.NEEDS_INTERPRETER]

def ir_clients(self) -> list[Client]:
return [c for c in self.clients.values() if not c.NEEDS_INTERPRETER]
return [c for c in self.clients if not c.NEEDS_INTERPRETER]

def compile_groups(self) -> list[CompileGroup]:
"""The IR clients grouped by the target their kernels are compiled
Expand Down Expand Up @@ -626,7 +650,7 @@ def begin_launch(self, call: LaunchCall) -> None:
self._reset_launch_state()
begun: list[Client] = []
try:
for client in self.clients.values():
for client in self.clients:
begun.append(client)
client.begin_launch(call)
except BaseException as exc:
Expand Down Expand Up @@ -670,7 +694,7 @@ def abort_launch(self, exc: BaseException) -> None:
# Held until every client got the abort, so no other thread's
# begin_launch resets the state in between.
self._launch_owner = thread
self._abort_clients(list(self.clients.values()), exc)
self._abort_clients(list(self.clients), exc)

def _abort_clients(self, clients: list[Client], exc: BaseException) -> None:
interrupt: BaseException | None = None
Expand Down Expand Up @@ -747,7 +771,7 @@ def _warmup_by_vote(
# not short-circuit on the first True.
votes = [
client.pre_warmup_callback(jit_fn, *args, **kwargs)
for client in self.clients.values()
for client in self.clients
]
if not any(votes):
return None
Expand All @@ -757,7 +781,7 @@ def _warmup_by_vote(
args, kwargs = real_args(jit_fn, args, kwargs)
with compile_context():
ret = warmup(*args, **kwargs)
for client in self.clients.values():
for client in self.clients:
client.post_warmup_callback(jit_fn, ret)
return ret

Expand Down Expand Up @@ -1112,7 +1136,7 @@ def finalize(self) -> None:
# Finalize every client even if a peer raises (e.g. SystemExit
# from an abort), then re-raise the first failure.
first_exc: BaseException | None = None
for client in self.clients.values():
for client in self.clients:
try:
# client may introduce tensors not declared in kernel args (e.g. tracer recording a tensor allocation)
self.launch.tensors.update(getattr(client, "tensors", []) or [])
Expand Down
Loading