File size: 24,258 Bytes
51e3123
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7c251e6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
51e3123
 
7c251e6
 
 
 
 
 
 
51e3123
 
 
 
 
 
 
 
 
 
 
 
7c251e6
 
51e3123
 
 
 
 
7c251e6
51e3123
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7c251e6
 
 
 
51e3123
 
7c251e6
51e3123
 
 
 
 
 
 
 
 
 
7c251e6
 
51e3123
 
 
 
 
 
 
 
7c251e6
51e3123
 
 
 
 
 
 
 
 
 
 
 
 
 
7c251e6
 
51e3123
 
 
 
 
 
 
7c251e6
51e3123
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7c251e6
 
51e3123
 
 
 
 
7c251e6
51e3123
 
 
 
 
 
 
 
 
 
 
 
7c251e6
 
51e3123
 
 
7c251e6
 
51e3123
 
 
 
 
 
7c251e6
51e3123
 
 
 
 
 
 
 
 
 
 
 
7c251e6
 
51e3123
7c251e6
51e3123
7c251e6
 
51e3123
 
 
 
 
 
 
 
 
 
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
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
/*
 * Standalone trit GEMV kernel — no PyTorch dependency.
 * Compiles with nvcc to a shared library (.so/.dll).
 * Called from Python via ctypes, or from C/C++ directly.
 *
 * Compile:
 *   nvcc -O3 --use_fast_math -shared -Xcompiler -fPIC \
 *     -gencode=arch=compute_70,code=sm_70 \
 *     -gencode=arch=compute_75,code=sm_75 \
 *     -gencode=arch=compute_80,code=sm_80 \
 *     -gencode=arch=compute_86,code=sm_86 \
 *     -gencode=arch=compute_89,code=sm_89 \
 *     -gencode=arch=compute_90,code=sm_90 \
 *     -gencode=arch=compute_100,code=sm_100 \
 *     -gencode=arch=compute_120,code=sm_120 \
 *     -o libtrit_gemv.so trit_gemv_standalone.cu
 *
 * Supports: Volta(V100), Turing(2080), Ampere(3090/A100), Ada(4080/4090),
 *           Hopper(H100), Blackwell(5070/5090) — all in one binary.
 *
 * API: C functions with extern "C" — callable from any language.
 */

#include <cuda_runtime.h>
#include <stdint.h>

#define GROUP_SIZE 64
#define WARP_SIZE 32
#define TRIT_POS 1
#define TRIT_NEG 2

// Forward declarations
static void set_l2_persist(void* ptr, size_t bytes);
static void clear_l2_persist();

// ============================================================
// V27: D2 int4-packed + dp4a + L2 persist (champion kernel)
// ============================================================

// ============================================================
// V28: Branchless interleaved nibble decode (7 instructions for 8 weights)
//
// The trick: extract even nibbles (0,2,4,6) and odd nibbles (1,3,5,7)
// as separate byte vectors using mask + shift. Sign-extend all 4 bytes
// simultaneously with XOR + SUB (zero branches).
//
// x activations are pre-interleaved to match: x_evens has values at
// positions 0,2,4,6 and x_odds has 1,3,5,7. Pre-interleave is done
// once at activation quantization time (negligible cost).
//
// Instructions per 8 weights:
//   v27: 32 (loop with branches)
//   v28: 14 (7 expand + 2 dp4a + 3 load + 2 scale)
//   Balance BW shift: 3.4 → 5.6 TB/s on A100 → crosses into memory-bound!
// ============================================================

#define V28_RPB 16
#define V28_WPG 8
#define V28_BS (V28_RPB * WARP_SIZE)

