Skip to main content

qualia_core_db/wgsl_forge/
validate.rs

1use serde::{Deserialize, Serialize};
2
3use crate::wgsl_forge::ForgeError;
4use crate::wgsl_forge::KernelSpec;
5use crate::wgsl_forge::TargetBackend;
6
7#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
8pub struct ValidationReport {
9    pub source_hash: String,
10    pub entry_points: Vec<String>,
11    pub binding_count: usize,
12    pub naga_validated: bool,
13    pub native_tool_validated: Option<String>,
14}
15
16pub fn validate_wgsl(source: &str) -> Result<ValidationReport, ForgeError> {
17    let module = naga::front::wgsl::parse_str(source)
18        .map_err(|error| ForgeError::WgslParse(error.emit_to_string(source)))?;
19    // Enabling capabilities only widens what is accepted, so kernels that do not
20    // use ray-query or cooperative-matrix features are unaffected.
21    let mut validator = naga::valid::Validator::new(
22        naga::valid::ValidationFlags::all(),
23        naga::valid::Capabilities::RAY_QUERY
24            | naga::valid::Capabilities::COOPERATIVE_MATRIX
25            | naga::valid::Capabilities::SHADER_FLOAT16,
26    );
27    validator
28        .validate(&module)
29        .map_err(|error| ForgeError::WgslValidation(format!("{error:?}")))?;
30
31    let mut entry_points = module
32        .entry_points
33        .iter()
34        .map(|entry| entry.name.clone())
35        .collect::<Vec<_>>();
36    entry_points.sort();
37    let binding_count = module
38        .global_variables
39        .iter()
40        .filter(|(_, variable)| variable.binding.is_some())
41        .count();
42    Ok(ValidationReport {
43        source_hash: blake3::hash(source.as_bytes()).to_hex().to_string(),
44        entry_points,
45        binding_count,
46        naga_validated: true,
47        native_tool_validated: None,
48    })
49}
50
51/// Validate a non-WGSL shader by spawning the target's offline compiler. When the
52/// source was generated from a known [`KernelSpec`], pass it so the report carries
53/// that kernel's real entry point and binding count; for an opaque external source
54/// (`kernel = None`) those fields are left empty/zero rather than guessed.
55pub fn validate_native(
56    source: &str,
57    target: TargetBackend,
58    kernel: Option<&KernelSpec>,
59) -> Result<ValidationReport, ForgeError> {
60    use std::process::Command;
61
62    let tool_name = match target {
63        TargetBackend::Ptx => "ptxas",
64        TargetBackend::Hlsl => "dxc",
65        TargetBackend::Msl => "xcrun",
66        _ => {
67            return Err(ForgeError::WgslValidation(
68                "Not a native target supported by offline validation".to_string(),
69            ))
70        }
71    };
72
73    let temp_file = tempfile::Builder::new()
74        .suffix(match target {
75            TargetBackend::Ptx => ".ptx",
76            TargetBackend::Hlsl => ".hlsl",
77            TargetBackend::Msl => ".metal",
78            _ => "",
79        })
80        .tempfile()
81        .map_err(|e| ForgeError::Io(format!("Failed to create temp file: {}", e)))?;
82
83    let temp_path = temp_file.path().to_path_buf();
84
85    // Write source to temp file
86    std::fs::write(&temp_path, source)
87        .map_err(|e| ForgeError::Io(format!("Failed to write to temp file: {}", e)))?;
88
89    let mut cmd = Command::new(tool_name);
90    match target {
91        TargetBackend::Ptx => {
92            // Check if ptxas exists first, fallback gently
93            if Command::new("ptxas").arg("--version").output().is_err() {
94                return Err(ForgeError::WgslValidation(
95                    "ptxas not found in PATH".to_string(),
96                ));
97            }
98            cmd.arg(&temp_path).arg("-c"); // Compile only
99        }
100        TargetBackend::Hlsl => {
101            if Command::new("dxc").arg("--help").output().is_err() {
102                return Err(ForgeError::WgslValidation(
103                    "dxc not found in PATH".to_string(),
104                ));
105            }
106            cmd.arg("-T").arg("cs_6_0").arg(&temp_path);
107        }
108        TargetBackend::Msl => {
109            if Command::new("xcrun").arg("--version").output().is_err() {
110                return Err(ForgeError::WgslValidation(
111                    "xcrun not found in PATH".to_string(),
112                ));
113            }
114            cmd.arg("-sdk")
115                .arg("macosx")
116                .arg("metal")
117                .arg("-c")
118                .arg(&temp_path);
119        }
120        _ => {}
121    }
122
123    let output = cmd
124        .output()
125        .map_err(|e| ForgeError::Io(format!("Failed to execute validation tool: {}", e)))?;
126
127    if !output.status.success() {
128        let stderr = String::from_utf8_lossy(&output.stderr);
129        return Err(ForgeError::WgslValidation(format!(
130            "{} validation failed: {}",
131            tool_name, stderr
132        )));
133    }
134
135    // The native compiler validated the source; report the kernel's real entry
136    // point + binding count when known, else leave them empty for an opaque file.
137    let (entry_points, binding_count) = match kernel {
138        Some(k) => (vec![k.entry_point.clone()], k.buffers.len()),
139        None => (Vec::new(), 0),
140    };
141    Ok(ValidationReport {
142        source_hash: blake3::hash(source.as_bytes()).to_hex().to_string(),
143        entry_points,
144        binding_count,
145        naga_validated: false,
146        native_tool_validated: Some(tool_name.to_string()),
147    })
148}
149
150#[cfg(test)]
151mod tests {
152    use super::*;
153    use crate::wgsl_forge::{generate_builtin, BuiltinKernel, Schedule, TargetBackend};
154
155    #[test]
156    fn generated_schedules_pass_full_naga_validation() {
157        for vector_width in [1, 2, 4] {
158            let generated = generate_builtin(
159                BuiltinKernel::AffineF32,
160                Schedule {
161                    vector_width,
162                    ..Schedule::default()
163                },
164                TargetBackend::Wgsl,
165            )
166            .unwrap();
167            let report = validate_wgsl(&generated.source).expect("Naga validation");
168            assert_eq!(report.entry_points, vec!["affine_f32"]);
169            assert_eq!(report.binding_count, 3);
170            assert_eq!(report.source_hash, generated.source_hash);
171        }
172    }
173
174    #[test]
175    fn generated_topk_passes_full_naga_validation() {
176        // Exercises workgroup-shared memory + barrier uniformity in the IR path.
177        for workgroup_size in [32u32, 64, 128, 256] {
178            let generated = generate_builtin(
179                BuiltinKernel::TopK,
180                Schedule {
181                    workgroup_size,
182                    items_per_invocation: 1,
183                    vector_width: 1,
184                    ..Schedule::default()
185                },
186                TargetBackend::Wgsl,
187            )
188            .unwrap();
189            assert!(
190                generated.source.contains("workgroupBarrier()"),
191                "top-k must emit barriers"
192            );
193            assert!(
194                generated
195                    .source
196                    .contains(&format!("array<f32, {workgroup_size}>")),
197                "shared arrays sized to workgroup size"
198            );
199            let report = validate_wgsl(&generated.source).expect("Naga validation of top-k");
200            assert_eq!(report.entry_points, vec!["topk"]);
201            assert_eq!(report.binding_count, 3);
202        }
203    }
204
205    #[test]
206    fn generated_p64_passes_naga_validation() {
207        let generated = generate_builtin(
208            BuiltinKernel::P64Project,
209            Schedule {
210                workgroup_size: 64,
211                items_per_invocation: 1,
212                vector_width: 1,
213                ..Schedule::default()
214            },
215            TargetBackend::Wgsl,
216        )
217        .unwrap();
218        assert!(generated.source.contains("struct P64Words64"));
219        assert!(generated.source.contains("arrayLength(&output)"));
220        let report = validate_wgsl(&generated.source).expect("Naga validation of p64-project");
221        assert_eq!(report.entry_points, vec!["p64_project"]);
222        assert_eq!(report.binding_count, 3);
223    }
224
225    #[test]
226    fn generated_ffn_passes_naga_validation() {
227        let generated = generate_builtin(
228            BuiltinKernel::FusedFfn,
229            Schedule {
230                workgroup_size: 64,
231                items_per_invocation: 1,
232                vector_width: 1,
233                ..Schedule::default()
234            },
235            TargetBackend::Wgsl,
236        )
237        .unwrap();
238        assert!(generated.source.contains("struct FfnParams"));
239        assert!(generated.source.contains("tanh("));
240        let report = validate_wgsl(&generated.source).expect("Naga validation of fused-ffn");
241        assert_eq!(report.entry_points, vec!["fused_ffn"]);
242        assert_eq!(report.binding_count, 5);
243    }
244
245    #[test]
246    fn generated_ray_probe_passes_naga_validation() {
247        // Exercises the acceleration_structure binding and the ray_query lowering.
248        let generated = generate_builtin(
249            BuiltinKernel::RayProbe,
250            Schedule {
251                workgroup_size: 64,
252                items_per_invocation: 1,
253                vector_width: 1,
254                ..Schedule::default()
255            },
256            TargetBackend::Wgsl,
257        )
258        .unwrap();
259        assert!(generated
260            .source
261            .contains("var scene: acceleration_structure;"));
262        assert!(generated.source.contains("rayQueryInitialize"));
263        assert!(generated
264            .source
265            .contains("rayQueryGetCommittedIntersection"));
266        let report = validate_wgsl(&generated.source).expect("Naga validation of ray-probe");
267        assert_eq!(report.entry_points, vec!["ray_probe"]);
268        assert_eq!(report.binding_count, 3);
269    }
270
271    #[test]
272    fn generated_ternary_gemv_passes_naga_validation() {
273        let generated = generate_builtin(
274            BuiltinKernel::TernaryGemv,
275            Schedule {
276                workgroup_size: 64,
277                items_per_invocation: 1,
278                vector_width: 1,
279                ..Schedule::default()
280            },
281            TargetBackend::Wgsl,
282        )
283        .unwrap();
284        assert!(generated.source.contains("struct TernaryGemvParams"));
285        // The dequant unpacks 2-bit lanes and maps the codes to {0, +1, -1}.
286        assert!(generated.source.contains("& 3u"));
287        let report = validate_wgsl(&generated.source).expect("Naga validation of ternary-gemv");
288        assert_eq!(report.entry_points, vec!["ternary_gemv"]);
289        assert_eq!(report.binding_count, 5);
290    }
291
292    #[test]
293    fn generated_gemm_passes_naga_validation() {
294        let generated = generate_builtin(
295            BuiltinKernel::Gemm,
296            Schedule {
297                workgroup_size: 64,
298                items_per_invocation: 1,
299                vector_width: 1,
300                ..Schedule::default()
301            },
302            TargetBackend::Wgsl,
303        )
304        .unwrap();
305        assert!(generated.source.contains("struct GemmParams"));
306        // The K-loop accumulates A-row * B-column into the output element.
307        assert!(generated.source.contains("kk * params.n + col"));
308        let report = validate_wgsl(&generated.source).expect("Naga validation of gemm");
309        assert_eq!(report.entry_points, vec!["gemm"]);
310        assert_eq!(report.binding_count, 4);
311    }
312
313    #[test]
314    fn generated_gemv_passes_naga_validation() {
315        let generated = generate_builtin(
316            BuiltinKernel::Gemv,
317            Schedule {
318                workgroup_size: 64,
319                items_per_invocation: 1,
320                vector_width: 1,
321                ..Schedule::default()
322            },
323            TargetBackend::Wgsl,
324        )
325        .unwrap();
326        assert!(generated.source.contains("struct GemvParams"));
327        // The N-loop accumulates A-row * x into the output row.
328        assert!(generated.source.contains("a[a_row + j] * x[j]"));
329        let report = validate_wgsl(&generated.source).expect("Naga validation of gemv");
330        assert_eq!(report.entry_points, vec!["gemv"]);
331        assert_eq!(report.binding_count, 4);
332    }
333
334    #[test]
335    fn generated_fft_passes_naga_validation() {
336        // The FFT uses workgroup-shared memory + barriers and reverseBits, so the
337        // schedule's workgroup_size must be a power of two (it is the transform
338        // length N). Exercise a representative power-of-two size.
339        let generated = generate_builtin(
340            BuiltinKernel::Fft,
341            Schedule {
342                workgroup_size: 256,
343                items_per_invocation: 1,
344                vector_width: 1,
345                ..Schedule::default()
346            },
347            TargetBackend::Wgsl,
348        )
349        .unwrap();
350        assert!(generated.source.contains("struct FftParams"));
351        // Bit-reversal load + barriers + the forward twiddle are the heart of it.
352        assert!(generated.source.contains("reverseBits("));
353        assert!(generated.source.contains("workgroupBarrier()"));
354        assert!(
355            generated.source.contains("array<f32, 256>"),
356            "shared arrays sized to workgroup size"
357        );
358        let report = validate_wgsl(&generated.source).expect("Naga validation of fft");
359        assert_eq!(report.entry_points, vec!["fft"]);
360        assert_eq!(report.binding_count, 3);
361    }
362
363    #[test]
364    fn cooperative_matrix_tile_validates() {
365        // Single 8x8x8 tensor-core tile: C = A * B (one subgroup cooperative).
366        let source = crate::wgsl_forge::matmul_tc_wgsl();
367        let report = validate_wgsl(&source).expect("coopmat tile should validate");
368        assert_eq!(report.entry_points, vec!["matmul_tc"]);
369    }
370
371    #[test]
372    fn cooperative_matrix_tiled_gemm_validates() {
373        // Tiled coopmat GEMM (loops the 8x8x8 tile over arbitrary m/n/k via workgroup_id +
374        // a K-loop, dims = [m,n,k] storage buffer). The WGSL is correct + naga-valid today;
375        // its *execution* is dormant until wgpu #9741 (the coopmat multiply returns zeros on
376        // 29.0.3) — `dispatch::coopmat_usable()` gates it until then.
377        let source = crate::wgsl_forge::emit::matmul_tc_wgsl_tiled();
378        let report = validate_wgsl(&source).expect("tiled coopmat GEMM should validate");
379        assert_eq!(
380            report.entry_points,
381            vec![crate::wgsl_forge::emit::MATMUL_TC_TILED_ENTRY]
382        );
383        assert_eq!(report.binding_count, 4);
384    }
385
386    #[test]
387    fn semantic_errors_are_rejected() {
388        let source = "@compute @workgroup_size(64) fn broken() { let x: u32 = 1.0; }";
389        assert!(validate_wgsl(source).is_err());
390    }
391
392    /// Structural check: every native emitter for the newly-implemented kernels
393    /// must produce a non-empty body with the expected native-specific keywords.
394    /// This catches regressions where a dispatch falls through to an empty
395    /// generic path or an "not implemented" error.
396    #[test]
397    fn native_emitters_produce_non_empty_bodies() {
398        use crate::wgsl_forge::BuiltinKernel as K;
399        let schedule = Schedule::default();
400        let cases: &[(K, TargetBackend, &[&str])] = &[
401            // HLSL
402            (
403                K::Gemm,
404                TargetBackend::Hlsl,
405                &["RWStructuredBuffer", "for (uint kk"],
406            ),
407            (
408                K::Gemv,
409                TargetBackend::Hlsl,
410                &["RWStructuredBuffer", "for (uint j"],
411            ),
412            (K::Fft, TargetBackend::Hlsl, &["groupshared", "reversebits"]),
413            (
414                K::TernaryGemv,
415                TargetBackend::Hlsl,
416                &["StructuredBuffer<uint>", "tern"],
417            ),
418            (
419                K::P64Project,
420                TargetBackend::Hlsl,
421                &["StructuredBuffer<P64Words64>"],
422            ),
423            // MSL
424            (K::Gemm, TargetBackend::Msl, &["device", "for (uint kk"]),
425            (K::Gemv, TargetBackend::Msl, &["device", "for (uint j"]),
426            (K::Fft, TargetBackend::Msl, &["threadgroup", "reverse_bits"]),
427            (
428                K::TernaryGemv,
429                TargetBackend::Msl,
430                &["device const uint", "tern"],
431            ),
432            (
433                K::P64Project,
434                TargetBackend::Msl,
435                &["P64Words64", "record_count"],
436            ),
437            // CUDA-C
438            (
439                K::TernaryGemv,
440                TargetBackend::CudaC,
441                &["__global__", "tern"],
442            ),
443            (
444                K::P64Project,
445                TargetBackend::CudaC,
446                &["__global__", "P64Words64"],
447            ),
448            (K::Fft, TargetBackend::CudaC, &["__shared__", "__brev"]),
449        ];
450        for (kernel, target, expected) in cases {
451            let generated = generate_builtin(*kernel, schedule, *target)
452                .unwrap_or_else(|e| panic!("{kernel:?}/{target:?} emission failed: {e}"));
453            let src = &generated.source;
454            assert!(
455                src.lines().count() > 10,
456                "{kernel:?}/{target:?} body too short ({} lines)",
457                src.lines().count()
458            );
459            for kw in *expected {
460                assert!(
461                    src.contains(kw),
462                    "{kernel:?}/{target:?} missing expected keyword `{kw}`"
463                );
464            }
465        }
466    }
467}