Files
project_6/ex_engine/fla_kernels/utils/__init__.py
claude 8d75652949 feat: import CUDA kernels from xllm/CCCL/FLA upstream repos
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)
2026-08-14 07:48:52 +00:00

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",
]