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;
}
|