use burn::prelude::*;
use burn::nn::{
conv::{Conv2d, Conv2dConfig},
GroupNorm, GroupNormConfig,
};
use burn::tensor::activation::gelu;
#[derive(Module, Debug)]
pub struct PatchEmbedNetwork<B: Backend> {
pub conv1: Conv2d<B>,
pub gn1: GroupNorm<B>,
pub conv2: Conv2d<B>,
pub gn2: GroupNorm<B>,
pub conv3: Conv2d<B>,
pub gn3: GroupNorm<B>,
pub patch_size: usize,
pub embed_dim: usize,
}
impl<B: Backend> PatchEmbedNetwork<B> {
pub fn new(embed_dim: usize, patch_size: usize, device: &B::Device) -> Self {
let out_channels = embed_dim / 4;
let groups = 4;
let kernel = patch_size / 2;
let conv1 = Conv2dConfig::new([1, out_channels], [1, kernel - 1])
.with_stride([1, kernel / 2])
.with_padding(burn::nn::PaddingConfig2d::Valid)
.with_bias(true)
.init(device);
let gn1 = GroupNormConfig::new(groups, out_channels)
.with_epsilon(1e-5)
.init(device);
let conv2 = Conv2dConfig::new([out_channels, out_channels], [1, 3])
.with_stride([1, 1])
.with_padding(burn::nn::PaddingConfig2d::Explicit(0, 1))
.with_bias(true)
.init(device);
let gn2 = GroupNormConfig::new(groups, out_channels)
.with_epsilon(1e-5)
.init(device);
let conv3 = Conv2dConfig::new([out_channels, out_channels], [1, 3])
.with_stride([1, 1])
.with_padding(burn::nn::PaddingConfig2d::Explicit(0, 1))
.with_bias(true)
.init(device);
let gn3 = GroupNormConfig::new(groups, out_channels)
.with_epsilon(1e-5)
.init(device);
Self {
conv1, gn1, conv2, gn2, conv3, gn3,
patch_size, embed_dim,
}
}
pub fn forward(&self, x: Tensor<B, 3>) -> Tensor<B, 3> {
let [b, c, t] = x.dims();
let s = t / self.patch_size;
let x = x.reshape([b, c * s, self.patch_size]).unsqueeze_dim::<4>(1);
let kernel = self.patch_size / 2;
let w_pad = kernel / 2 - 1;
let cs = c * s;
let pad_left = Tensor::zeros([b, 1, cs, w_pad], &x.device());
let pad_right = Tensor::zeros([b, 1, cs, w_pad], &x.device());
let x_padded = Tensor::cat(vec![pad_left, x, pad_right], 3);
let x = gelu(self.gn1.forward(self.conv1.forward(x_padded)));
let x = gelu(self.gn2.forward(self.conv2.forward(x)));
let x = gelu(self.gn3.forward(self.conv3.forward(x)));
let [b2, e, cs, d_prime] = x.dims();
x.swap_dims(1, 2) .swap_dims(2, 3) .reshape([b2, cs, d_prime * e]) }
}