Files
qwen38-flash-next-dgx-spark/patches/spark_ngram_adapter.py
T

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