4 gathers with same embedding dim → should fuse into 1 batched gather
| 38 | |
| 39 | // 4 gathers with same embedding dim → should fuse into 1 batched gather |
| 40 | TEST_CASE(gather_horiz_fusion_basic) |
| 41 | { |
| 42 | migraphx::module m1; |
| 43 | { |
| 44 | auto emb1 = |
| 45 | m1.add_literal(migraphx::generate_literal({migraphx::shape::float_type, {3, 2}}, 0)); |
| 46 | auto emb2 = |
| 47 | m1.add_literal(migraphx::generate_literal({migraphx::shape::float_type, {4, 2}}, 1)); |
| 48 | auto emb3 = |
| 49 | m1.add_literal(migraphx::generate_literal({migraphx::shape::float_type, {2, 2}}, 2)); |
| 50 | auto emb4 = |
| 51 | m1.add_literal(migraphx::generate_literal({migraphx::shape::float_type, {5, 2}}, 3)); |
| 52 | |
| 53 | auto idx1 = m1.add_parameter("idx1", {migraphx::shape::int32_type, {2}}); |
| 54 | auto idx2 = m1.add_parameter("idx2", {migraphx::shape::int32_type, {3}}); |
| 55 | auto idx3 = m1.add_parameter("idx3", {migraphx::shape::int32_type, {1}}); |
| 56 | auto idx4 = m1.add_parameter("idx4", {migraphx::shape::int32_type, {2}}); |
| 57 | |
| 58 | auto g1 = m1.add_instruction(migraphx::make_op("gather", {{"axis", 0}}), emb1, idx1); |
| 59 | auto g2 = m1.add_instruction(migraphx::make_op("gather", {{"axis", 0}}), emb2, idx2); |
| 60 | auto g3 = m1.add_instruction(migraphx::make_op("gather", {{"axis", 0}}), emb3, idx3); |
| 61 | auto g4 = m1.add_instruction(migraphx::make_op("gather", {{"axis", 0}}), emb4, idx4); |
| 62 | |
| 63 | // Combine all outputs so every gather stays live through DCE |
| 64 | m1.add_instruction(migraphx::make_op("concat", {{"axis", 0}}), |
| 65 | std::vector<migraphx::instruction_ref>{g1, g2, g3, g4}); |
| 66 | } |
| 67 | run_pass(m1); |
| 68 | |
| 69 | migraphx::module m2; |
| 70 | { |
| 71 | // Embedding literals (added first → pushed to front → end up at the back of no-dep list) |
| 72 | auto emb1 = |
| 73 | m2.add_literal(migraphx::generate_literal({migraphx::shape::float_type, {3, 2}}, 0)); |
| 74 | auto emb2 = |
| 75 | m2.add_literal(migraphx::generate_literal({migraphx::shape::float_type, {4, 2}}, 1)); |
| 76 | auto emb3 = |
| 77 | m2.add_literal(migraphx::generate_literal({migraphx::shape::float_type, {2, 2}}, 2)); |
| 78 | auto emb4 = |
| 79 | m2.add_literal(migraphx::generate_literal({migraphx::shape::float_type, {5, 2}}, 3)); |
| 80 | |
| 81 | // Parameters (added second → in middle of no-dep list) |
| 82 | auto idx1 = m2.add_parameter("idx1", {migraphx::shape::int32_type, {2}}); |
| 83 | auto idx2 = m2.add_parameter("idx2", {migraphx::shape::int32_type, {3}}); |
| 84 | auto idx3 = m2.add_parameter("idx3", {migraphx::shape::int32_type, {1}}); |
| 85 | auto idx4 = m2.add_parameter("idx4", {migraphx::shape::int32_type, {2}}); |
| 86 | |
| 87 | // Offset literals (added last → pushed to very front of no-dep list, |
| 88 | // matching order of add_literal calls inside the pass's fuse loop) |
| 89 | auto offset2 = m2.add_literal( |
| 90 | migraphx::literal{migraphx::shape{migraphx::shape::int32_type}, {std::size_t(3)}}); |
| 91 | auto offset3 = m2.add_literal( |
| 92 | migraphx::literal{migraphx::shape{migraphx::shape::int32_type}, {std::size_t(7)}}); |
| 93 | auto offset4 = m2.add_literal( |
| 94 | migraphx::literal{migraphx::shape{migraphx::shape::int32_type}, {std::size_t(9)}}); |
| 95 | |
| 96 | // Concatenated embedding table: [3+4+2+5, 2] = [14, 2] |
| 97 | auto concat_emb = |
nothing calls this directly
no test coverage detected