aclnnMixedQuantSparseFlashMla

产品支持情况

产品 是否支持
Ascend 950PR/Ascend 950DT
Atlas A3 训练系列产品/Atlas A3 推理系列产品 ×
Atlas A2 训练系列产品/Atlas A2 推理系列产品 ×
Atlas 200I/500 A2 推理产品 ×
Atlas 推理系列产品 ×
Atlas 训练系列产品 ×

功能说明

  • 接口功能:

    aclnnMixedQuantSparseFlashMla算子实现基于共享KV(Key=Value)的稀疏注意力计算,支持SWA(Sliding Window Attention)、CSA(Compressed Sparse Attention)、HCA(Heavily Compressed Attention)三类Attention计算场景。与SparseFlashMla的区别在于,本算子支持KV的per-token-group量化输入。该算子适用于大语言模型推理场景,通过滑动窗口和KV压缩机制大幅降低长序列注意力计算的开销。调用时需要使用aclnnMixedQuantSparseFlashMlaMetadata生成的任务列表metadata

    该算子不建议单独使用,建议与aclnnMixedQuantSparseFlashMlaMetadata算子配合使用,形成完整的工作流。

    典型调用流程如下:

    1. 准备qori_kvcmp_kv、序列长度、block tablesinks等输入。
    2. 调用aclnnMixedQuantSparseFlashMlaMetadata生成metadata
    3. 调用aclnnMixedQuantSparseFlashMla,将上一步得到的metadata传入主算子。
  • 计算公式:

    O=softmax(Q⋅K~T⋅softmax_scale)⋅V~O = \text{softmax}(Q \cdot \tilde{K}^T \cdot \text{softmax\_scale}) \cdot \tilde{V}

    其中K~=V~\tilde{K} = \tilde{V}(共享KV),K~\tilde{K}由滑动窗口内的原始KV和因果边界内的压缩KV拼接而成,具体参与计算的KV范围由模板模式和mask参数决定:

    • 滑动窗口部分(oriKv):对第iS1i_{S1}个Query token,其因果对角线位置为ori_threshold=S2act−S1act+iS1+1\text{ori\_threshold} = S2_{act} - S1_{act} + i_{S1} + 1,窗口范围为[max⁡(ori_threshold−ori_win_left−1,0),ori_threshold+ori_win_right)[\max(\text{ori\_threshold} - \text{ori\_win\_left} - 1, 0), \text{ori\_threshold} + \text{ori\_win\_right})

    • 压缩KV部分(cmpKv):因果边界阈值为cmp_threshold=⌊ori_thresholdcmp_ratio⌋\text{cmp\_threshold} = \lfloor \frac{\text{ori\_threshold}}{\text{cmp\_ratio}} \rfloor。HCA场景取[0,cmp_threshold)[0, \text{cmp\_threshold})内的连续压缩KV;CSA场景通过TopK索引从压缩KV中按需收集,仅保留begin_idx<cmp_threshold\text{begin\_idx} < \text{cmp\_threshold}的块。

    注意力计算采用Online Softmax(Flash Attention V2),S2方向按512分块循环,sinks作为每行softmax的初始最大值:

    row_max(0)=sinks[g],row_sum(0)=1.0,O(0)=0\text{row\_max}^{(0)} = \text{sinks}[g], \quad \text{row\_sum}^{(0)} = 1.0, \quad O^{(0)} = 0

    S(t)=Q⋅Ktile(t)T⋅softmax_scaleS^{(t)} = Q \cdot K_{tile}^{(t)T} \cdot \text{softmax\_scale}

    row_max(t+1)=max⁡(row_max(t),max⁡(S(t),dim=−1))\text{row\_max}^{(t+1)} = \max(\text{row\_max}^{(t)}, \max(S^{(t)}, \text{dim}=-1))

    row_sum(t+1)=exp⁡(row_max(t)−row_max(t+1))⋅row_sum(t)+∑exp⁡(S(t)−row_max(t+1))\text{row\_sum}^{(t+1)} = \exp(\text{row\_max}^{(t)} - \text{row\_max}^{(t+1)}) \cdot \text{row\_sum}^{(t)} + \sum \exp(S^{(t)} - \text{row\_max}^{(t+1)})

    O(t+1)=exp⁡(row_max(t)−row_max(t+1))⋅O(t)+exp⁡(S(t)−row_max(t+1))⋅Vtile(t)O^{(t+1)} = \exp(\text{row\_max}^{(t)} - \text{row\_max}^{(t+1)}) \cdot O^{(t)} + \exp(S^{(t)} - \text{row\_max}^{(t+1)}) \cdot V_{tile}^{(t)}

    Ofinal=O(Ts2)/row_sum(Ts2)O_{final} = O^{(T_{s2})} / \text{row\_sum}^{(T_{s2})}

  • 符号说明

    符号 含义
    Q Query输入,形状为[G, D](单行)
    K_tile_t 第t个S2分块的KV数据,K=V(共享KV)
    S_t 第t个分块的QK缩放注意力分数
    P_t 第t个分块的softmax概率
    O_t 第t个分块后的累加输出
    softmax_scale 缩放系数,通常取每个注意力头维度的倒数平方根
    B Batch Size
    S1/S1_act Query序列长度/实际有效长度
    S2/S2_act 原始KV序列长度/实际有效长度
    N1 Query头数
    N2 KV头数
    G GQA分组比,G=N1/N2
    D 每个注意力头的维度
    sinks 注意力汇点,形状为[N1]
    cmp_ratio cmpKv的压缩倍率,用于换算cmp侧mask的压缩前KV长度

