(
self,
hidden_states: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.LongTensor] = None,
past_key_value: Optional[Cache] = None,
output_attentions: Optional[bool] = False,
output_docs_score: Optional[bool] = False,
use_cache: Optional[bool] = False,
cache_position: Optional[torch.LongTensor] = None,
position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, # necessary, but kept here for BC
doc_ids: Optional[torch.Tensor] = None,
input_ids: Optional[torch.LongTensor] = None,
**kwargs: Unpack[FlashAttentionKwargs],
)
| 65 | config.sliding_window = False |
| 66 | |
| 67 | def forward( |
| 68 | self, |
| 69 | hidden_states: torch.Tensor, |
| 70 | attention_mask: Optional[torch.Tensor] = None, |
| 71 | position_ids: Optional[torch.LongTensor] = None, |
| 72 | past_key_value: Optional[Cache] = None, |
| 73 | output_attentions: Optional[bool] = False, |
| 74 | output_docs_score: Optional[bool] = False, |
| 75 | use_cache: Optional[bool] = False, |
| 76 | cache_position: Optional[torch.LongTensor] = None, |
| 77 | position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, # necessary, but kept here for BC |
| 78 | doc_ids: Optional[torch.Tensor] = None, |
| 79 | input_ids: Optional[torch.LongTensor] = None, |
| 80 | **kwargs: Unpack[FlashAttentionKwargs], |
| 81 | ) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]]: |
| 82 | residual = hidden_states |
| 83 | |
| 84 | hidden_states = self.input_layernorm(hidden_states) |
| 85 | |
| 86 | # Self Attention |
| 87 | if self.attn_type == "full_attention": |
| 88 | hidden_states, self_attn_weights = self.self_attn( |
| 89 | hidden_states=hidden_states, |
| 90 | attention_mask=attention_mask, |
| 91 | position_ids=position_ids, |
| 92 | past_key_value=past_key_value, |
| 93 | output_attentions=output_attentions, |
| 94 | use_cache=use_cache, |
| 95 | cache_position=cache_position, |
| 96 | position_embeddings=position_embeddings, |
| 97 | doc_ids=doc_ids, |
| 98 | input_ids=input_ids, |
| 99 | **kwargs, |
| 100 | ) |
| 101 | else: |
| 102 | hidden_states, self_attn_weights = self.self_attn( |
| 103 | hidden_states=hidden_states, |
| 104 | attention_mask=attention_mask, |
| 105 | position_ids=position_ids, |
| 106 | past_key_value=past_key_value, |
| 107 | output_attentions=output_attentions, |
| 108 | use_cache=use_cache, |
| 109 | cache_position=cache_position, |
| 110 | position_embeddings=position_embeddings, |
| 111 | doc_ids=doc_ids, |
| 112 | input_ids=input_ids, |
| 113 | **kwargs, |
| 114 | ) |
| 115 | |
| 116 | if isinstance(hidden_states, tuple): |
| 117 | hidden_states, docs_score = hidden_states |
| 118 | else: |
| 119 | docs_score = None |
| 120 | hidden_states = residual + hidden_states |
| 121 | |
| 122 | # Fully Connected |
| 123 | residual = hidden_states |
| 124 | hidden_states = self.post_attention_layernorm(hidden_states) |
nothing calls this directly
no outgoing calls
no test coverage detected