Skip to content

vllm.models.qwen4_exp.common.ngram_embedding

Shared Qwen4Exp n-gram embedding storage with device and pinned-host backends.

Both the NVIDIA and AMD Qwen4Exp implementations use these classes so the (large) n-gram embedding table can be kept in pinned host memory and looked up through Unified Virtual Addressing on any CUDA-alike platform.

Classes:

Qwen4ExpPLEDeviceEmbedding

Bases: Qwen4ExpPLEEmbedding

PLE table allocated on the active model device.

Methods:

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
class Qwen4ExpPLEDeviceEmbedding(Qwen4ExpPLEEmbedding):
    """PLE table allocated on the active model device."""

    def allocate_embedding_weight(
        self,
        num_embeddings: int,
        embedding_dim: int,
        dtype: torch.dtype,
    ) -> torch.Tensor:
        """Allocate the complete PLE weight on the active device."""
        return torch.empty(num_embeddings, embedding_dim, dtype=dtype)

    def start_prefetch(
        self,
        hidden_states: torch.Tensor,
        ngram_ids: torch.Tensor,
    ) -> None:
        """Resident embedding prefetch is a no-op."""
        return None

    def forward(self, ngram_ids: torch.Tensor) -> torch.Tensor:
        """Gather ETP inputs, look up embeddings, and select local rows."""
        slot_size, slot_offset = self._get_dp_gather_slot(ngram_ids.shape[0])
        gathered_ids = self._gather_dp_ids(ngram_ids, slot_size)
        embeddings = super().forward(gathered_ids)
        return self._select_embeddings(
            embeddings,
            ngram_ids.shape[0],
            slot_offset,
        )

allocate_embedding_weight(num_embeddings, embedding_dim, dtype)

Allocate the complete PLE weight on the active device.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
def allocate_embedding_weight(
    self,
    num_embeddings: int,
    embedding_dim: int,
    dtype: torch.dtype,
) -> torch.Tensor:
    """Allocate the complete PLE weight on the active device."""
    return torch.empty(num_embeddings, embedding_dim, dtype=dtype)

forward(ngram_ids)

Gather ETP inputs, look up embeddings, and select local rows.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
def forward(self, ngram_ids: torch.Tensor) -> torch.Tensor:
    """Gather ETP inputs, look up embeddings, and select local rows."""
    slot_size, slot_offset = self._get_dp_gather_slot(ngram_ids.shape[0])
    gathered_ids = self._gather_dp_ids(ngram_ids, slot_size)
    embeddings = super().forward(gathered_ids)
    return self._select_embeddings(
        embeddings,
        ngram_ids.shape[0],
        slot_offset,
    )

start_prefetch(hidden_states, ngram_ids)

Resident embedding prefetch is a no-op.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
def start_prefetch(
    self,
    hidden_states: torch.Tensor,
    ngram_ids: torch.Tensor,
) -> None:
    """Resident embedding prefetch is a no-op."""
    return None

Qwen4ExpPLEEmbedding

Bases: PLEVocabParallelEmbedding, ABC

ETP-sharded PLE table shared by device and pinned-host backends.

