70 lines
2.8 KiB
Python
70 lines
2.8 KiB
Python
"""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
|