| 62 | } |
| 63 | |
| 64 | void FusionChecker::ensure_init_graph() { |
| 65 | if (m_jit_y.node()) |
| 66 | return; |
| 67 | |
| 68 | m_graph = ComputingGraph::make(); |
| 69 | SymbolVarArray inputs(m_nr_input); |
| 70 | for (size_t i = 0; i < m_nr_input; ++i) { |
| 71 | inputs[i] = opr::Host2DeviceCopy::make(*m_graph, m_inputs_val[i]) |
| 72 | .rename(ssprintf("inp%zu", i)); |
| 73 | |
| 74 | auto dt = m_idx2dtype.find(i); |
| 75 | if (dt != m_idx2dtype.end()) { |
| 76 | inputs[i] = opr::TypeCvt::make(inputs[i], dt->second); |
| 77 | } |
| 78 | } |
| 79 | m_truth_y = m_exp_func(inputs); |
| 80 | |
| 81 | SymbolVar jit_y; |
| 82 | if (m_direct_build) { |
| 83 | auto ig_gen = |
| 84 | std::make_unique<InternalGraphGenerator>(m_truth_y.node()->owner_opr()); |
| 85 | ThinHashSet<VarNode*> endpoints_set; |
| 86 | for (size_t i = 0; i < m_nr_input; ++i) { |
| 87 | endpoints_set.insert(inputs[i].node()); |
| 88 | } |
| 89 | for (auto&& opr : get_rev_topo_order(m_truth_y, endpoints_set)) |
| 90 | ig_gen->add_opr(opr); |
| 91 | jit_y = JITExecutor::make(ig_gen->generate(), cg::to_var_node_array(inputs)); |
| 92 | } else { |
| 93 | ComputingGraph::Options opt; |
| 94 | opt.graph_opt_level = 3; |
| 95 | opt.graph_opt.jit = m_jit_level; |
| 96 | unpack_vector( |
| 97 | gopt::GraphOptimizer{} |
| 98 | .add_preset_passes(true, nullptr, &opt) |
| 99 | .apply({{m_truth_y}}) |
| 100 | .endpoint_vars(), |
| 101 | jit_y); |
| 102 | |
| 103 | size_t nr_jit_opr = 0; |
| 104 | cg::DepOprIter{[&nr_jit_opr, this](cg::OperatorNodeBase* opr) { |
| 105 | if (opr->same_type<JITExecutor>()) { |
| 106 | ++nr_jit_opr; |
| 107 | } else { |
| 108 | static const ThinHashSet<Typeinfo*> allowed_types{ |
| 109 | opr::Host2DeviceCopy::typeinfo(), opr::GetVarShape::typeinfo()}; |
| 110 | mgb_throw_if( |
| 111 | m_check_opr_type && !allowed_types.count(opr->dyn_typeinfo()), |
| 112 | InternalError, "encountered non-JIT opr after fusion: %s{%s}", |
| 113 | opr->cname(), opr->dyn_typeinfo()->name); |
| 114 | } |
| 115 | }}.add(jit_y.node()); |
| 116 | mgb_assert(nr_jit_opr == 1); |
| 117 | } |
| 118 | |
| 119 | SymbolVar loss_var0, loss_var1; |
| 120 | SmallVector<std::tuple<size_t, SymbolVar, SymbolVar>> grad_vars; |
| 121 | for (size_t i = 0; i < m_nr_input; ++i) { |
nothing calls this directly
no test coverage detected