#![allow(missing_docs)]
use alloc::vec::Vec;
use crate::{
DType, Distribution, Shape, Slice, SliceOps, calculate_matmul_output,
ops::{
conv::{
calculate_conv_output_shape, calculate_conv_transpose_output_shape,
calculate_pool_output_shape,
},
unfold::calculate_unfold_shape,
},
quantization::QuantScheme,
tensor::IndexingUpdateOp,
};
use crate::graph::{ScalarIr, TensorId, TensorIr};
use super::operation::*;
impl CreationOpIr {
pub fn create(shape: Shape, dtype: DType, new_id: impl FnOnce() -> TensorId) -> Self {
let out = TensorIr::uninit(new_id(), shape, dtype);
CreationOpIr { out }
}
}
impl AvgPool3dOpIr {
pub fn create_with_output_size(x: TensorIr, kernel_size: [usize; 3], stride: [usize; 3],
padding: [usize; 3], count_include_pad: bool, ceil_mode: bool,
output_size: [usize; 3], new_id: impl FnOnce() -> TensorId) -> Self {
assert_eq!(x.shape.rank(), 5, "volume pooling input rank differs");
let shape = Shape::new([x.shape[0], x.shape[1], output_size[0], output_size[1], output_size[2]]);
let out = TensorIr::uninit(new_id(), shape, x.dtype);
Self { x, kernel_size, stride, padding, count_include_pad, ceil_mode, out }
}
}
impl MaxPool3dOpIr {
pub fn create(x: TensorIr, kernel_size: [usize; 3], stride: [usize; 3],
padding: [usize; 3], dilation: [usize; 3], ceil_mode: bool,
new_id: impl FnOnce() -> TensorId) -> Self {
let shape = calculate_pool_output_shape(&x.shape, &kernel_size, &stride,
&padding, &dilation, ceil_mode).unwrap();
let out = TensorIr::uninit(new_id(), shape, x.dtype);
Self { x, kernel_size, stride, padding, dilation, ceil_mode, out }
}
}
impl MaxPool3dWithIndicesOpIr {
pub fn create(x: TensorIr, kernel_size: [usize; 3], stride: [usize; 3],
padding: [usize; 3], dilation: [usize; 3], ceil_mode: bool,
mut new_id: impl FnMut() -> TensorId) -> Self {
let shape = calculate_pool_output_shape(&x.shape, &kernel_size, &stride,
&padding, &dilation, ceil_mode).unwrap();
let out = TensorIr::uninit(new_id(), shape.clone(), x.dtype);
let out_indices = TensorIr::uninit(new_id(), shape, DType::I64);
Self { x, kernel_size, stride, padding, dilation, ceil_mode, out, out_indices }
}
}
impl MaxPool3dWithIndicesBackwardOpIr {
pub fn create(x: TensorIr, grad: TensorIr, indices: TensorIr, kernel_size: [usize; 3],
stride: [usize; 3], padding: [usize; 3], dilation: [usize; 3], ceil_mode: bool,
new_id: impl FnOnce() -> TensorId) -> Self {
let out = TensorIr::uninit(new_id(), x.shape.clone(), x.dtype);
Self { x, grad, indices, kernel_size, stride, padding, dilation, ceil_mode, out }
}
}
impl InitOperationIr {
pub fn create(shape: Shape, dtype: DType, new_id: impl FnOnce() -> TensorId) -> Self {
let out = TensorIr::uninit(new_id(), shape, dtype);
InitOperationIr { out }
}
}
impl RandomOpIr {
pub fn create(
shape: Shape,
dtype: DType,
distribution: Distribution,
new_id: impl FnOnce() -> TensorId,
) -> Self {
let out = TensorIr::uninit(new_id(), shape, dtype);
RandomOpIr { out, distribution }
}
}
impl FullOpIr {
pub fn create(
shape: Shape,
dtype: DType,
value: ScalarIr,
new_id: impl FnOnce() -> TensorId,
) -> Self {
let out = TensorIr::uninit(new_id(), shape, dtype);
FullOpIr { out, value }
}
}
impl CastOpIr {
pub fn create(input: TensorIr, dtype: DType, new_id: impl FnOnce() -> TensorId) -> Self {
let out = TensorIr::uninit(new_id(), input.shape.clone(), dtype);
CastOpIr { input, out }
}
}
impl ShapeOpIr {
pub fn expand(input: TensorIr, shape: Shape, new_id: impl FnOnce() -> TensorId) -> Self {
let shape = input.shape.expand(shape).unwrap();
Self::create(input, shape, new_id)
}
pub fn reshape(input: TensorIr, shape: Shape, new_id: impl FnOnce() -> TensorId) -> Self {
let shape = input.shape.reshape(shape).unwrap();
Self::create(input, shape, new_id)
}
fn create(input: TensorIr, shape: Shape, new_id: impl FnOnce() -> TensorId) -> Self {
let out = TensorIr::uninit(new_id(), shape, input.dtype);
ShapeOpIr { input, out }
}
}
impl From<MatmulOpIr> for BinaryOpIr {
fn from(value: MatmulOpIr) -> Self {
Self {
lhs: value.lhs,
rhs: value.rhs,
out: value.out,
}
}
}
impl From<ReduceOpIr> for UnaryOpIr {
fn from(value: ReduceOpIr) -> Self {
Self {
input: value.input,
out: value.out,
}
}
}
#[derive(Debug)]
#[allow(missing_docs)]
pub enum IrError {
DTypeMismatch,
}
fn dtype_compat(lhs: &DType, rhs: &DType) -> bool {
let lhs_qfloat = matches!(lhs, DType::QFloat(_));
let rhs_qfloat = matches!(rhs, DType::QFloat(_));
if lhs_qfloat && (rhs_qfloat || rhs.is_float())
|| lhs.is_float() && (rhs_qfloat || rhs.is_float())
{
true
} else {
lhs == rhs
}
}
fn output_check<'a, I>(inputs: I, compat: impl Fn(&DType, &DType) -> bool) -> Result<DType, IrError>
where
I: IntoIterator<Item = &'a DType>,
{
let mut iter = inputs.into_iter();
let first = iter.next().unwrap();
for d in iter {
if !compat(first, d) {
return Err(IrError::DTypeMismatch);
}
}
Ok(*first)
}
fn output_dtype<'a, I: IntoIterator<Item = &'a DType>>(inputs: I) -> Result<DType, IrError> {
output_check(inputs, |a, b| a == b)
}
fn output_dtype_mixed<'a, I: IntoIterator<Item = &'a DType>>(inputs: I) -> Result<DType, IrError> {
output_check(inputs, dtype_compat)
}
macro_rules! impl_ir_create {
(@create_fn $op:ident { $( $field:ident : $ty:ty ),* $(,)? } , $shape:expr, $dtype:expr) => {
#[doc = "Create a new operation IR from the given inputs."]
#[doc = "`new_id` should generate a unique `TensorId` for the uninitialized output tensor."]
#[allow(clippy::too_many_arguments)]
pub fn create($( $field : $ty ),*, new_id: impl FnOnce() -> crate::graph::TensorId) -> $op {
let shape = $shape;
let dtype = $dtype;
let out = TensorIr::uninit(new_id(), shape, dtype);
$op { $( $field ),*, out }
}
};
(
$op:ident { $( $field:ident : $ty:ty ),* $(,)? },
shape = $shape:expr,
dtype = $dtype:expr
) => {
impl $op {
impl_ir_create!(@create_fn $op { $( $field : $ty ),* }, $shape, $dtype);
}
};
(
$op:ident { $( $field:ident : $ty:ty ),* $(,)? },
shape = $shape:expr,
dtype = $dtype:expr,
$fn_name:ident ( $extra:ident : $extra_ty:ty )
) => {
impl $op {
impl_ir_create!(@create_fn $op { $( $field : $ty ),* }, $shape, $dtype);
#[doc = "Create a new operation IR from the given inputs and the given output dtype."]
#[allow(clippy::too_many_arguments)]
pub fn $fn_name($( $field : $ty ),*, $extra: $extra_ty, new_id: impl FnOnce() -> crate::graph::TensorId) -> Self {
let shape = $shape;
let _ = $dtype; let out = TensorIr::uninit(new_id(), shape, $extra);
$op { $( $field ),*, out }
}
}
};
}
impl_ir_create!(
UnaryOpIr { input: TensorIr },
shape = input.shape.clone(),
dtype = input.dtype,
create_comparison(bool_dtype: DType)
);
impl_ir_create!(
BinaryOpIr {
lhs: TensorIr,
rhs: TensorIr
},
shape = lhs.shape.broadcast(&rhs.shape).unwrap(),
dtype = output_dtype([&lhs.dtype, &rhs.dtype]).unwrap(),
create_comparison(bool_dtype: DType)
);
impl_ir_create!(
ScalarOpIr {
lhs: TensorIr,
rhs: ScalarIr
},
shape = lhs.shape.clone(),
dtype = lhs.dtype,
create_comparison(bool_dtype: DType)
);
impl_ir_create!(
MatmulOpIr {
lhs: TensorIr,
rhs: TensorIr
},
shape = calculate_matmul_output(&lhs.shape, &rhs.shape).unwrap(),
dtype = output_dtype_mixed([&lhs.dtype, &rhs.dtype]).unwrap(),
create_mixed(out_dtype: DType)
);
impl_ir_create!(
SwapDimsOpIr {
input: TensorIr,
dim1: usize,
dim2: usize
},
shape = input.shape.clone().swapped(dim1, dim2).unwrap(),
dtype = input.dtype
);
impl_ir_create!(
PermuteOpIr { input: TensorIr, axes: Vec<usize> },
shape = input.shape.clone().permuted(&axes).unwrap(),
dtype = input.dtype
);
impl_ir_create!(
RepeatDimOpIr {
tensor: TensorIr,
dim: usize,
times: usize
},
shape = tensor.shape.clone().repeat(dim, times).unwrap(),
dtype = tensor.dtype
);
impl_ir_create!(
FlipOpIr { input: TensorIr, axes: Vec<usize> },
shape = input.shape.clone(), dtype = input.dtype
);
impl_ir_create!(
CatOpIr { tensors: Vec<TensorIr>, dim: usize },
shape = Shape::cat(tensors.iter().map(|t| &t.shape), dim).unwrap(),
dtype = output_dtype(tensors.iter().map(|t| &t.dtype)).unwrap()
);
#[cfg(feature = "graph-distributed")]
impl_ir_create!(
AllReduceOpIr { tensor: TensorIr, op: crate::distributed::ReduceOperation, device_ids: Vec<(u16, u16)> },
shape = tensor.shape.clone(),
dtype = tensor.dtype
);
impl_ir_create!(
GatherOpIr {
tensor: TensorIr,
dim: usize,
indices: TensorIr
},
shape = indices.shape.clone(), dtype = tensor.dtype
);
impl_ir_create!(
ScatterOpIr {
tensor: TensorIr,
dim: usize,
indices: TensorIr,
value: TensorIr,
update: IndexingUpdateOp
},
shape = tensor.shape.clone(), dtype = output_dtype([&tensor.dtype, &value.dtype]).unwrap()
);
impl_ir_create!(
ScatterNdOpIr {
data: TensorIr,
indices: TensorIr,
values: TensorIr,
reduction: IndexingUpdateOp
},
shape = data.shape.clone(),
dtype = output_dtype([&data.dtype, &values.dtype]).unwrap()
);
impl GatherNdOpIr {
pub fn create(
data: TensorIr,
indices: TensorIr,
new_id: impl FnOnce() -> crate::graph::TensorId,
) -> Self {
let m = indices.shape.num_dims();
let k = indices.shape[m - 1];
let mut dims = indices.shape.as_slice()[..m - 1].to_vec();
dims.extend_from_slice(&data.shape.as_slice()[k..]);
let shape = Shape::from(dims);
let dtype = data.dtype;
let out = TensorIr::uninit(new_id(), shape, dtype);
GatherNdOpIr { data, indices, out }
}
}
impl_ir_create!(
ReduceOpIr { input: TensorIr },
shape = [1].into(),
dtype = input.dtype
);
fn reduce_output_shape(mut output_shape: Shape, axis: usize, accumulator_len: usize) -> Shape {
assert!(output_shape.rank() > axis);
output_shape[axis] = accumulator_len;
output_shape
}
impl_ir_create!(
ReduceDimOpIr {
input: TensorIr,
axis: usize,
accumulator_len: usize,
},
shape = reduce_output_shape(input.shape.clone(), axis, accumulator_len),
dtype = input.dtype,
create_arg(ind_dtype: DType)
);
impl_ir_create!(
DimOpIr {
input: TensorIr,
axis: usize
},
shape = input.shape.clone(), dtype = input.dtype
);
impl_ir_create!(
SelectOpIr {
tensor: TensorIr,
dim: usize,
indices: TensorIr
},
shape = {
let mut s = tensor.shape.clone();
s[dim] = indices.shape[0];
s
},
dtype = tensor.dtype
);
impl_ir_create!(
SelectAssignOpIr {
tensor: TensorIr,
dim: usize,
indices: TensorIr,
value: TensorIr,
update: IndexingUpdateOp
},
shape = tensor.shape.clone(),
dtype = output_dtype([&tensor.dtype, &value.dtype]).unwrap()
);
impl_ir_create!(
SliceOpIr {
tensor: TensorIr,
ranges: Vec<Slice>,
},
shape = tensor.shape.clone().slice(&ranges).unwrap(),
dtype = tensor.dtype
);
impl_ir_create!(
SliceAssignOpIr {
tensor: TensorIr,
ranges: Vec<Slice>,
value: TensorIr
},
shape = tensor.shape.clone(),
dtype = output_dtype([&tensor.dtype, &value.dtype]).unwrap()
);
impl_ir_create!(
MaskWhereOpIr {
tensor: TensorIr,
mask: TensorIr,
value: TensorIr
},
shape = Shape::broadcast_many([&tensor.shape, &mask.shape, &value.shape]).unwrap(),
dtype = output_dtype([&tensor.dtype, &value.dtype]).unwrap()
);
impl_ir_create!(
MaskFillOpIr {
tensor: TensorIr,
mask: TensorIr,
value: ScalarIr
},
shape = tensor.shape.broadcast(&mask.shape).unwrap(),
dtype = tensor.dtype
);
impl_ir_create!(
ClampOpIr {
tensor: TensorIr,
min: ScalarIr,
max: ScalarIr
},
shape = tensor.shape.clone(),
dtype = tensor.dtype
);
impl_ir_create!(
AvgPool1dOpIr {
x: TensorIr,
kernel_size: usize,
stride: usize,
padding: usize,
count_include_pad: bool,
ceil_mode: bool
},
shape = calculate_pool_output_shape(
&x.shape,
&[kernel_size],
&[stride],
&[padding],
&[1],
ceil_mode
)
.unwrap(),
dtype = x.dtype
);
impl_ir_create!(
AvgPool1dBackwardOpIr {
x: TensorIr,
grad: TensorIr,
kernel_size: usize,
stride: usize,
padding: usize,
count_include_pad: bool,
ceil_mode: bool
},
shape = x.shape.clone(),
dtype = x.dtype
);
impl_ir_create!(
AvgPool2dOpIr {
x: TensorIr,
kernel_size: [usize; 2],
stride: [usize; 2],
padding: [usize; 2],
count_include_pad: bool,
ceil_mode: bool
},
shape = calculate_pool_output_shape(
&x.shape,
&kernel_size,
&stride,
&padding,
&[1, 1],
ceil_mode
)
.unwrap(),
dtype = x.dtype
);
impl_ir_create!(
AvgPool2dBackwardOpIr {
x: TensorIr,
grad: TensorIr,
kernel_size: [usize; 2],
stride: [usize; 2],
padding: [usize; 2],
count_include_pad: bool,
ceil_mode: bool
},
shape = x.shape.clone(),
dtype = x.dtype
);
impl_ir_create!(
MaxPool1dOpIr {
x: TensorIr,
kernel_size: usize,
stride: usize,
padding: usize,
dilation: usize,
ceil_mode: bool
},
shape = calculate_pool_output_shape(
&x.shape,
&[kernel_size],
&[stride],
&[padding],
&[dilation],
ceil_mode
)
.unwrap(),
dtype = x.dtype
);
impl_ir_create!(
MaxPool2dOpIr {
x: TensorIr,
kernel_size: [usize; 2],
stride: [usize; 2],
padding: [usize; 2],
dilation: [usize; 2],
ceil_mode: bool
},
shape = calculate_pool_output_shape(
&x.shape,
&kernel_size,
&stride,
&padding,
&dilation,
ceil_mode
)
.unwrap(),
dtype = x.dtype
);
impl_ir_create!(
MaxPool1dWithIndicesBackwardOpIr {
x: TensorIr,
grad: TensorIr,
indices: TensorIr,
kernel_size: usize,
stride: usize,
padding: usize,
dilation: usize,
ceil_mode: bool
},
shape = x.shape.clone(),
dtype = x.dtype
);
impl_ir_create!(
MaxPool2dWithIndicesBackwardOpIr {
x: TensorIr,
grad: TensorIr,
indices: TensorIr,
kernel_size: [usize; 2],
stride: [usize; 2],
padding: [usize; 2],
dilation: [usize; 2],
ceil_mode: bool
},
shape = x.shape.clone(),
dtype = x.dtype
);
impl_ir_create!(
AdaptiveAvgPool1dOpIr {
x: TensorIr,
output_size: usize
},
shape = Shape::new([x.shape[0], x.shape[1], output_size]),
dtype = x.dtype
);
impl_ir_create!(
AdaptiveAvgPool2dOpIr {
x: TensorIr,
output_size: [usize; 2]
},
shape = Shape::new([x.shape[0], x.shape[1], output_size[0], output_size[1]]),
dtype = x.dtype
);
impl_ir_create!(
AdaptiveAvgPool1dBackwardOpIr {
x: TensorIr,
grad: TensorIr,
},
shape = x.shape.clone(),
dtype = x.dtype
);
impl_ir_create!(
AdaptiveAvgPool2dBackwardOpIr {
x: TensorIr,
grad: TensorIr,
},
shape = x.shape.clone(),
dtype = x.dtype
);
impl_ir_create!(
AdaptiveAvgPool3dOpIr {
x: TensorIr,
output_size: [usize; 3]
},
shape = Shape::new([x.shape[0], x.shape[1], output_size[0], output_size[1], output_size[2]]),
dtype = x.dtype
);
impl_ir_create!(
AvgPool3dBackwardOpIr {
x: TensorIr,
grad: TensorIr,
kernel_size: [usize; 3],
stride: [usize; 3],
padding: [usize; 3],
count_include_pad: bool,
ceil_mode: bool,
},
shape = x.shape.clone(),
dtype = x.dtype
);
impl_ir_create!(
AdaptiveAvgPool3dBackwardOpIr {
x: TensorIr,
grad: TensorIr,
},
shape = x.shape.clone(),
dtype = x.dtype
);
impl_ir_create!(
InterpolateOpIr {
x: TensorIr,
output_size: [usize; 2],
options: InterpolateOptionsIr
},
shape = Shape::new([x.shape[0], x.shape[1], output_size[0], output_size[1]]),
dtype = x.dtype
);
impl_ir_create!(
InterpolateBackwardOpIr {
x: TensorIr,
grad: TensorIr,
output_size: [usize; 2],
options: InterpolateOptionsIr
},
shape = x.shape.clone(),
dtype = x.dtype
);
impl_ir_create!(
Interpolate1dOpIr { x: TensorIr, output_size: usize, options: InterpolateOptionsIr },
shape = Shape::new([x.shape[0], x.shape[1], output_size]), dtype = x.dtype
);
impl_ir_create!(
Interpolate1dBackwardOpIr { x: TensorIr, grad: TensorIr, output_size: usize, options: InterpolateOptionsIr },
shape = x.shape.clone(), dtype = x.dtype
);
impl_ir_create!(
Interpolate3dOpIr { x: TensorIr, output_size: [usize; 3], options: InterpolateOptionsIr },
shape = Shape::new([x.shape[0], x.shape[1], output_size[0], output_size[1], output_size[2]]), dtype = x.dtype
);
impl_ir_create!(
Interpolate3dBackwardOpIr { x: TensorIr, grad: TensorIr, output_size: [usize; 3], options: InterpolateOptionsIr },
shape = x.shape.clone(), dtype = x.dtype
);
impl_ir_create!(
GridSample2dOpIr {
tensor: TensorIr,
grid: TensorIr,
options: GridSampleOptionsIr
},
shape = Shape::new([
tensor.shape[0],
tensor.shape[1],
grid.shape[1],
grid.shape[2]
]),
dtype = tensor.dtype
);
impl_ir_create!(
LinearOpIr {
x: TensorIr,
weight: TensorIr,
bias: Option<TensorIr>
},
shape = {
let n = x.shape.num_dims();
let mut dims: Vec<usize> = (0..n).map(|i| x.shape[i]).collect();
dims[n - 1] = weight.shape[1];
Shape::from(dims)
},
dtype = output_dtype(
[
Some(&x.dtype),
Some(&weight.dtype),
bias.as_ref().map(|b| &b.dtype),
]
.iter()
.filter_map(|&d| d),
)
.unwrap()
);
impl_ir_create!(
LinearXBackwardOpIr {
weight: TensorIr,
output_grad: TensorIr,
},
shape = {
let n = output_grad.shape.num_dims();
let mut dims: Vec<usize> = (0..n).map(|i| output_grad.shape[i]).collect();
dims[n - 1] = weight.shape[0];
Shape::from(dims)
},
dtype = output_grad.dtype
);
impl_ir_create!(
LinearWeightBackwardOpIr {
x: TensorIr,
output_grad: TensorIr,
},
shape = {
let d_input = x.shape[x.shape.num_dims() - 1];
let d_output = output_grad.shape[output_grad.shape.num_dims() - 1];
Shape::from(alloc::vec![d_input, d_output])
},
dtype = output_grad.dtype
);
impl_ir_create!(
LinearBiasBackwardOpIr {
output_grad: TensorIr,
},
shape = {
let d_output = output_grad.shape[output_grad.shape.num_dims() - 1];
Shape::from(alloc::vec![d_output])
},
dtype = output_grad.dtype
);
impl_ir_create!(
Conv1dOpIr {
x: TensorIr,
weight: TensorIr,
bias: Option<TensorIr>,
options: Conv1dOptionsIr
},
shape = calculate_conv_output_shape(
&x.shape,
&weight.shape,
&options.stride,
&options.padding,
&options.dilation,
)
.unwrap(),
dtype = output_dtype(
[
Some(&x.dtype),
Some(&weight.dtype),
bias.as_ref().map(|b| &b.dtype),
]
.iter()
.filter_map(|&d| d),
)
.unwrap()
);
impl_ir_create!(
Conv1dXBackwardOpIr {
x: TensorIr,
weight: TensorIr,
output_grad: TensorIr,
options: Conv1dOptionsIr
},
shape = x.shape.clone(),
dtype = output_grad.dtype
);
impl_ir_create!(
Conv1dWeightBackwardOpIr {
x: TensorIr,
weight: TensorIr,
output_grad: TensorIr,
options: Conv1dOptionsIr
},
shape = weight.shape.clone(),
dtype = output_grad.dtype
);
impl_ir_create!(
Conv1dBiasBackwardOpIr {
x: TensorIr,
bias: TensorIr,
output_grad: TensorIr,
},
shape = bias.shape.clone(),
dtype = output_grad.dtype
);
impl_ir_create!(
Conv2dOpIr {
x: TensorIr,
weight: TensorIr,
bias: Option<TensorIr>,
options: Conv2dOptionsIr
},
shape = calculate_conv_output_shape(
&x.shape,
&weight.shape,
&options.stride,
&options.padding,
&options.dilation,
)
.unwrap(),
dtype = output_dtype(
[
Some(&x.dtype),
Some(&weight.dtype),
bias.as_ref().map(|b| &b.dtype),
]
.iter()
.filter_map(|&d| d),
)
.unwrap()
);
impl_ir_create!(
Conv2dXBackwardOpIr {
x: TensorIr,
weight: TensorIr,
output_grad: TensorIr,
options: Conv2dOptionsIr
},
shape = x.shape.clone(),
dtype = output_grad.dtype
);
impl_ir_create!(
Conv2dWeightBackwardOpIr {
x: TensorIr,
weight: TensorIr,
output_grad: TensorIr,
options: Conv2dOptionsIr
},
shape = weight.shape.clone(),
dtype = output_grad.dtype
);
impl_ir_create!(
Conv2dBiasBackwardOpIr {
x: TensorIr,
bias: TensorIr,
output_grad: TensorIr,
},
shape = bias.shape.clone(),
dtype = output_grad.dtype
);
impl_ir_create!(
Conv3dOpIr {
x: TensorIr,
weight: TensorIr,
bias: Option<TensorIr>,
options: Conv3dOptionsIr
},
shape = calculate_conv_output_shape(
&x.shape,
&weight.shape,
&options.stride,
&options.padding,
&options.dilation,
)
.unwrap(),
dtype = output_dtype(
[
Some(&x.dtype),
Some(&weight.dtype),
bias.as_ref().map(|b| &b.dtype),
]
.iter()
.filter_map(|&d| d),
)
.unwrap()
);
impl_ir_create!(
Conv3dXBackwardOpIr {
x: TensorIr,
weight: TensorIr,
output_grad: TensorIr,
options: Conv3dOptionsIr
},
shape = x.shape.clone(),
dtype = output_grad.dtype
);
impl_ir_create!(
Conv3dWeightBackwardOpIr {
x: TensorIr,
weight: TensorIr,
output_grad: TensorIr,
options: Conv3dOptionsIr
},
shape = weight.shape.clone(),
dtype = output_grad.dtype
);
impl_ir_create!(
Conv3dBiasBackwardOpIr {
x: TensorIr,
bias: TensorIr,
output_grad: TensorIr,
},
shape = bias.shape.clone(),
dtype = output_grad.dtype
);
impl_ir_create!(
DeformConv2dOpIr {
x: TensorIr,
offset: TensorIr,
weight: TensorIr,
mask: Option<TensorIr>,
bias: Option<TensorIr>,
options: DeformableConv2dOptionsIr
},
shape = {
assert_eq!(x.shape.rank(), 4, "deform_conv2d input must have rank 4");
assert_eq!(weight.shape.rank(), 4, "deform_conv2d weight must have rank 4");
let options: crate::ops::DeformConvOptions<2> = options.clone().into();
let [height, width] = options.output_size(
[x.shape[2], x.shape[3]], [weight.shape[2], weight.shape[3]],
);
Shape::new([x.shape[0], weight.shape[0], height, width])
},
dtype = output_dtype(
[
Some(&x.dtype),
Some(&offset.dtype),
Some(&weight.dtype),
mask.as_ref().map(|m| &m.dtype),
bias.as_ref().map(|b| &b.dtype),
]
.iter()
.filter_map(|&d| d),
)
.unwrap()
);
impl_ir_create!(
ConvTranspose1dOpIr {
x: TensorIr,
weight: TensorIr,
bias: Option<TensorIr>,
options: ConvTranspose1dOptionsIr
},
shape = calculate_conv_transpose_output_shape(
&x.shape,
&weight.shape,
&options.stride,
&options.padding,
&options.padding_out,
&options.dilation,
options.groups,
)
.unwrap(),
dtype = output_dtype(
[
Some(&x.dtype),
Some(&weight.dtype),
bias.as_ref().map(|b| &b.dtype),
]
.iter()
.filter_map(|&d| d),
)
.unwrap()
);
impl_ir_create!(
ConvTranspose2dOpIr {
x: TensorIr,
weight: TensorIr,
bias: Option<TensorIr>,
options: ConvTranspose2dOptionsIr
},
shape = calculate_conv_transpose_output_shape(
&x.shape,
&weight.shape,
&options.stride,
&options.padding,
&options.padding_out,
&options.dilation,
options.groups,
)
.unwrap(),
dtype = output_dtype(
[
Some(&x.dtype),
Some(&weight.dtype),
bias.as_ref().map(|b| &b.dtype),
]
.iter()
.filter_map(|&d| d),
)
.unwrap()
);
impl_ir_create!(
ConvTranspose3dOpIr {
x: TensorIr,
weight: TensorIr,
bias: Option<TensorIr>,
options: ConvTranspose3dOptionsIr
},
shape = calculate_conv_transpose_output_shape(
&x.shape,
&weight.shape,
&options.stride,
&options.padding,
&options.padding_out,
&options.dilation,
options.groups,
)
.unwrap(),
dtype = output_dtype(
[
Some(&x.dtype),
Some(&weight.dtype),
bias.as_ref().map(|b| &b.dtype),
]
.iter()
.filter_map(|&d| d),
)
.unwrap()
);
impl_ir_create!(
UnfoldOpIr {
input: TensorIr,
dim: usize,
size: usize,
step: usize
},
shape = calculate_unfold_shape(input.shape.clone(), dim, size, step),
dtype = input.dtype
);
impl_ir_create!(
CrossOpIr {
lhs: TensorIr,
rhs: TensorIr,
dim: usize
},
shape = lhs.shape.broadcast(&rhs.shape).unwrap(),
dtype = output_dtype([&lhs.dtype, &rhs.dtype]).unwrap()
);
impl_ir_create!(
QuantizeOpIr {
tensor: TensorIr,
qparams: QuantizationParametersIr,
scheme: QuantScheme
},
shape = tensor.shape.clone(),
dtype = DType::QFloat(scheme)
);
impl_ir_create!(
AttentionOpIr {
query: TensorIr,
key: TensorIr,
value: TensorIr,
mask: Option<TensorIr>,
attn_bias: Option<TensorIr>,
options: AttentionOptionsIr,
},
shape = Shape::new([query.shape[0], query.shape[1], query.shape[2], value.shape[3]]),
dtype = query.dtype
);
impl_ir_create!(
CtcLossOpIr {
log_probs: TensorIr,
targets: TensorIr,
input_lengths: TensorIr,
target_lengths: TensorIr,
blank: usize,
},
shape = Shape::new([log_probs.shape[1]]),
dtype = log_probs.dtype
);
impl_ir_create!(
CtcLossBackwardOpIr {
log_probs: TensorIr,
targets: TensorIr,
input_lengths: TensorIr,
target_lengths: TensorIr,
grad_loss: TensorIr,
blank: usize,
},
shape = log_probs.shape.clone(),
dtype = log_probs.dtype
);
impl DequantizeOpIr {
pub fn create(input: TensorIr, dtype: DType, new_id: impl FnOnce() -> TensorId) -> Self {
let out = TensorIr::uninit(new_id(), input.shape.clone(), dtype);
DequantizeOpIr { input, out }
}
}
impl ExponentialReluOpIr {
pub fn create(x: TensorIr, alpha: f64, continuous: bool, new_id: impl FnOnce() -> TensorId) -> Self {
let out = TensorIr::uninit(new_id(), x.shape.clone(), x.dtype);
Self { x, alpha: ScalarIr::Float(alpha), continuous, out }
}
}
impl ExponentialReluBackwardOpIr {
pub fn create(x: TensorIr, grad: TensorIr, alpha: f64, continuous: bool, new_id: impl FnOnce() -> TensorId) -> Self {
assert_eq!(x.shape, grad.shape, "ELU/CELU gradient shape differs");
let out = TensorIr::uninit(new_id(), x.shape.clone(), x.dtype);
Self { x, grad, alpha: ScalarIr::Float(alpha), continuous, out }
}
}
impl LeakyReluOpIr {
pub fn create(x: TensorIr, negative_slope: f64, new_id: impl FnOnce() -> TensorId) -> Self {
let out = TensorIr::uninit(new_id(), x.shape.clone(), x.dtype);
Self { x, negative_slope: ScalarIr::Float(negative_slope), out }
}
}
impl LeakyReluBackwardOpIr {
pub fn create(x: TensorIr, grad: TensorIr, negative_slope: f64, new_id: impl FnOnce() -> TensorId) -> Self {
assert_eq!(x.shape, grad.shape, "LeakyReLU gradient shape differs");
let out = TensorIr::uninit(new_id(), x.shape.clone(), x.dtype);
Self { x, grad, negative_slope: ScalarIr::Float(negative_slope), out }
}
}
impl PreluOpIr {
pub fn create(x: TensorIr, alpha: TensorIr, new_id: impl FnOnce() -> TensorId) -> Self {
crate::ops::prelu_training::geometry(&x.shape, &alpha.shape);
let out = TensorIr::uninit(new_id(), x.shape.clone(), x.dtype);
Self { x, alpha, out }
}
}
impl PreluBackwardSelectOpIr {
pub fn create(x: TensorIr, alpha: TensorIr, grad: TensorIr, mask: [bool; 2], mut new_id: impl FnMut() -> TensorId) -> Self {
crate::ops::prelu_training::geometry(&x.shape, &alpha.shape);
assert_eq!(x.shape, grad.shape, "PReLU gradient shape differs");
let input_grad = mask[0].then(|| TensorIr::uninit(new_id(), x.shape.clone(), x.dtype));
let weight_grad = mask[1].then(|| TensorIr::uninit(new_id(), alpha.shape.clone(), alpha.dtype));
Self { x, alpha, grad, input_grad, weight_grad }
}
}
impl GroupNormOpIr {
pub fn create(x: TensorIr, gamma: Option<TensorIr>, beta: Option<TensorIr>, groups: usize, epsilon: f64,
mut new_id: impl FnMut() -> TensorId) -> Self {
let info = crate::ops::group_normalization::geometry(&x.shape, groups);
for value in gamma.iter().chain(beta.iter()) {
assert_eq!(value.shape, Shape::new([info.channels]), "GroupNorm affine shape differs");
}
let dtype = if x.dtype == DType::F64 { DType::F64 } else { DType::F32 };
let out = TensorIr::uninit(new_id(), x.shape.clone(), x.dtype);
let mean = TensorIr::uninit(new_id(), Shape::new([info.batch, groups]), dtype);
let rstd = TensorIr::uninit(new_id(), mean.shape.clone(), dtype);
Self { x, gamma, beta, groups, epsilon: ScalarIr::Float(epsilon), out, mean, rstd }
}
}
impl GroupNormBackwardSelectOpIr {
pub fn create(x: TensorIr, gamma: Option<TensorIr>, grad: TensorIr, mean: TensorIr, rstd: TensorIr,
groups: usize, mask: [bool; 3], mut new_id: impl FnMut() -> TensorId) -> Self {
let info = crate::ops::group_normalization::geometry(&x.shape, groups);
if let Some(gamma) = &gamma { assert_eq!(gamma.shape, Shape::new([info.channels]), "GroupNorm weight shape differs"); }
assert!(!mask[1] || gamma.is_some(), "GroupNorm weight gradient requires an actual weight");
assert_eq!(grad.shape, x.shape, "GroupNorm gradient shape differs");
assert_eq!(mean.shape, Shape::new([info.batch, groups]), "GroupNorm mean shape differs");
assert_eq!(rstd.shape, mean.shape, "GroupNorm reciprocal deviation shape differs");
let input_grad = mask[0].then(|| TensorIr::uninit(new_id(), x.shape.clone(), x.dtype));
let weight_grad = mask[1].then(|| {
let gamma = gamma.as_ref().expect("requested GroupNorm weight");
TensorIr::uninit(new_id(), gamma.shape.clone(), gamma.dtype)
});
let bias_dtype = if [&x, &grad, &mean, &rstd].into_iter().chain(gamma.iter()).any(|value| value.dtype == DType::F64) {
DType::F64
} else { DType::F32 };
let bias_grad = mask[2].then(|| TensorIr::uninit(new_id(), Shape::new([info.channels]), bias_dtype));
Self { x, gamma, grad, mean, rstd, groups, input_grad, weight_grad, bias_grad }
}
}
impl GeluOpIr {
pub fn create(x: TensorIr, approximate: bool, new_id: impl FnOnce() -> TensorId) -> Self {
let out = TensorIr::uninit(new_id(), x.shape.clone(), x.dtype);
Self { x, approximate, out }
}
}
impl GeluBackwardOpIr {
pub fn create(x: TensorIr, grad: TensorIr, approximate: bool, new_id: impl FnOnce() -> TensorId) -> Self {
assert_eq!(x.shape, grad.shape, "GELU gradient shape differs");
let out = TensorIr::uninit(new_id(), x.shape.clone(), x.dtype);
Self { x, grad, approximate, out }
}
}
impl SiluBackwardOpIr {
pub fn create(x: TensorIr, grad: TensorIr, new_id: impl FnOnce() -> TensorId) -> Self {
assert_eq!(x.shape, grad.shape, "SiLU gradient shape differs");
let out = TensorIr::uninit(new_id(), x.shape.clone(), x.dtype);
Self { x, grad, out }
}
}
impl SoftmaxOpIr {
pub fn create(x: TensorIr, dim: usize, logarithmic: bool, mut new_id: impl FnMut() -> TensorId) -> Self {
assert!(dim < x.shape.num_dims(), "softmax axis out of bounds");
assert!(x.shape[dim] > 0, "softmax axis must be nonempty");
let dtype = if x.dtype == DType::F64 { DType::F64 } else { DType::F32 };
let out = TensorIr::uninit(new_id(), x.shape.clone(), x.dtype);
let working = TensorIr::uninit(new_id(), x.shape.clone(), dtype);
Self { x, dim, logarithmic, out, working }
}
}
impl SoftmaxBackwardOpIr {
pub fn create(working: TensorIr, grad: TensorIr, dim: usize, logarithmic: bool,
mut new_id: impl FnMut() -> TensorId) -> Self {
assert!(dim < working.shape.num_dims(), "softmax backward axis out of bounds");
assert!(working.shape[dim] > 0, "softmax backward axis must be nonempty");
assert_eq!(grad.shape, working.shape, "softmax gradient shape differs");
assert!(matches!(working.dtype, DType::F32 | DType::F64), "softmax saved output must use working storage");
let dtype = if working.dtype == DType::F64 || grad.dtype == DType::F64 { DType::F64 } else { DType::F32 };
let out = TensorIr::uninit(new_id(), working.shape.clone(), dtype);
Self { working, grad, dim, logarithmic, out }
}
}
impl RmsNormBackwardSelectOpIr {
pub fn create(x: TensorIr, gamma: TensorIr, grad: TensorIr, rstd: TensorIr, mask: [bool; 2],
mut new_id: impl FnMut() -> TensorId) -> Self {
let width = *x.shape.last().expect("RMSNorm requires an axis");
assert!(width > 0, "RMSNorm final axis must be nonempty");
assert_eq!(gamma.shape, Shape::new([width]), "RMSNorm weight shape differs");
assert_eq!(grad.shape, x.shape, "RMSNorm gradient shape differs");
assert_eq!(rstd.shape, Shape::new([x.shape.num_elements() / width]), "RMSNorm reciprocal norm shape differs");
let input_grad = mask[0].then(|| TensorIr::uninit(new_id(), x.shape.clone(), x.dtype));
let weight_grad = mask[1].then(|| TensorIr::uninit(new_id(), gamma.shape.clone(), gamma.dtype));
Self { x, gamma, grad, rstd, input_grad, weight_grad }
}
}
impl LayerNormBackwardSelectOpIr {
pub fn create(x: TensorIr, gamma: TensorIr, grad: TensorIr, mean: TensorIr, rstd: TensorIr, mask: [bool; 3],
mut new_id: impl FnMut() -> TensorId) -> Self {
let width = *x.shape.last().expect("LayerNorm input must have an axis");
assert!(width > 0, "LayerNorm final axis must be nonempty");
assert_eq!(gamma.shape, Shape::new([width]), "LayerNorm weight shape differs");
assert_eq!(grad.shape, x.shape, "LayerNorm gradient shape differs");
assert_eq!(mean.shape, Shape::new([x.shape.num_elements() / width]), "LayerNorm mean shape differs");
assert_eq!(rstd.shape, mean.shape, "LayerNorm reciprocal deviation shape differs");
let input_grad = mask[0].then(|| TensorIr::uninit(new_id(), x.shape.clone(), x.dtype));
let weight_grad = mask[1].then(|| TensorIr::uninit(new_id(), gamma.shape.clone(), gamma.dtype));
let bias_dtype = if [&x, &gamma, &grad, &mean, &rstd].iter().any(|value| value.dtype == DType::F64) {
DType::F64
} else { DType::F32 };
let bias_grad = mask[2].then(|| TensorIr::uninit(new_id(), gamma.shape.clone(), bias_dtype));
Self { x, gamma, grad, mean, rstd, input_grad, weight_grad, bias_grad }
}
}
impl RmsNormOpIr {
pub fn create(x: TensorIr, gamma: TensorIr, epsilon: f64, mut new_id: impl FnMut() -> TensorId) -> Self {
let width = *x.shape.last().expect("RMSNorm requires an axis");
assert!(width > 0, "RMSNorm final axis must be nonempty");
assert_eq!(gamma.shape, Shape::new([width]), "RMSNorm weight shape differs");
assert!(epsilon.is_finite() && epsilon > 0.0, "RMSNorm epsilon must be finite and positive");
let dtype = if x.dtype == DType::F64 { DType::F64 } else { DType::F32 };
let rstd = TensorIr::uninit(new_id(), Shape::new([x.shape.num_elements() / width]), dtype);
let out = TensorIr::uninit(new_id(), x.shape.clone(), x.dtype);
Self { x, gamma, epsilon: ScalarIr::Float(epsilon), out, rstd }
}
}
impl RmsNormBackwardOpIr {
pub fn create(x: TensorIr, gamma: TensorIr, grad: TensorIr, rstd: TensorIr,
mut new_id: impl FnMut() -> TensorId) -> Self {
let width = *x.shape.last().expect("RMSNorm requires an axis");
assert!(width > 0, "RMSNorm final axis must be nonempty");
assert_eq!(gamma.shape, Shape::new([width]), "RMSNorm weight shape differs");
assert_eq!(grad.shape, x.shape, "RMSNorm gradient shape differs");
assert_eq!(rstd.shape, Shape::new([x.shape.num_elements() / width]), "RMSNorm reciprocal norm shape differs");
let input_grad = TensorIr::uninit(new_id(), x.shape.clone(), x.dtype);
let weight_grad = TensorIr::uninit(new_id(), gamma.shape.clone(), gamma.dtype);
Self { x, gamma, grad, rstd, input_grad, weight_grad }
}
}
impl LayerNormOpIr {
pub fn create(x: TensorIr, gamma: TensorIr, beta: Option<TensorIr>, epsilon: f64,
mut new_id: impl FnMut() -> TensorId) -> Self {
let width = *x.shape.last().expect("LayerNorm input must have an axis");
assert!(width > 0, "LayerNorm final axis must be nonempty");
assert_eq!(gamma.shape, Shape::new([width]), "LayerNorm weight shape differs");
if let Some(beta) = &beta { assert_eq!(beta.shape, gamma.shape, "LayerNorm bias shape differs"); }
assert!(epsilon.is_finite() && epsilon > 0.0, "LayerNorm epsilon must be finite and positive");
let stats_shape = Shape::new([x.shape.num_elements() / width]);
let stats_dtype = if x.dtype == DType::F64 { DType::F64 } else { DType::F32 };
let out = TensorIr::uninit(new_id(), x.shape.clone(), x.dtype);
let mean = TensorIr::uninit(new_id(), stats_shape.clone(), stats_dtype);
let rstd = TensorIr::uninit(new_id(), stats_shape, stats_dtype);
Self { x, gamma, beta, epsilon: ScalarIr::Float(epsilon), out, mean, rstd }
}
}
impl LayerNormBackwardOpIr {
pub fn create(x: TensorIr, gamma: TensorIr, grad: TensorIr, mean: TensorIr, rstd: TensorIr,
mut new_id: impl FnMut() -> TensorId) -> Self {
let width = *x.shape.last().expect("LayerNorm input must have an axis");
assert!(width > 0, "LayerNorm final axis must be nonempty");
assert_eq!(gamma.shape, Shape::new([width]), "LayerNorm weight shape differs");
assert_eq!(grad.shape, x.shape, "LayerNorm gradient shape differs");
assert_eq!(mean.shape, Shape::new([x.shape.num_elements() / width]), "LayerNorm mean shape differs");
assert_eq!(rstd.shape, mean.shape, "LayerNorm reciprocal deviation shape differs");
let input_grad = TensorIr::uninit(new_id(), x.shape.clone(), x.dtype);
let weight_grad = TensorIr::uninit(new_id(), gamma.shape.clone(), gamma.dtype);
let bias_dtype = if [&x, &gamma, &grad, &mean, &rstd].iter().any(|value| value.dtype == DType::F64) {
DType::F64
} else { DType::F32 };
let bias_grad = TensorIr::uninit(new_id(), gamma.shape.clone(), bias_dtype);
Self { x, gamma, grad, mean, rstd, input_grad, weight_grad, bias_grad }
}
}
impl ReduceDimWithIndicesOpIr {
pub fn create(
tensor: TensorIr,
dim: usize,
dtype_indices: DType,
mut new_id: impl FnMut() -> TensorId,
) -> Self {
let mut shape = tensor.shape.clone();
shape[dim] = 1;
let out = TensorIr::uninit(new_id(), shape.clone(), tensor.dtype);
let out_indices = TensorIr::uninit(new_id(), shape.clone(), dtype_indices);
ReduceDimWithIndicesOpIr {
tensor,
dim,
out,
out_indices,
}
}
}
impl DeformConv2dBackwardOpIr {
#[allow(clippy::too_many_arguments)]
pub fn create(
x: TensorIr,
offset: TensorIr,
weight: TensorIr,
mask: Option<TensorIr>,
bias: Option<TensorIr>,
out_grad: TensorIr,
options: DeformableConv2dOptionsIr,
mut new_id: impl FnMut() -> TensorId,
) -> Self {
let dtype = output_dtype(
[
Some(&x.dtype),
Some(&weight.dtype),
mask.as_ref().map(|m| &m.dtype),
bias.as_ref().map(|b| &b.dtype),
]
.iter()
.filter_map(|&d| d),
)
.unwrap();
let input_grad = TensorIr::uninit(new_id(), x.shape.clone(), dtype);
let offset_grad = TensorIr::uninit(new_id(), offset.shape.clone(), dtype);
let weight_grad = TensorIr::uninit(new_id(), weight.shape.clone(), dtype);
let mask_grad = mask
.as_ref()
.map(|t| TensorIr::uninit(new_id(), t.shape.clone(), dtype));
let bias_grad = bias
.as_ref()
.map(|t| TensorIr::uninit(new_id(), t.shape.clone(), dtype));
DeformConv2dBackwardOpIr {
x,
offset,
weight,
mask,
bias,
out_grad,
options,
input_grad,
offset_grad,
weight_grad,
mask_grad,
bias_grad,
}
}
}
impl MaxPool1dWithIndicesOpIr {
#[allow(clippy::too_many_arguments)]
pub fn create(
x: TensorIr,
kernel_size: usize,
stride: usize,
padding: usize,
dilation: usize,
ceil_mode: bool,
dtype_indices: DType,
mut new_id: impl FnMut() -> TensorId,
) -> Self {
let shape = calculate_pool_output_shape(
&x.shape,
&[kernel_size],
&[stride],
&[padding],
&[dilation],
ceil_mode,
)
.unwrap();
let out = TensorIr::uninit(new_id(), shape.clone(), x.dtype);
let out_indices = TensorIr::uninit(new_id(), shape, dtype_indices);
MaxPool1dWithIndicesOpIr {
x,
kernel_size,
stride,
padding,
dilation,
ceil_mode,
out,
out_indices,
}
}
}
impl MaxPool2dWithIndicesOpIr {
#[allow(clippy::too_many_arguments)]
pub fn create(
x: TensorIr,
kernel_size: [usize; 2],
stride: [usize; 2],
padding: [usize; 2],
dilation: [usize; 2],
ceil_mode: bool,
dtype_indices: DType,
mut new_id: impl FnMut() -> TensorId,
) -> Self {
let shape = calculate_pool_output_shape(
&x.shape,
&kernel_size,
&stride,
&padding,
&dilation,
ceil_mode,
)
.unwrap();
let out = TensorIr::uninit(new_id(), shape.clone(), x.dtype);
let out_indices = TensorIr::uninit(new_id(), shape, dtype_indices);
MaxPool2dWithIndicesOpIr {
x,
kernel_size,
stride,
padding,
dilation,
ceil_mode,
out,
out_indices,
}
}
}