use candle_core::{Result, Tensor};
use candle_nn::{linear, ops, Linear, Module, VarBuilder};
pub struct SymmetricPairHeadArgs {
pub code_dim: usize,
pub out_dim: usize,
pub n_experts: usize,
pub extra_dim: usize,
}
pub struct SymmetricPairHead {
gate: Option<Linear>,
experts: Linear,
code_dim: usize,
out_dim: usize,
n_experts: usize,
extra_dim: usize,
}
impl SymmetricPairHead {
pub fn new(args: SymmetricPairHeadArgs, vb: VarBuilder) -> Result<Self> {
let k = args.n_experts.max(1);
let in_dim = 2 * args.code_dim + args.extra_dim;
let gate = if k > 1 {
Some(linear(in_dim, k, vb.pp("gate"))?)
} else {
None
};
let experts = linear(in_dim, k * args.out_dim, vb.pp("experts"))?;
Ok(Self {
gate,
experts,
code_dim: args.code_dim,
out_dim: args.out_dim,
n_experts: k,
extra_dim: args.extra_dim,
})
}
pub fn code_dim(&self) -> usize {
self.code_dim
}
pub fn out_dim(&self) -> usize {
self.out_dim
}
pub fn n_experts(&self) -> usize {
self.n_experts
}
pub fn extra_dim(&self) -> usize {
self.extra_dim
}
pub fn features(&self, h_u: &Tensor, h_v: &Tensor, extra: Option<&Tensor>) -> Result<Tensor> {
let m = ((h_u + h_v)? * 0.5)?;
let p = (h_u * h_v)?;
match (extra, self.extra_dim) {
(None, 0) => Tensor::cat(&[&m, &p], 1),
(Some(x), e) if e > 0 => Tensor::cat(&[&m, &p, x], 1),
(Some(_), _) => {
candle_core::bail!("pair head: extra features given to a head built without them")
}
(None, e) => {
candle_core::bail!("pair head: built with {e} extra features but none were given")
}
}
}
pub fn gate(&self, features: &Tensor) -> Result<Tensor> {
match &self.gate {
Some(g) => ops::softmax(&g.forward(features)?, 1),
None => Tensor::ones((features.dim(0)?, 1), features.dtype(), features.device()),
}
}
pub fn forward(&self, h_u: &Tensor, h_v: &Tensor, extra: Option<&Tensor>) -> Result<Tensor> {
let f = self.features(h_u, h_v, extra)?;
let b = f.dim(0)?;
let e = self
.experts
.forward(&f)?
.reshape((b, self.n_experts, self.out_dim))?; if self.n_experts == 1 {
return e.squeeze(1);
}
let pi = self.gate(&f)?.unsqueeze(2)?; e.broadcast_mul(&pi)?.sum(1)
}
}
#[cfg(test)]
#[path = "pair_head_tests.rs"]
mod tests;