Skip to content

vllm.models.qwen4_exp.amd.ple_layer

GPU-resident Qwen4Exp position-learning enhancement layers.

Classes:

Qwen4ExpNGramEmbedding

Bases: Module

Methods:

  • load_weights –

    Load hash buffers and checkpoint-split embedding rows.

Source code in vllm/models/qwen4_exp/amd/ple_layer.py
class Qwen4ExpNGramEmbedding(nn.Module):
    _MASK64 = (1 << 64) - 1
    _SPLITMIX_GAMMA = 0x9E3779B97F4A7C15
    _SPLITMIX_M1 = 0xBF58476D1CE4E5B9
    _SPLITMIX_M2 = 0x94D049BB133111EB
    _PLE_LAYER_PRIME = 10007

    @classmethod
    def _splitmix64(cls, value: int) -> int:
        """Mix an integer into a deterministic unsigned 64-bit value."""
        value = (value + cls._SPLITMIX_GAMMA) & cls._MASK64
        value = ((value ^ (value >> 30)) * cls._SPLITMIX_M1) & cls._MASK64
        value = ((value ^ (value >> 27)) * cls._SPLITMIX_M2) & cls._MASK64
        return (value ^ (value >> 31)) & cls._MASK64

    @staticmethod
    def _is_prime_64(value: int) -> bool:
        """Return whether a 64-bit integer is prime."""
        if value < 2:
            return False
        for prime in (2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37):
            if value % prime == 0:
                return value == prime
        exponent = value - 1
        shifts = 0
        while exponent % 2 == 0:
            exponent //= 2
            shifts += 1
        for base in (2, 325, 9375, 28178, 450775, 9780504, 1795265022):
            if base % value == 0:
                continue
            witness = pow(base, exponent, value)
            if witness in (1, value - 1):
                continue
            for _ in range(shifts - 1):
                witness = pow(witness, 2, value)
                if witness == value - 1:
                    break
            else:
                return False
        return True

    @classmethod
    def _nth_prime_after(cls, start: int, count: int) -> int:
        """Return the ``count``-th prime strictly greater than ``start``."""
        prime = int(start)
        for _ in range(count):
            candidate = prime + 1
            if candidate <= 2:
                prime = 2
                continue
            if candidate % 2 == 0:
                candidate += 1
            while not cls._is_prime_64(candidate):
                candidate += 2
            prime = candidate
        return prime

    @classmethod
    def _make_layer_multipliers(
        cls,
        *,
        ngram_size: int,
        unigram_vocab_size: int,
        seed: int,
        ple_dense_layer_id: int,
    ) -> list[int]:
        """Build deterministic hash multipliers for one PLE layer."""
        max_multiplier = ((1 << 63) - 1) // unigram_vocab_size
        half_bound = max(1, max_multiplier // 2)
        base_seed = seed + cls._PLE_LAYER_PRIME * ple_dense_layer_id
        multipliers = []
        for index in range(ngram_size):
            value = base_seed + cls._SPLITMIX_GAMMA * (index + 1)
            multipliers.append(2 * (cls._splitmix64(value) % half_bound) + 1)
        return multipliers

    @classmethod
    def _make_vocab_layout(
        cls,
        *,
        ngram_vocab_size_base: int,
        ngram_heads: int,
        ple_dense_layer_id: int,
    ) -> tuple[list[int], list[int], int]:
        """Build per-head vocabulary sizes, offsets, and total row count."""
        sizes: list[int] = []
        offsets: list[int] = []
        offset = 0
        for local_head in range(ngram_heads):
            global_head = ple_dense_layer_id * ngram_heads + local_head
            size = cls._nth_prime_after(ngram_vocab_size_base - 1, global_head + 1)
            sizes.append(size)
            offsets.append(offset)
            offset += size
        return sizes, offsets, offset

    def __init__(
        self,
        config: Qwen4ExpTextConfig,
        embedding_dim: int,
        ple_dense_layer_id: int,
        max_total_tokens: int,
        prefix: str,
        layer_name: str,
        *,
        data_parallel_rank: int = 0,
        quant_config: QuantizationConfig | None = None,
        params_dtype: torch.dtype | None = None,
    ) -> None:
        super().__init__()
        self.embedding_dim = embedding_dim
        self.layer_name = layer_name
        self.ngram_size = int(config.ngram_size)
        self.heads_per_ngram = int(config.heads_per_ngram)
        self.ngram_heads = (self.ngram_size - 1) * self.heads_per_ngram
        if self.ngram_size < 2:
            raise ValueError(f"ngram_size must be >= 2, got {self.ngram_size}")
        if self.heads_per_ngram <= 0:
            raise ValueError(f"heads_per_ngram must be > 0, got {self.heads_per_ngram}")
        if embedding_dim % self.ngram_heads:
            raise ValueError(
                "ple_embed_dim must be divisible by total ngram heads: "
                f"{embedding_dim} % {self.ngram_heads} != 0"
            )
        self.head_dim = embedding_dim // self.ngram_heads
        self.eos_token_id = int(config.eos_token_id)
        self.unigram_vocab_size = int(config.vocab_size)
        self.split_ngram_parts = int(getattr(config, "split_ngram_parts", 512))
        if self.split_ngram_parts <= 0:
            raise ValueError("split_ngram_parts must be positive")

        multipliers = self._make_layer_multipliers(
            ngram_size=self.ngram_size,
            unigram_vocab_size=self.unigram_vocab_size,
            seed=int(getattr(config, "seed", 1234)),
            ple_dense_layer_id=ple_dense_layer_id,
        )
        self.register_buffer(
            "layer_multipliers",
            torch.tensor(multipliers, dtype=torch.long),
            persistent=True,
        )

        sizes, offsets, total_vocab_size = self._make_vocab_layout(
            ngram_vocab_size_base=int(config.ngram_vocab_size_base),
            ngram_heads=self.ngram_heads,
            ple_dense_layer_id=ple_dense_layer_id,
        )
        self.register_buffer(
            "ngram_heads_vocab_sizes",
            torch.tensor(sizes, dtype=torch.long),
            persistent=True,
        )
        self.register_buffer(
            "ngram_heads_offsets",
            torch.tensor(offsets, dtype=torch.long),
            persistent=True,
        )
        divisor = int(config.make_ngram_vocab_size_divisible_by)
        padded_vocab_size = ((total_vocab_size + divisor - 1) // divisor) * divisor
        if params_dtype is None:
            params_dtype = torch.get_default_dtype()
        embedding_prefix = f"{prefix}.ngram_embedding"
        embedding_quant_method = Qwen4ExpPLEEmbeddingMethod.from_quant_config(
            quant_config,
            embedding_prefix,
            getattr(config, "ple_embedding_dtype", None),
        )
        engram_config = get_current_vllm_config().engram_config
        embedding_cls = (
            Qwen4ExpPLEPinnedHostEmbedding
            if engram_config is not None and engram_config.cpu_offload
            else Qwen4ExpPLEDeviceEmbedding
        )
        self.ngram_embedding = embedding_cls(
            padded_vocab_size,
            self.head_dim,
            params_dtype=params_dtype,
            padding_size=divisor,
            prefix=embedding_prefix,
            embedding_method=embedding_quant_method,
            num_ngram_heads=self.ngram_heads,
            max_total_tokens=max_total_tokens,
            data_parallel_rank=data_parallel_rank,
        )
        logger.info(
            "Initialized AMD PLE embedding %s: quantization_method=%s, "
            "weight_dtype=%s, weight_device=%s, pinned=%s",
            embedding_prefix,
            type(embedding_quant_method).__name__,
            self.ngram_embedding.weight.dtype,
            self.ngram_embedding.weight.device,
            self.ngram_embedding.weight.is_pinned(),
        )

    def forward(
        self,
        input_ids: torch.Tensor,
        query_start_loc: torch.Tensor,
        ngram_context: torch.Tensor,
    ) -> torch.Tensor:
        input_ids = input_ids.reshape(-1)
        ngram_ids = input_ids.new_empty(
            (input_ids.shape[0], self.ngram_heads), dtype=torch.int64
        )
        torch.ops.vllm.qwen4_exp_amd_ple_ngram_ids(
            input_ids,
            query_start_loc,
            ngram_context,
            ngram_ids,
            self.layer_name,
        )
        embedding = self.ngram_embedding
        if embedding.supports_prefetch:
            output = ngram_ids.new_empty(
                (ngram_ids.shape[0], self.embedding_dim),
                dtype=embedding.weight.dtype,
            )
            torch.ops.vllm.qwen4_exp_amd_ple_ngram_embedding_pinned(
                ngram_ids,
                output,
                self.layer_name,
            )
            return output
        output = ngram_ids.new_empty(
            (ngram_ids.shape[0], self.embedding_dim),
            dtype=embedding.params_dtype,
        )
        torch.ops.vllm.qwen4_exp_amd_ple_ngram_embedding(
            ngram_ids,
            output,
            self.layer_name,
        )
        return output

    def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
        """Load hash buffers and checkpoint-split embedding rows."""
        persistent_buffers = {
            "layer_multipliers": self.layer_multipliers,
            "ngram_heads_offsets": self.ngram_heads_offsets,
            "ngram_heads_vocab_sizes": self.ngram_heads_vocab_sizes,
        }
        loaded: set[str] = set()
        regular_weights: list[tuple[str, torch.Tensor]] = []
        shard_prefix = "ngram_embedding.shard_"

        for name, loaded_weight in weights:
            leaf_name = name.rsplit(".", 1)[-1]
            if leaf_name.startswith("hashstats_") or leaf_name == "token_lookup":
                continue
            if name in persistent_buffers:
                buffer = persistent_buffers[name]
                if buffer.shape != loaded_weight.shape:
                    raise ValueError(
                        f"Shape mismatch for {name}: expected "
                        f"{tuple(buffer.shape)}, got {tuple(loaded_weight.shape)}"
                    )
                buffer.copy_(loaded_weight.to(device=buffer.device, dtype=buffer.dtype))
                loaded.add(name)
                continue
            if name.startswith(shard_prefix) and name.endswith(".weight"):
                shard_text = name[len(shard_prefix) : -len(".weight")]
                if not shard_text.isdigit():
                    regular_weights.append((name, loaded_weight))
                    continue
                shard_index = int(shard_text)
                if shard_index >= self.split_ngram_parts:
                    raise ValueError(
                        f"PLE embedding shard index {shard_index} exceeds "
                        f"split_ngram_parts={self.split_ngram_parts}"
                    )
                embedding = self.ngram_embedding
                shard_size = (
                    embedding.org_vocab_size + self.split_ngram_parts - 1
                ) // self.split_ngram_parts
                checkpoint_start = shard_index * shard_size
                expected_rows = max(
                    0,
                    min(shard_size, embedding.org_vocab_size - checkpoint_start),
                )
                expected_shape = (expected_rows, embedding.embedding_dim)
                if tuple(loaded_weight.shape) != expected_shape:
                    raise ValueError(
                        f"Shape mismatch for PLE embedding shard {shard_index}: "
                        f"expected {expected_shape}, got "
                        f"{tuple(loaded_weight.shape)}"
                    )
                embedding.weight.weight_loader(
                    embedding.weight,
                    loaded_weight,
                    checkpoint_start=checkpoint_start,
                )
                loaded.add("ngram_embedding.weight")
                continue
            regular_weights.append((name, loaded_weight))

        if regular_weights:
            loaded.update(AutoWeightsLoader(self).load_weights(regular_weights))
        return loaded

_is_prime_64(value) staticmethod

Return whether a 64-bit integer is prime.

Source code in vllm/models/qwen4_exp/amd/ple_layer.py
@staticmethod
def _is_prime_64(value: int) -> bool:
    """Return whether a 64-bit integer is prime."""
    if value < 2:
        return False
    for prime in (2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37):
        if value % prime == 0:
            return value == prime
    exponent = value - 1
    shifts = 0
    while exponent % 2 == 0:
        exponent //= 2
        shifts += 1
    for base in (2, 325, 9375, 28178, 450775, 9780504, 1795265022):
        if base % value == 0:
            continue
        witness = pow(base, exponent, value)
        if witness in (1, value - 1):
            continue
        for _ in range(shifts - 1):
            witness = pow(witness, 2, value)
            if witness == value - 1:
                break
        else:
            return False
    return True

_make_layer_multipliers(*, ngram_size, unigram_vocab_size, seed, ple_dense_layer_id) classmethod

Build deterministic hash multipliers for one PLE layer.

Source code in vllm/models/qwen4_exp/amd/ple_layer.py
@classmethod
def _make_layer_multipliers(
    cls,
    *,
    ngram_size: int,
    unigram_vocab_size: int,
    seed: int,
    ple_dense_layer_id: int,
) -> list[int]:
    """Build deterministic hash multipliers for one PLE layer."""
    max_multiplier = ((1 << 63) - 1) // unigram_vocab_size
    half_bound = max(1, max_multiplier // 2)
    base_seed = seed + cls._PLE_LAYER_PRIME * ple_dense_layer_id
    multipliers = []
    for index in range(ngram_size):
        value = base_seed + cls._SPLITMIX_GAMMA * (index + 1)
        multipliers.append(2 * (cls._splitmix64(value) % half_bound) + 1)
    return multipliers

_make_vocab_layout(*, ngram_vocab_size_base, ngram_heads, ple_dense_layer_id) classmethod

Build per-head vocabulary sizes, offsets, and total row count.

Source code in vllm/models/qwen4_exp/amd/ple_layer.py
@classmethod
def _make_vocab_layout(
    cls,
    *,
    ngram_vocab_size_base: int,
    ngram_heads: int,
    ple_dense_layer_id: int,
) -> tuple[list[int], list[int], int]:
    """Build per-head vocabulary sizes, offsets, and total row count."""
    sizes: list[int] = []
    offsets: list[int] = []
    offset = 0
    for local_head in range(ngram_heads):
        global_head = ple_dense_layer_id * ngram_heads + local_head
        size = cls._nth_prime_after(ngram_vocab_size_base - 1, global_head + 1)
        sizes.append(size)
        offsets.append(offset)
        offset += size
    return sizes, offsets, offset

_nth_prime_after(start, count) classmethod

Return the count-th prime strictly greater than start.

Source code in vllm/models/qwen4_exp/amd/ple_layer.py
@classmethod
def _nth_prime_after(cls, start: int, count: int) -> int:
    """Return the ``count``-th prime strictly greater than ``start``."""
    prime = int(start)
    for _ in range(count):
        candidate = prime + 1
        if candidate <= 2:
            prime = 2
            continue
        if candidate % 2 == 0:
            candidate += 1
        while not cls._is_prime_64(candidate):
            candidate += 2
        prime = candidate
    return prime

_splitmix64(value) classmethod

Mix an integer into a deterministic unsigned 64-bit value.

Source code in vllm/models/qwen4_exp/amd/ple_layer.py
@classmethod
def _splitmix64(cls, value: int) -> int:
    """Mix an integer into a deterministic unsigned 64-bit value."""
    value = (value + cls._SPLITMIX_GAMMA) & cls._MASK64
    value = ((value ^ (value >> 30)) * cls._SPLITMIX_M1) & cls._MASK64
    value = ((value ^ (value >> 27)) * cls._SPLITMIX_M2) & cls._MASK64
    return (value ^ (value >> 31)) & cls._MASK64

load_weights(weights)

Load hash buffers and checkpoint-split embedding rows.

Source code in vllm/models/qwen4_exp/amd/ple_layer.py
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
    """Load hash buffers and checkpoint-split embedding rows."""
    persistent_buffers = {
        "layer_multipliers": self.layer_multipliers,
        "ngram_heads_offsets": self.ngram_heads_offsets,
        "ngram_heads_vocab_sizes": self.ngram_heads_vocab_sizes,
    }
    loaded: set[str] = set()
    regular_weights: list[tuple[str, torch.Tensor]] = []
    shard_prefix = "ngram_embedding.shard_"

    for name, loaded_weight in weights:
        leaf_name = name.rsplit(".", 1)[-1]
        if leaf_name.startswith("hashstats_") or leaf_name == "token_lookup":
            continue
        if name in persistent_buffers:
            buffer = persistent_buffers[name]
            if buffer.shape != loaded_weight.shape:
                raise ValueError(
                    f"Shape mismatch for {name}: expected "
                    f"{tuple(buffer.shape)}, got {tuple(loaded_weight.shape)}"
                )
            buffer.copy_(loaded_weight.to(device=buffer.device, dtype=buffer.dtype))
            loaded.add(name)
            continue
        if name.startswith(shard_prefix) and name.endswith(".weight"):
            shard_text = name[len(shard_prefix) : -len(".weight")]
            if not shard_text.isdigit():
                regular_weights.append((name, loaded_weight))
                continue
            shard_index = int(shard_text)
            if shard_index >= self.split_ngram_parts:
                raise ValueError(
                    f"PLE embedding shard index {shard_index} exceeds "
                    f"split_ngram_parts={self.split_ngram_parts}"
                )
            embedding = self.ngram_embedding
            shard_size = (
                embedding.org_vocab_size + self.split_ngram_parts - 1
            ) // self.split_ngram_parts
            checkpoint_start = shard_index * shard_size
            expected_rows = max(
                0,
                min(shard_size, embedding.org_vocab_size - checkpoint_start),
            )
            expected_shape = (expected_rows, embedding.embedding_dim)
            if tuple(loaded_weight.shape) != expected_shape:
                raise ValueError(
                    f"Shape mismatch for PLE embedding shard {shard_index}: "
                    f"expected {expected_shape}, got "
                    f"{tuple(loaded_weight.shape)}"
                )
            embedding.weight.weight_loader(
                embedding.weight,
                loaded_weight,
                checkpoint_start=checkpoint_start,
            )
            loaded.add("ngram_embedding.weight")
            continue
        regular_weights.append((name, loaded_weight))

    if regular_weights:
        loaded.update(AutoWeightsLoader(self).load_weights(regular_weights))
    return loaded

qwen4_exp_amd_ple_ngram_embedding(ngram_ids, output, layer_name)

Run the large PLE embedding lookup outside Inductor's FX graph.

Keeping the embedding weight in static_forward_context prevents AOT compile-time autotuning from materializing a synthetic copy of the weight.

Source code in vllm/models/qwen4_exp/amd/ple_layer.py
def qwen4_exp_amd_ple_ngram_embedding(
    ngram_ids: torch.Tensor,
    output: torch.Tensor,
    layer_name: str,
) -> None:
    """Run the large PLE embedding lookup outside Inductor's FX graph.

    Keeping the embedding weight in ``static_forward_context`` prevents AOT
    compile-time autotuning from materializing a synthetic copy of the weight.
    """
    layer = get_forward_context().no_compile_layers[layer_name]
    if not isinstance(layer, Qwen4ExpPLELayer):
        raise TypeError(f"{layer_name} is not a Qwen4Exp PLE owner")
    result = layer.ple_embedding.ngram_embedding(ngram_ids).flatten(-2)
    output.copy_(result)

qwen4_exp_amd_ple_ngram_embedding_pinned(ngram_ids, output, layer_name)

Run the pinned PLE UVA lookup outside Inductor's FX graph.

Same rationale as the device-path escape: keeping the large embedding weight out of the graph prevents AOT compile-time autotuning from materializing a synthetic copy of the weight.

Source code in vllm/models/qwen4_exp/amd/ple_layer.py
def qwen4_exp_amd_ple_ngram_embedding_pinned(
    ngram_ids: torch.Tensor,
    output: torch.Tensor,
    layer_name: str,
) -> None:
    """Run the pinned PLE UVA lookup outside Inductor's FX graph.

    Same rationale as the device-path escape: keeping the large embedding weight
    out of the graph prevents AOT compile-time autotuning from materializing a
    synthetic copy of the weight.
    """
    layer = get_forward_context().no_compile_layers[layer_name]
    if not isinstance(layer, Qwen4ExpPLELayer):
        raise TypeError(f"{layer_name} is not a Qwen4Exp PLE owner")
    result = layer.ple_embedding.ngram_embedding.sync_lookup(ngram_ids).flatten(-2)
    output.copy_(result)

qwen4_exp_amd_ple_ngram_ids(input_ids, query_start_loc, ngram_context, output, layer_name)

Hash the current request layout into n-gram embedding indices.

The launch sizes its request binary search from the request count, which is symbolic under Dynamo, so it runs outside Inductor's FX graph.

Source code in vllm/models/qwen4_exp/amd/ple_layer.py
def qwen4_exp_amd_ple_ngram_ids(
    input_ids: torch.Tensor,
    query_start_loc: torch.Tensor,
    ngram_context: torch.Tensor,
    output: torch.Tensor,
    layer_name: str,
) -> None:
    """Hash the current request layout into n-gram embedding indices.

    The launch sizes its request binary search from the request count, which
    is symbolic under Dynamo, so it runs outside Inductor's FX graph.
    """
    layer = get_forward_context().no_compile_layers[layer_name]
    if not isinstance(layer, Qwen4ExpPLELayer):
        raise TypeError(f"{layer_name} is not a Qwen4Exp PLE owner")
    embedding = layer.ple_embedding
    ple_ngram_ids(
        input_ids=input_ids,
        query_start_loc=query_start_loc,
        ngram_context=ngram_context,
        layer_multipliers=embedding.layer_multipliers,
        ngram_heads_vocab_sizes=embedding.ngram_heads_vocab_sizes,
        ngram_heads_offsets=embedding.ngram_heads_offsets,
        eos_token_id=embedding.eos_token_id,
        heads_per_ngram=embedding.heads_per_ngram,
        output=output,
    )