__global__ void k_v28(
    const uint32_t* __restrict__ pt,    // int4 packed: 8 weights per uint32
    const float* __restrict__ ws,       // weight scales [rows * ng]
    const uint32_t* __restrict__ xt_e,  // x int8 EVEN positions [ng * 8]
    const uint32_t* __restrict__ xt_o,  // x int8 ODD positions [ng * 8]
    const float* __restrict__ xs,       // x scales [ng]
    float* __restrict__ y,
    int cols, int rows, int num_groups
) {
    int wid = threadIdx.x / WARP_SIZE;
    int lane = threadIdx.x % WARP_SIZE;
    int row = blockIdx.x * V28_RPB + wid;
    if (row >= rows) return;

    const uint32_t* row_w = &pt[row * num_groups * V28_WPG];
    const float* row_ws = &ws[row * num_groups];

    float acc = 0.0f;
    int total_words = num_groups * V28_WPG;

    for (int base = 0; base < total_words; base += WARP_SIZE) {
        int w = base + lane;
        if (w < total_words) {
            // COALESCED load
            uint32_t word = __ldg(&row_w[w]);
            int g = w >> 3;            // group index (shift, 1 cycle)
            int word_in_group = w & 7; // word within group (mask, 1 cycle)

            // === BRANCHLESS INT4→INT8 EXPANSION (7 instructions) ===
            // Extract even nibbles (weights 0,2,4,6) into bytes
            uint32_t evens = word & 0x0F0F0F0F;                    // AND (1 op)
            evens = (evens ^ 0x08080808) - 0x08080808;              // XOR+SUB (2 ops)

            // Extract odd nibbles (weights 1,3,5,7) into bytes
            uint32_t odds = (word >> 4) & 0x0F0F0F0F;              // SHR+AND (2 ops)
            odds = (odds ^ 0x08080808) - 0x08080808;                // XOR+SUB (2 ops)
            // Total: 7 instructions for 8 sign-extended int8 values

            // dp4a against pre-interleaved x
            // x_evens[g*8 + word_in_group] has activations at even positions
            // x_odds[g*8 + word_in_group] has activations at odd positions
            int x_idx = g * 8 + word_in_group;
            uint32_t xe = __ldg(&xt_e[x_idx]);
            uint32_t xo = __ldg(&xt_o[x_idx]);

            int dp_e = __dp4a((int)evens, (int)xe, 0);
            int dp_o = __dp4a((int)odds, (int)xo, 0);

            float combined_scale = __ldg(&row_ws[g]) * __ldg(&xs[g]);
            acc += (float)(dp_e + dp_o) * combined_scale;
        }
    }

    #pragma unroll
    for (int o = 16; o > 0; o >>= 1)
        acc += __shfl_down_sync(0xFFFFFFFF, acc, o);

    if (lane == 0) y[row] = acc;
}

// ============================================================
// V29: BIAS TRICK — zero sign extension, unsigned weights + correction
//
// Store weights as (level + 4) → range 0-8 (unsigned).
// Zero-extend nibbles: AND only, no XOR, no SUB.
// dp4a gives biased result. Subtract precomputed correction.
//
// Decode: 3 instructions (AND, SHR, AND) vs v28's 7
// Correction: 1 SUB + 1 LD per word (precomputed x_bias)
// Net: 10 instructions per word vs v28's 14
// ============================================================

#define V29_RPB 16
#define V29_WPG 8
#define V29_BS (V29_RPB * WARP_SIZE)

__global__ void k_v29(
    const uint32_t* __restrict__ pt,    // UNSIGNED int4: (level+4) packed, 0-8 per nibble
    const float* __restrict__ ws,       // weight scales [rows * ng]
    const uint32_t* __restrict__ xt_e,  // x int8 EVEN [ng * 8]
    const uint32_t* __restrict__ xt_o,  // x int8 ODD [ng * 8]
    const int* __restrict__ x_bias,     // precomputed 4×(sum of 8 x values) per word position [ng * 8]
    const float* __restrict__ xs,       // x scales [ng]
    float* __restrict__ y,
    int cols, int rows, int num_groups
) {
    int wid = threadIdx.x / WARP_SIZE;
    int lane = threadIdx.x % WARP_SIZE;
    int row = blockIdx.x * V29_RPB + wid;
    if (row >= rows) return;

    const uint32_t* row_w = &pt[row * num_groups * V29_WPG];
    const float* row_ws = &ws[row * num_groups];

    float acc = 0.0f;
    int total_words = num_groups * V29_WPG;

    for (int base = 0; base < total_words; base += WARP_SIZE) {
        int w = base + lane;
        if (w < total_words) {
            uint32_t word = __ldg(&row_w[w]);
            int g = w >> 3;
            int wig = w & 7;

            // BIAS DECODE: 3 instructions total (no XOR, no SUB)
            uint32_t evens = word & 0x0F0F0F0F;               // AND
            uint32_t odds = (word >> 4) & 0x0F0F0F0F;          // SHR + AND
            // Values 0-8 are valid positive int8 — no sign extension needed

            int x_idx = g * 8 + wig;
            uint32_t xe = __ldg(&xt_e[x_idx]);
            uint32_t xo = __ldg(&xt_o[x_idx]);

            // dp4a: biased result (includes +4 per weight)
            int dp = __dp4a((int)evens, (int)xe, 0)
                   + __dp4a((int)odds, (int)xo, 0);

            // Subtract precomputed bias: 4 × sum of 8 x values for this word
            int bias = __ldg(&x_bias[x_idx]);
            dp -= bias;

            float combined_scale = __ldg(&row_ws[g]) * __ldg(&xs[g]);
            acc += (float)dp * combined_scale;
        }
    }

    #pragma unroll
    for (int o = 16; o > 0; o >>= 1)
        acc += __shfl_down_sync(0xFFFFFFFF, acc, o);

    if (lane == 0) y[row] = acc;
}

