use crate::StencilTable;
use crate::csr::CsrVec;
#[derive(Debug, Clone)]
pub struct InverseStencilMap {
transpose: CsrVec,
output_count: usize,
}
impl<'a> From<&'a StencilTable> for InverseStencilMap {
fn from(table: &'a StencilTable) -> Self {
let output_count = table.output_count();
let input_count = table
.indices
.iter()
.copied()
.max()
.map_or(0, |m| m as usize + 1);
let mut buckets: Vec<Vec<u32>> = vec![Vec::new(); input_count];
for row in 0..output_count {
let start = table.offsets[row] as usize;
let end = table.offsets[row + 1] as usize;
for &input in &table.indices[start..end] {
buckets[input as usize].push(row as u32);
}
}
Self {
transpose: CsrVec::from_jagged_u32(&buckets),
output_count,
}
}
}
impl InverseStencilMap {
#[inline]
pub fn output_count(&self) -> usize {
self.output_count
}
pub fn affected_outputs(&self, changed_inputs: &[u32]) -> Vec<u32> {
let mut seen = vec![false; self.output_count];
let mut affected = Vec::new();
for &input in changed_inputs {
let i = input as usize;
if i < self.transpose.len() {
for &row in self.transpose.row(i) {
let r = row as usize;
if !seen[r] {
seen[r] = true;
affected.push(row);
}
}
}
}
affected.sort_unstable();
affected
}
}
#[derive(Debug, Clone)]
pub struct InverseStencilChain {
levels: Vec<InverseStencilMap>,
}
impl<'a> From<&'a [StencilTable]> for InverseStencilChain {
fn from(level_stencils: &'a [StencilTable]) -> Self {
Self {
levels: level_stencils.iter().map(InverseStencilMap::from).collect(),
}
}
}
impl InverseStencilChain {
#[inline]
pub fn level_count(&self) -> usize {
self.levels.len()
}
pub fn affected_outputs(&self, changed_base_inputs: &[u32]) -> Vec<u32> {
let mut scratch = AffectedScratch::default();
let mut out = Vec::new();
self.affected_outputs_into(changed_base_inputs, &mut scratch, &mut out);
out
}
pub fn affected_outputs_into(
&self,
changed_base_inputs: &[u32],
scratch: &mut AffectedScratch,
out: &mut Vec<u32>,
) {
let AffectedScratch {
stamp,
generation,
front,
back,
} = scratch;
out.clear();
if self.levels.is_empty() {
out.extend_from_slice(changed_base_inputs);
out.sort_unstable();
out.dedup();
return;
}
let max_output = self
.levels
.iter()
.map(|level| level.output_count)
.max()
.unwrap_or(0);
if stamp.len() < max_output {
stamp.resize(max_output, 0);
}
front.clear();
front.extend_from_slice(changed_base_inputs);
for level in &self.levels {
*generation = generation.wrapping_add(1);
if *generation == 0 {
stamp.iter_mut().for_each(|s| *s = 0);
*generation = 1;
}
let g = *generation;
back.clear();
for &input in front.iter() {
let i = input as usize;
if i < level.transpose.len() {
for &row in level.transpose.row(i) {
let r = row as usize;
if stamp[r] != g {
stamp[r] = g;
back.push(row);
}
}
}
}
back.sort_unstable();
std::mem::swap(front, back);
}
std::mem::swap(out, front);
}
}
#[derive(Debug, Clone, Default)]
pub struct AffectedScratch {
stamp: Vec<u32>,
generation: u32,
front: Vec<u32>,
back: Vec<u32>,
}
#[cfg(test)]
mod tests {
use super::*;
fn sample_table() -> StencilTable {
StencilTable {
offsets: vec![0, 1, 3, 4, 6],
indices: vec![0, 0, 1, 1, 1, 2],
weights: vec![1.0, 0.5, 0.5, 1.0, 0.5, 0.5],
}
}
fn brute_force_affected(table: &StencilTable, changed: &[u32]) -> Vec<u32> {
let set: std::collections::HashSet<u32> = changed.iter().copied().collect();
(0..table.output_count() as u32)
.filter(|&r| {
let s = table.offsets[r as usize] as usize;
let e = table.offsets[r as usize + 1] as usize;
table.indices[s..e].iter().any(|i| set.contains(i))
})
.collect()
}
#[test]
fn affected_outputs_for_single_input() {
let inv = InverseStencilMap::from(&sample_table());
assert_eq!(inv.affected_outputs(&[0]), vec![0, 1]);
assert_eq!(inv.affected_outputs(&[1]), vec![1, 2, 3]);
assert_eq!(inv.affected_outputs(&[2]), vec![3]);
}
#[test]
fn affected_outputs_is_sorted_unique_union() {
let inv = InverseStencilMap::from(&sample_table());
assert_eq!(inv.affected_outputs(&[2, 0, 0]), vec![0, 1, 3]);
}
#[test]
fn affected_outputs_matches_brute_force() {
let table = sample_table();
let inv = InverseStencilMap::from(&table);
for changed in [
vec![],
vec![0],
vec![1],
vec![2],
vec![0, 1],
vec![0, 2],
vec![1, 2],
vec![0, 1, 2],
] {
assert_eq!(
inv.affected_outputs(&changed),
brute_force_affected(&table, &changed),
"changed = {changed:?}"
);
}
}
#[test]
fn out_of_range_or_empty_input_affects_nothing() {
let inv = InverseStencilMap::from(&sample_table());
assert!(inv.affected_outputs(&[99]).is_empty());
assert!(inv.affected_outputs(&[]).is_empty());
}
fn two_level_tables() -> (StencilTable, StencilTable) {
let t0 = sample_table(); let t1 = StencilTable {
offsets: vec![0, 1, 3, 4],
indices: vec![0, 1, 2, 3],
weights: vec![1.0, 0.5, 0.5, 1.0],
};
(t0, t1)
}
#[test]
fn chain_affected_outputs_matches_composed_table() {
let (t0, t1) = two_level_tables();
let chain = InverseStencilChain::from([t0.clone(), t1.clone()].as_slice());
let composed = t0.compose(&t1);
for changed in [vec![], vec![0], vec![1], vec![2], vec![0, 2], vec![0, 1, 2]] {
assert_eq!(
chain.affected_outputs(&changed),
brute_force_affected(&composed, &changed),
"changed = {changed:?}"
);
}
}
#[test]
fn chain_single_level_matches_single_map() {
let t0 = sample_table();
let chain = InverseStencilChain::from(std::slice::from_ref(&t0));
let map = InverseStencilMap::from(&t0);
for changed in [vec![0], vec![1], vec![2], vec![0, 1, 2]] {
assert_eq!(
chain.affected_outputs(&changed),
map.affected_outputs(&changed)
);
}
}
#[test]
fn chain_with_no_levels_is_identity() {
let chain = InverseStencilChain::from(&[] as &[StencilTable]);
assert_eq!(chain.affected_outputs(&[2, 0, 0]), vec![0, 2]);
assert!(chain.affected_outputs(&[]).is_empty());
}
#[test]
fn affected_outputs_into_matches_oracle_and_reuses_scratch() {
let (t0, t1) = two_level_tables();
let chain = InverseStencilChain::from([t0.clone(), t1.clone()].as_slice());
let composed = t0.compose(&t1);
let mut scratch = AffectedScratch::default();
let mut out = Vec::new();
for changed in [vec![0u32], vec![2], vec![], vec![0, 1, 2], vec![1]] {
chain.affected_outputs_into(&changed, &mut scratch, &mut out);
assert_eq!(
out,
brute_force_affected(&composed, &changed),
"changed = {changed:?}"
);
}
}
#[test]
fn affected_outputs_into_no_levels_is_identity_and_clears_out() {
let chain = InverseStencilChain::from(&[] as &[StencilTable]);
let mut scratch = AffectedScratch::default();
let mut out = vec![999]; chain.affected_outputs_into(&[2, 0, 0], &mut scratch, &mut out);
assert_eq!(out, vec![0, 2]);
}
}