From e0e8977658462d491c263c9f9638e7f70528c9cb Mon Sep 17 00:00:00 2001 From: Nick Wall <46641379+walln@users.noreply.github.com> Date: Mon, 11 Sep 2023 17:40:26 -0500 Subject: [PATCH 1/4] feat: refactor ops and implement backend modules --- .vscode/settings.json | 2 +- Cargo.toml | 7 + src/backend/backend.rs | 95 ++++++ src/backend/cpu_backend.rs | 642 ++++++++++++++++++++++++++++++++++--- src/backend/mod.rs | 5 + src/backend/mps_backend.rs | 145 +++++++++ src/backprop.rs | 108 +++---- src/device.rs | 75 ++++- src/dtype.rs | 46 ++- src/error.rs | 37 ++- src/gradient_store.rs | 47 +++ src/index.rs | 26 +- src/layout.rs | 138 ++++++++ src/lib.rs | 5 +- src/operation.rs | 88 ++++- src/shape.rs | 71 +++- src/storage.rs | 229 ++++--------- src/tensor.rs | 402 +++++++++++++---------- src/utils.rs | 12 + tests/gradient_tests.rs | 23 +- tests/tensor_tests.rs | 46 +-- 21 files changed, 1722 insertions(+), 527 deletions(-) create mode 100644 src/backend/backend.rs create mode 100644 src/backend/mps_backend.rs create mode 100644 src/gradient_store.rs create mode 100644 src/layout.rs create mode 100644 src/utils.rs diff --git a/.vscode/settings.json b/.vscode/settings.json index 0ee0666..77cc3a7 100644 --- a/.vscode/settings.json +++ b/.vscode/settings.json @@ -1,3 +1,3 @@ { - "rust-analyzer.linkedProjects": ["./Cargo.toml"] + "rust-analyzer.linkedProjects": ["./Cargo.toml", "./Cargo.toml"] } diff --git a/Cargo.toml b/Cargo.toml index aa867f3..32d6a05 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -9,4 +9,11 @@ readme = "README.md" [dependencies] anyhow = "1.0.75" +num-traits = "0.2.16" thiserror = "1" + +# TODO: Switch back to the official gemm implementation once something similar to +# https://github.com/sarah-ek/gemm/pull/8 is available. +gemm = { git = "https://github.com/LaurentMazare/gemm.git", branch = "f16-vectorize-pack" } +rand = "0.8.5" +num_cpus = "1.16.0" diff --git a/src/backend/backend.rs b/src/backend/backend.rs new file mode 100644 index 0000000..44d2d75 --- /dev/null +++ b/src/backend/backend.rs @@ -0,0 +1,95 @@ +use crate::operation::{BinaryOperation, UnaryOperation}; +use crate::{CPUStorage, DType, Layout, Result, Shape}; + +pub(crate) trait BackendStorage: Sized { + type Device: BackendDevice; + + fn to_cpu(&self) -> Result; + + fn dtype(&self) -> DType; + + fn to_dtype(&self, layout: &Layout, dtype: DType) -> Result; + + fn device(&self) -> &Self::Device; + + fn try_clone(&self, layout: &Layout) -> Result; + + fn copy_strided_source( + &self, + destination: &mut Self, + destination_offset: usize, + source_layout: &Layout, + ) -> Result<()>; + + fn unary_operation(&self, layout: &Layout) -> Result; + + fn binary_operation( + &self, + rhs: &Self, + lhs_layout: &Layout, + rhs_layout: &Layout, + ) -> Result; + + fn affine(&self, layout: &Layout, add: f64, mul: f64) -> Result; + + fn matmul( + &self, + rhs: &Self, + bmnk: (usize, usize, usize, usize), + lhs_layout: &Layout, + rhs_layout: &Layout, + ) -> Result; + + fn where_condition( + &self, + _: &Layout, + _: &Self, + _: &Layout, + _: &Self, + _: &Layout, + ) -> Result; + + // fn conv1d( + // &self, + // _l: &Layout, + // _kernel: &Self, + // _kernel_l: &Layout, + // _params: &crate::conv::ParamsConv1D, + // ) -> Result; + + fn embedding(&self, _: &Layout, _: &Self, _: &Layout) -> Result; + + fn sum(&self, _: &Layout, _: &[usize]) -> Result; +} + +pub(crate) trait BackendDevice: Sized + std::fmt::Debug + Clone { + type Storage: BackendStorage; + + fn new(_: usize) -> Result; + + fn location(&self) -> crate::device::DeviceLocation; + + fn same_device(&self, _: &Self) -> bool; + + fn zeros_impl(&self, shape: &Shape, dtype: DType) -> Result; + + fn ones_impl(&self, shape: &Shape, dtype: DType) -> Result; + + fn rand_uniform( + &self, + shape: &Shape, + dtype: DType, + lower_bound: f64, + upper_bound: f64, + ) -> Result; + + fn rand_normal( + &self, + shape: &Shape, + dtype: DType, + mean: f64, + std: f64, + ) -> Result; + + fn from_cpu(&self, storage: &CPUStorage) -> Result; +} diff --git a/src/backend/cpu_backend.rs b/src/backend/cpu_backend.rs index 1bfd9c4..48bcbb9 100644 --- a/src/backend/cpu_backend.rs +++ b/src/backend/cpu_backend.rs @@ -1,95 +1,633 @@ -use crate::storage::{BinaryOperation, UnaryOperation}; -use crate::{index::StridedIndex, DType, Error, Result, Shape}; +use rand::distributions::uniform; + +use crate::backend::backend::BackendStorage; +use crate::operation::{BinaryOperation, UnaryOperation}; +use crate::{DType, Error, Layout, Result, Shape, WithDType}; + +use super::backend::BackendDevice; #[derive(Debug, Clone)] pub enum CPUStorage { F32(Vec), F64(Vec), + U32(Vec), } -impl CPUStorage { - pub(crate) fn dtype(&self) -> DType { +#[derive(Debug, Clone)] +pub struct CPUDevice; + +impl BackendStorage for CPUStorage { + type Device = CPUDevice; + + fn device(&self) -> &Self::Device { + &CPUDevice + } + + fn dtype(&self) -> DType { match self { - CPUStorage::F32(_) => DType::F32, - CPUStorage::F64(_) => DType::F64, + Self::F32(_) => DType::F32, + Self::F64(_) => DType::F64, + Self::U32(_) => DType::U32, } } - pub(crate) fn affine( - &self, - shape: &Shape, - stride: &[usize], - mul: f64, - add: f64, - ) -> Result { - match self { - Self::F32(storage) => { - let index = StridedIndex::new(shape.dims(), stride); - let mul = mul as f32; - let add = add as f32; - let data = index.map(|i| storage[i] * mul + add).collect(); + fn to_dtype(&self, layout: &Layout, dtype: DType) -> Result { + match (self, dtype) { + (Self::F32(storage), DType::F32) => { + let data = unary_map(storage, layout, |v| v); Ok(Self::F32(data)) } - Self::F64(storage) => { - let index = StridedIndex::new(shape.dims(), stride); - let data = index.map(|i| storage[i] * mul + add).collect(); + (Self::F32(storage), DType::F64) => { + let data = unary_map(storage, layout, |v| v as f64); + Ok(Self::F64(data)) + } + (Self::F32(storage), DType::U32) => { + let data = unary_map(storage, layout, |v| v as u32); + Ok(Self::U32(data)) + } + (Self::F64(storage), DType::F32) => { + let data = unary_map(storage, layout, |v| v as f32); + Ok(Self::F32(data)) + } + (Self::F64(storage), DType::F64) => { + let data = unary_map(storage, layout, |v| v); + Ok(Self::F64(data)) + } + (Self::F64(storage), DType::U32) => { + let data = unary_map(storage, layout, |v| v as u32); + Ok(Self::U32(data)) + } + + (Self::U32(storage), DType::U32) => { + let data = unary_map(storage, layout, |v| v); + Ok(Self::U32(data)) + } + (Self::U32(storage), DType::F32) => { + let data = unary_map(storage, layout, |v| v as f32); + Ok(Self::F32(data)) + } + (Self::U32(storage), DType::F64) => { + let data = unary_map(storage, layout, |v| v as f64); Ok(Self::F64(data)) } } } - pub(crate) fn unary_impl( + fn to_cpu(&self) -> Result { + Ok(self.clone()) + } + + fn try_clone(&self, layout: &Layout) -> Result { + Ok(self.clone()) + } + + fn binary_operation( &self, - shape: &Shape, - stride: &[usize], + rhs: &Self, + lhs_layout: &Layout, + rhs_layout: &Layout, ) -> Result { + match (self, rhs) { + (Self::F32(lhs), Self::F32(rhs)) => { + let data = binary_map(lhs_layout, rhs_layout, lhs, rhs, T::f32); + Ok(Self::F32(data)) + } + (Self::F64(lhs), Self::F64(rhs)) => { + let data = binary_map(lhs_layout, rhs_layout, lhs, rhs, T::f64); + Ok(Self::F64(data)) + } + _ => Err(Error::BinaryOperationDTypeMismatch { + lhs: self.dtype(), + rhs: rhs.dtype(), + op: T::NAME, + }), + } + } + + fn unary_operation(&self, layout: &Layout) -> Result { match self { Self::F32(storage) => { - let index = StridedIndex::new(shape.dims(), stride); - let data = index.map(|i| T::f32(storage[i])).collect(); + let data = unary_map(storage, layout, T::f32); Ok(Self::F32(data)) } Self::F64(storage) => { - let index = StridedIndex::new(shape.dims(), stride); - let data = index.map(|i| T::f64(storage[i])).collect(); + let data = unary_map(storage, layout, T::f64); Ok(Self::F64(data)) } + Self::U32(storage) => { + let data = unary_map(storage, layout, T::u32); + Ok(Self::U32(data)) + } } } - pub(crate) fn binary_operation( + fn affine(&self, layout: &Layout, add: f64, mul: f64) -> Result { + Affine(mul, add).map(self, layout) + } + + fn sum(&self, layout: &Layout, sum_dims: &[usize]) -> Result { + let source_dims = layout.dims(); + let mut destination_dims = source_dims.to_vec(); + for &sum_dim in sum_dims.iter() { + destination_dims[sum_dim] = 1; + } + let destination_shape = Shape::from(destination_dims); + let mut sum_dims = sum_dims.to_vec(); + + // When converting the indicies sort the sum dims as they are processed + // the dimensions are processed from left to right. + sum_dims.sort(); + let sum_dims_and_stride: Vec<_> = sum_dims + .iter() + .map(|&d| { + ( + source_dims[d], + source_dims[d + 1..].iter().product::(), + ) + }) + .collect(); + Sum { + destination_shape: &destination_shape, + sum_dims_and_stride, + } + .map(self, layout) + } + + fn matmul( &self, rhs: &Self, - shape: &Shape, - lhs_stride: &[usize], - rhs_stride: &[usize], + bmnk: (usize, usize, usize, usize), + lhs_layout: &Layout, + rhs_layout: &Layout, ) -> Result { - match (self, rhs) { - (CPUStorage::F32(lhs), CPUStorage::F32(rhs)) => { - let lhs_index = StridedIndex::new(shape.dims(), lhs_stride); - let rhs_index = StridedIndex::new(shape.dims(), rhs_stride); - let data = lhs_index - .zip(rhs_index) - .map(|(lhs_offset, rhs_offset)| T::f32(lhs[lhs_offset], rhs[rhs_offset])) - .collect(); + MatMul(bmnk).map(self, lhs_layout, rhs, rhs_layout) + } - Ok(Self::F32(data)) + fn embedding(&self, lhs_layout: &Layout, rhs: &Self, rhs_layout: &Layout) -> Result { + let lhs = self.as_slice::()?; + let (vocab_size, hidden_size) = rhs_layout.shape().rank_two()?; + Embedding { + vocab_size, + hidden_size, + ids: lhs, + ids_layout: lhs_layout, + } + .map(rhs, rhs_layout) + } + + fn where_condition( + &self, + layout: &Layout, + t: &Self, + t_layout: &Layout, + f: &Self, + f_layout: &Layout, + ) -> Result { + let pred = self.as_slice::()?; + WhereCondition(pred, layout).map(t, t_layout, f, f_layout) + } + + fn copy_strided_source( + &self, + destination: &mut Self, + destination_offset: usize, + source_layout: &Layout, + ) -> Result<()> { + match (self, destination) { + (Self::F32(src), Self::F32(dst)) => { + copy_strided_source_(src, dst, destination_offset, source_layout) } - (CPUStorage::F64(lhs), CPUStorage::F64(rhs)) => { - let lhs_index = StridedIndex::new(shape.dims(), lhs_stride); - let rhs_index = StridedIndex::new(shape.dims(), rhs_stride); - let data = lhs_index - .zip(rhs_index) - .map(|(lhs_offset, rhs_offset)| T::f64(lhs[lhs_offset], rhs[rhs_offset])) - .collect(); + (Self::F64(src), Self::F64(dst)) => { + copy_strided_source_(src, dst, destination_offset, source_layout) + } + (_, destination) => { + // This should be covered by the dtype check above. + return Err(Error::BinaryOperationDTypeMismatch { + lhs: self.dtype(), + rhs: destination.dtype(), + op: "copy_strided", + }); + } + } + Ok(()) + } +} - Ok(Self::F64(data)) +impl BackendDevice for CPUDevice { + type Storage = CPUStorage; + + fn new(_: usize) -> Result { + Ok(Self) + } + + fn location(&self) -> crate::device::DeviceLocation { + crate::device::DeviceLocation::CPU + } + + fn same_device(&self, _: &Self) -> bool { + true + } + + fn from_cpu(&self, storage: &CPUStorage) -> Result { + Ok(storage.clone()) + } + + fn zeros_impl(&self, shape: &Shape, dtype: DType) -> Result { + let elem_count = shape.elem_count(); + match dtype { + DType::F32 => Ok(Self::Storage::F32(vec![0.0; elem_count])), + DType::F64 => Ok(Self::Storage::F64(vec![0.0; elem_count])), + DType::U32 => Ok(Self::Storage::U32(vec![0; elem_count])), + } + } + + fn ones_impl(&self, shape: &Shape, dtype: DType) -> Result { + let elem_count = shape.elem_count(); + match dtype { + DType::F32 => Ok(Self::Storage::F32(vec![1.0; elem_count])), + DType::F64 => Ok(Self::Storage::F64(vec![1.0; elem_count])), + DType::U32 => Ok(Self::Storage::U32(vec![1; elem_count])), + } + } + + fn rand_uniform( + &self, + shape: &Shape, + dtype: DType, + lower_bound: f64, + upper_bound: f64, + ) -> Result { + use rand::prelude::*; + let elem_count = shape.elem_count(); + let mut rng = rand::thread_rng(); + + match dtype { + DType::F32 => { + let mut data = Vec::new(); + data.reserve(elem_count); + let uniform = + rand::distributions::Uniform::new(lower_bound as f32, upper_bound as f32); + for _ in 0..elem_count { + data.push(rng.sample::(uniform)) + } + Ok(CPUStorage::F32(data)) + } + DType::F64 => { + let mut data = Vec::new(); + data.reserve(elem_count); + let uniform = uniform::Uniform::new(lower_bound, upper_bound); + for _ in 0..elem_count { + data.push(rng.sample::(uniform)) + } + Ok(CPUStorage::F64(data)) + } + _ => Err(Error::UnsupportedDTypeForOperation { + dtype, + op: "rand_normal", + }), + } + } + + fn rand_normal( + &self, + shape: &Shape, + dtype: DType, + mean: f64, + std: f64, + ) -> Result { + use rand::prelude::*; + let elem_count = shape.elem_count(); + let mut rng = rand::thread_rng(); + + match dtype { + DType::F32 => { + let mut data = Vec::new(); + data.reserve(elem_count); + let mean = mean as f32; + let std = std as f32; + for _ in 0..elem_count { + data.push(rng.sample::(rand::distributions::Standard) * std + mean) + } + Ok(CPUStorage::F32(data)) + } + DType::F64 => { + let mut data = Vec::new(); + data.reserve(elem_count); + for _ in 0..elem_count { + data.push(rng.sample::(rand::distributions::Standard) * std + mean) + } + Ok(CPUStorage::F64(data)) + } + _ => Err(Error::UnsupportedDTypeForOperation { + dtype, + op: "rand_uniform", + }), + } + } +} + +impl CPUStorage { + pub fn as_slice(&self) -> Result<&[D]> { + D::cpu_storage_slice(self) + } +} + +/// Perform a mapping function of a unary operation on a contiguous slice of data. +fn unary_map U>( + values: &[T], + layout: &Layout, + mut f: F, +) -> Vec { + match layout.contiguous_offsets() { + Some((o1, o2)) => values[o1..o2].iter().map(|&v| f(v)).collect(), + None => layout.strided_index().map(|i| f(values[i])).collect(), + } +} + +/// Perform a mapping function of a binary operation on contiguous slices of data. +fn binary_map T>( + lhs_layout: &Layout, + rhs_layout: &Layout, + lhs: &[T], + rhs: &[T], + mut f: F, +) -> Vec { + match ( + lhs_layout.contiguous_offsets(), + rhs_layout.contiguous_offsets(), + ) { + (Some((o_l1, o_l2)), Some((o_r1, o_r2))) => lhs[o_l1..o_l2] + .iter() + .zip(rhs[o_r1..o_r2].iter()) + .map(|(&l, &r)| f(l, r)) + .collect(), + _ => lhs_layout + .strided_index() + .zip(rhs_layout.strided_index()) + .map(|(lhs_i, rhs_i)| f(lhs[lhs_i], rhs[rhs_i])) + .collect(), + } +} + +trait UnaryMappable { + fn f(&self, values: &[T], layout: &Layout) -> Result>; + + fn map(&self, storage: &CPUStorage, layout: &Layout) -> Result { + match storage { + CPUStorage::F32(values) => Ok(CPUStorage::F32(self.f(values, layout)?)), + CPUStorage::F64(values) => Ok(CPUStorage::F64(self.f(values, layout)?)), + CPUStorage::U32(values) => Ok(CPUStorage::U32(self.f(values, layout)?)), + } + } +} + +struct Affine(f64, f64); + +impl UnaryMappable for Affine { + fn f(&self, values: &[T], layout: &Layout) -> Result> { + let mul = T::from_f64(self.0); + let add = T::from_f64(self.1); + Ok(unary_map(values, layout, |v| v * mul + add)) + } +} + +struct Sum<'a> { + destination_shape: &'a Shape, + sum_dims_and_stride: Vec<(usize, usize)>, +} + +impl<'a> UnaryMappable for Sum<'a> { + fn f(&self, source: &[T], source_layout: &Layout) -> Result> { + let mut destination = vec![T::zero(); self.destination_shape.elem_count()]; + for (unstrided_index, source_index) in source_layout.strided_index().enumerate() { + let mut destination_index = unstrided_index; + // Set the sum_dims indexes to 0. + for &(dim, stride) in self.sum_dims_and_stride.iter() { + // The compiler is able to optimize the following in a single divmod op. + let (pre, post) = (destination_index / stride, destination_index % stride); + destination_index = (pre / dim) * stride + post; + } + destination[destination_index] += source[source_index]; + } + Ok(destination) + } +} + +struct Embedding<'a> { + vocab_size: usize, + hidden_size: usize, + ids: &'a [u32], + ids_layout: &'a Layout, +} + +impl<'a> UnaryMappable for Embedding<'a> { + fn f(&self, values: &[T], layout: &Layout) -> Result> { + // TODO: We assume that values is contiguous here. + let values = &values[layout.start_offset()..]; + let mut vals = Vec::with_capacity(self.ids_layout.shape().elem_count() * self.hidden_size); + // TODO: Optimize for the case where ids are contiguous. + for index in self.ids_layout.strided_index() { + let index = self.ids[index]; + let index = index.try_into().map_err(|_| Error::InvalidIndex { + index: index.try_into().unwrap(), + vocab_size: self.vocab_size, + op: "embedding", + })?; + if index >= self.vocab_size { + return Err(Error::InvalidIndex { + index, + vocab_size: self.vocab_size, + op: "take", + }); + } else { + let hidden_size = self.hidden_size; + vals.extend(&values[hidden_size * index..hidden_size * (index + 1)]); + } + } + Ok(vals) + } +} + +trait BinaryMappable { + const OP: &'static str; + fn f(&self, v1: &[T], l1: &Layout, v2: &[T], l2: &Layout) -> Result>; + + fn map( + &self, + lhs: &CPUStorage, + lhs_layout: &Layout, + rhs: &CPUStorage, + rhs_layout: &Layout, + ) -> Result { + match (lhs, rhs) { + (CPUStorage::F32(lhs), CPUStorage::F32(rhs)) => { + Ok(CPUStorage::F32(self.f(lhs, lhs_layout, rhs, rhs_layout)?)) + } + (CPUStorage::F64(lhs), CPUStorage::F64(rhs)) => { + Ok(CPUStorage::F64(self.f(lhs, lhs_layout, rhs, rhs_layout)?)) } _ => Err(Error::BinaryOperationDTypeMismatch { - lhs: self.dtype(), + lhs: lhs.dtype(), rhs: rhs.dtype(), - op: T::NAME, + op: Self::OP, }), } } } + +struct MatMul((usize, usize, usize, usize)); + +impl MatMul { + fn striding_error(&self, lhs_layout: &Layout, rhs_layout: &Layout, msg: &'static str) -> Error { + Error::MatMulUnexpectedStride { + lhs_layout: lhs_layout.clone(), + rhs_layout: rhs_layout.clone(), + bmnk: self.0, + msg, + } + } +} + +impl BinaryMappable for MatMul { + const OP: &'static str = "matmul"; + + fn f( + &self, + lhs: &[T], + lhs_layout: &Layout, + rhs: &[T], + rhs_layout: &Layout, + ) -> Result> { + use gemm::{gemm, Parallelism}; + let (b, m, n, k) = self.0; + let lhs = &lhs[lhs_layout.start_offset()..]; + let rhs = &rhs[rhs_layout.start_offset()..]; + + let lhs_stride = lhs_layout.stride(); + let rhs_stride = rhs_layout.stride(); + let rank = lhs_stride.len(); + let lhs_cs = lhs_stride[rank - 1]; + let lhs_rs = lhs_stride[rank - 2]; + + let rhs_cs = rhs_stride[rank - 1]; + let rhs_rs = rhs_stride[rank - 2]; + + let a_skip: usize = match lhs_stride[..rank - 2] { + [s1, stride] if s1 == stride * lhs_layout.dims()[1] => stride, + [stride] => stride, + [] => m * k, + _ => Err(self.striding_error(lhs_layout, rhs_layout, "non-contiguous lhs"))?, + }; + let b_skip: usize = match rhs_stride[..rank - 2] { + [s1, stride] if s1 == stride * rhs_layout.dims()[1] => stride, + [stride] => stride, + [] => n * k, + _ => Err(self.striding_error(lhs_layout, rhs_layout, "non-contiguous rhs"))?, + }; + let c_skip: usize = m * n; + + let dst_shape: Shape = (m, n).into(); + let dst_strides = dst_shape.stride_contiguous(); + let dst_rs = dst_strides[0]; + let dst_cs = dst_strides[1]; + + let mut dst = vec![T::zero(); b * m * n]; + let num_threads = crate::utils::get_num_threads(); + let parallelism = if num_threads > 1 { + Parallelism::Rayon(num_threads) + } else { + Parallelism::None + }; + for step in 0..b { + let lhs_p = &lhs[step * a_skip..]; + let rhs_p = &rhs[step * b_skip..]; + let dst_p = &mut dst[step * c_skip..]; + unsafe { + gemm( + /* m: usize = */ m, + /* n: usize = */ n, + /* k: usize = */ k, + /* dst: *mut T = */ dst_p.as_mut_ptr(), + /* dst_cs: isize = */ dst_cs as isize, + /* dst_rs: isize = */ dst_rs as isize, + /* read_dst: bool = */ false, + /* lhs: *const T = */ lhs_p.as_ptr(), + /* lhs_cs: isize = */ lhs_cs as isize, + /* lhs_rs: isize = */ lhs_rs as isize, + /* rhs: *const T = */ rhs_p.as_ptr(), + /* rhs_cs: isize = */ rhs_cs as isize, + /* rhs_rs: isize = */ rhs_rs as isize, + /* alpha: T = */ T::zero(), + /* beta: T = */ T::one(), + /* conj_dst: bool = */ false, + /* conj_lhs: bool = */ false, + /* conj_rhs: bool = */ false, + parallelism, + ) + } + } + Ok(dst) + } +} + +struct WhereCondition<'a>(&'a [u32], &'a Layout); + +impl<'a> BinaryMappable for WhereCondition<'a> { + const OP: &'static str = "where_condition"; + + fn f(&self, v1: &[T], l1: &Layout, v2: &[T], l2: &Layout) -> Result> { + let values = match ( + self.1.contiguous_offsets(), + l1.contiguous_offsets(), + l2.contiguous_offsets(), + ) { + (Some((o1, o2)), Some((o_t1, o_t2)), Some((o_f1, o_f2))) => { + let pred = &self.0[o1..o2]; + let v1 = &v1[o_t1..o_t2]; + let v2 = &v2[o_f1..o_f2]; + pred.iter() + .zip(v1.iter().zip(v2.iter())) + .map(|(&p, (&t, &f))| if p > 0 { t } else { f }) + .collect::>() + } + _ => self + .1 + .strided_index() + .zip(l1.strided_index().zip(l2.strided_index())) + .map(|(i_p, (v1_index, v2_index))| { + if self.0[i_p] > 0 { + v1[v1_index] + } else { + v2[v2_index] + } + }) + .collect::>(), + }; + Ok(values) + } +} + +fn copy_strided_source_( + source: &[T], + destination: &mut [T], + destination_offset: usize, + source_layout: &Layout, +) { + match source_layout.contiguous_offsets() { + Some((o_destination1, o_destination2)) => { + let elem_to_copy = + (destination.len() - destination_offset).min(o_destination2 - o_destination1); + destination[destination_offset..destination_offset + elem_to_copy] + .copy_from_slice(&source[o_destination1..o_destination2]) + } + None => { + for (destination_index, source_index) in source_layout.strided_index().enumerate() { + let destination_index = destination_index + destination_offset; + if destination_index >= destination.len() { + break; + } + destination[destination_index] = source[source_index] + } + } + } +} diff --git a/src/backend/mod.rs b/src/backend/mod.rs index c885a22..7b3749c 100644 --- a/src/backend/mod.rs +++ b/src/backend/mod.rs @@ -1 +1,6 @@ +pub(crate) mod backend; pub(crate) mod cpu_backend; +pub(crate) mod mps_backend; + +pub use cpu_backend::CPUStorage; +pub use mps_backend::MPSStorage; diff --git a/src/backend/mps_backend.rs b/src/backend/mps_backend.rs new file mode 100644 index 0000000..82c7652 --- /dev/null +++ b/src/backend/mps_backend.rs @@ -0,0 +1,145 @@ +use crate::device::DeviceLocation; +use crate::operation::{BinaryOperation, UnaryOperation}; +use crate::CPUStorage; +use crate::{DType, Layout, Result, Shape}; + +#[derive(Debug, Clone)] +pub enum MPSStorage { + F32(Vec), + F64(Vec), +} + +#[derive(Debug, Clone)] +pub struct MPSDevice; + +impl crate::backend::backend::BackendStorage for MPSStorage { + type Device = MPSDevice; + + fn device(&self) -> &Self::Device { + todo!() + } + + fn dtype(&self) -> DType { + match self { + MPSStorage::F32(_) => DType::F32, + MPSStorage::F64(_) => DType::F64, + } + } + + fn to_dtype(&self, layout: &Layout, dtype: DType) -> Result { + todo!() + } + + fn to_cpu(&self) -> Result { + todo!() + } + + fn try_clone(&self, layout: &Layout) -> Result { + todo!() + } + + fn copy_strided_source( + &self, + destination: &mut Self, + destination_offset: usize, + destination_layout: &Layout, + ) -> Result<()> { + todo!() + } + + fn binary_operation( + &self, + rhs: &Self, + lhs_layout: &Layout, + rhs_layout: &Layout, + ) -> Result { + todo!() + } + + fn unary_operation(&self, layout: &Layout) -> Result { + todo!() + } + + fn affine(&self, layout: &Layout, add: f64, mul: f64) -> Result { + todo!() + } + + fn matmul( + &self, + _: &Self, + _: (usize, usize, usize, usize), + _: &Layout, + _: &Layout, + ) -> Result { + todo!() + } + + fn where_condition( + &self, + _: &Layout, + _: &Self, + _: &Layout, + _: &Self, + _: &Layout, + ) -> Result { + todo!() + } + + fn embedding(&self, _: &Layout, _: &Self, _: &Layout) -> Result { + todo!() + } + + fn sum(&self, _: &Layout, _: &[usize]) -> Result { + todo!() + } +} + +impl crate::backend::backend::BackendDevice for MPSDevice { + type Storage = MPSStorage; + + fn new(_: usize) -> Result { + todo!("MPSDevice::new") + } + + fn location(&self) -> DeviceLocation { + todo!("MPSDevice::location") + } + + fn same_device(&self, rhs: &Self) -> bool { + rhs.location() == self.location() + } + + fn from_cpu(&self, storage: &CPUStorage) -> Result { + todo!("MPSDevice::from_cpu") + } + + fn zeros_impl(&self, shape: &Shape, dtype: DType) -> Result { + todo!("MPSDevice::zeros_impl") + } + + fn ones_impl(&self, shape: &Shape, dtype: DType) -> Result { + todo!("MPSDevice::ones_impl") + } + + fn rand_uniform( + &self, + shape: &Shape, + dtype: DType, + low: f64, + high: f64, + ) -> Result { + todo!("MPSDevice::rand_uniform") + } + + fn rand_normal( + &self, + shape: &Shape, + dtype: DType, + mean: f64, + std: f64, + ) -> Result { + todo!("MPSDevice::rand_normal") + } +} + +impl MPSStorage {} diff --git a/src/backprop.rs b/src/backprop.rs index 3227e75..32a589b 100644 --- a/src/backprop.rs +++ b/src/backprop.rs @@ -1,5 +1,7 @@ +use crate::gradient_store::GradientStore; +use crate::operation::Operation; use crate::tensor::{Tensor, TensorID}; -use crate::Operation; +use crate::Error; use crate::Result; use std::collections::HashMap; @@ -19,7 +21,7 @@ impl Tensor { } let mut tracked = false; - let mut nodes = if node.variable() { + let mut nodes = if node.is_variable() { tracked = true; nodes } else if let Some(op) = node.op() { @@ -34,7 +36,11 @@ impl Tensor { tracked |= target; nodes } - Operation::Sqr(node) | Operation::Sqrt(node) | Operation::Neg(node) => { + Operation::Sqr(node) + | Operation::Sqrt(node) + | Operation::Neg(node) + | Operation::Broadcast(node) + | Operation::ToDType(node) => { let (target, nodes) = walk(node, nodes, seen); tracked |= target; nodes @@ -74,100 +80,84 @@ impl Tensor { /// ```rust /// use phantom::{Tensor, Device}; /// - /// let x = Tensor::new(&[[2f32, 2.], [1f32, 2.]], Device::CPU)?; - /// let y = Tensor::new(&[[2f32, 2.], [5f32, 6.]], Device::CPU)?; + /// let x = Tensor::new(&[[2f32, 2.], [1f32, 2.]], &Device::CPU)?; + /// let y = Tensor::new(&[[2f32, 2.], [5f32, 6.]], &Device::CPU)?; /// let z = x.add(&y)?; /// let gradients = z.backward()?; /// assert_eq!(gradients.len(), 1); /// /// # Ok::<(), phantom::Error>(()) /// ``` - pub fn backward(&self) -> Result> { + pub fn backward(&self) -> Result { let sorted_nodes = self.sorted_nodes(); - let mut gradients = HashMap::new(); + let mut gradients = GradientStore::new(); - gradients.insert(self.id(), self.ones_like()); + gradients.insert(self, self.ones_like()?.contiguous()?); for node in sorted_nodes.iter() { - if node.variable() { + if node.is_variable() { continue; } - let gradient = gradients.remove(&node.id()).unwrap(); + let gradient = gradients.remove(node).unwrap(); if let Some(op) = node.op() { match op { Operation::Add(lhs, rhs) => { - let lhs_gradient_sum = gradients - .entry(lhs.id()) - .or_insert_with(|| lhs.zeros_like()); + let lhs_gradient_sum = gradients.or_insert(lhs)?; *lhs_gradient_sum = lhs_gradient_sum.add(&gradient)?; - let rhs_gradient_sum = gradients - .entry(rhs.id()) - .or_insert_with(|| rhs.zeros_like()); + let rhs_gradient_sum = gradients.or_insert(rhs)?; *rhs_gradient_sum = rhs_gradient_sum.add(&gradient)?; } Operation::Sub(lhs, rhs) => { - let lhs_gradient_sum = gradients - .entry(lhs.id()) - .or_insert_with(|| lhs.zeros_like()); + let lhs_gradient_sum = gradients.or_insert(lhs)?; *lhs_gradient_sum = lhs_gradient_sum.add(&gradient)?; - let rhs_gradient_sum = gradients - .entry(rhs.id()) - .or_insert_with(|| rhs.zeros_like()); - *rhs_gradient_sum = rhs_gradient_sum.add(&gradient.neg()?)?; + let rhs_gradient_sum = gradients.or_insert(rhs)?; + *rhs_gradient_sum = rhs_gradient_sum.sub(&gradient)?; } Operation::Mul(lhs, rhs) => { let lhs_gradient = gradient.mul(rhs)?; - let lhs_gradient_sum = gradients - .entry(lhs.id()) - .or_insert_with(|| lhs.zeros_like()); + let lhs_gradient_sum = gradients.or_insert(lhs)?; *lhs_gradient_sum = lhs_gradient_sum.add(&lhs_gradient)?; let rhs_gradient = gradient.mul(lhs)?; - let rhs_gradient_sum = gradients - .entry(rhs.id()) - .or_insert_with(|| rhs.zeros_like()); + let rhs_gradient_sum = gradients.or_insert(rhs)?; *rhs_gradient_sum = rhs_gradient_sum.add(&rhs_gradient)?; } Operation::Div(lhs, rhs) => { let lhs_gradient = gradient.div(rhs)?; - let lhs_gradient_sum = gradients - .entry(lhs.id()) - .or_insert_with(|| lhs.zeros_like()); + let lhs_gradient_sum = gradients.or_insert(lhs)?; *lhs_gradient_sum = lhs_gradient_sum.add(&lhs_gradient)?; let rhs_gradient = gradient.mul(lhs)?.div(&rhs.sqr()?)?; - let rhs_gradient_sum = gradients - .entry(rhs.id()) - .or_insert_with(|| rhs.zeros_like()); + let rhs_gradient_sum = gradients.or_insert(rhs)?; *rhs_gradient_sum = rhs_gradient_sum.add(&rhs_gradient)?; } - Operation::Affine { node, mul, .. } => { - let node_gradient = gradient.affine(*mul, 0.)?; - let gradient_sum = gradients - .entry(node.id()) - .or_insert_with(|| node.zeros_like()); - *gradient_sum = gradient_sum.add(&node_gradient)? + Operation::Affine { node: arg, mul, .. } => { + let gradient_arg = gradient.affine(*mul, 0.)?; + let gradient_sum = gradients.or_insert(arg)?; + *gradient_sum = gradient_sum.add(&gradient_arg)? + } + Operation::Sqr(arg) => { + let gradient_arg = arg.mul(&gradient)?.affine(2., 0.)?; + let gradient_sum = gradients.or_insert(node)?; + *gradient_sum = gradient_sum.add(&gradient_arg)? + } + Operation::Sqrt(arg) => { + let gradient_arg = gradient.div(arg)?.affine(0.5, 0.)?; + let gradient_sum = gradients.or_insert(arg)?; + *gradient_sum = gradient_sum.add(&gradient_arg)? } - Operation::Sqr(node) => { - let node_gradient = node.mul(&gradient)?.affine(2., 0.)?; - let gradient_sum = gradients - .entry(node.id()) - .or_insert_with(|| node.zeros_like()); - *gradient_sum = gradient_sum.add(&node_gradient)? + Operation::Neg(arg) => { + let gradient_sum = gradients.or_insert(arg)?; + *gradient_sum = gradient_sum.sub(&gradient)? } - Operation::Sqrt(node) => { - let node_gradient = gradient.div(node)?.affine(0.5, 0.)?; - let gradient_sum = gradients - .entry(node.id()) - .or_insert_with(|| node.zeros_like()); - *gradient_sum = gradient_sum.add(&node_gradient)? + Operation::Broadcast(_) => { + return Err(Error::BackwardUnsupported { + operation: "broadcast", + }) } - Operation::Neg(node) => { - let node_gradient = gradient.neg()?; - let gradient_sum = gradients - .entry(node.id()) - .or_insert_with(|| node.zeros_like()); - *gradient_sum = gradient_sum.add(&node_gradient)? + Operation::ToDType(arg) => { + let gradient_sum = gradients.or_insert(arg)?; + *gradient_sum = gradient_sum.add(&gradient.to_dtype(node.dtype())?)? } } } diff --git a/src/device.rs b/src/device.rs index c8458a0..713f6cf 100644 --- a/src/device.rs +++ b/src/device.rs @@ -1,41 +1,86 @@ -use crate::backend::cpu_backend::CPUStorage; +use crate::backend::backend::BackendDevice; +use crate::backend::cpu_backend::{CPUDevice, CPUStorage}; +use crate::WithDType; use crate::{storage::Storage, DType, Result, Shape}; +/// A device location is an actual physical device while a device is a logical +/// device. For example a GPU device means the tensor is loaded on a GPU while a +/// GPU location refers to the specific GPU that the tensor is loaded on. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub enum DeviceLocation { + CPU, + MPS, +} + #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] pub enum Device { CPU, + MPS, } impl Device { - pub fn zeros(&self, shape: &Shape, dtype: DType) -> Storage { - let elem_count: usize = shape.elem_count(); + pub fn same_id(&self, rhs: &Self) -> bool { + match (self, rhs) { + (Self::CPU, Self::CPU) => true, + (Self::MPS, Self::MPS) => true, + // For enums that carry values then call .same_device to compare + _ => false, + } + } + + pub fn same_device(&self, rhs: &Self) -> bool { + match (self, rhs) { + (Self::CPU, Self::CPU) => true, + (Self::MPS, Self::MPS) => true, + _ => false, + } + } + + pub fn location(&self) -> DeviceLocation { + match self { + Self::CPU => DeviceLocation::CPU, + Self::MPS => DeviceLocation::MPS, + } + } + + pub fn zeros(&self, shape: &Shape, dtype: DType) -> Result { match self { Device::CPU => { - let storage = match dtype { - DType::F32 => CPUStorage::F32(vec![0f32; elem_count]), - DType::F64 => CPUStorage::F64(vec![0f64; elem_count]), - }; - Storage::CPU(storage) + let storage = CPUDevice.zeros_impl(shape, dtype)?; + Ok(Storage::CPU(storage)) } + Device::MPS => todo!(), } } - pub fn ones(&self, shape: &Shape, dtype: DType) -> Storage { - let elem_count: usize = shape.elem_count(); + pub fn ones(&self, shape: &Shape, dtype: DType) -> Result { match self { Device::CPU => { - let storage = match dtype { - DType::F32 => CPUStorage::F32(vec![1f32; elem_count]), - DType::F64 => CPUStorage::F64(vec![1f64; elem_count]), - }; - Storage::CPU(storage) + let storage = CPUDevice.ones_impl(shape, dtype)?; + Ok(Storage::CPU(storage)) } + Device::MPS => todo!(), } } pub fn tensor(&self, data: A) -> Storage { match self { Device::CPU => Storage::CPU(data.to_cpu()), + Device::MPS => todo!(), + } + } + + pub fn storage(&self, array: A) -> Result { + match self { + Device::CPU => Ok(Storage::CPU(array.to_cpu())), + Device::MPS => todo!(), + } + } + + pub fn storage_owned(&self, data: Vec) -> Result { + match self { + Device::CPU => Ok(Storage::CPU(S::to_cpu_owned(data))), + Device::MPS => todo!(), } } } diff --git a/src/dtype.rs b/src/dtype.rs index a8ac79b..28a3f4a 100644 --- a/src/dtype.rs +++ b/src/dtype.rs @@ -1,32 +1,29 @@ +use crate::backend::backend::BackendStorage; use crate::{CPUStorage, Error, Result}; #[derive(Debug, Copy, Clone, PartialEq, Eq, Hash)] pub enum DType { F32, F64, + U32, } -impl DType { - pub fn size(&self) -> usize { - match self { - DType::F32 => 4, - DType::F64 => 8, - } - } -} - -pub trait WithDType: Sized + Copy { +pub trait WithDType: Sized + Copy + 'static + num_traits::NumAssign { const DTYPE: DType; fn to_cpu_owned(data: Vec) -> CPUStorage; fn to_cpu(data: &[Self]) -> CPUStorage { Self::to_cpu_owned(data.to_vec()) } - fn storage_slice(storage: &CPUStorage) -> Result<&[Self]>; + fn cpu_storage_slice(storage: &CPUStorage) -> Result<&[Self]>; + fn cpu_storage_data(storage: CPUStorage) -> Result>; + + fn to_f64(self) -> f64; + fn from_f64(value: f64) -> Self; } macro_rules! with_dtype { - ($type: ty, $dtype:ident) => { + ($type: ty, $dtype:ident, $from_f64: expr, $to_f64: expr) => { impl WithDType for $type { const DTYPE: DType = DType::$dtype; @@ -34,7 +31,17 @@ macro_rules! with_dtype { CPUStorage::$dtype(data) } - fn storage_slice(storage: &CPUStorage) -> Result<&[Self]> { + fn cpu_storage_slice(storage: &CPUStorage) -> Result<&[Self]> { + match storage { + CPUStorage::$dtype(data) => Ok(data), + _ => Err(Error::UnexpectedDType { + expected: DType::$dtype, + actual: storage.dtype(), + }), + } + } + + fn cpu_storage_data(storage: CPUStorage) -> Result> { match storage { CPUStorage::$dtype(data) => Ok(data), _ => Err(Error::UnexpectedDType { @@ -43,9 +50,18 @@ macro_rules! with_dtype { }), } } + + fn to_f64(self) -> f64 { + $to_f64(self) + } + + fn from_f64(value: f64) -> Self { + $from_f64(value) + } } }; } -with_dtype!(f32, F32); -with_dtype!(f64, F64); +with_dtype!(f32, F32, |v: f64| v as f32, |v: f32| v as f64); +with_dtype!(f64, F64, |v: f64| v, |v: f64| v); +with_dtype!(u32, U32, |v: f64| v as u32, |v: u32| v as f64); diff --git a/src/error.rs b/src/error.rs index 67aa9f0..0fbb01b 100644 --- a/src/error.rs +++ b/src/error.rs @@ -1,4 +1,5 @@ -use crate::{DType, Device, Shape}; +use crate::device::DeviceLocation; +use crate::{DType, Layout, Shape}; #[derive(thiserror::Error, Debug)] pub enum Error { @@ -14,8 +15,8 @@ pub enum Error { #[error("unexpected device in {op}, lhs: {lhs:?}, rhs: {rhs:?}")] BinaryOperationDeviceMismatch { - lhs: Device, - rhs: Device, + lhs: DeviceLocation, + rhs: DeviceLocation, op: &'static str, }, @@ -32,6 +33,36 @@ pub enum Error { rhs: Shape, op: &'static str, }, + + #[error("cannot broadcast {source_shape:?} to {destination_shape:?}")] + BroadcastIncompatibleShapes { + source_shape: Shape, + destination_shape: Shape, + }, + + #[error("Shape mismatch, got a buffer of size {buffer_size} which is incompatible with the shape {shape:?}")] + ShapeMismatch { buffer_size: usize, shape: Shape }, + + #[error("backward is not supported for {operation}")] + BackwardUnsupported { operation: &'static str }, + + #[error("unexpected stride in matmul, lhs: {lhs_layout:?}, rhs: {rhs_layout:?}, bmnk: {bmnk:?}, {msg}")] + MatMulUnexpectedStride { + lhs_layout: Layout, + rhs_layout: Layout, + bmnk: (usize, usize, usize, usize), + msg: &'static str, + }, + + #[error("{op} invalid index {index} with vocab {vocab_size}")] + InvalidIndex { + op: &'static str, + index: usize, + vocab_size: usize, + }, + + #[error("unsupported dtype {dtype:?} for {op}")] + UnsupportedDTypeForOperation { dtype: DType, op: &'static str }, } pub type Result = std::result::Result; diff --git a/src/gradient_store.rs b/src/gradient_store.rs new file mode 100644 index 0000000..87d9c4a --- /dev/null +++ b/src/gradient_store.rs @@ -0,0 +1,47 @@ +use crate::tensor::{Tensor, TensorID}; +use crate::Result; +use std::collections::HashMap; + +pub struct GradientStore(HashMap); + +impl GradientStore { + pub fn new() -> Self { + Self(HashMap::new()) + } + + pub fn get(&self, tensor: &Tensor) -> Option<&Tensor> { + self.0.get(&tensor.id()) + } + + pub fn get_id(&self, id: TensorID) -> Option<&Tensor> { + self.0.get(&id) + } + + pub fn insert(&mut self, tensor: &Tensor, gradient: Tensor) { + self.0.insert(tensor.id(), gradient); + } + + pub fn or_insert(&mut self, tensor: &Tensor) -> Result<&mut Tensor> { + use std::collections::hash_map::Entry; + let grad = match self.0.entry(tensor.id()) { + Entry::Occupied(entry) => entry.into_mut(), + Entry::Vacant(entry) => { + let grad = tensor.zeros_like()?; + entry.insert(grad) + } + }; + Ok(grad) + } + + pub fn remove(&mut self, tensor: &Tensor) -> Option { + self.0.remove(&tensor.id()) + } + + pub fn remove_id(&mut self, id: TensorID) -> Option { + self.0.remove(&id) + } + + pub fn len(&self) -> usize { + self.0.len() + } +} diff --git a/src/index.rs b/src/index.rs index 9665145..fdad632 100644 --- a/src/index.rs +++ b/src/index.rs @@ -1,3 +1,5 @@ +use crate::layout::Layout; + /// A strided index acts as an iterator for the elements of an N-dimensional array stored in a /// flat buffer using some potential strides. The iterator yields the offset position of each /// element in the buffer. @@ -5,24 +7,23 @@ pub struct StridedIndex<'a> { next_index: Option, multi_index: Vec, - dims: &'a [usize], - stride: &'a [usize], + layout: &'a Layout, } impl<'a> StridedIndex<'a> { - pub(crate) fn new(dims: &'a [usize], stride: &'a [usize]) -> Self { + pub(crate) fn new(layout: &'a Layout) -> Self { + let dims = layout.dims(); let elem_count: usize = dims.iter().product(); let next_index = if elem_count == 0 { None } else { // This applies to the scalar case. - Some(0) + Some(layout.start_offset()) }; StridedIndex { next_index, multi_index: vec![0; dims.len()], - dims, - stride, + layout, } } } @@ -36,7 +37,13 @@ impl<'a> Iterator for StridedIndex<'a> { Some(storage_index) => storage_index, }; let mut updated = false; - for (multi_i, max_i) in self.multi_index.iter_mut().zip(self.dims.iter()).rev() { + + for (multi_i, max_i) in self + .multi_index + .iter_mut() + .zip(self.layout.dims().iter()) + .rev() + { let next_i = *multi_i + 1; if next_i < *max_i { *multi_i = next_i; @@ -50,9 +57,10 @@ impl<'a> Iterator for StridedIndex<'a> { let next_storage_index = self .multi_index .iter() - .zip(self.stride.iter()) + .zip(self.layout.stride().iter()) .map(|(&x, &y)| x * y) - .sum(); + .sum::() + + self.layout.start_offset(); Some(next_storage_index) } else { None diff --git a/src/layout.rs b/src/layout.rs new file mode 100644 index 0000000..debbdc7 --- /dev/null +++ b/src/layout.rs @@ -0,0 +1,138 @@ +use crate::{Error, Result, Shape, StridedIndex}; + +#[derive(Clone, PartialEq, Eq, Debug)] +pub struct Layout { + shape: Shape, + /// Element-wise stride rather than byte-wise stride + stride: Vec, + start_offset: usize, +} + +impl Layout { + pub fn contiguous>(shape: S) -> Self { + Self::contiguous_with_offset(shape, 0) + } + + pub fn contiguous_with_offset>(shape: S, start_offset: usize) -> Self { + let shape = shape.into(); + let stride = shape.stride_contiguous(); + Self { + shape, + stride, + start_offset, + } + } + + /// Returns the appropriate start and stop offset if the data is stored in a C + /// contiguous (row major) way. + pub fn contiguous_offsets(&self) -> Option<(usize, usize)> { + if self.is_contiguous() { + let start_o = self.start_offset; + Some((start_o, start_o + self.shape.elem_count())) + } else { + None + } + } + + pub fn dims(&self) -> &[usize] { + self.shape.dims() + } + + pub fn shape(&self) -> &Shape { + &self.shape + } + + pub fn stride(&self) -> &[usize] { + &self.stride + } + + pub fn start_offset(&self) -> usize { + self.start_offset + } + + /// Returns true if the data is stored in a contiguous block of memory + /// (i.e. no padding between dimensions) in a Row-Major order. + pub fn is_contiguous(&self) -> bool { + self.shape.is_contiguous(&self.stride) + } + + pub(crate) fn narrow(&self, dim: usize, start: usize, length: usize) -> Result { + let dims = self.shape().dims(); + if dim >= dims.len() { + Err(Error::UnexpectedRank { + expected: dims.len(), + actual: dim, + shape: self.shape().clone(), + })? + } + + if start + length > dims[dim] { + todo!("add a proper error: out of bounds for narrow {dim} {start} {length} {dims:?}") + } + + let mut dims = dims.to_vec(); + dims[dim] = length; + Ok(Self { + shape: Shape::from(dims), + stride: self.stride.clone(), + start_offset: self.start_offset + self.stride[dim] * start, + }) + } + + pub(crate) fn transpose(&self, dim1: usize, dim2: usize) -> Result { + let rank = self.shape.rank(); + if rank <= dim1 || rank <= dim2 { + return Err(Error::UnexpectedRank { + expected: usize::max(dim1, dim2), + actual: rank, + shape: self.shape().clone(), + }); + } + let mut stride = self.stride().to_vec(); + let mut dims = self.shape().dims().to_vec(); + dims.swap(dim1, dim2); + stride.swap(dim1, dim2); + Ok(Self { + shape: Shape::from(dims), + stride, + start_offset: self.start_offset, + }) + } + + pub(crate) fn strided_index(&self) -> crate::StridedIndex { + StridedIndex::new(self) + } + + pub fn broadcast_as>(&self, shape: S) -> Result { + let shape = shape.into(); + if shape.rank() < self.shape().rank() { + Err(Error::BroadcastIncompatibleShapes { + source_shape: self.shape().clone(), + destination_shape: shape.clone(), + })? + } + let added_dims = shape.rank() - self.shape().rank(); + let mut stride = vec![0; added_dims]; + for (&destination_dim, (&source_dim, &source_stride)) in shape.dims()[added_dims..] + .iter() + .zip(self.dims().iter().zip(self.stride())) + { + let s = if destination_dim == source_dim { + source_stride + } else if source_dim != 1 { + return Err(Error::BroadcastIncompatibleShapes { + source_shape: self.shape().clone(), + destination_shape: shape, + }); + } else { + 0 + }; + stride.push(s) + } + Ok(Self { + shape, + stride, + start_offset: self.start_offset, + }) + } +} diff --git a/src/lib.rs b/src/lib.rs index 8a5cd25..29b3a3d 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -3,18 +3,21 @@ mod backprop; mod device; mod dtype; mod error; +mod gradient_store; mod index; +mod layout; mod operation; mod shape; mod storage; mod tensor; +mod utils; pub use backend::cpu_backend::CPUStorage; pub use device::Device; pub use dtype::{DType, WithDType}; pub use error::{Error, Result}; pub use index::StridedIndex; -pub use operation::Operation; +pub use layout::Layout; pub use shape::Shape; pub use storage::Storage; pub use tensor::{Tensor, TensorID}; diff --git a/src/operation.rs b/src/operation.rs index 66c9d76..4481c6b 100644 --- a/src/operation.rs +++ b/src/operation.rs @@ -1,14 +1,100 @@ use crate::Tensor; -pub enum Operation { +#[derive(Debug, Clone)] +pub(crate) enum Operation { + // Binary Operations Add(Tensor, Tensor), Sub(Tensor, Tensor), Mul(Tensor, Tensor), Div(Tensor, Tensor), + // Unary Operations Sqr(Tensor), Sqrt(Tensor), Neg(Tensor), + // Casting and complex ops Affine { node: Tensor, mul: f64, add: f64 }, + Broadcast(Tensor), + + ToDType(Tensor), +} + +pub(crate) trait UnaryOperation { + const NAME: &'static str; + const KERNEL: &'static str; + const V: Self; + fn f32(v1: f32) -> f32; + fn f64(v1: f64) -> f64; + fn u32(v1: u32) -> u32; +} + +pub(crate) trait BinaryOperation { + const NAME: &'static str; + const KERNEL: &'static str; + const V: Self; + fn f32(v1: f32, v2: f32) -> f32; + fn f64(v1: f64, v2: f64) -> f64; + fn u32(v1: u32, v2: u32) -> u32; +} + +pub(crate) struct Add; +pub(crate) struct Sub; +pub(crate) struct Mul; +pub(crate) struct Div; +pub(crate) struct Sqr; +pub(crate) struct Sqrt; +pub(crate) struct Neg; + +macro_rules! bin_op { + ($op:ident, $name: literal, $e: expr) => { + impl BinaryOperation for $op { + const NAME: &'static str = $name; + const KERNEL: &'static str = concat!("b", $name); + const V: Self = $op; + + fn f32(v1: f32, v2: f32) -> f32 { + $e(v1, v2) + } + + fn f64(v1: f64, v2: f64) -> f64 { + $e(v1, v2) + } + + fn u32(v1: u32, v2: u32) -> u32 { + $e(v1, v2) + } + } + }; } + +bin_op!(Add, "add", |v1, v2| v1 + v2); +bin_op!(Sub, "sub", |v1, v2| v1 - v2); +bin_op!(Mul, "mul", |v1, v2| v1 * v2); +bin_op!(Div, "div", |v1, v2| v1 / v2); + +macro_rules! unary_op { + ($op: ident, $name: literal, $a: ident, $e: expr) => { + impl UnaryOperation for $op { + const NAME: &'static str = $name; + const KERNEL: &'static str = concat!("u", $name); + const V: Self = $op; + + fn f32($a: f32) -> f32 { + $e + } + + fn f64($a: f64) -> f64 { + $e + } + + fn u32($a: u32) -> u32 { + todo!("no unary function for u32") + } + } + }; +} + +unary_op!(Sqr, "sqr", a, a * a); +unary_op!(Sqrt, "sqrt", a, a.sqrt()); +unary_op!(Neg, "neg", a, -a); diff --git a/src/shape.rs b/src/shape.rs index 5f4e0f0..3cb86eb 100644 --- a/src/shape.rs +++ b/src/shape.rs @@ -1,7 +1,9 @@ use crate::{Error, Result}; #[derive(Clone, PartialEq, Eq)] -pub struct Shape(pub(crate) Vec); +pub struct Shape(Vec); + +pub const SCALAR: Shape = Shape(vec![]); impl From<()> for Shape { fn from(_: ()) -> Self { @@ -21,6 +23,12 @@ impl From<(usize, usize)> for Shape { } } +impl From<(usize, usize, usize)> for Shape { + fn from((dim_0, dim_1, dim_2): (usize, usize, usize)) -> Self { + Self(vec![dim_0, dim_1, dim_2]) + } +} + impl From<&[usize; 1]> for Shape { fn from(dims: &[usize; 1]) -> Self { Self(dims.to_vec()) @@ -33,6 +41,12 @@ impl From<&[usize; 2]> for Shape { } } +impl From<&[usize; 3]> for Shape { + fn from(dims: &[usize; 3]) -> Self { + Self(dims.to_vec()) + } +} + impl From<&[usize]> for Shape { fn from(dims: &[usize]) -> Self { Self(dims.to_vec()) @@ -45,6 +59,12 @@ impl From<&Shape> for Shape { } } +impl From> for Shape { + fn from(dims: Vec) -> Self { + Self(dims) + } +} + macro_rules! get_rank { ($fn_name:ident, $cnt:tt, $dims:expr, $out_type:ty) => { pub fn $fn_name(&self) -> Result<$out_type> { @@ -74,6 +94,10 @@ impl Shape { &self.0 } + pub fn into_dims(self) -> Vec { + self.0 + } + pub fn elem_count(&self) -> usize { self.0.iter().product() } @@ -86,6 +110,12 @@ impl Shape { |dims: &[usize]| (dims[0], dims[1]), (usize, usize) ); + get_rank!( + rank_three, + 3, + |dims: &[usize]| (dims[0], dims[1], dims[2]), + (usize, usize, usize) + ); /// Stride over a contiguous n-dimensional array of this shape pub(crate) fn stride_contiguous(&self) -> Vec { @@ -103,6 +133,28 @@ impl Shape { stride.reverse(); stride } + + /// Returns true if the shape is contiguous with the given stride + /// (i.e. no padding between dimensions) in a Row-Major order. + pub fn is_contiguous(&self, stride: &[usize]) -> bool { + if self.0.len() != stride.len() { + return false; + } + + let mut accumulator = 1; + for (&stride, &dim) in stride.iter().zip(self.0.iter()).rev() { + if stride != accumulator { + return false; + } + accumulator *= dim; + } + true + } + + pub fn extend(mut self, additional_dims: &[usize]) -> Self { + self.0.extend(additional_dims); + self + } } impl std::fmt::Debug for Shape { @@ -110,3 +162,20 @@ impl std::fmt::Debug for Shape { write!(f, "{:?}", &self.dims()) } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn stride() { + let shape = Shape::from(()); + assert_eq!(shape.stride_contiguous(), Vec::::new()); + let shape = Shape::from(42); + assert_eq!(shape.stride_contiguous(), [1]); + let shape = Shape::from((42, 1337)); + assert_eq!(shape.stride_contiguous(), [1337, 1]); + let shape = Shape::from((299, 792, 458)); + assert_eq!(shape.stride_contiguous(), [458 * 792, 458, 1]); + } +} diff --git a/src/storage.rs b/src/storage.rs index cde6370..3c6bfff 100644 --- a/src/storage.rs +++ b/src/storage.rs @@ -1,116 +1,24 @@ -use crate::backend::cpu_backend::CPUStorage; -use crate::{DType, Device, Error, Result, Shape}; +use crate::backend::{backend::BackendStorage, CPUStorage, MPSStorage}; +use crate::operation::{BinaryOperation, UnaryOperation}; +use crate::{DType, Device, Error, Layout, Result}; pub enum Storage { CPU(CPUStorage), -} - -pub(crate) trait UnaryOperation { - const NAME: &'static str; - fn f32(value: f32) -> f32; - fn f64(value: f64) -> f64; -} - -pub(crate) trait BinaryOperation { - const NAME: &'static str; - fn f32(lhs: f32, rhs: f32) -> f32; - fn f64(lhs: f64, rhs: f64) -> f64; -} - -struct Add; - -impl BinaryOperation for Add { - const NAME: &'static str = "add"; - fn f32(lhs: f32, rhs: f32) -> f32 { - lhs + rhs - } - fn f64(lhs: f64, rhs: f64) -> f64 { - lhs + rhs - } -} - -struct Sub; - -impl BinaryOperation for Sub { - const NAME: &'static str = "sub"; - fn f32(lhs: f32, rhs: f32) -> f32 { - lhs - rhs - } - fn f64(lhs: f64, rhs: f64) -> f64 { - lhs - rhs - } -} - -struct Mul; - -impl BinaryOperation for Mul { - const NAME: &'static str = "mul"; - fn f32(lhs: f32, rhs: f32) -> f32 { - lhs * rhs - } - fn f64(lhs: f64, rhs: f64) -> f64 { - lhs * rhs - } -} - -struct Div; - -impl BinaryOperation for Div { - const NAME: &'static str = "div"; - fn f32(lhs: f32, rhs: f32) -> f32 { - lhs / rhs - } - fn f64(lhs: f64, rhs: f64) -> f64 { - lhs / rhs - } -} - -struct Sqr; - -impl UnaryOperation for Sqr { - const NAME: &'static str = "sqr"; - fn f32(value: f32) -> f32 { - value * value - } - fn f64(value: f64) -> f64 { - value * value - } -} - -struct Sqrt; - -impl UnaryOperation for Sqrt { - const NAME: &'static str = "sqrt"; - fn f32(value: f32) -> f32 { - value.sqrt() - } - fn f64(value: f64) -> f64 { - value.sqrt() - } -} - -struct Neg; - -impl UnaryOperation for Neg { - const NAME: &'static str = "neg"; - fn f32(value: f32) -> f32 { - -value - } - fn f64(value: f64) -> f64 { - -value - } + MPS(MPSStorage), } impl Storage { pub fn device(&self) -> Device { match self { Storage::CPU { .. } => Device::CPU, + Storage::MPS { .. } => Device::MPS, } } pub fn dtype(&self) -> DType { match self { Storage::CPU(storage) => storage.dtype(), + Storage::MPS(storage) => storage.dtype(), } } @@ -119,7 +27,11 @@ impl Storage { let rhs = rhs.device(); if lhs != rhs { - Err(Error::BinaryOperationDeviceMismatch { lhs, rhs, op }) + Err(Error::BinaryOperationDeviceMismatch { + lhs: lhs.location(), + rhs: rhs.location(), + op, + }) } else { Ok(()) } @@ -136,21 +48,24 @@ impl Storage { } } - fn unary_operation(&self, shape: &Shape, stride: &[usize]) -> Result { + pub(crate) fn unary_operation(&self, layout: &Layout) -> Result { match self { Storage::CPU(storage) => { - let storage = storage.unary_impl::(shape, stride)?; + let storage = storage.unary_operation::(layout)?; Ok(Self::CPU(storage)) } + Storage::MPS(storage) => { + let storage = storage.unary_operation::(layout)?; + Ok(Self::MPS(storage)) + } } } - fn binary_operation( + pub(crate) fn binary_operation( &self, rhs: &Self, - shape: &Shape, - lhs_stride: &[usize], - rhs_stride: &[usize], + lhs_layout: &Layout, + rhs_layout: &Layout, ) -> Result { // Check the operands are valid for this operation. self.matches_device(rhs, T::NAME)?; @@ -159,76 +74,66 @@ impl Storage { // This will need contiguous layout optimizations later match (self, rhs) { (Storage::CPU(lhs), Storage::CPU(rhs)) => { - let storage = lhs.binary_operation::(rhs, shape, lhs_stride, rhs_stride)?; + let storage = lhs.binary_operation::(rhs, lhs_layout, rhs_layout)?; Ok(Self::CPU(storage)) } + (Storage::MPS(lhs), Storage::MPS(rhs)) => { + let storage = lhs.binary_operation::(rhs, lhs_layout, rhs_layout)?; + Ok(Self::MPS(storage)) + } + (_, _) => Err(Error::BinaryOperationDeviceMismatch { + lhs: self.device().location(), + rhs: rhs.device().location(), + op: T::NAME, + }), } } - pub(crate) fn add( - &self, - rhs: &Self, - shape: &Shape, - lhs_stride: &[usize], - rhs_stride: &[usize], - ) -> Result { - self.binary_operation::(rhs, shape, lhs_stride, rhs_stride) - } - - pub(crate) fn sub( - &self, - rhs: &Self, - shape: &Shape, - lhs_stride: &[usize], - rhs_stride: &[usize], - ) -> Result { - self.binary_operation::(rhs, shape, lhs_stride, rhs_stride) - } - - pub(crate) fn mul( - &self, - rhs: &Self, - shape: &Shape, - lhs_stride: &[usize], - rhs_stride: &[usize], - ) -> Result { - self.binary_operation::(rhs, shape, lhs_stride, rhs_stride) - } - - pub(crate) fn div( - &self, - rhs: &Self, - shape: &Shape, - lhs_stride: &[usize], - rhs_stride: &[usize], - ) -> Result { - self.binary_operation::
(rhs, shape, lhs_stride, rhs_stride) - } - - pub(crate) fn affine( - &self, - shape: &Shape, - stride: &[usize], - mul: f64, - add: f64, - ) -> Result { + pub(crate) fn affine(&self, layout: &Layout, mul: f64, add: f64) -> Result { match self { Storage::CPU(storage) => { - let storage = storage.affine(shape, stride, mul, add)?; + let storage = storage.affine(layout, mul, add)?; Ok(Self::CPU(storage)) } + Storage::MPS(storage) => { + let storage = storage.affine(layout, mul, add)?; + Ok(Self::MPS(storage)) + } } } - pub(crate) fn sqr(&self, shape: &Shape, stride: &[usize]) -> Result { - self.unary_operation::(shape, stride) - } - - pub(crate) fn sqrt(&self, shape: &Shape, stride: &[usize]) -> Result { - self.unary_operation::(shape, stride) + pub(crate) fn to_dtype(&self, layout: &Layout, dtype: DType) -> Result { + match self { + Storage::CPU(storage) => { + let storage = storage.to_dtype(layout, dtype)?; + Ok(Self::CPU(storage)) + } + Storage::MPS(storage) => { + let storage = storage.to_dtype(layout, dtype)?; + Ok(Self::MPS(storage)) + } + } } - pub(crate) fn neg(&self, shape: &Shape, stride: &[usize]) -> Result { - self.unary_operation::(shape, stride) + /// The source is stridable and the destination is contiguous. + pub(crate) fn copy_strided_source( + &self, + destination: &mut Self, + destination_offset: usize, + source_layout: &Layout, + ) -> Result<()> { + match (self, destination) { + (Self::CPU(source), Self::CPU(destination)) => { + source.copy_strided_source(destination, destination_offset, source_layout) + } + (Self::MPS(source), Self::MPS(destination)) => { + source.copy_strided_source(destination, destination_offset, source_layout) + } + (lhs, rhs) => Err(Error::BinaryOperationDeviceMismatch { + lhs: lhs.device().location(), + rhs: rhs.device().location(), + op: "copy", + }), + } } } diff --git a/src/tensor.rs b/src/tensor.rs index dc98ade..aa1e5ed 100644 --- a/src/tensor.rs +++ b/src/tensor.rs @@ -3,9 +3,10 @@ use std::sync::Arc; use crate::device::{Device, NDArray}; use crate::index::StridedIndex; -use crate::storage::Storage; +use crate::operation::Operation; +use crate::storage::{self, Storage}; use crate::WithDType; -use crate::{DType, Error, Operation, Result, Shape}; +use crate::{DType, Error, Layout, Result, Shape}; /// Allow each tensor to be uniquely idenified. This makes it cheap to compute if /// a given tensor is a reference to the same underlying data as another tensor. @@ -21,10 +22,8 @@ impl TensorID { pub struct Tensor_ { id: TensorID, - storage: Storage, - shape: Shape, - /// Element-wise stride rather than byte-wise stride - stride: Vec, + storage: Arc, + layout: Layout, op: Option, variable: bool, } @@ -50,204 +49,187 @@ impl std::fmt::Debug for Tensor { } macro_rules! binary_operation { - ($fn_name:ident, $operation_name:ident, $storage_operation:ident) => { + ($fn_name:ident, $operation_name:ident) => { pub fn $fn_name(&self, rhs: &Self) -> Result { let shape = self.binary_operation_shape_matches(rhs, stringify!($fn_name))?; - let storage = self.storage.$storage_operation( - &rhs.storage, - shape, - self.stride(), - rhs.stride(), - )?; - let t = Tensor_ { - id: TensorID::new(), - storage, - shape: shape.clone(), - stride: shape.stride_contiguous(), - op: Some(Operation::$operation_name(self.clone(), rhs.clone())), - variable: false, + let storage = self + .storage + .binary_operation::( + &rhs.storage, + self.layout(), + rhs.layout(), + )?; + let op = if self.track_op() || rhs.track_op() { + Some(Operation::$operation_name(self.clone(), rhs.clone())) + } else { + None }; - Ok(Self(Arc::new(t))) + Ok(from_storage(storage, shape.clone(), op, false)) } }; } macro_rules! unary_operation { - ($fn_name:ident, $operation_name:ident, $storage_operation:ident) => { + ($fn_name:ident, $operation_name:ident) => { pub fn $fn_name(&self) -> Result { let shape = self.shape(); - let storage = self.storage.$storage_operation(shape, self.stride())?; - let t = Tensor_ { - id: TensorID::new(), - storage, - shape: shape.clone(), - stride: shape.stride_contiguous(), - op: Some(Operation::$operation_name(self.clone())), - variable: false, + let storage = self + .storage + .unary_operation::(self.layout())?; + let op = if self.track_op() { + Some(Operation::$operation_name(self.clone())) + } else { + None }; - Ok(Self(Arc::new(t))) + Ok(from_storage(storage, shape.clone(), op, false)) } }; } impl Tensor { - pub(crate) fn new_impl(array: A, device: Device, variable: bool) -> Result { - let shape: Shape = array.shape()?; - let storage: Storage = device.tensor(array); - let stride: Vec = shape.stride_contiguous(); - let id: TensorID = TensorID::new(); - - let t: Tensor_ = Tensor_ { - id, - storage, - shape, - stride, - op: None, - variable, - }; - Ok(Self(Arc::new(t))) + pub(crate) fn new_impl( + array: A, + shape: Shape, + device: &Device, + variable: bool, + ) -> Result { + let n: usize = shape.elem_count(); + let buffer_size: usize = array.shape()?.elem_count(); + if buffer_size != n { + return Err(Error::ShapeMismatch { buffer_size, shape }); + } + let storage = device.storage(array)?; + Ok(from_storage(storage, shape, None, variable)) } /// Creates a new tensor from a slice of data. /// ```rust /// use phantom::{Tensor, Device, Shape}; - /// let tensor = Tensor::new(&[0f32, 1., 2., 3., 4., 5.], Device::CPU)?; + /// let tensor = Tensor::new(&[0f32, 1., 2., 3., 4., 5.], &Device::CPU)?; /// assert_eq!(tensor.shape(), &Shape::from(&[6])); /// # Ok::<(), phantom::Error>(()) /// ``` - pub fn new(array: A, device: Device) -> Result { - Self::new_impl(array, device, false) + pub fn new(array: A, device: &Device) -> Result { + let shape = array.shape()?; + Self::new_impl(array, shape, device, false) } /// Creates a new variable tensor from a slice of data. /// ```rust /// use phantom::{Tensor, Device}; - /// let tensor = Tensor::var(&[0f32, 1., 2., 3., 4., 5.], Device::CPU)?; - /// assert_eq!(tensor.variable(), true); + /// let tensor = Tensor::var(&[0f32, 1., 2., 3., 4., 5.], &Device::CPU)?; + /// assert_eq!(tensor.is_variable(), true); /// # Ok::<(), phantom::Error>(()) /// ``` - pub fn var(array: A, device: Device) -> Result { - Self::new_impl(array, device, true) + pub fn var(array: A, device: &Device) -> Result { + let shape = array.shape()?; + Self::new_impl(array, shape, device, true) } pub(crate) fn zeros_impl>( shape: S, dtype: DType, - device: Device, + device: &Device, variable: bool, - ) -> Self { - let shape = shape.into(); - let storage = device.zeros(&shape, dtype); - let stride = shape.stride_contiguous(); - let id: TensorID = TensorID::new(); - - let t = Tensor_ { - id, - storage, - shape, - stride, - op: None, - variable, - }; - - Tensor(Arc::new(t)) + ) -> Result { + if variable { + let shape = shape.into(); + let storage = device.zeros(&shape, dtype)?; + Ok(from_storage(storage, shape, None, variable)) + } else { + let storage = device.zeros(&crate::shape::SCALAR, dtype)?; + from_storage(storage, crate::shape::SCALAR, None, variable).broadcast_as(shape) + } } /// Creates a new tensor of zeros with the given shape and data type. /// ```rust /// use phantom::{Tensor, Device, Shape}; - /// let tensor = Tensor::zeros(&[2, 2], phantom::DType::F32, Device::CPU); + /// let tensor = Tensor::zeros(&[2, 2], phantom::DType::F32, &Device::CPU)?; /// assert_eq!(tensor.shape(), &Shape::from(&[2, 2])); /// # Ok::<(), phantom::Error>(()) /// ``` - pub fn zeros>(shape: S, dtype: DType, device: Device) -> Self { + pub fn zeros>(shape: S, dtype: DType, device: &Device) -> Result { Self::zeros_impl(shape, dtype, device, false) } /// Creates a new variable tensor of zeros in the same shape and data type as the input tensor. /// ```rust /// use phantom::{Tensor, Device}; - /// let tensor = Tensor::var(&[[0f32, 1.], [2., 3.]], Device::CPU)?; - /// let zeros = tensor.zeros_like(); + /// let tensor = Tensor::var(&[[0f32, 1.], [2., 3.]], &Device::CPU)?; + /// let zeros = tensor.zeros_like()?; /// assert_eq!(zeros.shape(), tensor.shape()); /// # Ok::<(), phantom::Error>(()) /// ``` - pub fn zeros_like(&self) -> Self { - Tensor::zeros(self.shape(), self.dtype(), self.device()) + pub fn zeros_like(&self) -> Result { + Tensor::zeros(self.shape(), self.dtype(), &self.device()) } /// Creates a new variable tensor of zeros with the given shape and data type. /// ```rust /// use phantom::{Tensor, Device, Shape}; - /// let tensor = Tensor::zeros_var(&[2, 2], phantom::DType::F32, Device::CPU); - /// assert_eq!(tensor.variable(), true); + /// let tensor = Tensor::zeros_var(&[2, 2], phantom::DType::F32, &Device::CPU)?; + /// assert_eq!(tensor.is_variable(), true); /// # Ok::<(), phantom::Error>(()) /// ``` - pub fn zeros_var>(shape: S, dtype: DType, device: Device) -> Self { + pub fn zeros_var>(shape: S, dtype: DType, device: &Device) -> Result { Self::zeros_impl(shape, dtype, device, true) } pub fn ones_impl>( shape: S, dtype: DType, - device: Device, + device: &Device, variable: bool, - ) -> Self { - let shape = shape.into(); - let storage = device.ones(&shape, dtype); - let stride = shape.stride_contiguous(); - let id: TensorID = TensorID::new(); - - let t = Tensor_ { - id, - storage, - shape, - stride, - op: None, - variable, - }; - - Tensor(Arc::new(t)) + ) -> Result { + if variable { + let shape = shape.into(); + let storage = device.ones(&shape, dtype)?; + Ok(from_storage(storage, shape, None, variable)) + } else { + let storage = device.ones(&crate::shape::SCALAR, dtype)?; + from_storage(storage, crate::shape::SCALAR, None, variable).broadcast_as(shape) + } } /// Creates a new tensor of ones with the given shape and data type. /// ```rust /// use phantom::{Tensor, Device, Shape}; - /// let tensor = Tensor::ones(&[2, 2], phantom::DType::F32, Device::CPU); + /// let tensor = Tensor::ones(&[2, 2], phantom::DType::F32, &Device::CPU)?; /// assert_eq!(tensor.shape(), &Shape::from(&[2, 2])); /// # Ok::<(), phantom::Error>(()) /// ``` - pub fn ones>(shape: S, dtype: DType, device: Device) -> Self { + pub fn ones>(shape: S, dtype: DType, device: &Device) -> Result { Self::ones_impl(shape, dtype, device, false) } /// Creates a new variable tensor of ones in the same shape and data type as the input tensor. /// ```rust /// use phantom::{Tensor, Device}; - /// let tensor = Tensor::var(&[[0f32, 1.], [2., 3.]], Device::CPU)?; - /// let ones = tensor.ones_like(); + /// let tensor = Tensor::var(&[[0f32, 1.], [2., 3.]], &Device::CPU)?; + /// let ones = tensor.ones_like()?; /// assert_eq!(ones.shape(), tensor.shape()); /// # Ok::<(), phantom::Error>(()) /// ``` - pub fn ones_like(&self) -> Self { - Tensor::ones(self.shape(), self.dtype(), self.device()) + pub fn ones_like(&self) -> Result { + Tensor::ones(self.shape(), self.dtype(), &self.device()) } /// Creates a new variable tensor of ones with the given shape and data type. /// ```rust /// use phantom::{Tensor, Device, Shape}; - /// let tensor = Tensor::ones_var(&[2, 2], phantom::DType::F32, Device::CPU); - /// assert_eq!(tensor.variable(), true); + /// let tensor = Tensor::ones_var(&[2, 2], phantom::DType::F32, &Device::CPU)?; + /// assert_eq!(tensor.is_variable(), true); /// # Ok::<(), phantom::Error>(()) /// ``` - pub fn ones_var>(shape: S, dtype: DType, device: Device) -> Self { + pub fn ones_var>(shape: S, dtype: DType, device: &Device) -> Result { Self::ones_impl(shape, dtype, device, true) } /// Converts the tensor to a scalar if the tensor is rank 0. /// ```rust /// use phantom::{Tensor, Device}; - /// let tensor = Tensor::new(0f32, Device::CPU)?; + /// let tensor = Tensor::new(0f32, &Device::CPU)?; /// assert_eq!(tensor.to_scalar::()?, 0f32); /// # Ok::<(), phantom::Error>(()) /// ``` @@ -256,21 +238,25 @@ impl Tensor { return Err(Error::UnexpectedRank { expected: 0, actual: self.rank(), - shape: self.0.shape.clone(), + shape: self.shape().clone(), }); } - match &self.0.storage { - Storage::CPU(storage) => { - let data = S::storage_slice(storage)?; - Ok(data[0]) - } + + let from_cpu = |cpu_storage: &crate::CPUStorage| { + let data = S::cpu_storage_slice(cpu_storage)?; + Ok::<_, Error>(data[self.layout().start_offset()]) + }; + + match self.storage.as_ref() { + Storage::CPU(storage) => from_cpu(storage), + Storage::MPS(_) => todo!(), } } /// Returns the unique identifier for this tensor. /// ```rust /// use phantom::{Tensor, Device}; - /// let tensor = Tensor::new(&[0f32], Device::CPU)?; + /// let tensor = Tensor::new(&[0f32], &Device::CPU)?; /// assert_eq!(tensor.id(), tensor.id()); /// # Ok::<(), phantom::Error>(()) /// ``` @@ -281,7 +267,7 @@ impl Tensor { /// Returns the data type of this tensor used on the storage backend. /// ```rust /// use phantom::{Tensor, Device}; - /// let tensor = Tensor::new(&[0f32], Device::CPU)?; + /// let tensor = Tensor::new(&[0f32], &Device::CPU)?; /// assert_eq!(tensor.dtype(), phantom::DType::F32); /// # Ok::<(), phantom::Error>(()) /// ``` @@ -292,7 +278,7 @@ impl Tensor { /// Returns the device that this tensor is stored on. /// ```rust /// use phantom::{Tensor, Device}; - /// let tensor = Tensor::new(&[0f32], Device::CPU)?; + /// let tensor = Tensor::new(&[0f32], &Device::CPU)?; /// assert_eq!(tensor.device(), Device::CPU); /// # Ok::<(), phantom::Error>(()) /// ``` @@ -300,59 +286,64 @@ impl Tensor { self.storage.device() } + /// TODO: Docs + pub fn layout(&self) -> &Layout { + &self.layout + } + /// Returns the shape of the tensor. /// ```rust /// use phantom::{Tensor, Device, Shape}; - /// let tensor = Tensor::new(&[[0f32, 1.], [2., 3.]], Device::CPU)?; + /// let tensor = Tensor::new(&[[0f32, 1.], [2., 3.]], &Device::CPU)?; /// assert_eq!(tensor.shape(), &Shape::from(&[2, 2])); /// # Ok::<(), phantom::Error>(()) /// ``` pub fn shape(&self) -> &Shape { - &self.shape + &self.layout.shape() } /// Returns the rank of the tensor. /// ```rust /// use phantom::{Tensor, Device}; - /// let tensor = Tensor::new(&[[0f32, 1.], [2., 3.]], Device::CPU)?; + /// let tensor = Tensor::new(&[[0f32, 1.], [2., 3.]], &Device::CPU)?; /// assert_eq!(tensor.rank(), 2); /// # Ok::<(), phantom::Error>(()) /// ``` pub fn rank(&self) -> usize { - self.shape.rank() + self.layout.shape().rank() } /// Returns the dimension size for each axis of the tensor. /// ```rust /// use phantom::{Tensor, Device}; - /// let tensor = Tensor::new(&[[0f32, 1.], [2., 3.]], Device::CPU)?; + /// let tensor = Tensor::new(&[[0f32, 1.], [2., 3.]], &Device::CPU)?; /// assert_eq!(tensor.dims(), &[2, 2]); /// # Ok::<(), phantom::Error>(()) /// ``` pub fn dims(&self) -> &[usize] { - self.shape.dims() + self.layout.shape().dims() } /// Returns the total number of values in the tensor. /// ```rust /// use phantom::{Tensor, Device}; - /// let tensor = Tensor::new(&[[0f32, 1.], [2., 3.]], Device::CPU)?; + /// let tensor = Tensor::new(&[[0f32, 1.], [2., 3.]], &Device::CPU)?; /// assert_eq!(tensor.elem_count(), 4); /// # Ok::<(), phantom::Error>(()) /// ``` pub fn elem_count(&self) -> usize { - self.shape.elem_count() + self.layout.shape().elem_count() } /// Returns the element-wise stride of the tensor. /// ```rust /// use phantom::{Tensor, Device}; - /// let tensor = Tensor::new(&[[0f32, 1.], [2., 3.]], Device::CPU)?; + /// let tensor = Tensor::new(&[[0f32, 1.], [2., 3.]], &Device::CPU)?; /// assert_eq!(tensor.stride(), &[2, 1]); /// # Ok::<(), phantom::Error>(()) /// ``` pub fn stride(&self) -> &[usize] { - &self.stride + &self.layout.stride() } /// Returns the operation that created this tensor. @@ -360,14 +351,20 @@ impl Tensor { &self.op } + /// Returns true if the computation graph should track this operation or + /// if this is a variable or one of its dependencies is a variable. + pub(crate) fn track_op(&self) -> bool { + self.variable || self.op.is_some() + } + /// Returns true if the tensor is a variable that is tracked during backpropagation. /// ```rust /// use phantom::{Tensor, Device}; - /// let tensor = Tensor::var(&[[0f32, 1.], [2., 3.]], Device::CPU)?; - /// assert!(tensor.variable()); + /// let tensor = Tensor::var(&[[0f32, 1.], [2., 3.]], &Device::CPU)?; + /// assert!(tensor.is_variable()); /// # Ok::<(), phantom::Error>(()) /// ``` - pub fn variable(&self) -> bool { + pub fn is_variable(&self) -> bool { self.variable } @@ -375,7 +372,7 @@ impl Tensor { /// the elements to be iterated over in lexicographic order. /// ```rust /// use phantom::{Tensor, Device}; - /// let a = Tensor::new(&[[0f32], [2.]], Device::CPU)?; + /// let a = Tensor::new(&[[0f32], [2.]], &Device::CPU)?; /// let mut iter = a.strided_index(); /// assert_eq!(iter.next(), Some(0)); /// assert_eq!(iter.next(), Some(1)); @@ -383,20 +380,20 @@ impl Tensor { /// # Ok::<(), phantom::Error>(()) /// ``` pub fn strided_index(&self) -> StridedIndex { - StridedIndex::new(self.dims(), self.stride()) + self.layout.strided_index() } /// Returns true if the tensor is contiguous in memory. /// ```rust /// use phantom::{Tensor, Device}; /// // Contigious example - /// let a = Tensor::new(&[0f32], Device::CPU)?; - /// assert!(a.contiguous()); + /// let a = Tensor::new(&[0f32], &Device::CPU)?; + /// assert!(a.is_contiguous()); /// # Ok::<(), phantom::Error>(()) /// ``` - pub fn contiguous(&self) -> bool { + pub fn is_contiguous(&self) -> bool { let mut accumulated_stride = 1; - for (&dim, &stride) in self.shape.dims().iter().zip(self.stride.iter()).rev() { + for (&dim, &stride) in self.shape().dims().iter().zip(self.stride().iter()).rev() { if stride != accumulated_stride { return false; } @@ -405,10 +402,27 @@ impl Tensor { true } + pub fn contiguous(&self) -> Result { + if self.is_contiguous() { + Ok(self.clone()) + } else { + let shape = self.shape(); + let mut storage = self.device().zeros(shape, self.dtype())?; + self.storage + .copy_strided_source(&mut storage, 0, self.layout())?; + Ok(from_storage( + storage, + shape.clone(), + None, // TODO + false, + )) + } + } + /// Returns the contents of the rank 1 tensor as a vector. /// ```rust /// use phantom::{Tensor, Device}; - /// let a = Tensor::new(&[0f32, 1., 2., 3., 4., 5.], Device::CPU)?; + /// let a = Tensor::new(&[0f32, 1., 2., 3., 4., 5.], &Device::CPU)?; /// assert_eq!(a.to_vector_rank_one::()?, &[0., 1., 2., 3., 4., 5.]); /// # Ok::<(), phantom::Error>(()) /// ``` @@ -420,26 +434,27 @@ impl Tensor { shape: self.shape().clone(), }); } - match &self.storage { + match &self.storage.as_ref() { Storage::CPU(cpu_storage) => { - let data = S::storage_slice(cpu_storage)?; + let data = S::cpu_storage_slice(cpu_storage)?; Ok(self.strided_index().map(|i: usize| data[i]).collect()) } + Storage::MPS(_) => todo!(), } } /// Returns the contents of the rank 2 tensor as a vector of vectors in row-major order. /// ```rust /// use phantom::{Tensor, Device}; - /// let a = Tensor::new(&[[0f32, 1.], [2., 3.], [4., 5.]], Device::CPU)?; + /// let a = Tensor::new(&[[0f32, 1.], [2., 3.], [4., 5.]], &Device::CPU)?; /// assert_eq!(a.to_vector_rank_two::()?, &[[0., 1.], [2., 3.], [4., 5.]]); /// # Ok::<(), phantom::Error>(()) /// ``` pub fn to_vector_rank_two(&self) -> Result>> { let (dim_one, dim_two) = self.shape().rank_two()?; - match &self.storage { + match &self.storage.as_ref() { Storage::CPU(storage) => { - let data = S::storage_slice(storage)?; + let data = S::cpu_storage_slice(storage)?; let mut rows = vec![]; let mut index = self.strided_index(); for _idx_row in 0..dim_one { @@ -449,6 +464,7 @@ impl Tensor { assert!(index.next().is_none()); Ok(rows) } + Storage::MPS(_) => todo!(), } } @@ -456,8 +472,8 @@ impl Tensor { /// and the operation can be performed, returning the shape if successful. /// ```rust /// use phantom::{Tensor, Device}; - /// let a = Tensor::new(&[[0f32, 1.], [2., 3.]], Device::CPU)?; - /// let b = Tensor::new(&[[0f32, 1.], [2., 3.]], Device::CPU)?; + /// let a = Tensor::new(&[[0f32, 1.], [2., 3.]], &Device::CPU)?; + /// let b = Tensor::new(&[[0f32, 1.], [2., 3.]], &Device::CPU)?; /// assert_eq!(a.binary_operation_shape_matches(&b, "add")?, a.shape()); /// # Ok::<(), phantom::Error>(()) /// ``` @@ -488,40 +504,76 @@ impl Tensor { /// NOTE: This operation casts the input values to the appropriate type so some rounding might /// be performed if operating in mixed precision. /// - /// ```rust - /// use phantom::{Tensor, Device}; - /// let a = Tensor::new(&[[0f32, 1.], [2., 3.]], Device::CPU)?; - /// let a = a.affine(4., -2.)?; - /// assert_eq!(a.to_vector_rank_two::()?, &[[-2.0, 2.0], [6.0, 10.0]]); - /// # Ok::<(), phantom::Error>(()) - /// ``` + /// TODO: Add doctests pub fn affine(&self, mul: f64, add: f64) -> Result { - let shape = self.shape(); - let storage = self.storage.affine(self.shape(), self.stride(), mul, add)?; - - let t = Tensor_ { - id: TensorID::new(), - storage, - shape: shape.clone(), - stride: shape.stride_contiguous(), - op: Some(Operation::Affine { + let storage = self.storage.affine(self.layout(), mul, add)?; + let operation = if self.track_op() { + Some(Operation::Affine { node: self.clone(), mul, add, - }), + }) + } else { + None + }; + Ok(from_storage(storage, self.shape(), operation, false)) + } + + pub fn broadcast_as>(&self, shape: S) -> Result { + let operation = if self.track_op() { + Some(Operation::Broadcast(self.clone())) + } else { + None + }; + + let tensor = Tensor_ { + id: TensorID::new(), + storage: self.storage.clone(), + layout: self.layout.broadcast_as(shape)?, + op: operation, variable: false, }; - Ok(Self(Arc::new(t))) + + Ok(Tensor(Arc::new(tensor))) } - binary_operation!(add, Add, add); - binary_operation!(sub, Sub, sub); - binary_operation!(mul, Mul, mul); - binary_operation!(div, Div, div); + // Shorthand for broadcast_as + pub fn expand>(&self, shape: S) -> Result { + self.broadcast_as(shape) + } + + /// Returns a new tensor duplicating data from the original tensor. New dimensions are inserted + /// on the left. + pub fn broadcast_left>(&self, left_shape: S) -> Result { + let left_shape = left_shape.into(); + let mut dims = left_shape.into_dims(); + dims.extend(self.dims()); + self.broadcast_as(dims) + } - unary_operation!(sqr, Sqr, sqr); - unary_operation!(sqrt, Sqrt, sqrt); - unary_operation!(neg, Neg, neg); + pub fn to_dtype(&self, dtype: DType) -> Result { + if self.dtype() == dtype { + Ok(self.clone()) + } else { + let shape = self.shape(); + let storage = self.storage.to_dtype(&self.layout(), dtype)?; + let operation = if self.track_op() { + Some(Operation::ToDType(self.clone())) + } else { + None + }; + Ok(from_storage(storage, shape.clone(), operation, false)) + } + } + + binary_operation!(add, Add); + binary_operation!(sub, Sub); + binary_operation!(mul, Mul); + binary_operation!(div, Div); + + unary_operation!(sqr, Sqr); + unary_operation!(sqrt, Sqrt); + unary_operation!(neg, Neg); } /// Implement binary operations with operator shorthands. @@ -581,3 +633,19 @@ binary_trait!(Add, add, |_| 1., |v| v); binary_trait!(Sub, sub, |_| 1., |v: f64| -v); binary_trait!(Mul, mul, |v| v, |_| 0.); binary_trait!(Div, div, |v| 1. / v, |_| 0.); + +fn from_storage>( + storage: Storage, + shape: S, + operation: Option, + variable: bool, +) -> Tensor { + let tensor = Tensor_ { + id: TensorID::new(), + storage: Arc::new(storage), + layout: Layout::contiguous(shape), + op: operation, + variable, + }; + Tensor(Arc::new(tensor)) +} diff --git a/src/utils.rs b/src/utils.rs new file mode 100644 index 0000000..4b1e941 --- /dev/null +++ b/src/utils.rs @@ -0,0 +1,12 @@ +use std::str::FromStr; + +pub fn get_num_threads() -> usize { + // Respond to the same environment variable as rayon. + match std::env::var("RAYON_NUM_THREADS") + .ok() + .and_then(|s| usize::from_str(&s).ok()) + { + Some(x) if x > 0 => x, + Some(_) | None => num_cpus::get(), + } +} diff --git a/tests/gradient_tests.rs b/tests/gradient_tests.rs index 3b02102..81005d4 100644 --- a/tests/gradient_tests.rs +++ b/tests/gradient_tests.rs @@ -3,21 +3,21 @@ use phantom::{Device, Tensor}; #[test] fn simple_grad() -> Result<()> { - let five = Tensor::new(&[5f32, 5., 5.], Device::CPU)?; + let five = Tensor::new(&[5f32, 5., 5.], &Device::CPU)?; - let x = Tensor::var(&[3f32, 1., 4.], Device::CPU)?; + let x = Tensor::var(&[3f32, 1., 4.], &Device::CPU)?; let y = x.mul(&x)?.add(&x.mul(&five)?)?.add(&five)?; let gradients = y.backward()?; - let gradient_x = gradients.get(&x.id()).context("x has no gradient")?; + let gradient_x = gradients.get(&x).context("x has no gradient")?; assert_eq!(x.to_vector_rank_one::()?, [3., 1., 4.]); assert_eq!(y.to_vector_rank_one::()?, [29., 11., 41.]); assert_eq!(gradient_x.to_vector_rank_one::()?, [11., 7., 13.]); - let x = Tensor::var(&[4f32, 2., 8.], Device::CPU)?; + let x = Tensor::var(&[4f32, 2., 8.], &Device::CPU)?; let y = x.mul(&x)?.add(&x.mul(&five)?)?.add(&five)?; let gradients = y.backward()?; - let gradient_x = gradients.get(&x.id()).context("x has no gradient")?; + let gradient_x = gradients.get(&x).context("x has no gradient")?; assert_eq!(x.to_vector_rank_one::()?, [4., 2., 8.]); assert_eq!(y.to_vector_rank_one::()?, [41., 19., 109.]); @@ -25,16 +25,3 @@ fn simple_grad() -> Result<()> { Ok(()) } - -#[test] -fn simple_grad_constants() -> Result<()> { - let x = Tensor::var(&[3f32, 1., 4.], Device::CPU)?; - let y = (((&x * &x)? + &x * 5f64)? + 4f64)?; - let gradients = y.backward()?; - let gradient_x = gradients.get(&x.id()).context("x has no gradient")?; - - assert_eq!(x.to_vector_rank_one::()?, [3., 1., 4.]); - assert_eq!(y.to_vector_rank_one::()?, [28., 10., 40.]); - assert_eq!(gradient_x.to_vector_rank_one::()?, [11., 7., 13.]); - Ok(()) -} diff --git a/tests/tensor_tests.rs b/tests/tensor_tests.rs index 02adcee..dd5b2b1 100644 --- a/tests/tensor_tests.rs +++ b/tests/tensor_tests.rs @@ -2,7 +2,7 @@ use phantom::{DType, Device, Result, Tensor}; #[test] fn construct() -> Result<()> { - let tensor = Tensor::zeros(&[2, 3], DType::F32, Device::CPU); + let tensor = Tensor::zeros(&[2, 3], DType::F32, &Device::CPU)?; let rank = tensor.rank(); assert!(rank == 2); @@ -17,7 +17,7 @@ fn construct() -> Result<()> { #[test] fn zeros() -> Result<()> { - let tensor = Tensor::zeros((5, 2), DType::F32, Device::CPU); + let tensor = Tensor::zeros((5, 2), DType::F32, &Device::CPU)?; let (dim_one, dim_two) = tensor.shape().rank_two()?; assert_eq!(dim_one, 5); @@ -32,7 +32,7 @@ fn zeros() -> Result<()> { #[test] fn ones() -> Result<()> { - let tensor = Tensor::ones((5, 2), DType::F32, Device::CPU); + let tensor = Tensor::ones((5, 2), DType::F32, &Device::CPU)?; let (dim_one, dim_two) = tensor.shape().rank_two()?; assert_eq!(dim_one, 5); @@ -48,7 +48,7 @@ fn ones() -> Result<()> { #[test] fn rank_one() -> Result<()> { let data = &[1f32, 2f32, 3f32, 4f32, 5f32, 6f32]; - let tensor = Tensor::new(data, Device::CPU)?; + let tensor = Tensor::new(data, &Device::CPU)?; let dims = tensor.shape().rank_one()?; assert_eq!(dims, 6); @@ -65,7 +65,7 @@ fn rank_two() -> Result<()> { [1f32, 2f32, 3f32, 4f32, 5f32, 6f32], [7f32, 8f32, 9f32, 10f32, 11f32, 12f32], ]; - let tensor = Tensor::new(data, Device::CPU)?; + let tensor = Tensor::new(data, &Device::CPU)?; let dims = tensor.shape().rank_two()?; assert_eq!(dims, (2, 6)); @@ -78,8 +78,8 @@ fn rank_two() -> Result<()> { #[test] fn add_rank_one() -> Result<()> { - let a = Tensor::zeros(&[6], DType::F32, Device::CPU); - let b = Tensor::ones(&[6], DType::F32, Device::CPU); + let a = Tensor::zeros(&[6], DType::F32, &Device::CPU)?; + let b = Tensor::ones(&[6], DType::F32, &Device::CPU)?; let c = Tensor::add(&a, &b)?; @@ -91,8 +91,8 @@ fn add_rank_one() -> Result<()> { assert_eq!(content, vec![1f32; 6]); let data = &[1f32, 2f32, 3f32, 4f32, 5f32, 6f32]; - let a = Tensor::ones(&[6], DType::F32, Device::CPU); - let b = Tensor::new(data, Device::CPU)?; + let a = Tensor::ones(&[6], DType::F32, &Device::CPU)?; + let b = Tensor::new(data, &Device::CPU)?; let c = (&a + &b)?; let content: Vec = c.to_vector_rank_one()?; @@ -106,8 +106,8 @@ fn add_rank_one() -> Result<()> { #[test] fn add_rank_two() -> Result<()> { - let a = Tensor::zeros(&[2, 3], DType::F32, Device::CPU); - let b = Tensor::ones(&[2, 3], DType::F32, Device::CPU); + let a = Tensor::zeros(&[2, 3], DType::F32, &Device::CPU)?; + let b = Tensor::ones(&[2, 3], DType::F32, &Device::CPU)?; let c = Tensor::add(&a, &b)?; @@ -123,8 +123,8 @@ fn add_rank_two() -> Result<()> { [1f32, 2f32, 3f32, 4f32, 5f32, 6f32], [7f32, 8f32, 9f32, 10f32, 11f32, 12f32], ]; - let a = Tensor::ones(&[2, 6], DType::F32, Device::CPU); - let b = Tensor::new(data, Device::CPU)?; + let a = Tensor::ones(&[2, 6], DType::F32, &Device::CPU)?; + let b = Tensor::new(data, &Device::CPU)?; let c = (&a + &b)?; let content: Vec> = c.to_vector_rank_two()?; @@ -141,8 +141,8 @@ fn add_rank_two() -> Result<()> { #[test] fn mul_rank_one() -> Result<()> { - let a = Tensor::zeros(&[6], DType::F32, Device::CPU); - let b = Tensor::ones(&[6], DType::F32, Device::CPU); + let a = Tensor::zeros(&[6], DType::F32, &Device::CPU)?; + let b = Tensor::ones(&[6], DType::F32, &Device::CPU)?; let c = Tensor::mul(&a, &b)?; @@ -154,8 +154,8 @@ fn mul_rank_one() -> Result<()> { assert_eq!(content, vec![0f32; 6]); let data = &[1f32, 2f32, 3f32, 4f32, 5f32, 6f32]; - let a = Tensor::ones(&[6], DType::F32, Device::CPU); - let b = Tensor::new(data, Device::CPU)?; + let a = Tensor::ones(&[6], DType::F32, &Device::CPU)?; + let b = Tensor::new(data, &Device::CPU)?; let c = (&a * &b)?; let content: Vec = c.to_vector_rank_one()?; @@ -169,8 +169,8 @@ fn mul_rank_one() -> Result<()> { #[test] fn mul_rank_two() -> Result<()> { - let a = Tensor::zeros(&[2, 3], DType::F32, Device::CPU); - let b = Tensor::ones(&[2, 3], DType::F32, Device::CPU); + let a = Tensor::zeros(&[2, 3], DType::F32, &Device::CPU)?; + let b = Tensor::ones(&[2, 3], DType::F32, &Device::CPU)?; let c = Tensor::mul(&a, &b)?; @@ -186,8 +186,8 @@ fn mul_rank_two() -> Result<()> { [1f32, 2f32, 3f32, 4f32, 5f32, 6f32], [7f32, 8f32, 9f32, 10f32, 11f32, 12f32], ]; - let a = Tensor::ones(&[2, 6], DType::F32, Device::CPU); - let b = Tensor::new(data, Device::CPU)?; + let a = Tensor::ones(&[2, 6], DType::F32, &Device::CPU)?; + let b = Tensor::new(data, &Device::CPU)?; let c = (&a * &b)?; let content: Vec> = c.to_vector_rank_two()?; @@ -205,10 +205,10 @@ fn mul_rank_two() -> Result<()> { #[test] fn binary_chaining() -> Result<()> { let data_a = &[[3f32, 1., 4., 1., 5.], [2., 1., 7., 8., 2.]]; - let a = Tensor::new(data_a, Device::CPU)?; + let a = Tensor::new(data_a, &Device::CPU)?; let data_b = &[[5f32, 5., 5., 5., 5.], [2., 1., 7., 8., 2.]]; - let b = Tensor::new(data_b, Device::CPU)?; + let b = Tensor::new(data_b, &Device::CPU)?; let c = (&a + (&a * &a)? / (&a + &b))?; let dims = a.shape().rank_two()?; From a7d0af18dd6e9ae0a3fd18926933e6248059bef6 Mon Sep 17 00:00:00 2001 From: Nick Wall <46641379+walln@users.noreply.github.com> Date: Tue, 12 Sep 2023 00:07:53 -0500 Subject: [PATCH 2/4] feat: add simple matmul and simplify ops --- Cargo.toml | 2 +- src/backprop.rs | 65 +++++++++----- src/error.rs | 39 +++++++++ src/operation.rs | 80 ++++++++++++++--- src/shape.rs | 20 +++++ src/storage.rs | 27 ++++++ src/tensor.rs | 194 ++++++++++++++++++++++++++++++++---------- tests/tensor_tests.rs | 17 ++++ 8 files changed, 363 insertions(+), 81 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index 32d6a05..82f88fb 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -8,7 +8,7 @@ license = "MIT" readme = "README.md" [dependencies] -anyhow = "1.0.75" +anyhow = { version = "1.0.75", features = ["backtrace"]} num-traits = "0.2.16" thiserror = "1" diff --git a/src/backprop.rs b/src/backprop.rs index 32a589b..71f00ed 100644 --- a/src/backprop.rs +++ b/src/backprop.rs @@ -1,5 +1,5 @@ use crate::gradient_store::GradientStore; -use crate::operation::Operation; +use crate::operation::{BinaryOperations, Operation, UnaryOperations}; use crate::tensor::{Tensor, TensorID}; use crate::Error; use crate::Result; @@ -26,21 +26,11 @@ impl Tensor { nodes } else if let Some(op) = node.op() { match op { - Operation::Add(lhs, rhs) - | Operation::Sub(lhs, rhs) - | Operation::Mul(lhs, rhs) - | Operation::Div(lhs, rhs) => { - let (target, nodes) = walk(lhs, nodes, seen); - tracked |= target; - let (target, nodes) = walk(rhs, nodes, seen); - tracked |= target; - nodes - } - Operation::Sqr(node) - | Operation::Sqrt(node) - | Operation::Neg(node) - | Operation::Broadcast(node) - | Operation::ToDType(node) => { + Operation::Broadcast(node) + | Operation::ToDType(node) + | Operation::Transpose(node, _, _) + | Operation::Copy(node) + | Operation::Unary(node, _) => { let (target, nodes) = walk(node, nodes, seen); tracked |= target; nodes @@ -54,6 +44,14 @@ impl Tensor { nodes } } + + Operation::Binary(lhs, rhs, _) | Operation::Matmul(lhs, rhs) => { + let (target, nodes) = walk(lhs, nodes, seen); + tracked |= target; + let (target, nodes) = walk(rhs, nodes, seen); + tracked |= target; + nodes + } } } else { nodes @@ -103,19 +101,19 @@ impl Tensor { if let Some(op) = node.op() { match op { - Operation::Add(lhs, rhs) => { + Operation::Binary(lhs, rhs, BinaryOperations::Add) => { let lhs_gradient_sum = gradients.or_insert(lhs)?; *lhs_gradient_sum = lhs_gradient_sum.add(&gradient)?; let rhs_gradient_sum = gradients.or_insert(rhs)?; *rhs_gradient_sum = rhs_gradient_sum.add(&gradient)?; } - Operation::Sub(lhs, rhs) => { + Operation::Binary(lhs, rhs, BinaryOperations::Sub) => { let lhs_gradient_sum = gradients.or_insert(lhs)?; *lhs_gradient_sum = lhs_gradient_sum.add(&gradient)?; let rhs_gradient_sum = gradients.or_insert(rhs)?; *rhs_gradient_sum = rhs_gradient_sum.sub(&gradient)?; } - Operation::Mul(lhs, rhs) => { + Operation::Binary(lhs, rhs, BinaryOperations::Mul) => { let lhs_gradient = gradient.mul(rhs)?; let lhs_gradient_sum = gradients.or_insert(lhs)?; *lhs_gradient_sum = lhs_gradient_sum.add(&lhs_gradient)?; @@ -123,7 +121,7 @@ impl Tensor { let rhs_gradient_sum = gradients.or_insert(rhs)?; *rhs_gradient_sum = rhs_gradient_sum.add(&rhs_gradient)?; } - Operation::Div(lhs, rhs) => { + Operation::Binary(lhs, rhs, BinaryOperations::Div) => { let lhs_gradient = gradient.div(rhs)?; let lhs_gradient_sum = gradients.or_insert(lhs)?; *lhs_gradient_sum = lhs_gradient_sum.add(&lhs_gradient)?; @@ -136,17 +134,17 @@ impl Tensor { let gradient_sum = gradients.or_insert(arg)?; *gradient_sum = gradient_sum.add(&gradient_arg)? } - Operation::Sqr(arg) => { + Operation::Unary(arg, UnaryOperations::Sqr) => { let gradient_arg = arg.mul(&gradient)?.affine(2., 0.)?; let gradient_sum = gradients.or_insert(node)?; *gradient_sum = gradient_sum.add(&gradient_arg)? } - Operation::Sqrt(arg) => { + Operation::Unary(arg, UnaryOperations::Sqrt) => { let gradient_arg = gradient.div(arg)?.affine(0.5, 0.)?; let gradient_sum = gradients.or_insert(arg)?; *gradient_sum = gradient_sum.add(&gradient_arg)? } - Operation::Neg(arg) => { + Operation::Unary(arg, UnaryOperations::Neg) => { let gradient_sum = gradients.or_insert(arg)?; *gradient_sum = gradient_sum.sub(&gradient)? } @@ -159,6 +157,27 @@ impl Tensor { let gradient_sum = gradients.or_insert(arg)?; *gradient_sum = gradient_sum.add(&gradient.to_dtype(node.dtype())?)? } + Operation::Matmul(lhs, rhs) => { + // Skipping checks, the op went ok, we can skip + // the matmul size checks for now. + + let lhs_gradient = gradient.matmul(&rhs.t()?)?; + let lhs_gradient_sum = gradients.or_insert(lhs)?; + *lhs_gradient_sum = lhs_gradient_sum.add(&lhs_gradient)?; + + let rhs_gradient = lhs.t()?.matmul(&gradient)?; + let rhs_gradient_sum = gradients.or_insert(rhs)?; + *rhs_gradient_sum = rhs_gradient_sum.add(&rhs_gradient)?; + } + Operation::Transpose(arg, dim1, dim2) => { + let gradient_arg = gradient.transpose(*dim1, *dim2)?; + let gradient_sum = gradients.or_insert(arg)?; + *gradient_sum = gradient_sum.add(&gradient_arg)? + } + Operation::Copy(arg) => { + let gradient_sum = gradients.or_insert(arg)?; + *gradient_sum = gradient_sum.add(&gradient)? + } } } } diff --git a/src/error.rs b/src/error.rs index 0fbb01b..16809ea 100644 --- a/src/error.rs +++ b/src/error.rs @@ -3,6 +3,7 @@ use crate::{DType, Layout, Shape}; #[derive(thiserror::Error, Debug)] pub enum Error { + // TODO: consider renaming to unexpected dims despite the fact that rank is more accurate #[error("unexpected rank, expected: {expected}, actual: {actual}")] UnexpectedRank { expected: usize, @@ -10,6 +11,13 @@ pub enum Error { shape: Shape, }, + #[error("{op}: dimension index {dim} out of range for shape {shape:?}")] + DimOutOfRange { + shape: Shape, + dim: i32, + op: &'static str, + }, + #[error("unexpected dtype, expected: {expected:?}, actual: {actual:?}")] UnexpectedDType { expected: DType, actual: DType }, @@ -63,6 +71,37 @@ pub enum Error { #[error("unsupported dtype {dtype:?} for {op}")] UnsupportedDTypeForOperation { dtype: DType, op: &'static str }, + + #[error("{inner}\n{backtrace}")] + WithBacktrace { + inner: Box, + backtrace: Box, + }, + + /// User generated error message, typically created via `bail!`. + #[error("{0}")] + Message(String), +} + +impl Error { + pub fn backtrace(self) -> Self { + let backtrace = std::backtrace::Backtrace::capture(); + match backtrace.status() { + std::backtrace::BacktraceStatus::Disabled + | std::backtrace::BacktraceStatus::Unsupported => self, + _ => Self::WithBacktrace { + inner: Box::new(self), + backtrace: Box::new(backtrace), + }, + } + } + + pub fn message(err: T) -> Self + where + T: std::error::Error + Send + Sync + 'static, + { + Self::Message(err.to_string()).backtrace() + } } pub type Result = std::result::Result; diff --git a/src/operation.rs b/src/operation.rs index 4481c6b..cd5d1b8 100644 --- a/src/operation.rs +++ b/src/operation.rs @@ -1,22 +1,32 @@ use crate::Tensor; +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum BinaryOperations { + Add, + Mul, + Sub, + Div, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum UnaryOperations { + Sqr, + Sqrt, + Neg, +} + #[derive(Debug, Clone)] -pub(crate) enum Operation { - // Binary Operations - Add(Tensor, Tensor), - Sub(Tensor, Tensor), - Mul(Tensor, Tensor), - Div(Tensor, Tensor), - - // Unary Operations - Sqr(Tensor), - Sqrt(Tensor), - Neg(Tensor), +pub enum Operation { + Unary(Tensor, UnaryOperations), + Binary(Tensor, Tensor, BinaryOperations), // Casting and complex ops Affine { node: Tensor, mul: f64, add: f64 }, Broadcast(Tensor), + Transpose(Tensor, usize, usize), + Matmul(Tensor, Tensor), + Copy(Tensor), ToDType(Tensor), } @@ -98,3 +108,51 @@ macro_rules! unary_op { unary_op!(Sqr, "sqr", a, a * a); unary_op!(Sqrt, "sqrt", a, a.sqrt()); unary_op!(Neg, "neg", a, -a); + +#[derive(Clone, Debug)] +pub struct BackpropOperation(Option); + +impl BackpropOperation { + pub(crate) fn none() -> Self { + BackpropOperation(None) + } + + pub(crate) fn new>(args: &[A], f: impl Fn(Vec) -> Operation) -> Self { + let operation = if args.iter().any(|arg| arg.as_ref().track_op()) { + let args: Vec = args.iter().map(|arg| arg.as_ref().clone()).collect(); + Some(f(args)) + } else { + None + }; + Self(operation) + } + + pub(crate) fn new_unary(arg: &Tensor, f: impl Fn(Tensor) -> Operation) -> Self { + let operation = if arg.track_op() { + Some(f(arg.clone())) + } else { + None + }; + Self(operation) + } + + pub(crate) fn new_binary( + arg1: &Tensor, + arg2: &Tensor, + f: impl Fn(Tensor, Tensor) -> Operation, + ) -> Self { + let operation = if arg1.track_op() || arg2.track_op() { + Some(f(arg1.clone(), arg2.clone())) + } else { + None + }; + Self(operation) + } +} + +impl std::ops::Deref for BackpropOperation { + type Target = Option; + fn deref(&self) -> &Self::Target { + &self.0 + } +} diff --git a/src/shape.rs b/src/shape.rs index 3cb86eb..a9dd37b 100644 --- a/src/shape.rs +++ b/src/shape.rs @@ -163,6 +163,26 @@ impl std::fmt::Debug for Shape { } } +pub trait Dim { + fn to_index(&self, shape: &Shape, operation: &'static str) -> Result; +} + +impl Dim for usize { + fn to_index(&self, shape: &Shape, op: &'static str) -> Result { + let dim = *self; + if dim >= shape.dims().len() { + Err(Error::DimOutOfRange { + shape: shape.clone(), + dim: dim as i32, + op, + } + .backtrace())? + } else { + Ok(dim) + } + } +} + #[cfg(test)] mod tests { use super::*; diff --git a/src/storage.rs b/src/storage.rs index 3c6bfff..f8be545 100644 --- a/src/storage.rs +++ b/src/storage.rs @@ -102,6 +102,33 @@ impl Storage { } } + pub(crate) fn matmul( + &self, + rhs: &Self, + bmnk: (usize, usize, usize, usize), + lhs_layout: &Layout, + rhs_layout: &Layout, + ) -> Result { + self.matches_device(rhs, "matmul")?; + self.matches_dtype(rhs, "matmul")?; + + match (self, rhs) { + (Self::CPU(lhs), Self::CPU(rhs)) => { + let storage = lhs.matmul(rhs, bmnk, lhs_layout, rhs_layout)?; + Ok(Self::CPU(storage)) + } + (Self::MPS(lhs), Self::MPS(rhs)) => { + let storage = lhs.matmul(rhs, bmnk, lhs_layout, rhs_layout)?; + Ok(Self::MPS(storage)) + } + (lhs, rhs) => Err(Error::BinaryOperationDeviceMismatch { + lhs: lhs.device().location(), + rhs: rhs.device().location(), + op: "matmul", + }), + } + } + pub(crate) fn to_dtype(&self, layout: &Layout, dtype: DType) -> Result { match self { Storage::CPU(storage) => { diff --git a/src/tensor.rs b/src/tensor.rs index aa1e5ed..bd4de02 100644 --- a/src/tensor.rs +++ b/src/tensor.rs @@ -3,8 +3,9 @@ use std::sync::Arc; use crate::device::{Device, NDArray}; use crate::index::StridedIndex; -use crate::operation::Operation; -use crate::storage::{self, Storage}; +use crate::operation::{BackpropOperation, BinaryOperations, Operation, UnaryOperations}; +use crate::shape::Dim; +use crate::storage::Storage; use crate::WithDType; use crate::{DType, Error, Layout, Result, Shape}; @@ -24,8 +25,10 @@ pub struct Tensor_ { id: TensorID, storage: Arc, layout: Layout, - op: Option, + op: BackpropOperation, variable: bool, + dtype: DType, + device: Device, } /// Refcount tensors to make the construction of the graph cheap. Since tensors @@ -59,11 +62,9 @@ macro_rules! binary_operation { self.layout(), rhs.layout(), )?; - let op = if self.track_op() || rhs.track_op() { - Some(Operation::$operation_name(self.clone(), rhs.clone())) - } else { - None - }; + let op = BackpropOperation::new_binary(self, rhs, |a, b| { + Operation::Binary(a, b, BinaryOperations::$operation_name) + }); Ok(from_storage(storage, shape.clone(), op, false)) } }; @@ -76,11 +77,9 @@ macro_rules! unary_operation { let storage = self .storage .unary_operation::(self.layout())?; - let op = if self.track_op() { - Some(Operation::$operation_name(self.clone())) - } else { - None - }; + let op = BackpropOperation::new_unary(self, |arg| { + Operation::Unary(arg, UnaryOperations::$operation_name) + }); Ok(from_storage(storage, shape.clone(), op, false)) } }; @@ -99,7 +98,12 @@ impl Tensor { return Err(Error::ShapeMismatch { buffer_size, shape }); } let storage = device.storage(array)?; - Ok(from_storage(storage, shape, None, variable)) + Ok(from_storage( + storage, + shape, + BackpropOperation::none(), + variable, + )) } /// Creates a new tensor from a slice of data. @@ -126,6 +130,14 @@ impl Tensor { Self::new_impl(array, shape, device, true) } + pub fn from_slice, D: crate::WithDType>( + array: &[D], + shape: S, + device: &Device, + ) -> Result { + Self::new_impl(array, shape.into(), device, false) + } + pub(crate) fn zeros_impl>( shape: S, dtype: DType, @@ -135,10 +147,21 @@ impl Tensor { if variable { let shape = shape.into(); let storage = device.zeros(&shape, dtype)?; - Ok(from_storage(storage, shape, None, variable)) + Ok(from_storage( + storage, + shape, + BackpropOperation::none(), + variable, + )) } else { let storage = device.zeros(&crate::shape::SCALAR, dtype)?; - from_storage(storage, crate::shape::SCALAR, None, variable).broadcast_as(shape) + from_storage( + storage, + crate::shape::SCALAR, + BackpropOperation::none(), + variable, + ) + .broadcast_as(shape) } } @@ -185,10 +208,21 @@ impl Tensor { if variable { let shape = shape.into(); let storage = device.ones(&shape, dtype)?; - Ok(from_storage(storage, shape, None, variable)) + Ok(from_storage( + storage, + shape, + BackpropOperation::none(), + variable, + )) } else { let storage = device.ones(&crate::shape::SCALAR, dtype)?; - from_storage(storage, crate::shape::SCALAR, None, variable).broadcast_as(shape) + from_storage( + storage, + crate::shape::SCALAR, + BackpropOperation::none(), + variable, + ) + .broadcast_as(shape) } } @@ -410,12 +444,8 @@ impl Tensor { let mut storage = self.device().zeros(shape, self.dtype())?; self.storage .copy_strided_source(&mut storage, 0, self.layout())?; - Ok(from_storage( - storage, - shape.clone(), - None, // TODO - false, - )) + let operation = BackpropOperation::new_unary(self, Operation::Copy); + Ok(from_storage(storage, shape.clone(), operation, false)) } } @@ -507,31 +537,20 @@ impl Tensor { /// TODO: Add doctests pub fn affine(&self, mul: f64, add: f64) -> Result { let storage = self.storage.affine(self.layout(), mul, add)?; - let operation = if self.track_op() { - Some(Operation::Affine { - node: self.clone(), - mul, - add, - }) - } else { - None - }; + let operation = + BackpropOperation::new_unary(self, |node| Operation::Affine { node, mul, add }); Ok(from_storage(storage, self.shape(), operation, false)) } pub fn broadcast_as>(&self, shape: S) -> Result { - let operation = if self.track_op() { - Some(Operation::Broadcast(self.clone())) - } else { - None - }; - let tensor = Tensor_ { id: TensorID::new(), storage: self.storage.clone(), layout: self.layout.broadcast_as(shape)?, - op: operation, + op: BackpropOperation::new_unary(self, Operation::Broadcast), variable: false, + dtype: self.dtype, + device: self.device, }; Ok(Tensor(Arc::new(tensor))) @@ -557,15 +576,93 @@ impl Tensor { } else { let shape = self.shape(); let storage = self.storage.to_dtype(&self.layout(), dtype)?; - let operation = if self.track_op() { - Some(Operation::ToDType(self.clone())) - } else { - None - }; + let operation = BackpropOperation::new_unary(self, Operation::ToDType); Ok(from_storage(storage, shape.clone(), operation, false)) } } + pub fn matmul(&self, rhs: &Self) -> Result { + let a_dims = self.dims(); + let b_dims = rhs.dims(); + + if a_dims.len() < 2 || b_dims.len() != a_dims.len() { + Err(Error::BinaryOperationShapeMismatch { + lhs: self.shape().clone(), + rhs: rhs.shape().clone(), + op: "matmul", + } + .backtrace())? + } + + let dim = a_dims.len(); + + let m = a_dims[dim - 2]; + let n = b_dims[dim - 1]; + let k = a_dims[dim - 1]; + let k2 = b_dims[dim - 2]; + + let c_shape = Shape::from(&a_dims[..dim - 2]).extend(&[m, n]); + let batching = a_dims[..dim - 2].iter().product(); + let batching_b = b_dims[..dim - 2].iter().product(); + if k != k2 || batching != batching_b { + Err(Error::BinaryOperationShapeMismatch { + lhs: self.shape().clone(), + rhs: rhs.shape().clone(), + op: "matmul", + } + .backtrace())? + } + + let storage = self.storage.matmul( + &rhs.storage, + (batching, m, n, k), + self.layout(), + rhs.layout(), + )?; + + let operation = BackpropOperation::new_binary(self, rhs, Operation::Matmul); + Ok(from_storage(storage, c_shape, operation, false)) + } + + /// Transpose the input tesnor by swapping the dimensions + /// + /// ```rust + /// use phantom::{Tensor, Device}; + /// let tensor = Tensor::new(&[[0f32, 1.], [2., 3.], [4., 5.]], &Device::CPU)?; + /// let tensor = tensor.t()?; + /// assert_eq!(tensor.to_vector_rank_two::()?, &[[0.0, 2.0, 4.0], [1.0, 3.0, 5.0]]); + /// # Ok::<(), phantom::Error>(()) + /// ``` + pub fn t(&self) -> Result { + let rank = self.rank(); + if rank < 2 { + Err(Error::UnexpectedRank { + expected: 2, + actual: rank, + shape: self.shape().clone(), + } + .backtrace())? + } + self.transpose(rank - 2, rank - 1) + } + + /// Transpose the input tesnor by swapping the dimensions + pub fn transpose(&self, dim1: D1, dim2: D2) -> Result { + let dim1 = dim1.to_index(self.shape(), "transpose")?; + let dim2 = dim2.to_index(self.shape(), "transpose")?; + let op = BackpropOperation::new_unary(self, |t| Operation::Transpose(t, dim1, dim2)); + let tensor = Tensor_ { + id: TensorID::new(), + storage: self.storage.clone(), + layout: self.layout.transpose(dim1, dim2)?, + op, + variable: false, + dtype: self.dtype, + device: self.device.clone(), + }; + Ok(Tensor(Arc::new(tensor))) + } + binary_operation!(add, Add); binary_operation!(sub, Sub); binary_operation!(mul, Mul); @@ -637,15 +734,20 @@ binary_trait!(Div, div, |v| 1. / v, |_| 0.); fn from_storage>( storage: Storage, shape: S, - operation: Option, + operation: BackpropOperation, variable: bool, ) -> Tensor { + let dtype = storage.dtype(); + let device = storage.device(); + let tensor = Tensor_ { id: TensorID::new(), storage: Arc::new(storage), layout: Layout::contiguous(shape), op: operation, variable, + dtype, + device, }; Tensor(Arc::new(tensor)) } diff --git a/tests/tensor_tests.rs b/tests/tensor_tests.rs index dd5b2b1..ed08d45 100644 --- a/tests/tensor_tests.rs +++ b/tests/tensor_tests.rs @@ -225,3 +225,20 @@ fn binary_chaining() -> Result<()> { Ok(()) } + +#[test] +fn matmul() -> Result<()> { + let device = &Device::CPU; + let data = vec![1.0f32, 2.0, 3.0, 4.0]; + let a = Tensor::from_slice(&data, (2, 2), device)?; + let data = vec![1.0f32, 2.0, 3.0, 4.0]; + let b = Tensor::from_slice(&data, (2, 2), device)?; + + let c = a.matmul(&b)?; + assert_eq!( + c.to_vector_rank_two::()?, + &[[7.0f32, 10.0], [15.0, 22.0]] + ); + + Ok(()) +} From 5205129f20b26c5dd669d04261ec1c4b1ad4770d Mon Sep 17 00:00:00 2001 From: Nick Wall <46641379+walln@users.noreply.github.com> Date: Tue, 12 Sep 2023 12:27:09 -0500 Subject: [PATCH 3/4] feat: add matmul tests --- tests/tensor_tests.rs | 18 ++++++++++++++++++ 1 file changed, 18 insertions(+) diff --git a/tests/tensor_tests.rs b/tests/tensor_tests.rs index ed08d45..fa446e0 100644 --- a/tests/tensor_tests.rs +++ b/tests/tensor_tests.rs @@ -240,5 +240,23 @@ fn matmul() -> Result<()> { &[[7.0f32, 10.0], [15.0, 22.0]] ); + let data = vec![1.0f32, 2.0]; + let a = Tensor::from_slice(&data, (2, 1), device)?; + let data = vec![3.0f32, 4.0]; + let b = Tensor::from_slice(&data, (1, 2), device)?; + let c = a.matmul(&b)?; + assert_eq!(c.to_vector_rank_two::()?, &[&[3.0, 4.0], &[6.0, 8.0]]); + + let data: Vec<_> = (0..6).map(|i| i as f32).collect(); + let a = Tensor::from_slice(&data, (2, 3), device)?; + let data: Vec<_> = (0..6).map(|i| (i + 2) as f32).collect(); + let b = Tensor::from_slice(&data, (3, 2), device)?; + let c = a.matmul(&b)?; + assert_eq!(c.to_vector_rank_two::()?, &[&[16., 19.], &[52., 64.]]); + + // TODO: test matmul with broadcasting + // TODO: tests with higher ranks + // TODO: tests on contigious transposed tensors + Ok(()) } From 0d597986b9830cbf9e0c4e6dd64fe496c5f33a5c Mon Sep 17 00:00:00 2001 From: Nick Wall <46641379+walln@users.noreply.github.com> Date: Wed, 13 Sep 2023 20:37:41 -0500 Subject: [PATCH 4/4] refactor: move to cargo workspace --- Cargo.toml | 23 +--- phantom-core/Cargo.toml | 19 ++++ {src => phantom-core/src}/backend/backend.rs | 0 .../src}/backend/cpu_backend.rs | 0 {src => phantom-core/src}/backend/mod.rs | 0 .../src}/backend/mps_backend.rs | 0 {src => phantom-core/src}/backprop.rs | 4 +- {src => phantom-core/src}/device.rs | 0 {src => phantom-core/src}/dtype.rs | 0 {src => phantom-core/src}/error.rs | 0 {src => phantom-core/src}/gradient_store.rs | 0 {src => phantom-core/src}/index.rs | 0 {src => phantom-core/src}/layout.rs | 0 {src => phantom-core/src}/lib.rs | 0 {src => phantom-core/src}/operation.rs | 0 {src => phantom-core/src}/shape.rs | 0 {src => phantom-core/src}/storage.rs | 0 {src => phantom-core/src}/tensor.rs | 106 +++++++++--------- {src => phantom-core/src}/utils.rs | 0 .../tests}/gradient_tests.rs | 2 +- {tests => phantom-core/tests}/tensor_tests.rs | 2 +- 21 files changed, 81 insertions(+), 75 deletions(-) create mode 100644 phantom-core/Cargo.toml rename {src => phantom-core/src}/backend/backend.rs (100%) rename {src => phantom-core/src}/backend/cpu_backend.rs (100%) rename {src => phantom-core/src}/backend/mod.rs (100%) rename {src => phantom-core/src}/backend/mps_backend.rs (100%) rename {src => phantom-core/src}/backprop.rs (98%) rename {src => phantom-core/src}/device.rs (100%) rename {src => phantom-core/src}/dtype.rs (100%) rename {src => phantom-core/src}/error.rs (100%) rename {src => phantom-core/src}/gradient_store.rs (100%) rename {src => phantom-core/src}/index.rs (100%) rename {src => phantom-core/src}/layout.rs (100%) rename {src => phantom-core/src}/lib.rs (100%) rename {src => phantom-core/src}/operation.rs (100%) rename {src => phantom-core/src}/shape.rs (100%) rename {src => phantom-core/src}/storage.rs (100%) rename {src => phantom-core/src}/tensor.rs (90%) rename {src => phantom-core/src}/utils.rs (100%) rename {tests => phantom-core/tests}/gradient_tests.rs (96%) rename {tests => phantom-core/tests}/tensor_tests.rs (99%) diff --git a/Cargo.toml b/Cargo.toml index 82f88fb..6d31da6 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,19 +1,6 @@ -[package] -name = "phantom" -version = "0.1.0" -edition = "2021" +[workspace] +resolver = "2" -description = "A forward mode autodiff and deep learning library for Rust" -license = "MIT" -readme = "README.md" - -[dependencies] -anyhow = { version = "1.0.75", features = ["backtrace"]} -num-traits = "0.2.16" -thiserror = "1" - -# TODO: Switch back to the official gemm implementation once something similar to -# https://github.com/sarah-ek/gemm/pull/8 is available. -gemm = { git = "https://github.com/LaurentMazare/gemm.git", branch = "f16-vectorize-pack" } -rand = "0.8.5" -num_cpus = "1.16.0" +members = [ + "phantom-core", +] \ No newline at end of file diff --git a/phantom-core/Cargo.toml b/phantom-core/Cargo.toml new file mode 100644 index 0000000..73a0d2e --- /dev/null +++ b/phantom-core/Cargo.toml @@ -0,0 +1,19 @@ +[package] +name = "phantom-core" +version = "0.1.0" +edition = "2021" + +description = "A forward mode autodiff and deep learning library for Rust" +license = "MIT" +readme = "README.md" + +[dependencies] +anyhow = { version = "1.0.75", features = ["backtrace"]} +num-traits = "0.2.16" +thiserror = "1" + +# TODO: Switch back to the official gemm implementation once something similar to +# https://github.com/sarah-ek/gemm/pull/8 is available. +gemm = { git = "https://github.com/LaurentMazare/gemm.git", branch = "f16-vectorize-pack" } +rand = "0.8.5" +num_cpus = "1.16.0" diff --git a/src/backend/backend.rs b/phantom-core/src/backend/backend.rs similarity index 100% rename from src/backend/backend.rs rename to phantom-core/src/backend/backend.rs diff --git a/src/backend/cpu_backend.rs b/phantom-core/src/backend/cpu_backend.rs similarity index 100% rename from src/backend/cpu_backend.rs rename to phantom-core/src/backend/cpu_backend.rs diff --git a/src/backend/mod.rs b/phantom-core/src/backend/mod.rs similarity index 100% rename from src/backend/mod.rs rename to phantom-core/src/backend/mod.rs diff --git a/src/backend/mps_backend.rs b/phantom-core/src/backend/mps_backend.rs similarity index 100% rename from src/backend/mps_backend.rs rename to phantom-core/src/backend/mps_backend.rs diff --git a/src/backprop.rs b/phantom-core/src/backprop.rs similarity index 98% rename from src/backprop.rs rename to phantom-core/src/backprop.rs index 71f00ed..9ffaafd 100644 --- a/src/backprop.rs +++ b/phantom-core/src/backprop.rs @@ -76,7 +76,7 @@ impl Tensor { /// The gradient of a node is the sum of the gradients of all the nodes that depend on it. /// The gradient of a node is computed by applying the chain rule to the node's operation. /// ```rust - /// use phantom::{Tensor, Device}; + /// use phantom_core::{Tensor, Device}; /// /// let x = Tensor::new(&[[2f32, 2.], [1f32, 2.]], &Device::CPU)?; /// let y = Tensor::new(&[[2f32, 2.], [5f32, 6.]], &Device::CPU)?; @@ -84,7 +84,7 @@ impl Tensor { /// let gradients = z.backward()?; /// assert_eq!(gradients.len(), 1); /// - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn backward(&self) -> Result { let sorted_nodes = self.sorted_nodes(); diff --git a/src/device.rs b/phantom-core/src/device.rs similarity index 100% rename from src/device.rs rename to phantom-core/src/device.rs diff --git a/src/dtype.rs b/phantom-core/src/dtype.rs similarity index 100% rename from src/dtype.rs rename to phantom-core/src/dtype.rs diff --git a/src/error.rs b/phantom-core/src/error.rs similarity index 100% rename from src/error.rs rename to phantom-core/src/error.rs diff --git a/src/gradient_store.rs b/phantom-core/src/gradient_store.rs similarity index 100% rename from src/gradient_store.rs rename to phantom-core/src/gradient_store.rs diff --git a/src/index.rs b/phantom-core/src/index.rs similarity index 100% rename from src/index.rs rename to phantom-core/src/index.rs diff --git a/src/layout.rs b/phantom-core/src/layout.rs similarity index 100% rename from src/layout.rs rename to phantom-core/src/layout.rs diff --git a/src/lib.rs b/phantom-core/src/lib.rs similarity index 100% rename from src/lib.rs rename to phantom-core/src/lib.rs diff --git a/src/operation.rs b/phantom-core/src/operation.rs similarity index 100% rename from src/operation.rs rename to phantom-core/src/operation.rs diff --git a/src/shape.rs b/phantom-core/src/shape.rs similarity index 100% rename from src/shape.rs rename to phantom-core/src/shape.rs diff --git a/src/storage.rs b/phantom-core/src/storage.rs similarity index 100% rename from src/storage.rs rename to phantom-core/src/storage.rs diff --git a/src/tensor.rs b/phantom-core/src/tensor.rs similarity index 90% rename from src/tensor.rs rename to phantom-core/src/tensor.rs index bd4de02..5d18da5 100644 --- a/src/tensor.rs +++ b/phantom-core/src/tensor.rs @@ -108,10 +108,10 @@ impl Tensor { /// Creates a new tensor from a slice of data. /// ```rust - /// use phantom::{Tensor, Device, Shape}; + /// use phantom_core::{Tensor, Device, Shape}; /// let tensor = Tensor::new(&[0f32, 1., 2., 3., 4., 5.], &Device::CPU)?; /// assert_eq!(tensor.shape(), &Shape::from(&[6])); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn new(array: A, device: &Device) -> Result { let shape = array.shape()?; @@ -120,10 +120,10 @@ impl Tensor { /// Creates a new variable tensor from a slice of data. /// ```rust - /// use phantom::{Tensor, Device}; + /// use phantom_core::{Tensor, Device}; /// let tensor = Tensor::var(&[0f32, 1., 2., 3., 4., 5.], &Device::CPU)?; /// assert_eq!(tensor.is_variable(), true); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn var(array: A, device: &Device) -> Result { let shape = array.shape()?; @@ -167,10 +167,10 @@ impl Tensor { /// Creates a new tensor of zeros with the given shape and data type. /// ```rust - /// use phantom::{Tensor, Device, Shape}; - /// let tensor = Tensor::zeros(&[2, 2], phantom::DType::F32, &Device::CPU)?; + /// use phantom_core::{Tensor, Device, Shape}; + /// let tensor = Tensor::zeros(&[2, 2], phantom_core::DType::F32, &Device::CPU)?; /// assert_eq!(tensor.shape(), &Shape::from(&[2, 2])); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn zeros>(shape: S, dtype: DType, device: &Device) -> Result { Self::zeros_impl(shape, dtype, device, false) @@ -178,11 +178,11 @@ impl Tensor { /// Creates a new variable tensor of zeros in the same shape and data type as the input tensor. /// ```rust - /// use phantom::{Tensor, Device}; + /// use phantom_core::{Tensor, Device}; /// let tensor = Tensor::var(&[[0f32, 1.], [2., 3.]], &Device::CPU)?; /// let zeros = tensor.zeros_like()?; /// assert_eq!(zeros.shape(), tensor.shape()); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn zeros_like(&self) -> Result { Tensor::zeros(self.shape(), self.dtype(), &self.device()) @@ -190,10 +190,10 @@ impl Tensor { /// Creates a new variable tensor of zeros with the given shape and data type. /// ```rust - /// use phantom::{Tensor, Device, Shape}; - /// let tensor = Tensor::zeros_var(&[2, 2], phantom::DType::F32, &Device::CPU)?; + /// use phantom_core::{Tensor, Device, Shape}; + /// let tensor = Tensor::zeros_var(&[2, 2], phantom_core::DType::F32, &Device::CPU)?; /// assert_eq!(tensor.is_variable(), true); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn zeros_var>(shape: S, dtype: DType, device: &Device) -> Result { Self::zeros_impl(shape, dtype, device, true) @@ -228,10 +228,10 @@ impl Tensor { /// Creates a new tensor of ones with the given shape and data type. /// ```rust - /// use phantom::{Tensor, Device, Shape}; - /// let tensor = Tensor::ones(&[2, 2], phantom::DType::F32, &Device::CPU)?; + /// use phantom_core::{Tensor, Device, Shape}; + /// let tensor = Tensor::ones(&[2, 2], phantom_core::DType::F32, &Device::CPU)?; /// assert_eq!(tensor.shape(), &Shape::from(&[2, 2])); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn ones>(shape: S, dtype: DType, device: &Device) -> Result { Self::ones_impl(shape, dtype, device, false) @@ -239,11 +239,11 @@ impl Tensor { /// Creates a new variable tensor of ones in the same shape and data type as the input tensor. /// ```rust - /// use phantom::{Tensor, Device}; + /// use phantom_core::{Tensor, Device}; /// let tensor = Tensor::var(&[[0f32, 1.], [2., 3.]], &Device::CPU)?; /// let ones = tensor.ones_like()?; /// assert_eq!(ones.shape(), tensor.shape()); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn ones_like(&self) -> Result { Tensor::ones(self.shape(), self.dtype(), &self.device()) @@ -251,10 +251,10 @@ impl Tensor { /// Creates a new variable tensor of ones with the given shape and data type. /// ```rust - /// use phantom::{Tensor, Device, Shape}; - /// let tensor = Tensor::ones_var(&[2, 2], phantom::DType::F32, &Device::CPU)?; + /// use phantom_core::{Tensor, Device, Shape}; + /// let tensor = Tensor::ones_var(&[2, 2], phantom_core::DType::F32, &Device::CPU)?; /// assert_eq!(tensor.is_variable(), true); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn ones_var>(shape: S, dtype: DType, device: &Device) -> Result { Self::ones_impl(shape, dtype, device, true) @@ -262,10 +262,10 @@ impl Tensor { /// Converts the tensor to a scalar if the tensor is rank 0. /// ```rust - /// use phantom::{Tensor, Device}; + /// use phantom_core::{Tensor, Device}; /// let tensor = Tensor::new(0f32, &Device::CPU)?; /// assert_eq!(tensor.to_scalar::()?, 0f32); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn to_scalar(&self) -> Result { if self.rank() != 0 { @@ -289,10 +289,10 @@ impl Tensor { /// Returns the unique identifier for this tensor. /// ```rust - /// use phantom::{Tensor, Device}; + /// use phantom_core::{Tensor, Device}; /// let tensor = Tensor::new(&[0f32], &Device::CPU)?; /// assert_eq!(tensor.id(), tensor.id()); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn id(&self) -> TensorID { self.id @@ -300,10 +300,10 @@ impl Tensor { /// Returns the data type of this tensor used on the storage backend. /// ```rust - /// use phantom::{Tensor, Device}; + /// use phantom_core::{Tensor, Device}; /// let tensor = Tensor::new(&[0f32], &Device::CPU)?; - /// assert_eq!(tensor.dtype(), phantom::DType::F32); - /// # Ok::<(), phantom::Error>(()) + /// assert_eq!(tensor.dtype(), phantom_core::DType::F32); + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn dtype(&self) -> DType { self.storage.dtype() @@ -311,10 +311,10 @@ impl Tensor { /// Returns the device that this tensor is stored on. /// ```rust - /// use phantom::{Tensor, Device}; + /// use phantom_core::{Tensor, Device}; /// let tensor = Tensor::new(&[0f32], &Device::CPU)?; /// assert_eq!(tensor.device(), Device::CPU); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn device(&self) -> Device { self.storage.device() @@ -327,10 +327,10 @@ impl Tensor { /// Returns the shape of the tensor. /// ```rust - /// use phantom::{Tensor, Device, Shape}; + /// use phantom_core::{Tensor, Device, Shape}; /// let tensor = Tensor::new(&[[0f32, 1.], [2., 3.]], &Device::CPU)?; /// assert_eq!(tensor.shape(), &Shape::from(&[2, 2])); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn shape(&self) -> &Shape { &self.layout.shape() @@ -338,10 +338,10 @@ impl Tensor { /// Returns the rank of the tensor. /// ```rust - /// use phantom::{Tensor, Device}; + /// use phantom_core::{Tensor, Device}; /// let tensor = Tensor::new(&[[0f32, 1.], [2., 3.]], &Device::CPU)?; /// assert_eq!(tensor.rank(), 2); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn rank(&self) -> usize { self.layout.shape().rank() @@ -349,10 +349,10 @@ impl Tensor { /// Returns the dimension size for each axis of the tensor. /// ```rust - /// use phantom::{Tensor, Device}; + /// use phantom_core::{Tensor, Device}; /// let tensor = Tensor::new(&[[0f32, 1.], [2., 3.]], &Device::CPU)?; /// assert_eq!(tensor.dims(), &[2, 2]); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn dims(&self) -> &[usize] { self.layout.shape().dims() @@ -360,10 +360,10 @@ impl Tensor { /// Returns the total number of values in the tensor. /// ```rust - /// use phantom::{Tensor, Device}; + /// use phantom_core::{Tensor, Device}; /// let tensor = Tensor::new(&[[0f32, 1.], [2., 3.]], &Device::CPU)?; /// assert_eq!(tensor.elem_count(), 4); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn elem_count(&self) -> usize { self.layout.shape().elem_count() @@ -371,10 +371,10 @@ impl Tensor { /// Returns the element-wise stride of the tensor. /// ```rust - /// use phantom::{Tensor, Device}; + /// use phantom_core::{Tensor, Device}; /// let tensor = Tensor::new(&[[0f32, 1.], [2., 3.]], &Device::CPU)?; /// assert_eq!(tensor.stride(), &[2, 1]); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn stride(&self) -> &[usize] { &self.layout.stride() @@ -393,10 +393,10 @@ impl Tensor { /// Returns true if the tensor is a variable that is tracked during backpropagation. /// ```rust - /// use phantom::{Tensor, Device}; + /// use phantom_core::{Tensor, Device}; /// let tensor = Tensor::var(&[[0f32, 1.], [2., 3.]], &Device::CPU)?; /// assert!(tensor.is_variable()); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn is_variable(&self) -> bool { self.variable @@ -405,13 +405,13 @@ impl Tensor { /// Creates an iterator that yields the offset position of each element in the buffer. Allowing /// the elements to be iterated over in lexicographic order. /// ```rust - /// use phantom::{Tensor, Device}; + /// use phantom_core::{Tensor, Device}; /// let a = Tensor::new(&[[0f32], [2.]], &Device::CPU)?; /// let mut iter = a.strided_index(); /// assert_eq!(iter.next(), Some(0)); /// assert_eq!(iter.next(), Some(1)); /// assert_eq!(iter.next(), None); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn strided_index(&self) -> StridedIndex { self.layout.strided_index() @@ -419,11 +419,11 @@ impl Tensor { /// Returns true if the tensor is contiguous in memory. /// ```rust - /// use phantom::{Tensor, Device}; + /// use phantom_core::{Tensor, Device}; /// // Contigious example /// let a = Tensor::new(&[0f32], &Device::CPU)?; /// assert!(a.is_contiguous()); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn is_contiguous(&self) -> bool { let mut accumulated_stride = 1; @@ -451,10 +451,10 @@ impl Tensor { /// Returns the contents of the rank 1 tensor as a vector. /// ```rust - /// use phantom::{Tensor, Device}; + /// use phantom_core::{Tensor, Device}; /// let a = Tensor::new(&[0f32, 1., 2., 3., 4., 5.], &Device::CPU)?; /// assert_eq!(a.to_vector_rank_one::()?, &[0., 1., 2., 3., 4., 5.]); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn to_vector_rank_one(&self) -> Result> { if self.rank() != 1 { @@ -475,10 +475,10 @@ impl Tensor { /// Returns the contents of the rank 2 tensor as a vector of vectors in row-major order. /// ```rust - /// use phantom::{Tensor, Device}; + /// use phantom_core::{Tensor, Device}; /// let a = Tensor::new(&[[0f32, 1.], [2., 3.], [4., 5.]], &Device::CPU)?; /// assert_eq!(a.to_vector_rank_two::()?, &[[0., 1.], [2., 3.], [4., 5.]]); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn to_vector_rank_two(&self) -> Result>> { let (dim_one, dim_two) = self.shape().rank_two()?; @@ -501,11 +501,11 @@ impl Tensor { /// Checks to see if the shapes of two tensors attempting a binary operation match /// and the operation can be performed, returning the shape if successful. /// ```rust - /// use phantom::{Tensor, Device}; + /// use phantom_core::{Tensor, Device}; /// let a = Tensor::new(&[[0f32, 1.], [2., 3.]], &Device::CPU)?; /// let b = Tensor::new(&[[0f32, 1.], [2., 3.]], &Device::CPU)?; /// assert_eq!(a.binary_operation_shape_matches(&b, "add")?, a.shape()); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn binary_operation_shape_matches( &self, @@ -627,11 +627,11 @@ impl Tensor { /// Transpose the input tesnor by swapping the dimensions /// /// ```rust - /// use phantom::{Tensor, Device}; + /// use phantom_core::{Tensor, Device}; /// let tensor = Tensor::new(&[[0f32, 1.], [2., 3.], [4., 5.]], &Device::CPU)?; /// let tensor = tensor.t()?; /// assert_eq!(tensor.to_vector_rank_two::()?, &[[0.0, 2.0, 4.0], [1.0, 3.0, 5.0]]); - /// # Ok::<(), phantom::Error>(()) + /// # Ok::<(), phantom_core::Error>(()) /// ``` pub fn t(&self) -> Result { let rank = self.rank(); diff --git a/src/utils.rs b/phantom-core/src/utils.rs similarity index 100% rename from src/utils.rs rename to phantom-core/src/utils.rs diff --git a/tests/gradient_tests.rs b/phantom-core/tests/gradient_tests.rs similarity index 96% rename from tests/gradient_tests.rs rename to phantom-core/tests/gradient_tests.rs index 81005d4..e20a3c3 100644 --- a/tests/gradient_tests.rs +++ b/phantom-core/tests/gradient_tests.rs @@ -1,5 +1,5 @@ use anyhow::{Context, Result}; -use phantom::{Device, Tensor}; +use phantom_core::{Device, Tensor}; #[test] fn simple_grad() -> Result<()> { diff --git a/tests/tensor_tests.rs b/phantom-core/tests/tensor_tests.rs similarity index 99% rename from tests/tensor_tests.rs rename to phantom-core/tests/tensor_tests.rs index fa446e0..96c1c4b 100644 --- a/tests/tensor_tests.rs +++ b/phantom-core/tests/tensor_tests.rs @@ -1,4 +1,4 @@ -use phantom::{DType, Device, Result, Tensor}; +use phantom_core::{DType, Device, Result, Tensor}; #[test] fn construct() -> Result<()> {