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

Method try_transpose

dnn/src/cuda/relayout/opr_impl.cpp:19–63  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

17}
18
19bool 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
65bool RelayoutForwardImpl::Param::try_copy_contig() {
66 auto &&lsrc = m_src.layout, &&ldst = m_dst.layout;

Callers 1

execMethod · 0.80

Calls 8

is_low_bitMethod · 0.80
cublas_handleMethod · 0.80
concrete_handleFunction · 0.50
sizeMethod · 0.45
handleMethod · 0.45
one_deviceMethod · 0.45
zero_deviceMethod · 0.45
streamMethod · 0.45

Tested by

no test coverage detected