Skip to content
Merged
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
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,7 @@ classifiers = [
]
dependencies = [
"setuptools",
"triton>=3.6.0",
"triton>=3.8.0",
"pyarrow",
"pre-commit",
"z3-solver==4.15.3.0",
Expand Down
41 changes: 5 additions & 36 deletions tests/end_to_end/test_gluon.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,31 +6,23 @@
from triton import knobs
from triton.experimental import gluon
from triton.experimental.gluon import language as gl
from triton.experimental.gluon.language.amd.cdna4 import async_copy as amd_cdna4_cp
from triton.experimental.gluon.language.nvidia import blackwell, hopper
from triton.experimental.gluon.language.nvidia.blackwell import tma as blackwell_tma
from triton.experimental.gluon.language.nvidia.hopper import (
mbarrier,
tma,
)
from triton.experimental.gluon.nvidia.hopper import TensorDescriptor
from triton.experimental.gluon.nvidia.hopper import (
TensorDescriptor,
TensorDescriptorIm2Col,
)

import tilelens
from tilelens.clients.sanitizer.sanitizer import SymbolicSanitizer
from tilelens.core.data import Load
from tilelens.core.simulation.gluon import GluonInterpretedFunction, gluon_builder

try:
from triton.experimental.gluon.language.amd.cdna4 import (
async_copy as amd_cdna4_cp,
)
except ImportError:
amd_cdna4_cp = None

try:
from triton.experimental.gluon.nvidia.hopper import TensorDescriptorIm2Col
except ImportError:
TensorDescriptorIm2Col = None

try:
from triton.experimental.gluon.language.nvidia.ampere import async_copy as cp
except ImportError:
Expand All @@ -39,11 +31,6 @@
_HAS_AMPERE_ASYNC_COPY = (
cp is not None and getattr(cp, "async_copy_global_to_local", None) is not None
)
_HAS_AMD_CDNA4_ASYNC_COPY = (
amd_cdna4_cp is not None
and getattr(amd_cdna4_cp, "global_load_to_shared", None) is not None
)
_HAS_TMA_IM2COL = TensorDescriptorIm2Col is not None


def _run_gluon_on_cpu(fn, grid, *args, **kwargs):
Expand Down Expand Up @@ -942,10 +929,6 @@ def test_gluon_async_copy_runs_masked_1d_copy_on_cpu():
torch.testing.assert_close(out, inp, atol=0, rtol=0)


@pytest.mark.skipif(
not _HAS_AMD_CDNA4_ASYNC_COPY,
reason="Gluon AMD CDNA4 async copy builtins are unavailable",
)
def test_gluon_amd_async_copy_preserves_masked_other_on_cpu():
inp = torch.arange(40, dtype=torch.float32)
out = torch.full((64,), -1, dtype=torch.float32)
Expand Down Expand Up @@ -1058,10 +1041,6 @@ def test_gluon_tma_runs_float_atomics_on_cpu():
torch.testing.assert_close(max_dst, expected_max, atol=0, rtol=0)


@pytest.mark.skipif(
not _HAS_TMA_IM2COL,
reason="Gluon TensorDescriptorIm2Col is unavailable in this Triton build",
)
def test_gluon_tma_im2col_runs_simple_tile_on_cpu():
inp = torch.arange(1, 17, dtype=torch.float32).unsqueeze(1).repeat(1, 32)
inp = inp.reshape(1, 4, 4, 32)
Expand All @@ -1071,10 +1050,6 @@ def test_gluon_tma_im2col_runs_simple_tile_on_cpu():
torch.testing.assert_close(out, inp.reshape(16, 32), atol=0, rtol=0)


@pytest.mark.skipif(
not _HAS_TMA_IM2COL,
reason="Gluon TensorDescriptorIm2Col is unavailable in this Triton build",
)
def test_gluon_tma_im2col_zero_fills_padded_pixels_on_cpu():
inp = torch.arange(1, 17, dtype=torch.float32).unsqueeze(1).repeat(1, 32)
inp = inp.reshape(1, 4, 4, 32)
Expand All @@ -1088,10 +1063,6 @@ def test_gluon_tma_im2col_zero_fills_padded_pixels_on_cpu():
torch.testing.assert_close(out[:, 0], expected_first_channel, atol=0, rtol=0)


@pytest.mark.skipif(
not _HAS_TMA_IM2COL,
reason="Gluon TensorDescriptorIm2Col is unavailable in this Triton build",
)
def test_gluon_tma_im2col_honors_runtime_offsets_on_cpu():
inp = torch.arange(1, 17, dtype=torch.float32).unsqueeze(1).repeat(1, 32)
inp = inp.reshape(1, 4, 4, 32)
Expand Down Expand Up @@ -1125,8 +1096,6 @@ def test_gluon_builder_preserves_tensor_memory_fp4_padding():
False,
True,
)
if not hasattr(layout, "fp4_padded"):
pytest.skip("Gluon TensorMemoryLayout has no fp4_padded field")
assert layout.fp4_padded is True


