Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 9 additions & 5 deletions python/sglang/kernels/ops/activation/activation.py
Original file line number Diff line number Diff line change
Expand Up @@ -134,17 +134,19 @@ def run_activation(

@register_custom_op(mutates_args=["out"])
def _run_unary_activation_inplace(
op_name: str, input: torch.Tensor, out: torch.Tensor
op_name: str, input: torch.Tensor, out: torch.Tensor, fast_math: bool = True
) -> None:
last = input.shape[-1]
module = activation_module(input.dtype)
module = activation_module(input.dtype, fast_math=fast_math)
module.run_unary_activation(input.view(-1, last), out.view(-1, last), op_name)


def run_unary_activation(
op_name: str,
input: torch.Tensor,
out: Optional[torch.Tensor] = None,
*,
fast_math: bool = True,
) -> torch.Tensor:
"""Apply a standalone (non-gated) element-wise activation: ``out = act(input)``.

Expand All @@ -156,16 +158,18 @@ def run_unary_activation(
)
if out is None:
out = torch.empty_like(input)
_run_unary_activation_inplace(op_name, input, out)
_run_unary_activation_inplace(op_name, input, out, fast_math)
return out


def relu2(
input: torch.Tensor,
out: Optional[torch.Tensor] = None,
*,
fast_math: bool = True,
) -> torch.Tensor:
"""Squared ReLU: ``out = max(0, input) ** 2`` (element-wise)."""
return run_unary_activation("relu2", input, out)
"""Squared ReLU; disable fast math to preserve BF16 subnormal results."""
return run_unary_activation("relu2", input, out, fast_math=fast_math)


def silu_and_mul(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
import torch.nn as nn
import torch.nn.functional as F

from sglang.kernels.ops.activation.activation import relu2
from sglang.kernels.ops.diffusion import (
can_use_fused_inplace_qknorm_rope,
fused_qknorm_rope_pack_kv,
Expand Down Expand Up @@ -81,12 +82,18 @@ def _can_enable_t1_fused_qk_norm_rope(
tp_size: int,
sp_size: int,
is_compiled: bool,
hidden_size: int = 0,
) -> bool:
if is_compiled:
return False
if is_blackwell:
return True
return is_hopper and hidden_act != "relu2" and tp_size == 1 and sp_size == 1
return (
is_hopper
and (hidden_act != "relu2" or hidden_size == 2048)
and tp_size == 1
and sp_size == 1
)


# -----------------------------------------------------------------------------
Expand Down Expand Up @@ -540,8 +547,18 @@ def __init__(

def forward(self, x: torch.Tensor) -> torch.Tensor:
up, _ = self.up_proj(x)
up = F.relu(up)
out, _ = self.down_proj(up * up)
if (
up.is_cuda
and up.dtype == torch.bfloat16
and up.is_contiguous()
and not torch.is_grad_enabled()
and not torch.compiler.is_compiling()
):
up = relu2(up, fast_math=False)
else:
up = F.relu(up)
up = up * up
out, _ = self.down_proj(up)
return out


Expand Down Expand Up @@ -1710,13 +1727,14 @@ def forward(
self._ensure_cache_dicts()

# The T=1 fused path is faster on Blackwell. It also benefits the
# single-GPU Hopper Nano (SwiGLU) workload, while the Hopper
# Cosmos3-Super (dense MLP) multi-GPU workload remains on the split
# single-GPU Hopper Nano (SwiGLU) and Edge (2048-wide dense) workloads,
# while the Hopper Cosmos3-Super (dense MLP) multi-GPU workload remains on the split
# path because that shape regresses with the fusion.
enable_t1_fused_qk_norm_rope = T == 1 and _can_enable_t1_fused_qk_norm_rope(
is_blackwell=current_platform.is_blackwell(),
is_hopper=current_platform.is_hopper(),
hidden_act=self.hidden_act,
hidden_size=self.hidden_size,
tp_size=get_tp_world_size(),
sp_size=get_sp_world_size(),
is_compiled=self._gen_layers_torch_compiled,
Expand Down
24 changes: 24 additions & 0 deletions python/sglang/multimodal_gen/test/unit/test_cosmos3.py
Original file line number Diff line number Diff line change
Expand Up @@ -112,6 +112,7 @@ def _can_enable(
tp_size=1,
sp_size=1,
is_compiled=False,
hidden_size=0,
):
return _can_enable_t1_fused_qk_norm_rope(
is_blackwell=is_blackwell,
Expand All @@ -120,6 +121,7 @@ def _can_enable(
tp_size=tp_size,
sp_size=sp_size,
is_compiled=is_compiled,
hidden_size=hidden_size,
)

def test_blackwell_remains_enabled(self):
Expand All @@ -138,6 +140,28 @@ def test_hopper_swiglu_single_gpu_enabled(self):
def test_hopper_dense_mlp_disabled(self):
self.assertFalse(self._can_enable(is_hopper=True, hidden_act="relu2"))

def test_hopper_edge_single_gpu_enabled(self):
self.assertTrue(
self._can_enable(is_hopper=True, hidden_act="relu2", hidden_size=2048)
)

def test_hopper_edge_parallel_and_compiled_disabled(self):
for override in ({"tp_size": 2}, {"sp_size": 2}, {"is_compiled": True}):
with self.subTest(override=override):
self.assertFalse(
self._can_enable(
is_hopper=True,
hidden_act="relu2",
hidden_size=2048,
**override,
)
)

def test_hopper_larger_dense_mlp_disabled(self):
self.assertFalse(
self._can_enable(is_hopper=True, hidden_act="relu2", hidden_size=4096)
)

def test_hopper_tensor_parallel_disabled(self):
self.assertFalse(self._can_enable(is_hopper=True, tp_size=2))

Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,67 @@
import torch
import torch.nn.functional as F

from sglang.kernels.jit.benchmark import marker
from sglang.kernels.jit.benchmark.utils import create_random
from sglang.kernels.ops.activation.activation import relu2
from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm
from sglang.multimodal_gen.runtime.models.dits.cosmos3video import (
_apply_qwen3_qk_norm_rope_pack_kv,
_apply_qwen3_qk_norm_rope_split,
)
from sglang.test.ci.ci_register import register_cuda_ci

register_cuda_ci(
est_time=15, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)


def eager_relu2(x):
x = F.relu(x)
return x * x


@marker.parametrize("case", ["qk_400", "qk_1024", "relu2_400", "relu2_8190"])
@marker.benchmark("provider", ["eager", "fused"])
def benchmark(case, provider):
operation, tokens = case.split("_")
tokens = int(tokens)
if operation == "relu2":
x = create_random(1, tokens, 9216)
fn = eager_relu2 if provider == "eager" else lambda x: relu2(x, fast_math=False)
return marker.do_bench(fn, input_args=(x,))

qkv = create_random(1, tokens, 32, 128)
k_und, v_und = create_random(1, 32, 8, 128), create_random(1, 32, 8, 128)
angles = torch.randn(tokens, 64, device=qkv.device)
cache = torch.cat((angles.cos(), angles.sin()), -1).to(qkv.dtype)
if provider == "eager":
cache = cache.float()
positions = torch.arange(tokens, device=qkv.device)
q_norm = RMSNorm(128, eps=1e-6).to(device=qkv.device, dtype=qkv.dtype)
k_norm = RMSNorm(128, eps=1e-6).to(device=qkv.device, dtype=qkv.dtype)

def fn(qkv, k_und, v_und, cache, positions):
q, k, v = qkv[:, :, :16], qkv[:, :, 16:24], qkv[:, :, 24:]
if provider == "eager":
q, k = _apply_qwen3_qk_norm_rope_split(q, k, q_norm, k_norm, 128, cache)
return q, torch.cat((k_und, k), dim=1), torch.cat((v_und, v), dim=1)
return _apply_qwen3_qk_norm_rope_pack_kv(
q,
k,
v,
k_und,
v_und,
q_norm,
k_norm,
128,
cache,
positions,
round_norm_before_rope=True,
)

return marker.do_bench(fn, input_args=(qkv, k_und, v_und, cache, positions))


if __name__ == "__main__":
benchmark.run()
131 changes: 131 additions & 0 deletions test/registered/kernels/ops/diffusion/test_cosmos3_edge_fusions.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,131 @@
"""Edge's strided GQA views preserve the split BF16 norm/RoPE and UND cache."""

import sys
from unittest.mock import patch

import pytest
import torch
import torch.nn as nn
import torch.nn.functional as F

import sglang.multimodal_gen.runtime.models.dits.cosmos3video as cosmos3
from sglang.kernels.ops.activation.activation import relu2
from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm
from sglang.multimodal_gen.runtime.models.dits.cosmos3video import (
_apply_qwen3_qk_norm_rope_pack_kv,
_apply_qwen3_qk_norm_rope_split,
)
from sglang.test.ci.ci_register import register_cuda_ci

register_cuda_ci(est_time=10, stage="base-b-kernel-unit", runner_config="4-gpu-b200")
register_cuda_ci(est_time=10, stage="base-b-kernel-unit", runner_config="1-gpu-large")


@pytest.mark.parametrize(
"batch,tokens,prefix", [(1, 400, 32), (1, 1024, 64), (2, 401, 33)]
)
@torch.inference_mode()
def test_edge_qk_rope_pack_matches_split(batch, tokens, prefix):
if torch.cuda.get_device_capability()[0] != 9:
pytest.skip("Split-path bitwise parity covers the newly enabled Hopper path")
torch.manual_seed(42)
qkv = torch.randn(batch, tokens, 32, 128, device="cuda", dtype=torch.bfloat16)
q, k, v = qkv[:, :, :16], qkv[:, :, 16:24], qkv[:, :, 24:]
k_und = torch.randn(batch, prefix, 8, 128, device="cuda", dtype=torch.bfloat16)
v_und = torch.randn_like(k_und)
before = qkv.clone(), k_und.clone(), v_und.clone()
q_norm = RMSNorm(128, eps=1e-6).to(device="cuda", dtype=torch.bfloat16)
k_norm = RMSNorm(128, eps=1e-6).to(device="cuda", dtype=torch.bfloat16)
q_norm.weight.copy_(torch.rand_like(q_norm.weight) + 0.5)
k_norm.weight.copy_(torch.rand_like(k_norm.weight) + 0.5)
angles = torch.randn(batch * tokens, 64, device="cuda")
cache = torch.cat((angles.cos(), angles.sin()), -1).to(torch.bfloat16)
positions = torch.arange(batch * tokens, device="cuda")
q_ref, k_ref = _apply_qwen3_qk_norm_rope_split(
q, k, q_norm, k_norm, 128, cache.float()
)
k_ref = torch.cat((k_und, k_ref), dim=1)
v_ref = torch.cat((v_und, v), dim=1)
q_out, k_out, v_out = _apply_qwen3_qk_norm_rope_pack_kv(
q,
k,
v,
k_und,
v_und,
q_norm,
k_norm,
128,
cache,
positions,
round_norm_before_rope=True,
)
assert torch.equal(q_out, q_ref)
assert torch.equal(k_out, k_ref)
assert torch.equal(v_out, v_ref)
# The in-place Q/K suffix belongs to this forward; V and cached UND K/V
# must remain reusable, including Edge's separately normalized UND keys.
assert torch.equal(v, before[0][:, :, 24:])
assert torch.equal(k_und, before[1])
assert torch.equal(v_und, before[2])


class _IdentityProjection(nn.Module):
def forward(self, x):
return x, None


def _activation_only_mlp():
mlp = cosmos3.Cosmos3DenseMLP.__new__(cosmos3.Cosmos3DenseMLP)
nn.Module.__init__(mlp)
mlp.up_proj = _IdentityProjection()
mlp.down_proj = _IdentityProjection()
return mlp


@torch.inference_mode()
def test_edge_relu2_all_finite_bfloat16_encodings():
values = torch.arange(65536, dtype=torch.int32).to(torch.int16).view(torch.bfloat16)
values = values[torch.isfinite(values)].cuda().reshape(1, -1)
expected = F.relu(values)
expected = expected * expected
with patch.object(cosmos3, "relu2", wraps=relu2) as fused:
actual = _activation_only_mlp()(values)
fused.assert_called_once()
assert torch.equal(actual.view(torch.int16), expected.view(torch.int16))
mlp = _activation_only_mlp()
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
replayed = mlp(values)
values.neg_()
graph.replay()
expected = F.relu(values)
expected = expected * expected
assert torch.equal(replayed.view(torch.int16), expected.view(torch.int16))


@pytest.mark.parametrize("tokens", [400, 8190])
@torch.inference_mode()
def test_edge_relu2_native_shapes(tokens):
values = torch.randn(1, tokens, 9216, device="cuda", dtype=torch.bfloat16)
before = values.clone()
expected = F.relu(values)
expected = expected * expected
with patch.object(cosmos3, "relu2", wraps=relu2) as fused:
actual = _activation_only_mlp()(values)
fused.assert_called_once()
assert torch.equal(actual, expected)
assert torch.equal(values, before)


def test_edge_relu2_grad_falls_back_to_torch():
values = torch.tensor([-2.0, 0.0, 3.0], requires_grad=True)
with patch.object(
cosmos3, "relu2", side_effect=AssertionError("unexpected fusion")
):
actual = _activation_only_mlp()(values)
actual.sum().backward()
torch.testing.assert_close(values.grad, torch.tensor([0.0, 0.0, 6.0]))


if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v"]))
Loading