#include "quantization_q8align.h"
#include <cmath>
namespace cpu_inference {
void Q8Align::q8align_quantize(
const float16_t* const base_src_data, int8_t* base_dst_data_val, float* base_dst_data_d,
size_t block_id, size_t block_cnt, size_t block_in_row, bool is_input
) {
static const int elem_in_merge_block = elem_cnt_in_block * 2;
int block_end = block_id + block_cnt;
for (int i = block_id; i < block_end; i++) {
int merge_block_id = i / block_in_row / 2 * block_in_row + (i % block_in_row);
const float16_t* const src_data = base_src_data + i * elem_cnt_in_block;
int8_t* dst_data_val = base_dst_data_val + merge_block_id * elem_in_merge_block;
float* dst_data_d = base_dst_data_d + merge_block_id * 4;
int re_id = (i / block_in_row) % 2;
if(re_id == 1){
dst_data_val = dst_data_val + 8;
if(is_input){
dst_data_d = dst_data_d + 2;
}
else {
dst_data_d = dst_data_d + 1;
}
}
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 + 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;
if(is_input){
dst_data_d[0] = dst_data_d[1] = (float16_t)(d);
}
else {
dst_data_d[0] = dst_data_d[2] = (float16_t)(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_val + 16 * j, vi8);
}
}
}
};