/* * Copyright (c) Meta Platforms, Inc. and affiliates. * All rights reserved. * * This source code is licensed under the BSD-style license found in the * LICENSE file in the root directory of this source tree. */ #include #include #include #include #include #include namespace impl { namespace reference { namespace kernels { // Quantize a fp32 value to an int8_t/uint8_t value template T quantize(const float x, float scale, int32_t zero_point) { constexpr float min_val = std::numeric_limits::min(); constexpr float max_val = std::numeric_limits::max(); float tmp = roundf(x * scale + zero_point); return std::max(std::min(tmp, max_val), min_val); } // Quantize an fp32 array to an int8_t/uint8_t array template void quantize( T* __restrict__ y, const float* __restrict__ x, float inv_scale, int32_t zero_point, size_t size) { for (size_t i = 0; i < size; ++i) { y[i] = quantize(x[i], inv_scale, zero_point); } } // Dequantize an int8_t/uint8_t value to an fp32 value template float dequantize(const T x, float scale, int32_t zero_point) { return scale * (x - zero_point); } // Dequantize an int8_t/uint8_t/int16_t array to an fp32 array template void dequantize( float* __restrict__ y, const T* __restrict__ x, float scale, int32_t zero_point, size_t size) { for (size_t i = 0; i < size; ++i) { y[i] = dequantize(x[i], scale, zero_point); } } // Requantize the int8_t/uint8_t in value to a uint8_t/int8_t out value. // The scale and zero_point for requantization are in the args. template OT requantize( const IT in, float in_scale, int32_t in_zero_point, float inv_out_scale, int32_t out_zero_point) { float dequant = dequantize(in, in_scale, in_zero_point); return quantize(dequant, inv_out_scale, out_zero_point); } // Requantize the int8_t/uint8_t in array to a uint8_t/int8_t out array. // The scale and zero_point for requantization are in the args. template void requantize( OT* __restrict__ out, const IT* __restrict__ in, float in_scale, int32_t in_zero_point, float inv_out_scale, int32_t out_zero_point, size_t size) { for (size_t i = 0; i < size; ++i) { out[i] = requantize( in[i], in_scale, in_zero_point, inv_out_scale, out_zero_point); } } // explicit template instantiation #define typed_quantize_val(dtype) \ template dtype quantize(const float x, float inv_scale, int32_t zero_point); typed_quantize_val(int8_t); typed_quantize_val(uint8_t); typed_quantize_val(int16_t); typed_quantize_val(uint16_t); typed_quantize_val(int32_t); #undef typed_quantize_val #define typed_quantize_vec(dtype) \ template void quantize( \ dtype* __restrict__ y, \ const float* __restrict__ x, \ float inv_scale, \ int32_t zero_point, \ size_t size); typed_quantize_vec(int8_t); typed_quantize_vec(uint8_t); typed_quantize_vec(int16_t); typed_quantize_vec(uint16_t); typed_quantize_vec(int32_t); #undef typed_quantize_vec #define typed_dequantize_val(dtype) \ template float dequantize(const dtype x, float scale, int32_t zero_point); typed_dequantize_val(int8_t); typed_dequantize_val(uint8_t); typed_dequantize_val(int16_t); typed_dequantize_val(uint16_t); typed_dequantize_val(int32_t); #undef typed_dequantize_val #define typed_dequantize_vec(dtype) \ template void dequantize( \ float* __restrict__ y, \ const dtype* __restrict__ x, \ float scale, \ int32_t zero_point, \ size_t size); typed_dequantize_vec(int8_t); typed_dequantize_vec(uint8_t); typed_dequantize_vec(int16_t); typed_dequantize_vec(uint16_t); typed_dequantize_vec(int32_t); #undef typed_dequantize_vec #define typed_requantize_val(itype, otype) \ template otype requantize( \ const itype in, \ float in_scale, \ int32_t in_zero_point, \ float inv_out_scale, \ int32_t out_zero_point); typed_requantize_val(int8_t, int8_t); typed_requantize_val(int8_t, uint8_t); typed_requantize_val(int8_t, int16_t); typed_requantize_val(int8_t, uint16_t); typed_requantize_val(uint8_t, int8_t); typed_requantize_val(uint8_t, uint8_t); typed_requantize_val(uint8_t, int16_t); typed_requantize_val(uint8_t, uint16_t); typed_requantize_val(int16_t, int8_t); typed_requantize_val(int16_t, uint8_t); typed_requantize_val(int16_t, int16_t); typed_requantize_val(int16_t, uint16_t); typed_requantize_val(uint16_t, int8_t); typed_requantize_val(uint16_t, uint8_t); typed_requantize_val(uint16_t, int16_t); typed_requantize_val(uint16_t, uint16_t); #undef typed_requantize_val #define typed_requantize_vec(itype, otype) \ template void requantize( \ otype* __restrict__ out, \ const itype* __restrict__ in, \ float in_scale, \ int32_t in_zero_point, \ float inv_out_scale, \ int32_t out_zero_point, \ size_t size); typed_requantize_vec(int8_t, int8_t); typed_requantize_vec(int8_t, uint8_t); typed_requantize_vec(int8_t, int16_t); typed_requantize_vec(int8_t, uint16_t); typed_requantize_vec(uint8_t, int8_t); typed_requantize_vec(uint8_t, uint8_t); typed_requantize_vec(uint8_t, int16_t); typed_requantize_vec(uint8_t, uint16_t); typed_requantize_vec(int16_t, int8_t); typed_requantize_vec(int16_t, uint8_t); typed_requantize_vec(int16_t, int16_t); typed_requantize_vec(int16_t, uint16_t); typed_requantize_vec(uint16_t, int8_t); typed_requantize_vec(uint16_t, uint8_t); typed_requantize_vec(uint16_t, int16_t); typed_requantize_vec(uint16_t, uint16_t); #undef typed_requantize_vec }; // namespace kernels }; // namespace reference }; // namespace impl