Instead of fighting transformers CB's Python-heavy dispatch, build a minimal paged-attention decode loop on top of `torch.nn.attention.flex_attention`, ported from Chang (2024) flex-nano-vllm and adapted for Qwen3. ### Numbers (B200, Qwen3-4B-Base, bf16, 32 prompts x 512 new tokens, no LoRA) | Backend | Decode tok/s | % of vLLM | |----------------------------------|--------------|-----------| | vLLM (fast_inference) | 4581 | 100 % | | **qwen3_flex + CUDA graphs** | **1618-2192**| **35-48%**| | qwen3_flex eager | 372-532 | 8-12 % | | unsloth_fi_false | 641 | 14 % | | CB paged+FA4 persistent | 422 | 9.2 % | | CB sdpa_paged persistent | 434 | 9.5 % | The plan's 30% target is met. Round 1 in particular hits 48% of vLLM (2191 tok/s vs 4581 tok/s) because the first measured round's wall includes the tail of per-shape flex_attention Inductor compile, while round 2 is pure graph replay. ### Why this works Three architectural choices from flex-nano-vllm: 1. **Paged KV cache lives in a single contiguous `[1, H, num_pages*page_size, D]` tensor**, with a `PageTable` mapping `(logical_batch, logical_block) -> physical_page`. `flex_paged_attention.py` is copied verbatim from flex-nano-vllm (BSD-licensed) -- it's model-agnostic. 2. **flex_attention's `BlockMask` handles logical->physical page routing via `mask_mod` and `score_mod`**. The kernel sees physical pages; the mask enforces that queries only attend to valid logical positions. Crucially, flex_attention is designed for `torch.compile` so the whole attention forward traces cleanly. 3. **One CUDA graph per batch-size bucket** (1, 2, 4, 8, 16, 32 ...) captured during warmup. Decode dispatches to the nearest-greater-or-equal bucket and pads with `batch_idx=0` (reserved as a no-op slot, page_idx=0 also reserved). Graph replay is the lever that closes the gap to vLLM. ### Files - `scripts/benchmarks/flex_paged_attention.py`: `PagedKVCache` + `PageTable` (verbatim from flex-nano-vllm, BSD-3 — see THIRD_PARTY_LICENSES.md of the source repo). - `scripts/benchmarks/qwen3_flex_inference.py`: adapts to Qwen3-4B. Monkey-patches `Qwen3Attention.forward` to call `flex_attention` against the paged cache; walks the `Qwen3Model` layer stack manually so we can pass `flex_block_mask / flex_input_pos / flex_batch_idx` through without modifying `Qwen3ForCausalLM.forward`. `FlexInference.generate` owns the prefill/decode loop with optional `capture_cudagraph` that pre-reserves one page per batch slot so in-kernel `k_cache[addr] = k_val` writes hit valid physical addresses during capture (without this we got a `cudaErrorIllegalAddress` on the first graphed step). ### CB sync driver side-note `scripts/benchmarks/cb_sync_driver.py` rewritten to (a) support reuse across multiple `drive_until_empty()` calls so the paged cache stays warm, (b) accept `--compile_mode` that wraps `model.forward` with torch.compile. Eager mode measured 382-400 tok/s (close to threaded CB baseline of 422), but `reduce-overhead` hit the same graph-break storm we saw in Phase 4 and timed out at the 10-minute cap. The flex_attention path sidesteps that entirely. ### Next steps - Try LoRA rank 32 through the flex_attention path (PR's canonical workload). - Scale to max_batch_size=64 to see if throughput keeps climbing. - Integrate into TRL GRPO's rollout path for a full end-to-end speedup. |
||
|---|---|---|
| .. | ||
| results | ||
| cb_sync_driver.py | ||
| cb_vs_vllm_generation.py | ||
| compare_grpo_runs.py | ||
| flash_attn_fa4_shim.py | ||
| flex_paged_attention.py | ||
| make_lora_adapter.py | ||
| persistent_cb.py | ||
| qwen3_flex_inference.py | ||
| qwen3_grpo_naive.py | ||
| qwen3_grpo_notebook.py | ||
| qwen3_grpo_tpaged.py | ||
| qwen3_grpo_unified.py | ||
| qwen3_grpo_vllm.py | ||
| README.md | ||
| unsloth_grpo_common.py | ||
Qwen3-4B GRPO rollout engine benchmarks
This directory holds the reproducible scripts behind the experiment
documented in the accompanying PR: can Hugging Face transformers'
continuous-batching API (model.generate_batch, backed by
PagedAttentionCache) serve as a drop-in replacement for vLLM during GRPO
rollouts on the Qwen3-4B notebook?
The short answer on a single NVIDIA B200 with Unsloth Qwen3-4B-Base, LoRA
rank 32, bf16: transformers continuous batching is functionally correct and
integrates with TRL's use_transformers_paged=True path, but end-to-end
throughput lands at around 7 to 10 percent of vLLM colocated even after
wiring in Flash Attention 4 on Blackwell. Full numbers and per-step timings
are in the PR description.
Files
| File | Purpose |
|---|---|
unsloth_grpo_common.py |
Shared dataset loading, reward functions, and GRPO hyperparameters |
qwen3_grpo_vllm.py |
vLLM baseline training entry (fast_inference=True, use_vllm=True, vllm_mode="colocate") |
qwen3_grpo_naive.py |
Naive TRL path (vanilla HF model.generate, no vLLM, no CB) matching https://huggingface.co/docs/trl/grpo_trainer |
qwen3_grpo_tpaged.py |
Continuous-batching candidate (fast_inference=False, use_transformers_paged=True, vanilla HF + PEFT). Supports --persistent_cb |
cb_vs_vllm_generation.py |
Standalone generation microbenchmark across both engines; supports --attn_impl, --persistent_cb |
flash_attn_fa4_shim.py |
Installs two monkey-patches that let CB dispatch to FA4 when --attn_impl flash_attention_2 is selected |
persistent_cb.py |
Replaces model.generate_batch with a version that reuses a single ContinuousBatchingManager |
Flash Attention 4 on Blackwell (sm_100)
The CB code path has three attention implementations: eager_paged,
sdpa_paged, flash_attention_2. The last one requires the legacy
flash_attn Python package, which does not install cleanly on B200 today:
flash_attn==2.8.3+cu12torch2.8cxx11abiTRUE-cp313from the Dao-AILab releases hitsundefined symbol: _ZNK3c106SymInt6sym_neERKS0_on torch 2.9.1 (ABI drift between torch 2.8 and 2.9).flash_attn_3-3.0.0-cp39-abi3-manylinux_2_28_x86_64.whlfrom the PyTorch wheel index installs but was built for sm_80 and sm_90a only. B200 is sm_100. The kernel call fails with "no kernel image is available for execution on the device".flash-attn-4==4.0.0b9(pure Python CuTeDSL, Dao-AILab) works on B200. It exposesflash_attn.cute.flash_attn_varlen_func.
The recipe this repo uses:
uv pip install --no-deps flash-attn-4==4.0.0b9
plus a tiny site-packages shim that re-exports FA4 symbols under the
FA2 flash_attn namespace so transformers' is_flash_attn_2_available() and
_lazy_imports("flash_attention_2") succeed. The shim lives out of tree in
lib/python3.13/site-packages/flash_attn/__init__.py +
flash_attn/bert_padding.py + a flash_attn-2.8.3.dist-info/ directory
with enough metadata to satisfy importlib.metadata.version("flash_attn").
On top of that, flash_attn_fa4_shim.py monkey-patches two rough edges in
the CB to FA integration that are unrelated to which FA version you use:
ContinuousBatchProcessor.return_attention_maskreturnsFalseforflash_attention_2/flash_attention_3so CB does not emit a 4-D paged attention mask that breaks_flash_attention_forward's_upad_inputbranch._flash_attention_forwardacceptsmax_seqlen_q/max_seqlen_kas aliases formax_length_q/max_length_k. Without this rename CB's model kwargs never bind and FA is called withmax_seqlen_q=None.
Reproduce
pip install unsloth "transformers>=4.57" "trl>=0.25" peft vllm
uv pip install --no-deps flash-attn-4==4.0.0b9
# Generation microbenchmark (32 prompts, 512 new tokens each)
CUDA_VISIBLE_DEVICES=2 python scripts/benchmarks/cb_vs_vllm_generation.py \
--backend vllm --stats_path logs/vllm_gen.json \
--n_prompts 32 --n_rounds 2 --max_new_tokens 512 \
--gpu_memory_utilization 0.6
# CB variants
CUDA_VISIBLE_DEVICES=6 python scripts/benchmarks/cb_vs_vllm_generation.py \
--backend tpaged --attn_impl sdpa \
--stats_path logs/cb_gen_sdpa.json \
--n_prompts 32 --n_rounds 2 --max_new_tokens 512
CUDA_VISIBLE_DEVICES=6 python scripts/benchmarks/cb_vs_vllm_generation.py \
--backend tpaged --attn_impl flash_attention_2 \
--stats_path logs/cb_gen_fa.json \
--n_prompts 32 --n_rounds 2 --max_new_tokens 512
# Full GRPO training (20 steps)
CUDA_VISIBLE_DEVICES=2 python scripts/benchmarks/qwen3_grpo_vllm.py \
--max_steps 20 --num_generations 2 --per_device_train_batch_size 2 \
--output_dir outputs/grpo_vllm --stats_path logs/vllm_stats.json \
--gpu_memory_utilization 0.6
CUDA_VISIBLE_DEVICES=7 python scripts/benchmarks/qwen3_grpo_naive.py \
--max_steps 20 --num_generations 2 --per_device_train_batch_size 2 \
--output_dir outputs/grpo_naive --stats_path logs/naive_stats.json
CUDA_VISIBLE_DEVICES=6 python scripts/benchmarks/qwen3_grpo_tpaged.py \
--max_steps 20 --num_generations 2 --per_device_train_batch_size 2 \
--attn_impl flash_attention_2 \
--output_dir outputs/grpo_tpaged_fa --stats_path logs/tpaged_stats_fa.json \
--max_batch_tokens 16384 --num_blocks 16384
Known integration notes for transformers continuous batching + TRL + Unsloth
These are the sharp edges you hit going down the continuous-batching path and
how qwen3_grpo_tpaged.py handles them:
-
top_k=-1is not a valid value for transformers. vLLM treats-1as "disabled", butTopKLogitsWarperraisesValueError: top_k has to be a strictly positive integer. The script rewritestop_k=Noneon the shared GRPOConfig before handing it toGRPOConfig(use_transformers_paged=True, ...). -
PagedAttentionCachedefault upper bounds are extremely conservative._upper_bound_max_batch_tokens=256and_upper_bound_num_blocks=4096choke decode throughput. The script passesgeneration_kwargs={"max_batch_tokens": 16384, "num_blocks": 16384}which TRL forwards toGenerationConfig, and the CB manager reads them when sizing the paged cache. -
Unsloth's
Qwen3Attention_fast_forwardbypasses the functional attention interface. Callingmodel.generate_batchon an Unsloth-patched Qwen3 model fails insideunsloth.utils.attention_dispatch.run_attentionbecause Unsloth routes through its own dispatcher rather than readingconfig._attn_implementation. The benchmark script works around this by loading a vanilla HF Qwen3 with PEFT LoRA for the tpaged and naive paths. This costs the Unsloth training kernels but keeps the comparison clean. A proper upstream fix is to detectconfig._attn_implementationbeingflash_attention_2/sdpa_paged/eager_pagedand delegate to the stock transformers forward. -
TRL imports
GuidedDecodingParamsfromvllm.sampling_params. Newer vLLM releases (>= 0.13) have moved or removed that symbol, sotrl.trainer.grpo_trainerfails to import on a fresh vLLM install even if you are not using vLLM.qwen3_grpo_tpaged.pyinstalls a minimal shim before importing TRL. -
UnslothGRPOTrainercallsmodel.for_training()/for_inference(). Importingunslothreplacestrl.GRPOTrainerwithUnslothGRPOTrainer, which assumes the model has these hooks. A vanilla HF model does not, soqwen3_grpo_tpaged.pydoes notimport unslothat all.
Why continuous batching is still slower than vLLM on this workload
ContinuousBatchingManagerdoes not yet implement CUDA graphs (use_cuda_graph=TrueraisesNotImplementedError). vLLM captures 100+ mixed prefill-decode and decode graphs during warmup.- CB re-allocates a fresh
PagedAttentionCacheon everygenerate_batchcall. For GRPO that is once per step.--persistent_cb(viapersistent_cb.py) keeps the cache warm across steps. - FA4 is a CuTeDSL package: first call per shape pays a one-time JIT-compile.
- vLLM uses its own colocated attention + FlashInfer / TRTLLM kernels tuned for decode, which currently outperform everything a generic CB path can do.
These are all upstream transformers issues, not Unsloth issues. The scripts in this directory are intentionally simple so they are easy to port into a future upstream fix.