函数原型

每个算子分为两段式接口,必须先调用aclnnMixedQuantSparseFlashMlaGetWorkspaceSize接口获取计算所需workspace大小以及包含了算子计算流程的执行器,再调用aclnnMixedQuantSparseFlashMla执行实际计算。

aclnnStatus aclnnMixedQuantSparseFlashMlaGetWorkspaceSize(
    const aclTensor *q,
    const aclTensor *oriKvOptional,
    const aclTensor *cmpKvOptional,
    const aclTensor *oriSparseIndicesOptional,
    const aclTensor *cmpSparseIndicesOptional,
    const aclTensor *oriBlockTableOptional,
    const aclTensor *cmpBlockTableOptional,
    const aclTensor *cuSeqlensQOptional,
    const aclTensor *cuSeqlensOriKvOptional,
    const aclTensor *cuSeqlensCmpKvOptional,
    const aclTensor *sequsedQOptional,
    const aclTensor *sequsedOriKvOptional,
    const aclTensor *sequsedCmpKvOptional,
    const aclTensor *cmpResidualKvOptional,
    const aclTensor *oriTopkLengthOptional,
    const aclTensor *cmpTopkLengthOptional,
    const aclTensor *sinksOptional,
    const aclTensor *metadataOptional,
    int64_t          quantMode,
    int64_t          ropeHeadDim,
    double           softmaxScale,
    int64_t          cmpRatio,
    int64_t          oriMaskMode,
    int64_t          cmpMaskMode,
    int64_t          oriWinLeft,
    int64_t          oriWinRight,
    char            *layoutQOptional,
    char            *layoutKvOptional,
    int64_t          topkValueMode,
    bool             returnSoftmaxLse,
    const aclTensor *attnOutOut,
    const aclTensor *softmaxLseOutOptional,
    uint64_t        *workspaceSize,
    aclOpExecutor  **executor)
aclnnStatus aclnnMixedQuantSparseFlashMla(
    void          *workspace,
    uint64_t       workspaceSize,
    aclOpExecutor *executor,
    aclrtStream    stream)

