use candle_core::{Module, Tensor};
use candle_nn::VarBuilder;
use candle_transformers::models::debertav2::{Config as DebertaConfig, DebertaV2Model};
pub(crate) struct DebertaV2TokenClassifier {
deberta: DebertaV2Model,
dropout: candle_nn::Dropout,
classifier: candle_nn::Linear,
}
impl DebertaV2TokenClassifier {
pub(crate) fn load(
vb: &VarBuilder,
config: &DebertaConfig,
id2label_len: usize,
) -> candle_core::Result<Self> {
let deberta = DebertaV2Model::load(vb.clone(), config)?;
#[allow(clippy::cast_possible_truncation)]
let dropout = candle_nn::Dropout::new(config.hidden_dropout_prob as f32);
let classifier =
candle_nn::linear(config.hidden_size, id2label_len, vb.root().pp("classifier"))?;
Ok(Self {
deberta,
dropout,
classifier,
})
}
pub(crate) fn forward(
&self,
input_ids: &Tensor,
token_type_ids: Option<Tensor>,
attention_mask: Option<Tensor>,
) -> candle_core::Result<Tensor> {
let output = self
.deberta
.forward(input_ids, token_type_ids, attention_mask)?;
let output = self.dropout.forward(&output, false)?;
self.classifier.forward(&output)
}
}