File size: 3,179 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
#include "ling3/quantization.h"

#include <array>
#include <cmath>
#include <cstddef>
#include <cstdint>
#include <iostream>
#include <random>
#include <vector>

namespace {

std::byte Pack(std::int8_t low, std::int8_t high) {
    return static_cast<std::byte>(
        (static_cast<std::uint8_t>(low) & 0x0FU) |
        ((static_cast<std::uint8_t>(high) & 0x0FU) << 4U));
}

} // namespace

int main() {
    std::array<std::int8_t, 256> all_codes {};
    for (int value = -128; value <= 127; ++value) all_codes[value + 128] = value;
    std::array<std::int8_t, 256> high {};
    std::array<std::int8_t, 256> low {};
    ling3::SplitInt8ForInt4Matmul(all_codes, high, low);
    for (std::size_t index = 0; index < all_codes.size(); ++index) {
        const int reconstructed = 16 * static_cast<int>(high[index]) +
                                  static_cast<int>(low[index]) + 8;
        if (reconstructed != all_codes[index] || high[index] < -8 || high[index] > 7 ||
            low[index] < -8 || low[index] > 7) {
            std::cerr << "INT8 split identity failed at " << index << '\n';
            return 1;
        }
    }

    constexpr std::size_t k = 64;
    constexpr std::size_t n = 192;
    std::mt19937 generator(0x4C334E50U);
    std::uniform_int_distribution<int> activation_distribution(-127, 127);
    std::uniform_int_distribution<int> weight_distribution(-8, 7);
    std::vector<std::int8_t> activation(k);
    std::vector<std::int8_t> weights(k * n);
    for (auto & value : activation) value = static_cast<std::int8_t>(activation_distribution(generator));
    for (auto & value : weights) value = static_cast<std::int8_t>(weight_distribution(generator));
    std::vector<std::byte> packed(weights.size() / 2);
    for (std::size_t index = 0; index < weights.size(); index += 2) {
        packed[index / 2] = Pack(weights[index], weights[index + 1]);
    }

    std::vector<std::int32_t> direct(n);
    ling3::ReferenceW4Linear(activation, packed, n, direct);
    std::vector<std::int8_t> activation_high(k), activation_low(k);
    ling3::SplitInt8ForInt4Matmul(activation, activation_high, activation_low);
    std::vector<std::int32_t> decomposed(n, 0);
    for (std::size_t column = 0; column < n; ++column) {
        std::int32_t value = 0;
        std::int32_t correction = 0;
        for (std::size_t row = 0; row < k; ++row) {
            const int weight = weights[row * n + column];
            value += (16 * activation_high[row] + activation_low[row]) * weight;
            correction += 8 * weight;
        }
        decomposed[column] = value + correction;
    }
    if (direct != decomposed) {
        std::cerr << "two-pass W4 accumulator does not match INT8 x INT4\n";
        return 1;
    }

    const std::array<float, 6> input = {-2.0F, -0.5F, 0.0F, 0.25F, 1.0F, 2.0F};
    std::array<std::int8_t, 6> quantized {};
    const auto quant = ling3::QuantizeSymmetricInt8(input, quantized);
    if (std::abs(quant.scale - 2.0F / 127.0F) > 1.0e-7F || quantized.front() != -127 ||
        quantized.back() != 127 || quant.clipped != 0) {
        std::cerr << "dynamic symmetric quantization failed\n";
        return 1;
    }
    return 0;
}