MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / ensure_init_graph

Method ensure_init_graph

src/jit/test/helper.cpp:64–153  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

62}
63
64void 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) {

Callers

nothing calls this directly

Calls 15

unpack_vectorFunction · 0.85
gradFunction · 0.85
make_callback_copyFunction · 0.85
renameMethod · 0.80
symshapeMethod · 0.80
emplace_backMethod · 0.80
resizeMethod · 0.80
makeFunction · 0.50
nodeMethod · 0.45
findMethod · 0.45
endMethod · 0.45
owner_oprMethod · 0.45

Tested by

no test coverage detected