use std::cell::UnsafeCell;
#[derive(Debug)]
pub struct TriMatrix {
size: usize,
linkage_same: Vec<UnsafeCell<f64>>,
linkage_diff: Vec<UnsafeCell<f64>>,
}
unsafe impl Sync for TriMatrix {}
impl Clone for TriMatrix {
fn clone(&self) -> Self {
let linkage_same: Vec<UnsafeCell<f64>> = self
.linkage_same
.iter()
.map(|cell| {
unsafe { UnsafeCell::new(*cell.get()) }
})
.collect();
let linkage_diff: Vec<UnsafeCell<f64>> = self
.linkage_diff
.iter()
.map(|cell| {
unsafe { UnsafeCell::new(*cell.get()) }
})
.collect();
Self {
size: self.size,
linkage_same,
linkage_diff,
}
}
}
impl TriMatrix {
pub fn new(size: usize) -> Self {
let capacity = if size > 0 { (size * (size - 1)) / 2 } else { 0 };
Self {
size,
linkage_same: (0..capacity).map(|_| UnsafeCell::new(0.0)).collect(),
linkage_diff: (0..capacity).map(|_| UnsafeCell::new(0.0)).collect(),
}
}
#[inline(always)]
fn index(&self, i: usize, j: usize) -> usize {
debug_assert!(i < self.size && j < self.size);
debug_assert!(i != j, "Cannot access diagonal elements");
let (row, col) = if i > j { (i, j) } else { (j, i) };
((row * (row - 1)) >> 1) + col
}
#[inline]
pub fn write(&self, i: usize, j: usize, value: (f64, f64)) {
if i == j {
return; }
let idx = self.index(i, j);
unsafe {
*self.linkage_same[idx].get() = value.0;
*self.linkage_diff[idx].get() = value.1;
}
}
#[inline]
pub fn read(&self, i: usize, j: usize) -> (f64, f64) {
if i == j {
return (1.0, 1.0);
}
let idx = self.index(i, j);
unsafe { (*self.linkage_same[idx].get(), *self.linkage_diff[idx].get()) }
}
pub fn size(&self) -> usize {
self.size
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_tri_matrix() {
let matrix = TriMatrix::new(4);
matrix.write(0, 1, (0.5, 0.3));
matrix.write(2, 3, (0.8, 0.2));
assert_eq!(matrix.read(0, 1), (0.5, 0.3));
assert_eq!(matrix.read(1, 0), (0.5, 0.3)); assert_eq!(matrix.read(2, 3), (0.8, 0.2));
assert_eq!(matrix.read(0, 0), (1.0, 1.0)); }
}