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

Function parse_static_split

src/onnx/parse_split.cpp:90–160  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

88}
89
90static auto parse_static_split(const onnx_parser::node_info& info,
91 const onnx_parser& parser,
92 const std::vector<instruction_ref>& args,
93 int64_t tuned_axis)
94{
95 const auto& input_shape = args[0]->get_shape();
96 // either static shape or fixed dynamic_dimension for split axis
97 auto tuned_axis_len = input_shape.to_static(0).lens().at(tuned_axis);
98 std::vector<int64_t> vec_splits;
99 if(contains(info.attributes, "split"))
100 {
101 literal s = parser.parse_value(info.attributes.at("split"));
102 s.visit([&](auto v) { vec_splits.assign(v.begin(), v.end()); });
103 }
104 else if(args.size() == 2)
105 {
106 auto s = args[1]->eval();
107 check_arg_empty(s, "PARSE_SPLIT: non-constant `split` input is not supported");
108 s.visit([&](auto v) { vec_splits.assign(v.begin(), v.end()); });
109 }
110 // no split attribute, input is equally divided
111 else
112 {
113 std::size_t num_outputs = info.num_outputs;
114 // the num_outputs attribute seems to be redundant since we already have
115 // node_info::num_outputs, but we can still perform an error check
116 if(contains(info.attributes, "num_outputs"))
117 {
118 num_outputs = parser.parse_value(info.attributes.at("num_outputs")).at<std::size_t>();
119 if(num_outputs != info.num_outputs)
120 {
121 MIGRAPHX_THROW("PARSE_SPLIT: num_outputs attribute " + std::to_string(num_outputs) +
122 " doesn't match actual number of outputs " +
123 std::to_string(info.num_outputs) + "!");
124 }
125 }
126 if(tuned_axis_len % num_outputs == 0)
127 {
128 std::size_t chunk_size = tuned_axis_len / num_outputs;
129 vec_splits.resize(num_outputs, chunk_size);
130 }
131 else
132 {
133 std::size_t chunk_size = tuned_axis_len / num_outputs + 1;
134 std::size_t last_chunk_size = tuned_axis_len - chunk_size * (num_outputs - 1);
135 vec_splits.resize(num_outputs - 1, chunk_size);
136 vec_splits.push_back(last_chunk_size);
137 }
138 }
139
140 if(std::accumulate(vec_splits.begin(), vec_splits.end(), int64_t(0)) !=
141 static_cast<int64_t>(tuned_axis_len))
142 {
143 MIGRAPHX_THROW(
144 "PARSE_SPLIT: sum of split attribute unequal to dim size of axis! tuned axis:" +
145 std::to_string(tuned_axis_len) + " Output " + to_string_range(vec_splits) + " Rank " +
146 std::to_string(input_shape.ndim()));
147 }

Callers 1

parseMethod · 0.85

Calls 15

containsFunction · 0.85
check_arg_emptyFunction · 0.85
accumulateFunction · 0.85
to_string_rangeFunction · 0.85
atMethod · 0.80
lensMethod · 0.80
to_staticMethod · 0.80
parse_valueMethod · 0.80
resizeMethod · 0.80
ndimMethod · 0.80
to_stringFunction · 0.50
make_opFunction · 0.50

Tested by

no test coverage detected