#![cfg_attr(not(feature = "unchecked"), forbid(unsafe_code))]
#![cfg_attr(feature = "unchecked", deny(unsafe_code))]
#![allow(unused_imports)]
#![allow(clippy::too_many_arguments)]
#[cfg(target_arch = "aarch64")]
use core::arch::aarch64::*;
#[cfg(target_arch = "aarch64")]
use archmage::{Arm64, arcane, rite};
#[cfg(target_arch = "aarch64")]
use safe_unaligned_simd::aarch64 as safe_simd;
use std::cmp;
use std::ffi::c_int;
use crate::include::common::bitdepth::BitDepth;
use crate::include::common::bitdepth::LeftPixelRow;
use crate::include::dav1d::picture::PicOffset;
use crate::src::align::AlignedVec64;
use crate::src::disjoint_mut::DisjointMut;
use crate::src::looprestoration::{LooprestorationParams, LrEdgeFlags};
#[cfg(target_arch = "aarch64")]
use crate::include::common::bitdepth::BitDepth8;
#[cfg(target_arch = "aarch64")]
use crate::include::common::bitdepth::BitDepth16;
#[cfg(target_arch = "aarch64")]
use crate::include::common::intops::iclip;
#[cfg(target_arch = "aarch64")]
use crate::src::looprestoration::padding;
#[cfg(target_arch = "aarch64")]
use crate::src::strided::Strided as _;
#[cfg(target_arch = "aarch64")]
use crate::src::tables::dav1d_sgr_x_by_x;
#[cfg(target_arch = "aarch64")]
const S: usize = 256 * 3 / 2 + 3 + 3; #[cfg(target_arch = "aarch64")]
const MAXW: usize = 256 * 3 / 2; #[cfg(target_arch = "aarch64")]
const TMP_LEN: usize = (64 + 3 + 3) * S;
#[cfg(target_arch = "aarch64")]
const BOX_LEN: usize = (64 + 2 + 2) * S;
#[cfg(target_arch = "aarch64")]
const DST_LEN: usize = 64 * MAXW;
#[cfg(target_arch = "aarch64")]
fn boxed_zeroed<T: Copy + Default, const N: usize>() -> Box<[T; N]> {
vec![T::default(); N]
.into_boxed_slice()
.try_into()
.ok()
.unwrap()
}
#[cfg(target_arch = "aarch64")]
struct Scratch8 {
tmp: Box<[u8; TMP_LEN]>,
hor: Box<[u16; TMP_LEN]>,
sumsq: Box<[i32; BOX_LEN]>,
sum: Box<[i16; BOX_LEN]>,
d0: Box<[i16; DST_LEN]>,
d1: Box<[i16; DST_LEN]>,
}
#[cfg(target_arch = "aarch64")]
struct Scratch16 {
tmp: Box<[u16; TMP_LEN]>,
hor: Box<[u16; TMP_LEN]>,
sumsq: Box<[i32; BOX_LEN]>,
sum: Box<[i32; BOX_LEN]>,
d0: Box<[i32; DST_LEN]>,
d1: Box<[i32; DST_LEN]>,
}
#[cfg(target_arch = "aarch64")]
impl Scratch8 {
fn new() -> Self {
Self {
tmp: boxed_zeroed(),
hor: boxed_zeroed(),
sumsq: boxed_zeroed(),
sum: boxed_zeroed(),
d0: boxed_zeroed(),
d1: boxed_zeroed(),
}
}
#[cfg(feature = "__lrpoison")]
fn poison(&mut self) {
self.tmp.fill(0xA5);
self.hor.fill(0xA5A5);
self.sumsq.fill(0x5A5A_5A5A);
self.sum.fill(0x5A5A);
self.d0.fill(0x5A5A);
self.d1.fill(0x5A5A);
}
}
#[cfg(target_arch = "aarch64")]
impl Scratch16 {
fn new() -> Self {
Self {
tmp: boxed_zeroed(),
hor: boxed_zeroed(),
sumsq: boxed_zeroed(),
sum: boxed_zeroed(),
d0: boxed_zeroed(),
d1: boxed_zeroed(),
}
}
#[cfg(feature = "__lrpoison")]
fn poison(&mut self) {
self.tmp.fill(0xA5A5);
self.hor.fill(0xA5A5);
self.sumsq.fill(0x5A5A_5A5A);
self.sum.fill(0x5A5A_5A5A);
self.d0.fill(0x5A5A_5A5A);
self.d1.fill(0x5A5A_5A5A);
}
}
#[cfg(target_arch = "aarch64")]
thread_local! {
static SCRATCH8: std::cell::RefCell<Option<Scratch8>> = const { std::cell::RefCell::new(None) };
static SCRATCH16: std::cell::RefCell<Option<Scratch16>> = const { std::cell::RefCell::new(None) };
}
#[cfg(target_arch = "aarch64")]
fn with_scratch8<R>(f: impl FnOnce(&mut Scratch8) -> R) -> R {
SCRATCH8.with(|c| {
let mut b = c.borrow_mut();
let s = b.get_or_insert_with(Scratch8::new);
#[cfg(feature = "__lrpoison")]
s.poison();
f(s)
})
}
#[cfg(target_arch = "aarch64")]
fn with_scratch16<R>(f: impl FnOnce(&mut Scratch16) -> R) -> R {
SCRATCH16.with(|c| {
let mut b = c.borrow_mut();
let s = b.get_or_insert_with(Scratch16::new);
#[cfg(feature = "__lrpoison")]
s.poison();
f(s)
})
}
#[cfg(target_arch = "aarch64")]
#[arcane]
fn wiener_hor_8bpc(
_token: Arm64,
tmp: &[u8; TMP_LEN],
hor: &mut [u16; TMP_LEN],
w: usize,
h: usize,
taps: &[i16; 8],
) {
const BIAS: i32 = (1 << 14) + (1 << 2);
let vbias = vdupq_n_s32(BIAS);
let vmax = vdupq_n_s32((1 << 13) - 1);
let vzero = vdupq_n_s32(0);
let t: [int16x8_t; 7] = core::array::from_fn(|k| vdupq_n_s16(taps[k]));
for row in 0..h + 6 {
let base = row * S;
let src = &tmp[base..base + w + 6];
let dst = &mut hor[base..base + w];
let mut x = 0;
while x + 8 <= w {
let mut lo = vbias;
let mut hi = vbias;
for k in 0..7 {
let v = safe_simd::vld1_u8(src[x + k..][..8].try_into().unwrap());
let v16 = vreinterpretq_s16_u16(vmovl_u8(v));
lo = vmlal_s16(lo, vget_low_s16(v16), vget_low_s16(t[k]));
hi = vmlal_high_s16(hi, v16, t[k]);
}
let lo = vminq_s32(vmaxq_s32(vshrq_n_s32::<3>(lo), vzero), vmax);
let hi = vminq_s32(vmaxq_s32(vshrq_n_s32::<3>(hi), vzero), vmax);
let packed = vcombine_u16(
vmovn_u32(vreinterpretq_u32_s32(lo)),
vmovn_u32(vreinterpretq_u32_s32(hi)),
);
safe_simd::vst1q_u16((&mut dst[x..x + 8]).try_into().unwrap(), packed);
x += 8;
}
while x < w {
let mut sum = BIAS;
for k in 0..7 {
sum += src[x + k] as i32 * taps[k] as i32;
}
dst[x] = iclip(sum >> 3, 0, (1 << 13) - 1) as u16;
x += 1;
}
}
}
#[cfg(target_arch = "aarch64")]
#[arcane]
fn wiener_ver_8bpc(
_token: Arm64,
hor: &[u16; TMP_LEN],
p: PicOffset,
w: usize,
h: usize,
taps: &[i16; 8],
) {
const BIAS: i32 = -(1 << 18) + (1 << 10);
let vbias = vdupq_n_s32(BIAS);
let t: [int16x8_t; 7] = core::array::from_fn(|k| vdupq_n_s16(taps[k]));
let stride = p.pixel_stride::<BitDepth8>();
for j in 0..h {
let mut dst = (p + (j as isize * stride)).slice_mut::<BitDepth8>(w);
let mut x = 0;
while x + 8 <= w {
let mut lo = vbias;
let mut hi = vbias;
for k in 0..7 {
let v = safe_simd::vld1q_u16(hor[(j + k) * S + x..][..8].try_into().unwrap());
let v16 = vreinterpretq_s16_u16(v);
lo = vmlal_s16(lo, vget_low_s16(v16), vget_low_s16(t[k]));
hi = vmlal_high_s16(hi, v16, t[k]);
}
let packed = vqmovun_s16(vcombine_s16(
vqmovn_s32(vshrq_n_s32::<11>(lo)),
vqmovn_s32(vshrq_n_s32::<11>(hi)),
));
safe_simd::vst1_u8((&mut dst[x..x + 8]).try_into().unwrap(), packed);
x += 8;
}
while x < w {
let mut sum = BIAS;
for k in 0..7 {
sum += hor[(j + k) * S + x] as i32 * taps[k] as i32;
}
dst[x] = iclip(sum >> 11, 0, 255) as u8;
x += 1;
}
}
}
#[cfg(target_arch = "aarch64")]
fn wiener_8bpc(
token: Arm64,
p: PicOffset,
left: &[LeftPixelRow<u8>],
lpf: &DisjointMut<AlignedVec64<u8>>,
lpf_off: isize,
w: usize,
h: usize,
params: &LooprestorationParams,
edges: LrEdgeFlags,
) {
with_scratch8(|sc| {
padding::<BitDepth8>(&mut sc.tmp, p, left, lpf, lpf_off, w, h, edges);
let mut hf = params.filter[0];
hf[3] += 128;
wiener_hor_8bpc(token, &sc.tmp, &mut sc.hor, w, h, &hf);
wiener_ver_8bpc(token, &sc.hor, p, w, h, ¶ms.filter[1]);
});
}
#[cfg(target_arch = "aarch64")]
#[arcane]
fn wiener_hor_16bpc(
_token: Arm64,
tmp: &[u16; TMP_LEN],
hor: &mut [u16; TMP_LEN],
w: usize,
h: usize,
taps: &[i16; 8],
bitdepth: i32,
) {
let round_bits_h = 3 + (bitdepth == 12) as i32 * 2;
let bias = (1 << (bitdepth + 6)) + (1 << (round_bits_h - 1));
let clip_max = (1 << (bitdepth + 1 + 7 - round_bits_h)) - 1;
let vbias = vdupq_n_s32(bias);
let vsh = vdupq_n_s32(-round_bits_h);
let vmax = vdupq_n_s32(clip_max);
let vzero = vdupq_n_s32(0);
let t: [int16x8_t; 7] = core::array::from_fn(|k| vdupq_n_s16(taps[k]));
for row in 0..h + 6 {
let base = row * S;
let src = &tmp[base..base + w + 6];
let dst = &mut hor[base..base + w];
let mut x = 0;
while x + 8 <= w {
let mut lo = vbias;
let mut hi = vbias;
for k in 0..7 {
let v = safe_simd::vld1q_u16(src[x + k..][..8].try_into().unwrap());
let v16 = vreinterpretq_s16_u16(v);
lo = vmlal_s16(lo, vget_low_s16(v16), vget_low_s16(t[k]));
hi = vmlal_high_s16(hi, v16, t[k]);
}
let lo = vminq_s32(vmaxq_s32(vshlq_s32(lo, vsh), vzero), vmax);
let hi = vminq_s32(vmaxq_s32(vshlq_s32(hi, vsh), vzero), vmax);
let packed = vcombine_u16(
vmovn_u32(vreinterpretq_u32_s32(lo)),
vmovn_u32(vreinterpretq_u32_s32(hi)),
);
safe_simd::vst1q_u16((&mut dst[x..x + 8]).try_into().unwrap(), packed);
x += 8;
}
while x < w {
let mut sum = bias;
for k in 0..7 {
sum += src[x + k] as i32 * taps[k] as i32;
}
dst[x] = iclip(sum >> round_bits_h, 0, clip_max) as u16;
x += 1;
}
}
}
#[cfg(target_arch = "aarch64")]
#[arcane]
fn wiener_ver_16bpc(
_token: Arm64,
hor: &[u16; TMP_LEN],
p: PicOffset,
w: usize,
h: usize,
taps: &[i16; 8],
bitdepth: i32,
bitdepth_max: i32,
) {
let round_bits_v = 11 - (bitdepth == 12) as i32 * 2;
let bias = -(1 << (bitdepth + round_bits_v - 1)) + (1 << (round_bits_v - 1));
let vbias = vdupq_n_s32(bias);
let vsh = vdupq_n_s32(-round_bits_v);
let vmax = vdupq_n_s32(bitdepth_max);
let vzero = vdupq_n_s32(0);
let t: [int16x8_t; 7] = core::array::from_fn(|k| vdupq_n_s16(taps[k]));
let stride = p.pixel_stride::<BitDepth16>();
for j in 0..h {
let mut dst = (p + (j as isize * stride)).slice_mut::<BitDepth16>(w);
let mut x = 0;
while x + 8 <= w {
let mut lo = vbias;
let mut hi = vbias;
for k in 0..7 {
let v = safe_simd::vld1q_u16(hor[(j + k) * S + x..][..8].try_into().unwrap());
let v16 = vreinterpretq_s16_u16(v);
lo = vmlal_s16(lo, vget_low_s16(v16), vget_low_s16(t[k]));
hi = vmlal_high_s16(hi, v16, t[k]);
}
let lo = vminq_s32(vmaxq_s32(vshlq_s32(lo, vsh), vzero), vmax);
let hi = vminq_s32(vmaxq_s32(vshlq_s32(hi, vsh), vzero), vmax);
let packed = vcombine_u16(
vmovn_u32(vreinterpretq_u32_s32(lo)),
vmovn_u32(vreinterpretq_u32_s32(hi)),
);
safe_simd::vst1q_u16((&mut dst[x..x + 8]).try_into().unwrap(), packed);
x += 8;
}
while x < w {
let mut sum = bias;
for k in 0..7 {
sum += hor[(j + k) * S + x] as i32 * taps[k] as i32;
}
dst[x] = iclip(sum >> round_bits_v, 0, bitdepth_max) as u16;
x += 1;
}
}
}
#[cfg(target_arch = "aarch64")]
fn wiener_16bpc(
token: Arm64,
p: PicOffset,
left: &[LeftPixelRow<u16>],
lpf: &DisjointMut<AlignedVec64<u8>>,
lpf_off: isize,
w: usize,
h: usize,
params: &LooprestorationParams,
edges: LrEdgeFlags,
bitdepth_max: i32,
) {
with_scratch16(|sc| {
padding::<BitDepth16>(&mut sc.tmp, p, left, lpf, lpf_off, w, h, edges);
let bitdepth = if bitdepth_max == 1023 { 10 } else { 12 };
wiener_hor_16bpc(
token,
&sc.tmp,
&mut sc.hor,
w,
h,
¶ms.filter[0],
bitdepth,
);
wiener_ver_16bpc(
token,
&sc.hor,
p,
w,
h,
¶ms.filter[1],
bitdepth,
bitdepth_max,
);
});
}
#[cfg(target_arch = "aarch64")]
#[rite]
fn box_row_8bpc<const N: usize>(
_token: Arm64,
src: &[u8; TMP_LEN],
r: usize,
bw: usize,
vs: &mut [u16; S],
vq: &mut [u32; S],
out_sum: &mut [i16],
out_sq: &mut [i32],
) {
let top = r - (N == 5) as usize;
let mut x = 0;
while x + 8 <= bw {
let mut s = vdupq_n_u16(0);
let mut q0 = vdupq_n_u32(0);
let mut q1 = vdupq_n_u32(0);
for dy in 0..N {
let v = safe_simd::vld1_u8(src[(top + dy) * S + x..][..8].try_into().unwrap());
let v16 = vmovl_u8(v);
s = vaddq_u16(s, v16);
q0 = vmlal_u16(q0, vget_low_u16(v16), vget_low_u16(v16));
q1 = vmlal_high_u16(q1, v16, v16);
}
safe_simd::vst1q_u16((&mut vs[x..x + 8]).try_into().unwrap(), s);
safe_simd::vst1q_u32((&mut vq[x..x + 4]).try_into().unwrap(), q0);
safe_simd::vst1q_u32((&mut vq[x + 4..x + 8]).try_into().unwrap(), q1);
x += 8;
}
while x < bw {
let mut s = 0u16;
let mut q = 0u32;
for dy in 0..N {
let v = src[(top + dy) * S + x] as u16;
s += v;
q += v as u32 * v as u32;
}
vs[x] = s;
vq[x] = q;
x += 1;
}
let half = N / 2;
let mut x = 2;
while x + 8 <= bw - 2 {
let mut s = vdupq_n_u16(0);
let mut q0 = vdupq_n_u32(0);
let mut q1 = vdupq_n_u32(0);
for dx in 0..N {
s = vaddq_u16(
s,
safe_simd::vld1q_u16(vs[x + dx - half..][..8].try_into().unwrap()),
);
q0 = vaddq_u32(
q0,
safe_simd::vld1q_u32(vq[x + dx - half..][..4].try_into().unwrap()),
);
q1 = vaddq_u32(
q1,
safe_simd::vld1q_u32(vq[x + dx - half + 4..][..4].try_into().unwrap()),
);
}
safe_simd::vst1q_s16(
(&mut out_sum[x..x + 8]).try_into().unwrap(),
vreinterpretq_s16_u16(s),
);
safe_simd::vst1q_s32(
(&mut out_sq[x..x + 4]).try_into().unwrap(),
vreinterpretq_s32_u32(q0),
);
safe_simd::vst1q_s32(
(&mut out_sq[x + 4..x + 8]).try_into().unwrap(),
vreinterpretq_s32_u32(q1),
);
x += 8;
}
while x < bw - 2 {
let mut s = 0u16;
let mut q = 0u32;
for dx in 0..N {
s = s.wrapping_add(vs[x + dx - half]);
q = q.wrapping_add(vq[x + dx - half]);
}
out_sum[x] = s as i16;
out_sq[x] = q as i32;
x += 1;
}
}
#[cfg(target_arch = "aarch64")]
#[rite]
fn box_row_16bpc<const N: usize>(
_token: Arm64,
src: &[u16; TMP_LEN],
r: usize,
bw: usize,
vs: &mut [u32; S],
vq: &mut [u32; S],
out_sum: &mut [i32],
out_sq: &mut [i32],
) {
let top = r - (N == 5) as usize;
let mut x = 0;
while x + 8 <= bw {
let mut s0 = vdupq_n_u32(0);
let mut s1 = vdupq_n_u32(0);
let mut q0 = vdupq_n_u32(0);
let mut q1 = vdupq_n_u32(0);
for dy in 0..N {
let v = safe_simd::vld1q_u16(src[(top + dy) * S + x..][..8].try_into().unwrap());
s0 = vaddw_u16(s0, vget_low_u16(v));
s1 = vaddw_high_u16(s1, v);
q0 = vmlal_u16(q0, vget_low_u16(v), vget_low_u16(v));
q1 = vmlal_high_u16(q1, v, v);
}
safe_simd::vst1q_u32((&mut vs[x..x + 4]).try_into().unwrap(), s0);
safe_simd::vst1q_u32((&mut vs[x + 4..x + 8]).try_into().unwrap(), s1);
safe_simd::vst1q_u32((&mut vq[x..x + 4]).try_into().unwrap(), q0);
safe_simd::vst1q_u32((&mut vq[x + 4..x + 8]).try_into().unwrap(), q1);
x += 8;
}
while x < bw {
let mut s = 0u32;
let mut q = 0u32;
for dy in 0..N {
let v = src[(top + dy) * S + x] as u32;
s += v;
q += v * v;
}
vs[x] = s;
vq[x] = q;
x += 1;
}
let half = N / 2;
let mut x = 2;
while x + 4 <= bw - 2 {
let mut s = vdupq_n_u32(0);
let mut q = vdupq_n_u32(0);
for dx in 0..N {
s = vaddq_u32(
s,
safe_simd::vld1q_u32(vs[x + dx - half..][..4].try_into().unwrap()),
);
q = vaddq_u32(
q,
safe_simd::vld1q_u32(vq[x + dx - half..][..4].try_into().unwrap()),
);
}
safe_simd::vst1q_s32(
(&mut out_sum[x..x + 4]).try_into().unwrap(),
vreinterpretq_s32_u32(s),
);
safe_simd::vst1q_s32(
(&mut out_sq[x..x + 4]).try_into().unwrap(),
vreinterpretq_s32_u32(q),
);
x += 4;
}
while x < bw - 2 {
let mut s = 0u32;
let mut q = 0u32;
for dx in 0..N {
s = s.wrapping_add(vs[x + dx - half]);
q = q.wrapping_add(vq[x + dx - half]);
}
out_sum[x] = s as i32;
out_sq[x] = q as i32;
x += 1;
}
}
#[cfg(target_arch = "aarch64")]
#[rite]
fn sgr_lut16(_token: Arm64, idx: uint8x16_t) -> uint8x16_t {
let t: &[u8; 256] = &dav1d_sgr_x_by_x;
let q0 = uint8x16x4_t(
safe_simd::vld1q_u8(t[0..16].try_into().unwrap()),
safe_simd::vld1q_u8(t[16..32].try_into().unwrap()),
safe_simd::vld1q_u8(t[32..48].try_into().unwrap()),
safe_simd::vld1q_u8(t[48..64].try_into().unwrap()),
);
let q1 = uint8x16x4_t(
safe_simd::vld1q_u8(t[64..80].try_into().unwrap()),
safe_simd::vld1q_u8(t[80..96].try_into().unwrap()),
safe_simd::vld1q_u8(t[96..112].try_into().unwrap()),
safe_simd::vld1q_u8(t[112..128].try_into().unwrap()),
);
let q2 = uint8x16x4_t(
safe_simd::vld1q_u8(t[128..144].try_into().unwrap()),
safe_simd::vld1q_u8(t[144..160].try_into().unwrap()),
safe_simd::vld1q_u8(t[160..176].try_into().unwrap()),
safe_simd::vld1q_u8(t[176..192].try_into().unwrap()),
);
let q3 = uint8x16x4_t(
safe_simd::vld1q_u8(t[192..208].try_into().unwrap()),
safe_simd::vld1q_u8(t[208..224].try_into().unwrap()),
safe_simd::vld1q_u8(t[224..240].try_into().unwrap()),
safe_simd::vld1q_u8(t[240..256].try_into().unwrap()),
);
let r0 = vqtbl4q_u8(q0, idx);
let r1 = vqtbl4q_u8(q1, vsubq_u8(idx, vdupq_n_u8(64)));
let r2 = vqtbl4q_u8(q2, vsubq_u8(idx, vdupq_n_u8(128)));
let r3 = vqtbl4q_u8(q3, vsubq_u8(idx, vdupq_n_u8(192)));
vorrq_u8(vorrq_u8(r0, r1), vorrq_u8(r2, r3))
}
#[cfg(target_arch = "aarch64")]
#[rite]
fn sgr_z(_token: Arm64, a: int32x4_t, b: int32x4_t, n: i32, s: u32) -> uint32x4_t {
let p = vmaxq_s32(
vsubq_s32(vmulq_n_s32(a, n), vmulq_s32(b, b)),
vdupq_n_s32(0),
);
let p = vreinterpretq_u32_s32(p);
let z = vshrq_n_u32::<20>(vaddq_u32(
vmulq_u32(p, vdupq_n_u32(s)),
vdupq_n_u32(1 << 19),
));
vminq_u32(z, vdupq_n_u32(255))
}
#[cfg(target_arch = "aarch64")]
#[rite]
fn sgr_aa(_token: Arm64, x: uint32x4_t, b: int32x4_t, one_by_x: u32) -> int32x4_t {
let prod = vmulq_u32(
vmulq_u32(x, vreinterpretq_u32_s32(b)),
vdupq_n_u32(one_by_x),
);
vreinterpretq_s32_u32(vshrq_n_u32::<12>(vaddq_u32(prod, vdupq_n_u32(1 << 11))))
}
#[cfg(target_arch = "aarch64")]
#[rite]
fn sgr_pack_idx(_token: Arm64, z: [uint32x4_t; 4]) -> uint8x16_t {
let a = vcombine_u16(vmovn_u32(z[0]), vmovn_u32(z[1]));
let b = vcombine_u16(vmovn_u32(z[2]), vmovn_u32(z[3]));
vcombine_u8(vmovn_u16(a), vmovn_u16(b))
}
#[cfg(target_arch = "aarch64")]
#[rite]
fn sgr_unpack_x(_token: Arm64, x: uint8x16_t) -> [uint32x4_t; 4] {
let lo = vmovl_u8(vget_low_u8(x));
let hi = vmovl_high_u8(x);
[
vmovl_u16(vget_low_u16(lo)),
vmovl_high_u16(lo),
vmovl_u16(vget_low_u16(hi)),
vmovl_high_u16(hi),
]
}
#[cfg(target_arch = "aarch64")]
#[arcane]
fn sgr_ab_8bpc(
token: Arm64,
sumsq: &mut [i32; BOX_LEN],
sum: &mut [i16; BOX_LEN],
w: usize,
h: usize,
n: i32,
s: u32,
one_by_x: u32,
step: usize,
) {
let cols = w + 2;
let mut row = 0;
while row < h + 2 {
let base = (row + 1) * S + 2;
#[cfg(debug_assertions)]
for (i, (&a_val, &b_val)) in sumsq[base..base + cols]
.iter()
.zip(sum[base..base + cols].iter())
.enumerate()
{
let b_val = b_val as i32;
debug_assert!(
(0..=n * 255 * 255).contains(&a_val) && (0..=n * 255).contains(&b_val),
"sgr_ab_8bpc: box sums out of range at row {row} col {i} (w={w} h={h} n={n}): \
a={a_val} (max {}), b={b_val} (max {})",
n * 255 * 255,
n * 255,
);
}
let mut i = 0;
while i + 16 <= cols {
let mut z = [vdupq_n_u32(0); 4];
let mut bs = [vdupq_n_s32(0); 4];
for g in 0..4 {
let a = safe_simd::vld1q_s32(sumsq[base + i + g * 4..][..4].try_into().unwrap());
let b = vmovl_s16(safe_simd::vld1_s16(
sum[base + i + g * 4..][..4].try_into().unwrap(),
));
bs[g] = b;
z[g] = sgr_z(token, a, b, n, s);
}
let xs = sgr_unpack_x(token, sgr_lut16(token, sgr_pack_idx(token, z)));
for g in 0..4 {
let aa = sgr_aa(token, xs[g], bs[g], one_by_x);
safe_simd::vst1q_s32(
(&mut sumsq[base + i + g * 4..][..4]).try_into().unwrap(),
aa,
);
safe_simd::vst1_s16(
(&mut sum[base + i + g * 4..][..4]).try_into().unwrap(),
vmovn_s32(vreinterpretq_s32_u32(xs[g])),
);
}
i += 16;
}
while i < cols {
let idx = base + i;
let a_val = sumsq[idx];
let b_val = sum[idx] as i32;
let p = cmp::max(a_val * n - b_val * b_val, 0) as u32;
let z = (p.wrapping_mul(s).wrapping_add(1 << 19)) >> 20;
let x = dav1d_sgr_x_by_x[cmp::min(z, 255) as usize] as u32;
sumsq[idx] = ((x.wrapping_mul(b_val as u32).wrapping_mul(one_by_x))
.wrapping_add(1 << 11)
>> 12) as i32;
sum[idx] = x as i16;
i += 1;
}
row += step;
}
}
#[cfg(target_arch = "aarch64")]
#[rite]
fn six_i32(_token: Arm64, p: &[i32], i: usize) -> int32x4_t {
let up = safe_simd::vld1q_s32(p[i - S..][..4].try_into().unwrap());
let dn = safe_simd::vld1q_s32(p[i + S..][..4].try_into().unwrap());
let upl = safe_simd::vld1q_s32(p[i - S - 1..][..4].try_into().unwrap());
let upr = safe_simd::vld1q_s32(p[i - S + 1..][..4].try_into().unwrap());
let dnl = safe_simd::vld1q_s32(p[i + S - 1..][..4].try_into().unwrap());
let dnr = safe_simd::vld1q_s32(p[i + S + 1..][..4].try_into().unwrap());
vmlaq_n_s32(
vmulq_n_s32(vaddq_s32(up, dn), 6),
vaddq_s32(vaddq_s32(upl, upr), vaddq_s32(dnl, dnr)),
5,
)
}
#[cfg(target_arch = "aarch64")]
#[rite]
fn mid_i32(_token: Arm64, p: &[i32], i: usize) -> int32x4_t {
let c = safe_simd::vld1q_s32(p[i..][..4].try_into().unwrap());
let l = safe_simd::vld1q_s32(p[i - 1..][..4].try_into().unwrap());
let r = safe_simd::vld1q_s32(p[i + 1..][..4].try_into().unwrap());
vmlaq_n_s32(vmulq_n_s32(c, 6), vaddq_s32(l, r), 5)
}
#[cfg(target_arch = "aarch64")]
#[rite]
fn eight_i32(_token: Arm64, p: &[i32], i: usize) -> int32x4_t {
let c = safe_simd::vld1q_s32(p[i..][..4].try_into().unwrap());
let l = safe_simd::vld1q_s32(p[i - 1..][..4].try_into().unwrap());
let r = safe_simd::vld1q_s32(p[i + 1..][..4].try_into().unwrap());
let up = safe_simd::vld1q_s32(p[i - S..][..4].try_into().unwrap());
let dn = safe_simd::vld1q_s32(p[i + S..][..4].try_into().unwrap());
let upl = safe_simd::vld1q_s32(p[i - S - 1..][..4].try_into().unwrap());
let upr = safe_simd::vld1q_s32(p[i - S + 1..][..4].try_into().unwrap());
let dnl = safe_simd::vld1q_s32(p[i + S - 1..][..4].try_into().unwrap());
let dnr = safe_simd::vld1q_s32(p[i + S + 1..][..4].try_into().unwrap());
vmlaq_n_s32(
vmulq_n_s32(
vaddq_s32(vaddq_s32(vaddq_s32(c, l), vaddq_s32(r, up)), dn),
4,
),
vaddq_s32(vaddq_s32(upl, upr), vaddq_s32(dnl, dnr)),
3,
)
}
#[cfg(target_arch = "aarch64")]
#[rite]
fn ld4_i16(_token: Arm64, p: &[i16], i: usize) -> int32x4_t {
vmovl_s16(safe_simd::vld1_s16(p[i..][..4].try_into().unwrap()))
}
#[cfg(target_arch = "aarch64")]
#[rite]
fn six_i16(token: Arm64, p: &[i16], i: usize) -> int32x4_t {
let up = ld4_i16(token, p, i - S);
let dn = ld4_i16(token, p, i + S);
let upl = ld4_i16(token, p, i - S - 1);
let upr = ld4_i16(token, p, i - S + 1);
let dnl = ld4_i16(token, p, i + S - 1);
let dnr = ld4_i16(token, p, i + S + 1);
vmlaq_n_s32(
vmulq_n_s32(vaddq_s32(up, dn), 6),
vaddq_s32(vaddq_s32(upl, upr), vaddq_s32(dnl, dnr)),
5,
)
}
#[cfg(target_arch = "aarch64")]
#[rite]
fn mid_i16(token: Arm64, p: &[i16], i: usize) -> int32x4_t {
let c = ld4_i16(token, p, i);
let l = ld4_i16(token, p, i - 1);
let r = ld4_i16(token, p, i + 1);
vmlaq_n_s32(vmulq_n_s32(c, 6), vaddq_s32(l, r), 5)
}
#[cfg(target_arch = "aarch64")]
#[rite]
fn eight_i16(token: Arm64, p: &[i16], i: usize) -> int32x4_t {
let c = ld4_i16(token, p, i);
let l = ld4_i16(token, p, i - 1);
let r = ld4_i16(token, p, i + 1);
let up = ld4_i16(token, p, i - S);
let dn = ld4_i16(token, p, i + S);
let upl = ld4_i16(token, p, i - S - 1);
let upr = ld4_i16(token, p, i - S + 1);
let dnl = ld4_i16(token, p, i + S - 1);
let dnr = ld4_i16(token, p, i + S + 1);
vmlaq_n_s32(
vmulq_n_s32(
vaddq_s32(vaddq_s32(vaddq_s32(c, l), vaddq_s32(r, up)), dn),
4,
),
vaddq_s32(vaddq_s32(upl, upr), vaddq_s32(dnl, dnr)),
3,
)
}
#[cfg(target_arch = "aarch64")]
fn selfguided_8bpc(
token: Arm64,
dst: &mut [i16; DST_LEN],
src: &[u8; TMP_LEN],
w: usize,
h: usize,
n: i32,
s: u32,
sumsq: &mut [i32; BOX_LEN],
sum: &mut [i16; BOX_LEN],
) {
let one_by_x: u32 = if n == 25 { 164 } else { 455 };
let step = if n == 25 { 2 } else { 1 };
let (bw, bh) = (w + 6, h + 6);
boxsum_8bpc(token, sumsq, sum, src, bw, bh, n);
sgr_ab_8bpc(token, sumsq, sum, w, h, n, s, one_by_x, step);
sgr_out_8bpc(token, dst, src, sumsq, sum, w, h, n);
}
#[cfg(target_arch = "aarch64")]
#[arcane]
fn boxsum_8bpc(
token: Arm64,
sumsq: &mut [i32; BOX_LEN],
sum: &mut [i16; BOX_LEN],
src: &[u8; TMP_LEN],
bw: usize,
bh: usize,
n: i32,
) {
let mut vs = [0u16; S];
let mut vq = [0u32; S];
for r in 1..=bh - 4 {
let (os, oq) = (&mut sum[r * S..r * S + bw], &mut sumsq[r * S..r * S + bw]);
if n == 25 {
box_row_8bpc::<5>(token, src, r, bw, &mut vs, &mut vq, os, oq);
} else {
box_row_8bpc::<3>(token, src, r, bw, &mut vs, &mut vq, os, oq);
}
}
}
#[cfg(target_arch = "aarch64")]
#[arcane]
fn sgr_out_8bpc(
token: Arm64,
dst: &mut [i16; DST_LEN],
src: &[u8; TMP_LEN],
sumsq: &[i32; BOX_LEN],
sum: &[i16; BOX_LEN],
w: usize,
h: usize,
n: i32,
) {
let base = 2 * S + 3;
let src_base = 3 * S + 3;
macro_rules! emit {
($bv:expr, $av:expr, $sidx:expr, $didx:expr, $rnd:expr, $sh:expr) => {{
let px = vmovl_u16(vget_low_u16(vmovl_u8(safe_simd::vld1_u8(
src[$sidx..][..8].try_into().unwrap(),
))));
let v = vaddq_s32(
vsubq_s32($bv, vmulq_s32($av, vreinterpretq_s32_u32(px))),
vdupq_n_s32($rnd),
);
safe_simd::vst1_s16(
(&mut dst[$didx..][..4]).try_into().unwrap(),
vmovn_s32(vshrq_n_s32::<$sh>(v)),
);
}};
}
if n == 25 {
let mut j = 0;
while j + 1 < h {
for phase in 0..2 {
let rowa = base + (j + phase) * S;
let sidx0 = src_base + (j + phase) * S;
let didx0 = (j + phase) * MAXW;
let mut i = 0;
while i + 4 <= w {
let (bv, av) = if phase == 0 {
(
six_i32(token, sumsq, rowa + i),
six_i16(token, sum, rowa + i),
)
} else {
(
mid_i32(token, sumsq, rowa + i),
mid_i16(token, sum, rowa + i),
)
};
if phase == 0 {
emit!(bv, av, sidx0 + i, didx0 + i, 1 << 8, 9);
} else {
emit!(bv, av, sidx0 + i, didx0 + i, 1 << 7, 8);
}
i += 4;
}
while i < w {
let (b, a) = if phase == 0 {
(six_s(sumsq, rowa + i), six_s16(sum, rowa + i))
} else {
(mid_s(sumsq, rowa + i), mid_s16(sum, rowa + i))
};
let px = src[sidx0 + i] as i32;
dst[didx0 + i] = if phase == 0 {
((b - a * px + (1 << 8)) >> 9) as i16
} else {
((b - a * px + (1 << 7)) >> 8) as i16
};
i += 1;
}
}
j += 2;
}
if j + 1 == h {
let rowa = base + j * S;
let sidx0 = src_base + j * S;
let didx0 = j * MAXW;
let mut i = 0;
while i + 4 <= w {
let bv = six_i32(token, sumsq, rowa + i);
let av = six_i16(token, sum, rowa + i);
emit!(bv, av, sidx0 + i, didx0 + i, 1 << 8, 9);
i += 4;
}
while i < w {
let b = six_s(sumsq, rowa + i);
let a = six_s16(sum, rowa + i);
dst[didx0 + i] = ((b - a * src[sidx0 + i] as i32 + (1 << 8)) >> 9) as i16;
i += 1;
}
}
} else {
for j in 0..h {
let rowa = base + j * S;
let sidx0 = src_base + j * S;
let didx0 = j * MAXW;
let mut i = 0;
while i + 4 <= w {
let bv = eight_i32(token, sumsq, rowa + i);
let av = eight_i16(token, sum, rowa + i);
emit!(bv, av, sidx0 + i, didx0 + i, 1 << 8, 9);
i += 4;
}
while i < w {
let b = eight_s(sumsq, rowa + i);
let a = eight_s16(sum, rowa + i);
dst[didx0 + i] = ((b - a * src[sidx0 + i] as i32 + (1 << 8)) >> 9) as i16;
i += 1;
}
}
}
}
#[cfg(target_arch = "aarch64")]
fn six_s(p: &[i32], i: usize) -> i32 {
(p[i - S] + p[i + S]) * 6 + (p[i - S - 1] + p[i - S + 1] + p[i + S - 1] + p[i + S + 1]) * 5
}
#[cfg(target_arch = "aarch64")]
fn mid_s(p: &[i32], i: usize) -> i32 {
p[i] * 6 + (p[i - 1] + p[i + 1]) * 5
}
#[cfg(target_arch = "aarch64")]
fn eight_s(p: &[i32], i: usize) -> i32 {
(p[i] + p[i - 1] + p[i + 1] + p[i - S] + p[i + S]) * 4
+ (p[i - S - 1] + p[i - S + 1] + p[i + S - 1] + p[i + S + 1]) * 3
}
#[cfg(target_arch = "aarch64")]
fn six_s16(p: &[i16], i: usize) -> i32 {
(p[i - S] as i32 + p[i + S] as i32) * 6
+ (p[i - S - 1] as i32 + p[i - S + 1] as i32 + p[i + S - 1] as i32 + p[i + S + 1] as i32)
* 5
}
#[cfg(target_arch = "aarch64")]
fn mid_s16(p: &[i16], i: usize) -> i32 {
p[i] as i32 * 6 + (p[i - 1] as i32 + p[i + 1] as i32) * 5
}
#[cfg(target_arch = "aarch64")]
fn eight_s16(p: &[i16], i: usize) -> i32 {
(p[i] as i32 + p[i - 1] as i32 + p[i + 1] as i32 + p[i - S] as i32 + p[i + S] as i32) * 4
+ (p[i - S - 1] as i32 + p[i - S + 1] as i32 + p[i + S - 1] as i32 + p[i + S + 1] as i32)
* 3
}
#[cfg(target_arch = "aarch64")]
#[arcane]
fn sgr_apply_8bpc(
_token: Arm64,
p: PicOffset,
w: usize,
h: usize,
d0: &[i16; DST_LEN],
w0: i32,
d1: Option<&[i16; DST_LEN]>,
w1: i32,
) {
let stride = p.pixel_stride::<BitDepth8>();
for j in 0..h {
let mut row = (p + (j as isize * stride)).slice_mut::<BitDepth8>(w);
let mut i = 0;
while i + 8 <= w {
let a = safe_simd::vld1q_s16(d0[j * MAXW + i..][..8].try_into().unwrap());
let mut lo = vmulq_n_s32(vmovl_s16(vget_low_s16(a)), w0);
let mut hi = vmulq_n_s32(vmovl_high_s16(a), w0);
if let Some(d1) = d1 {
let b = safe_simd::vld1q_s16(d1[j * MAXW + i..][..8].try_into().unwrap());
lo = vmlaq_n_s32(lo, vmovl_s16(vget_low_s16(b)), w1);
hi = vmlaq_n_s32(hi, vmovl_high_s16(b), w1);
}
let lo = vshrq_n_s32::<11>(vaddq_s32(lo, vdupq_n_s32(1 << 10)));
let hi = vshrq_n_s32::<11>(vaddq_s32(hi, vdupq_n_s32(1 << 10)));
let add = vcombine_s16(vmovn_s32(lo), vmovn_s32(hi));
let px = vreinterpretq_s16_u16(vmovl_u8(safe_simd::vld1_u8(
row[i..][..8].try_into().unwrap(),
)));
safe_simd::vst1_u8(
(&mut row[i..i + 8]).try_into().unwrap(),
vqmovun_s16(vaddq_s16(px, add)),
);
i += 8;
}
while i < w {
let mut v = w0 * d0[j * MAXW + i] as i32;
if let Some(d1) = d1 {
v += w1 * d1[j * MAXW + i] as i32;
}
row[i] = iclip(row[i] as i32 + ((v + (1 << 10)) >> 11), 0, 255) as u8;
i += 1;
}
}
}
#[cfg(target_arch = "aarch64")]
fn sgr_8bpc(
token: Arm64,
p: PicOffset,
left: &[LeftPixelRow<u8>],
lpf: &DisjointMut<AlignedVec64<u8>>,
lpf_off: isize,
w: usize,
h: usize,
params: &LooprestorationParams,
edges: LrEdgeFlags,
variant: usize,
) {
with_scratch8(|sc| {
padding::<BitDepth8>(&mut sc.tmp, p, left, lpf, lpf_off, w, h, edges);
let sgr = params.sgr();
match variant {
2 => {
selfguided_8bpc(
token,
&mut sc.d0,
&sc.tmp,
w,
h,
25,
sgr.s0,
&mut sc.sumsq,
&mut sc.sum,
);
sgr_apply_8bpc(token, p, w, h, &sc.d0, sgr.w0 as i32, None, 0);
}
3 => {
selfguided_8bpc(
token,
&mut sc.d0,
&sc.tmp,
w,
h,
9,
sgr.s1,
&mut sc.sumsq,
&mut sc.sum,
);
sgr_apply_8bpc(token, p, w, h, &sc.d0, sgr.w1 as i32, None, 0);
}
_ => {
selfguided_8bpc(
token,
&mut sc.d0,
&sc.tmp,
w,
h,
25,
sgr.s0,
&mut sc.sumsq,
&mut sc.sum,
);
selfguided_8bpc(
token,
&mut sc.d1,
&sc.tmp,
w,
h,
9,
sgr.s1,
&mut sc.sumsq,
&mut sc.sum,
);
sgr_apply_8bpc(
token,
p,
w,
h,
&sc.d0,
sgr.w0 as i32,
Some(&sc.d1),
sgr.w1 as i32,
);
}
}
});
}
#[cfg(target_arch = "aarch64")]
#[arcane]
fn sgr_ab_16bpc(
token: Arm64,
sumsq: &mut [i32; BOX_LEN],
sum: &mut [i32; BOX_LEN],
w: usize,
h: usize,
n: i32,
s: u32,
one_by_x: u32,
step: usize,
bdm8: i32,
) {
let cols = w + 2;
let va_rnd = vdupq_n_s32((1 << (2 * bdm8)) >> 1);
let vb_rnd = vdupq_n_s32((1 << bdm8) >> 1);
let va_sh = vdupq_n_s32(-2 * bdm8);
let vb_sh = vdupq_n_s32(-bdm8);
let mut row = 0;
while row < h + 2 {
let base = (row + 1) * S + 2;
#[cfg(debug_assertions)]
{
let px_max = (1i32 << (8 + bdm8)) - 1;
for (i, (&a_raw, &b_raw)) in sumsq[base..base + cols]
.iter()
.zip(sum[base..base + cols].iter())
.enumerate()
{
debug_assert!(
(0..=n * px_max * px_max).contains(&a_raw) && (0..=n * px_max).contains(&b_raw),
"sgr_ab_16bpc: box sums out of range at row {row} col {i} \
(w={w} h={h} n={n} bdm8={bdm8}): a={a_raw} (max {}), b={b_raw} (max {})",
n * px_max * px_max,
n * px_max,
);
}
}
let mut i = 0;
while i + 16 <= cols {
let mut z = [vdupq_n_u32(0); 4];
let mut braw = [vdupq_n_s32(0); 4];
for g in 0..4 {
let a_raw =
safe_simd::vld1q_s32(sumsq[base + i + g * 4..][..4].try_into().unwrap());
let b_raw = safe_simd::vld1q_s32(sum[base + i + g * 4..][..4].try_into().unwrap());
braw[g] = b_raw;
let a = vshlq_s32(vaddq_s32(a_raw, va_rnd), va_sh);
let b = vshlq_s32(vaddq_s32(b_raw, vb_rnd), vb_sh);
z[g] = sgr_z(token, a, b, n, s);
}
let xs = sgr_unpack_x(token, sgr_lut16(token, sgr_pack_idx(token, z)));
for g in 0..4 {
let aa = sgr_aa(token, xs[g], braw[g], one_by_x);
safe_simd::vst1q_s32(
(&mut sumsq[base + i + g * 4..][..4]).try_into().unwrap(),
aa,
);
safe_simd::vst1q_s32(
(&mut sum[base + i + g * 4..][..4]).try_into().unwrap(),
vreinterpretq_s32_u32(xs[g]),
);
}
i += 16;
}
while i < cols {
let idx = base + i;
let a_raw = sumsq[idx];
let b_raw = sum[idx];
let a = (a_raw + ((1 << (2 * bdm8)) >> 1)) >> (2 * bdm8);
let b = (b_raw + ((1 << bdm8) >> 1)) >> bdm8;
let p = cmp::max(a * n - b * b, 0) as u32;
let z = (p.wrapping_mul(s).wrapping_add(1 << 19)) >> 20;
let x = dav1d_sgr_x_by_x[cmp::min(z, 255) as usize] as u32;
sumsq[idx] = ((x.wrapping_mul(b_raw as u32).wrapping_mul(one_by_x))
.wrapping_add(1 << 11)
>> 12) as i32;
sum[idx] = x as i32;
i += 1;
}
row += step;
}
}
#[cfg(target_arch = "aarch64")]
#[arcane]
fn sgr_out_16bpc(
token: Arm64,
dst: &mut [i32; DST_LEN],
src: &[u16; TMP_LEN],
sumsq: &[i32; BOX_LEN],
sum: &[i32; BOX_LEN],
w: usize,
h: usize,
n: i32,
) {
let base = 2 * S + 3;
let src_base = 3 * S + 3;
macro_rules! emit {
($bv:expr, $av:expr, $sidx:expr, $didx:expr, $rnd:expr, $sh:expr) => {{
let px = vreinterpretq_s32_u32(vmovl_u16(safe_simd::vld1_u16(
src[$sidx..][..4].try_into().unwrap(),
)));
let v = vaddq_s32(vsubq_s32($bv, vmulq_s32($av, px)), vdupq_n_s32($rnd));
safe_simd::vst1q_s32(
(&mut dst[$didx..][..4]).try_into().unwrap(),
vshrq_n_s32::<$sh>(v),
);
}};
}
if n == 25 {
let mut j = 0;
while j + 1 < h {
for phase in 0..2 {
let rowa = base + (j + phase) * S;
let sidx0 = src_base + (j + phase) * S;
let didx0 = (j + phase) * MAXW;
let mut i = 0;
while i + 4 <= w {
let (bv, av) = if phase == 0 {
(
six_i32(token, sumsq, rowa + i),
six_i32(token, sum, rowa + i),
)
} else {
(
mid_i32(token, sumsq, rowa + i),
mid_i32(token, sum, rowa + i),
)
};
if phase == 0 {
emit!(bv, av, sidx0 + i, didx0 + i, 1 << 8, 9);
} else {
emit!(bv, av, sidx0 + i, didx0 + i, 1 << 7, 8);
}
i += 4;
}
while i < w {
let (b, a) = if phase == 0 {
(six_s(sumsq, rowa + i), six_s(sum, rowa + i))
} else {
(mid_s(sumsq, rowa + i), mid_s(sum, rowa + i))
};
let px = src[sidx0 + i] as i32;
dst[didx0 + i] = if phase == 0 {
(b - a * px + (1 << 8)) >> 9
} else {
(b - a * px + (1 << 7)) >> 8
};
i += 1;
}
}
j += 2;
}
if j + 1 == h {
let rowa = base + j * S;
let sidx0 = src_base + j * S;
let didx0 = j * MAXW;
let mut i = 0;
while i + 4 <= w {
let bv = six_i32(token, sumsq, rowa + i);
let av = six_i32(token, sum, rowa + i);
emit!(bv, av, sidx0 + i, didx0 + i, 1 << 8, 9);
i += 4;
}
while i < w {
let b = six_s(sumsq, rowa + i);
let a = six_s(sum, rowa + i);
dst[didx0 + i] = (b - a * src[sidx0 + i] as i32 + (1 << 8)) >> 9;
i += 1;
}
}
} else {
for j in 0..h {
let rowa = base + j * S;
let sidx0 = src_base + j * S;
let didx0 = j * MAXW;
let mut i = 0;
while i + 4 <= w {
let bv = eight_i32(token, sumsq, rowa + i);
let av = eight_i32(token, sum, rowa + i);
emit!(bv, av, sidx0 + i, didx0 + i, 1 << 8, 9);
i += 4;
}
while i < w {
let b = eight_s(sumsq, rowa + i);
let a = eight_s(sum, rowa + i);
dst[didx0 + i] = (b - a * src[sidx0 + i] as i32 + (1 << 8)) >> 9;
i += 1;
}
}
}
}
#[cfg(target_arch = "aarch64")]
#[arcane]
fn sgr_apply_16bpc(
_token: Arm64,
p: PicOffset,
w: usize,
h: usize,
d0: &[i32; DST_LEN],
w0: i32,
d1: Option<&[i32; DST_LEN]>,
w1: i32,
bitdepth_max: i32,
) {
let stride = p.pixel_stride::<BitDepth16>();
let vmax = vdupq_n_s32(bitdepth_max);
let vzero = vdupq_n_s32(0);
for j in 0..h {
let mut row = (p + (j as isize * stride)).slice_mut::<BitDepth16>(w);
let mut i = 0;
while i + 4 <= w {
let a = safe_simd::vld1q_s32(d0[j * MAXW + i..][..4].try_into().unwrap());
let mut v = vmulq_n_s32(a, w0);
if let Some(d1) = d1 {
let b = safe_simd::vld1q_s32(d1[j * MAXW + i..][..4].try_into().unwrap());
v = vmlaq_n_s32(v, b, w1);
}
let add = vshrq_n_s32::<11>(vaddq_s32(v, vdupq_n_s32(1 << 10)));
let px = vreinterpretq_s32_u32(vmovl_u16(safe_simd::vld1_u16(
row[i..][..4].try_into().unwrap(),
)));
let out = vminq_s32(vmaxq_s32(vaddq_s32(px, add), vzero), vmax);
safe_simd::vst1_u16(
(&mut row[i..i + 4]).try_into().unwrap(),
vmovn_u32(vreinterpretq_u32_s32(out)),
);
i += 4;
}
while i < w {
let mut v = w0 * d0[j * MAXW + i];
if let Some(d1) = d1 {
v += w1 * d1[j * MAXW + i];
}
row[i] = iclip(row[i] as i32 + ((v + (1 << 10)) >> 11), 0, bitdepth_max) as u16;
i += 1;
}
}
}
#[cfg(target_arch = "aarch64")]
fn selfguided_16bpc(
token: Arm64,
dst: &mut [i32; DST_LEN],
src: &[u16; TMP_LEN],
w: usize,
h: usize,
n: i32,
s: u32,
bdm8: i32,
sumsq: &mut [i32; BOX_LEN],
sum: &mut [i32; BOX_LEN],
) {
let one_by_x: u32 = if n == 25 { 164 } else { 455 };
let step = if n == 25 { 2 } else { 1 };
let (bw, bh) = (w + 6, h + 6);
boxsum_16bpc(token, sumsq, sum, src, bw, bh, n);
sgr_ab_16bpc(token, sumsq, sum, w, h, n, s, one_by_x, step, bdm8);
sgr_out_16bpc(token, dst, src, sumsq, sum, w, h, n);
}
#[cfg(target_arch = "aarch64")]
#[arcane]
fn boxsum_16bpc(
token: Arm64,
sumsq: &mut [i32; BOX_LEN],
sum: &mut [i32; BOX_LEN],
src: &[u16; TMP_LEN],
bw: usize,
bh: usize,
n: i32,
) {
let mut vs = [0u32; S];
let mut vq = [0u32; S];
for r in 1..=bh - 4 {
let (os, oq) = (&mut sum[r * S..r * S + bw], &mut sumsq[r * S..r * S + bw]);
if n == 25 {
box_row_16bpc::<5>(token, src, r, bw, &mut vs, &mut vq, os, oq);
} else {
box_row_16bpc::<3>(token, src, r, bw, &mut vs, &mut vq, os, oq);
}
}
}
#[cfg(target_arch = "aarch64")]
fn sgr_16bpc(
token: Arm64,
p: PicOffset,
left: &[LeftPixelRow<u16>],
lpf: &DisjointMut<AlignedVec64<u8>>,
lpf_off: isize,
w: usize,
h: usize,
params: &LooprestorationParams,
edges: LrEdgeFlags,
variant: usize,
bitdepth_max: i32,
) {
with_scratch16(|sc| {
padding::<BitDepth16>(&mut sc.tmp, p, left, lpf, lpf_off, w, h, edges);
let sgr = params.sgr();
let bdm8 = if bitdepth_max == 1023 { 2 } else { 4 };
match variant {
2 => {
selfguided_16bpc(
token,
&mut sc.d0,
&sc.tmp,
w,
h,
25,
sgr.s0,
bdm8,
&mut sc.sumsq,
&mut sc.sum,
);
sgr_apply_16bpc(token, p, w, h, &sc.d0, sgr.w0 as i32, None, 0, bitdepth_max);
}
3 => {
selfguided_16bpc(
token,
&mut sc.d0,
&sc.tmp,
w,
h,
9,
sgr.s1,
bdm8,
&mut sc.sumsq,
&mut sc.sum,
);
sgr_apply_16bpc(token, p, w, h, &sc.d0, sgr.w1 as i32, None, 0, bitdepth_max);
}
_ => {
selfguided_16bpc(
token,
&mut sc.d0,
&sc.tmp,
w,
h,
25,
sgr.s0,
bdm8,
&mut sc.sumsq,
&mut sc.sum,
);
selfguided_16bpc(
token,
&mut sc.d1,
&sc.tmp,
w,
h,
9,
sgr.s1,
bdm8,
&mut sc.sumsq,
&mut sc.sum,
);
sgr_apply_16bpc(
token,
p,
w,
h,
&sc.d0,
sgr.w0 as i32,
Some(&sc.d1),
sgr.w1 as i32,
bitdepth_max,
);
}
}
});
}
#[cfg(target_arch = "aarch64")]
pub fn lr_filter_dispatch<BD: BitDepth>(
variant: usize,
dst: PicOffset,
left: &[LeftPixelRow<BD::Pixel>],
lpf: &DisjointMut<AlignedVec64<u8>>,
lpf_off: isize,
w: c_int,
h: c_int,
params: &LooprestorationParams,
edges: LrEdgeFlags,
bd: BD,
) -> bool {
use crate::include::common::bitdepth::BPC;
use crate::src::safe_simd::pixel_access::reinterpret_slice;
use archmage::SimdToken as _;
crate::src::ablate::note(
crate::src::ablate::Family::LoopRestoration,
(w as i64 * h as i64).unsigned_abs(),
);
if crate::src::ablate::is_off(crate::src::ablate::Family::LoopRestoration) {
return false;
}
let Some(token) = Arm64::summon() else {
return false;
};
#[cfg(feature = "__lrvarcov")]
{
use std::sync::atomic::{AtomicU8, Ordering};
static SEEN: [AtomicU8; 10] = [const { AtomicU8::new(0) }; 10];
let cell =
(BD::BPC == crate::include::common::bitdepth::BPC::BPC16) as usize * 5 + variant.min(4);
if SEEN[cell].swap(1, Ordering::Relaxed) == 0 {
let name = ["wiener7", "wiener5", "sgr_5x5", "sgr_3x3", "sgr_mix"][variant.min(4)];
let bpc = if cell >= 5 { "16bpc" } else { "8bpc" };
eprintln!("LRVAR\t{bpc}\t{name}");
}
}
let w = w as usize;
let h = h as usize;
let bd_c = bd.into_c();
match BD::BPC {
BPC::BPC8 => {
let left: &[LeftPixelRow<u8>] =
reinterpret_slice(left).expect("BD::Pixel layout matches u8");
match variant {
0 | 1 => wiener_8bpc(token, dst, left, lpf, lpf_off, w, h, params, edges),
v => sgr_8bpc(token, dst, left, lpf, lpf_off, w, h, params, edges, v),
}
}
BPC::BPC16 => {
let left: &[LeftPixelRow<u16>] =
reinterpret_slice(left).expect("BD::Pixel layout matches u16");
match variant {
0 | 1 => wiener_16bpc(token, dst, left, lpf, lpf_off, w, h, params, edges, bd_c),
v => sgr_16bpc(token, dst, left, lpf, lpf_off, w, h, params, edges, v, bd_c),
}
}
}
true
}
#[cfg(not(target_arch = "aarch64"))]
pub fn lr_filter_dispatch<BD: BitDepth>(
_variant: usize,
_dst: PicOffset,
_left: &[LeftPixelRow<BD::Pixel>],
_lpf: &DisjointMut<AlignedVec64<u8>>,
_lpf_off: isize,
w: c_int,
h: c_int,
_params: &LooprestorationParams,
_edges: LrEdgeFlags,
_bd: BD,
) -> bool {
crate::src::ablate::note(
crate::src::ablate::Family::LoopRestoration,
(w as i64 * h as i64).unsigned_abs(),
);
false
}