1use otf_pixels_core::{PixelsError, Result};
19
20const SM_WEIGHTS_4: [i32; 4] = [255, 149, 85, 64];
22const SM_WEIGHTS_8: [i32; 8] = [255, 197, 146, 105, 73, 50, 37, 32];
24const SM_WEIGHTS_16: [i32; 16] = [
26 255, 225, 196, 170, 145, 123, 102, 84, 68, 54, 43, 33, 26, 20, 17, 16,
27];
28const SM_WEIGHTS_32: [i32; 32] = [
30 255, 240, 225, 210, 196, 182, 169, 157, 145, 133, 122, 111, 101, 92, 83, 74, 66, 59, 52, 45,
31 39, 34, 29, 25, 21, 17, 14, 12, 10, 9, 8, 8,
32];
33const SM_WEIGHTS_64: [i32; 64] = [
35 255, 248, 240, 233, 225, 218, 210, 203, 196, 189, 182, 176, 169, 163, 156, 150, 144, 138, 133,
36 127, 121, 116, 111, 106, 101, 96, 91, 86, 82, 77, 73, 69, 65, 61, 57, 54, 50, 47, 44, 41, 38,
37 35, 32, 29, 27, 25, 22, 20, 18, 16, 15, 13, 12, 10, 9, 8, 7, 6, 6, 5, 5, 4, 4, 4,
38];
39
40fn sm_weights(dim: usize) -> &'static [i32] {
42 match dim {
43 8 => &SM_WEIGHTS_8,
44 16 => &SM_WEIGHTS_16,
45 32 => &SM_WEIGHTS_32,
46 64 => &SM_WEIGHTS_64,
47 _ => &SM_WEIGHTS_4,
48 }
49}
50
51#[derive(Debug, Clone, Copy, PartialEq, Eq)]
53pub enum IntraMode {
54 Dc,
56 V,
58 H,
60 D45,
62 D135,
64 D113,
66 D157,
68 D203,
70 D67,
72 Smooth,
74 SmoothV,
76 SmoothH,
78 Paeth,
80}
81
82impl IntraMode {
83 #[must_use]
85 pub fn from_index(index: u8) -> Option<Self> {
86 Some(match index {
87 0 => Self::Dc,
88 1 => Self::V,
89 2 => Self::H,
90 3 => Self::D45,
91 4 => Self::D135,
92 5 => Self::D113,
93 6 => Self::D157,
94 7 => Self::D203,
95 8 => Self::D67,
96 9 => Self::Smooth,
97 10 => Self::SmoothV,
98 11 => Self::SmoothH,
99 12 => Self::Paeth,
100 _ => return None,
101 })
102 }
103
104 #[must_use]
107 pub fn is_directional(self) -> bool {
108 matches!(
109 self,
110 Self::V
111 | Self::H
112 | Self::D45
113 | Self::D135
114 | Self::D113
115 | Self::D157
116 | Self::D203
117 | Self::D67
118 )
119 }
120}
121
122fn round2(x: i32, n: u32) -> i32 {
124 if n == 0 { x } else { (x + (1 << (n - 1))) >> n }
125}
126
127#[derive(Debug, Clone, Copy)]
131pub struct Neighbours {
132 pub above: [i32; 8],
134 pub left: [i32; 8],
136 pub corner: i32,
138 pub have_above: bool,
140 pub have_left: bool,
142}
143
144#[derive(Debug, Clone, Copy)]
149pub struct PredBlock<'a> {
150 pub above: &'a [i32],
152 pub left: &'a [i32],
154 pub corner: i32,
156 pub have_above: bool,
158 pub have_left: bool,
160 pub w: usize,
162 pub h: usize,
164}
165
166pub fn predict_intra_block(mode: IntraMode, b: &PredBlock<'_>, bit_depth: u8) -> Result<Vec<u16>> {
175 let (w, h) = (b.w, b.h);
176 let max = (1_i32 << bit_depth) - 1;
177 let clip1 = |v: i32| v.clamp(0, max) as u16;
178 let a = |j: usize| b.above.get(j).copied().unwrap_or(0);
179 let l = |i: usize| b.left.get(i).copied().unwrap_or(0);
180
181 let mut pred = vec![0_u16; w * h];
182 let put = |pred: &mut Vec<u16>, i: usize, j: usize, v: u16| {
183 if let Some(cell) = pred.get_mut(i * w + j) {
184 *cell = v;
185 }
186 };
187
188 match mode {
189 IntraMode::Dc => {
190 let value = dc_value(b, bit_depth);
191 pred.fill(value);
192 }
193 IntraMode::V => {
194 for i in 0..h {
196 for j in 0..w {
197 put(&mut pred, i, j, clip1(a(j)));
198 }
199 }
200 }
201 IntraMode::H => {
202 for i in 0..h {
204 for j in 0..w {
205 put(&mut pred, i, j, clip1(l(i)));
206 }
207 }
208 }
209 IntraMode::Paeth => {
210 for i in 0..h {
211 for j in 0..w {
212 let base = a(j) + l(i) - b.corner;
213 let p_left = (base - l(i)).abs();
214 let p_top = (base - a(j)).abs();
215 let p_corner = (base - b.corner).abs();
216 let v = if p_left <= p_top && p_left <= p_corner {
217 l(i)
218 } else if p_top <= p_corner {
219 a(j)
220 } else {
221 b.corner
222 };
223 put(&mut pred, i, j, clip1(v));
224 }
225 }
226 }
227 IntraMode::Smooth => {
228 let wx = sm_weights(w);
229 let wy = sm_weights(h);
230 let below_left = l(h - 1);
231 let above_right = a(w - 1);
232 for i in 0..h {
233 let wyi = wy.get(i).copied().unwrap_or(0);
234 for j in 0..w {
235 let wxj = wx.get(j).copied().unwrap_or(0);
236 let smooth = wyi * a(j)
237 + (256 - wyi) * below_left
238 + wxj * l(i)
239 + (256 - wxj) * above_right;
240 put(&mut pred, i, j, clip1(round2(smooth, 9)));
241 }
242 }
243 }
244 IntraMode::SmoothV => {
245 let wy = sm_weights(h);
246 let below_left = l(h - 1);
247 for i in 0..h {
248 let wyi = wy.get(i).copied().unwrap_or(0);
249 for j in 0..w {
250 let smooth = wyi * a(j) + (256 - wyi) * below_left;
251 put(&mut pred, i, j, clip1(round2(smooth, 8)));
252 }
253 }
254 }
255 IntraMode::SmoothH => {
256 let wx = sm_weights(w);
257 let above_right = a(w - 1);
258 for i in 0..h {
259 for j in 0..w {
260 let wxj = wx.get(j).copied().unwrap_or(0);
261 let smooth = wxj * l(i) + (256 - wxj) * above_right;
262 put(&mut pred, i, j, clip1(round2(smooth, 8)));
263 }
264 }
265 }
266 IntraMode::D45
267 | IntraMode::D135
268 | IntraMode::D113
269 | IntraMode::D157
270 | IntraMode::D203
271 | IntraMode::D67 => {
272 return Err(PixelsError::unsupported(
273 "avif: slanted directional intra prediction is not implemented yet",
274 ));
275 }
276 }
277 Ok(pred)
278}
279
280pub fn predict_intra_4x4(mode: IntraMode, n: &Neighbours, bit_depth: u8) -> Result<[[u16; 4]; 4]> {
288 let block = PredBlock {
289 above: &n.above,
290 left: &n.left,
291 corner: n.corner,
292 have_above: n.have_above,
293 have_left: n.have_left,
294 w: 4,
295 h: 4,
296 };
297 let flat = predict_intra_block(mode, &block, bit_depth)?;
298 let mut pred = [[0_u16; 4]; 4];
299 for (i, row) in pred.iter_mut().enumerate() {
300 for (j, cell) in row.iter_mut().enumerate() {
301 *cell = flat.get(i * 4 + j).copied().unwrap_or(0);
302 }
303 }
304 Ok(pred)
305}
306
307fn dc_value(b: &PredBlock<'_>, bit_depth: u8) -> u16 {
309 let max = (1_i32 << bit_depth) - 1;
310 let clip1 = |v: i32| v.clamp(0, max) as u16;
311 let (w, h) = (b.w, b.h);
312 let left_sum: i32 = b.left.iter().take(h).sum();
313 let above_sum: i32 = b.above.iter().take(w).sum();
314 match (b.have_left, b.have_above) {
315 (true, true) => {
316 let sum = left_sum + above_sum + ((w + h) >> 1) as i32;
318 (sum / (w + h) as i32) as u16
319 }
320 (true, false) => clip1((left_sum + (h >> 1) as i32) >> h.trailing_zeros()),
321 (false, true) => clip1((above_sum + (w >> 1) as i32) >> w.trailing_zeros()),
322 (false, false) => 1_u16 << (bit_depth - 1),
323 }
324}
325
326#[cfg(test)]
327#[allow(
328 clippy::unwrap_used,
329 clippy::indexing_slicing,
330 clippy::panic,
331 reason = "tests operate on known-good values and assert shapes directly"
332)]
333mod tests {
334 use super::*;
335
336 fn neighbours(above: [i32; 8], left: [i32; 8], corner: i32) -> Neighbours {
337 Neighbours {
338 above,
339 left,
340 corner,
341 have_above: true,
342 have_left: true,
343 }
344 }
345
346 #[test]
347 fn mode_indexing_and_directional_classification() {
348 assert_eq!(IntraMode::from_index(0), Some(IntraMode::Dc));
349 assert_eq!(IntraMode::from_index(12), Some(IntraMode::Paeth));
350 assert_eq!(IntraMode::from_index(13), None);
351 assert!(IntraMode::V.is_directional());
352 assert!(!IntraMode::Dc.is_directional());
353 assert!(!IntraMode::Smooth.is_directional());
354 assert!(IntraMode::D45.is_directional());
355 }
356
357 #[test]
358 fn dc_with_no_neighbours_is_the_midpoint() {
359 let n = Neighbours {
360 above: [0; 8],
361 left: [0; 8],
362 corner: 0,
363 have_above: false,
364 have_left: false,
365 };
366 let pred = predict_intra_4x4(IntraMode::Dc, &n, 8).unwrap();
367 assert_eq!(pred, [[128; 4]; 4]);
368 }
369
370 #[test]
371 fn dc_averages_both_edges() {
372 let n = neighbours([100; 8], [60; 8], 100);
374 let pred = predict_intra_4x4(IntraMode::Dc, &n, 8).unwrap();
375 assert_eq!(pred, [[80; 4]; 4]);
376 }
377
378 #[test]
379 fn v_copies_the_above_row_down_each_column() {
380 let n = neighbours([10, 20, 30, 40, 0, 0, 0, 0], [99; 8], 5);
381 let pred = predict_intra_4x4(IntraMode::V, &n, 8).unwrap();
382 for row in &pred {
383 assert_eq!(row, &[10, 20, 30, 40]);
384 }
385 }
386
387 #[test]
388 fn h_copies_the_left_column_across_each_row() {
389 let n = neighbours([99; 8], [10, 20, 30, 40, 0, 0, 0, 0], 5);
390 let pred = predict_intra_4x4(IntraMode::H, &n, 8).unwrap();
391 for (i, row) in pred.iter().enumerate() {
392 assert!(row.iter().all(|&v| v == [10, 20, 30, 40][i]));
393 }
394 }
395
396 #[test]
397 fn paeth_picks_the_closest_predictor() {
398 let n = neighbours([50; 8], [50; 8], 50);
400 let pred = predict_intra_4x4(IntraMode::Paeth, &n, 8).unwrap();
401 assert_eq!(pred, [[50; 4]; 4]);
402 }
403
404 #[test]
405 fn smooth_of_a_flat_edge_is_that_value() {
406 let n = neighbours([128; 8], [128; 8], 128);
408 for mode in [IntraMode::Smooth, IntraMode::SmoothV, IntraMode::SmoothH] {
409 let pred = predict_intra_4x4(mode, &n, 8).unwrap();
410 assert_eq!(pred, [[128; 4]; 4], "mode {mode:?}");
411 }
412 }
413
414 #[test]
415 fn slanted_directional_modes_are_unsupported_for_now() {
416 let n = neighbours([100; 8], [100; 8], 100);
417 assert!(predict_intra_4x4(IntraMode::D45, &n, 8).is_err());
418 }
419
420 #[test]
421 fn dc_averages_both_edges_at_a_rectangular_size() {
422 let above = [100; 8];
425 let left = [60; 8];
426 let b = PredBlock {
427 above: &above,
428 left: &left,
429 corner: 100,
430 have_above: true,
431 have_left: true,
432 w: 8,
433 h: 4,
434 };
435 let pred = predict_intra_block(IntraMode::Dc, &b, 8).unwrap();
436 assert_eq!(pred.len(), 32);
437 assert!(pred.iter().all(|&v| v == 87));
438 }
439
440 #[test]
441 fn smooth_of_a_flat_edge_is_that_value_at_8x8() {
442 let above = [128; 8];
443 let left = [128; 8];
444 for mode in [IntraMode::Smooth, IntraMode::SmoothV, IntraMode::SmoothH] {
445 let b = PredBlock {
446 above: &above,
447 left: &left,
448 corner: 128,
449 have_above: true,
450 have_left: true,
451 w: 8,
452 h: 8,
453 };
454 let pred = predict_intra_block(mode, &b, 8).unwrap();
455 assert_eq!(pred.len(), 64);
456 assert!(pred.iter().all(|&v| v == 128), "mode {mode:?}");
457 }
458 }
459
460 #[test]
461 fn v_copies_the_above_row_at_8x8() {
462 let above = [10, 20, 30, 40, 50, 60, 70, 80];
463 let left = [0; 8];
464 let b = PredBlock {
465 above: &above,
466 left: &left,
467 corner: 0,
468 have_above: true,
469 have_left: true,
470 w: 8,
471 h: 8,
472 };
473 let pred = predict_intra_block(IntraMode::V, &b, 8).unwrap();
474 let expected: [u16; 8] = [10, 20, 30, 40, 50, 60, 70, 80];
476 for row in pred.chunks(8) {
477 assert_eq!(row, &expected[..]);
478 }
479 }
480}