/**
 * 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.
 *
 * Generated By CANNBot
 */

#ifndef EXPINT_KERNEL_H
#define EXPINT_KERNEL_H

#include "kernel_operator.h"
#include "kernel_tiling/kernel_tiling.h"
#include "arch35/expint_tiling_data.h"
#include "arch35/expint_tiling_key.h"
#include <type_traits>
#include <cmath>
#include <limits>

namespace NsExpint {

using namespace AscendC;

constexpr float EULER_GAMMA = 0.5772156649015329f;
constexpr float FP16_MAX_VALUE = 65504.0f;
constexpr float BF16_MAX_VALUE = 3.38953138927157e+38f;

static constexpr float A1[6] = {
    -5.350447357812542947283E0f, 2.185049168816613393830E2f,  -4.176572384826693777058E3f,
    5.541176756393557601232E4f,  -3.313381331178144034309E5f, 1.592627163384945414220E6f,
};
static constexpr float B1[7] = {
    1.0f,
    -5.250547959112862969197E1f,
    1.259616186786790571525E3f,
    -1.756549581973534652631E4f,
    1.493062117002725991967E5f,
    -7.294949239640527645655E5f,
    1.592627163384945429726E6f,
};

static constexpr float A2[8] = {
    1.981808503259689673238E-2f,  -1.271645625984917501326E0f,  -2.088160335681228318920E0f,
    2.755544509187936721172E0f,   -4.409507048701600257171E-1f, 4.665623805935891391017E-2f,
    -1.545042679673485262580E-3f, 7.059980605299617478514E-5f,
};
static constexpr float B2[8] = {
    1.0f,
    1.476498670914921440652E0f,
    5.629177174822436244827E-1f,
    1.699017897879307263248E-1f,
    2.291647179034212017463E-2f,
    4.450150439728752875043E-3f,
    1.727439612206521482874E-4f,
    3.953167195549672482304E-5f,
};

static constexpr float A3[8] = {
    -1.373215375871208729803E0f,  -7.084559133740838761406E-1f, 1.580806855547941010501E0f,
    -2.601500427425622944234E-1f, 2.994674694113713763365E-2f,  -1.038086040188744005513E-3f,
    4.371064420753005429514E-5f,  2.141783679522602903795E-6f,
};
static constexpr float B3[9] = {
    1.0f,
    8.585231423622028380768E-1f,
    4.483285822873995129957E-1f,
    7.687932158124475434091E-2f,
    2.449868241021887685904E-2f,
    8.832165941927796567926E-4f,
    4.590952299511353531215E-4f,
    -4.729848351866523044863E-6f,
    2.665195537390710170105E-6f,
};

static constexpr float A4[10] = {
    -2.106934601691916512584E0f,  1.732733869664688041885E0f,   -2.423619178935841904839E-1f,
    2.322724180937565842585E-2f,  2.372880440493179832059E-4f,  -8.343219561192552752335E-5f,
    1.363408795605250394881E-5f,  -3.655412321999253963714E-7f, 1.464941733975961318456E-8f,
    6.176407863710360207074E-10f,
};
static constexpr float B4[10] = {
    1.0f,
    -2.298062239901678075778E-1f,
    1.105077041474037862347E-1f,
    -1.566542966630792353556E-2f,
    2.761106850817352773874E-3f,
    -2.089148012284048449115E-4f,
    1.708528938807675304186E-5f,
    -4.459311796356686423199E-7f,
    1.394634930353847498145E-8f,
    6.150865933977338354138E-10f,
};

static constexpr float A5[8] = {
    -2.458119367674020323359E-1f, -1.483382253322077687183E-1f, 7.248291795735551591813E-2f,
    -1.348315687380940523823E-2f, 1.342775069788636972294E-3f,  -7.942465637159712264564E-5f,
    2.644179518984235952241E-6f,  -4.239473659313765177195E-8f,
};
static constexpr float B5[9] = {
    1.0f,
    -1.044225908443871106315E-1f,
    -2.676453128101402655055E-1f,
    9.695000254621984627876E-2f,
    -1.601745692712991078208E-2f,
    1.496414899205908021882E-3f,
    -8.462452563778485013756E-5f,
    2.728938403476726394024E-6f,
    -4.239462431819542051337E-8f,
};

static constexpr float A6[6] = {
    1.212561118105456670844E-1f,  -5.823133179043894485122E-1f, 2.348887314557016779211E-1f,
    -3.040034318113248237280E-2f, 1.510082146865190661777E-3f,  -2.523137095499571377122E-5f,
};
static constexpr float B6[6] = {
    1.0f,
    -1.002252150365854016662E0f,
    2.928709694872224144953E-1f,
    -3.337004338674007801307E-2f,
    1.560544881127388842819E-3f,
    -2.523137093603234562648E-5f,
};

static constexpr float A7[9] = {
    -7.657847078286127362028E-1f, 6.886192415566705051750E-1f,  -2.132598113545206124553E-1f,
    3.346107552384193813594E-2f,  -3.076541477344756050249E-3f, 1.747119316454907477380E-4f,
    -6.103711682274170530369E-6f, 1.218032765428652199087E-7f,  -1.086076102793290233007E-9f,
};
static constexpr float B7[10] = {
    1.0f,
    -1.888802868662308731041E0f,
    1.066691687211408896850E0f,
    -2.751915982306380647738E-1f,
    3.930852688233823569726E-2f,
    -3.414684558602365085394E-3f,
    1.866844370703555398195E-4f,
    -6.345146083130515357861E-6f,
    1.239754287483206878024E-7f,
    -1.086076102793126632978E-9f,
};

template <typename T>
class ExpintKernel {
    static constexpr int BUFFER_NUM = 1;
    static constexpr bool NEED_CAST = !std::is_same_v<T, float>;
    static constexpr int64_t CMP_ALIGN = 64;          // 256 bytes / 4 bytes per float
    static constexpr int64_t MASK_ELEM_PER_FLOAT = 8; // sizeof(float) / sizeof(uint8_t)
    static constexpr int64_t MASK_ALIGN = 32;         // mask buffer alignment in bytes
    static constexpr float EXP_CLAMP = 88.0f;         // exp(88) < FLT_MAX, exp(89) > FLT_MAX
    static constexpr float ONE = 1.0f;                // multiplicative identity / reciprocal numerator
    static constexpr float ZERO = 0.0f;               // additive identity / zero boundary
    static constexpr float INTERVAL_BOUND_2 = 2.0f;   // Interval 1/2 boundary
    static constexpr float INTERVAL_BOUND_4 = 4.0f;   // Interval 2/3 boundary
    static constexpr float INTERVAL_BOUND_8 = 8.0f;   // Interval 3/4 boundary
    static constexpr float INTERVAL_BOUND_16 = 16.0f; // Interval 4/5 boundary
    static constexpr float INTERVAL_BOUND_32 = 32.0f; // Interval 5/6 boundary
    static constexpr float INTERVAL_BOUND_64 = 64.0f; // Interval 6/7 boundary

public:
    __aicore__ inline ExpintKernel() {}

