#![allow(dead_code)]
use candle_core::{Result, Tensor};
use candle_nn::{ops, Activation, Linear, Module};
use crate::candle::loss::gaussian_kl_loss;
pub struct IAFLayers<M, N>
where
M: Module,
N: Module,
{
z_mean: M,
z_lnvar: M,
autoregressive_mean: Vec<N>, autoregressive_sigmoid: Vec<N>, }
impl<M, N> IAFLayers<M, N>
where
M: Module,
N: Module,
{
fn iaf_step(&self, h: &Tensor, depth: usize, train: bool) -> Result<(Tensor, Tensor)> {
if depth == 0 {
let clamp_lo = -8.;
let clamp_hi = 8.;
let z_mean = self.z_mean.forward(h)?.clamp(clamp_lo, clamp_hi)?;
let z_lnvar = self.z_lnvar.forward(h)?.clamp(clamp_lo, clamp_hi)?;
let kl = gaussian_kl_loss(&z_mean, &z_lnvar)?;
let z = if train {
let eps = Tensor::randn_like(&z_mean, 0., 1.)?;
(z_mean + (z_lnvar * 0.5)?.exp()?.mul(&eps)?)?
} else {
z_mean
};
Ok((z, kl))
} else {
let (z_prev, kl) = self.iaf_step(h, depth - 1, train)?;
let zh = Tensor::cat(&[&z_prev, h], h.rank() - 1)?;
let z_new = self.autoregressive_mean[depth - 1].forward(&zh)?;
let s = self.autoregressive_sigmoid[depth - 1].forward(&zh)?;
let eps = 1e-8;
let p_stay = ((ops::sigmoid(&s)? * (1.0 - 2.0 * eps))? + eps)?;
let p_explore = ((ops::sigmoid(&s.neg()?)? * (1.0 - 2.0 * eps))? + eps)?;
let z = (z_new.mul(&p_explore)? + z_prev.mul(&p_stay)?)?;
let kl = kl.sub(&p_stay.log()?.sum(p_stay.rank() - 1)?)?;
Ok((z, kl))
}
}
pub fn flow(&self, h: &Tensor, train: bool) -> Result<(Tensor, Tensor)> {
self.iaf_step(h, self.autoregressive_mean.len(), train)
}
pub fn new(z_mean: M, z_lnvar: M, mean_layers: Vec<N>, sigmoid_layers: Vec<N>) -> Self {
debug_assert_eq!(mean_layers.len(), sigmoid_layers.len());
Self {
z_mean,
z_lnvar,
autoregressive_mean: mean_layers,
autoregressive_sigmoid: sigmoid_layers,
}
}
}
pub fn iaf_stack_linear(
in_dim: usize,
out_dim: usize,
n_layers: usize,
n_stack_layers: &[usize],
vb: candle_nn::VarBuilder,
) -> Result<IAFLayers<Linear, StackLayers<Linear>>> {
let z_mean = candle_nn::linear(in_dim, out_dim, vb.pp("iaf.z.mean"))?;
let z_lnvar = candle_nn::linear(in_dim, out_dim, vb.pp("iaf.z.lnvar"))?;
let mut mean_layers = vec![];
let mut sigmoid_layers = vec![];
for j in 0..n_layers {
let mean = stack_relu_linear(
in_dim + out_dim,
out_dim,
n_stack_layers,
vb.pp(format!("iaf.mean.{}", j)),
)?;
let sigmoid = stack_relu_linear(
in_dim + out_dim,
out_dim,
n_stack_layers,
vb.pp(format!("iaf.sigmoid.{}", j)),
)?;
mean_layers.push(mean);
sigmoid_layers.push(sigmoid);
}
let iaf_layers = IAFLayers::new(z_mean, z_lnvar, mean_layers, sigmoid_layers);
Ok(iaf_layers)
}
pub struct StackLayers<M>
where
M: Module,
{
module_layers: Vec<M>,
activation_layers: Vec<Option<Activation>>,
}
impl<M> Module for StackLayers<M>
where
M: Module,
{
fn forward(&self, input: &Tensor) -> Result<Tensor> {
let mut x = input.clone();
for (module, activation) in self.module_layers.iter().zip(self.activation_layers.iter()) {
x = module.forward(&x)?;
if let Some(activation) = activation {
x = activation.forward(&x)?;
}
}
Ok(x)
}
}
impl<M> StackLayers<M>
where
M: Module,
{
pub fn new() -> Self {
Self {
module_layers: Vec::new(),
activation_layers: Vec::new(),
}
}
pub fn push_with_act(&mut self, layer: M, activation: Activation) {
self.module_layers.push(layer);
self.activation_layers.push(Some(activation));
}
pub fn push(&mut self, layer: M) {
self.module_layers.push(layer);
self.activation_layers.push(None);
}
}
impl<M> Default for StackLayers<M>
where
M: Module,
{
fn default() -> Self {
Self::new()
}
}
pub fn stack_relu_linear(
in_dim: usize,
out_dim: usize,
intermediate_dims: &[usize],
vb: candle_nn::VarBuilder,
) -> Result<StackLayers<Linear>> {
let mut prev_dim = in_dim;
let mut ret = StackLayers::<Linear>::new();
for (k, &next_dim) in intermediate_dims.iter().enumerate() {
let _name = format!("relu_linear_stack.{}", k);
ret.push_with_act(
candle_nn::linear(prev_dim, next_dim, vb.pp(_name))?,
candle_nn::Activation::Relu,
);
prev_dim = next_dim;
}
let k = intermediate_dims.len();
let next_dim = out_dim;
let _name = format!("relu_linear_stack.{}", k);
ret.push_with_act(
candle_nn::linear(prev_dim, next_dim, vb.pp(_name))?,
candle_nn::Activation::Relu,
);
Ok(ret)
}
pub fn sparsemax(z: &Tensor) -> Result<Tensor> {
let z = z.contiguous()?; let dim = z.rank() - 1;
let (z_sorted, _indices) = z.sort_last_dim(false)?; let k = z.dim(dim)?;
let device = z.device();
let dtype = z.dtype();
let cumsum = z_sorted.cumsum(dim)?;
let range = Tensor::arange(1f32, (k + 1) as f32, device)?.to_dtype(dtype)?;
let shape: Vec<usize> = (0..z.rank())
.map(|i| if i == dim { k } else { 1 })
.collect();
let range = range.reshape(shape.as_slice())?;
let bound = (z_sorted.broadcast_mul(&range)? + 1.0)?;
let support = bound.gt(&cumsum)?;
let support_f = support.to_dtype(dtype)?;
let support_size = support_f.sum_keepdim(dim)?;
let z_sorted_masked = z_sorted.broadcast_mul(&support_f)?;
let z_sum_support = z_sorted_masked.sum_keepdim(dim)?;
let tau = (z_sum_support - 1.0)?.broadcast_div(&support_size.clamp(1.0, f64::INFINITY)?)?;
z.broadcast_sub(&tau)?.clamp(0.0, f64::INFINITY)
}