use burn::prelude::*;
use burn::module::{Param, ParamId};
use burn::nn::{Linear, LinearConfig, Embedding, EmbeddingConfig};
use crate::model::patch_embed::PatchEmbedNetwork;
use crate::model::freq_embed::FrequencyFeatureEmbedder;
use crate::model::cross_attention::CrossAttentionBlock;
use crate::model::encoder_block::RotaryEncoderBlock;
use crate::model::norm::LunaLayerNorm;
use crate::model::reconstruction_head::ReconstructionHead;
use crate::model::classification_head::ClassificationHead;
use crate::model::rope::RotaryEmbedding;
pub fn nerf_positional_encoding<B: Backend>(
coords: Tensor<B, 3>, embed_size: usize,
device: &B::Device,
) -> Tensor<B, 3> {
let [n, c, dim] = coords.dims(); let freqs = embed_size / (2 * dim);
let leftover = embed_size - freqs * 2 * dim;
let freq_data: Vec<f32> = (0..freqs).map(|i| 2.0_f32.powi(i as i32)).collect();
let freq_bands = Tensor::<B, 1>::from_data(
TensorData::new(freq_data, vec![freqs]), device,
);
let coords_4d = coords.unsqueeze_dim::<4>(3); let freq_4d = freq_bands.reshape([1, 1, 1, freqs]);
let scaled = coords_4d * freq_4d;
let sin_enc = scaled.clone().sin(); let cos_enc = scaled.cos();
let stacked = Tensor::stack::<5>(vec![sin_enc, cos_enc], 4); let encoded = stacked
.swap_dims(2, 3) .reshape([n, c, freqs * dim * 2]);
if leftover > 0 {
let pad = Tensor::zeros([n, c, leftover], device);
Tensor::cat(vec![encoded, pad], 2)
} else {
encoded
}
}
#[derive(Module, Debug)]
pub struct Luna<B: Backend> {
pub patch_embed: PatchEmbedNetwork<B>,
pub freq_embed: FrequencyFeatureEmbedder<B>,
pub chan_loc_fc1: Linear<B>,
pub chan_loc_norm: LunaLayerNorm<B>,
pub chan_loc_fc2: Linear<B>,
pub mask_token: Param<Tensor<B, 3>>,
pub cross_attn: CrossAttentionBlock<B>,
pub blocks: Vec<RotaryEncoderBlock<B>>,
pub norm: LunaLayerNorm<B>,
pub channel_emb: Option<Embedding<B>>,
pub decoder_head: Option<ReconstructionHead<B>>,
pub classifier: Option<ClassificationHead<B>>,
pub embed_dim: usize,
pub num_queries: usize,
pub patch_size: usize,
pub num_heads: usize,
pub num_classes: usize,
pub n_channel_names: usize,
}
impl<B: Backend> Luna<B> {
pub fn new(
patch_size: usize,
num_queries: usize,
embed_dim: usize,
depth: usize,
num_heads: usize,
mlp_ratio: f64,
norm_eps: f64,
num_classes: usize,
n_channel_names: usize,
device: &B::Device,
) -> Self {
let hidden_dim = embed_dim * num_queries;
let ff_dim = (embed_dim as f64 * mlp_ratio) as usize;
let patch_embed = PatchEmbedNetwork::new(embed_dim, patch_size, device);
let freq_embed = FrequencyFeatureEmbedder::new(embed_dim, patch_size, device);
let chan_loc_fc1 = LinearConfig::new(embed_dim, embed_dim * 2).with_bias(true).init(device);
let chan_loc_norm = LunaLayerNorm::new(embed_dim * 2, norm_eps, device);
let chan_loc_fc2 = LinearConfig::new(embed_dim * 2, embed_dim).with_bias(true).init(device);
let mask_token = Param::initialized(
ParamId::new(),
Tensor::zeros([1, 1, embed_dim], device),
);
let cross_attn = CrossAttentionBlock::new(
num_queries, embed_dim, embed_dim, num_heads, ff_dim, norm_eps, device,
);
let total_heads = num_heads * num_queries;
let blocks = (0..depth)
.map(|_| RotaryEncoderBlock::new(
hidden_dim, total_heads, mlp_ratio, true, norm_eps, device,
))
.collect();
let norm = LunaLayerNorm::new(hidden_dim, norm_eps, device);
let (channel_emb, decoder_head, classifier) = if num_classes == 0 {
let emb = EmbeddingConfig::new(n_channel_names, embed_dim).init(device);
let head = ReconstructionHead::new(
patch_size, embed_dim, num_heads, num_queries, device,
);
(Some(emb), Some(head), None)
} else {
let cls = ClassificationHead::new(
embed_dim, num_queries, num_heads, num_classes, device,
);
(None, None, Some(cls))
};
Self {
patch_embed,
freq_embed,
chan_loc_fc1, chan_loc_norm, chan_loc_fc2,
mask_token,
cross_attn,
blocks,
norm,
channel_emb,
decoder_head,
classifier,
embed_dim,
num_queries,
patch_size,
num_heads,
num_classes,
n_channel_names,
}
}
pub fn prepare_tokens(
&self,
x_signal: Tensor<B, 3>,
channel_locations: Tensor<B, 3>,
mask: Option<Tensor<B, 3>>,
) -> (Tensor<B, 3>, Tensor<B, 3>) {
let [b, num_channels, t] = x_signal.dims();
let num_patches = t / self.patch_size;
let device = x_signal.device();
let x_patched = self.patch_embed.forward(x_signal.clone()); let freq_embed = self.freq_embed.forward(x_signal); let x_patched = x_patched + freq_embed;
let x_masked = if let Some(ref m) = mask {
let mask_tokens = self.mask_token.val()
.expand([b, num_channels * num_patches, self.embed_dim]);
let m = m.clone().reshape([b, num_channels * num_patches, self.patch_size]);
let m = m.sum_dim(2); let m_bool = m.greater_elem(0.0); let m_expanded = m_bool.float()
.expand([b, num_channels * num_patches, self.embed_dim]);
x_patched.clone() * (Tensor::ones_like(&m_expanded) - m_expanded.clone())
+ mask_tokens * m_expanded
} else {
x_patched
};
let chan_locs_normed = {
let mins = channel_locations.clone().min_dim(1); let maxs = channel_locations.clone().max_dim(1); let range = maxs - mins.clone();
(channel_locations - mins) / (range + 1e-8)
};
let chan_locs_encoded = nerf_positional_encoding(
chan_locs_normed, self.embed_dim, &device,
);
let chan_loc_emb = burn::tensor::activation::gelu(
self.chan_loc_fc1.forward(chan_locs_encoded)
);
let chan_loc_emb = self.chan_loc_norm.forward(chan_loc_emb);
let chan_loc_emb = self.chan_loc_fc2.forward(chan_loc_emb);
let x_tokenized = x_masked
.reshape([b, num_channels, num_patches, self.embed_dim])
.swap_dims(1, 2) .reshape([b * num_patches, num_channels, self.embed_dim]);
let chan_loc_emb = chan_loc_emb.repeat_dim(0, num_patches);
let x_tokenized = x_tokenized + chan_loc_emb.clone();
(x_tokenized, chan_loc_emb)
}
pub fn forward(
&self,
x_signal: Tensor<B, 3>,
channel_locations: Tensor<B, 3>,
mask: Option<Tensor<B, 3>>,
channel_names: Option<Tensor<B, 2, Int>>,
rope: &RotaryEmbedding<B>,
) -> LunaOutput<B> {
let x_original = x_signal.clone();
let [b, num_channels, t] = x_signal.dims();
let num_patches = t / self.patch_size;
let (x_tokenized, chan_loc_emb) = self.prepare_tokens(
x_signal, channel_locations, mask,
);
let (x_unified, attention_scores) = self.cross_attn.forward(x_tokenized);
let x = x_unified.reshape([b, num_patches, self.num_queries * self.embed_dim]);
let freqs = rope.get_freqs(num_patches);
let mut x = x;
for blk in &self.blocks {
x = blk.forward(x, freqs.clone());
}
let x_latent = self.norm.forward(x);
if let Some(ref classifier) = self.classifier {
let logits = classifier.forward(x_latent); LunaOutput::Classification {
logits,
x_original,
}
} else if let Some(ref decoder_head) = self.decoder_head {
let mut decoder_queries = chan_loc_emb; if let (Some(ref emb), Some(names)) = (&self.channel_emb, channel_names) {
let ch_emb = emb.forward(names); let ch_emb = ch_emb.repeat_dim(0, num_patches); decoder_queries = decoder_queries + ch_emb;
}
let x_reconstructed = decoder_head.forward(
x_latent, decoder_queries, num_channels,
);
LunaOutput::Reconstruction {
x_reconstructed,
x_original,
attention_scores,
}
} else {
panic!("LUNA model has neither classifier nor decoder head");
}
}
}
pub enum LunaOutput<B: Backend> {
Classification {
logits: Tensor<B, 3>,
x_original: Tensor<B, 3>,
},
Reconstruction {
x_reconstructed: Tensor<B, 3>,
x_original: Tensor<B, 3>,
attention_scores: Tensor<B, 3>,
},
}