| 698 | return make_bundle(handle, sizes_in_bytes); |
| 699 | } |
| 700 | WorkspaceBundle 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( |
no test coverage detected