/**
 * 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 load_data_perf.asc
 * \brief LoadData 性能测试示例
 */
#ifdef ASCENDC_CPU_DEBUG
#include "cpu_debug_launch.h"
#endif
#include "acl/acl.h"
#include "kernel_operator.h"
#include <cstdio>
#include <cstdlib>

constexpr uint32_t BLOCK_CUBE = 16;
constexpr uint32_t C0_SIZE = 16;
constexpr uint64_t ADDR_0 = 0;
constexpr uint32_t SCALE_BASE_FACTOR = 64;
constexpr uint32_t SCALE_EVEN_NUMBER = 2;
constexpr uint32_t SCALE_CEIL_NUMBER = 32;

template <class T>
class KernelPerf {
public:
    __aicore__ inline KernelPerf(uint32_t scenario, uint32_t m, uint32_t k, uint32_t n)
    {
        mSize = m;
        kSize = k;
        nSize = n;
        scenarioNum = scenario;
        fractalShape[0] = 16;
        fractalShape[1] = 32 / sizeof(T);
        c0Size = C0_SIZE;
        fractalNum = 1;

        aSizeAlignL1 = CeilAlign(mSize, BLOCK_CUBE) * CeilAlign(kSize, c0Size);
        aSizeAlignL0 = CeilAlign(mSize, BLOCK_CUBE) * CeilAlign(kSize, c0Size);

        bSizeAlignL1 = CeilAlign(kSize, BLOCK_CUBE) * CeilAlign(nSize, c0Size);
        bSizeAlignL0 = CeilAlign(kSize, c0Size) * CeilAlign(nSize, BLOCK_CUBE);

        if (scenarioNum == 3) {
            aSizeAlignL1 = CeilAlign(kSize, BLOCK_CUBE) * CeilAlign(mSize, C0_SIZE);
            aSizeAlignL0 = CeilAlign(mSize, BLOCK_CUBE) * CeilAlign(kSize, C0_SIZE);
        }
        if (scenarioNum == 4) {
            bSizeAlignL1 = CeilAlign(nSize, BLOCK_CUBE) * CeilAlign(kSize, c0Size);
            bSizeAlignL0 = CeilAlign(kSize, c0Size) * CeilAlign(nSize, BLOCK_CUBE);
        }

        scaleK = CeilDiv(kSize, SCALE_BASE_FACTOR) * SCALE_EVEN_NUMBER;

        fractalSize = BLOCK_CUBE * c0Size;
        a1Addr = ADDR_0;
        b1Addr = a1Addr + aSizeAlignL1 * sizeof(T);
        c1Addr = a1Addr + b1Addr + bSizeAlignL1 * sizeof(T);
#if defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3510)
        // Ascend 950PR/Ascend 950DTC1:C1 到 C2 块长度为 32B,b32 对齐
        biasSizeAlign = CeilAlign(CeilDiv(nSize * sizeof(float), 32), 2) * 32;
#else
        // Atlas A3/Atlas A2:C1 到 C2 块长度为 64B
        biasSizeAlign = CeilAlign(nSize * sizeof(float), 64);
#endif

        // L1/L0 Zz 格式 scaleA 对齐
        scaleMAlignL1 = CeilAlign(mSize, fractalShape[0]);
        scaleASizeAlignL1 = scaleMAlignL1 * scaleK;

        // L1/L0 Nn 格式 scaleB 对齐
        scaleNAlignL1 = CeilAlign(nSize, fractalShape[0]);
        scaleBSizeAlignL1 = scaleK * scaleNAlignL1;
    }

