use candle_core::{Result as CandleResult, Tensor};
use candle_nn::{Conv2d, Conv2dConfig, Linear, Module, VarBuilder, conv2d, linear};
use super::bilstm::BiLstm;
use super::ops::max_pool2d_padded;
const FEATURE_CHANNELS: [(usize, usize); 7] = [
(1, 32),
(32, 64),
(64, 128),
(128, 128),
(128, 256),
(256, 256),
(256, 256),
];
const HIDDEN_SIZE: usize = 256;
const FEATURE_SIZE: usize = 256;
pub(super) struct CrnnNet {
features: Vec<Conv2d>,
recurrent: [(BiLstm, Linear); 2],
prediction: Linear,
}
impl CrnnNet {
pub(super) fn new(vb: VarBuilder) -> CandleResult<Self> {
let padded = Conv2dConfig {
padding: 1,
..Conv2dConfig::default()
};
let mut features = Vec::with_capacity(FEATURE_CHANNELS.len());
for (index, (in_channels, out_channels)) in FEATURE_CHANNELS.iter().enumerate() {
let (kernel, config) = if index == FEATURE_CHANNELS.len() - 1 {
(2, Conv2dConfig::default())
} else {
(3, padded)
};
features.push(conv2d(
*in_channels,
*out_channels,
kernel,
config,
vb.pp(format!("conv.{index}")),
)?);
}
let recurrent = [
(
BiLstm::new(FEATURE_SIZE, HIDDEN_SIZE, vb.pp("rnn.0"))?,
linear(2 * HIDDEN_SIZE, FEATURE_SIZE, vb.pp("linear.0"))?,
),
(
BiLstm::new(FEATURE_SIZE, HIDDEN_SIZE, vb.pp("rnn.1"))?,
linear(2 * HIDDEN_SIZE, FEATURE_SIZE, vb.pp("linear.1"))?,
),
];
let prediction = linear_from_shape(&vb, "linear.2")?;
Ok(Self {
features,
recurrent,
prediction,
})
}
pub(super) fn forward(&self, input: &Tensor) -> CandleResult<Tensor> {
let sequence = self.extract_features(input)?;
let mut hidden = sequence;
for (recurrent, projection) in &self.recurrent {
hidden = projection.forward(&recurrent.forward(&hidden)?)?;
}
self.prediction.forward(&hidden)
}
fn extract_features(&self, input: &Tensor) -> CandleResult<Tensor> {
let pool = |tensor: &Tensor, window: (usize, usize)| max_pool2d_padded(tensor, window, window, (0, 0));
let mut hidden = self.features[0].forward(input)?.relu()?;
hidden = pool(&hidden, (2, 2))?;
hidden = self.features[1].forward(&hidden)?.relu()?;
hidden = pool(&hidden, (2, 2))?;
hidden = self.features[2].forward(&hidden)?.relu()?;
hidden = self.features[3].forward(&hidden)?.relu()?;
hidden = pool(&hidden, (2, 1))?;
hidden = self.features[4].forward(&hidden)?.relu()?;
hidden = self.features[5].forward(&hidden)?.relu()?;
hidden = pool(&hidden, (2, 1))?;
hidden = self.features[6].forward(&hidden)?.relu()?;
hidden
.permute((0, 3, 1, 2))?
.contiguous()?
.mean_keepdim(3)?
.squeeze(3)?
.contiguous()
}
}
fn linear_from_shape(vb: &VarBuilder, path: &str) -> CandleResult<Linear> {
let weight = vb.pp(path).get_unchecked("weight")?;
let classes = weight.dim(0)?;
linear(FEATURE_SIZE, classes, vb.pp(path))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn should_describe_seven_feature_convolutions() {
assert_eq!(
FEATURE_CHANNELS.len(),
7,
"the exported recognizers carry seven Conv nodes"
);
for window in FEATURE_CHANNELS.windows(2) {
assert_eq!(
window[0].1, window[1].0,
"each convolution must consume the previous one's channels"
);
}
assert_eq!(
FEATURE_CHANNELS[0].0, 1,
"the recognizer reads a single grayscale plane"
);
assert_eq!(
FEATURE_CHANNELS[FEATURE_CHANNELS.len() - 1].1,
FEATURE_SIZE,
"the feature width entering the recurrent stack must match the last convolution"
);
}
}