    __aicore__ inline void Init(GM_ADDR x, GM_ADDR y, const ExpintTilingData* tilingData);
    __aicore__ inline void Process();

private:
    __aicore__ inline void CopyIn(int64_t progress, int64_t currentNum);
    __aicore__ inline void Compute(LocalTensor<float> xInput, LocalTensor<float> yLocal, int64_t count);
    template <size_t N>
    __aicore__ inline void HornerEval(LocalTensor<float> acc, LocalTensor<float> scratch, LocalTensor<float> var,
                                      const float (&coeffs)[N], int64_t n);
    __aicore__ inline void CopyOut(int64_t progress, int64_t currentNum);
    __aicore__ inline void ProcessFp32(int64_t loopCount);
    __aicore__ inline void ProcessFp16Bf16(int64_t loopCount);

    TPipe pipe;
    TQue<TPosition::VECIN, BUFFER_NUM> inputQueue;
    TQue<TPosition::VECOUT, BUFFER_NUM> outputQueue;

    GlobalTensor<T> inputGM;
    GlobalTensor<T> outputGM;

    TBuf<TPosition::VECCALC> tmpBuf1;
    TBuf<TPosition::VECCALC> tmpBuf2;
    TBuf<TPosition::VECCALC> oldResultBuf;
    TBuf<TPosition::VECCALC> scratchBuf;
    TBuf<TPosition::VECCALC> maskBuf;
    TBuf<TPosition::VECCALC> initBuf;
    TBuf<TPosition::VECCALC> xClampedBuf;

