Compiling rope while preserving true on policy (#12161)
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user