acyclib 0.3.0

ML library for directed acyclic tensor graphs.
Documentation
use crate::device::{
    Device, OperationError,
    function::DeviceOperation,
    operation::{BaseOperations, CoreDeviceOps, DiffableFromOutput},
    tensor::TensorRef,
};

#[derive(Clone, Copy, Debug, PartialEq)]
pub enum UnaryOp {
    DiffableFromOutput(DiffableFromOutput),
    Add(f32),
    Mul(f32),
    AbsPow(f32),
}

#[derive(Clone, Copy, Debug, PartialEq)]
pub enum Reduce {
    Sum,
    Avg,
}

#[derive(Clone)]
pub struct MaybeUpdateBatchSize<D: Device> {
    pub input: TensorRef<D>,
    pub output: TensorRef<D>,
}

impl<D: Device> DeviceOperation<D> for MaybeUpdateBatchSize<D> {
    fn opname(&self) -> String {
        "MaybeUpdateBatchSize".to_string()
    }

    fn execute(&self) -> Result<(), OperationError<D::DeviceError>> {
        let input = self.input.borrow();
        let mut output = self.output.dense_mut();

        if output.batch_size() != input.values.batch_size() {
            output.set_batch_size(input.values.batch_size())?;
        }

        Ok(())
    }
}

#[derive(Clone)]
pub struct ReduceAcrossBatch<D: Device> {
    pub input: TensorRef<D>,
    pub output: TensorRef<D>,
    pub input_mul: f32,
    pub output_mul: f32,
    pub reduction: Reduce,
}

impl<D: Device> DeviceOperation<D> for ReduceAcrossBatch<D> {
    fn opname(&self) -> String {
        format!("ReduceAcrossBatch({:?})", self.reduction)
    }

    fn execute(&self) -> Result<(), OperationError<D::DeviceError>> {
        let input = self.input.dense();
        let mut output = self.output.dense_mut();

        if input.batch_size().is_none() || output.batch_size().is_some() {
            return Err(OperationError::MismatchedBatchSizes);
        }

        if input.single_size() != output.single_size() {
            return Err(OperationError::InvalidTensorFormat);
        }

        let bs = input.batch_size().unwrap_or(1);

        let scale = match self.reduction {
            Reduce::Avg => 1.0 / bs as f32,
            Reduce::Sum => 1.0,
        };

        output.buf.reduce_across_batch(input.single_size(), bs, self.output_mul, self.input_mul * scale, &input.buf)?;

        Ok(())
    }
}

#[derive(Clone)]
pub struct SplatAcrossBatch<D: Device> {
    pub input: TensorRef<D>,
    pub output: TensorRef<D>,
    pub input_mul: f32,
    pub output_mul: f32,
    pub reduction: Reduce,
}

impl<D: Device> DeviceOperation<D> for SplatAcrossBatch<D> {
    fn opname(&self) -> String {
        format!("SplatAcrossBatch({:?})", self.reduction)
    }

    fn execute(&self) -> Result<(), OperationError<D::DeviceError>> {
        let input = self.input.dense();
        let mut output = self.output.dense_mut();

        if input.batch_size().is_some() || output.batch_size().is_none() {
            return Err(OperationError::MismatchedBatchSizes);
        }

        if input.single_size() != output.single_size() {
            return Err(OperationError::InvalidTensorFormat);
        }

        let bs = output.batch_size().unwrap_or(1);

        let scale = match self.reduction {
            Reduce::Avg => 1.0 / bs as f32,
            Reduce::Sum => 1.0,
        };

        output.buf.linear_comb_splat(input.single_size(), bs, self.output_mul, self.input_mul * scale, &input.buf)?;

        Ok(())
    }
}

#[derive(Clone)]
pub struct LinearCombination<D: Device> {
    pub input_mul: f32,
    pub output_mul: f32,
    pub input: TensorRef<D>,
    pub output: TensorRef<D>,
}

impl<D: Device> DeviceOperation<D> for LinearCombination<D> {
    fn opname(&self) -> String {
        "LinearCombination".to_string()
    }

    fn execute(&self) -> Result<(), OperationError<D::DeviceError>> {
        let input = self.input.dense();
        let mut output = self.output.dense_mut();

        if input.batch_size() != output.batch_size() {
            return Err(OperationError::MismatchedBatchSizes);
        }

        if input.single_size() != output.single_size() {
            return Err(OperationError::InvalidTensorFormat);
        }

        output.buf.linear_comb(input.size(), self.output_mul, self.input_mul, &input.buf)?;

        Ok(())
    }
}

#[derive(Clone)]
pub struct LinearCombinationSplat<D: Device> {
    pub input_mul: f32,
    pub output_mul: f32,
    pub input: TensorRef<D>,
    pub output: TensorRef<D>,
}

impl<D: Device> DeviceOperation<D> for LinearCombinationSplat<D> {
    fn opname(&self) -> String {
        "LinearCombinationSplat".to_string()
    }

    fn execute(&self) -> Result<(), OperationError<D::DeviceError>> {
        let input = self.input.dense();
        let mut output = self.output.dense_mut();

        if input.batch_size().is_some() || output.batch_size().is_none() {
            println!("{:?} {:?}", input.batch_size(), output.batch_size());
            return Err(OperationError::MismatchedBatchSizes);
        }

        if input.single_size() != output.single_size() {
            return Err(OperationError::InvalidTensorFormat);
        }

        let bs = output.batch_size().unwrap_or(1);
        output.buf.linear_comb_splat(input.size(), bs, self.output_mul, self.input_mul, &input.buf)?;

        Ok(())
    }
}

