Skip to main content

qualia_core_db/wgsl_forge/emit/
hlsl.rs

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    // Wave-intrinsic emitters: used when workgroup is a multiple of the
44    // typical wave size (32). Falls back to scalar emitters otherwise.
45    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"), // Fallback
137        };
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            // Vectorized affine (affine-f32 is the only kernel that sets vector_width>1):
182            // an unrolled fast path when the whole VECTOR_WIDTH span is in bounds, plus a
183            // bounds-checked tail for the final partial span. Correct for affine-f32 (its sole
184            // op is out = in*scale + bias). Native HLSL float4 SIMD loads would be a throughput
185            // optimization (needs DXC to validate — absent on this host), not a correctness gap.
186            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
237/// Top-k reduction in HLSL (compute shader 6.0): one thread group per block,
238/// `k` largest values per block in descending order, using `groupshared` arrays
239/// (driven by the IR) and `GroupMemoryBarrierWithGroupSync`.
240fn 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
323/// Fused FFN in HLSL (cs_6_0): one thread per output element (see the WGSL
324/// emitter for the math).
325fn 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
367/// P64 descriptor projection in HLSL: one thread per record, bound via the
368/// structured buffer's GetDimensions (no length uniform needed).
369fn 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
406/// Dense row-major GEMM in HLSL (cs_6_0): one thread per output element
407/// `o = i*N + j` computes `C[i][j] = sum_k A[i*K+k] * B[k*N+j]`. Same binding order,
408/// params layout and accumulation order as the certified WGSL `gemm`.
409fn 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
449/// Dense row-major GEMV in HLSL (cs_6_0): one thread per output ROW `i` computes
450/// `y[i] = sum_j A[i*N+j] * x[j]` — same order as the certified WGSL `gemv`.
451fn 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
489/// BitNet-style ternary GEMV in HLSL: one thread per output row `o` computes
490/// `out[o] = scale[o] * sum_i ternary(w[o,i]) * x[i]`. 2-bit codes, 16 per `uint`
491/// (low-to-high lanes; `0->0.0, 1->+1.0, 2->-1.0, 3->0.0`), `k_words` per row.
492/// `w_packed` is a `StructuredBuffer<uint>` — the generic path wrongly typed it float.
493fn 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
541/// Forward radix-2 DIT FFT in HLSL over ONE thread group of `N = workgroup_size`
542/// threads. Interleaved complex f32 (`input[2*j]`, `input[2*j+1]`), bit-reversal load
543/// into `groupshared`, then `log2(N)` butterfly stages with
544/// `GroupMemoryBarrierWithGroupSync()`. Same `exp(-2*pi*i*k/m)` convention as the
545/// WGSL kernel and the CPU DFT oracle. `reversebits` is the HLSL intrinsic
546/// (WGSL spells it `reverseBits`).
547fn 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                // No scalar HLSL lowering for a dense GEMM op; fail loudly rather
719                // than silently emit nothing (tensor-core GEMM is delivered elsewhere).
720                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
738/// HLSL WaveMatrix GEMV using SM 6.8+ tensor-core intrinsics.
739///
740/// Uses `WaveMatrixA` (f16), `WaveMatrixB` (f16), `WaveMatrixC` (f32) for
741/// 16×16 tile matrix multiply. DXC compiles this to SPIR-V
742/// `CooperativeMatrixKHR` when targeting `vulkan1.2`. Requires adapter
743/// support for `cooperative_matrix` — gate with `coopmat_usable()`.
744///
745/// Binding ABI: same as scalar GEMV (a, x, y, params).
746/// Dispatch: one wave per output row tile (16 rows per wave).
747pub 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
835/// HLSL fused QKV + RoPE kernel: f32 GEMV for Q, K, V projections with
836/// RoPE rotation applied to Q and K outputs before writing to global memory.
837/// V is written without rotation. Uses `groupshared` memory for the RoPE
838/// pair buffer.
839///
840/// Bindings: x, Wq, Wk, Wv, yq, yk, yv, dims={n_in, n_q, n_kv, n_head, head_dim, pos},
841/// rope_params={base_bits, scale_bits}.
842/// Dispatch: grid = ceil(n_q / ROWS_PER_BLOCK), block = 256.
843fn 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}