30
csrc/common/stub/op_tiling/CMakeLists.txt
Normal file
30
csrc/common/stub/op_tiling/CMakeLists.txt
Normal file
@@ -0,0 +1,30 @@
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
# 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.
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
|
||||
set(CMAKE_LIBRARY_OUTPUT_DIRECTORY ${PROJECT_SOURCE_DIR}/build)
|
||||
|
||||
file(GLOB_RECURSE OP_TILING_FILES "*.cpp")
|
||||
|
||||
add_library(optiling SHARED ${OP_TILING_FILES})
|
||||
|
||||
add_dependencies(optiling json)
|
||||
|
||||
target_compile_definitions(optiling PRIVATE
|
||||
_GLIBCXX_USE_CXX11_ABI=0
|
||||
LOG_CPP
|
||||
)
|
||||
|
||||
target_include_directories(optiling PRIVATE
|
||||
${OP_TILING_INCLUDE}
|
||||
)
|
||||
|
||||
target_link_libraries(optiling PRIVATE
|
||||
$<BUILD_INTERFACE:dlog_headers>
|
||||
)
|
||||
198
csrc/common/stub/op_tiling/op_cache_def_tiling.h
Normal file
198
csrc/common/stub/op_tiling/op_cache_def_tiling.h
Normal file
@@ -0,0 +1,198 @@
|
||||
/**
|
||||
* 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 op_cache_tiling.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef OPS_BUILT_IN_OP_TILING_OP_CACHE_DEF_TILING_H
|
||||
#define OPS_BUILT_IN_OP_TILING_OP_CACHE_DEF_TILING_H
|
||||
|
||||
#include <array>
|
||||
#include "exe_graph/runtime/tiling_context.h"
|
||||
|
||||
namespace optiling {
|
||||
struct BatchmatmulCompileParas {
|
||||
bool binary_mode_flag = false;
|
||||
bool bias_flag = false;
|
||||
bool at_l1_flag = true;
|
||||
bool split_k_flag = false;
|
||||
bool pattern_flag = false;
|
||||
bool zero_flag = false;
|
||||
bool sparse_4to2_flag = false;
|
||||
bool binary_constant_flag = false;
|
||||
bool vector_pre_conv_mode = false;
|
||||
float fused_double_operand_num = 0;
|
||||
float aub_double_num = 0;
|
||||
float bub_double_num = 0;
|
||||
int64_t quant_scale = 0;
|
||||
int64_t eltwise_src = 0;
|
||||
int8_t enable_pad = 0;
|
||||
bool enable_nz_fusion = false;
|
||||
bool enable_rt_bank_cache = false;
|
||||
};
|
||||
|
||||
struct BatchmatmulRunParas {
|
||||
bool nd_flag = false;
|
||||
bool use_pre_ub = false;
|
||||
bool trans_a_flag = false;
|
||||
bool trans_b_flag = false;
|
||||
bool format_a_nd = false;
|
||||
bool format_b_nd = false;
|
||||
bool format_out_nd = false;
|
||||
ge::Format format_a = ge::FORMAT_ND;
|
||||
ge::Format format_b = ge::FORMAT_ND;
|
||||
ge::Format format_out = ge::FORMAT_ND;
|
||||
bool reserved_bool = false;
|
||||
bool b_have_batch = false; // dim num > 2
|
||||
bool is_batch_matmul_mode = false; // dynamic_mode == "dynamic_mknb"
|
||||
bool is_batch_matmul_op = false; // BatchMatMulV2 or BatchMatMul
|
||||
bool used_aligned_pattern = false;
|
||||
bool non_factor_k = false;
|
||||
bool non_factor_bmn = false;
|
||||
bool bias_flag = false;
|
||||
bool pattern_flag = false;
|
||||
bool do_not_multi_batch = false;
|
||||
bool performance_flag = false;
|
||||
bool unaligned_flag = false;
|
||||
bool zero_flag = false;
|
||||
bool is_compress_quant = false;
|
||||
bool is_bmm_fixp = false;
|
||||
bool enable_nz_fusion = false;
|
||||
bool weight_nz_flag = false;
|
||||
int8_t enable_pad = 0;
|
||||
int8_t hf32_flag = 1;
|
||||
int8_t pad_flag = 0;
|
||||
int8_t nz_fusion_flag = 0;
|
||||
int32_t dtype_a = 0;
|
||||
int32_t dtype_b = 0;
|
||||
int32_t dtype_out = 0;
|
||||
int32_t dtype_bias = 0;
|
||||
int64_t m_mapped = 1;
|
||||
int64_t k_mapped = 1;
|
||||
int64_t n_mapped = 1;
|
||||
int64_t batch_mapped = 1;
|
||||
int64_t m = 1;
|
||||
int64_t k = 1;
|
||||
int64_t n = 1;
|
||||
int64_t batch = 1;
|
||||
int64_t ori_shape_m = 1;
|
||||
int64_t ori_shape_k = 1;
|
||||
int64_t ori_shape_n = 1;
|
||||
int64_t m_pad = 0;
|
||||
int64_t k_pad = 0;
|
||||
int64_t n_pad = 0;
|
||||
int64_t nl0 = 1;
|
||||
int64_t kl0 = 1;
|
||||
int64_t dim0_a = 0;
|
||||
int64_t dim1_a = 0;
|
||||
int64_t dim2_a = 0;
|
||||
int64_t dim0_b = 0;
|
||||
int64_t dim1_b = 0;
|
||||
int64_t dim2_b = 0;
|
||||
int64_t batch_a1 = 1;
|
||||
int64_t batch_a2 = 1;
|
||||
int64_t batch_a3 = 1;
|
||||
int64_t batch_a4 = 1;
|
||||
int64_t batch_b1 = 1;
|
||||
int64_t batch_b2 = 1;
|
||||
int64_t batch_b3 = 1;
|
||||
int64_t batch_b4 = 1;
|
||||
int64_t batch_c1 = 1;
|
||||
int64_t batch_c2 = 1;
|
||||
int64_t batch_c3 = 1;
|
||||
int64_t batch_c4 = 1;
|
||||
int32_t offset_x = 0;
|
||||
int32_t index_size = 0;
|
||||
bool m_quant_check = false;
|
||||
bool n_quant_check = false;
|
||||
bool is_weight_quant_bmm = false;
|
||||
bool vector_pre_conv_mode = false;
|
||||
bool is_quant_batch_matmul_v3 = false;
|
||||
bool is_weight_quant_batch_matmul_v2 = false;
|
||||
bool is_pertoken = false;
|
||||
// 3 is perm_a dim
|
||||
std::array<size_t, 3> perm_a = {0, 0, 0};
|
||||
// 3 is perm_b dim
|
||||
std::array<size_t, 3> perm_b = {0, 0, 0};
|
||||
ge::DataType bias_dtype = ge::DT_FLOAT16;
|
||||
};
|
||||
|
||||
class CacheTilingData
|
||||
{
|
||||
public:
|
||||
uint64_t tiling_id;
|
||||
int64_t n_cub = 1;
|
||||
int64_t db_cub = 1;
|
||||
int64_t m_l0 = 1;
|
||||
int64_t k_l0 = 1;
|
||||
int64_t n_l0 = 1;
|
||||
int64_t batch_dim = 1;
|
||||
int64_t n_dim = 1;
|
||||
int64_t m_dim = 1;
|
||||
int64_t k_dim = 1;
|
||||
int64_t kal1_16 = 1;
|
||||
int64_t kbl1_16 = 1;
|
||||
int64_t kal1_factor = 1;
|
||||
int64_t kbl1_factor = 1;
|
||||
int64_t m_al1 = 1;
|
||||
int64_t n_bl1 = 1;
|
||||
int64_t db_al1 = 1;
|
||||
int64_t db_bl1 = 1;
|
||||
int64_t k_aub = 1;
|
||||
int64_t m_aub = 1;
|
||||
int64_t db_aub = 1;
|
||||
int64_t k_bub = 1;
|
||||
int64_t n_bub = 1;
|
||||
int64_t db_bub = 1;
|
||||
int64_t aub_dim = 1;
|
||||
int64_t bub_dim = 1;
|
||||
int64_t m1_aub = 1;
|
||||
int64_t n1_bub = 1;
|
||||
int64_t k1_aub = 1;
|
||||
int64_t k1_bub = 1;
|
||||
int64_t m_aub_dim = 1;
|
||||
int64_t n_bub_dim = 1;
|
||||
int64_t k_aub_dim = 1;
|
||||
int64_t k_bub_dim = 1;
|
||||
int64_t k_org_dim = 1;
|
||||
int64_t db_l0c = 1;
|
||||
int64_t batch_l0 = 1;
|
||||
int64_t batch_aub = 1;
|
||||
int64_t batch_bub = 1;
|
||||
int64_t batch_cub = 1;
|
||||
int32_t out_branch_flag = 1;
|
||||
int32_t bias_flag = 0;
|
||||
int32_t aub_multi_flag = 0;
|
||||
int32_t bub_multi_flag = 0;
|
||||
int64_t a_align_value = 1;
|
||||
int64_t b_align_value = 1;
|
||||
int64_t aub_align_bound = 0;
|
||||
int64_t bub_align_bound = 0;
|
||||
int64_t min_kl1_cmp_kl0 = 0;
|
||||
int32_t al1_attach_flag = 0;
|
||||
int32_t bl1_attach_flag = 0;
|
||||
int32_t abkl1_attach_flag = 0;
|
||||
int32_t l0c_multi_batch = 0;
|
||||
int64_t m_single_core = 1;
|
||||
int64_t n_single_core = 1;
|
||||
bool flag_cub_solving_bank_conflict = false;
|
||||
bool al1_full_load = false;
|
||||
bool bl1_full_load = false;
|
||||
int8_t hf32_flag = 1;
|
||||
int32_t zero_flag = 0;
|
||||
bool datatype_bf16 = false;
|
||||
uint64_t deq_scale_var = 0x3F800000;
|
||||
uint32_t l2_cache_flag = 0;
|
||||
};
|
||||
} // namespace optiling
|
||||
|
||||
#endif
|
||||
35
csrc/common/stub/op_tiling/op_cache_tiling.cpp
Normal file
35
csrc/common/stub/op_tiling/op_cache_tiling.cpp
Normal file
@@ -0,0 +1,35 @@
|
||||
/**
|
||||
* 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 op_cache_tiling.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "op_cache_tiling.h"
|
||||
|
||||
namespace optiling {
|
||||
bool TilingPrepareForOpCache(gert::TilingContext* /*context*/)
|
||||
{
|
||||
return true;
|
||||
}
|
||||
|
||||
bool TilingPrepareForOpCache(gert::TilingParseContext* /*context*/)
|
||||
{
|
||||
return true;
|
||||
}
|
||||
|
||||
bool GenTiling(
|
||||
const std::string& /*op_type*/, const BatchmatmulCompileParas& /*compile_params*/,
|
||||
BatchmatmulRunParas& /*run_params*/, CacheTilingData& /*tiling*/, gert::TilingContext* /*context*/)
|
||||
{
|
||||
return true;
|
||||
}
|
||||
|
||||
} // namespace optiling
|
||||
34
csrc/common/stub/op_tiling/op_cache_tiling.h
Normal file
34
csrc/common/stub/op_tiling/op_cache_tiling.h
Normal file
@@ -0,0 +1,34 @@
|
||||
/**
|
||||
* 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 cop_ache_tiling.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef OPS_BUILT_IN_OP_TILING_OP_CACHE_TILING_H
|
||||
#define OPS_BUILT_IN_OP_TILING_OP_CACHE_TILING_H
|
||||
|
||||
#include <array>
|
||||
#include "exe_graph/runtime/tiling_context.h"
|
||||
#include "exe_graph/runtime/tiling_parse_context.h"
|
||||
#include "op_cache_def_tiling.h"
|
||||
|
||||
namespace optiling {
|
||||
|
||||
bool TilingPrepareForOpCache(gert::TilingContext* context);
|
||||
bool TilingPrepareForOpCache(gert::TilingParseContext* context);
|
||||
|
||||
bool GenTiling(
|
||||
const std::string& op_type, const BatchmatmulCompileParas& compile_params, BatchmatmulRunParas& run_params,
|
||||
CacheTilingData& tiling, gert::TilingContext* context);
|
||||
} // namespace optiling
|
||||
|
||||
#endif
|
||||
209
csrc/common/stub/op_tiling/register/tuning_bank_key_registry.h
Normal file
209
csrc/common/stub/op_tiling/register/tuning_bank_key_registry.h
Normal file
@@ -0,0 +1,209 @@
|
||||
/**
|
||||
* 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.
|
||||
*/
|
||||
|
||||
#ifndef __INC_REGISTER_TUNING_BANK_KEY_REGISTRY_HEADER__
|
||||
#define __INC_REGISTER_TUNING_BANK_KEY_REGISTRY_HEADER__
|
||||
#include <memory>
|
||||
#include <unordered_map>
|
||||
#include <nlohmann/json.hpp>
|
||||
#include <string>
|
||||
#include "graph/ascend_string.h"
|
||||
#include "register/register_types.h"
|
||||
#include "exe_graph/runtime/tiling_context.h"
|
||||
|
||||
// v1 stub
|
||||
#define REGISTER_OP_BANK_KEY_CONVERT_FUN(op, opfunc) REGISTER_OP_BANK_KEY_CONVERT_FUN_UNIQ_HELPER(op, (opfunc))
|
||||
|
||||
#define REGISTER_OP_BANK_KEY_CONVERT_FUN_UNIQ_HELPER(optype, opfunc) REGISTER_OP_BANK_KEY_UNIQ(optype, (opfunc))
|
||||
|
||||
#define REGISTER_OP_BANK_KEY_UNIQ(optype, opfunc) \
|
||||
static tuningtiling::OpBankKeyFuncRegistry g_##optype##BankKeyRegistryInterf(#optype, (opfunc))
|
||||
|
||||
#define REGISTER_OP_BANK_KEY_PARSE_FUN(op, parse_func, load_func) \
|
||||
REGISTER_OP_BANK_KEY_PARSE_FUN_UNIQ_HELPER(op, (parse_func), (load_func))
|
||||
|
||||
#define REGISTER_OP_BANK_KEY_PARSE_FUN_UNIQ_HELPER(optype, parse_func, load_func) \
|
||||
REGISTER_OP_BANK_KEY_PARSE_UNIQ(optype, (parse_func), (load_func))
|
||||
|
||||
#define REGISTER_OP_BANK_KEY_PARSE_UNIQ(optype, parse_func, load_func) \
|
||||
static tuningtiling::OpBankKeyFuncRegistry g_##optype##BankParseInterf(#optype, (parse_func), (load_func))
|
||||
|
||||
// v2
|
||||
#define REGISTER_OP_BANK_KEY_CONVERT_FUN_V2(op, opfunc) REGISTER_OP_BANK_KEY_CONVERT_FUN_UNIQ_HELPER_V2(op, (opfunc))
|
||||
|
||||
#define REGISTER_OP_BANK_KEY_CONVERT_FUN_UNIQ_HELPER_V2(optype, opfunc) REGISTER_OP_BANK_KEY_UNIQ_V2(optype, (opfunc))
|
||||
|
||||
#define REGISTER_OP_BANK_KEY_UNIQ_V2(optype, opfunc) \
|
||||
static tuningtiling::OpBankKeyFuncRegistryV2 g_##optype##BankKeyRegistryInterf(#optype, (opfunc))
|
||||
|
||||
#define REGISTER_OP_BANK_KEY_PARSE_FUN_V2(op, parse_func, load_func) \
|
||||
REGISTER_OP_BANK_KEY_PARSE_FUN_UNIQ_HELPER_V2(op, (parse_func), (load_func))
|
||||
|
||||
#define REGISTER_OP_BANK_KEY_PARSE_FUN_UNIQ_HELPER_V2(optype, parse_func, load_func) \
|
||||
REGISTER_OP_BANK_KEY_PARSE_UNIQ_V2(optype, (parse_func), (load_func))
|
||||
|
||||
#define REGISTER_OP_BANK_KEY_PARSE_UNIQ_V2(optype, parse_func, load_func) \
|
||||
static tuningtiling::OpBankKeyFuncRegistryV2 g_##optype##BankParseInterf(#optype, (parse_func), (load_func))
|
||||
|
||||
#define TUNING_TILING_MAKE_SHARED(exec_expr0, exec_expr1) \
|
||||
do { \
|
||||
try { \
|
||||
exec_expr0; \
|
||||
} catch (...) { \
|
||||
exec_expr1; \
|
||||
} \
|
||||
} while (0)
|
||||
|
||||
// v1 stub
|
||||
#define DECLARE_STRUCT_RELATE_WITH_OP(op, bank_key, ...) \
|
||||
do { \
|
||||
NLOHMANN_DEFINE_TYPE_NON_INTRUSIVE(bank_key, __VA_ARGS__); \
|
||||
static bool ParseFunc##op##bank_key( \
|
||||
const std::shared_ptr<void>& in_args, size_t len, ge::AscendString& bank_key_str) \
|
||||
{ \
|
||||
if (sizeof(bank_key_str) != len || in_args == nullptr) { \
|
||||
return false; \
|
||||
} \
|
||||
return false; \
|
||||
} \
|
||||
static bool LoadFunc##op##bank_key( \
|
||||
std::shared_ptr<void>& in_args, size_t& len, const ge::AscendString& bank_key_str) \
|
||||
{ \
|
||||
len = sizeof(bank_key_str); \
|
||||
TUNING_TILING_MAKE_SHARED(in_args = std::make_shared<bank_key>(), return false); \
|
||||
auto op_ky = std::static_pointer_cast<bank_key>(in_args); \
|
||||
return false; \
|
||||
} \
|
||||
REGISTER_OP_BANK_KEY_PARSE_FUN(op, ParseFunc##op##bank_key, LoadFunc##op##bank_key) \
|
||||
} while (0)
|
||||
|
||||
|
||||
// v2
|
||||
#define DECLARE_STRUCT_RELATE_WITH_OP_V2(op, bank_key, ...) \
|
||||
do { \
|
||||
NLOHMANN_DEFINE_TYPE_NON_INTRUSIVE(bank_key, __VA_ARGS__); \
|
||||
static bool ParseFuncV2##op##bank_key( \
|
||||
const std::shared_ptr<void>& in_args, size_t len, ge::AscendString& bank_key_json_str) \
|
||||
{ \
|
||||
if (sizeof(bank_key) != len || in_args == nullptr) { \
|
||||
return false; \
|
||||
} \
|
||||
nlohmann::json bank_key_json; \
|
||||
bank_key_json = *(std::static_pointer_cast<bank_key>(in_args)); \
|
||||
try { \
|
||||
std::string json_dump_str = bank_key_json.dump(); \
|
||||
bank_key_json_str = ge::AscendString(json_dump_str.c_str()); \
|
||||
} catch (std::exception & e) { \
|
||||
return false; \
|
||||
} \
|
||||
return true; \
|
||||
} \
|
||||
static bool LoadFuncV2##op##bank_key( \
|
||||
std::shared_ptr<void>& in_args, size_t& len, const ge::AscendString& bank_key_json_str) \
|
||||
{ \
|
||||
len = sizeof(bank_key); \
|
||||
TUNING_TILING_MAKE_SHARED(in_args = std::make_shared<bank_key>(), return false); \
|
||||
nlohmann::json bank_key_json; \
|
||||
try { \
|
||||
bank_key_json = nlohmann::json::parse(bank_key_json_str.GetString()); \
|
||||
auto op_ky = std::static_pointer_cast<bank_key>(in_args); \
|
||||
*op_ky = bank_key_json.get<bank_key>(); \
|
||||
} catch (std::exception & e) { \
|
||||
return false; \
|
||||
} \
|
||||
return true; \
|
||||
} \
|
||||
REGISTER_OP_BANK_KEY_PARSE_FUN_V2(op, ParseFuncV2##op##bank_key, LoadFuncV2##op##bank_key) \
|
||||
} while (0)
|
||||
|
||||
|
||||
namespace tuningtiling {
|
||||
// v1兼容老版本om
|
||||
using OpBankKeyConvertFun = std::function<bool(const gert::TilingContext*, std::shared_ptr<void>&, size_t&)>;
|
||||
using OpBankParseFun = std::function<bool(const std::shared_ptr<void>&, size_t, ge::AscendString&)>;
|
||||
using OpBankLoadFun = std::function<bool(std::shared_ptr<void>&, size_t&, const ge::AscendString&)>;
|
||||
|
||||
// v2
|
||||
using OpBankKeyConvertFunV2 = std::function<bool(const gert::TilingContext*, std::shared_ptr<void>&, size_t&)>;
|
||||
using OpBankParseFunV2 = std::function<bool(const std::shared_ptr<void>&, size_t, ge::AscendString&)>;
|
||||
using OpBankLoadFunV2 = std::function<bool(std::shared_ptr<void>&, size_t&, const ge::AscendString&)>;
|
||||
// v1兼容老版本om
|
||||
class FMK_FUNC_HOST_VISIBILITY OpBankKeyFuncInfo
|
||||
{
|
||||
public:
|
||||
explicit OpBankKeyFuncInfo(const ge::AscendString& optype);
|
||||
OpBankKeyFuncInfo() = default;
|
||||
~OpBankKeyFuncInfo() = default;
|
||||
void SetOpConvertFunc(const OpBankKeyConvertFun& convert_func);
|
||||
void SetOpParseFunc(const OpBankParseFun& parse_func);
|
||||
void SetOpLoadFunc(const OpBankLoadFun& load_func);
|
||||
const OpBankKeyConvertFun& GetBankKeyConvertFunc() const;
|
||||
const OpBankParseFun& GetBankKeyParseFunc() const;
|
||||
const OpBankLoadFun& GetBankKeyLoadFunc() const;
|
||||
const ge::AscendString& GetOpType() const
|
||||
{
|
||||
return optype_;
|
||||
}
|
||||
|
||||
private:
|
||||
ge::AscendString optype_;
|
||||
OpBankKeyConvertFun convert_func_;
|
||||
OpBankParseFun parse_func_;
|
||||
OpBankLoadFun load_func_;
|
||||
};
|
||||
|
||||
// v2
|
||||
class FMK_FUNC_HOST_VISIBILITY OpBankKeyFuncInfoV2
|
||||
{
|
||||
public:
|
||||
explicit OpBankKeyFuncInfoV2(const ge::AscendString& optypeV2);
|
||||
OpBankKeyFuncInfoV2() = default;
|
||||
~OpBankKeyFuncInfoV2() = default;
|
||||
void SetOpConvertFuncV2(const OpBankKeyConvertFunV2& convert_funcV2);
|
||||
void SetOpParseFuncV2(const OpBankParseFunV2& parse_funcV2);
|
||||
void SetOpLoadFuncV2(const OpBankLoadFunV2& load_funcV2);
|
||||
const OpBankKeyConvertFunV2& GetBankKeyConvertFuncV2() const;
|
||||
const OpBankParseFunV2& GetBankKeyParseFuncV2() const;
|
||||
const OpBankLoadFunV2& GetBankKeyLoadFuncV2() const;
|
||||
const ge::AscendString& GetOpTypeV2() const
|
||||
{
|
||||
return optypeV2_;
|
||||
}
|
||||
|
||||
private:
|
||||
ge::AscendString optypeV2_;
|
||||
OpBankKeyConvertFunV2 convert_funcV2_;
|
||||
OpBankParseFunV2 parse_funcV2_;
|
||||
OpBankLoadFunV2 load_funcV2_;
|
||||
};
|
||||
|
||||
// v1兼容老版本om
|
||||
class FMK_FUNC_HOST_VISIBILITY OpBankKeyFuncRegistry
|
||||
{
|
||||
public:
|
||||
OpBankKeyFuncRegistry(const ge::AscendString& optype, const OpBankKeyConvertFun& convert_func);
|
||||
OpBankKeyFuncRegistry(
|
||||
const ge::AscendString& optype, const OpBankParseFun& parse_func, const OpBankLoadFun& load_func);
|
||||
~OpBankKeyFuncRegistry() = default;
|
||||
static std::unordered_map<ge::AscendString, OpBankKeyFuncInfo>& RegisteredOpFuncInfo();
|
||||
};
|
||||
|
||||
// v2
|
||||
class FMK_FUNC_HOST_VISIBILITY OpBankKeyFuncRegistryV2
|
||||
{
|
||||
public:
|
||||
OpBankKeyFuncRegistryV2(const ge::AscendString& optype, const OpBankKeyConvertFunV2& convert_funcV2);
|
||||
OpBankKeyFuncRegistryV2(
|
||||
const ge::AscendString& optype, const OpBankParseFunV2& parse_funcV2, const OpBankLoadFunV2& load_funcV2);
|
||||
~OpBankKeyFuncRegistryV2() = default;
|
||||
static std::unordered_map<ge::AscendString, OpBankKeyFuncInfoV2>& RegisteredOpFuncInfoV2();
|
||||
};
|
||||
} // namespace tuningtiling
|
||||
#endif
|
||||
@@ -0,0 +1,191 @@
|
||||
/**
|
||||
* 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.
|
||||
*/
|
||||
|
||||
#ifndef __INC_REGISTER_TUNING_TILING_REFLECTION_UTILS_HEADER__
|
||||
#define __INC_REGISTER_TUNING_TILING_REFLECTION_UTILS_HEADER__
|
||||
#include <string>
|
||||
#include <type_traits>
|
||||
#include <tuple>
|
||||
#include <nlohmann/json.hpp>
|
||||
|
||||
namespace tuningtiling {
|
||||
// implement for std c++11
|
||||
template <class T>
|
||||
using decay_t = typename std::decay<T>::type;
|
||||
|
||||
template <bool B, class T = void>
|
||||
using enable_if_t = typename std::enable_if<B, T>::type;
|
||||
|
||||
template <typename T, T... Ints>
|
||||
struct integer_sequence {
|
||||
using value_type = T;
|
||||
static constexpr std::size_t size()
|
||||
{
|
||||
return sizeof...(Ints);
|
||||
}
|
||||
};
|
||||
|
||||
template <std::size_t... Ints>
|
||||
using index_sequence = integer_sequence<std::size_t, Ints...>;
|
||||
|
||||
template <typename T, std::size_t N, T... Is>
|
||||
struct make_integer_sequence : make_integer_sequence<T, N - 1U, N - 1U, Is...> {
|
||||
};
|
||||
|
||||
template <typename T, T... Is>
|
||||
struct make_integer_sequence<T, 0, Is...> : integer_sequence<T, Is...> {
|
||||
};
|
||||
|
||||
template <std::size_t N>
|
||||
using make_index_sequence = make_integer_sequence<std::size_t, N>;
|
||||
|
||||
template <typename T>
|
||||
struct StructInfo {
|
||||
static std::tuple<> Info()
|
||||
{
|
||||
return std::make_tuple();
|
||||
}
|
||||
};
|
||||
|
||||
#define DECLARE_SCHEMA(Struct, ...) \
|
||||
template <> \
|
||||
struct StructInfo<Struct> { \
|
||||
static decltype(std::make_tuple(__VA_ARGS__)) Info() \
|
||||
{ \
|
||||
return std::make_tuple(__VA_ARGS__); \
|
||||
} \
|
||||
};
|
||||
|
||||
#define FIELD(class, FieldName) std::make_tuple(#FieldName, &class ::FieldName)
|
||||
|
||||
template <typename Fn, typename Tuple, typename Field, std::size_t... Is>
|
||||
void ForEachTuple(Tuple&& tuple, Field&& fields, Fn&& fn, index_sequence<Is...>)
|
||||
{
|
||||
(void)std::initializer_list<size_t>{
|
||||
(fn(std::get<0>(std::get<Is>(fields)), tuple.*std::get<1>(std::get<Is>(fields))), Is)...};
|
||||
}
|
||||
|
||||
template <typename Fn, typename Tuple>
|
||||
void ForEachTuple(Tuple&& tuple, Fn&& fn)
|
||||
{
|
||||
const auto fields = StructInfo<decay_t<Tuple>>::Info();
|
||||
ForEachTuple(
|
||||
std::forward<Tuple>(tuple), fields, std::forward<Fn>(fn),
|
||||
make_index_sequence<std::tuple_size<decltype(fields)>::value>{});
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
struct is_optional : std::false_type {
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
struct is_optional<std::unique_ptr<T>> : std::true_type {
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
bool is_optional_v()
|
||||
{
|
||||
return is_optional<decay_t<T>>::value;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
decltype(std::begin(T()), std::true_type{}) containable(size_t);
|
||||
|
||||
template <typename T>
|
||||
std::false_type containable(...);
|
||||
|
||||
template <typename T>
|
||||
using is_containable = decltype(containable<T>(0U));
|
||||
|
||||
template <typename T>
|
||||
constexpr bool IsSerializeType()
|
||||
{
|
||||
return ((!std::is_class<decay_t<T>>::value) || is_containable<decay_t<T>>());
|
||||
}
|
||||
|
||||
template <typename T, typename Fn>
|
||||
void ForEachField(T&& value, Fn&& fn)
|
||||
{
|
||||
ForEachTuple(std::forward<T>(value), std::forward<Fn>(fn));
|
||||
}
|
||||
|
||||
template <typename Fn>
|
||||
struct DumpFunctor;
|
||||
|
||||
template <typename T, typename Js, enable_if_t<!IsSerializeType<T>()>* = nullptr>
|
||||
void DumpObj(T&& obj, const std::string& field_name, Js& j)
|
||||
{
|
||||
if (field_name.empty()) {
|
||||
ForEachField(std::forward<T>(obj), DumpFunctor<Js>(j));
|
||||
return;
|
||||
}
|
||||
ForEachField(std::forward<T>(obj), DumpFunctor<Js>(j[field_name]));
|
||||
}
|
||||
|
||||
template <typename T, typename Js, enable_if_t<IsSerializeType<T>()>* = nullptr>
|
||||
void DumpObj(T&& obj, const std::string& field_name, Js& j)
|
||||
{
|
||||
if (field_name.empty()) {
|
||||
return;
|
||||
}
|
||||
j[field_name] = std::forward<T>(obj);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
struct DumpFunctor {
|
||||
explicit DumpFunctor(T& j) : js(j)
|
||||
{}
|
||||
template <typename Name, typename Field>
|
||||
void operator()(Name&& name, Field&& field) const
|
||||
{
|
||||
DumpObj(std::forward<Field>(field), std::forward<Name>(name), js);
|
||||
}
|
||||
T& js;
|
||||
};
|
||||
|
||||
template <typename Fn>
|
||||
struct FromJsonFunctor;
|
||||
|
||||
template <typename T, typename Js, enable_if_t<!IsSerializeType<T>()>* = nullptr>
|
||||
void FromJsonImpl(T&& obj, const std::string& field_name, const Js& j)
|
||||
{
|
||||
if (field_name.empty()) {
|
||||
ForEachField(std::forward<T>(obj), FromJsonFunctor<Js>(j));
|
||||
return;
|
||||
}
|
||||
if (j.find(field_name) == j.cend()) {
|
||||
return;
|
||||
}
|
||||
ForEachField(std::forward<T>(obj), FromJsonFunctor<Js>(j[field_name]));
|
||||
}
|
||||
|
||||
template <typename T, typename Js, enable_if_t<IsSerializeType<T>()>* = nullptr>
|
||||
void FromJsonImpl(T&& obj, const std::string& field_name, const Js& j)
|
||||
{
|
||||
// ignore missing field of optional
|
||||
if ((tuningtiling::is_optional_v<decltype(obj)>()) || (j.find(field_name) == j.cend())) {
|
||||
return;
|
||||
}
|
||||
j.at(field_name).get_to(std::forward<T>(obj));
|
||||
}
|
||||
|
||||
template <typename Js>
|
||||
struct FromJsonFunctor {
|
||||
explicit FromJsonFunctor(const Js& j) : js(j)
|
||||
{}
|
||||
template <typename Name, typename Field>
|
||||
void operator()(Name&& name, Field&& field) const
|
||||
{
|
||||
FromJsonImpl(std::forward<Field>(field), std::forward<Name>(name), js);
|
||||
}
|
||||
const Js& js;
|
||||
};
|
||||
} // namespace tuningtiling
|
||||
#endif
|
||||
111
csrc/common/stub/op_tiling/register/tuning_tiling_registry.h
Normal file
111
csrc/common/stub/op_tiling/register/tuning_tiling_registry.h
Normal file
@@ -0,0 +1,111 @@
|
||||
/**
|
||||
* 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.
|
||||
*/
|
||||
|
||||
#ifndef __INC_REGISTER_TUNING_TILING_REGISTRY_HEADER__
|
||||
#define __INC_REGISTER_TUNING_TILING_REGISTRY_HEADER__
|
||||
#include <vector>
|
||||
#include <map>
|
||||
#include <memory>
|
||||
#include <nlohmann/json.hpp>
|
||||
#include "graph/ascend_string.h"
|
||||
#include "register/tuning_tiling_reflection_utils.h"
|
||||
namespace tuningtiling {
|
||||
struct TilingItem {
|
||||
ge::AscendString dtype_;
|
||||
ge::AscendString name_;
|
||||
};
|
||||
|
||||
class TuningTilingDef
|
||||
{
|
||||
public:
|
||||
virtual void FromJson(const nlohmann::json& j) = 0;
|
||||
virtual void ToJson(nlohmann::json& j) = 0;
|
||||
ge::AscendString GetClassName() const;
|
||||
virtual std::vector<TilingItem> GetItemInfo() const = 0;
|
||||
|
||||
protected:
|
||||
TuningTilingDef() = default;
|
||||
virtual ~TuningTilingDef() = default;
|
||||
// dtype , name
|
||||
std::vector<TilingItem> field_info_;
|
||||
ge::AscendString class_name_;
|
||||
};
|
||||
|
||||
#define BEGIN_TUNING_TILING_DEF(class_name) \
|
||||
class class_name : public TuningTilingDef \
|
||||
{ \
|
||||
public: \
|
||||
virtual void FromJson(const nlohmann::json& j) \
|
||||
{ \
|
||||
FromJsonImpl(*this, "", j); \
|
||||
} \
|
||||
\
|
||||
virtual void ToJson(nlohmann::json& j) \
|
||||
{ \
|
||||
DumpObj(*this, "", j); \
|
||||
} \
|
||||
\
|
||||
std::vector<TilingItem> GetItemInfo() const \
|
||||
{ \
|
||||
return field_info_; \
|
||||
} \
|
||||
\
|
||||
class FieldHandler \
|
||||
{ \
|
||||
public: \
|
||||
FieldHandler(class_name* pinstance, const ge::AscendString& dtype, const ge::AscendString& name) \
|
||||
{ \
|
||||
pinstance->field_info_.push_back({dtype, name}); \
|
||||
} \
|
||||
}; \
|
||||
friend class FieldHandler; \
|
||||
\
|
||||
public: \
|
||||
class_name() \
|
||||
{ \
|
||||
class_name_ = #class_name; \
|
||||
};
|
||||
|
||||
#define TUNING_TILING_DATA_FIELD_DEF(data_type, field_name) \
|
||||
public: \
|
||||
data_type field_name; \
|
||||
FieldHandler field_name##_handler_ = FieldHandler(this, #data_type, #field_name);
|
||||
|
||||
#define END_TUNING_TILING_DEF \
|
||||
} \
|
||||
;
|
||||
|
||||
using TuningTilingDefConstructor = std::shared_ptr<TuningTilingDef> (*)();
|
||||
class TuningTilingClassFactory
|
||||
{
|
||||
public:
|
||||
static std::map<ge::AscendString, TuningTilingDefConstructor>& RegisterInfo();
|
||||
static void RegisterTilingData(const ge::AscendString& optype, TuningTilingDefConstructor const constructor);
|
||||
static std::shared_ptr<TuningTilingDef> CreateTilingDataInstance(const ge::AscendString& optype);
|
||||
};
|
||||
|
||||
#define REGISTER_TUNING_TILING_CLASS(optype, class_name) \
|
||||
class optype##Helper \
|
||||
{ \
|
||||
public: \
|
||||
optype##Helper() \
|
||||
{ \
|
||||
TuningTilingClassFactory::RegisterTilingData(#optype, optype##Helper::CreateTilingDataInstance); \
|
||||
} \
|
||||
static std::shared_ptr<TuningTilingDef> CreateTilingDataInstance() \
|
||||
{ \
|
||||
return std::make_shared<class_name>(); \
|
||||
} \
|
||||
}; \
|
||||
optype##Helper g_tuning_tiling_##optype##Helper;
|
||||
using TuningTilingDefPtr = std::shared_ptr<TuningTilingDef>;
|
||||
} // namespace tuningtiling
|
||||
|
||||
#endif
|
||||
20
csrc/common/stub/op_tiling/runtime_kb_api.cpp
Normal file
20
csrc/common/stub/op_tiling/runtime_kb_api.cpp
Normal file
@@ -0,0 +1,20 @@
|
||||
/**
|
||||
* 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.
|
||||
*/
|
||||
|
||||
#include "runtime_kb_api.h"
|
||||
|
||||
namespace RuntimeKb {
|
||||
uint32_t QueryBank(
|
||||
const void* /*src*/, size_t /*src_len*/, const std::string& /*op_type*/, const std::string& /*soc_version*/,
|
||||
uint32_t /*core_num*/, tuningtiling::TuningTilingDefPtr& /*tiling*/)
|
||||
{
|
||||
return 0;
|
||||
}
|
||||
} // namespace RuntimeKb
|
||||
23
csrc/common/stub/op_tiling/runtime_kb_api.h
Normal file
23
csrc/common/stub/op_tiling/runtime_kb_api.h
Normal file
@@ -0,0 +1,23 @@
|
||||
/**
|
||||
* 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.
|
||||
*/
|
||||
|
||||
#ifndef RUNTIME_KB_RUNTIME_KB_API_H
|
||||
#define RUNTIME_KB_RUNTIME_KB_API_H
|
||||
|
||||
#include <string>
|
||||
#include "exe_graph/runtime/tiling_context.h"
|
||||
#include "register/tuning_tiling_registry.h"
|
||||
|
||||
namespace RuntimeKb {
|
||||
uint32_t QueryBank(const void *src, size_t src_len, const std::string &op_type, const std::string &soc_version,
|
||||
uint32_t core_num, tuningtiling::TuningTilingDefPtr &tiling);
|
||||
}
|
||||
|
||||
#endif
|
||||
39
csrc/common/stub/op_tiling/tbe_tiling_api.cpp
Normal file
39
csrc/common/stub/op_tiling/tbe_tiling_api.cpp
Normal file
@@ -0,0 +1,39 @@
|
||||
/**
|
||||
* 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.
|
||||
*/
|
||||
#include "tbe_tiling_api.h"
|
||||
|
||||
using namespace optiling;
|
||||
|
||||
namespace optiling {
|
||||
bool GetTbeTiling(const gert::TilingContext* context, Conv3dBpFilterV2RunInfo& runInfoForV2, Conv3dBackpropV2TBETilingData& tbeTilingForV2)
|
||||
{
|
||||
(void)context;
|
||||
(void)runInfoForV2;
|
||||
(void)tbeTilingForV2;
|
||||
return true;
|
||||
}
|
||||
bool GetTbeTiling(gert::TilingContext* context, Conv3dBpInputV2RunInfo& runInfoV2,
|
||||
Conv3dBackpropV2TBETilingData& tbeTilingForV2, const optiling::OpTypeV2 opType)
|
||||
{
|
||||
(void)context;
|
||||
(void)runInfoV2;
|
||||
(void)tbeTilingForV2;
|
||||
(void)opType;
|
||||
return true;
|
||||
}
|
||||
|
||||
bool GetTbeTiling(gert::TilingContext* context, Conv3dBackpropV2TBETilingData& tbeTilingForV2, const optiling::OpTypeV2 opType)
|
||||
{
|
||||
(void)context;
|
||||
(void)tbeTilingForV2;
|
||||
(void)opType;
|
||||
return true;
|
||||
}
|
||||
}
|
||||
193
csrc/common/stub/op_tiling/tbe_tiling_api.h
Normal file
193
csrc/common/stub/op_tiling/tbe_tiling_api.h
Normal file
@@ -0,0 +1,193 @@
|
||||
/**
|
||||
* 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 tbe_tiling_api.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef TBE_TILING_API_H
|
||||
#define TBE_TILING_API_H
|
||||
|
||||
#include <cstdint>
|
||||
#include <exe_graph/runtime/tiling_context.h>
|
||||
#include <tiling/platform/platform_ascendc.h>
|
||||
#include "graph/utils/type_utils.h"
|
||||
#include "platform/platform_infos_def.h"
|
||||
|
||||
namespace optiling {
|
||||
struct Conv3dBackpropV2TBETilingData {
|
||||
// L0 tiling parameters
|
||||
int32_t m_l0; // Base M dimension at L0
|
||||
int32_t k_l0; // Base K dimension at L0
|
||||
int32_t n_l0; // Base N dimension at L0
|
||||
|
||||
// L1 tiling parameters
|
||||
int32_t m_al1; // Step M dimension at L1
|
||||
int32_t n_bl1; // Step N dimension at L1
|
||||
int32_t k_al1; // Step K dimension A at L1
|
||||
int32_t k_bl1; // Step K dimension B at L1
|
||||
|
||||
// Buffer parameters
|
||||
int32_t db_l0c; // L0C buffer size
|
||||
int32_t db_al1; // AL1 buffer size
|
||||
int32_t db_bl1; // BL1 buffer size
|
||||
|
||||
// Dimension parameters
|
||||
int32_t batch_dim; // Batch dimension
|
||||
int32_t d_dim; // Depth dimension
|
||||
int32_t group_dim; // Group dimension
|
||||
int32_t m_dim; // M dimension
|
||||
int32_t n_dim; // N dimension
|
||||
int32_t k_dim; // K dimension
|
||||
};
|
||||
|
||||
struct Conv3dBpFilterV2RunInfo {
|
||||
int32_t batch;
|
||||
int32_t co; // output channels
|
||||
int32_t ci; // input channels
|
||||
int32_t cout1_g; // output channels per group
|
||||
int32_t cin1_g; // input channels per group
|
||||
int32_t dout; // output depth // codespell:ignore dout
|
||||
int32_t wo; // output width
|
||||
int32_t ho; // output height
|
||||
int32_t wi; // input width
|
||||
int32_t hi; // input height
|
||||
int32_t di; // input depth
|
||||
int32_t kw; // kernel width
|
||||
int32_t kh; // kernel height
|
||||
int32_t kd; // kernel depth
|
||||
int32_t real_g; // actual groups
|
||||
int32_t stride_w;
|
||||
int32_t stride_h;
|
||||
int32_t stride_d;
|
||||
int32_t pad_l; // left padding
|
||||
int32_t pad_r; // right padding
|
||||
int32_t pad_u; // up padding
|
||||
int32_t pad_d; // down padding
|
||||
int32_t pad_f; // front padding
|
||||
int32_t pad_b; // back padding
|
||||
int32_t dilation_w;
|
||||
int32_t dilation_h;
|
||||
int32_t dilation_d;
|
||||
int32_t ci1; // another input channels parameter
|
||||
uint64_t bl1_bound; // buffer limit 1 bound, tiling结果的衍生参数,不建议放在这里
|
||||
int32_t batch_dout_single_core; // batch*dout per core, tiling结果的衍生参数,不建议放在这里 // codespell:ignore dout
|
||||
|
||||
// Tiling parameters
|
||||
uint32_t k0;
|
||||
uint32_t m0;
|
||||
uint32_t n0;
|
||||
uint32_t hf32Flag;
|
||||
|
||||
ge::DataType a_dtype = ge::DT_FLOAT16;
|
||||
ge::DataType b_dtype = ge::DT_FLOAT16;
|
||||
ge::DataType c_dtype = ge::DT_FLOAT16;
|
||||
int32_t a_dtype_bytes = 2;
|
||||
int32_t b_dtype_bytes = 2;
|
||||
int32_t c_dtype_bytes = 2;
|
||||
uint32_t core_num;
|
||||
};
|
||||
|
||||
struct Conv3dBpInputV2RunInfo {
|
||||
// Batch and group related
|
||||
int32_t batch_n; // Batch size
|
||||
int32_t real_g; // Number of groups
|
||||
|
||||
// Input dimensions (dedx)
|
||||
int32_t dedx_d; // Input depth
|
||||
int32_t dedx_cin; // Input channels per group
|
||||
int32_t dedx_cin1; // Input channels per group
|
||||
int32_t dedx_cin1_g; // Input channels per group (grouped)
|
||||
int32_t dedx_h; // Input height
|
||||
int32_t dedx_w; // Input width
|
||||
|
||||
// Output dimensions (dedy)
|
||||
int32_t dedy_d; // Output depth
|
||||
int32_t dedy_cout; // Output channels per group
|
||||
int32_t dedy_cout1; // Output channels per group
|
||||
int32_t dedy_cout1_g; // Output channels per group (grouped)
|
||||
int32_t dedy_h; // Output height
|
||||
int32_t dedy_w; // Output width
|
||||
|
||||
// Kernel dimensions
|
||||
int32_t kernel_d; // Kernel depth
|
||||
int32_t kernel_h; // Kernel height
|
||||
int32_t kernel_w; // Kernel width
|
||||
|
||||
// Strides
|
||||
int32_t stride_d; // Stride depth
|
||||
int32_t stride_h; // Stride height
|
||||
int32_t stride_w; // Stride width
|
||||
|
||||
// Padding
|
||||
int32_t pad_h; // Padding height
|
||||
int32_t pad_t; // Padding top
|
||||
int32_t pad_u; // Padding up
|
||||
int32_t pad_d; // Padding down
|
||||
int32_t pad_l; // Padding left
|
||||
int32_t pad_r; // Padding right
|
||||
|
||||
// Dilation
|
||||
int32_t dilation_d; // Dilation depth
|
||||
int32_t dilation_h; // Dilation height
|
||||
int32_t dilation_w; // Dilation width
|
||||
|
||||
// Backprop padding
|
||||
int32_t backprop_pad_h; // Backprop padding height
|
||||
int32_t backprop_pad_t; // Backprop padding top
|
||||
int32_t backprop_pad_u; // Backprop padding up
|
||||
int32_t backprop_pad_d; // Backprop padding down
|
||||
int32_t backprop_pad_l; // Backprop padding left
|
||||
int32_t backprop_pad_r; // Backprop padding right
|
||||
|
||||
// Other flags
|
||||
int32_t hf32_flag; // Flag for FP32 handling
|
||||
int32_t a_dtype_bytes = 2;
|
||||
int32_t b_dtype_bytes = 2;
|
||||
int32_t c_dtype_bytes = 2;
|
||||
int32_t initOutputFlag = 0;
|
||||
};
|
||||
|
||||
struct Conv3DBackpropV2CompileInfo {
|
||||
std::string soc_version = "";
|
||||
platform_ascendc::SocVersion shortSocVersion = platform_ascendc::SocVersion::ASCEND910B;
|
||||
|
||||
uint32_t core_num = 0;
|
||||
uint64_t ub_size = 0;
|
||||
uint64_t l1_size = 0;
|
||||
uint64_t l2_size = 0;
|
||||
uint64_t l0a_size = 0;
|
||||
uint64_t l0b_size = 0;
|
||||
uint64_t l0c_size = 0;
|
||||
uint64_t bt_size = 0;
|
||||
int32_t cube_freq = 0;
|
||||
bool load3d_constraints = true;
|
||||
bool intrinsic_data_move_l12ub = true;
|
||||
bool intrinsic_matmul_ub_to_ub = false;
|
||||
bool intrinsic_conv_ub_to_ub = false;
|
||||
bool intrinsic_data_move_l0c2ub = true;
|
||||
bool intrinsic_fix_pipe_l0c2out = false;
|
||||
bool intrinsic_fix_pipe_l0c2ub = false;
|
||||
bool intrinsic_data_move_out2l1_nd2nz = false;
|
||||
bool intrinsic_data_move_l12bt_bf16 = false;
|
||||
};
|
||||
|
||||
enum OpTypeV2 : size_t {
|
||||
kConv3DBackpropFilterV2,
|
||||
kConv3DBackpropInputV2,
|
||||
kConv3DTransposeV2,
|
||||
};
|
||||
|
||||
bool GetTbeTiling(const gert::TilingContext* context, Conv3dBpFilterV2RunInfo& runInfoForV2, Conv3dBackpropV2TBETilingData& tbeTilingForV2);
|
||||
bool GetTbeTiling(gert::TilingContext* context, Conv3dBpInputV2RunInfo& runInfoV2,
|
||||
Conv3dBackpropV2TBETilingData& tbeTilingForV2, const optiling::OpTypeV2 opType);
|
||||
bool GetTbeTiling(gert::TilingContext* context, Conv3dBackpropV2TBETilingData& tbeTilingForV2, const optiling::OpTypeV2 opType);
|
||||
}
|
||||
#endif // TBE_TILING_API_H
|
||||
Reference in New Issue
Block a user