Skip to main content

candle_transformers/models/
marian.rs

1//! Marian Neural Machine Translation
2//!
3//! See "Marian: Fast Neural Machine Translation in C++" Junczys-Dowmunt et al. 2018
4//! - [ACL Anthology](https://aclanthology.org/P18-4020/)
5//! - [GitHub](https://github.com/marian-nmt/marian)
6//!
7use super::with_tracing::{linear, Embedding, Linear};
8use candle::{Result, Tensor};
9use candle_nn::{layer_norm, LayerNorm, VarBuilder};
10
11#[derive(Debug, Clone, serde::Deserialize)]
12pub struct Config {
13    pub vocab_size: usize,
14    pub decoder_vocab_size: Option<usize>,
15    pub max_position_embeddings: usize,
16    pub encoder_layers: usize,
17    pub encoder_ffn_dim: usize,
18    pub encoder_attention_heads: usize,
19    pub decoder_layers: usize,
20    pub decoder_ffn_dim: usize,
21    pub decoder_attention_heads: usize,
22    pub use_cache: bool,
23    pub is_encoder_decoder: bool,
24    pub activation_function: candle_nn::Activation,
25    pub d_model: usize,
26    pub decoder_start_token_id: u32,
27    pub scale_embedding: bool,
28    pub pad_token_id: u32,
29    pub eos_token_id: u32,
30    pub forced_eos_token_id: u32,
31    pub share_encoder_decoder_embeddings: bool,
32}
33
34impl Config {
35    // https://huggingface.co/Helsinki-NLP/opus-mt-tc-big-fr-en/blob/main/config.json
36    pub fn opus_mt_tc_big_fr_en() -> Self {
37        Self {
38            activation_function: candle_nn::Activation::Relu,
39            d_model: 1024,
40            decoder_attention_heads: 16,
41            decoder_ffn_dim: 4096,
42            decoder_layers: 6,
43            decoder_start_token_id: 53016,
44            decoder_vocab_size: Some(53017),
45            encoder_attention_heads: 16,
46            encoder_ffn_dim: 4096,
47            encoder_layers: 6,
48            eos_token_id: 43311,
49            forced_eos_token_id: 43311,
50            is_encoder_decoder: true,
51            max_position_embeddings: 1024,
52            pad_token_id: 53016,
53            scale_embedding: true,
54            share_encoder_decoder_embeddings: true,
55            use_cache: true,
56            vocab_size: 53017,
57        }
58    }
59
60    // https://huggingface.co/Helsinki-NLP/opus-mt-fr-en/blob/main/config.json
61    pub fn opus_mt_fr_en() -> Self {
62        Self {
63            activation_function: candle_nn::Activation::Swish,
64            d_model: 512,
65            decoder_attention_heads: 8,
66            decoder_ffn_dim: 2048,
67            decoder_layers: 6,
68            decoder_start_token_id: 59513,
69            decoder_vocab_size: Some(59514),
70            encoder_attention_heads: 8,
71            encoder_ffn_dim: 2048,
72            encoder_layers: 6,
73            eos_token_id: 0,
74            forced_eos_token_id: 0,
75            is_encoder_decoder: true,
76            max_position_embeddings: 512,
77            pad_token_id: 59513,
78            scale_embedding: true,
79            share_encoder_decoder_embeddings: true,
80            use_cache: true,
81            vocab_size: 59514,
82        }
83    }
84
85    pub fn opus_mt_en_zh() -> Self {
86        Self {
87            activation_function: candle_nn::Activation::Swish,
88            d_model: 512,
89            decoder_attention_heads: 8,
90            decoder_ffn_dim: 2048,
91            decoder_layers: 6,
92            decoder_start_token_id: 65000,
93            decoder_vocab_size: Some(65001),
94            encoder_attention_heads: 8,
95            encoder_ffn_dim: 2048,
96            encoder_layers: 6,
97            eos_token_id: 0,
98            forced_eos_token_id: 0,
99            is_encoder_decoder: true,
100            max_position_embeddings: 512,
101            pad_token_id: 65000,
102            scale_embedding: true,
103            share_encoder_decoder_embeddings: true,
104            use_cache: true,
105            vocab_size: 65001,
106        }
107    }
108
109    pub fn opus_mt_en_hi() -> Self {
110        Self {
111            activation_function: candle_nn::Activation::Swish,
112            d_model: 512,
113            decoder_attention_heads: 8,
114            decoder_ffn_dim: 2048,
115            decoder_layers: 6,
116            decoder_start_token_id: 61949,
117            decoder_vocab_size: Some(61950),
118            encoder_attention_heads: 8,
119            encoder_ffn_dim: 2048,
120            encoder_layers: 6,
121            eos_token_id: 0,
122            forced_eos_token_id: 0,
123            is_encoder_decoder: true,
124            max_position_embeddings: 512,
125            pad_token_id: 61949,
126            scale_embedding: true,
127            share_encoder_decoder_embeddings: true,
128            use_cache: true,
129            vocab_size: 61950,
130        }
131    }
132
133    pub fn opus_mt_en_es() -> Self {
134        Self {
135            activation_function: candle_nn::Activation::Swish,
136            d_model: 512,
137            decoder_attention_heads: 8,
138            decoder_ffn_dim: 2048,
139            decoder_layers: 6,
140            decoder_start_token_id: 65000,
141            decoder_vocab_size: Some(65001),
142            encoder_attention_heads: 8,
143            encoder_ffn_dim: 2048,
144            encoder_layers: 6,
145            eos_token_id: 0,
146            forced_eos_token_id: 0,
147            is_encoder_decoder: true,
148            max_position_embeddings: 512,
149            pad_token_id: 65000,
150            scale_embedding: true,
151            share_encoder_decoder_embeddings: true,
152            use_cache: true,
153            vocab_size: 65001,
154        }
155    }
156
157    pub fn opus_mt_en_fr() -> Self {
158        Self {
159            activation_function: candle_nn::Activation::Swish,
160            d_model: 512,
161            decoder_attention_heads: 8,
162            decoder_ffn_dim: 2048,
163            decoder_layers: 6,
164            decoder_start_token_id: 59513,
165            decoder_vocab_size: Some(59514),
166            encoder_attention_heads: 8,
167            encoder_ffn_dim: 2048,
168            encoder_layers: 6,
169            eos_token_id: 0,
170            forced_eos_token_id: 0,
171            is_encoder_decoder: true,
172            max_position_embeddings: 512,
173            pad_token_id: 59513,
174            scale_embedding: true,
175            share_encoder_decoder_embeddings: true,
176            use_cache: true,
177            vocab_size: 59514,
178        }
179    }
180
181    pub fn opus_mt_en_ru() -> Self {
182        Self {
183            activation_function: candle_nn::Activation::Swish,
184            d_model: 512,
185            decoder_attention_heads: 8,
186            decoder_ffn_dim: 2048,
187            decoder_layers: 6,
188            decoder_start_token_id: 62517,
189            decoder_vocab_size: Some(62518),
190            encoder_attention_heads: 8,
191            encoder_ffn_dim: 2048,
192            encoder_layers: 6,
193            eos_token_id: 0,
194            forced_eos_token_id: 0,
195            is_encoder_decoder: true,
196            max_position_embeddings: 512,
197            pad_token_id: 62517,
198            scale_embedding: true,
199            share_encoder_decoder_embeddings: true,
200            use_cache: true,
201            vocab_size: 62518,
202        }
203    }
204}
205
206#[derive(Debug, Clone)]
207struct SinusoidalPositionalEmbedding {
208    emb: Embedding,
209}
210
211impl SinusoidalPositionalEmbedding {
212    fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
213        let dev = vb.device();
214        let dtype = vb.dtype();
215        let num_positions = cfg.max_position_embeddings;
216        let dim = cfg.d_model;
217        let inv_freq: Vec<_> = (0..dim)
218            .step_by(2)
219            .map(|i| 1f32 / 10000f32.powf(i as f32 / dim as f32))
220            .collect();
221        let inv_freq_len = inv_freq.len();
222        let inv_freq = Tensor::from_vec(inv_freq, (1, inv_freq_len), dev)?.to_dtype(dtype)?;
223        let t = Tensor::arange(0u32, num_positions as u32, dev)?
224            .to_dtype(dtype)?
225            .reshape((num_positions, 1))?;
226        let freqs = t.matmul(&inv_freq)?;
227        let sin = freqs.sin()?;
228        let cos = freqs.cos()?;
229        let weights = Tensor::cat(&[&sin, &cos], 1)?.contiguous()?;
230        let emb = Embedding::from_weights(weights)?;
231        Ok(Self { emb })
232    }
233
234    fn forward(&self, input_ids: &Tensor, past_kv_len: usize) -> Result<Tensor> {
235        let seq_len = input_ids.dim(1)?;
236        Tensor::arange(
237            past_kv_len as u32,
238            (past_kv_len + seq_len) as u32,
239            input_ids.device(),
240        )?
241        .apply(&self.emb)
242    }
243}
244
245#[derive(Debug, Clone)]
246struct Attention {
247    q_proj: Linear,
248    k_proj: Linear,
249    v_proj: Linear,
250    out_proj: Linear,
251    scaling: f64,
252    num_heads: usize,
253    head_dim: usize,
254    kv_cache: Option<(Tensor, Tensor)>,
255    is_decoder: bool,
256}
257
258impl Attention {
259    fn new(cfg: &Config, is_decoder: bool, vb: VarBuilder) -> Result<Self> {
260        let num_heads = if is_decoder {
261            cfg.decoder_attention_heads
262        } else {
263            cfg.encoder_attention_heads
264        };
265        let embed_dim = cfg.d_model;
266        let head_dim = embed_dim / num_heads;
267        let scaling = (head_dim as f64).powf(-0.5);
268        let q_proj = linear(embed_dim, embed_dim, vb.pp("q_proj"))?;
269        let k_proj = linear(embed_dim, embed_dim, vb.pp("k_proj"))?;
270        let v_proj = linear(embed_dim, embed_dim, vb.pp("v_proj"))?;
271        let out_proj = linear(embed_dim, embed_dim, vb.pp("out_proj"))?;
272        Ok(Self {
273            q_proj,
274            k_proj,
275            v_proj,
276            out_proj,
277            scaling,
278            num_heads,
279            head_dim,
280            kv_cache: None,
281            is_decoder,
282        })
283    }
284
285    fn _shape(&self, tensor: &Tensor, bsz: usize) -> Result<Tensor> {
286        tensor
287            .reshape((bsz, (), self.num_heads, self.head_dim))?
288            .transpose(1, 2)?
289            .contiguous()
290    }
291
292    fn forward(
293        &mut self,
294        xs: &Tensor,
295        kv_states: Option<&Tensor>,
296        attn_mask: Option<&Tensor>,
297    ) -> Result<Tensor> {
298        let (b_sz, tgt_len, _) = xs.dims3()?;
299        let query_states = (xs.apply(&self.q_proj)? * self.scaling)?;
300        let (key_states, value_states) = match kv_states {
301            None => {
302                let key_states = self._shape(&xs.apply(&self.k_proj)?, b_sz)?;
303                let value_states = self._shape(&xs.apply(&self.v_proj)?, b_sz)?;
304                if self.is_decoder {
305                    let kv_states = match &self.kv_cache {
306                        None => (key_states, value_states),
307                        Some((p_key_states, p_value_states)) => {
308                            let key_states = Tensor::cat(&[p_key_states, &key_states], 2)?;
309                            let value_states = Tensor::cat(&[p_value_states, &value_states], 2)?;
310                            (key_states, value_states)
311                        }
312                    };
313                    self.kv_cache = Some(kv_states.clone());
314                    kv_states
315                } else {
316                    (key_states, value_states)
317                }
318            }
319            Some(kv_states) => {
320                let key_states = self._shape(&kv_states.apply(&self.k_proj)?, b_sz)?;
321                let value_states = self._shape(&kv_states.apply(&self.v_proj)?, b_sz)?;
322                (key_states, value_states)
323            }
324        };
325        let proj_shape = (b_sz * self.num_heads, (), self.head_dim);
326        let query_states = self._shape(&query_states, b_sz)?.reshape(proj_shape)?;
327        let key_states = key_states.reshape(proj_shape)?;
328        let value_states = value_states.reshape(proj_shape)?;
329        let attn_weights = query_states.matmul(&key_states.transpose(1, 2)?)?;
330        let attn_weights = match attn_mask {
331            None => attn_weights,
332            Some(attn_mask) => attn_weights.broadcast_add(attn_mask)?,
333        };
334        let attn_probs = candle_nn::ops::softmax_last_dim(&attn_weights)?;
335        let attn_output = attn_probs.matmul(&value_states)?;
336        attn_output
337            .reshape((b_sz, self.num_heads, tgt_len, self.head_dim))?
338            .transpose(1, 2)?
339            .reshape((b_sz, tgt_len, self.head_dim * self.num_heads))?
340            .apply(&self.out_proj)
341    }
342
343    fn reset_kv_cache(&mut self) {
344        self.kv_cache = None
345    }
346}
347
348#[derive(Debug, Clone)]
349struct EncoderLayer {
350    self_attn: Attention,
351    self_attn_layer_norm: LayerNorm,
352    activation_fn: candle_nn::Activation,
353    fc1: Linear,
354    fc2: Linear,
355    final_layer_norm: LayerNorm,
356}
357
358impl EncoderLayer {
359    fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
360        let self_attn = Attention::new(cfg, true, vb.pp("self_attn"))?;
361        let self_attn_layer_norm = layer_norm(cfg.d_model, 1e-5, vb.pp("self_attn_layer_norm"))?;
362        let fc1 = linear(cfg.d_model, cfg.encoder_ffn_dim, vb.pp("fc1"))?;
363        let fc2 = linear(cfg.encoder_ffn_dim, cfg.d_model, vb.pp("fc2"))?;
364        let final_layer_norm = layer_norm(cfg.d_model, 1e-5, vb.pp("final_layer_norm"))?;
365        Ok(Self {
366            self_attn,
367            self_attn_layer_norm,
368            activation_fn: cfg.activation_function,
369            fc1,
370            fc2,
371            final_layer_norm,
372        })
373    }
374
375    fn forward(&mut self, xs: &Tensor) -> Result<Tensor> {
376        let residual = xs;
377        let xs = (self.self_attn.forward(xs, None, None)? + residual)?
378            .apply(&self.self_attn_layer_norm)?;
379        let residual = &xs;
380        let xs = xs
381            .apply(&self.fc1)?
382            .apply(&self.activation_fn)?
383            .apply(&self.fc2)?;
384        (xs + residual)?.apply(&self.final_layer_norm)
385    }
386
387    fn reset_kv_cache(&mut self) {
388        self.self_attn.reset_kv_cache()
389    }
390}
391
392#[derive(Debug, Clone)]
393struct DecoderLayer {
394    self_attn: Attention,
395    self_attn_layer_norm: LayerNorm,
396    activation_fn: candle_nn::Activation,
397    encoder_attn: Attention,
398    encoder_attn_layer_norm: LayerNorm,
399    fc1: Linear,
400    fc2: Linear,
401    final_layer_norm: LayerNorm,
402}
403
404impl DecoderLayer {
405    fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
406        let self_attn = Attention::new(cfg, true, vb.pp("self_attn"))?;
407        let self_attn_layer_norm = layer_norm(cfg.d_model, 1e-5, vb.pp("self_attn_layer_norm"))?;
408        let encoder_attn = Attention::new(cfg, true, vb.pp("encoder_attn"))?;
409        let encoder_attn_layer_norm =
410            layer_norm(cfg.d_model, 1e-5, vb.pp("encoder_attn_layer_norm"))?;
411        let fc1 = linear(cfg.d_model, cfg.decoder_ffn_dim, vb.pp("fc1"))?;
412        let fc2 = linear(cfg.decoder_ffn_dim, cfg.d_model, vb.pp("fc2"))?;
413        let final_layer_norm = layer_norm(cfg.d_model, 1e-5, vb.pp("final_layer_norm"))?;
414        Ok(Self {
415            self_attn,
416            self_attn_layer_norm,
417            activation_fn: cfg.activation_function,
418            encoder_attn,
419            encoder_attn_layer_norm,
420            fc1,
421            fc2,
422            final_layer_norm,
423        })
424    }
425
426    fn forward(
427        &mut self,
428        xs: &Tensor,
429        encoder_xs: Option<&Tensor>,
430        attn_mask: &Tensor,
431    ) -> Result<Tensor> {
432        let residual = xs;
433        let xs = (self.self_attn.forward(xs, None, Some(attn_mask))? + residual)?
434            .apply(&self.self_attn_layer_norm)?;
435        let xs = match encoder_xs {
436            None => xs,
437            Some(encoder_xs) => {
438                let residual = &xs;
439                let xs = self.encoder_attn.forward(&xs, Some(encoder_xs), None)?;
440                (residual + xs)?.apply(&self.encoder_attn_layer_norm)?
441            }
442        };
443        let residual = &xs;
444        let xs = xs
445            .apply(&self.fc1)?
446            .apply(&self.activation_fn)?
447            .apply(&self.fc2)?;
448        let xs = (xs + residual)?.apply(&self.final_layer_norm)?;
449        Ok(xs)
450    }
451
452    fn reset_kv_cache(&mut self) {
453        self.self_attn.reset_kv_cache();
454        self.encoder_attn.reset_kv_cache()
455    }
456}
457
458#[derive(Debug, Clone)]
459pub struct Encoder {
460    embed_tokens: Embedding,
461    embed_positions: SinusoidalPositionalEmbedding,
462    layers: Vec<EncoderLayer>,
463    embed_scale: Option<f64>,
464}
465
466impl Encoder {
467    fn new(cfg: &Config, embed_tokens: &Embedding, vb: VarBuilder) -> Result<Self> {
468        let embed_positions = SinusoidalPositionalEmbedding::new(cfg, vb.pp("embed_positions"))?;
469        let mut layers = Vec::with_capacity(cfg.encoder_layers);
470        let vb_l = vb.pp("layers");
471        for idx in 0..cfg.encoder_layers {
472            let layer = EncoderLayer::new(cfg, vb_l.pp(idx))?;
473            layers.push(layer)
474        }
475        let embed_scale = if cfg.scale_embedding {
476            Some((cfg.d_model as f64).sqrt())
477        } else {
478            None
479        };
480        Ok(Self {
481            embed_tokens: embed_tokens.clone(),
482            embed_positions,
483            layers,
484            embed_scale,
485        })
486    }
487
488    pub fn forward(&mut self, xs: &Tensor, past_kv_len: usize) -> Result<Tensor> {
489        let xs = xs.apply(&self.embed_tokens)?;
490        let xs = match self.embed_scale {
491            None => xs,
492            Some(scale) => (xs * scale)?,
493        };
494        let embed_pos = self
495            .embed_positions
496            .forward(&xs, past_kv_len)?
497            .unsqueeze(0)?;
498        let mut xs = xs.broadcast_add(&embed_pos)?;
499        for layer in self.layers.iter_mut() {
500            xs = layer.forward(&xs)?
501        }
502        Ok(xs)
503    }
504
505    pub fn reset_kv_cache(&mut self) {
506        for layer in self.layers.iter_mut() {
507            layer.reset_kv_cache()
508        }
509    }
510}
511
512#[derive(Debug, Clone)]
513pub struct Decoder {
514    embed_tokens: Embedding,
515    embed_positions: SinusoidalPositionalEmbedding,
516    layers: Vec<DecoderLayer>,
517    embed_scale: Option<f64>,
518}
519
520impl Decoder {
521    fn new(cfg: &Config, embed_tokens: &Embedding, vb: VarBuilder) -> Result<Self> {
522        let embed_positions = SinusoidalPositionalEmbedding::new(cfg, vb.pp("embed_positions"))?;
523        let mut layers = Vec::with_capacity(cfg.decoder_layers);
524        let vb_l = vb.pp("layers");
525        for idx in 0..cfg.decoder_layers {
526            let layer = DecoderLayer::new(cfg, vb_l.pp(idx))?;
527            layers.push(layer)
528        }
529        let embed_scale = if cfg.scale_embedding {
530            Some((cfg.d_model as f64).sqrt())
531        } else {
532            None
533        };
534        Ok(Self {
535            embed_tokens: embed_tokens.clone(),
536            embed_positions,
537            layers,
538            embed_scale,
539        })
540    }
541
542    pub fn forward(
543        &mut self,
544        xs: &Tensor,
545        encoder_xs: Option<&Tensor>,
546        past_kv_len: usize,
547        attn_mask: &Tensor,
548    ) -> Result<Tensor> {
549        let xs = xs.apply(&self.embed_tokens)?;
550        let xs = match self.embed_scale {
551            None => xs,
552            Some(scale) => (xs * scale)?,
553        };
554        let embed_pos = self
555            .embed_positions
556            .forward(&xs, past_kv_len)?
557            .unsqueeze(0)?;
558        let mut xs = xs.broadcast_add(&embed_pos)?;
559        for layer in self.layers.iter_mut() {
560            xs = layer.forward(&xs, encoder_xs, attn_mask)?;
561        }
562        Ok(xs)
563    }
564
565    pub fn reset_kv_cache(&mut self) {
566        for layer in self.layers.iter_mut() {
567            layer.reset_kv_cache()
568        }
569    }
570}
571
572#[derive(Debug, Clone)]
573struct Model {
574    shared: Embedding,
575    encoder: Encoder,
576    decoder: Decoder,
577}
578
579impl Model {
580    fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
581        let shared = Embedding::new(cfg.vocab_size, cfg.d_model, vb.pp("shared"))?;
582        let encoder = Encoder::new(cfg, &shared, vb.pp("encoder"))?;
583        let decoder = Decoder::new(cfg, &shared, vb.pp("decoder"))?;
584        Ok(Self {
585            shared,
586            encoder,
587            decoder,
588        })
589    }
590
591    fn reset_kv_cache(&mut self) {
592        self.encoder.reset_kv_cache();
593        self.decoder.reset_kv_cache();
594    }
595}
596
597#[derive(Debug, Clone)]
598pub struct MTModel {
599    model: Model,
600    lm_head: Linear,
601    final_logits_bias: Tensor,
602}
603
604impl MTModel {
605    pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
606        let target_vocab_size = cfg.decoder_vocab_size.unwrap_or(cfg.vocab_size);
607        let final_logits_bias = vb.get((1, target_vocab_size), "final_logits_bias")?;
608        let model = Model::new(cfg, vb.pp("model"))?;
609        let lm_head = Linear::from_weights(model.shared.embeddings().clone(), None);
610        Ok(Self {
611            model,
612            lm_head,
613            final_logits_bias,
614        })
615    }
616
617    pub fn encoder(&mut self) -> &mut Encoder {
618        &mut self.model.encoder
619    }
620
621    pub fn decoder(&mut self) -> &mut Decoder {
622        &mut self.model.decoder
623    }
624
625    pub fn decode(
626        &mut self,
627        xs: &Tensor,
628        encoder_xs: &Tensor,
629        past_kv_len: usize,
630    ) -> Result<Tensor> {
631        let seq_len = xs.dim(1)?;
632        let mask: Vec<_> = (0..seq_len)
633            .flat_map(|i| (0..seq_len).map(move |j| if j > i { f32::NEG_INFINITY } else { 0f32 }))
634            .collect();
635        let mask = Tensor::from_vec(mask, (seq_len, seq_len), xs.device())?;
636        self.model
637            .decoder
638            .forward(xs, Some(encoder_xs), past_kv_len, &mask)?
639            .apply(&self.lm_head)?
640            .broadcast_add(&self.final_logits_bias)
641    }
642
643    pub fn reset_kv_cache(&mut self) {
644        self.model.reset_kv_cache();
645    }
646}