    // 测试场景编号说明:
    // 场景 1-9:适用于 Atlas A3/A2 训练/推理平台(dav-2201 架构)
    // 场景 11-19:适用于 Ascend 950PR/950DT 平台(dav-3510 架构)
    //
    // 场景 1:LoadData(2D矩阵搬运)A - 从 L1 加载到 L0A
    // 场景 2:LoadData(2D矩阵搬运)B - 从 L1 加载到 L0B
    // 场景 3:LoadDataWithTranspose A - 从 L1 加载到 L0A(转置)
    // 场景 4:LoadDataWithTranspose B - 从 L1 加载到 L0B(转置)
    // 场景 5:LoadData(卷积数据搬运)v2 A - 从 L1 加载到 L0A
    // 场景 6:LoadData(卷积数据搬运)v2 B - 从 L1 加载到 L0B
    // 场景 7:LoadSparse - LoadDataWithSparse 从 L1 加载到 L0B(稀疏加载)
    // 场景 8:LoadBias - 从 L1 加载到 bias buffer
    // 场景 9:LoadFixBuffer - 从 L1 加载到 fixpipe buffer
    //
    // 场景 11:LoadData(2D矩阵搬运V2)A - 从 L1 加载到 L0A
    // 场景 12:LoadData(2D矩阵搬运V2)B - 从 L1 加载到 L0B
    // 场景 13:LoadData(MX矩阵搬运)A - 从 L1 加载到 L0A 和 L0A_MX(带 scale)
    // 场景 14:LoadData(MX矩阵搬运)B - 从 L1 加载到 L0B 和 L0B_MX(带 scale)
    // 场景 15:LoadData(卷积数据搬运)v2 A - 从 L1 加载到 L0A
    // 场景 16:LoadData(卷积数据搬运)v2 B - 从 L1 加载到 L0B
    // 场景 17:LoadDataWithTranspose B - 从 L1 加载到 L0B(转置)
    // 场景 18:LoadBias - 从 L1 加载到 bias buffer
    // 场景 19:LoadFixBuffer - 从 L1 加载到 fixpipe buffer

    __aicore__ inline void Process()
    {
        AscendC::LocalTensor<T> a1Local(AscendC::TPosition::A1, a1Addr, aSizeAlignL1);
        AscendC::LocalTensor<T> b1Local(AscendC::TPosition::B1, b1Addr, bSizeAlignL1);
        AscendC::LocalTensor<T> a2Local(AscendC::TPosition::A2, ADDR_0, aSizeAlignL0);
        AscendC::LocalTensor<T> b2Local(AscendC::TPosition::B2, ADDR_0, bSizeAlignL0);
        AscendC::PipeBarrier<PIPE_ALL>();
        if (scenarioNum == 1 || scenarioNum == 11) {
            Load2DA(a2Local, a1Local);
        } else if (scenarioNum == 2 || scenarioNum == 12) {
            Load2DB(b2Local, b1Local);
        } else if (scenarioNum == 3) {
            Load2DAtranspose(a2Local, a1Local);
        } else if (scenarioNum == 4 || scenarioNum == 17) {
            Load2DBtranspose(b2Local, b1Local);
        } else if (scenarioNum == 5 || scenarioNum == 15) {
            Load3DA(a2Local, a1Local);
        } else if (scenarioNum == 6 || scenarioNum == 16) {
            Load3DB(b2Local, b1Local);
        } else if (scenarioNum == 7) {
            AscendC::LocalTensor<int8_t> dst = b2Local.template ReinterpretCast<int8_t>();
            AscendC::LocalTensor<uint8_t> idxB1Local(
                AscendC::TPosition::B1, 256 * 1024, bSizeAlignL1 / 4); // from L1 256k
            AscendC::LocalTensor<int8_t> src = b1Local.template ReinterpretCast<int8_t>();
            LoadSparse(dst, src, idxB1Local);
        } else if (scenarioNum == 8 || scenarioNum == 18) {
            AscendC::LocalTensor<float> bias1Local(AscendC::TPosition::C1, c1Addr, biasSizeAlign / sizeof(float));
            AscendC::LocalTensor<float> bias2Local(AscendC::TPosition::C2, ADDR_0, biasSizeAlign / sizeof(float));
            LoadBias(bias2Local, bias1Local);
        } else if (scenarioNum == 9 || scenarioNum == 19) {
            AscendC::LocalTensor<uint64_t> quantAlphaTensor(AscendC::TPosition::C1, c1Addr, nSize);
            LoadFixBuffer(quantAlphaTensor);
        } else if (scenarioNum == 13) {
            // Ascend 950PR/950DT 平台 LoadData(MX矩阵搬运) a 操作
            AscendC::LocalTensor<fp8_e4m3fn_t> src = a1Local.template ReinterpretCast<fp8_e4m3fn_t>();
            AscendC::LocalTensor<fp8_e4m3fn_t> dst = a2Local.template ReinterpretCast<fp8_e4m3fn_t>();
            AscendC::LocalTensor<fp8_e8m0_t> scaleA1Local(AscendC::TPosition::A1, c1Addr, scaleASizeAlignL1);
            Load2DMxA(dst, src, scaleA1Local);
        } else if (scenarioNum == 14) {
            // Ascend 950PR/950DT 平台 LoadData(MX矩阵搬运) b 操作
            AscendC::LocalTensor<fp8_e4m3fn_t> src = b1Local.template ReinterpretCast<fp8_e4m3fn_t>();
            AscendC::LocalTensor<fp8_e4m3fn_t> dst = b2Local.template ReinterpretCast<fp8_e4m3fn_t>();
            AscendC::LocalTensor<fp8_e8m0_t> scaleB1Local(AscendC::TPosition::B1, c1Addr, scaleBSizeAlignL1);
            Load2DMxB(dst, src, scaleB1Local);
        }

        AscendC::PipeBarrier<PIPE_ALL>();
    }

private:
    // 向上取整除法
    __aicore__ inline uint32_t CeilDiv(uint32_t numerator, uint32_t denominator)
    {
        return (numerator + denominator - 1) / denominator;
    }
    // 向上对齐
    __aicore__ inline uint32_t CeilAlign(uint32_t numerator, uint32_t denominator)
    {
        return (numerator + denominator - 1) / denominator * denominator;
    }