Methods:

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
class Qwen4ExpPLEEmbedding(PLEVocabParallelEmbedding, ABC):
    """ETP-sharded PLE table shared by device and pinned-host backends."""

    supports_prefetch: ClassVar[bool] = False

    def __init__(
        self,
        num_embeddings: int,
        embedding_dim: int,
        *,
        params_dtype: torch.dtype,
        padding_size: int,
        prefix: str,
        embedding_method: "Qwen4ExpPLEEmbeddingMethod",
        num_ngram_heads: int = 1,
        max_total_tokens: int = 0,
        data_parallel_rank: int = 0,
    ) -> None:
        del num_ngram_heads, max_total_tokens
        super().__init__(
            num_embeddings,
            embedding_dim,
            params_dtype=params_dtype,
            padding_size=padding_size,
            prefix=prefix,
            quant_method=embedding_method,
            parallel_group=get_etp_group(),
        )
        self.embedding_method = embedding_method
        self.data_parallel_rank = data_parallel_rank
        tp_size = get_tp_group().world_size
        if self.tp_size % tp_size:
            raise ValueError(
                "ETP size must be divisible by TP size, but got "
                f"ETP={self.tp_size} and TP={tp_size}"
            )
        self.etp_data_parallel_size = self.tp_size // tp_size

    @abstractmethod
    def allocate_embedding_weight(
        self,
        num_embeddings: int,
        embedding_dim: int,
        dtype: torch.dtype,
    ) -> torch.Tensor:
        """Allocate storage for the complete embedding weight."""
        raise NotImplementedError

    def dequantize(
        self,
        embeddings: torch.Tensor,
        output_dtype: torch.dtype,
    ) -> torch.Tensor:
        """Delegate storage-format conversion to the embedding method."""
        return self.embedding_method.dequantize(self, embeddings, output_dtype)

    def _get_dp_gather_slot(self, local_num_tokens: int) -> tuple[int, int]:
        """Return the per-DP slot size and this rank's slot offset."""
        if self.etp_data_parallel_size == 1:
            return local_num_tokens, 0
        dp_metadata: DPMetadata | None = get_forward_context().dp_metadata
        if dp_metadata is None:
            raise RuntimeError("ETP spanning DP requires DP token metadata")
        group_start = (self.data_parallel_rank // self.etp_data_parallel_size) * (
            self.etp_data_parallel_size
        )
        group_end = group_start + self.etp_data_parallel_size
        token_counts = dp_metadata.num_tokens_across_dp_cpu.tolist()
        group_counts = token_counts[group_start:group_end]
        slot_size = max(group_counts)
        dp_rank = get_dp_group().rank_in_group
        return slot_size, dp_rank * slot_size

    def _gather_dp_ids(
        self,
        ngram_ids: torch.Tensor,
        slot_size: int,
    ) -> torch.Tensor:
        """Gather DP-local IDs that share one ETP-sharded PLE table."""
        if self.etp_data_parallel_size == 1:
            return ngram_ids
        if ngram_ids.shape[0] < slot_size:
            padding = ngram_ids.new_zeros(
                slot_size - ngram_ids.shape[0], ngram_ids.shape[1]
            )
            ngram_ids = torch.cat((ngram_ids, padding), dim=0)
        return get_dp_group().all_gather(ngram_ids, dim=0)

    def _select_embeddings(
        self,
        embeddings: torch.Tensor,
        local_num_tokens: int,
        slot_offset: int,
    ) -> torch.Tensor:
        """Select this DP rank's rows from the ETP-reduced embeddings."""
        if self.etp_data_parallel_size == 1:
            return embeddings
        return embeddings.narrow(0, slot_offset, local_num_tokens)

    @abstractmethod
    def start_prefetch(
        self,
        hidden_states: torch.Tensor,
        ngram_ids: torch.Tensor,
    ) -> None:
        """Start an asynchronous lookup when supported."""
        raise NotImplementedError

_gather_dp_ids(ngram_ids, slot_size)

Gather DP-local IDs that share one ETP-sharded PLE table.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
def _gather_dp_ids(
    self,
    ngram_ids: torch.Tensor,
    slot_size: int,
) -> torch.Tensor:
    """Gather DP-local IDs that share one ETP-sharded PLE table."""
    if self.etp_data_parallel_size == 1:
        return ngram_ids
    if ngram_ids.shape[0] < slot_size:
        padding = ngram_ids.new_zeros(
            slot_size - ngram_ids.shape[0], ngram_ids.shape[1]
        )
        ngram_ids = torch.cat((ngram_ids, padding), dim=0)
    return get_dp_group().all_gather(ngram_ids, dim=0)

_get_dp_gather_slot(local_num_tokens)

Return the per-DP slot size and this rank's slot offset.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
def _get_dp_gather_slot(self, local_num_tokens: int) -> tuple[int, int]:
    """Return the per-DP slot size and this rank's slot offset."""
    if self.etp_data_parallel_size == 1:
        return local_num_tokens, 0
    dp_metadata: DPMetadata | None = get_forward_context().dp_metadata
    if dp_metadata is None:
        raise RuntimeError("ETP spanning DP requires DP token metadata")
    group_start = (self.data_parallel_rank // self.etp_data_parallel_size) * (
        self.etp_data_parallel_size
    )
    group_end = group_start + self.etp_data_parallel_size
    token_counts = dp_metadata.num_tokens_across_dp_cpu.tolist()
    group_counts = token_counts[group_start:group_end]
    slot_size = max(group_counts)
    dp_rank = get_dp_group().rank_in_group
    return slot_size, dp_rank * slot_size

_select_embeddings(embeddings, local_num_tokens, slot_offset)

Select this DP rank's rows from the ETP-reduced embeddings.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
def _select_embeddings(
    self,
    embeddings: torch.Tensor,
    local_num_tokens: int,
    slot_offset: int,
) -> torch.Tensor:
    """Select this DP rank's rows from the ETP-reduced embeddings."""
    if self.etp_data_parallel_size == 1:
        return embeddings
    return embeddings.narrow(0, slot_offset, local_num_tokens)

allocate_embedding_weight(num_embeddings, embedding_dim, dtype) abstractmethod

