unsloth/scripts/benchmarks
Daniel Han 314ab6ae86 flex: fuse LoRA refresh into a single torch.addmm per layer
Replace the copy+merge_adapter pair in refresh_lora_merge_from_pristine
with one torch.addmm(pristine, B, A, alpha=scaling, out=W_inf) per
LoraLayer, then set merged_adapters directly so PEFT's forward
short-circuits to base_layer(x).

Previously each refresh did two passes per weight: a bf16 copy from
pristine, then PEFT merge_adapter which materialises a full [out, in]
fp32 delta via get_delta_weight and in-place adds it back. The fused
path skips the transient delta allocation and runs one cuBLAS GEMM
instead. cuBLAS accumulates the bf16 matmul in fp32 internally, so the
numerical result stays within 1 bf16 ULP of PEFT's path (verified on
the rank-32 Qwen3-4B adapter: max abs diff 1.22e-04).

DoRA, fan_in_fan_out, and lora_bias=True layers fall back to PEFT's
get_delta_weight/merge path via a single trailing merge_adapter call
after restoring their base_layer.weight from pristine. rslora is not a
fallback -- PEFT folds alpha/sqrt(r) into module.scaling[adapter], so
the fused addmm picks it up transparently via alpha=.

Drift verification still passes: base bit-identical across 10
perturb+refresh cycles, inference state deterministic after LoRA
restore.
2026-04-21 04:53:06 +00:00
..
results flex: double-copy LoRA rollout to avoid bf16 merge/unmerge drift 2026-04-21 04:52:49 +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 flex: support --load_in_4bit with PEFT adapter (bnb-4bit shard) 2026-04-21 02:13:06 +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-20 15:57:55 +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 flex: fuse LoRA refresh into a single torch.addmm per layer 2026-04-21 04:53:06 +00:00
qwen3_grpo_naive.py [pre-commit.ci] auto fixes from pre-commit.com hooks 2026-04-20 01:38:38 +00:00
qwen3_grpo_notebook.py [pre-commit.ci] auto fixes from pre-commit.com hooks 2026-04-20 14:24:24 +00:00
qwen3_grpo_tpaged.py [pre-commit.ci] auto fixes from pre-commit.com hooks 2026-04-20 01:38:38 +00:00
qwen3_grpo_unified.py [pre-commit.ci] auto fixes from pre-commit.com hooks 2026-04-20 14:24:24 +00:00
qwen3_grpo_vllm.py [pre-commit.ci] auto fixes from pre-commit.com hooks 2026-04-19 14:45:48 +00:00
README.md Add FA4 + persistent CB benchmarks and a naive TRL baseline 2026-04-20 01:38:12 +00:00
unsloth_grpo_common.py [pre-commit.ci] auto fixes from pre-commit.com hooks 2026-04-19 14:45:48 +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.

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:

  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.