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

Function parse_dyn_split

src/onnx/parse_split.cpp:37–88  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

35namespace onnx {
36
37static 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
90static auto parse_static_split(const onnx_parser::node_info& info,
91 const onnx_parser& parser,

Callers 1

parseMethod · 0.85

Calls 7

containsFunction · 0.85
atMethod · 0.80
make_opFunction · 0.50
maxClass · 0.50
sizeMethod · 0.45
add_instructionMethod · 0.45
add_literalMethod · 0.45

Tested by

no test coverage detected