| 88 | } |
| 89 | |
| 90 | static 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 | } |
no test coverage detected