#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);
        }
    }
}

};  // namespace cpu_inference