git.lucas.co / cce-ui
GPU-accelerated UI toolkit (Vulkan)
git clone https://git.lucas.co/cce-ui.git

examples/compute_probe/jobs.rs (7.6K)

  1 //! The compute probe's jobs, shared by both halves: each runs on a device
  2 //! and is held to a CPU reference, and prints a digest of its result bytes,
  3 //! so the native run (`compute_native`, Vulkan) and the browser's
  4 //! (`compute_web`, WebGPU) can be compared line for line. The arithmetic is
  5 //! exact in f32 (small integers, power-of-two weights), so a conformant
  6 //! device must match the reference — and the other device — to the bit.
  7 
  8 use cce_ui::compute::{Binding, Kernel};
  9 
 10 pub const DOUBLE: &str = "
 11 @group(0) @binding(0) var<storage, read> a: array<f32>;
 12 @group(0) @binding(1) var<storage, read_write> b: array<f32>;
 13 @compute @workgroup_size(64) fn main(@builtin(global_invocation_id) id: vec3<u32>) {
 14     if id.x < arrayLength(&b) { b[id.x] = a[id.x] * 2.0; }
 15 }";
 16 
 17 pub const SAXPY: &str = "
 18 struct P { a: f32, n: u32 }
 19 @group(0) @binding(0) var<uniform> p: P;
 20 @group(0) @binding(1) var<storage, read> x: array<f32>;
 21 @group(0) @binding(2) var<storage, read_write> y: array<f32>;
 22 @compute @workgroup_size(32) fn main(@builtin(global_invocation_id) id: vec3<u32>) {
 23     if id.x < p.n { y[id.x] = p.a * x[id.x] + y[id.x]; }
 24 }";
 25 
 26 pub const BLUR: &str = "
 27 @group(0) @binding(0) var<storage, read> src: array<f32>;
 28 @group(0) @binding(1) var<storage, read_write> dst: array<f32>;
 29 @compute @workgroup_size(64) fn main(@builtin(global_invocation_id) id: vec3<u32>) {
 30     let n = arrayLength(&dst);
 31     let i = id.x;
 32     if i >= n { return; }
 33     let l = src[max(i, 1u) - 1u];
 34     let r = src[min(i + 1u, n - 1u)];
 35     dst[i] = 0.25 * l + 0.5 * src[i] + 0.25 * r;
 36 }";
 37 
 38 pub const GRID: &str = "
 39 @group(0) @binding(0) var<storage, read_write> g: array<f32>;
 40 @compute @workgroup_size(8, 8) fn main(@builtin(global_invocation_id) id: vec3<u32>) {
 41     if id.x < 32u && id.y < 24u { g[id.y * 32u + id.x] = f32(id.x) + 1000.0 * f32(id.y); }
 42 }";
 43 
 44 /// A kernel over `n` bindings: `n - 1` inputs summed into the last.
 45 pub fn many_source(n: usize) -> String {
 46     let mut s = String::new();
 47     for i in 0..n - 1 {
 48         s += &format!("@group(0) @binding({i}) var<storage, read> in{i}: array<f32>;\n");
 49     }
 50     s += &format!("@group(0) @binding({}) var<storage, read_write> out: array<f32>;\n", n - 1);
 51     s += "@compute @workgroup_size(16) fn main(@builtin(global_invocation_id) id: vec3<u32>) {\n";
 52     s += "    let i = id.x; if i >= arrayLength(&out) { return; }\n    var t = 0.0;\n";
 53     for i in 0..n - 1 {
 54         s += &format!("    t += in{i}[i];\n");
 55     }
 56     s += "    out[i] = t;\n}\n";
 57     s
 58 }
 59 
 60 /// FNV-1a over a result's bytes: what two devices are compared by.
 61 pub fn digest(v: &[f32]) -> String {
 62     let mut h: u64 = 0xcbf29ce484222325;
 63     for b in bytemuck::cast_slice::<f32, u8>(v) {
 64         h ^= *b as u64;
 65         h = h.wrapping_mul(0x100000001b3);
 66     }
 67     format!("{h:016x}")
 68 }
 69 
 70 pub fn report(name: &str, got: &[f32], want: &[f32]) -> String {
 71     let worst = got.iter().zip(want).map(|(g, w)| (g - w).abs()).fold(0.0f32, f32::max);
 72     let ok = got.len() == want.len() && got.iter().zip(want).all(|(g, w)| g.to_bits() == w.to_bits());
 73     format!("{name}: {} n={} max|d|={worst} digest={}", if ok { "exact" } else { "DIFFERS" }, got.len(), digest(got))
 74 }
 75 
 76 /// The jobs, written once over whichever device runs them. A macro rather
 77 /// than a generic function: one device's `run` is synchronous and the
 78 /// other's async, and `$await` is the one word between them.
 79 #[macro_export]
 80 macro_rules! compute_jobs {
 81     ($dev:expr, $($await:tt)*) => {{
 82         use cce_ui::compute::{Binding, Kernel};
 83         use jobs::*;
 84         let mut lines: Vec<String> = Vec::new();
 85 
 86         // A map over 1000 elements, sized from the entry's @workgroup_size.
 87         let a: Vec<f32> = (0..1000).map(|i| i as f32 - 500.0).collect();
 88         let mut b = vec![0.0f32; 1000];
 89         let r = $dev.run_over(&Kernel::new(DOUBLE, "main"), &mut [Binding::input(&a), Binding::rw(&mut b)], 1000)$($await)*;
 90         let want: Vec<f32> = a.iter().map(|v| v * 2.0).collect();
 91         lines.push(match r { Ok(()) => report("double", &b, &want), Err(e) => format!("double: ERR {e}") });
 92 
 93         // A uniform block beside storage.
 94         #[repr(C)]
 95         #[derive(Clone, Copy, bytemuck::Pod, bytemuck::Zeroable)]
 96         struct P { a: f32, n: u32, _pad: [u32; 2] }
 97         let x: Vec<f32> = (0..777).map(|i| (i % 13) as f32).collect();
 98         let mut y: Vec<f32> = (0..777).map(|i| (i % 7) as f32).collect();
 99         let want: Vec<f32> = x.iter().zip(&y).map(|(x, y)| 2.5 * x + y).collect();
100         let p = P { a: 2.5, n: 777, _pad: [0; 2] };
101         let r = $dev.run_over(&Kernel::new(SAXPY, "main"), &mut [Binding::uniform(&p), Binding::input(&x), Binding::rw(&mut y)], 777)$($await)*;
102         lines.push(match r { Ok(()) => report("saxpy", &y, &want), Err(e) => format!("saxpy: ERR {e}") });
103 
104         // Ping-pong passes, odd and even counts: the result lands in the
105         // output binding either way. 4096 elements: a length the 16-byte
106         // padding leaves alone, since `arrayLength` counts the padding.
107         for passes in [33u32, 34] {
108             let src: Vec<f32> = (0..4096).map(|i| if i % 512 == 256 { 4096.0 } else { 0.0 }).collect();
109             let mut dst = vec![0.0f32; 4096];
110             let mut want = src.clone();
111             for _ in 0..passes {
112                 let n = want.len();
113                 want = (0..n).map(|i| 0.25 * want[i.saturating_sub(1)] + 0.5 * want[i] + 0.25 * want[(i + 1).min(n - 1)]).collect();
114             }
115             let r = $dev
116                 .run_passes_over(&Kernel::new(BLUR, "main"), &mut [Binding::input(&src), Binding::rw(&mut dst)], 4096, passes, Some((0, 1)))
117                 $($await)*;
118             lines.push(match r { Ok(()) => report(&format!("blur x{passes}"), &dst, &want), Err(e) => format!("blur x{passes}: ERR {e}") });
119         }
120 
121         // A 2D dispatch.
122         let mut g = vec![-1.0f32; 32 * 24];
123         let r = $dev.run(&Kernel::new(GRID, "main"), &mut [Binding::rw(&mut g)], [4, 3, 1])$($await)*;
124         let want: Vec<f32> = (0..32 * 24).map(|i| (i % 32) as f32 + 1000.0 * (i / 32) as f32).collect();
125         lines.push(match r { Ok(()) => report("grid", &g, &want), Err(e) => format!("grid: ERR {e}") });
126 
127         // Ten bindings: past WebGPU's default of eight storage buffers a
128         // stage (the browser device asks the adapter for its own limit),
129         // within what every adapter offers — SwiftShader's is ten.
130         let ins: Vec<Vec<f32>> = (0..9).map(|k| (0..100).map(|i| (k * 100 + i) as f32).collect()).collect();
131         let mut out = vec![0.0f32; 100];
132         let want: Vec<f32> = (0..100).map(|i| ins.iter().map(|v| v[i]).sum()).collect();
133         let mut binds: Vec<Binding> = ins.iter().map(|v| Binding::input(v)).collect();
134         binds.push(Binding::rw(&mut out));
135         let src = many_source(10);
136         let r = $dev.run_over(&Kernel::new(src.as_str(), "main"), &mut binds, 100)$($await)*;
137         drop(binds);
138         lines.push(match r { Ok(()) => report("ten bindings", &out, &want), Err(e) => format!("ten bindings: ERR {e}") });
139 
140         // A user's bad kernel is an Err, with naga's word for what is wrong.
141         let mut z = vec![0.0f32; 4];
142         let r = $dev.run(&Kernel::new("fn main( {", "main"), &mut [Binding::rw(&mut z)], [1, 1, 1])$($await)*;
143         lines.push(format!("bad wgsl: {}", match r { Err(e) if e.contains("parse error") => "Err(parse error)".to_string(), other => format!("{:?}", other) }));
144         let r = $dev.run(&Kernel::new(DOUBLE, "nope"), &mut [Binding::input(&a), Binding::rw(&mut z)], [1, 1, 1])$($await)*;
145         lines.push(format!("no entry: {}", match r { Err(e) => e, Ok(()) => "Ok".into() }));
146         lines
147     }};
148 }
149 
150 #[allow(dead_code)]
151 fn _uses(_: Binding, _: Kernel) {}