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 }