From 8d587d96f03737dac65c14525c82a992ad21deb5 Mon Sep 17 00:00:00 2001 From: FengfengST <1322434659@qq.com> Date: Thu, 16 Jul 2026 20:48:34 +0800 Subject: [PATCH] feat: add GammaAddRmsNorm custom operator --- test/python_test/RegisterOps.cpp | 34 + test/python_test/custom_ops.py | 6 + test/python_test/test_gamma_add_rms_norm.py | 88 ++ xllm_ops/build_aclnn.sh | 3 + xllm_ops/gamma_add_rms_norm/CMakeLists.txt | 8 + .../norm_common/reduce_common_regbase.h | 1047 ++++++++++++ .../gamma_add_rms_norm/op_host/CMakeLists.txt | 25 + .../op_host/gamma_add_rms_norm_def.cpp | 74 + .../op_host/gamma_add_rms_norm_error_log.h | 37 + .../op_host/gamma_add_rms_norm_infershape.cpp | 85 + .../op_host/gamma_add_rms_norm_tiling.cpp | 563 +++++++ .../op_host/gamma_add_rms_norm_tiling.h | 89 ++ .../gamma_add_rms_norm_tiling_arch35.cpp | 256 +++ .../op_api/aclnn_gamma_add_rms_norm.cpp | 252 +++ .../op_host/op_api/aclnn_gamma_add_rms_norm.h | 75 + .../op_host/op_api/gamma_add_rms_norm.cpp | 70 + .../op_host/op_api/gamma_add_rms_norm.h | 32 + .../arch35/gamma_add_rms_norm_regbase.h | 361 +++++ .../gamma_add_rms_norm_regbase_common.h | 768 +++++++++ .../gamma_add_rms_norm_regbase_split_d.h | 336 ++++ .../op_kernel/gamma_add_rms_norm.cpp | 129 ++ .../op_kernel/gamma_add_rms_norm.h | 366 +++++ .../op_kernel/gamma_add_rms_norm_apt.cpp | 39 + .../op_kernel/gamma_add_rms_norm_base.h | 321 ++++ .../op_kernel/gamma_add_rms_norm_merge_n.h | 429 +++++ .../op_kernel/gamma_add_rms_norm_multi_n.h | 339 ++++ .../op_kernel/gamma_add_rms_norm_single_n.h | 370 +++++ .../op_kernel/gamma_add_rms_norm_split_d.h | 425 +++++ .../op_kernel/inc/platform.h | 73 + .../op_kernel/reduce_common.h | 180 +++ .../rms_norm/arch35/rms_norm_regbase_common.h | 1401 +++++++++++++++++ 31 files changed, 8281 insertions(+) create mode 100644 test/python_test/test_gamma_add_rms_norm.py create mode 100644 xllm_ops/gamma_add_rms_norm/CMakeLists.txt create mode 100644 xllm_ops/gamma_add_rms_norm/norm_common/reduce_common_regbase.h create mode 100644 xllm_ops/gamma_add_rms_norm/op_host/CMakeLists.txt create mode 100644 xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_def.cpp create mode 100644 xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_error_log.h create mode 100644 xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_infershape.cpp create mode 100644 xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_tiling.cpp create mode 100644 xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_tiling.h create mode 100644 xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_tiling_arch35.cpp create mode 100644 xllm_ops/gamma_add_rms_norm/op_host/op_api/aclnn_gamma_add_rms_norm.cpp create mode 100644 xllm_ops/gamma_add_rms_norm/op_host/op_api/aclnn_gamma_add_rms_norm.h create mode 100644 xllm_ops/gamma_add_rms_norm/op_host/op_api/gamma_add_rms_norm.cpp create mode 100644 xllm_ops/gamma_add_rms_norm/op_host/op_api/gamma_add_rms_norm.h create mode 100644 xllm_ops/gamma_add_rms_norm/op_kernel/arch35/gamma_add_rms_norm_regbase.h create mode 100644 xllm_ops/gamma_add_rms_norm/op_kernel/arch35/gamma_add_rms_norm_regbase_common.h create mode 100644 xllm_ops/gamma_add_rms_norm/op_kernel/arch35/gamma_add_rms_norm_regbase_split_d.h create mode 100644 xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm.cpp create mode 100644 xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm.h create mode 100644 xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_apt.cpp create mode 100644 xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_base.h create mode 100644 xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_merge_n.h create mode 100644 xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_multi_n.h create mode 100644 xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_single_n.h create mode 100644 xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_split_d.h create mode 100644 xllm_ops/gamma_add_rms_norm/op_kernel/inc/platform.h create mode 100644 xllm_ops/gamma_add_rms_norm/op_kernel/reduce_common.h create mode 100644 xllm_ops/gamma_add_rms_norm/rms_norm/arch35/rms_norm_regbase_common.h diff --git a/test/python_test/RegisterOps.cpp b/test/python_test/RegisterOps.cpp index bb5516d..2c8f54d 100644 --- a/test/python_test/RegisterOps.cpp +++ b/test/python_test/RegisterOps.cpp @@ -601,6 +601,37 @@ std::tuple add_rms_norm_bias_impl_npu( return std::make_tuple(y, rstd, x); } +std::tuple +gamma_add_rms_norm_impl_npu(const at::Tensor& x1, + const at::Tensor& x2, + const at::Tensor& gamma, + double epsilon, + bool add_gamma_offset) { + auto sizes = x1.sizes().vec(); + const int64_t dim_x = static_cast(sizes.size()); + const int64_t dim_gamma = gamma.dim(); + std::vector rstd_shape(sizes.begin(), sizes.end()); + for (int64_t i = dim_x - dim_gamma; i < dim_x && i >= 0; ++i) { + rstd_shape[i] = 1; + } + + at::Tensor y = at::empty(sizes, x1.options()); + at::Tensor rstd = + at::empty(rstd_shape, x1.options().dtype(at::ScalarType::Float)); + at::Tensor x = at::empty(sizes, x1.options()); + + EXEC_NPU_CMD(aclnnGammaAddRmsNorm, + x1, + x2, + gamma, + epsilon, + add_gamma_offset, + y, + rstd, + x); + return std::make_tuple(y, rstd, x); +} + std::tuple moe_init_routing_custom_impl_npu(const at::Tensor& x, const at::Tensor& expert_idx, @@ -1171,6 +1202,9 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("scatter_nd_update_v2", &scatter_nd_update_v2_impl_npu, "scatter_nd_update_v2"); m.def("hc_post", &hc_post_impl_npu, "hc_post"); m.def("add_rms_norm_bias", &add_rms_norm_bias_impl_npu, "add_rms_norm_bias"); + m.def("gamma_add_rms_norm", + &gamma_add_rms_norm_impl_npu, + "gamma_add_rms_norm"); m.def("moe_init_routing_custom", &moe_init_routing_custom_impl_npu, "moe_init_routing_custom"); m.def("moe_init_routing_v3", &moe_init_routing_v3_impl_npu, "moe_init_routing_v3"); m.def("hc_pre_sinkhorn", &hc_pre_sinkhorn_impl_npu, "hc_pre_sinkhorn"); diff --git a/test/python_test/custom_ops.py b/test/python_test/custom_ops.py index ae8c747..160769a 100644 --- a/test/python_test/custom_ops.py +++ b/test/python_test/custom_ops.py @@ -181,6 +181,12 @@ def add_rms_norm_bias_npu(x1, x2, gamma, beta=None, eps=1e-6): return custom_ops_lib.add_rms_norm_bias(x1, x2, gamma, beta, eps) +def gamma_add_rms_norm_npu(x1, x2, gamma, eps=1e-6, + add_gamma_offset=False): + return custom_ops_lib.gamma_add_rms_norm( + x1, x2, gamma, eps, add_gamma_offset) + + # hc_pre_sinkhorn (per-token: pre/post gating + sinkhorn-normalized comb_frag) # mixes last dim = 2*hc_mult + hc_mult^2 ([pre | post | comb]). # returns (y[bs,d] bf16, post[bs,hc_mult] fp32, comb_frag[bs,hc_mult,hc_mult] fp32) diff --git a/test/python_test/test_gamma_add_rms_norm.py b/test/python_test/test_gamma_add_rms_norm.py new file mode 100644 index 0000000..a204622 --- /dev/null +++ b/test/python_test/test_gamma_add_rms_norm.py @@ -0,0 +1,88 @@ +#!/usr/bin/env python3 +# Copyright 2026 The xLLM Authors. 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. +# ============================================================================== + +import pytest +import torch + + +torch_npu = pytest.importorskip("torch_npu") +custom_ops = pytest.importorskip("custom_ops") + + +def gamma_add_rms_norm_golden(x1, x2, gamma, eps, add_gamma_offset): + dtype = x1.dtype + if dtype == torch.float16: + x = x1 + x2 + else: + x = (x1.float() + x2.float()).to(dtype) + + adjusted_gamma = gamma + if add_gamma_offset: + adjusted_gamma = (gamma + 1).to(dtype) + + x_float = x.float() + variance = (x_float * x_float).mean( + dim=tuple(range(x.dim() - gamma.dim(), x.dim())), keepdim=True) + rstd = torch.rsqrt(variance + eps) + normalized = (x_float * rstd).to(dtype) + y = (normalized * adjusted_gamma).to(dtype) + return y, rstd.float(), x + + +# (x shape, gamma shape) covers decode, prefill, and multi-dimensional gamma. +CASES = [ + ((1, 1024), (1024,)), + ((4, 2048), (2048,)), + ((16, 5120), (5120,)), + ((2, 3, 256), (3, 256)), +] + + +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16, torch.float32]) +@pytest.mark.parametrize("add_gamma_offset", [False, True]) +@pytest.mark.parametrize("x_shape,gamma_shape", CASES) +def test_gamma_add_rms_norm(x_shape, gamma_shape, add_gamma_offset, dtype): + torch.manual_seed(20260716) + eps = 1e-6 + x1 = (torch.rand(x_shape, dtype=torch.float32) - 0.5).to(dtype) + x2 = (torch.rand(x_shape, dtype=torch.float32) - 0.5).to(dtype) + gamma = (torch.rand(gamma_shape, dtype=torch.float32) - 0.5).to(dtype) + + y_ref, rstd_ref, x_ref = gamma_add_rms_norm_golden( + x1, x2, gamma, eps, add_gamma_offset) + y_npu, rstd_npu, x_npu = custom_ops.gamma_add_rms_norm_npu( + x1.npu(), + x2.npu(), + gamma.npu(), + eps, + add_gamma_offset, + ) + torch.npu.synchronize() + + if dtype == torch.float16: + atol = rtol = 1e-3 + elif dtype == torch.bfloat16: + atol = rtol = 5e-3 + else: + atol = rtol = 1e-5 + + assert tuple(rstd_npu.shape) == tuple(rstd_ref.shape) + torch.testing.assert_close( + x_npu.cpu().float(), x_ref.float(), atol=atol, rtol=rtol) + torch.testing.assert_close( + rstd_npu.cpu().float(), rstd_ref, atol=1e-5, rtol=1e-5) + torch.testing.assert_close( + y_npu.cpu().float(), y_ref.float(), atol=atol, rtol=rtol) diff --git a/xllm_ops/build_aclnn.sh b/xllm_ops/build_aclnn.sh index 6bddae7..56769c1 100644 --- a/xllm_ops/build_aclnn.sh +++ b/xllm_ops/build_aclnn.sh @@ -110,6 +110,7 @@ elif [[ "$SOC_VERSION" =~ ^(ascend)?910b ]]; then "moe_init_routing_custom" "moe_gating_top_k_hash" "add_rms_norm_bias" + "gamma_add_rms_norm" "lightning_indexer_quant" "compressor" "quant_lightning_indexer" ## 已在 CANN 中内置,见 opp/built-in/op_impl/ai_core/tbe/impl/ops_transformer/ascendc/quant_lightning_indexer @@ -218,6 +219,7 @@ elif [[ "$SOC_VERSION" =~ ^ascend910_93 ]]; then "moe_init_routing_custom" "moe_gating_top_k_hash" "add_rms_norm_bias" + "gamma_add_rms_norm" "lightning_indexer_quant" "lightning_indexer_quant_metadata" "compressor" @@ -290,6 +292,7 @@ elif [[ "$SOC_VERSION" =~ ^ascend950 ]]; then "moe_init_routing_custom" "moe_gating_top_k_hash" "add_rms_norm_bias" + "gamma_add_rms_norm" "lightning_indexer_quant" "lightning_indexer_quant_metadata" "compressor" diff --git a/xllm_ops/gamma_add_rms_norm/CMakeLists.txt b/xllm_ops/gamma_add_rms_norm/CMakeLists.txt new file mode 100644 index 0000000..f751482 --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/CMakeLists.txt @@ -0,0 +1,8 @@ +# Copyright 2026 The xLLM Authors. All Rights Reserved. + +file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) +foreach(SUB_DIR ${CURRENT_DIRS}) + if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") + add_subdirectory(${SUB_DIR}) + endif() +endforeach() diff --git a/xllm_ops/gamma_add_rms_norm/norm_common/reduce_common_regbase.h b/xllm_ops/gamma_add_rms_norm/norm_common/reduce_common_regbase.h new file mode 100644 index 0000000..35ea76c --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/norm_common/reduce_common_regbase.h @@ -0,0 +1,1047 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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 reduce_common_regbase.h + * \brief reduce common regbase file + */ +#ifndef REDUCE_COMMON_REGBASE_H_RMS_NORM +#define REDUCE_COMMON_REGBASE_H_RMS_NORM +#include "kernel_operator.h" + +namespace NormCommon { +using namespace AscendC; +using AscendC::MicroAPI::CreateMask; +using AscendC::MicroAPI::LoadDist; +using AscendC::MicroAPI::LocalMemBar; +using AscendC::MicroAPI::MaskPattern; +using AscendC::MicroAPI::MaskReg; +using AscendC::MicroAPI::RegTensor; +using AscendC::MicroAPI::MemType; +using AscendC::MicroAPI::UpdateMask; +using AscendC::MicroAPI::StoreDist; + +namespace NormCommonRegbase { +__aicore__ inline constexpr uint32_t GetVRegSize() +{ +#if __CCE_AICORE__ == 310 + return AscendC::VECTOR_REG_WIDTH; +#else + return 256U; +#endif +} + +template +__aicore__ inline T CeilDiv(T a, T b) +{ + using type = typename std::conditional< + sizeof(T) == sizeof(uint8_t) || sizeof(T) == sizeof(uint16_t), uint32_t, uint64_t>::type; + type res = (static_cast(a) + static_cast(b) - 1) / static_cast(b); + return static_cast(res); +} + +template +__aicore__ inline T CeilAlign(T a, T b) +{ + using type = typename std::conditional< + sizeof(T) == sizeof(uint8_t) || sizeof(T) == sizeof(uint16_t), uint32_t, uint64_t>::type; + type res = (static_cast(a) + static_cast(b) - 1) / static_cast(b) * static_cast(b); + return static_cast(res); +} + +template +__aicore__ inline T Aligned(T value, T alignment) +{ + if (alignment == 0) { + return value; + } + return (value + alignment - 1) / alignment * alignment; +} + +} // namespace + +constexpr int32_t VL_SIZE = NormCommonRegbase::GetVRegSize(); +constexpr int32_t V_LENGTH = (VL_SIZE / static_cast(sizeof(float))); +constexpr uint32_t ONCE_VECTOR_SIZE = 256; +constexpr uint16_t DICHOTOMY_ADD_COEFF = 2; + +constexpr AscendC::MicroAPI::CastTrait castTraitB162B32 = { + AscendC::MicroAPI::RegLayout::ZERO, + AscendC::MicroAPI::SatMode::UNKNOWN, + AscendC::MicroAPI::MaskMergeMode::ZEROING, + AscendC::RoundMode::UNKNOWN, +}; + +constexpr AscendC::MicroAPI::CastTrait castTraitB322B16 = { + AscendC::MicroAPI::RegLayout::ZERO, + AscendC::MicroAPI::SatMode::NO_SAT, + AscendC::MicroAPI::MaskMergeMode::ZEROING, + AscendC::RoundMode::CAST_RINT, +}; + +__aicore__ inline void DichotomyAdd( + RegTensor& dstReg, __local_mem__ float* src, uint16_t outerLoop, uint16_t innerLoop, uint32_t lastNum) +{ + RegTensor tmpReg1; + RegTensor tmpReg2; + RegTensor tmpReg3; + LocalMemBar(); + MaskReg pregMain = CreateMask(); + for (uint16_t k = 0; k < outerLoop; k++) { + innerLoop = innerLoop / DICHOTOMY_ADD_COEFF; + for (uint16_t i = 0; i < innerLoop; i++) { + DataCopy(tmpReg1, src + i * V_LENGTH); + DataCopy(tmpReg2, src + (i + innerLoop) * V_LENGTH); + Add(tmpReg3, tmpReg1, tmpReg2, pregMain); + DataCopy(src + i * V_LENGTH, tmpReg3, pregMain); + } + LocalMemBar(); + } + uint32_t sreg0 = lastNum; + MaskReg pregLoop = UpdateMask(sreg0); + DataCopy(tmpReg3, src); + ReduceSum(dstReg, tmpReg3, pregLoop); +} + +template +__aicore__ inline void LoadTwoCloseRegVF( + RegTensor& dstA, RegTensor& dstB, __local_mem__ U* srcAddr, uint16_t offset) +{ + if constexpr (IsSameType::value) { + DataCopy(dstA, srcAddr + offset); + DataCopy(dstB, srcAddr + offset + V_LENGTH); + } else { + DataCopy(dstA, srcAddr + offset); + DataCopy(dstB, srcAddr + offset + V_LENGTH); + } +} + +template +__aicore__ inline void CastAddVF( + RegTensor& dstReg, RegTensor& src1Reg, RegTensor& src2Reg, MaskReg& pregLoop) +{ + if constexpr (IsSameType::value) { + Add(dstReg, src1Reg, src2Reg, pregLoop); + } else { + RegTensor src1RegFp32, src2RegFp32; + Cast(src1RegFp32, src1Reg, pregLoop); + Cast(src2RegFp32, src2Reg, pregLoop); + Add(dstReg, src1RegFp32, src2RegFp32, pregLoop); + } +} + +/** + * @brief Load and cast to fp32 reg. + * @param offset idx of VF loop. + */ +template +__aicore__ inline void LoadCastRegVF( + RegTensor& dstTensor, __local_mem__ T* srcAddr, uint16_t offset, MaskReg& pregLoop) +{ + if constexpr (IsSameType::value) { + DataCopy(dstTensor, srcAddr + offset * V_LENGTH); + } else { + RegTensor loadTmp; + DataCopy(loadTmp, srcAddr + offset * V_LENGTH); + Cast(dstTensor, loadTmp, pregLoop); + } +} + +template +__aicore__ inline void CastStoreTwoCloseRegVF( + __local_mem__ T* dstAddr, RegTensor& srcA, RegTensor& srcB, uint16_t offset, MaskReg& pregLoop) +{ + if constexpr (IsSameType::value) { + DataCopy(dstAddr + offset, srcA, pregLoop); + DataCopy(dstAddr + offset + V_LENGTH, srcB, pregLoop); + } else { + RegTensor srcATmp, srcBTmp; + Cast(srcATmp, srcA, pregLoop); + Cast(srcBTmp, srcB, pregLoop); + DataCopy(dstAddr + offset, srcATmp, pregLoop); + DataCopy(dstAddr + offset + V_LENGTH, srcBTmp, pregLoop); + } +} + +/** + * @brief Use VF to Compute reduceSum. + * dstLocal = reduceSum((x1+x2)^2) + * If HAS_XOUT is true, return xOut = (x1.to(float) + x2.to(float)).to(dtype). + * If HAS_XOUT_FP32 is true, return xOutFp32 = x1.to(float) + x2.to(float). + * If IS_RSTD is true, dstLocal = 1.0 / sqrt(avgFactor * reduceSum((x1+x2)^2) + epsilon) + * Use float32 VL_LENGTH + */ +template +__aicore__ inline void ReduceSumRstd(LocalTensor& dstLocal, LocalTensor& xOutLocal, + LocalTensor& xOutFp32Local, LocalTensor& x1Local, LocalTensor& x2Local, LocalTensor& workLocal, + uint32_t dstOffset, uint32_t count, uint32_t powerSplit, float avgFactor = 1.0f, float epsilon = 0.0f) +{ + uint32_t remainTile = count - powerSplit; + uint32_t remainSreg = remainTile; + uint16_t remainRepeats = remainTile / (2 * V_LENGTH); + + uint32_t masterTile = powerSplit - remainTile; + uint32_t masterSreg = masterTile; + uint16_t masterRepeats = masterTile / (2 * V_LENGTH); + + uint32_t mergeTile = powerSplit / (2 * V_LENGTH); + uint32_t mergeSreg = mergeTile; + uint16_t mergeRepeats = mergeTile / (2 * V_LENGTH); + + uint32_t meanTile = mergeRepeats == 0 ? mergeTile : mergeRepeats; + uint32_t meanSreg = meanTile; + + __local_mem__ U* x1MainAddr = (__ubuf__ U*)x1Local.GetPhyAddr(); + __local_mem__ U* x1TailAddr = (__ubuf__ U*)x1Local.GetPhyAddr() + int64_t(powerSplit); + __local_mem__ U* x1MasterAddr = (__ubuf__ U*)x1Local.GetPhyAddr() + int64_t(remainTile); + __local_mem__ U* x2MainAddr = (__ubuf__ U*)x2Local.GetPhyAddr(); + __local_mem__ U* x2TailAddr = (__ubuf__ U*)x2Local.GetPhyAddr() + int64_t(powerSplit); + __local_mem__ U* x2MasterAddr = (__ubuf__ U*)x2Local.GetPhyAddr() + int64_t(remainTile); + __local_mem__ U* xOutMainAddr; + __local_mem__ U* xOutTailAddr; + __local_mem__ U* xOutMasterAddr; + if constexpr (HAS_XOUT) { + xOutMainAddr = (__ubuf__ U*)xOutLocal.GetPhyAddr(); + xOutTailAddr = (__ubuf__ U*)xOutLocal.GetPhyAddr() + int64_t(powerSplit); + xOutMasterAddr = (__ubuf__ U*)xOutLocal.GetPhyAddr() + int64_t(remainTile); + } + __local_mem__ float* xOutFp32MainAddr; + __local_mem__ float* xOutFp32TailAddr; + __local_mem__ float* xOutFp32MasterAddr; + if constexpr (HAS_XOUT_FP32) { + xOutFp32MainAddr = (__ubuf__ float*)xOutFp32Local.GetPhyAddr(); + xOutFp32TailAddr = (__ubuf__ float*)xOutFp32Local.GetPhyAddr() + int64_t(powerSplit); + xOutFp32MasterAddr = (__ubuf__ float*)xOutFp32Local.GetPhyAddr() + int64_t(remainTile); + } + __local_mem__ float* workAddr = (__ubuf__ float*)workLocal.GetPhyAddr(); + __local_mem__ float* dstAddr = (__ubuf__ float*)dstLocal.GetPhyAddr(); + + __VEC_SCOPE__ + { + RegTensor mainA, mainB, tailA, tailB, vSum, vDupReg; + RegTensor x1MainA, x1MainB, x1TailA, x1TailB; + RegTensor x2MainA, x2MainB, x2TailA, x2TailB; + MaskReg pregMerge = CreateMask(); + MaskReg pregLoop; + + for (uint16_t i = 0; i < (uint16_t)remainRepeats; ++i) { + pregLoop = UpdateMask(remainSreg); + uint16_t offset = i * 2 * V_LENGTH; + // 1. Copy in reg + LoadTwoCloseRegVF(x1MainA, x1MainB, x1MainAddr, offset); + LoadTwoCloseRegVF(x1TailA, x1TailB, x1TailAddr, offset); + LoadTwoCloseRegVF(x2MainA, x2MainB, x2MainAddr, offset); + LoadTwoCloseRegVF(x2TailA, x2TailB, x2TailAddr, offset); + // 2. Cast add + CastAddVF(mainA, x1MainA, x2MainA, pregLoop); + CastAddVF(tailA, x1TailA, x2TailA, pregLoop); + CastAddVF(mainB, x1MainB, x2MainB, pregLoop); + CastAddVF(tailB, x1TailB, x2TailB, pregLoop); + if constexpr (HAS_XOUT) { + CastStoreTwoCloseRegVF(xOutMainAddr, mainA, mainB, offset, pregLoop); + CastStoreTwoCloseRegVF(xOutTailAddr, tailA, tailB, offset, pregLoop); + } + if constexpr (HAS_XOUT_FP32) { + CastStoreTwoCloseRegVF(xOutFp32MainAddr, mainA, mainB, offset, pregLoop); + CastStoreTwoCloseRegVF(xOutFp32TailAddr, tailA, tailB, offset, pregLoop); + } + // 3. Cal x^2 + Mul(mainA, mainA, mainA, pregLoop); + Mul(tailA, tailA, tailA, pregLoop); + Mul(mainB, mainB, mainB, pregLoop); + Mul(tailB, tailB, tailB, pregLoop); + Add(mainA, mainA, tailA, pregLoop); + Add(mainB, mainB, tailB, pregLoop); + Add(mainA, mainA, mainB, pregLoop); + ReduceSum(vSum, mainA, pregLoop); + DataCopy(workAddr + i, vSum, pregMerge); + } + for (uint16_t i = 0; i < (uint16_t)masterRepeats; ++i) { + uint16_t offset = i * 2 * V_LENGTH; + pregLoop = UpdateMask(masterSreg); + // 1. Copy in reg + LoadTwoCloseRegVF(x1MainA, x1MainB, x1MasterAddr, offset); + LoadTwoCloseRegVF(x2MainA, x2MainB, x2MasterAddr, offset); + // 2. Cast add + CastAddVF(mainA, x1MainA, x2MainA, pregLoop); + CastAddVF(mainB, x1MainB, x2MainB, pregLoop); + if constexpr (HAS_XOUT) { + CastStoreTwoCloseRegVF(xOutMasterAddr, mainA, mainB, offset, pregLoop); + } + if constexpr (HAS_XOUT_FP32) { + CastStoreTwoCloseRegVF(xOutFp32MasterAddr, mainA, mainB, offset, pregLoop); + } + // 3. Cal x^2 + Mul(mainA, mainA, mainA, pregLoop); + Mul(mainB, mainB, mainB, pregLoop); + Add(mainA, mainA, mainB, pregLoop); + ReduceSum(vSum, mainA, pregLoop); + DataCopy(workAddr + remainRepeats + i, vSum, pregMerge); + } + LocalMemBar(); + for (uint16_t i = 0; i < (uint16_t)mergeRepeats; ++i) { + pregLoop = UpdateMask(mergeSreg); + uint16_t offset = i * 2 * V_LENGTH; + LoadTwoCloseRegVF(mainA, mainB, workAddr, offset); + Add(mainA, mainA, mainB, pregLoop); + ReduceSum(vSum, mainA, pregLoop); + DataCopy(workAddr + i, vSum, pregMerge); + } + LocalMemBar(); + pregLoop = UpdateMask(meanSreg); + DataCopy(mainA, workAddr); + ReduceSum(vSum, mainA, pregLoop); + if constexpr (IS_RSTD) { + Muls(vSum, vSum, avgFactor, pregMerge); + Adds(vSum, vSum, epsilon, pregMerge); + Sqrt(vSum, vSum, pregMerge); + Duplicate(vDupReg, float(1.0), pregMerge); + Div(vSum, vDupReg, vSum, pregMerge); + } + DataCopy(dstAddr + dstOffset, vSum, pregMerge); + } +} + +/** + * @brief Use VF to Compute reduceSum(multi line). + * dstLocal = reduceSum((x1+x2)^2) + * If HAS_XOUT_FP32 is true, return xOutFp32 = x1.to(float) + x2.to(float). + * If IS_RSTD is true, dstLocal = 1.0 / sqrt(avgFactor * reduceSum((x1+x2)^2) + epsilon) + * Use float32 VL_LENGTH + */ +template +__aicore__ inline void ReduceSumRstdMulti( + LocalTensor& rstdLocal, LocalTensor& xOutLocal, LocalTensor& xOutFp32Local, + LocalTensor& x1Local, LocalTensor& x2Local, LocalTensor& workLocal, uint32_t rstdOffsetStart, + uint32_t count, uint32_t powerSplit, uint32_t repeatTimes, float avgFactor = 1.0f, float epsilon = 0.0f) +{ + uint32_t rstdOffset = rstdOffsetStart; + uint32_t remainTile = count - powerSplit; + uint16_t remainRepeats = remainTile / (2 * V_LENGTH); + + uint32_t masterTile = powerSplit - remainTile; + uint16_t masterRepeats = masterTile / (2 * V_LENGTH); + + uint32_t mergeTile = powerSplit / (2 * V_LENGTH); + uint16_t mergeRepeats = mergeTile / (2 * V_LENGTH); + + uint32_t meanTile = mergeRepeats == 0 ? mergeTile : mergeRepeats; + + __local_mem__ U* x1MainAddr = (__ubuf__ U*)x1Local.GetPhyAddr(); + __local_mem__ U* x1TailAddr = (__ubuf__ U*)x1Local.GetPhyAddr() + int64_t(powerSplit); + __local_mem__ U* x1MasterAddr = (__ubuf__ U*)x1Local.GetPhyAddr() + int64_t(remainTile); + __local_mem__ U* x2MainAddr = (__ubuf__ U*)x2Local.GetPhyAddr(); + __local_mem__ U* x2TailAddr = (__ubuf__ U*)x2Local.GetPhyAddr() + int64_t(powerSplit); + __local_mem__ U* x2MasterAddr = (__ubuf__ U*)x2Local.GetPhyAddr() + int64_t(remainTile); + __local_mem__ U *xOutMainAddr, *xOutTailAddr, *xOutMasterAddr; + if constexpr (HAS_XOUT) { + xOutMainAddr = (__ubuf__ U*)xOutLocal.GetPhyAddr(); + xOutTailAddr = (__ubuf__ U*)xOutLocal.GetPhyAddr() + int64_t(powerSplit); + xOutMasterAddr = (__ubuf__ U*)xOutLocal.GetPhyAddr() + int64_t(remainTile); + } + __local_mem__ float *xOutFp32MainAddr, *xOutFp32TailAddr, *xOutFp32MasterAddr; + if constexpr (HAS_XOUT_FP32) { + xOutFp32MainAddr = (__ubuf__ float*)xOutFp32Local.GetPhyAddr(); + xOutFp32TailAddr = (__ubuf__ float*)xOutFp32Local.GetPhyAddr() + int64_t(powerSplit); + xOutFp32MasterAddr = (__ubuf__ float*)xOutFp32Local.GetPhyAddr() + int64_t(remainTile); + } + __local_mem__ float* workAddr = (__ubuf__ float*)workLocal.GetPhyAddr(); + __local_mem__ float* rstdAddr = (__ubuf__ float*)rstdLocal.GetPhyAddr(); + + __VEC_SCOPE__ + { + MaskReg pregMerge = CreateMask(); + + for (uint16_t row = 0; row < (uint16_t)repeatTimes; row++) { + uint32_t remainSreg = remainTile; + uint32_t masterSreg = masterTile; + uint32_t mergeSreg = mergeTile; + uint32_t meanSreg = meanTile; + RegTensor x1MainA, x1MainB, x1TailA, x1TailB; + RegTensor x2MainA, x2MainB, x2TailA, x2TailB; + RegTensor mainA, mainB, tailA, tailB, vSum, vDupReg, rstdReg; + MaskReg pregLoop; + + for (uint16_t i = 0; i < (uint16_t)remainRepeats; ++i) { + pregLoop = UpdateMask(remainSreg); + uint16_t offset = i * 2 * V_LENGTH; + // 1. Copy in reg + LoadTwoCloseRegVF(x1MainA, x1MainB, x1MainAddr, offset); + LoadTwoCloseRegVF(x1TailA, x1TailB, x1TailAddr, offset); + LoadTwoCloseRegVF(x2MainA, x2MainB, x2MainAddr, offset); + LoadTwoCloseRegVF(x2TailA, x2TailB, x2TailAddr, offset); + // 2. Cast add + CastAddVF(mainA, x1MainA, x2MainA, pregLoop); + CastAddVF(tailA, x1TailA, x2TailA, pregLoop); + CastAddVF(mainB, x1MainB, x2MainB, pregLoop); + CastAddVF(tailB, x1TailB, x2TailB, pregLoop); + if constexpr (HAS_XOUT) { + CastStoreTwoCloseRegVF(xOutMainAddr, mainA, mainB, offset, pregLoop); + CastStoreTwoCloseRegVF(xOutTailAddr, tailA, tailB, offset, pregLoop); + } + if constexpr (HAS_XOUT_FP32) { + CastStoreTwoCloseRegVF(xOutFp32MainAddr, mainA, mainB, offset, pregLoop); + CastStoreTwoCloseRegVF(xOutFp32TailAddr, tailA, tailB, offset, pregLoop); + } + // 3. Cal x^2 + Mul(mainA, mainA, mainA, pregLoop); + Mul(tailA, tailA, tailA, pregLoop); + Mul(mainB, mainB, mainB, pregLoop); + Mul(tailB, tailB, tailB, pregLoop); + Add(mainA, mainA, tailA, pregLoop); + Add(mainB, mainB, tailB, pregLoop); + Add(mainA, mainA, mainB, pregLoop); + ReduceSum(vSum, mainA, pregLoop); + DataCopy(workAddr + i, vSum, pregMerge); + } + for (uint16_t i = 0; i < (uint16_t)masterRepeats; ++i) { + pregLoop = UpdateMask(masterSreg); + uint16_t offset = i * 2 * V_LENGTH; + // 1. Copy in reg + LoadTwoCloseRegVF(x1MainA, x1MainB, x1MasterAddr, offset); + LoadTwoCloseRegVF(x2MainA, x2MainB, x2MasterAddr, offset); + // 2. Cast add + CastAddVF(mainA, x1MainA, x2MainA, pregLoop); + CastAddVF(mainB, x1MainB, x2MainB, pregLoop); + if constexpr (HAS_XOUT) { + CastStoreTwoCloseRegVF(xOutMasterAddr, mainA, mainB, offset, pregLoop); + } + if constexpr (HAS_XOUT_FP32) { + CastStoreTwoCloseRegVF(xOutFp32MasterAddr, mainA, mainB, offset, pregLoop); + } + // 3. Cal x^2 + Mul(mainA, mainA, mainA, pregLoop); + Mul(mainB, mainB, mainB, pregLoop); + Add(mainA, mainA, mainB, pregLoop); + ReduceSum(vSum, mainA, pregLoop); + DataCopy(workAddr + remainRepeats + i, vSum, pregMerge); + } + LocalMemBar(); + for (uint16_t i = 0; i < (uint16_t)mergeRepeats; ++i) { + pregLoop = UpdateMask(mergeSreg); + uint16_t offset = i * 2 * V_LENGTH; + LoadTwoCloseRegVF(mainA, mainB, workAddr, offset); + Add(mainA, mainA, mainB, pregLoop); + ReduceSum(vSum, mainA, pregLoop); + DataCopy(workAddr + i, vSum, pregMerge); + } + LocalMemBar(); + pregLoop = UpdateMask(meanSreg); + DataCopy(mainA, workAddr); + ReduceSum(vSum, mainA, pregLoop); + if constexpr (IS_RSTD) { + Muls(vSum, vSum, avgFactor, pregMerge); + Adds(vSum, vSum, epsilon, pregMerge); + Sqrt(vSum, vSum, pregMerge); + Duplicate(vDupReg, float(1.0), pregMerge); + Div(rstdReg, vDupReg, vSum, pregMerge); + } + DataCopy(rstdAddr + rstdOffset, rstdReg, pregMerge); + + rstdOffset++; + x1MainAddr += int64_t(count); + x1TailAddr += int64_t(count); + x1MasterAddr += int64_t(count); + x2MainAddr += int64_t(count); + x2TailAddr += int64_t(count); + x2MasterAddr += int64_t(count); + if constexpr (HAS_XOUT) { + xOutMainAddr += int64_t(count); + xOutTailAddr += int64_t(count); + xOutMasterAddr += int64_t(count); + } + if constexpr (HAS_XOUT_FP32) { + xOutFp32MainAddr += int64_t(count); + xOutFp32TailAddr += int64_t(count); + xOutFp32MasterAddr += int64_t(count); + } + } + } +} + +template +__aicore__ inline void ComputeRstdNewtonRaphsonReg( + RegTensor& var, RegTensor& rstd, MaskReg& preg, float epsilon) +{ + static constexpr float POS_INF = 3.40282366920938E+38; + static constexpr float SCALAR1 = -0.5; + static constexpr float SCALAR2 = 1.5; + static constexpr float SCALAR3 = 0.5; + static constexpr float SCALAR0 = -99.99; + + RegTensor r; + RegTensor y; + RegTensor s; + RegTensor t; + RegTensor one; + RegTensor scalar1; + RegTensor t1; + RegTensor t3; + RegTensor t4; + RegTensor scalarInf; + RegTensor scalarZero; + MaskReg cmpRegZero; + MaskReg cmpRegInf; + + Duplicate(scalarInf, POS_INF, preg); + Duplicate(scalarZero, float(0.0), preg); + Duplicate(one, float(1.0), preg); + Duplicate(scalar1, SCALAR3, preg); + Duplicate(t1, SCALAR2, preg); + Duplicate(s, float(1.0), preg); + + Adds(var, var, epsilon, preg); + if constexpr (NEED_MAX) { + Maxs(var, var, SCALAR0, preg); + } + Div(r, one, var, preg); + Sqrt(y, r, preg); + Muls(t, var, SCALAR1, preg); + Mul(t, t, y, preg); + Mula(t1, t, y, preg); + Mul(rstd, y, t1, preg); + Muls(t3, var, float(-1.0), preg); + Mula(s, t3, r, preg); + Muls(t4, rstd, float(-1.0), preg); + Mula(r, t4, rstd, preg); + Mula(s, var, r, preg); + Mul(s, s, rstd, preg); + Mula(rstd, s, scalar1, preg); + CompareScalar(cmpRegZero, var, POS_INF, preg); + Select(rstd, scalarZero, rstd, cmpRegZero); + CompareScalar(cmpRegInf, var, float(0.0), preg); + Select(rstd, scalarInf, rstd, cmpRegInf); +} + +template +__aicore__ inline void LoadTensorUnAlignForDtypeT(__local_mem__ T*& src, RegTensor& dst, + AscendC::MicroAPI::UnalignReg& uSrc, MaskReg& preg, uint32_t postUpdateStride) +{ + if constexpr (IsSameType::value) { + AscendC::MicroAPI::DataCopyUnAlign( + dst, uSrc, src, postUpdateStride); + } else { + RegTensor xB16; + RegTensor xB16Unpack; + AscendC::MicroAPI::DataCopyUnAlign( + xB16, uSrc, src, postUpdateStride); + UnPack((RegTensor&)xB16Unpack, (RegTensor&)xB16); + Cast(dst, xB16Unpack, preg); + } +} + +template +__aicore__ inline void StoreTensorUnAlignForDtypeT(__local_mem__ T*& dst, RegTensor& src, + AscendC::MicroAPI::UnalignReg& uDst, MaskReg& preg, uint32_t postUpdateStride) +{ + if constexpr (IsSameType::value) { + AscendC::MicroAPI::DataCopyUnAlign( + dst, src, uDst, postUpdateStride); + } else { + RegTensor xB16; + RegTensor xB16Pack; + Cast(xB16, src, preg); + Pack((RegTensor&)xB16Pack, (RegTensor&)xB16); + AscendC::MicroAPI::DataCopyUnAlign( + dst, xB16Pack, uDst, postUpdateStride); + } +} + +template +__aicore__ inline void LoadTensorUnAlignForDtypeT( + __local_mem__ T* src, RegTensor& dst, MaskReg& preg, uint32_t postUpdateStride) +{ + AscendC::MicroAPI::UnalignReg uSrc; + __local_mem__ T* srcTmp = src; + AscendC::MicroAPI::DataCopyUnAlignPre(uSrc, srcTmp); + LoadTensorUnAlignForDtypeT(srcTmp, dst, uSrc, preg, postUpdateStride); +} + +template +__aicore__ inline void StoreTensorUnAlignForDtypeT( + __local_mem__ T* dst, RegTensor& src, MaskReg& preg, uint32_t postUpdateStride) +{ + AscendC::MicroAPI::UnalignReg uDst; + __local_mem__ T* dstTmp = dst; + StoreTensorUnAlignForDtypeT(dstTmp, src, uDst, preg, postUpdateStride); + AscendC::MicroAPI::DataCopyUnAlignPost(dstTmp, uDst, 0); +} + +// NOTE: x is overwritten in place (x = (x - mean) * scale * rstd); only y is the +// downstream-usable result. Callers must not rely on the original x after this call. +__aicore__ inline void NormalizeWithScaleBiasReg(RegTensor& x, RegTensor& scale, + RegTensor& bias, RegTensor& mean, RegTensor& rstd, RegTensor& y, MaskReg& preg) +{ + Sub(x, x, mean, preg); + Mul(x, x, scale, preg); + Mul(x, x, rstd, preg); + Add(y, x, bias, preg); +} + +template +__aicore__ inline void ComputeRstdNewtonRaphson( + __local_mem__ float* src, __local_mem__ float* dst, uint32_t rowCount, float epsilon, + float avgFactor = 1.0f, uint32_t vectorLen = V_LENGTH) +{ + uint16_t loopRows = static_cast((rowCount + vectorLen - 1) / vectorLen); + __VEC_SCOPE__ + { + RegTensor var; + RegTensor rstd; + MaskReg pregLoop; + + uint32_t sreg = rowCount; + for (uint16_t i = 0; i < loopRows; ++i) { + pregLoop = UpdateMask(sreg); + DataCopy(var, src + i * vectorLen); + if constexpr (NEED_AVG_FACTOR) { + Muls(var, var, avgFactor, pregLoop); + } + ComputeRstdNewtonRaphsonReg(var, rstd, pregLoop, epsilon); + DataCopy(dst + i * vectorLen, rstd, pregLoop); + } + } +} + +template +__aicore__ inline void ComputeRstdNewtonRaphson( + LocalTensor srcLocal, LocalTensor dstLocal, uint32_t rowCount, float epsilon, + float avgFactor = 1.0f, uint32_t vectorLen = V_LENGTH) +{ + __local_mem__ float* src = (__local_mem__ float*)srcLocal.GetPhyAddr(); + __local_mem__ float* dst = (__local_mem__ float*)dstLocal.GetPhyAddr(); + ComputeRstdNewtonRaphson(src, dst, rowCount, epsilon, avgFactor, vectorLen); +} + +/*! + * @brief Compute ReduceSum mean + * IS_RSTD: if True, will cal rstd otherwise sum. + * @param dstLocal dst levelTensor + * @param srcLocal src LevelTensor + * @param offset dst offset + * @param count src level size, must be ONCE_VECTOR_SIZE + * @param avgFactor avgFactor for cal rstd + * @param epsilon epsilon for cal rstd + * @return + */ +template +__aicore__ inline void LevelMergeRstd( + LocalTensor& dstLocal, LocalTensor srcLocal, uint64_t offset, uint32_t count, float avgFactor = 1.0f, + float epsilon = 0.0f) +{ + uint64_t calCount = count / 4; // Div 4 for VF parallel execution. + uint32_t sreg = (uint32_t)(calCount); + uint16_t repeatTimes = CeilDivision(calCount, V_LENGTH); + uint32_t meanTile = repeatTimes; + + __local_mem__ float* src1Addr = (__ubuf__ float*)srcLocal.GetPhyAddr() + 0 * calCount; + __local_mem__ float* src2Addr = (__ubuf__ float*)srcLocal.GetPhyAddr() + 1 * calCount; + __local_mem__ float* src3Addr = (__ubuf__ float*)srcLocal.GetPhyAddr() + 2 * calCount; + __local_mem__ float* src4Addr = (__ubuf__ float*)srcLocal.GetPhyAddr() + 3 * calCount; + __local_mem__ float* dstAddr = (__ubuf__ float*)dstLocal.GetPhyAddr(); + + __VEC_SCOPE__ + { + RegTensor vRegA, vRegB, vRegC, vRegD, dstReg, vSum, vDupReg; + MaskReg pregMerge = CreateMask(); + MaskReg pregLoop; + for (uint16_t i = 0; i < repeatTimes; ++i) { + pregLoop = UpdateMask(sreg); + DataCopy(vRegA, src1Addr + i * V_LENGTH); + DataCopy(vRegB, src2Addr + i * V_LENGTH); + DataCopy(vRegC, src3Addr + i * V_LENGTH); + DataCopy(vRegD, src4Addr + i * V_LENGTH); + Add(vRegA, vRegA, vRegB, pregLoop); + Add(vRegC, vRegC, vRegD, pregLoop); + Add(dstReg, vRegA, vRegC, pregLoop); + ReduceSum(vSum, dstReg, pregLoop); + if constexpr (IS_RSTD) { + Muls(vSum, vSum, avgFactor, pregMerge); + Adds(vSum, vSum, epsilon, pregMerge); + Sqrt(vSum, vSum, pregMerge); + Duplicate(vDupReg, float(1.0), pregMerge); + Div(vSum, vDupReg, vSum, pregMerge); + } + DataCopy(dstAddr + offset, vSum, pregMerge); + } + } +} + +/*! + * @brief compute final ReduceSum result + * IS_RSTD: if True, will cal rstd otherwise sum. + * @param dstLocal dst Tensor + * @param offset dst offset + * @param level1Local level1 Tensor + * @param level2Local level2 Tensor + * @param level3Local level3 Tensor + * @param level1 level1 elements + * @param level2 level2 elements + * @param level3 level3 elements + * @param avgFactor avgFactor for cal rstd + * @param epsilon epsilon for cal rstd + * @return + */ +template +__aicore__ inline void ComputeMultiLevelRstd( + LocalTensor& dstLocal, uint32_t offset, LocalTensor& level1Local, LocalTensor& level2Local, + LocalTensor& level3Local, uint32_t& level1, uint32_t& level2, float avgFactor = 1.0f, float epsilon = 0.0f) +{ + if (level1 > 0 && level1 < ONCE_VECTOR_SIZE) { + LevelMergeRstd(dstLocal, level1Local, offset, ONCE_VECTOR_SIZE, avgFactor, epsilon); + } else if (level2 > 0 && level2 < ONCE_VECTOR_SIZE) { + LevelMergeRstd(dstLocal, level2Local, offset, ONCE_VECTOR_SIZE, avgFactor, epsilon); + } else { + LevelMergeRstd(dstLocal, level3Local, offset, ONCE_VECTOR_SIZE, avgFactor, epsilon); + } +} + +namespace NormCommonRegbase { + +template +__aicore__ inline void LoadRegForDtype( + __local_mem__ T* src, RegTensor& dst, MaskReg& preg, uint32_t offset) +{ + if constexpr (IsSameType::value) { + DataCopy(dst, src + offset); + } else { + RegTensor srcReg; + DataCopy(srcReg, src + offset); + Cast(dst, srcReg, preg); + } +} + +template +__aicore__ inline void StoreRegForDtype( + __local_mem__ T* dst, RegTensor& src, MaskReg& preg, uint32_t offset) +{ + if constexpr (IsSameType::value) { + DataCopy(dst + offset, src, preg); + } else { + RegTensor dstReg; + Cast(dstReg, src, preg); + DataCopy(dst + offset, dstReg, preg); + } +} + +template +__aicore__ inline void CalculateSquareReduceSumLessThanVL( + __local_mem__ T* xPtr, __local_mem__ float* dstPtr, uint16_t rows, uint32_t rowStride, uint32_t reduceNum) +{ + __VEC_SCOPE__ + { + RegTensor xReg; + RegTensor sumReg; + MaskReg pregLoop = UpdateMask(reduceNum); + MaskReg pregOne = CreateMask(); + for (uint16_t i = 0; i < rows; ++i) { + LoadRegForDtype(xPtr, xReg, pregLoop, static_cast(i) * rowStride); + Mul(xReg, xReg, xReg, pregLoop); + ReduceSum(sumReg, xReg, pregLoop); + DataCopy(dstPtr + i, sumReg, pregOne); + } + } +} + +template +__aicore__ inline void CalculateSquareReduceSumLessThanTwoVL( + __local_mem__ T* xPtr, __local_mem__ float* dstPtr, uint16_t rows, uint32_t rowStride, uint32_t reduceNum) +{ + uint32_t tailLen = reduceNum - V_LENGTH; + __VEC_SCOPE__ + { + RegTensor xReg; + RegTensor xFoldReg; + RegTensor sumReg; + RegTensor reduceReg; + MaskReg pregFull = CreateMask(); + MaskReg pregOne = CreateMask(); + MaskReg pregTail = UpdateMask(tailLen); + for (uint16_t i = 0; i < rows; ++i) { + uint32_t baseOffset = static_cast(i) * rowStride; + LoadRegForDtype(xPtr, xReg, pregFull, baseOffset); + LoadRegForDtype(xPtr + V_LENGTH, xFoldReg, pregTail, baseOffset); + Mul(xReg, xReg, xReg, pregFull); + Mul(xFoldReg, xFoldReg, xFoldReg, pregTail); + ShiftLefts( + (RegTensor&)xFoldReg, (RegTensor&)xFoldReg, static_cast(0), pregTail); + Add(sumReg, xReg, xFoldReg, pregFull); + ReduceSum(reduceReg, sumReg, pregFull); + DataCopy(dstPtr + i, reduceReg, pregOne); + } + } +} + +template +__aicore__ inline void CalculateSquareReduceSumCommon(__local_mem__ T* xPtr, __local_mem__ float* dstPtr, + __local_mem__ float* tmpPtr, uint16_t rows, uint32_t rowStride, uint32_t reduceNum, uint32_t foldPoint, + uint32_t tmpStride) +{ + uint16_t foldLoops = static_cast((foldPoint + V_LENGTH - 1) / V_LENGTH); + uint32_t lastNum = foldPoint / V_LENGTH; + uint32_t tail = (reduceNum > foldPoint) ? reduceNum - foldPoint : 0; + uint16_t tailCeilLoops = static_cast((tail + V_LENGTH - 1) / V_LENGTH); + uint16_t tailFullLoops = static_cast(tail / V_LENGTH); + + __VEC_SCOPE__ + { + RegTensor xReg; + RegTensor xFoldReg; + RegTensor sumReg; + RegTensor reduceReg; + MaskReg pregFull = CreateMask(); + MaskReg pregOne = CreateMask(); + MaskReg pregLoop; + + for (uint16_t i = 0; i < rows; ++i) { + uint32_t baseOffset = static_cast(i) * rowStride; + uint32_t tmpOffset = static_cast(i) * tmpStride; + for (uint16_t r = 0; r < tailFullLoops; ++r) { + uint32_t offset = static_cast(r) * V_LENGTH + baseOffset; + LoadRegForDtype(xPtr, xReg, pregFull, offset); + LoadRegForDtype(xPtr + foldPoint, xFoldReg, pregFull, offset); + Mul(xReg, xReg, xReg, pregFull); + Mul(xFoldReg, xFoldReg, xFoldReg, pregFull); + Add(sumReg, xReg, xFoldReg, pregFull); + ReduceSum(reduceReg, sumReg, pregFull); + DataCopy(tmpPtr + tmpOffset + r, reduceReg, pregOne); + } + uint32_t tailRemain = tail - static_cast(tailFullLoops) * V_LENGTH; + if (tailRemain != 0) { + pregLoop = UpdateMask(tailRemain); + uint32_t offset = static_cast(tailFullLoops) * V_LENGTH + baseOffset; + LoadRegForDtype(xPtr, xReg, pregFull, offset); + LoadRegForDtype(xPtr + foldPoint, xFoldReg, pregLoop, offset); + Mul(xReg, xReg, xReg, pregFull); + Mul(xFoldReg, xFoldReg, xFoldReg, pregLoop); + ShiftLefts( + (RegTensor&)xFoldReg, (RegTensor&)xFoldReg, static_cast(0), pregLoop); + Add(sumReg, xReg, xFoldReg, pregFull); + ReduceSum(reduceReg, sumReg, pregFull); + DataCopy( + tmpPtr + tmpOffset + tailFullLoops, reduceReg, pregOne); + } + for (uint16_t r = tailCeilLoops; r < foldLoops; ++r) { + uint32_t offset = static_cast(r) * V_LENGTH + baseOffset; + LoadRegForDtype(xPtr, xReg, pregFull, offset); + Mul(xReg, xReg, xReg, pregFull); + ReduceSum(reduceReg, xReg, pregFull); + DataCopy(tmpPtr + tmpOffset + r, reduceReg, pregOne); + } + } + LocalMemBar(); + if constexpr (LAST_LOOP_NUMS == 1) { + MaskReg pregLast = UpdateMask(lastNum); + for (uint16_t i = 0; i < rows; ++i) { + DataCopy(xReg, tmpPtr + static_cast(i) * tmpStride); + ReduceSum(reduceReg, xReg, pregLast); + DataCopy(dstPtr + i, reduceReg, pregOne); + } + } else if constexpr (LAST_LOOP_NUMS == DICHOTOMY_ADD_COEFF) { + lastNum -= V_LENGTH; + MaskReg pregLast = UpdateMask(lastNum); + for (uint16_t i = 0; i < rows; ++i) { + uint32_t tmpOffset = static_cast(i) * tmpStride; + DataCopy(xReg, tmpPtr + tmpOffset); + DataCopy(xFoldReg, tmpPtr + tmpOffset + V_LENGTH); + ShiftLefts( + (RegTensor&)xFoldReg, (RegTensor&)xFoldReg, static_cast(0), pregLast); + Add(sumReg, xReg, xFoldReg, pregFull); + ReduceSum(reduceReg, sumReg, pregFull); + DataCopy(dstPtr + i, reduceReg, pregOne); + } + } + } +} + +template +// Squares input values inside this function, then reduces each row. +__aicore__ inline void CalculateSquareReduceSum(__local_mem__ T* xPtr, __local_mem__ float* dstPtr, + __local_mem__ float* tmpPtr, uint16_t rows, uint32_t rowStride, uint32_t reduceNum, uint32_t foldPoint, + uint32_t tmpStride, uint32_t branchNum = 0) +{ + uint32_t reduceBranchNum = branchNum == 0 ? reduceNum : branchNum; + if (reduceBranchNum <= V_LENGTH) { + CalculateSquareReduceSumLessThanVL(xPtr, dstPtr, rows, rowStride, reduceNum); + } else if (reduceBranchNum <= V_LENGTH + V_LENGTH) { + CalculateSquareReduceSumLessThanTwoVL(xPtr, dstPtr, rows, rowStride, reduceNum); + } else if (reduceBranchNum <= V_LENGTH * V_LENGTH * DICHOTOMY_ADD_COEFF) { + CalculateSquareReduceSumCommon(xPtr, dstPtr, tmpPtr, rows, rowStride, reduceNum, foldPoint, tmpStride); + } else { + CalculateSquareReduceSumCommon( + xPtr, dstPtr, tmpPtr, rows, rowStride, reduceNum, foldPoint, tmpStride); + } +} + +template +__aicore__ inline void CalculateSquareReduceSum(LocalTensor& xLocal, LocalTensor& dstLocal, + LocalTensor& tmpLocal, uint16_t rows, uint32_t rowStride, uint32_t reduceNum, uint32_t foldPoint, + uint32_t blockAlign, uint32_t branchNum = 0) +{ + __local_mem__ T* xPtr = (__local_mem__ T*)xLocal.GetPhyAddr(); + __local_mem__ float* dstPtr = (__local_mem__ float*)dstLocal.GetPhyAddr(); + __local_mem__ float* tmpPtr = (__local_mem__ float*)tmpLocal.GetPhyAddr(); + uint32_t foldLoops = (foldPoint + V_LENGTH - 1) / V_LENGTH; + uint32_t tmpStride = (foldLoops + blockAlign - 1) / blockAlign * blockAlign; + CalculateSquareReduceSum(xPtr, dstPtr, tmpPtr, rows, rowStride, reduceNum, foldPoint, tmpStride, branchNum); +} + +template +__aicore__ inline void CalculateSquareReduceSum(LocalTensor& xLocal, LocalTensor& dstLocal, + TBuf& tmpBuf, uint16_t rows, uint32_t rowStride, uint32_t reduceNum, uint32_t foldPoint, + uint32_t blockAlign, uint32_t branchNum = 0) +{ + LocalTensor tmpLocal = tmpBuf.Get(); + CalculateSquareReduceSum( + xLocal, dstLocal, tmpLocal, rows, rowStride, reduceNum, foldPoint, blockAlign, branchNum); +} + +__aicore__ inline void CalculateReduceSumLessThanVL( + __local_mem__ float* xPtr, __local_mem__ float* dstPtr, uint32_t reduceNum) +{ + __VEC_SCOPE__ + { + RegTensor xReg; + RegTensor sumReg; + MaskReg pregLoop = UpdateMask(reduceNum); + MaskReg pregOne = CreateMask(); + DataCopy(xReg, xPtr); + ReduceSum(sumReg, xReg, pregLoop); + DataCopy(dstPtr, sumReg, pregOne); + } +} + +__aicore__ inline void CalculateReduceSumLessThanTwoVL( + __local_mem__ float* xPtr, __local_mem__ float* dstPtr, uint32_t reduceNum) +{ + uint32_t tailLen = reduceNum - V_LENGTH; + __VEC_SCOPE__ + { + RegTensor xReg; + RegTensor xFoldReg; + RegTensor sumReg; + RegTensor reduceReg; + MaskReg pregFull = CreateMask(); + MaskReg pregTail = UpdateMask(tailLen); + MaskReg pregOne = CreateMask(); + DataCopy(xReg, xPtr); + DataCopy(xFoldReg, xPtr + V_LENGTH); + ShiftLefts((RegTensor&)xFoldReg, (RegTensor&)xFoldReg, static_cast(0), pregTail); + Add(sumReg, xReg, xFoldReg, pregFull); + ReduceSum(reduceReg, sumReg, pregFull); + DataCopy(dstPtr, reduceReg, pregOne); + } +} + +template +__aicore__ inline void CalculateReduceSumCommon( + __local_mem__ float* xPtr, __local_mem__ float* dstPtr, __local_mem__ float* tmpPtr, uint32_t reduceNum, + uint32_t foldPoint) +{ + uint16_t foldLoops = static_cast((foldPoint + V_LENGTH - 1) / V_LENGTH); + uint32_t lastNum = foldPoint / V_LENGTH; + uint32_t tail = reduceNum - foldPoint; + uint16_t tailCeilLoops = static_cast((tail + V_LENGTH - 1) / V_LENGTH); + uint16_t tailFullLoops = static_cast(tail / V_LENGTH); + + __VEC_SCOPE__ + { + RegTensor xReg; + RegTensor xFoldReg; + RegTensor sumReg; + RegTensor reduceReg; + MaskReg pregFull = CreateMask(); + MaskReg pregOne = CreateMask(); + MaskReg pregLoop; + + for (uint16_t r = 0; r < tailFullLoops; ++r) { + uint32_t offset = static_cast(r) * V_LENGTH; + DataCopy(xReg, xPtr + offset); + DataCopy(xFoldReg, xPtr + foldPoint + offset); + Add(sumReg, xReg, xFoldReg, pregFull); + ReduceSum(reduceReg, sumReg, pregFull); + DataCopy(tmpPtr + r, reduceReg, pregOne); + } + uint32_t tailRemain = tail - static_cast(tailFullLoops) * V_LENGTH; + if (tailRemain != 0) { + pregLoop = UpdateMask(tailRemain); + uint32_t offset = static_cast(tailFullLoops) * V_LENGTH; + DataCopy(xReg, xPtr + offset); + DataCopy(xFoldReg, xPtr + foldPoint + offset); + ShiftLefts( + (RegTensor&)xFoldReg, (RegTensor&)xFoldReg, static_cast(0), pregLoop); + Add(sumReg, xReg, xFoldReg, pregFull); + ReduceSum(reduceReg, sumReg, pregFull); + DataCopy(tmpPtr + tailFullLoops, reduceReg, pregOne); + } + // Fix the original local implementations' fixed-offset bug in the remaining reduce blocks. + for (uint16_t r = 0; r < static_cast(foldLoops - tailCeilLoops); ++r) { + uint32_t offset = static_cast(tailCeilLoops + r); + DataCopy(xReg, xPtr + offset * V_LENGTH); + ReduceSum(reduceReg, xReg, pregFull); + DataCopy(tmpPtr + offset, reduceReg, pregOne); + } + LocalMemBar(); + if constexpr (LAST_LOOP_NUMS == 1) { + MaskReg pregLast = UpdateMask(lastNum); + DataCopy(xReg, tmpPtr); + ReduceSum(reduceReg, xReg, pregLast); + DataCopy(dstPtr, reduceReg, pregOne); + } else if constexpr (LAST_LOOP_NUMS == DICHOTOMY_ADD_COEFF) { + lastNum -= V_LENGTH; + MaskReg pregLast = UpdateMask(lastNum); + DataCopy(xReg, tmpPtr); + DataCopy(xFoldReg, tmpPtr + V_LENGTH); + ShiftLefts( + (RegTensor&)xFoldReg, (RegTensor&)xFoldReg, static_cast(0), pregLast); + Add(sumReg, xReg, xFoldReg, pregFull); + ReduceSum(reduceReg, sumReg, pregFull); + DataCopy(dstPtr, reduceReg, pregOne); + } + } +} + +// Reduces an fp32 buffer whose values have already been squared by the caller. +__aicore__ inline void CalculateReduceSum( + __local_mem__ float* xPtr, __local_mem__ float* dstPtr, __local_mem__ float* tmpPtr, uint32_t reduceNum, + uint32_t foldPoint) +{ + if (reduceNum <= V_LENGTH) { + CalculateReduceSumLessThanVL(xPtr, dstPtr, reduceNum); + } else if (reduceNum <= V_LENGTH + V_LENGTH) { + CalculateReduceSumLessThanTwoVL(xPtr, dstPtr, reduceNum); + } else if (reduceNum <= V_LENGTH * V_LENGTH * DICHOTOMY_ADD_COEFF) { + CalculateReduceSumCommon<1>(xPtr, dstPtr, tmpPtr, reduceNum, foldPoint); + } else { + CalculateReduceSumCommon(xPtr, dstPtr, tmpPtr, reduceNum, foldPoint); + } +} + +__aicore__ inline void CalculateReduceSum(LocalTensor& xLocal, LocalTensor& dstLocal, + LocalTensor& tmpLocal, uint32_t reduceNum, uint32_t foldPoint) +{ + __local_mem__ float* xPtr = (__local_mem__ float*)xLocal.GetPhyAddr(); + __local_mem__ float* dstPtr = (__local_mem__ float*)dstLocal.GetPhyAddr(); + __local_mem__ float* tmpPtr = (__local_mem__ float*)tmpLocal.GetPhyAddr(); + CalculateReduceSum(xPtr, dstPtr, tmpPtr, reduceNum, foldPoint); +} + +__aicore__ inline void CalculateReduceSum(LocalTensor& xLocal, LocalTensor& dstLocal, + TBuf& tmpBuf, uint32_t reduceNum, uint32_t foldPoint) +{ + LocalTensor tmpLocal = tmpBuf.Get(); + CalculateReduceSum(xLocal, dstLocal, tmpLocal, reduceNum, foldPoint); +} +} // namespace NormCommonRegbase +} // namespace NormCommon + +#endif // REDUCE_COMMON_REGBASE_H_RMS_NORM diff --git a/xllm_ops/gamma_add_rms_norm/op_host/CMakeLists.txt b/xllm_ops/gamma_add_rms_norm/op_host/CMakeLists.txt new file mode 100644 index 0000000..0a44d3e --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_host/CMakeLists.txt @@ -0,0 +1,25 @@ +# Copyright 2026 The xLLM Authors. All Rights Reserved. + +add_op_to_compiled_list() + +if (BUILD_OPEN_PROJECT) + target_sources(op_host_aclnnExc PRIVATE + gamma_add_rms_norm_def.cpp + ) +endif() + +add_ops_compile_options( + OP_NAME GammaAddRmsNorm + OPTIONS + --cce-auto-sync=off + -Wno-deprecated-declarations + -Wno-error +) + +if (NOT BUILD_OPS_RTY_KERNEL) + add_modules_sources(OPTYPE gamma_add_rms_norm ACLNNTYPE aclnn_exclude) + target_include_directories(${OPHOST_NAME}_tiling_obj PRIVATE + ${CMAKE_CURRENT_SOURCE_DIR} + ${ASCEND_CANN_PACKAGE_PATH}/${SYSTEM_PREFIX}/pkg_inc + ) +endif() diff --git a/xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_def.cpp b/xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_def.cpp new file mode 100644 index 0000000..0373250 --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_def.cpp @@ -0,0 +1,74 @@ +/** + * This program is free software, you can redistribute it and/or modify. + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * This file is a part of the CANN Open Software. + * Licensed under 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 gamma_add_rms_norm_def.cpp + * \brief + */ +#include "register/op_def_registry.h" + +namespace ops { +class GammaAddRmsNorm : public OpDef { +public: + explicit GammaAddRmsNorm(const char* name) : OpDef(name) + { + this->Input("x1") + .ParamType(REQUIRED) + .DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16}) + .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) + .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("x2") + .ParamType(REQUIRED) + .DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16}) + .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) + .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("gamma") + .ParamType(REQUIRED) + .DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16}) + .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) + .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) + .AutoContiguous(); + this->Output("y") + .ParamType(REQUIRED) + .DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16}) + .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) + .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) + .AutoContiguous(); + this->Output("rstd") + .ParamType(REQUIRED) + .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT}) + .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) + .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) + .AutoContiguous(); + this->Output("x") + .ParamType(REQUIRED) + .DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16}) + .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) + .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) + .AutoContiguous(); + this->Attr("epsilon").AttrType(OPTIONAL).Float(1e-6f); + this->Attr("addGammaOffset").AttrType(OPTIONAL).Bool(false); + + this->AICore().AddConfig("ascend910b"); + this->AICore().AddConfig("ascend910_93"); + + OpAICoreConfig regbaseCfg; + regbaseCfg.DynamicCompileStaticFlag(true) + .DynamicRankSupportFlag(true) + .DynamicShapeSupportFlag(true) + .ExtendCfgInfo("opFile.value", "gamma_add_rms_norm_apt"); + this->AICore().AddConfig("ascend950", regbaseCfg); + } +}; +OP_ADD(GammaAddRmsNorm); +} // namespace ops diff --git a/xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_error_log.h b/xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_error_log.h new file mode 100644 index 0000000..0af748f --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_error_log.h @@ -0,0 +1,37 @@ +/** + * Copyright (c) 2026 The xLLM Authors. All Rights Reserved. + */ + +#pragma once + +#include "log/log.h" + +#ifndef OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON +#define OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(opname, param, actual, reason) \ + OP_LOGE(opname, "Invalid shape for %s, actual: %s, reason: %s", param, actual, reason) +#endif + +#ifndef OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON +#define OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(opname, param, actual, reason) \ + OP_LOGE(opname, "Invalid shapes for %s, actual: %s, reason: %s", param, actual, reason) +#endif + +#ifndef OP_LOGE_FOR_INVALID_SHAPEDIM +#define OP_LOGE_FOR_INVALID_SHAPEDIM(opname, param, actual, expected) \ + OP_LOGE(opname, "Invalid shape dim for %s, actual: %s, expected: %s", param, actual, expected) +#endif + +#ifndef OP_LOGE_FOR_INVALID_SHAPEDIMS_WITH_REASON +#define OP_LOGE_FOR_INVALID_SHAPEDIMS_WITH_REASON(opname, param, actual, reason) \ + OP_LOGE(opname, "Invalid shape dims for %s, actual: %s, reason: %s", param, actual, reason) +#endif + +#ifndef OP_LOGE_FOR_INVALID_VALUE +#define OP_LOGE_FOR_INVALID_VALUE(opname, param, actual, expected) \ + OP_LOGE(opname, "Invalid value for %s, actual: %s, expected: %s", param, actual, expected) +#endif + +#ifndef OP_LOGE_FOR_INVALID_VALUE_WITH_REASON +#define OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opname, param, actual, reason) \ + OP_LOGE(opname, "Invalid value for %s, actual: %s, reason: %s", param, actual, reason) +#endif diff --git a/xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_infershape.cpp b/xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_infershape.cpp new file mode 100644 index 0000000..6474eed --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_infershape.cpp @@ -0,0 +1,85 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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 gamma_add_rms_norm_infershape.cpp + * \brief + */ +#include "log/log.h" +#include "util/shape_util.h" +#include "register/op_impl_registry.h" + +static constexpr int IDX_0 = 0; +static constexpr int IDX_1 = 1; +static constexpr int IDX_2 = 2; + +using namespace ge; +using namespace Ops::Base; + +namespace ops { + +static ge::graphStatus InferShape4GammaAddRmsNorm(gert::InferShapeContext* context) +{ + OP_LOGD(context, "Begin to do InferShape4GammaAddRmsNorm"); + + // get input shapes + const gert::Shape* x1Shape = context->GetInputShape(IDX_0); + OP_CHECK_NULL_WITH_CONTEXT(context, x1Shape); + const gert::Shape* gammaShape = context->GetInputShape(IDX_2); + OP_CHECK_NULL_WITH_CONTEXT(context, gammaShape); + // get output shapes + gert::Shape* yShape = context->GetOutputShape(IDX_0); + gert::Shape* rstdShape = context->GetOutputShape(IDX_1); + gert::Shape* xShape = context->GetOutputShape(IDX_2); + OP_CHECK_NULL_WITH_CONTEXT(context, yShape); + OP_CHECK_NULL_WITH_CONTEXT(context, rstdShape); + OP_CHECK_NULL_WITH_CONTEXT(context, xShape); + *yShape = *x1Shape; + *xShape = *x1Shape; + + size_t xDimNum = x1Shape->GetDimNum(); + size_t gammaDimNum = gammaShape->GetDimNum(); + + if (IsUnknownRank(*x1Shape) || IsUnknownRank(*gammaShape)) { + SetUnknownRank(*rstdShape); + OP_LOGD(context, "End to do InferShape4GammaAddRmsNorm with unknown rank."); + return GRAPH_SUCCESS; + } + + OP_CHECK_IF( + xDimNum < gammaDimNum, OP_LOGE(context, "x dim num should not be smaller than gamma dim num."), + return GRAPH_FAILED); + + rstdShape->SetDimNum(xDimNum); + for (size_t rmsIdx = 0; rmsIdx < xDimNum; rmsIdx++) { + if (rmsIdx < xDimNum - gammaDimNum) { + rstdShape->SetDim(rmsIdx, x1Shape->GetDim(rmsIdx)); + } else { + rstdShape->SetDim(rmsIdx, 1); + } + } + + OP_LOGD(context, "End to do InferShape4GammaAddRmsNorm"); + return GRAPH_SUCCESS; +} + +static graphStatus InferDataType4GammaAddRmsNorm(gert::InferDataTypeContext* context) +{ + OP_LOGD(context, "Begin to do InferDataType4GammaAddRmsNorm"); + context->SetOutputDataType(IDX_0, context->GetInputDataType(IDX_0)); + context->SetOutputDataType(IDX_1, DT_FLOAT); + context->SetOutputDataType(IDX_2, context->GetInputDataType(IDX_0)); + OP_LOGD(context, "End to do InferDataType4GammaAddRmsNorm"); + return GRAPH_SUCCESS; +} + +IMPL_OP_INFERSHAPE(GammaAddRmsNorm).InferShape(InferShape4GammaAddRmsNorm).InferDataType(InferDataType4GammaAddRmsNorm); +} // namespace ops diff --git a/xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_tiling.cpp b/xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_tiling.cpp new file mode 100644 index 0000000..17decc1 --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_tiling.cpp @@ -0,0 +1,563 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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 gamma_add_rms_norm_tiling.cpp + * \brief + */ + +#include "op_common/op_host/util/math_util.h" +#include "tiling_base/tiling_util.h" +#include "gamma_add_rms_norm_tiling.h" + +namespace optiling { +constexpr uint32_t RMS_NORM_KEY = 0; +constexpr uint32_t PRE_RMS_NORM = 100; +constexpr uint32_t POST_RMS_NORM = 1000; +constexpr uint32_t DTYPE_KEY_FP16 = 1; +constexpr uint32_t DTYPE_KEY_FP32 = 2; +constexpr uint32_t DTYPE_KEY_BF16 = 3; +constexpr uint32_t UB_USED = 1024; +constexpr uint32_t UB_FACTOR_B16 = 12288; +constexpr uint32_t UB_FACTOR_B32 = 10240; +constexpr uint32_t UB_FACTOR_B16_CUTD = 12096; +constexpr uint32_t UB_FACTOR_B32_CUTD = 9696; +constexpr uint32_t BLOCK_ALIGN_NUM = 16; +constexpr uint32_t FLOAT_BLOCK_ALIGN_NUM = 8; +constexpr uint32_t SMALL_REDUCE_NUM = 2000; +constexpr uint32_t MODE_NORMAL = 0; +constexpr uint32_t MODE_SPLIT_D = 1; +constexpr uint32_t MODE_MERGE_N = 2; +constexpr uint32_t MODE_SINGLE_N = 3; +constexpr uint32_t MODE_MULTI_N = 4; +constexpr int32_t RMS_INPUT_X1_INDEX = 0; +constexpr int32_t RMS_INPUT_X2_INDEX = 1; +constexpr int32_t RMS_INPUT_GAMMA_INDEX = 2; +constexpr int32_t RMS_OUTPUT_Y_INDEX = 0; +constexpr int32_t RMS_OUTPUT_RSTD_INDEX = 1; +constexpr int32_t RMS_OUTPUT_X_INDEX = 2; +constexpr size_t MAX_DIM_NUM = 8; +constexpr size_t MIN_DIM_X = 1; +constexpr size_t MIN_DIM_GAMMA = 1; +constexpr size_t FP32_WEIGHT = 24; +constexpr size_t OTHER_WEIGHT = 18; +constexpr size_t DIV_FACTOR = 260; +constexpr size_t FLOAT_PER_REPEAT = 64; +constexpr size_t USE_SIZE = 256; +constexpr size_t NUM = 2; +constexpr int32_t TEN = 10; + +constexpr int32_t PERFORMANC_DIM_ZERO = 0; +constexpr int32_t PERFORMANC_DIM_ONE = 1; +constexpr int32_t PERFORMANC_DIM_TWO = 2; +constexpr int32_t PERFORMANC_DIM_THREE = 3; +constexpr int32_t PERFORMANC_DIM_ONE_MAX = 512; +constexpr int32_t PERFORMANC_DIM_TWO_MAX = 8; +constexpr int32_t PERFORMANC_DIM_THREE_MAX = 5120; + +static uint8_t getPerformanceFlag(uint32_t num_col, const gert::Shape& x_shape, const gert::Shape& gamma_shape, + uint32_t xDtypeKey, platform_ascendc::SocVersion socVersion) +{ + uint8_t isPerformance = 0; + if(socVersion != platform_ascendc::SocVersion::ASCEND910B) { + return isPerformance; + } + size_t xDimNum = x_shape.GetDimNum(); + size_t gammaDimNum = gamma_shape.GetDimNum(); + bool dimOK = ((xDimNum == PERFORMANC_DIM_TWO || xDimNum == PERFORMANC_DIM_THREE) && gammaDimNum == PERFORMANC_DIM_ONE); + bool sizeOk = num_col <= PERFORMANC_DIM_THREE_MAX && + ((xDimNum == PERFORMANC_DIM_TWO && x_shape.GetDim(PERFORMANC_DIM_ZERO) <= PERFORMANC_DIM_ONE_MAX) || + (xDimNum == PERFORMANC_DIM_THREE && x_shape.GetDim(PERFORMANC_DIM_ZERO) <= PERFORMANC_DIM_ONE_MAX && x_shape.GetDim(PERFORMANC_DIM_ONE) <= PERFORMANC_DIM_TWO_MAX)); + bool dtypeOk = (xDtypeKey == DTYPE_KEY_FP16 || xDtypeKey == DTYPE_KEY_BF16); + if(dimOK && sizeOk && dtypeOk) { + isPerformance = 1; + } + return isPerformance; +} + +static void SetByDtype(ge::DataType dataType, uint32_t& dtypeKey, uint32_t& dataPerBlock) +{ + switch (dataType) { + case ge::DT_FLOAT16: + dtypeKey = DTYPE_KEY_FP16; + dataPerBlock = BLOCK_ALIGN_NUM; + break; + case ge::DT_BF16: + dtypeKey = DTYPE_KEY_BF16; + dataPerBlock = BLOCK_ALIGN_NUM; + break; + default: + dtypeKey = DTYPE_KEY_FP32; + dataPerBlock = FLOAT_BLOCK_ALIGN_NUM; + break; + } +} +static bool CheckNullptr(const gert::TilingContext* context, uint32_t& normKey) +{ + const gert::StorageShape* x1_shape = context->GetInputShape(RMS_INPUT_X1_INDEX); + const gert::StorageShape* x2_shape = context->GetInputShape(RMS_INPUT_X2_INDEX); + const gert::StorageShape* gamma_shape = context->GetInputShape(RMS_INPUT_GAMMA_INDEX); + const gert::StorageShape* y_shape = context->GetOutputShape(RMS_OUTPUT_Y_INDEX); + const gert::StorageShape* rstd_shape = context->GetOutputShape(RMS_OUTPUT_RSTD_INDEX); + const gert::StorageShape* x_shape = context->GetOutputShape(RMS_OUTPUT_X_INDEX); + + OP_CHECK_NULL_WITH_CONTEXT(context, x1_shape); + OP_CHECK_NULL_WITH_CONTEXT(context, x2_shape); + OP_CHECK_NULL_WITH_CONTEXT(context, gamma_shape); + OP_CHECK_NULL_WITH_CONTEXT(context, y_shape); + OP_CHECK_NULL_WITH_CONTEXT(context, rstd_shape); + OP_CHECK_NULL_WITH_CONTEXT(context, x_shape); + + normKey = RMS_NORM_KEY; + if (rstd_shape->GetOriginShape().GetShapeSize() <= 0 && x_shape->GetOriginShape().GetShapeSize() <= 0) { + normKey = POST_RMS_NORM; + } else if (rstd_shape->GetOriginShape().GetShapeSize() <= 0) { + normKey = PRE_RMS_NORM; + } + return true; +} +static bool CheckInputOutputDim(const gert::TilingContext* context, uint32_t normKey) +{ + const gert::StorageShape* x1_shape = context->GetInputShape(RMS_INPUT_X1_INDEX); + const gert::StorageShape* x2_shape = context->GetInputShape(RMS_INPUT_X2_INDEX); + const gert::StorageShape* gamma_shape = context->GetInputShape(RMS_INPUT_GAMMA_INDEX); + const gert::StorageShape* y_shape = context->GetOutputShape(RMS_OUTPUT_Y_INDEX); + const gert::StorageShape* rstd_shape = context->GetOutputShape(RMS_OUTPUT_RSTD_INDEX); + const gert::StorageShape* x_shape = context->GetOutputShape(RMS_OUTPUT_X_INDEX); + + size_t x1DimNum = x1_shape->GetStorageShape().GetDimNum(); + size_t x2DimNum = x2_shape->GetStorageShape().GetDimNum(); + size_t gammaDimNum = gamma_shape->GetStorageShape().GetDimNum(); + size_t yDimNum = y_shape->GetStorageShape().GetDimNum(); + size_t rstdDimNum = rstd_shape->GetStorageShape().GetDimNum(); + size_t xDimNum = x_shape->GetStorageShape().GetDimNum(); + + OP_CHECK_IF( + x1DimNum > MAX_DIM_NUM || x1DimNum < MIN_DIM_X, + OP_LOGE_FOR_INVALID_SHAPEDIM( + context->GetNodeName(), "x1", std::to_string(x1DimNum).c_str(), "within the range [1, 8]"), + return false); + if (normKey == RMS_NORM_KEY) { + OP_CHECK_IF( + gammaDimNum > MAX_DIM_NUM || gammaDimNum < MIN_DIM_GAMMA, + OP_LOGE_FOR_INVALID_SHAPEDIM( + context->GetNodeName(), "gamma", std::to_string(gammaDimNum).c_str(), "within the range [1, 8]"), + return false); + OP_CHECK_IF( + x1DimNum < gammaDimNum, + OP_LOGE_FOR_INVALID_SHAPEDIMS_WITH_REASON( + context->GetNodeName(), "x1 and gamma", + (std::to_string(x1DimNum) + " and " + std::to_string(gammaDimNum)).c_str(), + "The shape dim of x1 should be greater than or equal to the shape dim of gamma"), + return false); + } else if (normKey == PRE_RMS_NORM || normKey == POST_RMS_NORM) { + OP_CHECK_IF( + gammaDimNum != 2, + OP_LOGE_FOR_INVALID_SHAPEDIM(context->GetNodeName(), "gamma", std::to_string(gammaDimNum).c_str(), "2"), + return false); + } + OP_CHECK_IF( + x1DimNum != yDimNum, + OP_LOGE_FOR_INVALID_SHAPEDIMS_WITH_REASON( + context->GetNodeName(), "x1 and y", (std::to_string(x1DimNum) + " and " + std::to_string(yDimNum)).c_str(), + "The shape dims of x1 and y should be the same"), + return false); + + OP_CHECK_IF( + x1DimNum != x2DimNum, + OP_LOGE_FOR_INVALID_SHAPEDIMS_WITH_REASON( + context->GetNodeName(), "x1 and x2", + (std::to_string(x1DimNum) + " and " + std::to_string(x2DimNum)).c_str(), + "The shape dims of x1 and x2 should be the same"), + return false); + + if (normKey == RMS_NORM_KEY) { + OP_CHECK_IF( + (yDimNum != xDimNum) || (xDimNum != x1DimNum) || (rstdDimNum != x1DimNum), + OP_LOGE_FOR_INVALID_SHAPEDIMS_WITH_REASON( + context->GetNodeName(), "y, x, rstd and x1", + (std::to_string(yDimNum) + ", " + std::to_string(xDimNum) + ", " + std::to_string(rstdDimNum) + + " and " + std::to_string(x1DimNum)).c_str(), + "The shape dims of y, x, rstd and x1 should be the same"), + return false); + } else if (normKey == PRE_RMS_NORM) { + OP_CHECK_IF( + (yDimNum != xDimNum) || (xDimNum != x1DimNum), + OP_LOGE_FOR_INVALID_SHAPEDIMS_WITH_REASON( + context->GetNodeName(), "y, x and x1", + (std::to_string(yDimNum) + ", " + std::to_string(xDimNum) + " and " + std::to_string(x1DimNum)).c_str(), + "The shape dims of y, x and x1 should be the same"), + return false); + } + return true; +} + +static bool CheckInputOutputShape(const gert::TilingContext* context, uint32_t normKey) +{ + OP_CHECK_IF(!CheckInputOutputDim(context, normKey), OP_LOGE(context, "Input Dim invalid."), return false); + const gert::StorageShape* x1_shape = context->GetInputShape(RMS_INPUT_X1_INDEX); + const gert::StorageShape* x2_shape = context->GetInputShape(RMS_INPUT_X2_INDEX); + const gert::StorageShape* gamma_shape = context->GetInputShape(RMS_INPUT_GAMMA_INDEX); + const gert::StorageShape* y_shape = context->GetOutputShape(RMS_OUTPUT_Y_INDEX); + const gert::StorageShape* rstd_shape = context->GetOutputShape(RMS_OUTPUT_RSTD_INDEX); + const gert::StorageShape* x_shape = context->GetOutputShape(RMS_OUTPUT_X_INDEX); + + size_t x1DimNum = x1_shape->GetStorageShape().GetDimNum(); + size_t gammaDimNum = gamma_shape->GetStorageShape().GetDimNum(); + + for (uint32_t i = 0; i < x1DimNum; i++) { + OP_CHECK_IF( + x1_shape->GetStorageShape().GetDim(i) == 0, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + context->GetNodeName(), "x1", Ops::Base::ToString(x1_shape->GetStorageShape()).c_str(), + "x1 cannot be an empty tensor"), + return false); + OP_CHECK_IF( + x2_shape->GetStorageShape().GetDim(i) != x1_shape->GetStorageShape().GetDim(i), + OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( + context->GetNodeName(), "x2 and x1", + (Ops::Base::ToString(x2_shape->GetStorageShape()) + " and " + + Ops::Base::ToString(x1_shape->GetStorageShape())).c_str(), + "The shapes of x2 and x1 should be the same"), + return false); + OP_CHECK_IF( + (y_shape->GetStorageShape().GetDim(i) != x1_shape->GetStorageShape().GetDim(i)), + OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( + context->GetNodeName(), "y and x1", + (Ops::Base::ToString(y_shape->GetStorageShape()) + " and " + + Ops::Base::ToString(x1_shape->GetStorageShape())).c_str(), + "The shapes of y and x1 should be the same"), + return false); + // x out shape check by mode + if (normKey == RMS_NORM_KEY || normKey == PRE_RMS_NORM) { + OP_CHECK_IF( + (x_shape->GetStorageShape().GetDim(i) != x1_shape->GetStorageShape().GetDim(i)), + OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( + context->GetNodeName(), "x and x1", + (Ops::Base::ToString(x_shape->GetStorageShape()) + " and " + + Ops::Base::ToString(x1_shape->GetStorageShape())).c_str(), + "The shapes of x and x1 should be the same"), + return false); + } + } + // rstd out shape check by mode + if (normKey == RMS_NORM_KEY) { + for (uint32_t i = 0; i < x1DimNum - gammaDimNum; i++) { + OP_CHECK_IF( + rstd_shape->GetStorageShape().GetDim(i) != x2_shape->GetStorageShape().GetDim(i), + OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( + context->GetNodeName(), "rstd and x1", + (Ops::Base::ToString(rstd_shape->GetStorageShape()) + " and " + + Ops::Base::ToString(x1_shape->GetStorageShape())).c_str(), + ("The shape of rstd should be the same as the first " + std::to_string(x1DimNum - gammaDimNum) + + " dim of x1").c_str()), + return false); + } + for (uint32_t i = 0; i < gammaDimNum; i++) { + OP_CHECK_IF( + gamma_shape->GetStorageShape().GetDim(i) != + x1_shape->GetStorageShape().GetDim(x1DimNum - gammaDimNum + i), + OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( + context->GetNodeName(), "gamma and x1", + (Ops::Base::ToString(gamma_shape->GetStorageShape()) + " and " + + Ops::Base::ToString(x1_shape->GetStorageShape())).c_str(), + ("The shape of gamma should be equal to the last " + std::to_string(gammaDimNum) + " dim of x1") + .c_str()), + return false); + OP_CHECK_IF( + rstd_shape->GetStorageShape().GetDim(x1DimNum - 1 - i) != 1, + OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON( + context->GetNodeName(), "rstd", + Ops::Base::ToString(rstd_shape->GetStorageShape()).c_str(), + ("The " + std::to_string(x1DimNum - 1 - i) + "th dimension of rstd must be 1").c_str()), + return false); + } + } else if (normKey == PRE_RMS_NORM || normKey == POST_RMS_NORM) { + OP_CHECK_IF( + (gamma_shape->GetStorageShape().GetDim(0) != 1 || + gamma_shape->GetStorageShape().GetDim(gammaDimNum - 1) != x1_shape->GetStorageShape().GetDim(x1DimNum - 1)), + OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON( + context->GetNodeName(), "gamma and x1", + (Ops::Base::ToString(gamma_shape->GetStorageShape()) + " and " + + Ops::Base::ToString(x1_shape->GetStorageShape())).c_str(), + "The first dim of gamma should be 1 and the last dim of gamma and x1 must be the same"), + return false); + } + return true; +} + +static void GetCompileParameters( + gert::TilingContext* context, uint32_t& numCore, uint64_t& ubSize, + platform_ascendc::SocVersion& socVersion) +{ + auto ptrCompileInfo = reinterpret_cast(context->GetCompileInfo()); + if (ptrCompileInfo == nullptr) { + auto ascendc_platform = platform_ascendc::PlatformAscendC(context->GetPlatformInfo()); + socVersion = ascendc_platform.GetSocVersion(); + numCore = ascendc_platform.GetCoreNumAiv(); + ascendc_platform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize); + } else { + numCore = ptrCompileInfo->totalCoreNum; + ubSize = ptrCompileInfo->totalUbSize; + socVersion = ptrCompileInfo->socVersion; + } + ubSize -= UB_USED; +} + +static void CalculateRowAndColParameters( + gert::TilingContext* context, uint32_t normKey, uint32_t& numRow, uint32_t& numCol) +{ + const gert::Shape x1_shape = context->GetInputShape(0)->GetStorageShape(); + const size_t gammaIndex = 2; + const gert::Shape gamma_shape = context->GetInputShape(gammaIndex)->GetStorageShape(); + numCol = gamma_shape.GetShapeSize(); + + const size_t x1DimNum = x1_shape.GetDimNum(); + size_t gammaDimNum = gamma_shape.GetDimNum(); + if (normKey == PRE_RMS_NORM || normKey == POST_RMS_NORM) { + gammaDimNum = gamma_shape.GetDimNum() - 1; + } + numRow = 1U; + for (size_t i = 0; i < x1DimNum - gammaDimNum; ++i) { + numRow *= x1_shape.GetDim(i); + } +} + +static ge::graphStatus GetEpsilonParameter(gert::TilingContext* context, float& epsilon) +{ + auto attrs = context->GetAttrs(); + OP_CHECK_NULL_WITH_CONTEXT(context, attrs); + epsilon = *attrs->GetFloat(0); + OP_CHECK_IF(epsilon < 0, + OP_LOGE_FOR_INVALID_VALUE(context->GetNodeName(), "epsilon", std::to_string(epsilon).c_str(), + "greater than or equal to zero"), return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +static ge::graphStatus GetAddGammaOffsetParameter(gert::TilingContext* context, uint32_t& addGammaOffset) +{ + auto attrs = context->GetAttrs(); + OP_CHECK_NULL_WITH_CONTEXT(context, attrs); + const bool* addGammaOffsetPtr = attrs->GetBool(1); + OP_CHECK_NULL_WITH_CONTEXT(context, addGammaOffsetPtr); + addGammaOffset = *addGammaOffsetPtr ? 1U : 0U; + return ge::GRAPH_SUCCESS; +} + +static void CalculateBlockParameters( + uint32_t numRow, uint32_t numCore, uint32_t& blockFactor, uint32_t& latsBlockFactor, uint32_t& useCoreNum) +{ + blockFactor = 1U; + uint32_t tileNum = Ops::Base::CeilDiv(numRow, numCore * blockFactor); + blockFactor *= tileNum; + useCoreNum = Ops::Base::CeilDiv(numRow, blockFactor); + latsBlockFactor = numRow - blockFactor * (useCoreNum - 1); +} + +static ge::DataType SetDataTypeParameters(gert::TilingContext* context, uint32_t& dtype_key, uint32_t& data_per_block) +{ + auto data_type = context->GetInputDesc(0)->GetDataType(); + dtype_key = DTYPE_KEY_FP16; + SetByDtype(data_type, dtype_key, data_per_block); + return data_type; +} + +static void DetermineModeParameters( + GammaAddRMSNormTilingData* tiling, + uint32_t numCol, uint32_t& ubFactor, uint32_t& rowFactor, uint32_t blockFactor, + uint32_t latsBlockFactor, ge::DataType dataType, uint32_t dtypKey, uint64_t ubSize, + uint32_t dataPerBlock, uint32_t numColAlign, uint32_t& modeKey, uint32_t isPerformance) +{ + if (numCol > ubFactor) { + modeKey = MODE_SPLIT_D; + ubFactor = (dataType == ge::DT_FLOAT) ? UB_FACTOR_B32_CUTD : UB_FACTOR_B16_CUTD; + uint32_t colTileNum = Ops::Base::CeilDiv(numCol, ubFactor); + ubFactor = Ops::Base::CeilDiv(numCol, colTileNum * dataPerBlock) * dataPerBlock; + } else if (blockFactor == 1) { + modeKey = MODE_SINGLE_N; + } else if (numColAlign <= SMALL_REDUCE_NUM) { + modeKey = MODE_MERGE_N; + uint64_t numColAlignWeight = (dtypKey == DTYPE_KEY_FP32) ? FP32_WEIGHT : OTHER_WEIGHT; + rowFactor = static_cast(ubSize) / + (numColAlign * static_cast(numColAlignWeight) + static_cast(DIV_FACTOR)); + ubFactor = rowFactor * numColAlign; + + uint32_t mulLoopFp32 = numColAlign / 64; + uint32_t mulTailFp32 = numColAlign - mulLoopFp32 * 64; + uint8_t dstRepStrideFp32 = numColAlign / 8; + + uint32_t mulLoopFp16 = numColAlign / 128; + uint32_t mulTailFp16 = numColAlign - mulLoopFp16 * 128; + uint8_t dstRepStrideFp16 = numColAlign / 16; + + tiling->set_is_performance(isPerformance); + tiling->set_mul_loop_fp32(mulLoopFp32); + tiling->set_mul_tail_fp32(mulTailFp32); + tiling->set_dst_rep_stride_fp32(dstRepStrideFp32); + tiling->set_mul_loop_fp16(mulLoopFp16); + tiling->set_mul_tail_fp16(mulTailFp16); + tiling->set_dst_rep_stride_fp16(dstRepStrideFp16); + } else if ((dataType == ge::DT_FLOAT16) && numCol == numColAlign) { + modeKey = MODE_MULTI_N; + rowFactor = (static_cast(ubSize) - static_cast(USE_SIZE) - + numColAlign * static_cast(NUM)) / + (numColAlign * BLOCK_ALIGN_NUM + static_cast(FLOAT_PER_REPEAT)); + ubFactor = rowFactor * numColAlign; + if (rowFactor == 0U) { + modeKey = MODE_NORMAL; + rowFactor = FLOAT_PER_REPEAT; + ubFactor = UB_FACTOR_B16; + } + } + uint32_t rowLoop = Ops::Base::CeilDiv(blockFactor, rowFactor); + uint32_t lastBlockRowLoop = Ops::Base::CeilDiv(latsBlockFactor, rowFactor); + uint32_t rowTail = blockFactor - (rowLoop - 1) * rowFactor; + uint32_t lastBlockRowTail = latsBlockFactor - (lastBlockRowLoop - 1) * rowFactor; + tiling->set_row_loop(rowLoop); + tiling->set_last_block_row_loop(lastBlockRowLoop); + tiling->set_row_tail(rowTail); + tiling->set_last_block_row_tail(lastBlockRowTail); +} + +static void SetTilingParameters( + GammaAddRMSNormTilingData* tiling, uint32_t num_row, uint32_t num_col, uint32_t numColAlign, + uint32_t block_factor, uint32_t latsBlockFactor, uint32_t row_factor, + uint32_t ub_factor, float epsilon, uint32_t addGammaOffset) +{ + const float avg_factor = (num_col == 0) ? 0 : 1.0f / num_col; + tiling->set_num_row(num_row); + tiling->set_num_col(num_col); + tiling->set_num_col_align(numColAlign); + tiling->set_block_factor(block_factor); + tiling->set_last_block_factor(latsBlockFactor); + tiling->set_row_factor(row_factor); + tiling->set_ub_factor(ub_factor); + tiling->set_epsilon(epsilon); + tiling->set_avg_factor(avg_factor); + tiling->set_add_gamma_offset(addGammaOffset); +} + +static void SaveTilingData( + gert::TilingContext* context, GammaAddRMSNormTilingData* tiling, uint32_t dtype_key, uint32_t mode_key, + uint32_t normKey) +{ + const uint32_t tiling_key = (dtype_key * 10 + mode_key) + normKey; + context->SetTilingKey(tiling_key); + tiling->SaveToBuffer(context->GetRawTilingData()->GetData(), context->GetRawTilingData()->GetCapacity()); + context->GetRawTilingData()->SetDataSize(tiling->GetDataSize()); +} + +static void SetWorkspaceSize(gert::TilingContext* context) +{ + size_t sysWorkspaceSize = 16 * 1024 * 1024; + constexpr size_t usrSize = 256; + size_t* currentWorkspace = context->GetWorkspaceSizes(1); + currentWorkspace[0] = usrSize + sysWorkspaceSize; +} + +static void LogTilingResults( + gert::TilingContext* context, GammaAddRMSNormTilingData* tiling, uint32_t mode_key, uint32_t dtype_key, + uint32_t use_core_num, float epsilon, uint32_t normKey) +{ + OP_LOGI(context, "Tiling Key: %u", (dtype_key * TEN + mode_key) + normKey); + OP_LOGI(context, "Block Dim: %u", use_core_num); + OP_LOGI(context, "usr Workspace: 256"); + OP_LOGI( + context, + "num_row: %d, num_col: %d, block_factor: %d, row_factor: %d, ub_factor: %d, epsilon: %f, avg_factor: %f", + tiling->get_num_row(), tiling->get_num_col(), tiling->get_block_factor(), tiling->get_row_factor(), + tiling->get_ub_factor(), epsilon, tiling->get_avg_factor()); +} + +static ge::graphStatus Tiling4GammaAddRmsNorm(gert::TilingContext* context) +{ + OP_LOGI("Tiling4GammaAddRmsNorm", "Enter Tiling4GammaAddRmsNorm"); + uint32_t normKey = RMS_NORM_KEY; + OP_CHECK_IF(!CheckNullptr(context, normKey), OP_LOGE(context, "Input shape invalid (nullptr)."), + return ge::GRAPH_FAILED); + OP_CHECK_IF(!CheckInputOutputShape(context, normKey), OP_LOGE(context, "Input shape invalid."), + return ge::GRAPH_FAILED); + + GammaAddRMSNormTilingData tiling; + uint32_t num_core; + uint64_t ub_size; + platform_ascendc::SocVersion socVersion; + + GetCompileParameters(context, num_core, ub_size, socVersion); + if (Ops::Xllm::OpTiling::IsRegbaseSocVersion(context)) { + return optiling::gammaAddRmsNormRegbase::TilingGammaAddRmsNormRegbase(context); + } + + uint32_t num_row; + uint32_t num_col; + CalculateRowAndColParameters(context, normKey, num_row, num_col); + + float epsilon = 0; + if (GetEpsilonParameter(context, epsilon) != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + uint32_t addGammaOffset = 0; + if (GetAddGammaOffsetParameter(context, addGammaOffset) != ge::GRAPH_SUCCESS) { + return ge::GRAPH_FAILED; + } + + uint32_t block_factor; + uint32_t latsBlockFactor; + uint32_t use_core_num; + CalculateBlockParameters(num_row, num_core, block_factor, latsBlockFactor, use_core_num); + context->SetBlockDim(use_core_num); + + uint32_t dtype_key; + uint32_t data_per_block; + ge::DataType data_type = SetDataTypeParameters(context, dtype_key, data_per_block); + + uint32_t mode_key = MODE_NORMAL; + uint32_t row_factor = 64; + uint32_t ub_factor = (dtype_key == DTYPE_KEY_FP32) ? UB_FACTOR_B32 : UB_FACTOR_B16; + uint32_t numColAlign = Ops::Base::CeilDiv(num_col, data_per_block) * data_per_block; + const gert::Shape x1_shape = context->GetInputShape(0)->GetStorageShape(); + const gert::Shape gamma_shape = context->GetInputShape(2)->GetStorageShape(); + uint8_t isPerformance = getPerformanceFlag(num_col, x1_shape, gamma_shape, dtype_key, socVersion); + DetermineModeParameters( + &tiling, + num_col, ub_factor, row_factor, block_factor, latsBlockFactor, + data_type, dtype_key, ub_size, data_per_block, + numColAlign, mode_key, isPerformance); + + SetTilingParameters(&tiling, num_row, num_col, numColAlign, block_factor, latsBlockFactor, row_factor, ub_factor, + epsilon, addGammaOffset); + SaveTilingData(context, &tiling, dtype_key, mode_key, normKey); + + SetWorkspaceSize(context); + + LogTilingResults(context, &tiling, mode_key, dtype_key, use_core_num, epsilon, normKey); + return ge::GRAPH_SUCCESS; +} + +static ge::graphStatus TilingPrepare4GammaAddRmsNorm(gert::TilingParseContext* context) +{ + OP_LOGI(context, "TilingPrepare4GammaAddRmsNorm running."); + auto compileInfo = context->GetCompiledInfo(); + OP_CHECK_NULL_WITH_CONTEXT(context, compileInfo); + auto platformInfo = context->GetPlatformInfo(); + OP_CHECK_NULL_WITH_CONTEXT(context, platformInfo); + auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo); + + compileInfo->socVersion = ascendcPlatform.GetSocVersion(); + compileInfo->totalCoreNum = ascendcPlatform.GetCoreNumAiv(); + ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, compileInfo->totalUbSize); + + return ge::GRAPH_SUCCESS; +} + +IMPL_OP_OPTILING(GammaAddRmsNorm).Tiling(Tiling4GammaAddRmsNorm).TilingParse(TilingPrepare4GammaAddRmsNorm); + +} // namespace optiling diff --git a/xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_tiling.h b/xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_tiling.h new file mode 100644 index 0000000..38bc641 --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_tiling.h @@ -0,0 +1,89 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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. + */ + +#ifndef OPS_BUILT_IN_OP_TILING_RUNTIME_GAMMA_ADD_RMS_NORM_H_ +#define OPS_BUILT_IN_OP_TILING_RUNTIME_GAMMA_ADD_RMS_NORM_H_ +#include "register/tilingdata_base.h" +#include "log/log.h" +#include "register/op_impl_registry.h" +#include "tiling/platform/platform_ascendc.h" +#include "platform/platform_infos_def.h" +#include "gamma_add_rms_norm_error_log.h" + +namespace optiling { +BEGIN_TILING_DATA_DEF(GammaAddRMSNormTilingData) +TILING_DATA_FIELD_DEF(uint32_t, num_row); +TILING_DATA_FIELD_DEF(uint32_t, num_col); +TILING_DATA_FIELD_DEF(uint32_t, block_factor); +TILING_DATA_FIELD_DEF(uint32_t, row_factor); +TILING_DATA_FIELD_DEF(uint32_t, ub_factor); +TILING_DATA_FIELD_DEF(float, epsilon); +TILING_DATA_FIELD_DEF(float, avg_factor); +TILING_DATA_FIELD_DEF(uint32_t, num_col_align); +TILING_DATA_FIELD_DEF(uint32_t, last_block_factor); +TILING_DATA_FIELD_DEF(uint32_t, row_loop); +TILING_DATA_FIELD_DEF(uint32_t, last_block_row_loop); +TILING_DATA_FIELD_DEF(uint32_t, row_tail); +TILING_DATA_FIELD_DEF(uint32_t, last_block_row_tail); +TILING_DATA_FIELD_DEF(uint32_t, mul_loop_fp32); +TILING_DATA_FIELD_DEF(uint32_t, mul_tail_fp32); +TILING_DATA_FIELD_DEF(uint32_t, dst_rep_stride_fp32); +TILING_DATA_FIELD_DEF(uint32_t, mul_loop_fp16); +TILING_DATA_FIELD_DEF(uint32_t, mul_tail_fp16); +TILING_DATA_FIELD_DEF(uint32_t, dst_rep_stride_fp16); +TILING_DATA_FIELD_DEF(uint32_t, is_performance); +TILING_DATA_FIELD_DEF(uint32_t, add_gamma_offset); +END_TILING_DATA_DEF; + +BEGIN_TILING_DATA_DEF(GammaAddRMSNormRegbaseTilingData) +TILING_DATA_FIELD_DEF(uint32_t, numRow); +TILING_DATA_FIELD_DEF(uint32_t, numCol); +TILING_DATA_FIELD_DEF(uint32_t, numColAlign); +TILING_DATA_FIELD_DEF(uint32_t, blockFactor); +TILING_DATA_FIELD_DEF(uint32_t, rowFactor); +TILING_DATA_FIELD_DEF(uint32_t, ubFactor); +TILING_DATA_FIELD_DEF(float, epsilon); +TILING_DATA_FIELD_DEF(float, avgFactor); +TILING_DATA_FIELD_DEF(uint32_t, ubLoop); +TILING_DATA_FIELD_DEF(uint32_t, colBuferLength); +TILING_DATA_FIELD_DEF(uint32_t, multiNNum); +TILING_DATA_FIELD_DEF(uint32_t, isNddma); +TILING_DATA_FIELD_DEF(uint32_t, addGammaOffset); +END_TILING_DATA_DEF; + +BEGIN_TILING_DATA_DEF(GammaAddRMSNormRegbaseRFullLoadTilingData) +TILING_DATA_FIELD_DEF(uint64_t, numRow); +TILING_DATA_FIELD_DEF(uint64_t, numCol); +TILING_DATA_FIELD_DEF(uint64_t, numColAlign); +TILING_DATA_FIELD_DEF(uint64_t, blockFactor); +TILING_DATA_FIELD_DEF(uint64_t, rowFactor); +TILING_DATA_FIELD_DEF(uint64_t, binAddQuotient); +TILING_DATA_FIELD_DEF(float, epsilon); +TILING_DATA_FIELD_DEF(float, avgFactor); +TILING_DATA_FIELD_DEF(uint32_t, addGammaOffset); +END_TILING_DATA_DEF; + +struct GammaAddRmsNormCompileInfo { + uint32_t totalCoreNum = 0; + uint64_t totalUbSize = 0; + platform_ascendc::SocVersion socVersion = platform_ascendc::SocVersion::ASCEND910B; +}; + +namespace gammaAddRmsNormRegbase { + ge::graphStatus TilingGammaAddRmsNormRegbase(gert::TilingContext* context); +} + +REGISTER_TILING_DATA_CLASS(GammaAddRmsNorm, GammaAddRMSNormTilingData) +REGISTER_TILING_DATA_CLASS(GammaAddRmsNorm_1000, GammaAddRMSNormRegbaseRFullLoadTilingData) +REGISTER_TILING_DATA_CLASS(GammaAddRmsNorm_2000, GammaAddRMSNormRegbaseTilingData) +} // namespace optiling + +#endif // OPS_BUILT_IN_OP_TILING_RUNTIME_GAMMA_ADD_RMS_NORM_H_ diff --git a/xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_tiling_arch35.cpp b/xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_tiling_arch35.cpp new file mode 100644 index 0000000..ce71867 --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_host/gamma_add_rms_norm_tiling_arch35.cpp @@ -0,0 +1,256 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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 gamma_add_rms_norm_tiling_arch35.cpp + * \brief + */ +#include +#include "register/op_impl_registry.h" +#include "gamma_add_rms_norm_tiling.h" +#include "op_common/op_host/util/math_util.h" +#include "op_common/op_host/util/platform_util.h" + +namespace optiling { +namespace gammaAddRmsNormRegbase { +constexpr uint32_t ULONG_BIT_LEN = 64; +constexpr uint32_t DTYPE_KEY_FP16 = 1; +constexpr uint32_t DTYPE_KEY_FP32 = 2; +constexpr uint32_t DTYPE_KEY_BF16 = 3; +constexpr uint32_t FLOAT_BLOCK_ALIGN_NUM = 8; +constexpr uint32_t FLOAT_PER_REAPEAT = 64; +constexpr uint32_t BYTE_SIZE_2_BLOCK_ALIGN_NUM = 16; +constexpr uint32_t X_INDEX = 0; +constexpr uint32_t GAMMA_INDEX = 2; +constexpr uint32_t FLOAT_BYTE_SIZE = 4; +constexpr uint32_t UB_USED = 1024; +constexpr uint32_t UB_RESERVE_FOR_RSTDALIGN = 1024; +constexpr uint32_t MODE_NORMAL = 1000; +constexpr uint32_t MODE_SPLIT_D = 2000; +constexpr uint32_t QUE_NUM = 5; +constexpr uint32_t QUE_MODE_NORMAL_NUM = 4; +constexpr uint64_t ALING_FACTOR_256 = 256; +constexpr uint64_t ALING_FACTOR_512 = 512; +constexpr uint32_t RETAINED_SIZE = 5120; // 256 * 5 * 4; +constexpr uint32_t DOUBLE_BUFFER_NUM = 2; +constexpr uint32_t MULTI_FACTOR_2 = 2; +constexpr uint32_t NUM_2 = 2; +constexpr uint32_t NDDMA_BETTER_STAGE = 512; + +const std::map dTypeByteMap = { + {ge::DT_FLOAT16, 2}, + {ge::DT_FLOAT, 4}, + {ge::DT_BF16, 2}, +}; + +template +auto CeilDiv(T x, T y) -> T +{ + return y == 0 ? x : (x + y - 1) / y; +} + +void SetByDtype(ge::DataType dataType, uint32_t& dtypeKey, uint32_t& dataPerBlock) +{ + switch (dataType) { + case ge::DT_FLOAT16: + dtypeKey = DTYPE_KEY_FP16; + dataPerBlock = BYTE_SIZE_2_BLOCK_ALIGN_NUM; + break; + case ge::DT_BF16: + dtypeKey = DTYPE_KEY_BF16; + dataPerBlock = BYTE_SIZE_2_BLOCK_ALIGN_NUM; + break; + default: + dtypeKey = DTYPE_KEY_FP32; + dataPerBlock = FLOAT_BLOCK_ALIGN_NUM; + break; + } +} + +uint32_t ComputeTotalBufSize(uint32_t bufferNum, ge::DataType dtype, uint32_t dtypeSize, uint32_t length, bool split) +{ + // queBufSize: UB space required for data movement. + uint32_t queBufSize = bufferNum * length * dtypeSize * QUE_NUM + FLOAT_PER_REAPEAT * bufferNum * FLOAT_BYTE_SIZE; + uint32_t tmpBufSzie = 0; // tmpBufSzie: temporary UB space required for computation. + if (split) { + // Split-D case. + tmpBufSzie = (dtype == ge::DT_FLOAT) ? 0 : length * FLOAT_BYTE_SIZE * NUM_2; + } else { + // Normal case: float16 and bfloat16 require an additional buffer for conversion to FP32. + tmpBufSzie = length * FLOAT_BYTE_SIZE; + } + return queBufSize + tmpBufSzie + RETAINED_SIZE; +} + +ge::graphStatus TilingGammaAddRmsNormRegbase(gert::TilingContext* context) +{ + OP_LOGD(context, " TilingGammaAddRmsNormRegbase"); + auto ptrCompileInfo = reinterpret_cast(context->GetCompileInfo()); + uint32_t numCore; + uint64_t ubSize; + if (nullptr == ptrCompileInfo) { + auto ascendcPlatform = platform_ascendc::PlatformAscendC(context->GetPlatformInfo()); + numCore = ascendcPlatform.GetCoreNumAiv(); + ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize); + } else { + numCore = ptrCompileInfo->totalCoreNum; + ubSize = ptrCompileInfo->totalUbSize; + } + const gert::Shape xShape = context->GetInputShape(X_INDEX)->GetStorageShape(); + + const gert::Shape gammaShape = context->GetInputShape(GAMMA_INDEX)->GetStorageShape(); + std::string opType(context->GetNodeType()); + auto attrs = context->GetAttrs(); + OP_CHECK_NULL_WITH_CONTEXT(context, attrs); + const float* epsilon = attrs->GetFloat(0); + OP_CHECK_NULL_WITH_CONTEXT(context, epsilon); + const bool* addGammaOffsetPtr = attrs->GetBool(1); + OP_CHECK_NULL_WITH_CONTEXT(context, addGammaOffsetPtr); + const uint32_t addGammaOffset = *addGammaOffsetPtr ? 1U : 0U; + OP_CHECK_IF( + *epsilon < 0, + OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(context->GetNodeName(), "epsilon", std::to_string(*epsilon).c_str(), + "epsilon should not be less than zero"), + return ge::GRAPH_FAILED); + uint64_t numCol = gammaShape.GetShapeSize(); + float avgFactor = (numCol == 0U) ? 0.0f : 1.0f / static_cast(numCol); + size_t xDimNum = xShape.GetDimNum(); + size_t gammaDimNum = gammaShape.GetDimNum(); + uint64_t numRow = 1; + for (size_t i = 0; i < xDimNum - gammaDimNum; i++) { + numRow *= xShape.GetDim(i); + } + for (size_t i = 0; i < xDimNum; i++) { + OP_LOGD(context, " TilingGammaAddRmsNormRegbase x shape:%ld", xShape.GetDim(i)); + } + for (size_t i = 0; i < gammaDimNum; i++) { + OP_LOGD(context, " TilingGammaAddRmsNormRegbase gama shape:%ld", gammaShape.GetDim(i)); + } + auto dataType = context->GetInputDesc(0)->GetDataType(); + uint32_t dtypeKey = DTYPE_KEY_FP16; + size_t usrSize = 256; + size_t sysWorkspaceSize = 16UL * 1024UL * 1024UL; + size_t* currentWorkspace = context->GetWorkspaceSizes(1); + currentWorkspace[0] = usrSize + sysWorkspaceSize; + uint64_t numColAlign = 0; + uint64_t ubBlockSize = Ops::Base::GetUbBlockSize(context); + uint64_t ubfp32 = ubBlockSize / sizeof(float); + uint64_t vlfp32 = Ops::Base::GetVRegSize(context) / sizeof(float); + uint64_t binaryAddElemtMaxLen = vlfp32 * vlfp32 * NUM_2 * NUM_2; + uint64_t blockFactor; + uint64_t ubFactor; + uint64_t rowFactor = 0; + uint32_t ubLoop{0}; + uint64_t colBuferLength{0}; + uint64_t multiNNum{0}; + + ubSize = ubSize - UB_USED; + uint32_t dataPerBlock; + SetByDtype(dataType, dtypeKey, dataPerBlock); + + blockFactor = static_cast(1); + uint64_t tileNum = CeilDiv(numRow, static_cast(numCore)); + blockFactor *= tileNum; + uint32_t useCoreNum = CeilDiv(numRow, blockFactor); + context->SetBlockDim(useCoreNum); + + auto dtypeByteIterator = dTypeByteMap.find(dataType); + OP_CHECK_IF( + dtypeByteIterator == dTypeByteMap.end(), OP_LOGE(context, "Fail to get dtype factor."), + return ge::GRAPH_FAILED); + uint32_t curElementByte = dtypeByteIterator->second; + numColAlign = CeilDiv(numCol * curElementByte, ubBlockSize) * ubBlockSize / curElementByte; + + // Calculate the boundary for binary-tree accumulation. + uint64_t binAddQuotient = numColAlign == 0 ? 1 : (1L << (ULONG_BIT_LEN - 1 - __builtin_clzl(numColAlign))); + binAddQuotient = (binAddQuotient == numColAlign) ? binAddQuotient / NUM_2 : binAddQuotient; + uint64_t binAddBufferOneline = Ops::Base::CeilAlign((binAddQuotient + vlfp32 - 1) / vlfp32, ubfp32); + + // Number of rows that can be fully loaded into UB. + int64_t tmpSize = static_cast(ubSize) - UB_RESERVE_FOR_RSTDALIGN - + (numColAlign * curElementByte); + if (tmpSize > 0 && numColAlign <= binaryAddElemtMaxLen) { + rowFactor = tmpSize / (numColAlign * curElementByte * DOUBLE_BUFFER_NUM * QUE_MODE_NORMAL_NUM + + numColAlign * sizeof(float) + sizeof(float) * (DOUBLE_BUFFER_NUM + 1) + + binAddBufferOneline * sizeof(float)); + } + if (rowFactor >= 1) { + // The reduction dimension can be fully loaded into UB. + rowFactor = std::min(rowFactor, blockFactor); // Actual number of rows to load. + GammaAddRMSNormRegbaseRFullLoadTilingData tiling; + tiling.set_numRow(numRow); + tiling.set_numCol(numCol); + tiling.set_numColAlign(numColAlign); + tiling.set_blockFactor(blockFactor); + tiling.set_rowFactor(rowFactor); + tiling.set_binAddQuotient(binAddQuotient); + tiling.set_epsilon(*epsilon); + tiling.set_avgFactor(avgFactor); + tiling.set_addGammaOffset(addGammaOffset); + OP_LOGI( + context, + "TilingData numCore: %u, ubSize: %lu, numRow: %u, numCol: %u, numColAlign: %u, " + "blockFactor: %u, rowFactor: %u, binAddQuotient: %u, " + "epsilon: %f, avgFactor: %f", + numCore, ubSize, tiling.get_numRow(), tiling.get_numCol(), tiling.get_numColAlign(), + tiling.get_blockFactor(), tiling.get_rowFactor(), tiling.get_binAddQuotient(), + tiling.get_epsilon(), tiling.get_avgFactor()); + + tiling.SaveToBuffer(context->GetRawTilingData()->GetData(), context->GetRawTilingData()->GetCapacity()); + context->SetTilingKey(MODE_NORMAL); + context->GetRawTilingData()->SetDataSize(tiling.GetDataSize()); + } else { + numColAlign = CeilDiv(numCol * curElementByte, ALING_FACTOR_512) * ALING_FACTOR_512 / curElementByte; + rowFactor = FLOAT_PER_REAPEAT; + ubFactor = 1U; + while (ComputeTotalBufSize(DOUBLE_BUFFER_NUM, dataType, curElementByte, ubFactor * MULTI_FACTOR_2, true) < + ubSize) { + ubFactor *= MULTI_FACTOR_2; + } + ubLoop = 1U; + while (ubLoop * MULTI_FACTOR_2 * ubFactor <= numCol) { + ubLoop *= MULTI_FACTOR_2; + } + colBuferLength = ubFactor; + uint32_t isNddma = numCol >= NDDMA_BETTER_STAGE ? 0U : 1U; + + GammaAddRMSNormRegbaseTilingData tiling; + tiling.set_numRow(numRow); + tiling.set_numCol(numCol); + tiling.set_numColAlign(numColAlign); + tiling.set_blockFactor(blockFactor); + tiling.set_rowFactor(rowFactor); + tiling.set_ubFactor(ubFactor); + tiling.set_epsilon(*epsilon); + tiling.set_avgFactor(avgFactor); + tiling.set_ubLoop(ubLoop); + tiling.set_colBuferLength(colBuferLength); + tiling.set_multiNNum(multiNNum); + tiling.set_isNddma(isNddma); + tiling.set_addGammaOffset(addGammaOffset); + OP_LOGI( + context, + "TilingData numCore: %u, ubSize: %lu, numRow: %u, numCol: %u, numColAlign: %u, colBuferLength: %u, " + "blockFactor: %u, rowFactor: %u, ubFactor: %u, " + "epsilon: %f, avgFactor: %f, ubLoop: %u, multiNNum: %u, isNddma: %u.", + numCore, ubSize, tiling.get_numRow(), tiling.get_numCol(), tiling.get_numColAlign(), + tiling.get_colBuferLength(), tiling.get_blockFactor(), tiling.get_rowFactor(), tiling.get_ubFactor(), + tiling.get_epsilon(), tiling.get_avgFactor(), tiling.get_ubLoop(), tiling.get_multiNNum(), + tiling.get_isNddma()); + + tiling.SaveToBuffer(context->GetRawTilingData()->GetData(), context->GetRawTilingData()->GetCapacity()); + context->SetTilingKey(MODE_SPLIT_D); + context->GetRawTilingData()->SetDataSize(tiling.GetDataSize()); + } + return ge::GRAPH_SUCCESS; +} +} // namespace gammaAddRmsNormRegbase +} // namespace optiling diff --git a/xllm_ops/gamma_add_rms_norm/op_host/op_api/aclnn_gamma_add_rms_norm.cpp b/xllm_ops/gamma_add_rms_norm/op_host/op_api/aclnn_gamma_add_rms_norm.cpp new file mode 100644 index 0000000..0bb5e7e --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_host/op_api/aclnn_gamma_add_rms_norm.cpp @@ -0,0 +1,252 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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. + */ +#include "aclnn/aclnn_base.h" +#include "op_api_def.h" + +#include "opdev/common_types.h" +#include "opdev/data_type_utils.h" +#include "opdev/op_executor.h" +#include "opdev/op_log.h" +#include "opdev/format_utils.h" +#include "opdev/tensor_view_utils.h" +#include "opdev/op_dfx.h" +#include "opdev/shape_utils.h" +#include "aclnn_kernels/common/op_error_check.h" +#include "aclnn_kernels/cast.h" +#include "aclnn_kernels/contiguous.h" +#include "aclnn_kernels/reshape.h" +#include "gamma_add_rms_norm.h" +#include "aclnn_gamma_add_rms_norm.h" + +using namespace op; +#ifdef __cplusplus +extern "C" { +#endif + +namespace GammaAddRmsNormACLNN { +constexpr int IDX_2 = 2; +constexpr int IDX_1 = 1; +constexpr int IDX_0 = 0; +constexpr int GAMMA_ADD_RMS_NORM_MODE = 0; +constexpr int PRE_RMS_NORM_MODE = 1; +constexpr int POST_RMS_NORM_MODE = 2; +const size_t MIN_SUPPORT_DIMS_NUMS = 1; +const size_t DIM_TWO = 2; + +struct GammaAddRmsNormInputTensor { + const aclTensor* x1; + const aclTensor* x2; + const aclTensor* gamma; +}; + +struct GammaAddRmsNormOutputTensor { + aclTensor* yOut; + aclTensor* rstdOut; + aclTensor* xOut; +}; + +static const std::initializer_list EXTEND_ATB_DTYPE_SUPPORT_LIST = { + op::DataType::DT_FLOAT16, op::DataType::DT_BF16}; + +static const std::initializer_list NORMAL_DTYPE_SUPPORT_LIST = { + op::DataType::DT_FLOAT, op::DataType::DT_FLOAT16, op::DataType::DT_BF16}; + +static bool CheckNotNull(GammaAddRmsNormInputTensor& inputTensor, GammaAddRmsNormOutputTensor& outputTensor, int64_t mode) +{ + OP_CHECK_NULL(inputTensor.x1, return false); + OP_CHECK_NULL(inputTensor.x2, return false); + OP_CHECK_NULL(inputTensor.gamma, return false); + OP_CHECK_NULL(outputTensor.yOut, return false); + if (mode == GammaAddRmsNormACLNN::GAMMA_ADD_RMS_NORM_MODE) { + OP_CHECK_NULL(outputTensor.rstdOut, return false); + OP_CHECK_NULL(outputTensor.xOut, return false); + } else if (mode == GammaAddRmsNormACLNN::PRE_RMS_NORM_MODE) { + OP_CHECK_NULL(outputTensor.xOut, return false); + } + return true; +} + +static bool CheckDtypeValid(GammaAddRmsNormInputTensor& inputTensor, GammaAddRmsNormOutputTensor& outputTensor, int64_t mode) +{ + std::initializer_list DTYPE_SUPPORT_LIST = NORMAL_DTYPE_SUPPORT_LIST; + if (mode == GammaAddRmsNormACLNN::PRE_RMS_NORM_MODE || mode == GammaAddRmsNormACLNN::POST_RMS_NORM_MODE) { + DTYPE_SUPPORT_LIST = EXTEND_ATB_DTYPE_SUPPORT_LIST; + } + + OP_CHECK_DTYPE_NOT_SUPPORT(inputTensor.x1, DTYPE_SUPPORT_LIST, return false); + OP_CHECK_DTYPE_NOT_SUPPORT(inputTensor.x2, DTYPE_SUPPORT_LIST, return false); + OP_CHECK_DTYPE_NOT_SUPPORT(inputTensor.gamma, DTYPE_SUPPORT_LIST, return false); + + OP_CHECK_DTYPE_NOT_SAME(inputTensor.x1, inputTensor.x2, return false); + OP_CHECK_DTYPE_NOT_SAME(inputTensor.x1, inputTensor.gamma, return false); + + OP_CHECK_DTYPE_NOT_SUPPORT(outputTensor.yOut, DTYPE_SUPPORT_LIST, return false); + OP_CHECK_DTYPE_NOT_SAME(outputTensor.yOut, inputTensor.x1, return false); + + if (mode == GammaAddRmsNormACLNN::GAMMA_ADD_RMS_NORM_MODE) { + OP_CHECK_DTYPE_NOT_SUPPORT(outputTensor.xOut, DTYPE_SUPPORT_LIST, return false); + OP_CHECK_DTYPE_NOT_SAME(outputTensor.xOut, outputTensor.yOut, return false); + OP_CHECK_DTYPE_NOT_MATCH(outputTensor.rstdOut, op::DataType::DT_FLOAT, return false); + } + if (mode == GammaAddRmsNormACLNN::PRE_RMS_NORM_MODE) { + OP_CHECK_DTYPE_NOT_SUPPORT(outputTensor.xOut, DTYPE_SUPPORT_LIST, return false); + OP_CHECK_DTYPE_NOT_SAME(outputTensor.xOut, outputTensor.yOut, return false); + } + return true; +} + +static bool CheckShapeDim(GammaAddRmsNormInputTensor& inputTensor, GammaAddRmsNormOutputTensor& outputTensor, int64_t mode) +{ + OP_CHECK_MAX_DIM(inputTensor.x1, MAX_SUPPORT_DIMS_NUMS, return false); + OP_CHECK_MAX_DIM(inputTensor.x2, MAX_SUPPORT_DIMS_NUMS, return false); + OP_CHECK_MAX_DIM(inputTensor.gamma, MAX_SUPPORT_DIMS_NUMS, return false); + + OP_CHECK_MIN_DIM(inputTensor.x1, MIN_SUPPORT_DIMS_NUMS, return false); + OP_CHECK_MIN_DIM(inputTensor.x2, MIN_SUPPORT_DIMS_NUMS, return false); + OP_CHECK_MIN_DIM(inputTensor.gamma, MIN_SUPPORT_DIMS_NUMS, return false); + + OP_CHECK_SHAPE_NOT_EQUAL(inputTensor.x1, inputTensor.x2, return false); + + OP_CHECK_MAX_DIM(outputTensor.yOut, MAX_SUPPORT_DIMS_NUMS, return false); + OP_CHECK_SHAPE_NOT_EQUAL(inputTensor.x1, outputTensor.yOut, return false); + + if (mode == GammaAddRmsNormACLNN::GAMMA_ADD_RMS_NORM_MODE) { + OP_CHECK_MAX_DIM(outputTensor.rstdOut, MAX_SUPPORT_DIMS_NUMS, return false); + + OP_CHECK_MAX_DIM(outputTensor.xOut, MAX_SUPPORT_DIMS_NUMS, return false); + OP_CHECK_SHAPE_NOT_EQUAL(inputTensor.x1, outputTensor.xOut, return false); + } + if (mode == GammaAddRmsNormACLNN::PRE_RMS_NORM_MODE) { + OP_CHECK_MAX_DIM(inputTensor.gamma, DIM_TWO, return false); + OP_CHECK_MAX_DIM(outputTensor.xOut, MAX_SUPPORT_DIMS_NUMS, return false); + OP_CHECK_SHAPE_NOT_EQUAL(inputTensor.x1, outputTensor.xOut, return false); + } + if(mode == GammaAddRmsNormACLNN::POST_RMS_NORM_MODE) { + OP_CHECK_MAX_DIM(inputTensor.gamma, DIM_TWO, return false); + } + return true; +} + +static aclnnStatus CheckParams(GammaAddRmsNormInputTensor& inputTensor, GammaAddRmsNormOutputTensor& outputTensor, int64_t& mode) +{ + if (outputTensor.xOut != nullptr && outputTensor.rstdOut == nullptr) { + mode = GammaAddRmsNormACLNN::PRE_RMS_NORM_MODE; // Two outputs: pre-RMSNorm mode. + } else if (outputTensor.xOut == nullptr && outputTensor.rstdOut == nullptr) { + mode = GammaAddRmsNormACLNN::POST_RMS_NORM_MODE; // One output: post-RMSNorm mode. + } + // 1. Check required input and output pointers. + CHECK_RET(CheckNotNull(inputTensor, outputTensor, mode), ACLNN_ERR_PARAM_NULLPTR); + + // 2. Validate input and output data types. + CHECK_RET(CheckDtypeValid(inputTensor, outputTensor, mode), ACLNN_ERR_PARAM_INVALID); + + // 3. Validate input and output shapes. + CHECK_RET(CheckShapeDim(inputTensor, outputTensor, mode), ACLNN_ERR_PARAM_INVALID); + + return ACLNN_SUCCESS; +} + +aclnnStatus ComputeGammaAddRmsNorm( + GammaAddRmsNormInputTensor& inputTensor, GammaAddRmsNormOutputTensor& outputTensor, double& epsilon, + bool addGammaOffset, int64_t& mode, aclOpExecutor* executor) +{ + aclTensor* yComputeOut = nullptr; + aclTensor* rstdComputeOut = nullptr; + aclTensor* xComputeOut = nullptr; + + auto GammaAddRmsNormOuts = + l0op::GammaAddRmsNorm(inputTensor.x1, inputTensor.x2, inputTensor.gamma, epsilon, addGammaOffset, mode, executor); + yComputeOut = std::get(GammaAddRmsNormOuts); + rstdComputeOut = std::get(GammaAddRmsNormOuts); + xComputeOut = std::get(GammaAddRmsNormOuts); + CHECK_RET(yComputeOut != nullptr, ACLNN_ERR_INNER_NULLPTR); + + // Copy yComputeOut to yOut. + auto viewCopyYResult = l0op::ViewCopy(yComputeOut, outputTensor.yOut, executor); + CHECK_RET(viewCopyYResult != nullptr, ACLNN_ERR_INNER_NULLPTR); + if (mode == GammaAddRmsNormACLNN::GAMMA_ADD_RMS_NORM_MODE) { + // Copy rstdComputeOut to rstdOut. + auto viewCopyXResult = l0op::ViewCopy(rstdComputeOut, outputTensor.rstdOut, executor); + CHECK_RET(viewCopyXResult != nullptr, ACLNN_ERR_INNER_NULLPTR); + // Copy xComputeOut to xOut. + auto viewCopyRstdResult = l0op::ViewCopy(xComputeOut, outputTensor.xOut, executor); + CHECK_RET(viewCopyRstdResult != nullptr, ACLNN_ERR_INNER_NULLPTR); + } + if (mode == GammaAddRmsNormACLNN::PRE_RMS_NORM_MODE) { + // Copy xComputeOut to xOut. + auto viewCopyRstdResult = l0op::ViewCopy(xComputeOut, outputTensor.xOut, executor); + CHECK_RET(viewCopyRstdResult != nullptr, ACLNN_ERR_INNER_NULLPTR); + } + return ACLNN_SUCCESS; +} +} // namespace GammaAddRmsNormACLNN + + +aclnnStatus aclnnGammaAddRmsNormGetWorkspaceSize( + const aclTensor* x1, const aclTensor* x2, const aclTensor* gamma, double epsilon, bool addGammaOffset, aclTensor* yOut, + aclTensor* rstdOut, aclTensor* xOut, uint64_t* workspaceSize, aclOpExecutor** executor) +{ + OP_LOGD("Enter aclnnGammaAddRmsNormGetWorkspaceSize."); + L2_DFX_PHASE_1(aclnnGammaAddRmsNorm, DFX_IN(x1, x2, gamma, epsilon, addGammaOffset), DFX_OUT(yOut, rstdOut, xOut)); + + // Create the operator executor. + auto uniqueExecutor = CREATE_EXECUTOR(); + CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR); + + // Validate parameters. + GammaAddRmsNormACLNN::GammaAddRmsNormInputTensor inputTensorOri = {x1, x2, gamma}; + GammaAddRmsNormACLNN::GammaAddRmsNormOutputTensor outputTensor = {yOut, rstdOut, xOut}; + + int64_t mode = GammaAddRmsNormACLNN::GAMMA_ADD_RMS_NORM_MODE; // 0: AddRmsNorm, 1: pre-RMSNorm, 2: post-RMSNorm. + auto ret = CheckParams(inputTensorOri, outputTensor, mode); + CHECK_RET(ret == ACLNN_SUCCESS, ret); + + // Support empty tensors. + bool anyEmptyTensor = x1->IsEmpty() || gamma->IsEmpty(); + if (anyEmptyTensor) { + OP_LOGW("Got empty tensor in aclnnGammaAddRmsNorm!"); + *workspaceSize = 0; + uniqueExecutor.ReleaseTo(executor); + return ACLNN_SUCCESS; + } + + // Convert inputs to contiguous tensors. Optional inputs do not require null checks here. + auto x1Cont = l0op::Contiguous(x1, uniqueExecutor.get()); + auto x2Cont = l0op::Contiguous(x2, uniqueExecutor.get()); + auto gammaCont = l0op::Contiguous(gamma, uniqueExecutor.get()); + + CHECK_RET(x1Cont != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(x2Cont != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(gammaCont != nullptr, ACLNN_ERR_INNER_NULLPTR); + + GammaAddRmsNormACLNN::GammaAddRmsNormInputTensor inputTensor = {x1Cont, x2Cont, gammaCont}; + + ret = GammaAddRmsNormACLNN::ComputeGammaAddRmsNorm(inputTensor, outputTensor, epsilon, addGammaOffset, mode, + uniqueExecutor.get()); + CHECK_RET(ret == ACLNN_SUCCESS, ret); + + // Obtain the workspace size required for computation. + *workspaceSize = uniqueExecutor->GetWorkspaceSize(); + uniqueExecutor.ReleaseTo(executor); + OP_LOGD("Finish aclnnGammaAddRmsNormGetWorkspaceSize."); + return ACLNN_SUCCESS; +} + +aclnnStatus aclnnGammaAddRmsNorm(void* workspace, uint64_t workspaceSize, aclOpExecutor* executor, aclrtStream stream) +{ + L2_DFX_PHASE_2(aclnnGammaAddRmsNorm); + // Execute the operator through the framework executor. + return CommonOpExecutorRun(workspace, workspaceSize, executor, stream); +} + +#ifdef __cplusplus +} +#endif diff --git a/xllm_ops/gamma_add_rms_norm/op_host/op_api/aclnn_gamma_add_rms_norm.h b/xllm_ops/gamma_add_rms_norm/op_host/op_api/aclnn_gamma_add_rms_norm.h new file mode 100644 index 0000000..3175ead --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_host/op_api/aclnn_gamma_add_rms_norm.h @@ -0,0 +1,75 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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. + */ +#ifndef OP_API_INC_LEVEL2_GAMMA_ADD_RMS_NORM_H_ +#define OP_API_INC_LEVEL2_GAMMA_ADD_RMS_NORM_H_ + +#include "aclnn/aclnn_base.h" +#include "aclnn_util.h" + +#ifdef __cplusplus +extern "C" { +#endif + +/** + * @brief First-stage aclnnGammaAddRmsNorm API that calculates the required workspace size. + * @domain aclnn_ops_infer + * + * Fuses Add and RMSNorm, returning the sum, normalized result, and reciprocal root mean square. + * Formula: + * x = x1 + x2 + * y = (x/RMS(x))*(gamma + (addGammaOffset ? 1 : 0)) + * rstd = 1/RMS(x) + * + * @param [in] x1: + * Input `x1` in the formula. Supports BFLOAT16, FLOAT, and FLOAT16 with at most eight dimensions. + * Non-contiguous tensors and the ND format are supported. + * @param [in] x2: + * Input `x2` in the formula. Supports BFLOAT16, FLOAT, and FLOAT16 with at most eight dimensions. + * Non-contiguous tensors and the ND format are supported. + * @param [in] gamma: + * Input `gamma` in the formula. Supports BFLOAT16, FLOAT, and FLOAT16 with at most eight dimensions. + * Non-contiguous tensors and the ND format are supported. + * @param [in] epsilon: Double-precision value used to prevent division by zero during normalization. + * @param [in] addGammaOffset: Whether to apply the Gemma-style `gamma + 1` inside the kernel. + * @param [in] yOut: + * Output `y` in the formula. Supports BFLOAT16, FLOAT, and FLOAT16 and must have the same shape as x1. + * Non-contiguous tensors and the ND format are supported. + * @param [in] rstdOut: + * Output `rstd` in the formula. Supports FLOAT with at most eight dimensions. + * Non-contiguous tensors and the ND format are supported. + * @param [in] xOut: + * Output `x` in the formula. Its data type and shape must match x1. + * Non-contiguous tensors and the ND format are supported. + * @param [out] workspaceSize: Workspace size that the caller must allocate on the NPU device. + * @param [out] executor: Operator executor containing the computation flow. + * @return aclnnStatus: Status code. + */ +ACLNN_API aclnnStatus aclnnGammaAddRmsNormGetWorkspaceSize( + const aclTensor* x1, const aclTensor* x2, const aclTensor* gamma, double epsilon, bool addGammaOffset, aclTensor* yOut, + aclTensor* rstdOut, aclTensor* xOut, uint64_t* workspaceSize, aclOpExecutor** executor); + +/** + * @brief Second-stage aclnnGammaAddRmsNorm API that executes the computation. + * + * @param [in] workspace: Start address of the workspace allocated on the NPU device. + * @param [in] workspaceSize: Workspace size returned by aclnnGammaAddRmsNormGetWorkspaceSize. + * @param [in] executor: Operator executor containing the computation flow. + * @param [in] stream: ACL stream used to execute the operator. + * @return aclnnStatus: Status code. + */ +ACLNN_API aclnnStatus +aclnnGammaAddRmsNorm(void* workspace, uint64_t workspaceSize, aclOpExecutor* executor, aclrtStream stream); + +#ifdef __cplusplus +} +#endif + +#endif // OP_API_INC_LEVEL2_GAMMA_ADD_RMS_NORM_H_ diff --git a/xllm_ops/gamma_add_rms_norm/op_host/op_api/gamma_add_rms_norm.cpp b/xllm_ops/gamma_add_rms_norm/op_host/op_api/gamma_add_rms_norm.cpp new file mode 100644 index 0000000..05ce6f6 --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_host/op_api/gamma_add_rms_norm.cpp @@ -0,0 +1,70 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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 gamma_add_rms_norm.cpp + * \brief + */ +#include "gamma_add_rms_norm.h" +#include "opdev/data_type_utils.h" +#include "opdev/format_utils.h" +#include "opdev/make_op_executor.h" +#include "opdev/op_def.h" +#include "opdev/op_dfx.h" +#include "opdev/op_executor.h" +#include "opdev/op_log.h" +#include "opdev/shape_utils.h" +#include "opdev/common_types.h" +#include "aclnn_kernels/cast.h" + +using namespace op; + +namespace l0op { +OP_TYPE_REGISTER(GammaAddRmsNorm); + +const std::array GammaAddRmsNorm( + const aclTensor* x1, const aclTensor* x2, const aclTensor* gamma, double epsilon, bool addGammaOffset, int64_t mode, + aclOpExecutor* executor) +{ + L0_DFX(GammaAddRmsNorm, x1, x2, gamma, epsilon, addGammaOffset); + Shape dummyShape({0}); + if(mode == GAMMA_ADD_RMS_NORM_MODE) { + Shape rstdShape; + size_t x1DimNum = x1->GetViewShape().GetDimNum(); + size_t gammaDimNum = gamma->GetViewShape().GetDimNum(); + for (uint32_t i = 0; i < x1DimNum - gammaDimNum; i++) { + rstdShape.AppendDim(x1->GetViewShape().GetDim(i)); + } + for (uint32_t i = 0; i < gammaDimNum; i++) { + rstdShape.AppendDim(1); + } + dummyShape = rstdShape; + } + + auto yOut = executor->AllocTensor(x1->GetViewShape(), x1->GetDataType(), x1->GetViewFormat()); + auto rstdOut = executor->AllocTensor(dummyShape, DataType::DT_FLOAT, x1->GetViewFormat()); + auto xOut = executor->AllocTensor(x1->GetViewShape(), x1->GetDataType(), x1->GetViewFormat()); + if (mode == PRE_RMS_NORM_MODE) { + xOut = executor->AllocTensor(x1->GetViewShape(), x1->GetDataType(), x1->GetViewFormat()); + } else if (mode == POST_RMS_NORM_MODE) { + xOut = executor->AllocTensor(dummyShape, x1->GetDataType(), x1->GetViewFormat()); + } + + auto ret = ADD_TO_LAUNCHER_LIST_AICORE( + GammaAddRmsNorm, OP_INPUT(x1, x2, gamma), OP_OUTPUT(yOut, rstdOut, xOut), + OP_ATTR(static_cast(epsilon), addGammaOffset)); + if (ret != ACL_SUCCESS) { + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "GammaAddRmsNorm ADD_TO_LAUNCHER_LIST_AICORE failed."); + return {nullptr, nullptr, nullptr}; + } + return {yOut, rstdOut, xOut}; +} +} // namespace l0op diff --git a/xllm_ops/gamma_add_rms_norm/op_host/op_api/gamma_add_rms_norm.h b/xllm_ops/gamma_add_rms_norm/op_host/op_api/gamma_add_rms_norm.h new file mode 100644 index 0000000..8ac82de --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_host/op_api/gamma_add_rms_norm.h @@ -0,0 +1,32 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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 gamma_add_rms_norm.h + * \brief + */ + +#ifndef OP_API_INC_LEVEL0_GAMMA_ADD_RMS_NORM_H_ +#define OP_API_INC_LEVEL0_GAMMA_ADD_RMS_NORM_H_ + +#include "opdev/op_executor.h" + +namespace l0op { +constexpr size_t GAMMA_ADD_RMS_NORM_OUT_NUM = 3; +constexpr int GAMMA_ADD_RMS_NORM_MODE = 0; +constexpr int PRE_RMS_NORM_MODE = 1; +constexpr int POST_RMS_NORM_MODE = 2; +const std::array GammaAddRmsNorm( + const aclTensor* x1, const aclTensor* x2, const aclTensor* gamma, double epsilon, bool addGammaOffset, int64_t mode, + aclOpExecutor* executor); +} // namespace l0op + +#endif // OP_API_INC_LEVEL0_GAMMA_ADD_RMS_NORM_H_ diff --git a/xllm_ops/gamma_add_rms_norm/op_kernel/arch35/gamma_add_rms_norm_regbase.h b/xllm_ops/gamma_add_rms_norm/op_kernel/arch35/gamma_add_rms_norm_regbase.h new file mode 100644 index 0000000..0c4807b --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_kernel/arch35/gamma_add_rms_norm_regbase.h @@ -0,0 +1,361 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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 gamma_add_rms_norm_regbase.h + * \brief + */ +#ifndef GAMMA_ADD_RMS_NORM_REGBASE_H +#define GAMMA_ADD_RMS_NORM_REGBASE_H +#include "gamma_add_rms_norm_regbase_common.h" +#include "../gamma_add_rms_norm_base.h" +#include "../inc/platform.h" +#include "kernel_operator.h" +#include "../../norm_common/reduce_common_regbase.h" + +namespace GammaAddRmsNorm { +using namespace AscendC; +constexpr uint64_t ALIGN_32_FACTOR = 32; +constexpr int32_t CONST_FACTOR_2 = 2; +constexpr int32_t NDDMA_DIM = 5; +constexpr int32_t UNROLL_NUM = 2; + +constexpr int32_t NUM_ONE = 1; +constexpr int32_t NUM_TWO = 2; + +using RmsNorm::DataCopyCustom; +using RmsNorm::DataCopyImpl; + +using AscendC::MicroAPI::LoadDist; +using AscendC::MicroAPI::MaskReg; +using AscendC::MicroAPI::RegTensor; + +constexpr static uint32_t BLOCK_SIZE = platform::GetUbBlockSize(); +constexpr static uint32_t VL_FP32 = platform::GetVRegSize() / sizeof(float); +constexpr static uint32_t BLK_B32 = BLOCK_SIZE / sizeof(float); + +template +__aicore__ inline T Min(T a, T b) +{ + return a > b ? b : a; +} +template +class KernelGammaAddRmsNormRegBase { +public: + __aicore__ inline KernelGammaAddRmsNormRegBase(TPipe* pipe) + { + pPipe = pipe; + } + __aicore__ inline void Init( + GM_ADDR x1, GM_ADDR x2, GM_ADDR gamma, GM_ADDR y, GM_ADDR rstd, GM_ADDR x, + const GammaAddRMSNormRegbaseRFullLoadTilingData* tiling) + { + ASSERT(GetBlockNum() != 0 && "Block dim can not be zero!"); + numRow = tiling->numRow; + numCol = tiling->numCol; + blockFactor = tiling->blockFactor; + binAddQuotient = tiling->binAddQuotient; + rowFactor = tiling->rowFactor; + epsilon = tiling->epsilon; + numColAlign = tiling->numColAlign; + avgFactor = tiling->avgFactor; + addGammaOffset = tiling->addGammaOffset; + rowWork = (GetBlockIdx() < GetBlockNum() - 1) ? blockFactor : numRow - (GetBlockNum() - 1) * blockFactor; + uint64_t rstdUbSizeAlignSize = CeilAlign(rowFactor, static_cast(VL_FP32)) * sizeof(float); + uint16_t binaryAddQuotientLoop = (binAddQuotient + VL_FP32 - 1) / VL_FP32; + uint32_t binaryAddBufLen = + (binaryAddQuotientLoop + BLK_B32 - 1) / BLK_B32 * BLK_B32 * sizeof(float) * rowFactor; + + xGm1.SetGlobalBuffer((__gm__ T*)x1 + GetBlockIdx() * blockFactor * numCol, rowWork * numCol); + xGm2.SetGlobalBuffer((__gm__ T*)x2 + GetBlockIdx() * blockFactor * numCol, rowWork * numCol); + gammaGm.SetGlobalBuffer((__gm__ T*)gamma, numCol); + yGm.SetGlobalBuffer((__gm__ T*)y + GetBlockIdx() * blockFactor * numCol, rowWork * numCol); + rstdGm.SetGlobalBuffer((__gm__ float*)rstd + GetBlockIdx() * blockFactor, blockFactor); + xOutGm.SetGlobalBuffer((__gm__ T*)x + GetBlockIdx() * blockFactor * numCol, rowWork * numCol); + + pPipe->InitBuffer(inQueueX1, DOUBLE_BUFFER_NUM, numColAlign * sizeof(T) * rowFactor); + pPipe->InitBuffer(inQueueX2, DOUBLE_BUFFER_NUM, numColAlign * sizeof(T) * rowFactor); + pPipe->InitBuffer(inQueueGamma, BUFFER_NUM, numColAlign * sizeof(T)); + pPipe->InitBuffer(outQueueY, DOUBLE_BUFFER_NUM, numColAlign * sizeof(T) * rowFactor); + pPipe->InitBuffer(outQueueX, DOUBLE_BUFFER_NUM, numColAlign * sizeof(T) * rowFactor); + pPipe->InitBuffer(outQueueRstd, DOUBLE_BUFFER_NUM, rstdUbSizeAlignSize); + pPipe->InitBuffer(xReduceBuff, rstdUbSizeAlignSize); + pPipe->InitBuffer(xFp32Buff, numColAlign * sizeof(float) * rowFactor); + pPipe->InitBuffer(binaryAddBuf, binaryAddBufLen); + } + + __aicore__ inline void Process() + { + CopyInGamma(); + LocalTensor gammaLocal = inQueueGamma.DeQue(); + if (addGammaOffset != 0U) { + AddGammaOffset(gammaLocal, static_cast(numCol)); + } + uint32_t rowLoopCount = CeilDiv(rowWork, rowFactor); + for (uint32_t rowLoopIdx = 0; rowLoopIdx < rowLoopCount; rowLoopIdx++) { + uint64_t rowLoopOffset = rowLoopIdx * rowFactor * numCol; + uint32_t curRows = Min(rowWork - rowLoopIdx * rowFactor, rowFactor); + Compute(rowLoopIdx, gammaLocal, curRows, rowLoopOffset); + } + inQueueGamma.FreeTensor(gammaLocal); + } + +private: + __aicore__ inline void Compute( + uint32_t rowLoopIdx, LocalTensor gammaLocal, uint32_t curRows, uint64_t rowLoopOffset) + { + CopyInXMutiMoveAlign(rowLoopOffset, numColAlign, curRows); + LocalTensor xLocal1 = inQueueX1.DeQue(); + LocalTensor xLocal2 = inQueueX2.DeQue(); + LocalTensor xOutLocal = outQueueX.AllocTensor(); + LocalTensor xFp32Local = xFp32Buff.Get(); + + CalculateXAdd(xLocal1, xLocal2, xOutLocal, xFp32Local, curRows, numColAlign); + inQueueX1.FreeTensor(xLocal1); + inQueueX2.FreeTensor(xLocal2); + outQueueX.EnQue(xOutLocal); + CopyOutX(rowLoopOffset, curRows, numColAlign); + + LocalTensor rstdLocal = outQueueRstd.AllocTensor(); + LocalTensor xReduceLocal = xReduceBuff.Get(); + NormCommon::NormCommonRegbase::CalculateSquareReduceSum( + xFp32Local, xReduceLocal, binaryAddBuf, static_cast(curRows), numColAlign, + numCol, static_cast(binAddQuotient), static_cast(BLK_B32)); + NormCommon::ComputeRstdNewtonRaphson( + xReduceLocal, rstdLocal, curRows, epsilon, avgFactor, VL_FP32); + outQueueRstd.EnQue(rstdLocal); + + rstdLocal = outQueueRstd.DeQue(); + DataCopyExtParams rstdCopyParams{ + static_cast(1), + static_cast(curRows * sizeof(float)), + static_cast(0), + static_cast(0), + 0 + }; + DataCopyPad(rstdGm[rowLoopIdx * rowFactor], rstdLocal, rstdCopyParams); + + LocalTensor yLocal = outQueueY.AllocTensor(); + CalculateY(xFp32Local, gammaLocal, yLocal, rstdLocal, curRows, numColAlign, numCol); + outQueueRstd.FreeTensor(rstdLocal); + outQueueY.EnQue(yLocal); + CopyOutY(rowLoopOffset, curRows, numColAlign); + } + + __aicore__ inline void CalculateXAdd( + LocalTensor& xLocal1, LocalTensor& xLocal2, LocalTensor& xOutLocal, LocalTensor& xFp32Local, + uint32_t curRows, uint32_t numColAlign) + { + __local_mem__ T* x1InUb = (__local_mem__ T*)xLocal1.GetPhyAddr(); + __local_mem__ T* x2InUb = (__local_mem__ T*)xLocal2.GetPhyAddr(); + __local_mem__ T* xOutInUb = (__local_mem__ T*)xOutLocal.GetPhyAddr(); + __local_mem__ float* xFp32Tmp = (__local_mem__ float*)xFp32Local.GetPhyAddr(); + + uint32_t sreg = curRows * numColAlign; + uint16_t loopCount = (sreg + VL_FP32 - 1) / VL_FP32; + + __VEC_SCOPE__ + { + RegTensor x1; + RegTensor x2; + RegTensor xSum; + MaskReg pregLoop; + for (uint16_t i = 0; i < loopCount; ++i) { + uint32_t offset = i * VL_FP32; + pregLoop = UpdateMask(sreg); + LoadRegForDtype(x1InUb, x1, pregLoop, offset); + LoadRegForDtype(x2InUb, x2, pregLoop, offset); + Add(xSum, x1, x2, pregLoop); + StoreRegForDtype(xOutInUb, xSum, pregLoop, offset); + DataCopy(xFp32Tmp + offset, xSum, pregLoop); + } + } + } + + __aicore__ inline void AddGammaOffset(LocalTensor& gammaLocal, uint32_t elementNum) + { + if constexpr (is_same::value) { + LocalTensor gammaFp32 = xFp32Buff.Get(); + Cast(gammaFp32, gammaLocal, RoundMode::CAST_NONE, elementNum); + PipeBarrier(); + Adds(gammaFp32, gammaFp32, static_cast(1.0), elementNum); + PipeBarrier(); + Cast(gammaLocal, gammaFp32, RoundMode::CAST_RINT, elementNum); + } else { + Adds(gammaLocal, gammaLocal, static_cast(1.0), elementNum); + } + PipeBarrier(); + } + + __aicore__ inline void CalculateY( + LocalTensor& xFp32Local, LocalTensor& gammaLocal, LocalTensor& yLocal, + LocalTensor& rstdLocal, uint32_t curRows, uint32_t numColAlign, uint32_t reduceNum) + { + __local_mem__ float* xFp32Tmp = (__local_mem__ float*)xFp32Local.GetPhyAddr(); + __local_mem__ T* gammaInUb = (__local_mem__ T*)gammaLocal.GetPhyAddr(); + __local_mem__ T* yInUb = (__local_mem__ T*)yLocal.GetPhyAddr(); + __local_mem__ float* rstdInUb = (__local_mem__ float*)rstdLocal.GetPhyAddr(); + + uint16_t loopRows = static_cast(curRows); + uint16_t loopCols = static_cast((reduceNum + VL_FP32 - 1) / VL_FP32); + uint16_t loopRowsFold = loopRows / 2; + uint16_t loopRowsHasLast = loopRows % 2; + + __VEC_SCOPE__ { + RegTensor x1Reg; + RegTensor x2Reg; + RegTensor gammaReg; + RegTensor rstd1Reg; + RegTensor rstd2Reg; + RegTensor mul1Reg; + RegTensor mul1UnrollReg; + RegTensor mul2Reg; + RegTensor mul2UnrollReg; + + for (uint16_t i = 0; i < loopRowsFold; ++i) { + uint32_t sregCount = reduceNum; + DataCopy(rstd1Reg, rstdInUb + 2 * i); + DataCopy(rstd2Reg, rstdInUb + (2 * i + 1)); + for (uint16_t r = 0; r < loopCols; ++r) { + uint32_t offset1 = (2 * i) * numColAlign + r * VL_FP32; + uint32_t offset2 = (2 * i + 1) * numColAlign + r * VL_FP32; + MaskReg regCurLoop = UpdateMask(sregCount); + LoadRegForDtype(xFp32Tmp, x1Reg, regCurLoop, offset1); + LoadRegForDtype(xFp32Tmp, x2Reg, regCurLoop, offset2); + Mul(mul1Reg, x1Reg, rstd1Reg, regCurLoop); + Mul(mul1UnrollReg, x2Reg, rstd2Reg, regCurLoop); + LoadRegForDtype(gammaInUb, gammaReg, regCurLoop, r * VL_FP32); + Mul(mul2Reg, mul1Reg, gammaReg, regCurLoop); + Mul(mul2UnrollReg, mul1UnrollReg, gammaReg, regCurLoop); + StoreRegForDtype(yInUb, mul2Reg, regCurLoop, offset1); + StoreRegForDtype(yInUb, mul2UnrollReg, regCurLoop, offset2); + } + } + for (uint16_t i = 0; i < loopRowsHasLast; ++i) { + uint32_t sregCount = reduceNum; + DataCopy(rstd1Reg, rstdInUb + 2 * loopRowsFold); + for (uint16_t r = 0; r < loopCols; ++r) { + uint32_t offset = (2 * loopRowsFold) * numColAlign + r * VL_FP32; + MaskReg regCurLoop = UpdateMask(sregCount); + LoadRegForDtype(xFp32Tmp, x1Reg, regCurLoop, offset); + Mul(mul1Reg, x1Reg, rstd1Reg, regCurLoop); + LoadRegForDtype(gammaInUb, gammaReg, regCurLoop, r * VL_FP32); + Mul(mul2Reg, mul1Reg, gammaReg, regCurLoop); + StoreRegForDtype(yInUb, mul2Reg, regCurLoop, offset); + } + } + } + } + + __aicore__ inline void CopyInXMutiMoveAlign(uint64_t offset, uint32_t curCols, uint32_t curRows = 0) + { + LocalTensor xLocal1 = inQueueX1.AllocTensor(); + LocalTensor xLocal2 = inQueueX2.AllocTensor(); + DataCopyExtParams extParams{ + static_cast(curRows), // blockCount + static_cast(numCol * sizeof(T)), // blockLen + static_cast(0), // srcStride + static_cast((numColAlign - curCols) * sizeof(T) / ALIGN_32_FACTOR), // dstStride + 0 // rsv + }; + DataCopyPadExtParams padParams{ + false, // isPad + static_cast(0), // leftPadding + static_cast(0), // rightPadding + static_cast(0.0) // paddingValue + }; + DataCopyPad(xLocal1, xGm1[offset], extParams, padParams); + DataCopyPad(xLocal2, xGm2[offset], extParams, padParams); + inQueueX1.EnQue(xLocal1); + inQueueX2.EnQue(xLocal2); + } + + __aicore__ inline void CopyInGamma() + { + LocalTensor gammaLocal = inQueueGamma.AllocTensor(); + DataCopyExtParams copyParams{ + static_cast(1), // blockCount + static_cast(numCol * sizeof(T)), // blockLen + static_cast(0), // srcStride + static_cast(0), // dstStride + 0 // rsv + }; + DataCopyPadExtParams padParams{ + false, // isPad + static_cast(0), // leftPadding + static_cast(0), // rightPadding + static_cast(0.0) // paddingValue + }; + DataCopyPad(gammaLocal, gammaGm, copyParams, padParams); + inQueueGamma.EnQue(gammaLocal); + } + + __aicore__ inline void CopyOutY(uint64_t offset, uint32_t curRows, uint32_t colAlign) + { + LocalTensor yLocal = outQueueY.DeQue(); + uint32_t srcStride = (numColAlign - colAlign) * sizeof(T) / ALIGN_32_FACTOR; + DataCopyExtParams copyParams{ + static_cast(curRows), // blockCount + static_cast(numCol * sizeof(T)), // blockLen + static_cast(srcStride), // srcStride + static_cast(0), // dstStride + 0 // rsv + }; + DataCopyPad(yGm[offset], yLocal, copyParams); + outQueueY.FreeTensor(yLocal); + } + + __aicore__ inline void CopyOutX(uint64_t offset, uint32_t curRows, uint32_t colAlign) + { + LocalTensor xLocal = outQueueX.DeQue(); + uint32_t srcStride = (numColAlign - colAlign) * sizeof(T) / ALIGN_32_FACTOR; + DataCopyExtParams copyParams{ + static_cast(curRows), // blockCount + static_cast(numCol * sizeof(T)), // blockLen + static_cast(srcStride), // srcStride + static_cast(0), // dstStride + 0 // rsv + }; + DataCopyPad(xOutGm[offset], xLocal, copyParams); + outQueueX.FreeTensor(xLocal); + } + +private: + TPipe* pPipe = nullptr; + TQue inQueueX1; + TQue inQueueX2; + TQue inQueueGamma; + TQue outQueueY; + TQue outQueueRstd; + TQue outQueueX; + TBuf xReduceBuff; + TBuf xFp32Buff; + TBuf binaryAddBuf; + MultiCopyParams dmaParam_; + GlobalTensor xGm1; + GlobalTensor xGm2; + GlobalTensor gammaGm; + GlobalTensor yGm; + GlobalTensor rstdGm; + GlobalTensor xOutGm; + uint64_t numRow; + uint64_t numCol; + uint64_t numColAlign; + uint64_t blockFactor; + uint64_t rowFactor; + uint64_t binAddQuotient; + float epsilon; + float avgFactor; + uint32_t addGammaOffset{0}; + uint64_t rowWork{1}; +}; +} // namespace GammaAddRmsNorm +#endif // GAMMA_ADD_RMS_NORM_REGBASE_H diff --git a/xllm_ops/gamma_add_rms_norm/op_kernel/arch35/gamma_add_rms_norm_regbase_common.h b/xllm_ops/gamma_add_rms_norm/op_kernel/arch35/gamma_add_rms_norm_regbase_common.h new file mode 100644 index 0000000..37a2b5a --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_kernel/arch35/gamma_add_rms_norm_regbase_common.h @@ -0,0 +1,768 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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 gamma_add_rms_norm_regbase_common.h + * \brief GammaAddRmsNorm regbase common + */ +#ifndef GAMMA_ADD_RMS_NORM_REGBASE_COMMON_H +#define GAMMA_ADD_RMS_NORM_REGBASE_COMMON_H + +#include "kernel_operator.h" +#include "kernel_tiling/kernel_tiling.h" +#include "../inc/platform.h" +#include "../gamma_add_rms_norm_base.h" +#include "../../rms_norm/arch35/rms_norm_regbase_common.h" +#include "../../norm_common/reduce_common_regbase.h" +namespace GammaAddRmsNorm { +using namespace AscendC; +using namespace AscendC::MicroAPI; +using namespace NormCommon; +using namespace NormCommon::NormCommonRegbase; +using NormCommon::NormCommonRegbase::LoadRegForDtype; +using NormCommon::NormCommonRegbase::StoreRegForDtype; +using AscendC::MicroAPI::Add; +using AscendC::MicroAPI::CreateMask; +using AscendC::MicroAPI::LoadDist; +using AscendC::MicroAPI::LocalMemBar; +using AscendC::MicroAPI::MaskPattern; +using AscendC::MicroAPI::MaskReg; +using AscendC::MicroAPI::MemType; +using AscendC::MicroAPI::RegTensor; +using AscendC::MicroAPI::StoreDist; +using AscendC::MicroAPI::UpdateMask; +using NormCommon::V_LENGTH; +using RmsNorm::castTraitB162B32; +using RmsNorm::castTraitB322B16; +using RmsNorm::CeilDiv; +using RmsNorm::BUFFER_NUM; +using RmsNorm::DOUBLE_BUFFER_NUM; +using RmsNorm::is_same; +using RmsNorm::ONCE_VECTOR_SIZE; + +template +__aicore__ inline void LoadForHandleRemainV1( + __local_mem__ T* mainAddr, __local_mem__ T* tailAddr, uint16_t offset1, uint16_t offset2, RegTensor& mainA, + RegTensor& mainB, RegTensor& tailA, RegTensor& tailB, MaskReg& pregLoop, + __local_mem__ float* xFp32MainAddr, __local_mem__ float* xFp32TailAddr, __local_mem__ T* mainAddr2, + __local_mem__ T* tailAddr2) +{ + if constexpr (IsSameType::value) { + // x1 load and cast + RegTensor xFp16MainA, xFp16MainB, xFp16TailA, xFp16TailB; + DataCopy(xFp16MainA, mainAddr + offset1); + DataCopy(xFp16MainB, mainAddr + offset2); + DataCopy(xFp16TailA, tailAddr + offset1); + DataCopy(xFp16TailB, tailAddr + offset2); + Cast(mainA, xFp16MainA, pregLoop); + Cast(mainB, xFp16MainB, pregLoop); + Cast(tailA, xFp16TailA, pregLoop); + Cast(tailB, xFp16TailB, pregLoop); + // x2 load and cast + RegTensor xFp16MainA2, xFp16MainB2, xFp16TailA2, xFp16TailB2; + DataCopy(xFp16MainA2, mainAddr2 + offset1); + DataCopy(xFp16MainB2, mainAddr2 + offset2); + DataCopy(xFp16TailA2, tailAddr2 + offset1); + DataCopy(xFp16TailB2, tailAddr2 + offset2); + RegTensor mainA2, mainB2, tailA2, tailB2; + Cast(mainA2, xFp16MainA2, pregLoop); + Cast(mainB2, xFp16MainB2, pregLoop); + Cast(tailA2, xFp16TailA2, pregLoop); + Cast(tailB2, xFp16TailB2, pregLoop); + // add x1 + x2 + Add(mainA, mainA, mainA2, pregLoop); + Add(mainB, mainB, mainB2, pregLoop); + Add(tailA, tailA, tailA2, pregLoop); + Add(tailB, tailB, tailB2, pregLoop); + DataCopy(xFp32MainAddr + offset1, mainA, pregLoop); + DataCopy(xFp32MainAddr + offset2, mainB, pregLoop); + DataCopy(xFp32TailAddr + offset1, tailA, pregLoop); + DataCopy(xFp32TailAddr + offset2, tailB, pregLoop); + Mul(mainA, mainA, mainA, pregLoop); + Mul(mainB, mainB, mainB, pregLoop); + Mul(tailA, tailA, tailA, pregLoop); + Mul(tailB, tailB, tailB, pregLoop); + } else if constexpr (IsSameType::value) { + // x1 load and cast + RegTensor xBFp16MainA, xBFp16MainB, xBFp16TailA, xBFp16TailB; + DataCopy(xBFp16MainA, mainAddr + offset1); + DataCopy(xBFp16MainB, mainAddr + offset2); + DataCopy(xBFp16TailA, tailAddr + offset1); + DataCopy(xBFp16TailB, tailAddr + offset2); + Cast(mainA, xBFp16MainA, pregLoop); + Cast(mainB, xBFp16MainB, pregLoop); + Cast(tailA, xBFp16TailA, pregLoop); + Cast(tailB, xBFp16TailB, pregLoop); + // x2 load and cast + RegTensor xBFp16MainA2, xBFp16MainB2, xBFp16TailA2, xBFp16TailB2; + DataCopy(xBFp16MainA2, mainAddr2 + offset1); + DataCopy(xBFp16MainB2, mainAddr2 + offset2); + DataCopy(xBFp16TailA2, tailAddr2 + offset1); + DataCopy(xBFp16TailB2, tailAddr2 + offset2); + // x2 cast + RegTensor mainA2, mainB2, tailA2, tailB2; + Cast(mainA2, xBFp16MainA2, pregLoop); + Cast(mainB2, xBFp16MainB2, pregLoop); + Cast(tailA2, xBFp16TailA2, pregLoop); + Cast(tailB2, xBFp16TailB2, pregLoop); + // add x1 + x2 + Add(mainA, mainA, mainA2, pregLoop); + Add(mainB, mainB, mainB2, pregLoop); + Add(tailA, tailA, tailA2, pregLoop); + Add(tailB, tailB, tailB2, pregLoop); + DataCopy(xFp32MainAddr + offset1, mainA, pregLoop); + DataCopy(xFp32MainAddr + offset2, mainB, pregLoop); + DataCopy(xFp32TailAddr + offset1, tailA, pregLoop); + DataCopy(xFp32TailAddr + offset2, tailB, pregLoop); + Mul(mainA, mainA, mainA, pregLoop); + Mul(mainB, mainB, mainB, pregLoop); + Mul(tailA, tailA, tailA, pregLoop); + Mul(tailB, tailB, tailB, pregLoop); + } else { + DataCopy(mainA, mainAddr + offset1); + DataCopy(mainB, mainAddr + offset2); + DataCopy(tailA, tailAddr + offset1); + DataCopy(tailB, tailAddr + offset2); + // load x2 + RegTensor mainA2, mainB2, tailA2, tailB2; + DataCopy(mainA2, mainAddr2 + offset1); + DataCopy(mainB2, mainAddr2 + offset2); + DataCopy(tailA2, tailAddr2 + offset1); + DataCopy(tailB2, tailAddr2 + offset2); + // add x1 + x2 + Add(mainA, mainA, mainA2, pregLoop); + Add(mainB, mainB, mainB2, pregLoop); + Add(tailA, tailA, tailA2, pregLoop); + Add(tailB, tailB, tailB2, pregLoop); + DataCopy(xFp32MainAddr + offset1, mainA, pregLoop); + DataCopy(xFp32MainAddr + offset2, mainB, pregLoop); + DataCopy(xFp32TailAddr + offset1, tailA, pregLoop); + DataCopy(xFp32TailAddr + offset2, tailB, pregLoop); + // x * x + Mul(mainA, mainA, mainA, pregLoop); + Mul(mainB, mainB, mainB, pregLoop); + Mul(tailA, tailA, tailA, pregLoop); + Mul(tailB, tailB, tailB, pregLoop); + } +} + +template +__aicore__ inline void LoadForHandleMasterV1( + __local_mem__ T* masterAddr, uint16_t offset1, uint16_t offset2, RegTensor& mainA, RegTensor& mainB, + MaskReg& pregLoop, __local_mem__ float* xFp32MasterAddr, __local_mem__ T* masterAddr2) +{ + if constexpr (IsSameType::value) { + // x1/x2 load and cast + RegTensor xFp16MainA, xFp16MainB; + DataCopy(xFp16MainA, masterAddr + offset1); + DataCopy(xFp16MainB, masterAddr + offset2); + Cast(mainA, xFp16MainA, pregLoop); + Cast(mainB, xFp16MainB, pregLoop); + RegTensor xFp16MainA2, xFp16MainB2; + DataCopy(xFp16MainA2, masterAddr2 + offset1); + DataCopy(xFp16MainB2, masterAddr2 + offset2); + RegTensor mainA2, mainB2; + Cast(mainA2, xFp16MainA2, pregLoop); + Cast(mainB2, xFp16MainB2, pregLoop); + // add x1 + x2 + Add(mainA, mainA, mainA2, pregLoop); + Add(mainB, mainB, mainB2, pregLoop); + DataCopy(xFp32MasterAddr + offset1, mainA, pregLoop); + DataCopy(xFp32MasterAddr + offset2, mainB, pregLoop); + Mul(mainA, mainA, mainA, pregLoop); + Mul(mainB, mainB, mainB, pregLoop); + } else if constexpr (IsSameType::value) { + // x1/x2 load and cast + RegTensor xBFp16MainA, xBFp16MainB; + DataCopy(xBFp16MainA, masterAddr + offset1); + DataCopy(xBFp16MainB, masterAddr + offset2); + Cast(mainA, xBFp16MainA, pregLoop); + Cast(mainB, xBFp16MainB, pregLoop); + RegTensor xBFp16MainA2, xBFp16MainB2; + DataCopy(xBFp16MainA2, masterAddr2 + offset1); + DataCopy(xBFp16MainB2, masterAddr2 + offset2); + RegTensor mainA2, mainB2; + Cast(mainA2, xBFp16MainA2, pregLoop); + Cast(mainB2, xBFp16MainB2, pregLoop); + // add x1 + x2 + Add(mainA, mainA, mainA2, pregLoop); + Add(mainB, mainB, mainB2, pregLoop); + DataCopy(xFp32MasterAddr + offset1, mainA, pregLoop); + DataCopy(xFp32MasterAddr + offset2, mainB, pregLoop); + Mul(mainA, mainA, mainA, pregLoop); + Mul(mainB, mainB, mainB, pregLoop); + } else { + DataCopy(mainA, masterAddr + offset1); + DataCopy(mainB, masterAddr + offset2); + RegTensor mainA2, mainB2; + DataCopy(mainA2, masterAddr2 + offset1); + DataCopy(mainB2, masterAddr2 + offset2); + // add x1 + x2 + Add(mainA, mainA, mainA2, pregLoop); + Add(mainB, mainB, mainB2, pregLoop); + DataCopy(xFp32MasterAddr + offset1, mainA, pregLoop); + DataCopy(xFp32MasterAddr + offset2, mainB, pregLoop); + // x * x + Mul(mainA, mainA, mainA, pregLoop); + Mul(mainB, mainB, mainB, pregLoop); + } +} + +template +__aicore__ inline void LoadForHandleRemainV2( + __local_mem__ T* mainAddr, __local_mem__ T* tailAddr, uint16_t offset1, uint16_t offset2, RegTensor& mainA, + RegTensor& mainB, RegTensor& tailA, RegTensor& tailB, MaskReg& pregMask, + __local_mem__ T* mainAddr2, __local_mem__ T* tailAddr2) +{ + if constexpr (IsSameType::value) { + // x2 load and cast + RegTensor xFp16MainA, xFp16MainB, xFp16TailA, xFp16TailB; + DataCopy(xFp16MainA, mainAddr + offset1); + DataCopy(xFp16MainB, mainAddr + offset2); + DataCopy(xFp16TailA, tailAddr + offset1); + DataCopy(xFp16TailB, tailAddr + offset2); + Cast(mainA, xFp16MainA, pregMask); + Cast(mainB, xFp16MainB, pregMask); + Cast(tailA, xFp16TailA, pregMask); + Cast(tailB, xFp16TailB, pregMask); + RegTensor xFp16MainA2, xFp16MainB2, xFp16TailA2, xFp16TailB2; + DataCopy(xFp16MainA2, mainAddr2 + offset1); + DataCopy(xFp16MainB2, mainAddr2 + offset2); + DataCopy(xFp16TailA2, tailAddr2 + offset1); + DataCopy(xFp16TailB2, tailAddr2 + offset2); + RegTensor mainA2, mainB2, tailA2, tailB2; + Cast(mainA2, xFp16MainA2, pregMask); + Cast(mainB2, xFp16MainB2, pregMask); + Cast(tailA2, xFp16TailA2, pregMask); + Cast(tailB2, xFp16TailB2, pregMask); + // add x1 + x2 + Add(mainA, mainA, mainA2, pregMask); + Add(mainB, mainB, mainB2, pregMask); + Add(tailA, tailA, tailA2, pregMask); + Add(tailB, tailB, tailB2, pregMask); + Mul(mainA, mainA, mainA, pregMask); + Mul(mainB, mainB, mainB, pregMask); + Mul(tailA, tailA, tailA, pregMask); + Mul(tailB, tailB, tailB, pregMask); + } else if constexpr (IsSameType::value) { + // x2 load and cast + RegTensor xBFp16MainA, xBFp16MainB, xBFp16TailA, xBFp16TailB; + DataCopy(xBFp16MainA, mainAddr + offset1); + DataCopy(xBFp16MainB, mainAddr + offset2); + DataCopy(xBFp16TailA, tailAddr + offset1); + DataCopy(xBFp16TailB, tailAddr + offset2); + Cast(mainA, xBFp16MainA, pregMask); + Cast(mainB, xBFp16MainB, pregMask); + Cast(tailA, xBFp16TailA, pregMask); + Cast(tailB, xBFp16TailB, pregMask); + RegTensor xBFp16MainA2, xBFp16MainB2, xBFp16TailA2, xBFp16TailB2; + DataCopy(xBFp16MainA2, mainAddr2 + offset1); + DataCopy(xBFp16MainB2, mainAddr2 + offset2); + DataCopy(xBFp16TailA2, tailAddr2 + offset1); + DataCopy(xBFp16TailB2, tailAddr2 + offset2); + RegTensor mainA2, mainB2, tailA2, tailB2; + Cast(mainA2, xBFp16MainA2, pregMask); + Cast(mainB2, xBFp16MainB2, pregMask); + Cast(tailA2, xBFp16TailA2, pregMask); + Cast(tailB2, xBFp16TailB2, pregMask); + // add x1 + x2 + Add(mainA, mainA, mainA2, pregMask); + Add(mainB, mainB, mainB2, pregMask); + Add(tailA, tailA, tailA2, pregMask); + Add(tailB, tailB, tailB2, pregMask); + Mul(mainA, mainA, mainA, pregMask); + Mul(mainB, mainB, mainB, pregMask); + Mul(tailA, tailA, tailA, pregMask); + Mul(tailB, tailB, tailB, pregMask); + } else { + DataCopy(mainA, mainAddr + offset1); + DataCopy(mainB, mainAddr + offset2); + DataCopy(tailA, tailAddr + offset1); + DataCopy(tailB, tailAddr + offset2); + // load x2 + RegTensor mainA2, mainB2, tailA2, tailB2; + DataCopy(mainA2, mainAddr2 + offset1); + DataCopy(mainB2, mainAddr2 + offset2); + DataCopy(tailA2, tailAddr2 + offset1); + DataCopy(tailB2, tailAddr2 + offset2); + // add x1 + x2 + Add(mainA, mainA, mainA2, pregMask); + Add(mainB, mainB, mainB2, pregMask); + Add(tailA, tailA, tailA2, pregMask); + Add(tailB, tailB, tailB2, pregMask); + // x * x + Mul(mainA, mainA, mainA, pregMask); + Mul(mainB, mainB, mainB, pregMask); + Mul(tailA, tailA, tailA, pregMask); + Mul(tailB, tailB, tailB, pregMask); + } +} + +template +__aicore__ inline void LoadForHandleMasterV2( + __local_mem__ T* masterAddr, uint16_t offset1, uint16_t offset2, RegTensor& mainA, RegTensor& mainB, + MaskReg& pregMask, __local_mem__ T* masterAddr2) +{ + if constexpr (IsSameType::value) { + // x1/x2 load and cast + RegTensor xFp16MainA, xFp16MainB; + DataCopy(xFp16MainA, masterAddr + offset1); + DataCopy(xFp16MainB, masterAddr + offset2); + Cast(mainA, xFp16MainA, pregMask); + Cast(mainB, xFp16MainB, pregMask); + RegTensor xFp16MainA2, xFp16MainB2; + DataCopy(xFp16MainA2, masterAddr2 + offset1); + DataCopy(xFp16MainB2, masterAddr2 + offset2); + RegTensor mainA2, mainB2; + Cast(mainA2, xFp16MainA2, pregMask); + Cast(mainB2, xFp16MainB2, pregMask); + // add x1 + x2 + Add(mainA, mainA, mainA2, pregMask); + Add(mainB, mainB, mainB2, pregMask); + + Mul(mainA, mainA, mainA, pregMask); + Mul(mainB, mainB, mainB, pregMask); + } else if constexpr (IsSameType::value) { + // x1/x2 load and cast + RegTensor xBFp16MainA, xBFp16MainB; + DataCopy(xBFp16MainA, masterAddr + offset1); + DataCopy(xBFp16MainB, masterAddr + offset2); + Cast(mainA, xBFp16MainA, pregMask); + Cast(mainB, xBFp16MainB, pregMask); + RegTensor xBFp16MainA2, xBFp16MainB2; + DataCopy(xBFp16MainA2, masterAddr2 + offset1); + DataCopy(xBFp16MainB2, masterAddr2 + offset2); + RegTensor mainA2, mainB2; + Cast(mainA2, xBFp16MainA2, pregMask); + Cast(mainB2, xBFp16MainB2, pregMask); + // add x1 + x2 + Add(mainA, mainA, mainA2, pregMask); + Add(mainB, mainB, mainB2, pregMask); + Mul(mainA, mainA, mainA, pregMask); + Mul(mainB, mainB, mainB, pregMask); + } else { + DataCopy(mainA, masterAddr + offset1); + DataCopy(mainB, masterAddr + offset2); + // x2 load + RegTensor mainA2, mainB2; + DataCopy(mainA2, masterAddr2 + offset1); + DataCopy(mainB2, masterAddr2 + offset2); + // add x1 + x2 + Add(mainA, mainA, mainA2, pregMask); + Add(mainB, mainB, mainB2, pregMask); + Mul(mainA, mainA, mainA, pregMask); + Mul(mainB, mainB, mainB, pregMask); + } +} + +template +__aicore__ inline void ComputeFormerImplV2( + LocalTensor& dstLocal, LocalTensor& xLocal1, LocalTensor& xLocal2, LocalTensor& workLocal, + uint32_t offset, uint32_t count, uint32_t powerSplit) +{ + uint32_t remainTile = count - powerSplit; + uint32_t masterTile = powerSplit - remainTile; + uint32_t remainSreg = remainTile; + uint16_t remainRepeats = remainTile / (2 * V_LENGTH); + uint32_t masterSreg = masterTile; + uint16_t masterRepeats = masterTile / (2 * V_LENGTH); + + uint32_t mergeTile = powerSplit / (2 * V_LENGTH); + uint32_t mergeSreg = mergeTile; + uint16_t mergeRepeats = mergeTile / (2 * V_LENGTH); + + uint32_t meanTile = mergeRepeats == 0 ? mergeTile : mergeRepeats; + uint32_t meanSreg = meanTile; + + __local_mem__ T* mainAddr = (__ubuf__ T*)xLocal1.GetPhyAddr(); + __local_mem__ T* tailAddr = (__ubuf__ T*)xLocal1.GetPhyAddr() + int64_t(powerSplit); + __local_mem__ T* masterAddr = (__ubuf__ T*)xLocal1.GetPhyAddr() + int64_t(remainTile); + __local_mem__ T* mainAddr2 = (__ubuf__ T*)xLocal2.GetPhyAddr(); + __local_mem__ T* tailAddr2 = (__ubuf__ T*)xLocal2.GetPhyAddr() + int64_t(powerSplit); + __local_mem__ T* masterAddr2 = (__ubuf__ T*)xLocal2.GetPhyAddr() + int64_t(remainTile); + + __local_mem__ float* workAddr = (__ubuf__ float*)workLocal.GetPhyAddr(); + __local_mem__ float* dstAddr = (__ubuf__ float*)dstLocal.GetPhyAddr(); + + __VEC_SCOPE__ + { + RegTensor mainA, mainB, tailA, tailB, vMean, vDupReg, rstdReg; + MaskReg pregMerge = CreateMask(); + MaskReg pregMask; + + for (uint16_t i = 0; i < (uint16_t)remainRepeats; ++i) { + pregMask = UpdateMask(remainSreg); + LoadForHandleRemainV2( + mainAddr, tailAddr, (i * 2 + 0) * V_LENGTH, (i * 2 + 1) * V_LENGTH, mainA, mainB, tailA, tailB, + pregMask, mainAddr2, tailAddr2); + Add(mainA, mainA, tailA, pregMask); + Add(mainB, mainB, tailB, pregMask); + Add(mainA, mainA, mainB, pregMask); + ReduceSum(vMean, mainA, pregMask); + DataCopy(workAddr + i, vMean, pregMerge); + } + for (uint16_t i = 0; i < (uint16_t)masterRepeats; ++i) { + pregMask = UpdateMask(masterSreg); + LoadForHandleMasterV2( + masterAddr, (i * 2 + 0) * V_LENGTH, (i * 2 + 1) * V_LENGTH, mainA, mainB, pregMask, masterAddr2); + Add(mainA, mainA, mainB, pregMask); + ReduceSum(vMean, mainA, pregMask); + DataCopy(workAddr + remainRepeats + i, vMean, pregMerge); + } + LocalMemBar(); + for (uint16_t i = 0; i < (uint16_t)mergeRepeats; ++i) { + pregMask = UpdateMask(mergeSreg); + DataCopy(mainA, workAddr + (i * 2 + 0) * V_LENGTH); + DataCopy(mainB, workAddr + (i * 2 + 1) * V_LENGTH); + Add(mainA, mainA, mainB, pregMask); + ReduceSum(vMean, mainA, pregMask); + DataCopy(workAddr + i, vMean, pregMerge); + } + LocalMemBar(); + { + pregMask = UpdateMask(meanSreg); + DataCopy(mainA, workAddr + 0); + ReduceSum(vMean, mainA, pregMask); + DataCopy(dstAddr + offset, vMean, pregMerge); + } + } +} + +template +__aicore__ inline void ComputeFormerImplV1MultiN( + LocalTensor& xLocal1, LocalTensor& xLocal2, LocalTensor& xFp32, LocalTensor& workLocal, + LocalTensor& rstdLocal, float avgFactor, float epsilon, uint32_t offset, uint32_t count, uint32_t powerSplit, + uint32_t curRows) +{ + uint32_t remainTile = count - powerSplit; + uint16_t remainRepeats = remainTile / (2 * V_LENGTH); + + uint32_t masterTile = powerSplit - remainTile; + uint16_t masterRepeats = masterTile / (2 * V_LENGTH); + + uint32_t mergeTile = powerSplit / (2 * V_LENGTH); + uint16_t mergeRepeats = mergeTile / (2 * V_LENGTH); + + uint32_t meanTile = mergeRepeats == 0 ? mergeTile : mergeRepeats; + + __local_mem__ T* mainAddr = (__ubuf__ T*)xLocal1.GetPhyAddr(); + __local_mem__ T* tailAddr = (__ubuf__ T*)xLocal1.GetPhyAddr() + int64_t(powerSplit); + __local_mem__ T* masterAddr = (__ubuf__ T*)xLocal1.GetPhyAddr() + int64_t(remainTile); + __local_mem__ T* mainAddr2 = (__ubuf__ T*)xLocal2.GetPhyAddr(); + __local_mem__ T* tailAddr2 = (__ubuf__ T*)xLocal2.GetPhyAddr() + int64_t(powerSplit); + __local_mem__ T* masterAddr2 = (__ubuf__ T*)xLocal2.GetPhyAddr() + int64_t(remainTile); + __local_mem__ float *xFp32MainAddr, *xFp32TailAddr, *xFp32MasterAddr; + + xFp32MainAddr = (__ubuf__ float*)xFp32.GetPhyAddr(); + xFp32TailAddr = (__ubuf__ float*)xFp32.GetPhyAddr() + int64_t(powerSplit); + xFp32MasterAddr = (__ubuf__ float*)xFp32.GetPhyAddr() + int64_t(remainTile); + + __local_mem__ float* workAddr = (__ubuf__ float*)workLocal.GetPhyAddr(); + __local_mem__ float* rstdAddr = (__ubuf__ float*)rstdLocal.GetPhyAddr(); + + uint32_t curRowsAlign = CeilDiv((int32_t)curRows, 2); + int64_t unrollOffset = (curRows / 2) * count; + bool isWithTail = curRowsAlign - (curRows / 2); + uint32_t tailOffset = offset + curRows / 2; + + __local_mem__ T* mainAddr3 = (__ubuf__ T*)xLocal1.GetPhyAddr() + unrollOffset; + __local_mem__ T* tailAddr3 = (__ubuf__ T*)xLocal1.GetPhyAddr() + int64_t(powerSplit) + unrollOffset; + __local_mem__ T* masterAddr3 = (__ubuf__ T*)xLocal1.GetPhyAddr() + int64_t(remainTile) + unrollOffset; + __local_mem__ T* mainAddr4 = (__ubuf__ T*)xLocal2.GetPhyAddr() + unrollOffset; + __local_mem__ T* tailAddr4 = (__ubuf__ T*)xLocal2.GetPhyAddr() + int64_t(powerSplit) + unrollOffset; + __local_mem__ T* masterAddr4 = (__ubuf__ T*)xLocal2.GetPhyAddr() + int64_t(remainTile) + unrollOffset; + __local_mem__ float *xFp32MainAddr3, *xFp32TailAddr3, *xFp32MasterAddr3; + + xFp32MainAddr3 = (__ubuf__ float*)xFp32.GetPhyAddr() + unrollOffset; + xFp32TailAddr3 = (__ubuf__ float*)xFp32.GetPhyAddr() + int64_t(powerSplit) + unrollOffset; + xFp32MasterAddr3 = (__ubuf__ float*)xFp32.GetPhyAddr() + int64_t(remainTile) + unrollOffset; + + __local_mem__ float* workAddr1 = (__ubuf__ float*)workLocal.GetPhyAddr() + ONCE_VECTOR_SIZE; + __local_mem__ float* rstdAddr1 = (__ubuf__ float*)rstdLocal.GetPhyAddr() + curRows / 2; + + __VEC_SCOPE__ + { + for (uint16_t row = 0; row < static_cast(curRows / 2); row++) { + uint32_t remainSreg = remainTile; + uint32_t masterSreg = masterTile; + uint32_t meanSreg = meanTile; + uint32_t mergeSreg = mergeTile; + RegTensor mainA, mainB, tailA, tailB, vMean, vDupReg, rstdReg; + MaskReg pregMerge = CreateMask(); + MaskReg pregMask; + + uint32_t remainSreg1 = remainTile; + uint32_t masterSreg1 = masterTile; + uint32_t mergeSreg1 = mergeTile; + uint32_t meanSreg1 = meanTile; + RegTensor mainA3, mainB3, tailA3, tailB3, vMean3, vDupReg3, rstdReg3; + MaskReg pregMain1 = CreateMask(); + MaskReg pregMerge1 = CreateMask(); + MaskReg pregMask1; + + for (uint16_t i = 0; i < (uint16_t)remainRepeats; ++i) { + pregMask = UpdateMask(remainSreg); + LoadForHandleRemainV1( + mainAddr, tailAddr, (i * 2 + 0) * V_LENGTH, (i * 2 + 1) * V_LENGTH, mainA, mainB, tailA, tailB, + pregMask, xFp32MainAddr, xFp32TailAddr, mainAddr2, tailAddr2); + Add(mainA, mainA, tailA, pregMask); + Add(mainB, mainB, tailB, pregMask); + Add(mainA, mainA, mainB, pregMask); + ReduceSum(vMean, mainA, pregMask); + DataCopy(workAddr + i, vMean, pregMerge); + } + for (uint16_t i = 0; i < (uint16_t)masterRepeats; ++i) { + pregMask = UpdateMask(masterSreg); + LoadForHandleMasterV1( + masterAddr, (i * 2 + 0) * V_LENGTH, (i * 2 + 1) * V_LENGTH, mainA, mainB, pregMask, xFp32MasterAddr, + masterAddr2); + Add(mainA, mainA, mainB, pregMask); + ReduceSum(vMean, mainA, pregMask); + DataCopy(workAddr + remainRepeats + i, vMean, pregMerge); + } + // unroll part + for (uint16_t i = 0; i < (uint16_t)remainRepeats; ++i) { + pregMask1 = UpdateMask(remainSreg1); + LoadForHandleRemainV1( + mainAddr3, tailAddr3, (i * 2 + 0) * V_LENGTH, (i * 2 + 1) * V_LENGTH, mainA3, mainB3, tailA3, + tailB3, pregMask1, xFp32MainAddr3, xFp32TailAddr3, mainAddr4, tailAddr4); + Add(mainA3, mainA3, tailA3, pregMask1); + Add(mainB3, mainB3, tailB3, pregMask1); + Add(mainA3, mainA3, mainB3, pregMask1); + ReduceSum(vMean3, mainA3, pregMask1); + DataCopy(workAddr1 + i, vMean3, pregMerge1); + } + for (uint16_t i = 0; i < (uint16_t)masterRepeats; ++i) { + pregMask1 = UpdateMask(masterSreg1); + LoadForHandleMasterV1( + masterAddr3, (i * 2 + 0) * V_LENGTH, (i * 2 + 1) * V_LENGTH, mainA3, mainB3, pregMask1, + xFp32MasterAddr3, masterAddr4); + Add(mainA3, mainA3, mainB3, pregMask1); + ReduceSum(vMean3, mainA3, pregMask1); + DataCopy(workAddr1 + remainRepeats + i, vMean3, pregMerge1); + } + LocalMemBar(); + for (uint16_t i = 0; i < (uint16_t)mergeRepeats; ++i) { + pregMask = UpdateMask(mergeSreg); + DataCopy(mainA, workAddr + (i * 2 + 0) * V_LENGTH); + DataCopy(mainB, workAddr + (i * 2 + 1) * V_LENGTH); + Add(mainA, mainA, mainB, pregMask); + ReduceSum(vMean, mainA, pregMask); + DataCopy(workAddr + i, vMean, pregMerge); + } + // unroll part + for (uint16_t i = 0; i < (uint16_t)mergeRepeats; ++i) { + pregMask1 = UpdateMask(mergeSreg1); + DataCopy(mainA3, workAddr1 + (i * 2 + 0) * V_LENGTH); + DataCopy(mainB3, workAddr1 + (i * 2 + 1) * V_LENGTH); + Add(mainA3, mainA3, mainB3, pregMask1); + ReduceSum(vMean3, mainA3, pregMask1); + DataCopy(workAddr1 + i, vMean3, pregMerge1); + } + LocalMemBar(); + { + pregMask = UpdateMask(meanSreg); + DataCopy(mainA, workAddr + 0); + ReduceSum(vMean, mainA, pregMask); + Muls(vMean, vMean, avgFactor, pregMerge); + Adds(vMean, vMean, epsilon, pregMerge); + Sqrt(vMean, vMean, pregMerge); + Duplicate(vDupReg, float(1.0), pregMerge); + Div(rstdReg, vDupReg, vMean, pregMerge); + DataCopy(rstdAddr + offset, rstdReg, pregMerge); + } + // unroll part + { + pregMask1 = UpdateMask(meanSreg1); + DataCopy(mainA3, workAddr1 + 0); + ReduceSum(vMean3, mainA3, pregMask1); + Muls(vMean3, vMean3, avgFactor, pregMerge1); + Adds(vMean3, vMean3, epsilon, pregMerge1); + Sqrt(vMean3, vMean3, pregMerge1); + Duplicate(vDupReg3, float(1.0), pregMerge1); + Div(rstdReg3, vDupReg3, vMean3, pregMerge1); + DataCopy(rstdAddr1 + offset, rstdReg3, pregMerge1); + } + offset += 1; + mainAddr += int64_t(count); + tailAddr += int64_t(count); + masterAddr += int64_t(count); + mainAddr2 += int64_t(count); + tailAddr2 += int64_t(count); + masterAddr2 += int64_t(count); + xFp32MainAddr += int64_t(count); + xFp32TailAddr += int64_t(count); + xFp32MasterAddr += int64_t(count); + + mainAddr3 += int64_t(count); + tailAddr3 += int64_t(count); + masterAddr3 += int64_t(count); + mainAddr4 += int64_t(count); + tailAddr4 += int64_t(count); + masterAddr4 += int64_t(count); + xFp32MainAddr3 += int64_t(count); + xFp32TailAddr3 += int64_t(count); + xFp32MasterAddr3 += int64_t(count); + } + } + uint32_t tailDataOffset = unrollOffset + (curRows / 2) * count; + __local_mem__ T* mainAddr5 = (__ubuf__ T*)xLocal1.GetPhyAddr() + tailDataOffset; + __local_mem__ T* tailAddr5 = (__ubuf__ T*)xLocal1.GetPhyAddr() + int64_t(powerSplit) + tailDataOffset; + __local_mem__ T* masterAddr5 = (__ubuf__ T*)xLocal1.GetPhyAddr() + int64_t(remainTile) + tailDataOffset; + __local_mem__ T* mainAddr6 = (__ubuf__ T*)xLocal2.GetPhyAddr() + tailDataOffset; + __local_mem__ T* tailAddr6 = (__ubuf__ T*)xLocal2.GetPhyAddr() + int64_t(powerSplit) + tailDataOffset; + __local_mem__ T* masterAddr6 = (__ubuf__ T*)xLocal2.GetPhyAddr() + int64_t(remainTile) + tailDataOffset; + __local_mem__ float *xFp32MainAddr5, *xFp32TailAddr5, *xFp32MasterAddr5; + + xFp32MainAddr5 = (__ubuf__ float*)xFp32.GetPhyAddr() + tailDataOffset; + xFp32TailAddr5 = (__ubuf__ float*)xFp32.GetPhyAddr() + int64_t(powerSplit) + tailDataOffset; + xFp32MasterAddr5 = (__ubuf__ float*)xFp32.GetPhyAddr() + int64_t(remainTile) + tailDataOffset; + + if (isWithTail) { + __VEC_SCOPE__ + { + uint32_t remainSreg1 = remainTile; + uint32_t masterSreg1 = masterTile; + uint32_t meanSreg1 = meanTile; + uint32_t mergeSreg1 = mergeTile; + RegTensor mainA1, mainB1, tailA1, tailB1, vMean1, vDupReg1, rstdReg1; + MaskReg pregMerge1 = CreateMask(); + MaskReg pregMain1 = CreateMask(); + MaskReg pregMask1; + + for (uint16_t i = 0; i < (uint16_t)remainRepeats; ++i) { + pregMask1 = UpdateMask(remainSreg1); + LoadForHandleRemainV1( + mainAddr5, tailAddr5, (i * 2 + 0) * V_LENGTH, (i * 2 + 1) * V_LENGTH, mainA1, mainB1, tailA1, + tailB1, pregMask1, xFp32MainAddr5, xFp32TailAddr5, mainAddr6, tailAddr6); + Add(mainA1, mainA1, tailA1, pregMask1); + Add(mainB1, mainB1, tailB1, pregMask1); + Add(mainA1, mainA1, mainB1, pregMask1); + ReduceSum(vMean1, mainA1, pregMask1); + DataCopy(workAddr1 + i, vMean1, pregMerge1); + } + for (uint16_t i = 0; i < (uint16_t)masterRepeats; ++i) { + pregMask1 = UpdateMask(masterSreg1); + LoadForHandleMasterV1( + masterAddr5, (i * 2 + 0) * V_LENGTH, (i * 2 + 1) * V_LENGTH, mainA1, mainB1, pregMask1, + xFp32MasterAddr5, masterAddr6); + Add(mainA1, mainA1, mainB1, pregMask1); + ReduceSum(vMean1, mainA1, pregMask1); + DataCopy(workAddr1 + remainRepeats + i, vMean1, pregMerge1); + } + LocalMemBar(); + for (uint16_t i = 0; i < (uint16_t)mergeRepeats; ++i) { + pregMask1 = UpdateMask(mergeSreg1); + DataCopy(mainA1, workAddr1 + (i * 2 + 0) * V_LENGTH); + DataCopy(mainB1, workAddr1 + (i * 2 + 1) * V_LENGTH); + Add(mainA1, mainA1, mainB1, pregMask1); + ReduceSum(vMean1, mainA1, pregMask1); + DataCopy(workAddr1 + i, vMean1, pregMerge1); + } + LocalMemBar(); + { + pregMask1 = UpdateMask(meanSreg1); + DataCopy(mainA1, workAddr1 + 0); + ReduceSum(vMean1, mainA1, pregMask1); + Muls(vMean1, vMean1, avgFactor, pregMerge1); + Adds(vMean1, vMean1, epsilon, pregMerge1); + Sqrt(vMean1, vMean1, pregMerge1); + Duplicate(vDupReg1, float(1.0), pregMerge1); + Div(rstdReg1, vDupReg1, vMean1, pregMerge1); + DataCopy(rstdAddr1 + tailOffset, rstdReg1, pregMerge1); + } + } + } +} + +template +__aicore__ inline void ComputeLatterY( + LocalTensor& xFp32, LocalTensor& gammaLocal, LocalTensor& yLocal, LocalTensor& rstdLocal, + uint32_t offset, uint32_t count, LocalTensor xOutLocal) +{ + uint32_t calCount = count / 2; + uint32_t sreg = (uint32_t)calCount; + uint16_t repeatTimes = CeilDivision(calCount, V_LENGTH); + + __local_mem__ float* xAddr1 = (__ubuf__ float*)xFp32.GetPhyAddr(); + __local_mem__ float* xAddr2 = (__ubuf__ float*)xFp32.GetPhyAddr() + calCount; + __local_mem__ T* gammaAddr1 = (__ubuf__ T*)gammaLocal.GetPhyAddr(); + __local_mem__ T* gammaAddr2 = (__ubuf__ T*)gammaLocal.GetPhyAddr() + calCount; + __local_mem__ float* srcAddr2 = (__ubuf__ float*)rstdLocal.GetPhyAddr(); + __local_mem__ T* yAddr1 = (__ubuf__ T*)yLocal.GetPhyAddr(); + __local_mem__ T* yAddr2 = (__ubuf__ T*)yLocal.GetPhyAddr() + calCount; + __local_mem__ T* xOutAddr1 = (__ubuf__ T*)xOutLocal.GetPhyAddr(); + __local_mem__ T* xOutAddr2 = (__ubuf__ T*)xOutLocal.GetPhyAddr() + calCount; + + if constexpr (IsSameType::value || IsSameType::value) { + __VEC_SCOPE__ + { + RegTensor yB16Reg1, yB16Reg2; + RegTensor xB32Reg1, xB32Reg2; + RegTensor gammaReg1, gammaReg2; + RegTensor rstdReg; + RegTensor dst1Reg, gammaFp32Reg1, yReg1; + RegTensor dst2Reg, gammaFp32Reg2, yReg2; + RegTensor xout16Reg1, xout16Reg2; + MaskReg maskReg; + DataCopy(rstdReg, srcAddr2 + offset); + for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) { + maskReg = UpdateMask(sreg); + DataCopy(xB32Reg1, xAddr1 + i * V_LENGTH); + DataCopy(xB32Reg2, xAddr2 + i * V_LENGTH); + DataCopy(gammaReg1, gammaAddr1 + i * V_LENGTH); + DataCopy(gammaReg2, gammaAddr2 + i * V_LENGTH); + Cast(gammaFp32Reg1, gammaReg1, maskReg); + Cast(gammaFp32Reg2, gammaReg2, maskReg); + Mul(dst1Reg, xB32Reg1, rstdReg, maskReg); + Mul(dst2Reg, xB32Reg2, rstdReg, maskReg); + Mul(yReg1, dst1Reg, gammaFp32Reg1, maskReg); + Mul(yReg2, dst2Reg, gammaFp32Reg2, maskReg); + Cast(yB16Reg1, yReg1, maskReg); + Cast(yB16Reg2, yReg2, maskReg); + DataCopy(yAddr1 + i * V_LENGTH, yB16Reg1, maskReg); + DataCopy(yAddr2 + i * V_LENGTH, yB16Reg2, maskReg); + // outX + Cast(xout16Reg1, xB32Reg1, maskReg); + Cast(xout16Reg2, xB32Reg2, maskReg); + DataCopy(xOutAddr1 + i * V_LENGTH, xout16Reg1, maskReg); + DataCopy(xOutAddr2 + i * V_LENGTH, xout16Reg2, maskReg); + } + } + } else if constexpr (IsSameType::value) { + __VEC_SCOPE__ + { + RegTensor rstdReg; + RegTensor xReg1, gammaReg1, yReg1, vRegTmp1; + RegTensor xReg2, gammaReg2, yReg2, vRegTmp2; + MaskReg maskReg; + DataCopy(rstdReg, srcAddr2 + offset); + for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) { + maskReg = UpdateMask(sreg); + DataCopy(xReg1, xAddr1 + i * V_LENGTH); + DataCopy(xReg2, xAddr2 + i * V_LENGTH); + DataCopy(gammaReg1, gammaAddr1 + i * V_LENGTH); + DataCopy(gammaReg2, gammaAddr2 + i * V_LENGTH); + Mul(vRegTmp1, xReg1, rstdReg, maskReg); + Mul(vRegTmp2, xReg2, rstdReg, maskReg); + Mul(yReg1, vRegTmp1, gammaReg1, maskReg); + Mul(yReg2, vRegTmp2, gammaReg2, maskReg); + DataCopy(yAddr1 + i * V_LENGTH, yReg1, maskReg); + DataCopy(yAddr2 + i * V_LENGTH, yReg2, maskReg); + // outX + DataCopy(xOutAddr1 + i * V_LENGTH, xReg1, maskReg); + DataCopy(xOutAddr2 + i * V_LENGTH, xReg2, maskReg); + } + } + } +} +} // namespace GammaAddRmsNorm +#endif // GAMMA_ADD_RMS_NORM_REGBASE_COMMON_H diff --git a/xllm_ops/gamma_add_rms_norm/op_kernel/arch35/gamma_add_rms_norm_regbase_split_d.h b/xllm_ops/gamma_add_rms_norm/op_kernel/arch35/gamma_add_rms_norm_regbase_split_d.h new file mode 100644 index 0000000..107aee7 --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_kernel/arch35/gamma_add_rms_norm_regbase_split_d.h @@ -0,0 +1,336 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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 gamma_add_rms_norm_regbase_split_d.h + * \brief + */ +#ifndef GAMMA_ADD_RMS_NORM_REGBASE_SPLIT_D_H +#define GAMMA_ADD_RMS_NORM_REGBASE_SPLIT_D_H + +#include "gamma_add_rms_norm_regbase_common.h" +#include "../gamma_add_rms_norm_base.h" +namespace GammaAddRmsNorm { +using namespace AscendC; +using GammaAddRmsNorm::ALIGN_32_FACTOR; +using GammaAddRmsNorm::ComputeLatterY; +using GammaAddRmsNorm::CONST_FACTOR_2; +using GammaAddRmsNorm::Min; +using NormCommon::ComputeMultiLevelRstd; +using RmsNorm::ComputeMultiLevelReduce; +using RmsNorm::ComputeRstd; +using RmsNorm::ComputeSum; +using RmsNorm::DataCopyImpl; +constexpr uint64_t ALIGN_512_FACTOR = 512; +constexpr uint32_t SUM_COUNT = 2; +template +class KernelGammaAddRmsNormRegBaseSplitD { +public: + __aicore__ inline KernelGammaAddRmsNormRegBaseSplitD(TPipe* pipe) + { + pPipe = pipe; + } + __aicore__ inline void Init( + GM_ADDR x1, GM_ADDR x2, GM_ADDR gamma, GM_ADDR y, GM_ADDR rstd, GM_ADDR x, + const GammaAddRMSNormRegbaseTilingData* tiling) + { + ASSERT(GetBlockNum() != 0 && "Block dim can not be zero!"); + numRow = tiling->numRow; + numCol = tiling->numCol; + blockFactor = tiling->blockFactor; + ubFactor = tiling->ubFactor; + ubLoop = tiling->ubLoop; + rowFactor = tiling->rowFactor; + epsilon = tiling->epsilon; + colBufferLength = tiling->colBuferLength; + avgFactor = tiling->avgFactor; + addGammaOffset = tiling->addGammaOffset; + rowWork = (GetBlockIdx() < GetBlockNum() - 1) ? blockFactor : numRow - (GetBlockNum() - 1) * blockFactor; + xGm1.SetGlobalBuffer((__gm__ T*)x1 + GetBlockIdx() * blockFactor * numCol, rowWork * numCol); + xGm2.SetGlobalBuffer((__gm__ T*)x2 + GetBlockIdx() * blockFactor * numCol, rowWork * numCol); + gammaGm.SetGlobalBuffer((__gm__ T*)gamma, numCol); + yGm.SetGlobalBuffer((__gm__ T*)y + GetBlockIdx() * blockFactor * numCol, rowWork * numCol); + rstdGm.SetGlobalBuffer((__gm__ float*)rstd + GetBlockIdx() * blockFactor, blockFactor); + xOutGm.SetGlobalBuffer((__gm__ T*)x + GetBlockIdx() * blockFactor * numCol, rowWork * numCol); + pPipe->InitBuffer(inQueueX1, DOUBLE_BUFFER_NUM, colBufferLength * sizeof(T)); + pPipe->InitBuffer(inQueueX2, DOUBLE_BUFFER_NUM, colBufferLength * sizeof(T)); + pPipe->InitBuffer(inQueueGamma, DOUBLE_BUFFER_NUM, colBufferLength * sizeof(T)); + pPipe->InitBuffer(outQueueY, DOUBLE_BUFFER_NUM, colBufferLength * sizeof(T)); + pPipe->InitBuffer(outQueueX, DOUBLE_BUFFER_NUM, colBufferLength * sizeof(T)); + pPipe->InitBuffer(outQueueRstd, DOUBLE_BUFFER_NUM, rowFactor * sizeof(float)); + if constexpr (!is_same::value) { + pPipe->InitBuffer(xFp32Buf, ubFactor * sizeof(float)); + pPipe->InitBuffer(xFp32Buf1, ubFactor * sizeof(float)); + } + pPipe->InitBuffer(workLocalBuf, ONCE_VECTOR_SIZE * sizeof(float)); + pPipe->InitBuffer(level1Buf, ONCE_VECTOR_SIZE * sizeof(float)); + pPipe->InitBuffer(level2Buf, ONCE_VECTOR_SIZE * sizeof(float)); + pPipe->InitBuffer(level3Buf, ONCE_VECTOR_SIZE * sizeof(float)); + pPipe->InitBuffer(tempBuf, V_LENGTH * sizeof(float)); + } + + __aicore__ inline void Process() + { + uint32_t repeatTimes = CeilDiv(rowWork, rowFactor); + for (uint32_t repeat = 0; repeat < repeatTimes; repeat++) { + uint32_t remain = rowWork - repeat * rowFactor; + uint32_t calRowNum = Min(remain, rowFactor); + SubProcess(repeat, calRowNum); + } + } + + __aicore__ inline void SubProcess(uint32_t rowRepeat, uint32_t calRowNum) + { + LocalTensor rstdLocal = outQueueRstd.AllocTensor(); + Duplicate(rstdLocal, (float)0.0, rowFactor); + uint32_t colRepeats = CeilDiv(numCol, ubFactor); + + for (uint32_t row = 0; row < calRowNum; row++) { + uint32_t split = ubLoop * ubFactor; + uint32_t colTail = numCol - split; + uint64_t offsets = (rowRepeat * rowFactor + row) * numCol; + uint32_t tail = colTail % ubFactor; + uint32_t tailLoop = colTail / ubFactor; + uint32_t masterLoop = tail != 0 ? 1 : 0; + masterLoop = ubLoop - tailLoop - masterLoop; + ComputeFormer(offsets, rstdLocal, row, masterLoop, tailLoop, tail); + } + + ComputeRstd(rstdLocal, epsilon, avgFactor, calRowNum); + + for (uint32_t repeat = 0; repeat < colRepeats; repeat++) { + uint32_t remain = numCol - repeat * ubFactor; + uint32_t calColNum = Min(remain, ubFactor); + ComputeLatter(rowRepeat, calRowNum, repeat, rstdLocal, calColNum); + } + + outQueueRstd.EnQue(rstdLocal); + CopyOutRstd(rowRepeat, calRowNum); + } + +private: + __aicore__ inline void CopyInX(uint64_t offset, uint32_t count, uint32_t left = 0, uint32_t right = 0) + { + LocalTensor xLocal1 = inQueueX1.AllocTensor(); + LocalTensor xLocal2 = inQueueX2.AllocTensor(); + DataCopyPadExtParams padParams{ + true, // isPad + static_cast(left), // leftPadding + static_cast(right), // rightPadding + static_cast(0.0) // paddingValue + }; + DataCopyImpl(xLocal1, xGm1[offset], 1, count, 0, 0, padParams); + inQueueX1.EnQue(xLocal1); + + DataCopyImpl(xLocal2, xGm2[offset], 1, count, 0, 0, padParams); + inQueueX2.EnQue(xLocal2); + } + + __aicore__ inline void CopyInGamma(uint32_t colRepeat, uint32_t calColNum) + { + LocalTensor gammaLocal = inQueueGamma.AllocTensor(); + DataCopyImpl(gammaLocal, gammaGm[colRepeat * ubFactor], 1, calColNum, 0, 0); + inQueueGamma.EnQue(gammaLocal); + } + + __aicore__ inline void ComputeFormerHandle( + LocalTensor& dstLocal, uint64_t srcOffset, uint64_t dstOffset, uint32_t count, uint32_t power) + { + uint32_t calCount = CeilAlign((uint64_t)(count * sizeof(T)), ALIGN_32_FACTOR) / sizeof(T); + CopyInX(srcOffset, count, 0, calCount - count); + LocalTensor xLocal1 = inQueueX1.DeQue(); + LocalTensor xLocal2 = inQueueX2.DeQue(); + LocalTensor workLocal = workLocalBuf.Get(); + uint32_t calNum = CeilAlign((uint64_t)(count * sizeof(T)), ALIGN_512_FACTOR) / sizeof(T); + if (calNum - calCount > 0) { + Duplicate(xLocal1[calCount], (T)0.0, calNum - calCount); + Duplicate(xLocal2[calCount], (T)0.0, calNum - calCount); + } + ComputeFormerImplV2(dstLocal, xLocal1, xLocal2, workLocal, dstOffset, calNum, power); + inQueueX1.FreeTensor(xLocal1); + inQueueX2.FreeTensor(xLocal2); + } + + __aicore__ inline void ComputeFormer( + uint32_t curRow, LocalTensor dstLocal, uint32_t position, uint32_t masterLoop, uint32_t tailLoop, + uint32_t tail) + { + uint64_t offset{curRow}; + uint32_t level1{0}; + uint32_t level2{0}; + uint32_t level3{0}; + LocalTensor level1Local = level1Buf.Get(); + LocalTensor level2Local = level2Buf.Get(); + LocalTensor level3Local = level3Buf.Get(); + LocalTensor tempLocal = tempBuf.Get(); + Duplicate(level1Local, (float)0.0, ONCE_VECTOR_SIZE); + Duplicate(level2Local, (float)0.0, ONCE_VECTOR_SIZE); + Duplicate(level3Local, (float)0.0, ONCE_VECTOR_SIZE); + // Stage 1: process complete tail blocks. + for (uint32_t repeat = 0; repeat < tailLoop; repeat++) { + ComputeFormerHandle(tempLocal, offset, 0, ubFactor, ubFactor); + offset += ubFactor; + + ComputeFormerHandle(tempLocal, offset, 1, ubFactor, ubFactor); + offset += ubFactor; + + ComputeSum(level1Local, tempLocal, level1, SUM_COUNT); + level1 += 1; + ComputeMultiLevelReduce(level1Local, level2Local, level3Local, level1, level2, level3); + } + // Stage 2: process the partial tail block. + if (tail > 0 && tail <= ubFactor / CONST_FACTOR_2) { + ComputeFormerHandle(tempLocal, offset, 0, ubFactor / CONST_FACTOR_2 + tail, ubFactor / CONST_FACTOR_2); + offset += ubFactor / CONST_FACTOR_2 + tail; + + ComputeFormerHandle(tempLocal, offset, 1, ubFactor / CONST_FACTOR_2, ubFactor / CONST_FACTOR_2); + offset += ubFactor / CONST_FACTOR_2; + + ComputeSum(level1Local, tempLocal, level1, SUM_COUNT); + level1 += 1; + ComputeMultiLevelReduce(level1Local, level2Local, level3Local, level1, level2, level3); + } else if (tail > ubFactor / CONST_FACTOR_2) { + ComputeFormerHandle(tempLocal, offset, 0, ubFactor, ubFactor); + offset += ubFactor; + + ComputeFormerHandle(tempLocal, offset, 1, tail, ubFactor / CONST_FACTOR_2); + offset += tail; + + ComputeSum(level1Local, tempLocal, level1, SUM_COUNT); + level1 += 1; + ComputeMultiLevelReduce(level1Local, level2Local, level3Local, level1, level2, level3); + } + // Stage 3: process the main blocks. + for (uint32_t repeat = 0; repeat < masterLoop; repeat++) { + ComputeFormerHandle(level1Local, offset, level1, ubFactor, ubFactor); + offset += ubFactor; + level1 += 1; + ComputeMultiLevelReduce(level1Local, level2Local, level3Local, level1, level2, level3); + } + ComputeMultiLevelRstd(dstLocal, position, level1Local, level2Local, level3Local, level1, level2, level3); + } + + __aicore__ inline void ComputeLatter( + uint32_t rowRepeat, uint32_t calRowNum, uint32_t colRepeat, LocalTensor& rstdLocal, uint32_t calColNum) + { + CopyInGamma(colRepeat, calColNum); + LocalTensor gammaLocal = inQueueGamma.DeQue(); + if (addGammaOffset != 0U) { + AddGammaOffset(gammaLocal, calColNum); + } + for (uint32_t row = 0; row < calRowNum; row++) { + uint64_t offset = (rowRepeat * rowFactor + row) * numCol + colRepeat * ubFactor; + CopyInX(offset, calColNum); + LocalTensor xLocal1 = inQueueX1.DeQue(); + LocalTensor xLocal2 = inQueueX2.DeQue(); + LocalTensor xFp32; + LocalTensor xFp32Other; + if constexpr (!is_same::value) { + xFp32 = xFp32Buf.Get(); + xFp32Other = xFp32Buf1.Get(); + Cast(xFp32, xLocal1, AscendC::RoundMode::CAST_NONE, calColNum); + Cast(xFp32Other, xLocal2, AscendC::RoundMode::CAST_NONE, calColNum); + PipeBarrier(); + AscendC::Add(xFp32, xFp32, xFp32Other, calColNum); + PipeBarrier(); + } else { + AscendC::Add(xLocal1, xLocal1, xLocal2, calColNum); + PipeBarrier(); + } + LocalTensor yLocal = outQueueY.AllocTensor(); + LocalTensor xOutLocal = outQueueX.AllocTensor(); + uint32_t calCount = CeilAlign((uint64_t)(calColNum * sizeof(T)), ALIGN_512_FACTOR) / sizeof(T); + if constexpr (!is_same::value) { + ComputeLatterY(xFp32, gammaLocal, yLocal, rstdLocal, row, calCount, xOutLocal); + } else { + ComputeLatterY(xLocal1, gammaLocal, yLocal, rstdLocal, row, calCount, xOutLocal); + } + inQueueX1.FreeTensor(xLocal1); + inQueueX2.FreeTensor(xLocal2); + outQueueY.EnQue(yLocal); + outQueueX.EnQue(xOutLocal); + CopyOutY(rowRepeat * rowFactor + row, colRepeat, calColNum); + CopyOutX(rowRepeat * rowFactor + row, colRepeat, calColNum); + } + inQueueGamma.FreeTensor(gammaLocal); + } + + __aicore__ inline void AddGammaOffset(LocalTensor& gammaLocal, uint32_t elementNum) + { + if constexpr (is_same::value) { + LocalTensor gammaFp32 = xFp32Buf.Get(); + Cast(gammaFp32, gammaLocal, RoundMode::CAST_NONE, elementNum); + PipeBarrier(); + Adds(gammaFp32, gammaFp32, static_cast(1.0), elementNum); + PipeBarrier(); + Cast(gammaLocal, gammaFp32, RoundMode::CAST_RINT, elementNum); + } else { + Adds(gammaLocal, gammaLocal, static_cast(1.0), elementNum); + } + PipeBarrier(); + } + + __aicore__ inline void CopyOutY(uint32_t curRow, uint32_t curCol, uint32_t calColNum) + { + LocalTensor yLocal = outQueueY.DeQue(); + DataCopyImpl(yGm[curRow * numCol + curCol * ubFactor], yLocal, 1, calColNum, 0, 0); + outQueueY.FreeTensor(yLocal); + } + + __aicore__ inline void CopyOutRstd(uint32_t rowRepeat, uint32_t calRowNum) + { + LocalTensor rstdLocal = outQueueRstd.DeQue(); + DataCopyImpl(rstdGm[rowRepeat * rowFactor], rstdLocal, 1, calRowNum, 0, 0); + outQueueRstd.FreeTensor(rstdLocal); + } + + __aicore__ inline void CopyOutX(uint32_t curRow, uint32_t curCol, uint32_t calColNum) + { + LocalTensor xOutLocal = outQueueX.DeQue(); + DataCopyImpl(xOutGm[curRow * numCol + curCol * ubFactor], xOutLocal, 1, calColNum, 0, 0); + outQueueX.FreeTensor(xOutLocal); + } + +private: + TPipe* pPipe = nullptr; + TQue inQueueX1; + TQue inQueueX2; + TQue inQueueGamma; + TQue outQueueY; + TQue outQueueX; + TQue outQueueRstd; + TBuf xFp32Buf; + TBuf xFp32Buf1; + TBuf workLocalBuf; + TBuf level1Buf; + TBuf level2Buf; + TBuf level3Buf; + TBuf tempBuf; + GlobalTensor xGm1; + GlobalTensor xGm2; + GlobalTensor gammaGm; + GlobalTensor yGm; + GlobalTensor xOutGm; + GlobalTensor rstdGm; + uint32_t numRow; + uint32_t numCol; + uint32_t blockFactor; + uint32_t ubFactor; + uint32_t colBufferLength; + uint32_t ubLoop; + uint32_t rowFactor; + float epsilon; + float avgFactor; + uint32_t addGammaOffset{0}; + uint32_t rowWork{1}; +}; +} // namespace GammaAddRmsNorm +#endif // _GAMMA_ADD_RMS_NORM_REGBASE_SPLIT_D_H diff --git a/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm.cpp b/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm.cpp new file mode 100644 index 0000000..3ccc1ee --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm.cpp @@ -0,0 +1,129 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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 gamma_add_rms_norm.cpp + * \brief + */ +#include "gamma_add_rms_norm.h" +#include "gamma_add_rms_norm_split_d.h" +#include "gamma_add_rms_norm_merge_n.h" +#include "gamma_add_rms_norm_multi_n.h" +#include "gamma_add_rms_norm_single_n.h" + +using namespace AscendC; + +#define GENERAL_OP_IMPL(templateClass, ...) \ + do { \ + templateClass<__VA_ARGS__> op(&pipe); \ + op.Init(x1, x2, gamma, y, rstd, x, workspace, &tilingData); \ + op.Process(); \ + } while (0) + +extern "C" __global__ __aicore__ void gamma_add_rms_norm( + GM_ADDR x1, GM_ADDR x2, GM_ADDR gamma, GM_ADDR y, GM_ADDR rstd, GM_ADDR x, GM_ADDR workspace, GM_ADDR tiling) +{ + TPipe pipe; + GET_TILING_DATA(tilingData, tiling); + if (TILING_KEY_IS(10)) { + GENERAL_OP_IMPL(KernelGammaAddRmsNorm, half, 1); + } else if (TILING_KEY_IS(20)) { + GENERAL_OP_IMPL(KernelGammaAddRmsNorm, float, 1); + } else if (TILING_KEY_IS(30)) { +#if !(defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113)) + GENERAL_OP_IMPL(KernelGammaAddRmsNorm, bfloat16_t, 1); +#endif + } else if (TILING_KEY_IS(11)) { + GENERAL_OP_IMPL(KernelGammaAddRmsNormSplitD, half, 1); + } else if (TILING_KEY_IS(21)) { + GENERAL_OP_IMPL(KernelGammaAddRmsNormSplitD, float, 1); + } else if (TILING_KEY_IS(31)) { +#if !(defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113)) + GENERAL_OP_IMPL(KernelGammaAddRmsNormSplitD, bfloat16_t, 1); +#endif + } else if (TILING_KEY_IS(12)) { + GENERAL_OP_IMPL(KernelGammaAddRmsNormMergeN, half, 1); + } else if (TILING_KEY_IS(22)) { + GENERAL_OP_IMPL(KernelGammaAddRmsNormMergeN, float, 1); + } else if (TILING_KEY_IS(32)) { +#if !(defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113)) + GENERAL_OP_IMPL(KernelGammaAddRmsNormMergeN, bfloat16_t, 1); +#endif + } else if (TILING_KEY_IS(13)) { + GENERAL_OP_IMPL(KernelGammaAddRmsNormSingleN, half, 1); + } else if (TILING_KEY_IS(23)) { + GENERAL_OP_IMPL(KernelGammaAddRmsNormSingleN, float, 1); + } else if (TILING_KEY_IS(33)) { +#if !(defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113)) + GENERAL_OP_IMPL(KernelGammaAddRmsNormSingleN, bfloat16_t, 1); +#endif + } else if (TILING_KEY_IS(14)) { + GENERAL_OP_IMPL(KernelGammaAddRmsNormMultiN, half, 1); + } else if (TILING_KEY_IS(34)) { + GENERAL_OP_IMPL(KernelGammaAddRmsNormMultiN, bfloat16_t, 1); + } else if (TILING_KEY_IS(110)) { + GENERAL_OP_IMPL(KernelGammaAddRmsNorm, half, 2); + } else if (TILING_KEY_IS(130)) { +#if !(defined(__NPU_ARCH__) && __NPU_ARCH__ == 3003) + GENERAL_OP_IMPL(KernelGammaAddRmsNorm, bfloat16_t, 2); +#endif + } else if (TILING_KEY_IS(111)) { + GENERAL_OP_IMPL(KernelGammaAddRmsNormSplitD, half, 2); + } else if (TILING_KEY_IS(131)) { +#if !(defined(__NPU_ARCH__) && __NPU_ARCH__ == 3003) + GENERAL_OP_IMPL(KernelGammaAddRmsNormSplitD, bfloat16_t, 2); +#endif + } else if (TILING_KEY_IS(112)) { + GENERAL_OP_IMPL(KernelGammaAddRmsNormMergeN, half, 2); + } else if (TILING_KEY_IS(132)) { +#if !(defined(__NPU_ARCH__) && __NPU_ARCH__ == 3003) + GENERAL_OP_IMPL(KernelGammaAddRmsNormMergeN, bfloat16_t, 2); +#endif + } else if (TILING_KEY_IS(113)) { + GENERAL_OP_IMPL(KernelGammaAddRmsNormSingleN, half, 2); + } else if (TILING_KEY_IS(133)) { +#if !(defined(__NPU_ARCH__) && __NPU_ARCH__ == 3003) + GENERAL_OP_IMPL(KernelGammaAddRmsNormSingleN, bfloat16_t, 2); +#endif + } else if (TILING_KEY_IS(114)) { + GENERAL_OP_IMPL(KernelGammaAddRmsNormMultiN, half, 2); + } else if (TILING_KEY_IS(134)) { + GENERAL_OP_IMPL(KernelGammaAddRmsNormMultiN, bfloat16_t, 2); + } else if (TILING_KEY_IS(1010)) { + GENERAL_OP_IMPL(KernelGammaAddRmsNorm, half, 3); + } else if (TILING_KEY_IS(1030)) { +#if !(defined(__NPU_ARCH__) && __NPU_ARCH__ == 3003) + GENERAL_OP_IMPL(KernelGammaAddRmsNorm, bfloat16_t, 3); +#endif + } else if (TILING_KEY_IS(1011)) { + GENERAL_OP_IMPL(KernelGammaAddRmsNormSplitD, half, 3); + } else if (TILING_KEY_IS(1031)) { +#if !(defined(__NPU_ARCH__) && __NPU_ARCH__ == 3003) + GENERAL_OP_IMPL(KernelGammaAddRmsNormSplitD, bfloat16_t, 3); +#endif + } else if (TILING_KEY_IS(1012)) { + GENERAL_OP_IMPL(KernelGammaAddRmsNormMergeN, half, 3); + } else if (TILING_KEY_IS(1032)) { +#if !(defined(__NPU_ARCH__) && __NPU_ARCH__ == 3003) + GENERAL_OP_IMPL(KernelGammaAddRmsNormMergeN, bfloat16_t, 3); +#endif + } else if (TILING_KEY_IS(1013)) { + GENERAL_OP_IMPL(KernelGammaAddRmsNormSingleN, half, 3); + } else if (TILING_KEY_IS(1033)) { +#if !(defined(__NPU_ARCH__) && __NPU_ARCH__ == 3003) + GENERAL_OP_IMPL(KernelGammaAddRmsNormSingleN, bfloat16_t, 3); +#endif + } else if (TILING_KEY_IS(1014)) { + GENERAL_OP_IMPL(KernelGammaAddRmsNormMultiN, half, 3); + } else if (TILING_KEY_IS(1034)) { + GENERAL_OP_IMPL(KernelGammaAddRmsNormMultiN, bfloat16_t, 3); + } +} diff --git a/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm.h b/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm.h new file mode 100644 index 0000000..ae0a989 --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm.h @@ -0,0 +1,366 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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 gamma_add_rms_norm.h + * \brief add rms norm file + */ +#ifndef GAMMA_ADD_RMS_NORM_H_ +#define GAMMA_ADD_RMS_NORM_H_ +#include "gamma_add_rms_norm_base.h" + +using namespace AscendC; +using namespace RmsNorm; + +template +class KernelGammaAddRmsNorm { +public: + __aicore__ inline KernelGammaAddRmsNorm(TPipe* pipe) + { + Ppipe = pipe; + } + __aicore__ inline void Init( + GM_ADDR x1, GM_ADDR x2, GM_ADDR gamma, GM_ADDR y, GM_ADDR rstd, GM_ADDR x, GM_ADDR workspace, const GammaAddRMSNormTilingData* tiling) + { + ASSERT(GetBlockNum() != 0 && "Block dim can not be zero!"); + this->numRow = tiling->num_row; + this->numCol = tiling->num_col; + this->blockFactor = tiling->block_factor; + this->rowFactor = tiling->row_factor; + this->ubFactor = tiling->ub_factor; + this->epsilon = tiling->epsilon; + this->avgFactor = (this->numCol != 0) ? (float)1.0 / this->numCol : 0; + this->addGammaOffset = tiling->add_gamma_offset; + + blockIdx_ = GetBlockIdx(); + if (blockIdx_ < GetBlockNum() - 1) { + this->rowWork = this->blockFactor; + } else if (blockIdx_ == GetBlockNum() - 1) { + this->rowWork = this->numRow - (GetBlockNum() - 1) * this->blockFactor; + } + // get start index for current core, core parallel + x1Gm.SetGlobalBuffer((__gm__ T*)x1 + blockIdx_ * this->blockFactor * this->numCol, this->rowWork * this->numCol); + x2Gm.SetGlobalBuffer((__gm__ T*)x2 + blockIdx_ * this->blockFactor * this->numCol, this->rowWork * this->numCol); + gammaGm.SetGlobalBuffer((__gm__ T*)gamma, this->numCol); + yGm.SetGlobalBuffer((__gm__ T*)y + blockIdx_ * this->blockFactor * this->numCol, this->rowWork * this->numCol); + + if constexpr (MODE == GAMMA_ADD_RMS_NORM_MODE) { + rstdGm.SetGlobalBuffer((__gm__ float*)rstd + blockIdx_ * this->blockFactor, this->blockFactor); + xGm.SetGlobalBuffer((__gm__ T*)x + blockIdx_ * this->blockFactor * this->numCol, this->rowWork * this->numCol); + } + if constexpr (MODE == PRE_RMS_NORM_MODE) { + xGm.SetGlobalBuffer((__gm__ T*)x + blockIdx_ * this->blockFactor * this->numCol, this->rowWork * this->numCol); + } + + // pipe alloc memory to queue, the unit is Bytes + Ppipe->InitBuffer(inQueueX, BUFFER_NUM, ubFactor * sizeof(T)); + Ppipe->InitBuffer(inQueueGamma, BUFFER_NUM, ubFactor * sizeof(T)); + Ppipe->InitBuffer(outQueueY, BUFFER_NUM, ubFactor * sizeof(T)); + Ppipe->InitBuffer(outQueueRstd, BUFFER_NUM, rowFactor * sizeof(float)); + + if constexpr (is_same::value || is_same::value) { + Ppipe->InitBuffer(xFp32Buf, ubFactor * sizeof(float)); + } + Ppipe->InitBuffer(sqxBuf, ubFactor * sizeof(float)); + Ppipe->InitBuffer(reduceFp32Buf, NUM_PER_REP_FP32 * sizeof(float)); + } + + __aicore__ inline void Process() + { + CopyInGamma(); + LocalTensor gammaLocal = inQueueGamma.DeQue(); + if (addGammaOffset != 0U) { + AddGammaOffset(gammaLocal, numCol); + } + + uint32_t i_o_max = RmsNorm::CeilDiv(this->rowWork, this->rowFactor); + uint32_t row_tail = this->rowWork - (i_o_max - 1) * this->rowFactor; + + for (uint32_t i_o = 0; i_o < i_o_max - 1; i_o++) { + SubProcess(i_o, this->rowFactor, gammaLocal); + } + SubProcess(i_o_max - 1, row_tail, gammaLocal); + inQueueGamma.FreeTensor(gammaLocal); + } + + __aicore__ inline void SubProcess(uint32_t i_o, uint32_t calc_row_num, LocalTensor& gammaLocal) + { + LocalTensor rstdLocal = outQueueRstd.AllocTensor(); + for (uint32_t i_i = 0; i_i < calc_row_num; i_i++) { + uint64_t gm_bias = (static_cast(i_o) * static_cast(this->rowFactor) + static_cast(i_i)) * static_cast(this->numCol); + CopyIn(gm_bias); + Compute(i_i, gammaLocal, rstdLocal); + CopyOutY(gm_bias); + } + if constexpr (MODE == GAMMA_ADD_RMS_NORM_MODE) { + outQueueRstd.EnQue(rstdLocal); + CopyOutRstd(i_o, calc_row_num); + } else { + outQueueRstd.FreeTensor(rstdLocal); + } + } + +private: + __aicore__ inline void CopyIn(uint32_t gm_bias) + { + LocalTensor x1Local_in = inQueueX.AllocTensor(); + LocalTensor x2Local = sqxBuf.Get(); + LocalTensor xLocal = outQueueY.AllocTensor(); + + if constexpr (is_same::value || is_same::value) { + x2Local = x2Local[ubFactor]; + } + + DataCopyCustom(x1Local_in, x1Gm[gm_bias], numCol); + DataCopyCustom(x2Local, x2Gm[gm_bias], numCol); + inQueueX.EnQue(x1Local_in); + auto x1Local = inQueueX.DeQue(); + + if constexpr (is_same::value) { + LocalTensor x1_fp32 = xFp32Buf.Get(); + Add(xLocal, x1Local, x2Local, numCol); + PipeBarrier(); + Cast(x1_fp32, xLocal, RoundMode::CAST_NONE, numCol); + PipeBarrier(); + } else if constexpr (is_same::value) { + LocalTensor x1_fp32 = xFp32Buf.Get(); + LocalTensor x2_fp32 = sqxBuf.Get(); + Cast(x1_fp32, x1Local, RoundMode::CAST_NONE, numCol); + Cast(x2_fp32, x2Local, RoundMode::CAST_NONE, numCol); + PipeBarrier(); + Add(x1_fp32, x1_fp32, x2_fp32, numCol); + PipeBarrier(); + Cast(xLocal, x1_fp32, RoundMode::CAST_RINT, numCol); + PipeBarrier(); + } else { + Add(x1Local, x1Local, x2Local, numCol); + PipeBarrier(); + Adds(xLocal, x1Local, (float)0, numCol); + } + inQueueX.FreeTensor(x1Local); + + // CopyOut x1 + x2 + outQueueY.EnQue(xLocal); + auto x_out = outQueueY.DeQue(); + if constexpr (MODE == GAMMA_ADD_RMS_NORM_MODE || MODE == PRE_RMS_NORM_MODE) { + DataCopyCustom(xGm[gm_bias], x_out, numCol); + } + outQueueY.FreeTensor(x_out); + } + + __aicore__ inline void CopyInGamma() + { + LocalTensor gammaLocal = inQueueGamma.AllocTensor(); + DataCopyCustom(gammaLocal, gammaGm, numCol); + inQueueGamma.EnQue(gammaLocal); + } + + __aicore__ inline void AddGammaOffset(LocalTensor& gammaLocal, uint32_t elementNum) + { + if constexpr (is_same::value) { + LocalTensor gammaFp32 = xFp32Buf.Get(); + Cast(gammaFp32, gammaLocal, RoundMode::CAST_NONE, elementNum); + PipeBarrier(); + Adds(gammaFp32, gammaFp32, static_cast(1.0), elementNum); + PipeBarrier(); + Cast(gammaLocal, gammaFp32, RoundMode::CAST_RINT, elementNum); + } else { + Adds(gammaLocal, gammaLocal, static_cast(1.0), elementNum); + } + PipeBarrier(); + } + + __aicore__ inline void Compute(uint32_t inner_progress, LocalTensor gammaLocal, LocalTensor rstdLocal) + { + LocalTensor xLocal = inQueueX.AllocTensor(); + LocalTensor sqx = sqxBuf.Get(); + LocalTensor reduce_buf_local = reduceFp32Buf.Get(); + Mul(sqx, xLocal, xLocal, numCol); + PipeBarrier(); + + Muls(sqx, sqx, avgFactor, numCol); + PipeBarrier(); + + ReduceSumCustom(sqx, sqx, reduce_buf_local, numCol); + PipeBarrier(); + Adds(sqx, sqx, epsilon, 1); + PipeBarrier(); + + Sqrt(sqx, sqx, 1); + Duplicate(reduce_buf_local, ONE, 1); + PipeBarrier(); + Div(sqx, reduce_buf_local, sqx, 1); + PipeBarrier(); + event_t event_v_s = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_S)); + SetFlag(event_v_s); + WaitFlag(event_v_s); + float rstdValue = sqx.GetValue(0); + event_t event_s_v = static_cast(GetTPipePtr()->FetchEventID(HardEvent::S_V)); + SetFlag(event_s_v); + WaitFlag(event_s_v); + rstdLocal.SetValue(inner_progress, rstdValue); + PipeBarrier(); + LocalTensor yLocal = outQueueY.AllocTensor(); + Muls(yLocal, xLocal, rstdValue, numCol); + inQueueX.FreeTensor(xLocal); + PipeBarrier(); + Mul(yLocal, gammaLocal, yLocal, numCol); + PipeBarrier(); + outQueueY.EnQue(yLocal); + } + + __aicore__ inline void Compute( + uint32_t inner_progress, LocalTensor gammaLocal, LocalTensor rstdLocal) + { + LocalTensor x_fp32 = xFp32Buf.Get(); + LocalTensor sqx = sqxBuf.Get(); + LocalTensor reduce_buf_local = reduceFp32Buf.Get(); + + Mul(sqx, x_fp32, x_fp32, numCol); + PipeBarrier(); + + Muls(sqx, sqx, this->avgFactor, this->numCol); + PipeBarrier(); + ReduceSumCustom(sqx, sqx, reduce_buf_local, this->numCol); + PipeBarrier(); + + Adds(sqx, sqx, this->epsilon, 1); + PipeBarrier(); + + Sqrt(sqx, sqx, 1); + Duplicate(reduce_buf_local, ONE, 1); + PipeBarrier(); + Div(sqx, reduce_buf_local, sqx, 1); + PipeBarrier(); + event_t event_v_s2 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_S)); + SetFlag(event_v_s2); + WaitFlag(event_v_s2); + float rstdValue2 = sqx.GetValue(0); + event_t event_s_v2 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::S_V)); + SetFlag(event_s_v2); + WaitFlag(event_s_v2); + rstdLocal.SetValue(inner_progress, rstdValue2); + PipeBarrier(); + Muls(x_fp32, x_fp32, rstdValue2, this->numCol); + PipeBarrier(); + LocalTensor yLocal = outQueueY.AllocTensor(); + Cast(yLocal, x_fp32, RoundMode::CAST_RINT, this->numCol); + PipeBarrier(); + Cast(x_fp32, yLocal, RoundMode::CAST_NONE, this->numCol); + PipeBarrier(); + Cast(sqx, gammaLocal, RoundMode::CAST_NONE, numCol); // gamma_fp32 reuse sqx + PipeBarrier(); + Mul(x_fp32, x_fp32, sqx, numCol); + PipeBarrier(); + Cast(yLocal, x_fp32, RoundMode::CAST_RINT, numCol); + PipeBarrier(); + + event_t event_v_mte = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_MTE2)); + SetFlag(event_v_mte); + WaitFlag(event_v_mte); + + outQueueY.EnQue(yLocal); + } + + __aicore__ inline void Compute(uint32_t inner_progress, LocalTensor gammaLocal, LocalTensor rstdLocal) + { + LocalTensor x_fp32 = xFp32Buf.Get(); + LocalTensor sqx = sqxBuf.Get(); + LocalTensor reduce_buf_local = reduceFp32Buf.Get(); + + Mul(sqx, x_fp32, x_fp32, this->numCol); + PipeBarrier(); + + Muls(sqx, sqx, this->avgFactor, this->numCol); + PipeBarrier(); + + ReduceSumCustom(sqx, sqx, reduce_buf_local, this->numCol); + PipeBarrier(); + + Adds(sqx, sqx, this->epsilon, 1); + PipeBarrier(); + + Sqrt(sqx, sqx, 1); + Duplicate(reduce_buf_local, ONE, 1); + PipeBarrier(); + Div(sqx, reduce_buf_local, sqx, 1); + PipeBarrier(); + event_t event_v_s3 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_S)); + SetFlag(event_v_s3); + WaitFlag(event_v_s3); + float rstdValue3 = sqx.GetValue(0); + event_t event_s_v3 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::S_V)); + SetFlag(event_s_v3); + WaitFlag(event_s_v3); + rstdLocal.SetValue(inner_progress, rstdValue3); + PipeBarrier(); + Muls(x_fp32, x_fp32, rstdValue3, this->numCol); + PipeBarrier(); + LocalTensor yLocal = outQueueY.AllocTensor(); + Cast(yLocal, x_fp32, RoundMode::CAST_NONE, this->numCol); + + event_t event_v_mte = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_MTE2)); + SetFlag(event_v_mte); + WaitFlag(event_v_mte); + + PipeBarrier(); + Mul(yLocal, gammaLocal, yLocal, numCol); + PipeBarrier(); + outQueueY.EnQue(yLocal); + } + + __aicore__ inline void CopyOutY(uint32_t progress) + { + LocalTensor yLocal = outQueueY.DeQue(); + DataCopyCustom(yGm[progress], yLocal, numCol); + outQueueY.FreeTensor(yLocal); + } + + __aicore__ inline void CopyOutRstd(uint32_t outer_progress, uint32_t num) + { + LocalTensor rstdLocal = outQueueRstd.DeQue(); +#if __CCE_AICORE__ == 220 || (defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113)) + DataCopyCustom(rstdGm[outer_progress * this->rowFactor], rstdLocal, num); +#endif + outQueueRstd.FreeTensor(rstdLocal); + } + +private: + TPipe* Ppipe = nullptr; + // create queues for input, in this case depth is equal to buffer num + TQue inQueueX; + TQue inQueueGamma; + // create queues for output, in this case depth is equal to buffer num + TQue outQueueY; + TQue outQueueRstd; + + TBuf xFp32Buf; + TBuf sqxBuf; + TBuf reduceFp32Buf; + GlobalTensor x1Gm; + GlobalTensor x2Gm; + GlobalTensor gammaGm; + GlobalTensor yGm; + GlobalTensor rstdGm; + GlobalTensor xGm; + + uint32_t numRow; + uint32_t numCol; + uint32_t blockFactor; + uint32_t rowFactor; + uint32_t ubFactor; + float epsilon; + float avgFactor; + int32_t blockIdx_; + uint32_t rowWork = 1; + uint32_t addGammaOffset = 0; +}; +#endif // GAMMA_ADD_RMS_NORM_H_ diff --git a/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_apt.cpp b/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_apt.cpp new file mode 100644 index 0000000..efdc8a6 --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_apt.cpp @@ -0,0 +1,39 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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 gamma_add_rms_norm.cpp + * \brief + */ + +#include "arch35/gamma_add_rms_norm_regbase.h" +#include "arch35/gamma_add_rms_norm_regbase_split_d.h" + +using namespace AscendC; +using namespace GammaAddRmsNorm; + +extern "C" __global__ __aicore__ void gamma_add_rms_norm( + GM_ADDR x1, GM_ADDR x2, GM_ADDR gamma, GM_ADDR y, GM_ADDR rstd, GM_ADDR x, GM_ADDR workspace, GM_ADDR tiling) +{ + KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY); + TPipe aptPipe; + if (TILING_KEY_IS(1000)) { + GET_TILING_DATA_WITH_STRUCT(GammaAddRMSNormRegbaseRFullLoadTilingData, aptTilingDataIn, tiling); + KernelGammaAddRmsNormRegBase op(&aptPipe); + op.Init(x1, x2, gamma, y, rstd, x, &aptTilingDataIn); + op.Process(); + } else if (TILING_KEY_IS(2000)) { + GET_TILING_DATA_WITH_STRUCT(GammaAddRMSNormRegbaseTilingData, aptTilingDataIn, tiling); + KernelGammaAddRmsNormRegBaseSplitD op(&aptPipe); + op.Init(x1, x2, gamma, y, rstd, x, &aptTilingDataIn); + op.Process(); + } +} diff --git a/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_base.h b/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_base.h new file mode 100644 index 0000000..cd71ca6 --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_base.h @@ -0,0 +1,321 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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. + */ + +#ifndef GAMMA_ADD_RMS_NORM_BASE_H_ +#define GAMMA_ADD_RMS_NORM_BASE_H_ +#include "kernel_operator.h" +#include "reduce_common.h" + +namespace RmsNorm { +using namespace AscendC; + + +/** + * Get the block size of unified buffer in bytes + */ +__aicore__ inline constexpr uint32_t GetUbBlockSize() +{ + return 32U; +} + +/** + * Get the size of vector registers in bytes + */ +__aicore__ inline constexpr uint32_t GetVRegSize() +{ +#if __CCE_AICORE__ == 310 + return AscendC::VECTOR_REG_WIDTH; +#else + return 256U; +#endif +} + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ != 220 && __CCE_AICORE__ != 220 && __CCE_AICORE__ != 310 +#define bfloat16_t int16_t +#endif +constexpr int32_t BUFFER_NUM = 1; // tensor num for each queue +constexpr int32_t NUM_PER_REP_FP32 = 64; // ONE_REPEAT_BYTE_SIZE / sizeof(float); +constexpr int32_t NUM_PER_BLK_FP32 = 8; +constexpr int32_t DOUBLE_BUFFER_NUM = 2; +constexpr int32_t UNROLL_NUM = 2; +constexpr int32_t FLOAT_BTYPE_SIZE = 4; +constexpr int32_t NUM_PER_BLK_FP16 = 16; +constexpr int32_t CONTINUE_STRIDE = 8; +constexpr int32_t BLOCK_SIZE = 32; +constexpr uint32_t ONCE_VECTOR_SIZE = 256; +constexpr float MINUS_HALF = -0.5f; +constexpr uint32_t ZERO_UINT = 0; +constexpr uint32_t ONE_UINT = 1; +constexpr uint32_t TWO_UINT = 2; +constexpr uint32_t THREE_UINT = 3; +constexpr float ONE = 1; +constexpr int32_t SECOND_LOOP = 2; +constexpr int32_t HALf_INTERVAL = 2; +constexpr int32_t MAX_REAPEAT = 255; +constexpr int32_t DIM_NUM = 2; +constexpr int32_t NDDMA_DIM = 5; + +constexpr uint32_t V_LENGTH = GetVRegSize() / sizeof(float); +constexpr uint64_t ALIGN_512_FACTOR = 512; +constexpr uint64_t ALIGN_32_FACTOR = 32; +constexpr int32_t CONST_FACTOR_2 = 2; +constexpr uint32_t SUM_COUNT = 2; +constexpr int32_t MOV_2 = 2; +constexpr int32_t MOV_4 = 4; +constexpr int32_t MOV_8 = 8; +constexpr int32_t MOV_16 = 16; +constexpr int32_t GAMMA_ADD_RMS_NORM_MODE = 1; +constexpr int32_t ADD_RMS_NORM_MODE = 1; +constexpr int32_t PRE_RMS_NORM_MODE = 2; +constexpr int32_t POST_RMS_NORM_MODE = 3; + +template +__aicore__ inline T CeilDiv(T x, T y) +{ + return y == 0 ? x : (x + y - 1) / y; +} + +template +__aicore__ inline T Min(T left, T right) +{ + return (left < right ? left : right); +} + +template +struct integral_constant { + static constexpr Tp value = v; +}; +using true_type = integral_constant; +using false_type = integral_constant; +template +struct is_same : public false_type {}; +template +struct is_same : public true_type {}; + +template +class KernelRmsNormBase { +#define IS_X_FP32 (is_same::value) +#define IS_GAMMA_FP32 (is_same::value) +#define IS_MIX_DTYPE ((!IS_X_FP32) && IS_GAMMA_FP32) +}; + +__aicore__ inline int32_t findPowerTwo(int32_t n) +{ + // find max power of 2 no more than n (32 bit) + n |= n >> 1; // Set the first digit of n's binary to 1 + n |= n >> MOV_2; + n |= n >> MOV_4; + n |= n >> MOV_8; + n |= n >> MOV_16; + return (n + 1) >> 1; +} + +__aicore__ inline void ReduceSumHalfIntervalToRepeat( + const LocalTensor& dst_local, const LocalTensor& src_local, int32_t count, int32_t left) +{ + // count need smaller than 255 repeat + if (likely(count > NUM_PER_BLK_FP32)) { + int32_t bodyCount = count - left; + int32_t tailCount = left; + if (tailCount > 0) { + Add(src_local, src_local, src_local[bodyCount], tailCount); + PipeBarrier(); + } + while (bodyCount > SECOND_LOOP * NUM_PER_BLK_FP32) { + bodyCount = bodyCount / HALf_INTERVAL; + Add(src_local, src_local, src_local[bodyCount], bodyCount); + PipeBarrier(); + } + bodyCount = bodyCount / HALf_INTERVAL; + Add(dst_local, src_local, src_local[bodyCount], bodyCount); + PipeBarrier(); + } +} + +__aicore__ inline void ReduceSumFP32( + const LocalTensor& dst_local, const LocalTensor& src_local, const LocalTensor& work_local, + int32_t count) +{ + // count need smaller than 255 repeat + uint64_t mask = NUM_PER_REP_FP32; + int32_t repeatTimes = count / NUM_PER_REP_FP32; + int32_t tailCount = count % NUM_PER_REP_FP32; + int32_t bodyCount = repeatTimes * NUM_PER_REP_FP32; + BinaryRepeatParams repeatParams; + repeatParams.src0RepStride = ONE_REPEAT_BYTE_SIZE / ONE_BLK_SIZE; + repeatParams.src0BlkStride = 1; + repeatParams.src1RepStride = 0; + repeatParams.src1BlkStride = 1; + repeatParams.dstRepStride = 0; + repeatParams.dstBlkStride = 1; + Duplicate(work_local, ZERO, NUM_PER_REP_FP32); + PipeBarrier(); + if (likely(repeatTimes > 0)) { + Add(work_local, src_local, work_local, mask, repeatTimes, repeatParams); + PipeBarrier(); + } + if (unlikely(tailCount != 0)) { + Add(work_local, src_local[bodyCount], work_local, tailCount, 1, repeatParams); + PipeBarrier(); + } + AscendCUtils::SetMask(NUM_PER_REP_FP32); +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 220 + if (g_coreType == AIV) { + WholeReduceSum(dst_local, work_local, MASK_PLACEHOLDER, 1, 0, 1, 0); + } +#elif !(defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113)) + WholeReduceSum(dst_local, work_local, MASK_PLACEHOLDER, 1, 1, 1, DEFAULT_REPEAT_STRIDE); +#endif + PipeBarrier(); +} + +__aicore__ inline void ReduceSumCustom( + const LocalTensor& dst_local, const LocalTensor& src_local, const LocalTensor& work_local_v1, + int32_t count) +{ + ReduceSumFP32(dst_local, src_local, work_local_v1, count); +} +__aicore__ inline void ReduceSumFP32ToBlock( + const LocalTensor& dst_local, const LocalTensor& src_local, const LocalTensor& work_local_v1, + int32_t count) +{ + // count need smaller than 255 repeat + uint64_t mask = NUM_PER_REP_FP32; + int32_t repeatTimes = count / NUM_PER_REP_FP32; + int32_t tailCount = count % NUM_PER_REP_FP32; + int32_t bodyCount = repeatTimes * NUM_PER_REP_FP32; + BinaryRepeatParams repeatParams; + repeatParams.src0RepStride = ONCE_VECTOR_SIZE / BLOCK_SIZE; + repeatParams.dstBlkStride = 1; + repeatParams.src0BlkStride = 1; + repeatParams.src1RepStride = 0; + repeatParams.src1BlkStride = 1; + repeatParams.dstRepStride = 0; + Duplicate(work_local_v1, ZERO, NUM_PER_REP_FP32); + PipeBarrier(); + if (likely(repeatTimes > 0)) { + Add(work_local_v1, src_local, work_local_v1, mask, repeatTimes, repeatParams); + PipeBarrier(); + } + if (unlikely(tailCount != 0)) { + Add(work_local_v1, src_local[bodyCount], work_local_v1, tailCount, 1, repeatParams); + PipeBarrier(); + } + BlockReduceSum(dst_local, work_local_v1, 1, mask, 1, 1, DEFAULT_REPEAT_STRIDE); + PipeBarrier(); +} + +__aicore__ inline void BlockReduceSumFP32( + const LocalTensor& dst_local_v1, const LocalTensor& src_local, int32_t count) +{ + // count need multiple of 8 + int32_t repeatTimes = count / NUM_PER_REP_FP32; + int32_t tailCount = count % NUM_PER_REP_FP32; + int32_t dstAddr = repeatTimes * 8; + int32_t srcAddr = repeatTimes * NUM_PER_REP_FP32; + if (likely(repeatTimes > 0)) { + BlockReduceSum(dst_local_v1, src_local, repeatTimes, NUM_PER_REP_FP32, 1, 1, DEFAULT_REPEAT_STRIDE); + PipeBarrier(); + } + if (tailCount != 0) { + BlockReduceSum(dst_local_v1[dstAddr], src_local[srcAddr], 1, tailCount, 1, 1, DEFAULT_REPEAT_STRIDE); + PipeBarrier(); + } +} + +template +__aicore__ inline void DataCopyCustom(const U& dstTensorV1, const R& srcTensor, const uint32_t count) +{ +#if (defined(__CCE_AICORE__) && __CCE_AICORE__ == 220) || (defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113)) + DataCopyParams copyParams; + copyParams.blockLen = count * sizeof(T); + copyParams.blockCount = 1; + if constexpr (is_same>::value) { + DataCopyPadParams padParams; + DataCopyPad(dstTensorV1, srcTensor, copyParams, padParams); + } else { + DataCopyPad(dstTensorV1, srcTensor, copyParams); + } +#else + // only support count greater than 32byte + int32_t numPerBlock = ONE_BLK_SIZE / sizeof(T); + if (count % numPerBlock == 0) { + DataCopy(dstTensorV1, srcTensor, count); + } else { + if constexpr (is_same>::value) { + int32_t num = AlignUp(count, numPerBlock); + DataCopy(dstTensorV1, srcTensor, num); + } else { + if (count < numPerBlock) { + DataCopy(dstTensorV1, srcTensor, numPerBlock); + } else { + int32_t num = count / numPerBlock * numPerBlock; + DataCopy(dstTensorV1, srcTensor, num); + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + for (int32_t i = 0; i < numPerBlock; i++) { + T tensorValue = srcTensor.GetValue(count - numPerBlock + i); + srcTensor.SetValue(i, tensorValue); + } + SetFlag(EVENT_ID0); + WaitFlag(EVENT_ID0); + DataCopy(dstTensorV1[count - numPerBlock], srcTensor, numPerBlock); + } + } + } +#endif +} + +template +__aicore__ inline void DataCopyCustom( + const LocalTensor& dstTensor, const GlobalTensor& srcTensor, const uint32_t numRow, const uint32_t numCol) +{ +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 220 + DataCopyParams copyParams; + copyParams.blockLen = numCol * sizeof(T); + copyParams.blockCount = numRow; + DataCopyPadParams padParams; + DataCopyPad(dstTensor, srcTensor, copyParams, padParams); +#endif +} + +template +__aicore__ inline void DataCopyCustom( + const GlobalTensor& dstTensor, const LocalTensor& srcTensor, const uint32_t numRow, const uint32_t numCol) +{ +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 220 + DataCopyParams copyParams; + copyParams.blockLen = numCol * sizeof(T); + copyParams.blockCount = numRow; + DataCopyPad(dstTensor, srcTensor, copyParams); +#endif +} + +__aicore__ inline void RoundFloat2Int8(LocalTensor& dstTensorV1, LocalTensor& srcTensor, int32_t size) +{ + Cast(srcTensor.ReinterpretCast(), srcTensor, RoundMode::CAST_RINT, size); + PipeBarrier(); + SetDeqScale((half)1.000000e+00f); + PipeBarrier(); + Cast(srcTensor.ReinterpretCast(), srcTensor.ReinterpretCast(), RoundMode::CAST_NONE, size); + PipeBarrier(); + Cast(dstTensorV1, srcTensor.ReinterpretCast(), RoundMode::CAST_TRUNC, size); +} + +__aicore__ inline uint32_t ROUND_UP(uint32_t x, uint32_t block_number) +{ + if (block_number > 0) { + return (x + block_number - 1) / block_number * block_number; + } + return 0; +} +} // namespace RmsNorm +#endif // GAMMA_ADD_RMS_NORM_BASE_H_ diff --git a/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_merge_n.h b/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_merge_n.h new file mode 100644 index 0000000..c026963 --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_merge_n.h @@ -0,0 +1,429 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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 gamma_add_rms_norm_merge_n.h + * \brief add rms norm merge n file + */ +#ifndef GAMMA_ADD_RMS_NORM_MERGE_N_H_ +#define GAMMA_ADD_RMS_NORM_MERGE_N_H_ +#include "gamma_add_rms_norm_base.h" + +using namespace AscendC; +using namespace RmsNorm; + +template +class KernelGammaAddRmsNormMergeN { +public: + __aicore__ inline KernelGammaAddRmsNormMergeN(TPipe* pipe) + { + Ppipe = pipe; + } + + __aicore__ inline void InitParams(const GammaAddRMSNormTilingData* tiling) + { + this->numRow = tiling->num_row; + this->numCol = tiling->num_col; + this->numColAlign = tiling->num_col_align; + this->blockFactor = tiling->block_factor; + this->rowFactor = tiling->row_factor; + this->ubFactor = tiling->ub_factor; + this->epsilon = tiling->epsilon; + this->avgFactor = tiling->avg_factor; + this->addGammaOffset = tiling->add_gamma_offset; + + blockIdx_ = GetBlockIdx(); + if (blockIdx_ < GetBlockNum() - 1) { + this->rowWork = blockFactor; + this->rowLoop = tiling->row_loop; + this->rowTail = tiling->row_tail; + } else if (blockIdx_ == GetBlockNum() - 1) { + this->rowWork = tiling->last_block_factor; + this->rowLoop = tiling->last_block_row_loop; + this->rowTail = tiling->last_block_row_tail; + } + this->mulLoopFp32 = tiling->mul_loop_fp32; + this->mulTailFp32 = tiling->mul_tail_fp32; + this->dstRepStrideFp32 = tiling->dst_rep_stride_fp32; + this->mulLoopFp16 = tiling->mul_loop_fp16; + this->mulTailFp16 = tiling->mul_tail_fp16; + this->dstRepStrideFp16 = tiling->dst_rep_stride_fp16; + this->isPerformance = tiling->is_performance; + } + + __aicore__ inline void Init( + GM_ADDR x1, GM_ADDR x2, GM_ADDR gamma, GM_ADDR y, GM_ADDR rstd, GM_ADDR x, GM_ADDR workspace, const GammaAddRMSNormTilingData* tiling) + { + ASSERT(GetBlockNum() != 0 && "Block dim can not be zero!"); + this->InitParams(tiling); + // get start index for current core, core parallel + x1Gm.SetGlobalBuffer((__gm__ T*)x1 + blockIdx_ * this->blockFactor * this->numCol, this->rowWork * this->numCol); + x2Gm.SetGlobalBuffer((__gm__ T*)x2 + blockIdx_ * this->blockFactor * this->numCol, this->rowWork * this->numCol); + gammaGm.SetGlobalBuffer((__gm__ T*)gamma, this->numCol); + yGm.SetGlobalBuffer((__gm__ T*)y + blockIdx_ * this->blockFactor * this->numCol, this->rowWork * this->numCol); + if constexpr (MODE == GAMMA_ADD_RMS_NORM_MODE) { + rstdGm.SetGlobalBuffer((__gm__ float*)rstd + blockIdx_ * this->blockFactor, this->blockFactor); + xGm.SetGlobalBuffer((__gm__ T*)x + blockIdx_ * this->blockFactor * this->numCol, this->rowWork * this->numCol); + } + if constexpr (MODE == PRE_RMS_NORM_MODE) { + xGm.SetGlobalBuffer((__gm__ T*)x + blockIdx_ * this->blockFactor * this->numCol, this->rowWork * this->numCol); + } + + // pipe alloc memory to queue, the unit is Bytes + Ppipe->InitBuffer(inQueueX, DOUBLE_BUFFER_NUM, this->ubFactor * sizeof(T)); + Ppipe->InitBuffer(inQueueGamma, BUFFER_NUM, this->ubFactor * sizeof(T)); + Ppipe->InitBuffer(outQueueY, DOUBLE_BUFFER_NUM, ubFactor * sizeof(T)); +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 220 || (defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113)) + Ppipe->InitBuffer(outQueueRstd, BUFFER_NUM, rowFactor * sizeof(float)); +#else + Ppipe->InitBuffer(rstdBuf, rowFactor * sizeof(float)); +#endif + if constexpr (is_same::value || is_same::value) { + Ppipe->InitBuffer(xFp32Buf, ubFactor * sizeof(float)); + } + Ppipe->InitBuffer(sqxBuf, ubFactor * sizeof(float)); + Ppipe->InitBuffer(tmpBuf, rowFactor * NUM_PER_REP_FP32 * sizeof(float)); + } + + __aicore__ inline void Process() + { + CopyInGamma(); + LocalTensor gammaLocal = inQueueGamma.DeQue(); + if (addGammaOffset != 0U) { + AddGammaOffset(gammaLocal, numCol); + } + for (uint32_t i_o = 0; i_o < rowLoop - 1; i_o++) { + MainCompute(i_o, rowFactor, gammaLocal); + } + MainCompute(rowLoop - 1, rowTail, gammaLocal); + inQueueGamma.FreeTensor(gammaLocal); + } + + __aicore__ inline void MainCompute(uint32_t i_o, uint32_t calc_row_num, LocalTensor& gammaLocal) + { + uint64_t gm_bias = static_cast(i_o) * static_cast(rowFactor) * static_cast(numCol); + uint32_t elementNum = calc_row_num * numColAlign; + CopyInX(gm_bias, calc_row_num); + LocalTensor xLocal = ComputeX(elementNum); + if constexpr (MODE == GAMMA_ADD_RMS_NORM_MODE || MODE == PRE_RMS_NORM_MODE) { + CopyOutX(gm_bias, calc_row_num); + } else { + outQueueY.FreeTensor(xLocal); + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 220 || (defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113)) + LocalTensor rstdLocal = outQueueRstd.AllocTensor(); + ComputeRstd(xLocal, rstdLocal, calc_row_num, elementNum); + if constexpr (MODE == GAMMA_ADD_RMS_NORM_MODE) { + outQueueRstd.EnQue(rstdLocal); + CopyOutRstd(i_o, calc_row_num); + } else { + outQueueRstd.FreeTensor(rstdLocal); + } +#else + LocalTensor rstdLocal = rstdBuf.Get(); + ComputeRstd(xLocal, rstdLocal, calc_row_num, elementNum); +#endif + ComputeY(xLocal, gammaLocal, rstdLocal, calc_row_num, elementNum); + CopyOutY(gm_bias, calc_row_num); + } + +private: + __aicore__ inline void CopyInX(uint32_t gm_bias, uint32_t calc_row_num) + { + LocalTensor x1Local = inQueueX.AllocTensor(); + if (isNumColAlign) { + DataCopyCustom(x1Local, x1Gm[gm_bias], calc_row_num * numCol); + } else { + DataCopyCustom(x1Local, x1Gm[gm_bias], calc_row_num, numCol); + } + inQueueX.EnQue(x1Local); + LocalTensor x2Local = inQueueX.AllocTensor(); + if (isNumColAlign) { + DataCopyCustom(x2Local, x2Gm[gm_bias], calc_row_num * numCol); + } else { + DataCopyCustom(x2Local, x2Gm[gm_bias], calc_row_num, numCol); + } + inQueueX.EnQue(x2Local); + } + + __aicore__ inline LocalTensor ComputeX(uint32_t elementNum) + { + LocalTensor x1Local = inQueueX.DeQue(); + LocalTensor x2Local = inQueueX.DeQue(); + LocalTensor xLocal = outQueueY.AllocTensor(); + if constexpr (!is_same::value) { + Add(xLocal, x1Local, x2Local, elementNum); + } else { + LocalTensor x1Fp32 = xFp32Buf.Get(); + LocalTensor x2Fp32 = sqxBuf.Get(); + Cast(x1Fp32, x1Local, RoundMode::CAST_NONE, elementNum); + Cast(x2Fp32, x2Local, RoundMode::CAST_NONE, elementNum); + PipeBarrier(); + Add(x1Fp32, x1Fp32, x2Fp32, elementNum); + PipeBarrier(); + Cast(xLocal, x1Fp32, RoundMode::CAST_RINT, elementNum); + } + PipeBarrier(); + inQueueX.FreeTensor(x1Local); + inQueueX.FreeTensor(x2Local); + if constexpr (MODE == GAMMA_ADD_RMS_NORM_MODE || MODE == PRE_RMS_NORM_MODE) { + outQueueY.EnQue(xLocal); + } + return xLocal; + } + + __aicore__ inline void CopyOutX(uint32_t gm_bias, uint32_t calc_row_num) + { + // CopyOut x1 + x2 + auto xOut = outQueueY.DeQue(); + event_t eventVMTE3_BF16_0 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_MTE3)); + SetFlag(eventVMTE3_BF16_0); + WaitFlag(eventVMTE3_BF16_0); + if (isNumColAlign) { + DataCopyCustom(xGm[gm_bias], xOut, calc_row_num * numCol); + } else { + DataCopyCustom(xGm[gm_bias], xOut, calc_row_num, numCol); + } + event_t eventMTE3V_BF16_0 = static_cast(GetTPipePtr()->AllocEventID()); + SetFlag(eventMTE3V_BF16_0); + WaitFlag(eventMTE3V_BF16_0); + GetTPipePtr()->ReleaseEventID(eventMTE3V_BF16_0); + outQueueY.FreeTensor(xOut); + } + + __aicore__ inline void CopyInGamma() + { + LocalTensor gammaLocal = inQueueGamma.AllocTensor(); + DataCopyCustom(gammaLocal, gammaGm, numCol); + inQueueGamma.EnQue(gammaLocal); + } + + __aicore__ inline void AddGammaOffset(LocalTensor& gammaLocal, uint32_t elementNum) + { + if constexpr (is_same::value) { + LocalTensor gammaFp32 = xFp32Buf.Get(); + Cast(gammaFp32, gammaLocal, RoundMode::CAST_NONE, elementNum); + PipeBarrier(); + Adds(gammaFp32, gammaFp32, static_cast(1.0), elementNum); + PipeBarrier(); + Cast(gammaLocal, gammaFp32, RoundMode::CAST_RINT, elementNum); + } else { + Adds(gammaLocal, gammaLocal, static_cast(1.0), elementNum); + } + PipeBarrier(); + } + + __aicore__ inline void ComputeRstd(LocalTensor xLocal, LocalTensor rstdLocal, uint32_t calc_row_num, uint32_t elementNum) + { + LocalTensor sqx = sqxBuf.Get(); + LocalTensor tmpLocal = tmpBuf.Get(); + if constexpr (!is_same::value) { + LocalTensor x_fp32 = xFp32Buf.Get(); + Cast(x_fp32, xLocal, RoundMode::CAST_NONE, elementNum); + PipeBarrier(); + Mul(sqx, x_fp32, x_fp32, elementNum); + } else { + Mul(sqx, xLocal, xLocal, elementNum); + } + PipeBarrier(); + + Muls(sqx, sqx, avgFactor, elementNum); + PipeBarrier(); + + ReduceSumMultiN(rstdLocal, sqx, tmpLocal, calc_row_num, numCol, numColAlign); + PipeBarrier(); + Adds(rstdLocal, rstdLocal, epsilon, calc_row_num); + PipeBarrier(); + + Sqrt(rstdLocal, rstdLocal, calc_row_num); + Duplicate(tmpLocal, ONE, calc_row_num); + PipeBarrier(); + + Div(rstdLocal, tmpLocal, rstdLocal, calc_row_num); + PipeBarrier(); + } + + __aicore__ inline void ComputeY( + LocalTensor xLocal, LocalTensor gammaLocal, LocalTensor rstdLocal, uint32_t calc_row_num, uint32_t elementNum) + { + LocalTensor tmpLocal = tmpBuf.Get(); + uint32_t splidRow = 240; + uint32_t rowRepeatLoop1 = calc_row_num / splidRow; + uint32_t rowRepeatTail1 = calc_row_num - rowRepeatLoop1 * splidRow; + for(uint32_t r_i = 0; r_i < rowRepeatLoop1; r_i ++) { + Brcb(tmpLocal[r_i * splidRow * MOV_8], rstdLocal[r_i * splidRow], splidRow, {1, 8}); + } + PipeBarrier(); + + if(rowRepeatTail1 > 0) { + Brcb(tmpLocal[rowRepeatLoop1 * splidRow * MOV_8], rstdLocal[rowRepeatLoop1 * splidRow], rowRepeatTail1, {1, 8}); + PipeBarrier(); + } + LocalTensor yLocal = outQueueY.AllocTensor(); + if constexpr (!is_same::value) { + LocalTensor x_fp32 = xFp32Buf.Get(); + repeatByRow(x_fp32, x_fp32, tmpLocal, calc_row_num, ONE_UINT); + if constexpr (is_same::value) { + Cast(yLocal, x_fp32, RoundMode::CAST_NONE, elementNum); + } else { + Cast(yLocal, x_fp32, RoundMode::CAST_RINT, elementNum); + } + } else { + repeatByRow(yLocal, xLocal, tmpLocal, calc_row_num, ONE_UINT); + } + PipeBarrier(); + if constexpr (is_same::value) { + repeatByRow(yLocal, yLocal, gammaLocal, calc_row_num, TWO_UINT); + } else if constexpr (is_same::value) { + // Cast BF16 values to FP32 before multiplication. + LocalTensor sqx = sqxBuf.Get(); + LocalTensor x_fp32 = xFp32Buf.Get(); + Cast(x_fp32, yLocal, RoundMode::CAST_NONE, elementNum); + Cast(sqx, gammaLocal, RoundMode::CAST_NONE, elementNum); + PipeBarrier(); + repeatByRow(x_fp32, x_fp32, sqx, calc_row_num, THREE_UINT); + Cast(yLocal, x_fp32, RoundMode::CAST_RINT, elementNum); + } else { + repeatByRow(yLocal, yLocal, gammaLocal, calc_row_num, THREE_UINT); + } + PipeBarrier(); + outQueueY.EnQue(yLocal); + } + + __aicore__ inline void CopyOutY(uint32_t progress, uint32_t calc_row_num) + { + LocalTensor yLocal = outQueueY.DeQue(); + if (isNumColAlign) { + DataCopyCustom(yGm[progress], yLocal, calc_row_num * numCol); + } else { + DataCopyCustom(yGm[progress], yLocal, calc_row_num, numCol); + } + outQueueY.FreeTensor(yLocal); + } + +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 220 || (defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113)) + __aicore__ inline void CopyOutRstd(uint32_t outer_progress, uint32_t num) + { + LocalTensor rstdLocal = outQueueRstd.DeQue(); + DataCopyCustom(rstdGm[outer_progress * rowFactor], rstdLocal, num); + outQueueRstd.FreeTensor(rstdLocal); + } +#endif + + template + __aicore__ inline void repeatByRow(const LocalTensor& dstLocal, const LocalTensor& src1Local, const LocalTensor& src2Local, uint32_t calc_row_num, uint32_t type) + { + // TWO_UINT=gammaFp16 ONE_UINT=rstd + uint32_t strideParams[6] = {mulLoopFp32, mulTailFp32, 64, 1, dstRepStrideFp32, 0}; + if (type == TWO_UINT) { + strideParams[0] = mulLoopFp16; + strideParams[1] = mulTailFp16; + strideParams[2] = 128; + strideParams[4] = dstRepStrideFp16; + } else if (type == ONE_UINT) { + strideParams[3] = 0; + strideParams[5] = 1; + } + uint32_t singlT = 255; + uint32_t rowRepeatLoop = calc_row_num / singlT; + uint32_t rowRepeatTail = calc_row_num - rowRepeatLoop * singlT; + uint32_t offset2 = 0; + for(uint32_t r_i = 0; r_i < rowRepeatLoop; r_i ++) { + offset2 = type == 1 ? (r_i * singlT * MOV_8) : 0; + mulRepeat(dstLocal[r_i * singlT * numColAlign], src1Local[r_i * singlT * numColAlign], src2Local[offset2], singlT, strideParams); + } + if(rowRepeatTail > 0) { + offset2 = type == 1 ? (rowRepeatLoop * singlT * MOV_8) : 0; + uint32_t offset1 = rowRepeatLoop * singlT * numColAlign; + mulRepeat(dstLocal[offset1], src1Local[offset1], src2Local[offset2], rowRepeatTail, strideParams); + } + } + + template + __aicore__ inline void mulRepeat(const LocalTensor& dstLocal, const LocalTensor& src1Local, const LocalTensor& src2Local, uint32_t calcRowNum, uint32_t strideParams[6]) + { + uint32_t mulLoop = strideParams[0]; + uint32_t mulTail = strideParams[1]; + uint32_t strideNum = strideParams[2]; + uint8_t src1BlkStride = static_cast(strideParams[3]); + uint8_t dstRepStride = static_cast(strideParams[4]); + uint8_t src1RepStride = static_cast(strideParams[5]); + if(src1BlkStride == 0) { + for (uint32_t m_i = 0; m_i < mulLoop; m_i++) { + Mul(dstLocal[m_i * strideNum], src1Local[m_i * strideNum], src2Local, strideNum, calcRowNum, {1, 1, src1BlkStride, dstRepStride, dstRepStride, src1RepStride}); + } + PipeBarrier(); + if(mulTail > 0) { + Mul(dstLocal[mulLoop * strideNum], src1Local[mulLoop * strideNum], src2Local, mulTail, calcRowNum, {1, 1, src1BlkStride, dstRepStride, dstRepStride, src1RepStride}); + } + PipeBarrier(); + } else { + for (uint32_t m_i = 0; m_i < mulLoop; m_i++) { + Mul(dstLocal[m_i * strideNum], src1Local[m_i * strideNum], src2Local[m_i * strideNum], strideNum, calcRowNum, {1, 1, src1BlkStride, dstRepStride, dstRepStride, src1RepStride}); + } + PipeBarrier(); + if(mulTail > 0) { + Mul(dstLocal[mulLoop * strideNum], src1Local[mulLoop * strideNum], src2Local[mulLoop * strideNum], mulTail, calcRowNum, {1, 1, src1BlkStride, dstRepStride, dstRepStride, src1RepStride}); + } + PipeBarrier(); + } + } + +private: + TPipe* Ppipe = nullptr; + // create queues for input, in this case depth is equal to buffer num + TQue inQueueGamma; + TQue inQueueX; + // create queues for output, in this case depth is equal to buffer num +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 220 || (defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113)) + TQue outQueueRstd; +#else + TBuf rstdBuf; +#endif + TQue outQueueY; + + TBuf xFp32Buf; + TBuf sqxBuf; + TBuf tmpBuf; + GlobalTensor x1Gm; + GlobalTensor x2Gm; + GlobalTensor gammaGm; + GlobalTensor yGm; + GlobalTensor rstdGm; + GlobalTensor xGm; + + uint32_t numRow; + uint32_t numCol; + uint32_t numColAlign; + uint32_t blockFactor; + uint32_t rowFactor; + uint32_t ubFactor; + float epsilon; + float avgFactor; + int32_t blockIdx_; + uint32_t rowWork = 1; +#if (defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113)) + bool isNumColAlign = true; +#else + bool isNumColAlign = false; +#endif + uint8_t isPerformance = 0; + uint32_t rowLoop = 1; + uint32_t rowTail = 0; + uint32_t mulLoopFp32; + uint32_t mulTailFp32; + uint8_t dstRepStrideFp32; + uint32_t mulLoopFp16; + uint32_t mulTailFp16; + uint8_t dstRepStrideFp16; + uint32_t addGammaOffset = 0; +}; +#endif // _GAMMA_ADD_RMS_NORM_MERGE_N_H_ diff --git a/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_multi_n.h b/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_multi_n.h new file mode 100644 index 0000000..8b8a0a0 --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_multi_n.h @@ -0,0 +1,339 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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 gamma_add_rms_norm_multi_n.h + * \brief add rms norm multi n file + */ +#ifndef GAMMA_ADD_RMS_NORM_MULTI_N_H_ +#define GAMMA_ADD_RMS_NORM_MULTI_N_H_ +#include "gamma_add_rms_norm_base.h" + +using namespace AscendC; +using namespace RmsNorm; + +template +class KernelGammaAddRmsNormMultiN { +public: + __aicore__ inline KernelGammaAddRmsNormMultiN(TPipe* pipe) + { + Ppipe = pipe; + } + __aicore__ inline void Init( + GM_ADDR x1, GM_ADDR x2, GM_ADDR gamma, GM_ADDR y, GM_ADDR rstd, GM_ADDR x, GM_ADDR workspace, const GammaAddRMSNormTilingData* tiling) + { + ASSERT(GetBlockNum() != 0 && "Block dim can not be zero!"); + this->numRow = tiling->num_row; + this->numCol = tiling->num_col; + this->numColAlign = tiling->num_col_align; + this->blockFactor = tiling->block_factor; + this->rowFactor = tiling->row_factor; + this->ubFactor = tiling->ub_factor; + this->epsilon = tiling->epsilon; + this->avgFactor = tiling->avg_factor; + this->addGammaOffset = tiling->add_gamma_offset; + + blockIdx_ = GetBlockIdx(); + if (blockIdx_ < GetBlockNum() - 1) { + this->rowWork = this->blockFactor; + this->rowLoop = tiling->row_loop; + this->rowTail = tiling->row_tail; + } else if (blockIdx_ == GetBlockNum() - 1) { + this->rowWork = tiling->last_block_factor; + this->rowLoop = tiling->last_block_row_loop; + this->rowTail = tiling->last_block_row_tail; + } + // get start index for current core, core parallel + x1Gm.SetGlobalBuffer((__gm__ T*)x1 + blockIdx_ * this->blockFactor * this->numCol, this->rowWork * this->numCol); + x2Gm.SetGlobalBuffer((__gm__ T*)x2 + blockIdx_ * this->blockFactor * this->numCol, this->rowWork * this->numCol); + gammaGm.SetGlobalBuffer((__gm__ T*)gamma, this->numCol); + yGm.SetGlobalBuffer((__gm__ T*)y + blockIdx_ * this->blockFactor * this->numCol, this->rowWork * this->numCol); + if constexpr (MODE == GAMMA_ADD_RMS_NORM_MODE) { + rstdGm.SetGlobalBuffer((__gm__ float*)rstd + blockIdx_ * blockFactor, blockFactor); + xGm.SetGlobalBuffer((__gm__ T*)x + blockIdx_ * blockFactor * numCol, rowWork * numCol); + } + if constexpr (MODE == PRE_RMS_NORM_MODE) { + xGm.SetGlobalBuffer((__gm__ T*)x + blockIdx_ * blockFactor * numCol, rowWork * numCol); + } + + // pipe alloc memory to queue, the unit is Bytes + Ppipe->InitBuffer(inQueueX, DOUBLE_BUFFER_NUM, this->ubFactor * sizeof(T)); + Ppipe->InitBuffer(inQueueGamma, BUFFER_NUM, this->numColAlign * sizeof(T)); + Ppipe->InitBuffer(outQueueY, DOUBLE_BUFFER_NUM, this->ubFactor * sizeof(T)); +#if __CCE_AICORE__ == 220 || (defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113)) + Ppipe->InitBuffer(outQueueRstd, BUFFER_NUM, this->rowFactor * NUM_PER_BLK_FP32 * sizeof(float)); +#else + Ppipe->InitBuffer(rstdBuf, this->rowFactor * NUM_PER_BLK_FP32 * sizeof(float)); +#endif + if constexpr (is_same::value || is_same::value) { + Ppipe->InitBuffer(xFp32Buf, this->ubFactor * sizeof(float)); + } + Ppipe->InitBuffer(sqxBuf, this->ubFactor * sizeof(float)); + Ppipe->InitBuffer(reduceFp32Buf, NUM_PER_REP_FP32 * sizeof(float)); + Ppipe->InitBuffer(offsetBuf, this->rowFactor * NUM_PER_BLK_FP32 * sizeof(uint32_t)); + } + __aicore__ inline void Process() + { + CopyInGamma(); + LocalTensor gammaLocal = inQueueGamma.DeQue(); + if (addGammaOffset != 0U) { + AddGammaOffset(gammaLocal, numCol); + } + LocalTensor offsetLocal = offsetBuf.Get(); + for (uint32_t i = 0; i < this->rowFactor; i++) { + Duplicate(offsetLocal[i * NUM_PER_BLK_FP32], i * ONE_BLK_SIZE, NUM_PER_BLK_FP32); + } + for (uint32_t i_o = 0; i_o < this->rowLoop - 1; i_o++) { + SubProcessHalf(i_o, this->rowFactor, gammaLocal); + } + SubProcessHalf(this->rowLoop - 1, this->rowTail, gammaLocal); + inQueueGamma.FreeTensor(gammaLocal); + } + + __aicore__ inline void SubProcessHalf(uint32_t i_o, uint32_t calc_row_num, LocalTensor& gammaLocal) + { + uint64_t gm_bias = static_cast(i_o) * static_cast(rowFactor) * static_cast(numCol); + CopyInX(gm_bias, calc_row_num); + LocalTensor xLocal = ComputeX(calc_row_num); + if constexpr (MODE == GAMMA_ADD_RMS_NORM_MODE || MODE == PRE_RMS_NORM_MODE) { + CopyOutX(gm_bias, calc_row_num); + } else { + outQueueY.FreeTensor(xLocal); + } +#if __CCE_AICORE__ == 220 || (defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113)) + LocalTensor rstdLocal = outQueueRstd.AllocTensor(); + ComputeRstd(xLocal, rstdLocal, calc_row_num); + if constexpr (MODE == GAMMA_ADD_RMS_NORM_MODE) { + outQueueRstd.EnQue(rstdLocal); + CopyOutRstd(i_o * rowFactor, calc_row_num); + } else { + outQueueRstd.FreeTensor(rstdLocal); + } +#else + LocalTensor rstdLocal = rstdBuf.Get(); + ComputeRstd(xLocal, rstdLocal, calc_row_num); +#endif + ComputeY(xLocal, gammaLocal, rstdLocal, calc_row_num); + CopyOutY(gm_bias, calc_row_num); + } + +private: + __aicore__ inline void CopyInX(uint32_t gm_bias, uint32_t calc_row_num) + { + LocalTensor x1Local = inQueueX.AllocTensor(); + DataCopyCustom(x1Local, x1Gm[gm_bias], calc_row_num * this->numCol); + inQueueX.EnQue(x1Local); + LocalTensor x2Local = inQueueX.AllocTensor(); + DataCopyCustom(x2Local, x2Gm[gm_bias], calc_row_num * this->numCol); + inQueueX.EnQue(x2Local); + } + + __aicore__ inline LocalTensor ComputeX(uint32_t calc_row_num) + { + uint32_t calc_num = calc_row_num * this->numColAlign; + LocalTensor x1Local = inQueueX.DeQue(); + LocalTensor x2Local = inQueueX.DeQue(); + LocalTensor xLocal = outQueueY.AllocTensor(); + if constexpr (!is_same::value) { + Add(xLocal, x1Local, x2Local, calc_num); + } else { + LocalTensor x1Fp32 = xFp32Buf.Get(); + LocalTensor x2Fp32 = sqxBuf.Get(); + Cast(x1Fp32, x1Local, RoundMode::CAST_NONE, calc_num); + Cast(x2Fp32, x2Local, RoundMode::CAST_NONE, calc_num); + PipeBarrier(); + Add(x1Fp32, x1Fp32, x2Fp32, calc_num); + PipeBarrier(); + Cast(xLocal, x1Fp32, RoundMode::CAST_RINT, calc_num); + } + inQueueX.FreeTensor(x1Local); + inQueueX.FreeTensor(x2Local); + if constexpr (MODE == PRE_RMS_NORM_MODE || MODE == GAMMA_ADD_RMS_NORM_MODE) { + outQueueY.EnQue(xLocal); + } + PipeBarrier(); + return xLocal; + } + + __aicore__ inline void CopyOutX(uint32_t gm_bias, uint32_t calc_row_num) + { + // CopyOut x1 + x2 + auto x_out = outQueueY.DeQue(); + DataCopyCustom(xGm[gm_bias], x_out, calc_row_num * numCol); + outQueueY.FreeTensor(x_out); + } + + __aicore__ inline void CopyInGamma() + { + LocalTensor gammaLocal = inQueueGamma.AllocTensor(); + DataCopyCustom(gammaLocal, gammaGm, numCol); + inQueueGamma.EnQue(gammaLocal); + } + + __aicore__ inline void AddGammaOffset(LocalTensor& gammaLocal, uint32_t elementNum) + { + if constexpr (is_same::value) { + LocalTensor gammaFp32 = xFp32Buf.Get(); + Cast(gammaFp32, gammaLocal, RoundMode::CAST_NONE, elementNum); + PipeBarrier(); + Adds(gammaFp32, gammaFp32, static_cast(1.0), elementNum); + PipeBarrier(); + Cast(gammaLocal, gammaFp32, RoundMode::CAST_RINT, elementNum); + } else { + Adds(gammaLocal, gammaLocal, static_cast(1.0), elementNum); + } + PipeBarrier(); + } + + __aicore__ inline void ComputeRstd(LocalTensor xLocal, LocalTensor rstdLocal, uint32_t calc_row_num) + { + LocalTensor x_fp32 = xFp32Buf.Get(); + LocalTensor sqx = sqxBuf.Get(); + LocalTensor reduce_buf_local = reduceFp32Buf.Get(); + Cast(x_fp32, xLocal, RoundMode::CAST_NONE, calc_row_num * numColAlign); + PipeBarrier(); + + Mul(sqx, x_fp32, x_fp32, calc_row_num * numColAlign); + PipeBarrier(); + + Muls(sqx, sqx, avgFactor, calc_row_num * numColAlign); + PipeBarrier(); + + for (uint32_t i_i = 0; i_i < calc_row_num; i_i++) { + ReduceSumCustom(rstdLocal[i_i * NUM_PER_BLK_FP32], sqx[i_i * numColAlign], reduce_buf_local, numCol); + } + Adds(rstdLocal, rstdLocal, epsilon, calc_row_num * NUM_PER_BLK_FP32); + PipeBarrier(); + + Sqrt(rstdLocal, rstdLocal, calc_row_num * NUM_PER_BLK_FP32); + Duplicate(reduce_buf_local, ONE, NUM_PER_BLK_FP32); + PipeBarrier(); + + int32_t repeatTimes = calc_row_num * NUM_PER_BLK_FP32 / NUM_PER_REP_FP32; + int32_t tailCount = calc_row_num * NUM_PER_BLK_FP32 % NUM_PER_REP_FP32; + int32_t bodyCount = repeatTimes * NUM_PER_REP_FP32; + + if (likely(repeatTimes > 0)) { + Div(rstdLocal, reduce_buf_local, rstdLocal, NUM_PER_REP_FP32, repeatTimes, {1, 0, 1, DEFAULT_REPEAT_STRIDE, 0, DEFAULT_REPEAT_STRIDE}); + } + if (unlikely(tailCount != 0)) { + Div(rstdLocal[bodyCount], reduce_buf_local, rstdLocal[bodyCount], tailCount, 1, {1, 0, 1, DEFAULT_REPEAT_STRIDE, 0, DEFAULT_REPEAT_STRIDE}); + } + PipeBarrier(); + } + + __aicore__ inline void ComputeY( + LocalTensor xLocal, LocalTensor gammaLocal, LocalTensor rstdLocal, uint32_t calc_row_num) + { + LocalTensor x_fp32 = xFp32Buf.Get(); + LocalTensor offsetLocal = offsetBuf.Get(); + Gather(rstdLocal, rstdLocal, offsetLocal, ZERO_UINT, calc_row_num * NUM_PER_BLK_FP32); + PipeBarrier(); + int32_t repeatTimes = numCol / NUM_PER_REP_FP32; + int32_t tailCount = numCol % NUM_PER_REP_FP32; + int32_t bodyCount = repeatTimes * NUM_PER_REP_FP32; + for (uint32_t i_i = 0; i_i < calc_row_num; i_i++) { + if (likely(repeatTimes > 0)) { + Mul(x_fp32[i_i * numColAlign], x_fp32[i_i * numColAlign], rstdLocal[i_i * NUM_PER_BLK_FP32], + NUM_PER_REP_FP32, repeatTimes, {1, 1, 0, DEFAULT_REPEAT_STRIDE, DEFAULT_REPEAT_STRIDE, 0}); + } + if (unlikely(tailCount != 0)) { + Mul(x_fp32[i_i * numColAlign + bodyCount], x_fp32[i_i * numColAlign + bodyCount], + rstdLocal[i_i * NUM_PER_BLK_FP32], tailCount, 1, + {1, 1, 0, DEFAULT_REPEAT_STRIDE, DEFAULT_REPEAT_STRIDE, 0}); + } + } + PipeBarrier(); + LocalTensor yLocal = outQueueY.AllocTensor(); + if constexpr (is_same::value) { + Cast(yLocal, x_fp32, RoundMode::CAST_NONE, calc_row_num * numColAlign); + PipeBarrier(); + + for (uint32_t i_i = 0; i_i < calc_row_num; i_i++) { + Mul(yLocal[i_i * numColAlign], gammaLocal, yLocal[i_i * numColAlign], numCol); + } + } else { + Cast(yLocal, x_fp32, RoundMode::CAST_RINT, calc_row_num * numColAlign); + PipeBarrier(); + LocalTensor yfp32 = xFp32Buf.Get(); + Cast(yfp32, yLocal, RoundMode::CAST_NONE, calc_row_num * numColAlign); + PipeBarrier(); + LocalTensor gammaFp32 = sqxBuf.Get(); + Cast(gammaFp32, gammaLocal, RoundMode::CAST_NONE, numCol); + PipeBarrier(); + for (uint32_t i_i = 0; i_i < calc_row_num; i_i++) { + Mul(yfp32[i_i * numColAlign], gammaFp32, yfp32[i_i * numColAlign], numCol); + } + PipeBarrier(); + Cast(yLocal, yfp32, RoundMode::CAST_RINT, calc_row_num * numColAlign); + } + PipeBarrier(); + outQueueY.EnQue(yLocal); + } + + __aicore__ inline void CopyOutY(uint32_t progress, uint32_t calc_row_num) + { + LocalTensor yLocal = outQueueY.DeQue(); + DataCopyCustom(yGm[progress], yLocal, calc_row_num * numCol); + outQueueY.FreeTensor(yLocal); + } + +#if __CCE_AICORE__ == 220 || (defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113)) + __aicore__ inline void CopyOutRstd(uint32_t outer_progress, uint32_t num) + { + LocalTensor rstdLocal = outQueueRstd.DeQue(); + DataCopyParams copyParams; + copyParams.blockLen = sizeof(float); + copyParams.blockCount = num; + DataCopyPad(rstdGm[outer_progress], rstdLocal, copyParams); + outQueueRstd.FreeTensor(rstdLocal); + } +#endif + +private: + TPipe* Ppipe = nullptr; + // create queues for input, in this case depth is equal to buffer num + TQue inQueueGamma; + TQue inQueueX; + // create queues for output, in this case depth is equal to buffer num +#if __CCE_AICORE__ == 220 || (defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113)) + TQue outQueueRstd; +#else + TBuf rstdBuf; +#endif + TQue outQueueY; + + TBuf xFp32Buf; + TBuf sqxBuf; + TBuf reduceFp32Buf; + TBuf offsetBuf; + GlobalTensor x1Gm; + GlobalTensor x2Gm; + GlobalTensor gammaGm; + GlobalTensor yGm; + GlobalTensor rstdGm; + GlobalTensor xGm; + + uint32_t numRow; + uint32_t numCol; + uint32_t blockFactor; // number of calculations rows on each core + uint32_t rowFactor; + uint32_t ubFactor; + float epsilon; + float avgFactor; + uint32_t numColAlign; + int32_t blockIdx_; + uint32_t rowWork = 1; + uint32_t rowLoop = 1; + uint32_t rowTail = 0; + uint32_t addGammaOffset = 0; +}; +#endif // GAMMA_ADD_RMS_NORM_H_ diff --git a/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_single_n.h b/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_single_n.h new file mode 100644 index 0000000..a44e031 --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_single_n.h @@ -0,0 +1,370 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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 gamma_add_rms_norm_single_n.h + * \brief add rms norm single n file + */ +#ifndef GAMMA_ADD_RMS_NORM_SINGLE_N_H_ +#define GAMMA_ADD_RMS_NORM_SINGLE_N_H_ +#include "gamma_add_rms_norm_base.h" + +using namespace AscendC; +using namespace RmsNorm; + +template +class KernelGammaAddRmsNormSingleN { + static constexpr int32_t MAXBUFFER = 195584; +public: + __aicore__ inline KernelGammaAddRmsNormSingleN(TPipe* pipe) + { + Ppipe = pipe; + } + __aicore__ inline void Init( + GM_ADDR x1, GM_ADDR x2, GM_ADDR gamma, GM_ADDR y, GM_ADDR rstd, GM_ADDR x, GM_ADDR workspace, const GammaAddRMSNormTilingData* tiling) + { + ASSERT(GetBlockNum() != 0 && "Block dim can not be zero!"); + + this->numCol = tiling->num_col; + this->blockFactor = 1; + this->ubFactor = tiling->ub_factor; + this->epsilon = tiling->epsilon; + this->avgFactor = (this->numCol != 0) ? (float)1.0 / this->numCol : 0; + this->addGammaOffset = tiling->add_gamma_offset; + + this->rowWork = 1; + blockIdx_ = GetBlockIdx(); + // get start index for current core, core parallel + x1Gm.SetGlobalBuffer((__gm__ T*)x1 + blockIdx_ * this->numCol, this->numCol); + x2Gm.SetGlobalBuffer((__gm__ T*)x2 + blockIdx_ * this->numCol, this->numCol); + gammaGm.SetGlobalBuffer((__gm__ T*)gamma, this->numCol); + yGm.SetGlobalBuffer((__gm__ T*)y + blockIdx_ * this->numCol, this->numCol); + + if constexpr (MODE == GAMMA_ADD_RMS_NORM_MODE) { + rstdGm.SetGlobalBuffer((__gm__ float*)rstd + blockIdx_, 1); + xGm.SetGlobalBuffer((__gm__ T*)x + blockIdx_ * this->numCol, this->numCol); + } + if constexpr (MODE == PRE_RMS_NORM_MODE) { + xGm.SetGlobalBuffer((__gm__ T*)x + blockIdx_ * numCol, numCol); + } + + Ppipe->InitBuffer(unitBuf, MAXBUFFER); // (192 - 1) * 1024 byte + } + + __aicore__ inline void Process() + { + if constexpr (is_same::value) { + ProcessFp16(); + } else if constexpr (is_same::value) { + ProcessFp32(); + } else { + ProcessBf16(); + } + } + +private: + __aicore__ inline void ProcessFp16() + { + LocalTensor ubLocal = unitBuf.Get(); + LocalTensor xLocal = ubLocal.template ReinterpretCast(); + LocalTensor x1Local = xLocal[0]; + LocalTensor x2Local = xLocal[ubFactor]; + LocalTensor xFp32Local = ubLocal[ubFactor]; + LocalTensor sqxLocal = ubLocal[ubFactor * 2]; + LocalTensor tmpLocal = ubLocal[ubFactor * 3]; + + DataCopyCustom(x1Local, x1Gm, numCol); + event_t eventMTE2V1 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::MTE2_V)); + SetFlag(eventMTE2V1); + DataCopyCustom(x2Local, x2Gm, numCol); + event_t eventMTE2V2 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::MTE2_V)); + WaitFlag(eventMTE2V1); + SetFlag(eventMTE2V2); + WaitFlag(eventMTE2V2); + Add(x1Local, x1Local, x2Local, numCol); + PipeBarrier(); + + // copy gamma + event_t eventVMTE2 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_MTE2)); + SetFlag(eventVMTE2); + WaitFlag(eventVMTE2); + + DataCopyCustom(x2Local, gammaGm, numCol); // gammaLocal use x2Local + SetFlag(eventMTE2V2); + + // copy x out + event_t eventVMTE3 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_MTE3)); + SetFlag(eventVMTE3); + WaitFlag(eventVMTE3); + if constexpr (MODE == GAMMA_ADD_RMS_NORM_MODE || MODE == PRE_RMS_NORM_MODE) { + DataCopyCustom(xGm, x1Local, numCol); + } + event_t eventMTE3V = static_cast(GetTPipePtr()->FetchEventID(HardEvent::MTE3_V)); + SetFlag(eventMTE3V); + + Cast(xFp32Local, x1Local, RoundMode::CAST_NONE, this->numCol); + PipeBarrier(); + Mul(sqxLocal, xFp32Local, xFp32Local, this->numCol); + PipeBarrier(); + Muls(sqxLocal, sqxLocal, this->avgFactor, this->numCol); + PipeBarrier(); + ReduceSumCustom(sqxLocal, sqxLocal, tmpLocal, this->numCol); + PipeBarrier(); + Adds(sqxLocal, sqxLocal, this->epsilon, 1); + PipeBarrier(); + Sqrt(sqxLocal, sqxLocal, 1); + Duplicate(tmpLocal, ONE, 1); + PipeBarrier(); + Div(sqxLocal, tmpLocal, sqxLocal, 1); + PipeBarrier(); + + // copyout rstd +#if (defined(__CCE_AICORE__) && __CCE_AICORE__ == 220) || (defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113)) + SetFlag(eventVMTE3); + WaitFlag(eventVMTE3); + if constexpr (MODE == GAMMA_ADD_RMS_NORM_MODE) { + DataCopyCustom(rstdGm, sqxLocal, 1); + } +#endif + event_t eventVS_FP32 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_S)); + SetFlag(eventVS_FP32); + WaitFlag(eventVS_FP32); + float rstdValueFp32 = sqxLocal.GetValue(0); + event_t eventSV_FP32 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::S_V)); + SetFlag(eventSV_FP32); + WaitFlag(eventSV_FP32); + + Muls(xFp32Local, xFp32Local, rstdValueFp32, this->numCol); + PipeBarrier(); + WaitFlag(eventMTE3V); + Cast(x1Local, xFp32Local, RoundMode::CAST_NONE, this->numCol); + PipeBarrier(); + WaitFlag(eventMTE2V2); + if (addGammaOffset != 0U) { + Adds(x2Local, x2Local, static_cast(1.0), this->numCol); + PipeBarrier(); + } + Mul(x1Local, x1Local, x2Local, this->numCol); + SetFlag(eventVMTE3); + WaitFlag(eventVMTE3); + DataCopyCustom(yGm, x1Local, this->numCol); + } + + __aicore__ inline void ProcessFp32() + { + LocalTensor ubLocal = unitBuf.Get(); + LocalTensor x1Local = ubLocal[0]; + LocalTensor x2Local = ubLocal[ubFactor]; + LocalTensor sqxLocal = ubLocal[ubFactor * 2]; + LocalTensor tmpLocal = ubLocal[ubFactor * 3]; + + DataCopyCustom(x1Local, x1Gm, numCol); + event_t eventMTE2V1 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::MTE2_V)); + SetFlag(eventMTE2V1); + DataCopyCustom(x2Local, x2Gm, numCol); + event_t eventMTE2V2 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::MTE2_V)); + SetFlag(eventMTE2V2); + WaitFlag(eventMTE2V1); + WaitFlag(eventMTE2V2); + Add(x1Local, x1Local, x2Local, numCol); + PipeBarrier(); + + // copy gamma + event_t eventVMTE2 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_MTE2)); + SetFlag(eventVMTE2); + WaitFlag(eventVMTE2); + + DataCopyCustom(x2Local, gammaGm, numCol); // gammaLocal use x2Local + SetFlag(eventMTE2V2); + + // copy x out + event_t eventVMTE3 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_MTE3)); + SetFlag(eventVMTE3); + WaitFlag(eventVMTE3); + DataCopyCustom(xGm, x1Local, numCol); + event_t eventMTE3V = static_cast(GetTPipePtr()->FetchEventID(HardEvent::MTE3_V)); + SetFlag(eventMTE3V); + + Mul(sqxLocal, x1Local, x1Local, numCol); + PipeBarrier(); + Muls(sqxLocal, sqxLocal, this->avgFactor, this->numCol); + PipeBarrier(); + ReduceSumCustom(sqxLocal, sqxLocal, tmpLocal, this->numCol); + PipeBarrier(); + Adds(sqxLocal, sqxLocal, this->epsilon, 1); + PipeBarrier(); + Sqrt(sqxLocal, sqxLocal, 1); + Duplicate(tmpLocal, ONE, 1); + PipeBarrier(); + Div(sqxLocal, tmpLocal, sqxLocal, 1); + PipeBarrier(); + + // copyout rstd +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 220 || (defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113)) + SetFlag(eventVMTE3); + WaitFlag(eventVMTE3); + DataCopyCustom(rstdGm, sqxLocal, 1); +#endif + event_t eventVS_FP16 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_S)); + SetFlag(eventVS_FP16); + WaitFlag(eventVS_FP16); + float rstdValue = sqxLocal.GetValue(0); + event_t eventSV_FP16 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::S_V)); + SetFlag(eventSV_FP16); + WaitFlag(eventSV_FP16); + WaitFlag(eventMTE3V); + Muls(x1Local, x1Local, rstdValue, numCol); + PipeBarrier(); + WaitFlag(eventMTE2V2); + if (addGammaOffset != 0U) { + Adds(x2Local, x2Local, static_cast(1.0), numCol); + PipeBarrier(); + } + Mul(x1Local, x1Local, x2Local, numCol); + SetFlag(eventVMTE3); + WaitFlag(eventVMTE3); + DataCopyCustom(yGm, x1Local, numCol); + } + + __aicore__ inline void ProcessBf16() + { + LocalTensor ubLocal = unitBuf.Get(); + LocalTensor xLocal = ubLocal.template ReinterpretCast(); + LocalTensor x1Local = xLocal[0]; + LocalTensor x2Local = xLocal[ubFactor]; + LocalTensor xFp32Local = ubLocal[ubFactor]; + LocalTensor sqxLocal = ubLocal[ubFactor * 2]; + LocalTensor tmpLocal = ubLocal[ubFactor * 3]; + + DataCopyCustom(x1Local, x1Gm, numCol); + event_t eventMTE2V1_BF16_0 = static_cast(GetTPipePtr()->AllocEventID()); + SetFlag(eventMTE2V1_BF16_0); + DataCopyCustom(x2Local, x2Gm, numCol); + event_t eventMTE2V2_BF16_0 = static_cast(GetTPipePtr()->AllocEventID()); + SetFlag(eventMTE2V2_BF16_0); + WaitFlag(eventMTE2V1_BF16_0); + GetTPipePtr()->ReleaseEventID(eventMTE2V1_BF16_0); + Cast(xFp32Local, x1Local, RoundMode::CAST_NONE, numCol); + WaitFlag(eventMTE2V2_BF16_0); + GetTPipePtr()->ReleaseEventID(eventMTE2V2_BF16_0); + Cast(sqxLocal, x2Local, RoundMode::CAST_NONE, numCol); + PipeBarrier(); + Add(xFp32Local, xFp32Local, sqxLocal, numCol); + PipeBarrier(); + Cast(x1Local, xFp32Local, RoundMode::CAST_RINT, numCol); + PipeBarrier(); + // copy gamma + event_t eventVMTE2_BF16_0 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_MTE2)); + SetFlag(eventVMTE2_BF16_0); + WaitFlag(eventVMTE2_BF16_0); + + DataCopyCustom(x2Local, gammaGm, numCol); // gammaLocal use x2Local + event_t eventMTE2V2_BF16_1 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::MTE2_V)); + SetFlag(eventMTE2V2_BF16_1); + + // copy x out + event_t eventVMTE3_BF16_0 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_MTE3)); + SetFlag(eventVMTE3_BF16_0); + WaitFlag(eventVMTE3_BF16_0); + if constexpr (MODE == GAMMA_ADD_RMS_NORM_MODE || MODE == PRE_RMS_NORM_MODE) { + DataCopyCustom(xGm, x1Local, numCol); + } + event_t eventMTE3V_BF16_0 = static_cast(GetTPipePtr()->AllocEventID()); + SetFlag(eventMTE3V_BF16_0); + + Cast(xFp32Local, x1Local, RoundMode::CAST_NONE, numCol); + PipeBarrier(); + Mul(sqxLocal, xFp32Local, xFp32Local, numCol); + PipeBarrier(); + Muls(sqxLocal, sqxLocal, avgFactor, numCol); + PipeBarrier(); + ReduceSumCustom(sqxLocal, sqxLocal, tmpLocal, numCol); + PipeBarrier(); + Adds(sqxLocal, sqxLocal, epsilon, 1); + PipeBarrier(); + Sqrt(sqxLocal, sqxLocal, 1); + Duplicate(tmpLocal, ONE, 1); + PipeBarrier(); + Div(sqxLocal, tmpLocal, sqxLocal, 1); + PipeBarrier(); + event_t eventVS_BF16_0 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_S)); + SetFlag(eventVS_BF16_0); + WaitFlag(eventVS_BF16_0); + float rstdValue = sqxLocal.GetValue(0); + event_t eventSV_BF16_0 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::S_V)); + SetFlag(eventSV_BF16_0); + WaitFlag(eventSV_BF16_0); + // copyout rstd +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 220 || (defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113)) + event_t eventVMTE3_BF16_1 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_MTE3)); + SetFlag(eventVMTE3_BF16_1); + WaitFlag(eventVMTE3_BF16_1); + if constexpr (MODE == GAMMA_ADD_RMS_NORM_MODE) { + DataCopyCustom(rstdGm, sqxLocal, 1); + } + event_t eventMTE3V2_BF16_0 = static_cast(GetTPipePtr()->AllocEventID()); + SetFlag(eventMTE3V2_BF16_0); +#endif + + Muls(xFp32Local, xFp32Local, rstdValue, numCol); + PipeBarrier(); + WaitFlag(eventMTE3V_BF16_0); + GetTPipePtr()->ReleaseEventID(eventMTE3V_BF16_0); + Cast(x1Local, xFp32Local, RoundMode::CAST_RINT, numCol); + PipeBarrier(); + Cast(xFp32Local, x1Local, RoundMode::CAST_NONE, numCol); + PipeBarrier(); + WaitFlag(eventMTE2V2_BF16_1); +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 220 || (defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113)) + WaitFlag(eventMTE3V2_BF16_0); + GetTPipePtr()->ReleaseEventID(eventMTE3V2_BF16_0); +#endif + Cast(sqxLocal, x2Local, RoundMode::CAST_NONE, numCol); + PipeBarrier(); + if (addGammaOffset != 0U) { + Adds(sqxLocal, sqxLocal, static_cast(1.0), numCol); + PipeBarrier(); + Cast(x2Local, sqxLocal, RoundMode::CAST_RINT, numCol); + PipeBarrier(); + Cast(sqxLocal, x2Local, RoundMode::CAST_NONE, numCol); + PipeBarrier(); + } + Mul(xFp32Local, xFp32Local, sqxLocal, numCol); + PipeBarrier(); + Cast(x1Local, xFp32Local, RoundMode::CAST_RINT, numCol); + event_t eventVMTE3_BF16_2 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_MTE3)); + SetFlag(eventVMTE3_BF16_2); + WaitFlag(eventVMTE3_BF16_2); + DataCopyCustom(yGm, x1Local, numCol); + } + +private: + TPipe* Ppipe = nullptr; + + TBuf unitBuf; + GlobalTensor x1Gm; + GlobalTensor x2Gm; + GlobalTensor gammaGm; + GlobalTensor yGm; + GlobalTensor rstdGm; + GlobalTensor xGm; + + uint32_t numRow; + uint32_t numCol; + uint32_t blockFactor; // number of calculations rows on each core + uint32_t ubFactor; + float epsilon; + float avgFactor; + int32_t blockIdx_; + uint32_t rowWork = 1; + uint32_t addGammaOffset = 0; +}; +#endif // _GAMMA_ADD_RMS_NORM_SINGLE_N_H_ diff --git a/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_split_d.h b/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_split_d.h new file mode 100644 index 0000000..c034d30 --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_kernel/gamma_add_rms_norm_split_d.h @@ -0,0 +1,425 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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 gamma_add_rms_norm_split_d.h + * \brief add rms norm split d file + */ +#ifndef GAMMA_ADD_RMS_NORM_SPLIT_D_H_ +#define GAMMA_ADD_RMS_NORM_SPLIT_D_H_ +#include "gamma_add_rms_norm_base.h" + +using namespace AscendC; +using namespace RmsNorm; + +template +class KernelGammaAddRmsNormSplitD { +public: + __aicore__ inline KernelGammaAddRmsNormSplitD(TPipe* pipe) + { + Ppipe = pipe; + } + __aicore__ inline void Init( + GM_ADDR x1, GM_ADDR x2, GM_ADDR gamma, GM_ADDR y, GM_ADDR rstd, GM_ADDR x, GM_ADDR workspace, const GammaAddRMSNormTilingData* tiling) + { + ASSERT(GetBlockNum() != 0 && "Block dim can not be zero!"); + this->numRow = tiling->num_row; + this->numCol = tiling->num_col; + this->blockFactor = tiling->block_factor; + this->rowFactor = tiling->row_factor; + this->ubFactor = tiling->ub_factor; + this->epsilon = tiling->epsilon; + this->avgFactor = (this->numCol != 0) ? (float)1.0 / this->numCol : 0; + this->addGammaOffset = tiling->add_gamma_offset; + + blockIdx_ = GetBlockIdx(); + if (blockIdx_ < GetBlockNum() - 1) { + this->rowWork = this->blockFactor; + } else if (blockIdx_ == GetBlockNum() - 1) { + this->rowWork = this->numRow - (GetBlockNum() - 1) * this->blockFactor; + } else { + } + // get start index for current core, core parallel + x1Gm.SetGlobalBuffer((__gm__ T*)x1 + blockIdx_ * this->blockFactor * this->numCol, this->rowWork * this->numCol); + x2Gm.SetGlobalBuffer((__gm__ T*)x2 + blockIdx_ * this->blockFactor * this->numCol, this->rowWork * this->numCol); + gammaGm.SetGlobalBuffer((__gm__ T*)gamma, this->numCol); + yGm.SetGlobalBuffer((__gm__ T*)y + blockIdx_ * this->blockFactor * this->numCol, this->rowWork * this->numCol); + if constexpr (MODE == GAMMA_ADD_RMS_NORM_MODE) { + rstdGm.SetGlobalBuffer((__gm__ float*)rstd + blockIdx_ * this->blockFactor, this->blockFactor); + xGm.SetGlobalBuffer((__gm__ T*)x + blockIdx_ * this->blockFactor * this->numCol, this->rowWork * this->numCol); + } + if constexpr (MODE == PRE_RMS_NORM_MODE) { + xGm.SetGlobalBuffer((__gm__ T*)x + blockIdx_ * blockFactor * numCol, rowWork * numCol); + } + + // pipe alloc memory to queue, the unit is Bytes. + // We need 2 buffers here for both x1 and x2. + Ppipe->InitBuffer(inQueueX, BUFFER_NUM, 2 * ubFactor * sizeof(T)); + Ppipe->InitBuffer(inQueueGamma, BUFFER_NUM, ubFactor * sizeof(T)); + Ppipe->InitBuffer(outQueueY, BUFFER_NUM, ubFactor * sizeof(T)); + Ppipe->InitBuffer(outQueueRstd, BUFFER_NUM, rowFactor * sizeof(float)); + + if constexpr (is_same::value || is_same::value) { + Ppipe->InitBuffer(xFp32Buf, ubFactor * sizeof(float)); + } + Ppipe->InitBuffer(sqxBuf, ubFactor * sizeof(float)); + Ppipe->InitBuffer(sumBuf, rowFactor * NUM_PER_BLK_FP32 * sizeof(float)); + Ppipe->InitBuffer(reduceFp32Buf, NUM_PER_REP_FP32 * sizeof(float)); + } + + __aicore__ inline void Process() + { + uint32_t i_o_max = RmsNorm::CeilDiv(rowWork, rowFactor); + uint32_t row_tail = rowWork - (i_o_max - 1) * rowFactor; + uint32_t j_max = RmsNorm::CeilDiv(numCol, ubFactor); + uint32_t col_tail = numCol - (j_max - 1) * ubFactor; + for (uint32_t i_o = 0; i_o < i_o_max - 1; i_o++) { + SubProcess(i_o, rowFactor, j_max, col_tail); + } + SubProcess(i_o_max - 1, row_tail, j_max, col_tail); + } + + __aicore__ inline void SubProcess(uint32_t i_o, uint32_t calc_row_num, uint32_t j_max, uint32_t col_tail) + { + LocalTensor sumLocal = sumBuf.Get(); + + LocalTensor rstdLocal = outQueueRstd.AllocTensor(); + Duplicate(rstdLocal, (float)0.0, calc_row_num); + PipeBarrier(); + for (uint32_t j = 0; j < j_max - 1; j++) { + ComputeFormer(i_o, calc_row_num, j, rstdLocal, sumLocal, ubFactor); + } + // do tail + ComputeFormer(i_o, calc_row_num, j_max - 1, rstdLocal, sumLocal, col_tail); + ComputeRstd(rstdLocal, calc_row_num); + + for (uint32_t j = 0; j < j_max - 1; j++) { + ComputeLatter(i_o, calc_row_num, j, rstdLocal, ubFactor); + } + ComputeLatter(i_o, calc_row_num, j_max - 1, rstdLocal, col_tail); + + if constexpr (MODE == GAMMA_ADD_RMS_NORM_MODE) { + outQueueRstd.EnQue(rstdLocal); + CopyOutRstd(i_o, calc_row_num); + } else { + outQueueRstd.FreeTensor(rstdLocal); + } + } + +private: + __aicore__ inline void CopyInX1X2(uint32_t i_idx, uint32_t j_idx, uint32_t num) + { + LocalTensor splitX1X2In = inQueueX.AllocTensor(); + LocalTensor splitX1In = splitX1X2In[0]; + LocalTensor splitX2In = splitX1X2In[this->ubFactor]; + DataCopyCustom(splitX1In, x1Gm[i_idx * this->numCol + j_idx * this->ubFactor], num); + DataCopyCustom(splitX2In, x2Gm[i_idx * this->numCol + j_idx * this->ubFactor], num); + inQueueX.EnQue(splitX1X2In); + } + + __aicore__ inline void AddX(uint32_t i_idx, uint32_t j_idx, uint32_t num) + { + CopyInX1X2(i_idx, j_idx, num); + LocalTensor splitX1X2Local = inQueueX.DeQue(); + + auto splitX1Local = splitX1X2Local[0]; + auto splitX2Local = splitX1X2Local[this->ubFactor]; + if constexpr (is_same::value) { + LocalTensor x1_fp32 = xFp32Buf.Get(); + LocalTensor x2_fp32 = splitX1X2Local.template ReinterpretCast(); + Cast(x1_fp32, splitX1Local, RoundMode::CAST_NONE, num); + PipeBarrier(); + Cast(x2_fp32, splitX2Local, RoundMode::CAST_NONE, num); + PipeBarrier(); + Add(x1_fp32, x1_fp32, x2_fp32, num); + PipeBarrier(); + Cast(splitX1X2Local, x1_fp32, RoundMode::CAST_RINT, num); + } else { + Add(splitX1X2Local, splitX1Local, splitX2Local, num); + } + PipeBarrier(); + inQueueX.EnQue(splitX1X2Local); + } + + __aicore__ inline void CopyInAndAdd(uint32_t i_idx, uint32_t j_idx, uint32_t num) + { + CopyInX1X2(i_idx, j_idx, num); + LocalTensor splitX1X2Local = inQueueX.DeQue(); + auto splitX1Local = splitX1X2Local[0]; + auto splitX2Local = splitX1X2Local[this->ubFactor]; + + LocalTensor splitXLocal = outQueueY.AllocTensor(); + + if constexpr (is_same::value) { + LocalTensor splitX1Fp32 = xFp32Buf.Get(); + + Add(splitXLocal, splitX1Local, splitX2Local, num); + PipeBarrier(); + Cast(splitX1Fp32, splitXLocal, RoundMode::CAST_NONE, num); + PipeBarrier(); + // x1+x2 saved in x1_fp32 + } else if constexpr (is_same::value) { + LocalTensor x1_fp32 = xFp32Buf.Get(); + LocalTensor x2_fp32 = splitX1X2Local.template ReinterpretCast(); + + Cast(x1_fp32, splitX1Local, RoundMode::CAST_NONE, num); + PipeBarrier(); + Cast(x2_fp32, splitX2Local, RoundMode::CAST_NONE, num); + PipeBarrier(); + + Add(x1_fp32, x1_fp32, x2_fp32, num); + PipeBarrier(); + Cast(splitXLocal, x1_fp32, RoundMode::CAST_RINT, num); + PipeBarrier(); + // x1+x2 saved in x1_fp32 + } else { + Add(splitX1Local, splitX1Local, splitX2Local, num); + PipeBarrier(); + Adds(splitXLocal, splitX1Local, (float)0.0, num); + // x1+x2 saved in inQueueX + } + inQueueX.FreeTensor(splitX1X2Local); + + // copy out to workspace && x_out + outQueueY.EnQue(splitXLocal); + auto x_out = outQueueY.DeQue(); + if constexpr (MODE == GAMMA_ADD_RMS_NORM_MODE || MODE == PRE_RMS_NORM_MODE) { + DataCopyCustom(xGm[i_idx * numCol + j_idx * ubFactor], x_out, num); + } + outQueueY.FreeTensor(x_out); + } + + __aicore__ inline void ComputeFormer( + uint32_t i_o_idx, uint32_t calc_row_num, uint32_t j_idx, LocalTensor& rstdLocal, + LocalTensor& sumLocal, uint32_t num) + { + for (uint32_t i_i = 0; i_i < calc_row_num; i_i++) { + CopyInAndAdd(i_o_idx * rowFactor + i_i, j_idx, num); + ComputeSum(i_i, sumLocal, num); + } + BlockReduceSumFP32(sumLocal, sumLocal, calc_row_num * NUM_PER_BLK_FP32); + Add(rstdLocal, rstdLocal, sumLocal, calc_row_num); + PipeBarrier(); + } + + __aicore__ inline void ComputeSum(uint32_t i_i_idx, LocalTensor& sumLocal, uint32_t num) + { + LocalTensor sqx = sqxBuf.Get(); + LocalTensor reduce_buf_local = reduceFp32Buf.Get(); + if constexpr (is_same::value || is_same::value) { + LocalTensor x_fp32 = xFp32Buf.Get(); + PipeBarrier(); + Mul(sqx, x_fp32, x_fp32, num); + } else { + LocalTensor xLocal = inQueueX.AllocTensor(); + PipeBarrier(); + Mul(sqx, xLocal, xLocal, num); + inQueueX.FreeTensor(xLocal); + } + PipeBarrier(); + Muls(sqx, sqx, avgFactor, num); + PipeBarrier(); + // 8 means 8 fp32 pre block + ReduceSumFP32ToBlock(sumLocal[i_i_idx * 8], sqx, reduce_buf_local, num); + } + + __aicore__ inline void ComputeRstd(LocalTensor rstdLocal, uint32_t num) + { + LocalTensor splitReduceBufLocal = reduceFp32Buf.Get(); + Adds(rstdLocal, rstdLocal, this->epsilon, num); + PipeBarrier(); + Sqrt(rstdLocal, rstdLocal, num); + Duplicate(splitReduceBufLocal, ONE, num); + PipeBarrier(); + Div(rstdLocal, splitReduceBufLocal, rstdLocal, num); + PipeBarrier(); + } + + __aicore__ inline void ComputeLatter( + uint32_t i_o_idx, uint32_t calc_row_num, uint32_t j_idx, LocalTensor& rstdLocal, uint32_t num) + { + CopyInGamma(j_idx, num); + LocalTensor splitGammaLocal = inQueueGamma.DeQue(); + if (addGammaOffset != 0U) { + AddGammaOffset(splitGammaLocal, num); + } + for (uint32_t i_i = 0; i_i < calc_row_num; i_i++) { + CopyInX(i_o_idx * rowFactor + i_i, j_idx, num); + ComputeY(i_i, splitGammaLocal, rstdLocal, num); + CopyOutY(i_o_idx * rowFactor + i_i, j_idx, num); + } + inQueueGamma.FreeTensor(splitGammaLocal); + } + + __aicore__ inline void CopyInGamma(uint32_t j_idx, uint32_t num) + { + LocalTensor gammaLocal = inQueueGamma.AllocTensor(); + DataCopyCustom(gammaLocal, gammaGm[j_idx * ubFactor], num); + inQueueGamma.EnQue(gammaLocal); + } + + __aicore__ inline void AddGammaOffset(LocalTensor& gammaLocal, uint32_t num) + { + if constexpr (is_same::value) { + LocalTensor gammaFp32 = xFp32Buf.Get(); + Cast(gammaFp32, gammaLocal, RoundMode::CAST_NONE, num); + PipeBarrier(); + Adds(gammaFp32, gammaFp32, static_cast(1.0), num); + PipeBarrier(); + Cast(gammaLocal, gammaFp32, RoundMode::CAST_RINT, num); + } else { + Adds(gammaLocal, gammaLocal, static_cast(1.0), num); + } + PipeBarrier(); + } + + __aicore__ inline void CopyInX(uint32_t i_idx, uint32_t j_idx, uint32_t num) + { + if constexpr (MODE == GAMMA_ADD_RMS_NORM_MODE || MODE == PRE_RMS_NORM_MODE) { + LocalTensor xLocal = inQueueX.AllocTensor(); + DataCopyCustom(xLocal, xGm[i_idx * numCol + j_idx * ubFactor], num); + inQueueX.EnQue(xLocal); + } + if constexpr (MODE == POST_RMS_NORM_MODE) { + AddX(i_idx, j_idx, num); + } + if constexpr (is_same::value || is_same::value) { + LocalTensor splitXFp32 = xFp32Buf.Get(); + LocalTensor splitXLocalDeq = inQueueX.DeQue(); + Cast(splitXFp32, splitXLocalDeq, RoundMode::CAST_NONE, num); + PipeBarrier(); + inQueueX.FreeTensor(splitXLocalDeq); + } + } + + __aicore__ inline void ComputeY( + uint32_t i_i_idx, LocalTensor& splitGammaLocal, LocalTensor& splitRstdLocal, uint32_t num) + { + LocalTensor splitXFp32 = xFp32Buf.Get(); + LocalTensor splitSqx = sqxBuf.Get(); + event_t splitEventVS = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_S)); + SetFlag(splitEventVS); + WaitFlag(splitEventVS); + float splitRstdValue = splitRstdLocal.GetValue(i_i_idx); + event_t splitEventSV = static_cast(GetTPipePtr()->FetchEventID(HardEvent::S_V)); + SetFlag(splitEventSV); + WaitFlag(splitEventSV); + PipeBarrier(); + Muls(splitXFp32, splitXFp32, splitRstdValue, num); + PipeBarrier(); + LocalTensor splitYLocal = outQueueY.AllocTensor(); + Cast(splitYLocal, splitXFp32, RoundMode::CAST_NONE, num); + PipeBarrier(); + Mul(splitYLocal, splitGammaLocal, splitYLocal, num); + PipeBarrier(); + outQueueY.EnQue(splitYLocal); + } + + __aicore__ inline void ComputeY( + uint32_t i_i_idx, LocalTensor& gammaLocal, LocalTensor& rstdLocal, uint32_t num) + { + LocalTensor xLocal = inQueueX.DeQue(); + LocalTensor sqx = sqxBuf.Get(); + event_t event_v_s = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_S)); + SetFlag(event_v_s); + WaitFlag(event_v_s); + float rstdValue = rstdLocal.GetValue(i_i_idx); + event_t event_s_v = static_cast(GetTPipePtr()->FetchEventID(HardEvent::S_V)); + SetFlag(event_s_v); + WaitFlag(event_s_v); + LocalTensor yLocal = outQueueY.AllocTensor(); + Muls(yLocal, xLocal, rstdValue, num); + inQueueX.FreeTensor(xLocal); + PipeBarrier(); + Mul(yLocal, gammaLocal, yLocal, num); + PipeBarrier(); + outQueueY.EnQue(yLocal); + } + + __aicore__ inline void ComputeY( + uint32_t i_i_idx, LocalTensor& gammaLocal, LocalTensor& rstdLocal, uint32_t num) + { + LocalTensor splitXFp32Bf16 = xFp32Buf.Get(); + LocalTensor splitSqxBf16 = sqxBuf.Get(); + event_t splitEventVSBf16 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_S)); + SetFlag(splitEventVSBf16); + WaitFlag(splitEventVSBf16); + float splitRstdValueBf16 = rstdLocal.GetValue(i_i_idx); + event_t splitEventSVBf16 = static_cast(GetTPipePtr()->FetchEventID(HardEvent::S_V)); + SetFlag(splitEventSVBf16); + WaitFlag(splitEventSVBf16); + PipeBarrier(); + Muls(splitXFp32Bf16, splitXFp32Bf16, splitRstdValueBf16, num); + PipeBarrier(); + LocalTensor splitYLocalBf16 = outQueueY.AllocTensor(); + Cast(splitYLocalBf16, splitXFp32Bf16, RoundMode::CAST_RINT, num); + PipeBarrier(); + Cast(splitXFp32Bf16, splitYLocalBf16, RoundMode::CAST_NONE, num); + PipeBarrier(); + Cast(splitSqxBf16, gammaLocal, RoundMode::CAST_NONE, num); + PipeBarrier(); + Mul(splitXFp32Bf16, splitXFp32Bf16, splitSqxBf16, num); + PipeBarrier(); + Cast(splitYLocalBf16, splitXFp32Bf16, RoundMode::CAST_RINT, num); + PipeBarrier(); + outQueueY.EnQue(splitYLocalBf16); + } + + __aicore__ inline void CopyOutY(uint32_t i_idx, uint32_t j_idx, uint32_t num) + { + LocalTensor splitYLocalOut = outQueueY.DeQue(); + DataCopyCustom(yGm[i_idx * this->numCol + j_idx * this->ubFactor], splitYLocalOut, num); + outQueueY.FreeTensor(splitYLocalOut); + } + + __aicore__ inline void CopyOutRstd(uint32_t i_o_idx, uint32_t num) + { + LocalTensor splitRstdLocal = outQueueRstd.DeQue(); +#if __CCE_AICORE__ == 220 || (defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113)) + DataCopyCustom(rstdGm[i_o_idx * this->rowFactor], splitRstdLocal, num); +#endif + outQueueRstd.FreeTensor(splitRstdLocal); + } + +private: + TPipe* Ppipe = nullptr; + // create input queues for split_d + TQue inQueueX; + TQue inQueueGamma; + // create output queues for split_d + TQue outQueueY; + TQue outQueueRstd; + TBuf xFp32Buf; + TBuf sqxBuf; + TBuf sumBuf; + TBuf reduceFp32Buf; + + GlobalTensor x1Gm; + GlobalTensor x2Gm; + GlobalTensor gammaGm; + GlobalTensor yGm; + GlobalTensor rstdGm; + GlobalTensor xGm; + + uint32_t numRow; + uint32_t numCol; + uint32_t blockFactor; + uint32_t rowFactor; + uint32_t ubFactor; + float epsilon; + float avgFactor; + int32_t blockIdx_; + uint32_t rowWork = 1; + uint32_t addGammaOffset = 0; + + int tempbufNum; +}; +#endif // _GAMMA_ADD_RMS_NORM_SPLIT_D_H_ diff --git a/xllm_ops/gamma_add_rms_norm/op_kernel/inc/platform.h b/xllm_ops/gamma_add_rms_norm/op_kernel/inc/platform.h new file mode 100644 index 0000000..67b2957 --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_kernel/inc/platform.h @@ -0,0 +1,73 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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 platform.h + * \brief platform apator + */ +#ifndef OPS_BUILT_IN_OP_ASCENDC_PLATFORM_INFO_H_ +#define OPS_BUILT_IN_OP_ASCENDC_PLATFORM_INFO_H_ + +#if ASC_DEVKIT_MAJOR >= 9 +#include "kernel_basic_intf.h" +#else +#include "kernel_operator.h" +#endif +#include "kernel_tiling/kernel_tiling.h" + +#ifndef KERNEL_API +#define KERNEL_API extern "C" __global__ __aicore__ +#endif + +namespace platform { + +#define MID_THREAD_NUM 1024 + +__aicore__ inline constexpr bool IsDataCopyPadSupport() +{ +#if __CCE_AICORE__ == 220 + return true; +#else + return false; +#endif +} + +/** + * Get the block size of unified buffer in bytes + */ +__aicore__ inline constexpr uint32_t GetUbBlockSize() +{ + return 32U; +} + +/** + * Get the size of vector registers in bytes + */ +__aicore__ inline constexpr uint32_t GetVRegSize() +{ +#if __CCE_AICORE__ == 310 + return AscendC::VECTOR_REG_WIDTH; +#else + return 256U; +#endif +} + +} // namespace platform + +namespace PlatformSocInfo { +__aicore__ inline constexpr bool IsDataCopyPadSupport() +{ + return platform::IsDataCopyPadSupport(); +} + +} + +#endif // OPS_BUILT_IN_OP_ASCENDC_PLATFORM_INFO_H_ diff --git a/xllm_ops/gamma_add_rms_norm/op_kernel/reduce_common.h b/xllm_ops/gamma_add_rms_norm/op_kernel/reduce_common.h new file mode 100644 index 0000000..5346493 --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/op_kernel/reduce_common.h @@ -0,0 +1,180 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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 reduce_common.h + */ +#ifndef REDUCE_COMMON_H_RMS_NORM +#define REDUCE_COMMON_H_RMS_NORM +#include "kernel_operator.h" +using namespace AscendC; + +constexpr uint32_t ELEM_PER_REP_FP32 = 64; +constexpr uint32_t ELEM_PER_BLK_FP32 = 8; +constexpr int32_t INDEX_SIXTEEN = 16; +constexpr float ZERO = 0; +constexpr int32_t HALf_INTERVAL = 2; +constexpr uint32_t MAX_REP_NUM = 255; +constexpr int32_t INDEX_TWO = 2; +constexpr int32_t INDEX_FOUR = 4; +constexpr int32_t INDEX_EIGHT = 8; + +__aicore__ inline void ReduceSumForSmallReduceDimPreRepeat( + const LocalTensor& dstLocal, const LocalTensor& srcLocal, const LocalTensor& tmpLocal, + const uint32_t elemNum, const uint32_t numLastDim, const uint32_t tailCount, const uint32_t repeat, + const uint8_t repStride) +{ + uint32_t elemIndex4 = 0; + for (; elemIndex4 + ELEM_PER_REP_FP32 <= numLastDim; elemIndex4 += ELEM_PER_REP_FP32) { + Add(tmpLocal, srcLocal[elemIndex4], tmpLocal, elemNum, repeat, + {1, 1, 1, ELEM_PER_BLK_FP32, repStride, ELEM_PER_BLK_FP32}); + PipeBarrier(); + } + if (unlikely(tailCount != 0)) { + Add(tmpLocal, srcLocal[elemIndex4], tmpLocal, tailCount, repeat, + {1, 1, 1, ELEM_PER_BLK_FP32, repStride, ELEM_PER_BLK_FP32}); + } + PipeBarrier(); + AscendCUtils::SetMask(ELEM_PER_REP_FP32); // set mask = 64 +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 220 + if ASCEND_IS_AIV { + WholeReduceSum(dstLocal, tmpLocal, elemNum, repeat, 1, 1, ELEM_PER_BLK_FP32); + } +#elif defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113) + WholeReduceSum(dstLocal, tmpLocal, elemNum, repeat, 1, 1, ELEM_PER_BLK_FP32); +#else + WholeReduceSum(dstLocal, tmpLocal, elemNum, repeat, 1, 1, ELEM_PER_BLK_FP32); +#endif +} + +/* + * reduce dim form (N, D) to (N, 1) + * this reduce sum is for small reduce dim. + */ +__aicore__ inline void ReduceSumForSmallReduceDim( + const LocalTensor& dstLocal2, const LocalTensor& srcLocal, const LocalTensor& tmpLocal, + const uint32_t numLastDimAligned, const uint32_t numLastDim, const uint32_t tailCount, const uint32_t repeat, + const uint8_t repStride) +{ + uint32_t repeatTimes = repeat / MAX_REP_NUM; + if (repeatTimes == 0) { + ReduceSumForSmallReduceDimPreRepeat( + dstLocal2, srcLocal, tmpLocal, ELEM_PER_REP_FP32, numLastDim, tailCount, repeat, repStride); + } else { + uint32_t repTailNum = repeat % MAX_REP_NUM; + uint32_t repIndex = 0; + uint32_t repElem; + for (; repIndex + MAX_REP_NUM <= repeat; repIndex += MAX_REP_NUM) { + ReduceSumForSmallReduceDimPreRepeat( + dstLocal2[repIndex], srcLocal[repIndex * numLastDimAligned], tmpLocal[repIndex * ELEM_PER_REP_FP32], + ELEM_PER_REP_FP32, numLastDim, tailCount, MAX_REP_NUM, repStride); + } + if (repTailNum != 0) { + ReduceSumForSmallReduceDimPreRepeat( + dstLocal2[repIndex], srcLocal[repIndex * numLastDimAligned], tmpLocal[repIndex * ELEM_PER_REP_FP32], + ELEM_PER_REP_FP32, numLastDim, tailCount, repTailNum, repStride); + } + } +} + +/* + * reduce dim form (N, D) to (N, 1) + * this reduce sum is for small reduce dim, require D < 255 * 8. + * size of tmpLocal: (N, 64) + */ +__aicore__ inline void ReduceSumMultiN( + const LocalTensor& dstLocal, const LocalTensor& srcLocal, const LocalTensor& tmpLocal4, + const uint32_t numRow, const uint32_t numCol, const uint32_t numColAlign) +{ + const uint32_t tailCount = numCol % ELEM_PER_REP_FP32; + const uint32_t repeat = numRow; + const uint8_t repStride = numColAlign / ELEM_PER_BLK_FP32; + Duplicate(tmpLocal4, ZERO, numRow * ELEM_PER_REP_FP32); + PipeBarrier(); + ReduceSumForSmallReduceDim(dstLocal, srcLocal, tmpLocal4, numColAlign, numCol, tailCount, repeat, repStride); +} + +__aicore__ inline int32_t findPowerTwo(int32_t n1) +{ + // find max power of 2 no more than n (32 bit) + n1 |= n1 >> 1; // Set the first digit of n's binary to 1 + n1 |= n1 >> INDEX_TWO; + n1 |= n1 >> INDEX_FOUR; + n1 |= n1 >> INDEX_EIGHT; + n1 |= n1 >> INDEX_SIXTEEN; + return (n1 + 1) >> 1; +} + +__aicore__ inline void ReduceSumHalfInterval( + const LocalTensor& dst_local, const LocalTensor& src_local6, int32_t count) +{ + if (likely(count > ELEM_PER_REP_FP32)) { + int32_t bodyCount = findPowerTwo(count); + int32_t tailCount = count - bodyCount; + if (tailCount > 0) { + Add(src_local6, src_local6, src_local6[bodyCount], tailCount); + PipeBarrier(); + } + while (bodyCount > ELEM_PER_REP_FP32) { + bodyCount = bodyCount / HALf_INTERVAL; + Add(src_local6, src_local6, src_local6[bodyCount], bodyCount); + PipeBarrier(); + } + + AscendCUtils::SetMask(ELEM_PER_REP_FP32); + } else { + AscendCUtils::SetMask(count); + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 220 + if (g_coreType == AIV) { + WholeReduceSum(dst_local, src_local6, ELEM_PER_REP_FP32, 1, 0, 1, 0); + } +#elif defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113) + WholeReduceSum(dst_local, src_local6, ELEM_PER_REP_FP32, 1, 1, 1, ELEM_PER_BLK_FP32); +#else + WholeReduceSum(dst_local, src_local6, ELEM_PER_REP_FP32, 1, 1, 1, DEFAULT_REPEAT_STRIDE); +#endif + PipeBarrier(); +} + +__aicore__ inline float ReduceSumHalfInterval(const LocalTensor& src_local4, int32_t count) +{ + if (likely(count > ELEM_PER_REP_FP32)) { + int32_t bodyCount = findPowerTwo(count); + int32_t tailCount = count - bodyCount; + if (tailCount > 0) { + Add(src_local4, src_local4, src_local4[bodyCount], tailCount); + PipeBarrier(); + } + while (bodyCount > ELEM_PER_REP_FP32) { + bodyCount = bodyCount / HALf_INTERVAL; + Add(src_local4, src_local4, src_local4[bodyCount], bodyCount); + PipeBarrier(); + } + + AscendCUtils::SetMask(ELEM_PER_REP_FP32); + } else { + AscendCUtils::SetMask(count); + } +#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 220 + if (g_coreType == AIV) { + WholeReduceSum(src_local4, src_local4, ELEM_PER_REP_FP32, 1, 0, 1, 0); + } +#elif defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113) + WholeReduceSum(src_local4, src_local4, ELEM_PER_REP_FP32, 1, 1, 1, ELEM_PER_BLK_FP32); +#else + WholeReduceSum(src_local4, src_local4, ELEM_PER_REP_FP32, 1, 1, 1, DEFAULT_REPEAT_STRIDE); +#endif + event_t event_v_s = static_cast(GetTPipePtr()->FetchEventID(HardEvent::V_S)); + SetFlag(event_v_s); + WaitFlag(event_v_s); + return src_local4.GetValue(0); +} +#endif // _REDUCE_COMMON_H_ diff --git a/xllm_ops/gamma_add_rms_norm/rms_norm/arch35/rms_norm_regbase_common.h b/xllm_ops/gamma_add_rms_norm/rms_norm/arch35/rms_norm_regbase_common.h new file mode 100644 index 0000000..c317284 --- /dev/null +++ b/xllm_ops/gamma_add_rms_norm/rms_norm/arch35/rms_norm_regbase_common.h @@ -0,0 +1,1401 @@ +/** + * Copyright (c) 2025 Huawei Technologies Co., Ltd. + * Copyright 2026 The xLLM Authors. All Rights Reserved. + * 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 rms_norm_regbase_common.h + * \brief RmsNorm regbase common + */ +#ifndef OPS_BUILT_IN_TBE_IMPL_ASCENDC_RMS_NORM_REGBASE_COMMON_H +#define OPS_BUILT_IN_TBE_IMPL_ASCENDC_RMS_NORM_REGBASE_COMMON_H +#include "kernel_operator.h" +#include "kernel_tiling/kernel_tiling.h" +#include "../../op_kernel/gamma_add_rms_norm_base.h" +#include "../../norm_common/reduce_common_regbase.h" + +namespace RmsNorm { +using namespace AscendC; +using namespace AscendC::MicroAPI; +using namespace NormCommon; +using namespace NormCommon::NormCommonRegbase; + +#ifndef FLOAT_OVERFLOW_MODE_CTRL +#define FLOAT_OVERFLOW_MODE_CTRL 60 +#endif + +constexpr AscendC::MicroAPI::CastTrait castTraitFp322Fp8 = { + AscendC::MicroAPI::RegLayout::ZERO, + AscendC::MicroAPI::SatMode::SAT, + AscendC::MicroAPI::MaskMergeMode::ZEROING, + RoundMode::CAST_RINT, +}; + +constexpr AscendC::MicroAPI::CastTrait castTraitFp322Hifp8 = { + AscendC::MicroAPI::RegLayout::ZERO, + AscendC::MicroAPI::SatMode::SAT, + AscendC::MicroAPI::MaskMergeMode::ZEROING, + RoundMode::CAST_ROUND, +}; + +template +__aicore__ inline uint64_t GetOverflowMode() +{ +#if (__NPU_ARCH__ == 3510) + if constexpr (IsSameType::value || IsSameType::value || + IsSameType::value) { + return AscendC::GetCtrlSpr(); + } +#endif + return 0; +} + +template +__aicore__ inline void SetOverflowMode(uint64_t mode) +{ +#if (__NPU_ARCH__ == 3510) + if constexpr (IsSameType::value || IsSameType::value || + IsSameType::value) { + AscendC::SetCtrlSpr(mode); + } +#endif +} + +template +__aicore__ inline void YCopyOutImpl( + const U& dstTensor, const R& srcTensor, uint32_t blockCount, uint32_t blockLen, uint32_t srcStride = 0, + uint32_t dstStride = 0) +{ + DataCopyExtParams extParams{ + static_cast(blockCount), // blockCount + static_cast(blockLen * sizeof(T)), // blockLen + srcStride, // srcStride + dstStride, // dstStride + 0 // rsv + }; + DataCopyPad(dstTensor, srcTensor, extParams); +} + +/*! + * DataCopy custom implement + * @tparam T DataCopy data type + * @tparam U Destination Operand type + * @tparam R Source Operand type + * @param dstTensor Destination Tensor + * @param srcTensor Source Tensor + * @param blockCount burst + * @param blockLen burst length + * @param padParams pad params + * @return void + */ +template +__aicore__ inline void DataCopyImpl( + const U& dstTensor, const R& srcTensor, uint32_t blockCount, uint32_t blockLen, uint32_t srcStride = 0, + uint32_t dstStride = 0, const DataCopyPadExtParams padParams = {false, 0, 0, 0}) +{ + DataCopyExtParams extParams{ + static_cast(blockCount), // blockCount + static_cast(blockLen * sizeof(T)), // blockLen + srcStride, // srcStride + dstStride, // dstStride + 0 // rsv + }; + if constexpr (is_same>::value) { + DataCopyPad(dstTensor, srcTensor, extParams, padParams); + } else { + DataCopyPad(dstTensor, srcTensor, extParams); + } +} + +template +__aicore__ inline void CopyInX( + TQue& inQueueX, GlobalTensor& srcGm, uint64_t gmOffset, uint32_t blockLen, + uint32_t left = 0, uint32_t right = 0) +{ + LocalTensor xLocal = inQueueX.AllocTensor(); + DataCopyPadExtParams padParams{ + true, // isPad + static_cast(left), // leftPadding + static_cast(right), // rightPadding + static_cast(0.0) // paddingValue + }; + DataCopyImpl(xLocal, srcGm[gmOffset], 1, blockLen, 0, 0, padParams); + inQueueX.EnQue(xLocal); +} + +template +__aicore__ inline void CopyOutX( + GlobalTensor& xGm, TQue& outQueueX, uint64_t gmOffset, uint32_t blockLen) +{ + LocalTensor xLocal = outQueueX.DeQue(); + DataCopyImpl(xGm[gmOffset], xLocal, 1, blockLen); + outQueueX.FreeTensor(xLocal); +} + +template +__aicore__ inline void CopyOutY( + GlobalTensor& yGm, TQue& outQueueY, uint64_t gmOffset, uint32_t blockLen) +{ + LocalTensor yLocal = outQueueY.DeQue(); + YCopyOutImpl(yGm[gmOffset], yLocal, 1, blockLen); + outQueueY.FreeTensor(yLocal); +} + +/*! + * x = x * x + scalar + * @param dstLocal dst Tensor + * @param srcLocal src Tensor + * @param scalar scalar + * @param count vector size + * @return + */ +__aicore__ inline void ComputeFormer( + LocalTensor& dstLocal, LocalTensor& srcLocal, float scalar, uint32_t count) +{ + __local_mem__ float* srcAddr = (__ubuf__ float*)srcLocal.GetPhyAddr(); + __local_mem__ float* dstAddr = (__ubuf__ float*)dstLocal.GetPhyAddr(); + + uint32_t calCount = count; + uint32_t sreg = (uint32_t)calCount; + uint16_t repeatTimes = CeilDivision(calCount, V_LENGTH); + __VEC_SCOPE__ + { + RegTensor vReg, vRegTmp; + MaskReg maskReg; + for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) { + maskReg = UpdateMask(sreg); + DataCopy(vReg, srcAddr + i * V_LENGTH); + Mul(vRegTmp, vReg, vReg, maskReg); + Muls(vReg, vRegTmp, scalar, maskReg); + DataCopy(dstAddr + i * V_LENGTH, vReg, maskReg); + } + } +} + +/*! + * x = x * x * scalar + * @param dstLocal dst Tensor + * @param xFp32 src.as(float32) Tensor + * @param srcLocal src Tensor + * @param scalar scalar + * @param count vector size + * @return + */ +template +__aicore__ inline void ComputeFormer( + LocalTensor& dstLocal, LocalTensor& xFp32, LocalTensor& srcLocal, float scalar, uint32_t count) +{ + uint32_t calCount = count / 2; + uint32_t sreg = (uint32_t)calCount; + uint16_t repeatTimes = CeilDivision(calCount, V_LENGTH); + + __local_mem__ T* xAddr1 = (__ubuf__ T*)srcLocal.GetPhyAddr(); + __local_mem__ T* xAddr2 = (__ubuf__ T*)srcLocal.GetPhyAddr() + calCount; + __local_mem__ float* dstAddr1 = (__ubuf__ float*)dstLocal.GetPhyAddr(); + __local_mem__ float* dstAddr2 = (__ubuf__ float*)dstLocal.GetPhyAddr() + calCount; + __local_mem__ float* xFp32Addr1 = (__ubuf__ float*)xFp32.GetPhyAddr(); + __local_mem__ float* xFp32Addr2 = (__ubuf__ float*)xFp32.GetPhyAddr() + calCount; + + if constexpr (IsSameType::value) { + __VEC_SCOPE__ + { + RegTensor xFp16Reg1, xFp16Reg2; + RegTensor xFp32Reg1, vReg1, vRegTmp1, xFp32Reg2, vReg2, vRegTmp2; + MaskReg maskReg; + for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) { + maskReg = UpdateMask(sreg); + DataCopy(xFp16Reg1, xAddr1 + i * V_LENGTH); + DataCopy(xFp16Reg2, xAddr2 + i * V_LENGTH); + Cast(xFp32Reg1, xFp16Reg1, maskReg); + Cast(xFp32Reg2, xFp16Reg2, maskReg); + Mul(vRegTmp1, xFp32Reg1, xFp32Reg1, maskReg); + Mul(vRegTmp2, xFp32Reg2, xFp32Reg2, maskReg); + Muls(vReg1, vRegTmp1, scalar, maskReg); + Muls(vReg2, vRegTmp2, scalar, maskReg); + DataCopy(dstAddr1 + i * V_LENGTH, vReg1, maskReg); + DataCopy(dstAddr2 + i * V_LENGTH, vReg2, maskReg); + DataCopy(xFp32Addr1 + i * V_LENGTH, xFp32Reg1, maskReg); + DataCopy(xFp32Addr2 + i * V_LENGTH, xFp32Reg2, maskReg); + } + } + } else if constexpr (IsSameType::value) { + __VEC_SCOPE__ + { + RegTensor xBFp16Reg1, xBFp16Reg2; + RegTensor xFp32Reg1, vReg1, vRegTmp1, xFp32Reg2, vReg2, vRegTmp2; + MaskReg maskReg; + for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) { + maskReg = UpdateMask(sreg); + DataCopy(xBFp16Reg1, xAddr1 + i * V_LENGTH); + DataCopy(xBFp16Reg2, xAddr2 + i * V_LENGTH); + Cast(xFp32Reg1, xBFp16Reg1, maskReg); + Cast(xFp32Reg2, xBFp16Reg2, maskReg); + Mul(vRegTmp1, xFp32Reg1, xFp32Reg1, maskReg); + Mul(vRegTmp2, xFp32Reg2, xFp32Reg2, maskReg); + Muls(vReg1, vRegTmp1, scalar, maskReg); + Muls(vReg2, vRegTmp2, scalar, maskReg); + DataCopy(dstAddr1 + i * V_LENGTH, vReg1, maskReg); + DataCopy(dstAddr2 + i * V_LENGTH, vReg2, maskReg); + DataCopy(xFp32Addr1 + i * V_LENGTH, xFp32Reg1, maskReg); + DataCopy(xFp32Addr2 + i * V_LENGTH, xFp32Reg2, maskReg); + } + } + } else { + __VEC_SCOPE__ + { + RegTensor vReg1, vRegTmp1, vReg2, vRegTmp2; + MaskReg maskReg; + for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) { + maskReg = UpdateMask(sreg); + DataCopy(vReg1, xAddr1 + i * V_LENGTH); + DataCopy(vReg2, xAddr2 + i * V_LENGTH); + Mul(vRegTmp1, vReg1, vReg1, maskReg); + Mul(vRegTmp2, vReg2, vReg2, maskReg); + Muls(vReg1, vRegTmp1, scalar, maskReg); + Muls(vReg2, vRegTmp2, scalar, maskReg); + DataCopy(dstAddr1 + i * V_LENGTH, vReg1, maskReg); + DataCopy(dstAddr2 + i * V_LENGTH, vReg2, maskReg); + } + } + } +} + +/*! + * dst = srcLocal * srcLocal + scalar + * + * @tparam T src dtype + * @param dstLocal output Local Tensor + * @param srcLocal input Local Tensor + * @param scalar average num + * @param count vector compute length + * @return void + */ +template +__aicore__ inline void ComputeInit(LocalTensor& dstLocal, LocalTensor& srcLocal, float scalar, uint32_t count) +{ + uint32_t calCount = count / 2; + uint32_t sreg = (uint32_t)calCount; + uint16_t repeatTimes = CeilDivision(calCount, V_LENGTH); + + __local_mem__ T* xAddr1 = (__ubuf__ T*)srcLocal.GetPhyAddr(); + __local_mem__ T* xAddr2 = (__ubuf__ T*)srcLocal.GetPhyAddr() + calCount; + __local_mem__ float* dstAddr1 = (__ubuf__ float*)dstLocal.GetPhyAddr(); + __local_mem__ float* dstAddr2 = (__ubuf__ float*)dstLocal.GetPhyAddr() + calCount; + + if constexpr (IsSameType::value) { + __VEC_SCOPE__ + { + RegTensor xFp16Reg1, xFp16Reg2; + RegTensor xFp32Reg1, vReg1, vRegTmp1, xFp32Reg2, vReg2, vRegTmp2; + MaskReg maskReg; + for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) { + maskReg = UpdateMask(sreg); + DataCopy(xFp16Reg1, xAddr1 + i * V_LENGTH); + DataCopy(xFp16Reg2, xAddr2 + i * V_LENGTH); + Cast(xFp32Reg1, xFp16Reg1, maskReg); + Cast(xFp32Reg2, xFp16Reg2, maskReg); + Mul(vRegTmp1, xFp32Reg1, xFp32Reg1, maskReg); + Mul(vRegTmp2, xFp32Reg2, xFp32Reg2, maskReg); + Muls(vReg1, vRegTmp1, scalar, maskReg); + Muls(vReg2, vRegTmp2, scalar, maskReg); + DataCopy(dstAddr1 + i * V_LENGTH, vReg1, maskReg); + DataCopy(dstAddr2 + i * V_LENGTH, vReg2, maskReg); + } + } + } else if constexpr (IsSameType::value) { + __VEC_SCOPE__ + { + RegTensor xBFp16Reg1, xBFp16Reg2; + RegTensor xFp32Reg1, vReg1, vRegTmp1, xFp32Reg2, vReg2, vRegTmp2; + MaskReg maskReg; + for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) { + maskReg = UpdateMask(sreg); + DataCopy(xBFp16Reg1, xAddr1 + i * V_LENGTH); + DataCopy(xBFp16Reg2, xAddr2 + i * V_LENGTH); + Cast(xFp32Reg1, xBFp16Reg1, maskReg); + Cast(xFp32Reg2, xBFp16Reg2, maskReg); + Mul(vRegTmp1, xFp32Reg1, xFp32Reg1, maskReg); + Mul(vRegTmp2, xFp32Reg2, xFp32Reg2, maskReg); + Muls(vReg1, vRegTmp1, scalar, maskReg); + Muls(vReg2, vRegTmp2, scalar, maskReg); + DataCopy(dstAddr1 + i * V_LENGTH, vReg1, maskReg); + DataCopy(dstAddr2 + i * V_LENGTH, vReg2, maskReg); + } + } + } else { + __VEC_SCOPE__ + { + RegTensor vReg1, vRegTmp1, vReg2, vRegTmp2; + MaskReg maskReg; + for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) { + maskReg = UpdateMask(sreg); + DataCopy(vReg1, xAddr1 + i * V_LENGTH); + DataCopy(vReg2, xAddr2 + i * V_LENGTH); + Mul(vRegTmp1, vReg1, vReg1, maskReg); + Mul(vRegTmp2, vReg2, vReg2, maskReg); + Muls(vReg1, vRegTmp1, scalar, maskReg); + Muls(vReg2, vRegTmp2, scalar, maskReg); + DataCopy(dstAddr1 + i * V_LENGTH, vReg1, maskReg); + DataCopy(dstAddr2 + i * V_LENGTH, vReg2, maskReg); + } + } + } +} + +/*! + * rstd = 1 / sqrt(mean / n + epsilon) + * @param rstdLocal store mean value Tensor + * @param epsilon RmsNorm's attr + * @param count The num of mean + * @return void + */ +__aicore__ inline void ComputeRstd(LocalTensor& rstdLocal, float epsilon, float avgFactor, uint32_t count) +{ + __local_mem__ float* rstdLocalAddr = (__ubuf__ float*)rstdLocal.GetPhyAddr(); + + uint32_t calCount = count; + uint32_t sreg = (uint32_t)calCount; + uint16_t repeatTimes = CeilDivision(calCount, V_LENGTH); + __VEC_SCOPE__ + { + RegTensor vReg, srcReg, dstReg; + MaskReg maskReg; + for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) { + maskReg = UpdateMask(sreg); + DataCopy(srcReg, rstdLocalAddr + i * V_LENGTH); + Muls(srcReg, srcReg, avgFactor, maskReg); + Adds(dstReg, srcReg, epsilon, maskReg); + Sqrt(vReg, dstReg, maskReg); + Duplicate(srcReg, float(1.0), maskReg); + Div(dstReg, srcReg, vReg, maskReg); + DataCopy(rstdLocalAddr + i * V_LENGTH, dstReg, maskReg); + } + } +} + +/*! + * compute yLocal = xLocal * rstd * gammaLocal + * + * @param xLocal input xLocal + * @param gammaLocal input gammaLocal + * @param yLocal output yLocal + * @param rstd input rstd + * @param count vector commpute length + * @return void + */ +template +__aicore__ inline void ComputeY( + LocalTensor& xLocal, LocalTensor& gammaLocal, LocalTensor& yLocal, LocalTensor& rstdLocal, + uint32_t offset, uint32_t count) +{ + uint32_t calCount = count / 2; + uint32_t sreg = (uint32_t)calCount; + uint16_t repeatTimes = CeilDivision(calCount, V_LENGTH); + + __local_mem__ float* xAddr1 = (__ubuf__ float*)xLocal.GetPhyAddr(); + __local_mem__ float* xAddr2 = (__ubuf__ float*)xLocal.GetPhyAddr() + calCount; + __local_mem__ DG* gammaAddr1 = (__ubuf__ DG*)gammaLocal.GetPhyAddr(); + __local_mem__ DG* gammaAddr2 = (__ubuf__ DG*)gammaLocal.GetPhyAddr() + calCount; + __local_mem__ float* rstdAddr = (__ubuf__ float*)rstdLocal.GetPhyAddr(); + __local_mem__ DX* yAddr1 = (__ubuf__ DX*)yLocal.GetPhyAddr(); + __local_mem__ DX* yAddr2 = (__ubuf__ DX*)yLocal.GetPhyAddr() + calCount; + + if constexpr (!IsSameType::value && !IsSameType::value) { + __VEC_SCOPE__ + { + RegTensor yB16Reg1, yB16Reg2; + RegTensor gammaReg1, gammaReg2; + RegTensor rstdReg; + RegTensor xReg1, dst1Reg, gammaFp32Reg1, yReg1; + RegTensor xReg2, dst2Reg, gammaFp32Reg2, yReg2; + MaskReg maskReg; + DataCopy(rstdReg, rstdAddr + offset); + for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) { + maskReg = UpdateMask(sreg); + DataCopy(xReg1, xAddr1 + i * V_LENGTH); + DataCopy(xReg2, xAddr2 + i * V_LENGTH); + DataCopy(gammaReg1, gammaAddr1 + i * V_LENGTH); + DataCopy(gammaReg2, gammaAddr2 + i * V_LENGTH); + Cast(gammaFp32Reg1, gammaReg1, maskReg); + Cast(gammaFp32Reg2, gammaReg2, maskReg); + Mul(dst1Reg, xReg1, rstdReg, maskReg); + Mul(dst2Reg, xReg2, rstdReg, maskReg); + Mul(yReg1, dst1Reg, gammaFp32Reg1, maskReg); + Mul(yReg2, dst2Reg, gammaFp32Reg2, maskReg); + Cast(yB16Reg1, yReg1, maskReg); + Cast(yB16Reg2, yReg2, maskReg); + DataCopy(yAddr1 + i * V_LENGTH, yB16Reg1, maskReg); + DataCopy(yAddr2 + i * V_LENGTH, yB16Reg2, maskReg); + } + } + } else if constexpr (!IsSameType::value and IsSameType::value) { + __VEC_SCOPE__ + { + RegTensor yB16Reg1, yB16Reg2; + RegTensor gammaReg1, gammaReg2; + RegTensor rstdReg; + RegTensor xReg1, dst1Reg, yReg1; + RegTensor xReg2, dst2Reg, yReg2; + MaskReg maskReg; + DataCopy(rstdReg, rstdAddr + offset); + for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) { + maskReg = UpdateMask(sreg); + DataCopy(xReg1, xAddr1 + i * V_LENGTH); + DataCopy(xReg2, xAddr2 + i * V_LENGTH); + DataCopy(gammaReg1, gammaAddr1 + i * V_LENGTH); + DataCopy(gammaReg2, gammaAddr2 + i * V_LENGTH); + Mul(dst1Reg, xReg1, rstdReg, maskReg); + Mul(dst2Reg, xReg2, rstdReg, maskReg); + Mul(yReg1, dst1Reg, gammaReg1, maskReg); + Mul(yReg2, dst2Reg, gammaReg2, maskReg); + Cast(yB16Reg1, yReg1, maskReg); + Cast(yB16Reg2, yReg2, maskReg); + DataCopy(yAddr1 + i * V_LENGTH, yB16Reg1, maskReg); + DataCopy(yAddr2 + i * V_LENGTH, yB16Reg2, maskReg); + } + } + } else { + __VEC_SCOPE__ + { + RegTensor rstdReg; + RegTensor xReg1, gammaReg1, yReg1, vRegTmp1; + RegTensor xReg2, gammaReg2, yReg2, vRegTmp2; + MaskReg maskReg; + DataCopy(rstdReg, rstdAddr + offset); + for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) { + maskReg = UpdateMask(sreg); + DataCopy(xReg1, xAddr1 + i * V_LENGTH); + DataCopy(xReg2, xAddr2 + i * V_LENGTH); + DataCopy(gammaReg1, gammaAddr1 + i * V_LENGTH); + DataCopy(gammaReg2, gammaAddr2 + i * V_LENGTH); + Mul(vRegTmp1, xReg1, rstdReg, maskReg); + Mul(vRegTmp2, xReg2, rstdReg, maskReg); + Mul(yReg1, vRegTmp1, gammaReg1, maskReg); + Mul(yReg2, vRegTmp2, gammaReg2, maskReg); + DataCopy(yAddr1 + i * V_LENGTH, yReg1, maskReg); + DataCopy(yAddr2 + i * V_LENGTH, yReg2, maskReg); + } + } + } +} +/*! + * compute multi N yLocal = xLocal * rstd * gammaLocal + * + * @param xLocal input xLocal + * @param gammaLocal input gammaLocal + * @param yLocal output yLocal + * @param rstd input rstd + * @param count vector commpute length + * @return void + */ +template +__aicore__ inline void ComputeYMultiN( + LocalTensor& xLocal, LocalTensor& gammaLocal, LocalTensor& yLocal, LocalTensor& rstdLocal, + uint32_t offset, uint32_t count, uint32_t curRows) +{ + uint32_t calCount = count / 2; + uint16_t repeatTimes = CeilDivision(calCount, V_LENGTH); + + __local_mem__ float* xAddr1 = (__ubuf__ float*)xLocal.GetPhyAddr(); + __local_mem__ float* xAddr2 = (__ubuf__ float*)xLocal.GetPhyAddr() + calCount; + __local_mem__ DG* gammaAddr1 = (__ubuf__ DG*)gammaLocal.GetPhyAddr(); + __local_mem__ DG* gammaAddr2 = (__ubuf__ DG*)gammaLocal.GetPhyAddr() + calCount; + __local_mem__ float* rstdAddr = (__ubuf__ float*)rstdLocal.GetPhyAddr(); + __local_mem__ DX* yAddr1 = (__ubuf__ DX*)yLocal.GetPhyAddr(); + __local_mem__ DX* yAddr2 = (__ubuf__ DX*)yLocal.GetPhyAddr() + calCount; + + if constexpr (!IsSameType::value && !IsSameType::value) { + __VEC_SCOPE__ + { + for (uint16_t row = 0; row < static_cast(curRows); row++) { + uint32_t sreg = (uint32_t)calCount; + RegTensor yB16Reg1, yB16Reg2; + RegTensor gammaReg1, gammaReg2; + RegTensor rstdReg; + RegTensor xReg1, dst1Reg, gammaFp32Reg1, yReg1; + RegTensor xReg2, dst2Reg, gammaFp32Reg2, yReg2; + MaskReg pregMask; + DataCopy(rstdReg, rstdAddr + offset); + for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) { + pregMask = UpdateMask(sreg); + DataCopy(xReg1, xAddr1 + i * V_LENGTH); + DataCopy(xReg2, xAddr2 + i * V_LENGTH); + DataCopy(gammaReg1, gammaAddr1 + i * V_LENGTH); + DataCopy(gammaReg2, gammaAddr2 + i * V_LENGTH); + Cast(gammaFp32Reg1, gammaReg1, pregMask); + Cast(gammaFp32Reg2, gammaReg2, pregMask); + Mul(dst1Reg, xReg1, rstdReg, pregMask); + Mul(dst2Reg, xReg2, rstdReg, pregMask); + Mul(yReg1, dst1Reg, gammaFp32Reg1, pregMask); + Mul(yReg2, dst2Reg, gammaFp32Reg2, pregMask); + Cast(yB16Reg1, yReg1, pregMask); + Cast(yB16Reg2, yReg2, pregMask); + DataCopy(yAddr1 + i * V_LENGTH, yB16Reg1, pregMask); + DataCopy(yAddr2 + i * V_LENGTH, yB16Reg2, pregMask); + } + offset++; + xAddr1 += count; + xAddr2 += count; + yAddr1 += count; + yAddr2 += count; + } + } + } else if constexpr (!IsSameType::value and IsSameType::value) { + __VEC_SCOPE__ + { + for (uint16_t row = 0; row < static_cast(curRows); row++) { + uint32_t sreg = (uint32_t)calCount; + RegTensor yB16Reg1, yB16Reg2; + RegTensor gammaReg1, gammaReg2; + RegTensor rstdReg; + RegTensor xReg1, dst1Reg, yReg1; + RegTensor xReg2, dst2Reg, yReg2; + MaskReg maskReg; + DataCopy(rstdReg, rstdAddr + offset); + for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) { + maskReg = UpdateMask(sreg); + DataCopy(xReg1, xAddr1 + i * V_LENGTH); + DataCopy(xReg2, xAddr2 + i * V_LENGTH); + DataCopy(gammaReg1, gammaAddr1 + i * V_LENGTH); + DataCopy(gammaReg2, gammaAddr2 + i * V_LENGTH); + Mul(dst1Reg, xReg1, rstdReg, maskReg); + Mul(dst2Reg, xReg2, rstdReg, maskReg); + Mul(yReg1, dst1Reg, gammaReg1, maskReg); + Mul(yReg2, dst2Reg, gammaReg2, maskReg); + Cast(yB16Reg1, yReg1, maskReg); + Cast(yB16Reg2, yReg2, maskReg); + DataCopy(yAddr1 + i * V_LENGTH, yB16Reg1, maskReg); + DataCopy(yAddr2 + i * V_LENGTH, yB16Reg2, maskReg); + } + offset++; + xAddr1 += count; + xAddr2 += count; + yAddr1 += count; + yAddr2 += count; + } + } + } else { + __VEC_SCOPE__ + { + for (uint16_t row = 0; row < static_cast(curRows); row++) { + uint32_t sreg = (uint32_t)calCount; + RegTensor rstdReg; + RegTensor xReg1, gammaReg1, yReg1, vRegTmp1; + RegTensor xReg2, gammaReg2, yReg2, vRegTmp2; + MaskReg maskReg; + DataCopy(rstdReg, rstdAddr + offset); + for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) { + maskReg = UpdateMask(sreg); + DataCopy(xReg1, xAddr1 + i * V_LENGTH); + DataCopy(xReg2, xAddr2 + i * V_LENGTH); + DataCopy(gammaReg1, gammaAddr1 + i * V_LENGTH); + DataCopy(gammaReg2, gammaAddr2 + i * V_LENGTH); + Mul(vRegTmp1, xReg1, rstdReg, maskReg); + Mul(vRegTmp2, xReg2, rstdReg, maskReg); + Mul(yReg1, vRegTmp1, gammaReg1, maskReg); + Mul(yReg2, vRegTmp2, gammaReg2, maskReg); + DataCopy(yAddr1 + i * V_LENGTH, yReg1, maskReg); + DataCopy(yAddr2 + i * V_LENGTH, yReg2, maskReg); + } + offset++; + xAddr1 += count; + xAddr2 += count; + yAddr1 += count; + yAddr2 += count; + } + } + } +} + +/*! + * compute yLocal = xLocal * rstd * gammaLocal + * + * @param xLocal input xLocal + * @param gammaLocal input gammaLocal + * @param yLocal output yLocal + * @param rstd input rstd + * @param count vector commpute length + * @return void + */ +template +__aicore__ inline void ComputeLatterY( + LocalTensor& xLocal, LocalTensor& gammaLocal, LocalTensor& yLocal, LocalTensor& rstdLocal, + uint32_t offset, uint32_t count) +{ + uint32_t calCount = count / 2; + uint32_t sreg = (uint32_t)calCount; + uint16_t repeatTimes = CeilDivision(calCount, V_LENGTH); + + __local_mem__ DX* xAddr1 = (__ubuf__ DX*)xLocal.GetPhyAddr(); + __local_mem__ DX* xAddr2 = (__ubuf__ DX*)xLocal.GetPhyAddr() + calCount; + __local_mem__ DG* gammaAddr1 = (__ubuf__ DG*)gammaLocal.GetPhyAddr(); + __local_mem__ DG* gammaAddr2 = (__ubuf__ DG*)gammaLocal.GetPhyAddr() + calCount; + __local_mem__ float* srcAddr2 = (__ubuf__ float*)rstdLocal.GetPhyAddr(); + __local_mem__ DX* yAddr1 = (__ubuf__ DX*)yLocal.GetPhyAddr(); + __local_mem__ DX* yAddr2 = (__ubuf__ DX*)yLocal.GetPhyAddr() + calCount; + + if constexpr (!IsSameType::value and !IsSameType::value) { + __VEC_SCOPE__ + { + RegTensor xB16Reg1, yB16Reg1, xB16Reg2, yB16Reg2; + RegTensor gammaReg1, gammaReg2; + RegTensor rstdReg; + RegTensor xReg1, dst1Reg, gammaFp32Reg1, yReg1; + RegTensor xReg2, dst2Reg, gammaFp32Reg2, yReg2; + MaskReg maskReg; + DataCopy(rstdReg, srcAddr2 + offset); + for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) { + maskReg = UpdateMask(sreg); + DataCopy(xB16Reg1, xAddr1 + i * V_LENGTH); + DataCopy(xB16Reg2, xAddr2 + i * V_LENGTH); + DataCopy(gammaReg1, gammaAddr1 + i * V_LENGTH); + DataCopy(gammaReg2, gammaAddr2 + i * V_LENGTH); + Cast(gammaFp32Reg1, gammaReg1, maskReg); + Cast(gammaFp32Reg2, gammaReg2, maskReg); + Cast(xReg1, xB16Reg1, maskReg); + Cast(xReg2, xB16Reg2, maskReg); + Mul(dst1Reg, xReg1, rstdReg, maskReg); + Mul(dst2Reg, xReg2, rstdReg, maskReg); + Mul(yReg1, dst1Reg, gammaFp32Reg1, maskReg); + Mul(yReg2, dst2Reg, gammaFp32Reg2, maskReg); + Cast(yB16Reg1, yReg1, maskReg); + Cast(yB16Reg2, yReg2, maskReg); + DataCopy(yAddr1 + i * V_LENGTH, yB16Reg1, maskReg); + DataCopy(yAddr2 + i * V_LENGTH, yB16Reg2, maskReg); + } + } + } else if constexpr (!IsSameType::value and IsSameType::value) { + __VEC_SCOPE__ + { + RegTensor xB16Reg1, yB16Reg1, xB16Reg2, yB16Reg2; + RegTensor rstdReg; + RegTensor xReg1, dst1Reg, gammaFp32Reg1, yReg1; + RegTensor xReg2, dst2Reg, gammaFp32Reg2, yReg2; + MaskReg maskReg; + DataCopy(rstdReg, srcAddr2 + offset); + for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) { + maskReg = UpdateMask(sreg); + DataCopy(xB16Reg1, xAddr1 + i * V_LENGTH); + DataCopy(xB16Reg2, xAddr2 + i * V_LENGTH); + DataCopy(gammaFp32Reg1, gammaAddr1 + i * V_LENGTH); + DataCopy(gammaFp32Reg2, gammaAddr2 + i * V_LENGTH); + Cast(xReg1, xB16Reg1, maskReg); + Cast(xReg2, xB16Reg2, maskReg); + Mul(dst1Reg, xReg1, rstdReg, maskReg); + Mul(dst2Reg, xReg2, rstdReg, maskReg); + Mul(yReg1, dst1Reg, gammaFp32Reg1, maskReg); + Mul(yReg2, dst2Reg, gammaFp32Reg2, maskReg); + Cast(yB16Reg1, yReg1, maskReg); + Cast(yB16Reg2, yReg2, maskReg); + DataCopy(yAddr1 + i * V_LENGTH, yB16Reg1, maskReg); + DataCopy(yAddr2 + i * V_LENGTH, yB16Reg2, maskReg); + } + } + } else { + __VEC_SCOPE__ + { + RegTensor rstdReg; + RegTensor xReg1, gammaReg1, yReg1, vRegTmp1; + RegTensor xReg2, gammaReg2, yReg2, vRegTmp2; + MaskReg maskReg; + DataCopy(rstdReg, srcAddr2 + offset); + for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) { + maskReg = UpdateMask(sreg); + DataCopy(xReg1, xAddr1 + i * V_LENGTH); + DataCopy(xReg2, xAddr2 + i * V_LENGTH); + DataCopy(gammaReg1, gammaAddr1 + i * V_LENGTH); + DataCopy(gammaReg2, gammaAddr2 + i * V_LENGTH); + Mul(vRegTmp1, xReg1, rstdReg, maskReg); + Mul(vRegTmp2, xReg2, rstdReg, maskReg); + Mul(yReg1, vRegTmp1, gammaReg1, maskReg); + Mul(yReg2, vRegTmp2, gammaReg2, maskReg); + DataCopy(yAddr1 + i * V_LENGTH, yReg1, maskReg); + DataCopy(yAddr2 + i * V_LENGTH, yReg2, maskReg); + } + } + } +} + +/*! + * The num of each level elements is 256, ReduceSum these elements and store to the next level. + * @param level1Local level1Tensor + * @param level2Local level2Tensor + * @param level3Local level3Tensor + * @param level1 level1 elements + * @param level2 level2 elements + * @param level3 level3 elements + * @return void + */ +__aicore__ inline void ComputeMultiLevelReduce( + LocalTensor& level1Local, LocalTensor& level2Local, LocalTensor& level3Local, uint32_t& level1, + uint32_t& level2, uint32_t& level3) +{ + if (level1 == NormCommon::ONCE_VECTOR_SIZE) { + LevelMergeRstd(level2Local, level1Local, level2, NormCommon::ONCE_VECTOR_SIZE); + level1 = 0; + level2 += 1; + } + if (level2 == NormCommon::ONCE_VECTOR_SIZE) { + LevelMergeRstd(level3Local, level2Local, level3, NormCommon::ONCE_VECTOR_SIZE); + level2 = 0; + level3 += 1; + } +} + +__aicore__ inline void ComputeSum( + LocalTensor& dstLocal, LocalTensor& srcLocal, uint32_t offset, uint32_t count) +{ + uint32_t meanTile = count; + uint32_t meanSreg = meanTile; + + __local_mem__ float* srcAddr = (__ubuf__ float*)srcLocal.GetPhyAddr(); + __local_mem__ float* dstAddr = (__ubuf__ float*)dstLocal.GetPhyAddr(); + + __VEC_SCOPE__ + { + RegTensor vReg, vMean; + MaskReg pregLoop; + MaskReg pregMerge = CreateMask(); + { + pregLoop = UpdateMask(meanSreg); + DataCopy(vReg, srcAddr + 0); + ReduceSum(vMean, vReg, pregLoop); + DataCopy(dstAddr + offset, vMean, pregMerge); + } + } +} + +/*! + * ReduceSum impl by half add. + * @param dstLocal dst Tensor + * @param srcLocal src Tensor + * @param workLocal temp Tensor + * @param offset dst offset + * @param count count aligned compute elements. + * @param powerSplit 2 ** k = powerSplit + * @return void + */ +__aicore__ inline void ReduceSumImpl( + LocalTensor& dstLocal, LocalTensor& srcLocal, LocalTensor& workLocal, uint32_t offset, + uint32_t count, uint32_t powerSplit) +{ + uint32_t remainTile = count - powerSplit; + uint32_t remainSreg = remainTile; + uint32_t remainRepeats = remainTile / (2 * V_LENGTH); + + uint32_t masterTile = powerSplit - remainTile; + uint32_t masterSreg = masterTile; + uint16_t masterRepeats = masterTile / (2 * V_LENGTH); + + uint32_t mergeTile = powerSplit / (2 * V_LENGTH); + uint32_t mergeSreg = mergeTile; + uint32_t mergeRepeats = mergeTile / (2 * V_LENGTH); + + uint32_t meanTile = mergeRepeats == 0 ? mergeTile : mergeRepeats; + uint32_t meanSreg = meanTile; + + __local_mem__ float* mainAddr = (__ubuf__ float*)srcLocal.GetPhyAddr(); + __local_mem__ float* tailAddr = (__ubuf__ float*)srcLocal.GetPhyAddr() + int64_t(powerSplit); + __local_mem__ float* masterAddr = (__ubuf__ float*)srcLocal.GetPhyAddr() + int64_t(remainTile); + __local_mem__ float* workAddr = (__ubuf__ float*)workLocal.GetPhyAddr(); + __local_mem__ float* dstAddr = (__ubuf__ float*)dstLocal.GetPhyAddr(); + + __VEC_SCOPE__ + { + RegTensor mainA, mainB, tailA, tailB, vMean; + MaskReg pregMerge = CreateMask(); + MaskReg pregLoop; + + for (uint16_t i = 0; i < (uint16_t)remainRepeats; ++i) { + pregLoop = UpdateMask(remainSreg); + DataCopy(mainA, mainAddr + (i * 2 + 0) * V_LENGTH); + DataCopy(mainB, mainAddr + (i * 2 + 1) * V_LENGTH); + DataCopy(tailA, tailAddr + (i * 2 + 0) * V_LENGTH); + DataCopy(tailB, tailAddr + (i * 2 + 1) * V_LENGTH); + + Add(mainA, mainA, tailA, pregLoop); + Add(mainB, mainB, tailB, pregLoop); + Add(mainA, mainA, mainB, pregLoop); + ReduceSum(vMean, mainA, pregLoop); + DataCopy(workAddr + i, vMean, pregMerge); + } + for (uint16_t i = 0; i < (uint16_t)masterRepeats; ++i) { + pregLoop = UpdateMask(masterSreg); + DataCopy(mainA, masterAddr + (i * 2 + 0) * V_LENGTH); + DataCopy(mainB, masterAddr + (i * 2 + 1) * V_LENGTH); + Add(mainA, mainA, mainB, pregLoop); + ReduceSum(vMean, mainA, pregLoop); + DataCopy(workAddr + remainRepeats + i, vMean, pregMerge); + } + LocalMemBar(); + for (uint16_t i = 0; i < (uint16_t)mergeRepeats; ++i) { + pregLoop = UpdateMask(mergeSreg); + DataCopy(mainA, workAddr + (i * 2 + 0) * V_LENGTH); + DataCopy(mainB, workAddr + (i * 2 + 1) * V_LENGTH); + Add(mainA, mainA, mainB, pregLoop); + ReduceSum(vMean, mainA, pregLoop); + DataCopy(workAddr + i, vMean, pregMerge); + } + LocalMemBar(); + { + pregLoop = UpdateMask(meanSreg); + DataCopy(mainA, workAddr + 0); + ReduceSum(vMean, mainA, pregLoop); + DataCopy(dstAddr + offset, vMean, pregMerge); + } + } +} + +template +__aicore__ inline void LoadForHandleRemainV1( + __local_mem__ T* mainAddr, __local_mem__ T* tailAddr, uint16_t offset1, uint16_t offset2, RegTensor& mainA, + RegTensor& mainB, RegTensor& tailA, RegTensor& tailB, MaskReg& pregLoop, + __local_mem__ float* xFp32MainAddr, __local_mem__ float* xFp32TailAddr) +{ + if constexpr (IsSameType::value) { + RegTensor xFp16MainA, xFp16MainB, xFp16TailA, xFp16TailB; + DataCopy(xFp16MainA, mainAddr + offset1); + DataCopy(xFp16MainB, mainAddr + offset2); + DataCopy(xFp16TailA, tailAddr + offset1); + DataCopy(xFp16TailB, tailAddr + offset2); + Cast(mainA, xFp16MainA, pregLoop); + Cast(mainB, xFp16MainB, pregLoop); + Cast(tailA, xFp16TailA, pregLoop); + Cast(tailB, xFp16TailB, pregLoop); + DataCopy(xFp32MainAddr + offset1, mainA, pregLoop); + DataCopy(xFp32MainAddr + offset2, mainB, pregLoop); + DataCopy(xFp32TailAddr + offset1, tailA, pregLoop); + DataCopy(xFp32TailAddr + offset2, tailB, pregLoop); + Mul(mainA, mainA, mainA, pregLoop); + Mul(mainB, mainB, mainB, pregLoop); + Mul(tailA, tailA, tailA, pregLoop); + Mul(tailB, tailB, tailB, pregLoop); + } else if constexpr (IsSameType::value) { + RegTensor xBFp16MainA, xBFp16MainB, xBFp16TailA, xBFp16TailB; + DataCopy(xBFp16MainA, mainAddr + offset1); + DataCopy(xBFp16MainB, mainAddr + offset2); + DataCopy(xBFp16TailA, tailAddr + offset1); + DataCopy(xBFp16TailB, tailAddr + offset2); + Cast(mainA, xBFp16MainA, pregLoop); + Cast(mainB, xBFp16MainB, pregLoop); + Cast(tailA, xBFp16TailA, pregLoop); + Cast(tailB, xBFp16TailB, pregLoop); + DataCopy(xFp32MainAddr + offset1, mainA, pregLoop); + DataCopy(xFp32MainAddr + offset2, mainB, pregLoop); + DataCopy(xFp32TailAddr + offset1, tailA, pregLoop); + DataCopy(xFp32TailAddr + offset2, tailB, pregLoop); + Mul(mainA, mainA, mainA, pregLoop); + Mul(mainB, mainB, mainB, pregLoop); + Mul(tailA, tailA, tailA, pregLoop); + Mul(tailB, tailB, tailB, pregLoop); + } else { + DataCopy(mainA, mainAddr + offset1); + DataCopy(mainB, mainAddr + offset2); + DataCopy(tailA, tailAddr + offset1); + DataCopy(tailB, tailAddr + offset2); + Mul(mainA, mainA, mainA, pregLoop); + Mul(mainB, mainB, mainB, pregLoop); + Mul(tailA, tailA, tailA, pregLoop); + Mul(tailB, tailB, tailB, pregLoop); + } +} + +template +__aicore__ inline void LoadForHandleMasterV1( + __local_mem__ T* masterAddr, uint16_t offset1, uint16_t offset2, RegTensor& mainA, RegTensor& mainB, + MaskReg& pregLoop, __local_mem__ float* xFp32MasterAddr) +{ + if constexpr (IsSameType::value) { + RegTensor xFp16MainA, xFp16MainB; + DataCopy(xFp16MainA, masterAddr + offset1); + DataCopy(xFp16MainB, masterAddr + offset2); + Cast(mainA, xFp16MainA, pregLoop); + Cast(mainB, xFp16MainB, pregLoop); + DataCopy(xFp32MasterAddr + offset1, mainA, pregLoop); + DataCopy(xFp32MasterAddr + offset2, mainB, pregLoop); + Mul(mainA, mainA, mainA, pregLoop); + Mul(mainB, mainB, mainB, pregLoop); + } else if constexpr (IsSameType::value) { + RegTensor xBFp16MainA, xBFp16MainB; + DataCopy(xBFp16MainA, masterAddr + offset1); + DataCopy(xBFp16MainB, masterAddr + offset2); + Cast(mainA, xBFp16MainA, pregLoop); + Cast(mainB, xBFp16MainB, pregLoop); + DataCopy(xFp32MasterAddr + offset1, mainA, pregLoop); + DataCopy(xFp32MasterAddr + offset2, mainB, pregLoop); + Mul(mainA, mainA, mainA, pregLoop); + Mul(mainB, mainB, mainB, pregLoop); + } else { + DataCopy(mainA, masterAddr + offset1); + DataCopy(mainB, masterAddr + offset2); + Mul(mainA, mainA, mainA, pregLoop); + Mul(mainB, mainB, mainB, pregLoop); + } +} + +template +__aicore__ inline void ComputeFormerImplV1( + LocalTensor& xLocal, LocalTensor& xFp32, LocalTensor& workLocal, LocalTensor& rstdLocal, + float avgFactor, float epsilon, uint32_t offset, uint32_t count, uint32_t powerSplit) +{ + uint32_t remainTile = count - powerSplit; + uint32_t remainSreg = remainTile; + uint16_t remainRepeats = remainTile / (2 * V_LENGTH); + + uint32_t masterTile = powerSplit - remainTile; + uint32_t masterSreg = masterTile; + uint16_t masterRepeats = masterTile / (2 * V_LENGTH); + + uint32_t mergeTile = powerSplit / (2 * V_LENGTH); + uint32_t mergeSreg = mergeTile; + uint16_t mergeRepeats = mergeTile / (2 * V_LENGTH); + + uint32_t meanTile = mergeRepeats == 0 ? mergeTile : mergeRepeats; + uint32_t meanSreg = meanTile; + + __local_mem__ T* mainAddr = (__ubuf__ T*)xLocal.GetPhyAddr(); + __local_mem__ T* tailAddr = (__ubuf__ T*)xLocal.GetPhyAddr() + int64_t(powerSplit); + __local_mem__ T* masterAddr = (__ubuf__ T*)xLocal.GetPhyAddr() + int64_t(remainTile); + __local_mem__ float *xFp32MainAddr, *xFp32TailAddr, *xFp32MasterAddr; + if constexpr (is_same::value || is_same::value) { + xFp32MainAddr = (__ubuf__ float*)xFp32.GetPhyAddr(); + xFp32TailAddr = (__ubuf__ float*)xFp32.GetPhyAddr() + int64_t(powerSplit); + xFp32MasterAddr = (__ubuf__ float*)xFp32.GetPhyAddr() + int64_t(remainTile); + } + __local_mem__ float* workAddr = (__ubuf__ float*)workLocal.GetPhyAddr(); + __local_mem__ float* rstdAddr = (__ubuf__ float*)rstdLocal.GetPhyAddr(); + + __VEC_SCOPE__ + { + RegTensor mainA, mainB, tailA, tailB, vMean, vDupReg, rstdReg; + MaskReg pregMerge = CreateMask(); + MaskReg pregLoop; + + for (uint16_t i = 0; i < (uint16_t)remainRepeats; ++i) { + pregLoop = UpdateMask(remainSreg); + LoadForHandleRemainV1( + mainAddr, tailAddr, (i * 2 + 0) * V_LENGTH, (i * 2 + 1) * V_LENGTH, mainA, mainB, tailA, tailB, + pregLoop, xFp32MainAddr, xFp32TailAddr); + Add(mainA, mainA, tailA, pregLoop); + Add(mainB, mainB, tailB, pregLoop); + Add(mainA, mainA, mainB, pregLoop); + ReduceSum(vMean, mainA, pregLoop); + DataCopy(workAddr + i, vMean, pregMerge); + } + for (uint16_t i = 0; i < (uint16_t)masterRepeats; ++i) { + pregLoop = UpdateMask(masterSreg); + LoadForHandleMasterV1( + masterAddr, (i * 2 + 0) * V_LENGTH, (i * 2 + 1) * V_LENGTH, mainA, mainB, pregLoop, xFp32MasterAddr); + Add(mainA, mainA, mainB, pregLoop); + ReduceSum(vMean, mainA, pregLoop); + DataCopy(workAddr + remainRepeats + i, vMean, pregMerge); + } + LocalMemBar(); + for (uint16_t i = 0; i < (uint16_t)mergeRepeats; ++i) { + pregLoop = UpdateMask(mergeSreg); + DataCopy(mainA, workAddr + (i * 2 + 0) * V_LENGTH); + DataCopy(mainB, workAddr + (i * 2 + 1) * V_LENGTH); + Add(mainA, mainA, mainB, pregLoop); + ReduceSum(vMean, mainA, pregLoop); + DataCopy(workAddr + i, vMean, pregMerge); + } + LocalMemBar(); + { + pregLoop = UpdateMask(meanSreg); + DataCopy(mainA, workAddr + 0); + ReduceSum(vMean, mainA, pregLoop); + Muls(vMean, vMean, avgFactor, pregMerge); + Adds(vMean, vMean, epsilon, pregMerge); + Sqrt(vMean, vMean, pregMerge); + Duplicate(vDupReg, float(1.0), pregMerge); + Div(rstdReg, vDupReg, vMean, pregMerge); + DataCopy(rstdAddr + offset, rstdReg, pregMerge); + } + } +} + +template +__aicore__ inline void ComputeFormerImplV1MultiN( + LocalTensor& xLocal, LocalTensor& xFp32, LocalTensor& workLocal, LocalTensor& rstdLocal, + float avgFactor, float epsilon, uint32_t offset, uint32_t count, uint32_t powerSplit, uint32_t curRows) +{ + uint32_t remainTile = count - powerSplit; + uint16_t remainRepeats = remainTile / (2 * V_LENGTH); + + uint32_t masterTile = powerSplit - remainTile; + uint16_t masterRepeats = masterTile / (2 * V_LENGTH); + + uint32_t mergeTile = powerSplit / (2 * V_LENGTH); + uint16_t mergeRepeats = mergeTile / (2 * V_LENGTH); + + uint32_t meanTile = mergeRepeats == 0 ? mergeTile : mergeRepeats; + + __local_mem__ T* mainAddr = (__ubuf__ T*)xLocal.GetPhyAddr(); + __local_mem__ T* tailAddr = (__ubuf__ T*)xLocal.GetPhyAddr() + int64_t(powerSplit); + __local_mem__ T* masterAddr = (__ubuf__ T*)xLocal.GetPhyAddr() + int64_t(remainTile); + __local_mem__ float *xFp32MainAddr, *xFp32TailAddr, *xFp32MasterAddr; + if constexpr (is_same::value || is_same::value) { + xFp32MainAddr = (__ubuf__ float*)xFp32.GetPhyAddr(); + xFp32TailAddr = (__ubuf__ float*)xFp32.GetPhyAddr() + int64_t(powerSplit); + xFp32MasterAddr = (__ubuf__ float*)xFp32.GetPhyAddr() + int64_t(remainTile); + } + uint32_t curRowsAlign = CeilDiv((int32_t)curRows, 2); + int64_t unrollOffset = (curRows / 2) * count; + bool isWithTail = curRowsAlign - (curRows / 2); + uint32_t tailOffset = offset + curRows / 2; + + __local_mem__ T* mainAddr1 = (__ubuf__ T*)xLocal.GetPhyAddr() + unrollOffset; + __local_mem__ T* tailAddr1 = (__ubuf__ T*)xLocal.GetPhyAddr() + int64_t(powerSplit) + unrollOffset; + __local_mem__ T* masterAddr1 = (__ubuf__ T*)xLocal.GetPhyAddr() + int64_t(remainTile) + unrollOffset; + __local_mem__ float *xFp32MainAddr1, *xFp32TailAddr1, *xFp32MasterAddr1; + if constexpr (is_same::value || is_same::value) { + xFp32MainAddr1 = (__ubuf__ float*)xFp32.GetPhyAddr() + unrollOffset; + xFp32TailAddr1 = (__ubuf__ float*)xFp32.GetPhyAddr() + int64_t(powerSplit) + unrollOffset; + xFp32MasterAddr1 = (__ubuf__ float*)xFp32.GetPhyAddr() + int64_t(remainTile) + unrollOffset; + } + + __local_mem__ float* workAddr = (__ubuf__ float*)workLocal.GetPhyAddr(); + __local_mem__ float* rstdAddr = (__ubuf__ float*)rstdLocal.GetPhyAddr(); + + __local_mem__ float* workAddr1 = + (__ubuf__ float*)workLocal.GetPhyAddr() + NormCommon::ONCE_VECTOR_SIZE; + __local_mem__ float* rstdAddr1 = (__ubuf__ float*)rstdLocal.GetPhyAddr() + curRows / 2; + + __VEC_SCOPE__ + { + for (uint16_t row = 0; row < static_cast(curRows / 2); row++) { + uint32_t remainSreg = remainTile; + uint32_t masterSreg = masterTile; + uint32_t mergeSreg = mergeTile; + uint32_t meanSreg = meanTile; + RegTensor mainA, mainB, tailA, tailB, vMean, vDupReg, rstdReg; + MaskReg pregMerge = CreateMask(); + MaskReg pregLoop; + + uint32_t remainSreg1 = remainTile; + uint32_t masterSreg1 = masterTile; + uint32_t mergeSreg1 = mergeTile; + uint32_t meanSreg1 = meanTile; + RegTensor mainA1, mainB1, tailA1, tailB1, vMean1, vDupReg1, rstdReg1; + MaskReg pregMain1 = CreateMask(); + MaskReg pregMerge1 = CreateMask(); + MaskReg pregLoop1; + + for (uint16_t i = 0; i < (uint16_t)remainRepeats; ++i) { + pregLoop = UpdateMask(remainSreg); + LoadForHandleRemainV1( + mainAddr, tailAddr, (i * 2 + 0) * V_LENGTH, (i * 2 + 1) * V_LENGTH, mainA, mainB, tailA, tailB, + pregLoop, xFp32MainAddr, xFp32TailAddr); + Add(mainA, mainA, tailA, pregLoop); + Add(mainB, mainB, tailB, pregLoop); + Add(mainA, mainA, mainB, pregLoop); + ReduceSum(vMean, mainA, pregLoop); + DataCopy(workAddr + i, vMean, pregMerge); + } + for (uint16_t i = 0; i < (uint16_t)masterRepeats; ++i) { + pregLoop = UpdateMask(masterSreg); + LoadForHandleMasterV1( + masterAddr, (i * 2 + 0) * V_LENGTH, (i * 2 + 1) * V_LENGTH, mainA, mainB, pregLoop, + xFp32MasterAddr); + Add(mainA, mainA, mainB, pregLoop); + ReduceSum(vMean, mainA, pregLoop); + DataCopy(workAddr + remainRepeats + i, vMean, pregMerge); + } + // unroll part + for (uint16_t i = 0; i < (uint16_t)remainRepeats; ++i) { + pregLoop1 = UpdateMask(remainSreg1); + LoadForHandleRemainV1( + mainAddr1, tailAddr1, (i * 2 + 0) * V_LENGTH, (i * 2 + 1) * V_LENGTH, mainA1, mainB1, tailA1, + tailB1, pregLoop1, xFp32MainAddr1, xFp32TailAddr1); + Add(mainA1, mainA1, tailA1, pregLoop1); + Add(mainB1, mainB1, tailB1, pregLoop1); + Add(mainA1, mainA1, mainB1, pregLoop1); + ReduceSum(vMean1, mainA1, pregLoop1); + DataCopy(workAddr1 + i, vMean1, pregMerge1); + } + for (uint16_t i = 0; i < (uint16_t)masterRepeats; ++i) { + pregLoop1 = UpdateMask(masterSreg1); + LoadForHandleMasterV1( + masterAddr1, (i * 2 + 0) * V_LENGTH, (i * 2 + 1) * V_LENGTH, mainA1, mainB1, pregLoop1, + xFp32MasterAddr1); + Add(mainA1, mainA1, mainB1, pregLoop1); + ReduceSum(vMean1, mainA1, pregLoop1); + DataCopy(workAddr1 + remainRepeats + i, vMean1, pregMerge1); + } + LocalMemBar(); + for (uint16_t i = 0; i < (uint16_t)mergeRepeats; ++i) { + pregLoop = UpdateMask(mergeSreg); + DataCopy(mainA, workAddr + (i * 2 + 0) * V_LENGTH); + DataCopy(mainB, workAddr + (i * 2 + 1) * V_LENGTH); + Add(mainA, mainA, mainB, pregLoop); + ReduceSum(vMean, mainA, pregLoop); + DataCopy(workAddr + i, vMean, pregMerge); + } + // unroll part + for (uint16_t i = 0; i < (uint16_t)mergeRepeats; ++i) { + pregLoop1 = UpdateMask(mergeSreg1); + DataCopy(mainA1, workAddr1 + (i * 2 + 0) * V_LENGTH); + DataCopy(mainB1, workAddr1 + (i * 2 + 1) * V_LENGTH); + Add(mainA1, mainA1, mainB1, pregLoop1); + ReduceSum(vMean1, mainA1, pregLoop1); + DataCopy(workAddr1 + i, vMean1, pregMerge1); + } + LocalMemBar(); + { + pregLoop = UpdateMask(meanSreg); + DataCopy(mainA, workAddr + 0); + ReduceSum(vMean, mainA, pregLoop); + Muls(vMean, vMean, avgFactor, pregMerge); + Adds(vMean, vMean, epsilon, pregMerge); + Sqrt(vMean, vMean, pregMerge); + Duplicate(vDupReg, float(1.0), pregMerge); + Div(rstdReg, vDupReg, vMean, pregMerge); + DataCopy(rstdAddr + offset, rstdReg, pregMerge); + } + // unroll part + { + pregLoop1 = UpdateMask(meanSreg1); + DataCopy(mainA1, workAddr1 + 0); + ReduceSum(vMean1, mainA1, pregLoop1); + Muls(vMean1, vMean1, avgFactor, pregMerge1); + Adds(vMean1, vMean1, epsilon, pregMerge1); + Sqrt(vMean1, vMean1, pregMerge1); + Duplicate(vDupReg1, float(1.0), pregMerge1); + Div(rstdReg1, vDupReg1, vMean1, pregMerge1); + DataCopy(rstdAddr1 + offset, rstdReg1, pregMerge1); + } + offset += 1; + mainAddr += int64_t(count); + tailAddr += int64_t(count); + masterAddr += int64_t(count); + xFp32MainAddr += int64_t(count); + xFp32TailAddr += int64_t(count); + xFp32MasterAddr += int64_t(count); + + mainAddr1 += int64_t(count); + tailAddr1 += int64_t(count); + masterAddr1 += int64_t(count); + xFp32MainAddr1 += int64_t(count); + xFp32TailAddr1 += int64_t(count); + xFp32MasterAddr1 += int64_t(count); + } + } + uint32_t tailDataOffset = unrollOffset + (curRows / 2) * count; + __local_mem__ T* mainAddr2 = (__ubuf__ T*)xLocal.GetPhyAddr() + tailDataOffset; + __local_mem__ T* tailAddr2 = (__ubuf__ T*)xLocal.GetPhyAddr() + int64_t(powerSplit) + tailDataOffset; + __local_mem__ T* masterAddr2 = (__ubuf__ T*)xLocal.GetPhyAddr() + int64_t(remainTile) + tailDataOffset; + __local_mem__ float *xFp32MainAddr2, *xFp32TailAddr2, *xFp32MasterAddr2; + if constexpr (is_same::value || is_same::value) { + xFp32MainAddr2 = (__ubuf__ float*)xFp32.GetPhyAddr() + tailDataOffset; + xFp32TailAddr2 = (__ubuf__ float*)xFp32.GetPhyAddr() + int64_t(powerSplit) + tailDataOffset; + xFp32MasterAddr2 = (__ubuf__ float*)xFp32.GetPhyAddr() + int64_t(remainTile) + tailDataOffset; + } + if (isWithTail) { + __VEC_SCOPE__ + { + uint32_t remainSreg1 = remainTile; + uint32_t masterSreg1 = masterTile; + uint32_t mergeSreg1 = mergeTile; + uint32_t meanSreg1 = meanTile; + RegTensor mainA1, mainB1, tailA1, tailB1, vMean1, vDupReg1, rstdReg1; + MaskReg pregMain1 = CreateMask(); + MaskReg pregMerge1 = CreateMask(); + MaskReg pregLoop1; + + for (uint16_t i = 0; i < (uint16_t)remainRepeats; ++i) { + pregLoop1 = UpdateMask(remainSreg1); + LoadForHandleRemainV1( + mainAddr2, tailAddr2, (i * 2 + 0) * V_LENGTH, (i * 2 + 1) * V_LENGTH, mainA1, mainB1, tailA1, + tailB1, pregLoop1, xFp32MainAddr2, xFp32TailAddr2); + Add(mainA1, mainA1, tailA1, pregLoop1); + Add(mainB1, mainB1, tailB1, pregLoop1); + Add(mainA1, mainA1, mainB1, pregLoop1); + ReduceSum(vMean1, mainA1, pregLoop1); + DataCopy(workAddr1 + i, vMean1, pregMerge1); + } + for (uint16_t i = 0; i < (uint16_t)masterRepeats; ++i) { + pregLoop1 = UpdateMask(masterSreg1); + LoadForHandleMasterV1( + masterAddr2, (i * 2 + 0) * V_LENGTH, (i * 2 + 1) * V_LENGTH, mainA1, mainB1, pregLoop1, + xFp32MasterAddr2); + Add(mainA1, mainA1, mainB1, pregLoop1); + ReduceSum(vMean1, mainA1, pregLoop1); + DataCopy(workAddr1 + remainRepeats + i, vMean1, pregMerge1); + } + LocalMemBar(); + for (uint16_t i = 0; i < (uint16_t)mergeRepeats; ++i) { + pregLoop1 = UpdateMask(mergeSreg1); + DataCopy(mainA1, workAddr1 + (i * 2 + 0) * V_LENGTH); + DataCopy(mainB1, workAddr1 + (i * 2 + 1) * V_LENGTH); + Add(mainA1, mainA1, mainB1, pregLoop1); + ReduceSum(vMean1, mainA1, pregLoop1); + DataCopy(workAddr1 + i, vMean1, pregMerge1); + } + LocalMemBar(); + { + pregLoop1 = UpdateMask(meanSreg1); + DataCopy(mainA1, workAddr1 + 0); + ReduceSum(vMean1, mainA1, pregLoop1); + Muls(vMean1, vMean1, avgFactor, pregMerge1); + Adds(vMean1, vMean1, epsilon, pregMerge1); + Sqrt(vMean1, vMean1, pregMerge1); + Duplicate(vDupReg1, float(1.0), pregMerge1); + Div(rstdReg1, vDupReg1, vMean1, pregMerge1); + DataCopy(rstdAddr1 + tailOffset, rstdReg1, pregMerge1); + } + } + } +} + +template +__aicore__ inline void LoadForHandleRemainV2( + __local_mem__ T* mainAddr, __local_mem__ T* tailAddr, uint16_t offset1, uint16_t offset2, RegTensor& mainA, + RegTensor& mainB, RegTensor& tailA, RegTensor& tailB, MaskReg& pregLoop) +{ + if constexpr (IsSameType::value) { + RegTensor xFp16MainA, xFp16MainB, xFp16TailA, xFp16TailB; + DataCopy(xFp16MainA, mainAddr + offset1); + DataCopy(xFp16MainB, mainAddr + offset2); + DataCopy(xFp16TailA, tailAddr + offset1); + DataCopy(xFp16TailB, tailAddr + offset2); + Cast(mainA, xFp16MainA, pregLoop); + Cast(mainB, xFp16MainB, pregLoop); + Cast(tailA, xFp16TailA, pregLoop); + Cast(tailB, xFp16TailB, pregLoop); + Mul(mainA, mainA, mainA, pregLoop); + Mul(mainB, mainB, mainB, pregLoop); + Mul(tailA, tailA, tailA, pregLoop); + Mul(tailB, tailB, tailB, pregLoop); + } else if constexpr (IsSameType::value) { + RegTensor xBFp16MainA, xBFp16MainB, xBFp16TailA, xBFp16TailB; + DataCopy(xBFp16MainA, mainAddr + offset1); + DataCopy(xBFp16MainB, mainAddr + offset2); + DataCopy(xBFp16TailA, tailAddr + offset1); + DataCopy(xBFp16TailB, tailAddr + offset2); + Cast(mainA, xBFp16MainA, pregLoop); + Cast(mainB, xBFp16MainB, pregLoop); + Cast(tailA, xBFp16TailA, pregLoop); + Cast(tailB, xBFp16TailB, pregLoop); + Mul(mainA, mainA, mainA, pregLoop); + Mul(mainB, mainB, mainB, pregLoop); + Mul(tailA, tailA, tailA, pregLoop); + Mul(tailB, tailB, tailB, pregLoop); + } else { + DataCopy(mainA, mainAddr + offset1); + DataCopy(mainB, mainAddr + offset2); + DataCopy(tailA, tailAddr + offset1); + DataCopy(tailB, tailAddr + offset2); + Mul(mainA, mainA, mainA, pregLoop); + Mul(mainB, mainB, mainB, pregLoop); + Mul(tailA, tailA, tailA, pregLoop); + Mul(tailB, tailB, tailB, pregLoop); + } +} + +template +__aicore__ inline void LoadForHandleMasterV2( + __local_mem__ T* masterAddr, uint16_t offset1, uint16_t offset2, RegTensor& mainA, RegTensor& mainB, + MaskReg& pregLoop) +{ + if constexpr (IsSameType::value) { + RegTensor xFp16MainA, xFp16MainB; + DataCopy(xFp16MainA, masterAddr + offset1); + DataCopy(xFp16MainB, masterAddr + offset2); + Cast(mainA, xFp16MainA, pregLoop); + Cast(mainB, xFp16MainB, pregLoop); + Mul(mainA, mainA, mainA, pregLoop); + Mul(mainB, mainB, mainB, pregLoop); + } else if constexpr (IsSameType::value) { + RegTensor xBFp16MainA, xBFp16MainB; + DataCopy(xBFp16MainA, masterAddr + offset1); + DataCopy(xBFp16MainB, masterAddr + offset2); + Cast(mainA, xBFp16MainA, pregLoop); + Cast(mainB, xBFp16MainB, pregLoop); + Mul(mainA, mainA, mainA, pregLoop); + Mul(mainB, mainB, mainB, pregLoop); + } else { + DataCopy(mainA, masterAddr + offset1); + DataCopy(mainB, masterAddr + offset2); + Mul(mainA, mainA, mainA, pregLoop); + Mul(mainB, mainB, mainB, pregLoop); + } +} + +template +__aicore__ inline void ComputeFormerImplV2( + LocalTensor& dstLocal, LocalTensor& xLocal, LocalTensor& workLocal, uint32_t offset, + uint32_t count, uint32_t powerSplit) +{ + uint32_t remainTile = count - powerSplit; + uint32_t remainSreg = remainTile; + uint16_t remainRepeats = remainTile / (2 * V_LENGTH); + + uint32_t masterTile = powerSplit - remainTile; + uint32_t masterSreg = masterTile; + uint16_t masterRepeats = masterTile / (2 * V_LENGTH); + + uint32_t mergeTile = powerSplit / (2 * V_LENGTH); + uint32_t mergeSreg = mergeTile; + uint16_t mergeRepeats = mergeTile / (2 * V_LENGTH); + + uint32_t meanTile = mergeRepeats == 0 ? mergeTile : mergeRepeats; + uint32_t meanSreg = meanTile; + + __local_mem__ T* mainAddr = (__ubuf__ T*)xLocal.GetPhyAddr(); + __local_mem__ T* tailAddr = (__ubuf__ T*)xLocal.GetPhyAddr() + int64_t(powerSplit); + __local_mem__ T* masterAddr = (__ubuf__ T*)xLocal.GetPhyAddr() + int64_t(remainTile); + + __local_mem__ float* workAddr = (__ubuf__ float*)workLocal.GetPhyAddr(); + __local_mem__ float* dstAddr = (__ubuf__ float*)dstLocal.GetPhyAddr(); + + __VEC_SCOPE__ + { + RegTensor mainA, mainB, tailA, tailB, vMean, vDupReg, rstdReg; + MaskReg pregMerge = CreateMask(); + MaskReg pregLoop; + + for (uint16_t i = 0; i < (uint16_t)remainRepeats; ++i) { + pregLoop = UpdateMask(remainSreg); + LoadForHandleRemainV2( + mainAddr, tailAddr, (i * 2 + 0) * V_LENGTH, (i * 2 + 1) * V_LENGTH, mainA, mainB, tailA, tailB, + pregLoop); + Add(mainA, mainA, tailA, pregLoop); + Add(mainB, mainB, tailB, pregLoop); + Add(mainA, mainA, mainB, pregLoop); + ReduceSum(vMean, mainA, pregLoop); + DataCopy(workAddr + i, vMean, pregMerge); + } + for (uint16_t i = 0; i < (uint16_t)masterRepeats; ++i) { + pregLoop = UpdateMask(masterSreg); + LoadForHandleMasterV2(masterAddr, (i * 2 + 0) * V_LENGTH, (i * 2 + 1) * V_LENGTH, mainA, mainB, pregLoop); + Add(mainA, mainA, mainB, pregLoop); + ReduceSum(vMean, mainA, pregLoop); + DataCopy(workAddr + remainRepeats + i, vMean, pregMerge); + } + LocalMemBar(); + for (uint16_t i = 0; i < (uint16_t)mergeRepeats; ++i) { + pregLoop = UpdateMask(mergeSreg); + DataCopy(mainA, workAddr + (i * 2 + 0) * V_LENGTH); + DataCopy(mainB, workAddr + (i * 2 + 1) * V_LENGTH); + Add(mainA, mainA, mainB, pregLoop); + ReduceSum(vMean, mainA, pregLoop); + DataCopy(workAddr + i, vMean, pregMerge); + } + LocalMemBar(); + { + pregLoop = UpdateMask(meanSreg); + DataCopy(mainA, workAddr + 0); + ReduceSum(vMean, mainA, pregLoop); + DataCopy(dstAddr + offset, vMean, pregMerge); + } + } +} + +} // namespace RmsNorm +#endif // OPS_BUILT_IN_TBE_IMPL_ASCENDC_RMS_NORM_REGBASE_COMMON_H