    __aicore__ inline void Load2DA(AscendC::LocalTensor<T>& a2Local, AscendC::LocalTensor<T>& a1Local)
    {
#if defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3510)
        // Nz 转 Nz,NPU_ARCH 为 3510 时 A2 为 Nz 格式
        AscendC::LoadData2DParamsV2 loadDataParams;
        loadDataParams.mStartPosition = 0;
        loadDataParams.kStartPosition = 0;
        loadDataParams.mStep = CeilDiv(mSize, BLOCK_CUBE);
        loadDataParams.kStep = CeilDiv(kSize, c0Size);
        loadDataParams.srcStride = CeilDiv(mSize, BLOCK_CUBE);
        loadDataParams.dstStride = CeilDiv(mSize, BLOCK_CUBE);
        loadDataParams.ifTranspose = false;
        AscendC::LoadData(a2Local, a1Local, loadDataParams);
#else
        // Nz 转 Zz,NPU_ARCH 为 2201 时 A2 为 Zz 格式
        uint32_t dstOffset = CeilDiv(kSize, c0Size) * fractalSize;
        uint32_t srcOffset = fractalSize;
        AscendC::LoadData2DParams loadDataParams;
        loadDataParams.startIndex = 0;
        loadDataParams.repeatTimes = CeilDiv(kSize, c0Size);
        loadDataParams.srcStride = CeilDiv(mSize, BLOCK_CUBE);
        loadDataParams.dstGap = 0;
        loadDataParams.ifTranspose = false;
        for (int i = 0; i < CeilDiv(mSize, BLOCK_CUBE); ++i) {
            AscendC::LoadData(a2Local[i * dstOffset], a1Local[i * srcOffset], loadDataParams);
        }
#endif
    }

    __aicore__ inline void Load2DAtranspose(AscendC::LocalTensor<T>& a2Local, AscendC::LocalTensor<T>& a1Local)
    {
#if defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3510)
        // Nz 转 Nz,NPU_ARCH 为 3510 时 A2 为 Nz 格式
        AscendC::LoadData2DParamsV2 loadDataParams;
        loadDataParams.mStartPosition = 0;
        loadDataParams.kStartPosition = 0;
        // b4 时 mStep 为 4,b8 时 mStep 为 2,b32 时 kStep 为 2
        loadDataParams.mStep = CeilDiv(kSize, BLOCK_CUBE);
        loadDataParams.kStep = CeilAlign(CeilDiv(mSize, c0Size), 2);
        loadDataParams.srcStride = CeilDiv(kSize, BLOCK_CUBE);
        loadDataParams.dstStride = CeilDiv(mSize, c0Size * fractalNum);
        loadDataParams.ifTranspose = true;
        AscendC::LoadData(a2Local, a1Local, loadDataParams);
#else
        // Nz 转 Zz,NPU_ARCH 为 2201 时 A2 为 Zz 格式
        uint32_t dstOffset = CeilDiv(kSize, c0Size * fractalNum) * fractalSize * fractalNum;
        uint32_t srcOffset = fractalSize * fractalNum;

        AscendC::LoadData2dTransposeParams loadDataParams;
        loadDataParams.startIndex = 0;
        loadDataParams.repeatTimes = CeilDiv(kSize, c0Size * fractalNum);
        loadDataParams.srcStride = CeilDiv(mSize, c0Size * fractalNum);
        loadDataParams.dstGap = 0;
        loadDataParams.dstFracGap = 0;
        for (int i = 0; i < CeilDiv(mSize, c0Size * fractalNum); ++i) {
            AscendC::LoadDataWithTranspose(a2Local[i * dstOffset], a1Local[i * srcOffset], loadDataParams);
        }
#endif
    }

