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

Function TEST_CASE

test/fuse_horizontal_test.cpp:40–138  ·  view source on GitHub ↗

4 gathers with same embedding dim → should fuse into 1 batched gather

Source from the content-addressed store, hash-verified

38
39// 4 gathers with same embedding dim → should fuse into 1 batched gather
40TEST_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 =

Callers

nothing calls this directly

Calls 6

generate_literalFunction · 0.85
add_parameterMethod · 0.80
run_passFunction · 0.70
make_opFunction · 0.50
add_literalMethod · 0.45
add_instructionMethod · 0.45

Tested by

no test coverage detected