Skip to main content

wmv_decoder/
na_wmv2dsp.rs

1#[inline(always)]
2fn clip_u8(v: i32) -> u8 {
3    if v < 0 {
4        0
5    } else if v > 255 {
6        255
7    } else {
8        v as u8
9    }
10}
11
12const W0: i32 = 2048;
13const W1: i32 = 2841;
14const W2: i32 = 2676;
15const W3: i32 = 2408;
16const W4: i32 = 2048;
17const W5: i32 = 1609;
18const W6: i32 = 1108;
19const W7: i32 = 565;
20
21#[inline(always)]
22fn wmv2_idct_row(b: &mut [i16]) {
23    debug_assert!(b.len() == 8);
24    let (b0, b1, b2, b3, b4, b5, b6, b7) = (
25        b[0] as i32,
26        b[1] as i32,
27        b[2] as i32,
28        b[3] as i32,
29        b[4] as i32,
30        b[5] as i32,
31        b[6] as i32,
32        b[7] as i32,
33    );
34
35    // step 1
36    let a1 = W1 * b1 + W7 * b7;
37    let a7 = W7 * b1 - W1 * b7;
38    let a5 = W5 * b5 + W3 * b3;
39    let a3 = W3 * b5 - W5 * b3;
40    let a2 = W2 * b2 + W6 * b6;
41    let a6 = W6 * b2 - W2 * b6;
42    let a0 = W0 * b0 + W0 * b4;
43    let a4 = W0 * b0 - W0 * b4;
44
45    // step 2
46    let s1 = ((181i32 * (a1 - a5 + a7 - a3) + 128) >> 8) as i32;
47    let s2 = ((181i32 * (a1 - a5 - a7 + a3) + 128) >> 8) as i32;
48
49    // step 3
50    b[0] = ((a0 + a2 + a1 + a5 + (1 << 7)) >> 8) as i16;
51    b[1] = ((a4 + a6 + s1 + (1 << 7)) >> 8) as i16;
52    b[2] = ((a4 - a6 + s2 + (1 << 7)) >> 8) as i16;
53    b[3] = ((a0 - a2 + a7 + a3 + (1 << 7)) >> 8) as i16;
54    b[4] = ((a0 - a2 - a7 - a3 + (1 << 7)) >> 8) as i16;
55    b[5] = ((a4 - a6 - s2 + (1 << 7)) >> 8) as i16;
56    b[6] = ((a4 + a6 - s1 + (1 << 7)) >> 8) as i16;
57    b[7] = ((a0 + a2 - a1 - a5 + (1 << 7)) >> 8) as i16;
58}
59
60#[inline(always)]
61fn wmv2_idct_col(block: &mut [i16; 64], col: usize) {
62    // step 1, with extended precision
63    let b1 = block[8 * 1 + col] as i32;
64    let b7 = block[8 * 7 + col] as i32;
65    let b5 = block[8 * 5 + col] as i32;
66    let b3 = block[8 * 3 + col] as i32;
67    let b2 = block[8 * 2 + col] as i32;
68    let b6 = block[8 * 6 + col] as i32;
69    let b0 = block[8 * 0 + col] as i32;
70    let b4 = block[8 * 4 + col] as i32;
71
72    let a1 = (W1 * b1 + W7 * b7 + 4) >> 3;
73    let a7 = (W7 * b1 - W1 * b7 + 4) >> 3;
74    let a5 = (W5 * b5 + W3 * b3 + 4) >> 3;
75    let a3 = (W3 * b5 - W5 * b3 + 4) >> 3;
76    let a2 = (W2 * b2 + W6 * b6 + 4) >> 3;
77    let a6 = (W6 * b2 - W2 * b6 + 4) >> 3;
78    let a0 = (W0 * b0 + W0 * b4) >> 3;
79    let a4 = (W0 * b0 - W0 * b4) >> 3;
80
81    // step 2
82    let s1 = (181i32 * (a1 - a5 + a7 - a3) + 128) >> 8;
83    let s2 = (181i32 * (a1 - a5 - a7 + a3) + 128) >> 8;
84
85    // step 3
86    block[8 * 0 + col] = ((a0 + a2 + a1 + a5 + (1 << 13)) >> 14) as i16;
87    block[8 * 1 + col] = ((a4 + a6 + s1 + (1 << 13)) >> 14) as i16;
88    block[8 * 2 + col] = ((a4 - a6 + s2 + (1 << 13)) >> 14) as i16;
89    block[8 * 3 + col] = ((a0 - a2 + a7 + a3 + (1 << 13)) >> 14) as i16;
90
91    block[8 * 4 + col] = ((a0 - a2 - a7 - a3 + (1 << 13)) >> 14) as i16;
92    block[8 * 5 + col] = ((a4 - a6 - s2 + (1 << 13)) >> 14) as i16;
93    block[8 * 6 + col] = ((a4 + a6 - s1 + (1 << 13)) >> 14) as i16;
94    block[8 * 7 + col] = ((a0 + a2 - a1 - a5 + (1 << 13)) >> 14) as i16;
95}
96
97pub fn wmv2_idct_add(dest: &mut [u8], dest_off: usize, stride: usize, block: &mut [i16; 64]) {
98    // row pass
99    for i in (0..64).step_by(8) {
100        wmv2_idct_row(&mut block[i..i + 8]);
101    }
102    // col pass
103    for c in 0..8 {
104        wmv2_idct_col(block, c);
105    }
106
107    // add
108    for r in 0..8 {
109        let d = dest_off + r * stride;
110        let b = r * 8;
111        for c in 0..8 {
112            let idx = d + c;
113            if idx >= dest.len() {
114                continue;
115            }
116            let v = dest[idx] as i32 + block[b + c] as i32;
117            dest[idx] = clip_u8(v);
118        }
119    }
120}
121
122pub fn wmv2_idct_put(dest: &mut [u8], dest_off: usize, stride: usize, block: &mut [i16; 64]) {
123    // row pass
124    for i in (0..64).step_by(8) {
125        wmv2_idct_row(&mut block[i..i + 8]);
126    }
127    // col pass
128    for c in 0..8 {
129        wmv2_idct_col(block, c);
130    }
131
132    // put
133    for r in 0..8 {
134        let d = dest_off + r * stride;
135        let b = r * 8;
136        for c in 0..8 {
137            let idx = d + c;
138            if idx >= dest.len() {
139                continue;
140            }
141            dest[idx] = clip_u8(block[b + c] as i32);
142        }
143    }
144}