use core::arch::aarch64::{
int64x2_t, vaddq_s64, vaddq_u64, vcgtq_u64, vdupq_n_s64, vld1q_s64, vorrq_u64, vreinterpretq_s64_u64, vreinterpretq_u64_s64,
vshlq_s64, vshlq_u64, vst1q_s64, vsubq_s64, vsubq_u64, vuzp1q_s64, vuzp2q_s64, vzip1q_s64, vzip2q_s64,
};
use poulpy_cpu_ref::NTT4x30Ref;
use poulpy_cpu_ref::reference::ntt4x30::{I128NormalizeOps, vec_znx_big::AssignOp};
struct NfcShifts {
sll_b2klsh: int64x2_t,
sra_b2klsh: int64x2_t,
srl_b2klsh: int64x2_t,
sra_b2klsh_co_hi: int64x2_t,
sll_lsh: int64x2_t,
sll_b2k: int64x2_t,
sra_b2k: int64x2_t,
srl_b2k: int64x2_t,
sra_b2k_carry: int64x2_t,
}
impl NfcShifts {
#[inline(always)]
fn new(base2k: u32, lsh: u32) -> Self {
let b2klsh = base2k - lsh;
unsafe {
Self {
sll_b2klsh: vdupq_n_s64((64 - b2klsh) as i64),
sra_b2klsh: vdupq_n_s64(-((64 - b2klsh) as i64)),
srl_b2klsh: vdupq_n_s64(-(b2klsh as i64)),
sra_b2klsh_co_hi: vdupq_n_s64(-(b2klsh as i64)),
sll_lsh: vdupq_n_s64(lsh as i64),
sll_b2k: vdupq_n_s64((64 - base2k) as i64),
sra_b2k: vdupq_n_s64(-((64 - base2k) as i64)),
srl_b2k: vdupq_n_s64(-(base2k as i64)),
sra_b2k_carry: vdupq_n_s64(-(base2k as i64)),
}
}
}
}
#[inline(always)]
unsafe fn load2_split_i128(p: *const i128) -> (int64x2_t, int64x2_t) {
unsafe {
let v0 = vld1q_s64(p as *const i64); let v1 = vld1q_s64((p as *const i64).add(2)); let lo = vuzp1q_s64(v0, v1); let hi = vuzp2q_s64(v0, v1); (lo, hi)
}
}
#[inline(always)]
unsafe fn store2_split_i128(p: *mut i128, lo: int64x2_t, hi: int64x2_t) {
unsafe {
vst1q_s64(p as *mut i64, vzip1q_s64(lo, hi)); vst1q_s64((p as *mut i64).add(2), vzip2q_s64(lo, hi)); }
}
#[inline(always)]
unsafe fn load2_i64_as_split_i128(r_ptr: *const i64) -> (int64x2_t, int64x2_t) {
unsafe {
let lo = vld1q_s64(r_ptr); let hi = vshlq_s64(lo, vdupq_n_s64(-63));
(lo, hi)
}
}
#[inline(always)]
unsafe fn store2_i64(r_ptr: *mut i64, lo: int64x2_t) {
unsafe { vst1q_s64(r_ptr, lo) }
}
#[inline(always)]
unsafe fn nfc_middle_chunk(
s: &NfcShifts,
lo_a: int64x2_t,
hi_a: int64x2_t,
lo_c: int64x2_t,
hi_c: int64x2_t,
) -> (int64x2_t, int64x2_t, int64x2_t) {
unsafe {
let lo_dig = vshlq_s64(vshlq_s64(lo_a, s.sll_b2klsh), s.sra_b2klsh);
let hi_dig = vshlq_s64(lo_dig, vdupq_n_s64(-63));
let diff_lo_u = vsubq_u64(vreinterpretq_u64_s64(lo_a), vreinterpretq_u64_s64(lo_dig));
let borrow_mask = vcgtq_u64(vreinterpretq_u64_s64(lo_dig), vreinterpretq_u64_s64(lo_a));
let borrow_s = vreinterpretq_s64_u64(borrow_mask); let diff_hi = vaddq_s64(vsubq_s64(hi_a, hi_dig), borrow_s);
let co_lo_u = vorrq_u64(
vshlq_u64(diff_lo_u, s.srl_b2klsh),
vshlq_u64(vreinterpretq_u64_s64(diff_hi), s.sll_b2klsh),
);
let co_lo = vreinterpretq_s64_u64(co_lo_u);
let co_hi = vshlq_s64(diff_hi, s.sra_b2klsh_co_hi);
let lo_dig_sh = vshlq_s64(lo_dig, s.sll_lsh);
let hi_dig_sh = vshlq_s64(lo_dig_sh, vdupq_n_s64(-63));
let lo_dpc = vaddq_s64(lo_dig_sh, lo_c);
let carry1_mask = vcgtq_u64(vreinterpretq_u64_s64(lo_dig_sh), vreinterpretq_u64_s64(lo_dpc));
let carry1_s = vreinterpretq_s64_u64(carry1_mask);
let hi_dpc = vsubq_s64(vaddq_s64(hi_dig_sh, hi_c), carry1_s);
let lo_out = vshlq_s64(vshlq_s64(lo_dpc, s.sll_b2k), s.sra_b2k);
let hi_out = vshlq_s64(lo_out, vdupq_n_s64(-63));
let diff2_lo_u = vsubq_u64(vreinterpretq_u64_s64(lo_dpc), vreinterpretq_u64_s64(lo_out));
let borrow2_mask = vcgtq_u64(vreinterpretq_u64_s64(lo_out), vreinterpretq_u64_s64(lo_dpc));
let diff2_hi = vaddq_s64(vsubq_s64(hi_dpc, hi_out), vreinterpretq_s64_u64(borrow2_mask));
let carry2_lo_u = vorrq_u64(
vshlq_u64(diff2_lo_u, s.srl_b2k),
vshlq_u64(vreinterpretq_u64_s64(diff2_hi), s.sll_b2k),
);
let carry2_lo = vreinterpretq_s64_u64(carry2_lo_u);
let carry2_hi = vshlq_s64(diff2_hi, s.sra_b2k_carry);
let new_lo_c_u = vaddq_u64(vreinterpretq_u64_s64(co_lo), vreinterpretq_u64_s64(carry2_lo));
let cmask = vcgtq_u64(vreinterpretq_u64_s64(co_lo), new_lo_c_u);
let new_lo_c = vreinterpretq_s64_u64(new_lo_c_u);
let new_hi_c = vsubq_s64(vaddq_s64(co_hi, carry2_hi), vreinterpretq_s64_u64(cmask));
(lo_out, new_lo_c, new_hi_c)
}
}
#[inline(always)]
unsafe fn nfc_final_chunk(s: &NfcShifts, lo_a: int64x2_t, lo_c: int64x2_t) -> int64x2_t {
unsafe {
let lo_dig = vshlq_s64(vshlq_s64(lo_a, s.sll_b2klsh), s.sra_b2klsh);
let lo_dpc = vaddq_s64(vshlq_s64(lo_dig, s.sll_lsh), lo_c);
vshlq_s64(vshlq_s64(lo_dpc, s.sll_b2k), s.sra_b2k)
}
}
pub(crate) fn nfc_middle_step_neon(base2k: usize, lsh: usize, res: &mut [i64], a: &[i128], carry: &mut [i128]) {
if base2k > 64 || res.len() < 2 {
<NTT4x30Ref as I128NormalizeOps>::nfc_middle_step(base2k, lsh, res, a, carry);
return;
}
let n = res.len();
let chunks = n >> 1;
unsafe {
let s = NfcShifts::new(base2k as u32, lsh as u32);
let mut a_ptr = a.as_ptr();
let mut c_ptr = carry.as_mut_ptr();
let mut r_ptr = res.as_mut_ptr();
for _ in 0..chunks {
let (lo_a, hi_a) = load2_split_i128(a_ptr);
let (lo_c, hi_c) = load2_split_i128(c_ptr as *const i128);
let (lo_out, new_lo_c, new_hi_c) = nfc_middle_chunk(&s, lo_a, hi_a, lo_c, hi_c);
store2_i64(r_ptr, lo_out);
store2_split_i128(c_ptr, new_lo_c, new_hi_c);
a_ptr = a_ptr.add(2);
c_ptr = c_ptr.add(2);
r_ptr = r_ptr.add(2);
}
}
let tail = chunks << 1;
if tail < n {
<NTT4x30Ref as I128NormalizeOps>::nfc_middle_step(base2k, lsh, &mut res[tail..], &a[tail..], &mut carry[tail..]);
}
}
pub(crate) fn nfc_middle_step_into_neon<O: AssignOp>(base2k: usize, lsh: usize, res: &mut [i64], a: &[i128], carry: &mut [i128]) {
if base2k > 64 || res.len() < 2 {
<NTT4x30Ref as I128NormalizeOps>::nfc_middle_step_into::<O>(base2k, lsh, res, a, carry);
return;
}
let n = res.len();
let chunks = n >> 1;
unsafe {
let s = NfcShifts::new(base2k as u32, lsh as u32);
let mut a_ptr = a.as_ptr();
let mut c_ptr = carry.as_mut_ptr();
let mut r_ptr = res.as_mut_ptr();
for _ in 0..chunks {
let (lo_a, hi_a) = load2_split_i128(a_ptr);
let (lo_c, hi_c) = load2_split_i128(c_ptr as *const i128);
let (lo_out, new_lo_c, new_hi_c) = nfc_middle_chunk(&s, lo_a, hi_a, lo_c, hi_c);
let lo_res = vld1q_s64(r_ptr);
let combined = if O::SUB {
vsubq_s64(lo_res, lo_out)
} else {
vaddq_s64(lo_res, lo_out)
};
vst1q_s64(r_ptr, combined);
store2_split_i128(c_ptr, new_lo_c, new_hi_c);
a_ptr = a_ptr.add(2);
c_ptr = c_ptr.add(2);
r_ptr = r_ptr.add(2);
}
}
let tail = chunks << 1;
if tail < n {
<NTT4x30Ref as I128NormalizeOps>::nfc_middle_step_into::<O>(
base2k,
lsh,
&mut res[tail..],
&a[tail..],
&mut carry[tail..],
);
}
}
pub(crate) fn nfc_middle_step_assign_neon(base2k: usize, lsh: usize, res: &mut [i64], carry: &mut [i128]) {
if base2k > 64 || res.len() < 2 {
<NTT4x30Ref as I128NormalizeOps>::nfc_middle_step_assign(base2k, lsh, res, carry);
return;
}
let n = res.len();
let chunks = n >> 1;
unsafe {
let s = NfcShifts::new(base2k as u32, lsh as u32);
let mut c_ptr = carry.as_mut_ptr();
let mut r_ptr = res.as_mut_ptr();
for _ in 0..chunks {
let (lo_a, hi_a) = load2_i64_as_split_i128(r_ptr);
let (lo_c, hi_c) = load2_split_i128(c_ptr as *const i128);
let (lo_out, new_lo_c, new_hi_c) = nfc_middle_chunk(&s, lo_a, hi_a, lo_c, hi_c);
store2_i64(r_ptr, lo_out);
store2_split_i128(c_ptr, new_lo_c, new_hi_c);
c_ptr = c_ptr.add(2);
r_ptr = r_ptr.add(2);
}
}
let tail = chunks << 1;
if tail < n {
<NTT4x30Ref as I128NormalizeOps>::nfc_middle_step_assign(base2k, lsh, &mut res[tail..], &mut carry[tail..]);
}
}
pub(crate) fn nfc_final_step_assign_neon(base2k: usize, lsh: usize, res: &mut [i64], carry: &mut [i128]) {
if base2k > 64 || res.len() < 2 {
<NTT4x30Ref as I128NormalizeOps>::nfc_final_step_assign(base2k, lsh, res, carry);
return;
}
let n = res.len();
let chunks = n >> 1;
unsafe {
let s = NfcShifts::new(base2k as u32, lsh as u32);
let mut c_ptr = carry.as_ptr();
let mut r_ptr = res.as_mut_ptr();
for _ in 0..chunks {
let (lo_a, _hi_a) = load2_i64_as_split_i128(r_ptr);
let c0 = vld1q_s64(c_ptr as *const i64);
let c1 = vld1q_s64((c_ptr as *const i64).add(2));
let lo_c = vuzp1q_s64(c0, c1);
let lo_out = nfc_final_chunk(&s, lo_a, lo_c);
store2_i64(r_ptr, lo_out);
c_ptr = c_ptr.add(2);
r_ptr = r_ptr.add(2);
}
}
let tail = chunks << 1;
if tail < n {
<NTT4x30Ref as I128NormalizeOps>::nfc_final_step_assign(base2k, lsh, &mut res[tail..], &mut carry[tail..]);
}
}
pub(crate) fn nfc_final_step_into_neon<O: AssignOp>(base2k: usize, lsh: usize, res: &mut [i64], carry: &mut [i128]) {
if base2k > 64 || res.len() < 2 {
<NTT4x30Ref as I128NormalizeOps>::nfc_final_step_into::<O>(base2k, lsh, res, carry);
return;
}
let n = res.len();
let chunks = n >> 1;
unsafe {
let s = NfcShifts::new(base2k as u32, lsh as u32);
let mut c_ptr = carry.as_ptr();
let mut r_ptr = res.as_mut_ptr();
for _ in 0..chunks {
let lo_res = vld1q_s64(r_ptr);
let (lo_a, _hi_a) = load2_i64_as_split_i128(r_ptr);
let c0 = vld1q_s64(c_ptr as *const i64);
let c1 = vld1q_s64((c_ptr as *const i64).add(2));
let lo_c = vuzp1q_s64(c0, c1);
let lo_out = nfc_final_chunk(&s, lo_a, lo_c);
let combined = if O::SUB {
vsubq_s64(lo_res, lo_out)
} else {
vaddq_s64(lo_res, lo_out)
};
vst1q_s64(r_ptr, combined);
c_ptr = c_ptr.add(2);
r_ptr = r_ptr.add(2);
}
}
let tail = chunks << 1;
if tail < n {
<NTT4x30Ref as I128NormalizeOps>::nfc_final_step_into::<O>(base2k, lsh, &mut res[tail..], &mut carry[tail..]);
}
}
#[cfg(test)]
mod tests {
use super::*;
use rand::{RngExt, SeedableRng};
use rand_chacha::ChaCha8Rng;
const LENGTHS: &[usize] = &[0, 1, 2, 3, 4, 5, 7, 8, 16, 17, 64, 65];
const SHIFTS: &[(usize, usize)] = &[(12, 0), (50, 0), (50, 7), (60, 0), (60, 30), (64, 0), (64, 17)];
fn rng() -> ChaCha8Rng {
ChaCha8Rng::seed_from_u64(0xb00b_b00b_b00b_b00b)
}
fn random_i128(rng: &mut ChaCha8Rng, n: usize) -> Vec<i128> {
(0..n)
.map(|_| {
let lo: u64 = rng.random();
let hi: u64 = rng.random();
(((hi as u128) << 64) | lo as u128) as i128
})
.collect()
}
fn random_i64(rng: &mut ChaCha8Rng, n: usize) -> Vec<i64> {
(0..n).map(|_| rng.random::<i64>()).collect()
}
use poulpy_cpu_ref::reference::ntt4x30::vec_znx_big::{AddOp, SubOp};
#[test]
fn nfc_middle_step_matches_scalar() {
let mut rng = rng();
for &n in LENGTHS {
for &(b, l) in SHIFTS {
if l >= b {
continue;
}
let a = random_i128(&mut rng, n);
let c0 = random_i128(&mut rng, n);
let mut got_r = vec![0i64; n];
let mut got_c = c0.clone();
let mut want_r = vec![0i64; n];
let mut want_c = c0;
nfc_middle_step_neon(b, l, &mut got_r, &a, &mut got_c);
<NTT4x30Ref as I128NormalizeOps>::nfc_middle_step(b, l, &mut want_r, &a, &mut want_c);
assert_eq!(got_r, want_r, "res mismatch n={n} base2k={b} lsh={l}");
assert_eq!(got_c, want_c, "carry mismatch n={n} base2k={b} lsh={l}");
}
}
}
#[test]
fn nfc_middle_step_assign_matches_scalar() {
let mut rng = rng();
for &n in LENGTHS {
for &(b, l) in SHIFTS {
if l >= b {
continue;
}
let r0 = random_i64(&mut rng, n);
let c0 = random_i128(&mut rng, n);
let mut got_r = r0.clone();
let mut got_c = c0.clone();
let mut want_r = r0;
let mut want_c = c0;
nfc_middle_step_assign_neon(b, l, &mut got_r, &mut got_c);
<NTT4x30Ref as I128NormalizeOps>::nfc_middle_step_assign(b, l, &mut want_r, &mut want_c);
assert_eq!(got_r, want_r, "res mismatch n={n} base2k={b} lsh={l}");
assert_eq!(got_c, want_c, "carry mismatch n={n} base2k={b} lsh={l}");
}
}
}
#[test]
fn nfc_middle_step_into_add_matches_scalar() {
let mut rng = rng();
for &n in LENGTHS {
for &(b, l) in SHIFTS {
if l >= b {
continue;
}
let r0 = random_i64(&mut rng, n);
let a = random_i128(&mut rng, n);
let c0 = random_i128(&mut rng, n);
let mut got_r = r0.clone();
let mut got_c = c0.clone();
let mut want_r = r0;
let mut want_c = c0;
nfc_middle_step_into_neon::<AddOp>(b, l, &mut got_r, &a, &mut got_c);
<NTT4x30Ref as I128NormalizeOps>::nfc_middle_step_into::<AddOp>(b, l, &mut want_r, &a, &mut want_c);
assert_eq!(got_r, want_r, "res mismatch n={n} base2k={b} lsh={l}");
assert_eq!(got_c, want_c, "carry mismatch n={n} base2k={b} lsh={l}");
}
}
}
#[test]
fn nfc_middle_step_into_sub_matches_scalar() {
let mut rng = rng();
for &n in LENGTHS {
for &(b, l) in SHIFTS {
if l >= b {
continue;
}
let r0 = random_i64(&mut rng, n);
let a = random_i128(&mut rng, n);
let c0 = random_i128(&mut rng, n);
let mut got_r = r0.clone();
let mut got_c = c0.clone();
let mut want_r = r0;
let mut want_c = c0;
nfc_middle_step_into_neon::<SubOp>(b, l, &mut got_r, &a, &mut got_c);
<NTT4x30Ref as I128NormalizeOps>::nfc_middle_step_into::<SubOp>(b, l, &mut want_r, &a, &mut want_c);
assert_eq!(got_r, want_r, "res mismatch n={n} base2k={b} lsh={l}");
assert_eq!(got_c, want_c, "carry mismatch n={n} base2k={b} lsh={l}");
}
}
}
#[test]
fn nfc_final_step_assign_matches_scalar() {
let mut rng = rng();
for &n in LENGTHS {
for &(b, l) in SHIFTS {
if l >= b {
continue;
}
let r0 = random_i64(&mut rng, n);
let c0 = random_i128(&mut rng, n);
let mut got_r = r0.clone();
let mut got_c = c0.clone();
let mut want_r = r0;
let mut want_c = c0;
nfc_final_step_assign_neon(b, l, &mut got_r, &mut got_c);
<NTT4x30Ref as I128NormalizeOps>::nfc_final_step_assign(b, l, &mut want_r, &mut want_c);
assert_eq!(got_r, want_r, "res mismatch n={n} base2k={b} lsh={l}");
}
}
}
#[test]
fn nfc_final_step_into_add_matches_scalar() {
let mut rng = rng();
for &n in LENGTHS {
for &(b, l) in SHIFTS {
if l >= b {
continue;
}
let r0 = random_i64(&mut rng, n);
let c0 = random_i128(&mut rng, n);
let mut got_r = r0.clone();
let mut got_c = c0.clone();
let mut want_r = r0;
let mut want_c = c0;
nfc_final_step_into_neon::<AddOp>(b, l, &mut got_r, &mut got_c);
<NTT4x30Ref as I128NormalizeOps>::nfc_final_step_into::<AddOp>(b, l, &mut want_r, &mut want_c);
assert_eq!(got_r, want_r, "res mismatch n={n} base2k={b} lsh={l}");
}
}
}
#[test]
fn nfc_final_step_into_sub_matches_scalar() {
let mut rng = rng();
for &n in LENGTHS {
for &(b, l) in SHIFTS {
if l >= b {
continue;
}
let r0 = random_i64(&mut rng, n);
let c0 = random_i128(&mut rng, n);
let mut got_r = r0.clone();
let mut got_c = c0.clone();
let mut want_r = r0;
let mut want_c = c0;
nfc_final_step_into_neon::<SubOp>(b, l, &mut got_r, &mut got_c);
<NTT4x30Ref as I128NormalizeOps>::nfc_final_step_into::<SubOp>(b, l, &mut want_r, &mut want_c);
assert_eq!(got_r, want_r, "res mismatch n={n} base2k={b} lsh={l}");
}
}
}
}