[Diffusion][CPU] Init CPU platform support for SGLang Diffusion (#20816)
This commit is contained in:
@@ -71,6 +71,26 @@ dependencies = [
|
|||||||
]
|
]
|
||||||
|
|
||||||
[project.optional-dependencies]
|
[project.optional-dependencies]
|
||||||
|
diffusion = [
|
||||||
|
"PyYAML==6.0.1",
|
||||||
|
"cloudpickle==3.1.2",
|
||||||
|
"diffusers==0.37.0",
|
||||||
|
"imageio==2.36.0",
|
||||||
|
"imageio-ffmpeg==0.5.1",
|
||||||
|
"moviepy>=2.0.0",
|
||||||
|
"opencv-python-headless==4.10.0.84",
|
||||||
|
"remote-pdb==2.1.0",
|
||||||
|
"st_attn==0.0.7 ; platform_machine != 'aarch64' and platform_machine != 'arm64'",
|
||||||
|
"vsa==0.0.4 ; platform_machine != 'aarch64' and platform_machine != 'arm64'",
|
||||||
|
"runai_model_streamer>=0.15.5",
|
||||||
|
"cache-dit==1.3.0",
|
||||||
|
"addict==2.4.0",
|
||||||
|
"av==16.1.0",
|
||||||
|
"scikit-image==0.25.2",
|
||||||
|
"trimesh>=4.0.0",
|
||||||
|
"xatlas",
|
||||||
|
]
|
||||||
|
|
||||||
tracing = [
|
tracing = [
|
||||||
"opentelemetry-sdk",
|
"opentelemetry-sdk",
|
||||||
"opentelemetry-api",
|
"opentelemetry-api",
|
||||||
|
|||||||
@@ -17,143 +17,25 @@ from torch import Tensor
|
|||||||
|
|
||||||
from sglang.srt.utils.tensor_bridge import mlx_to_torch, torch_to_mlx, use_mlx
|
from sglang.srt.utils.tensor_bridge import mlx_to_torch, torch_to_mlx, use_mlx
|
||||||
|
|
||||||
|
from .torch_fallback import (
|
||||||
|
apply_rotary_embedding_native,
|
||||||
|
fuse_scale_shift_kernel_native,
|
||||||
|
norm_infer_native,
|
||||||
|
rms_norm_fn_native,
|
||||||
|
triton_one_pass_rms_norm_native,
|
||||||
|
)
|
||||||
|
|
||||||
_use_mlx = use_mlx()
|
_use_mlx = use_mlx()
|
||||||
|
|
||||||
if _use_mlx:
|
if _use_mlx:
|
||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
|
|
||||||
|
# use the common torch native version form torch_fallback
|
||||||
def fuse_scale_shift_kernel_native(
|
fuse_scale_shift_kernel_native = fuse_scale_shift_kernel_native
|
||||||
x: torch.Tensor,
|
apply_rotary_embedding_native = apply_rotary_embedding_native
|
||||||
scale: torch.Tensor,
|
norm_infer_native = norm_infer_native
|
||||||
shift: torch.Tensor,
|
triton_one_pass_rms_norm_native = triton_one_pass_rms_norm_native
|
||||||
scale_constant: float = 1.0,
|
rms_norm_fn_native = rms_norm_fn_native
|
||||||
block_l: int = 128,
|
|
||||||
block_c: int = 128,
|
|
||||||
):
|
|
||||||
"""Native fallback for fuse_scale_shift_kernel with scale_constant support."""
|
|
||||||
B, L, C = x.shape
|
|
||||||
|
|
||||||
def _expand(t: torch.Tensor) -> torch.Tensor:
|
|
||||||
if t.dim() == 4:
|
|
||||||
# [B, F, 1, C] -> [B, L, C]
|
|
||||||
num_frames = t.shape[1]
|
|
||||||
frame_seqlen = L // num_frames
|
|
||||||
return (
|
|
||||||
t.squeeze(2)
|
|
||||||
.unsqueeze(2)
|
|
||||||
.expand(-1, -1, frame_seqlen, -1)
|
|
||||||
.reshape(B, L, C)
|
|
||||||
)
|
|
||||||
elif t.dim() == 2:
|
|
||||||
# [B, C] -> [B, 1, C]
|
|
||||||
return t.unsqueeze(1)
|
|
||||||
return t
|
|
||||||
|
|
||||||
scale = _expand(scale)
|
|
||||||
shift = _expand(shift)
|
|
||||||
|
|
||||||
return x * (scale_constant + scale) + shift
|
|
||||||
|
|
||||||
|
|
||||||
def apply_rotary_embedding_native(
|
|
||||||
x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, interleaved: bool = False
|
|
||||||
) -> torch.Tensor:
|
|
||||||
"""Native fallback for rotary embedding (shared with NPU implementation)."""
|
|
||||||
cos = cos.unsqueeze(-2).to(x.dtype)
|
|
||||||
sin = sin.unsqueeze(-2).to(x.dtype)
|
|
||||||
x1 = x[..., ::2]
|
|
||||||
x2 = x[..., 1::2]
|
|
||||||
o1 = x1 * cos - x2 * sin
|
|
||||||
o2 = x2 * cos + x1 * sin
|
|
||||||
return torch.stack((o1, o2), dim=-1).flatten(-2)
|
|
||||||
|
|
||||||
|
|
||||||
def norm_infer_native(
|
|
||||||
x: Tensor,
|
|
||||||
weight: Optional[Tensor],
|
|
||||||
bias: Optional[Tensor],
|
|
||||||
eps: float,
|
|
||||||
is_rms_norm: bool = False,
|
|
||||||
out: Optional[Tensor] = None,
|
|
||||||
) -> Tensor:
|
|
||||||
"""Native fallback for norm_infer (layer norm / rms norm inference)."""
|
|
||||||
orig_dtype = x.dtype
|
|
||||||
x = x.contiguous().float()
|
|
||||||
if is_rms_norm:
|
|
||||||
variance = x.pow(2).mean(dim=-1, keepdim=True)
|
|
||||||
x_hat = x * torch.rsqrt(variance + eps)
|
|
||||||
else:
|
|
||||||
mean = x.mean(dim=-1, keepdim=True)
|
|
||||||
variance = (x - mean).pow(2).mean(dim=-1, keepdim=True)
|
|
||||||
x_hat = (x - mean) * torch.rsqrt(variance + eps)
|
|
||||||
if weight is not None:
|
|
||||||
x_hat = x_hat * weight.float()
|
|
||||||
if bias is not None:
|
|
||||||
x_hat = x_hat + bias.float()
|
|
||||||
result = x_hat.to(orig_dtype)
|
|
||||||
if out is not None:
|
|
||||||
out.copy_(result)
|
|
||||||
return out
|
|
||||||
return result
|
|
||||||
|
|
||||||
|
|
||||||
def triton_one_pass_rms_norm_native(
|
|
||||||
x: torch.Tensor, w: torch.Tensor, eps: float = 1e-6
|
|
||||||
) -> torch.Tensor:
|
|
||||||
"""Native fallback for triton_one_pass_rms_norm."""
|
|
||||||
shape = x.shape
|
|
||||||
orig_dtype = x.dtype
|
|
||||||
x = x.contiguous().float()
|
|
||||||
variance = x.pow(2).mean(dim=-1, keepdim=True)
|
|
||||||
x_hat = x * torch.rsqrt(variance + eps)
|
|
||||||
return (x_hat * w.float()).to(orig_dtype).view(shape)
|
|
||||||
|
|
||||||
|
|
||||||
def rms_norm_fn_native(
|
|
||||||
x,
|
|
||||||
weight,
|
|
||||||
bias,
|
|
||||||
residual=None,
|
|
||||||
x1=None,
|
|
||||||
weight1=None,
|
|
||||||
bias1=None,
|
|
||||||
eps=1e-6,
|
|
||||||
dropout_p=0.0,
|
|
||||||
rowscale=None,
|
|
||||||
prenorm=False,
|
|
||||||
residual_in_fp32=False,
|
|
||||||
zero_centered_weight=False,
|
|
||||||
return_dropout_mask=False,
|
|
||||||
out_dtype=None,
|
|
||||||
out=None,
|
|
||||||
residual_out=None,
|
|
||||||
):
|
|
||||||
"""Native fallback for rms_norm_fn (inference only, no dropout/x1 support)."""
|
|
||||||
x_shape_og = x.shape
|
|
||||||
orig_dtype = x.dtype
|
|
||||||
x = x.reshape(-1, x.shape[-1]).float()
|
|
||||||
if residual is not None:
|
|
||||||
residual = residual.reshape(-1, residual.shape[-1]).float()
|
|
||||||
x = x + residual
|
|
||||||
residual_out_val = x.to(torch.float32 if residual_in_fp32 else orig_dtype)
|
|
||||||
else:
|
|
||||||
residual_out_val = None
|
|
||||||
variance = x.pow(2).mean(dim=-1, keepdim=True)
|
|
||||||
x_hat = x * torch.rsqrt(variance + eps)
|
|
||||||
if weight is not None:
|
|
||||||
w = weight.float()
|
|
||||||
if zero_centered_weight:
|
|
||||||
w = w + 1.0
|
|
||||||
x_hat = x_hat * w
|
|
||||||
if bias is not None:
|
|
||||||
x_hat = x_hat + bias.float()
|
|
||||||
final_dtype = out_dtype if out_dtype is not None else orig_dtype
|
|
||||||
y = x_hat.to(final_dtype).reshape(x_shape_og)
|
|
||||||
if residual is not None and residual_out_val is not None:
|
|
||||||
return y, residual_out_val.reshape(x_shape_og)
|
|
||||||
return y
|
|
||||||
|
|
||||||
|
|
||||||
# MLX-accelerated norm ops (1.4x–2.9x faster than torch native on MPS)
|
# MLX-accelerated norm ops (1.4x–2.9x faster than torch native on MPS)
|
||||||
# Uses mx.fast.rms_norm / mx.fast.layer_norm — single fused Metal kernels
|
# Uses mx.fast.rms_norm / mx.fast.layer_norm — single fused Metal kernels
|
||||||
|
|||||||
@@ -653,3 +653,9 @@ if current_platform.is_mps():
|
|||||||
|
|
||||||
norm_infer = norm_infer_native
|
norm_infer = norm_infer_native
|
||||||
rms_norm_fn = rms_norm_fn_native
|
rms_norm_fn = rms_norm_fn_native
|
||||||
|
|
||||||
|
if current_platform.is_cpu():
|
||||||
|
from .torch_fallback import norm_infer_native, rms_norm_fn_native
|
||||||
|
|
||||||
|
norm_infer = norm_infer_native
|
||||||
|
rms_norm_fn = rms_norm_fn_native
|
||||||
|
|||||||
@@ -75,3 +75,9 @@ if current_platform.is_mps():
|
|||||||
@debug_kernel_api
|
@debug_kernel_api
|
||||||
def triton_one_pass_rms_norm(x: torch.Tensor, w: torch.Tensor, eps: float = 1e-6):
|
def triton_one_pass_rms_norm(x: torch.Tensor, w: torch.Tensor, eps: float = 1e-6):
|
||||||
return triton_one_pass_rms_norm_native(x, w, eps)
|
return triton_one_pass_rms_norm_native(x, w, eps)
|
||||||
|
|
||||||
|
|
||||||
|
if current_platform.is_cpu():
|
||||||
|
from .torch_fallback import triton_one_pass_rms_norm_native
|
||||||
|
|
||||||
|
triton_one_pass_rms_norm = triton_one_pass_rms_norm_native
|
||||||
|
|||||||
@@ -134,3 +134,8 @@ if current_platform.is_mps():
|
|||||||
from .mps_fallback import apply_rotary_embedding_native
|
from .mps_fallback import apply_rotary_embedding_native
|
||||||
|
|
||||||
apply_rotary_embedding = apply_rotary_embedding_native
|
apply_rotary_embedding = apply_rotary_embedding_native
|
||||||
|
|
||||||
|
if current_platform.is_cpu():
|
||||||
|
from .torch_fallback import apply_rotary_embedding_native
|
||||||
|
|
||||||
|
apply_rotary_embedding = apply_rotary_embedding_native
|
||||||
|
|||||||
@@ -663,3 +663,10 @@ if current_platform.is_mps():
|
|||||||
from .mps_fallback import fuse_scale_shift_kernel_native
|
from .mps_fallback import fuse_scale_shift_kernel_native
|
||||||
|
|
||||||
fuse_scale_shift_kernel = fuse_scale_shift_kernel_native
|
fuse_scale_shift_kernel = fuse_scale_shift_kernel_native
|
||||||
|
|
||||||
|
if current_platform.is_cpu():
|
||||||
|
from .torch_fallback import (
|
||||||
|
fuse_scale_shift_kernel_native,
|
||||||
|
)
|
||||||
|
|
||||||
|
fuse_scale_shift_kernel = fuse_scale_shift_kernel_native
|
||||||
|
|||||||
@@ -0,0 +1,143 @@
|
|||||||
|
"""Pytorch native based fallbacks for Triton diffusion kernels.
|
||||||
|
|
||||||
|
Triton is not available on some platforms, so these pure-PyTorch
|
||||||
|
implementations replace the Triton kernels
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
|
||||||
|
def fuse_scale_shift_kernel_native(
|
||||||
|
x: torch.Tensor,
|
||||||
|
scale: torch.Tensor,
|
||||||
|
shift: torch.Tensor,
|
||||||
|
scale_constant: float = 1.0,
|
||||||
|
block_l: int = 128,
|
||||||
|
block_c: int = 128,
|
||||||
|
):
|
||||||
|
"""Native fallback for fuse_scale_shift_kernel with scale_constant support."""
|
||||||
|
B, L, C = x.shape
|
||||||
|
|
||||||
|
def _expand(t: torch.Tensor) -> torch.Tensor:
|
||||||
|
if t.dim() == 4:
|
||||||
|
# [B, F, 1, C] -> [B, L, C]
|
||||||
|
num_frames = t.shape[1]
|
||||||
|
frame_seqlen = L // num_frames
|
||||||
|
return (
|
||||||
|
t.squeeze(2)
|
||||||
|
.unsqueeze(2)
|
||||||
|
.expand(-1, -1, frame_seqlen, -1)
|
||||||
|
.reshape(B, L, C)
|
||||||
|
)
|
||||||
|
elif t.dim() == 2:
|
||||||
|
# [B, C] -> [B, 1, C]
|
||||||
|
return t.unsqueeze(1)
|
||||||
|
return t
|
||||||
|
|
||||||
|
scale = _expand(scale)
|
||||||
|
shift = _expand(shift)
|
||||||
|
|
||||||
|
return x * (scale_constant + scale) + shift
|
||||||
|
|
||||||
|
|
||||||
|
def apply_rotary_embedding_native(
|
||||||
|
x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, interleaved: bool = False
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Native fallback for rotary embedding (shared with NPU implementation)."""
|
||||||
|
cos = cos.unsqueeze(-2).to(x.dtype)
|
||||||
|
sin = sin.unsqueeze(-2).to(x.dtype)
|
||||||
|
x1 = x[..., ::2]
|
||||||
|
x2 = x[..., 1::2]
|
||||||
|
o1 = x1 * cos - x2 * sin
|
||||||
|
o2 = x2 * cos + x1 * sin
|
||||||
|
return torch.stack((o1, o2), dim=-1).flatten(-2)
|
||||||
|
|
||||||
|
|
||||||
|
def norm_infer_native(
|
||||||
|
x: Tensor,
|
||||||
|
weight: Optional[Tensor],
|
||||||
|
bias: Optional[Tensor],
|
||||||
|
eps: float,
|
||||||
|
is_rms_norm: bool = False,
|
||||||
|
out: Optional[Tensor] = None,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Native fallback for norm_infer (layer norm / rms norm inference)."""
|
||||||
|
orig_dtype = x.dtype
|
||||||
|
x = x.contiguous().float()
|
||||||
|
if is_rms_norm:
|
||||||
|
variance = x.pow(2).mean(dim=-1, keepdim=True)
|
||||||
|
x_hat = x * torch.rsqrt(variance + eps)
|
||||||
|
else:
|
||||||
|
mean = x.mean(dim=-1, keepdim=True)
|
||||||
|
variance = (x - mean).pow(2).mean(dim=-1, keepdim=True)
|
||||||
|
x_hat = (x - mean) * torch.rsqrt(variance + eps)
|
||||||
|
if weight is not None:
|
||||||
|
x_hat = x_hat * weight.float()
|
||||||
|
if bias is not None:
|
||||||
|
x_hat = x_hat + bias.float()
|
||||||
|
result = x_hat.to(orig_dtype)
|
||||||
|
if out is not None:
|
||||||
|
out.copy_(result)
|
||||||
|
return out
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def triton_one_pass_rms_norm_native(
|
||||||
|
x: torch.Tensor, w: torch.Tensor, eps: float = 1e-6
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Native fallback for triton_one_pass_rms_norm."""
|
||||||
|
shape = x.shape
|
||||||
|
orig_dtype = x.dtype
|
||||||
|
x = x.contiguous().float()
|
||||||
|
variance = x.pow(2).mean(dim=-1, keepdim=True)
|
||||||
|
x_hat = x * torch.rsqrt(variance + eps)
|
||||||
|
return (x_hat * w.float()).to(orig_dtype).view(shape)
|
||||||
|
|
||||||
|
|
||||||
|
def rms_norm_fn_native(
|
||||||
|
x,
|
||||||
|
weight,
|
||||||
|
bias,
|
||||||
|
residual=None,
|
||||||
|
x1=None,
|
||||||
|
weight1=None,
|
||||||
|
bias1=None,
|
||||||
|
eps=1e-6,
|
||||||
|
dropout_p=0.0,
|
||||||
|
rowscale=None,
|
||||||
|
prenorm=False,
|
||||||
|
residual_in_fp32=False,
|
||||||
|
zero_centered_weight=False,
|
||||||
|
return_dropout_mask=False,
|
||||||
|
out_dtype=None,
|
||||||
|
out=None,
|
||||||
|
residual_out=None,
|
||||||
|
):
|
||||||
|
"""Native fallback for rms_norm_fn (inference only, no dropout/x1 support)."""
|
||||||
|
x_shape_og = x.shape
|
||||||
|
orig_dtype = x.dtype
|
||||||
|
x = x.reshape(-1, x.shape[-1]).float()
|
||||||
|
if residual is not None:
|
||||||
|
residual = residual.reshape(-1, residual.shape[-1]).float()
|
||||||
|
x = x + residual
|
||||||
|
residual_out_val = x.to(torch.float32 if residual_in_fp32 else orig_dtype)
|
||||||
|
else:
|
||||||
|
residual_out_val = None
|
||||||
|
variance = x.pow(2).mean(dim=-1, keepdim=True)
|
||||||
|
x_hat = x * torch.rsqrt(variance + eps)
|
||||||
|
if weight is not None:
|
||||||
|
w = weight.float()
|
||||||
|
if zero_centered_weight:
|
||||||
|
w = w + 1.0
|
||||||
|
x_hat = x_hat * w
|
||||||
|
if bias is not None:
|
||||||
|
x_hat = x_hat + bias.float()
|
||||||
|
final_dtype = out_dtype if out_dtype is not None else orig_dtype
|
||||||
|
y = x_hat.to(final_dtype).reshape(x_shape_og)
|
||||||
|
if residual is not None and residual_out_val is not None:
|
||||||
|
return y, residual_out_val.reshape(x_shape_og)
|
||||||
|
return y
|
||||||
@@ -27,6 +27,7 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import (
|
|||||||
init_logger,
|
init_logger,
|
||||||
suppress_stdout,
|
suppress_stdout,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.utils import is_shm_available
|
||||||
|
|
||||||
try:
|
try:
|
||||||
import torch_musa # noqa: F401
|
import torch_musa # noqa: F401
|
||||||
@@ -186,7 +187,6 @@ class GroupCoordinator:
|
|||||||
self.device = get_local_torch_device()
|
self.device = get_local_torch_device()
|
||||||
|
|
||||||
self.use_device_communicator = use_device_communicator
|
self.use_device_communicator = use_device_communicator
|
||||||
|
|
||||||
self.device_communicator: DeviceCommunicatorBase = None # type: ignore
|
self.device_communicator: DeviceCommunicatorBase = None # type: ignore
|
||||||
if use_device_communicator and self.world_size > 1:
|
if use_device_communicator and self.world_size > 1:
|
||||||
# Platform-aware device communicator selection
|
# Platform-aware device communicator selection
|
||||||
@@ -323,6 +323,16 @@ class GroupCoordinator:
|
|||||||
# Bypass the function if we are using only 1 GPU.
|
# Bypass the function if we are using only 1 GPU.
|
||||||
if self.world_size == 1:
|
if self.world_size == 1:
|
||||||
return input_
|
return input_
|
||||||
|
else:
|
||||||
|
if (
|
||||||
|
current_platform.is_cpu()
|
||||||
|
and is_shm_available(input_.dtype, self.world_size, len(self.ranks))
|
||||||
|
and op is torch.distributed.ReduceOp.SUM
|
||||||
|
):
|
||||||
|
# for CPU platform, intra-node case we could speedup with shared memory based comm ops
|
||||||
|
torch.ops.sgl_kernel.shm_allreduce(
|
||||||
|
input_, int(torch.distributed.ReduceOp.SUM)
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
torch.distributed.all_reduce(
|
torch.distributed.all_reduce(
|
||||||
input_, op=op, group=self.device_group, async_op=async_op
|
input_, op=op, group=self.device_group, async_op=async_op
|
||||||
@@ -348,10 +358,17 @@ class GroupCoordinator:
|
|||||||
output_tensor = torch.empty(
|
output_tensor = torch.empty(
|
||||||
input_size, dtype=input_.dtype, device=input_.device
|
input_size, dtype=input_.dtype, device=input_.device
|
||||||
)
|
)
|
||||||
|
|
||||||
# All-gather.
|
# All-gather.
|
||||||
|
if current_platform.is_cpu() and is_shm_available(
|
||||||
|
input_.dtype, self.world_size, len(self.ranks)
|
||||||
|
):
|
||||||
|
return torch.ops.sgl_kernel.shm_allgather(input_, dim)
|
||||||
|
else:
|
||||||
torch.distributed.all_gather_into_tensor(
|
torch.distributed.all_gather_into_tensor(
|
||||||
output_tensor, input_, group=self.device_group
|
output_tensor, input_, group=self.device_group
|
||||||
)
|
)
|
||||||
|
|
||||||
if dim != 0:
|
if dim != 0:
|
||||||
input_size[0] //= world_size
|
input_size[0] //= world_size
|
||||||
output_tensor = output_tensor.reshape(
|
output_tensor = output_tensor.reshape(
|
||||||
|
|||||||
@@ -219,6 +219,7 @@ def init_distributed_environment(
|
|||||||
current_platform.is_mps()
|
current_platform.is_mps()
|
||||||
or current_platform.is_musa()
|
or current_platform.is_musa()
|
||||||
or current_platform.is_npu()
|
or current_platform.is_npu()
|
||||||
|
or current_platform.is_cpu()
|
||||||
or current_platform.is_xpu()
|
or current_platform.is_xpu()
|
||||||
)
|
)
|
||||||
else dict(device_id=device_id)
|
else dict(device_id=device_id)
|
||||||
|
|||||||
@@ -15,7 +15,6 @@ from sglang.jit_kernel.diffusion.qknorm_rope import (
|
|||||||
can_use_fused_inplace_qknorm_rope,
|
can_use_fused_inplace_qknorm_rope,
|
||||||
fused_inplace_qknorm_rope,
|
fused_inplace_qknorm_rope,
|
||||||
)
|
)
|
||||||
from sglang.jit_kernel.diffusion.triton.norm import norm_infer, rms_norm_fn
|
|
||||||
from sglang.jit_kernel.diffusion.triton.rmsnorm_onepass import triton_one_pass_rms_norm
|
from sglang.jit_kernel.diffusion.triton.rmsnorm_onepass import triton_one_pass_rms_norm
|
||||||
from sglang.jit_kernel.diffusion.triton.scale_shift import fuse_scale_shift_kernel
|
from sglang.jit_kernel.diffusion.triton.scale_shift import fuse_scale_shift_kernel
|
||||||
from sglang.jit_kernel.norm import can_use_fused_inplace_qknorm, fused_inplace_qknorm
|
from sglang.jit_kernel.norm import can_use_fused_inplace_qknorm, fused_inplace_qknorm
|
||||||
@@ -31,7 +30,9 @@ from sglang.multimodal_gen.runtime.utils.common import get_bool_env_var
|
|||||||
_is_cuda = current_platform.is_cuda()
|
_is_cuda = current_platform.is_cuda()
|
||||||
_is_npu = current_platform.is_npu()
|
_is_npu = current_platform.is_npu()
|
||||||
_is_musa = current_platform.is_musa()
|
_is_musa = current_platform.is_musa()
|
||||||
|
_is_cpu = current_platform.is_cpu()
|
||||||
_is_xpu = current_platform.is_xpu()
|
_is_xpu = current_platform.is_xpu()
|
||||||
|
|
||||||
if _is_cuda or _is_xpu:
|
if _is_cuda or _is_xpu:
|
||||||
from sgl_kernel import fused_add_rmsnorm, rmsnorm
|
from sgl_kernel import fused_add_rmsnorm, rmsnorm
|
||||||
|
|
||||||
@@ -40,6 +41,8 @@ if _is_npu:
|
|||||||
|
|
||||||
if _is_musa:
|
if _is_musa:
|
||||||
from sgl_kernel import fused_add_rmsnorm
|
from sgl_kernel import fused_add_rmsnorm
|
||||||
|
if not _is_cpu:
|
||||||
|
from sglang.jit_kernel.diffusion.triton.norm import norm_infer, rms_norm_fn
|
||||||
|
|
||||||
|
|
||||||
# Copied and adapted from sglang
|
# Copied and adapted from sglang
|
||||||
|
|||||||
@@ -276,10 +276,15 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
# Determine CPU offload behavior and target device
|
# Determine CPU offload behavior and target device
|
||||||
|
|
||||||
local_torch_device = get_local_torch_device()
|
local_torch_device = get_local_torch_device()
|
||||||
|
|
||||||
|
if not current_platform.is_cpu():
|
||||||
fsdp_cpu_offload = self.should_offload(server_args, model_config)
|
fsdp_cpu_offload = self.should_offload(server_args, model_config)
|
||||||
should_offload = (
|
should_offload = (
|
||||||
cpu_offload_flag if cpu_offload_flag is not None else fsdp_cpu_offload
|
cpu_offload_flag if cpu_offload_flag is not None else fsdp_cpu_offload
|
||||||
)
|
)
|
||||||
|
else:
|
||||||
|
fsdp_cpu_offload = False
|
||||||
|
should_offload = False
|
||||||
|
|
||||||
if should_offload and not current_platform.is_mps():
|
if should_offload and not current_platform.is_mps():
|
||||||
model_device = torch.device("cpu")
|
model_device = torch.device("cpu")
|
||||||
|
|||||||
@@ -0,0 +1,80 @@
|
|||||||
|
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||||
|
|
||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
import os
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
|
from sglang.multimodal_gen.runtime.utils.logging_utils import (
|
||||||
|
init_logger,
|
||||||
|
)
|
||||||
|
from sglang.srt.utils import cpu_has_amx_support, get_cpu_ids_by_node
|
||||||
|
|
||||||
|
from .gpu_worker import GPUWorker
|
||||||
|
|
||||||
|
_is_cpu_amx_available = cpu_has_amx_support()
|
||||||
|
|
||||||
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class CPUWorker(GPUWorker):
|
||||||
|
"""
|
||||||
|
A worker that executes the model on pure CPU platforms
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
local_rank: int,
|
||||||
|
rank: int,
|
||||||
|
master_port: int,
|
||||||
|
server_args: ServerArgs,
|
||||||
|
):
|
||||||
|
super().__init__(local_rank, rank, master_port, server_args)
|
||||||
|
if _is_cpu_amx_available:
|
||||||
|
self.init_cpu_threads_binding()
|
||||||
|
|
||||||
|
def init_cpu_threads_binding(self):
|
||||||
|
omp_cpuids = os.environ.get("SGLANG_CPU_OMP_THREADS_BIND", "all")
|
||||||
|
cpu_ids_by_node = get_cpu_ids_by_node()
|
||||||
|
n_numa_node = len(cpu_ids_by_node)
|
||||||
|
if omp_cpuids == "all":
|
||||||
|
assert self.server_args.tp_size <= n_numa_node, (
|
||||||
|
f"SGLANG_CPU_OMP_THREADS_BIND is not set, in this case, "
|
||||||
|
f"tp_size {self.server_args.tp_size} should be smaller than or equal to number of numa node on the machine {n_numa_node}. "
|
||||||
|
f"If you need tp_size to be larger than number of numa node, please set the CPU cores for each tp rank via SGLANG_CPU_OMP_THREADS_BIND explicitly. "
|
||||||
|
f"For example, on a machine with 2 numa nodes, where core 0-31 are on numa node 0 and core 32-63 are on numa node 1, "
|
||||||
|
f"it is suggested to use -tp 2 and bind tp rank 0 to core 0-31 and tp rank 1 to core 32-63. "
|
||||||
|
f"This is the default behavior if SGLANG_CPU_OMP_THREADS_BIND is not set and it is the same as setting SGLANG_CPU_OMP_THREADS_BIND=0-31|32-63. "
|
||||||
|
f"If you do need tp_size to be larger than the number of numa nodes, you could set SGLANG_CPU_OMP_THREADS_BIND explicitly for example SGLANG_CPU_OMP_THREADS_BIND=0-15|16-31|32-47|48-63 and run with -tp 4. "
|
||||||
|
f"If you don't want each tp rank to use all the cores on one numa node, you could set for example SGLANG_CPU_OMP_THREADS_BIND=0-15|32-47 and run with -tp 2."
|
||||||
|
)
|
||||||
|
if self.server_args.tp_size < n_numa_node:
|
||||||
|
logger.warning(
|
||||||
|
f"Detected the current machine has {n_numa_node} numa nodes available, but tp_size is set to {self.server_args.tp_size}, so only {self.server_args.tp_size} numa nodes are used."
|
||||||
|
)
|
||||||
|
self.local_omp_cpuid = cpu_ids_by_node[self.rank]
|
||||||
|
else:
|
||||||
|
threads_bind_list = omp_cpuids.split("|")
|
||||||
|
assert self.server_args.tp_size == len(threads_bind_list), (
|
||||||
|
f"SGLANG_CPU_OMP_THREADS_BIND setting must be aligned with TP size parameter ({self.server_args.tp_size}). "
|
||||||
|
f"Please double check your settings."
|
||||||
|
)
|
||||||
|
self.local_omp_cpuid = threads_bind_list[self.rank]
|
||||||
|
if self.server_args.tp_size > n_numa_node:
|
||||||
|
logger.warning(
|
||||||
|
f"TP size ({self.server_args.tp_size})is larger than numa node number ({n_numa_node}), "
|
||||||
|
f"in this case the available memory amount of each rank cannot be determined in prior. "
|
||||||
|
f"Please set proper `--max-total-tokens` to avoid the out-of-memory error."
|
||||||
|
)
|
||||||
|
|
||||||
|
# Bind OpenMP threads to CPU cores
|
||||||
|
torch.ops.sgl_kernel.init_cpu_threads_env(self.local_omp_cpuid)
|
||||||
|
|
||||||
|
# Set local size to hint SGLang to use shared memory based AllReduce
|
||||||
|
os.environ["LOCAL_SIZE"] = str(self.server_args.tp_size)
|
||||||
|
torch.ops.sgl_kernel.initialize(self.server_args.tp_size, self.rank)
|
||||||
|
|
||||||
|
@torch.library.register_fake("sgl_kernel::shm_allgather")
|
||||||
|
def _(data, dim):
|
||||||
|
return torch.cat([data] * self.server_args.tp_size, dim=dim)
|
||||||
@@ -223,13 +223,13 @@ class GPUWorker:
|
|||||||
req = batch[0]
|
req = batch[0]
|
||||||
output_batch = None
|
output_batch = None
|
||||||
try:
|
try:
|
||||||
if self.rank == 0:
|
if self.rank == 0 and not current_platform.is_cpu():
|
||||||
torch.get_device_module().reset_peak_memory_stats()
|
torch.get_device_module().reset_peak_memory_stats()
|
||||||
|
|
||||||
start_time = time.monotonic()
|
start_time = time.monotonic()
|
||||||
|
|
||||||
# capture memory baseline before forward
|
# capture memory baseline before forward
|
||||||
if self.rank == 0 and req.metrics:
|
if self.rank == 0 and req.metrics and not current_platform.is_cpu():
|
||||||
baseline_snapshot = capture_memory_snapshot()
|
baseline_snapshot = capture_memory_snapshot()
|
||||||
req.metrics.record_memory_snapshot("before_forward", baseline_snapshot)
|
req.metrics.record_memory_snapshot("before_forward", baseline_snapshot)
|
||||||
|
|
||||||
@@ -259,7 +259,11 @@ class GPUWorker:
|
|||||||
output_batch = result
|
output_batch = result
|
||||||
|
|
||||||
# capture memory after forward (peak)
|
# capture memory after forward (peak)
|
||||||
if self.rank == 0 and output_batch.metrics:
|
if (
|
||||||
|
self.rank == 0
|
||||||
|
and output_batch.metrics
|
||||||
|
and not current_platform.is_cpu()
|
||||||
|
):
|
||||||
peak_snapshot = capture_memory_snapshot()
|
peak_snapshot = capture_memory_snapshot()
|
||||||
output_batch.metrics.record_memory_snapshot(
|
output_batch.metrics.record_memory_snapshot(
|
||||||
"after_forward", peak_snapshot
|
"after_forward", peak_snapshot
|
||||||
@@ -268,6 +272,7 @@ class GPUWorker:
|
|||||||
if (
|
if (
|
||||||
self.rank == 0
|
self.rank == 0
|
||||||
and not req.suppress_logs
|
and not req.suppress_logs
|
||||||
|
and not current_platform.is_cpu()
|
||||||
and logger.isEnabledFor(logging.DEBUG)
|
and logger.isEnabledFor(logging.DEBUG)
|
||||||
):
|
):
|
||||||
self.do_mem_analysis(output_batch)
|
self.do_mem_analysis(output_batch)
|
||||||
|
|||||||
@@ -31,6 +31,7 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import (
|
|||||||
ShutdownReq,
|
ShutdownReq,
|
||||||
UnmergeLoraWeightsReq,
|
UnmergeLoraWeightsReq,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.managers.cpu_worker import CPUWorker
|
||||||
from sglang.multimodal_gen.runtime.managers.gpu_worker import GPUWorker
|
from sglang.multimodal_gen.runtime.managers.gpu_worker import GPUWorker
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core import Req
|
from sglang.multimodal_gen.runtime.pipelines_core import Req
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
|
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
|
||||||
@@ -86,8 +87,10 @@ class Scheduler(SchedulerDisaggMixin):
|
|||||||
logger.info(f"Scheduler bind at endpoint: {actual_endpoint}")
|
logger.info(f"Scheduler bind at endpoint: {actual_endpoint}")
|
||||||
else:
|
else:
|
||||||
self.receiver = None
|
self.receiver = None
|
||||||
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
|
|
||||||
worker = GPUWorker(
|
Exec_worker = CPUWorker if current_platform.is_cpu() else GPUWorker
|
||||||
|
worker = Exec_worker(
|
||||||
local_rank=local_rank,
|
local_rank=local_rank,
|
||||||
master_port=port_args.master_port,
|
master_port=port_args.master_port,
|
||||||
rank=gpu_id,
|
rank=gpu_id,
|
||||||
|
|||||||
@@ -11,10 +11,14 @@ import psutil
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.platforms.interface import (
|
from sglang.multimodal_gen.runtime.platforms.interface import (
|
||||||
|
AttentionBackendEnum,
|
||||||
CpuArchEnum,
|
CpuArchEnum,
|
||||||
Platform,
|
Platform,
|
||||||
PlatformEnum,
|
PlatformEnum,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
|
||||||
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class CpuPlatform(Platform):
|
class CpuPlatform(Platform):
|
||||||
@@ -34,6 +38,10 @@ class CpuPlatform(Platform):
|
|||||||
else:
|
else:
|
||||||
return CpuArchEnum.UNSPECIFIED
|
return CpuArchEnum.UNSPECIFIED
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_local_torch_device(cls) -> torch.device:
|
||||||
|
return torch.device("cpu")
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_device_name(cls, device_id: int = 0) -> str:
|
def get_device_name(cls, device_id: int = 0) -> str:
|
||||||
return platform.processor()
|
return platform.processor()
|
||||||
@@ -86,3 +94,21 @@ class CpuPlatform(Platform):
|
|||||||
@classmethod
|
@classmethod
|
||||||
def get_device_communicator_cls(cls) -> str:
|
def get_device_communicator_cls(cls) -> str:
|
||||||
return "sglang.multimodal_gen.runtime.distributed.device_communicators.cpu_communicator.CpuCommunicator"
|
return "sglang.multimodal_gen.runtime.distributed.device_communicators.cpu_communicator.CpuCommunicator"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_attn_backend_cls_str(
|
||||||
|
cls,
|
||||||
|
selected_backend: AttentionBackendEnum | None,
|
||||||
|
head_size: int,
|
||||||
|
dtype: torch.dtype,
|
||||||
|
) -> str:
|
||||||
|
|
||||||
|
logger.info("Using Torch SDPA backend")
|
||||||
|
return (
|
||||||
|
"sglang.multimodal_gen.runtime.layers.attention.backends.sdpa.SDPABackend"
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def enable_dit_layerwise_offload_for_wan_by_default(cls) -> bool:
|
||||||
|
"""Whether to enable DIT layerwise offload by default on the current platform."""
|
||||||
|
return False
|
||||||
|
|||||||
@@ -303,6 +303,8 @@ class Platform:
|
|||||||
return "mccl"
|
return "mccl"
|
||||||
elif self.is_mps():
|
elif self.is_mps():
|
||||||
return "gloo"
|
return "gloo"
|
||||||
|
elif self.is_cpu():
|
||||||
|
return "gloo"
|
||||||
elif self.is_xpu():
|
elif self.is_xpu():
|
||||||
return "xccl"
|
return "xccl"
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -318,6 +318,7 @@ class ServerArgs(DisaggArgsMixin):
|
|||||||
"""check consistency and raise errors for invalid configs"""
|
"""check consistency and raise errors for invalid configs"""
|
||||||
self._validate_pipeline()
|
self._validate_pipeline()
|
||||||
self._validate_offload()
|
self._validate_offload()
|
||||||
|
if not current_platform.is_cpu():
|
||||||
self._validate_parallelism()
|
self._validate_parallelism()
|
||||||
self._validate_cfg_parallel()
|
self._validate_cfg_parallel()
|
||||||
|
|
||||||
@@ -365,6 +366,10 @@ class ServerArgs(DisaggArgsMixin):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def _adjust_offload(self):
|
def _adjust_offload(self):
|
||||||
|
if current_platform.is_cpu():
|
||||||
|
# CPU platform does not need offload
|
||||||
|
return
|
||||||
|
|
||||||
# TODO: to be handled by each platform
|
# TODO: to be handled by each platform
|
||||||
if current_platform.get_device_total_memory() / BYTES_PER_GB < 30:
|
if current_platform.get_device_total_memory() / BYTES_PER_GB < 30:
|
||||||
logger.info("Enabling all offloading for GPU with low device memory")
|
logger.info("Enabling all offloading for GPU with low device memory")
|
||||||
@@ -549,6 +554,10 @@ class ServerArgs(DisaggArgsMixin):
|
|||||||
ring_unspecified = self.ring_degree is None
|
ring_unspecified = self.ring_degree is None
|
||||||
cfg_unspecified = self.enable_cfg_parallel is None
|
cfg_unspecified = self.enable_cfg_parallel is None
|
||||||
|
|
||||||
|
if current_platform.is_cpu() and self.tp_size > 1:
|
||||||
|
# CPU platform reuse num_gpus to represent num cpu numa nodes as devices
|
||||||
|
self.num_gpus = self.tp_size
|
||||||
|
|
||||||
if self.hsdp_shard_dim is None:
|
if self.hsdp_shard_dim is None:
|
||||||
self.hsdp_shard_dim = self.num_gpus
|
self.hsdp_shard_dim = self.num_gpus
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user