[Diffusion] Fuse Qwen-Image final adaptive LayerNorm (#37144)
This commit is contained in:
@@ -16,9 +16,13 @@ from diffusers.models.modeling_outputs import Transformer2DModelOutput
|
|||||||
from diffusers.models.normalization import AdaLayerNormContinuous
|
from diffusers.models.normalization import AdaLayerNormContinuous
|
||||||
|
|
||||||
from sglang.kernels.ops.diffusion import (
|
from sglang.kernels.ops.diffusion import (
|
||||||
|
BitExactFusionGate,
|
||||||
|
can_use_fused_layernorm_modulate,
|
||||||
can_use_linear_gelu,
|
can_use_linear_gelu,
|
||||||
fused_gelu_active,
|
fused_gelu_active,
|
||||||
|
fused_layernorm_modulate_raw,
|
||||||
fused_linear_gelu_tanh,
|
fused_linear_gelu_tanh,
|
||||||
|
is_plain_layer_norm,
|
||||||
mark_fused_gelu_site,
|
mark_fused_gelu_site,
|
||||||
try_fused_bias_mul_add,
|
try_fused_bias_mul_add,
|
||||||
try_fused_bias_scale_residual_norm_scale_shift,
|
try_fused_bias_scale_residual_norm_scale_shift,
|
||||||
@@ -93,6 +97,72 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import
|
|||||||
)
|
)
|
||||||
|
|
||||||
logger = init_logger(__name__) # pylint: disable=invalid-name
|
logger = init_logger(__name__) # pylint: disable=invalid-name
|
||||||
|
_QWEN_NORM_OUT = BitExactFusionGate("Qwen-Image fused norm_out", per_signature=True)
|
||||||
|
_QWEN_NORM_OUT_SIGS = _QWEN_NORM_OUT.verified_sigs
|
||||||
|
assert _QWEN_NORM_OUT_SIGS is not None
|
||||||
|
|
||||||
|
|
||||||
|
def _qwen_norm_out(
|
||||||
|
norm_out: AdaLayerNormContinuous,
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
conditioning_embedding: torch.Tensor,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Bit-exact final LayerNorm and adaLN scale/shift fusion."""
|
||||||
|
if torch.compiler.is_compiling():
|
||||||
|
return norm_out(hidden_states, conditioning_embedding)
|
||||||
|
|
||||||
|
# Keep the projected modulation available for direct eager dispatch. The
|
||||||
|
# public custom-op wrapper costs more than this small final norm itself.
|
||||||
|
emb = norm_out.linear(norm_out.silu(conditioning_embedding).to(hidden_states.dtype))
|
||||||
|
scale, shift = torch.chunk(emb, 2, dim=1)
|
||||||
|
if (
|
||||||
|
_QWEN_NORM_OUT.disabled
|
||||||
|
or not is_plain_layer_norm(norm_out.norm, hidden_states.shape[-1])
|
||||||
|
or not can_use_fused_layernorm_modulate(hidden_states, scale, shift)
|
||||||
|
):
|
||||||
|
return (
|
||||||
|
norm_out.norm(hidden_states) * (1 + scale)[:, None, :] + shift[:, None, :]
|
||||||
|
)
|
||||||
|
|
||||||
|
sig = (
|
||||||
|
hidden_states.dtype,
|
||||||
|
hidden_states.device,
|
||||||
|
hidden_states.shape[0],
|
||||||
|
hidden_states.shape[-1],
|
||||||
|
hidden_states.stride(-1),
|
||||||
|
scale.stride(0) if scale.shape[0] > 1 else hidden_states.shape[-1],
|
||||||
|
shift.stride(0) if shift.shape[0] > 1 else hidden_states.shape[-1],
|
||||||
|
norm_out.norm.eps,
|
||||||
|
)
|
||||||
|
verified = sig in _QWEN_NORM_OUT_SIGS
|
||||||
|
if not verified and torch.cuda.is_current_stream_capturing():
|
||||||
|
return (
|
||||||
|
norm_out.norm(hidden_states) * (1 + scale)[:, None, :] + shift[:, None, :]
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
fused = fused_layernorm_modulate_raw(
|
||||||
|
hidden_states, scale, shift, norm_out.norm.eps
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
_QWEN_NORM_OUT.on_exception(exc, logger=logger)
|
||||||
|
return (
|
||||||
|
norm_out.norm(hidden_states) * (1 + scale)[:, None, :] + shift[:, None, :]
|
||||||
|
)
|
||||||
|
if verified:
|
||||||
|
return fused
|
||||||
|
reference = (
|
||||||
|
norm_out.norm(hidden_states) * (1 + scale)[:, None, :] + shift[:, None, :]
|
||||||
|
)
|
||||||
|
return _QWEN_NORM_OUT.accept_or_fallback(
|
||||||
|
fused,
|
||||||
|
reference,
|
||||||
|
sig=sig,
|
||||||
|
logger=logger,
|
||||||
|
mismatch_msg=(
|
||||||
|
"Qwen-Image fused norm_out is not bit-exact on this platform; "
|
||||||
|
"falling back to eager"
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _attn_mask_meta_local_pad(attn_mask_meta) -> int:
|
def _attn_mask_meta_local_pad(attn_mask_meta) -> int:
|
||||||
@@ -2069,7 +2139,7 @@ class QwenImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
+ controlnet_block_samples[index_block // interval_control]
|
+ controlnet_block_samples[index_block // interval_control]
|
||||||
)
|
)
|
||||||
# Use only the image part (hidden_states) from the dual-stream blocks
|
# Use only the image part (hidden_states) from the dual-stream blocks
|
||||||
hidden_states = self.norm_out(hidden_states, temb_txt)
|
hidden_states = _qwen_norm_out(self.norm_out, hidden_states, temb_txt)
|
||||||
|
|
||||||
output, _ = self.proj_out(hidden_states)
|
output, _ = self.proj_out(hidden_states)
|
||||||
return output
|
return output
|
||||||
|
|||||||
@@ -7,20 +7,23 @@ from sglang.kernels.jit.benchmark.utils import run_benchmark_no_cudagraph
|
|||||||
from sglang.kernels.ops.diffusion import (
|
from sglang.kernels.ops.diffusion import (
|
||||||
fuse_layernorm_scale_shift_gate_select01_kernel,
|
fuse_layernorm_scale_shift_gate_select01_kernel,
|
||||||
fuse_residual_layernorm_scale_shift_gate_select01_kernel,
|
fuse_residual_layernorm_scale_shift_gate_select01_kernel,
|
||||||
|
fused_layernorm_modulate_raw,
|
||||||
norm_infer,
|
norm_infer,
|
||||||
)
|
)
|
||||||
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||||
from sglang.utils import is_in_ci
|
from sglang.utils import is_in_ci
|
||||||
|
|
||||||
register_cuda_ci(
|
register_cuda_ci(
|
||||||
est_time=13, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
|
est_time=15, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
|
||||||
)
|
)
|
||||||
register_amd_ci(est_time=13, stage="jit-kernel-benchmark", runner_config="amd")
|
register_amd_ci(est_time=13, stage="jit-kernel-benchmark", runner_config="amd")
|
||||||
|
|
||||||
if is_in_ci():
|
if is_in_ci():
|
||||||
B_RANGE, S_RANGE, D_RANGE = [1], [128], [3072]
|
B_RANGE, S_RANGE, D_RANGE = [1], [128], [3072]
|
||||||
|
NORM_OUT_S_RANGE = [128]
|
||||||
else:
|
else:
|
||||||
B_RANGE, S_RANGE, D_RANGE = [1, 2], [128, 512, 2048], [1024, 1536, 3072]
|
B_RANGE, S_RANGE, D_RANGE = [1, 2], [128, 512, 2048], [1024, 1536, 3072]
|
||||||
|
NORM_OUT_S_RANGE = [128, 512, 2048, 4096, 4608]
|
||||||
|
|
||||||
DTYPE = torch.bfloat16
|
DTYPE = torch.bfloat16
|
||||||
DEVICE = "cuda"
|
DEVICE = "cuda"
|
||||||
@@ -29,6 +32,7 @@ LINE_VALS = ["split", "fused"]
|
|||||||
LINE_NAMES = ["Triton Norm + Torch Select", "Fused Triton"]
|
LINE_NAMES = ["Triton Norm + Torch Select", "Fused Triton"]
|
||||||
STYLES = [("red", "-"), ("blue", "--")]
|
STYLES = [("red", "-"), ("blue", "--")]
|
||||||
CONFIG = [(b, s, d) for b in B_RANGE for s in S_RANGE for d in D_RANGE]
|
CONFIG = [(b, s, d) for b in B_RANGE for s in S_RANGE for d in D_RANGE]
|
||||||
|
NORM_OUT_CONFIG = [(1, s, 3072) for s in NORM_OUT_S_RANGE]
|
||||||
|
|
||||||
|
|
||||||
def _make_common_inputs(batch_size: int, seq_len: int, hidden_size: int):
|
def _make_common_inputs(batch_size: int, seq_len: int, hidden_size: int):
|
||||||
@@ -62,6 +66,40 @@ def _apply_select01_modulation(
|
|||||||
return x * (1 + scale) + shift, gate
|
return x * (1 + scale) + shift, gate
|
||||||
|
|
||||||
|
|
||||||
|
@triton.testing.perf_report(
|
||||||
|
triton.testing.Benchmark(
|
||||||
|
x_names=["B", "S", "D"],
|
||||||
|
x_vals=NORM_OUT_CONFIG,
|
||||||
|
line_arg="provider",
|
||||||
|
line_vals=LINE_VALS,
|
||||||
|
line_names=["Torch LayerNorm + Modulate", "Fused Triton"],
|
||||||
|
styles=STYLES,
|
||||||
|
ylabel="us",
|
||||||
|
plot_name="qwen_image_norm_out",
|
||||||
|
args={},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
def bench_norm_out(B: int, S: int, D: int, provider: str) -> Tuple[float, float, float]:
|
||||||
|
x = torch.randn(B, S, D, dtype=DTYPE, device=DEVICE)
|
||||||
|
scale = torch.randn(B, D, dtype=DTYPE, device=DEVICE)
|
||||||
|
shift = torch.randn(B, D, dtype=DTYPE, device=DEVICE)
|
||||||
|
norm = torch.nn.LayerNorm(
|
||||||
|
D, eps=EPS, elementwise_affine=False, device=DEVICE, dtype=DTYPE
|
||||||
|
)
|
||||||
|
|
||||||
|
if provider == "split":
|
||||||
|
|
||||||
|
def fn():
|
||||||
|
return norm(x) * (1 + scale)[:, None, :] + shift[:, None, :]
|
||||||
|
|
||||||
|
else:
|
||||||
|
|
||||||
|
def fn():
|
||||||
|
return fused_layernorm_modulate_raw(x, scale, shift, EPS)
|
||||||
|
|
||||||
|
return run_benchmark_no_cudagraph(fn)
|
||||||
|
|
||||||
|
|
||||||
@triton.testing.perf_report(
|
@triton.testing.perf_report(
|
||||||
triton.testing.Benchmark(
|
triton.testing.Benchmark(
|
||||||
x_names=["B", "S", "D"],
|
x_names=["B", "S", "D"],
|
||||||
@@ -176,6 +214,14 @@ def bench_residual_layernorm_scale_shift_gate_select01(
|
|||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
# The bit-exact LayerNorm kernel uses CUDA inline PTX; keep the shared
|
||||||
|
# Qwen modulation benchmark runnable on its registered ROCm lane.
|
||||||
|
if torch.version.hip is None:
|
||||||
|
print(f"\n{'=' * 80}")
|
||||||
|
print("Benchmark: qwen_image norm_out")
|
||||||
|
print(f"{'=' * 80}\n")
|
||||||
|
bench_norm_out.run(print_data=True)
|
||||||
|
|
||||||
print(f"\n{'=' * 80}")
|
print(f"\n{'=' * 80}")
|
||||||
print("Benchmark: qwen_image layernorm + scale_shift_gate_select01")
|
print("Benchmark: qwen_image layernorm + scale_shift_gate_select01")
|
||||||
print(f"{'=' * 80}\n")
|
print(f"{'=' * 80}\n")
|
||||||
|
|||||||
@@ -36,6 +36,7 @@ import sglang.multimodal_gen.runtime.models.dits.flux_2 as flux2
|
|||||||
import sglang.multimodal_gen.runtime.models.dits.glm_image as glm_image
|
import sglang.multimodal_gen.runtime.models.dits.glm_image as glm_image
|
||||||
import sglang.multimodal_gen.runtime.models.dits.longcat_image as longcat_image
|
import sglang.multimodal_gen.runtime.models.dits.longcat_image as longcat_image
|
||||||
import sglang.multimodal_gen.runtime.models.dits.ltx_2 as ltx2_module
|
import sglang.multimodal_gen.runtime.models.dits.ltx_2 as ltx2_module
|
||||||
|
import sglang.multimodal_gen.runtime.models.dits.qwen_image as qwen_image
|
||||||
import sglang.multimodal_gen.runtime.models.dits.sana as sana
|
import sglang.multimodal_gen.runtime.models.dits.sana as sana
|
||||||
from sglang.kernels.ops.diffusion import (
|
from sglang.kernels.ops.diffusion import (
|
||||||
can_use_fused_layernorm_modulate,
|
can_use_fused_layernorm_modulate,
|
||||||
@@ -98,6 +99,7 @@ from sglang.multimodal_gen.runtime.models.dits.ltx_2 import _ltx2_rms_norm_modul
|
|||||||
from sglang.multimodal_gen.runtime.models.dits.qwen_image import (
|
from sglang.multimodal_gen.runtime.models.dits.qwen_image import (
|
||||||
QwenImageTransformerBlock,
|
QwenImageTransformerBlock,
|
||||||
_qwen_modulation_cache_key,
|
_qwen_modulation_cache_key,
|
||||||
|
_qwen_norm_out,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.models.dits.sana import (
|
from sglang.multimodal_gen.runtime.models.dits.sana import (
|
||||||
_eager_ln_modulate as _sana_eager_ln_modulate,
|
_eager_ln_modulate as _sana_eager_ln_modulate,
|
||||||
@@ -151,6 +153,81 @@ def test_bitexact_norm_guards_follow_platform():
|
|||||||
assert can_use_fused_rmsnorm_scale_shift(x, weight, vec, vec) is is_cuda()
|
assert can_use_fused_rmsnorm_scale_shift(x, weight, vec, vec) is is_cuda()
|
||||||
|
|
||||||
|
|
||||||
|
# -------------------------------------------------------------------------
|
||||||
|
# Qwen-Image -- final LayerNorm + adaLN scale/shift
|
||||||
|
# -------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@requires_inline_ptx
|
||||||
|
def test_qwen_norm_out_matches_adaln_reference():
|
||||||
|
qwen_image._QWEN_NORM_OUT.disabled = False
|
||||||
|
qwen_image._QWEN_NORM_OUT.verified = False
|
||||||
|
qwen_image._QWEN_NORM_OUT_SIGS.clear()
|
||||||
|
torch.manual_seed(0)
|
||||||
|
norm_out = (
|
||||||
|
qwen_image.AdaLayerNormContinuous(
|
||||||
|
3072, 3072, elementwise_affine=False, eps=1e-6
|
||||||
|
)
|
||||||
|
.cuda()
|
||||||
|
.bfloat16()
|
||||||
|
)
|
||||||
|
hidden_states = torch.randn(1, 257, 3072, device="cuda", dtype=torch.bfloat16)
|
||||||
|
conditioning = torch.randn(1, 3072, device="cuda", dtype=torch.bfloat16)
|
||||||
|
|
||||||
|
expected = norm_out(hidden_states, conditioning)
|
||||||
|
actual = _qwen_norm_out(norm_out, hidden_states, conditioning)
|
||||||
|
|
||||||
|
assert torch.equal(actual, expected)
|
||||||
|
assert qwen_image._QWEN_NORM_OUT.verified
|
||||||
|
assert not qwen_image._QWEN_NORM_OUT.disabled
|
||||||
|
|
||||||
|
|
||||||
|
def test_qwen_norm_out_preserves_compile_path(monkeypatch):
|
||||||
|
norm_out = (
|
||||||
|
qwen_image.AdaLayerNormContinuous(16, 16, elementwise_affine=False, eps=1e-6)
|
||||||
|
.cuda()
|
||||||
|
.bfloat16()
|
||||||
|
)
|
||||||
|
hidden_states = torch.randn(1, 3, 16, device="cuda", dtype=torch.bfloat16)
|
||||||
|
conditioning = torch.randn(1, 16, device="cuda", dtype=torch.bfloat16)
|
||||||
|
expected = norm_out(hidden_states, conditioning)
|
||||||
|
|
||||||
|
monkeypatch.setattr(torch.compiler, "is_compiling", lambda: True)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
qwen_image,
|
||||||
|
"fused_layernorm_modulate_raw",
|
||||||
|
lambda *args, **kwargs: pytest.fail("compile path must not dispatch kernel"),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert torch.equal(_qwen_norm_out(norm_out, hidden_states, conditioning), expected)
|
||||||
|
|
||||||
|
|
||||||
|
def test_qwen_norm_out_does_not_verify_during_graph_capture(monkeypatch):
|
||||||
|
qwen_image._QWEN_NORM_OUT.disabled = False
|
||||||
|
qwen_image._QWEN_NORM_OUT.verified = False
|
||||||
|
qwen_image._QWEN_NORM_OUT_SIGS.clear()
|
||||||
|
norm_out = (
|
||||||
|
qwen_image.AdaLayerNormContinuous(
|
||||||
|
3072, 3072, elementwise_affine=False, eps=1e-6
|
||||||
|
)
|
||||||
|
.cuda()
|
||||||
|
.bfloat16()
|
||||||
|
)
|
||||||
|
hidden_states = torch.randn(1, 17, 3072, device="cuda", dtype=torch.bfloat16)
|
||||||
|
conditioning = torch.randn(1, 3072, device="cuda", dtype=torch.bfloat16)
|
||||||
|
expected = norm_out(hidden_states, conditioning)
|
||||||
|
|
||||||
|
monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: True)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
qwen_image,
|
||||||
|
"fused_layernorm_modulate_raw",
|
||||||
|
lambda *args, **kwargs: pytest.fail("capture must not verify a new layout"),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert torch.equal(_qwen_norm_out(norm_out, hidden_states, conditioning), expected)
|
||||||
|
assert not qwen_image._QWEN_NORM_OUT_SIGS
|
||||||
|
|
||||||
|
|
||||||
# -------------------------------------------------------------------------
|
# -------------------------------------------------------------------------
|
||||||
# FLUX.1 -- LayerNorm + adaLN modulate, and the shared-FF GELU site
|
# FLUX.1 -- LayerNorm + adaLN modulate, and the shared-FF GELU site
|
||||||
# -------------------------------------------------------------------------
|
# -------------------------------------------------------------------------
|
||||||
|
|||||||
Reference in New Issue
Block a user