use burn_backend::ops::ModuleOps;
use burn_dispatch::Dispatch;
use burn_std::{MatmulTransformAction, MatmulTransformAnalysis, MatmulTransformPolicy};
use crate::{
Bool, DType, Int, Tensor, check,
check::TensorCheck,
ops::{
AttentionModuleOptions, BridgeTensor, ConvOptions, ConvTransposeOptions, DeformConvOptions,
InterpolateOptions, PadMode, PaddedConvOptions, UnfoldOptions,
},
};
pub fn ctc_loss(
log_probs: Tensor<3>,
targets: Tensor<2, Int>,
input_lengths: Tensor<1, Int>,
target_lengths: Tensor<1, Int>,
blank: usize,
) -> Tensor<1> {
Tensor::new(BridgeTensor::float(Dispatch::ctc_loss(
log_probs.primitive.into_float(),
targets.primitive.into(),
input_lengths.primitive.into(),
target_lengths.primitive.into(),
blank,
)))
}
pub fn embedding(weights: Tensor<2>, indices: Tensor<2, Int>) -> Tensor<3> {
Tensor::new(BridgeTensor::float(Dispatch::embedding(
weights.primitive.into_float(),
indices.primitive.into(),
)))
}
pub fn conv1d(
x: Tensor<3>,
weight: Tensor<3>,
bias: Option<Tensor<1>>,
options: impl Into<PaddedConvOptions<1>>,
) -> Tensor<3> {
let padded_options = options.into();
check!(TensorCheck::conv(
"conv1d",
x.dims(),
weight.dims(),
padded_options.options.groups,
));
if let Some(padding_end) = padded_options.padding_end {
let left = padded_options.options.padding[0];
let right = padding_end[0];
let padded = x.pad((left, right, 0, 0), PadMode::Constant(0.0));
let zero_options = ConvOptions::new(
padded_options.options.stride,
[0],
padded_options.options.dilation,
padded_options.options.groups,
);
Tensor::new(BridgeTensor::float(Dispatch::conv1d(
padded.primitive.into_float(),
weight.primitive.into_float(),
bias.map(|b| b.primitive.into_float()),
zero_options,
)))
} else {
Tensor::new(BridgeTensor::float(Dispatch::conv1d(
x.primitive.into_float(),
weight.primitive.into_float(),
bias.map(|b| b.primitive.into_float()),
padded_options.options,
)))
}
}
pub fn conv2d(
x: Tensor<4>,
weight: Tensor<4>,
bias: Option<Tensor<1>>,
options: impl Into<PaddedConvOptions<2>>,
) -> Tensor<4> {
let padded_options = options.into();
check!(TensorCheck::conv(
"conv2d",
x.dims(),
weight.dims(),
padded_options.options.groups,
));
if let Some(padding_end) = padded_options.padding_end {
let top = padded_options.options.padding[0];
let left = padded_options.options.padding[1];
let bottom = padding_end[0];
let right = padding_end[1];
let padded = x.pad((left, right, top, bottom), PadMode::Constant(0.0));
let zero_options = ConvOptions::new(
padded_options.options.stride,
[0, 0],
padded_options.options.dilation,
padded_options.options.groups,
);
Tensor::new(BridgeTensor::float(Dispatch::conv2d(
padded.primitive.into_float(),
weight.primitive.into_float(),
bias.map(|b| b.primitive.into_float()),
zero_options,
)))
} else {
Tensor::new(BridgeTensor::float(Dispatch::conv2d(
x.primitive.into_float(),
weight.primitive.into_float(),
bias.map(|b| b.primitive.into_float()),
padded_options.options,
)))
}
}
pub fn conv3d(
x: Tensor<5>,
weight: Tensor<5>,
bias: Option<Tensor<1>>,
options: impl Into<PaddedConvOptions<3>>,
) -> Tensor<5> {
let padded_options = options.into();
check!(TensorCheck::conv(
"conv3d",
x.dims(),
weight.dims(),
padded_options.options.groups,
));
if padded_options.is_asymmetric() {
panic!("Asymmetric padding is not yet supported for conv3d");
}
Tensor::new(BridgeTensor::float(Dispatch::conv3d(
x.primitive.into_float(),
weight.primitive.into_float(),
bias.map(|b| b.primitive.into_float()),
padded_options.options,
)))
}
pub fn deform_conv2d(
x: Tensor<4>,
offset: Tensor<4>,
weight: Tensor<4>,
mask: Option<Tensor<4>>,
bias: Option<Tensor<1>>,
options: DeformConvOptions<2>,
) -> Tensor<4> {
check!(TensorCheck::conv(
"deform_conv2d",
x.dims(),
weight.dims(),
options.weight_groups,
));
Tensor::new(BridgeTensor::float(Dispatch::deform_conv2d(
x.primitive.into_float(),
offset.primitive.into_float(),
weight.primitive.into_float(),
mask.map(|m| m.primitive.into_float()),
bias.map(|b| b.primitive.into_float()),
options,
)))
}
pub fn conv_transpose1d(
x: Tensor<3>,
weight: Tensor<3>,
bias: Option<Tensor<1>>,
options: ConvTransposeOptions<1>,
) -> Tensor<3> {
check!(TensorCheck::conv_transpose(
"conv_transpose1d",
x.dims(),
weight.dims(),
));
Tensor::new(BridgeTensor::float(Dispatch::conv_transpose1d(
x.primitive.into_float(),
weight.primitive.into_float(),
bias.map(|b| b.primitive.into_float()),
options,
)))
}
pub fn conv_transpose2d(
x: Tensor<4>,
weight: Tensor<4>,
bias: Option<Tensor<1>>,
options: ConvTransposeOptions<2>,
) -> Tensor<4> {
check!(TensorCheck::conv_transpose(
"conv_transpose2d",
x.dims(),
weight.dims(),
));
Tensor::new(BridgeTensor::float(Dispatch::conv_transpose2d(
x.primitive.into_float(),
weight.primitive.into_float(),
bias.map(|b| b.primitive.into_float()),
options,
)))
}
pub fn conv_transpose3d(
x: Tensor<5>,
weight: Tensor<5>,
bias: Option<Tensor<1>>,
options: ConvTransposeOptions<3>,
) -> Tensor<5> {
check!(TensorCheck::conv_transpose(
"conv_transpose3d",
x.dims(),
weight.dims(),
));
Tensor::new(BridgeTensor::float(Dispatch::conv_transpose3d(
x.primitive.into_float(),
weight.primitive.into_float(),
bias.map(|b| b.primitive.into_float()),
options,
)))
}
pub fn unfold4d(x: Tensor<4>, kernel_size: [usize; 2], options: UnfoldOptions) -> Tensor<3> {
Tensor::new(BridgeTensor::float(Dispatch::unfold4d(
x.primitive.into_float(),
kernel_size,
options,
)))
}
pub fn max_pool1d(
x: Tensor<3>,
kernel_size: usize,
stride: usize,
padding: usize,
dilation: usize,
ceil_mode: bool,
) -> Tensor<3> {
Tensor::new(BridgeTensor::float(Dispatch::max_pool1d(
x.primitive.into_float(),
kernel_size,
stride,
padding,
dilation,
ceil_mode,
)))
}
pub fn max_pool2d(
x: Tensor<4>,
kernel_size: [usize; 2],
stride: [usize; 2],
padding: [usize; 2],
dilation: [usize; 2],
ceil_mode: bool,
) -> Tensor<4> {
Tensor::new(BridgeTensor::float(Dispatch::max_pool2d(
x.primitive.into_float(),
kernel_size,
stride,
padding,
dilation,
ceil_mode,
)))
}
pub fn avg_pool2d(
x: Tensor<4>,
kernel_size: [usize; 2],
stride: [usize; 2],
padding: [usize; 2],
count_include_pad: bool,
ceil_mode: bool,
) -> Tensor<4> {
Tensor::new(BridgeTensor::float(Dispatch::avg_pool2d(
x.primitive.into_float(),
kernel_size,
stride,
padding,
count_include_pad,
ceil_mode,
)))
}
pub fn avg_pool1d(
x: Tensor<3>,
kernel_size: usize,
stride: usize,
padding: usize,
count_include_pad: bool,
ceil_mode: bool,
) -> Tensor<3> {
Tensor::new(BridgeTensor::float(Dispatch::avg_pool1d(
x.primitive.into_float(),
kernel_size,
stride,
padding,
count_include_pad,
ceil_mode,
)))
}
pub fn max_pool1d_with_indices(
x: Tensor<3>,
kernel_size: usize,
stride: usize,
padding: usize,
dilation: usize,
ceil_mode: bool,
) -> (Tensor<3>, Tensor<3, Int>) {
let indices_dtype = x.device().settings().int_dtype;
let output = Dispatch::max_pool1d_with_indices(
x.primitive.into_float(),
kernel_size,
stride,
padding,
dilation,
ceil_mode,
indices_dtype,
);
(
Tensor::new(BridgeTensor::float(output.output)),
Tensor::new(BridgeTensor::int(output.indices)),
)
}
pub fn max_pool2d_with_indices(
x: Tensor<4>,
kernel_size: [usize; 2],
stride: [usize; 2],
padding: [usize; 2],
dilation: [usize; 2],
ceil_mode: bool,
) -> (Tensor<4>, Tensor<4, Int>) {
let indices_dtype = x.device().settings().int_dtype;
let output = Dispatch::max_pool2d_with_indices(
x.primitive.into_float(),
kernel_size,
stride,
padding,
dilation,
ceil_mode,
indices_dtype,
);
(
Tensor::new(BridgeTensor::float(output.output)),
Tensor::new(BridgeTensor::int(output.indices)),
)
}
pub fn adaptive_avg_pool2d(x: Tensor<4>, output_size: [usize; 2]) -> Tensor<4> {
Tensor::new(BridgeTensor::float(Dispatch::adaptive_avg_pool2d(
x.primitive.into_float(),
output_size,
)))
}
pub fn adaptive_avg_pool1d(x: Tensor<3>, output_size: usize) -> Tensor<3> {
Tensor::new(BridgeTensor::float(Dispatch::adaptive_avg_pool1d(
x.primitive.into_float(),
output_size,
)))
}
pub fn interpolate(
x: Tensor<4>,
output_size: [usize; 2],
options: InterpolateOptions,
) -> Tensor<4> {
Tensor::new(BridgeTensor::float(Dispatch::interpolate(
x.primitive.into_float(),
output_size,
options,
)))
}
pub fn linear<const D: usize>(
input: Tensor<D>,
weight: Tensor<2>,
bias: Option<Tensor<1>>,
) -> Tensor<D> {
if D == 1 {
let input = input.unsqueeze::<2>();
let output = linear(input, weight, bias);
return output.squeeze_dim(0);
}
if let DType::QFloat(_) = weight.dtype() {
let dims = input.dims();
let analysis = MatmulTransformAnalysis::from_shapes(&input.shape(), &weight.shape());
let output = match MatmulTransformPolicy::default().action(&analysis) {
MatmulTransformAction::MergeBatches { rows } => {
let d_in = dims[D - 1];
let d_out = weight.dims()[1];
let folded = input.reshape([rows, d_in]).matmul(weight);
let mut out_dims = dims;
out_dims[D - 1] = d_out;
folded.reshape(out_dims)
}
MatmulTransformAction::Keep => input.matmul(weight.unsqueeze::<D>()),
};
return match bias {
Some(bias) => output + bias.unsqueeze(),
None => output,
};
}
Tensor::new(linear_impl(
input.primitive,
weight.primitive,
bias.map(|b| b.primitive),
))
}
fn linear_impl(
input: BridgeTensor,
weight: BridgeTensor,
bias: Option<BridgeTensor>,
) -> BridgeTensor {
BridgeTensor::float(Dispatch::linear(
input.into_float(),
weight.into_float(),
bias.map(|b| b.into_float()),
))
}
pub fn attention(
query: Tensor<4>,
key: Tensor<4>,
value: Tensor<4>,
mask: Option<Tensor<4, Bool>>,
attn_bias: Option<Tensor<4>>,
options: AttentionModuleOptions,
) -> Tensor<4> {
Tensor::new(BridgeTensor::float(Dispatch::attention(
query.primitive.into_float(),
key.primitive.into_float(),
value.primitive.into_float(),
mask.map(|mask| mask.primitive.into()),
attn_bias.map(|bias| bias.primitive.into_float()),
options,
)))
}
pub fn attention_fallback(
query: Tensor<4>,
key: Tensor<4>,
value: Tensor<4>,
mask: Option<Tensor<4, Bool>>,
attn_bias: Option<Tensor<4>>,
options: AttentionModuleOptions,
) -> Tensor<4> {
Tensor::new(BridgeTensor::float(
burn_backend::ops::attention::attention_fallback::<Dispatch>(
query.primitive.into_float(),
key.primitive.into_float(),
value.primitive.into_float(),
mask.map(|mask| mask.primitive.into()),
attn_bias.map(|bias| bias.primitive.into_float()),
options,
),
))
}
pub fn conv2d_weight_backward(
x: Tensor<4>,
weight: Tensor<4>,
output_grad: Tensor<4>,
options: ConvOptions<2>,
) -> Tensor<4> {
Tensor::new(BridgeTensor::float(Dispatch::conv2d_weight_backward(
x.primitive.into_float(),
weight.primitive.into_float(),
output_grad.primitive.into_float(),
options,
)))
}
pub fn avg_pool2d_backward(
x: Tensor<4>,
grad: Tensor<4>,
kernel_size: [usize; 2],
stride: [usize; 2],
padding: [usize; 2],
count_include_pad: bool,
ceil_mode: bool,
) -> Tensor<4> {
Tensor::new(BridgeTensor::float(Dispatch::avg_pool2d_backward(
x.primitive.into_float(),
grad.primitive.into_float(),
kernel_size,
stride,
padding,
count_include_pad,
ceil_mode,
)))
}
#[allow(clippy::too_many_arguments)]
pub fn max_pool2d_with_indices_backward(
x: Tensor<4>,
kernel_size: [usize; 2],
stride: [usize; 2],
padding: [usize; 2],
dilation: [usize; 2],
ceil_mode: bool,
output_grad: Tensor<4>,
indices: Tensor<4, Int>,
) -> Tensor<4> {
Tensor::new(BridgeTensor::float(
Dispatch::max_pool2d_with_indices_backward(
x.primitive.into_float(),
kernel_size,
stride,
padding,
dilation,
ceil_mode,
output_grad.primitive.into_float(),
indices.primitive.into(),
)
.x_grad,
))
}
pub fn layer_norm<const D: usize>(
input: Tensor<D>,
gamma: Tensor<1>,
beta: Option<Tensor<1>>,
epsilon: f64,
) -> Tensor<D> {
Tensor::new(layer_norm_impl(
input.primitive,
gamma.primitive,
beta.map(|b| b.primitive),
epsilon,
))
}
fn layer_norm_impl(
input: BridgeTensor,
gamma: BridgeTensor,
beta: Option<BridgeTensor>,
epsilon: f64,
) -> BridgeTensor {
BridgeTensor::float(Dispatch::layer_norm(
input.into_float(),
gamma.into_float(),
beta.map(|b| b.into_float()),
epsilon,
))
}