Ling-3.0-tiny-RKNN / src /quantization.cpp
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
5.42 kB
#include "ling3/quantization.h"
#include <algorithm>
#include <cmath>
#include <limits>
#include <stdexcept>
#if defined(__aarch64__)
#include <arm_neon.h>
#endif
namespace ling3 {
namespace {
DynamicQuantization QuantizeSymmetric(
std::span<const float> input,
std::span<std::int8_t> output,
int limit) noexcept {
if (output.size() != input.size()) return {0.0F, input.size()};
float maximum = 0.0F;
for (float value : input) {
if (std::isfinite(value)) maximum = std::max(maximum, std::abs(value));
}
const float scale = maximum > 0.0F ? maximum / static_cast<float>(limit) : 1.0F;
const float inverse_scale = 1.0F / scale;
std::size_t clipped = 0;
std::size_t index = 0;
#if defined(__aarch64__)
const float32x4_t inverse = vdupq_n_f32(inverse_scale);
const float32x4_t largest = vdupq_n_f32(std::numeric_limits<float>::max());
const float32x4_t zero = vdupq_n_f32(0.0F);
const int32x4_t minimum = vdupq_n_s32(-limit);
const int32x4_t maximum_code = vdupq_n_s32(limit);
for (; index + 16 <= input.size(); index += 16) {
const auto convert = [&](std::size_t offset) {
const float32x4_t value = vld1q_f32(input.data() + index + offset);
const uint32x4_t finite = vcleq_f32(vabsq_f32(value), largest);
const float32x4_t scaled = vbslq_f32(finite, vmulq_f32(value, inverse), zero);
return vmaxq_s32(minimum, vminq_s32(maximum_code, vcvtnq_s32_f32(scaled)));
};
const int16x8_t codes01 = vcombine_s16(
vqmovn_s32(convert(0)), vqmovn_s32(convert(4)));
const int16x8_t codes23 = vcombine_s16(
vqmovn_s32(convert(8)), vqmovn_s32(convert(12)));
vst1q_s8(
output.data() + index,
vcombine_s8(vqmovn_s16(codes01), vqmovn_s16(codes23)));
}
#endif
for (; index < input.size(); ++index) {
const float scaled = std::isfinite(input[index]) ? input[index] * inverse_scale : 0.0F;
const long rounded = std::lrint(scaled);
clipped += rounded < -limit || rounded > limit;
output[index] = static_cast<std::int8_t>(
std::clamp(rounded, static_cast<long>(-limit), static_cast<long>(limit)));
}
return {scale, clipped};
}
} // namespace
DynamicQuantization QuantizeSymmetricInt8(
std::span<const float> input,
std::span<std::int8_t> output) noexcept {
return QuantizeSymmetric(input, output, 127);
}
DynamicQuantization QuantizeSymmetricInt4(
std::span<const float> input,
std::span<std::int8_t> output) noexcept {
return QuantizeSymmetric(input, output, 7);
}
std::int8_t DecodeInt4LowFirst(
std::span<const std::byte> packed,
std::size_t index) {
if (index / 2 >= packed.size()) throw std::out_of_range("INT4 index is out of range");
const auto byte = std::to_integer<std::uint8_t>(packed[index / 2]);
const std::uint8_t nibble = (index & 1U) == 0 ? byte & 0x0FU : byte >> 4U;
return static_cast<std::int8_t>(nibble >= 8U ? static_cast<int>(nibble) - 16 : nibble);
}
void SplitInt8ForInt4Matmul(
std::span<const std::int8_t> input,
std::span<std::int8_t> high,
std::span<std::int8_t> low) {
if (high.size() != input.size() || low.size() != input.size()) {
throw std::invalid_argument("INT8 split buffers have different sizes");
}
for (std::size_t index = 0; index < input.size(); ++index) {
const int value = input[index];
const int high_value = value < 0 ? -((-value + 15) / 16) : value / 16;
const int low_nibble = (value & 0x0F) ^ 0x08;
const int low_value = low_nibble >= 8 ? low_nibble - 16 : low_nibble;
high[index] = static_cast<std::int8_t>(high_value);
low[index] = static_cast<std::int8_t>(low_value);
}
}
void ReferenceW4Linear(
std::span<const std::int8_t> input,
std::span<const std::byte> packed_weights,
std::size_t output_channels,
std::span<std::int32_t> output) {
if (output_channels == 0 || output.size() != output_channels ||
packed_weights.size() * 2 != input.size() * output_channels) {
throw std::invalid_argument("W4 reference tensor sizes are inconsistent");
}
std::fill(output.begin(), output.end(), 0);
for (std::size_t k = 0; k < input.size(); ++k) {
for (std::size_t n = 0; n < output_channels; ++n) {
const auto weight = DecodeInt4LowFirst(packed_weights, k * output_channels + n);
output[n] += static_cast<std::int32_t>(input[k]) * weight;
}
}
}
void DequantizePerChannel(
std::span<const std::int32_t> input,
float input_scale,
std::span<const float> weight_scales,
std::span<float> output) noexcept {
if (input.size() != weight_scales.size() || output.size() != input.size()) return;
std::size_t index = 0;
#if defined(__aarch64__)
const float32x4_t activation = vdupq_n_f32(input_scale);
for (; index + 4 <= input.size(); index += 4) {
const float32x4_t values = vcvtq_f32_s32(vld1q_s32(input.data() + index));
const float32x4_t scales = vld1q_f32(weight_scales.data() + index);
vst1q_f32(output.data() + index, vmulq_f32(vmulq_f32(values, activation), scales));
}
#endif
for (; index < input.size(); ++index) {
output[index] = static_cast<float>(input[index]) * input_scale * weight_scales[index];
}
}
} // namespace ling3