* 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 test_aclnn_mixed_quant_sparse_flash_mla_metadata.cpp
*/
#include <iostream>
#include <vector>
#include <cmath>
#include <cstring>
#include <limits>
#include <functional>
#include <utility>
#include "acl/acl.h"
#include "aclnnop/aclnn_mixed_quant_sparse_flash_mla_metadata.h"
#define CHECK_LOG_RET(cond, ret_val, fmt, ...) \
do { \
if (!(cond)) { \
printf(fmt "\n", ##__VA_ARGS__); \
return (ret_val); \
} \
} while (0)
constexpr uint32_t AIC_CORE_MAX_NUM = 36;
constexpr uint32_t AIV_CORE_MAX_NUM = 72;
constexpr uint32_t MQSMLA_METADATA_TOTAL_SIZE = 1024;
constexpr uint32_t FA_METADATA_SIZE = 9;
constexpr uint32_t FD_METADATA_SIZE = 8;
constexpr uint32_t FA_CORE_ENABLE_INDEX = 0;
constexpr uint32_t FA_BN2_START_INDEX = 1;
constexpr uint32_t FA_M_START_INDEX = 2;
constexpr uint32_t FA_S2_START_INDEX = 3;
constexpr uint32_t FA_BN2_END_INDEX = 4;
constexpr uint32_t FA_M_END_INDEX = 5;
constexpr uint32_t FA_S2_END_INDEX = 6;
constexpr uint32_t FA_FIRST_FD_DATA_WORKSPACE_IDX_INDEX = 7;
constexpr uint32_t FA_S2_MAX_NUM = 8;
constexpr uint32_t FD_CORE_ENABLE_INDEX = 0;
constexpr uint32_t FD_BN2_IDX_INDEX = 1;
constexpr uint32_t FD_M_IDX_INDEX = 2;
constexpr uint32_t FD_WORKSPACE_IDX_INDEX = 3;
constexpr uint32_t FD_WORKSPACE_NUM_INDEX = 4;
constexpr uint32_t FD_M_START_INDEX = 5;
constexpr uint32_t FD_M_NUM_INDEX = 6;
struct MqsmlaMetadata {
uint32_t faMetadata[AIC_CORE_MAX_NUM][FA_METADATA_SIZE];
uint32_t fdMetadata[AIV_CORE_MAX_NUM][FD_METADATA_SIZE];
};
struct ScopeGuard
{
explicit ScopeGuard(std::function<void()> onExitScope) : m_exitFunc(std::move(onExitScope)),
m_isDismissed(false) {}
ScopeGuard(const ScopeGuard&) = delete;
ScopeGuard& operator=(const ScopeGuard&) = delete;
~ScopeGuard()
{
if (!m_isDismissed) {
m_exitFunc();
}
}
void Dismiss()
{
m_isDismissed = true;
}
std::function<void()> m_exitFunc;
bool m_isDismissed;
};
struct Tensor {
void *hostAddr { nullptr };
void *deviceAddr { nullptr };
aclTensor *data { nullptr };
};
struct ArgScenario {
bool hasCuSeq { false };
bool hasSeqused { false };
};
struct ArgContext {
int64_t numHeadsQ { 0 };
int64_t numHeadsKv { 0 };
int64_t headDim { 0 };
int64_t quantMode { 1 };
Tensor cuSeqlensQOptional {};
Tensor cuSeqlensOriKvOptional {};
Tensor cuSeqlensCmpKvOptional {};
Tensor sequsedQOptional {};
Tensor sequsedOriKvOptional {};
Tensor sequsedCmpKvOptional {};
Tensor cmpResidualKvOptional {};
Tensor oriTopkLengthOptional {};
Tensor cmpTopkLengthOptional {};
int64_t batchSize { 0 };
int64_t maxSeqlenQ { 0 };
int64_t maxSeqlenOriKv { 0 };
int64_t maxSeqlenCmpKv { 0 };
int64_t oriTopk { 0 };
int64_t cmpTopk { 0 };
int64_t ropeHeadDim { 64 };
int64_t cmpRatio { 0 };
int64_t oriMaskMode { 0 };
int64_t cmpMaskMode { 0 };
int64_t oriWinLeft { -1 };
int64_t oriWinRight { -1 };
char *layoutQOptional { nullptr };
char *layoutKvOptional { nullptr };
bool hasOriKv { true };
bool hasCmpKv { true };
Tensor metadata {};
};
int64_t GetShapeSize(const std::vector<int64_t>& shape)
{
int64_t shapeSize = 1;
for (auto i : shape) {
shapeSize *= i;
}
return shapeSize;
}
aclnnStatus Init(int32_t deviceId, aclrtStream* stream)
{
auto ret = aclInit(nullptr);
CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclInit failed. ERROR: %d", ret);
ret = aclrtSetDevice(deviceId);
CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtSetDevice failed. ERROR: %d", ret);
ret = aclrtCreateStream(stream);
CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtCreateStream failed. ERROR: %d", ret);
return ACL_SUCCESS;
}
void Finalize(int32_t deviceId, aclrtStream stream)
{
aclrtDestroyStream(stream);
aclrtResetDevice(deviceId);
aclFinalize();
}
aclnnStatus CreateTensor(aclDataType dataType, const std::vector<int64_t> &shape, Tensor &tensor)
{
auto size = GetShapeSize(shape) * aclDataTypeSize(dataType);
auto ret = aclrtMallocHost(&(tensor.hostAddr), size);
CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMallocHost failed. ERROR: %d", ret);
memset(tensor.hostAddr, 0, size);
ret = aclrtMalloc(&(tensor.deviceAddr), size, ACL_MEM_MALLOC_HUGE_FIRST);
CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMalloc failed. ERROR: %d", ret);
tensor.data = aclCreateTensor(shape.data(), shape.size(), dataType, nullptr, 0, aclFormat::ACL_FORMAT_ND,
shape.data(), shape.size(), tensor.deviceAddr);
CHECK_LOG_RET(tensor.data != nullptr, ACL_ERROR_FAILURE, "aclCreateTensor failed");
ret = aclrtMemcpy(tensor.deviceAddr, size, tensor.hostAddr, size, ACL_MEMCPY_HOST_TO_DEVICE);
CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMemcpy failed. ERROR: %d", ret);
return ACL_SUCCESS;
}
void DestroyTensor(Tensor &tensor)
{
if (tensor.data != nullptr) {
aclDestroyTensor(tensor.data);
tensor.data = nullptr;
}
if (tensor.deviceAddr != nullptr) {
aclrtFree(tensor.deviceAddr);
tensor.deviceAddr = nullptr;
}
if (tensor.hostAddr != nullptr) {
aclrtFreeHost(tensor.hostAddr);
tensor.hostAddr = nullptr;
}
}
void DestroyArgs(ArgContext &context)
{
DestroyTensor(context.metadata);
DestroyTensor(context.cuSeqlensQOptional);
DestroyTensor(context.cuSeqlensOriKvOptional);
DestroyTensor(context.cuSeqlensCmpKvOptional);
DestroyTensor(context.sequsedQOptional);
DestroyTensor(context.sequsedOriKvOptional);
DestroyTensor(context.sequsedCmpKvOptional);
DestroyTensor(context.cmpResidualKvOptional);
DestroyTensor(context.oriTopkLengthOptional);
DestroyTensor(context.cmpTopkLengthOptional);
if (context.layoutQOptional != nullptr) {
free(context.layoutQOptional);
context.layoutQOptional = nullptr;
}
if (context.layoutKvOptional != nullptr) {
free(context.layoutKvOptional);
context.layoutKvOptional = nullptr;
}
}
aclnnStatus CreateArgs(const ArgScenario &scenario, ArgContext &context)
{
ScopeGuard argsGuard([&] { DestroyArgs(context); });
aclnnStatus ret;
context.numHeadsQ = 64;
context.numHeadsKv = 1;
context.headDim = 512;
context.quantMode = 1;
ret = CreateTensor(aclDataType::ACL_INT32, { MQSMLA_METADATA_TOTAL_SIZE }, context.metadata);
CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create metadata failed. Error: %d", ret);
context.oriTopk = 0;
context.cmpTopk = 0;
context.ropeHeadDim = 64;
context.cmpRatio = 128;
context.oriMaskMode = 4;
context.cmpMaskMode = 3;
context.oriWinLeft = 127;
context.oriWinRight = 0;
context.layoutQOptional = (char *)malloc(sizeof(char) * 16);
context.layoutKvOptional = (char *)malloc(sizeof(char) * 16);
CHECK_LOG_RET(context.layoutQOptional != nullptr, ACL_ERROR_FAILURE, "Create layoutQOptional failed");
CHECK_LOG_RET(context.layoutKvOptional != nullptr, ACL_ERROR_FAILURE, "Create layoutKvOptional failed");
strcpy(context.layoutQOptional, "BSND");
strcpy(context.layoutKvOptional, "BSND");
context.hasOriKv = true;
context.hasCmpKv = true;
context.batchSize = 4;
context.maxSeqlenOriKv = 1024;
context.maxSeqlenCmpKv = 1024;
context.maxSeqlenQ = 1024;
if (scenario.hasCuSeq) {
ret = CreateTensor(aclDataType::ACL_INT32, { context.batchSize + 1 }, context.cuSeqlensQOptional);
CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create cuSeqlensQOptional failed. Error: %d", ret);
ret = CreateTensor(aclDataType::ACL_INT32, { context.batchSize + 1 }, context.cuSeqlensOriKvOptional);
CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create cuSeqlensOriKvOptional failed. Error: %d", ret);
ret = CreateTensor(aclDataType::ACL_INT32, { context.batchSize + 1 }, context.cuSeqlensCmpKvOptional);
CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create cuSeqlensCmpKvOptional failed. Error: %d", ret);
}
if (scenario.hasSeqused) {
ret = CreateTensor(aclDataType::ACL_INT32, { context.batchSize }, context.sequsedQOptional);
CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create sequsedQOptional failed. Error: %d", ret);
ret = CreateTensor(aclDataType::ACL_INT32, { context.batchSize }, context.sequsedOriKvOptional);
CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create sequsedOriKvOptional failed. Error: %d", ret);
ret = CreateTensor(aclDataType::ACL_INT32, { context.batchSize }, context.sequsedCmpKvOptional);
CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create sequsedCmpKvOptional failed. Error: %d", ret);
}
if (context.hasCmpKv && context.cmpRatio != 1 && context.cmpMaskMode == 3) {
ret = CreateTensor(aclDataType::ACL_INT32, { context.batchSize }, context.cmpResidualKvOptional);
CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create cmpResidualKvOptional failed. Error: %d", ret);
}
argsGuard.Dismiss();
return ACL_SUCCESS;
}
int main() {
int32_t deviceId = 0;
aclrtStream stream;
auto ret = Init(deviceId, &stream);
CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Init acl failed. ERROR: %d", ret);
ScopeGuard sysGuard([&] { Finalize(deviceId, stream); });
ArgScenario scenario {};
scenario.hasCuSeq = false;
scenario.hasSeqused = false;
ArgContext context {};
ret = CreateArgs(scenario, context);
CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create input arguments failed. ERROR: %d", ret);
ScopeGuard argsGuard([&] { DestroyArgs(context); });
uint64_t workspaceSize = 0;
aclOpExecutor *executor = nullptr;
void *workspaceAddr = nullptr;
ret = aclnnMixedQuantSparseFlashMlaMetadataGetWorkspaceSize(
context.cuSeqlensQOptional.data, context.cuSeqlensOriKvOptional.data, context.cuSeqlensCmpKvOptional.data,
context.sequsedQOptional.data, context.sequsedOriKvOptional.data, context.sequsedCmpKvOptional.data,
context.cmpResidualKvOptional.data, context.oriTopkLengthOptional.data, context.cmpTopkLengthOptional.data,
context.numHeadsQ, context.numHeadsKv, context.headDim, context.quantMode, context.batchSize,
context.maxSeqlenQ, context.maxSeqlenOriKv, context.maxSeqlenCmpKv, context.oriTopk, context.cmpTopk,
context.ropeHeadDim, context.cmpRatio, context.oriMaskMode, context.cmpMaskMode, context.oriWinLeft,
context.oriWinRight, context.layoutQOptional, context.layoutKvOptional, context.hasOriKv, context.hasCmpKv,
context.metadata.data, &workspaceSize, &executor);
CHECK_LOG_RET(ret == ACL_SUCCESS, ret,
"aclnnMixedQuantSparseFlashMlaMetadataGetWorkspaceSize failed. ERROR: %d\n", ret);
if (workspaceSize > static_cast<uint64_t>(0)) {
ret = aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST);
CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "allocate workspace failed. ERROR: %d\n", ret);
}
ScopeGuard workspaceGuard([&] {
if (workspaceAddr != nullptr) {
aclrtFree(workspaceAddr);
workspaceAddr = nullptr;
}
});
ret = aclnnMixedQuantSparseFlashMlaMetadata(workspaceAddr, workspaceSize, executor, stream);
CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclnnMixedQuantSparseFlashMlaMetadata failed. ERROR: %d\n", ret);
ret = aclrtSynchronizeStream(stream);
CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtSynchronizeStream failed. ERROR: %d\n", ret);
MqsmlaMetadata result {};
ret = aclrtMemcpy(&result, sizeof(result), context.metadata.deviceAddr, sizeof(result), ACL_MEMCPY_DEVICE_TO_HOST);
CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMemcpy failed. ERROR: %d\n", ret);
for (uint32_t i = 0; i < AIC_CORE_MAX_NUM; ++i) {
printf("AIC Core%u\n", i);
printf(" Core Enable : %u\n", result.faMetadata[i][FA_CORE_ENABLE_INDEX]);
printf(" Start BN2 : %u\n", result.faMetadata[i][FA_BN2_START_INDEX]);
printf(" Start M : %u\n", result.faMetadata[i][FA_M_START_INDEX]);
printf(" Start S2 : %u\n", result.faMetadata[i][FA_S2_START_INDEX]);
printf(" End BN2 : %u\n", result.faMetadata[i][FA_BN2_END_INDEX]);
printf(" End M : %u\n", result.faMetadata[i][FA_M_END_INDEX]);
printf(" End S2 : %u\n", result.faMetadata[i][FA_S2_END_INDEX]);
printf(" First Worksapce Index : %u\n", result.faMetadata[i][FA_FIRST_FD_DATA_WORKSPACE_IDX_INDEX]);
printf(" Max S2 Block Num : %u\n", result.faMetadata[i][FA_S2_MAX_NUM]);
}
for (uint32_t i = 0; i < AIV_CORE_MAX_NUM; ++i) {
printf("AIV Core%u\n", i);
printf(" Core Enable : %u\n", result.fdMetadata[i][FD_CORE_ENABLE_INDEX]);
printf(" FD Task BN2 Idx : %u\n", result.fdMetadata[i][FD_BN2_IDX_INDEX]);
printf(" FD Task M Idx : %u\n", result.fdMetadata[i][FD_M_IDX_INDEX]);
printf(" FD Task S2 Idx : %u\n", result.fdMetadata[i][FD_WORKSPACE_IDX_INDEX]);
printf(" FD Task Workspace Num : %u\n", result.fdMetadata[i][FD_WORKSPACE_NUM_INDEX]);
printf(" FD Subtask M Start : %u\n", result.fdMetadata[i][FD_M_START_INDEX]);
printf(" FD Subtask M Num : %u\n", result.fdMetadata[i][FD_M_NUM_INDEX]);
}
return 0;
}