use std::collections::HashMap;
use crate::instruction::Instr;
pub fn optimize(mut instructions: Vec<Instr>, temp_count: usize) -> (Vec<Instr>, usize) {
instructions = simplify(instructions);
instructions = cse(instructions);
instructions = eliminate_dead_code(instructions);
(instructions, temp_count)
}
fn simplify(instructions: Vec<Instr>) -> Vec<Instr> {
let mut result = Vec::with_capacity(instructions.len());
for instr in instructions {
match &instr {
Instr::Add { dst, srcs } if srcs.len() == 1 => {
result.push(Instr::Copy {
dst: *dst,
src: srcs[0],
});
}
Instr::Mul { dst, srcs } if srcs.len() == 1 => {
result.push(Instr::Copy {
dst: *dst,
src: srcs[0],
});
}
_ => result.push(instr),
}
}
result
}
fn cse(instructions: Vec<Instr>) -> Vec<Instr> {
let mut seen: HashMap<InstrKey, usize> = HashMap::new();
let mut result = Vec::with_capacity(instructions.len());
for instr in instructions {
let key = InstrKey::from_instr(&instr);
match key {
Some(k) => {
if let Some(&existing_dst) = seen.get(&k) {
result.push(Instr::Copy {
dst: instr.dst(),
src: existing_dst,
});
} else {
seen.insert(k, instr.dst());
result.push(instr);
}
}
None => {
result.push(instr);
}
}
}
result
}
fn eliminate_dead_code(instructions: Vec<Instr>) -> Vec<Instr> {
if instructions.is_empty() {
return instructions;
}
use std::collections::HashSet;
let mut live_slots: HashSet<usize> = HashSet::new();
if let Some(last) = instructions.last() {
live_slots.insert(last.dst());
}
let mut keep: Vec<bool> = vec![false; instructions.len()];
for (i, instr) in instructions.iter().enumerate().rev() {
if live_slots.contains(&instr.dst()) || is_side_effect(instr) {
keep[i] = true;
for src in instr_srcs(instr) {
live_slots.insert(src);
}
}
}
instructions
.into_iter()
.enumerate()
.filter(|(i, _)| keep[*i])
.map(|(_, instr)| instr)
.collect()
}
fn is_side_effect(instr: &Instr) -> bool {
matches!(instr, Instr::ExternalFun { .. })
}
fn instr_srcs(instr: &Instr) -> Vec<usize> {
match instr {
Instr::Add { srcs, .. } => srcs.clone(),
Instr::Mul { srcs, .. } => srcs.clone(),
Instr::Pow { base, .. } => vec![*base],
Instr::Powf { base, exp, .. } => vec![*base, *exp],
Instr::BuiltinOp { src, .. } => vec![*src],
Instr::ExternalFun { srcs, .. } => srcs.clone(),
Instr::Copy { src, .. } => vec![*src],
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
enum InstrKey {
Add(Vec<usize>),
Mul(Vec<usize>),
Pow(usize, i64),
Powf(usize, usize),
BuiltinOp(crate::instruction::BuiltinOp, usize),
Copy(usize),
}
impl InstrKey {
fn from_instr(instr: &Instr) -> Option<Self> {
match instr {
Instr::Add { srcs, .. } => Some(InstrKey::Add(srcs.clone())),
Instr::Mul { srcs, .. } => Some(InstrKey::Mul(srcs.clone())),
Instr::Pow { base, exp, .. } => Some(InstrKey::Pow(*base, *exp)),
Instr::Powf { base, exp, .. } => Some(InstrKey::Powf(*base, *exp)),
Instr::BuiltinOp { op, src, .. } => Some(InstrKey::BuiltinOp(*op, *src)),
Instr::Copy { src, .. } => Some(InstrKey::Copy(*src)),
Instr::ExternalFun { .. } => None,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn simplify_single_add() {
let instrs = vec![Instr::Add {
dst: 5,
srcs: vec![2],
}];
let result = simplify(instrs);
assert!(matches!(result[0], Instr::Copy { dst: 5, src: 2 }));
}
#[test]
fn simplify_single_mul() {
let instrs = vec![Instr::Mul {
dst: 5,
srcs: vec![2],
}];
let result = simplify(instrs);
assert!(matches!(result[0], Instr::Copy { dst: 5, src: 2 }));
}
#[test]
fn cse_duplicate_add() {
let instrs = vec![
Instr::Add {
dst: 3,
srcs: vec![0, 1],
},
Instr::Add {
dst: 4,
srcs: vec![0, 1],
},
];
let result = cse(instrs);
assert_eq!(result.len(), 2);
assert!(matches!(result[0], Instr::Add { dst: 3, .. }));
assert!(matches!(result[1], Instr::Copy { dst: 4, src: 3 }));
}
#[test]
fn cse_non_duplicate() {
let instrs = vec![
Instr::Add {
dst: 3,
srcs: vec![0, 1],
},
Instr::Mul {
dst: 4,
srcs: vec![0, 2],
},
];
let result = cse(instrs);
assert_eq!(result.len(), 2);
assert!(matches!(result[0], Instr::Add { .. }));
assert!(matches!(result[1], Instr::Mul { .. }));
}
#[test]
fn cse_copy() {
let instrs = vec![
Instr::Copy { dst: 3, src: 0 },
Instr::Copy { dst: 4, src: 0 },
];
let result = cse(instrs);
assert!(matches!(result[1], Instr::Copy { dst: 4, src: 3 }));
}
#[test]
fn dead_code_elimination_simple() {
let instrs = vec![
Instr::Copy { dst: 5, src: 0 }, Instr::Add {
dst: 6,
srcs: vec![0, 1],
}, ];
let result = eliminate_dead_code(instrs);
assert_eq!(result.len(), 1);
assert!(matches!(result[0], Instr::Add { dst: 6, .. }));
}
#[test]
fn dead_code_elimination_chain() {
let instrs = vec![
Instr::Add {
dst: 3,
srcs: vec![0, 1],
},
Instr::Mul {
dst: 4,
srcs: vec![3, 2],
},
Instr::Copy { dst: 5, src: 4 },
];
let result = eliminate_dead_code(instrs);
assert_eq!(result.len(), 3); }
#[test]
fn dead_code_elimination_unused_branch() {
let instrs = vec![
Instr::Add {
dst: 3,
srcs: vec![0, 1],
}, Instr::Mul {
dst: 4,
srcs: vec![0, 2],
}, Instr::Copy { dst: 5, src: 0 }, ];
let result = eliminate_dead_code(instrs);
assert_eq!(result.len(), 1);
assert!(matches!(result[0], Instr::Copy { dst: 5, src: 0 }));
}
#[test]
fn external_fun_not_eliminated() {
let instrs = vec![
Instr::ExternalFun {
dst: 3,
fn_idx: 0,
srcs: vec![0],
},
Instr::Copy { dst: 5, src: 3 },
];
let result = eliminate_dead_code(instrs);
assert_eq!(result.len(), 2);
}
}