use burn::{
config::Config,
module::Module,
nn::{
PaddingConfig2d,
activation::{
Activation,
ActivationConfig,
},
conv::{
Conv2d,
Conv2dConfig,
},
norm::{
Normalization,
NormalizationConfig,
},
},
prelude::{
Backend,
Tensor,
},
};
use crate::{
burner::module::ModuleInit,
errors::{
BunsenError,
BunsenResult,
},
ops::conv::maybe_conv_output_shape,
};
#[derive(Config, Debug)]
pub struct AbstractConvBlock2dConfig {
pub norm: Option<NormalizationConfig>,
#[config(default = "Some(ActivationConfig::Relu)")]
pub act: Option<ActivationConfig>,
}
impl AbstractConvBlock2dConfig {
pub fn build_config(
&self,
conv: Conv2dConfig,
) -> ConvBlock2dConfig {
ConvBlock2dConfig {
conv,
norm: self.norm.clone(),
act: self.act.clone(),
}
.match_norm_features()
}
}
pub trait ConvBlock2dMeta {
fn in_channels(&self) -> usize;
fn out_channels(&self) -> usize;
fn groups(&self) -> usize;
fn stride(&self) -> [usize; 2];
fn kernel_size(&self) -> [usize; 2];
fn dilation(&self) -> [usize; 2];
fn padding(&self) -> PaddingConfig2d;
fn try_output_resolution(
&self,
input_resolution: [usize; 2],
) -> BunsenResult<[usize; 2]> {
let stride = self.stride();
let kernel_size = self.kernel_size();
let dilation = self.dilation();
let total_padding = match self.padding() {
PaddingConfig2d::Valid => [0, 0],
PaddingConfig2d::Explicit(top, left, bottom, right) => [top + bottom, left + right],
PaddingConfig2d::Same => {
let mut pads = [0; 2];
for d in 0..2 {
let out = input_resolution[d].div_ceil(stride[d]);
pads[d] = (out.saturating_sub(1) * stride[d] + kernel_size[d])
.saturating_sub(input_resolution[d]);
}
pads
}
};
let effective = [
input_resolution[0] + total_padding[0],
input_resolution[1] + total_padding[1],
];
maybe_conv_output_shape(effective, kernel_size, stride, [0, 0], dilation).ok_or_else(|| {
BunsenError::Invalid(format!(
"ConvBlock2d has no legal output resolution for input resolution \
({input_resolution:?})"
))
})
}
}
#[derive(Config, Debug)]
pub struct ConvBlock2dConfig {
pub conv: Conv2dConfig,
pub norm: Option<NormalizationConfig>,
#[config(default = "Some(ActivationConfig::Relu)")]
pub act: Option<ActivationConfig>,
}
impl ConvBlock2dMeta for ConvBlock2dConfig {
fn in_channels(&self) -> usize {
self.conv.channels[0]
}
fn out_channels(&self) -> usize {
self.conv.channels[1]
}
fn groups(&self) -> usize {
self.conv.groups
}
fn stride(&self) -> [usize; 2] {
self.conv.stride
}
fn kernel_size(&self) -> [usize; 2] {
self.conv.kernel_size
}
fn dilation(&self) -> [usize; 2] {
self.conv.dilation
}
fn padding(&self) -> PaddingConfig2d {
self.conv.padding.clone()
}
}
impl ConvBlock2dConfig {
pub fn match_norm_features(self) -> Self {
let features = self.out_channels();
let norm = self.norm.map(|config| config.with_num_features(features));
Self { norm, ..self }
}
}
impl<B: Backend> ModuleInit<B, ConvBlock2d<B>> for ConvBlock2dConfig {
fn try_init(
&self,
device: &B::Device,
) -> BunsenResult<ConvBlock2d<B>> {
let out_channels = self.out_channels();
Ok(ConvBlock2d {
conv: self.conv.init(device),
norm: self
.norm
.as_ref()
.map(|config| config.clone().with_num_features(out_channels).init(device)),
act: self.act.as_ref().map(|config| config.init(device)),
})
}
}
#[derive(Module, Debug)]
pub struct ConvBlock2d<B: Backend> {
pub conv: Conv2d<B>,
pub norm: Option<Normalization<B>>,
pub act: Option<Activation<B>>,
}
impl<B: Backend> ConvBlock2dMeta for ConvBlock2d<B> {
fn in_channels(&self) -> usize {
self.conv.weight.dims()[1] * self.groups()
}
fn out_channels(&self) -> usize {
self.conv.weight.dims()[0]
}
fn groups(&self) -> usize {
self.conv.groups
}
fn stride(&self) -> [usize; 2] {
self.conv.stride
}
fn kernel_size(&self) -> [usize; 2] {
self.conv.kernel_size
}
fn dilation(&self) -> [usize; 2] {
self.conv.dilation
}
fn padding(&self) -> PaddingConfig2d {
self.conv.padding.clone()
}
}
impl<B: Backend> ConvBlock2d<B> {
pub fn forward(
&self,
input: Tensor<B, 4>,
) -> Tensor<B, 4> {
self.map_forward(input, |x| x)
}
pub fn map_forward<F>(
&self,
input: Tensor<B, 4>,
f: F,
) -> Tensor<B, 4>
where
F: FnOnce(Tensor<B, 4>) -> Tensor<B, 4>,
{
#[cfg(debug_assertions)]
use crate::{
contracts::{
assert_shape_contract_periodically,
unpack_shape_contract,
},
errors::WithOkOrPanic,
};
#[cfg(debug_assertions)]
let [batch, in_height, in_width] = unpack_shape_contract!(
["batch", "in_channels", "in_height", "in_width"],
&input.dims(),
&["batch", "in_height", "in_width"],
&[("in_channels", self.in_channels())]
);
#[cfg(debug_assertions)]
let [out_height, out_width] = self
.try_output_resolution([in_height, in_width])
.ok_or_panic();
let x = self.conv.forward(input);
#[cfg(debug_assertions)]
assert_shape_contract_periodically!(
["batch", "out_channels", "out_height", "out_width"],
&x.dims(),
&[
("batch", batch),
("out_channels", self.out_channels()),
("out_height", out_height),
("out_width", out_width)
]
);
let x = match &self.norm {
Some(norm) => norm.forward(x),
None => x,
};
let x = f(x);
let x = match &self.act {
Some(act) => act.forward(x),
None => x,
};
#[cfg(debug_assertions)]
assert_shape_contract_periodically!(
["batch", "out_channels", "out_height", "out_width"],
&x.dims(),
&[
("batch", batch),
("out_channels", self.out_channels()),
("out_height", out_height),
("out_width", out_width)
]
);
x
}
}
#[cfg(test)]
mod tests {
use burn::{
backend::Autodiff,
nn::{
BatchNormConfig,
PaddingConfig2d,
activation::ActivationConfig,
norm::NormalizationConfig,
},
tensor::Distribution,
};
use super::*;
use crate::support::testing::CpuBackend;
#[test]
fn test_conv_norm_config() {
let abstract_config = AbstractConvBlock2dConfig::new()
.with_norm(Some(NormalizationConfig::Batch(BatchNormConfig::new(0))));
let conv_config = Conv2dConfig::new([2, 4], [3, 3])
.with_stride([2, 2])
.with_padding(PaddingConfig2d::Explicit(1, 1, 1, 1))
.with_bias(false);
let config: ConvBlock2dConfig = abstract_config.build_config(conv_config.clone());
assert_eq!(config.in_channels(), 2);
assert_eq!(config.out_channels(), 4);
assert_eq!(config.groups(), 1);
assert_eq!(config.stride(), [2, 2]);
}
#[test]
fn test_output_resolution() {
let block = |conv: Conv2dConfig| ConvBlock2dConfig::new(conv);
let same = block(
Conv2dConfig::new([2, 4], [3, 3])
.with_dilation([2, 2])
.with_padding(PaddingConfig2d::Explicit(2, 2, 2, 2)),
);
assert_eq!(same.try_output_resolution([10, 12]).unwrap(), [10, 12]);
let valid_dilated = block(
Conv2dConfig::new([2, 4], [3, 3])
.with_dilation([2, 2])
.with_padding(PaddingConfig2d::Valid),
);
assert_eq!(
valid_dilated.try_output_resolution([10, 12]).unwrap(),
[6, 8]
);
let valid = block(Conv2dConfig::new([2, 4], [3, 3]).with_padding(PaddingConfig2d::Valid));
assert_eq!(valid.try_output_resolution([10, 12]).unwrap(), [8, 10]);
let strided = block(
Conv2dConfig::new([2, 4], [3, 3])
.with_stride([2, 2])
.with_padding(PaddingConfig2d::Explicit(1, 1, 1, 1)),
);
assert_eq!(strided.try_output_resolution([10, 12]).unwrap(), [5, 6]);
let too_big = block(Conv2dConfig::new([2, 4], [5, 5]).with_padding(PaddingConfig2d::Valid));
assert!(matches!(
too_big.try_output_resolution([3, 3]),
Err(BunsenError::Invalid(_))
));
}
#[test]
fn test_dilated_forward_shape() {
type I = CpuBackend;
type B = Autodiff<I>;
let device = Default::default();
let config = ConvBlock2dConfig::new(
Conv2dConfig::new([2, 4], [3, 3])
.with_dilation([2, 2])
.with_padding(PaddingConfig2d::Valid)
.with_bias(false),
)
.with_norm(None)
.with_act(None);
let layer: ConvBlock2d<B> = config.init(&device);
let input = Tensor::random([2, 2, 10, 12], Distribution::Default, &device);
let output = layer.forward(input);
assert_eq!(output.dims(), [2, 4, 6, 8]);
assert_eq!(
[output.dims()[2], output.dims()[3]],
layer.try_output_resolution([10, 12]).unwrap()
);
}
#[test]
fn test_cb() {
type I = CpuBackend;
type B = Autodiff<I>;
let device = Default::default();
let config = ConvBlock2dConfig::new(
Conv2dConfig::new([2, 4], [3, 3])
.with_stride([2, 2])
.with_padding(PaddingConfig2d::Explicit(1, 1, 1, 1))
.with_bias(false),
)
.with_norm(Some(NormalizationConfig::Batch(BatchNormConfig::new(0))))
.with_act(Some(ActivationConfig::Relu));
let layer: ConvBlock2d<B> = config.init(&device);
assert_eq!(layer.in_channels(), 2);
assert_eq!(layer.out_channels(), 4);
assert_eq!(layer.groups(), 1);
assert_eq!(layer.stride(), [2, 2]);
let batch_size = 2;
let height = 10;
let width = 10;
let channels = 2;
let input = Tensor::random(
[batch_size, channels, height, width],
Distribution::Default,
&device,
);
{
let output = layer.forward(input.clone());
let expected = {
let x = layer.conv.forward(input.clone());
let x = layer.norm.as_ref().unwrap().forward(x);
let x = layer.act.as_ref().unwrap().forward(x);
x
};
output.to_data().assert_eq(&expected.to_data(), true);
}
{
let hook = |x| x * 2.0;
let output = layer.map_forward(input.clone(), hook);
let expected = {
let x = layer.conv.forward(input.clone());
let x = layer.norm.as_ref().unwrap().forward(x);
let x = hook(x);
let x = layer.act.as_ref().unwrap().forward(x);
x
};
output.to_data().assert_eq(&expected.to_data(), true);
}
}
}