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) {}