use candle_core::{DType, Device, Tensor as CandleTensor};
use candle_nn::{Module, VarBuilder};
use crate::runtime::{
error::RuntimeError,
session::RuntimeSession,
tensor::{Shape, Tensor, TensorData},
};
use super::conformer::ConformerEncoder;
const PRED_HIDDEN: usize = 320;
const ENC_DIM: usize = 768;
const VOCAB: usize = 34;
fn backend_err(e: impl std::fmt::Display) -> RuntimeError {
RuntimeError::InferenceFailed(e.to_string())
}
pub struct EncoderSession {
enc: ConformerEncoder,
device: Device,
}
impl EncoderSession {
pub(crate) fn new(enc: ConformerEncoder, device: Device) -> Self {
Self { enc, device }
}
}
impl RuntimeSession for EncoderSession {
fn run(&self, inputs: &[Tensor]) -> Result<Vec<Tensor>, RuntimeError> {
if inputs.is_empty() {
return Err(RuntimeError::InvalidInputCount {
expected: 1,
got: inputs.len(),
});
}
let mel = super::tensor::to_candle(&inputs[0], &self.device)?;
let out = self
.enc
.forward(&mel)
.map_err(|e| RuntimeError::InferenceFailed(e.to_string()))?;
let enc_len = out.dims().get(2).copied().ok_or_else(|| {
RuntimeError::InferenceFailed(format!(
"encoder output has unexpected rank {}",
out.rank()
))
})?;
Ok(vec![
super::tensor::from_candle(&out)?,
Tensor::new(Shape::new(vec![1]), TensorData::I64(vec![enc_len as i64]))?,
])
}
}
pub struct DecoderSession {
embed: CandleTensor,
w_ih: CandleTensor,
w_hh: CandleTensor,
b_ih: CandleTensor,
b_hh: CandleTensor,
device: Device,
}
impl DecoderSession {
pub(crate) fn load(vb: VarBuilder, device: Device) -> Result<Self, RuntimeError> {
let embed = vb
.get((VOCAB, PRED_HIDDEN), "embed.weight")
.map_err(backend_err)?;
let w_ih = vb
.get((4 * PRED_HIDDEN, PRED_HIDDEN), "lstm.w_ih")
.map_err(backend_err)?;
let w_hh = vb
.get((4 * PRED_HIDDEN, PRED_HIDDEN), "lstm.w_hh")
.map_err(backend_err)?;
let b_ih = vb.get(4 * PRED_HIDDEN, "lstm.b_ih").map_err(backend_err)?;
let b_hh = vb.get(4 * PRED_HIDDEN, "lstm.b_hh").map_err(backend_err)?;
Ok(Self {
embed,
w_ih,
w_hh,
b_ih,
b_hh,
device,
})
}
fn lstm_step(
&self,
x: &CandleTensor, h_prev: &CandleTensor, c_prev: &CandleTensor, ) -> Result<(CandleTensor, CandleTensor), RuntimeError> {
let x_col = x.reshape((PRED_HIDDEN, 1)).map_err(backend_err)?;
let h_col = h_prev.reshape((PRED_HIDDEN, 1)).map_err(backend_err)?;
let gates = self
.w_ih
.matmul(&x_col)
.map_err(backend_err)?
.add(&self.w_hh.matmul(&h_col).map_err(backend_err)?)
.map_err(backend_err)?
.reshape(4 * PRED_HIDDEN) .map_err(backend_err)?
.add(&self.b_ih)
.map_err(backend_err)?
.add(&self.b_hh)
.map_err(backend_err)?;
let i = gates.narrow(0, 0, PRED_HIDDEN).map_err(backend_err)?;
let o = gates
.narrow(0, PRED_HIDDEN, PRED_HIDDEN)
.map_err(backend_err)?;
let f = gates
.narrow(0, 2 * PRED_HIDDEN, PRED_HIDDEN)
.map_err(backend_err)?;
let cg = gates
.narrow(0, 3 * PRED_HIDDEN, PRED_HIDDEN)
.map_err(backend_err)?;
let i = candle_nn::ops::sigmoid(&i).map_err(backend_err)?;
let o = candle_nn::ops::sigmoid(&o).map_err(backend_err)?;
let f = candle_nn::ops::sigmoid(&f).map_err(backend_err)?;
let cg = cg.tanh().map_err(backend_err)?;
let c_new = f
.mul(c_prev)
.map_err(backend_err)?
.add(&i.mul(&cg).map_err(backend_err)?)
.map_err(backend_err)?;
let h_new = o
.mul(&c_new.tanh().map_err(backend_err)?)
.map_err(backend_err)?;
Ok((h_new, c_new))
}
}
impl RuntimeSession for DecoderSession {
fn run(&self, inputs: &[Tensor]) -> Result<Vec<Tensor>, RuntimeError> {
if inputs.len() != 3 {
return Err(RuntimeError::InvalidInputCount {
expected: 3,
got: inputs.len(),
});
}
let token = inputs[0]
.view()
.data()
.as_i64()
.and_then(|s| s.first().copied())
.ok_or_else(|| {
RuntimeError::InferenceFailed("decoder prev_token must be i64 [1,1]".to_string())
})?;
if !(0..VOCAB as i64).contains(&token) {
return Err(RuntimeError::InferenceFailed(format!(
"decoder prev_token {token} out of range [0,{VOCAB})"
)));
}
let h_in = super::tensor::to_candle(&inputs[1], &self.device)?
.reshape(PRED_HIDDEN)
.map_err(backend_err)?;
let c_in = super::tensor::to_candle(&inputs[2], &self.device)?
.reshape(PRED_HIDDEN)
.map_err(backend_err)?;
let x = self
.embed
.narrow(0, token as usize, 1)
.map_err(backend_err)?
.reshape(PRED_HIDDEN)
.map_err(backend_err)?;
let (h_new, c_new) = self.lstm_step(&x, &h_in, &c_in)?;
let dec = to_runtime_vec(&h_new)?;
let new_h = to_runtime_vec(&h_new)?;
let new_c = to_runtime_vec(&c_new)?;
Ok(vec![
runtime_tensor_3d(dec)?,
runtime_tensor_3d(new_h)?,
runtime_tensor_3d(new_c)?,
])
}
}
pub struct JoinerSession {
enc_proj: candle_nn::Linear, dec_proj: candle_nn::Linear, out: candle_nn::Linear, device: Device,
}
impl JoinerSession {
pub(crate) fn load(vb: VarBuilder, device: Device) -> Result<Self, RuntimeError> {
let enc_proj =
candle_nn::linear(ENC_DIM, PRED_HIDDEN, vb.pp("enc_proj")).map_err(backend_err)?;
let dec_proj =
candle_nn::linear(PRED_HIDDEN, PRED_HIDDEN, vb.pp("dec_proj")).map_err(backend_err)?;
let out = candle_nn::linear(PRED_HIDDEN, VOCAB, vb.pp("out")).map_err(backend_err)?;
Ok(Self {
enc_proj,
dec_proj,
out,
device,
})
}
}
impl RuntimeSession for JoinerSession {
fn run(&self, inputs: &[Tensor]) -> Result<Vec<Tensor>, RuntimeError> {
if inputs.len() != 2 {
return Err(RuntimeError::InvalidInputCount {
expected: 2,
got: inputs.len(),
});
}
let enc = super::tensor::to_candle(&inputs[0], &self.device)?
.reshape((1, ENC_DIM))
.map_err(backend_err)?;
let dec = super::tensor::to_candle(&inputs[1], &self.device)?
.reshape((1, PRED_HIDDEN))
.map_err(backend_err)?;
let e = self.enc_proj.forward(&enc).map_err(backend_err)?; let d = self.dec_proj.forward(&dec).map_err(backend_err)?; let j = e
.add(&d)
.map_err(backend_err)?
.relu()
.map_err(backend_err)?; let logits = self.out.forward(&j).map_err(backend_err)?; let log_probs = candle_nn::ops::log_softmax(&logits, 1).map_err(backend_err)?;
let data = to_runtime_vec(&log_probs)?;
Ok(vec![Tensor::new(
Shape::new(vec![1, 1, 1, VOCAB]),
TensorData::F32(data),
)?])
}
}
fn to_runtime_vec(t: &CandleTensor) -> Result<Vec<f32>, RuntimeError> {
t.to_dtype(DType::F32)
.map_err(backend_err)?
.flatten_all()
.map_err(backend_err)?
.to_vec1::<f32>()
.map_err(backend_err)
}
fn runtime_tensor_3d(data: Vec<f32>) -> Result<Tensor, RuntimeError> {
Tensor::new(Shape::new(vec![1, 1, PRED_HIDDEN]), TensorData::F32(data))
}