// v29 wrapper moved to extern "C" block below

// Keep v27 as fallback (doesn't need interleaved x)
__device__ __forceinline__ uint32_t extract_int4x4_to_int8x4(uint32_t word, int start) {
    uint32_t result = 0;
    #pragma unroll
    for (int i = 0; i < 4; i++) {
        int shift = (start + i) * 4;
        int nibble = (word >> shift) & 0xF;
        int val = (nibble & 0x8) ? (nibble | 0xFFFFFFF0) : nibble;
        result |= ((uint32_t)(val & 0xFF)) << (i * 8);
    }
    return result;
}

#define V27_RPB 16
#define V27_WPG 8
#define V27_BS (V27_RPB * WARP_SIZE)

__global__ void k_v27(
    const uint32_t* __restrict__ pt,
    const float* __restrict__ ws,
    const uint32_t* __restrict__ xt,
    const float* __restrict__ xs,
    float* __restrict__ y,
    int cols, int rows, int num_groups
) {
    int wid = threadIdx.x / WARP_SIZE;
    int lane = threadIdx.x % WARP_SIZE;
    int row = blockIdx.x * V28_RPB + wid;
    if (row >= rows) return;

    const uint32_t* row_w = &pt[row * num_groups * V27_WPG];
    const float* row_ws = &ws[row * num_groups];

    float acc = 0.0f;
    int total_words = num_groups * V27_WPG;

    for (int base = 0; base < total_words; base += WARP_SIZE) {
        int w = base + lane;
        if (w < total_words) {
            uint32_t word = __ldg(&row_w[w]);
            int g = w >> 3;
            int word_in_group = w & 7;

            uint32_t lo = extract_int4x4_to_int8x4(word, 0);
            uint32_t hi = extract_int4x4_to_int8x4(word, 4);

            int x_base = g * 16 + word_in_group * 2;
            uint32_t x_lo = __ldg(&xt[x_base]);
            uint32_t x_hi = __ldg(&xt[x_base + 1]);

            int dp_lo = __dp4a((int)lo, (int)x_lo, 0);
            int dp_hi = __dp4a((int)hi, (int)x_hi, 0);

            float combined_scale = __ldg(&row_ws[g]) * __ldg(&xs[g]);
            acc += (float)(dp_lo + dp_hi) * combined_scale;
        }
    }

    #pragma unroll
    for (int o = 16; o > 0; o >>= 1)
        acc += __shfl_down_sync(0xFFFFFFFF, acc, o);

    if (lane == 0) y[row] = acc;
}

// ============================================================
// V9-style: trit-packed d3 (for models stored in trit format)
// ============================================================

#define V9R 4
#define V9W 2
#define V9BS (V9R * V9W * WARP_SIZE)

__device__ __forceinline__ float mac_wide_d3(
    const uint32_t* __restrict__ p, const float* __restrict__ x, int tid
) {
    float acc = 0.0f;
    #pragma unroll
    for (int i = 0; i < 4; i++) {
        int idx = tid * 4 + i;
        int w = idx / 5, pos = idx % 5;
        uint32_t bits = (__ldg(&p[w]) >> (pos * 6)) & 0x3F;
        int t0 = bits & 3, t1 = (bits >> 2) & 3, t2 = (bits >> 4) & 3;
        int lv = ((t2==TRIT_POS)-(t2==TRIT_NEG))*9 + ((t1==TRIT_POS)-(t1==TRIT_NEG))*3
               + ((t0==TRIT_POS)-(t0==TRIT_NEG));
        acc += lv * __ldg(&x[idx]);
    }
    return acc;
}

