use std::fmt::Debug;
use std::{collections::BTreeSet, hash::Hash};
use crate::{
HashMap, InstancePath, SourceAddr as AbsoluteAddr, SourceVarId,
symbolic::artifact::{
RelocationModule, SimModule, SymbolicGlueAddr as GlueAddr, SymbolicGlueBlock as GlueBlock,
},
};
use celox_design::{BinaryOp, BitAccess, InstanceId, UnaryOp, VarAtomBase};
use celox_slt::{
CombObserver, LogicPath, LogicPathTarget, NodeId, SLTNode, SLTNodeArena, SLTNodeFactsError,
get_width,
};
pub struct FlattenedModule {
pub relocation: RelocationModule,
pub pre_atomized_comb_blocks: Vec<LogicPath<AbsoluteAddr>>,
}
pub fn flatten_module(
module: &SimModule,
path: &InstancePath,
instance_ids: &HashMap<InstancePath, InstanceId>,
global_boundaries: &HashMap<AbsoluteAddr, BTreeSet<usize>>,
unpacked_element_widths: &HashMap<AbsoluteAddr, usize>,
arena: &mut SLTNodeArena<AbsoluteAddr>,
) -> Result<FlattenedModule, SLTNodeFactsError> {
let instance_id = instance_ids[path];
let cv = &|id: &SourceVarId| AbsoluteAddr {
instance_id,
var_id: *id,
};
let mut comb_cache = HashMap::default();
let mut comb_blocks: Vec<_> = module
.comb_blocks
.iter()
.map(|e| convert_logic_path(e, &module.arena, arena, &mut comb_cache, &cv))
.collect::<Result<_, _>>()?;
let mut observer_cache = HashMap::default();
let comb_observers: Vec<_> = module
.comb_observers
.iter()
.map(|observer| {
convert_comb_observer(observer, &module.arena, arena, &mut observer_cache, &cv)
})
.collect::<Result<_, _>>()?;
for (child_instance_name, gbs) in &module.glue_blocks {
for (idx, gb) in gbs.iter().enumerate() {
let mut glue_cache = HashMap::default();
let mut child_path = path.0.clone();
child_path.push((child_instance_name.clone(), idx));
let child_id = instance_ids[&InstancePath(child_path)];
comb_blocks.extend(convert_glue_block(
gb,
instance_id,
child_id,
&gb.arena,
arena,
&mut glue_cache,
)?);
}
}
let atomized_comb_blocks = atomize_logic_paths(
&comb_blocks,
global_boundaries,
unpacked_element_widths,
arena,
)?;
Ok(FlattenedModule {
pre_atomized_comb_blocks: comb_blocks,
relocation: RelocationModule {
eval_apply_ff_blocks: HashMap::default(),
eval_only_ff_blocks: HashMap::default(),
apply_ff_blocks: HashMap::default(),
comb_blocks: atomized_comb_blocks,
comb_observers,
},
})
}
fn atomize_logic_paths(
paths: &Vec<LogicPath<AbsoluteAddr>>,
boundaries: &HashMap<AbsoluteAddr, BTreeSet<usize>>,
unpacked_element_widths: &HashMap<AbsoluteAddr, usize>,
arena: &mut SLTNodeArena<AbsoluteAddr>,
) -> Result<Vec<LogicPath<AbsoluteAddr>>, SLTNodeFactsError> {
let mut atomized_paths = Vec::new();
for path in paths {
let Some(target_var) = path.target.var() else {
atomized_paths.push(path.clone());
continue;
};
let element_width = unpacked_element_widths.get(&target_var.id).copied();
let mut effective_boundaries = boundaries.get(&target_var.id).cloned().unwrap_or_default();
if let Some(element_width) = element_width {
let mut boundary =
(target_var.access.lsb / element_width + 1).saturating_mul(element_width);
while boundary <= target_var.access.msb {
effective_boundaries.insert(boundary);
let Some(next) = boundary.checked_add(element_width) else {
break;
};
boundary = next;
}
}
if !effective_boundaries.is_empty() {
let atoms = target_var.access.calculate_atoms(&effective_boundaries);
let original_source_ids: crate::HashSet<_> =
path.sources.iter().map(|s| s.id).collect();
let original_previous_sources = path.previous_sources.clone();
let mut atom_infos: Vec<(
BitAccess,
crate::HashSet<VarAtomBase<AbsoluteAddr>>,
crate::HashSet<AbsoluteAddr>,
)> = Vec::new();
for atom_access in &atoms {
let relative_atom_access = BitAccess::new(
atom_access.lsb - target_var.access.lsb,
atom_access.msb - target_var.access.lsb,
);
let new_expr = project_logic_path_expr(path.expr, relative_atom_access, arena)?;
let mut expr_inputs = crate::HashSet::default();
collect_inputs(new_expr, arena, &mut expr_inputs);
let filtered_sources: crate::HashSet<_> = expr_inputs
.into_iter()
.filter(|input_atom| original_source_ids.contains(&input_atom.id))
.collect();
let filtered_source_ids = filtered_sources.iter().map(|source| source.id).collect();
atom_infos.push((*atom_access, filtered_sources, filtered_source_ids));
}
let mut i = 0;
while i < atom_infos.len() {
let group_start = i;
while i + 1 < atom_infos.len() {
let current = &atom_infos[i];
let next = &atom_infos[i + 1];
let exact_sources_match = next.1 == current.1;
let source_objects_match = next.2 == current.2;
let pointwise_single_bits =
current.0.lsb == current.0.msb && next.0.lsb == next.0.msb;
let contiguous_unpacked_elements =
element_width.is_some_and(|width| width.is_multiple_of(8));
let crosses_strided_element = element_width.is_some_and(|width| {
!width.is_multiple_of(8) && next.0.lsb.is_multiple_of(width)
});
let may_recover_coarse_range = source_objects_match
&& (pointwise_single_bits || contiguous_unpacked_elements);
if !(exact_sources_match || may_recover_coarse_range) || crosses_strided_element
{
break;
}
i += 1;
}
let group_end = i;
i += 1;
let merged_lsb = atom_infos[group_start].0.lsb;
let merged_msb = atom_infos[group_end].0.msb;
let relative_access = BitAccess::new(
merged_lsb - target_var.access.lsb,
merged_msb - target_var.access.lsb,
);
let merged_width = merged_msb - merged_lsb + 1;
let original_width = target_var.access.msb - target_var.access.lsb + 1;
let merged_expr = if merged_width == original_width {
path.expr
} else {
project_logic_path_expr(path.expr, relative_access, arena)?
};
let mut merged_sources = crate::HashSet::default();
collect_inputs(merged_expr, arena, &mut merged_sources);
let filtered_sources: crate::HashSet<_> = merged_sources
.iter()
.copied()
.filter(|input_atom| original_source_ids.contains(&input_atom.id))
.collect();
let filtered_previous_sources: crate::HashSet<_> = merged_sources
.iter()
.copied()
.filter(|input_atom| {
original_previous_sources.iter().any(|previous| {
previous.id == input_atom.id
&& previous.access.overlaps(&input_atom.access)
})
})
.collect();
let filtered_address_sources: crate::HashSet<_> = merged_sources
.into_iter()
.filter(|input_atom| {
path.address_sources.iter().any(|address| {
address.id == input_atom.id
&& address.access.overlaps(&input_atom.access)
})
})
.collect();
let target = VarAtomBase::new(target_var.id, merged_lsb, merged_msb);
atomized_paths.push(LogicPath {
target: LogicPathTarget::Var(target),
sources: filtered_sources,
previous_sources: filtered_previous_sources,
address_sources: filtered_address_sources,
local_inputs: path.local_inputs.clone(),
order_before: path.order_before.clone(),
comb_capture_enable_sites: path.comb_capture_enable_sites.clone(),
comb_capture_enable_always: path.comb_capture_enable_always,
pre_lower_nodes: path.pre_lower_nodes.clone(),
expr: merged_expr,
});
}
} else {
atomized_paths.push(path.clone());
}
}
Ok(atomized_paths)
}
fn project_logic_path_expr(
expression: NodeId,
access: BitAccess,
arena: &mut SLTNodeArena<AbsoluteAddr>,
) -> Result<NodeId, SLTNodeFactsError> {
match arena.get(expression).clone() {
SLTNode::Input {
variable,
signed,
index,
access: input_access,
} if access.msb <= input_access.msb - input_access.lsb => arena.alloc(SLTNode::Input {
variable,
signed,
index,
access: BitAccess::new(input_access.lsb + access.lsb, input_access.lsb + access.msb),
}),
SLTNode::Slice {
expr: inner,
access: inner_access,
} if access.msb <= inner_access.msb - inner_access.lsb => project_logic_path_expr(
inner,
BitAccess::new(inner_access.lsb + access.lsb, inner_access.lsb + access.msb),
arena,
),
_ => arena.alloc(SLTNode::Slice {
expr: expression,
access,
}),
}
}
pub fn collect_inputs<A: Hash + Eq + Clone + Debug>(
expr: NodeId,
arena: &SLTNodeArena<A>,
set: &mut crate::HashSet<VarAtomBase<A>>,
) {
let mut visited = HashMap::default();
collect_inputs_with_window(expr, None, arena, set, &mut visited);
}
fn collect_inputs_with_window<A: Hash + Eq + Clone + Debug>(
expr: NodeId,
window: Option<BitAccess>,
arena: &SLTNodeArena<A>,
set: &mut crate::HashSet<VarAtomBase<A>>,
visited: &mut HashMap<NodeId, Vec<BitAccess>>,
) {
let requested = window.unwrap_or_else(|| BitAccess::new(0, get_width(expr, arena) - 1));
let uncovered = claim_uncovered_window(visited.entry(expr).or_default(), requested);
for window in uncovered.into_iter().map(Some) {
match arena.get(expr) {
SLTNode::Input {
variable,
access,
index,
..
} => {
if !index.is_empty() {
let element_width = get_width(expr, arena);
let full_width = access.msb - access.lsb + 1;
let mut max_reachable_elements = 1usize;
for idx in index {
let idx_width = get_width(idx.node, arena);
let reachable = 1usize.checked_shl(idx_width as u32).unwrap_or(usize::MAX);
max_reachable_elements = max_reachable_elements.saturating_mul(reachable);
}
let actual_elements = full_width / element_width;
let effective_elements = std::cmp::min(max_reachable_elements, actual_elements);
let reachable_lsb = access.lsb;
let reachable_msb = access.lsb + (effective_elements * element_width) - 1;
set.insert(VarAtomBase::new(
variable.clone(),
reachable_lsb,
std::cmp::min(reachable_msb, access.msb),
));
} else {
let full_width = access.msb - access.lsb + 1;
let win = window.unwrap_or(BitAccess::new(0, full_width - 1));
set.insert(VarAtomBase::new(
variable.clone(),
access.lsb + win.lsb,
access.lsb + win.msb,
));
}
for idx in index {
collect_inputs_with_window(idx.node, None, arena, set, visited);
}
}
SLTNode::Slice { expr, access } => {
let composed = if let Some(win) = window {
BitAccess::new(access.lsb + win.lsb, access.lsb + win.msb)
} else {
*access
};
collect_inputs_with_window(*expr, Some(composed), arena, set, visited)
}
SLTNode::Concat(parts) => {
if let Some(win) = window {
let mut part_lsb = 0usize;
for (part, width) in parts.iter().rev() {
let part_msb = part_lsb + width - 1;
if win.overlaps(&BitAccess::new(part_lsb, part_msb)) {
let ov_lsb = std::cmp::max(win.lsb, part_lsb);
let ov_msb = std::cmp::min(win.msb, part_msb);
let local = BitAccess::new(ov_lsb - part_lsb, ov_msb - part_lsb);
collect_inputs_with_window(*part, Some(local), arena, set, visited);
}
part_lsb += width;
}
} else {
for (part, _) in parts {
collect_inputs_with_window(*part, None, arena, set, visited);
}
}
}
SLTNode::Binary(lhs, op, rhs) => {
let pointwise = matches!(op, BinaryOp::And | BinaryOp::Or | BinaryOp::Xor);
let lhs_window = pointwise
.then(|| dependency_window(window, *lhs, arena))
.flatten();
let rhs_window = pointwise
.then(|| dependency_window(window, *rhs, arena))
.flatten();
collect_inputs_with_window(*lhs, lhs_window, arena, set, visited);
collect_inputs_with_window(*rhs, rhs_window, arena, set, visited);
}
SLTNode::Unary(op, inner) => {
let pointwise =
matches!(op, UnaryOp::Ident | UnaryOp::ToTwoState | UnaryOp::BitNot);
let inner_window = pointwise
.then(|| dependency_window(window, *inner, arena))
.flatten();
collect_inputs_with_window(*inner, inner_window, arena, set, visited);
}
SLTNode::Capture { expr, .. } => {
let inner_window = dependency_window(window, *expr, arena);
collect_inputs_with_window(*expr, inner_window, arena, set, visited);
}
SLTNode::Mux {
cond,
then_expr,
else_expr,
} => {
collect_inputs_with_window(*cond, None, arena, set, visited);
let then_window = dependency_window(window, *then_expr, arena);
let else_window = dependency_window(window, *else_expr, arena);
collect_inputs_with_window(*then_expr, then_window, arena, set, visited);
collect_inputs_with_window(*else_expr, else_window, arena, set, visited);
}
SLTNode::ForFold {
loop_var,
start,
end,
result,
initials,
updates,
effects,
continue_cond,
..
} => {
if let celox_slt::SLTLoopBound::Expr(node) = start {
collect_inputs_with_window(*node, None, arena, set, visited);
}
if let celox_slt::SLTLoopBound::Expr(node) = end {
collect_inputs_with_window(*node, None, arena, set, visited);
}
if let celox_slt::SLTForFoldResult::Transient { initial, update } = result {
collect_inputs_with_window(*initial, None, arena, set, visited);
collect_inputs_with_window(*update, None, arena, set, visited);
}
for init in initials {
collect_inputs_with_window(init.expr, None, arena, set, visited);
}
for update in updates {
collect_inputs_with_window(update.expr, None, arena, set, visited);
}
for effect in effects {
match effect {
celox_slt::SLTForEffect::Event { guard, args, .. } => {
if let Some(guard) = guard {
collect_inputs_with_window(*guard, None, arena, set, visited);
}
for arg in args {
collect_inputs_with_window(*arg, None, arena, set, visited);
}
}
celox_slt::SLTForEffect::Runner(runner) => {
collect_inputs_with_window(*runner, None, arena, set, visited);
}
}
}
collect_inputs_with_window(*continue_cond, None, arena, set, visited);
set.retain(|atom| atom.id != *loop_var);
}
SLTNode::ForFoldGroup {
loop_var,
entry_guard,
states,
..
} => {
let mut group_inputs = crate::HashSet::default();
let mut group_visited = HashMap::default();
collect_inputs_with_window(
*entry_guard,
None,
arena,
&mut group_inputs,
&mut group_visited,
);
for state in states {
collect_inputs_with_window(
state.initial,
None,
arena,
&mut group_inputs,
&mut group_visited,
);
}
let mut update_inputs = crate::HashSet::default();
let mut update_visited = HashMap::default();
for state in states {
collect_inputs_with_window(
state.update,
None,
arena,
&mut update_inputs,
&mut update_visited,
);
}
update_inputs.retain(|atom| {
atom.id != *loop_var && !carried_states_cover_atom(atom, states)
});
group_inputs.extend(update_inputs);
set.extend(group_inputs);
}
SLTNode::Constant(_, _, _, _) => {}
}
}
}
enum UncoveredWindows {
Empty,
One(BitAccess),
Multiple(Vec<BitAccess>),
}
impl UncoveredWindows {
fn push(&mut self, window: BitAccess) {
match std::mem::replace(self, Self::Empty) {
Self::Empty => *self = Self::One(window),
Self::One(first) => *self = Self::Multiple(vec![first, window]),
Self::Multiple(mut windows) => {
windows.push(window);
*self = Self::Multiple(windows);
}
}
}
fn is_empty(&self) -> bool {
matches!(self, Self::Empty)
}
}
enum UncoveredWindowsIter {
Empty,
One(std::option::IntoIter<BitAccess>),
Multiple(std::vec::IntoIter<BitAccess>),
}
impl Iterator for UncoveredWindowsIter {
type Item = BitAccess;
fn next(&mut self) -> Option<Self::Item> {
match self {
Self::Empty => None,
Self::One(window) => window.next(),
Self::Multiple(windows) => windows.next(),
}
}
}
impl IntoIterator for UncoveredWindows {
type Item = BitAccess;
type IntoIter = UncoveredWindowsIter;
fn into_iter(self) -> Self::IntoIter {
match self {
Self::Empty => UncoveredWindowsIter::Empty,
Self::One(window) => UncoveredWindowsIter::One(Some(window).into_iter()),
Self::Multiple(windows) => UncoveredWindowsIter::Multiple(windows.into_iter()),
}
}
}
fn claim_uncovered_window(covered: &mut Vec<BitAccess>, requested: BitAccess) -> UncoveredWindows {
let mut uncovered = UncoveredWindows::Empty;
let mut cursor = requested.lsb;
for range in covered.iter().copied() {
if range.msb < cursor {
continue;
}
if range.lsb > requested.msb {
break;
}
if range.lsb > cursor {
uncovered.push(BitAccess::new(cursor, requested.msb.min(range.lsb - 1)));
}
cursor = cursor.max(range.msb.saturating_add(1));
if cursor > requested.msb {
break;
}
}
if cursor <= requested.msb {
uncovered.push(BitAccess::new(cursor, requested.msb));
}
if uncovered.is_empty() {
return uncovered;
}
let mut merged = requested;
let start = covered.partition_point(|range| range.msb.saturating_add(1) < merged.lsb);
let mut end = start;
while end < covered.len() && covered[end].lsb <= merged.msb.saturating_add(1) {
merged.lsb = merged.lsb.min(covered[end].lsb);
merged.msb = merged.msb.max(covered[end].msb);
end += 1;
}
covered.splice(start..end, std::iter::once(merged));
uncovered
}
fn dependency_window<A: Hash + Eq + Clone + Debug>(
window: Option<BitAccess>,
operand: NodeId,
arena: &SLTNodeArena<A>,
) -> Option<BitAccess> {
window.filter(|window| window.msb < get_width(operand, arena))
}
fn carried_states_cover_atom<A: Hash + Eq + Clone>(
atom: &VarAtomBase<A>,
states: &[celox_slt::SLTForFoldGroupState<A>],
) -> bool {
let mut ranges = states
.iter()
.filter(|state| state.target.id == atom.id)
.map(|state| state.target.access)
.collect::<Vec<_>>();
ranges.sort_unstable_by_key(|access| (access.lsb, access.msb));
let mut next = atom.access.lsb;
for range in ranges {
if range.msb < next {
continue;
}
if range.lsb > next {
return false;
}
if range.msb >= atom.access.msb {
return true;
}
let Some(after) = range.msb.checked_add(1) else {
return false;
};
next = after;
}
false
}
fn convert_logic_path<
A: Hash + Eq + Clone + std::fmt::Debug + std::fmt::Display,
B: Hash + Eq + Clone,
>(
lp: &LogicPath<A>,
arena: &SLTNodeArena<A>,
target_arena: &mut SLTNodeArena<B>,
cache: &mut HashMap<NodeId, NodeId>,
f: &impl Fn(&A) -> B,
) -> Result<LogicPath<B>, SLTNodeFactsError> {
lp.map_addr(arena, target_arena, cache, f)
}
fn convert_comb_observer<
A: Hash + Eq + Clone + std::fmt::Debug + std::fmt::Display,
B: Hash + Eq + Clone,
>(
observer: &CombObserver<A>,
arena: &SLTNodeArena<A>,
target_arena: &mut SLTNodeArena<B>,
cache: &mut HashMap<NodeId, NodeId>,
f: &impl Fn(&A) -> B,
) -> Result<CombObserver<B>, SLTNodeFactsError> {
let mut map_node = |node| {
arena
.get(node)
.map_addr(node, arena, target_arena, cache, f)
};
Ok(CombObserver {
site_id: observer.site_id,
activation_group: observer.activation_group,
guard: observer.guard.map(&mut map_node).transpose()?,
args: observer
.args
.iter()
.copied()
.map(&mut map_node)
.collect::<Result<_, _>>()?,
loop_runner: observer.loop_runner.map(&mut map_node).transpose()?,
sensitivity: observer
.sensitivity
.iter()
.map(|v| VarAtomBase::new(f(&v.id), v.access.lsb, v.access.msb))
.collect(),
local_inputs: observer
.local_inputs
.iter()
.map(|(id, node)| {
Ok((
f(id),
arena
.get(*node)
.map_addr(*node, arena, target_arena, cache, f)?,
))
})
.collect::<Result<_, SLTNodeFactsError>>()?,
observed_inputs: observer
.observed_inputs
.iter()
.map(|v| VarAtomBase::new(f(&v.id), v.access.lsb, v.access.msb))
.collect(),
position_inputs: observer
.position_inputs
.iter()
.map(|v| VarAtomBase::new(f(&v.id), v.access.lsb, v.access.msb))
.collect(),
preceding_writes: observer
.preceding_writes
.iter()
.map(|v| VarAtomBase::new(f(&v.id), v.access.lsb, v.access.msb))
.collect(),
written_before: observer
.written_before
.iter()
.map(|v| VarAtomBase::new(f(&v.id), v.access.lsb, v.access.msb))
.collect(),
written_input_atoms: observer
.written_input_atoms
.iter()
.map(|v| VarAtomBase::new(f(&v.id), v.access.lsb, v.access.msb))
.collect(),
written_inputs: observer.written_inputs.iter().map(f).collect(),
captured_in_loop: observer.captured_in_loop,
})
}
fn convert_glue_block(
gb: &GlueBlock,
parent_id: InstanceId,
child_id: InstanceId,
arena: &SLTNodeArena<GlueAddr>,
target_arena: &mut SLTNodeArena<AbsoluteAddr>,
cache: &mut HashMap<NodeId, NodeId>,
) -> Result<Vec<LogicPath<AbsoluteAddr>>, SLTNodeFactsError> {
let GlueBlock {
module_id: _,
input_ports,
output_ports,
arena: _,
} = gb;
let cv = &|addr: &GlueAddr| match addr {
GlueAddr::Parent(v) => AbsoluteAddr {
instance_id: parent_id,
var_id: *v,
},
GlueAddr::Child(v) => AbsoluteAddr {
instance_id: child_id,
var_id: *v,
},
};
let mut res = Vec::new();
for (_ports, abb) in input_ports {
res.push(convert_logic_path(abb, arena, target_arena, cache, cv)?);
}
for (_ports, abb) in output_ports {
res.push(convert_logic_path(abb, arena, target_arena, cache, cv)?);
}
Ok(res)
}