Skip to main content

candle_transformers/models/
xlm_roberta.rs

1use crate::models::with_tracing::{linear, Linear};
2use candle::{DType, Module, Result, Tensor};
3use candle_nn::{
4    embedding, layer_norm, ops::softmax_last_dim, Activation, Embedding, LayerNorm, VarBuilder,
5};
6
7#[derive(Debug, Clone, serde::Deserialize)]
8pub struct Config {
9    pub hidden_size: usize,
10    pub layer_norm_eps: f64,
11    pub attention_probs_dropout_prob: f32,
12    pub hidden_dropout_prob: f32,
13    pub num_attention_heads: usize,
14    pub position_embedding_type: String,
15    pub intermediate_size: usize,
16    pub hidden_act: Activation,
17    pub num_hidden_layers: usize,
18    pub vocab_size: usize,
19    pub max_position_embeddings: usize,
20    pub type_vocab_size: usize,
21    pub pad_token_id: u32,
22}
23
24struct XLMRobertaEmbeddings {
25    word_embeddings: Embedding,
26    position_embeddings: Option<Embedding>,
27    token_type_embeddings: Embedding,
28    layer_norm: LayerNorm,
29    padding_idx: u32,
30    span: tracing::Span,
31}
32
33impl XLMRobertaEmbeddings {
34    fn load(vb: VarBuilder, config: &Config) -> Result<Self> {
35        let word_embeddings = embedding(
36            config.vocab_size,
37            config.hidden_size,
38            vb.pp("word_embeddings"),
39        )?;
40        let position_embeddings = embedding(
41            config.max_position_embeddings,
42            config.hidden_size,
43            vb.pp("position_embeddings"),
44        )?;
45        let token_type_embeddings = embedding(
46            config.type_vocab_size,
47            config.hidden_size,
48            vb.pp("token_type_embeddings"),
49        )?;
50        let layer_norm = layer_norm(
51            config.hidden_size,
52            config.layer_norm_eps,
53            vb.pp("LayerNorm"),
54        )?;
55        Ok(Self {
56            word_embeddings,
57            position_embeddings: Some(position_embeddings),
58            token_type_embeddings,
59            layer_norm,
60            padding_idx: config.pad_token_id,
61            span: tracing::span!(tracing::Level::TRACE, "embeddings"),
62        })
63    }
64
65    fn forward(&self, input_ids: &Tensor, token_type_ids: &Tensor) -> Result<Tensor> {
66        let _enter = self.span.enter();
67        let (_bsize, _) = input_ids.dims2()?;
68        let input_embeddings = self.word_embeddings.forward(input_ids)?;
69        let token_type_embeddings = self.token_type_embeddings.forward(token_type_ids)?;
70        let mut embeddings = (&input_embeddings + token_type_embeddings)?;
71        if let Some(position_embeddings) = &self.position_embeddings {
72            let mask = input_ids
73                .ne(self.padding_idx)?
74                .to_dtype(input_embeddings.dtype())?;
75            let cumsum = mask.cumsum(1)?;
76            let position_ids = (cumsum * mask)?
77                .broadcast_add(
78                    &Tensor::try_from(self.padding_idx)?
79                        .to_dtype(input_embeddings.dtype())?
80                        .to_device(input_embeddings.device())?,
81                )?
82                .to_dtype(candle::DType::U32)?;
83            embeddings = embeddings.broadcast_add(&position_embeddings.forward(&position_ids)?)?;
84        }
85        let embeddings = self.layer_norm.forward(&embeddings)?;
86        Ok(embeddings)
87    }
88}
89
90struct XLMRobertaSelfAttention {
91    num_attention_heads: usize,
92    attention_head_size: usize,
93    all_head_size: usize,
94    query: Linear,
95    key: Linear,
96    value: Linear,
97}
98
99impl XLMRobertaSelfAttention {
100    fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
101        let attention_head_size = cfg.hidden_size / cfg.num_attention_heads;
102        let all_head_size = cfg.num_attention_heads * attention_head_size;
103        Ok(Self {
104            num_attention_heads: cfg.num_attention_heads,
105            attention_head_size,
106            all_head_size,
107            query: linear(cfg.hidden_size, all_head_size, vb.pp("query"))?,
108            key: linear(cfg.hidden_size, all_head_size, vb.pp("key"))?,
109            value: linear(cfg.hidden_size, all_head_size, vb.pp("value"))?,
110        })
111    }
112
113    fn transpose_for_scores(&self, x: &Tensor) -> Result<Tensor> {
114        let mut new_x_shape = x.dims().to_vec();
115        new_x_shape[2] = self.num_attention_heads;
116        new_x_shape.push(self.attention_head_size);
117        let x = x.reshape(new_x_shape)?;
118        x.permute((0, 2, 1, 3))?.contiguous()
119    }
120
121    fn forward(
122        &self,
123        hidden_states: &Tensor,
124        encoder_hidden_states: Option<&Tensor>,
125        attention_mask: &Tensor,
126        past_key_value: Option<(&Tensor, &Tensor)>,
127        encoder_attention_mask: Option<&Tensor>,
128    ) -> Result<Tensor> {
129        let mixed_query_layer = self.query.forward(hidden_states)?;
130        let is_cross_attention = encoder_hidden_states.is_some();
131        let (key_layer, value_layer, attention_mask) = if is_cross_attention {
132            if let Some((past_key, past_value)) = past_key_value {
133                let key_layer = past_key.clone();
134                let value_layer = past_value.clone();
135                let attention_mask = encoder_attention_mask.unwrap().clone();
136                (key_layer, value_layer, Some(attention_mask))
137            } else {
138                let key_layer =
139                    self.transpose_for_scores(&self.key.forward(encoder_hidden_states.unwrap())?)?;
140                let value_layer = self
141                    .transpose_for_scores(&self.value.forward(encoder_hidden_states.unwrap())?)?;
142                let attention_mask = encoder_attention_mask.unwrap();
143                (key_layer, value_layer, Some(attention_mask.clone()))
144            }
145        } else if let Some((past_key, past_value)) = past_key_value {
146            let mut key_layer = self.transpose_for_scores(&self.key.forward(hidden_states)?)?;
147            let mut value_layer = self.transpose_for_scores(&self.value.forward(hidden_states)?)?;
148            key_layer = Tensor::cat(&[past_key.clone(), key_layer], 2)?;
149            value_layer = Tensor::cat(&[past_value.clone(), value_layer], 2)?;
150            (key_layer, value_layer, Some(attention_mask.clone()))
151        } else {
152            let key_layer = self.transpose_for_scores(&self.key.forward(hidden_states)?)?;
153            let value_layer = self.transpose_for_scores(&self.value.forward(hidden_states)?)?;
154            (key_layer, value_layer, Some(attention_mask.clone()))
155        };
156
157        let query_layer = self.transpose_for_scores(&mixed_query_layer)?;
158        let mut attention_scores = query_layer.matmul(&key_layer.transpose(2, 3)?)?;
159        let scale = 1f64 / f64::sqrt(self.attention_head_size as f64);
160
161        attention_scores = (attention_scores * scale)?;
162        attention_scores = match attention_mask {
163            None => attention_scores,
164            Some(mask) => {
165                attention_scores.broadcast_add(&mask.to_dtype(attention_scores.dtype())?)?
166            }
167        };
168        let attention_probs = softmax_last_dim(&attention_scores)?;
169
170        let context_layer = attention_probs
171            .matmul(&value_layer)?
172            .permute((0, 2, 1, 3))?
173            .contiguous()?;
174        let mut new_context_layer_shape =
175            context_layer.dims()[..context_layer.dims().len() - 2].to_vec();
176        new_context_layer_shape.push(self.all_head_size);
177        let context_layer = context_layer.reshape(new_context_layer_shape)?;
178
179        Ok(context_layer)
180    }
181}
182
183struct XLMRobertaSelfOutput {
184    dense: Linear,
185    layernorm: LayerNorm,
186}
187
188impl XLMRobertaSelfOutput {
189    fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
190        let dense = linear(cfg.hidden_size, cfg.hidden_size, vb.pp("dense"))?;
191        let layernorm =
192            candle_nn::layer_norm(cfg.hidden_size, cfg.layer_norm_eps, vb.pp("LayerNorm"))?;
193        Ok(Self { dense, layernorm })
194    }
195
196    fn forward(&self, hidden_states: &Tensor, input_tensor: &Tensor) -> Result<Tensor> {
197        let hidden_states = self.dense.forward(hidden_states)?;
198        let hidden_states = self.layernorm.forward(&(hidden_states + input_tensor)?)?;
199        Ok(hidden_states)
200    }
201}
202
203struct XLMRobertaAttention {
204    output: XLMRobertaSelfOutput,
205    self_attention: XLMRobertaSelfAttention,
206}
207
208impl XLMRobertaAttention {
209    fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
210        let output = XLMRobertaSelfOutput::new(cfg, vb.pp("output"))?;
211        let self_attention = XLMRobertaSelfAttention::new(cfg, vb.pp("self"))?;
212        Ok(Self {
213            output,
214            self_attention,
215        })
216    }
217
218    fn forward(
219        &self,
220        hidden_states: &Tensor,
221        attention_mask: &Tensor,
222        encoder_hidden_states: Option<&Tensor>,
223        encoder_attention_mask: Option<&Tensor>,
224        past_key_value: Option<(&Tensor, &Tensor)>,
225    ) -> Result<(Tensor, Tensor)> {
226        let self_outputs = self.self_attention.forward(
227            hidden_states,
228            encoder_hidden_states,
229            attention_mask,
230            past_key_value,
231            encoder_attention_mask,
232        )?;
233        let attention_output = self.output.forward(&self_outputs, hidden_states)?;
234        Ok((attention_output, self_outputs))
235    }
236}
237
238struct XLMRobertaOutput {
239    dense: Linear,
240    layernorm: LayerNorm,
241}
242
243impl XLMRobertaOutput {
244    fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
245        let dense = linear(cfg.intermediate_size, cfg.hidden_size, vb.pp("dense"))?;
246        let layernorm =
247            candle_nn::layer_norm(cfg.hidden_size, cfg.layer_norm_eps, vb.pp("LayerNorm"))?;
248        Ok(Self { dense, layernorm })
249    }
250
251    fn forward(&self, hidden_states: &Tensor, input_tensor: &Tensor) -> Result<Tensor> {
252        let hidden_states = self.dense.forward(hidden_states)?;
253        let hidden_states = self.layernorm.forward(&(hidden_states + input_tensor)?)?;
254        Ok(hidden_states)
255    }
256}
257
258struct XLMRobertaIntermediate {
259    dense: Linear,
260    intermediate_act_fn: Activation,
261}
262
263impl XLMRobertaIntermediate {
264    fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
265        let dense = linear(cfg.hidden_size, cfg.intermediate_size, vb.pp("dense"))?;
266        let intermediate_act_fn = cfg.hidden_act;
267        Ok(Self {
268            dense,
269            intermediate_act_fn,
270        })
271    }
272
273    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
274        let hidden_states = self.dense.forward(hidden_states)?;
275        let hidden_states = self.intermediate_act_fn.forward(&hidden_states)?;
276        Ok(hidden_states)
277    }
278}
279
280struct XLMRobertaLayer {
281    attention: XLMRobertaAttention,
282    intermediate: XLMRobertaIntermediate,
283    output: XLMRobertaOutput,
284}
285
286impl XLMRobertaLayer {
287    fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
288        let attention = XLMRobertaAttention::new(cfg, vb.pp("attention"))?;
289        let intermediate = XLMRobertaIntermediate::new(cfg, vb.pp("intermediate"))?;
290        let output = XLMRobertaOutput::new(cfg, vb.pp("output"))?;
291        Ok(Self {
292            attention,
293            intermediate,
294            output,
295        })
296    }
297
298    fn forward(
299        &self,
300        hidden_states: &Tensor,
301        attention_mask: &Tensor,
302        encoder_hidden_states: Option<&Tensor>,
303        encoder_attention_mask: Option<&Tensor>,
304        past_key_value: Option<(&Tensor, &Tensor)>,
305    ) -> Result<(Tensor, Tensor)> {
306        let self_attention_outputs = self.attention.forward(
307            hidden_states,
308            attention_mask,
309            encoder_hidden_states,
310            encoder_attention_mask,
311            past_key_value,
312        )?;
313        let attention_output = self_attention_outputs.0;
314        let outputs = self_attention_outputs.1;
315        let intermediate_output = self.intermediate.forward(&attention_output)?;
316        let layer_output = self
317            .output
318            .forward(&intermediate_output, &attention_output)?;
319        Ok((layer_output, outputs))
320    }
321}
322
323struct XLMRobertaEncoder {
324    layers: Vec<XLMRobertaLayer>,
325}
326
327impl XLMRobertaEncoder {
328    fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
329        let layers = (0..cfg.num_hidden_layers)
330            .map(|i| XLMRobertaLayer::new(cfg, vb.pp(format!("layer.{i}"))))
331            .collect::<Result<Vec<_>>>()?;
332        Ok(Self { layers })
333    }
334
335    fn forward(
336        &self,
337        hidden_states: &Tensor,
338        attention_mask: &Tensor,
339        encoder_hidden_states: Option<&Tensor>,
340        encoder_attention_mask: Option<&Tensor>,
341        past_key_value: Option<(&Tensor, &Tensor)>,
342    ) -> Result<Tensor> {
343        let mut hidden_states = hidden_states.clone();
344        for layer_module in self.layers.iter() {
345            let layer_outputs = layer_module.forward(
346                &hidden_states,
347                attention_mask,
348                encoder_hidden_states,
349                encoder_attention_mask,
350                past_key_value,
351            )?;
352            hidden_states = layer_outputs.0;
353        }
354        Ok(hidden_states)
355    }
356}
357
358pub struct XLMRobertaModel {
359    encoder: XLMRobertaEncoder,
360    embeddings: XLMRobertaEmbeddings,
361}
362
363impl XLMRobertaModel {
364    pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
365        let encoder = XLMRobertaEncoder::new(cfg, vb.pp("encoder"))?;
366        let embeddings = XLMRobertaEmbeddings::load(vb.pp("embeddings"), cfg)?;
367        Ok(Self {
368            encoder,
369            embeddings,
370        })
371    }
372
373    pub fn forward(
374        &self,
375        input_ids: &Tensor,
376        attention_mask: &Tensor,
377        token_type_ids: &Tensor,
378        past_key_value: Option<(&Tensor, &Tensor)>,
379        encoder_hidden_states: Option<&Tensor>,
380        encoder_attention_mask: Option<&Tensor>,
381    ) -> Result<Tensor> {
382        let hidden_states = self.embeddings.forward(input_ids, token_type_ids)?;
383        let attention_mask = prepare_4d_attention_mask(attention_mask, DType::F32, None)?
384            .to_device(hidden_states.device())?;
385        let hidden_states = self.encoder.forward(
386            &hidden_states,
387            &attention_mask,
388            encoder_hidden_states,
389            encoder_attention_mask,
390            past_key_value,
391        )?;
392        Ok(hidden_states)
393    }
394}
395
396struct XLMRobertaLMHead {
397    dense: Linear,
398    layer_norm: LayerNorm,
399}
400
401impl XLMRobertaLMHead {
402    fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
403        let dense = linear(cfg.hidden_size, cfg.hidden_size, vb.pp("dense"))?;
404        let layer_norm =
405            candle_nn::layer_norm(cfg.hidden_size, cfg.layer_norm_eps, vb.pp("layer_norm"))?;
406        Ok(Self { dense, layer_norm })
407    }
408
409    fn forward(&self, hidden_states: &Tensor, shared_embeddings: &Tensor) -> Result<Tensor> {
410        let hidden_states = self.dense.forward(hidden_states)?;
411        let hidden_states = candle_nn::Activation::Gelu.forward(&hidden_states)?;
412        let hidden_states = self.layer_norm.forward(&hidden_states)?;
413        let hidden_states = hidden_states.broadcast_matmul(shared_embeddings)?;
414        Ok(hidden_states)
415    }
416}
417
418pub struct XLMRobertaForMaskedLM {
419    roberta: XLMRobertaModel,
420    lm_head: XLMRobertaLMHead,
421}
422
423impl XLMRobertaForMaskedLM {
424    pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
425        let roberta = XLMRobertaModel::new(cfg, vb.pp("roberta"))?;
426        let lm_head = XLMRobertaLMHead::new(cfg, vb.pp("lm_head"))?;
427        Ok(Self { roberta, lm_head })
428    }
429
430    pub fn forward(
431        &self,
432        input_ids: &Tensor,
433        attention_mask: &Tensor,
434        token_type_ids: &Tensor,
435        past_key_value: Option<(&Tensor, &Tensor)>,
436        encoder_hidden_states: Option<&Tensor>,
437        encoder_attention_mask: Option<&Tensor>,
438    ) -> Result<Tensor> {
439        let hidden_states = self.roberta.forward(
440            input_ids,
441            attention_mask,
442            token_type_ids,
443            past_key_value,
444            encoder_hidden_states,
445            encoder_attention_mask,
446        )?;
447        let lm_logits = self.lm_head.forward(
448            &hidden_states,
449            &self
450                .roberta
451                .embeddings
452                .word_embeddings
453                .embeddings()
454                .t()?
455                .unsqueeze(0)?,
456        )?;
457        Ok(lm_logits)
458    }
459}
460
461struct XLMRobertaClassificationHead {
462    dense: Linear,
463    out_proj: Linear,
464}
465
466impl XLMRobertaClassificationHead {
467    fn new(num_labels: usize, cfg: &Config, vb: VarBuilder) -> Result<Self> {
468        let dense = linear(cfg.hidden_size, cfg.hidden_size, vb.pp("dense"))?;
469        let out_proj = linear(cfg.hidden_size, num_labels, vb.pp("out_proj"))?;
470        Ok(Self { dense, out_proj })
471    }
472
473    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
474        let cls_states = hidden_states.get_on_dim(1, 0)?.contiguous()?;
475        let hidden_states = self.dense.forward(&cls_states)?;
476        // The activation used in the classification head is tanh, as per the original
477        // implementation.
478        // https://github.com/huggingface/transformers/blob/6e3063422c4b1c014aa60c32b9254fd2902f0f28/src/transformers/models/xlm_roberta/modeling_xlm_roberta.py#L1454
479        let hidden_states = self.out_proj.forward(&hidden_states.tanh()?)?;
480        Ok(hidden_states)
481    }
482}
483
484pub struct XLMRobertaForSequenceClassification {
485    roberta: XLMRobertaModel,
486    classifier: XLMRobertaClassificationHead,
487}
488
489impl XLMRobertaForSequenceClassification {
490    pub fn new(num_labels: usize, cfg: &Config, vb: VarBuilder) -> Result<Self> {
491        let roberta = XLMRobertaModel::new(cfg, vb.pp("roberta"))?;
492        let classifier = XLMRobertaClassificationHead::new(num_labels, cfg, vb.pp("classifier"))?;
493        Ok(Self {
494            roberta,
495            classifier,
496        })
497    }
498
499    pub fn forward(
500        &self,
501        input_ids: &Tensor,
502        attention_mask: &Tensor,
503        token_type_ids: &Tensor,
504    ) -> Result<Tensor> {
505        let hidden_states =
506            self.roberta
507                .forward(input_ids, attention_mask, token_type_ids, None, None, None)?;
508        self.classifier.forward(&hidden_states)
509    }
510}
511
512fn prepare_4d_attention_mask(
513    mask: &Tensor,
514    dtype: DType,
515    tgt_len: Option<usize>,
516) -> Result<Tensor> {
517    let bsz = mask.dim(0)?;
518    let src_len = mask.dim(1)?;
519    let tgt_len = tgt_len.unwrap_or(src_len);
520
521    let expanded_mask = mask
522        .unsqueeze(1)?
523        .unsqueeze(2)?
524        .expand((bsz, 1, tgt_len, src_len))?
525        .to_dtype(dtype)?;
526
527    let inverted_mask = (1.0 - expanded_mask)?;
528
529    (inverted_mask * get_dtype_min_val(dtype))?.to_dtype(dtype)
530}
531
532fn get_dtype_min_val(dtype: DType) -> f64 {
533    match dtype {
534        DType::F32 => f32::MIN as f64,
535        DType::F64 => f64::MIN,
536        _ => panic!("Unsupported data type"),
537    }
538}