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

Function exec_ternary

dnn/src/cambricon/elemwise/opr_impl.cpp:508–642  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

506}
507
508void 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>(

Callers 1

execMethod · 0.85

Calls 11

to_broadcastFunction · 0.85
to_contiguousFunction · 0.85
cnnl_handleMethod · 0.80
access_bytesMethod · 0.80
enumvMethod · 0.45
setMethod · 0.45
get_workspaceMethod · 0.45
push_backMethod · 0.45
descMethod · 0.45
raw_ptrMethod · 0.45
sizeMethod · 0.45

Tested by

no test coverage detected