The sidecar design in the previous cut stored K/V at
[max_batch, n_kv, max_seq, D] layout, but the flex_attention block mask
is built for the paged cache's [1, H, n_pages*page_size, D] layout.
Shared layers running through that mismatch either needed a parallel
block mask (expensive to build per call, per layer) or had to fall back
to SDPA, which breaks the single-CUDA-graph capture story and silently
dropped the sliding-window mask on shared sliding layers.
This commit drops the sidecar entirely. Shared layers now reference the
store layer's `PagedKVCache` directly:
- `patch_gemma4_attention_forwards` allocates a cache on every
non-shared layer, then walks shared layers and points
`shared._paged_cache = store._paged_cache`, plus stashes the store
attention module itself on `shared._store_attn`.
- The store layer keeps its post-rotary `k`, `v` on
`self._last_k_val`, `self._last_v_val` before its own paged update
so shared successors can read the same packed prefill tensors.
- Shared-layer forward reads `_last_k_val` / `_last_v_val` on prefill
(q_len > 1) and `_paged_cache.k_cache` / `.v_cache` on decode. The
block_mask dispatched by `self.layer_type` works for both regimes
uniformly -- one block mask builder, one kernel compile per regime,
one CUDA graph per batch-size bucket.
Also fixes the 4-bit path: `AutoModelForCausalLM.from_pretrained` on
`unsloth/gemma-4-E2B-it-unsloth-bnb-4bit` resolves to
`Gemma4ForConditionalGeneration`, so `.model.embed_tokens` does not
exist. The loader now detects the multimodal wrapper, drops the vision
and audio towers, and moves the language_model into a
`Gemma4ForCausalLM` shell -- mirroring the bf16 path.
Benchmarks on a single B200 (sm_100), CUDA_VISIBLE_DEVICES=2, Gemma-4
E2B-it, n_prompts=64 n_rounds=5 max_new_tokens=512 max_batch_size=64
capture_cudagraph no-fa4_prefill BLOCK_M=32 BLOCK_N=32 (prefill) /
BLOCK_M=16 BLOCK_N=16 (decode):
| Config | Peak GB | Median tok/s | Best tok/s |
|-----------------|---------|--------------|------------|
| bf16 | 14.2 | 2794 | 2797 |
| bf16 + LoRA r32 | 23.0 | 2798 | 2801 |
| 4bit + LoRA r32 | 13.0 | 1659 | 1821 |
Drift verification (10 perturb+refresh cycles, noise_scale=0.01):
`base_bit_identical = true`, `inference_deterministic = true`.
Sample completions are coherent math reasoning ("Let the isosceles
trapezoid be $ABCD$ with bases ...").
|
||
|---|---|---|
| .. | ||
| results | ||
| cb_sync_driver.py | ||
| cb_vs_vllm_generation.py | ||
| compare_grpo_runs.py | ||
| flash_attn_fa4_shim.py | ||
| flex_autotune_replay.py | ||
| flex_paged_attention.py | ||
| gemma4_flex_inference.py | ||
| make_lora_adapter.py | ||
| persistent_cb.py | ||
| qwen3_flex_inference.py | ||
| qwen3_grpo_naive.py | ||
| qwen3_grpo_tpaged.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.
Install
Prereqs: torch >= 2.5 for torch.nn.attention.flex_attention. The
Triton backend that flex_attention uses by default runs on Ampere,
Hopper, and Blackwell -- no separate install.
FA4 (CuTeDSL) targets Hopper and Blackwell only. qwen3_flex_inference.py
auto-enables FA4 on supported GPUs and falls back to the Triton
flex_attention backend elsewhere. --fa4_prefill forces on (warns +
falls back if the GPU does not support it); --no-fa4_prefill forces off.
CUDA 13 (recommended, used on B200 / RTX 50xx):
pip install --index-url https://download.pytorch.org/whl/cu130 torch
pip install "flash-attn-4[cu13]"
CUDA 12 (H100 boxes still on cu12):
pip install torch # default index is cu12
pip install flash-attn-4
Pin flash-attn-4==4.0.0b9 to match this benchmark. The [cu13]
extra pulls in nvidia-cutlass-dsl built for CUDA 13.
qwen3_flex_inference.py runs on both Qwen3 and Llama-3.2 (the only
arch-specific branch is Qwen3's per-head QK RMSNorm; the rest of the
flex_attention + paged KV + CUDA graphs stack is identical). Pass
--model_name unsloth/Llama-3.2-3B-Instruct to target Llama, along with
--chat_template native to use Llama's shipped Instruct template instead
of the Qwen3 GRPO template.
gemma4_flex_inference.py extends the engine to unsloth/gemma-4-E2B-it.
Gemma-4 is not a drop-in: its text backbone has KV-sharing layers
(layers 15-34 consume full-sequence K/V produced by a store layer), two
attention regimes (full_attention with head_dim=512 / rope_theta=1e6
and sliding_attention with head_dim=256 / sliding_window=512),
per-layer input embeddings, four norms per block with double residuals,
and a final logit softcap. The new file keeps the shared helpers
(PagedKVCache, PageTable, LoRA double-copy, drift verification)
imported from qwen3_flex_inference.py and adds:
- a KV-sharing sidecar dict, sized
[max_batch, n_kv, max_seq, head_dim]per store layer, populated at prefill and read by the paired shared layers through eager SDPA (shared layers don't fit the paged-cache block-mask shape); - dual RoPE precomputation —
rotary_emb(x, pos, layer_type)called once per unique layer type, indexed byself.layer_type; - a walker that threads
per_layer_inputsfromget_per_layer_inputs+project_per_layer_inputsinto each layer and applies thelayer_scalarmultiply at layer end; tanh(logits / final_logit_softcapping) * final_logit_softcappingapplied on the lm_head output.
Requires transformers>=5.5.0 for the gemma4 module; if absent the
script exits with a clear install hint. Gemma-4 head_dim=256 exceeds FA4
on sm_100 (B200), so pass --no-fa4_prefill and small Triton blocks:
CUDA_VISIBLE_DEVICES=6 python scripts/benchmarks/gemma4_flex_inference.py \
--model_name unsloth/gemma-4-E2B-it \
--n_prompts 64 --max_new_tokens 512 --capture_cudagraph \
--no-fa4_prefill \
--prefill_kernel_options '{"FORCE_USE_FLEX_ATTENTION": true, "BLOCK_M": 32, "BLOCK_N": 32}' \
--decode_kernel_options '{"BLOCK_M": 16, "BLOCK_N": 16}' \
--stats_path logs/flex_gemma4_bf16.json
| GPU | arch | sm | Auto FA4 | Triton flex_attention |
|---|---|---|---|---|
| A100 | Ampere | sm_80 | off (uses Triton) | Works |
| H100 / H200 | Hopper | sm_90 | on | Works |
| RTX 50xx | Blackwell | sm_120 | on | Works |
| B200 / GB200 | Blackwell | sm_100 | on | Works |
The transformers continuous-batching path's FA4 wiring (the
flash_attn_fa4_shim.py monkey-patches and the
site-packages/flash_attn/__init__.py namespace shim that makes FA4
visible under the FA2 import name) is covered below under "Known
integration notes".
Reproduce
pip install unsloth "transformers>=4.57" "trl>=0.25" peft vllm
# 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
qwen3_grpo_naive.py and qwen3_grpo_tpaged.py accept an optional
--compile_mode {default,reduce-overhead,max-autotune-no-cudagraphs} flag.
When set, trainer.model.forward (and trainer.ref_model.forward, if present)
are wrapped with torch.compile after trainer construction. The vLLM driver
has no such flag because vLLM owns its own inference graph. --compile_dynamic
(default on) toggles dynamic-shape compilation.
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.