File size: 5,166 Bytes
daae291
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
// satu1_avx2.cpp β€” IKNN-Rl1-A1 β€” SatU1 AVX2 Kernel (Zen3 LUT Emulation, No VPOPCNTDQ)
// Version: v1.0
// Created: 2026-09-03T16:55:00+07:00
// Last Updated: 2026-09-03T16:55:00+07:00
// Status: PUBLISHABLE β€” EN ONLY β€” M1 Kernel Validation
// Repo: IKNN-Rl1-A1 β€” Integrated Knowledge-phase Neural Network β€” Recursive Language Iteration 1 β€” Architecture 1
// Hardware Target: AMD Ryzen 5 5650U β€” 6C/12T Zen3, AVX2 (NO AVX-512, NO VPOPCNTDQ), DDR4 35-45GB/s real, Vega 7 shared-RAM
// Description: SatU1 1-bit XNOR+POPCOUNT via AVX2 vpshufb LUT emulation (12 cycles per 256-bit, not 4)
//              Canonical mapping same as AVX-512 version
//              This kernel is for Ryzen5 target, but can be tested on Xeon AVX-512 in AVX2 mode

#include "satu1_common.h"
#include <cstdint>
#include <immintrin.h>

namespace iknn {
namespace satu1 {
namespace avx2 {

// AVX2 popcount emulation via nibble LUT (vpshufb) β€” 8 instructions per 256-bit
// LUT for 4-bit popcount: 0->0,1->1,2->1,3->2,...15->4
inline __m256i popcount256_avx2(__m256i v) {
    const __m256i lut = _mm256_setr_epi8(
        0,1,1,2,1,2,2,3,1,2,2,3,2,3,3,4,
        0,1,1,2,1,2,2,3,1,2,2,3,2,3,3,4
    );
    const __m256i low_mask = _mm256_set1_epi8(0x0F);
    __m256i lo = _mm256_and_si256(v, low_mask);
    __m256i hi = _mm256_and_si256(_mm256_srli_epi16(v, 4), low_mask);
    __m256i cnt1 = _mm256_shuffle_epi8(lut, lo);
    __m256i cnt2 = _mm256_shuffle_epi8(lut, hi);
    __m256i cnt = _mm256_add_epi8(cnt1, cnt2);
    // Sum bytes via SAD
    __m256i sum = _mm256_sad_epu8(cnt, _mm256_setzero_si256());
    return sum; // contains 4x 64-bit partial sums in 256-bit
}

inline float compute_satu1_block_avx2(
    const uint64_t* __restrict__ activation_bits,
    const uint64_t* __restrict__ weight_bits,
    float alpha,
    int block_size_words) // 1 word = 64 weights
{
    const __m256i zero = _mm256_setzero_si256();
    __m256i acc0 = _mm256_setzero_si256(); // will hold SAD results

    int i = 0;
    int32_t total_pop = 0;

    // Process 4 words = 256 bits per iteration
    for (; i + 3 < block_size_words; i += 4) {
        __m256i a = _mm256_loadu_si256(reinterpret_cast<const __m256i*>(&activation_bits[i]));
        __m256i w = _mm256_loadu_si256(reinterpret_cast<const __m256i*>(&weight_bits[i]));
        __m256i x = _mm256_xor_si256(a, w);
        __m256i xnor = _mm256_xor_si256(x, _mm256_set1_epi8(0xFF)); // NOT via XOR 0xFF per byte

        // Popcount emulation
        const __m256i lut = _mm256_setr_epi8(
            0,1,1,2,1,2,2,3,1,2,2,3,2,3,3,4,
            0,1,1,2,1,2,2,3,1,2,2,3,2,3,3,4
        );
        const __m256i low_mask = _mm256_set1_epi8(0x0F);
        __m256i lo = _mm256_and_si256(xnor, low_mask);
        __m256i hi = _mm256_and_si256(_mm256_srli_epi16(xnor, 4), low_mask);
        __m256i cnt1 = _mm256_shuffle_epi8(lut, lo);
        __m256i cnt2 = _mm256_shuffle_epi8(lut, hi);
        __m256i cnt = _mm256_add_epi8(cnt1, cnt2);
        __m256i sad = _mm256_sad_epu8(cnt, zero);

        // Extract 4x 64-bit sums
        // sad contains: [sum0(0-7), sum1(8-15), sum2(16-23), sum3(24-31)] each in 64-bit lane
        alignas(32) uint64_t tmp[4];
        _mm256_store_si256(reinterpret_cast<__m256i*>(tmp), sad);
        total_pop += tmp[0] + tmp[1] + tmp[2] + tmp[3];
    }

    // Tail
    for (; i < block_size_words; ++i) {
        uint64_t xnor = ~(activation_bits[i] ^ weight_bits[i]);
        total_pop += __builtin_popcountll(xnor);
    }

    int32_t signed_sum = (2 * total_pop) - (block_size_words * 64);
    return static_cast<float>(signed_sum) * alpha;
}

} // namespace avx2
} // namespace satu1
} // namespace iknn

#ifdef SATU1_AVX2_TEST
#include <random>
#include <iostream>
#include <chrono>
#include <cmath>

int main() {
    using namespace iknn::satu1;
    using namespace iknn::satu1::avx2;

    std::cout << "[SatU1 AVX2 Test] Zen3 LUT emulation" << std::endl;

    std::mt19937_64 rng(42);
    const int WORDS = 4;
    uint64_t act[WORDS], w[WORDS];
    for (int i = 0; i < WORDS; ++i) {
        act[i] = rng();
        w[i] = rng();
    }
    float alpha = 0.5f;

    float naive = compute_satu1_block_naive(act, w, alpha, WORDS);
    float avx2_res = compute_satu1_block_avx2(act, w, alpha, WORDS);

    std::cout << "Naive: " << naive << " AVX2: " << avx2_res << std::endl;
    if (std::abs(naive - avx2_res) < 1e-3f) {
        std::cout << "[PASS] Correctness AVX2" << std::endl;
    } else {
        std::cout << "[FAIL] AVX2 mismatch" << std::endl;
        return 1;
    }

    const int ITERS = 1000000;
    auto start = std::chrono::high_resolution_clock::now();
    float sum = 0;
    for (int it = 0; it < ITERS; ++it) {
        act[0] ^= it;
        sum += compute_satu1_block_avx2(act, w, alpha, WORDS);
    }
    auto end = std::chrono::high_resolution_clock::now();
    auto ms = std::chrono::duration_cast<std::chrono::milliseconds>(end - start).count();
    double giga_pop = (double)ITERS * WORDS * 64 / 1e9;
    double sec = ms / 1000.0;
    std::cout << "[BENCH AVX2] Time: " << ms << " ms, Giga-popcnt/s: " << giga_pop/sec << " Sum: " << sum << std::endl;
    return 0;
}
#endif