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 }