Add reproducible DGX Spark deployment for Qwen3.8 Flash Next NVFP4
This commit is contained in:
@@ -0,0 +1,49 @@
|
||||
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')
|
||||
Reference in New Issue
Block a user