1use std::cmp::Reverse;
28use std::collections::BTreeMap;
29
30use rucc_mir::{Func, Reg};
31use rucc_target::RegClass;
32
33use crate::assign;
34use crate::backtrack;
35use crate::live::{Area, Live};
36use crate::pressure::Pressure;
37
38#[derive(Debug, Clone, Copy)]
40struct Candidate<'a> {
41 reg: Reg,
42 area: Area<'a>,
43 weight: u128,
44 size: u32,
45}
46
47#[must_use]
55pub fn choose(func: &Func, live: &Live, pressure: &Pressure) -> Vec<Reg> {
56 let costs = backtrack::costs(func);
57 let mut forced = vec![false; func.vregs()];
58 for reg in assign::forced(func) {
59 forced[index(reg)] = true;
60 }
61 let mut classes: BTreeMap<RegClass, Vec<Candidate<'_>>> = BTreeMap::new();
62 for (number, &forced) in forced.iter().enumerate() {
63 let reg = Reg::virtual_reg(u32::try_from(number).expect("a register number"));
64 let (Some(area), Some(class)) = (live.area(reg), func.class_of(reg)) else {
65 continue;
66 };
67 if forced {
68 continue;
69 }
70 let size = backtrack::size(area);
71 let weight = costs[number] * 1024 / u128::from(size + 8);
72 classes.entry(class).or_default().push(Candidate { reg, area, weight, size });
73 }
74
75 let mut pressure = pressure.clone();
76 let mut chosen = Vec::new();
77 for (class, mut candidates) in classes {
78 candidates.sort_by_key(|candidate| candidate.area.hull().start);
79 let over: Vec<_> = pressure.over(class).collect();
80 let mut open: Vec<usize> = Vec::new();
82 let mut gone = vec![false; candidates.len()];
84 let mut next = 0;
85 for point in over {
86 while next < candidates.len() && candidates[next].area.hull().start <= point {
87 open.push(next);
88 next += 1;
89 }
90 if pressure.excess(class, point) == 0 {
94 continue;
95 }
96 let mut here = Vec::new();
100 open.retain(|&one| {
101 let candidate = &candidates[one];
102 if gone[one] || candidate.area.hull().end < point {
103 return false;
104 }
105 if candidate.area.covers(point) {
106 here.push(one);
107 }
108 true
109 });
110 while pressure.excess(class, point) > 0 {
111 let lightest = (0..here.len()).min_by_key(|&at| {
112 let one = here[at];
113 (candidates[one].weight, Reverse(candidates[one].size), one)
114 });
115 let Some(at) = lightest else { break };
116 let one = here.swap_remove(at);
117 gone[one] = true;
118 pressure.lift(class, candidates[one].area);
119 chosen.push(candidates[one].reg);
120 }
121 }
122 }
123 chosen
124}
125
126fn index(reg: Reg) -> usize {
127 usize::try_from(reg.number().expect("a virtual register")).expect("a register number")
128}
129
130#[cfg(test)]
131mod tests {
132 use rucc_base::Interner;
133 use rucc_mir::{BlockCall, Opcode, Weight};
134 use rucc_target::x86_64::{GPR, SYSV};
135
136 use super::*;
137 use crate::assign::Env;
138 use crate::order::Order;
139
140 fn narrow(count: usize) -> Env {
141 Env::new().with(GPR, &SYSV.int_order[..count], &SYSV.int_order[count..count + 1])
142 }
143
144 fn chosen(func: &Func, env: &Env) -> Vec<Reg> {
145 let order = Order::of(func);
146 let live = Live::of(func, &order);
147 let pressure = Pressure::of(func, &order, &live, env);
148 choose(func, &live, &pressure)
149 }
150
151 #[test]
152 fn nothing_goes_where_everything_fits() {
153 let mut names = Interner::new();
154 let mut func = Func::new(names.intern("f"));
155 let opcode = Opcode::new(names.intern("x64.nop"));
156 let block = func.create_block();
157 let first = func.new_vreg(GPR);
158 let second = func.new_vreg(GPR);
159 func.build(block, opcode).def(first, GPR).finish();
160 func.build(block, opcode).def(second, GPR).finish();
161 func.build(block, opcode).uses(first, GPR).uses(second, GPR).finish();
162
163 assert!(chosen(&func, &narrow(2)).is_empty());
164 }
165
166 #[test]
167 fn the_value_read_least_is_the_one_that_goes() {
168 let mut names = Interner::new();
169 let mut func = Func::new(names.intern("f"));
170 let opcode = Opcode::new(names.intern("x64.nop"));
171 let block = func.create_block();
172 let busy = func.new_vreg(GPR);
173 let once = func.new_vreg(GPR);
174 let other = func.new_vreg(GPR);
175 func.build(block, opcode).def(busy, GPR).finish();
176 func.build(block, opcode).def(once, GPR).finish();
177 func.build(block, opcode).def(other, GPR).finish();
178 for _ in 0..3 {
179 func.build(block, opcode).uses(busy, GPR).uses(other, GPR).finish();
180 }
181 func.build(block, opcode).uses(once, GPR).uses(busy, GPR).uses(other, GPR).finish();
182
183 assert_eq!(chosen(&func, &narrow(2)), [once]);
184 }
185
186 #[test]
187 fn a_value_read_in_a_loop_stays_and_one_read_outside_it_goes() {
188 let mut names = Interner::new();
189 let mut func = Func::new(names.intern("f"));
190 let opcode = Opcode::new(names.intern("x64.nop"));
191 let entry = func.create_block();
192 let body = func.create_block();
193 let out = func.create_block();
194 let step = func.new_vreg(GPR);
195 let cold = func.new_vreg(GPR);
196 func.build(entry, opcode).def(cold, GPR).finish();
197 func.build(entry, opcode).def(step, GPR).finish();
198 *func.succs_mut(entry) = vec![BlockCall::to(body)];
199 func.build(body, opcode).uses(step, GPR).finish();
200 *func.succs_mut(body) = vec![BlockCall::to(body), BlockCall::to(out)];
201 func.set_weight(body, Weight::parts(100 * Weight::SCALE));
202 func.build(out, opcode).uses(cold, GPR).finish();
203 func.build(out, opcode).uses(step, GPR).finish();
204
205 assert_eq!(chosen(&func, &narrow(1)), [cold]);
206 }
207
208 #[test]
209 fn one_long_light_value_settles_two_points_that_are_each_one_over() {
210 let mut names = Interner::new();
211 let mut func = Func::new(names.intern("f"));
212 let opcode = Opcode::new(names.intern("x64.nop"));
213 let block = func.create_block();
214 let long = func.new_vreg(GPR);
215 let first = func.new_vreg(GPR);
216 let second = func.new_vreg(GPR);
217 func.build(block, opcode).def(long, GPR).finish();
218 func.build(block, opcode).def(first, GPR).finish();
219 func.build(block, opcode).uses(first, GPR).finish();
220 for _ in 0..4 {
221 func.build(block, opcode).finish();
222 }
223 func.build(block, opcode).def(second, GPR).finish();
224 func.build(block, opcode).uses(second, GPR).finish();
225 func.build(block, opcode).uses(long, GPR).finish();
226
227 assert_eq!(chosen(&func, &narrow(1)), [long]);
230 }
231}