Allocate storage for the complete embedding weight.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
@abstractmethod
def allocate_embedding_weight(
    self,
    num_embeddings: int,
    embedding_dim: int,
    dtype: torch.dtype,
) -> torch.Tensor:
    """Allocate storage for the complete embedding weight."""
    raise NotImplementedError

dequantize(embeddings, output_dtype)

Delegate storage-format conversion to the embedding method.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
def dequantize(
    self,
    embeddings: torch.Tensor,
    output_dtype: torch.dtype,
) -> torch.Tensor:
    """Delegate storage-format conversion to the embedding method."""
    return self.embedding_method.dequantize(self, embeddings, output_dtype)

start_prefetch(hidden_states, ngram_ids) abstractmethod

Start an asynchronous lookup when supported.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
@abstractmethod
def start_prefetch(
    self,
    hidden_states: torch.Tensor,
    ngram_ids: torch.Tensor,
) -> None:
    """Start an asynchronous lookup when supported."""
    raise NotImplementedError

Qwen4ExpPLEEmbeddingMethod

Bases: QuantizeMethodBase

Quantization interface shared by resident and pinned PLE tables.

Methods:

  • dequantize –

    Convert looked-up PLE rows to the activation dtype.

  • from_quant_config –

    Select the concrete PLE embedding format for a layer.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
class Qwen4ExpPLEEmbeddingMethod(QuantizeMethodBase):
    """Quantization interface shared by resident and pinned PLE tables."""

    # PLE post-load processing only validates scales in their current storage.
    requires_device_loading: bool = False

    @staticmethod
    def from_quant_config(
        quant_config: QuantizationConfig | None,
        prefix: str,
        embedding_dtype: str | None = None,
    ) -> "Qwen4ExpPLEEmbeddingMethod":
        """Select the concrete PLE embedding format for a layer."""
        if embedding_dtype == "float8_e4m3fn":
            return Qwen4ExpPLEFp8EmbeddingMethod()
        if quant_config is None:
            return Qwen4ExpPLEUnquantizedEmbeddingMethod()
        if isinstance(quant_config, ModelOptMixedPrecisionConfig):
            if quant_config._resolve_quant_algo(prefix) == "FP8":
                return Qwen4ExpPLEFp8EmbeddingMethod()
            return Qwen4ExpPLEUnquantizedEmbeddingMethod()
        if isinstance(
            quant_config, ModelOptQuantConfigBase
        ) and quant_config.is_layer_excluded(prefix):
            return Qwen4ExpPLEUnquantizedEmbeddingMethod()
        if (
            isinstance(quant_config, CompressedTensorsConfig)
            and quant_config.get_scheme_dict(None, layer_name=prefix) is None
        ):
            return Qwen4ExpPLEUnquantizedEmbeddingMethod()
        if (
            isinstance(quant_config, INCConfig)
            and not quant_config.config_parser.resolve(None, prefix).quantized
        ):
            return Qwen4ExpPLEUnquantizedEmbeddingMethod()
        # Quark quantizes only Linear and MoE layers; PLE tables stay BF16.
        if isinstance(quant_config, QuarkConfig):
            return Qwen4ExpPLEUnquantizedEmbeddingMethod()
        if not isinstance(quant_config, Fp8Config):
            raise NotImplementedError(
                "Qwen4Exp PLE embedding does not support quantization config "
                f"{type(quant_config).__name__}"
            )

        ignored_layers = quant_config.ignored_layers
        if is_layer_skipped(
            prefix,
            ignored_layers,
            quant_config.packed_modules_mapping,
            match_mode=quant_config.ignored_layers_match_mode,
        ):
            return Qwen4ExpPLEUnquantizedEmbeddingMethod()
        # PLE checkpoint shards form one runtime embedding parameter.
        shard_prefix = f"{prefix}.shard_"
        if any(name.startswith(shard_prefix) for name in ignored_layers):
            return Qwen4ExpPLEUnquantizedEmbeddingMethod()
        if not quant_config.is_checkpoint_fp8_serialized:
            raise NotImplementedError(
                "Qwen4Exp PLE embedding only supports serialized FP8 checkpoints"
            )
        return Qwen4ExpPLEFp8EmbeddingMethod()

    def apply(
        self,
        layer: nn.Module,
        x: torch.Tensor,
        bias: torch.Tensor | None = None,
    ) -> torch.Tensor:
        raise NotImplementedError("PLE weights only support embedding lookup")

    def embedding(self, layer: nn.Module, input_: torch.Tensor) -> torch.Tensor:
        return F.embedding(input_, layer.weight)

    @abstractmethod
    def dequantize(
        self,
        layer: nn.Module,
        embeddings: torch.Tensor,
        output_dtype: torch.dtype,
    ) -> torch.Tensor:
        """Convert looked-up PLE rows to the activation dtype."""
        raise NotImplementedError

