| 506 | } |
| 507 | |
| 508 | void exec_ternary( |
| 509 | HandleImpl* handle, const TensorND& src0, const TensorND& src1, |
| 510 | const TensorND& src2, _megdnn_tensor_out dst, const param::Elemwise::Mode& mode, |
| 511 | const WorkspaceBundle& wk_bundle) { |
| 512 | auto cnnl_handler = handle->cnnl_handle(); |
| 513 | auto dtype_dest = dst.layout.dtype.enumv(); |
| 514 | |
| 515 | CnnlTensorDescriptor src0_desc, src1_desc, src2_desc, output_desc; |
| 516 | src0_desc.set(src0.layout); |
| 517 | src1_desc.set(src1.layout); |
| 518 | src2_desc.set(src2.layout); |
| 519 | output_desc.set(dst.layout); |
| 520 | |
| 521 | auto convert_mode_to_logicOp = [mode]() { |
| 522 | if (mode == Mode::COND_LEQ_MOV) { |
| 523 | return cnnlLogicOp_t::CNNL_LOGIC_OP_LE; |
| 524 | } else if (mode == Mode::COND_LT_MOV) { |
| 525 | return cnnlLogicOp_t::CNNL_LOGIC_OP_LT; |
| 526 | } else { |
| 527 | megdnn_throw("unsupport elemwise ternary opr"); |
| 528 | } |
| 529 | }; |
| 530 | |
| 531 | switch (mode) { |
| 532 | case Mode::COND_LT_MOV: |
| 533 | case Mode::COND_LEQ_MOV: { // float, half, int32 |
| 534 | megdnn_assert( |
| 535 | check_dtype_float_ieee(dtype_dest) || |
| 536 | dtype_dest == megdnn::DTypeEnum::Int32, |
| 537 | "Cambricon unsupport elemwise mode:%d with dtype:%d", |
| 538 | static_cast<int>(mode), static_cast<int>(dtype_dest)); |
| 539 | Workspace logic_wk, src0_wk, logic_res_wk, optensor_wk; |
| 540 | logic_wk = wk_bundle.get_workspace(0); |
| 541 | src0_wk = wk_bundle.get_workspace(1); |
| 542 | logic_res_wk = wk_bundle.get_workspace(2); |
| 543 | optensor_wk = wk_bundle.get_workspace(3); |
| 544 | TensorShapeArray src0_1; |
| 545 | src0_1.push_back(src0.layout); |
| 546 | src0_1.push_back(src1.layout); |
| 547 | TensorShape logic_res_shape; |
| 548 | Elemwise::deduce_shape(src0_1, logic_res_shape); |
| 549 | TensorLayout logic_res_layout(logic_res_shape, src0.layout.dtype); |
| 550 | void* src0_ptr = to_broadcast( |
| 551 | cnnl_handler, src0, logic_res_layout, src0_desc, src0_wk); |
| 552 | CnnlTensorDescriptor logic_res_desc; |
| 553 | logic_res_desc.set(logic_res_layout); |
| 554 | cnnl_check(cnnlLogicOp( |
| 555 | cnnl_handler, convert_mode_to_logicOp(), logic_res_desc.desc(), |
| 556 | src0_ptr, src1_desc.desc(), src1.raw_ptr(), logic_wk.raw_ptr, |
| 557 | logic_wk.size, logic_res_desc.desc(), logic_res_wk.raw_ptr)); |
| 558 | if (dtype_dest == megdnn::DTypeEnum::Int32) { |
| 559 | opTensorRun<int32_t>( |
| 560 | cnnl_handler, cnnlOpTensorDesc_t::CNNL_OP_TENSOR_MUL, |
| 561 | logic_res_desc, logic_res_wk.raw_ptr, src2_desc, src2.raw_ptr(), |
| 562 | output_desc, dst.raw_ptr(), dst.layout.dtype.enumv(), |
| 563 | optensor_wk); |
| 564 | } else { |
| 565 | opTensorRun<float>( |
no test coverage detected