"""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 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): ids = self.compute_ngram_ids(input_ids, query_start_loc, ngram_context) output = torch.empty((ids.shape[0], self.embedding_dim), dtype=torch.float8_e4m3fn, device=ids.device) torch.ops.vllm.ple_mmap_lookup_ids(ids, output, self._ple_mmap_prefix) return output _register_op() cls.__init__, cls.load_weights, cls.forward = init, load, forward