321 lines
16 KiB
C++
321 lines
16 KiB
C++
/**
|
||
* 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 <vector>
|
||
#include "transform/graph_ir/op_declare/selection_ops_declare.h"
|
||
|
||
namespace mindspore::transform {
|
||
// CumsumD
|
||
INPUT_MAP(CumsumD) = {{1, INPUT_DESC(x)}};
|
||
//一个输入映射x
|
||
INPUT_ATTR_MAP(CumsumD) = {{2, ATTR_DESC(axis, AnyTraits<int64_t>())}};
|
||
//CumsumD操作的输入属性映射,有一个名为"axis"的属性,类型为int64_t
|
||
ATTR_MAP(CumsumD) = {{"exclusive", ATTR_DESC(exclusive, AnyTraits<bool>())},
|
||
{"reverse", ATTR_DESC(reverse, AnyTraits<bool>())}};
|
||
// CumsumD操作的属性映射,列出了"exclusive"和"reverse"两个属性,类型分别为bool
|
||
OUTPUT_MAP(CumsumD) = {{0, OUTPUT_DESC(y)}};
|
||
// CumsumD操作的输出映射,有一个输出"y",索引为0
|
||
REG_ADPT_DESC(CumsumD, kNameCumSum, ADPT_DESC(CumsumD))
|
||
// 注册CumsumD操作的适配器描述kNameCumSum
|
||
//
|
||
// GatherV2
|
||
INPUT_MAP(GatherV2) = {{1, INPUT_DESC(x)}, {2, INPUT_DESC(indices)}, {3, INPUT_DESC(axis)}};
|
||
// GatherV2操作的输入映射,有三个输入:"x"、"indices"、"axis",索引分别为1、2、3
|
||
ATTR_MAP(GatherV2) = EMPTY_ATTR_MAP;
|
||
// GatherV2操作没有属性,为空的属性映射
|
||
OUTPUT_MAP(GatherV2) = {{0, OUTPUT_DESC(y)}};
|
||
// GatherV2操作的输出映射,有一个输出"y",索引为0
|
||
//
|
||
// CumprodD
|
||
INPUT_MAP(CumprodD) = {{1, INPUT_DESC(x)}};
|
||
// CumprodD操作的输入映射,有一个输入"x",索引为1
|
||
INPUT_ATTR_MAP(CumprodD) = {{2, ATTR_DESC(axis, AnyTraits<int64_t>())}};
|
||
// CumprodD操作的输入属性映射,有一个名为"axis"的属性,类型为int64_tz
|
||
ATTR_MAP(CumprodD) = {{"exclusive", ATTR_DESC(exclusive, AnyTraits<bool>())},
|
||
{"reverse", ATTR_DESC(reverse, AnyTraits<bool>())}};
|
||
// CumprodD操作的属性映射,列出了"exclusive"和"reverse"两个属性,类型分别为bool
|
||
OUTPUT_MAP(CumprodD) = {{0, OUTPUT_DESC(y)}};
|
||
// CumprodD操作的输出映射,有一个输出"y",索引为0
|
||
REG_ADPT_DESC(CumprodD, kNameCumProd, ADPT_DESC(CumprodD))
|
||
// 注册CumprodD操作的适配器描述kNameCumProd
|
||
|
||
INPUT_MAP(SliceD) = {{1, INPUT_DESC(x)}};
|
||
// SliceD操作的输入映射,有一个输入"x",索引为1
|
||
INPUT_ATTR_MAP(SliceD) = {{2, ATTR_DESC(offsets, AnyTraits<int64_t>(), AnyTraits<std::vector<int64_t>>())},
|
||
{3, ATTR_DESC(size, AnyTraits<int64_t>(), AnyTraits<std::vector<int64_t>>())}};
|
||
// 有两个输入属性,分别是"offsets"和"size",对应的类型分别是int64_t和std::vector<int64_t>。
|
||
ATTR_MAP(SliceD) = EMPTY_ATTR_MAP;
|
||
//SliceD操作没有属性,为空的属性映射
|
||
OUTPUT_MAP(SliceD) = {{0, OUTPUT_DESC(y)}};
|
||
// SliceD操作的输出映射
|
||
REG_ADPT_DESC(SliceD, kNameSlice, ADPT_DESC(SliceD))
|
||
//注册SliceD操作的适配器描述kNameSlice
|
||
|
||
// TopK
|
||
INPUT_MAP(TopK) = {{1, INPUT_DESC(x)}, {2, INPUT_DESC(k)}};
|
||
// TopK操作的输入映射,有两个输入:"x"、"k",索引分别为1、2
|
||
ATTR_MAP(TopK) = {{"sorted", ATTR_DESC(sorted, AnyTraits<bool>())}};
|
||
//属性映射,有一个属性sort类型是bool
|
||
OUTPUT_MAP(TopK) = {{0, OUTPUT_DESC(values)}, {1, OUTPUT_DESC(indices)}};
|
||
//输入映射,有两个输出:"values"、"indices",索引分别为0、1
|
||
REG_ADPT_DESC(TopK, kNameTopK, ADPT_DESC(TopK))
|
||
// 注册TopK操作的适配器描述kNameTopK
|
||
|
||
// InTopK
|
||
INPUT_MAP(InTopKD) = {{1, INPUT_DESC(x1)}, {2, INPUT_DESC(x2)}};
|
||
// InTopKD操作的输入映射,有两个输入:"x1"、"x2",索引分别为1、2
|
||
ATTR_MAP(InTopKD) = {{"k", ATTR_DESC(k, AnyTraits<int64_t>())}};
|
||
//属性映射,有一个属性k类型是int64_t
|
||
OUTPUT_MAP(InTopKD) = {{0, OUTPUT_DESC(y)}};
|
||
//输出映射,有一个输出y,索引是0
|
||
REG_ADPT_DESC(InTopKD, kNameInTopKD, ADPT_DESC(InTopKD))
|
||
//注册InTopKD操作的适配器kNameInTopK
|
||
|
||
// TileD
|
||
INPUT_MAP(TileD) = {{1, INPUT_DESC(x)}};
|
||
//输入映射,有一个输入x,索引是1
|
||
INPUT_ATTR_MAP(TileD) = {{2, ATTR_DESC(multiples, AnyTraits<int64_t>(), AnyTraits<std::vector<int64_t>>())}};
|
||
// 输入属性multiples,对应的类型是int64_t和std::vector<int64_t>,索引是2
|
||
ATTR_MAP(TileD) = EMPTY_ATTR_MAP;
|
||
//属性映射,空
|
||
OUTPUT_MAP(TileD) = {{0, OUTPUT_DESC(y)}};
|
||
//输出映射,有一个输出y,索引时0
|
||
REG_ADPT_DESC(TileD, kNameTile, ADPT_DESC(TileD))
|
||
// 注册TileD操作的适配器kNameTile
|
||
|
||
// OneHot
|
||
INPUT_MAP(OneHot) = {{1, INPUT_DESC(x)}, {2, INPUT_DESC(depth)}, {3, INPUT_DESC(on_value)}, {4, INPUT_DESC(off_value)}};
|
||
// 输入映射,有四个输入x、depth、on_value、off_value,索引分别是1、2、3、4
|
||
ATTR_MAP(OneHot) = {{"axis", ATTR_DESC(axis, AnyTraits<int64_t>())}};
|
||
//属性映射,一个属性axis,类型为int64_t
|
||
OUTPUT_MAP(OneHot) = {{0, OUTPUT_DESC(y)}};
|
||
//输出映射,一个输出y,索引为0
|
||
REG_ADPT_DESC(OneHot, prim::kPrimOneHot->name(), ADPT_DESC(OneHot))
|
||
//将名为 "OneHot" 的操作与适配器描述进行注册。
|
||
//将"OneHot" 的名称(prim::kPrimOneHot->name())和适配器描述(ADPT_DESC(OneHot))关联起来
|
||
|
||
// GatherV2D
|
||
INPUT_MAP(GatherV2D) = {{1, INPUT_DESC(x)}, {2, INPUT_DESC(indices)}};
|
||
//输入映射,有两个输入x、indices,索引分别为1、2
|
||
INPUT_ATTR_MAP(GatherV2D) = {{3, ATTR_DESC(axis, AnyTraits<int64_t>())}};
|
||
//输入属性映射,int64_t型的axis,索引为3
|
||
ATTR_MAP(GatherV2D) = EMPTY_ATTR_MAP;
|
||
//属性映射,空
|
||
OUTPUT_MAP(GatherV2D) = {{0, OUTPUT_DESC(y)}};
|
||
//有一个输出映射y,索引为0
|
||
REG_ADPT_DESC(GatherV2D, prim::kPrimGather->name(), ADPT_DESC(GatherV2D))
|
||
//这行代码是将名为 "GatherV2D" 的操作与适配器描述进行注册。
|
||
//它将 "GatherV2D" 的名称(prim::kPrimGather->name())和适配器描述(ADPT_DESC(GatherV2D))关联起来,以便在特定计算框架或引擎中能够正确地执行和优化 "GatherV2D" 操作。
|
||
REG_ADPT_DESC(Gather, kNameGather, ADPT_DESC(GatherV2D))
|
||
//将名为 "Gather" 的操作与适配器描述进行注册。
|
||
//将 "Gather" 的名称(kNameGather)和适配器描述(ADPT_DESC(GatherV2D))关联起来,以便在特定计算框架或引擎中能够正确地执行和优化 "Gather" 操作。
|
||
|
||
// ScatterNdD
|
||
INPUT_MAP(ScatterNdD) = {{1, INPUT_DESC(indices)}, {2, INPUT_DESC(x)}};
|
||
//输入映射,有两个,indices和x,索引为1、2
|
||
INPUT_ATTR_MAP(ScatterNdD) = {
|
||
{3, ATTR_DESC(shape, AnyTraits<std::vector<int64_t>>(), AnyTraits<std::vector<int64_t>>())}};
|
||
//输入属性映射,int64_t型的shape,索引为3
|
||
ATTR_MAP(ScatterNdD) = EMPTY_ATTR_MAP;
|
||
//属性映射为空
|
||
OUTPUT_MAP(ScatterNdD) = {{0, OUTPUT_DESC(y)}};
|
||
//输出映射y,索引为0
|
||
REG_ADPT_DESC(ScatterNdD, kNameScatterNdD, ADPT_DESC(ScatterNdD))
|
||
// 将名为 "ScatterNdD" 的操作与适配器描述进行注册。
|
||
// 将 "ScatterNdD" 的名称(kNameScatterNdD)和适配器描述(ADPT_DESC(ScatterNdD))关联起来,以便在特定计算框架或引擎中能够正确地执行和优化 "ScatterNdD" 操作。
|
||
//
|
||
// ScatterNonAliasingAdd
|
||
INPUT_MAP(ScatterNonAliasingAdd) = {{1, INPUT_DESC(x)}, {2, INPUT_DESC(indices)}, {3, INPUT_DESC(updates)}};
|
||
//输入映射,共三个,x索引为1,indices索引为2,updates索引为3
|
||
ATTR_MAP(ScatterNonAliasingAdd) = EMPTY_ATTR_MAP;
|
||
//属性映射,空
|
||
OUTPUT_MAP(ScatterNonAliasingAdd) = {{0, OUTPUT_DESC(y)}};
|
||
//输出映射,y索引为0
|
||
REG_ADPT_DESC(ScatterNonAliasingAdd, kNameScatterNonAliasingAdd, ADPT_DESC(ScatterNonAliasingAdd))
|
||
// 注册ScatterNonAliasingAdd操作的适配器描述kNameScatterNonAliasingAdd
|
||
//
|
||
// GatherNd
|
||
INPUT_MAP(GatherNd) = {{1, INPUT_DESC(x)}, {2, INPUT_DESC(indices)}};
|
||
//输入映射,共两个,x索引为1,indices索引为2
|
||
ATTR_MAP(GatherNd) = EMPTY_ATTR_MAP;
|
||
//属性映射,空
|
||
OUTPUT_MAP(GatherNd) = {{0, OUTPUT_DESC(y)}};
|
||
//输出映射,y索引为0
|
||
REG_ADPT_DESC(GatherNd, kNameGatherNd, ADPT_DESC(GatherNd))
|
||
// 注册GatherNd操作的适配器描述 kNameGatherNd
|
||
|
||
// GatherD
|
||
INPUT_MAP(GatherD) = {{1, INPUT_DESC(x)}, {2, INPUT_DESC(dim)}, {3, INPUT_DESC(index)}};
|
||
//输入映射,共三个,x索引为1,dim索引为2,index索引为3
|
||
ATTR_MAP(GatherD) = EMPTY_ATTR_MAP;
|
||
//属性映射,空
|
||
OUTPUT_MAP(GatherD) = {{0, OUTPUT_DESC(y)}};
|
||
//输出映射,y索引为0
|
||
REG_ADPT_DESC(GatherD, kNameGatherD, ADPT_DESC(GatherD))
|
||
// 注册kNameGatherD操作的适配器描述 GatherD
|
||
|
||
// Range
|
||
INPUT_MAP(RangeD) = {{1, INPUT_DESC(x)}};
|
||
//输入映射,x索引为1
|
||
ATTR_MAP(RangeD) = {{"start", ATTR_DESC(start, AnyTraits<float>())},
|
||
{"limit", ATTR_DESC(limit, AnyTraits<float>())},
|
||
{"delta", ATTR_DESC(delta, AnyTraits<float>())}};
|
||
//属性映射,列出了"start"、"limit"、"delta"三个属性,类型分别为float
|
||
REG_ADPT_DESC(RangeD, kNameRange, ADPT_DESC(RangeD))
|
||
//注册RangeD操作的适配器描述 kNameRange
|
||
|
||
// InplaceAddD
|
||
INPUT_MAP(InplaceAddD) = {{1, INPUT_DESC(x)}, {2, INPUT_DESC(v)}};
|
||
// 输入映射,x索引为1,y索引为2
|
||
ATTR_MAP(InplaceAddD) = {{"indices", ATTR_DESC(indices, AnyTraits<std::vector<int64_t>>())}};
|
||
//属性映射,属性indices类型为int64_t
|
||
OUTPUT_MAP(InplaceAddD) = {{0, OUTPUT_DESC(y)}};
|
||
//输出映射,y索引为0
|
||
REG_ADPT_DESC(InplaceAddD, kNameInplaceAddD, ADPT_DESC(InplaceAddD))
|
||
// 注册RangeD操作的适配器描述 kNameRange
|
||
|
||
// InplaceSubD
|
||
INPUT_MAP(InplaceSubD) = {{1, INPUT_DESC(x)}, {2, INPUT_DESC(v)}};
|
||
//输入映射,共两个,x索引为1,y索引为2
|
||
ATTR_MAP(InplaceSubD) = {{"indices", ATTR_DESC(indices, AnyTraits<std::vector<int64_t>>())}};
|
||
// 属性映射,属性indices类型为int64_t
|
||
OUTPUT_MAP(InplaceSubD) = {{0, OUTPUT_DESC(y)}};
|
||
// 输出映射,y索引为0
|
||
REG_ADPT_DESC(InplaceSubD, kNameInplaceSubD, ADPT_DESC(InplaceSubD))
|
||
// 注册InplaceSubD操作的适配器描述kNameInplaceSubD
|
||
|
||
// InplaceUpdateD
|
||
INPUT_MAP(InplaceUpdateD) = {{1, INPUT_DESC(x)}, {2, INPUT_DESC(v)}};
|
||
// 输入映射,共两个,x索引为1,y索引为2
|
||
ATTR_MAP(InplaceUpdateD) = {{"indices", ATTR_DESC(indices, AnyTraits<std::vector<int64_t>>())}};
|
||
// 属性映射,属性indices类型为int64_t
|
||
OUTPUT_MAP(InplaceUpdateD) = {{0, OUTPUT_DESC(y)}};
|
||
// 输出映射,y索引为0
|
||
REG_ADPT_DESC(InplaceUpdateD, kNameInplaceUpdateD, ADPT_DESC(InplaceUpdateD))
|
||
// 注册InplaceUpdateD操作的适配器描述kNameInplaceUpdateD
|
||
|
||
// Select
|
||
INPUT_MAP(Select) = {{1, INPUT_DESC(condition)}, {2, INPUT_DESC(x1)}, {3, INPUT_DESC(x2)}};
|
||
// 输入映射,共三个,condition索引为1,x1索引为2,x2索引为3
|
||
ATTR_MAP(Select) = EMPTY_ATTR_MAP;
|
||
//属性映射,空
|
||
OUTPUT_MAP(Select) = {{0, OUTPUT_DESC(y)}};
|
||
//输出映射,y索引为0
|
||
REG_ADPT_DESC(Select, prim::kPrimSelect->name(), ADPT_DESC(Select))
|
||
// 注册InplaceUpdateD操作的适配器描述prim::kPrimSelect->name()
|
||
|
||
// StridedSliceGrad
|
||
INPUT_MAP(StridedSliceGrad) = {
|
||
{1, INPUT_DESC(dy)}, {2, INPUT_DESC(shape)}, {3, INPUT_DESC(begin)}, {4, INPUT_DESC(end)}, {5, INPUT_DESC(strides)}};
|
||
// 输入映射,共五个,dy索引为1,shape索引为2,begin索引为3,end索引为4,strides索引为5
|
||
ATTR_MAP(StridedSliceGrad) = {{"begin_mask", ATTR_DESC(begin_mask, AnyTraits<int64_t>())},
|
||
{"end_mask", ATTR_DESC(end_mask, AnyTraits<int64_t>())},
|
||
{"ellipsis_mask", ATTR_DESC(ellipsis_mask, AnyTraits<int64_t>())},
|
||
{"new_axis_mask", ATTR_DESC(new_axis_mask, AnyTraits<int64_t>())},
|
||
{"shrink_axis_mask", ATTR_DESC(shrink_axis_mask, AnyTraits<int64_t>())}};
|
||
// 属性映射,列出了"begin_mask"、"end_mask"、"ellipsis_mask"、"new_axis_mask"、"shrink_axis_mask"五个属性,类型分别为int64_t
|
||
OUTPUT_MAP(StridedSliceGrad) = {{0, OUTPUT_DESC(output)}};
|
||
// 输出映射,output索引为0
|
||
REG_ADPT_DESC(StridedSliceGrad, kNameStridedSliceGrad, ADPT_DESC(StridedSliceGrad))
|
||
//注册StridedSliceGradD操作的适配器描述kNameStridedSliceGrad
|
||
|
||
// StridedSlice
|
||
INPUT_MAP(StridedSlice) = {{1, INPUT_DESC(x)}, {2, INPUT_DESC(begin)}, {3, INPUT_DESC(end)}, {4, INPUT_DESC(strides)}};
|
||
// 输入映射,共四个,x索引为1,begin索引为2,end索引为3,strides索引为4
|
||
ATTR_MAP(StridedSlice) = {{"begin_mask", ATTR_DESC(begin_mask, AnyTraits<int64_t>())},
|
||
{"end_mask", ATTR_DESC(end_mask, AnyTraits<int64_t>())},
|
||
{"ellipsis_mask", ATTR_DESC(ellipsis_mask, AnyTraits<int64_t>())},
|
||
{"new_axis_mask", ATTR_DESC(new_axis_mask, AnyTraits<int64_t>())},
|
||
{"shrink_axis_mask", ATTR_DESC(shrink_axis_mask, AnyTraits<int64_t>())}};
|
||
// 属性映射,列出了"begin_mask"、"end_mask"、"ellipsis_mask"、"new_axis_mask"、"shrink_axis_mask"五个属性,类型分别为int64_t
|
||
OUTPUT_MAP(StridedSlice) = {{0, OUTPUT_DESC(y)}};
|
||
// 输出映射,y索引为0
|
||
REG_ADPT_DESC(StridedSlice, kNameStridedSlice, ADPT_DESC(StridedSlice))
|
||
// 注册StridedSlice操作的适配器描述kNameStridedSlice
|
||
|
||
// StridedSliceV2
|
||
INPUT_MAP(StridedSliceV2) = {
|
||
{1, INPUT_DESC(x)}, {2, INPUT_DESC(begin)}, {3, INPUT_DESC(end)}, {4, INPUT_DESC(axes)}, {5, INPUT_DESC(strides)}};
|
||
// 输入映射,共五个,x索引为1,begin索引为2,end索引为3,axes索引为4,strides索引为5
|
||
ATTR_MAP(StridedSliceV2) = {{"begin_mask", ATTR_DESC(begin_mask, AnyTraits<int64_t>())},
|
||
{"end_mask", ATTR_DESC(end_mask, AnyTraits<int64_t>())},
|
||
{"ellipsis_mask", ATTR_DESC(ellipsis_mask, AnyTraits<int64_t>())},
|
||
{"new_axis_mask", ATTR_DESC(new_axis_mask, AnyTraits<int64_t>())},
|
||
{"shrink_axis_mask", ATTR_DESC(shrink_axis_mask, AnyTraits<int64_t>())}};
|
||
// 属性映射,列出了"begin_mask"、"end_mask"、"ellipsis_mask"、"new_axis_mask"、"shrink_axis_mask"五个属性,类型分别为int64_t
|
||
OUTPUT_MAP(StridedSliceV2) = {{0, OUTPUT_DESC(y)}};
|
||
// 输出映射,y索引为0
|
||
REG_ADPT_DESC(StridedSliceV2, kNameStridedSliceV2, ADPT_DESC(StridedSliceV2))
|
||
// 注册StridedSliceV2操作的适配器描述StridedSliceV2
|
||
|
||
// UnsortedSegmentSum
|
||
INPUT_MAP(UnsortedSegmentSumD) = {{1, INPUT_DESC(x)}, {2, INPUT_DESC(segment_ids)}};
|
||
// 输入映射,共两个,x索引为1,segment_ids索引为2
|
||
INPUT_ATTR_MAP(UnsortedSegmentSumD) = {{3, ATTR_DESC(num_segments, AnyTraits<int64_t>())}};
|
||
// 输入属性映射,num_segments类型是int64_t,索引是3
|
||
ATTR_MAP(UnsortedSegmentSumD) = EMPTY_ATTR_MAP;
|
||
// 属性映射,空
|
||
OUTPUT_MAP(UnsortedSegmentSumD) = {{0, OUTPUT_DESC(y)}};
|
||
//输出映射,y索引为0
|
||
REG_ADPT_DESC(UnsortedSegmentSumD, prim::kPrimUnsortedSegmentSum->name(), ADPT_DESC(UnsortedSegmentSumD))
|
||
//注册UnsortedSegmentSumD操作的适配器描述prim::kPrimUnsortedSegmentSum->name()
|
||
|
||
// UnsortedSegmentProdD
|
||
INPUT_MAP(UnsortedSegmentProdD) = {{1, INPUT_DESC(x)}, {2, INPUT_DESC(segment_ids)}};
|
||
// 输入映射,共两个,x索引为1,segment_ids索引为2
|
||
INPUT_ATTR_MAP(UnsortedSegmentProdD) = {{3, ATTR_DESC(num_segments, AnyTraits<int64_t>())}};
|
||
// 输入属性映射,num_segments类型是int64_t,索引是3
|
||
ATTR_MAP(UnsortedSegmentProdD) = EMPTY_ATTR_MAP;
|
||
// 属性映射,空
|
||
OUTPUT_MAP(UnsortedSegmentProdD) = {{0, OUTPUT_DESC(y)}};
|
||
// 输出映射,y索引为0
|
||
REG_ADPT_DESC(UnsortedSegmentProdD, kNameUnsortedSegmentProdD, ADPT_DESC(UnsortedSegmentProdD))
|
||
// 注册UnsortedSegmentSumD操作的适配器描述kNameUnsortedSegmentProdD
|
||
|
||
|
||
// UnsortedSegmentMaxD
|
||
INPUT_MAP(UnsortedSegmentMaxD) = {{1, INPUT_DESC(x)}, {2, INPUT_DESC(segment_ids)}};
|
||
// 输入映射,共两个,x索引为1,segment_ids索引为2
|
||
INPUT_ATTR_MAP(UnsortedSegmentMaxD) = {{3, ATTR_DESC(num_segments, AnyTraits<int64_t>())}};
|
||
// 输入属性映射,num_segments类型是int64_t,索引是3
|
||
ATTR_MAP(UnsortedSegmentMaxD) = EMPTY_ATTR_MAP;
|
||
// 属性映射,空
|
||
OUTPUT_MAP(UnsortedSegmentMaxD) = {{0, OUTPUT_DESC(y)}};
|
||
// 输出映射,y索引为0
|
||
REG_ADPT_DESC(UnsortedSegmentMaxD, kNameUnsortedSegmentMaxD, ADPT_DESC(UnsortedSegmentMaxD))
|
||
// 注册UnsortedSegmentMaxD操作的适配器描述kNameUnsortedSegmentMaxD
|
||
|
||
// UnsortedSegmentMin
|
||
INPUT_MAP(UnsortedSegmentMin) = {{1, INPUT_DESC(x)}, {2, INPUT_DESC(segment_ids)}, {3, INPUT_DESC(num_segments)}};
|
||
// 输入映射,共三个,x索引为1,segment_ids索引为2,num_segments索引为3
|
||
ATTR_MAP(UnsortedSegmentMin) = EMPTY_ATTR_MAP;
|
||
// 属性映射,空
|
||
OUTPUT_MAP(UnsortedSegmentMin) = {{0, OUTPUT_DESC(y)}};
|
||
// 输出映射,y索引为0
|
||
REG_ADPT_DESC(UnsortedSegmentMin, prim::kPrimUnsortedSegmentMin->name(), ADPT_DESC(UnsortedSegmentMin))
|
||
// 注册UnsortedSegmentMin操作的适配器描述prim::kPrimUnsortedSegmentMin->name()
|
||
|
||
// ReverseV2
|
||
INPUT_MAP(ReverseV2D) = {{1, INPUT_DESC(x)}};
|
||
// 输入映射,x索引为1
|
||
ATTR_MAP(ReverseV2D) = {{"axis", ATTR_DESC(axis, AnyTraits<int64_t>(), AnyTraits<std::vector<int64_t>>())}};
|
||
//属性映射,axis类型为int64_t和std::vector<int64_t>
|
||
OUTPUT_MAP(ReverseV2D) = {{0, OUTPUT_DESC(y)}};
|
||
//输出映射,y索引为0
|
||
REG_ADPT_DESC(ReverseV2D, kNameReverseV2, ADPT_DESC(ReverseV2D))
|
||
// 注册ReverseV2D操作的适配器描述kNameReverseV2
|
||
} // namespace mindspore::transform
|