(
self,
hidden_states: Optional[torch.FloatTensor],
layer_past: Optional[Tuple[torch.Tensor]] = None,
attention_mask: Optional[torch.FloatTensor] = None,
position_ids: Optional[torch.LongTensor] = None,
head_mask: Optional[torch.FloatTensor] = None,
use_cache: Optional[bool] = False,
output_attentions: Optional[bool] = False,
)
| 149 | return attn_output, attn_weights |
| 150 | |
| 151 | def forward( |
| 152 | self, |
| 153 | hidden_states: Optional[torch.FloatTensor], |
| 154 | layer_past: Optional[Tuple[torch.Tensor]] = None, |
| 155 | attention_mask: Optional[torch.FloatTensor] = None, |
| 156 | position_ids: Optional[torch.LongTensor] = None, |
| 157 | head_mask: Optional[torch.FloatTensor] = None, |
| 158 | use_cache: Optional[bool] = False, |
| 159 | output_attentions: Optional[bool] = False, |
| 160 | ) -> Union[ |
| 161 | Tuple[torch.Tensor, Tuple[torch.Tensor]], |
| 162 | Optional[Tuple[torch.Tensor, Tuple[torch.Tensor], Tuple[torch.Tensor, ...]]], |
| 163 | ]: |
| 164 | qkv = self.qkv_proj(hidden_states) |
| 165 | # TODO(enijkamp): factor out number of logical TPU-v4 cores or make forward pass agnostic |
| 166 | mp_num = 4 |
| 167 | qkv_split = qkv.reshape(qkv.shape[:-1] + (mp_num, -1)) |
| 168 | |
| 169 | local_dim = self.head_dim * self.num_attention_heads // mp_num |
| 170 | query, value, key = torch.split(qkv_split, local_dim, dim=-1) |
| 171 | query = self._split_heads(query, self.num_attention_heads, self.head_dim, mp_num=mp_num) |
| 172 | key = self._split_heads(key, self.num_attention_heads, self.head_dim, mp_num=mp_num) |
| 173 | |
| 174 | value = self._split_heads(value, self.num_attention_heads, self.head_dim, mp_num=mp_num) |
| 175 | value = value.permute(0, 2, 1, 3) |
| 176 | |
| 177 | embed_positions = self.embed_positions |
| 178 | if embed_positions.device != position_ids.device: |
| 179 | embed_positions = embed_positions.to(position_ids.device) |
| 180 | self.embed_positions = embed_positions |
| 181 | |
| 182 | sincos = embed_positions[position_ids] |
| 183 | sin, cos = torch.split(sincos, sincos.shape[-1] // 2, dim=-1) |
| 184 | |
| 185 | if self.rotary_dim is not None: |
| 186 | k_rot = key[:, :, :, : self.rotary_dim] |
| 187 | k_pass = key[:, :, :, self.rotary_dim :] |
| 188 | |
| 189 | q_rot = query[:, :, :, : self.rotary_dim] |
| 190 | q_pass = query[:, :, :, self.rotary_dim :] |
| 191 | |
| 192 | k_rot = apply_rotary_pos_emb(k_rot, sin, cos) |
| 193 | q_rot = apply_rotary_pos_emb(q_rot, sin, cos) |
| 194 | |
| 195 | key = torch.cat([k_rot, k_pass], dim=-1) |
| 196 | query = torch.cat([q_rot, q_pass], dim=-1) |
| 197 | else: |
| 198 | key = apply_rotary_pos_emb(key, sin, cos) |
| 199 | query = apply_rotary_pos_emb(query, sin, cos) |
| 200 | |
| 201 | key = key.permute(0, 2, 1, 3) |
| 202 | query = query.permute(0, 2, 1, 3) |
| 203 | |
| 204 | if layer_past is not None: |
| 205 | past_key = layer_past[0] |
| 206 | past_value = layer_past[1] |
| 207 | key = torch.cat((past_key, key), dim=-2) |
| 208 | value = torch.cat((past_value, value), dim=-2) |
nothing calls this directly
no test coverage detected