dequantize(layer, embeddings, output_dtype) abstractmethod

Convert looked-up PLE rows to the activation dtype.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
@abstractmethod
def dequantize(
    self,
    layer: nn.Module,
    embeddings: torch.Tensor,
    output_dtype: torch.dtype,
) -> torch.Tensor:
    """Convert looked-up PLE rows to the activation dtype."""
    raise NotImplementedError

from_quant_config(quant_config, prefix, embedding_dtype=None) staticmethod

Select the concrete PLE embedding format for a layer.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
@staticmethod
def from_quant_config(
    quant_config: QuantizationConfig | None,
    prefix: str,
    embedding_dtype: str | None = None,
) -> "Qwen4ExpPLEEmbeddingMethod":
    """Select the concrete PLE embedding format for a layer."""
    if embedding_dtype == "float8_e4m3fn":
        return Qwen4ExpPLEFp8EmbeddingMethod()
    if quant_config is None:
        return Qwen4ExpPLEUnquantizedEmbeddingMethod()
    if isinstance(quant_config, ModelOptMixedPrecisionConfig):
        if quant_config._resolve_quant_algo(prefix) == "FP8":
            return Qwen4ExpPLEFp8EmbeddingMethod()
        return Qwen4ExpPLEUnquantizedEmbeddingMethod()
    if isinstance(
        quant_config, ModelOptQuantConfigBase
    ) and quant_config.is_layer_excluded(prefix):
        return Qwen4ExpPLEUnquantizedEmbeddingMethod()
    if (
        isinstance(quant_config, CompressedTensorsConfig)
        and quant_config.get_scheme_dict(None, layer_name=prefix) is None
    ):
        return Qwen4ExpPLEUnquantizedEmbeddingMethod()
    if (
        isinstance(quant_config, INCConfig)
        and not quant_config.config_parser.resolve(None, prefix).quantized
    ):
        return Qwen4ExpPLEUnquantizedEmbeddingMethod()
    # Quark quantizes only Linear and MoE layers; PLE tables stay BF16.
    if isinstance(quant_config, QuarkConfig):
        return Qwen4ExpPLEUnquantizedEmbeddingMethod()
    if not isinstance(quant_config, Fp8Config):
        raise NotImplementedError(
            "Qwen4Exp PLE embedding does not support quantization config "
            f"{type(quant_config).__name__}"
        )

    ignored_layers = quant_config.ignored_layers
    if is_layer_skipped(
        prefix,
        ignored_layers,
        quant_config.packed_modules_mapping,
        match_mode=quant_config.ignored_layers_match_mode,
    ):
        return Qwen4ExpPLEUnquantizedEmbeddingMethod()
    # PLE checkpoint shards form one runtime embedding parameter.
    shard_prefix = f"{prefix}.shard_"
    if any(name.startswith(shard_prefix) for name in ignored_layers):
        return Qwen4ExpPLEUnquantizedEmbeddingMethod()
    if not quant_config.is_checkpoint_fp8_serialized:
        raise NotImplementedError(
            "Qwen4Exp PLE embedding only supports serialized FP8 checkpoints"
        )
    return Qwen4ExpPLEFp8EmbeddingMethod()

Qwen4ExpPLEFp8EmbeddingMethod

Bases: Qwen4ExpPLEEmbeddingMethod

FP8 PLE embedding with one global checkpoint scale.

Methods:

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
class Qwen4ExpPLEFp8EmbeddingMethod(Qwen4ExpPLEEmbeddingMethod):
    """FP8 PLE embedding with one global checkpoint scale."""

    def create_weights(
        self,
        layer: Qwen4ExpPLEEmbedding,
        input_size_per_partition: int,
        output_partition_sizes: list[int],
        input_size: int,
        output_size: int,
        params_dtype: torch.dtype,
        **extra_weight_attrs,
    ) -> None:
        del input_size, output_size, params_dtype
        weight_loader = extra_weight_attrs.get("weight_loader")
        weight = ModelWeightParameter(
            data=layer.allocate_embedding_weight(
                sum(output_partition_sizes),
                input_size_per_partition,
                torch.float8_e4m3fn,
            ),
            input_dim=1,
            output_dim=0,
            weight_loader=weight_loader,
        )
        layer.register_parameter("weight", weight)

        weight_scale = create_fp8_scale_parameter(
            PerTensorScaleParameter,
            output_partition_sizes,
            input_size_per_partition,
            None,
            weight_loader,
            scale_dtype=torch.float32,
        )
        layer.register_parameter("weight_scale", weight_scale)

    def process_weights_after_loading(self, layer: nn.Module) -> None:
        """Reject FP8 PLE checkpoints without a global scale."""
        sentinel = torch.finfo(torch.float32).min
        if torch.any(layer.weight_scale == sentinel):
            raise ValueError("FP8 PLE checkpoint is missing its global scale")

    def dequantize(
        self,
        layer: nn.Module,
        embeddings: torch.Tensor,
        output_dtype: torch.dtype,
    ) -> torch.Tensor:
        weight_scale = getattr(layer, "weight_scale", None)
        if weight_scale is None:
            raise RuntimeError("FP8 PLE embedding is missing its global scale")
        if weight_scale.device != embeddings.device:
            raise RuntimeError("FP8 PLE embedding scale must be on the output device")
        return embeddings.to(output_dtype) * weight_scale.to(output_dtype)

