| 32 | |
| 33 | template <typename T, int BITS> |
| 34 | void gemv_4bit_inference( |
| 35 | int m, int n, int k, T* A, unsigned char* B, float* absmax, float* datatype, T* out, int lda, int ldb, int ldc, |
| 36 | int blocksize, sycl::queue* stream |
| 37 | ) { |
| 38 | |
| 39 | auto& queue = *stream; |
| 40 | |
| 41 | const size_t GROUP_SIZE = 128; // workgroup_size |
| 42 | const size_t SUBG_SIZE = 32; // subgroup_size |
| 43 | const size_t NUM_PER_THREAD = GROUP_SIZE / SUBG_SIZE; |
| 44 | size_t workgroup_num = (n + NUM_PER_THREAD - 1) / NUM_PER_THREAD; |
| 45 | |
| 46 | kgemv_4bit_inference<T, GROUP_SIZE, NUM_PER_THREAD, SUBG_SIZE, BITS> kfn( |
| 47 | m, n, k, A, B, absmax, datatype, out, lda, ldb, ldc, blocksize |
| 48 | ); |
| 49 | |
| 50 | sycl_comp_kernel_submit<decltype(kfn), 1, SUBG_SIZE>( |
| 51 | sycl::nd_range<1>(sycl::range<1>(GROUP_SIZE * workgroup_num), sycl::range<1>(GROUP_SIZE)), queue, kfn |
| 52 | ); |
| 53 | } |
| 54 | |
| 55 | //============================================================== |
| 56 | // TEMPLATE DEFINITIONS |
nothing calls this directly
no outgoing calls
no test coverage detected