#![allow(dead_code)]
use candle_core::{Result, Tensor};
use candle_nn::{Activation, Linear, Module};
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)
}