(x, interleaved=False)
| 11 | |
| 12 | |
| 13 | def rotate_half(x, interleaved=False): |
| 14 | if not interleaved: |
| 15 | x1, x2 = x.chunk(2, dim=-1) |
| 16 | return torch.cat((-x2, x1), dim=-1) |
| 17 | else: |
| 18 | x1, x2 = x[..., ::2], x[..., 1::2] |
| 19 | return rearrange(torch.stack((-x2, x1), dim=-1), "... d two -> ... (d two)", two=2) |
| 20 | |
| 21 | |
| 22 | def apply_rotary_emb_torch(x, cos, sin, interleaved=False): |
no outgoing calls
no test coverage detected