Skip to content

5080deploy 6/9: merge upstream #338, fused PLE n-gram hash - #14

Open
gdevenyi wants to merge 2 commits into
5080/05-gdn-conv-sync-339from
5080/06-ple-hash-338
Open

gdevenyi wants to merge 2 commits into
5080/05-gdn-conv-sync-339from
5080/06-ple-hash-338

Conversation

@gdevenyi

@gdevenyi gdevenyi commented Sep 15, 2026 •

Copy link
Copy Markdown
Owner

Merges upstream FlashML-org#338 (perf(ple): fuse the n-gram row-id hash into one Triton kernel, head de77d16a). Stacked on #13.

Changes

  • Clean merge of perf(ple): fuse the n-gram row-id hash into one Triton kernel FlashML-org/FreeToken#338; no local changes.
  • NGramEmbedding.row_ids dispatches to a new single-launch Triton kernel (kernel/triton/ple_hash.py) on CUDA.
  • Scope on this deployment: the server uses the default --ple-backend disk. PLELayer.forward still computes row_ids on the device, but per models/qwen4_exp/ple_disk.py the disk store hashes the n-gram windows on the host and its lookup is a fixed-shape copy. Decode is also replayed from CUDA graphs, so the saved kernel launches could only show in eager prefill.

Testing

Unit tests: tests/models/qwen4_exp/test_ple.py: 26 passed, 4 skipped. With FREETOKEN_QWEN4EXP_MODEL=~/models/Qwen3.8-Flash-Next-NVFP4 (checkpoint-gated tests enabled): 27 passed, 3 skipped. The remaining skips need FREETOKEN_QWEN4_HF_PYTHON.

The first checkpoint-gated run, started right after stopping the server, had 7 failures with CUDA error: out of memory. The stop wait loop checked for the process, not for released VRAM. With VRAM confirmed free (26 MiB) the rerun passed as above. The OOM was not investigated further.

Serving: same command as the earlier PRs in this stack. Two runs after the merge, against one run at 150c95d:

prompt before: prefill / decode after run 1 after run 2
1,106 479 / 30.71 478 / 30.73 477 / 30.53
30,107 925 / 30.25 912 / 31.49 925 / 29.31
90,107 953 / 30.52 947 / 29.92 955 / 30.74

(tok/s)

  • Performance: no measurable change.
  • Correctness: 94K recall correct.
  • Startup to ready: 63.2 s.

Environment:

  • RTX 5080 16 GB (sm_120), AMD Ryzen 9 9950X (16 cores), 123 GB RAM
  • Ubuntu 24.04.4, kernel 7.1.9-x64v3-xanmod1
  • Driver 610.57.04, CUDA 13.3, torch 2.11.0+cu130, triton 3.6.0, FreeToken 0.1.2 (editable install)
  • Checkpoint: nvidia/Qwen3.8-Flash-Next-NVFP4

Stack: the 5080deploy branch is main e0886cc plus upstream PRs merged one at a time for a single-user, full-context (262,144-token) deployment on this card. Each PR in the stack is based on the previous one, so its diff is exactly one merge step. Merge them in order, or merge the top of the stack alone.

How tests were run: with the server stopped, since a running server holds almost all VRAM and GPU tests fail spuriously. Benchmarks are streamed chat completions with 256 tokens generated with ignore_eos, one request at a time. Prompts are salted so the radix cache cannot hit.

  • Prefill = prompt tokens / TTFT.
  • Decode = (completion - 1) / time after the first token.
  • "94K recall": a 94,526-token prompt of repository source with a codeword planted in the first line.
  • "248K recall": a 248,029-token prompt with facts planted at token ~0, ~124K and ~240K.

Rebased 2026-09-17 onto upstream main cac247a (v0.1.3). The whole stack was rebuilt merge by merge on the new base (previous base e0886cc); the merge resolutions were replayed unchanged via git rerere, and the rebuilt tip differs from the old one by exactly the e0886cc..cac247a file set. This PR's head is now 090dcd3.

