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)
66 lines
1.8 KiB
Python
66 lines
1.8 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
|
|
|
|
from .csr import prepare_block_csr
|
|
from .cumsum import (
|
|
chunk_global_cumsum,
|
|
chunk_global_cumsum_scalar,
|
|
chunk_global_cumsum_vector,
|
|
chunk_local_cumsum,
|
|
chunk_local_cumsum_scalar,
|
|
chunk_local_cumsum_vector,
|
|
)
|
|
from .index import (
|
|
get_max_num_splits,
|
|
prepare_chunk_indices,
|
|
prepare_chunk_offsets,
|
|
prepare_cu_seqlens_from_lens,
|
|
prepare_cu_seqlens_from_mask,
|
|
prepare_lens,
|
|
prepare_lens_from_mask,
|
|
prepare_position_ids,
|
|
prepare_sequence_ids,
|
|
prepare_token_indices,
|
|
)
|
|
from .logsumexp import logsumexp_fwd
|
|
from .matmul import addmm, matmul
|
|
from .pack import pack_sequence, unpack_sequence
|
|
from .pooling import mean_pooling
|
|
from .softmax import softmax_bwd, softmax_fwd
|
|
from .softplus import softplus
|
|
from .solve_tril import solve_tril
|
|
|
|
__all__ = [
|
|
"addmm",
|
|
"chunk_global_cumsum",
|
|
"chunk_global_cumsum_scalar",
|
|
"chunk_global_cumsum_vector",
|
|
"chunk_local_cumsum",
|
|
"chunk_local_cumsum_scalar",
|
|
"chunk_local_cumsum_vector",
|
|
"get_max_num_splits",
|
|
"logsumexp_fwd",
|
|
"matmul",
|
|
"mean_pooling",
|
|
"pack_sequence",
|
|
"prepare_block_csr",
|
|
"prepare_chunk_indices",
|
|
"prepare_chunk_offsets",
|
|
"prepare_cu_seqlens_from_lens",
|
|
"prepare_cu_seqlens_from_mask",
|
|
"prepare_lens",
|
|
"prepare_lens_from_mask",
|
|
"prepare_position_ids",
|
|
"prepare_sequence_ids",
|
|
"prepare_token_indices",
|
|
"softmax_bwd",
|
|
"softmax_fwd",
|
|
"softplus",
|
|
"solve_tril",
|
|
"unpack_sequence",
|
|
]
|