MCPcopy Create free account
hub / github.com/ROCm/AMDMIGraphX / parse

Method parse

src/tf/parse_split.cpp:38–116  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

36 std::vector<op_desc> operators() const { return {{"Split"}, {"SplitV"}}; }
37
38 std::vector<instruction_ref> parse(const op_desc& /*opd*/,
39 const tf_parser& /*parser*/,
40 tf_parser::node_info info,
41 std::vector<instruction_ref> args) const
42 {
43 bool vector_as_input = args.size() == 3;
44 int num_outputs = 1;
45 auto axis_arg = args[0];
46 auto input_arg = args[1];
47 if(vector_as_input)
48 {
49 input_arg = args[0];
50 axis_arg = args[2];
51 }
52
53 if(contains(info.attributes, "num_split"))
54 num_outputs = info.attributes.at("num_split").i();
55
56 std::vector<int> splits(num_outputs);
57 std::vector<int> slice_pos{0};
58 if(vector_as_input)
59 {
60 splits = args[1]->eval().get<int32_t>().to_vector();
61 num_outputs = splits.size();
62 }
63
64 assert(num_outputs > 0);
65
66 if(num_outputs == 1)
67 return std::vector<instruction_ref>{
68 info.add_instruction(make_op("identity"), input_arg)};
69
70 auto lens = input_arg->get_shape().lens();
71 auto num_dims = lens.size();
72 int axis = axis_arg->eval().at<int32_t>();
73
74 // ensure split is made evenly if "num_split" is used
75 assert(vector_as_input or lens[axis] % num_outputs == 0);
76
77 auto split_size = lens[axis] / num_outputs;
78
79 // push back first end point of slice
80 if(vector_as_input)
81 {
82 slice_pos.push_back(splits[0]);
83 }
84 else
85 {
86 slice_pos.push_back(split_size);
87 }
88
89 // calculate remaining end points for each slice
90 for(auto i = 1; i < num_outputs; i++)
91 {
92 if(vector_as_input)
93 {
94 splits[i] += splits[i - 1];
95 slice_pos.push_back(splits[i]);

Callers

nothing calls this directly

Calls 13

containsFunction · 0.85
iotaFunction · 0.85
atMethod · 0.80
lensMethod · 0.80
make_opFunction · 0.50
sizeMethod · 0.45
to_vectorMethod · 0.45
evalMethod · 0.45
add_instructionMethod · 0.45
get_shapeMethod · 0.45
push_backMethod · 0.45
beginMethod · 0.45

Tested by

no test coverage detected