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

Method gen_func_body

src/jit/impl/mlir/mlir_gen.cpp:119–162  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

117 }
118
119 mlir::LogicalResult gen_func_body(
120 const InternalGraph& internal_graph, const JITExecutor::Args& args) {
121 llvm::ScopedHashTableScope<llvm::StringRef, mlir::Value> var_scope(
122 m_symbol_table);
123 cg::DepOprIter{[&](cg::OperatorNodeBase* opr) {
124 if (opr->same_type<JITPlaceholder>()) {
125 return;
126 } else if (opr->same_type<opr::ImmutableTensor>()) {
127 auto imm = SymbolVar{opr->output(0)}.as_immutable_scalar();
128 if (imm.valid()) {
129 auto dtype = imm->dtype();
130 float scalar_value;
131 if (dtype == dtype::Float32()) {
132 scalar_value = imm->get<float>();
133 } else {
134 mgb_throw(
135 InternalError,
136 "mlir backend currently only support f32 "
137 "dtype, but got %s",
138 dtype.name());
139 }
140 auto&& out = m_builder.create<dialect::ConstantScalarOp>(
141 m_builder.getUnknownLoc(), m_builder.getF32Type(),
142 m_builder.getF32FloatAttr(scalar_value));
143 mgb_assert(mlir::succeeded(declare(opr->output(0)->name(), out)));
144 }
145 } else if (opr->same_type<opr::Elemwise>()) {
146 auto&& out = gen_elemwise(opr->cast_final<opr::Elemwise>());
147 mgb_assert(mlir::succeeded(declare(opr->output(0)->name(), out)));
148 return;
149 } else if (opr->same_type<opr::Dimshuffle>()) {
150 auto&& out = gen_dimshuffle(opr->cast_final<opr::Dimshuffle>());
151 mgb_assert(mlir::succeeded(declare(opr->output(0)->name(), out)));
152 } else if (opr->same_type<opr::TypeCvt>()) {
153 auto&& out = gen_typecvt(opr->cast_final<opr::TypeCvt>());
154 mgb_assert(mlir::succeeded(declare(opr->output(0)->name(), out)));
155 }
156 }}.add(internal_graph.output());
157 m_builder.create<dialect::AssignOp>(
158 m_builder.getUnknownLoc(), get(internal_graph.output()),
159 get(args.outputs[0].from));
160
161 return mlir::success();
162 }
163
164 mlir::Value gen_elemwise(const opr::Elemwise& opr) {
165 llvm::SmallVector<mlir::Value, 4> operands;

Callers

nothing calls this directly

Calls 7

getFunction · 0.85
as_immutable_scalarMethod · 0.80
addMethod · 0.45
outputMethod · 0.45
validMethod · 0.45
dtypeMethod · 0.45
nameMethod · 0.45

Tested by

no test coverage detected