use std::sync::Arc;
use shape_value::heap_value::MatrixData;
pub const MATRIX_DATA_OFFSET: i32 = 0;
pub const MATRIX_ROWS_OFFSET: i32 = 8;
pub const MATRIX_COLS_OFFSET: i32 = 12;
pub const MATRIX_TOTAL_LEN_OFFSET: i32 = 16;
pub const MATRIX_OWNER_OFFSET: i32 = 24;
#[repr(C)]
pub struct JitMatrix {
pub data: *const f64,
pub rows: u32,
pub cols: u32,
pub total_len: u64,
owner: *const MatrixData,
}
impl JitMatrix {
pub fn from_arc(arc: &Arc<MatrixData>) -> Self {
let mat = arc.as_ref();
let data = mat.data.as_slice().as_ptr();
let rows = mat.rows;
let cols = mat.cols;
let total_len = mat.data.len() as u64;
let owner = Arc::into_raw(Arc::clone(arc));
Self {
data,
rows,
cols,
total_len,
owner,
}
}
pub fn to_arc(&self) -> Arc<MatrixData> {
assert!(!self.owner.is_null(), "JitMatrix has null owner");
let arc = unsafe { Arc::from_raw(self.owner) };
let cloned = Arc::clone(&arc);
std::mem::forget(arc);
cloned
}
}
impl Drop for JitMatrix {
fn drop(&mut self) {
if !self.owner.is_null() {
unsafe {
let _ = Arc::from_raw(self.owner);
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use shape_value::aligned_vec::AlignedVec;
fn make_test_matrix(rows: u32, cols: u32) -> Arc<MatrixData> {
let n = (rows as usize) * (cols as usize);
let mut data = AlignedVec::with_capacity(n);
for i in 0..n {
data.push(i as f64);
}
Arc::new(MatrixData::from_flat(data, rows, cols))
}
#[test]
fn test_layout() {
assert_eq!(std::mem::offset_of!(JitMatrix, data), MATRIX_DATA_OFFSET as usize);
assert_eq!(std::mem::offset_of!(JitMatrix, rows), MATRIX_ROWS_OFFSET as usize);
assert_eq!(std::mem::offset_of!(JitMatrix, cols), MATRIX_COLS_OFFSET as usize);
assert_eq!(std::mem::offset_of!(JitMatrix, total_len), MATRIX_TOTAL_LEN_OFFSET as usize);
assert_eq!(std::mem::offset_of!(JitMatrix, owner), MATRIX_OWNER_OFFSET as usize);
assert_eq!(std::mem::size_of::<JitMatrix>(), 32);
}
#[test]
fn test_round_trip() {
let arc = make_test_matrix(3, 4);
let jm = JitMatrix::from_arc(&arc);
assert_eq!(jm.rows, 3);
assert_eq!(jm.cols, 4);
assert_eq!(jm.total_len, 12);
let slice = unsafe { std::slice::from_raw_parts(jm.data, jm.total_len as usize) };
assert_eq!(slice[0], 0.0);
assert_eq!(slice[11], 11.0);
let recovered = jm.to_arc();
assert_eq!(recovered.rows, 3);
assert_eq!(recovered.cols, 4);
assert_eq!(recovered.data[0], 0.0);
assert_eq!(recovered.data[11], 11.0);
assert_eq!(arc.data[5], 5.0);
}
#[test]
fn test_arc_refcount() {
let arc = make_test_matrix(2, 2);
assert_eq!(Arc::strong_count(&arc), 1);
let jm = JitMatrix::from_arc(&arc);
assert_eq!(Arc::strong_count(&arc), 2);
let recovered = jm.to_arc();
assert_eq!(Arc::strong_count(&arc), 3);
drop(recovered);
assert_eq!(Arc::strong_count(&arc), 2);
drop(jm);
assert_eq!(Arc::strong_count(&arc), 1); }
}