Copied from upstream_ref (NOT rewritten — exact upstream code):
ixformer C++ API (the authoritative header):
include/ixformer.h — ixformer::infer namespace: topk_softmax,
moe_compute_token_index_api, moe_w16a16_group_gemm, moe_expand_input,
moe_output_reduce_sum, silu_and_mul, rms_norm, xllm_paged_attention, etc.
include/ilu_ops_api.h — xllm::kernel::ilu namespace: moe_active_topk,
moe_gen_idx, moe_expand_input, group_gemm, moe_combine_result,
batch_prefill, batch_decode, rms_norm, matmul, act_and_mul, etc.
ILU kernel wrappers (call ixformer::infer directly):
csrc/ilu_kernel_fused_moe.cpp — topk routing + gen_idx + expand + combine
csrc/ilu_kernel_group_gemm.cpp — batched expert GEMM
csrc/ilu_kernel_{activation,norm,rope,matmul,attention}.cpp
ILU layer implementations (full pipeline):
csrc/ilu_layer_fused_moe.{cpp,h} — 797 lines, the complete MoE pipeline
that competitor 168 ran as corex_moe.py
csrc/ilu_layer_attention.{cpp,h} — prefill/decode attention dispatch
CUDA MoE kernels (from xllm + ds_vllm):
csrc/moe/moe_topk_softmax_kernels.cuh — CUB BlockReduce + warp topk
csrc/moe/moe_topk_sigmoid_kernels.cuh — sigmoid scoring variant
csrc/moe/moe_topk.cuh + moe_fused_topk.cu — entry points
csrc/moe/moeTopKFuncs.cuh — TRT-LLM derived vllm-compatible topk
csrc/moe/moe_ops.h + moe_align_sum_kernels.cu — alignment kernels
Common layer headers:
csrc/common_fused_moe{,_base}.h + common_moe_fused_topk.{cpp,h}
190 lines
7.9 KiB
C++
190 lines
7.9 KiB
C++
/* Copyright 2025 The xLLM Authors. All Rights Reserved.
|
|
|
|
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 "attention.h"
|
|
|
|
#include "kernels/ilu/ilu_ops_api.h"
|
|
#include "kernels/ops_api.h"
|
|
|
|
namespace xllm {
|
|
namespace layer {
|
|
AttentionImpl::AttentionImpl(int64_t num_heads,
|
|
int64_t head_size,
|
|
float scale,
|
|
int64_t num_kv_heads,
|
|
int64_t sliding_window)
|
|
: num_heads_(num_heads),
|
|
head_size_(head_size),
|
|
scale_(scale),
|
|
num_kv_heads_(num_kv_heads),
|
|
v_head_dim_(head_size),
|
|
use_fused_mla_qkv_(false),
|
|
enable_lighting_indexer_(false),
|
|
enable_mla_(false),
|
|
sliding_window_(sliding_window) {
|
|
if (sliding_window_ > -1) {
|
|
sliding_window_ = sliding_window_ - 1;
|
|
}
|
|
}
|
|
|
|
AttentionImpl::AttentionImpl(int64_t num_heads,
|
|
int64_t head_size,
|
|
int64_t num_kv_heads,
|
|
int64_t v_head_dim,
|
|
int64_t sliding_window,
|
|
float scale,
|
|
bool use_fused_mla_qkv,
|
|
bool enable_lighting_indexer,
|
|
bool enable_mla)
|
|
: num_heads_(num_heads),
|
|
head_size_(head_size),
|
|
scale_(scale),
|
|
num_kv_heads_(num_kv_heads),
|
|
v_head_dim_(v_head_dim),
|
|
use_fused_mla_qkv_(use_fused_mla_qkv),
|
|
enable_lighting_indexer_(enable_lighting_indexer),
|
|
enable_mla_(enable_mla),
|
|
sliding_window_(sliding_window) {
|
|
if (sliding_window_ > -1) {
|
|
sliding_window_ = sliding_window_ - 1;
|
|
}
|
|
}
|
|
|
|
std::tuple<torch::Tensor, std::optional<torch::Tensor>> AttentionImpl::forward(
|
|
const AttentionMetadata& attn_metadata,
|
|
torch::Tensor& query,
|
|
torch::Tensor& key,
|
|
torch::Tensor& value,
|
|
KVCache& kv_cache) {
|
|
std::optional<torch::Tensor> output_lse = std::nullopt;
|
|
torch::Tensor output;
|
|
if (enable_mla_) {
|
|
output = torch::empty({query.size(0), num_heads_ * v_head_dim_},
|
|
query.options());
|
|
} else {
|
|
output = torch::empty_like(query);
|
|
}
|
|
if (attn_metadata.is_dummy) {
|
|
return std::make_tuple(output, output_lse);
|
|
}
|
|
|
|
bool only_prefill =
|
|
attn_metadata.is_prefill || attn_metadata.is_chunked_prefill;
|
|
int64_t num_kv_heads = (enable_mla_ && !only_prefill) ? 1 : num_kv_heads_;
|
|
torch::Tensor k_cache = kv_cache.get_k_cache();
|
|
std::optional<torch::Tensor> v_cache;
|
|
std::optional<torch::Tensor> v;
|
|
if (!enable_mla_) {
|
|
v = value.view({-1, num_kv_heads, head_size_});
|
|
v_cache = kv_cache.get_v_cache();
|
|
}
|
|
|
|
bool skip_process_cache = enable_mla_ && (only_prefill || use_fused_mla_qkv_);
|
|
if (!skip_process_cache) {
|
|
xllm::kernel::ReshapePagedCacheParams reshape_paged_cache_params;
|
|
reshape_paged_cache_params.key = key.view({-1, num_kv_heads, head_size_});
|
|
reshape_paged_cache_params.value = v;
|
|
reshape_paged_cache_params.k_cache = k_cache;
|
|
reshape_paged_cache_params.v_cache = v_cache;
|
|
reshape_paged_cache_params.slot_mapping = attn_metadata.slot_mapping;
|
|
xllm::kernel::reshape_paged_cache(reshape_paged_cache_params);
|
|
}
|
|
|
|
if (enable_lighting_indexer_ || !only_prefill) {
|
|
decoder_forward(query, output, k_cache, v_cache, attn_metadata);
|
|
} else {
|
|
prefill_forward(query, key, value, output, k_cache, v_cache, attn_metadata);
|
|
}
|
|
|
|
int64_t head_size = enable_mla_ ? v_head_dim_ : head_size_;
|
|
output = output.view({-1, num_heads_ * head_size});
|
|
return {output, output_lse};
|
|
}
|
|
|
|
void AttentionImpl::prefill_forward(torch::Tensor& query,
|
|
torch::Tensor& key,
|
|
torch::Tensor& value,
|
|
torch::Tensor& output,
|
|
const torch::Tensor& k_cache,
|
|
const std::optional<torch::Tensor>& v_cache,
|
|
const AttentionMetadata& attn_metadata) {
|
|
int64_t head_size_v = enable_mla_ ? v_head_dim_ : head_size_;
|
|
std::optional<torch::Tensor> output_lse = std::nullopt;
|
|
query = query.view({-1, num_heads_, head_size_});
|
|
output = output.view({-1, num_heads_, head_size_v});
|
|
// torch::Tensor k_cache_ = k_cache;
|
|
// torch::Tensor v_cache_ = v_cache.value();
|
|
xllm::kernel::ilu::batch_prefill(query,
|
|
k_cache,
|
|
v_cache,
|
|
output,
|
|
output_lse,
|
|
attn_metadata.q_cu_seq_lens,
|
|
attn_metadata.kv_cu_seq_lens,
|
|
/*alibi_slope=*/std::nullopt,
|
|
/*attn_bias=*/std::nullopt,
|
|
/*q_quant_scale=*/std::nullopt,
|
|
/*k_quant_scale=*/std::nullopt,
|
|
/*v_quant_scale=*/std::nullopt,
|
|
attn_metadata.block_table,
|
|
attn_metadata.max_query_len,
|
|
attn_metadata.max_seq_len,
|
|
scale_,
|
|
attn_metadata.is_causal,
|
|
sliding_window_,
|
|
/*window_size_right=*/-1,
|
|
attn_metadata.compute_dtype,
|
|
/*return_lse=*/false);
|
|
}
|
|
|
|
void AttentionImpl::decoder_forward(torch::Tensor& query,
|
|
torch::Tensor& output,
|
|
const torch::Tensor& k_cache,
|
|
const std::optional<torch::Tensor>& v_cache,
|
|
const AttentionMetadata& attn_metadata) {
|
|
int64_t head_size_v = enable_mla_ ? v_head_dim_ : head_size_;
|
|
query = query.view({-1, 1, num_heads_, head_size_});
|
|
output = output.view({-1, 1, num_heads_, head_size_v});
|
|
std::optional<torch::Tensor> output_lse = std::nullopt;
|
|
|
|
int64_t block_aligned_max_seq_len =
|
|
attn_metadata.block_table.size(-1) * k_cache.size(2);
|
|
|
|
xllm::kernel::ilu::batch_decode(query,
|
|
k_cache,
|
|
output,
|
|
attn_metadata.block_table,
|
|
attn_metadata.kv_seq_lens,
|
|
v_cache,
|
|
output_lse,
|
|
/*q_quant_scale=*/std::nullopt,
|
|
/*k_quant_scale=*/std::nullopt,
|
|
/*v_quant_scale=*/std::nullopt,
|
|
/*out_quant_scale=*/std::nullopt,
|
|
/*alibi_slope=*/std::nullopt,
|
|
attn_metadata.attn_mask,
|
|
attn_metadata.compute_dtype,
|
|
block_aligned_max_seq_len,
|
|
sliding_window_,
|
|
/*window_size_right=*/-1,
|
|
scale_,
|
|
/*return_lse=*/false,
|
|
attn_metadata.is_causal,
|
|
/*kv_cache_quant_bit_size=*/-1);
|
|
}
|
|
|
|
} // namespace layer
|
|
} // namespace xllm
|