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
|
||||
Reference in New Issue
Block a user