process_weights_after_loading(layer)

Reject FP8 PLE checkpoints without a global scale.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
def process_weights_after_loading(self, layer: nn.Module) -> None:
    """Reject FP8 PLE checkpoints without a global scale."""
    sentinel = torch.finfo(torch.float32).min
    if torch.any(layer.weight_scale == sentinel):
        raise ValueError("FP8 PLE checkpoint is missing its global scale")

Qwen4ExpPLEPinnedHostEmbedding

Bases: Qwen4ExpPLEEmbedding

PLE table loaded into pinned CPU memory and looked up through UVA.

Methods:

  • allocate_embedding_weight –

    Allocate the complete PLE weight directly in pinned CPU memory.

  • forward –

    Finish the pinned lookup into graph-owned output storage.

  • start_prefetch –

    Gather ETP IDs and launch their UVA lookup on the side stream.

  • sync_lookup –

    Synchronous UVA lookup for platforms without prefetch wiring.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
class Qwen4ExpPLEPinnedHostEmbedding(Qwen4ExpPLEEmbedding):
    """PLE table loaded into pinned CPU memory and looked up through UVA."""

    supports_prefetch: ClassVar[bool] = True

    def __init__(
        self,
        num_embeddings: int,
        embedding_dim: int,
        *,
        params_dtype: torch.dtype,
        padding_size: int,
        prefix: str,
        embedding_method: Qwen4ExpPLEEmbeddingMethod,
        num_ngram_heads: int = 1,
        max_total_tokens: int = 0,
        data_parallel_rank: int = 0,
    ) -> None:
        if not is_uva_available():
            raise RuntimeError("Engram CPU offload requires UVA support")
        super().__init__(
            num_embeddings,
            embedding_dim,
            params_dtype=params_dtype,
            padding_size=padding_size,
            prefix=prefix,
            embedding_method=embedding_method,
            num_ngram_heads=num_ngram_heads,
            max_total_tokens=max_total_tokens,
            data_parallel_rank=data_parallel_rank,
        )
        self._uva_weight = get_accelerator_view_from_cpu_tensor(self.weight)
        self._row_bytes = self.embedding_dim * self.weight.element_size()
        self._block_d = triton.next_power_of_2(self._row_bytes)
        self._prefetch_stream: torch.cuda.Stream | None = None
        self._prefetch_buffer: torch.Tensor | None = None
        self._prefetch_alloc_lock = threading.Lock()
        self._prefetch_rows = max_total_tokens * self.etp_data_parallel_size
        self._num_ngram_heads = num_ngram_heads
        self._output_dim = num_ngram_heads * self.embedding_dim

    def allocate_embedding_weight(
        self,
        num_embeddings: int,
        embedding_dim: int,
        dtype: torch.dtype,
    ) -> torch.Tensor:
        """Allocate the complete PLE weight directly in pinned CPU memory."""
        return torch.empty(
            num_embeddings,
            embedding_dim,
            dtype=dtype,
            device="cpu",
            pin_memory=True,
        )

    def _lookup(
        self,
        input_ids: torch.Tensor,
        output: torch.Tensor | None = None,
    ) -> torch.Tensor:
        """Look up local ETP rows while preserving the weight storage dtype."""
        expected_shape = (*input_ids.shape, self.embedding_dim)
        if output is None:
            output = torch.empty(
                expected_shape,
                dtype=self.weight.dtype,
                device=input_ids.device,
            )
        elif (
            tuple(output.shape) != expected_shape
            or output.dtype != self.weight.dtype
            or output.device != input_ids.device
        ):
            raise ValueError(
                "PLE prefetch output must match the input shape, weight dtype, "
                "and input device"
            )

        flat_ids = input_ids.reshape(-1).long()
        if flat_ids.numel():
            _lookup_ple_embedding_from_pinned_kernel[(flat_ids.numel(),)](
                self._uva_weight,
                flat_ids,
                output,
                self._row_bytes,
                self.shard_indices.org_vocab_start_index,
                self.shard_indices.org_vocab_end_index,
                BLOCK_D=self._block_d,
            )
        return output

    def sync_lookup(self, ngram_ids: torch.Tensor) -> torch.Tensor:
        """Synchronous UVA lookup for platforms without prefetch wiring."""
        slot_size, slot_offset = self._get_dp_gather_slot(ngram_ids.shape[0])
        gathered_ids = self._gather_dp_ids(ngram_ids, slot_size)
        embeddings = self._lookup(gathered_ids)
        embeddings = self._reduce_etp_embeddings(embeddings)
        return self._select_embeddings(
            embeddings,
            ngram_ids.shape[0],
            slot_offset,
        )

    def _reduce_etp_embeddings(self, embeddings: torch.Tensor) -> torch.Tensor:
        """Combine pinned lookup results owned by different ETP ranks."""
        if self.tp_size == 1:
            return embeddings
        assert self.parallel_group is not None
        if embeddings.dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
            # Each vocabulary row has one owner, so reduce the raw FP8 bytes.
            reduced = self.parallel_group.all_reduce(embeddings.view(torch.int8))
            return reduced.view(embeddings.dtype)
        return self.parallel_group.all_reduce(embeddings)

    @eager_break_during_capture
    def start_prefetch(
        self,
        hidden_states: torch.Tensor,
        ngram_ids: torch.Tensor,
    ) -> None:
        """Gather ETP IDs and launch their UVA lookup on the side stream."""
        buffer = self._prefetch_buffer
        if buffer is None:
            # First use allocates. The eager profile run always precedes
            # cudagraph capture, so allocation never happens mid-capture;
            # the lock keeps concurrent first callers from tearing the
            # stream/buffer pair.
            with self._prefetch_alloc_lock:
                buffer = self._prefetch_buffer
                if buffer is None:
                    if torch.cuda.is_current_stream_capturing():
                        raise RuntimeError(
                            "pinned PLE prefetch buffer must be allocated "
                            "eagerly, before cudagraph capture"
                        )
                    self._prefetch_stream = torch.cuda.Stream(
                        device=self._uva_weight.device
                    )
                    buffer = torch.empty(
                        self._prefetch_rows,
                        self._num_ngram_heads,
                        self.embedding_dim,
                        dtype=self.weight.dtype,
                        device=self._uva_weight.device,
                    )
                    self._prefetch_buffer = buffer
        prefetch_stream = self._prefetch_stream
        if prefetch_stream is None:
            raise RuntimeError("pinned PLE prefetch stream was not allocated")
        slot_size, _ = self._get_dp_gather_slot(ngram_ids.shape[0])
        gathered_ids = self._gather_dp_ids(ngram_ids, slot_size)
        if gathered_ids.shape[0] > buffer.shape[0]:
            raise ValueError(
                f"pinned PLE prefetch buffer holds {buffer.shape[0]} rows, "
                f"but the batch needs {gathered_ids.shape[0]}"
            )
        active_output = buffer[: gathered_ids.shape[0]]
        prefetch_stream.wait_stream(torch.cuda.current_stream())
        gathered_ids.record_stream(prefetch_stream)
        with torch.cuda.stream(prefetch_stream):
            self._lookup(gathered_ids, output=active_output)

    @eager_break_during_capture
    def _finalize_prefetch(
        self,
        prefetch_output: torch.Tensor,
        output: torch.Tensor,
    ) -> None:
        """Join the side stream, reduce ETP shards, and select local rows."""
        prefetch_stream = self._prefetch_stream
        if prefetch_stream is None:
            raise RuntimeError("pinned PLE finalize requires a prior start_prefetch")
        torch.cuda.current_stream().wait_stream(prefetch_stream)
        slot_size, slot_offset = self._get_dp_gather_slot(output.shape[0])
        active_output = prefetch_output[: slot_size * self.etp_data_parallel_size]
        embeddings = self._reduce_etp_embeddings(active_output)
        embeddings = self._select_embeddings(
            embeddings,
            output.shape[0],
            slot_offset,
        )
        output.copy_(embeddings.flatten(-2))

    def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
        """Finish the pinned lookup into graph-owned output storage."""
        buffer = self._prefetch_buffer
        if buffer is None:
            raise RuntimeError("pinned PLE lookup requires a prior start_prefetch")
        output = buffer.new_empty((hidden_states.shape[0], self._output_dim))
        self._finalize_prefetch(buffer, output)
        return output