__global__ void k_v9(
    const uint32_t* __restrict__ pt, const float* __restrict__ sc,
    const float* __restrict__ x, float* __restrict__ y,
    int in_f, int out_f, int depth
) {
    __shared__ float parts[V9R * V9W];
    int base = blockIdx.x * V9R;
    int wid = threadIdx.x / WARP_SIZE, lane = threadIdx.x % WARP_SIZE;
    int lr = wid / V9W, rw = wid % V9W, row = base + lr;
    int ng = in_f / GROUP_SIZE;
    const int w = 13;

    int half = lane / 16;
    int tid_in_group = lane % 16;

    float partial = 0.0f;
    if (row < out_f) {
        for (int g_pair = rw; g_pair < (ng + 1) / 2; g_pair += V9W) {
            int g = g_pair * 2 + half;
            if (g < ng) {
                float ga = mac_wide_d3(&pt[(row * ng + g) * w],
                                       &x[g * GROUP_SIZE], tid_in_group);
                unsigned mask = half ? 0xFFFF0000u : 0x0000FFFFu;
                #pragma unroll
                for (int o = 8; o > 0; o >>= 1)
                    ga += __shfl_down_sync(mask, ga, o);
                if (tid_in_group == 0)
                    partial += ga * __ldg(&sc[row * ng + g]);
            }
        }
    }

    float my_partial = (lane == 0 || lane == 16) ? partial : 0.0f;
    my_partial += __shfl_xor_sync(0xFFFFFFFF, my_partial, 16);
    if (lane == 0) parts[wid] = my_partial;
    __syncthreads();
    if (lane == 0 && rw == 0 && row < out_f) {
        float s = 0;
        for (int i = 0; i < V9W; i++) s += parts[lr * V9W + i];
        y[row] = s;
    }
}

// ============================================================
// L2 persistence helpers
// ============================================================

static void set_l2_persist(void* ptr, size_t bytes) {
    cudaStreamAttrValue attr;
    attr.accessPolicyWindow.base_ptr = ptr;
    attr.accessPolicyWindow.num_bytes = bytes;
    attr.accessPolicyWindow.hitRatio = 1.0f;
    attr.accessPolicyWindow.hitProp = cudaAccessPropertyPersisting;
    attr.accessPolicyWindow.missProp = cudaAccessPropertyStreaming;
    cudaStreamSetAttribute(0, cudaStreamAttributeAccessPolicyWindow, &attr);
}

static void clear_l2_persist() {
    cudaStreamAttrValue attr;
    memset(&attr, 0, sizeof(attr));
    cudaStreamSetAttribute(0, cudaStreamAttributeAccessPolicyWindow, &attr);
}

// ============================================================
// C API — callable from any language via dlopen/ctypes/FFI
// ============================================================

// Error codes for the last_error reporting channel.
//   0 = success
//   negative = host-side argument validation failure (no kernel was launched)
//   positive = cudaError_t value from a kernel launch or runtime call
#define TRIT_OK              0
#define TRIT_ERR_NULL_PTR   -1
#define TRIT_ERR_BAD_DIM    -2
#define TRIT_ERR_BAD_GROUP  -3   // num_groups != cols / GROUP_SIZE
#define TRIT_ERR_BAD_BUFFER -4   // buf too small / invalid

// Last-error slot. Set by every public entrypoint; read via trit_gemv_get_last_error().
static int g_last_error = TRIT_OK;

// Host-side argument validation. Returns 0 on success, negative on failure.
// Sets g_last_error and returns 1 (truthy) on failure for use in `if (validate(...)) return;`.
static inline int trit_validate_gemv(
    const void* pt, const void* ws, const void* y,
    int cols, int rows, int num_groups
) {
    if (!pt || !ws || !y) { g_last_error = TRIT_ERR_NULL_PTR; return 1; }
    if (cols <= 0 || rows <= 0 || num_groups <= 0) { g_last_error = TRIT_ERR_BAD_DIM; return 1; }
    if (cols % GROUP_SIZE != 0) { g_last_error = TRIT_ERR_BAD_DIM; return 1; }
    if (cols / GROUP_SIZE != num_groups) { g_last_error = TRIT_ERR_BAD_GROUP; return 1; }
    return 0;
}

