use burn::prelude::*;
use burn::module::{Param, ParamId};
use burn::nn::{Linear, LinearConfig};
use burn::tensor::activation::{gelu, softmax};
use crate::model::norm::LunaLayerNorm;
#[derive(Module, Debug)]
pub struct FusedMultiheadAttention<B: Backend> {
pub in_proj: Linear<B>, pub out_proj: Linear<B>, pub n_heads: usize,
pub head_dim: usize,
}
impl<B: Backend> FusedMultiheadAttention<B> {
pub fn new(dim: usize, n_heads: usize, device: &B::Device) -> Self {
let head_dim = dim / n_heads;
Self {
in_proj: LinearConfig::new(dim, dim * 3).with_bias(true).init(device),
out_proj: LinearConfig::new(dim, dim).with_bias(true).init(device),
n_heads,
head_dim,
}
}
pub fn forward(
&self,
q_input: Tensor<B, 3>,
k_input: Tensor<B, 3>,
v_input: Tensor<B, 3>,
) -> (Tensor<B, 3>, Tensor<B, 3>) {
let [b, s_q, _] = q_input.dims();
let s_kv = k_input.dims()[1];
let (h, dh) = (self.n_heads, self.head_dim);
let dim = h * dh;
let qkv_q = self.in_proj.forward(q_input); let q = qkv_q.narrow(2, 0, dim).reshape([b, s_q, h, dh]).swap_dims(1, 2);
let qkv_k = self.in_proj.forward(k_input);
let k = qkv_k.narrow(2, dim, dim).reshape([b, s_kv, h, dh]).swap_dims(1, 2);
let qkv_v = self.in_proj.forward(v_input);
let v = qkv_v.narrow(2, dim * 2, dim).reshape([b, s_kv, h, dh]).swap_dims(1, 2);
let scale = (dh as f64).powf(-0.5) as f32;
let attn_weights = softmax(q.matmul(k.transpose()).mul_scalar(scale), 3);
let out = attn_weights.clone().matmul(v); let out = self.out_proj.forward(
out.swap_dims(1, 2).reshape([b, s_q, dim])
);
let avg_attn = attn_weights.mean_dim(1).squeeze::<3>();
(out, avg_attn)
}
}
#[derive(Module, Debug)]
pub struct TransformerEncoderLayer<B: Backend> {
pub norm1: LunaLayerNorm<B>,
pub self_attn: FusedMultiheadAttention<B>,
pub norm2: LunaLayerNorm<B>,
pub linear1: Linear<B>,
pub linear2: Linear<B>,
}
impl<B: Backend> TransformerEncoderLayer<B> {
pub fn new(dim: usize, n_heads: usize, ff_dim: usize, norm_eps: f64, device: &B::Device) -> Self {
Self {
norm1: LunaLayerNorm::new(dim, norm_eps, device),
self_attn: FusedMultiheadAttention::new(dim, n_heads, device),
norm2: LunaLayerNorm::new(dim, norm_eps, device),
linear1: LinearConfig::new(dim, ff_dim).with_bias(true).init(device),
linear2: LinearConfig::new(ff_dim, dim).with_bias(true).init(device),
}
}
pub fn forward(&self, x: Tensor<B, 3>) -> Tensor<B, 3> {
let normed = self.norm1.forward(x.clone());
let (attn_out, _) = self.self_attn.forward(normed.clone(), normed.clone(), normed);
let x = x + attn_out;
let normed = self.norm2.forward(x.clone());
let ff_out = self.linear2.forward(gelu(self.linear1.forward(normed)));
x + ff_out
}
}
#[derive(Module, Debug)]
pub struct CrossAttentionBlock<B: Backend> {
pub query_embed: Param<Tensor<B, 3>>,
pub temperature: Param<Tensor<B, 1>>,
pub cross_attention: FusedMultiheadAttention<B>,
pub ffn_fc1: Linear<B>,
pub ffn_norm: LunaLayerNorm<B>,
pub ffn_fc2: Linear<B>,
pub queries_norm: LunaLayerNorm<B>,
pub keys_norm: LunaLayerNorm<B>,
pub values_norm: LunaLayerNorm<B>,
pub self_attn_layers: Vec<TransformerEncoderLayer<B>>,
pub num_queries: usize,
}
impl<B: Backend> CrossAttentionBlock<B> {
pub fn new(
num_queries: usize,
input_embed_dim: usize,
output_embed_dim: usize,
num_heads: usize,
ff_dim: usize,
norm_eps: f64,
device: &B::Device,
) -> Self {
let self_attn_layers = (0..3)
.map(|_| TransformerEncoderLayer::new(
input_embed_dim, num_heads, ff_dim, norm_eps, device,
))
.collect();
Self {
query_embed: Param::initialized(
ParamId::new(),
Tensor::zeros([1, num_queries, input_embed_dim], device),
),
temperature: Param::initialized(
ParamId::new(),
Tensor::ones([1], device),
),
cross_attention: FusedMultiheadAttention::new(input_embed_dim, num_heads, device),
ffn_fc1: LinearConfig::new(input_embed_dim, ff_dim).with_bias(true).init(device),
ffn_norm: LunaLayerNorm::new(ff_dim, norm_eps, device),
ffn_fc2: LinearConfig::new(ff_dim, output_embed_dim).with_bias(true).init(device),
queries_norm: LunaLayerNorm::new(input_embed_dim, norm_eps, device),
keys_norm: LunaLayerNorm::new(input_embed_dim, norm_eps, device),
values_norm: LunaLayerNorm::new(input_embed_dim, norm_eps, device),
self_attn_layers,
num_queries,
}
}
pub fn forward(&self, x: Tensor<B, 3>) -> (Tensor<B, 3>, Tensor<B, 3>) {
let [batch_size, _num_channels, dim] = x.dims();
let queries = self.query_embed.val()
.expand([batch_size, self.num_queries, dim]);
let queries = self.queries_norm.forward(queries);
let keys = self.keys_norm.forward(x.clone());
let values = self.values_norm.forward(x);
let (attention_out, attention_scores) =
self.cross_attention.forward(queries, keys, values);
let ffn_out = self.ffn_fc2.forward(
self.ffn_norm.forward(gelu(self.ffn_fc1.forward(attention_out.clone())))
);
let attention_out = ffn_out + attention_out;
let mut out = attention_out;
for layer in &self.self_attn_layers {
out = layer.forward(out);
}
(out, attention_scores)
}
}