aclnnMixedQuantSparseFlashMlaGetWorkspaceSize

  • 参数说明

    参数名 输入/输出 描述 使用说明 数据类型 数据格式 维度(shape) 非连续Tensor
    q(aclTensor*) 输入 Query输入张量。 不支持空Tensor。N1/N2仅支持2、4、8、16、32、64、128范围内的2的幂;D仅支持512。 BFLOAT16 ND
    • layoutQ为BSND时:(B, S1, N1, D)
    • layoutQ为TND时:(T1, N1, D)
    oriKvOptional(aclTensor*) 输入 原始KV输入张量,Key与Value共享同一份数据。 SWA/CSA/HCA场景必须传入。量化KV布局由quantMode决定:quant_mode为1时,依次由rope(64,bfloat16)、nope(448,FLOAT8_E4M3FN)、scale(7,bfloat16)、pad(18B)拼接而成;quant_mode为2时,依次由nope(448,FLOAT8_E4M3FN)、rope(64,bfloat16)、scale(7,FLOAT8_E8M0)、pad(1B)拼接而成。当前仅支持1和2,量化模式2仅支持layout_kv为PA_BBND。 详见quantMode ND
    • layoutKv为PA_BBND时:(ori_block_num, ori_block_size, N2, D_KV),ori_block_size支持1到1024
    • layoutKv为BSND时:(B, S2, N2, D_KV)
    • layoutKv为TND时:(T2, N2, D_KV)
    N2仅支持1,D_KV由quantMode决定。
    cmpKvOptional(aclTensor*) 输入 压缩KV输入张量,Key与Value共享同一份数据。 CSA/HCA场景必须传入,SWA场景不传入。量化KV布局由quantMode决定,同oriKvOptional。 详见quantMode ND
    • layoutKv为PA_BBND时:(cmp_block_num, cmp_block_size, N2, D_KV),cmp_block_size支持1到1024
    • layoutKv为BSND时:(B, S3, N2, D_KV)
    • layoutKv为TND时:(T3, N2, D_KV)
    N2仅支持1,D_KV由quantMode决定。
    oriSparseIndicesOptional(aclTensor*) 输入 代表离散取oriKvCache的索引。 当前暂不支持,必须传入nullptr。 INT32 ND -
    cmpSparseIndicesOptional(aclTensor*) 输入 代表离散取cmpKvCache的TopK索引。 CSA场景必须传入,SWA/HCA场景不传入。 INT32 ND
    • layoutQ为BSND时:(B, S1, N2, K2)
    • layoutQ为TND时:(T1, N2, K2)
    其中K2为cmpKv的TopK稀疏选择数。
    oriBlockTableOptional(aclTensor*) 输入 PageAttention中oriKvCache存储使用的block映射表。 layoutKv为PA_BBND时必须传入。第二维长度不小于所有batch中最大的S2对应的block数量。 INT32 ND (B, ori_max_block_num_per_batch)
    cmpBlockTableOptional(aclTensor*) 输入 PageAttention中cmpKvCache存储使用的block映射表。 CSA/HCA场景且layoutKv为PA_BBND时必须传入。 INT32 ND (B, cmp_max_block_num_per_batch)
    cuSeqlensQOptional(aclTensor*) 输入 表示不同Batch中q的有效token数(前缀和形式)。 layoutQOptional为TND时必须传入。每个元素表示当前batch与之前所有batch的token数总和。 INT32 ND (B+1,)
    cuSeqlensOriKvOptional(aclTensor*) 输入 表示不同Batch中oriKv的有效token数(前缀和形式)。 layoutKvOptional为TND时必须传入。 INT32 ND (B+1,)
    cuSeqlensCmpKvOptional(aclTensor*) 输入 表示不同Batch中cmpKv的有效token数(前缀和形式)。 layoutKvOptional为TND且存在cmpKvOptional时必须传入。 INT32 ND (B+1,)
    sequsedQOptional(aclTensor*) 输入 表示不同Batch中q实际参与运算的token数。 当前暂不支持指定该参数。 INT32 ND (B,)
    sequsedOriKvOptional(aclTensor*) 输入 表示不同Batch中oriKv实际参与运算的token数。 layoutKvOptional为PA_BBND时必须传入;layoutKvOptional为BSND时可选传入,用于指定每个batch的oriKv有效长度;layoutKvOptional为TND时使用cuSeqlensOriKvOptional表达序列边界。 INT32 ND (B,)
    sequsedCmpKvOptional(aclTensor*) 输入 表示不同Batch中cmpKv实际参与运算的token数。 可选输入。传入时shape必须为(B,),作为每个batch的cmp逻辑有效长度,优先于cmpKvOptional shape、cuSeqlensCmpKvOptional或PA block table推导;layoutKvOptional为BSND、TND、PA_BBND时均可使用。 INT32 ND (B,)
    cmpResidualKvOptional(aclTensor*) 输入 压缩KV余数,用于恢复cmp侧mask使用的压缩前KV长度。 可选输入。传入时shape必须为(B,),第b个batch按cmp_len * cmpRatio + cmpResidualKvOptional[b]恢复压缩前KV长度;在CSA/HCA、cmpRatio不等于1且cmpMaskMode为3场景必传。该参数是主算子和aclnnMixedQuantSparseFlashMlaMetadata的可选入参,layoutKvOptional为BSND、TND、PA_BBND时均可使用。 INT32 ND (B,)
    oriTopkLengthOptional(aclTensor*) 输入 预留输入,当前版本不支持传入非空Tensor。 必须传入nullptr或空Tensor;传入非空Tensor会返回参数错误。 INT32 ND -
    cmpTopkLengthOptional(aclTensor*) 输入 预留输入,当前版本不支持传入非空Tensor。 必须传入nullptr或空Tensor;传入非空Tensor会返回参数错误。 INT32 ND -
    sinksOptional(aclTensor*) 输入 注意力汇点tensor,作为每行softmax的初始最大值。 必须传入。 FLOAT32 ND (N1,)
    metadataOptional(aclTensor*) 输入 AICPU算子aclnnMixedQuantSparseFlashMlaMetadata的分核结果。 必须传入。由aclnnMixedQuantSparseFlashMlaMetadata算子生成。 INT32 ND (1024,)
    quantMode(int64_t) 输入 表示量化模式。 量化模式1表示K、V nope为per-token-group量化,scale类型为bfloat16,量化模式2表示K、V nope为per-token-group量化,scale类型为FLOAT8_E8M0。当前仅支持1和2。 - - - -
    ropeHeadDim(int64_t) 输入 表示rope头的维度。 仅支持64。 - - - -
    softmaxScale(double) 输入 缩放系数,对应公式中的softmaxScale。 建议值为1/√D,其中D为每个注意力头的维度。 - - - -
    cmpRatio(int64_t) 输入 cmpKv相对于压缩前KV长度的压缩倍率,用于恢复cmp侧mask使用的压缩前KV长度。 支持1到128。 - - - -
    oriMaskMode(int64_t) 输入 q和oriKv计算的mask模式。 仅支持4: Band模式。 - - - -
    cmpMaskMode(int64_t) 输入 q和cmpKv计算的mask模式。 仅支持3: RightDownCausal模式。SWA场景下该参数不生效。 - - - -
    oriWinLeft(int64_t) 输入 q和oriKv计算中,在因果边界基础上向左多看的token数。 仅支持127。 - - - -
    oriWinRight(int64_t) 输入 q和oriKv计算中,在因果边界基础上向右多看的token数。 仅支持0。 - - - -
    layoutQOptional(char*) 输入 标识输入q的数据排布格式。 支持"BSND"和"TND"。 - - - -
    layoutKvOptional(char*) 输入 标识输入oriKvOptional和cmpKvOptional的数据排布格式。 支持"PA_BBND"、"BSND"和"TND"。 - - - -
    topkValueMode(int64_t) 输入 topk索引取值模式。 当前支持1。 - - - -
    returnSoftmaxLse(bool) 输入 是否返回softmaxLse。 支持true或false。 - - - -
    attnOutOut(aclTensor*) 输出 注意力计算输出。 - BFLOAT16 ND 与q的shape一致 ×
    softmaxLseOutOptional(aclTensor*) 输出 softmax的log-sum-exp结果。 returnSoftmaxLse为false时返回占位Tensor;returnSoftmaxLse为true时返回softmax的log-sum-exp结果。 FLOAT32 ND
    • layoutQ为BSND时:(B, N2, S1, N1/N2)
    • layoutQ为TND时:(N2, T1, N1/N2)
    • returnSoftmaxLse为false时:占位Tensor
    ×
    workspaceSize(uint64_t*) 输出 返回需要在Device侧申请的workspace大小。 - - - - -
    executor(aclOpExecutor**) 输出 返回op执行器,包含了算子计算流程。 - - - - -
    • Ascend 950PR/Ascend 950DT:N1/N2支持2、4、8、16、32、64、128,不支持1。
  • 返回值

    aclnnStatus:返回状态码,具体参见aclnn返回码

    第一段接口完成入参校验,出现以下场景时报错:

    返回值 错误码 描述
    ACLNN_ERR_PARAM_NULLPTR 161001 传入参数是必选输入、输出或者必选属性,且是空指针。
    ACLNN_ERR_PARAM_INVALID 161002 输入变量的数据类型和数据格式不在支持的范围内。
    N1不在[2,4,8,16,32,64,128]范围内,或N2不为1。
    D不为512。
    quantMode不为1或2。
    oriMaskMode不为4,或cmpMaskMode不为3。
    oriWinLeft不为127,或oriWinRight不为0。
    SWA场景cmpRatio不为1,或cmpRatio与CSA/HCA场景不匹配。
    layoutQOptional、layoutKvOptional、topkValueMode、cmpSparseIndicesOptional、metadataOptional、sinksOptional、cuSeqlens或seqused相关参数规格不在支持范围内。

