| 35 | namespace onnx { |
| 36 | |
| 37 | static auto parse_dyn_split(const onnx_parser::node_info& info, |
| 38 | const std::vector<instruction_ref>& args, |
| 39 | int64_t tuned_axis) |
| 40 | { |
| 41 | if(contains(info.attributes, "split")) |
| 42 | { |
| 43 | MIGRAPHX_THROW("PARSE_SPLIT: dynamic input and non-fixed split axis and `split` " |
| 44 | "attribute not supported"); |
| 45 | } |
| 46 | if(args.size() == 2) |
| 47 | { |
| 48 | MIGRAPHX_THROW("PARSE_SPLIT: dynamic input and non-fixed split axis and `split` " |
| 49 | "input not supported"); |
| 50 | } |
| 51 | |
| 52 | std::size_t num_outputs = info.num_outputs; |
| 53 | std::vector<instruction_ref> ret_ins(num_outputs); |
| 54 | |
| 55 | // Doing shape calculations for the splits in the graph |
| 56 | auto split_dim = info.add_instruction( |
| 57 | make_op("dimensions_of", {{"start", tuned_axis}, {"end", tuned_axis + 1}}), args[0]); |
| 58 | shape int64_scalar_shape{shape::int64_type, {1}, {0}}; |
| 59 | auto num_outputs_lit = info.add_literal(literal{int64_scalar_shape, {num_outputs}}); |
| 60 | auto num_outputs_minus_1_lit = info.add_literal(literal{int64_scalar_shape, {num_outputs - 1}}); |
| 61 | // (A + (B - 1)) / B == ceil(A / B) |
| 62 | auto chunk_size = info.add_instruction( |
| 63 | make_op("div"), |
| 64 | info.add_instruction(make_op("add"), split_dim, num_outputs_minus_1_lit), |
| 65 | num_outputs_lit); |
| 66 | for(int n = 0; n < num_outputs - 1; ++n) |
| 67 | { |
| 68 | // slice(input, starts = {n * chunk_size}, ends = {(n+1) * chunk_size}); axes = |
| 69 | // {tuned_axis} |
| 70 | ret_ins.at(n) = info.add_instruction( |
| 71 | make_op("slice", {{"axes", {tuned_axis}}}), |
| 72 | args[0], |
| 73 | info.add_instruction( |
| 74 | make_op("mul"), chunk_size, info.add_literal(literal{int64_scalar_shape, {n}})), |
| 75 | info.add_instruction(make_op("mul"), |
| 76 | chunk_size, |
| 77 | info.add_literal(literal{int64_scalar_shape, {n + 1}}))); |
| 78 | } |
| 79 | // last slice: slice(input, starts = {n * chunk_size}); ends = max_int, axes = |
| 80 | // {tuned_axis} |
| 81 | ret_ins.at(num_outputs - 1) = info.add_instruction( |
| 82 | make_op("slice", {{"axes", {tuned_axis}}, {"ends", {std::numeric_limits<int64_t>::max()}}}), |
| 83 | args[0], |
| 84 | info.add_instruction(make_op("mul"), |
| 85 | chunk_size, |
| 86 | info.add_literal(literal{int64_scalar_shape, {num_outputs - 1}}))); |
| 87 | return ret_ins; |
| 88 | } |
| 89 | |
| 90 | static auto parse_static_split(const onnx_parser::node_info& info, |
| 91 | const onnx_parser& parser, |
no test coverage detected