#include "ling3/quantization.h" #include #include #include #include #if defined(__aarch64__) #include #endif namespace ling3 { namespace { DynamicQuantization QuantizeSymmetric( std::span input, std::span 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(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::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::clamp(rounded, static_cast(-limit), static_cast(limit))); } return {scale, clipped}; } } // namespace DynamicQuantization QuantizeSymmetricInt8( std::span input, std::span output) noexcept { return QuantizeSymmetric(input, output, 127); } DynamicQuantization QuantizeSymmetricInt4( std::span input, std::span output) noexcept { return QuantizeSymmetric(input, output, 7); } std::int8_t DecodeInt4LowFirst( std::span 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(packed[index / 2]); const std::uint8_t nibble = (index & 1U) == 0 ? byte & 0x0FU : byte >> 4U; return static_cast(nibble >= 8U ? static_cast(nibble) - 16 : nibble); } void SplitInt8ForInt4Matmul( std::span input, std::span high, std::span 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(high_value); low[index] = static_cast(low_value); } } void ReferenceW4Linear( std::span input, std::span packed_weights, std::size_t output_channels, std::span 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(input[k]) * weight; } } } void DequantizePerChannel( std::span input, float input_scale, std::span weight_scales, std::span 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(input[index]) * input_scale * weight_scales[index]; } } } // namespace ling3