IKNN-Rl1-A1 / kernels /satu1_avx2.cpp
deeprcurs-staff's picture
Upload kernels/satu1_avx2.cpp with huggingface_hub
daae291 verified
Raw History Blame Contribute Delete
5.17 kB
// 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