(freqs_cis: torch.Tensor, x: torch.Tensor)
| 44 | |
| 45 | |
| 46 | def reshape_for_broadcast(freqs_cis: torch.Tensor, x: torch.Tensor): |
| 47 | ndim = x.ndim |
| 48 | assert 0 <= 1 < ndim |
| 49 | assert freqs_cis.shape == (x.shape[1], x.shape[-1]) |
| 50 | shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)] |
| 51 | return freqs_cis.view(*shape) |
| 52 | |
| 53 | |
| 54 | def apply_rotary_emb( |