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 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}