use crate::tensor::Tensor;
pub struct SubjectLayers {
pub weights: Tensor,
pub bias: Option<Tensor>,
pub average_subjects: bool,
pub n_subjects: usize,
}
impl SubjectLayers {
pub fn forward(&self, x: &Tensor, subject_ids: &[usize]) -> Tensor {
assert_eq!(x.ndim(), 3);
let (b, c_in, t) = (x.shape[0], x.shape[1], x.shape[2]);
let c_out = self.weights.shape[2];
let mut out = Tensor::zeros(&[b, c_out, t]);
for bi in 0..b {
let (w, bias) = if self.average_subjects {
let sid = self.n_subjects;
(
self.slice_weight(sid.min(self.weights.shape[0] - 1)),
self.bias
.as_ref()
.map(|bias| self.slice_bias(sid.min(bias.shape[0] - 1))),
)
} else {
let sid = subject_ids.get(bi).copied().unwrap_or(0);
let bias = self.bias.as_ref().map(|_| self.slice_bias(sid));
(self.slice_weight(sid), bias)
};
for ti in 0..t {
for co in 0..c_out {
let mut sum = 0.0f32;
for ci in 0..c_in {
sum += x.data[(bi * c_in + ci) * t + ti] * w.data[ci * c_out + co];
}
if let Some(ref bias) = bias {
sum += bias.data[co];
}
out.data[(bi * c_out + co) * t + ti] = sum;
}
}
}
out
}
fn slice_weight(&self, sid: usize) -> Tensor {
let c_in = self.weights.shape[1];
let c_out = self.weights.shape[2];
let base = sid * c_in * c_out;
Tensor::from_vec(
self.weights.data[base..base + c_in * c_out].to_vec(),
vec![c_in, c_out],
)
}
fn slice_bias(&self, sid: usize) -> Tensor {
let c_out = self.weights.shape[2];
let base = sid * c_out;
Tensor::from_vec(
self.bias.as_ref().unwrap().data[base..base + c_out].to_vec(),
vec![c_out],
)
}
}