lib/accy/src/profiling/reaction/kernel.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const accy = @import("accy");
  3 
  4 const Allocator = std.mem.Allocator;
  5 const Kernel = accy.kernel;
  6 
  7 pub const diffusion_u: f32 = 0.16;
  8 pub const diffusion_v: f32 = 0.08;
  9 pub const feed: f32 = 0.037;
 10 pub const kill: f32 = 0.060;
 11 pub const dt: f32 = 1.0;
 12 
 13 pub const seed_blobs = 14;
 14 pub const seed_u: f32 = 0.45;
 15 pub const seed_v: f32 = 0.28;
 16 
 17 pub fn buildStepGraph(allocator: Allocator, side: u32) !Kernel.Graph {
 18     var builder = try Kernel.Builder.init(allocator, Kernel.Builder.Limits.standard, "accy_reaction_step", &.{
 19         Kernel.dynamicBuffer(.f32),
 20         Kernel.dynamicBuffer(.f32),
 21         Kernel.dynamicBuffer(.f32),
 22         Kernel.dynamicBuffer(.f32),
 23     });
 24     errdefer builder.deinit();
 25 
 26     const u_field = builder.argument(0);
 27     const v_field = builder.argument(1);
 28     const next_u = builder.argument(2);
 29     const next_v = builder.argument(3);
 30 
 31     const gid = try builder.globalId(.x);
 32     const zero = try builder.constantIndex(0);
 33     const limit = try builder.constantIndex(@as(i64, side) - 1);
 34     const one = try builder.constantIndex(1);
 35     const stride = try builder.constantIndex(side);
 36     const row = try builder.div(gid, stride);
 37     const col = try builder.sub(gid, try builder.mul(row, stride));
 38 
 39     const col_left = try builder.select(
 40         try builder.compare(.gt, col, zero),
 41         try builder.sub(col, one),
 42         zero,
 43     );
 44     const col_right = try builder.min(try builder.add(col, one), limit);
 45     const row_up = try builder.select(
 46         try builder.compare(.gt, row, zero),
 47         try builder.sub(row, one),
 48         zero,
 49     );
 50     const row_down = try builder.min(try builder.add(row, one), limit);
 51 
 52     const center = try cellIndex(&builder, stride, row, col);
 53     const left = try cellIndex(&builder, stride, row, col_left);
 54     const right = try cellIndex(&builder, stride, row, col_right);
 55     const up = try cellIndex(&builder, stride, row_up, col);
 56     const down = try cellIndex(&builder, stride, row_down, col);
 57 
 58     const u_center = try builder.load(u_field, center);
 59     const v_center = try builder.load(v_field, center);
 60     const lap_u = try emitLaplacian(&builder, u_field, u_center, left, right, up, down);
 61     const lap_v = try emitLaplacian(&builder, v_field, v_center, left, right, up, down);
 62 
 63     const uvv = try builder.mul(u_center, try builder.mul(v_center, v_center));
 64     const one_f = try builder.constantFloat(.f32, 1.0);
 65     const zero_f = try builder.constantFloat(.f32, 0.0);
 66     const feed_term = try builder.mul(
 67         try builder.sub(one_f, u_center),
 68         try builder.constantFloat(.f32, feed),
 69     );
 70     const kill_term = try builder.mul(
 71         v_center,
 72         try builder.constantFloat(.f32, feed + kill),
 73     );
 74 
 75     const du = try builder.add(
 76         try builder.sub(
 77             try builder.mul(lap_u, try builder.constantFloat(.f32, diffusion_u)),
 78             uvv,
 79         ),
 80         feed_term,
 81     );
 82     const dv = try builder.sub(
 83         try builder.add(
 84             try builder.mul(lap_v, try builder.constantFloat(.f32, diffusion_v)),
 85             uvv,
 86         ),
 87         kill_term,
 88     );
 89 
 90     const dt_value = try builder.constantFloat(.f32, dt);
 91     const stepped_u = try builder.add(u_center, try builder.mul(du, dt_value));
 92     const stepped_v = try builder.add(v_center, try builder.mul(dv, dt_value));
 93     const clamped_u = try builder.min(try builder.max(stepped_u, zero_f), one_f);
 94     const clamped_v = try builder.min(try builder.max(stepped_v, zero_f), one_f);
 95 
 96     try builder.store(clamped_u, next_u, center);
 97     try builder.store(clamped_v, next_v, center);
 98     try builder.return_();
 99 
100     return builder.finish();
101 }
102 
103 fn cellIndex(builder: *Kernel.Builder, stride: Kernel.Value, row: Kernel.Value, col: Kernel.Value) !Kernel.Value {
104     return builder.add(try builder.mul(row, stride), col);
105 }
106 
107 fn emitLaplacian(
108     builder: *Kernel.Builder,
109     field: Kernel.Value,
110     center: Kernel.Value,
111     left: Kernel.Value,
112     right: Kernel.Value,
113     up: Kernel.Value,
114     down: Kernel.Value,
115 ) !Kernel.Value {
116     const left_value = try builder.load(field, left);
117     const right_value = try builder.load(field, right);
118     const up_value = try builder.load(field, up);
119     const down_value = try builder.load(field, down);
120     const neighbors = try builder.add(
121         try builder.add(left_value, right_value),
122         try builder.add(up_value, down_value),
123     );
124     return builder.sub(
125         neighbors,
126         try builder.mul(center, try builder.constantFloat(.f32, 4.0)),
127     );
128 }
129 
130 pub fn seedState(u: []f32, v: []f32, side: u32) void {
131     std.debug.assert(u.len == @as(usize, side) * side);
132     std.debug.assert(v.len == u.len);
133     @memset(u, 1.0);
134     @memset(v, 0.0);
135 
136     const side_f: f32 = @floatFromInt(side);
137     const radius = @max(side_f / 96.0, 2.0);
138     for (0..seed_blobs) |blob| {
139         const cx = hashUnit(blob * 2 + 1) * 0.84 + 0.08;
140         const cy = hashUnit(blob * 2 + 2) * 0.84 + 0.08;
141         stampBlob(u, v, side, cx * side_f, cy * side_f, radius * (0.7 + 0.6 * hashUnit(blob + 97)));
142     }
143 }
144 
145 fn stampBlob(u: []f32, v: []f32, side: u32, cx: f32, cy: f32, radius: f32) void {
146     const lo_row: usize = @intFromFloat(@max(cy - radius - 1.0, 0.0));
147     const hi_row: usize = @min(@as(usize, @intFromFloat(cy + radius + 1.0)), side - 1);
148     const lo_col: usize = @intFromFloat(@max(cx - radius - 1.0, 0.0));
149     const hi_col: usize = @min(@as(usize, @intFromFloat(cx + radius + 1.0)), side - 1);
150     var row = lo_row;
151     while (row <= hi_row) : (row += 1) {
152         var col = lo_col;
153         while (col <= hi_col) : (col += 1) {
154             const dx = @as(f32, @floatFromInt(col)) - cx;
155             const dy = @as(f32, @floatFromInt(row)) - cy;
156             if (dx * dx + dy * dy <= radius * radius) {
157                 const index = row * side + col;
158                 u[index] = seed_u;
159                 v[index] = seed_v;
160             }
161         }
162     }
163 }
164 
165 fn hashUnit(seed: usize) f32 {
166     var state: u64 = @as(u64, @intCast(seed)) +% 0x9e3779b97f4a7c15;
167     state = (state ^ (state >> 30)) *% 0xbf58476d1ce4e5b9;
168     state = (state ^ (state >> 27)) *% 0x94d049bb133111eb;
169     state ^= state >> 31;
170     return @as(f32, @floatFromInt(state & 0xffffff)) / 16777215.0;
171 }
172 
173 pub fn referenceStep(u: []const f32, v: []const f32, next_u: []f32, next_v: []f32, side: u32) void {
174     const cells = @as(usize, side) * side;
175     std.debug.assert(u.len == cells);
176     std.debug.assert(v.len == cells);
177     std.debug.assert(next_u.len == cells);
178     std.debug.assert(next_v.len == cells);
179     for (0..side) |row| {
180         for (0..side) |col| {
181             const index = row * side + col;
182             const u_center = u[index];
183             const v_center = v[index];
184             const lap_u = hostLaplacian(u, row, col, side);
185             const lap_v = hostLaplacian(v, row, col, side);
186             const uvv = u_center * v_center * v_center;
187             const du = diffusion_u * lap_u - uvv + feed * (1.0 - u_center);
188             const dv = diffusion_v * lap_v + uvv - (feed + kill) * v_center;
189             next_u[index] = clamp01(u_center + dt * du);
190             next_v[index] = clamp01(v_center + dt * dv);
191         }
192     }
193 }
194 
195 fn hostLaplacian(field: []const f32, row: usize, col: usize, side: u32) f32 {
196     const stride: usize = side;
197     const center = field[row * stride + col];
198     const left = field[row * stride + clampOffset(col, -1, side)];
199     const right = field[row * stride + clampOffset(col, 1, side)];
200     const up = field[clampOffset(row, -1, side) * stride + col];
201     const down = field[clampOffset(row, 1, side) * stride + col];
202     return left + right + up + down - 4.0 * center;
203 }
204 
205 fn clampOffset(coord: usize, delta: i64, side: u32) usize {
206     const shifted = @as(i64, @intCast(coord)) + delta;
207     const limit = @as(i64, side) - 1;
208     return @intCast(@max(0, @min(shifted, limit)));
209 }
210 
211 fn clamp01(value: f32) f32 {
212     return @max(0.0, @min(value, 1.0));
213 }
214 
215 test "buildStepGraph verifies" {
216     var graph = try buildStepGraph(std.testing.allocator, 64);
217     defer graph.deinit();
218     try graph.verify();
219 }
220 
221 test "seedState plants blobs inside the field" {
222     const side: u32 = 96;
223     const cells = @as(usize, side) * side;
224     const u = try std.testing.allocator.alloc(f32, cells);
225     defer std.testing.allocator.free(u);
226     const v = try std.testing.allocator.alloc(f32, cells);
227     defer std.testing.allocator.free(v);
228 
229     seedState(u, v, side);
230 
231     var seeded: usize = 0;
232     for (u, v) |u_value, v_value| {
233         if (u_value < 1.0) {
234             try std.testing.expectApproxEqAbs(seed_u, u_value, 0.0001);
235             try std.testing.expectApproxEqAbs(seed_v, v_value, 0.0001);
236             seeded += 1;
237         } else {
238             try std.testing.expectApproxEqAbs(@as(f32, 0.0), v_value, 0.0001);
239         }
240     }
241     try std.testing.expect(seeded > 100);
242     try std.testing.expect(seeded < cells / 8);
243 }
244 
245 test "referenceStep grows v inside a fresh blob and stays bounded" {
246     const side: u32 = 64;
247     const cells = @as(usize, side) * side;
248     var buffers: [4][]f32 = undefined;
249     for (&buffers) |*buffer| buffer.* = try std.testing.allocator.alloc(f32, cells);
250     defer for (buffers) |buffer| std.testing.allocator.free(buffer);
251 
252     seedState(buffers[0], buffers[1], side);
253     var v_before: f64 = 0.0;
254     for (buffers[1]) |value| v_before += value;
255 
256     var flip = false;
257     for (0..64) |_| {
258         const u = if (flip) buffers[2] else buffers[0];
259         const v = if (flip) buffers[3] else buffers[1];
260         const nu = if (flip) buffers[0] else buffers[2];
261         const nv = if (flip) buffers[1] else buffers[3];
262         referenceStep(u, v, nu, nv, side);
263         flip = !flip;
264     }
265 
266     const v_final = if (flip) buffers[3] else buffers[1];
267     var v_after: f64 = 0.0;
268     for (v_final) |value| {
269         try std.testing.expect(value >= 0.0 and value <= 1.0);
270         v_after += value;
271     }
272     try std.testing.expect(v_after > v_before);
273 }