_finalize_prefetch(prefetch_output, output)

Join the side stream, reduce ETP shards, and select local rows.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
@eager_break_during_capture
def _finalize_prefetch(
    self,
    prefetch_output: torch.Tensor,
    output: torch.Tensor,
) -> None:
    """Join the side stream, reduce ETP shards, and select local rows."""
    prefetch_stream = self._prefetch_stream
    if prefetch_stream is None:
        raise RuntimeError("pinned PLE finalize requires a prior start_prefetch")
    torch.cuda.current_stream().wait_stream(prefetch_stream)
    slot_size, slot_offset = self._get_dp_gather_slot(output.shape[0])
    active_output = prefetch_output[: slot_size * self.etp_data_parallel_size]
    embeddings = self._reduce_etp_embeddings(active_output)
    embeddings = self._select_embeddings(
        embeddings,
        output.shape[0],
        slot_offset,
    )
    output.copy_(embeddings.flatten(-2))

_lookup(input_ids, output=None)

Look up local ETP rows while preserving the weight storage dtype.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
def _lookup(
    self,
    input_ids: torch.Tensor,
    output: torch.Tensor | None = None,
) -> torch.Tensor:
    """Look up local ETP rows while preserving the weight storage dtype."""
    expected_shape = (*input_ids.shape, self.embedding_dim)
    if output is None:
        output = torch.empty(
            expected_shape,
            dtype=self.weight.dtype,
            device=input_ids.device,
        )
    elif (
        tuple(output.shape) != expected_shape
        or output.dtype != self.weight.dtype
        or output.device != input_ids.device
    ):
        raise ValueError(
            "PLE prefetch output must match the input shape, weight dtype, "
            "and input device"
        )

    flat_ids = input_ids.reshape(-1).long()
    if flat_ids.numel():
        _lookup_ple_embedding_from_pinned_kernel[(flat_ids.numel(),)](
            self._uva_weight,
            flat_ids,
            output,
            self._row_bytes,
            self.shard_indices.org_vocab_start_index,
            self.shard_indices.org_vocab_end_index,
            BLOCK_D=self._block_d,
        )
    return output

