lib/accy/src/kernel/model/program/builder/surface/domain.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 
  3 const core = @import("../../../core/root.zig");
  4 const builder_schedule = @import("schedule.zig");
  5 const builder = core.builder;
  6 const domain_mod = core.domain;
  7 const schedule_mod = core.schedule;
  8 
  9 pub const Index1D = domain_mod.Index1D;
 10 pub const VectorIndex1D = domain_mod.VectorIndex1D;
 11 pub const DomainAxis = domain_mod.DomainAxis;
 12 pub const Domain2D = domain_mod.Domain2D;
 13 pub const Domain3D = domain_mod.Domain3D;
 14 pub const Index2D = domain_mod.Index2D;
 15 pub const Index3D = domain_mod.Index3D;
 16 
 17 pub fn index1D(owner: anytype, name: []const u8, extent: u64, threads_per_block: u32) !Index1D {
 18     return indexDimension(owner, domain_mod.domainAxis(name, extent, threads_per_block), .x);
 19 }
 20 
 21 pub fn vectorIndex1D(owner: anytype, name: []const u8, extent: u64, width: u32, threads_per_block: u32) !VectorIndex1D {
 22     if (extent == 0 or width == 0 or threads_per_block == 0) return error.InvalidFactor;
 23     const lane_count: u64 = @intCast(width);
 24     if (extent % lane_count != 0) return error.ExtentNotDivisible;
 25     const packet_extent = extent / lane_count;
 26     const thread_packets: u64 = @intCast(threads_per_block);
 27     const logical_thread_span = std.math.mul(u64, thread_packets, lane_count) catch return error.LaunchDimensionOverflow;
 28     const axis_id = try owner.axis(name, extent);
 29     var tiled_axis: ?schedule_mod.Split = null;
 30     if (packet_extent <= thread_packets) {
 31         try owner.vectorize(axis_id, width);
 32         try owner.bind(axis_id, threadTarget(.x));
 33     } else {
 34         const tiled = try builder_schedule.tile(owner, axis_id, logical_thread_span);
 35         tiled_axis = tiled;
 36         try owner.vectorize(tiled.inner, width);
 37         try owner.bind(tiled.outer, blockTarget(.x));
 38         try owner.bind(tiled.inner, threadTarget(.x));
 39     }
 40 
 41     const packet = try owner.globalId(.x);
 42     const width_value = try constantIndexFromU64(owner, lane_count);
 43     const base = try owner.mul(packet, width_value);
 44     const packet_bound = try constantIndexFromU64(owner, packet_extent);
 45     const bound = try constantIndexFromU64(owner, extent);
 46     return .{
 47         .domain = axis_id,
 48         .tile = tiled_axis,
 49         .packet = packet,
 50         .packet_bound = packet_bound,
 51         .base = base,
 52         .bound = bound,
 53         .extent = extent,
 54         .width = width,
 55         .packets = packet_extent,
 56     };
 57 }
 58 
 59 pub fn index2D(owner: anytype, domain: Domain2D) !Index2D {
 60     return .{
 61         .x = try indexDimension(owner, domain.x, .x),
 62         .y = try indexDimension(owner, domain.y, .y),
 63     };
 64 }
 65 
 66 pub fn index3D(owner: anytype, domain: Domain3D) !Index3D {
 67     return .{
 68         .x = try indexDimension(owner, domain.x, .x),
 69         .y = try indexDimension(owner, domain.y, .y),
 70         .z = try indexDimension(owner, domain.z, .z),
 71     };
 72 }
 73 
 74 pub fn guardIndexDo(owner: anytype, index: Index1D, context: anytype, comptime body: anytype) !void {
 75     var active_guard = try owner.guardIndex(index);
 76     errdefer active_guard.abort();
 77     try body(owner, index, context);
 78     try active_guard.leave();
 79 }
 80 
 81 pub fn guardVectorIndexDo(owner: anytype, index: VectorIndex1D, context: anytype, comptime body: anytype) !void {
 82     const in_bounds = try owner.compare(.lt, index.packet, index.packet_bound);
 83     var active_guard = try owner.guard(in_bounds);
 84     errdefer active_guard.abort();
 85     try body(owner, index, context);
 86     try active_guard.leave();
 87 }
 88 
 89 pub fn guardIndex2DDo(owner: anytype, index: Index2D, context: anytype, comptime body: anytype) !void {
 90     var x_guard = try owner.guardIndex(index.x);
 91     errdefer x_guard.abort();
 92     var y_guard = try owner.guardIndex(index.y);
 93     errdefer y_guard.abort();
 94     try body(owner, index, context);
 95     try y_guard.leave();
 96     try x_guard.leave();
 97 }
 98 
 99 pub fn guardIndex3DDo(owner: anytype, index: Index3D, context: anytype, comptime body: anytype) !void {
100     var x_guard = try owner.guardIndex(index.x);
101     errdefer x_guard.abort();
102     var y_guard = try owner.guardIndex(index.y);
103     errdefer y_guard.abort();
104     var z_guard = try owner.guardIndex(index.z);
105     errdefer z_guard.abort();
106     try body(owner, index, context);
107     try z_guard.leave();
108     try y_guard.leave();
109     try x_guard.leave();
110 }
111 
112 pub fn forEach1D(owner: anytype, name: []const u8, extent: u64, threads_per_block: u32, context: anytype, comptime body: anytype) !Index1D {
113     const index = try index1D(owner, name, extent, threads_per_block);
114     try guardIndexDo(owner, index, context, body);
115     return index;
116 }
117 
118 pub fn forEachVector1D(owner: anytype, name: []const u8, extent: u64, width: u32, threads_per_block: u32, context: anytype, comptime body: anytype) !VectorIndex1D {
119     const index = try vectorIndex1D(owner, name, extent, width, threads_per_block);
120     try guardVectorIndexDo(owner, index, context, body);
121     return index;
122 }
123 
124 pub fn forEach2D(owner: anytype, domain: Domain2D, context: anytype, comptime body: anytype) !Index2D {
125     const index = try index2D(owner, domain);
126     try guardIndex2DDo(owner, index, context, body);
127     return index;
128 }
129 
130 pub fn forEach3D(owner: anytype, domain: Domain3D, context: anytype, comptime body: anytype) !Index3D {
131     const index = try index3D(owner, domain);
132     try guardIndex3DDo(owner, index, context, body);
133     return index;
134 }
135 
136 fn indexDimension(owner: anytype, domain_axis: DomainAxis, dimension: builder.Dimension) !Index1D {
137     if (domain_axis.extent == 0 or domain_axis.threads_per_block == 0) return error.InvalidFactor;
138     const axis_id = try owner.axis(domain_axis.name, domain_axis.extent);
139     var tiled_axis: ?schedule_mod.Split = null;
140     if (domain_axis.extent <= domain_axis.threads_per_block) {
141         try owner.bind(axis_id, threadTarget(dimension));
142     } else {
143         const tiled = try builder_schedule.tile(owner, axis_id, domain_axis.threads_per_block);
144         tiled_axis = tiled;
145         try owner.bind(tiled.outer, blockTarget(dimension));
146         try owner.bind(tiled.inner, threadTarget(dimension));
147     }
148 
149     const index = try owner.globalId(dimension);
150     const bound_value = try owner.constantIndex(std.math.cast(i64, domain_axis.extent) orelse return error.LaunchDimensionOverflow);
151     return .{
152         .domain = axis_id,
153         .tile = tiled_axis,
154         .index = index,
155         .bound = bound_value,
156         .extent = domain_axis.extent,
157     };
158 }
159 
160 fn constantIndexFromU64(owner: anytype, value: u64) !builder.Value {
161     return owner.constantIndex(std.math.cast(i64, value) orelse return error.LaunchDimensionOverflow);
162 }
163 
164 fn blockTarget(dimension: builder.Dimension) schedule_mod.BindTarget {
165     return switch (dimension) {
166         .x => .block_x,
167         .y => .block_y,
168         .z => .block_z,
169     };
170 }
171 
172 fn threadTarget(dimension: builder.Dimension) schedule_mod.BindTarget {
173     return switch (dimension) {
174         .x => .thread_x,
175         .y => .thread_y,
176         .z => .thread_z,
177     };
178 }