| 17 | } |
| 18 | |
| 19 | bool RelayoutForwardImpl::Param::try_transpose() { |
| 20 | if (m_dst.layout.dtype.is_low_bit()) |
| 21 | return false; |
| 22 | relayout::TransposeParam transp; |
| 23 | bool trans = relayout::is_transpose(m_src.layout, m_dst.layout, transp); |
| 24 | if (!trans) |
| 25 | return false; |
| 26 | size_t dsize = transp.c * m_src.layout.dtype.size(); |
| 27 | if (dsize != 1 && dsize != 2 && dsize != 4) |
| 28 | return false; |
| 29 | |
| 30 | if (m_src.layout.dtype == dtype::Float32() && transp.batch == 1 && transp.c == 1) { |
| 31 | auto handle = concrete_handle(m_opr->handle()); |
| 32 | cublas_check(cublasSgeam( |
| 33 | handle->cublas_handle(), CUBLAS_OP_T, CUBLAS_OP_T, transp.m, transp.n, |
| 34 | handle->one_device(), m_src.ptr<dt_float32>(), transp.n, |
| 35 | handle->zero_device(), m_src.ptr<dt_float32>(), transp.n, |
| 36 | m_dst.ptr<dt_float32>(), transp.m)); |
| 37 | return true; |
| 38 | } |
| 39 | float square_ratio = static_cast<float>(transp.m) / static_cast<float>(transp.n); |
| 40 | if (transp.m < 32 || transp.n < 32 || square_ratio < 0.5f || square_ratio > 2.f) |
| 41 | return false; |
| 42 | size_t batch = transp.batch, m = transp.m, n = transp.n; |
| 43 | size_t lda = n, ldb = m, stride_A = m * n, stride_B = m * n; |
| 44 | auto&& stream = m_opr->stream(); |
| 45 | #define RUN(_dt) \ |
| 46 | do { \ |
| 47 | typedef DTypeTrait<dtype::_dt>::ctype ctype; \ |
| 48 | copy_by_transpose<ctype>( \ |
| 49 | reinterpret_cast<const ctype*>(m_src.raw_ptr()), \ |
| 50 | reinterpret_cast<ctype*>(m_dst.raw_ptr()), batch, m, n, lda, ldb, \ |
| 51 | stride_A, stride_B, stream); \ |
| 52 | return true; \ |
| 53 | } while (0) |
| 54 | switch (dsize) { |
| 55 | case 1: |
| 56 | RUN(Int8); |
| 57 | case 2: |
| 58 | RUN(Float16); |
| 59 | case 4: |
| 60 | RUN(Int32); |
| 61 | } |
| 62 | megdnn_assert(0, "bad dtype size"); |
| 63 | } |
| 64 | |
| 65 | bool RelayoutForwardImpl::Param::try_copy_contig() { |
| 66 | auto &&lsrc = m_src.layout, &&ldst = m_dst.layout; |
no test coverage detected