aclnnMixedQuantSparseFlashMla

  • 参数说明

    参数名 输入/输出 描述
    workspace 输入 在Device侧申请的workspace内存地址。
    workspaceSize 输入 在Device侧申请的workspace大小,由第一段接口aclnnMixedQuantSparseFlashMlaGetWorkspaceSize获取。
    executor 输入 op执行器,包含了算子计算流程。
    stream 输入 指定执行任务的Stream。
  • 返回值

    返回aclnnStatus状态码,具体参见aclnn返回码

约束说明

  • 确定性计算

    • aclnnMixedQuantSparseFlashMla默认采用确定性实现,相同输入多次调用结果一致。
  • 使用约束

    • 本算子仅支持推理场景。
    • oriTopkLengthOptionalcmpTopkLengthOptional等预留输入可传入nullptr或空Tensor外,其余已传入Tensor不支持为空。
    • metadataOptional参数必须传入,由aclnnMixedQuantSparseFlashMlaMetadata算子生成,shape固定为(1024,)。
    • cmpResidualKvOptional为主算子和aclnnMixedQuantSparseFlashMlaMetadata的可选入参;传入后用于按cmp_len * cmpRatio + residual恢复cmp侧mask使用的压缩前长度。
    • ropeHeadDim仅支持64。
    • oriMaskMode仅支持4,cmpMaskMode仅支持3,oriWinLeft仅支持127,oriWinRight仅支持0。
  • 三种Attention场景输入要求

    场景 oriKvOptional cmpKvOptional cmpSparseIndicesOptional 说明
    SWA 必须传入 不传入 不传入 仅滑动窗口注意力
    CSA 必须传入 必须传入 必须传入 滑动窗口 + TopK稀疏压缩KV
    HCA 必须传入 必须传入 不传入 滑动窗口 + 稠密压缩KV
  • cmpRatio约束:SWA场景仅支持1。

  • Layout约束

    • layoutQOptionallayoutKvOptional组合仅支持"BSND"/"BSND"、"TND"/"TND"、"BSND"/"PA_BBND"、"TND"/"PA_BBND";非PA_BBND场景下layoutQOptionallayoutKvOptional必须一致。
    • layoutQOptional为TND时,cuSeqlensQOptional必须传入。
    • layoutKvOptional为PA_BBND时,sequsedOriKvOptional必须传入,oriBlockTableOptional必须传入。BSND场景可选传入sequsedOriKvOptional覆盖每个batch的oriKv有效长度;TND场景使用cuSeqlensOriKvOptional表达oriKv序列边界。
    • layoutKvOptional为TND时,cuSeqlensOriKvOptional必须传入。
    • layoutKvOptional为TND且存在cmpKvOptional时,cuSeqlensCmpKvOptional必须传入。
    • sequsedCmpKvOptional为所有layoutKvOptional下的可选输入,显式传入时用于覆盖cmp侧逻辑有效长度。

调用示例

调用示例代码如下,仅供参考,具体编译和执行过程请参考编译与运行样例

