Summary of sweep (all with CUDA graph capture):
| Batch | flex tps | vLLM tps | flex / vLLM | flex mem |
|------:|---------:|---------:|------------:|---------:|
| 8 | 680 | 1900 | 35.8 % | 44 GB |
| 16 | 1626 | 3698 | 44.0 % | 44 GB |
| 32 | 3134 | 6318 | 49.6 % | 44 GB |
| 64 | 5474 | 10459 | **52.3 %** | 44 GB |
| 128 | 5565 | 14996 | 37.1 % | 81 GB |
| 256 | 5812 | 21170 | 27.5 % | 154 GB |
Canonical GRPO (batch 64 + LoRA rank 32):
- vLLM: 7775 tok/s / 156 GB
- flex: **5616 tok/s / 44 GB** = **72% of vLLM at 3.5x less memory**
Up from 9 % (transformers CB) at the start of this work.
Best FlexKernelOptions after sweep:
decode: PRESCALE_QK, USE_TMA, BLOCKS_ARE_CONTIGUOUS, num_warps=8, num_stages=3
prefill: FORCE_USE_FLEX_ATTENTION, PRESCALE_QK, USE_TMA
Biggest single win: `num_warps=8` (+28% at batch 64). Inductor's default
picks 4 on small Triton blocks; 8 is better for our decode shapes.
`BLOCKS_ARE_CONTIGUOUS` adds +10% (safe in our setup because
PageTable.reserve allocates pages sequentially on a fresh batch).
TMA adds 2-3%.
Items documented that broke correctness or didn't help:
- ROWS_GUARANTEED_SAFE=true NaNs the softmax on padded batch slots that
only attend to reserved page 0 (mask returns False for every kv_idx).
- BACKEND="TRITON_DECODE" from the docs raises
NameError('TRITON_DECODE is not defined') inside Inductor.
- USE_TMA + torch.compile(call_model_with_flex_kwargs) -> misaligned
address at runtime (compile breaks TMA alignment assumptions).
- torch.compile(flex_attention, mode="max-autotune") nests cudagraph_trees
inside our raw CUDA graph -> "Cannot prepare for replay during
capturing stage". max-autotune-no-cudagraphs works but same throughput
as default mode.
- compile on call_model_with_flex_kwargs: same as eager walker (CUDA
graph capture already fuses every op in the walker).
- num_warps=4 / 16 both slower than num_warps=8.
CLI surface added to qwen3_flex_inference.py:
--decode_kernel_options JSON (FlexKernelOptions for decode)
--prefill_kernel_options JSON (same for prefill)
--compile_model_forward MODE (optional torch.compile on the walker)
25 lines
No EOL
1.2 KiB
JSON
25 lines
No EOL
1.2 KiB
JSON
{
|
|
"backend": "qwen3_flex",
|
|
"capture_cudagraph": true,
|
|
"lora_adapter": null,
|
|
"n_prompts": 64,
|
|
"n_decoded_tokens": 27985,
|
|
"wall_times_s": [
|
|
8.14494420203846,
|
|
6.53967270400608,
|
|
5.112081662984565,
|
|
5.535065989010036,
|
|
6.894235474988818
|
|
],
|
|
"median_wall_s": 6.53967270400608,
|
|
"best_wall_s": 5.112081662984565,
|
|
"decode_tps_median": 4279.266144750168,
|
|
"decode_tps_best": 5474.286571482827,
|
|
"max_new_tokens": 512,
|
|
"peak_memory_gb": 44.21115064620972,
|
|
"sample_completions": [
|
|
"First, let's find the sum of the numbers in Amanda's list. The sum of the first n even numbers is given by the formula n(n+1). In this case, n = 50 (since there are 50 even numbers from 2 to 100). So, the sum of Amanda's list is 50(50+1) = ",
|
|
"Let the roots of the equation be $r, r^2, r^3, r^4, r^5$ in geometric progression. By Vieta's formulas, the sum of the roots is $r + r^2 + r^3 + r^4 + r^5 = 180$. Dividing both sides by $r^5$, we get $1",
|
|
"First, let's find the angle \\( \\angle AOB into three equal parts. The area of each smaller triangle is:\n\\[ \\frac{\\sqrt{3}}{12} \\text{ triangle} = \\frac{\\sqrt{3}/4 \\]\n\nNow, let's find the value of \\( k + m + n \\). We have:\n\\[ k = 1 \\]\n\\["
|
|
]
|
|
} |