lib/accy/src/validation/conformance/run.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const gpu = @import("gpu");
  3 const accy = @import("accy");
  4 const pretty = @import("pretty");
  5 const sys = @import("sys");
  6 const ptx = @import("accy_validation_ptx");
  7 const conformance = @import("root.zig");
  8 
  9 const cases = conformance.cases;
 10 const harness = conformance.harness;
 11 const records = conformance.records;
 12 
 13 const BackendChoice = enum {
 14     cuda,
 15     metal,
 16     vulkan,
 17 };
 18 
 19 pub fn main(init: sys.process.Init) !void {
 20     sys.env.installProcessEnvironment(init.minimal.environ);
 21 
 22     var gpa: std.heap.DebugAllocator(.{}) = .{};
 23     defer std.debug.assert(gpa.deinit() == .ok);
 24     const allocator = gpa.allocator();
 25 
 26     var stdout_buffer: [4096]u8 = undefined;
 27     var stdout_writer = sys.stdio.stdout().writer(sys.stdio.debugIo(), &stdout_buffer);
 28     const writer = &stdout_writer.interface;
 29     defer writer.flush() catch {};
 30 
 31     const backend = try requestedBackend();
 32     switch (backend) {
 33         .cuda => try runCuda(allocator, writer),
 34         .metal => try runMetal(allocator, writer),
 35         .vulkan => try runVulkan(allocator, writer),
 36     }
 37 }
 38 
 39 fn runCuda(allocator: std.mem.Allocator, writer: *std.Io.Writer) !void {
 40     const ptxas_target = try ptx.assemble.probeTarget(allocator);
 41     const assemble_only = forcedAssembleMode();
 42 
 43     var cuda_state: ?harness.CudaState = if (assemble_only) null else harness.CudaState.initDevice(allocator, 0) catch |err| switch (err) {
 44         error.RuntimeUnavailable => null,
 45         else => return err,
 46     };
 47     defer if (cuda_state) |*state| state.deinit();
 48 
 49     if (cuda_state == null and ptxas_target == null) {
 50         try writeJsonLine(writer, records.Skip{ .reason = "cuda runtime and ptxas unavailable", .backend = "cuda" });
 51         return;
 52     }
 53 
 54     const scratch: ?[]u8 = if (ptxas_target != null) try ptx.assemble.scratchDir(allocator) else null;
 55     defer if (scratch) |dir| ptx.assemble.removeScratchDir(allocator, dir);
 56 
 57     const assembler: ?harness.Assembler = if (ptxas_target) |target_name| .{
 58         .dir = scratch.?,
 59         .target = target_name,
 60     } else null;
 61 
 62     var emission_state = ptx.backend.State.init(allocator);
 63     const mode: harness.Mode = if (cuda_state != null) .execute else .assemble;
 64     const handle = if (cuda_state) |*state| state.handle() else emission_state.handle();
 65     const gate = harness.Gate{ .mode = mode, .assembler = assembler };
 66 
 67     try runAllCases(allocator, writer, "cuda", handle, gate);
 68 }
 69 
 70 fn runMetal(allocator: std.mem.Allocator, writer: *std.Io.Writer) !void {
 71     if (!accy.validation.gating.enabledByBuildOptions(.metal)) {
 72         try writeJsonLine(writer, records.Skip{ .reason = "metal build flag disabled", .backend = "metal" });
 73         return;
 74     }
 75     if (!accy.validation.gating.appleMetalPlatform()) {
 76         try writeJsonLine(writer, records.Skip{ .reason = "unsupported platform", .backend = "metal" });
 77         return;
 78     }
 79 
 80     var state = harness.MetalState.initDevice(allocator) catch |err| switch (err) {
 81         error.RuntimeUnavailable => {
 82             try writeJsonLine(writer, records.Skip{ .reason = "metal runtime unavailable", .backend = "metal" });
 83             return;
 84         },
 85         else => return err,
 86     };
 87     defer state.deinit();
 88 
 89     try runAllCases(allocator, writer, "metal", state.handle(), .{ .mode = .execute, .assembler = null });
 90 }
 91 
 92 /// Runs on device 0 of whichever Vulkan loader the process opens. On Linux a vendor driver, such
 93 /// as NVIDIA's, runs only in a process built for the host's C library:
 94 /// `-Dtarget=native-linux-gnu` links glibc through the standard dynamic linker, so the system
 95 /// loader and its driver manifests resolve as they do for a shipped binary. The default build links
 96 /// the pinned Nix glibc and reaches only drivers built against it, such as the pinned Mesa's
 97 /// lavapipe or RADV named through `VK_DRIVER_FILES`, with the pinned loader on `LD_LIBRARY_PATH`.
 98 /// Either mismatch ends in the skip below.
 99 fn runVulkan(allocator: std.mem.Allocator, writer: *std.Io.Writer) !void {
100     var state = harness.VulkanState.initDevice(allocator, 0) catch |err| switch (err) {
101         error.RuntimeUnavailable => {
102             try writeJsonLine(writer, records.Skip{ .reason = "vulkan runtime unavailable", .backend = "vulkan" });
103             return;
104         },
105         else => return err,
106     };
107     defer state.deinit();
108 
109     try runAllCases(allocator, writer, "vulkan", state.handle(), .{ .mode = .execute, .assembler = null });
110 }
111 
112 fn runAllCases(
113     allocator: std.mem.Allocator,
114     writer: *std.Io.Writer,
115     backend_name: []const u8,
116     handle: harness.BackendHandle,
117     gate: harness.Gate,
118 ) !void {
119     const capabilities = try handle.queryCapabilities();
120     var failures: usize = 0;
121     var unsupported: usize = 0;
122     var invalid: usize = 0;
123     var assembled: usize = 0;
124     var total_compile_ns: u64 = 0;
125     var total_launch_read_ns: u64 = 0;
126 
127     inline for (cases.all) |Spec| {
128         var case = caseResult(Spec, allocator, handle, gate, capabilities);
129         case.backend = backend_name;
130         try writeJsonLine(writer, case);
131         total_compile_ns += case.compile_ns;
132         total_launch_read_ns += case.launch_read_ns;
133         if (std.mem.eql(u8, case.status, "unsupported")) {
134             unsupported += 1;
135         } else if (std.mem.eql(u8, case.status, "invalid")) {
136             invalid += 1;
137         } else if (std.mem.eql(u8, case.status, "assembled")) {
138             assembled += 1;
139         } else if (!std.mem.eql(u8, case.status, "pass")) {
140             failures += 1;
141         }
142     }
143 
144     try writeJsonLine(writer, records.Summary{
145         .backend = backend_name,
146         .status = if (failures == 0) "pass" else "fail",
147         .mode = @tagName(gate.mode),
148         .cases = cases.all.len,
149         .failures = failures,
150         .unsupported = unsupported,
151         .invalid = invalid,
152         .assembled = assembled,
153         .total_compile_ns = total_compile_ns,
154         .total_launch_read_ns = total_launch_read_ns,
155     });
156 
157     if (failures != 0) {
158         try writer.flush();
159         sys.process.exit(1);
160     }
161 }
162 
163 fn requestedBackend() !BackendChoice {
164     return parseBackend(sys.env.get("ACCY_CONFORMANCE_BACKEND"));
165 }
166 
167 fn parseBackend(raw: ?[]const u8) !BackendChoice {
168     const value = raw orelse return .cuda;
169     if (std.mem.eql(u8, value, "cuda")) return .cuda;
170     if (std.mem.eql(u8, value, "metal")) return .metal;
171     if (std.mem.eql(u8, value, "vulkan")) return .vulkan;
172     return error.InvalidConformanceBackend;
173 }
174 
175 test "conformance backend selection defaults to cuda" {
176     try std.testing.expectEqual(BackendChoice.cuda, try parseBackend(null));
177     try std.testing.expectEqual(BackendChoice.cuda, try parseBackend("cuda"));
178     try std.testing.expectEqual(BackendChoice.metal, try parseBackend("metal"));
179     try std.testing.expectEqual(BackendChoice.vulkan, try parseBackend("vulkan"));
180     try std.testing.expectError(error.InvalidConformanceBackend, parseBackend("webgpu"));
181 }
182 
183 fn caseResult(
184     comptime Spec: type,
185     allocator: std.mem.Allocator,
186     handle: harness.BackendHandle,
187     gate: harness.Gate,
188     capabilities: gpu.BackendCapabilities,
189 ) records.Case {
190     if (comptime Spec.expectation != .invalid) {
191         const required_dtypes = comptime harness.requiredDTypes(Spec);
192         if (!capabilities.dtypes.containsAll(required_dtypes)) return unsupportedCase(Spec, "BackendDTypeUnsupported");
193         const required_features = comptime harness.requiredFeatures(Spec);
194         if (!capabilities.supportsFeatures(required_features)) return unsupportedCase(Spec, "BackendFeatureUnsupported");
195         const required_subgroup = comptime harness.requiredSubgroup(Spec);
196         if (!capabilities.supportsSubgroup(required_subgroup)) return unsupportedCase(Spec, "BackendSubgroupUnsupported");
197     }
198 
199     const run_fn = comptime if (@hasDecl(Spec, "buildArtifact")) harness.runFamilyCase else harness.runCase;
200     const case = run_fn(Spec, allocator, handle, gate) catch |err| {
201         const status: []const u8 = switch (Spec.expectation) {
202             .verified => "fail",
203             .unsupported => "unsupported",
204             .invalid => "invalid",
205         };
206         return .{
207             .name = Spec.name,
208             .status = status,
209             .detail = @errorName(err),
210             .kernels = 0,
211             .compile_ns = 0,
212             .launch_read_ns = 0,
213             .checksum = 0,
214             .max_abs_error = std.math.floatMax(f32),
215             .tolerance = Spec.tolerance,
216         };
217     };
218     if (Spec.expectation == .unsupported) {
219         var promoted = case;
220         promoted.status = "fail";
221         promoted.detail = "expected unsupported case compiled; promote it to verified";
222         return promoted;
223     }
224     if (Spec.expectation == .invalid) {
225         var promoted = case;
226         promoted.status = "fail";
227         promoted.detail = "expected semantically invalid case built a module; promote it";
228         return promoted;
229     }
230     return case;
231 }
232 
233 fn unsupportedCase(comptime Spec: type, detail: []const u8) records.Case {
234     return .{
235         .name = Spec.name,
236         .status = "unsupported",
237         .detail = detail,
238         .kernels = 0,
239         .compile_ns = 0,
240         .launch_read_ns = 0,
241         .checksum = 0,
242         .max_abs_error = std.math.floatMax(f32),
243         .tolerance = Spec.tolerance,
244     };
245 }
246 
247 fn forcedAssembleMode() bool {
248     const mode = sys.env.get("ACCY_CONFORMANCE_MODE") orelse return false;
249     return std.mem.eql(u8, mode, "assemble");
250 }
251 
252 fn writeJsonLine(writer: *std.Io.Writer, value: anytype) !void {
253     try pretty.json.writeMinified(writer, value);
254     try writer.writeByte('\n');
255 }