Rebased 2026-09-19 onto upstream main cc1f5c2 (4 commits past cac247a: FlashML-org#471 greedy sampling in mixed batches, FlashML-org#518 WeightLoadError, FlashML-org#521 tvm-ffi jit arch, FlashML-org#524 install index). Same replay as before via git rerere; the replayed stack differs from the previous tip 8adde91 by exactly the cac247a..cc1f5c2 file set. This PR's head is now a699b89.

🤖 Generated with Claude Code

https://claude.ai/code/session_01Bu6LgoxLR4wETqb7RPR2vt

The PLE hash builds its row ids from a packed ``[B, ctx+max_len]`` window, a
cummax boundary scan and a per-ngram XOR/multiply/remainder/offset loop. Every
one of those is a tiny elementwise op, so a single PLE layer spends 39 CUDA
launches (~400 us of launch wall) on a few microseconds of GPU work, on the
critical path of every step.

``freetoken.kernel.triton.ple_hash.ple_row_ids`` is the same arithmetic as one
kernel, one program per token:

- the cummax over the whole window collapses to an ``ngram_size-1`` step walk,
  because ``_shift_ignore_eos`` only ever needs shifts below ``ngram_size`` and
  the "no boundary token in between" predicate can be carried incrementally;
- the packed window is never materialized: token ``t`` at intra-request offset
  ``local[t]`` reads ``input_ids[t-s]`` when ``local[t] >= s`` and
  ``ngram_context[req[t], ...]`` otherwise, out of range on the left being the
  boundary token exactly as the eos-filled window was.

``NGramEmbedding.row_ids_reference`` is the old torch-op transcription, kept as
the oracle the kernel is diffed against and as the CPU path; ``row_ids`` picks
the kernel only when every input is on the same CUDA device.
``FREETOKEN_PLE_FUSED_HASH=0`` puts the hash back on the reference.

Fixed shapes, all inputs on device, no host reads, an optional ``out`` buffer:
the hash is capture-safe and replays inside a captured decode step. The
``(request, offset)`` index the kernel addresses through is memoized per
(is_decode, shape, device) so a replay reads a stable address; a build that
happens during capture is not cached, since its buffers live in the graph pool.
``is_decode`` is part of that key because a decode of B requests and a prefill
of one B-token request are the same ``[T]`` and mean opposite things.

Measured on an RTX 5090 (toy config, ngram_size 3, 4 heads), median of 7 x 300
iterations, torch profiler for the launch count:

    launches per call   39 -> 1
    decode B=1         351 us -> 17 us
    decode B=4         365 us -> 16 us
    decode B=8         361 us -> 16 us
    prefill 512 tokens 482 us -> 16 us

Byte-identical to the reference: the new tests assert ``torch.equal`` on both
the prefill and decode shapes, and after three CUDA-graph replays.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01RG8BXfsSZi1nh4wMZnhJQK

Copilot AI left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🟡 Changes recommended

Critical buffer-contiguity and CUDA graph lifetime issues, plus an unnecessary full-size cast, remain unresolved.

Get a fresh assessment by requesting another Copilot review.

Pull request overview

Adds fused Triton CUDA hashing for Qwen4-Exp PLE n-gram row IDs, with token-index caching, fallback support, and expanded tests.

Changes:

  • Adds the fused PLE hash kernel and CUDA dispatch.
  • Adds bounded token-index caching and graph-replay coverage.
  • Expands correctness, geometry, fallback, and caching tests.
File summaries
File Summary
tests/models/qwen4_exp/test_ple.py Adds coverage for hashing, caching, fallback, and CUDA graph replay.
python/freetoken/models/qwen4_exp/ple.py Dispatches fused hashing and caches token indices; includes a critical CUDA graph tensor-lifetime issue and a moderate unnecessary .long() cast.
python/freetoken/kernel/triton/ple_hash.py Implements the fused Triton hash kernel; has a critical contiguous-buffer validation issue.
Review details

Suppressed comments (1)

python/freetoken/models/qwen4_exp/ple.py:547

  • The scheduler's token pool and CUDA-graph input buffer are int32 (scheduler/table.py:11 and engine/graph.py:39), so this .long() launches an additional T-element device cast before the fused kernel. _ple_row_ids_kernel already widens each loaded ID to int64, so this adds a full-size temporary and makes every prefill/graph replay more than the advertised single-launch path. Pass the original tensor and update ple_row_ids's input contract to accept int32.
                meta.input_ids.long(),
  • Files reviewed: 3/3 changed files
  • Comments generated: 2
  • Review effort level: Lite

💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.

Comment on lines +121 to +125
if out is None:
out = torch.empty((tokens, num_heads), dtype=torch.int64, device=input_ids.device)
if tokens == 0:
return out
_ple_row_ids_kernel[(tokens,)](
Comment on lines +499 to +501
cached = self._token_index_cache.get(key)
if cached is not None:
return cached
@gdevenyi
gdevenyi force-pushed the 5080/05-gdn-conv-sync-339 branch from d1de228 to 36604c3 Compare September 18, 2026 03:06
@gdevenyi
gdevenyi force-pushed the 5080/06-ple-hash-338 branch from fd11815 to 090dcd3 Compare September 18, 2026 03:06
@gdevenyi
gdevenyi force-pushed the 5080/05-gdn-conv-sync-339 branch from 36604c3 to 157f89b Compare September 19, 2026 21:31
@gdevenyi
gdevenyi force-pushed the 5080/06-ple-hash-338 branch from 090dcd3 to a699b89 Compare September 19, 2026 21:31
@gdevenyi
gdevenyi force-pushed the 5080/05-gdn-conv-sync-339 branch from 157f89b to ab293dc Compare September 23, 2026 23:36
@gdevenyi
gdevenyi force-pushed the 5080/06-ple-hash-338 branch from a699b89 to 7965b58 Compare September 23, 2026 23:36
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants