#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];
for (int j = 0; j < 4; j++) {
src[j] = vld1q_f16(src_data + i*qk + 8*j);
}
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]));
}
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));
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));
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);
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;
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);
}
}
}
};