    __aicore__ inline void Load2DB(AscendC::LocalTensor<T>& b2Local, AscendC::LocalTensor<T>& b1Local)
    {
        // B2 为 Zn 格式
        AscendC::LoadData2DParams loadDataParams;
        loadDataParams.repeatTimes = CeilDiv(nSize, BLOCK_CUBE) * CeilDiv(kSize, c0Size);
        loadDataParams.srcStride = 1;
        loadDataParams.dstGap = 0;
        loadDataParams.ifTranspose = false;
        AscendC::LoadData(b2Local, b1Local, loadDataParams);
    }

    __aicore__ inline void Load2DBtranspose(AscendC::LocalTensor<T>& b2Local, AscendC::LocalTensor<T>& b1Local)
    {
        uint32_t dstOffset = CeilDiv(nSize, BLOCK_CUBE * fractalNum) * fractalSize * fractalNum;
        uint32_t srcOffset = fractalSize * fractalNum;
        AscendC::LoadData2dTransposeParams loadDataParams;
        loadDataParams.startIndex = 0;
        loadDataParams.repeatTimes = CeilDiv(nSize, c0Size);
        loadDataParams.srcStride = CeilDiv(kSize, BLOCK_CUBE * fractalNum);
        loadDataParams.dstGap = 0;
        loadDataParams.dstFracGap = 0;
        for (int i = 0; i < CeilDiv(kSize, BLOCK_CUBE * fractalNum); ++i) {
            AscendC::LoadDataWithTranspose(b2Local[i * dstOffset], b1Local[i * srcOffset], loadDataParams);
        }
    }

    // A1 -> A2: LoadData A 到 L0A 和 scaleA 到 L0A_MX,A 默认不转置 [M,K]
    // LoadData2DParamsV2 用于 A,LoadData2DMxParams 用于 scaleA
    __aicore__ inline void Load2DMxA(
        AscendC::LocalTensor<fp8_e4m3fn_t> a2Local, AscendC::LocalTensor<fp8_e4m3fn_t> a1Local,
        AscendC::LocalTensor<fp8_e8m0_t> scaleA1Local)
    {
#if defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3510)
        fractalShape[1] = 32;
        AscendC::LoadData2DParamsV2 loadDataParams;
        loadDataParams.sid = 0;
        // A 在 L1 的起始位置,row0 和 col0 对齐 32B
        loadDataParams.mStartPosition = 0;
        loadDataParams.kStartPosition = 0;

        AscendC::LoadData2DMxParams loadMxDataParams;
        // scaleA 在 L1 的起始位置,row0 和 col0 对齐 32B
        loadMxDataParams.xStartPosition = 0;
        loadMxDataParams.yStartPosition = 0;

        // A[m, k] 从 L1 加载到 L0A
        // mStep/kStep 为 row/col 方向步进 32B
        loadDataParams.mStep = CeilDiv(mSize, fractalShape[0]);
        loadDataParams.kStep = CeilDiv(kSize, fractalShape[1]);
        // srcStride/dstStride 为 col 方向步进 512B
        loadDataParams.srcStride = CeilDiv(mSize, fractalShape[0]);
        loadDataParams.dstStride = CeilDiv(mSize, fractalShape[0]);
        loadDataParams.ifTranspose = false;

        // xStep/yStep 为 scaleA 的 row/col 方向 stride,row 方向步进
        loadMxDataParams.xStep = CeilDiv(scaleMAlignL1, fractalShape[0]);
        loadMxDataParams.yStep = scaleK / SCALE_EVEN_NUMBER;
        loadMxDataParams.srcStride = scaleK / SCALE_EVEN_NUMBER;
        loadMxDataParams.dstStride = scaleK / SCALE_EVEN_NUMBER;
        AscendC::LoadData(a2Local, a1Local, scaleA1Local, loadDataParams, loadMxDataParams);
#endif
    }

