Sources cloned and tree'd (no --depth):
- jd-opensource/xllm: ILU kernels, CUDA kernels, MoE kernels
- NVIDIA/cccl: CUB tuning/dispatch headers (block-level primitives)
- fla-org/flash-linear-attention: Triton GDN kernels
- NVIDIA/cutlass: grouped GEMM reference (read, not copied)
- Dao-AILab/flash-attention: attention kernel reference (SM80+, read only)
New CUDA kernels (from xllm, SM-agnostic, portable to BI-V100):
ex_engine/xllm_kernels/cuda/activation.cu (188 lines) — silu_and_mul, gelu
ex_engine/xllm_kernels/cuda/norm.cu (600 lines) — rms_norm, fused_add_rms_norm
ex_engine/xllm_kernels/cuda/rope.cu (258 lines) — rotary_embedding
ex_engine/xllm_kernels/cuda/block_copy.cu (209 lines) — copy_blocks, swap_blocks
ex_engine/xllm_kernels/cuda/reshape_paged_cache.cu (101 lines) — KV cache ops
ex_engine/xllm_kernels/cuda/headers/ (5 headers for compilation)
ILU bridge kernel sources (from xllm, verified SAME as upstream):
ex_engine/xllm_kernels/ilu/ (10 files, 925 lines total)
— activation.cpp, attention.cpp, fused_moe.cpp, group_gemm.cpp,
matmul.cpp, norm.cpp, rope.cpp, ilu_ops_api.h, ixformer.h, utils.h
FLA Triton GDN kernels (for GatedDeltaNet without SM90+ FlashQLA):
ex_engine/fla_kernels/gated_delta_rule/ (7 files, 2370 lines)
— chunk_fwd.py (428), chunk.py (487), wy_fast.py (409),
fused_recurrent.py (392), naive.py (161), gate.py (380)
CCCL sync (12 tuning + 14 dispatch headers updated from NVIDIA/cccl):
cccl_upstream/cub/cub/device/dispatch/tuning/ — 12 changed files synced
cccl_upstream/cub/cub/device/dispatch/ — 14 changed dispatch files synced
Compilation targets for real machine (ivcore10):
1. CUDA kernels: --cuda-gpu-arch=ivcore10 via corex clang/16
2. ILU bridges: torch.utils.cpp_extension linking ixformer .so
3. FLA kernels: Triton JIT (if Triton works on BI-V100)
102 lines
3.1 KiB
Python
102 lines
3.1 KiB
Python
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
|
|
#
|
|
# This source code is licensed under the MIT license found in the
|
|
# LICENSE file in the root directory of this source tree.
|
|
# For a list of all contributors, visit:
|
|
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
|
|
|
|
import os
|
|
|
|
import triton
|
|
import triton.language as tl
|
|
import triton.language.extra.libdevice as tldevice
|
|
|
|
from fla.utils import IS_GATHER_SUPPORTED, IS_NVIDIA_BLACKWELL
|
|
|
|
if os.environ.get('FLA_USE_FAST_OPS', '0') == '1':
|
|
@triton.jit
|
|
def exp(x): return tldevice.fast_expf(x.to(tl.float32))
|
|
@triton.jit
|
|
def exp2(x): return tldevice.exp2(x.to(tl.float32))
|
|
@triton.jit
|
|
def log(x): return tldevice.fast_logf(x.to(tl.float32))
|
|
@triton.jit
|
|
def log2(x): return tldevice.fast_log2f(x.to(tl.float32))
|
|
@triton.jit
|
|
def tanh(x): return tldevice.fast_tanhf(x.to(tl.float32))
|
|
else:
|
|
@triton.jit
|
|
def exp(x): return tl.exp(x.to(tl.float32))
|
|
@triton.jit
|
|
def exp2(x): return tl.math.exp2(x.to(tl.float32))
|
|
@triton.jit
|
|
def log(x): return tl.log(x.to(tl.float32))
|
|
@triton.jit
|
|
def log2(x): return tl.log2(x.to(tl.float32))
|
|
@triton.jit
|
|
def tanh(x): return tldevice.tanh(x.to(tl.float32))
|
|
|
|
|
|
if IS_NVIDIA_BLACKWELL:
|
|
"""
|
|
Compute tl.dot with Blackwell workaround.
|
|
|
|
On SM100 datacenter and SM120 consumer Blackwell GPUs, wraps the result in
|
|
inline assembly to prevent the TritonGPUHoistTMEMAlloc pass from incorrectly
|
|
fusing add and dot operations.
|
|
See: https://github.com/fla-org/flash-linear-attention/issues/638
|
|
|
|
TODO: Remove this workaround once the Triton compiler bug is fixed.
|
|
Track upstream issue at: https://github.com/triton-lang/triton/issues/8695
|
|
"""
|
|
@triton.jit
|
|
def safe_dot(a, b, allow_tf32: tl.constexpr = None):
|
|
return tl.inline_asm_elementwise(
|
|
asm="mov.f32 $0, $1;",
|
|
constraints="=r,r",
|
|
args=[tl.dot(a, b, allow_tf32=allow_tf32)],
|
|
dtype=tl.float32,
|
|
is_pure=True,
|
|
pack=1,
|
|
)
|
|
else:
|
|
@triton.jit
|
|
def safe_dot(a, b, allow_tf32: tl.constexpr = None):
|
|
return tl.dot(a, b, allow_tf32=allow_tf32)
|
|
|
|
|
|
if not IS_GATHER_SUPPORTED:
|
|
@triton.jit
|
|
def gather(src, index, axis, _builder=None):
|
|
"""
|
|
Gather operation that works when tl.gather is not supported.
|
|
This is a fallback implementation that returns None.
|
|
Just to make triton compiler happy.
|
|
"""
|
|
return None
|
|
else:
|
|
gather = tl.gather
|
|
|
|
|
|
if hasattr(triton.language, '_experimental_make_tensor_descriptor'):
|
|
# For Triton 3.3.x
|
|
make_tensor_descriptor = triton.language._experimental_make_tensor_descriptor
|
|
elif hasattr(triton.language, 'make_tensor_descriptor'):
|
|
# For Triton 3.4.x and later
|
|
make_tensor_descriptor = triton.language.make_tensor_descriptor
|
|
else:
|
|
"""
|
|
Fallback implementation when TMA is not supported.
|
|
Returns None to indicate TMA descriptors are unavailable.
|
|
Just make triton compiler happy.
|
|
"""
|
|
@triton.jit
|
|
def make_tensor_descriptor(
|
|
base,
|
|
shape,
|
|
strides,
|
|
block_shape,
|
|
_builder=None,
|
|
):
|
|
return None
|