/**
 * Copyright (c) 2026 Huawei Technologies Co., Ltd.
 * This file is a part of the CANN Open Software.
 * Licensed under CANN Open Software License Agreement Version 1.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 test_aclnn_mixed_quant_sparse_flash_mla.cpp
 * \brief
 */

#include <algorithm>
#include <cmath>
#include <cstdint>
#include <cstring>
#include <iostream>
#include <numeric>
#include <random>
#include <vector>
#include "acl/acl.h"
#include "aclnnop/aclnn_mixed_quant_sparse_flash_mla.h"
#include "aclnnop/aclnn_mixed_quant_sparse_flash_mla_metadata.h"

#define CHECK_RET(cond, return_expr) \
  do {                               \
    if (!(cond)) {                   \
      return_expr;                   \
    }                                \
  } while (0)

#define LOG_PRINT(message, ...)     \
  do {                              \
    printf(message, ##__VA_ARGS__); \
  } while (0)

namespace {

int64_t GetShapeSize(const std::vector<int64_t>& shape)
{
  int64_t shapeSize = 1;
  for (auto i : shape) {
    shapeSize *= i;
  }
  return shapeSize;
}

uint16_t FloatToBf16(float f)
{
  uint32_t bits;
  std::memcpy(&bits, &f, sizeof(bits));
  uint32_t lsb = (bits >> 16) & 1u;
  uint32_t roundingBias = 0x7fffu + lsb;
  bits += roundingBias;
  return static_cast<uint16_t>(bits >> 16);
}

float Bf16ToFloat(uint16_t h)
{
  uint32_t bits = static_cast<uint32_t>(h) << 16;
  float result;
  std::memcpy(&result, &bits, sizeof(result));
  return result;
}

void PrintOutResult(const std::vector<int64_t>& shape, void** deviceAddr)
{
  auto size = GetShapeSize(shape);
  std::vector<uint16_t> resultData(size, 0);
  auto ret = aclrtMemcpy(resultData.data(), resultData.size() * sizeof(resultData[0]),
                         *deviceAddr, size * sizeof(resultData[0]), ACL_MEMCPY_DEVICE_TO_HOST);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("copy result from device to host failed. ERROR: %d\n", ret); return);
  for (int64_t i = 0; i < size && i < 10; i++) {
    LOG_PRINT("result[%ld] is: %f\n", i, Bf16ToFloat(resultData[i]));
  }
}

int Init(int32_t deviceId, aclrtContext* context, aclrtStream* stream)
{
  auto ret = aclInit(nullptr);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclInit failed. ERROR: %d\n", ret); return ret);
  ret = aclrtSetDevice(deviceId);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSetDevice failed. ERROR: %d\n", ret); return ret);
  ret = aclrtCreateContext(context, deviceId);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtCreateContext failed. ERROR: %d\n", ret); return ret);
  ret = aclrtSetCurrentContext(*context);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSetCurrentContext failed. ERROR: %d\n", ret); return ret);
  ret = aclrtCreateStream(stream);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtCreateStream failed. ERROR: %d\n", ret); return ret);
  return 0;
}

template <typename T>
int CreateAclTensor(const std::vector<T>& hostData, const std::vector<int64_t>& shape, void** deviceAddr,
                    aclDataType dataType, aclTensor** tensor)
{
  auto size = GetShapeSize(shape) * sizeof(T);
  if (size > 0) {
    auto ret = aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST);
    CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtMalloc failed. ERROR: %d\n", ret); return ret);
    ret = aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE);
    CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtMemcpy failed. ERROR: %d\n", ret); return ret);
  } else {
    *deviceAddr = nullptr;
  }

  std::vector<int64_t> strides(shape.size(), 1);
  for (int64_t i = static_cast<int64_t>(shape.size()) - 2; i >= 0; i--) {
    strides[i] = shape[i + 1] * strides[i + 1];
  }

  *tensor = aclCreateTensor(shape.data(), shape.size(), dataType, strides.data(), 0, aclFormat::ACL_FORMAT_ND,
                            shape.data(), shape.size(), *deviceAddr);
  return 0;
}

std::vector<uint16_t> MakeBf16Data(int64_t size, float value)
{
  std::vector<uint16_t> data(static_cast<size_t>(size), FloatToBf16(value));
  return data;
}

std::vector<uint8_t> MakeFp8Data(int64_t size, uint8_t value)
{
  std::vector<uint8_t> data(static_cast<size_t>(size), value);
  return data;
}

}  // namespace

