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 }