_reduce_etp_embeddings(embeddings)

Combine pinned lookup results owned by different ETP ranks.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
def _reduce_etp_embeddings(self, embeddings: torch.Tensor) -> torch.Tensor:
    """Combine pinned lookup results owned by different ETP ranks."""
    if self.tp_size == 1:
        return embeddings
    assert self.parallel_group is not None
    if embeddings.dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
        # Each vocabulary row has one owner, so reduce the raw FP8 bytes.
        reduced = self.parallel_group.all_reduce(embeddings.view(torch.int8))
        return reduced.view(embeddings.dtype)
    return self.parallel_group.all_reduce(embeddings)

allocate_embedding_weight(num_embeddings, embedding_dim, dtype)

Allocate the complete PLE weight directly in pinned CPU memory.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
def allocate_embedding_weight(
    self,
    num_embeddings: int,
    embedding_dim: int,
    dtype: torch.dtype,
) -> torch.Tensor:
    """Allocate the complete PLE weight directly in pinned CPU memory."""
    return torch.empty(
        num_embeddings,
        embedding_dim,
        dtype=dtype,
        device="cpu",
        pin_memory=True,
    )

forward(hidden_states)

Finish the pinned lookup into graph-owned output storage.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
    """Finish the pinned lookup into graph-owned output storage."""
    buffer = self._prefetch_buffer
    if buffer is None:
        raise RuntimeError("pinned PLE lookup requires a prior start_prefetch")
    output = buffer.new_empty((hidden_states.shape[0], self._output_dim))
    self._finalize_prefetch(buffer, output)
    return output

start_prefetch(hidden_states, ngram_ids)

Gather ETP IDs and launch their UVA lookup on the side stream.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
@eager_break_during_capture
def start_prefetch(
    self,
    hidden_states: torch.Tensor,
    ngram_ids: torch.Tensor,
) -> None:
    """Gather ETP IDs and launch their UVA lookup on the side stream."""
    buffer = self._prefetch_buffer
    if buffer is None:
        # First use allocates. The eager profile run always precedes
        # cudagraph capture, so allocation never happens mid-capture;
        # the lock keeps concurrent first callers from tearing the
        # stream/buffer pair.
        with self._prefetch_alloc_lock:
            buffer = self._prefetch_buffer
            if buffer is None:
                if torch.cuda.is_current_stream_capturing():
                    raise RuntimeError(
                        "pinned PLE prefetch buffer must be allocated "
                        "eagerly, before cudagraph capture"
                    )
                self._prefetch_stream = torch.cuda.Stream(
                    device=self._uva_weight.device
                )
                buffer = torch.empty(
                    self._prefetch_rows,
                    self._num_ngram_heads,
                    self.embedding_dim,
                    dtype=self.weight.dtype,
                    device=self._uva_weight.device,
                )
                self._prefetch_buffer = buffer
    prefetch_stream = self._prefetch_stream
    if prefetch_stream is None:
        raise RuntimeError("pinned PLE prefetch stream was not allocated")
    slot_size, _ = self._get_dp_gather_slot(ngram_ids.shape[0])
    gathered_ids = self._gather_dp_ids(ngram_ids, slot_size)
    if gathered_ids.shape[0] > buffer.shape[0]:
        raise ValueError(
            f"pinned PLE prefetch buffer holds {buffer.shape[0]} rows, "
            f"but the batch needs {gathered_ids.shape[0]}"
        )
    active_output = buffer[: gathered_ids.shape[0]]
    prefetch_stream.wait_stream(torch.cuda.current_stream())
    gathered_ids.record_stream(prefetch_stream)
    with torch.cuda.stream(prefetch_stream):
        self._lookup(gathered_ids, output=active_output)

