Skip to content

Commit 5b8aacb

Browse files
committed
feat(jit): refine integer ranges from branches (#56)
1 parent 1fc7e9c commit 5b8aacb

1 file changed

Lines changed: 173 additions & 23 deletions

File tree

‎src/jit/compiler.zig‎

Lines changed: 173 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -72,10 +72,24 @@ const IntegerRange = struct {
7272
}
7373
};
7474

75+
const RangeOperand = struct {
76+
integer: IntegerFact,
77+
range: IntegerRange,
78+
local: ?u8,
79+
};
80+
81+
const Condition = struct {
82+
op: bc.Op,
83+
lhs: RangeOperand,
84+
rhs: RangeOperand,
85+
};
86+
7587
const StackValue = struct {
7688
kind: Kind,
7789
integer: IntegerFact,
7890
range: IntegerRange,
91+
local: ?u8 = null,
92+
condition: ?Condition = null,
7993
};
8094

8195
const State = struct {
@@ -85,13 +99,17 @@ const State = struct {
8599
stack: [jit.numeric_scratch_capacity]Kind = @splat(.undefined),
86100
stack_integers: [jit.numeric_scratch_capacity]IntegerFact = @splat(.unknown),
87101
stack_ranges: [jit.numeric_scratch_capacity]IntegerRange = @splat(.unknown),
102+
stack_locals: [jit.numeric_scratch_capacity]?u8 = @splat(null),
103+
stack_conditions: [jit.numeric_scratch_capacity]?Condition = @splat(null),
88104
depth: u8 = 0,
89105

90-
fn push(self: *State, kind: Kind, integer: IntegerFact, range: IntegerRange) error{UnsupportedChunk}!void {
106+
fn push(self: *State, value_state: StackValue) error{UnsupportedChunk}!void {
91107
if (self.depth == jit.numeric_scratch_capacity) return error.UnsupportedChunk;
92-
self.stack[self.depth] = kind;
93-
self.stack_integers[self.depth] = integer;
94-
self.stack_ranges[self.depth] = range;
108+
self.stack[self.depth] = value_state.kind;
109+
self.stack_integers[self.depth] = value_state.integer;
110+
self.stack_ranges[self.depth] = value_state.range;
111+
self.stack_locals[self.depth] = value_state.local;
112+
self.stack_conditions[self.depth] = value_state.condition;
95113
self.depth += 1;
96114
}
97115

@@ -101,10 +119,14 @@ const State = struct {
101119
const kind = self.stack[self.depth];
102120
const integer = self.stack_integers[self.depth];
103121
const range = self.stack_ranges[self.depth];
122+
const local = self.stack_locals[self.depth];
123+
const condition = self.stack_conditions[self.depth];
104124
self.stack[self.depth] = .undefined;
105125
self.stack_integers[self.depth] = .unknown;
106126
self.stack_ranges[self.depth] = .unknown;
107-
return .{ .kind = kind, .integer = integer, .range = range };
127+
self.stack_locals[self.depth] = null;
128+
self.stack_conditions[self.depth] = null;
129+
return .{ .kind = kind, .integer = integer, .range = range, .local = local, .condition = condition };
108130
}
109131
};
110132

@@ -276,6 +298,28 @@ fn joinKind(left: Kind, right: Kind) Kind {
276298
return .unknown;
277299
}
278300

301+
fn joinCondition(left: ?Condition, right: ?Condition, widen_ranges: bool) ?Condition {
302+
const lhs = left orelse return null;
303+
const rhs = right orelse return null;
304+
if (lhs.op != rhs.op or
305+
!std.meta.eql(lhs.lhs.local, rhs.lhs.local) or
306+
!std.meta.eql(lhs.rhs.local, rhs.rhs.local))
307+
return null;
308+
return .{
309+
.op = lhs.op,
310+
.lhs = .{
311+
.integer = joinIntegerFact(lhs.lhs.integer, rhs.lhs.integer),
312+
.range = joinIntegerRange(lhs.lhs.range, rhs.lhs.range, widen_ranges),
313+
.local = lhs.lhs.local,
314+
},
315+
.rhs = .{
316+
.integer = joinIntegerFact(lhs.rhs.integer, rhs.rhs.integer),
317+
.range = joinIntegerRange(lhs.rhs.range, rhs.rhs.range, widen_ranges),
318+
.local = lhs.rhs.local,
319+
},
320+
};
321+
}
322+
279323
fn mergeState(existing: *State, incoming: State, widen_ranges: bool) error{UnsupportedChunk}!bool {
280324
if (existing.depth != incoming.depth) return error.UnsupportedChunk;
281325
for (existing.stack[0..existing.depth], incoming.stack[0..incoming.depth]) |current, next|
@@ -317,6 +361,19 @@ fn mergeState(existing: *State, incoming: State, widen_ranges: bool) error{Unsup
317361
changed = true;
318362
}
319363
}
364+
for (existing.stack_locals[0..existing.depth], incoming.stack_locals[0..incoming.depth]) |*current, next| {
365+
if (!std.meta.eql(current.*, next) and current.* != null) {
366+
current.* = null;
367+
changed = true;
368+
}
369+
}
370+
for (existing.stack_conditions[0..existing.depth], incoming.stack_conditions[0..incoming.depth]) |*current, next| {
371+
const merged = joinCondition(current.*, next, widen_ranges);
372+
if (!std.meta.eql(current.*, merged)) {
373+
current.* = merged;
374+
changed = true;
375+
}
376+
}
320377
return changed;
321378
}
322379

@@ -371,6 +428,63 @@ fn arithmeticIntegerRange(op: bc.Op, lhs: StackValue, rhs: StackValue) IntegerRa
371428
};
372429
}
373430

