From 326c84c4935caa7ffc8968366c606cb3e31cdd96 Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Tue, 28 Oct 2025 08:02:17 +0800 Subject: [PATCH] Compiling rope while preserving true on policy (#12161) --- python/sglang/srt/layers/rotary_embedding.py | 11 +++++++++-- 1 file changed, 9 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/layers/rotary_embedding.py b/python/sglang/srt/layers/rotary_embedding.py index d48dc1b88..64ac00ff1 100644 --- a/python/sglang/srt/layers/rotary_embedding.py +++ b/python/sglang/srt/layers/rotary_embedding.py @@ -125,8 +125,13 @@ class RotaryEmbedding(CustomOp): self.cos_sin_cache: torch.Tensor self.register_buffer("cos_sin_cache", cache, persistent=False) + self._apply_rotary_emb_wrapped = _apply_rotary_emb + if get_global_server_args().rl_on_policy_target == "fsdp": self._forward_method = self.forward_native + self._apply_rotary_emb_wrapped = torch.compile(dynamic=True)( + self._apply_rotary_emb_wrapped + ) def _compute_inv_freq(self, base: Union[int, float]) -> torch.Tensor: """Compute the inverse frequency.""" @@ -185,14 +190,16 @@ class RotaryEmbedding(CustomOp): query = query.view(num_tokens, -1, self.head_size) query_rot = query[..., : self.rotary_dim] query_pass = query[..., self.rotary_dim :] - query_rot = _apply_rotary_emb(query_rot, cos, sin, self.is_neox_style) + query_rot = self._apply_rotary_emb_wrapped( + query_rot, cos, sin, self.is_neox_style + ) query = torch.cat((query_rot, query_pass), dim=-1).reshape(query_shape) key_shape = key.shape key = key.view(num_tokens, -1, self.head_size) key_rot = key[..., : self.rotary_dim] key_pass = key[..., self.rotary_dim :] - key_rot = _apply_rotary_emb(key_rot, cos, sin, self.is_neox_style) + key_rot = self._apply_rotary_emb_wrapped(key_rot, cos, sin, self.is_neox_style) key = torch.cat((key_rot, key_pass), dim=-1).reshape(key_shape) return query, key