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

Function get_binary_ws

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

Source from the content-addressed store, hash-verified

698 return make_bundle(handle, sizes_in_bytes);
699}
700WorkspaceBundle get_binary_ws(
701 HandleImpl* handle, const TensorLayoutArray& src, const TensorLayout& dst,
702 const param::Elemwise::Mode& mode) {
703 auto cnnl_handler = handle->cnnl_handle();
704 bool mode_need_ws = mode == Mode::TRUE_DIV || mode == Mode::FLOOR_DIV ||
705 mode == Mode::MOD || mode == Mode::MAX || mode == Mode::MIN ||
706 mode == Mode::POW || mode == Mode::AND || mode == Mode::ADD ||
707 mode == Mode::SUB || mode == Mode::MUL ||
708 mode == Mode::SWITCH_GT0 || mode == Mode::SOFTPLUS_GRAD ||
709 mode == Mode::SIGMOID_GRAD;
710 if (!mode_need_ws)
711 return {nullptr, {}, handle->alignment_requirement()};
712 CnnlTensorDescriptor lhs_desc, rhs_desc, output_desc;
713 lhs_desc.set(src[0]);
714 rhs_desc.set(src[1]);
715 output_desc.set(dst);
716 // 1st workspace is cnnl workspace, 2st and 3st are handle un-contiguous
717 SmallVector<size_t> sizes_in_bytes(3, 0);
718 // ADD,SUB,MUL,MOD,FLOOR_DIV,MIN,MAX,POW,SWITCH_GT0 need handle uncontig
719 auto handle_uncontig_wk = [&]() {
720 size_t lhs_wk = !src[0].is_contiguous() ? src[0].access_bytes() : 0;
721 size_t rhs_wk = !src[1].is_contiguous() ? src[1].access_bytes() : 0;
722 sizes_in_bytes[1] = lhs_wk;
723 sizes_in_bytes[2] = rhs_wk;
724 };
725
726 switch (mode) {
727 case Mode::TRUE_DIV:
728 cnnl_check(cnnlGetDivWorkspaceSize(
729 cnnl_handler, lhs_desc.desc(), rhs_desc.desc(), output_desc.desc(),
730 &sizes_in_bytes[0]));
731 break;
732 case Mode::FLOOR_DIV:
733 cnnl_check(cnnlGetFloorDivWorkspaceSize(
734 cnnl_handler, lhs_desc.desc(), rhs_desc.desc(), output_desc.desc(),
735 &sizes_in_bytes[0]));
736 handle_uncontig_wk();
737 break;
738 case Mode::MOD:
739 if (dst.dtype.enumv() == megdnn::DTypeEnum::Int32) {
740 cnnl_check(cnnlGetFloorModWorkspaceSize(
741 cnnl_handler, lhs_desc.desc(), rhs_desc.desc(),
742 output_desc.desc(), &sizes_in_bytes[0]));
743 } else {
744 cnnl_check(cnnlGetFloorModTruncWorkspaceSize(
745 cnnl_handler, lhs_desc.desc(), rhs_desc.desc(),
746 output_desc.desc(), &sizes_in_bytes[0]));
747 }
748 handle_uncontig_wk();
749 break;
750 case Mode::POW:
751 cnnl_check(cnnlGetPowWorkspaceSize(
752 cnnl_handler, lhs_desc.desc(), rhs_desc.desc(), output_desc.desc(),
753 &sizes_in_bytes[0]));
754 handle_uncontig_wk();
755 break;
756 case Mode::AND:
757 cnnl_check(cnnlGetLogicOpWorkspaceSize(

Callers 1

alloc_cnnl_workspaceMethod · 0.85

Calls 10

make_bundleFunction · 0.85
cnnl_handleMethod · 0.80
access_bytesMethod · 0.80
eq_layoutMethod · 0.80
alignment_requirementMethod · 0.45
setMethod · 0.45
is_contiguousMethod · 0.45
descMethod · 0.45
enumvMethod · 0.45
push_backMethod · 0.45

Tested by

no test coverage detected