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

Function generic_eval

src/program.cpp:464–551  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

462
463template <class F>
464static 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 }

Callers 5

eval_with_contextMethod · 0.85
evalMethod · 0.85
markMethod · 0.85
perf_reportMethod · 0.85
dry_runMethod · 0.85

Calls 15

iterator_forFunction · 0.85
traceFunction · 0.85
containsFunction · 0.85
is_compatible_shapeFunction · 0.85
emplaceMethod · 0.80
get_argumentMethod · 0.80
atMethod · 0.80
any_of_dynamicMethod · 0.80
resizeMethod · 0.80
normalized_operatorMethod · 0.80
get_target_idMethod · 0.80
get_main_moduleMethod · 0.80

Tested by

no test coverage detected