File size: 3,914 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
#include "ling3/quantization.h"
#include "ling3/w4_linear.h"
#include <algorithm>
#include <array>
#include <cmath>
#include <cstdlib>
#include <exception>
#include <iostream>
#include <stdexcept>
#include <thread>
#include <vector>

// Hardware regression: compare against integer reference math, including
// non-power-of-two batches, split K, indexed input, and lazy ordinary input.
int main() {
    using namespace ling3;
    unsetenv("LING3_PREFILL_W4A4");
    constexpr int k = 128, n = 576, source_rows = 131;
    std::vector<std::byte> packed(k * n / 2);
    std::vector<float> scales(n), input(source_rows * k), activation_scales(source_rows);
    std::vector<std::int32_t> correction(n);
    std::vector<std::int8_t> quantized(input.size());
    for (int col = 0; col < n; ++col) scales[col] = .01F * (1 + col % 7);
    for (int i = 0; i < k * n; ++i) {
        const int value = ((i * 13 + i / n * 7) % 9) - 4;
        packed[i / 2] |= std::byte((value & 15) << (4 * (i % 2)));
        correction[i % n] += 8 * value;
    }
    for (std::size_t i = 0; i < input.size(); ++i) input[i] = std::sin(float(i) * .17F) * 2;
    for (int row = 0; row < source_rows; ++row)
        activation_scales[row] = QuantizeSymmetricInt8(
            std::span<const float>(input).subspan(row * k, k),
            std::span<std::int8_t>(quantized).subspan(row * k, k)).scale;
    auto verify = [&](int parts, std::vector<int> cores, int offset) {
        std::vector<std::int32_t> integer(n);
        DynamicW4Linear linear({k, n, parts, std::move(cores), 2}, packed, scales, correction);
        for (std::size_t rows : {1, 2, 3, 4, 5, 8, 9, 16, 17, 32, 33, 64, 65, 128, 3, 1}) {
            std::vector<std::size_t> indices(rows);
            std::vector<float> gathered(rows * k), expected(rows * n), actual(rows * n);
            for (std::size_t row = 0; row < rows; ++row) {
                const auto index = indices[row] = (row * 17 + 3 + offset) % source_rows;
                std::copy_n(input.begin() + index * k, k, gathered.begin() + row * k);
                ReferenceW4Linear(std::span<const std::int8_t>(quantized).subspan(index * k, k),
                                  packed, n, integer);
                DequantizePerChannel(integer, activation_scales[index], scales,
                                     std::span<float>(expected).subspan(row * n, n));
            }
            linear.PrepareBatch(rows, true);
            linear.RunBatchQuantizedRows(quantized, activation_scales, indices, linear, actual);
            if (actual != expected) throw std::runtime_error("indexed/reference mismatch");
            // A previously indexed-only workspace must safely allocate private
            // quantizer scratch if an ordinary call later uses the same bucket.
            linear.RunBatch(gathered, rows, actual);
            if (actual != expected) throw std::runtime_error("ordinary/reference mismatch");
            linear.RunBatchQuantizedRows(quantized, activation_scales, indices, linear, actual);
            if (actual != expected) throw std::runtime_error("indexed after ordinary mismatch");
        }
    };
    for (int parts : {1, 2}) verify(parts, {0, 1, 2}, 0);
    // Independent expert runners execute concurrently, with different inputs.
    // Their A/C owners must remain isolated even when their shapes match.
    std::array<std::thread, 3> workers;
    std::array<std::exception_ptr, 3> errors {};
    for (int core = 0; core < 3; ++core) workers[core] = std::thread([&, core] {
        try { verify(1 + core % 2, {core}, core * 23); }
        catch (...) { errors[core] = std::current_exception(); }
    });
    for (auto & worker : workers) worker.join();
    for (const auto & error : errors) if (error) std::rethrow_exception(error);
    std::cout << "PASS: split/unsplit K, eight buckets, indexed/ordinary transitions, three concurrent cores match reference exactly\n";
}