authorgravatar for robin@voetter.nlRobin Voetter <robin@voetter.nl> 2024-01-21 22:24:53+01:00
committergravatar for robin@voetter.nlRobin Voetter <robin@voetter.nl> 2024-02-04 19:09:32+01:00
log76d5696434095e39d9aaae92c1533b2d016c1a31
treeca84126ca1b3c5256dc7828adbf351cef694faea
parent631d1b63a8027c49073995e28aab489534f01efa
signaturebadge-check Signed by SSH key SHA256:ZS52FNyUv2WUXvO4njmVaFVO46RHojFuOrxRc4LuKzg

spirv: air abs


5 files changed, 73 insertions(+), 10 deletions(-)

src/codegen/spirv.zig+71
...@@ -678,6 +678,18 @@ const DeclGen = struct {...@@ -678,6 +678,18 @@ const DeclGen = struct {
678 }678 }
679 }679 }
680680
681 /// Emits a float constant
682 fn constFloat(self: *DeclGen, ty_ref: CacheRef, value: f128) !IdRef {
683 const ty = self.spv.cache.lookup(ty_ref).float_type;
684 return switch (ty.bits) {
685 16 => try self.spv.resolveId(.{ .float = .{ .ty = ty_ref, .value = .{ .float16 = @floatCast(value) } } }),
686 32 => try self.spv.resolveId(.{ .float = .{ .ty = ty_ref, .value = .{ .float32 = @floatCast(value) } } }),
687 64 => try self.spv.resolveId(.{ .float = .{ .ty = ty_ref, .value = .{ .float64 = @floatCast(value) } } }),
688 80, 128 => unreachable, // TODO
689 else => unreachable,
690 };
691 }
692
681 /// Construct a struct at runtime.693 /// Construct a struct at runtime.
682 /// ty must be a struct type.694 /// ty must be a struct type.
683 /// Constituents should be in `indirect` representation (as the elements of a struct should be).695 /// Constituents should be in `indirect` representation (as the elements of a struct should be).
...@@ -2164,6 +2176,8 @@ const DeclGen = struct {...@@ -2164,6 +2176,8 @@ const DeclGen = struct {
2164 .sub, .sub_wrap, .sub_optimized => try self.airArithOp(inst, .OpFSub, .OpISub, .OpISub),2176 .sub, .sub_wrap, .sub_optimized => try self.airArithOp(inst, .OpFSub, .OpISub, .OpISub),
2165 .mul, .mul_wrap, .mul_optimized => try self.airArithOp(inst, .OpFMul, .OpIMul, .OpIMul),2177 .mul, .mul_wrap, .mul_optimized => try self.airArithOp(inst, .OpFMul, .OpIMul, .OpIMul),
21662178
2179 .abs => try self.airAbs(inst),
2180
2167 .div_float,2181 .div_float,
2168 .div_float_optimized,2182 .div_float_optimized,
2169 // TODO: Check that this is the right operation.2183 // TODO: Check that this is the right operation.
...@@ -2562,6 +2576,63 @@ const DeclGen = struct {...@@ -2562,6 +2576,63 @@ const DeclGen = struct {
2562 return try wip.finalize();2576 return try wip.finalize();
2563 }2577 }
25642578
2579 fn airAbs(self: *DeclGen, inst: Air.Inst.Index) !?IdRef {
2580 if (self.liveness.isUnused(inst)) return null;
2581
2582 const mod = self.module;
2583 const ty_op = self.air.instructions.items(.data)[@intFromEnum(inst)].ty_op;
2584 const operand_id = try self.resolve(ty_op.operand);
2585 // Note: operand_ty may be signed, while ty is always unsigned!
2586 const operand_ty = self.typeOf(ty_op.operand);
2587 const ty = self.typeOfIndex(inst);
2588 const info = self.arithmeticTypeInfo(ty);
2589 const operand_scalar_ty = operand_ty.scalarType(mod);
2590 const operand_scalar_ty_ref = try self.resolveType(operand_scalar_ty, .direct);
2591
2592 var wip = try self.elementWise(ty);
2593 defer wip.deinit();
2594
2595 const zero_id = switch (info.class) {
2596 .float => try self.constFloat(operand_scalar_ty_ref, 0),
2597 .integer, .strange_integer => try self.constInt(operand_scalar_ty_ref, 0),
2598 .composite_integer => unreachable, // TODO
2599 .bool => unreachable,
2600 };
2601 for (wip.results, 0..) |*result_id, i| {
2602 const elem_id = try wip.elementAt(operand_ty, operand_id, i);
2603 // Idk why spir-v doesn't have a dedicated abs() instruction in the base
2604 // instruction set. For now we're just going to negate and check to avoid
2605 // importing the extinst.
2606 const neg_id = self.spv.allocId();
2607 const args = .{
2608 .id_result_type = self.typeId(operand_scalar_ty_ref),
2609 .id_result = neg_id,
2610 .operand_1 = zero_id,
2611 .operand_2 = elem_id,
2612 };
2613 switch (info.class) {
2614 .float => try self.func.body.emit(self.spv.gpa, .OpFSub, args),
2615 .integer, .strange_integer => try self.func.body.emit(self.spv.gpa, .OpISub, args),
2616 .composite_integer => unreachable, // TODO
2617 .bool => unreachable,
2618 }
2619 const neg_norm_id = try self.normalize(wip.scalar_ty_ref, neg_id, info);
2620
2621 const gt_zero_id = try self.cmp(.gt, Type.bool, operand_scalar_ty, elem_id, zero_id);
2622 const abs_id = self.spv.allocId();
2623 try self.func.body.emit(self.spv.gpa, .OpSelect, .{
2624 .id_result_type = self.typeId(operand_scalar_ty_ref),
2625 .id_result = abs_id,
2626 .condition = gt_zero_id,
2627 .object_1 = elem_id,
2628 .object_2 = neg_norm_id,
2629 });
2630 // For Shader, we may need to cast from signed to unsigned here.
2631 result_id.* = try self.bitCast(wip.scalar_ty, operand_scalar_ty, abs_id);
2632 }
2633 return try wip.finalize();
2634 }
2635
2565 fn airAddSubOverflow(2636 fn airAddSubOverflow(
2566 self: *DeclGen,2637 self: *DeclGen,
2567 inst: Air.Inst.Index,2638 inst: Air.Inst.Index,
test/behavior/abs.zig+2-5
...@@ -7,7 +7,6 @@ test "@abs integers" {...@@ -7,7 +7,6 @@ test "@abs integers" {
7 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO7 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
8 if (builtin.zig_backend == .stage2_riscv64) return error.SkipZigTest; // TODO8 if (builtin.zig_backend == .stage2_riscv64) return error.SkipZigTest; // TODO
9 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO9 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
10 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
1110
12 try comptime testAbsIntegers();11 try comptime testAbsIntegers();
13 try testAbsIntegers();12 try testAbsIntegers();
...@@ -95,7 +94,6 @@ test "@abs floats" {...@@ -95,7 +94,6 @@ test "@abs floats" {
95 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO94 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
96 if (builtin.zig_backend == .stage2_riscv64) return error.SkipZigTest; // TODO95 if (builtin.zig_backend == .stage2_riscv64) return error.SkipZigTest; // TODO
97 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO96 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
98 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
99 if (builtin.zig_backend == .stage2_x86_64 and builtin.target.ofmt != .elf) return error.SkipZigTest;97 if (builtin.zig_backend == .stage2_x86_64 and builtin.target.ofmt != .elf) return error.SkipZigTest;
10098
101 try comptime testAbsFloats(f16);99 try comptime testAbsFloats(f16);
...@@ -105,9 +103,9 @@ test "@abs floats" {...@@ -105,9 +103,9 @@ test "@abs floats" {
105 try comptime testAbsFloats(f64);103 try comptime testAbsFloats(f64);
106 try testAbsFloats(f64);104 try testAbsFloats(f64);
107 try comptime testAbsFloats(f80);105 try comptime testAbsFloats(f80);
108 if (builtin.zig_backend != .stage2_wasm) try testAbsFloats(f80);106 if (builtin.zig_backend != .stage2_wasm and builtin.zig_backend != .stage2_spirv64) try testAbsFloats(f80);
109 try comptime testAbsFloats(f128);107 try comptime testAbsFloats(f128);
110 if (builtin.zig_backend != .stage2_wasm) try testAbsFloats(f128);108 if (builtin.zig_backend != .stage2_wasm and builtin.zig_backend != .stage2_spirv64) try testAbsFloats(f128);
111}109}
112110
113fn testAbsFloats(comptime T: type) !void {111fn testAbsFloats(comptime T: type) !void {
...@@ -155,7 +153,6 @@ test "@abs int vectors" {...@@ -155,7 +153,6 @@ test "@abs int vectors" {
155 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO153 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
156 if (builtin.zig_backend == .stage2_riscv64) return error.SkipZigTest; // TODO154 if (builtin.zig_backend == .stage2_riscv64) return error.SkipZigTest; // TODO
157 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO155 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
158 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
159156
160 try comptime testAbsIntVectors(1);157 try comptime testAbsIntVectors(1);
161 try testAbsIntVectors(1);158 try testAbsIntVectors(1);
test/behavior/cast.zig-1
...@@ -2465,7 +2465,6 @@ test "@as does not corrupt values with incompatible representations" {...@@ -2465,7 +2465,6 @@ test "@as does not corrupt values with incompatible representations" {
2465 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO2465 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
2466 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO2466 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
2467 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO2467 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
2468 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
2469 if (builtin.zig_backend == .stage2_x86_64 and builtin.target.ofmt != .elf) return error.SkipZigTest;2468 if (builtin.zig_backend == .stage2_x86_64 and builtin.target.ofmt != .elf) return error.SkipZigTest;
24702469
2471 const x: f32 = @as(f16, blk: {2470 const x: f32 = @as(f16, blk: {
test/behavior/floatop.zig-3
...@@ -969,7 +969,6 @@ test "@abs f16" {...@@ -969,7 +969,6 @@ test "@abs f16" {
969 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO969 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
970 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO970 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
971 if (builtin.zig_backend == .stage2_x86_64 and builtin.target.ofmt != .elf) return error.SkipZigTest;971 if (builtin.zig_backend == .stage2_x86_64 and builtin.target.ofmt != .elf) return error.SkipZigTest;
972 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
973972
974 try testFabs(f16);973 try testFabs(f16);
975 try comptime testFabs(f16);974 try comptime testFabs(f16);
...@@ -979,7 +978,6 @@ test "@abs f32/f64" {...@@ -979,7 +978,6 @@ test "@abs f32/f64" {
979 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO978 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
980 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO979 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
981 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO980 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
982 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
983981
984 try testFabs(f32);982 try testFabs(f32);
985 try comptime testFabs(f32);983 try comptime testFabs(f32);
...@@ -1070,7 +1068,6 @@ fn testFabs(comptime T: type) !void {...@@ -1070,7 +1068,6 @@ fn testFabs(comptime T: type) !void {
1070test "@abs with vectors" {1068test "@abs with vectors" {
1071 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO1069 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
1072 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO1070 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
1073 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
1074 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest; // TODO1071 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest; // TODO
10751072
1076 try testFabsWithVectors();1073 try testFabsWithVectors();
test/behavior/math.zig-1
...@@ -1687,7 +1687,6 @@ test "absFloat" {...@@ -1687,7 +1687,6 @@ test "absFloat" {
1687 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO1687 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
1688 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO1688 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
1689 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO1689 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
1690 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
16911690
1692 try testAbsFloat();1691 try testAbsFloat();
1693 try comptime testAbsFloat();1692 try comptime testAbsFloat();