def apply_rotary_pos_emb(x, cos, sin): # x shape: [batch, heads, seq_len, head_dim] x1 = x[..., :x.shape[-1] // 2] x2 = x[..., x.shape[-1] // 2:] rotated_x = torch.cat((-x2, x1), dim=-1) return (x * cos) + (rotated_x * sin)