"""PLE mmap adapter for pinned nightly 0bfc7a15, single GPU only. Uses the reviewed blazux mmap reader, retaining upstream hashing and dequantization. """ import torch from vllm_ple_mmap import _MmapNgramEmbedding, _setup_table_v029, _REGISTRY, _register_op def _lookup_with_hash(input_ids: torch.Tensor, query_start_loc: torch.Tensor, ngram_context: torch.Tensor, output: torch.Tensor, layer_name: str) -> None: # Capture records preceding kernels without executing them. CPU lookup must # therefore compute its hash IDs in this eager segment as well. from vllm_ple_mmap import _lookup_ids_impl layer = _REGISTRY[layer_name] ids = layer.compute_ngram_ids(input_ids, query_start_loc, ngram_context) _lookup_ids_impl(ids, output, layer_name) def _lookup_with_hash_fake(input_ids: torch.Tensor, query_start_loc: torch.Tensor, ngram_context: torch.Tensor, output: torch.Tensor, layer_name: str) -> None: pass class DiskEmbedding(_MmapNgramEmbedding): supports_prefetch = False def __init__(self, n, d, **kwargs): super().__init__(n, d) self.register_buffer("weight", torch.empty(0, dtype=torch.float8_e4m3fn), persistent=False) def dequantize(self, embeddings, output_dtype): if self.table is None: raise RuntimeError("PLE disk table not loaded") return embeddings.to(output_dtype) * self.weight_scale.to(output_dtype) def start_prefetch(self, *args): pass def apply(cls): import sys from vllm.config import get_current_vllm_config from vllm.distributed import get_etp_group mod = sys.modules[cls.__module__] original_init, original_load = cls.__init__, cls.load_weights def init(self, *args, **kwargs): if get_etp_group().world_size != 1: raise RuntimeError("Spark disk adapter supports only ETP=1") device_cls = mod.Qwen4ExpPLEDeviceEmbedding host_cls = mod.Qwen4ExpPLEPinnedHostEmbedding mod.Qwen4ExpPLEDeviceEmbedding = mod.Qwen4ExpPLEPinnedHostEmbedding = DiskEmbedding try: original_init(self, *args, **kwargs) finally: mod.Qwen4ExpPLEDeviceEmbedding, mod.Qwen4ExpPLEPinnedHostEmbedding = device_cls, host_cls self._ple_mmap_prefix = kwargs["prefix"] self._ple_mmap_model_path = get_current_vllm_config().model_config.model _REGISTRY[self._ple_mmap_prefix] = self def load(self, weights): loaded = set() def filtered(): for name, tensor in weights: if name.startswith("ngram_embedding.shard_") and name.endswith(".weight"): continue if name == "ngram_embedding.weight_scale": self.register_buffer("_offload_weight_scale", tensor.detach().to("cuda"), persistent=False) continue yield name, tensor loaded.update(original_load(self, filtered())) _setup_table_v029(self) self.ngram_embedding.weight_scale = self._offload_weight_scale return loaded def forward(self, hidden_states, input_ids, query_start_loc, ngram_context): output = torch.empty((input_ids.shape[0], self.embedding_dim), dtype=torch.float8_e4m3fn, device=input_ids.device) torch.ops.vllm.spark_ple_lookup(input_ids, query_start_loc, ngram_context, output, self._ple_mmap_prefix) return output from vllm.compilation.breakable_cudagraph import eager_break_during_capture from vllm.utils.torch_utils import direct_register_custom_op _register_op() if not hasattr(torch.ops.vllm, "spark_ple_lookup"): direct_register_custom_op( op_name="spark_ple_lookup", op_func=eager_break_during_capture(_lookup_with_hash), mutates_args=["output"], fake_impl=_lookup_with_hash_fake, ) cls.__init__, cls.load_weights, cls.forward = init, load, forward