| 69 | } |
| 70 | |
| 71 | bool try_transpose( |
| 72 | _megdnn_tensor_in src, _megdnn_tensor_out dst, RelayoutForwardImpl* opr, |
| 73 | bool cross_dev) { |
| 74 | if (cross_dev) { |
| 75 | return false; |
| 76 | } |
| 77 | |
| 78 | auto&& src_layout = src.layout; |
| 79 | auto&& dst_layout = dst.layout; |
| 80 | |
| 81 | std::vector<size_t> re_permute, permute; |
| 82 | if (check_tensor_for_transpose(src_layout, dst_layout, re_permute)) { |
| 83 | permute.resize(re_permute.size()); |
| 84 | for (size_t i = 0; i < re_permute.size(); ++i) { |
| 85 | permute[re_permute[i]] = i; |
| 86 | } |
| 87 | } else { |
| 88 | return false; |
| 89 | } |
| 90 | |
| 91 | if (!check_param_transpose(src.layout, dst.layout, permute)) { |
| 92 | return false; |
| 93 | } |
| 94 | |
| 95 | size_t ndim = static_cast<size_t>(src_layout.ndim); |
| 96 | CnnlTransposeDescriptor cnnl_transpose_dsc; |
| 97 | cnnl_transpose_dsc.set(ndim, permute.data()); |
| 98 | |
| 99 | auto handle = concrete_handle(opr->handle()); |
| 100 | CnnlTensorDescriptor cnnl_src_dsc, cnnl_dst_dsc; |
| 101 | TensorLayout origin_src_layout = src_layout.dimshuffle(re_permute); |
| 102 | cnnl_src_dsc.set(origin_src_layout); |
| 103 | cnnl_dst_dsc.set(dst_layout); |
| 104 | cnnl_check(cnnlTranspose( |
| 105 | handle->cnnl_handle(), cnnl_transpose_dsc.desc(), cnnl_src_dsc.desc(), |
| 106 | src.raw_ptr(), cnnl_dst_dsc.desc(), dst.raw_ptr())); |
| 107 | return true; |
| 108 | } |
| 109 | |
| 110 | bool try_boradcast( |
| 111 | _megdnn_tensor_in src, _megdnn_tensor_out dst, RelayoutForwardImpl* opr, |
no test coverage detected