| 202 | } |
| 203 | |
| 204 | void 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()); |
no test coverage detected