transform/split_combination_ops_decla...

95 lines
5.1 KiB
C++
Raw Permalink Blame History

This file contains ambiguous Unicode characters

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 2019-2021 Huawei Technologies Co., Ltd
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
#include "transform/graph_ir/op_declare/split_combination_ops_declare.h"
#include <vector>
namespace mindspore::transform {
// SplitD
INPUT_MAP(SplitD) = {{1, INPUT_DESC(x)}};
//输入映射将输入索引1映射为名为x的输入描述(INPUT_DESC)
ATTR_MAP(SplitD) = {{"axis", ATTR_DESC(split_dim, AnyTraits<int64_t>())},//指定维度
{"output_num", ATTR_DESC(num_split, AnyTraits<int64_t>())}};//指定输出数量
// 属性映射
DYN_OUTPUT_MAP(SplitD) = {{0, DYN_OUTPUT_DESC(y)}};
//动态输出映射 //将输出索引0映射为名为 "y" 的动态输出描述(DYN_OUTPUT_DESC)
REG_ADPT_DESC(SplitD, kNameSplitD, ADPT_DESC(SplitD))
//注册 "SplitD" 运算符的适配器描述(REG_ADPT_DESC),适配器描述中包含运算符名称 "kNameSplitD" 和适配器描述(ADPT_DESC)。
// 这将把 "SplitD" 运算符与其在框架中的实现关联起来。
//
// Pack
INPUT_MAP(Pack) = EMPTY_INPUT_MAP;
//定义 "Pack" 运算符的输入映射(INPUT_MAP)为空。这表示 "Pack" 运算符没有显式的输入,因此没有任何输入描述。
DYN_INPUT_MAP(Pack) = {{1, DYN_INPUT_DESC(x)}};
//定义 "Pack" 运算符的动态输入映射(DYN_INPUT_MAP)将输入索引1映射为名为 "x" 的动态输入描述(DYN_INPUT_DESC)
ATTR_MAP(Pack) = {{"num", ATTR_DESC(N, AnyTraits<int64_t>())}, {"axis", ATTR_DESC(axis, AnyTraits<int64_t>())}};
//定义 "Pack" 运算符的属性映射(ATTR_MAP) 两个属性:数量/维度
OUTPUT_MAP(Pack) = {{0, OUTPUT_DESC(y)}};//输出映射(OUTPUT_MAP)将输出索引0映射为名为 "y" 的输出描述(OUTPUT_DESC)
REG_ADPT_DESC(Pack, prim::kStack, ADPT_DESC(Pack))//注册 "Pack" 运算符的适配器描述(REG_ADPT_DESC)
// ParallelConcat
INPUT_MAP(ParallelConcat) = EMPTY_INPUT_MAP;
//定义 "ParallelConcat" 运算符的输入映射(INPUT_MAP)为空。
//这表示 "ParallelConcat" 运算符没有显式的输入,因此没有任何输入描述
DYN_INPUT_MAP(ParallelConcat) = {{1, DYN_INPUT_DESC(values)}};
//定义 "ParallelConcat" 运算符的动态输入映射(DYN_INPUT_MAP)将输入索引1映射为名为 "values" 的动态输入描述(DYN_INPUT_DESC)。
//这表示 "ParallelConcat" 运算符会接收一个动态数量的输入,而每个输入都可以使用 "values" 来标识
ATTR_MAP(ParallelConcat) = {//属性映射(ATTR_MAP),包含两个属性
{"shape", ATTR_DESC(shape, AnyTraits<std::vector<int64_t>>())},//连接操作时的形状
{"N", ATTR_DESC(N, AnyTraits<int64_t>())},//连接的数量
};
OUTPUT_MAP(ParallelConcat) = {{0, OUTPUT_DESC(output_data)}};
//输出映射(OUTPUT_MAP)
REG_ADPT_DESC(ParallelConcat, kNameParallelConcat, ADPT_DESC(ParallelConcat))
//注册 "ParallelConcat" 运算符的适配器描述(REG_ADPT_DESC)。
//适配器描述中包含运算符名称 "kNameParallelConcat" 和适配器描述(ADPT_DESC)。
//这将把 "ParallelConcat" 运算符与其在框架中的实现关联起来,并使用 "kNameParallelConcat" 作为标识来调用适配器。
// ConcatD
INPUT_MAP(ConcatD) = EMPTY_INPUT_MAP;
//定义 "ConcatD" 运算符的输入映射(INPUT_MAP)为空。
//这表示 "ConcatD" 运算符没有显式的输入,因此没有任何输入描述
DYN_INPUT_MAP(ConcatD) = {{1, DYN_INPUT_DESC(x)}};
//定义 "ConcatD" 运算符的动态输入映射(DYN_INPUT_MAP)将输入索引1映射为名为 "x" 的动态输入描述(DYN_INPUT_DESC)。
//这表示 "ConcatD" 运算符会接收一个动态数量的输入,而每个输入都可以使用 "x" 来标识。
ATTR_MAP(ConcatD) = {//属性映射
{"axis", ATTR_DESC(concat_dim, AnyTraits<int64_t>())},//连接操作时的维度
{"inputNums", ATTR_DESC(N, AnyTraits<int64_t>())},//输入数量
};
OUTPUT_MAP(ConcatD) = {{0, OUTPUT_DESC(y)}};
//定义 "ConcatD" 运算符的输出映射(OUTPUT_MAP)将输出索引0映射为名为 "y" 的输出描述(OUTPUT_DESC)。
//这表示 "ConcatD" 运算符会产生一个输出结果,并使用 "y" 来标识该输出。
REG_ADPT_DESC(ConcatD, prim::kPrimConcat->name(), ADPT_DESC(ConcatD))
//注册 "ConcatD" 运算符的适配器描述(REG_ADPT_DESC)。适配器描述中包含运算符名称 "prim::kPrimConcat->name()" 和适配器描述(ADPT_DESC)。
// 这将把 "ConcatD" 运算符与其在框架中的实现关联起来,并使用 "prim::kPrimConcat->name()" 作为标识来调用适配器。
//
// ConcatV2D Inference for tf
INPUT_MAP(ConcatV2D) = EMPTY_INPUT_MAP;//输入映射为空
DYN_INPUT_MAP(ConcatV2D) = {{1, DYN_INPUT_DESC(x)}};//动态输入映射,可以有多个数量的映射
ATTR_MAP(ConcatV2D) = {
{"axis", ATTR_DESC(concat_dim, AnyTraits<int64_t>())},//连接操作时的维度
{"N", ATTR_DESC(N, AnyTraits<int64_t>())},//连接的数量
};
OUTPUT_MAP(ConcatV2D) = {{0, OUTPUT_DESC(y)}};
//定义 "ConcatV2D" 运算符的输出映射(OUTPUT_MAP)将输出索引0映射为名为 "y" 的输出描述(OUTPUT_DESC)。
//这表示 "ConcatV2D" 运算符会产生一个输出结果,并使用 "y" 来标识该输出。
REG_ADPT_DESC(ConcatV2D, kNameConcatV2D, ADPT_DESC(ConcatV2D))
//注册 "ConcatV2D" 运算符的适配器描述(REG_ADPT_DESC)。适配器描述中包含运算符名称 "kNameConcatV2D" 和适配器描述(ADPT_DESC)。
//这将把 "ConcatV2D" 运算符与其在框架中的实现关联起来,并使用 "kNameConcatV2D" 作为标识来调用适配器。
} // namespace mindspore::transform