lib/accy/src/kernel/library/sdf.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 
  3 const kernel = @import("../root.zig");
  4 
  5 pub const grid_sample_family_version: u32 = 1;
  6 
  7 pub const GridSample = struct {
  8     gradient: bool = false,
  9     threads: u32 = 256,
 10 };
 11 
 12 pub fn gridSampleInstanceValid(instance: GridSample) bool {
 13     return instance.threads > 0 and instance.threads <= 1024;
 14 }
 15 
 16 const AxisState = struct {
 17     lo: kernel.Value,
 18     frac: kernel.Value,
 19     extent: kernel.Value,
 20 };
 21 
 22 fn axisState(
 23     inner: anytype,
 24     coordinate: kernel.Value,
 25     origin: kernel.Value,
 26     inv_spacing: kernel.Value,
 27     vertex_count: kernel.Value,
 28 ) !AxisState {
 29     const zero = try inner.constantFloat(.f32, 0);
 30     const one = try inner.constantFloat(.f32, 1);
 31     const negative_origin = try inner.mul(origin, try inner.constantFloat(.f32, -1));
 32     const local = try inner.mul(try inner.add(coordinate, negative_origin), inv_spacing);
 33     const last = try inner.sub(try inner.cast(vertex_count, .f32), one);
 34     const clamped = try inner.min(try inner.max(local, zero), last);
 35     const lo_f = try inner.floor(clamped);
 36     const lo_capped = try inner.min(lo_f, try inner.sub(last, one));
 37     const lo_safe = try inner.max(lo_capped, zero);
 38     const frac = try inner.sub(clamped, lo_safe);
 39     return .{
 40         .lo = try inner.castIndex(try inner.cast(lo_safe, .u32)),
 41         .frac = frac,
 42         .extent = try inner.castIndex(vertex_count),
 43     };
 44 }
 45 
 46 fn lerp(inner: anytype, from: kernel.Value, to: kernel.Value, t: kernel.Value) !kernel.Value {
 47     return inner.fma(try inner.sub(to, from), t, from);
 48 }
 49 
 50 fn gridSampleBody2Plain(k: anytype, spec: GridSample, args: anytype) !void {
 51     return gridSampleBody2(k, spec, args, false);
 52 }
 53 
 54 fn gridSampleBody2Gradient(k: anytype, spec: GridSample, args: anytype) !void {
 55     return gridSampleBody2(k, spec, args, true);
 56 }
 57 
 58 fn grid_sample_body2_active(inner: anytype, ctx: anytype) !void {
 59     const one_index = try inner.constantIndex(1);
 60     const x = (try ctx.args.param(.xs).load(inner, ctx.point)).raw();
 61     const y = (try ctx.args.param(.ys).load(inner, ctx.point)).raw();
 62     const sx = try axisState(inner, x, ctx.args.param(.origin_x).raw(), ctx.args.param(.inv_dx).raw(), ctx.args.param(.nx).raw());
 63     const sy = try axisState(inner, y, ctx.args.param(.origin_y).raw(), ctx.args.param(.inv_dy).raw(), ctx.args.param(.ny).raw());
 64 
 65     const row0 = try inner.mul(sy.lo, sx.extent);
 66     const row1 = try inner.mul(try inner.add(sy.lo, one_index), sx.extent);
 67     const x1 = try inner.add(sx.lo, one_index);
 68     const v00 = (try ctx.args.param(.values).load(inner, try inner.add(row0, sx.lo))).raw();
 69     const v10 = (try ctx.args.param(.values).load(inner, try inner.add(row0, x1))).raw();
 70     const v01 = (try ctx.args.param(.values).load(inner, try inner.add(row1, sx.lo))).raw();
 71     const v11 = (try ctx.args.param(.values).load(inner, try inner.add(row1, x1))).raw();
 72 
 73     const bottom = try lerp(inner, v00, v10, sx.frac);
 74     const top = try lerp(inner, v01, v11, sx.frac);
 75     const distance = try lerp(inner, bottom, top, sy.frac);
 76     try ctx.args.param(.dist).store(inner, distance, ctx.point);
 77 
 78     if (comptime ctx.gradient) {
 79         const dx_bottom = try inner.sub(v10, v00);
 80         const dx_top = try inner.sub(v11, v01);
 81         const dx = try lerp(inner, dx_bottom, dx_top, sy.frac);
 82         const dy_left = try inner.sub(v01, v00);
 83         const dy_right = try inner.sub(v11, v10);
 84         const dy = try lerp(inner, dy_left, dy_right, sx.frac);
 85         try ctx.args.param(.grad_x).store(inner, try inner.mul(dx, ctx.args.param(.inv_dx).raw()), ctx.point);
 86         try ctx.args.param(.grad_y).store(inner, try inner.mul(dy, ctx.args.param(.inv_dy).raw()), ctx.point);
 87     }
 88 }
 89 
 90 fn gridSampleBody2(k: anytype, spec: GridSample, args: anytype, comptime gradient: bool) !void {
 91     if (!gridSampleInstanceValid(spec)) return error.UnsupportedGridSampleInstance;
 92     const point = try k.globalId(.x);
 93     const count = try k.castIndex(args.param(.count).raw());
 94     const active = try k.compare(.lt, point, count);
 95     try k.guardDo(active, .{ .args = args, .point = point, .gradient = gradient }, grid_sample_body2_active);
 96 }
 97 
 98 fn gridSampleBody3Plain(k: anytype, spec: GridSample, args: anytype) !void {
 99     return gridSampleBody3(k, spec, args, false);
100 }
101 
102 fn gridSampleBody3Gradient(k: anytype, spec: GridSample, args: anytype) !void {
103     return gridSampleBody3(k, spec, args, true);
104 }
105 
106 fn grid_sample_body3_active(inner: anytype, ctx: anytype) !void {
107     const one_index = try inner.constantIndex(1);
108     const x = (try ctx.args.param(.xs).load(inner, ctx.point)).raw();
109     const y = (try ctx.args.param(.ys).load(inner, ctx.point)).raw();
110     const z = (try ctx.args.param(.zs).load(inner, ctx.point)).raw();
111     const sx = try axisState(inner, x, ctx.args.param(.origin_x).raw(), ctx.args.param(.inv_dx).raw(), ctx.args.param(.nx).raw());
112     const sy = try axisState(inner, y, ctx.args.param(.origin_y).raw(), ctx.args.param(.inv_dy).raw(), ctx.args.param(.ny).raw());
113     const sz = try axisState(inner, z, ctx.args.param(.origin_z).raw(), ctx.args.param(.inv_dz).raw(), ctx.args.param(.nz).raw());
114 
115     const x1 = try inner.add(sx.lo, one_index);
116     const plane = try inner.mul(sy.extent, sx.extent);
117     const slab0 = try inner.mul(sz.lo, plane);
118     const slab1 = try inner.mul(try inner.add(sz.lo, one_index), plane);
119     const row00 = try inner.add(slab0, try inner.mul(sy.lo, sx.extent));
120     const row01 = try inner.add(slab0, try inner.mul(try inner.add(sy.lo, one_index), sx.extent));
121     const row10 = try inner.add(slab1, try inner.mul(sy.lo, sx.extent));
122     const row11 = try inner.add(slab1, try inner.mul(try inner.add(sy.lo, one_index), sx.extent));
123 
124     const v000 = (try ctx.args.param(.values).load(inner, try inner.add(row00, sx.lo))).raw();
125     const v100 = (try ctx.args.param(.values).load(inner, try inner.add(row00, x1))).raw();
126     const v010 = (try ctx.args.param(.values).load(inner, try inner.add(row01, sx.lo))).raw();
127     const v110 = (try ctx.args.param(.values).load(inner, try inner.add(row01, x1))).raw();
128     const v001 = (try ctx.args.param(.values).load(inner, try inner.add(row10, sx.lo))).raw();
129     const v101 = (try ctx.args.param(.values).load(inner, try inner.add(row10, x1))).raw();
130     const v011 = (try ctx.args.param(.values).load(inner, try inner.add(row11, sx.lo))).raw();
131     const v111 = (try ctx.args.param(.values).load(inner, try inner.add(row11, x1))).raw();
132 
133     const bottom0 = try lerp(inner, v000, v100, sx.frac);
134     const top0 = try lerp(inner, v010, v110, sx.frac);
135     const slab0_value = try lerp(inner, bottom0, top0, sy.frac);
136     const bottom1 = try lerp(inner, v001, v101, sx.frac);
137     const top1 = try lerp(inner, v011, v111, sx.frac);
138     const slab1_value = try lerp(inner, bottom1, top1, sy.frac);
139     const distance = try lerp(inner, slab0_value, slab1_value, sz.frac);
140     try ctx.args.param(.dist).store(inner, distance, ctx.point);
141 
142     if (comptime ctx.gradient) {
143         const dx00 = try inner.sub(v100, v000);
144         const dx01 = try inner.sub(v110, v010);
145         const dx10 = try inner.sub(v101, v001);
146         const dx11 = try inner.sub(v111, v011);
147         const dx0 = try lerp(inner, dx00, dx01, sy.frac);
148         const dx1 = try lerp(inner, dx10, dx11, sy.frac);
149         const dx = try lerp(inner, dx0, dx1, sz.frac);
150 
151         const dy00 = try inner.sub(v010, v000);
152         const dy10 = try inner.sub(v110, v100);
153         const dy01 = try inner.sub(v011, v001);
154         const dy11 = try inner.sub(v111, v101);
155         const dy0 = try lerp(inner, dy00, dy10, sx.frac);
156         const dy1 = try lerp(inner, dy01, dy11, sx.frac);
157         const dy = try lerp(inner, dy0, dy1, sz.frac);
158 
159         const dz00 = try inner.sub(v001, v000);
160         const dz10 = try inner.sub(v101, v100);
161         const dz01 = try inner.sub(v011, v010);
162         const dz11 = try inner.sub(v111, v110);
163         const dz0 = try lerp(inner, dz00, dz10, sx.frac);
164         const dz1 = try lerp(inner, dz01, dz11, sx.frac);
165         const dz = try lerp(inner, dz0, dz1, sy.frac);
166 
167         try ctx.args.param(.grad_x).store(inner, try inner.mul(dx, ctx.args.param(.inv_dx).raw()), ctx.point);
168         try ctx.args.param(.grad_y).store(inner, try inner.mul(dy, ctx.args.param(.inv_dy).raw()), ctx.point);
169         try ctx.args.param(.grad_z).store(inner, try inner.mul(dz, ctx.args.param(.inv_dz).raw()), ctx.point);
170     }
171 }
172 
173 fn gridSampleBody3(k: anytype, spec: GridSample, args: anytype, comptime gradient: bool) !void {
174     if (!gridSampleInstanceValid(spec)) return error.UnsupportedGridSampleInstance;
175     const point = try k.globalId(.x);
176     const count = try k.castIndex(args.param(.count).raw());
177     const active = try k.compare(.lt, point, count);
178     try k.guardDo(active, .{ .args = args, .point = point, .gradient = gradient }, grid_sample_body3_active);
179 }
180 
181 fn gridSampleSchedule(instance: GridSample) kernel.logical.schedule.ThreadBlocks {
182     return kernel.logical.schedule.threadBlocks(.{ .x = instance.threads });
183 }
184 
185 pub const GridSampleFamily2D = kernel.logical.Family(.{
186     .name = "accy_kernel_sdf_grid_sample_2d_f32",
187     .parameters = .{
188         .dist = kernel.dynamicBuffer(.f32),
189         .values = kernel.dynamicBuffer(.f32),
190         .xs = kernel.dynamicBuffer(.f32),
191         .ys = kernel.dynamicBuffer(.f32),
192         .count = kernel.scalar(.i32),
193         .origin_x = kernel.scalar(.f32),
194         .origin_y = kernel.scalar(.f32),
195         .inv_dx = kernel.scalar(.f32),
196         .inv_dy = kernel.scalar(.f32),
197         .nx = kernel.scalar(.i32),
198         .ny = kernel.scalar(.i32),
199     },
200     .Instance = GridSample,
201     .schedule = gridSampleSchedule,
202     .body = gridSampleBody2Plain,
203 });
204 
205 pub const GridSampleGradientFamily2D = kernel.logical.Family(.{
206     .name = "accy_kernel_sdf_grid_sample_gradient_2d_f32",
207     .parameters = .{
208         .dist = kernel.dynamicBuffer(.f32),
209         .grad_x = kernel.dynamicBuffer(.f32),
210         .grad_y = kernel.dynamicBuffer(.f32),
211         .values = kernel.dynamicBuffer(.f32),
212         .xs = kernel.dynamicBuffer(.f32),
213         .ys = kernel.dynamicBuffer(.f32),
214         .count = kernel.scalar(.i32),
215         .origin_x = kernel.scalar(.f32),
216         .origin_y = kernel.scalar(.f32),
217         .inv_dx = kernel.scalar(.f32),
218         .inv_dy = kernel.scalar(.f32),
219         .nx = kernel.scalar(.i32),
220         .ny = kernel.scalar(.i32),
221     },
222     .Instance = GridSample,
223     .schedule = gridSampleSchedule,
224     .body = gridSampleBody2Gradient,
225 });
226 
227 pub const GridSampleFamily3D = kernel.logical.Family(.{
228     .name = "accy_kernel_sdf_grid_sample_3d_f32",
229     .parameters = .{
230         .dist = kernel.dynamicBuffer(.f32),
231         .values = kernel.dynamicBuffer(.f32),
232         .xs = kernel.dynamicBuffer(.f32),
233         .ys = kernel.dynamicBuffer(.f32),
234         .zs = kernel.dynamicBuffer(.f32),
235         .count = kernel.scalar(.i32),
236         .origin_x = kernel.scalar(.f32),
237         .origin_y = kernel.scalar(.f32),
238         .origin_z = kernel.scalar(.f32),
239         .inv_dx = kernel.scalar(.f32),
240         .inv_dy = kernel.scalar(.f32),
241         .inv_dz = kernel.scalar(.f32),
242         .nx = kernel.scalar(.i32),
243         .ny = kernel.scalar(.i32),
244         .nz = kernel.scalar(.i32),
245     },
246     .Instance = GridSample,
247     .schedule = gridSampleSchedule,
248     .body = gridSampleBody3Plain,
249 });
250 
251 pub const GridSampleGradientFamily3D = kernel.logical.Family(.{
252     .name = "accy_kernel_sdf_grid_sample_gradient_3d_f32",
253     .parameters = .{
254         .dist = kernel.dynamicBuffer(.f32),
255         .grad_x = kernel.dynamicBuffer(.f32),
256         .grad_y = kernel.dynamicBuffer(.f32),
257         .grad_z = kernel.dynamicBuffer(.f32),
258         .values = kernel.dynamicBuffer(.f32),
259         .xs = kernel.dynamicBuffer(.f32),
260         .ys = kernel.dynamicBuffer(.f32),
261         .zs = kernel.dynamicBuffer(.f32),
262         .count = kernel.scalar(.i32),
263         .origin_x = kernel.scalar(.f32),
264         .origin_y = kernel.scalar(.f32),
265         .origin_z = kernel.scalar(.f32),
266         .inv_dx = kernel.scalar(.f32),
267         .inv_dy = kernel.scalar(.f32),
268         .inv_dz = kernel.scalar(.f32),
269         .nx = kernel.scalar(.i32),
270         .ny = kernel.scalar(.i32),
271         .nz = kernel.scalar(.i32),
272     },
273     .Instance = GridSample,
274     .schedule = gridSampleSchedule,
275     .body = gridSampleBody3Gradient,
276 });
277 
278 pub fn gridSampleBlockCount(count: u64, threads: u32) u64 {
279     return (count + threads - 1) / threads;
280 }
281 
282 pub const ReferenceGrid2 = struct {
283     values: []const f32,
284     nx: u32,
285     ny: u32,
286     origin_x: f32,
287     origin_y: f32,
288     inv_dx: f32,
289     inv_dy: f32,
290 };
291 
292 pub const ReferenceGrid3 = struct {
293     values: []const f32,
294     nx: u32,
295     ny: u32,
296     nz: u32,
297     origin_x: f32,
298     origin_y: f32,
299     origin_z: f32,
300     inv_dx: f32,
301     inv_dy: f32,
302     inv_dz: f32,
303 };
304 
305 const ReferenceAxis = struct {
306     lo: usize,
307     frac: f32,
308 };
309 
310 fn referenceAxis(coordinate: f32, origin: f32, inv_spacing: f32, vertex_count: u32) ReferenceAxis {
311     const local = (coordinate + origin * -1) * inv_spacing;
312     const last = @as(f32, @floatFromInt(vertex_count)) - 1;
313     const clamped = @max(@as(f32, 0), @min(local, last));
314     const lo_f = @floor(clamped);
315     const lo_safe = @max(@as(f32, 0), @min(lo_f, last - 1));
316     return .{
317         .lo = @intFromFloat(lo_safe),
318         .frac = clamped - lo_safe,
319     };
320 }
321 
322 fn referenceLerp(from: f32, to: f32, t: f32) f32 {
323     return @mulAdd(f32, to - from, t, from);
324 }
325 
326 pub fn referenceGridSample2(grid: ReferenceGrid2, x: f32, y: f32, gradient: ?*[2]f32) f32 {
327     const sx = referenceAxis(x, grid.origin_x, grid.inv_dx, grid.nx);
328     const sy = referenceAxis(y, grid.origin_y, grid.inv_dy, grid.ny);
329     const row0 = sy.lo * grid.nx;
330     const row1 = (sy.lo + 1) * grid.nx;
331     const v00 = grid.values[row0 + sx.lo];
332     const v10 = grid.values[row0 + sx.lo + 1];
333     const v01 = grid.values[row1 + sx.lo];
334     const v11 = grid.values[row1 + sx.lo + 1];
335     const bottom = referenceLerp(v00, v10, sx.frac);
336     const top = referenceLerp(v01, v11, sx.frac);
337     if (gradient) |out| {
338         out[0] = referenceLerp(v10 - v00, v11 - v01, sy.frac) * grid.inv_dx;
339         out[1] = referenceLerp(v01 - v00, v11 - v10, sx.frac) * grid.inv_dy;
340     }
341     return referenceLerp(bottom, top, sy.frac);
342 }
343 
344 pub fn referenceGridSample3(grid: ReferenceGrid3, x: f32, y: f32, z: f32, gradient: ?*[3]f32) f32 {
345     const sx = referenceAxis(x, grid.origin_x, grid.inv_dx, grid.nx);
346     const sy = referenceAxis(y, grid.origin_y, grid.inv_dy, grid.ny);
347     const sz = referenceAxis(z, grid.origin_z, grid.inv_dz, grid.nz);
348     const plane = @as(usize, grid.ny) * grid.nx;
349     const slab0 = sz.lo * plane;
350     const slab1 = (sz.lo + 1) * plane;
351     const row00 = slab0 + sy.lo * grid.nx;
352     const row01 = slab0 + (sy.lo + 1) * grid.nx;
353     const row10 = slab1 + sy.lo * grid.nx;
354     const row11 = slab1 + (sy.lo + 1) * grid.nx;
355     const v000 = grid.values[row00 + sx.lo];
356     const v100 = grid.values[row00 + sx.lo + 1];
357     const v010 = grid.values[row01 + sx.lo];
358     const v110 = grid.values[row01 + sx.lo + 1];
359     const v001 = grid.values[row10 + sx.lo];
360     const v101 = grid.values[row10 + sx.lo + 1];
361     const v011 = grid.values[row11 + sx.lo];
362     const v111 = grid.values[row11 + sx.lo + 1];
363     const bottom0 = referenceLerp(v000, v100, sx.frac);
364     const top0 = referenceLerp(v010, v110, sx.frac);
365     const slab0_value = referenceLerp(bottom0, top0, sy.frac);
366     const bottom1 = referenceLerp(v001, v101, sx.frac);
367     const top1 = referenceLerp(v011, v111, sx.frac);
368     const slab1_value = referenceLerp(bottom1, top1, sy.frac);
369     if (gradient) |out| {
370         const dx0 = referenceLerp(v100 - v000, v110 - v010, sy.frac);
371         const dx1 = referenceLerp(v101 - v001, v111 - v011, sy.frac);
372         out[0] = referenceLerp(dx0, dx1, sz.frac) * grid.inv_dx;
373         const dy0 = referenceLerp(v010 - v000, v110 - v100, sx.frac);
374         const dy1 = referenceLerp(v011 - v001, v111 - v101, sx.frac);
375         out[1] = referenceLerp(dy0, dy1, sz.frac) * grid.inv_dy;
376         const dz0 = referenceLerp(v001 - v000, v101 - v100, sx.frac);
377         const dz1 = referenceLerp(v011 - v010, v111 - v110, sx.frac);
378         out[2] = referenceLerp(dz0, dz1, sy.frac) * grid.inv_dz;
379     }
380     return referenceLerp(slab0_value, slab1_value, sz.frac);
381 }
382 
383 const testing = std.testing;
384 
385 fn diskDistance(x: f32, y: f32) f32 {
386     return @sqrt(x * x + y * y) - 0.75;
387 }
388 
389 fn fillDiskGrid(values: []f32, nx: u32, ny: u32, origin: f32, spacing: f32) void {
390     var iy: u32 = 0;
391     while (iy < ny) : (iy += 1) {
392         var ix: u32 = 0;
393         while (ix < nx) : (ix += 1) {
394             const x = origin + @as(f32, @floatFromInt(ix)) * spacing;
395             const y = origin + @as(f32, @floatFromInt(iy)) * spacing;
396             values[@as(usize, iy) * nx + ix] = diskDistance(x, y);
397         }
398     }
399 }
400 
401 test "2d grid sample with gradient matches the reference on the interpreter" {
402     const allocator = testing.allocator;
403     const nx: u32 = 17;
404     const ny: u32 = 13;
405     const origin: f32 = -1.0;
406     const spacing: f32 = 0.125;
407     const count: u32 = 64;
408 
409     const values = try allocator.alloc(f32, @as(usize, nx) * ny);
410     defer allocator.free(values);
411     fillDiskGrid(values, nx, ny, origin, spacing);
412 
413     const xs = try allocator.alloc(f32, count);
414     defer allocator.free(xs);
415     const ys = try allocator.alloc(f32, count);
416     defer allocator.free(ys);
417     for (0..count) |point| {
418         xs[point] = -1.4 + @as(f32, @floatFromInt(point % 11)) * 0.27;
419         ys[point] = -1.2 + @as(f32, @floatFromInt(point % 7)) * 0.31;
420     }
421 
422     const grid = ReferenceGrid2{
423         .values = values,
424         .nx = nx,
425         .ny = ny,
426         .origin_x = origin,
427         .origin_y = origin,
428         .inv_dx = 1.0 / spacing,
429         .inv_dy = 1.0 / spacing,
430     };
431 
432     const expected_dist = try allocator.alloc(f32, count);
433     defer allocator.free(expected_dist);
434     const expected_gx = try allocator.alloc(f32, count);
435     defer allocator.free(expected_gx);
436     const expected_gy = try allocator.alloc(f32, count);
437     defer allocator.free(expected_gy);
438     for (0..count) |point| {
439         var gradient: [2]f32 = undefined;
440         expected_dist[point] = referenceGridSample2(grid, xs[point], ys[point], &gradient);
441         expected_gx[point] = gradient[0];
442         expected_gy[point] = gradient[1];
443     }
444 
445     const dist = try allocator.alloc(f32, count);
446     defer allocator.free(dist);
447     const gx = try allocator.alloc(f32, count);
448     defer allocator.free(gx);
449     const gy = try allocator.alloc(f32, count);
450     defer allocator.free(gy);
451     @memset(dist, 0);
452     @memset(gx, 0);
453     @memset(gy, 0);
454 
455     const instance = GridSample{ .gradient = true, .threads = 32 };
456     var graph = try GridSampleGradientFamily2D.build(allocator, GridSampleGradientFamily2D.Limits.testing, instance);
457     defer graph.deinit();
458     try graph.runCpuWithLaunch(allocator, &.{
459         kernel.argumentBuffer(f32, dist),
460         kernel.argumentBuffer(f32, gx),
461         kernel.argumentBuffer(f32, gy),
462         kernel.argumentBuffer(f32, values),
463         kernel.argumentBuffer(f32, xs),
464         kernel.argumentBuffer(f32, ys),
465         kernel.argumentI32(@intCast(count)),
466         kernel.argumentF32(origin),
467         kernel.argumentF32(origin),
468         kernel.argumentF32(1.0 / spacing),
469         kernel.argumentF32(1.0 / spacing),
470         kernel.argumentI32(@intCast(nx)),
471         kernel.argumentI32(@intCast(ny)),
472     }, .{
473         .grid = .{ @intCast(gridSampleBlockCount(count, instance.threads)), 1, 1 },
474         .block = .{ instance.threads, 1, 1 },
475     });
476 
477     try testing.expectEqualSlices(f32, expected_dist, dist);
478     try testing.expectEqualSlices(f32, expected_gx, gx);
479     try testing.expectEqualSlices(f32, expected_gy, gy);
480 }
481 
482 test "3d grid sample matches the reference on the interpreter" {
483     const allocator = testing.allocator;
484     const nx: u32 = 9;
485     const ny: u32 = 7;
486     const nz: u32 = 6;
487     const origin: f32 = -1.0;
488     const spacing: f32 = 0.3;
489     const count: u32 = 48;
490 
491     const values = try allocator.alloc(f32, @as(usize, nx) * ny * nz);
492     defer allocator.free(values);
493     for (values, 0..) |*value, index| {
494         const iz = index / (@as(usize, nx) * ny);
495         const rem = index % (@as(usize, nx) * ny);
496         const iy = rem / nx;
497         const ix = rem % nx;
498         const x = origin + @as(f32, @floatFromInt(ix)) * spacing;
499         const y = origin + @as(f32, @floatFromInt(iy)) * spacing;
500         const z = origin + @as(f32, @floatFromInt(iz)) * spacing;
501         value.* = @sqrt(x * x + y * y + z * z) - 0.9;
502     }
503 
504     const xs = try allocator.alloc(f32, count);
505     defer allocator.free(xs);
506     const ys = try allocator.alloc(f32, count);
507     defer allocator.free(ys);
508     const zs = try allocator.alloc(f32, count);
509     defer allocator.free(zs);
510     for (0..count) |point| {
511         xs[point] = -1.5 + @as(f32, @floatFromInt(point % 9)) * 0.36;
512         ys[point] = -1.1 + @as(f32, @floatFromInt(point % 5)) * 0.44;
513         zs[point] = -1.3 + @as(f32, @floatFromInt(point % 6)) * 0.41;
514     }
515 
516     const grid = ReferenceGrid3{
517         .values = values,
518         .nx = nx,
519         .ny = ny,
520         .nz = nz,
521         .origin_x = origin,
522         .origin_y = origin,
523         .origin_z = origin,
524         .inv_dx = 1.0 / spacing,
525         .inv_dy = 1.0 / spacing,
526         .inv_dz = 1.0 / spacing,
527     };
528 
529     const expected = try allocator.alloc(f32, count);
530     defer allocator.free(expected);
531     for (0..count) |point| {
532         expected[point] = referenceGridSample3(grid, xs[point], ys[point], zs[point], null);
533     }
534 
535     const dist = try allocator.alloc(f32, count);
536     defer allocator.free(dist);
537     @memset(dist, 0);
538 
539     const instance = GridSample{ .threads = 32 };
540     var graph = try GridSampleFamily3D.build(allocator, GridSampleFamily3D.Limits.testing, instance);
541     defer graph.deinit();
542     try graph.runCpuWithLaunch(allocator, &.{
543         kernel.argumentBuffer(f32, dist),
544         kernel.argumentBuffer(f32, values),
545         kernel.argumentBuffer(f32, xs),
546         kernel.argumentBuffer(f32, ys),
547         kernel.argumentBuffer(f32, zs),
548         kernel.argumentI32(@intCast(count)),
549         kernel.argumentF32(origin),
550         kernel.argumentF32(origin),
551         kernel.argumentF32(origin),
552         kernel.argumentF32(1.0 / spacing),
553         kernel.argumentF32(1.0 / spacing),
554         kernel.argumentF32(1.0 / spacing),
555         kernel.argumentI32(@intCast(nx)),
556         kernel.argumentI32(@intCast(ny)),
557         kernel.argumentI32(@intCast(nz)),
558     }, .{
559         .grid = .{ @intCast(gridSampleBlockCount(count, instance.threads)), 1, 1 },
560         .block = .{ instance.threads, 1, 1 },
561     });
562 
563     try testing.expectEqualSlices(f32, expected, dist);
564 }
565 
566 test "grid sample clamps queries outside the grid to the boundary" {
567     const grid = ReferenceGrid2{
568         .values = &.{ 0, 1, 2, 3 },
569         .nx = 2,
570         .ny = 2,
571         .origin_x = 0,
572         .origin_y = 0,
573         .inv_dx = 1,
574         .inv_dy = 1,
575     };
576     try testing.expectEqual(@as(f32, 0), referenceGridSample2(grid, -5, -5, null));
577     try testing.expectEqual(@as(f32, 3), referenceGridSample2(grid, 5, 5, null));
578 }
579 
580 test "grid sample gradient recovers a linear field exactly" {
581     const spacing: f32 = 0.5;
582     var values: [9]f32 = undefined;
583     for (&values, 0..) |*value, index| {
584         const ix = index % 3;
585         const iy = index / 3;
586         value.* = 2.0 * (@as(f32, @floatFromInt(ix)) * spacing) + 3.0 * (@as(f32, @floatFromInt(iy)) * spacing);
587     }
588     const grid = ReferenceGrid2{
589         .values = values[0..],
590         .nx = 3,
591         .ny = 3,
592         .origin_x = 0,
593         .origin_y = 0,
594         .inv_dx = 1.0 / spacing,
595         .inv_dy = 1.0 / spacing,
596     };
597     var gradient: [2]f32 = undefined;
598     _ = referenceGridSample2(grid, 0.6, 0.4, &gradient);
599     try testing.expectApproxEqAbs(@as(f32, 2), gradient[0], 0.0001);
600     try testing.expectApproxEqAbs(@as(f32, 3), gradient[1], 0.0001);
601 }