@@ -0,0 +1,19 @@
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
# CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
# Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
# See LICENSE in the root of the software repository for the full text of the License.
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
|
||||
file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*)
|
||||
if(NOT ENABLE_TEST AND NOT BENCHMARK)
|
||||
list(REMOVE_ITEM CURRENT_DIRS tests)
|
||||
endif()
|
||||
foreach(SUB_DIR ${CURRENT_DIRS})
|
||||
if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt")
|
||||
add_subdirectory(${SUB_DIR})
|
||||
endif()
|
||||
endforeach()
|
||||
222
csrc/attention/kv_quant_sparse_flash_attention/README.md
Normal file
222
csrc/attention/kv_quant_sparse_flash_attention/README.md
Normal file
@@ -0,0 +1,222 @@
|
||||
# KvQuantSparseFlashAttention
|
||||
|
||||
## 产品支持情况
|
||||
|
||||
|产品 | 是否支持 |
|
||||
|:----------------------------|:-----------:|
|
||||
|<term>Ascend 950PR/Ascend 950DT</term>| √ |
|
||||
|<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>| √ |
|
||||
|<term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>| √ |
|
||||
|<term>Atlas 200I/500 A2 推理产品</term>| × |
|
||||
|<term>Atlas 推理系列加速卡产品</term>| × |
|
||||
|<term>Atlas 训练系列产品</term>| × |
|
||||
|
||||
## 功能说明
|
||||
|
||||
- API功能:`kv_quant_sparse_flash_attention`在`sparse_flash_attention`的基础上支持了[Per-Token-Head-Tile-128量化]输入。随着大模型上下文长度的增加,Sparse Attention的重要性与日俱增,这一技术通过“只计算关键部分”大幅减少计算量,然而会引入大量的离散访存,造成数据搬运时间增加,进而影响整体性能。
|
||||
|
||||
- 计算公式:
|
||||
|
||||
$$
|
||||
Attention=\text{softmax}(\frac{Q @ \text{Dequant}({\tilde{K}^{INT8}},{Scale_K})^T}{\sqrt{d_k}})@\text{Dequant}(\tilde{V}^{INT8},{Scale_V}),
|
||||
$$
|
||||
|
||||
其中$\tilde{K},\tilde{V}$为基于某种选择算法(如`LightningIndexer`)得到的重要性较高的Key和Value,一般具有稀疏或分块稀疏的特征,$d_k$为$Q,\tilde{K}$每一个头的维度,$\text{Dequant}(\cdot,\cdot)$为反量化函数。
|
||||
本次公布的`kv_quant_sparse_flash_attention`是面向Sparse Attention的全新算子,针对离散访存进行了指令缩减及搬运聚合的细致优化。
|
||||
|
||||
## 参数说明
|
||||
|
||||
> **说明:**<br>
|
||||
> 参数维度含义:B表示Batch Size、Q_S和KV_S分别表示query和key/value的Sequence Length、Q_N和KV_N分别表示query和key/value的Head Num、Q_D和KV_D分别表示query和key/value的Head Dim、Q_T和KV_T分别表示query和key/value的Total Tokens、sparse_size表示一次离散选取的block数、block_num和block_size分别表示PageAttention场景下的block总数和每个block的token数。
|
||||
|
||||
<table style="undefined;table-layout: fixed; width: 1080px"><colgroup>
|
||||
<col style="width: 200px">
|
||||
<col style="width: 150px">
|
||||
<col style="width: 280px">
|
||||
<col style="width: 330px">
|
||||
<col style="width: 120px">
|
||||
</colgroup>
|
||||
<thead>
|
||||
<tr>
|
||||
<th>参数名</th>
|
||||
<th>输入/输出/属性</th>
|
||||
<th>描述</th>
|
||||
<th>数据类型</th>
|
||||
<th>数据格式</th>
|
||||
</tr></thead>
|
||||
<tbody>
|
||||
<tr>
|
||||
<td>query</td>
|
||||
<td>输入</td>
|
||||
<td>attention结构的Q输入,不支持非连续。query由相同数据类型的q_nope和q_rope按D维度拼接得到。layout_query为"BSND"时shape为[B, Q_S, Q_N, Q_D]。layout_query为"TND"时shape为[Q_T, Q_N, Q_D]。其中Q_D值仅支持576,即q_nope+q_rope=512+64;Q_N值支持1/2/4/8/16/32/48/64/128。</td>
|
||||
<td>FLOAT16、BFLOAT16</td>
|
||||
<td>ND</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>key</td>
|
||||
<td>输入</td>
|
||||
<td>attention结构的K输入,不支持非连续。k_nope、query相同数据类型的k_rope和float32的量化参数按D维度拼接得到。layout_kv为"BSND"时shape为[B, KV_S, KV_N, KV_D]。layout_kv为"TND"时shape为[KV_T, KV_N, KV_D]。layout_kv为"PA_BSND"时shape为[block_num, block_size, KV_N, KV_D],其中block_num为PageAttention时block总数,block_size为一个block的token数,block_size取值为16的整数倍,最大支持到1024。KV_N仅支持1;KV_D值仅支持656,即nope+rope*2+dequant_scale*4=512+64*2+4*4。</td>
|
||||
<td>FLOAT8_E4M3、INT8、HIFLOAT8</td>
|
||||
<td>ND</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>value</td>
|
||||
<td>输入</td>
|
||||
<td>attention结构的V输入,不支持非连续。</td>
|
||||
<td>FLOAT8_E4M3、INT8、HIFLOAT8</td>
|
||||
<td>ND</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>sparse_indices</td>
|
||||
<td>输入</td>
|
||||
<td>代表离散取kvCache的索引,不支持非连续。layout_query为"BSND"时shape为[B, Q_S, KV_N, sparse_size]。layout_query为"TND"时shape为[Q_T, KV_N, sparse_size]。其中sparse_size为一次离散选取的block数,需要保证每行有效值均在前半部分,无效值均在后半部分,且需要满足sparse_size大于0。当key和value的数据类型为hifloat8时,sparse_size仅支持2048。</td>
|
||||
<td>INT32</td>
|
||||
<td>ND</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>scale_value</td>
|
||||
<td>属性</td>
|
||||
<td>公式中d<sub>k</sub>开根号的倒数,代表缩放系数,作为query和key矩阵乘后Muls的scalar值。</td>
|
||||
<td>FLOAT</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>key_quant_mode</td>
|
||||
<td>属性</td>
|
||||
<td>代表key的量化模式,仅支持传入2,代表per_tile量化模式。</td>
|
||||
<td>INT64</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>value_quant_mode</td>
|
||||
<td>属性</td>
|
||||
<td>代表value的量化模式,仅支持传入2,代表per_tile量化模式。</td>
|
||||
<td>INT64</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>key_dequant_scale</td>
|
||||
<td>输入</td>
|
||||
<td>预留参数。</td>
|
||||
<td>-</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>value_dequant_scale</td>
|
||||
<td>输入</td>
|
||||
<td>预留参数。</td>
|
||||
<td>-</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>block_table</td>
|
||||
<td>输入</td>
|
||||
<td>表示PageAttention中kvCache存储使用的block映射表。shape为[B, KV_S_max/block_size],其中第一维长度为B,第二维长度不小于所有batch中最大的KV_S对应的block数量,即KV_S_max / block_size向上取整。</td>
|
||||
<td>INT32</td>
|
||||
<td>ND</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>actual_seq_lengths_query</td>
|
||||
<td>输入</td>
|
||||
<td>表示不同Batch中query的有效token数。如果不指定seqlen可传入None,表示和query shape的Q_S长度相同。shape为[B,]。每个Batch的有效token数不超过query中的Q_S大小且不小于0。当layout_query为"TND"时,该入参必须传入,且以该入参元素的数量作为B值,该入参中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。</td>
|
||||
<td>INT32</td>
|
||||
<td>ND</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>actual_seq_lengths_kv</td>
|
||||
<td>输入</td>
|
||||
<td>表示不同Batch中key和value的有效token数。如果不指定None,表示和key的shape的KV_S长度相同。shape为[B,]。每个Batch的有效token数不超过key/value中的KV_S大小且不小于0。当layout_kv为"TND"或"PA_BSND"时,该入参必须传入,layout_kv为"TND"时,该参数中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。</td>
|
||||
<td>INT32</td>
|
||||
<td>ND</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>sparse_block_size</td>
|
||||
<td>属性</td>
|
||||
<td>代表sparse阶段的block大小。sparse_block_size为1时,为Token-wise稀疏化场景;sparse_block_size大于1且小于等于128时,为Block-wise稀疏化场景,块内token共享相同的稀疏化决策。</td>
|
||||
<td>INT64</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>layout_query</td>
|
||||
<td>属性</td>
|
||||
<td>用于标识输入query的数据排布格式,默认值"BSND",支持传入BSND和TND。</td>
|
||||
<td>STRING</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>layout_kv</td>
|
||||
<td>属性</td>
|
||||
<td>用于标识输入key的数据排布格式,默认值"BSND",支持传入BSND、TND和PA_BSND,PA_BSND在开启PageAttention时使用。</td>
|
||||
<td>STRING</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>sparse_mode</td>
|
||||
<td>属性</td>
|
||||
<td>表示sparse的模式。sparse_mode为0时,代表全部计算。sparse_mode为3时,代表rightDownCausal模式的mask,对应以右下顶点往左上为划分线的下三角场景。</td>
|
||||
<td>INT64</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>pre_tokens</td>
|
||||
<td>属性</td>
|
||||
<td>用于稀疏计算,表示attention需要和前几个Token计算关联,仅支持2^63-1。</td>
|
||||
<td>INT64</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>next_tokens</td>
|
||||
<td>属性</td>
|
||||
<td>用于稀疏计算,表示attention需要和后几个Token计算关联,仅支持2^63-1。</td>
|
||||
<td>INT64</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>attention_mode</td>
|
||||
<td>属性</td>
|
||||
<td>表示attention的模式,仅支持传入2,表示MLA-absorb模式,即QK的D包含rope和nope两部分,且KV是同一份。</td>
|
||||
<td>INT64</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>quant_scale_repo_mode</td>
|
||||
<td>属性</td>
|
||||
<td>表示量化参数的存放模式,仅支持传入1,表示combine模式,即量化参数和数据混合存放。</td>
|
||||
<td>INT64</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>tile_size</td>
|
||||
<td>属性</td>
|
||||
<td>表示per_tile时每个参数对应的数据块大小,仅在per_tile时有效,仅支持128。</td>
|
||||
<td>INT64</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>rope_head_dim</td>
|
||||
<td>属性</td>
|
||||
<td>表示MLA架构下的rope_head_dim大小,仅在attention_mode为2时有效,仅支持64。</td>
|
||||
<td>INT64</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>output</td>
|
||||
<td>输出</td>
|
||||
<td>代表公式中的输出Attention。输出shape与入参query的shape保持一致,layout_query为"BSND"时shape为[B, Q_S, Q_N, Q_out_D],layout_query为"TND"时shape为[Q_T, Q_N, Q_out_D],其中Q_out_D = Q_D - rope_head_dim。</td>
|
||||
<td>FLOAT16、BFLOAT16</td>
|
||||
<td>ND</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
## 约束说明
|
||||
|
||||
- 该接口支持图模式。
|
||||
- 参数query shape中:<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>、<term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:Q_N不支持48。
|
||||
- 参数key、value数据类型要求:
|
||||
- <term>Ascend 950PR/Ascend 950DT</term>:仅支持float8_e4m3、int8、hifloat8数据类型。
|
||||
- <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>、<term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:仅支持int8数据类型。
|
||||
- 参数sparse\_block\_size:
|
||||
- <term>Ascend 950PR/Ascend 950DT</term>:只支持sparse\_block\_size为1。
|
||||
- <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>、<term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:支持[1,16],且要求是2的幂次方,在PageAttention场景下要求sparse\_block\_size整除block\_size
|
||||
- 非PageAttention场景layout\_query和layout\_kv取值需要保持一致。
|
||||
@@ -0,0 +1,168 @@
|
||||
/*
|
||||
* Copyright (c) Huawei Technologies Co., Ltd. 2026. 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
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* 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.
|
||||
*/
|
||||
#ifndef KV_QUANT_SPARSE_FLASH_ATTENTION_TORCH_ADPT_H
|
||||
#define KV_QUANT_SPARSE_FLASH_ATTENTION_TORCH_ADPT_H
|
||||
|
||||
namespace vllm_ascend {
|
||||
|
||||
namespace {
|
||||
|
||||
std::tuple<at::Tensor, at::Tensor, at::Tensor>
|
||||
construct_kv_quant_sparse_flash_attention_output_tensor(
|
||||
const at::Tensor &query,
|
||||
const at::Tensor &key,
|
||||
const std::string &layout_query_str,
|
||||
const std::string &layout_kv_str,
|
||||
int64_t rope_head_dim,
|
||||
bool return_softmax_lse)
|
||||
{
|
||||
constexpr int64_t SIZE = 8;
|
||||
constexpr int64_t DIM_0 = 0;
|
||||
constexpr int64_t DIM_1 = 1;
|
||||
constexpr int64_t DIM_2 = 2;
|
||||
constexpr int64_t DIM_3 = 3;
|
||||
constexpr int64_t DIM_4 = 4;
|
||||
|
||||
TORCH_CHECK(layout_query_str == "BSND" || layout_query_str == "TND",
|
||||
"The layout of query only support BSND and TND, but got ",
|
||||
layout_query_str);
|
||||
for (size_t i = 0; i < query.sizes().size(); i++) {
|
||||
TORCH_CHECK(query.size(i) > 0,
|
||||
"All values within query's shape should be greater than 0, but shape[",
|
||||
i, "] is ", query.size(i));
|
||||
}
|
||||
|
||||
at::SmallVector<int64_t, SIZE> output_size;
|
||||
if (layout_query_str == "BSND") {
|
||||
TORCH_CHECK(query.dim() == DIM_4,
|
||||
"When the layout of query is BSND, the query dimension must be 4, but got ",
|
||||
query.dim());
|
||||
output_size = {query.size(DIM_0), query.size(DIM_1), query.size(DIM_2),
|
||||
query.size(DIM_3) - rope_head_dim};
|
||||
} else {
|
||||
TORCH_CHECK(query.dim() == DIM_3,
|
||||
"When the layout of query is TND, the query dimension must be 3, but got ",
|
||||
query.dim());
|
||||
output_size = {query.size(DIM_0), query.size(DIM_1),
|
||||
query.size(DIM_2) - rope_head_dim};
|
||||
}
|
||||
|
||||
at::Tensor attention_output =
|
||||
at::empty(output_size, query.options().dtype(query.dtype()));
|
||||
at::SmallVector<int64_t, SIZE> softmax_size;
|
||||
if (return_softmax_lse) {
|
||||
if (query.dim() == DIM_3) {
|
||||
const int64_t kv_head_dim =
|
||||
layout_kv_str == "PA_BSND" ? key.size(DIM_2) : key.size(DIM_1);
|
||||
softmax_size = {kv_head_dim, query.size(DIM_0),
|
||||
query.size(DIM_1) / kv_head_dim};
|
||||
} else {
|
||||
softmax_size = {query.size(DIM_0), key.size(DIM_2),
|
||||
query.size(DIM_1),
|
||||
query.size(DIM_2) / key.size(DIM_2)};
|
||||
}
|
||||
} else {
|
||||
softmax_size = {0};
|
||||
}
|
||||
|
||||
at::Tensor softmax_max =
|
||||
at::empty(softmax_size, query.options().dtype(at::kFloat));
|
||||
at::Tensor softmax_sum =
|
||||
at::empty(softmax_size, query.options().dtype(at::kFloat));
|
||||
return std::tuple<at::Tensor, at::Tensor, at::Tensor>(
|
||||
attention_output, softmax_max, softmax_sum);
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
std::tuple<at::Tensor, at::Tensor, at::Tensor>
|
||||
npu_kv_quant_sparse_flash_attention(
|
||||
const at::Tensor &query,
|
||||
const at::Tensor &key,
|
||||
const at::Tensor &value,
|
||||
const at::Tensor &sparse_indices,
|
||||
double scale_value,
|
||||
int64_t key_quant_mode,
|
||||
int64_t value_quant_mode,
|
||||
const c10::optional<at::Tensor> &key_dequant_scale,
|
||||
const c10::optional<at::Tensor> &value_dequant_scale,
|
||||
const c10::optional<at::Tensor> &block_table,
|
||||
const c10::optional<at::Tensor> &actual_seq_lengths_query,
|
||||
const c10::optional<at::Tensor> &actual_seq_lengths_kv,
|
||||
int64_t sparse_block_size,
|
||||
c10::string_view layout_query,
|
||||
c10::string_view layout_kv,
|
||||
int64_t sparse_mode,
|
||||
int64_t pre_tokens,
|
||||
int64_t next_tokens,
|
||||
int64_t attention_mode,
|
||||
int64_t quant_scale_repo_mode,
|
||||
int64_t tile_size,
|
||||
int64_t rope_head_dim,
|
||||
bool return_softmax_lse)
|
||||
{
|
||||
TORCH_CHECK(query.numel() > 0, "Tensor query is empty.");
|
||||
TORCH_CHECK(key.numel() > 0, "Tensor key is empty.");
|
||||
TORCH_CHECK(value.numel() > 0, "Tensor value is empty.");
|
||||
TORCH_CHECK(sparse_indices.numel() > 0, "Tensor sparse_indices is empty.");
|
||||
|
||||
std::string layout_query_str = std::string(layout_query);
|
||||
std::string layout_kv_str = std::string(layout_kv);
|
||||
|
||||
auto output = construct_kv_quant_sparse_flash_attention_output_tensor(
|
||||
query, key, layout_query_str, layout_kv_str, rope_head_dim,
|
||||
return_softmax_lse);
|
||||
at::Tensor attention_output = std::get<0>(output);
|
||||
at::Tensor softmax_max = std::get<1>(output);
|
||||
at::Tensor softmax_sum = std::get<2>(output);
|
||||
|
||||
char *layout_query_ptr = const_cast<char *>(layout_query_str.c_str());
|
||||
char *layout_kv_ptr = const_cast<char *>(layout_kv_str.c_str());
|
||||
|
||||
EXEC_NPU_CMD(
|
||||
aclnnKvQuantSparseFlashAttention,
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
sparse_indices,
|
||||
key_dequant_scale,
|
||||
value_dequant_scale,
|
||||
block_table,
|
||||
actual_seq_lengths_query,
|
||||
actual_seq_lengths_kv,
|
||||
scale_value,
|
||||
key_quant_mode,
|
||||
value_quant_mode,
|
||||
sparse_block_size,
|
||||
layout_query_ptr,
|
||||
layout_kv_ptr,
|
||||
sparse_mode,
|
||||
pre_tokens,
|
||||
next_tokens,
|
||||
attention_mode,
|
||||
quant_scale_repo_mode,
|
||||
tile_size,
|
||||
rope_head_dim,
|
||||
return_softmax_lse,
|
||||
attention_output,
|
||||
softmax_max,
|
||||
softmax_sum);
|
||||
return std::tuple<at::Tensor, at::Tensor, at::Tensor>(
|
||||
attention_output, softmax_max, softmax_sum);
|
||||
}
|
||||
} // namespace vllm_ascend
|
||||
|
||||
#endif // KV_QUANT_SPARSE_FLASH_ATTENTION_TORCH_ADPT_H
|
||||
@@ -0,0 +1,32 @@
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
# CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
# Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
# See LICENSE in the root of the software repository for the full text of the License.
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
add_op_to_compiled_list()
|
||||
|
||||
set(KV_QUANT_SPARSE_FLASH_ATTENTION_SKIP_HEADER TRUE CACHE INTERNAL "Skip packaging header for this operator")
|
||||
|
||||
if (BUILD_OPEN_PROJECT)
|
||||
set(kv_quant_sparse_flash_attention_depends attention/common attention/sparse_flash_attention CACHE INTERNAL "Dependencies for kv_quant_sparse_flash_attention")
|
||||
target_sources(op_host_aclnn PRIVATE
|
||||
kv_quant_sparse_flash_attention_def.cpp
|
||||
)
|
||||
endif()
|
||||
|
||||
add_ops_compile_options(
|
||||
OP_NAME KvQuantSparseFlashAttention
|
||||
OPTIONS --cce-auto-sync=off
|
||||
-Wno-deprecated-declarations
|
||||
-mllvm -cce-vf-remove-membar=false
|
||||
-mllvm -cce-aicore-hoist-movemask=false
|
||||
)
|
||||
|
||||
if (NOT BUILD_OPS_RTY_KERNEL)
|
||||
add_modules_sources(OPTYPE kv_quant_sparse_flash_attention ACLNNTYPE aclnn)
|
||||
endif()
|
||||
|
||||
@@ -0,0 +1,171 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_flash_attention_def.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "register/op_def_registry.h"
|
||||
|
||||
namespace ops {
|
||||
class KvQuantSparseFlashAttention : public OpDef {
|
||||
public:
|
||||
explicit KvQuantSparseFlashAttention(const char *name) : OpDef(name)
|
||||
{
|
||||
this->Input("query")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT16, ge::DT_BF16})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("key")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT8, ge::DT_INT8})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("value")
|
||||
.ParamType(REQUIRED)
|
||||
.Follow("key")
|
||||
.AutoContiguous();
|
||||
this->Input("sparse_indices")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT32, ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("key_dequant_scale")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("value_dequant_scale")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("block_table")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_INT32, ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("actual_seq_lengths_query")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_INT32, ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("actual_seq_lengths_kv")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_INT32, ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Output("attention_out")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT16, ge::DT_BF16})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Output("softmax_max")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Output("softmax_sum")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Attr("scale_value").AttrType(REQUIRED).Float(1.0);
|
||||
this->Attr("key_quant_mode").AttrType(REQUIRED).Int(1);
|
||||
this->Attr("value_quant_mode").AttrType(REQUIRED).Int(1);
|
||||
this->Attr("sparse_block_size").AttrType(OPTIONAL).Int(1);
|
||||
this->Attr("layout_query").AttrType(OPTIONAL).String("BSND");
|
||||
this->Attr("layout_kv").AttrType(OPTIONAL).String("BSND");
|
||||
this->Attr("sparse_mode").AttrType(OPTIONAL).Int(3); // 3:默认值,只计算下三角
|
||||
this->Attr("pre_tokens").AttrType(OPTIONAL).Int(INT64_MAX);
|
||||
this->Attr("next_tokens").AttrType(OPTIONAL).Int(INT64_MAX);
|
||||
this->Attr("attention_mode").AttrType(OPTIONAL).Int(0);
|
||||
this->Attr("quant_scale_repo_mode").AttrType(OPTIONAL).Int(1);
|
||||
this->Attr("tile_size").AttrType(OPTIONAL).Int(128); // 128:默认值
|
||||
this->Attr("rope_head_dim").AttrType(OPTIONAL).Int(64); // 64:默认值
|
||||
this->Attr("return_softmax_lse").AttrType(OPTIONAL).Bool(false);
|
||||
OpAICoreConfig aicore_config;
|
||||
aicore_config.DynamicCompileStaticFlag(true)
|
||||
.DynamicFormatFlag(true)
|
||||
.DynamicRankSupportFlag(true)
|
||||
.DynamicShapeSupportFlag(true)
|
||||
.NeedCheckSupportFlag(false)
|
||||
.PrecisionReduceFlag(true);
|
||||
this->AICore().AddConfig("ascend910b", aicore_config);
|
||||
this->AICore().AddConfig("ascend910_93", aicore_config);
|
||||
|
||||
OpAICoreConfig aicore_config_95;
|
||||
aicore_config_95.Input("query")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16})
|
||||
.FormatList({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
aicore_config_95.Input("key")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_INT8})
|
||||
.FormatList({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
aicore_config_95.Input("value")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_INT8})
|
||||
.FormatList({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
aicore_config_95.Input("sparse_indices")
|
||||
.ParamType(REQUIRED)
|
||||
.DataTypeList({ge::DT_INT32})
|
||||
.FormatList({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
aicore_config_95.Input("key_dequant_scale")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataTypeList({ge::DT_FLOAT})
|
||||
.FormatList({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
aicore_config_95.Input("value_dequant_scale")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataTypeList({ge::DT_FLOAT})
|
||||
.FormatList({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
aicore_config_95.Input("block_table")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataTypeList({ge::DT_INT32})
|
||||
.FormatList({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
aicore_config_95.Input("actual_seq_lengths_query")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataTypeList({ge::DT_INT32})
|
||||
.FormatList({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
aicore_config_95.Input("actual_seq_lengths_kv")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataTypeList({ge::DT_INT32})
|
||||
.FormatList({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
aicore_config_95.Output("attention_out")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16})
|
||||
.FormatList({ge::FORMAT_ND});
|
||||
aicore_config_95.Output("softmax_max")
|
||||
.ParamType(REQUIRED)
|
||||
.DataTypeList({ge::DT_FLOAT})
|
||||
.FormatList({ge::FORMAT_ND});
|
||||
aicore_config_95.Output("softmax_sum")
|
||||
.ParamType(REQUIRED)
|
||||
.DataTypeList({ge::DT_FLOAT})
|
||||
.FormatList({ge::FORMAT_ND});
|
||||
aicore_config_95.DynamicCompileStaticFlag(true)
|
||||
.DynamicFormatFlag(true)
|
||||
.DynamicRankSupportFlag(true)
|
||||
.DynamicShapeSupportFlag(true)
|
||||
.NeedCheckSupportFlag(false)
|
||||
.PrecisionReduceFlag(true);
|
||||
this->AICore().AddConfig("ascend950", aicore_config_95);
|
||||
}
|
||||
};
|
||||
OP_ADD(KvQuantSparseFlashAttention);
|
||||
} // namespace ops
|
||||
@@ -0,0 +1,153 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_flash_attention_infershape.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include <graph/utils/type_utils.h>
|
||||
#include <register/op_impl_registry.h>
|
||||
#include "err/ops_err.h"
|
||||
|
||||
using namespace ge;
|
||||
|
||||
#ifdef OP_LOGE_WITH_INVALID_INPUT
|
||||
#undef OP_LOGE_WITH_INVALID_INPUT
|
||||
#endif
|
||||
#define OP_LOGE_WITH_INVALID_INPUT(opname, param) \
|
||||
OP_LOGE(opname, "Invalid input: %s.", param)
|
||||
|
||||
namespace ops {
|
||||
constexpr size_t QUERY_INPUT_INDEX = 0;
|
||||
constexpr size_t KEY_INPUT_INDEX = 1;
|
||||
constexpr uint32_t LAYOUT_QUERY_ATTR_INDEX = 4;
|
||||
constexpr uint32_t LAYOUT_KV_ATTR_INDEX = 5;
|
||||
constexpr uint32_t ROPE_HEAD_DIM_ATTR_INDEX = 12;
|
||||
constexpr uint32_t RETURN_SOFTMAX_LSE_INDEX = 13;
|
||||
constexpr uint32_t DIM_INDEX_0 = 0;
|
||||
constexpr uint32_t DIM_INDEX_1 = 1;
|
||||
constexpr uint32_t DIM_INDEX_2 = 2;
|
||||
constexpr uint32_t DIM_INDEX_3 = 3;
|
||||
constexpr uint32_t DIM_NUM_1 = 1;
|
||||
constexpr uint32_t DIM_NUM_3 = 3;
|
||||
constexpr uint32_t DIM_NUM_4 = 4;
|
||||
constexpr uint32_t OUTPUT_INDEX_0 = 0;
|
||||
constexpr uint32_t OUTPUT_INDEX_1 = 1;
|
||||
constexpr uint32_t OUTPUT_INDEX_2 = 2;
|
||||
|
||||
ge::graphStatus InferShapeKvQuantSparseFlashAttention(gert::InferShapeContext *context)
|
||||
{
|
||||
OP_CHECK_IF(context == nullptr, OP_LOGE_WITH_INVALID_INPUT("KvQuantSparseFlashAttention", "InferShapeContext"),
|
||||
return ge::GRAPH_FAILED);
|
||||
const gert::Shape *queryShape = context->GetInputShape(QUERY_INPUT_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, queryShape);
|
||||
const gert::Shape *keyShape = context->GetInputShape(KEY_INPUT_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, keyShape);
|
||||
|
||||
gert::Shape *attentionOutShape = context->GetOutputShape(0);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, attentionOutShape);
|
||||
gert::Shape *softmaxMaxShape = context->GetOutputShape(OUTPUT_INDEX_1);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, softmaxMaxShape);
|
||||
gert::Shape *softmaxSumShape = context->GetOutputShape(OUTPUT_INDEX_2);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, softmaxSumShape);
|
||||
|
||||
auto attrs = context->GetAttrs();
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, attrs);
|
||||
const char *inputLayoutQueryPtr = attrs->GetAttrPointer<char>(LAYOUT_QUERY_ATTR_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, inputLayoutQueryPtr);
|
||||
std::string inputLayoutQueryPtrStr = std::string(inputLayoutQueryPtr);
|
||||
const char *inputLayoutKvPtr = attrs->GetAttrPointer<char>(LAYOUT_KV_ATTR_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, inputLayoutKvPtr);
|
||||
std::string inputLayoutKvPtrStr = std::string(inputLayoutKvPtr);
|
||||
const int64_t ropeHeadDim = *attrs->GetAttrPointer<int64_t>(ROPE_HEAD_DIM_ATTR_INDEX);
|
||||
const bool *lse_flag = attrs->GetAttrPointer<bool>(RETURN_SOFTMAX_LSE_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, lse_flag);
|
||||
bool return_softmax_lse = (lse_flag != nullptr) ? *lse_flag : false;
|
||||
|
||||
*attentionOutShape = *queryShape;
|
||||
if (inputLayoutQueryPtrStr == "BSND") {
|
||||
attentionOutShape->SetDimNum(DIM_NUM_4);
|
||||
attentionOutShape->SetDim(DIM_INDEX_0, queryShape->GetDim(DIM_INDEX_0));
|
||||
attentionOutShape->SetDim(DIM_INDEX_1, queryShape->GetDim(DIM_INDEX_1));
|
||||
attentionOutShape->SetDim(DIM_INDEX_2, queryShape->GetDim(DIM_INDEX_2)); // 2:dim2
|
||||
if(queryShape->GetDim(DIM_INDEX_3) != -1){
|
||||
attentionOutShape->SetDim(DIM_INDEX_3, queryShape->GetDim(DIM_INDEX_3) - ropeHeadDim); // 3:dim3
|
||||
}
|
||||
} else { // TND
|
||||
attentionOutShape->SetDimNum(DIM_NUM_3);
|
||||
attentionOutShape->SetDim(DIM_INDEX_0, queryShape->GetDim(DIM_INDEX_0));
|
||||
attentionOutShape->SetDim(DIM_INDEX_1, queryShape->GetDim(DIM_INDEX_1));
|
||||
if(queryShape->GetDim(DIM_INDEX_2) != -1){
|
||||
attentionOutShape->SetDim(DIM_INDEX_2, queryShape->GetDim(DIM_INDEX_2) - ropeHeadDim); // 2:dim2
|
||||
}
|
||||
}
|
||||
|
||||
if (return_softmax_lse) {
|
||||
if (queryShape->GetDimNum() == DIM_NUM_3) {
|
||||
if (inputLayoutKvPtrStr == "PA_BSND") {
|
||||
softmaxMaxShape->SetDimNum(DIM_NUM_3);
|
||||
softmaxMaxShape->SetDim(DIM_INDEX_0, keyShape->GetDim(DIM_INDEX_2));
|
||||
softmaxMaxShape->SetDim(DIM_INDEX_1, queryShape->GetDim(DIM_INDEX_0));
|
||||
softmaxMaxShape->SetDim(DIM_INDEX_2, queryShape->GetDim(DIM_INDEX_1) / keyShape->GetDim(DIM_INDEX_2));
|
||||
|
||||
softmaxSumShape->SetDimNum(DIM_NUM_3);
|
||||
softmaxSumShape->SetDim(DIM_INDEX_0, keyShape->GetDim(DIM_INDEX_2));
|
||||
softmaxSumShape->SetDim(DIM_INDEX_1, queryShape->GetDim(DIM_INDEX_0));
|
||||
softmaxSumShape->SetDim(DIM_INDEX_2, queryShape->GetDim(DIM_INDEX_1) / keyShape->GetDim(DIM_INDEX_2));
|
||||
} else {
|
||||
softmaxMaxShape->SetDimNum(DIM_NUM_3);
|
||||
softmaxMaxShape->SetDim(DIM_INDEX_0, keyShape->GetDim(DIM_INDEX_1));
|
||||
softmaxMaxShape->SetDim(DIM_INDEX_1, queryShape->GetDim(DIM_INDEX_0));
|
||||
softmaxMaxShape->SetDim(DIM_INDEX_2, queryShape->GetDim(DIM_INDEX_1) / keyShape->GetDim(DIM_INDEX_1));
|
||||
|
||||
softmaxSumShape->SetDimNum(DIM_NUM_3);
|
||||
softmaxSumShape->SetDim(DIM_INDEX_0, keyShape->GetDim(DIM_INDEX_1));
|
||||
softmaxSumShape->SetDim(DIM_INDEX_1, queryShape->GetDim(DIM_INDEX_0));
|
||||
softmaxSumShape->SetDim(DIM_INDEX_2, queryShape->GetDim(DIM_INDEX_1) / keyShape->GetDim(DIM_INDEX_1));
|
||||
}
|
||||
} else {
|
||||
softmaxMaxShape->SetDimNum(DIM_NUM_4);
|
||||
softmaxMaxShape->SetDim(DIM_INDEX_0, queryShape->GetDim(DIM_INDEX_0));
|
||||
softmaxMaxShape->SetDim(DIM_INDEX_1, keyShape->GetDim(DIM_INDEX_2));
|
||||
softmaxMaxShape->SetDim(DIM_INDEX_2, queryShape->GetDim(DIM_INDEX_1));
|
||||
softmaxMaxShape->SetDim(DIM_INDEX_3, queryShape->GetDim(DIM_INDEX_2) / keyShape->GetDim(DIM_INDEX_2));
|
||||
|
||||
softmaxSumShape->SetDimNum(DIM_NUM_4);
|
||||
softmaxSumShape->SetDim(DIM_INDEX_0, queryShape->GetDim(DIM_INDEX_0));
|
||||
softmaxSumShape->SetDim(DIM_INDEX_1, keyShape->GetDim(DIM_INDEX_2));
|
||||
softmaxSumShape->SetDim(DIM_INDEX_2, queryShape->GetDim(DIM_INDEX_1));
|
||||
softmaxSumShape->SetDim(DIM_INDEX_3, queryShape->GetDim(DIM_INDEX_2) / keyShape->GetDim(DIM_INDEX_2));
|
||||
}
|
||||
} else {
|
||||
softmaxMaxShape->SetDimNum(DIM_NUM_1);
|
||||
softmaxMaxShape->SetDim(DIM_INDEX_0, 0);
|
||||
softmaxSumShape->SetDimNum(DIM_NUM_1);
|
||||
softmaxSumShape->SetDim(DIM_INDEX_0, 0);
|
||||
}
|
||||
return GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus InferDataTypeKvQuantSparseFlashAttention(gert::InferDataTypeContext *context)
|
||||
{
|
||||
OP_CHECK_IF(context == nullptr, OP_LOGE_WITH_INVALID_INPUT("KvQuantSparseFlashAttention", "InferShapeContext"),
|
||||
return ge::GRAPH_FAILED);
|
||||
const auto inputDataType = context->GetInputDataType(QUERY_INPUT_INDEX);
|
||||
context->SetOutputDataType(OUTPUT_INDEX_0, inputDataType);
|
||||
context->SetOutputDataType(OUTPUT_INDEX_1, ge::DT_FLOAT);
|
||||
context->SetOutputDataType(OUTPUT_INDEX_2, ge::DT_FLOAT);
|
||||
context->SetOutputDataType(0, inputDataType);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_INFERSHAPE(KvQuantSparseFlashAttention)
|
||||
.InferShape(InferShapeKvQuantSparseFlashAttention)
|
||||
.InferDataType(InferDataTypeKvQuantSparseFlashAttention);
|
||||
} // namespace ops
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,614 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_flash_attention_tiling.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef KV_QUANT_SPARSE_FLASH_ATTENTION_TILING_H
|
||||
#define KV_QUANT_SPARSE_FLASH_ATTENTION_TILING_H
|
||||
|
||||
#include <sstream>
|
||||
#include <graph/utils/type_utils.h>
|
||||
#include <tiling/platform/platform_ascendc.h>
|
||||
#include <exe_graph/runtime/tiling_context.h>
|
||||
#include "register/tilingdata_base.h"
|
||||
#include "exe_graph/runtime/tiling_context.h"
|
||||
#include "platform/soc_spec.h"
|
||||
namespace optiling {
|
||||
// ------------------算子原型索引常量定义----------------
|
||||
// Inputs Index
|
||||
constexpr uint32_t QUERY_INPUT_INDEX = 0;
|
||||
constexpr uint32_t KEY_INPUT_INDEX = 1;
|
||||
constexpr uint32_t VALUE_INPUT_INDEX = 2;
|
||||
constexpr uint32_t SPARSE_INDICES_INPUT_INDEX = 3;
|
||||
constexpr uint32_t KEY_DEQUANT_SCALE_INPUT_INDEX = 3;
|
||||
constexpr uint32_t VALUE_DEQUANT_SCALE_INPUT_INDEX = 3;
|
||||
constexpr uint32_t BLOCK_TABLE_INPUT_INDEX = 6;
|
||||
constexpr uint32_t ACT_SEQ_LEN_Q_INPUT_INDEX = 7;
|
||||
constexpr uint32_t ACT_SEQ_LEN_KV_INPUT_INDEX = 8;
|
||||
// Outputs Index
|
||||
constexpr uint32_t OUTPUT_INDEX = 0;
|
||||
constexpr uint32_t SOFTMAXMAX_INDEX = 1;
|
||||
constexpr uint32_t SOFTMAXSUM_INDEX = 2;
|
||||
// Attributes Index
|
||||
constexpr uint32_t SCALE_VALUE_ATTR_INDEX = 0;
|
||||
constexpr uint32_t KEY_QUANT_MODE_ATTR_INDEX = 1;
|
||||
constexpr uint32_t VALUE_QUANT_MODE_ATTR_INDEX = 2;
|
||||
constexpr uint32_t SPARSE_BLOCK_SIZE_ATTR_INDEX = 3;
|
||||
constexpr uint32_t LAYOUT_QUERY_ATTR_INDEX = 4;
|
||||
constexpr uint32_t LAYOUT_KV_ATTR_INDEX = 5;
|
||||
constexpr uint32_t SPARSE_MODE_ATTR_INDEX = 6;
|
||||
constexpr uint32_t PRE_TOKENS_ATTR_INDEX = 7;
|
||||
constexpr uint32_t NEXT_TOKENS_ATTR_INDEX = 8;
|
||||
constexpr uint32_t ATTENTION_MODE_ATTR_INDEX = 9;
|
||||
constexpr uint32_t QUANT_SCALE_REPO_MODE_ATTR_INDEX = 10;
|
||||
constexpr uint32_t TILE_SIZE_ATTR_INDEX = 11;
|
||||
constexpr uint32_t ROPE_HEAD_DIM_ATTR_INDEX = 12;
|
||||
constexpr uint32_t RETURN_SOFTMAX_LSE_ATTR_INDEX = 13;
|
||||
// Dim Num
|
||||
constexpr size_t DIM_NUM_TWO = 2;
|
||||
constexpr size_t DIM_NUM_THREE = 3;
|
||||
constexpr size_t DIM_NUM_FOUR = 4;
|
||||
// 常量
|
||||
constexpr uint32_t MAX_BLOCK_SIZE = 1024;
|
||||
constexpr uint32_t COPYND2NZ_SRC_STRIDE_LIMITATION = 65535;
|
||||
constexpr uint32_t NUM_BYTES_FLOAT = 4;
|
||||
constexpr uint32_t NUM_BYTES_FLOAT16 = 2;
|
||||
constexpr uint32_t NUM_BYTES_BF16 = 2;
|
||||
constexpr uint32_t BYTE_BLOCK = 32;
|
||||
const uint32_t QSFA_MAX_AIC_CORE_NUM = 26; // 25 + 1 保证数组8字节对齐
|
||||
|
||||
// ------------------公共定义--------------------------
|
||||
enum class QSFALayout : uint32_t {
|
||||
BSND = 0,
|
||||
TND = 1,
|
||||
PA_BSND = 2,
|
||||
};
|
||||
|
||||
struct QSFATilingShapeCompareParam {
|
||||
int64_t B = 1;
|
||||
int64_t S = 1;
|
||||
int64_t N = 1;
|
||||
int64_t D = 1;
|
||||
int64_t T = 1;
|
||||
// PA
|
||||
int64_t Bs = 1;
|
||||
int64_t Bn = 1;
|
||||
};
|
||||
|
||||
enum class KvStorageMode : uint32_t {
|
||||
BATCH_CONTINUOUS = 0,
|
||||
PAGE_ATTENTION = 1
|
||||
};
|
||||
|
||||
enum class QSFAPerfMode : uint32_t {
|
||||
C_TEMPLATE_MODE = 0,
|
||||
V_TEMPLATE_MODE
|
||||
};
|
||||
|
||||
enum class QSFAAxis : uint32_t {
|
||||
B = 0,
|
||||
S = 1,
|
||||
N = 2,
|
||||
D = 3,
|
||||
K = 3, // sparse_indices的K和key的D枚举值相同,表达相同位置, 最后一维
|
||||
T = 5,
|
||||
Bn = 6, // block number
|
||||
Bs = 7, // block size
|
||||
};
|
||||
|
||||
struct QSFARequiredParaInfo {
|
||||
const gert::CompileTimeTensorDesc *desc;
|
||||
const gert::StorageShape *shape;
|
||||
};
|
||||
|
||||
struct QSFAOptionalParaInfo {
|
||||
const gert::CompileTimeTensorDesc *desc;
|
||||
const gert::Tensor *tensor;
|
||||
};
|
||||
|
||||
// -----------算子Tiling入参结构体定义---------------
|
||||
struct QSFAParaInfo {
|
||||
QSFARequiredParaInfo query = {nullptr, nullptr};
|
||||
QSFARequiredParaInfo key = {nullptr, nullptr};
|
||||
QSFARequiredParaInfo value = {nullptr, nullptr};
|
||||
QSFARequiredParaInfo sparseIndices = {nullptr, nullptr};
|
||||
QSFAOptionalParaInfo blockTable = {nullptr, nullptr};
|
||||
QSFAOptionalParaInfo actualSeqLengthsQ = {nullptr, nullptr};
|
||||
QSFAOptionalParaInfo actualSeqLengths = {nullptr, nullptr};
|
||||
QSFAOptionalParaInfo queryRope = {nullptr, nullptr};
|
||||
QSFAOptionalParaInfo keyRope = {nullptr, nullptr};
|
||||
QSFAOptionalParaInfo keyDequantScale = {nullptr, nullptr};
|
||||
QSFAOptionalParaInfo valueDequantScale = {nullptr, nullptr};
|
||||
QSFARequiredParaInfo attenOut = {nullptr, nullptr};
|
||||
QSFARequiredParaInfo softmaxMax = {nullptr, nullptr};
|
||||
QSFARequiredParaInfo softmaxSum = {nullptr, nullptr};
|
||||
|
||||
const char *layoutQuery = nullptr;
|
||||
const char *layoutKV = nullptr;
|
||||
const int64_t *sparseBlockSize = nullptr;
|
||||
const uint32_t *sparseBlockCount = nullptr;
|
||||
const uint32_t *blockSize = nullptr;
|
||||
const float *scaleValue = nullptr;
|
||||
const int64_t *sparseMode = nullptr;
|
||||
const int64_t *attentionMode = nullptr;
|
||||
const int64_t *keyQuantMode = nullptr;
|
||||
const int64_t *valueQuantMode = nullptr;
|
||||
const int64_t *quantScaleRepoMode = nullptr;
|
||||
const int64_t *tileSize = nullptr;
|
||||
const int64_t *ropeHeadDim = nullptr;
|
||||
const int64_t *preTokens = nullptr;
|
||||
const int64_t *nextTokens = nullptr;
|
||||
const bool *returnSoftmaxLse = nullptr;
|
||||
};
|
||||
|
||||
struct InnerSplitParams {
|
||||
uint32_t s1GBaseSize = 1;
|
||||
uint32_t s2BaseSize = 1;
|
||||
};
|
||||
|
||||
// -----------算子TilingData定义---------------
|
||||
BEGIN_TILING_DATA_DEF(KvQuantSparseFlashAttentionBaseParamsMla)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, batchSize)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, seqSize)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, qSeqSize)
|
||||
TILING_DATA_FIELD_DEF(int64_t, blockSize)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, maxBlockNumPerBatch)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, actualLenDimsQ)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, actualLenDimsKV)
|
||||
TILING_DATA_FIELD_DEF(float, scaleValue)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, nNumOfQInOneGroup)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, outputLayout)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, sparseMode)
|
||||
TILING_DATA_FIELD_DEF(int64_t, sparseBlockSize)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, sparseBlockCount)
|
||||
TILING_DATA_FIELD_DEF(int64_t, dSizeVInput)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, isActualLenDimsNull)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, isActualLenDimsKVNull)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, returnSoftmaxLse)
|
||||
END_TILING_DATA_DEF
|
||||
|
||||
REGISTER_TILING_DATA_CLASS(KvQuantSparseFlashAttentionBaseParamsMlaOp, KvQuantSparseFlashAttentionBaseParamsMla)
|
||||
|
||||
BEGIN_TILING_DATA_DEF(KvQuantSparseFlashAttentionSingleCoreParamsMla)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, usedCoreNum);
|
||||
END_TILING_DATA_DEF
|
||||
REGISTER_TILING_DATA_CLASS(KvQuantSparseFlashAttentionSingleCoreParamsMlaOp,
|
||||
KvQuantSparseFlashAttentionSingleCoreParamsMla)
|
||||
|
||||
BEGIN_TILING_DATA_DEF(KvQuantSparseFlashAttentionSingleCoreTensorSizeMla)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, mmResUbSize);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, bmm2ResUbSize);
|
||||
END_TILING_DATA_DEF
|
||||
REGISTER_TILING_DATA_CLASS(KvQuantSparseFlashAttentionSingleCoreTensorSizeMlaOp,
|
||||
KvQuantSparseFlashAttentionSingleCoreTensorSizeMla)
|
||||
|
||||
BEGIN_TILING_DATA_DEF(KvQuantSparseFlashAttentionSplitKVParamsMla)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, s2) // S2切分份数
|
||||
TILING_DATA_FIELD_DEF(uint32_t, accumOutSize) // FD workspace
|
||||
TILING_DATA_FIELD_DEF(uint32_t, logSumExpSize) // FD workspace
|
||||
END_TILING_DATA_DEF
|
||||
REGISTER_TILING_DATA_CLASS(KvQuantSparseFlashAttentionSplitKVParamsMlaOp,
|
||||
KvQuantSparseFlashAttentionSplitKVParamsMla)
|
||||
|
||||
// 内切基本块参数
|
||||
BEGIN_TILING_DATA_DEF(KvQuantSparseFlashAttentionInnerSplitParams)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, mBaseSize)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, s2BaseSize)
|
||||
END_TILING_DATA_DEF
|
||||
REGISTER_TILING_DATA_CLASS(KvQuantSparseFlashAttentionInnerSplitParamsOp,
|
||||
KvQuantSparseFlashAttentionInnerSplitParams)
|
||||
|
||||
BEGIN_TILING_DATA_DEF(KvQuantSparseFlashAttentionTilingDataMla)
|
||||
TILING_DATA_FIELD_DEF_STRUCT(KvQuantSparseFlashAttentionBaseParamsMla, baseParams);
|
||||
TILING_DATA_FIELD_DEF_STRUCT(KvQuantSparseFlashAttentionSplitKVParamsMla, splitKVParams);
|
||||
TILING_DATA_FIELD_DEF_STRUCT(KvQuantSparseFlashAttentionSingleCoreParamsMla, singleCoreParams);
|
||||
TILING_DATA_FIELD_DEF_STRUCT(KvQuantSparseFlashAttentionSingleCoreTensorSizeMla, singleCoreTensorSize);
|
||||
TILING_DATA_FIELD_DEF_STRUCT(KvQuantSparseFlashAttentionInnerSplitParams, innerSplitParams);
|
||||
END_TILING_DATA_DEF
|
||||
REGISTER_TILING_DATA_CLASS(KvQuantSparseFlashAttention, KvQuantSparseFlashAttentionTilingDataMla)
|
||||
|
||||
template <typename T> inline T Align(T num, T rnd)
|
||||
{
|
||||
return (((rnd) == 0) ? 0 : (((num) + (rnd) - 1) / (rnd) * (rnd)));
|
||||
}
|
||||
|
||||
static std::string QSFADataTypeToSerialString(ge::DataType type);
|
||||
std::string QSFATensorDesc2String(const gert::StorageShape *shape, const gert::CompileTimeTensorDesc *tensor);
|
||||
std::string QSFADebugTilingContext(const gert::TilingContext *context);
|
||||
std::string QSFALayoutToSerialString(QSFALayout layout);
|
||||
|
||||
// -----------算子Tiling入参信息类---------------
|
||||
struct QSFATilingInfo {
|
||||
const char *opName = nullptr;
|
||||
fe::PlatFormInfos *platformInfo = nullptr;
|
||||
QSFAParaInfo opParamInfo;
|
||||
|
||||
// Base Param
|
||||
NpuArch npuArch = NpuArch::DAV_2201;
|
||||
bool isA5 = false;
|
||||
uint32_t bSize = 0;
|
||||
uint32_t n1Size = 0;
|
||||
uint32_t n2Size = 0;
|
||||
uint32_t s1Size = 0;
|
||||
int64_t s2Size = 0;
|
||||
uint32_t qHeadDim = 0;
|
||||
uint32_t kHeadDim = 0;
|
||||
uint32_t vHeadDim = 0;
|
||||
uint32_t gSize = 0;
|
||||
uint32_t ropeHeadDim = 0;
|
||||
uint32_t qTSize = 0; // 仅TND时生效
|
||||
uint32_t kvTSize = 0; // 仅TND时生效
|
||||
float scaleValue = 0;
|
||||
uint32_t innerPrecise = 0;
|
||||
uint32_t l2CacheOffFlag = 0;
|
||||
int64_t sparseBlockSize = 0;
|
||||
int64_t sparseBlockCount = 0;
|
||||
|
||||
bool pageAttentionFlag = false;
|
||||
int64_t blockSize = 0;
|
||||
uint32_t blockTypeSize = 0;
|
||||
uint32_t maxBlockNumPerBatch = 0;
|
||||
uint32_t totalBlockNum = 0;
|
||||
|
||||
uint32_t actualLenDimsQ = 0;
|
||||
uint32_t maxActualseq = 0;
|
||||
|
||||
bool actualQSeqLenFlag = false;
|
||||
bool actualSeqLenFlag = false;
|
||||
bool isSameSeqAllKVTensor = true;
|
||||
bool isSameActualseq = true;
|
||||
uint32_t actualLenDimsKV = 0;
|
||||
std::vector<int64_t> kvListSeqLens {};
|
||||
|
||||
uint32_t sparseMode = 0;
|
||||
bool returnSoftmaxLse = false;
|
||||
|
||||
int64_t attentionMode = 0;
|
||||
int64_t keyQuantMode = 0;
|
||||
int64_t valueQuantMode = 0;
|
||||
int64_t quantScaleRepoMode = 0;
|
||||
int64_t tileSize = 0;
|
||||
int64_t preTokens = 0;
|
||||
int64_t nextTokens = 0;
|
||||
|
||||
ge::DataType inputQType = ge::DT_FLOAT16;
|
||||
ge::DataType inputKvType = ge::DT_FLOAT16;
|
||||
ge::DataType outputType = ge::DT_FLOAT16;
|
||||
|
||||
KvStorageMode kvStorageMode = KvStorageMode::BATCH_CONTINUOUS;
|
||||
|
||||
QSFALayout qLayout = QSFALayout::BSND;
|
||||
QSFALayout topkLayout = QSFALayout::BSND;
|
||||
QSFALayout outLayout = QSFALayout::BSND;
|
||||
QSFALayout kvLayout = QSFALayout::BSND;
|
||||
|
||||
ge::DataType inputQRopeType = ge::DT_FLOAT16;
|
||||
ge::DataType inputKRopeType = ge::DT_FLOAT16;
|
||||
|
||||
uint64_t l2CacheSize = 0;
|
||||
int64_t dSizeVInput = 0;
|
||||
};
|
||||
|
||||
// ---------------算子Tiling类---------------
|
||||
class QSFAMlaTiling {
|
||||
public:
|
||||
explicit QSFAMlaTiling(gert::TilingContext *context) : context_(context) {}
|
||||
ge::graphStatus DoOpTiling(QSFATilingInfo *qsfaInfo);
|
||||
|
||||
private:
|
||||
ge::graphStatus SetBlockDim(uint32_t blockDim) const;
|
||||
ge::graphStatus SetTilingKey(uint64_t tilingKey) const;
|
||||
ge::graphStatus SetWorkspaceSize(uint64_t workspaceSize) const;
|
||||
ge::graphStatus SetTilingData(TilingDef &tilingData) const;
|
||||
gert::TilingContext *context_ = nullptr;
|
||||
ge::graphStatus GetPlatformInfo();
|
||||
void GenTilingKey();
|
||||
bool DealSameSeqEachBatch();
|
||||
|
||||
void ZeroTensorProcess() const;
|
||||
void InitParams();
|
||||
|
||||
void Split();
|
||||
bool IsBalanceSplitCore();
|
||||
|
||||
void SplitBalanced();
|
||||
void CalcInnerSize(uint32_t qsfaS2Size);
|
||||
|
||||
bool IsFlashDecode(uint32_t coreNum);
|
||||
|
||||
void FillTilingBaseParamsMla();
|
||||
void FillTilingSplitKVMla();
|
||||
|
||||
void FillTilingSingleCoreParamsMla();
|
||||
void FillTilingSingleCoreTensorSizeMla();
|
||||
void FillTiling();
|
||||
|
||||
void CalcUbBmm();
|
||||
void CheckUbSpace();
|
||||
void NormalCalcFDWorkSpace(const uint32_t actCoreNum);
|
||||
void CalcFDWorkSpace(const uint32_t actCoreNum);
|
||||
void GetWorkspaceSize();
|
||||
|
||||
uint32_t CalcBalanceFDParamNums(const uint32_t actCoreNum) const;
|
||||
|
||||
void CalcBlockDim();
|
||||
|
||||
bool balanceModeFlag_ = false;
|
||||
bool splitKVFlag_ = false;
|
||||
|
||||
uint32_t coreNum_ = 0;
|
||||
QSFAPerfMode perfMode_ = QSFAPerfMode::V_TEMPLATE_MODE;
|
||||
uint32_t kvSplitPart_ = 1;
|
||||
size_t mmResUbSize_ = 0;
|
||||
size_t bmm2ResUbSize_ = 0;
|
||||
size_t qPreSizeMla_ = 0;
|
||||
uint32_t sInnerLoopTimes_ = 0;
|
||||
uint32_t sInnerSize_ = 0;
|
||||
uint32_t sInnerSizeTail_ = 0;
|
||||
uint32_t sInnerSizeAlign_ = 0;
|
||||
uint32_t kvSplit_ = 0;
|
||||
uint32_t usedCoreNum_ = 0;
|
||||
uint32_t formerCoreNum_ = 0;
|
||||
uint32_t blockSplitBn2Range_ = 0;
|
||||
uint32_t tailSplitedBatchRange_ = 0;
|
||||
|
||||
uint32_t aicNum_ = 0;
|
||||
uint32_t aivNum_ = 0;
|
||||
size_t libapiSize_ = 0;
|
||||
|
||||
KvQuantSparseFlashAttentionTilingDataMla tilingData_;
|
||||
uint32_t blockDim_{0};
|
||||
uint64_t workspaceSize_{0};
|
||||
uint64_t tilingKey_{0};
|
||||
|
||||
uint32_t headDimAlign_ = 0;
|
||||
uint32_t mBaseSize_ = 128;
|
||||
uint32_t mFdBaseSize_ = 8;
|
||||
|
||||
QSFATilingInfo *qsfaInfo_ = nullptr;
|
||||
};
|
||||
|
||||
// -----------算子Tiling入参信息解析及Check类---------------
|
||||
class QSFATilingCheck {
|
||||
public:
|
||||
explicit QSFATilingCheck(const QSFATilingInfo &qsfaInfo) : qsfaInfo_(qsfaInfo) {};
|
||||
~QSFATilingCheck() = default;
|
||||
ge::graphStatus Process();
|
||||
private:
|
||||
void Init();
|
||||
void LogErrorDtypeSupport(const std::vector<ge::DataType> &expectDtypeList,
|
||||
const ge::DataType &actualDtype, const std::string &name) const;
|
||||
ge::graphStatus CheckDtypeSupport(const gert::CompileTimeTensorDesc *qsfaDesc,
|
||||
const std::string &name) const;
|
||||
template <typename T> void LogErrorNumberSupport(const std::vector<T> &expectNumberList,
|
||||
const T &actualValue, const std::string &name, const std::string subName) const;
|
||||
template <typename T> void LogErrorDimNumSupport(const std::vector<T> &expectNumberList,
|
||||
const T &actualValue, const std::string &name) const;
|
||||
ge::graphStatus CheckDimNumSupport(const gert::StorageShape *shape,
|
||||
const std::vector<size_t> &qsfaExpectDimNumList, const std::string &name) const;
|
||||
ge::graphStatus CheckDimNumInLayoutSupport(const QSFALayout &layout,
|
||||
const gert::StorageShape *shape, const std::string &name) const;
|
||||
void LogErrorLayoutSupport(const std::vector<QSFALayout> &expectLayoutList,
|
||||
const QSFALayout &actualLayout, const std::string &name) const;
|
||||
ge::graphStatus GetExpectedShape(gert::Shape &shapeExpected,
|
||||
const QSFATilingShapeCompareParam ¶m, const QSFALayout &layout) const;
|
||||
ge::graphStatus CompareShape(QSFATilingShapeCompareParam ¶m,
|
||||
const gert::Shape &shape, const QSFALayout &layout, const std::string &name) const;
|
||||
ge::graphStatus CheckLayoutSupport(const QSFALayout &actualLayout, const std::string &name) const;
|
||||
ge::graphStatus CheckSingleParaQuery() const;
|
||||
ge::graphStatus CheckSingleParaKey() const;
|
||||
ge::graphStatus CheckSingleParaValue() const;
|
||||
ge::graphStatus CheckSingleParaAttenOut() const;
|
||||
ge::graphStatus CheckSingleParaNumHeads() const;
|
||||
ge::graphStatus CheckSingleParaKvHeadNums() const;
|
||||
ge::graphStatus CheckSingleParaLayout() const;
|
||||
ge::graphStatus CheckSingleParaSparseMode() const;
|
||||
ge::graphStatus CheckSingleParaSparseBlockSize() const;
|
||||
ge::graphStatus CheckSingleParaSparseIndices() const;
|
||||
ge::graphStatus CheckSinglePara() const;
|
||||
ge::graphStatus CheckMultiParaConsistency() const;
|
||||
ge::graphStatus CheckDequantScaleNotExistence();
|
||||
template <typename T> ge::graphStatus CheckAttrValueByMap(
|
||||
std::map<std::string, std::pair<const T *, T>> &attrMap) const;
|
||||
ge::graphStatus CheckParaExistenceMlaAntiquant() const;
|
||||
ge::graphStatus CheckParaExistenceGqaAntiquant() const;
|
||||
ge::graphStatus CheckParaExistenceMla() const;
|
||||
ge::graphStatus CheckParaExistence();
|
||||
void SetQSFAShapeCompare();
|
||||
ge::graphStatus CheckKVDType();
|
||||
ge::graphStatus CheckKVShapeForBatchContinuous();
|
||||
ge::graphStatus CheckKVShapeForPageAttention();
|
||||
ge::graphStatus CheckKVShape();
|
||||
ge::graphStatus CheckKV();
|
||||
ge::graphStatus CheckTopK();
|
||||
ge::graphStatus CheckTopkShape();
|
||||
ge::graphStatus CheckBlockTable() const;
|
||||
ge::graphStatus CheckDTypeConsistency(const ge::DataType &actualDtype,
|
||||
const ge::DataType &expectDtype, const std::string &name) const;
|
||||
|
||||
ge::graphStatus CheckAttenOut();
|
||||
ge::graphStatus CheckAttenOutShape();
|
||||
ge::graphStatus CheckActualSeqLensQ();
|
||||
ge::graphStatus CheckActualSeqLensQShape();
|
||||
ge::graphStatus CheckActualSeqLensQDType();
|
||||
ge::graphStatus CheckActualSeqLens();
|
||||
ge::graphStatus CheckActualSeqLensDType();
|
||||
ge::graphStatus CheckActualSeqLensShape();
|
||||
ge::graphStatus CheckMultiParaConsistency();
|
||||
|
||||
ge::graphStatus CheckFeatureMlaAntiquantShape() const;
|
||||
ge::graphStatus CheckFeatureMlaAntiquantShapeSizes() const;
|
||||
ge::graphStatus CheckFeatureMlaAntiquantShapeSparseAndHeadDim() const;
|
||||
ge::graphStatus CheckFeatureMlaAntiquantLayout() const;
|
||||
ge::graphStatus CheckFeatureMlaAntiquantDtype() const;
|
||||
ge::graphStatus CheckFeatureMlaAntiquantAttr() const;
|
||||
ge::graphStatus CheckFeatureMlaAntiquantPa() const;
|
||||
ge::graphStatus CheckFeatureMlaAntiquant() const;
|
||||
ge::graphStatus CheckFeatureMla() const;
|
||||
ge::graphStatus CheckFeature() const;
|
||||
|
||||
private:
|
||||
const char *opName_;
|
||||
fe::PlatFormInfos *platformInfo_;
|
||||
QSFAParaInfo opParamInfo_;
|
||||
const QSFATilingInfo &qsfaInfo_;
|
||||
|
||||
uint32_t bSize_ = 0;
|
||||
uint32_t n1Size_ = 0;
|
||||
uint32_t n2Size_ = 0;
|
||||
uint32_t gSize_ = 0;
|
||||
uint32_t s1Size_ = 0;
|
||||
int64_t s2Size_ = 0;
|
||||
uint32_t qHeadDim_ = 0;
|
||||
uint32_t kHeadDim_ = 0;
|
||||
uint32_t vHeadDim_ = 0;
|
||||
uint32_t qTSize_ = 0; // 仅TND时生效
|
||||
uint32_t kvTSize_ = 0; // 仅TND时生效
|
||||
KvStorageMode kvStorageMode_ = KvStorageMode::BATCH_CONTINUOUS;
|
||||
uint32_t sparseBlockCount_ = 0;
|
||||
int64_t sparseBlockSize_ = 0;
|
||||
int32_t attentionMode_ = 0;
|
||||
int32_t keyQuantMode_ = 0;
|
||||
int32_t valueQuantMode_ = 0;
|
||||
int32_t quantScaleRepoMode_ = 0;
|
||||
int64_t tileSize_ = 0;
|
||||
int64_t preTokens_ = 0;
|
||||
int64_t nextTokens_ = 0;
|
||||
int32_t ropeHeadDim_ = 0;
|
||||
|
||||
QSFALayout qLayout_ = QSFALayout::BSND;
|
||||
QSFALayout topkLayout_ = QSFALayout::BSND;
|
||||
QSFALayout outLayout_ = QSFALayout::BSND;
|
||||
QSFALayout kvLayout_ = QSFALayout::BSND;
|
||||
|
||||
uint32_t maxBlockNumPerBatch_ = 0;
|
||||
int64_t blockSize_ = 0;
|
||||
|
||||
uint32_t aicNum_ = 0;
|
||||
uint32_t aivNum_ = 0;
|
||||
NpuArch npuArch_ = NpuArch::DAV_2201;
|
||||
bool isA5_ = false;
|
||||
uint64_t l2CacheSize_ = 0;
|
||||
|
||||
ge::DataType inputQType_ = ge::DT_FLOAT16;
|
||||
ge::DataType inputKvType_ = ge::DT_FLOAT16;
|
||||
ge::DataType outputType_ = ge::DT_FLOAT16;
|
||||
|
||||
gert::Shape queryShapeCmp_{};
|
||||
gert::Shape keyShapeCmp_{};
|
||||
gert::Shape valueShapeCmp_{};
|
||||
gert::Shape topkShapeCmp_{};
|
||||
gert::Shape attenOutShapeCmp_{};
|
||||
};
|
||||
|
||||
class QSFAInfoParser {
|
||||
public:
|
||||
explicit QSFAInfoParser(const gert::TilingContext *context) : context_(context) {}
|
||||
~QSFAInfoParser() = default;
|
||||
|
||||
ge::graphStatus CheckRequiredInOutExistence() const;
|
||||
ge::graphStatus CheckRequiredAttrExistence() const;
|
||||
ge::graphStatus CheckRequiredParaExistence() const;
|
||||
|
||||
ge::graphStatus GetActualSeqLenQSize(uint32_t &size);
|
||||
ge::graphStatus GetNpuInfo();
|
||||
ge::graphStatus GetOpName();
|
||||
void GetOptionalInputParaInfo();
|
||||
void GetInputParaInfo();
|
||||
void GetOutputParaInfo();
|
||||
ge::graphStatus GetAttrParaInfo();
|
||||
ge::graphStatus GetOpParaInfo();
|
||||
ge::graphStatus GetKvCache();
|
||||
|
||||
ge::graphStatus GetInOutDataType();
|
||||
ge::graphStatus GetQTSize();
|
||||
ge::graphStatus GetBatchSize();
|
||||
ge::graphStatus GetKVTSize();
|
||||
ge::graphStatus GetQHeadDim();
|
||||
ge::graphStatus GetKHeadDim();
|
||||
ge::graphStatus GetS1Size();
|
||||
ge::graphStatus GetKvStorageMode();
|
||||
ge::graphStatus GetKvLayout();
|
||||
void SetQSFAShape();
|
||||
ge::graphStatus GetS2SizeForBatchContinuous();
|
||||
ge::graphStatus GetMaxBlockNumPerBatch();
|
||||
ge::graphStatus GetBlockSize();
|
||||
ge::graphStatus GetS2SizeForPageAttention();
|
||||
ge::graphStatus GetS2Size();
|
||||
ge::graphStatus GetValueHeadDim();
|
||||
ge::graphStatus GetDSizeKV();
|
||||
ge::graphStatus GetRopeHeadDim();
|
||||
ge::graphStatus GetQueryAndOutLayout();
|
||||
ge::graphStatus GetTopkLayout();
|
||||
ge::graphStatus GetN1Size();
|
||||
ge::graphStatus GetN2Size();
|
||||
ge::graphStatus GetGSize();
|
||||
ge::graphStatus GetSparseBlockCount();
|
||||
ge::graphStatus GetActualseqInfo();
|
||||
ge::graphStatus GetShapeAndSizeInfo();
|
||||
void GenerateInfo(QSFATilingInfo &qsfaInfo);
|
||||
void FillTilingInfoAttrsAndLayouts(QSFATilingInfo &qsfaInfo);
|
||||
ge::graphStatus Parse(QSFATilingInfo &qsfaInfo);
|
||||
|
||||
const gert::TilingContext *context_ = nullptr;
|
||||
|
||||
const char *opName_;
|
||||
fe::PlatFormInfos *platformInfo_;
|
||||
QSFAParaInfo opParamInfo_;
|
||||
|
||||
uint32_t bSize_ = 0;
|
||||
uint32_t n1Size_ = 0;
|
||||
uint32_t n2Size_ = 0;
|
||||
uint32_t gSize_ = 0;
|
||||
uint32_t s1Size_ = 0;
|
||||
int64_t s2Size_ = 0;
|
||||
uint32_t qHeadDim_ = 0;
|
||||
uint32_t kHeadDim_ = 0;
|
||||
uint32_t vHeadDim_ = 0;
|
||||
int32_t ropeHeadDim_ = 0;
|
||||
int64_t dSizeKV_ = 0;
|
||||
uint32_t qTSize_ = 0; // 仅TND时生效
|
||||
uint32_t kvTSize_ = 0; // 仅TND时生效
|
||||
KvStorageMode kvStorageMode_ = KvStorageMode::BATCH_CONTINUOUS;
|
||||
uint32_t sparseBlockCount_ = 0;
|
||||
|
||||
QSFALayout qLayout_ = QSFALayout::BSND;
|
||||
QSFALayout topkLayout_ = QSFALayout::BSND;
|
||||
QSFALayout outLayout_ = QSFALayout::BSND;
|
||||
QSFALayout kvLayout_ = QSFALayout::BSND;
|
||||
|
||||
uint32_t maxBlockNumPerBatch_ = 0;
|
||||
uint32_t blockSize_ = 0;
|
||||
|
||||
NpuArch npuArch_ = NpuArch::DAV_2201;
|
||||
bool isA5_ = false;
|
||||
|
||||
ge::DataType inputQType_ = ge::DT_FLOAT16;
|
||||
ge::DataType inputKvType_ = ge::DT_FLOAT16;
|
||||
ge::DataType outputType_ = ge::DT_FLOAT16;
|
||||
|
||||
uint64_t l2CacheSize_ = 0;
|
||||
|
||||
bool isSameSeqAllKVTensor_ = true;
|
||||
bool isSameActualseq_ = true;
|
||||
uint32_t maxActualseq_ = 0;
|
||||
|
||||
uint32_t actualLenDimsQ_ = 0;
|
||||
uint32_t actualLenDimsKV_ = 0;
|
||||
|
||||
gert::Shape queryShape_{};
|
||||
gert::Shape keyShape_{};
|
||||
gert::Shape valueShape_{};
|
||||
gert::Shape sparseIndicesShape_{};
|
||||
};
|
||||
} // namespace optiling
|
||||
#endif // KV_QUANT_SPARSE_FLASH_ATTENTION_TILING_H
|
||||
@@ -0,0 +1,130 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_flash_attention_common_arch35.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef KV_QUANT_SPARSE_FLASH_ATTENTION_COMMON_ARCH35_H
|
||||
#define KV_QUANT_SPARSE_FLASH_ATTENTION_COMMON_ARCH35_H
|
||||
#include <type_traits>
|
||||
#include "kernel_tiling/kernel_tiling.h"
|
||||
|
||||
#if __has_include("../../sparse_flash_attention/arch35/common/util_regbase.h")
|
||||
#include "../../sparse_flash_attention/arch35/common/util_regbase.h"
|
||||
#else
|
||||
#include "../../../sparse_flash_attention/op_kernel/arch35/common/util_regbase.h"
|
||||
#endif
|
||||
|
||||
#if __has_include("../../common/op_kernel/buffer.h")
|
||||
#include "../../common/op_kernel/buffer.h"
|
||||
#else
|
||||
#include "../../common/buffer.h"
|
||||
#endif
|
||||
#if __has_include("../../common/op_kernel/buffer_manager.h")
|
||||
#include "../../common/op_kernel/buffer_manager.h"
|
||||
#else
|
||||
#include "../../common/buffer_manager.h"
|
||||
#endif
|
||||
#if __has_include("../../common/op_kernel/buffers_policy.h")
|
||||
#include "../../common/op_kernel/buffers_policy.h"
|
||||
#else
|
||||
#include "../../common/buffers_policy.h"
|
||||
#endif
|
||||
|
||||
constexpr uint64_t BLOCK_BYTE = 32;
|
||||
constexpr uint32_t NEGATIVE_MIN_VALUE_FP32 = 0xFF7FFFFF;
|
||||
|
||||
constexpr uint32_t BUFFER_SIZE_16K = 16384; // 16384表示16 * 1024
|
||||
constexpr uint32_t BUFFER_SIZE_32K = 32768; // 32768表示32 * 1024
|
||||
constexpr uint32_t BUFFER_SIZE_128K = 131072; // 131072表示128 * 1024
|
||||
|
||||
constexpr uint32_t L0AB_SHARED_SIZE_64K = 65536; // 65536表示64*1024
|
||||
constexpr uint32_t L0C_SHARED_SIZE_256K = 262144; // 262144表示256 * 1024
|
||||
|
||||
constexpr uint32_t CV_RATIO = 2;
|
||||
constexpr uint64_t SYNC_MODE = 4;
|
||||
|
||||
static constexpr uint32_t QSFA_SYNC_MODE0 = 0;
|
||||
|
||||
enum class QSFA_LAYOUT {
|
||||
BSND = 0,
|
||||
TND = 1,
|
||||
PA_BSND = 2,
|
||||
};
|
||||
|
||||
enum class QSFATemplateMode {
|
||||
SWA_TEMPLATE_MODE = 0,
|
||||
CFA_TEMPLATE_MODE = 1,
|
||||
SCFA_TEMPLATE_MODE = 2
|
||||
};
|
||||
|
||||
namespace BaseApi {
|
||||
__aicore__ constexpr uint64_t Align2Func(uint64_t data) {
|
||||
return (data + 1UL) >> 1UL << 1UL; // 向上2对齐, +1移位2
|
||||
}
|
||||
|
||||
__aicore__ constexpr uint64_t Align8Func(uint64_t data) {
|
||||
return (data + 7UL) >> 3UL << 3UL; // 向上8对齐, +7移位3
|
||||
}
|
||||
|
||||
__aicore__ constexpr uint64_t Align16Func(uint64_t data) {
|
||||
return (data + 15UL) >> 4UL << 4UL; // 向上16对齐, +15移位4
|
||||
}
|
||||
|
||||
__aicore__ constexpr uint64_t Align64Func(uint64_t data) {
|
||||
return (data + 63UL) >> 6UL << 6UL; // 向上64对齐, +63移位6
|
||||
}
|
||||
}
|
||||
|
||||
#define TEMPLATE_INTF \
|
||||
template <typename Q_T, typename KV_T, typename T, typename OUTPUT_T, bool isFd, bool isPa, QSFA_LAYOUT LAYOUT_T, \
|
||||
QSFA_LAYOUT KV_LAYOUT_T, QSFATemplateMode TEMPLATE_MODE, bool IS_SPLIT_G>
|
||||
|
||||
#define TEMPLATE_INTF_ARGS \
|
||||
Q_T, KV_T, T, OUTPUT_T, isFd, isPa, LAYOUT_T, KV_LAYOUT_T, TEMPLATE_MODE, IS_SPLIT_G
|
||||
|
||||
#define QSFA_CUBE_BLOCK_TRAITS_TYPE_FIELDS(X) \
|
||||
X(Q_T) \
|
||||
X(KV_T) \
|
||||
X(T) \
|
||||
X(OUTPUT_T) \
|
||||
|
||||
#define QSFA_CUBE_BLOCK_TRAITS_CONST_FIELDS(X) \
|
||||
X(isFd, bool, false) \
|
||||
X(isPa, bool, true) \
|
||||
X(LAYOUT_T, QSFA_LAYOUT, QSFA_LAYOUT::BSND) \
|
||||
X(KV_LAYOUT_T, QSFA_LAYOUT, QSFA_LAYOUT::PA_BSND) \
|
||||
X(TEMPLATE_MODE, QSFATemplateMode, QSFATemplateMode::SCFA_TEMPLATE_MODE) \
|
||||
X(IS_SPLIT_G, bool, false)
|
||||
|
||||
|
||||
/* 1. 生成带默认值的模版Template */
|
||||
#define GEN_TYPE_PARAM(name) typename name,
|
||||
#define GEN_CONST_PARAM(name, type, default_val) type name = default_val,
|
||||
|
||||
#define TEMPLATES_DEF \
|
||||
template <QSFA_CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_TYPE_PARAM) \
|
||||
QSFA_CUBE_BLOCK_TRAITS_CONST_FIELDS(GEN_CONST_PARAM) bool end = true>
|
||||
|
||||
/* 2. 生成不带默认值的模版Template */
|
||||
#define GEN_TEMPLATE_TYPE_NODEF(name) typename name,
|
||||
#define GEN_TEMPLATE_CONST_NODEF(name, type, default_val) type name,
|
||||
#define TEMPLATES_DEF_NO_DEFAULT \
|
||||
template <QSFA_CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_TEMPLATE_TYPE_NODEF) \
|
||||
QSFA_CUBE_BLOCK_TRAITS_CONST_FIELDS(GEN_TEMPLATE_CONST_NODEF) bool end>
|
||||
|
||||
/* 3. 生成有默认值的Args */
|
||||
#define GEN_ARG_NAME(name, ...) name,
|
||||
#define TEMPLATE_ARGS \
|
||||
QSFA_CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_ARG_NAME) \
|
||||
QSFA_CUBE_BLOCK_TRAITS_CONST_FIELDS(GEN_ARG_NAME) end
|
||||
|
||||
#endif //KV_QUANT_SPARSE_FLASH_ATTENTION_COMMON_ARCH35_H
|
||||
@@ -0,0 +1,707 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_flash_attention_kernel_mla.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef KV_QUANT_SPARSE_FLASH_ATTENTION_KERNEL_MLA_H
|
||||
#define KV_QUANT_SPARSE_FLASH_ATTENTION_KERNEL_MLA_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "kernel_operator_list_tensor_intf.h"
|
||||
#include "kernel_tiling/kernel_tiling.h"
|
||||
#include "lib/matmul_intf.h"
|
||||
#include "lib/matrix/matmul/tiling.h"
|
||||
#include "kv_quant_sparse_flash_attention_service_cube_mla.h"
|
||||
#include "kv_quant_sparse_flash_attention_service_vector_mla.h"
|
||||
#include "kv_quant_sparse_flash_attention_common_arch35.h"
|
||||
#include "kv_quant_sparse_flash_attention_kvcache.h"
|
||||
#if __has_include("../../common/op_kernel/CopyInL1.h")
|
||||
#include "../../common/op_kernel/CopyInL1.h"
|
||||
#else
|
||||
#include "../common/CopyInL1.h"
|
||||
#endif
|
||||
#if __has_include("../../common/op_kernel/matmul.h")
|
||||
#include "../../common/op_kernel/matmul.h"
|
||||
#else
|
||||
#include "../common/matmul.h"
|
||||
#endif
|
||||
#if __has_include("../../common/op_kernel/FixpipeOut.h")
|
||||
#include "../../common/op_kernel/FixpipeOut.h"
|
||||
#else
|
||||
#include "../common/FixpipeOut.h"
|
||||
#endif
|
||||
|
||||
using matmul::MatmulType;
|
||||
using namespace AscendC;
|
||||
using namespace AscendC::Impl::Detail;
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace BaseApi {
|
||||
template <typename CubeBlockType, typename VecBlockType> class KvQuantSparseFlashAttentionMla {
|
||||
public:
|
||||
ARGS_TRAITS;
|
||||
|
||||
__aicore__ inline KvQuantSparseFlashAttentionMla(){};
|
||||
__aicore__ inline void Init(__gm__ uint8_t *query, __gm__ uint8_t *key, __gm__ uint8_t *value,
|
||||
__gm__ uint8_t *sparseIndices, __gm__ uint8_t* keyScale,
|
||||
__gm__ uint8_t* valueScale, __gm__ uint8_t *blockTable,
|
||||
__gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengths,
|
||||
__gm__ uint8_t *attentionOut, __gm__ uint8_t *workspace,
|
||||
const KvQuantSparseFlashAttentionTilingDataMla *__restrict tiling,
|
||||
TPipe *tPipe);
|
||||
__aicore__ inline void Process();
|
||||
|
||||
private:
|
||||
__aicore__ inline void ProcessMainLoop();
|
||||
__aicore__ inline void InitGlobalBuffer(__gm__ uint8_t *query, __gm__ uint8_t *key, __gm__ uint8_t *value,
|
||||
__gm__ uint8_t *sparseIndices, __gm__ uint8_t *blockTable, __gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengths,
|
||||
__gm__ uint8_t *workspace, const KvQuantSparseFlashAttentionTilingDataMla *__restrict tiling, TPipe *tPipe);
|
||||
__aicore__ inline void InitLocalBuffer();
|
||||
__aicore__ inline void ComputeConstexpr();
|
||||
__aicore__ inline void InitMMResBuf(__gm__ uint8_t *workspace);
|
||||
__aicore__ inline void SetRunInfo(RunInfo &runInfo, RunParamStr &runParam, int64_t taskId, int64_t s2LoopCount,
|
||||
int64_t s2LoopLimit, int64_t multiCoreInnerIdx);
|
||||
__aicore__ inline void ComputeBmm1Tail(RunInfo &runInfo, RunParamStr &runParam);
|
||||
__aicore__ inline void InitUniqueConstInfo();
|
||||
__aicore__ inline void InitUniqueRunInfo(const RunParamStr &runParam, RunInfo &runInfo);
|
||||
__aicore__ inline void ComputeAxisIdxByBnAndGs1(int64_t bnIndex, int64_t gS1Index, RunParamStr &runParam);
|
||||
__aicore__ inline void InitCalcParamsEach();
|
||||
__aicore__ inline uint64_t GetBalanceActualSeqLengths(GlobalTensor<int32_t> &actualSeqLengths, uint32_t bIdx);
|
||||
__aicore__ inline void GetAxisStartIdx(uint32_t bN2EndPrev, uint32_t s1GEndPrev, uint32_t s2EndPrev);
|
||||
|
||||
TPipe *pipe;
|
||||
|
||||
const KvQuantSparseFlashAttentionTilingDataMla *__restrict tilingData;
|
||||
static constexpr uint64_t SYNC_MODE = 4;
|
||||
static constexpr uint32_t PRELOAD_NUM = 2;
|
||||
/* 核间通道 */
|
||||
BufferManager<BufferType::UB> ubBufferManager;
|
||||
BuffersPolicyDB<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> bmm1Buffers;
|
||||
BuffersPolicySingleBuffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> bmm2Buffers;
|
||||
BufferManager<BufferType::GM> gmBufferManager;
|
||||
|
||||
// mm2左矩阵P
|
||||
BufferManager<BufferType::L1> l1BufferManager;
|
||||
BuffersPolicy3buff<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> l1RightBuffers;
|
||||
CVSharedParams sharedParams;
|
||||
/* GM信息 */
|
||||
__gm__ int32_t *actualSeqKvlenAddr = nullptr;
|
||||
__gm__ int32_t *actualSeqQlenAddr = nullptr;
|
||||
|
||||
GlobalTensor<int32_t> actualSeqLengthsQGm;
|
||||
uint32_t usedCoreNum = 0U;
|
||||
|
||||
/* workspace 空间 */
|
||||
BuffersPolicy3buff<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> v0ResGmBuffers;
|
||||
|
||||
/* 核Index信息 */
|
||||
int32_t aicIdx;
|
||||
|
||||
/* 切G时最大s2Loop */
|
||||
int64_t maxS2LoopCnt;
|
||||
|
||||
/* 初始化后不变的信息 */
|
||||
ConstInfo constInfo;
|
||||
|
||||
/* 模板库Block */
|
||||
CubeBlockType cubeBlock;
|
||||
VecBlockType vecBlock;
|
||||
|
||||
uint32_t crossCoreSyncBufId = 0;
|
||||
};
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::Init(
|
||||
__gm__ uint8_t *query,
|
||||
__gm__ uint8_t *key, __gm__ uint8_t *value,
|
||||
__gm__ uint8_t *sparseIndices, __gm__ uint8_t* keyScale,
|
||||
__gm__ uint8_t* valueScale, __gm__ uint8_t *blockTable, __gm__ uint8_t *actualSeqLengthsQ,
|
||||
__gm__ uint8_t *actualSeqLengths, __gm__ uint8_t *attentionOut, __gm__ uint8_t *workspace,
|
||||
const KvQuantSparseFlashAttentionTilingDataMla *__restrict tiling,
|
||||
TPipe *tPipe)
|
||||
{
|
||||
fa_base_matmul::idCounterNum = 0;
|
||||
constInfo.subBlockIdx = GetSubBlockIdx();
|
||||
if ASCEND_IS_AIC {
|
||||
this->aicIdx = GetBlockIdx();
|
||||
constInfo.aivIdx = 0;
|
||||
} else {
|
||||
constInfo.aivIdx = GetBlockIdx();
|
||||
this->aicIdx = constInfo.aivIdx >> 1;
|
||||
this->tilingData = tiling;
|
||||
}
|
||||
|
||||
constInfo.s1BaseSize = 64;
|
||||
constInfo.s2BaseSize = 128;
|
||||
|
||||
this->pipe = tPipe;
|
||||
vecBlock.InitVecBlock(tPipe, this->tilingData, this->sharedParams, this->aicIdx, constInfo.subBlockIdx, actualSeqLengthsQ, actualSeqLengths);
|
||||
if ASCEND_IS_AIV {
|
||||
constInfo.bSize = this->sharedParams.bSize;
|
||||
constInfo.gSize = this->sharedParams.gSize;
|
||||
constInfo.s1Size = this->sharedParams.s1Size;
|
||||
constInfo.needInit = this->sharedParams.needInit;
|
||||
constInfo.dSizeV = 512;
|
||||
}
|
||||
vecBlock.CleanOutput(attentionOut, constInfo);
|
||||
/* cube侧不依赖sharedParams的scalar前置 */
|
||||
InitMMResBuf(workspace);
|
||||
if ASCEND_IS_AIC {
|
||||
cubeBlock.InitCubeBlock(pipe, &l1BufferManager, query);
|
||||
/* wait kfc message */
|
||||
CrossCoreWaitFlag<SYNC_MODE, PIPE_S>(15);
|
||||
auto tempTilingSSbuf = reinterpret_cast<__ssbuf__ uint32_t*>(0); // 从ssbuf的0地址开始拷贝
|
||||
auto tempTiling = reinterpret_cast<uint32_t *>(&sharedParams);
|
||||
#pragma unroll
|
||||
for (int i = 0; i < sizeof(CVSharedParams) / sizeof(uint32_t); ++i, ++tempTilingSSbuf, ++tempTiling) {
|
||||
*tempTiling = *tempTilingSSbuf;
|
||||
}
|
||||
}
|
||||
this->ComputeConstexpr();
|
||||
this->InitGlobalBuffer(query, key, value, sparseIndices, blockTable, actualSeqLengthsQ, actualSeqLengths,
|
||||
workspace, tiling, tPipe); // gm设置
|
||||
this->InitCalcParamsEach();
|
||||
this->InitLocalBuffer();
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::InitCalcParamsEach()
|
||||
{
|
||||
// 计算总的基本块
|
||||
maxS2LoopCnt = 0; // 所有核中最大累计s2Loop
|
||||
uint32_t qsfaTotalBaseNum = 0;
|
||||
uint32_t actBatchS2 = 1;
|
||||
uint32_t coreNum = GetBlockNum(); // G128时相邻两个cube核处理一个s1,coreNum减半
|
||||
uint32_t currCoreIdx = aicIdx;
|
||||
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
currCoreIdx = currCoreIdx >> 1;
|
||||
coreNum = coreNum >> 1;
|
||||
}
|
||||
|
||||
uint32_t actBatchS1 = 1;
|
||||
for (uint32_t bIdx = 0; bIdx < constInfo.bSize; bIdx++) {
|
||||
uint32_t actBatchS1 = GetBalanceActualSeqLengths(actualSeqLengthsQGm, bIdx); //不切S2,只关注S1
|
||||
qsfaTotalBaseNum += actBatchS1 * actBatchS2;
|
||||
}
|
||||
|
||||
uint32_t avgBaseNum = 1;
|
||||
if (qsfaTotalBaseNum > coreNum) {
|
||||
avgBaseNum = (qsfaTotalBaseNum + coreNum - 1) / coreNum;
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
usedCoreNum = ((qsfaTotalBaseNum + avgBaseNum - 1) / avgBaseNum) << 1;
|
||||
}
|
||||
} else {
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
usedCoreNum = qsfaTotalBaseNum << 1;
|
||||
} else {
|
||||
usedCoreNum = qsfaTotalBaseNum;
|
||||
}
|
||||
}
|
||||
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
maxS2LoopCnt = avgBaseNum * (Min(constInfo.sparseBlockCount, constInfo.s2Size) +
|
||||
constInfo.s2BaseSize - 1) / constInfo.s2BaseSize;
|
||||
}
|
||||
|
||||
if (aicIdx >= usedCoreNum) {
|
||||
return;
|
||||
}
|
||||
// 计算当前核的基本块
|
||||
uint32_t qsfaAccumBaseNum = 0; // qsfa当前累积的基本块数
|
||||
uint32_t targetBaseNum = 0;
|
||||
uint32_t qsfaLastValidBIdx = 0;
|
||||
uint32_t lastValidactBatchS1 = 0;
|
||||
bool setStart = false;
|
||||
targetBaseNum = (currCoreIdx + 1) * avgBaseNum; // 计算当前的目标权重
|
||||
uint32_t targetStartBaseNum = targetBaseNum - avgBaseNum;
|
||||
for (uint32_t bN2Idx = 0; bN2Idx < constInfo.bSize * constInfo.n2Size; bN2Idx++) {
|
||||
uint32_t bIdx = bN2Idx / constInfo.n2Size;
|
||||
actBatchS1 = GetBalanceActualSeqLengths(actualSeqLengthsQGm, bIdx);
|
||||
for (uint32_t s1GIdx = 0; s1GIdx < actBatchS1; s1GIdx++) {
|
||||
qsfaAccumBaseNum += 1;
|
||||
if (!setStart && qsfaAccumBaseNum >= targetStartBaseNum) {
|
||||
constInfo.bN2Start = bN2Idx;
|
||||
constInfo.gS1Start = s1GIdx;
|
||||
setStart = true;
|
||||
}
|
||||
if (qsfaAccumBaseNum >= targetBaseNum) {
|
||||
// 更新当前核的End分核信息
|
||||
constInfo.s2End = 0;
|
||||
constInfo.bN2End = bN2Idx;
|
||||
constInfo.gS1End = s1GIdx;
|
||||
|
||||
if (currCoreIdx != 0) {
|
||||
GetAxisStartIdx(constInfo.bN2Start, constInfo.gS1Start, 0);
|
||||
}
|
||||
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
if ((actBatchS1 > 0) && (actBatchS2 > 0)) {
|
||||
qsfaLastValidBIdx = bIdx;
|
||||
lastValidactBatchS1 = actBatchS1;
|
||||
}
|
||||
}
|
||||
if (!setStart) {
|
||||
constInfo.bN2Start = qsfaLastValidBIdx;
|
||||
constInfo.gS1Start = lastValidactBatchS1 - 1;
|
||||
}
|
||||
if (qsfaAccumBaseNum < targetBaseNum) {
|
||||
// 更新最后一个核的End分核信息
|
||||
constInfo.bN2End = qsfaLastValidBIdx;
|
||||
constInfo.gS1End = lastValidactBatchS1 - 1;
|
||||
constInfo.s2End = 0;
|
||||
if (currCoreIdx != 0) {
|
||||
GetAxisStartIdx(constInfo.bN2Start, constInfo.gS1Start, 0);
|
||||
}
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline uint64_t KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::\
|
||||
GetBalanceActualSeqLengths(GlobalTensor<int32_t> &actualSeqLengths, uint32_t bIdx)
|
||||
{
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
if (bIdx == 0) {
|
||||
return actualSeqQlenAddr[0];
|
||||
} else if (bIdx > 0) {
|
||||
return actualSeqQlenAddr[bIdx] - actualSeqQlenAddr[bIdx - 1];
|
||||
} else {
|
||||
return 0;
|
||||
}
|
||||
} else {
|
||||
if (constInfo.isActualLenDimsNull == 0) {
|
||||
return actualSeqQlenAddr[bIdx];
|
||||
} else {
|
||||
return constInfo.s1Size;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::GetAxisStartIdx(uint32_t bN2EndPrev,
|
||||
uint32_t s1GEndPrev,
|
||||
uint32_t s2EndPrev)
|
||||
{
|
||||
uint32_t qsfaBEndPrev = bN2EndPrev / constInfo.n2Size;
|
||||
uint32_t actualSeqQPrev = GetBalanceActualSeqLengths(actualSeqLengthsQGm, qsfaBEndPrev);
|
||||
uint32_t s1GPrevBaseNum = actualSeqQPrev;
|
||||
constInfo.bN2Start = bN2EndPrev;
|
||||
constInfo.gS1Start = s1GEndPrev;
|
||||
constInfo.s2Start = 0;
|
||||
if (s1GEndPrev >= s1GPrevBaseNum - 1) { // 上个核把S1G处理完了
|
||||
constInfo.bN2Start++;
|
||||
constInfo.gS1Start = 0;
|
||||
} else {
|
||||
constInfo.gS1Start++;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::InitGlobalBuffer(
|
||||
__gm__ uint8_t *query, __gm__ uint8_t *key, __gm__ uint8_t *value, __gm__ uint8_t *sparseIndices,
|
||||
__gm__ uint8_t *blockTable, __gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengths,
|
||||
__gm__ uint8_t *workspace, const KvQuantSparseFlashAttentionTilingDataMla *__restrict tiling, TPipe *tPipe)
|
||||
{
|
||||
if (actualSeqLengthsQ != nullptr) {
|
||||
actualSeqQlenAddr = (__gm__ int32_t *)actualSeqLengthsQ;
|
||||
}
|
||||
|
||||
if (actualSeqLengths != nullptr) {
|
||||
actualSeqKvlenAddr = (__gm__ int32_t *)actualSeqLengths;
|
||||
}
|
||||
|
||||
vecBlock.InitGlobalBuffer(key, value, sparseIndices, blockTable);
|
||||
cubeBlock.InitCubeInput(actualSeqLengthsQ, constInfo);
|
||||
}
|
||||
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::InitMMResBuf(
|
||||
__gm__ uint8_t *workspace)
|
||||
{
|
||||
uint32_t mm1RightSize = constInfo.s2BaseSize * 576 * sizeof(Q_T);
|
||||
l1BufferManager.Init(pipe, 524288); // 512 * 1024
|
||||
l1RightBuffers.Init(l1BufferManager, mm1RightSize);
|
||||
l1RightBuffers.Get().SetCrossCoreID(crossCoreSyncBufId, INVALID_CROSS_CORE_EVENT_ID);
|
||||
crossCoreSyncBufId++;
|
||||
l1RightBuffers.Get().SetCrossCoreID(crossCoreSyncBufId, INVALID_CROSS_CORE_EVENT_ID);
|
||||
crossCoreSyncBufId++;
|
||||
l1RightBuffers.Get().SetCrossCoreID(crossCoreSyncBufId, INVALID_CROSS_CORE_EVENT_ID);
|
||||
crossCoreSyncBufId++;
|
||||
|
||||
if ASCEND_IS_AIC {
|
||||
l1RightBuffers.Get().SetCrossCore();
|
||||
l1RightBuffers.Get().SetCrossCore();
|
||||
l1RightBuffers.Get().SetCrossCore();
|
||||
}
|
||||
uint32_t mm1ResultSize = constInfo.s1BaseSize / CV_RATIO * constInfo.s2BaseSize * sizeof(T);
|
||||
uint32_t mm2ResultSize = constInfo.s1BaseSize / CV_RATIO * 512 * sizeof(T);
|
||||
ubBufferManager.Init(pipe, mm1ResultSize * 2 + mm2ResultSize);
|
||||
|
||||
bmm1Buffers.Init(ubBufferManager, mm1ResultSize);
|
||||
bmm1Buffers.Get().SetCrossCoreID(crossCoreSyncBufId, crossCoreSyncBufId);
|
||||
crossCoreSyncBufId++;
|
||||
bmm1Buffers.Get().SetCrossCoreID(crossCoreSyncBufId, crossCoreSyncBufId);
|
||||
crossCoreSyncBufId++;
|
||||
if ASCEND_IS_AIV {
|
||||
bmm1Buffers.Get().SetCrossCore();
|
||||
bmm1Buffers.Get().SetCrossCore();
|
||||
}
|
||||
|
||||
bmm2Buffers.Init(ubBufferManager, mm2ResultSize);
|
||||
bmm2Buffers.Get().SetCrossCoreID(crossCoreSyncBufId, crossCoreSyncBufId);
|
||||
crossCoreSyncBufId++;
|
||||
|
||||
if ASCEND_IS_AIV {
|
||||
bmm2Buffers.Get().SetCrossCore();
|
||||
}
|
||||
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
uint32_t v0ResSize = constInfo.s2BaseSize * 576U * sizeof(Q_T);
|
||||
int64_t totalOffset = v0ResSize * 3 * (aicIdx >> 1U);
|
||||
gmBufferManager.Init(workspace + totalOffset);
|
||||
v0ResGmBuffers.Init(gmBufferManager, v0ResSize);
|
||||
v0ResGmBuffers.Get().SetCrossCoreID(INVALID_CROSS_CORE_EVENT_ID, crossCoreSyncBufId);
|
||||
crossCoreSyncBufId++;
|
||||
v0ResGmBuffers.Get().SetCrossCoreID(INVALID_CROSS_CORE_EVENT_ID, crossCoreSyncBufId);
|
||||
crossCoreSyncBufId++;
|
||||
v0ResGmBuffers.Get().SetCrossCoreID(INVALID_CROSS_CORE_EVENT_ID, crossCoreSyncBufId);
|
||||
crossCoreSyncBufId++;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::InitLocalBuffer()
|
||||
{
|
||||
vecBlock.InitLocalBuffer(pipe, constInfo);
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::ComputeConstexpr()
|
||||
{
|
||||
// 计算轴的乘积
|
||||
usedCoreNum = sharedParams.usedCoreNum;
|
||||
|
||||
if ASCEND_IS_AIC {
|
||||
constInfo.bSize = this->sharedParams.bSize;
|
||||
constInfo.gSize = this->sharedParams.gSize;
|
||||
constInfo.s1Size = this->sharedParams.s1Size;
|
||||
constInfo.needInit = this->sharedParams.needInit;
|
||||
constInfo.dSizeV = 512;
|
||||
}
|
||||
constInfo.n2Size = sharedParams.n2Size;
|
||||
constInfo.s2Size = sharedParams.s2Size;
|
||||
constInfo.dSize = sharedParams.dSize;
|
||||
constInfo.dSizeVInput = sharedParams.dSizeVInput;
|
||||
constInfo.dSizeRope = sharedParams.dSizeRope;
|
||||
constInfo.dSizeNope = constInfo.dSize - constInfo.dSizeRope;
|
||||
constInfo.tileSize = sharedParams.tileSize;
|
||||
constInfo.sparseBlockCount = sharedParams.sparseBlockCount;
|
||||
constInfo.sparseBlockSize = 1;
|
||||
|
||||
constInfo.sparseMode = sharedParams.maskMode;
|
||||
constInfo.n2G = constInfo.n2Size * constInfo.gSize;
|
||||
|
||||
constInfo.s1Dv = constInfo.s1Size * constInfo.dSizeV;
|
||||
constInfo.s2Dv = constInfo.s2Size * constInfo.dSizeV;
|
||||
constInfo.n2Dv = constInfo.n2Size * constInfo.dSizeV;
|
||||
|
||||
constInfo.gDv = constInfo.gSize * constInfo.dSizeV;
|
||||
constInfo.n2S2Dv = constInfo.n2Size * constInfo.s2Dv;
|
||||
constInfo.n2GDv = constInfo.n2Size * constInfo.gDv;
|
||||
constInfo.s2BaseN2Dv = constInfo.s2BaseSize * constInfo.n2Dv;
|
||||
constInfo.layoutType = sharedParams.layoutType;
|
||||
|
||||
constInfo.isActualLenDimsNull = sharedParams.isActualSeqLengthsNull;
|
||||
constInfo.isActualLenDimsKVNull = sharedParams.isActualSeqLengthsKVNull;
|
||||
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
// (BS)ND
|
||||
constInfo.s1BaseN2GDv = constInfo.s1BaseSize * constInfo.n2GDv;
|
||||
constInfo.mm1Ka = constInfo.n2Size * constInfo.dSize;
|
||||
if ASCEND_IS_AIV {
|
||||
constInfo.attentionOutStride = \
|
||||
(constInfo.n2G - constInfo.gSize) * constInfo.dSizeV * sizeof(OUTPUT_T);
|
||||
}
|
||||
} else if constexpr (LAYOUT_T == QSFA_LAYOUT::BSND) {
|
||||
// BSH/BSNGD
|
||||
constInfo.s1BaseN2GDv = constInfo.s1BaseSize * constInfo.n2GDv;
|
||||
constInfo.mm1Ka = constInfo.n2Size * constInfo.dSize;
|
||||
|
||||
if ASCEND_IS_AIV {
|
||||
constInfo.attentionOutStride = \
|
||||
(constInfo.n2G - constInfo.gSize) * constInfo.dSizeV * sizeof(OUTPUT_T);
|
||||
}
|
||||
}
|
||||
|
||||
if ASCEND_IS_AIV {
|
||||
constInfo.blockSize = sharedParams.blockSize;
|
||||
constInfo.softmaxScale = sharedParams.softmaxScale;
|
||||
constInfo.maxBlockNumPerBatch = sharedParams.maxBlockNumPerBatch;
|
||||
}
|
||||
|
||||
InitUniqueConstInfo();
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::InitUniqueConstInfo()
|
||||
{
|
||||
// bsize + 1-> bsize
|
||||
this->constInfo.actualSeqLenSize = this->sharedParams.bSize;
|
||||
this->constInfo.actualSeqLenKVSize = this->sharedParams.bSize;
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::Process()
|
||||
{
|
||||
// SyncAll Cube和Vector都需要调用
|
||||
if (this->sharedParams.needInit) {
|
||||
SyncAll<false>();
|
||||
}
|
||||
|
||||
ProcessMainLoop();
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::ProcessMainLoop()
|
||||
{
|
||||
bool hasLoad = aicIdx < usedCoreNum;
|
||||
if (!hasLoad) {
|
||||
if ASCEND_IS_AIV {
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
for (int64_t loopCnt = 0; loopCnt < maxS2LoopCnt; loopCnt++) {
|
||||
CrossCoreSetFlag<QSFA_SYNC_MODE0, PIPE_MTE3>(15);
|
||||
CrossCoreWaitFlag<QSFA_SYNC_MODE0, PIPE_MTE3>(15);
|
||||
}
|
||||
}
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
// 适配分核左闭右开
|
||||
uint32_t bIdx = constInfo.bN2End / constInfo.n2Size;
|
||||
uint32_t qsfaActS1Size = GetBalanceActualSeqLengths(actualSeqLengthsQGm, bIdx);
|
||||
uint32_t gS1max = qsfaActS1Size;
|
||||
if (constInfo.gS1End + 1 < gS1max) {
|
||||
/* constInfo.gS1End != gS1max时,gS1End需要往后加一格, bN2End不变 */
|
||||
constInfo.gS1End = constInfo.gS1End + 1;
|
||||
} else {
|
||||
/* constInfo.gS1End == gS1max,bN2End需要往后加一格,bN2End变为0,以代表末尾 */
|
||||
constInfo.bN2End = constInfo.bN2End + 1;
|
||||
constInfo.gS1End = 0;
|
||||
}
|
||||
|
||||
// 分核信息
|
||||
uint32_t qsfaBN2StartIdx = constInfo.bN2Start;
|
||||
uint32_t bN2EndIdx = constInfo.bN2End;
|
||||
uint32_t gS1StartIdx = constInfo.gS1Start;
|
||||
uint32_t nextGs1Idx = constInfo.gS1End;
|
||||
uint32_t s2StartIdx = 0;
|
||||
uint32_t s2EndIdx = 0;
|
||||
|
||||
uint32_t s2LoopLimit = 0;
|
||||
if (nextGs1Idx != 0) {
|
||||
bN2EndIdx++;
|
||||
}
|
||||
|
||||
RunInfo runInfo[3];
|
||||
RunParamStr runParam;
|
||||
int64_t taskId = 0;
|
||||
bool notLast = true;
|
||||
int64_t multiCoreInnerIdx = 1;
|
||||
for (int64_t qsfaBnIdx = qsfaBN2StartIdx; qsfaBnIdx < bN2EndIdx; qsfaBnIdx++) {
|
||||
bool lastBN = (qsfaBnIdx == bN2EndIdx - 1);
|
||||
runParam.boIdx = qsfaBnIdx;
|
||||
runParam.n2oIdx = 0;
|
||||
ComputeParamBatch<TEMPLATE_INTF_ARGS>(runParam, this->constInfo,
|
||||
this->actualSeqQlenAddr, this->actualSeqKvlenAddr);
|
||||
ComputeS1LoopInfo<TEMPLATE_INTF_ARGS>(runParam, this->constInfo, lastBN, nextGs1Idx, gS1StartIdx);
|
||||
|
||||
int64_t gS1LoopEnd = lastBN ? (runParam.gs1LoopEndIdx + PRELOAD_NUM) : runParam.gs1LoopEndIdx;
|
||||
for (int64_t gS1Index = runParam.gs1LoopStartIdx; gS1Index < gS1LoopEnd; gS1Index++) {
|
||||
bool notLastTwoLoop = true;
|
||||
if (lastBN) {
|
||||
int32_t qsfaExtraGS1 = gS1Index - runParam.gs1LoopEndIdx;
|
||||
switch (qsfaExtraGS1) {
|
||||
case 0:
|
||||
notLastTwoLoop = false;
|
||||
break;
|
||||
case 1:
|
||||
notLastTwoLoop = false;
|
||||
notLast = false;
|
||||
break;
|
||||
default:
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
if (notLastTwoLoop) {
|
||||
this->ComputeAxisIdxByBnAndGs1(qsfaBnIdx, gS1Index, runParam);
|
||||
bool s1NoNeedCalc = ComputeParamS1<TEMPLATE_INTF_ARGS>(
|
||||
runParam, this->constInfo, gS1Index, this->actualSeqQlenAddr);
|
||||
bool s2NoNeedCalc =
|
||||
ComputeS2LoopInfo<TEMPLATE_INTF_ARGS>(runParam, this->constInfo);
|
||||
// s1和s2有任意一个不需要算, 则continue, 如果是当前核最后一次循环,则补充计算taskIdx+2的部分
|
||||
if (s1NoNeedCalc || s2NoNeedCalc) {
|
||||
continue;
|
||||
}
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
maxS2LoopCnt -= runParam.s2LoopEndIdx;
|
||||
}
|
||||
s2LoopLimit = runParam.s2LoopEndIdx - 1;
|
||||
} else {
|
||||
s2LoopLimit = 0;
|
||||
}
|
||||
|
||||
for (int64_t s2LoopCount = 0; s2LoopCount <= s2LoopLimit; ++s2LoopCount) {
|
||||
if (notLastTwoLoop) {
|
||||
RunInfo &runInfo1 = runInfo[taskId % 3];
|
||||
this->SetRunInfo(runInfo1, runParam, taskId, s2LoopCount, s2LoopLimit, multiCoreInnerIdx);
|
||||
if ASCEND_IS_AIC {
|
||||
this->cubeBlock.IterateBmm1(this->bmm1Buffers.Get(), this->l1RightBuffers.Get(),
|
||||
this->v0ResGmBuffers.Get(), runInfo1, this->constInfo);
|
||||
} else {
|
||||
this->vecBlock.ProcessVec0(this->l1RightBuffers.Get(), this->v0ResGmBuffers.Get(),
|
||||
runInfo1, this->constInfo);
|
||||
}
|
||||
} else {
|
||||
if ASCEND_IS_AIV {
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
if (maxS2LoopCnt > 0) {
|
||||
maxS2LoopCnt--;
|
||||
CrossCoreSetFlag<QSFA_SYNC_MODE0, PIPE_MTE3>(15);
|
||||
CrossCoreWaitFlag<QSFA_SYNC_MODE0, PIPE_MTE3>(15);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
if (taskId > 0 && notLast) {
|
||||
auto &runInfo2 = runInfo[(taskId + 2) % 3];
|
||||
if ASCEND_IS_AIV {
|
||||
this->vecBlock.ProcessVec1(this->l1RightBuffers.GetReused(), this->bmm1Buffers.Get(), runInfo2,
|
||||
this->constInfo);
|
||||
} else {
|
||||
RunInfo &runInfo2 = runInfo[(taskId + 2) % 3];
|
||||
this->cubeBlock.IterateBmm2(this->bmm2Buffers.Get(), this->l1RightBuffers, this->l1RightBuffers.GetReused(), runInfo2,
|
||||
this->constInfo);
|
||||
}
|
||||
}
|
||||
if (taskId > 1) {
|
||||
if ASCEND_IS_AIV {
|
||||
RunInfo &qsfaRunInfo3 = runInfo[(taskId + 1) % 3];
|
||||
this->vecBlock.ProcessVec2(this->bmm2Buffers.Get(), qsfaRunInfo3, this->constInfo);
|
||||
}
|
||||
}
|
||||
++taskId;
|
||||
}
|
||||
++multiCoreInnerIdx;
|
||||
}
|
||||
gS1StartIdx = 0;
|
||||
}
|
||||
|
||||
if ASCEND_IS_AIV {
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
for (int64_t qsfaLoopCnt = 0; qsfaLoopCnt < maxS2LoopCnt; qsfaLoopCnt++) {
|
||||
CrossCoreSetFlag<QSFA_SYNC_MODE0, PIPE_MTE3>(15);
|
||||
CrossCoreWaitFlag<QSFA_SYNC_MODE0, PIPE_MTE3>(15);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::ComputeAxisIdxByBnAndGs1(
|
||||
int64_t bnIndex, int64_t gS1Index, RunParamStr &runParam)
|
||||
{
|
||||
// GS1合轴, 不切G, 只切S1
|
||||
runParam.s1oIdx = gS1Index * runParam.qSNumInOneBlock;
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
runParam.goIdx = (aicIdx % 2 == 0) ? 0 : 64; // N1=128场景,相邻cube核处理一个s1,第一个cube核承担0-63行g,第二个cube核承担后64行g
|
||||
} else {
|
||||
runParam.goIdx = 0;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::SetRunInfo(
|
||||
RunInfo &runInfo, RunParamStr &runParam, int64_t taskId, int64_t s2LoopCount, int64_t s2LoopLimit, int64_t multiCoreInnerIdx)
|
||||
{
|
||||
if (s2LoopCount < runParam.kvLoopEndIdx) {
|
||||
runInfo.s2StartIdx = runParam.s2LineStartIdx;
|
||||
runInfo.s2EndIdx = runParam.s2LineEndIdx;
|
||||
}
|
||||
|
||||
runInfo.s2LoopCount = s2LoopCount;
|
||||
|
||||
if (runInfo.multiCoreInnerIdx != multiCoreInnerIdx) {
|
||||
runInfo.boIdx = runParam.boIdx;
|
||||
runInfo.s1oIdx = runParam.s1oIdx;
|
||||
runInfo.n2oIdx = runParam.n2oIdx;
|
||||
runInfo.goIdx = runParam.goIdx;
|
||||
|
||||
runInfo.multiCoreInnerIdx = multiCoreInnerIdx;
|
||||
runInfo.multiCoreIdxMod2 = multiCoreInnerIdx & 1;
|
||||
runInfo.multiCoreIdxMod3 = multiCoreInnerIdx % 3;
|
||||
}
|
||||
|
||||
runInfo.s2LoopLimit = s2LoopLimit;
|
||||
runInfo.taskId = taskId;
|
||||
runInfo.taskIdMod2 = taskId & 1;
|
||||
runInfo.taskIdMod3 = taskId % 3;
|
||||
|
||||
runInfo.sOuterOffset = runParam.sOuterOffset;
|
||||
runInfo.actualS1Size = runParam.actualS1Size;
|
||||
runInfo.actualS2Size = runParam.actualS2Size;
|
||||
runInfo.attentionOutOffset = runParam.attentionOutOffset;
|
||||
this->ComputeBmm1Tail(runInfo, runParam);
|
||||
InitUniqueRunInfo(runParam, runInfo);
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::InitUniqueRunInfo(
|
||||
const RunParamStr &runParam, RunInfo &runInfo)
|
||||
{
|
||||
InitTaskParamByRun<TEMPLATE_INTF_ARGS>(runParam, runInfo);
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::ComputeBmm1Tail(
|
||||
RunInfo &runInfo, RunParamStr &runParam)
|
||||
{
|
||||
// ------------------------S1 Base Related---------------------------
|
||||
runInfo.s1RealSize = runParam.s1RealSize;
|
||||
runInfo.halfS1RealSize = runParam.halfS1RealSize;
|
||||
runInfo.firstHalfS1RealSize = runParam.firstHalfS1RealSize;
|
||||
|
||||
runInfo.halfMRealSize = runParam.halfMRealSize;
|
||||
runInfo.firstHalfMRealSize = runParam.firstHalfMRealSize;
|
||||
runInfo.mRealSize = runParam.mRealSize;
|
||||
|
||||
runInfo.vec2S1BaseSize = runInfo.halfS1RealSize;
|
||||
runInfo.vec2MBaseSize = runInfo.halfMRealSize;
|
||||
|
||||
// ------------------------S2 Base Related----------------------------
|
||||
runInfo.s2RealSize = constInfo.s2BaseSize;
|
||||
runInfo.s2AlignedSize = runInfo.s2RealSize;
|
||||
|
||||
if (runInfo.s2StartIdx + (runInfo.s2LoopCount + 1) * runInfo.s2RealSize > runInfo.s2EndIdx) {
|
||||
runInfo.s2RealSize = runInfo.s2EndIdx - runInfo.s2LoopCount * runInfo.s2RealSize - runInfo.s2StartIdx;
|
||||
runInfo.s2AlignedSize = Align(runInfo.s2RealSize);
|
||||
}
|
||||
}
|
||||
}
|
||||
#endif // KV_QUANT_SPARSE_FLASH_ATTENTION_KERNEL_MLA_H
|
||||
@@ -0,0 +1,256 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_flash_attention_kvcache.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef KV_QUANT_SPARSE_FLASH_ATTENTION_KVCACHE_H
|
||||
#define KV_QUANT_SPARSE_FLASH_ATTENTION_KVCACHE_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "kernel_operator_list_tensor_intf.h"
|
||||
#include "kv_quant_sparse_flash_attention_common_arch35.h"
|
||||
|
||||
using namespace matmul;
|
||||
using namespace regbaseutil;
|
||||
using namespace AscendC;
|
||||
using namespace AscendC::Impl::Detail;
|
||||
static constexpr uint32_t sparseModeThree = 3;
|
||||
static constexpr uint32_t sparseModeZero = 0;
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline void GetSingleCoreParam(RunParamStr& runParam, const ConstInfo &constInfo,
|
||||
__gm__ int32_t *actualSeqQlenAddr, __gm__ int32_t * actualSeqKvlenAddr)
|
||||
{
|
||||
int32_t qsfaActualS1Size = 0;
|
||||
int32_t qsfaActualS2Size = 0;
|
||||
int32_t actualSeqMin = 1;
|
||||
int32_t actualSeqKVMin = 1;
|
||||
int32_t sIdx = runParam.boIdx;
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
// actual seq length first
|
||||
if (actualSeqQlenAddr != nullptr) {
|
||||
qsfaActualS1Size = (sIdx == 0) ? actualSeqQlenAddr[0] :
|
||||
actualSeqQlenAddr[sIdx] - actualSeqQlenAddr[sIdx - 1];
|
||||
} else {
|
||||
qsfaActualS1Size = constInfo.s1Size;
|
||||
}
|
||||
} else {
|
||||
qsfaActualS1Size = (actualSeqQlenAddr == nullptr) ? constInfo.s1Size :
|
||||
actualSeqQlenAddr[sIdx];
|
||||
}
|
||||
|
||||
if (constInfo.isActualLenDimsKVNull) {
|
||||
qsfaActualS2Size = constInfo.s2Size;
|
||||
} else {
|
||||
if constexpr (isPa) {
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
qsfaActualS2Size = actualSeqKvlenAddr[sIdx];
|
||||
} else {
|
||||
qsfaActualS2Size = (constInfo.actualSeqLenKVSize == actualSeqKVMin) ?
|
||||
actualSeqKvlenAddr[0] : actualSeqKvlenAddr[sIdx];
|
||||
}
|
||||
} else {
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
qsfaActualS2Size = (sIdx == 0) ? actualSeqKvlenAddr[0] :
|
||||
actualSeqKvlenAddr[sIdx] - actualSeqKvlenAddr[sIdx - 1];
|
||||
} else {
|
||||
qsfaActualS2Size = (constInfo.actualSeqLenKVSize == actualSeqKVMin) ?
|
||||
actualSeqKvlenAddr[0] : actualSeqKvlenAddr[sIdx];
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
runParam.actualS1Size = qsfaActualS1Size;
|
||||
runParam.actualS2Size = qsfaActualS2Size;
|
||||
runParam.preTokensPerBatch = runParam.actualS1Size;
|
||||
if (constInfo.sparseMode == sparseModeZero) {
|
||||
runParam.nextTokensPerBatch = MAX_PRE_NEXT_TOKENS;
|
||||
} else {
|
||||
runParam.nextTokensPerBatch = runParam.actualS2Size - runParam.actualS1Size;
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline void ComputeParamBatch(RunParamStr& runParam,
|
||||
const ConstInfo &constInfo, __gm__ int32_t *actualSeqQlenAddr, __gm__ int32_t *actualSeqKvlenAddr)
|
||||
{
|
||||
GetSingleCoreParam<TEMPLATE_INTF_ARGS>(runParam, constInfo, actualSeqQlenAddr, actualSeqKvlenAddr);
|
||||
}
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline void ComputeS1LoopInfo(RunParamStr& runParam, const ConstInfo &constInfo,
|
||||
bool lastBN, int64_t nextGs1Idx, int64_t gS1StartIdx)
|
||||
{
|
||||
runParam.gs1LoopStartIdx = gS1StartIdx;
|
||||
runParam.qSNumInOneBlock = 1; // qsfa 不切G轴, 计算每个基本块可以拷贝多少行s
|
||||
|
||||
if (runParam.nextTokensPerBatch < 0) {
|
||||
uint64_t invalidTokenCount = static_cast<uint64_t>(-(runParam.nextTokensPerBatch + 1)) + 1ULL;
|
||||
int64_t gs1LoopStartIdx =
|
||||
invalidTokenCount / runParam.qSNumInOneBlock * runParam.qSNumInOneBlock;
|
||||
if (gs1LoopStartIdx > gS1StartIdx) {
|
||||
runParam.gs1LoopStartIdx = gs1LoopStartIdx;
|
||||
}
|
||||
}
|
||||
|
||||
int32_t qsfaGs1LoopEndIdx = runParam.actualS1Size; // qsfa 不切G轴, 每次拷贝一行的topk,只算一行的qs
|
||||
|
||||
// 不是最后一个bn, 赋值souterBlockNum
|
||||
if (!lastBN) {
|
||||
runParam.gs1LoopEndIdx = qsfaGs1LoopEndIdx;
|
||||
} else { // 最后一个bn, 从数组下一个元素取值
|
||||
runParam.gs1LoopEndIdx = nextGs1Idx == 0 ? qsfaGs1LoopEndIdx : nextGs1Idx;
|
||||
}
|
||||
|
||||
if (runParam.gs1LoopStartIdx > runParam.gs1LoopEndIdx) {
|
||||
runParam.gs1LoopStartIdx = runParam.gs1LoopEndIdx;
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline void ComputeSouterParam(RunParamStr& runParam, const ConstInfo &constInfo,
|
||||
uint32_t sOuterLoopIdx)
|
||||
{
|
||||
int64_t qsfaCubeSOuterOffset = sOuterLoopIdx * runParam.qSNumInOneBlock;
|
||||
if (runParam.actualS1Size == 0) {
|
||||
runParam.s1RealSize = 0;
|
||||
runParam.mRealSize = 0;
|
||||
} else {
|
||||
runParam.s1RealSize = Min(runParam.qSNumInOneBlock, runParam.actualS1Size - qsfaCubeSOuterOffset);
|
||||
runParam.mRealSize = runParam.s1RealSize * constInfo.gSize;
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
runParam.mRealSize = runParam.mRealSize >> 1;
|
||||
}
|
||||
}
|
||||
|
||||
runParam.cubeMOuterOffset = qsfaCubeSOuterOffset * constInfo.gSize;
|
||||
runParam.halfMRealSize = (runParam.mRealSize + 1) >> 1;
|
||||
runParam.firstHalfMRealSize = runParam.halfMRealSize;
|
||||
if (constInfo.subBlockIdx == 0) {
|
||||
runParam.mOuterOffset = runParam.cubeMOuterOffset;
|
||||
} else {
|
||||
runParam.halfMRealSize = runParam.mRealSize - runParam.halfMRealSize;
|
||||
runParam.mOuterOffset = runParam.cubeMOuterOffset + runParam.firstHalfMRealSize;
|
||||
}
|
||||
runParam.halfS1RealSize = (runParam.s1RealSize + 1) >> 1;
|
||||
runParam.firstHalfS1RealSize = runParam.halfS1RealSize;
|
||||
|
||||
if (constInfo.subBlockIdx == 1) {
|
||||
runParam.halfS1RealSize = runParam.s1RealSize - runParam.halfS1RealSize;
|
||||
runParam.sOuterOffset = qsfaCubeSOuterOffset + runParam.halfMRealSize / constInfo.gSize;
|
||||
} else {
|
||||
runParam.sOuterOffset = qsfaCubeSOuterOffset;
|
||||
}
|
||||
runParam.cubeSOuterOffset = qsfaCubeSOuterOffset;
|
||||
}
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline void LoopSOuterOffsetInit(RunParamStr& runParam, const ConstInfo &constInfo,
|
||||
int32_t sIdx, __gm__ int32_t *cuSeqlensQAddr)
|
||||
{
|
||||
if ASCEND_IS_AIV {
|
||||
int64_t qsfaSeqOffset = 0;
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
qsfaSeqOffset = sIdx == 0 ? 0 : cuSeqlensQAddr[sIdx - 1];
|
||||
} else {
|
||||
qsfaSeqOffset = sIdx * constInfo.s1Size;
|
||||
}
|
||||
|
||||
int64_t attentionOutSeqOffset = qsfaSeqOffset * constInfo.n2GDv;
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::BSND || LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
runParam.attentionOutOffset = attentionOutSeqOffset +
|
||||
runParam.sOuterOffset * constInfo.n2GDv + runParam.n2oIdx * constInfo.gDv +
|
||||
runParam.goIdx * constInfo.dSizeV;
|
||||
}
|
||||
if (constInfo.subBlockIdx == 1) {
|
||||
runParam.attentionOutOffset += runParam.firstHalfMRealSize * constInfo.dSizeV;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline bool ComputeParamS1(RunParamStr& runParam, const ConstInfo &constInfo,
|
||||
uint32_t sOuterLoopIdx, __gm__ int32_t *cuSeqlensQAddr)
|
||||
{
|
||||
if (runParam.nextTokensPerBatch < 0) {
|
||||
uint64_t invalidTokenCount = static_cast<uint64_t>(-(runParam.nextTokensPerBatch + 1)) + 1ULL;
|
||||
if (runParam.s1oIdx <
|
||||
invalidTokenCount / runParam.qSNumInOneBlock * runParam.qSNumInOneBlock) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
ComputeSouterParam<TEMPLATE_INTF_ARGS>(runParam, constInfo, sOuterLoopIdx);
|
||||
LoopSOuterOffsetInit<TEMPLATE_INTF_ARGS>(runParam, constInfo,
|
||||
runParam.boIdx, cuSeqlensQAddr);
|
||||
return false;
|
||||
}
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline bool ComputeLastBN(RunParamStr& runParam, __gm__ int32_t *cuSeqlensQAddr)
|
||||
{
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
// TND格式下 相邻Batch中当actualSeqQlen相等时则返回true
|
||||
if (runParam.boIdx > 0 && ((runParam.boIdx == 0 && cuSeqlensQAddr[runParam.boIdx] == 0) || (cuSeqlensQAddr[runParam.boIdx] - cuSeqlensQAddr[runParam.boIdx - 1] == 0))) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline int64_t ClipSInnerTokenCube(int64_t qsfaSInnerToken, int64_t minValue, int64_t maxValue)
|
||||
{
|
||||
qsfaSInnerToken = qsfaSInnerToken > minValue ? qsfaSInnerToken : minValue;
|
||||
qsfaSInnerToken = qsfaSInnerToken < maxValue ? qsfaSInnerToken : maxValue;
|
||||
return qsfaSInnerToken;
|
||||
}
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline bool ComputeS2LoopInfo(RunParamStr& runParam, const ConstInfo &constInfo)
|
||||
{
|
||||
if (runParam.actualS2Size == 0) {
|
||||
runParam.kvLoopEndIdx = 0;
|
||||
runParam.s2LoopEndIdx = 0;
|
||||
return true;
|
||||
}
|
||||
uint32_t qsfaS2BaseSize = constInfo.s2BaseSize;
|
||||
|
||||
if (constInfo.sparseMode == sparseModeZero) {
|
||||
runParam.s2LineStartIdx = 0;
|
||||
runParam.s2LineEndIdx = Min(runParam.actualS2Size, constInfo.sparseBlockCount);
|
||||
} else if (constInfo.sparseMode == sparseModeThree) {
|
||||
runParam.s2LineStartIdx = ClipSInnerTokenCube<TEMPLATE_INTF_ARGS>(runParam.cubeSOuterOffset - runParam.preTokensPerBatch,
|
||||
0, runParam.actualS2Size);
|
||||
runParam.s2LineEndIdx = ClipSInnerTokenCube<TEMPLATE_INTF_ARGS>(runParam.cubeSOuterOffset + runParam.nextTokensPerBatch +
|
||||
runParam.s1RealSize, 0, runParam.actualS2Size);
|
||||
runParam.s2LineEndIdx = Min(runParam.s2LineEndIdx, constInfo.sparseBlockCount); // 当前LI输出的block size只可能是1
|
||||
}
|
||||
|
||||
runParam.kvLoopEndIdx = (runParam.s2LineEndIdx + qsfaS2BaseSize - 1) / qsfaS2BaseSize;
|
||||
runParam.s2LoopEndIdx = runParam.kvLoopEndIdx;
|
||||
return false;
|
||||
}
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline void InitTaskParamByRun(const RunParamStr& runParam, RunInfo &runInfo)
|
||||
{
|
||||
runInfo.boIdx = runParam.boIdx;
|
||||
runInfo.actualS1Size = runParam.actualS1Size;
|
||||
runInfo.actualS2Size = runParam.actualS2Size;
|
||||
runInfo.preTokensPerBatch = runParam.preTokensPerBatch;
|
||||
runInfo.nextTokensPerBatch = runParam.nextTokensPerBatch;
|
||||
runInfo.softmaxLseOffset = runParam.softmaxLseOffset;
|
||||
runInfo.qSNumInOneBlock = runParam.qSNumInOneBlock;
|
||||
runInfo.kvLoopEndIdx = runParam.kvLoopEndIdx;
|
||||
}
|
||||
|
||||
#endif // KV_QUANT_SPARSE_FLASH_ATTENTION_KVCACHE_H
|
||||
@@ -0,0 +1,388 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_flash_attention_service_cube_mla.h
|
||||
*/
|
||||
#ifndef KV_QUANT_SPARSE_FLASH_ATTENTION_SERVICE_CUBE_MLA_H
|
||||
#define KV_QUANT_SPARSE_FLASH_ATTENTION_SERVICE_CUBE_MLA_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "kernel_operator_list_tensor_intf.h"
|
||||
#include "kernel_tiling/kernel_tiling.h"
|
||||
#include "lib/matmul_intf.h"
|
||||
#include "lib/matrix/matmul/tiling.h"
|
||||
#include "kv_quant_sparse_flash_attention_common_arch35.h"
|
||||
|
||||
#if __has_include("../../common/op_kernel/offset_calculator.h")
|
||||
#include "../../common/op_kernel/offset_calculator.h"
|
||||
#else
|
||||
#include "../common/offset_calculator.h"
|
||||
#endif
|
||||
|
||||
#if __has_include("../../common/op_kernel/matmul.h")
|
||||
#include "../../common/op_kernel/matmul.h"
|
||||
#else
|
||||
#include "../common/matmul.h"
|
||||
#endif
|
||||
#if __has_include("../../common/op_kernel/CopyInL1.h")
|
||||
#include "../../common/op_kernel/CopyInL1.h"
|
||||
#else
|
||||
#include "../common/CopyInL1.h"
|
||||
#endif
|
||||
#if __has_include("../../common/op_kernel/FixpipeOut.h")
|
||||
#include "../../common/op_kernel/FixpipeOut.h"
|
||||
#else
|
||||
#include "../common/FixpipeOut.h"
|
||||
#endif
|
||||
|
||||
using namespace AscendC;
|
||||
using namespace AscendC::Impl::Detail;
|
||||
|
||||
using namespace fa_base_matmul;
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace BaseApi {
|
||||
struct CubeCoordInfo {
|
||||
uint32_t curBIdx;
|
||||
uint32_t s1Coord;
|
||||
uint32_t s2Coord;
|
||||
};
|
||||
|
||||
template <QSFA_LAYOUT LAYOUT>
|
||||
__aicore__ inline constexpr GmFormat GetQueryGmFormat()
|
||||
{
|
||||
if constexpr (LAYOUT == QSFA_LAYOUT::BSND) {
|
||||
return GmFormat::BSNGD;
|
||||
} else {
|
||||
return GmFormat::TNGD;
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF
|
||||
class QSFAMatmulService {
|
||||
public:
|
||||
/* =================编译期常量的基本块信息================= */
|
||||
static constexpr uint32_t s1BaseSize = 64;
|
||||
static constexpr uint32_t s2BaseSize = 128;
|
||||
static constexpr uint32_t dBaseSize = 576;
|
||||
static constexpr uint32_t dBaseMatmulSize = 128;
|
||||
|
||||
__aicore__ inline QSFAMatmulService() {};
|
||||
__aicore__ inline void InitCubeBlock(TPipe *pipe, BufferManager<BufferType::L1> *qsfaL1BufferManagerPtr,
|
||||
__gm__ uint8_t *query);
|
||||
__aicore__ inline void InitCubeInput(__gm__ uint8_t *cuSeqlensQ, const ConstInfo& constInfo);
|
||||
__aicore__ inline void IterateBmm1(Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &output,
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf,
|
||||
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm,
|
||||
RunInfo &runInfo, ConstInfo &constInfo);
|
||||
|
||||
__aicore__ inline void IterateBmm2(Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
|
||||
BuffersPolicy3buff<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputLeftBuffers,
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf, RunInfo &runInfo,
|
||||
ConstInfo &constInfo);
|
||||
|
||||
private:
|
||||
__aicore__ inline void InitLocalBuffer();
|
||||
__aicore__ inline void InitGmTensor(__gm__ uint8_t *cuSeqlensQ, const ConstInfo& constInfo);
|
||||
__aicore__ inline void CalcS1Coord(RunInfo &runInfo, ConstInfo &constInfo);
|
||||
|
||||
__aicore__ inline void IterateBmm1QSFA(Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf,
|
||||
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm,
|
||||
RunInfo &runInfo, ConstInfo &constInfo);
|
||||
__aicore__ inline void PrepareLeftMatrixBmm1QSFA(Buffer<BufferType::L1> &inputLeftBuf,
|
||||
RunInfo &runInfo, ConstInfo &constInfo);
|
||||
|
||||
// --------------------Bmm2--------------------------
|
||||
__aicore__ inline void IterateBmm2QSFA(Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
|
||||
BuffersPolicy3buff<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputLeftBuffers,
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf, RunInfo &runInfo,
|
||||
ConstInfo &constInfo);
|
||||
TPipe *tPipe;
|
||||
/* =====================GM变量==================== */
|
||||
static constexpr GmFormat Q_FORMAT = GetQueryGmFormat<LAYOUT_T>();
|
||||
FaGmTensor<Q_T, Q_FORMAT, int32_t> queryGm;
|
||||
|
||||
/* =====================运行时变量==================== */
|
||||
CubeCoordInfo coordInfo[3];
|
||||
TEventID mte1ToMte2Id[3];
|
||||
TEventID mte2ToMte1Id[3];
|
||||
|
||||
/* =====================LocalBuffer变量==================== */
|
||||
// D小于等于256 mm1左矩阵Q,GS1循环内左矩阵复用, GS1循环间开pingpong;D大于256使用单块Buffer,S1循环间驻留;fp32场景单块不驻留
|
||||
BuffersPolicySingleBuffer<BufferType::L1> l1QBuffers;
|
||||
// L0空间buffer manager
|
||||
BufferManager<BufferType::L1> *qsfaL1BufferManagerPtr;
|
||||
BufferManager<BufferType::L0A> l0aBufferManager;
|
||||
BufferManager<BufferType::L0B> l0bBufferManager;
|
||||
BufferManager<BufferType::L0C> l0cBufferManager;
|
||||
// L0A
|
||||
BuffersPolicyDB<BufferType::L0A> mmL0ABuffers;
|
||||
// L0B
|
||||
BuffersPolicyDB<BufferType::L0B> mmL0BBuffers;
|
||||
// L0C
|
||||
BuffersPolicyDB<BufferType::L0C> mmL0CBuffers;
|
||||
};
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAMatmulService<TEMPLATE_ARGS>::InitCubeBlock(
|
||||
TPipe *pipe, BufferManager<BufferType::L1> *qsfaL1BuffMgr, __gm__ uint8_t *query)
|
||||
{
|
||||
if ASCEND_IS_AIC {
|
||||
tPipe = pipe;
|
||||
qsfaL1BufferManagerPtr = qsfaL1BuffMgr;
|
||||
this->queryGm.gmTensor.SetGlobalBuffer((__gm__ Q_T *)query);
|
||||
InitLocalBuffer();
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void
|
||||
QSFAMatmulService<TEMPLATE_ARGS>::InitCubeInput(__gm__ uint8_t *qsfaActualSeqLengthsQ, const ConstInfo& constInfo)
|
||||
{
|
||||
if ASCEND_IS_AIC {
|
||||
InitGmTensor(qsfaActualSeqLengthsQ, constInfo);
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
mte1ToMte2Id[0] = GetTPipePtr()->AllocEventID<HardEvent::MTE2_MTE1>();
|
||||
mte1ToMte2Id[1] = GetTPipePtr()->AllocEventID<HardEvent::MTE2_MTE1>();
|
||||
mte1ToMte2Id[2] = GetTPipePtr()->AllocEventID<HardEvent::MTE2_MTE1>();
|
||||
mte2ToMte1Id[0] = GetTPipePtr()->AllocEventID<HardEvent::MTE1_MTE2>();
|
||||
mte2ToMte1Id[1] = GetTPipePtr()->AllocEventID<HardEvent::MTE1_MTE2>();
|
||||
mte2ToMte1Id[2] = GetTPipePtr()->AllocEventID<HardEvent::MTE1_MTE2>();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void
|
||||
QSFAMatmulService<TEMPLATE_ARGS>::InitLocalBuffer()
|
||||
{
|
||||
constexpr uint32_t mm1LeftSize = s1BaseSize * dBaseSize * sizeof(Q_T);
|
||||
l1QBuffers.Init((*qsfaL1BufferManagerPtr), mm1LeftSize);
|
||||
|
||||
// L0A B C 当前写死,能否通过基础api获取
|
||||
l0aBufferManager.Init(tPipe, L0AB_SHARED_SIZE_64K);
|
||||
l0bBufferManager.Init(tPipe, L0AB_SHARED_SIZE_64K);
|
||||
l0cBufferManager.Init(tPipe, L0C_SHARED_SIZE_256K);
|
||||
|
||||
mmL0ABuffers.Init(l0aBufferManager, BUFFER_SIZE_16K); // db类型,填入数值是总大小的一半
|
||||
mmL0BBuffers.Init(l0bBufferManager, BUFFER_SIZE_32K);
|
||||
mmL0CBuffers.Init(l0cBufferManager, BUFFER_SIZE_128K);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void
|
||||
QSFAMatmulService<TEMPLATE_ARGS>::InitGmTensor(__gm__ uint8_t *qsfaActualSeqLengthsQ, const ConstInfo& constInfo)
|
||||
{
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::BSND) {
|
||||
this->queryGm.offsetCalculator.Init(constInfo.bSize, constInfo.n2Size, constInfo.gSize,
|
||||
constInfo.s1Size, constInfo.dSize);
|
||||
} else { // QSFA_LAYOUT::TND
|
||||
GlobalTensor<int32_t> actualSeqQLen;
|
||||
actualSeqQLen.SetGlobalBuffer((__gm__ int32_t *)qsfaActualSeqLengthsQ);
|
||||
this->queryGm.offsetCalculator.Init(constInfo.n2Size, constInfo.gSize, constInfo.dSize,
|
||||
actualSeqQLen, constInfo.actualSeqLenSize);
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAMatmulService<TEMPLATE_ARGS>::CalcS1Coord(RunInfo &runInfo,
|
||||
ConstInfo &constInfo)
|
||||
{
|
||||
// 计算s1方向偏移
|
||||
coordInfo[runInfo.taskIdMod3].s1Coord = runInfo.s1oIdx * runInfo.qSNumInOneBlock;
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAMatmulService<TEMPLATE_ARGS>::IterateBmm1(
|
||||
Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf,
|
||||
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm,
|
||||
RunInfo &runInfo, ConstInfo &constInfo)
|
||||
{
|
||||
CalcS1Coord(runInfo, constInfo);
|
||||
|
||||
IterateBmm1QSFA(outputBuf, inputRightBuf, v0ResGm, runInfo, constInfo);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAMatmulService<TEMPLATE_ARGS>::IterateBmm2(Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
|
||||
BuffersPolicy3buff<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputLeftBuffers,
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf, RunInfo &runInfo,
|
||||
ConstInfo &constInfo)
|
||||
{
|
||||
IterateBmm2QSFA(outputBuf, inputLeftBuffers, inputRightBuf, runInfo, constInfo);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAMatmulService<TEMPLATE_ARGS>::IterateBmm1QSFA(
|
||||
Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf,
|
||||
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm,
|
||||
RunInfo &runInfo, ConstInfo &constInfo)
|
||||
{
|
||||
Buffer<BufferType::L1> inputLeftBuf;
|
||||
PrepareLeftMatrixBmm1QSFA(inputLeftBuf, runInfo, constInfo);
|
||||
|
||||
// 加载当前轮的右矩阵到L1
|
||||
inputRightBuf.WaitCrossCore(); // 核间同步,这里需要根据V0操作处理同步,确保取tensor时,数据已经准备好
|
||||
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
SetFlag<HardEvent::MTE1_MTE2>(mte2ToMte1Id[runInfo.taskIdMod3]);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(mte2ToMte1Id[runInfo.taskIdMod3]);
|
||||
LocalTensor<Q_T> dst = inputRightBuf.GetTensor<Q_T>();
|
||||
v0ResGm.WaitCrossCore();
|
||||
GlobalTensor<Q_T> v0ResGmTensor = v0ResGm.template GetTensor<Q_T>();
|
||||
DataCopy(dst, v0ResGmTensor, Align16Func(runInfo.s2RealSize) * constInfo.dSize);
|
||||
SetFlag<HardEvent::MTE2_MTE1>(mte1ToMte2Id[runInfo.taskIdMod3]);
|
||||
WaitFlag<HardEvent::MTE2_MTE1>(mte1ToMte2Id[runInfo.taskIdMod3]);
|
||||
}
|
||||
|
||||
inputLeftBuf.Wait<HardEvent::MTE2_MTE1>(); // 等待L1A
|
||||
Buffer<BufferType::L0C> mm1ResL0C = mmL0CBuffers.Get();
|
||||
mm1ResL0C.Wait<HardEvent::FIX_M>(); // 占用
|
||||
|
||||
MMParam param = {static_cast<uint32_t>(runInfo.mRealSize), // singleM
|
||||
static_cast<uint32_t>(runInfo.s2RealSize), // singleN
|
||||
static_cast<uint32_t>(constInfo.dSize), // singleK
|
||||
0, // isLeftTranspose
|
||||
1 // isRightTranspose
|
||||
};
|
||||
|
||||
MatmulK<Q_T, Q_T, T, s1BaseSize, s2BaseSize, dBaseMatmulSize, ABLayout::MK, ABLayout::KN>( // m,n不切,k切128
|
||||
inputLeftBuf.GetTensor<Q_T>(), inputRightBuf.GetTensor<Q_T>(), // mm1B直接用tensor的数据
|
||||
mmL0ABuffers, mmL0BBuffers, mm1ResL0C.GetTensor<T>(), param);
|
||||
|
||||
if (unlikely(runInfo.s2LoopCount == runInfo.s2LoopLimit)) {
|
||||
inputLeftBuf.Set<HardEvent::MTE1_MTE2>(); // 释放L1A
|
||||
}
|
||||
|
||||
mm1ResL0C.Set<HardEvent::M_FIX>(); // 通知
|
||||
mm1ResL0C.Wait<HardEvent::M_FIX>(); // 等待L0C
|
||||
|
||||
outputBuf.WaitCrossCore();
|
||||
FixpipeParamsC310<CO2Layout::ROW_MAJOR> fixpipeParams; // L0C→UB
|
||||
fixpipeParams.mSize = Align2Func(runInfo.mRealSize); // 有效数据不足16行,只需要输出部分行即可;
|
||||
fixpipeParams.nSize = Align8Func(runInfo.s2RealSize); // L0C上的bmm1结果矩阵N方向的size大小; 同mmadParams.n; 为什么要8个元素对齐(32B对齐) // 128
|
||||
fixpipeParams.srcStride = Align16Func(fixpipeParams.mSize); // L0C上bmm1结果相邻连续数据片段间隔(前面一个数据块的头与后面数据块的头的间隔), 单位为16*sizeof(T) // 源Nz矩阵中相邻大Z排布的起始地址偏移
|
||||
fixpipeParams.dstStride = s2BaseSize; // mmResUb上两行之间的间隔,单位:element。 // 128:根据比对dump文件得到, ND方案(S1*S2)时脏数据用mask剔除
|
||||
fixpipeParams.dualDstCtl = 1; // 双目标模式,按M维度拆分,M / 2 * N写入每个UB, M必须为2的倍数
|
||||
fixpipeParams.params.srcNdStride = 0;
|
||||
fixpipeParams.params.dstNdStride = 0;
|
||||
fixpipeParams.params.ndNum = 1;
|
||||
|
||||
Fixpipe<T, T, PFA_CFG_ROW_MAJOR_UB>(outputBuf.template GetTensor<T>(), mm1ResL0C.GetTensor<T>(), fixpipeParams); // 将matmul结果从L0C搬运到UB
|
||||
mm1ResL0C.Set<HardEvent::FIX_M>(); // 释放L0C
|
||||
outputBuf.SetCrossCore();
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAMatmulService<TEMPLATE_ARGS>::PrepareLeftMatrixBmm1QSFA(
|
||||
Buffer<BufferType::L1> &inputLeftBuf,
|
||||
RunInfo &runInfo, ConstInfo &constInfo)
|
||||
{
|
||||
// 左矩阵复用,S2的第一次循环加载左矩阵
|
||||
// 加载左矩阵到L1, 全载
|
||||
if (unlikely(runInfo.s2LoopCount == 0)) { // sOuter循环第一个基本块:搬运Q
|
||||
inputLeftBuf = l1QBuffers.Get();
|
||||
inputLeftBuf.Wait<HardEvent::MTE1_MTE2>(); // 占用L1A
|
||||
LocalTensor<Q_T> inputLeftTensor = inputLeftBuf.GetTensor<Q_T>();
|
||||
uint64_t gmOffset = this->queryGm.offsetCalculator.GetOffset(runInfo.boIdx, runInfo.n2oIdx, runInfo.goIdx,
|
||||
coordInfo[runInfo.taskIdMod3].s1Coord, 0);
|
||||
CopyToL1Nd2Nz<Q_T>(inputLeftTensor, this->queryGm.gmTensor[gmOffset], runInfo.mRealSize, constInfo.dSize,
|
||||
constInfo.mm1Ka);
|
||||
|
||||
inputLeftBuf.Set<HardEvent::MTE2_MTE1>(); // 通知
|
||||
} else { // 非S2的第一次循环直接复用Q
|
||||
inputLeftBuf = l1QBuffers.GetPre();
|
||||
// 左矩阵复用时,sinner循环内不需要MTE2同步等待
|
||||
inputLeftBuf.Set<HardEvent::MTE2_MTE1>(); // 通知
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAMatmulService<TEMPLATE_ARGS>::IterateBmm2QSFA(
|
||||
Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
|
||||
BuffersPolicy3buff<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputLeftBuffers,
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf,
|
||||
RunInfo &runInfo, ConstInfo &constInfo)
|
||||
{
|
||||
inputRightBuf.WaitCrossCore();
|
||||
Buffer<BufferType::L0C> mm2ResL0C = mmL0CBuffers.Get();
|
||||
mm2ResL0C.Wait<HardEvent::FIX_M>(); // 占用
|
||||
|
||||
MMParam qsfaParam = {static_cast<uint32_t>(runInfo.mRealSize), // singleM 64
|
||||
static_cast<uint32_t>(constInfo.dSizeNope), // singleN 576->512
|
||||
static_cast<uint32_t>(runInfo.s2RealSize), // singleK 128
|
||||
0, 0};
|
||||
|
||||
MatmulN<Q_T, Q_T, T, s1BaseSize, s2BaseSize, dBaseMatmulSize, ABLayout::MK, ABLayout::KN>(
|
||||
inputRightBuf.GetTensor<Q_T>(s2BaseSize * constInfo.dSizeNope), // 左矩阵P 来自rope位置
|
||||
inputRightBuf.GetTensor<Q_T>(), // 右矩阵V nope
|
||||
mmL0ABuffers, mmL0BBuffers,
|
||||
mm2ResL0C.GetTensor<T>(), qsfaParam);
|
||||
|
||||
inputRightBuf.SetCrossCore(); // bmm2才释放KV,在这里释放
|
||||
mm2ResL0C.Set<HardEvent::M_FIX>(); // 通知
|
||||
mm2ResL0C.Wait<HardEvent::M_FIX>(); // 等待
|
||||
|
||||
outputBuf.WaitCrossCore(); //占用
|
||||
|
||||
FixpipeParamsC310<CO2Layout::ROW_MAJOR> fixpipeParams; // L0C→UB;FixpipeParamsM300:L0C→UB
|
||||
fixpipeParams.mSize = Align2Func(runInfo.mRealSize); // 有效数据不足16行,只需要输出部分行即可;
|
||||
fixpipeParams.nSize = Align8Func(constInfo.dSizeNope); // L0C上的bmm1结果矩阵N方向的size大小, 分档计算且vector2中通过mask筛选出实际有效值
|
||||
fixpipeParams.srcStride = Align16Func(fixpipeParams.mSize); // L0C上bmm1结果相邻连续数据片段间隔(前面一个数据块的头与后面数据块的头的间隔)
|
||||
fixpipeParams.dstStride = Align16Func(constInfo.dSizeNope);
|
||||
fixpipeParams.dualDstCtl = 1;
|
||||
fixpipeParams.params.srcNdStride = 0;
|
||||
fixpipeParams.params.dstNdStride = 0;
|
||||
fixpipeParams.params.ndNum = 1;
|
||||
|
||||
Fixpipe<T, T, PFA_CFG_ROW_MAJOR_UB>(outputBuf.template GetTensor<T>(), mm2ResL0C.GetTensor<T>(), fixpipeParams); // 将matmul结果从L0C搬运到UB
|
||||
mm2ResL0C.Set<HardEvent::FIX_M>(); // 释放
|
||||
|
||||
outputBuf.SetCrossCore();
|
||||
}
|
||||
|
||||
TEMPLATES_DEF
|
||||
class QSFAMatmulServiceDummy {
|
||||
public:
|
||||
__aicore__ inline QSFAMatmulServiceDummy() {};
|
||||
__aicore__ inline void InitCubeBlock(TPipe *pipe, BufferManager<BufferType::L1> *qsfaL1BufferManagerPtr,
|
||||
__gm__ uint8_t *query) {}
|
||||
__aicore__ inline void InitCubeInput(__gm__ uint8_t *cuSeqlensQ, const ConstInfo& constInfo) {}
|
||||
__aicore__ inline void IterateBmm1(Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf, RunInfo &runInfo,
|
||||
ConstInfo &constInfo) {}
|
||||
|
||||
__aicore__ inline void IterateBmm2(Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
|
||||
BuffersPolicyDB<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputLeftBuffers,
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf, RunInfo &runInfo,
|
||||
ConstInfo &constInfo) {}
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
struct CubeBlockTraits; // 声明
|
||||
/* 生成CubeBlockTraits */
|
||||
#define GEN_TRAIT_TYPE(name, ...) using name##_TRAITS = name;
|
||||
#define GEN_TRAIT_CONST(name, type, ...) static constexpr type name##Traits = name;
|
||||
|
||||
#define DEFINE_QSFA_CUBE_BLOCK_TRAITS(CUBE_BLOCK_CLASS) \
|
||||
TEMPLATES_DEF_NO_DEFAULT \
|
||||
struct CubeBlockTraits<CUBE_BLOCK_CLASS<TEMPLATE_ARGS>> { \
|
||||
QSFA_CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_TRAIT_TYPE) \
|
||||
QSFA_CUBE_BLOCK_TRAITS_CONST_FIELDS(GEN_TRAIT_CONST) \
|
||||
}
|
||||
|
||||
DEFINE_QSFA_CUBE_BLOCK_TRAITS(QSFAMatmulService);
|
||||
DEFINE_QSFA_CUBE_BLOCK_TRAITS(QSFAMatmulServiceDummy);
|
||||
|
||||
// /* 生成Arg Traits, kernel中只需要调用ARGS_TRAITS就可以获取所有CubeBlock中的模板参数 */
|
||||
#define GEN_ARGS_TYPE(name, ...) using name = typename CubeBlockTraits<CubeBlockType>::name##_TRAITS;
|
||||
#define GEN_ARGS_CONST(name, type, ...) static constexpr type name = CubeBlockTraits<CubeBlockType>::name##Traits;
|
||||
#define ARGS_TRAITS \
|
||||
QSFA_CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_ARGS_TYPE) \
|
||||
QSFA_CUBE_BLOCK_TRAITS_CONST_FIELDS(GEN_ARGS_CONST)
|
||||
}
|
||||
#endif // KV_QUANT_SPARSE_FLASH_ATTENTION_SERVICE_CUBE_MLA_H
|
||||
@@ -0,0 +1,894 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_flash_attention_service_vector_mla.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef KV_QUANT_SPARSE_FLASH_ATTENTION_SERVICE_VECTOR_MLA_H
|
||||
#define KV_QUANT_SPARSE_FLASH_ATTENTION_SERVICE_VECTOR_MLA_H
|
||||
|
||||
#include "kv_quant_sparse_flash_attention_common_arch35.h"
|
||||
#include "kernel_operator_list_tensor_intf.h"
|
||||
#include "lib/matrix/matmul/tiling.h"
|
||||
#include "lib/matmul_intf.h"
|
||||
#if __has_include("../../common/op_kernel/arch35/vf/vf_mul_sel_softmaxflashv2_cast_nz_sfa.h")
|
||||
#include "../../common/op_kernel/arch35/vf/vf_mul_sel_softmaxflashv2_cast_nz_sfa.h"
|
||||
#else
|
||||
#include "../../common/arch35/vf/vf_mul_sel_softmaxflashv2_cast_nz_sfa.h"
|
||||
#endif
|
||||
#if __has_include("../../common/op_kernel/arch35/vf/vf_flashupdate_new.h")
|
||||
#include "../../common/op_kernel/arch35/vf/vf_flashupdate_new.h"
|
||||
#else
|
||||
#include "../../common/arch35/vf/vf_flashupdate_new.h"
|
||||
#endif
|
||||
|
||||
using namespace AscendC;
|
||||
using namespace FaVectorApi;
|
||||
using namespace AscendC::Impl::Detail;
|
||||
using namespace regbaseutil;
|
||||
using namespace matmul;
|
||||
|
||||
namespace BaseApi {
|
||||
|
||||
TEMPLATES_DEF
|
||||
class QSFAVectorService {
|
||||
public:
|
||||
// BUFFER的字节数
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_32B = 32;
|
||||
/* =================编译期常量的基本块信息================= */
|
||||
static constexpr uint32_t s1BaseSize = 64;
|
||||
static constexpr uint32_t s2BaseSize = 128;
|
||||
static constexpr uint32_t vec1Srcstride = (s1BaseSize >> 1) + 1;
|
||||
static constexpr uint32_t dVTemplateType = 512;
|
||||
static constexpr uint32_t qsfaDTemplateAlign64 = Align64Func(dVTemplateType);
|
||||
static constexpr uint32_t dVTemplateTypeInput = 672;
|
||||
static constexpr float R0 = 1.0f;
|
||||
static constexpr uint64_t SYNC_SINKS_BUF_FLAG = 6;
|
||||
|
||||
// ==================== Functions ======================
|
||||
__aicore__ inline QSFAVectorService() {};
|
||||
__aicore__ inline void InitVecBlock(TPipe *pipe, const KvQuantSparseFlashAttentionTilingDataMla *__restrict tiling,
|
||||
CVSharedParams &sharedParams, int32_t aicIdx, uint8_t subBlockIdx, __gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengths)
|
||||
{
|
||||
if ASCEND_IS_AIV {
|
||||
tilingData = tiling;
|
||||
tPipe = pipe;
|
||||
if (actualSeqLengths != nullptr) {
|
||||
actualSeqLengthsKVGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengths);
|
||||
}
|
||||
if (actualSeqLengthsQ != nullptr) {
|
||||
cuSeqlensQGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsQ);
|
||||
}
|
||||
|
||||
this->InitCubeVecSharedParams(sharedParams, aicIdx, subBlockIdx);
|
||||
this->GetExtremeValue(this->negativeFloatScalar);
|
||||
}
|
||||
}
|
||||
|
||||
// 初始化LocalTensor
|
||||
__aicore__ inline void InitLocalBuffer(TPipe *pipe, ConstInfo &constInfo);
|
||||
// 初始化attentionOutGM
|
||||
__aicore__ inline void CleanOutput(__gm__ uint8_t *attentionOut, ConstInfo &constInfo);
|
||||
__aicore__ inline void InitGlobalBuffer(__gm__ uint8_t *key, __gm__ uint8_t *value, __gm__ uint8_t *sparseIndices,
|
||||
__gm__ uint8_t *blockTable);
|
||||
__aicore__ inline void InitOutputSingleCore(ConstInfo &constInfo);
|
||||
__aicore__ inline void ProcessVec0(Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &outputL1,
|
||||
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm,
|
||||
const RunInfo &runInfo, ConstInfo &constInfo);
|
||||
__aicore__ inline void ProcessVec1(Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &outputBuf,
|
||||
Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &bmm1ResBuf, RunInfo &runInfo,
|
||||
ConstInfo &constInfo);
|
||||
using mm2ResPos = Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH>;
|
||||
__aicore__ inline void ProcessVec2(mm2ResPos &bmm2ResBuf, RunInfo &runInfo,
|
||||
ConstInfo &constInfo);
|
||||
|
||||
private:
|
||||
__aicore__ inline void ProcessVec1SoftmaxDispatchQSFA(LocalTensor<Q_T> &stage1CastTensor,
|
||||
LocalTensor<T> &mmRes, LocalTensor<float> &sumUb, LocalTensor<float> &maxUb,
|
||||
LocalTensor<T> &apiTmpBuffer, RunInfo &runInfo, ConstInfo &constInfo);
|
||||
__aicore__ inline void ProcessSparseKv(Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &outputL1,
|
||||
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm,
|
||||
const RunInfo &runInfo, ConstInfo &constInfo);
|
||||
__aicore__ inline void CalSparseCalSize(const RunInfo &runInfo, ConstInfo &constInfo);
|
||||
__aicore__ inline int64_t GetkeyOffset(int64_t s2Idx, const RunInfo &runInfo, ConstInfo &constInfo);
|
||||
__aicore__ inline void GetRealCmpS2Idx(int64_t &token0Idx, int64_t &token1Idx, int64_t s2IdxInBase,
|
||||
const RunInfo &runInfo, ConstInfo &constInfo);
|
||||
__aicore__ inline void CopyInKvNotSparse(LocalTensor<KV_T> kvMergUb, int64_t v0Loop, int64_t dealRow,
|
||||
int64_t s2StartIdx, const RunInfo &runInfo, ConstInfo &constInfo);
|
||||
__aicore__ inline uint32_t CopyInKvSparse(LocalTensor<KV_T> kvInUb , int64_t startRow, int64_t token0Idx,
|
||||
int64_t token1Idx, const RunInfo &runInfo, ConstInfo &constInfo);
|
||||
__aicore__ inline void DequantKv(LocalTensor<Q_T> antiKvTensorAsB16, LocalTensor<KV_T> srcTensor, int64_t dealRow,
|
||||
ConstInfo &constInfo);
|
||||
__aicore__ inline void CopyOutKvUb2L1(Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &outputL1,
|
||||
LocalTensor<Q_T> antiKvTensorAsB16, int64_t dealRow, int64_t s2StartIdx,
|
||||
const RunInfo &runInfo, ConstInfo &constInfo);
|
||||
__aicore__ inline void CopyOutKvUb2Gm(Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm,
|
||||
LocalTensor<Q_T> antiKvTensorAsB16, int64_t dealRow, int64_t s2StartIdx, const RunInfo &runInfo,
|
||||
ConstInfo &constInfo);
|
||||
__aicore__ inline void CopyOutMrgeResult(Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &outputL1,
|
||||
int64_t mte2Size, int64_t mte3Size, int64_t s2keyOffset, int64_t mergeMte3Idx, const RunInfo &runInfo);
|
||||
__aicore__ inline void CopyInSingleKv(LocalTensor<KV_T> kvInUb, int64_t startRow,
|
||||
int64_t keyOffset, uint32_t combineBytes);
|
||||
/* VEC2_RES_T 表示bmm2ResUb当前的类型,VEC2_RES_T = Q_T那么不需要做Cast。另外,无效行场景当前默认需要做Cast */
|
||||
using VEC2_RES_T = T;
|
||||
template <typename VEC2_RES_T>
|
||||
__aicore__ inline void Bmm2DataCopyOut(RunInfo &runInfo, ConstInfo &constInfo,
|
||||
LocalTensor<VEC2_RES_T> &vec2ResUb, int64_t vec2S1Idx, int64_t qsfaVec2CalcSize = 0);
|
||||
template <typename VEC2_RES_T>
|
||||
__aicore__ inline void CopyOutAttentionOut(
|
||||
RunInfo &runInfo, ConstInfo &constInfo, LocalTensor<VEC2_RES_T> &vec2ResUb, int64_t vec2S1Idx,
|
||||
int64_t qsfaVec2CalcSize);
|
||||
__aicore__ inline void SoftmaxInitBuffer();
|
||||
__aicore__ inline void InitCubeVecSharedParams(CVSharedParams &sharedParams, int32_t aicIdx, uint8_t subBlockIdx);
|
||||
__aicore__ inline void ComputeNeedInitQSFA(CVSharedParams &sharedParams) const;
|
||||
__aicore__ inline void GetExtremeValue(T &negativeScalar);
|
||||
|
||||
TPipe *tPipe;
|
||||
const KvQuantSparseFlashAttentionTilingDataMla *__restrict tilingData;
|
||||
|
||||
GlobalTensor<OUTPUT_T> attentionOutGm;
|
||||
GlobalTensor<KV_T> keyGm;
|
||||
GlobalTensor<int32_t> SparseIndicesGm;
|
||||
GlobalTensor<int32_t> blockTableGm;
|
||||
GlobalTensor<int32_t> cuSeqlensQGm;
|
||||
GlobalTensor<int32_t> actualSeqLengthsKVGm;
|
||||
|
||||
TBuf<> commonTBuf; // common的复用空间
|
||||
TQue<QuePosition::VECOUT, 1> stage1OutQue[2]; // 2份表示可能存在pingpong
|
||||
TQue<QuePosition::VECIN, 2> stage0InQue; // for v0 input, 2份表示可能存在pingpong
|
||||
TQue<QuePosition::VECOUT, 2> stage0OutQue; // for v0 output, 2份表示可能存在pingpong
|
||||
TBuf<> stage2OutBuf;
|
||||
TEventID mte3ToVId[2]; // 存放MTE3_V的eventId, 2份表示可能存在pingpong
|
||||
TEventID vToMte3Id[2]; // 存放V_MTE3的eventId, 2份表示可能存在pingpong
|
||||
TBuf<> softmaxMaxBuf[2];
|
||||
TBuf<> softmaxSumBuf[2];
|
||||
TBuf<> softmaxExpBuf[2];
|
||||
|
||||
T negativeFloatScalar;
|
||||
uint32_t maxBlockNumPerBatch;
|
||||
uint32_t blockSize;
|
||||
int64_t qsfaSparseCalSize;
|
||||
int64_t sparseS2Start;
|
||||
int64_t sparseS2End;
|
||||
};
|
||||
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::GetRealCmpS2Idx(int64_t &token0Idx, int64_t &token1Idx,
|
||||
int64_t s2IdxInBase, const RunInfo &runInfo, ConstInfo &constInfo)
|
||||
{
|
||||
int64_t topkBS1Idx = 0;
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
uint64_t actualSeqQPrefixSum = runInfo.boIdx == 0 ? 0 : cuSeqlensQGm.GetValue(runInfo.boIdx - 1);
|
||||
topkBS1Idx += (actualSeqQPrefixSum + runInfo.s1oIdx) * constInfo.sparseBlockCount; // T, N2(1), K
|
||||
} else {
|
||||
topkBS1Idx += runInfo.boIdx * constInfo.s1Size * constInfo.sparseBlockCount +
|
||||
runInfo.s1oIdx * constInfo.sparseBlockCount; // B, S1, N2(1), K
|
||||
}
|
||||
|
||||
int64_t qsfaCmpS2LoopCnt = runInfo.s2LoopCount;
|
||||
int64_t qsfaTopkIdx = s2IdxInBase + qsfaCmpS2LoopCnt * constInfo.s2BaseSize;
|
||||
|
||||
if (unlikely(qsfaTopkIdx >= constInfo.sparseBlockCount)) {
|
||||
token0Idx = -1;
|
||||
} else {
|
||||
token0Idx = SparseIndicesGm.GetValue(topkBS1Idx + qsfaTopkIdx) + runInfo.s2StartIdx;
|
||||
}
|
||||
qsfaTopkIdx += 1;
|
||||
if (unlikely((qsfaTopkIdx >= constInfo.sparseBlockCount) || (s2IdxInBase + 1 >= sparseS2End))) {
|
||||
token1Idx = -1;
|
||||
} else {
|
||||
token1Idx = SparseIndicesGm.GetValue(topkBS1Idx + qsfaTopkIdx) + runInfo.s2StartIdx;
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline int64_t QSFAVectorService<TEMPLATE_ARGS>::GetkeyOffset(int64_t s2Idx, const RunInfo &runInfo, ConstInfo &constInfo)
|
||||
{
|
||||
if (s2Idx < 0) {
|
||||
return -1;
|
||||
}
|
||||
int64_t realkeyOffset = 0;
|
||||
if constexpr (isPa) {
|
||||
int64_t blkTableIdx = s2Idx / blockSize;
|
||||
int64_t blkTableOffset = s2Idx % blockSize;
|
||||
realkeyOffset = blockTableGm.GetValue(runInfo.boIdx * maxBlockNumPerBatch + blkTableIdx) *
|
||||
static_cast<int64_t>(blockSize) * constInfo.dSizeVInput +
|
||||
blkTableOffset * constInfo.dSizeVInput; // BlockNum, BlockSize, N(1), D
|
||||
} else {
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::BSND) {
|
||||
realkeyOffset = (runInfo.boIdx * constInfo.s2Size + s2Idx) * constInfo.dSizeVInput; // BSN(1)D
|
||||
} else if constexpr (LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
int64_t batchKvStart = (runInfo.boIdx == 0) ? 0 : actualSeqLengthsKVGm.GetValue(runInfo.boIdx - 1);
|
||||
realkeyOffset = (batchKvStart + s2Idx) * constInfo.dSizeVInput;
|
||||
}
|
||||
}
|
||||
return realkeyOffset;
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void
|
||||
QSFAVectorService<TEMPLATE_ARGS>::CopyInSingleKv(LocalTensor<KV_T> kvInUb, int64_t startRow,
|
||||
int64_t keyOffset, uint32_t combineBytes)
|
||||
{
|
||||
if (keyOffset < 0) {
|
||||
return;
|
||||
}
|
||||
DataCopyExtParams intriParams;
|
||||
|
||||
intriParams.blockCount = 1;
|
||||
intriParams.dstStride = 0;
|
||||
intriParams.srcStride = 0;
|
||||
DataCopyPadExtParams<KV_T> padParams;
|
||||
// 当前仅支持COMBINE模式
|
||||
intriParams.blockLen = combineBytes;
|
||||
uint32_t combineDim = combineBytes / sizeof(KV_T);
|
||||
uint32_t combineDimAlign = CeilAlign(combineBytes, BUFFER_SIZE_BYTE_32B) / sizeof(KV_T);
|
||||
padParams.isPad = true;
|
||||
padParams.leftPadding = 0;
|
||||
padParams.rightPadding = combineDimAlign - combineDim;
|
||||
padParams.paddingValue = 0;
|
||||
DataCopyPad(kvInUb[startRow * combineDimAlign], keyGm[keyOffset], intriParams, padParams);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline uint32_t QSFAVectorService<TEMPLATE_ARGS>::CopyInKvSparse(LocalTensor<KV_T> kvInUb , int64_t startRow,
|
||||
int64_t token0Idx, int64_t token1Idx, const RunInfo &runInfo, ConstInfo &constInfo)
|
||||
{
|
||||
int64_t keyOffset0 = GetkeyOffset(token0Idx, runInfo, constInfo);
|
||||
int64_t keyOffset1 = GetkeyOffset(token1Idx, runInfo, constInfo);
|
||||
if (unlikely(keyOffset0 < 0 && keyOffset1 < 0)) {
|
||||
return 0;
|
||||
}
|
||||
uint32_t combineBytes = constInfo.dSizeVInput * sizeof(KV_T);
|
||||
int64_t keySrcStride = (keyOffset0 > keyOffset1 ? (keyOffset0 - keyOffset1) * sizeof(KV_T):
|
||||
(keyOffset1 - keyOffset0)) * sizeof(KV_T) - combineBytes;
|
||||
if (keySrcStride >= INT32_MAX || keySrcStride < 0 || constInfo.sparseBlockSize > 1) {
|
||||
// stride溢出、stride为负数、s2超长等异常场景,还原成2条搬运指令
|
||||
CopyInSingleKv(kvInUb, startRow, keyOffset0, combineBytes);
|
||||
CopyInSingleKv(kvInUb, startRow + 1, keyOffset1, combineBytes);
|
||||
} else {
|
||||
DataCopyExtParams intriParams;
|
||||
intriParams.blockCount = (keyOffset0 >= 0) + (keyOffset1 >= 0);
|
||||
intriParams.blockLen = combineBytes;
|
||||
intriParams.dstStride = 0;
|
||||
intriParams.srcStride = keySrcStride;
|
||||
DataCopyPadExtParams<KV_T> padParams;
|
||||
|
||||
int64_t keyOffset = keyOffset0 > -1 ? keyOffset0 : keyOffset1;
|
||||
if (keyOffset1 > -1 && keyOffset1 < keyOffset0) {
|
||||
keyOffset = keyOffset1;
|
||||
}
|
||||
|
||||
// 当前仅支持COMBINE模式
|
||||
uint32_t combineDim = combineBytes / sizeof(KV_T);
|
||||
uint32_t combineDimAlign = CeilAlign(combineBytes, BUFFER_SIZE_BYTE_32B) / sizeof(KV_T);
|
||||
padParams.isPad = true;
|
||||
padParams.leftPadding = 0;
|
||||
padParams.rightPadding = combineDimAlign - combineDim;
|
||||
padParams.paddingValue = 0;
|
||||
DataCopyPad(kvInUb[startRow * combineDimAlign], keyGm[keyOffset], intriParams, padParams);
|
||||
}
|
||||
return (keyOffset0 > -1) + (keyOffset1 > -1);
|
||||
}
|
||||
|
||||
// fp8->fp32
|
||||
static constexpr MicroAPI::CastTrait castTraitFp8_1 = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN,
|
||||
MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
|
||||
// fp32->fp16
|
||||
static constexpr MicroAPI::CastTrait castTraitFp8_3 = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::NO_SAT,
|
||||
MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_RINT};
|
||||
|
||||
// int8->half
|
||||
static constexpr MicroAPI::CastTrait castTraitint8_1 = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN,
|
||||
MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
|
||||
// half->fp32
|
||||
static constexpr MicroAPI::CastTrait castTraithalf_1 = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN,
|
||||
MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
|
||||
template <typename Q_T, typename KV_T>
|
||||
__simd_vf__ void AntiquantVFImplFp8D448(__ubuf__ int8_t* ubSrcAddr, __ubuf__ Q_T* ubDstAddr, // output first
|
||||
__ubuf__ float* ubScaleSrcAddr, uint32_t dealRowCount)
|
||||
{
|
||||
uint32_t combineDim = 672; // 128对齐 640->672
|
||||
MicroAPI::RegTensor<KV_T> vKvData0;
|
||||
MicroAPI::RegTensor<KV_T> vKvData1;
|
||||
MicroAPI::RegTensor<half> vKvDataHalf0;
|
||||
MicroAPI::RegTensor<half> vKvDataHalf1;
|
||||
MicroAPI::RegTensor<half> vCastHalfRes0;
|
||||
MicroAPI::RegTensor<half> vCastHalfRes1;
|
||||
MicroAPI::RegTensor<float> vCastFp32Res0;
|
||||
MicroAPI::RegTensor<float> vCastFp32Res1;
|
||||
MicroAPI::RegTensor<float> vMulRes0;
|
||||
MicroAPI::RegTensor<float> vMulRes1;
|
||||
MicroAPI::RegTensor<float> vScale0;
|
||||
MicroAPI::RegTensor<float> vScale1;
|
||||
MicroAPI::RegTensor<Q_T> vCastRes0;
|
||||
MicroAPI::RegTensor<Q_T> vCastRes1;
|
||||
MicroAPI::RegTensor<Q_T> vCastResPack0;
|
||||
MicroAPI::RegTensor<Q_T> vCastResPack1;
|
||||
|
||||
MicroAPI::MaskReg kvTypeMaskAll = MicroAPI::CreateMask<KV_T, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg kvRopeTypeMaskAll = MicroAPI::CreateMask<Q_T, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg int8MaskAll = MicroAPI::CreateMask<half, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg fp32MaskAll = MicroAPI::CreateMask<float, MicroAPI::MaskPattern::ALL>();
|
||||
uint32_t blockStride = 17; // +1 to solve bank conflict
|
||||
uint32_t repeatStride = 1;
|
||||
const uint32_t nopeDim = 512; // 448->512 64
|
||||
const uint32_t kvNumPerLoop = 128;
|
||||
const uint32_t scaleNumPerLoop = 1;
|
||||
const uint32_t tileSize = 128;
|
||||
static constexpr bool isKvInt8 = (IsSameType<KV_T, int8_t>::value);
|
||||
// tilesize is 128, deal 128 b8 kv, deal 1 fp32 scale
|
||||
for (uint16_t j = 0; j < (nopeDim / kvNumPerLoop); j++) {
|
||||
__ubuf__ int8_t* ubSrcTemp = ubSrcAddr + j * kvNumPerLoop;
|
||||
__ubuf__ float* ubScaleSrcAddrTemp = ubScaleSrcAddr + j * scaleNumPerLoop;
|
||||
__ubuf__ Q_T* ubDstAddrTmp = ubDstAddr + j * kvNumPerLoop * blockStride;
|
||||
for (uint16_t i = 0; i < static_cast<uint16_t>(dealRowCount); i++) {
|
||||
// load scale
|
||||
MicroAPI::LoadAlign<int8_t, MicroAPI::PostLiteral::POST_MODE_UPDATE, MicroAPI::LoadDist::DIST_UNPACK4_B8>(
|
||||
(MicroAPI::RegTensor<int8_t>&)vKvData0, ubSrcTemp, tileSize / 2);
|
||||
MicroAPI::LoadAlign<int8_t, MicroAPI::PostLiteral::POST_MODE_UPDATE, MicroAPI::LoadDist::DIST_UNPACK4_B8>(
|
||||
(MicroAPI::RegTensor<int8_t>&)vKvData1, ubSrcTemp, combineDim - tileSize / 2);
|
||||
|
||||
MicroAPI::LoadAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE, MicroAPI::LoadDist::DIST_BRC_B32>(
|
||||
(MicroAPI::RegTensor<float>&)vScale0, ubScaleSrcAddrTemp, combineDim / 4);
|
||||
|
||||
if constexpr (isKvInt8) {
|
||||
// int8 -> half
|
||||
MicroAPI::Cast<half, KV_T, castTraitint8_1>(vCastHalfRes0, vKvData0, int8MaskAll);
|
||||
MicroAPI::Cast<half, KV_T, castTraitint8_1>(vCastHalfRes1, vKvData1, int8MaskAll);
|
||||
// half -> float
|
||||
MicroAPI::Cast<float, half, castTraithalf_1>(vCastFp32Res0, vCastHalfRes0, fp32MaskAll);
|
||||
MicroAPI::Cast<float, half, castTraithalf_1>(vCastFp32Res1, vCastHalfRes1, fp32MaskAll);
|
||||
} else {
|
||||
MicroAPI::Cast<float, KV_T, castTraitFp8_1>(vCastFp32Res0, vKvData0, fp32MaskAll);
|
||||
MicroAPI::Cast<float, KV_T, castTraitFp8_1>(vCastFp32Res1, vKvData1, fp32MaskAll);
|
||||
}
|
||||
|
||||
MicroAPI::Mul<float, MicroAPI::MaskMergeMode::ZEROING>(vMulRes0, vCastFp32Res0, vScale0, fp32MaskAll);
|
||||
MicroAPI::Mul<float, MicroAPI::MaskMergeMode::ZEROING>(vMulRes1, vCastFp32Res1, vScale0, fp32MaskAll);
|
||||
|
||||
MicroAPI::Cast<Q_T, float, castTraitFp8_3>(vCastRes0, vMulRes0, fp32MaskAll);
|
||||
MicroAPI::Cast<Q_T, float, castTraitFp8_3>(vCastRes1, vMulRes1, fp32MaskAll);
|
||||
|
||||
MicroAPI::DeInterleave(vCastResPack0, vCastResPack1, vCastRes0, vCastRes1);
|
||||
|
||||
MicroAPI::StoreAlign<Q_T, MicroAPI::DataCopyMode::DATA_BLOCK_COPY, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
ubDstAddrTmp, vCastResPack0, blockStride, repeatStride, kvRopeTypeMaskAll);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename Q_T, typename KV_T>
|
||||
__aicore__ inline void AntiquantVFFp8D448(LocalTensor<Q_T>& outputUb, LocalTensor<KV_T>& inputUb, uint32_t dealRowCount)
|
||||
{
|
||||
__ubuf__ int8_t* ubSrcAddr = (__ubuf__ int8_t*)(inputUb.GetPhyAddr()); // nope改成在左,所以起始位置是0
|
||||
__ubuf__ Q_T* ubDstAddr = (__ubuf__ Q_T*)(outputUb.GetPhyAddr());
|
||||
__ubuf__ float* ubScaleAddr = (__ubuf__ float*)(inputUb[512 + 64 * 2].GetPhyAddr());
|
||||
|
||||
AntiquantVFImplFp8D448<Q_T, KV_T>(ubSrcAddr, ubDstAddr, ubScaleAddr, dealRowCount);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::DequantKv(LocalTensor<Q_T> antiKvTensorAsB16,
|
||||
LocalTensor<KV_T> srcTensor, int64_t dealRow, ConstInfo &constInfo)
|
||||
{
|
||||
// srcTensor是nope(512) + nope(64) + scale + pad, dstTensor是nope(512) + rope(64)
|
||||
AntiquantVFFp8D448<Q_T, KV_T>(antiKvTensorAsB16, srcTensor, dealRow);
|
||||
|
||||
LocalTensor<Q_T> kRopeUb = srcTensor[constInfo.dSizeNope].template ReinterpretCast<Q_T>();
|
||||
LocalTensor<Q_T> kRopeUbNz = antiKvTensorAsB16[constInfo.dSizeNope * (16 + 1)]; // V0单次处理16行数据
|
||||
Copy(kRopeUbNz, kRopeUb,
|
||||
constInfo.dSizeRope, // mask 处理多少列数据
|
||||
static_cast<uint8_t>(dealRow), // repeatTime, 每次处理多少个block
|
||||
{
|
||||
17, // dst stride
|
||||
1, // src stride
|
||||
1, // dst repeat stride
|
||||
21 // src repeat stride, 640 / 32 // 640 -> 672 : 20 -> 21
|
||||
});
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::CopyOutKvUb2L1(
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &outputL1,
|
||||
LocalTensor<Q_T> antiKvTensorAsB16, int64_t dealRow, int64_t s2StartIdx,
|
||||
const RunInfo &runInfo, ConstInfo &constInfo)
|
||||
{
|
||||
uint64_t blockElementNum = 16;
|
||||
DataCopyParams dataCopyParams;
|
||||
dataCopyParams.blockCount = (constInfo.dSizeNope + constInfo.dSizeRope) / blockElementNum;
|
||||
dataCopyParams.blockLen = dealRow;
|
||||
dataCopyParams.srcGap = blockElementNum + 1 - dealRow;
|
||||
dataCopyParams.dstGap = Align16Func(runInfo.s2RealSize) - dealRow;
|
||||
|
||||
LocalTensor<Q_T> dst = outputL1.GetTensor<Q_T>();
|
||||
DataCopy(dst[s2StartIdx * blockElementNum], antiKvTensorAsB16, dataCopyParams);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::CopyOutKvUb2Gm(
|
||||
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm, LocalTensor<Q_T> antiKvTensorAsB16,
|
||||
int64_t dealRow, int64_t s2StartIdx, const RunInfo &runInfo, ConstInfo &constInfo)
|
||||
{
|
||||
GlobalTensor<Q_T> v0ResGmTensor = v0ResGm.template GetTensor<Q_T>();
|
||||
uint64_t blockElementNum = 16;
|
||||
DataCopyParams dataCopyParams;
|
||||
dataCopyParams.blockCount = (constInfo.dSizeNope + constInfo.dSizeRope) / blockElementNum;
|
||||
dataCopyParams.blockLen = dealRow;
|
||||
dataCopyParams.srcGap = blockElementNum + 1 - dealRow;
|
||||
dataCopyParams.dstGap = Align16Func(runInfo.s2RealSize) - dealRow;
|
||||
DataCopy(v0ResGmTensor[s2StartIdx * blockElementNum], antiKvTensorAsB16, dataCopyParams);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::CalSparseCalSize(const RunInfo &runInfo, ConstInfo &constInfo)
|
||||
{
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
uint32_t aicIdx = constInfo.aivIdx >> 1U;
|
||||
uint32_t v0S2SizeFirstCore = CeilDiv(runInfo.s2RealSize, 2);
|
||||
uint32_t v0S2SizeSecondCore = runInfo.s2RealSize - v0S2SizeFirstCore;
|
||||
if (aicIdx % 2U == 0) {
|
||||
if (GetSubBlockIdx() == 0) {
|
||||
qsfaSparseCalSize = CeilDiv(v0S2SizeFirstCore, 2); // 2: Vector split size for first core (first half)
|
||||
sparseS2Start = 0;
|
||||
} else {
|
||||
// 2: Vector split size for first core (second half)
|
||||
qsfaSparseCalSize = v0S2SizeFirstCore - CeilDiv(v0S2SizeFirstCore, 2);
|
||||
sparseS2Start = CeilDiv(v0S2SizeFirstCore, 2); // 2: Start offset for second half of first core
|
||||
}
|
||||
} else {
|
||||
if (GetSubBlockIdx() == 0) {
|
||||
qsfaSparseCalSize = CeilDiv(v0S2SizeSecondCore, 2); // 2: Same as above
|
||||
sparseS2Start = v0S2SizeFirstCore;
|
||||
} else {
|
||||
qsfaSparseCalSize = v0S2SizeSecondCore - CeilDiv(v0S2SizeSecondCore, 2); // 2: Same as above
|
||||
sparseS2Start = v0S2SizeFirstCore + CeilDiv(v0S2SizeSecondCore, 2); // 2: Same as above
|
||||
}
|
||||
}
|
||||
sparseS2End = sparseS2Start + qsfaSparseCalSize;
|
||||
} else {
|
||||
int64_t s2PerVecLoop = 2LL;
|
||||
int64_t vecNum = 2LL;
|
||||
int64_t s2Loops = CeilDiv(CeilDiv(runInfo.s2RealSize, vecNum), s2PerVecLoop);
|
||||
sparseS2Start = GetSubBlockIdx() == 0 ? 0 : s2Loops * s2PerVecLoop;
|
||||
sparseS2End = GetSubBlockIdx() == 0 ? s2Loops * s2PerVecLoop : runInfo.s2RealSize;
|
||||
qsfaSparseCalSize = sparseS2End - sparseS2Start;
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::ProcessVec0(
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &outputL1,
|
||||
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm,
|
||||
const RunInfo &runInfo, ConstInfo &constInfo)
|
||||
{
|
||||
outputL1.WaitCrossCore(); // 核间同步
|
||||
blockSize = constInfo.blockSize;
|
||||
maxBlockNumPerBatch = constInfo.maxBlockNumPerBatch;
|
||||
|
||||
CalSparseCalSize(runInfo, constInfo);
|
||||
ProcessSparseKv(outputL1, v0ResGm, runInfo, constInfo);
|
||||
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
CrossCoreSetFlag<QSFA_SYNC_MODE0, PIPE_MTE3>(15); // 15: 跨核同步标志位值
|
||||
CrossCoreWaitFlag<QSFA_SYNC_MODE0, PIPE_MTE3>(15); // 15: 跨核同步标志位值
|
||||
}
|
||||
|
||||
outputL1.SetCrossCore(); // 核间同步
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
v0ResGm.SetCrossCore();
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::ProcessSparseKv(
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &outputL1,
|
||||
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm, const RunInfo &runInfo, ConstInfo &constInfo)
|
||||
{
|
||||
if (qsfaSparseCalSize == 0) {
|
||||
return;
|
||||
}
|
||||
// Left-closed, right-open interval
|
||||
// 4x = 2x + 2x
|
||||
// 4x + 1 = (2x + 2) + (2x - 1)
|
||||
// 4x + 2 = (2x + 2) + (2x)
|
||||
// 4x + 3 = (2x + 2) + (2x + 1)
|
||||
int64_t s2Start = sparseS2Start;
|
||||
int64_t s2 = sparseS2Start;
|
||||
bool meetEnd = false;
|
||||
int64_t token0Idx, token1Idx; // 拷贝进入的两个token的index
|
||||
// 处理一个s2的base块
|
||||
while ((s2 < sparseS2End) && !meetEnd) { // 拷贝到s2End或者遇到-1
|
||||
int64_t dealRow = 0;
|
||||
// 1、copy kv in, gm ->ub
|
||||
LocalTensor<KV_T> kvInUb = stage0InQue.AllocTensor<KV_T>();
|
||||
while (dealRow < Min(16, qsfaSparseCalSize) && s2<sparseS2End) { // 拷贝满16行或者遇到-1
|
||||
GetRealCmpS2Idx(token0Idx, token1Idx, s2, runInfo, constInfo);
|
||||
s2 += 2; // 每次搬运2行
|
||||
if (token0Idx== -1 && token1Idx == -1) {
|
||||
meetEnd = true;
|
||||
break;
|
||||
}
|
||||
dealRow += CopyInKvSparse(kvInUb, dealRow, token0Idx, token1Idx, runInfo, constInfo);
|
||||
if (token1Idx == -1) {
|
||||
meetEnd = true;
|
||||
break;
|
||||
}
|
||||
}
|
||||
if (dealRow == 0) {
|
||||
stage0InQue.FreeTensor(kvInUb);
|
||||
return;
|
||||
}
|
||||
stage0InQue.EnQue(kvInUb);
|
||||
kvInUb = stage0InQue.DeQue<KV_T>();
|
||||
|
||||
// 2、dequant by vf
|
||||
LocalTensor<Q_T> kvDequantOutUb = stage0OutQue.AllocTensor<Q_T>();
|
||||
DequantKv(kvDequantOutUb, kvInUb, dealRow, constInfo);
|
||||
stage0InQue.FreeTensor(kvInUb);
|
||||
stage0OutQue.EnQue(kvDequantOutUb);
|
||||
kvDequantOutUb = stage0OutQue.DeQue<Q_T>();
|
||||
|
||||
// 3、copy kv out, ub -> l1
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
CopyOutKvUb2Gm(v0ResGm, kvDequantOutUb, dealRow, s2Start, runInfo, constInfo);
|
||||
} else {
|
||||
CopyOutKvUb2L1(outputL1, kvDequantOutUb, dealRow, s2Start, runInfo, constInfo);
|
||||
}
|
||||
s2Start += dealRow;
|
||||
stage0OutQue.FreeTensor(kvDequantOutUb);
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::ProcessVec1(
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &outputBuf,
|
||||
Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &bmm1ResBuf, RunInfo &runInfo,
|
||||
ConstInfo &constInfo)
|
||||
{
|
||||
bmm1ResBuf.WaitCrossCore();
|
||||
|
||||
LocalTensor<float> sumUb = this->softmaxSumBuf[runInfo.multiCoreIdxMod2].template Get<float>();
|
||||
LocalTensor<float> maxUb = this->softmaxMaxBuf[runInfo.multiCoreIdxMod2].template Get<float>();
|
||||
LocalTensor<float> qsfaExpUb = this->softmaxExpBuf[runInfo.taskIdMod2].template Get<T>();
|
||||
int64_t stage1Offset = runInfo.taskIdMod2;
|
||||
auto stage1CastTensor = this->stage1OutQue[stage1Offset].template AllocTensor<Q_T>();
|
||||
|
||||
LocalTensor<T> apiTmpBuffer = this->commonTBuf.template Get<T>();
|
||||
LocalTensor<T> mmRes = bmm1ResBuf.template GetTensor<T>();
|
||||
|
||||
ProcessVec1SoftmaxDispatchQSFA(stage1CastTensor, mmRes, sumUb, maxUb, apiTmpBuffer, runInfo, constInfo);
|
||||
|
||||
bmm1ResBuf.SetCrossCore();
|
||||
// ===================DataCopy to L1 ====================
|
||||
this->stage1OutQue[stage1Offset].template EnQue(stage1CastTensor);
|
||||
this->stage1OutQue[stage1Offset].template DeQue<Q_T>();
|
||||
|
||||
LocalTensor<Q_T> mm2AL1Tensor =
|
||||
outputBuf.GetTensor<Q_T>(s2BaseSize * constInfo.dSizeV);
|
||||
|
||||
if (likely(runInfo.halfMRealSize != 0)) {
|
||||
DataCopy(mm2AL1Tensor[constInfo.subBlockIdx * (BLOCK_BYTE / sizeof(Q_T)) * (runInfo.mRealSize - runInfo.halfMRealSize)],
|
||||
stage1CastTensor, {s2BaseSize / 16, (uint16_t)runInfo.halfMRealSize,
|
||||
(uint16_t)(vec1Srcstride - runInfo.halfMRealSize),
|
||||
(uint16_t)(Align16Func(runInfo.mRealSize) - runInfo.halfMRealSize)});
|
||||
}
|
||||
|
||||
this->stage1OutQue[stage1Offset].template FreeTensor(stage1CastTensor);
|
||||
|
||||
outputBuf.SetCrossCore();
|
||||
if (runInfo.s2LoopCount != 0) {
|
||||
SFAUpdateExpSumAndExpMax<T>(sumUb, maxUb, qsfaExpUb, sumUb, maxUb, apiTmpBuffer, runInfo.halfMRealSize);
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::ProcessVec1SoftmaxDispatchQSFA(
|
||||
LocalTensor<Q_T> &stage1CastTensor, LocalTensor<T> &mmRes, LocalTensor<float> &sumUb,
|
||||
LocalTensor<float> &maxUb, LocalTensor<T> &apiTmpBuffer, RunInfo &runInfo,
|
||||
ConstInfo &constInfo)
|
||||
{
|
||||
if (runInfo.s2LoopCount == 0) {
|
||||
if (likely(runInfo.s2RealSize == 128)) { // s2RealSize等于128分档, VF内常量化减少if判断
|
||||
ProcessVec1Vf<T, Q_T, false, s1BaseSize, s2BaseSize, FaVectorApi::OriginNRange::EQ_128_SFA>(
|
||||
stage1CastTensor, mmRes, sumUb, maxUb, maxUb, apiTmpBuffer, runInfo.halfMRealSize, runInfo.s2RealSize,
|
||||
static_cast<T>(constInfo.softmaxScale), negativeFloatScalar);
|
||||
} else if (runInfo.s2RealSize <= 64) { // s2RealSize小于等于64分档, VF内常量化减少if判断
|
||||
ProcessVec1Vf<T, Q_T, false, s1BaseSize, s2BaseSize, FaVectorApi::OriginNRange::GT_0_AND_LTE_64_SFA>(
|
||||
stage1CastTensor, mmRes, sumUb, maxUb, maxUb, apiTmpBuffer, runInfo.halfMRealSize, runInfo.s2RealSize,
|
||||
static_cast<T>(constInfo.softmaxScale), negativeFloatScalar);
|
||||
} else if (runInfo.s2RealSize < 128) { // s2RealSize小于128分档, VF内常量化减少if判断
|
||||
ProcessVec1Vf<T, Q_T, false, s1BaseSize, s2BaseSize, FaVectorApi::OriginNRange::GT_64_AND_LTE_128_SFA>(
|
||||
stage1CastTensor, mmRes, sumUb, maxUb, maxUb, apiTmpBuffer, runInfo.halfMRealSize,
|
||||
runInfo.s2RealSize, static_cast<T>(constInfo.softmaxScale), negativeFloatScalar);
|
||||
}
|
||||
} else {
|
||||
if (likely(runInfo.s2RealSize == 128)) { // s2RealSize等于128分档, VF内常量化减少if判断
|
||||
ProcessVec1Vf<T, Q_T, true, s1BaseSize, s2BaseSize, FaVectorApi::OriginNRange::EQ_128_SFA>(
|
||||
stage1CastTensor, mmRes, sumUb, maxUb, maxUb, apiTmpBuffer, runInfo.halfMRealSize,
|
||||
runInfo.s2RealSize, static_cast<T>(constInfo.softmaxScale), negativeFloatScalar);
|
||||
} else if (runInfo.s2RealSize <= 64) { // s2RealSize小于等于64分档, VF内常量化减少if判断
|
||||
ProcessVec1Vf<T, Q_T, true, s1BaseSize, s2BaseSize, FaVectorApi::OriginNRange::GT_0_AND_LTE_64_SFA>(
|
||||
stage1CastTensor, mmRes, sumUb, maxUb, maxUb, apiTmpBuffer, runInfo.halfMRealSize,
|
||||
runInfo.s2RealSize, static_cast<T>(constInfo.softmaxScale), negativeFloatScalar);
|
||||
} else if (runInfo.s2RealSize < 128) { // s2RealSize小于128分档, VF内常量化减少if判断
|
||||
ProcessVec1Vf<T, Q_T, true, s1BaseSize, s2BaseSize, FaVectorApi::OriginNRange::GT_64_AND_LTE_128_SFA>(
|
||||
stage1CastTensor, mmRes, sumUb, maxUb, maxUb, apiTmpBuffer, runInfo.halfMRealSize,
|
||||
runInfo.s2RealSize, static_cast<T>(constInfo.softmaxScale), negativeFloatScalar);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::ProcessVec2(
|
||||
Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &bmm2ResBuf, RunInfo &runInfo,
|
||||
ConstInfo &constInfo)
|
||||
{
|
||||
bmm2ResBuf.WaitCrossCore();
|
||||
|
||||
if (unlikely(runInfo.vec2MBaseSize == 0)) {
|
||||
bmm2ResBuf.SetCrossCore();
|
||||
return;
|
||||
}
|
||||
runInfo.vec2MRealSize = runInfo.vec2MBaseSize;
|
||||
runInfo.vec2S1RealSize = runInfo.vec2S1BaseSize;
|
||||
int64_t qsfaVec2CalcSize = runInfo.vec2MRealSize * qsfaDTemplateAlign64;
|
||||
|
||||
LocalTensor<T> vec2ResUb = this->stage2OutBuf.template Get<T>();
|
||||
LocalTensor<T> mmRes = bmm2ResBuf.template GetTensor<T>();
|
||||
|
||||
WaitFlag<HardEvent::MTE3_V>(mte3ToVId[0]);
|
||||
if (unlikely(runInfo.s2LoopCount == 0)) {
|
||||
DataCopy(vec2ResUb, mmRes, qsfaVec2CalcSize);
|
||||
} else {
|
||||
LocalTensor<T> qsfaExpUb = softmaxExpBuf[runInfo.taskIdMod2].template Get<T>();
|
||||
if (runInfo.s2LoopCount < runInfo.s2LoopLimit) {
|
||||
FlashUpdateNew<T, Q_T, OUTPUT_T, qsfaDTemplateAlign64, false, false>(
|
||||
vec2ResUb, mmRes, vec2ResUb, qsfaExpUb, qsfaExpUb, runInfo.vec2MRealSize,
|
||||
qsfaDTemplateAlign64, 1.0, 1.0);
|
||||
} else {
|
||||
LocalTensor<float> sumUb = this->softmaxSumBuf[runInfo.multiCoreIdxMod2].template Get<float>();
|
||||
FlashUpdateLastNew<T, Q_T, OUTPUT_T, qsfaDTemplateAlign64, false, false>(
|
||||
vec2ResUb, mmRes, vec2ResUb, qsfaExpUb, qsfaExpUb, sumUb, runInfo.vec2MRealSize,
|
||||
qsfaDTemplateAlign64, 1.0, 1.0);
|
||||
}
|
||||
}
|
||||
|
||||
bmm2ResBuf.SetCrossCore();
|
||||
if (runInfo.s2LoopCount == runInfo.s2LoopLimit) {
|
||||
if (unlikely(runInfo.s2LoopCount == 0)) {
|
||||
LocalTensor<float> sumUb = this->softmaxSumBuf[runInfo.multiCoreIdxMod2].template Get<float>();
|
||||
LastDivNew<T, Q_T, OUTPUT_T, qsfaDTemplateAlign64, false>(
|
||||
vec2ResUb, vec2ResUb, sumUb, runInfo.vec2MRealSize, qsfaDTemplateAlign64, 1.0);
|
||||
}
|
||||
|
||||
this->CopyOutAttentionOut(runInfo, constInfo, vec2ResUb, 0, qsfaVec2CalcSize);
|
||||
}
|
||||
SetFlag<HardEvent::MTE3_V>(mte3ToVId[0]);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
template <typename VEC2_RES_T>
|
||||
__aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::Bmm2DataCopyOut (RunInfo &runInfo, ConstInfo &constInfo,
|
||||
LocalTensor<VEC2_RES_T> &vec2ResUb, int64_t vec2S1Idx, int64_t qsfaVec2CalcSize)
|
||||
{
|
||||
LocalTensor<OUTPUT_T> attenOut;
|
||||
int64_t dSizeAligned64 = (int64_t)qsfaDTemplateAlign64;
|
||||
|
||||
attenOut.SetAddr(vec2ResUb.address_);
|
||||
Cast(attenOut, vec2ResUb, RoundMode::CAST_ROUND, qsfaVec2CalcSize);
|
||||
SetFlag<HardEvent::V_MTE3>(vToMte3Id[0]);
|
||||
WaitFlag<HardEvent::V_MTE3>(vToMte3Id[0]);
|
||||
|
||||
DataCopyExtParams dataCopyParams;
|
||||
dataCopyParams.blockLen = constInfo.dSizeV * sizeof(OUTPUT_T);
|
||||
dataCopyParams.srcStride = (dSizeAligned64 - constInfo.dSizeV) >> 4; // 以32B为单位偏移,bf16类型即偏移16个数,右移4
|
||||
dataCopyParams.dstStride = constInfo.attentionOutStride;
|
||||
dataCopyParams.blockCount = runInfo.vec2MRealSize;
|
||||
|
||||
DataCopyPad(this->attentionOutGm[runInfo.attentionOutOffset], attenOut, dataCopyParams);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
template <typename VEC2_RES_T>
|
||||
__aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::CopyOutAttentionOut(
|
||||
RunInfo &runInfo, ConstInfo &constInfo, LocalTensor<VEC2_RES_T> &vec2ResUb,
|
||||
int64_t vec2S1Idx, int64_t qsfaVec2CalcSize)
|
||||
{
|
||||
this->Bmm2DataCopyOut(runInfo, constInfo, vec2ResUb, vec2S1Idx, qsfaVec2CalcSize);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::InitOutputSingleCore(ConstInfo &constInfo)
|
||||
{
|
||||
uint32_t coreNum = GetBlockNum();
|
||||
uint64_t totalOutputSize = 0;
|
||||
|
||||
// n2 = 1, n1 = gn2 = gSize
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::BSND) {
|
||||
totalOutputSize = constInfo.bSize * constInfo.gSize * constInfo.s1Size * constInfo.dSizeV;
|
||||
} else if constexpr (LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
totalOutputSize = constInfo.s1Size * constInfo.gSize * constInfo.dSizeV;
|
||||
}
|
||||
|
||||
if (coreNum != 0) {
|
||||
uint64_t singleCoreSize = (totalOutputSize + (CV_RATIO * coreNum) - 1) / (CV_RATIO * coreNum);
|
||||
uint64_t tailSize = totalOutputSize - constInfo.aivIdx * singleCoreSize;
|
||||
uint64_t singleInitOutputSize = tailSize < singleCoreSize ? tailSize : singleCoreSize;
|
||||
if (singleInitOutputSize > 0) {
|
||||
matmul::InitOutput<OUTPUT_T>(this->attentionOutGm[constInfo.aivIdx * singleCoreSize], singleInitOutputSize, 0);
|
||||
}
|
||||
}
|
||||
SyncAll();
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::CleanOutput(__gm__ uint8_t *attentionOut, ConstInfo &constInfo)
|
||||
{
|
||||
if ASCEND_IS_AIV {
|
||||
this->attentionOutGm.SetGlobalBuffer((__gm__ OUTPUT_T *)attentionOut);
|
||||
if (constInfo.needInit == 1) {
|
||||
InitOutputSingleCore(constInfo);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::InitGlobalBuffer(__gm__ uint8_t *key,
|
||||
__gm__ uint8_t *value, __gm__ uint8_t *sparseIndices, __gm__ uint8_t *blockTable)
|
||||
{
|
||||
keyGm.SetGlobalBuffer((__gm__ KV_T *)(key));
|
||||
SparseIndicesGm.SetGlobalBuffer((__gm__ int32_t *)sparseIndices);
|
||||
if constexpr (isPa) {
|
||||
blockTableGm.SetGlobalBuffer((__gm__ int32_t *)blockTable);
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::SoftmaxInitBuffer()
|
||||
{
|
||||
constexpr uint32_t softmaxBufSize = 256; // VF单次操作256Byte
|
||||
tPipe->InitBuffer(softmaxSumBuf[0], softmaxBufSize);
|
||||
tPipe->InitBuffer(softmaxSumBuf[1], softmaxBufSize);
|
||||
tPipe->InitBuffer(softmaxMaxBuf[0], softmaxBufSize);
|
||||
tPipe->InitBuffer(softmaxMaxBuf[1], softmaxBufSize);
|
||||
tPipe->InitBuffer(softmaxExpBuf[0], softmaxBufSize);
|
||||
tPipe->InitBuffer(softmaxExpBuf[1], softmaxBufSize);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::InitLocalBuffer(TPipe *pipe, ConstInfo &constInfo)
|
||||
{
|
||||
// ub buffer
|
||||
SoftmaxInitBuffer();
|
||||
|
||||
tPipe->InitBuffer(commonTBuf, 512); // commonTBuf内存申请512B
|
||||
tPipe->InitBuffer(stage0InQue, 2, dVTemplateTypeInput * 16 * sizeof(KV_T)); // V0阶段每次处理16个seq, 开2 buffer
|
||||
// 576: 模型特征维度(dSize)
|
||||
tPipe->InitBuffer(stage0OutQue, 2, 576 * (16 + 1) * sizeof(Q_T)); // kv输入D轴640, V0阶段每次处理16个seq, 开2 buffer
|
||||
|
||||
tPipe->InitBuffer(stage1OutQue[0], 1, vec1Srcstride * s2BaseSize * sizeof(Q_T));
|
||||
tPipe->InitBuffer(stage1OutQue[1], 1, vec1Srcstride * s2BaseSize * sizeof(Q_T));
|
||||
tPipe->InitBuffer(stage2OutBuf, (s1BaseSize / CV_RATIO) * qsfaDTemplateAlign64 * sizeof(T));
|
||||
|
||||
mte3ToVId[0] = GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>();
|
||||
mte3ToVId[1] = GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>();
|
||||
|
||||
vToMte3Id[0] = GetTPipePtr()->AllocEventID<HardEvent::V_MTE3>();
|
||||
vToMte3Id[1] = GetTPipePtr()->AllocEventID<HardEvent::V_MTE3>();
|
||||
SetFlag<HardEvent::MTE3_V>(mte3ToVId[0]);
|
||||
SetFlag<HardEvent::MTE3_V>(mte3ToVId[1]);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::InitCubeVecSharedParams(
|
||||
CVSharedParams &sharedParams, int32_t aicIdx, uint8_t subBlockIdx)
|
||||
{
|
||||
auto &sparseAttnSharedkvBaseParams = this->tilingData->baseParams;
|
||||
sharedParams.bSize = sparseAttnSharedkvBaseParams.batchSize;
|
||||
sharedParams.n2Size = 1;
|
||||
sharedParams.s1Size = sparseAttnSharedkvBaseParams.qSeqSize;
|
||||
sharedParams.s2Size = sparseAttnSharedkvBaseParams.seqSize;
|
||||
sharedParams.gSize = sparseAttnSharedkvBaseParams.nNumOfQInOneGroup;
|
||||
|
||||
sharedParams.sparseBlockCount = sparseAttnSharedkvBaseParams.sparseBlockCount;
|
||||
sharedParams.maskMode = sparseAttnSharedkvBaseParams.sparseMode;
|
||||
sharedParams.layoutType = sparseAttnSharedkvBaseParams.outputLayout;
|
||||
sharedParams.dSizeRope = 64; // 64: 编码维度
|
||||
sharedParams.softmaxScale = sparseAttnSharedkvBaseParams.scaleValue;
|
||||
sharedParams.dSize = 576; // 576: 模型特征维度(dSize)
|
||||
sharedParams.dSizeVInput = sparseAttnSharedkvBaseParams.dSizeVInput;
|
||||
sharedParams.usedCoreNum = this->tilingData->singleCoreParams.usedCoreNum;
|
||||
if constexpr (isPa) {
|
||||
sharedParams.blockSize = sparseAttnSharedkvBaseParams.blockSize;
|
||||
sharedParams.maxBlockNumPerBatch = sparseAttnSharedkvBaseParams.maxBlockNumPerBatch;
|
||||
}
|
||||
|
||||
sharedParams.isActualSeqLengthsNull = sparseAttnSharedkvBaseParams.isActualLenDimsNull;
|
||||
sharedParams.isActualSeqLengthsKVNull = sparseAttnSharedkvBaseParams.isActualLenDimsKVNull;
|
||||
|
||||
ComputeNeedInitQSFA(sharedParams);
|
||||
|
||||
if ASCEND_IS_AIV {
|
||||
if (subBlockIdx == 0) {
|
||||
auto qsfaTempTilingSSbuf = reinterpret_cast<__ssbuf__ uint32_t*>(0); // 从ssbuf的0地址开始拷贝
|
||||
auto tempTiling = reinterpret_cast<uint32_t *>(&sharedParams);
|
||||
|
||||
#pragma unroll
|
||||
for (int i = 0; i < sizeof(CVSharedParams) / sizeof(uint32_t); ++i, ++qsfaTempTilingSSbuf, ++tempTiling) {
|
||||
*qsfaTempTilingSSbuf = *tempTiling;
|
||||
}
|
||||
|
||||
CrossCoreSetFlag<SYNC_MODE, PIPE_S>(15);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::ComputeNeedInitQSFA(
|
||||
CVSharedParams &sharedParams) const
|
||||
{
|
||||
sharedParams.needInit = 0;
|
||||
for (uint32_t bIdx = 0; bIdx < sharedParams.bSize; bIdx++) {
|
||||
int64_t s2Size;
|
||||
if constexpr (KV_LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
s2Size = (bIdx == 0) ? actualSeqLengthsKVGm.GetValue(bIdx) : \
|
||||
(actualSeqLengthsKVGm.GetValue(bIdx) - actualSeqLengthsKVGm.GetValue(bIdx - 1));
|
||||
} else {
|
||||
if (sharedParams.isActualSeqLengthsKVNull) {
|
||||
s2Size = sharedParams.s2Size;
|
||||
} else {
|
||||
s2Size = actualSeqLengthsKVGm.GetValue(bIdx);
|
||||
}
|
||||
}
|
||||
|
||||
int64_t s1Size;
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
s1Size = (bIdx == 0) ? cuSeqlensQGm.GetValue(bIdx) : \
|
||||
(cuSeqlensQGm.GetValue(bIdx) - cuSeqlensQGm.GetValue(bIdx - 1));
|
||||
} else {
|
||||
if (sharedParams.isActualSeqLengthsNull) {
|
||||
s1Size = sharedParams.s1Size;
|
||||
} else {
|
||||
s1Size = cuSeqlensQGm.GetValue(bIdx);
|
||||
}
|
||||
}
|
||||
if (s1Size > s2Size || (LAYOUT_T == QSFA_LAYOUT::BSND && s1Size < sharedParams.s1Size)) {
|
||||
sharedParams.needInit = 1;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::GetExtremeValue(
|
||||
T &negativeScalar)
|
||||
{
|
||||
uint32_t tmp1 = NEGATIVE_MIN_VALUE_FP32;
|
||||
negativeScalar = *((float *)&tmp1);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF class QSFAVectorServiceDummy {
|
||||
public:
|
||||
__aicore__ inline QSFAVectorServiceDummy() {};
|
||||
__aicore__ inline void CleanOutput(__gm__ uint8_t *attentionOut, ConstInfo &constInfo) {}
|
||||
__aicore__ inline void InitGlobalBuffer(__gm__ uint8_t *key, __gm__ uint8_t *value, __gm__ uint8_t *sparseIndices,
|
||||
__gm__ uint8_t *blockTable) {}
|
||||
__aicore__ inline void InitVecBlock(TPipe *pipe, const KvQuantSparseFlashAttentionTilingDataMla *__restrict tiling,
|
||||
CVSharedParams &sharedParams, int32_t aicIdx, uint8_t subBlockIdx, __gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengths) {};
|
||||
__aicore__ inline void InitLocalBuffer(TPipe *pipe, ConstInfo &constInfo) {}
|
||||
|
||||
__aicore__ inline void ProcessVec1(Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &outputBuf,
|
||||
Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &bmm1ResBuf,
|
||||
RunInfo &runInfo,
|
||||
ConstInfo &constInfo) {}
|
||||
using mm2ResPos = Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH>;
|
||||
__aicore__ inline void ProcessVec2(mm2ResPos &bmm2ResBuf, RunInfo &runInfo,
|
||||
ConstInfo &constInfo) {}
|
||||
};
|
||||
}
|
||||
#endif // KV_QUANT_SPARSE_FLASH_ATTENTION_SERVICE_VECTOR_MLA_H
|
||||
@@ -0,0 +1,147 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_flash_attention.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "kv_quant_sparse_flash_attention_template_tiling_key.h"
|
||||
#if (__CCE_AICORE__ == 310)
|
||||
#include "arch35/kv_quant_sparse_flash_attention_kernel_mla.h"
|
||||
#else
|
||||
#include "kv_quant_sparse_flash_attention_kernel_mla.h"
|
||||
#endif
|
||||
|
||||
using namespace AscendC;
|
||||
|
||||
#if (__CCE_AICORE__ == 310)
|
||||
#if defined(__DAV_C310_CUBE__)
|
||||
#define QSFA_OP_IMPL(templateClass, tilingdataClass, ...) \
|
||||
do { \
|
||||
using CubeBlockType = typename std::conditional<g_coreType == AscendC::AIC, \
|
||||
BaseApi::QSFAMatmulService<__VA_ARGS__>, BaseApi::QSFAMatmulServiceDummy<__VA_ARGS__>>::type; \
|
||||
using VecBlockType = typename std::conditional<g_coreType == AscendC::AIC, \
|
||||
BaseApi::QSFAVectorServiceDummy<__VA_ARGS__>, BaseApi::QSFAVectorService<__VA_ARGS__>>::type; \
|
||||
templateClass<CubeBlockType, VecBlockType> op; \
|
||||
op.Init(query, key, value, sparseIndices, keyScale, valueScale, blocktable, \
|
||||
actualSeqLengthsQuery, actualSeqLengthsKV, \
|
||||
attentionOut, user, nullptr, &tPipe); \
|
||||
op.Process(); \
|
||||
} while (0)
|
||||
#else
|
||||
#define QSFA_OP_IMPL(templateClass, tilingdataClass, ...) \
|
||||
do { \
|
||||
using CubeBlockType = typename std::conditional<g_coreType == AscendC::AIC, \
|
||||
BaseApi::QSFAMatmulService<__VA_ARGS__>, BaseApi::QSFAMatmulServiceDummy<__VA_ARGS__>>::type; \
|
||||
using VecBlockType = typename std::conditional<g_coreType == AscendC::AIC, \
|
||||
BaseApi::QSFAVectorServiceDummy<__VA_ARGS__>, BaseApi::QSFAVectorService<__VA_ARGS__>>::type; \
|
||||
templateClass<CubeBlockType, VecBlockType> op; \
|
||||
GET_TILING_DATA_WITH_STRUCT(tilingdataClass, tilingDataIn, tiling); \
|
||||
const tilingdataClass *__restrict tilingData = &tilingDataIn; \
|
||||
op.Init(query, key, value, sparseIndices, keyScale, valueScale, blocktable, \
|
||||
actualSeqLengthsQuery, actualSeqLengthsKV, \
|
||||
attentionOut, user, tilingData, &tPipe); \
|
||||
op.Process(); \
|
||||
} while (0)
|
||||
#endif
|
||||
#else
|
||||
#define QSFA_OP_IMPL(templateClass, tilingdataClass, ...) \
|
||||
do { \
|
||||
templateClass<QSFAType<__VA_ARGS__>> op; \
|
||||
GET_TILING_DATA_WITH_STRUCT(tilingdataClass, tiling_data_in, tiling); \
|
||||
const tilingdataClass *__restrict tiling_data = &tiling_data_in; \
|
||||
op.Init(query, key, value, sparseIndices, keyScale, valueScale, blocktable, \
|
||||
actualSeqLengthsQuery, actualSeqLengthsKV, \
|
||||
attentionOut, softmaxMax, softmaxSum, user, tiling_data, tiling, &tPipe); \
|
||||
op.Process(); \
|
||||
} while (0)
|
||||
#endif
|
||||
|
||||
#if (__CCE_AICORE__ == 310)
|
||||
template<int FLASH_DECODE, int PAGE_ATTENTION, int LAYOUT_T, int KV_LAYOUT_T, int TEMPLATE_MODE, int IS_SPLIT_G>
|
||||
__aicore__ inline void DispatchKernelDtype310(
|
||||
__gm__ uint8_t *query, __gm__ uint8_t *key, __gm__ uint8_t *value,
|
||||
__gm__ uint8_t *sparseIndices, __gm__ uint8_t *keyScale, __gm__ uint8_t *valueScale,
|
||||
__gm__ uint8_t *blocktable, __gm__ uint8_t *actualSeqLengthsQuery,
|
||||
__gm__ uint8_t *actualSeqLengthsKV, __gm__ uint8_t *attentionOut,
|
||||
__gm__ uint8_t *user, __gm__ uint8_t *tiling, TPipe &tPipe)
|
||||
{
|
||||
if constexpr (ORIG_DTYPE_QUERY == DT_BF16 && ORIG_DTYPE_KEY == DT_FLOAT8_E4M3FN &&
|
||||
ORIG_DTYPE_ATTENTION_OUT == DT_BF16) {
|
||||
QSFA_OP_IMPL(BaseApi::KvQuantSparseFlashAttentionMla, KvQuantSparseFlashAttentionTilingDataMla,
|
||||
bfloat16_t, fp8_e4m3fn_t, float, bfloat16_t, FLASH_DECODE, PAGE_ATTENTION,
|
||||
static_cast<QSFA_LAYOUT>(LAYOUT_T), static_cast<QSFA_LAYOUT>(KV_LAYOUT_T),
|
||||
static_cast<QSFATemplateMode>(TEMPLATE_MODE), IS_SPLIT_G);
|
||||
} else if constexpr (ORIG_DTYPE_QUERY == DT_BF16 && ORIG_DTYPE_KEY == DT_HIFLOAT8 &&
|
||||
ORIG_DTYPE_ATTENTION_OUT == DT_BF16) {
|
||||
QSFA_OP_IMPL(BaseApi::KvQuantSparseFlashAttentionMla, KvQuantSparseFlashAttentionTilingDataMla,
|
||||
bfloat16_t, hifloat8_t, float, bfloat16_t, FLASH_DECODE, PAGE_ATTENTION,
|
||||
static_cast<QSFA_LAYOUT>(LAYOUT_T), static_cast<QSFA_LAYOUT>(KV_LAYOUT_T),
|
||||
static_cast<QSFATemplateMode>(TEMPLATE_MODE), IS_SPLIT_G);
|
||||
} else if constexpr (ORIG_DTYPE_QUERY == DT_BF16 && ORIG_DTYPE_KEY == DT_INT8 &&
|
||||
ORIG_DTYPE_ATTENTION_OUT == DT_BF16) {
|
||||
QSFA_OP_IMPL(BaseApi::KvQuantSparseFlashAttentionMla, KvQuantSparseFlashAttentionTilingDataMla,
|
||||
bfloat16_t, int8_t, float, bfloat16_t, FLASH_DECODE, PAGE_ATTENTION,
|
||||
static_cast<QSFA_LAYOUT>(LAYOUT_T), static_cast<QSFA_LAYOUT>(KV_LAYOUT_T),
|
||||
static_cast<QSFATemplateMode>(TEMPLATE_MODE), IS_SPLIT_G);
|
||||
} else if constexpr (ORIG_DTYPE_QUERY == DT_FLOAT16 && ORIG_DTYPE_KEY == DT_FLOAT8_E4M3FN &&
|
||||
ORIG_DTYPE_ATTENTION_OUT == DT_FLOAT16) {
|
||||
QSFA_OP_IMPL(BaseApi::KvQuantSparseFlashAttentionMla, KvQuantSparseFlashAttentionTilingDataMla,
|
||||
half, fp8_e4m3fn_t, float, half, FLASH_DECODE, PAGE_ATTENTION,
|
||||
static_cast<QSFA_LAYOUT>(LAYOUT_T), static_cast<QSFA_LAYOUT>(KV_LAYOUT_T),
|
||||
static_cast<QSFATemplateMode>(TEMPLATE_MODE), IS_SPLIT_G);
|
||||
} else if constexpr (ORIG_DTYPE_QUERY == DT_FLOAT16 && ORIG_DTYPE_KEY == DT_HIFLOAT8 &&
|
||||
ORIG_DTYPE_ATTENTION_OUT == DT_FLOAT16) {
|
||||
QSFA_OP_IMPL(BaseApi::KvQuantSparseFlashAttentionMla, KvQuantSparseFlashAttentionTilingDataMla,
|
||||
half, hifloat8_t, float, half, FLASH_DECODE, PAGE_ATTENTION,
|
||||
static_cast<QSFA_LAYOUT>(LAYOUT_T), static_cast<QSFA_LAYOUT>(KV_LAYOUT_T),
|
||||
static_cast<QSFATemplateMode>(TEMPLATE_MODE), IS_SPLIT_G);
|
||||
} else if constexpr (ORIG_DTYPE_QUERY == DT_FLOAT16 && ORIG_DTYPE_KEY == DT_INT8 &&
|
||||
ORIG_DTYPE_ATTENTION_OUT == DT_FLOAT16) {
|
||||
QSFA_OP_IMPL(BaseApi::KvQuantSparseFlashAttentionMla, KvQuantSparseFlashAttentionTilingDataMla,
|
||||
half, int8_t, float, half, FLASH_DECODE, PAGE_ATTENTION,
|
||||
static_cast<QSFA_LAYOUT>(LAYOUT_T), static_cast<QSFA_LAYOUT>(KV_LAYOUT_T),
|
||||
static_cast<QSFATemplateMode>(TEMPLATE_MODE), IS_SPLIT_G);
|
||||
}
|
||||
}
|
||||
#endif
|
||||
|
||||
template<int FLASH_DECODE, int PAGE_ATTENTION, int LAYOUT_T, int KV_LAYOUT_T, int TEMPLATE_MODE, int IS_SPLIT_G>
|
||||
__global__ __aicore__ void
|
||||
kv_quant_sparse_flash_attention(__gm__ uint8_t *query, __gm__ uint8_t *key, __gm__ uint8_t *value,
|
||||
__gm__ uint8_t *sparseIndices, __gm__ uint8_t* keyScale, __gm__ uint8_t* valueScale,
|
||||
__gm__ uint8_t *blocktable, __gm__ uint8_t *actualSeqLengthsQuery,
|
||||
__gm__ uint8_t *actualSeqLengthsKV, __gm__ uint8_t *attentionOut,
|
||||
__gm__ uint8_t *softmaxMax, __gm__ uint8_t *softmaxSum,
|
||||
__gm__ uint8_t *workspace, __gm__ uint8_t *tiling)
|
||||
{
|
||||
KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIC_1_2);
|
||||
|
||||
TPipe tPipe;
|
||||
__gm__ uint8_t *user = GetUserWorkspace(workspace);
|
||||
#if (__CCE_AICORE__ == 310)
|
||||
DispatchKernelDtype310<FLASH_DECODE, PAGE_ATTENTION, LAYOUT_T, KV_LAYOUT_T, TEMPLATE_MODE, IS_SPLIT_G>(
|
||||
query, key, value, sparseIndices, keyScale, valueScale, blocktable,
|
||||
actualSeqLengthsQuery, actualSeqLengthsKV, attentionOut, user, tiling, tPipe);
|
||||
#else
|
||||
if constexpr (ORIG_DTYPE_QUERY == DT_FLOAT16 && ORIG_DTYPE_KEY == DT_INT8 &&
|
||||
ORIG_DTYPE_ATTENTION_OUT == DT_FLOAT16) {
|
||||
QSFA_OP_IMPL(KvQuantSparseFlashAttentionMla, KvQuantSparseFlashAttentionTilingDataMla, half, int8_t,
|
||||
half, FLASH_DECODE, static_cast<QSFA_LAYOUT>(LAYOUT_T), static_cast<QSFA_LAYOUT>(KV_LAYOUT_T),
|
||||
TEMPLATE_MODE);
|
||||
} else { // bf16
|
||||
QSFA_OP_IMPL(KvQuantSparseFlashAttentionMla, KvQuantSparseFlashAttentionTilingDataMla, bfloat16_t, int8_t,
|
||||
bfloat16_t, FLASH_DECODE, static_cast<QSFA_LAYOUT>(LAYOUT_T), static_cast<QSFA_LAYOUT>(KV_LAYOUT_T),
|
||||
TEMPLATE_MODE);
|
||||
}
|
||||
#endif
|
||||
}
|
||||
@@ -0,0 +1,225 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_flash_attention_common.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef KV_QUANT_SPARSE_FLASH_ATTENTION_COMMON_H
|
||||
#define KV_QUANT_SPARSE_FLASH_ATTENTION_COMMON_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "lib/matmul_intf.h"
|
||||
#include "lib/matrix/matmul/tiling.h"
|
||||
|
||||
using namespace AscendC;
|
||||
// 将isCheckTiling设置为false, 输入输出的max&sum&exp的shape为(m, 1)
|
||||
constexpr SoftmaxConfig QSFA_SOFTMAX_FLASHV2_CFG_WITHOUT_BRC = {false, 0, 0, SoftmaxMode::SOFTMAX_OUTPUT_WITHOUT_BRC};
|
||||
|
||||
enum class QSFA_LAYOUT {
|
||||
BSND = 0,
|
||||
TND = 1,
|
||||
PA_BSND = 2,
|
||||
};
|
||||
|
||||
enum class QUANT_MODE {
|
||||
PER_CHANNEL = 0, // GQA支持
|
||||
PER_TOKEN_HEAD = 1, // GQA支持
|
||||
PER_TILE = 2, // MLA支持
|
||||
};
|
||||
|
||||
enum class ATTENTION_MODE {
|
||||
GQA_MHA = 0, // QKV headDim相等
|
||||
MLA_NATIVE = 1, // Dn=128, Dr=64
|
||||
MLA_ABSORB = 2, // Dn=512, Dr=64
|
||||
};
|
||||
|
||||
enum class QUANT_SCALE_REPO_MODE {
|
||||
SEPARATE = 0, // 分开存储
|
||||
COMBINE = 1, // 合并存储,量化模式是PER_TOKEN_HEAD/PER_TILE时支持COMBINE模式,参数顺序为:Nope+Rope+DequantScale
|
||||
};
|
||||
|
||||
template <typename Q_T, typename KV_T, typename OUT_T, const bool FLASH_DECODE = false,
|
||||
QSFA_LAYOUT LAYOUT_T = QSFA_LAYOUT::BSND, QSFA_LAYOUT KV_LAYOUT_T = QSFA_LAYOUT::BSND,
|
||||
const int TEMPLATE_MODE = C_TEMPLATE, typename... Args>
|
||||
struct QSFAType {
|
||||
using queryType = Q_T;
|
||||
using kvType = KV_T;
|
||||
using kRopeType = Q_T;
|
||||
using outputType = OUT_T;
|
||||
static constexpr bool flashDecode = FLASH_DECODE;
|
||||
static constexpr QSFA_LAYOUT layout = LAYOUT_T;
|
||||
static constexpr QSFA_LAYOUT kvLayout = KV_LAYOUT_T;
|
||||
static constexpr int templateMode = TEMPLATE_MODE;
|
||||
static constexpr bool pageAttention = (KV_LAYOUT_T == QSFA_LAYOUT::PA_BSND);
|
||||
};
|
||||
|
||||
// ================================Util functions==================================
|
||||
template <typename T> __aicore__ inline T QSFAAlign(T num, T rnd)
|
||||
{
|
||||
return (((rnd) == 0) ? 0 : (((num) + (rnd) - 1) / (rnd) * (rnd)));
|
||||
}
|
||||
|
||||
template <typename T> __aicore__ inline size_t BlockAlign(size_t s)
|
||||
{
|
||||
if constexpr (IsSameType<T, int4b_t>::value) {
|
||||
return (s + 63) / 64 * 64;
|
||||
}
|
||||
size_t n = (32 / sizeof(T));
|
||||
return (s + n - 1) / n * n;
|
||||
}
|
||||
|
||||
template <typename T1, typename T2> __aicore__ inline T1 Min(T1 a, T2 b)
|
||||
{
|
||||
return (a > b) ? (b) : (a);
|
||||
}
|
||||
|
||||
struct RunInfo {
|
||||
uint32_t loop;
|
||||
uint32_t bIdx;
|
||||
uint32_t gIdx;
|
||||
uint32_t s1Idx;
|
||||
uint32_t s2Idx;
|
||||
uint32_t bn2IdxInCurCore;
|
||||
uint32_t curSInnerLoopTimes;
|
||||
uint32_t s2BatchOffset;
|
||||
|
||||
uint64_t tndBIdxOffsetForQ;
|
||||
uint64_t tndBIdxOffsetForKV;
|
||||
uint64_t tensorAOffset;
|
||||
uint64_t tensorBOffset;
|
||||
uint64_t tensorARopeOffset;
|
||||
uint64_t tensorBRopeOffset;
|
||||
uint64_t attenOutOffset;
|
||||
uint64_t topKBaseOffset;
|
||||
uint64_t attenMaskOffset;
|
||||
|
||||
uint32_t actualSingleProcessSInnerSize;
|
||||
uint32_t actualSingleProcessSInnerSizeAlign;
|
||||
uint32_t gSize;
|
||||
uint32_t s1Size;
|
||||
uint32_t s2Size;
|
||||
uint32_t mSize;
|
||||
uint32_t mSizeV;
|
||||
uint32_t mSizeVStart;
|
||||
uint32_t tndIsS2SplitCore;
|
||||
uint32_t tndCoreStartKVSplitPos;
|
||||
bool isBmm2Output;
|
||||
bool isValid = false;
|
||||
bool isFirstSInnerLoop;
|
||||
bool isChangeBatch;
|
||||
static constexpr uint32_t n2Idx = 0;
|
||||
|
||||
uint64_t actS1Size = 1;
|
||||
uint64_t curActualSeqLenOri = 0ULL;
|
||||
uint64_t actS2Size = 1;
|
||||
|
||||
uint32_t gS1Idx;
|
||||
uint32_t actMBaseSize;
|
||||
int32_t nextTokensPerBatch = 0;
|
||||
bool isLastS2Loop;
|
||||
uint8_t resv[3];
|
||||
int64_t threshold;
|
||||
};
|
||||
|
||||
struct ConstInfo {
|
||||
// CUBE与VEC核间同步的模式
|
||||
static constexpr uint32_t QSFA_SYNC_MODE2 = 2;
|
||||
// BUFFER的字节数
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_1K = 1024;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_2K = 2048;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_4K = 4096;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_8K = 8192;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_16K = 16384;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_32K = 32768;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_32B = 32;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_64B = 64;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_256B = 256;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_512B = 512;
|
||||
// FP32的0值和极大值
|
||||
static constexpr float FLOAT_ZERO = 0;
|
||||
static constexpr float FLOAT_MAX = 3.402823466e+38F;
|
||||
|
||||
// preLoad的总次数
|
||||
uint32_t preLoadNum = 0U;
|
||||
uint32_t nBufferMBaseSize = 0U;
|
||||
// CUBE和VEC的核间同步EventID
|
||||
uint32_t syncV0C1 = 0U;
|
||||
uint32_t syncC1V1 = 0U;
|
||||
uint32_t syncV1C2 = 0U;
|
||||
uint32_t syncC2V2 = 0U;
|
||||
uint32_t syncC2V1 = 0U;
|
||||
uint32_t syncV1NupdateC2 = 0U;
|
||||
|
||||
uint32_t mmResUbSize = 0U; // Matmul1输出结果GM上的大小
|
||||
uint32_t vec1ResUbSize = 0U; // Vector1输出结果GM上的大小
|
||||
uint32_t bmm2ResUbSize = 0U; // Matmul2输出结果GM上的大小
|
||||
uint64_t gSize = 0ULL;
|
||||
uint64_t batchSize = 0ULL;
|
||||
uint64_t qHeadNum = 0ULL;
|
||||
uint64_t kvHeadNum;
|
||||
uint64_t headDim;
|
||||
uint64_t headDimRope;
|
||||
uint64_t combineHeadDim; // quantScaleRepoMode为Combine模式时=headDim+headDimRope, 否则=headDim
|
||||
uint64_t kvSeqSize = 0ULL; // kv最大S长度
|
||||
uint64_t qSeqSize = 1ULL; // q最大S长度
|
||||
int64_t kvCacheBlockSize = 0; // PA场景的block size
|
||||
uint32_t maxBlockNumPerBatch = 0; // PA场景的最大单batch block number
|
||||
uint32_t splitKVNum = 0U; // S2核间切分的切分份数
|
||||
QSFA_LAYOUT outputLayout; // 输出的Transpose格式
|
||||
uint32_t sparseMode = 0;
|
||||
bool returnSoftmaxLse = false;
|
||||
bool needInit = false;
|
||||
|
||||
// FlashDecoding
|
||||
uint64_t combineLseOffset = 0ULL;
|
||||
uint64_t combineAccumOutOffset = 0ULL;
|
||||
uint32_t actualCombineLoopSize = 0U; // FlashDecoding场景, S2在核间切分的最大份数
|
||||
|
||||
uint32_t actualLenDimsQ = 0U; // query的actualSeqLength 的维度
|
||||
uint32_t actualLenDimsKV = 0U; // KV 的actualSeqLength 的维度
|
||||
|
||||
// TND
|
||||
uint32_t s2Start = 0U; // TND场景下,S2的起始位置
|
||||
uint32_t s2End = 0U; // 单核TND场景下S2循环index上限
|
||||
|
||||
uint32_t bN2Start = 0U;
|
||||
uint32_t bN2End = 0U;
|
||||
uint32_t gS1Start = 0U;
|
||||
uint32_t gS1End = 0U;
|
||||
|
||||
uint32_t mBaseSize = 1ULL;
|
||||
uint32_t s2BaseSize = 1ULL;
|
||||
|
||||
uint32_t tndFDCoreArrLen = 0U; // TNDFlashDecoding相关分核信息array的长度
|
||||
uint32_t coreStartKVSplitPos = 0U; // TNDFlashDecoding kv起始位置
|
||||
|
||||
// sparse attr
|
||||
uint32_t sparseBlockCount = 0;
|
||||
int64_t sparseBlockSize = 0;
|
||||
|
||||
// attention模式与量化模式
|
||||
ATTENTION_MODE attentionMode = ATTENTION_MODE::MLA_ABSORB;
|
||||
QUANT_MODE keyQuantMode = QUANT_MODE::PER_TILE;
|
||||
QUANT_MODE valueQuantMode = QUANT_MODE::PER_TILE;
|
||||
QUANT_SCALE_REPO_MODE quantScaleRepoMode = QUANT_SCALE_REPO_MODE::COMBINE;
|
||||
uint64_t tileSize = 128ULL;
|
||||
};
|
||||
|
||||
struct MSplitInfo {
|
||||
uint32_t nBufferIdx = 0U;
|
||||
uint32_t nBufferStartM = 0U;
|
||||
uint32_t nBufferDealM = 0U;
|
||||
uint32_t vecStartM = 0U;
|
||||
uint32_t vecDealM = 0U;
|
||||
};
|
||||
|
||||
#endif // KV_QUANT_SPARSE_FLASH_ATTENTION_COMMON_H
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,943 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_flash_attention_service_cube_mla.h
|
||||
* \brief use 7 buffer for matmul l1, better pipeline
|
||||
*/
|
||||
#ifndef KV_QUANT_SPARSE_FLASH_ATTENTION_SERVICE_CUBE_MLA_H
|
||||
#define KV_QUANT_SPARSE_FLASH_ATTENTION_SERVICE_CUBE_MLA_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "kernel_operator_list_tensor_intf.h"
|
||||
#include "kernel_tiling/kernel_tiling.h"
|
||||
#include "lib/matmul_intf.h"
|
||||
#include "lib/matrix/matmul/tiling.h"
|
||||
#include "kv_quant_sparse_flash_attention_common.h"
|
||||
|
||||
struct Position {
|
||||
uint32_t bIdx;
|
||||
uint32_t n2Idx;
|
||||
uint32_t s2Idx;
|
||||
uint32_t dIdx;
|
||||
};
|
||||
|
||||
struct PAShape {
|
||||
uint32_t blockSize;
|
||||
uint32_t headNum; // 一般为kv的head num,对应n2
|
||||
uint32_t headDim; // mla下rope为64,nope为512, 对应d
|
||||
uint32_t maxblockNumPerBatch; // block table 每一行的最大个数
|
||||
uint32_t actHeadDim; // 实际拷贝col大小,考虑到N切块 s*d, 对应d
|
||||
uint32_t copyRowNum; // 总共要拷贝的行数
|
||||
uint32_t copyRowNumAlign;
|
||||
};
|
||||
|
||||
// 场景:query、queryRope、key、value GM to L1
|
||||
// GM按ND格式存储
|
||||
// L1按NZ格式存储
|
||||
// GM的行、列、列的stride
|
||||
template <typename T>
|
||||
__aicore__ inline void DataCopyGmNDToL1(LocalTensor<T> &l1Tensor, GlobalTensor<T> &gmTensor,
|
||||
uint32_t rowAct, uint32_t rowAlign,
|
||||
uint32_t col, // D
|
||||
uint32_t colStride) // D or N*D
|
||||
{
|
||||
Nd2NzParams nd2nzPara;
|
||||
nd2nzPara.ndNum = 1;
|
||||
nd2nzPara.nValue = rowAct; // nd矩阵的行数
|
||||
// T为int4场景下,dValue = col / 2,srcDValue = colStride / 2
|
||||
nd2nzPara.srcDValue = colStride; // 同一nd矩阵相邻行起始地址间的偏移
|
||||
nd2nzPara.dValue = col; // nd矩阵的列数
|
||||
nd2nzPara.dstNzC0Stride = rowAlign;
|
||||
nd2nzPara.dstNzNStride = 1;
|
||||
nd2nzPara.dstNzMatrixStride = 0;
|
||||
nd2nzPara.srcNdMatrixStride = 0;
|
||||
DataCopy(l1Tensor, gmTensor, nd2nzPara);
|
||||
}
|
||||
|
||||
/*
|
||||
适用PA数据从GM拷贝到L1,支持ND、NZ数据;
|
||||
PA的layout分 BNBD(blockNum,N,blockSize,D) BBH(blockNum,blockSize,N*D
|
||||
BSH\BSND\TND 为BBH
|
||||
shape.copyRowNumAlign 需要16字节对齐,如拷贝k矩阵,一次拷贝128*512,遇到尾块 10*512 需对齐到16*512
|
||||
*/
|
||||
template <typename T, QSFA_LAYOUT SRC_LAYOUT>
|
||||
__aicore__ inline void DataCopyPA(LocalTensor<T> &dstTensor, // l1
|
||||
GlobalTensor<T> &srcTensor, // gm
|
||||
GlobalTensor<int32_t> &blockTableGm,
|
||||
const PAShape &shape, // blockSize, headNum, headDim
|
||||
const Position &startPos) // bacthIdx nIdx curSeqIdx
|
||||
{
|
||||
uint32_t copyFinishRowCnt = 0;
|
||||
uint64_t blockTableBaseOffset = startPos.bIdx * shape.maxblockNumPerBatch;
|
||||
uint32_t curS2Idx = startPos.s2Idx;
|
||||
uint32_t blockElementCnt = 32 / sizeof(T);
|
||||
while (copyFinishRowCnt < shape.copyRowNum) {
|
||||
uint64_t blockIdOffset = curS2Idx / shape.blockSize; // 获取block table上的索引
|
||||
uint64_t reaminRowCnt = curS2Idx % shape.blockSize; // 获取在单个块上超出的行数
|
||||
// 从block table上的获取编号
|
||||
uint64_t idInBlockTable = blockTableGm.GetValue(blockTableBaseOffset + blockIdOffset);
|
||||
// 计算可以拷贝行数
|
||||
uint32_t copyRowCnt = shape.blockSize - reaminRowCnt; // 一次只能处理一个Block
|
||||
if (copyFinishRowCnt + copyRowCnt > shape.copyRowNum) {
|
||||
copyRowCnt = shape.copyRowNum - copyFinishRowCnt; // 一个block未拷满
|
||||
}
|
||||
uint64_t offset = idInBlockTable * shape.blockSize * shape.headNum * shape.headDim ; // PA的偏移
|
||||
|
||||
uint64_t dStride = shape.headDim;
|
||||
if constexpr (SRC_LAYOUT == QSFA_LAYOUT::BSND || SRC_LAYOUT == QSFA_LAYOUT::TND) {
|
||||
offset += (uint64_t)(startPos.n2Idx * shape.headDim) +
|
||||
reaminRowCnt * shape.headDim * shape.headNum + startPos.dIdx;
|
||||
dStride = shape.headDim * shape.headNum;
|
||||
} else {
|
||||
offset += (uint64_t)(startPos.n2Idx * shape.headDim * shape.blockSize) +
|
||||
reaminRowCnt * shape.headDim + startPos.dIdx;
|
||||
}
|
||||
|
||||
uint32_t srcDValue = dStride;
|
||||
uint32_t dValue = shape.actHeadDim;
|
||||
LocalTensor<T> tmpDstTensor = dstTensor[copyFinishRowCnt * blockElementCnt];
|
||||
GlobalTensor<T> tmpSrcTensor = srcTensor[offset];
|
||||
|
||||
DataCopyGmNDToL1<T>(tmpDstTensor, tmpSrcTensor, copyRowCnt, shape.copyRowNumAlign, dValue, srcDValue);
|
||||
copyFinishRowCnt += copyRowCnt;
|
||||
curS2Idx += copyRowCnt;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QSFAT> class QSFAMatmulService {
|
||||
public:
|
||||
// 中间计算数据类型为float, 高精度模式
|
||||
using T = float;
|
||||
using Q_T = typename QSFAT::queryType;
|
||||
using KV_T = typename QSFAT::kvType;
|
||||
using K_ROPE_T = typename QSFAT::kRopeType;
|
||||
using OUT_T = typename QSFAT::outputType;
|
||||
using MM_OUT_T = T;
|
||||
|
||||
__aicore__ inline QSFAMatmulService(){};
|
||||
__aicore__ inline void InitParams(const ConstInfo &constInfo);
|
||||
__aicore__ inline void InitMm1GlobalTensor(GlobalTensor<Q_T> queryGm, GlobalTensor<Q_T> qRopeGm,
|
||||
GlobalTensor<KV_T> keyGm, GlobalTensor<K_ROPE_T> kRopeGm,
|
||||
GlobalTensor<MM_OUT_T> mm1ResGm);
|
||||
__aicore__ inline void InitMm2GlobalTensor(GlobalTensor<K_ROPE_T> vec1ResGm, GlobalTensor<KV_T> valueGm,
|
||||
GlobalTensor<MM_OUT_T> mm2ResGm, GlobalTensor<OUT_T> attentionOutGm);
|
||||
__aicore__ inline void InitPageAttentionInfo(const GlobalTensor<K_ROPE_T>& kvMergeGm,
|
||||
GlobalTensor<int32_t> blockTableGm, GlobalTensor<int32_t> topKGm,
|
||||
uint32_t blockSize, uint32_t maxBlockNumPerBatch);
|
||||
__aicore__ inline void InitBuffers(TPipe *pipe);
|
||||
__aicore__ inline void UpdateKey(GlobalTensor<KV_T> keyGm);
|
||||
__aicore__ inline void UpdateValue(GlobalTensor<KV_T> valueGm);
|
||||
|
||||
__aicore__ inline void AllocEventID();
|
||||
__aicore__ inline void FreeEventID();
|
||||
__aicore__ inline void CalcTopKBlockInfo(const RunInfo &info, uint32_t &curTopKIdx,
|
||||
uint64_t &curOffsetInSparseBlock, uint32_t curSeqIdx,
|
||||
uint32_t ©RowCnt, uint64_t &idInTopK);
|
||||
__aicore__ inline void ComputeMm1(const RunInfo &info, const MSplitInfo mSplitInfo);
|
||||
__aicore__ inline void ComputeMm2(const RunInfo &info, const MSplitInfo mSplitInfo);
|
||||
|
||||
private:
|
||||
static constexpr bool PAGE_ATTENTION = QSFAT::pageAttention;
|
||||
static constexpr int TEMPLATE_MODE = QSFAT::templateMode;
|
||||
static constexpr bool FLASH_DECODE = QSFAT::flashDecode;
|
||||
static constexpr QSFA_LAYOUT LAYOUT_T = QSFAT::layout;
|
||||
static constexpr QSFA_LAYOUT KV_LAYOUT_T = QSFAT::kvLayout;
|
||||
|
||||
static constexpr uint32_t M_SPLIT_SIZE = 128; // m方向切分
|
||||
static constexpr uint32_t N_SPLIT_SIZE = 128; // n方向切分
|
||||
static constexpr uint32_t N_WORKSPACE_SIZE = 512; // n方向切分
|
||||
static constexpr uint32_t K_SPLIT_SIZE = 288; // K方向切分
|
||||
|
||||
static constexpr uint32_t L1_BLOCK_SIZE = (64 * (512 + 64) * sizeof(Q_T));
|
||||
static constexpr uint32_t L1_BLOCK_OFFSET = 64 * (512 + 64); // 72K的元素个数
|
||||
|
||||
static constexpr uint32_t L0A_PP_SIZE = (32 * 1024);
|
||||
static constexpr uint32_t L0B_PP_SIZE = (32 * 1024);
|
||||
static constexpr uint32_t L0C_PP_SIZE = (64 * 1024);
|
||||
|
||||
// m <> mte1 EventID
|
||||
static constexpr uint32_t L0AB_EVENT0 = EVENT_ID3;
|
||||
static constexpr uint32_t L0AB_EVENT1 = EVENT_ID4;
|
||||
|
||||
// mte2 <> mte1 EventID
|
||||
// L1 3buf, 使用3个eventId
|
||||
static constexpr uint32_t L1_EVENT0 = EVENT_ID2;
|
||||
static constexpr uint32_t L1_EVENT1 = EVENT_ID3;
|
||||
static constexpr uint32_t L1_EVENT2 = EVENT_ID4;
|
||||
static constexpr uint32_t L1_EVENT3 = EVENT_ID5;
|
||||
static constexpr uint32_t L1_EVENT4 = EVENT_ID6;
|
||||
static constexpr uint32_t L1_EVENT5 = EVENT_ID7;
|
||||
static constexpr uint32_t L1_EVENT6 = EVENT_ID1;
|
||||
|
||||
static constexpr IsResetLoad3dConfig LOAD3DV2_CONFIG = {true, true}; // isSetFMatrix isSetPadding;
|
||||
static constexpr uint32_t mte21QPIds[4] = {L1_EVENT0, L1_EVENT1, L1_EVENT2, L1_EVENT3}; // mte12复用
|
||||
static constexpr uint32_t mte21KVIds[3] = {L1_EVENT4, L1_EVENT5, L1_EVENT6};
|
||||
|
||||
static constexpr uint32_t BLOCK_ELEMENT_NUM = ConstInfo::BUFFER_SIZE_BYTE_32B / sizeof(K_ROPE_T);
|
||||
|
||||
uint32_t kvCacheBlockSize = 0;
|
||||
uint32_t maxBlockNumPerBatch = 0;
|
||||
ConstInfo constInfo{};
|
||||
|
||||
// L1分成3块buf, 用于记录
|
||||
uint32_t qpL1BufIter = 0;
|
||||
uint32_t kvL1BufIter = -1;
|
||||
uint32_t abL0BufIter = 0;
|
||||
uint32_t cL0BufIter = 0;
|
||||
|
||||
// mm1
|
||||
GlobalTensor<Q_T> queryGm;
|
||||
GlobalTensor<Q_T> qRopeGm;
|
||||
GlobalTensor<KV_T> keyGm;
|
||||
GlobalTensor<K_ROPE_T> kRopeGm;
|
||||
GlobalTensor<MM_OUT_T> mm1ResGm;
|
||||
GlobalTensor<K_ROPE_T> kvMergeGm_;
|
||||
|
||||
// mm2
|
||||
GlobalTensor<K_ROPE_T> vec1ResGm;
|
||||
GlobalTensor<KV_T> valueGm;
|
||||
GlobalTensor<MM_OUT_T> mm2ResGm;
|
||||
GlobalTensor<OUT_T> attentionOutGm;
|
||||
|
||||
// block_table
|
||||
GlobalTensor<int32_t> topKGm;
|
||||
GlobalTensor<int32_t> blockTableGm;
|
||||
|
||||
TBuf<TPosition::A1> bufQPL1;
|
||||
TBuf<TPosition::A1> bufKVL1;
|
||||
TBuf<TPosition::A2> tmpBufL0A;
|
||||
TBuf<TPosition::B2> tmpBufL0B;
|
||||
TBuf<TPosition::CO1> tmpBufL0C;
|
||||
|
||||
LocalTensor<K_ROPE_T> aL0TensorPingPong;
|
||||
LocalTensor<K_ROPE_T> bL0TensorPingPong;
|
||||
LocalTensor<MM_OUT_T> cL0TensorPingPong;
|
||||
LocalTensor<Q_T> l1QPTensor;
|
||||
LocalTensor<Q_T> l1KVTensor;
|
||||
|
||||
// L0AB m <> mte1 EventID
|
||||
__aicore__ inline uint32_t Mte1MmABEventId(uint32_t idx)
|
||||
{
|
||||
return (L0AB_EVENT0 + idx);
|
||||
}
|
||||
|
||||
__aicore__ inline uint32_t GetQPL1RealIdx(uint32_t mIdx, uint32_t k1Idx)
|
||||
{
|
||||
uint32_t idxMap[] = {0, 2}; // 确保0块和1块连在一起, 2和3块连在一起, 来保证同一m块的地址相连
|
||||
return idxMap[mIdx % 2] + k1Idx;
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyGmToL1(LocalTensor<K_ROPE_T> &l1Tensor, GlobalTensor<K_ROPE_T> &gmSrcTensor,
|
||||
uint32_t srcN, uint32_t srcD, uint32_t srcDstride);
|
||||
__aicore__ inline void CopyInMm1AToL1(LocalTensor<K_ROPE_T> &aL1Tensor, const RunInfo &info, uint32_t mSeqIdx,
|
||||
uint32_t mSizeAct, uint32_t headSize, uint32_t headOffset);
|
||||
__aicore__ inline void CopyInMm1ARopeToL1(LocalTensor<K_ROPE_T> &aL1Tensor, const RunInfo &info, uint32_t mSeqIdx,
|
||||
uint32_t mSizeAct);
|
||||
__aicore__ inline void CopyInMm1BToL1(LocalTensor<K_ROPE_T> &bL1Tensor, const uint64_t keyGmBaseOffset,
|
||||
uint32_t copyTotalRowCntAlign, uint32_t copyStartRowCnt,
|
||||
uint32_t nActCopyRowCount, uint32_t headSize);
|
||||
__aicore__ inline void CopyInMm1BRopeToL1(LocalTensor<K_ROPE_T> &bL1Tensor, const uint64_t keyGmBaseOffset,
|
||||
uint32_t copyTotalRowCntAlign, uint32_t copyStartRowCnt,
|
||||
uint32_t nActCopyRowCount, uint32_t headSize);
|
||||
__aicore__ inline void CopyInMm2AToL1(LocalTensor<K_ROPE_T> &aL1Tensor, const RunInfo &info, uint32_t mSeqIdx,
|
||||
uint32_t subMSizeAct, uint32_t nSize, uint32_t nOffset);
|
||||
__aicore__ inline void CopyInMm2BToL1(LocalTensor<K_ROPE_T> &bL1Tensor, const uint64_t valueGmBaseOffset,
|
||||
uint32_t copyTotalRowCntAlign, uint32_t copyStartRowCnt,
|
||||
uint32_t nActCopyRowCount, uint32_t copyStartColumnCount,
|
||||
uint32_t copyColumnCount);
|
||||
__aicore__ inline void LoadDataMm1A(LocalTensor<K_ROPE_T> &aL0Tensor, LocalTensor<K_ROPE_T> &aL1Tensor,
|
||||
uint32_t idx, uint32_t kSplitSize, uint32_t mSize, uint32_t kSize);
|
||||
__aicore__ inline void LoadDataMm1B(LocalTensor<K_ROPE_T> &bL0Tensor, LocalTensor<K_ROPE_T> &bL1Tensor,
|
||||
uint32_t idx, uint32_t kSplitSize, uint32_t kSize, uint32_t nSize);
|
||||
};
|
||||
|
||||
template <typename QSFAT> __aicore__ inline void QSFAMatmulService<QSFAT>::InitParams(const ConstInfo &constInfo)
|
||||
{
|
||||
this->constInfo = constInfo;
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void
|
||||
QSFAMatmulService<QSFAT>::InitMm1GlobalTensor(GlobalTensor<Q_T> queryGm, GlobalTensor<Q_T> qRopeGm,
|
||||
GlobalTensor<KV_T> keyGm, GlobalTensor<K_ROPE_T> kRopeGm,
|
||||
GlobalTensor<MM_OUT_T> mm1ResGm)
|
||||
{
|
||||
// mm1
|
||||
this->queryGm = queryGm;
|
||||
this->qRopeGm = qRopeGm;
|
||||
this->keyGm = keyGm;
|
||||
this->kRopeGm = kRopeGm;
|
||||
this->mm1ResGm = mm1ResGm;
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void
|
||||
QSFAMatmulService<QSFAT>::InitMm2GlobalTensor(GlobalTensor<K_ROPE_T> vec1ResGm, GlobalTensor<KV_T> valueGm,
|
||||
GlobalTensor<MM_OUT_T> mm2ResGm, GlobalTensor<OUT_T> attentionOutGm)
|
||||
{
|
||||
// mm2
|
||||
this->vec1ResGm = vec1ResGm;
|
||||
this->valueGm = valueGm;
|
||||
this->mm2ResGm = mm2ResGm;
|
||||
this->attentionOutGm = attentionOutGm;
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void
|
||||
QSFAMatmulService<QSFAT>::InitPageAttentionInfo(const GlobalTensor<K_ROPE_T>& kvMergeGm,
|
||||
GlobalTensor<int32_t> blockTableGm, GlobalTensor<int32_t> topKGm,
|
||||
uint32_t blockSize, uint32_t maxBlockNumPerBatch)
|
||||
{
|
||||
this->blockTableGm = blockTableGm;
|
||||
this->topKGm = topKGm;
|
||||
this->kvCacheBlockSize = blockSize;
|
||||
this->maxBlockNumPerBatch = maxBlockNumPerBatch;
|
||||
this->kvMergeGm_ = kvMergeGm;
|
||||
}
|
||||
|
||||
template <typename QSFAT> __aicore__ inline void QSFAMatmulService<QSFAT>::InitBuffers(TPipe *pipe)
|
||||
{
|
||||
pipe->InitBuffer(bufQPL1, L1_BLOCK_SIZE * 4); // (64K + 8K) * 4
|
||||
l1QPTensor = bufQPL1.Get<Q_T>();
|
||||
pipe->InitBuffer(bufKVL1, L1_BLOCK_SIZE * 3); // (64K + 8K) * 3
|
||||
l1KVTensor = bufKVL1.Get<K_ROPE_T>();
|
||||
|
||||
// L0A
|
||||
pipe->InitBuffer(tmpBufL0A, L0A_PP_SIZE * 2); // 64K
|
||||
aL0TensorPingPong = tmpBufL0A.Get<K_ROPE_T>();
|
||||
// L0B
|
||||
pipe->InitBuffer(tmpBufL0B, L0B_PP_SIZE * 2); // 64K
|
||||
bL0TensorPingPong = tmpBufL0B.Get<K_ROPE_T>();
|
||||
// L0C
|
||||
pipe->InitBuffer(tmpBufL0C, L0C_PP_SIZE * 2); // 128K
|
||||
cL0TensorPingPong = tmpBufL0C.Get<MM_OUT_T>();
|
||||
}
|
||||
|
||||
template <typename QSFAT> __aicore__ inline void QSFAMatmulService<QSFAT>::UpdateKey(GlobalTensor<KV_T> keyGm)
|
||||
{
|
||||
this->keyGm = keyGm;
|
||||
}
|
||||
|
||||
template <typename QSFAT> __aicore__ inline void QSFAMatmulService<QSFAT>::UpdateValue(GlobalTensor<KV_T> valueGm)
|
||||
{
|
||||
this->valueGm = valueGm;
|
||||
}
|
||||
|
||||
template <typename QSFAT> __aicore__ inline void QSFAMatmulService<QSFAT>::AllocEventID()
|
||||
{
|
||||
SetFlag<HardEvent::M_MTE1>(L0AB_EVENT0);
|
||||
SetFlag<HardEvent::M_MTE1>(L0AB_EVENT1);
|
||||
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT0);
|
||||
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT1);
|
||||
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT2);
|
||||
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT3);
|
||||
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT4);
|
||||
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT5);
|
||||
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT6);
|
||||
}
|
||||
|
||||
template <typename QSFAT> __aicore__ inline void QSFAMatmulService<QSFAT>::FreeEventID()
|
||||
{
|
||||
WaitFlag<HardEvent::M_MTE1>(L0AB_EVENT0);
|
||||
WaitFlag<HardEvent::M_MTE1>(L0AB_EVENT1);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT0);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT1);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT2);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT3);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT4);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT5);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT6);
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void QSFAMatmulService<QSFAT>::CopyGmToL1(LocalTensor<K_ROPE_T> &l1Tensor,
|
||||
GlobalTensor<K_ROPE_T> &gmSrcTensor, uint32_t srcN,
|
||||
uint32_t srcD, uint32_t srcDstride)
|
||||
{
|
||||
Nd2NzParams nd2nzPara;
|
||||
nd2nzPara.ndNum = 1;
|
||||
nd2nzPara.dValue = srcD;
|
||||
nd2nzPara.nValue = srcN; // 行数
|
||||
nd2nzPara.srcDValue = srcDstride;
|
||||
nd2nzPara.dstNzC0Stride = (srcN + 15) / 16 * 16; // 对齐到16 单位block
|
||||
nd2nzPara.dstNzNStride = 1;
|
||||
nd2nzPara.dstNzMatrixStride = 0;
|
||||
nd2nzPara.srcNdMatrixStride = 0;
|
||||
DataCopy(l1Tensor, gmSrcTensor, nd2nzPara);
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void QSFAMatmulService<QSFAT>::CopyInMm1AToL1(LocalTensor<K_ROPE_T> &l1Tensor, const RunInfo &info,
|
||||
uint32_t mSeqIdx, uint32_t mSizeAct,
|
||||
uint32_t headSize, uint32_t headOffset)
|
||||
{
|
||||
auto srcGm = queryGm[info.tensorAOffset + mSeqIdx * constInfo.combineHeadDim + headOffset];
|
||||
CopyGmToL1(l1Tensor, srcGm, mSizeAct, headSize, headSize);
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void QSFAMatmulService<QSFAT>::CopyInMm1ARopeToL1(LocalTensor<K_ROPE_T> &l1Tensor,
|
||||
const RunInfo &info, uint32_t mSeqIdx,
|
||||
uint32_t mSizeAct)
|
||||
{
|
||||
auto srcGm = qRopeGm[info.tensorARopeOffset + mSeqIdx * constInfo.headDimRope];
|
||||
CopyGmToL1(l1Tensor, srcGm, mSizeAct, constInfo.headDimRope, constInfo.headDimRope);
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void
|
||||
QSFAMatmulService<QSFAT>::CopyInMm1BToL1(LocalTensor<K_ROPE_T> &bL1Tensor, const uint64_t keyGmBaseOffset,
|
||||
uint32_t copyTotalRowCntAlign, uint32_t copyStartRowCnt,
|
||||
uint32_t nActCopyRowCount, uint32_t headSize)
|
||||
{
|
||||
uint64_t dStride = constInfo.headDim;
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::BSND || LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
dStride = constInfo.headDim * constInfo.kvHeadNum;
|
||||
}
|
||||
|
||||
uint32_t blockElementCnt = 32 / sizeof(K_ROPE_T);
|
||||
|
||||
Nd2NzParams mm1Nd2NzParamsForB;
|
||||
mm1Nd2NzParamsForB.ndNum = 1;
|
||||
mm1Nd2NzParamsForB.nValue = nActCopyRowCount;
|
||||
mm1Nd2NzParamsForB.dValue = headSize;
|
||||
mm1Nd2NzParamsForB.srcDValue = dStride;
|
||||
mm1Nd2NzParamsForB.dstNzNStride = 1;
|
||||
mm1Nd2NzParamsForB.dstNzC0Stride = copyTotalRowCntAlign;
|
||||
mm1Nd2NzParamsForB.srcNdMatrixStride = 0;
|
||||
mm1Nd2NzParamsForB.dstNzMatrixStride = 0;
|
||||
DataCopy(bL1Tensor[copyStartRowCnt * blockElementCnt], keyGm[keyGmBaseOffset], mm1Nd2NzParamsForB);
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void
|
||||
QSFAMatmulService<QSFAT>::CopyInMm1BRopeToL1(LocalTensor<K_ROPE_T> &bL1Tensor, const uint64_t kRopeGmBaseOffset,
|
||||
uint32_t copyTotalRowCntAlign, uint32_t copyStartRowCnt,
|
||||
uint32_t nActCopyRowCount, uint32_t headSize)
|
||||
{
|
||||
uint64_t dStride = constInfo.headDimRope;
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::BSND || LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
dStride = constInfo.headDimRope * constInfo.kvHeadNum;
|
||||
}
|
||||
|
||||
uint32_t blockElementCnt = 32 / sizeof(K_ROPE_T);
|
||||
|
||||
Nd2NzParams mm1Nd2NzParamsForB;
|
||||
mm1Nd2NzParamsForB.nValue = nActCopyRowCount;
|
||||
mm1Nd2NzParamsForB.dValue = headSize;
|
||||
mm1Nd2NzParamsForB.ndNum = 1;
|
||||
mm1Nd2NzParamsForB.srcDValue = dStride;
|
||||
mm1Nd2NzParamsForB.srcNdMatrixStride = 0;
|
||||
mm1Nd2NzParamsForB.dstNzMatrixStride = 0;
|
||||
mm1Nd2NzParamsForB.dstNzNStride = 1;
|
||||
mm1Nd2NzParamsForB.dstNzC0Stride = copyTotalRowCntAlign;
|
||||
DataCopy(bL1Tensor[copyStartRowCnt * blockElementCnt], kRopeGm[kRopeGmBaseOffset], mm1Nd2NzParamsForB);
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void QSFAMatmulService<QSFAT>::LoadDataMm1A(LocalTensor<K_ROPE_T> &aL0Tensor,
|
||||
LocalTensor<K_ROPE_T> &aL1Tensor, uint32_t idx,
|
||||
uint32_t kSplitSize, uint32_t mSize, uint32_t kSize)
|
||||
{
|
||||
LocalTensor<K_ROPE_T> srcTensor = aL1Tensor[mSize * kSplitSize * idx];
|
||||
LoadData3DParamsV2<K_ROPE_T> loadData3DParams;
|
||||
// SetFmatrixParams
|
||||
loadData3DParams.l1H = mSize / 16; // Hin=M1=8
|
||||
loadData3DParams.l1W = 16; // Win=M0
|
||||
loadData3DParams.padList[0] = 0;
|
||||
loadData3DParams.padList[1] = 0;
|
||||
loadData3DParams.padList[2] = 0;
|
||||
loadData3DParams.padList[3] = 255; // 尾部数据不影响滑窗的结果
|
||||
|
||||
// SetLoadToA0Params
|
||||
loadData3DParams.mExtension = mSize; // M
|
||||
loadData3DParams.kExtension = kSize; // K
|
||||
loadData3DParams.mStartPt = 0;
|
||||
loadData3DParams.kStartPt = 0;
|
||||
loadData3DParams.strideH = 1;
|
||||
loadData3DParams.strideW = 1;
|
||||
loadData3DParams.filterW = 1;
|
||||
loadData3DParams.filterSizeW = (1 >> 8) & 255;
|
||||
loadData3DParams.filterH = 1;
|
||||
loadData3DParams.filterSizeH = (1 >> 8) & 255;
|
||||
loadData3DParams.dilationFilterH = 1;
|
||||
loadData3DParams.dilationFilterW = 1;
|
||||
loadData3DParams.fMatrixCtrl = 0;
|
||||
loadData3DParams.channelSize = kSize; // Cin=K
|
||||
loadData3DParams.enTranspose = 0;
|
||||
LoadData<K_ROPE_T, LOAD3DV2_CONFIG>(aL0Tensor, srcTensor, loadData3DParams);
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void QSFAMatmulService<QSFAT>::LoadDataMm1B(LocalTensor<K_ROPE_T> &l0Tensor,
|
||||
LocalTensor<K_ROPE_T> &l1Tensor, uint32_t idx,
|
||||
uint32_t kSplitSize, uint32_t kSize, uint32_t nSize)
|
||||
{
|
||||
// N 方向全载
|
||||
LocalTensor<K_ROPE_T> srcTensor = l1Tensor[nSize * kSplitSize * idx];
|
||||
|
||||
LoadData2DParams loadData2DParams;
|
||||
loadData2DParams.startIndex = 0;
|
||||
loadData2DParams.repeatTimes = (nSize + 15) / 16 * kSize / (32 / sizeof(K_ROPE_T));
|
||||
loadData2DParams.srcStride = 1;
|
||||
loadData2DParams.dstGap = 0;
|
||||
loadData2DParams.ifTranspose = false;
|
||||
LoadData(l0Tensor, srcTensor, loadData2DParams);
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void QSFAMatmulService<QSFAT>::CopyInMm2AToL1(LocalTensor<K_ROPE_T> &aL1Tensor, const RunInfo &info,
|
||||
uint32_t mSeqIdx, uint32_t subMSizeAct,
|
||||
uint32_t nSize, uint32_t nOffset)
|
||||
{
|
||||
auto srcGm = vec1ResGm[(info.loop % constInfo.preLoadNum) * constInfo.mmResUbSize +
|
||||
mSeqIdx * info.actualSingleProcessSInnerSizeAlign + nOffset];
|
||||
CopyGmToL1(aL1Tensor, srcGm, subMSizeAct, nSize, info.actualSingleProcessSInnerSizeAlign);
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void QSFAMatmulService<QSFAT>::CopyInMm2BToL1(
|
||||
LocalTensor<K_ROPE_T> &bL1Tensor, const uint64_t valueGmBaseOffset, uint32_t copyTotalRowCntAlign,
|
||||
uint32_t copyStartRowCnt, uint32_t nActCopyRowCount, uint32_t copyStartColumnCount, uint32_t copyColumnCount)
|
||||
{
|
||||
uint64_t step = constInfo.headDim;
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::BSND || LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
step = constInfo.headDim * constInfo.kvHeadNum;
|
||||
}
|
||||
|
||||
uint32_t blockElementCnt = 32 / sizeof(K_ROPE_T);
|
||||
|
||||
Nd2NzParams mm1Nd2NzParamsForB;
|
||||
mm1Nd2NzParamsForB.ndNum = 1;
|
||||
mm1Nd2NzParamsForB.nValue = nActCopyRowCount;
|
||||
mm1Nd2NzParamsForB.dValue = copyColumnCount;
|
||||
mm1Nd2NzParamsForB.srcDValue = step;
|
||||
mm1Nd2NzParamsForB.dstNzNStride = 1;
|
||||
mm1Nd2NzParamsForB.dstNzC0Stride = copyTotalRowCntAlign;
|
||||
mm1Nd2NzParamsForB.srcNdMatrixStride = 0;
|
||||
mm1Nd2NzParamsForB.dstNzMatrixStride = 0;
|
||||
DataCopy(bL1Tensor[copyStartRowCnt * blockElementCnt], valueGm[valueGmBaseOffset + copyStartColumnCount],
|
||||
mm1Nd2NzParamsForB);
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void QSFAMatmulService<QSFAT>::CalcTopKBlockInfo(
|
||||
const RunInfo &info, uint32_t &curTopKIdx, uint64_t &curOffsetInSparseBlock,
|
||||
uint32_t curSeqIdx, uint32_t ©RowCnt, uint64_t &idInTopK)
|
||||
{
|
||||
if (curTopKIdx == 0 && curOffsetInSparseBlock == 0 && copyRowCnt == 0) {
|
||||
uint64_t sparseLen = 0;
|
||||
for (uint64_t qsfaTopkidx = 0; qsfaTopkidx < constInfo.sparseBlockCount; qsfaTopkidx++) {
|
||||
int32_t qsfaSparseIndices = topKGm.GetValue(info.topKBaseOffset + qsfaTopkidx);
|
||||
if (qsfaSparseIndices == -1) {
|
||||
break;
|
||||
}
|
||||
uint64_t qsfaBlockBegin = qsfaSparseIndices * constInfo.sparseBlockSize;
|
||||
if (qsfaBlockBegin >= info.threshold) {
|
||||
continue;
|
||||
}
|
||||
uint64_t qsfaBlockEnd = (qsfaBlockBegin + constInfo.sparseBlockSize > info.curActualSeqLenOri) ?
|
||||
info.curActualSeqLenOri : qsfaBlockBegin + constInfo.sparseBlockSize;
|
||||
uint64_t qsfaBlockLen = (qsfaBlockEnd <= info.threshold) ? \
|
||||
qsfaBlockEnd - qsfaBlockBegin : info.threshold - qsfaBlockBegin;
|
||||
sparseLen += qsfaBlockLen;
|
||||
if (sparseLen >= curSeqIdx + 1) {
|
||||
curTopKIdx = qsfaTopkidx;
|
||||
idInTopK = qsfaSparseIndices;
|
||||
curOffsetInSparseBlock = qsfaBlockLen - (sparseLen - curSeqIdx);
|
||||
copyRowCnt = sparseLen - curSeqIdx;
|
||||
break;
|
||||
}
|
||||
}
|
||||
return;
|
||||
}
|
||||
uint64_t qsfaBlockBegin = idInTopK * constInfo.sparseBlockSize;
|
||||
uint64_t qsfaBlockEnd = (qsfaBlockBegin + constInfo.sparseBlockSize > info.threshold) ?
|
||||
info.threshold : qsfaBlockBegin + constInfo.sparseBlockSize;
|
||||
uint64_t qsfaBlockLen = qsfaBlockEnd - qsfaBlockBegin;
|
||||
if (curOffsetInSparseBlock + copyRowCnt < qsfaBlockLen) {
|
||||
curOffsetInSparseBlock += copyRowCnt;
|
||||
copyRowCnt = qsfaBlockLen - curOffsetInSparseBlock;
|
||||
} else {
|
||||
for (uint64_t qsfaTopkidx = curTopKIdx + 1; qsfaTopkidx < constInfo.sparseBlockCount; qsfaTopkidx++) {
|
||||
int64_t qsfaSparseIndices = topKGm.GetValue(info.topKBaseOffset + qsfaTopkidx);
|
||||
if (qsfaSparseIndices == -1) {
|
||||
break;
|
||||
}
|
||||
|
||||
uint64_t qsfaBlockBegin = qsfaSparseIndices * constInfo.sparseBlockSize;
|
||||
if (qsfaBlockBegin >= info.threshold) {
|
||||
continue;
|
||||
}
|
||||
uint64_t qsfaBlockEnd = (qsfaBlockBegin + constInfo.sparseBlockSize > info.threshold) ?
|
||||
info.threshold : qsfaBlockBegin + constInfo.sparseBlockSize;
|
||||
uint64_t qsfaBlockLen = qsfaBlockEnd - qsfaBlockBegin;
|
||||
curTopKIdx = qsfaTopkidx;
|
||||
idInTopK = qsfaSparseIndices;
|
||||
curOffsetInSparseBlock = 0;
|
||||
copyRowCnt = qsfaBlockLen;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void QSFAMatmulService<QSFAT>::ComputeMm1(const RunInfo &info, const MSplitInfo mSplitInfo)
|
||||
{
|
||||
// 最外层还需要一层m的循环
|
||||
uint32_t mSize = mSplitInfo.nBufferDealM;
|
||||
uint32_t mL1Size = M_SPLIT_SIZE;
|
||||
uint32_t mL1SizeAlign = QSFAAlign(M_SPLIT_SIZE, 16U);
|
||||
uint32_t mL1Loops = (mSize + M_SPLIT_SIZE - 1) / M_SPLIT_SIZE;
|
||||
|
||||
uint32_t nSize = info.actualSingleProcessSInnerSize;
|
||||
uint32_t nL1Size = N_SPLIT_SIZE;
|
||||
uint32_t nL1SizeAlign = QSFAAlign(N_SPLIT_SIZE, 16U);
|
||||
uint32_t nL1Loops = (nSize + N_SPLIT_SIZE - 1) / N_SPLIT_SIZE;
|
||||
|
||||
uint32_t kSize = 576;
|
||||
uint32_t kL1Size = 288;
|
||||
uint32_t kL1Loops = 2; // 2 : 576/288, mla专用 这里不考虑d泛化
|
||||
|
||||
uint32_t kL0Size = 96;
|
||||
uint32_t kL0Loops = (kL1Size + kL0Size - 1) / kL0Size; // 288 / 96 = 3 kloops
|
||||
|
||||
// ka表示左矩阵4buf选择哪一块buf, kb表示右矩阵3buf选择哪一块buf
|
||||
uint32_t ka = 0, kb = 0;
|
||||
for (uint32_t mL1 = 0; mL1 < mL1Loops; mL1++) {
|
||||
mL1Size = M_SPLIT_SIZE;
|
||||
mL1SizeAlign = QSFAAlign(M_SPLIT_SIZE, 16U);
|
||||
if (mL1 == (mL1Loops - 1)) {
|
||||
// 尾块重新计算size
|
||||
mL1Size = mSize - (mL1Loops - 1) * M_SPLIT_SIZE;
|
||||
mL1SizeAlign = QSFAAlign(mL1Size, 16U);
|
||||
}
|
||||
|
||||
// 左矩阵L1选择12块还是34块的index, 由m l1 index决定
|
||||
// 左矩阵L1选择12块或34块的前一块还是后一块, 由k l1 index决定
|
||||
uint32_t mIdx = qpL1BufIter + mL1;
|
||||
ka = GetQPL1RealIdx(mIdx, 0);
|
||||
LocalTensor<Q_T> aL1Tensor = l1QPTensor[ka * L1_BLOCK_OFFSET];
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(mte21QPIds[ka]);
|
||||
CopyInMm1AToL1(aL1Tensor, info, mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE, mL1Size, 576, 0);
|
||||
SetFlag<HardEvent::MTE2_MTE1>(mte21QPIds[ka]);
|
||||
WaitFlag<HardEvent::MTE2_MTE1>(mte21QPIds[ka]);
|
||||
for (uint32_t nL1 = 0; nL1 < nL1Loops; nL1++) { // L1切n, 512/128=4
|
||||
if (nL1 == (nL1Loops - 1)) {
|
||||
// 尾块重新计算size
|
||||
nL1Size = nSize - (nL1Loops - 1) * N_SPLIT_SIZE;
|
||||
nL1SizeAlign = QSFAAlign(nL1Size, 16U);
|
||||
}
|
||||
|
||||
// 使用unitflag同步
|
||||
// 需要保证cL0BufIter和m步调一致
|
||||
LocalTensor cL0Tensor = cL0TensorPingPong[(cL0BufIter % 2) * (L0C_PP_SIZE / sizeof(MM_OUT_T))];
|
||||
for (uint32_t kL1 = 0; kL1 < kL1Loops; kL1++) { // L1切k, 576/288, 这里不考虑d泛化
|
||||
kvL1BufIter++;
|
||||
uint32_t kb = kvL1BufIter % 3;
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(mte21KVIds[kb]);
|
||||
// 从k当中取当前的块
|
||||
LocalTensor<K_ROPE_T> bL1Tensor = l1KVTensor[kb * L1_BLOCK_OFFSET];
|
||||
if constexpr (TEMPLATE_MODE == V_TEMPLATE) {
|
||||
if (kL1 == 0) {
|
||||
DataCopyParams copyParams;
|
||||
copyParams.blockCount = 288 / BLOCK_ELEMENT_NUM;
|
||||
copyParams.blockLen = nL1Size;
|
||||
copyParams.srcStride = constInfo.s2BaseSize - nL1Size;
|
||||
copyParams.dstStride = nL1SizeAlign - nL1Size;
|
||||
DataCopy(bL1Tensor, kvMergeGm_[info.loop % 4 * N_WORKSPACE_SIZE * kSize +
|
||||
nL1 * N_SPLIT_SIZE * BLOCK_ELEMENT_NUM], copyParams);
|
||||
} else {
|
||||
DataCopyParams copyParams;
|
||||
copyParams.blockCount = 224 / BLOCK_ELEMENT_NUM;
|
||||
copyParams.blockLen = nL1Size;
|
||||
copyParams.srcStride = constInfo.s2BaseSize - nL1Size;
|
||||
copyParams.dstStride = nL1SizeAlign - nL1Size;
|
||||
DataCopy(bL1Tensor, kvMergeGm_[info.loop % 4 * N_WORKSPACE_SIZE * kSize +
|
||||
288 * constInfo.s2BaseSize + nL1 * N_SPLIT_SIZE * BLOCK_ELEMENT_NUM], copyParams);
|
||||
copyParams.blockCount = constInfo.headDimRope / BLOCK_ELEMENT_NUM;
|
||||
DataCopy(
|
||||
bL1Tensor[224 * nL1SizeAlign],
|
||||
kvMergeGm_[info.loop % 4 * N_WORKSPACE_SIZE * kSize + N_WORKSPACE_SIZE * constInfo.headDim +
|
||||
nL1 * N_SPLIT_SIZE * BLOCK_ELEMENT_NUM],
|
||||
copyParams);
|
||||
}
|
||||
}
|
||||
SetFlag<HardEvent::MTE2_MTE1>(mte21KVIds[kb]);
|
||||
WaitFlag<HardEvent::MTE2_MTE1>(mte21KVIds[kb]);
|
||||
|
||||
aL1Tensor = l1QPTensor[ka * L1_BLOCK_OFFSET + kL1 * mL1SizeAlign * K_SPLIT_SIZE];
|
||||
for (uint32_t kL0 = 0; kL0 < kL0Loops; kL0++) {
|
||||
WaitFlag<HardEvent::M_MTE1>(Mte1MmABEventId(abL0BufIter % 2));
|
||||
LocalTensor<K_ROPE_T> aL0Tensor = aL0TensorPingPong[(abL0BufIter % 2) * (L0A_PP_SIZE /
|
||||
sizeof(K_ROPE_T))];
|
||||
LoadDataMm1A(aL0Tensor, aL1Tensor, kL0, kL0Size, mL1SizeAlign, kL0Size);
|
||||
LocalTensor<K_ROPE_T> bL0Tensor = bL0TensorPingPong[(abL0BufIter % 2) * (L0B_PP_SIZE /
|
||||
sizeof(K_ROPE_T))];
|
||||
LoadDataMm1B(bL0Tensor, bL1Tensor, kL0, kL0Size, kL0Size, nL1SizeAlign);
|
||||
SetFlag<HardEvent::MTE1_M>(Mte1MmABEventId(abL0BufIter % 2));
|
||||
WaitFlag<HardEvent::MTE1_M>(Mte1MmABEventId(abL0BufIter % 2));
|
||||
|
||||
// m == 1的时候需要特殊处理
|
||||
MmadParams mmadParams;
|
||||
mmadParams.m = mL1SizeAlign;
|
||||
mmadParams.n = nL1SizeAlign;
|
||||
mmadParams.k = kL0Size;
|
||||
mmadParams.cmatrixSource = false;
|
||||
mmadParams.cmatrixInitVal = (kL1 == 0 && kL0 == 0);
|
||||
mmadParams.unitFlag =
|
||||
(kL1 == 1 && kL0 == (kL0Loops - 1)) ? 0b11 : 0b10; // 累加最后一次翻转flag, 表示可以搬出
|
||||
Mmad(cL0Tensor, aL0Tensor, bL0Tensor, mmadParams);
|
||||
|
||||
if ((mmadParams.m / 16) * (mmadParams.n / 16) < 10) {
|
||||
PipeBarrier<PIPE_M>();
|
||||
}
|
||||
SetFlag<HardEvent::M_MTE1>(Mte1MmABEventId(abL0BufIter % 2));
|
||||
abL0BufIter++;
|
||||
}
|
||||
SetFlag<HardEvent::MTE1_MTE2>(mte21KVIds[kb]); // 反向同步, 表示L1已经被mte1消费完
|
||||
}
|
||||
FixpipeParamsV220 fixParams;
|
||||
fixParams.mSize = mL1SizeAlign;
|
||||
fixParams.nSize = nL1SizeAlign;
|
||||
fixParams.srcStride = mL1SizeAlign;
|
||||
fixParams.ndNum = 1; // 输出ND
|
||||
// 改成nSizeAlign
|
||||
fixParams.dstStride = info.actualSingleProcessSInnerSizeAlign; // mm1ResGm两行之间的间隔
|
||||
fixParams.unitFlag = 0b11;
|
||||
|
||||
// 输出偏移info.loop % (constInfo.preLoadNum)) * mmResUbSize是否在matmul里计算
|
||||
Fixpipe(mm1ResGm[(info.loop % (constInfo.preLoadNum)) * constInfo.mmResUbSize + nL1 * N_SPLIT_SIZE +
|
||||
(mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE) *
|
||||
info.actualSingleProcessSInnerSizeAlign],
|
||||
cL0Tensor, fixParams);
|
||||
cL0BufIter++;
|
||||
}
|
||||
SetFlag<HardEvent::MTE1_MTE2>(mte21QPIds[ka]); // 反向同步, 表示L1中的A已经被mte1消费完
|
||||
}
|
||||
qpL1BufIter += mL1Loops;
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void QSFAMatmulService<QSFAT>::ComputeMm2(const RunInfo &info, const MSplitInfo mSplitInfo)
|
||||
{
|
||||
uint32_t mSize = mSplitInfo.nBufferDealM;
|
||||
uint32_t mSizeAlign = (mSize + 16 - 1) / 16;
|
||||
uint32_t mL1Loops = (mSize + M_SPLIT_SIZE - 1) / M_SPLIT_SIZE;
|
||||
uint32_t mL1SizeAlign = M_SPLIT_SIZE; // 16对齐
|
||||
uint32_t mL1Size = M_SPLIT_SIZE; // m的实际大小
|
||||
|
||||
uint32_t nSize = BlockAlign<K_ROPE_T>(constInfo.headDim);
|
||||
uint32_t nL1Loops = (nSize + N_SPLIT_SIZE - 1) / N_SPLIT_SIZE;
|
||||
uint32_t nL1SizeAlign = N_SPLIT_SIZE; // 16对齐
|
||||
uint32_t nL1Size = N_SPLIT_SIZE; // n的实际大小
|
||||
|
||||
uint32_t kSize = info.actualSingleProcessSInnerSize;
|
||||
uint32_t kL1Size = 256;
|
||||
uint32_t kL1SizeAlign = QSFAAlign(kL1Size, 16U);
|
||||
uint32_t kL1Loops = (kSize + kL1Size - 1) / kL1Size;
|
||||
uint32_t kL0Size = 128;
|
||||
uint32_t kL0Loops = (kL1Size + kL0Size - 1) / kL0Size;
|
||||
uint32_t kL0SizeAlign = kL0Size;
|
||||
LocalTensor<K_ROPE_T> bL1Tensor;
|
||||
LocalTensor<K_ROPE_T> subvTensor;
|
||||
|
||||
// ka表示左矩阵4buf选择哪一块buf, kb表示右矩阵3buf选择哪一块buf
|
||||
uint32_t ka = 0, qsfaKb = 0;
|
||||
uint32_t mBaseIdx = qpL1BufIter;
|
||||
for (uint32_t nL1 = 0; nL1 < nL1Loops; nL1++) { // n切L1
|
||||
if (nL1 == (nL1Loops - 1)) {
|
||||
// 尾块
|
||||
nL1Size = nSize - (nL1Loops - 1) * N_SPLIT_SIZE;
|
||||
nL1SizeAlign = QSFAAlign(nL1Size, 16U);
|
||||
}
|
||||
|
||||
// k l1写成一个循环, 和mm1保持一致
|
||||
kL1Size = 256;
|
||||
kL1SizeAlign = QSFAAlign(kL1Size, 16U);
|
||||
for (uint32_t k1 = 0; k1 < kL1Loops; k1++) { // k切L1, 这里套了一层l0来操作
|
||||
if (k1 == (kL1Loops - 1)) {
|
||||
// 尾块
|
||||
kL1Size = kSize - (kL1Loops - 1) * 256;
|
||||
kL1SizeAlign = QSFAAlign(kL1Size, 16U);
|
||||
}
|
||||
kvL1BufIter++;
|
||||
uint32_t qsfaKb = kvL1BufIter % 3;
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(mte21KVIds[qsfaKb]);
|
||||
bL1Tensor = l1KVTensor[qsfaKb * L1_BLOCK_OFFSET];
|
||||
uint32_t qsfaKOffset = k1 * kL0Loops;
|
||||
kL0Size = 128;
|
||||
// 此处必须先初始化kL0Size, 再求kL0Loops, 否则由于循环会改变kL0Size大小, 导致kL0Loops错误
|
||||
kL0Loops = (kL1Size + kL0Size - 1) / kL0Size;
|
||||
kL0SizeAlign = kL0Size;
|
||||
for (uint32_t qsfaKL1 = qsfaKOffset; qsfaKL1 < kL0Loops + qsfaKOffset; qsfaKL1++) { // 128 循环搬pa
|
||||
if (qsfaKL1 == qsfaKOffset + kL0Loops - 1) {
|
||||
// 尾块
|
||||
kL0Size = kL1Size - (kL0Loops - 1) * kL0Size;
|
||||
kL0SizeAlign = QSFAAlign(kL0Size, 16U);
|
||||
}
|
||||
if constexpr (TEMPLATE_MODE == V_TEMPLATE) {
|
||||
DataCopyParams copyParams;
|
||||
copyParams.blockLen = kL0Size;
|
||||
copyParams.blockCount = nL1Size / BLOCK_ELEMENT_NUM;
|
||||
copyParams.srcStride = constInfo.s2BaseSize - kL0Size;
|
||||
copyParams.dstStride = kL0SizeAlign - kL0Size;
|
||||
DataCopy(bL1Tensor[(qsfaKL1 - qsfaKOffset) * 128 * N_SPLIT_SIZE], kvMergeGm_[info.loop % 4 *
|
||||
N_WORKSPACE_SIZE * 576 + qsfaKL1 * 128 * BLOCK_ELEMENT_NUM + nL1 * N_SPLIT_SIZE *
|
||||
constInfo.s2BaseSize], copyParams);
|
||||
}
|
||||
}
|
||||
SetFlag<HardEvent::MTE2_MTE1>(mte21KVIds[qsfaKb]);
|
||||
WaitFlag<HardEvent::MTE2_MTE1>(mte21KVIds[qsfaKb]);
|
||||
mL1SizeAlign = M_SPLIT_SIZE;
|
||||
mL1Size = M_SPLIT_SIZE; // m的实际大小
|
||||
for (uint32_t qsfaML1 = 0; qsfaML1 < mL1Loops; qsfaML1++) {
|
||||
if (qsfaML1 == (mL1Loops - 1)) {
|
||||
// 尾块
|
||||
mL1Size = mSize - (mL1Loops - 1) * M_SPLIT_SIZE;
|
||||
mL1SizeAlign = QSFAAlign(mL1Size, 16U);
|
||||
}
|
||||
|
||||
uint32_t mIdx = mBaseIdx + qsfaML1;
|
||||
ka = GetQPL1RealIdx(mIdx, k1);
|
||||
LocalTensor<K_ROPE_T> aL1Tensor = l1QPTensor[ka * L1_BLOCK_OFFSET];
|
||||
if (nL1 == 0) {
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(mte21QPIds[ka]);
|
||||
CopyInMm2AToL1(aL1Tensor, info, mSplitInfo.nBufferStartM + qsfaML1 * M_SPLIT_SIZE, mL1Size, kL1Size,
|
||||
256 * k1);
|
||||
SetFlag<HardEvent::MTE2_MTE1>(mte21QPIds[ka]);
|
||||
WaitFlag<HardEvent::MTE2_MTE1>(mte21QPIds[ka]);
|
||||
}
|
||||
|
||||
LocalTensor cL0Tensor =
|
||||
cL0TensorPingPong[(cL0BufIter % 2) *
|
||||
(L0C_PP_SIZE / sizeof(MM_OUT_T))]; // 需要保证cL0BufIter和m步调一致
|
||||
uint32_t qsfaBaseK = 128;
|
||||
uint32_t qsfaBaseN = 128;
|
||||
kL0Size = 128;
|
||||
kL0SizeAlign = kL0Size;
|
||||
for (uint32_t qsfaKL0 = 0; qsfaKL0 < kL0Loops; qsfaKL0++) {
|
||||
if (qsfaKL0 + 1 == kL0Loops) {
|
||||
kL0Size = kL1Size - (kL0Loops - 1) * kL0Size;
|
||||
kL0SizeAlign = QSFAAlign(kL0Size, 16U);
|
||||
}
|
||||
WaitFlag<HardEvent::M_MTE1>(Mte1MmABEventId(abL0BufIter % 2));
|
||||
LocalTensor<K_ROPE_T> bL0Tensor = bL0TensorPingPong[(abL0BufIter % 2) * (L0B_PP_SIZE /
|
||||
sizeof(K_ROPE_T))];
|
||||
LoadData3DParamsV2<K_ROPE_T> loadData3DParamsForB;
|
||||
loadData3DParamsForB.l1H = kL0SizeAlign / 16; // 源操作数height
|
||||
loadData3DParamsForB.l1W = 16; // 源操作数weight=16,目的height=l1H*L1W
|
||||
loadData3DParamsForB.padList[0] = 0;
|
||||
loadData3DParamsForB.padList[1] = 0;
|
||||
loadData3DParamsForB.padList[2] = 0;
|
||||
loadData3DParamsForB.padList[3] = 255; // 尾部数据不影响滑窗的结果
|
||||
|
||||
loadData3DParamsForB.mExtension = kL0SizeAlign; // 在目的操作数height维度的传输长度
|
||||
loadData3DParamsForB.kExtension = nL1SizeAlign; // 在目的操作数width维度的传输长度
|
||||
loadData3DParamsForB.mStartPt = 0; // 卷积核在目的操作数width维度的起点
|
||||
loadData3DParamsForB.kStartPt = 0; // 卷积核在目的操作数height维度的起点
|
||||
loadData3DParamsForB.strideH = 1;
|
||||
loadData3DParamsForB.strideW = 1;
|
||||
loadData3DParamsForB.filterW = 1;
|
||||
loadData3DParamsForB.filterSizeW = false; // 是否在filterW的基础上将卷积核width增加256个元素
|
||||
loadData3DParamsForB.filterH = 1;
|
||||
loadData3DParamsForB.filterSizeH = false; // 是否在filterH的基础上将卷积核height增加256个元素
|
||||
loadData3DParamsForB.dilationFilterH = 1; // 卷积核height膨胀系数
|
||||
loadData3DParamsForB.dilationFilterW = 1; // 卷积核width膨胀系数
|
||||
loadData3DParamsForB.enTranspose = 1; // 是否启用转置功能
|
||||
// 使用FMATRIX_LEFT还是使用FMATRIX_RIGHT,=0使用FMATRIX_LEFT,=1使用FMATRIX_RIGHT 1
|
||||
loadData3DParamsForB.fMatrixCtrl = 0;
|
||||
// 源操作数的通道数。膨胀系数为1时,目的weight为filterW*filterH*channelSize
|
||||
loadData3DParamsForB.channelSize = nL1SizeAlign;
|
||||
LoadData<K_ROPE_T, LOAD3DV2_CONFIG>(bL0Tensor, bL1Tensor[qsfaKL0 * qsfaBaseK * qsfaBaseN],
|
||||
loadData3DParamsForB);
|
||||
LocalTensor<K_ROPE_T> aL0Tensor = aL0TensorPingPong[(abL0BufIter % 2) * (L0A_PP_SIZE /
|
||||
sizeof(K_ROPE_T))];
|
||||
LoadData3DParamsV2<K_ROPE_T> loadData3DParamsForA;
|
||||
loadData3DParamsForA.l1H = mL1SizeAlign / 16; // 源操作数height
|
||||
loadData3DParamsForA.l1W = 16; // 源操作数weight
|
||||
loadData3DParamsForA.padList[0] = 0;
|
||||
loadData3DParamsForA.padList[1] = 0;
|
||||
loadData3DParamsForA.padList[2] = 0;
|
||||
loadData3DParamsForA.padList[3] = 255; // 尾部数据不影响滑窗的结果
|
||||
|
||||
loadData3DParamsForA.mExtension = mL1SizeAlign; // 在目的操作数height维度的传输长度
|
||||
loadData3DParamsForA.kExtension = kL0SizeAlign; // 在目的操作数width维度的传输长度
|
||||
loadData3DParamsForA.mStartPt = 0; // 卷积核在目的操作数width维度的起点
|
||||
loadData3DParamsForA.kStartPt = 0; // 卷积核在目的操作数height维度的起点
|
||||
loadData3DParamsForA.strideW = 1; // 卷积核在源操作数width维度滑动的步长
|
||||
loadData3DParamsForA.strideH = 1; // 卷积核在源操作数height维度滑动的步长
|
||||
loadData3DParamsForA.filterW = 1; // 卷积核width
|
||||
loadData3DParamsForA.filterSizeW = false; // 是否在filterW的基础上将卷积核width增加256个元素
|
||||
loadData3DParamsForA.filterH = 1; // 卷积核height
|
||||
loadData3DParamsForA.filterSizeH = false; // 是否在filterH的基础上将卷积核height增加256个元素
|
||||
loadData3DParamsForA.dilationFilterW = 1; // 卷积核width膨胀系数
|
||||
loadData3DParamsForA.dilationFilterH = 1; // 卷积核height膨胀系数
|
||||
loadData3DParamsForA.enTranspose = 0; // 是否启用转置功能,对整个目标矩阵进行转置
|
||||
loadData3DParamsForA.fMatrixCtrl = 0;
|
||||
// 源操作数的通道数。膨胀系数为1时,目的weight为filterW*filterH*channelSize
|
||||
loadData3DParamsForA.channelSize = kL0SizeAlign;
|
||||
LoadData<K_ROPE_T, LOAD3DV2_CONFIG>(aL0Tensor, aL1Tensor[qsfaKL0 * qsfaBaseK * mL1SizeAlign],
|
||||
loadData3DParamsForA);
|
||||
SetFlag<HardEvent::MTE1_M>(Mte1MmABEventId(abL0BufIter % 2));
|
||||
WaitFlag<HardEvent::MTE1_M>(Mte1MmABEventId(abL0BufIter % 2));
|
||||
|
||||
MmadParams mmadParams;
|
||||
mmadParams.m = mL1SizeAlign;
|
||||
mmadParams.n = nL1SizeAlign;
|
||||
mmadParams.k = kL0Size;
|
||||
mmadParams.cmatrixInitVal = (qsfaKL0 == 0 && k1 == 0);
|
||||
mmadParams.cmatrixSource = false;
|
||||
mmadParams.unitFlag = ((k1 == (kL1Loops - 1)) && (qsfaKL0 == (kL0Loops - 1))) ? 0b11 : 0b10;
|
||||
|
||||
Mmad(cL0Tensor, aL0Tensor, bL0Tensor, mmadParams);
|
||||
if ((mmadParams.m / 16) * (mmadParams.n / 16) < 10) {
|
||||
PipeBarrier<PIPE_M>();
|
||||
}
|
||||
SetFlag<HardEvent::M_MTE1>(Mte1MmABEventId(abL0BufIter % 2));
|
||||
abL0BufIter++;
|
||||
}
|
||||
|
||||
if (nL1 == (nL1Loops - 1)) { // nL1最后一轮, 需要将B驻留在L1中, 用于下一轮的计算?
|
||||
SetFlag<HardEvent::MTE1_MTE2>(mte21QPIds[ka]); // 反向同步, 表示L1中的A已经被mte1消费完
|
||||
}
|
||||
|
||||
if (k1 == (kL1Loops - 1)) {
|
||||
// ND
|
||||
FixpipeParamsV220 fixParams;
|
||||
fixParams.nSize = nL1SizeAlign;
|
||||
fixParams.mSize = mL1SizeAlign;
|
||||
fixParams.srcStride = mL1SizeAlign;
|
||||
fixParams.dstStride = nSize; // mm2ResGm两行之间的间隔
|
||||
fixParams.ndNum = 1; // 输出ND
|
||||
fixParams.unitFlag = 0b11;
|
||||
|
||||
uint64_t qsfaMm2Offset = (mSplitInfo.nBufferStartM + qsfaML1 * M_SPLIT_SIZE) * nSize +
|
||||
nL1 * N_SPLIT_SIZE;
|
||||
Fixpipe(mm2ResGm[(info.loop % (constInfo.preLoadNum)) *
|
||||
constInfo.bmm2ResUbSize + qsfaMm2Offset], cL0Tensor, fixParams);
|
||||
}
|
||||
|
||||
if (mL1Loops == 2) {
|
||||
cL0BufIter++;
|
||||
}
|
||||
}
|
||||
SetFlag<HardEvent::MTE1_MTE2>(mte21KVIds[qsfaKb]); // 反向同步, 表示L1已经被mte1消费完
|
||||
}
|
||||
// cL0BufIter已经不在使用
|
||||
if (mL1Loops == 1) {
|
||||
cL0BufIter++;
|
||||
}
|
||||
}
|
||||
qpL1BufIter += mL1Loops;
|
||||
}
|
||||
|
||||
#endif // KV_QUANT_SPARSE_FLASH_ATTENTION_SERVICE_CUBE_MLA_H
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,82 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_flash_attention_template_tiling_key.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef KV_QUANT_SPARSE_FLASH_ATTENTION_TEMPLATE_TILING_KEY_H
|
||||
#define KV_QUANT_SPARSE_FLASH_ATTENTION_TEMPLATE_TILING_KEY_H
|
||||
|
||||
#include "ascendc/host_api/tiling/template_argument.h"
|
||||
|
||||
#define QSFA_LAYOUT_BSND 0
|
||||
#define QSFA_LAYOUT_TND 1
|
||||
#define QSFA_LAYOUT_PA_BSND 2
|
||||
|
||||
#define ASCENDC_TPL_4_BW 4
|
||||
|
||||
#define C_TEMPLATE 0
|
||||
#define V_TEMPLATE 1
|
||||
|
||||
// 模板参数支持的范围定义
|
||||
ASCENDC_TPL_ARGS_DECL(KvQuantSparseFlashAttention, // 算子OpType
|
||||
ASCENDC_TPL_BOOL_DECL(FLASH_DECODE, 0, 1),
|
||||
ASCENDC_TPL_BOOL_DECL(PAGE_ATTENTION, 0, 1),
|
||||
ASCENDC_TPL_UINT_DECL(LAYOUT_T, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST,
|
||||
QSFA_LAYOUT_BSND, QSFA_LAYOUT_TND),
|
||||
ASCENDC_TPL_UINT_DECL(KV_LAYOUT_T, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST,
|
||||
QSFA_LAYOUT_BSND, QSFA_LAYOUT_TND, QSFA_LAYOUT_PA_BSND),
|
||||
ASCENDC_TPL_UINT_DECL(TEMPLATE_MODE, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST, C_TEMPLATE, V_TEMPLATE),
|
||||
ASCENDC_TPL_BOOL_DECL(IS_SPLIT_G, 0, 1),
|
||||
);
|
||||
|
||||
// 支持的模板参数组合
|
||||
// 用于调用GET_TPL_TILING_KEY获取TilingKey时,接口内部校验TilingKey是否合法
|
||||
ASCENDC_TPL_SEL(
|
||||
ASCENDC_TPL_ARGS_SEL(
|
||||
ASCENDC_TPL_BOOL_SEL(FLASH_DECODE, 0),
|
||||
ASCENDC_TPL_BOOL_SEL(PAGE_ATTENTION, 0),
|
||||
ASCENDC_TPL_UINT_SEL(LAYOUT_T, ASCENDC_TPL_UI_LIST, QSFA_LAYOUT_BSND),
|
||||
ASCENDC_TPL_UINT_SEL(KV_LAYOUT_T, ASCENDC_TPL_UI_LIST, QSFA_LAYOUT_BSND),
|
||||
ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, V_TEMPLATE),
|
||||
ASCENDC_TPL_BOOL_SEL(IS_SPLIT_G, 0, 1),
|
||||
),
|
||||
|
||||
ASCENDC_TPL_ARGS_SEL(
|
||||
ASCENDC_TPL_BOOL_SEL(FLASH_DECODE, 0),
|
||||
ASCENDC_TPL_BOOL_SEL(PAGE_ATTENTION, 0),
|
||||
ASCENDC_TPL_UINT_SEL(LAYOUT_T, ASCENDC_TPL_UI_LIST, QSFA_LAYOUT_TND),
|
||||
ASCENDC_TPL_UINT_SEL(KV_LAYOUT_T, ASCENDC_TPL_UI_LIST, QSFA_LAYOUT_TND),
|
||||
ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, V_TEMPLATE),
|
||||
ASCENDC_TPL_BOOL_SEL(IS_SPLIT_G, 0, 1),
|
||||
),
|
||||
|
||||
ASCENDC_TPL_ARGS_SEL(
|
||||
ASCENDC_TPL_BOOL_SEL(FLASH_DECODE, 0),
|
||||
ASCENDC_TPL_BOOL_SEL(PAGE_ATTENTION, 1),
|
||||
ASCENDC_TPL_UINT_SEL(LAYOUT_T, ASCENDC_TPL_UI_LIST, QSFA_LAYOUT_BSND),
|
||||
ASCENDC_TPL_UINT_SEL(KV_LAYOUT_T, ASCENDC_TPL_UI_LIST, QSFA_LAYOUT_PA_BSND),
|
||||
ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, V_TEMPLATE),
|
||||
ASCENDC_TPL_BOOL_SEL(IS_SPLIT_G, 0, 1),
|
||||
),
|
||||
|
||||
ASCENDC_TPL_ARGS_SEL(
|
||||
ASCENDC_TPL_BOOL_SEL(FLASH_DECODE, 0),
|
||||
ASCENDC_TPL_BOOL_SEL(PAGE_ATTENTION, 1),
|
||||
ASCENDC_TPL_UINT_SEL(LAYOUT_T, ASCENDC_TPL_UI_LIST, QSFA_LAYOUT_TND),
|
||||
ASCENDC_TPL_UINT_SEL(KV_LAYOUT_T, ASCENDC_TPL_UI_LIST, QSFA_LAYOUT_PA_BSND),
|
||||
ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, V_TEMPLATE),
|
||||
ASCENDC_TPL_BOOL_SEL(IS_SPLIT_G, 0, 1),
|
||||
),
|
||||
);
|
||||
|
||||
#endif // TEMPLATE_TILING_KEY
|
||||
Reference in New Issue
Block a user