use burn::prelude::*;
use burn::module::{Param, ParamId};
use burn::nn::{Linear, LinearConfig};
use burn::tensor::activation::gelu;
use crate::model::cross_attention::FusedMultiheadAttention;
#[derive(Module, Debug)]
pub struct ClassificationHead<B: Backend> {
pub learned_agg: Param<Tensor<B, 3>>,
pub decoder_attn: FusedMultiheadAttention<B>,
pub ffn_fc1: Linear<B>,
pub ffn_fc2: Linear<B>,
pub full_dim: usize,
}
impl<B: Backend> ClassificationHead<B> {
pub fn new(
embed_dim: usize,
num_queries: usize,
num_heads: usize,
num_classes: usize,
device: &B::Device,
) -> Self {
let full_dim = embed_dim * num_queries;
Self {
learned_agg: Param::initialized(
ParamId::new(),
Tensor::zeros([1, 1, full_dim], device), ),
decoder_attn: FusedMultiheadAttention::new(full_dim, num_heads, device),
ffn_fc1: LinearConfig::new(full_dim, full_dim * 4).with_bias(true).init(device),
ffn_fc2: LinearConfig::new(full_dim * 4, num_classes).with_bias(true).init(device),
full_dim,
}
}
pub fn forward(&self, x: Tensor<B, 3>) -> Tensor<B, 3> {
let [b, _n, _d] = x.dims();
let queries = self.learned_agg.val()
.expand([b, 1, self.full_dim]);
let (attn_out, _) = self.decoder_attn.forward(queries, x.clone(), x);
let out = attn_out.narrow(1, 0, 1).reshape([b, self.full_dim]);
let out = gelu(self.ffn_fc1.forward(out));
let out = self.ffn_fc2.forward(out);
out.unsqueeze_dim::<3>(1) }
}