#include "quantization_q8_0.h"
#include <cmath>
namespace cpu_inference {
void QuantizationBase<block_q8_0>::quantize(
const void* const src_data, block_q8_0* dst_data, size_t element_cnt, std::string ori_dtype)
{
if (ori_dtype == "fp16") {
const float16_t* const src_data_fp16 = static_cast<const float16_t*>(src_data);
quantize_from_fp16(src_data_fp16, dst_data, element_cnt);
} else {
std::cerr << "Unsupported original data type: " << ori_dtype << std::endl;
}
}
void QuantizationBase<block_q8_0>::quantize_from_fp16(
const float16_t* const src_data, block_q8_0* dst_data, size_t element_cnt)
{
#if defined(__ARM_NEON)
quantize_from_fp16_neon(src_data, dst_data, element_cnt);
#else
quantize_from_fp16_default(src_data, dst_data, element_cnt);
#endif
}
void QuantizationBase<block_q8_0>::quantize_from_fp16_neon(
const float16_t* const src_data, block_q8_0* dst_data, size_t element_cnt)
{
static const int qk = QK8_0;
const int n_blocks = element_cnt / qk;
for (int i = 0; i < n_blocks; i++) {
float16x8_t srcv [4];
float16x8_t asrcv[4];
float16x8_t amaxv[4];
for (int j = 0; j < 4; j++) srcv[j] = vld1q_f16(src_data + i*qk + 8*j);
for (int j = 0; j < 4; j++) asrcv[j] = vabsq_f16(srcv[j]);
for (int j = 0; j < 2; j++) amaxv[2*j] = vmaxq_f16(asrcv[2*j], asrcv[2*j+1]);
for (int j = 0; j < 1; j++) amaxv[4*j] = vmaxq_f16(amaxv[4*j], amaxv[4*j+2]);
const float amax = (float)vmaxvq_f16(amaxv[0]);
const float d = amax / ((1 << 7) - 1);
const float id = d ? 1.0f/d : 0.0f;
dst_data[i].d = fp32_to_fp16(d);
for (int j = 0; j < 4; j++) {
const float16x8_t v = vmulq_n_f16(srcv[j], id);
const int16x8_t vi = vcvtnq_s16_f16(v);
const int8x8_t vi8 = vqmovn_s16(vi);
vst1_s8(&dst_data[i].qs[8*j], vi8);
}
}
}
void QuantizationBase<block_q8_0>::quantize_from_fp16_default(
const float16_t* const src_data, block_q8_0* dst_data, size_t element_cnt)
{
static const int qk = QK8_0;
const int n_blocks = element_cnt / qk;
for (int i = 0; i < n_blocks; i++) {
float amax = 0.0f;
for (int j = 0; j < qk; j++) {
const float v = fp16_to_fp32(src_data[i*qk + j]);
amax = std::max(amax, std::abs(v));
}
const float d = amax / ((1 << 7) - 1);
const float id = d ? 1.0f/d : 0.0f;
dst_data[i].d = fp32_to_fp16(d);
for (int j = 0; j < qk; ++j) {
const float x0 = src_data[i*qk + j]*id;
dst_data[i].qs[j] = std::round(x0);
}
}
}
void QuantizationBase<block_q8_0>::dequantize(
const block_q8_0* const src_data, void* dst_data, size_t element_cnt, std::string aim_dtype)
{
if(aim_dtype == "fp16") {
float16_t* dst_data_fp16 = static_cast<float16_t*>(dst_data);
dequantize_to_fp16(src_data, dst_data_fp16, element_cnt);
} else {
std::cerr << "Unsupported aim data type: " << aim_dtype << std::endl;
}
}
void QuantizationBase<block_q8_0>::dequantize_to_fp16(
const block_q8_0* const src_data, float16_t* dst_data, size_t element_cnt)
{
static const int qk = QK8_0;
const int n_blocks = element_cnt / qk;
for (int i = 0; i < n_blocks; i++) {
const float d = fp16_to_fp32(src_data[i].d);
for (int j = 0; j < qk; ++j) {
dst_data[i*qk + j] = fp32_to_fp16(src_data[i].qs[j]*d);
}
}
}
};