chore: update torch v2.5.1 (#1849)
This commit is contained in:
@@ -38,6 +38,7 @@ from sglang.srt.utils import set_weight_attrs
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@CustomOp.register("silu_and_mul")
|
||||
class SiluAndMul(CustomOp):
|
||||
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
|
||||
d = x.shape[-1] // 2
|
||||
@@ -51,6 +52,7 @@ class SiluAndMul(CustomOp):
|
||||
return out
|
||||
|
||||
|
||||
@CustomOp.register("gelu_and_mul")
|
||||
class GeluAndMul(CustomOp):
|
||||
def __init__(self, approximate="tanh"):
|
||||
super().__init__()
|
||||
|
||||
@@ -36,6 +36,7 @@ from vllm.model_executor.custom_op import CustomOp
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@CustomOp.register("rmsnorm")
|
||||
class RMSNorm(CustomOp):
|
||||
def __init__(
|
||||
self,
|
||||
@@ -78,6 +79,7 @@ class RMSNorm(CustomOp):
|
||||
return x, residual
|
||||
|
||||
|
||||
@CustomOp.register("gemma_rmsnorm")
|
||||
class GemmaRMSNorm(CustomOp):
|
||||
def __init__(
|
||||
self,
|
||||
|
||||
Reference in New Issue
Block a user