sync_lookup(ngram_ids)

Synchronous UVA lookup for platforms without prefetch wiring.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
def sync_lookup(self, ngram_ids: torch.Tensor) -> torch.Tensor:
    """Synchronous UVA lookup for platforms without prefetch wiring."""
    slot_size, slot_offset = self._get_dp_gather_slot(ngram_ids.shape[0])
    gathered_ids = self._gather_dp_ids(ngram_ids, slot_size)
    embeddings = self._lookup(gathered_ids)
    embeddings = self._reduce_etp_embeddings(embeddings)
    return self._select_embeddings(
        embeddings,
        ngram_ids.shape[0],
        slot_offset,
    )

Qwen4ExpPLEUnquantizedEmbeddingMethod

Bases: Qwen4ExpPLEEmbeddingMethod

Unquantized PLE embedding storage and lookup semantics.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
class Qwen4ExpPLEUnquantizedEmbeddingMethod(Qwen4ExpPLEEmbeddingMethod):
    """Unquantized PLE embedding storage and lookup semantics."""

    def create_weights(
        self,
        layer: Qwen4ExpPLEEmbedding,
        input_size_per_partition: int,
        output_partition_sizes: list[int],
        input_size: int,
        output_size: int,
        params_dtype: torch.dtype,
        **extra_weight_attrs,
    ) -> None:
        del input_size, output_size
        weight = nn.Parameter(
            layer.allocate_embedding_weight(
                sum(output_partition_sizes),
                input_size_per_partition,
                params_dtype,
            ),
            requires_grad=False,
        )
        set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0})
        set_weight_attrs(weight, extra_weight_attrs)
        layer.register_parameter("weight", weight)

    def dequantize(
        self,
        layer: nn.Module,
        embeddings: torch.Tensor,
        output_dtype: torch.dtype,
    ) -> torch.Tensor:
        del layer, output_dtype
        return embeddings

_lookup_ple_embedding_from_pinned_kernel(weight_ptr, ids_ptr, output_ptr, row_bytes, tp_vocab_start, tp_vocab_end, BLOCK_D)

Copy TP-owned PLE rows as raw bytes from a CUDA view of pinned host memory.

Byte pointers keep the storage dtype out of the kernel signature, so dtypes Triton cannot lower on the GPU (FP8 E4M3FN before SM89) still work.

Source code in vllm/models/qwen4_exp/common/ngram_embedding.py
@triton.jit
def _lookup_ple_embedding_from_pinned_kernel(
    weight_ptr: tl.pointer_type(tl.uint8),  # type: ignore[valid-type]
    ids_ptr,
    output_ptr: tl.pointer_type(tl.uint8),  # type: ignore[valid-type]
    row_bytes,
    tp_vocab_start,
    tp_vocab_end,
    BLOCK_D: tl.constexpr,
):
    """Copy TP-owned PLE rows as raw bytes from a CUDA view of pinned host memory.

    Byte pointers keep the storage dtype out of the kernel signature, so
    dtypes Triton cannot lower on the GPU (FP8 E4M3FN before SM89) still work.
    """
    row_id = tl.program_id(0).to(tl.int64)
    global_idx = tl.load(ids_ptr + row_id)
    in_range = (global_idx >= tp_vocab_start) & (global_idx < tp_vocab_end)
    local_idx = tl.where(in_range, global_idx - tp_vocab_start, 0)
    offsets = tl.arange(0, BLOCK_D)
    store_mask = offsets < row_bytes
    load_mask = store_mask & in_range
    values = tl.load(
        weight_ptr + local_idx * row_bytes + offsets,
        mask=load_mask,
        other=0,
    )
    tl.store(
        output_ptr + row_id * row_bytes + offsets,
        values,
        mask=store_mask,
    )