use vyre_foundation::ir::{BufferAccess, BufferDecl, DataType, Node, Program};
use crate::fixpoint::persistent_fixpoint::{
persistent_fixpoint, persistent_fixpoint_grid, PERSISTENT_FIXPOINT_WORKGROUP_SIZE,
};
use crate::math::semiring_gemm::{semiring_gemm, Semiring};
use crate::math::sinkhorn::sinkhorn_scale;
pub const OP_ID: &str = "vyre-primitives::math::sinkhorn_iterate";
#[must_use]
#[allow(clippy::too_many_arguments)]
pub fn sinkhorn_iterate(
k: &str,
k_t: &str,
a: &str,
b: &str,
u_curr: &str,
u_next: &str,
v: &str,
kv: &str,
ktu: &str,
changed: &str,
m: u32,
n: u32,
max_iterations: u32,
) -> Program {
if m == 0 {
return crate::invalid_output_program(
OP_ID,
u_curr,
DataType::U32,
"Fix: sinkhorn_iterate requires m > 0, got 0.".to_string(),
);
}
if n == 0 {
return crate::invalid_output_program(
OP_ID,
u_curr,
DataType::U32,
"Fix: sinkhorn_iterate requires n > 0, got 0.".to_string(),
);
}
let Some(matrix_cells) = m.checked_mul(n) else {
return crate::invalid_output_program(
OP_ID,
u_curr,
DataType::U32,
format!("Fix: sinkhorn_iterate m*n overflows u32: {m}*{n}."),
);
};
let transfer_body = sinkhorn_transfer_body(k, k_t, a, b, u_next, v, kv, ktu, m, n);
let needs_grid_sync = matrix_cells > PERSISTENT_FIXPOINT_WORKGROUP_SIZE[0];
let inner = if needs_grid_sync {
persistent_fixpoint_grid(transfer_body, u_curr, u_next, changed, m, max_iterations)
} else {
persistent_fixpoint(transfer_body, u_curr, u_next, changed, m, max_iterations)
};
let changed_words = if needs_grid_sync {
max_iterations.max(1)
} else {
1
};
sinkhorn_wrap(
&inner,
k,
k_t,
a,
b,
u_curr,
u_next,
v,
kv,
ktu,
changed,
m,
n,
matrix_cells,
changed_words,
)
}
#[allow(clippy::too_many_arguments)]
fn sinkhorn_transfer_body(
k: &str,
k_t: &str,
a: &str,
b: &str,
u_next: &str,
v: &str,
kv: &str,
ktu: &str,
m: u32,
n: u32,
) -> Vec<Node> {
let extract_body = |p: Program| -> Vec<Node> {
let mut body = Vec::new();
for node in p.entry() {
if let Node::Region {
body: region_body, ..
} = node
{
body.extend(region_body.iter().cloned());
}
}
body
};
let seq_cst = || Node::Barrier {
ordering: vyre_foundation::MemoryOrdering::SeqCst,
};
let mut transfer_body = Vec::new();
transfer_body.extend(extract_body(semiring_gemm(
k,
v,
kv,
m,
1,
n,
Semiring::Real,
)));
transfer_body.push(seq_cst());
transfer_body.extend(extract_body(sinkhorn_scale(a, kv, u_next, m)));
transfer_body.push(seq_cst());
transfer_body.extend(extract_body(semiring_gemm(
k_t,
u_next,
ktu,
n,
1,
m,
Semiring::Real,
)));
transfer_body.push(seq_cst());
transfer_body.extend(extract_body(sinkhorn_scale(b, ktu, v, n)));
transfer_body.push(seq_cst());
transfer_body
}
#[allow(clippy::too_many_arguments)]
fn sinkhorn_wrap(
inner: &Program,
k: &str,
k_t: &str,
a: &str,
b: &str,
u_curr: &str,
u_next: &str,
v: &str,
kv: &str,
ktu: &str,
changed: &str,
m: u32,
n: u32,
matrix_cells: u32,
changed_words: u32,
) -> Program {
super::wrap_fixpoint_program(
OP_ID,
inner,
vec![
BufferDecl::storage(u_curr, 0, BufferAccess::ReadWrite, DataType::U32).with_count(m),
BufferDecl::storage(u_next, 1, BufferAccess::ReadWrite, DataType::U32).with_count(m),
BufferDecl::storage(changed, 2, BufferAccess::ReadWrite, DataType::U32)
.with_count(changed_words),
BufferDecl::storage(k, 3, BufferAccess::ReadOnly, DataType::U32)
.with_count(matrix_cells),
BufferDecl::storage(k_t, 4, BufferAccess::ReadOnly, DataType::U32)
.with_count(matrix_cells),
BufferDecl::storage(a, 5, BufferAccess::ReadOnly, DataType::U32).with_count(m),
BufferDecl::storage(b, 6, BufferAccess::ReadOnly, DataType::U32).with_count(n),
BufferDecl::storage(v, 7, BufferAccess::ReadWrite, DataType::U32).with_count(n),
BufferDecl::storage(kv, 8, BufferAccess::ReadWrite, DataType::U32).with_count(m),
BufferDecl::storage(ktu, 9, BufferAccess::ReadWrite, DataType::U32).with_count(n),
],
)
}
#[cfg(test)]
#[allow(clippy::too_many_arguments)]
fn sinkhorn_single_word_harness(
k: &str,
k_t: &str,
a: &str,
b: &str,
u_curr: &str,
u_next: &str,
v: &str,
kv: &str,
ktu: &str,
changed: &str,
m: u32,
n: u32,
max_iterations: u32,
) -> Program {
let matrix_cells = m
.checked_mul(n)
.expect("Fix: the divergence fixture must use non-overflowing extents.");
let transfer_body = sinkhorn_transfer_body(k, k_t, a, b, u_next, v, kv, ktu, m, n);
let inner = persistent_fixpoint(transfer_body, u_curr, u_next, changed, m, max_iterations);
sinkhorn_wrap(
&inner,
k,
k_t,
a,
b,
u_curr,
u_next,
v,
kv,
ktu,
changed,
m,
n,
matrix_cells,
1,
)
}
#[cfg(any(test, feature = "cpu-parity"))]
#[must_use]
#[allow(clippy::too_many_arguments)]
pub fn cpu_ref(
k: &[u32],
k_t: &[u32],
a: &[u32],
b: &[u32],
u_curr: &[u32],
v: &[u32],
m: u32,
n: u32,
max_iterations: u32,
) -> (Vec<u32>, Vec<u32>, u32) {
let mut u = Vec::new();
let mut v_mut = Vec::new();
let mut u_old = Vec::new();
let iters = try_cpu_ref_into(
k,
k_t,
a,
b,
u_curr,
v,
m,
n,
max_iterations,
&mut u,
&mut v_mut,
&mut u_old,
)
.expect("Fix: replace expect with fallible API or document caller precondition; panic only on programmer error - sinkhorn_iterate cpu_ref failed: invalid fixed-point Sinkhorn buffers");
(u, v_mut, iters)
}
#[cfg(any(test, feature = "cpu-parity"))]
#[allow(clippy::too_many_arguments)]
pub fn try_cpu_ref(
k: &[u32],
k_t: &[u32],
a: &[u32],
b: &[u32],
u_curr: &[u32],
v: &[u32],
m: u32,
n: u32,
max_iterations: u32,
) -> Result<(Vec<u32>, Vec<u32>, u32), String> {
let mut u = Vec::new();
let mut v_mut = Vec::new();
let mut u_old = Vec::new();
let iters = try_cpu_ref_into(
k,
k_t,
a,
b,
u_curr,
v,
m,
n,
max_iterations,
&mut u,
&mut v_mut,
&mut u_old,
)?;
Ok((u, v_mut, iters))
}
#[cfg(any(test, feature = "cpu-parity"))]
#[allow(clippy::too_many_arguments)]
pub fn cpu_ref_into(
k: &[u32],
k_t: &[u32],
a: &[u32],
b: &[u32],
u_curr: &[u32],
v: &[u32],
m: u32,
n: u32,
max_iterations: u32,
u_out: &mut Vec<u32>,
v_out: &mut Vec<u32>,
u_old: &mut Vec<u32>,
) -> u32 {
try_cpu_ref_into(
k,
k_t,
a,
b,
u_curr,
v,
m,
n,
max_iterations,
u_out,
v_out,
u_old,
)
.expect("Fix: replace expect with fallible API or document caller precondition; panic only on programmer error - sinkhorn_iterate cpu_ref_into failed: invalid fixed-point Sinkhorn buffers")
}
#[cfg(any(test, feature = "cpu-parity"))]
#[allow(clippy::too_many_arguments)]
pub fn try_cpu_ref_into(
k: &[u32],
k_t: &[u32],
a: &[u32],
b: &[u32],
u_curr: &[u32],
v: &[u32],
m: u32,
n: u32,
max_iterations: u32,
u_out: &mut Vec<u32>,
v_out: &mut Vec<u32>,
u_old: &mut Vec<u32>,
) -> Result<u32, String> {
let (m_usize, n_usize, matrix_cells) = checked_fixed_sinkhorn_shape(m, n)?;
require_fixed_len("k", k.len(), matrix_cells)?;
require_fixed_len("k_t", k_t.len(), matrix_cells)?;
require_fixed_len("a", a.len(), m_usize)?;
require_fixed_len("b", b.len(), n_usize)?;
require_fixed_len("u_curr", u_curr.len(), m_usize)?;
require_fixed_len("v", v.len(), n_usize)?;
reserve_u32_vec(u_out, m_usize, "u output")?;
reserve_u32_vec(v_out, n_usize, "v output")?;
reserve_u32_vec(u_old, m_usize, "u convergence scratch")?;
u_out.clear();
u_out.extend_from_slice(&u_curr[..m_usize]);
v_out.clear();
v_out.extend_from_slice(&v[..n_usize]);
let mut iters = 0;
for iter in 0..max_iterations {
u_old.clear();
u_old.extend_from_slice(u_out);
for i in 0..m_usize {
let mut sum = 0u32;
for j in 0..n_usize {
sum = sum.wrapping_add(k[i * n_usize + j].wrapping_mul(v_out[j]));
}
let divisor = if sum == 0 { 1 } else { sum };
u_out[i] = a[i] / divisor;
}
for j in 0..n_usize {
let mut sum = 0u32;
for i in 0..m_usize {
sum = sum.wrapping_add(k_t[j * m_usize + i].wrapping_mul(u_out[i]));
}
let divisor = if sum == 0 { 1 } else { sum };
v_out[j] = b[j] / divisor;
}
if u_out == u_old {
return Ok(iter);
}
iters = iter + 1;
}
Ok(iters)
}
#[cfg(any(test, feature = "cpu-parity"))]
fn checked_fixed_sinkhorn_shape(m: u32, n: u32) -> Result<(usize, usize, usize), String> {
if m == 0 || n == 0 {
return Err(format!(
"sinkhorn_iterate CPU oracle requires non-zero dimensions, got m={m}, n={n}."
));
}
let m_usize =
usize::try_from(m).map_err(|_| format!("sinkhorn_iterate m={m} does not fit usize."))?;
let n_usize =
usize::try_from(n).map_err(|_| format!("sinkhorn_iterate n={n} does not fit usize."))?;
let matrix_cells = m_usize.checked_mul(n_usize).ok_or_else(|| {
format!("sinkhorn_iterate CPU oracle matrix cells overflow: m={m}, n={n}.")
})?;
Ok((m_usize, n_usize, matrix_cells))
}
#[cfg(any(test, feature = "cpu-parity"))]
fn require_fixed_len(name: &str, got: usize, need: usize) -> Result<(), String> {
if got < need {
Err(format!(
"sinkhorn_iterate CPU oracle buffer `{name}` is too short: got {got}, need {need}."
))
} else {
Ok(())
}
}
crate::graph::scratch::define_reserve_graph_capacity!(
reserve_u32_vec,
u32,
"Sinkhorn iterate CPU oracle"
);
#[cfg(feature = "inventory-registry")]
inventory::submit! {
crate::harness::OpEntry::new(
OP_ID,
|| sinkhorn_iterate("k", "kt", "a", "b", "uc", "un", "v", "kv", "ktu", "c", 2, 2, 5),
Some(|| {
let to_bytes = |w: &[u32]| crate::wire::pack_u32_slice(w);
vec![vec![
to_bytes(&[65536, 65536]), to_bytes(&[0, 0]), to_bytes(&[0]), to_bytes(&[65536, 65536, 65536, 65536]), to_bytes(&[65536, 65536, 65536, 65536]), to_bytes(&[32768, 32768]), to_bytes(&[32768, 32768]), to_bytes(&[65536, 65536]), to_bytes(&[0, 0]), to_bytes(&[0, 0]), ]]
}),
Some(|| {
let to_bytes = |w: &[u32]| crate::wire::pack_u32_slice(w);
vec![vec![
to_bytes(&[32768, 32768]), to_bytes(&[32768, 32768]), to_bytes(&[0]), to_bytes(&[32768, 32768]), to_bytes(&[0, 0]), to_bytes(&[0, 0]), ]]
}),
)
}
#[cfg(test)]
mod tests;
#[must_use]
#[cfg(any(test, feature = "cpu-parity"))]
pub fn sinkhorn_iterate_f64(
k: &[f64],
a: &[f64],
b: &[f64],
tolerance: f64,
max_iterations: u32,
) -> (Vec<f64>, Vec<f64>, u32) {
let mut u = Vec::new();
let mut v = Vec::new();
let mut u_old = Vec::new();
let iters = sinkhorn_iterate_f64_into(
k,
a,
b,
tolerance,
max_iterations,
&mut u,
&mut v,
&mut u_old,
);
(u, v, iters)
}
#[cfg(any(test, feature = "cpu-parity"))]
pub fn try_sinkhorn_iterate_f64(
k: &[f64],
a: &[f64],
b: &[f64],
tolerance: f64,
max_iterations: u32,
) -> Result<(Vec<f64>, Vec<f64>, u32), String> {
let mut u = Vec::new();
let mut v = Vec::new();
let mut u_old = Vec::new();
let iters = try_sinkhorn_iterate_f64_into(
k,
a,
b,
tolerance,
max_iterations,
&mut u,
&mut v,
&mut u_old,
)?;
Ok((u, v, iters))
}
#[allow(clippy::too_many_arguments)]
#[cfg(any(test, feature = "cpu-parity"))]
pub fn sinkhorn_iterate_f64_into(
k: &[f64],
a: &[f64],
b: &[f64],
tolerance: f64,
max_iterations: u32,
u: &mut Vec<f64>,
v: &mut Vec<f64>,
u_old: &mut Vec<f64>,
) -> u32 {
match try_sinkhorn_iterate_f64_into(k, a, b, tolerance, max_iterations, u, v, u_old) {
Ok(iters) => iters,
Err(error) => panic!("vyre-primitives Sinkhorn iterate CPU reference failed: {error}"),
}
}
#[allow(clippy::too_many_arguments)]
#[cfg(any(test, feature = "cpu-parity"))]
pub fn try_sinkhorn_iterate_f64_into(
k: &[f64],
a: &[f64],
b: &[f64],
tolerance: f64,
max_iterations: u32,
u: &mut Vec<f64>,
v: &mut Vec<f64>,
u_old: &mut Vec<f64>,
) -> Result<u32, String> {
let m = a.len();
let n = b.len();
if k.len() != m * n || tolerance <= 0.0 || !tolerance.is_finite() {
return Err(format!(
"sinkhorn_iterate_f64 requires k.len()==a.len()*b.len() and finite positive tolerance, got k={}, m={m}, n={n}, tolerance={tolerance}.",
k.len()
));
}
reserve_f64_vec(u, m, "u output")?;
reserve_f64_vec(v, n, "v output")?;
reserve_f64_vec(u_old, m, "u convergence scratch")?;
u.clear();
v.clear();
u_old.clear();
u.resize(m, 1.0_f64);
v.resize(n, 1.0_f64);
for iter in 0..max_iterations {
u_old.clear();
u_old.extend_from_slice(u);
for i in 0..m {
let mut sum = 0.0_f64;
for j in 0..n {
sum += k[i * n + j] * v[j];
}
u[i] = if sum == 0.0 { 0.0 } else { a[i] / sum };
}
for j in 0..n {
let mut sum = 0.0_f64;
for i in 0..m {
sum += k[i * n + j] * u[i];
}
v[j] = if sum == 0.0 { 0.0 } else { b[j] / sum };
}
let max_delta = u
.iter()
.zip(u_old.iter())
.map(|(new, old)| (new - old).abs())
.fold(0.0_f64, f64::max);
if max_delta < tolerance {
return Ok(iter + 1);
}
}
Ok(max_iterations)
}
crate::graph::scratch::define_reserve_graph_capacity!(
reserve_f64_vec,
f64,
"Sinkhorn iterate f64 CPU oracle"
);
#[cfg(any(test, feature = "cpu-parity"))]
fn max_residual(target: &[f64], sum_at: impl Fn(usize) -> f64) -> f64 {
target
.iter()
.enumerate()
.map(|(index, expected)| (sum_at(index) - expected).abs())
.fold(0.0_f64, f64::max)
}
#[cfg(any(test, feature = "cpu-parity"))]
#[derive(Clone, Copy)]
enum ResidualAxis {
Row,
Column,
}
#[cfg(any(test, feature = "cpu-parity"))]
fn sinkhorn_residual(k: &[f64], u: &[f64], v: &[f64], target: &[f64], axis: ResidualAxis) -> f64 {
let m = u.len();
let n = v.len();
assert_eq!(k.len(), m * n);
match axis {
ResidualAxis::Row => {
assert_eq!(target.len(), m);
max_residual(target, |i| (0..n).map(|j| u[i] * k[i * n + j] * v[j]).sum())
}
ResidualAxis::Column => {
assert_eq!(target.len(), n);
max_residual(target, |j| (0..m).map(|i| u[i] * k[i * n + j] * v[j]).sum())
}
}
}
#[must_use]
#[cfg(any(test, feature = "cpu-parity"))]
pub fn sinkhorn_row_residual(k: &[f64], u: &[f64], v: &[f64], a: &[f64]) -> f64 {
sinkhorn_residual(k, u, v, a, ResidualAxis::Row)
}
#[must_use]
#[cfg(any(test, feature = "cpu-parity"))]
pub fn sinkhorn_col_residual(k: &[f64], u: &[f64], v: &[f64], b: &[f64]) -> f64 {
sinkhorn_residual(k, u, v, b, ResidualAxis::Column)
}
#[cfg(test)]
mod f64_tests;