| 792 | } |
| 793 | |
| 794 | WorkspaceBundle 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; |
no test coverage detected