已合并
修复引入mhc算子后与libopapi.so中的aclnn函数产生冲突的问题 #6244
修复引入mhc算子后与libopapi.so中的aclnn函数产生冲突的问题 #6244
已合并
AmadeusAlex创建于 6月1日
共 2 个文件变更+19-18
@@ -329,6 +329,7 @@ add_subdirectory(ffn)
329list(APPEND OP_LIST "ffn")329list(APPEND OP_LIST "ffn")
330list(APPEND OP_DIR_LIST ${CMAKE_CURRENT_SOURCE_DIR}/ffn/ffn)330list(APPEND OP_DIR_LIST ${CMAKE_CURRENT_SOURCE_DIR}/ffn/ffn)
331add_subdirectory(attention)331add_subdirectory(attention)
332+add_subdirectory(mhc)
332list(APPEND OP_LIST ${COMPILED_OPS})333list(APPEND OP_LIST ${COMPILED_OPS})
333list(REMOVE_DUPLICATES OP_LIST)334list(REMOVE_DUPLICATES OP_LIST)
334list(APPEND OP_DIR_LIST ${COMPILED_OP_DIRS})335list(APPEND OP_DIR_LIST ${COMPILED_OP_DIRS})
@@ -44,11 +44,11 @@ constexpr int64_t D_ALIGNMENT = 16;
44constexpr int64_t ALPHA_DIM_SIZE = 3;44constexpr int64_t ALPHA_DIM_SIZE = 3;
45constexpr int64_t PHI_DIM_OFFSET = 2;45constexpr int64_t PHI_DIM_OFFSET = 2;
46 46 
47-bool CheckAlphaShape(const aclTensor *alphaTensor);47+static bool CheckAlphaShape(const aclTensor *alphaTensor);
48-bool ValidateNDParams(int64_t n, int64_t d);48+static bool ValidateNDParams(int64_t n, int64_t d);
49-bool CheckPhiShape(const aclTensor *phiTensor, int64_t n2Plus2n, int64_t nD);49+static bool CheckPhiShape(const aclTensor *phiTensor, int64_t n2Plus2n, int64_t nD);
50-bool CheckBiasShape(const aclTensor *biasTensor, int64_t n2Plus2n);50+static bool CheckBiasShape(const aclTensor *biasTensor, int64_t n2Plus2n);
51-bool CheckGammaShape(const aclTensor *gammaOptional, int64_t n, int64_t d);51+static bool CheckGammaShape(const aclTensor *gammaOptional, int64_t n, int64_t d);
52 52 
53struct MhcParamsBase {53struct MhcParamsBase {
54 const aclTensor *x = nullptr;54 const aclTensor *x = nullptr;
@@ -126,7 +126,7 @@ private:
126 MhcParamsBase obj_;126 MhcParamsBase obj_;
127};127};
128 128 
129-bool CheckNotNull(const MhcParamsBase &params)129+static bool CheckNotNull(const MhcParamsBase &params)
130{130{
131 if (params.x == nullptr) {131 if (params.x == nullptr) {
132 OP_LOGE(ACLNN_ERR_PARAM_NULLPTR, "X tensor is nullptr");132 OP_LOGE(ACLNN_ERR_PARAM_NULLPTR, "X tensor is nullptr");
@@ -159,7 +159,7 @@ bool CheckNotNull(const MhcParamsBase &params)
159 return true;159 return true;
160}160}
161 161 
162-bool CheckEmptyTensor(const MhcParamsBase &params)162+static bool CheckEmptyTensor(const MhcParamsBase &params)
163{163{
164 if (params.x->IsEmpty()) {164 if (params.x->IsEmpty()) {
165 OP_LOGE(ACLNN_ERR_PARAM_INVALID, "X tensor is empty");165 OP_LOGE(ACLNN_ERR_PARAM_INVALID, "X tensor is empty");
@@ -180,7 +180,7 @@ bool CheckEmptyTensor(const MhcParamsBase &params)
180 return true;180 return true;
181}181}
182 182 
183-bool CheckInputOutDims(const MhcParamsBase &params)183+static bool CheckInputOutDims(const MhcParamsBase &params)
184{184{
185 auto xDimNum = params.x->GetViewShape().GetDimNum();185 auto xDimNum = params.x->GetViewShape().GetDimNum();
186 if (xDimNum != DIM_NUM_3 && xDimNum != DIM_NUM_4) {186 if (xDimNum != DIM_NUM_3 && xDimNum != DIM_NUM_4) {
@@ -217,7 +217,7 @@ bool CheckInputOutDims(const MhcParamsBase &params)
217 return true;217 return true;
218}218}
219 219 
220-bool CheckInputOutShape(const MhcParamsBase &params)220+static bool CheckInputOutShape(const MhcParamsBase &params)
221{221{
222 auto xShape = params.x->GetViewShape();222 auto xShape = params.x->GetViewShape();
223 223 
@@ -260,7 +260,7 @@ bool CheckInputOutShape(const MhcParamsBase &params)
260 return true;260 return true;
261}261}
262 262 
263-bool CheckAlphaShape(const aclTensor *alphaTensor)263+static bool CheckAlphaShape(const aclTensor *alphaTensor)
264{264{
265 auto alphaShape = alphaTensor->GetViewShape();265 auto alphaShape = alphaTensor->GetViewShape();
266 if (alphaShape.GetDim(0) != ALPHA_DIM_SIZE) {266 if (alphaShape.GetDim(0) != ALPHA_DIM_SIZE) {
@@ -270,7 +270,7 @@ bool CheckAlphaShape(const aclTensor *alphaTensor)
270 return true;270 return true;
271}271}
272 272 
273-bool ValidateNDParams(int64_t n, int64_t d)273+static bool ValidateNDParams(int64_t n, int64_t d)
274{274{
275 if (n <= 0 || d <= 0) {275 if (n <= 0 || d <= 0) {
276 OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Invalid X tensor shape: n=%ld, d=%ld", n, d);276 OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Invalid X tensor shape: n=%ld, d=%ld", n, d);
@@ -297,7 +297,7 @@ bool ValidateNDParams(int64_t n, int64_t d)
297 return true;297 return true;
298}298}
299 299 
300-bool CheckPhiShape(const aclTensor *phiTensor, int64_t n2Plus2n, int64_t nD)300+static bool CheckPhiShape(const aclTensor *phiTensor, int64_t n2Plus2n, int64_t nD)
301{301{
302 auto phiShape = phiTensor->GetViewShape();302 auto phiShape = phiTensor->GetViewShape();
303 if (phiShape.GetDim(0) != n2Plus2n) {303 if (phiShape.GetDim(0) != n2Plus2n) {
@@ -312,7 +312,7 @@ bool CheckPhiShape(const aclTensor *phiTensor, int64_t n2Plus2n, int64_t nD)
312 return true;312 return true;
313}313}
314 314 
315-bool CheckBiasShape(const aclTensor *biasTensor, int64_t n2Plus2n)315+static bool CheckBiasShape(const aclTensor *biasTensor, int64_t n2Plus2n)
316{316{
317 auto biasShape = biasTensor->GetViewShape();317 auto biasShape = biasTensor->GetViewShape();
318 if (biasShape.GetDim(0) != n2Plus2n) {318 if (biasShape.GetDim(0) != n2Plus2n) {
@@ -323,7 +323,7 @@ bool CheckBiasShape(const aclTensor *biasTensor, int64_t n2Plus2n)
323 return true;323 return true;
324}324}
325 325 
326-bool CheckGammaShape(const aclTensor *gammaOptional, int64_t n, int64_t d)326+static bool CheckGammaShape(const aclTensor *gammaOptional, int64_t n, int64_t d)
327{327{
328 if (gammaOptional != nullptr) {328 if (gammaOptional != nullptr) {
329 auto gammaShape = gammaOptional->GetViewShape();329 auto gammaShape = gammaOptional->GetViewShape();
@@ -341,7 +341,7 @@ bool CheckGammaShape(const aclTensor *gammaOptional, int64_t n, int64_t d)
341 return true;341 return true;
342}342}
343 343 
344-bool CheckDtypeValid(const MhcParamsBase &params)344+static bool CheckDtypeValid(const MhcParamsBase &params)
345{345{
346 const std::initializer_list<DataType> X_SUPPORT_DTYPE_LIST = {DataType::DT_BF16, DataType::DT_FLOAT16};346 const std::initializer_list<DataType> X_SUPPORT_DTYPE_LIST = {DataType::DT_BF16, DataType::DT_FLOAT16};
347 347 
@@ -390,7 +390,7 @@ static bool IsPrivateFormat(ge::Format format)
390 return false;390 return false;
391}391}
392 392 
393-bool CheckFormat(const MhcParamsBase &params)393+static bool CheckFormat(const MhcParamsBase &params)
394{394{
395 if (IsPrivateFormat(params.x->GetViewFormat())) {395 if (IsPrivateFormat(params.x->GetViewFormat())) {
396 OP_LOGE(ACLNN_ERR_PARAM_INVALID, "X tensor format must be ND");396 OP_LOGE(ACLNN_ERR_PARAM_INVALID, "X tensor format must be ND");
@@ -422,7 +422,7 @@ bool CheckFormat(const MhcParamsBase &params)
422 return true;422 return true;
423}423}
424 424 
425-aclnnStatus CheckParams(const MhcParamsBase &params)425+static aclnnStatus CheckParams(const MhcParamsBase &params)
426{426{
427 // 1. 检查参数是否为空指针、空tensor427 // 1. 检查参数是否为空指针、空tensor
428 CHECK_RET(CheckNotNull(params), ACLNN_ERR_PARAM_NULLPTR);428 CHECK_RET(CheckNotNull(params), ACLNN_ERR_PARAM_NULLPTR);
@@ -443,7 +443,7 @@ aclnnStatus CheckParams(const MhcParamsBase &params)
443 return ACLNN_SUCCESS;443 return ACLNN_SUCCESS;
444}444}
445 445 
446-aclnnStatus ConvertDataContiguous(MhcParamsBase &params, aclOpExecutor *executor)446+static aclnnStatus ConvertDataContiguous(MhcParamsBase &params, aclOpExecutor *executor)
447{447{
448 // 将输入tensor转换为连续格式448 // 将输入tensor转换为连续格式
449 params.xContiguous = l0op::Contiguous(params.x, executor);449 params.xContiguous = l0op::Contiguous(params.x, executor);