use crate::model::fourier::FourierEmb;
use crate::tensor::Tensor;
pub const INVALID_POS: f32 = -0.1;
pub struct ChannelMerger {
pub embedding: FourierEmb,
pub heads: Tensor,
pub per_subject: bool,
pub n_virtual: usize,
pub invalid_value: f32,
}
impl ChannelMerger {
pub fn forward(&self, meg: &Tensor, subject_ids: &[usize], positions: &Tensor) -> Tensor {
assert_eq!(meg.ndim(), 3);
let (b, c_in, t) = (meg.shape[0], meg.shape[1], meg.shape[2]);
let embedding = self.embedding.forward(positions);
let pos_dim = self.embedding.total_dim;
let mut out = Tensor::zeros(&[b, self.n_virtual, t]);
for bi in 0..b {
let sid = subject_ids.get(bi).copied().unwrap_or(0);
let heads = if self.per_subject {
self.gather_heads(sid)
} else {
self.heads.clone()
};
for vo in 0..self.n_virtual {
let mut scores = vec![f32::NEG_INFINITY; c_in];
for ci in 0..c_in {
let pos_base = (bi * c_in + ci) * 2;
if positions.data[pos_base] == self.invalid_value
&& positions.data[pos_base + 1] == self.invalid_value
{
continue;
}
let emb_base = (bi * c_in + ci) * pos_dim;
let head_base = vo * pos_dim;
let mut dot = 0.0f32;
for j in 0..pos_dim {
dot += embedding.data[emb_base + j] * heads.data[head_base + j];
}
scores[ci] = dot;
}
let max_s = scores.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0f32;
let mut weights = vec![0.0f32; c_in];
for ci in 0..c_in {
if scores[ci].is_finite() {
weights[ci] = (scores[ci] - max_s).exp();
sum += weights[ci];
}
}
if sum > 0.0 {
for w in &mut weights {
*w /= sum;
}
}
for ti in 0..t {
let mut v = 0.0f32;
for ci in 0..c_in {
v += weights[ci] * meg.data[(bi * c_in + ci) * t + ti];
}
out.data[(bi * self.n_virtual + vo) * t + ti] =
if v.is_nan() { 0.0 } else { v };
}
}
}
out
}
fn gather_heads(&self, subject_id: usize) -> Tensor {
let pos_dim = self.embedding.total_dim;
let n_subjects = self.heads.shape[0];
let sid = subject_id.min(n_subjects.saturating_sub(1));
let n_virt = self.heads.shape[1];
let mut data = vec![0.0f32; n_virt * pos_dim];
let src_base = sid * n_virt * pos_dim;
data.copy_from_slice(&self.heads.data[src_base..src_base + n_virt * pos_dim]);
Tensor::from_vec(data, vec![n_virt, pos_dim])
}
}