This file provides guidance to Claude Code (claude.ai/code) when working with code in this repository.
JetEngine is a lightweight inference engine for block diffusion language models (SDAR, LLaDa, dLLM-Var). It supports dense and MoE architectures, tensor parallelism, CUDA graph capture, paged KV caching, and flashinfer acceleration. Built from nanovllm.
# Install (requires flash-attn, CUDA GPU)
pip install flash-attn --no-build-isolation
pip install .
# Single GPU (must use accelerate or torchrun for distributed init)
CUDA_VISIBLE_DEVICES='0' accelerate launch --multi_gpu example.py
# or
CUDA_VISIBLE_DEVICES='0' torchrun --nproc_per_node=1 example.py
# Multi-GPU (data parallel by default, uses all visible GPUs)
accelerate launch --multi_gpu example.pyNo test suite exists in this repo. Correctness is verified through end-to-end MATH-500 evaluation (see Evaluation section below).
LLM (in jetengine/llm.py) is a thin subclass of LLMEngine (jetengine/engine/llm_engine.py). The engine orchestrates the full inference loop:
add_request()— tokenizes prompt, creates aSequence, hands it to theSchedulerstep()— callsscheduler.schedule()to get prefill/denoise batches, runs them throughModelRunner, then postprocesses logits back through the schedulergenerate()— batched generation where all prompts are added upfront, loopstep()until donegenerate_streaming()— streaming generation with backpressure:max_activecontrols how many sequences run concurrently, new prompts are added as old ones finish
JetEngine operates in two distinct regimes depending on the relationship between total requests and max_active:
All sequences fit in memory simultaneously. The flow is:
- Prefill phase: All sequences get prefilled in one batch
- Pure denoise phase: All sequences denoise together with no interruptions
- Drain phase: Sequences finish at different times, batch size shrinks
In this mode, there is no interleaving of prefill and denoise. The chain mechanism (chain=5 when no pending) completes a full block (4 denoising steps + SAVING) in a single scheduler step. This maximizes GPU utilization since the batch stays large.
Profile with: tests/bench_throughput.py or tests/profile_realistic.py
More sequences than can fit in memory. The engine uses generate_streaming():
- Initial
max_activesequences get prefilled - Denoise steps run on active sequences
- As sequences finish and free KV cache blocks, new sequences from the waiting queue get prefilled
- Prefill and denoise alternate in each
step()call
In this mode, prefill of new sequences interleaves with denoise of active sequences. The chain depth is limited (chain=2 when prefills are pending) to let the scheduler refill slots. This introduces more scheduling overhead but keeps GPU utilization high.
Profile with: tests/profile_streaming.py
The MATH-500 evaluation is an instance of Mode 2: 500 problems × num_generations (e.g. 4) = 2000 total sequences, but max_active is typically 64-128.
Manages sequence lifecycle through states: WAITING → PREFILLING → DENOISING → SAVING → FINISHED. Key responsibilities:
- Block management — allocates/deallocates paged KV cache blocks via
BlockManager - Batching — prepares separate prefill and denoise batches each step
- Postprocessing — implements all remasking strategies (sampling + token selection). Has two paths:
postprocess()— general path, supports per-sequence sampling params and mixed strategiespostprocess_unify()— optimized batched path when all sequences share the same sampling params
The chain mechanism runs multiple denoising steps within a single step() call, avoiding scheduler overhead between steps. Adaptive chain depth:
chain=5when no pending prefills (full block in one step: 4 denoise + 1 SAVING)chain=4when batch_size ≤ 8 (streaming tail)chain=2when prefills are pending (let scheduler refill slots)
Key optimizations in the chain path:
- Context reuse: positions/block_tables are cached and reused for pure DENOISING chain steps
- Zero-sync fast path: for intermediate chain steps (step < denoising_steps - 1), GPU→CPU synchronization is eliminated by predicting SAVING transitions from CPU-side state
- Chain batch tokens: batch tensor from postprocess is passed directly to prepare_denoise, avoiding per-sequence Python loop + torch.cat
Represents a single generation request. Tracks block diffusion state: intermediate_block_tokens (current block being denoised as a tensor), block_trajectory, block_logprobs, block_entropies. The commit_block() method finalizes a denoised block into the sequence's token stream.
Handles model execution: weight loading, KV cache allocation, input preparation (prefill vs denoise have different attention patterns), and CUDA graph capture/replay.
Init order: load_model() → warmup_model() → allocate_kv_cache() → _init_flashinfer() → capture_cudagraph()
Key features:
- flashinfer:
BatchPrefillWithPagedKVCacheWrapperwithuse_cuda_graph=True, one wrapper per batch size (1..128). Falls back toflash_attn_with_kvcachefor larger batches. - CUDA graphs: Captured for batch sizes 1 to min(max_num_seqs, 128). Graph outputs hidden states;
compute_logits(LM head) runs outside the graph to allow selective computation. - Selective logits: Only computes LM head for DENOISING sequences (not SAVING ones), saving ~25% of the largest GEMM.
Models register via @register_model("name") decorator. Currently registered:
sdar— dense SDAR models (Qwen-based architecture)sdar_moe— SDAR MoE variant with fused MoE kernelsllada— LLaDa diffusion models
Defined in SamplingParams.remasking_strategy and implemented in the scheduler's postprocess methods:
sequential— unmask left-to-rightlow_confidence_static— unmask top-k by confidence (default for evaluation)low_confidence_dynamic— unmask above threshold, fallback to sequentialentropy_bounded— unmask by cumulative entropy budgetrandom— random selection
Config(jetengine/config.py):mask_token_id(required),block_length,kvcache_block_size(must be multiple of 256),enforce_eager(disable CUDA graphs),torch_compile,quantize_fp8SamplingParams(jetengine/sampling_params.py):block_length,denoising_steps,remasking_strategy,dynamic_threshold,eb_threshold,repetition_penalty
Custom tensor-parallel layers: QKVParallelLinear, MergedColumnParallelLinear, RowParallelLinear, VocabParallelEmbedding, ParallelLMHead.
Optimized operations:
RMSNorm—@torch.compilewith fused residual addSiluAndMul—@torch.compilewith liger-kernel- Attention — FA3 (Flash Attention 3, Hopper SM90) for denoise (paged KV, CUDA graph compatible), flashinfer fallback, Triton sparse attention for prefill
Custom Triton kernels for block prefill attention (sparse_attn_varlen) and store_kvcache_kernel. Also includes fused MoE kernel.
The primary correctness and performance benchmark. Uses open-r1/scripts/evaluate.py with the SDAR-4B-Chat model.
Config: open-r1/recipes/evaluation/eval_SDAR-4B_math.yaml
Launch command (single GPU):
cd /mnt/shared-storage-user/liudawei/home/open-r1
conda activate /mnt/shared-storage-user/liudawei/envs/open-r1
CUDA_VISIBLE_DEVICES=0 accelerate launch --multi_gpu scripts/evaluate.py \
--config recipes/evaluation/eval_SDAR-4B_math.yaml \
--dataset_name /mnt/shared-storage-user/liudawei/MyDatasets/MATH-500/ \
--num_generations 4 \
--temperature 0.6 \
--top_p 0.95 \
--max_active 128 \
--pass_k 1 4 \
--output_dataset_name <test_name>Important notes:
- Must use
accelerate launch --multi_gpu(not--num_processes=1) — it initializes the distributed process group required by JetEngine --pass_kvalues must be ≤--num_generations(e.g.--pass_k 1 4requires--num_generations≥ 4)- This is a Mode 2 (streaming) workload: 500 × 4 = 2000 sequences through max_active=128 slots
- The eval config in the yaml has
max_num_seqs: 512(max capacity), butmax_activecontrols actual concurrency
Results location: /mnt/shared-storage-user/liudawei/work_dirs_ckpt/grpo_checkpoints/open-r1-eval/<test_name>/results.json
Quality threshold: pass@1 ≥ 0.678 (baseline). Current best: pass@1 ≈ 0.705 at temp=0.6/top_p=0.95.
cd /mnt/shared-storage-user/liudawei/home/JetEngine
conda activate /mnt/shared-storage-user/liudawei/envs/open-r1
# Quick throughput measurement
CUDA_VISIBLE_DEVICES=0 torchrun --nproc_per_node=1 tests/bench_throughput.py
# Detailed profiling (monkey-patched step components)
CUDA_VISIBLE_DEVICES=0 python /tmp/profile_opt20.py # or write inline scriptMeasures peak throughput with homogeneous prompts. All sequences start together, denoise together. Chain=5 (full block per step). Typical results on H200: 4114 tok/s at bs=128.
CUDA_VISIBLE_DEVICES=0 accelerate launch --multi_gpu tests/profile_streaming.py \
--model /mnt/shared-storage-user/liudawei/Models/SDAR/SDAR-4B-Chat \
--max_active 128 --num_prompts 256 --mask_id 151669Measures throughput with diverse prompts streaming through the engine. Shows prefill/denoise interleaving, batch fragmentation, and chain hit rates. Typical chain hit rate: ~59% with diverse prompts (vs 99%+ with homogeneous).
Component breakdown (from profiling with torch.cuda.synchronize barriers):
| Component | Time | % of Wall | Calls | Per-call |
|---|---|---|---|---|
| denoise_fwd | 2.655s | 58.8% | 167 | 15.9ms |
| postprocess | 0.649s | 14.4% | 168 | 3.9ms |
| prefill_fwd | 1.170s | 25.9% | 1 | 1170ms |
| schedule | 0.003s | 0.1% | 35 | 0.09ms |
Note: Profiling sync overhead inflates total time (4.5s vs 3.3s unprofiled). Proportions are informative, not absolute.
- Block diffusion generates tokens in fixed-size blocks (
block_length=4), each block goes throughdenoising_steps=4steps before being committed. - Each denoising step transfers a fixed number of tokens from masked to unmasked (1 per step for block_length=4, denoising_steps=4).
- The SAVING state runs the forward pass to write final token representations into KV cache — skipping SAVING forward corrupts KV cache (discovered: pass@1 dropped from 0.68 → 0.14).
- CUDA graphs are captured for batch sizes 1 to 128. Hidden states output from graph, LM head runs outside for selective logits.
consistent_sampling_paramsflag enables batched postprocessing. Set automatically when all sequences share the sameSamplingParams.- The
build/directory contains stale copies — always work in the top-leveljetengine/package.
| Opt | Description | pass@1 | tok/s (bs=128) |
|---|---|---|---|
| baseline | Original | 0.683 | ~1500 |
| opt13 | chain across SAVING (100% chain hit) | 0.694 | - |
| opt17 | CUDA graph 1-128, selective logits | - | 2677 (bs=64) |
| opt18 | flashinfer paged attention | - | 4109 |
| opt19 | bs=128 verified | 0.708 | 4109 |
| opt20 | GPU sync elimination, chain batch tokens | 0.705 | 4114 |
| opt21 | Sparse sampling (only masked positions) | 0.722 | 4180 |
| opt22 | Lazy entropy (skip in chain intermediate) | 0.709 | 4242 |
| opt23 | Sparse logits (LM head only masked positions) | 0.711 | 4268 |
| opt24 | FA3 (Flash Attention 3, Hopper SM90) | 0.722 | 5677 |
Throughput: 5677 tok/s (bench_throughput.py, bs=128, max_tokens=128)
FA3 replaces flashinfer for denoise attention, combining KV cache write + paged attention in one kernel. Forward pass dominates, with cuBLAS GEMM as the main bottleneck.
- FP8 weight quantization: 35% per-layer relative error compounds over 36 layers → NaN. FP8 e4m3fn has only 3 mantissa bits. Tested with
torch._scaled_mm(1.4-1.9x GEMM speedup) but quality loss is fundamental. - INT8 via torch._int_mm: 7-8x SLOWER than BF16 on H200. PyTorch's INT8 kernel is not optimized for Hopper architecture.
- torch.compile on decoder layers: Incompatible with CUDA graph capture — attention layer's
get_context()dynamic dispatch causes graph breaks. - denoising_steps < 4: Quality drops catastrophically (0.683 → 0.441 with steps=3).
- Skip SAVING forward pass: Corrupts KV cache (pass@1 → 0.14).
- FP8 dynamic input quantization: Overhead > GEMM savings for M ≤ 256.
- Fused sampling from logits: flashinfer's
top_k_top_p_sampling_from_logitsis slower than separate pipe+sampling.
- Tensor parallelism: Near-linear scaling for multi-GPU (most impactful remaining option)
- NVIDIA Transformer Engine FP8: Proper FP8 with delayed scaling and amax history (not installed, needs
pip install transformer-engine) - Async postprocess overlap: Run non-critical postprocess work on secondary CUDA stream while next forward starts (~1ms overlap potential)
- Fused Triton postprocess kernel: Fuse softmax+sampling into single kernel (diminishing returns after sparse sampling)
- Larger block_length (8): 2x tokens/step but needs model retraining
- INT4 GPTQ/AWQ: Needs calibration data + AutoGPTQ/AutoAWQ (not installed), major quality risk
- Implement the optimization (new kernel, algorithm change, etc.)
- Unit test — every new operator or algorithm must have a minimal standalone test verifying correctness against a reference implementation
- Throughput benchmark —
tests/bench_throughput.pyfor Mode 1 - End-to-end eval — MATH-500 pass@1 ≥ 0.678 (Mode 2, streaming)
- Git commit on success — each successful speedup gets its own commit
- Git revert on failure — revert to last good commit and try a different approach
- Each optimization that passes both unit tests and end-to-end eval gets committed
- Commit messages:
opt<N>: <description>(e.g.opt20: GPU sync elimination in chain fast path) - If an optimization breaks quality or doesn't improve throughput,
git revertto the last good state before continuing
- New Triton kernels: must have a test comparing output against a pure PyTorch reference implementation, checking both correctness (max relative error < threshold) and edge cases
- New sampling algorithms (e.g. sparse sampling): must have a test verifying the output distribution matches the dense reference, and that token selection is identical
- Postprocess changes: must verify against
postprocess()(the general path) as ground truth - Unit tests go in
tests/test_*.py, benchmarks intests/bench_*.py
| File | Purpose |
|---|---|
tests/bench_throughput.py |
Mode 1 throughput benchmark (ideal decode) |
tests/profile_streaming.py |
Mode 2 profiler (streaming, diverse prompts) |
tests/profile_realistic.py |
Mode 1 profiler (homogeneous prompts, detailed) |
tests/test_correctness.py |
Unit tests for entropy, commit_block, postprocess |
tests/bench_fp8_static.py |
FP8 vs BF16 GEMM benchmark |
tests/test_fp8_static.py |
FP8 linear layer correctness test |
tests/bench_compile.py |
torch.compile A/B benchmark |
tests/test_sparse_sampling.py |
Sparse sampling correctness + benchmark |
- Remasking Strategy Exploration:
research/remasking_strategies.md— systematic study of commit-order policies, hybrid strategies, and portfolio approaches for block diffusion decoding. Phases: diagnostic (exhaustive 24 orders) → signal isolation → hybrid design → speed optimization → validation.
jetengine/engine/scheduler.py— postprocess_unify, chain_state, zero-sync path, try_allocate_chain_blocksjetengine/engine/model_runner.py— flashinfer init, CUDA graph capture, selective logits, chain batch tokensjetengine/engine/sequence.py— _commit_block_from_cpu, tensor-based block statejetengine/engine/llm_engine.py— chain mechanism, adaptive chain depth, generate_streamingjetengine/layers/linear.py— FP8 quantize_to_fp8 method, _fp8_linear functionjetengine/layers/attention.py— FA3 path (priority) + flashinfer fallback in BlockAttention.forwardjetengine/config.py— quantize_fp8, torch_compile flags