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 }