Fix backwards compatability for older versions of triton

This commit is contained in:
Jaret Burkett
2026-08-04 21:07:57 -06:00
parent a8d67ecd90
commit d870e9b68a

View File

@@ -287,6 +287,21 @@ def _triton_available() -> bool:
return _triton_ok
def _import_triton():
"""Lazy triton import for the kernel builders. The names are ALSO published
as module globals: older triton versions resolve a jit kernel's free
variables through ``fn.__globals__`` only (no closure capture), so ``tl``
referenced inside a kernel defined in a builder function must exist at
module scope or those versions die with ``NameError: 'tl' is not defined``
at first compile."""
import triton as _triton
import triton.language as _tl
globals()["triton"] = _triton
globals()["tl"] = _tl
return _triton, _tl
_kernel = None
@@ -294,8 +309,7 @@ def _get_kernel():
global _kernel
if _kernel is not None:
return _kernel
import triton
import triton.language as tl
triton, tl = _import_triton()
@triton.jit
def nvfp4_act_quant_kernel(
@@ -467,8 +481,7 @@ def _get_dequant_kernel():
global _dequant_kernel
if _dequant_kernel is not None:
return _dequant_kernel
import triton
import triton.language as tl
triton, tl = _import_triton()
@triton.jit
def nvfp4_dequant_kernel(
@@ -835,8 +848,7 @@ def _get_int8_kernels():
global _int8_kernels
if _int8_kernels is not None:
return _int8_kernels
import triton
import triton.language as tl
triton, tl = _import_triton()
from triton.language.extra import libdevice
@triton.jit
@@ -1631,8 +1643,7 @@ def _get_intn_kernel():
global _intn_kernel
if _intn_kernel is not None:
return _intn_kernel
import triton
import triton.language as tl
triton, tl = _import_triton()
@triton.jit
def intn_unpack_kernel(
@@ -1675,8 +1686,7 @@ def _get_intn_grouped_kernel():
global _intn_grouped_kernel
if _intn_grouped_kernel is not None:
return _intn_grouped_kernel
import triton
import triton.language as tl
triton, tl = _import_triton()
from triton.language.extra import libdevice
@triton.jit
@@ -1728,8 +1738,7 @@ def _get_bitnet_kernel():
global _bitnet_kernel
if _bitnet_kernel is not None:
return _bitnet_kernel
import triton
import triton.language as tl
triton, tl = _import_triton()
from triton.language.extra import libdevice
@triton.jit
@@ -1860,8 +1869,7 @@ def _get_int_gemv_kernel():
global _int_gemv_kernel
if _int_gemv_kernel is not None:
return _int_gemv_kernel
import triton
import triton.language as tl
triton, tl = _import_triton()
from triton.language.extra import libdevice
@triton.jit