Skip to main content

otf_pixels_codec_avif/av1/
predict.rs

1//! Intra prediction (spec §7.11.2).
2//!
3//! An intra block is predicted from the already-reconstructed samples in the
4//! row above and the column to its left, then the residual is added. This
5//! module owns the prediction; it takes the assembled `AboveRow`/`LeftCol`
6//! neighbour arrays and the mode, and returns the predicted block. Building the
7//! neighbour arrays from the plane (with the frame-edge and availability rules)
8//! is the tile driver's job, which keeps these predictors pure and testable.
9//!
10//! The modes that need no edge filtering are implemented at any transform size
11//! ([`predict_intra_block`]): DC (§7.11.2.5), Paeth (§7.11.2.2), the three
12//! Smooth variants (§7.11.2.6), and the two axis-aligned directional modes (V,
13//! H, which at exactly 90 and 180 degrees are plain copies). The slanted
14//! directional modes need the edge-filter/upsample machinery and report
15//! [`PixelsError::unsupported`] until that lands. [`predict_intra_4x4`] is a thin
16//! wrapper over the general path, which the lossless 4x4 tile drives.
17
18use otf_pixels_core::{PixelsError, Result};
19
20/// `Sm_Weights_Tx_4x4` (§9.3): the smooth-prediction interpolation weights.
21const SM_WEIGHTS_4: [i32; 4] = [255, 149, 85, 64];
22/// `Sm_Weights_Tx_8x8` (§9.3).
23const SM_WEIGHTS_8: [i32; 8] = [255, 197, 146, 105, 73, 50, 37, 32];
24/// `Sm_Weights_Tx_16x16` (§9.3).
25const SM_WEIGHTS_16: [i32; 16] = [
26    255, 225, 196, 170, 145, 123, 102, 84, 68, 54, 43, 33, 26, 20, 17, 16,
27];
28/// `Sm_Weights_Tx_32x32` (§9.3).
29const 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];
33/// `Sm_Weights_Tx_64x64` (§9.3).
34const 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
40/// `Sm_Weights_Tx[dim]` (§9.3): the interpolation weights for one side length.
41fn 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/// The 13 intra prediction modes (§6.10.2), in their coded order.
52#[derive(Debug, Clone, Copy, PartialEq, Eq)]
53pub enum IntraMode {
54    /// `DC_PRED` (0).
55    Dc,
56    /// `V_PRED` (1).
57    V,
58    /// `H_PRED` (2).
59    H,
60    /// `D45_PRED` (3).
61    D45,
62    /// `D135_PRED` (4).
63    D135,
64    /// `D113_PRED` (5).
65    D113,
66    /// `D157_PRED` (6).
67    D157,
68    /// `D203_PRED` (7).
69    D203,
70    /// `D67_PRED` (8).
71    D67,
72    /// `SMOOTH_PRED` (9).
73    Smooth,
74    /// `SMOOTH_V_PRED` (10).
75    SmoothV,
76    /// `SMOOTH_H_PRED` (11).
77    SmoothH,
78    /// `PAETH_PRED` (12).
79    Paeth,
80}
81
82impl IntraMode {
83    /// The mode for a coded index, or `None` if out of range.
84    #[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    /// Whether this is a directional mode (`is_directional_mode`, §7.11.2): the
105    /// eight modes V through D67.
106    #[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
122/// `Round2(x, n)` (§4.7).
123fn round2(x: i32, n: u32) -> i32 {
124    if n == 0 { x } else { (x + (1 << (n - 1))) >> n }
125}
126
127/// The neighbour samples a 4x4 intra predictor reads: the row above and the
128/// column to the left, each `w + h = 8` samples long, plus the shared
129/// top-left corner. Assembled by the tile driver per §7.11.2 general process.
130#[derive(Debug, Clone, Copy)]
131pub struct Neighbours {
132    /// `AboveRow[0..8]`.
133    pub above: [i32; 8],
134    /// `LeftCol[0..8]`.
135    pub left: [i32; 8],
136    /// `AboveRow[-1]` (equal to `LeftCol[-1]`).
137    pub corner: i32,
138    /// Whether the above row holds real reconstructed samples.
139    pub have_above: bool,
140    /// Whether the left column holds real reconstructed samples.
141    pub have_left: bool,
142}
143
144/// One intra block's assembled neighbours and geometry: the row above, the
145/// column to the left, the shared top-left `corner` (`AboveRow[-1]`), the
146/// availability flags, and the block size `w` x `h` in samples. `above` must
147/// hold at least `w` entries and `left` at least `h`.
148#[derive(Debug, Clone, Copy)]
149pub struct PredBlock<'a> {
150    /// `AboveRow[0..w]`.
151    pub above: &'a [i32],
152    /// `LeftCol[0..h]`.
153    pub left: &'a [i32],
154    /// `AboveRow[-1]` (equal to `LeftCol[-1]`).
155    pub corner: i32,
156    /// Whether the above row holds real reconstructed samples.
157    pub have_above: bool,
158    /// Whether the left column holds real reconstructed samples.
159    pub have_left: bool,
160    /// The block width in samples.
161    pub w: usize,
162    /// The block height in samples.
163    pub h: usize,
164}
165
166/// Predict an intra block of any size for the modes that need no edge filtering
167/// (§7.11.2): DC, Paeth, the three Smooth variants, and the axis-aligned V/H
168/// copies. The result is `w * h` samples in row-major order.
169///
170/// # Errors
171///
172/// Returns [`PixelsError::unsupported`] for the slanted directional modes, which
173/// need the edge-filter and upsample machinery not yet implemented.
174pub 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            // pAngle 90: a plain copy of the (unfiltered) above row.
195            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            // pAngle 180: a plain copy of the left column.
203            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
280/// Predict a 4x4 intra block (§7.11.2 for the 4x4 case): a thin wrapper over the
281/// size-general [`predict_intra_block`], reshaping the flat result to `[[_; 4];
282/// 4]`.
283///
284/// # Errors
285///
286/// Returns [`PixelsError::unsupported`] for the slanted directional modes.
287pub 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
307/// The DC prediction value (§7.11.2.5) for a block of any size.
308fn 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            // (sum + (w + h) / 2) / (w + h); already in range, so no Clip1.
317            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        // Above all 100, left all 60: avg = (400 + 240 + 4) / 8 = 80.
373        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        // Flat above=left=corner=50 -> base=50, all distances 0, picks left=50.
399        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        // All neighbours 128: every weighted combination is 128.
407        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        // 8 wide, 4 tall. Above all 100 (8 samples), left all 60 (4 samples):
423        // sum = 800 + 240 + ((8+4)>>1) = 1046; avg = 1046 / 12 = 87.
424        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        // Every one of the 8 rows repeats the above row.
475        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}