(self, x: Tensor, rope=None, drop_ratio: Optional[float] = None)
| 199 | return sin, cos |
| 200 | |
| 201 | def _forward(self, x: Tensor, rope=None, drop_ratio: Optional[float] = None) -> Tensor: |
| 202 | b, _, _ = x.shape |
| 203 | effective_drop_ratio = drop_ratio if drop_ratio is not None else self.sample_drop_ratio |
| 204 | if self.training and effective_drop_ratio > 0.0: |
| 205 | indices_1, residual_scale_factor = get_branges_scales(x, effective_drop_ratio) |
| 206 | |
| 207 | x_subset_1 = x[indices_1] |
| 208 | rope_subset = self._maybe_index_rope(rope, indices_1) |
| 209 | residual_1 = self.attn(self.norm1(x_subset_1), rope=rope_subset) |
| 210 | |
| 211 | x_attn = torch.index_add( |
| 212 | x, |
| 213 | dim=0, |
| 214 | source=self.ls1(residual_1), |
| 215 | index=indices_1, |
| 216 | alpha=residual_scale_factor, |
| 217 | ) |
| 218 | indices_2, residual_scale_factor = get_branges_scales(x_attn, effective_drop_ratio) |
| 219 | x_subset_2 = x_attn[indices_2] |
| 220 | residual_2 = self.mlp(self.norm2(x_subset_2)) |
| 221 | |
| 222 | x_ffn = torch.index_add( |
| 223 | x_attn, |
| 224 | dim=0, |
| 225 | source=self.ls2(residual_2), |
| 226 | index=indices_2, |
| 227 | alpha=residual_scale_factor, |
| 228 | ) |
| 229 | else: |
| 230 | x_attn = x + self.ls1(self.attn(self.norm1(x), rope=rope)) |
| 231 | x_ffn = x_attn + self.ls2(self.mlp(self.norm2(x_attn))) |
| 232 | |
| 233 | return x_ffn |
| 234 | |
| 235 | def _forward_list(self, x_list: List[Tensor], rope_list=None, drop_ratio: Optional[float] = None) -> List[Tensor]: |
| 236 | b_list = [x.shape[0] for x in x_list] |
nothing calls this directly
no test coverage detected