#![allow(dead_code)]
use crate::candle::nn::layers::sparsemax;
use candle_core::{Result, Tensor};
use candle_nn::{ops, Module};
#[derive(Clone, Debug)]
pub struct NonNegLinear {
in_dim: usize,
out_dim: usize,
log_weight_dk: Tensor,
bias_d: Option<Tensor>,
}
impl NonNegLinear {
pub fn new(
in_dim: usize,
out_dim: usize,
log_weight_dk: Tensor,
bias_d: Option<Tensor>,
) -> Self {
Self {
in_dim,
out_dim,
log_weight_dk,
bias_d,
}
}
pub fn weight(&self) -> Result<Tensor> {
Ok(self.log_weight_dk.clone())
}
}
pub fn non_neg_linear(
in_dim: usize,
out_dim: usize,
vb: candle_nn::VarBuilder,
) -> Result<NonNegLinear> {
let ws = vb.get_with_hints((out_dim, in_dim), "weight", candle_nn::init::ZERO)?;
let bs_d = vb.get_with_hints((1, out_dim), "bias.d", candle_nn::init::ZERO)?;
Ok(NonNegLinear::new(in_dim, out_dim, ws, Some(bs_d)))
}
impl Module for NonNegLinear {
fn forward(&self, h_nk: &Tensor) -> Result<Tensor> {
let log_w_kd = match *h_nk.dims() {
[b1, b2, _, _] => self.log_weight_dk.broadcast_left((b1, b2))?.t()?,
[bsize, _, _] => self.log_weight_dk.broadcast_left(bsize)?.t()?,
_ => self.log_weight_dk.t()?,
};
let log_w_kd = match &self.bias_d {
None => log_w_kd,
Some(bias) => log_w_kd.broadcast_add(bias)?,
};
let eps = 1e-4;
let w_kd = (log_w_kd.relu()? + eps)?;
h_nk.matmul(&w_kd)
}
}
#[derive(Clone, Debug)]
pub struct AggregateLinear {
weight_dk: Tensor,
use_hard_assignment: bool,
}
impl AggregateLinear {
pub fn membership(&self) -> Result<Tensor> {
let eps = 1e-8;
(ops::sigmoid(&self.weight_dk)? * (1.0 - eps))? + eps
}
pub fn get_assignments(&self) -> Result<Tensor> {
self.weight_dk.argmax(1) }
fn forward_hard(&self, x_nd: &Tensor) -> Result<Tensor> {
let k = self.weight_dk.dim(1)?;
let assignments_d = self.get_assignments()?;
let mut columns = Vec::with_capacity(k);
for module_idx in 0..k {
let mask_d = assignments_d.eq(module_idx as f64)?;
let mask_float = mask_d.to_dtype(x_nd.dtype())?;
columns.push(mask_float.unsqueeze(1)?);
}
let assignment_matrix_dk = Tensor::cat(&columns, 1)?;
x_nd.matmul(&assignment_matrix_dk)
}
fn forward_soft(&self, x_nd: &Tensor) -> Result<Tensor> {
let c_dk = self.membership()?;
x_nd.matmul(&c_dk)
}
}
impl Module for AggregateLinear {
fn forward(&self, x_nd: &Tensor) -> Result<Tensor> {
if self.use_hard_assignment {
self.forward_hard(x_nd)
} else {
self.forward_soft(x_nd)
}
}
}
pub fn aggregate_linear(
in_dim: usize,
out_dim: usize,
vb: candle_nn::VarBuilder,
) -> Result<AggregateLinear> {
let init_ws = candle_nn::init::DEFAULT_KAIMING_NORMAL;
let weight_dk = vb.get_with_hints((in_dim, out_dim), "logits", init_ws)?;
Ok(AggregateLinear {
weight_dk,
use_hard_assignment: false,
})
}
pub fn aggregate_linear_hard(
in_dim: usize,
out_dim: usize,
vb: candle_nn::VarBuilder,
) -> Result<AggregateLinear> {
let init_ws = candle_nn::init::DEFAULT_KAIMING_NORMAL;
let weight_dk = vb.get_with_hints((in_dim, out_dim), "logits", init_ws)?;
Ok(AggregateLinear {
weight_dk,
use_hard_assignment: true,
})
}
#[derive(Clone, Debug)]
pub struct SoftmaxLinear {
weight_kd: Tensor,
bias_1d: Option<Tensor>,
}
impl SoftmaxLinear {
pub fn new(weight_kd: Tensor, bias_1d: Option<Tensor>) -> Self {
Self { weight_kd, bias_1d }
}
pub fn weight_dk(&self) -> Result<Tensor> {
ops::log_softmax(&self.weight_kd, self.weight_kd.rank() - 1)?
.transpose(0, 1)?
.contiguous()
}
pub(crate) fn raw_biased_logits_kd(&self) -> Result<Tensor> {
match &self.bias_1d {
Some(bias) => self.weight_kd.broadcast_add(bias),
_ => Ok(self.weight_kd.clone()),
}
}
pub(crate) fn biased_weight_ks_conditional(
&self,
union_indices: &Tensor,
log_q_s: &Tensor,
) -> Result<Tensor> {
let w_ks = self.weight_kd.index_select(union_indices, 1)?;
let w_ks = match &self.bias_1d {
Some(bias) => {
let bias_s = bias.index_select(union_indices, 1)?;
w_ks.broadcast_add(&bias_s)?
}
_ => w_ks,
};
let w_ks = w_ks.broadcast_sub(log_q_s)?;
ops::log_softmax(&w_ks, w_ks.rank() - 1)
}
pub(crate) fn biased_weight_kd(&self) -> Result<Tensor> {
match &self.bias_1d {
Some(bias) => ops::log_softmax(
&self.weight_kd.broadcast_add(bias)?,
self.weight_kd.rank() - 1,
),
_ => ops::log_softmax(&self.weight_kd, self.weight_kd.rank() - 1),
}?
.contiguous()
}
pub fn forward_log(&self, log_h_nk: &Tensor) -> Result<Tensor> {
logsumexp_forward(log_h_nk, &self.biased_weight_kd()?)
}
pub fn forward_log_slice(
&self,
log_h_nk: &Tensor,
log_w_kd: Option<&Tensor>,
start: usize,
len: usize,
) -> Result<Tensor> {
match log_w_kd {
Some(w) => logsumexp_forward(log_h_nk, &w.narrow(1, start, len)?),
None => logsumexp_forward(log_h_nk, &self.biased_weight_kd()?.narrow(1, start, len)?),
}
}
pub fn log_weight_kd(&self) -> Result<Tensor> {
self.biased_weight_kd()
}
}
impl Module for SoftmaxLinear {
fn forward(&self, log_h_nk: &Tensor) -> Result<Tensor> {
self.forward_log(log_h_nk)?.exp()
}
}
pub fn logsumexp_forward(log_h_nk: &Tensor, log_w_kd: &Tensor) -> Result<Tensor> {
let log_w_kd = match *log_h_nk.dims() {
[b1, b2, _, _] => log_w_kd.broadcast_left((b1, b2))?,
[bsize, _, _] => log_w_kd.broadcast_left(bsize)?,
_ => log_w_kd.clone(),
};
let log_h = log_h_nk.unsqueeze(2)?;
let log_w = log_w_kd.unsqueeze(0)?;
log_h.broadcast_add(&log_w)?.log_sum_exp(1)
}
pub fn log_softmax_linear(
in_dim: usize,
out_dim: usize,
vb: candle_nn::VarBuilder,
) -> Result<SoftmaxLinear> {
let init_ws = candle_nn::init::DEFAULT_KAIMING_NORMAL;
let ws_kd = vb.get_with_hints((in_dim, out_dim), "logits", init_ws)?;
let b_1d = vb.get_with_hints((1, out_dim), "logit_bias", candle_nn::init::ZERO)?;
Ok(SoftmaxLinear::new(ws_kd, Some(b_1d)))
}
pub fn log_softmax_linear_nobias(
in_dim: usize,
out_dim: usize,
vb: candle_nn::VarBuilder,
) -> Result<SoftmaxLinear> {
let init_ws = candle_nn::init::DEFAULT_KAIMING_NORMAL;
let ws_kd = vb.get_with_hints((in_dim, out_dim), "logits", init_ws)?;
Ok(SoftmaxLinear::new(ws_kd, None))
}
#[derive(Clone, Debug)]
pub struct SparsemaxLinear {
weight_kd: Tensor,
bias_1d: Option<Tensor>,
}
impl SparsemaxLinear {
pub fn new(weight_kd: Tensor, bias_1d: Option<Tensor>) -> Self {
Self { weight_kd, bias_1d }
}
pub fn weight_dk(&self) -> Result<Tensor> {
sparsemax(&self.weight_kd)?.transpose(0, 1)?.contiguous()
}
fn biased_weight_kd(&self) -> Result<Tensor> {
match &self.bias_1d {
Some(bias) => sparsemax(&self.weight_kd.broadcast_add(bias)?),
_ => sparsemax(&self.weight_kd),
}
}
pub fn forward_log(&self, log_h_nk: &Tensor) -> Result<Tensor> {
let h_nk = log_h_nk.exp()?;
let w_kd = match *h_nk.dims() {
[b1, b2, _, _] => self.biased_weight_kd()?.broadcast_left((b1, b2))?,
[bsize, _, _] => self.biased_weight_kd()?.broadcast_left(bsize)?,
_ => self.biased_weight_kd()?,
};
let eps = 1e-20;
(h_nk.matmul(&w_kd)? + eps)?.log()
}
}
impl Module for SparsemaxLinear {
fn forward(&self, log_h_nk: &Tensor) -> Result<Tensor> {
let h_nk = log_h_nk.exp()?;
let w_kd = match *h_nk.dims() {
[b1, b2, _, _] => self.biased_weight_kd()?.broadcast_left((b1, b2))?,
[bsize, _, _] => self.biased_weight_kd()?.broadcast_left(bsize)?,
_ => self.biased_weight_kd()?,
};
h_nk.matmul(&w_kd)
}
}
pub fn sparsemax_linear(
in_dim: usize,
out_dim: usize,
vb: candle_nn::VarBuilder,
) -> Result<SparsemaxLinear> {
let init_ws = candle_nn::init::DEFAULT_KAIMING_NORMAL;
let ws_kd = vb.get_with_hints((in_dim, out_dim), "logits", init_ws)?;
let b_1d = vb.get_with_hints((1, out_dim), "logit_bias", candle_nn::init::ZERO)?;
Ok(SparsemaxLinear::new(ws_kd, Some(b_1d)))
}