namespace optiling {
struct ReduceAllCompileInfo {};
static inline uint32_t AlignUp(uint32_t value, uint32_t alignment)
{
return ((value + alignment - 1) / alignment) * alignment;
}
static inline uint32_t AlignDown(uint32_t value, uint32_t alignment)
{
return (value / alignment) * alignment;
}
static inline uint64_t AlignUp64(uint64_t value, uint64_t alignment)
{
return ((value + alignment - 1) / alignment) * alignment;
}
* @brief 从 dim 属性读取 axes 数据
*/
static ge::graphStatus GetReduceAxesFromAttr(
gert::TilingContext* context,
const gert::Shape& inputShape,
std::vector<int64_t>& normalizedAxes)
{
int64_t rank = static_cast<int64_t>(inputShape.GetDimNum());
normalizedAxes.clear();
auto attrs = context->GetAttrs();
if (attrs == nullptr) {
for (int64_t i = 0; i < rank; i++) {
normalizedAxes.push_back(i);
}
return ge::GRAPH_SUCCESS;
}
auto dimListPtr = attrs->GetListInt(0);
if (dimListPtr == nullptr) {
for (int64_t i = 0; i < rank; i++) {
normalizedAxes.push_back(i);
}
return ge::GRAPH_SUCCESS;
}
size_t axesNum = dimListPtr->GetSize();
if (axesNum == 0) {
for (int64_t i = 0; i < rank; i++) {
normalizedAxes.push_back(i);
}
return ge::GRAPH_SUCCESS;
}
const int64_t* axesData = dimListPtr->GetData();
if (axesData == nullptr) {
for (int64_t i = 0; i < rank; i++) {
normalizedAxes.push_back(i);
}
return ge::GRAPH_SUCCESS;
}
normalizedAxes.reserve(axesNum);
for (size_t i = 0; i < axesNum; ++i) {
int64_t axis = axesData[i];
if (axis < 0) {
axis += rank;
}
OP_CHECK_IF(axis < 0 || axis >= rank,
OP_LOGE(context, "[ReduceAll] axis %ld out of range [0, %ld)", axis, rank),
return ge::GRAPH_FAILED);
normalizedAxes.push_back(axis);
}
std::sort(normalizedAxes.begin(), normalizedAxes.end());
normalizedAxes.erase(std::unique(normalizedAxes.begin(), normalizedAxes.end()), normalizedAxes.end());
return ge::GRAPH_SUCCESS;
}
* @brief 解析并验证 axes;计算 totalInputSize;输出规范化 axes;标记空/全规约
*/
static bool ParseAndValidate(const gert::Shape& inputShape,
const std::vector<int64_t>& normalizedAxes,
uint64_t& totalInputSize,
bool& hasZeroDim,
bool& isFullReduce)
{
int64_t ndim = static_cast<int64_t>(inputShape.GetDimNum());
totalInputSize = 1;
hasZeroDim = false;
for (size_t i = 0; i < inputShape.GetDimNum(); ++i) {
int64_t dimVal = inputShape.GetDim(i);
if (dimVal == 0) {
hasZeroDim = true;
}
totalInputSize *= static_cast<uint64_t>(dimVal);
}
isFullReduce = (normalizedAxes.size() == static_cast<size_t>(ndim));
return true;
}
* @brief 轴融合:将连续的 R/K 轴合并
*/
static void FuseAxes(const gert::Shape& inputShape,
const std::vector<int64_t>& normalizedAxes,
std::vector<uint32_t>& fusedDims,
std::vector<uint32_t>& fusedIsReduce)
{
size_t ndim = inputShape.GetDimNum();
fusedDims.clear();
fusedIsReduce.clear();
if (ndim == 0) {
return;
}
std::vector<bool> isReduceAxis(ndim, false);
for (int64_t axis : normalizedAxes) {
isReduceAxis[static_cast<size_t>(axis)] = true;
}
bool currentIsReduce = isReduceAxis[ndim - 1];
uint64_t currentSize = static_cast<uint64_t>(inputShape.GetDim(ndim - 1));
for (int64_t i = static_cast<int64_t>(ndim) - 2; i >= 0; --i) {
size_t idx = static_cast<size_t>(i);
if (isReduceAxis[idx] == currentIsReduce) {
currentSize *= static_cast<uint64_t>(inputShape.GetDim(idx));
} else {
fusedDims.insert(fusedDims.begin(), static_cast<uint32_t>(currentSize));
fusedIsReduce.insert(fusedIsReduce.begin(), currentIsReduce ? 1 : 0);
currentIsReduce = isReduceAxis[idx];
currentSize = static_cast<uint64_t>(inputShape.GetDim(idx));
}
}
fusedDims.insert(fusedDims.begin(), static_cast<uint32_t>(currentSize));
fusedIsReduce.insert(fusedIsReduce.begin(), currentIsReduce ? 1 : 0);
}
* @brief 计算 fused 维的 inputStrides
*/
static void ComputeFusedInputStrides(const std::vector<uint32_t>& fusedDims,
std::vector<uint32_t>& inputStrides)
{
inputStrides.clear();
inputStrides.resize(fusedDims.size());
uint64_t stride = 1;
for (int64_t i = static_cast<int64_t>(fusedDims.size()) - 1; i >= 0; --i) {
inputStrides[static_cast<size_t>(i)] = static_cast<uint32_t>(stride);
stride *= static_cast<uint64_t>(fusedDims[static_cast<size_t>(i)]);
}
}
* @brief 计算 totalOutputSize + outputStrides
*/
static void ComputeOutputMetaByFused(const std::vector<uint32_t>& fusedDims,
const std::vector<uint32_t>& fusedIsReduce,
uint64_t& totalOutputSize,
std::vector<uint32_t>& outputStridesByFused)
{
totalOutputSize = 1;
outputStridesByFused.clear();
outputStridesByFused.resize(fusedDims.size(), 0);
std::vector<uint32_t> kDims;
std::vector<size_t> kPos;
for (size_t i = 0; i < fusedDims.size(); ++i) {
if (fusedIsReduce[i] == 0) {
kDims.push_back(fusedDims[i]);
kPos.push_back(i);
totalOutputSize *= static_cast<uint64_t>(fusedDims[i]);
}
}
if (kDims.empty()) {
totalOutputSize = 1;
return;
}
uint32_t stride = 1;
for (int64_t i = static_cast<int64_t>(kDims.size()) - 1; i >= 0; --i) {
size_t fusedIdx = kPos[static_cast<size_t>(i)];
outputStridesByFused[fusedIdx] = stride;
stride *= kDims[static_cast<size_t>(i)];
}
}
* @brief 确定规约模式
*/
static uint32_t DetermineReduceMode(const std::vector<uint32_t>& fusedIsReduce,
bool isFullReduce)
{
if (isFullReduce) {
return MODE_FULL_REDUCE;
}
size_t dimCount = fusedIsReduce.size();
if (dimCount == 0) {
return MODE_FULL_REDUCE;
}
if (dimCount == 2) {
if (fusedIsReduce[0] == 0 && fusedIsReduce[1] == 1) {
return MODE_KR;
}
if (fusedIsReduce[0] == 1 && fusedIsReduce[1] == 0) {
return MODE_RK;
}
}
if (dimCount == 3) {
if (fusedIsReduce[0] == 1 && fusedIsReduce[1] == 0 && fusedIsReduce[2] == 1) {
return MODE_RKR;
}
if (fusedIsReduce[0] == 0 && fusedIsReduce[1] == 1 && fusedIsReduce[2] == 0) {
return MODE_KRK;
}
}
return MODE_GENERAL;
}
static uint32_t ComputeMaxTileSize(uint64_t ubSize)
{
uint32_t maxTileBytsSize = static_cast<uint32_t>((ubSize - FIXED_EXPENSES) / UB_BUFFER_FACTOR);
maxTileBytsSize = AlignDown(maxTileBytsSize, BLOCK_SIZE);
if (maxTileBytsSize < MIN_TILE_SIZE) {
maxTileBytsSize = MIN_TILE_SIZE;
}
return maxTileBytsSize;
}
static void ComputeCoreAllocation(uint64_t totalElements,
uint32_t tileSize,
uint32_t availableCores,
uint32_t& coreNum,
uint64_t& elementsPerCore,
uint32_t& largeCoreCount)
{
if (totalElements == 0) {
coreNum = 1;
elementsPerCore = 0;
largeCoreCount = 0;
return;
}
uint64_t maxCores = (totalElements + tileSize - 1) / tileSize;
coreNum = (maxCores < availableCores) ? static_cast<uint32_t>(maxCores) : availableCores;
if (coreNum == 0) {
coreNum = 1;
}
elementsPerCore = totalElements / coreNum;
largeCoreCount = static_cast<uint32_t>(totalElements % coreNum);
}
static ge::graphStatus ReduceAllTilingFunc(gert::TilingContext* context)
{
OP_CHECK_IF(context == nullptr, OP_LOGE(context, "context is nullptr"), return ge::GRAPH_FAILED);
const gert::StorageShape* selfShape = context->GetInputShape(0);
OP_CHECK_IF(selfShape == nullptr, OP_LOGE(context, "selfShape is nullptr"), return ge::GRAPH_FAILED);
const gert::Shape& inputShape = selfShape->GetStorageShape();
std::vector<int64_t> normalizedAxes;
ge::graphStatus axesStatus = GetReduceAxesFromAttr(context, inputShape, normalizedAxes);
OP_CHECK_IF(axesStatus != ge::GRAPH_SUCCESS,
OP_LOGE(context, "GetReduceAxesFromAttr failed"),
return ge::GRAPH_FAILED);
auto ascendcPlatform = platform_ascendc::PlatformAscendC(context->GetPlatformInfo());
uint64_t ubSize = 0;
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize);
uint32_t availableCores = ascendcPlatform.GetCoreNum();
ReduceAllTilingData* tiling = context->GetTilingData<ReduceAllTilingData>();
OP_CHECK_NULL_WITH_CONTEXT(context, tiling);
OP_CHECK_IF(memset_s(tiling, sizeof(ReduceAllTilingData), 0, sizeof(ReduceAllTilingData)) != EOK,
OP_LOGE(context, "set tiling data error"), return ge::GRAPH_FAILED);
uint64_t totalInputSize = 0;
bool hasZeroDim = false;
bool isFullReduce = false;
if (!ParseAndValidate(inputShape, normalizedAxes, totalInputSize, hasZeroDim, isFullReduce)) {
OP_LOGE(context, "ParseAndValidate failed");
return ge::GRAPH_FAILED;
}
tiling->totalInputSize = totalInputSize;
if (hasZeroDim || totalInputSize == 0) {
tiling->totalOutputSize = isFullReduce ? 1 : 0;
tiling->workGmSize = AlignUp64(static_cast<uint64_t>(tiling->totalOutputSize), 64);
tiling->coreNum = 1;
tiling->ubSize = ubSize;
tiling->tileSize = ComputeMaxTileSize(ubSize);
tiling->elementsPerCore = 0;
tiling->largeCoreCount = 0;
tiling->reduceMode = MODE_FULL_REDUCE;
context->SetBlockDim(tiling->coreNum);
context->SetTilingKey(GET_TPL_TILING_KEY(REDUCE_ALL_KEY_BOOL));
uint32_t sysWorkspaceSize = ascendcPlatform.GetLibApiWorkSpaceSize();
size_t* currentWorkspace = context->GetWorkspaceSizes(1);
size_t workGmSize = static_cast<size_t>(tiling->coreNum) * tiling->workGmSize;
currentWorkspace[0] = sysWorkspaceSize + workGmSize + 64;
return ge::GRAPH_SUCCESS;
}
std::vector<uint32_t> fusedDims;
std::vector<uint32_t> fusedIsReduce;
FuseAxes(inputShape, normalizedAxes, fusedDims, fusedIsReduce);
tiling->fusedDimCount = static_cast<uint32_t>(fusedDims.size());
OP_CHECK_IF(tiling->fusedDimCount > MAX_DIMS, OP_LOGE(context, "fusedDimCount overflow"), return ge::GRAPH_FAILED);
for (size_t i = 0; i < fusedDims.size(); ++i) {
tiling->fusedDims[i] = fusedDims[i];
tiling->fusedAxes[i] = fusedIsReduce[i];
}
std::vector<uint32_t> inputStrides;
ComputeFusedInputStrides(fusedDims, inputStrides);
for (size_t i = 0; i < inputStrides.size(); ++i) {
tiling->inputStrides[i] = inputStrides[i];
}
uint64_t totalOutputSize = 0;
std::vector<uint32_t> outputStridesByFused;
ComputeOutputMetaByFused(fusedDims, fusedIsReduce, totalOutputSize, outputStridesByFused);
tiling->totalOutputSize = totalOutputSize;
for (size_t i = 0; i < outputStridesByFused.size(); ++i) {
tiling->outputStrides[i] = outputStridesByFused[i];
}
tiling->reduceMode = DetermineReduceMode(fusedIsReduce, isFullReduce);
tiling->ubSize = ubSize;
tiling->tileSize = ComputeMaxTileSize(ubSize);
ComputeCoreAllocation(totalInputSize, tiling->tileSize, availableCores,
tiling->coreNum, tiling->elementsPerCore, tiling->largeCoreCount);
tiling->workGmSize = AlignUp64(totalOutputSize, 64);
context->SetBlockDim(tiling->coreNum);
context->SetTilingKey(GET_TPL_TILING_KEY(REDUCE_ALL_KEY_BOOL));
size_t usrSize = 0;
uint32_t sysWorkspaceSize = ascendcPlatform.GetLibApiWorkSpaceSize();
size_t* currentWorkspace = context->GetWorkspaceSizes(1);
currentWorkspace[0] = usrSize + sysWorkspaceSize + static_cast<size_t>(tiling->coreNum) * tiling->workGmSize;
return ge::GRAPH_SUCCESS;
}
static ge::graphStatus TilingParseForReduceAll([[maybe_unused]] gert::TilingParseContext* context)
{
return ge::GRAPH_SUCCESS;
}
IMPL_OP_OPTILING(ReduceAll).Tiling(ReduceAllTilingFunc).TilingParse<ReduceAllCompileInfo>(TilingParseForReduceAll);
}