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)
}
}