use std::collections::BTreeSet;
use ocas_core::FastHashMap as HashMap;
use crate::instruction::Instr;
pub fn optimize(
mut instructions: Vec<Instr>,
temp_base: usize,
temp_count: usize,
live_roots: &[usize],
) -> (Vec<Instr>, usize, Vec<usize>) {
instructions = simplify(instructions);
instructions = cse(instructions);
instructions = eliminate_dead_code(instructions, live_roots);
compact_stack(instructions, temp_base, temp_count, live_roots)
}
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::default();
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>, live_roots: &[usize]) -> Vec<Instr> {
if instructions.is_empty() {
return instructions;
}
use std::collections::HashSet;
let mut live_slots: HashSet<usize> = live_roots.iter().copied().collect();
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 compact_stack(
instructions: Vec<Instr>,
temp_base: usize,
_temp_count: usize,
live_roots: &[usize],
) -> (Vec<Instr>, usize, Vec<usize>) {
let mut temps: BTreeSet<usize> = BTreeSet::new();
for instr in &instructions {
if instr.dst() >= temp_base {
temps.insert(instr.dst());
}
for src in instr_srcs(instr) {
if src >= temp_base {
temps.insert(src);
}
}
}
for &root in live_roots {
if root >= temp_base {
temps.insert(root);
}
}
let map: HashMap<usize, usize> = temps
.iter()
.enumerate()
.map(|(i, &slot)| (slot, temp_base + i))
.collect();
let remap = |slot: usize| -> usize { if slot >= temp_base { map[&slot] } else { slot } };
let new_instructions: Vec<Instr> = instructions
.into_iter()
.map(|instr| match instr {
Instr::Add { dst, srcs } => Instr::Add {
dst: remap(dst),
srcs: srcs.into_iter().map(remap).collect(),
},
Instr::Mul { dst, srcs } => Instr::Mul {
dst: remap(dst),
srcs: srcs.into_iter().map(remap).collect(),
},
Instr::Pow { dst, base, exp } => Instr::Pow {
dst: remap(dst),
base: remap(base),
exp,
},
Instr::Powf { dst, base, exp } => Instr::Powf {
dst: remap(dst),
base: remap(base),
exp: remap(exp),
},
Instr::BuiltinOp { dst, op, src } => Instr::BuiltinOp {
dst: remap(dst),
op,
src: remap(src),
},
Instr::ExternalFun { dst, fn_idx, srcs } => Instr::ExternalFun {
dst: remap(dst),
fn_idx,
srcs: srcs.into_iter().map(remap).collect(),
},
Instr::Copy { dst, src } => Instr::Copy {
dst: remap(dst),
src: remap(src),
},
})
.collect();
let new_roots: Vec<usize> = live_roots.iter().map(|&r| remap(r)).collect();
(new_instructions, temps.len(), new_roots)
}
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, &[6]);
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, &[5]);
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, &[5]);
assert_eq!(result.len(), 1);
assert!(matches!(result[0], Instr::Copy { dst: 5, src: 0 }));
}
#[test]
fn dead_code_elimination_multi_root() {
let instrs = vec![
Instr::Add {
dst: 3,
srcs: vec![0, 1],
},
Instr::Mul {
dst: 4,
srcs: vec![0, 1],
},
];
let result = eliminate_dead_code(instrs, &[3, 4]);
assert_eq!(result.len(), 2);
}
#[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, &[5]);
assert_eq!(result.len(), 2);
}
#[test]
fn compact_stack_removes_holes() {
let instrs = vec![
Instr::Add {
dst: 5,
srcs: vec![0, 1],
},
Instr::Mul {
dst: 9,
srcs: vec![5, 5],
},
];
let (result, temp_count, roots) = compact_stack(instrs, 2, 10, &[9]);
assert_eq!(temp_count, 2);
assert_eq!(roots, vec![3]);
assert!(matches!(result[0], Instr::Add { dst: 2, .. }));
assert!(matches!(result[1], Instr::Mul { dst: 3, ref srcs } if srcs == &vec![2, 2]));
}
#[test]
fn compact_stack_preserves_params_and_consts() {
let instrs = vec![Instr::Copy { dst: 4, src: 1 }];
let (result, _, roots) = compact_stack(instrs, 2, 3, &[4]);
assert!(matches!(result[0], Instr::Copy { dst: 2, src: 1 }));
assert_eq!(roots, vec![2]);
}
#[test]
fn optimize_full_pipeline() {
let instrs = vec![
Instr::Add {
dst: 2,
srcs: vec![0, 1],
},
Instr::Add {
dst: 3,
srcs: vec![0, 1],
}, Instr::Copy { dst: 4, src: 3 }, Instr::Mul {
dst: 5,
srcs: vec![0, 0],
}, ];
let (result, temp_count, roots) = optimize(instrs, 2, 4, &[4]);
assert_eq!(roots.len(), 1);
assert!(temp_count <= 3);
for instr in &result {
for src in instr_srcs(instr) {
assert!(src < 2 + temp_count, "src {src} out of compacted range");
}
assert!(instr.dst() < 2 + temp_count);
}
}
}