diff --git a/Dockerfile b/Dockerfile index f314852..c7fed2d 100644 --- a/Dockerfile +++ b/Dockerfile @@ -260,7 +260,7 @@ RUN set -eu; cd /opt/patches; \ patch_gdn_wmma patch_gdn_aiter_prefill patch_preshuffle install_radiance_hooks \ patch_unpad patch_mtp_mm_mask patch_mtp_loopbreak patch_qwen3_toolparse patch_from_json_filter \ patch_dynamo_metrics patch_conv1d_blockn patch_r4d patch_dflash_base \ - patch_dflash_fused_kv_fp8 patch_dflash_logits_cache_stride patch_dflash_w4 \ + patch_dflash_fused_kv_fp8 patch_dflash_logits_cache_stride patch_dflash_sampling_rng patch_dflash_w4 \ patch_dflash_selector_topk patch_gdn_metadata patch_gdn_shared_build \ patch_topk_triton_rows patch_topk_composite patch_rocm_cudagraph_current_stream \ patch_quark_mxfp4 patch_quark_bf16_mtp patch_ar_maxbytes patch_ar_geometry \ diff --git a/Dockerfile.dev b/Dockerfile.dev index a7ffd55..53eb804 100644 --- a/Dockerfile.dev +++ b/Dockerfile.dev @@ -76,7 +76,7 @@ RUN set -eu; cd /opt/patches; \ patch_gdn_wmma patch_gdn_aiter_prefill patch_preshuffle install_radiance_hooks \ patch_unpad patch_mtp_mm_mask patch_mtp_loopbreak patch_qwen3_toolparse patch_from_json_filter \ patch_dynamo_metrics patch_conv1d_blockn patch_r4d patch_dflash_base \ - patch_dflash_fused_kv_fp8 patch_dflash_logits_cache_stride patch_dflash_w4 \ + patch_dflash_fused_kv_fp8 patch_dflash_logits_cache_stride patch_dflash_sampling_rng patch_dflash_w4 \ patch_dflash_selector_topk patch_gdn_metadata patch_gdn_shared_build \ patch_topk_triton_rows patch_topk_composite patch_rocm_cudagraph_current_stream \ patch_quark_mxfp4 patch_quark_bf16_mtp patch_ar_maxbytes patch_ar_geometry \ diff --git a/Dockerfile.patch b/Dockerfile.patch index 97f4c40..d6a5147 100644 --- a/Dockerfile.patch +++ b/Dockerfile.patch @@ -38,7 +38,7 @@ RUN set -eu; cd /opt/patches; \ patch_gdn_wmma patch_gdn_aiter_prefill patch_preshuffle install_radiance_hooks \ patch_unpad patch_mtp_mm_mask patch_mtp_loopbreak patch_qwen3_toolparse patch_from_json_filter \ patch_dynamo_metrics patch_conv1d_blockn patch_r4d patch_dflash_base \ - patch_dflash_fused_kv_fp8 patch_dflash_logits_cache_stride patch_dflash_w4 \ + patch_dflash_fused_kv_fp8 patch_dflash_logits_cache_stride patch_dflash_sampling_rng patch_dflash_w4 \ patch_dflash_selector_topk patch_gdn_metadata patch_gdn_shared_build \ patch_topk_triton_rows patch_topk_composite patch_rocm_cudagraph_current_stream \ patch_quark_mxfp4 patch_quark_bf16_mtp patch_ar_maxbytes patch_ar_geometry \ diff --git a/README.md b/README.md index 8826ebe..ae92dfd 100644 --- a/README.md +++ b/README.md @@ -343,6 +343,7 @@ and an earlier mismatched combination caused sustained TP hangs. ## Documentation +- [DFlash sampling and GDN prefill numerical corrections](docs/NUMERICAL_CORRECTIONS.md) - [Upgrade and reproducibility history](https://gitlab.sayou.io/lance-wright/vllm-radiance/-/blob/main/docs/UPGRADE_PROGRESS.md) - [Stable vLLM v0.28 upgrade and qualification](https://gitlab.sayou.io/lance-wright/vllm-radiance/-/blob/main/docs/V028_UPGRADE.md) - [Radiance 0.9.3 / libr4d 0.5.0 qualification](https://gitlab.sayou.io/lance-wright/vllm-radiance/-/blob/main/docs/RADIANCE_093_R4D050_MXFP4.md) diff --git a/benchmarks/bin/check_dflash_sampling_rng.py b/benchmarks/bin/check_dflash_sampling_rng.py new file mode 100644 index 0000000..54f5ed4 --- /dev/null +++ b/benchmarks/bin/check_dflash_sampling_rng.py @@ -0,0 +1,107 @@ +#!/usr/bin/env python3 +"""Exercise the installed DFlash selector/rejection kernels with fixed probabilities. + +The control subtracts the proposal salt to reproduce the old shared-noise bug. +No checkpoint or conversation is needed. Run after patch_dflash_sampling_rng.py. +""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import time +import sysconfig +from pathlib import Path + + +def run(output: Path, samples: int): + import torch + from vllm.v1.worker.gpu.sample.gumbel import gumbel_sample + from vllm.v1.worker.gpu.spec_decode.dflash2.speculator import _selector_walk_kernel + from vllm.v1.worker.gpu.spec_decode.rejection_sampler_utils import rejection_sample + + device = "cuda" + target = torch.tensor([0.1, 0.5, 0.4], device=device) + draft = torch.tensor([0.5, 0.3, 0.2], device=device) + n = samples + scores = draft.log().view(1, 1, 1, 3).expand(n, 1, 3, 3).contiguous() + candidate = torch.arange(3, device=device, dtype=torch.int64).view(1, 1, 3).expand(n, 1, 3).contiguous() + state = torch.arange(n, device=device, dtype=torch.int32) + temperature = torch.ones(n, device=device, dtype=torch.float32) + seeds = torch.arange(n, device=device, dtype=torch.int64) + 1234567 + logits = target.log().view(1, 3).expand(2 * n, 3).contiguous() + expanded_state = state.repeat_interleave(2) + local_pos = torch.tensor([0, 1], device=device, dtype=torch.int32).repeat(n) + cu_logits = torch.arange(0, 2 * n + 1, 2, device=device, dtype=torch.int32) + tokens = torch.empty((n, 1), device=device, dtype=torch.int64) + realized = torch.empty((n, 1, 3), device=device, dtype=torch.float32) + rows = [] + native_rows = [] + started = time.monotonic() + for position in (37, 86142, 132739): + sample_pos = torch.full((n,), position + 1, device=device, dtype=torch.int64) + positions = torch.tensor([position, position + 1], device=device, dtype=torch.int64).repeat(n) + native = gumbel_sample(logits[:n], state, temperature, seeds, + sample_pos - 1, apply_temperature=True) + native_observed = torch.bincount(native.long(), minlength=3).float() / n + native_rows.append({"position": position, "observed": native_observed.tolist(), + "max_error": float((native_observed - target).abs().max())}) + for independent in (False, True): + proposal_positions = sample_pos if independent else sample_pos - (1 << 30) + _selector_walk_kernel[(n,)]( + scores, candidate, proposal_positions, state, temperature, + seeds, tokens, realized, + num_steps=1, top_k=3, BLOCK_K=4, + SAMPLE_PROBABILISTIC=True, USE_FP64=False, num_warps=1, + ) + draft_inputs = torch.zeros((n, 2), device=device, dtype=torch.int64) + draft_inputs[:, 1] = tokens[:, 0] + sampled, counts = rejection_sample( + logits, realized, draft_inputs.flatten(), cu_logits, + positions, state, expanded_state, local_pos, temperature, + seeds, 1, use_fp64=False, + ) + observed = torch.bincount(sampled[:, 0].long(), minlength=3).float() / n + proposal = torch.bincount(tokens[:, 0], minlength=3).float() / n + row = {"position": position, "independent_proposal_noise": independent, + "samples": n, "target": target.tolist(), "draft": draft.tolist(), + "observed": observed.tolist(), "proposal_observed": proposal.tolist(), + "max_error": float((observed - target).abs().max()), + "accepted_fraction": float((counts > 1).float().mean())} + rows.append(row) + print(json.dumps(row), flush=True) + # Greedy decoding is unaffected by the proposal RNG stream. + temperature.zero_() + greedy = [] + for independent in (False, True): + _selector_walk_kernel[(n,)]( + scores, candidate, sample_pos if independent else sample_pos - (1 << 30), state, temperature, + seeds, tokens, realized, + num_steps=1, top_k=3, BLOCK_K=4, + SAMPLE_PROBABILISTIC=True, USE_FP64=False, num_warps=1, + ) + draft_inputs[:, 1] = tokens[:, 0] + sampled, _ = rejection_sample(logits, realized, draft_inputs.flatten(), cu_logits, + positions, state, expanded_state, local_pos, temperature, seeds, 1, use_fp64=False) + greedy.append(bool((sampled[:, 0] == 1).all())) + package = Path(sysconfig.get_paths()["purelib"]) + sources = ["vllm/v1/worker/gpu/spec_decode/dflash2/speculator.py", + "vllm/v1/worker/gpu/spec_decode/rejection_sampler_utils.py", + "vllm/v1/worker/gpu/sample/gumbel.py"] + value = {"rows": rows, "native_target_rows": native_rows, + "greedy_exact": all(greedy), "seconds": time.monotonic() - started, + "source_hashes": {name: hashlib.sha256((package / name).read_bytes()).hexdigest() for name in sources}} + output.write_text(json.dumps(value, indent=2) + "\n") + assert all(row["max_error"] < 0.004 for row in rows if row["independent_proposal_noise"]) + assert all(row["max_error"] > 0.01 for row in rows if not row["independent_proposal_noise"]) + assert all(greedy) + assert all(row["max_error"] < 0.004 for row in native_rows) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--output", type=Path, required=True) + parser.add_argument("--samples", type=int, default=200000) + args = parser.parse_args() + run(args.output, args.samples) diff --git a/benchmarks/bin/check_gdn_extreme_decay.py b/benchmarks/bin/check_gdn_extreme_decay.py new file mode 100644 index 0000000..7bacde0 --- /dev/null +++ b/benchmarks/bin/check_gdn_extreme_decay.py @@ -0,0 +1,156 @@ +#!/usr/bin/env python3 +"""Check native GDN prefill against an independent FP64 sequential recurrence. + +Random tensors only: no checkpoint or conversation is needed. The optional +--library selects a separately compiled scan translation unit for qualification. +--expect-failure verifies that an uncorrected library reproduces the defect. +""" + +from __future__ import annotations + +import argparse +import ctypes +import hashlib +import json +import math +import time +from pathlib import Path + + +def probe(output_path: Path, library_path: Path | None, heads: int): + import torch + import radiance_gdn as native + import r4d + + torch.set_grad_enabled(False) + torch.manual_seed(14929) + device = "cuda" + query_heads, width = heads // 3, 128 + assert heads in (24, 48) + scale = width ** -0.5 + observed = hashlib.sha256(Path(r4d.__file__).read_bytes()).hexdigest() + assert native.ENABLED and native.CHUNK == 64 + library_sha256 = None + if library_path is not None: + library_sha256 = hashlib.sha256(library_path.read_bytes()).hexdigest() + library = ctypes.CDLL(str(library_path.resolve())) + scan = library.r4d_gdn_chunk_scan_k128_v128_c64_bf16 + scan.argtypes = ([ctypes.c_void_p] * 10 + [ctypes.c_int] * 6 + + [ctypes.c_float, ctypes.c_void_p]) + scan.restype = ctypes.c_int + native._CHUNK_SCAN = scan + started = time.monotonic() + rows = [] + cases = [(65, 0.02, 1, [65], False), (128, 1.0, 1, [128], False), + (128, 2.5, 1, [128], False), (128, 3.2, 1, [128], False), + (128, 8.0, 1, [128], False), (257, 0.02, 1, [257], False), + (257, 3.2, 1, [257], False), (1024, 0.02, 1, [1024], False), + (128, 32.0, 1, [128], False), (128, 2.5, 10000, [128], False), + (195, 0.02, 1, [65, 130], False), (195, 3.2, 1, [65, 130], True)] + for tokens, decay, amplitude, lengths, mixed_heads in cases: + q = torch.randn(tokens, query_heads, width, device=device) + k = torch.randn_like(q) + q = torch.nn.functional.normalize(q, dim=-1).to(torch.bfloat16) + k = torch.nn.functional.normalize(k, dim=-1).to(torch.bfloat16) + v = torch.randn(tokens, heads, width, device=device, dtype=torch.bfloat16) + v *= amplitude + steps = torch.full((tokens, heads), -decay, device=device, dtype=torch.float32) + if mixed_heads: + steps[:, 1:-1] = -0.02 + beta = torch.full_like(steps, 0.5) + boundaries = [0] + for length in lengths: + boundaries.append(boundaries[-1] + length) + assert boundaries[-1] == tokens + cumulative = torch.cat([ + steps[i:min(i + native.CHUNK, end)].cumsum(0) + for begin, end in zip(boundaries[:-1], boundaries[1:], strict=True) + for i in range(begin, end, native.CHUNK) + ], dim=0).contiguous() + cu = torch.tensor(boundaries, device=device, dtype=torch.int32) + initial = torch.randn(len(lengths), heads, width, width, device=device) * 0.01 + matrix = native.kkt_solve(k, beta, cumulative, cu, len(lengths), tokens, heads, query_heads) + # Guard the caller-owned output, including non-full final chunks. + count = v.numel() + slab = torch.full((count + 512,), 37.0, device=device, dtype=torch.bfloat16) + output = slab[256:256 + count].view(1, tokens, heads, width) + actual, final = native.fused_prefill( + q.unsqueeze(0), k.unsqueeze(0), v.unsqueeze(0), matrix.unsqueeze(0), + cumulative.unsqueeze(0), beta.unsqueeze(0), scale, initial, True, cu, + None, out=output, + ) + torch.cuda.synchronize() + # No WY factorization, chunk scan, kernel output or native intermediate + # participates in this reference. It follows the defining recurrence. + qr = q.double().repeat_interleave(heads // query_heads, dim=1) + kr = k.double().repeat_interleave(heads // query_heads, dim=1) + vr, br, gr = v.double(), beta.double(), steps.double().exp() + oracle = torch.empty_like(vr) + oracle_final = torch.empty_like(initial, dtype=torch.float64) + for sequence, (begin, end) in enumerate(zip(boundaries[:-1], boundaries[1:], strict=True)): + state = initial[sequence].double().clone() + for position in range(begin, end): + state *= gr[position, :, None, None] + residual = vr[position] - (state * kr[position, :, None, :]).sum(-1) + state += (br[position, :, None] * residual)[:, :, None] * kr[position, :, None, :] + oracle[position] = (state * qr[position, :, None, :]).sum(-1) * scale + oracle_final[sequence] = state + out_error = ((actual[0].double() - oracle).norm() / oracle.norm()).item() + state_error = ((final.double() - oracle_final).norm() / oracle_final.norm()).item() + finite = bool(torch.isfinite(actual).all() and torch.isfinite(final).all()) + guards = bool((slab[:256] == 37).all() and (slab[-256:] == 37).all()) + row = {"tokens": tokens, "heads": heads, "query_heads": query_heads, + "sequence_lengths": lengths, "mixed_head_decay": mixed_heads, + "value_amplitude": amplitude, + "negative_log_decay_per_token": decay, + "maximum_chunk_decay_span": min(tokens - 1, 63) * decay, + "output_relative_error": out_error if math.isfinite(out_error) else None, + "state_relative_error": state_error if math.isfinite(state_error) else None, + "finite": finite, "guards_intact": guards, + "within_one_percent": finite and guards and max(out_error, state_error) < 0.01} + rows.append(row) + output_path.write_text(json.dumps({"complete": False, "cases": rows}, indent=2)) + print(json.dumps(row), flush=True) + # The last case has two unequal sequence lengths and mixed affected / + # unaffected heads. Capture both launches and replay with changed q data. + graph = torch.cuda.CUDAGraph() + torch.cuda.synchronize() + with torch.cuda.graph(graph): + graph_output, graph_final = native.fused_prefill( + q.unsqueeze(0), k.unsqueeze(0), v.unsqueeze(0), matrix.unsqueeze(0), + cumulative.unsqueeze(0), beta.unsqueeze(0), scale, initial, True, cu, + None, out=output, + ) + q.mul_(-1) + for _ in range(3): + graph.replay() + torch.cuda.synchronize() + graph_out_error = ((graph_output[0].double() + oracle).norm() / oracle.norm()).item() + graph_state_error = ((graph_final.double() - oracle_final).norm() / oracle_final.norm()).item() + graph_report = {"replays": 3, "changed_query_inputs": True, + "output_relative_error": graph_out_error, + "state_relative_error": graph_state_error, + "within_one_percent": max(graph_out_error, graph_state_error) < 0.01} + report = {"complete": True, "reference": "independent FP64 sequential gated delta recurrence", + "r4d_sha256": observed, "elapsed_seconds": time.monotonic() - started, + "all_within_one_percent": all(row["within_one_percent"] for row in rows), + "scan_library_sha256": library_sha256, + "graph_capture": graph_report, "cases": rows} + output_path.write_text(json.dumps(report, indent=2)) + return report + + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--output", type=Path, required=True) + parser.add_argument("--library", type=Path) + parser.add_argument("--heads", type=int, choices=(24, 48), default=48) + parser.add_argument("--expect-failure", action="store_true") + args = parser.parse_args() + result = probe(args.output, args.library, args.heads) + print(json.dumps({key: value for key, value in result.items() if key != "cases"})) + passed = result["all_within_one_percent"] and result["graph_capture"]["within_one_percent"] + if args.expect_failure: + assert not passed, "uncorrected control did not reproduce the defect" + else: + assert passed, "native GDN disagrees with the independent recurrence" diff --git a/benchmarks/results/20260914-numerical-corrections/qualification.json b/benchmarks/results/20260914-numerical-corrections/qualification.json new file mode 100644 index 0000000..06b4b49 --- /dev/null +++ b/benchmarks/results/20260914-numerical-corrections/qualification.json @@ -0,0 +1,83 @@ +{ + "platform": "AMD Radeon R9700 (gfx1201)", + "radiance_image_id": "4adbeb66552ed76b329734ee5e830d79ed22961662b886e5029c9920d1bdb7cf", + "libr4d_commit": "e8de4bc1f3dbd608dcb8d3ffceb6b48acdf83bb7", + "scope": "Native numerical regressions; no end-to-end quality or throughput claim", + "sampling": { + "samples_per_arm_and_position": 200000, + "greedy_exact": true, + "positions": [ + { + "position": 37, + "shared_noise_max_probability_error": 0.01952502131462097, + "corrected_max_probability_error": 0.001145005226135254 + }, + { + "position": 86142, + "shared_noise_max_probability_error": 0.01780998706817627, + "corrected_max_probability_error": 0.0009650290012359619 + }, + { + "position": 132739, + "shared_noise_max_probability_error": 0.018975019454956055, + "corrected_max_probability_error": 0.001180022954940796 + } + ] + }, + "gdn": [ + { + "arm": "gdn-original", + "cases": 12, + "all_within_one_percent": false, + "max_output_relative_error": 0.9444960157038121, + "max_state_relative_error": 1.0, + "graph_replays": { + "replays": 3, + "changed_query_inputs": true, + "output_relative_error": 0.02229825229335498, + "state_relative_error": 0.002728025914061443, + "within_one_percent": false + }, + "scan_library_sha256": "bff8fb07da5ac191326ca655b9b589934f459e9fc52502d8aeb0e1ae5e677af1", + "report_sha256": "5f532e9ba42c2742345c5d1250173aca37a1de5159309f7dd7b71238c078549b" + }, + { + "arm": "gdn-corrected-48", + "cases": 12, + "all_within_one_percent": true, + "max_output_relative_error": 0.0034449567651698796, + "max_state_relative_error": 0.0025163873233518637, + "graph_replays": { + "replays": 3, + "changed_query_inputs": true, + "output_relative_error": 0.003416871299398062, + "state_relative_error": 0.0024383175819910753, + "within_one_percent": true + }, + "scan_library_sha256": "ec6681c1b5b524fb8d9f7bf9d7637084abe03978f30bee94a51f4c72d7dbf3e8", + "report_sha256": "7291986df89dd1f18d8c9f9369d2ea706191b249348ed5eb192e350977f3352d" + }, + { + "arm": "gdn-corrected-24", + "cases": 12, + "all_within_one_percent": true, + "max_output_relative_error": 0.0034455666192448127, + "max_state_relative_error": 0.002501668141343721, + "graph_replays": { + "replays": 3, + "changed_query_inputs": true, + "output_relative_error": 0.003408335996166081, + "state_relative_error": 0.002441776999705826, + "within_one_percent": true + }, + "scan_library_sha256": "ec6681c1b5b524fb8d9f7bf9d7637084abe03978f30bee94a51f4c72d7dbf3e8", + "report_sha256": "5bb70ce9afff329d892fe16b59ea6e3c202d69b63ee1da12b5fac37ae343c742" + } + ], + "input_source_sha256": "64b60b16ed503e44e4efef746669f9a008741357c3698bec5e23c9f1c4a60402", + "corrected_source_sha256": "8fbe32751567c2ce0dd4a0378b7aa885c08b81e016e733c8ffc5b046e514998d", + "test_source_sha256": { + "check_dflash_sampling_rng.py": "902e5a31b65a07165f824876610556b29653049e71648024c25069793ec771fc", + "check_gdn_extreme_decay.py": "31e018b580d8e87125e330b03ef22f1ea80483cb6c4ba8aa096516bc7f4070f4" + } +} diff --git a/docs/NUMERICAL_CORRECTIONS.md b/docs/NUMERICAL_CORRECTIONS.md new file mode 100644 index 0000000..14c8199 --- /dev/null +++ b/docs/NUMERICAL_CORRECTIONS.md @@ -0,0 +1,91 @@ +# DFlash sampling and GDN prefill corrections + +These changes correct two independent numerical defects. The qualification uses +synthetic probabilities and random tensors; it requires no model or conversation. +It establishes the arithmetic corrections, not an end-to-end answer-quality or +repetition-rate improvement. + +## DFlash selector proposal noise + +The selector walk samples a proposal with the same `(seed, position)` noise that +target rejection sampling uses to choose a replacement. Conditioning on the +rejected proposal biases that replacement distribution. Offset the selector's +local Philox index by `1 << 30`, adapting the independent draft stream in +[vLLM #54282](https://github.com/vllm-project/vllm/pull/54282), commit +`fe755c88995ad468882517b6c4bdd60138d46a3a`. + +This affects probabilistic drafting. It does not change model positions, cache +indices, greedy selection or the proposal probabilities passed to verification. +Upstream #54282 already fixes this same DFlash2 selector by passing +`IS_DRAFTING=True` to its updated `gumbel_noised_argmax` API. Radiance's pinned +vLLM v0.28.0 predates that API, so this backport adds the same salt directly to +the selector's local RNG index. It is not an additional correction to upstream +vLLM main. The selector's sampling-position buffer is already unclamped. + +With target probabilities `[0.1, 0.5, 0.4]`, draft probabilities `[0.5, 0.3, 0.2]` +and 200,000 draws at each of three positions, the largest absolute probability +error is 0.01781–0.01953 under shared noise and 0.000965–0.001180 after correction. +Target-only controls pass, and greedy outputs remain exact. + +## Extreme GDN decay spans + +The pinned libr4d scan factorizes bounded decay products around a chunk midpoint. +It clamps growing exponentials at 80. A chunk decay span above 160 can attenuate +valid diagonal contributions and erase the carried state while leaving every +result finite. The source is +`r4d_gdn_chunk_scan_k128_v128_c64_bf16.hip` at libr4d +`e8de4bc1f3dbd608dcb8d3ffceb6b48acdf83bb7`. + +After the fast scan, a second GPU kernel detects affected sequence/head pairs +using a conservative span threshold of 128. It recomputes those pairs from the +bounded FP32 recurrence, starting from the original initial state: + +```text +S = exp(g_step) * S +residual = beta * (v - S @ k) +S = S + residual @ k.T +output = scale * (S @ q) +``` + +Other heads retain the fast result. Both launches use the caller's stream; there +is no host readback, synchronization, temporary tensor allocation or decode-path +change. The existing scan ABI and tensor layouts are preserved. The correction +is included in the existing libr4d build patch, so all three image build paths +receive it. + +On 128-token inputs, negative log decay 3.2 produces 46.4% relative output error +and 100% final-state error before correction; decay 8 produces 82.8% and 100%. +After correction, those output errors are approximately 0.17%, and their state +errors are below 0.000014%. Across 12 cases for each of the 48/16-head and +24/8-head layouts, maximum corrected output error is 0.345% and state error is +0.252% against an independent FP64 recurrence. Unequal sequences, partial chunks, +mixed affected heads, large values, and three graph replays with changed inputs +are covered. These are single-GPU tests of both layouts, not distributed TP2 +qualification. + +The extra launch and recurrence add prefill work. End-to-end throughput impact +has not been qualified. Previously computed context state must be rebuilt when +adopting the corrected math; persisted state made by the old kernel should not be +silently reused. No sampler penalty or repetition detector is introduced. + +## Reproduction and recorded evidence + +Build a corrected image using the usual image workflow, then run inside it with +this repository mounted as the working directory: + +```bash +python benchmarks/bin/check_dflash_sampling_rng.py --output /tmp/sampling.json +python benchmarks/bin/check_gdn_extreme_decay.py --output /tmp/gdn-48.json +python benchmarks/bin/check_gdn_extreme_decay.py --heads 24 --output /tmp/gdn-24.json +``` + +The GDN checker also accepts `--library` for an independently compiled scan +translation unit. Use `--expect-failure` with an original unit to verify that the +control reproduces the defect. The sampler checker subtracts the proposal salt +in its control arm to reproduce the old coupling through the same native kernels. + +The tests ran on gfx1201/R9700 using the pinned Radiance 1.0.16 image. The scan +translation units were built with its compiler and the libr4d build's +`-O3 -std=c++17 -fPIC --offload-arch=gfx1201 -mcumode` flags. This qualifies the +changed unit, not a full rebuild of every image variant. Compact numeric evidence +is in [qualification.json](../benchmarks/results/20260914-numerical-corrections/qualification.json). diff --git a/patch_dflash_sampling_rng.py b/patch_dflash_sampling_rng.py new file mode 100644 index 0000000..36c811e --- /dev/null +++ b/patch_dflash_sampling_rng.py @@ -0,0 +1,43 @@ +#!/usr/bin/env python3 +"""Separate DFlash2 selector proposal noise from target replacement noise. + +Adapt the independent draft Philox stream from vLLM PR #54282 (commit +fe755c88995ad468882517b6c4bdd60138d46a3a) to the selector walk's direct call to +gumbel_noised_argmax. Sharing (seed, position) conditions the replacement draw +on the rejected proposal and biases probabilistic rejection sampling. + +The selector uses an unclamped sampling-position buffer. Offset only its local +RNG index: model positions, cache addressing, and greedy sampling are unchanged. +""" + +import sysconfig +from pathlib import Path + +from _patchlib import apply + + +def main(): + path = (Path(sysconfig.get_paths()["purelib"]) + / "vllm/v1/worker/gpu/spec_decode/dflash2/speculator.py") + apply( + path, + " position = tl.load(sample_pos_ptr + flat) - 1\n", + " # vLLM #54282: proposal and target replacement need independent noise.\n" + " # This local RNG index does not change model/cache positions.\n" + " position = tl.load(sample_pos_ptr + flat) - 1 + (1 << 30)\n", + "position = tl.load(sample_pos_ptr + flat) - 1 + (1 << 30)", + "dflash2: separate selector proposal RNG stream", + ) + apply( + path, + "# Candidate ids key the noise, matching the target's own sampling.", + "# Candidate ids key the proposal's independent sampling noise.", + "# Candidate ids key the proposal's independent sampling noise.", + "dflash2: describe independent proposal noise", + ) + # The direct helper has no implicit drafting salt in the pinned vLLM. A + # future upstream sampler API change requires reviewing this adaptation. + + +if __name__ == "__main__": + main() diff --git a/r4d_radiance_extras.patch b/r4d_radiance_extras.patch index cf5e4d1..ec116a0 100644 --- a/r4d_radiance_extras.patch +++ b/r4d_radiance_extras.patch @@ -832,3 +832,109 @@ index 0bd04ac..17ec780 100644 {"gdn_recurrent_update_k128_v128_bf16_fp32state", "gdn", "gdn_recurrent_update", "gdn decode recurrent update: gating, qk l2norm, delta-rule state update, gated rms norm " "and output", +diff --git a/r4d_gdn_chunk_scan_k128_v128_c64_bf16.hip b/r4d_gdn_chunk_scan_k128_v128_c64_bf16.hip +--- a/r4d_gdn_chunk_scan_k128_v128_c64_bf16.hip ++++ b/r4d_gdn_chunk_scan_k128_v128_c64_bf16.hip +@@ -856,6 +856,89 @@ + } + } + ++// The midpoint factorization above clamps positive exponents at 80. A chunk ++// decay span over 160 can therefore erase valid diagonal contributions or the ++// carried state while leaving all outputs finite. Use a conservative span of ++// 128 and recompute affected sequence/head pairs from the defining recurrence: ++// S <- exp(g_step) S; d <- beta (v - S k); S <- S + d k^T; o <- scale S q. ++// Per-token negative decay is bounded. Ordinary heads retain the fast result. ++// This runs on the same stream, with no host readback or synchronization. ++static __device__ __forceinline__ float r4d_gdn_reference_from_bf16(uint16_t value) { ++ return __uint_as_float(static_cast(value) << 16); ++} ++ ++static __device__ __forceinline__ uint16_t r4d_gdn_reference_to_bf16(float value) { ++ const uint32_t bits = __float_as_uint(value); ++ if ((bits & 0x7fffffffU) > 0x7f800000U) ++ return static_cast((bits >> 16) | 0x40U); ++ return static_cast((bits + 0x7fffU + ((bits >> 16) & 1U)) >> 16); ++} ++ ++__global__ __launch_bounds__(256) void r4d_gdn_repair_extreme_decay( ++ const uint16_t *__restrict__ q, const uint16_t *__restrict__ k, ++ const uint16_t *__restrict__ v, const float *__restrict__ cumulative, ++ const float *__restrict__ beta, const float *__restrict__ initial, ++ uint16_t *__restrict__ output, float *__restrict__ final, ++ const int *__restrict__ cu, int heads, int query_heads, int chunk, ++ float scale, float threshold) { ++ const int sequence = blockIdx.x, head = blockIdx.y; ++ const int begin = cu[sequence], end = cu[sequence + 1]; ++ const int thread = threadIdx.x; ++ __shared__ int required; ++ if (thread == 0) { ++ required = 0; ++ for (int first = begin; first < end; first += chunk) { ++ const int last = min(first + chunk, end) - 1; ++ const float span = fabsf(cumulative[(size_t)last * heads + head] - ++ cumulative[(size_t)first * heads + head]); ++ if (span > threshold) { ++ required = 1; ++ break; ++ } ++ } ++ } ++ __syncthreads(); ++ if (!required || begin == end) ++ return; ++ const int row = thread >> 1, half = thread & 1, start_k = half * 64; ++ const int query_head = head / (heads / query_heads); ++ const size_t state_offset = ((size_t)sequence * heads + head) * 128 * 128 + ++ row * 128 + start_k; ++ float state[64]; ++#pragma unroll ++ for (int j = 0; j < 64; ++j) ++ state[j] = initial[state_offset + j]; ++ for (int token = begin; token < end; ++token) { ++ const size_t gate_offset = (size_t)token * heads + head; ++ const float previous = ((token - begin) % chunk == 0) ++ ? 0.0f : cumulative[gate_offset - heads]; ++ const float decay = expf(cumulative[gate_offset] - previous); ++ const size_t qk_offset = ((size_t)token * query_heads + query_head) * 128 + start_k; ++ float keys[64], prediction = 0.0f; ++#pragma unroll ++ for (int j = 0; j < 64; ++j) { ++ keys[j] = r4d_gdn_reference_from_bf16(k[qk_offset + j]); ++ state[j] *= decay; ++ prediction = fmaf(state[j], keys[j], prediction); ++ } ++ prediction += __shfl_xor(prediction, 1, 32); ++ const float residual = beta[gate_offset] * ++ (r4d_gdn_reference_from_bf16(v[gate_offset * 128 + row]) - prediction); ++ float value = 0.0f; ++#pragma unroll ++ for (int j = 0; j < 64; ++j) { ++ state[j] = fmaf(residual, keys[j], state[j]); ++ value = fmaf(state[j], r4d_gdn_reference_from_bf16(q[qk_offset + j]), value); ++ } ++ value += __shfl_xor(value, 1, 32); ++ if (half == 0) ++ output[gate_offset * 128 + row] = r4d_gdn_reference_to_bf16(value * scale); ++ } ++#pragma unroll ++ for (int j = 0; j < 64; ++j) ++ final[state_offset + j] = state[j]; ++} ++ + extern "C" int r4d_gdn_chunk_scan_k128_v128_c64_bf16( + const void* q, const void* k, const void* v, const void* A, const void* g, + const void* beta, const void* h0, void* o, void* ht, const void* cu, +@@ -867,6 +950,12 @@ + (const unsigned short*)q, (const unsigned short*)k, (const unsigned short*)v, + (const unsigned short*)A, (const float*)g, (const float*)beta, (const float*)h0, + (unsigned short*)o, (float*)ht, (const int*)cu, H, Hg, scale); ++ const hipError_t scan_error = hipGetLastError(); ++ if (scan_error != hipSuccess) return (int)scan_error; ++ r4d_gdn_repair_extreme_decay<<>>( ++ (const uint16_t*)q, (const uint16_t*)k, (const uint16_t*)v, ++ (const float*)g, (const float*)beta, (const float*)h0, ++ (uint16_t*)o, (float*)ht, (const int*)cu, H, Hg, BT, scale, 128.0f); + return (int)hipGetLastError(); + } +