    // B1 -> B2: LoadData B 到 L0B 和 scaleB 到 L0B_MX,B 默认转置 [K,N]
    // LoadData2DParamsV2 用于 B,LoadData2DMxParams 用于 scaleB
    __aicore__ inline void Load2DMxB(
        AscendC::LocalTensor<fp8_e4m3fn_t> b2Local, AscendC::LocalTensor<fp8_e4m3fn_t> b1Local,
        AscendC::LocalTensor<fp8_e8m0_t> scaleB1Local)
    {
#if defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3510)
        fractalShape[1] = 32;
        AscendC::LoadData2DParamsV2 loadDataParams;
        loadDataParams.sid = 0;
        // B 在 L1 的起始位置,row0 和 col0 对齐 32B
        loadDataParams.mStartPosition = 0;
        loadDataParams.kStartPosition = 0;

        AscendC::LoadData2DMxParams loadMxDataParams;
        // scaleB 在 L1 的起始位置,row0 和 col0 对齐 32B
        loadMxDataParams.xStartPosition = 0;
        loadMxDataParams.yStartPosition = 0;

        // L1 加载到 L0B
        // mStep/kStep 为 row/col 方向步进 32B
        loadDataParams.mStep = CeilDiv(nSize, fractalShape[0]);
        loadDataParams.kStep = CeilDiv(kSize, fractalShape[1]);
        // srcStride/dstStride 为 col 方向步进 512B
        loadDataParams.srcStride = CeilDiv(nSize, fractalShape[0]);
        loadDataParams.dstStride = CeilDiv(nSize, fractalShape[0]);
        loadDataParams.ifTranspose = false;

        // xStep/yStep 为 scaleB 的 row/col 方向 stride,row 方向步进
        loadMxDataParams.xStep = CeilDiv(scaleNAlignL1, fractalShape[0]);
        loadMxDataParams.yStep = scaleK / SCALE_EVEN_NUMBER;
        loadMxDataParams.srcStride = scaleK / SCALE_EVEN_NUMBER;
        loadMxDataParams.dstStride = scaleK / SCALE_EVEN_NUMBER;
        AscendC::LoadData(b2Local, b1Local, scaleB1Local, loadDataParams, loadMxDataParams);
#endif
    }

    // A 使用 LoadData(卷积数据搬运)v2 加载
    __aicore__ inline void Load3DA(AscendC::LocalTensor<T>& a2Local, AscendC::LocalTensor<T>& a1Local)
    {
        AscendC::LoadData3DParamsV2<T> loadDataParams;
        // 高度方向
        loadDataParams.l1H = 1;
        // 宽度方向
        loadDataParams.l1W = CeilAlign(mSize, fractalShape[0]);
        // img2col 的 ho * wo,ho=wo=ho * wo = loadDataParams.l1H * loadDataParams.l1w
        // img2col 的 ci * kh * kw,kh=1,kw=1,ci=loadDataParams.channelSize = m
        loadDataParams.channelSize = CeilAlign(kSize, fractalShape[1]);
        // 宽度扩展,half 为 16,int8_t/uint8_t 为 32
        loadDataParams.kExtension = CeilAlign(kSize, fractalShape[1]);
        // 高度扩展,half/int8_t/uint8_t 为 16
        loadDataParams.mExtension = CeilAlign(mSize, fractalShape[0]);
        // 宽度方向步长
        loadDataParams.strideW = 1;
        // 高度方向步长
        loadDataParams.strideH = 1;
        // 卷积核宽度
        loadDataParams.filterW = 1;
        // 卷积核高度
        loadDataParams.filterH = 1;
        // 卷积核宽度方向膨胀
        loadDataParams.dilationFilterW = 1;
        // 卷积核高度方向膨胀
        loadDataParams.dilationFilterH = 1;
        loadDataParams.filterSizeW = false;
        loadDataParams.filterSizeH = false;
        loadDataParams.enTranspose = false;
        loadDataParams.fMatrixCtrl = false;
#if defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3510)
        // 使用 SetLoadDataRepeatWithStride 设置 LoadData(卷积数据搬运)v2 的 repeat 参数
        AscendC::LoadDataRepeatParamWithStride repeatParams;
        repeatParams.repeatTime = 1;                              // 重复次数(高度/宽度)
        repeatParams.repeatStride = 1;                            // 重复步进(高度/宽度)
        repeatParams.repeatMode = 0;                              // 0:高度方向;1:宽度方向
        repeatParams.dstStride = CeilDiv(mSize, fractalShape[0]); // K 方向步进
        AscendC::SetLoadDataRepeatWithStride(repeatParams);
        AscendC::LoadDataWithStride(a2Local, a1Local, loadDataParams);
#else
        AscendC::LoadData(a2Local, a1Local, loadDataParams);
#endif
    }

