forked from ccf-ai-infra/TileOPs-Metax
3668 lines
131 KiB
Python
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()
|