431+
fn operandLowerBound(operand: RangeOperand) ?i64 {
432+
if (operand.range.isKnown()) return operand.range.min;
433+
return switch (operand.integer) {
434+
.positive => 1,
435+
.nonnegative => 0,
436+
.signed, .unknown => null,
437+
};
438+
}
439+
440+
fn operandUpperBound(operand: RangeOperand) ?i64 {
441+
return if (operand.range.isKnown()) operand.range.max else null;
442+
}
443+
444+
fn refineOperand(state: *State, operand: RangeOperand, minimum: ?i64, maximum: ?i64) void {
445+
const slot = operand.local orelse return;
446+
var lower = operandLowerBound(operand);
447+
var upper = operandUpperBound(operand);
448+
if (minimum) |bound| lower = if (lower) |current| @max(current, bound) else bound;
449+
if (maximum) |bound| upper = if (upper) |current| @min(current, bound) else bound;
450+
if (lower == null or upper == null or lower.? > upper.?) return;
451+
state.local_ranges[slot] = .{ .min = lower.?, .max = upper.? };
452+
}
453+
454+
fn refineOrdered(state: *State, lhs: RangeOperand, rhs: RangeOperand, inclusive: bool) void {
455+
const delta: i64 = if (inclusive) 0 else 1;
456+
const lhs_max = if (operandUpperBound(rhs)) |bound| bound - delta else null;
457+
const rhs_min = if (operandLowerBound(lhs)) |bound| bound + delta else null;
458+
refineOperand(state, lhs, null, lhs_max);
459+
refineOperand(state, rhs, rhs_min, null);
460+
}
461+
462+
fn refineEqual(state: *State, lhs: RangeOperand, rhs: RangeOperand) void {
463+
const minimum = if (operandLowerBound(lhs)) |lhs_min|
464+
if (operandLowerBound(rhs)) |rhs_min| @max(lhs_min, rhs_min) else lhs_min
465+
else
466+
operandLowerBound(rhs);
467+
const maximum = if (operandUpperBound(lhs)) |lhs_max|
468+
if (operandUpperBound(rhs)) |rhs_max| @min(lhs_max, rhs_max) else lhs_max
469+
else
470+
operandUpperBound(rhs);
471+
refineOperand(state, lhs, minimum, maximum);
472+
refineOperand(state, rhs, minimum, maximum);
473+
}
474+
475+
fn refineComparison(state: *State, condition: Condition, truth: bool) void {
476+
if (condition.lhs.integer == .unknown or condition.rhs.integer == .unknown) return;
477+
switch (condition.op) {
478+
.lt => if (truth) refineOrdered(state, condition.lhs, condition.rhs, false) else refineOrdered(state, condition.rhs, condition.lhs, true),
479+
.le => if (truth) refineOrdered(state, condition.lhs, condition.rhs, true) else refineOrdered(state, condition.rhs, condition.lhs, false),
480+
.gt => if (truth) refineOrdered(state, condition.rhs, condition.lhs, false) else refineOrdered(state, condition.lhs, condition.rhs, true),
481+
.ge => if (truth) refineOrdered(state, condition.rhs, condition.lhs, true) else refineOrdered(state, condition.lhs, condition.rhs, false),
482+
.eq, .eq_strict => if (truth) refineEqual(state, condition.lhs, condition.rhs),
483+
.neq, .neq_strict => if (!truth) refineEqual(state, condition.lhs, condition.rhs),
484+
else => unreachable,
485+
}
486+
}
487+
374488
fn enqueueState(
375489
states: []?State,
376490
worklist: *std.ArrayListUnmanaged(u32),
@@ -427,23 +541,28 @@ fn analyzeNumeric(chunk: *const Chunk, integer_parameters: bool) !Analysis {
427541
.load_const => {
428542
if (inst.a >= chunk.consts.items.len) return error.UnsupportedChunk;
429543
const constant = chunk.consts.items[inst.a];
430-
try state.push(
431-
classify(constant) orelse return error.UnsupportedChunk,
432-
classifyInteger(constant),
433-
classifyIntegerRange(constant),
434-
);
544+
try state.push(.{
545+
.kind = classify(constant) orelse return error.UnsupportedChunk,
546+
.integer = classifyInteger(constant),
547+
.range = classifyIntegerRange(constant),
548+
});
435549
},
436-
.load_undefined => try state.push(.undefined, .unknown, .unknown),
437-
.load_null => try state.push(.null, .unknown, .unknown),
438-
.load_true, .load_false => try state.push(.boolean, .unknown, .unknown),
550+
.load_undefined => try state.push(.{ .kind = .undefined, .integer = .unknown, .range = .unknown }),
551+
.load_null => try state.push(.{ .kind = .null, .integer = .unknown, .range = .unknown }),
552+
.load_true, .load_false => try state.push(.{ .kind = .boolean, .integer = .unknown, .range = .unknown }),
439553
.pop => _ = try state.pop(),
440554
.load_local => {
441555
if (inst.a >= chunk.local_count) return error.UnsupportedChunk;
442556
// A local whose representation differs across predecessors is
443557
// safe only after an unconditional overwrite. Reconciliation
444558
// at a live mixed-kind merge is outside this numeric tier.
445559
if (state.locals[inst.a] == .unknown) return error.UnsupportedChunk;
446-
try state.push(state.locals[inst.a], state.local_integers[inst.a], state.local_ranges[inst.a]);
560+
try state.push(.{
561+
.kind = state.locals[inst.a],
562+
.integer = state.local_integers[inst.a],
563+
.range = state.local_ranges[inst.a],
564+
.local = @intCast(inst.a),
565+
});
447566
},
448567
.store_local => {
449568
if (inst.a >= chunk.local_count or state.depth == 0) return error.UnsupportedChunk;
@@ -455,23 +574,40 @@ fn analyzeNumeric(chunk: *const Chunk, integer_parameters: bool) !Analysis {
455574
const rhs = try state.pop();
456575
const lhs = try state.pop();
457576
if (rhs.kind != .number or lhs.kind != .number) return error.UnsupportedChunk;
458-
try state.push(
459-
.number,
460-
arithmeticIntegerFact(inst.op, lhs.integer, rhs.integer),
461-
arithmeticIntegerRange(inst.op, lhs, rhs),
462-
);
577+
try state.push(.{
578+
.kind = .number,
579+
.integer = arithmeticIntegerFact(inst.op, lhs.integer, rhs.integer),
580+
.range = arithmeticIntegerRange(inst.op, lhs, rhs),
581+
});
463582
},
464583
.lt, .le, .gt, .ge, .eq, .neq, .eq_strict, .neq_strict => {
465-
if ((try state.pop()).kind != .number or (try state.pop()).kind != .number) return error.UnsupportedChunk;
466-
try state.push(.boolean, .unknown, .unknown);
584+
const rhs = try state.pop();
585+
const lhs = try state.pop();
586+
if (rhs.kind != .number or lhs.kind != .number) return error.UnsupportedChunk;
587+
try state.push(.{
588+
.kind = .boolean,
589+
.integer = .unknown,
590+
.range = .unknown,
591+
.condition = .{
592+
.op = inst.op,
593+
.lhs = .{ .integer = lhs.integer, .range = lhs.range, .local = lhs.local },
594+
.rhs = .{ .integer = rhs.integer, .range = rhs.range, .local = rhs.local },
595+
},
596+
});
467597
},
468598
.jump => {
469599
try enqueueState(states, &worklist, allocator, inst.a, state, inst.a <= ip);
470600
fallthrough = false;
471601
},
472602
.jump_if_false => {
473-
if ((try state.pop()).kind != .boolean) return error.UnsupportedChunk;
474-
try enqueueState(states, &worklist, allocator, inst.a, state, inst.a <= ip);
603+
const condition_value = try state.pop();
604+
if (condition_value.kind != .boolean) return error.UnsupportedChunk;
605+
var false_state = state;
606+
if (condition_value.condition) |condition| {
607+
refineComparison(&false_state, condition, false);
608+
refineComparison(&state, condition, true);
609+
}
610+
try enqueueState(states, &worklist, allocator, inst.a, false_state, inst.a <= ip);
475611
},
476612
.ret => {
477613
_ = try state.pop();
@@ -1067,9 +1203,22 @@ test "integer provenance converges through benchmark-shaped loops" {
10671203
var analysis = try analyzeNumeric(chunk, true);
10681204
defer analysis.deinit();
10691205

1206+
var saw_inner_bound = false;
10701207
var saw_remainder = false;
10711208
var saw_return = false;
10721209
for (chunk.code.items, 0..) |inst, ip| switch (inst.op) {
1210+
.lt => {
1211+
const state = analysis.states[ip].?;
1212+
if (state.depth >= 2 and
1213+
std.meta.eql(state.stack_ranges[state.depth - 1], IntegerRange{ .min = 100_000, .max = 100_000 }))
1214+
{
1215+
const induction_slot = state.stack_locals[state.depth - 2].?;
1216+
try std.testing.expectEqual(bc.Op.jump_if_false, chunk.code.items[ip + 1].op);
1217+
const body_state = analysis.states[ip + 2].?;
1218+
try std.testing.expectEqual(IntegerRange{ .min = 0, .max = 99_999 }, body_state.local_ranges[induction_slot]);
1219+
saw_inner_bound = true;
1220+
}
1221+
},
10731222
.mod => {
10741223
const state = analysis.states[ip].?;
10751224
try std.testing.expect(state.depth >= 2);
@@ -1088,6 +1237,7 @@ test "integer provenance converges through benchmark-shaped loops" {
10881237
},
10891238
else => {},
10901239
};
1240+
try std.testing.expect(saw_inner_bound);
10911241
try std.testing.expect(saw_remainder);
10921242
try std.testing.expect(saw_return);
10931243
}

0 commit comments

Comments
 (0)