#include "ling3/quantization.h" #include #include #include #include #include #include #include namespace { std::byte Pack(std::int8_t low, std::int8_t high) { return static_cast( (static_cast(low) & 0x0FU) | ((static_cast(high) & 0x0FU) << 4U)); } } // namespace int main() { std::array all_codes {}; for (int value = -128; value <= 127; ++value) all_codes[value + 128] = value; std::array high {}; std::array low {}; ling3::SplitInt8ForInt4Matmul(all_codes, high, low); for (std::size_t index = 0; index < all_codes.size(); ++index) { const int reconstructed = 16 * static_cast(high[index]) + static_cast(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 activation_distribution(-127, 127); std::uniform_int_distribution weight_distribution(-8, 7); std::vector activation(k); std::vector weights(k * n); for (auto & value : activation) value = static_cast(activation_distribution(generator)); for (auto & value : weights) value = static_cast(weight_distribution(generator)); std::vector 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 direct(n); ling3::ReferenceW4Linear(activation, packed, n, direct); std::vector activation_high(k), activation_low(k); ling3::SplitInt8ForInt4Matmul(activation, activation_high, activation_low); std::vector 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 input = {-2.0F, -0.5F, 0.0F, 0.25F, 1.0F, 2.0F}; std::array 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; }