    TBuf<TPosition::VECCALC> castInBuf;
    TBuf<TPosition::VECCALC> resultFp32Buf;

    int64_t blockLength_ = 0;
    int64_t ubLength_ = 0;
    int64_t alignedUbLength_ = 0;
};

template <typename T>
__aicore__ inline void ExpintKernel<T>::Init(GM_ADDR x, GM_ADDR y, const ExpintTilingData* tilingData)
{
    int64_t remainder = tilingData->totalNum - tilingData->blockFactor * GetBlockIdx();
    blockLength_ = (remainder > tilingData->blockFactor) ? tilingData->blockFactor : remainder;
    ubLength_ = tilingData->ubFactor;

    alignedUbLength_ = ((ubLength_ + CMP_ALIGN - 1) / CMP_ALIGN) * CMP_ALIGN;

    inputGM.SetGlobalBuffer((__gm__ T*)x + tilingData->blockFactor * GetBlockIdx(), blockLength_);
    outputGM.SetGlobalBuffer((__gm__ T*)y + tilingData->blockFactor * GetBlockIdx(), blockLength_);

    pipe.InitBuffer(inputQueue, BUFFER_NUM, alignedUbLength_ * sizeof(T));
    pipe.InitBuffer(outputQueue, BUFFER_NUM, alignedUbLength_ * sizeof(T));

    pipe.InitBuffer(tmpBuf1, alignedUbLength_ * sizeof(float));
    pipe.InitBuffer(tmpBuf2, alignedUbLength_ * sizeof(float));
    pipe.InitBuffer(oldResultBuf, alignedUbLength_ * sizeof(float));
    pipe.InitBuffer(scratchBuf, alignedUbLength_ * sizeof(float));
    pipe.InitBuffer(xClampedBuf, alignedUbLength_ * sizeof(float));
    pipe.InitBuffer(maskBuf, ((alignedUbLength_ / MASK_ELEM_PER_FLOAT + MASK_ALIGN - 1) / MASK_ALIGN) * MASK_ALIGN);

    if constexpr (std::is_same_v<T, float>) {
        pipe.InitBuffer(initBuf, alignedUbLength_ * sizeof(float));
    }

    if constexpr (NEED_CAST) {
        pipe.InitBuffer(castInBuf, alignedUbLength_ * sizeof(float));
        pipe.InitBuffer(resultFp32Buf, alignedUbLength_ * sizeof(float));
    }
}

template <typename T>
__aicore__ inline void ExpintKernel<T>::CopyIn(int64_t progress, int64_t currentNum)
{
    LocalTensor<T> xLocal = inputQueue.template AllocTensor<T>();
    DataCopyParams copyParams;
    copyParams.blockCount = 1;
    copyParams.blockLen = currentNum * sizeof(T);
    copyParams.srcStride = 0;
    copyParams.dstStride = 0;
    DataCopyPad(xLocal, inputGM[progress * ubLength_], copyParams, {false, 0, 0, 0});
    inputQueue.EnQue(xLocal);
}

template <typename T>
__aicore__ inline void ExpintKernel<T>::CopyOut(int64_t progress, int64_t currentNum)
{
    LocalTensor<T> yLocal = outputQueue.template DeQue<T>();
    DataCopyParams copyParams;
    copyParams.blockCount = 1;
    copyParams.blockLen = currentNum * sizeof(T);
    copyParams.srcStride = 0;
    copyParams.dstStride = 0;
    DataCopyPad(outputGM[progress * ubLength_], yLocal, copyParams);
    outputQueue.FreeTensor(yLocal);
}

template <typename T>
template <size_t N>
__aicore__ inline void ExpintKernel<T>::HornerEval(LocalTensor<float> acc, LocalTensor<float> scratch,
                                                   LocalTensor<float> var, const float (&coeffs)[N], int64_t n)
{
    Duplicate(acc, coeffs[0], n);
    for (size_t i = 1; i < N; i++) {
        Mul(scratch, acc, var, n);
        Adds(acc, scratch, coeffs[i], n);
    }
}

template <typename T>
__aicore__ inline void ExpintKernel<T>::Compute(LocalTensor<float> xInput, LocalTensor<float> yLocal, int64_t count)
{
    LocalTensor<float> tmp1 = tmpBuf1.Get<float>();
    LocalTensor<float> tmp2 = tmpBuf2.Get<float>();
    LocalTensor<float> scratch = scratchBuf.Get<float>();

    int64_t n = ((count + CMP_ALIGN - 1) / CMP_ALIGN) * CMP_ALIGN;

    // Use yLocal directly as result buffer (separate result TBuf has issues on Ascend950)
    LocalTensor<float>& result = yLocal;
    LocalTensor<float> oldResult = oldResultBuf.Get<float>();
    LocalTensor<uint8_t> mask = maskBuf.Get<uint8_t>();

    // Clamp x to EXP_CLAMP to prevent exp(x) overflow in float32
    LocalTensor<float> xFp32 = xClampedBuf.Get<float>();
    Duplicate(tmp1, EXP_CLAMP, n);
    Compare(mask, xInput, tmp1, CMPMODE::GT, n);
    Select(xFp32, mask, tmp1, xInput, SELMODE::VSEL_TENSOR_TENSOR_MODE, n);

    // Interval 7 (x >= 64): asymptotic expansion
    Duplicate(tmp1, ONE, n);
    Div(result, tmp1, xFp32, n);              // w = 1/x
    HornerEval(tmp1, tmp2, result, A7, n);    // P7(w)
    HornerEval(tmp2, scratch, result, B7, n); // Q7(w)
    Div(scratch, tmp1, tmp2, n);              // P7/Q7
    Exp(tmp2, xFp32, n);                      // e^x
    Div(tmp1, tmp2, xFp32, n);                // e^x/x
    Mul(tmp2, result, scratch, n);            // w * P7/Q7
    Adds(tmp2, tmp2, ONE, n);                 // w * P7/Q7 + 1
    Mul(result, tmp1, tmp2, n);               // e^x/x * (w*P7/Q7 + 1)

    // Intervals 6-2: rational approximation with Select merge
    // Interval 6 (32 <= x < 64)
    Adds(oldResult, result, ZERO, n);
    Duplicate(tmp1, ONE, n);
    Div(result, tmp1, xFp32, n); // w = 1/x
    HornerEval(tmp1, tmp2, result, A6, n);
    HornerEval(tmp2, scratch, result, B6, n);
    Div(scratch, tmp1, tmp2, n);
    Exp(tmp2, xFp32, n);
    Div(tmp1, tmp2, xFp32, n);
    Mul(tmp2, result, scratch, n);
    Adds(tmp2, tmp2, ONE, n);
    Mul(result, tmp1, tmp2, n);
    Duplicate(tmp1, INTERVAL_BOUND_64, n);
    Compare(mask, xFp32, tmp1, CMPMODE::LT, n);
    Select(result, mask, result, oldResult, SELMODE::VSEL_TENSOR_TENSOR_MODE, n);

    // Interval 5 (16 <= x < 32)
    Adds(oldResult, result, ZERO, n);
    Duplicate(tmp1, ONE, n);
    Div(result, tmp1, xFp32, n);
    HornerEval(tmp1, tmp2, result, A5, n);
    HornerEval(tmp2, scratch, result, B5, n);
    Div(scratch, tmp1, tmp2, n);
    Exp(tmp2, xFp32, n);
    Div(tmp1, tmp2, xFp32, n);
    Mul(tmp2, result, scratch, n);
    Adds(tmp2, tmp2, ONE, n);
    Mul(result, tmp1, tmp2, n);
    Duplicate(tmp1, INTERVAL_BOUND_32, n);
    Compare(mask, xFp32, tmp1, CMPMODE::LT, n);
    Select(result, mask, result, oldResult, SELMODE::VSEL_TENSOR_TENSOR_MODE, n);

    // Interval 4 (8 <= x < 16)
    Adds(oldResult, result, ZERO, n);
    Duplicate(tmp1, ONE, n);
    Div(result, tmp1, xFp32, n);
    HornerEval(tmp1, tmp2, result, A4, n);
    HornerEval(tmp2, scratch, result, B4, n);
    Div(scratch, tmp1, tmp2, n);
    Exp(tmp2, xFp32, n);
    Div(tmp1, tmp2, xFp32, n);
    Mul(tmp2, result, scratch, n);
    Adds(tmp2, tmp2, ONE, n);
    Mul(result, tmp1, tmp2, n);
    Duplicate(tmp1, INTERVAL_BOUND_16, n);
    Compare(mask, xFp32, tmp1, CMPMODE::LT, n);
    Select(result, mask, result, oldResult, SELMODE::VSEL_TENSOR_TENSOR_MODE, n);

    // Interval 3 (4 <= x < 8)
    Adds(oldResult, result, ZERO, n);
    Duplicate(tmp1, ONE, n);
    Div(result, tmp1, xFp32, n);
    HornerEval(tmp1, tmp2, result, A3, n);
    HornerEval(tmp2, scratch, result, B3, n);
    Div(scratch, tmp1, tmp2, n);
    Exp(tmp2, xFp32, n);
    Div(tmp1, tmp2, xFp32, n);
    Mul(tmp2, result, scratch, n);
    Adds(tmp2, tmp2, ONE, n);
    Mul(result, tmp1, tmp2, n);
    Duplicate(tmp1, INTERVAL_BOUND_8, n);
    Compare(mask, xFp32, tmp1, CMPMODE::LT, n);
    Select(result, mask, result, oldResult, SELMODE::VSEL_TENSOR_TENSOR_MODE, n);

    // Interval 2 (2 <= x < 4)
    Adds(oldResult, result, ZERO, n);
    Duplicate(tmp1, ONE, n);
    Div(result, tmp1, xFp32, n);
    HornerEval(tmp1, tmp2, result, A2, n);
    HornerEval(tmp2, scratch, result, B2, n);
    Div(scratch, tmp1, tmp2, n);
    Exp(tmp2, xFp32, n);
    Div(tmp1, tmp2, xFp32, n);
    Mul(tmp2, result, scratch, n);
    Adds(tmp2, tmp2, ONE, n);
    Mul(result, tmp1, tmp2, n);
    Duplicate(tmp1, INTERVAL_BOUND_4, n);
    Compare(mask, xFp32, tmp1, CMPMODE::LT, n);
    Select(result, mask, result, oldResult, SELMODE::VSEL_TENSOR_TENSOR_MODE, n);

    // Interval 1 (0 < x < 2): Ei(x) = gamma + ln(x) + x * P1(x)/Q1(x)
    HornerEval(tmp1, tmp2, xFp32, A1, n);
    Adds(scratch, result, ZERO, n);
    HornerEval(tmp2, oldResult, xFp32, B1, n);
    Div(oldResult, tmp1, tmp2, n);
    Mul(tmp1, oldResult, xFp32, n);
    Adds(tmp1, tmp1, EULER_GAMMA, n);
    Ln(tmp2, xFp32, n);
    Add(oldResult, tmp1, tmp2, n);
    Duplicate(tmp1, INTERVAL_BOUND_2, n);
    Compare(mask, xFp32, tmp1, CMPMODE::LT, n);
    Select(result, mask, oldResult, scratch, SELMODE::VSEL_TENSOR_TENSOR_MODE, n);

    // ApplyBoundaries
    Duplicate(tmp1, ZERO, n);
    Compare(mask, xFp32, tmp1, CMPMODE::LT, n);
    Duplicate(tmp1, std::numeric_limits<float>::quiet_NaN(), n);
    Select(result, mask, tmp1, result, SELMODE::VSEL_TENSOR_TENSOR_MODE, n);

    Duplicate(tmp1, std::numeric_limits<float>::infinity(), n);
    Compare(mask, xFp32, tmp1, CMPMODE::EQ, n);
    Select(result, mask, tmp1, result, SELMODE::VSEL_TENSOR_TENSOR_MODE, n);

    // x == 0 -> -inf (exact comparison)
    Duplicate(tmp1, ZERO, n);
    Compare(mask, xFp32, tmp1, CMPMODE::EQ, n);
    Duplicate(tmp1, -std::numeric_limits<float>::infinity(), n);
    Select(result, mask, tmp1, result, SELMODE::VSEL_TENSOR_TENSOR_MODE, n);
}

template <typename T>
__aicore__ inline void ExpintKernel<T>::ProcessFp32(int64_t loopCount)
{
    for (int64_t i = 0; i < loopCount; i++) {
        int64_t currentNum = (i == (loopCount - 1)) ? (blockLength_ - ubLength_ * i) : ubLength_;
        int64_t alignedCount = ((currentNum + CMP_ALIGN - 1) / CMP_ALIGN) * CMP_ALIGN;

        LocalTensor<float> xFp32 = initBuf.Get<float>();
        Duplicate(xFp32, ONE, alignedCount);
        SetFlag<HardEvent::V_MTE2>(0);
        WaitFlag<HardEvent::V_MTE2>(0);
        DataCopyParams copyParams;
        copyParams.blockCount = 1;
        copyParams.blockLen = currentNum * sizeof(float);
        copyParams.srcStride = 0;
        copyParams.dstStride = 0;
        DataCopyPad(xFp32, inputGM[i * ubLength_], copyParams, {false, 0, 0, 0});
        SetFlag<HardEvent::MTE2_V>(0);
        WaitFlag<HardEvent::MTE2_V>(0);

        LocalTensor<float> yLocal = outputQueue.template AllocTensor<float>();
        Compute(xFp32, yLocal, currentNum);
        outputQueue.template EnQue<float>(yLocal);

        CopyOut(i, currentNum);
    }
}

template <typename T>
__aicore__ inline void ExpintKernel<T>::ProcessFp16Bf16(int64_t loopCount)
{
    for (int64_t i = 0; i < loopCount; i++) {
        int64_t currentNum = (i == (loopCount - 1)) ? (blockLength_ - ubLength_ * i) : ubLength_;
        CopyIn(i, currentNum);

        LocalTensor<T> xInput = inputQueue.template DeQue<T>();

        LocalTensor<float> xFp32 = castInBuf.Get<float>();
        int64_t alignedCount = ((currentNum + CMP_ALIGN - 1) / CMP_ALIGN) * CMP_ALIGN;
        Duplicate(xFp32, ONE, alignedCount);
        Cast<float, T>(xFp32, xInput, RoundMode::CAST_NONE, currentNum);

        LocalTensor<float> resultFp32 = resultFp32Buf.Get<float>();
        Compute(xFp32, resultFp32, currentNum);

        // Clamp float32 result to target dtype range to prevent Cast overflow to inf
        LocalTensor<float> clampVal = castInBuf.Get<float>();
        LocalTensor<uint8_t> clampMask = maskBuf.Get<uint8_t>();
        float maxVal;
        if constexpr (std::is_same_v<T, half>) {
            maxVal = FP16_MAX_VALUE;
        } else {
            maxVal = BF16_MAX_VALUE;
        }
        Duplicate(clampVal, maxVal, alignedCount);
        Compare(clampMask, resultFp32, clampVal, CMPMODE::GT, alignedCount);
        LocalTensor<float> clampTmp = oldResultBuf.Get<float>();
        Select(clampTmp, clampMask, clampVal, resultFp32, SELMODE::VSEL_TENSOR_TENSOR_MODE, alignedCount);
        Adds(resultFp32, clampTmp, ZERO, alignedCount);

        LocalTensor<T> yOutput = outputQueue.template AllocTensor<T>();
        Cast<T, float>(yOutput, resultFp32, RoundMode::CAST_ROUND, currentNum);
        outputQueue.template EnQue<T>(yOutput);

        inputQueue.FreeTensor(xInput);
        CopyOut(i, currentNum);
    }
}

template <typename T>
__aicore__ inline void ExpintKernel<T>::Process()
{
    if (ubLength_ == 0 || blockLength_ == 0) {
        return;
    }

    int64_t loopCount = (blockLength_ + ubLength_ - 1) / ubLength_;

    if constexpr (std::is_same_v<T, float>) {
        ProcessFp32(loopCount);
    } else {
        ProcessFp16Bf16(loopCount);
    }
}

} // namespace NsExpint

#endif // EXPINT_KERNEL_H