#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; // absolute max

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

};  // namespace cpu_inference