[fix] recover auto-dispatch for rmsnorm and rope (#6745)
This commit is contained in:
@@ -11,7 +11,20 @@ class CustomOp(nn.Module):
|
||||
super().__init__()
|
||||
self._forward_method = self.dispatch_forward()
|
||||
|
||||
# States for torch.compile
|
||||
self._original_forward_method = None
|
||||
self.is_torch_compile = False
|
||||
|
||||
def enter_torch_compile(self, num_tokens: int):
|
||||
# Skip if Op is already entered compile mode.
|
||||
# NOTE(alcanderian): Some Ops(for example RotaryEmbedding) will be reused
|
||||
# among layers and `enter_torch_compile` will be called many times.
|
||||
# We should prevent `self._original_forward_method` from being overridden when
|
||||
# it is not the first time `enter_torch_compile` called.
|
||||
if self.is_torch_compile:
|
||||
return
|
||||
|
||||
self._original_forward_method = self._forward_method
|
||||
# NOTE: Temporarily workaround MoE
|
||||
if "FusedMoE" in self.__class__.__name__:
|
||||
if num_tokens == 1:
|
||||
@@ -27,7 +40,12 @@ class CustomOp(nn.Module):
|
||||
self.is_torch_compile = True
|
||||
|
||||
def leave_torch_compile(self):
|
||||
self._forward_method = self.forward_cuda
|
||||
# Skip if Op is already exited compile mode.
|
||||
if not self.is_torch_compile:
|
||||
return
|
||||
|
||||
self._forward_method = self._original_forward_method
|
||||
self._original_forward_method = None
|
||||
self.is_torch_compile = False
|
||||
|
||||
# Please do not override this method, because `self._forward_method` can change when in torch compile mode
|
||||
|
||||
Reference in New Issue
Block a user