lib/accy/src/preparation/kernelization/lowering/scan.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const gpu = @import("gpu");
3 const choir_abi = @import("choir_abi");
4 const alloc_fixed = @import("alloc_fixed");
5 const generated_abi = @import("abi.zig");
6 const generated_builder = @import("builder.zig");
7 const common = @import("common.zig");
8 const generated_guard = @import("guard.zig");
9 const generated_name = @import("name.zig");
10 const generated_schedule = @import("schedule.zig");
11 const preparation = @import("../../root.zig");
12
13 const ir = common.ir;
14 const dialect_mod = common.dialect_mod;
15 const kernel_root = common.kernel_root;
16 const bufferization = common.bufferization;
17 const kernelization_model = @import("../model/root.zig");
18 const schedule_planning = common.schedule_planning;
19 const target_facts = preparation.target;
20 const bufferSlotById = common.bufferSlotById;
21 const externalInputIndex = common.externalInputIndex;
22 const mapKernelBuildError = common.mapKernelBuildError;
23 const isName = common.isName;
24
25 const LoweredKernel = kernelization_model.LoweredKernel;
26 const Schedule = target_facts.GeneratedScanSchedule;
27
28 const block_schedules = [_]Schedule{
29 .{ .threads = 512, .items = 16 },
30 .{ .threads = 256, .items = 16 },
31 };
32
33 pub const schedule_version: u32 = 1;
34 pub const max_schedule_candidates = block_schedules.len;
35
36 pub fn scheduleCandidates(
37 total: u64,
38 format: ?gpu.ArtifactFormat,
39 buffer: *[max_schedule_candidates]Schedule,
40 ) []const Schedule {
41 if (format != .cuda_ptx) return buffer[0..0];
42 var count: usize = 0;
43 for (block_schedules) |schedule| {
44 if (!scheduleViable(schedule, total)) continue;
45 buffer[count] = schedule;
46 count += 1;
47 }
48 return buffer[0..count];
49 }
50
51 fn scheduleViable(schedule: Schedule, total: u64) bool {
52 const tile = @as(u64, schedule.threads) * schedule.items;
53 if (total % tile != 0) return false;
54 return std.math.cast(u32, total / tile) != null;
55 }
56
57 const ScanDescription = struct {
58 total: u64,
59 blocks: u32,
60 input: *ir.Value,
61 scratch: ?*ir.Value,
62 };
63
64 fn BlockScanBodyType(comptime scan_threads: u32, comptime scan_items: u32) type {
65 return struct {
66 fn emit(logical: anytype, ctx: anytype) !void {
67 try emitBlockScanBodyScheduled(scan_threads, scan_items, logical, ctx);
68 }
69 };
70 }
71
72 pub fn lower(
73 allocator: std.mem.Allocator,
74 ir_ctx: *ir.Context,
75 outline: kernelization_model.KernelOutline,
76 work: schedule_planning.ScheduleWorkItem,
77 buffer_plan: *const bufferization.BufferPlanAnalysis,
78 format: ?gpu.ArtifactFormat,
79 scan_schedules: ?[]const u8,
80 ) common.LoweringError!LoweredKernel {
81 const desc = try scanDescriptionForWork(work);
82
83 const input_dtypes = allocator.alloc(choir_abi.DType, outline.inputCount()) catch return error.OutOfMemory;
84 defer allocator.free(input_dtypes);
85 for (outline.input_slot_ids, 0..) |slot_id, index| {
86 const slot = bufferSlotById(buffer_plan, slot_id) orelse return error.InvalidArtifact;
87 input_dtypes[index] = slot.dtype;
88 }
89
90 var abi = try generated_abi.flatTyped(allocator, .f32, input_dtypes);
91 defer abi.deinit(allocator);
92
93 if (format == .cuda_ptx) {
94 if (scan_schedules) |encoded| {
95 if (try target_facts.resolveGeneratedScanSchedule(encoded, desc.total)) |schedule| {
96 if (scheduleViable(schedule, desc.total)) {
97 if (try lowerBlockScanSchedule(schedule, allocator, ir_ctx, outline, work, buffer_plan, desc, abi)) |lowered| {
98 return lowered;
99 }
100 }
101 }
102 }
103 var candidate_buffer: [max_schedule_candidates]Schedule = undefined;
104 for (scheduleCandidates(desc.total, format, &candidate_buffer)) |schedule| {
105 if (try lowerBlockScanSchedule(schedule, allocator, ir_ctx, outline, work, buffer_plan, desc, abi)) |lowered| {
106 return lowered;
107 }
108 }
109 }
110
111 const entry_name = try generated_name.scanSerial(allocator, desc.total, work.id);
112 errdefer allocator.free(entry_name);
113 return generated_builder.withoutLaunch(allocator, ir_ctx, work.id, entry_name, abi.params(), generated_schedule.flat(), .{
114 .abi = abi,
115 .allocator = allocator,
116 .desc = desc,
117 .outline = outline,
118 .buffer_plan = buffer_plan,
119 }, emitSerialScanBody);
120 }
121
122 fn scanDescriptionForWork(
123 work: schedule_planning.ScheduleWorkItem,
124 ) common.LoweringError!ScanDescription {
125 if (work.kind != .scan) return error.UnsupportedOperation;
126 if (work.ops.len != 1) return error.UnsupportedOperation;
127 const op = work.ops[0];
128 if (!isName(op.name.name, dialect_mod.AccyDialect.CumsumOp.operation_name)) return error.UnsupportedOperation;
129
130 const axis_attr = op.getAttrAs(ir.Attribute.IntegerAttr, "axis") orelse return error.InvalidArtifact;
131 if (axis_attr.getValue() != 0) return error.UnsupportedOperation;
132
133 const input = op.getOperand(0) orelse return error.InvalidArtifact;
134 var dims_arena_buffer: [256]u8 = undefined;
135 var dims_arena = alloc_fixed.FixedBuffer.init(dims_arena_buffer[0..]);
136 const input_type = dialect_mod.decodeTensorType(dims_arena.allocator(), input.type) catch return error.InvalidArtifact;
137 if (input_type.dtype != .f32 or input_type.dims.len != 1) return error.UnsupportedOperation;
138 if (input_type.dims[0] < 1) return error.UnsupportedOperation;
139 const total: u64 = @intCast(input_type.dims[0]);
140 return .{
141 .total = total,
142 .blocks = 0,
143 .input = input,
144 .scratch = op.getOperand(1),
145 };
146 }
147
148 fn lowerBlockScanSchedule(
149 schedule: Schedule,
150 allocator: std.mem.Allocator,
151 ir_ctx: *ir.Context,
152 outline: kernelization_model.KernelOutline,
153 work: schedule_planning.ScheduleWorkItem,
154 buffer_plan: *const bufferization.BufferPlanAnalysis,
155 desc: ScanDescription,
156 abi: generated_abi.Flat,
157 ) common.LoweringError!?LoweredKernel {
158 inline for (block_schedules) |candidate| {
159 if (candidate.threads == schedule.threads and candidate.items == schedule.items) {
160 return lowerBlockScan(
161 candidate.threads,
162 candidate.items,
163 BlockScanBodyType(candidate.threads, candidate.items).emit,
164 allocator,
165 ir_ctx,
166 outline,
167 work,
168 buffer_plan,
169 desc,
170 abi,
171 );
172 }
173 }
174 return null;
175 }
176
177 fn lowerBlockScan(
178 comptime scan_threads: u32,
179 comptime scan_items: u32,
180 comptime emit_body: anytype,
181 allocator: std.mem.Allocator,
182 ir_ctx: *ir.Context,
183 outline: kernelization_model.KernelOutline,
184 work: schedule_planning.ScheduleWorkItem,
185 buffer_plan: *const bufferization.BufferPlanAnalysis,
186 desc: ScanDescription,
187 abi: generated_abi.Flat,
188 ) common.LoweringError!?LoweredKernel {
189 const scan_tile = scan_threads * scan_items;
190 if (desc.total % scan_tile != 0) return null;
191 const blocks = std.math.cast(u32, desc.total / scan_tile) orelse return error.UnsupportedOperation;
192 var scheduled_desc = desc;
193 scheduled_desc.blocks = blocks;
194 if (blocks == 1) {
195 scheduled_desc.scratch = null;
196 } else if (scheduled_desc.scratch == null) {
197 return null;
198 }
199
200 const entry_name = try generated_name.scanLookback(allocator, scheduled_desc.total, scheduled_desc.blocks, work.id);
201 errdefer allocator.free(entry_name);
202 var lowered = try generated_builder.withoutLaunch(allocator, ir_ctx, work.id, entry_name, abi.params(), generated_schedule.flatThreads(scan_threads), .{
203 .abi = abi,
204 .allocator = allocator,
205 .desc = scheduled_desc,
206 .outline = outline,
207 .buffer_plan = buffer_plan,
208 }, emit_body);
209 lowered.body = .{ .scan = .{
210 .blocks = scheduled_desc.blocks,
211 .threads = scan_threads,
212 .items = scan_items,
213 } };
214 if (scheduled_desc.scratch != null) lowered.scratch_fill_pattern = 0;
215 return lowered;
216 }
217
218 fn claim_scan_block(inner: anytype, guard_ctx: anytype) !void {
219 const bid = try inner.atomicRmw(
220 .add,
221 guard_ctx.one,
222 guard_ctx.scratch,
223 guard_ctx.zero,
224 );
225 try inner.storeIndex(bid, guard_ctx.bid_share, guard_ctx.zero);
226 }
227
228 fn initialize_single_scan_block(inner: anytype, guard_ctx: anytype) !void {
229 try inner.storeIndex(guard_ctx.zero_u, guard_ctx.bid_share, guard_ctx.zero);
230 }
231
232 fn store_scan_warp_sum(inner: anytype, guard_ctx: anytype) !void {
233 try inner.storeIndex(guard_ctx.value, guard_ctx.warp_sums, guard_ctx.warp);
234 }
235
236 fn publish_scan_lookback(inner: anytype, guard_ctx: anytype) !void {
237 const total_bits = try inner.bitcast(guard_ctx.block_total, .u32);
238 try inner.storeIndex(total_bits, guard_ctx.scratch, try inner.add(guard_ctx.partial_base, guard_ctx.bid));
239 try inner.fence(.device);
240 _ = try inner.atomicRmw(.add, guard_ctx.one_u, guard_ctx.scratch, try inner.add(guard_ctx.status_base, guard_ctx.bid));
241
242 var lookback = try inner.whileScope(
243 &.{ guard_ctx.bid, guard_ctx.zero_f },
244 &.{ guard_ctx.bid.valueType(), guard_ctx.zero_f.valueType() },
245 );
246 errdefer lookback.abort();
247 const remaining = lookback.beforeArg(0) orelse return error.UnsupportedOperation;
248 const prefix = lookback.beforeArg(1) orelse return error.UnsupportedOperation;
249 const keep_going = try inner.compare(.gt, remaining, guard_ctx.zero);
250 try lookback.condition(keep_going, &.{ remaining, prefix });
251
252 const body_remaining = lookback.afterArg(0) orelse return error.UnsupportedOperation;
253 const body_prefix = lookback.afterArg(1) orelse return error.UnsupportedOperation;
254 const body_j = try inner.sub(body_remaining, try inner.constantIndex(1));
255 const status_index = try inner.add(guard_ctx.status_base, body_j);
256
257 var spin = try inner.whileScope(&.{guard_ctx.zero_u}, &.{guard_ctx.zero_u.valueType()});
258 errdefer spin.abort();
259 _ = spin.beforeArg(0) orelse return error.UnsupportedOperation;
260 const observed = try inner.atomicRmw(.add, guard_ctx.zero_u, guard_ctx.scratch, status_index);
261 const still_waiting = try inner.compare(.eq, observed, guard_ctx.zero_u);
262 try spin.condition(still_waiting, &.{observed});
263 const spin_after = spin.afterArg(0) orelse return error.UnsupportedOperation;
264 try spin.leave(&.{spin_after});
265 const status = spin.result(0) orelse return error.UnsupportedOperation;
266
267 try inner.fence(.device);
268 const done = try inner.compare(.eq, status, try inner.constantInt(.u32, 2));
269 const partial_bits = try inner.loadIndex(guard_ctx.scratch, try inner.add(guard_ctx.partial_base, body_j));
270 const inclusive_bits = try inner.loadIndex(guard_ctx.scratch, try inner.add(guard_ctx.inclusive_base, body_j));
271 const partial = try inner.bitcast(partial_bits, .f32);
272 const inclusive = try inner.bitcast(inclusive_bits, .f32);
273 const take = try inner.select(done, inclusive, partial);
274 const next_prefix = try inner.add(body_prefix, take);
275 const next_remaining = try inner.select(done, guard_ctx.zero, body_j);
276 try lookback.leave(&.{ next_remaining, next_prefix });
277
278 const final_prefix = lookback.result(1) orelse return error.UnsupportedOperation;
279 const inclusive_total = try inner.add(final_prefix, guard_ctx.block_total);
280 const inclusive_total_bits = try inner.bitcast(inclusive_total, .u32);
281 try inner.storeIndex(inclusive_total_bits, guard_ctx.scratch, try inner.add(guard_ctx.inclusive_base, guard_ctx.bid));
282 try inner.fence(.device);
283 _ = try inner.atomicRmw(.add, guard_ctx.one_u, guard_ctx.scratch, try inner.add(guard_ctx.status_base, guard_ctx.bid));
284
285 try inner.storeIndex(final_prefix, guard_ctx.prefix_share, guard_ctx.zero);
286 }
287
288 fn emitBlockScanBodyScheduled(
289 comptime scan_threads: u32,
290 comptime scan_items: u32,
291 logical: anytype,
292 ctx: anytype,
293 ) !void {
294 const scan_tile = scan_threads * scan_items;
295 const scan_warp_count = scan_threads / 32;
296 const desc: ScanDescription = ctx.desc;
297
298 _ = try logical.index1D("lane", @as(u64, desc.blocks) * scan_threads);
299 const tid = try logical.threadId(.x);
300
301 const input_slot = ctx.buffer_plan.getSlot(desc.input) orelse return error.UnsupportedOperation;
302 const input_index = externalInputIndex(ctx.outline, input_slot.id) orelse return error.UnsupportedOperation;
303 const input = ctx.abi.input(logical, input_index);
304 const out = ctx.abi.output(logical);
305
306 var scratch_value: ?kernel_root.Value = null;
307 if (desc.scratch) |scratch_ir| {
308 const scratch_slot = ctx.buffer_plan.getSlot(scratch_ir) orelse return error.UnsupportedOperation;
309 const scratch_index = externalInputIndex(ctx.outline, scratch_slot.id) orelse return error.UnsupportedOperation;
310 scratch_value = ctx.abi.input(logical, scratch_index);
311 }
312
313 const bid_share = try logical.sharedBuffer(.u32, 1);
314 const prefix_share = try logical.sharedBuffer(.f32, 1);
315 const warp_sums = try logical.sharedBuffer(.f32, scan_warp_count);
316 const tile_share = try logical.sharedBuffer(.f32, scan_tile);
317
318 const zero_index = try logical.constantIndex(0);
319 const one_u32 = try logical.constantInt(.u32, 1);
320 const zero_u32 = try logical.constantInt(.u32, 0);
321 const zero_f = try logical.constantFloat(.f32, 0.0);
322
323 const tid_zero = try logical.compare(.eq, tid, zero_index);
324
325 if (scratch_value) |scratch| {
326 try logical.guardDo(
327 tid_zero,
328 .{ .scratch = scratch, .bid_share = bid_share, .one = one_u32, .zero = zero_index },
329 claim_scan_block,
330 );
331 try logical.barrier(.block);
332 } else {
333 try logical.guardDo(
334 tid_zero,
335 .{ .bid_share = bid_share, .zero = zero_index, .zero_u = zero_u32 },
336 initialize_single_scan_block,
337 );
338 try logical.barrier(.block);
339 }
340
341 const bid_u = try logical.loadIndex(bid_share, zero_index);
342 const bid = try logical.castIndex(bid_u);
343
344 const tile_extent = try logical.constantIndex(scan_tile);
345 const items_extent = try logical.constantIndex(scan_items);
346
347 const tile_base = try logical.mul(bid, tile_extent);
348 {
349 var chunk: u32 = 0;
350 while (chunk < scan_items / 4) : (chunk += 1) {
351 const flat_quad = try logical.add(try logical.mul(try logical.constantIndex(@as(i64, chunk) * scan_threads), try logical.constantIndex(4)), try logical.mul(tid, try logical.constantIndex(4)));
352 const loaded = try logical.loadVector(input, try logical.add(tile_base, flat_quad), 4);
353 try logical.storeIndex(loaded, tile_share, flat_quad);
354 }
355 }
356 try logical.barrier(.block);
357
358 var values: [scan_items]kernel_root.Value = undefined;
359 {
360 const local_base = try logical.mul(tid, items_extent);
361 var quad: u32 = 0;
362 while (quad < scan_items / 4) : (quad += 1) {
363 const loaded = try logical.loadVector(tile_share, try logical.add(local_base, try logical.constantIndex(@as(i64, quad) * 4)), 4);
364 for (0..4) |lane_index| {
365 values[quad * 4 + lane_index] = try logical.extractLane(loaded, @intCast(lane_index), .f32);
366 }
367 }
368 }
369
370 var running: [scan_items]kernel_root.Value = undefined;
371 running[0] = values[0];
372 for (1..scan_items) |i| {
373 running[i] = try logical.add(running[i - 1], values[i]);
374 }
375 const thread_total = running[scan_items - 1];
376
377 const warp_inclusive = try logical.warpScan(.add, .inclusive, thread_total);
378 const warp_exclusive = try logical.sub(warp_inclusive, thread_total);
379
380 const warp_extent = try logical.constantIndex(32);
381 const warp = try logical.div(tid, warp_extent);
382 const lane = try logical.sub(tid, try logical.mul(warp, warp_extent));
383 const lane_last = try logical.compare(.eq, lane, try logical.constantIndex(31));
384 try logical.guardDo(
385 lane_last,
386 .{ .warp_sums = warp_sums, .warp = warp, .value = warp_inclusive },
387 store_scan_warp_sum,
388 );
389 try logical.barrier(.block);
390
391 var warp_offset = zero_f;
392 var block_total = zero_f;
393 for (0..scan_warp_count) |w| {
394 const w_extent = try logical.constantIndex(@as(i64, @intCast(w)));
395 const w_sum = try logical.loadIndex(warp_sums, w_extent);
396 const before = try logical.compare(.lt, w_extent, warp);
397 const contribution = try logical.select(before, w_sum, zero_f);
398 warp_offset = try logical.add(warp_offset, contribution);
399 block_total = try logical.add(block_total, w_sum);
400 }
401
402 var block_prefix = zero_f;
403 if (scratch_value) |scratch| {
404 const blocks_extent = try logical.constantIndex(desc.blocks);
405 const status_base = try logical.constantIndex(1);
406 const partial_base = try logical.add(status_base, blocks_extent);
407 const inclusive_base = try logical.add(partial_base, blocks_extent);
408
409 try logical.guardDo(tid_zero, .{
410 .scratch = scratch,
411 .prefix_share = prefix_share,
412 .bid = bid,
413 .block_total = block_total,
414 .status_base = status_base,
415 .partial_base = partial_base,
416 .inclusive_base = inclusive_base,
417 .one_u = one_u32,
418 .zero_u = zero_u32,
419 .zero = zero_index,
420 .zero_f = zero_f,
421 }, publish_scan_lookback);
422 try logical.barrier(.block);
423 block_prefix = try logical.loadIndex(prefix_share, zero_index);
424 }
425
426 const base_value = try logical.add(block_prefix, try logical.add(warp_offset, warp_exclusive));
427 try logical.barrier(.block);
428 {
429 const local_base = try logical.mul(tid, items_extent);
430 var quad: u32 = 0;
431 while (quad < scan_items / 4) : (quad += 1) {
432 var lanes: [4]kernel_root.Value = undefined;
433 for (0..4) |lane_index| {
434 lanes[lane_index] = try logical.add(base_value, running[quad * 4 + lane_index]);
435 }
436 const packed_values = try logical.packVector(lanes);
437 try logical.storeIndex(packed_values, tile_share, try logical.add(local_base, try logical.constantIndex(@as(i64, quad) * 4)));
438 }
439 }
440 try logical.barrier(.block);
441 {
442 var chunk: u32 = 0;
443 while (chunk < scan_items / 4) : (chunk += 1) {
444 const flat_quad = try logical.add(try logical.mul(try logical.constantIndex(@as(i64, chunk) * scan_threads), try logical.constantIndex(4)), try logical.mul(tid, try logical.constantIndex(4)));
445 const quad_values = try logical.loadVector(tile_share, flat_quad, 4);
446 try logical.storeIndex(quad_values, out, try logical.add(tile_base, flat_quad));
447 }
448 }
449 }
450
451 fn emit_serial_scan_index(
452 inner: anytype,
453 domain_index: kernel_root.Index1D,
454 guard_ctx: anytype,
455 ) !void {
456 _ = domain_index;
457 const desc_inner: ScanDescription = guard_ctx.desc;
458 const input_slot = guard_ctx.buffer_plan.getSlot(desc_inner.input) orelse
459 return error.UnsupportedOperation;
460 const input_index = externalInputIndex(guard_ctx.outline, input_slot.id) orelse
461 return error.UnsupportedOperation;
462 const input = guard_ctx.abi.input(inner, input_index);
463 const out = guard_ctx.abi.output(inner);
464
465 const zero_f = try inner.constantFloat(.f32, 0.0);
466 const total_extent = try inner.constantIndex(@as(i64, @intCast(desc_inner.total)));
467 const one = try inner.constantIndex(1);
468 var scope = try inner.forScope(
469 try inner.constantIndex(0),
470 total_extent,
471 one,
472 &.{zero_f},
473 &.{zero_f.valueType()},
474 );
475 errdefer scope.abort();
476 const j = scope.inductionVar();
477 const running = scope.iterArg(0) orelse return error.UnsupportedOperation;
478 const value = try inner.loadIndex(input, j);
479 const next = try inner.add(running, value);
480 try inner.storeIndex(next, out, j);
481 try scope.leave(&.{next});
482 }
483
484 fn emitSerialScanBody(logical: anytype, ctx: anytype) !void {
485 const desc: ScanDescription = ctx.desc;
486 const domain = try logical.index1D("row", 1);
487 try generated_guard.countIndexDo(logical, domain, ctx.abi.count(logical), .{
488 .abi = ctx.abi,
489 .desc = desc,
490 .outline = ctx.outline,
491 .buffer_plan = ctx.buffer_plan,
492 }, emit_serial_scan_index);
493 }