Compare commits

..
105 Commits
Author SHA1 Message Date
fxmarty-amdandGitHub 53f6dd5c6f [CI][ROCm] Fix test_ocp_mx_wikitext_correctness reference value (#49690)
Signed-off-by: Felix Marty <Felix.Marty@amd.com>
2026-07-27 21:47:07 +00:00
Andreas KaratzasandGitHub 99115fcdcd [CI] Initialize DeepEP FP8 test weights (#49912)
Signed-off-by: Andreas Karatzas <akaratza@amd.com>
2026-07-28 05:29:08 +08:00
1053e248f0 [ROCm][Quantization][5/N] Refactor quark_moe w8a8-int8 w/ oracle (#46765)
Signed-off-by: amd-sourjya <amd-sourjya@users.noreply.github.com>
Co-authored-by: amd-sourjya <amd-sourjya@users.noreply.github.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: Andreas Karatzas <akaratza@amd.com>
2026-07-27 16:01:34 -05:00
Wentao YeandGitHub b5bcb3ce88 [Refactor] Remove dead code in multiple files (#49745)
Signed-off-by: yewentao256 <zhyanwentao@126.com>
2026-07-27 15:58:26 -04:00
Wentao YeandGitHub b2f9e4caa4 [DSv4 Perf] Adaptive topk width, 1.0% E2E throughput improvement (#50004)
Signed-off-by: yewentao256 <zhyanwentao@126.com>
2026-07-27 15:56:52 -04:00
831d3848f1 [Core] Fail fast when /dev/shm is too small for the shm ring buffer (#48879)
Signed-off-by: Dr Andrea Tassi <andrea@verticular.uk>
Co-authored-by: Dr Andrea Tassi <andrea@verticular.uk>
Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
2026-07-27 19:48:35 +00:00
fd10e8946d [Test] Regression test for hybrid-Mamba eagle cache-peek in Mooncake connector (#43559) (#48361)
Signed-off-by: Rishi Puri <riship@nvidia.com>
Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
2026-07-27 19:03:58 +00:00
ed13deb376 [Bugfix][CPU] Fall back to torch for unaligned swigluoai on NEON/vec MoE (#49985)
Signed-off-by: oops-oom <73481342@qq.com>
Co-authored-by: oops-oom <73481342@qq.com>
Co-authored-by: Claude <noreply@anthropic.com>
2026-07-27 18:35:57 +00:00
99de48e98f Fix MLA padding and grouped topk routing in the Transformers modelling backend (#49982)
Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
2026-07-27 18:32:32 +00:00
Matthew BonanniGitHubgemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
bf2b45b5d6 [Attention] Integrate FlashAttention 4 SM100 headdim 256 support (#42669)
Signed-off-by: Matthew Bonanni <mbonanni@redhat.com>
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2026-07-27 18:25:50 +00:00
8112b6c997 [MRV2] Always build attn metadata at capture time (#49364) (#49995)
Signed-off-by: Woosuk Kwon <woosuk@inferact.ai>
Co-authored-by: Woosuk Kwon <woosuk@inferact.ai>
2026-07-27 17:25:17 +00:00
TobyJBellandGitHub 15d65f8669 [Bugfix] Changed speech to text chunk timestamp to cumulative approach (#41131)
Signed-off-by: Toby Bell <toby.bell1702@hotmail.co.uk>
2026-07-27 17:06:46 +00:00
e3c2fc3b3c [Rust Frontend][gRPC] Add server and model discovery (#49491)
Signed-off-by: Connor Carpenter <connorc@nvidia.com>
Co-authored-by: Nick Hill <nickhill123@gmail.com>
2026-07-27 09:53:27 -07:00
Nicolò LucchesiandGitHub 2b465b2c42 [Misc][PD] Nixl cleanup get_backend_aware_kv_block_len and virtually_split_kv_in_blocks (#49988)
Signed-off-by: NickLucche <nicolo.lucchesi@mistral.ai>
2026-07-27 18:51:37 +02:00
yzong-rhandGitHub 3f47a8384d [Bugfix] Fix VLLM_ENFORCE_STRICT_TOOL_CALLING mutation in tests (#49846)
Signed-off-by: Yifan Zong <yzong@redhat.com>
2026-07-27 16:12:14 +00:00
Guan-Ming ChiuGitHubIsotr0pymergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
04502deca2 [Perf] Hash videos by source bytes (#49607)
Signed-off-by: Guan-Ming (Wesley) Chiu <105915352+guan404ming@users.noreply.github.com>
Signed-off-by: Guan-Ming Chiu <105915352+guan404ming@users.noreply.github.com>
Co-authored-by: Isotr0py <2037008807@qq.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-07-27 23:57:55 +08:00
d2ca3002d9 [MRV2][Performance] Skip no-op FP32 logits materialization (#47711)
Signed-off-by: jesse <szxfml@gmail.com>
Signed-off-by: Song Zhixin <szxfml@gmail.com>
Co-authored-by: Wentao Ye <44945378+yewentao256@users.noreply.github.com>
Co-authored-by: Jee Jee Li <pandaleefree@gmail.com>
2026-07-27 15:31:36 +00:00
Umut PolatandGitHub 27d7061ef6 [Bugfix] Restore truncate_prompt_tokens for Jina rerank/score online (#49963)
Signed-off-by: Umut Polat <52835619+umut-polat@users.noreply.github.com>
2026-07-27 22:52:59 +08:00
Guan-Ming ChiuandGitHub ef9975d021 [Bugfix] Reject pipeline parallelism for DiffusionGemma (#45828)
Signed-off-by: Guan-Ming (Wesley) Chiu <105915352+guan404ming@users.noreply.github.com>
2026-07-27 14:37:54 +00:00
Roberto L. CastroandGitHub 56c96b0d91 [Perf] Tune LL BF16 Router GEMM (#48774)
Signed-off-by: LopezCastroRoberto <rocastro@redhat.com>
Signed-off-by: Roberto L. Castro <38211239+LopezCastroRoberto@users.noreply.github.com>
2026-07-27 10:25:37 -04:00
Nick HillGitHubmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
59a6b0411d [Core] Fix internal LB load-balancing (#49204)
Signed-off-by: Nick Hill <nickhill123@gmail.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-07-27 14:00:46 +00:00
Rui "Garry" GaoandGitHub dbccc5ae32 [Model] Enable EVS for Qwen3.5 (#48912)
Signed-off-by: Rui "Garry" Gao <garrygaogg@gmail.com>
2026-07-27 13:42:35 +00:00
liminfei-amdandGitHub a89015c6df [Perf] Make merge attention context count a runtime argument (#48739)
Signed-off-by: liminfei-amd <91481003+liminfei-amd@users.noreply.github.com>
2026-07-27 13:24:42 +00:00
neweyesandGitHub 96fa3f42c9 [Perf] Skip ll_bf16 router GEMM warmup for non-MoE models (#49659)
Signed-off-by: neweyes <328719365@qq.com>
2026-07-27 05:16:42 -07:00
rongfu.lengGitHubmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
81962bb699 [Bugfix]Reject invalid FlashInfer MNNVL workspaces (#49043)
Signed-off-by: lengrongfu <lenronfu@gmail.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-07-27 08:12:23 -04:00
Harry MellorandGitHub 92e8518d37 Improve Transformers modelling backend fx tracer (#49957)
Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
2026-07-27 12:51:16 +01:00
Ronen SchafferandGitHub 77cba0259f [KV Offloading] Per-request tier filtering with TierFilter/TierMatcher (#48123)
Signed-off-by: Ronen Schaffer <ronen.schaffer@ibm.com>
2026-07-27 13:29:57 +03:00
Andreas KaratzasandGitHub 30fbd05537 [ROCm] Use backend-default dot precision for ReplaySSM (#49909)
Signed-off-by: Andreas Karatzas <akaratza@amd.com>
2026-07-27 18:11:06 +08:00
0906123953 [ROCm] [Model] Enable TML inkling (#48841)
Signed-off-by: tjtanaa <tunjian.tan@embeddedllm.com>
Co-authored-by: OpenAI Codex <codex@openai.com>
2026-07-27 10:05:46 +00:00
bc3629b1c4 [ROCm][CI] Skip three torchao tests of gfx950 until torchao==0.18 is released (#49732)
Signed-off-by: Felix Marty <Felix.Marty@amd.com>
Co-authored-by: Felix Marty <Felix.Marty@amd.com>
Co-authored-by: Claude <noreply@anthropic.com>
2026-07-27 09:36:38 +00:00
312ea82e75 [CI][ROCm] Make hf-xet reconstruction safe on shared NFS (#49837)
Signed-off-by: Andreas Karatzas <akaratza@amd.com>
Co-authored-by: OpenAI Codex <noreply@openai.com>
2026-07-27 17:23:12 +08:00
394beb633b [Bugfix][ROCm] Use batch DMA for CPU KV cache loads (#49843)
Signed-off-by: Andreas Karatzas <akaratza@amd.com>
Co-authored-by: OpenAI Codex <noreply@openai.com>
2026-07-27 02:18:01 -07:00
haoyangli0109GitHubtjtanaamergify[bot] <37929162+mergify[bot]@users.noreply.github.com>Douglas LehrAndreas Karatzas
7f599d7854 [communication] [bugfix] fix quickreduce acc error in cudagraph mode (#46913)
Signed-off-by: Haoyang Li <lihaoyang0109@gmail.com>
Signed-off-by: tjtanaa <tunjian.tan@embeddedllm.com>
Co-authored-by: tjtanaa <tunjian.tan@embeddedllm.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
Co-authored-by: Douglas Lehr <91553416+dllehr-amd@users.noreply.github.com>
Co-authored-by: Andreas Karatzas <akaratza@amd.com>
2026-07-27 01:38:45 -07:00
eb290ab673 [Bugfix][CPU] Zero-pad MoE intermediate size for grouped-gemm TP alignment (#49591)
Signed-off-by: jiang1.li <jiang1.li@intel.com>
Co-authored-by: Claude Sonnet 5 <noreply@anthropic.com>
2026-07-27 16:32:23 +08:00
Andreas KaratzasandGitHub cbc3a87200 [Tokenizer] Use HF config for HF tokenizers (#49907)
Signed-off-by: Andreas Karatzas <akaratza@amd.com>
2026-07-27 07:49:11 +00:00
liuzhenweiandGitHub afc94523c9 [XPU][CI] Use platform device in InputBatch V2 test (#49939)
Signed-off-by: zhenwei-intel <zhenwei.liu@intel.com>
2026-07-27 15:28:02 +08:00
8061dc26bd [Bugfix] Normalize sparse MLA warmup compression ratios (#49392)
Signed-off-by: Wu, Xiaochang <xiaochang.wu@intel.com>
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
2026-07-27 07:02:48 +00:00
fd9d2ede6f [Rust Frontend] Keep --max-model-len engine-owned (#49944)
Co-authored-by: OpenAI Codex <codex@openai.com>
Signed-off-by: Bugen Zhao <i@bugenzhao.com>
2026-07-27 14:50:08 +08:00
Andreas KaratzasandGitHub e09900436c [CI][ROCm] Reduce kernel test runtime (#49915)
Signed-off-by: Andreas Karatzas <akaratza@amd.com>
2026-07-27 06:44:21 +00:00
Zhenzhong XuGitHubYi Liumergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
5d07e268b1 [Quantization][INC]Add MXFP8 Linear Support (#47514)
Signed-off-by: Zhenzhong1 <zhenzhong.xu@intel.com>
Signed-off-by: Zhenzhong Xu <zhenzhong.xu@intel.com>
Co-authored-by: Yi Liu <yi4.liu@intel.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-07-27 14:26:31 +08:00
Dao007foreverGitHubmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
d742856610 [3/N][Core][KV Connector] Support reliable partial-tail KV offload for sub-block prompts (#49502)
Signed-off-by: Dao Le <Dao007forever@gmail.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-07-26 23:22:17 -07:00
Tianmu LiGitHubCodexLi, Jiang <jiang1.li@intel.com>
544cb724c8 [CPU][Spec Decode] Optimize GDN conv path for speculative decoding (#48577)
Signed-off-by: Li, Tianmu <tianmu.li@intel.com>
Co-authored-by: Codex <codex@openai.com>
Co-authored-by: Li, Jiang <jiang1.li@intel.com>
2026-07-27 06:20:14 +00:00
Fadi ArafehGitHubLi, Jiang <jiang1.li@intel.com>
c314af1abf [CPU][Perf] INT8 Fused MoE Kernel for Arm CPUs (#48637)
Signed-off-by: Fadi Arafeh <fadi.arafeh@arm.com>
Co-authored-by: Li, Jiang <jiang1.li@intel.com>
2026-07-27 05:53:09 +00:00
f19ee27e39 [Hardware][Power] Add FAST_EXP for Power (#49571)
Signed-off-by: Akash kaothalkar <akash.kaothalkar@ibm.com>
Co-authored-by: Akash kaothalkar <akash.kaothalkar@ibm.com>
2026-07-27 05:47:29 +00:00
Andreas KaratzasandGitHub 5f89a03dcb [CI] Explicitly tear down speculative decode runners (#49910)
Signed-off-by: Andreas Karatzas <akaratza@amd.com>
2026-07-27 05:32:46 +00:00
53397fbfac [Bugfix][KV Offload][P2P] Fix EngineCore crash reconnecting to a reaped peer (#49823)
Signed-off-by: Jason Yao <wsyjh8@gmail.com>
Co-authored-by: Or Ozeri <oro@il.ibm.com>
2026-07-27 08:07:00 +03:00
Andreas KaratzasandGitHub 49f31d7cee [ROCm] Make vllm_c RMSNorm output contiguous (#49913)
Signed-off-by: Andreas Karatzas <akaratza@amd.com>
2026-07-26 23:59:14 -05:00
Nick HillandGitHub 74d3b799e1 [Bugfix] Fix mHC block-M prenorm GEMM cross-row reduction carry-over (#49429)
Signed-off-by: Nick Hill <nickhill123@gmail.com>
2026-07-27 04:42:47 +00:00
29fdeab254 [XPU][CI] Add more test cases in Intel GPU CI (#49422)
Signed-off-by: zengxian <xiangdong.zeng@intel.com>
Co-authored-by: Kunshang Ji <kunshang.ji@intel.com>
2026-07-27 12:40:22 +08:00
ff6173997d [CI] Add kimi and k3 auto-labeling rules (#49895)
Signed-off-by: Joe Cotant <joe@inferact.ai>
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
Co-authored-by: Roger Wang <hey@rogerw.io>
2026-07-26 21:00:27 -07:00
8de50e46d4 [Docs] Document NVFP4 GEMM kernel selection and Marlin weight-only fallback (#49376)
Signed-off-by: harjoth <harjoth.khara@gmail.com>
Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
2026-07-26 20:53:28 -07:00
Andreas KaratzasandGitHub da99ffcc13 [ROCm][CI] Keep native datasets cache off shared NFS (#49516)
Signed-off-by: Andreas Karatzas <Andreas.Karatzas@amd.com>
2026-07-26 22:31:24 -05:00
Walter Beller-MoralesGitHubmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
8040ef2426 [Frontend] expose stream_interval as req sampling param (#49754)
Signed-off-by: walterbm <walter.beller.morales@gmail.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-07-27 11:30:37 +08:00
bf4f633b4c [XPU] Enable QK Norm + RoPE fusion pass on XPU (#49394)
Signed-off-by: Chaojun Zhang <chaojun.zhang@intel.com>
Co-authored-by: Kunshang Ji <kunshang.ji@intel.com>
2026-07-27 03:22:29 +00:00
Andreas KaratzasandGitHub 854c33f380 [CI][ROCm] Keep global GPU memory cleanup opt-in (#49911)
Signed-off-by: Andreas Karatzas <akaratza@amd.com>
2026-07-26 22:19:55 -05:00
Andreas KaratzasandGitHub ac87549cbd [CI][ROCm] Reduce V1 attention test runtime (#49916)
Signed-off-by: Andreas Karatzas <akaratza@amd.com>
2026-07-27 11:12:45 +08:00
439f336212 [Core] Fix gpu<->cpu syncs in MRV2 mamba_hybrid.py (#49736)
Signed-off-by: Nick Hill <nickhill123@gmail.com>
Co-authored-by: Benjamin Chislett <bchislett@nvidia.com>
2026-07-27 02:41:27 +00:00
limewardandGitHub ffc4f08c8e [Core][KV-transfer] MoRIIO: heterogeneous TP<->DP prefill/decode read routing (#46116)
Signed-off-by: Edwin Lim <edwin.lim@mangoboost.io>
2026-07-27 02:10:00 +00:00
Nick HillandGitHub 50aa830482 [BugFix][MRV2] Don't create dummy requests longer than max_model_len (#49751)
Signed-off-by: Nick Hill <nickhill123@gmail.com>
2026-07-27 02:04:20 +00:00
f0553889c0 [Bugfix] Prevent NaN poisoning in xpu_mla_sparse for fully-masked index chunks (#48366)
Signed-off-by: Nick Iusiumbeli <nickuspro@gmail.com>
Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
Co-authored-by: Kunshang Ji <kunshang.ji@intel.com>
2026-07-27 09:06:07 +08:00
fdaa0d9e59 [ModelRunner V2] Support encoder-only attention (#49331)
Signed-off-by: Nick Hill <nickhill123@gmail.com>
Signed-off-by: Taneem Ibrahim <taneem.ibrahim@gmail.com>
Co-authored-by: Taneem Ibrahim <taneem.ibrahim@gmail.com>
2026-07-27 00:38:44 +00:00
0934b26790 [CI/Build] Refresh tags before building macOS wheel (#49901)
Signed-off-by: khluu <khluu000@gmail.com>
Co-authored-by: OpenAI Codex <codex@openai.com>
2026-07-26 13:07:32 -07:00
9e50e1037e [Bugfix][CuMem] Make KV-cache wake cleanup tag-safe (#49857)
Signed-off-by: aoshen02 <aoshen02@users.noreply.github.com>
Co-authored-by: aoshen02 <aoshen02@users.noreply.github.com>
2026-07-26 12:09:13 -07:00
Schwinn SaereesitthipitakandGitHub b5b61c622c [Core][Distributed] Add process-checkpoint lifecycle hooks for communicators (starting with Flashinfer) (#46877)
Signed-off-by: Schwinn Saereesitthipitak <schwinns@nvidia.com>
2026-07-26 14:47:50 -04:00
b68d7ef262 [Bugfix][KV Offload] Namespace auto cache dtype by effective dtype (#49438)
Signed-off-by: Jonguk Cheong <jdal3031@snu.ac.kr>
Co-authored-by: OpenAI Codex <codex@openai.com>
Co-authored-by: Or Ozeri <oro@il.ibm.com>
2026-07-26 20:59:09 +03:00
7154856f3d [Bugfix] Fix handling 5D KV cache in kv_postprocess_layout_on_receive (#47791)
Signed-off-by: Daniel Socek <daniel.socek@intel.com>
Co-authored-by: Kunshang Ji <kunshang.ji@intel.com>
2026-07-26 22:22:42 +08:00
3f1d40960f [KV Offload] Fix num_tokens_after_batch for different termination types (#49285)
Signed-off-by: Alex <jihuihuang@example.com>
Signed-off-by: Alex <jihui.huang@daocloud.io>
Signed-off-by: Alex <alex.tech.lab@outlook.com>
Signed-off-by: Alex <jihuihuang@users.noreply.github.com>
Co-authored-by: Claude <noreply@anthropic.com>
Co-authored-by: Or Ozeri <oro@il.ibm.com>
2026-07-26 16:22:58 +03:00
Taneem IbrahimandGitHub 0da6e7f3d6 [Bugfix] Reject contradictory custom-op directives (#49134)
Signed-off-by: Taneem Ibrahim <taneem.ibrahim@gmail.com>
2026-07-26 08:42:14 -04:00
5559679229 [Bugfix][KV Offload] Bound unaligned SWA loads by physical GPU blocks (#49052)
Signed-off-by: Colton Ottley <colton@ottleyengineering.com>
Co-authored-by: Colton Ottley <colton@ottleyengineering.com>
Co-authored-by: jasl <jasl9187@hotmail.com>
Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
Co-authored-by: Or Ozeri <oro@il.ibm.com>
2026-07-26 14:53:32 +03:00
da3a252fd1 [KVOffload][P2P] Generic P2P secondary tier: peer lookup and serving via ParentManager (#48021)
Signed-off-by: Liran Schour <lirans@il.ibm.com>
Signed-off-by: liranschour <liranschour@users.noreply.github.com>
Co-authored-by: Or Ozeri <or@ozery.com>
Co-authored-by: Or Ozeri <oro@il.ibm.com>
2026-07-26 11:45:47 +03:00
Guan-Ming ChiuandGitHub 21fd9e85a0 [Model] Support top_k and top_p sampling for DiffusionGemma (#45429)
Signed-off-by: Guan-Ming (Wesley) Chiu <105915352+guan404ming@users.noreply.github.com>
2026-07-26 08:39:25 +00:00
Guan-Ming ChiuandGitHub 8d28b48d01 [Perf] Isolate MM preprocessing on its own executor (#49524)
Signed-off-by: Guan-Ming (Wesley) Chiu <105915352+guan404ming@users.noreply.github.com>
2026-07-26 08:04:17 +00:00
Taneem IbrahimandGitHub 0164022c90 [CI] Fix speech correctness check rejecting improved WER (#49853)
Signed-off-by: Taneem Ibrahim <taneem.ibrahim@gmail.com>
2026-07-26 06:24:07 +00:00
30b0714031 [Perf] DeepSeek-OCR-2 TTFT Optimize (#49531)
Signed-off-by: RED <outofthewoods@qq.com>
Signed-off-by: Isotr0py <Isotr0py@outlook.com>
Co-authored-by: Isotr0py <Isotr0py@outlook.com>
2026-07-26 05:53:06 +00:00
Nils MattesonandGitHub 2e860de498 [Doc] Add compile cache volume example to the Docker deployment page (#49782)
Signed-off-by: Nils Matteson <nilsmatteson@icloud.com>
2026-07-26 05:24:01 +00:00
7eca0e1a64 [KV Offload] Deduplicate replicated MLA KV in the shared CPU region (#48906)
Signed-off-by: Change72 <changg@nvidia.com>
Co-authored-by: OpenAI Codex <noreply@openai.com>
2026-07-26 08:22:30 +03:00
7a29a3c54c [Bugfix][KV Offload] Namespace persistent cache by model runner (#49440)
Signed-off-by: Jonguk Cheong <jdal3031@snu.ac.kr>
Co-authored-by: Or Ozeri <oro@il.ibm.com>
2026-07-26 08:21:52 +03:00
Athrael SojuandGitHub 1240c74c0a [Bugfix] Respect declared attention contract for ColQwen3.5 retrievers (#49372)
Signed-off-by: Athrael Soju <athrael.soju@gmail.com>
2026-07-26 04:08:07 +00:00
48ebd6f2f1 [Bugfix][KVConnector] Disable cross-layer KV blocks for per-token-head quant (#49226)
Signed-off-by: Achyuthan Sivasankar <achyuthan.sivasankar@gmail.com>
Co-authored-by: Or Ozeri <oro@il.ibm.com>
2026-07-26 05:24:27 +03:00
liuzhenweiandGitHub b153ae6089 [XPU][CI] add heterogeneous TP UT (#49651)
Signed-off-by: zhenwei-intel <zhenwei.liu@intel.com>
2026-07-26 02:01:15 +00:00
Chang GuoandGitHub 7a6a5b3667 [CI] Compute speech WER directly with jiwer (#49773) 2026-07-25 20:51:21 -04:00
0111002323 [Kernel] TD operand loads for batched MoE GEMM (moe_mmk) on XPU (#46340)
Signed-off-by: oonyshch <xonyshch@gmail.com>
Co-authored-by: Kunshang Ji <kunshang.ji@intel.com>
2026-07-26 08:50:07 +08:00
d30b1ecd1b [Bugfix][KV Offloading] Defer request finalization until final store (#49671)
Signed-off-by: Rui Yin <2260891073@qq.com>
Co-authored-by: Or Ozeri <oro@il.ibm.com>
2026-07-25 20:58:09 +00:00
Taneem IbrahimandGitHub dbd80cc031 [UX] DCP Topology Validation (#49777)
Signed-off-by: Taneem Ibrahim <taneem.ibrahim@gmail.com>
2026-07-25 16:53:57 -04:00
70009fb934 [MM][CG] Support ViT CUDA Graph for Gemma-4 (#46837)
Signed-off-by: Anthony Su <xsuanthony@gmail.com>
Co-authored-by: Linkun Chen <github@lkchen.net>
2026-07-25 15:02:09 -05:00
Tyler Michael SmithandGitHub ee1d996367 [Build] Fix for DeepEP manylinux pidfd sycall usage (#49814) 2026-07-25 15:29:38 -04:00
Taneem IbrahimandGitHub 6b0103d1c9 [CI] Stabilize Pooling Rerank Equivalence Test (#49822) 2026-07-25 14:36:50 -04:00
9321aff536 [Bugfix] Wait for the linear bias before layerwise online processing (#49805)
Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
2026-07-25 17:56:32 +00:00
Harry MellorandGitHub 26d725c334 [Model] Add VaultGemma via Transformers modeling backend (#49803)
Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
2026-07-25 16:54:15 +00:00
Wentao YeandGitHub 7fe6d3c76b [Perf] Fix moe reduce_scatter perf regression by removing additional comm, 5% E2E throughput gain back. (#48763)
Signed-off-by: yewentao256 <zhyanwentao@126.com>
2026-07-25 16:36:19 +00:00
Harry MellorandGitHub 2e0da24150 Mergify message not on cancelled (#45117)
Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
2026-07-25 16:20:24 +00:00
+1 0b0bd2b5f6 [Feature] Add fault tolerance framework (simplified) for DP+EP external LB deployments (#44428)
Signed-off-by: fangyuchu <fangyuchu@qq.com>
Signed-off-by: a798347923 <2645302020@qq.com>
Signed-off-by: TianZhuo <2770730562@qq.com>
Signed-off-by: a798347923 <39047817+a798347923@users.noreply.github.com>
Signed-off-by: 205150940 <112750056+205150940@users.noreply.github.com>
Signed-off-by: w00689259 <wangzhuo66@huawei.com>
Signed-off-by: zWaNg3 <37772915+zWaNg3@users.noreply.github.com>
Signed-off-by: zWaNg3 <389750525@qq.com>
Signed-off-by: yzchang-plus <1078477584@qq.com>
Signed-off-by: Jade Zheng <zheng.shoujian@outlook.com>
Co-authored-by: zWaNg3 <37772915+zWaNg3@users.noreply.github.com>
Co-authored-by: a798347923 <2645302020@qq.com>
Co-authored-by: TianZhuo <2770730562@qq.com>
Co-authored-by: 205150940 <112750056+205150940@users.noreply.github.com>
Co-authored-by: a798347923 <39047817+a798347923@users.noreply.github.com>
Co-authored-by: w00689259 <wangzhuo66@huawei.com>
Co-authored-by: zWaNg3 <389750525@qq.com>
Co-authored-by: yzchang-plus <1078477584@qq.com>
Co-authored-by: Jade Zheng <zheng.shoujian@outlook.com>
2026-07-25 11:49:09 -04:00
Canlin GuoandGitHub 33ef67e9fb [BugFix] Increase the max supported duration for MOSS-TD (#49403)
Signed-off-by: Canlin Guo <canlinguosdu@gmail.com>
2026-07-25 14:48:12 +00:00
Harry MellorandGitHub 3e74c60b9c [Docs] Use gen-files for generated docs content (#49587)
Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
2026-07-25 14:43:21 +00:00
rongfu.lengandGitHub d1a8ba63d9 [Bugfix][MiniMax-M3] Fix token-major top-k buffer handling in Triton … (#49149)
Signed-off-by: rongfu.leng <lenronfu@gmail.com>
2026-07-25 06:45:27 -07:00
1423569ff5 [Bugfix][Tool Parser] Fix dropped streaming arguments in Jamba and InternLM2 parsers (#48852)
Signed-off-by: mosya415 <263250241+mosya415@users.noreply.github.com>
Co-authored-by: mosya415 <263250241+mosya415@users.noreply.github.com>
2026-07-25 09:34:43 -04:00
Harry MellorandGitHub 9a50464698 [CI] Stop flaky test from downloading model every time (#49800)
Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
2026-07-25 13:29:36 +00:00
Harshal JanjaniGitHubHarry Mellormergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
b9b6306ebe feat[vLLM × v5]: Add audio support for the Transformers backend (#39330)
Signed-off-by: Harshal Janjani <harshaljanjani@gmail.com>
Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
Co-authored-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-07-25 04:20:58 -07:00
ca0defa343 Make bare hugging_face imports forbidden (#49726)
Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
Co-authored-by: Andreas Karatzas <akaratza@amd.com>
2026-07-25 04:18:39 -07:00
0b1a8bb1f6 [Bugfix][CI] Fix stale Mooncake lookup expectation broken by a merge race (#49802)
Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
2026-07-25 03:10:26 -07:00
fe5145765f [Core] Keep attention backends eligible for text-only serving of prefix-LM models (#48796)
Signed-off-by: qtris123 <voquangtri2021@gmail.com>
Co-authored-by: Cyrus Leung <tlleungac@connect.ust.hk>
2026-07-25 03:06:14 -07:00
dbcc1cdd0a [Model] Remove Ouro (#49786)
Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
2026-07-25 02:50:08 -07:00
a82f1b388f [Perf][V1] Skip LRU hash-split in free_blocks when prefix caching is off (#48017)
Signed-off-by: Dobrzyniewicz, Agata <agata.dobrzyniewicz@intel.com>
Signed-off-by: Agata Dobrzyniewicz <160237065+adobrzyn@users.noreply.github.com>
Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
2026-07-25 09:06:20 +00:00
Johnny-LiouandGitHub 190be7dad2 [Docs] Fix confusing docstring indentation in nemotron_h.py (#49781)
Signed-off-by: Johnny-Liou <a897111@gmail.com>
2026-07-25 07:33:43 +00:00
94682b79f4 [multimodal] Make PyNvVideoCodec decoder concurrency configurable (#49753)
Signed-off-by: Brandon Pelfrey <bpelfrey@nvidia.com>
Co-authored-by: OpenAI Codex <noreply@openai.com>
2026-07-24 23:31:41 -07:00
646 changed files with 29175 additions and 44609 deletions
@@ -0,0 +1,26 @@
group: Benchmarks
depends_on:
- image-build-xpu
steps:
- label: Benchmarks CLI Test
key: benchmarks-cli-test
timeout_in_minutes: 40
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
working_dir: "."
env:
REGISTRY: "public.ecr.aws/q9t5s3a7"
REPO: "vllm-ci-test-repo"
VLLM_TEST_DEVICE: "xpu"
source_file_dependencies:
- vllm/
- tests/benchmarks/
commands:
- >-
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
'cd tests &&
pytest -v -s benchmarks/'
+76
View File
@@ -2,6 +2,44 @@ group: Engine Intel
depends_on:
- image-build-xpu
steps:
- label: Engine
key: engine
timeout_in_minutes: 40
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
working_dir: "."
env:
REGISTRY: "public.ecr.aws/q9t5s3a7"
REPO: "vllm-ci-test-repo"
VLLM_TEST_DEVICE: "xpu"
source_file_dependencies:
- vllm/compilation/
- vllm/config/
- vllm/engine/
- vllm/entrypoints/logger.py
- vllm/envs.py
- vllm/logger.py
- vllm/logging_utils/
- vllm/platforms/
- vllm/sequence.py
- vllm/triton_utils/
- vllm/utils/
- tests/engine
- tests/test_sequence
- tests/test_config
- tests/test_logger
- tests/test_vllm_port
- tests/test_jit_monitor.py
commands:
- >-
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
'cd tests &&
pytest -v -s engine/test_arg_utils.py test_sequence.py test_logger.py test_vllm_port.py test_jit_monitor.py'
- label: Engine (1 GPU)
timeout_in_minutes: 30
device: intel_gpu
@@ -23,3 +61,41 @@ steps:
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
'cd tests &&
pytest -v -s v1/engine --ignore v1/engine/test_preprocess_error_handling.py'
- label: V1 e2e (2 GPUs)
timeout_in_minutes: 30
device: intel_gpu
agent_tags:
label: production
gpu: 2+
mem: 16+
no_plugin: true
working_dir: "."
env:
REGISTRY: "public.ecr.aws/q9t5s3a7"
REPO: "vllm-ci-test-repo"
VLLM_TEST_DEVICE: "xpu"
source_file_dependencies:
- vllm/compilation/
- vllm/config/
- vllm/distributed/
- vllm/engine/
- vllm/envs.py
- vllm/forward_context.py
- vllm/inputs/
- vllm/logger.py
- vllm/logging_utils/
- vllm/model_executor/
- vllm/multimodal/
- vllm/platforms/
- vllm/sampling_params.py
- vllm/transformers_utils/
- vllm/triton_utils/
- vllm/utils/
- vllm/v1/
- tests/v1/e2e/spec_decode
commands:
- >-
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
'cd tests &&
pytest -v -s v1/e2e/spec_decode/test_spec_decode.py -k "tensor_parallelism"'
+29 -4
View File
@@ -125,13 +125,13 @@ steps:
pytest -v -s v1/kv_offload &&
pytest -v -s v1/kv_connector/unit/test_offloading_connector.py'
- label: NixlConnector PD accuracy (2 GPUs)
- label: NixlConnector PD accuracy (4 GPUs)
timeout_in_minutes: 60
num_devices: 2
num_devices: 4
device: intel_gpu
agent_tags:
label: production
gpu: 2+
gpu: 4+
mem: 16+
no_plugin: true
working_dir: "."
@@ -148,7 +148,10 @@ steps:
- >-
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
'cd tests &&
bash v1/kv_connector/nixl_integration/run_xpu_disagg_accuracy_test.sh'
bash v1/kv_connector/nixl_integration/run_xpu_disagg_accuracy_test.sh &&
PREFILLER_TP_SIZE=2 DECODER_TP_SIZE=1 bash v1/kv_connector/nixl_integration/run_xpu_disagg_accuracy_test.sh &&
PREFILLER_TP_SIZE=1 DECODER_TP_SIZE=2 bash v1/kv_connector/nixl_integration/run_xpu_disagg_accuracy_test.sh &&
PREFILLER_TP_SIZE=2 DECODER_TP_SIZE=2 bash v1/kv_connector/nixl_integration/run_xpu_disagg_accuracy_test.sh'
- label: Regression
key: regression
@@ -259,3 +262,25 @@ steps:
pytest -v -s detokenizer &&
pytest -v -s -m "not cpu_test" ./multimodal &&
pytest -v -s utils_ --ignore=utils_/test_mem_utils.py'
- label: Fusion Unit Tests
timeout_in_minutes: 30
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 16+
no_plugin: true
working_dir: "."
env:
REGISTRY: "public.ecr.aws/q9t5s3a7"
REPO: "vllm-ci-test-repo"
VLLM_TEST_DEVICE: "xpu"
source_file_dependencies:
- vllm/compilation/
- tests/compile/passes/test_qk_norm_rope_fusion.py
commands:
- >-
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
'cd tests &&
pytest -v -s compile/passes/test_qk_norm_rope_fusion.py'
@@ -0,0 +1,33 @@
group: Model Executor Intel
depends_on:
- image-build-xpu
steps:
- label: Model Executor (Intel)
key: model-executor-intel
timeout_in_minutes: 45
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 24+
no_plugin: true
working_dir: "."
env:
REGISTRY: "public.ecr.aws/q9t5s3a7"
REPO: "vllm-ci-test-repo"
VLLM_TEST_DEVICE: "xpu"
source_file_dependencies:
- vllm/engine/arg_utils.py
- vllm/config/model.py
- vllm/model_executor
- tests/model_executor
commands:
- >-
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
'apt-get update && apt-get install -y curl libsodium23 &&
pip3 install tensorizer==2.10.1 &&
pip3 install runai-model-streamer[s3,gcs,azure]\>=0.15.7 &&
export VLLM_WORKER_MULTIPROC_METHOD=spawn &&
export PYTHONFAULTHANDLER=1 &&
cd tests &&
pytest -v -s model_executor -m "not slow_test" --ignore="model_executor/layers/test_rocm_unquantized_gemm.py" --deselect="tests/model_executor/model_loader/test_reload.py::test_kv_scale_reload"'
@@ -8,7 +8,7 @@ steps:
agent_tags:
label: production
gpu: 2+
mem: 16+
mem: 24+
no_plugin: true
working_dir: "."
env:
@@ -28,7 +28,9 @@ steps:
'export VLLM_USE_V2_MODEL_RUNNER=1 &&
cd tests &&
pytest -v -s v1/engine/test_llm_engine.py -k "not test_engine_metrics" &&
pytest -v -s v1/e2e/general/test_context_length.py &&
ENFORCE_EAGER=1 pytest -v -s v1/e2e/general/test_async_scheduling.py -k "not ngram" &&
pytest -v -s entrypoints/llm/test_struct_output_generate.py -k "xgrammar and not speculative_config6 and not speculative_config7 and not speculative_config8 and not speculative_config0" &&
pytest -v -s v1/e2e/general/test_min_tokens.py'
- label: Model Runner V2 Examples (Intel)
@@ -60,3 +62,55 @@ steps:
python3 basic/offline_inference/generate.py --model facebook/opt-125m &&
python3 generate/multimodal/vision_language_offline.py --seed 0 &&
python3 features/automatic_prefix_caching/prefix_caching_offline.py'
- label: Model Runner V2 Distributed (2 GPUs)
timeout_in_minutes: 50
device: intel_gpu
agent_tags:
label: production
gpu: 2+
mem: 16+
no_plugin: true
working_dir: "."
env:
REGISTRY: "public.ecr.aws/q9t5s3a7"
REPO: "vllm-ci-test-repo"
VLLM_TEST_DEVICE: "xpu"
source_file_dependencies:
- vllm/v1/worker/gpu/
- vllm/v1/worker/gpu_worker.py
- tests/basic_correctness/test_basic_correctness.py
- tests/v1/distributed/test_async_llm_dp.py
- tests/v1/distributed/test_eagle_dp.py
commands:
- >-
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
'export VLLM_USE_V2_MODEL_RUNNER=1 &&
cd tests &&
TARGET_TEST_SUITE=L4 pytest -v -s basic_correctness/test_basic_correctness.py -m "distributed\(num_gpus=2\)" -k "not ray and not True"'
- label: Model Runner V2 Spec Decode
timeout_in_minutes: 50
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 24+
no_plugin: true
working_dir: "."
env:
REGISTRY: "public.ecr.aws/q9t5s3a7"
REPO: "vllm-ci-test-repo"
VLLM_TEST_DEVICE: "xpu"
source_file_dependencies:
- vllm/v1/worker/gpu/
- vllm/v1/worker/gpu_worker.py
- tests/v1/spec_decode/test_max_len.py
- tests/v1/spec_decode/test_rejection_sampler_utils.py
- tests/v1/e2e/spec_decode/test_spec_decode.py
commands:
- >-
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
'export VLLM_USE_V2_MODEL_RUNNER=1 &&
cd tests &&
pytest -v -s v1/spec_decode/test_synthetic_rejection_sampler_utils.py'
+29
View File
@@ -0,0 +1,29 @@
group: Samplers Intel
depends_on:
- image-build-xpu
steps:
- label: Samplers Test (FlashInfer)
key: samplers-test-flashinfer-intel
timeout_in_minutes: 40
device: intel_gpu
agent_tags:
label: production
gpu: 1+
mem: 24+
no_plugin: true
working_dir: "."
env:
REGISTRY: "public.ecr.aws/q9t5s3a7"
REPO: "vllm-ci-test-repo"
VLLM_TEST_DEVICE: "xpu"
source_file_dependencies:
- vllm/model_executor/layers
- vllm/sampling_metadata.py
- tests/samplers
- tests/conftest.py
- vllm/entrypoints/generate/beam_search
commands:
- >-
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
'cd tests &&
VLLM_USE_FLASHINFER_SAMPLER=1 pytest -v -s samplers'
+3
View File
@@ -7,6 +7,9 @@
set -euo pipefail
# The macmini queue uses persistent checkouts, so refresh tags for setuptools-scm.
git fetch --tags --force origin
# The Rust frontend build needs protoc.
if ! command -v protoc >/dev/null 2>&1; then
brew install protobuf
+19 -2
View File
@@ -387,6 +387,7 @@ initialize_native_environment() {
local job_id="${BUILDKITE_JOB_ID:-${BUILDKITE_PARALLEL_JOB:-local}}"
local job_id_suffix=""
local native_root=""
local hf_fstype=""
local hf_mount=""
if [[ "$(id -u)" -ne 0 ]]; then
@@ -405,11 +406,14 @@ initialize_native_environment() {
VLLM_CACHE_ROOT="${native_root}/cache/vllm"
XDG_CACHE_HOME="${native_root}/cache/xdg"
: "${HF_HOME:=/home/buildkite-agent/huggingface}"
# datasets uses POSIX locks that are unsupported by the shared HF NFS cache.
# Keep processed datasets job-local while retaining the persistent Hub cache.
HF_DATASETS_CACHE="${native_root}/cache/huggingface/datasets"
: "${HF_HUB_DOWNLOAD_TIMEOUT:=300}"
: "${HF_HUB_ETAG_TIMEOUT:=60}"
export TMPDIR VLLM_RPC_BASE_PATH
export TORCHINDUCTOR_CACHE_DIR TRITON_CACHE_DIR VLLM_CACHE_ROOT XDG_CACHE_HOME
export HF_HOME HF_HUB_DOWNLOAD_TIMEOUT HF_HUB_ETAG_TIMEOUT
export HF_HOME HF_DATASETS_CACHE HF_HUB_DOWNLOAD_TIMEOUT HF_HUB_ETAG_TIMEOUT
export PYTORCH_ROCM_ARCH=""
mkdir -p "${TMPDIR}" \
@@ -417,7 +421,8 @@ initialize_native_environment() {
"${TRITON_CACHE_DIR}" \
"${VLLM_CACHE_ROOT}" \
"${XDG_CACHE_HOME}" \
"${HF_HOME}" || return 1
"${HF_HOME}" \
"${HF_DATASETS_CACHE}" || return 1
echo "Native compile caches: VLLM_CACHE_ROOT=${VLLM_CACHE_ROOT} TORCHINDUCTOR_CACHE_DIR=${TORCHINDUCTOR_CACHE_DIR}"
@@ -432,6 +437,18 @@ initialize_native_environment() {
return 1
fi
fi
if command -v findmnt >/dev/null 2>&1; then
hf_fstype=$(findmnt -n -T "${HF_HOME}" -o FSTYPE 2>/dev/null || true)
fi
if [[ "${hf_fstype}" == nfs || "${hf_fstype}" == nfs4 ]]; then
# Keep hf-xet state local and avoid vectored writes on shared NFS.
export HF_XET_CACHE="${native_root}/cache/hf-xet"
export HF_XET_HIGH_PERFORMANCE=0
export HF_XET_RECONSTRUCTION_USE_VECTORED_WRITE=0
mkdir -p "${HF_XET_CACHE}" || return 1
echo "Configured hf-xet for shared ${hf_fstype} cache at ${HF_HOME}"
fi
}
run_native_preflight() {
+25 -37
View File
@@ -169,20 +169,6 @@ steps:
- pip install helion==1.1.0
- pytest -v -s kernels/helion/
- label: Kernels Mamba Test # TBD
timeout_in_minutes: 180
mirror_hardwares: [amdexperimental, amdproduction, amdgfx90anightly, amdmi250]
agent_pool: mi250_1
optional: true
working_dir: "/vllm-workspace/tests"
source_file_dependencies:
- csrc/mamba/
- tests/kernels/mamba
- vllm/model_executor/layers/mamba/ops
- vllm/platforms/rocm.py
commands:
- pytest -v -s kernels/mamba
#------------------------------------------------------ mi250 · models / basic -------------------------------------------------------#
- label: Basic Models Test (Other CPU) # TBD
@@ -364,22 +350,6 @@ steps:
commands:
- pytest -v -s v1/e2e/spec_decode -k "speculators or mtp_correctness"
- label: V1 attention (H100-MI250) # TBD
timeout_in_minutes: 180
mirror_hardwares: [amdexperimental, amdproduction, amdgfx90anightly, amdmi250]
agent_pool: mi250_1
working_dir: "/vllm-workspace/tests"
source_file_dependencies:
- vllm/config/attention.py
- vllm/model_executor/layers/attention
- vllm/v1/attention
- tests/v1/attention
- vllm/_aiter_ops.py
- vllm/envs.py
- vllm/platforms/rocm.py
commands:
- pytest -v -s v1/attention
- label: V1 others (CPU) # TBD
timeout_in_minutes: 180
mirror_hardwares: [amdexperimental, amdproduction, amdgfx90anightly, amdmi250]
@@ -1577,11 +1547,12 @@ steps:
commands:
- pytest -v -s kernels/attention --shard-id=$$BUILDKITE_PARALLEL_JOB --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT
- label: Kernels Core Operation Test # TBD
- label: Kernels Core Operation Test %N # TBD
timeout_in_minutes: 180
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi300]
dind: false
agent_pool: mi300_1
parallelism: 3
working_dir: "/vllm-workspace/tests"
source_file_dependencies:
- csrc/
@@ -1592,7 +1563,7 @@ steps:
- vllm/_aiter_ops.py
- vllm/platforms/rocm.py
commands:
- pytest -v -s kernels/core --ignore=kernels/core/test_minimax_reduce_rms.py kernels/test_concat_mla_q.py kernels/test_top_k_per_row.py
- pytest -v -s kernels/core --ignore=kernels/core/test_minimax_reduce_rms.py kernels/test_concat_mla_q.py kernels/test_top_k_per_row.py --shard-id=$$BUILDKITE_PARALLEL_JOB --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT
- label: Kernels KDA Test # TBD
timeout_in_minutes: 180
@@ -1610,6 +1581,21 @@ steps:
commands:
- pytest -v -s kernels/test_kda.py
- label: Kernels Mamba Test # TBD
timeout_in_minutes: 180
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi300]
dind: false
agent_pool: mi300_1
optional: true
working_dir: "/vllm-workspace/tests"
source_file_dependencies:
- csrc/mamba/
- tests/kernels/mamba
- vllm/model_executor/layers/mamba/ops
- vllm/platforms/rocm.py
commands:
- pytest -v -s kernels/mamba
- label: Kernels MoE Test %N # TBD
timeout_in_minutes: 180
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi300]
@@ -2161,7 +2147,7 @@ steps:
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi300]
dind: false
agent_pool: mi300_1
parallelism: 4
parallelism: 8
optional: true
working_dir: "/vllm-workspace/"
source_file_dependencies:
@@ -2694,11 +2680,12 @@ steps:
- export VLLM_WORKER_MULTIPROC_METHOD=spawn
- pytest -v -s v1/kv_connector/extract_hidden_states_integration
- label: V1 attention (H100-MI300) # TBD
- label: V1 attention (H100-MI300) %N # TBD
timeout_in_minutes: 180
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi300]
dind: false
agent_pool: mi300_1
parallelism: 2
optional: true
working_dir: "/vllm-workspace/tests"
source_file_dependencies:
@@ -2710,7 +2697,7 @@ steps:
- vllm/envs.py
- vllm/platforms/rocm.py
commands:
- pytest -v -s v1/attention
- pytest -v -s v1/attention --shard-id=$$BUILDKITE_PARALLEL_JOB --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT
- label: V1 Core + KV + Metrics # TBD
timeout_in_minutes: 180
@@ -3676,11 +3663,12 @@ steps:
#------------------------------------------------------------ mi355 · v1 -------------------------------------------------------------#
- label: V1 attention (B200-MI355) # TBD
- label: V1 attention (B200-MI355) %N # TBD
timeout_in_minutes: 180
mirror_hardwares: [amdexperimental, amdproduction, amdgfx950nightly, amdmi355]
dind: false
agent_pool: mi355_1
parallelism: 2
working_dir: "/vllm-workspace/tests"
source_file_dependencies:
- vllm/config/attention.py
@@ -3691,7 +3679,7 @@ steps:
- vllm/envs.py
- vllm/platforms/rocm.py
commands:
- pytest -v -s v1/attention
- pytest -v -s v1/attention --shard-id=$$BUILDKITE_PARALLEL_JOB --num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT
- label: V1 Core + KV + Metrics # TBD
timeout_in_minutes: 180
@@ -0,0 +1,26 @@
group: Fault Tolerance
depends_on:
- image-build
steps:
- label: Fault Tolerance E2E (2xH100)
key: fault-tolerance-e2e-2xh100
timeout_in_minutes: 35
device: h100
num_devices: 2
working_dir: "/vllm-workspace/tests"
source_file_dependencies:
- vllm/v1/fault_tolerance/
- vllm/v1/worker/sentinel/
- vllm/entrypoints/serve/fault_tolerance/
- vllm/distributed/elastic_ep/
- vllm/distributed/device_communicators/
- vllm/v1/engine/
- vllm/v1/worker/
- tests/v1/fault_tolerance/
- tests/v1/distributed/test_external_lb_dp.py
commands:
# Base image has no nixl; install it or has_nixl_ep() skips the tests.
- bash /vllm-workspace/.buildkite/scripts/install-kv-connectors.sh
# https://github.com/NVIDIA/nccl/issues/1838
- export NCCL_CUMEM_HOST_ENABLE=0
- pytest -v -s v1/fault_tolerance/test_fault_tolerance_e2e.py
+2 -2
View File
@@ -337,7 +337,7 @@ steps:
- label: LM Eval KV-Offload (2xH100)
key: kv-offload-medium
timeout_in_minutes: 30
timeout_in_minutes: 45
device: h100
num_devices: 2
source_file_dependencies:
@@ -347,7 +347,7 @@ steps:
- vllm/v1/simple_kv_offload/
- tests/evals/gsm8k/test_gsm8k_offloading.py
commands:
- pytest -s -v evals/gsm8k/test_gsm8k_offloading.py -k "qwen3.5-35b"
- pytest -s -v evals/gsm8k/test_gsm8k_offloading.py -k "qwen3.5-35b or deepseek-v2-lite"
- label: LM Eval KV-Offload (4xH100)
key: kv-offload-large
-1
View File
@@ -3,7 +3,6 @@
dist
vllm/*.so
vllm/vllm-rs
.git
# Byte-compiled / optimized / DLL files
__pycache__/
+26
View File
@@ -19,6 +19,7 @@ pull_request_rules:
description: Comment on PR when pre-commit check fails
conditions:
- check-failure=pre-commit
- -check-cancelled=pre-commit
- -closed
- -draft
- or:
@@ -232,6 +233,31 @@ pull_request_rules:
add:
- gpt-oss
- name: label-kimi
description: Automatically apply kimi label
conditions:
- label != stale
- or:
- files~=(?i)kimi
- files~=(?i)moonshot
- title~=(?i)(?:kimi|moonshot)
actions:
label:
add:
- kimi
- name: label-k3
description: Automatically apply k3 label (launch triage; retire after ramp-down)
conditions:
- label != stale
- or:
- files~=(?i)kimi[-_]?k3
- title~=(?i)(?:kimi[-\s]?k3|\bk3\b)
actions:
label:
add:
- k3
- name: label-nvidia
description: Automatically apply nvidia label
conditions:
+19
View File
@@ -130,6 +130,25 @@ jobs:
},
],
},
kimi: {
keywords: [
{ term: "Kimi", searchIn: "both" },
{ term: "Moonshot", searchIn: "both" },
],
substrings: [
{ term: "moonshotai/", searchIn: "both" },
{ term: "kimi", searchIn: "title" },
],
},
k3: {
keywords: [
{ term: "Kimi K3", searchIn: "both" },
{ term: "K3", searchIn: "title" },
],
substrings: [
{ term: "moonshotai/kimi-k3", searchIn: "both" },
],
},
quantization: {
keywords: [
{
-3
View File
@@ -173,9 +173,6 @@ venv.bak/
# mkdocs documentation
/site
docs/argparse
docs/examples/*
!docs/examples/README.md
# mypy
.mypy_cache/
+3
View File
@@ -3,6 +3,9 @@ MD007:
MD013: false
MD024:
siblings_only: true
MD025:
# Allow front matter title to be different from the first heading in the document.
front_matter_title: ""
MD031:
list_items: false
MD033: false
+1 -5
View File
@@ -4,7 +4,7 @@ default_install_hook_types:
default_stages:
- pre-commit # Run locally
- manual # Run in CI
exclude: 'vllm/third_party/.*|vllm/models/kimi_k3/nvidia/ops/third_party/.*|vllm/models/kimi_k3/amd/ops/third_party/.*'
exclude: 'vllm/third_party/.*'
repos:
- repo: https://github.com/astral-sh/ruff-pre-commit
rev: v0.14.0
@@ -260,10 +260,6 @@ repos:
files: ^docker/(Dockerfile|versions\.json)$
pass_filenames: false
additional_dependencies: [dockerfile-parse]
- id: attention-backend-docs
name: Check attention backend documentation is up to date
entry: python tools/pre_commit/generate_attention_backend_docs.py --check
language: python
- id: check-boolean-context-manager
name: Check for boolean ops in with-statements
entry: python tools/pre_commit/check_boolean_context_manager.py
+1 -48
View File
@@ -416,11 +416,8 @@ if(VLLM_GPU_LANG STREQUAL "CUDA" OR VLLM_GPU_LANG STREQUAL "HIP")
"csrc/libtorch_stable/mamba/selective_scan_fwd.cu"
"csrc/libtorch_stable/cache_kernels.cu"
"csrc/libtorch_stable/cache_kernels_fused.cu"
"csrc/libtorch_stable/custom_all_gather_reduce_scatter.cu"
"csrc/libtorch_stable/custom_all_gather_reduce_scatter_ops.cpp"
"csrc/libtorch_stable/custom_all_reduce.cu"
"csrc/libtorch_stable/fused_deepseek_v4_qnorm_rope_kv_insert_kernel.cu"
"csrc/libtorch_stable/fused_kimi_k3_mla_key_concat_kv_cache_kernel.cu")
"csrc/libtorch_stable/fused_deepseek_v4_qnorm_rope_kv_insert_kernel.cu")
if(VLLM_GPU_LANG STREQUAL "CUDA" AND
DEFINED CMAKE_CUDA_COMPILER_VERSION AND
@@ -1077,41 +1074,6 @@ if(VLLM_GPU_LANG STREQUAL "CUDA" OR VLLM_GPU_LANG STREQUAL "HIP")
set(MLA_ARCHS)
endif()
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 13.0)
cuda_archs_loose_intersection(FUSED_KDA_DECODE_ARCHS
"9.0a;10.0f;12.0f" "${CUDA_ARCHS}")
endif()
if(FUSED_KDA_DECODE_ARCHS)
set(FUSED_KDA_DECODE_SRC
"csrc/libtorch_stable/kimi_k3/fused_kda_decode_kernel.cu")
set_gencode_flags_for_srcs(
SRCS "${FUSED_KDA_DECODE_SRC}"
CUDA_ARCHS "${FUSED_KDA_DECODE_ARCHS}")
set_property(SOURCE ${FUSED_KDA_DECODE_SRC} APPEND PROPERTY
COMPILE_OPTIONS "$<$<COMPILE_LANGUAGE:CUDA>:--use_fast_math>")
list(APPEND VLLM_STABLE_EXT_SRC "${FUSED_KDA_DECODE_SRC}")
message(STATUS
"Building fused KDA decode for archs: ${FUSED_KDA_DECODE_ARCHS}")
endif()
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 13.0)
cuda_archs_loose_intersection(KIMI_K3_ATTN_RES_ARCHS
"10.0f" "${CUDA_ARCHS}")
endif()
if(KIMI_K3_ATTN_RES_ARCHS)
set(KIMI_K3_ATTN_RES_SRC
"csrc/libtorch_stable/kimi_k3/attn_res_kernel.cu")
set_gencode_flags_for_srcs(
SRCS "${KIMI_K3_ATTN_RES_SRC}"
CUDA_ARCHS "${KIMI_K3_ATTN_RES_ARCHS}")
set_property(SOURCE ${KIMI_K3_ATTN_RES_SRC} APPEND PROPERTY
COMPILE_OPTIONS
"$<$<COMPILE_LANGUAGE:CUDA>:--expt-relaxed-constexpr;--expt-extended-lambda;--use_fast_math>")
list(APPEND VLLM_STABLE_EXT_SRC "${KIMI_K3_ATTN_RES_SRC}")
message(STATUS
"Building Kimi K3 AttnRes for archs: ${KIMI_K3_ATTN_RES_ARCHS}")
endif()
# Hadacore kernels
cuda_archs_loose_intersection(HADACORE_ARCHS "8.0+PTX;9.0+PTX" "${CUDA_ARCHS}")
if(HADACORE_ARCHS)
@@ -1153,14 +1115,6 @@ if(VLLM_GPU_LANG STREQUAL "CUDA" OR VLLM_GPU_LANG STREQUAL "HIP")
target_compile_definitions(_C_stable_libtorch PRIVATE
VLLM_ENABLE_COOPERATIVE_TOPK=1)
endif()
if(FUSED_KDA_DECODE_ARCHS)
target_compile_definitions(_C_stable_libtorch PRIVATE
VLLM_ENABLE_FUSED_KDA_DECODE=1)
endif()
if(KIMI_K3_ATTN_RES_ARCHS)
target_compile_definitions(_C_stable_libtorch PRIVATE
VLLM_ENABLE_KIMI_K3_ATTN_RES=1)
endif()
# Needed by CUTLASS kernels
target_compile_definitions(_C_stable_libtorch PRIVATE
CUTLASS_ENABLE_DIRECT_CUDA_DRIVER_CALL=1)
@@ -1458,7 +1412,6 @@ if (VLLM_GPU_LANG STREQUAL "CUDA")
include(cmake/external_projects/deepgemm.cmake)
include(cmake/external_projects/fmha_sm100.cmake)
include(cmake/external_projects/flashmla.cmake)
include(cmake/external_projects/flashkda.cmake)
include(cmake/external_projects/qutlass.cmake)
include(cmake/external_projects/tml_fa4.cmake)
@@ -1358,6 +1358,10 @@ def main():
profile_memory=args.profile_memory,
warmup_ms=args.warmup_ms,
prefill_backend=pb,
kv_lora_rank=args.kv_lora_rank,
qk_nope_head_dim=args.qk_nope_head_dim,
qk_rope_head_dim=args.qk_rope_head_dim,
v_head_dim=args.v_head_dim,
)
result = run_benchmark(config)
@@ -0,0 +1,176 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import statistics
import torch
from tabulate import tabulate
from vllm.models.inkling.nvidia.ops import qkvr_prep
from vllm.utils.argparse_utils import FlexibleArgumentParser
def make_inputs(tokens: int, tp_size: int, is_local: bool):
torch.manual_seed(0)
num_q_heads = 64 // tp_size
num_kv_heads = (16 if is_local else 8) // tp_size
head_dim = 128
d_rel = 16
rel_extent = 512 if is_local else 1024
page_size = 16
num_blocks = (tokens + page_size - 1) // page_size
q_width = num_q_heads * head_dim
kv_width = num_kv_heads * head_dim
r_width = num_q_heads * d_rel
device = "cuda"
qkvr = torch.randn(
tokens,
q_width + 2 * kv_width + r_width,
device=device,
dtype=torch.bfloat16,
)
k_weight = torch.randn(kv_width, 4, device=device, dtype=torch.bfloat16)
v_weight = torch.randn_like(k_weight)
q_norm_weight = torch.randn(head_dim, device=device, dtype=torch.bfloat16)
k_norm_weight = torch.randn_like(q_norm_weight)
rel_proj = torch.randn(d_rel, rel_extent, device=device, dtype=torch.bfloat16)
conv_cache = torch.zeros(
num_blocks,
num_kv_heads,
page_size,
2 * head_dim,
device=device,
dtype=torch.bfloat16,
)
key_cache = torch.empty(
num_blocks,
page_size,
num_kv_heads,
head_dim,
device=device,
dtype=torch.bfloat16,
)
value_cache = torch.empty_like(key_cache)
positions = torch.arange(tokens, device=device, dtype=torch.int64)
block_table = torch.arange(num_blocks, device=device, dtype=torch.int32)[None]
seq_idx = torch.zeros(tokens, device=device, dtype=torch.int32)
slots = torch.arange(tokens, device=device, dtype=torch.int64)
query_start = torch.zeros(tokens, device=device, dtype=torch.int32)
log_scaling = None
if not is_local:
effective_n = (positions + 1).to(torch.float32)
log_scaling = 1.0 + 0.1 * torch.log(torch.clamp(effective_n / 128000, min=1.0))
return (
qkvr,
k_weight,
v_weight,
q_norm_weight,
k_norm_weight,
rel_proj,
1e-6,
num_q_heads,
num_kv_heads,
head_dim,
d_rel,
conv_cache,
key_cache,
value_cache,
positions,
block_table,
seq_idx,
slots,
query_start,
slots,
0,
head_dim,
page_size,
log_scaling,
)
def capture(implementation, inputs):
outputs = []
def run():
outputs[:] = implementation.fused_qkvr_prep(*inputs)
stream = torch.cuda.Stream()
stream.wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(stream):
for _ in range(3):
run()
torch.cuda.current_stream().wait_stream(stream)
torch.accelerator.synchronize()
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
run()
torch.accelerator.synchronize()
return graph, outputs
def time_graph(graph: torch.cuda.CUDAGraph, warmup: int, repeats: int) -> float:
for _ in range(warmup):
graph.replay()
torch.accelerator.synchronize()
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(repeats):
graph.replay()
end.record()
end.synchronize()
return start.elapsed_time(end) * 1000 / repeats
def benchmark(inputs, args) -> float:
graph, _ = capture(qkvr_prep, inputs)
return statistics.median(
time_graph(graph, args.warmup, args.repeats) for _ in range(args.trials)
)
@torch.inference_mode()
def main(args):
rows = []
for tp_size in args.tp_sizes:
for tokens in args.tokens:
for is_local in (True, False):
triton_us = benchmark(make_inputs(tokens, tp_size, is_local), args)
rows.append(
[
tp_size,
tokens,
"local" if is_local else "global",
triton_us,
]
)
print("Inkling QKVR prep (CUDA graph, median latency)")
print(
tabulate(
rows,
headers=[
"TP",
"tokens",
"scope",
"Triton (us)",
],
floatfmt=("d", "d", "", ".2f"),
)
)
if __name__ == "__main__":
parser = FlexibleArgumentParser()
parser.add_argument(
"--tokens",
type=int,
nargs="+",
default=[1 << power for power in range(15)],
)
parser.add_argument("--tp-sizes", type=int, nargs="+", default=[4, 8])
parser.add_argument("--warmup", type=int, default=20)
parser.add_argument("--repeats", type=int, default=200)
parser.add_argument("--trials", type=int, default=5)
main(parser.parse_args())
@@ -1,367 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Benchmark the Kimi-K3 latent MoE addmm against CuTe residual GEMM.
The benchmark covers ``BF16[M, 3584] @ BF16[7168, 3584].T + BF16[M, 7168]``
with FP32 accumulation and BF16 output. Both backends execute through CUDA
Graph replay. Weights and residuals rotate across buffers exceeding L2 so the
comparison models the full latent MoE projection-and-add path.
"""
from __future__ import annotations
import argparse
import dataclasses
import importlib.util
import json
import math
import statistics
from collections.abc import Callable, Sequence
from pathlib import Path
from typing import Any
import cutlass
import cutlass.cute as cute
import torch
from cuda.bindings import driver as cuda
from cuda.bindings.driver import CUstream
from quack.compile_utils import make_fake_tensor
N = 7168
K = 3584
@dataclasses.dataclass(frozen=True, slots=True)
class Config:
block_size: int
outputs_per_block: int
k_unroll: int
vector_width: int = 8
def parse_config(value: str) -> Config:
try:
parts = [int(part) for part in value.split(",")]
except ValueError as error:
raise argparse.ArgumentTypeError(
"config must be BLOCK,OUTPUTS,K_UNROLL[,VECTOR_WIDTH]"
) from error
if len(parts) == 3:
return Config(*parts)
if len(parts) == 4:
return Config(*parts)
raise argparse.ArgumentTypeError(
"config must be BLOCK,OUTPUTS,K_UNROLL[,VECTOR_WIDTH]"
)
def production_residual_config(m: int) -> Config | None:
"""The measured Latent-MoE residual config for M, from the K3 table."""
from vllm.models.kimi_k3.nvidia.low_latency_gemm import KIMI_K3_PROJECTIONS
spec = KIMI_K3_PROJECTIONS.get((N, K))
config = spec.residual_config(m) if spec is not None else None
if config is None:
return None
return Config(
config.block_size,
config.outputs_per_block,
config.k_unroll,
config.vector_width,
)
def candidate_configs(mode: str, selected: Config | None, m: int) -> list[Config]:
if mode == "selected":
if selected is not None:
return [selected]
# No explicit --config: fall back to the production table for this M.
config = production_residual_config(m)
return [config] if config is not None else []
if mode == "baseline":
return [Config(224, 4, 2)]
return [
Config(block_size, outputs_per_block, k_unroll, vector_width)
for vector_width in (4, 8)
for block_size in (32, 64, 128, 224, 448)
if block_size % 32 == 0 and K % (block_size * vector_width) == 0
for outputs_per_block in (1, 2, 4, 7, 8)
if N % outputs_per_block == 0
for k_unroll in (1, 2, 4)
]
def load_kernel_class(path: Path):
spec = importlib.util.spec_from_file_location("cute_skinny_device", path)
if spec is None or spec.loader is None:
raise RuntimeError(f"cannot load CuTe kernel from {path}")
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module.CuteSkinnyGemm
def stream() -> CUstream:
return CUstream(torch.cuda.current_stream().cuda_stream)
def compile_kernel(kernel_class, m: int, config: Config, max_registers: int):
element_type = cutlass.BFloat16
n = cute.sym_int(divisibility=config.outputs_per_block)
k = cute.sym_int(divisibility=config.block_size * config.vector_width)
a = make_fake_tensor(element_type, (m, k), divisibility=config.vector_width)
b = make_fake_tensor(element_type, (n, k), divisibility=config.vector_width)
residual = make_fake_tensor(element_type, (m, n), divisibility=1)
c = make_fake_tensor(element_type, (m, n), divisibility=1)
kernel = kernel_class(
element_type=element_type,
num_rows=m,
block_size=config.block_size,
outputs_per_block=config.outputs_per_block,
vector_width=config.vector_width,
k_unroll=config.k_unroll,
has_residual=True,
use_pdl=True,
)
return cute.compile(
kernel,
a,
b,
residual,
c,
stream(),
options=(
"--enable-tvm-ffi --keep-cubin "
f"--ptxas-options -maxrregcount={max_registers} "
"--ptxas-options -lineinfo"
),
)
def resource_usage(compiled) -> dict[str, Any]:
executor = getattr(compiled, "_default_executor", None)
context = getattr(executor, "exec_context", None)
functions = getattr(context, "kernel_functions", None)
if not functions:
return {"resource_metrics_available": False}
def attribute(name, function) -> int:
error, value = cuda.cuFuncGetAttribute(name, function)
if error != cuda.CUresult.CUDA_SUCCESS:
raise RuntimeError(f"cuFuncGetAttribute failed with {error}")
return int(value)
registers = [
attribute(cuda.CUfunction_attribute.CU_FUNC_ATTRIBUTE_NUM_REGS, function)
for function in functions
]
local_bytes = [
attribute(
cuda.CUfunction_attribute.CU_FUNC_ATTRIBUTE_LOCAL_SIZE_BYTES,
function,
)
for function in functions
]
return {
"resource_metrics_available": True,
"registers_per_thread": max(registers, default=0),
"spill_bytes": max(local_bytes, default=0),
}
def rotating_buffer_count(m: int, multiplier: float, limit: int) -> int:
properties = torch.cuda.get_device_properties(0)
bytes_per_pair = (N * K + m * N) * 2
target = math.ceil(multiplier * properties.L2_cache_size)
return max(2, min(limit, math.ceil(target / bytes_per_pair)))
def graph_samples(
launch: Callable[[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor], None],
activation: torch.Tensor,
weights: Sequence[torch.Tensor],
residuals: Sequence[torch.Tensor],
repeats: int,
replays: int,
) -> tuple[list[float], list[torch.Tensor]]:
outputs = [torch.empty_like(residual) for residual in residuals]
for weight, residual, output in zip(weights, residuals, outputs):
launch(activation, weight, residual, output)
torch.cuda.synchronize()
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
for weight, residual, output in zip(weights, residuals, outputs):
launch(activation, weight, residual, output)
for _ in range(20):
graph.replay()
torch.cuda.synchronize()
samples = []
for _ in range(repeats):
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(replays):
graph.replay()
end.record()
end.synchronize()
samples.append(start.elapsed_time(end) * 1000.0 / (replays * len(weights)))
return samples, outputs
def summarize(samples: Sequence[float]) -> dict[str, Any]:
ordered = sorted(samples)
def percentile(fraction: float) -> float:
position = fraction * (len(ordered) - 1)
lower = math.floor(position)
upper = math.ceil(position)
if lower == upper:
return ordered[lower]
weight = position - lower
return ordered[lower] * (1.0 - weight) + ordered[upper] * weight
mean = statistics.mean(samples)
return {
"median_us": statistics.median(samples),
"p10_us": percentile(0.1),
"p90_us": percentile(0.9),
"mean_us": mean,
"cv_pct": statistics.pstdev(samples) / mean * 100.0,
"samples_us": list(samples),
}
def correctness(
output: torch.Tensor,
activation: torch.Tensor,
weight: torch.Tensor,
residual: torch.Tensor,
) -> dict[str, Any]:
actual = output.float()
reference = activation.float() @ weight.float().t() + residual.float()
error = (actual - reference).abs()
scaled_error = error / (reference.abs() + 1.0)
cosine = torch.nn.functional.cosine_similarity(
actual.flatten(), reference.flatten(), dim=0
).item()
return {
"valid": cosine > 0.999,
"cosine": cosine,
"max_abs_error": error.max().item(),
"max_scaled_error": scaled_error.max().item(),
}
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--kernel", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument(
"--mode", choices=("baseline", "sweep", "selected"), default="baseline"
)
parser.add_argument("--config", type=parse_config)
parser.add_argument("--m", type=int, action="append")
parser.add_argument("--config-shard", type=int, default=0)
parser.add_argument("--num-config-shards", type=int, default=1)
parser.add_argument("--repeats", type=int, default=21)
parser.add_argument("--replays", type=int, default=200)
parser.add_argument("--cache-multiplier", type=float, default=3.0)
parser.add_argument("--max-buffers", type=int, default=32)
parser.add_argument("--max-registers", type=int, default=64)
args = parser.parse_args()
token_counts = args.m or list(range(1, 17))
if any(not 1 <= m <= 16 for m in token_counts):
raise ValueError("expected 1 <= M <= 16")
if not 0 <= args.config_shard < args.num_config_shards:
raise ValueError("config shard must be in [0, num_config_shards)")
torch.cuda.set_device(0)
if torch.cuda.get_device_capability() != (10, 3):
raise RuntimeError("this benchmark requires SM103")
kernel_class = load_kernel_class(args.kernel)
properties = torch.cuda.get_device_properties(0)
metadata = {
"device": properties.name,
"compute_capability": list(torch.cuda.get_device_capability()),
"torch_version": torch.__version__,
"cuda_version": torch.version.cuda,
}
args.output.parent.mkdir(parents=True, exist_ok=True)
with args.output.open("w", encoding="utf-8") as output_file:
for m in token_counts:
configs = candidate_configs(args.mode, args.config, m)
torch.manual_seed(20260722 + m)
count = rotating_buffer_count(m, args.cache_multiplier, args.max_buffers)
activation = torch.randn((m, K), device="cuda", dtype=torch.bfloat16)
weights = [
torch.randn((N, K), device="cuda", dtype=torch.bfloat16)
for _ in range(count)
]
residuals = [
torch.randn((m, N), device="cuda", dtype=torch.bfloat16)
for _ in range(count)
]
candidates: list[tuple[str, Config | None]] = [("cublas_addmm", None)]
candidates.extend(
("cute_residual", config)
for index, config in enumerate(configs)
if index % args.num_config_shards == args.config_shard
)
for backend, config in candidates:
row: dict[str, Any] = {
"m": m,
"n": N,
"k": K,
"backend": backend,
"mode": args.mode,
"config": dataclasses.asdict(config) if config else {},
"num_buffers": count,
"cache_multiplier": args.cache_multiplier,
**metadata,
}
try:
if backend == "cublas_addmm":
launch = lambda a, b, residual, c: torch.addmm(
residual, a, b.t(), out=c
)
else:
if config is None:
raise AssertionError("missing CuTe config")
compiled = compile_kernel(
kernel_class, m, config, args.max_registers
)
launch = lambda a, b, residual, c, fn=compiled: fn(
a, b, residual, c, stream()
)
row.update(resource_usage(compiled))
samples, outputs = graph_samples(
launch,
activation,
weights,
residuals,
args.repeats,
args.replays,
)
row.update(
correctness(outputs[0], activation, weights[0], residuals[0])
)
row.update(summarize(samples))
except Exception as error: # noqa: BLE001
row.update(
{
"valid": False,
"error": f"{type(error).__name__}: {error}",
}
)
output_file.write(json.dumps(row, sort_keys=True) + "\n")
output_file.flush()
print(json.dumps(row, sort_keys=True), flush=True)
del activation, weights, residuals
torch.cuda.empty_cache()
if __name__ == "__main__":
main()
@@ -1,806 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Benchmark the Kimi K3 latent-MoE tail and its up-projection kernels.
The ``up-projection`` subcommand isolates the TP-local dynamic and static-M
skinny GEMMs. It rotates weights through a working set larger than L2 to model
successive model layers.
The ``whole-tail`` subcommand measures the distributed operator. Its reference
path includes two AllReduces, RMSNorm, the replicated up-projection, and the
final add. CUDA-event samples report the slowest rank so cross-rank skew is
included.
Examples:
.. code-block:: console
.venv/bin/python \
benchmarks/kernels/benchmark_kimi_k3_latent_moe_tail.py up-projection
torchrun --nproc-per-node=8 \
benchmarks/kernels/benchmark_kimi_k3_latent_moe_tail.py whole-tail
For multi-node runs, launch one ``torchrun`` agent per node and use a shared
rendezvous endpoint.
"""
from __future__ import annotations
import argparse
import json
import math
import os
import statistics
from collections.abc import Callable, Sequence
from dataclasses import asdict
from pathlib import Path
from typing import Any
import cutlass
import cutlass.utils as utils
import torch
import torch.distributed as dist
import torch.nn.functional as F
from cuda.bindings import driver as cuda
from vllm.distributed import get_tp_group
from vllm.distributed.parallel_state import (
init_distributed_environment,
initialize_model_parallel,
set_custom_all_reduce,
)
from vllm.model_executor.warmup.cutedsl_warmup import cutedsl_warmup
from vllm.models.kimi_k3.nvidia.ops import latent_moe_tail
from vllm.models.kimi_k3.nvidia.ops.cute_dsl.latent_moe_tail import (
fused_add_multicast_gemm,
fused_add_multicast_skinny_gemm,
)
HIDDEN_SIZE = 7168
LATENT_SIZE = 3584
RMS_EPS = 0.1
MAX_NUM_TOKENS = 16
MMA_TILER_MN = (64, 32)
CLUSTER_SHAPE_MN = (1, 8)
B_PRIME_STAGES = 2
def parse_up_projection_config(
value: str,
) -> fused_add_multicast_skinny_gemm.SkinnyConfig:
try:
values = [int(part) for part in value.split(",")]
except ValueError as error:
raise argparse.ArgumentTypeError(
"config must be BLOCK,OUTPUTS,K_UNROLL[,VECTOR_WIDTH[,PREFETCH_B]]"
) from error
if len(values) in (3, 4):
return fused_add_multicast_skinny_gemm.SkinnyConfig(*values)
if len(values) == 5 and values[4] in (0, 1):
return fused_add_multicast_skinny_gemm.SkinnyConfig(
*values[:4],
prefetch_b_before_pdl=bool(values[4]),
)
raise argparse.ArgumentTypeError(
"config must be BLOCK,OUTPUTS,K_UNROLL"
"[,VECTOR_WIDTH[,PREFETCH_B]], where PREFETCH_B is 0 or 1"
)
def parse_tail_skinny_config(
value: str,
) -> tuple[int, fused_add_multicast_skinny_gemm.SkinnyConfig]:
try:
values = [int(part) for part in value.split(",")]
except ValueError as error:
raise argparse.ArgumentTypeError(
"config must be M,BLOCK,OUTPUTS,K_UNROLL[,VECTOR_WIDTH[,PREFETCH_B]]"
) from error
if len(values) == 4:
num_tokens, *config = values
return num_tokens, fused_add_multicast_skinny_gemm.SkinnyConfig(*config)
if len(values) == 5:
num_tokens, *config = values
return num_tokens, fused_add_multicast_skinny_gemm.SkinnyConfig(*config)
if len(values) == 6 and values[5] in (0, 1):
num_tokens, block, outputs, unroll, vector_width, prefetch = values
return num_tokens, fused_add_multicast_skinny_gemm.SkinnyConfig(
block,
outputs,
unroll,
vector_width,
bool(prefetch),
)
raise argparse.ArgumentTypeError(
"config must be M,BLOCK,OUTPUTS,K_UNROLL"
"[,VECTOR_WIDTH[,PREFETCH_B]], where PREFETCH_B is 0 or 1"
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
subparsers = parser.add_subparsers(dest="scope", required=True)
up_projection = subparsers.add_parser(
"up-projection",
help="Benchmark the isolated TP-local up-projection kernels.",
)
up_projection.add_argument(
"--backend",
choices=("dynamic", "skinny", "both"),
default="both",
)
up_projection.add_argument("--tp-size", type=int, default=16)
up_projection.add_argument(
"--num-tokens",
type=int,
nargs="+",
default=[*range(1, 9), 16],
)
up_projection.add_argument(
"--skinny-config",
type=parse_up_projection_config,
action="append",
help="Benchmark a static-M config for every selected token count.",
)
up_projection.add_argument("--cache-multiplier", type=float, default=2.0)
up_projection.add_argument("--max-weights", type=int, default=64)
up_projection.add_argument("--warmup-replays", type=int, default=10)
up_projection.add_argument("--samples", type=int, default=31)
up_projection.add_argument("--output", type=Path)
whole_tail = subparsers.add_parser(
"whole-tail",
help="Benchmark the distributed latent-MoE tail operator.",
)
whole_tail.add_argument(
"--backend",
choices=("reference", "fused", "both"),
default="both",
)
whole_tail.add_argument(
"--num-tokens",
type=int,
nargs="+",
default=[1, 5, 8, 16],
)
whole_tail.add_argument("--warmup-replays", type=int, default=20)
whole_tail.add_argument("--samples", type=int, default=51)
whole_tail.add_argument(
"--skinny-max-num-tokens",
type=int,
nargs="+",
help="Override the fused operator's static-M cutoff; use 0 for dynamic-only.",
)
whole_tail.add_argument(
"--skinny-config",
type=parse_tail_skinny_config,
action="append",
help="Override one static-M config for tuning.",
)
whole_tail.add_argument("--output", type=Path)
return parser.parse_args()
def percentile(samples: Sequence[float], fraction: float) -> float:
ordered = sorted(samples)
position = fraction * (len(ordered) - 1)
lower = math.floor(position)
upper = math.ceil(position)
if lower == upper:
return ordered[lower]
upper_weight = position - lower
return ordered[lower] * (1.0 - upper_weight) + ordered[upper] * upper_weight
def summarize(samples_us: Sequence[float]) -> dict[str, Any]:
mean_us = statistics.mean(samples_us)
return {
"median_us": statistics.median(samples_us),
"p10_us": percentile(samples_us, 0.1),
"p90_us": percentile(samples_us, 0.9),
"mean_us": mean_us,
"cv_pct": statistics.pstdev(samples_us) / mean_us * 100.0,
"samples_us": list(samples_us),
}
def rotating_weight_count(
shard_size: int,
cache_multiplier: float,
limit: int,
) -> int:
properties = torch.cuda.get_device_properties(
torch.accelerator.current_device_index()
)
weight_bytes = shard_size * LATENT_SIZE * 2
target_bytes = math.ceil(properties.L2_cache_size * cache_multiplier)
return max(2, min(limit, math.ceil(target_bytes / weight_bytes)))
def capture_up_projection_graph(
launches: Sequence[Callable[[], None]],
) -> torch.cuda.CUDAGraph:
for launch in launches:
launch()
torch.accelerator.synchronize()
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
for launch in launches:
launch()
torch.accelerator.synchronize()
return graph
def benchmark_up_projection_graph(
graph: torch.cuda.CUDAGraph,
*,
operations_per_replay: int,
warmup_replays: int,
samples: int,
) -> dict[str, Any]:
for _ in range(warmup_replays):
graph.replay()
torch.accelerator.synchronize()
samples_us = []
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
for _ in range(samples):
start.record()
graph.replay()
end.record()
end.synchronize()
samples_us.append(start.elapsed_time(end) * 1000.0 / operations_per_replay)
return summarize(samples_us)
class DynamicKernel:
def __init__(
self,
shard_size: int,
mailbox: torch.Tensor,
shared_shard: torch.Tensor,
) -> None:
self.shard_size = shard_size
self.mailbox = mailbox
self.mailbox_c = fused_add_multicast_gemm._as_cute(mailbox)
compile_latent = torch.empty(
(1, MAX_NUM_TOKENS, LATENT_SIZE),
dtype=torch.bfloat16,
device=mailbox.device,
)
compile_weight = torch.empty(
(1, shard_size, LATENT_SIZE),
dtype=torch.bfloat16,
device=mailbox.device,
)
cluster_size = math.prod(CLUSTER_SHAPE_MN)
max_active_clusters = utils.HardwareInfo().get_max_active_clusters(cluster_size)
self.compiled = fused_add_multicast_gemm.compile_kernel(
(MAX_NUM_TOKENS, shard_size, LATENT_SIZE, 1),
fused_add_multicast_gemm._as_cute(
compile_latent,
dynamic_m=True,
),
fused_add_multicast_gemm._as_cute(compile_weight),
self.mailbox_c,
fused_add_multicast_gemm._as_cute(shared_shard),
HIDDEN_SIZE,
shard_size,
MMA_TILER_MN,
CLUSTER_SHAPE_MN,
max_active_clusters,
B_PRIME_STAGES,
)
def launch(
self,
latent: torch.Tensor,
weight: torch.Tensor,
shared_shard: torch.Tensor,
) -> None:
stream = cuda.CUstream(torch.cuda.current_stream().cuda_stream)
self.compiled(
fused_add_multicast_gemm._as_cute(
latent.unsqueeze(0),
dynamic_m=True,
),
fused_add_multicast_gemm._as_cute(weight.unsqueeze(0)),
self.mailbox_c,
fused_add_multicast_gemm._as_cute(shared_shard),
cutlass.Int64(latent.shape[0]),
cutlass.Int64(self.mailbox.data_ptr()),
stream,
)
class SkinnyKernel:
def __init__(
self,
num_tokens: int,
shard_size: int,
config: fused_add_multicast_skinny_gemm.SkinnyConfig,
) -> None:
self.compiled = fused_add_multicast_skinny_gemm.compile_kernel(
num_rows=num_tokens,
latent_dim=LATENT_SIZE,
hidden_dim=HIDDEN_SIZE,
shard_dim=shard_size,
config=config,
)
def launch(
self,
latent: torch.Tensor,
weight: torch.Tensor,
shared_shard: torch.Tensor,
mailbox: torch.Tensor,
) -> None:
self.compiled(
fused_add_multicast_skinny_gemm._as_cute(latent),
fused_add_multicast_skinny_gemm._as_cute(weight),
fused_add_multicast_skinny_gemm._as_cute(shared_shard),
cutlass.Int64(mailbox.data_ptr()),
cuda.CUstream(torch.cuda.current_stream().cuda_stream),
)
def check_up_projection_output(
actual: torch.Tensor,
latent: torch.Tensor,
weight: torch.Tensor,
shared_shard: torch.Tensor,
) -> None:
gemm = F.linear(latent.float(), weight.float()).to(torch.bfloat16)
expected = (gemm.float() + shared_shard.float()).to(torch.bfloat16)
torch.testing.assert_close(actual, expected, atol=8e-2, rtol=3e-2)
def make_up_projection_launches(
launch: Callable[[torch.Tensor, torch.Tensor, torch.Tensor], None],
latent: torch.Tensor,
weights: Sequence[torch.Tensor],
shared_shard: torch.Tensor,
) -> list[Callable[[], None]]:
return [
lambda weight=weight: launch(latent, weight, shared_shard) for weight in weights
]
def benchmark_up_projection(args: argparse.Namespace) -> None:
if args.tp_size <= 0 or HIDDEN_SIZE % args.tp_size:
raise ValueError("TP size must be positive and divide the hidden size")
if any(not 1 <= num_tokens <= MAX_NUM_TOKENS for num_tokens in args.num_tokens):
raise ValueError("--num-tokens values must be in [1, 16]")
if args.cache_multiplier <= 0 or args.max_weights <= 0:
raise ValueError("cache multiplier and max weights must be positive")
if args.warmup_replays < 0 or args.samples <= 0:
raise ValueError("warmup replays must be nonnegative and samples positive")
torch.accelerator.set_device_index(0)
device = torch.device("cuda", 0)
if torch.cuda.get_device_capability(device)[0] != 10:
raise RuntimeError("Kimi K3 latent-MoE tail requires SM100")
shard_size = HIDDEN_SIZE // args.tp_size
weight_count = rotating_weight_count(
shard_size,
args.cache_multiplier,
args.max_weights,
)
torch.manual_seed(20260726)
weights = [
torch.randn(
(shard_size, LATENT_SIZE),
dtype=torch.bfloat16,
device=device,
)
/ LATENT_SIZE**0.5
for _ in range(weight_count)
]
mailbox = torch.empty(
(1, MAX_NUM_TOKENS, HIDDEN_SIZE),
dtype=torch.bfloat16,
device=device,
)
shared = torch.randn(
(MAX_NUM_TOKENS, HIDDEN_SIZE),
dtype=torch.bfloat16,
device=device,
)
shared_shard = shared[:, :shard_size]
use_dynamic = args.backend in ("dynamic", "both")
use_skinny = args.backend in ("skinny", "both")
dynamic_kernel = (
DynamicKernel(shard_size, mailbox, shared_shard) if use_dynamic else None
)
results = []
for num_tokens in args.num_tokens:
latent = torch.randn(
(num_tokens, LATENT_SIZE),
dtype=torch.bfloat16,
device=device,
)
result: dict[str, Any] = {"num_tokens": num_tokens}
if dynamic_kernel is not None:
launches = make_up_projection_launches(
dynamic_kernel.launch,
latent,
weights,
shared_shard,
)
graph = capture_up_projection_graph(launches)
result["dynamic"] = benchmark_up_projection_graph(
graph,
operations_per_replay=len(launches),
warmup_replays=args.warmup_replays,
samples=args.samples,
)
check_up_projection_output(
mailbox[0, :num_tokens, :shard_size],
latent,
weights[-1],
shared_shard[:num_tokens],
)
if use_skinny:
configs = args.skinny_config or [
fused_add_multicast_skinny_gemm.config_for_m(
num_tokens,
shard_size,
)
]
skinny_results = []
for config in configs:
skinny_kernel = SkinnyKernel(num_tokens, shard_size, config)
def launch_skinny(
latent: torch.Tensor,
weight: torch.Tensor,
shared_shard: torch.Tensor,
*,
skinny_kernel: SkinnyKernel = skinny_kernel,
num_tokens: int = num_tokens,
) -> None:
skinny_kernel.launch(
latent,
weight,
shared_shard[:num_tokens],
mailbox,
)
launches = make_up_projection_launches(
launch_skinny,
latent,
weights,
shared_shard,
)
graph = capture_up_projection_graph(launches)
timing = benchmark_up_projection_graph(
graph,
operations_per_replay=len(launches),
warmup_replays=args.warmup_replays,
samples=args.samples,
)
check_up_projection_output(
mailbox[0, :num_tokens, :shard_size],
latent,
weights[-1],
shared_shard[:num_tokens],
)
skinny_results.append(
{
"config": asdict(config),
**timing,
}
)
result["skinny"] = skinny_results
results.append(result)
properties = torch.cuda.get_device_properties(device)
report = {
"scope": "up-projection",
"device": properties.name,
"compute_capability": list(torch.cuda.get_device_capability(device)),
"tp_size": args.tp_size,
"shard_size": shard_size,
"weight_count": weight_count,
"cache_multiplier": args.cache_multiplier,
"warmup_replays": args.warmup_replays,
"samples": args.samples,
"results": results,
}
rendered = json.dumps(report, indent=2)
print(rendered, flush=True)
if args.output is not None:
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(rendered + "\n", encoding="utf-8")
def capture_tail_graph(
operation: Callable[[], torch.Tensor],
cpu_group: dist.ProcessGroup,
) -> tuple[torch.cuda.CUDAGraph, torch.Tensor]:
for _ in range(3):
dist.barrier(group=cpu_group)
output = operation()
torch.accelerator.synchronize()
dist.barrier(group=cpu_group)
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
output = operation()
torch.accelerator.synchronize()
return graph, output
def benchmark_tail_graph(
graph: torch.cuda.CUDAGraph,
*,
warmup_replays: int,
samples: int,
device_group: dist.ProcessGroup,
cpu_group: dist.ProcessGroup,
) -> dict[str, Any]:
for _ in range(warmup_replays):
graph.replay()
torch.accelerator.synchronize()
dist.barrier(group=cpu_group)
starts = [torch.cuda.Event(enable_timing=True) for _ in range(samples + 1)]
ends = [torch.cuda.Event(enable_timing=True) for _ in range(samples + 1)]
for start, end in zip(starts, ends):
start.record()
graph.replay()
end.record()
torch.accelerator.synchronize()
samples_us = torch.tensor(
[start.elapsed_time(end) * 1000.0 for start, end in zip(starts, ends)],
dtype=torch.float64,
device=torch.accelerator.current_device_index(),
)
dist.all_reduce(samples_us, op=dist.ReduceOp.MAX, group=device_group)
return summarize(samples_us[1:].tolist())
def make_inputs(
num_tokens: int,
rank: int,
device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor]:
torch.manual_seed(20260726 + 100 * num_tokens + rank)
routed = torch.randn(
(num_tokens, LATENT_SIZE),
dtype=torch.bfloat16,
device=device,
).mul_(0.01)
shared = torch.randn(
(num_tokens, HIDDEN_SIZE),
dtype=torch.bfloat16,
device=device,
)
return routed, shared
def make_reference(
routed: torch.Tensor,
shared: torch.Tensor,
rms_weight: torch.Tensor,
up_weight: torch.Tensor,
device_group: dist.ProcessGroup,
) -> Callable[[], torch.Tensor]:
routed_workspace = torch.empty_like(routed)
shared_workspace = torch.empty_like(shared)
def reference() -> torch.Tensor:
routed_workspace.copy_(routed)
dist.all_reduce(routed_workspace, group=device_group)
normalized = F.rms_norm(
routed_workspace,
(LATENT_SIZE,),
rms_weight,
RMS_EPS,
)
projected = F.linear(normalized, up_weight)
shared_workspace.copy_(shared)
dist.all_reduce(shared_workspace, group=device_group)
return projected.add(shared_workspace)
return reference
def check_fused_output(
fused_output: torch.Tensor,
reference: Callable[[], torch.Tensor],
cpu_group: dist.ProcessGroup,
) -> None:
dist.barrier(group=cpu_group)
expected = reference()
torch.testing.assert_close(fused_output, expected, atol=8e-2, rtol=3e-2)
def benchmark_whole_tail(args: argparse.Namespace) -> None:
if any(not 1 <= num_tokens <= 16 for num_tokens in args.num_tokens):
raise ValueError("--num-tokens values must be in [1, 16]")
if args.warmup_replays < 0 or args.samples <= 0:
raise ValueError("warmup replays must be nonnegative and samples positive")
if args.skinny_max_num_tokens is not None and any(
not 0 <= cutoff <= 8 for cutoff in args.skinny_max_num_tokens
):
raise ValueError("--skinny-max-num-tokens must be in [0, 8]")
skinny_configs = dict(args.skinny_config or ())
if len(skinny_configs) != len(args.skinny_config or ()):
raise ValueError("--skinny-config must not repeat an M value")
if any(not 1 <= num_tokens <= 8 for num_tokens in skinny_configs):
raise ValueError("--skinny-config M values must be in [1, 8]")
if not {"RANK", "WORLD_SIZE", "LOCAL_RANK"} <= os.environ.keys():
raise RuntimeError("launch this benchmark with torchrun")
rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"])
local_rank = int(os.environ["LOCAL_RANK"])
device = torch.device("cuda", local_rank)
torch.accelerator.set_device_index(device)
init_distributed_environment()
if world_size > 8:
set_custom_all_reduce(False)
initialize_model_parallel(tensor_model_parallel_size=world_size)
device_group = get_tp_group().device_group
cpu_group = dist.new_group(backend="gloo")
if torch.cuda.get_device_capability(device)[0] != 10:
raise RuntimeError("Kimi K3 latent-MoE tail requires SM100")
torch.manual_seed(20260726)
rms_weight = 1 + 0.1 * torch.randn(
LATENT_SIZE,
dtype=torch.bfloat16,
device=device,
)
up_weight = (
torch.randn(
(HIDDEN_SIZE, LATENT_SIZE),
dtype=torch.bfloat16,
device=device,
)
/ LATENT_SIZE**0.5
)
use_reference = args.backend in ("reference", "both")
use_fused = args.backend in ("fused", "both")
fused_ops = []
if use_fused:
production_config_for_m = fused_add_multicast_skinny_gemm.config_for_m
def config_for_m(
num_rows: int,
shard_dim: int = 896,
) -> fused_add_multicast_skinny_gemm.SkinnyConfig:
config = skinny_configs.get(num_rows)
if config is not None:
return config
return production_config_for_m(num_rows, shard_dim)
fused_add_multicast_skinny_gemm.config_for_m = config_for_m
cutoffs = args.skinny_max_num_tokens or [latent_moe_tail._SKINNY_MAX_NUM_TOKENS]
for cutoff in cutoffs:
latent_moe_tail._SKINNY_MAX_NUM_TOKENS = cutoff
latent_moe_tail.KimiK3LatentMoETailOp._instances.clear()
fused_ops.append(
(
cutoff,
latent_moe_tail.KimiK3LatentMoETailOp.initialize(
hidden_size=HIDDEN_SIZE,
latent_size=LATENT_SIZE,
dtype=torch.bfloat16,
device=device,
rms_eps=RMS_EPS,
),
)
)
cutedsl_warmup()
results = []
for num_tokens in args.num_tokens:
routed, shared = make_inputs(num_tokens, rank, device)
reference = make_reference(
routed,
shared,
rms_weight,
up_weight,
device_group,
)
result: dict[str, Any] = {"num_tokens": num_tokens}
if use_reference:
reference_graph, _ = capture_tail_graph(reference, cpu_group)
result["reference"] = benchmark_tail_graph(
reference_graph,
warmup_replays=args.warmup_replays,
samples=args.samples,
device_group=device_group,
cpu_group=cpu_group,
)
for cutoff, fused_op in fused_ops:
def fused(
routed: torch.Tensor = routed,
shared: torch.Tensor = shared,
fused_op: latent_moe_tail.KimiK3LatentMoETailOp = fused_op,
) -> torch.Tensor:
return fused_op(routed, shared, rms_weight, up_weight)
fused_graph, fused_output = capture_tail_graph(fused, cpu_group)
fused_key = "fused" if len(fused_ops) == 1 else f"fused_skinny_max_{cutoff}"
result[fused_key] = benchmark_tail_graph(
fused_graph,
warmup_replays=args.warmup_replays,
samples=args.samples,
device_group=device_group,
cpu_group=cpu_group,
)
check_fused_output(fused_output, reference, cpu_group)
if "reference" in result:
speedup = (
result["reference"]["median_us"] / result[fused_key]["median_us"]
)
if len(fused_ops) == 1:
result["speedup"] = speedup
else:
result[f"{fused_key}_speedup"] = speedup
results.append(result)
properties = torch.cuda.get_device_properties(device)
report = {
"scope": "whole-tail",
"device": properties.name,
"compute_capability": list(torch.cuda.get_device_capability(device)),
"world_size": world_size,
"torch_version": torch.__version__,
"cuda_version": torch.version.cuda,
"warmup_replays": args.warmup_replays,
"samples": args.samples,
"skinny_max_num_tokens": [cutoff for cutoff, _ in fused_ops],
"skinny_configs": {
str(num_tokens): asdict(config)
for num_tokens, config in skinny_configs.items()
},
"timing_scope": {
"reference": (
"two input copies, two AllReduces, RMSNorm, full replicated "
"up-projection GEMM, and final add"
),
"fused": (
"routed AllReduce/RMSNorm plus shared ReduceScatter, sharded "
"up-projection/multicast, and Lamport copy"
),
},
"results": results,
}
if rank == 0:
rendered = json.dumps(report, indent=2)
print(rendered, flush=True)
if args.output is not None:
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(rendered + "\n", encoding="utf-8")
dist.barrier(group=cpu_group)
def main() -> None:
args = parse_args()
if args.scope == "up-projection":
benchmark_up_projection(args)
return
from vllm.config import VllmConfig, set_current_vllm_config
with set_current_vllm_config(VllmConfig()):
benchmark_whole_tail(args)
if __name__ == "__main__":
main()
@@ -1,239 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import argparse
import json
import os
import statistics
from collections.abc import Callable
import torch
import torch.distributed as dist
import vllm._custom_ops as ops
from vllm.distributed.device_communicators.custom_all_reduce import CustomAllreduce
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--tokens", type=int, nargs="+", default=[8, 32, 128, 1024])
parser.add_argument("--hidden-size", type=int, default=7168)
parser.add_argument("--graph-repeats", type=int, default=20)
parser.add_argument("--warmup-replays", type=int, default=5)
parser.add_argument("--samples", type=int, default=15)
return parser.parse_args()
def capture_graph(op: Callable[[], None], repeats: int) -> torch.cuda.CUDAGraph:
stream = torch.cuda.Stream()
stream.wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(stream):
for _ in range(3):
op()
stream.synchronize()
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph, stream=stream):
for _ in range(repeats):
op()
torch.cuda.current_stream().wait_stream(stream)
return graph
def max_rank_graph_time(
graph: torch.cuda.CUDAGraph,
repeats: int,
warmup_replays: int,
samples: int,
device_group: dist.ProcessGroup,
cpu_group: dist.ProcessGroup,
) -> float:
for _ in range(warmup_replays):
graph.replay()
torch.accelerator.synchronize()
timings = []
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
for _ in range(samples):
dist.barrier(group=cpu_group)
start.record()
graph.replay()
end.record()
end.synchronize()
elapsed = torch.tensor(
start.elapsed_time(end) / repeats,
dtype=torch.float64,
device=torch.accelerator.current_device_index(),
)
dist.all_reduce(elapsed, op=dist.ReduceOp.MAX, group=device_group)
timings.append(elapsed.item())
return statistics.median(timings)
def check_outputs(
comm: CustomAllreduce,
local: torch.Tensor,
reduce_input: torch.Tensor,
device_group: dist.ProcessGroup,
) -> None:
expected_gather = torch.empty(
(local.shape[0] * dist.get_world_size(), local.shape[1]),
dtype=local.dtype,
device=local.device,
)
dist.all_gather_into_tensor(expected_gather, local, group=device_group)
gathered = comm.custom_all_gather(local)
assert gathered is not None
torch.testing.assert_close(gathered, expected_gather)
expected_scatter = torch.empty_like(local)
dist.reduce_scatter_tensor(
expected_scatter,
reduce_input.clone(),
group=device_group,
)
scattered = comm.custom_reduce_scatter(reduce_input)
assert scattered is not None
torch.testing.assert_close(scattered, expected_scatter)
def benchmark_shape(
comm: CustomAllreduce,
global_tokens: int,
hidden_size: int,
graph_repeats: int,
warmup_replays: int,
samples: int,
device_group: dist.ProcessGroup,
cpu_group: dist.ProcessGroup,
) -> dict[str, float | int]:
world_size = dist.get_world_size()
rank = dist.get_rank()
padded_tokens = (global_tokens + world_size - 1) // world_size * world_size
local_tokens = padded_tokens // world_size
local = torch.full(
(local_tokens, hidden_size),
rank + 1,
dtype=torch.bfloat16,
device=torch.accelerator.current_device_index(),
)
reduce_input = torch.full(
(padded_tokens, hidden_size),
rank + 1,
dtype=torch.bfloat16,
device=local.device,
)
check_outputs(comm, local, reduce_input, device_group)
custom_gather_out = torch.empty(
(padded_tokens, hidden_size),
dtype=local.dtype,
device=local.device,
)
custom_scatter_out = torch.empty_like(local)
nccl_gather_out = torch.empty_like(custom_gather_out)
nccl_scatter_out = torch.empty_like(local)
def custom_ag() -> None:
ops.mnnvl_lamport_all_gather(
comm._ptr,
local,
custom_gather_out,
comm.mnnvl_lamport_ag_local_ptr,
comm.mnnvl_lamport_ag_multicast_ptr,
comm.mnnvl_lamport_ag_epoch_ptr,
comm.mnnvl_buffer_size,
)
def custom_rs() -> None:
ops.mnnvl_lamport_reduce_scatter(
comm._ptr,
reduce_input,
custom_scatter_out,
comm.mnnvl_lamport_rs_local_ptr,
comm.mnnvl_lamport_rs_epoch_ptr,
comm.mnnvl_buffer_size,
)
def nccl_ag() -> None:
dist.all_gather_into_tensor(nccl_gather_out, local, group=device_group)
def nccl_rs() -> None:
dist.reduce_scatter_tensor(
nccl_scatter_out,
reduce_input,
group=device_group,
)
graphs = {
"custom_ag_us": capture_graph(custom_ag, graph_repeats),
"nccl_ag_us": capture_graph(nccl_ag, graph_repeats),
"custom_rs_us": capture_graph(custom_rs, graph_repeats),
"nccl_rs_us": capture_graph(nccl_rs, graph_repeats),
}
times = {
name: max_rank_graph_time(
graph,
graph_repeats,
warmup_replays,
samples,
device_group,
cpu_group,
)
* 1000
for name, graph in graphs.items()
}
torch.testing.assert_close(custom_gather_out, nccl_gather_out)
torch.testing.assert_close(custom_scatter_out, nccl_scatter_out)
return {
"global_tokens": global_tokens,
"padded_tokens": padded_tokens,
"local_bytes": local.nbytes,
"full_bytes": reduce_input.nbytes,
**times,
"ag_speedup": times["nccl_ag_us"] / times["custom_ag_us"],
"rs_speedup": times["nccl_rs_us"] / times["custom_rs_us"],
}
def main() -> None:
args = parse_args()
local_rank = int(os.environ["LOCAL_RANK"])
torch.accelerator.set_device_index(local_rank)
dist.init_process_group("nccl")
device_group = dist.group.WORLD
cpu_group = dist.new_group(backend="gloo")
comm = CustomAllreduce(
group=cpu_group,
device=torch.device("cuda", local_rank),
)
assert not comm.disabled
assert comm.world_size == 16
assert comm.mnnvl_only
assert comm.mnnvl_multicast_ptr
results = [
benchmark_shape(
comm,
tokens,
args.hidden_size,
args.graph_repeats,
args.warmup_replays,
args.samples,
device_group,
cpu_group,
)
for tokens in args.tokens
]
if dist.get_rank() == 0:
print(json.dumps(results, indent=2), flush=True)
comm.close()
dist.destroy_process_group(cpu_group)
dist.destroy_process_group()
if __name__ == "__main__":
main()
+1 -1
View File
@@ -154,7 +154,7 @@ def main(
scale=scale,
causal=True,
alibi_slopes=None,
sliding_window=window_size,
sliding_window=window_size if sliding_window is not None else -1,
block_table=block_tables,
softcap=0,
scheduler_metadata=metadata,
+19 -1
View File
@@ -15,6 +15,7 @@ endif()
#
set(ENABLE_X86_ISA $ENV{VLLM_CPU_X86})
set(ENABLE_ARM_BF16 $ENV{VLLM_CPU_ARM_BF16})
set(ENABLE_ARM_I8MM $ENV{VLLM_CPU_ARM_I8MM})
set(ENABLE_RVV_BF16 $ENV{VLLM_CPU_RVV_BF16})
include_directories("${CMAKE_SOURCE_DIR}/csrc")
@@ -96,12 +97,14 @@ if (MACOSX_FOUND AND CMAKE_SYSTEM_PROCESSOR STREQUAL "arm64")
set(ENABLE_NUMA OFF)
check_sysctl(hw.optional.neon ASIMD_FOUND)
check_sysctl(hw.optional.arm.FEAT_BF16 ARM_BF16_FOUND)
check_sysctl(hw.optional.arm.FEAT_I8MM ARM_I8MM_FOUND)
else()
find_isa(${CPUINFO} "Power11" POWER11_FOUND)
find_isa(${CPUINFO} "POWER10" POWER10_FOUND)
find_isa(${CPUINFO} "POWER9" POWER9_FOUND)
find_isa(${CPUINFO} "asimd" ASIMD_FOUND) # Check for ARM NEON support
find_isa(${CPUINFO} "bf16" ARM_BF16_FOUND) # Check for ARM BF16 support
find_isa(${CPUINFO} "i8mm" ARM_I8MM_FOUND) # Check for ARM I8MM support
find_isa(${CPUINFO} "S390" S390_FOUND)
find_isa(${CPUINFO} "zvfhmin" RVV_FP16_FOUND) # Check for RISC-V Vector FP16 support
find_isa(${CPUINFO} "zvfbfmin" RVV_BF16_FOUND) # Check for RISC-V Vector BF16 support
@@ -111,6 +114,11 @@ else()
set(ARM_BF16_FOUND ON)
message(STATUS "ARM BF16 support enabled via VLLM_CPU_ARM_BF16 environment variable")
endif()
if (ENABLE_ARM_I8MM)
set(ARM_I8MM_FOUND ON)
message(STATUS
"ARM I8MM support enabled via VLLM_CPU_ARM_I8MM environment variable")
endif()
# Some kernels (e.g. Bianbu on Spacemit X100) do not report zvfbfmin
# in /proc/cpuinfo despite hardware support. VLLM_CPU_RVV_BF16=1
# overrides the detection result.
@@ -166,6 +174,11 @@ elseif (ASIMD_FOUND)
message(WARNING "BF16 functionality is not available")
set(MARCH_FLAGS "-march=armv8.2-a+dotprod+fp16")
endif()
if(ARM_I8MM_FOUND)
message(STATUS "I8MM extension detected")
string(APPEND MARCH_FLAGS "+i8mm")
add_compile_definitions(ARM_I8MM_SUPPORT)
endif()
list(APPEND CXX_COMPILE_FLAGS ${MARCH_FLAGS})
elseif (S390_FOUND)
message(STATUS "S390 detected")
@@ -447,8 +460,13 @@ if (ASIMD_FOUND AND NOT APPLE_SILICON_FOUND)
"csrc/cpu/shm.cpp"
"csrc/cpu/activation_lut_bf16.cpp"
"csrc/cpu/cpu_tanhf_neon.hpp"
"csrc/cpu/cpu_fused_moe.cpp"
${VLLM_EXT_SRC})
if (ARM_BF16_FOUND)
set(VLLM_EXT_SRC "csrc/cpu/cpu_fused_moe.cpp" ${VLLM_EXT_SRC})
if (ARM_I8MM_FOUND)
set(VLLM_EXT_SRC "csrc/cpu/cpu_fused_moe_int8.cpp" ${VLLM_EXT_SRC})
endif()
endif()
endif()
if (POWER9_FOUND OR POWER10_FOUND OR POWER11_FOUND)
+2 -2
View File
@@ -28,9 +28,9 @@ if(DEEPGEMM_SRC_DIR)
message(STATUS "DeepGEMM using local DEEPGEMM_SRC_DIR: ${deepgemm_SOURCE_DIR}")
else()
# Keep in sync with tools/install_deepgemm.sh
set(_DEEPGEMM_UPSTREAM_REPO "git@github.com:Inferact/DeepGEMM.git")
set(_DEEPGEMM_UPSTREAM_REPO "https://github.com/deepseek-ai/DeepGEMM.git")
# NOTE: This is currently targeting nv-dev branch due to sm120 support
set(_DEEPGEMM_UPSTREAM_TAG "f5a76426fa084087169693fd0cd815223576d6e9")
set(_DEEPGEMM_UPSTREAM_TAG "a6b593d2826719dcf4892609af7b84ee23aaf32a")
set(_deepgemm_fc_root "${FETCHCONTENT_BASE_DIR}")
if(NOT _deepgemm_fc_root)
-74
View File
@@ -1,74 +0,0 @@
include(FetchContent)
if(DEFINED ENV{FLASH_KDA_SRC_DIR})
set(FLASH_KDA_SRC_DIR $ENV{FLASH_KDA_SRC_DIR})
endif()
if(FLASH_KDA_SRC_DIR)
FetchContent_Declare(
flashkda
SOURCE_DIR ${FLASH_KDA_SRC_DIR}
)
else()
FetchContent_Declare(
flashkda
GIT_REPOSITORY git@github.com:Inferact/FlashKDA.git
GIT_TAG a3e42bbbece3bb38f7c426b880315294a336e82f
GIT_PROGRESS TRUE
GIT_SUBMODULES cutlass
)
endif()
FetchContent_MakeAvailable(flashkda)
message(STATUS "FlashKDA is available at ${flashkda_SOURCE_DIR}")
set(FLASH_KDA_SUPPORT_ARCHS)
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 12.0)
list(APPEND FLASH_KDA_SUPPORT_ARCHS "9.0a")
endif()
if(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 13.0)
list(APPEND FLASH_KDA_SUPPORT_ARCHS "10.0f" "12.0f")
elseif(${CMAKE_CUDA_COMPILER_VERSION} VERSION_GREATER_EQUAL 12.9)
list(APPEND FLASH_KDA_SUPPORT_ARCHS "10.0a" "10.3a" "12.0a")
endif()
cuda_archs_loose_intersection(
FLASH_KDA_ARCHS "${FLASH_KDA_SUPPORT_ARCHS}" "${CUDA_ARCHS}")
if(FLASH_KDA_ARCHS)
message(STATUS "FlashKDA CUDA architectures: ${FLASH_KDA_ARCHS}")
set(FLASH_KDA_SOURCES
csrc/flashkda_registration.cpp
${flashkda_SOURCE_DIR}/csrc/flash_kda.cpp
${flashkda_SOURCE_DIR}/csrc/smxx/fwd_launch.cu)
set(FLASH_KDA_INCLUDES
${flashkda_SOURCE_DIR}/csrc
${flashkda_SOURCE_DIR}/cutlass/include
${flashkda_SOURCE_DIR}/cutlass/examples/common
${flashkda_SOURCE_DIR}/cutlass/tools/util/include)
set_gencode_flags_for_srcs(
SRCS "${FLASH_KDA_SOURCES}"
CUDA_ARCHS "${FLASH_KDA_ARCHS}")
define_extension_target(
_flashkda_C
DESTINATION vllm
LANGUAGE ${VLLM_GPU_LANG}
SOURCES ${FLASH_KDA_SOURCES}
COMPILE_FLAGS ${VLLM_GPU_FLAGS}
ARCHITECTURES ${VLLM_GPU_ARCHES}
INCLUDE_DIRECTORIES ${FLASH_KDA_INCLUDES}
USE_SABI 3
WITH_SOABI)
target_compile_options(_flashkda_C PRIVATE
$<$<COMPILE_LANGUAGE:CUDA>:-UPy_LIMITED_API --expt-relaxed-constexpr --expt-extended-lambda --use_fast_math -O3>
$<$<COMPILE_LANGUAGE:CXX>:-UPy_LIMITED_API>)
else()
message(STATUS
"FlashKDA will not compile: CUDA >=12.0 and a supported architecture "
"(SM90, SM10x, or SM12x) are required")
add_custom_target(_flashkda_C)
endif()
+11
View File
@@ -172,4 +172,15 @@
#endif // __riscv_v
// Power VSX
#ifdef __powerpc__
// FP32Vec16::exp() in cpu_types_vsx.hpp delegates to FP32Vec8::exp(), which
// implements a vectorised 5-term minimax polynomial using VSX intrinsics.
#define DEFINE_FAST_EXP \
auto fast_exp = [&](const vec_op::FP32Vec16& vec) \
__attribute__((always_inline)) { return vec.exp(); }; \
auto fast_exp_f16 = fast_exp;
#endif // __powerpc__
#endif
+5 -187
View File
@@ -1,5 +1,6 @@
#include "cpu/cpu_types.hpp"
#include "cpu/utils.hpp"
#include "cpu/cpu_fused_moe_activations.hpp"
#include "cpu/micro_gemm/cpu_micro_gemm_vec.hpp"
#include "cpu/cpu_arch_macros.h"
@@ -43,193 +44,9 @@
}()
namespace {
enum class FusedMOEAct {
SiluAndMul,
SwigluOAIAndMul,
GeluAndMul,
GeluTanhAndMul,
};
FusedMOEAct get_act_type(const std::string& act) {
if (act == "silu") {
return FusedMOEAct::SiluAndMul;
} else if (act == "swigluoai") {
return FusedMOEAct::SwigluOAIAndMul;
} else if (act == "gelu") {
return FusedMOEAct::GeluAndMul;
} else if (act == "gelu_tanh") {
return FusedMOEAct::GeluTanhAndMul;
} else {
TORCH_CHECK(false, "Invalid act type: " + act);
}
}
template <typename scalar_t>
void swigluoai_and_mul(float* __restrict__ input, scalar_t* __restrict__ output,
const int32_t m_size, const int32_t n_size,
const int32_t input_stride,
const int32_t output_stride) {
using scalar_vec_t = typename cpu_utils::VecTypeTrait<scalar_t>::vec_t;
#if !defined(__aarch64__)
// For GPT-OSS interleaved gate-up weights
alignas(64) static int32_t index[16] = {0, 2, 4, 6, 8, 10, 12, 14,
16, 18, 20, 22, 24, 26, 28, 30};
vec_op::INT32Vec16 index_vec(index);
#endif
vec_op::FP32Vec16 gate_up_max_vec(7.0);
vec_op::FP32Vec16 up_min_vec(-7.0);
vec_op::FP32Vec16 alpha_vec(1.702);
vec_op::FP32Vec16 one_vec(1.0);
DEFINE_FAST_EXP
for (int32_t m = 0; m < m_size; ++m) {
for (int32_t n = 0; n < n_size; n += 32) {
// Note: AdvSIMD does not support gather loads
#if defined(__aarch64__)
vec_op::FP32Vec16 gate_vec(vec_op::uninit);
vec_op::FP32Vec16 up_vec(vec_op::uninit);
vec_op::FP32Vec16::load_even_odd(input + n, gate_vec, up_vec);
#else
vec_op::FP32Vec16 gate_vec(input + n, index_vec);
vec_op::FP32Vec16 up_vec(input + n + 1, index_vec);
#endif
gate_vec = gate_vec.min(gate_up_max_vec);
up_vec = up_vec.clamp(up_min_vec, gate_up_max_vec);
auto sigmoid_vec = one_vec / (one_vec + fast_exp(-gate_vec * alpha_vec));
auto glu = gate_vec * sigmoid_vec;
auto gated_output_fp32 = (one_vec + up_vec) * glu;
scalar_vec_t gated_output = scalar_vec_t(gated_output_fp32);
gated_output.save(output + n / 2);
}
input += input_stride;
output += output_stride;
}
}
template <typename scalar_t>
void silu_and_mul(float* __restrict__ input, scalar_t* __restrict__ output,
const int32_t m_size, const int32_t n_size,
const int32_t input_stride, const int32_t output_stride) {
using scalar_vec_t = typename cpu_utils::VecTypeTrait<scalar_t>::vec_t;
const int32_t dim = n_size / 2;
float* __restrict__ gate = input;
float* __restrict__ up = input + dim;
vec_op::FP32Vec16 one_vec(1.0);
DEFINE_FAST_EXP
for (int32_t m = 0; m < m_size; ++m) {
for (int32_t n = 0; n < dim; n += 16) {
vec_op::FP32Vec16 gate_vec(gate + n);
vec_op::FP32Vec16 up_vec(up + n);
auto sigmoid_vec = one_vec / (one_vec + fast_exp(-gate_vec));
auto silu = gate_vec * sigmoid_vec;
auto gated_output_fp32 = up_vec * silu;
scalar_vec_t gated_output = scalar_vec_t(gated_output_fp32);
gated_output.save(output + n);
}
gate += input_stride;
up += input_stride;
output += output_stride;
}
}
template <typename scalar_t>
void gelu_and_mul(float* __restrict__ input, scalar_t* __restrict__ output,
const int32_t m_size, const int32_t n_size,
const int32_t input_stride, const int32_t output_stride) {
using scalar_vec_t = typename cpu_utils::VecTypeTrait<scalar_t>::vec_t;
const int32_t dim = n_size / 2;
float* __restrict__ gate = input;
float* __restrict__ up = input + dim;
vec_op::FP32Vec16 one_vec(1.0);
vec_op::FP32Vec16 w1_vec(M_SQRT1_2);
vec_op::FP32Vec16 w2_vec(0.5);
alignas(64) float temp[16];
DEFINE_FAST_EXP
for (int32_t m = 0; m < m_size; ++m) {
for (int32_t n = 0; n < dim; n += 16) {
vec_op::FP32Vec16 gate_vec(gate + n);
vec_op::FP32Vec16 up_vec(up + n);
auto er_input_vec = gate_vec * w1_vec;
er_input_vec.save(temp);
for (int32_t i = 0; i < 16; ++i) {
temp[i] = std::erf(temp[i]);
}
vec_op::FP32Vec16 er_vec(temp);
auto gelu = gate_vec * w2_vec * (one_vec + er_vec);
auto gated_output_fp32 = up_vec * gelu;
scalar_vec_t gated_output = scalar_vec_t(gated_output_fp32);
gated_output.save(output + n);
}
gate += input_stride;
up += input_stride;
output += output_stride;
}
}
template <typename scalar_t>
void gelu_tanh_and_mul(float* __restrict__ input, scalar_t* __restrict__ output,
const int32_t m_size, const int32_t n_size,
const int32_t input_stride,
const int32_t output_stride) {
using scalar_vec_t = typename cpu_utils::VecTypeTrait<scalar_t>::vec_t;
const int32_t dim = n_size / 2;
float* __restrict__ gate = input;
float* __restrict__ up = input + dim;
vec_op::FP32Vec16 one_vec(1.0);
vec_op::FP32Vec16 w1_vec(0.7978845608028654);
vec_op::FP32Vec16 w2_vec(0.5);
vec_op::FP32Vec16 w3_vec(0.044715);
for (int32_t m = 0; m < m_size; ++m) {
for (int32_t n = 0; n < dim; n += 16) {
vec_op::FP32Vec16 gate_vec(gate + n);
vec_op::FP32Vec16 up_vec(up + n);
auto gate_pow3_vec = gate_vec * gate_vec * gate_vec;
auto inner_vec = w1_vec * (gate_vec + w3_vec * gate_pow3_vec);
// Note: can't use fast_exp form because diffusiongemma will generate
// wrong results
auto tanh_vec = inner_vec.tanh();
auto gelu_tanh = gate_vec * w2_vec * (one_vec + tanh_vec);
auto gated_output_fp32 = up_vec * gelu_tanh;
scalar_vec_t gated_output = scalar_vec_t(gated_output_fp32);
gated_output.save(output + n);
}
gate += input_stride;
up += input_stride;
output += output_stride;
}
}
template <typename scalar_t>
FORCE_INLINE void apply_gated_act(const FusedMOEAct act,
float* __restrict__ input,
scalar_t* __restrict__ output,
const int32_t m, const int32_t n,
const int32_t input_stride,
const int32_t output_stride) {
switch (act) {
case FusedMOEAct::SwigluOAIAndMul:
swigluoai_and_mul(input, output, m, n, input_stride, output_stride);
return;
case FusedMOEAct::SiluAndMul:
silu_and_mul(input, output, m, n, input_stride, output_stride);
return;
case FusedMOEAct::GeluAndMul:
gelu_and_mul(input, output, m, n, input_stride, output_stride);
return;
case FusedMOEAct::GeluTanhAndMul:
gelu_tanh_and_mul(input, output, m, n, input_stride, output_stride);
return;
default:
TORCH_CHECK(false, "Unsupported act type.");
}
}
using cpu_fused_moe_utils::apply_gated_act;
using cpu_fused_moe_utils::FusedMOEAct;
template <typename scalar_t, typename gemm_t>
void prepack_moe_weight_impl(scalar_t* __restrict__ weight_ptr,
@@ -817,6 +634,7 @@ void fused_moe_impl(scalar_t* __restrict__ output, scalar_t* __restrict__ input,
}
}
}
} // namespace
void prepack_moe_weight(
@@ -864,7 +682,7 @@ void cpu_fused_moe(
const int32_t input_size_2 = w2.size(2);
const int32_t output_size_2 = w2.size(1);
const int32_t topk_num = topk_id.size(1);
const FusedMOEAct act_type = get_act_type(act);
const FusedMOEAct act_type = cpu_fused_moe_utils::get_act_type(act);
cpu_utils::ISA isa_type = cpu_utils::get_isa(isa);
TORCH_CHECK(!skip_weighted || topk_num == 1,
"skip_weighted is only supported for topk=1 on CPU");
+204
View File
@@ -0,0 +1,204 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
#ifndef CPU_FUSED_MOE_ACTIVATIONS_HPP
#define CPU_FUSED_MOE_ACTIVATIONS_HPP
#include <cmath>
#include <cstdint>
#include <string>
#include "cpu/cpu_arch_macros.h"
#include "cpu/utils.hpp"
namespace cpu_fused_moe_utils {
enum class FusedMOEAct {
SiluAndMul,
SwigluOAIAndMul,
GeluAndMul,
GeluTanhAndMul,
};
inline FusedMOEAct get_act_type(const std::string& act) {
if (act == "silu") {
return FusedMOEAct::SiluAndMul;
} else if (act == "swigluoai") {
return FusedMOEAct::SwigluOAIAndMul;
} else if (act == "gelu") {
return FusedMOEAct::GeluAndMul;
} else if (act == "gelu_tanh") {
return FusedMOEAct::GeluTanhAndMul;
} else {
TORCH_CHECK(false, "Invalid act type: " + act);
}
}
template <typename scalar_t>
void swigluoai_and_mul(float* __restrict__ input, scalar_t* __restrict__ output,
const int32_t m_size, const int32_t n_size,
const int32_t input_stride,
const int32_t output_stride) {
using scalar_vec_t = typename cpu_utils::VecTypeTrait<scalar_t>::vec_t;
#if !defined(__aarch64__)
// For GPT-OSS interleaved gate-up weights
alignas(64) static int32_t index[16] = {0, 2, 4, 6, 8, 10, 12, 14,
16, 18, 20, 22, 24, 26, 28, 30};
vec_op::INT32Vec16 index_vec(index);
#endif
vec_op::FP32Vec16 gate_up_max_vec(7.0);
vec_op::FP32Vec16 up_min_vec(-7.0);
vec_op::FP32Vec16 alpha_vec(1.702);
vec_op::FP32Vec16 one_vec(1.0);
DEFINE_FAST_EXP
for (int32_t m = 0; m < m_size; ++m) {
for (int32_t n = 0; n < n_size; n += 32) {
// Note: AdvSIMD does not support gather loads
#if defined(__aarch64__)
vec_op::FP32Vec16 gate_vec(vec_op::uninit);
vec_op::FP32Vec16 up_vec(vec_op::uninit);
vec_op::FP32Vec16::load_even_odd(input + n, gate_vec, up_vec);
#else
vec_op::FP32Vec16 gate_vec(input + n, index_vec);
vec_op::FP32Vec16 up_vec(input + n + 1, index_vec);
#endif
gate_vec = gate_vec.min(gate_up_max_vec);
up_vec = up_vec.clamp(up_min_vec, gate_up_max_vec);
auto sigmoid_vec = one_vec / (one_vec + fast_exp(-gate_vec * alpha_vec));
auto glu = gate_vec * sigmoid_vec;
auto gated_output_fp32 = (one_vec + up_vec) * glu;
scalar_vec_t gated_output = scalar_vec_t(gated_output_fp32);
gated_output.save(output + n / 2);
}
input += input_stride;
output += output_stride;
}
}
template <typename scalar_t>
void silu_and_mul(float* __restrict__ input, scalar_t* __restrict__ output,
const int32_t m_size, const int32_t n_size,
const int32_t input_stride, const int32_t output_stride) {
using scalar_vec_t = typename cpu_utils::VecTypeTrait<scalar_t>::vec_t;
const int32_t dim = n_size / 2;
float* __restrict__ gate = input;
float* __restrict__ up = input + dim;
vec_op::FP32Vec16 one_vec(1.0);
DEFINE_FAST_EXP
for (int32_t m = 0; m < m_size; ++m) {
for (int32_t n = 0; n < dim; n += 16) {
vec_op::FP32Vec16 gate_vec(gate + n);
vec_op::FP32Vec16 up_vec(up + n);
auto sigmoid_vec = one_vec / (one_vec + fast_exp(-gate_vec));
auto silu = gate_vec * sigmoid_vec;
auto gated_output_fp32 = up_vec * silu;
scalar_vec_t gated_output = scalar_vec_t(gated_output_fp32);
gated_output.save(output + n);
}
gate += input_stride;
up += input_stride;
output += output_stride;
}
}
template <typename scalar_t>
void gelu_and_mul(float* __restrict__ input, scalar_t* __restrict__ output,
const int32_t m_size, const int32_t n_size,
const int32_t input_stride, const int32_t output_stride) {
using scalar_vec_t = typename cpu_utils::VecTypeTrait<scalar_t>::vec_t;
const int32_t dim = n_size / 2;
float* __restrict__ gate = input;
float* __restrict__ up = input + dim;
vec_op::FP32Vec16 one_vec(1.0);
vec_op::FP32Vec16 w1_vec(M_SQRT1_2);
vec_op::FP32Vec16 w2_vec(0.5);
alignas(64) float temp[16];
DEFINE_FAST_EXP
for (int32_t m = 0; m < m_size; ++m) {
for (int32_t n = 0; n < dim; n += 16) {
vec_op::FP32Vec16 gate_vec(gate + n);
vec_op::FP32Vec16 up_vec(up + n);
auto er_input_vec = gate_vec * w1_vec;
er_input_vec.save(temp);
for (int32_t i = 0; i < 16; ++i) {
temp[i] = std::erf(temp[i]);
}
vec_op::FP32Vec16 er_vec(temp);
auto gelu = gate_vec * w2_vec * (one_vec + er_vec);
auto gated_output_fp32 = up_vec * gelu;
scalar_vec_t gated_output = scalar_vec_t(gated_output_fp32);
gated_output.save(output + n);
}
gate += input_stride;
up += input_stride;
output += output_stride;
}
}
template <typename scalar_t>
void gelu_tanh_and_mul(float* __restrict__ input, scalar_t* __restrict__ output,
const int32_t m_size, const int32_t n_size,
const int32_t input_stride,
const int32_t output_stride) {
using scalar_vec_t = typename cpu_utils::VecTypeTrait<scalar_t>::vec_t;
const int32_t dim = n_size / 2;
float* __restrict__ gate = input;
float* __restrict__ up = input + dim;
vec_op::FP32Vec16 one_vec(1.0);
vec_op::FP32Vec16 w1_vec(0.7978845608028654);
vec_op::FP32Vec16 w2_vec(0.5);
vec_op::FP32Vec16 w3_vec(0.044715);
for (int32_t m = 0; m < m_size; ++m) {
for (int32_t n = 0; n < dim; n += 16) {
vec_op::FP32Vec16 gate_vec(gate + n);
vec_op::FP32Vec16 up_vec(up + n);
auto gate_pow3_vec = gate_vec * gate_vec * gate_vec;
auto inner_vec = w1_vec * (gate_vec + w3_vec * gate_pow3_vec);
// Note: can't use fast_exp form because diffusiongemma will generate
// wrong results
auto tanh_vec = inner_vec.tanh();
auto gelu_tanh = gate_vec * w2_vec * (one_vec + tanh_vec);
auto gated_output_fp32 = up_vec * gelu_tanh;
scalar_vec_t gated_output = scalar_vec_t(gated_output_fp32);
gated_output.save(output + n);
}
gate += input_stride;
up += input_stride;
output += output_stride;
}
}
template <typename scalar_t>
FORCE_INLINE void apply_gated_act(const FusedMOEAct act,
float* __restrict__ input,
scalar_t* __restrict__ output,
const int32_t m, const int32_t n,
const int32_t input_stride,
const int32_t output_stride) {
switch (act) {
case FusedMOEAct::SwigluOAIAndMul:
swigluoai_and_mul(input, output, m, n, input_stride, output_stride);
return;
case FusedMOEAct::SiluAndMul:
silu_and_mul(input, output, m, n, input_stride, output_stride);
return;
case FusedMOEAct::GeluAndMul:
gelu_and_mul(input, output, m, n, input_stride, output_stride);
return;
case FusedMOEAct::GeluTanhAndMul:
gelu_tanh_and_mul(input, output, m, n, input_stride, output_stride);
return;
default:
TORCH_CHECK(false, "Unsupported act type.");
}
}
} // namespace cpu_fused_moe_utils
#endif
+647
View File
@@ -0,0 +1,647 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
#include "cpu/cpu_arch_macros.h"
#include <algorithm>
#include <cstdint>
#include <cstring>
#include <optional>
#include <string>
#include "cpu/cpu_fused_moe_activations.hpp"
#include "cpu/cpu_types.hpp"
#include "cpu/micro_gemm/cpu_micro_gemm_impl.hpp"
#include "cpu/utils.hpp"
#if defined(ARM_I8MM_SUPPORT) && defined(ARM_BF16_SUPPORT)
#include "cpu/micro_gemm/cpu_micro_gemm_int8_neon.hpp"
#define NEON_DISPATCH(SCALAR_TYPE, ...) \
case cpu_utils::ISA::NEON: { \
using gemm_t = \
cpu_micro_gemm::MicroGemmINT8<cpu_utils::ISA::NEON, SCALAR_TYPE>; \
return __VA_ARGS__(); \
}
#else
#define NEON_DISPATCH(SCALAR_TYPE, ...) case cpu_utils::ISA::NEON:
#endif
#define CPU_INT8_ISA_DISPATCH_IMPL(ISA_TYPE, SCALAR_TYPE, ...) \
[&] { \
switch (ISA_TYPE) { \
NEON_DISPATCH(SCALAR_TYPE, __VA_ARGS__) \
default: { \
TORCH_CHECK(false, "Invalid CPU ISA type."); \
} \
} \
}()
namespace {
using cpu_fused_moe_utils::apply_gated_act;
using cpu_fused_moe_utils::FusedMOEAct;
template <typename gemm_t>
void prepack_moe_weight_int8_impl(const int8_t* __restrict__ weight_ptr,
int8_t* __restrict__ packed_weight_ptr,
const int32_t expert_num,
const int32_t output_size,
const int32_t input_size,
const int64_t expert_stride) {
#pragma omp parallel for
for (int32_t e_idx = 0; e_idx < expert_num; ++e_idx) {
gemm_t::pack_weight(weight_ptr + expert_stride * e_idx,
packed_weight_ptr + expert_stride * e_idx, output_size,
input_size);
}
}
// INT8 MoE kernel, based on the original BF16 kernel in cpu_fused_moe.cpp
template <typename scalar_t, typename gemm_t>
void fused_moe_int8_impl(
scalar_t* __restrict__ output, const scalar_t* __restrict__ input,
const int8_t* __restrict__ w13, const int8_t* __restrict__ w2,
const float* __restrict__ w13_scales, const float* __restrict__ w2_scales,
scalar_t* __restrict__ w13_bias, scalar_t* __restrict__ w2_bias,
const float* __restrict__ topk_weights, const int32_t* __restrict__ topk_id,
const FusedMOEAct act_type, const int32_t token_num,
const int32_t expert_num, const int32_t topk_num,
const int32_t input_size_13, const int32_t output_size_13,
const int32_t input_size_2, const int32_t output_size_2,
const bool skip_weighted) {
using scalar_vec_t = typename cpu_utils::VecTypeTrait<scalar_t>::vec_t;
constexpr int32_t gemm_n_tile_size = gemm_t::NSize;
constexpr int32_t gemm_m_tile_size = gemm_t::MaxMSize;
constexpr int32_t min_w13_n_tile_size = 2 * gemm_n_tile_size;
TORCH_CHECK_EQ(input_size_13 % gemm_t::K, 0);
TORCH_CHECK_EQ(input_size_2 % gemm_t::K, 0);
TORCH_CHECK_EQ(output_size_13 % min_w13_n_tile_size, 0);
TORCH_CHECK_EQ(output_size_2 % gemm_n_tile_size, 0);
TORCH_CHECK_EQ(output_size_13 / 2, input_size_2);
const int32_t thread_num = cpu_utils::get_max_threads();
const int32_t w13_input_buffer_size = cpu_utils::round_up<64>(
gemm_m_tile_size * input_size_13 * sizeof(int8_t));
const int32_t w2_input_buffer_size =
cpu_utils::round_up<64>(gemm_m_tile_size * input_size_2 * sizeof(int8_t));
const int32_t w13_n_tile_size = [&]() {
const int64_t cache_size = cpu_utils::get_available_l2_size();
const int32_t n_size_cache_limit =
(cache_size - w13_input_buffer_size) /
(gemm_m_tile_size * sizeof(float) + input_size_13 * sizeof(int8_t));
const int32_t n_size_thread_limit =
output_size_13 / std::max(1, thread_num / topk_num);
const int32_t n_size = cpu_utils::round_down<min_w13_n_tile_size>(
std::min(n_size_cache_limit, n_size_thread_limit));
return std::max(n_size, min_w13_n_tile_size);
}();
const int32_t w2_n_tile_size = [&]() {
const int64_t cache_size = cpu_utils::get_available_l2_size();
const int32_t n_size_cache_limit =
(cache_size - w2_input_buffer_size) / (input_size_2 * sizeof(int8_t));
const int32_t n_size_thread_limit =
output_size_2 / std::max(1, thread_num / topk_num);
const int32_t n_size = cpu_utils::round_down<gemm_n_tile_size>(
std::min(n_size_cache_limit, n_size_thread_limit));
return std::max(n_size, gemm_n_tile_size);
}();
int32_t common_buffer_offset = 0;
const int32_t token_num_per_group_buffer_offset = common_buffer_offset;
common_buffer_offset += cpu_utils::round_up<64>(expert_num * sizeof(int32_t));
const int32_t cu_token_num_per_group_buffer_offset = common_buffer_offset;
common_buffer_offset +=
cpu_utils::round_up<64>((expert_num + 1) * sizeof(int32_t));
const int32_t expanded_token_num = token_num * topk_num;
const int32_t expand_token_id_buffer_offset = common_buffer_offset;
common_buffer_offset +=
cpu_utils::round_up<64>(expanded_token_num * sizeof(int32_t));
const int32_t expand_token_id_index_buffer_offset = common_buffer_offset;
common_buffer_offset +=
cpu_utils::round_up<64>(expanded_token_num * sizeof(int32_t));
const int32_t input_quant_buffer_offset = common_buffer_offset;
common_buffer_offset +=
cpu_utils::round_up<64>(token_num * input_size_13 * sizeof(int8_t));
const int32_t input_scale_buffer_offset = common_buffer_offset;
common_buffer_offset += cpu_utils::round_up<64>(token_num * sizeof(float));
const int32_t w13_gemm_output_buffer_offset = common_buffer_offset;
common_buffer_offset += cpu_utils::round_up<64>(
expanded_token_num * input_size_2 * sizeof(scalar_t));
const int32_t w13_output_scale_buffer_offset = common_buffer_offset;
common_buffer_offset +=
cpu_utils::round_up<64>(expanded_token_num * sizeof(float));
const int32_t w2_gemm_output_buffer_offset = common_buffer_offset;
common_buffer_offset += cpu_utils::round_up<64>(
expanded_token_num * output_size_2 * sizeof(float));
int32_t gemm_thread_buffer_offset = 0;
const int32_t gemm_input_buffer_offset = gemm_thread_buffer_offset;
gemm_thread_buffer_offset +=
std::max(w13_input_buffer_size, w2_input_buffer_size);
const int32_t gemm_output_buffer_offset = gemm_thread_buffer_offset;
gemm_thread_buffer_offset += cpu_utils::round_up<64>(
gemm_m_tile_size * std::max(w13_n_tile_size, w2_n_tile_size) *
sizeof(int32_t));
const int32_t ws_output_buffer_offset = 0;
const int32_t ws_thread_buffer_size =
cpu_utils::round_up<64>(output_size_2 * sizeof(float));
const int32_t thread_buffer_size =
std::max(gemm_thread_buffer_offset, ws_thread_buffer_size);
const int32_t buffer_size =
common_buffer_offset + thread_buffer_size * thread_num;
cpu_utils::ScratchPadManager::get_scratchpad_manager()->realloc(buffer_size);
uint8_t* common_buffer_start =
cpu_utils::ScratchPadManager::get_scratchpad_manager()
->get_data<uint8_t>();
uint8_t* thread_buffer_start = common_buffer_start + common_buffer_offset;
int32_t* __restrict__ token_num_per_group_buffer = reinterpret_cast<int32_t*>(
common_buffer_start + token_num_per_group_buffer_offset);
int32_t* __restrict__ cu_token_num_per_group_buffer =
reinterpret_cast<int32_t*>(common_buffer_start +
cu_token_num_per_group_buffer_offset);
int32_t* __restrict__ expand_token_id_buffer = reinterpret_cast<int32_t*>(
common_buffer_start + expand_token_id_buffer_offset);
int32_t* __restrict__ expand_token_id_index_buffer =
reinterpret_cast<int32_t*>(common_buffer_start +
expand_token_id_index_buffer_offset);
int8_t* __restrict__ input_quant_buffer = reinterpret_cast<int8_t*>(
common_buffer_start + input_quant_buffer_offset);
float* __restrict__ input_scale_buffer =
reinterpret_cast<float*>(common_buffer_start + input_scale_buffer_offset);
std::memset(token_num_per_group_buffer, 0, expert_num * sizeof(int32_t));
for (int32_t i = 0; i < expanded_token_num; ++i) {
++token_num_per_group_buffer[topk_id[i]];
}
int32_t token_num_sum = 0;
cu_token_num_per_group_buffer[0] = 0;
int32_t* token_index_buffer = cu_token_num_per_group_buffer + 1;
for (int32_t i = 0; i < expert_num; ++i) {
token_index_buffer[i] = token_num_sum;
token_num_sum += token_num_per_group_buffer[i];
}
for (int32_t i = 0; i < token_num; ++i) {
const int32_t* curr_topk_id = topk_id + i * topk_num;
int32_t* curr_index_buffer = expand_token_id_index_buffer + i * topk_num;
for (int32_t j = 0; j < topk_num; ++j) {
const int32_t curr_expert_id = curr_topk_id[j];
const int32_t curr_index = token_index_buffer[curr_expert_id]++;
expand_token_id_buffer[curr_index] = i;
curr_index_buffer[j] = curr_index;
}
}
// quantize inputs
#pragma omp parallel for
for (int32_t token_idx = 0; token_idx < token_num; ++token_idx) {
gemm_t::quantize_row(input + token_idx * input_size_13,
input_quant_buffer + token_idx * input_size_13,
input_scale_buffer[token_idx], input_size_13);
}
{
alignas(64) cpu_utils::Counter counter;
cpu_utils::Counter* counter_ptr = &counter;
// w13 GEMM + act
#pragma omp parallel for schedule(static, 1)
for (int32_t thread_id = 0; thread_id < thread_num; ++thread_id) {
const int32_t task_num_per_expert =
(output_size_13 + w13_n_tile_size - 1) / w13_n_tile_size;
const int32_t task_num = task_num_per_expert * expert_num;
uint8_t* __restrict__ thread_buffer =
thread_buffer_start + thread_id * thread_buffer_size;
int8_t* __restrict__ gemm_input_buffer =
reinterpret_cast<int8_t*>(thread_buffer + gemm_input_buffer_offset);
float* __restrict__ gemm_output_buffer =
reinterpret_cast<float*>(thread_buffer + gemm_output_buffer_offset);
auto* __restrict__ w13_gemm_output_buffer = reinterpret_cast<scalar_t*>(
common_buffer_start + w13_gemm_output_buffer_offset);
gemm_t gemm;
const int32_t w13_n_group_stride =
gemm_t::WeightOCGroupSize * input_size_13;
const int32_t w13_n_tile_stride = gemm_n_tile_size * input_size_13;
for (;;) {
const int32_t task_id = counter_ptr->acquire_counter();
if (task_id >= task_num) {
break;
}
const int32_t curr_expert_id = task_id / task_num_per_expert;
const int32_t curr_output_group_id = task_id % task_num_per_expert;
const int32_t curr_token_num =
token_num_per_group_buffer[curr_expert_id];
if (curr_token_num == 0) {
continue;
}
const int32_t actual_n_tile_size =
std::min(w13_n_tile_size,
output_size_13 - curr_output_group_id * w13_n_tile_size);
const int32_t* __restrict__ curr_expand_token_id_buffer =
expand_token_id_buffer +
cu_token_num_per_group_buffer[curr_expert_id];
scalar_t* __restrict__ curr_w13_gemm_output_buffer =
w13_gemm_output_buffer +
cu_token_num_per_group_buffer[curr_expert_id] * input_size_2 +
curr_output_group_id * w13_n_tile_size / 2;
const int8_t* w13_weight_ptr_0 = nullptr;
const int8_t* w13_weight_ptr_1 = nullptr;
const float* w13_scale_ptr_0 = nullptr;
const float* w13_scale_ptr_1 = nullptr;
scalar_t* w13_bias_ptr_0 = nullptr;
scalar_t* w13_bias_ptr_1 = nullptr;
if (act_type == FusedMOEAct::SwigluOAIAndMul) {
const int32_t output_offset = curr_output_group_id * w13_n_tile_size;
w13_weight_ptr_0 = w13 +
curr_expert_id * input_size_13 * output_size_13 +
output_offset * input_size_13;
w13_weight_ptr_1 =
w13_weight_ptr_0 + actual_n_tile_size / 2 * input_size_13;
w13_scale_ptr_0 =
w13_scales + curr_expert_id * output_size_13 + output_offset;
w13_scale_ptr_1 = w13_scale_ptr_0 + actual_n_tile_size / 2;
if (w13_bias != nullptr) {
w13_bias_ptr_0 =
w13_bias + curr_expert_id * output_size_13 + output_offset;
w13_bias_ptr_1 = w13_bias_ptr_0 + actual_n_tile_size / 2;
}
} else {
const int32_t output_offset =
curr_output_group_id * (w13_n_tile_size / 2);
w13_weight_ptr_0 = w13 +
curr_expert_id * input_size_13 * output_size_13 +
output_offset * input_size_13;
w13_weight_ptr_1 =
w13_weight_ptr_0 + output_size_13 / 2 * input_size_13;
w13_scale_ptr_0 =
w13_scales + curr_expert_id * output_size_13 + output_offset;
w13_scale_ptr_1 = w13_scale_ptr_0 + output_size_13 / 2;
if (w13_bias != nullptr) {
w13_bias_ptr_0 =
w13_bias + curr_expert_id * output_size_13 + output_offset;
w13_bias_ptr_1 = w13_bias_ptr_0 + output_size_13 / 2;
}
}
for (int32_t token_idx = 0; token_idx < curr_token_num;
token_idx += gemm_m_tile_size) {
const int32_t actual_token_num =
std::min(gemm_m_tile_size, curr_token_num - token_idx);
const int8_t* input_rows[gemm_m_tile_size];
alignas(64) float input_scales[gemm_m_tile_size];
// gather and pack
for (int32_t i = 0; i < actual_token_num; ++i) {
const int32_t curr_token_id = curr_expand_token_id_buffer[i];
input_rows[i] = input_quant_buffer + curr_token_id * input_size_13;
input_scales[i] = input_scale_buffer[curr_token_id];
}
gemm_t::pack_input_from_rows(input_rows, gemm_input_buffer,
actual_token_num, input_size_13);
curr_expand_token_id_buffer += actual_token_num;
const int8_t* w13_weight_ptr_0_iter = w13_weight_ptr_0;
const int8_t* w13_weight_ptr_1_iter = w13_weight_ptr_1;
const float* w13_scale_ptr_0_iter = w13_scale_ptr_0;
const float* w13_scale_ptr_1_iter = w13_scale_ptr_1;
scalar_t* w13_bias_ptr_0_iter = w13_bias_ptr_0;
scalar_t* w13_bias_ptr_1_iter = w13_bias_ptr_1;
float* w13_output_buffer_0_iter = gemm_output_buffer;
float* w13_output_buffer_1_iter =
gemm_output_buffer + actual_n_tile_size / 2;
for (int32_t i = 0; i < actual_n_tile_size;
i += min_w13_n_tile_size) {
auto* output_0_int32 =
reinterpret_cast<int32_t*>(w13_output_buffer_0_iter);
gemm.gemm(gemm_input_buffer, w13_weight_ptr_0_iter, output_0_int32,
actual_token_num, input_size_13, w13_n_group_stride,
actual_n_tile_size);
gemm_t::dequantize_tile(output_0_int32, w13_output_buffer_0_iter,
input_scales, w13_scale_ptr_0_iter,
actual_token_num, gemm_n_tile_size,
actual_n_tile_size);
if (w13_bias != nullptr) {
cpu_micro_gemm::add_bias_epilogue<gemm_n_tile_size>(
w13_output_buffer_0_iter, w13_output_buffer_0_iter,
w13_bias_ptr_0_iter, actual_token_num, actual_n_tile_size,
actual_n_tile_size);
w13_bias_ptr_0_iter += gemm_n_tile_size;
}
auto* output_1_int32 =
reinterpret_cast<int32_t*>(w13_output_buffer_1_iter);
gemm.gemm(gemm_input_buffer, w13_weight_ptr_1_iter, output_1_int32,
actual_token_num, input_size_13, w13_n_group_stride,
actual_n_tile_size);
gemm_t::dequantize_tile(output_1_int32, w13_output_buffer_1_iter,
input_scales, w13_scale_ptr_1_iter,
actual_token_num, gemm_n_tile_size,
actual_n_tile_size);
if (w13_bias != nullptr) {
cpu_micro_gemm::add_bias_epilogue<gemm_n_tile_size>(
w13_output_buffer_1_iter, w13_output_buffer_1_iter,
w13_bias_ptr_1_iter, actual_token_num, actual_n_tile_size,
actual_n_tile_size);
w13_bias_ptr_1_iter += gemm_n_tile_size;
}
w13_weight_ptr_0_iter += w13_n_tile_stride;
w13_weight_ptr_1_iter += w13_n_tile_stride;
w13_scale_ptr_0_iter += gemm_n_tile_size;
w13_scale_ptr_1_iter += gemm_n_tile_size;
w13_output_buffer_0_iter += gemm_n_tile_size;
w13_output_buffer_1_iter += gemm_n_tile_size;
}
apply_gated_act(act_type, gemm_output_buffer,
curr_w13_gemm_output_buffer, actual_token_num,
actual_n_tile_size, actual_n_tile_size, input_size_2);
curr_w13_gemm_output_buffer += gemm_m_tile_size * input_size_2;
}
}
}
}
auto* __restrict__ w13_gemm_output_buffer = reinterpret_cast<scalar_t*>(
common_buffer_start + w13_gemm_output_buffer_offset);
float* __restrict__ w13_output_scale_buffer = reinterpret_cast<float*>(
common_buffer_start + w13_output_scale_buffer_offset);
// quantize w2 inputs - in place
#pragma omp parallel for
for (int32_t token_idx = 0; token_idx < expanded_token_num; ++token_idx) {
scalar_t* input_row = w13_gemm_output_buffer + token_idx * input_size_2;
int8_t* output_row = reinterpret_cast<int8_t*>(input_row);
gemm_t::quantize_row(input_row, output_row,
w13_output_scale_buffer[token_idx], input_size_2);
}
{
alignas(64) cpu_utils::Counter counter;
cpu_utils::Counter* counter_ptr = &counter;
// w2 gemm
#pragma omp parallel for schedule(static, 1)
for (int32_t thread_id = 0; thread_id < thread_num; ++thread_id) {
const int32_t task_num_per_expert =
(output_size_2 + w2_n_tile_size - 1) / w2_n_tile_size;
const int32_t task_num = task_num_per_expert * expert_num;
uint8_t* __restrict__ thread_buffer =
thread_buffer_start + thread_id * thread_buffer_size;
int8_t* __restrict__ gemm_input_buffer =
reinterpret_cast<int8_t*>(thread_buffer + gemm_input_buffer_offset);
float* __restrict__ gemm_output_buffer =
reinterpret_cast<float*>(thread_buffer + gemm_output_buffer_offset);
float* __restrict__ w2_gemm_output_buffer = reinterpret_cast<float*>(
common_buffer_start + w2_gemm_output_buffer_offset);
gemm_t gemm;
const int32_t w2_n_group_stride =
gemm_t::WeightOCGroupSize * input_size_2;
const int32_t w2_n_tile_stride = gemm_n_tile_size * input_size_2;
for (;;) {
const int32_t task_id = counter_ptr->acquire_counter();
if (task_id >= task_num) {
break;
}
const int32_t curr_expert_id = task_id / task_num_per_expert;
const int32_t curr_output_group_id = task_id % task_num_per_expert;
const int32_t curr_token_num =
token_num_per_group_buffer[curr_expert_id];
if (curr_token_num == 0) {
continue;
}
const int32_t actual_n_tile_size =
std::min(w2_n_tile_size,
output_size_2 - curr_output_group_id * w2_n_tile_size);
scalar_t* __restrict__ curr_w13_gemm_output_buffer =
w13_gemm_output_buffer +
cu_token_num_per_group_buffer[curr_expert_id] * input_size_2;
float* __restrict__ curr_w13_output_scale_buffer =
w13_output_scale_buffer +
cu_token_num_per_group_buffer[curr_expert_id];
float* __restrict__ curr_w2_gemm_output_buffer =
w2_gemm_output_buffer +
cu_token_num_per_group_buffer[curr_expert_id] * output_size_2 +
curr_output_group_id * w2_n_tile_size;
const int8_t* __restrict__ w2_weight_ptr =
w2 + curr_expert_id * output_size_2 * input_size_2 +
curr_output_group_id * w2_n_tile_size * input_size_2;
const float* __restrict__ w2_scale_ptr =
w2_scales + curr_expert_id * output_size_2 +
curr_output_group_id * w2_n_tile_size;
scalar_t* w2_bias_ptr = nullptr;
if (w2_bias != nullptr) {
w2_bias_ptr = w2_bias + curr_expert_id * output_size_2 +
curr_output_group_id * w2_n_tile_size;
}
for (int32_t token_idx = 0; token_idx < curr_token_num;
token_idx += gemm_m_tile_size) {
const int32_t actual_token_num =
std::min(gemm_m_tile_size, curr_token_num - token_idx);
const int8_t* input_rows[gemm_m_tile_size];
alignas(64) float input_scales[gemm_m_tile_size];
for (int32_t i = 0; i < actual_token_num; ++i) {
input_rows[i] = reinterpret_cast<const int8_t*>(
curr_w13_gemm_output_buffer + i * input_size_2);
input_scales[i] = curr_w13_output_scale_buffer[i];
}
gemm_t::pack_input_from_rows(input_rows, gemm_input_buffer,
actual_token_num, input_size_2);
const int8_t* w2_weight_ptr_iter = w2_weight_ptr;
const float* w2_scale_ptr_iter = w2_scale_ptr;
scalar_t* w2_bias_ptr_iter = w2_bias_ptr;
float* curr_w2_gemm_output_buffer_iter = curr_w2_gemm_output_buffer;
for (int32_t i = 0; i < actual_n_tile_size; i += gemm_n_tile_size) {
auto* output_int32 = reinterpret_cast<int32_t*>(gemm_output_buffer);
gemm.gemm(gemm_input_buffer, w2_weight_ptr_iter, output_int32,
actual_token_num, input_size_2, w2_n_group_stride,
gemm_n_tile_size);
gemm_t::dequantize_tile(output_int32, gemm_output_buffer,
input_scales, w2_scale_ptr_iter,
actual_token_num, gemm_n_tile_size,
gemm_n_tile_size);
if (w2_bias != nullptr) {
cpu_micro_gemm::add_bias_epilogue<gemm_n_tile_size>(
gemm_output_buffer, gemm_output_buffer, w2_bias_ptr_iter,
actual_token_num, gemm_n_tile_size, gemm_n_tile_size);
w2_bias_ptr_iter += gemm_n_tile_size;
}
for (int32_t m_idx = 0; m_idx < actual_token_num; ++m_idx) {
std::memcpy(
curr_w2_gemm_output_buffer_iter + m_idx * output_size_2,
gemm_output_buffer + m_idx * gemm_n_tile_size,
gemm_n_tile_size * sizeof(float));
}
w2_weight_ptr_iter += w2_n_tile_stride;
w2_scale_ptr_iter += gemm_n_tile_size;
curr_w2_gemm_output_buffer_iter += gemm_n_tile_size;
}
curr_w13_gemm_output_buffer += gemm_m_tile_size * input_size_2;
curr_w13_output_scale_buffer += gemm_m_tile_size;
curr_w2_gemm_output_buffer += gemm_m_tile_size * output_size_2;
}
}
}
}
{
alignas(64) cpu_utils::Counter counter;
cpu_utils::Counter* counter_ptr = &counter;
#pragma omp parallel for schedule(static, 1)
for (int32_t thread_id = 0; thread_id < thread_num; ++thread_id) {
uint8_t* __restrict__ thread_buffer =
thread_buffer_start + thread_id * thread_buffer_size;
float* __restrict__ ws_output_buffer =
reinterpret_cast<float*>(thread_buffer + ws_output_buffer_offset);
float* __restrict__ w2_gemm_output_buffer = reinterpret_cast<float*>(
common_buffer_start + w2_gemm_output_buffer_offset);
for (;;) {
const int32_t token_id = counter_ptr->acquire_counter();
if (token_id >= token_num) {
break;
}
int32_t* __restrict__ curr_expand_token_id_index_buffer =
expand_token_id_index_buffer + token_id * topk_num;
const float* __restrict__ curr_weight =
topk_weights + token_id * topk_num;
const float first_weight = skip_weighted ? 1.0f : curr_weight[0];
scalar_t* __restrict__ curr_output_buffer =
output + token_id * output_size_2;
if (topk_num > 1) {
int32_t w2_output_idx = curr_expand_token_id_index_buffer[0];
float* w2_output_iter =
w2_gemm_output_buffer + w2_output_idx * output_size_2;
float* ws_output_buffer_iter = ws_output_buffer;
vec_op::FP32Vec16 weight_vec(first_weight);
for (int32_t i = 0; i < output_size_2; i += 16) {
vec_op::FP32Vec16 vec(w2_output_iter);
(vec * weight_vec).save(ws_output_buffer_iter);
w2_output_iter += 16;
ws_output_buffer_iter += 16;
}
for (int32_t idx = 1; idx < topk_num - 1; ++idx) {
w2_output_idx = curr_expand_token_id_index_buffer[idx];
w2_output_iter =
w2_gemm_output_buffer + w2_output_idx * output_size_2;
ws_output_buffer_iter = ws_output_buffer;
weight_vec = vec_op::FP32Vec16(curr_weight[idx]);
for (int32_t i = 0; i < output_size_2; i += 16) {
vec_op::FP32Vec16 vec(w2_output_iter);
vec_op::FP32Vec16 sum(ws_output_buffer_iter);
(sum + vec * weight_vec).save(ws_output_buffer_iter);
w2_output_iter += 16;
ws_output_buffer_iter += 16;
}
}
const int32_t last_idx = topk_num - 1;
w2_output_idx = curr_expand_token_id_index_buffer[last_idx];
w2_output_iter =
w2_gemm_output_buffer + w2_output_idx * output_size_2;
ws_output_buffer_iter = ws_output_buffer;
scalar_t* curr_output_buffer_iter = curr_output_buffer;
weight_vec = vec_op::FP32Vec16(curr_weight[last_idx]);
for (int32_t i = 0; i < output_size_2; i += 16) {
vec_op::FP32Vec16 vec(w2_output_iter);
vec_op::FP32Vec16 sum(ws_output_buffer_iter);
scalar_vec_t(sum + vec * weight_vec).save(curr_output_buffer_iter);
w2_output_iter += 16;
ws_output_buffer_iter += 16;
curr_output_buffer_iter += 16;
}
} else {
const int32_t w2_output_idx = curr_expand_token_id_index_buffer[0];
float* w2_output_iter =
w2_gemm_output_buffer + w2_output_idx * output_size_2;
scalar_t* curr_output_buffer_iter = curr_output_buffer;
vec_op::FP32Vec16 weight_vec(first_weight);
for (int32_t i = 0; i < output_size_2; i += 16) {
vec_op::FP32Vec16 vec(w2_output_iter);
scalar_vec_t(vec * weight_vec).save(curr_output_buffer_iter);
w2_output_iter += 16;
curr_output_buffer_iter += 16;
}
}
}
}
}
}
} // namespace
void prepack_moe_weight_int8(
const torch::Tensor& weight, // [expert_num, output_size, input_size]
torch::Tensor& packed_weight, const std::string& isa) {
TORCH_CHECK(weight.is_contiguous());
const int32_t expert_num = weight.size(0);
const int32_t output_size = weight.size(1);
const int32_t input_size = weight.size(2);
const int64_t expert_stride = weight.stride(0);
const cpu_utils::ISA isa_type = cpu_utils::get_isa(isa);
TORCH_CHECK_EQ(output_size % 32, 0);
CPU_INT8_ISA_DISPATCH_IMPL(isa_type, c10::BFloat16, [&]() {
TORCH_CHECK_EQ(input_size % gemm_t::K, 0);
prepack_moe_weight_int8_impl<gemm_t>(
weight.data_ptr<int8_t>(), packed_weight.data_ptr<int8_t>(), expert_num,
output_size, input_size, expert_stride);
});
}
void cpu_fused_moe_int8(torch::Tensor& output, const torch::Tensor& input,
const torch::Tensor& w13, const torch::Tensor& w2,
const torch::Tensor& w13_scale,
const torch::Tensor& w2_scale,
const std::optional<torch::Tensor>& w13_bias,
const std::optional<torch::Tensor>& w2_bias,
const torch::Tensor& topk_weights,
const torch::Tensor& topk_id, const bool skip_weighted,
const std::string& act, const std::string& isa) {
const int32_t token_num = input.size(0);
const int32_t input_size_13 = input.size(1);
const int64_t input_stride = input.stride(0);
TORCH_CHECK_EQ(input_stride, input_size_13);
const int32_t expert_num = w13.size(0);
const int32_t output_size_13 = w13.size(1);
const int32_t input_size_2 = w2.size(2);
const int32_t output_size_2 = w2.size(1);
const int32_t topk_num = topk_id.size(1);
const FusedMOEAct act_type = cpu_fused_moe_utils::get_act_type(act);
const cpu_utils::ISA isa_type = cpu_utils::get_isa(isa);
TORCH_CHECK(!skip_weighted || topk_num == 1,
"skip_weighted is only supported for topk=1 on CPU");
VLLM_DISPATCH_FLOATING_TYPES(
input.scalar_type(), "cpu_fused_moe_int8", [&]() {
CPU_INT8_ISA_DISPATCH_IMPL(isa_type, scalar_t, [&]() {
fused_moe_int8_impl<scalar_t, gemm_t>(
output.data_ptr<scalar_t>(), input.data_ptr<scalar_t>(),
w13.data_ptr<int8_t>(), w2.data_ptr<int8_t>(),
w13_scale.data_ptr<float>(), w2_scale.data_ptr<float>(),
w13_bias.has_value() ? w13_bias->data_ptr<scalar_t>() : nullptr,
w2_bias.has_value() ? w2_bias->data_ptr<scalar_t>() : nullptr,
topk_weights.data_ptr<float>(), topk_id.data_ptr<int32_t>(),
act_type, token_num, expert_num, topk_num, input_size_13,
output_size_13, input_size_2, output_size_2, skip_weighted);
});
});
}
+12 -3
View File
@@ -287,7 +287,7 @@ struct FP32Vec4 : public Vec<FP32Vec4> {
explicit FP32Vec4(__vector float data) : reg(data) {}
explicit FP32Vec4(const FP32Vec4& data) : reg(data.reg) {}
FP32Vec4(const FP32Vec4& data) : reg(data.reg) {}
};
struct FP32Vec8 : public Vec<FP32Vec8> {
@@ -316,7 +316,7 @@ struct FP32Vec8 : public Vec<FP32Vec8> {
explicit FP32Vec8(f32x4x2_t data) : reg(data) {}
explicit FP32Vec8(const FP32Vec8& data) {
FP32Vec8(const FP32Vec8& data) {
reg.val[0] = data.reg.val[0];
reg.val[1] = data.reg.val[1];
}
@@ -593,7 +593,7 @@ struct FP32Vec16 : public Vec<FP32Vec16> {
explicit FP32Vec16(bool, const float* ptr) : FP32Vec16(ptr) {}
explicit FP32Vec16(f32x4x4_t data) : reg(data) {}
explicit FP32Vec16(const FP32Vec16& data) {
FP32Vec16(const FP32Vec16& data) {
reg.val[0] = data.reg.val[0];
reg.val[1] = data.reg.val[1];
reg.val[2] = data.reg.val[2];
@@ -747,6 +747,15 @@ struct FP32Vec16 : public Vec<FP32Vec16> {
vec_abs(reg.val[2]), vec_abs(reg.val[3])}));
}
FP32Vec16 exp() const {
FP32Vec8 lo(f32x4x2_t{reg.val[0], reg.val[1]});
FP32Vec8 hi(f32x4x2_t{reg.val[2], reg.val[3]});
auto lo_e = lo.exp();
auto hi_e = hi.exp();
return FP32Vec16(f32x4x4_t{lo_e.reg.val[0], lo_e.reg.val[1],
hi_e.reg.val[0], hi_e.reg.val[1]});
}
float reduce_max() {
__vector float max01 = vec_max(reg.val[0], reg.val[1]);
__vector float max23 = vec_max(reg.val[2], reg.val[3]);
@@ -31,6 +31,9 @@ class MicroGemm {
}
};
template <cpu_utils::ISA isa, typename scalar_t>
class MicroGemmINT8;
template <int32_t n_size, typename scalar_t>
FORCE_INLINE void default_epilogue(float* __restrict__ c_ptr,
scalar_t* __restrict__ d_ptr,
@@ -0,0 +1,424 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
#ifndef CPU_MICRO_GEMM_INT8_NEON_HPP
#define CPU_MICRO_GEMM_INT8_NEON_HPP
#include <algorithm>
#include <cstdint>
#include "cpu/micro_gemm/cpu_micro_gemm_impl.hpp"
#include <arm_bf16.h>
#include <arm_neon.h>
#include <c10/util/BFloat16.h>
#include <c10/util/Exception.h>
#include <c10/util/Half.h>
namespace cpu_micro_gemm {
namespace neon_smmla {
constexpr int32_t K = 8;
constexpr int32_t Cols = 2;
constexpr int32_t TileSize = K * Cols;
FORCE_INLINE float32x4x2_t load_as_f32(const float* input) {
float32x4x2_t result;
result.val[0] = vld1q_f32(input);
result.val[1] = vld1q_f32(input + 4);
return result;
}
FORCE_INLINE float32x4x2_t load_as_f32(const c10::Half* input) {
const auto input_vec = vld1q_f16(reinterpret_cast<const float16_t*>(input));
float32x4x2_t result;
result.val[0] = vcvt_f32_f16(vget_low_f16(input_vec));
result.val[1] = vcvt_f32_f16(vget_high_f16(input_vec));
return result;
}
FORCE_INLINE float32x4x2_t load_as_f32(const c10::BFloat16* input) {
const auto input_vec = vld1q_bf16(reinterpret_cast<const bfloat16_t*>(input));
float32x4x2_t result;
result.val[0] = vcvt_f32_bf16(vget_low_bf16(input_vec));
result.val[1] = vcvt_f32_bf16(vget_high_bf16(input_vec));
return result;
}
FORCE_INLINE void store_acc_rowpair(const int32x4_t acc01,
const int32x4_t acc23,
const int32x4_t acc45,
const int32x4_t acc67,
int32_t* __restrict__ c_ptr,
const int64_t ldc, const int32_t m_rows) {
if (m_rows == 0) {
return;
}
vst1q_s32(c_ptr, vcombine_s32(vget_low_s32(acc01), vget_low_s32(acc23)));
vst1q_s32(c_ptr + 4, vcombine_s32(vget_low_s32(acc45), vget_low_s32(acc67)));
if (m_rows == 2) {
vst1q_s32(c_ptr + ldc,
vcombine_s32(vget_high_s32(acc01), vget_high_s32(acc23)));
vst1q_s32(c_ptr + ldc + 4,
vcombine_s32(vget_high_s32(acc45), vget_high_s32(acc67)));
}
}
FORCE_INLINE void gemm_micro_smmla_8x8_packed_a(
const int8_t* __restrict__ a_packed, const int8_t* __restrict__ b_packed,
int32_t* __restrict__ c_ptr, const int32_t m, const int32_t k_size,
const int64_t ldc) {
const int32x4_t zero = vdupq_n_s32(0);
int32x4_t acc0101 = zero, acc0123 = zero, acc0145 = zero, acc0167 = zero;
int32x4_t acc2301 = zero, acc2323 = zero, acc2345 = zero, acc2367 = zero;
int32x4_t acc4501 = zero, acc4523 = zero, acc4545 = zero, acc4567 = zero;
int32x4_t acc6701 = zero, acc6723 = zero, acc6745 = zero, acc6767 = zero;
const int8_t* __restrict__ a_tile = a_packed;
const int8_t* __restrict__ b_tile = b_packed;
#pragma GCC unroll 8
for (int32_t k_idx = 0; k_idx < k_size; k_idx += K) {
const int8x16_t a_tile01 = vld1q_s8(a_tile);
const int8x16_t a_tile23 = vld1q_s8(a_tile + TileSize);
const int8x16_t a_tile45 = vld1q_s8(a_tile + 2 * TileSize);
const int8x16_t a_tile67 = vld1q_s8(a_tile + 3 * TileSize);
const int8x16_t b_tile01 = vld1q_s8(b_tile);
const int8x16_t b_tile23 = vld1q_s8(b_tile + TileSize);
const int8x16_t b_tile45 = vld1q_s8(b_tile + 2 * TileSize);
const int8x16_t b_tile67 = vld1q_s8(b_tile + 3 * TileSize);
acc0101 = vmmlaq_s32(acc0101, a_tile01, b_tile01);
acc2301 = vmmlaq_s32(acc2301, a_tile23, b_tile01);
acc4501 = vmmlaq_s32(acc4501, a_tile45, b_tile01);
acc6701 = vmmlaq_s32(acc6701, a_tile67, b_tile01);
acc0123 = vmmlaq_s32(acc0123, a_tile01, b_tile23);
acc2323 = vmmlaq_s32(acc2323, a_tile23, b_tile23);
acc4523 = vmmlaq_s32(acc4523, a_tile45, b_tile23);
acc6723 = vmmlaq_s32(acc6723, a_tile67, b_tile23);
acc0145 = vmmlaq_s32(acc0145, a_tile01, b_tile45);
acc2345 = vmmlaq_s32(acc2345, a_tile23, b_tile45);
acc4545 = vmmlaq_s32(acc4545, a_tile45, b_tile45);
acc6745 = vmmlaq_s32(acc6745, a_tile67, b_tile45);
acc0167 = vmmlaq_s32(acc0167, a_tile01, b_tile67);
acc2367 = vmmlaq_s32(acc2367, a_tile23, b_tile67);
acc4567 = vmmlaq_s32(acc4567, a_tile45, b_tile67);
acc6767 = vmmlaq_s32(acc6767, a_tile67, b_tile67);
a_tile += 4 * TileSize;
b_tile += 4 * TileSize;
}
store_acc_rowpair(acc0101, acc0123, acc0145, acc0167, c_ptr, ldc,
std::min(2, m));
store_acc_rowpair(acc2301, acc2323, acc2345, acc2367, c_ptr + 2 * ldc, ldc,
std::min(2, std::max(0, m - 2)));
store_acc_rowpair(acc4501, acc4523, acc4545, acc4567, c_ptr + 4 * ldc, ldc,
std::min(2, std::max(0, m - 4)));
store_acc_rowpair(acc6701, acc6723, acc6745, acc6767, c_ptr + 6 * ldc, ldc,
std::min(2, std::max(0, m - 6)));
}
FORCE_INLINE void gemm_micro_smmla_4x16_packed_a(
const int8_t* __restrict__ a_packed, const int8_t* __restrict__ b_packed,
int32_t* __restrict__ c_ptr, const int32_t m, const int32_t k_size,
const int64_t b_n_group_stride, const int64_t ldc) {
const int32_t m_rows_01 = std::min(2, m);
const int32_t m_rows_23 = std::min(2, std::max(0, m - 2));
const int32x4_t zero = vdupq_n_s32(0);
int32x4_t acc0101 = zero, acc0123 = zero, acc0145 = zero, acc0167 = zero;
int32x4_t acc2301 = zero, acc2323 = zero, acc2345 = zero, acc2367 = zero;
int32x4_t acc0189 = zero, acc011011 = zero, acc011213 = zero,
acc011415 = zero;
int32x4_t acc2389 = zero, acc231011 = zero, acc231213 = zero,
acc231415 = zero;
const int8_t* __restrict__ a_tile = a_packed;
// note: b packs 8 panels contiguously, so we need 2 b_tile ptrs
// for the 4x16 microkernel
const int8_t* __restrict__ b_tile0 = b_packed;
const int8_t* __restrict__ b_tile1 = b_packed + b_n_group_stride;
#pragma GCC unroll 8
for (int32_t k_idx = 0; k_idx < k_size; k_idx += K) {
const int8x16_t a_tile01 = vld1q_s8(a_tile);
const int8x16_t a_tile23 = vld1q_s8(a_tile + TileSize);
const int8x16_t b_tile01 = vld1q_s8(b_tile0);
const int8x16_t b_tile23 = vld1q_s8(b_tile0 + TileSize);
const int8x16_t b_tile45 = vld1q_s8(b_tile0 + 2 * TileSize);
const int8x16_t b_tile67 = vld1q_s8(b_tile0 + 3 * TileSize);
const int8x16_t b_tile89 = vld1q_s8(b_tile1);
const int8x16_t b_tile1011 = vld1q_s8(b_tile1 + TileSize);
const int8x16_t b_tile1213 = vld1q_s8(b_tile1 + 2 * TileSize);
const int8x16_t b_tile1415 = vld1q_s8(b_tile1 + 3 * TileSize);
acc0101 = vmmlaq_s32(acc0101, a_tile01, b_tile01);
acc2301 = vmmlaq_s32(acc2301, a_tile23, b_tile01);
acc0123 = vmmlaq_s32(acc0123, a_tile01, b_tile23);
acc2323 = vmmlaq_s32(acc2323, a_tile23, b_tile23);
acc0145 = vmmlaq_s32(acc0145, a_tile01, b_tile45);
acc2345 = vmmlaq_s32(acc2345, a_tile23, b_tile45);
acc0167 = vmmlaq_s32(acc0167, a_tile01, b_tile67);
acc2367 = vmmlaq_s32(acc2367, a_tile23, b_tile67);
acc0189 = vmmlaq_s32(acc0189, a_tile01, b_tile89);
acc2389 = vmmlaq_s32(acc2389, a_tile23, b_tile89);
acc011011 = vmmlaq_s32(acc011011, a_tile01, b_tile1011);
acc231011 = vmmlaq_s32(acc231011, a_tile23, b_tile1011);
acc011213 = vmmlaq_s32(acc011213, a_tile01, b_tile1213);
acc231213 = vmmlaq_s32(acc231213, a_tile23, b_tile1213);
acc011415 = vmmlaq_s32(acc011415, a_tile01, b_tile1415);
acc231415 = vmmlaq_s32(acc231415, a_tile23, b_tile1415);
a_tile += 2 * TileSize;
b_tile0 += 4 * TileSize;
b_tile1 += 4 * TileSize;
}
// rows 0-1, columns 0-7
store_acc_rowpair(acc0101, acc0123, acc0145, acc0167, c_ptr, ldc, m_rows_01);
// rows 0-1, columns 8-15
store_acc_rowpair(acc0189, acc011011, acc011213, acc011415, c_ptr + 8, ldc,
m_rows_01);
// rows 2-3, columns 0-7
store_acc_rowpair(acc2301, acc2323, acc2345, acc2367, c_ptr + 2 * ldc, ldc,
m_rows_23);
// rows 2-3, columns 8-15
store_acc_rowpair(acc2389, acc231011, acc231213, acc231415,
c_ptr + 2 * ldc + 8, ldc, m_rows_23);
}
} // namespace neon_smmla
template <typename scalar_t>
class MicroGemmINT8<cpu_utils::ISA::NEON, scalar_t> {
public:
static constexpr int32_t K = neon_smmla::K;
static constexpr int32_t Mr = 8;
static constexpr int32_t Nr = 8;
static constexpr int32_t NrGemv = 16;
static constexpr int32_t MaxMSize = 8;
static constexpr int32_t NSize = 32;
static constexpr int32_t WeightOCGroupSize = Nr;
static_assert(MaxMSize % Mr == 0);
static FORCE_INLINE void quantize_row(const scalar_t* input, int8_t* output,
float& scale, const int32_t size) {
TORCH_CHECK_EQ(size % K, 0);
float32x4_t max_vec = vdupq_n_f32(0.0f);
for (int32_t i = 0; i < size; i += K) {
const float32x4x2_t input_vec = neon_smmla::load_as_f32(input + i);
max_vec = vmaxq_f32(max_vec, vabsq_f32(input_vec.val[0]));
max_vec = vmaxq_f32(max_vec, vabsq_f32(input_vec.val[1]));
}
const float abs_max = std::max(vmaxvq_f32(max_vec), 1.0e-7f);
scale = abs_max / 127.0f;
const float32x4_t inv_scale_vec = vdupq_n_f32(127.0f / abs_max);
for (int32_t i = 0; i < size; i += K) {
const float32x4x2_t input_vec = neon_smmla::load_as_f32(input + i);
const int32x4_t output_low =
vcvtnq_s32_f32(vmulq_f32(input_vec.val[0], inv_scale_vec));
const int32x4_t output_high =
vcvtnq_s32_f32(vmulq_f32(input_vec.val[1], inv_scale_vec));
const int16x8_t output_s16 =
vcombine_s16(vqmovn_s32(output_low), vqmovn_s32(output_high));
vst1_s8(output + i, vqmovn_s16(output_s16));
}
}
// with current code, fusing this into the gemm micro kernel didn't move the
// needle
static FORCE_INLINE void dequantize_tile(
int32_t* input, float* output, const float* __restrict__ input_scales,
const float* __restrict__ weight_scales, const int32_t m, const int32_t n,
const int32_t stride) {
TORCH_CHECK_EQ(n % 4, 0);
for (int32_t m_idx = 0; m_idx < m; ++m_idx) {
const float32x4_t input_scale_vec = vdupq_n_f32(input_scales[m_idx]);
for (int32_t n_idx = 0; n_idx < n; n_idx += 4) {
const int32x4_t input_vec = vld1q_s32(input + m_idx * stride + n_idx);
const float32x4_t weight_scale_vec = vld1q_f32(weight_scales + n_idx);
const float32x4_t output_vec =
vmulq_f32(vcvtq_f32_s32(input_vec),
vmulq_f32(input_scale_vec, weight_scale_vec));
vst1q_f32(output + m_idx * stride + n_idx, output_vec);
}
}
}
// physical layout [
// M / (8 or 4); Mr is 8 or 4
// K / 8; K for smmla is 8
// 4, ; 4 row-pairs for each 8 rows
// 2, ; row-pair is 2 rows
// 4 ; 4 elements per row
// ]
static void pack_input_from_rows(const int8_t* const* __restrict__ rows,
int8_t* __restrict__ a_packed,
const int32_t m, const int32_t k) {
TORCH_CHECK(m > 0 && m <= MaxMSize);
TORCH_CHECK(k % K == 0);
const int8x8_t zero = vdup_n_s8(0);
for (int32_t row_base = 0; row_base < m; row_base += Mr) {
const int32_t panel_m = std::min(Mr, m - row_base);
const int8_t* const* panel_rows = rows + row_base;
int8_t* __restrict__ out = a_packed + row_base * k;
// fast path for full 8-row panels (fast path for 4-row panels didn't move
// the needle)
if (panel_m == Mr) {
const int8_t* __restrict__ row0 = panel_rows[0];
const int8_t* __restrict__ row1 = panel_rows[1];
const int8_t* __restrict__ row2 = panel_rows[2];
const int8_t* __restrict__ row3 = panel_rows[3];
const int8_t* __restrict__ row4 = panel_rows[4];
const int8_t* __restrict__ row5 = panel_rows[5];
const int8_t* __restrict__ row6 = panel_rows[6];
const int8_t* __restrict__ row7 = panel_rows[7];
int32_t k_idx = 0;
for (; k_idx + 2 * K <= k; k_idx += 2 * K) {
int8_t* __restrict__ block0 = out;
int8_t* __restrict__ block1 = out + 4 * neon_smmla::TileSize;
int8x16_t a0 = vld1q_s8(row0 + k_idx);
int8x16_t a1 = vld1q_s8(row1 + k_idx);
vst1q_s8(block0, vcombine_s8(vget_low_s8(a0), vget_low_s8(a1)));
vst1q_s8(block1, vcombine_s8(vget_high_s8(a0), vget_high_s8(a1)));
a0 = vld1q_s8(row2 + k_idx);
a1 = vld1q_s8(row3 + k_idx);
vst1q_s8(block0 + neon_smmla::TileSize,
vcombine_s8(vget_low_s8(a0), vget_low_s8(a1)));
vst1q_s8(block1 + neon_smmla::TileSize,
vcombine_s8(vget_high_s8(a0), vget_high_s8(a1)));
a0 = vld1q_s8(row4 + k_idx);
a1 = vld1q_s8(row5 + k_idx);
vst1q_s8(block0 + 2 * neon_smmla::TileSize,
vcombine_s8(vget_low_s8(a0), vget_low_s8(a1)));
vst1q_s8(block1 + 2 * neon_smmla::TileSize,
vcombine_s8(vget_high_s8(a0), vget_high_s8(a1)));
a0 = vld1q_s8(row6 + k_idx);
a1 = vld1q_s8(row7 + k_idx);
vst1q_s8(block0 + 3 * neon_smmla::TileSize,
vcombine_s8(vget_low_s8(a0), vget_low_s8(a1)));
vst1q_s8(block1 + 3 * neon_smmla::TileSize,
vcombine_s8(vget_high_s8(a0), vget_high_s8(a1)));
out += 8 * neon_smmla::TileSize;
}
for (; k_idx < k; k_idx += K) {
int8x8_t a0 = vld1_s8(row0 + k_idx);
int8x8_t a1 = vld1_s8(row1 + k_idx);
vst1q_s8(out, vcombine_s8(a0, a1));
a0 = vld1_s8(row2 + k_idx);
a1 = vld1_s8(row3 + k_idx);
vst1q_s8(out + neon_smmla::TileSize, vcombine_s8(a0, a1));
a0 = vld1_s8(row4 + k_idx);
a1 = vld1_s8(row5 + k_idx);
vst1q_s8(out + 2 * neon_smmla::TileSize, vcombine_s8(a0, a1));
a0 = vld1_s8(row6 + k_idx);
a1 = vld1_s8(row7 + k_idx);
vst1q_s8(out + 3 * neon_smmla::TileSize, vcombine_s8(a0, a1));
out += 4 * neon_smmla::TileSize;
}
continue;
}
const int32_t row_pairs = (panel_m <= 4) ? 2 : Mr / 2;
for (int32_t k_idx = 0; k_idx < k; k_idx += K) {
for (int32_t pair_idx = 0; pair_idx < row_pairs; ++pair_idx) {
const int32_t row_idx = pair_idx * 2;
const int8x8_t row0 =
(row_idx < panel_m) ? vld1_s8(panel_rows[row_idx] + k_idx) : zero;
const int8x8_t row1 = (row_idx + 1 < panel_m)
? vld1_s8(panel_rows[row_idx + 1] + k_idx)
: zero;
vst1q_s8(out, vcombine_s8(row0, row1));
out += neon_smmla::TileSize;
}
}
}
}
// physical layout [
// N / 8; Nr is 8
// K / 8; K for smmla is 8
// 4, ; 4 col-pairs for each 8 cols
// 2, ; col-pair is 2 cols
// 4 ; 4 elements per col
// ]
static void pack_weight(const int8_t* __restrict__ weight,
int8_t* __restrict__ packed_weight,
const int32_t output_size, const int32_t input_size) {
TORCH_CHECK(output_size % NSize == 0);
TORCH_CHECK(input_size % K == 0);
for (int32_t o_idx = 0; o_idx < output_size; o_idx += Nr) {
int8_t* __restrict__ dst = packed_weight + o_idx * input_size;
for (int32_t k_idx = 0; k_idx < input_size; k_idx += K) {
for (int32_t pair_idx = 0; pair_idx < Nr;
pair_idx += neon_smmla::Cols) {
const int8_t* __restrict__ row0 =
weight + (o_idx + pair_idx) * input_size + k_idx;
const int8_t* __restrict__ row1 = row0 + input_size;
vst1q_s8(dst, vcombine_s8(vld1_s8(row0), vld1_s8(row1)));
dst += neon_smmla::TileSize;
}
}
}
}
void gemm(const int8_t* __restrict__ a_packed,
const int8_t* __restrict__ b_packed, int32_t* __restrict__ c,
const int32_t m, const int32_t k, const int64_t b_n_group_stride,
const int64_t ldc) const {
TORCH_CHECK(m > 0 && m <= MaxMSize);
TORCH_CHECK(k % K == 0);
for (int32_t n_idx = 0; n_idx < NSize; n_idx += NrGemv) {
const int8_t* __restrict__ b_panel = b_packed + n_idx * k;
for (int32_t row_base = 0; row_base < m; row_base += Mr) {
const int32_t panel_m = std::min(Mr, m - row_base);
const int8_t* __restrict__ a_panel = a_packed + row_base * k;
int32_t* __restrict__ c_panel = c + row_base * ldc + n_idx;
if (panel_m <= 4) {
neon_smmla::gemm_micro_smmla_4x16_packed_a(
a_panel, b_panel, c_panel, panel_m, k, b_n_group_stride, ldc);
} else {
neon_smmla::gemm_micro_smmla_8x8_packed_a(a_panel, b_panel, c_panel,
panel_m, k, ldc);
neon_smmla::gemm_micro_smmla_8x8_packed_a(
a_panel, b_panel + b_n_group_stride, c_panel + Nr, panel_m, k,
ldc);
}
}
}
}
};
} // namespace cpu_micro_gemm
#endif
+14 -8
View File
@@ -1,3 +1,6 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
#ifndef CPU_MICRO_GEMM_NEON_HPP
#define CPU_MICRO_GEMM_NEON_HPP
@@ -16,9 +19,6 @@ namespace {
constexpr int32_t K = 4;
constexpr int32_t Cols = 2;
constexpr int32_t TileSize = K * Cols;
constexpr int32_t Mr = 8;
constexpr int32_t Nr = 8;
constexpr int32_t Nr_gemv = 16;
// a = [a0, a1, a2, a3], b = [b0, b1, b2, b3] -> [a0, a1, b0, b1]
FORCE_INLINE float32x4_t zip1_f32x4(const float32x4_t a, const float32x4_t b) {
@@ -132,7 +132,7 @@ FORCE_INLINE void gemm_micro_bfmmla_8x8_packed_a(
acc6767 = vbfmmlaq_f32(acc6767, a_tile67, b_tile67);
a_tile += 4 * TileSize;
b_tile += Nr * K;
b_tile += 4 * TileSize;
}
store_acc_rowpair(acc0101, acc0123, acc0145, acc0167, c_ptr, ldc,
@@ -205,8 +205,8 @@ FORCE_INLINE void gemm_micro_bfmmla_4x16_packed_a(
acc231415 = vbfmmlaq_f32(acc231415, a_tile23, b_tile1415);
a_tile += 2 * TileSize;
b_tile0 += Nr * K;
b_tile1 += Nr * K;
b_tile0 += 4 * TileSize;
b_tile1 += 4 * TileSize;
}
store_acc_rowpair(acc0101, acc0123, acc0145, acc0167, c_ptr, ldc, m_rows_01);
@@ -223,6 +223,9 @@ FORCE_INLINE void gemm_micro_bfmmla_4x16_packed_a(
template <typename scalar_t>
class MicroGemm<cpu_utils::ISA::NEON, scalar_t> {
public:
static constexpr int32_t Mr = 8;
static constexpr int32_t Nr = 8;
static constexpr int32_t NrGemv = 16;
static constexpr int32_t MaxMSize = 8;
static constexpr int32_t NSize = 32;
static constexpr int32_t WeightOCGroupSize = Nr;
@@ -246,6 +249,9 @@ class MicroGemm<cpu_utils::ISA::NEON, c10::BFloat16> {
public:
using scalar_t = c10::BFloat16;
static constexpr int32_t Mr = 8;
static constexpr int32_t Nr = 8;
static constexpr int32_t NrGemv = 16;
static constexpr int32_t MaxMSize = 8;
static constexpr int32_t NSize = 32;
static constexpr int32_t WeightOCGroupSize = Nr;
@@ -253,7 +259,7 @@ class MicroGemm<cpu_utils::ISA::NEON, c10::BFloat16> {
public:
// physical layout [
// M / 8; Mr is 8
// M / (8 or 4); Mr is 8 or 4
// K / 4; K for bfmmla is 4
// 4, ; 4 row-pairs for each 8 rows
// 2, ; row-pair is 2 rows
@@ -439,7 +445,7 @@ class MicroGemm<cpu_utils::ISA::NEON, c10::BFloat16> {
(void)lda; // A is packed, so lda is not needed
TORCH_CHECK_EQ(k % K, 0);
for (int32_t n_idx = 0; n_idx < NSize; n_idx += Nr_gemv) {
for (int32_t n_idx = 0; n_idx < NSize; n_idx += NrGemv) {
const bfloat16_t* __restrict__ b_panel =
reinterpret_cast<const bfloat16_t*>(b_ptr) + n_idx * k;
+173 -12
View File
@@ -451,6 +451,90 @@ void causal_conv1d_update_kernel_impl(
});
}
template <typename scalar_t>
void causal_conv1d_update_multi_kernel_impl(
scalar_t* __restrict__ out,
const scalar_t* __restrict__ input,
scalar_t* __restrict__ conv_states,
const scalar_t* __restrict__ weight,
const scalar_t* __restrict__ bias,
const int32_t* __restrict__ num_accepted_tokens,
const int32_t* __restrict__ conv_indices,
bool silu_activation,
int64_t batch,
int64_t dim,
int64_t seqlen,
int64_t width,
int64_t state_len,
int64_t conv_state_slot_stride) {
constexpr int64_t BLOCK_N = block_size_n() * 2;
const int64_t NB = div_up(dim, BLOCK_N);
AT_DISPATCH_BOOL2(bias != nullptr, has_bias, silu_activation, has_silu, [&] {
at::parallel_for(0, batch * NB, 0, [&](int64_t begin, int64_t end) {
int64_t bs{0}, nb{0};
data_index_init(begin, bs, batch, nb, NB);
for (int64_t i = begin; i < end; ++i) {
const int64_t nb_start = nb * BLOCK_N;
const int64_t nb_size = std::min(dim - nb_start, BLOCK_N);
const int32_t conv_state_index = conv_indices[bs];
const int32_t history_offset = num_accepted_tokens[bs] - 1;
switch (width << 4 | nb_size >> 4) {
case 0x42:
tinygemm_kernel<scalar_t, 4, 32, has_bias, has_silu>::apply(
input + bs * seqlen * dim + nb_start,
weight + nb_start * width,
out + bs * seqlen * dim + nb_start,
has_bias ? bias + nb_start : nullptr,
conv_states + conv_state_index * conv_state_slot_stride +
history_offset * dim + nb_start,
true,
seqlen,
dim,
true);
break;
case 0x44:
tinygemm_kernel<scalar_t, 4, 64, has_bias, has_silu>::apply(
input + bs * seqlen * dim + nb_start,
weight + nb_start * width,
out + bs * seqlen * dim + nb_start,
has_bias ? bias + nb_start : nullptr,
conv_states + conv_state_index * conv_state_slot_stride +
history_offset * dim + nb_start,
true,
seqlen,
dim,
true);
break;
default:
TORCH_CHECK(false, "Unexpected block size, ", width, " x ", nb_size);
}
data_index_step(bs, batch, nb, NB);
}
});
});
at::parallel_for(0, batch, 0, [&](int64_t begin, int64_t end) {
for (int64_t bs = begin; bs < end; ++bs) {
const int32_t conv_state_index = conv_indices[bs];
const int32_t num_accepted = num_accepted_tokens[bs];
scalar_t* state = conv_states + conv_state_index * conv_state_slot_stride;
std::memmove(
state,
state + num_accepted * dim,
(state_len - seqlen) * dim * sizeof(scalar_t));
std::memcpy(
state + (state_len - seqlen) * dim,
input + bs * seqlen * dim,
seqlen * dim * sizeof(scalar_t));
}
});
}
} // anonymous namespace
// from [dim, width] or [N, K]
@@ -545,7 +629,7 @@ at::Tensor get_block_indices(const std::optional<at::Tensor>& offsets, int64_t n
// query_start_loc: (batch + 1) int32
// cache_indices: (batch) int32
// has_initial_state: (batch) bool
// conv_states: (..., dim, width - 1) itype
// conv_states: (..., dim, state_len) itype, where state_len >= width - 1
// activation: either None or "silu" or "swish"
// pad_slot_id: int
//
@@ -586,11 +670,14 @@ at::Tensor causal_conv1d_fwd_cpu(
CHECK_EQ(conv_states_val.scalar_type(), scalar_type);
CHECK_GE(padded_batch, batch);
CHECK_EQ(conv_states_val.size(1), dim);
CHECK_EQ(conv_states_val.size(2), width - 1);
const int64_t state_len = conv_states_val.size(2);
CHECK_GE(state_len, width - 1);
// adjust `conv_states` to be contiguous on `dim`
// should happen only once
if (conv_states_val.stride(-2) != 1) {
TORCH_CHECK(state_len == width - 1,
"causal_conv1d_fwd_cpu: wide conv_states must be contiguous on dim.");
auto conv_states_copy = conv_states_val.clone();
conv_states_val.as_strided_({padded_batch, dim, width - 1}, {(width - 1) * dim, 1, dim});
conv_states_val.copy_(conv_states_copy);
@@ -651,14 +738,14 @@ at::Tensor causal_conv1d_fwd_cpu(
// API aligned with GPUs
//
// x: (batch, dim) or (batch, dim, seqlen)
// x: (batch, dim) or (batch, seqlen, dim)
// conv_state: (..., dim, state_len), where state_len >= width - 1
// weight: (dim, width)
// bias: (dim,)
// cache_seqlens: (batch,), dtype int32.
// num_accepted_tokens: (batch,), dtype int32.
// conv_state_indices: (batch,), dtype int32
// pad_slot_id: int
// out: (batch, dim) or (batch, dim, seqlen)
// out: (batch, dim) or (batch, seqlen, dim)
//
at::Tensor causal_conv1d_update_cpu(
const at::Tensor& x,
@@ -666,7 +753,7 @@ at::Tensor causal_conv1d_update_cpu(
const at::Tensor& weight,
const std::optional<at::Tensor>& bias,
bool silu_activation,
const std::optional<at::Tensor>& cache_seqlens,
const std::optional<at::Tensor>& num_accepted_tokens,
const std::optional<at::Tensor>& conv_state_indices,
int64_t pad_slot_id,
bool is_vnni) {
@@ -674,13 +761,13 @@ at::Tensor causal_conv1d_update_cpu(
CHECK_CONTIGUOUS(weight);
auto packed_w = is_vnni ? weight : causal_conv1d_weight_pack(weight);
// TODO: add multi-token prediction support
TORCH_CHECK(x.dim() == 2, "causal_conv1d_update_cpu: expect x to be 2D tensor.");
TORCH_CHECK(!cache_seqlens.has_value(), "causal_conv1d_update_cpu: don't support cache_seqlens.");
TORCH_CHECK(
x.dim() == 2 || x.dim() == 3,
"causal_conv1d_update_cpu: expect x to be 2D or 3D tensor.");
int64_t batch = x.size(0);
int64_t dim = x.size(1);
int64_t seqlen = 1;
int64_t dim = x.dim() == 2 ? x.size(1) : x.size(2);
int64_t seqlen = x.dim() == 2 ? 1 : x.size(1);
int64_t width = weight.size(-1);
const auto scalar_type = x.scalar_type();
@@ -690,10 +777,84 @@ at::Tensor causal_conv1d_update_cpu(
CHECK_EQ(conv_states.scalar_type(), scalar_type);
CHECK_EQ(conv_states.size(1), dim);
CHECK_EQ(conv_states.size(2), width - 1);
const int64_t state_len = conv_states.size(2);
CHECK_GE(state_len, width - 1);
if (x.dim() == 3) {
TORCH_CHECK(
num_accepted_tokens.has_value(),
"causal_conv1d_update_cpu: num_accepted_tokens is required for 3D x.");
TORCH_CHECK(
conv_state_indices.has_value(),
"causal_conv1d_update_cpu: conv_state_indices is required for 3D x.");
CHECK_OPTIONAL_SHAPE_DTYPE(num_accepted_tokens, batch, at::kInt);
TORCH_CHECK(
width == 4,
"causal_conv1d_update_cpu: support only width of 4 for 3D x.");
TORCH_CHECK(
seqlen > 0,
"causal_conv1d_update_cpu: expect non-empty sequence for 3D x.");
TORCH_CHECK(
state_len >= seqlen,
"causal_conv1d_update_cpu: state_len must be >= seqlen for 3D x.");
TORCH_CHECK(
conv_states.stride(-2) == 1 && conv_states.stride(-1) == dim,
"causal_conv1d_update_cpu: 3D x requires SD conv_states layout.");
const int32_t* accepted_counts =
num_accepted_tokens.value().data_ptr<int32_t>();
const int32_t* indices = conv_state_indices.value().data_ptr<int32_t>();
const int64_t num_slots = conv_states.size(0);
for (int64_t bs = 0; bs < batch; ++bs) {
const int32_t num_accepted = accepted_counts[bs];
const int32_t conv_state_index = indices[bs];
TORCH_CHECK(
conv_state_index != pad_slot_id,
"causal_conv1d_update_cpu: 3D x does not support pad slots.");
TORCH_CHECK(
conv_state_index >= 0 && conv_state_index < num_slots,
"causal_conv1d_update_cpu: conv_state_indices out of range.");
TORCH_CHECK(
num_accepted >= 1 && num_accepted <= seqlen,
"causal_conv1d_update_cpu: num_accepted_tokens must be in [1, "
"seqlen].");
TORCH_CHECK(
num_accepted - 1 + width - 1 <= state_len,
"causal_conv1d_update_cpu: history window exceeds conv_states.");
}
int64_t conv_state_slot_stride = conv_states.stride(0);
at::Tensor out = at::empty_like(x);
AT_DISPATCH_REDUCED_FLOATING_TYPES(
scalar_type, "causal_conv1d_update_multi_kernel_impl", [&] {
causal_conv1d_update_multi_kernel_impl<scalar_t>(
out.data_ptr<scalar_t>(),
x.data_ptr<scalar_t>(),
conv_states.data_ptr<scalar_t>(),
packed_w.data_ptr<scalar_t>(),
conditional_data_ptr<scalar_t>(bias),
accepted_counts,
indices,
silu_activation,
batch,
dim,
seqlen,
width,
state_len,
conv_state_slot_stride);
});
return out;
}
TORCH_CHECK(
!num_accepted_tokens.has_value(),
"causal_conv1d_update_cpu: num_accepted_tokens is only supported for 3D "
"x.");
// adjust `conv_states` to be contiguous on `dim`
if (conv_states.stride(-2) != 1) {
TORCH_CHECK(state_len == width - 1,
"causal_conv1d_update_cpu: wide conv_states must be contiguous on dim.");
int64_t num_cache_lines = conv_states.size(0);
auto conv_states_copy = conv_states.clone();
conv_states.as_strided_({num_cache_lines, dim, width - 1}, {(width - 1) * dim, 1, dim});
+34 -5
View File
@@ -147,7 +147,7 @@ at::Tensor causal_conv1d_fwd_cpu(
at::Tensor causal_conv1d_update_cpu(
const at::Tensor& x, const at::Tensor& conv_states,
const at::Tensor& weight, const std::optional<at::Tensor>& bias,
bool silu_activation, const std::optional<at::Tensor>& cache_seqlens,
bool silu_activation, const std::optional<at::Tensor>& num_accepted_tokens,
const std::optional<at::Tensor>& conv_state_indices, int64_t pad_slot_id,
bool is_vnni);
@@ -207,6 +207,20 @@ void cpu_fused_moe(torch::Tensor& output, const torch::Tensor& input,
const torch::Tensor& topk_id, const bool skip_weighted,
const std::string& act, const std::string& isa);
void prepack_moe_weight_int8(const torch::Tensor& weight,
torch::Tensor& packed_weight,
const std::string& isa);
void cpu_fused_moe_int8(torch::Tensor& output, const torch::Tensor& input,
const torch::Tensor& w13, const torch::Tensor& w2,
const torch::Tensor& w13_scale,
const torch::Tensor& w2_scale,
const std::optional<torch::Tensor>& w13_bias,
const std::optional<torch::Tensor>& w2_bias,
const torch::Tensor& topk_weights,
const torch::Tensor& topk_id, const bool skip_weighted,
const std::string& act, const std::string& isa);
void compute_slot_mapping_kernel_impl(const torch::Tensor query_start_loc,
const torch::Tensor positions,
const torch::Tensor block_table,
@@ -502,7 +516,8 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
ops.def(
"causal_conv1d_update_cpu(Tensor x, Tensor(a!) conv_states, Tensor "
"weight, Tensor? bias, bool silu_activation,"
"Tensor? cache_seqlens, Tensor? conv_state_indices, int pad_slot_id, "
"Tensor? num_accepted_tokens, Tensor? conv_state_indices, int "
"pad_slot_id, "
"bool is_vnni) -> Tensor");
ops.impl("causal_conv1d_update_cpu", torch::kCPU, &causal_conv1d_update_cpu);
#endif
@@ -596,8 +611,7 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
#endif
// fused moe
#if defined(__AVX512F__) || \
(defined(__aarch64__) && !defined(__APPLE__) && defined(ARM_BF16_SUPPORT))
#if defined(__AVX512F__) || (defined(ARM_BF16_SUPPORT) && !defined(__APPLE__))
ops.def(
"prepack_moe_weight(Tensor weight, Tensor(a1!) packed_weight, str isa) "
"-> ()");
@@ -608,7 +622,22 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
"bool skip_weighted, "
"str act, str isa) -> ()");
ops.impl("cpu_fused_moe", torch::kCPU, &cpu_fused_moe);
#endif
#endif // #if defined(__AVX512F__) || (defined(ARM_BF16_SUPPORT) &&
// !defined(__APPLE__))
#if defined(ARM_I8MM_SUPPORT) && defined(ARM_BF16_SUPPORT) && \
!defined(__APPLE__)
ops.def(
"prepack_moe_weight_int8(Tensor weight, Tensor(a1!) packed_weight, "
"str isa) -> ()");
ops.impl("prepack_moe_weight_int8", torch::kCPU, &prepack_moe_weight_int8);
ops.def(
"cpu_fused_moe_int8(Tensor(a0!) output, Tensor input, Tensor w13, "
"Tensor w2, Tensor w13_scale, Tensor w2_scale, Tensor? w13_bias, "
"Tensor? w2_bias, Tensor topk_weights, Tensor topk_id, bool "
"skip_weighted, str act, str isa) -> ()");
ops.impl("cpu_fused_moe_int8", torch::kCPU, &cpu_fused_moe_int8);
#endif // #if defined(ARM_I8MM_SUPPORT) && defined(ARM_BF16_SUPPORT) &&
// !defined(__APPLE__)
ops.def(
"mla_decode_kvcache("
" Tensor! out, Tensor query, Tensor kv_cache,"
-326
View File
@@ -1,326 +0,0 @@
#pragma once
#include "custom_collective_common.cuh"
namespace vllm {
constexpr int kMnnvlLamportAgThreads = 128;
constexpr int kMnnvlLamportRsThreads = 256;
constexpr int kMnnvlLamportConcurrentPollMaxPacks = 8192;
using CopyPack = array_t<uint64_t, 2>;
template <int ngpus>
__global__ void __launch_bounds__(512, 1)
cross_device_all_gather(RankData* _dp, RankSignals sg, Signal* self_sg,
CopyPack* __restrict__ result, int rank,
int size_per_rank) {
auto dp = *_dp;
int tid = blockIdx.x * blockDim.x + threadIdx.x;
int stride = gridDim.x * blockDim.x;
barrier_at_start<ngpus>(sg, self_sg, rank);
#pragma unroll
for (int src_rank = 0; src_rank < ngpus; ++src_rank) {
auto src = reinterpret_cast<const CopyPack*>(dp.ptrs[src_rank]);
auto dst = result + src_rank * size_per_rank;
for (int idx = tid; idx < size_per_rank; idx += stride) {
dst[idx] = src[idx];
}
}
barrier_at_end<ngpus, true>(sg, self_sg, rank);
}
template <typename T, int ngpus>
__global__ void __launch_bounds__(512, 1)
cross_device_reduce_scatter(RankData* _dp, RankSignals sg, Signal* self_sg,
T* __restrict__ result, int rank,
int size_per_rank) {
using P = typename packed_t<T>::P;
using A = typename packed_t<T>::A;
auto dp = *_dp;
auto offset = rank * size_per_rank;
barrier_at_start<ngpus>(sg, self_sg, rank);
for (int idx = blockIdx.x * blockDim.x + threadIdx.x; idx < size_per_rank;
idx += gridDim.x * blockDim.x) {
reinterpret_cast<P*>(result)[idx] =
packed_reduce<P, ngpus, A>((const P**)&dp.ptrs[0], offset + idx);
}
barrier_at_end<ngpus, true>(sg, self_sg, rank);
}
template <typename P>
union LamportPack {
P packed;
uint32_t words[sizeof(P) / sizeof(uint32_t)];
};
template <typename P>
DINLINE LamportPack<P> load_lamport_pack(const P* ptr) {
static_assert(sizeof(P) == 16);
LamportPack<P> value;
#if !defined(USE_ROCM)
asm volatile("ld.volatile.global.v4.u32 {%0, %1, %2, %3}, [%4];"
: "=r"(value.words[0]), "=r"(value.words[1]),
"=r"(value.words[2]), "=r"(value.words[3])
: "l"(ptr)
: "memory");
#else
const volatile uint32_t* src = reinterpret_cast<const volatile uint32_t*>(ptr);
#pragma unroll
for (int i = 0; i < sizeof(P) / sizeof(uint32_t); ++i) {
value.words[i] = src[i];
}
#endif
return value;
}
template <typename P>
DINLINE bool is_lamport_dirty(const LamportPack<P>& value) {
#pragma unroll
for (int i = 0; i < sizeof(P) / sizeof(uint32_t); ++i) {
if (value.words[i] == 0x80000000U) return true;
}
return false;
}
template <typename P>
DINLINE P lamport_sentinel() {
LamportPack<P> value;
#pragma unroll
for (int i = 0; i < sizeof(P) / sizeof(uint32_t); ++i) {
value.words[i] = 0x80000000U;
}
return value.packed;
}
template <typename P>
DINLINE P sanitize_lamport_payload(P packed) {
LamportPack<P> value{.packed = packed};
#pragma unroll
for (int i = 0; i < sizeof(P) / sizeof(uint32_t); ++i) {
if (value.words[i] == 0x80000000U) value.words[i] = 0;
}
return value.packed;
}
template <typename P>
DINLINE P wait_lamport_payload(const P* ptr) {
auto value = load_lamport_pack(ptr);
while (is_lamport_dirty(value)) value = load_lamport_pack(ptr);
return value.packed;
}
template <typename P, int ngpus>
DINLINE void wait_lamport_payloads(const P* base, int rank, int rank_stride,
P local_value, P (&values)[ngpus]) {
bool ready[ngpus];
#pragma unroll
for (int src_rank = 0; src_rank < ngpus; ++src_rank) {
ready[src_rank] = src_rank == rank;
if (src_rank == rank) values[src_rank] = local_value;
}
int remaining = ngpus - 1;
while (remaining != 0) {
#pragma unroll
for (int src_rank = 0; src_rank < ngpus; ++src_rank) {
if (!ready[src_rank]) {
auto value = load_lamport_pack(base + src_rank * rank_stride);
if (!is_lamport_dirty(value)) {
values[src_rank] = value.packed;
ready[src_rank] = true;
--remaining;
}
}
}
}
}
template <typename P, typename A, int ngpus>
DINLINE P reduce_lamport_payloads(const P* current_local, const P* packed_input,
int rank, int size_per_rank, int idx) {
P source_zero =
rank == 0 ? packed_input[idx] : wait_lamport_payload(current_local + idx);
A tmp = upcast(source_zero);
#pragma unroll
for (int src_rank = 1; src_rank < ngpus; ++src_rank) {
P value = src_rank == rank
? packed_input[rank * size_per_rank + idx]
: wait_lamport_payload(current_local +
src_rank * size_per_rank + idx);
packed_assign_add(tmp, upcast(value));
}
return sanitize_lamport_payload(downcast<P>(tmp));
}
DINLINE void lamport_cta_arrive(uint32_t* counter) {
#if !defined(USE_ROCM)
if (threadIdx.x < 32) {
asm volatile("barrier.cta.sync 1, %0;" : : "r"(blockDim.x) : "memory");
if (threadIdx.x == 0) {
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 1000
asm volatile("red.async.release.global.gpu.add.u32 [%0], 1;"
:
: "l"(counter)
: "memory");
#elif defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 700
asm volatile("red.release.global.gpu.add.u32 [%0], 1;"
:
: "l"(counter)
: "memory");
#else
atomicAdd(counter, 1);
#endif
}
} else {
asm volatile("barrier.cta.arrive 1, %0;" : : "r"(blockDim.x) : "memory");
}
#else
__syncthreads();
if (threadIdx.x == 0) atomicAdd(counter, 1);
#endif
}
template <typename T, int ngpus>
__global__ void __launch_bounds__(kMnnvlLamportAgThreads, 1)
mnnvl_lamport_all_gather(RankData* _dp, const T* __restrict__ input,
T* __restrict__ result,
T* __restrict__ multicast_buffer,
uint32_t* __restrict__ epochs, int rank,
int size_per_rank, int stage_size) {
using P = typename packed_t<T>::P;
#if !defined(USE_ROCM) && CUDA_VERSION >= 12000 && defined(__CUDA_ARCH__) && \
(__CUDA_ARCH__ >= 900)
cudaGridDependencySynchronize();
#endif
auto dp = *_dp;
int tid = blockIdx.x * blockDim.x + threadIdx.x;
int stride = gridDim.x * blockDim.x;
uint32_t epoch = epochs[0];
int current_stage = epoch % 3;
int dirty_stage = (epoch + 1) % 3;
int dirty_size = epochs[2 + dirty_stage];
auto local_buffer = reinterpret_cast<P*>(const_cast<void*>(dp.ptrs[rank]));
auto current_local = local_buffer + current_stage * stage_size;
auto dirty_local = local_buffer + dirty_stage * stage_size;
auto current_multicast =
reinterpret_cast<P*>(multicast_buffer) + current_stage * stage_size;
auto packed_input = reinterpret_cast<const P*>(input);
auto packed_result = reinterpret_cast<P*>(result);
int total_size = size_per_rank * ngpus;
P local_value;
if (tid < size_per_rank) {
local_value = packed_input[tid];
current_multicast[rank * size_per_rank + tid] =
sanitize_lamport_payload(local_value);
}
#if !defined(USE_ROCM) && CUDA_VERSION >= 12000 && defined(__CUDA_ARCH__) && \
(__CUDA_ARCH__ >= 900)
cudaTriggerProgrammaticLaunchCompletion();
#endif
lamport_cta_arrive(&epochs[1]);
for (int idx = tid; idx < dirty_size; idx += stride) {
dirty_local[idx] = lamport_sentinel<P>();
}
if (tid < size_per_rank) {
#pragma unroll
for (int src_rank = 0; src_rank < ngpus; ++src_rank) {
int output_idx = src_rank * size_per_rank + tid;
P value = src_rank == rank
? local_value
: wait_lamport_payload(current_local + output_idx);
packed_result[output_idx] = value;
}
}
if (tid == 0) {
while (*reinterpret_cast<volatile uint32_t*>(&epochs[1]) < gridDim.x);
epochs[2 + current_stage] = total_size;
epochs[0] = epoch + 1;
epochs[1] = 0;
}
}
template <typename T, int ngpus>
__global__ void __launch_bounds__(kMnnvlLamportRsThreads, 1)
mnnvl_lamport_reduce_scatter_kernel(RankData* _dp,
const T* __restrict__ input,
T* __restrict__ result,
uint32_t* __restrict__ epochs, int rank,
int size_per_rank, int stage_size) {
using P = typename packed_t<T>::P;
using A = typename packed_t<T>::A;
#if !defined(USE_ROCM) && CUDA_VERSION >= 12000 && defined(__CUDA_ARCH__) && \
(__CUDA_ARCH__ >= 900)
cudaGridDependencySynchronize();
#endif
auto dp = *_dp;
int dst_rank = blockIdx.x % ngpus;
int tile = blockIdx.x / ngpus;
int idx = tile * blockDim.x + threadIdx.x;
int tid = blockIdx.x * blockDim.x + threadIdx.x;
int stride = gridDim.x * blockDim.x;
uint32_t epoch = epochs[0];
int current_stage = epoch % 3;
int dirty_stage = (epoch + 1) % 3;
int dirty_size = epochs[2 + dirty_stage];
auto local_buffer = reinterpret_cast<P*>(const_cast<void*>(dp.ptrs[rank]));
auto current_local = local_buffer + current_stage * stage_size;
auto dirty_local = local_buffer + dirty_stage * stage_size;
auto packed_input = reinterpret_cast<const P*>(input);
if (idx < size_per_rank && dst_rank != rank) {
auto dst = reinterpret_cast<P*>(const_cast<void*>(dp.ptrs[dst_rank])) +
current_stage * stage_size + rank * size_per_rank;
auto src = packed_input + dst_rank * size_per_rank;
dst[idx] = sanitize_lamport_payload(src[idx]);
}
#if !defined(USE_ROCM) && CUDA_VERSION >= 12000 && defined(__CUDA_ARCH__) && \
(__CUDA_ARCH__ >= 900)
cudaTriggerProgrammaticLaunchCompletion();
#endif
lamport_cta_arrive(&epochs[1]);
for (int idx = tid; idx < dirty_size; idx += stride) {
dirty_local[idx] = lamport_sentinel<P>();
}
if (idx < size_per_rank && dst_rank == rank) {
if constexpr (ngpus == 4) {
if (size_per_rank > kMnnvlLamportConcurrentPollMaxPacks) {
reinterpret_cast<P*>(result)[idx] =
reduce_lamport_payloads<P, A, ngpus>(current_local, packed_input,
rank, size_per_rank, idx);
} else {
P values[ngpus];
wait_lamport_payloads<P, ngpus>(
current_local + idx, rank, size_per_rank,
packed_input[rank * size_per_rank + idx], values);
A tmp = upcast(values[0]);
#pragma unroll
for (int src_rank = 1; src_rank < ngpus; ++src_rank) {
packed_assign_add(tmp, upcast(values[src_rank]));
}
reinterpret_cast<P*>(result)[idx] =
sanitize_lamport_payload(downcast<P>(tmp));
}
} else {
reinterpret_cast<P*>(result)[idx] = reduce_lamport_payloads<P, A, ngpus>(
current_local, packed_input, rank, size_per_rank, idx);
}
}
if (tid == 0) {
while (*reinterpret_cast<volatile uint32_t*>(&epochs[1]) < gridDim.x);
epochs[2 + current_stage] = size_per_rank * ngpus;
epochs[0] = epoch + 1;
epochs[1] = 0;
}
}
} // namespace vllm
+296 -20
View File
@@ -1,8 +1,299 @@
#pragma once
#include "custom_collective_common.cuh"
#include <cuda.h>
#include <cuda_bf16.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#if defined(USE_ROCM)
typedef __hip_bfloat16 nv_bfloat16;
#endif
#include <iostream>
#include <array>
#include <limits>
#include <map>
#include <unordered_map>
#include <vector>
#include <cstdlib>
#include <cstring>
namespace vllm {
#define CUDACHECK(cmd) \
do { \
cudaError_t e = cmd; \
if (e != cudaSuccess) { \
printf("Failed: Cuda error %s:%d '%s'\n", __FILE__, __LINE__, \
cudaGetErrorString(e)); \
exit(EXIT_FAILURE); \
} \
} while (0)
// Maximal number of blocks in allreduce kernel.
constexpr int kMaxBlocks = 36;
// Default number of blocks in allreduce kernel.
#ifndef USE_ROCM
const int defaultBlockLimit = 36;
CUpointer_attribute rangeStartAddrAttr = CU_POINTER_ATTRIBUTE_RANGE_START_ADDR;
#else
const int defaultBlockLimit = 16;
hipPointer_attribute rangeStartAddrAttr =
HIP_POINTER_ATTRIBUTE_RANGE_START_ADDR;
#endif
// Counter may overflow, but it's fine since unsigned int overflow is
// well-defined behavior.
using FlagType = uint32_t;
// Two sets of peer counters are needed for two syncs: starting and ending an
// operation. The reason is that it's possible for peer GPU block to arrive at
// the second sync point while the current GPU block haven't passed the first
// sync point. Thus, peer GPU may write counter+1 while current GPU is busy
// waiting for counter. We use alternating counter array to avoid this
// possibility.
struct Signal {
alignas(128) FlagType start[kMaxBlocks][8];
alignas(128) FlagType end[kMaxBlocks][8];
alignas(128) FlagType _flag[kMaxBlocks]; // incremental flags for each rank
};
struct __align__(16) RankData {
const void* ptrs[8];
};
struct __align__(16) RankSignals {
Signal* signals[8];
};
// like std::array, but aligned
template <typename T, int sz>
struct __align__(alignof(T) * sz) array_t {
T data[sz];
using type = T;
static constexpr int size = sz;
};
// use packed type to maximize memory efficiency
// goal: generate ld.128 and st.128 instructions
template <typename T>
struct packed_t {
// the (P)acked type for load/store
using P = array_t<T, 16 / sizeof(T)>;
// the (A)ccumulator type for reduction
using A = array_t<float, 16 / sizeof(T)>;
};
#define DINLINE __device__ __forceinline__
// scalar cast functions
DINLINE float upcast_s(half val) { return __half2float(val); }
template <typename T>
DINLINE T downcast_s(float val);
template <>
DINLINE half downcast_s(float val) {
return __float2half(val);
}
// scalar add functions
// for some reason when compiling with Pytorch, the + operator for half and
// bfloat is disabled so we call the intrinsics directly
DINLINE half& assign_add(half& a, half b) {
a = __hadd(a, b);
return a;
}
DINLINE float& assign_add(float& a, float b) { return a += b; }
#if (__CUDA_ARCH__ >= 800 || !defined(__CUDA_ARCH__))
DINLINE float upcast_s(nv_bfloat16 val) { return __bfloat162float(val); }
template <>
DINLINE nv_bfloat16 downcast_s(float val) {
return __float2bfloat16(val);
}
DINLINE nv_bfloat16& assign_add(nv_bfloat16& a, nv_bfloat16 b) {
a = __hadd(a, b);
return a;
}
#endif
template <typename T, int N>
DINLINE array_t<T, N>& packed_assign_add(array_t<T, N>& a, array_t<T, N> b) {
#pragma unroll
for (int i = 0; i < N; i++) {
assign_add(a.data[i], b.data[i]);
}
return a;
}
template <typename T, int N>
DINLINE array_t<float, N> upcast(array_t<T, N> val) {
if constexpr (std::is_same<T, float>::value) {
return val;
} else {
array_t<float, N> out;
#pragma unroll
for (int i = 0; i < N; i++) {
out.data[i] = upcast_s(val.data[i]);
}
return out;
}
}
template <typename O>
DINLINE O downcast(array_t<float, O::size> val) {
if constexpr (std::is_same<typename O::type, float>::value) {
return val;
} else {
O out;
#pragma unroll
for (int i = 0; i < O::size; i++) {
out.data[i] = downcast_s<typename O::type>(val.data[i]);
}
return out;
}
}
#if !defined(USE_ROCM)
static DINLINE void st_flag_release(FlagType* flag_addr, FlagType flag) {
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 700
asm volatile("st.release.sys.global.u32 [%1], %0;" ::"r"(flag),
"l"(flag_addr));
#else
asm volatile("membar.sys; st.volatile.global.u32 [%1], %0;" ::"r"(flag),
"l"(flag_addr));
#endif
}
static DINLINE FlagType ld_flag_acquire(FlagType* flag_addr) {
FlagType flag;
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 700
asm volatile("ld.acquire.sys.global.u32 %0, [%1];"
: "=r"(flag)
: "l"(flag_addr));
#else
asm volatile("ld.volatile.global.u32 %0, [%1]; membar.gl;"
: "=r"(flag)
: "l"(flag_addr));
#endif
return flag;
}
static DINLINE void st_flag_volatile(FlagType* flag_addr, FlagType flag) {
asm volatile("st.volatile.global.u32 [%1], %0;" ::"r"(flag), "l"(flag_addr));
}
static DINLINE FlagType ld_flag_volatile(FlagType* flag_addr) {
FlagType flag;
asm volatile("ld.volatile.global.u32 %0, [%1];"
: "=r"(flag)
: "l"(flag_addr));
return flag;
}
// This function is meant to be used as the first synchronization in the all
// reduce kernel. Thus, it doesn't need to make any visibility guarantees for
// prior memory accesses. Note: volatile writes will not be reordered against
// other volatile writes.
template <int ngpus>
DINLINE void barrier_at_start(const RankSignals& sg, Signal* self_sg,
int rank) {
uint32_t flag = self_sg->_flag[blockIdx.x] + 1;
if (threadIdx.x < ngpus) {
auto peer_counter_ptr = &sg.signals[threadIdx.x]->start[blockIdx.x][rank];
auto self_counter_ptr = &self_sg->start[blockIdx.x][threadIdx.x];
// Write the expected counter value to peer and wait for correct value
// from peer.
st_flag_volatile(peer_counter_ptr, flag);
while (ld_flag_volatile(self_counter_ptr) != flag);
}
__syncthreads();
// use one thread to update flag
if (threadIdx.x == 0) self_sg->_flag[blockIdx.x] = flag;
}
// This function is meant to be used as the second or the final
// synchronization barrier in the all reduce kernel. If it's the final
// synchronization barrier, we don't need to make any visibility guarantees
// for prior memory accesses.
template <int ngpus, bool final_sync = false>
DINLINE void barrier_at_end(const RankSignals& sg, Signal* self_sg, int rank) {
__syncthreads();
uint32_t flag = self_sg->_flag[blockIdx.x] + 1;
if (threadIdx.x < ngpus) {
auto peer_counter_ptr = &sg.signals[threadIdx.x]->end[blockIdx.x][rank];
auto self_counter_ptr = &self_sg->end[blockIdx.x][threadIdx.x];
// Write the expected counter value to peer and wait for correct value from
// peer.
if constexpr (!final_sync) {
st_flag_release(peer_counter_ptr, flag);
while (ld_flag_acquire(self_counter_ptr) != flag);
} else {
st_flag_volatile(peer_counter_ptr, flag);
while (ld_flag_volatile(self_counter_ptr) != flag);
}
}
if constexpr (!final_sync) __syncthreads();
// use one thread to update flag
if (threadIdx.x == 0) self_sg->_flag[blockIdx.x] = flag;
}
#else
template <int ngpus>
DINLINE void barrier_at_start(const RankSignals& sg, Signal* self_sg,
int rank) {
uint32_t flag = self_sg->_flag[blockIdx.x] + 1;
if (threadIdx.x < ngpus) {
// simultaneously write to the corresponding flag of all ranks.
// Latency = 1 p2p write
__scoped_atomic_store_n(&sg.signals[threadIdx.x]->start[blockIdx.x][rank],
flag, __ATOMIC_RELAXED, __MEMORY_SCOPE_SYSTEM);
// wait until we got true from all ranks
while (__scoped_atomic_load_n(&self_sg->start[blockIdx.x][threadIdx.x],
__ATOMIC_RELAXED,
__MEMORY_SCOPE_DEVICE) < flag);
}
__syncthreads();
// use one thread to update flag
if (threadIdx.x == 0) self_sg->_flag[blockIdx.x] = flag;
}
template <int ngpus, bool final_sync = false>
DINLINE void barrier_at_end(const RankSignals& sg, Signal* self_sg, int rank) {
__syncthreads();
uint32_t flag = self_sg->_flag[blockIdx.x] + 1;
if (threadIdx.x < ngpus) {
// simultaneously write to the corresponding flag of all ranks.
// Latency = 1 p2p write
__scoped_atomic_store_n(&sg.signals[threadIdx.x]->end[blockIdx.x][rank],
flag,
final_sync ? __ATOMIC_RELAXED : __ATOMIC_RELEASE,
__MEMORY_SCOPE_SYSTEM);
// wait until we got true from all ranks
while (
__scoped_atomic_load_n(&self_sg->end[blockIdx.x][threadIdx.x],
final_sync ? __ATOMIC_RELAXED : __ATOMIC_ACQUIRE,
__MEMORY_SCOPE_DEVICE) < flag);
}
if constexpr (!final_sync) __syncthreads();
// use one thread to update flag
if (threadIdx.x == 0) self_sg->_flag[blockIdx.x] = flag;
}
#endif
template <typename P, int ngpus, typename A>
DINLINE P packed_reduce(const P* ptrs[], int idx) {
A tmp = upcast(ptrs[0][idx]);
#pragma unroll
for (int i = 1; i < ngpus; i++) {
packed_assign_add(tmp, upcast(ptrs[i][idx]));
}
return downcast<P>(tmp);
}
template <typename T, int ngpus>
__global__ void __launch_bounds__(512, 1)
@@ -325,21 +616,6 @@ class CustomAllreduce {
#undef KL
}
void allgather(cudaStream_t stream, void* input, void* output, int size_bytes,
int threads = 512, int block_limit = defaultBlockLimit);
template <typename T>
void mnnvl_lamport_allgather(cudaStream_t stream, T* input, T* output,
void* local_buffer, void* multicast_buffer,
uint32_t* epochs, int size_bytes,
int stage_size_bytes);
template <typename T>
void reduce_scatter(cudaStream_t stream, T* input, T* output, int size,
int threads = 512, int block_limit = defaultBlockLimit);
template <typename T>
void mnnvl_lamport_reduce_scatter(cudaStream_t stream, T* input, T* output,
void* local_buffer, uint32_t* epochs,
int size, int stage_size_bytes);
~CustomAllreduce() {
for (auto [_, ptr] : ipc_handles_) {
CUDACHECK(cudaIpcCloseMemHandle(ptr));
@@ -349,8 +625,8 @@ class CustomAllreduce {
/**
* To inspect PTX/SASS, copy paste this header file to compiler explorer and
* add a template instantiation:
add a template instantiation:
* template void vllm::CustomAllreduce::allreduce<half>(cudaStream_t, half *,
* half *, int, int, int);
*/
} // namespace vllm
half *, int, int, int);
*/
} // namespace vllm
-332
View File
@@ -1,332 +0,0 @@
#pragma once
#include <cuda.h>
#include <cuda_bf16.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#if defined(USE_ROCM)
typedef __hip_bfloat16 nv_bfloat16;
#endif
#include <iostream>
#include <array>
#include <limits>
#include <map>
#include <unordered_map>
#include <vector>
#include <cstdlib>
#include <cstring>
namespace vllm {
constexpr int kMaxCustomCollectiveRanks = 16;
#define CUDACHECK(cmd) \
do { \
cudaError_t e = cmd; \
if (e != cudaSuccess) { \
printf("Failed: Cuda error %s:%d '%s'\n", __FILE__, __LINE__, \
cudaGetErrorString(e)); \
exit(EXIT_FAILURE); \
} \
} while (0)
// Maximal number of blocks in allreduce kernel.
constexpr int kMaxBlocks = 36;
// Default number of blocks in allreduce kernel.
#ifndef USE_ROCM
inline constexpr int defaultBlockLimit = 36;
inline CUpointer_attribute rangeStartAddrAttr =
CU_POINTER_ATTRIBUTE_RANGE_START_ADDR;
#else
inline constexpr int defaultBlockLimit = 16;
inline hipPointer_attribute rangeStartAddrAttr =
HIP_POINTER_ATTRIBUTE_RANGE_START_ADDR;
#endif
// Counter may overflow, but it's fine since unsigned int overflow is
// well-defined behavior.
using FlagType = uint32_t;
// Two sets of peer counters are needed for two syncs: starting and ending an
// operation. The reason is that it's possible for peer GPU block to arrive at
// the second sync point while the current GPU block haven't passed the first
// sync point. Thus, peer GPU may write counter+1 while current GPU is busy
// waiting for counter. We use alternating counter array to avoid this
// possibility.
struct Signal {
alignas(128) FlagType start[kMaxBlocks][kMaxCustomCollectiveRanks];
alignas(128) FlagType end[kMaxBlocks][kMaxCustomCollectiveRanks];
alignas(128) FlagType _flag[kMaxBlocks]; // incremental flags for each rank
};
struct __align__(16) RankData {
const void* ptrs[kMaxCustomCollectiveRanks];
};
struct __align__(16) RankSignals {
Signal* signals[kMaxCustomCollectiveRanks];
};
// like std::array, but aligned
template <typename T, int sz>
struct __align__(alignof(T) * sz) array_t {
T data[sz];
using type = T;
static constexpr int size = sz;
};
// use packed type to maximize memory efficiency
// goal: generate ld.128 and st.128 instructions
template <typename T>
struct packed_t {
// the (P)acked type for load/store
using P = array_t<T, 16 / sizeof(T)>;
// the (A)ccumulator type for reduction
using A = array_t<float, 16 / sizeof(T)>;
};
#define DINLINE __device__ __forceinline__
// scalar cast functions
DINLINE float upcast_s(half val) { return __half2float(val); }
template <typename T>
DINLINE T downcast_s(float val);
template <>
DINLINE half downcast_s(float val) {
return __float2half(val);
}
// scalar add functions
// for some reason when compiling with Pytorch, the + operator for half and
// bfloat is disabled so we call the intrinsics directly
DINLINE half& assign_add(half& a, half b) {
a = __hadd(a, b);
return a;
}
DINLINE float& assign_add(float& a, float b) { return a += b; }
#if (__CUDA_ARCH__ >= 800 || !defined(__CUDA_ARCH__))
DINLINE float upcast_s(nv_bfloat16 val) { return __bfloat162float(val); }
template <>
DINLINE nv_bfloat16 downcast_s(float val) {
return __float2bfloat16(val);
}
DINLINE nv_bfloat16& assign_add(nv_bfloat16& a, nv_bfloat16 b) {
a = __hadd(a, b);
return a;
}
#endif
template <typename T, int N>
DINLINE array_t<T, N>& packed_assign_add(array_t<T, N>& a, array_t<T, N> b) {
#pragma unroll
for (int i = 0; i < N; i++) {
assign_add(a.data[i], b.data[i]);
}
return a;
}
template <typename T, int N>
DINLINE array_t<float, N> upcast(array_t<T, N> val) {
if constexpr (std::is_same<T, float>::value) {
return val;
} else {
array_t<float, N> out;
#pragma unroll
for (int i = 0; i < N; i++) {
out.data[i] = upcast_s(val.data[i]);
}
return out;
}
}
template <typename O>
DINLINE O downcast(array_t<float, O::size> val) {
if constexpr (std::is_same<typename O::type, float>::value) {
return val;
} else {
O out;
#pragma unroll
for (int i = 0; i < O::size; i++) {
out.data[i] = downcast_s<typename O::type>(val.data[i]);
}
return out;
}
}
#if !defined(USE_ROCM)
static DINLINE void st_flag_release(FlagType* flag_addr, FlagType flag) {
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 700
asm volatile("st.release.sys.global.u32 [%1], %0;" ::"r"(flag),
"l"(flag_addr));
#else
asm volatile("membar.sys; st.volatile.global.u32 [%1], %0;" ::"r"(flag),
"l"(flag_addr));
#endif
}
static DINLINE FlagType ld_flag_acquire(FlagType* flag_addr) {
FlagType flag;
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 700
asm volatile("ld.acquire.sys.global.u32 %0, [%1];"
: "=r"(flag)
: "l"(flag_addr));
#else
asm volatile("ld.volatile.global.u32 %0, [%1]; membar.gl;"
: "=r"(flag)
: "l"(flag_addr));
#endif
return flag;
}
static DINLINE void st_flag_volatile(FlagType* flag_addr, FlagType flag) {
asm volatile("st.volatile.global.u32 [%1], %0;" ::"r"(flag), "l"(flag_addr));
}
static DINLINE FlagType ld_flag_volatile(FlagType* flag_addr) {
FlagType flag;
asm volatile("ld.volatile.global.u32 %0, [%1];"
: "=r"(flag)
: "l"(flag_addr));
return flag;
}
// This function is meant to be used as the first synchronization in the all
// reduce kernel. Thus, it doesn't need to make any visibility guarantees for
// prior memory accesses. Note: volatile writes will not be reordered against
// other volatile writes.
template <int ngpus>
DINLINE void barrier_at_start(const RankSignals& sg, Signal* self_sg,
int rank) {
uint32_t flag = self_sg->_flag[blockIdx.x] + 1;
if (threadIdx.x < ngpus) {
auto peer_counter_ptr = &sg.signals[threadIdx.x]->start[blockIdx.x][rank];
auto self_counter_ptr = &self_sg->start[blockIdx.x][threadIdx.x];
// Write the expected counter value to peer and wait for correct value
// from peer.
st_flag_volatile(peer_counter_ptr, flag);
while (ld_flag_volatile(self_counter_ptr) != flag);
}
__syncthreads();
// use one thread to update flag
if (threadIdx.x == 0) self_sg->_flag[blockIdx.x] = flag;
}
template <int ngpus>
DINLINE void barrier_at_start_release(const RankSignals& sg, Signal* self_sg,
int rank) {
__syncthreads();
uint32_t flag = self_sg->_flag[blockIdx.x] + 1;
if (threadIdx.x < ngpus) {
auto peer_counter_ptr = &sg.signals[threadIdx.x]->start[blockIdx.x][rank];
auto self_counter_ptr = &self_sg->start[blockIdx.x][threadIdx.x];
st_flag_release(peer_counter_ptr, flag);
while (ld_flag_acquire(self_counter_ptr) != flag);
}
__syncthreads();
if (threadIdx.x == 0) self_sg->_flag[blockIdx.x] = flag;
}
// This function is meant to be used as the second or the final
// synchronization barrier in the all reduce kernel. If it's the final
// synchronization barrier, we don't need to make any visibility guarantees
// for prior memory accesses.
template <int ngpus, bool final_sync = false>
DINLINE void barrier_at_end(const RankSignals& sg, Signal* self_sg, int rank) {
__syncthreads();
uint32_t flag = self_sg->_flag[blockIdx.x] + 1;
if (threadIdx.x < ngpus) {
auto peer_counter_ptr = &sg.signals[threadIdx.x]->end[blockIdx.x][rank];
auto self_counter_ptr = &self_sg->end[blockIdx.x][threadIdx.x];
// Write the expected counter value to peer and wait for correct value from
// peer.
if constexpr (!final_sync) {
st_flag_release(peer_counter_ptr, flag);
while (ld_flag_acquire(self_counter_ptr) != flag);
} else {
st_flag_volatile(peer_counter_ptr, flag);
while (ld_flag_volatile(self_counter_ptr) != flag);
}
}
if constexpr (!final_sync) __syncthreads();
// use one thread to update flag
if (threadIdx.x == 0) self_sg->_flag[blockIdx.x] = flag;
}
#else
template <int ngpus>
DINLINE void barrier_at_start(const RankSignals& sg, Signal* self_sg,
int rank) {
uint32_t flag = self_sg->_flag[blockIdx.x] + 1;
if (threadIdx.x < ngpus) {
// simultaneously write to the corresponding flag of all ranks.
// Latency = 1 p2p write
__scoped_atomic_store_n(&sg.signals[threadIdx.x]->start[blockIdx.x][rank],
flag, __ATOMIC_RELAXED, __MEMORY_SCOPE_SYSTEM);
// wait until we got true from all ranks
while (__scoped_atomic_load_n(&self_sg->start[blockIdx.x][threadIdx.x],
__ATOMIC_RELAXED,
__MEMORY_SCOPE_DEVICE) < flag);
}
__syncthreads();
// use one thread to update flag
if (threadIdx.x == 0) self_sg->_flag[blockIdx.x] = flag;
}
template <int ngpus>
DINLINE void barrier_at_start_release(const RankSignals& sg, Signal* self_sg,
int rank) {
__syncthreads();
uint32_t flag = self_sg->_flag[blockIdx.x] + 1;
if (threadIdx.x < ngpus) {
__scoped_atomic_store_n(&sg.signals[threadIdx.x]->start[blockIdx.x][rank],
flag, __ATOMIC_RELEASE, __MEMORY_SCOPE_SYSTEM);
while (__scoped_atomic_load_n(&self_sg->start[blockIdx.x][threadIdx.x],
__ATOMIC_ACQUIRE,
__MEMORY_SCOPE_DEVICE) < flag);
}
__syncthreads();
if (threadIdx.x == 0) self_sg->_flag[blockIdx.x] = flag;
}
template <int ngpus, bool final_sync = false>
DINLINE void barrier_at_end(const RankSignals& sg, Signal* self_sg, int rank) {
__syncthreads();
uint32_t flag = self_sg->_flag[blockIdx.x] + 1;
if (threadIdx.x < ngpus) {
// simultaneously write to the corresponding flag of all ranks.
// Latency = 1 p2p write
__scoped_atomic_store_n(&sg.signals[threadIdx.x]->end[blockIdx.x][rank],
flag,
final_sync ? __ATOMIC_RELAXED : __ATOMIC_RELEASE,
__MEMORY_SCOPE_SYSTEM);
// wait until we got true from all ranks
while (
__scoped_atomic_load_n(&self_sg->end[blockIdx.x][threadIdx.x],
final_sync ? __ATOMIC_RELAXED : __ATOMIC_ACQUIRE,
__MEMORY_SCOPE_DEVICE) < flag);
}
if constexpr (!final_sync) __syncthreads();
// use one thread to update flag
if (threadIdx.x == 0) self_sg->_flag[blockIdx.x] = flag;
}
#endif
template <typename P, int ngpus, typename A>
DINLINE P packed_reduce(const P* ptrs[], int idx) {
A tmp = upcast(ptrs[0][idx]);
#pragma unroll
for (int i = 1; i < ngpus; i++) {
packed_assign_add(tmp, upcast(ptrs[i][idx]));
}
return downcast<P>(tmp);
}
} // namespace vllm
-17
View File
@@ -1,17 +0,0 @@
#include "core/registration.h"
#include "flash_kda.h"
TORCH_LIBRARY(_flashkda_C, m) {
m.def("get_workspace_size(int T_total, int H, int N=1) -> int",
&get_workspace_size);
m.def(
"fwd(Tensor q, Tensor k, Tensor v, Tensor g, Tensor beta, float scale, "
"Tensor(a!) out, Tensor workspace, Tensor A_log, Tensor dt_bias, "
"float lower_bound, "
"Tensor? initial_state=None, Tensor(b!)? final_state=None, "
"Tensor? cu_seqlens=None) -> ()");
}
TORCH_LIBRARY_IMPL(_flashkda_C, CUDA, m) { m.impl("fwd", &fwd); }
REGISTER_EXTENSION(_flashkda_C)
-108
View File
@@ -464,66 +464,6 @@ __global__ void swigluoai_and_mul_kernel(
}
}
// SITU (Kimi SituGLU) gated activation. Non-interleaved layout:
// input = [gate(d), up(d)] per token.
// gate_out = beta * tanh(gate / beta) * sigmoid(gate)
// up_out = (linear_beta > 0) ? linear_beta * tanh(up / linear_beta) : up
// out = gate_out * up_out
// Compute is done in fp32 and written straight to `out` -- no intermediate
// tensors and no full-tensor fp32 upcast (the pure-torch forward_native
// allocated ~8 fp32 temporaries per call, which blows up MoE profiling).
template <typename scalar_t>
__global__ void situ_and_mul_kernel(
scalar_t* __restrict__ out, // [..., d]
const scalar_t* __restrict__ input, // [..., 2, d]
const int d, const float beta, const float linear_beta) {
const int64_t row = blockIdx.x;
const scalar_t* gate_ptr = input + row * 2 * d;
const scalar_t* up_ptr = gate_ptr + d;
scalar_t* out_ptr = out + row * d;
const bool clamp_up = linear_beta > 0.0f;
const float inv_beta = 1.0f / beta;
const float inv_linear_beta = clamp_up ? 1.0f / linear_beta : 0.0f;
for (int64_t idx = threadIdx.x; idx < d; idx += blockDim.x) {
const float g = (float)VLLM_LDG(&gate_ptr[idx]);
const float u = (float)VLLM_LDG(&up_ptr[idx]);
const float gate_out = beta * tanhf(g * inv_beta) / (1.0f + expf(-g));
const float up_out =
clamp_up ? linear_beta * tanhf(u * inv_linear_beta) : u;
out_ptr[idx] = (scalar_t)(gate_out * up_out);
}
}
template <typename scalar_t>
__global__ void masked_situ_and_mul_kernel(
scalar_t* __restrict__ out, const scalar_t* __restrict__ input,
const int* __restrict__ expert_num_tokens, const int max_num_tokens,
const int d, const float beta, const float linear_beta) {
const int expert = blockIdx.y;
const int num_tokens = expert_num_tokens[expert];
const int idx = blockIdx.x * blockDim.x + threadIdx.x;
if (idx >= d || num_tokens == 0) {
return;
}
const bool clamp_up = linear_beta > 0.0f;
const float inv_beta = 1.0f / beta;
const float inv_linear_beta = clamp_up ? 1.0f / linear_beta : 0.0f;
const int64_t expert_row = static_cast<int64_t>(expert) * max_num_tokens;
for (int token = 0; token < num_tokens; ++token) {
const int64_t row = expert_row + token;
const scalar_t* gate_ptr = input + row * 2 * d;
const scalar_t* up_ptr = gate_ptr + d;
scalar_t* out_ptr = out + row * d;
const float g = (float)VLLM_LDG(&gate_ptr[idx]);
const float u = (float)VLLM_LDG(&up_ptr[idx]);
const float gate_out = beta * tanhf(g * inv_beta) / (1.0f + expf(-g));
const float up_out =
clamp_up ? linear_beta * tanhf(u * inv_linear_beta) : u;
out_ptr[idx] = (scalar_t)(gate_out * up_out);
}
}
} // namespace vllm
#define LAUNCH_ACTIVATION_GATE_KERNEL_WITH_PARAM(KERNEL, PACKED_KERNEL, PARAM) \
@@ -613,54 +553,6 @@ void swigluoai_and_mul(torch::stable::Tensor& out, // [..., d]
double alpha, double limit) {
LAUNCH_SIGLUOAI_AND_MUL(vllm::swigluoai_and_mul, alpha, limit);
}
// Kimi SITU gated activation. `linear_beta <= 0` means "unset" (up passed
// through), matching SituAndMul(linear_beta=None) on the Python side.
void situ_and_mul(torch::stable::Tensor& out, // [..., d]
torch::stable::Tensor& input, // [..., 2 * d]
double beta, double linear_beta) {
int d = input.size(-1) / 2;
int64_t num_tokens = input.numel() / input.size(-1);
if (num_tokens == 0) {
return;
}
dim3 grid(num_tokens);
dim3 block(std::min(d, 1024));
const torch::stable::accelerator::DeviceGuard device_guard(
input.get_device_index());
const cudaStream_t stream = get_current_cuda_stream();
VLLM_STABLE_DISPATCH_FLOATING_TYPES(
input.scalar_type(), "situ_and_mul_kernel", [&] {
vllm::situ_and_mul_kernel<scalar_t><<<grid, block, 0, stream>>>(
out.mutable_data_ptr<scalar_t>(), input.const_data_ptr<scalar_t>(),
d, (float)beta, (float)linear_beta);
});
}
void masked_situ_and_mul(torch::stable::Tensor& out, // [E, T, d]
torch::stable::Tensor& input, // [E, T, 2 * d]
const torch::stable::Tensor& expert_num_tokens,
double beta, double linear_beta) {
int num_experts = input.size(0);
int max_num_tokens = input.size(1);
int d = input.size(2) / 2;
if (num_experts == 0 || max_num_tokens == 0) {
return;
}
constexpr int block_size = 256;
dim3 grid((d + block_size - 1) / block_size, num_experts);
dim3 block(block_size);
const torch::stable::accelerator::DeviceGuard device_guard(
input.get_device_index());
const cudaStream_t stream = get_current_cuda_stream();
VLLM_STABLE_DISPATCH_FLOATING_TYPES(
input.scalar_type(), "masked_situ_and_mul_kernel", [&] {
vllm::masked_situ_and_mul_kernel<scalar_t><<<grid, block, 0, stream>>>(
out.mutable_data_ptr<scalar_t>(), input.const_data_ptr<scalar_t>(),
expert_num_tokens.const_data_ptr<int>(), max_num_tokens, d,
(float)beta, (float)linear_beta);
});
}
namespace vllm {
// Element-wise activation kernel template.
@@ -21,10 +21,7 @@ __global__ void merge_attn_states_kernel(
const float* prefix_lse, const scalar_t* suffix_output,
const float* suffix_lse, const uint num_tokens, const uint num_heads,
const uint head_size, const uint prefix_head_stride,
const uint output_head_stride, const uint prefix_lse_head_stride,
const uint prefix_lse_token_stride, const uint suffix_lse_head_stride,
const uint suffix_lse_token_stride, const uint output_lse_head_stride,
const uint output_lse_token_stride, const uint prefix_num_tokens,
const uint output_head_stride, const uint prefix_num_tokens,
const float* output_scale) {
// Inputs always load 128-bit packs (pack_size elements of scalar_t).
// Outputs store pack_size elements of output_t, which is smaller for FP8.
@@ -87,19 +84,15 @@ __global__ void merge_attn_states_kernel(
}
}
if (output_lse != nullptr && pack_idx == 0) {
float s_lse = suffix_lse[head_idx * suffix_lse_head_stride +
token_idx * suffix_lse_token_stride];
output_lse[head_idx * output_lse_head_stride +
token_idx * output_lse_token_stride] = s_lse;
float s_lse = suffix_lse[head_idx * num_tokens + token_idx];
output_lse[head_idx * num_tokens + token_idx] = s_lse;
}
return;
}
// For tokens within prefix range, merge prefix and suffix
float p_lse = prefix_lse[head_idx * prefix_lse_head_stride +
token_idx * prefix_lse_token_stride];
float s_lse = suffix_lse[head_idx * suffix_lse_head_stride +
token_idx * suffix_lse_token_stride];
float p_lse = prefix_lse[head_idx * num_tokens + token_idx];
float s_lse = suffix_lse[head_idx * num_tokens + token_idx];
p_lse = std::isinf(p_lse) ? -std::numeric_limits<float>::infinity() : p_lse;
s_lse = std::isinf(s_lse) ? -std::numeric_limits<float>::infinity() : s_lse;
@@ -139,8 +132,7 @@ __global__ void merge_attn_states_kernel(
}
// We only need to write to output_lse once per head.
if (output_lse != nullptr && pack_idx == 0) {
output_lse[head_idx * output_lse_head_stride +
token_idx * output_lse_token_stride] = max_lse;
output_lse[head_idx * num_tokens + token_idx] = max_lse;
}
return;
}
@@ -195,8 +187,7 @@ __global__ void merge_attn_states_kernel(
// We only need to write to output_lse once per head.
if (output_lse != nullptr && pack_idx == 0) {
float out_lse = logf(out_se) + max_lse;
output_lse[head_idx * output_lse_head_stride +
token_idx * output_lse_token_stride] = out_lse;
output_lse[head_idx * num_tokens + token_idx] = out_lse;
}
}
@@ -230,9 +221,6 @@ __global__ void merge_attn_states_kernel(
reinterpret_cast<scalar_t*>(suffix_output.data_ptr()), \
reinterpret_cast<float*>(suffix_lse.data_ptr()), num_tokens, \
num_heads, head_size, prefix_head_stride, output_head_stride, \
prefix_lse_head_stride, prefix_lse_token_stride, \
suffix_lse_head_stride, suffix_lse_token_stride, \
output_lse_head_stride, output_lse_token_stride, \
prefix_num_tokens, output_scale_ptr); \
}
@@ -271,19 +259,6 @@ void merge_attn_states_launcher(
const uint head_size = output.size(2);
const uint prefix_head_stride = prefix_output.stride(1);
const uint output_head_stride = output.stride(1);
// lse tensors are [NUM_HEADS, NUM_TOKENS] but may be non-contiguous views
// (e.g. a transpose of a backend's [NUM_TOKENS, NUM_HEADS] output), so index
// them by their actual strides rather than assuming a contiguous layout.
const uint prefix_lse_head_stride = prefix_lse.stride(0);
const uint prefix_lse_token_stride = prefix_lse.stride(1);
const uint suffix_lse_head_stride = suffix_lse.stride(0);
const uint suffix_lse_token_stride = suffix_lse.stride(1);
uint output_lse_head_stride = 0;
uint output_lse_token_stride = 0;
if (output_lse.has_value()) {
output_lse_head_stride = output_lse.value().stride(0);
output_lse_token_stride = output_lse.value().stride(1);
}
// Thread mapping is based on input BF16 pack_size
const uint pack_size = 16 / sizeof(scalar_t);
STD_TORCH_CHECK(head_size % pack_size == 0,
-96
View File
@@ -443,55 +443,6 @@ __global__ void concat_and_cache_mla_kernel(
copy(k_pe, kv_cache, k_pe_stride, block_stride, pe_dim, kv_lora_rank);
}
// Grouped variant of concat_and_cache_mla: inserts the context K/V for every
// draft layer in a single launch. Grid is (num_tokens, num_layers); each layer
// reads its own cache base pointer from kv_cache_ptrs (same pointer-array
// pattern as copy_blocks_kernel). bf16 only, so it is a raw 16-bit copy with no
// scaling or quantization; scalar_t is uint16_t for portability.
template <typename scalar_t>
__global__ void concat_and_cache_mla_grouped_kernel(
const scalar_t* __restrict__ kv_c, // [num_layers, num_tokens,
// kv_lora_rank]
const scalar_t* __restrict__ k_pe, // [num_layers, num_tokens, pe_dim]
const int64_t* __restrict__ kv_cache_ptrs, // [num_layers]
const int64_t* __restrict__ slot_mapping, // [num_layers, num_tokens]
const int64_t kv_c_layer_stride, const int64_t kv_c_token_stride,
const int64_t k_pe_layer_stride, const int64_t k_pe_token_stride,
const int64_t slot_layer_stride, const int64_t block_stride,
const int64_t entry_stride, const int kv_lora_rank, const int pe_dim,
const int block_size) {
const int64_t token_idx = blockIdx.x;
const int64_t layer_idx = blockIdx.y;
const int64_t slot_idx =
slot_mapping[layer_idx * slot_layer_stride + token_idx];
// NOTE: slot_idx can be -1 if the token is padded
if (slot_idx < 0) {
return;
}
const int64_t block_idx = slot_idx / block_size;
const int64_t block_offset = slot_idx % block_size;
scalar_t* __restrict__ kv_cache =
reinterpret_cast<scalar_t*>(kv_cache_ptrs[layer_idx]);
const scalar_t* __restrict__ kv_c_layer =
kv_c + layer_idx * kv_c_layer_stride;
const scalar_t* __restrict__ k_pe_layer =
k_pe + layer_idx * k_pe_layer_stride;
auto copy = [&](const scalar_t* __restrict__ src, int64_t src_token_stride,
int size, int offset) {
for (int i = threadIdx.x; i < size; i += blockDim.x) {
const int64_t src_idx = token_idx * src_token_stride + i;
const int64_t dst_idx =
block_idx * block_stride + block_offset * entry_stride + i + offset;
kv_cache[dst_idx] = src[src_idx];
}
};
copy(kv_c_layer, kv_c_token_stride, kv_lora_rank, 0);
copy(k_pe_layer, k_pe_token_stride, pe_dim, kv_lora_rank);
}
template <typename scalar_t, typename cache_t, Fp8KVCacheDataType kv_dt>
__global__ void concat_and_cache_ds_mla_kernel(
const scalar_t* __restrict__ kv_c, // [num_tokens, kv_lora_rank]
@@ -951,53 +902,6 @@ void concat_and_cache_mla(
}
}
void concat_and_cache_mla_grouped(
torch::stable::Tensor& kv_c, // [num_layers, num_tokens, kv_lora_rank]
torch::stable::Tensor& k_pe, // [num_layers, num_tokens, pe_dim]
torch::stable::Tensor& kv_cache_ptrs, // [num_layers] int64, on device
torch::stable::Tensor& slot_mapping, // [num_layers, num_tokens] int64
int64_t block_size, int64_t block_stride, int64_t entry_stride) {
int num_layers = kv_c.size(0);
int num_tokens = kv_c.size(1);
int kv_lora_rank = kv_c.size(2);
int pe_dim = k_pe.size(2);
STD_TORCH_CHECK(
kv_c.scalar_type() == torch::headeronly::ScalarType::BFloat16 &&
k_pe.scalar_type() == torch::headeronly::ScalarType::BFloat16,
"concat_and_cache_mla_grouped only supports a bf16 KV cache; got kv_c=",
kv_c.scalar_type(), ", k_pe=", k_pe.scalar_type());
STD_TORCH_CHECK(
kv_cache_ptrs.scalar_type() == torch::headeronly::ScalarType::Long,
"kv_cache_ptrs must be int64");
if (num_tokens == 0 || num_layers == 0) {
return;
}
const int64_t kv_c_layer_stride = kv_c.stride(0);
const int64_t kv_c_token_stride = kv_c.stride(1);
const int64_t k_pe_layer_stride = k_pe.stride(0);
const int64_t k_pe_token_stride = k_pe.stride(1);
const int64_t slot_layer_stride = slot_mapping.stride(0);
const torch::stable::accelerator::DeviceGuard device_guard(
kv_c.get_device_index());
const cudaStream_t stream = get_current_cuda_stream();
dim3 grid(num_tokens, num_layers);
dim3 block(std::min(kv_lora_rank, 512));
vllm::concat_and_cache_mla_grouped_kernel<uint16_t>
<<<grid, block, 0, stream>>>(
reinterpret_cast<const uint16_t*>(kv_c.data_ptr()),
reinterpret_cast<const uint16_t*>(k_pe.data_ptr()),
kv_cache_ptrs.const_data_ptr<int64_t>(),
slot_mapping.const_data_ptr<int64_t>(), kv_c_layer_stride,
kv_c_token_stride, k_pe_layer_stride, k_pe_token_stride,
slot_layer_stride, block_stride, entry_stride, kv_lora_rank, pe_dim,
block_size);
}
namespace vllm {
template <typename Tout, typename Tin, Fp8KVCacheDataType kv_dt>
@@ -1,362 +0,0 @@
#include "torch_utils.h"
#include <torch/csrc/stable/macros.h>
#include <torch/csrc/stable/accelerator.h>
#include <torch/csrc/stable/tensor.h>
#include <torch/headeronly/core/ScalarType.h>
#include "custom_all_reduce.cuh"
#include "custom_all_gather_reduce_scatter.cuh"
namespace vllm {
void CustomAllreduce::allgather(cudaStream_t stream, void* input, void* output,
int size_bytes, int threads, int block_limit) {
if (size_bytes % sizeof(CopyPack) != 0)
throw std::runtime_error(
"custom allgather requires input byte size to be a multiple of " +
std::to_string(sizeof(CopyPack)));
auto ptrs = buffers_.at(input);
int size_per_rank = size_bytes / sizeof(CopyPack);
int total_size = size_per_rank * world_size_;
int blocks = std::min(block_limit, (total_size + threads - 1) / threads);
#define AG_CASE(ngpus) \
case ngpus: \
cross_device_all_gather<ngpus><<<blocks, threads, 0, stream>>>( \
ptrs, sg_, self_sg_, reinterpret_cast<CopyPack*>(output), rank_, \
size_per_rank); \
break;
switch (world_size_) {
AG_CASE(2)
AG_CASE(4)
AG_CASE(6)
AG_CASE(8)
default:
throw std::runtime_error(
"custom allgather only supports num gpus in (2,4,6,8)");
}
#undef AG_CASE
}
template <typename T>
void CustomAllreduce::mnnvl_lamport_allgather(cudaStream_t stream, T* input,
T* output, void* local_buffer,
void* multicast_buffer,
uint32_t* epochs, int size_bytes,
int stage_size_bytes) {
if (size_bytes % sizeof(typename packed_t<T>::P) != 0 ||
stage_size_bytes % sizeof(typename packed_t<T>::P) != 0)
throw std::runtime_error(
"MNNVL Lamport allgather requires 16-byte aligned sizes");
auto ptrs = buffers_.at(local_buffer);
int size_per_rank = size_bytes / sizeof(typename packed_t<T>::P);
int stage_size = stage_size_bytes / sizeof(typename packed_t<T>::P);
int blocks =
(size_per_rank + kMnnvlLamportAgThreads - 1) / kMnnvlLamportAgThreads;
#if !defined(USE_ROCM) && CUDA_VERSION >= 12000
cudaLaunchAttribute attributes[1]{};
attributes[0].id = cudaLaunchAttributeProgrammaticStreamSerialization;
attributes[0].val.programmaticStreamSerializationAllowed = 1;
cudaLaunchConfig_t config{.gridDim = dim3(blocks),
.blockDim = dim3(kMnnvlLamportAgThreads),
.dynamicSmemBytes = 0,
.stream = stream,
.attrs = attributes,
.numAttrs = 1};
#define MNNVL_LAMPORT_AG_LAUNCH(ngpus) \
CUDACHECK(cudaLaunchKernelEx(&config, &mnnvl_lamport_all_gather<T, ngpus>, \
ptrs, input, output, \
reinterpret_cast<T*>(multicast_buffer), \
epochs, rank_, size_per_rank, stage_size))
#else
#define MNNVL_LAMPORT_AG_LAUNCH(ngpus) \
mnnvl_lamport_all_gather<T, ngpus> \
<<<blocks, kMnnvlLamportAgThreads, 0, stream>>>( \
ptrs, input, output, reinterpret_cast<T*>(multicast_buffer), \
epochs, rank_, size_per_rank, stage_size)
#endif
#define MNNVL_LAMPORT_AG_CASE(ngpus) \
case ngpus: \
MNNVL_LAMPORT_AG_LAUNCH(ngpus); \
break;
switch (world_size_) {
MNNVL_LAMPORT_AG_CASE(2)
MNNVL_LAMPORT_AG_CASE(4)
MNNVL_LAMPORT_AG_CASE(6)
MNNVL_LAMPORT_AG_CASE(8)
MNNVL_LAMPORT_AG_CASE(16)
default:
throw std::runtime_error(
"MNNVL Lamport allgather only supports num gpus in (2,4,6,8,16)");
}
#undef MNNVL_LAMPORT_AG_CASE
#undef MNNVL_LAMPORT_AG_LAUNCH
}
template <typename T>
void CustomAllreduce::reduce_scatter(cudaStream_t stream, T* input, T* output,
int size, int threads, int block_limit) {
auto packed_size = packed_t<T>::P::size;
if (size % (packed_size * world_size_) != 0)
throw std::runtime_error(
"custom reduce-scatter requires each output shard byte size to be "
"a multiple of 16");
auto ptrs = buffers_.at(input);
int size_per_rank = size / packed_size / world_size_;
int blocks = std::min(block_limit, (size_per_rank + threads - 1) / threads);
#define RS_CASE(ngpus) \
case ngpus: \
cross_device_reduce_scatter<T, ngpus><<<blocks, threads, 0, stream>>>( \
ptrs, sg_, self_sg_, output, rank_, size_per_rank); \
break;
switch (world_size_) {
RS_CASE(2)
RS_CASE(4)
RS_CASE(6)
RS_CASE(8)
default:
throw std::runtime_error(
"custom reduce-scatter only supports num gpus in (2,4,6,8)");
}
#undef RS_CASE
}
template <typename T>
void CustomAllreduce::mnnvl_lamport_reduce_scatter(cudaStream_t stream,
T* input, T* output,
void* local_buffer,
uint32_t* epochs, int size,
int stage_size_bytes) {
auto packed_size = packed_t<T>::P::size;
if (size % (packed_size * world_size_) != 0 ||
stage_size_bytes % sizeof(typename packed_t<T>::P) != 0)
throw std::runtime_error(
"MNNVL Lamport reduce-scatter requires 16-byte aligned sizes");
auto ptrs = buffers_.at(local_buffer);
int size_per_rank = size / packed_size / world_size_;
int stage_size = stage_size_bytes / sizeof(typename packed_t<T>::P);
int blocks_per_rank =
(size_per_rank + kMnnvlLamportRsThreads - 1) / kMnnvlLamportRsThreads;
int blocks = blocks_per_rank * world_size_;
#if !defined(USE_ROCM) && CUDA_VERSION >= 12000
cudaLaunchAttribute attributes[1]{};
attributes[0].id = cudaLaunchAttributeProgrammaticStreamSerialization;
attributes[0].val.programmaticStreamSerializationAllowed = 1;
cudaLaunchConfig_t config{.gridDim = dim3(blocks),
.blockDim = dim3(kMnnvlLamportRsThreads),
.dynamicSmemBytes = 0,
.stream = stream,
.attrs = attributes,
.numAttrs = 1};
#define MNNVL_LAMPORT_RS_LAUNCH(ngpus) \
CUDACHECK(cudaLaunchKernelEx( \
&config, &mnnvl_lamport_reduce_scatter_kernel<T, ngpus>, ptrs, input, \
output, epochs, rank_, size_per_rank, stage_size))
#else
#define MNNVL_LAMPORT_RS_LAUNCH(ngpus) \
mnnvl_lamport_reduce_scatter_kernel<T, ngpus> \
<<<blocks, kMnnvlLamportRsThreads, 0, stream>>>( \
ptrs, input, output, epochs, rank_, size_per_rank, stage_size)
#endif
#define MNNVL_LAMPORT_RS_CASE(ngpus) \
case ngpus: \
MNNVL_LAMPORT_RS_LAUNCH(ngpus); \
break;
switch (world_size_) {
MNNVL_LAMPORT_RS_CASE(2)
MNNVL_LAMPORT_RS_CASE(4)
MNNVL_LAMPORT_RS_CASE(6)
MNNVL_LAMPORT_RS_CASE(8)
MNNVL_LAMPORT_RS_CASE(16)
default:
throw std::runtime_error(
"MNNVL Lamport reduce-scatter only supports num gpus in "
"(2,4,6,8,16)");
}
#undef MNNVL_LAMPORT_RS_CASE
#undef MNNVL_LAMPORT_RS_LAUNCH
}
} // namespace vllm
using fptr_t = int64_t;
static_assert(sizeof(void*) == sizeof(fptr_t));
bool _is_weak_contiguous(torch::stable::Tensor& t);
void custom_all_gather(fptr_t _fa, torch::stable::Tensor& inp,
torch::stable::Tensor& out, fptr_t _reg_buffer,
int64_t reg_buffer_sz_bytes) {
auto fa = reinterpret_cast<vllm::CustomAllreduce*>(_fa);
const torch::stable::accelerator::DeviceGuard device_guard(
inp.get_device_index());
const cudaStream_t stream = get_current_cuda_stream(inp.get_device_index());
STD_TORCH_CHECK((inp.scalar_type()) == (out.scalar_type()));
STD_TORCH_CHECK((inp.numel() * fa->world_size_) == (out.numel()));
STD_TORCH_CHECK(_is_weak_contiguous(out));
STD_TORCH_CHECK(_is_weak_contiguous(inp));
auto input_size = inp.numel() * inp.element_size();
auto reg_buffer = reinterpret_cast<void*>(_reg_buffer);
STD_TORCH_CHECK(reg_buffer != nullptr);
STD_TORCH_CHECK((input_size) <= (reg_buffer_sz_bytes));
STD_CUDA_CHECK(cudaMemcpyAsync(reg_buffer, inp.const_data_ptr(), input_size,
cudaMemcpyDeviceToDevice, stream));
fa->allgather(stream, reg_buffer, out.mutable_data_ptr(), input_size);
}
void mnnvl_lamport_all_gather(fptr_t _fa, torch::stable::Tensor& inp,
torch::stable::Tensor& out, fptr_t _local_buffer,
fptr_t _multicast_buffer, fptr_t _epoch_buffer,
int64_t stage_sz_bytes) {
auto fa = reinterpret_cast<vllm::CustomAllreduce*>(_fa);
const torch::stable::accelerator::DeviceGuard device_guard(
inp.get_device_index());
const cudaStream_t stream = get_current_cuda_stream(inp.get_device_index());
STD_TORCH_CHECK((inp.scalar_type()) == (out.scalar_type()));
STD_TORCH_CHECK((inp.numel() * fa->world_size_) == (out.numel()));
STD_TORCH_CHECK(_is_weak_contiguous(out));
STD_TORCH_CHECK(_is_weak_contiguous(inp));
auto input_size = inp.numel() * inp.element_size();
STD_TORCH_CHECK((input_size * fa->world_size_) <= stage_sz_bytes);
auto local_buffer = reinterpret_cast<void*>(_local_buffer);
auto multicast_buffer = reinterpret_cast<void*>(_multicast_buffer);
auto epochs = reinterpret_cast<uint32_t*>(_epoch_buffer);
switch (out.scalar_type()) {
case torch::headeronly::ScalarType::Float: {
fa->mnnvl_lamport_allgather<float>(
stream, reinterpret_cast<float*>(inp.mutable_data_ptr()),
reinterpret_cast<float*>(out.mutable_data_ptr()), local_buffer,
multicast_buffer, epochs, input_size, stage_sz_bytes);
break;
}
case torch::headeronly::ScalarType::Half: {
fa->mnnvl_lamport_allgather<half>(
stream, reinterpret_cast<half*>(inp.mutable_data_ptr()),
reinterpret_cast<half*>(out.mutable_data_ptr()), local_buffer,
multicast_buffer, epochs, input_size, stage_sz_bytes);
break;
}
#if (__CUDA_ARCH__ >= 800 || !defined(__CUDA_ARCH__))
case torch::headeronly::ScalarType::BFloat16: {
fa->mnnvl_lamport_allgather<nv_bfloat16>(
stream, reinterpret_cast<nv_bfloat16*>(inp.mutable_data_ptr()),
reinterpret_cast<nv_bfloat16*>(out.mutable_data_ptr()), local_buffer,
multicast_buffer, epochs, input_size, stage_sz_bytes);
break;
}
#endif
default:
throw std::runtime_error(
"MNNVL Lamport allgather only supports float32, float16 and "
"bfloat16");
}
}
void custom_reduce_scatter(fptr_t _fa, torch::stable::Tensor& inp,
torch::stable::Tensor& out, fptr_t _reg_buffer,
int64_t reg_buffer_sz_bytes) {
auto fa = reinterpret_cast<vllm::CustomAllreduce*>(_fa);
const torch::stable::accelerator::DeviceGuard device_guard(
inp.get_device_index());
const cudaStream_t stream = get_current_cuda_stream(inp.get_device_index());
STD_TORCH_CHECK((inp.scalar_type()) == (out.scalar_type()));
STD_TORCH_CHECK((out.numel() * fa->world_size_) == (inp.numel()));
STD_TORCH_CHECK(_is_weak_contiguous(out));
STD_TORCH_CHECK(_is_weak_contiguous(inp));
auto input_size = inp.numel() * inp.element_size();
auto reg_buffer = reinterpret_cast<void*>(_reg_buffer);
STD_TORCH_CHECK(reg_buffer != nullptr);
STD_TORCH_CHECK((input_size) <= (reg_buffer_sz_bytes));
STD_CUDA_CHECK(cudaMemcpyAsync(reg_buffer, inp.const_data_ptr(), input_size,
cudaMemcpyDeviceToDevice, stream));
switch (out.scalar_type()) {
case torch::headeronly::ScalarType::Float: {
fa->reduce_scatter<float>(
stream, reinterpret_cast<float*>(reg_buffer),
reinterpret_cast<float*>(out.mutable_data_ptr()), inp.numel());
break;
}
case torch::headeronly::ScalarType::Half: {
fa->reduce_scatter<half>(stream, reinterpret_cast<half*>(reg_buffer),
reinterpret_cast<half*>(out.mutable_data_ptr()),
inp.numel());
break;
}
#if (__CUDA_ARCH__ >= 800 || !defined(__CUDA_ARCH__))
case torch::headeronly::ScalarType::BFloat16: {
fa->reduce_scatter<nv_bfloat16>(
stream, reinterpret_cast<nv_bfloat16*>(reg_buffer),
reinterpret_cast<nv_bfloat16*>(out.mutable_data_ptr()), inp.numel());
break;
}
#endif
default:
throw std::runtime_error(
"custom reduce-scatter only supports float32, float16 and bfloat16");
}
}
void mnnvl_lamport_reduce_scatter(fptr_t _fa, torch::stable::Tensor& inp,
torch::stable::Tensor& out,
fptr_t _local_buffer, fptr_t _epoch_buffer,
int64_t stage_sz_bytes) {
auto fa = reinterpret_cast<vllm::CustomAllreduce*>(_fa);
const torch::stable::accelerator::DeviceGuard device_guard(
inp.get_device_index());
const cudaStream_t stream = get_current_cuda_stream(inp.get_device_index());
STD_TORCH_CHECK((inp.scalar_type()) == (out.scalar_type()));
STD_TORCH_CHECK((out.numel() * fa->world_size_) == (inp.numel()));
STD_TORCH_CHECK(_is_weak_contiguous(out));
STD_TORCH_CHECK(_is_weak_contiguous(inp));
auto input_size = inp.numel() * inp.element_size();
STD_TORCH_CHECK(input_size <= stage_sz_bytes);
auto local_buffer = reinterpret_cast<void*>(_local_buffer);
auto epochs = reinterpret_cast<uint32_t*>(_epoch_buffer);
switch (out.scalar_type()) {
case torch::headeronly::ScalarType::Float: {
fa->mnnvl_lamport_reduce_scatter<float>(
stream, reinterpret_cast<float*>(inp.mutable_data_ptr()),
reinterpret_cast<float*>(out.mutable_data_ptr()), local_buffer,
epochs, inp.numel(), stage_sz_bytes);
break;
}
case torch::headeronly::ScalarType::Half: {
fa->mnnvl_lamport_reduce_scatter<half>(
stream, reinterpret_cast<half*>(inp.mutable_data_ptr()),
reinterpret_cast<half*>(out.mutable_data_ptr()), local_buffer, epochs,
inp.numel(), stage_sz_bytes);
break;
}
#if (__CUDA_ARCH__ >= 800 || !defined(__CUDA_ARCH__))
case torch::headeronly::ScalarType::BFloat16: {
fa->mnnvl_lamport_reduce_scatter<nv_bfloat16>(
stream, reinterpret_cast<nv_bfloat16*>(inp.mutable_data_ptr()),
reinterpret_cast<nv_bfloat16*>(out.mutable_data_ptr()), local_buffer,
epochs, inp.numel(), stage_sz_bytes);
break;
}
#endif
default:
throw std::runtime_error(
"MNNVL Lamport reduce-scatter only supports float32, float16 and "
"bfloat16");
}
}
@@ -1,29 +0,0 @@
#include "ops.h"
#include "core/registration.h"
#include <torch/csrc/stable/library.h>
STABLE_TORCH_LIBRARY_FRAGMENT(_C_custom_ar, custom_ag_rs) {
custom_ag_rs.def(
"custom_all_gather(int fa, Tensor inp, Tensor! out, int reg_buffer, "
"int reg_buffer_sz_bytes) -> ()");
custom_ag_rs.def(
"mnnvl_lamport_all_gather(int fa, Tensor inp, Tensor! out, int "
"local_buffer, int multicast_buffer, int epoch_buffer, int "
"stage_sz_bytes) -> ()");
custom_ag_rs.def(
"custom_reduce_scatter(int fa, Tensor inp, Tensor! out, int reg_buffer, "
"int reg_buffer_sz_bytes) -> ()");
custom_ag_rs.def(
"mnnvl_lamport_reduce_scatter(int fa, Tensor inp, Tensor! out, int "
"local_buffer, int epoch_buffer, int stage_sz_bytes) -> ()");
}
STABLE_TORCH_LIBRARY_IMPL(_C_custom_ar, CUDA, custom_ag_rs) {
custom_ag_rs.impl("custom_all_gather", TORCH_BOX(&custom_all_gather));
custom_ag_rs.impl("mnnvl_lamport_all_gather",
TORCH_BOX(&mnnvl_lamport_all_gather));
custom_ag_rs.impl("custom_reduce_scatter", TORCH_BOX(&custom_reduce_scatter));
custom_ag_rs.impl("mnnvl_lamport_reduce_scatter",
TORCH_BOX(&mnnvl_lamport_reduce_scatter));
}
+4 -4
View File
@@ -18,14 +18,14 @@ fptr_t init_custom_ar(const std::vector<fptr_t>& fake_ipc_ptrs,
torch::stable::Tensor& rank_data, int64_t rank,
bool fully_connected) {
int world_size = fake_ipc_ptrs.size();
if (world_size > vllm::kMaxCustomCollectiveRanks)
throw std::invalid_argument("world size > 16 is not supported");
if (world_size > 8)
throw std::invalid_argument("world size > 8 is not supported");
if (world_size % 2 != 0)
throw std::invalid_argument("Odd num gpus is not supported for now");
if (rank < 0 || rank >= world_size)
throw std::invalid_argument("invalid rank passed in");
vllm::Signal* ipc_ptrs[vllm::kMaxCustomCollectiveRanks];
vllm::Signal* ipc_ptrs[8];
for (int i = 0; i < world_size; i++) {
ipc_ptrs[i] = reinterpret_cast<vllm::Signal*>(fake_ipc_ptrs[i]);
}
@@ -124,7 +124,7 @@ int64_t meta_size() { return sizeof(vllm::Signal); }
void register_buffer(fptr_t _fa, const std::vector<fptr_t>& fake_ipc_ptrs) {
auto fa = reinterpret_cast<vllm::CustomAllreduce*>(_fa);
STD_TORCH_CHECK(fake_ipc_ptrs.size() == fa->world_size_);
void* ipc_ptrs[vllm::kMaxCustomCollectiveRanks];
void* ipc_ptrs[8];
for (int i = 0; i < fake_ipc_ptrs.size(); i++) {
ipc_ptrs[i] = reinterpret_cast<void*>(fake_ipc_ptrs[i]);
}
+35 -115
View File
@@ -647,17 +647,17 @@ __global__ __launch_bounds__(256, 1) void fused_a_gemm_kernel(
#endif
}
template <typename T, int kHdIn, int kHdOut, int kTileN, int kTileK = 256>
template <typename T, int kHdIn, int kHdOut, int kTileN>
void invokeFusedAGemm(T* output, T const* mat_a, T const* mat_b, int num_tokens,
cudaStream_t const stream, bool enable_pdl) {
constexpr int gemm_m = kHdOut;
int const gemm_n = num_tokens;
constexpr int gemm_k = kHdIn;
cudaStream_t const stream) {
constexpr int gemm_m = kHdOut; // 2112
int const gemm_n = num_tokens; // 1-16
constexpr int gemm_k = kHdIn; // 7168
constexpr int batch_size = 1;
std::swap(mat_a, mat_b);
constexpr int tile_m = 16;
constexpr int tile_n = kTileN;
constexpr int tile_k = kTileK;
constexpr int tile_n = kTileN; // 8 or 16
constexpr int tile_k = std::max(256, 1024 / tile_n); // 256
constexpr int max_stage_cnt =
1024 * 192 / ((tile_m + tile_n) * tile_k * sizeof(bf16_t));
constexpr int k_iter_cnt = gemm_k / tile_k;
@@ -679,8 +679,7 @@ void invokeFusedAGemm(T* output, T const* mat_a, T const* mat_b, int num_tokens,
config.stream = stream;
cudaLaunchAttribute attrs[1];
attrs[0].id = cudaLaunchAttributeProgrammaticStreamSerialization;
attrs[0].val.programmaticStreamSerializationAllowed =
enable_pdl || getEnvEnablePDL();
attrs[0].val.programmaticStreamSerializationAllowed = getEnvEnablePDL();
config.numAttrs = 1;
config.attrs = attrs;
if (smem_bytes >= (48 * 1024)) {
@@ -695,48 +694,36 @@ void invokeFusedAGemm(T* output, T const* mat_a, T const* mat_b, int num_tokens,
output, mat_a, mat_b, gemm_n);
}
template <typename T, int kHdIn, int kHdOut, int kTileK = 256>
void invokeFusedAGemmForTokens(T* output, T const* mat_a, T const* mat_b,
int num_tokens, cudaStream_t const stream,
bool enable_pdl) {
if (num_tokens <= 8) {
invokeFusedAGemm<T, kHdIn, kHdOut, 8, kTileK>(
output, mat_a, mat_b, num_tokens, stream, enable_pdl);
} else {
invokeFusedAGemm<T, kHdIn, kHdOut, 16, kTileK>(
output, mat_a, mat_b, num_tokens, stream, enable_pdl);
}
}
template void invokeFusedAGemm<__nv_bfloat16, 7168, 2112, 8>(
__nv_bfloat16*, __nv_bfloat16 const*, __nv_bfloat16 const*, int num_tokens,
cudaStream_t);
template void invokeFusedAGemm<__nv_bfloat16, 7168, 2112, 16>(
__nv_bfloat16*, __nv_bfloat16 const*, __nv_bfloat16 const*, int num_tokens,
cudaStream_t);
void dsv3_fused_a_gemm(torch::stable::Tensor& output,
torch::stable::Tensor const& mat_a,
torch::stable::Tensor const& mat_b, bool enable_pdl) {
torch::stable::Tensor const& mat_b) {
STD_TORCH_CHECK(mat_a.dim() == 2 && mat_b.dim() == 2 && output.dim() == 2);
int const num_tokens = mat_a.size(0);
int const hd_in = mat_a.size(1);
int const hd_out = mat_b.size(1);
constexpr int kHdIn = 7168;
constexpr int kHdOut = 2112;
STD_TORCH_CHECK(num_tokens >= 1 && num_tokens <= 16,
"required 1 <= mat_a.shape[0] <= 16");
STD_TORCH_CHECK(hd_in == kHdIn, "required mat_a.shape[1] == 7168");
STD_TORCH_CHECK(hd_out == kHdOut, "required mat_b.shape[1] == 2112");
STD_TORCH_CHECK(output.size(0) == num_tokens,
"required output.shape[0] == mat_a.shape[0]");
STD_TORCH_CHECK(output.size(1) == hd_out,
"required output.shape[1] == mat_b.shape[1]");
STD_TORCH_CHECK(mat_b.size(0) == hd_in,
"required mat_b.shape[0] == mat_a.shape[1]");
STD_TORCH_CHECK(mat_a.get_device_index() == mat_b.get_device_index() &&
mat_a.get_device_index() == output.get_device_index(),
"mat_a, mat_b, and output must be on the same device");
// The kernels index global memory with raw pointers and packed strides, so
// reject any padded or transposed view rather than reading out of bounds.
STD_TORCH_CHECK(mat_a.stride(0) == hd_in && mat_a.stride(1) == 1,
"mat_a must be a packed row-major [num_tokens, hd_in] tensor");
STD_TORCH_CHECK(output.stride(0) == hd_out && output.stride(1) == 1,
"output must be a packed row-major [num_tokens, hd_out] tensor");
STD_TORCH_CHECK(mat_b.stride(0) == 1 && mat_b.stride(1) == hd_in,
"mat_b must be a packed column-major [hd_in, hd_out] tensor");
STD_TORCH_CHECK(mat_a.stride(1) == 1, "mat_a must be a row major tensor");
STD_TORCH_CHECK(output.stride(1) == 1, "output must be a row major tensor");
STD_TORCH_CHECK(mat_b.stride(0) == 1, "mat_b must be a column major tensor");
STD_TORCH_CHECK(
mat_a.scalar_type() == torch::headeronly::ScalarType::BFloat16 &&
@@ -751,86 +738,19 @@ void dsv3_fused_a_gemm(torch::stable::Tensor& output,
STD_TORCH_CHECK(getSMVersion() >= 90, "required CUDA ARCH >= SM_90");
auto stream = get_current_cuda_stream(mat_a.get_device_index());
auto* output_ptr =
reinterpret_cast<__nv_bfloat16*>(output.mutable_data_ptr());
auto const* mat_a_ptr =
reinterpret_cast<__nv_bfloat16 const*>(mat_a.data_ptr());
auto const* mat_b_ptr =
reinterpret_cast<__nv_bfloat16 const*>(mat_b.data_ptr());
#define DISPATCH_DSV3_SHAPE(HD_IN, HD_OUT) \
if (hd_in == HD_IN && hd_out == HD_OUT) { \
invokeFusedAGemmForTokens<__nv_bfloat16, HD_IN, HD_OUT>( \
output_ptr, mat_a_ptr, mat_b_ptr, num_tokens, stream, \
enable_pdl); \
return; \
if (num_tokens <= 8) {
invokeFusedAGemm<__nv_bfloat16, kHdIn, kHdOut, 8>(
reinterpret_cast<__nv_bfloat16*>(output.mutable_data_ptr()),
reinterpret_cast<__nv_bfloat16 const*>(mat_a.data_ptr()),
reinterpret_cast<__nv_bfloat16 const*>(mat_b.data_ptr()), num_tokens,
stream);
} else {
invokeFusedAGemm<__nv_bfloat16, kHdIn, kHdOut, 16>(
reinterpret_cast<__nv_bfloat16*>(output.mutable_data_ptr()),
reinterpret_cast<__nv_bfloat16 const*>(mat_a.data_ptr()),
reinterpret_cast<__nv_bfloat16 const*>(mat_b.data_ptr()), num_tokens,
stream);
}
// Shapes the Kimi-K3 selector routes to dsv3_fused_a (see the dsv3 winners
// in KIMI_K3_PROJECTIONS) plus the DeepSeek V2/V3 QKV A-projection.
DISPATCH_DSV3_SHAPE(7168, 1536)
DISPATCH_DSV3_SHAPE(7168, 2112)
DISPATCH_DSV3_SHAPE(1536, 2304)
DISPATCH_DSV3_SHAPE(1536, 4608)
DISPATCH_DSV3_SHAPE(7168, 3584)
DISPATCH_DSV3_SHAPE(768, 7168)
// TP16 dsv3 winners, as (hd_in=K, hd_out=N). TP16 dense down_proj is absent
// because hd_in=2112 is not a multiple of any supported tile_k.
DISPATCH_DSV3_SHAPE(1536, 1152)
DISPATCH_DSV3_SHAPE(7168, 768)
DISPATCH_DSV3_SHAPE(7168, 3216)
DISPATCH_DSV3_SHAPE(7168, 4224)
#ifdef VLLM_K3_BENCH_SHAPES
// The selector routes these shapes to CuTe or the default GEMM, so they are
// never reached in production. They are compiled only for offline
// DSV3-vs-CuTe benchmarking.
DISPATCH_DSV3_SHAPE(7168, 6288)
DISPATCH_DSV3_SHAPE(1536, 7168)
DISPATCH_DSV3_SHAPE(3584, 7168)
DISPATCH_DSV3_SHAPE(7168, 8448)
DISPATCH_DSV3_SHAPE(7168, 20480)
DISPATCH_DSV3_SHAPE(7168, 3072)
DISPATCH_DSV3_SHAPE(7168, 12448)
DISPATCH_DSV3_SHAPE(3072, 7168)
DISPATCH_DSV3_SHAPE(8448, 7168)
DISPATCH_DSV3_SHAPE(7168, 16896)
DISPATCH_DSV3_SHAPE(7168, 40960)
#endif
#undef DISPATCH_DSV3_SHAPE
if (hd_in == 128 && hd_out == 1536) {
invokeFusedAGemmForTokens<__nv_bfloat16, 128, 1536, 128>(
output_ptr, mat_a_ptr, mat_b_ptr, num_tokens, stream, enable_pdl);
return;
}
if (hd_in == 128 && hd_out == 3072) {
invokeFusedAGemmForTokens<__nv_bfloat16, 128, 3072, 128>(
output_ptr, mat_a_ptr, mat_b_ptr, num_tokens, stream, enable_pdl);
return;
}
// TP16 KDA f_b_proj and shared_expert down_proj. Neither hd_in is a multiple
// of 256, so both need the 128 tile_k.
if (hd_in == 128 && hd_out == 768) {
invokeFusedAGemmForTokens<__nv_bfloat16, 128, 768, 128>(
output_ptr, mat_a_ptr, mat_b_ptr, num_tokens, stream, enable_pdl);
return;
}
if (hd_in == 384 && hd_out == 7168) {
invokeFusedAGemmForTokens<__nv_bfloat16, 384, 7168, 128>(
output_ptr, mat_a_ptr, mat_b_ptr, num_tokens, stream, enable_pdl);
return;
}
#ifdef VLLM_K3_BENCH_SHAPES
if (hd_in == 4224 && hd_out == 7168) {
invokeFusedAGemmForTokens<__nv_bfloat16, 4224, 7168, 128>(
output_ptr, mat_a_ptr, mat_b_ptr, num_tokens, stream, enable_pdl);
return;
}
#endif
STD_TORCH_CHECK(false, "unsupported DSV3 fused-A GEMM shape");
}
STABLE_TORCH_LIBRARY_IMPL(_C, CUDA, m) {
File diff suppressed because it is too large Load Diff
@@ -1,954 +0,0 @@
/*
* Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.
*/
// Production AttnRes forward for Blackwell (SM100).
//
// Warp-specialized online softmax + residual + RMSNorm:
// - 1 producer warp issues cp.async.bulk row loads into shared memory.
// - 8 consumer warps compute reductions and output.
// - Q=res_weight*rms_weight remains in registers across persistent tokens.
// - V rows are converted once and cached as FP32 in TMEM between passes.
//
// Integration contract: Kimi K3 H=7168, 1<=num_blocks<=8, and token-major
// block residual storage.
#include "../torch_utils.h"
#include <cfloat>
#include <cstdint>
#include <cstdio>
#include <cuda_runtime.h>
#include <type_traits>
using bf16_t = __nv_bfloat16;
namespace sm100 {
namespace fwd_prod_v2 {
constexpr int K_TILE = 1024;
constexpr int N_CHUNK_DEFAULT = 4;
constexpr int CHUNK_DEPTH = 2;
constexpr int BLK = 288; // 1 producer warp + 8 consumer warps
constexpr int CONSUMER_THREADS = BLK - 32; // 256
constexpr int CONSUMER_WARPS = CONSUMER_THREADS / 32;
constexpr int CONSUMER_GROUPS = 2; // two 128-thread consumer groups
constexpr int CONSUMER_THREADS_PER_GROUP = CONSUMER_THREADS / CONSUMER_GROUPS;
constexpr int FIRST_USER_NAMED_BARRIER = 8;
__device__ __forceinline__ const bf16_t* residual_addr(
const bf16_t* block_res, const bf16_t* layer_res, int source, int N,
int token, int block_stride_m, int block_stride_r, int H) {
if (source < N - 1) {
return block_res + static_cast<long long>(token) * block_stride_m +
source * block_stride_r;
}
return layer_res + static_cast<long long>(token) * H;
}
__device__ __forceinline__ uint32_t elect_one_sync() {
uint32_t pred = 0;
uint32_t laneid = 0;
asm volatile(
"{\n"
".reg .b32 %%rx;\n"
".reg .pred %%px;\n"
" elect.sync %%rx|%%px, %2;\n"
"@%%px mov.s32 %1, 1;\n"
" mov.s32 %0, %%rx;\n"
"}\n"
: "+r"(laneid), "+r"(pred)
: "r"(0xffffffff));
return pred;
}
__device__ __forceinline__ void mbarrier_init(uint64_t& barrier,
int thread_count) {
uint32_t const barrier_addr =
static_cast<uint32_t>(__cvta_generic_to_shared(&barrier));
asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;\n" ::"r"(barrier_addr),
"r"(thread_count));
}
__device__ __forceinline__ void mbarrier_expect_tx(uint64_t& barrier,
uint32_t bytes) {
uint32_t const barrier_addr =
static_cast<uint32_t>(__cvta_generic_to_shared(&barrier));
asm volatile("mbarrier.arrive.expect_tx.shared::cta.b64 _, [%0], %1;\n" ::"r"(
barrier_addr),
"r"(bytes));
}
__device__ __forceinline__ void mbarrier_wait(uint64_t& barrier, int phase) {
uint32_t const barrier_addr =
static_cast<uint32_t>(__cvta_generic_to_shared(&barrier));
asm volatile(
"{\n"
".reg .pred p;\n"
"WAIT:\n"
"mbarrier.try_wait.parity.shared::cta.b64 p, [%0], %1;\n"
"@p bra DONE;\n"
"bra WAIT;\n"
"DONE:\n"
"}\n" ::"r"(barrier_addr),
"r"(phase));
}
__device__ __forceinline__ void mbarrier_arrive(uint64_t& barrier) {
uint32_t const barrier_addr =
static_cast<uint32_t>(__cvta_generic_to_shared(&barrier));
asm volatile(
"{\n"
".reg .b64 state;\n"
"mbarrier.arrive.shared::cta.b64 state, [%0];\n"
"}\n" ::"r"(barrier_addr));
}
__device__ __forceinline__ void fence_mbarrier_init() {
asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");
}
__device__ __forceinline__ void named_barrier_sync(uint32_t num_threads,
uint32_t user_barrier_id) {
asm volatile(
"bar.sync %0, %1;" ::"r"(user_barrier_id + FIRST_USER_NAMED_BARRIER),
"r"(num_threads)
: "memory");
}
__device__ __forceinline__ void tmem_allocate(int num_columns, uint32_t* dst) {
uint32_t const dst_addr =
static_cast<uint32_t>(__cvta_generic_to_shared(dst));
asm volatile(
"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" ::"r"(
dst_addr),
"r"(num_columns));
}
__device__ __forceinline__ void tmem_free(uint32_t tmem_ptr, int num_columns) {
asm volatile(
"tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" ::"r"(tmem_ptr),
"r"(num_columns));
}
__device__ __forceinline__ void tmem_release_allocation_lock() {
asm volatile("tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;");
}
__device__ __forceinline__ void tmem_store_wait() {
asm volatile("tcgen05.wait::st.sync.aligned;" ::: "memory");
}
template <int N, typename T>
__device__ __forceinline__ void tmem_load(uint32_t src_addr, T* dst) {
uint32_t* values = reinterpret_cast<uint32_t*>(dst);
if constexpr (N == 8) {
asm volatile(
"tcgen05.ld.sync.aligned.32x32b.x8.b32"
"{%0, %1, %2, %3, %4, %5, %6, %7}, [%8];\n"
: "=r"(values[0]), "=r"(values[1]), "=r"(values[2]), "=r"(values[3]),
"=r"(values[4]), "=r"(values[5]), "=r"(values[6]), "=r"(values[7])
: "r"(src_addr));
} else {
static_assert(N == 4, "AttnRes TMEM helpers support x4 and x8");
asm volatile(
"tcgen05.ld.sync.aligned.32x32b.x4.b32"
"{%0, %1, %2, %3}, [%4];\n"
: "=r"(values[0]), "=r"(values[1]), "=r"(values[2]), "=r"(values[3])
: "r"(src_addr));
}
}
template <int N, typename T>
__device__ __forceinline__ void tmem_store(uint32_t dst_addr, T* src) {
uint32_t* values = reinterpret_cast<uint32_t*>(src);
if constexpr (N == 8) {
asm volatile(
"tcgen05.st.sync.aligned.32x32b.x8.b32"
"[%8], {%0, %1, %2, %3, %4, %5, %6, %7};\n" ::"r"(values[0]),
"r"(values[1]), "r"(values[2]), "r"(values[3]), "r"(values[4]),
"r"(values[5]), "r"(values[6]), "r"(values[7]), "r"(dst_addr));
} else {
static_assert(N == 4, "AttnRes TMEM helpers support x4 and x8");
asm volatile(
"tcgen05.st.sync.aligned.32x32b.x4.b32"
"[%4], {%0, %1, %2, %3};\n" ::"r"(values[0]),
"r"(values[1]), "r"(values[2]), "r"(values[3]), "r"(dst_addr));
}
}
__device__ __forceinline__ float2 float2_add(const float2& a, const float2& b) {
float2 result;
asm volatile("add.rn.f32x2 %0, %1, %2;\n"
: "=l"(reinterpret_cast<uint64_t&>(result))
: "l"(reinterpret_cast<uint64_t const&>(a)),
"l"(reinterpret_cast<uint64_t const&>(b)));
return result;
}
__device__ __forceinline__ float2 float2_mul(const float2& a, const float2& b) {
float2 result;
asm volatile("mul.f32x2 %0, %1, %2;\n"
: "=l"(reinterpret_cast<uint64_t&>(result))
: "l"(reinterpret_cast<uint64_t const&>(a)),
"l"(reinterpret_cast<uint64_t const&>(b)));
return result;
}
__device__ __forceinline__ float2 float2_fma(const float2& a, const float2& b,
const float2& c) {
float2 result;
asm volatile("fma.rn.f32x2 %0, %1, %2, %3;\n"
: "=l"(reinterpret_cast<uint64_t&>(result))
: "l"(reinterpret_cast<uint64_t const&>(a)),
"l"(reinterpret_cast<uint64_t const&>(b)),
"l"(reinterpret_cast<uint64_t const&>(c)));
return result;
}
template <int NC>
struct FwdSmemPlan {
alignas(16) uint64_t bar_ready[CHUNK_DEPTH];
alignas(16) uint64_t bar_consumed[CHUNK_DEPTH];
alignas(16) uint64_t bar_output_norm_ready;
alignas(16) float2 ws_stats[CONSUMER_WARPS][NC];
uint32_t tmem_base;
};
__device__ __forceinline__ void cp_async_bulk(void* smem_dst,
const void* gmem_src, int bytes,
uint64_t& mbar) {
uint32_t const s = static_cast<uint32_t>(__cvta_generic_to_shared(smem_dst));
uint32_t const m = static_cast<uint32_t>(__cvta_generic_to_shared(&mbar));
asm volatile(
"cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes [%0], "
"[%1], %2, [%3];\n" ::"r"(s),
"l"(gmem_src), "r"(bytes), "r"(m)
: "memory");
}
template <int H, int NC = N_CHUNK_DEFAULT, bool RELEASE_TMEM = false,
bool HAS_DELTA = false, bool HAS_OUTPUT_NORM = false,
bool OUTPUT_NORM_IN_SMEM = false>
__global__ void __launch_bounds__(BLK, 1) attn_res_fwd_online_v2_kernel(
const bf16_t* __restrict__ block_res, bf16_t* __restrict__ layer_res,
const bf16_t* __restrict__ delta, const bf16_t* __restrict__ res_w,
const bf16_t* __restrict__ rms_w, bf16_t* __restrict__ output, int N, int T,
int B, int block_stride_m, int block_stride_r, float rms_eps,
const bf16_t* __restrict__ output_norm_weight, float output_norm_eps) {
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 1000 && __CUDA_ARCH__ < 1100
constexpr float LOG2_E = 1.4426950408889634f;
constexpr int N_CHUNK = NC;
// The two-source specialization only consumes half of the TMEM columns.
constexpr int TMEM_COLS_ALLOC = NC == 2 ? 128 : 256;
constexpr int NUM_BUFS = CHUNK_DEPTH * NC;
constexpr int NHT = H / K_TILE;
constexpr int SLICES_PER_GROUP =
(NHT + CONSUMER_GROUPS - 1) / CONSUMER_GROUPS;
constexpr int VEC = 8;
constexpr int ACC_PER_THREAD = H == 7168 ? 28 : SLICES_PER_GROUP * VEC;
constexpr int TMEM_V_COLS_PER_GROUP = SLICES_PER_GROUP * N_CHUNK * VEC;
constexpr int TMEM_V_COLS_TOTAL = CONSUMER_GROUPS * TMEM_V_COLS_PER_GROUP;
static_assert(TMEM_V_COLS_TOTAL <= TMEM_COLS_ALLOC);
static_assert(H >= 4096 && H <= 8192);
static_assert(H % K_TILE == 0);
const int tid = threadIdx.x;
const int wid = tid >> 5;
const int lane = tid & 31;
const int TB = T * B;
const int num_ctas = gridDim.x;
const int num_chunks = (N + N_CHUNK - 1) / N_CHUNK;
const int comp_wid = wid - 1;
const int comp_tid = tid - 32;
const int group = (comp_wid >= 4) ? 1 : 0;
const int ct_in_group =
(comp_tid >= 0) ? (comp_tid & (CONSUMER_THREADS_PER_GROUP - 1)) : -1;
const int k_local = ct_in_group * VEC;
constexpr size_t V_BYTES = (size_t)NUM_BUFS * H * sizeof(bf16_t);
constexpr size_t DELTA_BYTES =
HAS_DELTA ? (size_t)CHUNK_DEPTH * H * sizeof(bf16_t) : 0;
constexpr size_t OUTPUT_NORM_BYTES =
OUTPUT_NORM_IN_SMEM ? (size_t)H * sizeof(bf16_t) : 0;
extern __shared__ __align__(16) char smem_raw[];
bf16_t* v_bufs = reinterpret_cast<bf16_t*>(smem_raw); // [NUM_BUFS][H]
bf16_t* delta_bufs = reinterpret_cast<bf16_t*>(smem_raw + V_BYTES);
bf16_t* output_norm_buf =
reinterpret_cast<bf16_t*>(smem_raw + V_BYTES + DELTA_BYTES);
FwdSmemPlan<NC>& plan = *reinterpret_cast<FwdSmemPlan<NC>*>(
smem_raw + V_BYTES + DELTA_BYTES + OUTPUT_NORM_BYTES);
auto slot_of = [](long long gci, int n) {
return (int)(gci % CHUNK_DEPTH) * N_CHUNK + n;
};
auto phase_of = [](long long gci) { return (int)((gci / CHUNK_DEPTH) & 1); };
auto buf_ptr = [&](int slot) -> bf16_t* { return v_bufs + slot * H; };
auto delta_buf_ptr = [&](int chunk_slot) -> bf16_t* {
return delta_bufs + chunk_slot * H;
};
if (wid == 0 && elect_one_sync()) {
#pragma unroll
for (int i = 0; i < CHUNK_DEPTH; i++) {
mbarrier_init(plan.bar_ready[i], 1);
mbarrier_init(plan.bar_consumed[i], CONSUMER_WARPS);
}
if constexpr (OUTPUT_NORM_IN_SMEM) {
mbarrier_init(plan.bar_output_norm_ready, 1);
}
fence_mbarrier_init();
}
// gdc wait BEFORE tmem alloc
cudaGridDependencySynchronize();
if (wid == 1) {
tmem_allocate(TMEM_COLS_ALLOC, &plan.tmem_base);
if constexpr (RELEASE_TMEM) {
tmem_release_allocation_lock();
}
}
__syncthreads();
if constexpr (OUTPUT_NORM_IN_SMEM) {
if (wid == 0 && elect_one_sync()) {
mbarrier_expect_tx(plan.bar_output_norm_ready, H * (int)sizeof(bf16_t));
cp_async_bulk(output_norm_buf, output_norm_weight, H * sizeof(bf16_t),
plan.bar_output_norm_ready);
}
}
const uint32_t my_v_tmem =
comp_tid >= 0 ? plan.tmem_base + group * TMEM_V_COLS_PER_GROUP : 0;
float q_cache[ACC_PER_THREAD];
if (comp_tid >= 0) {
#pragma unroll
for (int si = 0; si < SLICES_PER_GROUP; si++) {
if constexpr (H == 7168) {
if (si == SLICES_PER_GROUP - 1) {
int h_base = 6 * K_TILE + group * (K_TILE / 2) + ct_in_group * 4;
#pragma unroll
for (int j = 0; j < 4; j++) {
int h = h_base + j;
q_cache[si * VEC + j] =
__bfloat162float(rms_w[h]) * __bfloat162float(res_w[h]);
}
continue;
}
}
int dt = si * CONSUMER_GROUPS + group;
if (dt >= NHT) continue;
int h_base = dt * K_TILE + k_local;
#pragma unroll
for (int j = 0; j < VEC; j++) {
int h = h_base + j;
q_cache[si * VEC + j] =
__bfloat162float(rms_w[h]) * __bfloat162float(res_w[h]);
}
}
}
if (wid == 0) {
if (elect_one_sync()) {
long long gci = 0;
for (int tb = blockIdx.x; tb < TB; tb += num_ctas) {
const int t = tb / B;
for (int ci = 0; ci < num_chunks; ci++, gci++) {
int ns = ci * N_CHUNK;
int an = min(N_CHUNK, N - ns);
int chunk_slot = (int)(gci % CHUNK_DEPTH);
int pc = phase_of(gci);
mbarrier_wait(plan.bar_consumed[chunk_slot], pc ^ 1);
int transaction_bytes = an * H * (int)sizeof(bf16_t);
if constexpr (HAS_DELTA) {
int prefix_n = N - 1 - ns;
if (prefix_n >= 0 && prefix_n < an) {
transaction_bytes += H * (int)sizeof(bf16_t);
}
}
mbarrier_expect_tx(plan.bar_ready[chunk_slot], transaction_bytes);
#pragma unroll
for (int n = 0; n < N_CHUNK; n++) {
if (n >= an) continue;
int slot = slot_of(gci, n);
const bf16_t* src =
residual_addr(block_res, layer_res, ns + n, N, t,
block_stride_m, block_stride_r, H);
cp_async_bulk(buf_ptr(slot), src, H * sizeof(bf16_t),
plan.bar_ready[chunk_slot]);
}
if constexpr (HAS_DELTA) {
int prefix_n = N - 1 - ns;
if (prefix_n >= 0 && prefix_n < an) {
cp_async_bulk(delta_buf_ptr(chunk_slot),
delta + (long long)tb * H, H * sizeof(bf16_t),
plan.bar_ready[chunk_slot]);
}
}
}
}
}
} else {
float acc32[ACC_PER_THREAD] = {};
float eps_cache;
asm volatile("mov.b32 %0, %1;" : "=f"(eps_cache) : "f"(rms_eps));
long long gci = 0;
for (int tb = blockIdx.x; tb < TB; tb += num_ctas) {
float m_running = -FLT_MAX;
float s_running = 0.f;
#pragma unroll
for (int i = 0; i < ACC_PER_THREAD; i++) {
acc32[i] = 0.f;
}
for (int ci = 0; ci < num_chunks; ci++, gci++) {
int ns = ci * N_CHUNK;
int an = min(N_CHUNK, N - ns);
int chunk_slot = (int)(gci % CHUNK_DEPTH);
int pr = phase_of(gci);
mbarrier_wait(plan.bar_ready[chunk_slot], pr);
float2 sq_local[N_CHUNK] = {};
float2 dot_local[N_CHUNK] = {};
auto pass_A_body = [&](auto AN_TOK) {
constexpr int AN = decltype(AN_TOK)::value;
#pragma unroll
for (int si = 0; si < SLICES_PER_GROUP; si++) {
if constexpr (H == 7168) {
if (si == SLICES_PER_GROUP - 1) {
int h_base =
6 * K_TILE + group * (K_TILE / 2) + ct_in_group * 4;
const float* qv = &q_cache[si * VEC];
#pragma unroll
for (int n = 0; n < AN; n++) {
int slot = slot_of(gci, n);
int2 vp =
*reinterpret_cast<const int2*>(buf_ptr(slot) + h_base);
auto* v2 = reinterpret_cast<__nv_bfloat162*>(&vp);
if constexpr (HAS_DELTA) {
int prefix_n = N - 1 - ns;
if (n == prefix_n) {
const bf16_t* delta_ptr =
delta_buf_ptr(chunk_slot) + h_base;
#pragma unroll
for (int j = 0; j < 2; j++) {
auto delta2 = *reinterpret_cast<const __nv_bfloat162*>(
delta_ptr + 2 * j);
v2[j] = __hadd2(v2[j], delta2);
}
*reinterpret_cast<int2*>(layer_res + (long long)tb * H +
h_base) = vp;
}
}
float2 f[2] = {__bfloat1622float2(v2[0]),
__bfloat1622float2(v2[1])};
tmem_store<4>(my_v_tmem + (si * N_CHUNK + n) * VEC, f);
sq_local[n] = float2_fma(f[0], f[0], sq_local[n]);
sq_local[n] = float2_fma(f[1], f[1], sq_local[n]);
dot_local[n] =
float2_fma(f[0], make_float2(qv[0], qv[1]), dot_local[n]);
dot_local[n] =
float2_fma(f[1], make_float2(qv[2], qv[3]), dot_local[n]);
}
continue;
}
}
int dt = si * CONSUMER_GROUPS + group;
if (dt >= NHT) continue;
int h_base = dt * K_TILE + k_local;
const float* qv = &q_cache[si * VEC];
#pragma unroll
for (int n = 0; n < AN; n++) {
int slot = slot_of(gci, n);
int4 vp = *reinterpret_cast<const int4*>(buf_ptr(slot) + h_base);
auto* v2 = reinterpret_cast<__nv_bfloat162*>(&vp);
if constexpr (HAS_DELTA) {
int prefix_n = N - 1 - ns;
if (n == prefix_n) {
const bf16_t* delta_ptr = delta_buf_ptr(chunk_slot) + h_base;
#pragma unroll
for (int j = 0; j < VEC / 2; j++) {
auto delta2 = *reinterpret_cast<const __nv_bfloat162*>(
delta_ptr + 2 * j);
v2[j] = __hadd2(v2[j], delta2);
}
*reinterpret_cast<int4*>(layer_res + (long long)tb * H +
h_base) = vp;
}
}
float2 f[4] = {
__bfloat1622float2(v2[0]), __bfloat1622float2(v2[1]),
__bfloat1622float2(v2[2]), __bfloat1622float2(v2[3])};
tmem_store<VEC>(my_v_tmem + (si * N_CHUNK + n) * VEC, f);
#pragma unroll
for (int j = 0; j < VEC / 2; j++) {
sq_local[n] = float2_fma(f[j], f[j], sq_local[n]);
dot_local[n] = float2_fma(
f[j], make_float2(qv[2 * j], qv[2 * j + 1]), dot_local[n]);
}
}
}
};
if constexpr (NC == 4) {
switch (an) {
case 4:
pass_A_body(std::integral_constant<int, 4>{});
break;
case 3:
pass_A_body(std::integral_constant<int, 3>{});
break;
case 2:
pass_A_body(std::integral_constant<int, 2>{});
break;
case 1:
pass_A_body(std::integral_constant<int, 1>{});
break;
default:
__builtin_unreachable();
}
} else if constexpr (NC == 3) {
switch (an) {
case 3:
pass_A_body(std::integral_constant<int, 3>{});
break;
case 2:
pass_A_body(std::integral_constant<int, 2>{});
break;
case 1:
pass_A_body(std::integral_constant<int, 1>{});
break;
default:
__builtin_unreachable();
}
} else {
static_assert(NC == 2);
switch (an) {
case 2:
pass_A_body(std::integral_constant<int, 2>{});
break;
case 1:
pass_A_body(std::integral_constant<int, 1>{});
break;
default:
__builtin_unreachable();
}
}
if (lane == 0) {
mbarrier_arrive(plan.bar_consumed[chunk_slot]);
}
tmem_store_wait();
float2 reduce_pair[N_CHUNK];
#pragma unroll
for (int n = 0; n < N_CHUNK; n++) {
reduce_pair[n] = make_float2(sq_local[n].x + sq_local[n].y,
dot_local[n].x + dot_local[n].y);
}
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
#pragma unroll
for (int n = 0; n < N_CHUNK; n++) {
uint64_t packed = reinterpret_cast<uint64_t&>(reduce_pair[n]);
packed = __shfl_xor_sync(0xffffffff, packed, offset);
float2 other = reinterpret_cast<float2&>(packed);
reduce_pair[n] = float2_add(reduce_pair[n], other);
}
}
if (lane == 0) {
#pragma unroll
for (int n = 0; n < N_CHUNK; n++) {
plan.ws_stats[comp_wid][n] = reduce_pair[n];
}
}
named_barrier_sync(CONSUMER_THREADS, 0);
float local_rsig = 0.f;
float local_logit = 0.f;
int stat_n = lane / CONSUMER_WARPS;
int stat_w = lane % CONSUMER_WARPS;
float2 totals = {};
if (stat_n < N_CHUNK) {
totals = plan.ws_stats[stat_w][stat_n];
}
#pragma unroll
for (int offset = CONSUMER_WARPS / 2; offset > 0; offset >>= 1) {
totals.x +=
__shfl_down_sync(0xffffffff, totals.x, offset, CONSUMER_WARPS);
totals.y +=
__shfl_down_sync(0xffffffff, totals.y, offset, CONSUMER_WARPS);
}
if (stat_n < N_CHUNK && stat_w == 0) {
local_rsig = rsqrtf(totals.x / H + eps_cache);
local_logit = totals.y * local_rsig;
}
float logit_n[N_CHUNK];
#pragma unroll
for (int n = 0; n < N_CHUNK; n++) {
logit_n[n] = __shfl_sync(0xffffffff, local_logit, n * CONSUMER_WARPS);
}
float m_chunk = -FLT_MAX;
#pragma unroll
for (int n = 0; n < N_CHUNK; n++) {
if (n < an) m_chunk = fmaxf(m_chunk, logit_n[n]);
}
float m_new = fmaxf(m_running, m_chunk);
float corr = exp2f((m_running - m_new) * LOG2_E);
float w_n[N_CHUNK] = {};
float w_sum = 0.f;
#pragma unroll
for (int n = 0; n < N_CHUNK; n++) {
if (n < an) {
w_n[n] = exp2f((logit_n[n] - m_new) * LOG2_E);
w_sum += w_n[n];
}
}
auto pass_B_body = [&](auto AN_TOK) {
constexpr int AN = decltype(AN_TOK)::value;
#pragma unroll
for (int si = 0; si < SLICES_PER_GROUP; si++) {
if constexpr (H == 7168) {
if (si == SLICES_PER_GROUP - 1) {
float2 corr2 = make_float2(corr, corr);
float2 a[2];
#pragma unroll
for (int j = 0; j < 2; j++) {
float2 old = make_float2(acc32[si * VEC + 2 * j],
acc32[si * VEC + 2 * j + 1]);
a[j] = float2_mul(old, corr2);
}
float2 f_cache[AN][2];
#pragma unroll
for (int n = 0; n < AN; n++) {
tmem_load<4>(my_v_tmem + (si * N_CHUNK + n) * VEC,
f_cache[n]);
}
#pragma unroll
for (int n = 0; n < AN; n++) {
float2 wn = make_float2(w_n[n], w_n[n]);
#pragma unroll
for (int j = 0; j < 2; j++) {
a[j] = float2_fma(wn, f_cache[n][j], a[j]);
}
}
#pragma unroll
for (int j = 0; j < 2; j++) {
acc32[si * VEC + 2 * j] = a[j].x;
acc32[si * VEC + 2 * j + 1] = a[j].y;
}
continue;
}
}
int dt = si * CONSUMER_GROUPS + group;
if (dt >= NHT) continue;
float2 corr2 = make_float2(corr, corr);
float2 a[VEC / 2];
#pragma unroll
for (int j = 0; j < VEC / 2; j++) {
float2 old = make_float2(acc32[si * VEC + 2 * j],
acc32[si * VEC + 2 * j + 1]);
a[j] = float2_mul(old, corr2);
}
float2 f_cache[AN][VEC / 2];
#pragma unroll
for (int n = 0; n < AN; n++) {
tmem_load<VEC>(my_v_tmem + (si * N_CHUNK + n) * VEC, f_cache[n]);
}
#pragma unroll
for (int n = 0; n < AN; n++) {
float2 wn = make_float2(w_n[n], w_n[n]);
#pragma unroll
for (int j = 0; j < VEC / 2; j++) {
a[j] = float2_fma(wn, f_cache[n][j], a[j]);
}
}
#pragma unroll
for (int j = 0; j < VEC / 2; j++) {
acc32[si * VEC + 2 * j] = a[j].x;
acc32[si * VEC + 2 * j + 1] = a[j].y;
}
}
};
if constexpr (NC == 4) {
switch (an) {
case 4:
pass_B_body(std::integral_constant<int, 4>{});
break;
case 3:
pass_B_body(std::integral_constant<int, 3>{});
break;
case 2:
pass_B_body(std::integral_constant<int, 2>{});
break;
case 1:
pass_B_body(std::integral_constant<int, 1>{});
break;
default:
__builtin_unreachable();
}
} else if constexpr (NC == 3) {
switch (an) {
case 3:
pass_B_body(std::integral_constant<int, 3>{});
break;
case 2:
pass_B_body(std::integral_constant<int, 2>{});
break;
case 1:
pass_B_body(std::integral_constant<int, 1>{});
break;
default:
__builtin_unreachable();
}
} else {
static_assert(NC == 2);
switch (an) {
case 2:
pass_B_body(std::integral_constant<int, 2>{});
break;
case 1:
pass_B_body(std::integral_constant<int, 1>{});
break;
default:
__builtin_unreachable();
}
}
s_running = s_running * corr + w_sum;
m_running = m_new;
}
float inv_s = 1.f / s_running;
bf16_t* out_ptr = output + (long long)tb * H;
float2 output_sq_pair = {};
// When output RMSNorm is fused, the softmax denominator cancels:
// (acc / s) * rsqrt(mean((acc / s)^2) + eps)
// = acc * rsqrt(mean(acc^2) + eps * s^2).
#pragma unroll
for (int si = 0; si < SLICES_PER_GROUP; si++) {
if constexpr (H == 7168) {
if (si == SLICES_PER_GROUP - 1) {
int h_base = 6 * K_TILE + group * (K_TILE / 2) + ct_in_group * 4;
uint2 packed;
auto* ov2 = reinterpret_cast<__nv_bfloat162*>(&packed);
float2 inv2 = make_float2(inv_s, inv_s);
#pragma unroll
for (int j = 0; j < 2; j++) {
float2 old = make_float2(acc32[si * VEC + 2 * j],
acc32[si * VEC + 2 * j + 1]);
if constexpr (HAS_OUTPUT_NORM) {
output_sq_pair = float2_fma(old, old, output_sq_pair);
} else {
float2 mixed = float2_mul(old, inv2);
ov2[j] = __float22bfloat162_rn(mixed);
}
}
if constexpr (!HAS_OUTPUT_NORM) {
*reinterpret_cast<uint2*>(out_ptr + h_base) = packed;
}
continue;
}
}
int dt = si * CONSUMER_GROUPS + group;
if (dt >= NHT) continue;
int h_base = dt * K_TILE + k_local;
uint4 packed;
auto* ov2 = reinterpret_cast<__nv_bfloat162*>(&packed);
float2 inv2 = make_float2(inv_s, inv_s);
#pragma unroll
for (int j = 0; j < VEC / 2; j++) {
float2 old =
make_float2(acc32[si * VEC + 2 * j], acc32[si * VEC + 2 * j + 1]);
if constexpr (HAS_OUTPUT_NORM) {
output_sq_pair = float2_fma(old, old, output_sq_pair);
} else {
float2 mixed = float2_mul(old, inv2);
ov2[j] = __float22bfloat162_rn(mixed);
}
}
if constexpr (!HAS_OUTPUT_NORM) {
*reinterpret_cast<uint4*>(out_ptr + h_base) = packed;
}
}
if constexpr (HAS_OUTPUT_NORM) {
if constexpr (OUTPUT_NORM_IN_SMEM) {
// The immutable weight copy is acquired once, at its first use.
if (tb == blockIdx.x) {
mbarrier_wait(plan.bar_output_norm_ready, 0);
}
}
float output_sq = output_sq_pair.x + output_sq_pair.y;
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
output_sq += __shfl_xor_sync(0xffffffff, output_sq, offset);
}
if (lane == 0) {
plan.ws_stats[comp_wid][0] = make_float2(output_sq, 0.f);
}
named_barrier_sync(CONSUMER_THREADS, 0);
float total_sq = lane < CONSUMER_WARPS ? plan.ws_stats[lane][0].x : 0.f;
#pragma unroll
for (int offset = CONSUMER_WARPS / 2; offset > 0; offset >>= 1) {
total_sq +=
__shfl_down_sync(0xffffffff, total_sq, offset, CONSUMER_WARPS);
}
if (lane == 0) {
total_sq =
rsqrtf(total_sq / H + output_norm_eps * s_running * s_running);
}
float output_rsigma = __shfl_sync(0xffffffff, total_sq, 0);
#pragma unroll
for (int si = 0; si < SLICES_PER_GROUP; si++) {
if constexpr (H == 7168) {
if (si == SLICES_PER_GROUP - 1) {
int h_base = 6 * K_TILE + group * (K_TILE / 2) + ct_in_group * 4;
uint2 packed;
auto* values = reinterpret_cast<bf16_t*>(&packed);
#pragma unroll
for (int j = 0; j < 4; j++) {
const bf16_t* weight_ptr =
OUTPUT_NORM_IN_SMEM ? output_norm_buf : output_norm_weight;
float weight = __bfloat162float(weight_ptr[h_base + j]);
values[j] = __float2bfloat16(acc32[si * VEC + j] *
output_rsigma * weight);
}
*reinterpret_cast<uint2*>(out_ptr + h_base) = packed;
continue;
}
}
int dt = si * CONSUMER_GROUPS + group;
if (dt >= NHT) continue;
int h_base = dt * K_TILE + k_local;
uint4 packed;
auto* values = reinterpret_cast<bf16_t*>(&packed);
#pragma unroll
for (int j = 0; j < VEC; j++) {
const bf16_t* weight_ptr =
OUTPUT_NORM_IN_SMEM ? output_norm_buf : output_norm_weight;
float weight = __bfloat162float(weight_ptr[h_base + j]);
values[j] =
__float2bfloat16(acc32[si * VEC + j] * output_rsigma * weight);
}
*reinterpret_cast<uint4*>(out_ptr + h_base) = packed;
}
}
}
}
cudaTriggerProgrammaticLaunchCompletion();
__syncthreads();
if (wid == 1) {
tmem_free(plan.tmem_base, TMEM_COLS_ALLOC);
}
#else
if (threadIdx.x == 0) {
printf("attn_res_fwd_online_v2_kernel requires sm_10x\n");
}
#endif
}
template <int H, int NC = N_CHUNK_DEFAULT, bool RELEASE_TMEM = false,
bool HAS_DELTA = false, bool HAS_OUTPUT_NORM = false,
bool OUTPUT_NORM_IN_SMEM = false>
static void launch_fwd(const bf16_t* block_residual, bf16_t* layer_residual,
const bf16_t* delta, const bf16_t* res_weight,
const bf16_t* rms_weight, bf16_t* output, int N, int T,
int B, float rms_eps, int num_sm, cudaStream_t stream,
const bf16_t* output_norm_weight = nullptr,
float output_norm_eps = 0.f, int block_stride_m = 0,
int block_stride_r = 0) {
constexpr size_t smem_size =
((size_t)CHUNK_DEPTH * (NC + (HAS_DELTA ? 1 : 0)) * H * sizeof(bf16_t) +
(OUTPUT_NORM_IN_SMEM ? (size_t)H * sizeof(bf16_t) : 0) +
sizeof(FwdSmemPlan<NC>) + 15) &
~size_t(15);
auto kernel =
&attn_res_fwd_online_v2_kernel<H, NC, RELEASE_TMEM, HAS_DELTA,
HAS_OUTPUT_NORM, OUTPUT_NORM_IN_SMEM>;
static bool attrs_set = false;
if (!attrs_set) {
if (smem_size > 48 * 1024) {
cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize,
smem_size);
}
attrs_set = true;
}
int grid = RELEASE_TMEM ? num_sm * 2 : num_sm;
cudaLaunchConfig_t config{};
config.gridDim = grid;
config.blockDim = BLK;
config.dynamicSmemBytes = smem_size;
config.stream = stream;
cudaLaunchAttribute attrs[1];
attrs[0].id = cudaLaunchAttributeProgrammaticStreamSerialization;
attrs[0].val.programmaticStreamSerializationAllowed = 1;
config.attrs = attrs;
config.numAttrs = 1;
cudaLaunchKernelEx(&config, kernel, block_residual, layer_residual, delta,
res_weight, rms_weight, output, N, T, B, block_stride_m,
block_stride_r, rms_eps, output_norm_weight,
output_norm_eps);
}
} // namespace fwd_prod_v2
} // namespace sm100
void kimi_k3_attn_res(torch::stable::Tensor& prefix,
torch::stable::Tensor const& delta,
torch::stable::Tensor const& blocks,
torch::stable::Tensor const& norm_weight,
torch::stable::Tensor const& qk_weight,
torch::stable::Tensor const& output_norm_weight,
torch::stable::Tensor& output, int64_t num_blocks,
double eps, double output_norm_eps) {
int const num_tokens = static_cast<int>(prefix.size(0));
int const device = prefix.get_device_index();
torch::stable::accelerator::DeviceGuard const device_guard(device);
cudaDeviceProp const* properties = get_device_prop();
STD_TORCH_CHECK(properties->major == 10,
"Kimi K3 AttnRes requires the SM100 family");
using namespace sm100::fwd_prod_v2;
// Two-source chunks and two resident CTAs are beneficial once setup is
// amortized by the long, full eight-block prefill workload.
if (num_blocks == 8 && num_tokens >= 4096) {
launch_fwd<7168, 2, true, true, true, true>(
static_cast<bf16_t const*>(blocks.data_ptr()),
static_cast<bf16_t*>(prefix.data_ptr()),
static_cast<bf16_t const*>(delta.data_ptr()),
static_cast<bf16_t const*>(qk_weight.data_ptr()),
static_cast<bf16_t const*>(norm_weight.data_ptr()),
static_cast<bf16_t*>(output.data_ptr()),
static_cast<int>(num_blocks) + 1, num_tokens, 1,
static_cast<float>(eps), properties->multiProcessorCount,
get_current_cuda_stream(device),
static_cast<bf16_t const*>(output_norm_weight.data_ptr()),
static_cast<float>(output_norm_eps), static_cast<int>(blocks.stride(0)),
static_cast<int>(blocks.stride(1)));
} else {
launch_fwd<7168, 4, false, true, true, true>(
static_cast<bf16_t const*>(blocks.data_ptr()),
static_cast<bf16_t*>(prefix.data_ptr()),
static_cast<bf16_t const*>(delta.data_ptr()),
static_cast<bf16_t const*>(qk_weight.data_ptr()),
static_cast<bf16_t const*>(norm_weight.data_ptr()),
static_cast<bf16_t*>(output.data_ptr()),
static_cast<int>(num_blocks) + 1, num_tokens, 1,
static_cast<float>(eps), properties->multiProcessorCount,
get_current_cuda_stream(device),
static_cast<bf16_t const*>(output_norm_weight.data_ptr()),
static_cast<float>(output_norm_eps), static_cast<int>(blocks.stride(0)),
static_cast<int>(blocks.stride(1)));
}
cudaError_t const error = cudaGetLastError();
STD_TORCH_CHECK(
error == cudaSuccess,
"Kimi K3 AttnRes kernel launch failed: ", cudaGetErrorString(error));
}
File diff suppressed because it is too large Load Diff
@@ -25,7 +25,6 @@
#include "libtorch_stable/torch_utils.h"
#include <cmath>
#include <cstdint>
#include <tuple>
#include <cuda_fp16.h>
#include <cuda_bf16.h>
@@ -449,8 +448,7 @@ enum ScoringFunc {
SCORING_SIGMOID = 1 // apply sigmoid
};
// Adapted from
// https://github.com/NVIDIA/TensorRT-LLM/blob/v1.3.0rc2/cpp/tensorrt_llm/kernels/noAuxTcKernels.cu
// Efficient sigmoid approximation from TensorRT-LLM
__device__ inline float sigmoid_accurate(float x) {
return 0.5f * tanhf(0.5f * x) + 0.5f;
}
@@ -892,434 +890,6 @@ __global__ void grouped_topk_fused_small_expert_count_kernel(
#endif
}
// Adapted from
// https://github.com/flashinfer-ai/flashinfer/blob/06400d062a2d51564bbe781f6f811d0b75ca593e/include/flashinfer/trtllm/fused_moe/RoutingKernelTopK.cuh
namespace single_group_topk {
namespace detail {
static constexpr int BlockDim = 256;
static constexpr uint32_t FullWarpMask = 0xffffffffU;
static constexpr float InvalidScore = -INFINITY;
// TopK-only tuning: use wider workers and keep these tiers on the block path.
template <int MaxNumExperts, int MaxNumTopExperts>
static constexpr bool UseTunedBlockPath =
MaxNumTopExperts == 16 && (MaxNumExperts == 896 || MaxNumExperts == 1024);
template <typename T, typename BiasT, ScoringFunc SF>
__device__ __forceinline__ void preprocess_score(T input, BiasT correction_bias,
float& unbiased_score,
float& selection_score) {
unbiased_score = 0.0F;
selection_score = InvalidScore;
float const input_float = cuda_cast<float, T>(input);
float const bias = cuda_cast<float, BiasT>(correction_bias);
if (!is_finite(input_float) || !is_finite(bias)) {
return;
}
float const unbiased = apply_scoring<SF>(input_float);
float const biased = unbiased + bias;
if constexpr (SF == SCORING_NONE) {
if (!is_finite(biased)) {
return;
}
}
unbiased_score = unbiased;
selection_score = biased == 0.0F ? 0.0F : biased;
}
template <typename IdxT>
__device__ __forceinline__ void write_outputs(
cg::thread_block_tile<WARP_SIZE> const& warp, float lane_selection_score,
float lane_unbiased, int32_t lane_expert, int32_t lane, int32_t token,
int32_t topk, float* topk_values, IdxT* topk_indices, bool renormalize,
float routed_scaling_factor) {
bool const finite_selection =
lane < topk && lane_selection_score != InvalidScore;
lane_unbiased = finite_selection ? lane_unbiased : 0.0F;
unsigned const finite_mask = __ballot_sync(FullWarpMask, finite_selection);
float const sum = cg::reduce(warp, lane_unbiased, cg::plus<float>{});
if (lane < topk) {
float output = 0.0F;
if (finite_mask == 0) {
if (renormalize) {
output = 1.0F / static_cast<float>(topk);
}
} else if (finite_selection) {
float scale = routed_scaling_factor;
if (renormalize) {
scale /= sum + 1e-20F;
}
output = lane_unbiased * scale;
}
int64_t const output_index = int64_t{token} * topk + lane;
topk_values[output_index] = output;
topk_indices[output_index] = static_cast<IdxT>(lane_expert);
}
}
template <typename T, typename BiasT, typename IdxT, ScoringFunc SF,
int MaxNumExperts, int MaxNumTopExperts>
__global__ void __launch_bounds__(BlockDim)
single_group_topk_block_kernel(T const* scores, float* topk_values,
IdxT* topk_indices, BiasT const* bias,
int64_t num_experts, int64_t topk,
bool renormalize,
float routed_scaling_factor,
bool enable_pdl) {
static constexpr int NumChunks = (MaxNumExperts + WARP_SIZE - 1) / WARP_SIZE;
static constexpr int WorkerValuesPerLane =
UseTunedBlockPath<MaxNumExperts, MaxNumTopExperts> ? 8 : 4;
static constexpr int ExpertsPerWorkerWarp = WorkerValuesPerLane * WARP_SIZE;
using LaneOwnedRange =
reduce_topk::HighExpertLaneOwnedTopKRange<MaxNumExperts,
MaxNumTopExperts>;
static constexpr int NumWorkerWarps =
(MaxNumExperts + ExpertsPerWorkerWarp - 1) / ExpertsPerWorkerWarp;
static constexpr int NumIntermediate = NumWorkerWarps * MaxNumTopExperts;
static constexpr int MergeValuesPerLane =
(NumIntermediate + WARP_SIZE - 1) / WARP_SIZE;
static constexpr bool LaneOwnedResourcesFit =
NumWorkerWarps <= BlockDim / WARP_SIZE && MergeValuesPerLane <= 64;
static constexpr bool UseHierarchicalLaneTopK =
LaneOwnedRange::kEnabled && LaneOwnedResourcesFit;
static_assert(NumChunks <= 64);
static_assert(MaxNumTopExperts <= WARP_SIZE);
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)
if (enable_pdl) {
cudaGridDependencySynchronize();
}
#endif
__shared__ float __attribute((aligned(128))) biased_scores[MaxNumExperts];
__shared__ float __attribute((aligned(128))) unbiased_scores[MaxNumExperts];
int32_t const token = static_cast<int32_t>(blockIdx.x);
int32_t const lane = static_cast<int32_t>(threadIdx.x) % WARP_SIZE;
int32_t const warp_id = static_cast<int32_t>(threadIdx.x) / WARP_SIZE;
int32_t const num_experts_i32 = static_cast<int32_t>(num_experts);
int32_t const topk_i32 = static_cast<int32_t>(topk);
T const* token_scores = scores + int64_t{token} * num_experts;
for (int32_t expert = static_cast<int32_t>(threadIdx.x);
expert < num_experts_i32; expert += BlockDim) {
preprocess_score<T, BiasT, SF>(token_scores[expert], bias[expert],
unbiased_scores[expert],
biased_scores[expert]);
}
__syncthreads();
auto warp = cg::tiled_partition<WARP_SIZE>(cg::this_thread_block());
if constexpr (UseHierarchicalLaneTopK) {
__shared__ float
__attribute((aligned(128))) intermediate_scores[NumIntermediate];
__shared__ int32_t
__attribute((aligned(128))) intermediate_indices[NumIntermediate];
if (warp_id < NumWorkerWarps) {
float local_scores[WorkerValuesPerLane];
int32_t local_indices[WorkerValuesPerLane];
#pragma unroll
for (int index = 0; index < WorkerValuesPerLane; ++index) {
int32_t const expert =
warp_id * ExpertsPerWorkerWarp + index * WARP_SIZE + lane;
local_scores[index] =
expert < num_experts_i32 ? biased_scores[expert] : InvalidScore;
local_indices[index] = expert;
}
float lane_score;
int32_t lane_expert;
reduce_topk::reduceTopKForLane<MaxNumTopExperts>(
warp, lane_score, lane_expert, local_scores, local_indices,
InvalidScore, lane);
if (lane < MaxNumTopExperts) {
int32_t const intermediate = warp_id * MaxNumTopExperts + lane;
bool const active = lane < topk_i32;
intermediate_scores[intermediate] = active ? lane_score : InvalidScore;
intermediate_indices[intermediate] =
active ? lane_expert : MaxNumExperts;
}
}
__syncthreads();
if (warp_id != 0) {
return;
}
float merge_scores[MergeValuesPerLane];
int32_t merge_indices[MergeValuesPerLane];
#pragma unroll
for (int index = 0; index < MergeValuesPerLane; ++index) {
int32_t const intermediate = index * WARP_SIZE + lane;
bool const active = intermediate < NumIntermediate;
merge_scores[index] =
active ? intermediate_scores[intermediate] : InvalidScore;
merge_indices[index] =
active ? intermediate_indices[intermediate] : MaxNumExperts;
}
float lane_score;
int32_t lane_expert;
reduce_topk::reduceTopKForLane<MaxNumTopExperts>(
warp, lane_score, lane_expert, merge_scores, merge_indices,
InvalidScore, lane);
float const lane_unbiased =
lane < topk_i32 && lane_expert >= 0 && lane_expert < num_experts_i32
? unbiased_scores[lane_expert]
: 0.0F;
write_outputs(warp, lane_score, lane_unbiased, lane_expert, lane, token,
topk_i32, topk_values, topk_indices, renormalize,
routed_scaling_factor);
} else {
if (warp_id != 0) {
return;
}
float local_scores[NumChunks];
int32_t local_indices[NumChunks];
#pragma unroll
for (int index = 0; index < NumChunks; ++index) {
int32_t const expert = index * WARP_SIZE + lane;
local_scores[index] =
expert < num_experts_i32 ? biased_scores[expert] : InvalidScore;
local_indices[index] = expert;
}
float top_scores[MaxNumTopExperts];
int32_t top_experts[MaxNumTopExperts];
reduce_topk::reduceTopK(warp, top_scores, top_experts, local_scores,
local_indices, InvalidScore, topk_i32);
float const lane_score = lane < topk_i32 ? top_scores[lane] : InvalidScore;
int32_t const lane_expert = lane < topk_i32 ? top_experts[lane] : -1;
float const lane_unbiased =
lane < topk_i32 && lane_expert >= 0 && lane_expert < num_experts_i32
? unbiased_scores[lane_expert]
: 0.0F;
write_outputs(warp, lane_score, lane_unbiased, lane_expert, lane, token,
topk_i32, topk_values, topk_indices, renormalize,
routed_scaling_factor);
}
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)
if (enable_pdl) {
cudaTriggerProgrammaticLaunchCompletion();
}
#endif
}
template <int MaxNumExperts>
struct WarpTopKLaunchConfig {
static constexpr int DefaultBlockDim =
MaxNumExperts <= 1024 ? MaxNumExperts : 1024;
static constexpr int BlockDim = DefaultBlockDim > 256 ? 256 : DefaultBlockDim;
static constexpr int NumWarps = BlockDim / WARP_SIZE;
static constexpr int MaxBlockScale =
(DefaultBlockDim + BlockDim - 1) / BlockDim;
static constexpr int MaxBlocks = 1024 * MaxBlockScale;
static_assert(BlockDim % WARP_SIZE == 0);
static uint32_t grid_dim(int64_t num_tokens) {
int64_t const token_blocks = (num_tokens + NumWarps - 1) / NumWarps;
int64_t const selected =
token_blocks < MaxBlocks ? token_blocks : MaxBlocks;
return static_cast<uint32_t>(selected > 0 ? selected : 1);
}
};
template <typename T, typename BiasT, typename IdxT, ScoringFunc SF,
int MaxNumExperts, int MaxNumTopExperts>
__global__ void __launch_bounds__(WarpTopKLaunchConfig<MaxNumExperts>::BlockDim)
single_group_topk_warp_kernel(T const* scores, float* topk_values,
IdxT* topk_indices, BiasT const* bias,
int64_t num_tokens, int64_t num_experts,
int64_t topk, bool renormalize,
float routed_scaling_factor,
bool enable_pdl) {
static constexpr int NumChunks = (MaxNumExperts + WARP_SIZE - 1) / WARP_SIZE;
static constexpr int WarpBlockDim =
WarpTopKLaunchConfig<MaxNumExperts>::BlockDim;
using LaneOwnedRange =
reduce_topk::HighExpertLaneOwnedTopKRange<MaxNumExperts,
MaxNumTopExperts>;
static_assert(NumChunks <= 64);
static_assert(MaxNumTopExperts <= WARP_SIZE);
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)
if (enable_pdl) {
cudaGridDependencySynchronize();
}
#endif
int32_t const lane = static_cast<int32_t>(threadIdx.x) % WARP_SIZE;
int32_t const warp_id = static_cast<int32_t>(threadIdx.x) / WARP_SIZE;
int32_t const global_warp =
static_cast<int32_t>(blockIdx.x) * WarpBlockDim / WARP_SIZE + warp_id;
int32_t const global_warp_stride =
static_cast<int32_t>(gridDim.x) * WarpBlockDim / WARP_SIZE;
int32_t const num_experts_i32 = static_cast<int32_t>(num_experts);
int32_t const topk_i32 = static_cast<int32_t>(topk);
auto warp = cg::tiled_partition<WARP_SIZE>(cg::this_thread_block());
for (int32_t token = global_warp; token < num_tokens;
token += global_warp_stride) {
T const* token_scores = scores + int64_t{token} * num_experts;
float local_scores[NumChunks];
int32_t local_indices[NumChunks];
#pragma unroll
for (int index = 0; index < NumChunks; ++index) {
int32_t const expert = index * WARP_SIZE + lane;
float unbiased;
float selection;
if (expert < num_experts_i32) {
preprocess_score<T, BiasT, SF>(token_scores[expert], bias[expert],
unbiased, selection);
} else {
selection = InvalidScore;
}
local_scores[index] = selection;
local_indices[index] = expert;
}
float lane_score;
int32_t lane_expert;
if constexpr (LaneOwnedRange::kEnabled) {
reduce_topk::reduceTopKForLane<MaxNumTopExperts>(
warp, lane_score, lane_expert, local_scores, local_indices,
InvalidScore, lane);
} else {
float top_scores[MaxNumTopExperts];
int32_t top_experts[MaxNumTopExperts];
reduce_topk::reduceTopK(warp, top_scores, top_experts, local_scores,
local_indices, InvalidScore, topk_i32);
lane_score = lane < topk_i32 ? top_scores[lane] : InvalidScore;
lane_expert = lane < topk_i32 ? top_experts[lane] : -1;
}
float lane_unbiased = 0.0F;
if (lane < topk_i32 && lane_expert >= 0 && lane_expert < num_experts_i32) {
lane_unbiased = lane_score - cuda_cast<float, BiasT>(bias[lane_expert]);
}
write_outputs(warp, lane_score, lane_unbiased, lane_expert, lane, token,
topk_i32, topk_values, topk_indices, renormalize,
routed_scaling_factor);
}
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)
if (enable_pdl) {
cudaTriggerProgrammaticLaunchCompletion();
}
#endif
}
template <int Experts, int TopK>
struct Tier {
static constexpr int kExperts = Experts;
static constexpr int kTopK = TopK;
};
template <typename... Tiers>
struct TierList {};
using SigmoidBiasTiers =
TierList<Tier<128, 8>, Tier<256, 8>, Tier<384, 8>, Tier<512, 8>,
Tier<512, 22>, Tier<768, 16>, Tier<896, 16>, Tier<1024, 16>>;
using PrecomputedSoftmaxBiasTiers =
TierList<Tier<128, 4>, Tier<128, 8>, Tier<160, 8>, Tier<256, 8>,
Tier<256, 16>, Tier<512, 8>, Tier<512, 16>, Tier<512, 22>,
Tier<512, 32>, Tier<576, 8>, Tier<768, 16>, Tier<896, 16>,
Tier<1024, 16>>;
template <typename T, typename BiasT, typename IdxT, ScoringFunc SF,
int MaxNumExperts, int MaxNumTopExperts>
void launch(T* scores, float* topk_values, IdxT* topk_indices,
BiasT const* bias, int64_t num_tokens, int64_t num_experts,
int64_t topk, bool renormalize, double routed_scaling_factor,
bool enable_pdl, cudaLaunchConfig_t& config) {
config.dynamicSmemBytes = 0;
bool const use_block_kernel =
UseTunedBlockPath<MaxNumExperts, MaxNumTopExperts> ||
MaxNumExperts > 1024 || num_experts >= 1024 ||
(num_experts >= 256 && num_tokens <= 1024);
if (use_block_kernel) {
config.gridDim = static_cast<uint32_t>(num_tokens);
config.blockDim = BlockDim;
cudaLaunchKernelEx(
&config,
&single_group_topk_block_kernel<T, BiasT, IdxT, SF, MaxNumExperts,
MaxNumTopExperts>,
scores, topk_values, topk_indices, bias, num_experts, topk, renormalize,
static_cast<float>(routed_scaling_factor), enable_pdl);
} else {
using WarpConfig = WarpTopKLaunchConfig<MaxNumExperts>;
config.gridDim = WarpConfig::grid_dim(num_tokens);
config.blockDim = WarpConfig::BlockDim;
cudaLaunchKernelEx(
&config,
&single_group_topk_warp_kernel<T, BiasT, IdxT, SF, MaxNumExperts,
MaxNumTopExperts>,
scores, topk_values, topk_indices, bias, num_tokens, num_experts, topk,
renormalize, static_cast<float>(routed_scaling_factor), enable_pdl);
}
}
template <typename T, typename BiasT, typename IdxT, ScoringFunc SF>
bool dispatch(TierList<>*, T*, float*, IdxT*, BiasT const*, int64_t, int64_t,
int64_t, bool, double, bool, cudaLaunchConfig_t&) {
return false;
}
template <typename T, typename BiasT, typename IdxT, ScoringFunc SF,
typename First, typename... Rest>
bool dispatch(TierList<First, Rest...>*, T* scores, float* topk_values,
IdxT* topk_indices, BiasT const* bias, int64_t num_tokens,
int64_t num_experts, int64_t topk, bool renormalize,
double routed_scaling_factor, bool enable_pdl,
cudaLaunchConfig_t& config) {
if (num_experts <= First::kExperts && topk <= First::kTopK) {
launch<T, BiasT, IdxT, SF, First::kExperts, First::kTopK>(
scores, topk_values, topk_indices, bias, num_tokens, num_experts, topk,
renormalize, routed_scaling_factor, enable_pdl, config);
return true;
}
return dispatch<T, BiasT, IdxT, SF>(
static_cast<TierList<Rest...>*>(nullptr), scores, topk_values,
topk_indices, bias, num_tokens, num_experts, topk, renormalize,
routed_scaling_factor, enable_pdl, config);
}
} // namespace detail
template <typename T, typename BiasT, typename IdxT, ScoringFunc SF>
bool invoke(T* scores, float* topk_values, IdxT* topk_indices,
BiasT const* bias, int64_t num_tokens, int64_t num_experts,
int64_t topk, bool renormalize, double routed_scaling_factor,
bool enable_pdl, cudaLaunchConfig_t& config) {
static_assert(SF == SCORING_NONE || SF == SCORING_SIGMOID);
if constexpr (SF == SCORING_SIGMOID) {
return detail::dispatch<T, BiasT, IdxT, SF>(
static_cast<detail::SigmoidBiasTiers*>(nullptr), scores, topk_values,
topk_indices, bias, num_tokens, num_experts, topk, renormalize,
routed_scaling_factor, enable_pdl, config);
} else {
return detail::dispatch<T, BiasT, IdxT, SF>(
static_cast<detail::PrecomputedSoftmaxBiasTiers*>(nullptr), scores,
topk_values, topk_indices, bias, num_tokens, num_experts, topk,
renormalize, routed_scaling_factor, enable_pdl, config);
}
}
} // namespace single_group_topk
template <typename T, typename BiasT, typename IdxT, ScoringFunc SF>
void invokeNoAuxTc(T* scores, float* topk_values, IdxT* topk_indices,
BiasT const* bias, int64_t const num_tokens,
@@ -1335,12 +905,6 @@ void invokeNoAuxTc(T* scores, float* topk_values, IdxT* topk_indices,
attrs[0].val.programmaticStreamSerializationAllowed = enable_pdl;
config.numAttrs = 1;
config.attrs = attrs;
if (n_group == 1 && topk_group == 1 &&
single_group_topk::invoke<T, BiasT, IdxT, SF>(
scores, topk_values, topk_indices, bias, num_tokens, num_experts,
topk, renormalize, routed_scaling_factor, enable_pdl, config)) {
return;
}
// Check if we can use the optimized
// grouped_topk_fused_small_expert_count_kernel
+120 -231
View File
@@ -1,7 +1,6 @@
/*
* Adapted from
* https://github.com/NVIDIA/TensorRT-LLM/blob/v1.3.0rc2/cpp/tensorrt_llm/kernels/moeTopKFuncs.cuh
* https://github.com/flashinfer-ai/flashinfer/blob/06400d062a2d51564bbe781f6f811d0b75ca593e/include/flashinfer/trtllm/fused_moe/RoutingKernelTopK.cuh
* Copyright (c) 2026, The vLLM team.
* SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION. All rights
* reserved. SPDX-License-Identifier: Apache-2.0
@@ -24,9 +23,6 @@
#include <cooperative_groups/reduce.h>
#include <cub/cub.cuh>
#include <cstdint>
#include <type_traits>
namespace vllm {
namespace moe {
namespace reduce_topk {
@@ -42,10 +38,11 @@ struct TopKRedType {
"Top K reduction only implemented for int, float, float16 and bfloat16");
using TypeCmp = std::conditional_t<sizeof(T) == 4, uint64_t, uint32_t>;
using IdxT = std::conditional_t<sizeof(T) == 4, int32_t, int16_t>;
static constexpr int kMoveBits = (sizeof(T) == 4) ? 32 : 16;
static constexpr int kMaxIdx = 65535;
TypeCmp compVal;
TypeCmp compValIdx;
static __host__ __device__ inline TypeCmp makeCmpVal(T val, int32_t idx = 0) {
auto valueBits = cub::Traits<T>::TwiddleIn(
@@ -72,175 +69,69 @@ struct TopKRedType {
__host__ __device__ TopKRedType() = default;
__host__ __device__ TopKRedType(T val, int32_t idx)
: compVal(makeCmpVal(val, idx)) {}
: compValIdx(makeCmpVal(val, idx)) {}
__host__ __device__ operator TypeCmp() const noexcept { return compVal; }
__host__ __device__ operator TypeCmp() const noexcept { return compValIdx; }
__device__ inline TypeCmp reduce(
cg::thread_block_tile<kWARP_SIZE> const& warp) {
#ifdef __CUDA_ARCH__
static constexpr bool kHAS_FAST_REDUX = (__CUDA_ARCH__ / 100) >= 10;
#else
static constexpr bool kHAS_FAST_REDUX = false;
#endif
if constexpr (!kHAS_FAST_REDUX) {
return cg::reduce(warp, compVal, cg::greater<TypeCmp>{});
} else if constexpr (sizeof(TypeCmp) == 8) {
uint32_t hi = static_cast<uint32_t>(compVal >> 32);
uint32_t lo = static_cast<uint32_t>(compVal & 0xffffffffu);
uint32_t maxHi;
asm volatile("redux.sync.max.u32 %0, %1, 0xffffffff;\n"
: "=r"(maxHi)
: "r"(hi));
uint32_t loContrib = hi == maxHi ? lo : 0u;
uint32_t maxLo;
asm volatile("redux.sync.max.u32 %0, %1, 0xffffffff;\n"
: "=r"(maxLo)
: "r"(loContrib));
return (static_cast<TypeCmp>(maxHi) << 32) | static_cast<TypeCmp>(maxLo);
} else {
TypeCmp result;
asm volatile("redux.sync.max.u32 %0, %1, 0xffffffff;\n"
: "=r"(result)
: "r"(compVal));
return result;
}
return cg::reduce(warp, compValIdx, cg::greater<TypeCmp>{});
}
};
template <int N>
struct IsPowerOf2 {
static constexpr bool value = N > 0 && (N & (N - 1)) == 0;
////////////////////////////////////////////////////////////////////////////////////////////////////
template <int K_, bool Enable_>
struct TopKIdx {
// by default, empty
};
template <int N>
struct NextPow2 {
private:
static constexpr unsigned u = static_cast<unsigned>(N - 1);
static constexpr unsigned s1 = u | (u >> 1);
static constexpr unsigned s2 = s1 | (s1 >> 2);
static constexpr unsigned s3 = s2 | (s2 >> 4);
static constexpr unsigned s4 = s3 | (s3 >> 8);
static constexpr unsigned s5 = s4 | (s4 >> 16);
public:
static constexpr int value = N <= 1 ? 1 : static_cast<int>(s5 + 1);
template <int K_>
struct TopKIdx<K_, true> {
static constexpr int K = K_;
int32_t val[K];
};
template <int A, int B, int Size, typename T>
__device__ __forceinline__ void topkCompareSwap(T* a) {
if constexpr (A < Size && B < Size) {
if (a[A] < a[B]) {
T tmp = a[A];
a[A] = a[B];
a[B] = tmp;
}
} else {
(void)a;
}
}
////////////////////////////////////////////////////////////////////////////////////////////////////
template <int I, int End, int Step, int PairStride, int Size, typename T>
__device__ __forceinline__ void topkMergePairs(T* a) {
if constexpr (I + Step < End) {
topkCompareSwap<I, I + Step, Size, T>(a);
topkMergePairs<I + PairStride, End, Step, PairStride, Size, T>(a);
} else {
(void)a;
#define TOPK_SWAP(I, J) \
{ \
auto pairMin = min(topK[I].compValIdx, topK[J].compValIdx); \
auto pairMax = max(topK[I].compValIdx, topK[J].compValIdx); \
topK[I].compValIdx = pairMax; \
topK[J].compValIdx = pairMin; \
}
}
template <int Lo, int N, int R, int Size, typename T>
__device__ __forceinline__ void topkOEM(T* a) {
constexpr int M = R * 2;
if constexpr (M < N) {
topkOEM<Lo, N, M, Size, T>(a);
topkOEM<Lo + R, N - R, M, Size, T>(a);
topkMergePairs<Lo + R, Lo + N, R, M, Size, T>(a);
} else if constexpr (R < N) {
topkCompareSwap<Lo, Lo + R, Size, T>(a);
} else {
(void)a;
}
}
template <int Lo, int N, int Size, typename T>
__device__ __forceinline__ void topkSortBatcher(T* a) {
if constexpr (N > 1) {
constexpr int Half = N / 2;
topkSortBatcher<Lo, Half, Size, T>(a);
topkSortBatcher<Lo + Half, N - Half, Size, T>(a);
topkOEM<Lo, N, 1, Size, T>(a);
} else {
(void)a;
}
}
template <int N, typename RedType>
struct Sort {
static_assert(N > 0 && N <= 64, "Sort only supports N in range [1, 64]");
static __device__ void run(RedType* topK) {
if constexpr (IsPowerOf2<N>::value) {
#pragma unroll
for (int k = 2; k <= N; k *= 2) {
#pragma unroll
for (int j = k / 2; j > 0; j /= 2) {
#pragma unroll
for (int i = 0; i < N; ++i) {
int ixj = i ^ j;
if (ixj > i) {
if ((i & k) == 0) {
if (topK[i].compVal < topK[ixj].compVal) {
auto tmp = topK[i].compVal;
topK[i].compVal = topK[ixj].compVal;
topK[ixj].compVal = tmp;
}
} else {
if (topK[i].compVal > topK[ixj].compVal) {
auto tmp = topK[i].compVal;
topK[i].compVal = topK[ixj].compVal;
topK[ixj].compVal = tmp;
}
}
}
}
}
}
} else {
constexpr int P = NextPow2<N>::value;
topkSortBatcher<0, P, N, RedType>(topK);
}
}
};
struct Sort;
template <typename RedType>
struct Sort<1, RedType> {
static __device__ void run(RedType*) {}
static __device__ void run(RedType* topK) {}
};
template <typename RedType>
struct Sort<2, RedType> {
static __device__ void run(RedType* topK) { topkCompareSwap<0, 1, 2>(topK); }
static __device__ void run(RedType* topK) { TOPK_SWAP(0, 1); }
};
template <typename RedType>
struct Sort<3, RedType> {
static __device__ void run(RedType* topK) {
topkCompareSwap<0, 1, 3>(topK);
topkCompareSwap<1, 2, 3>(topK);
topkCompareSwap<0, 1, 3>(topK);
TOPK_SWAP(0, 1);
TOPK_SWAP(1, 2);
TOPK_SWAP(0, 1);
}
};
template <typename RedType>
struct Sort<4, RedType> {
static __device__ void run(RedType* topK) {
topkCompareSwap<0, 2, 4>(topK);
topkCompareSwap<1, 3, 4>(topK);
topkCompareSwap<0, 1, 4>(topK);
topkCompareSwap<2, 3, 4>(topK);
topkCompareSwap<1, 2, 4>(topK);
TOPK_SWAP(0, 2);
TOPK_SWAP(1, 3);
TOPK_SWAP(0, 1);
TOPK_SWAP(2, 3);
TOPK_SWAP(1, 2);
}
};
@@ -256,112 +147,110 @@ __forceinline__ __device__ void reduceTopK(
typename RedType::TypeCmp packedMax{};
#pragma unroll
for (int kk = 0; kk < actualK; ++kk) {
topK = kk > 0 && packedMax == topK.compVal ? RedType{minValue, idx} : topK;
topK =
kk > 0 && packedMax == topK.compValIdx ? RedType{minValue, idx} : topK;
// get the next largest value
packedMax = topK.reduce(warp);
RedType::unpack(out[kk], outIdx[kk], packedMax);
}
};
template <int K, typename Type, int N, bool IsSorted = false>
__device__ void reduceTopKFunc(cg::thread_block_tile<kWARP_SIZE> const& warp,
Type (&out)[K], int32_t (&outIdx)[K],
Type (&value)[N], int32_t (&idx)[N],
Type minValue, int actualK = K) {
static_assert(K > 0, "Top K must have K > 0");
static_assert(K < kWARP_SIZE, "Top K must have K < kWARP_SIZE");
static_assert(N > 0, "Top K must have N > 0");
static_assert(N < 5,
"Only support candidates number less than or equal to 128");
using RedType = TopKRedType<Type>;
RedType topK[N];
#pragma unroll
for (int nn = 0; nn < N; ++nn) {
topK[nn] = RedType{value[nn], idx[nn]};
}
if constexpr (!IsSorted) {
Sort<N, RedType>::run(topK);
}
typename RedType::TypeCmp packedMax{};
#pragma unroll
for (int kk = 0; kk < actualK; ++kk) {
bool update = kk > 0 && packedMax == topK[0].compValIdx;
#pragma unroll
for (int nn = 0; nn < N; ++nn) {
topK[nn] = update && nn == N - 1 ? RedType{minValue, idx[nn]}
: update ? topK[nn + 1]
: topK[nn];
}
// get the next largest value
packedMax = topK[0].reduce(warp);
RedType::unpack(out[kk], outIdx[kk], packedMax);
}
};
template <int K, typename Type, int N>
__forceinline__ __device__ void reduceTopK(
cg::thread_block_tile<kWARP_SIZE> const& warp, Type (&out)[K],
int32_t (&outIdx)[K], Type (&value)[N], int32_t (&idx)[N],
Type const minValue, int actualK = K) {
static_assert(K > 0, "Top K must have K > 0");
static_assert(K <= kWARP_SIZE, "Top K must have K <= kWARP_SIZE");
static_assert(K < kWARP_SIZE, "Top K must have K < kWARP_SIZE");
static_assert(N > 0, "Top K must have N > 0");
static_assert(N <= 64,
"Only support candidates number less than or equal to "
"64*32=2048");
static_assert(
N <= 16,
"Only support candidates number less than or equal to 16*32=512");
static_assert(N <= 4 || N % 4 == 0,
"Only support candidates number is a multiple of 4*32=128 or "
"less than or equal to 4");
using RedType = TopKRedType<Type>;
RedType topK[N];
#pragma unroll
for (int nn = 0; nn < N; ++nn) {
topK[nn] = RedType{value[nn], idx[nn]};
}
Sort<N, RedType>::run(topK);
typename RedType::TypeCmp packedMax{};
for (int kk = 0; kk < actualK; ++kk) {
bool update = kk > 0 && packedMax == topK[0].compVal;
#pragma unroll
for (int nn = 0; nn < N; ++nn) {
topK[nn] = update && nn == N - 1 ? RedType{minValue, idx[nn]}
: update ? topK[nn + 1]
: topK[nn];
}
packedMax = topK[0].reduce(warp);
RedType::unpack(out[kk], outIdx[kk], packedMax);
}
};
template <int NumExperts, int NumTopExperts, int MinExperts, int MaxExperts,
int MinTopExperts, int MaxTopExperts>
struct LaneOwnedTopKRange {
static_assert(MinExperts > 0 && MinExperts <= MaxExperts);
static_assert(MinTopExperts > 0 && MinTopExperts <= MaxTopExperts);
static constexpr bool kEnabled =
NumExperts >= MinExperts && NumExperts <= MaxExperts &&
NumTopExperts >= MinTopExperts && NumTopExperts <= MaxTopExperts;
};
static constexpr int kHIGH_EXPERT_LANE_OWNED_TOPK_MIN_EXPERTS = 512;
static constexpr int kHIGH_EXPERT_LANE_OWNED_TOPK_MAX_EXPERTS = 1024;
static constexpr int kHIGH_EXPERT_LANE_OWNED_TOPK_MIN_TOP_EXPERTS = 9;
static constexpr int kHIGH_EXPERT_LANE_OWNED_TOPK_MAX_TOP_EXPERTS = 16;
template <int NumExperts, int NumTopExperts>
using HighExpertLaneOwnedTopKRange =
LaneOwnedTopKRange<NumExperts, NumTopExperts,
kHIGH_EXPERT_LANE_OWNED_TOPK_MIN_EXPERTS,
kHIGH_EXPERT_LANE_OWNED_TOPK_MAX_EXPERTS,
kHIGH_EXPERT_LANE_OWNED_TOPK_MIN_TOP_EXPERTS,
kHIGH_EXPERT_LANE_OWNED_TOPK_MAX_TOP_EXPERTS>;
template <int K, typename Type, int N>
__forceinline__ __device__ void reduceTopKForLane(
cg::thread_block_tile<kWARP_SIZE> const& warp, Type& out, int32_t& outIdx,
Type (&value)[N], int32_t (&idx)[N], Type const minValue, int32_t laneIdx) {
static_assert(K > 0, "Top K must have K > 0");
static_assert(K <= kWARP_SIZE, "Top K must have K <= kWARP_SIZE");
static_assert(N > 0, "Top K must have N > 0");
static_assert(N <= 64,
"Only support candidates number less than or equal to "
"64*32=2048");
using RedType = TopKRedType<Type>;
RedType topK[N];
#pragma unroll
for (int nn = 0; nn < N; ++nn) {
topK[nn] = RedType{value[nn], idx[nn]};
}
Sort<N, RedType>::run(topK);
typename RedType::TypeCmp packedMax{};
typename RedType::TypeCmp lanePacked{};
#pragma unroll
for (int kk = 0; kk < K; ++kk) {
bool update = kk > 0 && packedMax == topK[0].compVal;
#pragma unroll
for (int nn = 0; nn < N; ++nn) {
topK[nn] = update && nn == N - 1 ? RedType{minValue, idx[nn]}
: update ? topK[nn + 1]
: topK[nn];
}
packedMax = topK[0].reduce(warp);
if (laneIdx == kk) {
lanePacked = packedMax;
}
}
if (laneIdx < K) {
RedType::unpack(out, outIdx, lanePacked);
if constexpr (N <= 4) {
reduceTopKFunc<K, Type, N>(warp, out, outIdx, value, idx, minValue,
actualK);
} else {
out = minValue;
outIdx = -1;
constexpr int numLoops = N / 4;
constexpr int numResults = (numLoops * K - 1) / kWARP_SIZE + 1;
Type topKBufferValue[numResults];
int32_t topKBufferIdx[numResults];
int32_t laneIdx = threadIdx.x % kWARP_SIZE;
for (int ii = 0; ii < numResults; ++ii) {
topKBufferValue[ii] = minValue;
topKBufferIdx[ii] = ii * kWARP_SIZE - 1;
}
for (int loop = 0; loop < numLoops; ++loop) {
int start = loop * 4;
Type topKValue[K];
int32_t topKIdx[K];
Type inValue[4];
int32_t inIdx[4];
for (int i = 0; i < 4; ++i) {
inValue[i] = value[start + i];
inIdx[i] = idx[start + i];
}
reduceTopKFunc<K, Type, 4>(warp, topKValue, topKIdx, inValue, inIdx,
minValue, actualK);
int inOffset = laneIdx % K;
if (laneIdx >= loop * K && laneIdx < (loop + 1) * K) {
topKBufferValue[0] = topKValue[inOffset];
topKBufferIdx[0] = topKIdx[inOffset];
}
if (loop == numLoops - 1 && (laneIdx < (numLoops * K - kWARP_SIZE))) {
topKBufferValue[1] = topKValue[inOffset];
topKBufferIdx[1] = topKIdx[inOffset];
}
}
reduceTopKFunc<K, Type, numResults>(warp, out, outIdx, topKBufferValue,
topKBufferIdx, minValue, actualK);
}
}
};
#undef TOPK_SWAP
} // namespace reduce_topk
} // namespace moe
@@ -1086,4 +1086,4 @@ void moe_lora_align_block_size(
has_expert_map);
}
});
}
}
-106
View File
@@ -276,61 +276,6 @@ void fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_bf16_insert(
torch::stable::Tensor const& cos_sin_cache, double eps,
int64_t cache_block_size);
void fused_kimi_k3_mla_key_concat_kv_cache_insert(
torch::stable::Tensor& q, torch::stable::Tensor const& k_nope,
torch::stable::Tensor const& k_pe, torch::stable::Tensor const& kv_c_normed,
torch::stable::Tensor& k_out, torch::stable::Tensor& k_cache,
torch::stable::Tensor const& slot_mapping, int64_t cache_block_size,
std::optional<torch::stable::Tensor> position_ids,
std::optional<torch::stable::Tensor> cos_sin_cache);
void fused_kimi_k3_mla_key_concat_ds_mla_insert(
torch::stable::Tensor& q, torch::stable::Tensor const& k_nope,
torch::stable::Tensor const& k_pe, torch::stable::Tensor const& kv_c_normed,
torch::stable::Tensor& k_out, torch::stable::Tensor& k_cache,
torch::stable::Tensor const& slot_mapping, int64_t cache_block_size,
std::optional<torch::stable::Tensor> position_ids,
std::optional<torch::stable::Tensor> cos_sin_cache);
void fused_kimi_k3_mla_qkv_quant_kv_cache_fp8_insert(
torch::stable::Tensor const& q, torch::stable::Tensor const& k_nope,
torch::stable::Tensor const& k_pe, torch::stable::Tensor const& kv_c_normed,
torch::stable::Tensor const& v, torch::stable::Tensor& q_fp8,
torch::stable::Tensor& k_fp8, torch::stable::Tensor& v_fp8,
torch::stable::Tensor& k_cache, torch::stable::Tensor const& slot_mapping,
torch::stable::Tensor const& q_scale_inv,
torch::stable::Tensor const& k_scale_inv,
torch::stable::Tensor const& v_scale_inv,
torch::stable::Tensor const& cache_scale_inv, int64_t cache_block_size,
std::optional<torch::stable::Tensor> position_ids,
std::optional<torch::stable::Tensor> cos_sin_cache);
void fused_kimi_k3_mla_decode_q_concat_kv_cache_insert(
torch::stable::Tensor const& ql_nope, torch::stable::Tensor const& q_pe,
torch::stable::Tensor const& kv_c_normed, torch::stable::Tensor const& k_pe,
torch::stable::Tensor& mqa_q, torch::stable::Tensor& k_cache,
torch::stable::Tensor const& slot_mapping, int64_t cache_block_size,
std::optional<torch::stable::Tensor> position_ids,
std::optional<torch::stable::Tensor> cos_sin_cache);
void fused_kimi_k3_mla_decode_q_concat_kv_cache_fp8_insert(
torch::stable::Tensor const& ql_nope, torch::stable::Tensor const& q_pe,
torch::stable::Tensor const& kv_c_normed, torch::stable::Tensor const& k_pe,
torch::stable::Tensor& mqa_q, torch::stable::Tensor& k_cache,
torch::stable::Tensor const& slot_mapping,
torch::stable::Tensor const& q_scale_inv,
torch::stable::Tensor const& cache_scale_inv, int64_t cache_block_size,
std::optional<torch::stable::Tensor> position_ids,
std::optional<torch::stable::Tensor> cos_sin_cache);
void fused_kimi_k3_mla_decode_q_concat_ds_mla_insert(
torch::stable::Tensor const& ql_nope, torch::stable::Tensor const& q_pe,
torch::stable::Tensor const& kv_c_normed, torch::stable::Tensor const& k_pe,
torch::stable::Tensor& mqa_q, torch::stable::Tensor& k_cache,
torch::stable::Tensor const& slot_mapping, int64_t cache_block_size,
std::optional<torch::stable::Tensor> position_ids,
std::optional<torch::stable::Tensor> cos_sin_cache);
void fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_fp8_insert(
torch::stable::Tensor const& q, torch::stable::Tensor const& kv,
torch::stable::Tensor& q_fp8, torch::stable::Tensor& k_cache,
@@ -370,30 +315,6 @@ void fused_minimax_m3_qknorm_rope_kv_insert(
std::optional<torch::stable::Tensor> index_q_out,
const std::string& kv_cache_dtype, bool skip_index_branch);
#ifdef VLLM_ENABLE_FUSED_KDA_DECODE
void fused_kda_decode(
torch::stable::Tensor const& x, torch::stable::Tensor const& weight,
std::optional<torch::stable::Tensor> bias,
torch::stable::Tensor& conv_state, torch::stable::Tensor const& raw_g,
torch::stable::Tensor const& raw_beta, torch::stable::Tensor const& a_log,
torch::stable::Tensor const& dt_bias,
torch::stable::Tensor const& state_indices, torch::stable::Tensor& state,
torch::stable::Tensor& out, std::optional<double> lower_bound,
std::optional<torch::stable::Tensor> output_gate,
std::optional<torch::stable::Tensor> norm_weight, double norm_eps);
#endif
#ifdef VLLM_ENABLE_KIMI_K3_ATTN_RES
void kimi_k3_attn_res(torch::stable::Tensor& prefix,
torch::stable::Tensor const& delta,
torch::stable::Tensor const& blocks,
torch::stable::Tensor const& norm_weight,
torch::stable::Tensor const& qk_weight,
torch::stable::Tensor const& output_norm_weight,
torch::stable::Tensor& output, int64_t num_blocks,
double eps, double output_norm_eps);
#endif
// Sampler kernels (shared CUDA/ROCm)
void apply_repetition_penalties_(
torch::stable::Tensor& logits, const torch::stable::Tensor& prompt_mask,
@@ -451,20 +372,6 @@ fptr_t init_custom_ar(const std::vector<int64_t>& fake_ipc_ptrs,
void all_reduce(fptr_t _fa, torch::stable::Tensor& inp,
torch::stable::Tensor& out, fptr_t reg_buffer,
int64_t reg_buffer_sz_bytes);
void custom_all_gather(fptr_t _fa, torch::stable::Tensor& inp,
torch::stable::Tensor& out, fptr_t reg_buffer,
int64_t reg_buffer_sz_bytes);
void mnnvl_lamport_all_gather(fptr_t _fa, torch::stable::Tensor& inp,
torch::stable::Tensor& out, fptr_t local_buffer,
fptr_t multicast_buffer, fptr_t epoch_buffer,
int64_t stage_sz_bytes);
void custom_reduce_scatter(fptr_t _fa, torch::stable::Tensor& inp,
torch::stable::Tensor& out, fptr_t reg_buffer,
int64_t reg_buffer_sz_bytes);
void mnnvl_lamport_reduce_scatter(fptr_t _fa, torch::stable::Tensor& inp,
torch::stable::Tensor& out,
fptr_t local_buffer, fptr_t epoch_buffer,
int64_t stage_sz_bytes);
void dispose(fptr_t _fa);
int64_t meta_size();
void register_buffer(fptr_t _fa, const std::vector<int64_t>& fake_ipc_ptrs);
@@ -503,12 +410,6 @@ void fatrelu_and_mul(torch::stable::Tensor& out, torch::stable::Tensor& input,
double threshold);
void swigluoai_and_mul(torch::stable::Tensor& out, torch::stable::Tensor& input,
double alpha = 1.702, double limit = 7.0);
void situ_and_mul(torch::stable::Tensor& out, torch::stable::Tensor& input,
double beta = 1.0, double linear_beta = -1.0);
void masked_situ_and_mul(torch::stable::Tensor& out,
torch::stable::Tensor& input,
const torch::stable::Tensor& expert_num_tokens,
double beta = 1.0, double linear_beta = -1.0);
void gelu_new(torch::stable::Tensor& out, torch::stable::Tensor& input);
void gelu_fast(torch::stable::Tensor& out, torch::stable::Tensor& input);
void gelu_quick(torch::stable::Tensor& out, torch::stable::Tensor& input);
@@ -585,13 +486,6 @@ void concat_and_cache_mla(torch::stable::Tensor& kv_c,
const std::string& kv_cache_dtype,
torch::stable::Tensor& scale);
void concat_and_cache_mla_grouped(torch::stable::Tensor& kv_c,
torch::stable::Tensor& k_pe,
torch::stable::Tensor& kv_cache_ptrs,
torch::stable::Tensor& slot_mapping,
int64_t block_size, int64_t block_stride,
int64_t entry_stride);
// NOTE: k_pe and kv_c order is flipped compared to concat_and_cache_mla
void concat_and_cache_mla_rope_fused(
torch::stable::Tensor& positions, torch::stable::Tensor& q_pe,
+1 -102
View File
@@ -324,8 +324,7 @@ STABLE_TORCH_LIBRARY_FRAGMENT(_C, ops) {
// DeepSeek V3 fused A GEMM (SM 9.0+, bf16 only, 1-16 tokens).
// conditionally compiled so impl registration is in source file
ops.def(
"dsv3_fused_a_gemm(Tensor! output, Tensor mat_a, Tensor mat_b, "
"bool enable_pdl=False) -> ()");
"dsv3_fused_a_gemm(Tensor! output, Tensor mat_a, Tensor mat_b) -> ()");
// BF16/FP32 x FP32 -> FP32 router GEMM for H=3072, E=256, M<=32 (SM90+).
// conditionally compiled so impl registration is in source file
@@ -448,48 +447,6 @@ STABLE_TORCH_LIBRARY_FRAGMENT(_C, ops) {
"Tensor fp8_scale, Tensor q_fp8_scale_inv, float eps, "
"int cache_block_size) -> ()");
// Kimi-K3 MLA epilogues: optional RoPE followed by concat/cache insertion.
ops.def(
"fused_kimi_k3_mla_key_concat_kv_cache_insert("
"Tensor! q, Tensor k_nope, Tensor k_pe, Tensor kv_c_normed, "
"Tensor! k_out, Tensor! k_cache, Tensor slot_mapping, "
"int cache_block_size, Tensor? position_ids=None, "
"Tensor? cos_sin_cache=None) -> ()");
ops.def(
"fused_kimi_k3_mla_key_concat_ds_mla_insert("
"Tensor! q, Tensor k_nope, Tensor k_pe, Tensor kv_c_normed, "
"Tensor! k_out, Tensor! k_cache, Tensor slot_mapping, "
"int cache_block_size, Tensor? position_ids=None, "
"Tensor? cos_sin_cache=None) -> ()");
ops.def(
"fused_kimi_k3_mla_qkv_quant_kv_cache_fp8_insert("
"Tensor q, Tensor k_nope, Tensor k_pe, Tensor kv_c_normed, Tensor v, "
"Tensor! q_fp8, Tensor! k_fp8, Tensor! v_fp8, Tensor! k_cache, "
"Tensor slot_mapping, Tensor q_scale_inv, Tensor k_scale_inv, "
"Tensor v_scale_inv, Tensor cache_scale_inv, int cache_block_size, "
"Tensor? position_ids=None, Tensor? cos_sin_cache=None) -> ()");
// Kimi-K3 MLA decode epilogue: concat mqa_q = [ql_nope | q_pe] and insert the
// latent [kv_c_normed | k_pe] into the paged cache (bf16 / fp8 / fp8_ds_mla).
ops.def(
"fused_kimi_k3_mla_decode_q_concat_kv_cache_insert("
"Tensor ql_nope, Tensor q_pe, Tensor kv_c_normed, Tensor k_pe, "
"Tensor! mqa_q, Tensor! k_cache, Tensor slot_mapping, "
"int cache_block_size, Tensor? position_ids=None, "
"Tensor? cos_sin_cache=None) -> ()");
ops.def(
"fused_kimi_k3_mla_decode_q_concat_kv_cache_fp8_insert("
"Tensor ql_nope, Tensor q_pe, Tensor kv_c_normed, Tensor k_pe, "
"Tensor! mqa_q, Tensor! k_cache, Tensor slot_mapping, "
"Tensor q_scale_inv, Tensor cache_scale_inv, int cache_block_size, "
"Tensor? position_ids=None, Tensor? cos_sin_cache=None) -> ()");
ops.def(
"fused_kimi_k3_mla_decode_q_concat_ds_mla_insert("
"Tensor ql_nope, Tensor q_pe, Tensor kv_c_normed, Tensor k_pe, "
"Tensor! mqa_q, Tensor! k_cache, Tensor slot_mapping, "
"int cache_block_size, Tensor? position_ids=None, "
"Tensor? cos_sin_cache=None) -> ()");
#ifndef USE_ROCM
ops.def(
"minimax_allreduce_rms_qk("
@@ -511,24 +468,6 @@ STABLE_TORCH_LIBRARY_FRAGMENT(_C, ops) {
"int block_size, Tensor!? q_out, Tensor!? index_q_out, "
"str kv_cache_dtype, bool skip_index_branch=False) -> ()");
#ifdef VLLM_ENABLE_FUSED_KDA_DECODE
ops.def(
"fused_kda_decode("
"Tensor x, Tensor weight, Tensor? bias, Tensor! conv_state, "
"Tensor raw_g, Tensor raw_beta, Tensor A_log, Tensor dt_bias, "
"Tensor state_indices, Tensor! state, Tensor! out, "
"float? lower_bound=None, Tensor? output_gate=None, "
"Tensor? norm_weight=None, float norm_eps=1e-5) -> ()");
#endif
#ifdef VLLM_ENABLE_KIMI_K3_ATTN_RES
ops.def(
"kimi_k3_attn_res("
"Tensor! prefix, Tensor delta, Tensor blocks, Tensor norm_weight, "
"Tensor qk_weight, Tensor output_norm_weight, Tensor! output, "
"int num_blocks, float eps, float output_norm_eps) -> ()");
#endif
// Apply repetition penalties to logits in-place.
ops.def(
"apply_repetition_penalties_(Tensor! logits, Tensor prompt_mask, "
@@ -591,14 +530,6 @@ STABLE_TORCH_LIBRARY_FRAGMENT(_C, ops) {
"limit=7.0) "
"-> ()");
// Kimi SITU (SituGLU) gated activation. linear_beta<=0 means unset.
ops.def(
"situ_and_mul(Tensor! out, Tensor input, float beta=1.0, float "
"linear_beta=-1.0) -> ()");
ops.def(
"masked_situ_and_mul(Tensor! out, Tensor input, Tensor "
"expert_num_tokens, float beta=1.0, float linear_beta=-1.0) -> ()");
// GELU implementation used in GPT-2.
ops.def("gelu_new(Tensor! out, Tensor input) -> ()");
@@ -757,30 +688,11 @@ STABLE_TORCH_LIBRARY_IMPL(_C, CUDA, ops) {
ops.impl(
"fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_fp8_insert",
TORCH_BOX(&fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_fp8_insert));
ops.impl("fused_kimi_k3_mla_key_concat_kv_cache_insert",
TORCH_BOX(&fused_kimi_k3_mla_key_concat_kv_cache_insert));
ops.impl("fused_kimi_k3_mla_key_concat_ds_mla_insert",
TORCH_BOX(&fused_kimi_k3_mla_key_concat_ds_mla_insert));
ops.impl("fused_kimi_k3_mla_qkv_quant_kv_cache_fp8_insert",
TORCH_BOX(&fused_kimi_k3_mla_qkv_quant_kv_cache_fp8_insert));
ops.impl("fused_kimi_k3_mla_decode_q_concat_kv_cache_insert",
TORCH_BOX(&fused_kimi_k3_mla_decode_q_concat_kv_cache_insert));
ops.impl("fused_kimi_k3_mla_decode_q_concat_kv_cache_fp8_insert",
TORCH_BOX(&fused_kimi_k3_mla_decode_q_concat_kv_cache_fp8_insert));
ops.impl("fused_kimi_k3_mla_decode_q_concat_ds_mla_insert",
TORCH_BOX(&fused_kimi_k3_mla_decode_q_concat_ds_mla_insert));
#ifndef USE_ROCM
ops.impl("minimax_allreduce_rms_qk", TORCH_BOX(&minimax_allreduce_rms_qk));
#endif
ops.impl("fused_minimax_m3_qknorm_rope_kv_insert",
TORCH_BOX(&fused_minimax_m3_qknorm_rope_kv_insert));
#ifdef VLLM_ENABLE_FUSED_KDA_DECODE
ops.impl("fused_kda_decode", TORCH_BOX(&fused_kda_decode));
#endif
#ifdef VLLM_ENABLE_KIMI_K3_ATTN_RES
ops.impl("kimi_k3_attn_res", TORCH_BOX(&kimi_k3_attn_res));
#endif
// Sampler kernels (shared CUDA/ROCm)
ops.impl("apply_repetition_penalties_",
@@ -803,8 +715,6 @@ STABLE_TORCH_LIBRARY_IMPL(_C, CUDA, ops) {
ops.impl("gelu_tanh_and_mul", TORCH_BOX(&gelu_tanh_and_mul));
ops.impl("fatrelu_and_mul", TORCH_BOX(&fatrelu_and_mul));
ops.impl("swigluoai_and_mul", TORCH_BOX(&swigluoai_and_mul));
ops.impl("situ_and_mul", TORCH_BOX(&situ_and_mul));
ops.impl("masked_situ_and_mul", TORCH_BOX(&masked_situ_and_mul));
ops.impl("gelu_new", TORCH_BOX(&gelu_new));
ops.impl("gelu_fast", TORCH_BOX(&gelu_fast));
ops.impl("gelu_quick", TORCH_BOX(&gelu_quick));
@@ -902,15 +812,6 @@ STABLE_TORCH_LIBRARY_FRAGMENT(_C_cache_ops, ops) {
" str kv_cache_dtype,"
" Tensor scale) -> ()");
// Grouped concat_and_cache_mla across all layers (bf16 only). Each
// layer's cache base pointer is read from kv_cache_ptrs.
ops.def(
"concat_and_cache_mla_grouped(Tensor kv_c, Tensor k_pe,"
" Tensor kv_cache_ptrs,"
" Tensor slot_mapping,"
" int block_size, int block_stride,"
" int entry_stride) -> ()");
// Rotate Q and K, then write to kv cache for MLA
ops.def(
"concat_and_cache_mla_rope_fused("
@@ -1009,8 +910,6 @@ STABLE_TORCH_LIBRARY_IMPL(_C_cache_ops, CUDA, ops) {
ops.impl("reshape_and_cache", TORCH_BOX(&reshape_and_cache));
ops.impl("reshape_and_cache_flash", TORCH_BOX(&reshape_and_cache_flash));
ops.impl("concat_and_cache_mla", TORCH_BOX(&concat_and_cache_mla));
ops.impl("concat_and_cache_mla_grouped",
TORCH_BOX(&concat_and_cache_mla_grouped));
ops.impl("concat_and_cache_mla_rope_fused",
TORCH_BOX(&concat_and_cache_mla_rope_fused));
ops.impl("convert_fp8", TORCH_BOX(&convert_fp8));
@@ -13,18 +13,10 @@ namespace vllm {
namespace fp8 {
#ifdef ENABLE_FP8
// Unspecialized conversions are a compile error: the old passthrough
// (`return x;`) silently skipped fp8 encoding for any (Tout, Tin) pair
// without a specialization below (e.g. the torch stable-ABI scalar types),
// corrupting quantized data with no runtime signal.
template <typename>
inline constexpr bool _no_conversion_specialization = false;
template <typename Tout, typename Tin>
__inline__ __device__ Tout vec_conversion(
const Tin& x, const __nv_fp8_interpretation_t fp8_type = __NV_E4M3) {
static_assert(_no_conversion_specialization<Tin>,
"no vec_conversion specialization for this (Tout, Tin) pair");
return x;
}
// float -> c10::Float8_e4m3fn
@@ -309,9 +301,7 @@ __inline__ __device__ bf16_8_t vec_conversion<bf16_8_t, Float8_>(
template <typename Tout, typename Tin>
__inline__ __device__ Tout scaled_vec_conversion(
const Tin& x, const float scale, const __nv_fp8_interpretation_t fp8_type) {
static_assert(
_no_conversion_specialization<Tin>,
"no scaled_vec_conversion specialization for this (Tout, Tin) pair");
return x;
}
// fp8 -> half
@@ -502,25 +492,6 @@ __inline__ __device__ uint8_t scaled_vec_conversion<uint8_t, __nv_bfloat16>(
__builtin_unreachable(); // Suppress missing return statement warning
}
// torch stable-ABI (headeronly) scalar types delegate to the CUDA-native
// conversions, so libtorch_stable kernels dispatched on c10::BFloat16 /
// c10::Half quantize correctly without manual casts.
template <>
__inline__ __device__ uint8_t scaled_vec_conversion<uint8_t, c10::BFloat16>(
const c10::BFloat16& a, const float scale,
const __nv_fp8_interpretation_t fp8_type) {
return scaled_vec_conversion<uint8_t, __nv_bfloat16>(
reinterpret_cast<const __nv_bfloat16&>(a), scale, fp8_type);
}
template <>
__inline__ __device__ uint8_t scaled_vec_conversion<uint8_t, c10::Half>(
const c10::Half& a, const float scale,
const __nv_fp8_interpretation_t fp8_type) {
return scaled_vec_conversion<uint8_t, uint16_t>(
reinterpret_cast<const uint16_t&>(a), scale, fp8_type);
}
// float -> fp8
template <>
__inline__ __device__ uint8_t scaled_vec_conversion<uint8_t, float>(
+30 -8
View File
@@ -22,17 +22,25 @@ template <typename AllReduceKernel, typename T>
__global__ __quickreduce_launch_bounds_two_shot__ static void
allreduce_prototype_twoshot(T const* A, T* B, uint32_t N, uint32_t num_blocks,
int rank, uint8_t** dbuffer_list,
uint32_t data_offset, uint32_t flag_color,
uint32_t data_offset, uint32_t* d_flag_counters,
int64_t data_size_per_phase) {
int block = blockIdx.x;
int grid = gridDim.x;
// Load this block's counter from device memory and advance it on-device,
// so the color keeps changing across graph replays instead of being frozen.
uint32_t flag_color = d_flag_counters[blockIdx.x];
while (block < num_blocks) {
AllReduceKernel::run(A, B, N, block, rank, dbuffer_list, data_offset,
flag_color, data_size_per_phase);
block += grid;
flag_color++;
}
// All threads compute the same final value; one writer per block is enough.
if (threadIdx.x == 0 && threadIdx.y == 0) {
d_flag_counters[blockIdx.x] = flag_color;
}
}
#define TWOSHOT_DISPATCH(__codec) \
@@ -42,21 +50,21 @@ allreduce_prototype_twoshot(T const* A, T* B, uint32_t N, uint32_t num_blocks,
hipLaunchKernelGGL((allreduce_prototype_twoshot<AllReduceKernel, T>), \
dim3(grid), dim3(kBlockTwoShot), 0, stream, A, B, N, \
num_blocks, rank, dbuffer_list, data_offset, \
flag_color, this->kMaxProblemSize); \
d_flag_counters, this->kMaxProblemSize); \
} else if (world_size == 4) { \
using LineCodec = __codec<T, 4>; \
using AllReduceKernel = AllReduceTwoshot<T, LineCodec, cast_bf2half>; \
hipLaunchKernelGGL((allreduce_prototype_twoshot<AllReduceKernel, T>), \
dim3(grid), dim3(kBlockTwoShot), 0, stream, A, B, N, \
num_blocks, rank, dbuffer_list, data_offset, \
flag_color, this->kMaxProblemSize); \
d_flag_counters, this->kMaxProblemSize); \
} else if (world_size == 8) { \
using LineCodec = __codec<T, 8>; \
using AllReduceKernel = AllReduceTwoshot<T, LineCodec, cast_bf2half>; \
hipLaunchKernelGGL((allreduce_prototype_twoshot<AllReduceKernel, T>), \
dim3(grid), dim3(kBlockTwoShot), 0, stream, A, B, N, \
num_blocks, rank, dbuffer_list, data_offset, \
flag_color, this->kMaxProblemSize); \
d_flag_counters, this->kMaxProblemSize); \
}
// INT3 only retains good performance on TP2 (world_size == 2). On TP4/TP8
@@ -69,7 +77,7 @@ allreduce_prototype_twoshot(T const* A, T* B, uint32_t N, uint32_t num_blocks,
hipLaunchKernelGGL((allreduce_prototype_twoshot<AllReduceKernel, T>), \
dim3(grid), dim3(kBlockTwoShot), 0, stream, A, B, N, \
num_blocks, rank, dbuffer_list, data_offset, \
flag_color, this->kMaxProblemSize); \
d_flag_counters, this->kMaxProblemSize); \
} else { \
throw std::runtime_error( \
"INT3 quick all-reduce is only supported for world_size == 2 " \
@@ -94,7 +102,7 @@ struct DeviceComms {
static int constexpr kMaxWorldSize = 8;
bool initialized = false;
uint32_t flag_color = 1;
uint32_t* d_flag_counters = nullptr;
int world_size;
int rank;
@@ -128,6 +136,16 @@ struct DeviceComms {
// Clear the flags buffer.
HIP_CHECK(hipMemset(dbuffer, 0, flags_buffer_size));
// One flag-color counter per block, advanced by the kernel. Start at 1
// to stay clear of the flags buffer we just zeroed.
HIP_CHECK(hipMalloc(&d_flag_counters, kMaxNumBlocks * sizeof(uint32_t)));
{
std::vector<uint32_t> init_color(kMaxNumBlocks, 1u);
HIP_CHECK(hipMemcpy(d_flag_counters, init_color.data(),
kMaxNumBlocks * sizeof(uint32_t),
hipMemcpyHostToDevice));
}
// Device-side list of IPC buffers.
buffer_list.resize(world_size);
HIP_CHECK(hipMalloc(&dbuffer_list, world_size * sizeof(uint8_t*)));
@@ -144,6 +162,12 @@ struct DeviceComms {
hipIpcMemHandle_t const get_handle() { return buffer_ipc_handle; }
void destroy() {
// Allocated before `initialized` flips true, so free it on its own guard
// to avoid a leak if init fails partway through.
if (d_flag_counters) {
HIP_CHECK(hipFree(d_flag_counters));
d_flag_counters = nullptr;
}
if (initialized) {
for (int i = 0; i < world_size; i++) {
if (i != rank) {
@@ -211,8 +235,6 @@ struct DeviceComms {
break;
}
HIP_CHECK(cudaGetLastError());
// Rotate the flag color.
flag_color += divceil(N, grid);
}
};
+3 -1
View File
@@ -56,7 +56,9 @@ nav:
- API Reference:
- api/README.md
- api/vllm
- CLI Reference: cli
- CLI Reference:
- cli/README.md
- vllm: cli
- Community:
- community/*
- Governance: governance
+7 -9
View File
@@ -1,10 +1,8 @@
nav:
- README.md
- serve.md
- chat.md
- complete.md
- run-batch.md
- vllm bench:
- bench/**/*.md
- vllm launch:
- launch/**/*.md
- "*.md"
- bench:
- bench/*.md
- sweep:
- bench/sweep/*.md
- launch:
- launch/*.md
-9
View File
@@ -1,9 +0,0 @@
# vllm bench latency
## JSON CLI Arguments
--8<-- "docs/cli/json_tip.inc.md"
## Arguments
--8<-- "docs/generated/argparse/bench_latency.inc.md"
-55
View File
@@ -1,55 +0,0 @@
# vllm bench mm-processor
## Overview
`vllm bench mm-processor` profiles the multimodal input processor pipeline of
vision-language models. It measures per-stage latency from the HuggingFace
processor through to the encoder forward pass, helping you identify
preprocessing bottlenecks and understand how different image resolutions or
item counts affect end-to-end request time.
The benchmark supports two data sources: synthetic random multimodal inputs
(`random-mm`) and HuggingFace datasets (`hf`). Warmup requests are run before
measurement to ensure stable results.
## Quick Start
```bash
vllm bench mm-processor \
--model Qwen/Qwen2-VL-7B-Instruct \
--dataset-name random-mm \
--num-prompts 50 \
--random-input-len 300 \
--random-output-len 40 \
--random-mm-base-items-per-request 2 \
--random-mm-limit-mm-per-prompt '{"image": 3, "video": 0}' \
--random-mm-bucket-config '{(256, 256, 1): 0.7, (720, 1280, 1): 0.3}'
```
## Measured Stages
| Stage | Description |
| ----- | ----------- |
| `get_mm_hashes_secs` | Time spent hashing multimodal inputs |
| `get_cache_missing_items_secs` | Time spent looking up the processor cache |
| `apply_hf_processor_secs` | Time spent in the HuggingFace processor |
| `merge_mm_kwargs_secs` | Time spent merging multimodal kwargs |
| `apply_prompt_updates_secs` | Time spent updating prompt tokens |
| `preprocessor_total_secs` | Total preprocessing time |
| `encoder_forward_secs` | Time spent in the encoder model forward pass |
| `num_encoder_calls` | Number of encoder invocations per request |
The benchmark also reports end-to-end latency (TTFT + decode time) per
request. Use `--metric-percentiles` to select which percentiles to report
(default: p99) and `--output-json` to save results.
For more examples (HF datasets, warmup, JSON output), see
[Benchmarking CLI — Multimodal Processor Benchmark](../../benchmarking/cli.md#multimodal-processor-benchmark).
## JSON CLI Arguments
--8<-- "docs/cli/json_tip.inc.md"
## Arguments
--8<-- "docs/generated/argparse/bench_mm_processor.inc.md"
-9
View File
@@ -1,9 +0,0 @@
# vllm bench serve
## JSON CLI Arguments
--8<-- "docs/cli/json_tip.inc.md"
## Arguments
--8<-- "docs/generated/argparse/bench_serve.inc.md"
-9
View File
@@ -1,9 +0,0 @@
# vllm bench sweep plot
## JSON CLI Arguments
--8<-- "docs/cli/json_tip.inc.md"
## Arguments
--8<-- "docs/generated/argparse/bench_sweep_plot.inc.md"
-9
View File
@@ -1,9 +0,0 @@
# vllm bench sweep plot_pareto
## JSON CLI Arguments
--8<-- "docs/cli/json_tip.inc.md"
## Arguments
--8<-- "docs/generated/argparse/bench_sweep_plot_pareto.inc.md"
-9
View File
@@ -1,9 +0,0 @@
# vllm bench sweep serve
## JSON CLI Arguments
--8<-- "docs/cli/json_tip.inc.md"
## Arguments
--8<-- "docs/generated/argparse/bench_sweep_serve.inc.md"
-9
View File
@@ -1,9 +0,0 @@
# vllm bench sweep serve_workload
## JSON CLI Arguments
--8<-- "docs/cli/json_tip.inc.md"
## Arguments
--8<-- "docs/generated/argparse/bench_sweep_serve_workload.inc.md"
-9
View File
@@ -1,9 +0,0 @@
# vllm bench throughput
## JSON CLI Arguments
--8<-- "docs/cli/json_tip.inc.md"
## Arguments
--8<-- "docs/generated/argparse/bench_throughput.inc.md"
-5
View File
@@ -1,5 +0,0 @@
# vllm chat
## Arguments
--8<-- "docs/generated/argparse/chat.inc.md"
-5
View File
@@ -1,5 +0,0 @@
# vllm complete
## Arguments
--8<-- "docs/generated/argparse/complete.inc.md"
-10
View File
@@ -1,10 +0,0 @@
<!-- markdownlint-disable MD041 -->
When passing JSON CLI arguments, the following sets of arguments are equivalent:
- `--json-arg '{"key1": "value1", "key2": {"key3": "value2"}}'`
- `--json-arg.key1 value1 --json-arg.key2.key3 value2`
Additionally, list elements can be passed individually using `+`:
- `--json-arg '{"key4": ["value3", "value4", "value5"]}'`
- `--json-arg.key4+ value3 --json-arg.key4+='value4,value5'`
-22
View File
@@ -1,22 +0,0 @@
# vllm launch render
## Overview
`vllm launch render` starts a GPU-less rendering server for preprocessing and
postprocessing only.
```bash
vllm launch render meta-llama/Llama-3.2-1B-Instruct --port 8100
```
This command reuses the standard serving parser, so model, frontend,
networking, and related CLI options follow the same conventions as
[`vllm serve`](../serve.md).
## JSON CLI Arguments
--8<-- "docs/cli/json_tip.inc.md"
## Arguments
--8<-- "docs/generated/argparse/launch_render.inc.md"
-9
View File
@@ -1,9 +0,0 @@
# vllm run-batch
## JSON CLI Arguments
--8<-- "docs/cli/json_tip.inc.md"
## Arguments
--8<-- "docs/generated/argparse/run-batch.inc.md"
-9
View File
@@ -1,9 +0,0 @@
# vllm serve
## JSON CLI Arguments
--8<-- "docs/cli/json_tip.inc.md"
## Arguments
--8<-- "docs/generated/argparse/serve.inc.md"
+1 -9
View File
@@ -11,12 +11,4 @@ Engine arguments control the behavior of the vLLM engine.
The engine argument classes, [EngineArgs][vllm.engine.arg_utils.EngineArgs] and [AsyncEngineArgs][vllm.engine.arg_utils.AsyncEngineArgs], are a combination of the configuration classes defined in [vllm.config][]. Therefore, if you are interested in developer documentation, we recommend looking at these configuration classes as they are the source of truth for types, defaults and docstrings.
--8<-- "docs/cli/json_tip.inc.md"
## `EngineArgs`
--8<-- "docs/generated/argparse/engine_args.inc.md"
## `AsyncEngineArgs`
--8<-- "docs/generated/argparse/async_engine_args.inc.md"
--8<-- "gen:engine-args"
+1 -1
View File
@@ -195,7 +195,7 @@ Provide a fast duration→token estimate to improve streaming usage statistics:
The API server takes care of basic audio I/O and optional chunking before building prompts:
- Resampling: Input audio is resampled to `SpeechToTextConfig.sample_rate` using `AudioResampler`.
- Chunking: If `SpeechToTextConfig.allow_audio_chunking` is True and the duration exceeds `max_audio_clip_s`, the server splits the audio into overlapping chunks and generates a prompt per chunk. Overlap is controlled by `overlap_chunk_second`.
- Chunking: If `SpeechToTextConfig.allow_audio_chunking` is True and the duration exceeds `max_audio_clip_s`, the server splits the audio into chunks and generates a prompt per chunk. There is no overlap between chunks, overlap_chunk_second controls the size of the search window used to find the split point.
- Energy-aware splitting: When `min_energy_split_window_size` is set, the server finds low-energy regions to minimize cutting within words.
Relevant server logic:
+20
View File
@@ -8,6 +8,26 @@ toc_depth: 2
--8<-- "docs/getting_started/installation/gpu.md:pre-built-images"
## Persist the compile cache across containers
Mounting the Hugging Face cache keeps model weights across containers, but each
new container still starts with an empty `VLLM_CACHE_ROOT` (default
`~/.cache/vllm`) and recompiles the model's `torch.compile` artifacts. Mount a
named volume at that path to reuse the inductor, Triton, and AOT artifacts from
the second container onward:
```bash
docker run --rm --gpus all \
-v ~/.cache/huggingface:/root/.cache/huggingface \
-v vllm-cache:/root/.cache/vllm \
-p 8000:8000 \
vllm/vllm-openai:latest \
meta-llama/Llama-3.1-8B-Instruct
```
See [Faster Startup](../configuration/optimization.md#faster-startup) for the
mechanism and for what invalidates the cache.
## Run as a non-root user
The CUDA `vllm/vllm-openai` image runs as root by default for backward
+14 -90
View File
@@ -1,15 +1,9 @@
# Attention Backend Feature Support
This document is auto-generated by `tools/pre_commit/generate_attention_backend_docs.py`.
It shows the feature support for each registered attention backend
based on the checks in `AttentionBackend.validate_configuration()`.
**Do not edit this file manually.** Run the following command to
regenerate it:
```bash
python tools/pre_commit/generate_attention_backend_docs.py
```
The priority and feature tables on this page are auto-generated from the
attention backend registry by
`docs/mkdocs/gen_files/generate_attention_backends.py`, based on the checks in
`AttentionBackend.validate_configuration()`.
## Setting the Attention Backend
@@ -98,40 +92,11 @@ Priority is **1 = highest** (tried first).
### Standard Attention (MHA, MQA, GQA)
**Blackwell (SM 10.x):**
| Priority | Backend |
| -------- | ------- |
| 1 | `FLASHINFER` |
| 2 | `FLASH_ATTN` |
| 3 | `TRITON_ATTN` |
| 4 | `FLEX_ATTENTION` |
| 5 | `TURBOQUANT` |
**Ampere/Hopper (SM 8.x-9.x):**
| Priority | Backend |
| -------- | ------- |
| 1 | `FLASH_ATTN` |
| 2 | `FLASHINFER` |
| 3 | `TRITON_ATTN` |
| 4 | `FLEX_ATTENTION` |
| 5 | `TURBOQUANT` |
--8<-- "gen:priority-standard"
### MLA Attention (DeepSeek-style)
**Blackwell (SM 10.x):**
| Priority | Backend |
| -------- | ------- |
| 1 | `FLASHINFER_MLA` |
| 2 | `TOKENSPEED_MLA` |
| 3 | `CUTLASS_MLA` |
| 4 | `FLASH_ATTN_MLA` |
| 5 | `FLASHMLA` |
| 6 | `TRITON_MLA` |
| 7 | `FLASHINFER_MLA_SPARSE`**\*** |
| 8 | `FLASHMLA_SPARSE` |
--8<-- "gen:priority-mla"
> **\*** For sparse MLA, FP8 KV cache always prefers `FLASHINFER_MLA_SPARSE`. With BF16 KV cache, `FLASHINFER_MLA_SPARSE` is preferred for low query-head counts (<= 16), while `FLASHMLA_SPARSE` is preferred otherwise.
>
@@ -157,24 +122,7 @@ Priority is **1 = highest** (tried first).
## Standard Attention (MHA, MQA, GQA) Backends
| Backend | Version | Dtypes | KV Dtypes | Block Sizes | Head Sizes | Sink | Non-Causal | MM Prefix | DCP | Attention Types | Compute Cap. |
| ------- | ------- | ------ | --------- | ----------- | ---------- | ---- | ---------- | --------- | --- | --------------- | ------------ |
| `CPU_ATTN` | | fp16, bf16, fp32 | `auto`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | %16 | 32, 64, 80, 96, 112, 128, 160, 192, 224, 256, 512 | ❌ | ✅ | ❌ | ❌ | All | N/A |
| `FLASHINFER` | Native† | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | 16, 32, 64, 128, 256, 512, 1024 | 64, 128, 256, 512 | ❌ | ✅ | ❌ | ✅ | Decoder | 8.x-9.x |
| `FLASHINFER` | XQA† | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | 16, 32, 64, 128, 256, 512, 1024 | 64, 128, 256, 512 | ❌ | ❌ | ❌ | ✅ | Decoder | 9.0 |
| `FLASHINFER` | trtllm-gen† | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2`, `nvfp4` | 16, 32, 64, 128, 256, 512, 1024 | 64, 128, 256, 512 | ✅ | ✅ | ❌ | ✅ | Decoder | 10.x |
| `FLASH_ATTN` | FA2* | fp16, bf16 | `auto`, `float16`, `bfloat16` | %16 | Any | ❌ | ✅ | ❌ | ✅ | All | ≥8.0 |
| `FLASH_ATTN` | FA3* | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | %16 | Any | ✅ | ✅ | ❌ | ✅ | All | 9.x |
| `FLASH_ATTN` | FA4* | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | %16 | Any | ✅ | ✅ | ❌ | ✅ | All | ≥10.0 |
| `FLASH_ATTN_DIFFKV` | | fp16, bf16 | `auto` | Any | Any | ❌ | ❌ | ❌ | ✅ | Decoder | Any |
| `FLEX_ATTENTION` | | fp16, bf16, fp32 | `auto`, `float16`, `bfloat16` | %16 | Any | ❌ | ✅ | ✅ | ❌ | Decoder, Encoder Only | Any |
| `HPC_ATTN` | | fp16, bf16 | `auto`, `bfloat16`, `fp8_e4m3` | 64 | 128 | ❌ | ❌ | ❌ | ❌ | Decoder | ≥9.0 |
| `ROCM_AITER_FA` | | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | 16, 32 | 64, 128, 256 | ✅ | ✅ | ❌ | ❌ | Decoder | N/A |
| `ROCM_AITER_UNIFIED_ATTN` | | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | %16 | Any | ✅ | ❌ | ✅ | ❌ | All | N/A |
| `ROCM_ATTN` | | fp16, bf16, fp32 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | %16 | 32, 64, 80, 96, 128, 160, 192, 224, 256 | ❌ | ✅ | ✅ | ❌ | Decoder, Encoder, Encoder Only | N/A |
| `TRITON_ATTN` | | fp16, bf16, fp32 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2`, `int4_per_token_head`, `int8_per_token_head`, `fp8_per_token_head` | %16 | Any | ✅ | ✅ | ✅ | ❌ | All | Any |
| `TRITON_ATTN_DIFFKV` | | fp16, bf16 | `auto`, `bfloat16` | Any | Any | ❌ | ❌ | ❌ | ❌ | Decoder | Any |
| `TURBOQUANT` | | fp16, bf16 | `turboquant_k8v4`, `turboquant_4bit_nc`, `turboquant_k3v4_nc`, `turboquant_3bit_nc` | 16, 32, 64, 128 | Any | ❌ | ❌ | ❌ | ❌ | Decoder | Any |
--8<-- "gen:table-standard"
> **†** FlashInfer Native is the regular FlashInfer path. XQA is the SM90 decode path exposed through FlashInfer's TRTLLM decode API. trtllm-gen is used on SM100 and supports sinks. Disable XQA/trtllm-gen via `--attention-config.use_trtllm_attention=0`.
>
@@ -188,9 +136,7 @@ automatic priority lists above. A lightning indexer scores KV blocks, the
top-k blocks (plus fixed init/local blocks) are selected, and attention
attends only to those blocks; index keys live in a separate side cache.
| Backend | Dtypes | KV Dtypes | Block Sizes | Head Sizes | Sink | Non-Causal | MM Prefix | DCP | Attention Types | Compute Cap. |
| ------- | ------ | --------- | ----------- | ---------- | ---- | ---------- | --------- | --- | --------------- | ------------ |
| `MINIMAX_M3_SPARSE` | bf16, fp16 | `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | 128 | 128 | ❌ | ❌ | ❌ | ❌ | Decoder | Any |
--8<-- "gen:table-minimax"
## MLA (Multi-head Latent Attention) Backends
@@ -203,38 +149,20 @@ To explicitly select a prefill backend, use
Otherwise, the prefill backend is selected automatically at runtime based on
hardware and configuration.
| Backend | Description | Dtypes | Compute Cap. | Notes |
| ------- | ----------- | ------ | ------------ | ----- |
| `FLASH_ATTN`‡ | FlashAttention varlen (FA2/FA3/FA4) | fp16, bf16 | Any | (qk_nope_head_dim=128, qk_rope_head_dim=64, v_head_dim=128) (FA2/FA3/FA4) or (qk_nope_head_dim=64, qk_rope_head_dim=64, v_head_dim=128) (FA2/FA3/FA4) or (qk_nope_head_dim=192, qk_rope_head_dim=64, v_head_dim=256) (FA2/FA3 only) |
| `TRTLLM_RAGGED` | TensorRT-LLM ragged attention | fp16, bf16 | 10.x | (qk_nope_head_dim=128, qk_rope_head_dim=64, v_head_dim=128) or (qk_nope_head_dim=192, qk_rope_head_dim=64, v_head_dim=256) only |
| `FLASHINFER` | FlashInfer CUTLASS backend | fp16, bf16 | 10.x | (qk_nope_head_dim=128, qk_rope_head_dim=64, v_head_dim=128) only |
| `TOKENSPEED_MLA` | | fp16, bf16 | 10.x | (qk_nope_head_dim=128, qk_rope_head_dim=64, v_head_dim=128) only |
--8<-- "gen:table-mla-prefill"
> **‡** Automatic selection tries FlashAttention first. On Blackwell
> (SM100), the fallback order is TRT-LLM Ragged, FlashInfer, then
> TokenSpeed MLA. On other GPUs, only FlashAttention is considered.
> TokenSpeed MLA; for (qk_nope_head_dim=192, qk_rope_head_dim=64,
> v_head_dim=256) TRT-LLM Ragged is tried before FlashAttention. On other
> GPUs, only FlashAttention is considered.
### Decode Backends
MLA decode backends are selected using the standard
`-ac.backend=<BACKEND>` argument (e.g., `FLASHMLA`, `TRITON_MLA`).
| Backend | Dtypes | KV Dtypes | Block Sizes | Head Sizes | Sink | Non-Causal | Sparse | MM Prefix | DCP | Attention Types | Compute Cap. |
| ------- | ------ | --------- | ----------- | ---------- | ---- | ---------- | ------ | --------- | --- | --------------- | ------------ |
| `CUTLASS_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 128 | Any | ❌ | ❌ | ❌ | ❌ | ✅ | Decoder | 10.x |
| `FLASHINFER_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 32, 64 | Any | ❌ | ✅ | ❌ | ❌ | ✅ | Decoder | 10.x |
| `FLASHINFER_MLA_SPARSE` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 32, 64 | Any | ❌ | ❌ | ❌ | ❌ | ✅ | Decoder | 10.x |
| `FLASHINFER_MLA_SPARSE_SM120` | bf16 | `auto`, `fp8`, `fp8_e4m3`, `fp8_ds_mla` | 64, 256 | Any | ❌ | ❌ | ❌ | ❌ | ❌ | Decoder | 12.x |
| `FLASHMLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 64 | Any | ❌ | ❌ | ❌ | ❌ | ✅ | Decoder | 9.x-10.x |
| `FLASHMLA_SPARSE` | bf16 | `auto`, `bfloat16`, `fp8_ds_mla` | 64 | 576 | ❌ | ❌ | ✅ | ❌ | ❌ | Decoder | 9.x-10.x |
| `FLASH_ATTN_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16` | %16 | Any | ❌ | ❌ | ❌ | ❌ | ✅ | Decoder | 9.x |
| `FLASH_ATTN_MLA_SPARSE` | fp16, bf16 | `auto`, `float16`, `bfloat16` | 64 | Any | ❌ | ❌ | ✅ | ❌ | ❌ | Decoder | 9.x |
| `ROCM_AITER_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | %1 | Any | ❌ | ❌ | ❌ | ❌ | ❌ | Decoder | N/A |
| `ROCM_AITER_MLA_SPARSE` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 1, 64 | Any | ❌ | ❌ | ✅ | ❌ | ❌ | Decoder | N/A |
| `ROCM_AITER_TRITON_MLA` | fp16, bf16 | `auto` | Any | Any | ❌ | ❌ | ❌ | ❌ | ❌ | Decoder | N/A |
| `TOKENSPEED_MLA` | fp16, bf16 | `fp8`, `fp8_e4m3` | 32, 64 | Any | ❌ | ❌ | ❌ | ❌ | ✅ | Decoder | 10.x |
| `TRITON_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | %16 | Any | ❌ | ❌ | ❌ | ❌ | ✅ | Decoder | Any |
| `XPU_MLA_SPARSE` | fp16, bf16 | `auto`, `float16`, `bfloat16` | Any | 576 | ❌ | ❌ | ✅ | ❌ | ❌ | Decoder | Any |
--8<-- "gen:table-mla-decode"
### DeepSeek V4 Decode Backends
@@ -245,8 +173,4 @@ pipeline (compressor + SWA + indexer, 256-token blocks, head 512);
default on NVIDIA is `FLASHINFER_MLA_SPARSE_DSV4` on SM12x and
`FLASHMLA_SPARSE_DSV4` on other supported CUDA architectures.
| Backend | Dtypes | KV Dtypes | Block Sizes | Head Sizes | Sink | Non-Causal | Sparse | MM Prefix | DCP | Attention Types | Compute Cap. |
| ------- | ------ | --------- | ----------- | ---------- | ---- | ---------- | ------ | --------- | --- | --------------- | ------------ |
| `FLASHINFER_MLA_SPARSE_DSV4` | bf16 | `auto`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_ds_mla` | 256 | 512 | ✅ | ❌ | ✅ | ❌ | ❌ | Decoder | 10.x, 12.x |
| `FLASHMLA_SPARSE_DSV4` | bf16 | `auto`, `fp8_ds_mla`, `fp8` | 256 | 512 | ✅ | ❌ | ✅ | ❌ | ❌ | Decoder | 9.x-10.x |
| `ROCM_FLASHMLA_SPARSE_DSV4` | fp16, bf16 | `auto` | Any | Any | ❌ | ❌ | ❌ | ❌ | ❌ | Decoder | N/A |
--8<-- "gen:table-mla-v4-decode"
+1
View File
@@ -129,6 +129,7 @@ Models opt-in to encoder CUDA Graphs by implementing the [SupportsEncoderCudaGra
| `DeepseekOCRForCausalLM` | `DeepSeek-OCR` | ✅︎ | ❌︎ | ✅︎ |
| `Gemma3ForConditionalGeneration` | `Gemma3` | ✅︎ | ❌︎ | ❌︎ |
| `Glm4vForConditionalGeneration` | `GLM-4.1V, GLM-4.6V-Flash` | ✅︎ | ✅︎ | ❌︎ |
| `Gemma4ForConditionalGeneration` | `Gemma-4` | ✅︎ | ✅︎ | ❌︎ |
| `InternVLChatModel` | `InternVL3.5`, `InternVL3`, `InternVL2.5`, `InternVL2` | ✅︎ | ✅︎ | ❌︎ |
| `KimiVLForConditionalGeneration` | `Kimi-VL` | ✅︎ | ❌︎ | ❌︎ |
| `Llama4ForConditionalGeneration` | `Llama 4` | ✅︎ | ❌︎ | ❌︎ |
+66 -1
View File
@@ -157,7 +157,7 @@ Object keys follow the same run-configuration digest scheme as the filesystem ti
The P2P tier (`type: "p2p"`) shares completed KV blocks between vLLM instances over RDMA via NIXL. Each instance binds a control socket on `host:port` and exchanges blocks directly with peers — no shared filesystem required.
PYTHONHASHSEED environment variable must be set to the same fixed value on all nodes.
The `PYTHONHASHSEED` environment variable must be set to the same fixed value (e.g. `"0"`) on all nodes so that block content hashes match across instances (see [Cross-Process Sharing](#cross-process-sharing)). This is enforced: a P2P instance started without `PYTHONHASHSEED` set fails at startup, and each peer's value is verified during the connect handshake — a peer advertising a different `PYTHONHASHSEED` is rejected.
| Key | Required | Default | Notes |
| --- | --- | --- | --- |
@@ -176,6 +176,71 @@ Rather than embedding `host`/`port` in each `secondary_tiers` entry, set them on
- `VLLM_P2P_SIDE_CHANNEL_HOST` (default `localhost`): address the P2P control socket binds to. It is used **verbatim** as both the bind address and the identity peers dial back — there is no auto-detection (this mirrors `VLLM_NIXL_SIDE_CHANNEL_HOST`). The default binds the loopback interface only, so peers on another host cannot reach it. **For any cross-host P2P deployment you must set this explicitly to the node's routable IP** (e.g. the pod IP) before launching `vllm serve` — otherwise remote peers will fail to connect. The NIXL agent name is a separate per-process identifier, so peers sharing a `host:port` never collide.
- `VLLM_P2P_SIDE_CHANNEL_PORT` (default `5710`): base port for the P2P control socket. The port actually bound is `VLLM_P2P_SIDE_CHANNEL_PORT + data_parallel_index` — one socket per DP replica, matching NIXL (for DP=1 the offset is 0). The peer's port is passed as `remote_port` in `kv_transfer_params`; the router/EPP that selects the DP rank (e.g. via the `X-data-parallel-rank` header) computes `remote_port = base + rank`. The DP-index offset separates replicas *within* one deployment; two co-located *deployments* (a prefiller and a decoder on the same host) still need distinct base ports (e.g. decoder base `5711`) to avoid a bind collision.
#### Orchestration-Layer Protocol
The P2P tier does not decide *which* peer to pull from — that is the orchestration layer's job (the router/EPP and its scheduler). The orchestrator drives every transfer through a request's `kv_transfer_params` dict: it picks the request's role, allocates a unique transaction ID, and supplies the remote peer's address. All block lookup, hash matching, and NIXL transfer happen at the tier level below; the orchestrator only sets the correct role keys and enforces the allowed combinations.
Every vLLM instance is a symmetric **peer**. Per request it acts as a **consumer** (pulls KV blocks from a remote peer's CPU cache instead of computing locally) or a **producer** (serves blocks from its own CPU cache to remote consumers) — or both, on the same session, for different requests. Roles are chosen per request by the keys below; there are no fixed prefiller/decoder processes.
Three role keys are defined, each mapping to a sub-dict. All are optional; a request with none of them uses the tier only as a local CPU cache.
Each key names the **remote counterpart** this peer transfers with (not this
peer's own role), so the name reads as "the remote ___ I transfer with".
| Key | Set on | Value fields | Meaning |
| --- | --- | --- | --- |
| `remote_decoder` | prefill producer request | `kv_request_id` | Peer computes KV and keeps it available in CPU cache for the remote decoder to pull. |
| `remote_prefiller` | decode consumer request | `kv_request_id`, `remote_host`, `remote_port` | Peer pulls KV from the remote prefiller at the given address (classic P/D disaggregation). |
| `remote_kv_source` | P2P consumer request | `kv_request_id`, `remote_host`, `remote_port` | Peer looks up and pulls whatever blocks the remote source currently holds in CPU cache. |
Field semantics:
- `kv_request_id` (str): unique transaction ID allocated by the orchestrator and pushed to every peer involved in the transfer; used to correlate the lookup, fetch, and transfer-done messages. The producer is implicit — it serves whatever block hashes it currently holds in its CPU cache for that ID.
- `remote_host` (str): IP/hostname of the remote peer's control socket to query. Must be the peer's routable node IP (see [Environment Variables](#environment-variables)).
- `remote_port` (int): the peer's bound control-socket port, i.e. `base + data_parallel_index` for the selected DP rank.
Allowed and forbidden combinations:
- **`remote_decoder` + `remote_kv_source`** is the only legal multi-key combination: a prefill producer may *also* act as a P2P consumer for the same request — skipping prefix prefill by pulling cached blocks from a source while still keeping its own computed blocks available for a downstream decoder.
- Forbidden: `remote_prefiller` + `remote_decoder` (contradictory roles), `remote_prefiller` + `remote_kv_source` (two competing fetch sources), and all three together.
Minimal examples (values that would appear in the request's `kv_transfer_params`):
```python
# Prefill producer — compute and keep KV for a remote decoder to pull
kv_transfer_params = {"remote_decoder": {"kv_request_id": "<unique-transfer-id>"}}
# Decode consumer — pull KV from a specific prefiller (classic P/D)
kv_transfer_params = {
"remote_prefiller": {
"kv_request_id": "<unique-transfer-id>",
"remote_host": "<prefiller-node-ip>",
"remote_port": 5710,
}
}
# P2P consumer — pull whatever the source already has cached
kv_transfer_params = {
"remote_kv_source": {
"kv_request_id": "<unique-transfer-id>",
"remote_host": "<source-node-ip>",
"remote_port": 5710,
}
}
```
Runtime handshake for a P2P (or P/D) pull, once the orchestrator has set the keys above:
1. Both peers already have listener threads on their control sockets (see [Environment Variables](#environment-variables)).
2. **Lookup.** The consumer's tiering manager does per-block lookups; in P2P mode the tier returns `None` and registers the key. At `on_schedule_end` the consumer sends one **`LookupMsg`** (`kv_request_id` + block hashes) to the peer, per request, per step.
3. The producer matches those hashes against its local CPU cache and replies with a **`LookupRespMsg`** carrying the hit block hashes.
4. **Resolve.** Retried lookups now return hit / miss / in-flight. The consumer calls `submit_load` for hits only, allocating CPU slots only for hits.
5. The consumer sends a **`FetchMsg`** (`kv_request_id`, block hashes, destination block indexes).
6. The producer performs the **NIXL WRITE** transfer and sends **`TransferDone`** with a success status.
7. On `get_finished`, hits are loaded into GPU as ordinary cache hits; misses are recomputed by the engine.
In classic **P/D mode** (`remote_prefiller` set, no `remote_kv_source`), the lookup phase (steps 24) is skipped: the decode consumer assumes the prefiller holds all of the request's blocks, so every block `lookup()` returns an immediate hit and the consumer jumps straight to the **`FetchMsg`** in step 5. The `LookupMsg`/`LookupRespMsg` round-trip only happens in P2P mode, where the consumer does not know in advance which blocks the peer has cached.
## Tuning Tips
- `cpu_bytes_to_use`: a bigger CPU tier means fewer trips to slower secondary tiers and a higher hit rate. The value is total across all workers, not per-worker. Leave headroom for the rest of the host workload.
+21 -4
View File
@@ -818,16 +818,18 @@ Full example: [examples/generate/multimodal/openai_chat_completion_client_for_mu
#### Video Decoding Backend
vLLM decodes video bytes into frames using a selectable decoding backend. Three
vLLM decodes video bytes into frames using a selectable decoding backend. Five
backends are supported:
- `opencv` (default): OpenCV-based decoder.
- `pyav`: PyAV decoder.
- `torchcodec`: TorchCodec (PyTorch-native) decoder.
- `pynvvideocodec`: NVIDIA NVDEC-based decoder.
- `deepstream`: NVIDIA DeepStream NVDEC-based decoder.
All three backends are ultimately backed by FFmpeg. `torchcodec` lets
you choose which FFmpeg version is used while `opencv` and `pyav` rely on
whichever FFmpeg build they were linked against.
The CPU backends are backed by FFmpeg. `torchcodec` lets you choose which FFmpeg
version is used while `opencv` and `pyav` rely on whichever FFmpeg build they
were linked against.
Select the backend by passing the `backend` parameter via `--media-io-kwargs`:
@@ -854,6 +856,21 @@ vllm serve Qwen/Qwen3-VL-30B-A3B-Instruct \
--media-io-kwargs '{"video": {"backend": "torchcodec", "seek_mode": "approximate", "num_ffmpeg_threads": 4}}'
```
**PyNvVideoCodec-specific parameters:**
- `hw_decoders`: Maximum number of concurrent hardware decoder slots retained
by each API server process. It must be a positive integer and defaults to `2`,
which is the recommended starting point for concurrent video workloads.
Because vLLM reserves GPU memory for these slots at startup, this value cannot
be overridden per request. Benchmark before increasing it because each
additional slot increases the GPU memory reservation.
```bash
# Example: explicitly use the recommended 2 hardware decoders
vllm serve Qwen/Qwen3-VL-30B-A3B-Instruct \
--media-io-kwargs '{"video": {"backend": "pynvvideocodec", "hw_decoders": 2}}'
```
#### Video Frame Recovery
For improved robustness when processing potentially corrupted or truncated video files, vLLM supports optional frame recovery using a dynamic window forward-scan approach. When enabled, if a target frame fails to load during sequential reading, the next successfully grabbed frame (before the next target frame) will be used in its place.
+14
View File
@@ -19,6 +19,20 @@ following `quantization.quant_algo` values:
- `NVFP4`: ModelOpt NVFP4 checkpoints (use `quantization="modelopt_fp4"`).
- `MXFP8`: ModelOpt MXFP8 checkpoints (use `quantization="modelopt_mxfp8"`).
!!! note
For NVFP4 checkpoints, vLLM selects a GEMM kernel automatically at load
time from the backends available on the current platform (CUTLASS,
FlashInfer, Marlin, and others). On GPUs without a supported native FP4
GEMM kernel, vLLM falls back to weight-only (W4A16) execution via Marlin
and logs a warning; this may reduce throughput for compute-heavy
workloads. Use `--linear-backend` to override the automatic selection
(this replaces the deprecated `VLLM_NVFP4_GEMM_BACKEND` environment
variable). Values relevant to NVFP4 include `cutlass`,
`flashinfer_cutlass`, `flashinfer_trtllm`, `flashinfer_cudnn`, and
`marlin`; the full list is documented under `KernelConfig` on the
[Engine Arguments](../../configuration/engine_args.md) page and shown by
`vllm serve --help=KernelConfig`.
## Quantizing HuggingFace Models with PTQ
You can quantize HuggingFace models using the example scripts provided in the Model Optimizer repository. The primary script for LLM PTQ is typically found within the `examples/llm_ptq` directory.
@@ -2,6 +2,7 @@
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import importlib.metadata
import importlib.util
import inspect
import logging
import sys
import textwrap
@@ -10,17 +11,21 @@ from argparse import SUPPRESS, Action, HelpFormatter
from collections.abc import Callable, Iterable
from importlib.machinery import ModuleSpec
from pathlib import Path
from typing import TYPE_CHECKING, Literal
from typing import TYPE_CHECKING
from unittest.mock import MagicMock, patch
import mkdocs_gen_files
import regex as re
from pydantic_core import core_schema
logger = logging.getLogger("mkdocs")
ROOT_DIR = Path(__file__).parent.parent.parent.parent
ARGPARSE_DOC_DIR = ROOT_DIR / "docs/generated/argparse"
sys.path.insert(0, str(ROOT_DIR))
sys.path.insert(0, str(Path(__file__).parent))
from generated_content import fill_markers # noqa: E402
def mock_if_no_torch(mock_module: str, mock: MagicMock):
@@ -132,8 +137,8 @@ def auto_mock(module_name: str, attr: str, max_mocks: int = 100):
bench_latency = auto_mock("vllm.benchmarks", "latency")
bench_mm_processor = auto_mock("vllm.benchmarks", "mm_processor")
bench_serve = auto_mock("vllm.benchmarks", "serve")
bench_startup = auto_mock("vllm.benchmarks", "startup")
bench_sweep_plot = auto_mock("vllm.benchmarks.sweep.plot", "SweepPlotArgs")
bench_sweep_plot_pareto = auto_mock(
"vllm.benchmarks.sweep.plot_pareto", "SweepPlotParetoArgs"
@@ -142,12 +147,28 @@ bench_sweep_serve = auto_mock("vllm.benchmarks.sweep.serve", "SweepServeArgs")
bench_sweep_serve_workload = auto_mock(
"vllm.benchmarks.sweep.serve_workload", "SweepServeWorkloadArgs"
)
bench_sweep_startup = auto_mock("vllm.benchmarks.sweep.startup", "SweepStartupArgs")
bench_throughput = auto_mock("vllm.benchmarks", "throughput")
AsyncEngineArgs = auto_mock("vllm.engine.arg_utils", "AsyncEngineArgs")
EngineArgs = auto_mock("vllm.engine.arg_utils", "EngineArgs")
ChatCommand = auto_mock("vllm.entrypoints.cli.openai", "ChatCommand")
CompleteCommand = auto_mock("vllm.entrypoints.cli.openai", "CompleteCommand")
BenchmarkSubcommand = auto_mock(
"vllm.entrypoints.cli.benchmark.main", "BenchmarkSubcommand"
)
import_bench_subcommands = auto_mock(
"vllm.entrypoints.cli.benchmark.main", "_import_bench_subcommand_modules"
)
BenchmarkSubcommandBase = auto_mock(
"vllm.entrypoints.cli.benchmark.base", "BenchmarkSubcommandBase"
)
BenchmarkMMProcessorSubcommand = auto_mock(
"vllm.entrypoints.cli.benchmark.mm_processor", "BenchmarkMMProcessorSubcommand"
)
LaunchSubcommandBase = auto_mock("vllm.entrypoints.cli.launch", "LaunchSubcommandBase")
launch_description = auto_mock("vllm.entrypoints.cli.launch", "DESCRIPTION")
RenderSubcommand = auto_mock("vllm.entrypoints.cli.launch", "RenderSubcommand")
sweep_subcommands = auto_mock("vllm.benchmarks.sweep.cli", "SUBCOMMANDS")
openai_cli_args = auto_mock("vllm.entrypoints.openai", "cli_args")
openai_run_batch = auto_mock("vllm.entrypoints.openai", "run_batch")
@@ -179,7 +200,7 @@ class MarkdownFormatter(HelpFormatter):
def add_text(self, text: str):
if text:
self._markdown_output.append(f"{text.strip()}\n\n")
self._markdown_output.append(f"{inspect.cleandoc(text)}\n\n")
def add_usage(self, usage, actions, groups, prefix=None):
pass
@@ -241,49 +262,163 @@ def create_parser(add_cli_args, **kwargs) -> FlexibleArgumentParser:
return _parser or parser
def on_startup(command: Literal["build", "gh-deploy", "serve"], dirty: bool):
logger.info("Generating argparse documentation")
logger.debug("Root directory: %s", ROOT_DIR.resolve())
logger.debug("Output directory: %s", ARGPARSE_DOC_DIR.resolve())
# Create the ARGPARSE_DOC_DIR if it doesn't exist
if not ARGPARSE_DOC_DIR.exists():
ARGPARSE_DOC_DIR.mkdir(parents=True)
# Create parsers to document
parsers = {
# Engine args
"engine_args": create_parser(EngineArgs.add_cli_args),
"async_engine_args": create_parser(
AsyncEngineArgs.add_cli_args, async_args_only=True
),
# CLI
"serve": create_parser(openai_cli_args.make_arg_parser),
"chat": create_parser(ChatCommand.add_cli_args),
"complete": create_parser(CompleteCommand.add_cli_args),
"launch_render": create_parser(RenderSubcommand.add_cli_args),
"run-batch": create_parser(openai_run_batch.make_arg_parser),
# Benchmark CLI
"bench_latency": create_parser(bench_latency.add_cli_args),
"bench_mm_processor": create_parser(bench_mm_processor.add_cli_args),
"bench_serve": create_parser(bench_serve.add_cli_args),
"bench_sweep_plot": create_parser(bench_sweep_plot.add_cli_args),
"bench_sweep_plot_pareto": create_parser(bench_sweep_plot_pareto.add_cli_args),
"bench_sweep_serve": create_parser(bench_sweep_serve.add_cli_args),
"bench_sweep_serve_workload": create_parser(
bench_sweep_serve_workload.add_cli_args
),
"bench_throughput": create_parser(bench_throughput.add_cli_args),
}
# Generate documentation for each parser
for stem, parser in parsers.items():
doc_path = ARGPARSE_DOC_DIR / f"{stem}.inc.md"
# Specify encoding for building on Windows
with open(doc_path, "w", encoding="utf-8") as f:
f.write(super(type(parser), parser).format_help())
logger.info("Argparse generated: %s", doc_path.relative_to(ROOT_DIR))
def format_help(parser: FlexibleArgumentParser) -> str:
"""Format a parser's help as markdown using `MarkdownFormatter`."""
return super(type(parser), parser).format_help()
if __name__ == "__main__":
on_startup("build", False)
# Absolute docs URLs are kept in the help text because they are useful in the
# terminal. Wrap them as markdown links so the `url_schemes` hook can rewrite
# them into doc-relative links / cross-references at render time.
_DOCS_URL = re.compile(r"https://docs\.vllm\.ai/en/[^/\s]+/[^\s)>]+")
def linkify_docs_urls(text: str) -> str:
"""Wrap bare docs.vllm.ai URLs in help text as markdown links."""
return _DOCS_URL.sub(lambda m: f"[{m.group()}]({m.group()})", text)
logger.info("Generating argparse documentation")
logger.debug("Root directory: %s", ROOT_DIR.resolve())
# The JSON tip is always rendered immediately before generated argument content,
# and the generator is its only consumer, so it lives here rather than in a
# separate snippet file. (The runtime terminal equivalent is
# `FlexibleArgumentParser._json_tip` in vllm/utils/argparse_utils.py.)
JSON_TIP = """## JSON CLI Arguments
When passing JSON CLI arguments, the following sets of arguments are equivalent:
- `--json-arg '{"key1": "value1", "key2": {"key3": "value2"}}'`
- `--json-arg.key1 value1 --json-arg.key2.key3 value2`
Additionally, list elements can be passed individually using `+`:
- `--json-arg '{"key4": ["value3", "value4", "value5"]}'`
- `--json-arg.key4+ value3 --json-arg.key4+='value4,value5'`
"""
# Argument sections filled into `gen:` markers on handwritten pages
engine_args = create_parser(EngineArgs.add_cli_args)
async_engine_args = create_parser(AsyncEngineArgs.add_cli_args, async_args_only=True)
fill_markers(
"configuration/engine_args.md",
{
"engine-args": (
f"{JSON_TIP}## `EngineArgs`\n\n"
f"{linkify_docs_urls(format_help(engine_args))}"
f"## `AsyncEngineArgs`\n\n"
f"{linkify_docs_urls(format_help(async_engine_args))}"
)
},
)
# CLI reference pages generated entirely from their parser: page -> (parser, JSON tip)
pages = {
"cli/serve.md": (create_parser(openai_cli_args.make_arg_parser), True),
"cli/chat.md": (create_parser(ChatCommand.add_cli_args), False),
"cli/complete.md": (create_parser(CompleteCommand.add_cli_args), False),
"cli/run-batch.md": (create_parser(openai_run_batch.make_arg_parser), True),
"cli/launch/render.md": (create_parser(RenderSubcommand.add_cli_args), True),
"cli/bench/latency.md": (create_parser(bench_latency.add_cli_args), True),
# URL kept as `mm_processor` for back-compat; command name is `mm-processor`
"cli/bench/mm_processor.md": (
create_parser(BenchmarkMMProcessorSubcommand.add_cli_args),
True,
),
"cli/bench/serve.md": (create_parser(bench_serve.add_cli_args), True),
"cli/bench/startup.md": (create_parser(bench_startup.add_cli_args), True),
"cli/bench/throughput.md": (create_parser(bench_throughput.add_cli_args), True),
"cli/bench/sweep/plot.md": (create_parser(bench_sweep_plot.add_cli_args), True),
"cli/bench/sweep/plot_pareto.md": (
create_parser(bench_sweep_plot_pareto.add_cli_args),
True,
),
"cli/bench/sweep/serve.md": (create_parser(bench_sweep_serve.add_cli_args), True),
"cli/bench/sweep/serve_workload.md": (
create_parser(bench_sweep_serve_workload.add_cli_args),
True,
),
"cli/bench/sweep/startup.md": (
create_parser(bench_sweep_startup.add_cli_args),
True,
),
}
# Command name for pages whose file stem differs (URL kept for back-compat).
COMMAND_NAMES = {"cli/bench/mm_processor.md": "mm-processor"}
for doc_path, (parser, json_tip) in pages.items():
segments = Path(doc_path).relative_to("cli").with_suffix("").parts
label = COMMAND_NAMES.get(doc_path, segments[-1])
command = " ".join([*segments[:-1], label])
# `title` frontmatter keeps the nav label to just this command's segment,
# while the H1 stays the full `vllm ...` command for the page heading.
content = f"---\ntitle: {label}\n---\n\n"
content += f"# vllm {command}\n\n"
if parser.description:
content += f"## Overview\n\n{parser.description}\n\n"
# Rendered above instead of at the top of the Arguments section
parser.description = None
if json_tip:
content += JSON_TIP
content += f"## Arguments\n\n{linkify_docs_urls(format_help(parser))}"
with mkdocs_gen_files.open(doc_path, "w") as f:
f.write(content)
logger.debug("CLI reference generated: %s", doc_path)
logger.info("Total argparse docs generated: %d", len(pages) + 2)
# --- Bare subcommand (group) pages -------------------------------------------
# Mirror `vllm <group> --help`: an overview plus a table of child subcommands,
# each linked to its reference page. Children are read from the CLI registries
# so the listing can never drift from the actual subcommands. Each page is the
# `README.md` of its command directory so it becomes that section's index and is
# picked up by the existing nav globs.
import_bench_subcommands() # populate BenchmarkSubcommandBase.__subclasses__()
bench_subcommands = BenchmarkSubcommandBase.__subclasses__()
bench_children = [(cmd.name, cmd.help) for cmd in bench_subcommands]
groups = {
"cli/bench/README.md": (BenchmarkSubcommand.help, bench_children),
"cli/launch/README.md": (
launch_description,
[(cmd.name, cmd.help) for cmd in LaunchSubcommandBase.__subclasses__()],
),
"cli/bench/sweep/README.md": (
dict(bench_children).get("sweep"),
[(args.parser_name, args.parser_help) for args, _ in sweep_subcommands],
),
}
# Doc paths that exist, so we only link a child that has a reference page.
existing_pages = set(pages) | set(groups)
def child_link(group_doc: str, name: str) -> str | None:
group_dir = Path(group_doc).parent # cli/bench/README.md -> cli/bench
for stem in (name, name.replace("-", "_")):
# A leaf page (bench/latency.md) or a nested group index (sweep/README.md)
for candidate in (group_dir / f"{stem}.md", group_dir / stem / "README.md"):
if candidate.as_posix() in existing_pages:
return candidate.relative_to(group_dir).as_posix()
return None
for doc_path, (overview, children) in groups.items():
title = "vllm " + Path(doc_path).parent.relative_to("cli").as_posix()
lines = [f"# {title.replace('/', ' ')}", ""]
if overview:
lines += ["## Overview", "", overview.strip(), ""]
lines += ["## Subcommands", "", "| Command | Description |", "| --- | --- |"]
for name, summary in children:
link = child_link(doc_path, name)
command = f"[`{name}`]({link})" if link else f"`{name}`"
lines.append(f"| {command} | {(summary or '').strip()} |")
with mkdocs_gen_files.open(doc_path, "w") as f:
f.write("\n".join(lines) + "\n")
logger.debug("CLI group reference generated: %s", doc_path)
logger.info("CLI group reference pages generated: %d", len(groups))
@@ -9,33 +9,28 @@ based on the checks in AttentionBackend.validate_configuration().
This approach avoids requiring CUDA/ROCm/GPU libraries to be installed.
When used as a pre-commit hook, this script receives filenames as arguments
and only runs the check if any of the relevant files were modified.
It runs as an mkdocs-gen-files script, so the page is generated at docs build
time rather than being committed to the repository.
"""
import argparse
import ast
import fnmatch
import logging
import sys
from collections.abc import Callable
from pathlib import Path
from typing import Any
sys.path.insert(0, str(Path(__file__).parent))
from generated_content import fill_markers # noqa: E402
logger = logging.getLogger("mkdocs")
# ---------------------------------------------------------------------------
# Constants and file paths
# ---------------------------------------------------------------------------
REPO_ROOT = Path(__file__).parent.parent.parent
RELEVANT_PATTERNS = [
"vllm/v1/attention/backends/*.py",
"vllm/v1/attention/backends/**/*.py",
"vllm/models/minimax_m3/common/sparse_attention.py",
"vllm/model_executor/layers/attention/mla_attention.py",
"vllm/platforms/cuda.py",
"tools/pre_commit/generate_attention_backend_docs.py",
"docs/design/attention_backends.md",
]
REPO_ROOT = Path(__file__).parent.parent.parent.parent
BACKENDS_DIR = REPO_ROOT / "vllm" / "v1" / "attention" / "backends"
REGISTRY_FILE = BACKENDS_DIR / "registry.py"
@@ -55,19 +50,6 @@ BACKEND_KV_DTYPE_EXCLUDES: dict[str, set[str]] = {
}
def is_relevant_file(filepath: str) -> bool:
"""Check if a file matches any of the relevant patterns."""
path = Path(filepath)
if path.is_absolute():
try:
path = path.relative_to(REPO_ROOT)
except ValueError:
return False
path_str = str(path)
return any(fnmatch.fnmatch(path_str, pattern) for pattern in RELEVANT_PATTERNS)
MLA_PREFILL_DIR = BACKENDS_DIR / "mla" / "prefill"
MLA_PREFILL_REGISTRY_FILE = MLA_PREFILL_DIR / "registry.py"
MLA_PREFILL_SELECTOR_FILE = MLA_PREFILL_DIR / "selector.py"
@@ -960,7 +942,7 @@ def analyze_backend(backend_name: str, class_path: str) -> dict[str, Any] | None
try:
tree = ast.parse(file_path.read_text())
except Exception as e:
print(f" Warning: Could not parse {file_path}: {e}", file=sys.stderr)
logger.warning("Could not parse %s: %s", file_path, e)
return None
class_name = class_path.rsplit(".", 1)[1]
@@ -1657,113 +1639,12 @@ def _render_table(
return lines
def generate_markdown_table(
backends: list[dict[str, Any]], title: str, is_mla_table: bool = False
) -> str:
"""Generate a titled markdown table from backend info."""
if not backends:
return f"## {title}\n\nNo backends found.\n"
has_versions = any(b.get("version") for b in backends)
columns = _build_columns(is_mla_table, has_versions)
lines = [f"## {title}", ""]
lines.extend(_render_table(columns, backends))
lines.append("")
return "\n".join(lines)
# ---------------------------------------------------------------------------
# Markdown section generators (usage, priority, legend, MLA)
# ---------------------------------------------------------------------------
def generate_usage_section() -> str:
"""Generate the usage documentation section."""
return """## Setting the Attention Backend
### Command Line
There are two ways to specify the backend from the command line:
**Option 1: Using `--attention-backend` (simple)**
```bash
vllm serve <model> --attention-backend FLASH_ATTN
```
**Option 2: Using `--attention-config.backend` / `-ac.backend` (structured config)**
```bash
# Dot notation
vllm serve <model> --attention-config.backend FLASH_ATTN
vllm serve <model> -ac.backend FLASH_ATTN
# JSON format
vllm serve <model> --attention-config '{"backend": "FLASH_ATTN"}'
vllm serve <model> -ac '{"backend": "FLASH_ATTN"}'
```
> **Note:** `--attention-backend` and `--attention-config.backend` are mutually
> exclusive. Use one or the other, not both.
### Python API
Use `AttentionConfig` with the `LLM` class:
```python
from vllm import LLM
from vllm.config import AttentionConfig
from vllm.v1.attention.backends.registry import AttentionBackendEnum
# Method 1: Using AttentionConfig with enum
llm = LLM(
model="Qwen/Qwen3-0.6B",
attention_config=AttentionConfig(backend=AttentionBackendEnum.FLASH_ATTN),
)
# Method 2: Using attention_backend parameter with string
llm = LLM(
model="Qwen/Qwen3-0.6B",
attention_backend="FLASH_ATTN",
)
```
## Backend Selection Behavior
### Manual Selection
When you explicitly set a backend via `--attention-backend` or `AttentionConfig`:
1. The backend is **validated** against your configuration (model dtype, head
size, compute capability, etc.)
2. If the backend **doesn't support** your configuration, an error is raised
with the specific reason
3. If valid, the backend is used
Example error when selecting an incompatible backend:
```text
ValueError: Selected backend FLASHMLA is not valid for this configuration.
Reason: ['compute capability not supported']
```
### Automatic Selection
When no backend is specified (the default):
1. vLLM iterates through backends in **priority order** (see tables below)
2. Each backend is validated against your configuration
3. The **first compatible backend** is selected
4. If no backend is compatible, an error is raised listing all backends and
their incompatibility reasons
"""
def _priority_table(
title: str,
backends: list[str],
annotations: dict[str, str] | None = None,
) -> list[str]:
"""Generate a priority table for a list of backends."""
"""Render a priority table for a list of backends."""
def _fmt(b: str) -> str:
suffix = annotations.get(b, "") if annotations else ""
@@ -1779,102 +1660,38 @@ def _priority_table(
]
def generate_priority_section(priorities: dict[str, list[str]]) -> str:
"""Generate the priority ranking section."""
lines = [
"## Backend Priority (CUDA)",
"",
"When no backend is explicitly selected, vLLM chooses the first",
"compatible backend from these priority-ordered lists.",
"",
"Priority is **1 = highest** (tried first).",
"",
"### Standard Attention (MHA, MQA, GQA)",
"",
]
sm100 = "Blackwell (SM 10.x)"
ampere = "Ampere/Hopper (SM 8.x-9.x)"
if "standard_sm100" in priorities:
lines.extend(_priority_table(sm100, priorities["standard_sm100"]))
if "standard_default" in priorities:
lines.extend(_priority_table(ampere, priorities["standard_default"]))
lines.extend(["### MLA Attention (DeepSeek-style)", ""])
mla_sm100_annotations = {
"FLASHINFER_MLA_SPARSE": "**\\***",
}
if "mla_sm100" in priorities:
lines.extend(
_priority_table(sm100, priorities["mla_sm100"], mla_sm100_annotations)
)
if "mla_default" in priorities:
lines.extend(_priority_table(ampere, priorities["mla_default"]))
if "mla_sm100" in priorities:
lines.append(
"> **\\*** For sparse MLA, FP8 KV cache always prefers "
"`FLASHINFER_MLA_SPARSE`. With BF16 KV cache, `FLASHINFER_MLA_SPARSE` "
"is preferred for low query-head counts (<= 16), while "
"`FLASHMLA_SPARSE` is preferred otherwise."
)
lines.append(">")
lines.append(
"> **Note:** ROCm and CPU platforms have their own selection logic. "
"See the platform-specific documentation for details."
)
lines.append("")
return "\n".join(lines)
_SM100 = "Blackwell (SM 10.x)"
_AMPERE = "Ampere/Hopper (SM 8.x-9.x)"
def generate_legend() -> str:
"""Generate a legend explaining the table columns."""
return """## Legend
| Column | Description |
| ------ | ----------- |
| **Dtypes** | Supported model data types (fp16, bf16, fp32) |
| **KV Dtypes** | Supported KV cache data types (`auto`, `fp8`, `fp8_e4m3`, etc.) |
| **Block Sizes** | Supported KV cache block sizes (%N means multiples of N) |
| **Head Sizes** | Supported attention head sizes |
| **Sink** | Attention sink support (for StreamingLLM) |
| **Non-Causal** | Non-causal (bidirectional) attention support for decoder models |
| **Sparse** | Sparse attention support (MLA only) |
| **MM Prefix** | Multimodal prefix full attention support |
| **DCP** | Decode Context Parallelism support (`--decode-context-parallel-size`) |
| **Attention Types** | Supported attention patterns (Decoder, Encoder, Enc-Dec) |
| **Compute Cap.** | Required CUDA compute capability (N/A for non-CUDA backends) |
**Symbols:** = Supported, = Not supported
"""
def generate_mla_section(
prefill_backends: list[dict[str, Any]],
decode_backends: list[dict[str, Any]],
v4_decode_backends: list[dict[str, Any]] | None = None,
def _priority_block(
priorities: dict[str, list[str]],
sm100_key: str,
default_key: str,
sm100_annotations: dict[str, str] | None = None,
) -> str:
"""Generate the complete MLA section with prefill and decode tables."""
"""Render whichever priority tables exist for one attention category."""
lines: list[str] = []
if sm100_key in priorities:
lines += _priority_table(_SM100, priorities[sm100_key], sm100_annotations)
if default_key in priorities:
lines += _priority_table(_AMPERE, priorities[default_key])
return "\n".join(lines).strip()
def _feature_table(backends: list[dict[str, Any]], is_mla: bool) -> str:
"""Render a backend feature table (header, separator, one row per backend)."""
has_versions = any(b.get("version") for b in backends)
columns = _build_columns(is_mla, has_versions)
return "\n".join(_render_table(columns, backends))
def _mla_prefill_table(prefill_backends: list[dict[str, Any]]) -> str:
"""Render the MLA prefill backend table."""
lines = [
"## MLA (Multi-head Latent Attention) Backends",
"",
"MLA uses separate backends for prefill and decode phases.",
"",
"### Prefill Backends",
"",
"To explicitly select a prefill backend, use",
"`-ac.mla_prefill_backend=<BACKEND>` (e.g., `FLASH_ATTN`, `FLASHINFER`).",
"Otherwise, the prefill backend is selected automatically at runtime based on",
"hardware and configuration.",
"",
"| Backend | Description | Dtypes | Compute Cap. | Notes |",
"| ------- | ----------- | ------ | ------------ | ----- |",
]
for backend in prefill_backends:
row = "| `{}`{} | {} | {} | {} | {} |".format(
backend["name"],
@@ -1885,87 +1702,21 @@ def generate_mla_section(
backend.get("notes", ""),
)
lines.append(row.replace(" ", " "))
lines.extend(
[
"",
"> **‡** Automatic selection tries FlashAttention first. On Blackwell",
"> (SM100), the fallback order is TRT-LLM Ragged, FlashInfer, then",
"> TokenSpeed MLA. On other GPUs, only FlashAttention is considered.",
"",
"### Decode Backends",
"",
"MLA decode backends are selected using the standard",
"`-ac.backend=<BACKEND>` argument (e.g., `FLASHMLA`, `TRITON_MLA`).",
"",
]
)
# Reuse data-driven table rendering for decode backends
columns = _build_columns(is_mla=True, has_versions=False)
lines.extend(_render_table(columns, decode_backends))
if v4_decode_backends:
lines.extend(
[
"",
"### DeepSeek V4 Decode Backends",
"",
"DeepSeek V4 sparse MLA uses its own decode backends, selected via",
"`--attention-backend=<BACKEND>` (e.g., `FLASHMLA_SPARSE_DSV4`,",
"`FLASHINFER_MLA_SPARSE_DSV4`). They share the V4 sparse-index",
"pipeline (compressor + SWA + indexer, 256-token blocks, head 512);",
"default on NVIDIA is `FLASHINFER_MLA_SPARSE_DSV4` on SM12x and",
"`FLASHMLA_SPARSE_DSV4` on other supported CUDA architectures.",
"",
]
)
lines.extend(_render_table(columns, v4_decode_backends))
lines.append("")
return "\n".join(lines)
def generate_minimax_section(backends: list[dict[str, Any]]) -> str:
"""Generate the MiniMax M3 sparse attention section."""
lines = [
"## MiniMax M3 Sparse Attention Backends",
"",
'Block-sparse GQA backend used by MiniMax M3 sparse ("lightning indexer")',
"layers. It is wired in directly by the model and is not part of the",
"automatic priority lists above. A lightning indexer scores KV blocks, the",
"top-k blocks (plus fixed init/local blocks) are selected, and attention",
"attends only to those blocks; index keys live in a separate side cache.",
"",
]
columns = _build_columns(is_mla=False, has_versions=False)
lines.extend(_render_table(columns, backends))
lines.append("")
return "\n".join(lines)
def build_blocks() -> dict[str, str]:
"""Build the generated table blocks keyed by their `gen:` marker name.
# ---------------------------------------------------------------------------
# Top-level orchestration
# ---------------------------------------------------------------------------
def generate_docs() -> str:
"""Generate the complete documentation."""
Only the tables are generated here; the surrounding prose lives in the
handwritten ``docs/design/attention_backends.md`` page.
"""
attention_backends_map = parse_registry()
# Parse priority lists from cuda.py
priorities = parse_cuda_priority_lists()
# Parse FlashAttention FA2/FA3 feature differences
fa_features = parse_flash_attn_features()
# Parse FlashInfer TRTLLM feature differences (native vs TRTLLM on Blackwell)
fi_features = parse_flashinfer_trtllm_features()
# Parse MLA prefill backends
mla_prefill_backends = parse_mla_prefill_backends()
# Collect backend info
all_backends = []
for backend_name, class_path in attention_backends_map.items():
if backend_name in SKIP_BACKENDS:
@@ -1973,17 +1724,14 @@ def generate_docs() -> str:
info = analyze_backend(backend_name, class_path)
if info:
all_backends.append(info)
# Expand backends into version variants
if fa_features:
all_backends = _expand_flash_attn_variants(all_backends, fa_features)
if fi_features:
all_backends = _expand_flashinfer_variants(all_backends, fi_features)
# DeepSeek V4 (*_DSV4) decode backends and MiniMax M3 sparse backends each
# get their own subsection rather than mixing into the main MLA / standard
# tables (the ROCm V4 backend isn't flagged is_mla by the AST heuristic, so
# filter purely on the name).
# DeepSeek V4 (*_DSV4) and MiniMax M3 sparse backends get their own tables
# rather than mixing into the main MLA / standard tables (the ROCm V4 backend
# isn't flagged is_mla by the AST heuristic, so filter purely on the name).
def _is_v4(b: dict[str, Any]) -> bool:
return b["name"].endswith("_DSV4")
@@ -1999,112 +1747,21 @@ def generate_docs() -> str:
if not b["is_mla"] and not _is_v4(b) and not _is_minimax(b)
]
# Generate documentation
script_path = "tools/pre_commit/generate_attention_backend_docs.py"
doc_lines = [
"# Attention Backend Feature Support",
"",
f"This document is auto-generated by `{script_path}`.",
"It shows the feature support for each registered attention backend",
"based on the checks in `AttentionBackend.validate_configuration()`.",
"",
"**Do not edit this file manually.** Run the following command to",
"regenerate it:",
"",
"```bash",
f"python {script_path}",
"```",
"",
]
# Add usage documentation
doc_lines.append(generate_usage_section())
# Add priority section
doc_lines.append(generate_priority_section(priorities))
# Add legend and feature tables
doc_lines.append(generate_legend())
standard_title = "Standard Attention (MHA, MQA, GQA) Backends"
doc_lines.append(
generate_markdown_table(non_mla_backends, standard_title, is_mla_table=False)
)
# Add footnotes for version/variant distinctions (in table order)
footnotes = []
if fi_features:
footnotes.append(
"> **†** FlashInfer Native is the regular FlashInfer path. XQA is the "
"SM90 decode path exposed through FlashInfer's TRTLLM decode API. "
"trtllm-gen is used on SM100 and supports sinks. Disable XQA/trtllm-gen "
"via `--attention-config.use_trtllm_attention=0`."
)
if fa_features:
footnotes.append(
"> **\\*** Specify the FlashAttention version via "
"`--attention-config.flash_attn_version=2`, `3`, or `4`. "
"Default is FA4 on SM100+ (Blackwell), FA3 on SM90 (Hopper), "
"FA2 otherwise."
)
if footnotes:
doc_lines.append("\n>\n".join(footnotes) + "\n")
# Add MiniMax M3 sparse section (separate category after standard GQA)
if minimax_backends:
doc_lines.append(generate_minimax_section(minimax_backends))
# Add MLA section with prefill and decode backends
doc_lines.append(
generate_mla_section(mla_prefill_backends, mla_backends, v4_decode_backends)
)
return "\n".join(doc_lines)
mla_sm100_annotations = {"FLASHINFER_MLA_SPARSE": "**\\***"}
return {
"priority-standard": _priority_block(
priorities, "standard_sm100", "standard_default"
),
"priority-mla": _priority_block(
priorities, "mla_sm100", "mla_default", mla_sm100_annotations
),
"table-standard": _feature_table(non_mla_backends, is_mla=False),
"table-minimax": _feature_table(minimax_backends, is_mla=False),
"table-mla-prefill": _mla_prefill_table(mla_prefill_backends),
"table-mla-decode": _feature_table(mla_backends, is_mla=True),
"table-mla-v4-decode": _feature_table(v4_decode_backends, is_mla=True),
}
def main():
parser = argparse.ArgumentParser(
description="Generate attention backend documentation table"
)
parser.add_argument(
"--output",
"-o",
type=str,
default=str(REPO_ROOT / "docs" / "design" / "attention_backends.md"),
help="Output file path (default: docs/design/attention_backends.md)",
)
parser.add_argument(
"--check",
action="store_true",
help="Check if the documentation is up to date (for pre-commit)",
)
parser.add_argument(
"files",
nargs="*",
help="Files to check (passed by pre-commit). If none are relevant, skip.",
)
args = parser.parse_args()
if args.files and not any(is_relevant_file(f) for f in args.files):
sys.exit(0)
output_path = Path(args.output)
new_content = generate_docs()
if args.check:
needs_update = (
not output_path.exists() or output_path.read_text() != new_content
)
if needs_update:
output_path.parent.mkdir(parents=True, exist_ok=True)
output_path.write_text(new_content)
print(f"🔄 Regenerated: {output_path}")
sys.exit(1)
print(f"✅ Up to date: {output_path}")
sys.exit(0)
output_path.parent.mkdir(parents=True, exist_ok=True)
output_path.write_text(new_content)
print(f"Generated: {output_path}")
if __name__ == "__main__":
main()
logger.info("Generating attention backend documentation")
fill_markers("design/attention_backends.md", build_blocks())
@@ -5,16 +5,15 @@ import logging
from dataclasses import dataclass
from functools import cached_property
from pathlib import Path
from typing import Literal
import mkdocs_awesome_nav.nav.directory as _nav_dir
import mkdocs_gen_files
import regex as re
logger = logging.getLogger("mkdocs")
ROOT_DIR = Path(__file__).parent.parent.parent.parent
ROOT_DIR_RELATIVE = "../../../../.."
EXAMPLE_DIR = ROOT_DIR / "examples"
EXAMPLE_DOC_DIR = ROOT_DIR / "docs/examples"
def title(text: str) -> str:
@@ -197,44 +196,38 @@ class Example:
return content
def on_startup(command: Literal["build", "gh-deploy", "serve"], dirty: bool):
# Monkey-patch dirname_to_title in awesome-nav so that sub-directory names are
# title-cased (e.g. "Offline Inference" instead of "Offline inference").
import mkdocs_awesome_nav.nav.directory as _nav_dir
# Monkey-patch dirname_to_title in awesome-nav so that sub-directory names are
# title-cased (e.g. "Offline Inference" instead of "Offline inference").
_nav_dir.dirname_to_title = title
logger.info("Generating example documentation")
logger.debug("Root directory: %s", ROOT_DIR.resolve())
logger.debug("Example directory: %s", EXAMPLE_DIR.resolve())
_nav_dir.dirname_to_title = title
logger.info("Generating example documentation")
logger.debug("Root directory: %s", ROOT_DIR.resolve())
logger.debug("Example directory: %s", EXAMPLE_DIR.resolve())
logger.debug("Example document directory: %s", EXAMPLE_DOC_DIR.resolve())
categories = sorted(
p for p in EXAMPLE_DIR.iterdir() if p.is_dir() and not p.name.startswith(".")
)
# Create the EXAMPLE_DOC_DIR if it doesn't exist
if not EXAMPLE_DOC_DIR.exists():
EXAMPLE_DOC_DIR.mkdir(parents=True)
examples = []
glob_patterns = ["*.py", "*.md", "*.sh"]
# Find categorised examples
for category in categories:
logger.info("Processing category: %s", category.stem)
globs = [category.glob(pattern) for pattern in glob_patterns]
for path in itertools.chain(*globs):
examples.append(Example(path, category.stem))
# Find examples in subdirectories
globs = [category.glob(f"*/{pattern}") for pattern in glob_patterns]
for path in itertools.chain(*globs):
examples.append(Example(path.parent, category.stem))
categories = sorted(p for p in EXAMPLE_DIR.iterdir() if p.is_dir())
examples = []
glob_patterns = ["*.py", "*.md", "*.sh"]
# Find categorised examples
for category in categories:
logger.info("Processing category: %s", category.stem)
globs = [category.glob(pattern) for pattern in glob_patterns]
for path in itertools.chain(*globs):
examples.append(Example(path, category.stem))
# Find examples in subdirectories
globs = [category.glob(f"*/{pattern}") for pattern in glob_patterns]
for path in itertools.chain(*globs):
examples.append(Example(path.parent, category.stem))
# Generate the example documentation
for example in sorted(examples, key=lambda e: e.path.stem):
example_name = f"{example.path.stem}.md"
doc_path = EXAMPLE_DOC_DIR / example.category / example_name
if not doc_path.parent.exists():
doc_path.parent.mkdir(parents=True)
# Specify encoding for building on Windows
with open(doc_path, "w+", encoding="utf-8") as f:
f.write(example.generate())
logger.debug("Example generated: %s", doc_path.relative_to(ROOT_DIR))
logger.info("Total examples generated: %d", len(examples))
# Generate the example documentation
for example in sorted(examples, key=lambda e: e.path.stem):
doc_path = f"examples/{example.category}/{example.path.stem}.md"
with mkdocs_gen_files.open(doc_path, "w") as f:
f.write(example.generate())
if example.main_file is not None:
# Point the edit button at the example's source file
edit_path = Path("..") / example.main_file.relative_to(ROOT_DIR)
mkdocs_gen_files.set_edit_path(doc_path, str(edit_path))
logger.debug("Example generated: %s", doc_path)
logger.info("Total examples generated: %d", len(examples))
@@ -2,27 +2,28 @@
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import ast
import logging
import sys
from pathlib import Path
from typing import Literal
sys.path.insert(0, str(Path(__file__).parent))
from generated_content import fill_markers # noqa: E402
logger = logging.getLogger("mkdocs")
ROOT_DIR = Path(__file__).parent.parent.parent.parent
DOCS_DIR = ROOT_DIR / "docs"
GENERATED_METRICS_DIR = DOCS_DIR / "generated" / "metrics"
# Files to scan for metric definitions - each will generate a separate table
# Files to scan for metric definitions - each fills a `gen:` marker in
# docs/usage/metrics.md with its table (the section heading and any preamble
# live in the tracked page next to the marker).
METRIC_SOURCE_FILES = [
{"path": "vllm/v1/metrics/loggers.py", "output": "general.inc.md"},
{
"path": "vllm/v1/spec_decode/metrics.py",
"output": "spec_decode.inc.md",
},
{"path": "vllm/v1/metrics/loggers.py", "key": "metrics-general"},
{"path": "vllm/v1/spec_decode/metrics.py", "key": "metrics-spec-decode"},
{
"path": "vllm/distributed/kv_transfer/kv_connector/v1/nixl/stats.py",
"output": "nixl_connector.inc.md",
"key": "metrics-nixl",
},
{"path": "vllm/v1/metrics/perf.py", "output": "perf.inc.md"},
{"path": "vllm/v1/metrics/perf.py", "key": "metrics-mfu"},
]
@@ -110,41 +111,27 @@ def generate_markdown_table(metrics: list[dict[str, str]]) -> str:
return "\n".join(lines) + "\n"
def on_startup(command: Literal["build", "gh-deploy", "serve"], dirty: bool):
"""Generate metrics documentation tables from source files."""
logger.info("Generating metrics documentation")
logger.info("Generating metrics documentation")
# Create generated directory if it doesn't exist
GENERATED_METRICS_DIR.mkdir(parents=True, exist_ok=True)
blocks = {}
total_metrics = 0
for source_config in METRIC_SOURCE_FILES:
source_path = source_config["path"]
total_metrics = 0
for source_config in METRIC_SOURCE_FILES:
source_path = source_config["path"]
output_file = source_config["output"]
filepath = ROOT_DIR / source_path
if not filepath.exists():
raise FileNotFoundError(f"Metrics source file not found: {filepath}")
filepath = ROOT_DIR / source_path
if not filepath.exists():
raise FileNotFoundError(f"Metrics source file not found: {filepath}")
logger.debug("Extracting metrics from: %s", source_path)
metrics = extract_metrics_from_file(filepath)
logger.debug("Found %d metrics in %s", len(metrics), source_path)
logger.debug("Extracting metrics from: %s", source_path)
metrics = extract_metrics_from_file(filepath)
logger.debug("Found %d metrics in %s", len(metrics), source_path)
blocks[source_config["key"]] = generate_markdown_table(metrics).strip()
total_metrics += len(metrics)
# Generate and write the markdown table for this source
table_content = generate_markdown_table(metrics)
output_path = GENERATED_METRICS_DIR / output_file
with open(output_path, "w", encoding="utf-8") as f:
f.write(table_content)
total_metrics += len(metrics)
logger.info(
"Generated metrics table: %s (%d metrics)",
output_path.relative_to(ROOT_DIR),
len(metrics),
)
logger.info(
"Total metrics generated: %d across %d files",
total_metrics,
len(METRIC_SOURCE_FILES),
)
fill_markers("usage/metrics.md", blocks)
logger.info(
"Total metrics generated: %d across %d files",
total_metrics,
len(METRIC_SOURCE_FILES),
)
@@ -0,0 +1,56 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Inline build-time generated content into existing docs pages.
Source pages mark where generated content goes with a snippet-style marker,
`--8<-- "gen:<key>"`, so the insertion point is explicit and readable. The
substitution happens here (at gen-files time, before mkdocs-gen-files shadows
the page), not via pymdownx.snippets, so the content can be generated at build
time without living in a real file on disk.
The `gen:` prefix keeps these markers distinct from real pymdownx.snippets
includes, and `fill_markers` fails loudly if a marker is missing or left behind
(pymdownx.snippets would otherwise silently drop an unsubstituted marker).
"""
from pathlib import Path
import mkdocs_gen_files
import regex as re
DOCS_DIR = Path(__file__).parent.parent.parent
_MARKER = '--8<-- "gen:{key}"'
_ANY_MARKER = re.compile(r'--8<-- "gen:[^"]*"')
def fill_markers(doc_path: str, blocks: dict[str, str]) -> None:
"""Replace `--8<-- "gen:<key>"` markers in a docs page with generated content.
Args:
doc_path: Docs-relative path of the source page to fill.
blocks: Mapping of marker key to the markdown to insert in its place.
Raises:
FileNotFoundError: If the source page does not exist.
ValueError: If an expected marker is missing, or any `gen:` marker is
left unsubstituted after filling.
"""
source = DOCS_DIR / doc_path
if not source.exists():
raise FileNotFoundError(f"Cannot fill markers in missing page: {doc_path}")
text = source.read_text()
for key, content in blocks.items():
marker = _MARKER.format(key=key)
if marker not in text:
raise ValueError(f"{doc_path}: missing marker {marker}")
text = text.replace(marker, content)
if leftover := _ANY_MARKER.search(text):
raise ValueError(f"{doc_path}: unsubstituted marker {leftover.group()}")
with mkdocs_gen_files.open(doc_path, "w") as f:
f.write(text)
# Keep the edit button pointing at the real source page
mkdocs_gen_files.set_edit_path(doc_path, doc_path)
+39 -4
View File
@@ -19,6 +19,7 @@ The on_page_markdown hook passes the current page context to the preprocessor be
each page is converted.
"""
import posixpath
from pathlib import Path
import regex as re
@@ -38,18 +39,22 @@ TITLE = r"(?P<title>[^\[\]<>]+?)"
REPO = r"(?P<repo>.+?/.+?)"
TYPE = r"(?P<type>issues|pull|projects)"
NUMBER = r"(?P<number>\d+)"
VERSION = r"[^/\s]+"
PATH = r"(?P<path>[^\s]+?)"
FRAGMENT = r"(?P<fragment>#[^\s]+)?"
URL = f"https://github.com/{REPO}/{TYPE}/{NUMBER}{FRAGMENT}"
URL_GITHUB = f"https://github.com/{REPO}/{TYPE}/{NUMBER}{FRAGMENT}"
RELATIVE = rf"(?!(https?|ftp)://|#){PATH}{FRAGMENT}"
URL_DOCS = f"https://docs.vllm.ai/en/{VERSION}/{PATH}{FRAGMENT}"
# Common titles to use for GitHub links when none is provided in the link.
TITLES = {"issues": "Issue ", "pull": "Pull Request ", "projects": "Project "}
# Regex to match GitHub issue, PR, and project links with optional titles.
github_link = re.compile(rf"(\[{TITLE}\]\(|<){URL}(\)|>)")
github_link = re.compile(rf"(\[{TITLE}\]\(|<){URL_GITHUB}(\)|>)")
# Regex to match relative file links with optional titles.
relative_link = re.compile(rf"\[{TITLE}\]\({RELATIVE}\)")
# Regex to match absolute docs.vllm.ai links (should only exist in CLI).
docs_link = re.compile(rf"\[{TITLE}\]\({URL_DOCS}\)")
class UrlSchemesPreprocessor(Preprocessor):
@@ -61,7 +66,8 @@ class UrlSchemesPreprocessor(Preprocessor):
def run(self, lines):
page = self.ext.page
if page is None or getattr(page.file, "abs_src_path", None) is None:
files = self.ext.files
if page is None:
return lines
def replace_relative_link(match: re.Match) -> str:
@@ -70,7 +76,7 @@ class UrlSchemesPreprocessor(Preprocessor):
"""
title = match.group("title")
path = match.group("path")
path = (Path(page.file.abs_src_path).parent / path).resolve()
path = ((DOC_DIR / page.file.src_uri).parent / path).resolve()
fragment = match.group("fragment") or ""
# Check if the path exists and is outside the docs dir
@@ -105,9 +111,36 @@ class UrlSchemesPreprocessor(Preprocessor):
url = f"https://github.com/{repo}/{type}/{number}{fragment}"
return f"[{gh_icon} {title}]({url})"
def replace_docs_link(match: re.Match) -> str:
"""Rewrite absolute docs.vllm.ai links as doc-relative links."""
title = match.group("title")
path = match.group("path").rstrip("/")
fragment = match.group("fragment") or ""
# vllm.config.<Class> API reference -> mkdocstrings cross-reference
if path == "api/vllm/config" and re.fullmatch(
r"#vllm\.config\.\w+", fragment
):
ident = fragment[1:]
return f"[`{ident}`][{ident}]"
# Other docs pages -> link relative to the current page, but only
# when the target is a known docs page (real or generated); leave
# unknown/external URLs untouched. This is correct even when the same
# docstring is also rendered on its API reference page.
src = f"{path.removesuffix('.html')}.md"
if files.get_file_from_path(src) is None:
return match.group(0)
rel = posixpath.relpath(src, posixpath.dirname(page.file.src_uri))
# Auto-wrapped bare URLs use the URL as their title; make it readable.
if title.startswith("http"):
title = path.removesuffix(".html")
return f"[{title}]({rel}{fragment})"
markdown = "\n".join(lines)
markdown = github_link.sub(replace_github_link, markdown)
markdown = relative_link.sub(replace_relative_link, markdown)
markdown = docs_link.sub(replace_docs_link, markdown)
return markdown.split("\n")
@@ -116,6 +149,7 @@ class UrlSchemesExtension(Extension):
def __init__(self, **kwargs):
self.page = None
self.files = None
super().__init__(**kwargs)
def extendMarkdown(self, md):
@@ -138,4 +172,5 @@ def on_page_markdown(
) -> str:
"""Pass the current page context to the preprocessor."""
_ext.page = page
_ext.files = files
return markdown
+4 -4
View File
@@ -19,7 +19,7 @@ vLLM also supports model implementations that are available in Transformers. We
Currently, the Transformers modeling backend works for the following:
- Modalities: embedding models, language models and vision-language models*
- Modalities: embedding models, language models, vision-language models* and audio-language models
- Architectures: encoder-only, decoder-only, mixture-of-experts
- Attention types: full attention and/or sliding attention
@@ -427,7 +427,6 @@ th {
| `OlmoeForCausalLM` | OLMoE | `allenai/OLMoE-1B-7B-0924`, `allenai/OLMoE-1B-7B-0924-Instruct`, etc. | | ✅︎ |
| `OPTForCausalLM` | OPT, OPT-IML | `facebook/opt-66b`, `facebook/opt-iml-max-30b`, etc. | ✅︎ | ✅︎ |
| `OrionForCausalLM` | Orion | `OrionStarAI/Orion-14B-Base`, `OrionStarAI/Orion-14B-Chat`, etc. | | ✅︎ |
| `OuroForCausalLM` | ouro | `ByteDance/Ouro-1.4B`, `ByteDance/Ouro-2.6B`, etc. | ✅︎ | |
| `PanguEmbeddedForCausalLM` | openPangu-Embedded-7B | `FreedomIntelligence/openPangu-Embedded-7B-V1.1` | ✅︎ | ✅︎ |
| `PanguProMoEV2ForCausalLM` | openpangu-pro-moe-v2 | | ✅︎ | ✅︎ |
| `PanguUltraMoEForCausalLM` | openpangu-ultra-moe-718b-model | `FreedomIntelligence/openPangu-Ultra-MoE-718B-V1.1` | ✅︎ | ✅︎ |
@@ -465,6 +464,7 @@ Some models are supported only via the [Transformers modeling backend](#transfor
| `Olmo2ForCausalLM` | OLMo2 | `allenai/OLMo-2-0425-1B`, etc. | ✅︎ | ✅︎ |
| `SmolLM3ForCausalLM` | SmolLM3 | `HuggingFaceTB/SmolLM3-3B` | ✅︎ | ✅︎ |
| `Starcoder2ForCausalLM` | Starcoder2 | `bigcode/starcoder2-3b`, `bigcode/starcoder2-7b`, `bigcode/starcoder2-15b`, etc. | ✅︎ | ✅︎ |
| `VaultGemmaForCausalLM` | VaultGemma | `google/vaultgemma-1b` | ✅︎ | ✅︎ |
!!! note
Currently, the ROCm version of vLLM supports Mistral and Mixtral only for context lengths up to 4096.
@@ -547,7 +547,6 @@ These models primarily accept the [`LLM.generate`](./generative_models.md#llmgen
| `KeyeVL1_5ForConditionalGeneration` | Keye-VL-1_5-8B | T + I<sup>E+</sup> + V<sup>E+</sup> | `Kwai-Keye/Keye-VL-1_5-8B` | ✅︎ | ✅︎ |
| `KimiAudioForConditionalGeneration` | Kimi-Audio | T + A<sup>+</sup> | `moonshotai/Kimi-Audio-7B-Instruct` | | ✅︎ |
| `KimiK25ForConditionalGeneration` | Kimi-K2.5 | T + I<sup>+</sup> | `moonshotai/Kimi-K2.5` | | ✅︎ |
| `KimiK3ForConditionalGeneration` | Kimi-K3 | T + I<sup>+</sup> | `moonshotai/Kimi-K3` | | ✅︎ |
| `KimiVLForConditionalGeneration` | Kimi-VL-A3B-Instruct, Kimi-VL-A3B-Thinking | T + I<sup>+</sup> | `moonshotai/Kimi-VL-A3B-Instruct`, `moonshotai/Kimi-VL-A3B-Thinking` | | ✅︎ |
| `LightOnOCRForConditionalGeneration` | LightOnOCR-1B | T + I<sup>+</sup> | `lightonai/LightOnOCR-1B`, etc | ✅︎ | ✅︎ |
| `Lfm2VlForConditionalGeneration` | LFM2-VL | T + I<sup>+</sup> | `LiquidAI/LFM2-VL-450M`, `LiquidAI/LFM2-VL-3B`, `LiquidAI/LFM2-VL-8B-A1B`, etc. | ✅︎ | ✅︎ |
@@ -608,7 +607,8 @@ Some models are supported only via the [Transformers modeling backend](#transfor
| Architecture | Models | Inputs | Example HF Models | [LoRA](../features/lora.md) | [PP](../serving/parallelism_scaling.md) |
| ------------ | ------ | ------ | ----------------- | --------------------------- | --------------------------------------- |
| `Emu3ForConditionalGeneration` | Emu3 | T + I | `BAAI/Emu3-Chat-hf` | ✅︎ | ✅︎ |
| `Emu3ForConditionalGeneration` | Emu3 | T + I<sup>+</sup> | `BAAI/Emu3-Chat-hf` | ✅︎ | ✅︎ |
| `VibeVoiceAsrForConditionalGeneration` | VibeVoice-ASR | T + A<sup>+</sup> | `microsoft/VibeVoice-ASR-HF` | ✅︎ | ✅︎ |
<sup>^</sup> You need to set the architecture name via `--hf-overrides` to match the one in vLLM.</br>
<sup>E</sup> Pre-computed embeddings can be inputted for this modality.</br>
+1 -1
View File
@@ -28,7 +28,7 @@ DOCS_PATHS=(
docs/ # Actual docs content
examples/ # Examples are rendered in docs
vllm/ # API & CLI reference
requirements/test/cuda.txt # CLI reference (see docs/mkdocs/hooks/generate_argparse.py)
requirements/test/cuda.txt # CLI reference (see docs/mkdocs/gen_files/generate_argparse.py)
mkdocs.yaml # Affects build process
.readthedocs.yaml # Affects build process
requirements/docs.txt # Affects build process
+4 -4
View File
@@ -35,21 +35,21 @@ The following metrics are exposed:
## General Metrics
--8<-- "docs/generated/metrics/general.inc.md"
--8<-- "gen:metrics-general"
## Speculative Decoding Metrics
--8<-- "docs/generated/metrics/spec_decode.inc.md"
--8<-- "gen:metrics-spec-decode"
## NIXL KV Connector Metrics
--8<-- "docs/generated/metrics/nixl_connector.inc.md"
--8<-- "gen:metrics-nixl"
## Model Flops Utilization (MFU) Performance Metrics
These metrics are available via `--enable-mfu-metrics`:
--8<-- "docs/generated/metrics/perf.inc.md"
--8<-- "gen:metrics-mfu"
## Deprecation Policy
@@ -503,6 +503,45 @@ def run_gemma3n(questions: list[str], modality: str) -> ModelRequestData:
)
# Gemma 4
def run_gemma4(questions: list[str], modality: str) -> ModelRequestData:
assert modality in ("image", "video")
model_name = "google/gemma-4-31B-it"
# NOTE: Gemma-4-31B is a large model. Users running into Out-Of-Memory (OOM)
# errors might need to set `tensor_parallel_size` to > 1.
engine_args = EngineArgs(
model=model_name,
max_model_len=4096,
max_num_seqs=2,
limit_mm_per_prompt={modality: 1},
)
if modality == "image":
prompts = [
(
"<bos><start_of_turn>user\n"
f"<|image|>\n{question}<end_of_turn>\n"
"<start_of_turn>model\n"
)
for question in questions
]
else: # video
prompts = [
(
"<bos><start_of_turn>user\n"
f"<|video|>\n{question}<end_of_turn>\n"
"<start_of_turn>model\n"
)
for question in questions
]
return ModelRequestData(
engine_args=engine_args,
prompts=prompts,
)
# GLM-4v
def run_glm4v(questions: list[str], modality: str) -> ModelRequestData:
assert modality == "image"
@@ -2303,6 +2342,7 @@ model_example_map = {
"exaone4_5": run_exaone4_5,
"gemma3": run_gemma3,
"gemma3n": run_gemma3n,
"gemma4": run_gemma4,
"glm4v": run_glm4v,
"glm4_1v": run_glm4_1v,
"glm4_5v": run_glm4_5v,
@@ -2374,6 +2414,7 @@ MODELS_NEED_VIDEO_METADATA = [
MODELS_SUPPORT_VIT_CUDA_GRAPH = [
"llama4",
"gemma4",
"qwen2_vl",
"qwen2_5_vl",
"qwen3_vl",
+7 -10
View File
@@ -3,7 +3,6 @@ site_url: !ENV READTHEDOCS_CANONICAL_URL
repo_url: https://github.com/vllm-project/vllm
edit_uri: edit/main/docs/
exclude_docs: |
argparse
*.inc.md
*.template.md
theme:
@@ -50,24 +49,22 @@ theme:
hooks:
- docs/mkdocs/hooks/remove_announcement.py
- docs/mkdocs/hooks/generate_examples.py
- docs/mkdocs/hooks/generate_argparse.py
- docs/mkdocs/hooks/generate_metrics.py
- docs/mkdocs/hooks/url_schemes.py
- docs/mkdocs/hooks/autoref_code.py
plugins:
- meta
- search
- gen-files:
scripts:
- docs/mkdocs/gen_files/generate_examples.py
- docs/mkdocs/gen_files/generate_argparse.py
- docs/mkdocs/gen_files/generate_metrics.py
- docs/mkdocs/gen_files/generate_attention_backends.py
- autorefs
- awesome-nav
- glightbox
- git-revision-date-localized:
# exclude autogenerated files
exclude:
- api/*
- examples/*
- generated/*
- git-revision-date-localized
- minify:
minify_html: true
minify_js: true
+1 -4
View File
@@ -17,7 +17,7 @@ PyNvVideoCodec==2.0.4
flashinfer-python==0.6.15.post1
flashinfer-cubin==0.6.15.post1
apache-tvm-ffi==0.1.10
tilelang==0.1.12
tilelang==0.1.9
nvidia-cudnn-frontend>=1.19.1
# Required for LLM_NVTX_SCOPES_FOR_PROFILING=1
nvtx==0.2.15
@@ -33,6 +33,3 @@ tokenspeed-mla==0.1.8; platform_system == "Linux"
# Humming kernels for quantization gemm
humming-kernels[cu13]==0.1.10
# KDA
flash-linear-attention==0.5.0
+2 -1
View File
@@ -2220,7 +2220,7 @@ checksum = "11d3d7f243d5c5a8b9bb5d6dd2b1602c0cb0b9db1621bafc7ed66e35ff9fe092"
[[package]]
name = "llm-multimodal"
version = "1.7.1"
source = "git+ssh://git@github.com/Inferact/llm-multimodal-internal.git?branch=k3-image#ceec43ec6beea5812d5ae59af712d9ff616496ef"
source = "git+https://github.com/smg-project/llm-multimodal?rev=5390032d6dc8a3e6fdc83acd320260367eb4b9b5#5390032d6dc8a3e6fdc83acd320260367eb4b9b5"
dependencies = [
"anyhow",
"base64 0.22.1",
@@ -2235,6 +2235,7 @@ dependencies = [
"once_cell",
"pkg-config",
"rayon",
"realfft",
"reqwest 0.13.4",
"rustfft",
"serde",
+1 -1
View File
@@ -58,7 +58,7 @@ indexmap = "2.13.0"
indicatif = "0.18.4"
itertools = "0.14.0"
libc = "0.2.177"
llm-multimodal = { git = "ssh://git@github.com/Inferact/llm-multimodal-internal.git", branch = "k3-image", default-features = false, features = ["native-tls"] }
llm-multimodal = { git = "https://github.com/smg-project/llm-multimodal", rev = "5390032d6dc8a3e6fdc83acd320260367eb4b9b5", default-features = false, features = ["native-tls"] }
mimalloc = "0.1.52"
minijinja = { version = "2.0", features = ["unstable_machinery", "json", "builtins", "loader", "loop_controls", "preserve_order"] }
minijinja-contrib = { version = "2.0", features = ["pycompat"] }
+53
View File
@@ -0,0 +1,53 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
syntax = "proto3";
package vllm;
service Control {
rpc GetServerInfo (GetServerInfoRequest) returns (ServerInfo) {}
rpc GetModelInfo (GetModelInfoRequest) returns (ModelInfo) {}
rpc Abort (AbortRequest) returns (AbortResponse) {}
}
message GetServerInfoRequest {}
message ServerInfo {
string engine_version = 1;
string api_version = 2;
string instance_id = 3;
ParallelismInfo parallelism = 4;
uint32 max_model_len = 5;
uint32 kv_block_size = 6;
uint64 total_kv_blocks = 7;
uint64 max_running_requests = 8;
uint64 max_batched_tokens = 9;
}
message ParallelismInfo {
uint32 tensor_parallel_size = 1;
uint32 pipeline_parallel_size = 2;
uint32 data_parallel_size = 3;
uint32 data_parallel_rank = 4;
uint32 decode_context_parallel_size = 5;
}
message GetModelInfoRequest {}
message ModelInfo {
string model_id = 1;
string served_model_name = 2;
repeated string served_model_aliases = 3;
bool supports_text_input = 20;
bool supports_token_ids_input = 21;
bool supports_multimodal = 23;
string reasoning_parser = 24;
string tool_call_parser = 25;
}
message AbortRequest {
repeated string request_ids = 1;
}
message AbortResponse {}
@@ -7,17 +7,13 @@ package vllm;
import "google/protobuf/struct.proto";
service Generate {
service Inference {
// Generates text given a prompt
rpc Generate (GenerateRequest) returns (GenerateResponse) {}
// Generates text given a prompt, streaming the outputs
rpc GenerateStream (GenerateRequest) returns (stream GenerateResponse) {}
}
service Control {
rpc Abort (AbortRequest) returns (AbortResponse) {}
}
// ======================================================================================
// Generate Request
// ======================================================================================
@@ -204,13 +200,3 @@ message CandidateTokenInfo {
message TokenIds {
repeated uint32 ids = 1;
}
// ======================================================================================
// Control
// ======================================================================================
message AbortRequest {
repeated string request_ids = 1;
}
message AbortResponse {}
+1 -2
View File
@@ -20,7 +20,7 @@ use crate::output::{
use crate::renderer::hf::{HfChatRenderer, MultimodalRenderInfo};
use crate::renderer::{
DeepSeekV4ChatRenderer, DeepSeekV32ChatRenderer, DynChatRenderer, HarmonyChatRenderer,
InklingChatRenderer, KimiK3ChatRenderer,
InklingChatRenderer,
};
use crate::request::ChatRequest;
use crate::{DynChatOutputProcessor, RendererSelection};
@@ -73,7 +73,6 @@ impl HfChatBackend {
RendererSelection::DeepSeekV4 => Arc::new(DeepSeekV4ChatRenderer::new()),
RendererSelection::Harmony => Arc::new(HarmonyChatRenderer::new()?),
RendererSelection::Inkling => Arc::new(InklingChatRenderer::new(tokenizer.clone())?),
RendererSelection::KimiK3 => Arc::new(KimiK3ChatRenderer::new()),
};
info!(
+30 -10
View File
@@ -33,8 +33,7 @@ pub use parser::tool::{ToolParser, ToolParserError, ToolParserFactory};
pub use renderer::hf::ChatTemplateContentFormatOption;
pub use renderer::{
ChatRenderer, DeepSeekV4ChatRenderer, DeepSeekV32ChatRenderer, DynChatRenderer,
HarmonyChatRenderer, InklingChatRenderer, KimiK3ChatRenderer, RenderedPrompt,
RendererSelection,
HarmonyChatRenderer, InklingChatRenderer, RenderedPrompt, RendererSelection,
};
pub use request::{
ChatContent, ChatContentPart, ChatMessage, ChatOptions, ChatRequest, ChatRole, ChatTool,
@@ -256,6 +255,33 @@ impl ChatLlm {
self.text.engine_core_client()
}
/// Whether the loaded backend has a registered multimodal processor.
pub fn supports_multimodal(&self) -> bool {
self.processor.backend.multimodal_model_info().is_some()
}
/// Effective tool-call parser name for this model, if parsing is enabled.
pub fn tool_call_parser_name(&self) -> Option<&str> {
match &self.tool_call_parser {
ParserSelection::Auto => {
ToolParserFactory::global().resolve_name_for_model(self.model_id())
}
ParserSelection::None => None,
ParserSelection::Explicit(name) => Some(name),
}
}
/// Effective reasoning parser name for this model, if parsing is enabled.
pub fn reasoning_parser_name(&self) -> Option<&str> {
match &self.reasoning_parser {
ParserSelection::Auto => {
ReasoningParserFactory::global().resolve_name_for_model(self.model_id())
}
ParserSelection::None => None,
ParserSelection::Explicit(name) => Some(name),
}
}
/// Render, tokenize, and submit one chat request.
pub async fn chat(&self, request: ChatRequest) -> Result<ChatEventStream> {
let (text_request, output_processor) = self
@@ -327,12 +353,6 @@ mod tests {
.unwrap();
}
#[test]
fn validate_parser_overrides_accepts_explicit_kimi_k3() {
let selection = ParserSelection::Explicit("kimi_k3".to_string());
validate_parser_overrides(&selection, &selection).unwrap();
}
#[test]
fn validate_parser_overrides_accepts_auto_and_none() {
validate_parser_overrides(&ParserSelection::Auto, &ParserSelection::None).unwrap();
@@ -346,7 +366,7 @@ mod tests {
)
.unwrap_err();
expect_test::expect!["tool parser `definitely_missing_tool_parser` is not registered (choose from: deepseek_v3, deepseek_v31, deepseek_v32, deepseek_v4, gemma4, glm45, glm47, granite4, hermes, hy_v3, inkling, internlm, kimi_k2, kimi_k3, llama3_json, llama4_json, minimax_m2, minimax_m3, mistral, phi4_mini_json, qwen3_coder, qwen3_xml, seed_oss)"].assert_eq(&error.to_report_string());
expect_test::expect!["tool parser `definitely_missing_tool_parser` is not registered (choose from: deepseek_v3, deepseek_v31, deepseek_v32, deepseek_v4, gemma4, glm45, glm47, granite4, hermes, hy_v3, inkling, internlm, kimi_k2, llama3_json, llama4_json, minimax_m2, minimax_m3, mistral, phi4_mini_json, qwen3_coder, qwen3_xml, seed_oss)"].assert_eq(&error.to_report_string());
}
#[test]
@@ -357,6 +377,6 @@ mod tests {
)
.unwrap_err();
expect_test::expect!["reasoning parser `definitely_missing_reasoning_parser` is not registered (choose from: cohere_cmd, deepseek_r1, deepseek_v3, deepseek_v4, gemma4, glm45, inkling, kimi, kimi_k2, kimi_k3, minimax_m2, minimax_m3, nemotron_v3, qwen3, seed_oss, step3, step3p5)"].assert_eq(&error.to_report_string());
expect_test::expect!["reasoning parser `definitely_missing_reasoning_parser` is not registered (choose from: cohere_cmd, deepseek_r1, deepseek_v3, deepseek_v4, gemma4, glm45, inkling, kimi, kimi_k2, minimax_m2, minimax_m3, nemotron_v3, qwen3, seed_oss, step3, step3p5)"].assert_eq(&error.to_report_string());
}
}

Some files were not shown because too many files have changed in this diff Show More