    // B 使用 LoadData(卷积数据搬运)v2 加载
    __aicore__ inline void Load3DB(AscendC::LocalTensor<T>& b2Local, AscendC::LocalTensor<T>& b1Local)
    {
        AscendC::LoadData3DParamsV2<T> loadDataParams;
        loadDataParams.l1H = 1;
        loadDataParams.l1W = CeilAlign(kSize, fractalShape[0]);
        loadDataParams.channelSize = CeilAlign(nSize, fractalShape[1]);
        loadDataParams.kExtension = CeilAlign(nSize, fractalShape[1]);
        loadDataParams.mExtension = CeilAlign(kSize, fractalShape[0]);
        loadDataParams.strideW = 1;
        loadDataParams.strideH = 1;
        loadDataParams.filterW = 1;
        loadDataParams.filterH = 1;
        loadDataParams.dilationFilterW = 1;
        loadDataParams.dilationFilterH = 1;
        loadDataParams.filterSizeW = false;
        loadDataParams.filterSizeH = false;
        // LoadData(卷积数据搬运)v2 加载到 L0B 时,b loadDataParams.enTranspose 为 false,L0A 为 true
        loadDataParams.enTranspose = true;
        loadDataParams.fMatrixCtrl = false;
#if defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3510)
        // 使用 SetLoadDataRepeatWithStride 设置 LoadData(卷积数据搬运)v2 的 repeat 参数
        AscendC::LoadDataRepeatParamWithStride repeatParams;
        repeatParams.repeatTime = 1;                              // 重复次数(高度/宽度)
        repeatParams.repeatStride = 1;                            // 重复步进(高度/宽度)
        repeatParams.repeatMode = 0;                              // 0:高度方向;1:宽度方向
        repeatParams.dstStride = CeilDiv(nSize, fractalShape[0]); // K 方向步进
        AscendC::SetLoadDataRepeatWithStride(repeatParams);
        AscendC::LoadDataWithStride(b2Local, b1Local, loadDataParams);
#else
        AscendC::LoadData(b2Local, b1Local, loadDataParams);
#endif
    }

    __aicore__ inline void LoadSparse(
        AscendC::LocalTensor<int8_t> b2Local, AscendC::LocalTensor<int8_t> b1Local,
        AscendC::LocalTensor<uint8_t> idxB1Local)
    {
#if defined(__NPU_ARCH__) && (__NPU_ARCH__ == 2201)
        uint32_t nBlocks = nSize / BLOCK_CUBE;
        uint32_t kBlocks = kSize / C0_SIZE;
        // zn 转 zn 格式
        AscendC::LoadData2DParams loadDataParams;
        loadDataParams.repeatTimes = kBlocks * nBlocks / 2;
        loadDataParams.srcStride = 0;
        loadDataParams.ifTranspose = false;
        AscendC::LoadDataWithSparse(b2Local, b1Local, idxB1Local, loadDataParams);
#endif
    }

    __aicore__ inline void LoadBias(AscendC::LocalTensor<float>& bias2Local, AscendC::LocalTensor<float>& bias1Local)
    {
#if defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3510)
        AscendC::DataCopyParams c12c2Params = {1, static_cast<uint16_t>(biasSizeAlign / 32), 0, 0};
#else
        AscendC::DataCopyParams c12c2Params = {1, static_cast<uint16_t>(biasSizeAlign / 64), 0, 0};
#endif
        AscendC::DataCopy(bias2Local, bias1Local, c12c2Params);
    }

