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