| 237 | } |
| 238 | |
| 239 | static migraphx::program create_scatter_elements_program_3x3(const std::string& reduction_mode, |
| 240 | int axis) |
| 241 | { |
| 242 | migraphx::program p; |
| 243 | auto* mm = p.get_main_module(); |
| 244 | migraphx::shape sd{migraphx::shape::float_type, {3, 3}}; |
| 245 | std::vector<float> vd(sd.elements(), 3.0f); |
| 246 | |
| 247 | migraphx::shape si{migraphx::shape::int32_type, {3, 2}}; |
| 248 | std::vector<int> vi = {1, 0, 0, 2, 2, 1}; |
| 249 | |
| 250 | migraphx::shape su{migraphx::shape::float_type, {3, 2}}; |
| 251 | std::vector<float> vu = {1.0, 7.0, 1.1, 7.1, 1.2, 7.2}; |
| 252 | |
| 253 | auto ld = mm->add_literal(migraphx::literal{sd, vd}); |
| 254 | auto li = mm->add_literal(migraphx::literal{si, vi}); |
| 255 | auto lu = mm->add_literal(migraphx::literal{su, vu}); |
| 256 | auto r = mm->add_instruction( |
| 257 | migraphx::make_op("scatter_" + reduction_mode, {{"axis", axis}}), ld, li, lu); |
| 258 | mm->add_return({r}); |
| 259 | return p; |
| 260 | } |
| 261 | |
| 262 | TEST_CASE(scatter_elements_none_axis_0_3x3_test) |
| 263 | { |
no test coverage detected