#![allow(dead_code)]
use std::marker::PhantomData;
use furiosa_mapping::*;
use crate::scalar::{MaterializableScalar, Scalar};
use crate::tensor::Tensor;
#[derive(Debug)]
pub struct MemTensor<D: Scalar, Buf: M> {
inner: Tensor<D, Buf>,
}
#[derive(Debug)]
pub struct StreamTensor<'l, D: Scalar, Time: M, Packet: M> {
inner: Tensor<D, Pair<Time, Packet>>,
_marker: PhantomData<&'l ()>,
}
impl<D: Scalar, Buf: M> MemTensor<D, Buf> {
pub fn read<'l, Time: M, Packet: M>(&'l self) -> StreamTensor<'l, D, Time, Packet> {
StreamTensor {
inner: self.inner.transpose(true),
_marker: PhantomData,
}
}
pub fn write<'l, Time: M, Packet: M>(&mut self, stream: StreamTensor<'l, D, Time, Packet>) {
self.inner = stream.inner.transpose(false);
}
}
impl<D: Scalar, Buf: M> MemTensor<D, Buf> {
pub fn from_vec(data: impl IntoIterator<Item = D>) -> Self {
Self {
inner: Tensor::from_vec(data),
}
}
pub fn into_vec(self) -> Vec<D>
where
D: MaterializableScalar,
{
self.inner.into_vec()
}
}
impl<'l, D: Scalar, Time: M, Packet: M> StreamTensor<'l, D, Time, Packet> {
pub fn into_vec(self) -> Vec<D>
where
D: MaterializableScalar,
{
self.inner.into_vec()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn read_identity_preserves_order() {
axes![A = 2, B = 3];
let buf = MemTensor::<i32, m![A, B]>::from_vec(vec![10, 11, 12, 20, 21, 22]);
let stream = buf.read::<m![A], m![B]>();
assert_eq!(
stream.into_vec(),
Tensor::<i32, m![A, B]>::from_vec(vec![10, 11, 12, 20, 21, 22]).into_vec()
);
}
#[test]
fn read_reorders_axes() {
axes![A = 2, B = 3];
let buf = MemTensor::<i32, m![A, B]>::from_vec(vec![10, 11, 12, 20, 21, 22]);
let stream = buf.read::<m![B], m![A]>();
assert_eq!(
stream.into_vec(),
Tensor::<i32, m![A, B]>::from_vec(vec![10, 20, 11, 21, 12, 22]).into_vec()
);
}
#[test]
fn write_inverts_read() {
axes![A = 2, B = 3];
let original = vec![10, 11, 12, 20, 21, 22];
let buf = MemTensor::<i32, m![A, B]>::from_vec(original.clone());
let stream = buf.read::<m![B], m![A]>();
let mut sink = MemTensor::<i32, m![A, B]>::from_vec(vec![0; 6]);
sink.write(stream);
assert_eq!(sink.into_vec(), Tensor::<i32, m![A, B]>::from_vec(original).into_vec());
}
#[test]
fn read_splits_axis() {
axes![A = 4];
let buf = MemTensor::<i32, m![A]>::from_vec(vec![0, 1, 2, 3]);
let stream = buf.read::<m![A % 2], m![A / 2]>();
assert_eq!(
stream.into_vec(),
Tensor::<i32, m![A]>::from_vec(vec![0, 2, 1, 3]).into_vec()
);
}
}