19
csrc/moe/hc_pre_sinkhorn/CMakeLists.txt
Normal file
19
csrc/moe/hc_pre_sinkhorn/CMakeLists.txt
Normal file
@@ -0,0 +1,19 @@
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
# CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
# Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
# See LICENSE in the root of the software repository for the full text of the License.
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
|
||||
file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*)
|
||||
if(NOT ENABLE_TEST AND NOT BENCHMARK)
|
||||
list(REMOVE_ITEM CURRENT_DIRS tests)
|
||||
endif()
|
||||
foreach(SUB_DIR ${CURRENT_DIRS})
|
||||
if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt")
|
||||
add_subdirectory(${SUB_DIR})
|
||||
endif()
|
||||
endforeach()
|
||||
62
csrc/moe/hc_pre_sinkhorn/op_host/CMakeLists.txt
Normal file
62
csrc/moe/hc_pre_sinkhorn/op_host/CMakeLists.txt
Normal file
@@ -0,0 +1,62 @@
|
||||
# This program is free software, you can redistribute it and/or modify it.
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
# 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.
|
||||
# ======================================================================================================================
|
||||
|
||||
# add_ops_compile_options(
|
||||
# OP_NAME HcPreSinkhorn
|
||||
# OPTIONS --cce-auto-sync=off
|
||||
# -Wno-deprecated-declarations
|
||||
# -Werror
|
||||
# -mllvm -cce-aicore-hoist-movemask=false
|
||||
# --op_relocatable_kernel_binary=true
|
||||
# )
|
||||
|
||||
# set(hc_pre_sinkhorn_depends transformer/attention/hc_pre_sinkhorn PARENT_SCOPE)
|
||||
|
||||
# target_sources(op_host_aclnn PRIVATE
|
||||
# op_host/hc_pre_sinkhorn_def.cpp
|
||||
# )
|
||||
|
||||
# target_sources(optiling PRIVATE
|
||||
# op_host/hc_pre_sinkhorn_tiling.cpp
|
||||
# )
|
||||
|
||||
# if (NOT BUILD_OPEN_PROJECT)
|
||||
# target_sources(opmaster_ct PRIVATE
|
||||
# op_host/hc_pre_sinkhorn_tiling.cpp
|
||||
# )
|
||||
# endif ()
|
||||
|
||||
# target_include_directories(optiling PRIVATE
|
||||
# ${CMAKE_CURRENT_SOURCE_DIR}/op_host
|
||||
# )
|
||||
|
||||
# target_sources(opsproto PRIVATE
|
||||
# op_host/hc_pre_sinkhorn_proto.cpp
|
||||
# )
|
||||
|
||||
add_op_to_compiled_list()
|
||||
|
||||
if (BUILD_OPEN_PROJECT)
|
||||
target_sources(op_host_aclnn PRIVATE
|
||||
hc_pre_sinkhorn_def.cpp
|
||||
)
|
||||
endif()
|
||||
|
||||
add_ops_compile_options(
|
||||
OP_NAME HcPreSinkhorn
|
||||
OPTIONS --cce-auto-sync=off
|
||||
-Wno-deprecated-declarations
|
||||
-mllvm -cce-aicore-hoist-movemask=false
|
||||
--op_relocatable_kernel_binary=true
|
||||
)
|
||||
|
||||
|
||||
if (NOT BUILD_OPS_RTY_KERNEL)
|
||||
add_modules_sources(OPTYPE hc_pre_sinkhorn ACLNNTYPE aclnn)
|
||||
endif()
|
||||
82
csrc/moe/hc_pre_sinkhorn/op_host/hc_pre_sinkhorn_def.cpp
Normal file
82
csrc/moe/hc_pre_sinkhorn/op_host/hc_pre_sinkhorn_def.cpp
Normal file
@@ -0,0 +1,82 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file hc_pre_sinkhorn_def.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include <cstdint>
|
||||
#include "register/op_def_registry.h"
|
||||
|
||||
namespace ops {
|
||||
class HcPreSinkhorn : public OpDef {
|
||||
public:
|
||||
explicit HcPreSinkhorn(const char* name) : OpDef(name)
|
||||
{
|
||||
this->Input("mixes")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND});
|
||||
this->Input("rsqrt")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND});
|
||||
this->Input("hc_scale")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND});
|
||||
this->Input("hc_base")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND});
|
||||
this->Input("x")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_BF16})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND});
|
||||
this->Output("y")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_BF16})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND});
|
||||
this->Output("post")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND});
|
||||
this->Output("comb_frag")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND});
|
||||
|
||||
this->Attr("hc_mult").AttrType(OPTIONAL).Int(4);
|
||||
this->Attr("hc_sinkhorn_iters").AttrType(OPTIONAL).Int(20);
|
||||
this->Attr("hc_eps").AttrType(OPTIONAL).Float(1e-6f); // default value
|
||||
|
||||
this->AICore().AddConfig("ascend910b");
|
||||
this->AICore().AddConfig("ascend910_93");
|
||||
|
||||
OpAICoreConfig regbaseCfg;
|
||||
regbaseCfg.DynamicCompileStaticFlag(true)
|
||||
.DynamicRankSupportFlag(true)
|
||||
.DynamicShapeSupportFlag(true)
|
||||
.ExtendCfgInfo("opFile.value", "hc_pre_sinkhorn");
|
||||
this->AICore().AddConfig("ascend950", regbaseCfg);
|
||||
}
|
||||
};
|
||||
|
||||
OP_ADD(HcPreSinkhorn);
|
||||
} // namespace ops
|
||||
91
csrc/moe/hc_pre_sinkhorn/op_host/hc_pre_sinkhorn_proto.cpp
Normal file
91
csrc/moe/hc_pre_sinkhorn/op_host/hc_pre_sinkhorn_proto.cpp
Normal file
@@ -0,0 +1,91 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file hc_pre_sinkhorn_proto.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include <graph/utils/type_utils.h>
|
||||
#include <register/op_impl_registry.h>
|
||||
|
||||
#include "error/ops_error.h"
|
||||
|
||||
using namespace ge;
|
||||
namespace ops {
|
||||
graphStatus InferShape4HcPreSinkhorn(gert::InferShapeContext* context)
|
||||
{
|
||||
OPS_LOG_D(context->GetNodeName(), "Begin to do InferShape4HcPreSinkhorn.");
|
||||
|
||||
const gert::Shape* mixShape = context->GetInputShape(0);
|
||||
OPS_LOG_E_IF_NULL(context, mixShape, ge::GRAPH_FAILED);
|
||||
auto mixDimNum = mixShape->GetDimNum();
|
||||
const gert::Shape* xShape = context->GetInputShape(4);
|
||||
OPS_LOG_E_IF_NULL(context, xShape, ge::GRAPH_FAILED);
|
||||
|
||||
gert::Shape* yShape = context->GetOutputShape(0);
|
||||
OPS_LOG_E_IF_NULL(context, yShape, ge::GRAPH_FAILED);
|
||||
gert::Shape* postShape = context->GetOutputShape(1);
|
||||
OPS_LOG_E_IF_NULL(context, postShape, ge::GRAPH_FAILED);
|
||||
gert::Shape* combFragShape = context->GetOutputShape(2);
|
||||
OPS_LOG_E_IF_NULL(context, combFragShape, ge::GRAPH_FAILED);
|
||||
|
||||
auto attrs = context->GetAttrs();
|
||||
auto *hcMult = attrs->GetAttrPointer<int>(0);
|
||||
|
||||
yShape->SetDimNum(mixDimNum);
|
||||
postShape->SetDimNum(mixDimNum);
|
||||
combFragShape->SetDimNum(mixDimNum + 1);
|
||||
if (mixDimNum == 2) {
|
||||
yShape->SetDim(0, mixShape->GetDim(0));
|
||||
yShape->SetDim(1, xShape->GetDim(2));
|
||||
|
||||
postShape->SetDim(0, mixShape->GetDim(0));
|
||||
postShape->SetDim(1, *hcMult);
|
||||
|
||||
combFragShape->SetDim(0, mixShape->GetDim(0));
|
||||
combFragShape->SetDim(1, *hcMult);
|
||||
combFragShape->SetDim(2, *hcMult);
|
||||
|
||||
} else {
|
||||
yShape->SetDim(0, mixShape->GetDim(0));
|
||||
yShape->SetDim(1, mixShape->GetDim(1));
|
||||
yShape->SetDim(2, xShape->GetDim(3));
|
||||
|
||||
postShape->SetDim(0, mixShape->GetDim(0));
|
||||
postShape->SetDim(1, mixShape->GetDim(1));
|
||||
postShape->SetDim(2, *hcMult);
|
||||
|
||||
combFragShape->SetDim(0, mixShape->GetDim(0));
|
||||
combFragShape->SetDim(1, mixShape->GetDim(1));
|
||||
combFragShape->SetDim(2, *hcMult);
|
||||
combFragShape->SetDim(3, *hcMult);
|
||||
}
|
||||
|
||||
OPS_LOG_D(context->GetNodeName(), "End to do InferShape4HcPreSinkhorn");
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
graphStatus InferDtype4HcPreSinkhorn(gert::InferDataTypeContext* context)
|
||||
{
|
||||
OPS_LOG_D(context->GetNodeName(), "InferDtype4HcPreSinkhorn enter");
|
||||
|
||||
const auto xDataType = context->GetInputDataType(4);
|
||||
context->SetOutputDataType(0, xDataType);
|
||||
context->SetOutputDataType(1, DT_FLOAT);
|
||||
context->SetOutputDataType(2, DT_FLOAT);
|
||||
OPS_LOG_D(context->GetNodeName(), "InferDtype4HcPreSinkhorn end");
|
||||
|
||||
return GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_INFERSHAPE(HcPreSinkhorn)
|
||||
.InferShape(InferShape4HcPreSinkhorn)
|
||||
.InferDataType(InferDtype4HcPreSinkhorn);
|
||||
} // namespace ops
|
||||
402
csrc/moe/hc_pre_sinkhorn/op_host/hc_pre_sinkhorn_tiling.cpp
Normal file
402
csrc/moe/hc_pre_sinkhorn/op_host/hc_pre_sinkhorn_tiling.cpp
Normal file
@@ -0,0 +1,402 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file hc_pre_sinkhorn_tiling.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include <sstream>
|
||||
#include "hc_pre_sinkhorn_tiling.h"
|
||||
|
||||
using namespace ge;
|
||||
namespace optiling {
|
||||
namespace {
|
||||
constexpr uint64_t WORKSPACE_SIZE = 32;
|
||||
int64_t CeilDiv(int64_t x, int64_t y)
|
||||
{
|
||||
if (y != 0) {
|
||||
return (x + y - 1) / y;
|
||||
}
|
||||
return x;
|
||||
}
|
||||
int64_t DownAlign(int64_t x, int64_t y) {
|
||||
if (y == 0) {
|
||||
return x;
|
||||
}
|
||||
return (x / y) * y;
|
||||
}
|
||||
int64_t RoundUp(int64_t x, int64_t y) {
|
||||
return CeilDiv(x, y) * y;
|
||||
}
|
||||
|
||||
constexpr int64_t BLOCK_SIZE = 32;
|
||||
constexpr int64_t REPEAT_SIZE = 256;
|
||||
constexpr int64_t DOUBLE_BUFFER = 2;
|
||||
}
|
||||
|
||||
ge::graphStatus HcPreSinkhornTiling::GetPlatformInfo()
|
||||
{
|
||||
auto platformInfo = context_->GetPlatformInfo();
|
||||
if (platformInfo == nullptr) {
|
||||
auto compileInfoPtr = context_->GetCompileInfo<HcPreSinkhornCompileInfo>();
|
||||
OPS_ERR_IF(compileInfoPtr == nullptr, OPS_LOG_E(context_, "compile info is null"),
|
||||
return ge::GRAPH_FAILED);
|
||||
coreNum_ = compileInfoPtr->coreNum;
|
||||
ubSize_ = compileInfoPtr->ubSize;
|
||||
} else {
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
|
||||
coreNum_ = ascendcPlatform.GetCoreNumAiv();
|
||||
uint64_t ubSizePlatForm;
|
||||
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSizePlatForm);
|
||||
ubSize_ = ubSizePlatForm;
|
||||
socVersion_ = ascendcPlatform.GetSocVersion();
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus HcPreSinkhornTiling::GetAttr()
|
||||
{
|
||||
auto* attrs = context_->GetAttrs();
|
||||
OPS_LOG_E_IF_NULL(context_, attrs, return ge::GRAPH_FAILED);
|
||||
|
||||
auto hcMultAttr = attrs->GetAttrPointer<int64_t>(0);
|
||||
hcMult_ = hcMultAttr == nullptr ? 4 : *hcMultAttr;
|
||||
|
||||
auto iterTimesAttr = attrs->GetAttrPointer<int64_t>(1);
|
||||
iterTimes_ = iterTimesAttr == nullptr ? 20 : *iterTimesAttr;
|
||||
|
||||
auto epsAttr = attrs->GetAttrPointer<float>(2);
|
||||
eps_ = epsAttr == nullptr ? 1e-5 : *epsAttr;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus HcPreSinkhornTiling::GetShapeAttrsInfoInner()
|
||||
{
|
||||
// (b, s, hc_mix) or (bs, hc_mix)
|
||||
auto shapeMixes = context_->GetInputShape(0);
|
||||
OPS_LOG_E_IF_NULL(context_, shapeMixes, return ge::GRAPH_FAILED);
|
||||
size_t mixerDimNum = shapeMixes->GetStorageShape().GetDimNum();
|
||||
if (mixerDimNum == 2) {
|
||||
bs_ = shapeMixes->GetStorageShape().GetDim(0);
|
||||
hcMix_ = shapeMixes->GetStorageShape().GetDim(1);
|
||||
} else if (mixerDimNum == 3) {
|
||||
int64_t b = shapeMixes->GetStorageShape().GetDim(0);
|
||||
int64_t s = shapeMixes->GetStorageShape().GetDim(1);
|
||||
bs_ = b * s;
|
||||
hcMix_ = shapeMixes->GetStorageShape().GetDim(2);
|
||||
}
|
||||
|
||||
auto shapeHcScale = context_->GetInputShape(2);
|
||||
int64_t scaleFirstDim = shapeHcScale->GetStorageShape().GetDim(0);
|
||||
OPS_ERR_IF(scaleFirstDim != 3,
|
||||
OPS_LOG_E(context_->GetNodeName(),
|
||||
"hc_scale size should be equal with 3, but is %ld", scaleFirstDim),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
auto shapeHcBase = context_->GetInputShape(3);
|
||||
int64_t baseFirstDim = shapeHcBase->GetStorageShape().GetDim(0);
|
||||
OPS_ERR_IF(baseFirstDim != hcMix_,
|
||||
OPS_LOG_E(context_->GetNodeName(),
|
||||
"hc_base size should be equal with mixhc, but is %ld", baseFirstDim),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
auto shapeX = context_->GetInputShape(4);
|
||||
d_ = (mixerDimNum == 2 ? shapeX->GetStorageShape().GetDim(2) : shapeX->GetStorageShape().GetDim(3));
|
||||
|
||||
OPS_ERR_IF(GetAttr() != ge::GRAPH_SUCCESS,
|
||||
OPS_LOG_E(context_->GetNodeName(), "get attr failed."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
|
||||
ge::graphStatus HcPreSinkhornTiling::CalcRegbaseOpTiling()
|
||||
{
|
||||
rowOfFormerBlock_ = CeilDiv(bs_, static_cast<int64_t>(coreNum_));
|
||||
usedCoreNums_ = std::min(CeilDiv(bs_, rowOfFormerBlock_), static_cast<int64_t>(coreNum_));
|
||||
rowOfTailBlock_ = bs_ - (usedCoreNums_ - 1) * rowOfFormerBlock_;
|
||||
|
||||
int64_t minRowPerCore = 1;
|
||||
int64_t rowOnceLoop = std::min(rowOfFormerBlock_, minRowPerCore);
|
||||
|
||||
hcMultAlign_ = RoundUp(hcMult_, BLOCK_SIZE / sizeof(float));
|
||||
int64_t mix0Size = rowOnceLoop * hcMultAlign_ * sizeof(float) * DOUBLE_BUFFER;
|
||||
int64_t mix1Size = rowOnceLoop * hcMultAlign_ * sizeof(float) * DOUBLE_BUFFER;
|
||||
int64_t mix2Size = rowOnceLoop * hcMult_ * hcMultAlign_ * sizeof(float) * DOUBLE_BUFFER;
|
||||
int64_t rsqrtSize = RoundUp(rowOnceLoop, BLOCK_SIZE / sizeof(float)) * sizeof(float) * DOUBLE_BUFFER;
|
||||
int64_t xSize = rowOnceLoop * hcMult_ * RoundUp(d_, 16) * 2 * DOUBLE_BUFFER; // x是bfloat16_t 类型
|
||||
int64_t ySize = rowOnceLoop * RoundUp(d_, 16) * 2 * DOUBLE_BUFFER;
|
||||
int64_t postSize = rowOnceLoop * hcMultAlign_ * sizeof(float) * DOUBLE_BUFFER;
|
||||
int64_t combFragSize = rowOnceLoop * hcMult_ * hcMultAlign_ * sizeof(float) * DOUBLE_BUFFER;
|
||||
int64_t base0Size = hcMultAlign_ * sizeof(float);
|
||||
int64_t base1Size = hcMultAlign_ * sizeof(float);
|
||||
int64_t base2Size = hcMult_ * hcMultAlign_ * sizeof(float);
|
||||
|
||||
int64_t totalSize = mix0Size + mix1Size + mix2Size + rsqrtSize + xSize + ySize + postSize + combFragSize +
|
||||
base0Size + base1Size + base2Size;
|
||||
rowFactor_ = rowOnceLoop;
|
||||
if (totalSize <= ubSize_) {
|
||||
// row和d均可以在ub内全载
|
||||
dLoop_ = 1;
|
||||
dFactor_ = d_;
|
||||
tailDFactor_ = dFactor_;
|
||||
} else {
|
||||
int64_t usedUbSize = mix0Size + mix1Size + mix2Size + rsqrtSize + postSize + combFragSize +
|
||||
base0Size + base1Size + base2Size;
|
||||
int64_t ubRemain = ubSize_ - usedUbSize;
|
||||
dFactor_ = d_;
|
||||
int64_t base = 2;
|
||||
while (1) {
|
||||
dFactor_ = CeilDiv(d_, base);
|
||||
xSize = rowOnceLoop * hcMult_ * RoundUp(dFactor_, 16) * 2 * DOUBLE_BUFFER; // x是bfloat16_t 类型
|
||||
ySize = rowOnceLoop * RoundUp(dFactor_, 16) * 2 * DOUBLE_BUFFER;
|
||||
int64_t targetSize = xSize + ySize;
|
||||
if (targetSize <= ubRemain) {
|
||||
break;
|
||||
}
|
||||
base++;
|
||||
}
|
||||
if (dFactor_ > 32) {
|
||||
dFactor_ = DownAlign(dFactor_, 32);
|
||||
}
|
||||
dLoop_ = CeilDiv(d_, dFactor_);
|
||||
tailDFactor_ = d_ % dFactor_ == 0 ? dFactor_ : d_ % dFactor_;
|
||||
}
|
||||
|
||||
// d全载,尝试搬入更多的bs
|
||||
if (dFactor_ == d_) {
|
||||
while (rowFactor_ <= rowOfFormerBlock_) {
|
||||
mix0Size = rowFactor_ * hcMultAlign_ * sizeof(float) * DOUBLE_BUFFER;
|
||||
mix1Size = rowFactor_ * hcMultAlign_ * sizeof(float) * DOUBLE_BUFFER;
|
||||
mix2Size = rowFactor_ * hcMult_ * hcMultAlign_ * sizeof(float) * DOUBLE_BUFFER;
|
||||
rsqrtSize = RoundUp(rowFactor_, BLOCK_SIZE / sizeof(float)) * sizeof(float) * DOUBLE_BUFFER;
|
||||
xSize = rowFactor_ * hcMult_ * RoundUp(d_, 16) * 2 * DOUBLE_BUFFER; // x是bfloat16_t 类型
|
||||
ySize = rowFactor_ * RoundUp(d_, 16) * 2 * DOUBLE_BUFFER;
|
||||
postSize = rowFactor_ * hcMultAlign_ * sizeof(float) * DOUBLE_BUFFER;
|
||||
combFragSize = rowFactor_ * hcMult_ * hcMultAlign_ * sizeof(float) * DOUBLE_BUFFER;
|
||||
totalSize = mix0Size + mix1Size + mix2Size + rsqrtSize + xSize + ySize + postSize + combFragSize +
|
||||
base0Size + base1Size + base2Size;
|
||||
if (totalSize > ubSize_) {
|
||||
rowFactor_ = rowFactor_ - 1;
|
||||
break;
|
||||
}
|
||||
rowFactor_ = rowFactor_ + 1;
|
||||
}
|
||||
rowFactor_ = rowFactor_ > rowOfFormerBlock_ ? rowFactor_ - 1 : rowFactor_;
|
||||
}
|
||||
rowLoopOfFormerBlock_ = CeilDiv(rowOfFormerBlock_, rowFactor_);
|
||||
rowLoopOfTailBlock_ = CeilDiv(rowOfTailBlock_, rowFactor_);
|
||||
tailRowFactorOfFormerBlock_ = rowOfFormerBlock_ % rowFactor_ == 0 ? rowFactor_ : rowOfFormerBlock_ % rowFactor_;
|
||||
tailRowFactorOfTailBlock_ = rowOfTailBlock_ % rowFactor_ == 0 ? rowFactor_ : rowOfTailBlock_ % rowFactor_;
|
||||
|
||||
tilingData_.set_bs(bs_);
|
||||
tilingData_.set_hcMix(hcMix_);
|
||||
tilingData_.set_hcMult(hcMult_);
|
||||
tilingData_.set_d(d_);
|
||||
tilingData_.set_hcMultAlign(hcMultAlign_);
|
||||
tilingData_.set_rowOfFormerBlock(rowOfFormerBlock_);
|
||||
tilingData_.set_rowOfTailBlock(rowOfTailBlock_);
|
||||
tilingData_.set_rowLoopOfFormerBlock(rowLoopOfFormerBlock_);
|
||||
tilingData_.set_rowLoopOfTailBlock(rowLoopOfTailBlock_);
|
||||
tilingData_.set_rowFactor(rowFactor_);
|
||||
tilingData_.set_tailRowFactorOfFormerBlock(tailRowFactorOfFormerBlock_);
|
||||
tilingData_.set_tailRowFactorOfTailBlock(tailRowFactorOfTailBlock_);
|
||||
tilingData_.set_dLoop(dLoop_);
|
||||
tilingData_.set_dFactor(dFactor_);
|
||||
tilingData_.set_tailDFactor(tailDFactor_);
|
||||
tilingData_.set_iterTimes(iterTimes_);
|
||||
tilingData_.set_eps(eps_);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus HcPreSinkhornTiling::CalcMembaseOpTiling()
|
||||
{
|
||||
rowOfFormerBlock_ = CeilDiv(bs_, static_cast<int64_t>(coreNum_));
|
||||
usedCoreNums_ = std::min(CeilDiv(bs_, rowOfFormerBlock_), static_cast<int64_t>(coreNum_));
|
||||
rowOfTailBlock_ = bs_ - (usedCoreNums_ - 1) * rowOfFormerBlock_;
|
||||
|
||||
int64_t minRowPerCore = 1;
|
||||
int64_t rowOnceLoop = std::min(rowOfFormerBlock_, minRowPerCore);
|
||||
|
||||
hcMultAlign_ = RoundUp(hcMult_, BLOCK_SIZE / sizeof(float));
|
||||
int64_t mix0Size = rowOnceLoop * hcMultAlign_ * sizeof(float) * DOUBLE_BUFFER;
|
||||
int64_t mix1Size = rowOnceLoop * hcMultAlign_ * sizeof(float) * DOUBLE_BUFFER;
|
||||
int64_t mix2Size = rowOnceLoop * hcMult_ * hcMultAlign_ * sizeof(float) * DOUBLE_BUFFER;
|
||||
int64_t rsqrtSize = RoundUp(rowOnceLoop, BLOCK_SIZE / sizeof(float)) * sizeof(float) * DOUBLE_BUFFER;
|
||||
int64_t xSize = rowOnceLoop * hcMult_ * RoundUp(d_, 16) * 2 * DOUBLE_BUFFER; // x是bfloat16_t 类型
|
||||
int64_t ySize = rowOnceLoop * RoundUp(d_, 16) * 2 * DOUBLE_BUFFER;
|
||||
int64_t postSize = rowOnceLoop * hcMultAlign_ * sizeof(float) * DOUBLE_BUFFER;
|
||||
int64_t combFragSize = rowOnceLoop * hcMult_ * hcMultAlign_ * sizeof(float) * DOUBLE_BUFFER;
|
||||
int64_t base0Size = hcMultAlign_ * sizeof(float);
|
||||
int64_t base1Size = hcMultAlign_ * sizeof(float);
|
||||
int64_t base2Size = hcMult_ * hcMultAlign_ * sizeof(float);
|
||||
int64_t xCastSize = rowOnceLoop * hcMult_ * RoundUp(d_, 8) * sizeof(float);
|
||||
int64_t yCastSize = rowOnceLoop * RoundUp(d_, 8) * sizeof(float);
|
||||
int64_t rowBrcb0Size = RoundUp(rowOnceLoop, 8) * BLOCK_SIZE;
|
||||
int64_t hcBrcb1Size = RoundUp(rowOnceLoop * hcMultAlign_, 8) * BLOCK_SIZE;
|
||||
int64_t reduceBufSize = rowOnceLoop * hcMultAlign_ * sizeof(float);
|
||||
|
||||
int64_t totalSize = mix0Size + mix1Size + mix2Size + rsqrtSize + xSize + ySize + postSize + combFragSize +
|
||||
base0Size + base1Size + base2Size + xCastSize + yCastSize + rowBrcb0Size + hcBrcb1Size + reduceBufSize;
|
||||
rowFactor_ = rowOnceLoop;
|
||||
if (totalSize <= ubSize_) {
|
||||
// row和d均可以在ub内全载
|
||||
dLoop_ = 1;
|
||||
dFactor_ = d_;
|
||||
tailDFactor_ = dFactor_;
|
||||
} else {
|
||||
int64_t usedUbSize = mix0Size + mix1Size + mix2Size + rsqrtSize + postSize + combFragSize +
|
||||
base0Size + base1Size + base2Size + rowBrcb0Size + hcBrcb1Size + reduceBufSize;
|
||||
int64_t ubRemain = ubSize_ - usedUbSize;
|
||||
dFactor_ = d_;
|
||||
int64_t base = 2;
|
||||
while (1) {
|
||||
dFactor_ = CeilDiv(d_, base);
|
||||
xSize = rowOnceLoop * hcMult_ * RoundUp(dFactor_, 16) * 2 * DOUBLE_BUFFER; // x是bfloat16_t 类型
|
||||
ySize = rowOnceLoop * RoundUp(dFactor_, 16) * 2 * DOUBLE_BUFFER;
|
||||
xCastSize = rowOnceLoop * hcMult_ * RoundUp(dFactor_, 8) * sizeof(float);
|
||||
yCastSize = rowOnceLoop * RoundUp(dFactor_, 8) * sizeof(float);
|
||||
int64_t targetSize = xSize + ySize + xCastSize + yCastSize;
|
||||
if (targetSize <= ubRemain) {
|
||||
break;
|
||||
}
|
||||
base++;
|
||||
}
|
||||
if (dFactor_ > 32) {
|
||||
dFactor_ = DownAlign(dFactor_, 32);
|
||||
}
|
||||
dLoop_ = CeilDiv(d_, dFactor_);
|
||||
tailDFactor_ = d_ % dFactor_ == 0 ? dFactor_ : d_ % dFactor_;
|
||||
}
|
||||
|
||||
// d全载,尝试搬入更多的bs
|
||||
if (dFactor_ == d_) {
|
||||
while (rowFactor_ <= rowOfFormerBlock_) {
|
||||
mix0Size = rowFactor_ * hcMultAlign_ * sizeof(float) * DOUBLE_BUFFER;
|
||||
mix1Size = rowFactor_ * hcMultAlign_ * sizeof(float) * DOUBLE_BUFFER;
|
||||
mix2Size = rowFactor_ * hcMult_ * hcMultAlign_ * sizeof(float) * DOUBLE_BUFFER;
|
||||
rsqrtSize = RoundUp(rowFactor_, BLOCK_SIZE / sizeof(float)) * sizeof(float) * DOUBLE_BUFFER;
|
||||
xSize = rowFactor_ * hcMult_ * RoundUp(d_, 16) * 2 * DOUBLE_BUFFER; // x是bfloat16_t 类型
|
||||
ySize = rowFactor_ * RoundUp(d_, 16) * 2 * DOUBLE_BUFFER;
|
||||
postSize = rowFactor_ * hcMultAlign_ * sizeof(float) * DOUBLE_BUFFER;
|
||||
combFragSize = rowFactor_ * hcMult_ * hcMultAlign_ * sizeof(float) * DOUBLE_BUFFER;
|
||||
xCastSize = rowFactor_ * hcMult_ * RoundUp(d_, 8) * sizeof(float);
|
||||
yCastSize = rowFactor_ * RoundUp(d_, 8) * sizeof(float);
|
||||
rowBrcb0Size = RoundUp(rowFactor_, 8) * BLOCK_SIZE;
|
||||
hcBrcb1Size = RoundUp(rowFactor_ * hcMultAlign_, 8) * BLOCK_SIZE;
|
||||
reduceBufSize = rowFactor_ * hcMultAlign_ * sizeof(float);
|
||||
totalSize = mix0Size + mix1Size + mix2Size + rsqrtSize + xSize + ySize + postSize + combFragSize +
|
||||
base0Size + base1Size + base2Size + xCastSize + yCastSize + rowBrcb0Size + hcBrcb1Size + reduceBufSize;;
|
||||
if (totalSize > ubSize_) {
|
||||
rowFactor_ = rowFactor_ - 1;
|
||||
break;
|
||||
}
|
||||
rowFactor_ = rowFactor_ + 1;
|
||||
}
|
||||
rowFactor_ = rowFactor_ > rowOfFormerBlock_ ? rowFactor_ - 1 : rowFactor_;
|
||||
}
|
||||
rowLoopOfFormerBlock_ = CeilDiv(rowOfFormerBlock_, rowFactor_);
|
||||
rowLoopOfTailBlock_ = CeilDiv(rowOfTailBlock_, rowFactor_);
|
||||
tailRowFactorOfFormerBlock_ = rowOfFormerBlock_ % rowFactor_ == 0 ? rowFactor_ : rowOfFormerBlock_ % rowFactor_;
|
||||
tailRowFactorOfTailBlock_ = rowOfTailBlock_ % rowFactor_ == 0 ? rowFactor_ : rowOfTailBlock_ % rowFactor_;
|
||||
|
||||
tilingData_.set_bs(bs_);
|
||||
tilingData_.set_hcMix(hcMix_);
|
||||
tilingData_.set_hcMult(hcMult_);
|
||||
tilingData_.set_d(d_);
|
||||
tilingData_.set_hcMultAlign(hcMultAlign_);
|
||||
tilingData_.set_rowOfFormerBlock(rowOfFormerBlock_);
|
||||
tilingData_.set_rowOfTailBlock(rowOfTailBlock_);
|
||||
tilingData_.set_rowLoopOfFormerBlock(rowLoopOfFormerBlock_);
|
||||
tilingData_.set_rowLoopOfTailBlock(rowLoopOfTailBlock_);
|
||||
tilingData_.set_rowFactor(rowFactor_);
|
||||
tilingData_.set_tailRowFactorOfFormerBlock(tailRowFactorOfFormerBlock_);
|
||||
tilingData_.set_tailRowFactorOfTailBlock(tailRowFactorOfTailBlock_);
|
||||
tilingData_.set_dLoop(dLoop_);
|
||||
tilingData_.set_dFactor(dFactor_);
|
||||
tilingData_.set_tailDFactor(tailDFactor_);
|
||||
tilingData_.set_iterTimes(iterTimes_);
|
||||
tilingData_.set_eps(eps_);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus HcPreSinkhornTiling::CalcOpTiling() {
|
||||
if (socVersion_ == platform_ascendc::SocVersion::ASCEND950) {
|
||||
return CalcRegbaseOpTiling();
|
||||
}
|
||||
return CalcMembaseOpTiling();
|
||||
}
|
||||
|
||||
|
||||
ge::graphStatus HcPreSinkhornTiling::DoOpTiling()
|
||||
{
|
||||
if (GetPlatformInfo() == ge::GRAPH_FAILED) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
if (GetShapeAttrsInfoInner() == ge::GRAPH_FAILED) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
if (CalcOpTiling() == ge::GRAPH_FAILED) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
if (GetWorkspaceSize() == ge::GRAPH_FAILED) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
if (PostTiling() == ge::GRAPH_FAILED) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
context_->SetTilingKey(0);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus HcPreSinkhornTiling::GetWorkspaceSize()
|
||||
{
|
||||
workspaceSize_ = WORKSPACE_SIZE;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus HcPreSinkhornTiling::PostTiling()
|
||||
{
|
||||
context_->SetTilingKey(0);
|
||||
context_->SetBlockDim(usedCoreNums_);
|
||||
size_t* workspaces = context_->GetWorkspaceSizes(1);
|
||||
workspaces[0] = workspaceSize_;
|
||||
tilingData_.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());
|
||||
context_->GetRawTilingData()->SetDataSize(tilingData_.GetDataSize());
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus TilingPrepareForHcPreSinkhorn(gert::TilingParseContext *context)
|
||||
{
|
||||
(void)context;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus TilingForHcPreSinkhorn(gert::TilingContext *context)
|
||||
{
|
||||
OPS_ERR_IF(context == nullptr, OPS_REPORT_VECTOR_INNER_ERR("HcPreSinkhorn", "Tiling context is null"),
|
||||
return ge::GRAPH_FAILED);
|
||||
HcPreSinkhornTiling hcPreTiling(context);
|
||||
return hcPreTiling.DoOpTiling();
|
||||
}
|
||||
|
||||
IMPL_OP_OPTILING(HcPreSinkhorn)
|
||||
.Tiling(TilingForHcPreSinkhorn)
|
||||
.TilingParse<HcPreSinkhornCompileInfo>(TilingPrepareForHcPreSinkhorn);
|
||||
|
||||
} // namespace optiling
|
||||
118
csrc/moe/hc_pre_sinkhorn/op_host/hc_pre_sinkhorn_tiling.h
Normal file
118
csrc/moe/hc_pre_sinkhorn/op_host/hc_pre_sinkhorn_tiling.h
Normal file
@@ -0,0 +1,118 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file hc_pre_sinkhorn_tiling.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef HC_PRE_SINKHORN_TILING_H
|
||||
#define HC_PRE_SINKHORN_TILING_H
|
||||
|
||||
|
||||
#include <vector>
|
||||
#include <iostream>
|
||||
#include "register/op_impl_registry.h"
|
||||
#include "platform/platform_infos_def.h"
|
||||
#include "exe_graph/runtime/tiling_context.h"
|
||||
#include "tiling/platform/platform_ascendc.h"
|
||||
#include "register/op_def_registry.h"
|
||||
#include "register/tilingdata_base.h"
|
||||
#include "tiling/tiling_api.h"
|
||||
#include "error/ops_error.h"
|
||||
#include "platform/platform_info.h"
|
||||
|
||||
namespace optiling {
|
||||
// ----------公共定义----------
|
||||
struct TilingRequiredParaInfo {
|
||||
const gert::CompileTimeTensorDesc *desc;
|
||||
const gert::StorageShape *shape;
|
||||
};
|
||||
|
||||
struct TilingOptionalParaInfo {
|
||||
const gert::CompileTimeTensorDesc *desc;
|
||||
const gert::Tensor *tensor;
|
||||
};
|
||||
// ----------算子TilingData定义----------
|
||||
BEGIN_TILING_DATA_DEF(HcPreSinkhornTilingData)
|
||||
TILING_DATA_FIELD_DEF(int64_t, bs);
|
||||
TILING_DATA_FIELD_DEF(int64_t, hcMix);
|
||||
TILING_DATA_FIELD_DEF(int64_t, hcMult);
|
||||
TILING_DATA_FIELD_DEF(int64_t, d);
|
||||
TILING_DATA_FIELD_DEF(int64_t, hcMultAlign);
|
||||
TILING_DATA_FIELD_DEF(int64_t, rowOfFormerBlock);
|
||||
TILING_DATA_FIELD_DEF(int64_t, rowOfTailBlock);
|
||||
TILING_DATA_FIELD_DEF(int64_t, rowLoopOfFormerBlock);
|
||||
TILING_DATA_FIELD_DEF(int64_t, rowLoopOfTailBlock);
|
||||
TILING_DATA_FIELD_DEF(int64_t, rowFactor);
|
||||
TILING_DATA_FIELD_DEF(int64_t, tailRowFactorOfFormerBlock);
|
||||
TILING_DATA_FIELD_DEF(int64_t, tailRowFactorOfTailBlock);
|
||||
TILING_DATA_FIELD_DEF(int64_t, dLoop);
|
||||
TILING_DATA_FIELD_DEF(int64_t, dFactor);
|
||||
TILING_DATA_FIELD_DEF(int64_t, tailDFactor);
|
||||
TILING_DATA_FIELD_DEF(int64_t, iterTimes);
|
||||
TILING_DATA_FIELD_DEF(float, eps);
|
||||
END_TILING_DATA_DEF;
|
||||
|
||||
REGISTER_TILING_DATA_CLASS(HcPreSinkhorn, HcPreSinkhornTilingData)
|
||||
|
||||
// ----------算子CompileInfo定义----------
|
||||
struct HcPreSinkhornCompileInfo {
|
||||
uint64_t coreNum = 0;
|
||||
uint64_t ubSize = 0;
|
||||
};
|
||||
|
||||
// ----------算子Tiling入参信息解析及check类----------
|
||||
class HcPreSinkhornTiling {
|
||||
public:
|
||||
explicit HcPreSinkhornTiling(gert::TilingContext* tilingContext) : context_(tilingContext)
|
||||
{
|
||||
}
|
||||
~HcPreSinkhornTiling() = default;
|
||||
|
||||
ge::graphStatus GetPlatformInfo();
|
||||
ge::graphStatus DoOpTiling();
|
||||
ge::graphStatus GetWorkspaceSize();
|
||||
ge::graphStatus PostTiling();
|
||||
ge::graphStatus GetAttr();
|
||||
ge::graphStatus GetShapeAttrsInfoInner();
|
||||
ge::graphStatus CalcOpTiling();
|
||||
ge::graphStatus CalcMembaseOpTiling();
|
||||
ge::graphStatus CalcRegbaseOpTiling();
|
||||
private:
|
||||
gert::TilingContext *context_ = nullptr;
|
||||
uint64_t tilingKey_ = 0;
|
||||
HcPreSinkhornTilingData tilingData_;
|
||||
uint64_t coreNum_ = 0;
|
||||
uint64_t workspaceSize_ = 0;
|
||||
uint64_t usedCoreNums_ = 0;
|
||||
uint64_t ubSize_ = 0;
|
||||
int64_t bs_ = 0;
|
||||
int64_t hcMix_ = 0;
|
||||
int64_t hcMult_ = 0;
|
||||
int64_t d_ = 0;
|
||||
int64_t hcMultAlign_ = 0;
|
||||
int64_t rowOfFormerBlock_ = 0;
|
||||
int64_t rowOfTailBlock_ = 0;
|
||||
int64_t rowLoopOfFormerBlock_ = 0;
|
||||
int64_t rowLoopOfTailBlock_ = 0;
|
||||
int64_t rowFactor_ = 0;
|
||||
int64_t tailRowFactorOfFormerBlock_ = 0;
|
||||
int64_t tailRowFactorOfTailBlock_= 0;
|
||||
int64_t dLoop_ = 0;
|
||||
int64_t dFactor_ = 0;
|
||||
int64_t tailDFactor_ = 0;
|
||||
int64_t iterTimes_ = 0;
|
||||
double eps_ = 0.0;
|
||||
platform_ascendc::SocVersion socVersion_ = platform_ascendc::SocVersion::ASCEND910B;
|
||||
};
|
||||
|
||||
} // namespace optiling
|
||||
#endif // HC_PRE_SINKHORN_TILING_H
|
||||
45
csrc/moe/hc_pre_sinkhorn/op_kernel/hc_pre_sinkhorn.cpp
Normal file
45
csrc/moe/hc_pre_sinkhorn/op_kernel/hc_pre_sinkhorn.cpp
Normal file
@@ -0,0 +1,45 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file hc_pre_sinkhorn.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#if defined(__DAV_C310__)
|
||||
#include "hc_pre_sinkhorn_regbase_perf.h"
|
||||
#include "hc_pre_sinkhorn_regbase_base.h"
|
||||
#else
|
||||
#include "hc_pre_sinkhorn_perf.h"
|
||||
#include "hc_pre_sinkhorn_base.h"
|
||||
#endif
|
||||
|
||||
using namespace HcPreSinkhorn;
|
||||
|
||||
extern "C" __global__ __aicore__ void hc_pre_sinkhorn(GM_ADDR mixes, GM_ADDR rsqrt, GM_ADDR hcScale, GM_ADDR hcBase,
|
||||
GM_ADDR x, GM_ADDR y, GM_ADDR post, GM_ADDR combFrag, GM_ADDR workspace,
|
||||
GM_ADDR tiling)
|
||||
{
|
||||
if (workspace == nullptr) {
|
||||
return;
|
||||
}
|
||||
|
||||
GM_ADDR userWs = GetUserWorkspace(workspace);
|
||||
if (userWs == nullptr) {
|
||||
return;
|
||||
}
|
||||
GET_TILING_DATA(tilingData, tiling);
|
||||
TPipe pipe;
|
||||
if (TILING_KEY_IS(0)) {
|
||||
HcPreSinkhorn::HcPreSinkhornPerf<DTYPE_X> op;
|
||||
op.Init(mixes, rsqrt, hcScale, hcBase, x, y, post, combFrag, userWs, &tilingData, &pipe);
|
||||
op.Process();
|
||||
}
|
||||
}
|
||||
623
csrc/moe/hc_pre_sinkhorn/op_kernel/hc_pre_sinkhorn_base.h
Normal file
623
csrc/moe/hc_pre_sinkhorn/op_kernel/hc_pre_sinkhorn_base.h
Normal file
@@ -0,0 +1,623 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file hc_pre_perf.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef HC_PRE_SINKHORN_BASE_H
|
||||
#define HC_PRE_SINKHORN_BASE_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
|
||||
namespace HcPreSinkhorn {
|
||||
using namespace AscendC;
|
||||
constexpr int32_t BLOCK_SIZE = 32;
|
||||
constexpr int32_t DEFAULT_BLOCK_STRIDE = 1;
|
||||
constexpr int32_t DEFAULT_REPEAT_STRIDE = 8;
|
||||
constexpr int32_t ONE_REPEAT_BLOCK_NUMS = 8;
|
||||
constexpr int32_t REPEAT_SIZE = 256;
|
||||
constexpr int32_t MAX_REPEAT_STRIDE = 255;
|
||||
|
||||
__aicore__ inline int32_t CeilDiv(int32_t a, int32_t b)
|
||||
{
|
||||
if (b == 0) {
|
||||
return a;
|
||||
}
|
||||
return (a + b - 1) / b;
|
||||
}
|
||||
|
||||
__aicore__ inline int32_t CeilAlign(int32_t a, int32_t b)
|
||||
{
|
||||
return CeilDiv(a, b) * b;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline int32_t RoundUp(int32_t num)
|
||||
{
|
||||
int32_t elemNum = BLOCK_SIZE / sizeof(T);
|
||||
return CeilAlign(num, elemNum);
|
||||
}
|
||||
|
||||
template <typename T, bool needBrc = true>
|
||||
__aicore__ inline void MulABLastDimBrcInline(const LocalTensor<T> &output, const LocalTensor<T> &input0,
|
||||
const LocalTensor<T> &input1, const LocalTensor<T> &tmpBuffer,
|
||||
const int32_t curRowNum, const int32_t curColNum)
|
||||
{
|
||||
if constexpr (needBrc) {
|
||||
uint32_t repeatTimes = CeilDiv(curRowNum, ONE_REPEAT_BLOCK_NUMS);
|
||||
Brcb(tmpBuffer, input1, repeatTimes, {DEFAULT_BLOCK_STRIDE, DEFAULT_REPEAT_STRIDE});
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
uint32_t elemInOneBlock = BLOCK_SIZE / sizeof(T);
|
||||
uint32_t elemInOneRepeat = REPEAT_SIZE / sizeof(T);
|
||||
uint32_t curColNumAlign = RoundUp<T>(curColNum);
|
||||
if (curColNum <= elemInOneBlock) {
|
||||
Mul(output, input0, tmpBuffer, curRowNum * curColNumAlign);
|
||||
} else {
|
||||
int32_t numRepeatPerLine = curColNum / elemInOneRepeat;
|
||||
int32_t numRemainPerLine = curColNum % elemInOneRepeat;
|
||||
int32_t dstRepStridePerLine = CeilDiv(curColNum, elemInOneBlock);
|
||||
BinaryRepeatParams instrParams;
|
||||
|
||||
if (numRepeatPerLine > 0) {
|
||||
if (dstRepStridePerLine > MAX_REPEAT_STRIDE || curRowNum < numRepeatPerLine) {
|
||||
// 在Col方向开Repeat, 并且Repeat小于255
|
||||
instrParams.dstBlkStride = 1;
|
||||
instrParams.src0BlkStride = 1;
|
||||
instrParams.src1BlkStride = 0;
|
||||
instrParams.dstRepStride = DEFAULT_REPEAT_STRIDE;
|
||||
instrParams.src0RepStride = DEFAULT_REPEAT_STRIDE;
|
||||
instrParams.src1RepStride = 0;
|
||||
for (uint32_t i = 0; i < curRowNum; i++) {
|
||||
Mul(output[i * curColNumAlign], input0[i * curColNumAlign], tmpBuffer[i * elemInOneBlock],
|
||||
elemInOneRepeat, numRepeatPerLine, instrParams);
|
||||
}
|
||||
} else {
|
||||
// 在Row方向开Repeat
|
||||
instrParams.dstBlkStride = 1;
|
||||
instrParams.src0BlkStride = 1;
|
||||
instrParams.src1BlkStride = 0;
|
||||
instrParams.dstRepStride = dstRepStridePerLine;
|
||||
instrParams.src0RepStride = dstRepStridePerLine;
|
||||
instrParams.src1RepStride = 1;
|
||||
for (uint32_t i = 0; i < numRepeatPerLine; i++) {
|
||||
Mul(output[i * elemInOneRepeat], input0[i * elemInOneRepeat], tmpBuffer, elemInOneRepeat, curRowNum,
|
||||
instrParams);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (numRemainPerLine > 0) {
|
||||
if (dstRepStridePerLine > MAX_REPEAT_STRIDE) {
|
||||
instrParams.dstBlkStride = 1;
|
||||
instrParams.src0BlkStride = 1;
|
||||
instrParams.src1BlkStride = 0;
|
||||
instrParams.dstRepStride = 0;
|
||||
instrParams.src0RepStride = 0;
|
||||
instrParams.src1RepStride = 0;
|
||||
for (uint32_t i = 0; i < curRowNum; i++) {
|
||||
Mul(output[numRepeatPerLine * elemInOneRepeat + i * curColNumAlign],
|
||||
input0[numRepeatPerLine * elemInOneRepeat + i * curColNumAlign], tmpBuffer[i * elemInOneBlock],
|
||||
numRemainPerLine, 1, instrParams);
|
||||
}
|
||||
} else {
|
||||
instrParams.dstBlkStride = 1;
|
||||
instrParams.src0BlkStride = 1;
|
||||
instrParams.src1BlkStride = 0;
|
||||
instrParams.dstRepStride = dstRepStridePerLine;
|
||||
instrParams.src0RepStride = dstRepStridePerLine;
|
||||
instrParams.src1RepStride = 0;
|
||||
Mul(output[numRepeatPerLine * elemInOneRepeat], input0[numRepeatPerLine * elemInOneRepeat], tmpBuffer,
|
||||
numRemainPerLine, curRowNum, instrParams);
|
||||
}
|
||||
}
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
template <typename T, bool needBrc = true>
|
||||
__aicore__ inline void SubABLastDimBrcInline(const LocalTensor<T> &output, const LocalTensor<T> &input0,
|
||||
const LocalTensor<T> &input1, const LocalTensor<T> &tmpBuffer,
|
||||
const int32_t curRowNum, const int32_t curColNum)
|
||||
{
|
||||
if constexpr (needBrc) {
|
||||
uint32_t repeatTimes = CeilDiv(curRowNum, ONE_REPEAT_BLOCK_NUMS);
|
||||
Brcb(tmpBuffer, input1, repeatTimes, {DEFAULT_BLOCK_STRIDE, DEFAULT_REPEAT_STRIDE});
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
uint32_t elemInOneBlock = BLOCK_SIZE / sizeof(T);
|
||||
uint32_t elemInOneRepeat = REPEAT_SIZE / sizeof(T);
|
||||
uint32_t curColNumAlign = RoundUp<T>(curColNum);
|
||||
if (curColNum <= elemInOneBlock) {
|
||||
Sub(output, input0, tmpBuffer, curRowNum * curColNumAlign);
|
||||
} else {
|
||||
int32_t numRepeatPerLine = curColNum / elemInOneRepeat;
|
||||
int32_t numRemainPerLine = curColNum % elemInOneRepeat;
|
||||
int32_t dstRepStridePerLine = CeilDiv(curColNum, elemInOneBlock);
|
||||
BinaryRepeatParams instrParams;
|
||||
|
||||
if (numRepeatPerLine > 0) {
|
||||
if (dstRepStridePerLine > MAX_REPEAT_STRIDE || curRowNum < numRepeatPerLine) {
|
||||
// 在Col方向开Repeat, 并且Repeat小于255
|
||||
instrParams.dstBlkStride = 1;
|
||||
instrParams.src0BlkStride = 1;
|
||||
instrParams.src1BlkStride = 0;
|
||||
instrParams.dstRepStride = DEFAULT_REPEAT_STRIDE;
|
||||
instrParams.src0RepStride = DEFAULT_REPEAT_STRIDE;
|
||||
instrParams.src1RepStride = 0;
|
||||
for (uint32_t i = 0; i < curRowNum; i++) {
|
||||
Sub(output[i * curColNumAlign], input0[i * curColNumAlign], tmpBuffer[i * elemInOneBlock],
|
||||
elemInOneRepeat, numRepeatPerLine, instrParams);
|
||||
}
|
||||
} else {
|
||||
// 在Row方向开Repeat
|
||||
instrParams.dstBlkStride = 1;
|
||||
instrParams.src0BlkStride = 1;
|
||||
instrParams.src1BlkStride = 0;
|
||||
instrParams.dstRepStride = dstRepStridePerLine;
|
||||
instrParams.src0RepStride = dstRepStridePerLine;
|
||||
instrParams.src1RepStride = 1;
|
||||
for (uint32_t i = 0; i < numRepeatPerLine; i++) {
|
||||
Sub(output[i * elemInOneRepeat], input0[i * elemInOneRepeat], tmpBuffer, elemInOneRepeat, curRowNum,
|
||||
instrParams);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (numRemainPerLine > 0) {
|
||||
if (dstRepStridePerLine > MAX_REPEAT_STRIDE) {
|
||||
instrParams.dstBlkStride = 1;
|
||||
instrParams.src0BlkStride = 1;
|
||||
instrParams.src1BlkStride = 0;
|
||||
instrParams.dstRepStride = 0;
|
||||
instrParams.src0RepStride = 0;
|
||||
instrParams.src1RepStride = 0;
|
||||
for (uint32_t i = 0; i < curRowNum; i++) {
|
||||
Sub(output[numRepeatPerLine * elemInOneRepeat + i * curColNumAlign],
|
||||
input0[numRepeatPerLine * elemInOneRepeat + i * curColNumAlign], tmpBuffer[i * elemInOneBlock],
|
||||
numRemainPerLine, 1, instrParams);
|
||||
}
|
||||
} else {
|
||||
instrParams.dstBlkStride = 1;
|
||||
instrParams.src0BlkStride = 1;
|
||||
instrParams.src1BlkStride = 0;
|
||||
instrParams.dstRepStride = dstRepStridePerLine;
|
||||
instrParams.src0RepStride = dstRepStridePerLine;
|
||||
instrParams.src1RepStride = 0;
|
||||
Sub(output[numRepeatPerLine * elemInOneRepeat], input0[numRepeatPerLine * elemInOneRepeat], tmpBuffer,
|
||||
numRemainPerLine, curRowNum, instrParams);
|
||||
}
|
||||
}
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
|
||||
template <typename T, bool needBrc = true>
|
||||
__aicore__ inline void DivABLastDimBrcInline(const LocalTensor<T> &output, const LocalTensor<T> &input0,
|
||||
const LocalTensor<T> &input1, const LocalTensor<T> &tmpBuffer,
|
||||
const int32_t curRowNum, const int32_t curColNum)
|
||||
{
|
||||
if constexpr (needBrc) {
|
||||
uint32_t repeatTimes = CeilDiv(curRowNum, ONE_REPEAT_BLOCK_NUMS);
|
||||
Brcb(tmpBuffer, input1, repeatTimes, {DEFAULT_BLOCK_STRIDE, DEFAULT_REPEAT_STRIDE});
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
uint32_t elemInOneBlock = BLOCK_SIZE / sizeof(T);
|
||||
uint32_t elemInOneRepeat = REPEAT_SIZE / sizeof(T);
|
||||
uint32_t curColNumAlign = RoundUp<T>(curColNum);
|
||||
if (curColNum <= elemInOneBlock) {
|
||||
Div(output, input0, tmpBuffer, curRowNum * curColNumAlign);
|
||||
} else {
|
||||
int32_t numRepeatPerLine = curColNum / elemInOneRepeat;
|
||||
int32_t numRemainPerLine = curColNum % elemInOneRepeat;
|
||||
int32_t dstRepStridePerLine = CeilDiv(curColNum, elemInOneBlock);
|
||||
BinaryRepeatParams instrParams;
|
||||
|
||||
if (numRepeatPerLine > 0) {
|
||||
if (dstRepStridePerLine > MAX_REPEAT_STRIDE || curRowNum < numRepeatPerLine) {
|
||||
// 在Col方向开Repeat, 并且Repeat小于255
|
||||
instrParams.dstBlkStride = 1;
|
||||
instrParams.src0BlkStride = 1;
|
||||
instrParams.src1BlkStride = 0;
|
||||
instrParams.dstRepStride = DEFAULT_REPEAT_STRIDE;
|
||||
instrParams.src0RepStride = DEFAULT_REPEAT_STRIDE;
|
||||
instrParams.src1RepStride = 0;
|
||||
for (uint32_t i = 0; i < curRowNum; i++) {
|
||||
Div(output[i * curColNumAlign], input0[i * curColNumAlign], tmpBuffer[i * elemInOneBlock],
|
||||
elemInOneRepeat, numRepeatPerLine, instrParams);
|
||||
}
|
||||
} else {
|
||||
// 在Row方向开Repeat
|
||||
instrParams.dstBlkStride = 1;
|
||||
instrParams.src0BlkStride = 1;
|
||||
instrParams.src1BlkStride = 0;
|
||||
instrParams.dstRepStride = dstRepStridePerLine;
|
||||
instrParams.src0RepStride = dstRepStridePerLine;
|
||||
instrParams.src1RepStride = 1;
|
||||
for (uint32_t i = 0; i < numRepeatPerLine; i++) {
|
||||
Div(output[i * elemInOneRepeat], input0[i * elemInOneRepeat], tmpBuffer, elemInOneRepeat, curRowNum,
|
||||
instrParams);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (numRemainPerLine > 0) {
|
||||
if (dstRepStridePerLine > MAX_REPEAT_STRIDE) {
|
||||
instrParams.dstBlkStride = 1;
|
||||
instrParams.src0BlkStride = 1;
|
||||
instrParams.src1BlkStride = 0;
|
||||
instrParams.dstRepStride = 0;
|
||||
instrParams.src0RepStride = 0;
|
||||
instrParams.src1RepStride = 0;
|
||||
for (uint32_t i = 0; i < curRowNum; i++) {
|
||||
Div(output[numRepeatPerLine * elemInOneRepeat + i * curColNumAlign],
|
||||
input0[numRepeatPerLine * elemInOneRepeat + i * curColNumAlign], tmpBuffer[i * elemInOneBlock],
|
||||
numRemainPerLine, 1, instrParams);
|
||||
}
|
||||
} else {
|
||||
instrParams.dstBlkStride = 1;
|
||||
instrParams.src0BlkStride = 1;
|
||||
instrParams.src1BlkStride = 0;
|
||||
instrParams.dstRepStride = dstRepStridePerLine;
|
||||
instrParams.src0RepStride = dstRepStridePerLine;
|
||||
instrParams.src1RepStride = 0;
|
||||
Div(output[numRepeatPerLine * elemInOneRepeat], input0[numRepeatPerLine * elemInOneRepeat], tmpBuffer,
|
||||
numRemainPerLine, curRowNum, instrParams);
|
||||
}
|
||||
}
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void AddBAFirstDimBrcInline(const LocalTensor<T> &output, const LocalTensor<T> &input0,
|
||||
const LocalTensor<T> &input1, const int32_t curRowNum,
|
||||
const int32_t curColNum)
|
||||
{
|
||||
uint32_t elemInOneBlock = BLOCK_SIZE / sizeof(T);
|
||||
uint32_t elemInOneRepeat = REPEAT_SIZE / sizeof(T);
|
||||
uint32_t curColNumAlign = RoundUp<T>(curColNum);
|
||||
int32_t numRepeatPerLine = curColNum / elemInOneRepeat;
|
||||
int32_t numRemainPerLine = curColNum % elemInOneRepeat;
|
||||
int32_t dstRepStridePerLine = CeilDiv(curColNum, elemInOneBlock);
|
||||
BinaryRepeatParams instrParams;
|
||||
if (numRepeatPerLine > 0) {
|
||||
if (dstRepStridePerLine > MAX_REPEAT_STRIDE || curRowNum < numRepeatPerLine) {
|
||||
// 在Col方向开Repeat, 并且Repeat小于255
|
||||
instrParams.dstBlkStride = 1;
|
||||
instrParams.src0BlkStride = 1;
|
||||
instrParams.src1BlkStride = 1;
|
||||
instrParams.dstRepStride = DEFAULT_REPEAT_STRIDE;
|
||||
instrParams.src0RepStride = DEFAULT_REPEAT_STRIDE;
|
||||
instrParams.src1RepStride = DEFAULT_REPEAT_STRIDE;
|
||||
for (uint32_t i = 0; i < curRowNum; i++) {
|
||||
Add(output[i * curColNumAlign], input0[i * curColNumAlign], input1, elemInOneRepeat, numRepeatPerLine,
|
||||
instrParams);
|
||||
}
|
||||
} else {
|
||||
// 在Row方向开Repeat
|
||||
instrParams.dstBlkStride = 1;
|
||||
instrParams.src0BlkStride = 1;
|
||||
instrParams.src1BlkStride = 1;
|
||||
instrParams.dstRepStride = dstRepStridePerLine;
|
||||
instrParams.src0RepStride = dstRepStridePerLine;
|
||||
instrParams.src1RepStride = 0;
|
||||
for (uint32_t i = 0; i < numRepeatPerLine; i++) {
|
||||
Add(output[i * elemInOneRepeat], input0[i * elemInOneRepeat], input1[i * elemInOneRepeat],
|
||||
elemInOneRepeat, curRowNum, instrParams);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (numRemainPerLine > 0) {
|
||||
if (dstRepStridePerLine > MAX_REPEAT_STRIDE) {
|
||||
instrParams.dstBlkStride = 1;
|
||||
instrParams.src0BlkStride = 1;
|
||||
instrParams.src1BlkStride = 1;
|
||||
instrParams.dstRepStride = 0;
|
||||
instrParams.src0RepStride = 0;
|
||||
instrParams.src1RepStride = 0;
|
||||
for (uint32_t i = 0; i < curRowNum; i++) {
|
||||
Add(output[numRepeatPerLine * elemInOneRepeat], input0[numRepeatPerLine * elemInOneRepeat], input1,
|
||||
numRemainPerLine, 1, instrParams);
|
||||
}
|
||||
} else {
|
||||
instrParams.dstBlkStride = 1;
|
||||
instrParams.src0BlkStride = 1;
|
||||
instrParams.src1BlkStride = 1;
|
||||
instrParams.dstRepStride = dstRepStridePerLine;
|
||||
instrParams.src0RepStride = dstRepStridePerLine;
|
||||
instrParams.src1RepStride = 0;
|
||||
Add(output[numRepeatPerLine * elemInOneRepeat], input0[numRepeatPerLine * elemInOneRepeat], input1,
|
||||
numRemainPerLine, curRowNum, instrParams);
|
||||
}
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void CalcDenominator(const LocalTensor<T> &output, const LocalTensor<T> &input,
|
||||
const uint32_t calCount)
|
||||
{
|
||||
Muls(output, input, static_cast<T>(-1.0), calCount);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Exp(output, output, calCount);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Adds(output, output, static_cast<T>(1.0), calCount);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
// 暂时不处理repeat超限场景
|
||||
template <typename T>
|
||||
__aicore__ inline void SigmoidPerf(const LocalTensor<T> &output, const LocalTensor<T> &input,
|
||||
const LocalTensor<T> &tmpBuffer, const int64_t calCount)
|
||||
{
|
||||
uint32_t elemInOneBlock = BLOCK_SIZE / sizeof(T);
|
||||
Duplicate(tmpBuffer, static_cast<T>(1.0), elemInOneBlock);
|
||||
CalcDenominator(output, input, calCount);
|
||||
uint32_t elemInOneRepeat = REPEAT_SIZE / sizeof(T);
|
||||
int32_t numRepeatPerLine = calCount / elemInOneRepeat;
|
||||
int32_t numRemainPerLine = calCount % elemInOneRepeat;
|
||||
BinaryRepeatParams instrParams;
|
||||
instrParams.dstBlkStride = 1;
|
||||
instrParams.src0BlkStride = 0;
|
||||
instrParams.src1BlkStride = 1;
|
||||
instrParams.dstRepStride = DEFAULT_REPEAT_STRIDE;
|
||||
instrParams.src0RepStride = 0;
|
||||
instrParams.src1RepStride = DEFAULT_REPEAT_STRIDE;
|
||||
Div(output, tmpBuffer, output, elemInOneRepeat, numRepeatPerLine, instrParams);
|
||||
if (numRemainPerLine != 0) {
|
||||
Div(output[numRepeatPerLine * elemInOneRepeat], tmpBuffer, output[numRepeatPerLine * elemInOneRepeat],
|
||||
numRemainPerLine, 1, instrParams);
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
__aicore__ inline void ProcessPre(const LocalTensor<float> &preLocal, const LocalTensor<float> &mixLocal,
|
||||
const LocalTensor<float> &hcBaseLocal, const LocalTensor<float> &rsqrtLocal,
|
||||
const LocalTensor<float> &tmpBuffer0, const LocalTensor<float> &tmpBuffer1,
|
||||
float scale, float eps, const int32_t curRowNum, const int32_t curColNum)
|
||||
{
|
||||
int32_t curColNumAlign = RoundUp<float>(curColNum);
|
||||
MulABLastDimBrcInline<float, true>(mixLocal, mixLocal, rsqrtLocal, tmpBuffer0, curRowNum, curColNum);
|
||||
Muls(mixLocal, mixLocal, scale, curRowNum * curColNumAlign);
|
||||
PipeBarrier<PIPE_V>();
|
||||
AddBAFirstDimBrcInline<float>(mixLocal, mixLocal, hcBaseLocal, curRowNum, curColNum);
|
||||
SigmoidPerf(preLocal, mixLocal, tmpBuffer1, curRowNum * curColNumAlign);
|
||||
Adds(preLocal, preLocal, eps, curRowNum * curColNumAlign);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
__aicore__ inline void ReduceSumARAPerf(const LocalTensor<float> &output, const LocalTensor<float> &input,
|
||||
const uint32_t dim0, const uint32_t dim1, const uint32_t dim2)
|
||||
{
|
||||
uint32_t elemInOneBlock = BLOCK_SIZE / sizeof(float);
|
||||
uint32_t elemInOneRepeat = REPEAT_SIZE / sizeof(float);
|
||||
uint32_t dim2Align = RoundUp<float>(dim2);
|
||||
|
||||
// 拷贝第一个R到output上
|
||||
DataCopyParams copyParams;
|
||||
copyParams.blockCount = dim0;
|
||||
copyParams.blockLen = dim2Align / elemInOneBlock;
|
||||
copyParams.srcStride = (dim1 - 1) * (dim2Align / elemInOneBlock);
|
||||
copyParams.dstStride = 0;
|
||||
DataCopy(output, input, copyParams);
|
||||
PipeBarrier<PIPE_V>();
|
||||
uint32_t dim2RepeatTimes = dim2 / elemInOneRepeat;
|
||||
uint32_t dim2Reminder = dim2 % elemInOneRepeat;
|
||||
// 沿着dim2方向开repeat
|
||||
BinaryRepeatParams instrParams;
|
||||
instrParams.dstBlkStride = 1;
|
||||
instrParams.src0BlkStride = 1;
|
||||
instrParams.src1BlkStride = 1;
|
||||
instrParams.dstRepStride = DEFAULT_REPEAT_STRIDE;
|
||||
instrParams.src0RepStride = DEFAULT_REPEAT_STRIDE;
|
||||
instrParams.src1RepStride = DEFAULT_REPEAT_STRIDE;
|
||||
for (uint32_t i = 0; i < dim0; i++) {
|
||||
for (uint32_t j = 1; j < dim1; j++) {
|
||||
Add(output[i * dim2Align], output[i * dim2Align], input[i * dim1 * dim2Align + j * dim2Align],
|
||||
elemInOneRepeat, dim2RepeatTimes, instrParams);
|
||||
if (dim2Reminder != 0) {
|
||||
Add(output[i * dim2Align + dim2RepeatTimes * elemInOneRepeat],
|
||||
output[i * dim2Align + dim2RepeatTimes * elemInOneRepeat],
|
||||
input[i * dim1 * dim2Align + j * dim2Align + +dim2RepeatTimes * elemInOneRepeat], dim2Reminder, 1,
|
||||
instrParams);
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
template <typename T0, typename T1>
|
||||
__aicore__ inline void CastTwoDim(const LocalTensor<T0> &output, const LocalTensor<T1> &input, const uint32_t dim0,
|
||||
const uint32_t dim1)
|
||||
{
|
||||
uint32_t dim1AlignT0 = RoundUp<T0>(dim1);
|
||||
uint32_t dim1AlignT1 = RoundUp<T1>(dim1);
|
||||
if constexpr (IsSameType<T1, bfloat16_t>::value && IsSameType<T0, float>::value) {
|
||||
for (uint32_t i = 0; i < dim0; i++) {
|
||||
Cast(output[i * dim1AlignT0], input[i * dim1AlignT1], AscendC::RoundMode::CAST_NONE, dim1);
|
||||
}
|
||||
} else {
|
||||
for (uint32_t i = 0; i < dim0; i++) {
|
||||
Cast(output[i * dim1AlignT0], input[i * dim1AlignT1], AscendC::RoundMode::CAST_RINT, dim1);
|
||||
}
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
|
||||
template <typename T>
|
||||
__aicore__ void inline ProcessY(const LocalTensor<T> &yLocal, const LocalTensor<T> &xLocal,
|
||||
const LocalTensor<float> &mix01Local, const LocalTensor<float> &hcBrcbLocal1,
|
||||
const LocalTensor<float> &xCastLocal, const LocalTensor<float> &yCastLocal,
|
||||
const uint32_t dim0, const uint32_t dim1, const uint32_t dim2)
|
||||
{
|
||||
CastTwoDim(xCastLocal, xLocal, dim0 * dim1, dim2);
|
||||
MulABLastDimBrcInline<float, true>(xCastLocal, xCastLocal, mix01Local, hcBrcbLocal1, dim0 * dim1, dim2);
|
||||
ReduceSumARAPerf(yCastLocal, xCastLocal, dim0, dim1, dim2);
|
||||
CastTwoDim(yLocal, yCastLocal, dim0, dim2);
|
||||
}
|
||||
|
||||
|
||||
__aicore__ inline void ProcessPost(const LocalTensor<float> &postLocal, const LocalTensor<float> &mixLocal,
|
||||
const LocalTensor<float> &hcBaseLocal, const LocalTensor<float> &rsqrtLocal,
|
||||
const LocalTensor<float> &tmpBuffer0, const LocalTensor<float> &tmpBuffer1,
|
||||
float scale, const int32_t curRowNum, const int32_t curColNum)
|
||||
{
|
||||
int32_t curColNumAlign = RoundUp<float>(curColNum);
|
||||
MulABLastDimBrcInline<float, false>(mixLocal, mixLocal, rsqrtLocal, tmpBuffer0, curRowNum, curColNum);
|
||||
Muls(mixLocal, mixLocal, scale, curRowNum * curColNumAlign);
|
||||
PipeBarrier<PIPE_V>();
|
||||
AddBAFirstDimBrcInline<float>(mixLocal, mixLocal, hcBaseLocal, curRowNum, curColNum);
|
||||
SigmoidPerf(postLocal, mixLocal, tmpBuffer1, curRowNum * curColNumAlign);
|
||||
Muls(postLocal, postLocal, static_cast<float>(2.0f), curRowNum * curColNumAlign);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
__aicore__ inline void LastDimReduceMaxPerf(const LocalTensor<float> &output, const LocalTensor<float> &input,
|
||||
const uint32_t curRowNum, const uint32_t curColNum)
|
||||
{
|
||||
uint32_t elemInOneBlock = BLOCK_SIZE / sizeof(float);
|
||||
WholeReduceMax(output, input, curColNum, curRowNum, 1, 1, CeilDiv(curColNum, elemInOneBlock),
|
||||
ReduceOrder::ORDER_ONLY_VALUE);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
__aicore__ inline void LastDimReduceSumPerf(const LocalTensor<float> &output, const LocalTensor<float> &input,
|
||||
const uint32_t curRowNum, const uint32_t curColNum)
|
||||
{
|
||||
uint32_t elemInOneBlock = BLOCK_SIZE / sizeof(float);
|
||||
WholeReduceSum(output, input, curColNum, curRowNum, 1, 1, CeilDiv(curColNum, elemInOneBlock));
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
|
||||
// 暂时只支持R轴小于64,既curColNum不能超过64
|
||||
__aicore__ inline void SoftmaxFP32Perf(const LocalTensor<float> &output, const LocalTensor<float> &input,
|
||||
const LocalTensor<float> &tmpReduceBuffer,
|
||||
const LocalTensor<float> tmpBrcbBuffer, const int32_t curRowNum,
|
||||
const int32_t curColNum, float eps)
|
||||
{
|
||||
LastDimReduceMaxPerf(tmpReduceBuffer, input, curRowNum, curColNum);
|
||||
SubABLastDimBrcInline<float, true>(output, input, tmpReduceBuffer, tmpBrcbBuffer, curRowNum, curColNum);
|
||||
uint32_t curColNumAlign = RoundUp<float>(curColNum);
|
||||
Exp(output, output, curRowNum * curColNumAlign);
|
||||
PipeBarrier<PIPE_V>();
|
||||
LastDimReduceSumPerf(tmpReduceBuffer, output, curRowNum, curColNum);
|
||||
DivABLastDimBrcInline<float, true>(output, output, tmpReduceBuffer, tmpBrcbBuffer, curRowNum, curColNum);
|
||||
Adds(output, output, eps, curRowNum * curColNumAlign);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
// (bs, hc_mult, hc_mult) = (bs, hc_mult, hc_mult) + (bs, 1, hc_mult)
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void DivABABrcInline(const LocalTensor<T> &output, const LocalTensor<T> &input0,
|
||||
const LocalTensor<T> &input1, const uint32_t dim0, const uint32_t dim1,
|
||||
const uint32_t dim2)
|
||||
{
|
||||
uint32_t elemInOneBlock = BLOCK_SIZE / sizeof(T);
|
||||
uint32_t elemInOneRepeat = REPEAT_SIZE / sizeof(T);
|
||||
uint32_t dim2Align = RoundUp<T>(dim2);
|
||||
uint32_t dim2RepeatTimes = dim2 / elemInOneRepeat;
|
||||
uint32_t dim2Reminder = dim2 % elemInOneRepeat;
|
||||
uint32_t dim2RepeatStride = CeilDiv(dim2, elemInOneBlock);
|
||||
// 在dim1方向开repeat
|
||||
BinaryRepeatParams instrParams;
|
||||
if (dim1 >= dim2RepeatTimes) {
|
||||
instrParams.dstBlkStride = 1;
|
||||
instrParams.src0BlkStride = 1;
|
||||
instrParams.src1BlkStride = 1;
|
||||
instrParams.dstRepStride = dim2RepeatStride;
|
||||
instrParams.src0RepStride = dim2RepeatStride;
|
||||
instrParams.src1RepStride = 0;
|
||||
for (uint32_t i = 0; i < dim0; i++) {
|
||||
for (uint32_t j = 0; j < dim2RepeatTimes; j++) {
|
||||
Div(output[i * dim1 * dim2Align + j * elemInOneRepeat],
|
||||
input0[i * dim1 * dim2Align + j * elemInOneRepeat], input1[i * dim2Align + j * elemInOneRepeat],
|
||||
elemInOneRepeat, dim1, instrParams);
|
||||
}
|
||||
if (dim2Reminder != 0) {
|
||||
Div(output[i * dim1 * dim2Align + dim2RepeatTimes * elemInOneRepeat],
|
||||
input0[i * dim1 * dim2Align + dim2RepeatTimes * elemInOneRepeat],
|
||||
input1[i * dim2Align + dim2RepeatTimes * elemInOneRepeat], dim2Reminder, dim1, instrParams);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// 在dim2方向开repeat
|
||||
instrParams.dstBlkStride = 1;
|
||||
instrParams.src0BlkStride = 1;
|
||||
instrParams.src1BlkStride = 1;
|
||||
instrParams.dstRepStride = DEFAULT_REPEAT_STRIDE;
|
||||
instrParams.src0RepStride = DEFAULT_REPEAT_STRIDE;
|
||||
instrParams.src1RepStride = DEFAULT_REPEAT_STRIDE;
|
||||
for (uint32_t i = 0; i < dim0; i++) {
|
||||
for (uint32_t j = 0; j < dim1; j++) {
|
||||
Div(output[i * dim1 * dim2Align + j * dim2Align], input0[i * dim1 * dim2Align + j * dim2Align],
|
||||
input1[i * dim2Align], dim2);
|
||||
}
|
||||
}
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void CopyIn(const GlobalTensor<T> &inputGm, const LocalTensor<T> &inputTensor, const uint16_t nBurst,
|
||||
const uint32_t copyLen, uint32_t srcStride = 0)
|
||||
{
|
||||
DataCopyPadExtParams<T> dataCopyPadExtParams;
|
||||
dataCopyPadExtParams.isPad = false;
|
||||
dataCopyPadExtParams.leftPadding = 0;
|
||||
dataCopyPadExtParams.rightPadding = 0;
|
||||
dataCopyPadExtParams.paddingValue = 0;
|
||||
|
||||
DataCopyExtParams dataCoptExtParams;
|
||||
dataCoptExtParams.blockCount = nBurst;
|
||||
dataCoptExtParams.blockLen = copyLen * sizeof(T);
|
||||
dataCoptExtParams.srcStride = srcStride * sizeof(T);
|
||||
dataCoptExtParams.dstStride = 0;
|
||||
DataCopyPad(inputTensor, inputGm, dataCoptExtParams, dataCopyPadExtParams);
|
||||
}
|
||||
|
||||
// (bs, hc_mult, hc_mult) --> (bs, hc_mult, hc_mult_align)
|
||||
template <typename T>
|
||||
__aicore__ inline void CopyInWithOuterFor(const GlobalTensor<T> &inputGm, const LocalTensor<T> &inputTensor,
|
||||
const uint16_t outerLoop, const uint16_t nBurst, const uint32_t copyLen,
|
||||
const uint32_t gmLastDim)
|
||||
{
|
||||
uint32_t elemInOneBlock = BLOCK_SIZE / sizeof(T);
|
||||
uint32_t ubLastDimAlign = RoundUp<T>(copyLen);
|
||||
|
||||
for (uint16_t i = 0; i < outerLoop; i++) {
|
||||
CopyIn(inputGm[i * nBurst * gmLastDim], inputTensor[i * nBurst * ubLastDimAlign], nBurst, copyLen);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void CopyOut(const LocalTensor<T> &outputTensor, const GlobalTensor<T> &outputGm,
|
||||
const uint16_t nBurst, const uint32_t copyLen, uint32_t dstStride = 0)
|
||||
{
|
||||
DataCopyExtParams dataCopyParams;
|
||||
dataCopyParams.blockCount = nBurst;
|
||||
dataCopyParams.blockLen = copyLen * sizeof(T);
|
||||
dataCopyParams.srcStride = 0;
|
||||
dataCopyParams.dstStride = dstStride * sizeof(T);
|
||||
DataCopyPad(outputGm, outputTensor, dataCopyParams);
|
||||
}
|
||||
|
||||
} // namespace HcPreSinkhorn
|
||||
|
||||
#endif
|
||||
256
csrc/moe/hc_pre_sinkhorn/op_kernel/hc_pre_sinkhorn_perf.h
Normal file
256
csrc/moe/hc_pre_sinkhorn/op_kernel/hc_pre_sinkhorn_perf.h
Normal file
@@ -0,0 +1,256 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file hc_pre_sinkhorn_perf.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef HC_PRE_SINKHORN_PERF_H
|
||||
#define HC_PRE_SINKHORN_PERF_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "hc_pre_sinkhorn_base.h"
|
||||
|
||||
namespace HcPreSinkhorn {
|
||||
using namespace AscendC;
|
||||
template <typename T>
|
||||
class HcPreSinkhornPerf {
|
||||
public:
|
||||
__aicore__ inline HcPreSinkhornPerf()
|
||||
{
|
||||
}
|
||||
|
||||
__aicore__ inline void Init(GM_ADDR mixes, GM_ADDR rsqrt, GM_ADDR hcScale, GM_ADDR hcBase, GM_ADDR x, GM_ADDR y,
|
||||
GM_ADDR post, GM_ADDR combFrag, GM_ADDR workspace,
|
||||
const HcPreSinkhornTilingData *tilingDataPtr, TPipe *pipePtr)
|
||||
{
|
||||
pipe = pipePtr;
|
||||
tilingData = tilingDataPtr;
|
||||
|
||||
mixesGm.SetGlobalBuffer((__gm__ float *)mixes);
|
||||
rsqrtGm.SetGlobalBuffer((__gm__ float *)rsqrt);
|
||||
hcScaleGm.SetGlobalBuffer((__gm__ float *)hcScale);
|
||||
hcBaseGm.SetGlobalBuffer((__gm__ float *)hcBase);
|
||||
xGm.SetGlobalBuffer((__gm__ T *)x);
|
||||
yGm.SetGlobalBuffer((__gm__ T *)y);
|
||||
postGm.SetGlobalBuffer((__gm__ float *)post);
|
||||
combFragGm.SetGlobalBuffer((__gm__ float *)combFrag);
|
||||
|
||||
// InQue
|
||||
int64_t mixesQue01Size = tilingData->rowFactor * tilingData->hcMultAlign * 2 * sizeof(float);
|
||||
pipe->InitBuffer(mixesQue01, 2, mixesQue01Size);
|
||||
pipe->InitBuffer(mixesQue2, 2,
|
||||
tilingData->rowFactor * tilingData->hcMult * tilingData->hcMultAlign * sizeof(float));
|
||||
pipe->InitBuffer(rsqrtQue, 2, RoundUp<float>(tilingData->rowFactor) * sizeof(float));
|
||||
pipe->InitBuffer(xQue, 2,
|
||||
tilingData->rowFactor * tilingData->hcMult * RoundUp<T>(tilingData->dFactor) * sizeof(T));
|
||||
|
||||
// OutQue
|
||||
pipe->InitBuffer(yQue, 2,
|
||||
tilingData->rowFactor * RoundUp<T>(tilingData->dFactor) * sizeof(T));
|
||||
pipe->InitBuffer(postQue, 2, tilingData->rowFactor * tilingData->hcMultAlign * sizeof(float));
|
||||
pipe->InitBuffer(combFragQue, 2,
|
||||
tilingData->rowFactor * tilingData->hcMult * tilingData->hcMultAlign * sizeof(float));
|
||||
|
||||
// TBuf
|
||||
pipe->InitBuffer(hcBaseBuf0, tilingData->hcMultAlign * sizeof(float));
|
||||
pipe->InitBuffer(hcBaseBuf1, tilingData->hcMultAlign * sizeof(float));
|
||||
pipe->InitBuffer(hcBaseBuf2, tilingData->hcMult * tilingData->hcMultAlign * sizeof(float));
|
||||
pipe->InitBuffer(rowBrcbBuf0, RoundUp<float>(tilingData->rowFactor) * BLOCK_SIZE);
|
||||
pipe->InitBuffer(hcBrcbBuf1, RoundUp<float>(tilingData->rowFactor * tilingData->hcMultAlign) * BLOCK_SIZE);
|
||||
pipe->InitBuffer(reduceBuf, tilingData->rowFactor * tilingData->hcMultAlign * sizeof(float));
|
||||
pipe->InitBuffer(xCastBuf, tilingData->rowFactor * tilingData->hcMult * RoundUp<float>(tilingData->dFactor) *
|
||||
sizeof(float));
|
||||
pipe->InitBuffer(yCastBuf, tilingData->rowFactor * RoundUp<float>(tilingData->dFactor) *
|
||||
sizeof(float));
|
||||
|
||||
hcBase0Local = hcBaseBuf0.Get<float>();
|
||||
hcBase1Local = hcBaseBuf1.Get<float>();
|
||||
hcBase2Local = hcBaseBuf2.Get<float>();
|
||||
rowBrcbLocal0 = rowBrcbBuf0.Get<float>();
|
||||
hcBrcbLocal1 = hcBrcbBuf1.Get<float>();
|
||||
reduceLocal = reduceBuf.Get<float>();
|
||||
xCastLocal = xCastBuf.Get<float>();
|
||||
yCastLocal = yCastBuf.Get<float>();
|
||||
}
|
||||
|
||||
__aicore__ inline void Process()
|
||||
{
|
||||
int64_t curBlockIdx = GetBlockIdx();
|
||||
int64_t totalBlockNum = GetBlockNum();
|
||||
|
||||
int64_t rowOuterLoop =
|
||||
(curBlockIdx == totalBlockNum - 1) ? tilingData->rowLoopOfTailBlock : tilingData->rowLoopOfFormerBlock;
|
||||
int64_t tailRowFactor = (curBlockIdx == totalBlockNum - 1) ? tilingData->tailRowFactorOfTailBlock :
|
||||
tilingData->tailRowFactorOfFormerBlock;
|
||||
|
||||
CopyIn(hcBaseGm, hcBase0Local, 1, tilingData->hcMult);
|
||||
CopyIn(hcBaseGm[tilingData->hcMult], hcBase1Local, 1, tilingData->hcMult);
|
||||
CopyIn(hcBaseGm[tilingData->hcMult * 2], hcBase2Local, tilingData->hcMult, tilingData->hcMult);
|
||||
event_t eventId = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE2_V));
|
||||
SetFlag<HardEvent::MTE2_V>(eventId);
|
||||
WaitFlag<HardEvent::MTE2_V>(eventId);
|
||||
int64_t mixGmBaseOffset = curBlockIdx * tilingData->rowOfFormerBlock * tilingData->hcMix;
|
||||
int64_t xGmBaseOffset = curBlockIdx * tilingData->rowOfFormerBlock * tilingData->hcMult * tilingData->d;
|
||||
for (int64_t rowOuterIdx = 0; rowOuterIdx < rowOuterLoop; rowOuterIdx++) {
|
||||
int64_t curRowFactor = (rowOuterIdx == rowOuterLoop - 1) ? tailRowFactor : tilingData->rowFactor;
|
||||
mixes01Local = mixesQue01.AllocTensor<float>();
|
||||
CopyIn(mixesGm[mixGmBaseOffset + rowOuterIdx * tilingData->rowFactor * tilingData->hcMix], mixes01Local,
|
||||
curRowFactor, tilingData->hcMult, tilingData->hcMix - tilingData->hcMult);
|
||||
CopyIn(
|
||||
mixesGm[mixGmBaseOffset + rowOuterIdx * tilingData->rowFactor * tilingData->hcMix + tilingData->hcMult],
|
||||
mixes01Local[tilingData->rowFactor * tilingData->hcMultAlign], curRowFactor, tilingData->hcMult,
|
||||
tilingData->hcMix - tilingData->hcMult);
|
||||
mixesQue01.EnQue(mixes01Local);
|
||||
rsqrtLocal = rsqrtQue.AllocTensor<float>();
|
||||
CopyIn(rsqrtGm[curBlockIdx * tilingData->rowOfFormerBlock + rowOuterIdx * tilingData->rowFactor],
|
||||
rsqrtLocal, 1, curRowFactor);
|
||||
rsqrtQue.EnQue(rsqrtLocal);
|
||||
|
||||
mixes01Local = mixesQue01.DeQue<float>();
|
||||
rsqrtLocal = rsqrtQue.DeQue<float>();
|
||||
ProcessPre(mixes01Local, mixes01Local, hcBase0Local, rsqrtLocal, rowBrcbLocal0, hcBrcbLocal1,
|
||||
hcScaleGm.GetValue(0), tilingData->eps, curRowFactor, tilingData->hcMult);
|
||||
for (int64_t dLoopIdx = 0; dLoopIdx < tilingData->dLoop; dLoopIdx++) {
|
||||
int64_t curDFactor =
|
||||
(dLoopIdx == tilingData->dLoop - 1) ? tilingData->tailDFactor : tilingData->dFactor;
|
||||
xLocal = xQue.template AllocTensor<T>();
|
||||
CopyIn(xGm[xGmBaseOffset + rowOuterIdx * tilingData->rowFactor * tilingData->hcMult * tilingData->d +
|
||||
dLoopIdx * tilingData->dFactor],
|
||||
xLocal, tilingData->rowFactor * tilingData->hcMult, curDFactor, tilingData->d - curDFactor);
|
||||
xQue.template EnQue(xLocal);
|
||||
xLocal = xQue.template DeQue<T>();
|
||||
yLocal = yQue.template AllocTensor<T>();
|
||||
ProcessY(yLocal, xLocal, mixes01Local, hcBrcbLocal1, xCastLocal, yCastLocal, curRowFactor,
|
||||
tilingData->hcMult, curDFactor);
|
||||
xQue.template FreeTensor(xLocal);
|
||||
yQue.template EnQue(yLocal);
|
||||
yLocal = yQue.template DeQue<T>();
|
||||
CopyOut(yLocal,
|
||||
yGm[curBlockIdx * tilingData->rowOfFormerBlock * tilingData->d +
|
||||
rowOuterIdx * tilingData->rowFactor * tilingData->d + dLoopIdx * tilingData->dFactor],
|
||||
curRowFactor, curDFactor, tilingData->d - curDFactor);
|
||||
yQue.template FreeTensor(yLocal);
|
||||
}
|
||||
// post
|
||||
postLocal = postQue.AllocTensor<float>();
|
||||
ProcessPost(postLocal, mixes01Local[tilingData->rowFactor * tilingData->hcMultAlign], hcBase1Local,
|
||||
rsqrtLocal, rowBrcbLocal0, hcBrcbLocal1, hcScaleGm.GetValue(1), curRowFactor,
|
||||
tilingData->hcMult);
|
||||
mixesQue01.template FreeTensor(mixes01Local);
|
||||
postQue.EnQue(postLocal);
|
||||
postLocal = postQue.DeQue<float>();
|
||||
CopyOut(postLocal,
|
||||
postGm[curBlockIdx * tilingData->rowOfFormerBlock * tilingData->hcMult +
|
||||
rowOuterIdx * tilingData->rowFactor * tilingData->hcMult],
|
||||
curRowFactor, tilingData->hcMult);
|
||||
postQue.FreeTensor(postLocal);
|
||||
|
||||
// combFrag
|
||||
mixes2Local = mixesQue2.AllocTensor<float>();
|
||||
CopyInWithOuterFor(mixesGm[mixGmBaseOffset + rowOuterIdx * tilingData->rowFactor * tilingData->hcMix +
|
||||
tilingData->hcMult * 2],
|
||||
mixes2Local, curRowFactor, tilingData->hcMult, tilingData->hcMult, tilingData->hcMix);
|
||||
mixesQue2.EnQue(mixes2Local);
|
||||
mixes2Local = mixesQue2.DeQue<float>();
|
||||
|
||||
combFragLocal = combFragQue.AllocTensor<float>();
|
||||
|
||||
MulABLastDimBrcInline<float, false>(mixes2Local, mixes2Local, rsqrtLocal, rowBrcbLocal0, curRowFactor,
|
||||
tilingData->hcMult * tilingData->hcMultAlign);
|
||||
Muls(mixes2Local, mixes2Local, hcScaleGm.GetValue(2),
|
||||
curRowFactor * tilingData->hcMult * tilingData->hcMultAlign);
|
||||
PipeBarrier<PIPE_V>();
|
||||
AddBAFirstDimBrcInline<float>(mixes2Local, mixes2Local, hcBase2Local, curRowFactor,
|
||||
tilingData->hcMult * tilingData->hcMultAlign);
|
||||
SoftmaxFP32Perf(mixes2Local, mixes2Local, reduceLocal, hcBrcbLocal1, curRowFactor * tilingData->hcMult,
|
||||
tilingData->hcMult, tilingData->eps);
|
||||
ReduceSumARAPerf(reduceLocal, mixes2Local, curRowFactor, tilingData->hcMult, tilingData->hcMult);
|
||||
Adds(reduceLocal, reduceLocal, tilingData->eps, curRowFactor * tilingData->hcMult);
|
||||
PipeBarrier<PIPE_V>();
|
||||
DivABABrcInline(combFragLocal, mixes2Local, reduceLocal, curRowFactor, tilingData->hcMult,
|
||||
tilingData->hcMult);
|
||||
for (int64_t iter = 0; iter < tilingData->iterTimes - 1; iter++) {
|
||||
LastDimReduceSumPerf(reduceLocal, combFragLocal, curRowFactor * tilingData->hcMult, tilingData->hcMult);
|
||||
Adds(reduceLocal, reduceLocal, tilingData->eps, curRowFactor * tilingData->hcMult);
|
||||
PipeBarrier<PIPE_V>();
|
||||
DivABLastDimBrcInline<float, true>(combFragLocal, combFragLocal, reduceLocal, hcBrcbLocal1,
|
||||
curRowFactor * tilingData->hcMult, tilingData->hcMult);
|
||||
ReduceSumARAPerf(reduceLocal, combFragLocal, curRowFactor, tilingData->hcMult, tilingData->hcMult);
|
||||
Adds(reduceLocal, reduceLocal, tilingData->eps, curRowFactor * tilingData->hcMult);
|
||||
PipeBarrier<PIPE_V>();
|
||||
DivABABrcInline(combFragLocal, combFragLocal, reduceLocal, curRowFactor, tilingData->hcMult,
|
||||
tilingData->hcMult);
|
||||
}
|
||||
mixesQue2.FreeTensor(mixes2Local);
|
||||
rsqrtQue.FreeTensor(rsqrtLocal);
|
||||
|
||||
combFragQue.EnQue(combFragLocal);
|
||||
combFragLocal = combFragQue.DeQue<float>();
|
||||
CopyOut(combFragLocal,
|
||||
combFragGm[curBlockIdx * tilingData->rowOfFormerBlock * tilingData->hcMult * tilingData->hcMult +
|
||||
rowOuterIdx * tilingData->rowFactor * tilingData->hcMult * tilingData->hcMult],
|
||||
curRowFactor * tilingData->hcMult, tilingData->hcMult);
|
||||
combFragQue.FreeTensor(combFragLocal);
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
TPipe *pipe;
|
||||
const HcPreSinkhornTilingData *tilingData;
|
||||
GlobalTensor<float> mixesGm;
|
||||
GlobalTensor<float> rsqrtGm;
|
||||
GlobalTensor<float> hcScaleGm;
|
||||
GlobalTensor<float> hcBaseGm;
|
||||
GlobalTensor<T> xGm;
|
||||
GlobalTensor<T> yGm;
|
||||
GlobalTensor<float> postGm;
|
||||
GlobalTensor<float> combFragGm;
|
||||
|
||||
TQue<QuePosition::VECIN, 1> mixesQue01;
|
||||
TQue<QuePosition::VECIN, 1> mixesQue2;
|
||||
TQue<QuePosition::VECIN, 1> rsqrtQue;
|
||||
TQue<QuePosition::VECIN, 1> xQue;
|
||||
TQue<QuePosition::VECOUT, 1> yQue;
|
||||
TQue<QuePosition::VECOUT, 1> postQue;
|
||||
TQue<QuePosition::VECOUT, 1> combFragQue;
|
||||
|
||||
TBuf<QuePosition::VECCALC> hcBaseBuf0;
|
||||
TBuf<QuePosition::VECCALC> hcBaseBuf1;
|
||||
TBuf<QuePosition::VECCALC> hcBaseBuf2;
|
||||
|
||||
TBuf<QuePosition::VECCALC> rowBrcbBuf0;
|
||||
TBuf<QuePosition::VECCALC> hcBrcbBuf1;
|
||||
TBuf<QuePosition::VECCALC> reduceBuf;
|
||||
|
||||
TBuf<QuePosition::VECCALC> xCastBuf;
|
||||
TBuf<QuePosition::VECCALC> yCastBuf;
|
||||
|
||||
LocalTensor<float> mixes01Local;
|
||||
LocalTensor<float> mixes2Local;
|
||||
LocalTensor<float> rsqrtLocal;
|
||||
LocalTensor<T> xLocal;
|
||||
LocalTensor<T> yLocal;
|
||||
LocalTensor<float> postLocal;
|
||||
LocalTensor<float> combFragLocal;
|
||||
LocalTensor<float> hcBase0Local;
|
||||
LocalTensor<float> hcBase1Local;
|
||||
LocalTensor<float> hcBase2Local;
|
||||
LocalTensor<float> rowBrcbLocal0;
|
||||
LocalTensor<float> hcBrcbLocal1;
|
||||
LocalTensor<float> reduceLocal;
|
||||
LocalTensor<float> xCastLocal;
|
||||
LocalTensor<float> yCastLocal;
|
||||
};
|
||||
|
||||
} // namespace HcPreSinkhorn
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,511 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file hc_pre_sinkhorn_regbase_base.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef HC_PRE_SINKHORN_RGEBASE_BASE_H
|
||||
#define HC_PRE_SINKHORN_RGEBASE_BASE_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
|
||||
namespace HcPreSinkhorn {
|
||||
using namespace AscendC;
|
||||
using namespace AscendC::MicroAPI;
|
||||
using AscendC::MicroAPI::MaskReg;
|
||||
using AscendC::MicroAPI::RegTensor;
|
||||
using AscendC::MicroAPI::UnalignReg;
|
||||
constexpr int32_t BLOCK_SIZE = 32;
|
||||
constexpr int32_t VL_FP32 = 64;
|
||||
|
||||
__aicore__ inline int32_t CeilDiv(int32_t a, int b)
|
||||
{
|
||||
if (b == 0) {
|
||||
return a;
|
||||
}
|
||||
return (a + b - 1) / b;
|
||||
}
|
||||
|
||||
__aicore__ inline int32_t CeilAlign(int32_t a, int b)
|
||||
{
|
||||
return CeilDiv(a, b) * b;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline int32_t RoundUp(int32_t num)
|
||||
{
|
||||
int32_t elemNum = BLOCK_SIZE / sizeof(T);
|
||||
return CeilAlign(num, elemNum);
|
||||
}
|
||||
|
||||
constexpr AscendC::MicroAPI::CastTrait castTraitB162B32Even = {
|
||||
AscendC::MicroAPI::RegLayout::ZERO,
|
||||
AscendC::MicroAPI::SatMode::UNKNOWN,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING,
|
||||
AscendC::RoundMode::UNKNOWN,
|
||||
};
|
||||
|
||||
constexpr AscendC::MicroAPI::CastTrait castTraitB322B16Even = {
|
||||
AscendC::MicroAPI::RegLayout::ZERO,
|
||||
AscendC::MicroAPI::SatMode::NO_SAT,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING,
|
||||
AscendC::RoundMode::CAST_RINT,
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void LoadInputData(RegTensor<float>& dst, __local_mem__ T* src, MaskReg pregLoop, uint32_t srcOffset)
|
||||
{
|
||||
if constexpr (IsSameType<T, float>::value) {
|
||||
DataCopy(dst, src + srcOffset);
|
||||
} else if constexpr (IsSameType<T, half>::value || IsSameType<T, bfloat16_t>::value) {
|
||||
RegTensor<T> tmp;
|
||||
DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_UNPACK_B16>(tmp, src + srcOffset);
|
||||
Cast<float, T, castTraitB162B32Even>(dst, tmp, pregLoop);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void StoreOutputData(
|
||||
__local_mem__ T* dst, RegTensor<float>& src, MaskReg pregLoop, uint32_t dstOffset)
|
||||
{
|
||||
if constexpr (IsSameType<T, float>::value) {
|
||||
DataCopy(dst + dstOffset, src, pregLoop);
|
||||
} else if constexpr (IsSameType<T, half>::value || IsSameType<T, bfloat16_t>::value) {
|
||||
RegTensor<T> tmp;
|
||||
Cast<T, float, castTraitB322B16Even>(tmp, src, pregLoop);
|
||||
DataCopy<T, AscendC::MicroAPI::StoreDist::DIST_PACK_B32>(dst + dstOffset, tmp, pregLoop);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void LoadInputDataWithBrc(
|
||||
RegTensor<float>& dst, __local_mem__ T* src, MaskReg pregLoop, uint32_t srcOffset)
|
||||
{
|
||||
if constexpr (IsSameType<T, float>::value) {
|
||||
DataCopy<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(dst, src + srcOffset);
|
||||
} else if constexpr (IsSameType<T, half>::value || IsSameType<T, bfloat16_t>::value) {
|
||||
RegTensor<T> tmp;
|
||||
DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_BRC_B16>(tmp, src + srcOffset);
|
||||
Cast<float, T, castTraitB162B32Even>(dst, tmp, pregLoop);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void VFSigmoid(
|
||||
RegTensor<float>& y, RegTensor<float>& x, RegTensor<float>& one, MaskReg pregLoop)
|
||||
{
|
||||
Muls(x, x, static_cast<float>(-1), pregLoop);
|
||||
Exp(x, x, pregLoop);
|
||||
Adds(x, x, static_cast<float>(1), pregLoop);
|
||||
Div(y, one, x, pregLoop);
|
||||
}
|
||||
|
||||
__aicore__ inline void VFProcessPre(
|
||||
const LocalTensor<float>& preLocal, const LocalTensor<float>& mixLocal, const LocalTensor<float>& hcBaseLocal,
|
||||
const LocalTensor<float>& rsqrtLocal, float scale, float eps, uint16_t curRowNum, uint16_t curColNum)
|
||||
{
|
||||
__local_mem__ float* preLocalAddr = (__local_mem__ float*)preLocal.GetPhyAddr();
|
||||
__local_mem__ float* mixLocalAddr = (__local_mem__ float*)mixLocal.GetPhyAddr();
|
||||
__local_mem__ float* hcBaseLocalAddr = (__local_mem__ float*)hcBaseLocal.GetPhyAddr();
|
||||
__local_mem__ float* rsqrtLocalAddr = (__local_mem__ float*)rsqrtLocal.GetPhyAddr();
|
||||
uint16_t loopCount = CeilDiv(curColNum, VL_FP32);
|
||||
uint32_t curColNumAlign = RoundUp<float>(curColNum);
|
||||
if (loopCount > 1) {
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
RegTensor<float> mix;
|
||||
RegTensor<float> base;
|
||||
RegTensor<float> rsqrt;
|
||||
RegTensor<float> one;
|
||||
MaskReg pregLoop = CreateMask<float>();
|
||||
uint32_t sreg = curColNum;
|
||||
Duplicate(one, static_cast<float>(1), pregLoop);
|
||||
for (uint16_t i = 0; i < loopCount; i++) {
|
||||
pregLoop = UpdateMask<float>(sreg);
|
||||
LoadInputData<float>(base, hcBaseLocalAddr, pregLoop, i * VL_FP32);
|
||||
for (uint16_t j = 0; j < curRowNum; j++) {
|
||||
LoadInputDataWithBrc<float>(rsqrt, rsqrtLocalAddr, pregLoop, j);
|
||||
LoadInputData<float>(mix, mixLocalAddr, pregLoop, i * VL_FP32 + j * curColNumAlign);
|
||||
Mul(mix, mix, rsqrt, pregLoop);
|
||||
Muls(mix, mix, scale, pregLoop);
|
||||
Add(mix, mix, base, pregLoop);
|
||||
VFSigmoid(mix, mix, one, pregLoop);
|
||||
Adds(mix, mix, eps, pregLoop);
|
||||
StoreOutputData(preLocalAddr, mix, pregLoop, i * VL_FP32 + j * curColNumAlign);
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
RegTensor<float> mix;
|
||||
RegTensor<float> base;
|
||||
RegTensor<float> rsqrt;
|
||||
RegTensor<float> one;
|
||||
uint32_t sreg = curColNum;
|
||||
MaskReg pregLoop = UpdateMask<float>(sreg);
|
||||
Duplicate(one, static_cast<float>(1), pregLoop);
|
||||
LoadInputData<float>(base, hcBaseLocalAddr, pregLoop, 0);
|
||||
for (uint16_t i = 0; i < curRowNum; i++) {
|
||||
LoadInputData<float>(mix, mixLocalAddr, pregLoop, i * curColNumAlign);
|
||||
LoadInputDataWithBrc<float>(rsqrt, rsqrtLocalAddr, pregLoop, i);
|
||||
Mul(mix, mix, rsqrt, pregLoop);
|
||||
Muls(mix, mix, scale, pregLoop);
|
||||
Add(mix, mix, base, pregLoop);
|
||||
VFSigmoid(mix, mix, one, pregLoop);
|
||||
Adds(mix, mix, eps, pregLoop);
|
||||
StoreOutputData(preLocalAddr, mix, pregLoop, i * curColNumAlign);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void VFProcessPost(
|
||||
const LocalTensor<float>& postLocal, const LocalTensor<float>& mixLocal, const LocalTensor<float>& hcBaseLocal,
|
||||
const LocalTensor<float>& rsqrtLocal, float scale, float eps, uint16_t curRowNum, uint16_t curColNum)
|
||||
{
|
||||
__local_mem__ float* postLocalAddr = (__local_mem__ float*)postLocal.GetPhyAddr();
|
||||
__local_mem__ float* mixLocalAddr = (__local_mem__ float*)mixLocal.GetPhyAddr();
|
||||
__local_mem__ float* hcBaseLocalAddr = (__local_mem__ float*)hcBaseLocal.GetPhyAddr();
|
||||
__local_mem__ float* rsqrtLocalAddr = (__local_mem__ float*)rsqrtLocal.GetPhyAddr();
|
||||
uint16_t loopCount = CeilDiv(curColNum, VL_FP32);
|
||||
uint32_t curColNumAlign = RoundUp<float>(curColNum);
|
||||
if (loopCount > 1) {
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
RegTensor<float> mix;
|
||||
RegTensor<float> base;
|
||||
RegTensor<float> rsqrt;
|
||||
RegTensor<float> one;
|
||||
MaskReg pregLoop = CreateMask<float>();
|
||||
uint32_t sreg = curColNum;
|
||||
Duplicate(one, static_cast<float>(1), pregLoop);
|
||||
for (uint16_t i = 0; i < loopCount; i++) {
|
||||
pregLoop = UpdateMask<float>(sreg);
|
||||
LoadInputData<float>(base, hcBaseLocalAddr, pregLoop, i * VL_FP32);
|
||||
for (uint16_t j = 0; j < curRowNum; j++) {
|
||||
LoadInputData<float>(mix, mixLocalAddr, pregLoop, i * VL_FP32 + j * curColNumAlign);
|
||||
LoadInputDataWithBrc<float>(rsqrt, rsqrtLocalAddr, pregLoop, i);
|
||||
Mul(mix, mix, rsqrt, pregLoop);
|
||||
Muls(mix, mix, scale, pregLoop);
|
||||
Add(mix, mix, base, pregLoop);
|
||||
VFSigmoid(mix, mix, one, pregLoop);
|
||||
Muls(mix, mix, static_cast<float>(2.0), pregLoop);
|
||||
StoreOutputData(postLocalAddr, mix, pregLoop, i * VL_FP32 + j * curColNumAlign);
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
RegTensor<float> mix;
|
||||
RegTensor<float> base;
|
||||
RegTensor<float> rsqrt;
|
||||
RegTensor<float> one;
|
||||
uint32_t sreg = curColNum;
|
||||
MaskReg pregLoop = UpdateMask<float>(sreg);
|
||||
Duplicate(one, static_cast<float>(1), pregLoop);
|
||||
LoadInputData<float>(base, hcBaseLocalAddr, pregLoop, 0);
|
||||
for (uint16_t i = 0; i < curRowNum; i++) {
|
||||
LoadInputData<float>(mix, mixLocalAddr, pregLoop, i * curColNumAlign);
|
||||
LoadInputDataWithBrc<float>(rsqrt, rsqrtLocalAddr, pregLoop, i);
|
||||
Mul(mix, mix, rsqrt, pregLoop);
|
||||
Muls(mix, mix, scale, pregLoop);
|
||||
Add(mix, mix, base, pregLoop);
|
||||
VFSigmoid(mix, mix, one, pregLoop);
|
||||
Muls(mix, mix, static_cast<float>(2.0), pregLoop);
|
||||
StoreOutputData(postLocalAddr, mix, pregLoop, i * curColNumAlign);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// dim2是R轴,R轴小于64, 不需要回写UB
|
||||
__aicore__ inline void VFProcessCombFragRLessVL(
|
||||
const LocalTensor<float>& combFragLocal, const LocalTensor<float>& mixLocal, const LocalTensor<float>& hcBaseLocal,
|
||||
const LocalTensor<float>& rsqrtLocal, float scale, float eps, uint16_t iters, uint16_t dim0, uint16_t dim1,
|
||||
uint16_t dim2)
|
||||
{
|
||||
__local_mem__ float* combFragLocalAddr = (__local_mem__ float*)combFragLocal.GetPhyAddr();
|
||||
__local_mem__ float* mixLocalAddr = (__local_mem__ float*)mixLocal.GetPhyAddr();
|
||||
__local_mem__ float* hcBaseLocalAddr = (__local_mem__ float*)hcBaseLocal.GetPhyAddr();
|
||||
__local_mem__ float* rsqrtLocalAddr = (__local_mem__ float*)rsqrtLocal.GetPhyAddr();
|
||||
uint32_t dim2Align = RoundUp<float>(dim2);
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
RegTensor<float> base;
|
||||
RegTensor<float> mix;
|
||||
RegTensor<float> rsqrt;
|
||||
RegTensor<float> max;
|
||||
RegTensor<float> sum;
|
||||
RegTensor<float> sum1;
|
||||
uint32_t sreg = dim2;
|
||||
MaskReg pregLoop = UpdateMask<float>(sreg);
|
||||
for (uint16_t i = 0; i < dim0; i++) {
|
||||
Duplicate(sum1, static_cast<float>(0), pregLoop);
|
||||
LoadInputDataWithBrc<float>(rsqrt, rsqrtLocalAddr, pregLoop, i);
|
||||
for (uint16_t j = 0; j < dim1; j++) {
|
||||
LoadInputData<float>(base, hcBaseLocalAddr, pregLoop, j * dim2Align);
|
||||
LoadInputData<float>(mix, mixLocalAddr, pregLoop, i * dim1 * dim2Align + j * dim2Align);
|
||||
Mul(mix, mix, rsqrt, pregLoop);
|
||||
Muls(mix, mix, scale, pregLoop);
|
||||
Add(mix, mix, base, pregLoop);
|
||||
ReduceMax(max, mix, pregLoop);
|
||||
Duplicate(max, max, pregLoop);
|
||||
Sub(mix, mix, max, pregLoop);
|
||||
Exp(mix, mix, pregLoop);
|
||||
ReduceSum(sum, mix, pregLoop);
|
||||
Duplicate(sum, sum, pregLoop);
|
||||
Div(mix, mix, sum, pregLoop);
|
||||
Adds(mix, mix, eps, pregLoop);
|
||||
Add(sum1, sum1, mix, pregLoop);
|
||||
StoreOutputData(combFragLocalAddr, mix, pregLoop, i * dim1 * dim2Align + j * dim2Align);
|
||||
}
|
||||
LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
Adds(sum1, sum1, eps, pregLoop);
|
||||
for (uint16_t j = 0; j < dim1; j++) {
|
||||
LoadInputData<float>(mix, combFragLocalAddr, pregLoop, i * dim1 * dim2Align + j * dim2Align);
|
||||
Div(mix, mix, sum1, pregLoop);
|
||||
StoreOutputData(combFragLocalAddr, mix, pregLoop, i * dim1 * dim2Align + j * dim2Align);
|
||||
}
|
||||
}
|
||||
for (uint16_t i = 0; i < iters; i++) {
|
||||
LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
for (uint16_t j = 0; j < dim0; j++) {
|
||||
Duplicate(sum1, static_cast<float>(0), pregLoop);
|
||||
for (uint16_t k = 0; k < dim1; k++) {
|
||||
LoadInputData<float>(mix, combFragLocalAddr, pregLoop, j * dim1 * dim2Align + k * dim2Align);
|
||||
ReduceSum(sum, mix, pregLoop);
|
||||
Duplicate(sum, sum, pregLoop);
|
||||
Adds(sum, sum, eps, pregLoop);
|
||||
Div(mix, mix, sum, pregLoop);
|
||||
Add(sum1, sum1, mix, pregLoop);
|
||||
StoreOutputData(combFragLocalAddr, mix, pregLoop, j * dim1 * dim2Align + k * dim2Align);
|
||||
}
|
||||
LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
Adds(sum1, sum1, eps, pregLoop);
|
||||
for (uint16_t k = 0; k < dim1; k++) {
|
||||
LoadInputData<float>(mix, combFragLocalAddr, pregLoop, j * dim1 * dim2Align + k * dim2Align);
|
||||
Div(mix, mix, sum1, pregLoop);
|
||||
StoreOutputData(combFragLocalAddr, mix, pregLoop, j * dim1 * dim2Align + k * dim2Align);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void VFProcessIteration(RegTensor<float>& sum0, RegTensor<float>& sum1, RegTensor<float>& mix, float eps, MaskReg pregLoop)
|
||||
{
|
||||
ReduceSum(sum1, mix, pregLoop);
|
||||
Duplicate(sum1, sum1, pregLoop);
|
||||
Adds(sum1, sum1, eps, pregLoop);
|
||||
Div(mix, mix, sum1, pregLoop);
|
||||
Add(sum0, sum0, mix, pregLoop);
|
||||
}
|
||||
|
||||
__aicore__ inline void VFProcessCombFragRLessVLUseFourUnfold(
|
||||
const LocalTensor<float>& combFragLocal, const LocalTensor<float>& mixLocal, const LocalTensor<float>& hcBaseLocal,
|
||||
const LocalTensor<float>& rsqrtLocal, float scale, float eps, uint16_t iters, uint16_t dim0, uint16_t dim1,
|
||||
uint16_t dim2)
|
||||
{
|
||||
__local_mem__ float* combFragLocalAddr = (__local_mem__ float*)combFragLocal.GetPhyAddr();
|
||||
__local_mem__ float* mixLocalAddr = (__local_mem__ float*)mixLocal.GetPhyAddr();
|
||||
__local_mem__ float* hcBaseLocalAddr = (__local_mem__ float*)hcBaseLocal.GetPhyAddr();
|
||||
__local_mem__ float* rsqrtLocalAddr = (__local_mem__ float*)rsqrtLocal.GetPhyAddr();
|
||||
uint32_t dim2Align = RoundUp<float>(dim2);
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
RegTensor<float> base;
|
||||
RegTensor<float> mix;
|
||||
RegTensor<float> mix1;
|
||||
RegTensor<float> mix2;
|
||||
RegTensor<float> mix3;
|
||||
RegTensor<float> mix4;
|
||||
RegTensor<float> rsqrt;
|
||||
RegTensor<float> max;
|
||||
RegTensor<float> sum;
|
||||
RegTensor<float> sum1;
|
||||
RegTensor<float> sum2;
|
||||
RegTensor<float> sum3;
|
||||
RegTensor<float> sum4;
|
||||
uint32_t sreg = dim2;
|
||||
MaskReg pregLoop = UpdateMask<float>(sreg);
|
||||
for (uint16_t i = 0; i < dim0; i++) {
|
||||
Duplicate(sum1, static_cast<float>(0), pregLoop);
|
||||
LoadInputDataWithBrc<float>(rsqrt, rsqrtLocalAddr, pregLoop, i);
|
||||
for (uint16_t j = 0; j < dim1; j++) {
|
||||
LoadInputData<float>(base, hcBaseLocalAddr, pregLoop, j * dim2Align);
|
||||
LoadInputData<float>(mix, mixLocalAddr, pregLoop, i * dim1 * dim2Align + j * dim2Align);
|
||||
Mul(mix, mix, rsqrt, pregLoop);
|
||||
Muls(mix, mix, scale, pregLoop);
|
||||
Add(mix, mix, base, pregLoop);
|
||||
ReduceMax(max, mix, pregLoop);
|
||||
Duplicate(max, max, pregLoop);
|
||||
Sub(mix, mix, max, pregLoop);
|
||||
Exp(mix, mix, pregLoop);
|
||||
ReduceSum(sum, mix, pregLoop);
|
||||
Duplicate(sum, sum, pregLoop);
|
||||
Div(mix, mix, sum, pregLoop);
|
||||
Adds(mix, mix, eps, pregLoop);
|
||||
Add(sum1, sum1, mix, pregLoop);
|
||||
StoreOutputData(combFragLocalAddr, mix, pregLoop, i * dim1 * dim2Align + j * dim2Align);
|
||||
}
|
||||
LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
Adds(sum1, sum1, eps, pregLoop);
|
||||
for (uint16_t j = 0; j < dim1; j++) {
|
||||
LoadInputData<float>(mix, combFragLocalAddr, pregLoop, i * dim1 * dim2Align + j * dim2Align);
|
||||
Div(mix, mix, sum1, pregLoop);
|
||||
StoreOutputData(combFragLocalAddr, mix, pregLoop, i * dim1 * dim2Align + j * dim2Align);
|
||||
}
|
||||
}
|
||||
LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
for (uint16_t i = 0; i < dim0; i++) {
|
||||
LoadInputData<float>(mix1, combFragLocalAddr, pregLoop, i * dim1 * dim2Align);
|
||||
LoadInputData<float>(mix2, combFragLocalAddr, pregLoop, i * dim1 * dim2Align + 1 * dim2Align);
|
||||
LoadInputData<float>(mix3, combFragLocalAddr, pregLoop, i * dim1 * dim2Align + 2 * dim2Align);
|
||||
LoadInputData<float>(mix4, combFragLocalAddr, pregLoop, i * dim1 * dim2Align + 3 * dim2Align);
|
||||
for (uint16_t j = 0; j < iters; j++) {
|
||||
Duplicate(sum, static_cast<float>(0), pregLoop);
|
||||
VFProcessIteration(sum, sum1, mix1, eps, pregLoop);
|
||||
VFProcessIteration(sum, sum2, mix2, eps, pregLoop);
|
||||
VFProcessIteration(sum, sum3, mix3, eps, pregLoop);
|
||||
VFProcessIteration(sum, sum4, mix4, eps, pregLoop);
|
||||
Adds(sum, sum, eps, pregLoop);
|
||||
Div(mix1, mix1, sum, pregLoop);
|
||||
Div(mix2, mix2, sum, pregLoop);
|
||||
Div(mix3, mix3, sum, pregLoop);
|
||||
Div(mix4, mix4, sum, pregLoop);
|
||||
}
|
||||
StoreOutputData(combFragLocalAddr, mix1, pregLoop, i * dim1 * dim2Align);
|
||||
StoreOutputData(combFragLocalAddr, mix2, pregLoop, i * dim1 * dim2Align + 1 * dim2Align);
|
||||
StoreOutputData(combFragLocalAddr, mix3, pregLoop, i * dim1 * dim2Align + 2 * dim2Align);
|
||||
StoreOutputData(combFragLocalAddr, mix4, pregLoop, i * dim1 * dim2Align + 3 * dim2Align);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void VFProcessY(
|
||||
const LocalTensor<T>& yLocal, const LocalTensor<float>& mixLocal, const LocalTensor<T>& xLocal, uint16_t bs,
|
||||
uint16_t hcMult, uint16_t d)
|
||||
{
|
||||
__local_mem__ T* yLocalAddr = (__local_mem__ T*)yLocal.GetPhyAddr();
|
||||
__local_mem__ float* mixLocalAddr = (__local_mem__ float*)mixLocal.GetPhyAddr();
|
||||
__local_mem__ T* xLocalAddr = (__local_mem__ T*)xLocal.GetPhyAddr();
|
||||
uint32_t dAlign = RoundUp<T>(d);
|
||||
uint16_t loopCount = CeilDiv(d, VL_FP32);
|
||||
uint32_t hcMultAlign = RoundUp<float>(hcMult);
|
||||
if (loopCount > 1) {
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
RegTensor<float> x;
|
||||
RegTensor<float> mix;
|
||||
RegTensor<float> sum;
|
||||
MaskReg pregLoop;
|
||||
for (uint16_t i = 0; i < bs; i++) {
|
||||
uint32_t sreg = d;
|
||||
for (uint16_t j = 0; j < loopCount; j++) {
|
||||
pregLoop = UpdateMask<float>(sreg);
|
||||
Duplicate(sum, static_cast<float>(0), pregLoop);
|
||||
for (uint16_t k = 0; k < hcMult; k++) {
|
||||
LoadInputDataWithBrc<float>(mix, mixLocalAddr, pregLoop, i * hcMultAlign + k);
|
||||
LoadInputData<T>(x, xLocalAddr, pregLoop, i * hcMult * dAlign + j * VL_FP32 + k * dAlign);
|
||||
Mul(x, mix, x, pregLoop);
|
||||
Add(sum, sum, x, pregLoop);
|
||||
}
|
||||
StoreOutputData(yLocalAddr, sum, pregLoop, i * dAlign + j * VL_FP32);
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
RegTensor<float> x;
|
||||
RegTensor<float> mix;
|
||||
RegTensor<float> sum;
|
||||
uint32_t sreg = d;
|
||||
MaskReg pregLoop = UpdateMask<float>(sreg);
|
||||
for (uint16_t i = 0; i < bs; i++) {
|
||||
Duplicate(sum, static_cast<float>(0), pregLoop);
|
||||
for (uint16_t j = 0; j < hcMult; j++) {
|
||||
LoadInputDataWithBrc<float>(mix, mixLocalAddr, pregLoop, i * hcMultAlign + j);
|
||||
LoadInputData<T>(x, xLocalAddr, pregLoop, i * hcMult * dAlign + j * dAlign);
|
||||
Mul(x, mix, x, pregLoop);
|
||||
Add(sum, sum, x, pregLoop);
|
||||
}
|
||||
StoreOutputData(yLocalAddr, sum, pregLoop, i * dAlign);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void CopyIn(
|
||||
const GlobalTensor<T>& inputGm, const LocalTensor<T>& inputTensor, const uint16_t nBurst, const uint32_t copyLen, uint32_t srcStride = 0)
|
||||
{
|
||||
DataCopyPadExtParams<T> dataCopyPadExtParams;
|
||||
dataCopyPadExtParams.isPad = false;
|
||||
dataCopyPadExtParams.leftPadding = 0;
|
||||
dataCopyPadExtParams.rightPadding = 0;
|
||||
dataCopyPadExtParams.paddingValue = 0;
|
||||
|
||||
DataCopyExtParams dataCoptExtParams;
|
||||
dataCoptExtParams.blockCount = nBurst;
|
||||
dataCoptExtParams.blockLen = copyLen * sizeof(T);
|
||||
dataCoptExtParams.srcStride = srcStride * sizeof(T);
|
||||
dataCoptExtParams.dstStride = 0;
|
||||
DataCopyPad(inputTensor, inputGm, dataCoptExtParams, dataCopyPadExtParams);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void CopyInWithLoopMode(
|
||||
const GlobalTensor<T>& inputGm, const LocalTensor<T>& inputTensor, const uint16_t outerLoop, const uint16_t nBurst, const uint32_t copyLen, const uint32_t gmLastDim, uint32_t srcStride = 0)
|
||||
{
|
||||
uint16_t copyLenAlign = RoundUp<T>(copyLen);
|
||||
LoopModeParams loopParams;
|
||||
loopParams.loop2Size = 1;
|
||||
loopParams.loop1Size = outerLoop;
|
||||
loopParams.loop2SrcStride = 0;
|
||||
loopParams.loop1SrcStride = gmLastDim * sizeof(T);
|
||||
loopParams.loop2DstStride = 0;
|
||||
loopParams.loop1DstStride = nBurst * copyLenAlign * sizeof(T);
|
||||
|
||||
DataCopyPadExtParams<T> dataCopyPadExtParams;
|
||||
dataCopyPadExtParams.isPad = false;
|
||||
dataCopyPadExtParams.leftPadding = 0;
|
||||
dataCopyPadExtParams.rightPadding = 0;
|
||||
dataCopyPadExtParams.paddingValue = 0;
|
||||
|
||||
DataCopyExtParams dataCoptExtParams;
|
||||
dataCoptExtParams.blockCount = nBurst;
|
||||
dataCoptExtParams.blockLen = copyLen * sizeof(T);
|
||||
dataCoptExtParams.srcStride = srcStride * sizeof(T);
|
||||
dataCoptExtParams.dstStride = 0;
|
||||
SetLoopModePara(loopParams, DataCopyMVType::OUT_TO_UB);
|
||||
DataCopyPad(inputTensor, inputGm, dataCoptExtParams, dataCopyPadExtParams);
|
||||
ResetLoopModePara(DataCopyMVType::OUT_TO_UB);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void CopyOut(
|
||||
const LocalTensor<T>& outputTensor, const GlobalTensor<T>& outputGm, const uint16_t nBurst, const uint32_t copyLen, uint32_t dstStride = 0)
|
||||
{
|
||||
DataCopyExtParams dataCopyParams;
|
||||
dataCopyParams.blockCount = nBurst;
|
||||
dataCopyParams.blockLen = copyLen * sizeof(T);
|
||||
dataCopyParams.srcStride = 0;
|
||||
dataCopyParams.dstStride = dstStride * sizeof(T);
|
||||
DataCopyPad(outputGm, outputTensor, dataCopyParams);
|
||||
}
|
||||
|
||||
} // namespace HCPreSinkhorn
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,204 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file hc_pre_sinkhorn_regbase_perf.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef HC_PRE_SINKHORN_REGBASE_PERF_H
|
||||
#define HC_PRE_SINKHORN_REGBASE_PERF_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "hc_pre_sinkhorn_regbase_base.h"
|
||||
|
||||
namespace HcPreSinkhorn {
|
||||
using namespace AscendC;
|
||||
template <typename T>
|
||||
class HcPreSinkhornPerf {
|
||||
public:
|
||||
__aicore__ inline HcPreSinkhornPerf()
|
||||
{}
|
||||
|
||||
__aicore__ inline void Init(
|
||||
GM_ADDR mixes, GM_ADDR rsqrt, GM_ADDR hcScale, GM_ADDR hcBase, GM_ADDR x, GM_ADDR y, GM_ADDR post,
|
||||
GM_ADDR combFrag, GM_ADDR workspace, const HcPreSinkhornTilingData* tilingDataPtr, TPipe* pipePtr)
|
||||
{
|
||||
pipe = pipePtr;
|
||||
tilingData = tilingDataPtr;
|
||||
|
||||
mixesGm.SetGlobalBuffer((__gm__ float*)mixes);
|
||||
rsqrtGm.SetGlobalBuffer((__gm__ float*)rsqrt);
|
||||
hcScaleGm.SetGlobalBuffer((__gm__ float*)hcScale);
|
||||
hcBaseGm.SetGlobalBuffer((__gm__ float*)hcBase);
|
||||
xGm.SetGlobalBuffer((__gm__ T*)x);
|
||||
yGm.SetGlobalBuffer((__gm__ T*)y);
|
||||
postGm.SetGlobalBuffer((__gm__ float*)post);
|
||||
combFragGm.SetGlobalBuffer((__gm__ float*)combFrag);
|
||||
|
||||
// InQue
|
||||
int64_t mixesQue01Size = tilingData->rowFactor * tilingData->hcMultAlign * 2 * sizeof(float);
|
||||
pipe->InitBuffer(mixesQue01, 2, mixesQue01Size);
|
||||
pipe->InitBuffer(
|
||||
mixesQue2, 2, tilingData->rowFactor * tilingData->hcMult * tilingData->hcMultAlign * sizeof(float));
|
||||
pipe->InitBuffer(rsqrtQue, 2, RoundUp<float>(tilingData->rowFactor) * sizeof(float));
|
||||
pipe->InitBuffer(
|
||||
xQue, 2, tilingData->rowFactor * tilingData->hcMult * RoundUp<T>(tilingData->dFactor) * sizeof(T));
|
||||
|
||||
// OutQue
|
||||
pipe->InitBuffer(
|
||||
yQue, 2, tilingData->rowFactor * RoundUp<T>(tilingData->dFactor) * sizeof(T));
|
||||
pipe->InitBuffer(postQue, 2, tilingData->rowFactor * tilingData->hcMultAlign * sizeof(float));
|
||||
pipe->InitBuffer(
|
||||
combFragQue, 2, tilingData->rowFactor * tilingData->hcMult * tilingData->hcMultAlign * sizeof(float));
|
||||
|
||||
// TBuf
|
||||
pipe->InitBuffer(hcBaseBuf0, tilingData->hcMultAlign * sizeof(float));
|
||||
pipe->InitBuffer(hcBaseBuf1, tilingData->hcMultAlign * sizeof(float));
|
||||
pipe->InitBuffer(hcBaseBuf2, tilingData->hcMult * tilingData->hcMultAlign * sizeof(float));
|
||||
|
||||
hcBase0Local = hcBaseBuf0.Get<float>();
|
||||
hcBase1Local = hcBaseBuf1.Get<float>();
|
||||
hcBase2Local = hcBaseBuf2.Get<float>();
|
||||
}
|
||||
|
||||
__aicore__ inline void Process()
|
||||
{
|
||||
int64_t curBlockIdx = GetBlockIdx();
|
||||
int64_t totalBlockNum = GetBlockNum();
|
||||
|
||||
int64_t rowOuterLoop =
|
||||
(curBlockIdx == totalBlockNum - 1) ? tilingData->rowLoopOfTailBlock : tilingData->rowLoopOfFormerBlock;
|
||||
int64_t tailRowFactor = (curBlockIdx == totalBlockNum - 1) ? tilingData->tailRowFactorOfTailBlock :
|
||||
tilingData->tailRowFactorOfFormerBlock;
|
||||
|
||||
CopyIn(hcBaseGm, hcBase0Local, 1, tilingData->hcMult);
|
||||
CopyIn(hcBaseGm[tilingData->hcMult], hcBase1Local, 1, tilingData->hcMult);
|
||||
CopyIn(hcBaseGm[tilingData->hcMult * 2], hcBase2Local, tilingData->hcMult, tilingData->hcMult);
|
||||
event_t eventId = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE2_V));
|
||||
SetFlag<HardEvent::MTE2_V>(eventId);
|
||||
WaitFlag<HardEvent::MTE2_V>(eventId);
|
||||
int64_t mixGmBaseOffset = curBlockIdx * tilingData->rowOfFormerBlock * tilingData->hcMix;
|
||||
int64_t xGmBaseOffset = curBlockIdx * tilingData->rowOfFormerBlock * tilingData->hcMult * tilingData->d;
|
||||
for (int64_t rowOuterIdx = 0; rowOuterIdx < rowOuterLoop; rowOuterIdx++) {
|
||||
int64_t curRowFactor = (rowOuterIdx == rowOuterLoop - 1) ? tailRowFactor : tilingData->rowFactor;
|
||||
mixes01Local = mixesQue01.AllocTensor<float>();
|
||||
CopyIn(
|
||||
mixesGm[mixGmBaseOffset + rowOuterIdx * tilingData->rowFactor * tilingData->hcMix], mixes01Local,
|
||||
curRowFactor, tilingData->hcMult, tilingData->hcMix - tilingData->hcMult);
|
||||
CopyIn(
|
||||
mixesGm[mixGmBaseOffset + rowOuterIdx * tilingData->rowFactor * tilingData->hcMix + tilingData->hcMult],
|
||||
mixes01Local[tilingData->rowFactor * tilingData->hcMultAlign], curRowFactor, tilingData->hcMult,
|
||||
tilingData->hcMix - tilingData->hcMult);
|
||||
mixesQue01.EnQue(mixes01Local);
|
||||
|
||||
rsqrtLocal = rsqrtQue.AllocTensor<float>();
|
||||
CopyIn(
|
||||
rsqrtGm[curBlockIdx * tilingData->rowOfFormerBlock + rowOuterIdx * tilingData->rowFactor], rsqrtLocal,
|
||||
1, curRowFactor);
|
||||
rsqrtQue.EnQue(rsqrtLocal);
|
||||
|
||||
mixes01Local = mixesQue01.DeQue<float>();
|
||||
rsqrtLocal = rsqrtQue.DeQue<float>();
|
||||
VFProcessPre(
|
||||
mixes01Local, mixes01Local, hcBase0Local, rsqrtLocal, hcScaleGm.GetValue(0), tilingData->eps,
|
||||
curRowFactor, tilingData->hcMult);
|
||||
for (int64_t dLoopIdx = 0; dLoopIdx < tilingData->dLoop; dLoopIdx++) {
|
||||
int64_t curDFactor =
|
||||
(dLoopIdx == tilingData->dLoop - 1) ? tilingData->tailDFactor : tilingData->dFactor;
|
||||
xLocal = xQue.template AllocTensor<T>();
|
||||
CopyIn(
|
||||
xGm[xGmBaseOffset + rowOuterIdx * tilingData->rowFactor * tilingData->hcMult * tilingData->d +
|
||||
dLoopIdx * tilingData->dFactor],
|
||||
xLocal, tilingData->rowFactor * tilingData->hcMult, curDFactor, tilingData->d - curDFactor);
|
||||
xQue.template EnQue(xLocal);
|
||||
xLocal = xQue.template DeQue<T>();
|
||||
yLocal = yQue.template AllocTensor<T>();
|
||||
VFProcessY(yLocal, mixes01Local, xLocal, curRowFactor, tilingData->hcMult, curDFactor);
|
||||
xQue.template FreeTensor(xLocal);
|
||||
yQue.template EnQue(yLocal);
|
||||
yLocal = yQue.template DeQue<T>();
|
||||
CopyOut(yLocal, yGm[curBlockIdx * tilingData->rowOfFormerBlock * tilingData->d + rowOuterIdx * tilingData->rowFactor * tilingData->d + dLoopIdx * tilingData->dFactor], curRowFactor, curDFactor, tilingData->d - curDFactor);
|
||||
yQue.template FreeTensor(yLocal);
|
||||
}
|
||||
|
||||
// post
|
||||
postLocal = postQue.AllocTensor<float>();
|
||||
VFProcessPost(
|
||||
postLocal, mixes01Local[tilingData->rowFactor * tilingData->hcMultAlign], hcBase1Local, rsqrtLocal,
|
||||
hcScaleGm.GetValue(1), tilingData->eps, curRowFactor, tilingData->hcMult);
|
||||
mixesQue01.template FreeTensor(mixes01Local);
|
||||
postQue.EnQue(postLocal);
|
||||
postLocal = postQue.DeQue<float>();
|
||||
CopyOut(postLocal, postGm[curBlockIdx * tilingData->rowOfFormerBlock * tilingData->hcMult + rowOuterIdx * tilingData->rowFactor * tilingData->hcMult], curRowFactor, tilingData->hcMult);
|
||||
postQue.FreeTensor(postLocal);
|
||||
|
||||
// combFrag
|
||||
mixes2Local = mixesQue2.AllocTensor<float>();
|
||||
CopyInWithLoopMode(
|
||||
mixesGm
|
||||
[mixGmBaseOffset + rowOuterIdx * tilingData->rowFactor * tilingData->hcMix + tilingData->hcMult * 2],
|
||||
mixes2Local, curRowFactor, tilingData->hcMult, tilingData->hcMult, tilingData->hcMix);
|
||||
mixesQue2.EnQue(mixes2Local);
|
||||
mixes2Local = mixesQue2.DeQue<float>();
|
||||
|
||||
combFragLocal = combFragQue.AllocTensor<float>();
|
||||
VFProcessCombFragRLessVLUseFourUnfold(
|
||||
combFragLocal, mixes2Local, hcBase2Local, rsqrtLocal, hcScaleGm.GetValue(2), tilingData->eps,
|
||||
tilingData->iterTimes - 1, curRowFactor, tilingData->hcMult, tilingData->hcMult);
|
||||
mixesQue2.FreeTensor(mixes2Local);
|
||||
rsqrtQue.FreeTensor(rsqrtLocal);
|
||||
|
||||
combFragQue.EnQue(combFragLocal);
|
||||
combFragLocal = combFragQue.DeQue<float>();
|
||||
CopyOut(combFragLocal, combFragGm[curBlockIdx * tilingData->rowOfFormerBlock * tilingData->hcMult * tilingData->hcMult + rowOuterIdx * tilingData->rowFactor * tilingData->hcMult * tilingData->hcMult], curRowFactor * tilingData->hcMult, tilingData->hcMult);
|
||||
combFragQue.FreeTensor(combFragLocal);
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
TPipe* pipe;
|
||||
const HcPreSinkhornTilingData* tilingData;
|
||||
GlobalTensor<float> mixesGm;
|
||||
GlobalTensor<float> rsqrtGm;
|
||||
GlobalTensor<float> hcScaleGm;
|
||||
GlobalTensor<float> hcBaseGm;
|
||||
GlobalTensor<T> xGm;
|
||||
GlobalTensor<T> yGm;
|
||||
GlobalTensor<float> postGm;
|
||||
GlobalTensor<float> combFragGm;
|
||||
|
||||
TQue<QuePosition::VECIN, 1> mixesQue01;
|
||||
TQue<QuePosition::VECIN, 1> mixesQue2;
|
||||
TQue<QuePosition::VECIN, 1> rsqrtQue;
|
||||
TQue<QuePosition::VECIN, 1> xQue;
|
||||
TQue<QuePosition::VECOUT, 1> yQue;
|
||||
TQue<QuePosition::VECOUT, 1> postQue;
|
||||
TQue<QuePosition::VECOUT, 1> combFragQue;
|
||||
|
||||
TBuf<QuePosition::VECCALC> hcBaseBuf0;
|
||||
TBuf<QuePosition::VECCALC> hcBaseBuf1;
|
||||
TBuf<QuePosition::VECCALC> hcBaseBuf2;
|
||||
|
||||
LocalTensor<float> mixes01Local;
|
||||
LocalTensor<float> mixes2Local;
|
||||
LocalTensor<float> rsqrtLocal;
|
||||
LocalTensor<T> xLocal;
|
||||
LocalTensor<T> yLocal;
|
||||
LocalTensor<float> postLocal;
|
||||
LocalTensor<float> combFragLocal;
|
||||
LocalTensor<float> hcBase0Local;
|
||||
LocalTensor<float> hcBase1Local;
|
||||
LocalTensor<float> hcBase2Local;
|
||||
};
|
||||
|
||||
} // namespace HCPreSinkhorn
|
||||
|
||||
#endif
|
||||
Reference in New Issue
Block a user