rust_bert/models/prophetnet/
attention.rs

1// Copyright 2020 The Microsoft Authors and The HuggingFace Inc. team.
2// Copyright 2020 Guillaume Becquin
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//     http://www.apache.org/licenses/LICENSE-2.0
7// Unless required by applicable law or agreed to in writing, software
8// distributed under the License is distributed on an "AS IS" BASIS,
9// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
10// See the License for the specific language governing permissions and
11// limitations under the License.
12
13use crate::common::activations::TensorFunction;
14use crate::common::dropout::Dropout;
15use crate::prophetnet::ProphetNetConfig;
16use crate::RustBertError;
17use std::borrow::Borrow;
18use tch::nn::ModuleT;
19use tch::{nn, Kind, Tensor};
20
21#[derive(Debug)]
22/// # Cache for ProphetNet attention layers
23/// Stores the cached value of key and value
24pub struct LayerState {
25    /// Cached keys
26    pub prev_key: Tensor,
27    /// Cached values
28    pub prev_value: Tensor,
29}
30
31impl Clone for LayerState {
32    fn clone(&self) -> Self {
33        LayerState {
34            prev_key: self.prev_key.copy(),
35            prev_value: self.prev_value.copy(),
36        }
37    }
38}
39
40impl LayerState {
41    pub(crate) fn reorder_cache(&mut self, new_indices: &Tensor) {
42        self.prev_key = self.prev_key.index_select(0, new_indices);
43        self.prev_value = self.prev_value.index_select(0, new_indices);
44    }
45}
46
47pub struct ProphetNetAttention {
48    key_proj: nn::Linear,
49    value_proj: nn::Linear,
50    query_proj: nn::Linear,
51    out_proj: nn::Linear,
52    dropout: Dropout,
53    attention_dropout: Dropout,
54    num_attention_heads: i64,
55    head_dim: i64,
56    output_attentions: bool,
57}
58
59impl ProphetNetAttention {
60    pub fn new<'p, P>(
61        p: P,
62        config: &ProphetNetConfig,
63        num_attention_heads: i64,
64    ) -> Result<ProphetNetAttention, RustBertError>
65    where
66        P: Borrow<nn::Path<'p>>,
67    {
68        let p = p.borrow();
69        let dropout = Dropout::new(config.dropout);
70        let attention_dropout = Dropout::new(config.attention_dropout);
71
72        if config.hidden_size % num_attention_heads != 0 {
73            return Err(RustBertError::InvalidConfigurationError(format!(
74                "Invalid number of heads for self attention, {} not a multiple of {}",
75                config.hidden_size, num_attention_heads
76            )));
77        }
78
79        let head_dim = config.hidden_size / num_attention_heads;
80
81        let key_proj = nn::linear(
82            p / "key_proj",
83            config.hidden_size,
84            config.hidden_size,
85            Default::default(),
86        );
87
88        let value_proj = nn::linear(
89            p / "value_proj",
90            config.hidden_size,
91            config.hidden_size,
92            Default::default(),
93        );
94
95        let query_proj = nn::linear(
96            p / "query_proj",
97            config.hidden_size,
98            config.hidden_size,
99            Default::default(),
100        );
101
102        let out_proj = nn::linear(
103            p / "out_proj",
104            config.hidden_size,
105            config.hidden_size,
106            Default::default(),
107        );
108
109        let output_attentions = config.output_attentions.unwrap_or(false);
110
111        Ok(ProphetNetAttention {
112            key_proj,
113            value_proj,
114            query_proj,
115            out_proj,
116            dropout,
117            attention_dropout,
118            num_attention_heads,
119            head_dim,
120            output_attentions,
121        })
122    }
123
124    fn flatten(&self, x: Tensor, dim_0: i64, bs: i64) -> Tensor {
125        x.contiguous()
126            .view((dim_0, bs * self.num_attention_heads, self.head_dim))
127            .transpose(0, 1)
128    }
129
130    pub fn forward_t(
131        &self,
132        hidden_states: &Tensor,
133        key_value_states: Option<&Tensor>,
134        attention_mask: Option<&Tensor>,
135        mut layer_state: Option<LayerState>,
136        train: bool,
137    ) -> (Tensor, Option<Tensor>, Option<LayerState>) {
138        let hidden_states_size = hidden_states.size();
139        let (sequence_length, batch_size, hidden_size) = (
140            hidden_states_size[0],
141            hidden_states_size[1],
142            hidden_states_size[2],
143        );
144        let is_cross_attention = key_value_states.is_some();
145        let query_states = hidden_states.apply(&self.query_proj) / (self.head_dim as f64).sqrt();
146        let query_states = self.flatten(query_states, sequence_length, batch_size);
147
148        let (key_states, value_states) = if !is_cross_attention {
149            let key_states = self.flatten(hidden_states.apply(&self.key_proj), -1, batch_size);
150            let value_states = self.flatten(hidden_states.apply(&self.value_proj), -1, batch_size);
151            (key_states, value_states)
152        } else if layer_state.is_none() {
153            let key_states = self.flatten(
154                key_value_states.unwrap().apply(&self.key_proj),
155                -1,
156                batch_size,
157            );
158            let value_states = self.flatten(
159                key_value_states.unwrap().apply(&self.value_proj),
160                -1,
161                batch_size,
162            );
163            (key_states, value_states)
164        } else {
165            let past_state = layer_state.as_ref().unwrap();
166            (
167                past_state.prev_key.view([
168                    batch_size * self.num_attention_heads,
169                    -1,
170                    self.head_dim,
171                ]),
172                past_state.prev_value.view([
173                    batch_size * self.num_attention_heads,
174                    -1,
175                    self.head_dim,
176                ]),
177            )
178        };
179
180        if is_cross_attention {
181            if layer_state.is_some() {
182                layer_state.as_mut().unwrap().prev_key =
183                    key_states.view([batch_size, self.num_attention_heads, -1, self.head_dim]);
184                layer_state.as_mut().unwrap().prev_value =
185                    value_states.view([batch_size, self.num_attention_heads, -1, self.head_dim]);
186            } else {
187                layer_state = Some(LayerState {
188                    prev_key: key_states.view([
189                        batch_size,
190                        self.num_attention_heads,
191                        -1,
192                        self.head_dim,
193                    ]),
194                    prev_value: value_states.view([
195                        batch_size,
196                        self.num_attention_heads,
197                        -1,
198                        self.head_dim,
199                    ]),
200                })
201            }
202        };
203
204        let key_sequence_key = key_states.size()[1];
205        let mut attention_weights = query_states.bmm(&key_states.transpose(1, 2));
206
207        if let Some(attention_mask) = attention_mask {
208            attention_weights = attention_weights + attention_mask;
209        };
210
211        let attention_weights_reshaped = attention_weights.view([
212            batch_size,
213            self.num_attention_heads,
214            sequence_length,
215            key_sequence_key,
216        ]);
217
218        let attention_probs = attention_weights_reshaped
219            .view([
220                batch_size * self.num_attention_heads,
221                sequence_length,
222                key_sequence_key,
223            ])
224            .softmax(-1, attention_weights_reshaped.kind())
225            .apply_t(&self.attention_dropout, train);
226
227        let attention_output = attention_probs
228            .bmm(&value_states)
229            .transpose(0, 1)
230            .contiguous()
231            .view([sequence_length, batch_size, hidden_size])
232            .apply(&self.out_proj)
233            .apply_t(&self.dropout, train);
234
235        let attention_weights = if self.output_attentions {
236            Some(attention_weights_reshaped)
237        } else {
238            None
239        };
240        (attention_output, attention_weights, layer_state)
241    }
242}
243
244#[derive(Debug)]
245pub struct ProphetNetFeedForward {
246    activation_function: TensorFunction,
247    intermediate: nn::Linear,
248    output: nn::Linear,
249    activation_dropout: Dropout,
250    dropout: Dropout,
251}
252
253impl ProphetNetFeedForward {
254    pub fn new<'p, P>(p: P, config: &ProphetNetConfig, ffn_dim: i64) -> ProphetNetFeedForward
255    where
256        P: Borrow<nn::Path<'p>>,
257    {
258        let p = p.borrow();
259
260        let activation_function = config.activation_function.get_function();
261        let intermediate = nn::linear(
262            p / "intermediate",
263            config.hidden_size,
264            ffn_dim,
265            Default::default(),
266        );
267        let output = nn::linear(
268            p / "output",
269            ffn_dim,
270            config.hidden_size,
271            Default::default(),
272        );
273        let activation_dropout = Dropout::new(config.activation_dropout);
274        let dropout = Dropout::new(config.dropout);
275        ProphetNetFeedForward {
276            activation_function,
277            intermediate,
278            output,
279            activation_dropout,
280            dropout,
281        }
282    }
283}
284
285impl ModuleT for ProphetNetFeedForward {
286    fn forward_t(&self, xs: &Tensor, train: bool) -> Tensor {
287        let hidden_states = (self.activation_function.get_fn())(&xs.apply(&self.intermediate));
288        hidden_states
289            .apply_t(&self.activation_dropout, train)
290            .apply(&self.output)
291            .apply_t(&self.dropout, train)
292    }
293}
294
295pub struct ProphetNetNgramAttention {
296    num_buckets: i64,
297    ngram: i64,
298    relative_max_distance: i64,
299    num_attention_heads: i64,
300    dropout: Dropout,
301    attention_dropout: Dropout,
302    head_dim: i64,
303    key_proj: nn::Linear,
304    value_proj: nn::Linear,
305    query_proj: nn::Linear,
306    out_proj: nn::Linear,
307    relative_pos_embeddings: nn::Linear,
308    output_attentions: bool,
309}
310
311impl ProphetNetNgramAttention {
312    pub fn new<'p, P>(p: P, config: &ProphetNetConfig) -> ProphetNetNgramAttention
313    where
314        P: Borrow<nn::Path<'p>>,
315    {
316        let p = p.borrow();
317
318        let num_buckets = config.num_buckets;
319        let ngram = config.ngram;
320        let relative_max_distance = config.relative_max_distance;
321        let num_attention_heads = config.num_decoder_attention_heads;
322        let dropout = Dropout::new(config.dropout);
323        let attention_dropout = Dropout::new(config.attention_dropout);
324        let head_dim = config.hidden_size / num_attention_heads;
325
326        let key_proj = nn::linear(
327            p / "key_proj",
328            config.hidden_size,
329            config.hidden_size,
330            Default::default(),
331        );
332
333        let value_proj = nn::linear(
334            p / "value_proj",
335            config.hidden_size,
336            config.hidden_size,
337            Default::default(),
338        );
339
340        let query_proj = nn::linear(
341            p / "query_proj",
342            config.hidden_size,
343            config.hidden_size,
344            Default::default(),
345        );
346
347        let out_proj = nn::linear(
348            p / "out_proj",
349            config.hidden_size,
350            config.hidden_size,
351            Default::default(),
352        );
353
354        let relative_pos_embeddings = nn::linear(
355            p / "relative_pos_embeddings",
356            config.hidden_size,
357            num_buckets * num_attention_heads,
358            Default::default(),
359        );
360
361        let output_attentions = config.output_attentions.unwrap_or(false);
362
363        ProphetNetNgramAttention {
364            num_buckets,
365            ngram,
366            relative_max_distance,
367            num_attention_heads,
368            dropout,
369            attention_dropout,
370            head_dim,
371            key_proj,
372            value_proj,
373            query_proj,
374            out_proj,
375            relative_pos_embeddings,
376            output_attentions,
377        }
378    }
379
380    fn flatten<T>(&self, x: T, dim_0: i64, bs: i64) -> Tensor
381    where
382        T: Borrow<Tensor>,
383    {
384        x.borrow()
385            .contiguous()
386            .view((dim_0, bs * self.num_attention_heads, self.head_dim))
387            .transpose(0, 1)
388    }
389
390    pub fn forward_t(
391        &self,
392        hidden_states: &Tensor,
393        mut layer_state: Option<LayerState>,
394        attention_mask: Option<&Tensor>,
395        extended_predict_attention_mask: Option<&Tensor>,
396        main_relative_position_buckets: Option<&Tensor>,
397        predict_relative_position_buckets: Option<&Tensor>,
398        position_ids: &Tensor,
399        train: bool,
400    ) -> (Tensor, Option<Tensor>, Option<Tensor>, Option<LayerState>) {
401        let hidden_states_size = hidden_states.size();
402        let (sequence_length, batch_size, hidden_size) = (
403            hidden_states_size[0],
404            hidden_states_size[1],
405            hidden_states_size[2],
406        );
407
408        let query_states = hidden_states.apply(&self.query_proj) / (self.head_dim as f64).sqrt();
409        let key_states = hidden_states.apply(&self.key_proj);
410        let value_states = hidden_states.apply(&self.value_proj);
411
412        let mut main_hidden_states = hidden_states.chunk(1 + self.ngram, 0);
413        let mut main_query_states = self
414            .flatten(query_states, sequence_length, batch_size)
415            .chunk(1 + self.ngram, 1);
416        let mut main_key_states = self
417            .flatten(key_states, -1, batch_size)
418            .chunk(1 + self.ngram, 1);
419        let mut main_value_states = self
420            .flatten(value_states, -1, batch_size)
421            .chunk(1 + self.ngram, 1);
422
423        let hidden_states_predict_list = main_hidden_states.split_off(1);
424        let predict_query_states_list = main_query_states.split_off(1);
425        let predict_key_states_list = main_key_states.split_off(1);
426        let predict_value_states_list = main_value_states.split_off(1);
427
428        let main_hidden_states = main_hidden_states.pop().unwrap();
429        let main_query_states = main_query_states.pop().unwrap();
430        let mut main_key_states = main_key_states.pop().unwrap();
431        let mut main_value_states = main_value_states.pop().unwrap();
432
433        if let Some(layer_state_value) = &layer_state {
434            let prev_main_key_states = layer_state_value.prev_key.view([
435                batch_size * self.num_attention_heads,
436                -1,
437                self.head_dim,
438            ]);
439            let prev_main_value_states = layer_state_value.prev_value.view([
440                batch_size * self.num_attention_heads,
441                -1,
442                self.head_dim,
443            ]);
444            main_key_states = Tensor::cat(&[prev_main_key_states, main_key_states], 1);
445            main_value_states = Tensor::cat(&[prev_main_value_states, main_value_states], 1);
446        };
447
448        if layer_state.is_some() {
449            layer_state.as_mut().unwrap().prev_key =
450                main_key_states.view([batch_size, self.num_attention_heads, -1, self.head_dim]);
451            layer_state.as_mut().unwrap().prev_value =
452                main_value_states.view([batch_size, self.num_attention_heads, -1, self.head_dim]);
453        } else {
454            layer_state = Some(LayerState {
455                prev_key: main_key_states.view([
456                    batch_size,
457                    self.num_attention_heads,
458                    -1,
459                    self.head_dim,
460                ]),
461                prev_value: main_value_states.view([
462                    batch_size,
463                    self.num_attention_heads,
464                    -1,
465                    self.head_dim,
466                ]),
467            })
468        };
469        let main_sequence_length = sequence_length / (1 + self.ngram);
470
471        let main_attention_weights = main_query_states.bmm(&main_key_states.transpose(1, 2));
472
473        let main_relative_pos_embeddings = self.get_main_relative_position_embeddings(
474            &main_hidden_states,
475            &main_attention_weights,
476            position_ids,
477            main_relative_position_buckets,
478        );
479        let mut main_attention_weights = main_attention_weights + main_relative_pos_embeddings;
480        if let Some(attention_mask_value) = attention_mask {
481            main_attention_weights = main_attention_weights + attention_mask_value;
482        };
483
484        let main_attention_probas = main_attention_weights
485            .softmax(-1, main_attention_weights.kind())
486            .apply_t(&self.attention_dropout, train);
487
488        let main_attention_output = main_attention_probas
489            .bmm(&main_value_states)
490            .transpose(0, 1)
491            .contiguous()
492            .view([-1, main_sequence_length, batch_size, hidden_size])
493            .apply(&self.out_proj);
494
495        let predict_hidden_states = Tensor::cat(hidden_states_predict_list.as_slice(), 0).view([
496            self.ngram,
497            main_sequence_length,
498            batch_size,
499            hidden_size,
500        ]);
501
502        let predict_query_states = Tensor::cat(predict_query_states_list.as_slice(), 0).view([
503            self.ngram,
504            -1,
505            main_sequence_length,
506            self.head_dim,
507        ]);
508
509        let predict_key_states = Tensor::cat(
510            predict_key_states_list
511                .iter()
512                .map(|predict_key_state| {
513                    Tensor::cat(&[&main_key_states, predict_key_state], 1).unsqueeze(0)
514                })
515                .collect::<Vec<Tensor>>()
516                .as_slice(),
517            0,
518        );
519
520        let predict_value_states = Tensor::cat(
521            predict_value_states_list
522                .iter()
523                .map(|predict_value_state| {
524                    Tensor::cat(&[&main_value_states, predict_value_state], 1).unsqueeze(0)
525                })
526                .collect::<Vec<Tensor>>()
527                .as_slice(),
528            0,
529        );
530
531        let predict_attention_weights = Tensor::einsum(
532            "nbtc,nbsc->nbts",
533            &[predict_query_states, predict_key_states],
534            None::<i64>,
535        );
536
537        let predict_relative_pos_embeddings = self.get_predict_relative_pos_embeddings(
538            &predict_hidden_states,
539            &predict_attention_weights,
540            position_ids,
541            predict_relative_position_buckets,
542        );
543
544        let mut predict_attention_weights =
545            predict_attention_weights + predict_relative_pos_embeddings;
546        if let Some(extended_predict_attention_mask_value) = extended_predict_attention_mask {
547            predict_attention_weights =
548                predict_attention_weights + extended_predict_attention_mask_value;
549        };
550
551        let predict_attention_probas = predict_attention_weights
552            .softmax(-1, predict_attention_weights.kind())
553            .apply_t(&self.attention_dropout, train);
554
555        let predict_attention_output = Tensor::einsum(
556            "nbts,nbsc->nbtc",
557            &[&predict_attention_probas, &predict_value_states],
558            None::<i64>,
559        )
560        .transpose(1, 2)
561        .contiguous()
562        .view([self.ngram, main_sequence_length, batch_size, hidden_size])
563        .apply(&self.out_proj);
564
565        let attention_output = Tensor::cat(&[main_attention_output, predict_attention_output], 0)
566            .view([-1, batch_size, hidden_size])
567            .apply_t(&self.dropout, train);
568
569        let (main_attention_probas, predict_attention_probas) = if self.output_attentions {
570            let main_attention_probas = main_attention_probas.view([
571                batch_size,
572                self.num_attention_heads,
573                main_sequence_length,
574                -1,
575            ]);
576            let predict_attention_probas = predict_attention_probas
577                .view([
578                    self.ngram,
579                    batch_size,
580                    self.num_attention_heads,
581                    main_sequence_length,
582                    -1,
583                ])
584                .transpose(0, 1);
585            (Some(main_attention_probas), Some(predict_attention_probas))
586        } else {
587            (None, None)
588        };
589        (
590            attention_output,
591            main_attention_probas,
592            predict_attention_probas,
593            layer_state,
594        )
595    }
596
597    fn get_main_relative_position_embeddings(
598        &self,
599        hidden_states: &Tensor,
600        attention_weights: &Tensor,
601        position_ids: &Tensor,
602        main_relative_position_buckets: Option<&Tensor>,
603    ) -> Tensor {
604        let hidden_states_size = hidden_states.size();
605        let (sequence_length, batch_size) = (hidden_states_size[0], hidden_states_size[1]);
606        let calc_main_relative_position_buckets = if main_relative_position_buckets.is_none() {
607            let relative_positions = Tensor::arange_start(
608                1,
609                attention_weights.size().last().unwrap() + 1,
610                (Kind::Int64, hidden_states.device()),
611            )
612            .unsqueeze(0)
613            .unsqueeze(0)
614            .repeat([batch_size, sequence_length, 1]);
615            let relative_positions = relative_positions
616                - position_ids
617                    .unsqueeze(0)
618                    .repeat([batch_size, sequence_length, 1]);
619            Some(compute_relative_buckets(
620                self.num_buckets,
621                self.relative_max_distance,
622                &relative_positions,
623                false,
624            ))
625        } else {
626            None
627        };
628        let main_relative_position_buckets = main_relative_position_buckets
629            .unwrap_or_else(|| calc_main_relative_position_buckets.as_ref().unwrap());
630
631        let rel_pos_embeddings = hidden_states
632            .transpose(0, 1)
633            .apply(&self.relative_pos_embeddings)
634            .view([
635                batch_size,
636                sequence_length,
637                self.num_buckets,
638                self.num_attention_heads,
639            ])
640            .permute([0, 3, 1, 2])
641            .reshape([-1, self.num_buckets]);
642
643        let main_relative_position_buckets = main_relative_position_buckets
644            .repeat([1, self.num_attention_heads, 1])
645            .view([-1, *main_relative_position_buckets.size().last().unwrap()]);
646
647        let mut new_shape = attention_weights
648            .size()
649            .into_iter()
650            .take(2)
651            .collect::<Vec<i64>>();
652        new_shape.push(-1);
653        rel_pos_embeddings
654            .gather(1, &main_relative_position_buckets, false)
655            .view(new_shape.as_slice())
656    }
657
658    fn get_predict_relative_pos_embeddings(
659        &self,
660        hidden_states: &Tensor,
661        attention_weights: &Tensor,
662        position_ids: &Tensor,
663        predict_relative_position_buckets: Option<&Tensor>,
664    ) -> Tensor {
665        let hidden_states_size = hidden_states.size();
666        let (sequence_length, batch_size) = (hidden_states_size[1], hidden_states_size[2]);
667
668        let calc_predict_relative_position_buckets = if predict_relative_position_buckets.is_none()
669        {
670            let key_sequence_length = *attention_weights.size().last().unwrap();
671            let relative_positions =
672                Tensor::arange(key_sequence_length, (Kind::Int64, hidden_states.device()))
673                    .unsqueeze(0)
674                    .unsqueeze(0)
675                    .repeat([batch_size, sequence_length, 1]);
676            let relative_positions = relative_positions
677                - position_ids
678                    .unsqueeze(0)
679                    .repeat([batch_size, sequence_length, 1]);
680            Some(compute_relative_buckets(
681                self.num_buckets,
682                self.relative_max_distance,
683                &relative_positions,
684                false,
685            ))
686        } else {
687            None
688        };
689
690        let predict_relative_position_buckets = predict_relative_position_buckets
691            .unwrap_or_else(|| calc_predict_relative_position_buckets.as_ref().unwrap());
692
693        let rel_pos_embeddings = hidden_states
694            .transpose(0, 1)
695            .apply(&self.relative_pos_embeddings)
696            .view([
697                self.ngram,
698                batch_size,
699                sequence_length,
700                self.num_buckets,
701                self.num_attention_heads,
702            ])
703            .permute([0, 1, 4, 2, 3])
704            .reshape([-1, self.num_buckets]);
705
706        let predict_relative_position_buckets = predict_relative_position_buckets
707            .unsqueeze(0)
708            .repeat([self.ngram, 1, self.num_attention_heads, 1])
709            .view([
710                -1,
711                *predict_relative_position_buckets.size().last().unwrap(),
712            ]);
713
714        rel_pos_embeddings
715            .gather(1, &predict_relative_position_buckets, false)
716            .view([
717                self.ngram,
718                batch_size * self.num_attention_heads,
719                sequence_length,
720                -1,
721            ])
722    }
723}
724
725pub(crate) fn compute_relative_buckets(
726    num_buckets: i64,
727    max_distance: i64,
728    relative_positions: &Tensor,
729    bidirectional: bool,
730) -> Tensor {
731    let inverse_relative_positions = -relative_positions;
732
733    let (num_buckets, relative_positions_bucket, inverse_relative_positions) = if bidirectional {
734        let num_buckets = num_buckets / 2;
735        let relative_position_bucket =
736            inverse_relative_positions.lt(0).totype(Kind::Int) * num_buckets;
737        let inverse_relative_position = inverse_relative_positions.abs();
738        (
739            num_buckets,
740            relative_position_bucket,
741            inverse_relative_position,
742        )
743    } else {
744        (
745            num_buckets,
746            relative_positions.zeros_like(),
747            inverse_relative_positions.max_other(&inverse_relative_positions.zeros_like()),
748        )
749    };
750    let max_exact = num_buckets / 2;
751    let is_small = inverse_relative_positions.lt(max_exact);
752    let max_exact_f64 = max_exact as f64;
753    let val_if_large = (inverse_relative_positions.totype(Kind::Float) / max_exact_f64).log2()
754        / (max_distance as f64 / max_exact_f64).log2()
755        * (num_buckets as f64 - max_exact_f64)
756        + max_exact_f64;
757
758    let val_if_large = val_if_large
759        .min_other(&(val_if_large.ones_like() * (num_buckets as f64 - 1.0)))
760        .totype(Kind::Int64);
761
762    relative_positions_bucket + inverse_relative_positions.where_self(&is_small, &val_if_large)
763}
764
765pub(crate) fn compute_all_stream_relative_buckets(
766    num_buckets: i64,
767    max_distance: i64,
768    position_ids: &Tensor,
769) -> (Tensor, Tensor) {
770    let main_stream_relative_positions =
771        position_ids
772            .unsqueeze(1)
773            .repeat([1, *position_ids.size().last().unwrap(), 1])
774            - position_ids.unsqueeze(-1);
775
776    let predicting_stream_relative_positions = Tensor::cat(&[&(position_ids - 1), position_ids], 1)
777        .unsqueeze(1)
778        .repeat([1, *position_ids.size().last().unwrap(), 1])
779        - position_ids.unsqueeze(-1);
780
781    let main_relative_position_buckets = compute_relative_buckets(
782        num_buckets,
783        max_distance,
784        &main_stream_relative_positions,
785        false,
786    );
787
788    let predict_relative_position_buckets = compute_relative_buckets(
789        num_buckets,
790        max_distance,
791        &predicting_stream_relative_positions,
792        false,
793    );
794
795    (
796        main_relative_position_buckets,
797        predict_relative_position_buckets,
798    )
799}