File size: 5,423 Bytes
3fd1a35 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 | #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
|