#include "quantization_q4_0.h"

#ifdef __ARM_NEON
#include <arm_neon.h>
#endif

#ifdef __ARM_FEATURE_SVE
#include <arm_sve.h>
#endif

namespace cpu_inference {

namespace {

constexpr uint8_t clamp_u4(int v) {
    return (v < 0) ? 0 : (v > 15) ? 15 : static_cast<uint8_t>(v);
}

};

void QuantizationBase<block_q4_0>::quantize(
    const void* const src_data, block_q4_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_q4_0>::quantize_from_fp16(
    const float16_t* const src_data, block_q4_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_q4_0>::quantize_from_fp16_neon(
    const float16_t* const src_data, block_q4_0* dst_data, size_t element_cnt)
{
    static const int qk = QK4_0;
    const int n_blocks = element_cnt / qk;
    for (int i = 0; i < n_blocks; i++) {
        float16x8_t src[4];                     // 加载32个float32到8个NEON寄存器,每个寄存器加载4个float32
        for (int j = 0; j < 4; j++) {
            src[j] = vld1q_f16(src_data + i*qk + 8*j);      // 处理第i个block,每个向量寄存器加载4个float32
        }

        float16x8_t maxv = src[0];              // 计算最大值和最小值
        float16x8_t minv = src[0];
        for (int j = 1; j < 4; j++) {
            maxv = vmaxq_f16(maxv, src[j]);
            minv = vminq_f16(minv, src[j]);
        }

        float max_val = fp16_to_fp32(vmaxvq_f16(maxv));
        float min_val = fp16_to_fp32(vminvq_f16(minv));

        // 确定绝对值最大值及其符号
        float selected_val = std::abs(max_val) > std::abs(min_val) ? max_val : min_val;

        // 计算量化参数
        const float d = selected_val / -8;
        const float id = 1.0f / d;

        dst_data[i].d = fp32_to_fp16(d);

        // 准备常量和结果寄存器
        const float32x4_t vid = vdupq_n_f32(id);
        const float32x4_t v8_5 = vdupq_n_f32(8.5f);
        uint8x16_t result;

        float32x4_t src1[8];
        for (int j = 0; j < 4; j++){
            src1[2*j  ] = vcvt_f32_f16(vget_low_f16 (src[j]));
            src1[2*j+1] = vcvt_f32_f16(vget_high_f16(src[j]));
        }
        // 处理前16个元素 (0-15)
        int32x4_t i0 = vcvtq_s32_f32(vaddq_f32(vmulq_f32(src1[0], vid), v8_5));
        int32x4_t i1 = vcvtq_s32_f32(vaddq_f32(vmulq_f32(src1[1], vid), v8_5));
        int32x4_t i2 = vcvtq_s32_f32(vaddq_f32(vmulq_f32(src1[2], vid), v8_5));
        int32x4_t i3 = vcvtq_s32_f32(vaddq_f32(vmulq_f32(src1[3], vid), v8_5));
        // 处理后16个元素 (16-31)
        int32x4_t i4 = vcvtq_s32_f32(vaddq_f32(vmulq_f32(src1[4], vid), v8_5));
        int32x4_t i5 = vcvtq_s32_f32(vaddq_f32(vmulq_f32(src1[5], vid), v8_5));
        int32x4_t i6 = vcvtq_s32_f32(vaddq_f32(vmulq_f32(src1[6], vid), v8_5));
        int32x4_t i7 = vcvtq_s32_f32(vaddq_f32(vmulq_f32(src1[7], vid), v8_5));

        // 将结果饱和到[0, 15]范围并打包
        uint8x8_t low = vqmovun_s16(vcombine_s16(
            vqmovn_s32(vmaxq_s32(vminq_s32(i0, vdupq_n_s32(15)), vdupq_n_s32(0))),
            vqmovn_s32(vmaxq_s32(vminq_s32(i1, vdupq_n_s32(15)), vdupq_n_s32(0)))
        ));
        uint8x8_t high = vqmovun_s16(vcombine_s16(
            vqmovn_s32(vmaxq_s32(vminq_s32(i2, vdupq_n_s32(15)), vdupq_n_s32(0))),
            vqmovn_s32(vmaxq_s32(vminq_s32(i3, vdupq_n_s32(15)), vdupq_n_s32(0)))
        ));
        result = vcombine_u8(low, high);

        low = vqmovun_s16(vcombine_s16(
            vqmovn_s32(vmaxq_s32(vminq_s32(i4, vdupq_n_s32(15)), vdupq_n_s32(0))),
            vqmovn_s32(vmaxq_s32(vminq_s32(i5, vdupq_n_s32(15)), vdupq_n_s32(0)))
        ));
        high = vqmovun_s16(vcombine_s16(
            vqmovn_s32(vmaxq_s32(vminq_s32(i6, vdupq_n_s32(15)), vdupq_n_s32(0))),
            vqmovn_s32(vmaxq_s32(vminq_s32(i7, vdupq_n_s32(15)), vdupq_n_s32(0)))
        ));
        uint8x16_t high_part = vcombine_u8(low, high);

        // 合并高低4位数据
        result = vorrq_u8(result, vshlq_n_u8(high_part, 4));

        // 存储结果
        vst1q_u8(dst_data[i].qs, result);
    }
}

void QuantizationBase<block_q4_0>::quantize_from_fp16_default(
    const float16_t* const src_data, block_q4_0* dst_data, size_t element_cnt)
{
    static const int qk = QK4_0;
    const int n_blocks = element_cnt / qk;
    for (int i = 0; i < n_blocks; i++) {
        float amax = 0.0f; // absolute max
        float max  = 0.0f;

        for (int j = 0; j < qk; j++) {
            const float v = fp16_to_fp32(src_data[i*qk + j]);
            if (amax < std::abs(v)) {
                amax = std::abs(v);
                max  = v;
            }
        }

        const float d  = max / -8;
        const float id = d ? 1.0f/d : 0.0f;

        dst_data[i].d = fp32_to_fp16(d);

        for (int j = 0; j < qk/2; ++j) {
            const float x0 = src_data[i*qk + 0    + j]*id;
            const float x1 = src_data[i*qk + qk/2 + j]*id;

            const uint8_t xi0 = clamp_u4((int8_t)(x0 + 8.5f));
            const uint8_t xi1 = clamp_u4((int8_t)(x1 + 8.5f));

            dst_data[i].qs[j] = xi0;
            dst_data[i].qs[j] |= xi1 << 4;
        }
    }
}

void QuantizationBase<block_q4_0>::dequantize(
    const block_q4_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_q4_0>::dequantize_to_fp16(
    const block_q4_0* const src_data, float16_t* dst_data, size_t element_cnt)
{
    static const int qk = QK4_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 / 2; ++j) {
            const int x0 = (src_data[i].qs[j] & 0x0F) - 8;
            const int x1 = (src_data[i].qs[j] >> 4) - 8;

            dst_data[i*qk + j + 0   ] = fp32_to_fp16(x0*d);
            dst_data[i*qk + j + qk/2] = fp32_to_fp16(x1*d);
        }
    }
}

};