Files
project_6/ex_engine/xllm_kernels/ilu/ixformer.h

148 lines
5.6 KiB
C
Raw Normal View History

/* Copyright 2025 The xLLM Authors. All Rights Reserved.
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
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
https://github.com/jd-opensource/xllm/blob/main/LICENSE
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#include <torch/all.h>
#include "ATen/Tensor.h"
#include "utils.h"
namespace ixformer::infer {
torch::Tensor ixinfer_flash_attn_unpad_with_block_tables(
torch::Tensor& query,
torch::Tensor& key_cache,
torch::Tensor& value_cache,
torch::Tensor& out,
torch::Tensor& block_tables,
torch::Tensor& cu_seq_q,
torch::Tensor& cu_seq_k,
int64_t max_seq_q,
int64_t max_seq_k,
bool is_causal,
int64_t window_left,
int64_t window_right,
double scale,
double softcap,
bool sqrt_alibi,
const std::optional<torch::Tensor>& alibi_slopes,
const std::optional<torch::Tensor>& sinks,
std::optional<torch::Tensor>& lse);
void silu_and_mul(torch::Tensor& input, torch::Tensor& output);
torch::Tensor xllm_paged_attention(
torch::Tensor& out,
torch::Tensor& query,
torch::Tensor& key_cache,
torch::Tensor& value_cache,
int64_t num_kv_heads,
double scale,
torch::Tensor& block_tables,
torch::Tensor& context_lens,
int64_t block_size,
int64_t max_context_len,
const std::optional<torch::Tensor>& alibi_slopes,
bool causal,
int32_t window_left,
int32_t window_right,
double softcap,
bool enable_cuda_graph,
bool use_sqrt_alibi,
const std::optional<torch::Tensor>& sinks);
torch::Tensor ixformer_linear(torch::Tensor& input,
torch::Tensor& weight,
int64_t act_type,
const std::optional<torch::Tensor>& bias,
const std::optional<torch::Tensor>& out,
const std::optional<bool> persistent);
torch::Tensor ixformer_linear_ex(torch::Tensor& input,
torch::Tensor& weight,
const c10::optional<torch::Tensor>& bias,
const c10::optional<torch::Tensor>& out);
void xllm_reshape_and_cache(torch::Tensor& key,
torch::Tensor& value,
torch::Tensor& key_cache,
torch::Tensor& value_cache,
torch::Tensor& slot_mapping,
int64_t key_token_stride,
int64_t value_token_stride);
void xllm_rotary_embedding(torch::Tensor& positions,
torch::Tensor& query,
torch::Tensor& key,
int64_t head_size,
torch::Tensor& cos_sin_cache,
bool is_neox);
void residual_rms_norm(torch::Tensor& input,
torch::Tensor& residual,
torch::Tensor& weight,
torch::Tensor& output,
torch::Tensor& residual_output,
const std::optional<torch::Tensor>& fused_bias,
double alpha,
double eps,
bool is_post);
void rms_norm(torch::Tensor& input,
torch::Tensor& weight,
torch::Tensor& output,
const std::optional<torch::Tensor>& fused_bias,
double eps);
void topk_softmax(torch::Tensor& topk_weights,
torch::Tensor& topk_indices,
torch::Tensor& token_expert_indices,
torch::Tensor& gating_output,
bool renormalize);
void moe_compute_token_index_api(
torch::Tensor& topk_ids,
torch::Tensor& src_dst,
torch::Tensor& dst_src,
torch::Tensor& expert_sizes_gpu,
const c10::optional<torch::Tensor>& expert_mask,
const c10::optional<torch::Tensor>& expert_sizes_cpu,
const c10::optional<torch::Tensor>& expand_tokens_gpu,
int64_t start_expert_id,
int64_t end_expert_id,
int64_t num_experts);
void moe_expand_input(torch::Tensor outputs,
torch::Tensor inputs,
torch::Tensor dst_to_src,
const c10::optional<torch::Tensor>& src_to_dst,
int64_t dst_tokens,
int64_t expand_factor);
void moe_w16a16_group_gemm(torch::Tensor output,
torch::Tensor inputs,
torch::Tensor weights,
torch::Tensor tokens_per_experts,
const c10::optional<torch::Tensor>& dst_to_src,
const c10::optional<torch::Tensor>& bias,
std::string format,
int64_t persistent,
int64_t output_n);
void moe_output_reduce_sum(torch::Tensor outputs,
torch::Tensor inputs,
const c10::optional<torch::Tensor>& mul_weight,
const c10::optional<torch::Tensor>& mask,
const c10::optional<torch::Tensor>& extra_residual,
double scaling_factor);
} // namespace ixformer::infer