1use 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)]
22pub struct LayerState {
25 pub prev_key: Tensor,
27 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}