import os os.environ["VLLM_USE_BREAKABLE_CUDAGRAPH"] = "1" 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') # Runtime graph capture must break around the CPU mmap operation, and replay # must consume changed inputs instead of reusing capture-time row values. from vllm.compilation.breakable_cudagraph import BreakableCUDAGraphCapture original_hash = layer.compute_ngram_ids def checked_hash(*args): assert not torch.cuda.is_current_stream_capturing(), "CPU-dependent hash captured" return original_hash(*args) layer.compute_ngram_ids = checked_hash stream = torch.cuda.Stream() stream.wait_stream(torch.cuda.current_stream()) with torch.cuda.stream(stream): layer(None, tokens, starts, context) torch.cuda.synchronize() capture = BreakableCUDAGraphCapture() with capture: captured = layer(None, tokens, starts, context) downstream = captured.to(torch.float32) * 0.25 assert capture.num_eager_breaks >= 1 for values in ([4, 6, 8], [9, 2, 5], [3, 7, 9]): tokens.copy_(torch.tensor(values, device="cuda")) capture.replay() expected_ids = layer.compute_ngram_ids(tokens, starts, context) expected_bytes = full.view(torch.uint8)[expected_ids.cpu()].cuda().flatten(-2) assert torch.equal(captured.view(torch.uint8), expected_bytes) expected_fp8 = expected_bytes.view(torch.float8_e4m3fn) assert torch.equal(downstream, expected_fp8.to(torch.float32) * 0.25) print("BREAKABLE_GRAPH_PASS: CPU lookup excluded, changed-input replay byte-exact")