use burn::prelude::*;
use burn::nn::{Linear, LinearConfig};
use burn::tensor::activation::gelu;
use crate::model::norm::LunaLayerNorm;
use crate::model::cross_attention::FusedMultiheadAttention;
#[derive(Module, Debug)]
pub struct ReconstructionHead<B: Backend> {
pub self_attn_norm: LunaLayerNorm<B>,
pub self_attn: FusedMultiheadAttention<B>,
pub cross_attn_norm: LunaLayerNorm<B>,
pub cross_attn: FusedMultiheadAttention<B>,
pub ffn_norm: LunaLayerNorm<B>,
pub ffn_linear1: Linear<B>,
pub ffn_linear2: Linear<B>,
pub output_norm: LunaLayerNorm<B>,
pub output_fc1: Linear<B>,
pub output_fc2: Linear<B>,
pub embed_dim: usize,
pub num_queries: usize,
pub patch_size: usize,
}
impl<B: Backend> ReconstructionHead<B> {
pub fn new(
input_dim: usize, embed_dim: usize,
num_heads: usize,
num_queries: usize,
device: &B::Device,
) -> Self {
let ff_dim = embed_dim * 4;
Self {
self_attn_norm: LunaLayerNorm::new(embed_dim, 1e-5, device),
self_attn: FusedMultiheadAttention::new(embed_dim, num_heads, device),
cross_attn_norm: LunaLayerNorm::new(embed_dim, 1e-5, device),
cross_attn: FusedMultiheadAttention::new(embed_dim, num_heads, device),
ffn_norm: LunaLayerNorm::new(embed_dim, 1e-5, device),
ffn_linear1: LinearConfig::new(embed_dim, ff_dim).with_bias(true).init(device),
ffn_linear2: LinearConfig::new(ff_dim, embed_dim).with_bias(true).init(device),
output_norm: LunaLayerNorm::new(embed_dim, 1e-5, device),
output_fc1: LinearConfig::new(embed_dim, ff_dim).with_bias(true).init(device),
output_fc2: LinearConfig::new(ff_dim, input_dim).with_bias(true).init(device),
embed_dim,
num_queries,
patch_size: input_dim,
}
}
pub fn forward(
&self,
enc: Tensor<B, 3>,
decoder_queries: Tensor<B, 3>,
num_channels: usize,
) -> Tensor<B, 3> {
let [b, num_patches, _qd] = enc.dims();
let memory = enc.reshape([b * num_patches, self.num_queries, self.embed_dim]);
let mut tgt = decoder_queries;
let normed = self.self_attn_norm.forward(tgt.clone());
let (sa_out, _) = self.self_attn.forward(normed.clone(), normed.clone(), normed);
tgt = tgt + sa_out;
let normed = self.cross_attn_norm.forward(tgt.clone());
let (ca_out, _) = self.cross_attn.forward(normed, memory.clone(), memory);
tgt = tgt + ca_out;
let normed = self.ffn_norm.forward(tgt.clone());
let ff_out = self.ffn_linear2.forward(gelu(self.ffn_linear1.forward(normed)));
let out = tgt + ff_out;
let out = self.output_norm.forward(out);
let out = gelu(self.output_fc1.forward(out));
let out = self.output_fc2.forward(out);
out.reshape([b, num_patches, num_channels, self.patch_size])
.swap_dims(1, 2) .reshape([b, num_channels, num_patches * self.patch_size])
}
}