use super::gate_controller::GateController;
use crate::activation::{Activation, ActivationConfig};
use ruda_model::config::Config;
use ruda_model::module::Initializer;
use ruda_model::module::Module;
use ruda_model::module::{Content, DisplaySettings, ModuleDisplay};
use ruda_model::tensor::Tensor;
use ruda_model::tensor::backend::Backend;
#[derive(Config, Debug)]
pub struct GruConfig {
pub d_input: usize,
pub d_hidden: usize,
pub bias: bool,
#[config(default = "true")]
pub reset_after: bool,
#[config(default = "Initializer::XavierNormal{gain:1.0}")]
pub initializer: Initializer,
#[config(default = "ActivationConfig::Sigmoid")]
pub gate_activation: ActivationConfig,
#[config(default = "ActivationConfig::Tanh")]
pub hidden_activation: ActivationConfig,
pub clip: Option<f64>,
}
#[derive(Module, Debug)]
#[module(custom_display)]
pub struct Gru<B: Backend> {
pub update_gate: GateController<B>,
pub reset_gate: GateController<B>,
pub new_gate: GateController<B>,
pub d_hidden: usize,
pub reset_after: bool,
pub gate_activation: Activation<B>,
pub hidden_activation: Activation<B>,
pub clip: Option<f64>,
}
impl<B: Backend> ModuleDisplay for Gru<B> {
fn custom_settings(&self) -> Option<DisplaySettings> {
DisplaySettings::new()
.with_new_line_after_attribute(false)
.optional()
}
fn custom_content(&self, content: Content) -> Option<Content> {
let [d_input, _] = self.update_gate.input_transform.weight.shape().dims();
let bias = self.update_gate.input_transform.bias.is_some();
content
.add("d_input", &d_input)
.add("d_hidden", &self.d_hidden)
.add("bias", &bias)
.add("reset_after", &self.reset_after)
.optional()
}
}
impl GruConfig {
pub fn init<B: Backend>(&self, device: &B::Device) -> Gru<B> {
let d_output = self.d_hidden;
let update_gate = GateController::new(
self.d_input,
d_output,
self.bias,
self.initializer.clone(),
device,
);
let reset_gate = GateController::new(
self.d_input,
d_output,
self.bias,
self.initializer.clone(),
device,
);
let new_gate = GateController::new(
self.d_input,
d_output,
self.bias,
self.initializer.clone(),
device,
);
Gru {
update_gate,
reset_gate,
new_gate,
d_hidden: self.d_hidden,
reset_after: self.reset_after,
gate_activation: self.gate_activation.init(device),
hidden_activation: self.hidden_activation.init(device),
clip: self.clip,
}
}
}
impl<B: Backend> Gru<B> {
pub fn forward(
&self,
batched_input: Tensor<B, 3>,
state: Option<Tensor<B, 2>>,
) -> Tensor<B, 3> {
let device = batched_input.device();
let [batch_size, seq_length, _] = batched_input.shape().dims();
self.forward_iter(
batched_input.iter_dim(1).zip(0..seq_length),
state,
batch_size,
seq_length,
&device,
)
.0
}
pub(crate) fn forward_iter<I: Iterator<Item = (Tensor<B, 3>, usize)>>(
&self,
input_timestep_iter: I,
state: Option<Tensor<B, 2>>,
batch_size: usize,
seq_length: usize,
device: &B::Device,
) -> (Tensor<B, 3>, Tensor<B, 2>) {
let mut batched_hidden_state =
Tensor::empty([batch_size, seq_length, self.d_hidden], device);
let mut hidden_t = match state {
Some(state) => state,
None => Tensor::zeros([batch_size, self.d_hidden], device),
};
for (input_t, t) in input_timestep_iter {
let input_t = input_t.squeeze_dim(1);
let biased_ug_input_sum =
self.gate_product(&input_t, &hidden_t, None, &self.update_gate);
let update_values = self.gate_activation.forward(biased_ug_input_sum);
let biased_rg_input_sum =
self.gate_product(&input_t, &hidden_t, None, &self.reset_gate);
let reset_values = self.gate_activation.forward(biased_rg_input_sum);
let biased_ng_input_sum = if self.reset_after {
self.gate_product(&input_t, &hidden_t, Some(&reset_values), &self.new_gate)
} else {
let reset_t = hidden_t.clone().mul(reset_values);
self.gate_product(&input_t, &reset_t, None, &self.new_gate)
};
let candidate_state = self.hidden_activation.forward(biased_ng_input_sum);
let one_minus_z = update_values.clone().neg().add_scalar(1.0);
hidden_t = candidate_state.mul(one_minus_z) + update_values.mul(hidden_t);
if let Some(clip) = self.clip {
hidden_t = hidden_t.clamp(-clip, clip);
}
let unsqueezed_hidden_state = hidden_t.clone().unsqueeze_dim(1);
batched_hidden_state = batched_hidden_state.slice_assign(
[0..batch_size, t..(t + 1), 0..self.d_hidden],
unsqueezed_hidden_state,
);
}
(batched_hidden_state, hidden_t)
}
fn gate_product(
&self,
input: &Tensor<B, 2>,
hidden: &Tensor<B, 2>,
reset: Option<&Tensor<B, 2>>,
gate: &GateController<B>,
) -> Tensor<B, 2> {
let input_product = input.clone().matmul(gate.input_transform.weight.val());
let hidden_product = hidden.clone().matmul(gate.hidden_transform.weight.val());
let input_part = match &gate.input_transform.bias {
Some(bias) => input_product + bias.val().unsqueeze(),
None => input_product,
};
let hidden_part = match &gate.hidden_transform.bias {
Some(bias) => hidden_product + bias.val().unsqueeze(),
None => hidden_product,
};
match reset {
Some(r) => input_part + r.clone().mul(hidden_part),
None => input_part + hidden_part,
}
}
}
#[derive(Config, Debug)]
pub struct BiGruConfig {
pub d_input: usize,
pub d_hidden: usize,
pub bias: bool,
#[config(default = "true")]
pub reset_after: bool,
#[config(default = "Initializer::XavierNormal{gain:1.0}")]
pub initializer: Initializer,
#[config(default = true)]
pub batch_first: bool,
#[config(default = "ActivationConfig::Sigmoid")]
pub gate_activation: ActivationConfig,
#[config(default = "ActivationConfig::Tanh")]
pub hidden_activation: ActivationConfig,
pub clip: Option<f64>,
}
#[derive(Module, Debug)]
#[module(custom_display)]
pub struct BiGru<B: Backend> {
pub forward: Gru<B>,
pub reverse: Gru<B>,
pub d_hidden: usize,
pub batch_first: bool,
}
impl<B: Backend> ModuleDisplay for BiGru<B> {
fn custom_settings(&self) -> Option<DisplaySettings> {
DisplaySettings::new()
.with_new_line_after_attribute(false)
.optional()
}
fn custom_content(&self, content: Content) -> Option<Content> {
let [d_input, _] = self
.forward
.update_gate
.input_transform
.weight
.shape()
.dims();
let bias = self.forward.update_gate.input_transform.bias.is_some();
content
.add("d_input", &d_input)
.add("d_hidden", &self.d_hidden)
.add("bias", &bias)
.optional()
}
}
impl BiGruConfig {
pub fn init<B: Backend>(&self, device: &B::Device) -> BiGru<B> {
let base_config = GruConfig::new(self.d_input, self.d_hidden, self.bias)
.with_initializer(self.initializer.clone())
.with_reset_after(self.reset_after)
.with_gate_activation(self.gate_activation.clone())
.with_hidden_activation(self.hidden_activation.clone())
.with_clip(self.clip);
BiGru {
forward: base_config.clone().init(device),
reverse: base_config.init(device),
d_hidden: self.d_hidden,
batch_first: self.batch_first,
}
}
}
impl<B: Backend> BiGru<B> {
pub fn forward(
&self,
batched_input: Tensor<B, 3>,
state: Option<Tensor<B, 3>>,
) -> (Tensor<B, 3>, Tensor<B, 3>) {
let batched_input = if self.batch_first {
batched_input
} else {
batched_input.swap_dims(0, 1)
};
let device = batched_input.clone().device();
let [batch_size, seq_length, _] = batched_input.shape().dims();
let [init_state_forward, init_state_reverse] = match state {
Some(state) => {
let hidden_state_forward = state
.clone()
.slice([0..1, 0..batch_size, 0..self.d_hidden])
.squeeze_dim(0);
let hidden_state_reverse = state
.slice([1..2, 0..batch_size, 0..self.d_hidden])
.squeeze_dim(0);
[Some(hidden_state_forward), Some(hidden_state_reverse)]
}
None => [None, None],
};
let (batched_hidden_state_forward, final_state_forward) = self.forward.forward_iter(
batched_input.clone().iter_dim(1).zip(0..seq_length),
init_state_forward,
batch_size,
seq_length,
&device,
);
let (batched_hidden_state_reverse, final_state_reverse) = self.reverse.forward_iter(
batched_input.iter_dim(1).rev().zip((0..seq_length).rev()),
init_state_reverse,
batch_size,
seq_length,
&device,
);
let output = Tensor::cat(
[batched_hidden_state_forward, batched_hidden_state_reverse].to_vec(),
2,
);
let output = if self.batch_first {
output
} else {
output.swap_dims(0, 1)
};
let state = Tensor::stack([final_state_forward, final_state_reverse].to_vec(), 0);
(output, state)
}
}
#[cfg(test)]
mod tests;