1use crate::{Image, ImageDescriptor, Node, NodeId, Result};
43use std::collections::HashMap;
44use std::sync::Arc;
45
46#[derive(Debug, Clone, Copy, PartialEq, Eq)]
48pub struct Reduction {
49 pub from: (u32, u32),
51 pub to: (u32, u32),
53}
54
55impl Reduction {
56 #[must_use]
58 pub fn factor(&self) -> f64 {
59 let before = f64::from(self.from.0) * f64::from(self.from.1);
60 let after = (f64::from(self.to.0) * f64::from(self.to.1)).max(1.0);
61 before / after
62 }
63}
64
65pub fn shrink_on_load(image: &Image) -> Result<(Image, Option<Reduction>)> {
78 let root = Arc::clone(image.node());
79 let order = topological_order(&root);
80
81 let mut rescaled: HashMap<NodeId, Arc<dyn crate::Op>> = HashMap::new();
87 for node in &order {
88 let Some(op) = node.op() else { continue };
89 let Some(copy) = op.rescaled() else {
90 return Ok((image.clone(), None));
91 };
92 rescaled.insert(node.id(), copy);
93 }
94
95 let sources: Vec<&Arc<Node>> = order
97 .iter()
98 .filter(|node| node.producer().is_some())
99 .collect();
100 let [source] = sources.as_slice() else {
101 return Ok((image.clone(), None));
102 };
103 let Some(producer) = source.producer() else {
104 return Ok((image.clone(), None));
105 };
106
107 let full = source.descriptor();
108 let target = (root.descriptor().width, root.descriptor().height);
109 let Some(reduced) = producer.reduced_descriptor(target) else {
110 return Ok((image.clone(), None));
111 };
112 if (reduced.width, reduced.height) == (full.width, full.height) {
113 return Ok((image.clone(), None));
114 }
115
116 let Some(simulated) = simulate(&order, source.id(), reduced, &rescaled) else {
119 return Ok((image.clone(), None));
120 };
121 let Some(&new_root) = simulated.get(&root.id()) else {
122 return Ok((image.clone(), None));
123 };
124 if (new_root.width, new_root.height) != (root.descriptor().width, root.descriptor().height) {
125 return Ok((image.clone(), None));
126 }
127
128 producer.reduce_to(reduced)?;
130 let rebuilt = rebuild(
131 &root,
132 source.id(),
133 producer,
134 image,
135 &rescaled,
136 &mut HashMap::new(),
137 )?;
138
139 Ok((
140 rebuilt,
141 Some(Reduction {
142 from: (full.width, full.height),
143 to: (reduced.width, reduced.height),
144 }),
145 ))
146}
147
148fn simulate(
154 order: &[Arc<Node>],
155 source: NodeId,
156 reduced: ImageDescriptor,
157 rescaled: &HashMap<NodeId, Arc<dyn crate::Op>>,
158) -> Option<HashMap<NodeId, ImageDescriptor>> {
159 let mut descriptors: HashMap<NodeId, ImageDescriptor> = HashMap::with_capacity(order.len());
160 for node in order {
161 if node.id() == source {
162 descriptors.insert(node.id(), reduced);
163 continue;
164 }
165 let op = rescaled.get(&node.id())?;
171 let inputs: Vec<ImageDescriptor> = node
172 .inputs()
173 .iter()
174 .filter_map(|input| descriptors.get(&input.id()).copied())
175 .collect();
176 if inputs.len() != node.inputs().len() {
177 return None;
178 }
179 descriptors.insert(node.id(), op.output_descriptor(&inputs).ok()?);
180 }
181 Some(descriptors)
182}
183
184fn rebuild(
190 node: &Arc<Node>,
191 source: NodeId,
192 producer: &Arc<dyn crate::Producer>,
193 original: &Image,
194 rescaled: &HashMap<NodeId, Arc<dyn crate::Op>>,
195 memo: &mut HashMap<NodeId, Image>,
196) -> Result<Image> {
197 if let Some(built) = memo.get(&node.id()) {
198 return Ok(built.clone());
199 }
200 let built = if node.id() == source {
201 Image::from_producer(Arc::clone(producer), original.metadata()?.format)
204 } else {
205 let op = rescaled
208 .get(&node.id())
209 .ok_or_else(|| crate::PixelsError::graph("an op was rebuilt without being rescaled"))?;
210 let inputs = node
211 .inputs()
212 .iter()
213 .map(|input| rebuild(input, source, producer, original, rescaled, memo))
214 .collect::<Result<Vec<_>>>()?;
215 Image::combine(&inputs, Arc::clone(op))?
216 };
217 memo.insert(node.id(), built.clone());
218 Ok(built)
219}
220
221fn topological_order(root: &Arc<Node>) -> Vec<Arc<Node>> {
223 let mut order = Vec::new();
224 let mut seen = std::collections::HashSet::new();
225 visit(root, &mut seen, &mut order);
226 order
227}
228
229fn visit(
230 node: &Arc<Node>,
231 seen: &mut std::collections::HashSet<NodeId>,
232 order: &mut Vec<Arc<Node>>,
233) {
234 if !seen.insert(node.id()) {
235 return;
236 }
237 for input in node.inputs() {
238 visit(input, seen, order);
239 }
240 order.push(Arc::clone(node));
241}
242
243#[cfg(test)]
244#[allow(
245 clippy::unwrap_used,
246 clippy::expect_used,
247 clippy::indexing_slicing,
248 clippy::panic,
249 reason = "tests operate on known-good values and assert shapes directly"
250)]
251mod tests {
252 use super::*;
253 use crate::{
254 AccessPattern, DecodeCapability, Format, Op, PixelFormat, PixelsError, Producer, Region,
255 TileMut,
256 };
257 use std::sync::Mutex;
258
259 #[derive(Debug)]
261 struct Shrinkable {
262 descriptor: Mutex<ImageDescriptor>,
263 reduced: Mutex<bool>,
264 }
265
266 impl Shrinkable {
267 fn new(width: u32, height: u32) -> Arc<Self> {
268 Arc::new(Self {
269 descriptor: Mutex::new(
270 ImageDescriptor::new(width, height, PixelFormat::Gray8).unwrap(),
271 ),
272 reduced: Mutex::new(false),
273 })
274 }
275
276 fn shape(&self) -> ImageDescriptor {
277 *self.descriptor.lock().unwrap()
278 }
279 }
280
281 impl Producer for Shrinkable {
282 fn name(&self) -> &'static str {
283 "shrinkable"
284 }
285 fn descriptor(&self) -> ImageDescriptor {
286 self.shape()
287 }
288 fn capability(&self) -> DecodeCapability {
289 DecodeCapability::Regions
290 }
291 fn produce(&self, _: Region, _: &mut TileMut<'_>) -> Result<()> {
292 Ok(())
293 }
294 fn reduced_descriptor(&self, target: (u32, u32)) -> Option<ImageDescriptor> {
295 let full = self.shape();
296 let (half_width, half_height) = (full.width / 2, full.height / 2);
297 if half_width >= target.0 && half_height >= target.1 {
298 ImageDescriptor::new(half_width, half_height, full.pixel).ok()
299 } else {
300 None
301 }
302 }
303 fn reduce_to(&self, descriptor: ImageDescriptor) -> Result<()> {
304 *self.descriptor.lock().unwrap() = descriptor;
305 *self.reduced.lock().unwrap() = true;
306 Ok(())
307 }
308 }
309
310 #[derive(Debug)]
312 struct FixedResize {
313 width: u32,
314 height: u32,
315 covariant: bool,
316 }
317
318 impl Op for FixedResize {
319 fn name(&self) -> &'static str {
320 "fixed-resize"
321 }
322 fn access_pattern(&self) -> AccessPattern {
323 AccessPattern::Spatial
324 }
325 fn output_descriptor(&self, inputs: &[ImageDescriptor]) -> Result<ImageDescriptor> {
326 let input = inputs
327 .first()
328 .ok_or_else(|| PixelsError::graph("no input"))?;
329 ImageDescriptor::new(self.width, self.height, input.pixel)
330 }
331 fn input_regions(&self, _: Region, inputs: &[ImageDescriptor]) -> Result<Vec<Region>> {
332 Ok(vec![inputs[0].region()])
333 }
334 fn compute(&self, _: &[Tile<'_>], _: &mut TileMut<'_>) -> Result<()> {
335 Ok(())
336 }
337 fn rescaled(&self) -> Option<Arc<dyn Op>> {
338 self.covariant.then(|| {
339 Arc::new(Self {
340 width: self.width,
341 height: self.height,
342 covariant: true,
343 }) as Arc<dyn Op>
344 })
345 }
346 }
347
348 use crate::Tile;
349
350 #[derive(Debug)]
352 struct SameShape {
353 covariant: bool,
354 }
355
356 impl Op for SameShape {
357 fn name(&self) -> &'static str {
358 "same-shape"
359 }
360 fn access_pattern(&self) -> AccessPattern {
361 AccessPattern::Sequential
362 }
363 fn output_descriptor(&self, inputs: &[ImageDescriptor]) -> Result<ImageDescriptor> {
364 inputs
365 .first()
366 .copied()
367 .ok_or_else(|| PixelsError::graph("no input"))
368 }
369 fn input_regions(&self, output: Region, _: &[ImageDescriptor]) -> Result<Vec<Region>> {
370 Ok(vec![output])
371 }
372 fn compute(&self, _: &[Tile<'_>], _: &mut TileMut<'_>) -> Result<()> {
373 Ok(())
374 }
375 fn rescaled(&self) -> Option<Arc<dyn Op>> {
376 self.covariant
377 .then(|| Arc::new(Self { covariant: true }) as Arc<dyn Op>)
378 }
379 }
380
381 fn source(width: u32, height: u32) -> (Image, Arc<Shrinkable>) {
382 let producer = Shrinkable::new(width, height);
383 let image = Image::from_producer(Arc::clone(&producer) as Arc<dyn Producer>, Format::Jpeg);
384 (image, producer)
385 }
386
387 #[test]
388 fn a_resize_pipeline_shrinks_its_source() {
389 let (image, producer) = source(800, 600);
390 let pipeline = image
391 .apply(Arc::new(FixedResize {
392 width: 100,
393 height: 75,
394 covariant: true,
395 }))
396 .unwrap();
397
398 let (rebuilt, reduction) = shrink_on_load(&pipeline).unwrap();
399 let reduction = reduction.expect("a resize to 1/8 should shrink the source");
400 assert_eq!(reduction.from, (800, 600));
401 assert_eq!(reduction.to, (400, 300));
402 assert!((reduction.factor() - 4.0).abs() < 0.001);
403
404 assert_eq!(producer.shape().width, 400);
407 assert_eq!(
408 (rebuilt.descriptor().width, rebuilt.descriptor().height),
409 (100, 75)
410 );
411 }
412
413 #[test]
414 fn an_op_that_is_not_scale_covariant_blocks_the_reduction() {
415 let (image, producer) = source(800, 600);
416 let pipeline = image
418 .apply(Arc::new(FixedResize {
419 width: 100,
420 height: 75,
421 covariant: false,
422 }))
423 .unwrap();
424
425 let (rebuilt, reduction) = shrink_on_load(&pipeline).unwrap();
426 assert!(reduction.is_none(), "a non-covariant op must block it");
427 assert_eq!(producer.shape().width, 800, "the source was reduced anyway");
428 assert_eq!(rebuilt.descriptor().width, 100);
429 }
430
431 #[test]
432 fn a_pipeline_with_no_resize_keeps_its_size() {
433 let (image, producer) = source(800, 600);
436 let pipeline = image
437 .apply(Arc::new(SameShape { covariant: true }))
438 .unwrap();
439
440 let (rebuilt, reduction) = shrink_on_load(&pipeline).unwrap();
441 assert!(
442 reduction.is_none(),
443 "reducing here would change the output size"
444 );
445 assert_eq!(producer.shape().width, 800);
446 assert_eq!(rebuilt.descriptor().width, 800);
447 }
448
449 #[test]
450 fn a_source_that_cannot_reduce_is_left_alone() {
451 let descriptor = ImageDescriptor::new(64, 64, PixelFormat::Gray8).unwrap();
452 let buffer = Arc::new(crate::TileBuf::for_image(&descriptor).unwrap());
453 let image = Image::from_producer(
454 Arc::new(crate::BufferSource::new(descriptor, buffer).unwrap()),
455 Format::Raw,
456 );
457 let pipeline = image
458 .apply(Arc::new(FixedResize {
459 width: 8,
460 height: 8,
461 covariant: true,
462 }))
463 .unwrap();
464
465 let (_, reduction) = shrink_on_load(&pipeline).unwrap();
466 assert!(reduction.is_none(), "pixels in memory have one resolution");
467 }
468
469 #[test]
470 fn the_reduction_never_goes_below_the_target() {
471 let (image, producer) = source(800, 600);
474 let pipeline = image
475 .apply(Arc::new(FixedResize {
476 width: 401,
477 height: 301,
478 covariant: true,
479 }))
480 .unwrap();
481
482 let (_, reduction) = shrink_on_load(&pipeline).unwrap();
483 assert!(reduction.is_none());
484 assert_eq!(producer.shape().width, 800);
485 }
486
487 #[test]
488 fn a_shared_subgraph_stays_shared_through_the_rebuild() {
489 let (image, _) = source(800, 600);
490 let resized = image
491 .apply(Arc::new(FixedResize {
492 width: 100,
493 height: 75,
494 covariant: true,
495 }))
496 .unwrap();
497 let branch = resized
500 .apply(Arc::new(SameShape { covariant: true }))
501 .unwrap();
502
503 let before = branch.node().node_count();
504 let (rebuilt, reduction) = shrink_on_load(&branch).unwrap();
505 assert!(reduction.is_some());
506 assert_eq!(
507 rebuilt.node().node_count(),
508 before,
509 "the rebuild changed the graph's shape"
510 );
511 }
512}