use super::*;
#[inline(always)]
pub unsafe fn transpose256_w32(i: [__m256; 8]) -> [__m256; 8] {
let t0 = _mm256_unpacklo_ps(i[0], i[1]);
let t1 = _mm256_unpackhi_ps(i[0], i[1]);
let t2 = _mm256_unpacklo_ps(i[2], i[3]);
let t3 = _mm256_unpackhi_ps(i[2], i[3]);
let t4 = _mm256_unpacklo_ps(i[4], i[5]);
let t5 = _mm256_unpackhi_ps(i[4], i[5]);
let t6 = _mm256_unpacklo_ps(i[6], i[7]);
let t7 = _mm256_unpackhi_ps(i[6], i[7]);
let s0 = _mm256_shuffle_ps::<0x44>(t0, t2);
let s1 = _mm256_shuffle_ps::<0xEE>(t0, t2);
let s2 = _mm256_shuffle_ps::<0x44>(t1, t3);
let s3 = _mm256_shuffle_ps::<0xEE>(t1, t3);
let s4 = _mm256_shuffle_ps::<0x44>(t4, t6);
let s5 = _mm256_shuffle_ps::<0xEE>(t4, t6);
let s6 = _mm256_shuffle_ps::<0x44>(t5, t7);
let s7 = _mm256_shuffle_ps::<0xEE>(t5, t7);
[
_mm256_permute2f128_ps::<0x20>(s0, s4),
_mm256_permute2f128_ps::<0x20>(s1, s5),
_mm256_permute2f128_ps::<0x20>(s2, s6),
_mm256_permute2f128_ps::<0x20>(s3, s7),
_mm256_permute2f128_ps::<0x31>(s0, s4),
_mm256_permute2f128_ps::<0x31>(s1, s5),
_mm256_permute2f128_ps::<0x31>(s2, s6),
_mm256_permute2f128_ps::<0x31>(s3, s7),
]
}
#[inline(always)]
pub unsafe fn transpose256_w64(i: [__m256d; 4]) -> [__m256d; 4] {
let t0 = _mm256_unpacklo_pd(i[0], i[1]);
let t1 = _mm256_unpackhi_pd(i[0], i[1]);
let t2 = _mm256_unpacklo_pd(i[2], i[3]);
let t3 = _mm256_unpackhi_pd(i[2], i[3]);
[
_mm256_permute2f128_pd(t0, t2, 0x20),
_mm256_permute2f128_pd(t1, t3, 0x20),
_mm256_permute2f128_pd(t0, t2, 0x31),
_mm256_permute2f128_pd(t1, t3, 0x31),
]
}
#[inline(always)]
pub unsafe fn radix_by_w32_si<const N: usize>(inputs: [__m256i; N]) -> [__m256i; N] {
let f = transpose256_w32([
_mm256_castsi256_ps(*inputs.get_unchecked(0)),
_mm256_castsi256_ps(*inputs.get_unchecked(1)),
_mm256_castsi256_ps(*inputs.get_unchecked(2)),
_mm256_castsi256_ps(*inputs.get_unchecked(3)),
_mm256_castsi256_ps(*inputs.get_unchecked(4)),
_mm256_castsi256_ps(*inputs.get_unchecked(5)),
_mm256_castsi256_ps(*inputs.get_unchecked(6)),
_mm256_castsi256_ps(*inputs.get_unchecked(7)),
]);
let mut out = inputs;
let mut k = 0;
while k < 8 {
*out.get_unchecked_mut(k) = _mm256_castps_si256(f[k]);
k += 1;
}
out
}
#[inline(always)]
pub unsafe fn radix_by_w64_si<const N: usize>(inputs: [__m256i; N]) -> [__m256i; N] {
let d = transpose256_w64([
_mm256_castsi256_pd(*inputs.get_unchecked(0)),
_mm256_castsi256_pd(*inputs.get_unchecked(1)),
_mm256_castsi256_pd(*inputs.get_unchecked(2)),
_mm256_castsi256_pd(*inputs.get_unchecked(3)),
]);
let mut out = inputs;
let mut k = 0;
while k < 4 {
*out.get_unchecked_mut(k) = _mm256_castpd_si256(d[k]);
k += 1;
}
out
}
pub const LADDER_MAX_N: usize = 32;
#[derive(Debug, Clone, Copy)]
pub struct LadderPlan {
pub ok: bool,
pub rounds: usize,
pub rung: [u8; 6],
pub style: [u8; 6],
pub relabel: [u8; LADDER_MAX_N],
}
const LADDER_NONE: LadderPlan = LadderPlan {
ok: false,
rounds: 0,
rung: [0; 6],
style: [0; 6],
relabel: [0; LADDER_MAX_N],
};
const fn sim_rung(rung: u8, x: [u16; 8], y: [u16; 8]) -> ([u16; 8], [u16; 8]) {
match rung {
0 => (
[x[0], y[0], x[1], y[1], x[4], y[4], x[5], y[5]],
[x[2], y[2], x[3], y[3], x[6], y[6], x[7], y[7]],
),
1 => (
[x[0], x[2], y[0], y[2], x[4], x[6], y[4], y[6]],
[x[1], x[3], y[1], y[3], x[5], x[7], y[5], y[7]],
),
2 => (
[x[0], x[1], y[0], y[1], x[4], x[5], y[4], y[5]],
[x[2], x[3], y[2], y[3], x[6], x[7], y[6], y[7]],
),
_ => (
[x[0], x[1], x[2], x[3], y[0], y[1], y[2], y[3]],
[x[4], x[5], x[6], x[7], y[4], y[5], y[6], y[7]],
),
}
}
const fn sim_plan(n: usize, g: usize, plan: &mut LadderPlan) -> bool {
let mut buf = [[0u16; 8]; LADDER_MAX_N];
let mut tmp = [[0u16; 8]; LADDER_MAX_N];
let mut r = 0;
while r < n {
let mut s = 0;
while s < 8 {
buf[r][s] = (r * 8 + s) as u16;
s += 1;
}
r += 1;
}
let mut round = 0;
let mut bsize = n;
while round < plan.rounds {
let half = bsize / 2;
let mut base = 0;
while base < n {
let mut i = 0;
while i < half {
let (a, b, da, db) = if plan.style[round] == 0 {
(base + 2 * i, base + 2 * i + 1, base + i, base + half + i)
} else {
(base + i, base + half + i, base + 2 * i, base + 2 * i + 1)
};
let (lo, hi) = sim_rung(plan.rung[round], buf[a], buf[b]);
tmp[da] = lo;
tmp[db] = hi;
i += 1;
}
base += bsize;
}
let mut c = 0;
while c < n {
buf[c] = tmp[c];
c += 1;
}
bsize = half;
round += 1;
}
let m = 8 / g;
let mut out_reg = 0;
while out_reg < n {
let mut found = usize::MAX;
let mut cand = 0;
while cand < n {
let mut matches = true;
let mut q = 0;
while q < m && matches {
let j = q * n + out_reg;
let mut s = 0;
while s < g && matches {
let want = ((j / m) * 8 + (j % m) * g + s) as u16;
if buf[cand][q * g + s] != want {
matches = false;
}
s += 1;
}
q += 1;
}
if matches {
found = cand;
break;
}
cand += 1;
}
if found == usize::MAX {
return false;
}
plan.relabel[out_reg] = found as u8;
out_reg += 1;
}
true
}
const fn ilog2(mut v: usize) -> usize {
let mut l = 0;
while v > 1 {
v /= 2;
l += 1;
}
l
}
pub const fn ladder_search(n: usize, g: usize) -> LadderPlan {
if n < 2 || n > LADDER_MAX_N || !n.is_power_of_two() {
return LADDER_NONE;
}
if g == 0 || g > 8 || !g.is_power_of_two() || 8 % g != 0 {
return LADDER_NONE;
}
let m = 8 / g;
if m < 2 {
return LADDER_NONE;
}
let min_rounds = if ilog2(n) < ilog2(m) { ilog2(n) } else { ilog2(m) };
if min_rounds == 0 || min_rounds > 4 {
return LADDER_NONE;
}
let mut pass = 0;
while pass < 3 {
let rounds = min_rounds + pass;
if rounds > 6 || rounds > ilog2(n) {
pass += 1;
continue;
}
let distinct = pass == 0;
if let Some(plan) = ladder_pass(n, g, rounds, distinct) {
return plan;
}
pass += 1;
}
LADDER_NONE
}
const fn ladder_pass(n: usize, g: usize, rounds: usize, distinct: bool) -> Option<LadderPlan> {
let allowed: [bool; 4] = [g <= 1, g <= 1, g <= 2, g <= 4];
let mut code = 0usize;
let total = {
let mut t = 2; let mut i = 0;
while i < rounds {
t *= 4;
i += 1;
}
t
};
while code < total {
let mut plan = LadderPlan {
ok: false,
rounds,
rung: [0; 6],
style: [0; 6],
relabel: [0; LADDER_MAX_N],
};
let mut c = code;
let style = (c % 2) as u8;
c /= 2;
let mut valid = true;
let mut used_width = [false; 3]; let mut i = 0;
while i < rounds {
let rung = (c % 4) as u8;
c /= 4;
if !allowed[rung as usize] {
valid = false;
break;
}
if rung == 3 && i + 1 != rounds && g < 4 {
valid = false;
break;
}
let w = if rung <= 1 { 0 } else { (rung - 1) as usize };
if distinct {
if used_width[w] {
valid = false;
break;
}
used_width[w] = true;
}
plan.rung[i] = rung;
plan.style[i] = style;
i += 1;
}
if valid && sim_plan(n, g, &mut plan) {
plan.ok = true;
return Some(plan);
}
code += 1;
}
None
}
pub const fn ladder_search_elem(n: usize, group: usize, elem_bytes: usize) -> LadderPlan {
let bytes = group * elem_bytes;
if !bytes.is_multiple_of(4) {
return LADDER_NONE;
}
ladder_search(n, bytes / 4)
}
pub const fn ladder_viable(n: usize, group: usize, elem_bytes: usize) -> bool {
ladder_search_elem(n, group, elem_bytes).ok
}
#[inline(always)]
unsafe fn rung_ps(rung: u8, inv: bool, x: __m256, y: __m256) -> (__m256, __m256) {
let r = if inv && rung <= 1 { 1 - rung } else { rung };
match r {
0 => (_mm256_unpacklo_ps(x, y), _mm256_unpackhi_ps(x, y)),
1 => (_mm256_shuffle_ps::<0x88>(x, y), _mm256_shuffle_ps::<0xDD>(x, y)),
2 => (_mm256_shuffle_ps::<0x44>(x, y), _mm256_shuffle_ps::<0xEE>(x, y)),
_ => (
_mm256_permute2f128_ps::<0x20>(x, y),
_mm256_permute2f128_ps::<0x31>(x, y),
),
}
}
#[inline(always)]
unsafe fn ladder_round_ps<const N: usize>(buf: &mut [__m256; N], rung: u8, inv: bool, style: u8, bsize: usize) {
let tmp = *buf;
let half = bsize / 2;
let mut base = 0;
while base < N {
let mut i = 0;
while i < half {
let (a, b, da, db) = if style == 0 {
(base + 2 * i, base + 2 * i + 1, base + i, base + half + i)
} else {
(base + i, base + half + i, base + 2 * i, base + 2 * i + 1)
};
unsafe {
let (lo, hi) = rung_ps(rung, inv, *tmp.get_unchecked(a), *tmp.get_unchecked(b));
*buf.get_unchecked_mut(da) = lo;
*buf.get_unchecked_mut(db) = hi;
}
i += 1;
}
base += bsize;
}
}
#[inline(always)]
pub unsafe fn ladder_radix_by_ps<const N: usize, const DEINT: bool>(
inputs: [__m256; N],
plan: LadderPlan,
) -> [__m256; N] {
debug_assert!(plan.ok && N <= LADDER_MAX_N);
let mut buf = inputs;
if !DEINT {
let tmp = buf;
let mut j = 0;
while j < N {
unsafe { *buf.get_unchecked_mut(plan.relabel[j] as usize) = *tmp.get_unchecked(j) };
j += 1;
}
}
macro_rules! round {
($step:literal) => {
if $step < plan.rounds {
let round = if DEINT { $step } else { plan.rounds - 1 - $step };
let style = if DEINT {
plan.style[round]
} else {
1 - plan.style[round]
};
unsafe { ladder_round_ps::<N>(&mut buf, plan.rung[round], !DEINT, style, N >> round) };
}
};
}
round!(0);
round!(1);
round!(2);
round!(3);
round!(4);
round!(5);
if DEINT {
let tmp = buf;
let mut j = 0;
while j < N {
unsafe { *buf.get_unchecked_mut(j) = *tmp.get_unchecked(plan.relabel[j] as usize) };
j += 1;
}
}
buf
}
#[inline(always)]
pub unsafe fn ladder_radix_by_si<const N: usize, const DEINT: bool>(
inputs: [__m256i; N],
plan: LadderPlan,
) -> [__m256i; N] {
let mut ps = [_mm256_setzero_ps(); N];
let mut k = 0;
while k < N {
ps[k] = _mm256_castsi256_ps(inputs[k]);
k += 1;
}
let out = ladder_radix_by_ps::<N, DEINT>(ps, plan);
let mut si = inputs;
let mut k = 0;
while k < N {
si[k] = _mm256_castps_si256(out[k]);
k += 1;
}
si
}
#[inline(always)]
pub unsafe fn ladder_radix_by_pd<const N: usize, const DEINT: bool>(
inputs: [__m256d; N],
plan: LadderPlan,
) -> [__m256d; N] {
let mut ps = [_mm256_setzero_ps(); N];
let mut k = 0;
while k < N {
ps[k] = _mm256_castpd_ps(inputs[k]);
k += 1;
}
let out = ladder_radix_by_ps::<N, DEINT>(ps, plan);
let mut pd = inputs;
let mut k = 0;
while k < N {
pd[k] = _mm256_castps_pd(out[k]);
k += 1;
}
pd
}
#[cfg(all(test, feature = "std"))]
mod tests {
use super::*;
#[test]
fn ladder_certification_coverage() {
let cases: &[(usize, usize, bool)] = &[
(8, 1, true),
(4, 2, true),
(2, 4, true),
(4, 1, false), (2, 2, false), (2, 1, false),
(16, 1, true),
(32, 1, true),
(8, 2, true),
(16, 2, true),
(32, 2, false),
(8, 4, true),
(16, 4, true),
(32, 4, true),
(4, 4, true),
(3, 2, false),
(12, 1, false),
(2, 8, false), ];
for &(n, g, expect) in cases {
let plan = ladder_search(n, g);
assert_eq!(
plan.ok, expect,
"ladder_search({n}, {g}): expected ok={expect}, got {plan:?}"
);
if plan.ok {
std::println!(
"ladder({n:2}, {g}): rounds={} rung={:?} style={:?} relabel={:?}",
plan.rounds,
&plan.rung[..plan.rounds],
&plan.style[..plan.rounds],
&plan.relabel[..n]
);
}
}
}
}