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

Function get_ternary_ws

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

Source from the content-addressed store, hash-verified

792}
793
794WorkspaceBundle get_ternary_ws(
795 HandleImpl* handle, const TensorLayoutArray& src, const TensorLayout& dst,
796 const param::Elemwise::Mode& mode) {
797 auto cnnl_handler = handle->cnnl_handle();
798 bool mode_need_ws = mode == Mode::CLIP || mode == Mode::COND_LEQ_MOV ||
799 mode == Mode::COND_LT_MOV || mode == Mode::FUSE_MUL_ADD3;
800 if (!mode_need_ws)
801 return {nullptr, {}, handle->alignment_requirement()};
802 CnnlTensorDescriptor src0_desc, src1_desc, src2_desc, output_desc;
803 src0_desc.set(src[0]);
804 src1_desc.set(src[1]);
805 src2_desc.set(src[2]);
806 output_desc.set(dst);
807 // 1st workspace is cnnl workspace
808 SmallVector<size_t> sizes_in_bytes(1, 0);
809 switch (mode) {
810 case Mode::CLIP: {
811 size_t dtype_size = src[0].dtype.size(1);
812 size_t src_wk =
813 !src[0].is_contiguous() ? src[0].total_nr_elems() * dtype_size : 0;
814 sizes_in_bytes.push_back(src_wk);
815 break;
816 }
817 case Mode::COND_LT_MOV:
818 case Mode::COND_LEQ_MOV: {
819 TensorShapeArray src0_1;
820 src0_1.push_back(src[0]);
821 src0_1.push_back(src[1]);
822 TensorShape logic_res_shape;
823 Elemwise::deduce_shape(src0_1, logic_res_shape);
824 TensorLayout logic_res_layout(logic_res_shape, src[0].dtype);
825 CnnlTensorDescriptor logic_res_desc;
826 logic_res_desc.set(logic_res_layout);
827 cnnl_check(cnnlGetLogicOpWorkspaceSize(
828 cnnl_handler, logic_res_desc.desc(), src1_desc.desc(),
829 logic_res_desc.desc(), &sizes_in_bytes[0]));
830 size_t src0_wk = 0, logic_res_wk = 0;
831 if (!src[0].eq_layout(dst)) {
832 src0_wk = logic_res_layout.access_bytes();
833 }
834 logic_res_wk = logic_res_layout.access_bytes();
835 sizes_in_bytes.push_back(src0_wk);
836 sizes_in_bytes.push_back(logic_res_wk);
837 size_t optensor_wk = 0;
838 cnnl_check(cnnlGetOpTensorWorkspaceSize(
839 cnnl_handler, logic_res_desc.desc(), src2_desc.desc(),
840 output_desc.desc(), &optensor_wk));
841 sizes_in_bytes.push_back(optensor_wk);
842 break;
843 }
844 case Mode::FUSE_MUL_ADD3: {
845 TensorShapeArray src0_1;
846 src0_1.push_back(src[0]);
847 src0_1.push_back(src[1]);
848 TensorShape mul_res_shape;
849 Elemwise::deduce_shape(src0_1, mul_res_shape);
850 TensorLayout mul_res_layout(mul_res_shape, src[0].dtype);
851 CnnlTensorDescriptor mul_res_desc;

Callers 1

alloc_cnnl_workspaceMethod · 0.85

Calls 11

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

Tested by

no test coverage detected