ADD file via upload
This commit is contained in:
parent
9dc45c6663
commit
404bf3d4af
|
|
@ -0,0 +1,94 @@
|
|||
/**
|
||||
* 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
|
||||
Loading…
Reference in New Issue