use std::collections::BTreeMap;
use acir::{
AcirField,
circuit::{Circuit, Opcode, brillig::BrilligFunctionId},
};
use itertools::Itertools;
mod common_subexpression;
mod general;
mod redundant_range;
mod unused_memory;
pub(crate) use general::GeneralOptimizer;
pub(crate) use redundant_range::RangeOptimizer;
use tracing::info;
use self::unused_memory::UnusedMemoryOptimizer;
use super::{AcirTransformationMap, transform_assert_messages};
pub fn optimize<F: AcirField>(
acir: Circuit<F>,
brillig_side_effects: &BTreeMap<BrilligFunctionId, bool>,
) -> (Circuit<F>, AcirTransformationMap) {
let acir_opcode_positions = (0..acir.opcodes.len()).collect();
let (mut acir, new_opcode_positions) =
optimize_internal(acir, acir_opcode_positions, brillig_side_effects);
let transformation_map = AcirTransformationMap::new(&new_opcode_positions);
acir.assert_messages = transform_assert_messages(acir.assert_messages, &transformation_map);
(acir, transformation_map)
}
#[tracing::instrument(level = "trace", name = "optimize_acir" skip(acir, acir_opcode_positions))]
pub(super) fn optimize_internal<F: AcirField>(
acir: Circuit<F>,
acir_opcode_positions: Vec<usize>,
brillig_side_effects: &BTreeMap<BrilligFunctionId, bool>,
) -> (Circuit<F>, Vec<usize>) {
if acir.opcodes.len() == 1 && matches!(acir.opcodes[0], Opcode::BrilligCall { .. }) {
info!("Program is fully unconstrained, skipping optimization pass");
return (acir, acir_opcode_positions);
}
info!("Number of opcodes before: {}", acir.opcodes.len());
let (opcodes, acir_opcode_positions): (Vec<_>, Vec<_>) = acir
.opcodes
.into_iter()
.zip_eq(acir_opcode_positions)
.filter_map(|(opcode, position)| {
if let Opcode::AssertZero(arith_expr) = opcode {
let optimized = GeneralOptimizer::optimize(arith_expr);
if optimized.is_zero() {
return None;
}
Some((Opcode::AssertZero(optimized), position))
} else {
Some((opcode, position))
}
})
.unzip();
let acir = Circuit { opcodes, ..acir };
let memory_optimizer = UnusedMemoryOptimizer::new(acir);
let (acir, acir_opcode_positions) =
memory_optimizer.remove_unused_memory_initializations(acir_opcode_positions);
let range_optimizer = RangeOptimizer::new(acir, brillig_side_effects);
let (acir, acir_opcode_positions) =
range_optimizer.replace_redundant_ranges(acir_opcode_positions);
let max_transformer_passes_or_default = None;
let (acir, acir_opcode_positions, _opcodes_hash_stabilized) =
common_subexpression::transform_internal(
acir,
acir_opcode_positions,
brillig_side_effects,
max_transformer_passes_or_default,
);
info!("Number of opcodes after: {}", acir.opcodes.len());
(acir, acir_opcode_positions)
}
#[cfg(test)]
mod tests {
use acir::{FieldElement, circuit::Circuit};
use std::collections::BTreeMap;
use crate::{assert_circuit_snapshot, compiler::optimizers::optimize_internal};
#[test]
fn removes_empty_assert_zero_opcodes() {
let src = "
private parameters: [w0, w1]
public parameters: []
return values: []
ASSERT w0*w1 - w1*w0 = 0
";
let circuit = Circuit::<FieldElement>::from_str(src).unwrap();
let acir_opcode_positions = (0..circuit.opcodes.len()).collect();
let (optimized, _) = optimize_internal(circuit, acir_opcode_positions, &BTreeMap::new());
assert_circuit_snapshot!(optimized, @r"
private parameters: [w0, w1]
public parameters: []
return values: []
");
}
}