use crate::HashMap;
use crate::HashSet;
use crate::PassOptions;
use crate::ir::*;
use std::sync::Arc;
use super::block_opt::aggregate_static_offset;
use super::pass_manager::ExecutionUnitPass;
pub(in crate::optimizer) struct CoalesceStoresPass {
pub element_widths: Arc<HashMap<RegionedAbsoluteAddr, usize>>,
pub max_store_width: usize,
}
impl Default for CoalesceStoresPass {
fn default() -> Self {
Self {
element_widths: Arc::default(),
max_store_width: 64,
}
}
}
impl ExecutionUnitPass for CoalesceStoresPass {
fn name(&self) -> &'static str {
"coalesce_stores"
}
fn run(&self, eu: &mut ExecutionUnit<RegionedAbsoluteAddr>, _options: &PassOptions) {
let mut reg_counter = eu.register_map.keys().map(|r| r.0).max().unwrap_or(0);
for block in eu.blocks.values_mut() {
coalesce_block(
block,
&mut eu.register_map,
&mut reg_counter,
&self.element_widths,
self.max_store_width,
);
}
}
}
struct StoreCandidate {
addr: RegionedAbsoluteAddr,
inst_index: usize,
offset: usize,
width: usize,
src_reg: RegisterId,
comb_capture_sites: Vec<u32>,
}
fn coalesce_block(
block: &mut BasicBlock<RegionedAbsoluteAddr>,
register_map: &mut HashMap<RegisterId, RegisterType>,
reg_counter: &mut usize,
element_widths: &HashMap<RegionedAbsoluteAddr, usize>,
max_store_width: usize,
) {
let mut groups: HashMap<RegionedAbsoluteAddr, Vec<StoreCandidate>> = HashMap::default();
let mut sealed_groups: Vec<Vec<StoreCandidate>> = Vec::new();
for (i, inst) in block.instructions.iter().enumerate() {
match inst {
SIRInstruction::Store(addr, SIROffset::Static(off), width, src, triggers, sites)
if triggers.is_empty() && sites.is_empty() =>
{
groups.entry(*addr).or_default().push(StoreCandidate {
addr: *addr,
inst_index: i,
offset: *off,
width: *width,
src_reg: *src,
comb_capture_sites: sites.clone(),
});
}
SIRInstruction::Load(_, addr, _, _) => {
seal_group(&mut groups, addr, &mut sealed_groups);
}
SIRInstruction::Commit(src, _, _, _, _) => {
seal_group(&mut groups, src, &mut sealed_groups);
}
SIRInstruction::Store(
addr,
SIROffset::Dynamic(_)
| SIROffset::Element { .. }
| SIROffset::PackedElements { .. },
_,
_,
_,
_,
) => {
seal_group(&mut groups, addr, &mut sealed_groups);
}
_ => {}
}
}
for (_, candidates) in groups.drain() {
if candidates.len() >= 2 {
sealed_groups.push(candidates);
}
}
struct Replacement {
removed_indices: Vec<usize>,
anchor_index: usize,
new_instructions: Vec<SIRInstruction<RegionedAbsoluteAddr>>,
}
let mut replacements: Vec<Replacement> = Vec::new();
for mut group in sealed_groups {
if group.len() < 2 {
continue;
}
group.sort_by_key(|c| c.offset);
let mut run_start = 0;
while run_start < group.len() {
let mut run_end = run_start;
let mut scan_end = run_start;
let mut expected = group[run_start].offset + group[run_start].width;
while scan_end + 1 < group.len() {
let next = &group[scan_end + 1];
if next.offset != expected {
break;
}
scan_end += 1;
expected += next.width;
let aggregate_width = expected - group[run_start].offset;
if aggregate_width > max_store_width {
break;
}
if aggregate_static_offset(
group[run_start].offset,
aggregate_width,
element_widths.get(&group[run_start].addr).copied(),
)
.is_some()
{
run_end = scan_end;
} else if element_widths
.get(&group[run_start].addr)
.is_some_and(|element_width| {
!group[run_start].offset.is_multiple_of(*element_width)
})
{
break;
}
}
let run_len = run_end - run_start + 1;
if run_len >= 2 {
let sub_run = &group[run_start..=run_end];
let merged_lsb = sub_run[0].offset;
let total_width: usize = sub_run.iter().map(|c| c.width).sum();
debug_assert!(total_width <= max_store_width);
let anchor_index = sub_run.iter().map(|c| c.inst_index).max().unwrap();
let removed_indices: Vec<usize> = sub_run.iter().map(|c| c.inst_index).collect();
*reg_counter += 1;
while register_map.contains_key(&RegisterId(*reg_counter)) {
*reg_counter += 1;
}
let concat_reg = RegisterId(*reg_counter);
register_map.insert(
concat_reg,
RegisterType::Bit {
width: total_width,
signed: false,
},
);
let concat_args: Vec<RegisterId> =
sub_run.iter().rev().map(|c| c.src_reg).collect();
let mut comb_capture_sites = Vec::new();
for site_id in sub_run
.iter()
.flat_map(|c| c.comb_capture_sites.iter().copied())
{
if !comb_capture_sites.contains(&site_id) {
comb_capture_sites.push(site_id);
}
}
let addr = if let SIRInstruction::Store(addr, _, _, _, _, _) =
&block.instructions[sub_run[0].inst_index]
{
*addr
} else {
unreachable!()
};
let Some(offset) = aggregate_static_offset(
merged_lsb,
total_width,
element_widths.get(&addr).copied(),
) else {
run_start = run_end + 1;
continue;
};
let new_instructions = vec![
SIRInstruction::Concat(concat_reg, concat_args),
SIRInstruction::Store(
addr,
offset,
total_width,
concat_reg,
Vec::new(),
comb_capture_sites,
),
];
replacements.push(Replacement {
removed_indices,
anchor_index,
new_instructions,
});
}
run_start = run_end + 1;
}
}
if replacements.is_empty() {
return;
}
let mut removed_set: HashSet<usize> = HashSet::default();
let mut insert_map: HashMap<usize, Vec<SIRInstruction<RegionedAbsoluteAddr>>> =
HashMap::default();
for repl in replacements {
for &idx in &repl.removed_indices {
removed_set.insert(idx);
}
insert_map
.entry(repl.anchor_index)
.or_default()
.extend(repl.new_instructions);
}
let old_instructions = std::mem::take(&mut block.instructions);
let mut new_instructions = Vec::with_capacity(old_instructions.len());
for (i, inst) in old_instructions.into_iter().enumerate() {
if let Some(replacement) = insert_map.remove(&i) {
new_instructions.extend(replacement);
} else if !removed_set.contains(&i) {
new_instructions.push(inst);
}
}
block.instructions = new_instructions;
}
fn seal_group(
groups: &mut HashMap<RegionedAbsoluteAddr, Vec<StoreCandidate>>,
addr: &RegionedAbsoluteAddr,
sealed_groups: &mut Vec<Vec<StoreCandidate>>,
) {
if let Some(candidates) = groups.remove(addr) {
if candidates.len() >= 2 {
sealed_groups.push(candidates);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn make_var_id(n: u32) -> celox_design::StateObjectId {
celox_design::StateObjectId(n)
}
fn make_addr(var_id: u32) -> RegionedAbsoluteAddr {
RegionedAbsoluteAddr {
region: 0,
instance_id: InstanceId(0),
var_id: make_var_id(var_id),
}
}
fn make_eu(
instructions: Vec<SIRInstruction<RegionedAbsoluteAddr>>,
register_map: HashMap<RegisterId, RegisterType>,
) -> ExecutionUnit<RegionedAbsoluteAddr> {
let mut blocks = HashMap::default();
blocks.insert(
BlockId(0),
BasicBlock {
id: BlockId(0),
params: Vec::new(),
instructions,
terminator: SIRTerminator::Return,
},
);
ExecutionUnit {
entry_block_id: BlockId(0),
blocks,
register_map,
}
}
#[test]
fn test_basic_contiguous_coalescing() {
let addr = make_addr(0);
let mut register_map = HashMap::default();
for i in 0..3 {
register_map.insert(
RegisterId(i),
RegisterType::Bit {
width: 1,
signed: false,
},
);
}
let instructions = vec![
SIRInstruction::Store(
addr,
SIROffset::Static(0),
1,
RegisterId(0),
Vec::new(),
Vec::new(),
),
SIRInstruction::Store(
addr,
SIROffset::Static(1),
1,
RegisterId(1),
Vec::new(),
Vec::new(),
),
SIRInstruction::Store(
addr,
SIROffset::Static(2),
1,
RegisterId(2),
Vec::new(),
Vec::new(),
),
];
let mut eu = make_eu(instructions, register_map);
let options = PassOptions::default();
CoalesceStoresPass::default().run(&mut eu, &options);
let block = eu.blocks.get(&BlockId(0)).unwrap();
assert_eq!(block.instructions.len(), 2);
match &block.instructions[0] {
SIRInstruction::Concat(dst, args) => {
assert_eq!(args.len(), 3);
assert_eq!(args[0], RegisterId(2));
assert_eq!(args[1], RegisterId(1));
assert_eq!(args[2], RegisterId(0));
let reg_type = eu.register_map.get(dst).unwrap();
assert_eq!(reg_type.width(), 3);
}
other => panic!("Expected Concat, got {:?}", other),
}
match &block.instructions[1] {
SIRInstruction::Store(_, SIROffset::Static(0), 3, _, triggers, _) => {
assert!(triggers.is_empty());
}
other => panic!("Expected wide Store, got {:?}", other),
}
}
#[test]
fn partitions_a_long_run_at_the_configured_memory_width() {
let addr = make_addr(0);
let register_map = (0..4)
.map(|index| {
(
RegisterId(index),
RegisterType::Bit {
width: 64,
signed: false,
},
)
})
.collect();
let instructions = (0..4)
.map(|index| {
SIRInstruction::Store(
addr,
SIROffset::Static(index * 64),
64,
RegisterId(index),
vec![],
vec![],
)
})
.collect();
let mut eu = make_eu(instructions, register_map);
CoalesceStoresPass {
element_widths: Arc::default(),
max_store_width: 128,
}
.run(&mut eu, &PassOptions::default());
let widths = eu.blocks[&BlockId(0)]
.instructions
.iter()
.filter_map(|instruction| match instruction {
SIRInstruction::Store(_, _, width, ..) => Some(*width),
_ => None,
})
.collect::<Vec<_>>();
assert_eq!(widths, vec![128, 128]);
}
#[test]
fn test_interleaved_stores_different_vars() {
let addr0 = make_addr(0);
let addr1 = make_addr(1);
let mut register_map = HashMap::default();
for i in 0..4 {
register_map.insert(
RegisterId(i),
RegisterType::Bit {
width: 1,
signed: false,
},
);
}
let instructions = vec![
SIRInstruction::Store(
addr0,
SIROffset::Static(0),
1,
RegisterId(0),
Vec::new(),
Vec::new(),
),
SIRInstruction::Store(
addr1,
SIROffset::Static(0),
1,
RegisterId(2),
Vec::new(),
Vec::new(),
),
SIRInstruction::Store(
addr0,
SIROffset::Static(1),
1,
RegisterId(1),
Vec::new(),
Vec::new(),
),
SIRInstruction::Store(
addr1,
SIROffset::Static(1),
1,
RegisterId(3),
Vec::new(),
Vec::new(),
),
];
let mut eu = make_eu(instructions, register_map);
let options = PassOptions::default();
CoalesceStoresPass::default().run(&mut eu, &options);
let block = eu.blocks.get(&BlockId(0)).unwrap();
assert_eq!(block.instructions.len(), 4);
}
#[test]
fn test_seal_on_load() {
let addr = make_addr(0);
let mut register_map = HashMap::default();
for i in 0..3 {
register_map.insert(
RegisterId(i),
RegisterType::Bit {
width: 1,
signed: false,
},
);
}
let instructions = vec![
SIRInstruction::Store(
addr,
SIROffset::Static(0),
1,
RegisterId(0),
Vec::new(),
Vec::new(),
),
SIRInstruction::Load(RegisterId(2), addr, SIROffset::Static(0), 1),
SIRInstruction::Store(
addr,
SIROffset::Static(1),
1,
RegisterId(1),
Vec::new(),
Vec::new(),
),
];
let mut eu = make_eu(instructions, register_map);
let options = PassOptions::default();
CoalesceStoresPass::default().run(&mut eu, &options);
let block = eu.blocks.get(&BlockId(0)).unwrap();
assert_eq!(block.instructions.len(), 3);
}
#[test]
fn test_non_contiguous_stores_not_coalesced() {
let addr = make_addr(0);
let mut register_map = HashMap::default();
for i in 0..2 {
register_map.insert(
RegisterId(i),
RegisterType::Bit {
width: 1,
signed: false,
},
);
}
let instructions = vec![
SIRInstruction::Store(
addr,
SIROffset::Static(0),
1,
RegisterId(0),
Vec::new(),
Vec::new(),
),
SIRInstruction::Store(
addr,
SIROffset::Static(2),
1,
RegisterId(1),
Vec::new(),
Vec::new(),
),
];
let mut eu = make_eu(instructions, register_map);
let options = PassOptions::default();
CoalesceStoresPass::default().run(&mut eu, &options);
let block = eu.blocks.get(&BlockId(0)).unwrap();
assert_eq!(block.instructions.len(), 2);
}
#[test]
fn test_stores_with_triggers_not_coalesced() {
let addr = make_addr(0);
let mut register_map = HashMap::default();
for i in 0..2 {
register_map.insert(
RegisterId(i),
RegisterType::Bit {
width: 1,
signed: false,
},
);
}
let trigger = TriggerIdWithKind {
id: 0,
kind: DomainKind::ClockPosedge,
};
let instructions = vec![
SIRInstruction::Store(
addr,
SIROffset::Static(0),
1,
RegisterId(0),
vec![trigger],
Vec::new(),
),
SIRInstruction::Store(
addr,
SIROffset::Static(1),
1,
RegisterId(1),
vec![trigger],
Vec::new(),
),
];
let mut eu = make_eu(instructions, register_map);
let options = PassOptions::default();
CoalesceStoresPass::default().run(&mut eu, &options);
let block = eu.blocks.get(&BlockId(0)).unwrap();
assert_eq!(block.instructions.len(), 2);
}
#[test]
fn test_partial_contiguous_run() {
let addr = make_addr(0);
let mut register_map = HashMap::default();
for i in 0..4 {
register_map.insert(
RegisterId(i),
RegisterType::Bit {
width: 1,
signed: false,
},
);
}
let instructions = vec![
SIRInstruction::Store(
addr,
SIROffset::Static(0),
1,
RegisterId(0),
Vec::new(),
Vec::new(),
),
SIRInstruction::Store(
addr,
SIROffset::Static(1),
1,
RegisterId(1),
Vec::new(),
Vec::new(),
),
SIRInstruction::Store(
addr,
SIROffset::Static(2),
1,
RegisterId(2),
Vec::new(),
Vec::new(),
),
SIRInstruction::Store(
addr,
SIROffset::Static(5),
1,
RegisterId(3),
Vec::new(),
Vec::new(),
),
];
let mut eu = make_eu(instructions, register_map);
let options = PassOptions::default();
CoalesceStoresPass::default().run(&mut eu, &options);
let block = eu.blocks.get(&BlockId(0)).unwrap();
assert_eq!(block.instructions.len(), 3);
}
#[test]
fn test_unpacked_cross_element_run_uses_packed_elements() {
let addr = make_addr(0);
let register_map = (0..4)
.map(|index| {
(
RegisterId(index),
RegisterType::Bit {
width: 6,
signed: false,
},
)
})
.collect();
let instructions = (0..4)
.map(|index| {
SIRInstruction::Store(
addr,
SIROffset::Static(index * 6),
6,
RegisterId(index),
vec![],
vec![],
)
})
.collect();
let mut eu = make_eu(instructions, register_map);
let pass = CoalesceStoresPass {
element_widths: Arc::new([(addr, 12usize)].into_iter().collect()),
max_store_width: 64,
};
pass.run(&mut eu, &PassOptions::default());
let stores = eu.blocks[&BlockId(0)]
.instructions
.iter()
.filter_map(|instruction| match instruction {
SIRInstruction::Store(
_,
SIROffset::PackedElements {
bit_offset,
element_width,
},
width,
..,
) => Some((*bit_offset, *element_width, *width)),
_ => None,
})
.collect::<Vec<_>>();
assert_eq!(stores, vec![(0, 12, 24)]);
}
}