init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View 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()

View 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()

View 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

View 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

View 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

View 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

View 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();
}
}

View 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

View 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

View File

@@ -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

View File

@@ -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