| 111 | } |
| 112 | |
| 113 | void code1x16_matvec( |
| 114 | const torch::Tensor& A, |
| 115 | const torch::Tensor& B, |
| 116 | torch::Tensor& C, |
| 117 | const torch::Tensor& codebook, |
| 118 | const bool use_bfloat16 |
| 119 | ) { |
| 120 | const at::cuda::OptionalCUDAGuard device_guard(device_of(A)); |
| 121 | int prob_m = C.size(0); |
| 122 | int prob_k = B.size(0); |
| 123 | |
| 124 | if (codebook.size(3) == 8) { |
| 125 | if (use_bfloat16) { |
| 126 | code1x16_matvec_cuda<true, 8>(A.data_ptr(), B.data_ptr(), C.data_ptr(), codebook.data_ptr(), prob_m, prob_k); |
| 127 | } else { |
| 128 | code1x16_matvec_cuda<false, 8>(A.data_ptr(), B.data_ptr(), C.data_ptr(), codebook.data_ptr(), prob_m, prob_k); |
| 129 | } |
| 130 | } else if (codebook.size(3) == 16) { |
| 131 | if (use_bfloat16) { |
| 132 | code1x16_matvec_cuda<true, 16>(A.data_ptr(), B.data_ptr(), C.data_ptr(), codebook.data_ptr(), prob_m, prob_k); |
| 133 | } else { |
| 134 | code1x16_matvec_cuda<false, 16>(A.data_ptr(), B.data_ptr(), C.data_ptr(), codebook.data_ptr(), prob_m, prob_k); |
| 135 | } |
| 136 | } else { |
| 137 | throw c10::NotImplementedError( |
| 138 | {__func__, __FILE__, static_cast<uint32_t>(__LINE__)}, |
| 139 | c10::str( |
| 140 | "AQLM CUDA kernels only support codebooks with 8 or 16 features. Got ", |
| 141 | codebook.size(3), |
| 142 | "." |
| 143 | ) |
| 144 | ); |
| 145 | } |
| 146 | } |
| 147 | |
| 148 | torch::Tensor code1x16_matmat( |
| 149 | const torch::Tensor& input, |