| 462 | |
| 463 | template <class F> |
| 464 | static std::vector<argument> generic_eval(const module* mod, |
| 465 | std::vector<context>& ctx, |
| 466 | const std::unordered_map<std::string, argument>& params, |
| 467 | pmr::unordered_map<instruction_ref, argument>& results, |
| 468 | F trace) |
| 469 | { |
| 470 | assert(mod->validate() == mod->end()); |
| 471 | std::vector<argument> values; |
| 472 | values.reserve(16); |
| 473 | for(auto ins : iterator_for(*mod)) |
| 474 | { |
| 475 | assert(mod->name() != "main" or results.find(ins) == results.end()); |
| 476 | #ifndef NDEBUG |
| 477 | results.emplace(ins, argument{}); |
| 478 | #endif |
| 479 | const auto& name = ins->name(); |
| 480 | if(name == "@literal") |
| 481 | { |
| 482 | results.insert_or_assign(ins, |
| 483 | trace(ins, [&] { return ins->get_literal().get_argument(); })); |
| 484 | } |
| 485 | else if(name == "@param") |
| 486 | { |
| 487 | results.insert_or_assign( |
| 488 | ins, trace(ins, [&] { |
| 489 | auto param_name = any_cast<builtin::param>(ins->get_operator()).parameter; |
| 490 | if(not contains(params, param_name)) |
| 491 | MIGRAPHX_THROW("Parameter not found: " + param_name); |
| 492 | auto param = params.at(param_name); |
| 493 | // TODO: may want to check correct number of dimensions and/or was within bounds |
| 494 | if(not ins->get_shape().any_of_dynamic() and |
| 495 | param.get_shape() != ins->get_shape()) |
| 496 | { |
| 497 | MIGRAPHX_THROW("Incorrect shape {" + to_string(param.get_shape()) + |
| 498 | "} for parameter: " + param_name + |
| 499 | " should be: " + to_string(ins->get_shape())); |
| 500 | } |
| 501 | return param; |
| 502 | })); |
| 503 | } |
| 504 | else if(name == "@outline") |
| 505 | { |
| 506 | results.insert_or_assign( |
| 507 | ins, trace(ins, [&] { return argument{ins->get_shape(), nullptr}; })); |
| 508 | } |
| 509 | else if(name == "@return") |
| 510 | { |
| 511 | std::vector<argument> prog_outputs; |
| 512 | std::transform(ins->inputs().begin(), |
| 513 | ins->inputs().end(), |
| 514 | std::back_inserter(prog_outputs), |
| 515 | [&](instruction_ref i) { |
| 516 | assert(results.find(i) != results.end()); |
| 517 | return results[i]; |
| 518 | }); |
| 519 | |
| 520 | return prog_outputs; |
| 521 | } |
no test coverage detected