#![allow(clippy::excessive_precision)]
use super::Reduction;
use ruda_model::config::Config;
use ruda_model::module::Module;
use ruda_model::tensor::{Int, Tensor, backend::Backend};
#[derive(Config, Debug)]
pub struct CTCLossConfig {
#[config(default = 0)]
pub blank: usize,
#[config(default = false)]
pub zero_infinity: bool,
}
impl CTCLossConfig {
pub fn init(&self) -> CTCLoss {
CTCLoss {
blank: self.blank,
zero_infinity: self.zero_infinity,
}
}
}
#[derive(Module, Clone, Debug)]
pub struct CTCLoss {
blank: usize,
zero_infinity: bool,
}
impl CTCLoss {
pub fn forward<B: Backend>(
&self,
log_probs: Tensor<B, 3>,
targets: Tensor<B, 2, Int>,
input_lengths: Tensor<B, 1, Int>,
target_lengths: Tensor<B, 1, Int>,
) -> Tensor<B, 1> {
let [max_input_length, batch_size, num_classes] = log_probs.dims();
let max_target_len = targets.dims()[1];
let input_lengths_len = input_lengths.dims()[0];
let target_lengths_len = target_lengths.dims()[0];
self.assertions(
batch_size,
num_classes,
targets.clone(),
input_lengths_len,
target_lengths_len,
);
self.length_assertions(
input_lengths.clone(),
target_lengths.clone(),
max_target_len,
max_input_length,
);
let mut loss = ruda_model::tensor::module::ctc_loss(
log_probs,
targets,
input_lengths,
target_lengths,
self.blank,
);
if self.zero_infinity {
let inf_mask = loss.clone().is_inf();
loss = loss.clone().mask_where(inf_mask, loss.clone().zeros_like());
}
loss
}
pub fn forward_with_reduction<B: Backend>(
&self,
log_probs: Tensor<B, 3>,
targets: Tensor<B, 2, Int>,
input_lengths: Tensor<B, 1, Int>,
target_lengths: Tensor<B, 1, Int>,
reduction: Reduction,
) -> Tensor<B, 1> {
let ctc_loss_tensor =
self.forward(log_probs, targets, input_lengths, target_lengths.clone());
match reduction {
Reduction::Auto | Reduction::Mean => {
let target_lengths_float = target_lengths.float();
ctc_loss_tensor.div(target_lengths_float).mean()
}
Reduction::Sum => ctc_loss_tensor.sum(),
other => panic!("{other:?} reduction is not supported"),
}
}
#[allow(unused_variables)]
fn length_assertions<B: Backend>(
&self,
input_lengths: Tensor<B, 1, Int>,
target_lengths: Tensor<B, 1, Int>,
max_target_len: usize,
max_input_length: usize,
) {
#[cfg(debug_assertions)]
{
let target_lengths_data = target_lengths.into_data();
let input_lengths_data = input_lengths.into_data();
let target_iter = target_lengths_data.iter::<i64>();
let input_iter = input_lengths_data.iter::<i64>();
for (i, (tl, il)) in target_iter.zip(input_iter).enumerate() {
assert!(tl >= 0, "target_lengths[{i}] = {tl} must be non-negative");
assert!(
tl as usize <= max_target_len,
"target_lengths[{i}] = {tl} exceeds the targets tensor width {max_target_len}"
);
assert!(
il >= tl,
"input_lengths[{i}] = {il} must be >= target_lengths[{i}] = {tl} \
(no valid CTC alignment otherwise)"
);
assert!(
il as usize <= max_input_length,
"input_lengths[{i}] = {il} exceeds the log_probs time dimension \
{max_input_length}"
);
}
}
}
fn assertions<B: Backend>(
&self,
batch_size: usize,
num_classes: usize,
targets: Tensor<B, 2, Int>,
input_lengths_len: usize,
target_lengths_len: usize,
) {
assert!(
self.blank < num_classes,
"blank index {} must be less than num_classes {}",
self.blank,
num_classes
);
assert_eq!(
targets.dims()[0],
batch_size,
"targets batch dimension {} must equal batch_size {}",
targets.dims()[0],
batch_size
);
assert_eq!(
input_lengths_len, batch_size,
"input_lengths length {} must equal batch_size {}",
input_lengths_len, batch_size
);
assert_eq!(
target_lengths_len, batch_size,
"target_lengths length {} must equal batch_size {}",
target_lengths_len, batch_size
);
}
}
#[cfg(test)]
mod tests;
#[cfg(test)]
mod pytorch_comparison_tests;