use crate::tensor::Tensor;
pub struct BahdanauAttention {
pub wa_w: Tensor,
pub wa_b: Tensor,
pub va_w: Tensor,
pub va_b: Tensor,
}
impl BahdanauAttention {
pub fn forward(&self, keys: &Tensor) -> Tensor {
assert_eq!(keys.ndim(), 3);
let (b, f, t) = (keys.shape[0], keys.shape[1], keys.shape[2]);
let keys_t = keys.transpose(&[0, 2, 1]);
let mut sum = keys_t.linear(&self.wa_w, Some(&self.wa_b));
let sum_data: Vec<f32> = sum.data.iter().map(|x| x.tanh()).collect();
sum = Tensor::from_vec(sum_data, sum.shape.clone());
let scores = sum.linear(&self.va_w, Some(&self.va_b));
assert_eq!(scores.shape, vec![b, t, 1]);
let mut weights = vec![0.0f32; b * t];
for bi in 0..b {
let mut max_s = f32::NEG_INFINITY;
for ti in 0..t {
max_s = max_s.max(scores.data[(bi * t + ti) * 1]);
}
let mut denom = 0.0f32;
for ti in 0..t {
let e = (scores.data[(bi * t + ti) * 1] - max_s).exp();
weights[bi * t + ti] = e;
denom += e;
}
for ti in 0..t {
weights[bi * t + ti] /= denom;
}
}
let mut context = vec![0.0f32; b * f];
for bi in 0..b {
for fi in 0..f {
let mut v = 0.0f32;
for ti in 0..t {
v += weights[bi * t + ti] * keys_t.data[(bi * t + ti) * f + fi];
}
context[bi * f + fi] = v;
}
}
Tensor::from_vec(context, vec![b, f, 1])
}
}