authorgravatar for robin@voetter.nlRobin Voetter <robin@voetter.nl> 2024-01-16 23:06:15+01:00
committergravatar for robin@voetter.nlRobin Voetter <robin@voetter.nl> 2024-02-04 19:09:26+01:00
log2f815853dcae49bbfd109675cde1f4097b75c8cc
treeef8790e6ef5dfc4a4ad40007a9c396dd0eb7eb80
parent15cf5f88c1bfb4b92285c515d294e1061d02672a
signaturebadge-check Signed by SSH key SHA256:ZS52FNyUv2WUXvO4njmVaFVO46RHojFuOrxRc4LuKzg

spirv: shlWithOverflow


3 files changed, 103 insertions(+), 32 deletions(-)

src/codegen/spirv.zig+103-25
......@@ -1782,19 +1782,6 @@ const DeclGen = struct {
17821782 wip.dg.gpa.free(wip.results);
17831783 }
17841784
1785 /// Return the scalar type of an input vector. This type is expected to be a vector
1786 /// if `wip.is_vector`, and a scalar otherwise.
1787 fn scalarType(wip: WipElementWise, ty: Type) Type {
1788 const mod = wip.dg.module;
1789 if (wip.is_vector) {
1790 assert(ty.isVector(mod));
1791 return ty.childType(mod);
1792 } else {
1793 assert(!ty.isVector(mod));
1794 return ty;
1795 }
1796 }
1797
17981785 /// Utility function to extract the element at a particular index in an
17991786 /// input vector. This type is expected to be a vector if `wip.is_vector`, and
18001787 /// a scalar otherwise.
......@@ -1844,7 +1831,7 @@ const DeclGen = struct {
18441831 const results = try self.gpa.alloc(IdRef, num_results);
18451832 for (results) |*result| result.* = undefined;
18461833
1847 const scalar_ty = if (is_vector) result_ty.childType(mod) else result_ty;
1834 const scalar_ty = result_ty.scalarType(mod);
18481835 const scalar_ty_ref = try self.resolveType(scalar_ty, .direct);
18491836
18501837 return .{
......@@ -2198,6 +2185,7 @@ const DeclGen = struct {
21982185
21992186 .add_with_overflow => try self.airAddSubOverflow(inst, .OpIAdd, .OpULessThan, .OpSLessThan),
22002187 .sub_with_overflow => try self.airAddSubOverflow(inst, .OpISub, .OpUGreaterThan, .OpSGreaterThan),
2188 .shl_with_overflow => try self.airShlOverflow(inst),
22012189
22022190 .shuffle => try self.airShuffle(inst),
22032191
......@@ -2343,23 +2331,30 @@ const DeclGen = struct {
23432331 const bin_op = self.air.instructions.items(.data)[@intFromEnum(inst)].bin_op;
23442332 const lhs_id = try self.resolve(bin_op.lhs);
23452333 const rhs_id = try self.resolve(bin_op.rhs);
2334
23462335 const result_ty = self.typeOfIndex(inst);
2336 const shift_ty = self.typeOf(bin_op.rhs);
2337 const scalar_shift_ty_ref = try self.resolveType(shift_ty.scalarType(mod), .direct);
23472338
2348 // Sometimes Zig doesn't make both of the arguments the same types here. SPIR-V expects that,
2349 // so just manually upcast it if required.
2350 // TODO(robin)
2339 const info = try self.arithmeticTypeInfo(result_ty);
2340 switch (info.class) {
2341 .composite_integer => return self.todo("shift ops for composite integers", .{}),
2342 .integer, .strange_integer => {},
2343 .float, .bool => unreachable,
2344 }
23512345
23522346 var wip = try self.elementWise(result_ty);
23532347 defer wip.deinit();
2354
2355 const shift_ty = wip.scalarType(self.typeOf(bin_op.rhs));
2356 const shift_ty_ref = try self.resolveType(shift_ty, .direct);
2357
23582348 for (0..wip.results.len) |i| {
23592349 const lhs_elem_id = try wip.elementAt(result_ty, lhs_id, i);
2360 const rhs_elem_id = try wip.elementAt(result_ty, rhs_id, i);
2350 const rhs_elem_id = try wip.elementAt(shift_ty, rhs_id, i);
2351
2352 // TODO: Can we omit normalizing lhs?
2353 const lhs_norm_id = try self.normalizeInt(wip.scalar_ty_ref, lhs_elem_id, info);
23612354
2362 const shift_id = if (shift_ty_ref != wip.result_ty_ref) blk: {
2355 // Sometimes Zig doesn't make both of the arguments the same types here. SPIR-V expects that,
2356 // so just manually upcast it if required.
2357 const shift_id = if (scalar_shift_ty_ref != wip.scalar_ty_ref) blk: {
23632358 const shift_id = self.spv.allocId();
23642359 try self.func.body.emit(self.spv.gpa, .OpUConvert, .{
23652360 .id_result_type = wip.scalar_ty_id,
......@@ -2368,12 +2363,13 @@ const DeclGen = struct {
23682363 });
23692364 break :blk shift_id;
23702365 } else rhs_elem_id;
2366 const shift_norm_id = try self.normalizeInt(wip.scalar_ty_ref, shift_id, info);
23712367
23722368 const args = .{
23732369 .id_result_type = wip.scalar_ty_id,
23742370 .id_result = wip.allocId(i),
2375 .base = lhs_elem_id,
2376 .shift = shift_id,
2371 .base = lhs_norm_id,
2372 .shift = shift_norm_id,
23772373 };
23782374
23792375 if (result_ty.isSignedInt(mod)) {
......@@ -2680,6 +2676,88 @@ const DeclGen = struct {
26802676 );
26812677 }
26822678
2679 fn airShlOverflow(self: *DeclGen, inst: Air.Inst.Index) !?IdRef {
2680 if (self.liveness.isUnused(inst)) return null;
2681 const mod = self.module;
2682 const ty_pl = self.air.instructions.items(.data)[@intFromEnum(inst)].ty_pl;
2683 const extra = self.air.extraData(Air.Bin, ty_pl.payload).data;
2684 const lhs = try self.resolve(extra.lhs);
2685 const rhs = try self.resolve(extra.rhs);
2686
2687 const result_ty = self.typeOfIndex(inst);
2688 const operand_ty = self.typeOf(extra.lhs);
2689 const shift_ty = self.typeOf(extra.rhs);
2690 const scalar_shift_ty_ref = try self.resolveType(shift_ty.scalarType(mod), .direct);
2691
2692 const ov_ty = result_ty.structFieldType(1, self.module);
2693
2694 const bool_ty_ref = try self.resolveType(Type.bool, .direct);
2695
2696 const info = try self.arithmeticTypeInfo(operand_ty);
2697 switch (info.class) {
2698 .composite_integer => return self.todo("overflow shift for composite integers", .{}),
2699 .integer, .strange_integer => {},
2700 .float, .bool => unreachable,
2701 }
2702
2703 var wip_result = try self.elementWise(operand_ty);
2704 defer wip_result.deinit();
2705 var wip_ov = try self.elementWise(ov_ty);
2706 defer wip_ov.deinit();
2707 for (0..wip_result.results.len, wip_ov.results) |i, *ov_id| {
2708 const lhs_elem_id = try wip_result.elementAt(operand_ty, lhs, i);
2709 const rhs_elem_id = try wip_result.elementAt(shift_ty, rhs, i);
2710
2711 // Normalize both so that we can shift back and check if the result is the same.
2712 const lhs_norm_id = try self.normalizeInt(wip_result.scalar_ty_ref, lhs_elem_id, info);
2713
2714 // Sometimes Zig doesn't make both of the arguments the same types here. SPIR-V expects that,
2715 // so just manually upcast it if required.
2716 const shift_id = if (scalar_shift_ty_ref != wip_result.scalar_ty_ref) blk: {
2717 const shift_id = self.spv.allocId();
2718 try self.func.body.emit(self.spv.gpa, .OpUConvert, .{
2719 .id_result_type = wip_result.scalar_ty_id,
2720 .id_result = shift_id,
2721 .unsigned_value = rhs_elem_id,
2722 });
2723 break :blk shift_id;
2724 } else rhs_elem_id;
2725 const shift_norm_id = try self.normalizeInt(wip_result.scalar_ty_ref, shift_id, info);
2726
2727 try self.func.body.emit(self.spv.gpa, .OpShiftLeftLogical, .{
2728 .id_result_type = wip_result.scalar_ty_id,
2729 .id_result = wip_result.allocId(i),
2730 .base = lhs_norm_id,
2731 .shift = shift_norm_id,
2732 });
2733
2734 // To check if overflow happened, just check if the right-shifted result is the same value.
2735 const right_shift_id = self.spv.allocId();
2736 try self.func.body.emit(self.spv.gpa, .OpShiftRightLogical, .{
2737 .id_result_type = wip_result.scalar_ty_id,
2738 .id_result = right_shift_id,
2739 .base = try self.normalizeInt(wip_result.scalar_ty_ref, wip_result.results[i], info),
2740 .shift = shift_norm_id,
2741 });
2742
2743 const overflowed_id = self.spv.allocId();
2744 try self.func.body.emit(self.spv.gpa, .OpINotEqual, .{
2745 .id_result_type = self.typeId(bool_ty_ref),
2746 .id_result = overflowed_id,
2747 .operand_1 = lhs_norm_id,
2748 .operand_2 = right_shift_id,
2749 });
2750
2751 ov_id.* = try self.intFromBool(wip_ov.scalar_ty_ref, overflowed_id);
2752 }
2753
2754 return try self.constructStruct(
2755 result_ty,
2756 &.{ operand_ty, ov_ty },
2757 &.{ try wip_result.finalize(), try wip_ov.finalize() },
2758 );
2759 }
2760
26832761 fn airShuffle(self: *DeclGen, inst: Air.Inst.Index) !?IdRef {
26842762 const mod = self.module;
26852763 if (self.liveness.isUnused(inst)) return null;
test/behavior/math.zig-2
......@@ -1328,8 +1328,6 @@ fn testShlTrunc(x: u16) !void {
13281328}
13291329
13301330test "exact shift left" {
1331 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
1332
13331331 try testShlExact(0b00110101);
13341332 try comptime testShlExact(0b00110101);
13351333
test/behavior/vector.zig-5
......@@ -179,7 +179,6 @@ test "array vector coercion - odd sizes" {
179179 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
180180 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
181181 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
182 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
183182 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest;
184183 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest;
185184
......@@ -219,7 +218,6 @@ test "array to vector with element type coercion" {
219218 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
220219 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
221220 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
222 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
223221 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest; // TODO
224222 if (builtin.zig_backend == .stage2_x86_64 and builtin.target.ofmt != .elf) return error.SkipZigTest;
225223
......@@ -659,7 +657,6 @@ test "vector shift operators" {
659657 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
660658 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
661659 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
662 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
663660
664661 const S = struct {
665662 fn doTheTestShift(x: anytype, y: anytype) !void {
......@@ -1168,7 +1165,6 @@ test "@shlWithOverflow" {
11681165 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
11691166 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
11701167 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
1171 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
11721168
11731169 const S = struct {
11741170 fn doTheTest() !void {
......@@ -1453,7 +1449,6 @@ test "compare vectors with different element types" {
14531449 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest; // TODO
14541450 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest; // TODO
14551451 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
1456 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest; // TODO
14571452
14581453 var a: @Vector(2, u8) = .{ 1, 2 };
14591454 var b: @Vector(2, u9) = .{ 3, 0 };