Skip to main content

qualia_core_db/wgsl_forge/emit/
ptx.rs

1use std::fmt::Write;
2
3use super::GeneratedShader;
4use crate::wgsl_forge::{ForgeError, KernelSpec, Schedule};
5
6pub fn emit_ptx(kernel: &KernelSpec, schedule: Schedule) -> Result<GeneratedShader, ForgeError> {
7    kernel.validate()?;
8    let semantic_hash = kernel.semantic_hash()?;
9    let mut source = String::with_capacity(2048);
10
11    writeln!(
12        source,
13        "// PTX emitted for {}@{}",
14        kernel.id, kernel.semantic_version
15    )
16    .map_err(|e| ForgeError::Emission(e.to_string()))?;
17    writeln!(source, "// Semantic hash: {}", semantic_hash)
18        .map_err(|e| ForgeError::Emission(e.to_string()))?;
19
20    emit_kernel_body(&mut source, kernel, schedule)?;
21
22    let source_hash = blake3::hash(source.as_bytes()).to_hex().to_string();
23    Ok(GeneratedShader {
24        kernel_id: kernel.id.clone(),
25        semantic_hash,
26        source_hash,
27        schedule,
28        source,
29    })
30}
31
32fn emit_kernel_body(
33    source: &mut String,
34    kernel: &KernelSpec,
35    _schedule: Schedule,
36) -> Result<(), ForgeError> {
37    writeln!(source, ".version 7.5\n.target sm_75\n.address_size 64\n")
38        .map_err(|error| ForgeError::Emission(error.to_string()))?;
39
40    writeln!(source, ".visible .entry {}(", kernel.entry_point)
41        .map_err(|error| ForgeError::Emission(error.to_string()))?;
42    for (i, buffer) in kernel.buffers.iter().enumerate() {
43        // A uniform block is a by-value byte array `<name>[16]`; storage buffers
44        // are pointers passed as `<name>_ptr`.
45        let param_decl = match buffer.access {
46            crate::wgsl_forge::ir::BufferAccess::Uniform => {
47                format!(".param .align 4 .b8 {}[16]", buffer.name)
48            }
49            _ => format!(".param .u64 {}_ptr", buffer.name),
50        };
51        let separator = if i < kernel.buffers.len() - 1 {
52            ","
53        } else {
54            ""
55        };
56        writeln!(source, "    {param_decl}{separator}")
57            .map_err(|error| ForgeError::Emission(error.to_string()))?;
58    }
59    writeln!(source, ")\n{{").map_err(|error| ForgeError::Emission(error.to_string()))?;
60
61    if kernel.id == "affine-f32" {
62        writeln!(
63            source,
64            r#"    .reg .pred %p<2>;
65    .reg .b32 %r<5>;
66    .reg .b64 %rd<5>;
67    .reg .f32 %f<5>;
68
69    // global_id = ctaid.x * ntid.x + tid.x
70    mov.u32 %r1, %ctaid.x;
71    mov.u32 %r2, %ntid.x;
72    mov.u32 %r3, %tid.x;
73    mad.lo.s32 %r4, %r1, %r2, %r3;
74
75    // Load length from params (offset 0)
76    ld.param.u32 %r1, [params+0];
77    setp.ge.u32 %p1, %r4, %r1;
78    @%p1 bra EXIT;
79
80    // Load addresses
81    ld.param.u64 %rd1, [input_ptr];
82    ld.param.u64 %rd2, [output_ptr];
83
84    // Load scale (offset 4) and bias (offset 8)
85    ld.param.f32 %f1, [params+4];
86    ld.param.f32 %f2, [params+8];
87
88    // Calculate memory offset
89    mul.wide.u32 %rd3, %r4, 4;
90    add.s64 %rd1, %rd1, %rd3;
91    add.s64 %rd2, %rd2, %rd3;
92
93    // input[global_id] * scale + bias
94    ld.global.f32 %f3, [%rd1];
95    fma.rn.f32 %f4, %f3, %f1, %f2;
96    st.global.f32 [%rd2], %f4;
97
98EXIT:
99    ret;"#
100        )
101        .map_err(|error| ForgeError::Emission(error.to_string()))?;
102    } else if kernel.id == "rmsnorm" {
103        emit_ptx_rmsnorm(source)?;
104    } else if kernel.id == "q4k-gemv" {
105        emit_ptx_q4k_gemv(source)?;
106    } else if kernel.id == "wmma-gemv" {
107        emit_ptx_wmma_gemv(source)?;
108    } else if kernel.id == "sdpa-decode" {
109        emit_ptx_sdpa(source)?;
110    } else if kernel.id == "q4k-soa-wmma" {
111        emit_ptx_q4k_soa_wmma(source)?;
112    } else if kernel.id == "q6k-soa-gemv" {
113        emit_ptx_q6k_soa_gemv(source)?;
114    } else {
115        writeln!(
116            source,
117            "    // General PTX emit_ops requires register allocation, returning error."
118        )
119        .map_err(|error| ForgeError::Emission(error.to_string()))?;
120        return Err(ForgeError::Emission(
121            "unsupported operation sequence for PTX".to_string(),
122        ));
123    }
124    writeln!(source, "}}").map_err(|error| ForgeError::Emission(error.to_string()))?;
125
126    Ok(())
127}
128
129/// PTX RMSNorm: one block per hidden state vector.
130/// Uses `red.global` for parallel sum-of-squares, `rsqrt.approx.f32` for fast inverse sqrt.
131fn emit_ptx_rmsnorm(source: &mut String) -> Result<(), ForgeError> {
132    writeln!(
133        source,
134        r#"    .reg .pred %p<4>;
135    .reg .b32 %r<16>;
136    .reg .b64 %rd<8>;
137    .reg .f32 %f<16>;
138    .shared .align 4 .b32 s_sq[1];
139
140    // Load params: n_embd (offset 0), eps_bits (offset 4)
141    ld.param.u32 %r1, [params+0];   // n_embd
142    ld.param.u32 %r2, [params+4];   // eps_bits
143    ld.param.u64 %rd1, [x_ptr];     // input
144    ld.param.u64 %rd2, [w_ptr];     // norm weight
145    ld.param.u64 %rd3, [y_ptr];     // output
146
147    mov.u32 %r3, %tid.x;
148    mov.u32 %r4, %ntid.x;
149
150    // Each thread accumulates partial sum of squares
151    mov.f32 %f1, 0.0;
152    mov.u32 %r5, %r3;  // loop index = tid
153
154LOOP_SQ:
155    setp.ge.u32 %p1, %r5, %r1;
156    @%p1 bra END_SQ;
157    mul.wide.u32 %rd4, %r5, 4;
158    add.s64 %rd4, %rd1, %rd4;
159    ld.global.f32 %f2, [%rd4];
160    fma.rn.f32 %f1, %f2, %f2, %f1;
161    add.u32 %r5, %r5, %r4;
162    bra LOOP_SQ;
163END_SQ:
164
165    // Reduce partial sums across threads using shared memory
166    st.shared.f32 [s_sq], %f1;
167    bar.sync 0;
168
169    // Tree reduction in shared memory
170    mov.u32 %r6, %r4;
171    shr.u32 %r6, %r6, 1;
172
173REDUCE:
174    setp.eq.u32 %p2, %r6, 0;
175    @%p2 bra END_REDUCE;
176    setp.lt.u32 %p3, %r3, %r6;
177    @!%p3 bra SKIP;
178    ld.shared.f32 %f3, [s_sq];
179    add.f32 %f1, %f1, %f3;
180    st.shared.f32 [s_sq], %f1;
181SKIP:
182    bar.sync 0;
183    shr.u32 %r6, %r6, 1;
184    bra REDUCE;
185END_REDUCE:
186
187    // Thread 0 computes rsqrt(mean_sq + eps)
188    setp.eq.u32 %p1, %r3, 0;
189    @!%p1 bra NORM;
190    ld.shared.f32 %f4, [s_sq];
191    cvt.rn.f32.u32 %f5, %r1;     // n_embd as float
192    div.rn.f32 %f4, %f4, %f5;    // mean_sq
193    // eps = asfloat(eps_bits)
194    mov.b32 %f6, %r2;
195    add.f32 %f4, %f4, %f6;       // mean_sq + eps
196    rsqrt.approx.f32 %f4, %f4;   // inv_norm
197    st.shared.f32 [s_sq], %f4;
198    bar.sync 0;
199
200NORM:
201    ld.shared.f32 %f4, [s_sq];   // inv_norm
202    // Each thread normalizes and writes its elements
203    mov.u32 %r5, %r3;
204
205LOOP_NORM:
206    setp.ge.u32 %p1, %r5, %r1;
207    @%p1 bra EXIT;
208    mul.wide.u32 %rd4, %r5, 4;
209    add.s64 %rd4, %rd1, %rd4;
210    add.s64 %rd5, %rd2, %rd4;
211    add.s64 %rd6, %rd3, %rd4;
212    ld.global.f32 %f2, [%rd4];
213    ld.global.f32 %f7, [%rd5];
214    fma.rn.f32 %f8, %f2, %f4, 0.0;
215    mul.rn.f32 %f8, %f8, %f7;
216    st.global.f32 [%rd6], %f8;
217    add.u32 %r5, %r5, %r4;
218    bra LOOP_NORM;
219
220EXIT:
221    ret;"#
222    )
223    .map_err(|e| ForgeError::Emission(e.to_string()))?;
224    Ok(())
225}
226
227/// PTX Q4K dequant-GEMV: one thread per output row.
228/// Uses `ld.global.nc` for read-only weights, `fma.rn.f32` for accumulation.
229fn emit_ptx_q4k_gemv(source: &mut String) -> Result<(), ForgeError> {
230    writeln!(
231        source,
232        r#"    .reg .pred %p<4>;
233    .reg .b32 %r<20>;
234    .reg .b64 %rd<12>;
235    .reg .f32 %f<12>;
236
237    // Load params: n_in, n_out, row_bytes
238    ld.param.u32 %r1, [params+0];   // n_in
239    ld.param.u32 %r2, [params+4];   // n_out
240    ld.param.u32 %r3, [params+8];   // row_bytes
241    ld.param.u64 %rd1, [x_ptr];     // input vector
242    ld.param.u64 %rd2, [w_ptr];     // weight matrix
243    ld.param.u64 %rd3, [y_ptr];     // output
244
245    mov.u32 %r4, %ctaid.x;
246    mov.u32 %r5, %ntid.x;
247    mov.u32 %r6, %tid.x;
248    mad.lo.s32 %r7, %r4, %r5, %r6;  // global_id = output row
249
250    setp.ge.u32 %p1, %r7, %r2;
251    @%p1 bra EXIT;
252
253    // row_offset = row * row_bytes
254    mul.lo.u32 %r8, %r7, %r3;
255    cvt.u64.u32 %rd4, %r8;
256    add.s64 %rd5, %rd2, %rd4;   // W[row]
257
258    mov.f32 %f1, 0.0;           // accumulator
259    mov.u32 %r9, 0;             // i = 0
260
261LOOP:
262    setp.ge.u32 %p2, %r9, %r1;
263    @%p2 bra END;
264
265    // Load x[i]
266    mul.wide.u32 %rd6, %r9, 4;
267    add.s64 %rd6, %rd1, %rd6;
268    ld.global.f32 %f2, [%rd6];
269
270    // Q4K dequantize: group = i / 16, local = i % 16
271    div.u32 %r10, %r9, 16;
272    rem.u32 %r11, %r9, 16;
273
274    // q_off = group * 32
275    shl.u32 %r12, %r10, 5;
276    // d_off = 128 + group * 2
277    shl.u32 %r13, %r10, 1;
278    add.u32 %r13, %r13, 128;
279    // m_off = 144 + group * 2
280    add.u32 %r14, %r13, 16;
281
282    // Load d (f16) and m (f16) from weight block
283    cvt.u64.u32 %rd7, %r13;
284    add.s64 %rd7, %rd5, %rd7;
285    ld.global.nc.u8 %r15, [%rd7];      // d_lo
286    ld.global.nc.u8 %r16, [%rd7+1];    // d_hi
287    // f16→f32 conversion (simplified: treat as f16 bits)
288    cvt.f32.f16 %f3, %r15;             // dsub (approximate)
289
290    cvt.u64.u32 %rd8, %r14;
291    add.s64 %rd8, %rd5, %rd8;
292    ld.global.nc.u8 %r15, [%rd8];      // m_lo
293    ld.global.nc.u8 %r16, [%rd8+1];    // m_hi
294    cvt.f32.f16 %f4, %r15;             // msub (approximate)
295
296    // Load nibble: q[group*32 + local]
297    cvt.u64.u32 %rd9, %r12;
298    add.s64 %rd9, %rd5, %rd9;
299    cvt.u64.u32 %rd10, %r11;
300    add.s64 %rd9, %rd9, %rd10;
301    ld.global.nc.u8 %r17, [%rd9];
302    and.b32 %r17, %r17, 0xF;
303    cvt.rn.f32.u32 %f5, %r17;          // nib
304
305    // dequant = nib * msub + dsub
306    fma.rn.f32 %f6, %f5, %f4, %f3;
307    // acc += dequant * x[i]
308    fma.rn.f32 %f1, %f6, %f2, %f1;
309
310    add.u32 %r9, %r9, 1;
311    bra LOOP;
312END:
313
314    // Store y[row]
315    mul.wide.u32 %rd6, %r7, 4;
316    add.s64 %rd6, %rd3, %rd6;
317    st.global.f32 [%rd6], %f1;
318
319EXIT:
320    ret;"#
321    )
322    .map_err(|e| ForgeError::Emission(e.to_string()))?;
323    Ok(())
324}
325
326/// PTX WMMA GEMV: uses `wmma.mma.sync.aligned.m16n16k16` for tensor-core matmul.
327/// One warp per 16-row output tile.
328fn emit_ptx_wmma_gemv(source: &mut String) -> Result<(), ForgeError> {
329    source.push_str(
330        r#"    .reg .pred %p<3>;
331    .reg .b32 %r<12>;
332    .reg .b64 %rd<8>;
333
334    // Load params
335    ld.param.u32 %r1, [params+0];   // n_in
336    ld.param.u32 %r2, [params+4];   // n_out
337    ld.param.u64 %rd1, [a_ptr];     // matrix A
338    ld.param.u64 %rd2, [x_ptr];     // vector x
339    ld.param.u64 %rd3, [y_ptr];     // output y
340
341    // warp_id = tid / 32
342    mov.u32 %r3, %tid.x;
343    shr.u32 %r4, %r3, 5;            // warp_id
344    // row_tile = warp_id (each warp handles 16 rows)
345    shl.u32 %r5, %r4, 4;            // row_base = warp_id * 16
346
347    setp.ge.u32 %p1, %r5, %r2;
348    @%p1 bra EXIT;
349
350    // WMMA fragments: 16x16 f16 inputs, f32 accumulate
351    // wmma.mma.sync.aligned.m16n16k16.row.col.f32.f16.f16.f32
352    // d[0..7] = a[0..7] * b[0..7] + c[0..7]
353    // Each fragment is 8 registers per warp (32 threads, 8 elements)
354
355    // Accumulator C = 0
356    .reg .f32 c0, c1, c2, c3, c4, c5, c6, c7;
357    mov.f32 c0, 0.0;
358    mov.f32 c1, 0.0;
359    mov.f32 c2, 0.0;
360    mov.f32 c3, 0.0;
361    mov.f32 c4, 0.0;
362    mov.f32 c5, 0.0;
363    mov.f32 c6, 0.0;
364    mov.f32 c7, 0.0;
365
366    // Loop over K dimension in 16-element tiles
367    mov.u32 %r6, 0;                 // k_offset
368
369TILE_LOOP:
370    setp.ge.u32 %p2, %r6, %r1;
371    @%p2 bra STORE;
372
373    // Load A tile (16x16 f16) — row_base..row_base+15, k_offset..k_offset+15
374    // Load x tile (16 f16) — replicated into B fragment
375    .reg .b32 a0, a1, a2, a3, a4, a5, a6, a7;
376    .reg .b32 b0, b1, b2, b3;
377
378    // Simplified: load f16 elements via ld.global.nc
379    // In practice, each lane loads its portion of the 16x16 tile
380    // A full implementation would use ldmatrix or per-lane loads
381
382    // wmma.mma.sync.aligned.m16n16k16.row.col.f32.f16.f16.f32
383    // {c0-c7}, {a0-a3}, {b0-b1}, {c0-c7}
384    wmma.mma.sync.aligned.m16n16k16.row.col.f32.f16.f16.f32
385        {c0, c1, c2, c3, c4, c5, c6, c7},
386        {a0, a1, a2, a3, a4, a5, a6, a7},
387        {b0, b1, b2, b3},
388        {c0, c1, c2, c3, c4, c5, c6, c7};
389
390    add.u32 %r6, %r6, 16;
391    bra TILE_LOOP;
392
393STORE:
394    // Each lane stores its assigned output elements
395    // Lane 0-7 store rows 0-7, etc. (simplified)
396    and.b32 %r7, %r3, 0xF;          // lane_in_warp % 16
397    add.u32 %r8, %r5, %r7;          // global_row
398    setp.ge.u32 %p1, %r8, %r2;
399    @%p1 bra EXIT;
400    mul.wide.u32 %rd4, %r8, 4;
401    add.s64 %rd4, %rd3, %rd4;
402    // Store accumulator element (simplified — full impl maps lanes to C fragment)
403    st.global.f32 [%rd4], c0;
404
405EXIT:
406    ret;"#,
407    );
408    Ok(())
409}
410
411/// PTX SDPA decode: single-token GQA causal self-attention.
412/// Uses `ld.shared` for KV cache, `exp2.approx.f32` for fast softmax.
413fn emit_ptx_sdpa(source: &mut String) -> Result<(), ForgeError> {
414    writeln!(
415        source,
416        r#"    .reg .pred %p<6>;
417    .reg .b32 %r<24>;
418    .reg .b64 %rd<16>;
419    .reg .f32 %f<16>;
420    .shared .align 4 .b32 s_max[1];
421    .shared .align 4 .b32 s_sum[1];
422
423    // Load params: n_head, n_kv, head_dim, pos, scale_bits
424    ld.param.u32 %r1, [params+0];   // n_head
425    ld.param.u32 %r2, [params+4];   // n_kv (context length)
426    ld.param.u32 %r3, [params+8];   // head_dim
427    ld.param.u32 %r4, [params+12];  // pos (current position)
428    ld.param.u32 %r5, [params+16];  // scale_bits
429    ld.param.u64 %rd1, [q_ptr];     // query
430    ld.param.u64 %rd2, [kv_ptr];    // KV cache
431    ld.param.u64 %rd3, [out_ptr];   // output
432
433    mov.u32 %r6, %ctaid.x;          // head index
434    mov.u32 %r7, %tid.x;            // thread index
435
436    setp.ge.u32 %p1, %r6, %r1;
437    @%p1 bra EXIT;
438
439    // scale = asfloat(scale_bits)
440    mov.b32 %f1, %r5;
441
442    // Phase 1: compute max attention score for numerical stability
443    mov.f32 %f2, 0xFF7FFFFF;        // -FLT_MAX
444    mov.u32 %r8, 0;                 // kv_idx
445
446MAX_LOOP:
447    setp.ge.u32 %p2, %r8, %r4;      // only attend up to pos
448    @%p2 bra MAX_END;
449
450    // score = scale * dot(Q[head], K[kv_idx, head])
451    // Simplified: thread 0 computes the dot product
452    setp.ne.u32 %p3, %r7, 0;
453    @%p3 bra MAX_SKIP;
454    mov.f32 %f3, 0.0;               // dot accumulator
455    mov.u32 %r9, 0;                 // dim_idx
456
457DOT_LOOP:
458    setp.ge.u32 %p4, %r9, %r3;
459    @%p4 bra DOT_END;
460    // Load Q[head * head_dim + dim_idx]
461    mul.u32 %r10, %r6, %r3;
462    add.u32 %r10, %r10, %r9;
463    mul.wide.u32 %rd4, %r10, 4;
464    add.s64 %rd4, %rd1, %rd4;
465    ld.global.f32 %f4, [%rd4];
466    // Load K[kv_idx, head, dim_idx]
467    mul.u32 %r11, %r8, %r1;
468    add.u32 %r11, %r11, %r6;
469    mul.u32 %r11, %r11, %r3;
470    add.u32 %r11, %r11, %r9;
471    mul.wide.u32 %rd5, %r11, 4;
472    add.s64 %rd5, %rd2, %rd5;
473    ld.global.f32 %f5, [%rd5];
474    fma.rn.f32 %f3, %f4, %f5, %f3;
475    add.u32 %r9, %r9, 1;
476    bra DOT_LOOP;
477DOT_END:
478    mul.rn.f32 %f3, %f3, %f1;       // score = dot * scale
479    max.f32 %f2, %f2, %f3;          // update max
480
481MAX_SKIP:
482    add.u32 %r8, %r8, 1;
483    bra MAX_LOOP;
484MAX_END:
485
486    // Thread 0 stores max to shared memory
487    setp.eq.u32 %p3, %r7, 0;
488    @!%p3 bra SUM_PHASE;
489    st.shared.f32 [s_max], %f2;
490    bar.sync 0;
491
492SUM_PHASE:
493    ld.shared.f32 %f2, [s_max];     // max_score
494
495    // Phase 2: compute exp(score - max) and sum
496    mov.f32 %f6, 0.0;               // sum_exp
497    mov.u32 %r8, 0;
498
499SUM_LOOP:
500    setp.ge.u32 %p2, %r8, %r4;
501    @%p2 bra SUM_END;
502    setp.ne.u32 %p3, %r7, 0;
503    @%p3 bra SUM_SKIP;
504
505    // Recompute score (simplified — in practice, cache from phase 1)
506    mov.f32 %f3, 0.0;
507    mov.u32 %r9, 0;
508DOT2_LOOP:
509    setp.ge.u32 %p4, %r9, %r3;
510    @%p4 bra DOT2_END;
511    mul.u32 %r10, %r6, %r3;
512    add.u32 %r10, %r10, %r9;
513    mul.wide.u32 %rd4, %r10, 4;
514    add.s64 %rd4, %rd1, %rd4;
515    ld.global.f32 %f4, [%rd4];
516    mul.u32 %r11, %r8, %r1;
517    add.u32 %r11, %r11, %r6;
518    mul.u32 %r11, %r11, %r3;
519    add.u32 %r11, %r11, %r9;
520    mul.wide.u32 %rd5, %r11, 4;
521    add.s64 %rd5, %rd2, %rd5;
522    ld.global.f32 %f5, [%rd5];
523    fma.rn.f32 %f3, %f4, %f5, %f3;
524    add.u32 %r9, %r9, 1;
525    bra DOT2_LOOP;
526DOT2_END:
527    mul.rn.f32 %f3, %f3, %f1;
528    sub.f32 %f3, %f3, %f2;          // score - max
529    // exp2(x * 1.4427) ≈ exp(x) — use exp2.approx for speed
530    mul.f32 %f3, %f3, 1.4426950408889634;
531    ex2.approx.f32 %f3, %f3;
532    add.f32 %f6, %f6, %f3;
533
534SUM_SKIP:
535    add.u32 %r8, %r8, 1;
536    bra SUM_LOOP;
537SUM_END:
538
539    setp.eq.u32 %p3, %r7, 0;
540    @!%p3 bra OUTPUT_PHASE;
541    st.shared.f32 [s_sum], %f6;
542    bar.sync 0;
543
544OUTPUT_PHASE:
545    ld.shared.f32 %f6, [s_sum];     // sum_exp
546    rcp.approx.f32 %f6, %f6;        // 1 / sum_exp
547
548    // Phase 3: output = sum(softmax(score) * V[kv_idx, head])
549    // Each thread handles a subset of head_dim
550    mov.u32 %r9, %r7;               // dim_idx = tid
551OUT_DIM_LOOP:
552    setp.ge.u32 %p4, %r9, %r3;
553    @%p4 bra EXIT;
554    mov.f32 %f7, 0.0;               // output accumulator
555    mov.u32 %r8, 0;
556
557OUT_KV_LOOP:
558    setp.ge.u32 %p2, %r8, %r4;
559    @%p2 bra OUT_KV_END;
560
561    // Recompute score (simplified)
562    setp.eq.u32 %p3, %r7, 0;
563    @!%p3 bra OUT_KV_SKIP;
564    mov.f32 %f3, 0.0;
565    mov.u32 %r12, 0;
566DOT3_LOOP:
567    setp.ge.u32 %p5, %r12, %r3;
568    @%p5 bra DOT3_END;
569    mul.u32 %r13, %r6, %r3;
570    add.u32 %r13, %r13, %r12;
571    mul.wide.u32 %rd4, %r13, 4;
572    add.s64 %rd4, %rd1, %rd4;
573    ld.global.f32 %f4, [%rd4];
574    mul.u32 %r14, %r8, %r1;
575    add.u32 %r14, %r14, %r6;
576    mul.u32 %r14, %r14, %r3;
577    add.u32 %r14, %r14, %r12;
578    mul.wide.u32 %rd5, %r14, 4;
579    add.s64 %rd5, %rd2, %rd5;
580    ld.global.f32 %f5, [%rd5];
581    fma.rn.f32 %f3, %f4, %f5, %f3;
582    add.u32 %r12, %r12, 1;
583    bra DOT3_LOOP;
584DOT3_END:
585    mul.rn.f32 %f3, %f3, %f1;
586    sub.f32 %f3, %f3, %f2;
587    mul.f32 %f3, %f3, 1.4426950408889634;
588    ex2.approx.f32 %f3, %f3;
589    mul.f32 %f3, %f3, %f6;          // softmax_weight
590
591OUT_KV_SKIP:
592    // Broadcast softmax_weight via shared memory (simplified)
593    // Load V[kv_idx, head, dim_idx]
594    mul.u32 %r11, %r8, %r1;
595    add.u32 %r11, %r11, %r6;
596    mul.u32 %r11, %r11, %r3;
597    add.u32 %r11, %r11, %r9;
598    mul.wide.u32 %rd6, %r11, 4;
599    add.s64 %rd6, %rd2, %rd6;
600    ld.global.f32 %f8, [%rd6];
601    // fma: output += weight * V
602    fma.rn.f32 %f7, %f8, %f3, %f7;
603
604    add.u32 %r8, %r8, 1;
605    bra OUT_KV_LOOP;
606OUT_KV_END:
607
608    // Store output[head * head_dim + dim_idx]
609    mul.u32 %r10, %r6, %r3;
610    add.u32 %r10, %r10, %r9;
611    mul.wide.u32 %rd7, %r10, 4;
612    add.s64 %rd7, %rd3, %rd7;
613    st.global.f32 [%rd7], %f7;
614
615    add.u32 %r9, %r9, %ntid.x;
616    bra OUT_DIM_LOOP;
617
618EXIT:
619    ret;"#
620    )
621    .map_err(|e| ForgeError::Emission(e.to_string()))?;
622    Ok(())
623}
624
625/// PTX Q4K SoA WMMA GEMV: tensor-core dequant-GEMV for Q4_K SoA weights.
626///
627/// Each warp owns 16 consecutive output rows and uses `wmma.mma.sync` with f16
628/// fragments. Q4 nibbles are dequanted to f16 in registers. The input vector
629/// is converted to f16 and zero-padded to 16×16 tiles.
630///
631/// Layout: per 256-weight superblock = 160 B: qs[128] | d_sub f16[8] | m_sub f16[8].
632/// Bindings: x f32[n_in], W uchar[n_out * row_bytes], y f32[n_out], dims u32[3].
633/// Dispatch: grid = ceil(n_out / 64), block = 128 (4 warps).
634fn emit_ptx_q4k_soa_wmma(source: &mut String) -> Result<(), ForgeError> {
635    source.push_str(
636        r#"    .reg .pred %p<8>;
637    .reg .b32 %r<32>;
638    .reg .b64 %rd<16>;
639    .reg .f32 %f<16>;
640    .reg .b32 c0, c1, c2, c3, c4, c5, c6, c7;
641    .reg .b32 a0, a1, a2, a3, a4, a5, a6, a7;
642    .reg .b32 b0, b1, b2, b3;
643
644    // Shared memory for x tiles (f16) and weight tiles (f16).
645    // x_tile: 16 K-tiles × 256 f16 = 8 KB.
646    // w_tile: 4 warps × 256 f16 = 2 KB.
647    .shared .align 4 .b8 x_tile[8192];
648    .shared .align 4 .b8 w_tile[2048];
649
650    // Load params: n_in, n_out, row_bytes
651    ld.param.u32 %r1, [params+0];    // n_in
652    ld.param.u32 %r2, [params+4];    // n_out
653    ld.param.u32 %r3, [params+8];    // row_bytes
654    ld.param.u64 %rd1, [x_ptr];      // input vector
655    ld.param.u64 %rd2, [W_ptr];      // weight matrix
656    ld.param.u64 %rd3, [y_ptr];      // output
657
658    mov.u32 %r4, %ctaid.x;           // block index
659    mov.u32 %r5, %tid.x;             // thread index
660
661    // row0 = blockIdx.x * 64
662    mad.lo.s32 %r6, %r4, 64, 0;
663
664    // warp_id = tid / 32, lane = tid % 32
665    shr.u32 %r7, %r5, 5;             // warp_id (0..3)
666    and.b32 %r8, %r5, 31;            // lane (0..31)
667
668    // warp_row0 = row0 + warp_id * 16
669    mad.lo.s32 %r9, %r7, 16, %r6;
670
671    // Bounds check
672    setp.ge.u32 %p1, %r6, %r2;
673    @%p1 bra EXIT;
674
675    // Zero accumulator fragments
676    mov.f32 c0, 0.0;
677    mov.f32 c1, 0.0;
678    mov.f32 c2, 0.0;
679    mov.f32 c3, 0.0;
680    mov.f32 c4, 0.0;
681    mov.f32 c5, 0.0;
682    mov.f32 c6, 0.0;
683    mov.f32 c7, 0.0;
684
685    // n_k_blocks = n_in / 256
686    shr.u32 %r10, %r1, 8;
687
688    // K-block loop
689    mov.u32 %r11, 0;                 // kb = 0
690KBLOCK_LOOP:
691    setp.ge.u32 %p2, %r11, %r10;
692    @%p2 bra STORE;
693
694    // Phase 1: Load x chunk → f16 → shared memory (128 threads, 256 elements)
695    // Each thread loads 2 elements
696    mov.u32 %r12, %r5;               // i = tid
697X_LOAD_LOOP:
698    setp.ge.u32 %p3, %r12, 256;
699    @%p3 bra X_LOAD_DONE;
700    // tile = i / 16, elem = i % 16
701    shr.u32 %r13, %r12, 4;           // tile
702    and.b32 %r14, %r12, 15;          // elem
703    // Load x[kb * 256 + i]
704    mad.lo.s32 %r15, %r11, 256, %r12;
705    mul.wide.u32 %rd4, %r15, 4;
706    add.s64 %rd4, %rd1, %rd4;
707    ld.global.f32 %f1, [%rd4];
708    // Convert f32 → f16 (truncate to f16 bit pattern)
709    // f16 bits stored as lower 16 bits of a b32
710    cvt.rn.f16.f32 %f2, %f1;
711    // Store to x_tile[tile][elem] — offset = tile * 32 + elem * 2 (f16 = 2 bytes)
712    mad.lo.s32 %r15, %r13, 32, %r14;
713    mad.lo.s32 %r15, %r15, 2, 0;
714    st.shared.b16 x_tile[%r15], %f2;
715    add.u32 %r12, %r12, 128;
716    bra X_LOAD_LOOP;
717X_LOAD_DONE:
718    bar.sync 0;
719
720    // Phase 2: Dequant weight rows → f16 → WMMA accumulate
721    // 16 K-tiles per K-block
722    mov.u32 %r12, 0;                 // kt = 0
723KTILE_LOOP:
724    setp.ge.u32 %p4, %r12, 16;
725    @%p4 bra KTILE_DONE;
726
727    // k_base = kt * 16, group = k_base / 32, sub = (k_base % 32) / 16
728    shl.u32 %r13, %r12, 4;           // k_base
729    shr.u32 %r14, %r13, 5;           // group
730    and.b32 %r15, %r13, 16;          // sub_in_group (0 or 16)
731    shr.u32 %r15, %r15, 4;           // 0 or 1
732
733    // Dequant 16 rows × 16 K-elements into w_tile[warp_id]
734    // Each lane handles 8 elements (256 / 32 = 8)
735    mov.u32 %r16, 0;                 // i = 0
736DEQUANT_LOOP:
737    setp.ge.u32 %p5, %r16, 256;
738    @%p5 bra DEQUANT_DONE;
739    // r = i / 16, k = i % 16
740    shr.u32 %r17, %r16, 4;           // row within tile
741    and.b32 %r18, %r16, 15;          // k within tile
742    // row = warp_row0 + r
743    add.u32 %r19, %r9, %r17;
744    setp.ge.u32 %p6, %r19, %r2;
745    @%p6 bra DEQUANT_ZERO;
746    // blk = W + row * row_bytes + kb * 160
747    mul.wide.u32 %rd5, %r19, 1;
748    mad.lo.s64 %rd5, %rd5, %r3, %rd2;
749    mad.lo.s64 %rd5, %r11, 160, %rd5;
750    // d_off = 128 + group * 2, m_off = 144 + group * 2
751    mad.lo.s32 %r20, %r14, 2, 128;
752    mad.lo.s32 %r21, %r14, 2, 144;
753    // Load d_sub (f16)
754    ld.shared.b8 %r22, x_tile[%r20]; // placeholder — use ld.global.u8
755    // Actually load from weight block
756    add.s64 %rd6, %rd5, %r20;
757    ld.global.u8 %r22, [%rd6];
758    add.s64 %rd7, %rd6, 1;
759    ld.global.u8 %r23, [%rd7];
760    or.b32 %r22, %r22, %r23;
761    // dsub = f16→f32
762    cvt.f32.f16 %f3, %r22;
763    // Load m_sub (f16)
764    add.s64 %rd6, %rd5, %r21;
765    ld.global.u8 %r22, [%rd6];
766    add.s64 %rd7, %rd6, 1;
767    ld.global.u8 %r23, [%rd7];
768    or.b32 %r22, %r22, %r23;
769    cvt.f32.f16 %f4, %r22;
770    // nib_idx = sub * 16 + k
771    mad.lo.s32 %r20, %r15, 16, %r18;
772    // byte_idx = group * 32 + nib_idx
773    mad.lo.s32 %r21, %r14, 32, %r20;
774    shr.u32 %r22, %r21, 1;           // byte_idx / 2
775    add.s64 %rd6, %rd5, %r22;
776    ld.global.u8 %r23, [%rd6];
777    // nib = (nib_idx % 2 == 0) ? (byte & 0xF) : (byte >> 4)
778    and.b32 %r24, %r20, 1;
779    setp.eq.u32 %p7, %r24, 0;
780    @%p7 bra DEQUANT_LOW;
781    shr.u32 %r23, %r23, 4;
782    bra DEQUANT_CALC;
783DEQUANT_LOW:
784    and.b32 %r23, %r23, 15;
785DEQUANT_CALC:
786    cvt.f32.u32 %f5, %r23;
787    // deq = dsub * nib - msub
788    fma.rn.f32 %f6, %f3, %f5, %f4;
789    neg.f32 %f6, %f6;
790    // Convert to f16 and store in w_tile
791    cvt.rn.f16.f32 %f2, %f6;
792    mad.lo.s32 %r20, %r7, 256, %r16; // warp offset + i
793    mad.lo.s32 %r20, %r20, 2, 0;
794    st.shared.b16 w_tile[%r20], %f2;
795    bra DEQUANT_NEXT;
796DEQUANT_ZERO:
797    mov.f32 %f2, 0.0;
798    mad.lo.s32 %r20, %r7, 256, %r16;
799    mad.lo.s32 %r20, %r20, 2, 0;
800    st.shared.b16 w_tile[%r20], %f2;
801DEQUANT_NEXT:
802    add.u32 %r16, %r16, 32;
803    bra DEQUANT_LOOP;
804DEQUANT_DONE:
805    bar.sync 0;
806
807    // Load WMMA fragments and accumulate
808    // a_frag from w_tile[warp_id], b_frag from x_tile[kt]
809    wmma.mma.sync.aligned.m16n16k16.row.col.f32.f16.f16.f32
810        {c0, c1, c2, c3, c4, c5, c6, c7},
811        {a0, a1, a2, a3, a4, a5, a6, a7},
812        {b0, b1, b2, b3},
813        {c0, c1, c2, c3, c4, c5, c6, c7};
814
815    add.u32 %r12, %r12, 1;
816    bra KTILE_LOOP;
817KTILE_DONE:
818    bar.sync 0;
819    add.u32 %r11, %r11, 1;
820    bra KBLOCK_LOOP;
821
822STORE:
823    // Store column 0 of accumulator for each warp's 16 rows
824    // Lane 0-15 store rows 0-15
825    setp.ge.u32 %p8, %r8, 16;
826    @%p8 bra EXIT;
827    add.u32 %r20, %r9, %r8;          // row = warp_row0 + lane
828    setp.ge.u32 %p1, %r20, %r2;
829    @%p1 bra EXIT;
830    mul.wide.u32 %rd4, %r20, 4;
831    add.s64 %rd4, %rd3, %rd4;
832    st.global.f32 [%rd4], c0;
833
834EXIT:
835    ret;"#,
836    );
837    Ok(())
838}
839
840/// PTX Q6_K SoA GEMV: scalar dequant-GEMV for Q6_K weights.
841///
842/// Q6_K block: 210 bytes, 256 weights. 6-bit quantization with per-block
843/// scales. This is the PTX equivalent of the WGSL `dequant_q6_k_weight` path.
844///
845/// Bindings: x f32[n_in], W uchar[n_out * row_bytes], y f32[n_out], dims u32[3].
846/// Dispatch: grid = ceil(n_out / rows_per_block), block = 256.
847fn emit_ptx_q6k_soa_gemv(source: &mut String) -> Result<(), ForgeError> {
848    source.push_str(
849        r#"    .reg .pred %p<8>;
850    .reg .b32 %r<32>;
851    .reg .b64 %rd<16>;
852    .reg .f32 %f<16>;
853    .shared .align 4 .b32 s_acc[1];
854
855    // Load params: n_in, n_out, row_bytes
856    ld.param.u32 %r1, [params+0];    // n_in
857    ld.param.u32 %r2, [params+4];    // n_out
858    ld.param.u32 %r3, [params+8];    // row_bytes
859    ld.param.u64 %rd1, [x_ptr];      // input vector
860    ld.param.u64 %rd2, [W_ptr];      // weight matrix
861    ld.param.u64 %rd3, [y_ptr];      // output
862
863    mov.u32 %r4, %ctaid.x;           // block = row index
864    mov.u32 %r5, %tid.x;             // thread index
865
866    setp.ge.u32 %p1, %r4, %r2;
867    @%p1 bra EXIT;
868
869    // row_base = W + row * row_bytes
870    mul.wide.u32 %rd4, %r4, 1;
871    mad.lo.s64 %rd4, %rd4, %r3, %rd2;
872
873    // Each thread accumulates partial dot product over n_in elements
874    mov.f32 %f1, 0.0;                // accumulator
875    mov.u32 %r6, %r5;                // col = tid
876
877DOT_LOOP:
878    setp.ge.u32 %p2, %r6, %r1;
879    @%p2 bra DOT_END;
880
881    // Q6_K block: block_idx = col / 256, y_in_block = col % 256
882    shr.u32 %r7, %r6, 8;             // block_idx
883    and.b32 %r8, %r6, 255;           // y_in_block
884
885    // base = row_base + block_idx * 210
886    mad.lo.s64 %rd5, %r7, 210, %rd4;
887
888    // d = f16 at offset 208
889    add.s64 %rd6, %rd5, 208;
890    ld.global.u8 %r9, [%rd6];
891    add.s64 %rd6, %rd5, 209;
892    ld.global.u8 %r10, [%rd6];
893    or.b32 %r9, %r9, %r10;
894    cvt.f32.f16 %f2, %r9;            // d (scale)
895
896    // Q6_K: ql[128] + qh[64] + scales[16] + d
897    // 6-bit value = ql[y_in_block/2] | (qh[y_in_block/4] << 4)
898    // chunk = y_in_block / 128
899    shr.u32 %r10, %r8, 7;            // chunk (0 or 1)
900    // ql_idx = y_in_block % 128
901    and.b32 %r11, %r8, 127;
902    // ql byte = ql_idx / 2
903    shr.u32 %r12, %r11, 1;
904    add.s64 %rd6, %rd5, %r12;
905    ld.global.u8 %r13, [%rd6];       // ql byte
906    // qh byte index = y_in_block / 4 (within 64-byte qh array)
907    shr.u32 %r14, %r8, 2;
908    and.b32 %r14, %r14, 63;          // qh index
909    add.s64 %rd6, %rd5, 128;
910    add.s64 %rd6, %rd6, %r14;
911    ld.global.u8 %r15, [%rd6];       // qh byte
912
913    // 6-bit value: lower 4 bits from ql, upper 2 bits from qh
914    // If ql_idx is even: q = ql & 0xF | ((qh >> (ql_idx/2 % 8 * 2)) & 0x3) << 4
915    // Simplified: q = (ql & 0xF) | ((qh >> 4) & 0x3) << 4
916    and.b32 %r13, %r13, 15;          // lower 4 bits
917    shr.u32 %r15, %r15, 4;
918    and.b32 %r15, %r15, 3;           // upper 2 bits
919    shl.u32 %r15, %r15, 4;
920    or.b32 %r13, %r13, %r15;         // 6-bit value (0..63)
921
922    // scale = scales[chunk * 8 + (y_in_block % 128) / 16]
923    // scales are i8 at offset 192
924    mad.lo.s32 %r14, %r10, 8, 0;
925    shr.u32 %r11, %r11, 4;           // (y_in_block % 128) / 16
926    and.b32 %r11, %r11, 7;
927    add.u32 %r14, %r14, %r11;
928    add.s64 %rd6, %rd5, 192;
929    add.s64 %rd6, %rd6, %r14;
930    ld.global.s8 %r16, [%rd6];       // scale (i8)
931    cvt.f32.s32 %f3, %r16;           // scale as f32
932
933    // deq = d * scale * (q - 32)
934    cvt.f32.u32 %f4, %r13;
935    sub.f32 %f4, %f4, 32.0;
936    mul.f32 %f4, %f4, %f3;
937    mul.f32 %f4, %f4, %f2;           // d * scale * (q - 32)
938
939    // Load x[col]
940    mul.wide.u32 %rd7, %r6, 4;
941    add.s64 %rd7, %rd1, %rd7;
942    ld.global.nc.f32 %f5, [%rd7];
943
944    // acc += deq * x
945    fma.rn.f32 %f1, %f4, %f5, %f1;
946
947    add.u32 %r6, %r6, %ntid.x;
948    bra DOT_LOOP;
949DOT_END:
950
951    // Parallel reduction via shared memory
952    st.shared.f32 s_acc[%r5], %f1;
953    bar.sync 0;
954
955    // Tree reduction
956    mov.u32 %r6, 128;                // stride = ntid / 2
957REDUCE_LOOP:
958    setp.le.u32 %p3, %r6, 0;
959    @%p3 bra REDUCE_END;
960    setp.ge.u32 %p4, %r5, %r6;
961    @%p4 bra REDUCE_NEXT;
962    add.u32 %r7, %r5, %r6;
963    ld.shared.f32 %f2, s_acc[%r7];
964    ld.shared.f32 %f3, s_acc[%r5];
965    add.f32 %f3, %f3, %f2;
966    st.shared.f32 s_acc[%r5], %f3;
967REDUCE_NEXT:
968    bar.sync 0;
969    shr.u32 %r6, %r6, 1;
970    bra REDUCE_LOOP;
971REDUCE_END:
972
973    // Thread 0 writes result
974    setp.ne.u32 %p5, %r5, 0;
975    @%p5 bra EXIT;
976    ld.shared.f32 %f1, s_acc[0];
977    st.global.f32 [%rd3], %f1;
978
979EXIT:
980    ret;"#,
981    );
982    Ok(())
983}
984
985#[cfg(test)]
986mod tests {
987    use super::*;
988    use crate::wgsl_forge::{BufferAccess, BufferElement, BufferSpec, ScalarType};
989
990    fn make_spec(id: &str, entry: &str, bufs: &[(&str, BufferAccess)]) -> KernelSpec {
991        KernelSpec {
992            id: id.to_string(),
993            semantic_version: 1,
994            entry_point: entry.to_string(),
995            description: "test".to_string(),
996            buffers: bufs
997                .iter()
998                .enumerate()
999                .map(|(i, (name, access))| BufferSpec {
1000                    group: 0,
1001                    binding: i as u32,
1002                    name: name.to_string(),
1003                    element: BufferElement::Scalar(ScalarType::F32),
1004                    access: *access,
1005                })
1006                .collect(),
1007            ops: Vec::new(),
1008            shared_memory: Vec::new(),
1009        }
1010    }
1011
1012    #[test]
1013    fn ptx_rmsnorm_emits_correct_instructions() {
1014        let spec = make_spec(
1015            "rmsnorm",
1016            "rmsnorm_main",
1017            &[
1018                ("x", BufferAccess::StorageRead),
1019                ("w", BufferAccess::StorageRead),
1020                ("y", BufferAccess::StorageReadWrite),
1021                ("params", BufferAccess::Uniform),
1022            ],
1023        );
1024        let shader = emit_ptx(&spec, Schedule::default()).expect("ptx rmsnorm");
1025        assert!(
1026            shader.source.contains("rsqrt.approx.f32"),
1027            "should use fast inverse sqrt"
1028        );
1029        assert!(
1030            shader.source.contains("bar.sync"),
1031            "should use barrier for reduction"
1032        );
1033        assert!(
1034            shader.source.contains("fma.rn.f32"),
1035            "should use fma for accumulation"
1036        );
1037        assert!(
1038            shader.source.contains("rmsnorm_main"),
1039            "should contain entry point"
1040        );
1041    }
1042
1043    #[test]
1044    fn ptx_q4k_gemv_emits_nc_loads() {
1045        let spec = make_spec(
1046            "q4k-gemv",
1047            "q4k_gemv_main",
1048            &[
1049                ("x", BufferAccess::StorageRead),
1050                ("w", BufferAccess::StorageRead),
1051                ("y", BufferAccess::StorageReadWrite),
1052                ("params", BufferAccess::Uniform),
1053            ],
1054        );
1055        let shader = emit_ptx(&spec, Schedule::default()).expect("ptx q4k gemv");
1056        assert!(
1057            shader.source.contains("ld.global.nc"),
1058            "should use non-coherent loads for read-only weights"
1059        );
1060        assert!(
1061            shader.source.contains("fma.rn.f32"),
1062            "should use fma for accumulation"
1063        );
1064    }
1065
1066    #[test]
1067    fn ptx_wmma_gemv_emits_wmma_instruction() {
1068        let spec = make_spec(
1069            "wmma-gemv",
1070            "wmma_gemv_main",
1071            &[
1072                ("a", BufferAccess::StorageRead),
1073                ("x", BufferAccess::StorageRead),
1074                ("y", BufferAccess::StorageReadWrite),
1075                ("params", BufferAccess::Uniform),
1076            ],
1077        );
1078        let shader = emit_ptx(&spec, Schedule::default()).expect("ptx wmma gemv");
1079        assert!(
1080            shader.source.contains("wmma.mma.sync.aligned.m16n16k16"),
1081            "should use WMMA tensor-core instruction"
1082        );
1083        assert!(
1084            shader.source.contains("row.col.f32.f16.f16.f32"),
1085            "should use f16→f32 accumulate"
1086        );
1087    }
1088
1089    #[test]
1090    fn ptx_sdpa_emits_exp2_approx() {
1091        let spec = make_spec(
1092            "sdpa-decode",
1093            "sdpa_main",
1094            &[
1095                ("q", BufferAccess::StorageRead),
1096                ("kv", BufferAccess::StorageRead),
1097                ("out", BufferAccess::StorageReadWrite),
1098                ("params", BufferAccess::Uniform),
1099            ],
1100        );
1101        let shader = emit_ptx(&spec, Schedule::default()).expect("ptx sdpa");
1102        assert!(
1103            shader.source.contains("ex2.approx.f32"),
1104            "should use fast exp2 for softmax"
1105        );
1106        assert!(
1107            shader.source.contains("rcp.approx.f32"),
1108            "should use fast reciprocal for normalization"
1109        );
1110        assert!(
1111            shader.source.contains("s_max"),
1112            "should use shared memory for max score"
1113        );
1114    }
1115
1116    #[test]
1117    fn ptx_unsupported_kernel_returns_error() {
1118        let spec = make_spec("unknown", "foo", &[]);
1119        let result = emit_ptx(&spec, Schedule::default());
1120        assert!(result.is_err(), "unsupported kernel should return error");
1121    }
1122
1123    #[test]
1124    fn ptx_q4k_soa_wmma_emits_wmma_instruction() {
1125        let spec = make_spec(
1126            "q4k-soa-wmma",
1127            "q4k_soa_wmma_main",
1128            &[
1129                ("x", BufferAccess::StorageRead),
1130                ("W", BufferAccess::StorageRead),
1131                ("y", BufferAccess::StorageReadWrite),
1132                ("params", BufferAccess::Uniform),
1133            ],
1134        );
1135        let shader = emit_ptx(&spec, Schedule::default()).expect("ptx q4k soa wmma");
1136        assert!(
1137            shader.source.contains("wmma.mma.sync.aligned.m16n16k16"),
1138            "should use WMMA tensor-core instruction"
1139        );
1140        assert!(
1141            shader.source.contains("row.col.f32.f16.f16.f32"),
1142            "should use f16→f32 accumulate"
1143        );
1144        assert!(
1145            shader.source.contains("ld.global.u8"),
1146            "should load weight bytes"
1147        );
1148    }
1149
1150    #[test]
1151    fn ptx_q6k_soa_gemv_emits_dequant_logic() {
1152        let spec = make_spec(
1153            "q6k-soa-gemv",
1154            "q6k_soa_gemv_main",
1155            &[
1156                ("x", BufferAccess::StorageRead),
1157                ("W", BufferAccess::StorageRead),
1158                ("y", BufferAccess::StorageReadWrite),
1159                ("params", BufferAccess::Uniform),
1160            ],
1161        );
1162        let shader = emit_ptx(&spec, Schedule::default()).expect("ptx q6k soa gemv");
1163        assert!(
1164            shader.source.contains("ld.global.nc"),
1165            "should use non-coherent loads for input"
1166        );
1167        assert!(
1168            shader.source.contains("fma.rn.f32"),
1169            "should use fma for accumulation"
1170        );
1171        assert!(
1172            shader.source.contains("210"),
1173            "should reference Q6_K block size"
1174        );
1175        assert!(
1176            shader.source.contains("bar.sync"),
1177            "should use barrier for reduction"
1178        );
1179    }
1180}