Add Wall-X serving and Turtle2 TCP WebSocket bridge
Pre-commit / pre-commit (push) Canceled after 0s
Pre-commit / pre-commit (push) Canceled after 0s
This commit is contained in:
@@ -5,6 +5,7 @@ Mixture-of-Experts processing.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
|
||||
import torch
|
||||
|
||||
@@ -13,6 +14,14 @@ from wall_x.model.core.ops.base import OpsProxy
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _force_pytorch_backend() -> bool:
|
||||
return os.environ.get("WALL_X_FORCE_TORCH_MOE_OPS", "").lower() in {
|
||||
"1",
|
||||
"true",
|
||||
"yes",
|
||||
}
|
||||
|
||||
|
||||
class PermuteOp(OpsProxy):
|
||||
"""Reorder tokens by expert assignment for MoE processing.
|
||||
|
||||
@@ -21,9 +30,11 @@ class PermuteOp(OpsProxy):
|
||||
|
||||
@property
|
||||
def _external_accel_name(self):
|
||||
return "permute"
|
||||
return None if _force_pytorch_backend() else "permute"
|
||||
|
||||
def _get_cuda_kernel(self):
|
||||
if _force_pytorch_backend():
|
||||
return None
|
||||
try:
|
||||
from wall_x.model.core.ops._cuda_wrappers import permute_kernel
|
||||
|
||||
@@ -64,9 +75,11 @@ class UnpermuteOp(OpsProxy):
|
||||
|
||||
@property
|
||||
def _external_accel_name(self):
|
||||
return "unpermute"
|
||||
return None if _force_pytorch_backend() else "unpermute"
|
||||
|
||||
def _get_cuda_kernel(self):
|
||||
if _force_pytorch_backend():
|
||||
return None
|
||||
try:
|
||||
from wall_x.model.core.ops._cuda_wrappers import unpermute_kernel
|
||||
|
||||
|
||||
Reference in New Issue
Block a user