| 494 | } |
| 495 | |
| 496 | void RelayoutForwardImpl::exec_after_preprocess( |
| 497 | const TensorND& src, const TensorND& dst, relayout::TransposeParam* transpose) { |
| 498 | if (transpose) { |
| 499 | bool is_bit4 = is_int4(src.layout); |
| 500 | auto kernel = [tparam = *transpose, src, dst, is_bit4]() { |
| 501 | auto t = tparam; |
| 502 | void (*kptr)(size_t, size_t, size_t, size_t, void*, void*, size_t) = |
| 503 | nullptr; |
| 504 | auto src_addr = reinterpret_cast<uintptr_t>(src.raw_ptr()), |
| 505 | dst_addr = reinterpret_cast<uintptr_t>(dst.raw_ptr()); |
| 506 | size_t dsize = 0; |
| 507 | if (is_bit4) { |
| 508 | dsize = t.c >> 1; |
| 509 | } else { |
| 510 | dsize = src.layout.dtype.size() * t.c; |
| 511 | } |
| 512 | if (is_bit4 && dsize == 0) { |
| 513 | kptr = call_transpose<dt_qint4>; |
| 514 | } else { |
| 515 | if (dsize == 1) { |
| 516 | megdnn_assert(t.c == 1); |
| 517 | kptr = call_transpose<uint8_t>; |
| 518 | } else if (dsize == 2) { |
| 519 | t.c = 1; |
| 520 | if (!((src_addr | dst_addr) & (alignof(uint16_t) - 1))) { |
| 521 | kptr = call_transpose<uint16_t>; |
| 522 | } else { |
| 523 | kptr = call_transpose<equiv_ctype_storage<2>>; |
| 524 | megdnn_log_error("unaligned addr in relayout"); |
| 525 | } |
| 526 | } else if (dsize == 3) { |
| 527 | t.c = 1; |
| 528 | kptr = call_transpose<equiv_ctype_storage<3>>; |
| 529 | } else if (dsize == 4) { |
| 530 | t.c = 1; |
| 531 | if (!((src_addr | dst_addr) & (alignof(uint32_t) - 1))) { |
| 532 | kptr = call_transpose<uint32_t>; |
| 533 | } else { |
| 534 | kptr = call_transpose<equiv_ctype_storage<4>>; |
| 535 | megdnn_log_error("unaligned addr in relayout"); |
| 536 | } |
| 537 | } else if (dsize == 12) { |
| 538 | t.c = 1; |
| 539 | if (!((src_addr | dst_addr) & (alignof(uint32_t) - 1))) { |
| 540 | kptr = call_transpose<equiv_ctype_storage<3, uint32_t>>; |
| 541 | } else { |
| 542 | kptr = call_transpose<equiv_ctype_storage<12>>; |
| 543 | megdnn_log_error("unaligned addr in relayout"); |
| 544 | } |
| 545 | } else if (dsize <= TRANSPOSE_CV_MAX_C) { |
| 546 | switch (dst.layout.dtype.enumv()) { |
| 547 | #define cb(_dt) \ |
| 548 | case DTypeTrait<dtype::_dt>::enumv: \ |
| 549 | kptr = transpose_cv<equiv_ctype<dtype::_dt>::type>; \ |
| 550 | break; |
| 551 | MEGDNN_FOREACH_DTYPE_NAME(cb) |
| 552 | MEGDNN_FOREACH_PARAMETERIZED_DTYPE(cb) |
| 553 | #undef cb |
nothing calls this directly
no test coverage detected