Files
enginex-ascend-910-vllm/csrc/common/stub/op_tiling/tbe_tiling_api.h
Sun Ruoxi 7f8a1b1f7a init v0.23.0
Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
2026-08-27 15:11:51 +08:00

194 lines
7.1 KiB
C++
Raw Blame History

This file contains invisible Unicode characters

This file contains invisible Unicode characters that are indistinguishable to humans but may be processed differently by a computer. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

/**
 * 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