// Capture cudaGetLastError() after a kernel launch into g_last_error.
static inline void trit_capture_launch_status() {
    cudaError_t e = cudaGetLastError();
    g_last_error = (e == cudaSuccess) ? TRIT_OK : (int)e;
}

extern "C" {

// Returns the error code from the most recent public-API call.
// 0 means success. Negative codes are host-side validation failures
// (TRIT_ERR_*); positive codes are cudaError_t values from CUDA itself.
int trit_gemv_get_last_error() {
    return g_last_error;
}

// v27: d2 int4-packed + dp4a (champion for GPU)
// pt: [rows * ng * 8] int32 (int4 packed weights)
// ws: [rows * ng] float32 (weight scales)
// xt: [ng * 16] int32 (int8 packed activations)
// xs: [ng] float32 (activation scales)
// y:  [rows] float32 (output)
void trit_gemv_d2_dp4a(
    const int32_t* pt, const float* ws,
    const int32_t* xt, const float* xs,
    float* y, int cols, int rows, int num_groups,
    int use_l2_persist
) {
    if (trit_validate_gemv(pt, ws, y, cols, rows, num_groups)) return;
    if (!xt || !xs) { g_last_error = TRIT_ERR_NULL_PTR; return; }
    if (use_l2_persist) {
        set_l2_persist((void*)pt, (size_t)rows * num_groups * 8 * sizeof(int32_t));
    }
    k_v27<<<(rows + V27_RPB - 1) / V27_RPB, V27_BS>>>(
        (const uint32_t*)pt, ws, (const uint32_t*)xt, xs, y, cols, rows, num_groups);
    trit_capture_launch_status();
    if (use_l2_persist) {
        clear_l2_persist();
    }
}

// v9: trit-packed d3 (for native trit format)
// pt: [rows * ng * 13] int32 (trit packed weights)
// sc: [rows * ng] float32 (scales)
// x:  [cols] float32 (activations)
// y:  [rows] float32 (output)
void trit_gemv_d3_native(
    const int32_t* pt, const float* sc,
    const float* x, float* y,
    int cols, int rows, int depth
) {
    if (!pt || !sc || !x || !y) { g_last_error = TRIT_ERR_NULL_PTR; return; }
    if (cols <= 0 || rows <= 0) { g_last_error = TRIT_ERR_BAD_DIM; return; }
    if (cols % GROUP_SIZE != 0) { g_last_error = TRIT_ERR_BAD_DIM; return; }
    if (depth < 1 || depth > 4) { g_last_error = TRIT_ERR_BAD_DIM; return; }
    k_v9<<<(rows + V9R - 1) / V9R, V9BS>>>(
        (const uint32_t*)pt, sc, x, y, cols, rows, depth);
    trit_capture_launch_status();
}

// v29: d2 unsigned int4 + bias trick (no sign extension)
void trit_gemv_d2_bias(
    const int32_t* pt, const float* ws,
    const int32_t* xt_e, const int32_t* xt_o,
    const int32_t* x_bias, const float* xs,
    float* y, int cols, int rows, int num_groups,
    int use_l2_persist
) {
    if (trit_validate_gemv(pt, ws, y, cols, rows, num_groups)) return;
    if (!xt_e || !xt_o || !x_bias || !xs) { g_last_error = TRIT_ERR_NULL_PTR; return; }
    if (use_l2_persist) {
        set_l2_persist((void*)pt, (size_t)rows * num_groups * 8 * sizeof(int32_t));
    }
    k_v29<<<(rows + V29_RPB - 1) / V29_RPB, V29_BS>>>(
        (const uint32_t*)pt, ws,
        (const uint32_t*)xt_e, (const uint32_t*)xt_o,
        (const int*)x_bias, xs,
        y, cols, rows, num_groups);
    trit_capture_launch_status();
    if (use_l2_persist) {
        clear_l2_persist();
    }
}

// v28: d2 int4 + branchless interleaved decode + dp4a
// xt_e/xt_o: pre-interleaved x (even/odd nibble positions)
// Each has ng*8 uint32 words (4 int8 values per word, 8 words per group)
void trit_gemv_d2_fast(
    const int32_t* pt, const float* ws,
    const int32_t* xt_e, const int32_t* xt_o, const float* xs,
    float* y, int cols, int rows, int num_groups,
    int use_l2_persist
) {
    if (trit_validate_gemv(pt, ws, y, cols, rows, num_groups)) return;
    if (!xt_e || !xt_o || !xs) { g_last_error = TRIT_ERR_NULL_PTR; return; }
    if (use_l2_persist) {
        set_l2_persist((void*)pt, (size_t)rows * num_groups * 8 * sizeof(int32_t));
    }
    k_v28<<<(rows + V28_RPB - 1) / V28_RPB, V28_BS>>>(
        (const uint32_t*)pt, ws,
        (const uint32_t*)xt_e, (const uint32_t*)xt_o, xs,
        y, cols, rows, num_groups);
    trit_capture_launch_status();
    if (use_l2_persist) {
        clear_l2_persist();
    }
}

// v21f: d3 int8-packed + dp4a (same format as v21f in wrapper)
// wt: [rows * ng * 16] int32 (int8 packed weight levels)
// ws: [rows * ng] float32 (weight scales)
// xt: [ng * 16] int32 (int8 packed activations)
// xs: [ng] float32 (activation scales)
#define V21F_RPB 4
#define V21F_BS (V21F_RPB * WARP_SIZE)

__global__ void k_v21f_standalone(
    const uint32_t* __restrict__ wt,
    const float* __restrict__ ws,
    const uint32_t* __restrict__ xt,
    const float* __restrict__ xs,
    float* __restrict__ y,
    int cols, int rows, int num_groups
) {
    int wid = threadIdx.x / WARP_SIZE;
    int lane = threadIdx.x % WARP_SIZE;
    int row = blockIdx.x * V21F_RPB + wid;
    if (row >= rows) return;

    const uint32_t* row_w = &wt[row * num_groups * 16];
    const float* row_ws = &ws[row * num_groups];
    float acc = 0.0f;
    int total_words = num_groups * 16;

    for (int base = 0; base < total_words; base += WARP_SIZE) {
        int w = base + lane;
        if (w < total_words) {
            uint32_t w_word = __ldg(&row_w[w]);
            uint32_t x_word = __ldg(&xt[w]);
            int dp = __dp4a((int)w_word, (int)x_word, 0);
            int g = w >> 4;  // 16 words per group
            acc += (float)dp * __ldg(&row_ws[g]) * __ldg(&xs[g]);
        }
    }

    #pragma unroll
    for (int o = 16; o > 0; o >>= 1)
        acc += __shfl_down_sync(0xFFFFFFFF, acc, o);

    if (lane == 0) y[row] = acc;
}

// ============================================================
// D3 HARDENED: int8 dp4a, 16 RPB, L2 persist, deferred reduction
//
// d3 levels: -13 to +13 (27 values), stored as int8 (1 byte each)
// 16 words per group, 4 int8 values per word = 64 values/group
// Division by 16 = shift (no div-by-13 problem!)
//
// This is the SIMPLEST kernel — no decode at all.
// The int8 values go DIRECTLY into dp4a.
// Pure memory-bound on every GPU.
// ============================================================

#define D3H_RPB 16
#define D3H_WPG 16   // 16 uint32 words per group (4 int8 each = 64 values)
#define D3H_BS (D3H_RPB * WARP_SIZE)

__global__ void k_d3_hardened(
    const uint32_t* __restrict__ wt,    // int8 packed: 16 words per group
    const float* __restrict__ ws,       // weight scales [rows * ng]
    const uint32_t* __restrict__ xt,    // x int8 packed: 16 words per group
    const float* __restrict__ xs,       // x scales [ng]
    float* __restrict__ y,
    int cols, int rows, int num_groups
) {
    int wid = threadIdx.x / WARP_SIZE;
    int lane = threadIdx.x % WARP_SIZE;
    int row = blockIdx.x * D3H_RPB + wid;
    if (row >= rows) return;

    const uint32_t* row_w = &wt[row * num_groups * D3H_WPG];
    const float* row_ws = &ws[row * num_groups];

    float acc = 0.0f;
    int total_words = num_groups * D3H_WPG;

    for (int base = 0; base < total_words; base += WARP_SIZE) {
        int w = base + lane;
        if (w < total_words) {
            // COALESCED load — 32 threads × 4 bytes = 128 bytes
            uint32_t w_word = __ldg(&row_w[w]);
            uint32_t x_word = __ldg(&xt[w]);

            // dp4a: 4× int8 multiply-accumulate — ZERO decode
            int dp = __dp4a((int)w_word, (int)x_word, 0);

            // Group index: SHIFT (16 is power of 2!)
            int g = w >> 4;

            // Deferred: accumulate per-thread, ONE reduction at end
            acc += (float)dp * __ldg(&row_ws[g]) * __ldg(&xs[g]);
        }
    }

    // ONE warp reduction
    #pragma unroll
    for (int o = 16; o > 0; o >>= 1)
        acc += __shfl_down_sync(0xFFFFFFFF, acc, o);

    if (lane == 0) y[row] = acc;
}

// d3 hardened: uses k_d3_hardened (16 RPB, deferred reduction)
void trit_gemv_d3_int8_dp4a(
    const int32_t* wt, const float* ws,
    const int32_t* xt, const float* xs,
    float* y, int cols, int rows, int num_groups,
    int use_l2_persist
) {
    if (trit_validate_gemv(wt, ws, y, cols, rows, num_groups)) return;
    if (!xt || !xs) { g_last_error = TRIT_ERR_NULL_PTR; return; }
    if (use_l2_persist) {
        set_l2_persist((void*)wt, (size_t)rows * num_groups * 16 * sizeof(int32_t));
    }
    k_d3_hardened<<<(rows + D3H_RPB - 1) / D3H_RPB, D3H_BS>>>(
        (const uint32_t*)wt, ws, (const uint32_t*)xt, xs, y, cols, rows, num_groups);
    trit_capture_launch_status();
    if (use_l2_persist) {
        clear_l2_persist();
    }
}

// Run the same layer N times back-to-back to measure pipeline / L2 reuse benefit
void trit_gemv_pipeline_bench(
    const int32_t* pt, const float* ws,
    const int32_t* xt_e, const int32_t* xt_o, const float* xs,
    float* y, int cols, int rows, int num_groups,
    int n_repeats, int use_l2_persist
) {
    if (trit_validate_gemv(pt, ws, y, cols, rows, num_groups)) return;
    if (!xt_e || !xt_o || !xs || n_repeats <= 0) { g_last_error = TRIT_ERR_NULL_PTR; return; }
    if (use_l2_persist) {
        set_l2_persist((void*)pt, (size_t)rows * num_groups * 8 * sizeof(int32_t));
    }
    // Launch n_repeats sequential v28 kernels in the SAME stream — measures
    // the L2-reuse benefit of back-to-back launches sharing weights.
    for (int i = 0; i < n_repeats; i++) {
        k_v28<<<(rows + V28_RPB - 1) / V28_RPB, V28_BS>>>(
            (const uint32_t*)pt, ws,
            (const uint32_t*)xt_e, (const uint32_t*)xt_o, xs,
            y, cols, rows, num_groups);
    }
    trit_capture_launch_status();
    if (use_l2_persist) {
        clear_l2_persist();
    }
}

// Query L2 cache size (for deciding whether to use L2 persist)
int get_l2_cache_bytes() {
    cudaDeviceProp prop;
    cudaGetDeviceProperties(&prop, 0);
    return prop.l2CacheSize;
}

// Query GPU name. `buf` must be a writable buffer of `buflen >= 1` bytes.
// On invalid input, the call is a no-op and g_last_error is set.
void get_gpu_name(char* buf, int buflen) {
    if (!buf || buflen <= 0) { g_last_error = TRIT_ERR_BAD_BUFFER; return; }
    cudaDeviceProp prop;
    cudaError_t e = cudaGetDeviceProperties(&prop, 0);
    if (e != cudaSuccess) { g_last_error = (int)e; buf[0] = '\0'; return; }
    strncpy(buf, prop.name, buflen - 1);
    buf[buflen - 1] = '\0';
}

// Synchronize (for timing from Python)
void cuda_sync() {
    cudaDeviceSynchronize();
}

}  // extern "C"