int main()
{
  int32_t deviceId = 0;
  aclrtContext context = nullptr;
  aclrtStream stream = nullptr;
  auto ret = Init(deviceId, &context, &stream);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("Init acl failed. ERROR: %d\n", ret); return ret);

  int64_t B = 1;
  int64_t S1 = 1;
  int64_t S2 = 1024;
  int64_t N1 = 64;
  int64_t N2 = 1;
  int64_t D = 512;
  int64_t K = 512;
  int64_t oriBlockSize = 128;
  int64_t cmpBlockSize = 128;
  int64_t s2Act = 1024;
  int64_t cmpRatio = 4;
  int64_t oriWinLeft = 127;
  int64_t oriWinRight = 0;
  int64_t oriMaskMode = 4;
  int64_t cmpMaskMode = 3;
  int64_t quantMode = 1;
  int64_t ropeHeadDim = 64;
  int64_t tileSize = 64;
  double softmaxScale = 1.0 / sqrt(static_cast<double>(D));

  int64_t nopeHeadDim = D - ropeHeadDim;
  int64_t quantScaleHeadDim = (nopeHeadDim + tileSize - 1) / tileSize;
  int64_t kvD = nopeHeadDim + ropeHeadDim * 2 + quantScaleHeadDim * 2 + 18;

  int64_t T1 = B * S1;
  int64_t cmpKvLen = s2Act / cmpRatio;
  int64_t oriBlockNum = ((s2Act + oriBlockSize - 1) / oriBlockSize) * B;
  int64_t cmpBlockNum = ((cmpKvLen + cmpBlockSize - 1) / cmpBlockSize) * B;

  std::vector<int64_t> qShape = {T1, N1, D};
  std::vector<int64_t> oriKvShape = {oriBlockNum, oriBlockSize, N2, kvD};
  std::vector<int64_t> cmpKvShape = {cmpBlockNum, cmpBlockSize, N2, kvD};
  std::vector<int64_t> cmpSparseIndicesShape = {T1, N2, K};
  std::vector<int64_t> oriBlockTableShape = {B, (s2Act + oriBlockSize - 1) / oriBlockSize};
  std::vector<int64_t> cmpBlockTableShape = {B, (cmpKvLen + cmpBlockSize - 1) / cmpBlockSize};
  std::vector<int64_t> cuSeqLensQShape = {B + 1};
  std::vector<int64_t> seqUsedOriKvShape = {B};
  std::vector<int64_t> seqUsedCmpKvShape = {B};
  std::vector<int64_t> cmpResidualKvShape = {B};
  std::vector<int64_t> sinksShape = {N1};
  std::vector<int64_t> metadataShape = {1024};
  std::vector<int64_t> attnOutShape = {T1, N1, D};
  std::vector<int64_t> softmaxLseShape = {T1, N1, 1};
  std::vector<int64_t> emptyShape = {0};

  void* qDeviceAddr = nullptr;
  void* oriKvDeviceAddr = nullptr;
  void* cmpKvDeviceAddr = nullptr;
  void* cmpSparseIndicesDeviceAddr = nullptr;
  void* oriBlockTableDeviceAddr = nullptr;
  void* cmpBlockTableDeviceAddr = nullptr;
  void* cuSeqLensQDeviceAddr = nullptr;
  void* cuSeqLensOriKvDeviceAddr = nullptr;
  void* cuSeqLensCmpKvDeviceAddr = nullptr;
  void* seqUsedQDeviceAddr = nullptr;
  void* seqUsedOriKvDeviceAddr = nullptr;
  void* seqUsedCmpKvDeviceAddr = nullptr;
  void* cmpResidualKvDeviceAddr = nullptr;
  void* sinksDeviceAddr = nullptr;
  void* metadataDeviceAddr = nullptr;
  void* attnOutDeviceAddr = nullptr;
  void* softmaxLseDeviceAddr = nullptr;

  aclTensor* q = nullptr;
  aclTensor* oriKv = nullptr;
  aclTensor* cmpKv = nullptr;
  aclTensor* cmpSparseIndices = nullptr;
  aclTensor* oriBlockTable = nullptr;
  aclTensor* cmpBlockTable = nullptr;
  aclTensor* cuSeqLensQ = nullptr;
  aclTensor* cuSeqLensOriKv = nullptr;
  aclTensor* cuSeqLensCmpKv = nullptr;
  aclTensor* seqUsedQ = nullptr;
  aclTensor* seqUsedOriKv = nullptr;
  aclTensor* seqUsedCmpKv = nullptr;
  aclTensor* cmpResidualKv = nullptr;
  aclTensor* sinks = nullptr;
  aclTensor* metadata = nullptr;
  aclTensor* attnOut = nullptr;
  aclTensor* softmaxLse = nullptr;

  int64_t qSize = GetShapeSize(qShape);
  int64_t oriKvSize = GetShapeSize(oriKvShape);
  int64_t cmpKvSize = GetShapeSize(cmpKvShape);
  int64_t cmpSparseIndicesSize = GetShapeSize(cmpSparseIndicesShape);
  int64_t oriBlockTableSize = GetShapeSize(oriBlockTableShape);
  int64_t cmpBlockTableSize = GetShapeSize(cmpBlockTableShape);
  int64_t attnOutSize = GetShapeSize(attnOutShape);
  int64_t softmaxLseSize = GetShapeSize(softmaxLseShape);

  std::vector<uint16_t> qHostData = MakeBf16Data(qSize, 1.0f);
  std::vector<uint8_t> oriKvHostData = MakeFp8Data(oriKvSize, 0x38);
  std::vector<uint8_t> cmpKvHostData = MakeFp8Data(cmpKvSize, 0x38);
  std::vector<int32_t> cmpSparseIndicesHostData(cmpSparseIndicesSize);
  std::vector<int32_t> oriBlockTableHostData(oriBlockTableSize);
  std::iota(oriBlockTableHostData.begin(), oriBlockTableHostData.end(), 0);
  std::vector<int32_t> cmpBlockTableHostData(cmpBlockTableSize);
  std::iota(cmpBlockTableHostData.begin(), cmpBlockTableHostData.end(), 0);
  std::vector<int32_t> cuSeqLensQHostData(B + 1);
  for (int64_t i = 0; i <= B; i++) {
    cuSeqLensQHostData[i] = static_cast<int32_t>(i * S1);
  }
  std::vector<int32_t> emptyHostData;
  std::vector<int32_t> seqUsedOriKvHostData(B, static_cast<int32_t>(s2Act));
  std::vector<int32_t> seqUsedCmpKvHostData(B, static_cast<int32_t>(cmpKvLen));
  std::vector<int32_t> cmpResidualKvHostData(B);
  for (int64_t i = 0; i < B; i++) {
    cmpResidualKvHostData[i] = seqUsedOriKvHostData[i] % static_cast<int32_t>(cmpRatio);
  }
  std::vector<float> sinksHostData(N1, 1.0f);
  std::vector<int32_t> metadataHostData(1024, 0);
  std::vector<uint16_t> attnOutHostData = MakeBf16Data(attnOutSize, 0.0f);
  std::vector<float> softmaxLseHostData(softmaxLseSize, 0.0f);

  std::mt19937 gen(42);
  for (int64_t t = 0; t < T1; t++) {
    for (int64_t n = 0; n < N2; n++) {
      for (int64_t k = 0; k < K; k++) {
        cmpSparseIndicesHostData[t * N2 * K + n * K + k] = static_cast<int32_t>(gen() % cmpKvLen);
      }
    }
  }

  ret = CreateAclTensor(qHostData, qShape, &qDeviceAddr, aclDataType::ACL_BF16, &q);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(oriKvHostData, oriKvShape, &oriKvDeviceAddr, aclDataType::ACL_FLOAT8_E4M3FN, &oriKv);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(cmpKvHostData, cmpKvShape, &cmpKvDeviceAddr, aclDataType::ACL_FLOAT8_E4M3FN, &cmpKv);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(cmpSparseIndicesHostData, cmpSparseIndicesShape, &cmpSparseIndicesDeviceAddr,
                        aclDataType::ACL_INT32, &cmpSparseIndices);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(oriBlockTableHostData, oriBlockTableShape, &oriBlockTableDeviceAddr, aclDataType::ACL_INT32,
                        &oriBlockTable);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(cmpBlockTableHostData, cmpBlockTableShape, &cmpBlockTableDeviceAddr, aclDataType::ACL_INT32,
                        &cmpBlockTable);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(cuSeqLensQHostData, cuSeqLensQShape, &cuSeqLensQDeviceAddr, aclDataType::ACL_INT32, &cuSeqLensQ);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(emptyHostData, emptyShape, &cuSeqLensOriKvDeviceAddr, aclDataType::ACL_INT32, &cuSeqLensOriKv);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(emptyHostData, emptyShape, &cuSeqLensCmpKvDeviceAddr, aclDataType::ACL_INT32, &cuSeqLensCmpKv);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(emptyHostData, emptyShape, &seqUsedQDeviceAddr, aclDataType::ACL_INT32, &seqUsedQ);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(seqUsedOriKvHostData, seqUsedOriKvShape, &seqUsedOriKvDeviceAddr, aclDataType::ACL_INT32, &seqUsedOriKv);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(seqUsedCmpKvHostData, seqUsedCmpKvShape, &seqUsedCmpKvDeviceAddr, aclDataType::ACL_INT32, &seqUsedCmpKv);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(cmpResidualKvHostData, cmpResidualKvShape, &cmpResidualKvDeviceAddr, aclDataType::ACL_INT32, &cmpResidualKv);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(sinksHostData, sinksShape, &sinksDeviceAddr, aclDataType::ACL_FLOAT, &sinks);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(metadataHostData, metadataShape, &metadataDeviceAddr, aclDataType::ACL_INT32, &metadata);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(attnOutHostData, attnOutShape, &attnOutDeviceAddr, aclDataType::ACL_BF16, &attnOut);
  CHECK_RET(ret == ACL_SUCCESS, return ret);
  ret = CreateAclTensor(softmaxLseHostData, softmaxLseShape, &softmaxLseDeviceAddr, aclDataType::ACL_FLOAT, &softmaxLse);
  CHECK_RET(ret == ACL_SUCCESS, return ret);

  char layoutQ[] = "TND";
  char layoutKv[] = "PA_BBND";

  uint64_t metadataWorkspaceSize = 0;
  aclOpExecutor* metadataExecutor = nullptr;

  ret = aclnnMixedQuantSparseFlashMlaMetadataGetWorkspaceSize(
      cuSeqLensQ, cuSeqLensOriKv, cuSeqLensCmpKv,
      seqUsedQ, seqUsedOriKv, seqUsedCmpKv,
      cmpResidualKv, nullptr, nullptr,
      N1, N2, D, quantMode, B,
      S1, S2, cmpKvLen,
      0, K, ropeHeadDim,
      cmpRatio, oriMaskMode, cmpMaskMode,
      oriWinLeft, oriWinRight,
      layoutQ, layoutKv,
      true, true,
      metadata,
      &metadataWorkspaceSize, &metadataExecutor);
  CHECK_RET(ret == ACL_SUCCESS,
            LOG_PRINT("aclnnMixedQuantSparseFlashMlaMetadataGetWorkspaceSize failed. ERROR: %d\n", ret); return ret);

  void* metadataWorkspaceAddr = nullptr;
  if (metadataWorkspaceSize > 0) {
    ret = aclrtMalloc(&metadataWorkspaceAddr, metadataWorkspaceSize, ACL_MEM_MALLOC_HUGE_FIRST);
    CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("allocate metadata workspace failed. ERROR: %d\n", ret); return ret);
  }

  ret = aclnnMixedQuantSparseFlashMlaMetadata(metadataWorkspaceAddr, metadataWorkspaceSize, metadataExecutor, stream);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnMixedQuantSparseFlashMlaMetadata failed. ERROR: %d\n", ret); return ret);

  ret = aclrtSynchronizeStream(stream);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSynchronizeStream after metadata failed. ERROR: %d\n", ret); return ret);

  uint64_t workspaceSize = 0;
  aclOpExecutor* executor = nullptr;

  ret = aclnnMixedQuantSparseFlashMlaGetWorkspaceSize(
      q, oriKv, cmpKv,
      nullptr, cmpSparseIndices,
      oriBlockTable, cmpBlockTable,
      cuSeqLensQ, nullptr, nullptr,
      nullptr, seqUsedOriKv, seqUsedCmpKv,
      cmpResidualKv, nullptr, nullptr,
      sinks, metadata,
      quantMode, ropeHeadDim,
      softmaxScale, cmpRatio,
      oriMaskMode, cmpMaskMode,
      oriWinLeft, oriWinRight,
      layoutQ, layoutKv,
      1, false,
      attnOut, softmaxLse,
      &workspaceSize, &executor);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnMixedQuantSparseFlashMlaGetWorkspaceSize failed. ERROR: %d\n", ret); return ret);

  void* workspaceAddr = nullptr;
  if (workspaceSize > 0) {
    ret = aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST);
    CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("allocate workspace failed. ERROR: %d\n", ret); return ret);
  }

  ret = aclnnMixedQuantSparseFlashMla(workspaceAddr, workspaceSize, executor, stream);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnMixedQuantSparseFlashMla failed. ERROR: %d\n", ret); return ret);

  ret = aclrtSynchronizeStream(stream);
  CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSynchronizeStream failed. ERROR: %d\n", ret); return ret);

  PrintOutResult(attnOutShape, &attnOutDeviceAddr);

  aclDestroyTensor(q);
  aclDestroyTensor(oriKv);
  aclDestroyTensor(cmpKv);
  aclDestroyTensor(cmpSparseIndices);
  aclDestroyTensor(oriBlockTable);
  aclDestroyTensor(cmpBlockTable);
  aclDestroyTensor(cuSeqLensQ);
  aclDestroyTensor(cuSeqLensOriKv);
  aclDestroyTensor(cuSeqLensCmpKv);
  aclDestroyTensor(seqUsedQ);
  aclDestroyTensor(seqUsedOriKv);
  aclDestroyTensor(seqUsedCmpKv);
  aclDestroyTensor(cmpResidualKv);
  aclDestroyTensor(sinks);
  aclDestroyTensor(metadata);
  aclDestroyTensor(attnOut);
  aclDestroyTensor(softmaxLse);

  aclrtFree(qDeviceAddr);
  aclrtFree(oriKvDeviceAddr);
  aclrtFree(cmpKvDeviceAddr);
  aclrtFree(cmpSparseIndicesDeviceAddr);
  aclrtFree(oriBlockTableDeviceAddr);
  aclrtFree(cmpBlockTableDeviceAddr);
  if (cuSeqLensQDeviceAddr != nullptr) {
    aclrtFree(cuSeqLensQDeviceAddr);
  }
  if (seqUsedOriKvDeviceAddr != nullptr) {
    aclrtFree(seqUsedOriKvDeviceAddr);
  }
  if (seqUsedCmpKvDeviceAddr != nullptr) {
    aclrtFree(seqUsedCmpKvDeviceAddr);
  }
  if (cmpResidualKvDeviceAddr != nullptr) {
    aclrtFree(cmpResidualKvDeviceAddr);
  }
  aclrtFree(sinksDeviceAddr);
  aclrtFree(metadataDeviceAddr);
  aclrtFree(attnOutDeviceAddr);
  aclrtFree(softmaxLseDeviceAddr);
  if (metadataWorkspaceSize > 0) {
    aclrtFree(metadataWorkspaceAddr);
  }
  if (workspaceSize > 0) {
    aclrtFree(workspaceAddr);
  }
  aclrtDestroyStream(stream);
  aclrtDestroyContext(context);
  aclrtResetDevice(deviceId);
  aclFinalize();

  return 0;
}