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