| 146 | } |
| 147 | |
| 148 | torch::Tensor code1x16_matmat( |
| 149 | const torch::Tensor& input, |
| 150 | const torch::Tensor& codes, |
| 151 | const torch::Tensor& codebooks, |
| 152 | const torch::Tensor& scales, |
| 153 | const std::optional<torch::Tensor>& bias |
| 154 | ) { |
| 155 | bool use_bfloat16 = check_use_bfloat16(input); |
| 156 | auto input_sizes = input.sizes(); |
| 157 | auto out_features = codes.size(0) * codebooks.size(2); |
| 158 | auto flat_input = input.reshape({-1, input.size(-1)}); |
| 159 | auto flat_output = torch::empty({flat_input.size(0), out_features}, |
| 160 | torch::TensorOptions() |
| 161 | .dtype(input.dtype()) |
| 162 | .device(input.device()) |
| 163 | ); |
| 164 | |
| 165 | for (int i = 0; i < flat_input.size(0); ++i) { |
| 166 | auto input_vec = flat_input.index({i}); |
| 167 | auto output_vec = flat_output.index({i}); |
| 168 | code1x16_matvec( |
| 169 | codes.squeeze(2), |
| 170 | input_vec, |
| 171 | output_vec, |
| 172 | codebooks, |
| 173 | use_bfloat16 |
| 174 | ); |
| 175 | } |
| 176 | return scale_bias_unflatten_output( |
| 177 | flat_output, |
| 178 | scales, |
| 179 | bias, |
| 180 | input_sizes |
| 181 | ); |
| 182 | } |
| 183 | |
| 184 | torch::Tensor code1x16_dequant( |
| 185 | const torch::Tensor& codes, |
nothing calls this directly
no test coverage detected