authorgravatar for robin@voetter.nlRobin Voetter <robin@voetter.nl> 2023-10-08 13:02:16+02:00
committergravatar for robin@voetter.nlRobin Voetter <robin@voetter.nl> 2023-10-15 14:00:07+02:00
log4f279078c8bd57e7ee47eebce2572ab25729850a
tree01b17bf5eead8d87b22b1e9d5b33e61a2f5a7cf8
parentf858bf161602d72584da9e950c5a9eeadfe8b29d
signature Signed by SSH key SHA256:CQ99aPxq+RueiL9u7z0FEki5Fm7V6T8q4PrEGmINrA4

spirv: air min/max


2 files changed, 57 insertions(+), 2 deletions(-)

src/codegen/spirv.zig+51
...@@ -1970,6 +1970,9 @@ const DeclGen = struct {...@@ -1970,6 +1970,9 @@ const DeclGen = struct {
19701970
1971 .shl => try self.airShift(inst, .OpShiftLeftLogical),1971 .shl => try self.airShift(inst, .OpShiftLeftLogical),
19721972
1973 .min => try self.airMinMax(inst, .lt),
1974 .max => try self.airMinMax(inst, .gt),
1975
1973 .bitcast => try self.airBitCast(inst),1976 .bitcast => try self.airBitCast(inst),
1974 .intcast, .trunc => try self.airIntCast(inst),1977 .intcast, .trunc => try self.airIntCast(inst),
1975 .int_from_ptr => try self.airIntFromPtr(inst),1978 .int_from_ptr => try self.airIntFromPtr(inst),
...@@ -2103,6 +2106,54 @@ const DeclGen = struct {...@@ -2103,6 +2106,54 @@ const DeclGen = struct {
2103 return result_id;2106 return result_id;
2104 }2107 }
21052108
2109 fn airMinMax(self: *DeclGen, inst: Air.Inst.Index, op: std.math.CompareOperator) !?IdRef {
2110 if (self.liveness.isUnused(inst)) return null;
2111
2112 const bin_op = self.air.instructions.items(.data)[inst].bin_op;
2113 const lhs_id = try self.resolve(bin_op.lhs);
2114 const rhs_id = try self.resolve(bin_op.rhs);
2115 const result_ty = self.typeOfIndex(inst);
2116 const result_ty_ref = try self.resolveType(result_ty, .direct);
2117
2118 const info = try self.arithmeticTypeInfo(result_ty);
2119 // TODO: Use fmin for OpenCL
2120 const cmp_id = try self.cmp(op, result_ty, lhs_id, rhs_id);
2121 const selection_id = switch (info.class) {
2122 .float => blk: {
2123 // cmp uses OpFOrd. When we have 0 [<>] nan this returns false,
2124 // but we want it to pick lhs. Therefore we also have to check if
2125 // rhs is nan. We don't need to care about the result when both
2126 // are nan.
2127 const rhs_is_nan_id = self.spv.allocId();
2128 const bool_ty_ref = try self.resolveType(Type.bool, .direct);
2129 try self.func.body.emit(self.spv.gpa, .OpIsNan, .{
2130 .id_result_type = self.typeId(bool_ty_ref),
2131 .id_result = rhs_is_nan_id,
2132 .x = rhs_id,
2133 });
2134 const float_cmp_id = self.spv.allocId();
2135 try self.func.body.emit(self.spv.gpa, .OpLogicalOr, .{
2136 .id_result_type = self.typeId(bool_ty_ref),
2137 .id_result = float_cmp_id,
2138 .operand_1 = cmp_id,
2139 .operand_2 = rhs_is_nan_id,
2140 });
2141 break :blk float_cmp_id;
2142 },
2143 else => cmp_id,
2144 };
2145
2146 const result_id = self.spv.allocId();
2147 try self.func.body.emit(self.spv.gpa, .OpSelect, .{
2148 .id_result_type = self.typeId(result_ty_ref),
2149 .id_result = result_id,
2150 .condition = selection_id,
2151 .object_1 = lhs_id,
2152 .object_2 = rhs_id,
2153 });
2154 return result_id;
2155 }
2156
2106 fn maskStrangeInt(self: *DeclGen, ty_ref: CacheRef, value_id: IdRef, bits: u16) !IdRef {2157 fn maskStrangeInt(self: *DeclGen, ty_ref: CacheRef, value_id: IdRef, bits: u16) !IdRef {
2107 const mask_value = if (bits == 64) 0xFFFF_FFFF_FFFF_FFFF else (@as(u64, 1) << @as(u6, @intCast(bits))) - 1;2158 const mask_value = if (bits == 64) 0xFFFF_FFFF_FFFF_FFFF else (@as(u64, 1) << @as(u6, @intCast(bits))) - 1;
2108 const result_id = self.spv.allocId();2159 const result_id = self.spv.allocId();
test/behavior/maximum_minimum.zig+6-2
...@@ -9,14 +9,16 @@ test "@max" {...@@ -9,14 +9,16 @@ test "@max" {
9 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO9 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
10 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO10 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
11 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO11 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
12 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
1312
14 const S = struct {13 const S = struct {
15 fn doTheTest() !void {14 fn doTheTest() !void {
16 var x: i32 = 10;15 var x: i32 = 10;
17 var y: f32 = 0.68;16 var y: f32 = 0.68;
17 var nan: f32 = std.math.nan(f32);
18 try expect(@as(i32, 10) == @max(@as(i32, -3), x));18 try expect(@as(i32, 10) == @max(@as(i32, -3), x));
19 try expect(@as(f32, 3.2) == @max(@as(f32, 3.2), y));19 try expect(@as(f32, 3.2) == @max(@as(f32, 3.2), y));
20 try expect(y == @max(nan, y));
21 try expect(y == @max(y, nan));
20 }22 }
21 };23 };
22 try S.doTheTest();24 try S.doTheTest();
...@@ -58,14 +60,16 @@ test "@min" {...@@ -58,14 +60,16 @@ test "@min" {
58 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO60 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
59 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO61 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
60 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO62 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
61 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
6263
63 const S = struct {64 const S = struct {
64 fn doTheTest() !void {65 fn doTheTest() !void {
65 var x: i32 = 10;66 var x: i32 = 10;
66 var y: f32 = 0.68;67 var y: f32 = 0.68;
68 var nan: f32 = std.math.nan(f32);
67 try expect(@as(i32, -3) == @min(@as(i32, -3), x));69 try expect(@as(i32, -3) == @min(@as(i32, -3), x));
68 try expect(@as(f32, 0.68) == @min(@as(f32, 3.2), y));70 try expect(@as(f32, 0.68) == @min(@as(f32, 3.2), y));
71 try expect(y == @min(nan, y));
72 try expect(y == @min(y, nan));
69 }73 }
70 };74 };
71 try S.doTheTest();75 try S.doTheTest();