| 807 | |
| 808 | template <typename scalar_t, int DATA_TYPE> |
| 809 | void tinygemm_kernel( |
| 810 | const scalar_t* __restrict__ A, const unsigned char* __restrict__ B, scalar_t* __restrict__ C, |
| 811 | const scalar_t* __restrict__ Bs, scalar_t* __restrict__ Btmp, float* __restrict__ Ctmp, int64_t M, int64_t N, |
| 812 | int64_t K, int group_size, int64_t lda, int64_t ldb, int64_t ldc, int64_t strideBz, int64_t strideBs |
| 813 | ) { |
| 814 | constexpr int64_t BLOCK_M = 4; |
| 815 | constexpr int64_t BLOCK_N = 64; |
| 816 | const int64_t MB = div_up(M, BLOCK_M); |
| 817 | const int64_t NB = div_up(N, BLOCK_N); |
| 818 | for (int mb = 0; mb < MB; ++mb) { |
| 819 | int64_t mb_start = mb * BLOCK_M; |
| 820 | int64_t mb_size = std::min(BLOCK_M, M - mb_start); |
| 821 | for (int64_t nb = 0; nb < NB; ++nb) { |
| 822 | int64_t nb_start = nb * BLOCK_N; |
| 823 | int64_t nb_size = std::min(BLOCK_N, N - nb_start); |
| 824 | |
| 825 | switch (mb_size << 4 | nb_size >> 4) { |
| 826 | // mb_size = 1 |
| 827 | case 0x12: |
| 828 | LAUNCH_TINYGEMM_KERNEL_NN(1, 32, DATA_TYPE); |
| 829 | break; |
| 830 | case 0x14: |
| 831 | LAUNCH_TINYGEMM_KERNEL_NN(1, 64, DATA_TYPE); |
| 832 | break; |
| 833 | // mb_size = 2 |
| 834 | case 0x22: |
| 835 | LAUNCH_TINYGEMM_KERNEL_NN(2, 32, DATA_TYPE); |
| 836 | break; |
| 837 | case 0x24: |
| 838 | LAUNCH_TINYGEMM_KERNEL_NN(2, 64, DATA_TYPE); |
| 839 | break; |
| 840 | // mb_size = 3 |
| 841 | case 0x32: |
| 842 | LAUNCH_TINYGEMM_KERNEL_NN(3, 32, DATA_TYPE); |
| 843 | break; |
| 844 | case 0x34: |
| 845 | LAUNCH_TINYGEMM_KERNEL_NN(3, 64, DATA_TYPE); |
| 846 | break; |
| 847 | // mb_size = 4 |
| 848 | case 0x42: |
| 849 | LAUNCH_TINYGEMM_KERNEL_NN(4, 32, DATA_TYPE); |
| 850 | break; |
| 851 | case 0x44: |
| 852 | LAUNCH_TINYGEMM_KERNEL_NN(4, 64, DATA_TYPE); |
| 853 | break; |
| 854 | default: { |
| 855 | std::fprintf( |
| 856 | stderr, "[bitsandbytes] Unexpected block size %lldx%lld\n", (long long)mb_size, (long long)nb_size |
| 857 | ); |
| 858 | std::abort(); // or return; if you prefer silent exit |
| 859 | } |
| 860 | } |
| 861 | } |
| 862 | } |
| 863 | } |
| 864 | |
| 865 | template <typename T, int DATA_TYPE> |
| 866 | void gemv_4bit_inference( |