#include "quantization_fp16.h"

namespace cpu_inference {

void QuantizationBase<float16_t>::quantize(
    const void* const src_data, float16_t* dst_data, size_t element_cnt, std::string ori_dtype)
{
    if (ori_dtype == "fp16") {
        std::memcpy(dst_data, src_data, element_cnt * sizeof(float16_t));
    } else if (ori_dtype == "fp32") {
        const float32_t* const src_data_fp32 = static_cast<const float32_t*>(src_data);
        for (size_t i = 0; i < element_cnt; i++) {
            dst_data[i] = fp32_to_fp16(src_data_fp32[i]);
        }
    } else {
        std::cerr << "Unsupported original data type: " << ori_dtype << std::endl;
    }
}

void QuantizationBase<float16_t>::dequantize(
    const float16_t* const src_data, void* dst_data, size_t element_cnt, std::string aim_dtype)
{
    if (aim_dtype == "fp16") {
        std::memcpy(dst_data, src_data, element_cnt * sizeof(float16_t));
    } else if (aim_dtype == "fp32") {
        float32_t* dst_data_fp32 = static_cast<float32_t*>(dst_data);
        for (size_t i = 0; i < element_cnt; i++) {
            dst_data_fp32[i] = fp16_to_fp32(src_data[i]);
        }
    } else {
        std::cerr << "Unsupported aim data type: " << aim_dtype << std::endl;
    }
}

};  // namespace cpu_inference