    __aicore__ inline void LoadFixBuffer(AscendC::LocalTensor<uint64_t> quantAlphaTensor)
    {
        AscendC::LocalTensor<uint64_t> fbTensor(AscendC::TPosition::C2PIPE2GM, 0, nSize);
        AscendC::DataCopy(fbTensor, quantAlphaTensor, nSize);
    }

private:
    uint32_t mSize, kSize, nSize, scenarioNum;
    uint32_t c0Size, fractalSize, fractalNum;
    uint32_t aSizeAlignL1, bSizeAlignL1, aSizeAlignL0, bSizeAlignL0, biasSizeAlign;
    uint32_t scaleASizeAlignL1, scaleBSizeAlignL1;
    uint16_t fractalShape[2] = {0, 0};
    uint64_t a1Addr, b1Addr, c1Addr;
    uint32_t scaleMAlignL1, scaleNAlignL1;
    uint32_t mAlignL1, kaAlignL1, nAlignL1, kbAlignL1;
    uint32_t mAlignL0, kaAlignL0, nAlignL0, kbAlignL0;
    uint32_t scaleK;
};

__global__ __cube__ void load_data_perf_custom(
    __gm__ uint8_t* a, __gm__ uint8_t* b, uint32_t scenario, uint32_t m, uint32_t k, uint32_t n)
{
    AscendC::InitSocState();
    KernelPerf<bfloat16_t> op(scenario, m, k, n);
    op.Process();
    AscendC::PipeBarrier<PIPE_ALL>();
}

int32_t main(int32_t argc, char* argv[])
{
    if (argc != 5) {
        std::printf("Usage: %s SCENARIO_NUM M K N\n", argv[0]);
        return 1;
    }

    uint32_t scenario = static_cast<uint32_t>(std::strtoul(argv[1], nullptr, 10));
    uint32_t m = static_cast<uint32_t>(std::strtoul(argv[2], nullptr, 10));
    uint32_t k = static_cast<uint32_t>(std::strtoul(argv[3], nullptr, 10));
    uint32_t n = static_cast<uint32_t>(std::strtoul(argv[4], nullptr, 10));
    if (scenario < 1 || scenario > 19) {
        std::printf("SCENARIO_NUM must be 1-19.\n");
        return 1;
    }
    if (m == 0 || k == 0 || n == 0) {
        std::printf("M, K and N must be positive integers.\n");
        return 1;
    }

    size_t aFileSize = static_cast<size_t>(m) * k * sizeof(bfloat16_t);
    size_t bFileSize = static_cast<size_t>(k) * n * sizeof(bfloat16_t);

    uint32_t numBlocks = 1;

    aclInit(nullptr);
    int32_t deviceId = 0;
    aclrtSetDevice(deviceId);
    aclrtStream stream = nullptr;
    aclrtCreateStream(&stream);

    uint8_t* aHost = nullptr;
    uint8_t* aDevice = nullptr;
    aclrtMallocHost((void**)(&aHost), aFileSize);
    aclrtMalloc((void**)&aDevice, aFileSize, ACL_MEM_MALLOC_HUGE_FIRST);
    for (size_t i = 0; i < aFileSize; ++i) {
        aHost[i] = static_cast<uint8_t>(i & 0xff);
    }
    aclrtMemcpy(aDevice, aFileSize, aHost, aFileSize, ACL_MEMCPY_HOST_TO_DEVICE);

    uint8_t* bHost = nullptr;
    uint8_t* bDevice = nullptr;
    aclrtMallocHost((void**)(&bHost), bFileSize);
    aclrtMalloc((void**)&bDevice, bFileSize, ACL_MEM_MALLOC_HUGE_FIRST);
    for (size_t i = 0; i < bFileSize; ++i) {
        bHost[i] = static_cast<uint8_t>((i + 1) & 0xff);
    }
    aclrtMemcpy(bDevice, bFileSize, bHost, bFileSize, ACL_MEMCPY_HOST_TO_DEVICE);

    load_data_perf_custom<<<numBlocks, 0, stream>>>(aDevice, bDevice, scenario, m, k, n);
    aclrtSynchronizeStream(stream);

    aclrtFree(aDevice);
    aclrtFreeHost(aHost);
    aclrtFree(bDevice);
    aclrtFreeHost(bHost);

    aclrtDestroyStream(stream);
    aclrtResetDevice(deviceId);
    aclFinalize();
    return 0;
}