| 569 | } |
| 570 | |
| 571 | void dispatch_custom_op( |
| 572 | std::shared_ptr<const CustomOp> op, const Param& param, |
| 573 | std::shared_ptr<::megdnn::SmallVector<::mgb::DeviceTensorND>> inputs, |
| 574 | std::shared_ptr<::megdnn::SmallVector<::mgb::DeviceTensorND>> outputs) { |
| 575 | if (outputs->size() == 0) { |
| 576 | return; |
| 577 | } |
| 578 | |
| 579 | auto compnode = outputs->at(0).comp_node(); |
| 580 | if (compnode.device_type() == CompNode::DeviceType::CPU) { |
| 581 | auto&& cpu_env = CompNodeEnv::from_comp_node(compnode).cpu_env(); |
| 582 | cpu_env.dispatch([op, param, inputs, outputs]() { |
| 583 | compute_impl(op, param, inputs, outputs); |
| 584 | }); |
| 585 | |
| 586 | } else { |
| 587 | mgb_assert( |
| 588 | compnode.device_type() == CompNode::DeviceType::CUDA, |
| 589 | "custom op only support cuda/cpu now, but get %s", |
| 590 | compnode.to_string().c_str()); |
| 591 | compute_impl(op, param, inputs, outputs); |
| 592 | } |
| 593 | } |
| 594 | |
| 595 | } // namespace custom |
| 596 | |