Files
qwen38-flash-next-dgx-spark/experiments/graph/test_disk_adapter.py
T

80 lines
4.3 KiB
Python

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")