[Diffusion][CPU] Init CPU platform support for SGLang Diffusion (#20816)

This commit is contained in:
jianan-gu
2026-04-21 14:25:54 +08:00
committed by GitHub
parent 2b2cad70d6
commit 2cf3ac515b
17 changed files with 369 additions and 149 deletions
@@ -17,143 +17,25 @@ from torch import Tensor
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()
if _use_mlx:
import mlx.core as mx
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
# use the common torch native version form torch_fallback
fuse_scale_shift_kernel_native = fuse_scale_shift_kernel_native
apply_rotary_embedding_native = apply_rotary_embedding_native
norm_infer_native = norm_infer_native
triton_one_pass_rms_norm_native = triton_one_pass_rms_norm_native
rms_norm_fn_native = rms_norm_fn_native
# 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
@@ -653,3 +653,9 @@ if current_platform.is_mps():
norm_infer = norm_infer_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
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)
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
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
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