use burn_core as burn;
use burn::config::Config;
use burn::module::{Initializer, Module, Param};
use burn::tensor::Device;
use burn::tensor::Tensor;
use burn_nn::conv::{Conv2d, Conv2dConfig};
use burn_nn::{LayerNorm, LayerNormConfig, Linear, LinearConfig};
use super::clip_attention::{ClipQkvAttention, ClipQkvAttentionConfig};
use super::quick_gelu::QuickGelu;
#[derive(Config, Debug)]
pub struct ClipVisualEncoderConfig {
#[config(default = "3")]
pub in_channels: usize,
#[config(default = "768")]
pub embed_dim: usize,
#[config(default = "32")]
pub patch_size: usize,
#[config(default = "12")]
pub num_layers: usize,
#[config(default = "12")]
pub num_heads: usize,
#[config(default = "3072")]
pub mlp_dim: usize,
#[config(default = "256")]
pub image_size: usize,
}
impl ClipVisualEncoderConfig {
pub fn init(&self, device: &Device) -> ClipVisualEncoder {
assert_eq!(
self.image_size % self.patch_size,
0,
"image_size ({}) must be a multiple of patch_size ({})",
self.image_size,
self.patch_size
);
let num_patches = (self.image_size / self.patch_size).pow(2);
let seq_len = num_patches + 1;
let patch_embed = Conv2dConfig::new(
[self.in_channels, self.embed_dim],
[self.patch_size, self.patch_size],
)
.with_stride([self.patch_size, self.patch_size])
.with_bias(false)
.init(device);
let init = Initializer::Normal {
mean: 0.0,
std: 0.02,
};
let class_token = init.init([self.embed_dim], device);
let positional_embedding = init.init([seq_len, self.embed_dim], device);
let blocks = (0..self.num_layers)
.map(|_| {
TransformerBlockConfig::new(self.embed_dim, self.num_heads, self.mlp_dim)
.init(device)
})
.collect();
ClipVisualEncoder {
patch_embed,
class_token,
positional_embedding,
ln_pre: LayerNormConfig::new(self.embed_dim).init(device),
blocks,
ln_post: LayerNormConfig::new(self.embed_dim).init(device),
embed_dim: self.embed_dim,
patch_size: self.patch_size,
image_size: self.image_size,
}
}
}
#[derive(Debug)]
pub struct ClipOutput {
pub features: Vec<Tensor<3>>,
pub cls: Option<Tensor<2>>,
}
#[derive(Module, Debug)]
pub struct ClipVisualEncoder {
pub(crate) patch_embed: Conv2d,
pub(crate) class_token: Param<Tensor<1>>,
pub(crate) positional_embedding: Param<Tensor<2>>,
pub(crate) ln_pre: LayerNorm,
pub(crate) blocks: Vec<TransformerBlock>,
pub(crate) ln_post: LayerNorm,
pub(crate) embed_dim: usize,
pub(crate) patch_size: usize,
pub(crate) image_size: usize,
}
impl ClipVisualEncoder {
pub fn forward(&self, image: Tensor<4>) -> Tensor<2> {
self.forward_with_features(image, true)
.cls
.expect("cls requested")
}
pub fn forward_with_features(&self, image: Tensor<4>, return_cls: bool) -> ClipOutput {
let [batch, _, height, width] = image.dims();
assert_eq!(
height % self.patch_size,
0,
"image height ({}) must be a multiple of patch_size ({})",
height,
self.patch_size
);
assert_eq!(
width % self.patch_size,
0,
"image width ({}) must be a multiple of patch_size ({})",
width,
self.patch_size
);
let embed = self.embed_dim;
let x = self.patch_embed.forward(image);
let [_, _, h_out, w_out] = x.dims();
let num_patches = h_out * w_out;
let seq_len = num_patches + 1;
let x = x.reshape([batch, embed, num_patches]).swap_dims(1, 2);
let cls = self
.class_token
.val()
.reshape([1, 1, embed])
.expand([batch, 1, embed]);
let x = Tensor::cat(vec![cls, x], 1);
let pos_seq = self.positional_embedding.dims()[0];
assert_eq!(
pos_seq, seq_len,
"positional_embedding length {} does not match runtime sequence length {}",
pos_seq, seq_len
);
let pos = self.positional_embedding.val().reshape([1, seq_len, embed]);
let x = x + pos;
let mut x = self.ln_pre.forward(x);
let mut features = Vec::with_capacity(self.blocks.len());
for block in &self.blocks {
x = block.forward(x);
let patch_features = x.clone().slice([0..batch, 1..seq_len, 0..embed]);
features.push(patch_features);
}
let cls = return_cls.then(|| {
let cls = x.slice([0..batch, 0..1, 0..embed]).reshape([batch, embed]);
self.ln_post.forward(cls)
});
ClipOutput { features, cls }
}
}
#[derive(Config, Debug)]
pub(crate) struct TransformerBlockConfig {
pub d_model: usize,
pub n_heads: usize,
pub mlp_dim: usize,
}
impl TransformerBlockConfig {
pub(crate) fn init(&self, device: &Device) -> TransformerBlock {
TransformerBlock {
ln_1: LayerNormConfig::new(self.d_model).init(device),
attn: ClipQkvAttentionConfig::new(self.d_model, self.n_heads).init(device),
ln_2: LayerNormConfig::new(self.d_model).init(device),
mlp: TransformerMlpConfig::new(self.d_model, self.mlp_dim).init(device),
}
}
}
#[derive(Module, Debug)]
pub(crate) struct TransformerBlock {
pub(crate) ln_1: LayerNorm,
pub(crate) attn: ClipQkvAttention,
pub(crate) ln_2: LayerNorm,
pub(crate) mlp: TransformerMlp,
}
impl TransformerBlock {
pub(crate) fn forward(&self, x: Tensor<3>) -> Tensor<3> {
let attn_out = self.attn.forward(self.ln_1.forward(x.clone()));
let x = x + attn_out;
let mlp_out = self.mlp.forward(self.ln_2.forward(x.clone()));
x + mlp_out
}
}
#[derive(Config, Debug)]
pub(crate) struct TransformerMlpConfig {
pub d_model: usize,
pub mlp_dim: usize,
}
impl TransformerMlpConfig {
pub(crate) fn init(&self, device: &Device) -> TransformerMlp {
TransformerMlp {
c_fc: LinearConfig::new(self.d_model, self.mlp_dim)
.with_bias(true)
.init(device),
activation: QuickGelu,
c_proj: LinearConfig::new(self.mlp_dim, self.d_model)
.with_bias(true)
.init(device),
}
}
}
#[derive(Module, Debug)]
pub(crate) struct TransformerMlp {
pub(crate) c_fc: Linear,
pub(crate) activation: QuickGelu,
pub(crate) c_proj: Linear,
}
impl TransformerMlp {
pub(crate) fn forward(&self, x: Tensor<3>) -> Tensor<3> {
let x = self.c_fc.forward(x);
let x = self.activation.forward(x);
self.c_proj.forward(x)
}
}
#[cfg(test)]
mod tests {
use super::*;
use burn::tensor::Distribution;
#[test]
fn clip_visual_encoder_forward_shape() {
let device = Default::default();
let encoder = ClipVisualEncoderConfig::new()
.with_image_size(64)
.with_num_layers(2)
.init(&device);
let image = Tensor::<4>::random([1, 3, 64, 64], Distribution::Default, &device);
let cls = encoder.forward(image);
assert_eq!(cls.dims(), [1, 768]);
}
#[test]
fn clip_visual_encoder_forward_with_features_shape() {
let device = Default::default();
let encoder = ClipVisualEncoderConfig::new()
.with_image_size(64)
.with_num_layers(3)
.init(&device);
let image = Tensor::<4>::random([2, 3, 64, 64], Distribution::Default, &device);
let ClipOutput { features, cls } = encoder.forward_with_features(image, true);
let cls = cls.expect("cls requested");
assert_eq!(cls.dims(), [2, 768]);
assert_eq!(features.len(), 3);
for level in &features {
assert_eq!(level.dims(), [2, 4, 768]);
}
}
}