#[derive(Clone)]
pub struct SparseToDense<D: Device> {
    pub input: TensorRef<D>,
    pub output: TensorRef<D>,
}

impl<D: Device> DeviceOperation<D> for SparseToDense<D> {
    fn opname(&self) -> String {
        "SparseToDense".to_string()
    }

    fn execute(&self) -> Result<(), OperationError<D::DeviceError>> {
        let input = self.input.sparse();
        let mut output = self.output.dense_mut();

        if input.batch_size() != output.batch_size() {
            return Err(OperationError::MismatchedBatchSizes);
        }

        if input.single_size() != output.single_size() {
            return Err(OperationError::InvalidTensorFormat);
        }

        D::sparse_to_dense(input.batch_size().unwrap_or(1), input.single_size, input.nnz, &input.buf, &mut output.buf)
    }
}

#[derive(Clone)]
pub struct PairwiseMul<D: Device> {
    pub offset: usize,
    pub input: TensorRef<D>,
    pub output: TensorRef<D>,
}

impl<D: Device> DeviceOperation<D> for PairwiseMul<D> {
    fn opname(&self) -> String {
        "PairwiseMul".to_string()
    }

    fn execute(&self) -> Result<(), OperationError<D::DeviceError>> {
        let input = self.input.dense();
        let mut output = self.output.dense_mut();

        if input.batch_size() != output.batch_size() {
            return Err(OperationError::MismatchedBatchSizes);
        }

        if input.single_size() > 2 * output.single_size() {
            return Err(OperationError::InvalidTensorFormat);
        }

        let single_size = input.single_size();
        let stride = output.single_size();
        let batch_size = input.batch_size().unwrap_or(1);

        output.buf.pairwise_fwd(self.offset, stride, single_size, batch_size, &input.buf)?;

        Ok(())
    }
}

#[derive(Clone)]
pub struct Unary<D: Device> {
    pub input: TensorRef<D>,
    pub output: TensorRef<D>,
    pub op: UnaryOp,
}

impl<D: Device> DeviceOperation<D> for Unary<D> {
    fn opname(&self) -> String {
        format!("Unary({:?})", self.op)
    }

    fn execute(&self) -> Result<(), OperationError<D::DeviceError>> {
        let input = self.input.dense();
        let mut output = self.output.dense_mut();

        if input.batch_size() != output.batch_size() {
            return Err(OperationError::MismatchedBatchSizes);
        }

        if input.single_size() != output.single_size() {
            return Err(OperationError::InvalidTensorFormat);
        }

        let size = input.size();

        match self.op {
            UnaryOp::AbsPow(p) => output.buf.abs_pow_scalar(size, p, &input.buf)?,
            UnaryOp::Add(x) => output.buf.add_scalar(size, x, &input.buf)?,
            UnaryOp::Mul(x) => output.buf.linear_comb(size, 0.0, x, &input.buf)?,
            UnaryOp::DiffableFromOutput(act) => output.buf.diffable_from_output_fwd(size, &input.buf, act)?,
        }

        Ok(())
    }
}

#[derive(Clone)]
pub struct CopyOrAddStrided<D: Device> {
    pub input: TensorRef<D>,
    pub output: TensorRef<D>,
    pub input_offset: usize,
    pub output_offset: usize,
    pub add: bool,
    pub len_is_out: bool,
}

impl<D: Device> DeviceOperation<D> for CopyOrAddStrided<D> {
    fn opname(&self) -> String {
        "CopyOrAddStrided".to_string()
    }

    fn execute(&self) -> Result<(), OperationError<<D as Device>::DeviceError>> {
        let input = self.input.dense();
        let mut output = self.output.dense_mut();

        if input.batch_size() != output.batch_size() {
            return Err(OperationError::MismatchedBatchSizes);
        }

        let output_size = output.single_size();
        let rows = if self.len_is_out { output_size } else { input.single_size() };

        output.buf.copy_or_add_strided(
            self.add,
            rows,
            input.batch_size().unwrap_or(1),
            self.output_offset,
            output_size,
            &input.buf,
            self.input_offset,
            input.single_size(),
        )?;

        Ok(())
    }
}

#[derive(Clone)]
pub struct Softmax<D: Device> {
    pub input: TensorRef<D>,
    pub output: TensorRef<D>,
}

impl<D: CoreDeviceOps> DeviceOperation<D> for Softmax<D> {
    fn opname(&self) -> String {
        "Softmax".to_string()
    }

    fn execute(&self) -> Result<(), OperationError<<D as Device>::DeviceError>> {
        let input = self.input.dense();
        let mut output = self.output.dense_mut();

        let batch_size = input.batch_size();
        let single_size = input.single_size();

        if batch_size != output.batch_size() {
            return Err(OperationError::MismatchedBatchSizes);
        }

        if single_size != output.single_size() {
            return Err(OperationError::InvalidTensorFormat);
        }

        D::softmax_across_batch(batch_size.unwrap_or(1), single_size, &input.buf, &mut output.buf)
    }
}