authorgravatar for kkhaike@gmail.comkkHAIKE <kkhaike@gmail.com> 2022-11-19 04:06:49+08:00
committergravatar for noreply@github.comGitHub <noreply@github.com> 2022-11-18 22:06:49+02:00
logea590ece4bcc26175d890dd82899ffc9ef64b6c3
treeb5d77cc7bbe683e54c9bf2b7d29b4681dc057ba5
parent04f3067a7921a51bd55870f3138220fd66426e33
signaturebadge-question-mark Signed by PGP key 4AEE18F83AFDEB23

Sema: optimize compare comptime float with int


3 files changed, 137 insertions(+), 15 deletions(-)

src/Sema.zig+46-10
......@@ -13896,9 +13896,25 @@ fn analyzeArithmetic(
1389613896 // because there is a possible value for which the addition would
1389713897 // overflow (max_int), causing illegal behavior.
1389813898 // For floats: either operand being undef makes the result undef.
13899 // If either of the operands are inf, and the other operand is zero,
13900 // the result is nan.
13901 // If either of the operands are nan, the result is nan.
1389913902 if (maybe_lhs_val) |lhs_val| {
1390013903 if (!lhs_val.isUndef()) {
13901 if (try lhs_val.compareAllWithZeroAdvanced(.eq, sema)) {
13904 if (lhs_val.isNan()) {
13905 return sema.addConstant(resolved_type, lhs_val);
13906 }
13907 if (try lhs_val.compareAllWithZeroAdvanced(.eq, sema)) lz: {
13908 if (maybe_rhs_val) |rhs_val| {
13909 if (rhs_val.isNan()) {
13910 return sema.addConstant(resolved_type, rhs_val);
13911 }
13912 if (rhs_val.isInf()) {
13913 return sema.addConstant(resolved_type, try Value.Tag.float_32.create(sema.arena, std.math.nan_f32));
13914 }
13915 } else if (resolved_type.isAnyFloat()) {
13916 break :lz;
13917 }
1390213918 const zero_val = if (is_vector) b: {
1390313919 break :b try Value.Tag.repeated.create(sema.arena, Value.zero);
1390413920 } else Value.zero;
......@@ -13918,7 +13934,17 @@ fn analyzeArithmetic(
1391813934 return sema.addConstUndef(resolved_type);
1391913935 }
1392013936 }
13921 if (try rhs_val.compareAllWithZeroAdvanced(.eq, sema)) {
13937 if (rhs_val.isNan()) {
13938 return sema.addConstant(resolved_type, rhs_val);
13939 }
13940 if (try rhs_val.compareAllWithZeroAdvanced(.eq, sema)) rz: {
13941 if (maybe_lhs_val) |lhs_val| {
13942 if (lhs_val.isInf()) {
13943 return sema.addConstant(resolved_type, try Value.Tag.float_32.create(sema.arena, std.math.nan_f32));
13944 }
13945 } else if (resolved_type.isAnyFloat()) {
13946 break :rz;
13947 }
1392213948 const zero_val = if (is_vector) b: {
1392313949 break :b try Value.Tag.repeated.create(sema.arena, Value.zero);
1392413950 } else Value.zero;
......@@ -28175,8 +28201,10 @@ fn cmpNumeric(
2817528201 else => return Air.Inst.Ref.bool_false,
2817628202 };
2817728203 if (lhs_val.isInf()) switch (op) {
28178 .gt, .neq => return Air.Inst.Ref.bool_true,
28179 .lt, .lte, .eq, .gte => return Air.Inst.Ref.bool_false,
28204 .neq => return Air.Inst.Ref.bool_true,
28205 .eq => return Air.Inst.Ref.bool_false,
28206 .gt, .gte => return if (lhs_val.isNegativeInf()) Air.Inst.Ref.bool_false else Air.Inst.Ref.bool_true,
28207 .lt, .lte => return if (lhs_val.isNegativeInf()) Air.Inst.Ref.bool_true else Air.Inst.Ref.bool_false,
2818028208 };
2818128209 if (!rhs_is_signed) {
2818228210 switch (lhs_val.orderAgainstZero()) {
......@@ -28193,14 +28221,17 @@ fn cmpNumeric(
2819328221 }
2819428222 }
2819528223 if (lhs_is_float) {
28196 var bigint = try float128IntPartToBigInt(sema.gpa, lhs_val.toFloat(f128));
28197 defer bigint.deinit();
2819828224 if (lhs_val.floatHasFraction()) {
2819928225 switch (op) {
2820028226 .eq => return Air.Inst.Ref.bool_false,
2820128227 .neq => return Air.Inst.Ref.bool_true,
2820228228 else => {},
2820328229 }
28230 }
28231
28232 var bigint = try float128IntPartToBigInt(sema.gpa, lhs_val.toFloat(f128));
28233 defer bigint.deinit();
28234 if (lhs_val.floatHasFraction()) {
2820428235 if (lhs_is_signed) {
2820528236 try bigint.addScalar(&bigint, -1);
2820628237 } else {
......@@ -28228,8 +28259,10 @@ fn cmpNumeric(
2822828259 else => return Air.Inst.Ref.bool_false,
2822928260 };
2823028261 if (rhs_val.isInf()) switch (op) {
28231 .lt, .neq => return Air.Inst.Ref.bool_true,
28232 .gt, .lte, .eq, .gte => return Air.Inst.Ref.bool_false,
28262 .neq => return Air.Inst.Ref.bool_true,
28263 .eq => return Air.Inst.Ref.bool_false,
28264 .gt, .gte => return if (rhs_val.isNegativeInf()) Air.Inst.Ref.bool_true else Air.Inst.Ref.bool_false,
28265 .lt, .lte => return if (rhs_val.isNegativeInf()) Air.Inst.Ref.bool_false else Air.Inst.Ref.bool_true,
2823328266 };
2823428267 if (!lhs_is_signed) {
2823528268 switch (rhs_val.orderAgainstZero()) {
......@@ -28246,14 +28279,17 @@ fn cmpNumeric(
2824628279 }
2824728280 }
2824828281 if (rhs_is_float) {
28249 var bigint = try float128IntPartToBigInt(sema.gpa, rhs_val.toFloat(f128));
28250 defer bigint.deinit();
2825128282 if (rhs_val.floatHasFraction()) {
2825228283 switch (op) {
2825328284 .eq => return Air.Inst.Ref.bool_false,
2825428285 .neq => return Air.Inst.Ref.bool_true,
2825528286 else => {},
2825628287 }
28288 }
28289
28290 var bigint = try float128IntPartToBigInt(sema.gpa, rhs_val.toFloat(f128));
28291 defer bigint.deinit();
28292 if (rhs_val.floatHasFraction()) {
2825728293 if (rhs_is_signed) {
2825828294 try bigint.addScalar(&bigint, -1);
2825928295 } else {
src/value.zig+25-5
......@@ -2081,6 +2081,15 @@ pub const Value = extern union {
20812081 op: std.math.CompareOperator,
20822082 opt_sema: ?*Sema,
20832083 ) Module.CompileError!bool {
2084 if (lhs.isInf()) {
2085 switch (op) {
2086 .neq => return true,
2087 .eq => return false,
2088 .gt, .gte => return !lhs.isNegativeInf(),
2089 .lt, .lte => return lhs.isNegativeInf(),
2090 }
2091 }
2092
20842093 switch (lhs.tag()) {
20852094 .repeated => return lhs.castTag(.repeated).?.data.compareAllWithZeroAdvanced(op, opt_sema),
20862095 .aggregate => {
......@@ -2089,11 +2098,11 @@ pub const Value = extern union {
20892098 }
20902099 return true;
20912100 },
2092 .float_16 => if (std.math.isNan(lhs.castTag(.float_16).?.data)) return op != .neq,
2093 .float_32 => if (std.math.isNan(lhs.castTag(.float_32).?.data)) return op != .neq,
2094 .float_64 => if (std.math.isNan(lhs.castTag(.float_64).?.data)) return op != .neq,
2095 .float_80 => if (std.math.isNan(lhs.castTag(.float_80).?.data)) return op != .neq,
2096 .float_128 => if (std.math.isNan(lhs.castTag(.float_128).?.data)) return op != .neq,
2101 .float_16 => if (std.math.isNan(lhs.castTag(.float_16).?.data)) return op == .neq,
2102 .float_32 => if (std.math.isNan(lhs.castTag(.float_32).?.data)) return op == .neq,
2103 .float_64 => if (std.math.isNan(lhs.castTag(.float_64).?.data)) return op == .neq,
2104 .float_80 => if (std.math.isNan(lhs.castTag(.float_80).?.data)) return op == .neq,
2105 .float_128 => if (std.math.isNan(lhs.castTag(.float_128).?.data)) return op == .neq,
20972106 else => {},
20982107 }
20992108 return (try orderAgainstZeroAdvanced(lhs, opt_sema)).compare(op);
......@@ -3817,6 +3826,17 @@ pub const Value = extern union {
38173826 };
38183827 }
38193828
3829 pub fn isNegativeInf(val: Value) bool {
3830 return switch (val.tag()) {
3831 .float_16 => std.math.isNegativeInf(val.castTag(.float_16).?.data),
3832 .float_32 => std.math.isNegativeInf(val.castTag(.float_32).?.data),
3833 .float_64 => std.math.isNegativeInf(val.castTag(.float_64).?.data),
3834 .float_80 => std.math.isNegativeInf(val.castTag(.float_80).?.data),
3835 .float_128 => std.math.isNegativeInf(val.castTag(.float_128).?.data),
3836 else => false,
3837 };
3838 }
3839
38203840 pub fn floatRem(lhs: Value, rhs: Value, float_type: Type, arena: Allocator, target: Target) !Value {
38213841 if (float_type.zigTypeTag() == .Vector) {
38223842 const result_data = try arena.alloc(Value, float_type.vectorLen());
test/behavior/bugs/12891.zig+66
......@@ -18,3 +18,69 @@ test "inf" {
1818 var i: usize = 0;
1919 try std.testing.expect(f > i);
2020}
21test "-inf < 0" {
22 const f = comptime -std.math.inf(f64);
23 var i: usize = 0;
24 try std.testing.expect(f < i);
25}
26test "inf >= 1" {
27 const f = comptime std.math.inf(f64);
28 var i: usize = 1;
29 try std.testing.expect(f >= i);
30}
31test "isNan(nan * 1)" {
32 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
33 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
34 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest; // TODO
35
36 const nan_times_one = comptime std.math.nan(f64) * 1;
37 try std.testing.expect(std.math.isNan(nan_times_one));
38}
39test "runtime isNan(nan * 1)" {
40 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
41 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
42 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest; // TODO
43
44 const nan_times_one = std.math.nan(f64) * 1;
45 try std.testing.expect(std.math.isNan(nan_times_one));
46}
47test "isNan(nan * 0)" {
48 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
49 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
50 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest; // TODO
51
52 const nan_times_zero = comptime std.math.nan(f64) * 0;
53 try std.testing.expect(std.math.isNan(nan_times_zero));
54 const zero_times_nan = 0 * comptime std.math.nan(f64);
55 try std.testing.expect(std.math.isNan(zero_times_nan));
56}
57test "isNan(inf * 0)" {
58 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
59 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
60 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest; // TODO
61
62 const inf_times_zero = comptime std.math.inf(f64) * 0;
63 try std.testing.expect(std.math.isNan(inf_times_zero));
64 const zero_times_inf = 0 * comptime std.math.inf(f64);
65 try std.testing.expect(std.math.isNan(zero_times_inf));
66}
67test "runtime isNan(nan * 0)" {
68 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
69 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
70 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest; // TODO
71
72 const nan_times_zero = std.math.nan(f64) * 0;
73 try std.testing.expect(std.math.isNan(nan_times_zero));
74 const zero_times_nan = 0 * std.math.nan(f64);
75 try std.testing.expect(std.math.isNan(zero_times_nan));
76}
77test "runtime isNan(inf * 0)" {
78 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
79 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
80 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest; // TODO
81
82 const inf_times_zero = std.math.inf(f64) * 0;
83 try std.testing.expect(std.math.isNan(inf_times_zero));
84 const zero_times_inf = 0 * std.math.inf(f64);
85 try std.testing.expect(std.math.isNan(zero_times_inf));
86}