import tempfile from types import SimpleNamespace from pathlib import Path import torch import vllm.config import vllm.distributed from safetensors.torch import save_file with tempfile.TemporaryDirectory() as folder: fake_config = SimpleNamespace(engram_config=None, model_config=SimpleNamespace(model=folder)) vllm.config.get_current_vllm_config = lambda: fake_config vllm.distributed.get_etp_group = lambda: SimpleNamespace(world_size=1) from vllm.models.qwen4_exp.nvidia.ngram_embedding import Qwen4ExpNGramEmbedding from vllm.model_executor.layers.quantization.fp8 import Fp8Config cfg = SimpleNamespace(ngram_size=3, heads_per_ngram=2, eos_token_id=0, vocab_size=100, split_ngram_parts=2, seed=1234, ngram_vocab_size_base=17, make_ngram_vocab_size_divisible_by=8) prefix = 'model.language_model.layers.0.ple.ple_embedding' with torch.device('cuda'): layer = Qwen4ExpNGramEmbedding(cfg, 8, 0, 32, data_parallel_rank=0, prefix=prefix, quant_config=Fp8Config(is_checkpoint_fp8_serialized=True)) count = layer.ngram_embedding.org_vocab_size full = (torch.arange(count * 2).reshape(count, 2) % 13).to(torch.float8_e4m3fn) tensors = {} for i, part in enumerate(full.chunk(2)): tensors[f'{prefix}.ngram_embedding.shard_{i}.weight'] = part.contiguous() tensors[f'{prefix}.ngram_embedding.weight_scale'] = torch.tensor(0.25) save_file(tensors, str(Path(folder) / 'ple.safetensors')) layer.load_weights((name.removeprefix(prefix + '.'), value) for name, value in tensors.items()) ids = torch.tensor([[0, 1, 1, count - 1], [5, 2, 9, 7]], device='cuda') out = layer.ngram_embedding(ids) expected = full.view(torch.uint8)[ids.cpu()].view(torch.float8_e4m3fn).cuda() assert torch.equal(out.view(torch.uint8), expected.view(torch.uint8)) dequant = layer.ngram_embedding.dequantize(out, torch.bfloat16) assert torch.equal(dequant, expected.to(torch.bfloat16) * 0.25) assert layer.ngram_embedding.weight.numel() == 0 tokens = torch.tensor([3, 7, 9], device='cuda') starts = torch.tensor([0, 3], dtype=torch.int32, device='cuda') context = torch.tensor([[0, 0]], device='cuda') hashed = layer.compute_ngram_ids(tokens, starts, context) actual = layer(None, tokens, starts, context) reference = full.view(torch.uint8)[hashed.cpu()].view(torch.float8_e4m3fn).cuda().flatten(-2) assert torch.equal(actual.view(torch.uint8), reference.view(torch.uint8)) print('PASS: disk rows, repeated IDs, boundary IDs, FP8 bytes and scaling; no resident PLE table') compiled = torch.compile(lambda t, s, c: layer(None, t, s, c), fullgraph=True) actual_compiled = compiled(tokens, starts, context) assert torch.equal(actual_compiled.view(torch.uint8), reference.view(torch.uint8)) print('COMPILE_PASS: full graph matches exact FP8 table lookup')