Expand Down
6 changes: 0 additions & 6 deletions tests/unit/test_patch_scope.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,3 @@
import pytest
import warnings
from types import SimpleNamespace

Expand Down Expand Up @@ -85,13 +84,8 @@ def test_scope_removes_tensor_magic_methods_added_by_interpreter():

def test_scope_restores_tensor_descriptor_base_builtins(monkeypatch):
"""Scope must restore descriptor builtins on tensor_descriptor_base."""
if not hasattr(tl.core, "tensor_descriptor_base"):
pytest.skip("tensor_descriptor_base is not available")

descriptor = tl.core.tensor_descriptor_base
attr = "load"
if not hasattr(descriptor, attr):
pytest.skip("descriptor load is not available")

original = getattr(descriptor, attr)
sentinel = object()
Expand Down
26 changes: 7 additions & 19 deletions tilelens/core/frontend/gluon.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,11 +9,18 @@
from triton.experimental.gluon.language import _semantic as gluon_semantic # type: ignore
from triton.experimental.gluon.language.amd import cdna3 as gluon_amd_cdna3 # type: ignore
from triton.experimental.gluon.language.amd import cdna4 as gluon_amd_cdna4 # type: ignore
from triton.experimental.gluon.language.amd import gfx1250 as gluon_amd_gfx1250 # type: ignore
from triton.experimental.gluon.language.amd import rdna3 as gluon_amd_rdna3 # type: ignore
from triton.experimental.gluon.language.amd import rdna4 as gluon_amd_rdna4 # type: ignore
from triton.experimental.gluon.language.amd.cdna4 import ( # type: ignore
async_copy as gluon_amd_cdna4_async_copy,
)
from triton.experimental.gluon.language.amd.gfx1250 import ( # type: ignore
async_copy as gluon_amd_async_copy,
)
from triton.experimental.gluon.language.amd.gfx1250 import ( # type: ignore
tdm as gluon_amd_tdm,
)
from triton.experimental.gluon.language.nvidia.blackwell import ( # type: ignore
tma as gluon_blackwell_tma,
)
Expand Down Expand Up @@ -52,23 +59,6 @@
from .base import AdapterResult, Frontend, _LangPatchScope, register_frontend
from .triton import TritonFrontend

try:
from triton.experimental.gluon.language.amd import gfx1250 as gluon_amd_gfx1250 # type: ignore
from triton.experimental.gluon.language.amd.gfx1250 import ( # type: ignore
async_copy as gluon_amd_async_copy,
)
from triton.experimental.gluon.language.amd.gfx1250 import ( # type: ignore
tdm as gluon_amd_tdm,
)
except ImportError as exc:
if "is_hip_gfx1250" not in str(exc) or "triton.language.target_info" not in str(
exc
):
raise
gluon_amd_gfx1250 = None
gluon_amd_async_copy = None
gluon_amd_tdm = None

_WARP_SPECIALIZE_SCHEDULER: Any = None
_MISSING = object()

Expand Down Expand Up @@ -309,8 +299,6 @@ def _gluon_make_range_adapter(start: Any, end: Any, *_args: Any, **_kwargs: Any)


def _existing_ops(namespace: Any, attrs: dict[str, type[Op]]) -> dict[str, type[Op]]:
if namespace is None:
return {}
namespace_attrs = vars(namespace)
return {attr: op_type for attr, op_type in attrs.items() if attr in namespace_attrs}

Expand Down
9 changes: 3 additions & 6 deletions tilelens/core/frontend/triton.py
Original file line number Diff line number Diff line change
Expand Up @@ -422,16 +422,13 @@ def _triton_language_attr_targets() -> list[tuple[Any, tuple[str, ...]]]:
targets.append((lang.dtype, ("to_ir",)))
if lang == tl:
targets.append((lang.math, ()))
if hasattr(tl.core, "tensor_descriptor_base"):
targets.append((tl.core.tensor_descriptor_base, ()))
targets.append((tl.core.tensor_descriptor_base, ()))
return targets

def _triton_snapshot_scope(self, fn: Callable[..., Any]) -> _LangPatchScope:
"""
Stores Triton attributes into a LangPatchScope for later unpatching.
This is to be run before patching with the interpreter.
This is equivalent to what triton>=3.6.0 does natively
but also works for triton<3.6.0.
"""

scope = _LangPatchScope()
Expand Down Expand Up @@ -905,8 +902,8 @@ def unpatch_for_loop(self) -> None:

def patch_lang(self, fn, client_manager=None) -> _LangPatchScope:
# Snapshot before calling Triton's patcher because Triton mutates many
# attributes in-place and older Triton versions do not retain enough
# restore metadata for nested/generated kernels.
# attributes in-place but only records the language modules visible
# from `fn`; nested/generated kernels without a `tl` global need both.
scope = self._triton_snapshot_scope(fn)
triton_patch_lang(fn)
for module in self._triton_extra_modules:
Expand Down
Loading
Loading