From e97614d10c8e2c90a72387276086dc33f21bba67 Mon Sep 17 00:00:00 2001 From: Jimmy Shong <69131491+Jiminator@users.noreply.github.com> Date: Sun, 20 Sep 2026 08:28:25 -0700 Subject: [PATCH] [Qwen4-Exp] Build the offloaded PLE table on the meta device so --ple-offload-embedding never materialises it on the accelerator (#39928) Co-authored-by: Yangmin Li --- python/sglang/srt/models/qwen4_exp.py | 44 ++++++++++------- .../ops/embeddings/test_qwen4_ple_offload.py | 48 +++++++++++++++++++ 2 files changed, 74 insertions(+), 18 deletions(-) diff --git a/python/sglang/srt/models/qwen4_exp.py b/python/sglang/srt/models/qwen4_exp.py index 34ac3e22d..9389968b5 100644 --- a/python/sglang/srt/models/qwen4_exp.py +++ b/python/sglang/srt/models/qwen4_exp.py @@ -508,20 +508,32 @@ class Qwen4ExpNGramEmbedding(nn.Module): and not self.use_attn_tp_ngram ) ngram_prefix = f"{prefix}.ngram_embedding" if prefix else "ngram_embedding" - self.ngram_embedding = VocabParallelEmbedding( - padded_vocab_size, - self.head_dim_per_ngram, - params_dtype=( - torch.float8_e4m3fn - if _ple_table_is_fp8(config, quant_config, ngram_prefix) - else torch.bfloat16 - ), - output_dtype=torch.bfloat16, - use_attn_tp_group=self.use_attn_tp_ngram, - ) - self.ngram_embedding.register_buffer( + offload_embedding = bool(config.ple_offload_embedding) + # Offload only needs this embedding's metadata: build it on meta so the + # shard is never allocated on the device. + with torch.device("meta") if offload_embedding else nullcontext(): + ngram_embedding = VocabParallelEmbedding( + padded_vocab_size, + self.head_dim_per_ngram, + params_dtype=( + torch.float8_e4m3fn + if _ple_table_is_fp8(config, quant_config, ngram_prefix) + else torch.bfloat16 + ), + output_dtype=torch.bfloat16, + use_attn_tp_group=self.use_attn_tp_ngram, + ) + # weight_scale stays a real device tensor. + ngram_embedding.register_buffer( "weight_scale", torch.ones(1, dtype=torch.bfloat16), persistent=True ) + if offload_embedding: + ngram_embedding = Qwen4ExpPinnedHostEmbedding( + ngram_embedding, + backend=getattr(config, "ple_offload_backend", "pinned"), + table_dir=getattr(config, "ple_offload_dir", None), + ) + self.ngram_embedding = ngram_embedding @classmethod def _splitmix64(cls, x: int) -> int: @@ -771,6 +783,8 @@ class Qwen4ExpPinnedHostEmbedding(VocabParallelEmbedding): The table stays in its checkpoint storage dtype (fp8 with a per-tensor weight_scale for fp8 checkpoints, bf16 otherwise); gathers emit bf16. + + The source weight may be on the meta device; only its metadata is used. """ _COPIED_ATTRIBUTES = ( @@ -931,12 +945,6 @@ class Qwen4ExpPLELayer(nn.Module): quant_config=quant_config, prefix=f"{prefix}.ple_embedding" if prefix else "ple_embedding", ) - if config.ple_offload_embedding: - self.ple_embedding.ngram_embedding = Qwen4ExpPinnedHostEmbedding( - self.ple_embedding.ngram_embedding, - backend=getattr(config, "ple_offload_backend", "pinned"), - table_dir=getattr(config, "ple_offload_dir", None), - ) self.short_conv_dilation = self.ple_embedding.ngram_size self.short_conv_state_len = ( self.conv_kernel_size - 1 diff --git a/test/registered/kernels/ops/embeddings/test_qwen4_ple_offload.py b/test/registered/kernels/ops/embeddings/test_qwen4_ple_offload.py index 2c4a9b03b..b95c4d8ae 100644 --- a/test/registered/kernels/ops/embeddings/test_qwen4_ple_offload.py +++ b/test/registered/kernels/ops/embeddings/test_qwen4_ple_offload.py @@ -6,6 +6,7 @@ import pytest import torch from torch import nn +from sglang.srt.configs.qwen4_exp import Qwen4ExpTextConfig from sglang.srt.layers.quantization.unquant import UnquantizedEmbeddingMethod from sglang.srt.layers.vocab_parallel_embedding import ( VocabParallelEmbeddingShardIndices, @@ -15,6 +16,7 @@ from sglang.srt.models.qwen4_exp import ( Qwen4ExpPinnedHostEmbedding, Qwen4ExpPLELayer, ) +from sglang.srt.runtime_context import get_context, get_parallel from sglang.srt.utils import set_weight_attrs from sglang.test.ci.ci_register import register_cuda_ci @@ -191,6 +193,52 @@ def test_qwen4_ple_prefetch_buffer_lifecycle(monkeypatch): assert set(layer._graph_prefetch_buffers) == {3, 5} +@pytest.fixture +def single_rank_runtime_context(): + """``Qwen4ExpPLELayer.__init__`` reads the TP topology through + ``VocabParallelEmbedding``; pin it to one rank without a process group.""" + override = get_context().override_server_args(tp_size=1) + override.install() + try: + with get_parallel().override( + tp_rank=0, tp_size=1, attn_tp_rank=0, attn_tp_size=1 + ): + yield + finally: + override.restore() + + +def test_qwen4_ple_offload_avoids_device_table(single_rank_runtime_context): + # sgl-project/sglang#39841: the table was built on the device before the + # host table existed, so the flag needed a full per-rank shard of free VRAM. + # Small everywhere except the n-gram table (16 heads x ~20k rows x 4 dims), + # which must dominate the layer's footprint for the peak check to bite. + config = Qwen4ExpTextConfig( + vocab_size=64, + hidden_size=16, + hc_count=2, + ple_embed_dim=64, + ngram_size=3, + heads_per_ngram=8, + ngram_vocab_size_base=20_000, + eos_token_id=1, + ple_offload_embedding=True, + ) + torch.cuda.synchronize() + torch.cuda.reset_peak_memory_stats() + base = torch.cuda.memory_allocated() + with torch.device("cuda"): # the model loader builds every layer this way + layer = Qwen4ExpPLELayer(config, prefix="ple", layer_id=0, ple_layer_index=0) + peak = torch.cuda.max_memory_allocated() - base + + emb = layer.ple_embedding.ngram_embedding + table_bytes = emb.weight.numel() * emb.weight.element_size() + assert peak < table_bytes // 2, (peak, table_bytes) + assert emb.weight.device.type == "cpu" and emb.weight.is_pinned() + assert emb.weight_scale.is_cuda + assert not any(t.is_meta for t in (*layer.parameters(), *layer.buffers())) + + def _file_backend_supported() -> bool: from sglang.srt.models.qwen4_exp_ple_table import device_uses_host_page_tables