Replaces all Python PyTorch fallback attention with native ixformer kernels: Decode path: - ≤32K: paged_attention_v1 (5D KV layout, x=8) — verified on real BI-V100 - >32K: paged_attention_v2 (5D→4D permute) — verified 65K+ on real BI-V100 - Removes _forward_decode_pytorch Python fallback entirely Prefill path (profiling): - _run_sdpa_fallback now uses ixformer.flash_attn_varlen_func - head_dim=256 verified correct (diff<0.004) and 1.7x faster than PyTorch - Falls back to Q-tiling pure-math if ixformer unavailable Also includes: MoE kernel integration, GDN C++ kernels, diagnostic scripts, xllm upstream layer/kernel references, .dockerignore cleanup. All changes verified on real BI-V100 hardware (single card).
153 lines
5.6 KiB
C++
153 lines
5.6 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/npu/npu_ops_api.h"
|
|
#include "kernels/ops_api.h"
|
|
|
|
DECLARE_bool(enable_chunked_prefill);
|
|
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),
|
|
num_kv_heads_(num_kv_heads),
|
|
sliding_window_(sliding_window),
|
|
scale_(scale) {
|
|
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 = 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;
|
|
|
|
torch::Tensor k_cache = kv_cache.get_k_cache();
|
|
torch::Tensor v = value.view({-1, num_kv_heads_, head_size_});
|
|
std::optional<torch::Tensor> v_cache = kv_cache.get_v_cache();
|
|
|
|
// Reshape and cache key/value
|
|
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 (only_prefill) {
|
|
prefill_forward(query, key, value, output, k_cache, v_cache, attn_metadata);
|
|
} else {
|
|
decoder_forward(query, output, k_cache, v_cache, attn_metadata);
|
|
}
|
|
|
|
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) {
|
|
query = query.view({-1, num_heads_, head_size_});
|
|
output = output.view({-1, num_heads_, head_size_});
|
|
|
|
if (attn_metadata.is_prefill) {
|
|
key = key.view({-1, num_kv_heads_, head_size_});
|
|
value = value.view({-1, num_kv_heads_, head_size_});
|
|
|
|
xllm::kernel::npu::batch_prefill(query,
|
|
key,
|
|
value,
|
|
attn_metadata.attn_mask,
|
|
attn_metadata.kv_seq_lens_host,
|
|
scale_,
|
|
output);
|
|
} else if (attn_metadata.is_chunked_prefill) {
|
|
xllm::kernel::npu::batch_prefill(query,
|
|
k_cache,
|
|
v_cache.value(),
|
|
attn_metadata.attn_mask,
|
|
attn_metadata.kv_seq_lens_host,
|
|
scale_,
|
|
output);
|
|
}
|
|
}
|
|
|
|
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) {
|
|
query = query.view({-1, 1, num_heads_, head_size_});
|
|
output = output.view({-1, 1, num_heads_, head_size_});
|
|
|
|
torch::Tensor kv_seq_lens;
|
|
if (attn_metadata.kv_seq_lens_host.defined()) {
|
|
kv_seq_lens = attn_metadata.kv_seq_lens_host;
|
|
} else {
|
|
// Fallback if host tensor isn't prepared.
|
|
kv_seq_lens = attn_metadata.kv_seq_lens;
|
|
}
|
|
|
|
if (attn_metadata.paged_attention_tiling_data.defined()) {
|
|
// Use CustomPagedAttention for ACL graph mode to avoid .to(kCPU) operations
|
|
|
|
xllm::kernel::npu::batch_decode_acl_graph(
|
|
query,
|
|
k_cache,
|
|
v_cache.value_or(torch::Tensor()),
|
|
scale_,
|
|
attn_metadata.block_table,
|
|
kv_seq_lens,
|
|
attn_metadata.paged_attention_tiling_data,
|
|
output);
|
|
} else {
|
|
// Standard PagedAttention path
|
|
xllm::kernel::npu::batch_decode(query,
|
|
k_cache,
|
|
v_cache.value_or(torch::Tensor()),
|
|
scale_,
|
|
attn_metadata.block_table,
|
|
kv_seq_lens,
|
|
output);
|
|
}
|
|
}
|
|
|
|
} // namespace layer
|
|
} // namespace xllm
|