1use std::fmt::Write;
2
3use super::GeneratedShader;
4use crate::wgsl_forge::{ForgeError, KernelSpec, Op, Schedule};
5
6pub fn emit_hlsl(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 "// HLSL 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 writeln!(
20 source,
21 "// Schedule: workgroup={}, items={}, vector={}",
22 schedule.workgroup_size, schedule.items_per_invocation, schedule.vector_width
23 )
24 .map_err(|e| ForgeError::Emission(e.to_string()))?;
25
26 emit_kernel_body(&mut source, kernel, schedule)?;
27
28 let source_hash = blake3::hash(source.as_bytes()).to_hex().to_string();
29 Ok(GeneratedShader {
30 kernel_id: kernel.id.clone(),
31 semantic_hash,
32 source_hash,
33 schedule,
34 source,
35 })
36}
37
38fn emit_kernel_body(
39 source: &mut String,
40 kernel: &KernelSpec,
41 schedule: Schedule,
42) -> Result<(), ForgeError> {
43 let use_wave = schedule.workgroup_size % 32 == 0 && schedule.workgroup_size >= 32;
46
47 if kernel.id == "topk" {
48 if use_wave {
49 return super::hlsl_wave::emit_topk_wave_hlsl(source, kernel, schedule);
50 }
51 return emit_topk_hlsl(source, kernel, schedule);
52 }
53 if kernel.id == "fused-ffn" {
54 if use_wave {
55 return super::hlsl_wave::emit_fused_ffn_wave_hlsl(source, kernel, schedule);
56 }
57 return emit_ffn_hlsl(source, kernel, schedule);
58 }
59 if kernel.id == "p64-project" {
60 return emit_p64_hlsl(source, kernel, schedule);
61 }
62 if kernel.id == "gemm" {
63 return emit_gemm_hlsl(source, kernel, schedule);
64 }
65 if kernel.id == "gemv" {
66 if use_wave {
67 return super::hlsl_wave::emit_gemv_wave_hlsl(source, kernel, schedule);
68 }
69 return emit_gemv_hlsl(source, kernel, schedule);
70 }
71 if kernel.id == "fused-qkv-rope" {
72 return emit_fused_qkv_rope_hlsl(source, kernel, schedule);
73 }
74 if kernel.id == "ternary-gemv" {
75 return emit_ternary_gemv_hlsl(source, kernel, schedule);
76 }
77 if kernel.id == "fft" {
78 return emit_fft_hlsl(source, kernel, schedule);
79 }
80 if kernel.id == "ray-probe" {
81 return Err(ForgeError::Emission(
82 "ray-query is only emitted for the WGSL target (HLSL RT uses a distinct API)"
83 .to_string(),
84 ));
85 }
86 if kernel.id == "affine-f32" {
87 writeln!(
88 source,
89 r#"struct AffineParams {{
90 uint length;
91 float scale;
92 float bias;
93 uint _pad;
94}};"#
95 )
96 .map_err(|error| ForgeError::Emission(error.to_string()))?;
97 }
98
99 if kernel.id == "p64-project" {
100 writeln!(
101 source,
102 r#"struct P64Words64 {{
103 uint4 lanes[4];
104}};"#
105 )
106 .map_err(|error| ForgeError::Emission(error.to_string()))?;
107 }
108
109 writeln!(source, "").map_err(|error| ForgeError::Emission(error.to_string()))?;
110
111 for buffer in &kernel.buffers {
112 let (type_decl, reg_type) = match (buffer.element, buffer.access) {
113 (crate::wgsl_forge::ir::BufferElement::AffineParams, _) => {
114 ("ConstantBuffer<AffineParams>", "b")
115 }
116 (
117 crate::wgsl_forge::ir::BufferElement::P64Words64,
118 crate::wgsl_forge::ir::BufferAccess::StorageRead,
119 ) => ("StructuredBuffer<P64Words64>", "t"),
120 (
121 crate::wgsl_forge::ir::BufferElement::P64Words64,
122 crate::wgsl_forge::ir::BufferAccess::StorageReadWrite,
123 ) => ("RWStructuredBuffer<P64Words64>", "u"),
124 (
125 crate::wgsl_forge::ir::BufferElement::Scalar(
126 crate::wgsl_forge::ir::ScalarType::F32,
127 ),
128 crate::wgsl_forge::ir::BufferAccess::StorageRead,
129 ) => ("StructuredBuffer<float>", "t"),
130 (
131 crate::wgsl_forge::ir::BufferElement::Scalar(
132 crate::wgsl_forge::ir::ScalarType::F32,
133 ),
134 crate::wgsl_forge::ir::BufferAccess::StorageReadWrite,
135 ) => ("RWStructuredBuffer<float>", "u"),
136 _ => ("StructuredBuffer<float>", "t"), };
138 writeln!(
139 source,
140 "{} {} : register({}{}, space{});",
141 type_decl, buffer.name, reg_type, buffer.binding, buffer.group
142 )
143 .map_err(|error| ForgeError::Emission(error.to_string()))?;
144 }
145
146 writeln!(source, "\n[numthreads({}, 1, 1)]", schedule.workgroup_size)
147 .map_err(|error| ForgeError::Emission(error.to_string()))?;
148 writeln!(
149 source,
150 "void {}(uint3 gid : SV_DispatchThreadID) {{",
151 kernel.entry_point
152 )
153 .map_err(|error| ForgeError::Emission(error.to_string()))?;
154
155 writeln!(
156 source,
157 " const uint ITEMS_PER_INVOCATION = {};\n const uint VECTOR_WIDTH = {};",
158 schedule.items_per_invocation, schedule.vector_width
159 )
160 .map_err(|error| ForgeError::Emission(error.to_string()))?;
161
162 if kernel.id == "affine-f32" {
163 writeln!(
164 source,
165 " for (uint item = 0; item < ITEMS_PER_INVOCATION; item++) {{"
166 )
167 .map_err(|error| ForgeError::Emission(error.to_string()))?;
168 writeln!(
169 source,
170 " uint global_id = (gid.x * ITEMS_PER_INVOCATION + item) * VECTOR_WIDTH;"
171 )
172 .map_err(|error| ForgeError::Emission(error.to_string()))?;
173
174 if schedule.vector_width == 1 {
175 writeln!(source, " if (global_id < params.length) {{")
176 .map_err(|error| ForgeError::Emission(error.to_string()))?;
177 emit_ops(source, &kernel.ops, " ")?;
178 writeln!(source, " }}")
179 .map_err(|error| ForgeError::Emission(error.to_string()))?;
180 } else {
181 writeln!(
187 source,
188 " if (global_id + {} < params.length) {{",
189 schedule.vector_width - 1
190 )
191 .map_err(|error| ForgeError::Emission(error.to_string()))?;
192 for index in 0..schedule.vector_width {
193 writeln!(source, " output[global_id + {index}] = input[global_id + {index}] * params.scale + params.bias;").map_err(|error| ForgeError::Emission(error.to_string()))?;
194 }
195 writeln!(source, " }} else {{")
196 .map_err(|error| ForgeError::Emission(error.to_string()))?;
197 writeln!(
198 source,
199 " for (uint component = 0; component < VECTOR_WIDTH; component++) {{"
200 )
201 .map_err(|error| ForgeError::Emission(error.to_string()))?;
202 writeln!(source, " uint base = global_id + component;")
203 .map_err(|error| ForgeError::Emission(error.to_string()))?;
204 writeln!(source, " if (base < params.length) {{")
205 .map_err(|error| ForgeError::Emission(error.to_string()))?;
206 writeln!(
207 source,
208 " output[base] = input[base] * params.scale + params.bias;"
209 )
210 .map_err(|error| ForgeError::Emission(error.to_string()))?;
211 writeln!(source, " }}")
212 .map_err(|error| ForgeError::Emission(error.to_string()))?;
213 writeln!(source, " }}")
214 .map_err(|error| ForgeError::Emission(error.to_string()))?;
215 writeln!(source, " }}")
216 .map_err(|error| ForgeError::Emission(error.to_string()))?;
217 }
218 writeln!(source, " }}\n}}").map_err(|error| ForgeError::Emission(error.to_string()))?;
219 } else {
220 writeln!(
221 source,
222 " for (uint item = 0; item < ITEMS_PER_INVOCATION; item++) {{"
223 )
224 .map_err(|error| ForgeError::Emission(error.to_string()))?;
225 writeln!(
226 source,
227 " uint global_id = gid.x * ITEMS_PER_INVOCATION + item;"
228 )
229 .map_err(|error| ForgeError::Emission(error.to_string()))?;
230 emit_ops(source, &kernel.ops, " ")?;
231 writeln!(source, " }}\n}}").map_err(|error| ForgeError::Emission(error.to_string()))?;
232 }
233
234 Ok(())
235}
236
237fn emit_topk_hlsl(
241 source: &mut String,
242 kernel: &KernelSpec,
243 schedule: Schedule,
244) -> Result<(), ForgeError> {
245 let wg = schedule.workgroup_size;
246 writeln!(
247 source,
248 "struct TopKParams {{\n uint length;\n uint k;\n uint block_size;\n uint _pad;\n}};\n"
249 )
250 .map_err(|error| ForgeError::Emission(error.to_string()))?;
251
252 writeln!(
253 source,
254 "StructuredBuffer<float> input : register(t0, space0);"
255 )
256 .map_err(|error| ForgeError::Emission(error.to_string()))?;
257 writeln!(
258 source,
259 "RWStructuredBuffer<float> output : register(u1, space0);"
260 )
261 .map_err(|error| ForgeError::Emission(error.to_string()))?;
262 writeln!(
263 source,
264 "ConstantBuffer<TopKParams> params : register(b2, space0);\n"
265 )
266 .map_err(|error| ForgeError::Emission(error.to_string()))?;
267
268 for shared in &kernel.shared_memory {
269 let ty = hlsl_scalar(shared.element);
270 writeln!(
271 source,
272 "groupshared {} {}[{}];",
273 ty,
274 shared.name,
275 shared.length.resolve(wg)
276 )
277 .map_err(|error| ForgeError::Emission(error.to_string()))?;
278 }
279
280 writeln!(
281 source,
282 r#"
283[numthreads({wg}, 1, 1)]
284void {entry}(uint tid : SV_GroupIndex, uint3 group_id : SV_GroupID) {{
285 uint block = group_id.x;
286 uint base = block * {wg}u;
287 uint gidx = base + tid;
288 float sentinel = asfloat(0xff7fffffu);
289 float v = sentinel;
290 if (gidx < params.length) {{ v = input[gidx]; }}
291 s_val[tid] = v;
292 s_idx[tid] = tid;
293 GroupMemoryBarrierWithGroupSync();
294
295 for (uint i = 0u; i < params.k; i++) {{
296 r_val[tid] = s_val[tid];
297 r_idx[tid] = s_idx[tid];
298 GroupMemoryBarrierWithGroupSync();
299 for (uint stride = {wg}u / 2u; stride > 0u; stride /= 2u) {{
300 if (tid < stride) {{
301 if (r_val[tid + stride] > r_val[tid]) {{
302 r_val[tid] = r_val[tid + stride];
303 r_idx[tid] = r_idx[tid + stride];
304 }}
305 }}
306 GroupMemoryBarrierWithGroupSync();
307 }}
308 if (tid == 0u) {{
309 output[block * params.k + i] = r_val[0];
310 s_val[r_idx[0]] = sentinel;
311 }}
312 GroupMemoryBarrierWithGroupSync();
313 }}
314}}"#,
315 wg = wg,
316 entry = kernel.entry_point
317 )
318 .map_err(|error| ForgeError::Emission(error.to_string()))?;
319
320 Ok(())
321}
322
323fn emit_ffn_hlsl(
326 source: &mut String,
327 kernel: &KernelSpec,
328 schedule: Schedule,
329) -> Result<(), ForgeError> {
330 let wg = schedule.workgroup_size;
331 writeln!(
332 source,
333 r#"struct FfnParams {{
334 uint input_size;
335 uint hidden_size;
336 uint output_size;
337 uint _pad;
338}};
339
340StructuredBuffer<float> input : register(t0, space0);
341StructuredBuffer<float> w1 : register(t1, space0);
342StructuredBuffer<float> w2 : register(t2, space0);
343RWStructuredBuffer<float> output : register(u3, space0);
344ConstantBuffer<FfnParams> params : register(b4, space0);
345
346[numthreads({wg}, 1, 1)]
347void {entry}(uint3 gid : SV_DispatchThreadID) {{
348 uint o = gid.x;
349 if (o >= params.output_size) {{ return; }}
350 float acc = 0.0;
351 for (uint h = 0; h < params.hidden_size; h++) {{
352 float hv = 0.0;
353 uint w1_row = h * params.input_size;
354 for (uint i = 0; i < params.input_size; i++) {{ hv += w1[w1_row + i] * input[i]; }}
355 float g = 0.5f * hv * (1.0f + tanh(0.7978845608f * (hv + 0.044715f * hv * hv * hv)));
356 acc += w2[o * params.hidden_size + h] * g;
357 }}
358 output[o] = acc;
359}}"#,
360 wg = wg,
361 entry = kernel.entry_point
362 )
363 .map_err(|error| ForgeError::Emission(error.to_string()))?;
364 Ok(())
365}
366
367fn emit_p64_hlsl(
370 source: &mut String,
371 kernel: &KernelSpec,
372 schedule: Schedule,
373) -> Result<(), ForgeError> {
374 let wg = schedule.workgroup_size;
375 writeln!(
376 source,
377 r#"struct P64Words64 {{
378 uint4 lanes[4];
379}};
380
381StructuredBuffer<P64Words64> input : register(t0, space0);
382StructuredBuffer<float> weights : register(t1, space0);
383RWStructuredBuffer<float> output : register(u2, space0);
384
385[numthreads({wg}, 1, 1)]
386void {entry}(uint3 gid : SV_DispatchThreadID) {{
387 uint r = gid.x;
388 uint count, stride;
389 output.GetDimensions(count, stride);
390 if (r >= count) {{ return; }}
391 P64Words64 rec = input[r];
392 float acc = 0.0;
393 for (uint w = 0; w < 16; w++) {{
394 uint word = rec.lanes[w / 4][w % 4];
395 acc += weights[w] * (float)word;
396 }}
397 output[r] = acc;
398}}"#,
399 wg = wg,
400 entry = kernel.entry_point
401 )
402 .map_err(|error| ForgeError::Emission(error.to_string()))?;
403 Ok(())
404}
405
406fn emit_gemm_hlsl(
410 source: &mut String,
411 kernel: &KernelSpec,
412 schedule: Schedule,
413) -> Result<(), ForgeError> {
414 let wg = schedule.workgroup_size;
415 writeln!(
416 source,
417 r#"struct GemmParams {{
418 uint m;
419 uint n;
420 uint k;
421 uint _pad;
422}};
423
424StructuredBuffer<float> a : register(t0, space0);
425StructuredBuffer<float> b : register(t1, space0);
426RWStructuredBuffer<float> c : register(u2, space0);
427ConstantBuffer<GemmParams> params : register(b3, space0);
428
429[numthreads({wg}, 1, 1)]
430void {entry}(uint3 gid : SV_DispatchThreadID) {{
431 uint o = gid.x;
432 if (o >= params.m * params.n) {{ return; }}
433 uint row = o / params.n;
434 uint col = o % params.n;
435 float acc = 0.0;
436 uint a_row = row * params.k;
437 for (uint kk = 0; kk < params.k; kk++) {{
438 acc += a[a_row + kk] * b[kk * params.n + col];
439 }}
440 c[o] = acc;
441}}"#,
442 wg = wg,
443 entry = kernel.entry_point
444 )
445 .map_err(|error| ForgeError::Emission(error.to_string()))?;
446 Ok(())
447}
448
449fn emit_gemv_hlsl(
452 source: &mut String,
453 kernel: &KernelSpec,
454 schedule: Schedule,
455) -> Result<(), ForgeError> {
456 let wg = schedule.workgroup_size;
457 writeln!(
458 source,
459 r#"struct GemvParams {{
460 uint m;
461 uint n;
462 uint _pad0;
463 uint _pad1;
464}};
465
466StructuredBuffer<float> a : register(t0, space0);
467StructuredBuffer<float> x : register(t1, space0);
468RWStructuredBuffer<float> y : register(u2, space0);
469ConstantBuffer<GemvParams> params : register(b3, space0);
470
471[numthreads({wg}, 1, 1)]
472void {entry}(uint3 gid : SV_DispatchThreadID) {{
473 uint i = gid.x;
474 if (i >= params.m) {{ return; }}
475 float acc = 0.0;
476 uint a_row = i * params.n;
477 for (uint j = 0; j < params.n; j++) {{
478 acc += a[a_row + j] * x[j];
479 }}
480 y[i] = acc;
481}}"#,
482 wg = wg,
483 entry = kernel.entry_point
484 )
485 .map_err(|error| ForgeError::Emission(error.to_string()))?;
486 Ok(())
487}
488
489fn emit_ternary_gemv_hlsl(
494 source: &mut String,
495 kernel: &KernelSpec,
496 schedule: Schedule,
497) -> Result<(), ForgeError> {
498 let wg = schedule.workgroup_size;
499 writeln!(
500 source,
501 r#"struct TernaryGemvParams {{
502 uint m;
503 uint k;
504 uint k_words;
505 uint _pad;
506}};
507
508StructuredBuffer<float> x : register(t0, space0);
509StructuredBuffer<uint> w_packed : register(t1, space0);
510StructuredBuffer<float> scale : register(t2, space0);
511RWStructuredBuffer<float> output : register(u3, space0);
512ConstantBuffer<TernaryGemvParams> params : register(b4, space0);
513
514[numthreads({wg}, 1, 1)]
515void {entry}(uint3 gid : SV_DispatchThreadID) {{
516 uint o = gid.x;
517 if (o >= params.m) {{ return; }}
518 float acc = 0.0;
519 uint row_base = o * params.k_words;
520 for (uint word_idx = 0; word_idx < params.k_words; word_idx++) {{
521 uint word = w_packed[row_base + word_idx];
522 uint lane_base = word_idx * 16u;
523 for (uint lane = 0; lane < 16u; lane++) {{
524 uint i = lane_base + lane;
525 if (i >= params.k) {{ break; }}
526 uint code = (word >> (lane * 2u)) & 3u;
527 float tern = 0.0;
528 if (code == 1u) {{ tern = 1.0; }} else if (code == 2u) {{ tern = -1.0; }}
529 acc += tern * x[i];
530 }}
531 }}
532 output[o] = scale[o] * acc;
533}}"#,
534 wg = wg,
535 entry = kernel.entry_point
536 )
537 .map_err(|error| ForgeError::Emission(error.to_string()))?;
538 Ok(())
539}
540
541fn emit_fft_hlsl(
548 source: &mut String,
549 kernel: &KernelSpec,
550 schedule: Schedule,
551) -> Result<(), ForgeError> {
552 let wg = schedule.workgroup_size;
553 writeln!(
554 source,
555 r#"struct FftParams {{
556 uint n;
557 uint log2n;
558 uint _pad0;
559 uint _pad1;
560}};
561
562StructuredBuffer<float> input : register(t0, space0);
563RWStructuredBuffer<float> output : register(u1, space0);
564ConstantBuffer<FftParams> params : register(b2, space0);
565
566groupshared float s_re[{wg}];
567groupshared float s_im[{wg}];
568
569[numthreads({wg}, 1, 1)]
570void {entry}(uint tid : SV_GroupIndex) {{
571 uint t = tid;
572 uint n = params.n;
573 uint logn = params.log2n;
574 uint rev = reversebits(t) >> (32u - logn);
575 s_re[rev] = input[2u * t];
576 s_im[rev] = input[2u * t + 1u];
577 GroupMemoryBarrierWithGroupSync();
578 for (uint s = 0u; s < logn; s++) {{
579 uint span = 1u << s;
580 uint m = span << 1u;
581 if (t < (n >> 1u)) {{
582 uint k = t & (span - 1u);
583 uint j = ((t >> s) << (s + 1u)) + k;
584 uint jp = j + span;
585 float ang = -6.28318548f * (float)k / (float)m;
586 float wr = cos(ang);
587 float wi = sin(ang);
588 float ur = s_re[j];
589 float ui = s_im[j];
590 float vr = s_re[jp];
591 float vi = s_im[jp];
592 float tr = vr * wr - vi * wi;
593 float ti = vr * wi + vi * wr;
594 s_re[j] = ur + tr;
595 s_im[j] = ui + ti;
596 s_re[jp] = ur - tr;
597 s_im[jp] = ui - ti;
598 }}
599 GroupMemoryBarrierWithGroupSync();
600 }}
601 output[2u * t] = s_re[t];
602 output[2u * t + 1u] = s_im[t];
603}}"#,
604 wg = wg,
605 entry = kernel.entry_point
606 )
607 .map_err(|error| ForgeError::Emission(error.to_string()))?;
608 Ok(())
609}
610
611fn hlsl_scalar(element: crate::wgsl_forge::ir::ScalarType) -> &'static str {
612 use crate::wgsl_forge::ir::ScalarType;
613 match element {
614 ScalarType::F32 => "float",
615 ScalarType::U32 => "uint",
616 ScalarType::I32 => "int",
617 ScalarType::U64Words => "uint2",
618 }
619}
620
621fn emit_ops(source: &mut String, ops: &[Op], indent: &str) -> Result<(), ForgeError> {
622 for op in ops {
623 match op {
624 Op::StructLoad {
625 buffer,
626 field,
627 destination,
628 } => {
629 writeln!(source, "{indent}float {destination} = {buffer}.{field};")
630 .map_err(|error| ForgeError::Emission(error.to_string()))?;
631 }
632 Op::Load {
633 buffer,
634 index,
635 destination,
636 } => {
637 writeln!(source, "{indent}float {destination} = {buffer}[{index}];")
638 .map_err(|error| ForgeError::Emission(error.to_string()))?;
639 }
640 Op::Store {
641 buffer,
642 index,
643 value,
644 } => {
645 writeln!(source, "{indent}{buffer}[{index}] = {value};")
646 .map_err(|error| ForgeError::Emission(error.to_string()))?;
647 }
648 Op::Fma {
649 a,
650 b,
651 c,
652 destination,
653 } => {
654 writeln!(source, "{indent}float {destination} = {a} * {b} + {c};")
655 .map_err(|error| ForgeError::Emission(error.to_string()))?;
656 }
657 Op::Mul {
658 left,
659 right,
660 destination,
661 } => {
662 writeln!(source, "{indent}float {destination} = {left} * {right};")
663 .map_err(|error| ForgeError::Emission(error.to_string()))?;
664 }
665 Op::Add {
666 left,
667 right,
668 destination,
669 } => {
670 writeln!(source, "{indent}float {destination} = {left} + {right};")
671 .map_err(|error| ForgeError::Emission(error.to_string()))?;
672 }
673 Op::DotProduct {
674 left_buffer,
675 left_base,
676 right_buffer,
677 right_base,
678 len,
679 destination,
680 } => {
681 writeln!(source, "{indent}float {destination} = 0.0;")
682 .map_err(|error| ForgeError::Emission(error.to_string()))?;
683 writeln!(source, "{indent}for (uint i = 0; i < {len}; i++) {{")
684 .map_err(|error| ForgeError::Emission(error.to_string()))?;
685 writeln!(source, "{indent} {destination} += {left_buffer}[{left_base} + i] * {right_buffer}[{right_base} + i];").map_err(|error| ForgeError::Emission(error.to_string()))?;
686 writeln!(source, "{indent}}}")
687 .map_err(|error| ForgeError::Emission(error.to_string()))?;
688 }
689 Op::Loop {
690 induction_var,
691 start,
692 end,
693 step,
694 body,
695 } => {
696 writeln!(source, "{indent}for (uint {induction_var} = {start}; {induction_var} < {end}; {induction_var} += {step}) {{").map_err(|error| ForgeError::Emission(error.to_string()))?;
697 emit_ops(source, body, &format!("{indent} "))?;
698 writeln!(source, "{indent}}}")
699 .map_err(|error| ForgeError::Emission(error.to_string()))?;
700 }
701 Op::Relu {
702 operand,
703 destination,
704 } => {
705 writeln!(
706 source,
707 "{indent}float {destination} = max(0.0f, {operand});"
708 )
709 .map_err(|error| ForgeError::Emission(error.to_string()))?;
710 }
711 Op::Gelu {
712 operand,
713 destination,
714 } => {
715 writeln!(source, "{indent}float {destination} = 0.5f * {operand} * (1.0f + tanh(0.7978845608f * ({operand} + 0.044715f * {operand} * {operand} * {operand})));").map_err(|error| ForgeError::Emission(error.to_string()))?;
716 }
717 Op::MatrixMultiply { .. } => {
718 return Err(ForgeError::Emission(
721 "Op::MatrixMultiply has no scalar HLSL lowering; use the cooperative-matrix / CUDA WMMA path".to_string(),
722 ));
723 }
724 Op::Barrier => {
725 writeln!(source, "{indent}GroupMemoryBarrierWithGroupSync();")
726 .map_err(|error| ForgeError::Emission(error.to_string()))?;
727 }
728 Op::Intrinsic(_) => {
729 return Err(ForgeError::Emission(
730 "Intrinsics not implemented for HLSL yet".to_string(),
731 ));
732 }
733 }
734 }
735 Ok(())
736}
737
738pub fn emit_gemv_wavematrix_hlsl(
748 source: &mut String,
749 kernel: &KernelSpec,
750 schedule: Schedule,
751) -> Result<(), ForgeError> {
752 let wg = schedule.workgroup_size;
753 writeln!(
754 source,
755 r#"struct GemvParams {{
756 uint m;
757 uint n;
758 uint _pad0;
759 uint _pad1;
760}};
761
762RWByteAddressBuffer a : register(u0, space0);
763RWByteAddressBuffer x : register(u1, space0);
764RWByteAddressBuffer y : register(u2, space0);
765ConstantBuffer<GemvParams> params : register(b3, space0);
766
767[numthreads({wg}, 1, 1)]
768void {entry}(uint3 gid : SV_DispatchThreadID) {{
769 uint wave_size = WaveGetLaneCount();
770 uint row_tile = gid.x / wave_size; // each wave processes 16 rows
771 uint row_base = row_tile * 16;
772 if (row_base >= params.m) {{ return; }}
773
774 // WaveMatrix fragments: 16x16 tiles
775 // A tile = 16 rows × 16 cols of matrix A (f16)
776 // B tile = 16 rows × 16 cols of vector x (f16, replicated)
777 // C tile = 16x16 accumulator (f32)
778 WaveMatrixA<float16_t> matA;
779 WaveMatrixB<float16_t> matB;
780 WaveMatrixC<float> matC;
781 WaveMatrixFill(matC, 0.0f);
782
783 uint n_tiles = (params.n + 15) / 16;
784 for (uint t = 0; t < n_tiles; t++) {{
785 uint col_base = t * 16;
786 // Load 16×16 tile of A (row_base..row_base+15, col_base..col_base+15)
787 for (uint i = WaveGetLaneIndex(); i < 256; i += wave_size) {{
788 uint local_row = i / 16;
789 uint local_col = i % 16;
790 uint global_row = row_base + local_row;
791 uint global_col = col_base + local_col;
792 float16_t val = 0.0h;
793 if (global_row < params.m && global_col < params.n) {{
794 val = float16_t(asfloat(a.Load2(global_row * params.n * 4 + global_col * 4)));
795 }}
796 WaveMatrixASetElement(matA, i, val);
797 }}
798 // Load 16×16 tile of x (replicated column vector into B tile)
799 for (uint i = WaveGetLaneIndex(); i < 256; i += wave_size) {{
800 uint local_row = i / 16;
801 uint global_col = col_base + local_row;
802 float16_t val = 0.0h;
803 if (global_col < params.n) {{
804 val = float16_t(asfloat(x.Load2(global_col * 4)));
805 }}
806 // Replicate across columns (each column of B gets same x value)
807 for (uint c = 0; c < 16; c++) {{
808 WaveMatrixBSetElement(matB, local_row * 16 + c, val);
809 }}
810 }}
811 WaveMatrixMultiply(matC, matA, matB);
812 }}
813
814 // Extract results: each lane writes its assigned output elements
815 uint lane = WaveGetLaneIndex();
816 for (uint i = lane; i < 16; i += wave_size) {{
817 uint global_row = row_base + i;
818 if (global_row < params.m) {{
819 // Sum across columns (GEMV: only one output per row)
820 float acc = 0.0f;
821 for (uint c = 0; c < 16; c++) {{
822 acc += WaveMatrixCGetElement(matC, i * 16 + c);
823 }}
824 y.Store2(global_row * 4, asuint(acc));
825 }}
826 }}
827}}"#,
828 wg = wg,
829 entry = kernel.entry_point
830 )
831 .map_err(|e| ForgeError::Emission(e.to_string()))?;
832 Ok(())
833}
834
835fn emit_fused_qkv_rope_hlsl(
844 source: &mut String,
845 kernel: &KernelSpec,
846 schedule: Schedule,
847) -> Result<(), ForgeError> {
848 let wg = schedule.workgroup_size.max(32);
849 let _ = kernel;
850 writeln!(
851 source,
852 r#"#define ROWS_PER_BLOCK 16u
853#define WG {wg}u
854
855RWStructuredBuffer<float> x : register(b0);
856RWStructuredBuffer<float> Wq : register(b1);
857RWStructuredBuffer<float> Wk : register(b2);
858RWStructuredBuffer<float> Wv : register(b3);
859RWStructuredBuffer<float> yq : register(b4);
860RWStructuredBuffer<float> yk : register(b5);
861RWStructuredBuffer<float> yv : register(b6);
862RWStructuredBuffer<uint> dims : register(b7);
863RWStructuredBuffer<uint> rope_params : register(b8);
864
865groupshared float s_red[ROWS_PER_BLOCK * WG];
866groupshared float s_rope_buf[ROWS_PER_BLOCK];
867
868[numthreads(WG, 1, 1)]
869void {entry}(uint3 dtid : SV_DispatchThreadID, uint3 gtid : SV_GroupThreadID, uint3 gid : SV_GroupID) {{
870 uint n_in = dims[0];
871 uint n_q = dims[1];
872 uint n_kv = dims[2];
873 uint n_head = dims[3];
874 uint head_dim = dims[4];
875 uint pos = dims[5];
876 uint row0 = gid.x * ROWS_PER_BLOCK;
877 uint t = gtid.x;
878 if (row0 >= n_q) return;
879
880 // RoPE parameters
881 uint base_bits = rope_params[0];
882 uint scale_bits = rope_params[1];
883 float base = asfloat(base_bits);
884 float scale = asfloat(scale_bits);
885 float inv_scale = 1.0 / scale;
886 float inv_head_dim = 1.0 / (float)head_dim;
887
888 // === Q projection with RoPE ===
889 float acc_q[ROWS_PER_BLOCK];
890 [unroll] for (uint r = 0u; r < ROWS_PER_BLOCK; r++) acc_q[r] = 0.0;
891 for (uint j = t; j < n_in; j += WG) {{
892 float xv = x[j];
893 [unroll] for (uint r = 0u; r < ROWS_PER_BLOCK; r++) {{
894 uint row = row0 + r;
895 if (row < n_q)
896 acc_q[r] += Wq[row * n_in + j] * xv;
897 }}
898 }}
899 // Tree reduction
900 [unroll] for (uint r = 0u; r < ROWS_PER_BLOCK; r++)
901 s_red[r * WG + t] = acc_q[r];
902 GroupMemoryBarrierWithGroupSync();
903 for (uint s = WG / 2u; s > 0u; s >>= 1u) {{
904 if (t < s) {{
905 [unroll] for (uint r = 0u; r < ROWS_PER_BLOCK; r++)
906 s_red[r * WG + t] += s_red[r * WG + t + s];
907 }}
908 GroupMemoryBarrierWithGroupSync();
909 }}
910 // Apply RoPE to Q and write
911 if (t < ROWS_PER_BLOCK) {{
912 uint row = row0 + t;
913 if (row < n_q) {{
914 uint head = row / head_dim;
915 uint d = row % head_dim;
916 uint half = head_dim / 2u;
917 if (half > 0u && head < n_head) {{
918 float val = s_red[t * WG];
919 uint i = d / 2u;
920 float theta = (float)pos * inv_scale * pow(base, -2.0 * (float)i * inv_head_dim);
921 float s_val = sin(theta);
922 float c_val = cos(theta);
923 float pair_val;
924 if (d % 2u == 0u) {{
925 pair_val = (t + 1u < ROWS_PER_BLOCK && (row0 + t + 1u) < n_q)
926 ? s_red[(t + 1u) * WG] : 0.0;
927 s_rope_buf[t] = val * c_val - pair_val * s_val;
928 }} else {{
929 pair_val = (t >= 1u) ? s_red[(t - 1u) * WG] : 0.0;
930 s_rope_buf[t] = pair_val * s_val + val * c_val;
931 }}
932 }} else {{
933 s_rope_buf[t] = s_red[t * WG];
934 }}
935 }}
936 }}
937 GroupMemoryBarrierWithGroupSync();
938 if (t < ROWS_PER_BLOCK) {{
939 uint row = row0 + t;
940 if (row < n_q) yq[row] = s_rope_buf[t];
941 }}
942 GroupMemoryBarrierWithGroupSync();
943
944 // === K projection with RoPE ===
945 float acc_k[ROWS_PER_BLOCK];
946 [unroll] for (uint r = 0u; r < ROWS_PER_BLOCK; r++) acc_k[r] = 0.0;
947 for (uint j = t; j < n_in; j += WG) {{
948 float xv = x[j];
949 [unroll] for (uint r = 0u; r < ROWS_PER_BLOCK; r++) {{
950 uint row = row0 + r;
951 if (row < n_kv)
952 acc_k[r] += Wk[row * n_in + j] * xv;
953 }}
954 }}
955 [unroll] for (uint r = 0u; r < ROWS_PER_BLOCK; r++)
956 s_red[r * WG + t] = acc_k[r];
957 GroupMemoryBarrierWithGroupSync();
958 for (uint s = WG / 2u; s > 0u; s >>= 1u) {{
959 if (t < s) {{
960 [unroll] for (uint r = 0u; r < ROWS_PER_BLOCK; r++)
961 s_red[r * WG + t] += s_red[r * WG + t + s];
962 }}
963 GroupMemoryBarrierWithGroupSync();
964 }}
965 if (t < ROWS_PER_BLOCK) {{
966 uint row = row0 + t;
967 if (row < n_kv) {{
968 uint head = row / head_dim;
969 uint d = row % head_dim;
970 uint half = head_dim / 2u;
971 if (half > 0u && head < n_head) {{
972 float val = s_red[t * WG];
973 uint i = d / 2u;
974 float theta = (float)pos * inv_scale * pow(base, -2.0 * (float)i * inv_head_dim);
975 float s_val = sin(theta);
976 float c_val = cos(theta);
977 float pair_val;
978 if (d % 2u == 0u) {{
979 pair_val = (t + 1u < ROWS_PER_BLOCK && (row0 + t + 1u) < n_kv)
980 ? s_red[(t + 1u) * WG] : 0.0;
981 s_rope_buf[t] = val * c_val - pair_val * s_val;
982 }} else {{
983 pair_val = (t >= 1u) ? s_red[(t - 1u) * WG] : 0.0;
984 s_rope_buf[t] = pair_val * s_val + val * c_val;
985 }}
986 }} else {{
987 s_rope_buf[t] = s_red[t * WG];
988 }}
989 }}
990 }}
991 GroupMemoryBarrierWithGroupSync();
992 if (t < ROWS_PER_BLOCK) {{
993 uint row = row0 + t;
994 if (row < n_kv) yk[row] = s_rope_buf[t];
995 }}
996 GroupMemoryBarrierWithGroupSync();
997
998 // === V projection (no RoPE) ===
999 float acc_v[ROWS_PER_BLOCK];
1000 [unroll] for (uint r = 0u; r < ROWS_PER_BLOCK; r++) acc_v[r] = 0.0;
1001 for (uint j = t; j < n_in; j += WG) {{
1002 float xv = x[j];
1003 [unroll] for (uint r = 0u; r < ROWS_PER_BLOCK; r++) {{
1004 uint row = row0 + r;
1005 if (row < n_kv)
1006 acc_v[r] += Wv[row * n_in + j] * xv;
1007 }}
1008 }}
1009 [unroll] for (uint r = 0u; r < ROWS_PER_BLOCK; r++)
1010 s_red[r * WG + t] = acc_v[r];
1011 GroupMemoryBarrierWithGroupSync();
1012 for (uint s = WG / 2u; s > 0u; s >>= 1u) {{
1013 if (t < s) {{
1014 [unroll] for (uint r = 0u; r < ROWS_PER_BLOCK; r++)
1015 s_red[r * WG + t] += s_red[r * WG + t + s];
1016 }}
1017 GroupMemoryBarrierWithGroupSync();
1018 }}
1019 if (t < ROWS_PER_BLOCK) {{
1020 uint row = row0 + t;
1021 if (row < n_kv) yv[row] = s_red[t * WG];
1022 }}
1023}}"#,
1024 wg = wg,
1025 entry = kernel.entry_point,
1026 )
1027 .map_err(|e| ForgeError::Emission(e.to_string()))?;
1028 Ok(())
1029}
1030
1031#[cfg(test)]
1032mod wavematrix_tests {
1033 use super::*;
1034
1035 #[test]
1036 fn emit_gemv_wavematrix_emits_intrinsics() {
1037 let kernel = KernelSpec {
1038 id: "gemv".to_string(),
1039 semantic_version: 1,
1040 entry_point: "gemv_wm".to_string(),
1041 description: "WaveMatrix GEMV".to_string(),
1042 buffers: Vec::new(),
1043 ops: Vec::new(),
1044 shared_memory: Vec::new(),
1045 };
1046 let schedule = Schedule {
1047 workgroup_size: 32,
1048 ..Default::default()
1049 };
1050 let mut source = String::new();
1051 emit_gemv_wavematrix_hlsl(&mut source, &kernel, schedule).expect("wavematrix emit");
1052 assert!(source.contains("WaveMatrixA"), "should use WaveMatrixA");
1053 assert!(source.contains("WaveMatrixB"), "should use WaveMatrixB");
1054 assert!(source.contains("WaveMatrixC"), "should use WaveMatrixC");
1055 assert!(
1056 source.contains("WaveMatrixMultiply"),
1057 "should call WaveMatrixMultiply"
1058 );
1059 assert!(source.contains("gemv_wm"), "should contain entry point");
1060 }
1061}