TileOPs-Metax/tileops/kernels/elementwise.py

3668 lines
131 KiB
Python

"""Elementwise kernel templates and strategy factories.
Three template base classes for 66 elementwise ops:
- UnaryKernel: 1-input → 1-output (relu, sigmoid, abs, ...)
- BinaryKernel: 2-input → 1-output with N-dim broadcast (add, mul, ...)
- FusedGatedKernel: fused gate+activation (silu_and_mul, gelu_and_mul, ...)
Each kernel uses one of three strategies (no shared memory):
Global → Register → Compute → Register → Global
Strategies:
- direct: 1 element per thread, simplest codegen
- explicit_parallel: N elements per thread via T.Parallel(threads, npt)
- register_copy: fragment load → compute → fragment store (unary only)
Binary register_copy is NOT supported (incompatible with stride-based access).
Boundary checks handled by TileLang LegalizeSafeMemoryAccess.
fp8 (e4m3fn, e5m2) accumulates in fp16 — direct fp8 arithmetic loses too much
precision for sigmoid/exp and friends. Defaults: num_per_thread=16 (128-bit
alignment) and explicit_parallel (register_copy is unreliable for fp8).
Saturation follows the NVIDIA spec: e4m3fn has no Inf, so the kernel's
saturating T.Cast clamping to ±448.0 is correct; e5m2 does, so the kernel emits
fp16 and the Op layer does the final non-saturating cast.
"""
import functools
import math
import warnings
import tilelang
import tilelang.language as T
import torch
from tileops.kernels.kernel_base import Kernel
__all__ = [
# --- base classes ---
"BinaryKernel",
"FusedGatedKernel",
"UnaryKernel",
# --- unary: existing ---
"ReluFwdKernel",
# --- unary: math (17) ---
"AbsFwdKernel",
"CeilFwdKernel",
"CosFwdKernel",
"ErfFwdKernel",
"ExpFwdKernel",
"Expm1FwdKernel",
"FloorFwdKernel",
"Log1pFwdKernel",
"LogFwdKernel",
"NegFwdKernel",
"ReciprocalFwdKernel",
"RoundFwdKernel",
"RsqrtFwdKernel",
"SignFwdKernel",
"SinFwdKernel",
"SqrtFwdKernel",
"TruncFwdKernel",
# --- unary: activations (9) ---
"GeluFwdKernel",
"GeluTanhFwdKernel",
"HardsigmoidFwdKernel",
"HardswishFwdKernel",
"MishFwdKernel",
"SeluFwdKernel",
"SigmoidFwdKernel",
"SiluFwdKernel",
"TanhFwdKernel",
# --- unary: logical / bitwise (2) ---
"BitwiseNotFwdKernel",
"LogicalNotBoolStorageFwdKernel",
"LogicalNotFwdKernel",
# --- unary: special predicates (3) ---
"IsfiniteFwdKernel",
"IsinfFwdKernel",
"IsnanFwdKernel",
# --- binary arithmetic ---
"AddFwdKernel",
"SubFwdKernel",
"MulFwdKernel",
"DivFwdKernel",
"DivTruncFwdKernel",
"RemainderFwdKernel",
"PowFwdKernel",
"FloorDivideFwdKernel",
"LerpFwdKernel",
"MaximumFwdKernel",
"MinimumFwdKernel",
# --- comparison (OUTPUT_DTYPE = torch.bool) ---
"EqBoolStorageFwdKernel",
"EqFwdKernel",
"GeBoolStorageFwdKernel",
"GeFwdKernel",
"GtBoolStorageFwdKernel",
"GtFwdKernel",
"LeBoolStorageFwdKernel",
"LeFwdKernel",
"LtFwdKernel",
"LtBoolStorageFwdKernel",
"NeBoolStorageFwdKernel",
"NeFwdKernel",
# --- logical (OUTPUT_DTYPE = torch.bool) ---
"LogicalAndBoolStorageFwdKernel",
"LogicalAndFwdKernel",
"LogicalOrBoolStorageFwdKernel",
"LogicalOrFwdKernel",
# --- bitwise ---
"BitwiseAndBoolStorageFwdKernel",
"BitwiseAndFwdKernel",
"BitwiseOrBoolStorageFwdKernel",
"BitwiseOrFwdKernel",
"BitwiseXorBoolStorageFwdKernel",
"BitwiseXorFwdKernel",
# --- fused gated ---
"SiluAndMulFwdKernel",
"GeluAndMulFwdKernel",
"GeluTanhAndMulFwdKernel",
# --- independent (custom-signature) ---
"LeakyReluFwdKernel",
"EluFwdKernel",
"HardtanhFwdKernel",
"SoftplusFwdKernel",
"PreluFwdKernel",
"WhereFwdKernel",
"LerpTensorFwdKernel",
"ClampFwdKernel",
"ClampTensorFwdKernel",
"MaskedFillFwdKernel",
"MaskedFillTensorValueFwdKernel",
"NanToNumFwdKernel",
"AlibiFwdKernel",
"SinusoidalFwdKernel",
]
_BITWISE_DTYPES = (
torch.bool,
torch.uint8,
torch.int8,
torch.int16,
torch.int32,
torch.int64,
)
_FP8_DTYPES = (
torch.float8_e4m3fn,
torch.float8_e5m2,
)
_FLOAT_DTYPES = (
torch.float16,
torch.bfloat16,
torch.float32,
)
_LOGICAL_DTYPES = _BITWISE_DTYPES + _FLOAT_DTYPES
# Binary arithmetic dtype unions, mirroring the manifest entries for
# torch.add / torch.sub. fp8 is excluded because PyTorch does not define
# add/sub for float8 storage; bool is excluded for sub because PyTorch
# rejects bool subtraction.
_BINARY_FULL_DTYPES = _BITWISE_DTYPES + (
torch.float16,
torch.bfloat16,
torch.float32,
)
_BINARY_NO_BOOL_DTYPES = tuple(
dt for dt in _BINARY_FULL_DTYPES if dt is not torch.bool
)
def _is_fp8(dtype: torch.dtype) -> bool:
"""Check if a torch dtype is an fp8 variant."""
return dtype in _FP8_DTYPES
def _strategy_npt(strategy: str, dtype: torch.dtype) -> int:
"""Return the default num_per_thread for a strategy + dtype pair.
Strategy-aware heuristic (from H200 benchmarks):
- explicit_parallel: npt=4 for fp16/bf16 (42% bandwidth gain vs npt=8)
- register_copy: npt=8 for fp16/bf16 (vectorized 128-bit loads)
- fp32: npt=4 for all strategies (4 bytes x 4 = 128-bit alignment)
- fp8: handled separately by callers (npt=16)
"""
if dtype == torch.float32:
return 4
# fp16 / bf16: strategy-dependent
if strategy == "explicit_parallel" and dtype in (torch.float16, torch.bfloat16):
return 4
return 8
def _fp8_needs_nonsaturating_cast(dtype: torch.dtype) -> bool:
"""Return True if the fp8 format supports Inf/NaN and needs non-saturating output.
e5m2 has Inf/NaN representation -- TileLang's T.Cast uses saturating conversion
which incorrectly clamps Inf to max-finite. For e5m2, the kernel must produce
fp16 output and let PyTorch do the final non-saturating cast.
e4m3fn has no Inf representation, so saturating T.Cast is correct.
"""
return dtype == torch.float8_e5m2
def _fp8_accum_dtype_str() -> str:
"""Return the TileLang dtype string used for fp8 intermediate accumulation."""
return "float16"
def _get_fp8_output_dtypes(dtype: torch.dtype):
"""Return (fp8_output_dtype, kernel_output_dtype) for fp8 handling.
For e5m2: kernel produces fp16 to preserve Inf/NaN; Op layer does the
final non-saturating cast to e5m2 via PyTorch.
For e4m3fn or non-fp8: kernel outputs directly in the input dtype.
Returns:
Tuple of (_fp8_output_dtype, output_dtype). _fp8_output_dtype is
the original fp8 dtype when a post-cast is needed, else None.
"""
if _is_fp8(dtype) and _fp8_needs_nonsaturating_cast(dtype):
return dtype, torch.float16
return None, dtype
def _clamp_to_dtype_range(value, dtype: torch.dtype):
"""Normalize *value* into the storage representation of *dtype*.
Mirrors PyTorch ``Tensor.masked_fill`` scalar coercion so the literal lands
as the same bit pattern PyTorch would write:
- bool: non-zero → ``1``, else ``0``.
- Signed int: truncate toward zero; ``+/-Inf`` maps to ``iinfo.max/min``
so a bypassed validator cannot raise ``OverflowError`` on ``int(inf)``.
- ``uint8``: negatives wrap via ``& 0xFF``, non-negatives truncate.
- ``fp16/bf16/fp32`` and ``fp8_e5m2``: ``NaN`` / ``+-Inf`` pass through,
finite values clamp to ``finfo``.
- ``fp8_e4m3fn`` has no Inf, so ``+-Inf`` saturates to ``finfo.max/min``
to avoid a TVM ``FloatImm`` overflow.
"""
if dtype == torch.bool:
return 1 if bool(value) else 0
if dtype in _BITWISE_DTYPES:
if isinstance(value, float) and math.isinf(value):
iinfo = torch.iinfo(dtype)
return iinfo.max if value > 0 else iinfo.min
if dtype == torch.uint8 and isinstance(value, int) and not isinstance(value, bool) and value < 0:
return value & 0xFF
return int(value)
fvalue = float(value)
if math.isnan(fvalue):
return fvalue
finfo = torch.finfo(dtype)
if math.isinf(fvalue):
if dtype in _FP8_DTYPES and not _fp8_needs_nonsaturating_cast(dtype):
return finfo.max if fvalue > 0 else finfo.min
return fvalue
return max(finfo.min, min(finfo.max, fvalue))
def _wrap_fp8_accumulation(base_op, dtype, dtype_str, arity=1):
"""Wrap an op function with fp8 accumulation logic if *dtype* is fp8.
Both fp8 dtypes cast inputs to fp16 and compute there. e4m3fn casts the
result back via saturating ``T.Cast`` (correct — it has no Inf); e5m2
leaves the result in fp16 and the Op layer does the final non-saturating
cast, which preserves Inf/NaN.
Non-fp8 dtypes get *base_op* back unchanged.
"""
if not _is_fp8(dtype):
return base_op
accum = _fp8_accum_dtype_str()
if _fp8_needs_nonsaturating_cast(dtype):
# e5m2: compute in fp16, leave result as fp16
if arity == 1:
def fp8_accum_op(x):
return base_op(T.cast(x, accum))
else:
def fp8_accum_op(a, b):
return base_op(T.cast(a, accum), T.cast(b, accum))
return fp8_accum_op
# e4m3fn: compute in fp16, saturating cast back
if arity == 1:
def fp8_accum_op(x):
return T.Cast(dtype_str, base_op(T.cast(x, accum)))
else:
def fp8_accum_op(a, b):
return T.Cast(dtype_str, base_op(T.cast(a, accum), T.cast(b, accum)))
return fp8_accum_op
# Strategy factory: Unary
@functools.lru_cache(maxsize=32)
def _make_unary_direct(N, dtype, op_func, output_dtype=None, threads=256):
"""Strategy 1: 1 element per thread."""
out_dtype = output_dtype or dtype
@tilelang.jit(out_idx=[1])
def kernel(threads_arg):
@T.prim_func
def main(x: T.Tensor((N,), dtype), y: T.Tensor((N,), out_dtype)):
with T.Kernel(T.ceildiv(N, threads_arg), threads=threads_arg) as bx:
for i in T.Parallel(threads_arg):
idx = bx * threads_arg + i
y[idx] = op_func(x[idx])
return main
return kernel
@functools.lru_cache(maxsize=32)
def _make_unary_explicit(N, dtype, op_func, output_dtype=None, threads=256, num_per_thread=8):
"""Strategy 2: N elements per thread via T.Parallel(threads, npt)."""
block_size = threads * num_per_thread
out_dtype = output_dtype or dtype
@tilelang.jit(out_idx=[1])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(x: T.Tensor((N,), dtype), y: T.Tensor((N,), out_dtype)):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
for i, j in T.Parallel(threads_arg, npt_arg):
idx = (bx * threads_arg + i) * npt_arg + j
y[idx] = op_func(x[idx])
return main
return kernel
@functools.lru_cache(maxsize=32)
def _make_unary_regcopy(N, dtype, op_func, output_dtype=None, threads=256, num_per_thread=8):
"""Strategy 3: fragment load → compute → fragment store."""
block_size = threads * num_per_thread
out_dtype = output_dtype or dtype
@tilelang.jit(out_idx=[1])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(x: T.Tensor((N,), dtype), y: T.Tensor((N,), out_dtype)):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
x_reg = T.alloc_fragment((block_size,), dtype)
y_reg = T.alloc_fragment((block_size,), out_dtype)
T.copy(x[bx * block_size : (bx + 1) * block_size], x_reg)
for i, j in T.Parallel(threads_arg, npt_arg):
y_reg[i * npt_arg + j] = op_func(x_reg[i * npt_arg + j])
T.copy(y_reg, y[bx * block_size : (bx + 1) * block_size])
return main
return kernel
# Strategy factory: Binary
def _compute_broadcast_offsets(flat_idx, ndim, divisors, a_strides, b_strides):
"""Compute a_off and b_off from flat_idx using compile-time unrolled divmod chain.
All arguments except flat_idx are Python-level constants, so the loop
unrolls at kernel build time.
"""
a_off = 0
b_off = 0
remaining = flat_idx
for d in range(ndim - 1):
coord = remaining // divisors[d]
remaining = remaining % divisors[d]
a_off = a_off + coord * a_strides[d]
b_off = b_off + coord * b_strides[d]
a_off = a_off + remaining * a_strides[ndim - 1]
b_off = b_off + remaining * b_strides[ndim - 1]
return a_off, b_off
def _is_contiguous_same_shape(coalesced_shape, a_strides, b_strides):
"""Return True when both inputs are contiguous with the same shape (no broadcast)."""
return (
len(coalesced_shape) == 1
and all(s == 1 for s in a_strides)
and all(s == 1 for s in b_strides)
)
@functools.lru_cache(maxsize=32)
def _make_binary_register_copy(
N_total, dtype, op_func, output_dtype=None, threads=256, num_per_thread=8,
):
"""Binary register_copy: fragment load -> compute -> fragment store.
Only available for same-shape contiguous inputs (no broadcast).
Uses T.alloc_fragment + T.copy for vectorized 128-bit memory access,
giving ~2-3x bandwidth vs scalar access for complex op_funcs that
prevent TVM's auto-vectorizer from kicking in.
"""
out_dtype = output_dtype or dtype
@tilelang.jit(out_idx=[2])
def kernel(threads, num_per_thread):
block_size = threads * num_per_thread
@T.prim_func
def main(
a: T.Tensor((N_total,), dtype),
b: T.Tensor((N_total,), dtype),
y: T.Tensor((N_total,), out_dtype),
):
with T.Kernel(T.ceildiv(N_total, block_size), threads=threads) as bx:
a_reg = T.alloc_fragment((block_size,), dtype)
b_reg = T.alloc_fragment((block_size,), dtype)
y_reg = T.alloc_fragment((block_size,), out_dtype)
T.copy(a[bx * block_size:(bx + 1) * block_size], a_reg)
T.copy(b[bx * block_size:(bx + 1) * block_size], b_reg)
for i, j in T.Parallel(threads, num_per_thread):
idx = i * num_per_thread + j
y_reg[idx] = op_func(a_reg[idx], b_reg[idx])
T.copy(y_reg, y[bx * block_size:(bx + 1) * block_size])
return main
return kernel
@functools.lru_cache(maxsize=32)
def _make_binary_direct(
N_total, dtype, op_func, coalesced_shape, a_strides, b_strides,
a_numel, b_numel, output_dtype=None, threads=256,
):
"""Binary direct: 1 element per thread with stride-based broadcast."""
out_dtype = output_dtype or dtype
# Fast path: same-shape contiguous inputs -- skip broadcast machinery
if _is_contiguous_same_shape(coalesced_shape, a_strides, b_strides):
@tilelang.jit(out_idx=[2])
def kernel(threads):
@T.prim_func
def main(
a: T.Tensor((N_total,), dtype),
b: T.Tensor((N_total,), dtype),
y: T.Tensor((N_total,), out_dtype),
):
with T.Kernel(T.ceildiv(N_total, threads), threads=threads) as bx:
for i in T.Parallel(threads):
idx = bx * threads + i
y[idx] = op_func(a[idx], b[idx])
return main
return kernel
ndim = len(coalesced_shape)
divisors = [1] * ndim
for i in range(ndim - 2, -1, -1):
divisors[i] = divisors[i + 1] * coalesced_shape[i + 1]
@tilelang.jit(out_idx=[2])
def kernel(threads):
@T.prim_func
def main(
a: T.Tensor((a_numel,), dtype),
b: T.Tensor((b_numel,), dtype),
y: T.Tensor((N_total,), out_dtype),
):
with T.Kernel(T.ceildiv(N_total, threads), threads=threads) as bx:
for i in T.Parallel(threads):
flat_idx = bx * threads + i
a_off, b_off = _compute_broadcast_offsets(
flat_idx, ndim, divisors, a_strides, b_strides,
)
y[flat_idx] = op_func(a[a_off], b[b_off])
return main
return kernel
@functools.lru_cache(maxsize=32)
def _make_binary_explicit(
N_total, dtype, op_func, coalesced_shape, a_strides, b_strides,
a_numel, b_numel, output_dtype=None, threads=256, num_per_thread=8,
):
"""Binary explicit_parallel: N elements per thread with stride-based broadcast."""
out_dtype = output_dtype or dtype
# Fast path: same-shape contiguous inputs -- skip broadcast machinery
if _is_contiguous_same_shape(coalesced_shape, a_strides, b_strides):
@tilelang.jit(out_idx=[2])
def kernel(threads, num_per_thread):
block_size = threads * num_per_thread
@T.prim_func
def main(
a: T.Tensor((N_total,), dtype),
b: T.Tensor((N_total,), dtype),
y: T.Tensor((N_total,), out_dtype),
):
with T.Kernel(T.ceildiv(N_total, block_size), threads=threads) as bx:
for i, j in T.Parallel(threads, num_per_thread):
idx = (bx * threads + i) * num_per_thread + j
y[idx] = op_func(a[idx], b[idx])
return main
return kernel
ndim = len(coalesced_shape)
divisors = [1] * ndim
for i in range(ndim - 2, -1, -1):
divisors[i] = divisors[i + 1] * coalesced_shape[i + 1]
@tilelang.jit(out_idx=[2])
def kernel(threads, num_per_thread):
block_size = threads * num_per_thread
@T.prim_func
def main(
a: T.Tensor((a_numel,), dtype),
b: T.Tensor((b_numel,), dtype),
y: T.Tensor((N_total,), out_dtype),
):
with T.Kernel(T.ceildiv(N_total, block_size), threads=threads) as bx:
for i, j in T.Parallel(threads, num_per_thread):
flat_idx = (bx * threads + i) * num_per_thread + j
a_off, b_off = _compute_broadcast_offsets(
flat_idx, ndim, divisors, a_strides, b_strides,
)
y[flat_idx] = op_func(a[a_off], b[b_off])
return main
return kernel
# Strategy factory: FusedGated
@functools.lru_cache(maxsize=32)
def _make_fused_gated_direct(M, N, dtype, op_func, threads=256, output_dtype=None):
"""FusedGated direct: 1 element per thread. x[:, :N] is gate, x[:, N:] is value.
``op_func(gate, value)`` is the compound operation that applies the
activation to *gate* and multiplies by *value*. For fp8 dtypes the
caller wraps it via ``_wrap_fp8_accumulation`` so this factory stays
fp8-agnostic.
Args:
output_dtype: TileLang dtype string for the output tensor. Defaults to dtype.
"""
out_dtype = output_dtype or dtype
@tilelang.jit(out_idx=[1])
def kernel(threads_arg):
@T.prim_func
def main(x: T.Tensor((M, 2 * N), dtype), y: T.Tensor((M, N), out_dtype)):
with T.Kernel(T.ceildiv(N, threads_arg), M, threads=threads_arg) as (bx, by):
for i in T.Parallel(threads_arg):
col = bx * threads_arg + i
gate = x[by, col]
value = x[by, N + col]
y[by, col] = op_func(gate, value)
return main
return kernel
@functools.lru_cache(maxsize=32)
def _make_fused_gated_explicit(M, N, dtype, op_func, threads=256, num_per_thread=8,
output_dtype=None):
"""FusedGated explicit_parallel: N elements per thread.
``op_func(gate, value)`` is the compound operation (see
``_make_fused_gated_direct``). fp8 accumulation is handled by the
caller wrapping ``op_func`` via ``_wrap_fp8_accumulation``, so this
factory no longer needs an ``fp8_accum`` parameter.
Args:
output_dtype: TileLang dtype string for the output tensor. Defaults to dtype.
"""
block_N = threads * num_per_thread
out_dtype = output_dtype or dtype
@tilelang.jit(out_idx=[1])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(x: T.Tensor((M, 2 * N), dtype), y: T.Tensor((M, N), out_dtype)):
with T.Kernel(T.ceildiv(N, block_N), M, threads=threads_arg) as (bx, by):
for i, j in T.Parallel(threads_arg, npt_arg):
col = (bx * threads_arg + i) * npt_arg + j
gate = x[by, col]
value = x[by, N + col]
y[by, col] = op_func(gate, value)
return main
return kernel
# Template base classes
class UnaryKernel(Kernel):
"""Template base class for unary elementwise kernels.
Subclass must override ``op_func`` with a static method implementing
the pointwise operation (e.g., relu, sigmoid).
Args:
N_total: Total number of elements (flattened).
dtype: Torch dtype for input.
config: Optional dict with "strategy", "threads" and "num_per_thread".
"strategy" is one of "direct", "explicit_parallel",
"register_copy"; it selects the kernel body at build time.
tune: Whether to autotune (sweeps "threads" / "num_per_thread"
within the resolved strategy).
"""
supported_archs: list[int] = [80, 86, 89, 90]
STRATEGIES = ["direct", "explicit_parallel", "register_copy"]
# Benchmark (H200): register_copy wins for fp16/bf16 across all tested shapes;
# fp32 small shapes show variance between register_copy and explicit_parallel.
DEFAULT_STRATEGY = "register_copy"
OUTPUT_DTYPE = None
SUPPORTED_DTYPES = None
@staticmethod
def op_func(x):
"""Pointwise operation. Must be overridden by subclass."""
raise NotImplementedError
def __init__(self, N_total, dtype, config=None, tune=False):
super().__init__()
if self.SUPPORTED_DTYPES is not None and dtype not in self.SUPPORTED_DTYPES:
supported = ", ".join(str(dt) for dt in self.SUPPORTED_DTYPES)
raise ValueError(
f"{self.__class__.__name__} only supports dtypes [{supported}], got {dtype}"
)
self.N_total = N_total
self.dtype = dtype
# For e5m2: kernel produces fp16 to preserve Inf/NaN; Op layer
# performs the final non-saturating cast to e5m2 via PyTorch.
# For e4m3fn: kernel produces e4m3fn via saturating T.Cast (correct,
# since e4m3fn has no Inf representation).
self._fp8_output_dtype = None
if _is_fp8(dtype) and self.OUTPUT_DTYPE is None and _fp8_needs_nonsaturating_cast(dtype):
self._fp8_output_dtype = dtype
self.output_dtype = torch.float16
else:
self.output_dtype = self.OUTPUT_DTYPE or dtype
# Validate a config-requested strategy up front so typos raise the
# same ValueError regardless of dtype (the bool coercion below would
# otherwise silently accept an unknown strategy for bool inputs).
requested = (config or {}).get("strategy")
if requested is not None and requested not in self.STRATEGIES:
raise ValueError(
f"Unknown strategy '{requested}', expected one of {self.STRATEGIES}"
)
# torch.bool maps to TileLang ``boolx<N>`` for vectorised loads, which
# the CUDA codegen cannot lower. Keep bool inputs on the scalar path.
bool_output = torch.bool == self.OUTPUT_DTYPE
bool_output_needs_scalar = bool_output and dtype in (
torch.uint8, torch.int8, torch.int16,
)
if dtype == torch.bool:
if requested is not None and requested != "direct":
warnings.warn(
f"UnaryKernel: dtype=torch.bool requires strategy="
f"'direct' (TileLang cannot lower vectorised boolx<N> "
f"loads); overriding requested strategy={requested!r}.",
RuntimeWarning,
stacklevel=2,
)
self.strategy = "direct"
elif bool_output_needs_scalar:
if requested is not None and requested != "direct":
warnings.warn(
f"UnaryKernel: dtype={dtype} with torch.bool output "
f"requires strategy='direct' (TileLang cannot lower "
f"vectorised boolx<N> stores for sub-32-bit integer "
f"inputs); overriding requested strategy={requested!r}.",
RuntimeWarning,
stacklevel=2,
)
self.strategy = "direct"
# fp8: register_copy may not reliably handle 8-bit fragments;
# default to explicit_parallel for fp8 dtypes
elif requested is None and _is_fp8(dtype):
self.strategy = "explicit_parallel"
else:
self.strategy = requested or self.DEFAULT_STRATEGY
if self.strategy not in self.STRATEGIES:
raise ValueError(
f"Unknown strategy '{self.strategy}', expected one of {self.STRATEGIES}"
)
self.kernel = self._build_kernel(self.strategy)
self.init_config(config, tune)
def _get_effective_op_func(self):
"""Return op_func wrapped with fp8->fp16 accumulation if needed.
Delegates to the shared ``_wrap_fp8_accumulation`` helper.
When ``OUTPUT_DTYPE`` is set (e.g. bool-output ops) fp8 wrapping is
skipped because the kernel already outputs a non-fp8 type.
"""
if self.OUTPUT_DTYPE is not None:
return self.op_func
return _wrap_fp8_accumulation(self.op_func, self.dtype, self.dtype_str, arity=1)
def _build_kernel(self, strategy):
cfg = self.default_config
effective_op = self._get_effective_op_func()
if strategy == "direct":
return _make_unary_direct(
self.N_total, self.dtype_str, effective_op,
output_dtype=self.output_dtype_str, threads=cfg["threads"],
)
elif strategy == "explicit_parallel":
return _make_unary_explicit(
self.N_total, self.dtype_str, effective_op,
output_dtype=self.output_dtype_str,
threads=cfg["threads"], num_per_thread=cfg["num_per_thread"],
)
elif strategy == "register_copy":
return _make_unary_regcopy(
self.N_total, self.dtype_str, effective_op,
output_dtype=self.output_dtype_str,
threads=cfg["threads"], num_per_thread=cfg["num_per_thread"],
)
else:
raise ValueError(f"Unknown strategy: {strategy}")
@property
def output_dtype_str(self) -> str:
return self.dtype_to_str(self.output_dtype)
@property
def default_config(self) -> dict:
if _is_fp8(self.dtype):
# fp8: 1 byte per element, 16 elements = 128-bit alignment
return {"strategy": self.strategy, "threads": 256, "num_per_thread": 16}
npt = _strategy_npt(self.strategy, self.dtype)
return {"strategy": self.strategy, "threads": 256, "num_per_thread": npt}
@property
def autotune_configs(self) -> list[dict]:
"""Search space: threads in {128, 256, 512} x num_per_thread in {2, 4, 8}.
Covers a range of occupancy/register-pressure tradeoffs for
bandwidth-bound unary elementwise kernels. "strategy" is a
build-time config key (it selects the kernel body, not a JIT
parameter), so it is excluded from the sweep.
"""
if _is_fp8(self.dtype):
# fp8 needs 128-bit alignment: npt >= 16 for 1-byte elements
threads_opts = [128, 256, 512]
npt_opts = [16, 32]
else:
# fp16 / bf16 / fp32
threads_opts = [128, 256, 512]
npt_opts = [2, 4, 8]
return [
{"threads": t, "num_per_thread": n}
for t in threads_opts
for n in npt_opts
]
def autotune(self, warmup: int = 10, rep: int = 10) -> None:
"""Override to handle serialization failures in the TileLang autotuner.
UnaryKernel JIT functions capture op_func closures that the autotuner
subprocess cannot serialize. Catch the error and fall back to the
default config so that ``tune=True`` never crashes.
"""
import warnings
try:
super().autotune(warmup=warmup, rep=rep)
except (AssertionError, Exception) as exc:
if "not serializable" in str(exc) or "pickle" in str(exc).lower():
warnings.warn(
f"{self.__class__.__name__} autotuning failed "
f"(op_func is not serializable); falling back to "
f"default_config.",
stacklevel=2,
)
self.config = dict(self.default_config)
else:
raise
def init_config(self, config=None, tune=False):
"""Override to cache the compiled kernel function after config is set."""
super().init_config(config, tune)
# Record the resolved strategy so ``self.config`` is the single
# source of truth (a coerced/downgraded request or an autotune
# result would otherwise leave the key stale or missing).
self.config["strategy"] = self.strategy
# Pre-compile and cache the kernel function for the chosen config
# to avoid JIT lookup overhead on every forward() call.
cfg = self.config
if self.strategy == "direct":
self._compiled_fn = self.kernel(cfg["threads"])
else:
self._compiled_fn = self.kernel(cfg["threads"], cfg["num_per_thread"])
def forward(self, x):
result = self._compiled_fn(x)
if self._fp8_output_dtype is not None:
result = result.to(self._fp8_output_dtype)
return result
class BinaryKernel(Kernel):
"""Template base class for binary elementwise kernels with N-dim broadcast.
Subclass must override ``op_func`` with a static method implementing
the pointwise operation (e.g., add, mul).
Args:
N_total: Total output elements.
dtype: Torch dtype for input.
coalesced_shape: Coalesced broadcast dimensions.
a_strides: Strides for input a (0 means broadcast).
b_strides: Strides for input b (0 means broadcast).
a_numel: Number of elements in a.
b_numel: Number of elements in b.
config: Optional dict with "strategy", "threads" and "num_per_thread".
"strategy" is one of "direct", "explicit_parallel",
"register_copy". If "register_copy" is requested but inputs
require broadcast, silently downgrades to "explicit_parallel".
tune: Whether to autotune (sweeps "threads" / "num_per_thread"
within the resolved strategy).
"""
supported_archs: list[int] = [80, 86, 89, 90]
STRATEGIES = ["direct", "explicit_parallel", "register_copy"]
DEFAULT_STRATEGY = "explicit_parallel"
OUTPUT_DTYPE = None # Subclass override for output dtype (e.g., torch.int8)
SUPPORTED_DTYPES = None # Subclass override to restrict input dtypes
@staticmethod
def op_func(a, b):
"""Pointwise operation. Must be overridden by subclass."""
raise NotImplementedError
def __init__(
self, N_total, dtype, coalesced_shape, a_strides, b_strides,
a_numel, b_numel, config=None, tune=False,
):
super().__init__()
if self.SUPPORTED_DTYPES is not None and dtype not in self.SUPPORTED_DTYPES:
supported = ", ".join(str(dt) for dt in self.SUPPORTED_DTYPES)
raise ValueError(
f"{self.__class__.__name__} only supports dtypes [{supported}], got {dtype}"
)
self.N_total = N_total
self.dtype = dtype
self._fp8_output_dtype = None
if _is_fp8(dtype) and self.OUTPUT_DTYPE is None and _fp8_needs_nonsaturating_cast(dtype):
self._fp8_output_dtype = dtype
self.output_dtype = torch.float16
else:
self.output_dtype = self.OUTPUT_DTYPE or dtype
self.coalesced_shape = coalesced_shape
self.a_strides = a_strides
self.b_strides = b_strides
self.a_numel = a_numel
self.b_numel = b_numel
self._same_shape = _is_contiguous_same_shape(
coalesced_shape, a_strides, b_strides,
)
# Validate a config-requested strategy up front so typos raise the
# same ValueError regardless of dtype (the bool override below
# otherwise silently accepts an unknown strategy for bool inputs).
requested = (config or {}).get("strategy")
if requested is not None and requested not in self.STRATEGIES:
raise ValueError(
f"Unknown strategy '{requested}', expected one of {self.STRATEGIES}"
)
# torch.bool maps to TileLang ``boolx<N>`` for vectorised loads /
# stores, which the CUDA codegen cannot lower. Force the scalar
# ``direct`` strategy for bool inputs regardless of caller request.
bool_input = dtype == torch.bool
bool_output = torch.bool == self.OUTPUT_DTYPE
bool_output_needs_scalar = bool_output and dtype in (
torch.uint8, torch.int8, torch.int16,
)
if bool_input:
if requested is not None and requested != "direct":
warnings.warn(
f"BinaryKernel: dtype=torch.bool requires strategy="
f"'direct' (TileLang cannot lower vectorised boolx<N> "
f"loads); overriding requested strategy={requested!r}.",
RuntimeWarning,
stacklevel=2,
)
self.strategy = "direct"
elif bool_output_needs_scalar:
if requested is not None and requested != "direct":
warnings.warn(
f"BinaryKernel: dtype={dtype} with torch.bool output "
f"requires strategy='direct' (TileLang cannot lower "
f"vectorised boolx<N> stores for sub-32-bit integer "
f"inputs); overriding requested strategy={requested!r}.",
RuntimeWarning,
stacklevel=2,
)
self.strategy = "direct"
elif requested is not None:
# register_copy requires same-shape contiguous inputs (no
# broadcast); silently downgrade to explicit_parallel when
# the caller requests register_copy on broadcast shapes.
if requested == "register_copy" and (not self._same_shape or bool_output):
self.strategy = "explicit_parallel"
else:
self.strategy = requested
elif self._same_shape and not bool_output:
# register_copy gives vectorized 128-bit loads, ~2-3x faster
# for complex op_funcs that block TVM's auto-vectorizer.
self.strategy = "register_copy"
else:
self.strategy = self.DEFAULT_STRATEGY
if self.strategy not in self.STRATEGIES:
raise ValueError(
f"Unknown strategy '{self.strategy}', expected one of {self.STRATEGIES}"
)
self.kernel = self._build_kernel(self.strategy)
self.init_config(config, tune)
def _get_effective_op_func(self):
"""Return op_func wrapped with fp8->fp16 accumulation if needed.
Delegates to the shared ``_wrap_fp8_accumulation`` helper (arity=2).
When ``OUTPUT_DTYPE`` is set (e.g. comparison/logical ops) fp8 wrapping
is skipped because the kernel already outputs a non-fp8 type.
"""
if self.OUTPUT_DTYPE is not None:
return self.op_func
return _wrap_fp8_accumulation(self.op_func, self.dtype, self.dtype_str, arity=2)
def _build_kernel(self, strategy):
cfg = self.default_config
effective_op = self._get_effective_op_func()
# For e5m2: kernel output is fp16 (non-saturating path)
kernel_output_dtype = (
self.dtype_to_str(self.OUTPUT_DTYPE) if self.OUTPUT_DTYPE is not None else None
)
if self._fp8_output_dtype is not None:
kernel_output_dtype = _fp8_accum_dtype_str()
if strategy == "direct":
return _make_binary_direct(
self.N_total, self.dtype_str, effective_op,
self.coalesced_shape, self.a_strides, self.b_strides,
self.a_numel, self.b_numel,
output_dtype=kernel_output_dtype, threads=cfg["threads"],
)
elif strategy == "explicit_parallel":
return _make_binary_explicit(
self.N_total, self.dtype_str, effective_op,
self.coalesced_shape, self.a_strides, self.b_strides,
self.a_numel, self.b_numel,
output_dtype=kernel_output_dtype,
threads=cfg["threads"], num_per_thread=cfg["num_per_thread"],
)
elif strategy == "register_copy":
return _make_binary_register_copy(
self.N_total, self.dtype_str, effective_op,
output_dtype=kernel_output_dtype,
threads=cfg["threads"], num_per_thread=cfg["num_per_thread"],
)
else:
raise ValueError(f"Unknown strategy: {strategy}")
@property
def default_config(self) -> dict:
if _is_fp8(self.dtype):
return {"strategy": self.strategy, "threads": 256, "num_per_thread": 16}
npt = _strategy_npt(self.strategy, self.dtype)
return {"strategy": self.strategy, "threads": 256, "num_per_thread": npt}
@property
def autotune_configs(self) -> list[dict]:
"""Search space: threads in {128, 256, 512} x num_per_thread in {2, 4, 8}.
Covers a range of occupancy/register-pressure tradeoffs for
bandwidth-bound binary elementwise kernels.
"""
if _is_fp8(self.dtype):
# fp8 needs 128-bit alignment: npt >= 16 for 1-byte elements
threads_opts = [128, 256, 512]
npt_opts = [16, 32]
else:
# fp16 / bf16 / fp32
threads_opts = [128, 256, 512]
npt_opts = [2, 4, 8]
return [
{"threads": t, "num_per_thread": n}
for t in threads_opts
for n in npt_opts
]
def autotune(self, warmup: int = 10, rep: int = 10) -> None:
"""Override to handle known TileLang autotuner fallback failures.
BinaryKernel JIT functions capture op_func closures that the autotuner
subprocess cannot serialize. Newer TileLang binders can also reject
the autotune wrapper signature. Catch these errors and fall back to
the default config so that ``tune=True`` never crashes.
"""
import warnings
try:
super().autotune(warmup=warmup, rep=rep)
except (AssertionError, Exception) as exc:
message = str(exc)
if (
"not serializable" in message
or "pickle" in message.lower()
or "missing a required argument" in message
):
warnings.warn( # noqa: B028
f"{self.__class__.__name__} autotuning failed "
f"({message}); falling back to default_config.")
self.config = dict(self.default_config)
else:
raise
def init_config(self, config=None, tune=False):
"""Override to cache the compiled kernel function after config is set."""
super().init_config(config, tune)
self.config["strategy"] = self.strategy
# Pre-compile and cache the kernel function for the chosen config
# to avoid JIT lookup overhead on every forward() call.
cfg = self.config
if self.strategy == "direct":
self._compiled_fn = self.kernel(cfg["threads"])
else:
self._compiled_fn = self.kernel(cfg["threads"], cfg["num_per_thread"])
def forward(self, a, b):
result = self._compiled_fn(a, b)
if self._fp8_output_dtype is not None:
result = result.to(self._fp8_output_dtype)
return result
class FusedGatedKernel(Kernel):
"""Template base class for fused gated elementwise kernels.
Input layout: x has shape (M, 2*N) where x[:, :N] is the gate
and x[:, N:] is the value. Output: y = activation(gate) * value.
Subclass must override ``activation_func`` with a static method.
Args:
M: Number of rows.
N: Half the column dimension (output width).
dtype: Torch dtype.
config: Optional dict with "strategy", "threads" and "num_per_thread".
"strategy" is one of "direct", "explicit_parallel"; it selects
the kernel body at build time.
tune: Whether to autotune (sweeps "threads" / "num_per_thread"
within the resolved strategy).
"""
supported_archs: list[int] = [80, 86, 89, 90]
STRATEGIES = ["direct", "explicit_parallel"]
# Benchmark (H200, 4096x4096 fp16): explicit_parallel ~2x faster than direct
# silu_and_mul: 3.04 TB/s explicit vs 1.50 TB/s direct
# gelu_and_mul: 2.72 TB/s explicit vs 1.47 TB/s direct
# gelu_tanh_and_mul: 3.38 TB/s explicit vs 1.51 TB/s direct
DEFAULT_STRATEGY = "explicit_parallel"
SUPPORTED_DTYPES = None # Subclass override to restrict input dtypes
@staticmethod
def activation_func(x):
"""Activation function. Must be overridden by subclass."""
raise NotImplementedError
def __init__(self, M, N, dtype, config=None, tune=False):
super().__init__()
if self.SUPPORTED_DTYPES is not None and dtype not in self.SUPPORTED_DTYPES:
supported = ", ".join(str(dt) for dt in self.SUPPORTED_DTYPES)
raise ValueError(
f"{self.__class__.__name__} only supports dtypes [{supported}], got {dtype}"
)
self.M = M
self.N = N
self.dtype = dtype
self._fp8_output_dtype = None
self._kernel_output_dtype = None
if _is_fp8(dtype) and _fp8_needs_nonsaturating_cast(dtype):
self._kernel_output_dtype = _fp8_accum_dtype_str()
self._fp8_output_dtype = dtype
self.output_dtype = torch.float16
else:
self.output_dtype = dtype
self.strategy = (config or {}).get("strategy") or self.DEFAULT_STRATEGY
if self.strategy not in self.STRATEGIES:
raise ValueError(
f"Unknown strategy '{self.strategy}', expected one of {self.STRATEGIES}"
)
self.kernel = self._build_kernel(self.strategy)
self.init_config(config, tune)
def _get_effective_op_func(self):
"""Return compound op ``(gate, value) -> activation(gate) * value``.
Delegates to the shared ``_wrap_fp8_accumulation`` helper (arity=2)
so that fp8 cast-in / cast-out logic is centralised.
"""
act = self.activation_func
def fused_op(gate, value):
return act(gate) * value
return _wrap_fp8_accumulation(fused_op, self.dtype, self.dtype_str, arity=2)
def _build_kernel(self, strategy):
cfg = self.default_config
effective_op = self._get_effective_op_func()
if strategy == "direct":
return _make_fused_gated_direct(
self.M, self.N, self.dtype_str, effective_op,
threads=cfg["threads"],
output_dtype=self._kernel_output_dtype,
)
elif strategy == "explicit_parallel":
return _make_fused_gated_explicit(
self.M, self.N, self.dtype_str, effective_op,
cfg["threads"], cfg["num_per_thread"],
output_dtype=self._kernel_output_dtype,
)
else:
raise ValueError(f"Unknown strategy: {strategy}")
@property
def default_config(self) -> dict:
if _is_fp8(self.dtype):
return {"strategy": self.strategy, "threads": 256, "num_per_thread": 16}
if self.strategy == "explicit_parallel" and self.dtype in (torch.float16, torch.bfloat16):
# 128x8 keeps block_N=1024 but widens loads to 128-bit and lifts occupancy.
# Only fp16/bf16 gain the width: fp32 npt=4 already saturates LDG.128.
return {"strategy": self.strategy, "threads": 128, "num_per_thread": 8}
npt = _strategy_npt(self.strategy, self.dtype)
return {"strategy": self.strategy, "threads": 256, "num_per_thread": npt}
@property
def autotune_configs(self) -> list[dict]:
"""Search space: threads in {128, 256, 512} x num_per_thread in {2, 4, 8}.
Covers a range of occupancy/register-pressure tradeoffs for
bandwidth-bound fused gated elementwise kernels.
"""
if _is_fp8(self.dtype):
# fp8 needs 128-bit alignment: npt >= 16 for 1-byte elements
threads_opts = [128, 256, 512]
npt_opts = [16, 32]
else:
# fp16 / bf16 / fp32
threads_opts = [128, 256, 512]
npt_opts = [2, 4, 8]
return [
{"threads": t, "num_per_thread": n}
for t in threads_opts
for n in npt_opts
]
def autotune(self, warmup: int = 10, rep: int = 10) -> None:
"""Override to handle serialization failures in the TileLang autotuner.
FusedGatedKernel JIT functions capture activation_func closures that
the autotuner subprocess cannot serialize. Catch the error and fall
back to the default config so that ``tune=True`` never crashes.
"""
import warnings
try:
super().autotune(warmup=warmup, rep=rep)
except (AssertionError, Exception) as exc:
if "not serializable" in str(exc) or "pickle" in str(exc).lower():
warnings.warn(
f"{self.__class__.__name__} autotuning failed "
f"(activation_func is not serializable); falling back to "
f"default_config.",
stacklevel=2,
)
self.config = dict(self.default_config)
else:
raise
def init_config(self, config=None, tune=False):
"""Override to cache the compiled kernel function after config is set."""
super().init_config(config, tune)
self.config["strategy"] = self.strategy
# Pre-compile and cache the kernel function for the chosen config
# to avoid JIT lookup overhead on every forward() call.
cfg = self.config
if self.strategy == "direct":
self._compiled_fn = self.kernel(cfg["threads"])
else:
self._compiled_fn = self.kernel(cfg["threads"], cfg["num_per_thread"])
def forward(self, x):
result = self._compiled_fn(x)
if self._fp8_output_dtype is not None:
result = result.to(self._fp8_output_dtype)
return result
# Concrete kernel subclasses
class FloatUnaryKernel(UnaryKernel):
"""Unary kernel base for float-only elementwise ops."""
SUPPORTED_DTYPES = _FLOAT_DTYPES
class FloatPredicateKernel(FloatUnaryKernel):
"""Unary kernel base for float predicates with bool output."""
DEFAULT_STRATEGY = "explicit_parallel"
OUTPUT_DTYPE = torch.bool
class LogicalUnaryKernel(UnaryKernel):
"""Unary kernel base for logical predicates with bool output."""
DEFAULT_STRATEGY = "explicit_parallel"
SUPPORTED_DTYPES = _LOGICAL_DTYPES
OUTPUT_DTYPE = torch.bool
class _Uint8StorageUnaryKernel(UnaryKernel):
"""Unary bool-storage kernel: public bool tensors are viewed as uint8."""
DEFAULT_STRATEGY = "register_copy"
SUPPORTED_DTYPES = (torch.uint8,)
@property
def default_config(self) -> dict:
return {"strategy": self.strategy, "threads": 256, "num_per_thread": 16}
class _Uint8StorageBinaryKernel(BinaryKernel):
"""Binary bool-storage kernel: public bool tensors are viewed as uint8."""
DEFAULT_STRATEGY = "explicit_parallel"
SUPPORTED_DTYPES = (torch.uint8,)
@property
def default_config(self) -> dict:
return {"strategy": self.strategy, "threads": 256, "num_per_thread": 16}
class ReluFwdKernel(FloatUnaryKernel):
"""ReLU: y = max(x, 0)."""
@staticmethod
def op_func(x):
return T.if_then_else(x > T.cast(0, x.dtype), x, T.cast(0, x.dtype))
class _AlphaScaledBinaryKernel(BinaryKernel):
"""Shared base for ``y = a (op) alpha * b`` kernels.
Subclasses set ``_combine`` to either addition or subtraction. ``alpha``
is baked in at kernel construction time (one specialization per distinct
``alpha`` value, matching the lru_cache key shape used by the binary
builders) so the kernel surface stays scalar-free.
"""
@staticmethod
def _combine(a_scaled, b_scaled):
raise NotImplementedError
@staticmethod
def op_func(a, b):
raise NotImplementedError(
"_AlphaScaledBinaryKernel uses a per-instance op_func built from "
"alpha; use the kernel via __init__ instead of calling op_func."
)
def __init__(
self, N_total, dtype, coalesced_shape, a_strides, b_strides,
a_numel, b_numel, config=None, tune=False, alpha=1,
):
# PyTorch rejects a floating alpha on an integral input; mirror that so
# the kernel cannot silently truncate alpha through an fp32 cast.
# Out-of-range integer alphas are NOT rejected — PyTorch wraps them via
# the input dtype (uint8 alpha=-1 → 255), which T.cast reproduces.
if dtype in _BITWISE_DTYPES and float(alpha) != float(int(alpha)):
raise ValueError(
"alpha must be an integer when input dtype is integral"
)
self._alpha = alpha
super().__init__(
N_total, dtype, coalesced_shape, a_strides, b_strides,
a_numel, b_numel, config=config, tune=tune,
)
def _alpha_op_func(self):
"""Build a binary op_func with ``alpha`` baked in.
Floating inputs route the scalar multiply through fp32 to dodge
narrow-type literal issues for fp16 / bf16; integer/bool inputs
keep native integer arithmetic. Following PyTorch, the integral
alpha is coerced via the input dtype, so out-of-range values
wrap silently (uint8 alpha=-1 -> 255; bool alpha=2 -> low-bit).
"""
alpha = self._alpha
combine = type(self)._combine
if alpha == 1:
# Identity multiplier: skip the scalar multiply so the kernel
# stays byte-identical to the pre-alpha fast path.
def op_func(a, b):
return combine(a, b)
return op_func
if self.dtype in _BITWISE_DTYPES:
# Native integer arithmetic. Coerce alpha into the input dtype's
# representable range in Python before T.cast: TVM rejects a
# negative literal cast to an unsigned dtype, so reproduce
# PyTorch's "scalar wraps via the input dtype" semantics here.
if self.dtype is torch.bool:
int_alpha = int(bool(alpha))
else:
info = torch.iinfo(self.dtype)
width = info.max - info.min + 1
int_alpha = int(alpha)
if int_alpha < info.min or int_alpha > info.max:
int_alpha = ((int_alpha - info.min) % width) + info.min
def op_func(a, b):
scaled_b = T.cast(int_alpha, a.dtype) * b
return combine(a, scaled_b)
return op_func
def op_func(a, b):
scaled_b = T.cast(T.cast(alpha, "float32") * T.cast(b, "float32"), a.dtype)
return combine(a, scaled_b)
return op_func
def _get_effective_op_func(self):
"""Inject the alpha-baked op_func into the parent build pipeline."""
op_func = self._alpha_op_func()
if self.OUTPUT_DTYPE is not None:
return op_func
return _wrap_fp8_accumulation(op_func, self.dtype, self.dtype_str, arity=2)
class AddFwdKernel(_AlphaScaledBinaryKernel):
"""Element-wise addition with scalar alpha: y = a + alpha * b."""
SUPPORTED_DTYPES = _BINARY_FULL_DTYPES
@staticmethod
def _combine(a, scaled_b):
return a + scaled_b
class SubFwdKernel(_AlphaScaledBinaryKernel):
"""Element-wise subtraction with scalar alpha: y = a - alpha * b."""
SUPPORTED_DTYPES = _BINARY_NO_BOOL_DTYPES
@staticmethod
def _combine(a, scaled_b):
return a - scaled_b
class MulFwdKernel(BinaryKernel):
"""Element-wise multiplication: y = a * b.
Supports the manifest dtype union (bool / unsigned / signed integer /
half / single precision floats). Bool multiplication is logical AND
(PyTorch semantics).
"""
SUPPORTED_DTYPES = _BINARY_FULL_DTYPES
@staticmethod
def op_func(a, b):
return a * b
class DivFwdKernel(BinaryKernel):
"""Element-wise division: y = a / b."""
SUPPORTED_DTYPES = _FLOAT_DTYPES
@staticmethod
def op_func(a, b):
return a / b
class DivTruncFwdKernel(BinaryKernel):
"""Element-wise truncated division: y = trunc(a / b).
Matches ``torch.div(a, b, rounding_mode="trunc")`` semantics: rounds
the quotient toward zero. Division and ``trunc`` are computed in fp32
to avoid two sources of error: (1) ``htrunc`` is not available for
``cutlass::half_t`` in CUDA, and (2) fp16 division rounds the
quotient before ``trunc`` sees it.
"""
SUPPORTED_DTYPES = _FLOAT_DTYPES
@staticmethod
def op_func(a, b):
a_f32 = T.cast(a, "float32")
b_f32 = T.cast(b, "float32")
return T.Cast(a.dtype, T.trunc(a_f32 / b_f32))
class RemainderFwdKernel(BinaryKernel):
"""Element-wise remainder: y = a - floor(a / b) * b.
Matches PyTorch remainder semantics for floating-point inputs.
Uses floor-based formula since T.FloorMod requires integer types.
Division and floor are computed in fp32 to avoid two sources of error:
(1) ``hfloor`` is not available for ``cutlass::half_t`` in CUDA, and
(2) fp16 division rounds the quotient before floor sees it (e.g.
2.999... rounds to 3.0 in fp16). The floored quotient is then cast
back to native dtype so the final ``a - floored * b`` matches PyTorch
semantics for the multiply-subtract step.
"""
SUPPORTED_DTYPES = _FLOAT_DTYPES
@staticmethod
def op_func(a, b):
a_f32 = T.cast(a, "float32")
b_f32 = T.cast(b, "float32")
floored = T.Cast(a.dtype, T.floor(a_f32 / b_f32))
return a - floored * b
class PowFwdKernel(BinaryKernel):
"""Element-wise power: y = a ** b."""
SUPPORTED_DTYPES = _FLOAT_DTYPES
@staticmethod
def op_func(a, b):
a_f32 = T.Cast("float32", a)
b_f32 = T.Cast("float32", b)
return T.Cast(a.dtype, T.pow(a_f32, b_f32))
class FloorDivideFwdKernel(BinaryKernel):
"""Element-wise floor division: y = floor(a / b).
Division and floor are computed in fp32 to avoid two sources of error:
(1) ``hfloor`` is not available for ``cutlass::half_t`` in CUDA, and
(2) fp16 division rounds the quotient before floor sees it (e.g.
2.999... rounds to 3.0 in fp16, giving floor=3 instead of 2).
"""
SUPPORTED_DTYPES = _FLOAT_DTYPES
@staticmethod
def op_func(a, b):
a_f32 = T.cast(a, "float32")
b_f32 = T.cast(b, "float32")
return T.Cast(a.dtype, T.floor(a_f32 / b_f32))
class LerpFwdKernel(BinaryKernel):
"""Element-wise lerp: y = a + weight * (b - a).
PyTorch lerp is ternary (a, b, weight). Here weight is a compile-time
constant passed at kernel construction, keeping the binary kernel template.
Args:
weight: Scalar interpolation weight (default 0.5).
"""
SUPPORTED_DTYPES = _FLOAT_DTYPES
@staticmethod
def op_func(a, b):
raise NotImplementedError("Use _make_lerp_op_func(weight) instead")
def __init__(
self, N_total, dtype, coalesced_shape, a_strides, b_strides,
a_numel, b_numel, config=None, tune=False, weight=0.5,
):
self._weight = weight
super().__init__(
N_total, dtype, coalesced_shape, a_strides, b_strides,
a_numel, b_numel, config=config, tune=tune,
)
def _build_kernel(self, strategy):
"""Override to inject compile-time weight into op_func."""
w = self._weight
def lerp_func(a, b):
return a + T.cast(w, a.dtype) * (b - a)
# Wrap with fp8 accumulation via shared helper
effective_op = _wrap_fp8_accumulation(
lerp_func, self.dtype, self.dtype_str, arity=2,
)
# For e5m2: kernel output is fp16 (non-saturating path)
kernel_output_dtype = (
self.dtype_to_str(self.OUTPUT_DTYPE) if self.OUTPUT_DTYPE is not None else None
)
if self._fp8_output_dtype is not None:
kernel_output_dtype = _fp8_accum_dtype_str()
cfg = self.default_config
if strategy == "direct":
return _make_binary_direct(
self.N_total, self.dtype_str, effective_op,
self.coalesced_shape, self.a_strides, self.b_strides,
self.a_numel, self.b_numel,
output_dtype=kernel_output_dtype, threads=cfg["threads"],
)
elif strategy == "explicit_parallel":
return _make_binary_explicit(
self.N_total, self.dtype_str, effective_op,
self.coalesced_shape, self.a_strides, self.b_strides,
self.a_numel, self.b_numel,
output_dtype=kernel_output_dtype,
threads=cfg["threads"], num_per_thread=cfg["num_per_thread"],
)
elif strategy == "register_copy":
return _make_binary_register_copy(
self.N_total, self.dtype_str, effective_op,
output_dtype=kernel_output_dtype,
threads=cfg["threads"], num_per_thread=cfg["num_per_thread"],
)
else:
raise ValueError(f"Unknown strategy: {strategy}")
def _is_float_dtype_str(dtype_str: str) -> bool:
"""Return True for floating-point TileLang dtype strings.
TileLang IR exposes operand dtypes only as strings (``"float16"``,
``"bfloat16"``, ``"float32"``, ``"float8_e4m3fn"`` ...), so prefix
matching is the established convention for float detection inside
``op_func`` kernel bodies. All TileLang float dtype names start
with ``"float"`` or ``"bfloat"``; integer / bool dtype names
(``"int*"``, ``"uint*"``, ``"bool"``) do not.
"""
return dtype_str.startswith(("float", "bfloat"))
class MaximumFwdKernel(BinaryKernel):
"""Element-wise maximum: y = max(a, b).
For float dtypes, matches torch.maximum semantics:
- If either operand is NaN, the result is NaN.
- maximum(+0.0, -0.0) = +0.0 (IEEE 754 signed-zero).
For integer / bool dtypes (no NaN representation), uses ``T.max``
directly without the NaN guards.
Performance (float path): uses T.max for the fast path (correct
signed-zero on CUDA -- fmaxf returns +0 for max(+0,-0)) plus two
isnan guards for NaN propagation. Total IR: 1 max + 2 fp32 casts +
2 isnan + 2 select.
"""
SUPPORTED_DTYPES = _BINARY_FULL_DTYPES
@staticmethod
def op_func(a, b):
result = T.max(a, b)
if not _is_float_dtype_str(str(a.dtype)):
# Integer / bool: no NaN representation, T.max is sufficient.
return result
# Float path: T.max handles signed-zero correctly but does NOT
# propagate NaN -- it returns the non-NaN operand. Cast to fp32
# for isnan (bfloat16 lacks native isnan).
a_is_nan = T.isnan(T.Cast("float32", a))
b_is_nan = T.isnan(T.Cast("float32", b))
result = T.if_then_else(b_is_nan, b, result)
result = T.if_then_else(a_is_nan, a, result)
return result
class MinimumFwdKernel(BinaryKernel):
"""Element-wise minimum: y = min(a, b).
For float dtypes, matches torch.minimum semantics:
- If either operand is NaN, the result is NaN.
- minimum(-0.0, +0.0) = -0.0 (IEEE 754 signed-zero).
For integer / bool dtypes (no NaN representation), uses ``T.min``
directly without the NaN guards.
Performance (float path): uses T.min for the fast path (correct
signed-zero on CUDA -- fminf returns -0 for min(-0,+0)) plus two
isnan guards for NaN propagation. See MaximumFwdKernel for full
rationale.
"""
SUPPORTED_DTYPES = _BINARY_FULL_DTYPES
@staticmethod
def op_func(a, b):
result = T.min(a, b)
if not _is_float_dtype_str(str(a.dtype)):
return result
a_is_nan = T.isnan(T.Cast("float32", a))
b_is_nan = T.isnan(T.Cast("float32", b))
result = T.if_then_else(b_is_nan, b, result)
result = T.if_then_else(a_is_nan, a, result)
return result
# Comparison kernel subclasses (bool output)
class EqFwdKernel(BinaryKernel):
"""Element-wise equality: y = (a == b)."""
SUPPORTED_DTYPES = _BINARY_FULL_DTYPES
OUTPUT_DTYPE = torch.bool
DEFAULT_STRATEGY = "explicit_parallel"
@staticmethod
def op_func(a, b):
return a == b
class EqBoolStorageFwdKernel(_Uint8StorageBinaryKernel):
"""Element-wise equality on uint8-backed bool storage."""
@staticmethod
def op_func(a, b):
return T.bitwise_xor(T.bitwise_xor(a, b), T.cast(1, "uint8"))
class NeFwdKernel(BinaryKernel):
"""Element-wise not-equal: y = (a != b)."""
SUPPORTED_DTYPES = _BINARY_FULL_DTYPES
OUTPUT_DTYPE = torch.bool
DEFAULT_STRATEGY = "explicit_parallel"
@staticmethod
def op_func(a, b):
return a != b
class NeBoolStorageFwdKernel(_Uint8StorageBinaryKernel):
"""Element-wise not-equal on uint8-backed bool storage."""
@staticmethod
def op_func(a, b):
return T.bitwise_xor(a, b)
class GtFwdKernel(BinaryKernel):
"""Element-wise greater-than: y = (a > b)."""
SUPPORTED_DTYPES = _BINARY_FULL_DTYPES
OUTPUT_DTYPE = torch.bool
DEFAULT_STRATEGY = "explicit_parallel"
@staticmethod
def op_func(a, b):
return a > b
class GtBoolStorageFwdKernel(_Uint8StorageBinaryKernel):
"""Element-wise greater-than on uint8-backed bool storage."""
@staticmethod
def op_func(a, b):
return T.bitwise_and(a, T.bitwise_xor(b, T.cast(1, "uint8")))
class LtFwdKernel(BinaryKernel):
"""Element-wise less-than: y = (a < b)."""
SUPPORTED_DTYPES = _BINARY_FULL_DTYPES
OUTPUT_DTYPE = torch.bool
DEFAULT_STRATEGY = "explicit_parallel"
@staticmethod
def op_func(a, b):
return a < b
class LtBoolStorageFwdKernel(_Uint8StorageBinaryKernel):
"""Element-wise less-than on uint8-backed bool storage."""
@staticmethod
def op_func(a, b):
return T.bitwise_and(T.bitwise_xor(a, T.cast(1, "uint8")), b)
class GeFwdKernel(BinaryKernel):
"""Element-wise greater-equal: y = (a >= b)."""
SUPPORTED_DTYPES = _BINARY_FULL_DTYPES
OUTPUT_DTYPE = torch.bool
DEFAULT_STRATEGY = "explicit_parallel"
@staticmethod
def op_func(a, b):
return a >= b
class GeBoolStorageFwdKernel(_Uint8StorageBinaryKernel):
"""Element-wise greater-equal on uint8-backed bool storage."""
@staticmethod
def op_func(a, b):
return T.bitwise_or(a, T.bitwise_xor(b, T.cast(1, "uint8")))
class LeFwdKernel(BinaryKernel):
"""Element-wise less-equal: y = (a <= b)."""
SUPPORTED_DTYPES = _BINARY_FULL_DTYPES
OUTPUT_DTYPE = torch.bool
DEFAULT_STRATEGY = "explicit_parallel"
@staticmethod
def op_func(a, b):
return a <= b
class LeBoolStorageFwdKernel(_Uint8StorageBinaryKernel):
"""Element-wise less-equal on uint8-backed bool storage."""
@staticmethod
def op_func(a, b):
return T.bitwise_or(T.bitwise_xor(a, T.cast(1, "uint8")), b)
# Logical kernel subclasses (bool output)
class LogicalAndFwdKernel(BinaryKernel):
"""Element-wise logical AND with non-zero truthiness."""
SUPPORTED_DTYPES = _LOGICAL_DTYPES
OUTPUT_DTYPE = torch.bool
DEFAULT_STRATEGY = "explicit_parallel"
@staticmethod
def op_func(a, b):
a_nonzero = a != T.cast(0, a.dtype)
b_nonzero = b != T.cast(0, b.dtype)
return a_nonzero & b_nonzero
class LogicalAndBoolStorageFwdKernel(_Uint8StorageBinaryKernel):
"""Element-wise logical AND on uint8-backed bool storage."""
@staticmethod
def op_func(a, b):
return T.bitwise_and(a, b)
class LogicalOrFwdKernel(BinaryKernel):
"""Element-wise logical OR with non-zero truthiness."""
SUPPORTED_DTYPES = _LOGICAL_DTYPES
OUTPUT_DTYPE = torch.bool
DEFAULT_STRATEGY = "explicit_parallel"
@staticmethod
def op_func(a, b):
a_nonzero = a != T.cast(0, a.dtype)
b_nonzero = b != T.cast(0, b.dtype)
return a_nonzero | b_nonzero
class LogicalOrBoolStorageFwdKernel(_Uint8StorageBinaryKernel):
"""Element-wise logical OR on uint8-backed bool storage."""
@staticmethod
def op_func(a, b):
return T.bitwise_or(a, b)
# Bitwise kernel subclasses
class BitwiseAndFwdKernel(BinaryKernel):
"""Element-wise bitwise AND: y = a & b (integer inputs)."""
SUPPORTED_DTYPES = _BITWISE_DTYPES
@staticmethod
def op_func(a, b):
return a & b
class BitwiseAndBoolStorageFwdKernel(_Uint8StorageBinaryKernel):
"""Element-wise bitwise AND on uint8-backed bool storage."""
@staticmethod
def op_func(a, b):
return T.bitwise_and(a, b)
class BitwiseOrFwdKernel(BinaryKernel):
"""Element-wise bitwise OR: y = a | b (integer inputs)."""
SUPPORTED_DTYPES = _BITWISE_DTYPES
@staticmethod
def op_func(a, b):
return a | b
class BitwiseOrBoolStorageFwdKernel(_Uint8StorageBinaryKernel):
"""Element-wise bitwise OR on uint8-backed bool storage."""
@staticmethod
def op_func(a, b):
return T.bitwise_or(a, b)
class BitwiseXorFwdKernel(BinaryKernel):
"""Element-wise bitwise XOR: y = a ^ b (integer inputs)."""
SUPPORTED_DTYPES = _BITWISE_DTYPES
@staticmethod
def op_func(a, b):
return a ^ b
class BitwiseXorBoolStorageFwdKernel(_Uint8StorageBinaryKernel):
"""Element-wise bitwise XOR on uint8-backed bool storage."""
@staticmethod
def op_func(a, b):
return T.bitwise_xor(a, b)
# Fused gated kernel subclasses
class SiluAndMulFwdKernel(FusedGatedKernel):
"""SiLU-and-Mul: y = silu(gate) * value = (gate * sigmoid(gate)) * value."""
SUPPORTED_DTYPES = _FLOAT_DTYPES
@staticmethod
def activation_func(x):
# exp2 form (fp32): exp2 lowers to one MUFU.EX2 vs expf's multi-op sequence.
g = T.Cast("float32", x)
one = T.cast(1.0, "float32")
log2e = T.cast(1.4426950408889634, "float32")
return g / (one + T.exp2(-g * log2e))
class GeluAndMulFwdKernel(FusedGatedKernel):
"""GELU-and-Mul: y = gelu(gate) * value.
Uses exact GELU: gelu(x) = x * 0.5 * (1 + erf(x / sqrt(2))).
erf is computed in float32 to avoid missing half-precision intrinsic.
"""
SUPPORTED_DTYPES = _FLOAT_DTYPES
@staticmethod
def activation_func(x):
inv_sqrt2 = T.cast(0.7071067811865476, "float32") # 1/sqrt(2)
half = T.cast(0.5, x.dtype)
one = T.cast(1.0, x.dtype)
x_f32 = T.Cast("float32", x)
erf_val = T.Cast(x.dtype, T.erf(x_f32 * inv_sqrt2))
return x * half * (one + erf_val)
class GeluTanhAndMulFwdKernel(FusedGatedKernel):
"""GELU-Tanh-and-Mul: y = gelu_tanh(gate) * value.
Uses tanh approximation: gelu(x) = 0.5 * x * (1 + tanh(sqrt(2/pi) * (x + 0.044715 * x^3))).
"""
SUPPORTED_DTYPES = _FLOAT_DTYPES
@staticmethod
def activation_func(x):
sqrt_2_over_pi = T.cast(0.7978845608028654, "float32") # sqrt(2/pi)
coeff = T.cast(0.044715, "float32") # GELU tanh approx coefficient
half = T.cast(0.5, x.dtype)
one = T.cast(1.0, x.dtype)
x_f32 = T.Cast("float32", x)
inner = sqrt_2_over_pi * (x_f32 + coeff * x_f32 * x_f32 * x_f32)
tanh_val = T.Cast(x.dtype, T.tanh(inner))
return half * x * (one + tanh_val)
# Concrete unary kernel subclasses -- math (17)
class ExpFwdKernel(FloatUnaryKernel):
"""Element-wise exp(x)."""
@staticmethod
def op_func(x):
return T.exp(T.cast(x, "float32"))
class LogFwdKernel(FloatUnaryKernel):
"""Element-wise log(x)."""
@staticmethod
def op_func(x):
return T.log(T.cast(x, "float32"))
class SqrtFwdKernel(FloatUnaryKernel):
"""Element-wise sqrt(x)."""
@staticmethod
def op_func(x):
return T.sqrt(T.cast(x, "float32"))
class RsqrtFwdKernel(FloatUnaryKernel):
"""Element-wise 1/sqrt(x)."""
@staticmethod
def op_func(x):
return T.rsqrt(T.cast(x, "float32"))
class AbsFwdKernel(FloatUnaryKernel):
"""Element-wise |x|."""
@staticmethod
def op_func(x):
return T.abs(x)
class NegFwdKernel(FloatUnaryKernel):
"""Element-wise -x."""
@staticmethod
def op_func(x):
return -x
class ReciprocalFwdKernel(FloatUnaryKernel):
"""Element-wise 1/x."""
@staticmethod
def op_func(x):
return T.cast(1.0, "float32") / x
class SignFwdKernel(FloatUnaryKernel):
"""Element-wise sign(x): -1, 0, or +1."""
@staticmethod
def op_func(x):
zero = T.cast(0.0, x.dtype)
one = T.cast(1.0, x.dtype)
neg_one = T.cast(-1.0, x.dtype)
return T.if_then_else(
x > zero,
one,
T.if_then_else(x < zero, neg_one, zero),
)
class SinFwdKernel(FloatUnaryKernel):
"""Element-wise sin(x)."""
@staticmethod
def op_func(x):
return T.sin(T.cast(x, "float32"))
class CosFwdKernel(FloatUnaryKernel):
"""Element-wise cos(x)."""
@staticmethod
def op_func(x):
return T.cos(T.cast(x, "float32"))
class FloorFwdKernel(FloatUnaryKernel):
"""Element-wise floor(x).
Casts to fp32 before calling ``T.floor`` because ``hfloor`` is not
available for ``cutlass::half_t`` in CUDA.
"""
@staticmethod
def op_func(x):
return T.floor(T.cast(x, "float32"))
class CeilFwdKernel(FloatUnaryKernel):
"""Element-wise ceil(x).
Casts to fp32 before calling ``T.ceil`` because ``hceil`` is not
available for ``cutlass::half_t`` in CUDA.
"""
@staticmethod
def op_func(x):
return T.ceil(T.cast(x, "float32"))
class RoundFwdKernel(FloatUnaryKernel):
"""Element-wise round(x) with banker's rounding (round-to-nearest-even).
Uses ``T.nearbyint`` (maps to ``nearbyintf`` in CUDA) to match
PyTorch's ``torch.round`` semantics. Casts to fp32 because
``hnearbyint`` is not available for ``cutlass::half_t``.
"""
@staticmethod
def op_func(x):
return T.nearbyint(T.cast(x, "float32"))
class TruncFwdKernel(FloatUnaryKernel):
"""Element-wise trunc(x) -- integer part toward zero.
Casts to fp32 before calling ``T.trunc`` because ``htrunc`` is not
available for ``cutlass::half_t`` in CUDA.
"""
@staticmethod
def op_func(x):
return T.trunc(T.cast(x, "float32"))
class ErfFwdKernel(FloatUnaryKernel):
"""Element-wise erf(x).
Casts to fp32 before calling ``T.erf`` because the half-precision
intrinsic ``herf`` is not a valid CUDA built-in.
"""
@staticmethod
def op_func(x):
return T.erf(T.cast(x, "float32"))
class Log1pFwdKernel(FloatUnaryKernel):
"""Element-wise log(1 + x).
Uses composite ``log(1 + x)`` because ``T.log1p`` is not lowered
by the TileLang compiler.
"""
@staticmethod
def op_func(x):
return T.log(T.cast(1.0, "float32") + x)
class Expm1FwdKernel(FloatUnaryKernel):
"""Element-wise exp(x) - 1."""
@staticmethod
def op_func(x):
return T.exp(T.cast(x, "float32")) - T.cast(1.0, "float32")
# Concrete unary kernel subclasses -- activations (9)
class GeluFwdKernel(FloatUnaryKernel):
"""Element-wise GELU using the standard erf formulation."""
@staticmethod
def op_func(x):
inv_sqrt_2 = T.cast(0.7071067811865476, "float32")
half = T.cast(0.5, "float32")
one = T.cast(1.0, "float32")
return half * x * (one + T.erf(T.cast(x, "float32") * inv_sqrt_2))
class GeluTanhFwdKernel(FloatUnaryKernel):
"""Element-wise GELU using the tanh approximation.
Computes ``0.5 * x * (1 + tanh(sqrt(2/pi) * (x + 0.044715 * x^3)))``,
matching ``torch.nn.functional.gelu(x, approximate='tanh')``.
"""
@staticmethod
def op_func(x):
sqrt_2_over_pi = T.cast(0.7978845608028654, "float32")
coeff = T.cast(0.044715, "float32")
half = T.cast(0.5, "float32")
one = T.cast(1.0, "float32")
x_f32 = T.cast(x, "float32")
inner = sqrt_2_over_pi * (x_f32 + coeff * x_f32 * x_f32 * x_f32)
return half * x_f32 * (one + T.tanh(inner))
class SiluFwdKernel(FloatUnaryKernel):
"""Element-wise SiLU (Swish): x * sigmoid(x)."""
@staticmethod
def op_func(x):
return x * T.sigmoid(x)
class SigmoidFwdKernel(FloatUnaryKernel):
"""Element-wise sigmoid(x)."""
@staticmethod
def op_func(x):
return T.sigmoid(x)
class TanhFwdKernel(FloatUnaryKernel):
"""Element-wise tanh(x)."""
@staticmethod
def op_func(x):
return T.tanh(T.cast(x, "float32"))
class HardswishFwdKernel(FloatUnaryKernel):
"""Element-wise HardSwish: x * clamp(x + 3, 0, 6) / 6."""
@staticmethod
def op_func(x):
three = T.cast(3.0, "float32")
six = T.cast(6.0, "float32")
zero = T.cast(0.0, "float32")
clamped = T.min(T.max(x + three, zero), six)
return x * clamped / six
class HardsigmoidFwdKernel(FloatUnaryKernel):
"""Element-wise HardSigmoid: clamp(x + 3, 0, 6) / 6."""
@staticmethod
def op_func(x):
three = T.cast(3.0, "float32")
six = T.cast(6.0, "float32")
zero = T.cast(0.0, "float32")
return T.min(T.max(x + three, zero), six) / six
class MishFwdKernel(FloatUnaryKernel):
"""Element-wise Mish: x * tanh(softplus(x)) = x * tanh(log(1 + exp(x)))."""
@staticmethod
def op_func(x):
one = T.cast(1.0, "float32")
return x * T.tanh(T.log(one + T.exp(x)))
class SeluFwdKernel(FloatUnaryKernel):
"""Element-wise SELU: scale * (max(0,x) + min(0, alpha*(exp(x)-1))).
alpha = 1.6732632423543772, scale = 1.0507009873554805
"""
@staticmethod
def op_func(x):
alpha = T.cast(1.6732632423543772, "float32")
scale = T.cast(1.0507009873554805, "float32")
one = T.cast(1.0, "float32")
zero = T.cast(0.0, "float32")
x32 = T.cast(x, "float32")
return scale * T.if_then_else(x32 > zero, x32, alpha * (T.exp(x32) - one))
# Concrete unary kernel subclasses -- logical / bitwise (2)
class LogicalNotFwdKernel(LogicalUnaryKernel):
"""Element-wise logical NOT with torch-style bool output."""
@staticmethod
def op_func(x):
return x == T.cast(0, x.dtype)
class LogicalNotBoolStorageFwdKernel(_Uint8StorageUnaryKernel):
"""Element-wise logical NOT on uint8-backed bool storage."""
@staticmethod
def op_func(x):
return T.bitwise_xor(x, T.cast(1, "uint8"))
class BitwiseNotFwdKernel(UnaryKernel):
"""Element-wise bitwise NOT (~x) for bool/integer inputs.
Uses XOR with ``-1`` (all-ones) because ``T.bitwise_not`` fails on
vectorized ``int4`` CUDA types.
"""
DEFAULT_STRATEGY = "direct"
SUPPORTED_DTYPES = _BITWISE_DTYPES
@staticmethod
def op_func(x):
if x.dtype == "bool":
return x == T.cast(0, "bool")
if x.dtype == "uint8":
return T.bitwise_xor(x, T.cast(255, "uint8"))
return T.bitwise_xor(x, T.cast(-1, x.dtype))
# Concrete unary kernel subclasses -- special predicates (3)
class IsnanFwdKernel(FloatPredicateKernel):
"""Element-wise isnan with torch-style bool output."""
@staticmethod
def op_func(x):
return T.isnan(T.cast(x, "float32"))
class IsinfFwdKernel(FloatPredicateKernel):
"""Element-wise isinf with torch-style bool output."""
@staticmethod
def op_func(x):
return T.isinf(T.cast(x, "float32"))
class IsfiniteFwdKernel(FloatPredicateKernel):
"""Element-wise isfinite with torch-style bool output."""
@staticmethod
def op_func(x):
return T.isfinite(T.cast(x, "float32"))
# Independent (custom-signature) kernel classes (11)
class ParametricUnaryKernel(Kernel):
"""Shared base for independent parametric elementwise kernels.
Subclasses must define:
- ``_builder_fn``: a ``@staticmethod`` returning the ``@lru_cache``-d
builder function (e.g. ``_make_leaky_relu_kernel``).
- ``_builder_args(self) -> tuple``: positional args for the builder
*between* ``N_total`` and the common ``output_dtype, is_fp8, threads,
npt`` suffix.
Optional overrides:
- ``_DEFAULT_THREADS``: class-level default thread count (default 256).
- ``_NPT_FP8``: npt when dtype is fp8 but not fp32 (default 16).
- ``_NPT_NON_FP32``: npt for non-fp32, non-fp8 (default 8).
- ``_skip_fp8_output``: set to ``True`` if the kernel should *not*
use ``_get_fp8_output_dtypes`` (e.g. Where, which is a pure selection
op). When True, ``_fp8_output_dtype`` is ``None``.
"""
supported_archs: list[int] = [80, 86, 89, 90]
SUPPORTED_DTYPES = _FLOAT_DTYPES
_DEFAULT_THREADS: int = 256
_NPT_FP8: int = 16
_NPT_NON_FP32: int = 8
_skip_fp8_output: bool = False
def __init__(self, N_total, dtype, config=None, tune=False):
super().__init__()
if dtype not in self.SUPPORTED_DTYPES:
supported = ", ".join(str(dt) for dt in self.SUPPORTED_DTYPES)
raise ValueError(
f"{self.__class__.__name__} only supports dtypes [{supported}], got {dtype}"
)
self.N_total = N_total
self.dtype = dtype
# fp8 output handling
if self._skip_fp8_output:
self._fp8_output_dtype = None
else:
self._fp8_output_dtype, self.output_dtype = _get_fp8_output_dtypes(dtype)
# Post-fp8 parameter processing (e.g. clamping scalars to output dtype range)
self._post_init_params()
# Build the kernel via the subclass-provided builder
cfg = self.default_config
builder_kwargs = {
"is_fp8": _is_fp8(dtype),
"threads": cfg["threads"],
"npt": cfg["num_per_thread"],
}
if not self._skip_fp8_output:
builder_kwargs["output_dtype"] = self.dtype_to_str(self.output_dtype)
self.kernel = self._builder_fn()(
*self._builder_positional_args(), **builder_kwargs,
)
self.init_config(config, tune)
@staticmethod
def _builder_fn():
"""Return the @lru_cache builder function for this kernel."""
raise NotImplementedError
def _builder_positional_args(self) -> tuple:
"""Return all positional args for the builder function.
Default: ``(N_total, dtype_str, *_builder_args())``.
Override if the builder has a different parameter order (e.g. PReLU).
"""
return (self.N_total, self.dtype_str, *self._builder_args())
def _builder_args(self) -> tuple:
"""Return op-specific positional args (after N_total, dtype_str)."""
return ()
def _post_init_params(self):
"""Hook called after fp8 output dtypes are set, before kernel build.
Override to clamp scalar parameters to the output dtype range (e.g.
MaskedFill, NanToNum).
"""
@property
def default_config(self):
if self.dtype == torch.float32:
npt = 4
elif _is_fp8(self.dtype):
npt = self._NPT_FP8
else:
npt = self._NPT_NON_FP32
return {"threads": self._DEFAULT_THREADS, "num_per_thread": npt}
def init_config(self, config=None, tune=False):
"""Override to cache the compiled kernel function after config is set."""
super().init_config(config, tune)
cfg = self.config
self._compiled_fn = self.kernel(cfg["threads"], cfg["num_per_thread"])
def forward(self, x):
result = self._compiled_fn(x)
if self._fp8_output_dtype is not None:
result = result.to(self._fp8_output_dtype)
return result
@functools.lru_cache(maxsize=32)
def _make_leaky_relu_kernel(N, dtype, negative_slope, output_dtype=None,
is_fp8=False, threads=256, npt=8):
"""Build leaky_relu kernel: y = x if x > 0 else negative_slope * x.
For non-fp8 dtypes, uses register_copy strategy: fragment load -> compute
-> fragment store for coalesced memory access.
For fp8 dtypes, uses explicit_parallel with fp16 accumulation (register_copy
is unreliable for 8-bit fragments).
"""
out_dtype = output_dtype or dtype
block_size = threads * npt
if is_fp8:
accum = _fp8_accum_dtype_str()
@tilelang.jit(out_idx=[1])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(x: T.Tensor((N,), dtype), y: T.Tensor((N,), out_dtype)):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
for i, j in T.Parallel(threads_arg, npt_arg):
idx = (bx * threads_arg + i) * npt_arg + j
if idx < N:
val = x[idx]
v = T.cast(val, accum)
zero = T.cast(0, accum)
slope = T.cast(negative_slope, accum)
result = T.if_then_else(v > zero, v, slope * v)
y[idx] = T.Cast(out_dtype, result)
return main
else:
@tilelang.jit(out_idx=[1])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(x: T.Tensor((N,), dtype), y: T.Tensor((N,), dtype)):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
x_reg = T.alloc_fragment((block_size,), dtype)
y_reg = T.alloc_fragment((block_size,), dtype)
T.copy(x[bx * block_size : (bx + 1) * block_size], x_reg)
for i, j in T.Parallel(threads_arg, npt_arg):
val = x_reg[i * npt_arg + j]
zero = T.cast(0, val.dtype)
slope = T.cast(negative_slope, val.dtype)
y_reg[i * npt_arg + j] = T.if_then_else(val > zero, val, slope * val)
T.copy(y_reg, y[bx * block_size : (bx + 1) * block_size])
return main
return kernel
class LeakyReluFwdKernel(ParametricUnaryKernel):
"""Leaky ReLU: y = x if x > 0 else negative_slope * x."""
def __init__(self, N_total, dtype, negative_slope=0.01, config=None, tune=False):
self.negative_slope = negative_slope
super().__init__(N_total, dtype, config=config, tune=tune)
@staticmethod
def _builder_fn():
return _make_leaky_relu_kernel
def _builder_args(self):
return (self.negative_slope,)
@functools.lru_cache(maxsize=32)
def _make_elu_kernel(N, dtype, alpha, output_dtype=None, is_fp8=False,
threads=256, npt=8):
"""Build ELU kernel: y = x if x > 0 else alpha * (exp(x) - 1).
"""
out_dtype = output_dtype or dtype
block_size = threads * npt
if is_fp8:
@tilelang.jit(out_idx=[1])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(x: T.Tensor((N,), dtype), y: T.Tensor((N,), out_dtype)):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
for i, j in T.Parallel(threads_arg, npt_arg):
idx = (bx * threads_arg + i) * npt_arg + j
if idx < N:
val = x[idx]
zero = T.cast(0, "float32")
a = T.cast(alpha, "float32")
one = T.cast(1.0, "float32")
v32 = T.cast(val, "float32")
y[idx] = T.if_then_else(v32 > zero, T.Cast(out_dtype, v32), T.Cast(out_dtype, a * (T.exp(v32) - one)))
return main
else:
@tilelang.jit(out_idx=[1])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(x: T.Tensor((N,), dtype), y: T.Tensor((N,), dtype)):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
x_reg = T.alloc_fragment((block_size,), dtype)
y_reg = T.alloc_fragment((block_size,), dtype)
T.copy(x[bx * block_size : (bx + 1) * block_size], x_reg)
for i, j in T.Parallel(threads_arg, npt_arg):
val = x_reg[i * npt_arg + j]
zero = T.cast(0, "float32")
a = T.cast(alpha, "float32")
one = T.cast(1.0, "float32")
v32 = T.cast(val, "float32")
y_reg[i * npt_arg + j] = T.if_then_else(
v32 > zero, val, T.Cast(val.dtype, a * (T.exp(v32) - one)),
)
T.copy(y_reg, y[bx * block_size : (bx + 1) * block_size])
return main
return kernel
class EluFwdKernel(ParametricUnaryKernel):
"""ELU: y = x if x > 0 else alpha * (exp(x) - 1)."""
def __init__(self, N_total, dtype, alpha=1.0, config=None, tune=False):
self.alpha = alpha
super().__init__(N_total, dtype, config=config, tune=tune)
@staticmethod
def _builder_fn():
return _make_elu_kernel
def _builder_args(self):
return (self.alpha,)
@functools.lru_cache(maxsize=32)
def _make_hardtanh_kernel(N, dtype, min_val, max_val, output_dtype=None,
is_fp8=False, threads=256, npt=8):
"""Build hardtanh kernel: y = clamp(x, min_val, max_val).
"""
out_dtype = output_dtype or dtype
block_size = threads * npt
if is_fp8:
@tilelang.jit(out_idx=[1])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(x: T.Tensor((N,), dtype), y: T.Tensor((N,), out_dtype)):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
for i, j in T.Parallel(threads_arg, npt_arg):
idx = (bx * threads_arg + i) * npt_arg + j
if idx < N:
val = x[idx]
lo = T.cast(min_val, "float32")
hi = T.cast(max_val, "float32")
v32 = T.cast(val, "float32")
y[idx] = T.Cast(out_dtype, T.min(T.max(v32, lo), hi))
return main
else:
@tilelang.jit(out_idx=[1])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(x: T.Tensor((N,), dtype), y: T.Tensor((N,), dtype)):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
x_reg = T.alloc_fragment((block_size,), dtype)
y_reg = T.alloc_fragment((block_size,), dtype)
T.copy(x[bx * block_size : (bx + 1) * block_size], x_reg)
for i, j in T.Parallel(threads_arg, npt_arg):
val = x_reg[i * npt_arg + j]
lo = T.cast(min_val, "float32")
hi = T.cast(max_val, "float32")
v32 = T.cast(val, "float32")
y_reg[i * npt_arg + j] = T.Cast(val.dtype, T.min(T.max(v32, lo), hi))
T.copy(y_reg, y[bx * block_size : (bx + 1) * block_size])
return main
return kernel
class HardtanhFwdKernel(ParametricUnaryKernel):
"""Hardtanh: y = clamp(x, min_val, max_val)."""
def __init__(self, N_total, dtype, min_val=-1.0, max_val=1.0, config=None, tune=False):
self.min_val = min_val
self.max_val = max_val
super().__init__(N_total, dtype, config=config, tune=tune)
@staticmethod
def _builder_fn():
return _make_hardtanh_kernel
def _builder_args(self):
return (self.min_val, self.max_val)
@functools.lru_cache(maxsize=32)
def _make_softplus_kernel(N, dtype, beta, threshold, output_dtype=None,
is_fp8=False, threads=256, npt=8):
"""Build softplus kernel: y = log(1 + exp(x*beta))/beta if x*beta <= threshold else x.
"""
out_dtype = output_dtype or dtype
block_size = threads * npt
if is_fp8:
@tilelang.jit(out_idx=[1])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(x: T.Tensor((N,), dtype), y: T.Tensor((N,), out_dtype)):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
for i, j in T.Parallel(threads_arg, npt_arg):
idx = (bx * threads_arg + i) * npt_arg + j
if idx < N:
val = x[idx]
v32 = T.cast(val, "float32")
b = T.cast(beta, "float32")
t = T.cast(threshold, "float32")
one = T.cast(1.0, "float32")
scaled = v32 * b
sp = T.log(one + T.exp(scaled)) / b
y[idx] = T.if_then_else(scaled > t, T.Cast(out_dtype, v32), T.Cast(out_dtype, sp))
return main
else:
@tilelang.jit(out_idx=[1])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(x: T.Tensor((N,), dtype), y: T.Tensor((N,), dtype)):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
x_reg = T.alloc_fragment((block_size,), dtype)
y_reg = T.alloc_fragment((block_size,), dtype)
T.copy(x[bx * block_size : (bx + 1) * block_size], x_reg)
for i, j in T.Parallel(threads_arg, npt_arg):
val = x_reg[i * npt_arg + j]
v32 = T.cast(val, "float32")
b = T.cast(beta, "float32")
t = T.cast(threshold, "float32")
one = T.cast(1.0, "float32")
scaled = v32 * b
sp = T.log(one + T.exp(scaled)) / b
y_reg[i * npt_arg + j] = T.if_then_else(
scaled > t, val, T.Cast(val.dtype, sp),
)
T.copy(y_reg, y[bx * block_size : (bx + 1) * block_size])
return main
return kernel
class SoftplusFwdKernel(ParametricUnaryKernel):
"""Softplus: y = log(1 + exp(x*beta))/beta if x*beta <= threshold else x."""
def __init__(self, N_total, dtype, beta=1.0, threshold=20.0, config=None, tune=False):
self.beta = beta
self.threshold = threshold
super().__init__(N_total, dtype, config=config, tune=tune)
@staticmethod
def _builder_fn():
return _make_softplus_kernel
def _builder_args(self):
return (self.beta, self.threshold)
@functools.lru_cache(maxsize=32)
def _make_prelu_kernel(N, C, inner_size, dtype, output_dtype=None,
is_fp8=False, threads=256, npt=8):
"""Build PReLU kernel: y = x if x > 0 else weight[channel] * x.
Weight is per-channel. Channel index follows PyTorch convention:
for flat index ``idx``, channel = (idx // inner_size) % C, where
``inner_size`` is the product of all dimensions after the channel dim.
For non-fp8 dtypes, uses register_copy strategy for input/output to
improve memory coalescing for the main data path.
"""
out_dtype = output_dtype or dtype
block_size = threads * npt
if is_fp8:
accum = _fp8_accum_dtype_str()
@tilelang.jit(out_idx=[2])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(
x: T.Tensor((N,), dtype),
weight: T.Tensor((C,), dtype),
y: T.Tensor((N,), out_dtype),
):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
for i, j in T.Parallel(threads_arg, npt_arg):
k = i * npt_arg + j
idx = bx * block_size + k
if idx < N:
val = x[idx]
ch = (idx // inner_size) % C
w = weight[ch]
v = T.cast(val, accum)
wf = T.cast(w, accum)
zero = T.cast(0, accum)
y[idx] = T.if_then_else(v > zero, T.Cast(out_dtype, v), T.Cast(out_dtype, wf * v))
return main
else:
@tilelang.jit(out_idx=[2])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(
x: T.Tensor((N,), dtype),
weight: T.Tensor((C,), dtype),
y: T.Tensor((N,), dtype),
):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
x_reg = T.alloc_fragment((block_size,), dtype)
y_reg = T.alloc_fragment((block_size,), dtype)
T.copy(x[bx * block_size : (bx + 1) * block_size], x_reg)
for i, j in T.Parallel(threads_arg, npt_arg):
k = i * npt_arg + j
idx = bx * block_size + k
val = x_reg[k]
ch = (idx // inner_size) % C
w = weight[ch]
zero = T.cast(0, val.dtype)
y_reg[k] = T.if_then_else(val > zero, val, w * val)
T.copy(y_reg, y[bx * block_size : (bx + 1) * block_size])
return main
return kernel
class PreluFwdKernel(ParametricUnaryKernel):
"""PReLU: y = x if x > 0 else weight[channel] * x."""
def __init__(self, N_total, C, inner_size, dtype, config=None, tune=False):
self.C = C
self.inner_size = inner_size
super().__init__(N_total, dtype, config=config, tune=tune)
@staticmethod
def _builder_fn():
return _make_prelu_kernel
def _builder_positional_args(self):
return (self.N_total, self.C, self.inner_size, self.dtype_str)
def forward(self, x, weight):
return self._compiled_fn(x, weight)
@functools.lru_cache(maxsize=32)
def _make_where_kernel(N, dtype, is_fp8=False, threads=256, npt=8):
"""Build where kernel: out = cond ? x : y.
The Op layer packs the bool condition as uint8 so that T.copy can
perform vectorized loads (TileLang does not vectorize bool tensors).
Each uint8 element is 0 or 1; the kernel loads it into a register
fragment and unpacks per-element with a != 0 comparison.
For non-fp8 dtypes, writes the result back into the x register fragment
(in-place) to reduce register pressure and avoid a fourth data-typed
fragment allocation.
"""
block_size = threads * npt
if is_fp8:
@tilelang.jit(out_idx=[3])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(
cond: T.Tensor((N,), "uint8"),
x: T.Tensor((N,), dtype),
y_in: T.Tensor((N,), dtype),
out: T.Tensor((N,), dtype),
):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
for i, j in T.Parallel(threads_arg, npt_arg):
idx = (bx * threads_arg + i) * npt_arg + j
if idx < N:
out[idx] = T.if_then_else(
cond[idx] != T.cast(0, "uint8"), x[idx], y_in[idx],
)
return main
else:
@tilelang.jit(out_idx=[3])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(
cond: T.Tensor((N,), "uint8"),
x: T.Tensor((N,), dtype),
y_in: T.Tensor((N,), dtype),
out: T.Tensor((N,), dtype),
):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
c_reg = T.alloc_fragment((block_size,), "uint8")
x_reg = T.alloc_fragment((block_size,), dtype)
y_reg = T.alloc_fragment((block_size,), dtype)
T.copy(cond[bx * block_size : (bx + 1) * block_size], c_reg)
T.copy(x[bx * block_size : (bx + 1) * block_size], x_reg)
T.copy(y_in[bx * block_size : (bx + 1) * block_size], y_reg)
for i, j in T.Parallel(threads_arg, npt_arg):
k = i * npt_arg + j
x_reg[k] = T.if_then_else(
c_reg[k] != T.cast(0, "uint8"), x_reg[k], y_reg[k],
)
T.copy(x_reg, out[bx * block_size : (bx + 1) * block_size])
return main
return kernel
class WhereFwdKernel(ParametricUnaryKernel):
"""Where: out = cond ? x : y."""
_DEFAULT_THREADS = 512
_skip_fp8_output = True
@staticmethod
def _builder_fn():
return _make_where_kernel
def forward(self, cond, x, y):
return self._compiled_fn(cond, x, y)
@functools.lru_cache(maxsize=32)
def _make_lerp_tensor_kernel(N, dtype, output_dtype=None, is_fp8=False,
threads=256, npt=8):
"""Build Tensor-weight lerp kernel: out = a + weight * (b - a).
The Op layer pre-broadcasts ``input`` / ``end`` / ``weight`` to the
flat output shape so the kernel sees three contiguous 1-D tensors of
size ``N``. Computation is performed in the input dtype for fp16 /
bfloat16 / float32 (the only dtypes the manifest declares); the fp8
path is unreachable here because the kernel's ``SUPPORTED_DTYPES``
excludes fp8.
Uses the register-fragment load -> compute -> fragment store strategy
(matches the non-fp8 ``_make_where_kernel`` layout) so all three
inputs and the output share the same vectorized memory access path.
"""
del is_fp8 # fp8 is not in the manifest contract for this op
out_dtype = output_dtype or dtype
block_size = threads * npt
@tilelang.jit(out_idx=[3])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(
a: T.Tensor((N,), dtype),
b: T.Tensor((N,), dtype),
w: T.Tensor((N,), dtype),
out: T.Tensor((N,), out_dtype),
):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
a_reg = T.alloc_fragment((block_size,), dtype)
b_reg = T.alloc_fragment((block_size,), dtype)
w_reg = T.alloc_fragment((block_size,), dtype)
T.copy(a[bx * block_size : (bx + 1) * block_size], a_reg)
T.copy(b[bx * block_size : (bx + 1) * block_size], b_reg)
T.copy(w[bx * block_size : (bx + 1) * block_size], w_reg)
for i, j in T.Parallel(threads_arg, npt_arg):
k = i * npt_arg + j
a_reg[k] = a_reg[k] + w_reg[k] * (b_reg[k] - a_reg[k])
T.copy(a_reg, out[bx * block_size : (bx + 1) * block_size])
return main
return kernel
class LerpTensorFwdKernel(ParametricUnaryKernel):
"""Tensor-weight lerp: out = input + weight * (end - input).
Implements the Tensor-weight overload of ``torch.lerp`` —
``torch.lerp(input, end, weight: Tensor)`` — where all three operands
are float tensors of the same dtype broadcast together by the Op
layer to a flat ``N``-element view.
Manifest declares ``float16 | bfloat16 | float32``; fp8 is rejected
at construction. The Op layer is responsible for broadcasting the
three inputs to ``N_total`` before dispatch.
"""
SUPPORTED_DTYPES = (torch.float16, torch.bfloat16, torch.float32)
_DEFAULT_THREADS = 512
_skip_fp8_output = True
@staticmethod
def _builder_fn():
return _make_lerp_tensor_kernel
def forward(self, a, b, w):
return self._compiled_fn(a, b, w)
@functools.lru_cache(maxsize=32)
def _make_clamp_kernel(N, dtype, has_min, has_max, min_val, max_val,
output_dtype=None, is_fp8=False, threads=256, npt=8):
"""Build clamp kernel: y = clamp(x, min_val, max_val) with optional bounds.
For non-fp8 dtypes, uses register_copy strategy: fragment load -> compute
-> fragment store for coalesced memory access. Computes in fp32 then
casts back to preserve precision.
"""
out_dtype = output_dtype or dtype
block_size = threads * npt
if is_fp8:
@tilelang.jit(out_idx=[1])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(x: T.Tensor((N,), dtype), y: T.Tensor((N,), out_dtype)):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
for i, j in T.Parallel(threads_arg, npt_arg):
idx = (bx * threads_arg + i) * npt_arg + j
if idx < N:
v32 = T.cast(x[idx], "float32")
if has_min:
lo = T.cast(min_val, "float32")
v32 = T.max(v32, lo)
if has_max:
hi = T.cast(max_val, "float32")
v32 = T.min(v32, hi)
y[idx] = T.Cast(out_dtype, v32)
return main
else:
@tilelang.jit(out_idx=[1])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(x: T.Tensor((N,), dtype), y: T.Tensor((N,), dtype)):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
x_reg = T.alloc_fragment((block_size,), dtype)
y_reg = T.alloc_fragment((block_size,), dtype)
T.copy(x[bx * block_size : (bx + 1) * block_size], x_reg)
for i, j in T.Parallel(threads_arg, npt_arg):
val = x_reg[i * npt_arg + j]
v32 = T.cast(val, "float32")
if has_min:
lo = T.cast(min_val, "float32")
v32 = T.max(v32, lo)
if has_max:
hi = T.cast(max_val, "float32")
v32 = T.min(v32, hi)
y_reg[i * npt_arg + j] = T.Cast(val.dtype, v32)
T.copy(y_reg, y[bx * block_size : (bx + 1) * block_size])
return main
return kernel
class ClampFwdKernel(ParametricUnaryKernel):
"""Clamp: y = clamp(x, min, max) with optional bounds."""
def __init__(self, N_total, dtype, min_val=None, max_val=None, config=None, tune=False):
self.min_val = min_val
self.max_val = max_val
super().__init__(N_total, dtype, config=config, tune=tune)
@staticmethod
def _builder_fn():
return _make_clamp_kernel
def _builder_args(self):
return (
self.min_val is not None,
self.max_val is not None,
self.min_val if self.min_val is not None else 0.0,
self.max_val if self.max_val is not None else 0.0,
)
@functools.lru_cache(maxsize=32)
def _make_clamp_tensor_kernel(N, dtype, has_min, has_max,
output_dtype=None, is_fp8=False,
threads=256, npt=8):
"""Build Tensor-bound clamp kernel.
Inputs (all flat, length N, pre-broadcast/expanded by the Op layer):
x: data tensor.
lo: lower-bound tensor (only present when ``has_min``).
hi: upper-bound tensor (only present when ``has_max``).
Output:
y: clamp result, same dtype as ``output_dtype`` (or ``dtype``).
For fp8 the cast/compute uses fp32 to preserve precision; for non-fp8
the kernel uses register_copy with fp32 accumulation.
NaN semantics: matches ``torch.clamp`` / ``torch.clamp_min`` /
``torch.clamp_max``. If ``x``, ``lo``, or ``hi`` is NaN at a position,
the output at that position is NaN. ``T.max`` / ``T.min`` on CUDA do
not propagate NaN by themselves (they return the non-NaN operand), so
we add explicit ``isnan`` guards in fp32 -- mirroring the pattern used
by ``MaximumFwdKernel`` / ``MinimumFwdKernel``.
"""
if not (has_min or has_max):
raise ValueError(
"_make_clamp_tensor_kernel requires has_min or has_max to be True",
)
out_dtype = output_dtype or dtype
block_size = threads * npt
if is_fp8:
if has_min and has_max:
@tilelang.jit(out_idx=[3])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(
x: T.Tensor((N,), dtype),
lo: T.Tensor((N,), dtype),
hi: T.Tensor((N,), dtype),
y: T.Tensor((N,), out_dtype),
):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
for i, j in T.Parallel(threads_arg, npt_arg):
idx = (bx * threads_arg + i) * npt_arg + j
if idx < N:
x32 = T.cast(x[idx], "float32")
lo32 = T.cast(lo[idx], "float32")
hi32 = T.cast(hi[idx], "float32")
r = T.max(x32, lo32)
r = T.min(r, hi32)
# NaN propagation (PyTorch semantics):
# if any of x/lo/hi is NaN -> output NaN.
r = T.if_then_else(T.isnan(hi32), hi32, r)
r = T.if_then_else(T.isnan(lo32), lo32, r)
r = T.if_then_else(T.isnan(x32), x32, r)
y[idx] = T.Cast(out_dtype, r)
return main
return kernel
if has_min:
@tilelang.jit(out_idx=[2])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(
x: T.Tensor((N,), dtype),
lo: T.Tensor((N,), dtype),
y: T.Tensor((N,), out_dtype),
):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
for i, j in T.Parallel(threads_arg, npt_arg):
idx = (bx * threads_arg + i) * npt_arg + j
if idx < N:
x32 = T.cast(x[idx], "float32")
lo32 = T.cast(lo[idx], "float32")
r = T.max(x32, lo32)
# NaN propagation (PyTorch clamp_min):
# if x or lo is NaN -> output NaN.
r = T.if_then_else(T.isnan(lo32), lo32, r)
r = T.if_then_else(T.isnan(x32), x32, r)
y[idx] = T.Cast(out_dtype, r)
return main
return kernel
# has_max only
@tilelang.jit(out_idx=[2])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(
x: T.Tensor((N,), dtype),
hi: T.Tensor((N,), dtype),
y: T.Tensor((N,), out_dtype),
):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
for i, j in T.Parallel(threads_arg, npt_arg):
idx = (bx * threads_arg + i) * npt_arg + j
if idx < N:
x32 = T.cast(x[idx], "float32")
hi32 = T.cast(hi[idx], "float32")
r = T.min(x32, hi32)
# NaN propagation (PyTorch clamp_max):
# if x or hi is NaN -> output NaN.
r = T.if_then_else(T.isnan(hi32), hi32, r)
r = T.if_then_else(T.isnan(x32), x32, r)
y[idx] = T.Cast(out_dtype, r)
return main
return kernel
# non-fp8 path (register_copy)
if has_min and has_max:
@tilelang.jit(out_idx=[3])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(
x: T.Tensor((N,), dtype),
lo: T.Tensor((N,), dtype),
hi: T.Tensor((N,), dtype),
y: T.Tensor((N,), dtype),
):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
x_reg = T.alloc_fragment((block_size,), dtype)
lo_reg = T.alloc_fragment((block_size,), dtype)
hi_reg = T.alloc_fragment((block_size,), dtype)
T.copy(x[bx * block_size : (bx + 1) * block_size], x_reg)
T.copy(lo[bx * block_size : (bx + 1) * block_size], lo_reg)
T.copy(hi[bx * block_size : (bx + 1) * block_size], hi_reg)
for i, j in T.Parallel(threads_arg, npt_arg):
k = i * npt_arg + j
x32 = T.cast(x_reg[k], "float32")
lo32 = T.cast(lo_reg[k], "float32")
hi32 = T.cast(hi_reg[k], "float32")
r = T.max(x32, lo32)
r = T.min(r, hi32)
# NaN propagation (PyTorch clamp):
# if any of x/lo/hi is NaN -> output NaN.
r = T.if_then_else(T.isnan(hi32), hi32, r)
r = T.if_then_else(T.isnan(lo32), lo32, r)
r = T.if_then_else(T.isnan(x32), x32, r)
x_reg[k] = T.Cast(dtype, r)
T.copy(x_reg, y[bx * block_size : (bx + 1) * block_size])
return main
return kernel
if has_min:
@tilelang.jit(out_idx=[2])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(
x: T.Tensor((N,), dtype),
lo: T.Tensor((N,), dtype),
y: T.Tensor((N,), dtype),
):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
x_reg = T.alloc_fragment((block_size,), dtype)
lo_reg = T.alloc_fragment((block_size,), dtype)
T.copy(x[bx * block_size : (bx + 1) * block_size], x_reg)
T.copy(lo[bx * block_size : (bx + 1) * block_size], lo_reg)
for i, j in T.Parallel(threads_arg, npt_arg):
k = i * npt_arg + j
x32 = T.cast(x_reg[k], "float32")
lo32 = T.cast(lo_reg[k], "float32")
r = T.max(x32, lo32)
# NaN propagation (PyTorch clamp_min):
# if x or lo is NaN -> output NaN.
r = T.if_then_else(T.isnan(lo32), lo32, r)
r = T.if_then_else(T.isnan(x32), x32, r)
x_reg[k] = T.Cast(dtype, r)
T.copy(x_reg, y[bx * block_size : (bx + 1) * block_size])
return main
return kernel
# has_max only
@tilelang.jit(out_idx=[2])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(
x: T.Tensor((N,), dtype),
hi: T.Tensor((N,), dtype),
y: T.Tensor((N,), dtype),
):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
x_reg = T.alloc_fragment((block_size,), dtype)
hi_reg = T.alloc_fragment((block_size,), dtype)
T.copy(x[bx * block_size : (bx + 1) * block_size], x_reg)
T.copy(hi[bx * block_size : (bx + 1) * block_size], hi_reg)
for i, j in T.Parallel(threads_arg, npt_arg):
k = i * npt_arg + j
x32 = T.cast(x_reg[k], "float32")
hi32 = T.cast(hi_reg[k], "float32")
r = T.min(x32, hi32)
# NaN propagation (PyTorch clamp_max):
# if x or hi is NaN -> output NaN.
r = T.if_then_else(T.isnan(hi32), hi32, r)
r = T.if_then_else(T.isnan(x32), x32, r)
x_reg[k] = T.Cast(dtype, r)
T.copy(x_reg, y[bx * block_size : (bx + 1) * block_size])
return main
return kernel
class ClampTensorFwdKernel(ParametricUnaryKernel):
"""Tensor-bound clamp kernel.
Computes ``y = clamp(x, lo, hi)`` over flat tensors of length
``N_total``. The Op layer broadcasts ``input`` / ``min`` / ``max``
to the output shape and flattens them before dispatch. ``has_min``
/ ``has_max`` select between the three forms used by the Tensor
clamp, clamp_min, and clamp_max ops.
"""
_DEFAULT_THREADS = 512
def __init__(self, N_total, dtype, has_min, has_max,
config=None, tune=False):
if not (has_min or has_max):
raise ValueError(
"ClampTensorFwdKernel requires has_min or has_max to be True",
)
self.has_min = bool(has_min)
self.has_max = bool(has_max)
super().__init__(N_total, dtype, config=config, tune=tune)
@staticmethod
def _builder_fn():
return _make_clamp_tensor_kernel
def _builder_args(self):
return (self.has_min, self.has_max)
def forward(self, x, lo=None, hi=None):
if self.has_min and self.has_max:
result = self._compiled_fn(x, lo, hi)
elif self.has_min:
result = self._compiled_fn(x, lo)
else:
result = self._compiled_fn(x, hi)
if self._fp8_output_dtype is not None:
result = result.to(self._fp8_output_dtype)
return result
@functools.lru_cache(maxsize=32)
def _make_masked_fill_kernel(N, dtype, fill_value, output_dtype=None,
is_fp8=False, threads=256, npt=8):
"""Build masked_fill kernel: out = mask ? fill_value : x.
The Op layer packs the bool mask as uint8 so that T.copy can
perform vectorized loads (TileLang does not vectorize bool tensors).
Each uint8 element is 0 or 1; the kernel loads it into a register
fragment and unpacks per-element with a != 0 comparison.
For non-fp8 dtypes, writes the result back into the x register fragment
(in-place) to reduce register pressure and avoid a third data-typed
fragment allocation.
For e5m2, the kernel outputs fp16 so the Op layer can do a
non-saturating cast to e5m2.
"""
out_dtype = output_dtype or dtype
block_size = threads * npt
if is_fp8:
@tilelang.jit(out_idx=[2])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(
x: T.Tensor((N,), dtype),
mask: T.Tensor((N,), "uint8"),
out: T.Tensor((N,), out_dtype),
):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
for i, j in T.Parallel(threads_arg, npt_arg):
idx = (bx * threads_arg + i) * npt_arg + j
if idx < N:
fv = T.cast(fill_value, out_dtype)
x_val = T.Cast(out_dtype, x[idx])
out[idx] = T.if_then_else(
mask[idx] != T.cast(0, "uint8"), fv, x_val,
)
return main
else:
@tilelang.jit(out_idx=[2])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(
x: T.Tensor((N,), dtype),
mask: T.Tensor((N,), "uint8"),
out: T.Tensor((N,), dtype),
):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
m_reg = T.alloc_fragment((block_size,), "uint8")
x_reg = T.alloc_fragment((block_size,), dtype)
T.copy(mask[bx * block_size : (bx + 1) * block_size], m_reg)
T.copy(x[bx * block_size : (bx + 1) * block_size], x_reg)
for i, j in T.Parallel(threads_arg, npt_arg):
k = i * npt_arg + j
fv = T.cast(fill_value, dtype)
x_reg[k] = T.if_then_else(
m_reg[k] != T.cast(0, "uint8"), fv, x_reg[k],
)
T.copy(x_reg, out[bx * block_size : (bx + 1) * block_size])
return main
return kernel
class MaskedFillFwdKernel(ParametricUnaryKernel):
"""MaskedFill: out = mask ? fill_value : x.
Supports the PyTorch ``Tensor.masked_fill(mask, value: Number)`` dtype
union of integer and floating-point input dtypes. The bool dtype path
is handled at the Op layer by viewing the input as uint8 and casting
the result back to bool, so the kernel itself only sees integer and
floating-point storage dtypes.
"""
_DEFAULT_THREADS = 512
SUPPORTED_DTYPES = _BITWISE_DTYPES[1:] + _FLOAT_DTYPES # uint8/intN + fp16/bf16/fp32
def __init__(self, N_total, dtype, fill_value, config=None, tune=False):
self._raw_fill_value = fill_value
super().__init__(N_total, dtype, config=config, tune=tune)
def _post_init_params(self):
self.fill_value = _clamp_to_dtype_range(self._raw_fill_value, self.output_dtype)
@staticmethod
def _builder_fn():
return _make_masked_fill_kernel
def _builder_args(self):
return (self.fill_value,)
def forward(self, x, mask):
return self._compiled_fn(x, mask)
@functools.lru_cache(maxsize=32)
def _make_masked_fill_tensor_value_kernel(N, dtype, output_dtype=None,
is_fp8=False, threads=256, npt=8):
"""Build masked_fill kernel with a 0-dim Tensor fill value.
Inputs (all flat, length N, pre-broadcast/expanded by the Op layer):
x: data tensor (length N).
mask: bool mask packed as uint8 (length N).
value: scalar fill value carried as a length-1 tensor (the Op
layer reshapes the 0-dim Tensor to ``(1,)``).
Output:
out: ``out[i] = value[0] if mask[i] else x[i]``.
"""
out_dtype = output_dtype or dtype
block_size = threads * npt
if is_fp8:
@tilelang.jit(out_idx=[3])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(
x: T.Tensor((N,), dtype),
mask: T.Tensor((N,), "uint8"),
value: T.Tensor((1,), dtype),
out: T.Tensor((N,), out_dtype),
):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
fv = T.Cast(out_dtype, value[0])
for i, j in T.Parallel(threads_arg, npt_arg):
idx = (bx * threads_arg + i) * npt_arg + j
if idx < N:
x_val = T.Cast(out_dtype, x[idx])
out[idx] = T.if_then_else(
mask[idx] != T.cast(0, "uint8"), fv, x_val,
)
return main
return kernel
@tilelang.jit(out_idx=[3])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(
x: T.Tensor((N,), dtype),
mask: T.Tensor((N,), "uint8"),
value: T.Tensor((1,), dtype),
out: T.Tensor((N,), dtype),
):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
m_reg = T.alloc_fragment((block_size,), "uint8")
x_reg = T.alloc_fragment((block_size,), dtype)
T.copy(mask[bx * block_size : (bx + 1) * block_size], m_reg)
T.copy(x[bx * block_size : (bx + 1) * block_size], x_reg)
fv = value[0]
for i, j in T.Parallel(threads_arg, npt_arg):
k = i * npt_arg + j
x_reg[k] = T.if_then_else(
m_reg[k] != T.cast(0, "uint8"), fv, x_reg[k],
)
T.copy(x_reg, out[bx * block_size : (bx + 1) * block_size])
return main
return kernel
class MaskedFillTensorValueFwdKernel(ParametricUnaryKernel):
"""MaskedFill kernel with 0-dim Tensor fill value.
Computes ``out = mask ? value : x`` over flat tensors of length
``N_total``. The Op layer broadcasts ``input`` and ``mask`` to the
output shape, flattens them, packs the mask as uint8, and reshapes
the 0-dim ``value`` to a length-1 tensor before dispatch. The bool
input dtype is routed through uint8 at the Op layer, so this kernel
only sees integer and floating-point storage dtypes.
"""
_DEFAULT_THREADS = 512
SUPPORTED_DTYPES = _BITWISE_DTYPES[1:] + _FLOAT_DTYPES # uint8/intN + fp16/bf16/fp32
@staticmethod
def _builder_fn():
return _make_masked_fill_tensor_value_kernel
def forward(self, x, mask, value):
result = self._compiled_fn(x, mask, value)
if self._fp8_output_dtype is not None:
result = result.to(self._fp8_output_dtype)
return result
@functools.lru_cache(maxsize=32)
def _make_nan_to_num_kernel(N, dtype, nan_val, posinf_val, neginf_val,
output_dtype=None, is_fp8=False, threads=256, npt=8):
"""Build nan_to_num kernel: replace NaN, +Inf, -Inf with given values.
For non-fp8 dtypes, uses register_copy strategy: fragment load -> compute
-> fragment store for coalesced memory access.
"""
out_dtype = output_dtype or dtype
block_size = threads * npt
if is_fp8:
@tilelang.jit(out_idx=[1])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(x: T.Tensor((N,), dtype), y: T.Tensor((N,), out_dtype)):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
for i, j in T.Parallel(threads_arg, npt_arg):
idx = (bx * threads_arg + i) * npt_arg + j
if idx < N:
val = x[idx]
v32 = T.cast(val, "float32")
nan_r = T.cast(nan_val, out_dtype)
pos_r = T.cast(posinf_val, out_dtype)
neg_r = T.cast(neginf_val, out_dtype)
pass_through = T.Cast(out_dtype, v32)
result = T.if_then_else(
T.isnan(v32),
nan_r,
T.if_then_else(
T.isinf(v32),
T.if_then_else(v32 > T.cast(0, "float32"), pos_r, neg_r),
pass_through,
),
)
y[idx] = result
return main
else:
@tilelang.jit(out_idx=[1])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(x: T.Tensor((N,), dtype), y: T.Tensor((N,), dtype)):
with T.Kernel(T.ceildiv(N, block_size), threads=threads_arg) as bx:
x_reg = T.alloc_fragment((block_size,), dtype)
y_reg = T.alloc_fragment((block_size,), dtype)
T.copy(x[bx * block_size : (bx + 1) * block_size], x_reg)
for i, j in T.Parallel(threads_arg, npt_arg):
k = i * npt_arg + j
val = x_reg[k]
v32 = T.cast(val, "float32")
nan_r = T.cast(nan_val, val.dtype)
pos_r = T.cast(posinf_val, val.dtype)
neg_r = T.cast(neginf_val, val.dtype)
result = T.if_then_else(
T.isnan(v32),
nan_r,
T.if_then_else(
T.isinf(v32),
T.if_then_else(v32 > T.cast(0, "float32"), pos_r, neg_r),
val,
),
)
y_reg[k] = result
T.copy(y_reg, y[bx * block_size : (bx + 1) * block_size])
return main
return kernel
class NanToNumFwdKernel(ParametricUnaryKernel):
"""NanToNum: replace NaN, +Inf, -Inf with specified values."""
def __init__(self, N_total, dtype, nan_val=0.0, posinf_val=1e4, neginf_val=-1e4,
config=None, tune=False):
self._raw_nan_val = nan_val
self._raw_posinf_val = posinf_val
self._raw_neginf_val = neginf_val
super().__init__(N_total, dtype, config=config, tune=tune)
def _post_init_params(self):
self.nan_val = _clamp_to_dtype_range(self._raw_nan_val, self.output_dtype)
self.posinf_val = _clamp_to_dtype_range(self._raw_posinf_val, self.output_dtype)
self.neginf_val = _clamp_to_dtype_range(self._raw_neginf_val, self.output_dtype)
@staticmethod
def _builder_fn():
return _make_nan_to_num_kernel
def _builder_args(self):
return (self.nan_val, self.posinf_val, self.neginf_val)
@functools.lru_cache(maxsize=32)
def _make_alibi_kernel(seq_len, num_heads, dtype, threads=256, npt=8):
"""Build ALiBi kernel: bias[h, i, j] = -slope_h * |i - j|.
Slopes: slope_h = 2^(-8*h/H) for head h in [0, H).
Output shape: (num_heads, seq_len, seq_len).
Total elements: num_heads * seq_len * seq_len.
"""
N_total = num_heads * seq_len * seq_len
block_size = threads * npt
S2 = seq_len * seq_len
@tilelang.jit(out_idx=[0])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(out: T.Tensor((N_total,), dtype)):
with T.Kernel(T.ceildiv(N_total, block_size), threads=threads_arg) as bx:
for i, j in T.Parallel(threads_arg, npt_arg):
flat = (bx * threads_arg + i) * npt_arg + j
if flat < N_total:
h = flat // S2
rem = flat % S2
row = rem // seq_len
col = rem % seq_len
# slope = 2^(-8 * (h+1) / num_heads)
exp_val = T.cast(-8.0, "float32") * T.cast(h + 1, "float32") / T.cast(num_heads, "float32")
slope = T.exp2(exp_val)
dist = T.cast(row - col, "float32")
# Use abs via if_then_else since T.abs may not handle int
abs_dist = T.if_then_else(dist > T.cast(0, "float32"), dist, -dist)
out[flat] = T.Cast(dtype, -slope * abs_dist)
return main
return kernel
class AlibiFwdKernel(Kernel):
"""ALiBi position encoding: bias[h, i, j] = -slope_h * |i - j|.
Generates the full (num_heads, seq_len, seq_len) bias tensor.
Slopes follow the ALiBi paper: slope_h = 2^(-8*(h+1)/H).
Args:
seq_len: Sequence length.
num_heads: Number of attention heads.
dtype: Torch dtype.
config: Optional config dict.
tune: Whether to autotune.
"""
supported_archs: list[int] = [80, 86, 89, 90]
SUPPORTED_DTYPES = _FLOAT_DTYPES
def __init__(self, seq_len, num_heads, dtype, config=None, tune=False):
super().__init__()
if dtype not in self.SUPPORTED_DTYPES:
supported = ", ".join(str(dt) for dt in self.SUPPORTED_DTYPES)
raise ValueError(
f"{self.__class__.__name__} only supports dtypes [{supported}], got {dtype}"
)
self.seq_len = seq_len
self.num_heads = num_heads
self.dtype = dtype
self._fp8_output_dtype, self.output_dtype = _get_fp8_output_dtypes(dtype)
cfg = self.default_config
self.kernel = _make_alibi_kernel(
seq_len, num_heads, self.dtype_to_str(self.output_dtype),
cfg["threads"], cfg["num_per_thread"],
)
self.init_config(config, tune)
@property
def default_config(self):
npt = 4 if self.dtype == torch.float32 else (16 if _is_fp8(self.dtype) else 8)
return {"threads": 256, "num_per_thread": npt}
def init_config(self, config=None, tune=False):
"""Override to cache the compiled kernel function after config is set."""
super().init_config(config, tune)
cfg = self.config
self._compiled_fn = self.kernel(cfg["threads"], cfg["num_per_thread"])
def forward(self):
return self._compiled_fn()
@functools.lru_cache(maxsize=32)
def _make_sinusoidal_kernel(seq_len, d_model, dtype, threads=256, npt=8):
"""Build sinusoidal positional encoding kernel.
PE(pos, 2i) = sin(pos / 10000^(2i/d_model))
PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))
Output shape: (seq_len, d_model).
"""
N_total = seq_len * d_model
block_size = threads * npt
@tilelang.jit(out_idx=[0])
def kernel(threads_arg, npt_arg):
@T.prim_func
def main(out: T.Tensor((N_total,), dtype)):
with T.Kernel(T.ceildiv(N_total, block_size), threads=threads_arg) as bx:
for i, j in T.Parallel(threads_arg, npt_arg):
flat = (bx * threads_arg + i) * npt_arg + j
if flat < N_total:
pos = flat // d_model
dim = flat % d_model
# dim_pair = dim // 2 (the "i" in the formula)
dim_pair = dim // 2
# angle = pos / 10000^(2*dim_pair / d_model)
base = T.cast(10000.0, "float32")
exp_frac = T.cast(dim_pair, "float32") * T.cast(2.0, "float32") / T.cast(d_model, "float32")
divisor = T.pow(base, exp_frac)
angle = T.cast(pos, "float32") / divisor
# Even dim -> sin, odd dim -> cos
is_even = dim % 2 == 0
result = T.if_then_else(is_even, T.sin(angle), T.cos(angle))
out[flat] = T.Cast(dtype, result)
return main
return kernel
class SinusoidalFwdKernel(Kernel):
"""Sinusoidal positional encoding from "Attention Is All You Need".
PE(pos, 2i) = sin(pos / 10000^(2i/d_model))
PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))
Args:
seq_len: Sequence length.
d_model: Model dimension (must be even).
dtype: Torch dtype.
config: Optional config dict.
tune: Whether to autotune.
"""
supported_archs: list[int] = [80, 86, 89, 90]
SUPPORTED_DTYPES = _FLOAT_DTYPES
def __init__(self, seq_len, d_model, dtype, config=None, tune=False):
super().__init__()
if dtype not in self.SUPPORTED_DTYPES:
supported = ", ".join(str(dt) for dt in self.SUPPORTED_DTYPES)
raise ValueError(
f"{self.__class__.__name__} only supports dtypes [{supported}], got {dtype}"
)
self.seq_len = seq_len
self.d_model = d_model
self.dtype = dtype
self._fp8_output_dtype, self.output_dtype = _get_fp8_output_dtypes(dtype)
cfg = self.default_config
self.kernel = _make_sinusoidal_kernel(
seq_len, d_model, self.dtype_to_str(self.output_dtype),
cfg["threads"], cfg["num_per_thread"],
)
self.init_config(config, tune)
@property
def default_config(self):
npt = 4 if self.dtype == torch.float32 else (16 if _is_fp8(self.dtype) else 8)
return {"threads": 256, "num_per_thread": npt}
def init_config(self, config=None, tune=False):
"""Override to cache the compiled kernel function after config is set."""
super().init_config(config, tune)
cfg = self.config
self._compiled_fn = self.kernel(cfg["threads"], cfg["num_per_thread"])
def forward(self):
return self._compiled_fn()