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 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
129fn 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
227fn 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
326fn 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
411fn 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
625fn 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
840fn 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}