unsloth/scripts/benchmarks
Daniel Han d5a4ee22ad benchmarks: gemma4_flex_inference -- per-layer-type sliding window mask
The first cut routed every non-shared layer through one causal block mask
and relied on SDPA with is_causal=True for shared layers. That gives the
right semantics for Gemma-4's full_attention layers but silently drops the
sliding_attention window, so sliding layers attend far beyond their
512-token window as soon as the prefix grows.

This commit builds a block mask per attention regime (full_attention is
pure causal; sliding_attention is causal AND q_pos - kv_pos < window) and
passes both into each attention call as a dict, letting the patched
forward select by self.layer_type. For the shared-KV sidecar path, the
SDPA call now receives an explicit attn_mask composed the same way, so
sliding shared layers also respect the window. Strict less-than
comparison matches Unsloth's flex-attention convention for GPT-OSS.

`_causal_blockmask_with_window` and `_prefill_blockmask_with_window` are
local to this file; the shared helpers in `flex_paged_attention.py` stay
untouched. `FlexGemma4Inference` now caches both logical decode masks,
slices both per-row in `_decode_block_mask`, and runs the PageTable's
logical->physical conversion on each before passing them down.
2026-04-21 08:59:54 +00:00
..
results flex: drop scripts/benchmarks/results/stats JSONs 2026-04-21 05:52:11 +00:00
cb_sync_driver.py [pre-commit.ci] auto fixes from pre-commit.com hooks 2026-04-20 15:57:55 +00:00
cb_vs_vllm_generation.py benchmarks: add --chat_template and --enforce_eager to cb_vs_vllm_generation 2026-04-21 07:52:50 +00:00
compare_grpo_runs.py [pre-commit.ci] auto fixes from pre-commit.com hooks 2026-04-20 14:37:42 +00:00
flash_attn_fa4_shim.py Add FA4 + persistent CB benchmarks and a naive TRL baseline 2026-04-20 01:38:12 +00:00
flex_autotune_replay.py [pre-commit.ci] auto fixes from pre-commit.com hooks 2026-04-21 00:25:27 +00:00
flex_paged_attention.py [pre-commit.ci] auto fixes from pre-commit.com hooks 2026-04-21 04:58:26 +00:00
gemma4_flex_inference.py benchmarks: gemma4_flex_inference -- per-layer-type sliding window mask 2026-04-21 08:59:54 +00:00
make_lora_adapter.py [pre-commit.ci] auto fixes from pre-commit.com hooks 2026-04-20 14:24:24 +00:00
persistent_cb.py [pre-commit.ci] auto fixes from pre-commit.com hooks 2026-04-20 01:38:38 +00:00
qwen3_flex_inference.py [pre-commit.ci] auto fixes from pre-commit.com hooks 2026-04-21 06:52:23 +00:00
qwen3_grpo_naive.py benchmarks: consolidate GRPO entrypoints and extract shared helpers 2026-04-21 06:18:23 +00:00
qwen3_grpo_tpaged.py benchmarks: consolidate GRPO entrypoints and extract shared helpers 2026-04-21 06:18:23 +00:00
qwen3_grpo_vllm.py benchmarks: consolidate GRPO entrypoints and extract shared helpers 2026-04-21 06:18:23 +00:00
README.md benchmarks: add gemma4_flex_inference for unsloth/gemma-4-E2B-it 2026-04-21 08:45:05 +00:00
unsloth_grpo_common.py [pre-commit.ci] auto fixes from pre-commit.com hooks 2026-04-21 06:19:56 +00:00

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-cp313 from the Dao-AILab releases hits undefined 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.whl from 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 exposes flash_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:

  1. ContinuousBatchProcessor.return_attention_mask returns False for flash_attention_2 / flash_attention_3 so CB does not emit a 4-D paged attention mask that breaks _flash_attention_forward's _upad_input branch.
  2. _flash_attention_forward accepts max_seqlen_q / max_seqlen_k as aliases for max_length_q / max_length_k. Without this rename CB's model kwargs never bind and FA is called with max_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 by self.layer_type;
  • a walker that threads per_layer_inputs from get_per_layer_inputs + project_per_layer_inputs into each layer and applies the layer_scalar multiply at layer end;
  • tanh(logits / final_logit_softcapping) * final_logit_softcapping applied 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:

  1. top_k=-1 is not a valid value for transformers. vLLM treats -1 as "disabled", but TopKLogitsWarper raises ValueError: top_k has to be a strictly positive integer. The script rewrites top_k=None on the shared GRPOConfig before handing it to GRPOConfig(use_transformers_paged=True, ...).

  2. PagedAttentionCache default upper bounds are extremely conservative. _upper_bound_max_batch_tokens=256 and _upper_bound_num_blocks=4096 choke decode throughput. The script passes generation_kwargs={"max_batch_tokens": 16384, "num_blocks": 16384} which TRL forwards to GenerationConfig, and the CB manager reads them when sizing the paged cache.

  3. Unsloth's Qwen3Attention_fast_forward bypasses the functional attention interface. Calling model.generate_batch on an Unsloth-patched Qwen3 model fails inside unsloth.utils.attention_dispatch.run_attention because Unsloth routes through its own dispatcher rather than reading config._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 detect config._attn_implementation being flash_attention_2 / sdpa_paged / eager_paged and delegate to the stock transformers forward.

  4. TRL imports GuidedDecodingParams from vllm.sampling_params. Newer vLLM releases (>= 0.13) have moved or removed that symbol, so trl.trainer.grpo_trainer fails to import on a fresh vLLM install even if you are not using vLLM. qwen3_grpo_tpaged.py installs a minimal shim before importing TRL.

  5. UnslothGRPOTrainer calls model.for_training() / for_inference(). Importing unsloth replaces trl.GRPOTrainer with UnslothGRPOTrainer, which assumes the model has these hooks. A vanilla HF model does not, so qwen3_grpo_tpaged.py does not import unsloth at all.

Why continuous batching is still slower than vLLM on this workload

  • ContinuousBatchingManager does not yet implement CUDA graphs (use_cuda_graph=True raises NotImplementedError). vLLM captures 100+ mixed prefill-decode and decode graphs during warmup.
  • CB re-allocates a fresh PagedAttentionCache on every generate_batch call. For GRPO that is once per step. --persistent_cb (via persistent_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.