Refactor CustomOp to avoid confusing bugs (#5382)
This commit is contained in:
@@ -1,6 +1,3 @@
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from sglang.srt.utils import is_cuda, is_hip
|
||||
@@ -14,6 +11,26 @@ class CustomOp(nn.Module):
|
||||
super().__init__()
|
||||
self._forward_method = self.dispatch_forward()
|
||||
|
||||
def enter_torch_compile(self, num_tokens: int):
|
||||
# NOTE: Temporarily workaround MoE
|
||||
if "FusedMoE" in self.__class__.__name__:
|
||||
if num_tokens == 1:
|
||||
from sglang.srt.layers.moe.fused_moe_native import (
|
||||
fused_moe_forward_native,
|
||||
)
|
||||
|
||||
# The performance of torch.compile on this layer is not always good when bs > 1,
|
||||
# so we decide to only use torch.compile when bs =1
|
||||
self._forward_method = fused_moe_forward_native
|
||||
else:
|
||||
self._forward_method = self.forward_native
|
||||
self.is_torch_compile = True
|
||||
|
||||
def leave_torch_compile(self):
|
||||
self._forward_method = self.forward_cuda
|
||||
self.is_torch_compile = False
|
||||
|
||||
# Please do not override this method, because `self._forward_method` can change when in torch compile mode
|
||||
def forward(self, *args, **kwargs):
|
||||
return self._forward_method(*args, **kwargs)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user