From f389f01714c9ebe84595565b5252ff97866b9617 Mon Sep 17 00:00:00 2001 From: Yuan Luo Date: Mon, 27 Oct 2025 23:49:22 +0800 Subject: [PATCH] Optimize triton_mrope with torch compile (#12112) Co-authored-by: luoyuan.luo --- python/sglang/srt/layers/rotary_embedding.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/python/sglang/srt/layers/rotary_embedding.py b/python/sglang/srt/layers/rotary_embedding.py index 4b3856fd1..48f564f83 100644 --- a/python/sglang/srt/layers/rotary_embedding.py +++ b/python/sglang/srt/layers/rotary_embedding.py @@ -1424,6 +1424,7 @@ class MRotaryEmbedding(RotaryEmbedding): else: return self._forward_native(positions, query, key) + @torch.compile(dynamic=True, backend=get_compiler_backend()) def _forward_triton( self, positions: torch.Tensor, @@ -1442,6 +1443,7 @@ class MRotaryEmbedding(RotaryEmbedding): if positions.ndim == 2: assert self.mrope_section + torch._dynamo.graph_break() q, k = triton_mrope( query, key, @@ -1453,6 +1455,7 @@ class MRotaryEmbedding(RotaryEmbedding): self.mrope_interleaved, self.is_neox_style, ) + torch._dynamo.graph_break() return q.reshape(query_shape), k.reshape(key_shape)