Skip to main content

alduin/backend/
reg_alloc.rs

1use std::marker::PhantomData;
2
3use bitvec::{bitvec, vec::BitVec};
4
5use crate::compiler::graph::cfg::{
6    liveness::{LiveInterval, LiveIntervalValueCategory},
7    CFG,
8};
9
10use super::{Reg, ISA};
11
12/// See: https://link.springer.com/content/pdf/10.1007/3-540-45937-5_17.pdf
13pub struct LinearScanRegisterAllocator<'cfg, Isa: ISA> {
14    cfg: &'cfg mut CFG<Isa::Op>,
15    unhandled: Vec<usize>,
16    active: Vec<usize>,
17    inactive: Vec<usize>,
18    handled: Vec<usize>,
19    free: BitVec,
20    weights: Vec<usize>,
21    _p: PhantomData<Isa>,
22}
23
24impl<'cfg, Isa: ISA> LinearScanRegisterAllocator<'cfg, Isa> {
25    pub fn new(cfg: &'cfg mut CFG<Isa::Op>) -> Self {
26        cfg.used_registers = bitvec![0; Isa::Reg::MAX_COUNT];
27        Self {
28            cfg,
29            unhandled: vec![],
30            active: vec![],
31            inactive: vec![],
32            handled: vec![],
33            free: bitvec![0; Isa::Reg::MAX_COUNT],
34            weights: vec![0; Isa::Reg::MAX_COUNT],
35            _p: PhantomData,
36        }
37    }
38
39    fn interval(&self, index: usize) -> &LiveInterval {
40        &self.cfg.liveness[index]
41    }
42
43    fn initialize(&mut self) {
44        let mut unhandled = vec![];
45        for (i, interval) in self.cfg.liveness.iter().enumerate() {
46            assert!(!interval.ranges.is_empty());
47            unhandled.push(i);
48        }
49        unhandled.sort_by_key(|i| self.interval(*i).ranges[0].start);
50        self.unhandled = unhandled;
51        self.active = vec![];
52        self.inactive = vec![];
53        self.handled = vec![];
54        for gpr in Isa::Reg::GPRS {
55            self.free.set((*gpr).into(), true);
56        }
57        for gpr in Isa::Reg::FPRS {
58            self.free.set((*gpr).into(), true);
59        }
60        self.assign_stack_location_for_all_intervals();
61    }
62
63    fn assign_stack_location_for_all_intervals(&mut self) {
64        let mut offset = 0;
65        for interval in &mut *self.cfg.liveness {
66            interval.mem = offset as i32;
67            offset += interval.max_mem_size;
68        }
69        let align_mask = (1 << 4) - 1;
70        self.cfg.stack_size = (offset + align_mask) & !align_mask;
71    }
72
73    fn allocate_mem_loc(&mut self, curr_index: usize) {
74        for i in 0..self.weights.len() {
75            self.weights[i] = 0;
76        }
77        let mut f = |i: usize| {
78            if self.cfg.liveness[i].intersects(&self.cfg.liveness[curr_index]) {
79                let i_reg = self.cfg.liveness[i].reg.unwrap();
80                self.weights[i_reg] += self.cfg.liveness[i].weight;
81            }
82        };
83        for i in 0..self.active.len() {
84            f(self.active[i]);
85        }
86        for i in 0..self.inactive.len() {
87            f(self.inactive[i]);
88        }
89        for i in 0..self.unhandled.len() {
90            let it = self.unhandled[i];
91            if self.cfg.liveness[it].reg.is_some() {
92                f(it);
93            }
94        }
95        // find r with minimum weights[r]
96        let mut min_index = 0;
97        let mut min_value = usize::MAX;
98        let v_cat = self.cfg.liveness[curr_index].value_category.unwrap();
99        for i in 0..self.weights.len() {
100            if v_cat == LiveIntervalValueCategory::Int
101                && (Isa::Reg::from(i).is_fpr()
102                    || Isa::Reg::RESERVED_GPRS.contains(&Isa::Reg::from(i)))
103            {
104                continue;
105            }
106            if v_cat == LiveIntervalValueCategory::Float
107                && (Isa::Reg::from(i).is_gpr()
108                    || Isa::Reg::RESERVED_FPRS.contains(&Isa::Reg::from(i)))
109            {
110                continue;
111            }
112            if self.weights[i] < min_value {
113                min_index = i;
114                min_value = self.weights[i];
115            }
116        }
117        assert_ne!(min_value, usize::MAX);
118        let selected_reg = min_index;
119        if self.interval(curr_index).weight < self.weights[selected_reg]
120            || self.interval(curr_index).reg.is_some()
121        {
122            // if let Some(selected_reg) = self.interval(curr_index).reg {
123            //     // move all active or inactive intervals to which r was assigned to handled
124            //     // and assign memory locations to them
125            //     for index in std::mem::take(&mut self.active) {
126            //         if self.interval(index).reg == Some(selected_reg) {
127            //             self.assign_mem(index);
128            //             self.handled.push(index);
129            //         } else {
130            //             self.active.push(index);
131            //         }
132            //     }
133            //     for index in std::mem::take(&mut self.inactive) {
134            //         if self.interval(index).reg == Some(selected_reg) {
135            //             self.assign_mem(index);
136            //             self.handled.push(index);
137            //         } else {
138            //             self.inactive.push(index);
139            //         }
140            //     }
141            // }
142            // assign a memory location to cur and move cur to handled
143            self.assign_mem(curr_index);
144            self.handled.push(curr_index);
145        } else {
146            // move all active or inactive intervals to which r was assigned to handled
147            // and assign memory locations to them
148            for index in std::mem::take(&mut self.active) {
149                if self.interval(index).reg == Some(selected_reg) {
150                    self.assign_mem(index);
151                    self.handled.push(index);
152                } else {
153                    self.active.push(index);
154                }
155            }
156            for index in std::mem::take(&mut self.inactive) {
157                if self.interval(index).reg == Some(selected_reg) {
158                    self.assign_mem(index);
159                    self.handled.push(index);
160                } else {
161                    self.inactive.push(index);
162                }
163            }
164            // assign selected_reg to curr
165            self.assign_reg(curr_index, Some(selected_reg));
166            // move curr to active
167            self.active.push(curr_index);
168        }
169    }
170
171    /// Assign a stack location to an interval.
172    /// This also removes the previously assigned register.
173    fn assign_mem(&mut self, interval_index: usize) {
174        // All intervals are pre-assigned with a stack location.
175        // We only need to un-assign the register here
176        self.assign_reg(interval_index, None);
177    }
178
179    /// Assign a register to an interval
180    fn assign_reg(&mut self, interval_index: usize, reg: Option<usize>) {
181        if let Some(reg) = reg {
182            self.cfg.used_registers.set(reg, true);
183        }
184        self.cfg.liveness[interval_index].reg = reg;
185    }
186
187    pub fn allocate(&mut self) {
188        self.initialize();
189        while !self.unhandled.is_empty() {
190            let curr_index = self.unhandled.remove(0);
191            let curr = self.interval(curr_index);
192            // check for internval in `active` that are handled or inactive
193            let curr_begin = curr.ranges[0].start;
194            for index in std::mem::take(&mut self.active) {
195                let it = self.interval(index);
196                let it_reg = it.reg.clone();
197                if it.ranges.last().unwrap().end <= curr_begin {
198                    self.handled.push(index);
199                    if let Some(reg) = it_reg {
200                        self.free.set(reg, true);
201                    }
202                } else if !it.covers(curr_begin) {
203                    self.inactive.push(index);
204                    if let Some(reg) = it_reg {
205                        self.free.set(reg, true);
206                    }
207                } else {
208                    self.active.push(index);
209                }
210            }
211            // check for intervals in `inactive` that are handled or active
212            for index in std::mem::take(&mut self.inactive) {
213                let it = self.interval(index);
214                let it_reg = it.reg.clone();
215                if it.ranges.last().unwrap().end <= curr_begin {
216                    self.handled.push(index);
217                } else if it.covers(curr_begin) {
218                    self.active.push(index);
219                    if let Some(reg) = it_reg {
220                        self.free.set(reg, false);
221                    }
222                } else {
223                    self.inactive.push(index);
224                }
225            }
226            // collect available registers
227            let curr = self.interval(curr_index);
228            let mut free_regs = self.free.clone();
229            for i in &self.inactive {
230                let it = self.interval(*i);
231                if it.intersects(curr) {
232                    if let Some(reg) = it.reg {
233                        free_regs.set(reg, false)
234                    }
235                }
236            }
237            for i in &self.unhandled {
238                let it = self.interval(*i);
239                if it.reg.is_some() && it.intersects(curr) {
240                    free_regs.set(it.reg.unwrap(), false);
241                }
242            }
243            if curr.value_category.unwrap() == LiveIntervalValueCategory::Float {
244                // remove all GPRs
245                for r in Isa::Reg::GPRS {
246                    free_regs.set((*r).into(), false);
247                }
248            } else {
249                // remove all FPRs
250                for r in Isa::Reg::FPRS {
251                    free_regs.set((*r).into(), false);
252                }
253            }
254            // select a register
255            if !free_regs.any() || curr.reg.map(|r| !free_regs[r]).unwrap_or(false) {
256                self.allocate_mem_loc(curr_index);
257            } else {
258                let reg = if curr.reg.is_none() {
259                    let r = free_regs.first_one().unwrap();
260                    self.assign_reg(curr_index, Some(r));
261                    r
262                } else {
263                    assert!(
264                        free_regs[curr.reg.unwrap()],
265                        "Ref for fixed interval #{} is not free",
266                        curr_index
267                    );
268                    curr.reg.unwrap()
269                };
270                self.free.set(reg, false);
271                self.active.push(curr_index);
272            }
273        }
274        // self.cfg.liveness.dump(self.cfg);
275    }
276}