use crate::{Image, ImageDescriptor, Node, NodeId, Result};
use std::collections::HashMap;
use std::sync::Arc;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Reduction {
pub from: (u32, u32),
pub to: (u32, u32),
}
impl Reduction {
#[must_use]
pub fn factor(&self) -> f64 {
let before = f64::from(self.from.0) * f64::from(self.from.1);
let after = (f64::from(self.to.0) * f64::from(self.to.1)).max(1.0);
before / after
}
}
pub fn shrink_on_load(image: &Image) -> Result<(Image, Option<Reduction>)> {
let root = Arc::clone(image.node());
let order = topological_order(&root);
let mut rescaled: HashMap<NodeId, Arc<dyn crate::Op>> = HashMap::new();
for node in &order {
let Some(op) = node.op() else { continue };
let Some(copy) = op.rescaled() else {
return Ok((image.clone(), None));
};
rescaled.insert(node.id(), copy);
}
let sources: Vec<&Arc<Node>> = order
.iter()
.filter(|node| node.producer().is_some())
.collect();
let [source] = sources.as_slice() else {
return Ok((image.clone(), None));
};
let Some(producer) = source.producer() else {
return Ok((image.clone(), None));
};
let full = source.descriptor();
let target = (root.descriptor().width, root.descriptor().height);
let Some(reduced) = producer.reduced_descriptor(target) else {
return Ok((image.clone(), None));
};
if (reduced.width, reduced.height) == (full.width, full.height) {
return Ok((image.clone(), None));
}
let Some(simulated) = simulate(&order, source.id(), reduced, &rescaled) else {
return Ok((image.clone(), None));
};
let Some(&new_root) = simulated.get(&root.id()) else {
return Ok((image.clone(), None));
};
if (new_root.width, new_root.height) != (root.descriptor().width, root.descriptor().height) {
return Ok((image.clone(), None));
}
producer.reduce_to(reduced)?;
let rebuilt = rebuild(
&root,
source.id(),
producer,
image,
&rescaled,
&mut HashMap::new(),
)?;
Ok((
rebuilt,
Some(Reduction {
from: (full.width, full.height),
to: (reduced.width, reduced.height),
}),
))
}
fn simulate(
order: &[Arc<Node>],
source: NodeId,
reduced: ImageDescriptor,
rescaled: &HashMap<NodeId, Arc<dyn crate::Op>>,
) -> Option<HashMap<NodeId, ImageDescriptor>> {
let mut descriptors: HashMap<NodeId, ImageDescriptor> = HashMap::with_capacity(order.len());
for node in order {
if node.id() == source {
descriptors.insert(node.id(), reduced);
continue;
}
let op = rescaled.get(&node.id())?;
let inputs: Vec<ImageDescriptor> = node
.inputs()
.iter()
.filter_map(|input| descriptors.get(&input.id()).copied())
.collect();
if inputs.len() != node.inputs().len() {
return None;
}
descriptors.insert(node.id(), op.output_descriptor(&inputs).ok()?);
}
Some(descriptors)
}
fn rebuild(
node: &Arc<Node>,
source: NodeId,
producer: &Arc<dyn crate::Producer>,
original: &Image,
rescaled: &HashMap<NodeId, Arc<dyn crate::Op>>,
memo: &mut HashMap<NodeId, Image>,
) -> Result<Image> {
if let Some(built) = memo.get(&node.id()) {
return Ok(built.clone());
}
let built = if node.id() == source {
Image::from_producer(Arc::clone(producer), original.metadata()?.format)
} else {
let op = rescaled
.get(&node.id())
.ok_or_else(|| crate::PixelsError::graph("an op was rebuilt without being rescaled"))?;
let inputs = node
.inputs()
.iter()
.map(|input| rebuild(input, source, producer, original, rescaled, memo))
.collect::<Result<Vec<_>>>()?;
Image::combine(&inputs, Arc::clone(op))?
};
memo.insert(node.id(), built.clone());
Ok(built)
}
fn topological_order(root: &Arc<Node>) -> Vec<Arc<Node>> {
let mut order = Vec::new();
let mut seen = std::collections::HashSet::new();
visit(root, &mut seen, &mut order);
order
}
fn visit(
node: &Arc<Node>,
seen: &mut std::collections::HashSet<NodeId>,
order: &mut Vec<Arc<Node>>,
) {
if !seen.insert(node.id()) {
return;
}
for input in node.inputs() {
visit(input, seen, order);
}
order.push(Arc::clone(node));
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
clippy::indexing_slicing,
clippy::panic,
reason = "tests operate on known-good values and assert shapes directly"
)]
mod tests {
use super::*;
use crate::{
AccessPattern, DecodeCapability, Format, Op, PixelFormat, PixelsError, Producer, Region,
TileMut,
};
use std::sync::Mutex;
#[derive(Debug)]
struct Shrinkable {
descriptor: Mutex<ImageDescriptor>,
reduced: Mutex<bool>,
}
impl Shrinkable {
fn new(width: u32, height: u32) -> Arc<Self> {
Arc::new(Self {
descriptor: Mutex::new(
ImageDescriptor::new(width, height, PixelFormat::Gray8).unwrap(),
),
reduced: Mutex::new(false),
})
}
fn shape(&self) -> ImageDescriptor {
*self.descriptor.lock().unwrap()
}
}
impl Producer for Shrinkable {
fn name(&self) -> &'static str {
"shrinkable"
}
fn descriptor(&self) -> ImageDescriptor {
self.shape()
}
fn capability(&self) -> DecodeCapability {
DecodeCapability::Regions
}
fn produce(&self, _: Region, _: &mut TileMut<'_>) -> Result<()> {
Ok(())
}
fn reduced_descriptor(&self, target: (u32, u32)) -> Option<ImageDescriptor> {
let full = self.shape();
let (half_width, half_height) = (full.width / 2, full.height / 2);
if half_width >= target.0 && half_height >= target.1 {
ImageDescriptor::new(half_width, half_height, full.pixel).ok()
} else {
None
}
}
fn reduce_to(&self, descriptor: ImageDescriptor) -> Result<()> {
*self.descriptor.lock().unwrap() = descriptor;
*self.reduced.lock().unwrap() = true;
Ok(())
}
}
#[derive(Debug)]
struct FixedResize {
width: u32,
height: u32,
covariant: bool,
}
impl Op for FixedResize {
fn name(&self) -> &'static str {
"fixed-resize"
}
fn access_pattern(&self) -> AccessPattern {
AccessPattern::Spatial
}
fn output_descriptor(&self, inputs: &[ImageDescriptor]) -> Result<ImageDescriptor> {
let input = inputs
.first()
.ok_or_else(|| PixelsError::graph("no input"))?;
ImageDescriptor::new(self.width, self.height, input.pixel)
}
fn input_regions(&self, _: Region, inputs: &[ImageDescriptor]) -> Result<Vec<Region>> {
Ok(vec![inputs[0].region()])
}
fn compute(&self, _: &[Tile<'_>], _: &mut TileMut<'_>) -> Result<()> {
Ok(())
}
fn rescaled(&self) -> Option<Arc<dyn Op>> {
self.covariant.then(|| {
Arc::new(Self {
width: self.width,
height: self.height,
covariant: true,
}) as Arc<dyn Op>
})
}
}
use crate::Tile;
#[derive(Debug)]
struct SameShape {
covariant: bool,
}
impl Op for SameShape {
fn name(&self) -> &'static str {
"same-shape"
}
fn access_pattern(&self) -> AccessPattern {
AccessPattern::Sequential
}
fn output_descriptor(&self, inputs: &[ImageDescriptor]) -> Result<ImageDescriptor> {
inputs
.first()
.copied()
.ok_or_else(|| PixelsError::graph("no input"))
}
fn input_regions(&self, output: Region, _: &[ImageDescriptor]) -> Result<Vec<Region>> {
Ok(vec![output])
}
fn compute(&self, _: &[Tile<'_>], _: &mut TileMut<'_>) -> Result<()> {
Ok(())
}
fn rescaled(&self) -> Option<Arc<dyn Op>> {
self.covariant
.then(|| Arc::new(Self { covariant: true }) as Arc<dyn Op>)
}
}
fn source(width: u32, height: u32) -> (Image, Arc<Shrinkable>) {
let producer = Shrinkable::new(width, height);
let image = Image::from_producer(Arc::clone(&producer) as Arc<dyn Producer>, Format::Jpeg);
(image, producer)
}
#[test]
fn a_resize_pipeline_shrinks_its_source() {
let (image, producer) = source(800, 600);
let pipeline = image
.apply(Arc::new(FixedResize {
width: 100,
height: 75,
covariant: true,
}))
.unwrap();
let (rebuilt, reduction) = shrink_on_load(&pipeline).unwrap();
let reduction = reduction.expect("a resize to 1/8 should shrink the source");
assert_eq!(reduction.from, (800, 600));
assert_eq!(reduction.to, (400, 300));
assert!((reduction.factor() - 4.0).abs() < 0.001);
assert_eq!(producer.shape().width, 400);
assert_eq!(
(rebuilt.descriptor().width, rebuilt.descriptor().height),
(100, 75)
);
}
#[test]
fn an_op_that_is_not_scale_covariant_blocks_the_reduction() {
let (image, producer) = source(800, 600);
let pipeline = image
.apply(Arc::new(FixedResize {
width: 100,
height: 75,
covariant: false,
}))
.unwrap();
let (rebuilt, reduction) = shrink_on_load(&pipeline).unwrap();
assert!(reduction.is_none(), "a non-covariant op must block it");
assert_eq!(producer.shape().width, 800, "the source was reduced anyway");
assert_eq!(rebuilt.descriptor().width, 100);
}
#[test]
fn a_pipeline_with_no_resize_keeps_its_size() {
let (image, producer) = source(800, 600);
let pipeline = image
.apply(Arc::new(SameShape { covariant: true }))
.unwrap();
let (rebuilt, reduction) = shrink_on_load(&pipeline).unwrap();
assert!(
reduction.is_none(),
"reducing here would change the output size"
);
assert_eq!(producer.shape().width, 800);
assert_eq!(rebuilt.descriptor().width, 800);
}
#[test]
fn a_source_that_cannot_reduce_is_left_alone() {
let descriptor = ImageDescriptor::new(64, 64, PixelFormat::Gray8).unwrap();
let buffer = Arc::new(crate::TileBuf::for_image(&descriptor).unwrap());
let image = Image::from_producer(
Arc::new(crate::BufferSource::new(descriptor, buffer).unwrap()),
Format::Raw,
);
let pipeline = image
.apply(Arc::new(FixedResize {
width: 8,
height: 8,
covariant: true,
}))
.unwrap();
let (_, reduction) = shrink_on_load(&pipeline).unwrap();
assert!(reduction.is_none(), "pixels in memory have one resolution");
}
#[test]
fn the_reduction_never_goes_below_the_target() {
let (image, producer) = source(800, 600);
let pipeline = image
.apply(Arc::new(FixedResize {
width: 401,
height: 301,
covariant: true,
}))
.unwrap();
let (_, reduction) = shrink_on_load(&pipeline).unwrap();
assert!(reduction.is_none());
assert_eq!(producer.shape().width, 800);
}
#[test]
fn a_shared_subgraph_stays_shared_through_the_rebuild() {
let (image, _) = source(800, 600);
let resized = image
.apply(Arc::new(FixedResize {
width: 100,
height: 75,
covariant: true,
}))
.unwrap();
let branch = resized
.apply(Arc::new(SameShape { covariant: true }))
.unwrap();
let before = branch.node().node_count();
let (rebuilt, reduction) = shrink_on_load(&branch).unwrap();
assert!(reduction.is_some());
assert_eq!(
rebuilt.node().node_count(),
before,
"the rebuild changed the graph's shape"
);
}
}