transform/selection_ops_declare.cc

321 lines
16 KiB
C++
Raw 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 <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索引为1indices索引为2updates索引为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索引为1indices索引为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索引为1dim索引为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索引为1y索引为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索引为1y索引为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索引为1y索引为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索引为1x1索引为2x2索引为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索引为1shape索引为2begin索引为3end索引为4strides索引为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索引为1begin索引为2end索引为3strides索引为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索引为1begin索引为2end索引为3axes索引为4strides索引为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索引为1segment_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索引为1segment_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索引为1segment_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索引为1segment_ids索引为2num_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