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

Method update_args

src/jit/impl/executor_opr.cpp:204–279  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

202}
203
204void JITExecutor::update_args() {
205 m_args.outputs.clear();
206 for (auto out : output()) {
207 m_args.outputs.push_back({out, out->layout(), -1});
208 }
209 m_args.inputs.resize(input().size());
210
211 auto is_host_value_shape_input = [this](size_t idx) {
212 return m_internal_graph->placeholders().at(idx)->is_host_value_shape_input();
213 };
214
215 for (size_t i = 0; i < input().size(); i++) {
216 auto&& dst_data = m_args.inputs[i];
217 dst_data.from = input(i);
218 dst_data.idx = i;
219 if (is_host_value_shape_input(i)) {
220 auto&& mgr = owner_graph()->static_infer_manager();
221 auto&& shpval_inp_val = &mgr.infer_value(input(i));
222 cg::copy_tensor_value_to_shape(dst_data.layout, *shpval_inp_val);
223 dst_data.layout.dtype = {};
224 for (size_t i = 0; i < dst_data.layout.ndim; ++i) {
225 dst_data.layout.stride[i] = 0;
226 }
227 } else {
228 dst_data.layout = input(i)->layout();
229 }
230 }
231
232 //! dimshuffle opr need to change the input.
233 if (has_dimshuffle()) {
234 do_dimshuffle();
235 }
236
237 if (m_compiler->property().contain_flag(CPFlag::NEED_INPUT_COLLAPSE)) {
238 // collective collapse datum layout, try to reduce the output ndim
239 opr::Elemwise::TensorLayoutPtrArray inp_layouts;
240 inp_layouts.reserve(m_args.inputs.size());
241 for (size_t i = 0; i < m_args.inputs.size(); i++) {
242 if (!is_host_value_shape_input(i)) {
243 inp_layouts.push_back(&m_args.inputs[i].layout);
244 }
245 }
246 opr::Elemwise::broadcast_collective_collapse(
247 inp_layouts, &m_args.outputs[0].layout);
248 }
249
250 // compute and update hash
251 XXHash hstate;
252
253 // update layout info
254 auto prop = m_compiler->property();
255 if (prop.contain_flag(CPFlag::BIND_NDIM | CPFlag::BIND_SHAPE)) {
256 mgb_assert(
257 prop.contain_flag(CPFlag::BIND_NDIM),
258 "BIND_NDIM must be set if bind_shape is set");
259 std::vector<size_t> buf;
260 buf.reserve(1024);
261 buf.push_back(m_args.inputs.size());

Callers 1

executor_opr.cppFile · 0.80

Calls 15

has_dimshuffleFunction · 0.85
resizeMethod · 0.80
infer_valueMethod · 0.80
clearMethod · 0.45
push_backMethod · 0.45
layoutMethod · 0.45
sizeMethod · 0.45
atMethod · 0.45
contain_flagMethod · 0.45
propertyMethod · 0.45
reserveMethod · 0.45

Tested by

no test coverage detected