| ... | @@ -4,6 +4,7 @@ const mem = std.mem; | ... | @@ -4,6 +4,7 @@ const mem = std.mem; |
| 4 | const assert = std.debug.assert; | 4 | const assert = std.debug.assert; |
| 5 | const Allocator = std.mem.Allocator; | 5 | const Allocator = std.mem.Allocator; |
| 6 | const Air = @import("Air.zig"); | 6 | const Air = @import("Air.zig"); |
| | 7 | const StaticBitSet = std.bit_set.StaticBitSet; |
| 7 | const Type = @import("type.zig").Type; | 8 | const Type = @import("type.zig").Type; |
| 8 | const Module = @import("Module.zig"); | 9 | const Module = @import("Module.zig"); |
| 9 | const expect = std.testing.expect; | 10 | const expect = std.testing.expect; |
| ... | @@ -41,66 +42,54 @@ pub fn RegisterManager( | ... | @@ -41,66 +42,54 @@ pub fn RegisterManager( |
| 41 | registers: [tracked_registers.len]Air.Inst.Index = undefined, | 42 | registers: [tracked_registers.len]Air.Inst.Index = undefined, |
| 42 | /// Tracks which registers are free (in which case the | 43 | /// Tracks which registers are free (in which case the |
| 43 | /// corresponding bit is set to 1) | 44 | /// corresponding bit is set to 1) |
| 44 | free_registers: FreeRegInt = math.maxInt(FreeRegInt), | 45 | free_registers: RegisterBitSet = RegisterBitSet.initFull(), |
| 45 | /// Tracks all registers allocated in the course of this | 46 | /// Tracks all registers allocated in the course of this |
| 46 | /// function | 47 | /// function |
| 47 | allocated_registers: FreeRegInt = 0, | 48 | allocated_registers: RegisterBitSet = RegisterBitSet.initEmpty(), |
| 48 | /// Tracks registers which are locked from being allocated | 49 | /// Tracks registers which are locked from being allocated |
| 49 | locked_registers: FreeRegInt = 0, | 50 | locked_registers: RegisterBitSet = RegisterBitSet.initEmpty(), |
| 50 | | 51 | |
| 51 | const Self = @This(); | 52 | const Self = @This(); |
| 52 | | 53 | |
| 53 | /// An integer whose bits represent all the registers and | 54 | pub const RegisterBitSet = StaticBitSet(tracked_registers.len); |
| 54 | /// whether they are free. | | |
| 55 | const FreeRegInt = std.meta.Int(.unsigned, tracked_registers.len); | | |
| 56 | const ShiftInt = math.Log2Int(FreeRegInt); | | |
| 57 | | 55 | |
| 58 | fn getFunction(self: *Self) *Function { | 56 | fn getFunction(self: *Self) *Function { |
| 59 | return @fieldParentPtr(Function, "register_manager", self); | 57 | return @fieldParentPtr(Function, "register_manager", self); |
| 60 | } | 58 | } |
| 61 | | 59 | |
| 62 | fn getRegisterMask(reg: Register) ?FreeRegInt { | | |
| 63 | const index = indexOfRegIntoTracked(reg) orelse return null; | | |
| 64 | const shift = @intCast(ShiftInt, index); | | |
| 65 | const mask = @as(FreeRegInt, 1) << shift; | | |
| 66 | return mask; | | |
| 67 | } | | |
| 68 | | | |
| 69 | fn excludeRegister(reg: Register, mask: FreeRegInt) bool { | | |
| 70 | const reg_mask = getRegisterMask(reg) orelse return true; | | |
| 71 | return reg_mask & mask == 0; | | |
| 72 | } | | |
| 73 | | | |
| 74 | fn markRegAllocated(self: *Self, reg: Register) void { | 60 | fn markRegAllocated(self: *Self, reg: Register) void { |
| 75 | const mask = getRegisterMask(reg) orelse return; | 61 | const index = indexOfRegIntoTracked(reg) orelse return; |
| 76 | self.allocated_registers |= mask; | 62 | self.allocated_registers.set(index); |
| 77 | } | 63 | } |
| 78 | | 64 | |
| 79 | fn markRegUsed(self: *Self, reg: Register) void { | 65 | fn markRegUsed(self: *Self, reg: Register) void { |
| 80 | const mask = getRegisterMask(reg) orelse return; | 66 | const index = indexOfRegIntoTracked(reg) orelse return; |
| 81 | self.free_registers &= ~mask; | 67 | self.free_registers.unset(index); |
| 82 | } | 68 | } |
| 83 | | 69 | |
| 84 | fn markRegFree(self: *Self, reg: Register) void { | 70 | fn markRegFree(self: *Self, reg: Register) void { |
| 85 | const mask = getRegisterMask(reg) orelse return; | 71 | const index = indexOfRegIntoTracked(reg) orelse return; |
| 86 | self.free_registers |= mask; | 72 | self.free_registers.set(index); |
| 87 | } | 73 | } |
| 88 | | 74 | |
| 89 | pub fn indexOfReg(comptime registers: []const Register, reg: Register) ?std.math.IntFittingRange(0, registers.len - 1) { | 75 | pub fn indexOfReg( |
| | 76 | comptime registers: []const Register, |
| | 77 | reg: Register, |
| | 78 | ) ?std.math.IntFittingRange(0, registers.len - 1) { |
| 90 | inline for (tracked_registers) |cpreg, i| { | 79 | inline for (tracked_registers) |cpreg, i| { |
| 91 | if (reg.id() == cpreg.id()) return i; | 80 | if (reg.id() == cpreg.id()) return i; |
| 92 | } | 81 | } |
| 93 | return null; | 82 | return null; |
| 94 | } | 83 | } |
| 95 | | 84 | |
| 96 | pub fn indexOfRegIntoTracked(reg: Register) ?ShiftInt { | 85 | pub fn indexOfRegIntoTracked(reg: Register) ?RegisterBitSet.ShiftInt { |
| 97 | return indexOfReg(tracked_registers, reg); | 86 | return indexOfReg(tracked_registers, reg); |
| 98 | } | 87 | } |
| 99 | | 88 | |
| 100 | /// Returns true when this register is not tracked | 89 | /// Returns true when this register is not tracked |
| 101 | pub fn isRegFree(self: Self, reg: Register) bool { | 90 | pub fn isRegFree(self: Self, reg: Register) bool { |
| 102 | const mask = getRegisterMask(reg) orelse return true; | 91 | const index = indexOfRegIntoTracked(reg) orelse return true; |
| 103 | return self.free_registers & mask != 0; | 92 | return self.free_registers.isSet(index); |
| 104 | } | 93 | } |
| 105 | | 94 | |
| 106 | /// Returns whether this register was allocated in the course | 95 | /// Returns whether this register was allocated in the course |
| ... | @@ -108,16 +97,16 @@ pub fn RegisterManager( | ... | @@ -108,16 +97,16 @@ pub fn RegisterManager( |
| 108 | /// | 97 | /// |
| 109 | /// Returns false when this register is not tracked | 98 | /// Returns false when this register is not tracked |
| 110 | pub fn isRegAllocated(self: Self, reg: Register) bool { | 99 | pub fn isRegAllocated(self: Self, reg: Register) bool { |
| 111 | const mask = getRegisterMask(reg) orelse return false; | 100 | const index = indexOfRegIntoTracked(reg) orelse return false; |
| 112 | return self.allocated_registers & mask != 0; | 101 | return self.allocated_registers.isSet(index); |
| 113 | } | 102 | } |
| 114 | | 103 | |
| 115 | /// Returns whether this register is locked | 104 | /// Returns whether this register is locked |
| 116 | /// | 105 | /// |
| 117 | /// Returns false when this register is not tracked | 106 | /// Returns false when this register is not tracked |
| 118 | pub fn isRegLocked(self: Self, reg: Register) bool { | 107 | pub fn isRegLocked(self: Self, reg: Register) bool { |
| 119 | const mask = getRegisterMask(reg) orelse return false; | 108 | const index = indexOfRegIntoTracked(reg) orelse return false; |
| 120 | return self.locked_registers & mask != 0; | 109 | return self.locked_registers.isSet(index); |
| 121 | } | 110 | } |
| 122 | | 111 | |
| 123 | pub const RegisterLock = struct { | 112 | pub const RegisterLock = struct { |
| ... | @@ -136,8 +125,8 @@ pub fn RegisterManager( | ... | @@ -136,8 +125,8 @@ pub fn RegisterManager( |
| 136 | log.debug(" register already locked", .{}); | 125 | log.debug(" register already locked", .{}); |
| 137 | return null; | 126 | return null; |
| 138 | } | 127 | } |
| 139 | const mask = getRegisterMask(reg) orelse return null; | 128 | const index = indexOfRegIntoTracked(reg) orelse return null; |
| 140 | self.locked_registers |= mask; | 129 | self.locked_registers.set(index); |
| 141 | return RegisterLock{ .register = reg }; | 130 | return RegisterLock{ .register = reg }; |
| 142 | } | 131 | } |
| 143 | | 132 | |
| ... | @@ -146,8 +135,8 @@ pub fn RegisterManager( | ... | @@ -146,8 +135,8 @@ pub fn RegisterManager( |
| 146 | pub fn lockRegAssumeUnused(self: *Self, reg: Register) RegisterLock { | 135 | pub fn lockRegAssumeUnused(self: *Self, reg: Register) RegisterLock { |
| 147 | log.debug("locking asserting free {}", .{reg}); | 136 | log.debug("locking asserting free {}", .{reg}); |
| 148 | assert(!self.isRegLocked(reg)); | 137 | assert(!self.isRegLocked(reg)); |
| 149 | const mask = getRegisterMask(reg) orelse unreachable; | 138 | const index = indexOfRegIntoTracked(reg) orelse unreachable; |
| 150 | self.locked_registers |= mask; | 139 | self.locked_registers.set(index); |
| 151 | return RegisterLock{ .register = reg }; | 140 | return RegisterLock{ .register = reg }; |
| 152 | } | 141 | } |
| 153 | | 142 | |
| ... | @@ -169,17 +158,17 @@ pub fn RegisterManager( | ... | @@ -169,17 +158,17 @@ pub fn RegisterManager( |
| 169 | /// Call `lockReg` to obtain the lock first. | 158 | /// Call `lockReg` to obtain the lock first. |
| 170 | pub fn unlockReg(self: *Self, lock: RegisterLock) void { | 159 | pub fn unlockReg(self: *Self, lock: RegisterLock) void { |
| 171 | log.debug("unlocking {}", .{lock.register}); | 160 | log.debug("unlocking {}", .{lock.register}); |
| 172 | const mask = getRegisterMask(lock.register) orelse return; | 161 | const index = indexOfRegIntoTracked(lock.register) orelse return; |
| 173 | self.locked_registers &= ~mask; | 162 | self.locked_registers.unset(index); |
| 174 | } | 163 | } |
| 175 | | 164 | |
| 176 | /// Returns true when at least one register is locked | 165 | /// Returns true when at least one register is locked |
| 177 | pub fn lockedRegsExist(self: Self) bool { | 166 | pub fn lockedRegsExist(self: Self) bool { |
| 178 | return self.locked_registers != 0; | 167 | return self.locked_registers.count() > 0; |
| 179 | } | 168 | } |
| 180 | | 169 | |
| 181 | const AllocOpts = struct { | 170 | const AllocOpts = struct { |
| 182 | selector_mask: ?FreeRegInt = null, | 171 | selector_mask: ?RegisterBitSet = null, |
| 183 | }; | 172 | }; |
| 184 | | 173 | |
| 185 | /// Allocates a specified number of registers, optionally | 174 | /// Allocates a specified number of registers, optionally |
| ... | @@ -193,17 +182,22 @@ pub fn RegisterManager( | ... | @@ -193,17 +182,22 @@ pub fn RegisterManager( |
| 193 | ) ?[count]Register { | 182 | ) ?[count]Register { |
| 194 | comptime assert(count > 0 and count <= tracked_registers.len); | 183 | comptime assert(count > 0 and count <= tracked_registers.len); |
| 195 | | 184 | |
| 196 | const selector_mask = if (opts.selector_mask) |mask| mask else ~@as(FreeRegInt, 0); | 185 | const available_registers = opts.selector_mask orelse RegisterBitSet.initFull(); |
| 197 | const free_registers = self.free_registers & selector_mask; | 186 | |
| 198 | const free_and_not_locked_registers = free_registers & ~self.locked_registers; | 187 | var free_and_not_locked_registers = self.free_registers; |
| 199 | const free_and_not_locked_registers_count = @popCount(FreeRegInt, free_and_not_locked_registers); | 188 | free_and_not_locked_registers.setIntersection(available_registers); |
| 200 | if (free_and_not_locked_registers_count < count) return null; | 189 | |
| | 190 | var unlocked_registers = self.locked_registers; |
| | 191 | unlocked_registers.toggleAll(); |
| | 192 | |
| | 193 | free_and_not_locked_registers.setIntersection(unlocked_registers); |
| | 194 | |
| | 195 | if (free_and_not_locked_registers.count() < count) return null; |
| 201 | | 196 | |
| 202 | var regs: [count]Register = undefined; | 197 | var regs: [count]Register = undefined; |
| 203 | var i: usize = 0; | 198 | var i: usize = 0; |
| 204 | for (tracked_registers) |reg| { | 199 | for (tracked_registers) |reg| { |
| 205 | if (i >= count) break; | 200 | if (i >= count) break; |
| 206 | if (excludeRegister(reg, selector_mask)) continue; | | |
| 207 | if (self.isRegLocked(reg)) continue; | 201 | if (self.isRegLocked(reg)) continue; |
| 208 | if (!self.isRegFree(reg)) continue; | 202 | if (!self.isRegFree(reg)) continue; |
| 209 | | 203 | |
| ... | @@ -244,11 +238,12 @@ pub fn RegisterManager( | ... | @@ -244,11 +238,12 @@ pub fn RegisterManager( |
| 244 | ) AllocateRegistersError![count]Register { | 238 | ) AllocateRegistersError![count]Register { |
| 245 | comptime assert(count > 0 and count <= tracked_registers.len); | 239 | comptime assert(count > 0 and count <= tracked_registers.len); |
| 246 | | 240 | |
| 247 | const selector_mask = if (opts.selector_mask) |mask| mask else ~@as(FreeRegInt, 0); | 241 | const available_registers = opts.selector_mask orelse RegisterBitSet.initFull(); |
| 248 | const available_registers_count = @popCount(FreeRegInt, selector_mask); | 242 | |
| 249 | const locked_registers = self.locked_registers & selector_mask; | 243 | var locked_registers = self.locked_registers; |
| 250 | const locked_registers_count = @popCount(FreeRegInt, locked_registers); | 244 | locked_registers.setIntersection(available_registers); |
| 251 | if (count > available_registers_count - locked_registers_count) return error.OutOfRegisters; | 245 | |
| | 246 | if (count > available_registers.count() - locked_registers.count()) return error.OutOfRegisters; |
| 252 | | 247 | |
| 253 | const result = self.tryAllocRegs(count, insts, opts) orelse blk: { | 248 | const result = self.tryAllocRegs(count, insts, opts) orelse blk: { |
| 254 | // We'll take over the first count registers. Spill | 249 | // We'll take over the first count registers. Spill |
| ... | @@ -258,7 +253,6 @@ pub fn RegisterManager( | ... | @@ -258,7 +253,6 @@ pub fn RegisterManager( |
| 258 | var i: usize = 0; | 253 | var i: usize = 0; |
| 259 | for (tracked_registers) |reg| { | 254 | for (tracked_registers) |reg| { |
| 260 | if (i >= count) break; | 255 | if (i >= count) break; |
| 261 | if (excludeRegister(reg, selector_mask)) continue; | | |
| 262 | if (self.isRegLocked(reg)) continue; | 256 | if (self.isRegLocked(reg)) continue; |
| 263 | | 257 | |
| 264 | regs[i] = reg; | 258 | regs[i] = reg; |