diff --git a/QSRT_PR227_PR238_MANIFEST.sha256 b/QSRT_PR227_PR238_MANIFEST.sha256 new file mode 100644 index 000000000..83440fefc --- /dev/null +++ b/QSRT_PR227_PR238_MANIFEST.sha256 @@ -0,0 +1,243 @@ +fade99495b02074246288e0ce5255bd14ce17a72f5a45c1dce7ce3638879483c b12x/__init__.py +bd6f8f4a0b0261c47788779104c21db1f6a1c7acf92297d7637dfe89370ce208 b12x/_lib/__init__.py +441b3974b84b42c4d9c40d5122f896c9ec69a48561b71579d5df3d20a990b7bf b12x/_lib/compiler.py +f3ea177103652486c82225dc9b91a8caa718f9345e9289c3bfc9e19e1bf42203 b12x/_lib/dense_gemm.py +2a16f399049476f278df910c4ed7d3baf4357caf4c34662f3568cf179dda7001 b12x/_lib/dense_gemm_mxfp6.py +5e7af97f1427917b8e840040b35b558e64882be9c0404e999d5bd0e1680c069d b12x/_lib/env.py +986f24c90f257d9ff9ebcb4debcf3d313e0fcec3a86b0300efe2c0ab008f2598 b12x/_lib/fp6.py +5bc885d507a66d61df73fa5dcacea201499dc91b727a0cffa74227fb3613e1b4 b12x/_lib/gating.py +e77e0b076dde663f395bd646e4e9232c6cf08c551b7a28a5653bf9d9bfa59122 b12x/_lib/intrinsics.py +85b5ee925cd2d9a055c386141ab88cb6e119481be46a46c831f9d3c5d2098e6b b12x/_lib/meta.py +bf66a3164d078ed176b08b0765a6ebb999f843c6c5bcac7fd6b6e6ae50310525 b12x/_lib/quant/__init__.py +8dd5ad7820217c0762df63d9c496dbd55c3b30c61e0ddccac2dce893ac9cf3c2 b12x/_lib/quant/mxfp8_rows.py +9087a6332f6bced9361797a8624e947a4abd110aaf627a186919f9c3bf107694 b12x/_lib/quant/sqg_e4m3.py +ad2311c45bd5318737b9b73a8046178943609f5a11428f7c3d7ffd3d67335bdf b12x/_lib/quant/sqg_fp16_d3l.py +0b3ab3370bca79f1b7bd28da228e4a51649673bbc3232f47486c67789633188d b12x/_lib/quant/x4t_scales.py +db9655ad961af80bc9a7b1c72dd1d22fe2a8d49c7e11f0e8abdb281a7b92617c b12x/_lib/runtime_control.py +53ce88b19bbb049e4d328e18f3a181e96595dfffc23279b87819bab70f852a9a b12x/_lib/runtime_patches.py +dbec8827bc4769d72a13e151ef076d86e1c7fa85a90994510d0491bf92545110 b12x/_lib/scratch.py +ae5f9485752b08d1f9d2b212de173217fbd56ecd776e836a8f9f2def6547102c b12x/_lib/scratch_layout.py +a5ea3a8925187ca5a9100c472ec59698bf801a18ac36672d5d11d9558a5aff4e b12x/_lib/smem.py +b72bb74c309357cb05b6d6ae012152c396fc51802316c25bdc2ffced0ff7a1dd b12x/_lib/utils.py +f06f026c44ba142f70a66b7c58303da49ca946555264fafd8366888d40a0fa1d b12x/attention/__init__.py +9292af7f053f094d2e25a1a7fcd7ac9eaf0cf937c5dfba48d794f0be33e1811b b12x/attention/_shared/__init__.py +8c001d985d88e4aa3f732d8e0d649c09505814dc77088921bf8f4f7d36f45e08 b12x/attention/_shared/contiguous/README.md +3c8bd108b3b266711bfe878eeccf250ea91b1e025262d84a2d379ff3ffb2527d b12x/attention/_shared/contiguous/__init__.py +447e7d8ef18ba3ea0fe3b918c0cd842574924648823a63c914f0aad1dfd9f18e b12x/attention/_shared/contiguous/api.py +6abc7e4a3b1ec61fc6e2dbb2f4005c5b91b71cdaf654fcdd4596d34535c020d2 b12x/attention/_shared/contiguous/block_info.py +200cdcbf7c13efd80fa40f4fdeb21d30c62e14e73ae217d1983f06d451684fb7 b12x/attention/_shared/contiguous/cute_dsl_utils.py +2629b36c5f0a11ba621db32e6cb8dd13cae75738663845aeebaa7a27b62d97c2 b12x/attention/_shared/contiguous/forward.py +03ffc597d4e643b8a3193aafb1a92df878bca845d035f5f366c7ca7686cf9647 b12x/attention/_shared/contiguous/layout_utils.py +c9e2631ffa54b683c31319acc094cef9adf0878f723f285002dd159be4203608 b12x/attention/_shared/contiguous/mask.py +76efb3ef35726da1a720a91dfaad8d2bc2b25798b5327b144b89a56231588d1b b12x/attention/_shared/contiguous/named_barrier.py +e26431b5a699e233c21305837475fca860387f9c8279f7a566978bb7f54619d3 b12x/attention/_shared/contiguous/pack_gqa.py +a8bd38c73a07e2ac58a65562089535c1da122832c6583d4fe0afeaf7999877a0 b12x/attention/_shared/contiguous/seqlen_info.py +c749212e36804661deae982f3889447c425b549f1aaf2700811e398b78ce15a0 b12x/attention/_shared/contiguous/softmax.py +8bccf6360b2cf13a4f15c7adf5ec1c919532335e789a0d3ce75910752a2a5ca5 b12x/attention/_shared/contiguous/tile_scheduler.py +370fc36e08e87c52093d0e58e218fea3bf3b5765d9ca3b77cc50dacd7a64edf8 b12x/attention/_shared/cute/__init__.py +a1ef9f60879bb0a38efb9bf01ea6266d0ef3071d3e114a7a6bab4c24fb112c23 b12x/attention/_shared/cute/copy.py +0d88ddb09c22737b39a625c3606f6a4b364c1250222faef9b3e36e4c9e4a7cbf b12x/attention/_shared/cute/ops.py +0e67505f62a96fa7f005e0206f9a3fc08b1fef1507dcbb59b344615b35f3b0c1 b12x/attention/_shared/cute/pipeline.py +7554405c4f6c69bda1645acf08c52ec18c0c1324fe3fb43894611b37a8c4c51d b12x/attention/_shared/mla/__init__.py +b1605b821bec3ec54f746cdb3208dd6e3d2e6a5f290e7def6328e787f2a19f5f b12x/attention/_shared/mla/api.py +79d792edbe5a85db7c15f85067346a6361ebca326ec91c4aa9fc2b82f5f6fa04 b12x/attention/_shared/mla/compressed_api.py +e33fd05486cd704bd003a65b06141851f94a2ac8c8f57d1bec8cf50f23251e6f b12x/attention/_shared/mla/compressed_config.py +6fa1fe3f1213d428b5718256c52fe2e0ac4ed4a65811a5fccfb808a6e2ac2321 b12x/attention/_shared/mla/compressed_reference.py +8768f329080496d46f420bb2791f1b84aac6cf21a69acf5d49e52b7d9f68f247 b12x/attention/_shared/mla/decode_math.py +f7c5b80e53fc6bdd5fb9188654e40250650d1ac9b1f4efd0e3634bb99a21a9d3 b12x/attention/_shared/mla/io.py +aa9b60d699da6e2a132d47cd6cdb7fb9602add2a24151e5909305b305eb29237 b12x/attention/_shared/mla/io_mg.py +dd782340f0ad33a06263119f5f334e5ae43a54d5737fb7a24ce89c0c7879bd26 b12x/attention/_shared/mla/kernel.py +f93bdfc8e8f8dc0f285968e1f0cb112c55bb8910583c72515ecd36e47d647c33 b12x/attention/_shared/mla/kv_cache.py +779b328400977cdcf077da8e28c76e078c526c63ca8ec56b954cfc3de967c73f b12x/attention/_shared/mla/merge.py +f90789ae26a997071032a1b791f434e8d341bcb56d583acf31eba539e8e1314d b12x/attention/_shared/mla/packed.py +a1b78f076db2dda66cd04acc37a952118f00f350670e80b2e1bf6528a9f1d956 b12x/attention/_shared/mla/prefill.py +d608e48061660d9c0df434261d1ede38b713b07915e5c26be9870ab4c8543168 b12x/attention/_shared/mla/prefill_mg.py +6f9735f2bcc3ddf1b6613e40da4b113ebd0e9b7e427c0cbf46ba5f9a8743d22c b12x/attention/_shared/mla/reference.py +c9d255c891c1912421dedeed77a9d7401d574d9734024f2a4896145c33818717 b12x/attention/_shared/mla/smem.py +c50fe413c5cfbca9c352e3d61eb3342624bfa59692b6ec59dd39e0c7cf059cca b12x/attention/_shared/mla/smem_mg.py +df197f9a5c5be2e9d004e9526b0e8f2f33e25ed869221181dbd7254356bab2a0 b12x/attention/_shared/mla/traits.py +93fd7acdaa2f595c5e726dd1a590e34a488147d67370b8e608ded6a9d4f62cbe b12x/attention/_shared/workspace.py +203c15b78e98117792d3df7fc090f7d8d739be786fac78c0f5436f664fce47ef b12x/attention/compressed_mla/__init__.py +c0ecf6e906be135f69b3e6db96df6ddea304e63fd65e816ecfa59588c81b5a49 b12x/attention/compressed_mla/_scratch.py +c9b44b3a1b004155461d0101d833a7ecc044f20aab23097913943a18c3038981 b12x/attention/compressed_mla/api.py +62d18d94c49447a2c93978ee95dba15e194515857fd1238474bf3da45c330569 b12x/attention/compressed_mla/reference.py +c16c1c1515dc3fad434f0e991aa0336107c90f99f99ef76717c4c12feb460e26 b12x/attention/dense_mla/__init__.py +6f6abda33505df59c4a6d05075b41a1d8452e828d7d3f0c893b00f8afa1034a7 b12x/attention/dense_mla/_forward.py +5c2d708142b00a5f81eaba725027dccba88ce68f565c562690949fb6e12d650c b12x/attention/dense_mla/_io.py +3c6de1868be854eb9e78c39cfc598d390a537ac1df38a707e9ca9b7a62644fa2 b12x/attention/dense_mla/_kernel.py +7aea41a322fdd6814894cd2d5c4f1cc0f315037ef0042eec32034de7752562df b12x/attention/dense_mla/_layout.py +9e8d7bcbb7fab50d1651b3731f853bf7dba348e5bfb02dcee6b01f4906cc7fde b12x/attention/dense_mla/_math.py +19649c12fe1e1e756b4644ca6a9ac37357941032a73068f3ec0781c5a5c8bf66 b12x/attention/dense_mla/_merge.py +496d2a26746f2fa254348662074cc87d4f66bcecbc617bb2ad438670f2aefa1e b12x/attention/dense_mla/_reference.py +161db3d2aa4c39a341ea13558cca29204595871f3092fecb7101d8802584f84f b12x/attention/dense_mla/_scratch.py +a731663c6a99b4bfd5fbf7ab579d150e2232fe5e9c1b4d545f4c724bc75f40b8 b12x/attention/dense_mla/api.py +f5c55f720322cd3cf3ffb3e5fd42b36f1fc28ef7a5e7ffa9cb04557816d4fec7 b12x/attention/dense_mla/planner.py +94c4f744f3d6c07cf5b6884efc7d8d7aa2be4d65684a8b4f859d2b2d9656cad6 b12x/attention/nsa_indexer/MSA.md +6e4528c26bad40736ad991f8c5877ef12dbffe1c642c06cfc97dff4245f71086 b12x/attention/nsa_indexer/__init__.py +e7ed71422c1d479a2dbb402ed8037b0f79e11c6f930698a14f2905b6b7e05e30 b12x/attention/nsa_indexer/_impl.py +d5f0912fae79d17be5621bf914495feb8bdf1df1c53b2d0642e59f669ece3135 b12x/attention/nsa_indexer/api.py +b45695e27634bc7aa0aebba097496437a6ba28b7ceae2c39b61203968db63513 b12x/attention/nsa_indexer/contiguous_kernel.py +69110dcf9d54d4e14ee4d501990245a2cbad7621d0a3d84f035f93d368add7af b12x/attention/nsa_indexer/fused_indexer.py +0d29ff511d04482dd302f49a9f138213361cbe1a1b44635045539d9219898dd0 b12x/attention/nsa_indexer/kernel.py +d8d4c028147dcfa0ac4769d2338a92be1718627a10a5bcba119cdc2481cbfe58 b12x/attention/nsa_indexer/msa_reference.py +b0eaded9823737d8057e890ea1eeca4a22b34bb686107ac63041af76d06b1e44 b12x/attention/nsa_indexer/paged.py +05d0682bf07adae8249ac7ad387c63613c8f91ee3e2995f927ab5190eb45dc67 b12x/attention/nsa_indexer/persistent_topk.py +d39b11a3b3eafe859bd1fc8a61d8b5f51ead1113af10ceecf6b378997e0e4809 b12x/attention/nsa_indexer/reference.py +9c6be2d5eb140b6321cd0bf191036625f828faba2df9e3b51b450d935a7c2ccc b12x/attention/nsa_indexer/schedule_metadata.py +88c8f4cbbc4d15ad154c338c0d79fe9873137731eff5abc947fcb3e03584f3b0 b12x/attention/nsa_indexer/scratch.py +722b1207b2be93db0caeb8083043a11f9e12dc5f24dcdca09e14baa713f7abf9 b12x/attention/nsa_indexer/tiled_topk.py +179470a782d4b334ed9129885cd642e2102ffafec114290f9770efe0a2d761dd b12x/attention/paged/__init__.py +0e4e945be97d8dbb1436f5db352c084440e32e2f7c7c0fc0e3be13381cdc9cad b12x/attention/paged/_forward.py +3d73bd0e1caa0db913b039c1e53348bf25bd287f0209f8cf36d21094021f1fe0 b12x/attention/paged/_scratch.py +77d7fe942e84aaaf323b6edd15f4fec801d9b33bba56adebfb994493987088ae b12x/attention/paged/api.py +7b7ffdf17201b3bd2c957302e7406a7180567a0ef6be1ae21964637a72e3f95b b12x/attention/paged/forward_extend_generic.py +345f8fe4e8d8b8836f8b6da4f325149458fddd313272974f651ea2dbcf4d2fe9 b12x/attention/paged/forward_paged.py +dcac8e52521e0641dc777578e3d9b313803d470a53b3a7c05aefc7f66162a8ed b12x/attention/paged/graph_replay.py +e067dd881357afbe153d57c1040a89fa0250fef69d62fce6d31e317322d1d1a2 b12x/attention/paged/merge.py +c232ef42b3653f537e9b8c9e8d8b51c952385e2ead6bb01403ff10fe4f5c14db b12x/attention/paged/planner.py +8aa689e840a1d61552fa4df5c1c0a10948b47970d917e9b9e2473735134aff8c b12x/attention/paged/reference.py +df6e1b1aa0110b5fff4989dcfe567431ba40333311787b60e5a87b6321c64781 b12x/attention/paged/traits.py +a82ed9e7c303aad17feb7df5c59b248dfe904cc657b65d14749be3df0d03aabd b12x/attention/paged/tuning/__init__.py +7c5559be7b71466e710b9fc9bdc5051237e0dad600fbcffeb9e8d0d9edb63c9c b12x/attention/paged/tuning/registry.py +4c75917c9fd5c518bceaf0f99396bda8352c1b009edfddcb5390997169462dcd b12x/attention/paged/workspace.py +2bf12190ef649bd958b33b8cf842e1dd55665137391a44848a66ea638bfb6f5c b12x/attention/sparse_mla/__init__.py +73fe60978007a3bf3d3b61040187167b5a9efc0ffc5d9973702ed16f35ddbe47 b12x/attention/sparse_mla/_scratch.py +56727154aef78eed8f23d7adb241a7dede200dc0a37f2dc8e7361c56709ad9de b12x/attention/sparse_mla/api.py +6a8159933971155b4466ca6535dff7c02520081eee6c6c2dc194864c12eab0ef b12x/attention/sparse_mla/reference.py +93f35a6197346e913542426d8b52134ddbef7c86ea9c84c6bd9bcff5b462f40e b12x/attention/varlen/__init__.py +2365506f5bf0764fc4c8ad91741e8351484dd1bb6c80390db8a76be6a4129122 b12x/attention/varlen/api.py +40e9a3d28152c461cf0cc8d3ae8a9c6f16ba8005ee3f38d772379e7bc607191e b12x/comm/__init__.py +7f50f8c5a9659dede06dec2f7c5901edc21d73e42128e532bb51447e89e8b72d b12x/comm/pcie/__init__.py +8ae32fff0537ef226253670142758731f57e75db45279ff41eca4712f34971bc b12x/comm/pcie/_cuda_ipc.py +2eaba37b27df510550a0dffddffb66f5674e50e6f15ed1552c73245901db2b41 b12x/comm/pcie/_cute_intrinsics.py +ed6fc1e2885d4bf482fccace45335b7248cd7486b7e6a129d50dc9303e7882db b12x/comm/pcie/_dcp_a2a_cute.py +4f6442f151c3ca826e755910901ebfef4f7739058e7eef23cfae0c0fb18cfa0f b12x/comm/pcie/_dcp_cute_common.py +559d2571df8f0f134e8300a7e812af5ceac3a1e38ff3a69277c4bd2a81a3668d b12x/comm/pcie/_dcp_topk_cute.py +6f173ce05183e28ae7cf121ff863620026e368a3eeb6cc07a14f83f4f9f6e3a1 b12x/comm/pcie/_dma_kernels.py +200d7437baf6d2f503fb134ec738ae61fa05d57e7d92b59b76fc32b032262db9 b12x/comm/pcie/_hierarchical_cute.py +03fabbf3d07cd9d606f55fe8a4f82d48cc13cc40cb0a38d5166052643c8ad3e9 b12x/comm/pcie/_island_rs_cute.py +50712d803840215514780a25ad8722479aa0be4a7dc1f7805131105346fbf864 b12x/comm/pcie/_launch_geometry.py +4f120f49fda819253bedc62e877325b776a7fe6d1eea6eb457543766285d54c7 b12x/comm/pcie/_oneshot_cute.py +0210e618b4d1fd791fa040e8fb80f6a3dc64185cf4e70feb07da4deb4f51f1b8 b12x/comm/pcie/_twoshot_cute.py +75232f0b8f2ac15886642b402a8ba572f50784733610b24c24f92cb8e6fe2b9a b12x/comm/pcie/_vocab_argmax_cute.py +cf9dfdba53601e3418a7db0b3b6b0856104cb078df5e06e97ae5c91acf486c44 b12x/comm/pcie/api.py +ed1f6fdabe46ced9af860fc5817ae8184ddbc73a5112de598db5f2c33fd3da75 b12x/comm/pcie/overlap_probe.py +4b3f6935437ed9b4283c2fc9c10162803f832f2899316ece531d46b7e8682082 b12x/comm/pcie/pcie_allreduce.py +9078b4f8866b5a621834599c760f1cd851015161278b527e1c31464d194da1f3 b12x/comm/pcie/pcie_dcp_a2a.py +dfc66858cef824f7a2f76ac5c1878e20f88a3dc2c4c82554b3a67313a5f7fe93 b12x/comm/pcie/pcie_dcp_topk.py +8a6c8b68870256e3cdcf663d4127923c69eabc3e46a18f8a7a3679e66532dd45 b12x/comm/pcie/pcie_dma.py +20f35fe16359633a71e9f4cbd46032e486162f5d8fa583d7cfa99354687516dc b12x/comm/pcie/pcie_hierarchical.py +8c06c1dcbe5d85bf32c1b447cdc3ea483d6cc2eca63eefc52152e4974bc592e4 b12x/comm/pcie/pcie_island_rs.py +44d98c104195a43a9c10a3059ac3d5b3b0ebe38a1e21efbd8231593f43567b4d b12x/comm/pcie/pcie_oneshot.py +0cf4fcc397f367c3eb44daa42d01c2ee004285ac7adb877bcdc7037f26b87a27 b12x/comm/pcie/pcie_twoshot.py +90e0cbab83877fff921e9c583ce595cba3e8828e0007501798cdcf6bdad45673 b12x/comm/pcie/pcie_vocab_argmax.py +aceac92faf7902c00eff2332b2553e7a2961f8d805e7233bf8aa00b4bfb852a9 b12x/gemm/__init__.py +f72c33fdd880c18aa8b52e7268c9d8777cb943a9eb80d95ce05ebe7088b002f2 b12x/gemm/_bmm/__init__.py +bc7f8f597c4a1a97009580bff2372fcba8b6d28ffbbc3fadb8bf538e0bf5e58f b12x/gemm/_bmm/api.py +1d73d0cd8f6ddf06228ff3dfea94affa0d5ddd6faff6ebe3dd1690326cbbf840 b12x/gemm/_shared/__init__.py +f837065140dcfd7200f11bfa7fa86dee5d5aa890379f1af4569b3cac36d5d7df b12x/gemm/_shared/block_fp8.py +d181e9a976c37e11eb2421774fb53a8c024a95fa8d1f3c51ee771e7f39d92aa4 b12x/gemm/_shared/mxfp8_bmm.py +84252b604a319d37f852a853c2b4a4bcf5a30554b6bffd172f9e0fc3602cffc6 b12x/gemm/_shared/wo_mxfp8.py +ba21573eaeac204d74037941ee5d90b7a7fbefcd08d8b427c82870eb7d7df658 b12x/gemm/bf16_gemv/__init__.py +9db881ec72b33081bacc901434d91ded967a0482dae1939f097487ba0a556fc2 b12x/gemm/bf16_gemv/_kernel.py +08df111d9cf77677c024be99a6383891cb1b8bb90a3a3b581a4ad49fd776f5d1 b12x/gemm/bf16_gemv/api.py +66484e9b4e9e697d1e6ee3c906084b50983d8ac09cadcb09d667c161392c0e6b b12x/gemm/block_fp8_linear/__init__.py +964881d9cb3ebadc5447801d02842dca7f01ff5166752dc444993ad95e90f1ff b12x/gemm/block_fp8_linear/api.py +70a641c742759916cfca2ec3c366f8936cdca4f311cbd00138d2f54ec8644076 b12x/gemm/blockscaled/__init__.py +82a451038a1dc83d71e8d47c8a69b70fa271847f05157e6fc013492242918eee b12x/gemm/blockscaled/api.py +b61ee841e24684f41dea178d1c88beac7792dd814781a240ae34a045647dc24a b12x/gemm/mla_query_projection/__init__.py +92dfdc75492e7a4a9a9c23bde865a6bac0264fcbce0f68e31fedb8fb7084cb20 b12x/gemm/mla_query_projection/_bf16.py +aabbd178ca27c78a59a348a30672a3b495de701e7da787c68a4dd2fdc173d4ef b12x/gemm/mla_query_projection/api.py +bf7a069d79a2396d9ea3d83233b75cb73279c9f928213c219b0105e702190641 b12x/gemm/mxfp8_linear/__init__.py +e79b15fa138d12fec43071469202fc5d3780265e13bc8705631a50488f2aab7f b12x/gemm/mxfp8_linear/_kernel.py +7777af3259663b5542b1607bbbac7a191f45b2ca669b257094eb81fd9c01b918 b12x/gemm/mxfp8_linear/api.py +08fca880b7c7be073e0a89f127a9ff1e65dba3b666fc6af0a501ddba8a96db62 b12x/gemm/tensor_fp8_linear/__init__.py +393d59ecaee5e1177c2a3b7c73b7b6e63f13a66056adb37670d3578444abd3b3 b12x/gemm/tensor_fp8_linear/_kernel.py +e0a390f9ee5cdbc63e456413a60ff57e40e0640184d95ba264096b1fe5794af8 b12x/gemm/tensor_fp8_linear/api.py +571868dd226d54a08e127f32e4764a5be7a18e65944f131cc398101dc8059d52 b12x/gemm/trellis_linear/__init__.py +ee498b925e892cf6c3d83542baf257ad3641bb282a1a370555452935ff247b30 b12x/gemm/trellis_linear/_small_m.py +04e00691a2966b0f32dbb58f6f7e0daa73a4e39284a07f64c1ee691f3456973d b12x/gemm/trellis_linear/api.py +f19215989d655b5db0855feaeab20dea625e01d879f169b838aee79d9f364710 b12x/gemm/trellis_linear/csrc/trellis_k6_small.cu +27a32b6263fcd96c79d3beeecf221c4366780bdf15ad51986f48650bd7369bff b12x/gemm/trellis_linear/csrc/vendor/LICENSE.exllamav3 +bda302eb7c36d54e26192be226ff9095278b12dbe46ac96e056293812659b51c b12x/gemm/trellis_linear/csrc/vendor/compat.cuh +efefd4fab99cfb62de4a5909198355f2ece2579dddbd1825beba535d828e7348 b12x/gemm/trellis_linear/csrc/vendor/ptx.cuh +699c6743fe9e03eb5a4d8b722ef78efcbe69422576e5ae67f8668c547cb820cc b12x/gemm/trellis_linear/csrc/vendor/quant/codebook.cuh +44a8a357c584093b186dde6885239b77d0cef33b2a3192782103495a98a95aee b12x/gemm/trellis_linear/csrc/vendor/quant/exl3_devctx.cuh +07a72440d42d805fc947da090aa61b7f6333ddf93b2becf46096172a28df373e b12x/gemm/trellis_linear/csrc/vendor/quant/exl3_dq.cuh +0f8438df76096e3824cd1a8a7da9fba6713d515528637605247b73f5ebbd4c9e b12x/gemm/trellis_linear/csrc/vendor/quant/exl3_gemm_inner.cuh +6084b5cd97865028af3e63e6d8bb0a1f8a189e5d1b7104f00929b3667b1e9cb4 b12x/gemm/trellis_linear/csrc/vendor/quant/exl3_gemm_kernel.cuh +082368c0852c9cb5679353a0682ef468d3ee6565fe93b79e80d3764d1d8d4cdf b12x/gemm/trellis_linear/csrc/vendor/quant/exl3_kernel_map.cuh +a2382a68f0bd2053a05d43f09751a3ed917a111e73f9f4f47a806f56118a7c1a b12x/gemm/trellis_linear/csrc/vendor/quant/hadamard_inner.cuh +2bc7c3a22414517450832c753fd13eb0cf9960f16f958e493b4ca29cf02bdc09 b12x/gemm/trellis_linear/csrc/vendor/util.cuh +d5ee6fbe6942918d2644ffeb431c391fe403d7c7e6a7cf5ba925ae231ec9ce18 b12x/gemm/trellis_linear/csrc/vendor/util.h +cb6eac5af17fa47e47c407ca2e494014877537fa1d62702cd031e282e584d184 b12x/gemm/wo_projection/__init__.py +92f1d09ec49197946918d4e9a93b404edf686a5b15f47fd5bf45a43fb72f1bf6 b12x/gemm/wo_projection/_quant_cute.py +ccc307761ea131c138d8111f73b4cfbb1401f6efaab1ecdcc10fe88c371adb6f b12x/gemm/wo_projection/api.py +adb536e2514af049a800cd78a5c147c32d9c3dcaee2d1fb95ea30c7c99f92520 b12x/integration/__init__.py +cad9c836b78133c304265430f9bb17cefd4c36823338e4eb8e772851c01ec957 b12x/integration/vllm/README.md +4928525128b483fdbfc87cfb62499cdf7cfe0dbb8356a98a9a3eeed6484b480d b12x/integration/vllm/__init__.py +aab476e35c53e0f4433a37b4bd289d4080edb1112d4bf56d81397fb08fa186b2 b12x/integration/vllm/fp6_serving.py +b96f127155d26e9edff179d7f9f3ce1c7934683d7c7e39ff97ec240333a5f371 b12x/integration/vllm/plugin.py +70f9957c76a4dcf08c0e683d6824896af67107df09259a309f946268984c4620 b12x/moe/__init__.py +0f30da977628d1a031bcdf559c247fe5ee0ce01425a524b0dbf5dd399e23a378 b12x/moe/_shared/__init__.py +39acd00227129e41f934e41b358edc52bffe15b567059ce2df84a1299238bcd8 b12x/moe/_shared/execution.py +138f7d07db901b511a820f981bb97041ad1bd646deec351893a0138aaed21301 b12x/moe/_shared/kernels/__init__.py +2d29c84341fcaa2df7859ad6dbcfb16ac5e3b4e48d282cc9586bffb43e906676 b12x/moe/_shared/kernels/activations.py +3b4a9bc7e45d95f5c85415c324fdb7d693f16157df938521ad97f0ed30772065 b12x/moe/_shared/kernels/dynamic.py +c2f1140afb0f17b31615aa3072c8fcec5dfbf67afa9193ee07eab19b40e7894d b12x/moe/_shared/kernels/micro.py +4525eefb70abd2625605b14f90cc6f8036347e7d276c75790d3ae00dd703a811 b12x/moe/_shared/kernels/mxfp6_moe.py +d3e6e44ebba3e7c9566a6e6f76b982324713528f95cf96303aa3d2c73af866c1 b12x/moe/_shared/kernels/reference.py +c60c687972121568a97a7af49b0dec4f338229b4ec7743ae552c920ed1f45cc3 b12x/moe/_shared/kernels/reference_flashinfer.py +d81c98d9a30b28ec2b94637b6132fd7aa499d5053760d26a697c4d6998ad0f87 b12x/moe/_shared/kernels/relu2.py +abb74591a75307180761af3b71232412c40b8bf50d94a44f57b2f87e0e0ba2f5 b12x/moe/_shared/kernels/silu.py +e497cc542ff757707b7cee0fbbdd0ca24284ed9bdc4d535b07e66c653132c2eb b12x/moe/_shared/kernels/situ.py +f3d004e4fe7e97b876522e8946d0282db27f99707f107ce3dfbdc31f5447a831 b12x/moe/_shared/kernels/tiny_decode.py +6e8d31d9095b4d8f94f02ffa8deb7da7cfa5348643ad99f1f11218a9622de66e b12x/moe/_shared/kernels/trellis_decode.py +3c012619efbf7c533501e7048bd9b3c6ac321859904f7a07b849631b86350357 b12x/moe/_shared/kernels/w4a16/__init__.py +893b5f1cb643c96ad72ee405e362ca9b3ac2b040f30d9ec7e1b9796445bbc60e b12x/moe/_shared/kernels/w4a16/host.py +6f4a08715ca0ad28560881d72398173b72103de919d428d2118f80ba806f8b6b b12x/moe/_shared/kernels/w4a16/kernel.py +dc03f8e71123ef5288264825fbacc0777227cf933c27d280de7c2f8db2a356c9 b12x/moe/_shared/kernels/w4a16/mixed_trellis.py +66bbab89e186dd57870349774d93b4f9844b1428d5ac38895e70b5bf8e944fd6 b12x/moe/_shared/kernels/w4a16/prepare.py +6fab8230d72e664415f1242d975de342d40b642173fa395b16e4ccb29d2b69f6 b12x/moe/_shared/kernels/w4a16/route_pack.py +299922369ef7fb23ae46d01626e97b3d5eaed8166f97ff5e50f86fb428370324 b12x/moe/_shared/kernels/w4a8/__init__.py +95378890e6d58befda88e5b81c6e9fdf92d3c5f8d598af5eb65e46c17c33243f b12x/moe/_shared/kernels/w4a8/weights.py +cb72c50dab933ee866103caf7c32f1a6cfb15df991fcdd8b753ac765551946b1 b12x/moe/_shared/kernels/w4a8_phase1.py +ce31085628a8423468a453275f18c3322f8029d35afccb35f5e7060cbab11a7e b12x/moe/_shared/kernels/w4a8_phase2.py +0a128b4f31f49ada5ac3ba3ac1aa1571670451176428d89ddd83abe3700f15ac b12x/moe/_shared/kernels/w4a8_trellis_decode.py +363cf5716975a1747014e7750bd90292f45879620dc3f9d10c9fd9120a1d812e b12x/moe/_shared/kernels/w6a8/__init__.py +e1883a2625d2e7cf59d31f357676d598858086b14f63e6719e98546823eb0756 b12x/moe/_shared/kernels/w6a8/weights.py +13c5a01a6f7631c2c98c11ff6122bf88b2b4edf292c1fa4bb1df7a71eeb45f80 b12x/moe/_shared/routing.py +efc68e246e42a25641de621a5b6867996ad89689015590cce3636080ca0a8040 b12x/moe/_shared/tuning/__init__.py +26afcd57a406bdb9c753375f5daba19d7d4d5c3679221e8c5aa48a54a9e41222 b12x/moe/_shared/tuning/decode_max_active_clusters.py +fdfef846b025284d40fc738160fc165fbc6d291c2c05b98500987922a5f6c83b b12x/moe/_shared/tuning/registry.py +02ef189de17bdcf2f8dcf4ab82d2c18b7f9518c859740088eb33c0772c2ae90b b12x/moe/calibration.py +84c3b337aa18515401c58fcba3002fa4b3ad8a20a373216405356b0196035a10 b12x/moe/ep_moe/__init__.py +6f882c86124031774a6a65b9d160f791cb8c2344e1c79c644016a2a7b9236cae b12x/moe/ep_moe/_impl.py +db00e13913e2739e85025326925d0de9acdc3bb3a1e57ebe38c076f1f265511f b12x/moe/ep_moe/api.py +b4fe2e37922a49068d9c719187345826561fdd86b7e3115187efca698d332de6 b12x/moe/fused_moe/__init__.py +7214ccddb2927e9af88cf0d6d520743d6bfc6d37cc52cabfb8f73478e72827b1 b12x/moe/fused_moe/_impl.py +d3266bda7016e723693a33159a824d90cbad32a84c1ea2d36194eff2457fe8a5 b12x/moe/fused_moe/api.py +33fe280c87f6098fbfa196fb664747e6643a8faac6e47e0af894551f0dba83c4 b12x/norm/__init__.py +3558836e0e17b61a148e3e49cbf0331e6afe912478e5a6a1d1e419951056a41d b12x/norm/mhc/__init__.py +c83481c6e367fb7d5c1f7e21892f1902a11c345934ca80afd1eda9fe4fca59f9 b12x/norm/mhc/_impl.py +decba110f9e817d2f5016f7255a01d4349381a435380eba559f3a50490954859 b12x/norm/mhc/_kernels.py +1a130dcfef7eedf596dd0f090a93d4fd119f630cc47803cef04889efe21422b0 b12x/norm/mhc/api.py +69a42dbb9d5a092062557e71e3f07b513c2298c9d9694cca9c1ab696a60c46a1 b12x/quantization/__init__.py +91ae99f19b190cf043fefd4c53bcbe9cf463b7dcb3fb862a74fb6a5414fb5aeb b12x/quantization/mxfp6/__init__.py +3f53454a64777a879d1184ef2dabbdbcf5fc9538af39d0345f806f65802ee121 b12x/quantization/mxfp6/bf16_to_fp6_small_m.py +86b34f05bd70c62255a4798d9ecd7c726c917793acc5803e66f24020df358f07 b12x/quantization/mxfp6/bf16_to_fp6_tma.py +8891f37a6929d8f3fbce6b0859e5b6ecddd15c7f07ac145e7bbb135e8cfd5ae9 b12x/quantization/mxfp6/calib_io.py +6b20888b97a991deac33f83ca027f701afb624b4684b81278107bfa246039597 b12x/quantization/mxfp6/fp6_checkpoint.py +af1f14ad3ff8959dc67944486df95f33b1e6e3ab7d1a998ea953cb719cd8447b b12x/quantization/mxfp6/fp6_dense_op.py +6457d3f792ee8cefef97fa7eebb97d82ce9b182e6a03db04ecc978ac01a5f34a b12x/quantization/mxfp6/fp6_dense_weights.py +15363047d54a6f4aa9cdffb3099680bbf94f8b67580c2da05e75cff7f215b94c b12x/quantization/mxfp6/fp6_moe_weights.py +977888b899ee6fa9d316777badcbea1f4ae06070d7f54a97bc954a77d3eaa091 b12x/quantization/mxfp6/fp6_row_gs.py +7074f9e6d897bd43988ad3d95e61f0c8601a36b82efd79de6564c6749634e6a9 b12x/quantization/mxfp6/fp6_safetensors_export.py +4041b7fbec8b3dc55bc274b6628148c8a536dc27f0ecbbec72f3c655640257f1 b12x/quantization/mxfp6/fp6_safetensors_load.py +42b2d7f611f2413f298d70ad6899a95a9dfe80d4ded6d7bdecaebcd57b6bc5f6 b12x/quantization/mxfp6/model_fp6.py +04b47b16c46c14f2b0247b79ec425c69d17baf5007bf7ccb65754962e8709f06 b12x/quantization/mxfp8/__init__.py +4343940a331240ebba3f222d6689faefaf9564bfd2cc4b63c574e9bb3df200f9 b12x/quantization/mxfp8/api.py +58d8fc674d2874cd303e11f44b417a644538fdbf2acc5fd96adb70ad24ad182d b12x/quantization/nvfp4/__init__.py +2f041a7b20c93d7644bb72ef5b1cf89405596c046855be1a39e62eb63470e0ca b12x/quantization/nvfp4/_impl.py +23efdfad34efc9bebb2ba337933feb93bf4d5123db1428c7a8f3157bdf273655 b12x/quantization/nvfp4/_kernel.py +bbce31ccc9cf38e0049d368af2f594a8368ae9d4e25037dc982865ddfe0a5499 b12x/quantization/nvfp4/api.py diff --git a/b12x/__init__.py b/b12x/__init__.py index 30af26d48..af965c49f 100644 --- a/b12x/__init__.py +++ b/b12x/__init__.py @@ -46,7 +46,6 @@ "attention.nsa_indexer", "attention.varlen", "comm.pcie", - "gemm.bf16_gemv", "gemm.blockscaled", "gemm.block_fp8_linear", "gemm.bmm", diff --git a/b12x/_lib/compiler.py b/b12x/_lib/compiler.py index 529985c86..28aa68029 100644 --- a/b12x/_lib/compiler.py +++ b/b12x/_lib/compiler.py @@ -1934,12 +1934,6 @@ def _compile_options_cache_key(compile_callable: Any) -> tuple[str, ...]: return tuple(serialized) -def _dsl_compile_options_kwargs_key(compile_callable: Any) -> tuple[str, ...]: - """Return the raw subscripted compile options for cache provenance.""" - - return _compile_options_cache_key(compile_callable) - - def _compile_disk_cache_payload( compile_callable: Any, func: Any, @@ -2645,9 +2639,7 @@ def compile( compile_callable = CompileCallable(dsl_compile_options) kwargs = dict(kwargs) - kwargs["__dsl_compile_options_key"] = _dsl_compile_options_kwargs_key( - compile_callable - ) + kwargs["__dsl_compile_options_key"] = _structural_cache_key(dsl_compile_options) memory_cache_key = _compile_memory_cache_key( compile_callable, func, args, kwargs, compile_spec ) diff --git a/b12x/_lib/dense_gemm.py b/b12x/_lib/dense_gemm.py index ad5790b6a..4cbc9022c 100644 --- a/b12x/_lib/dense_gemm.py +++ b/b12x/_lib/dense_gemm.py @@ -89,6 +89,7 @@ scatter_add_bf16, scatter_add_bf16x2, shared_ptr_to_u32, + st_global_u32, st_global_u64, st_shared_u16, st_shared_u8, @@ -124,7 +125,8 @@ def _dense_spark_policy_for_sm_count(sm_count: int) -> bool: _B12X_TIMING = ( - os.getenv("B12X_TIMING", "0") == "1" or os.getenv("VLLM_B12X_TIMING", "0") == "1" + os.getenv("B12X_TIMING", "0") == "1" + or os.getenv("VLLM_B12X_TIMING", "0") == "1" ) _B12X_TIMING_THRESHOLD_MS = float( os.getenv( @@ -132,12 +134,16 @@ def _dense_spark_policy_for_sm_count(sm_count: int) -> bool: os.getenv("VLLM_B12X_TIMING_THRESHOLD_MS", "0"), ) ) -_B12X_DENSE_SPLITK_TURBO = os.getenv("B12X_DENSE_SPLITK_TURBO", "1") == "1" +_B12X_DENSE_SPLITK_TURBO = ( + os.getenv("B12X_DENSE_SPLITK_TURBO", "1") == "1" +) # MX-FP6 decode uses at most three mainloop stages when two CTAs share an SM. _FP6_DECODE_TILE = (16, 64) _FP6_PREFILL_TILE = (128, 128) -_B12X_DENSE_ATOM_24 = os.getenv("B12X_DENSE_ATOM_24", "0") == "1" +_B12X_DENSE_ATOM_24 = ( + os.getenv("B12X_DENSE_ATOM_24", "0") == "1" +) _DENSE_LOAD_PATHS = ("tma", "cpasync") # Expand-ahead for packed-B: at k_block 0 the MMA warps wait for stage s+1 and @@ -154,11 +160,9 @@ def _dense_spark_policy_for_sm_count(sm_count: int) -> bool: # hot path). The GEMM's producer warp does a full-row amax scan, derives # gs/alpha, then quantizes each K-tile's 32-element blocks directly into # sA/sSFA smem. Distinct from the MXFP8 ``fused_quant_a`` machinery. -_DENSE_FUSED_QUANT = os.environ.get("B12X_DENSE_FUSED_QUANT", "0").lower() not in ( - "0", - "false", -) - +_DENSE_FUSED_QUANT = os.environ.get( + "B12X_DENSE_FUSED_QUANT", "0" +).lower() not in ("0", "false") @dataclass(frozen=True) class _DenseGemmPlan: @@ -639,9 +643,7 @@ def __init__( mxfp6_fmt_a = mxfp6_fmt mxfp6_fmt_b = mxfp6_fmt elif mxfp6_fmt_a is None or mxfp6_fmt_b is None: - raise ValueError( - "mxfp6_fmt_a and mxfp6_fmt_b must both be set or both None" - ) + raise ValueError("mxfp6_fmt_a and mxfp6_fmt_b must both be set or both None") self.mxfp6_fmt_a = mxfp6_fmt_a self.mxfp6_fmt_b = mxfp6_fmt_b self.block_fp8 = bool(block_fp8) @@ -848,7 +850,9 @@ def __init__( self.mma_register_requirement = 232 def _setup_attributes(self): - mma_sf_dtype = cutlass.Float8E8M0FNU if self.block_fp8 else self.sf_dtype + mma_sf_dtype = ( + cutlass.Float8E8M0FNU if self.block_fp8 else self.sf_dtype + ) if cutlass.const_expr(self.a_dtype == cutlass.Float8E4M3FN): mma_op = cute.nvgpu.warp.MmaMXF8Op( self.a_dtype, @@ -856,7 +860,8 @@ def _setup_attributes(self): mma_sf_dtype, ) elif cutlass.const_expr( - self.a_dtype == cutlass.Float6E3M2FN or self.a_dtype == cutlass.Float6E2M3FN + self.a_dtype == cutlass.Float6E3M2FN + or self.a_dtype == cutlass.Float6E2M3FN ): # MX-FP6 uses inline ``mxf8f6f4`` MMA in the mainloop. Build tiled_mma # with the MXFP8 op so smem/SF layouts match m16n8k32 geometry. @@ -920,7 +925,6 @@ def _setup_attributes(self): # shared memory. The explicit format is therefore the reliable policy # discriminator; ``a_dtype`` alone also matches ordinary MXFP8. if self.mxfp6_fmt_a is not None: - def _probe_stages(epi_tile: tuple, epi_stage_cap: int) -> tuple: return self._compute_stages( self.tile_shape_mnk, @@ -944,7 +948,9 @@ def _probe_stages(epi_tile: tuple, epi_stage_cap: int) -> tuple: _probe_stages, stages_through_smem=not self.use_m1_non_tma_c, ) - self.ab_stage, self.epi_stage = _probe_stages(self.epi_tile, epi_stage_cap) + self.ab_stage, self.epi_stage = _probe_stages( + self.epi_tile, epi_stage_cap + ) else: # Non-FP6 families use the generic stage policy. self.ab_stage, self.epi_stage = self._compute_stages( @@ -1319,11 +1325,17 @@ def _accumulate_block_fp8_stage( ) -> None: accum_mn = _reshape_acc_to_mn(accumulators) stage_accum_mn = _reshape_acc_to_mn(stage_accumulators) - scale_n = (tile_coord_mnl[1] * Int32(self.tile_shape_mnk[1])) // Int32(128) - scale_b = cutlass.Float32(sfb[(scale_n, k_tile_global, tile_coord_mnl[2])]) + scale_n = ( + tile_coord_mnl[1] * Int32(self.tile_shape_mnk[1]) + ) // Int32(128) + scale_b = cutlass.Float32( + sfb[(scale_n, k_tile_global, tile_coord_mnl[2])] + ) for acc_m in cutlass.range_constexpr(cute.size(accum_mn.shape[0])): coord = coord_mn[acc_m, 0] - m_coord = tile_coord_mnl[0] * Int32(self.tile_shape_mnk[0]) + coord[0] + m_coord = ( + tile_coord_mnl[0] * Int32(self.tile_shape_mnk[0]) + coord[0] + ) scale_a = cutlass.Float32(0.0) if m_coord < Int32(sfa.shape[0]): scale_a = cutlass.Float32( @@ -1331,7 +1343,9 @@ def _accumulate_block_fp8_stage( ) scale_ab = scale_a * scale_b for acc_n in cutlass.range_constexpr(cute.size(accum_mn.shape[1])): - accum_mn[acc_m, acc_n] += stage_accum_mn[acc_m, acc_n] * scale_ab + accum_mn[acc_m, acc_n] += ( + stage_accum_mn[acc_m, acc_n] * scale_ab + ) stage_accum_mn[acc_m, acc_n] = 0.0 @cute.jit @@ -2158,8 +2172,12 @@ def kernel( if cutlass.const_expr(not self.block_fp8): tCsSFA_p_filtered = cute.filter_zeros(tCsSFA_p) tCsSFB_p_filtered = cute.filter_zeros(tCsSFB_p) - tCrSFA_copy_view_filtered = cute.filter_zeros(tCrSFA_tile_copy_view) - tCrSFB_copy_view_filtered = cute.filter_zeros(tCrSFB_tile_copy_view) + tCrSFA_copy_view_filtered = cute.filter_zeros( + tCrSFA_tile_copy_view + ) + tCrSFB_copy_view_filtered = cute.filter_zeros( + tCrSFB_tile_copy_view + ) # Whole-stage SF copy: scale bytes for all k blocks of the # acquired stage load in one bulk copy. @@ -2213,8 +2231,10 @@ def kernel( # boundary. The matching write->ldmatrix fence # is the deferred barrier after the last # k-block's MMA below. - lookahead_peek = mainloop_pipeline.consumer_try_wait( - packed_b_lookahead_state + lookahead_peek = ( + mainloop_pipeline.consumer_try_wait( + packed_b_lookahead_state + ) ) mainloop_pipeline.consumer_wait( packed_b_lookahead_state, lookahead_peek @@ -2225,7 +2245,8 @@ def kernel( Int32(tidx), self.tile_shape_mnk[1], self.tile_shape_mnk[2], - self.num_mma_warps * self.num_threads_per_warp, + self.num_mma_warps + * self.num_threads_per_warp, self.mma_sync_barrier, ) packed_b_lookahead_state.advance() @@ -2278,7 +2299,8 @@ def kernel( Int32(tidx), self.tile_shape_mnk[1], self.tile_shape_mnk[2], - self.num_mma_warps * self.num_threads_per_warp, + self.num_mma_warps + * self.num_threads_per_warp, self.mma_sync_barrier, ) @@ -2324,9 +2346,7 @@ def kernel( else: mma_atom.set( WarpField.SFB, - tCrSFB_tile[ - None, _nt, k_block_idx - ].iterator, + tCrSFB_tile[None, _nt, k_block_idx].iterator, ) cute.gemm( mma_atom, @@ -2380,7 +2400,9 @@ def kernel( tCsSFA_p_filtered, tCrSFA_copy_view_filtered, ) - if cutlass.const_expr(self.direct_sfb_representative): + if cutlass.const_expr( + self.direct_sfb_representative + ): self._fill_replicated_sfb_fragment( tCrSFB_tile[None, None, 0], sSFB[ @@ -2645,7 +2667,9 @@ def kernel( elem_idx ] if cutlass.const_expr(self.row_scale): - tRS_cAcc_slice = tRS_cAcc[(None, mma_m, mma_n)] + tRS_cAcc_slice = tRS_cAcc[ + (None, mma_m, mma_n) + ] tRS_rRowScale_slice = tRS_rRowScale[ (None, mma_m_in_epi, mma_n_in_epi) ] @@ -2732,7 +2756,8 @@ def kernel( * split_k_acc_mn[acc_m, acc_n1], ) if cutlass.const_expr( - cute.size(split_k_acc_mn.shape[1]) % 2 == 1 + cute.size(split_k_acc_mn.shape[1]) % 2 + == 1 ): acc_n = ( cute.size(split_k_acc_mn.shape[1]) - 1 @@ -3204,12 +3229,20 @@ def kernel( ) for w in (w0, w1, w2, w3): hi = u32_as_f32(w & Uint32(0x7FFF0000)) - lo = u32_as_f32((w << Uint32(16)) & Uint32(0x7FFF0000)) - local_amax = fmax_f32(local_amax, fmax_f32(hi, lo)) + lo = u32_as_f32( + (w << Uint32(16)) & Uint32(0x7FFF0000) + ) + local_amax = fmax_f32( + local_amax, fmax_f32(hi, lo) + ) i_vec += Int32(self.num_threads_per_warp) fused_amax = warp_reduce(local_amax, fmax_f32) - fused_amax_c = fmax_f32(fused_amax, cutlass.Float32(1e-6)) - fused_gs = cutlass.Float32(self._fused_gs_num) / fused_amax_c + fused_amax_c = fmax_f32( + fused_amax, cutlass.Float32(1e-6) + ) + fused_gs = ( + cutlass.Float32(self._fused_gs_num) / fused_amax_c + ) _fq_k_base = Int32(0) _fq_sa_stage = Int32(0) _fq_packed_scales = Uint32(0) @@ -3232,7 +3265,7 @@ def kernel( _fq_sg = 0 _fq_si = 0 - for _k_tile in range(0, k_tile_iter_cnt, 1, unroll=2): + for k_tile in range(0, k_tile_iter_cnt, 1, unroll=2): mainloop_pipeline.producer_acquire(mainloop_producer_state) k_tile_global = k_tile_start + mainloop_producer_state.count @@ -3711,9 +3744,16 @@ def kernel( # quantize this K-tile's 32-element blocks from the # BF16 row straight into sA/sSFA using the row-wide # gs derived in the work-tile prologue above. - _fq_k_base = k_tile_global * Int32(self.tile_shape_mnk[2]) - _fq_sa_stage = mainloop_producer_state.index * Int32( - self.tile_shape_mnk[0] * self.tile_shape_mnk[2] + _fq_k_base = ( + k_tile_global + * Int32(self.tile_shape_mnk[2]) + ) + _fq_sa_stage = ( + mainloop_producer_state.index + * Int32( + self.tile_shape_mnk[0] + * self.tile_shape_mnk[2] + ) ) _fq_packed_scales = Uint32(0) @@ -3721,24 +3761,34 @@ def kernel( self.tile_shape_mnk[2] // self.sf_vec_size ): _fq_k_abs = ( - _fq_k_base + Int32(_fq_sg * self.sf_vec_size) + fq_lane + _fq_k_base + + Int32(_fq_sg * self.sf_vec_size) + + fq_lane ) _fq_val = cutlass.Float32( directX_bf16[(Int32(0), _fq_k_abs)] ) - _fq_bmax = warp_reduce(fabs_f32(_fq_val), fmax_f32) + _fq_bmax = warp_reduce( + fabs_f32(_fq_val), fmax_f32 + ) _fq_su32 = fp6_block_ue8m0_exact( _fq_bmax, fused_gs, cutlass.Float32(self._fused_fmt_max), ) - _fq_inv = ue8m0_output_scale_exact(_fq_su32, fused_gs) + _fq_inv = ue8m0_output_scale_exact( + _fq_su32, fused_gs + ) _fq_scaled = _fq_val * _fq_inv - if cutlass.const_expr(self._fused_act_fmt == "e4m3"): + if cutlass.const_expr( + self._fused_act_fmt == "e4m3" + ): _fq_pair = cvt_f32_to_e4m3x2( cutlass.Float32(0.0), _fq_scaled ) - elif cutlass.const_expr(self._fused_act_fmt == "e3m2"): + elif cutlass.const_expr( + self._fused_act_fmt == "e3m2" + ): _fq_pair = cvt_f32_to_e3m2x2( cutlass.Float32(0.0), _fq_scaled ) @@ -3758,7 +3808,8 @@ def kernel( ) _fq_packed_scales = _fq_packed_scales | ( - (_fq_su32 & Uint32(0xFF)) << Uint32(_fq_sg * 8) + (_fq_su32 & Uint32(0xFF)) + << Uint32(_fq_sg * 8) ) # Broadcast scale bytes to all 128 SFA M-rows. @@ -3766,8 +3817,12 @@ def kernel( # for (m_row, sg) = (m%32)*16 + (m//32)*4 + sg. _fq_sfa_sg = self.tile_shape_mnk[2] // self.sf_vec_size _fq_sfa_slots = self.sfa_tile_shape_mk[0] * _fq_sfa_sg - _fq_ssfa_stage = mainloop_producer_state.index * Int32( - (self.sfa_tile_shape_mk[0] // 128) * 128 * _fq_sfa_sg + _fq_ssfa_stage = ( + mainloop_producer_state.index + * Int32( + (self.sfa_tile_shape_mk[0] // 128) * 128 + * _fq_sfa_sg + ) ) for _fq_si in cutlass.range_constexpr( (_fq_sfa_slots + self.num_threads_per_warp - 1) @@ -3778,21 +3833,23 @@ def kernel( ) if _fq_lin < Int32(_fq_sfa_slots): _fq_m = _fq_lin // Int32(_fq_sfa_sg) - _fq_sg_idx = _fq_lin - _fq_m * Int32(_fq_sfa_sg) + _fq_sg_idx = _fq_lin - _fq_m * Int32( + _fq_sfa_sg + ) _fq_sf_off = ( (_fq_m & Int32(31)) * Int32(16) + (_fq_m >> Int32(5)) * Int32(4) + _fq_sg_idx ) _fq_sb = Uint8( - ( - _fq_packed_scales - >> (Uint32(_fq_sg_idx) * Uint32(8)) - ) + (_fq_packed_scales + >> (Uint32(_fq_sg_idx) * Uint32(8))) & Uint32(0xFF) ) st_shared_u8( - ssfa_base_addr + _fq_ssfa_stage + _fq_sf_off, + ssfa_base_addr + + _fq_ssfa_stage + + _fq_sf_off, _fq_sb, ) @@ -4318,7 +4375,9 @@ def _compute_grid( num_ctas_mnl, cluster_shape_mnl, swizzle_size=( - 16 if tile_shape_mnk == (128, 128, 64) and not large_m_unroll else 1 + 16 + if tile_shape_mnk == (128, 128, 64) and not large_m_unroll + else 1 ), ) if cutlass.const_expr(split_k_slices > 1): @@ -4471,7 +4530,9 @@ def can_implement( tile_k = mxfp6_tile_k() if is_mxfp6_ab_dtype(ab_dtype) else 128 else: tile_k = sf_vec_size * 8 - return k % tile_k == 0 + if k % tile_k != 0: + return False + return True class _DenseGemmLaunch: @@ -5012,7 +5073,9 @@ def _make_runtime_pointers( alpha_tensor_gpu.data_ptr(), ) x_bf16_data_ptr = ( - x_bf16_tensor_gpu.data_ptr() if x_bf16_tensor_gpu is not None else 16 + x_bf16_tensor_gpu.data_ptr() + if x_bf16_tensor_gpu is not None + else 16 ) w_gscale_data_ptr = ( w_gscale_tensor_gpu.data_ptr() @@ -6830,7 +6893,6 @@ def _select_default_dense_gemm_plan( is_mxfp8: bool, is_mxfp6: bool = False, expected_m: Optional[int] = None, - select_swapped_output_storage: bool = False, ) -> _DenseGemmPlan: tile = _select_default_mma_tiler_mn( m, @@ -6841,22 +6903,11 @@ def _select_default_dense_gemm_plan( expected_m=expected_m, k=k, ) - plan = _DenseGemmPlan( + return _DenseGemmPlan( mma_tiler_mn=tile, load_path="tma", swap_ab=(not is_mxfp8 and not is_mxfp6 and tile[1] < 64), ) - if not (is_mxfp8 and select_swapped_output_storage): - return plan - - # Swapping operands reverses the logical MMA tile axes. Transpose the - # tuned default so expected-M, N, K, and SM-count policy remains intact. - # A 64x32 tile gives the qualified narrow-output path more independent N - # tiles than the square default without changing the public output shape. - swapped_tile = ( - (64, 32) if n < 64 or (n <= 256 and tile == (64, 64)) else (tile[1], tile[0]) - ) - return _DenseGemmPlan(swapped_tile, plan.load_path, True) def dense_gemm_fused_quant_a( @@ -7133,7 +7184,9 @@ def dense_gemm( is_mxfp6 = False mma_k = 32 tile_k = ( - 128 if block_fp8 else _select_mxfp8_tile_k(m, n, k, expected_m, sm_count) + 128 + if block_fp8 + else _select_mxfp8_tile_k(m, n, k, expected_m, sm_count) ) elif ab_dtype in ("float6_e3m2fn", "float6_e2m3fn"): is_mxfp8 = False @@ -7264,10 +7317,6 @@ def dense_gemm( if not (b_preexpanded or b_packed): b_torch = _expand_packed_mxfp6_ab(b_torch, k) c_cutlass_dtype = get_cutlass_dtype(c_dtype) - c_row_stride_bytes = n * c_cutlass_dtype.width // 8 - output_requires_swapped_store = (m > 1 or l > 1) and c_row_stride_bytes % 16 != 0 - use_default_mma_tiler = mma_tiler_mn is None - use_default_output_storage = mma_tiler_mn is None and swap_ab is None if mma_tiler_mn is None or load_path is None or swap_ab is None: default_plan = _select_default_dense_gemm_plan( m, @@ -7277,43 +7326,15 @@ def dense_gemm( is_mxfp8=is_mxfp8, is_mxfp6=is_mxfp6, expected_m=expected_m, - select_swapped_output_storage=( - use_default_output_storage - and l == 1 - and (n < 64 or output_requires_swapped_store) - ), ) if mma_tiler_mn is None: mma_tiler_mn = default_plan.mma_tiler_mn if load_path is None: load_path = default_plan.load_path if swap_ab is None: - if use_default_mma_tiler: - swap_ab = default_plan.swap_ab - else: - swap_ab = default_plan.swap_ab if mma_tiler_mn[1] < 64 else False + swap_ab = default_plan.swap_ab if mma_tiler_mn[1] < 64 else False assert load_path is not None assert swap_ab is not None - if l > 1 and swap_ab: - raise ValueError( - "swapped dense_gemm output storage supports L=1 only; pad N for " - f"grouped output, got L={l}, N={n}" - ) - if output_requires_swapped_store and not swap_ab: - remedy = ( - "pad N; swapped output storage is unsupported when L > 1" - if l > 1 - else "use a supported swapped plan or pad N" - ) - raise ValueError( - "the unswapped dense_gemm TMA epilogue requires a 16-byte-aligned " - f"C row stride, but N={n} and c_dtype={c_dtype!r} produce " - f"{c_row_stride_bytes} bytes; {remedy}" - ) - if is_mxfp8 and swap_ab: - # BK64 packed-scale staging requires the weight operand to remain in - # the unswapped 128-row slot. Swapped storage therefore uses BK128. - tile_k = 128 if is_mxfp6: # Only the unswapped single-slice TMA mainloop is wired for the FP6 # byte-container path; fail loudly instead of silently miscomputing. @@ -7321,7 +7342,8 @@ def dense_gemm( raise ValueError("MX-FP6 dense_gemm does not support swap_ab") if load_path != "tma": raise ValueError( - f"MX-FP6 dense_gemm only supports load_path='tma', got {load_path!r}" + "MX-FP6 dense_gemm only supports load_path='tma', got " + f"{load_path!r}" ) if _quantized_c is not None: raise ValueError("MX-FP6 dense_gemm does not support quantized C output") @@ -7653,7 +7675,9 @@ def dense_gemm( alpha = _cached_alpha_one(a_torch.device) t0 = time.perf_counter() if _B12X_TIMING else 0.0 - cache_before = _get_compiled_dense_gemm.cache_info() if _B12X_TIMING else None + cache_before = ( + _get_compiled_dense_gemm.cache_info() if _B12X_TIMING else None + ) t_compiled = t0 kernel_c_dtype_name = ( "float32" if split_k_output and not split_k_atomic_bf16 else c_dtype diff --git a/b12x/_lib/intrinsics.py b/b12x/_lib/intrinsics.py index d4ca937e8..7b4c65e55 100644 --- a/b12x/_lib/intrinsics.py +++ b/b12x/_lib/intrinsics.py @@ -1028,6 +1028,39 @@ def ldmatrix_m8n8x4_right_half_b16( return Uint32(r0), Uint32(r1) +@dsl_user_op +def ldmatrix_m16n16x2_trans_b8( + smem_addr: Int32, *, loc=None, ip=None +) -> Tuple[Uint32, Uint32, Uint32, Uint32]: + """Issue ``ldmatrix.sync.aligned.m16n16.x2.trans.shared.b8`` (sm_120a). + + Two 16x16 byte matrices; lanes 0-15 supply the 16-byte row addresses of + matrix 0 and lanes 16-31 those of matrix 1. Each returned register holds + four bytes of one *column*: lane ``i`` receives column ``i // 4`` (r0, r2) + and column ``i // 4 + 8`` (r1, r3) at rows ``4 * (i % 4) .. +3`` of matrix + 0 (r0, r1) and matrix 1 (r2, r3). With rows = K (tokens) and columns = N + (dims) this is exactly the ``mma.m16n8k32`` B fragment: for dims + ``[d, d+8)`` b0 = r0 (K 0-15), b1 = r2 (K 16-31); for dims ``[d+8, d+16)`` + b0 = r1, b1 = r3. Row addresses must be 16-byte aligned. + """ + result = llvm.inline_asm( + llvm.StructType.get_literal([T.i32(), T.i32(), T.i32(), T.i32()]), + [Int32(smem_addr).ir_value(loc=loc, ip=ip)], + "ldmatrix.sync.aligned.m16n16.x2.trans.shared.b8 {$0, $1, $2, $3}, [$4];", + "=r,=r,=r,=r,r", + has_side_effects=True, + is_align_stack=False, + asm_dialect=llvm.AsmDialect.AD_ATT, + loc=loc, + ip=ip, + ) + r0 = llvm.extractvalue(T.i32(), result, [0], loc=loc, ip=ip) + r1 = llvm.extractvalue(T.i32(), result, [1], loc=loc, ip=ip) + r2 = llvm.extractvalue(T.i32(), result, [2], loc=loc, ip=ip) + r3 = llvm.extractvalue(T.i32(), result, [3], loc=loc, ip=ip) + return Uint32(r0), Uint32(r1), Uint32(r2), Uint32(r3) + + @dsl_user_op def ldmatrix_m8n8x4_trans_b16( smem_addr: Int32, *, loc=None, ip=None @@ -1546,6 +1579,37 @@ def red_add_global_f32(addr: Int64, val: Float32, *, loc=None, ip=None): ) +@dsl_user_op +def red_add_global_v4_f32( + addr: Int64, + val0: Float32, + val1: Float32, + val2: Float32, + val3: Float32, + *, + loc=None, + ip=None, +): + """Reduce-add four FP32 elements at a 16-byte-aligned address.""" + llvm.inline_asm( + None, + [ + Int64(addr).ir_value(loc=loc, ip=ip), + Float32(val0).ir_value(loc=loc, ip=ip), + Float32(val1).ir_value(loc=loc, ip=ip), + Float32(val2).ir_value(loc=loc, ip=ip), + Float32(val3).ir_value(loc=loc, ip=ip), + ], + "red.relaxed.gpu.global.v4.f32.add [$0], {$1, $2, $3, $4};", + "l,f,f,f,f", + has_side_effects=True, + is_align_stack=False, + asm_dialect=llvm.AsmDialect.AD_ATT, + loc=loc, + ip=ip, + ) + + @dsl_user_op def cvt_bf16x2_to_f16x2(packed: Uint32, *, loc=None, ip=None) -> Uint32: """Convert a u32 holding two bf16 (lo, hi) into an f16x2 u32 (lo, hi).""" diff --git a/b12x/attention/__init__.py b/b12x/attention/__init__.py index 40e6066fa..c5d4a69ce 100644 --- a/b12x/attention/__init__.py +++ b/b12x/attention/__init__.py @@ -2,9 +2,8 @@ - ``paged``: paged-KV self-attention (decode + extend, FP8 KV, MSA block-sparse variant) with on-device graph-replay metadata staging. -- ``dense_mla``: dense compressed-cache MLA with strided physical records and - optional causal sliding-window masking. -- ``sparse_mla``: top-k-selected MLA, including strided physical records. +- ``dense_mla``: dense compressed-cache MLA for Kimi K3 geometry. +- ``sparse_mla``: top-k-selected MLA decode/extend (DeepSeek-V3.2 / GLM NSA). - ``compressed_mla``: MLA decode directly from compressed KV pages (DSV4). - ``nsa_indexer``: the NSA index stage — quantize -> score -> select. - ``varlen``: contiguous batched/varlen attention (reduced-assurance tier). diff --git a/b12x/attention/_shared/contiguous/api.py b/b12x/attention/_shared/contiguous/api.py index b293c49ae..3b2bf678d 100644 --- a/b12x/attention/_shared/contiguous/api.py +++ b/b12x/attention/_shared/contiguous/api.py @@ -75,12 +75,6 @@ def _lse_shape(q_shape: tuple[int, ...]) -> tuple[int, ...]: return (batch, q_heads, seqlen_q) -def _output_shape( - q_shape: tuple[int, ...], v_shape: tuple[int, ...] -) -> tuple[int, ...]: - return (*q_shape[:-1], int(v_shape[-1])) - - def _seq_dims(shape: tuple[int, ...]) -> tuple[tuple[int, ...], int, int, int]: if len(shape) == 3: seqlen, num_heads, head_dim = shape @@ -230,10 +224,10 @@ def _validate_forward_inputs( batch_v, _, v_heads, v_head_dim = _seq_dims(v_shape) if batch_q != batch_k or batch_q != batch_v: raise ValueError("q, k, and v must have matching batch dimensions") - if q_head_dim != k_head_dim: - raise ValueError("q and k must have matching head dimensions") - if v_head_dim <= 0: - raise ValueError("v head dimension must be positive") + if q_head_dim != k_head_dim or q_head_dim != v_head_dim: + raise ValueError( + "q, k, and v must have matching head dimensions in the initial path" + ) if kv_heads != v_heads: raise ValueError("k and v must have the same number of KV heads") if q_heads % kv_heads != 0: @@ -736,13 +730,13 @@ def __init__( self._q_shape = q_shape self._k_shape = k_shape self._v_shape = v_shape - self._o_shape = _output_shape(q_shape, v_shape) + self._o_shape = q_shape self._lse_shape = _lse_shape(q_shape) self._attention_sink_bias_shape = (q_shape[-2],) self._q_stride = _contiguous_stride(q_shape) self._k_stride = _contiguous_stride(k_shape) self._v_stride = _contiguous_stride(v_shape) - self._o_stride = _contiguous_stride(self._o_shape) + self._o_stride = _contiguous_stride(q_shape) self._lse_stride = _contiguous_stride(self._lse_shape) self._attention_sink_bias_stride = _contiguous_stride( self._attention_sink_bias_shape @@ -876,7 +870,7 @@ def __init__( self._v_shape = v_shape self._cu_seqlens_q_shape = cu_seqlens_q_shape self._cu_seqlens_k_shape = cu_seqlens_k_shape - self._o_shape = _output_shape(q_shape, v_shape) + self._o_shape = q_shape self._lse_shape = _lse_shape(q_shape) self._attention_sink_bias_shape = (q_shape[-2],) self._q_stride = _contiguous_stride(q_shape) @@ -884,7 +878,7 @@ def __init__( self._v_stride = _contiguous_stride(v_shape) self._cu_seqlens_q_stride = _contiguous_stride(cu_seqlens_q_shape) self._cu_seqlens_k_stride = _contiguous_stride(cu_seqlens_k_shape) - self._o_stride = _contiguous_stride(self._o_shape) + self._o_stride = _contiguous_stride(q_shape) self._lse_stride = _contiguous_stride(self._lse_shape) self._attention_sink_bias_stride = _contiguous_stride( self._attention_sink_bias_shape @@ -942,7 +936,9 @@ def __init__( # existing unpacked-head kernel as the static fallback for those # shapes; this choice depends only on plan geometry and is graph # capture safe. - pack_gqa=(qhead_per_kvhead != 1 and tile_m % qhead_per_kvhead == 0), + pack_gqa=( + qhead_per_kvhead != 1 and tile_m % qhead_per_kvhead == 0 + ), tile_m=tile_m, tile_n=tile_n, ) @@ -1396,11 +1392,9 @@ def _validate_attention_output_lse( lse: torch.Tensor, plan: AttentionPlan | VarlenAttentionPlan, ) -> None: - expected_output_shape = _output_shape(plan.q_shape, plan.v_shape) - if output.shape != expected_output_shape: + if output.shape != plan.q_shape: raise ValueError( - "attention output must have shape " - f"{expected_output_shape}, got {tuple(output.shape)}" + f"attention output must have shape {plan.q_shape}, got {tuple(output.shape)}" ) if output.device != plan.device: raise ValueError( @@ -1428,13 +1422,12 @@ def _validate_attention_output_lse( def _attention_scratch_layout( *, q_shape: tuple[int, ...], - v_shape: tuple[int, ...], dtype: torch.dtype, ) -> _AttentionScratchLayout: cursor = 0 cursor = _align_up(cursor, _ARENA_ALIGN_BYTES) output_offset_bytes = cursor - cursor += _shape_numel(_output_shape(q_shape, v_shape)) * _dtype_nbytes(dtype) + cursor += _shape_numel(q_shape) * _dtype_nbytes(dtype) cursor = _align_up(cursor, _ARENA_ALIGN_BYTES) lse_offset_bytes = cursor cursor += _shape_numel(_lse_shape(q_shape)) * _dtype_nbytes(torch.float32) @@ -1493,7 +1486,7 @@ def _attention_scratch_views_from_arena( output = _arena_view( arena, offset_bytes=layout.output_offset_bytes, - shape=_output_shape(plan.q_shape, plan.v_shape), + shape=plan.q_shape, dtype=plan.dtype, ) lse = _arena_view( @@ -1547,7 +1540,7 @@ def _varlen_attention_scratch_views_from_arena( output = _arena_view( arena, offset_bytes=layout.output_offset_bytes, - shape=_output_shape(plan.q_shape, plan.v_shape), + shape=plan.q_shape, dtype=plan.dtype, ) lse = _arena_view( @@ -1962,9 +1955,7 @@ def _build_varlen_attention_binding_from_views( def plan_attention_scratch(plan: AttentionPlan) -> AttentionScratchPlan: - layout = _attention_scratch_layout( - q_shape=plan.q_shape, v_shape=plan.v_shape, dtype=plan.dtype - ) + layout = _attention_scratch_layout(q_shape=plan.q_shape, dtype=plan.dtype) return AttentionScratchPlan( plan=plan, _layout=layout, @@ -1981,9 +1972,7 @@ def plan_attention_scratch(plan: AttentionPlan) -> AttentionScratchPlan: def plan_varlen_attention_scratch( plan: VarlenAttentionPlan, ) -> VarlenAttentionScratchPlan: - layout = _attention_scratch_layout( - q_shape=plan.q_shape, v_shape=plan.v_shape, dtype=plan.dtype - ) + layout = _attention_scratch_layout(q_shape=plan.q_shape, dtype=plan.dtype) return VarlenAttentionScratchPlan( plan=plan, _layout=layout, @@ -1999,11 +1988,7 @@ def plan_varlen_attention_scratch( def allocate_attention_workspace_for_plan(plan: AttentionPlan) -> AttentionWorkspace: """Allocate reusable scratch for one exact contiguous attention plan.""" - output = torch.empty( - _output_shape(plan.q_shape, plan.v_shape), - dtype=plan.dtype, - device=plan.device, - ) + output = torch.empty(plan.q_shape, dtype=plan.dtype, device=plan.device) lse = torch.empty(_lse_shape(plan.q_shape), dtype=torch.float32, device=plan.device) return AttentionWorkspace( q_shape=plan.q_shape, @@ -2027,11 +2012,7 @@ def allocate_varlen_attention_workspace_for_plan( plan: VarlenAttentionPlan, ) -> VarlenAttentionWorkspace: """Allocate reusable scratch for one exact packed varlen attention plan.""" - output = torch.empty( - _output_shape(plan.q_shape, plan.v_shape), - dtype=plan.dtype, - device=plan.device, - ) + output = torch.empty(plan.q_shape, dtype=plan.dtype, device=plan.device) lse = torch.empty(_lse_shape(plan.q_shape), dtype=torch.float32, device=plan.device) return VarlenAttentionWorkspace( q_shape=plan.q_shape, diff --git a/b12x/attention/_shared/mla/api.py b/b12x/attention/_shared/mla/api.py index 9dc229ad5..cc4fe8e57 100644 --- a/b12x/attention/_shared/mla/api.py +++ b/b12x/attention/_shared/mla/api.py @@ -286,9 +286,10 @@ def _validate_split_workspace_views( if tmp_output.device != workspace.device or tmp_lse.device != workspace.device: raise ValueError("split MLA scratch buffers must be on the workspace device") - if tmp_output.dtype != workspace.dtype: + if tmp_output.dtype != workspace.dtype and tmp_output.dtype != torch.float32: raise TypeError( - f"split MLA tmp_output dtype {tmp_output.dtype} does not match workspace dtype {workspace.dtype}" + f"split MLA tmp_output dtype {tmp_output.dtype} must match workspace " + f"dtype {workspace.dtype} or be torch.float32" ) if tmp_lse.dtype != torch.float32: raise TypeError( @@ -389,6 +390,7 @@ def sparse_mla_decode_forward( scale_format: int | None = None, fp8_rope: bool | None = None, latent_scale_per_token: bool = False, + split_policy: Literal["static", "balanced"] = "static", ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: q_all, page_table_1, cache_seqlens_int32, nsa_cache_seqlens_int32, workspace = ( _resolve_sparse_mla_binding( @@ -421,6 +423,7 @@ def sparse_mla_decode_forward( scale_format=scale_format, fp8_rope=fp8_rope, latent_scale_per_token=latent_scale_per_token, + split_policy=split_policy, ) @@ -497,7 +500,12 @@ def _run_sparse_mla( scale_format: int | None = None, fp8_rope: bool | None = None, latent_scale_per_token: bool = False, + split_policy: Literal["static", "balanced"] = "static", ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: + if split_policy not in ("static", "balanced"): + raise ValueError( + f"split_policy must be 'static' or 'balanced', got {split_policy!r}" + ) if q_all.ndim != 3: raise ValueError(f"q_all must be rank-3, got {tuple(q_all.shape)}") if kv_cache.ndim != 3: @@ -557,9 +565,16 @@ def _run_sparse_mla( "nsa_cache_seqlens_int32 device " f"{active_token_counts.device} does not match workspace device {workspace.device}" ) - if q_all.dtype != workspace.dtype: + from b12x.attention.sparse_mla._scratch import ( + PACKED_QUERY_RECORD_BYTES, + is_packed_query, + ) + + q_packed = is_packed_query(q_all, head_dim=int(workspace.head_dim)) + if q_all.dtype != workspace.dtype and not q_packed: raise ValueError( - f"q_all dtype {q_all.dtype} does not match workspace dtype {workspace.dtype}" + f"q_all dtype {q_all.dtype} does not match workspace dtype {workspace.dtype} " + f"(or the uint8 packed {PACKED_QUERY_RECORD_BYTES}-byte query record)" ) if kv_cache.dtype != workspace.kv_dtype: raise ValueError( @@ -678,12 +693,18 @@ def _run_sparse_mla( raise ValueError( f"q_all num_heads {q_all.shape[1]} does not match workspace num_q_heads {workspace.num_q_heads}" ) - if q_all.shape[-1] != workspace.head_dim: + if q_all.shape[-1] != workspace.head_dim and not q_packed: raise ValueError( f"q_all head_dim {q_all.shape[-1]} does not match workspace head_dim {workspace.head_dim}" ) + if q_packed and ( + not _sm120_route or workspace.mode in ("extend", "verify", "draft_extend") + ): + raise ValueError( + "packed query records require the SM120 sparse MLA decode kernel path" + ) if _sm120_route: - q_head_dim = int(q_all.shape[-1]) + q_head_dim = int(workspace.head_dim) if q_packed else int(q_all.shape[-1]) if q_head_dim != _MLA_UNIFIED_GLM_Q_HEAD_DIM: raise ValueError( f"SM120 sparse MLA decode requires the GLM_NSA contract " @@ -724,6 +745,7 @@ def _run_sparse_mla( scale_format_override=scale_format_for_call, fp8_rope_override=fp8_rope_for_call, latent_scale_per_token=latent_scale_per_token, + split_policy=split_policy, ) if _is_cuda_graph_capture_active(q_all.device): raise RuntimeError( diff --git a/b12x/attention/_shared/mla/decode_math.py b/b12x/attention/_shared/mla/decode_math.py index f364046ba..449b1a387 100644 --- a/b12x/attention/_shared/mla/decode_math.py +++ b/b12x/attention/_shared/mla/decode_math.py @@ -71,8 +71,11 @@ ld_shared_f32, ld_shared_u8_offset, ld_shared_u32, + ld_global_nc_v4_u32, ldmatrix_m8n8x2_b16, ldmatrix_m8n8x4_b16, + ldmatrix_m16n16x2_trans_b8, + st_shared_v4_u32, mma_m16n8k16_f32_bf16, mma_m16n8k32_f32_e4m3, mxfp8_mma_m16n8k32_f32_e4m3, @@ -484,6 +487,83 @@ def s0_quantize_q_to_smem( cute.arch.barrier(**bar_kw) +@cute.jit +def s0_load_packed_q_to_smem( + q_token: cute.Tensor, # (NUM_HEADS, 656) u8 view for this token (packed query) + q_fp8_base_addr: Int32, # u32 smem addr of q_fp8 (HPB x Q_NOPE_STRIDE) + q_sc_view: cute.Tensor, # smem fp32 view (HPB*NUM_SCALES,) -- pow2 scales + q_rope_base_addr: Int32, # u32 smem addr of q_rope (HPB x q_rope_stride bf16) + head_base: Int32, # first head index of this CTA (h_start) + valid_hpb: Int32, # number of valid heads (<= HPB) + tid: Int32, # flat thread id in [0, MATH_THREADS) + *, + d_nope: cutlass.Constexpr, # 512 + d_rope: cutlass.Constexpr, # 64 + num_scales: cutlass.Constexpr, # 4 + hpb: cutlass.Constexpr, # 16 + q_nope_stride: cutlass.Constexpr, # 528 + q_rope_stride: cutlass.Constexpr, # 64 (+ optional bank-layout pad) + num_threads: cutlass.Constexpr, # 256 + barrier_id: cutlass.Constexpr, +): + """S0 (packed query): copy a pre-quantized query record into the Q stages. + + The record per head is ``[512 B E4M3 nope][16 B: NUM_SCALES fp32 pow2 + scales][128 B bf16 rope]``, produced by the host with the same algorithm + as :func:`s0_quantize_q_to_smem` (per-tile absmax, ``max(amax, 1e-4) / + FP8_MAX`` rounded up to a power of two, ``q * (1 / scale)`` clamped and + converted with ``cvt.rn.satfinite.e4m3``). The stages therefore hold the + same bytes S0 would have written, so the attention result is bit-identical + to the bf16-query path. Each thread copies 16-byte chunks; invalid tail + heads are zero-filled. + """ + bar_kw = dict(barrier_id=barrier_id, number_of_threads=num_threads) + nope_chunks = d_nope // 16 + scale_chunks = (num_scales * 4) // 16 + rope_chunks = (d_rope * 2) // 16 + chunks_per_head = nope_chunks + scale_chunks + rope_chunks + q_sc_addr = shared_ptr_to_u32(q_sc_view.iterator) + i = tid + while i < Int32(hpb * chunks_per_head): + h = i // Int32(chunks_per_head) + c = i - h * Int32(chunks_per_head) + v0 = Uint32(0) + v1 = Uint32(0) + v2 = Uint32(0) + v3 = Uint32(0) + if h < valid_hpb: + src = get_ptr_as_int64( + q_token, cute.crd2idx((head_base + h, c * Int32(16)), q_token.layout) + ) + v0, v1, v2, v3 = ld_global_nc_v4_u32(src) + if c < Int32(nope_chunks): + st_shared_v4_u32( + q_fp8_base_addr + h * Int32(q_nope_stride) + c * Int32(16), v0, v1, v2, v3 + ) + elif c < Int32(nope_chunks + scale_chunks): + st_shared_v4_u32( + q_sc_addr + + h * Int32(num_scales * 4) + + (c - Int32(nope_chunks)) * Int32(16), + v0, + v1, + v2, + v3, + ) + else: + st_shared_v4_u32( + q_rope_base_addr + + h * Int32(q_rope_stride * 2) + + (c - Int32(nope_chunks + scale_chunks)) * Int32(16), + v0, + v1, + v2, + v3, + ) + i += Int32(num_threads) + cute.arch.barrier(**bar_kw) + + @cute.jit def s0_load_q_bf16_to_smem( q_token: cute.Tensor, # (NUM_HEADS, D_QK) bf16 view for this token @@ -1019,9 +1099,14 @@ def s2_qk_rope_bf16( q_rope_stride: cutlass.Constexpr, valid_hpb: cutlass.Constexpr = 16, fp8_rope: cutlass.Constexpr = False, + kv_rope_stride_bytes: cutlass.Constexpr = 0, # 0 -> D_ROPE*2 (separate rope stage) ): """S2: accumulate Q_rope . K_rope into qk[0..3] via D_ROPE/16=4 bf16 MMAs. + ``kv_rope_stride_bytes`` is the byte stride between consecutive tokens' + rope rows: the default ``D_ROPE * 2`` for the separate rope stage, or the + packed record stride when the rope is staged inline after the nope bytes. + A (Q-rope, 16x16 bf16) via ldmatrix.x4 (ldmatrix_load_A_bf16); B (K-rope) via per-lane scalar u32 reads -- the N-major rope smem layout can't feed ldmatrix.x2 here (decode_dsv4 :401-425). @@ -1056,7 +1141,8 @@ def s2_qk_rope_bf16( d_rope=d_rope, ) else: - row_byte = entry * Int32(d_rope * 2) + ko * Int32(2) + rope_row_stride = kv_rope_stride_bytes if kv_rope_stride_bytes else d_rope * 2 + row_byte = entry * Int32(rope_row_stride) + ko * Int32(2) b0 = _ld_u32(kv_rope_base_addr, row_byte + tid * Int32(2) * Int32(2)) b1 = _ld_u32( kv_rope_base_addr, row_byte + (tid * Int32(2) + Int32(8)) * Int32(2) @@ -1431,9 +1517,27 @@ def s4_online_softmax( num_threads: cutlass.Constexpr, # 256 barrier_id: cutlass.Constexpr, # math-only named-barrier slot n_acc_tiles: cutlass.Constexpr = None, # len(acc_nope); defaults to n_v_chunks + skip_unit_rescale: cutlass.Constexpr = False, # warp-uniform skip when alpha == 1 + return_state: cutlass.Constexpr = False, # also return acc/max/sum containers ): """S4: per-warp + cross-warp max/sum, exp2(qk-max), cross-chunk rescale. + ``skip_unit_rescale`` skips the accumulator rescale for a chunk in which + every lane's cross-chunk factor is exactly 1.0 (the running maximum of + both of the lane's heads did not change), decided with a warp vote so the + branch is uniform. Multiplying by exactly 1.0 is the identity, so the + result is bit-identical; after the first few chunks most chunks skip. + + ``return_state`` additionally returns ``(acc_nope, acc_rope, global_max, + global_sum)`` and the unified decode kernel's chunk loop rebinds them from + the return value. The accumulator rescale below sits inside ``if rescale:``, + which the DSL lowers as a region even when ``rescale`` is the Python + constant ``True``; nested-list element assignments made inside such a + region do not reach the caller's list objects, so a caller that relies on + in-place mutation keeps un-rescaled accumulators whenever the running + maximum rises in a later chunk (research + ``fp8-ds-mla-perf-20260905/dsl_if_region_probe.py``). + Returns ``(p, warp_rescale0, warp_rescale1)``; mutates ``acc_nope`` / ``acc_rope`` (rescaled by the cross-chunk alpha) and ``global_max`` / ``global_sum`` in place. ``p`` is exp2(qk - local_max) (the warp-local-frame @@ -1523,15 +1627,20 @@ def s4_online_softmax( _n_acc = cutlass.const_expr(n_acc_tiles if n_acc_tiles is not None else n_v_chunks) if cutlass.const_expr(not is_first_chunk): - for vc in cutlass.range_constexpr(_n_acc): - acc_nope[vc][0] = acc_nope[vc][0] * alpha0 - acc_nope[vc][1] = acc_nope[vc][1] * alpha0 - acc_nope[vc][2] = acc_nope[vc][2] * alpha1 - acc_nope[vc][3] = acc_nope[vc][3] * alpha1 - acc_rope[0] = acc_rope[0] * alpha0 - acc_rope[1] = acc_rope[1] * alpha0 - acc_rope[2] = acc_rope[2] * alpha1 - acc_rope[3] = acc_rope[3] * alpha1 + rescale = True + if cutlass.const_expr(skip_unit_rescale): + unit = (alpha0 == Float32(1.0)) & (alpha1 == Float32(1.0)) + rescale = not cute.arch.vote_all_sync(unit) + if rescale: + for vc in cutlass.range_constexpr(_n_acc): + acc_nope[vc][0] = acc_nope[vc][0] * alpha0 + acc_nope[vc][1] = acc_nope[vc][1] * alpha0 + acc_nope[vc][2] = acc_nope[vc][2] * alpha1 + acc_nope[vc][3] = acc_nope[vc][3] * alpha1 + acc_rope[0] = acc_rope[0] * alpha0 + acc_rope[1] = acc_rope[1] * alpha0 + acc_rope[2] = acc_rope[2] * alpha1 + acc_rope[3] = acc_rope[3] * alpha1 global_sum[0] = global_sum[0] * alpha0 + block_local_sum0 * block_rescale0 global_sum[1] = global_sum[1] * alpha1 + block_local_sum1 * block_rescale1 else: @@ -1540,6 +1649,11 @@ def s4_online_softmax( global_max[0] = new_gmax0 global_max[1] = new_gmax1 + if cutlass.const_expr(return_state): + # The rescale inside ``if rescale:`` is invisible to the caller's list + # objects (see the docstring); hand the containers back explicitly so + # the chunk loop accumulates the next chunk onto rescaled values. + return p, warp_rescale0, warp_rescale1, acc_nope, acc_rope, global_max, global_sum return p, warp_rescale0, warp_rescale1 @@ -1564,6 +1678,7 @@ def s4_online_softmax_glm_h8_swap_ab( barrier_id: cutlass.Constexpr, rope_tiles_per_warp: cutlass.Constexpr = 0, barrier_threads: cutlass.Constexpr = 0, # barrier width override (0 -> num_threads) + return_state: cutlass.Constexpr = False, # also return acc/max/sum containers ): """Online softmax for the swapped 16-candidate x 8-head score tile. @@ -1674,6 +1789,8 @@ def s4_online_softmax_glm_h8_swap_ab( global_sum[1] = global_sum[1] * alpha1 + block_sum1 * block_rescale1 global_max[0] = new_gmax0 global_max[1] = new_gmax1 + if cutlass.const_expr(return_state): + return p, warp_rescale0, warp_rescale1, acc_nope, acc_rope, global_max, global_sum return p, warp_rescale0, warp_rescale1 @@ -1761,6 +1878,8 @@ def s6_xv_nope( sm_p_full_addr: Int32 = None, # NVFP4 BF16 PV only sm_p_stride: cutlass.Constexpr = 0, # NVFP4: bf16 elems per sm_p row (0 -> BI) latent_scale_per_token: cutlass.Constexpr = False, # NVFP4_E4M3 only + v_ldsm_b8: cutlass.Constexpr = False, # GLM 2-pass: V B-fragments via ldmatrix.b8 + w_hw_dequant: cutlass.Constexpr = False, # GLM 2-pass: residual from cvt.f16x2.e4m3x2 ): """S6: accumulate W . V_nope into acc_nope[vc*NT+nt][0..3] via PLAIN fp8 MMAs (14 DSV4 / 16 GLM = N_V_CHUNKS * NT_PER_WARP_XV * (BI/32)). @@ -1798,7 +1917,27 @@ def s6_xv_nope( ``scale_format == NVFP4_E4M3`` (const_expr) bypasses the FP8 W machinery entirely: BF16 probabilities staged in sm_p (``sm_p_full_addr``) are the A operand and V is dequantized in registers from packed E2M1 + E4M3 group-16 - scales (see ``s6_xv_nope_nvfp4_bf16``).""" + scales (see ``s6_xv_nope_nvfp4_bf16``). + + ``v_ldsm_b8`` (const_expr, GLM 2-pass branch only, ``nt_per_warp_xv == 2``) + loads the V B-fragments with ``ldmatrix.m16n16.x2.trans.b8`` straight from + the row-major smem stage instead of synthesizing them from 8 LDS.32 + 6 + PRMT per (n-tile, k-step) (``_d2_load_b_fp8``): one ldmatrix per k-step + yields both n-tiles, loaded once per V chunk and reused by both W passes. + Each warp then owns 16 consecutive dims per V chunk + (``vc * V_CHUNK + warp * 16 + nt * 8``) instead of two 8-dim tiles 64 dims + apart; the epilogue must use the same mapping (``warp_contiguous_dims``). + The loaded bytes and the MMA sequence are unchanged, so the result is + bit-identical. The same flag quantizes each lane's candidate pair with one + ``cvt.rn.satfinite.e4m3x2.f32`` and stores it as one 16-bit word (the + per-value conversion produces the same bytes). + + ``w_hw_dequant`` (const_expr, GLM 2-pass branch only) reconstructs the + HIGH byte for the residual with ``cvt.rn.f16x2.e4m3x2`` (exact for every + E4M3 value) instead of the scalar software expansion, which mis-decodes + E4M3 subnormals and -0 as normal numbers. The LOW byte therefore changes + for W values below 2^-6 of the head's chunk maximum; results are not + bit-identical to the software path but closer to the fp32 reference.""" if cutlass.const_expr(scale_format == 2): return s6_xv_nope_nvfp4_bf16( acc_nope, @@ -2011,6 +2150,15 @@ def _vsc(cand: Int32, vc: int): # SEPARATE const_expr branch, so DSV4 (above) is untouched. for vc in cutlass.range_constexpr(n_v_chunks): w_fp8_addr = w_fp8_base_addr + Int32(vc & 1) * Int32(hpb * w_fp8_stride) + # v_ldsm_b8: the HIGH byte always uses slot 0 and the LOW residual slot + # 1. Every warp writes HIGH(vc+1) only after the barrier that follows + # its LOW(vc) store, i.e. after every warp finished the HIGH(vc) MMAs, + # and LOW(vc+1) only after the barrier that follows HIGH(vc+1), i.e. + # after every warp finished the LOW(vc) MMAs; so neither store needs + # the extra serialization barrier of the vc-parity slot scheme. + if cutlass.const_expr(v_ldsm_b8): + w_fp8_hi_addr = w_fp8_base_addr + w_fp8_lo_addr = w_fp8_base_addr + Int32(hpb * w_fp8_stride) si0 = Float32(1.0) / w_head_sc_view[Int32(vc) * Int32(hpb) + gid] sc0 = w_head_sc_view[Int32(vc) * Int32(hpb) + gid] if cutlass.const_expr(hi): @@ -2031,42 +2179,98 @@ def _vsc(cand: Int32, vc: int): [Float32(0.0), Float32(0.0), Float32(0.0), Float32(0.0)] for _ in range(nt_per_warp_xv) ] + if cutlass.const_expr(v_ldsm_b8): + # V B-fragments for this warp's 16 consecutive dims, both k-steps, + # loaded once and shared by the two W passes. Lane ``l`` supplies + # the 16-byte-aligned row address of token ``ko + l``. + v_dim_base = Int32(vc) * Int32(v_chunk) + warp_id * Int32(8 * nt_per_warp_xv) + v_frag = [] + for kstep in cutlass.range_constexpr(bi // 32): + v_row_addr = ( + kv_fp8_base_addr + + (Int32(kstep) * Int32(32) + lane) * Int32(kv_smem_stride) + + v_dim_base + ) + v_frag.append(ldmatrix_m16n16x2_trans_b8(v_row_addr)) for wpass in cutlass.range_constexpr(2): - if cutlass.const_expr(wpass > 0): - # serialize: the prev pass's MMA reads of w_fp8 must finish before - # we overwrite it with the residual bytes (same double-buffer slot). - cute.arch.barrier(**bar_kw) - # LOW residual = e4m3(Wn - dequant(hi_byte)); halves W's quant error. - f00 = _quant_e4m3_residual_byte(wn00) - f01 = _quant_e4m3_residual_byte(wn01) + if cutlass.const_expr(v_ldsm_b8): + w_fp8_addr = w_fp8_lo_addr if wpass > 0 else w_fp8_hi_addr + if cutlass.const_expr(v_ldsm_b8): + # Pairwise quantization of (cand_e0, cand_e1): one conversion + # and one 16-bit store per row (cand_e0 is even). + if cutlass.const_expr(wpass == 0): + vc00 = fmax_f32(Float32(_FP8_MIN), fmin_f32(Float32(_FP8_MAX), wn00)) + vc01 = fmax_f32(Float32(_FP8_MIN), fmin_f32(Float32(_FP8_MAX), wn01)) + fhi2_0 = _cvt_f32x2_to_e4m3x2(vc00, vc01) + word0 = fhi2_0 + if cutlass.const_expr(hi): + vc10 = fmax_f32(Float32(_FP8_MIN), fmin_f32(Float32(_FP8_MAX), wn10)) + vc11 = fmax_f32(Float32(_FP8_MIN), fmin_f32(Float32(_FP8_MAX), wn11)) + fhi2_1 = _cvt_f32x2_to_e4m3x2(vc10, vc11) + word1 = fhi2_1 + else: + if cutlass.const_expr(w_hw_dequant): + h00, h01 = f16x2_to_f32x2(_cvt_e4m3x2_to_f16x2(fhi2_0)) + else: + h00 = fp8_e4m3_to_f32(fhi2_0 & Uint32(0xFF)) + h01 = fp8_e4m3_to_f32((fhi2_0 >> Uint32(8)) & Uint32(0xFF)) + r00 = fmax_f32(Float32(_FP8_MIN), fmin_f32(Float32(_FP8_MAX), vc00 - h00)) + r01 = fmax_f32(Float32(_FP8_MIN), fmin_f32(Float32(_FP8_MAX), vc01 - h01)) + word0 = _cvt_f32x2_to_e4m3x2(r00, r01) + if cutlass.const_expr(hi): + if cutlass.const_expr(w_hw_dequant): + h10, h11 = f16x2_to_f32x2(_cvt_e4m3x2_to_f16x2(fhi2_1)) + else: + h10 = fp8_e4m3_to_f32(fhi2_1 & Uint32(0xFF)) + h11 = fp8_e4m3_to_f32((fhi2_1 >> Uint32(8)) & Uint32(0xFF)) + r10 = fmax_f32(Float32(_FP8_MIN), fmin_f32(Float32(_FP8_MAX), vc10 - h10)) + r11 = fmax_f32(Float32(_FP8_MIN), fmin_f32(Float32(_FP8_MAX), vc11 - h11)) + word1 = _cvt_f32x2_to_e4m3x2(r10, r11) + _st_shared_u16(w_fp8_addr + gid * Int32(w_fp8_stride) + cand_e0, word0) if cutlass.const_expr(hi): - f10 = _quant_e4m3_residual_byte(wn10) - f11 = _quant_e4m3_residual_byte(wn11) + _st_shared_u16( + w_fp8_addr + (gid + Int32(8)) * Int32(w_fp8_stride) + cand_e0, + word1, + ) else: - f00 = _quant_e4m3_byte(wn00) - f01 = _quant_e4m3_byte(wn01) - if cutlass.const_expr(hi): - f10 = _quant_e4m3_byte(wn10) - f11 = _quant_e4m3_byte(wn11) - st_shared_u8( - w_fp8_addr + gid * Int32(w_fp8_stride) + cand_e0, f00.to(cutlass.Uint8) - ) - st_shared_u8( - w_fp8_addr + gid * Int32(w_fp8_stride) + cand_e1, f01.to(cutlass.Uint8) - ) - if cutlass.const_expr(hi): + if cutlass.const_expr(wpass > 0): + # serialize: the prev pass's MMA reads of w_fp8 must finish + # before we overwrite it with the residual bytes (same + # double-buffer slot). + cute.arch.barrier(**bar_kw) + # LOW residual = e4m3(Wn - dequant(hi_byte)); halves W's quant error. + f00 = _quant_e4m3_residual_byte(wn00) + f01 = _quant_e4m3_residual_byte(wn01) + if cutlass.const_expr(hi): + f10 = _quant_e4m3_residual_byte(wn10) + f11 = _quant_e4m3_residual_byte(wn11) + else: + f00 = _quant_e4m3_byte(wn00) + f01 = _quant_e4m3_byte(wn01) + if cutlass.const_expr(hi): + f10 = _quant_e4m3_byte(wn10) + f11 = _quant_e4m3_byte(wn11) st_shared_u8( - w_fp8_addr + (gid + Int32(8)) * Int32(w_fp8_stride) + cand_e0, - f10.to(cutlass.Uint8), + w_fp8_addr + gid * Int32(w_fp8_stride) + cand_e0, f00.to(cutlass.Uint8) ) st_shared_u8( - w_fp8_addr + (gid + Int32(8)) * Int32(w_fp8_stride) + cand_e1, - f11.to(cutlass.Uint8), + w_fp8_addr + gid * Int32(w_fp8_stride) + cand_e1, f01.to(cutlass.Uint8) ) + if cutlass.const_expr(hi): + st_shared_u8( + w_fp8_addr + (gid + Int32(8)) * Int32(w_fp8_stride) + cand_e0, + f10.to(cutlass.Uint8), + ) + st_shared_u8( + w_fp8_addr + (gid + Int32(8)) * Int32(w_fp8_stride) + cand_e1, + f11.to(cutlass.Uint8), + ) cute.arch.barrier(**bar_kw) for nt in cutlass.range_constexpr(nt_per_warp_xv): - # dim = vc*V_CHUNK + (nt*N_WARPS + warp_id)*8 (covers the full V_CHUNK). + # dim = vc*V_CHUNK + (nt*N_WARPS + warp_id)*8 (covers the full V_CHUNK); + # with v_ldsm_b8 the warp's tiles are the 16 consecutive dims + # vc*V_CHUNK + warp*16 + nt*8 (see ``v_frag``). dim = Int32(vc) * Int32(v_chunk) + ( Int32(nt) * Int32(n_warps) + warp_id ) * Int32(8) @@ -2078,9 +2282,14 @@ def _vsc(cand: Int32, vc: int): ko = Int32(kstep) * Int32(32) a_addr = w_fp8_addr + a_row * Int32(w_fp8_stride) + ko + a_col a0, a1, a2, a3 = ldmatrix_m8n8x4_b16(a_addr) - b0, b1 = _d2_load_b_fp8( - kv_fp8_base_addr, ko, dim, lane, kv_smem_stride=kv_smem_stride - ) + if cutlass.const_expr(v_ldsm_b8): + # r0/r1: K 0-15 of dims [0,8)/[8,16); r2/r3: K 16-31. + b0 = v_frag[kstep][nt] + b1 = v_frag[kstep][2 + nt] + else: + b0, b1 = _d2_load_b_fp8( + kv_fp8_base_addr, ko, dim, lane, kv_smem_stride=kv_smem_stride + ) xv0, xv1, xv2, xv3 = mma_m16n8k32_f32_e4m3( xv0, xv1, xv2, xv3, a0, a1, a2, a3, b0, b1 ) @@ -2744,6 +2953,7 @@ def s7_epilogue( num_threads: cutlass.Constexpr = 0, barrier_id: cutlass.Constexpr = 0, coalesced_output: cutlass.Constexpr = False, + warp_contiguous_dims: cutlass.Constexpr = False, ): """S7: normalized O + base-2 LSE epilogue. ``epilogue_mode`` (const_expr) selects the destination + normalizer convention; the (gid, d0) output-write @@ -2817,11 +3027,21 @@ def s7_epilogue( for vc in cutlass.range_constexpr(n_v_chunks): for nt in cutlass.range_constexpr(nt_per_warp_xv): at = vc * nt_per_warp_xv + nt - d0 = ( - Int32(vc) * Int32(v_chunk) - + (Int32(nt) * Int32(n_warps) + warp_id) * Int32(8) - + tid * Int32(2) - ) + if cutlass.const_expr(warp_contiguous_dims): + # S6 with v_ldsm_b8: tile nt of this warp holds dims + # vc*V_CHUNK + warp*(8*NT) + nt*8 .. +8. + d0 = ( + Int32(vc) * Int32(v_chunk) + + warp_id * Int32(8 * nt_per_warp_xv) + + Int32(nt) * Int32(8) + + tid * Int32(2) + ) + else: + d0 = ( + Int32(vc) * Int32(v_chunk) + + (Int32(nt) * Int32(n_warps) + warp_id) * Int32(8) + + tid * Int32(2) + ) if cutlass.const_expr(coalesced_output): st_shared_bf16_from_f32( staging_base_addr + (gid * Int32(d_v) + d0) * Int32(2), diff --git a/b12x/attention/_shared/mla/kernel.py b/b12x/attention/_shared/mla/kernel.py index 53716d07c..3b53cb4b9 100644 --- a/b12x/attention/_shared/mla/kernel.py +++ b/b12x/attention/_shared/mla/kernel.py @@ -35,6 +35,7 @@ from b12x._lib.intrinsics import shared_ptr_to_u32 from .decode_math import ( + s0_load_packed_q_to_smem, s0_load_q_bf16_to_smem, s0_quantize_q_to_smem, s1_qk_nope_block_scaled, @@ -77,6 +78,8 @@ _GLM_HEAD_DIM = 576 # GLM per-token packed cache record (reference.pack_mla_kv_cache_reference). _GLM_KV_GMEM_STRIDE = 656 +_GLM_NOPE_SCALE_BYTES_SMEM = 528 # rope offset inside a packed 656-byte staged row +_PACKED_QUERY_RECORD_BYTES = 656 # packed query: 512 e4m3 + 16 scale + 128 rope bytes # DSV4 H8 packs the contiguous 576-byte data record into a 592-byte smem row. # The 16-byte pad preserves KV_SMEM_STRIDE/4 % 32 == 20, matching the generic # 464-byte row's bank rotation while allowing one bulk copy per candidate. @@ -282,6 +285,71 @@ def _wave_balanced_num_splits( # land on chunk boundaries (a candidate is processed by exactly one split -> # multi-split is numerically identical to single-split). # --------------------------------------------------------------------------- +_MLA_SM120_BALANCED_WAVES_ENV = "B12X_MLA_SM120_BALANCED_WAVES" +_MLA_SM120_GLM_FASTPATH_ENV = "B12X_MLA_SM120_GLM_FASTPATH" +_MLA_SM120_GLM_W_HW_DEQUANT_ENV = "B12X_MLA_SM120_GLM_W_HW_DEQUANT" + + +def _env_glm_w_hw_dequant_enabled() -> bool: + """GLM fast path: reconstruct the W HIGH byte with cvt.rn.f16x2.e4m3x2. + + Exact for every E4M3 value; the software expansion it replaces mis-decodes + subnormals and -0, so LOW residual bytes (and results) change slightly. + Off by default; requires the fast path. + """ + raw = os.environ.get(_MLA_SM120_GLM_W_HW_DEQUANT_ENV) + return raw is not None and raw.strip().lower() in {"1", "true", "on", "yes"} + + +def _env_glm_fastpath_enabled() -> bool: + """GLM generic (HPB=16) per-token decode fast path (sm_120a). + + Packed 656-byte KV staging (one bulk copy per token), PV B-fragments via + ``ldmatrix.m16n16.x2.trans.b8`` and W hi/lo slots without the residual + serialization barrier. Bit-identical to the base path; off by default so + existing compile keys and PTX are unchanged. + """ + raw = os.environ.get(_MLA_SM120_GLM_FASTPATH_ENV) + return raw is not None and raw.strip().lower() in {"1", "true", "on", "yes"} + + +def _env_balanced_waves() -> float: + """CTA waves the balanced split policy fills (default 1.0).""" + raw = os.environ.get(_MLA_SM120_BALANCED_WAVES_ENV) + if raw is None: + return 1.0 + try: + waves = float(raw) + except ValueError as exc: + raise ValueError( + f"{_MLA_SM120_BALANCED_WAVES_ENV} must be a positive number, got {raw!r}" + ) from exc + if not waves > 0: + raise ValueError( + f"{_MLA_SM120_BALANCED_WAVES_ENV} must be a positive number, got {raw!r}" + ) + return waves + + +def balanced_split_target_for( + *, + num_splits: int, + rows: int, + h_blocks: int, + sm_count: int, +) -> int: + """Return the balanced policy's bound on active splits per row. + + ``floor(waves * sm_count / (rows * h_blocks))`` CTAs per row fill ``waves`` + CTA waves (one CTA per SM); the bound is clamped to ``[1, num_splits]``. + The kernel additionally never uses fewer splits than the static ranges + would for the same row. + """ + ctas_per_row = max(1, int(rows) * int(h_blocks)) + target = int(_env_balanced_waves() * int(sm_count)) // ctas_per_row + return max(1, min(int(num_splits), target)) + + def plan_unified_decode_splits( *, topk: int, @@ -393,11 +461,46 @@ def __init__( native_glm_h8=False, native_dsv4_h8=False, native_dsv4_h16=False, + balanced_splits=False, + glm_fastpath=False, + glm_w_hw_dequant=False, + q_packed=False, ): self.traits = traits self.layout = layout self.page_block_size = int(page_block_size) self.chunks_per_split = int(chunks_per_split) + # GLM generic (HPB=16, 8 math warps) per-token fast path: the IO warp + # stages each 656-byte record with one bulk copy (rope inline at +528, + # row stride 656 over the contiguous kv_fp8+kv_rope allocation), the PV + # stage loads V B-fragments with ldmatrix.m16n16.x2.trans.b8 and keeps + # W hi/lo in fixed slots (see decode_math.s6_xv_nope), and the epilogue + # maps each warp's 16 consecutive V dims. Bytes and MMA order are + # unchanged, so results are bit-identical to the base path. + self.glm_fastpath = bool( + glm_fastpath + and per_token_len + and int(traits.scale_format) == int(ScaleFormat.ARBITRARY_FP32) + and int(traits.nt_per_warp_xv) == 2 + and not (native_glm_h8 or native_dsv4_h8 or native_dsv4_h16) + ) + # Hardware E4M3 -> f16 reconstruction of the W HIGH byte for the LOW + # residual (fast path only; not bit-identical to the software expansion). + self.glm_w_hw_dequant = bool(glm_w_hw_dequant and self.glm_fastpath) + # Packed query: q_all is a uint8 (rows, heads, 656) record per head + # (E4M3 nope, four fp32 pow2 tile scales, bf16 rope) and S0 copies it + # into the Q stages instead of quantizing a bf16 query (GLM generic + # per-token entry; bit-identical to the bf16-query path). + self.q_packed = bool(q_packed and self.glm_fastpath) + # Balanced split policy (per-token single-cache entry only). False keeps + # the static chunk ranges ``[split * chunks_per_split, +chunks_per_split)`` + # and the existing entry points byte-identical. True selects + # ``kernel_pertok_balanced``, whose runtime ``split_target`` T makes every + # CTA derive its range from the row's live chunk count n at replay time: + # chunks per split = min(ceil(n / T), chunks_per_split), so at most + # max(T, planned splits) splits are active with contiguous ranges of + # near-equal length. The launch grid stays capacity-based. + self.balanced_splits = bool(balanced_splits) self.h_blocks = int(h_blocks) self.num_splits = int(num_splits) self.num_heads = int(num_heads) @@ -407,6 +510,7 @@ def __init__( self.q_stride_row = int(q_stride[0]) self.q_stride_head = int(q_stride[1]) self.q_stride_dim = int(q_stride[2]) + self.q_row_bytes = _PACKED_QUERY_RECORD_BYTES self.swa_indices_stride_row = int(swa_indices_stride0) self.extra_indices_stride_row = int(extra_indices_stride0) self.mid_out_stride_row = int(mid_out_stride[0]) @@ -591,6 +695,44 @@ def call_pertok( stream=stream, ) + @cute.jit + def call_pertok_balanced( + self, + q_all: cute.Tensor, # (rows, heads, D_QK) bf16 + kv_cache_u8: cute.Tensor, # flat (pages*page_nbytes,) u8 (MAIN cache) + swa_indices: cute.Tensor, # (rows, topk) int32 (MAIN indices) + mid_out: cute.Tensor, # (rows, heads, splits, D_V) bf16/f32 partials + mid_lse: cute.Tensor, # (rows, heads, splits) f32 base-2 LSE + sm_scale_log2: Float32, + latent_scale: Float32, + topk_length: cute.Tensor, # (rows,) int32 per-token MAIN valid length + stride_kv_block: Int64, # MAIN per-block byte stride + split_target: Int32, # balanced policy: active splits per row bound + num_tokens: Int32, + stream: cuda.CUstream, + ): + # SINGLE-CACHE PER-TOKEN entry with runtime-balanced chunk ranges. The + # grid stays capacity-based; ``split_target`` is a plain kernel argument, + # so one compiled kernel serves every row count and CUDA-graph replay + # keeps the value captured with the launch. + self.kernel_pertok_balanced( + q_all, + kv_cache_u8, + swa_indices, + mid_out, + mid_lse, + sm_scale_log2, + latent_scale, + topk_length, + stride_kv_block, + split_target, + ).launch( + grid=(num_tokens, self.h_blocks, self.num_splits), + block=[self.block_threads, 1, 1], + min_blocks_per_mp=1, + stream=stream, + ) + @cute.jit def call_extra_pertok( self, @@ -671,9 +813,6 @@ def kernel( # row, and retire a wholly empty CTA before it allocates/initializes the KV # pipeline. The merge treats LSE=-inf as a neutral partial and does not # read the corresponding (potentially stale) mid_out row. - cps = Int32(self.chunks_per_split) - split_first_chunk = split_idx * cps - split_last_chunk = split_first_chunk + cps main_valid_chunks = (section_len + Int32(_CAND_WINDOW - 1)) // Int32( _CAND_WINDOW ) @@ -682,6 +821,9 @@ def kernel( max_main_chunks = Int32((self.topk + _CAND_WINDOW - 1) // _CAND_WINDOW) if main_valid_chunks > max_main_chunks: main_valid_chunks = max_main_chunks + cps = Int32(self.chunks_per_split) + split_first_chunk = split_idx * cps + split_last_chunk = split_first_chunk + cps main_chunk_end = split_last_chunk if main_chunk_end > main_valid_chunks: main_chunk_end = main_valid_chunks @@ -980,7 +1122,7 @@ def kernel( lane, ) p = [Float32(0.0), Float32(0.0), Float32(0.0), Float32(0.0)] - p, wr0, wr1 = s4_online_softmax_glm_h8_swap_ab( + p, wr0, wr1, acc_nope, acc_rope, global_max, global_sum = s4_online_softmax_glm_h8_swap_ab( qk, p, acc_nope, @@ -998,6 +1140,7 @@ def kernel( num_threads=self.math_threads, barrier_id=3, rope_tiles_per_warp=(2 if self.native_dsv4_h8 else 0), + return_state=True, ) w_pre = [ p[0] * wr0, @@ -1091,7 +1234,7 @@ def kernel( lane, ) p = [Float32(0.0), Float32(0.0), Float32(0.0), Float32(0.0)] - p, wr0, wr1 = s4_online_softmax( + p, wr0, wr1, acc_nope, acc_rope, global_max, global_sum = s4_online_softmax( qk, p, acc_nope, @@ -1111,6 +1254,7 @@ def kernel( num_threads=self.math_threads, barrier_id=3, n_acc_tiles=n_acc_tiles, + return_state=True, ) w_pre = [ p[0] * wr0, @@ -1159,6 +1303,7 @@ def kernel( sm_p_full_addr=sm_p_full_addr, sm_p_stride=L.sm_p_full_stride, latent_scale_per_token=t.latent_scale_per_token, + v_ldsm_b8=False, ) # S6b (XV-RoPE) is DSV4-only (V_HAS_ROPE). const_expr-elided for GLM. @@ -1275,6 +1420,7 @@ def kernel( nt_per_warp_xv=t.nt_per_warp_xv, v_has_rope=t.v_has_rope, rope_tiles_per_warp=(2 if self.native_dsv4_h8 else 1), + warp_contiguous_dims=False, ) @cute.kernel @@ -1314,6 +1460,7 @@ def kernel_extra( stride_extra_kv_block, swa_indices, extra_indices, # length tensors unused (per_token_len=False) + Int32(0), has_extra=True, per_token_len=False, ) @@ -1354,8 +1501,49 @@ def kernel_pertok( stride_kv_block, topk_length, swa_indices, + Int32(0), + has_extra=False, + per_token_len=True, + ) + + @cute.kernel + def kernel_pertok_balanced( + self, + q_all: cute.Tensor, + kv_cache_u8: cute.Tensor, + swa_indices: cute.Tensor, + mid_out: cute.Tensor, + mid_lse: cute.Tensor, + sm_scale_log2: Float32, + latent_scale: Float32, + topk_length: cute.Tensor, + stride_kv_block: Int64, + split_target: Int32, + ): + # SINGLE-CACHE PER-TOKEN entry with the balanced split policy: the + # runtime ``split_target`` bounds the active splits per row (see + # __init__). A distinct mangled name keeps ``kernel_pertok`` unchanged. + self._kernel_body( + q_all, + kv_cache_u8, + swa_indices, + mid_out, + mid_lse, + sm_scale_log2, + latent_scale, + Int32(0), + stride_kv_block, + kv_cache_u8, + swa_indices, + Int32(0), + Int32(0), + stride_kv_block, + topk_length, + swa_indices, + split_target, has_extra=False, per_token_len=True, + balanced=True, ) @cute.kernel @@ -1397,6 +1585,7 @@ def kernel_extra_pertok( stride_extra_kv_block, topk_length, extra_topk_length, + Int32(0), has_extra=True, per_token_len=True, ) @@ -1420,9 +1609,11 @@ def _kernel_body( stride_extra_kv_block: Int64, topk_length: cute.Tensor, extra_topk_length: cute.Tensor, + split_target: Int32, *, has_extra: cutlass.Constexpr, per_token_len: cutlass.Constexpr, + balanced: cutlass.Constexpr = False, ): t = self.traits L = self.layout @@ -1462,10 +1653,6 @@ def _kernel_body( # and extra chunk prefixes. This also handles a short-main gap before the # fixed extra-section boundary. Producer and consumer use the same compact # order, so their mbarrier phases remain matched. - cps = Int32(self.chunks_per_split) - split_first_chunk = split_idx * cps - split_last_chunk = split_first_chunk + cps - main_valid_chunks = (section_len + Int32(_CAND_WINDOW - 1)) // Int32( _CAND_WINDOW ) @@ -1474,6 +1661,19 @@ def _kernel_body( max_main_chunks = Int32((self.topk + _CAND_WINDOW - 1) // _CAND_WINDOW) if main_valid_chunks > max_main_chunks: main_valid_chunks = max_main_chunks + cps = Int32(self.chunks_per_split) + if cutlass.const_expr(balanced): + # Balanced policy over the main section only (the launcher rejects + # it for dual-cache launches): chunks per split = + # min(ceil(live main chunks / split_target), chunks_per_split), so a + # row never uses fewer splits than the static ranges would. + balanced_cps = (main_valid_chunks + split_target - Int32(1)) // split_target + if balanced_cps < Int32(1): + balanced_cps = Int32(1) + if balanced_cps < cps: + cps = balanced_cps + split_first_chunk = split_idx * cps + split_last_chunk = split_first_chunk + cps main_chunk_end = split_last_chunk if main_chunk_end > main_valid_chunks: main_chunk_end = main_valid_chunks @@ -1542,7 +1742,7 @@ def _kernel_body( # Match the single-cache body's allocation-preserving packed H8 layout. staged_kv_stride = t.kv_smem_stride - if cutlass.const_expr(self.native_glm_h8): + if cutlass.const_expr(self.native_glm_h8 or self.glm_fastpath): staged_kv_stride = _GLM_KV_GMEM_STRIDE if cutlass.const_expr(self.native_dsv4_h8 or self.native_dsv4_h16): staged_kv_stride = _DSV4_PACKED_SMEM_STRIDE @@ -1553,6 +1753,12 @@ def _kernel_body( kv_rope_addr = kv_fp8_addr + Int32(_DSV4_PACKED_ROPE_OFFSET) kv_rope_buf = kv_fp8_buf kv_sc_addr = kv_fp8_addr + Int32(2) * kv_fp8_buf + if cutlass.const_expr(self.glm_fastpath): + # Packed GLM record staging: rope follows the 528-byte nope+scales + # inside each 656-byte row of the kv_fp8 stage (the two 64x656 + # stages exactly cover the kv_fp8 + kv_rope allocation). + kv_rope_addr = kv_fp8_addr + Int32(_GLM_NOPE_SCALE_BYTES_SMEM) + kv_rope_buf = kv_fp8_buf tok_buf_elems = Int32(L.token_idx_buf_bytes // 4) # mbarrier array: full[0], full[1], empty[0], empty[1] (u64 each). @@ -1598,11 +1804,12 @@ def _kernel_body( ) else: extra_row = topk_row - # q for THIS token row: a 2-D (heads, D_QK) view (s0 indexes [head_base+h, d]). + # q for THIS token row: a 2-D (heads, D_QK) view (s0 indexes [head_base+h, d]); + # with a packed query the row is (heads, 656) bytes. q_token = cute.make_tensor( q_all.iterator + token_idx.to(Int64) * Int64(self.q_stride_row), cute.make_layout( - (self.num_heads, self.q_head_dim), + (self.num_heads, self.q_row_bytes if self.q_packed else self.q_head_dim), stride=(self.q_stride_head, self.q_stride_dim), ), ) @@ -1655,7 +1862,7 @@ def _kernel_body( io_threads=self.io_threads, split_mbar_arrival=self.native_dsv4_h16, fp8_rope=t.fp8_rope, - packed_glm=self.native_glm_h8, + packed_glm=self.native_glm_h8 or self.glm_fastpath, packed_dsv4=self.native_dsv4_h8 or self.native_dsv4_h16, overlap_footer_gather=self.native_dsv4_h16, per_token_latent_scale=t.latent_scale_per_token, @@ -1781,7 +1988,25 @@ def _kernel_body( cute.make_layout(int(L.w_head_sc_bytes // 4)), ) - if cutlass.const_expr(t.scale_format == ScaleFormat.NVFP4_E4M3): + if cutlass.const_expr(self.q_packed): + s0_load_packed_q_to_smem( + q_token, + q_fp8_stage, + q_sc_stage_view, + q_rope_stage, + head_base_stage, + Int32(self.valid_hpb), + tid_sel, + d_nope=t.d_nope, + d_rope=t.d_rope, + num_scales=t.num_scales, + hpb=t.hpb, + q_nope_stride=t.q_nope_stride, + q_rope_stride=L.q_rope_stride, + num_threads=nt_stage, + barrier_id=2, + ) + elif cutlass.const_expr(t.scale_format == ScaleFormat.NVFP4_E4M3): s0_load_q_bf16_to_smem( q_token, q_fp8_stage, @@ -1947,7 +2172,7 @@ def _kernel_body( lane, ) p = [Float32(0.0), Float32(0.0), Float32(0.0), Float32(0.0)] - p, wr0, wr1 = s4_online_softmax_glm_h8_swap_ab( + p, wr0, wr1, acc_nope, acc_rope, global_max, global_sum = s4_online_softmax_glm_h8_swap_ab( qk, p, acc_nope, @@ -1968,6 +2193,7 @@ def _kernel_body( 2 if (self.native_dsv4_h8 or self.native_dsv4_h16) else 0 ), barrier_threads=bt_stage, + return_state=True, ) w_pre = [ p[0] * wr0, @@ -2036,7 +2262,7 @@ def _kernel_body( num_scales=t.num_scales, quant_tile=t.quant_tile, q_nope_stride=t.q_nope_stride, - kv_smem_stride=t.kv_smem_stride, + kv_smem_stride=staged_kv_stride, scale_bytes_per_token=8, scale_format=t.scale_format, latent_scale_per_token=t.latent_scale_per_token, @@ -2050,6 +2276,9 @@ def _kernel_body( d_rope=t.d_rope, q_rope_stride=L.q_rope_stride, fp8_rope=t.fp8_rope, + kv_rope_stride_bytes=( + _GLM_KV_GMEM_STRIDE if self.glm_fastpath else t.d_rope * 2 + ), ) # Per-chunk section dispatch for the S3 mask: compare the @@ -2089,7 +2318,7 @@ def _kernel_body( ) p = [Float32(0.0), Float32(0.0), Float32(0.0), Float32(0.0)] - p, wr0, wr1 = s4_online_softmax( + p, wr0, wr1, acc_nope, acc_rope, global_max, global_sum = s4_online_softmax( qk, p, acc_nope, @@ -2109,6 +2338,8 @@ def _kernel_body( num_threads=self.math_threads, barrier_id=3, n_acc_tiles=n_acc_tiles, + skip_unit_rescale=self.glm_fastpath, + return_state=True, ) w_pre = [ p[0] * wr0, @@ -2146,7 +2377,7 @@ def _kernel_body( v_chunk=t.quant_tile, hpb=t.hpb, bi=t.bi, - kv_smem_stride=t.kv_smem_stride, + kv_smem_stride=staged_kv_stride, w_fp8_stride=t.bi + 16, n_warps=8, scale_bytes_per_token=8, @@ -2157,6 +2388,8 @@ def _kernel_body( sm_p_full_addr=sm_p_full_addr, sm_p_stride=L.sm_p_full_stride, latent_scale_per_token=t.latent_scale_per_token, + v_ldsm_b8=self.glm_fastpath, + w_hw_dequant=self.glm_w_hw_dequant, ) # S6b (XV-RoPE) is DSV4-only (V_HAS_ROPE). const_expr-elided for GLM. @@ -2275,6 +2508,7 @@ def _kernel_body( rope_tiles_per_warp=( 2 if (self.native_dsv4_h8 or self.native_dsv4_h16) else 1 ), + warp_contiguous_dims=self.glm_fastpath, ) @@ -2329,9 +2563,7 @@ def _cache_block_stride_bytes( expected = int(page_size) * rec else: expected = int(page_size) * COMPRESSED_MLA_BYTES_PER_TOKEN - if int(model_type) == int(ModelType.GLM_NSA) and cache.is_contiguous(): - return expected - # Use the tensor's physical page stride for padded views. + # Use the tensor's physical page stride for exact-payload and padded views. if cache.ndim >= 2: stride = int(cache.stride(0)) * int(cache.element_size()) if stride < expected: @@ -2380,9 +2612,21 @@ def _sparse_mla_decode_grid_flat_launch( has_extra: bool, per_token_len: bool, latent_scale_per_token: bool = False, + balanced_split_target: int = 0, ) -> None: - q_head_dim = int(q_all.shape[-1]) + q_packed = bool( + q_all.dtype == torch.uint8 and int(q_all.shape[-1]) == _PACKED_QUERY_RECORD_BYTES + ) + q_head_dim = _GLM_HEAD_DIM if q_packed else int(q_all.shape[-1]) rows = int(q_all.shape[0]) + balanced_splits = int(balanced_split_target) > 0 + if balanced_splits and (has_extra or not per_token_len): + raise ValueError( + "SM120 sparse MLA decode balanced splits require the single-cache " + "per-token entry" + ) + glm_fastpath = _env_glm_fastpath_enabled() + glm_w_hw_dequant = _env_glm_w_hw_dequant_enabled() heads = int(q_all.shape[1]) native_glm_h8 = bool( int(model_type) == int(ModelType.GLM_NSA) @@ -2436,13 +2680,24 @@ def _sparse_mla_decode_grid_flat_launch( hpb = int(traits.hpb) d_v = int(traits.d_v) stream = cuda.CUstream(torch.cuda.current_stream().cuda_stream) + q_all_cute_dtype = cutlass.Uint8 if q_packed else cutlass.BFloat16 + # Split partials are bf16 by default; an fp32 partial workspace keeps every + # split result exact until the merge rounds once at the output. + if mid_out.dtype == torch.bfloat16: + mid_out_cute_dtype = cutlass.BFloat16 + elif mid_out.dtype == torch.float32: + mid_out_cute_dtype = cutlass.Float32 + else: + raise TypeError( + f"SM120 sparse MLA decode mid_out must be bf16 or fp32, got {mid_out.dtype}" + ) if per_token_len: pertok_base = ( - _to_cute(q_all, cutlass.BFloat16, dynamic_layout=True), + _to_cute(q_all, q_all_cute_dtype, dynamic_layout=True), _to_cute(kv_flat, cutlass.Uint8, align=16), _to_cute(swa_indices, cutlass.Int32, align=4, dynamic_layout=True), - _to_cute(mid_out, cutlass.BFloat16, align=16, dynamic_layout=True), + _to_cute(mid_out, mid_out_cute_dtype, align=16, dynamic_layout=True), _to_cute(mid_lse, cutlass.Float32, align=4, dynamic_layout=True), Float32(float(sm_scale) * LOG2_E), Float32(float(latent_scale)), @@ -2459,14 +2714,16 @@ def _sparse_mla_decode_grid_flat_launch( Int32(rows), stream, ) + elif balanced_split_target > 0: + args = pertok_base + (Int32(balanced_split_target), Int32(rows), stream) else: args = pertok_base + (Int32(rows), stream) else: base_args = ( - _to_cute(q_all, cutlass.BFloat16, dynamic_layout=True), + _to_cute(q_all, q_all_cute_dtype, dynamic_layout=True), _to_cute(kv_flat, cutlass.Uint8, align=16), _to_cute(swa_indices, cutlass.Int32, align=4, dynamic_layout=True), - _to_cute(mid_out, cutlass.BFloat16, align=16, dynamic_layout=True), + _to_cute(mid_out, mid_out_cute_dtype, align=16, dynamic_layout=True), _to_cute(mid_lse, cutlass.Float32, align=4, dynamic_layout=True), Float32(float(sm_scale) * LOG2_E), Float32(float(latent_scale)), @@ -2510,9 +2767,22 @@ def _sparse_mla_decode_grid_flat_launch( native_glm_h8=native_glm_h8, native_dsv4_h8=native_dsv4_h8, native_dsv4_h16=native_dsv4_h16, + balanced_splits=balanced_splits, + glm_fastpath=glm_fastpath, + glm_w_hw_dequant=glm_w_hw_dequant, + q_packed=q_packed, ) + if q_packed and not kernel.q_packed: + raise ValueError( + "SM120 sparse MLA decode packed query requires the GLM generic " + "per-token fast path (B12X_MLA_SM120_GLM_FASTPATH=1)" + ) spec_fields = [ key_field("model_type", traits.model_type), + key_field("balanced_splits", int(balanced_splits)), + key_field("glm_fastpath", int(kernel.glm_fastpath)), + key_field("glm_w_hw_dequant", int(kernel.glm_w_hw_dequant)), + key_field("q_packed", int(kernel.q_packed)), key_field("compute_mode", traits.compute_mode), key_field("scale_format", traits.scale_format), key_field("fp8_rope", int(traits.fp8_rope)), @@ -2539,7 +2809,7 @@ def _sparse_mla_decode_grid_flat_launch( dims=( DimKey.dynamic(), DimKey.exact(heads), - DimKey.exact(q_head_dim), + DimKey.exact(int(q_all.shape[-1])), ), ), tensor_key( @@ -2603,7 +2873,12 @@ def _sparse_mla_decode_grid_flat_launch( *spec_fields, ) if per_token_len: - entry = kernel.call_extra_pertok if has_extra else kernel.call_pertok + if has_extra: + entry = kernel.call_extra_pertok + elif balanced_splits: + entry = kernel.call_pertok_balanced + else: + entry = kernel.call_pertok else: entry = kernel.call_extra if has_extra else kernel b12x_launch( @@ -2649,6 +2924,7 @@ def _sparse_mla_decode_grid_op( has_extra: bool, per_token_len: bool, latent_scale_per_token: bool = False, + balanced_split_target: int = 0, ) -> None: _sparse_mla_decode_grid_flat_launch( q_all, @@ -2681,6 +2957,7 @@ def _sparse_mla_decode_grid_op( has_extra, per_token_len, latent_scale_per_token, + balanced_split_target, ) @@ -2716,6 +2993,7 @@ def _sparse_mla_decode_grid_fake( has_extra: bool, per_token_len: bool, latent_scale_per_token: bool = False, + balanced_split_target: int = 0, ) -> None: return None @@ -2743,9 +3021,26 @@ def run_unified_decode( scale_format_override: int | None = None, fp8_rope_override: bool | None = None, latent_scale_per_token: bool = False, + split_policy: str = "static", ): """Active SM120 sparse-MLA decode: kernel (split-K partials) + merge. + ``split_policy`` selects how the launched splits partition a row's live + 64-candidate chunks. ``"static"`` assigns split ``s`` the fixed range + ``[s * chunks_per_split, (s + 1) * chunks_per_split)`` of the planned + capacity, so a short row keeps only its leading splits busy while each of + them scans up to ``chunks_per_split`` chunks serially. ``"balanced"`` makes + each CTA derive the range from the row's live chunk count n: chunks per + split = ``min(ceil(n / T), chunks_per_split)`` with + ``T = min(num_splits, floor(waves * sm_count / (rows * head_blocks)))`` + (``B12X_MLA_SM120_BALANCED_WAVES``, default one CTA wave), so a short row + spreads over up to T splits with near-equal ranges while a long row keeps + the static ranges. T is a runtime kernel argument of a dedicated per-token + entry; the grid and the workspace stay capacity-based, so both policies + are CUDA-graph safe. ``"balanced"`` changes which chunks each partial + covers and therefore the merge rounding, not the attention math. It + requires per-token lengths and a single-cache launch. + Routes DSV4 (q_head_dim==512, UE8M0 footer) AND GLM_NSA (q_head_dim==576, ARBITRARY_FP32 inline scales) to the SAME warp-specialized kernel via the cute.constexpr traits branches (model_type/scale_format/v_has_rope). The @@ -2797,7 +3092,10 @@ def run_unified_decode( "(q_head_dim==512); GLM/DSV3.2 has no extra cache" ) - q_head_dim = int(q_all.shape[-1]) + q_packed = bool( + q_all.dtype == torch.uint8 and int(q_all.shape[-1]) == _PACKED_QUERY_RECORD_BYTES + ) + q_head_dim = _GLM_HEAD_DIM if q_packed else int(q_all.shape[-1]) if q_head_dim not in (_DSV4_HEAD_DIM, _GLM_HEAD_DIM): raise NotImplementedError( f"SM120 sparse MLA decode supports q_head_dim 512 (DSV4) or 576 (GLM); " @@ -3020,6 +3318,34 @@ def _length_tensor(lengths, name, cap): extra_topk=extra_topk, preferred_num_splits=preferred_num_splits, ) + if split_policy not in ("static", "balanced"): + raise ValueError( + f"SM120 sparse MLA decode split_policy must be 'static' or " + f"'balanced', got {split_policy!r}" + ) + balanced_split_target = 0 + if split_policy == "balanced": + if has_extra: + raise ValueError( + "SM120 sparse MLA decode split_policy='balanced' supports " + "single-cache launches only" + ) + if not per_token_len: + raise ValueError( + "SM120 sparse MLA decode split_policy='balanced' requires " + "per-token lengths (swa_topk_lengths on a CUDA device)" + ) + if sm_count is None: + raise ValueError( + "SM120 sparse MLA decode split_policy='balanced' requires a CUDA " + "device (SM count unavailable)" + ) + balanced_split_target = balanced_split_target_for( + num_splits=int(num_splits), + rows=rows, + h_blocks=int(h_blocks), + sm_count=int(sm_count), + ) # Side-channel record of the chosen split plan (benchmarks / AutoTuner read # LAST_DECODE_PLAN["num_splits"]). Informational only. native_glm_h8 = bool( @@ -3071,6 +3397,28 @@ def _length_tensor(lengths, name, cap): h_blocks=int(h_blocks), sm_count=(int(sm_count) if sm_count else None), per_token_len=bool(per_token_len), + split_policy=str(split_policy), + balanced_split_target=int(balanced_split_target), + glm_fastpath=bool( + _env_glm_fastpath_enabled() + and per_token_len + and int(traits.scale_format) == int(ScaleFormat.ARBITRARY_FP32) + and int(traits.nt_per_warp_xv) == 2 + and not (native_glm_h8 or native_dsv4_h8 or native_dsv4_h16) + ), + glm_w_hw_dequant=bool( + _env_glm_w_hw_dequant_enabled() + and _env_glm_fastpath_enabled() + and per_token_len + and int(traits.scale_format) == int(ScaleFormat.ARBITRARY_FP32) + and int(traits.nt_per_warp_xv) == 2 + and not (native_glm_h8 or native_dsv4_h8 or native_dsv4_h16) + ), + partial_dtype=( + str(workspace.tmp_output.dtype) + if workspace.tmp_output is not None + else None + ), ) # Workspace mid_out/mid_lse must hold num_splits partials per (token, head). if num_splits > max_chunks: @@ -3173,6 +3521,7 @@ def _launch_grid(grid_h_blocks: int, valid_hpb: int, head_block_offset: int): bool(has_extra), bool(per_token_len), bool(latent_scale_per_token), + int(balanced_split_target), ) if h_blocks_full > 0: diff --git a/b12x/attention/_shared/mla/merge.py b/b12x/attention/_shared/mla/merge.py index 38891e3a8..d2d4eaf62 100644 --- a/b12x/attention/_shared/mla/merge.py +++ b/b12x/attention/_shared/mla/merge.py @@ -715,9 +715,12 @@ def run_sparse_mla_split_decode_merge( raise ValueError("sparse MLA merge tensors must be on the same device") if tmp_lse.dtype != torch.float32: raise TypeError(f"tmp_lse must have dtype torch.float32, got {tmp_lse.dtype}") - if tmp_output.dtype != output.dtype: + if tmp_output.dtype != output.dtype and tmp_output.dtype != torch.float32: + # Partials are either the output element type or fp32 (exact split + # partials merged in fp32 and rounded once at the output). raise TypeError( - f"tmp_output dtype {tmp_output.dtype} must match output dtype {output.dtype}" + f"tmp_output dtype {tmp_output.dtype} must match output dtype " + f"{output.dtype} or be torch.float32" ) if tmp_output.ndim != 4: raise ValueError( diff --git a/b12x/attention/_shared/mla/prefill.py b/b12x/attention/_shared/mla/prefill.py index aeefc42ef..3c9738a65 100644 --- a/b12x/attention/_shared/mla/prefill.py +++ b/b12x/attention/_shared/mla/prefill.py @@ -58,9 +58,8 @@ def _cache_block_stride_bytes( expected = int(page_size) * rec else: expected = int(page_size) * COMPRESSED_MLA_BYTES_PER_TOKEN - if model_type == ModelType.GLM_NSA and cache.is_contiguous(): - return expected - # The runtime page stride is part of the packed/padded cache contract. + # The runtime page stride is part of the cache contract. It can be either + # the exact payload width or a larger packed/padded stride. if cache.ndim >= 2: stride = int(cache.stride(0)) * int(cache.element_size()) if stride < expected: diff --git a/b12x/attention/_shared/mla/prefill_mg.py b/b12x/attention/_shared/mla/prefill_mg.py index 5e67b8661..411da2426 100644 --- a/b12x/attention/_shared/mla/prefill_mg.py +++ b/b12x/attention/_shared/mla/prefill_mg.py @@ -208,9 +208,7 @@ def _cache_block_stride_bytes( expected = int(page_size) * rec else: expected = int(page_size) * COMPRESSED_MLA_BYTES_PER_TOKEN - if is_glm and cache.is_contiguous(): - return expected - # Use the tensor's physical page stride for padded views. + # Use the tensor's physical page stride for exact-payload and padded views. if cache.ndim >= 2: stride = int(cache.stride(0)) * int(cache.element_size()) if stride < expected: diff --git a/b12x/attention/_shared/static_fp8_quant.py b/b12x/attention/_shared/static_fp8_quant.py deleted file mode 100644 index d1714437d..000000000 --- a/b12x/attention/_shared/static_fp8_quant.py +++ /dev/null @@ -1,188 +0,0 @@ -"""Graph-safe BF16 to per-tensor E4M3 query quantization.""" - -from __future__ import annotations - -from dataclasses import dataclass -from threading import RLock - -import cuda.bindings.driver as cuda -import cutlass -import cutlass.cute as cute -import torch -from cutlass import Float32, Int32, Int64 -from cutlass.cute.runtime import from_dlpack - -from b12x._lib.compiler import KernelCompileSpec, compile as compile_cute -from b12x._lib.compiler import key_field, run_compiled -from b12x._lib.intrinsics import ( - bfloat2_to_float2_scaled, - cvt_f32x4_to_e4m3x4, - get_ptr_as_int64, - ld_global_v2_u32, - st_global_u32, -) - -FP8 = torch.float8_e4m3fn -_THREADS = 128 -_VALUES_PER_THREAD = 4 -_LOCK = RLock() -_CACHE: dict[tuple[int, int], object] = {} - - -def _byte_base_pointer(tensor: torch.Tensor) -> torch.Tensor: - byte_view = tensor.view(torch.uint8) - return torch.as_strided(byte_view, size=(1,), stride=(1,)) - - -def _to_cute(tensor: torch.Tensor, dtype, *, align: int): - converted = from_dlpack(tensor, assumed_align=align) - converted.element_type = dtype - return converted.mark_layout_dynamic(leading_dim=0) - - -class _StaticFp8QuantKernel: - def __init__(self, max_numel: int): - self.max_numel = int(max_numel) - - @cute.jit - def __call__( - self, - source_bytes: cute.Tensor, - output_bytes: cute.Tensor, - scale: cute.Tensor, - active_numel: Int32, - stream: cuda.CUstream, - ): - grid = (self.max_numel + _THREADS * _VALUES_PER_THREAD - 1) // ( - _THREADS * _VALUES_PER_THREAD - ) - self.kernel(source_bytes, output_bytes, scale, active_numel).launch( - grid=(grid, 1, 1), - block=(_THREADS, 1, 1), - stream=stream, - ) - - @cute.kernel - def kernel( - self, - source_bytes: cute.Tensor, - output_bytes: cute.Tensor, - scale: cute.Tensor, - active_numel: Int32, - ): - thread, _, _ = cute.arch.thread_idx() - block, _, _ = cute.arch.block_idx() - word = Int32(block) * Int32(_THREADS) + Int32(thread) - element = word * Int32(_VALUES_PER_THREAD) - if element < active_numel: - inverse_scale = Float32(1.0) / Float32(scale[0]) - lo, hi = ld_global_v2_u32( - get_ptr_as_int64(source_bytes, element.to(Int64) * Int64(2)) - ) - value0, value1 = bfloat2_to_float2_scaled(lo, inverse_scale) - value2, value3 = bfloat2_to_float2_scaled(hi, inverse_scale) - packed = cvt_f32x4_to_e4m3x4(value0, value1, value2, value3) - st_global_u32( - get_ptr_as_int64(output_bytes, element.to(Int64)), - packed, - ) - - -@dataclass(frozen=True) -class Binding: - source: torch.Tensor - output: torch.Tensor - scale: torch.Tensor - max_numel: int - - -def bind( - *, - source: torch.Tensor, - output: torch.Tensor, - scale: torch.Tensor, - max_numel: int, -) -> Binding: - if source.dtype != torch.bfloat16 or not source.is_contiguous(): - raise TypeError("query source must be contiguous BF16") - if output.dtype != FP8 or tuple(output.shape) != tuple(source.shape): - raise TypeError("query quant output must be E4M3 with the source shape") - if not output.is_contiguous() or output.device != source.device: - raise ValueError("query quant output must be contiguous on the source device") - if scale.dtype != torch.float32 or scale.numel() != 1: - raise TypeError("query quant scale must be a scalar float32 tensor") - if scale.device != source.device: - raise ValueError("query quant scale must be on the source device") - if source.numel() % _VALUES_PER_THREAD: - raise ValueError("query element count must be divisible by four") - if not 0 < source.numel() <= int(max_numel): - raise ValueError("query element count exceeds the planned quant capacity") - return Binding( - source=source.detach(), - output=output.detach(), - scale=scale.detach(), - max_numel=int(max_numel), - ) - - -def _launch(binding: Binding): - source_bytes = _byte_base_pointer(binding.source) - output_bytes = _byte_base_pointer(binding.output) - stream = cuda.CUstream(torch.cuda.current_stream().cuda_stream) - entry = _StaticFp8QuantKernel(binding.max_numel) - args = ( - _to_cute(source_bytes, cutlass.Uint8, align=16), - _to_cute(output_bytes, cutlass.Uint8, align=16), - _to_cute(binding.scale.reshape(1), cutlass.Float32, align=4), - Int32(binding.source.numel()), - stream, - ) - spec = KernelCompileSpec.from_fields( - "attention.static_fp8_query_quant", - 1, - key_field("max_numel", binding.max_numel), - ) - return entry, args, spec - - -def _signature(binding: Binding) -> tuple[int, int]: - index = binding.source.device.index - if index is None: - index = torch.cuda.current_device() - return int(index), int(binding.max_numel) - - -def compile(*, binding: Binding) -> None: - signature = _signature(binding) - with _LOCK: - compiled = _CACHE.get(signature) - if compiled is None: - entry, args, spec = _launch(binding) - compiled = compile_cute(entry, *args, compile_spec=spec) - with _LOCK: - _CACHE[signature] = compiled - - -def run(*, binding: Binding) -> torch.Tensor: - signature = _signature(binding) - with _LOCK: - compiled = _CACHE.get(signature) - if compiled is None: - if torch.cuda.is_current_stream_capturing(): - raise RuntimeError( - "query quant compile miss during CUDA graph capture; call compile first" - ) - compile(binding=binding) - with _LOCK: - compiled = _CACHE[signature] - _, args, _ = _launch(binding) - run_compiled(compiled, args) - return binding.output - - -def clear_caches() -> None: - with _LOCK: - _CACHE.clear() - - -__all__ = ["Binding", "bind", "clear_caches", "compile", "run"] diff --git a/b12x/attention/_shared/workspace.py b/b12x/attention/_shared/workspace.py index 92d91b4d9..b276bb4fa 100644 --- a/b12x/attention/_shared/workspace.py +++ b/b12x/attention/_shared/workspace.py @@ -2541,7 +2541,7 @@ def bind_compressed_mla( ) return build_compressed_mla_binding( - scratch=self, + workspace=self, q=q, swa_indices=swa_indices, swa_lengths=swa_lengths, diff --git a/b12x/attention/dense_mla/__init__.py b/b12x/attention/dense_mla/__init__.py index 46361faf5..8b633a067 100644 --- a/b12x/attention/dense_mla/__init__.py +++ b/b12x/attention/dense_mla/__init__.py @@ -1,25 +1,22 @@ -"""Paged dense Multi-head Latent Attention on SM12x. +"""Dense Multi-head Latent Attention for Kimi K3 on SM12x. -The operator consumes absorbed queries and combined compressed-cache records -directly. Logical Q/K width, logical value width, and physical-record width -are planned independently. Supported logical widths are ``(576, 512)`` and -``(1088, 1024)``. +The operator consumes the absorbed K3 query and the combined compressed-cache +record directly: -* query: ``[total_q, local_heads, logical_qk_width]``; -* cache: ``[pages, page_size, physical_record_width]`` (or a singleton-head - rank-4 view), where the logical key prefix is read directly; -* value: the logical value prefix of each record. +* query: ``[total_q, local_heads, 576]``; +* cache: ``[pages, page_size, 576]`` (or a singleton-head rank-4 view); +* key: all 576 record elements; +* value: the first 512 record elements. -``window_size=None`` attends to the full causal context. A positive -``window_size`` restricts query position ``p`` to the inclusive interval -``[max(0, p - window_size + 1), p]``. The attention scale is caller-provided. -BF16 and E4M3 records are supported. For E4M3 cache with BF16 query input, -``q_dtype=torch.bfloat16`` plans fixed query-quantization storage in the same -caller-owned scratch buffer. +The attention scale is based on K3's *raw* 192-wide QK head +(``1 / sqrt(192)``), not the 576-wide absorbed record. BF16 and standard +E4M3 combined records are supported; E4M3 uses one caller-provided FP32 scalar +scale for the whole record. Planned lifecycle: ``plan(Caps(...))`` -> caller allocates ``scratch_specs`` -> -``bind`` (views and validation only) -> ``compile``/``run``. Decode and extend -share the same graph-safe dense core. No selected-index array is materialized. +``bind`` (views and validation only) -> ``compile``/``run``. Decode and +single-request prefill share the same graph-safe dense core. No selected-index +array is materialized. """ from __future__ import annotations @@ -41,6 +38,7 @@ "plan", "bind", "compile", + "dynamic_sparse_chunk_indices", "run", "reference", "infer_mode", @@ -48,7 +46,7 @@ "clear_caches", ), dtypes=("bf16", "fp8_e4m3"), - recipes=("strided_paged_dense_mla",), + recipes=("kimi_k3_dense_mla",), requires=(), provenance=Provenance( repo="https://github.com/lukealonso/b12x", @@ -62,8 +60,9 @@ test_path="tests/attention/test_dense_mla.py", since="0.8.0", notes=( - "Physical page and record addressing is 64-bit. Optional BF16-to-E4M3 " - "query conversion uses planned caller-owned storage." + "K3 TP12 uses eight local heads, a 576-element absorbed Q/K record, " + "a 512-element latent value, and 1/sqrt(192) scaling. Pool-scaled " + "addressing is 64-bit and the runtime owns no allocations." ), ) @@ -77,6 +76,7 @@ bind, clear_caches, compile, + dynamic_sparse_chunk_indices, infer_mode, is_supported, plan, diff --git a/b12x/attention/dense_mla/_forward.py b/b12x/attention/dense_mla/_forward.py index 64316b371..947e0756f 100644 --- a/b12x/attention/dense_mla/_forward.py +++ b/b12x/attention/dense_mla/_forward.py @@ -59,10 +59,13 @@ def __init__( num_splits: int, chunks_per_split: int, query_tile: int, + uses_query_cache_seqlens: bool, + sparse_stride: int, + sparse_min_tokens: int, + sparse_sink_chunks: int, + sparse_recent_chunks: int, + sparse_refresh_interval: int, fp8: bool, - qk_dim: int, - value_dim: int, - window_size: int | None, ): self.layout = layout self.page_size = int(page_size) @@ -71,10 +74,13 @@ def __init__( self.num_splits = int(num_splits) self.chunks_per_split = int(chunks_per_split) self.query_tile = int(query_tile) + self.uses_query_cache_seqlens = bool(uses_query_cache_seqlens) + self.sparse_stride = int(sparse_stride) + self.sparse_min_tokens = int(sparse_min_tokens) + self.sparse_sink_chunks = int(sparse_sink_chunks) + self.sparse_recent_chunks = int(sparse_recent_chunks) + self.sparse_refresh_interval = int(sparse_refresh_interval) self.fp8 = bool(fp8) - self.qk_dim = int(qk_dim) - self.value_dim = int(value_dim) - self.window_size = None if window_size is None else int(window_size) self.kv_stages = int(layout.kv_stages) self.math_warps = MATH_WARPS_PER_QUERY * self.query_tile self.math_threads = self.math_warps * 32 @@ -88,6 +94,7 @@ def __call__( page_table: cute.Tensor, cache_seqlens: cute.Tensor, cu_seqlens_q: cute.Tensor, + query_cache_seqlens: cute.Tensor, output: cute.Tensor, final_lse: cute.Tensor, partial_output: cute.Tensor, @@ -98,7 +105,6 @@ def __call__( q_stride_row_bytes: Int64, q_stride_head_bytes: Int64, page_stride_bytes: Int64, - cache_record_stride_bytes: Int64, page_table_stride: Int64, total_q: Int32, batch: Int32, @@ -112,6 +118,7 @@ def __call__( page_table, cache_seqlens, cu_seqlens_q, + query_cache_seqlens, output, final_lse, partial_output, @@ -122,7 +129,6 @@ def __call__( q_stride_row_bytes, q_stride_head_bytes, page_stride_bytes, - cache_record_stride_bytes, page_table_stride, total_q, batch, @@ -134,6 +140,34 @@ def __call__( stream=stream, ) + @cute.jit + def _selected_chunk( + self, + selected_index: Int32, + valid_chunks: Int32, + sparse_active: Int32, + ) -> Int32: + chunk = selected_index + if sparse_active != Int32(0): + sink = Int32(self.sparse_sink_chunks) + if sink > valid_chunks: + sink = valid_chunks + recent = Int32(self.sparse_recent_chunks) + if recent > valid_chunks - sink: + recent = valid_chunks - sink + middle_end = valid_chunks - recent + middle_span = middle_end - sink + middle_count = (middle_span + Int32(self.sparse_stride - 1)) // Int32( + self.sparse_stride + ) + if selected_index < sink: + chunk = selected_index + elif selected_index < sink + middle_count: + chunk = sink + (selected_index - sink) * Int32(self.sparse_stride) + else: + chunk = middle_end + selected_index - sink - middle_count + return chunk + @cute.jit def _run_math_group( self, @@ -146,7 +180,8 @@ def _run_math_group( mbar_base, active_chunks: Int32, split_first_chunk: Int32, - visible_begin: Int32, + valid_chunks: Int32, + sparse_active: Int32, visible_end: Int32, query_row: Int32, head_base: Int32, @@ -165,10 +200,10 @@ def _run_math_group( *, barrier_id: cutlass.Constexpr, ): - accumulator_fragment = cute.make_rmem_tensor(self.value_dim // 16, Float32) + accumulator_fragment = cute.make_rmem_tensor(32, Float32) global_max_fragment = cute.make_rmem_tensor(2, Float32) global_sum_fragment = cute.make_rmem_tensor(2, Float32) - for idx in cutlass.range_constexpr(self.value_dim // 16): + for idx in cutlass.range_constexpr(32): accumulator_fragment[idx] = Float32(0.0) global_max_fragment[0] = Float32(-1.0e30) global_max_fragment[1] = Float32(-1.0e30) @@ -180,7 +215,12 @@ def _run_math_group( kv_buffer_bytes = Int32(CANDIDATES_PER_CHUNK * self.layout.record_stride_bytes) for local_chunk in cutlass.range(active_chunks, unroll=1): - chunk = split_first_chunk + Int32(local_chunk) + selected_index = split_first_chunk + Int32(local_chunk) + chunk = self._selected_chunk( + selected_index, + valid_chunks, + sparse_active, + ) chunk_begin = chunk * Int32(CANDIDATES_PER_CHUNK) buffer = Int32(0) if cutlass.const_expr(self.kv_stages == 2): @@ -201,7 +241,7 @@ def _run_math_group( accumulator_fragment[tile * 2], accumulator_fragment[tile * 2 + 1], ] - for tile in range(self.value_dim // 32) + for tile in range(16) ] global_max = [ global_max_fragment[0], @@ -225,7 +265,6 @@ def _run_math_group( local_warp, lane, record_stride_bytes=self.layout.record_stride_bytes, - qk_dim=self.qk_dim, ) else: qk = qk_bf16( @@ -235,12 +274,10 @@ def _run_math_group( local_warp, lane, record_stride_bytes=self.layout.record_stride_bytes, - qk_dim=self.qk_dim, ) qk = mask_and_scale( qk, chunk_begin, - visible_begin, visible_end, score_scale_log2, local_warp, @@ -264,7 +301,6 @@ def _run_math_group( lane, local_tid, barrier_id=barrier_id, - value_dim=self.value_dim, ) weights = [ probability[0] * warp_scale0, @@ -284,7 +320,6 @@ def _run_math_group( local_tid, record_stride_bytes=self.layout.record_stride_bytes, barrier_id=barrier_id, - value_dim=self.value_dim, ) else: accumulator = pv_bf16( @@ -296,9 +331,8 @@ def _run_math_group( lane, record_stride_bytes=self.layout.record_stride_bytes, barrier_id=barrier_id, - value_dim=self.value_dim, ) - for tile in cutlass.range_constexpr(self.value_dim // 32): + for tile in cutlass.range_constexpr(16): accumulator_fragment[tile * 2] = accumulator[tile][0] accumulator_fragment[tile * 2 + 1] = accumulator[tile][1] global_max_fragment[0] = global_max[0] @@ -324,7 +358,7 @@ def _run_math_group( accumulator_fragment[tile * 2], accumulator_fragment[tile * 2 + 1], ] - for tile in range(self.value_dim // 32) + for tile in range(16) ] global_max = [ global_max_fragment[0], @@ -353,7 +387,6 @@ def _run_math_group( has_splits=True, fp8=self.fp8, ln2=0.6931471805599453, - value_dim=self.value_dim, ) else: write_partial_or_final( @@ -373,7 +406,6 @@ def _run_math_group( has_splits=False, fp8=self.fp8, ln2=0.6931471805599453, - value_dim=self.value_dim, ) else: write_partial_or_final( @@ -393,7 +425,6 @@ def _run_math_group( has_splits=False, fp8=self.fp8, ln2=0.6931471805599453, - value_dim=self.value_dim, ) @cute.kernel @@ -404,6 +435,7 @@ def kernel( page_table: cute.Tensor, cache_seqlens: cute.Tensor, cu_seqlens_q: cute.Tensor, + query_cache_seqlens: cute.Tensor, output: cute.Tensor, final_lse: cute.Tensor, partial_output: cute.Tensor, @@ -414,7 +446,6 @@ def kernel( q_stride_row_bytes: Int64, q_stride_head_bytes: Int64, page_stride_bytes: Int64, - cache_record_stride_bytes: Int64, page_table_stride: Int64, total_q: Int32, batch: Int32, @@ -441,33 +472,41 @@ def kernel( else: upper = middle request = lower + else: + request = query_tile_index query_begin = Int32(cu_seqlens_q[request]) query_end = Int32(cu_seqlens_q[request + Int32(1)]) query_length = query_end - query_begin cache_length = Int32(cache_seqlens[request]) - first_valid_chunk = Int32(0) - visible_chunks_end = cache_length - if cutlass.const_expr(self.window_size is not None): - tile_local_query = query_start - query_begin - visible_chunks_end = ( - cache_length - query_length + tile_local_query + Int32(1) - ) - if visible_chunks_end < Int32(0): - visible_chunks_end = Int32(0) - if visible_chunks_end > cache_length: - visible_chunks_end = cache_length - visible_chunks_begin = visible_chunks_end - Int32(self.window_size) - if visible_chunks_begin < Int32(0): - visible_chunks_begin = Int32(0) - first_valid_chunk = visible_chunks_begin // Int32(CANDIDATES_PER_CHUNK) - valid_chunks = (visible_chunks_end + Int32(CANDIDATES_PER_CHUNK - 1)) // Int32( + valid_chunks = (cache_length + Int32(CANDIDATES_PER_CHUNK - 1)) // Int32( CANDIDATES_PER_CHUNK ) - split_first_chunk = first_valid_chunk + split * Int32(self.chunks_per_split) + sparse_active = Int32(0) + selected_chunks = valid_chunks + if cutlass.const_expr(self.sparse_stride > 1): + if cache_length > Int32(self.sparse_min_tokens): + sparse_active = Int32(1) + if cutlass.const_expr(self.sparse_refresh_interval > 0): + refresh_position = cache_length % Int32(self.sparse_refresh_interval) + if refresh_position < query_length: + sparse_active = Int32(0) + if sparse_active != Int32(0): + sink = Int32(self.sparse_sink_chunks) + if sink > valid_chunks: + sink = valid_chunks + recent = Int32(self.sparse_recent_chunks) + if recent > valid_chunks - sink: + recent = valid_chunks - sink + middle_span = valid_chunks - recent - sink + middle_count = (middle_span + Int32(self.sparse_stride - 1)) // Int32( + self.sparse_stride + ) + selected_chunks = sink + middle_count + recent + split_first_chunk = split * Int32(self.chunks_per_split) split_last_chunk = split_first_chunk + Int32(self.chunks_per_split) - if split_last_chunk > valid_chunks: - split_last_chunk = valid_chunks + if split_last_chunk > selected_chunks: + split_last_chunk = selected_chunks active_chunks = split_last_chunk - split_first_chunk if active_chunks < Int32(0): active_chunks = Int32(0) @@ -527,7 +566,12 @@ def kernel( CANDIDATES_PER_CHUNK * self.layout.record_stride_bytes ) for local_chunk in cutlass.range(active_chunks, unroll=1): - chunk = split_first_chunk + Int32(local_chunk) + selected_index = split_first_chunk + Int32(local_chunk) + chunk = self._selected_chunk( + selected_index, + valid_chunks, + sparse_active, + ) token_begin = chunk * Int32(CANDIDATES_PER_CHUNK) buffer = Int32(0) if cutlass.const_expr(self.kv_stages == 2): @@ -546,7 +590,6 @@ def kernel( cache_length, io_lane, page_stride_bytes, - cache_record_stride_bytes, page_table_stride, page_size=self.page_size, record_bytes=self.layout.record_bytes, @@ -585,15 +628,12 @@ def kernel( query_valid = Int32(0) local_query = query_row - query_begin visible_end = cache_length - query_length + local_query + Int32(1) + if cutlass.const_expr(self.uses_query_cache_seqlens): + visible_end = Int32(query_cache_seqlens[query_row]) if visible_end < Int32(0): visible_end = Int32(0) if visible_end > cache_length: visible_end = cache_length - visible_begin = Int32(0) - if cutlass.const_expr(self.window_size is not None): - visible_begin = visible_end - Int32(self.window_size) - if visible_begin < Int32(0): - visible_begin = Int32(0) score_scale = sm_scale_log2 value_scale = Float32(1.0) @@ -628,7 +668,8 @@ def kernel( mbar_base, active_chunks, split_first_chunk, - visible_begin, + valid_chunks, + sparse_active, visible_end, query_row, head_base, @@ -657,7 +698,8 @@ def kernel( mbar_base, active_chunks, split_first_chunk, - visible_begin, + valid_chunks, + sparse_active, visible_end, query_row, head_base, @@ -686,7 +728,8 @@ def kernel( mbar_base, active_chunks, split_first_chunk, - visible_begin, + valid_chunks, + sparse_active, visible_end, query_row, head_base, @@ -715,7 +758,8 @@ def kernel( mbar_base, active_chunks, split_first_chunk, - visible_begin, + valid_chunks, + sparse_active, visible_end, query_row, head_base, diff --git a/b12x/attention/dense_mla/_io.py b/b12x/attention/dense_mla/_io.py index f6e082bac..5f91aaf6d 100644 --- a/b12x/attention/dense_mla/_io.py +++ b/b12x/attention/dense_mla/_io.py @@ -26,7 +26,6 @@ def issue_dense_page_gather( token_end: Int32, io_lane: Int32, page_stride_bytes: Int64, - record_stride_bytes_global: Int64, page_table_stride: Int64, *, page_size: cutlass.Constexpr, @@ -61,10 +60,9 @@ def issue_dense_page_gather( # Load-bearing Int64 conversions. Neither multiplication is allowed to # occur in Int32, even when benchmark page ids happen to be small. - source_offset = ( - physical_page.to(Int64) * page_stride_bytes - + in_page.to(Int64) * record_stride_bytes_global - ) + source_offset = physical_page.to(Int64) * page_stride_bytes + in_page.to( + Int64 + ) * Int64(record_bytes) cp_async_bulk_g2s_mbar( kv_dst_addr + entry * Int32(record_stride_bytes), get_ptr_as_int64(cache_bytes, source_offset), diff --git a/b12x/attention/dense_mla/_kernel.py b/b12x/attention/dense_mla/_kernel.py index 6f14d2470..86ccb87ab 100644 --- a/b12x/attention/dense_mla/_kernel.py +++ b/b12x/attention/dense_mla/_kernel.py @@ -21,7 +21,6 @@ tensor_key, ) -from .._shared import static_fp8_quant from ._forward import DenseMlaForwardKernel from ._layout import make_smem_layout from ._merge import DenseMlaMergeKernel @@ -80,13 +79,15 @@ def _signature(binding: Binding) -> tuple[object, ...]: scratch.page_size, scratch.max_total_q, scratch.max_batch, + scratch.uses_query_cache_seqlens, + scratch.sparse_stride, + scratch.sparse_min_tokens, + scratch.sparse_sink_chunks, + scratch.sparse_recent_chunks, + scratch.sparse_refresh_interval, scratch.query_tile, scratch.num_splits, scratch.chunks_per_split, - scratch.physical_record_width, - scratch.head_dim, - scratch.v_head_dim, - scratch.window_size, tuple(int(value) for value in binding.q.stride()), tuple(int(value) for value in binding.kv_cache.stride()), tuple(int(value) for value in binding.output.stride()), @@ -118,7 +119,6 @@ def _forward_launch(binding: Binding) -> _ForwardLaunch: layout = make_smem_layout( query_tile=scratch.query_tile, fp8=fp8, - qk_dim=scratch.head_dim, ) entry = DenseMlaForwardKernel( layout=layout, @@ -127,10 +127,13 @@ def _forward_launch(binding: Binding) -> _ForwardLaunch: num_splits=scratch.num_splits, chunks_per_split=scratch.chunks_per_split, query_tile=scratch.query_tile, + uses_query_cache_seqlens=scratch.uses_query_cache_seqlens, + sparse_stride=scratch.sparse_stride, + sparse_min_tokens=scratch.sparse_min_tokens, + sparse_sink_chunks=scratch.sparse_sink_chunks, + sparse_recent_chunks=scratch.sparse_recent_chunks, + sparse_refresh_interval=scratch.sparse_refresh_interval, fp8=fp8, - qk_dim=scratch.head_dim, - value_dim=scratch.v_head_dim, - window_size=scratch.window_size, ) q_bytes = _byte_base_pointer(binding.q) @@ -179,6 +182,12 @@ def _forward_launch(binding: Binding) -> _ForwardLaunch: align=4, dynamic_layout=True, ), + _to_cute( + binding.query_cache_seqlens, + cutlass.Int32, + align=4, + dynamic_layout=True, + ), _to_cute( output, cutlass.BFloat16, @@ -194,7 +203,6 @@ def _forward_launch(binding: Binding) -> _ForwardLaunch: Int64(int(binding.q.stride(0)) * binding.q.element_size()), Int64(int(binding.q.stride(1)) * binding.q.element_size()), Int64(int(binding.kv_cache.stride(0)) * binding.kv_cache.element_size()), - Int64(int(binding.kv_cache.stride(1)) * binding.kv_cache.element_size()), Int64(int(binding.page_table.stride(0))), Int32(int(binding.q.shape[0])), Int32(int(binding.cache_seqlens.shape[0])), @@ -203,19 +211,24 @@ def _forward_launch(binding: Binding) -> _ForwardLaunch: ) spec = KernelCompileSpec.from_fields( "attention.dense_mla.forward", - 4, + 5, key_field("dtype", "fp8" if fp8 else "bf16"), key_field("heads", scratch.num_q_heads), key_field("page_size", scratch.page_size), key_field("max_total_q", scratch.max_total_q), key_field("max_batch", scratch.max_batch), key_field("query_tile", scratch.query_tile), + key_field( + "uses_query_cache_seqlens", + scratch.uses_query_cache_seqlens, + ), + key_field("sparse_stride", scratch.sparse_stride), + key_field("sparse_min_tokens", scratch.sparse_min_tokens), + key_field("sparse_sink_chunks", scratch.sparse_sink_chunks), + key_field("sparse_recent_chunks", scratch.sparse_recent_chunks), + key_field("sparse_refresh_interval", scratch.sparse_refresh_interval), key_field("num_splits", scratch.num_splits), key_field("chunks_per_split", scratch.chunks_per_split), - key_field("physical_record_width", scratch.physical_record_width), - key_field("head_dim", scratch.head_dim), - key_field("v_head_dim", scratch.v_head_dim), - key_field("window_size", scratch.window_size), key_field("record_stride_bytes", layout.record_stride_bytes), tensor_key( "q", @@ -223,7 +236,7 @@ def _forward_launch(binding: Binding) -> _ForwardLaunch: dims=( DimKey.dynamic(), DimKey.exact(scratch.num_q_heads), - DimKey.exact(scratch.head_dim), + DimKey.exact(576), ), ), tensor_key( @@ -232,7 +245,7 @@ def _forward_launch(binding: Binding) -> _ForwardLaunch: dims=( DimKey.dynamic(), DimKey.exact(scratch.page_size), - DimKey.exact(scratch.physical_record_width), + DimKey.exact(576), ), ), tensor_key( @@ -241,7 +254,7 @@ def _forward_launch(binding: Binding) -> _ForwardLaunch: dims=( DimKey.dynamic(), DimKey.exact(scratch.num_q_heads), - DimKey.exact(scratch.v_head_dim), + DimKey.exact(512), ), ), ) @@ -254,7 +267,7 @@ def _merge_launch(binding: Binding) -> _MergeLaunch | None: return None assert scratch.partial_output is not None assert scratch.partial_lse is not None - entry = DenseMlaMergeKernel(scratch.num_splits, scratch.v_head_dim) + entry = DenseMlaMergeKernel(scratch.num_splits) stream = cuda.CUstream(torch.cuda.current_stream().cuda_stream) args = ( _to_cute( @@ -278,7 +291,6 @@ def _merge_launch(binding: Binding) -> _MergeLaunch | None: 4, key_field("num_splits", scratch.num_splits), key_field("heads", scratch.num_q_heads), - key_field("v_head_dim", scratch.v_head_dim), tensor_key( "partial_output", scratch.partial_output, @@ -286,7 +298,7 @@ def _merge_launch(binding: Binding) -> _MergeLaunch | None: DimKey.capacity(scratch.max_total_q), DimKey.exact(scratch.num_q_heads), DimKey.exact(scratch.num_splits), - DimKey.exact(scratch.v_head_dim), + DimKey.exact(512), ), ), tensor_key( @@ -295,7 +307,7 @@ def _merge_launch(binding: Binding) -> _MergeLaunch | None: dims=( DimKey.dynamic(), DimKey.exact(scratch.num_q_heads), - DimKey.exact(scratch.v_head_dim), + DimKey.exact(512), ), ), ) @@ -333,8 +345,6 @@ def compile_dense_mla(*, binding: Binding) -> None: """Compile the exact native forward/merge entries without launching.""" if not binding.q.is_cuda: raise ValueError("dense MLA native compilation requires CUDA tensors") - if binding.query_quant is not None: - static_fp8_quant.compile(binding=binding.query_quant) _compile_entries(binding) @@ -345,8 +355,6 @@ def run_dense_mla( """Launch only previously planned storage; no device allocation occurs.""" if not binding.q.is_cuda: raise ValueError("dense MLA native execution requires CUDA tensors") - if binding.query_quant is not None: - static_fp8_quant.run(binding=binding.query_quant) signature = _signature(binding) with _LOCK: compiled_forward = _FORWARD_CACHE.get(signature) @@ -374,7 +382,6 @@ def run_dense_mla( def clear_dense_mla_kernel_caches() -> None: - static_fp8_quant.clear_caches() with _LOCK: _FORWARD_CACHE.clear() _MERGE_CACHE.clear() diff --git a/b12x/attention/dense_mla/_layout.py b/b12x/attention/dense_mla/_layout.py index ab6f417d7..4b5856771 100644 --- a/b12x/attention/dense_mla/_layout.py +++ b/b12x/attention/dense_mla/_layout.py @@ -33,9 +33,7 @@ class SmemLayout: total_bytes: int -def make_smem_layout( - *, query_tile: int, fp8: bool, qk_dim: int = K3_QK_DIM -) -> SmemLayout: +def make_smem_layout(*, query_tile: int, fp8: bool) -> SmemLayout: query_tile = int(query_tile) if query_tile not in (1, 2, 4): raise ValueError("dense MLA query_tile must be 1, 2, or 4") @@ -47,11 +45,8 @@ def make_smem_layout( # opt-in SMEM limit exposed by RTX PRO 6000 Blackwell. BF16 retains the # same producer/consumer transaction protocol with one stage; the primary # E4M3 serving path keeps the latency-hiding double buffer. - qk_dim = int(qk_dim) - if qk_dim <= 0 or qk_dim % 32: - raise ValueError("dense MLA qk_dim must be a positive multiple of 32") - kv_stages = KV_STAGES if fp8 and qk_dim <= K3_QK_DIM else 1 - record_bytes = qk_dim * element_bytes + kv_stages = KV_STAGES if fp8 else 1 + record_bytes = K3_QK_DIM * element_bytes # The extra 16 bytes rotate successive rows across banks while making # every row a legal cp.async.bulk destination. record_stride_bytes = record_bytes + 16 diff --git a/b12x/attention/dense_mla/_math.py b/b12x/attention/dense_mla/_math.py index ae8d92925..e40b02b4c 100644 --- a/b12x/attention/dense_mla/_math.py +++ b/b12x/attention/dense_mla/_math.py @@ -39,6 +39,8 @@ from ._layout import ( CANDIDATES_PER_CHUNK, HEADS_PER_TILE, + K3_QK_DIM, + K3_VALUE_DIM, MATH_WARPS_PER_QUERY, ) @@ -155,7 +157,6 @@ def qk_fp8( lane: Int32, *, record_stride_bytes: cutlass.Constexpr, - qk_dim: cutlass.Constexpr, ): """Raw E4M3 K-by-Q MMA for one 16-candidate warp tile.""" a_row = (lane & Int32(7)) + ((lane >> Int32(3)) & Int32(1)) * Int32(8) @@ -164,7 +165,7 @@ def qk_fp8( b_col = ((lane >> Int32(3)) & Int32(1)) * Int32(16) first_candidate = local_warp * Int32(16) - for ks in cutlass.range_constexpr(qk_dim // 32): + for ks in cutlass.range_constexpr(K3_QK_DIM // 32): ko = Int32(ks * 32) a_addr = ( kv_smem_addr @@ -199,7 +200,6 @@ def qk_bf16( lane: Int32, *, record_stride_bytes: cutlass.Constexpr, - qk_dim: cutlass.Constexpr, ): """Raw BF16 K-by-Q MMA for one 16-candidate warp tile.""" gid = lane >> Int32(2) @@ -208,7 +208,7 @@ def qk_bf16( a_col = (lane >> Int32(4)) * Int32(8) first_candidate = local_warp * Int32(16) - for ks in cutlass.range_constexpr(qk_dim // 16): + for ks in cutlass.range_constexpr(K3_QK_DIM // 16): ko = Int32(ks * 16) a_addr = ( kv_smem_addr @@ -238,7 +238,6 @@ def qk_bf16( def mask_and_scale( qk, chunk_begin: Int32, - visible_begin: Int32, visible_end: Int32, scale_log2: Float32, local_warp: Int32, @@ -248,10 +247,10 @@ def mask_and_scale( gid = lane >> Int32(2) candidate0 = chunk_begin + local_warp * Int32(16) + gid candidate1 = candidate0 + Int32(8) - if candidate0 < visible_begin or candidate0 >= visible_end: + if candidate0 >= visible_end: qk[0] = Float32(_MASK) qk[1] = Float32(_MASK) - if candidate1 < visible_begin or candidate1 >= visible_end: + if candidate1 >= visible_end: qk[2] = Float32(_MASK) qk[3] = Float32(_MASK) for idx in cutlass.range_constexpr(4): @@ -273,7 +272,6 @@ def online_softmax( local_tid: Int32, *, barrier_id: cutlass.Constexpr, - value_dim: cutlass.Constexpr, ): """Update a base-2 online softmax and rescale persistent PV registers.""" gid = lane >> Int32(2) @@ -368,7 +366,7 @@ def online_softmax( row_alpha = row_alpha0 if (gid & Int32(1)) != Int32(0): row_alpha = row_alpha1 - for tile in cutlass.range_constexpr(value_dim // 32): + for tile in cutlass.range_constexpr(16): acc[tile][0] = acc[tile][0] * row_alpha acc[tile][1] = acc[tile][1] * row_alpha @@ -474,7 +472,6 @@ def pv_fp8( *, record_stride_bytes: cutlass.Constexpr, barrier_id: cutlass.Constexpr, - value_dim: cutlass.Constexpr, ): """HIGH/LOW E4M3 probability MMA against the raw K3 value prefix.""" gid = lane >> Int32(2) @@ -566,7 +563,7 @@ def pv_fp8( a_row = (lane & Int32(7)) + ((lane >> Int32(3)) & Int32(1)) * Int32(8) a_col = (lane >> Int32(4)) * Int32(16) row_scale = ld_shared_f32(weight_scale_addr + gid * Int32(4)) - for value_chunk in cutlass.range_constexpr(value_dim // 64): + for value_chunk in cutlass.range_constexpr(K3_VALUE_DIM // 64): for n_tile in cutlass.range_constexpr(2): dimension = Int32(value_chunk * 64) + ( Int32(n_tile * MATH_WARPS_PER_QUERY) + local_warp @@ -615,7 +612,6 @@ def pv_bf16( *, record_stride_bytes: cutlass.Constexpr, barrier_id: cutlass.Constexpr, - value_dim: cutlass.Constexpr, ): """BF16 probability/value MMA for the K3 value prefix.""" gid = lane >> Int32(2) @@ -648,7 +644,7 @@ def pv_bf16( a_row = (lane & Int32(7)) + ((lane >> Int32(3)) & Int32(1)) * Int32(8) a_col = (lane >> Int32(4)) * Int32(8) - for value_chunk in cutlass.range_constexpr(value_dim // 64): + for value_chunk in cutlass.range_constexpr(K3_VALUE_DIM // 64): for n_tile in cutlass.range_constexpr(2): column = ( Int32(value_chunk * 64) @@ -725,7 +721,6 @@ def write_partial_or_final( has_splits: cutlass.Constexpr, fp8: cutlass.Constexpr, ln2: cutlass.Constexpr, - value_dim: cutlass.Constexpr, ): """Write a normalized split partial, or the final one-split result.""" gid = lane >> Int32(2) @@ -749,7 +744,7 @@ def write_partial_or_final( scale = scale * value_scale if query_valid != Int32(0) and head_base + gid < Int32(num_heads): - for value_chunk in cutlass.range_constexpr(value_dim // 64): + for value_chunk in cutlass.range_constexpr(K3_VALUE_DIM // 64): for n_tile in cutlass.range_constexpr(2): tile = value_chunk * 2 + n_tile dimension = ( diff --git a/b12x/attention/dense_mla/_merge.py b/b12x/attention/dense_mla/_merge.py index 5f26c5aa2..afe93603e 100644 --- a/b12x/attention/dense_mla/_merge.py +++ b/b12x/attention/dense_mla/_merge.py @@ -21,9 +21,8 @@ class DenseMlaMergeKernel: """Merge normalized BF16 split partials and emit natural-log LSE.""" - def __init__(self, num_splits: int, value_dim: int): + def __init__(self, num_splits: int): self.num_splits = int(num_splits) - self.value_dim = int(value_dim) def _get_shared_storage_cls(self): """Return the per-CTA split-weight storage. @@ -149,57 +148,55 @@ def kernel( cute.arch.barrier() - while dimension < Int32(self.value_dim): - accumulator = cute.make_rmem_tensor(4, Float32) - for idx in cutlass.range_constexpr(4): - accumulator[idx] = Float32(0.0) - - split = Int32(0) - while split < active_splits: - weight = Float32(split_weight[split]) - # Capture-static tail splits deliberately leave their partial - # vectors undefined and publish -inf LSE. Do not load those - # vectors: NaN * 0 would otherwise poison the merged output. - if weight > Float32(0.0): - partial_offset = ( - (Int64(query) * Int64(partial_output.shape[1]) + Int64(head)) - * Int64(self.num_splits) - + Int64(split) - ) * Int64(self.value_dim) + Int64(dimension) - packed0, packed1 = ld_global_v2_u32( - get_ptr_as_int64(partial_output, partial_offset) - ) - value0, value1 = bfloat2_to_float2_scaled( - packed0, - weight, - ) - value2, value3 = bfloat2_to_float2_scaled( - packed1, - weight, - ) - accumulator[0] += value0 - accumulator[1] += value1 - accumulator[2] += value2 - accumulator[3] += value3 - split += Int32(1) - - inverse = Float32(inverse_storage[0]) - output_offset = cute.crd2idx( - (query, head, dimension), - output.layout, - ) - st_global_v2_u32( - get_ptr_as_int64(output, output_offset), - pack_f32x2_to_bfloat2( - accumulator[0] * inverse, - accumulator[1] * inverse, - ), - pack_f32x2_to_bfloat2( - accumulator[2] * inverse, - accumulator[3] * inverse, - ), - ) - dimension += Int32(128 * 4) + accumulator = cute.make_rmem_tensor(4, Float32) + for idx in cutlass.range_constexpr(4): + accumulator[idx] = Float32(0.0) + + split = Int32(0) + while split < active_splits: + weight = Float32(split_weight[split]) + # Capture-static tail splits deliberately leave their partial + # vectors undefined and publish -inf LSE. Do not load those + # vectors: NaN * 0 would otherwise poison the merged output. + if weight > Float32(0.0): + partial_offset = ( + (Int64(query) * Int64(partial_output.shape[1]) + Int64(head)) + * Int64(self.num_splits) + + Int64(split) + ) * Int64(512) + Int64(dimension) + packed0, packed1 = ld_global_v2_u32( + get_ptr_as_int64(partial_output, partial_offset) + ) + value0, value1 = bfloat2_to_float2_scaled( + packed0, + weight, + ) + value2, value3 = bfloat2_to_float2_scaled( + packed1, + weight, + ) + accumulator[0] += value0 + accumulator[1] += value1 + accumulator[2] += value2 + accumulator[3] += value3 + split += Int32(1) + + inverse = Float32(inverse_storage[0]) + output_offset = cute.crd2idx( + (query, head, dimension), + output.layout, + ) + st_global_v2_u32( + get_ptr_as_int64(output, output_offset), + pack_f32x2_to_bfloat2( + accumulator[0] * inverse, + accumulator[1] * inverse, + ), + pack_f32x2_to_bfloat2( + accumulator[2] * inverse, + accumulator[3] * inverse, + ), + ) __all__ = ["DenseMlaMergeKernel"] diff --git a/b12x/attention/dense_mla/_reference.py b/b12x/attention/dense_mla/_reference.py index 21a10e6c7..536ca0e5f 100644 --- a/b12x/attention/dense_mla/_reference.py +++ b/b12x/attention/dense_mla/_reference.py @@ -1,4 +1,4 @@ -"""High-precision paged dense-MLA oracle.""" +"""High-precision paged dense-MLA oracle for the Kimi-K3 contract.""" from __future__ import annotations @@ -6,26 +6,25 @@ import torch +from .planner import dynamic_sparse_chunk_indices + K3_ABSORBED_DIM = 576 K3_VALUE_DIM = 512 K3_RAW_QK_DIM = 192 K3_SM_SCALE = 1.0 / math.sqrt(K3_RAW_QK_DIM) -def _cache_rank3(cache: torch.Tensor, *, head_dim: int) -> torch.Tensor: +def _cache_rank3(cache: torch.Tensor) -> torch.Tensor: if cache.ndim == 4: if int(cache.shape[2]) != 1: raise ValueError("rank-4 dense MLA cache must have one KV head") cache = cache[:, :, 0, :] if cache.ndim != 3: raise ValueError( - "cache must be [pages,page_size,physical_record_width] or its " - "singleton-head rank-4 form" - ) - if int(cache.shape[-1]) < head_dim: - raise ValueError( - f"dense MLA cache record must contain at least {head_dim} elements" + "cache must be [pages,page_size,576] or [pages,page_size,1,576]" ) + if int(cache.shape[-1]) != K3_ABSORBED_DIM: + raise ValueError(f"dense MLA cache record must be {K3_ABSORBED_DIM} wide") return cache @@ -46,32 +45,25 @@ def dense_mla_reference( cache_seqlens: torch.Tensor, cu_seqlens_q: torch.Tensor, *, + query_cache_seqlens: torch.Tensor | None = None, + sparse_stride: int = 1, + sparse_min_tokens: int = 0, + sparse_sink_chunks: int = 0, + sparse_recent_chunks: int = 0, + sparse_refresh_interval: int = 0, kv_scale: torch.Tensor | float | None = None, q_scale: torch.Tensor | float | None = None, sm_scale: float = K3_SM_SCALE, - v_head_dim: int | None = None, - window_size: int | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: - """Compute right-aligned causal dense MLA with an optional local window. + """Compute the exact right-aligned causal K3 dense MLA definition. The cache already contains the current query chunk. Query row ``i`` in a request of length ``Q`` therefore sees keys through ``cache_len - Q + i`` (inclusive). """ - if q.ndim != 3: - raise ValueError("q must have shape [total_q, heads, head_dim]") - head_dim = int(q.shape[-1]) - default_value_dims = {K3_ABSORBED_DIM: K3_VALUE_DIM, 1088: 1024} - if head_dim not in default_value_dims: - raise ValueError(f"unsupported dense MLA query width {head_dim}") - if v_head_dim is None: - v_head_dim = default_value_dims[head_dim] - v_head_dim = int(v_head_dim) - if not 0 < v_head_dim <= head_dim: - raise ValueError("v_head_dim must be in [1, head_dim]") - if window_size is not None and int(window_size) <= 0: - raise ValueError("window_size must be positive or None") - cache = _cache_rank3(kv_cache, head_dim=head_dim) + cache = _cache_rank3(kv_cache) + if q.ndim != 3 or tuple(q.shape[1:])[-1:] != (K3_ABSORBED_DIM,): + raise ValueError("q must have shape [total_q, heads, 576]") if q.dtype not in (torch.bfloat16, torch.float8_e4m3fn): raise TypeError("q must be BF16 or E4M3") if cache.dtype not in (torch.bfloat16, torch.float8_e4m3fn): @@ -87,8 +79,20 @@ def dense_mla_reference( raise ValueError("cache_seqlens shape must match page_table batch") if tuple(cu_seqlens_q.shape) != (batch + 1,): raise ValueError("cu_seqlens_q shape must be [batch + 1]") + if query_cache_seqlens is not None and ( + query_cache_seqlens.dtype != torch.int32 + or tuple(query_cache_seqlens.shape) != (int(q.shape[0]),) + ): + raise TypeError("query_cache_seqlens must be int32 with shape [total_q]") if any( - t.device != q.device for t in (cache, page_table, cache_seqlens, cu_seqlens_q) + t.device != q.device + for t in ( + cache, + page_table, + cache_seqlens, + cu_seqlens_q, + *(() if query_cache_seqlens is None else (query_cache_seqlens,)), + ) ): raise ValueError("dense MLA reference tensors must be on one device") @@ -98,22 +102,6 @@ def dense_mla_reference( raise ValueError("E4M3 kv_cache requires kv_scale") if q.dtype == torch.float8_e4m3fn and q_scale is None: raise ValueError("E4M3 q requires q_scale") - if cache.dtype == torch.bfloat16 and (kv_scale is not None or q_scale is not None): - raise ValueError("BF16 dense MLA does not accept quantization scales") - if q.dtype != cache.dtype and not ( - q.dtype == torch.bfloat16 and cache.dtype == torch.float8_e4m3fn - ): - raise TypeError( - "dense MLA reference supports matching query/cache dtypes or a " - "BF16 query with an E4M3 cache" - ) - if q.dtype == torch.bfloat16 and cache.dtype == torch.float8_e4m3fn: - if q_scale is None: - raise ValueError("BF16 query quantization for E4M3 cache requires q_scale") - quantized_q = (q.float() / q_mul).to(torch.float8_e4m3fn) - q_f32 = quantized_q.float() * q_mul - else: - q_f32 = q.float() * q_mul cu_host = [int(v) for v in cu_seqlens_q.detach().cpu().tolist()] lens_host = [int(v) for v in cache_seqlens.detach().cpu().tolist()] @@ -121,7 +109,7 @@ def dense_mla_reference( raise ValueError("cu_seqlens_q must span exactly q.shape[0] rows") page_size = int(cache.shape[1]) output = torch.empty( - (int(q.shape[0]), int(q.shape[1]), v_head_dim), + (int(q.shape[0]), int(q.shape[1]), K3_VALUE_DIM), dtype=torch.float32, device=q.device, ) @@ -143,24 +131,45 @@ def dense_mla_reference( if pages_needed > int(page_table.shape[1]): raise ValueError("page_table is too narrow for cache_seqlens") physical_pages = page_table[request, :pages_needed].to(torch.long) - records = cache.index_select(0, physical_pages).reshape( - -1, int(cache.shape[-1]) - ) - records = records[:, :head_dim] + records = cache.index_select(0, physical_pages).reshape(-1, K3_ABSORBED_DIM) records = records[:kv_len].float() * kv_mul + sparse_active = sparse_stride > 1 and kv_len > sparse_min_tokens + if sparse_refresh_interval > 0 and kv_len % sparse_refresh_interval < q_len: + sparse_active = False + if sparse_active: + selected_chunks = dynamic_sparse_chunk_indices( + (kv_len + 63) // 64, + stride=sparse_stride, + sink_chunks=sparse_sink_chunks, + recent_chunks=sparse_recent_chunks, + ) + selected_positions = torch.cat( + [ + torch.arange( + chunk * 64, + min((chunk + 1) * 64, kv_len), + device=q.device, + ) + for chunk in selected_chunks + ] + ) + else: + selected_positions = torch.arange(kv_len, device=q.device) - q_rows = q_f32[q_begin:q_end] + q_rows = q[q_begin:q_end].float() * q_mul for local_q in range(q_len): - visible = kv_len - q_len + local_q + 1 - visible_begin = ( - 0 if window_size is None else max(0, visible - int(window_size)) + visible = ( + int(query_cache_seqlens[q_begin + local_q]) + if query_cache_seqlens is not None + else kv_len - q_len + local_q + 1 ) - key = records[visible_begin:visible] + visible_positions = selected_positions[selected_positions < visible] + key = records.index_select(0, visible_positions) logits = torch.einsum("hd,kd->hk", q_rows[local_q], key) logits = logits * float(sm_scale) probs = torch.softmax(logits, dim=-1) output[q_begin + local_q] = torch.einsum( - "hk,kd->hd", probs, key[:, :v_head_dim] + "hk,kd->hd", probs, key[:, :K3_VALUE_DIM] ) lse[q_begin + local_q] = torch.logsumexp(logits, dim=-1) diff --git a/b12x/attention/dense_mla/_scratch.py b/b12x/attention/dense_mla/_scratch.py index 37235c1d4..9a32fca71 100644 --- a/b12x/attention/dense_mla/_scratch.py +++ b/b12x/attention/dense_mla/_scratch.py @@ -20,7 +20,6 @@ materialize_scratch_view, ) -from .._shared import static_fp8_quant from .planner import Budget, choose_num_splits from ._layout import make_smem_layout from ._reference import ( @@ -30,7 +29,7 @@ ) _FP8 = torch.float8_e4m3fn -_MAX_Q_ROWS = 65_536 +_MAX_Q_ROWS = 1_024 _MAX_CACHE_TOKENS = 1_048_576 @@ -42,20 +41,15 @@ def _canonical_device(device: torch.device | str) -> torch.device: def _query_tile(caps: Caps) -> int: - if caps.mode == "decode" or caps.max_batch != 1 or caps.window_size is not None: + if caps.mode == "verify" and caps.max_total_q == caps.max_batch * 4: + return 4 if caps.kv_dtype == _FP8 else 1 + if caps.mode == "decode" or caps.max_batch != 1: return 1 if caps.kv_dtype == _FP8 and caps.max_total_q >= 3: return 4 return 2 -def _max_attended_tokens(caps: Caps) -> int: - if caps.window_size is None: - return caps.max_cache_tokens - # A local interval can begin and end in partial physical pages. - return min(caps.max_cache_tokens, caps.window_size + caps.page_size - 1) - - @dataclass(frozen=True, kw_only=True) class Caps: device: torch.device | str @@ -69,12 +63,15 @@ class Caps: max_page_table_width: int num_cache_pages: int dtype: torch.dtype = torch.bfloat16 - q_dtype: torch.dtype | None = None head_dim: int = K3_ABSORBED_DIM v_head_dim: int = K3_VALUE_DIM - physical_record_width: int = K3_ABSORBED_DIM - window_size: int | None = None use_cuda_graph: bool = False + uses_query_cache_seqlens: bool = False + sparse_stride: int = 1 + sparse_min_tokens: int = 0 + sparse_sink_chunks: int = 0 + sparse_recent_chunks: int = 0 + sparse_refresh_interval: int = 0 budget: Budget | None = None def __post_init__(self) -> None: @@ -82,35 +79,19 @@ def __post_init__(self) -> None: if self.mode not in ("decode", "extend", "verify"): raise ValueError(f"unsupported dense MLA mode {self.mode!r}") if self.dtype != torch.bfloat16: - raise TypeError("dense MLA output must be torch.bfloat16") + raise TypeError("K3 dense MLA activations/output must be torch.bfloat16") if self.kv_dtype not in (torch.bfloat16, _FP8): - raise TypeError("dense MLA KV must be BF16 or E4M3") - q_dtype = self.kv_dtype if self.q_dtype is None else self.q_dtype - if q_dtype not in (torch.bfloat16, _FP8): - raise TypeError("dense MLA query must be BF16 or E4M3") - if q_dtype != self.kv_dtype and not ( - q_dtype == torch.bfloat16 and self.kv_dtype == _FP8 - ): - raise TypeError( - "dense MLA supports matching query/cache dtypes or a BF16 " - "query with an E4M3 cache" - ) - object.__setattr__(self, "q_dtype", q_dtype) - geometry = (int(self.head_dim), int(self.v_head_dim)) - if geometry not in ((K3_ABSORBED_DIM, K3_VALUE_DIM), (1088, 1024)): - raise ValueError( - "dense MLA supports logical (QK, V) widths (576, 512) and (1088, 1024)" - ) - if int(self.physical_record_width) < int(self.head_dim): - raise ValueError("physical_record_width must cover the logical key record") - if self.window_size is not None and int(self.window_size) <= 0: - raise ValueError("window_size must be positive or None") + raise TypeError("K3 dense MLA KV must be BF16 or E4M3") + if int(self.head_dim) != K3_ABSORBED_DIM: + raise ValueError(f"K3 dense MLA head_dim must be {K3_ABSORBED_DIM}") + if int(self.v_head_dim) != K3_VALUE_DIM: + raise ValueError(f"K3 dense MLA v_head_dim must be {K3_VALUE_DIM}") heads = int(self.num_q_heads) if heads <= 0: raise ValueError("num_q_heads must be positive") page_size = int(self.page_size) - if page_size <= 0 or (page_size != 1 and page_size % 16): - raise ValueError("page_size must be 1 or a positive multiple of 16") + if page_size <= 0 or page_size % 16: + raise ValueError("page_size must be a positive multiple of 16") max_total_q = int(self.max_total_q) if not 1 <= max_total_q <= _MAX_Q_ROWS: raise ValueError(f"max_total_q must be in [1, {_MAX_Q_ROWS}]") @@ -130,6 +111,22 @@ def __post_init__(self) -> None: raise ValueError("num_cache_pages must be positive") if self.budget is not None and not isinstance(self.budget, Budget): raise TypeError("budget must be dense_mla.Budget or None") + uses_query_cache_seqlens = bool(self.uses_query_cache_seqlens) + if uses_query_cache_seqlens and self.mode != "verify": + raise ValueError( + "per-query cache lengths are supported only by verify plans" + ) + for name, minimum in ( + ("sparse_stride", 1), + ("sparse_min_tokens", 0), + ("sparse_sink_chunks", 0), + ("sparse_recent_chunks", 0), + ("sparse_refresh_interval", 0), + ): + value = int(getattr(self, name)) + if value < minimum: + raise ValueError(f"{name} must be >= {minimum}") + object.__setattr__(self, name, value) for name in ( "num_q_heads", "page_size", @@ -140,12 +137,14 @@ def __post_init__(self) -> None: "num_cache_pages", "head_dim", "v_head_dim", - "physical_record_width", ): object.__setattr__(self, name, int(getattr(self, name))) object.__setattr__(self, "use_cuda_graph", bool(self.use_cuda_graph)) - if self.window_size is not None: - object.__setattr__(self, "window_size", int(self.window_size)) + object.__setattr__( + self, + "uses_query_cache_seqlens", + uses_query_cache_seqlens, + ) @dataclass(frozen=True) @@ -156,21 +155,18 @@ class _ScratchLayout: partial_output_offset_bytes: int | None partial_lse_offset_bytes: int | None final_lse_offset_bytes: int - quantized_q_offset_bytes: int | None def _dense_mla_scratch_layout( caps: Caps, ) -> _ScratchLayout: query_tile = _query_tile(caps) - max_attended_tokens = _max_attended_tokens(caps) if caps.device.type == "cuda": properties = torch.cuda.get_device_properties(caps.device) sm_count = int(properties.multi_processor_count) shared_layout = make_smem_layout( query_tile=query_tile, fp8=caps.kv_dtype == _FP8, - qk_dim=caps.head_dim, ) shared_limit = int(properties.shared_memory_per_block_optin) if shared_layout.total_bytes > shared_limit: @@ -182,14 +178,14 @@ def _dense_mla_scratch_layout( else: sm_count = 1 num_splits = choose_num_splits( - max_cache_tokens=max_attended_tokens, + max_cache_tokens=caps.max_cache_tokens, max_total_q=caps.max_total_q, num_q_heads=caps.num_q_heads, query_tile=query_tile, sm_count=sm_count, budget=caps.budget, ) - num_chunks = (max_attended_tokens + 63) // 64 + num_chunks = (caps.max_cache_tokens + 63) // 64 chunks_per_split = (num_chunks + num_splits - 1) // num_splits partial_rows = caps.max_total_q * num_splits if ( @@ -213,7 +209,7 @@ def _dense_mla_scratch_layout( caps.max_total_q * caps.num_q_heads * num_splits - * caps.v_head_dim + * K3_VALUE_DIM * dtype_nbytes(torch.bfloat16) ) cursor = align_up(cursor, SCRATCH_ALIGN_BYTES) @@ -228,13 +224,6 @@ def _dense_mla_scratch_layout( cursor = align_up(cursor, SCRATCH_ALIGN_BYTES) final_lse_offset_bytes = cursor cursor += caps.max_total_q * caps.num_q_heads * dtype_nbytes(torch.float32) - quantized_q_offset_bytes = None - if caps.q_dtype != caps.kv_dtype: - cursor = align_up(cursor, SCRATCH_ALIGN_BYTES) - quantized_q_offset_bytes = cursor - cursor += ( - caps.max_total_q * caps.num_q_heads * caps.head_dim * dtype_nbytes(_FP8) - ) cursor = align_up(cursor, SCRATCH_ALIGN_BYTES) return _ScratchLayout( nbytes=max(cursor, SCRATCH_ALIGN_BYTES), @@ -243,7 +232,6 @@ def _dense_mla_scratch_layout( partial_output_offset_bytes=partial_output_offset_bytes, partial_lse_offset_bytes=partial_lse_offset_bytes, final_lse_offset_bytes=final_lse_offset_bytes, - quantized_q_offset_bytes=quantized_q_offset_bytes, ) @@ -252,7 +240,6 @@ class Scratch: shared_scratch: torch.Tensor device: torch.device mode: str - q_dtype: torch.dtype kv_dtype: torch.dtype num_q_heads: int page_size: int @@ -261,18 +248,19 @@ class Scratch: max_cache_tokens: int max_page_table_width: int num_cache_pages: int - physical_record_width: int - head_dim: int - v_head_dim: int - window_size: int | None num_splits: int chunks_per_split: int query_tile: int use_cuda_graph: bool + uses_query_cache_seqlens: bool + sparse_stride: int + sparse_min_tokens: int + sparse_sink_chunks: int + sparse_recent_chunks: int + sparse_refresh_interval: int partial_output: torch.Tensor | None partial_lse: torch.Tensor | None final_lse: torch.Tensor - quantized_q: torch.Tensor | None @dataclass(frozen=True, kw_only=True) @@ -283,12 +271,12 @@ class Binding: output: torch.Tensor page_table: torch.Tensor cache_seqlens: torch.Tensor + query_cache_seqlens: torch.Tensor cu_seqlens_q: torch.Tensor kv_scale: torch.Tensor | None q_scale: torch.Tensor | None sm_scale: float active_splits: int - query_quant: static_fp8_quant.Binding | None def _storage_bounds(tensor: torch.Tensor, *, name: str) -> None: @@ -350,8 +338,7 @@ def _rank3_cache( cache = cache[:, :, 0, :] if cache.ndim != 3: raise ValueError( - "kv_cache must be [pages,page_size,physical_record_width] or " - "[pages,page_size,1,physical_record_width]" + "kv_cache must be [pages,page_size,576] or [pages,page_size,1,576]" ) if cache.dtype != scratch.kv_dtype: raise TypeError( @@ -359,26 +346,18 @@ def _rank3_cache( ) if cache.device != scratch.device: raise ValueError("kv_cache device does not match dense MLA plan") - if tuple(cache.shape[1:]) != ( - scratch.page_size, - scratch.physical_record_width, - ): + if tuple(cache.shape[1:]) != (scratch.page_size, K3_ABSORBED_DIM): raise ValueError( "kv_cache inner shape must be " - f"({scratch.page_size}, {scratch.physical_record_width}), got " - f"{tuple(cache.shape[1:])}" + f"({scratch.page_size}, {K3_ABSORBED_DIM}), got {tuple(cache.shape[1:])}" ) if int(cache.shape[0]) > scratch.num_cache_pages: raise ValueError("kv_cache page count exceeds planned capacity") - if ( - int(cache.stride(2)) != 1 - or int(cache.stride(1)) < scratch.physical_record_width - ): + if int(cache.stride(2)) != 1 or int(cache.stride(1)) != K3_ABSORBED_DIM: raise ValueError( - "kv_cache requires contiguous elements and a token stride covering " - "physical_record_width" + "kv_cache requires contiguous records/tokens (stride[-1]=1, stride[-2]=576)" ) - if int(cache.stride(0)) < scratch.page_size * int(cache.stride(1)): + if int(cache.stride(0)) < scratch.page_size * K3_ABSORBED_DIM: raise ValueError("kv_cache page stride is smaller than one page payload") _require_16_byte_records( cache, @@ -397,6 +376,7 @@ def _validate_binding( output: torch.Tensor, page_table: torch.Tensor, cache_seqlens: torch.Tensor, + query_cache_seqlens: torch.Tensor | None, cu_seqlens_q: torch.Tensor, kv_scale: torch.Tensor | None, q_scale: torch.Tensor | None, @@ -405,23 +385,21 @@ def _validate_binding( ) -> Binding: if q.ndim != 3 or tuple(q.shape[1:]) != ( scratch.num_q_heads, - scratch.head_dim, + K3_ABSORBED_DIM, ): raise ValueError( "q must have shape " - f"[total_q,{scratch.num_q_heads},{scratch.head_dim}], got {tuple(q.shape)}" - ) - if q.dtype != scratch.q_dtype: - raise TypeError( - f"q dtype {q.dtype} does not match planned query dtype {scratch.q_dtype}" + f"[total_q,{scratch.num_q_heads},{K3_ABSORBED_DIM}], got {tuple(q.shape)}" ) + if q.dtype not in (torch.bfloat16, _FP8): + raise TypeError("q must be BF16 or E4M3") if q.device != scratch.device: raise ValueError("q device does not match dense MLA plan") if not 1 <= int(q.shape[0]) <= scratch.max_total_q: raise ValueError("q rows exceed planned capacity") - if int(q.stride(2)) != 1 or int(q.stride(1)) != scratch.head_dim: + if int(q.stride(2)) != 1 or int(q.stride(1)) != K3_ABSORBED_DIM: raise ValueError("q requires contiguous absorbed head records") - if int(q.stride(0)) < scratch.num_q_heads * scratch.head_dim: + if int(q.stride(0)) < scratch.num_q_heads * K3_ABSORBED_DIM: raise ValueError("q row stride overlaps absorbed head records") _require_16_byte_records( q, @@ -431,21 +409,26 @@ def _validate_binding( _storage_bounds(q, name="q") cache = _rank3_cache(kv_cache, scratch=scratch) + if q.dtype != cache.dtype: + raise TypeError( + "q and kv_cache must use the same native dense MLA format " + "(BF16/BF16 or E4M3/E4M3)" + ) if output.ndim != 3 or tuple(output.shape) != ( int(q.shape[0]), scratch.num_q_heads, - scratch.v_head_dim, + K3_VALUE_DIM, ): raise ValueError( "output must have shape " - f"{(int(q.shape[0]), scratch.num_q_heads, scratch.v_head_dim)}, " + f"{(int(q.shape[0]), scratch.num_q_heads, K3_VALUE_DIM)}, " f"got {tuple(output.shape)}" ) if output.dtype != torch.bfloat16 or output.device != scratch.device: raise TypeError("output must be BF16 on the plan device") - if int(output.stride(2)) != 1 or int(output.stride(1)) != scratch.v_head_dim: + if int(output.stride(2)) != 1 or int(output.stride(1)) != K3_VALUE_DIM: raise ValueError("output requires contiguous latent head records") - if int(output.stride(0)) < scratch.num_q_heads * scratch.v_head_dim: + if int(output.stride(0)) < scratch.num_q_heads * K3_VALUE_DIM: raise ValueError("output row stride overlaps latent head records") _require_16_byte_records( output, @@ -465,6 +448,34 @@ def _validate_binding( raise TypeError(f"{name} must be contiguous int32 with shape {shape}") if tensor.device != scratch.device or not tensor.is_contiguous(): raise ValueError(f"{name} must be contiguous on the plan device") + if scratch.uses_query_cache_seqlens: + if query_cache_seqlens is None: + raise ValueError("verify plan requires per-query cache lengths") + if query_cache_seqlens.dtype != torch.int32 or tuple( + query_cache_seqlens.shape + ) != (int(q.shape[0]),): + raise TypeError( + "query_cache_seqlens must be contiguous int32 with shape [total_q]" + ) + if ( + query_cache_seqlens.device != scratch.device + or not query_cache_seqlens.is_contiguous() + ): + raise ValueError( + "query_cache_seqlens must be contiguous on the plan device" + ) + elif query_cache_seqlens is not None: + raise ValueError("decode plan does not accept per-query cache lengths") + else: + query_cache_seqlens = cache_seqlens + if ( + scratch.mode == "verify" + and scratch.query_tile > 1 + and int(q.shape[0]) != batch * scratch.query_tile + ): + raise ValueError( + "tiled verify plan requires one complete query tile per request" + ) if ( page_table.ndim != 2 or int(page_table.shape[0]) != batch @@ -488,24 +499,10 @@ def _validate_binding( q_scale, name="q_scale", device=scratch.device, - required=cache.dtype == _FP8, + required=q.dtype == _FP8, ) - if cache.dtype != _FP8 and (kv_scale is not None or q_scale is not None): + if q.dtype != _FP8 and (kv_scale is not None or q_scale is not None): raise ValueError("BF16 dense MLA does not accept quantization scales") - query_quant = None - native_q = q - if q.dtype != cache.dtype: - if scratch.quantized_q is None: - raise RuntimeError("dense MLA query quantization scratch is missing") - if q_scale is None: - raise RuntimeError("dense MLA query quantization scale is missing") - native_q = scratch.quantized_q[: int(q.shape[0])] - query_quant = static_fp8_quant.bind( - source=q, - output=native_q, - scale=q_scale, - max_numel=(scratch.max_total_q * scratch.num_q_heads * scratch.head_dim), - ) sm_scale = float(sm_scale) if not (sm_scale > 0.0): raise ValueError("sm_scale must be positive") @@ -517,17 +514,17 @@ def _validate_binding( ) return Binding( scratch=scratch, - q=native_q.detach(), + q=q.detach(), kv_cache=cache, output=output, page_table=page_table.detach(), cache_seqlens=cache_seqlens.detach(), + query_cache_seqlens=query_cache_seqlens.detach(), cu_seqlens_q=cu_seqlens_q.detach(), kv_scale=kv_scale, q_scale=q_scale, sm_scale=sm_scale, active_splits=active_splits, - query_quant=query_quant, ) @@ -548,7 +545,7 @@ def _materialize( caps.max_total_q, caps.num_q_heads, layout.num_splits, - caps.v_head_dim, + K3_VALUE_DIM, ), dtype=torch.bfloat16, ) @@ -564,20 +561,11 @@ def _materialize( shape=(caps.max_total_q, caps.num_q_heads), dtype=torch.float32, ) - quantized_q = None - if layout.quantized_q_offset_bytes is not None: - quantized_q, _ = materialize_scratch_view( - scratch_storage, - offset_bytes=layout.quantized_q_offset_bytes, - shape=(caps.max_total_q, caps.num_q_heads, caps.head_dim), - dtype=_FP8, - ) query_tile = _query_tile(caps) return Scratch( shared_scratch=scratch_storage, device=caps.device, mode=caps.mode, - q_dtype=caps.q_dtype, kv_dtype=caps.kv_dtype, num_q_heads=caps.num_q_heads, page_size=caps.page_size, @@ -586,18 +574,19 @@ def _materialize( max_cache_tokens=caps.max_cache_tokens, max_page_table_width=caps.max_page_table_width, num_cache_pages=caps.num_cache_pages, - physical_record_width=caps.physical_record_width, - head_dim=caps.head_dim, - v_head_dim=caps.v_head_dim, - window_size=caps.window_size, num_splits=layout.num_splits, chunks_per_split=layout.chunks_per_split, query_tile=query_tile, use_cuda_graph=caps.use_cuda_graph, + uses_query_cache_seqlens=caps.uses_query_cache_seqlens, + sparse_stride=caps.sparse_stride, + sparse_min_tokens=caps.sparse_min_tokens, + sparse_sink_chunks=caps.sparse_sink_chunks, + sparse_recent_chunks=caps.sparse_recent_chunks, + sparse_refresh_interval=caps.sparse_refresh_interval, partial_output=partial_output, partial_lse=partial_lse, final_lse=final_lse, - quantized_q=quantized_q, ) @@ -634,6 +623,7 @@ def bind( output: torch.Tensor, page_table: torch.Tensor, cache_seqlens: torch.Tensor, + query_cache_seqlens: torch.Tensor | None = None, cu_seqlens_q: torch.Tensor, kv_scale: torch.Tensor | None = None, q_scale: torch.Tensor | None = None, @@ -653,6 +643,7 @@ def bind( output=output, page_table=page_table, cache_seqlens=cache_seqlens, + query_cache_seqlens=query_cache_seqlens, cu_seqlens_q=cu_seqlens_q, kv_scale=kv_scale, q_scale=q_scale, diff --git a/b12x/attention/dense_mla/api.py b/b12x/attention/dense_mla/api.py index 198d69dcd..f8e977e8a 100644 --- a/b12x/attention/dense_mla/api.py +++ b/b12x/attention/dense_mla/api.py @@ -1,4 +1,4 @@ -"""Public planned API for paged dense MLA.""" +"""Public planned API for dense Kimi-K3 MLA.""" from __future__ import annotations @@ -17,6 +17,7 @@ ) from .planner import Budget from .planner import ( + dynamic_sparse_chunk_indices, infer_dense_mla_mode, ) from ._reference import dense_mla_reference @@ -43,7 +44,7 @@ def run(*, binding: Binding) -> tuple[torch.Tensor, torch.Tensor]: def reference(*args, **kwargs): - """Run the FP32 paged dense-MLA oracle.""" + """Run the FP32 paged K3 dense-MLA oracle.""" return dense_mla_reference(*args, **kwargs) @@ -55,10 +56,9 @@ def is_supported(device=None) -> bool: """True only for the fail-closed SM120/SM121 production envelope.""" if not default_is_supported(device, requires=META.requires): return False - if device is None: - device = torch.device("cuda", torch.cuda.current_device()) - else: - device = torch.device(device) + device = torch.device( + device if device is not None else ("cuda", torch.cuda.current_device()) + ) return tuple(torch.cuda.get_device_capability(device)) in ((12, 0), (12, 1)) @@ -75,6 +75,7 @@ def clear_caches() -> None: "bind", "clear_caches", "compile", + "dynamic_sparse_chunk_indices", "infer_mode", "is_supported", "plan", diff --git a/b12x/attention/dense_mla/planner.py b/b12x/attention/dense_mla/planner.py index c834b7ace..88bb6e02a 100644 --- a/b12x/attention/dense_mla/planner.py +++ b/b12x/attention/dense_mla/planner.py @@ -13,6 +13,26 @@ _MAX_CEIL_WAVES = 3 +def dynamic_sparse_chunk_indices( + num_chunks: int, + *, + stride: int, + sink_chunks: int, + recent_chunks: int, +) -> tuple[int, ...]: + """Return sorted sink/strided-history/recent chunk indices.""" + num_chunks = max(int(num_chunks), 0) + stride = max(int(stride), 1) + sink = min(max(int(sink_chunks), 0), num_chunks) + recent = min(max(int(recent_chunks), 0), num_chunks - sink) + middle_end = num_chunks - recent + return tuple( + [*range(sink)] + + [*range(sink, middle_end, stride)] + + [*range(middle_end, num_chunks)] + ) + + @dataclass(frozen=True, kw_only=True) class Budget: """Optional caller capacity clamps; B12X still chooses launch policy.""" @@ -119,5 +139,6 @@ def choose_num_splits( "Budget", "Mode", "choose_num_splits", + "dynamic_sparse_chunk_indices", "infer_dense_mla_mode", ] diff --git a/b12x/attention/paged/forward_paged.py b/b12x/attention/paged/forward_paged.py index ffa385440..c5c45971c 100644 --- a/b12x/attention/paged/forward_paged.py +++ b/b12x/attention/paged/forward_paged.py @@ -8227,7 +8227,15 @@ def __call__( tma_atom_K, tma_atom_V, ).launch( - grid=(mBlockValidMask.shape[0], mKCache.shape[2], 1), + grid=( + ( + self.analytic_verify_max_chunks, + mKCache.shape[2], + self.analytic_verify_batch, + ) + if self.use_q64_laguna_verifier + else (mBlockValidMask.shape[0], mKCache.shape[2], 1) + ), block=[32, 1, 4], # 67,584 B including barriers/alignment on SM120: one CTA/SM. min_blocks_per_mp=1, @@ -8998,7 +9006,7 @@ def __call__( tma_atom_K, tma_atom_V, ).launch( - grid=(mBlockValidMask.shape[0], mKCache.shape[2], 1), + grid=launch_grid, block=[32, 4, 1], # 99,328 B including barriers/alignment on SM120: one CTA/SM. min_blocks_per_mp=1, diff --git a/b12x/attention/sparse_mla/_paged_index_remap.py b/b12x/attention/sparse_mla/_paged_index_remap.py deleted file mode 100644 index e6d467425..000000000 --- a/b12x/attention/sparse_mla/_paged_index_remap.py +++ /dev/null @@ -1,376 +0,0 @@ -"""Native request-relative to physical-slot remapping for paged sparse MLA.""" - -from __future__ import annotations - -from dataclasses import dataclass -from threading import RLock - -import cuda.bindings.driver as cuda -import cutlass -import cutlass.cute as cute -import torch -from cutlass import Int32, Int64 -from cutlass.cute.runtime import from_dlpack - -from b12x._lib.compiler import KernelCompileSpec, compile as compile_cute -from b12x._lib.compiler import key_field, run_compiled -from b12x._lib.utils import current_cuda_stream - -_LOCK = RLock() -_CACHE: dict[tuple[int, int, bool], object] = {} -BLOCK_SIZE = 64 -TOPK = 2048 -THREADS = 32 -INDICES_PER_THREAD = TOPK // THREADS - - -def _flat_i32(tensor: torch.Tensor): - converted = from_dlpack(tensor.reshape(-1), assumed_align=4) - converted.element_type = cutlass.Int32 - return converted.mark_layout_dynamic(leading_dim=0) - - -class _RemapKernel: - def __init__(self, max_q_rows: int, request_relative: bool): - self.max_q_rows = int(max_q_rows) - self.request_relative = bool(request_relative) - - @cute.jit - def __call__( - self, - request_ids: cute.Tensor, - block_table: cute.Tensor, - logical_indices: cute.Tensor, - physical_indices: cute.Tensor, - input_counts: cute.Tensor, - selected_counts: cute.Tensor, - active_rows: Int32, - request_count: Int32, - table_width: Int32, - num_cache_blocks: Int32, - block_stride_records: Int64, - token_stride_records: Int64, - stream: cuda.CUstream, - ): - self.kernel( - request_ids, - block_table, - logical_indices, - physical_indices, - input_counts, - selected_counts, - active_rows, - request_count, - table_width, - num_cache_blocks, - block_stride_records, - token_stride_records, - ).launch( - grid=(self.max_q_rows, 1, 1), - block=(THREADS, 1, 1), - stream=stream, - ) - - @cute.kernel - def kernel( - self, - request_ids: cute.Tensor, - block_table: cute.Tensor, - logical_indices: cute.Tensor, - physical_indices: cute.Tensor, - input_counts: cute.Tensor, - selected_counts: cute.Tensor, - active_rows: Int32, - request_count: Int32, - table_width: Int32, - num_cache_blocks: Int32, - block_stride_records: Int64, - token_stride_records: Int64, - ): - row_idx, _, _ = cute.arch.block_idx() - row = Int32(row_idx) - if row < active_rows: - thread = Int32(cute.arch.thread_idx()[0]) - request = Int32(0) - if cutlass.const_expr(self.request_relative): - request = request_ids[row].to(Int32) - column_begin = thread * Int32(INDICES_PER_THREAD) - count = Int32(0) - for offset in cutlass.range(Int32(INDICES_PER_THREAD), unroll=1): - column = column_begin + offset - logical_offset = Int64(row) * Int64(TOPK) + Int64(column) - logical_slot = logical_indices[logical_offset].to(Int32) - logical_block = logical_slot // Int32(BLOCK_SIZE) - valid = (logical_slot >= Int32(0)) & (logical_block >= Int32(0)) - if cutlass.const_expr(self.request_relative): - valid = ( - valid - & (request >= Int32(0)) - & (request < request_count) - & (logical_block < table_width) - ) - else: - valid = valid & (column < input_counts[row].to(Int32)) - if valid: - physical_block = logical_block - if cutlass.const_expr(self.request_relative): - table_offset = request.to(Int64) * table_width.to( - Int64 - ) + logical_block.to(Int64) - physical_block = block_table[table_offset].to(Int32) - if (physical_block >= Int32(0)) & ( - physical_block < num_cache_blocks - ): - count += Int32(1) - - # The first warp of the output row is temporary storage for the - # per-thread counts. Every thread loads its stable prefix before - # any thread overwrites the temporary values. - row_offset = Int64(row) * Int64(TOPK) - physical_indices[row_offset + thread.to(Int64)] = count - cute.arch.sync_threads() - output_base = Int32(0) - total = Int32(0) - for peer in cutlass.range_constexpr(THREADS): - peer_count = physical_indices[row_offset + Int64(peer)].to(Int32) - if Int32(peer) < thread: - output_base += peer_count - total += peer_count - cute.arch.sync_threads() - if thread == Int32(0): - selected_counts[row] = total - - local_output = Int32(0) - for offset in cutlass.range(Int32(INDICES_PER_THREAD), unroll=1): - column = column_begin + offset - logical_offset = row_offset + column.to(Int64) - logical_slot = logical_indices[logical_offset].to(Int32) - logical_block = logical_slot // Int32(BLOCK_SIZE) - valid = (logical_slot >= Int32(0)) & (logical_block >= Int32(0)) - if cutlass.const_expr(self.request_relative): - valid = ( - valid - & (request >= Int32(0)) - & (request < request_count) - & (logical_block < table_width) - ) - else: - valid = valid & (column < input_counts[row].to(Int32)) - if valid: - physical_block = logical_block - if cutlass.const_expr(self.request_relative): - table_offset = request.to(Int64) * table_width.to( - Int64 - ) + logical_block.to(Int64) - physical_block = block_table[table_offset].to(Int32) - if (physical_block >= Int32(0)) & ( - physical_block < num_cache_blocks - ): - physical_record = ( - physical_block.to(Int64) * block_stride_records - + (logical_slot % Int32(BLOCK_SIZE)).to(Int64) - * token_stride_records - ) - output_offset = ( - row_offset + output_base.to(Int64) + local_output.to(Int64) - ) - physical_indices[output_offset] = physical_record.to(Int32) - local_output += Int32(1) - - -@dataclass(frozen=True) -class Binding: - request_ids: torch.Tensor - block_table: torch.Tensor - logical_indices: torch.Tensor - physical_indices: torch.Tensor - input_counts: torch.Tensor - selected_counts: torch.Tensor - max_q_rows: int - num_cache_blocks: int - block_stride_records: int - token_stride_records: int - request_relative: bool - - -def bind( - *, - request_ids: torch.Tensor, - block_table: torch.Tensor, - logical_indices: torch.Tensor, - physical_indices: torch.Tensor, - selected_counts: torch.Tensor, - max_q_rows: int, - num_cache_blocks: int, - block_stride_records: int, - token_stride_records: int, -) -> Binding: - rows = int(logical_indices.shape[0]) - tensors = ( - request_ids, - block_table, - logical_indices, - physical_indices, - selected_counts, - ) - if any(tensor.dtype != torch.int32 for tensor in tensors): - raise TypeError("sparse index metadata must use int32 tensors") - if any(not tensor.is_contiguous() for tensor in tensors): - raise ValueError("sparse index metadata must be contiguous") - if any(tensor.device != logical_indices.device for tensor in tensors): - raise ValueError("sparse index metadata must share one device") - if tuple(logical_indices.shape) != (rows, TOPK): - raise ValueError(f"logical_indices must have shape ({rows}, {TOPK})") - if tuple(physical_indices.shape) != (rows, TOPK): - raise ValueError(f"physical_indices must have shape ({rows}, {TOPK})") - if tuple(request_ids.shape) != (rows,): - raise ValueError(f"request_ids must have shape ({rows},)") - if tuple(selected_counts.shape) != (rows,): - raise ValueError(f"selected_counts must have shape ({rows},)") - if block_table.ndim != 2 or block_table.shape[0] <= 0 or block_table.shape[1] <= 0: - raise ValueError("block_table must be a non-empty two-dimensional tensor") - if not 0 < rows <= int(max_q_rows): - raise ValueError("logical index rows exceed the planned capacity") - if not 0 < int(num_cache_blocks) <= torch.iinfo(torch.int32).max // BLOCK_SIZE: - raise ValueError("num_cache_blocks exceeds the physical-slot int32 range") - if int(block_stride_records) <= 0 or int(token_stride_records) <= 0: - raise ValueError("physical record strides must be positive") - return Binding( - request_ids=request_ids.detach(), - block_table=block_table.detach(), - logical_indices=logical_indices.detach(), - physical_indices=physical_indices.detach(), - input_counts=selected_counts.detach(), - selected_counts=selected_counts.detach(), - max_q_rows=int(max_q_rows), - num_cache_blocks=int(num_cache_blocks), - block_stride_records=int(block_stride_records), - token_stride_records=int(token_stride_records), - request_relative=True, - ) - - -def bind_physical_slots( - *, - physical_slots: torch.Tensor, - input_counts: torch.Tensor, - physical_indices: torch.Tensor, - selected_counts: torch.Tensor, - max_q_rows: int, - num_cache_blocks: int, - block_stride_records: int, - token_stride_records: int, -) -> Binding: - rows = int(physical_slots.shape[0]) - tensors = (physical_slots, input_counts, physical_indices, selected_counts) - if any(tensor.dtype != torch.int32 for tensor in tensors): - raise TypeError("sparse index metadata must use int32 tensors") - if any(not tensor.is_contiguous() for tensor in tensors): - raise ValueError("sparse index metadata must be contiguous") - if any(tensor.device != physical_slots.device for tensor in tensors): - raise ValueError("sparse index metadata must share one device") - if tuple(physical_slots.shape) != (rows, TOPK): - raise ValueError(f"physical_slots must have shape ({rows}, {TOPK})") - if tuple(physical_indices.shape) != (rows, TOPK): - raise ValueError(f"physical_indices must have shape ({rows}, {TOPK})") - if tuple(input_counts.shape) != (rows,) or tuple(selected_counts.shape) != (rows,): - raise ValueError(f"sparse counts must have shape ({rows},)") - if not 0 < rows <= int(max_q_rows): - raise ValueError("physical slot rows exceed the planned capacity") - if not 0 < int(num_cache_blocks) <= torch.iinfo(torch.int32).max // BLOCK_SIZE: - raise ValueError("num_cache_blocks exceeds the physical-slot int32 range") - if int(block_stride_records) <= 0 or int(token_stride_records) <= 0: - raise ValueError("physical record strides must be positive") - placeholder = input_counts - return Binding( - request_ids=placeholder.detach(), - block_table=placeholder.detach(), - logical_indices=physical_slots.detach(), - physical_indices=physical_indices.detach(), - input_counts=input_counts.detach(), - selected_counts=selected_counts.detach(), - max_q_rows=int(max_q_rows), - num_cache_blocks=int(num_cache_blocks), - block_stride_records=int(block_stride_records), - token_stride_records=int(token_stride_records), - request_relative=False, - ) - - -def _signature(binding: Binding) -> tuple[int, int, bool]: - index = binding.logical_indices.device.index - if index is None: - index = torch.cuda.current_device() - return int(index), int(binding.max_q_rows), bool(binding.request_relative) - - -def _launch(binding: Binding): - entry = _RemapKernel(binding.max_q_rows, binding.request_relative) - request_count = int(binding.block_table.shape[0]) if binding.request_relative else 1 - table_width = int(binding.block_table.shape[1]) if binding.request_relative else 1 - args = ( - _flat_i32(binding.request_ids), - _flat_i32(binding.block_table), - _flat_i32(binding.logical_indices), - _flat_i32(binding.physical_indices), - _flat_i32(binding.input_counts), - _flat_i32(binding.selected_counts), - Int32(binding.logical_indices.shape[0]), - Int32(request_count), - Int32(table_width), - Int32(binding.num_cache_blocks), - Int64(binding.block_stride_records), - Int64(binding.token_stride_records), - current_cuda_stream(), - ) - spec = KernelCompileSpec.from_fields( - "attention.paged_sparse_index_remap", - 2, - key_field("max_q_rows", binding.max_q_rows), - key_field("request_relative", binding.request_relative), - ) - return entry, args, spec - - -def compile(*, binding: Binding) -> None: - signature = _signature(binding) - with _LOCK: - compiled = _CACHE.get(signature) - if compiled is None: - entry, args, spec = _launch(binding) - compiled = compile_cute(entry, *args, compile_spec=spec) - with _LOCK: - _CACHE[signature] = compiled - - -def run(*, binding: Binding) -> tuple[torch.Tensor, torch.Tensor]: - signature = _signature(binding) - with _LOCK: - compiled = _CACHE.get(signature) - if compiled is None: - if torch.cuda.is_current_stream_capturing(): - raise RuntimeError( - "sparse index remap compile miss during CUDA graph capture; " - "call compile first" - ) - compile(binding=binding) - with _LOCK: - compiled = _CACHE[signature] - _, args, _ = _launch(binding) - run_compiled(compiled, args) - return binding.physical_indices, binding.selected_counts - - -def clear_caches() -> None: - with _LOCK: - _CACHE.clear() - - -__all__ = [ - "Binding", - "bind", - "bind_physical_slots", - "clear_caches", - "compile", - "run", -] diff --git a/b12x/attention/sparse_mla/_scratch.py b/b12x/attention/sparse_mla/_scratch.py index 8cf34efcb..a727ee07a 100644 --- a/b12x/attention/sparse_mla/_scratch.py +++ b/b12x/attention/sparse_mla/_scratch.py @@ -54,6 +54,11 @@ class B12XSparseMLAScratchCaps: max_q_chunks: int | None = None page_size: int = 64 head_major_output: bool = False + # Element type of the split-K decode partials (``tmp_output``). ``dtype`` + # (bf16) rounds every split partial before the merge; ``torch.float32`` + # keeps the partials exact so the merged result is rounded once, at the + # output. Prefill-like modes have no partials and ignore this field. + partial_dtype: torch.dtype | None = None def __post_init__(self) -> None: device = torch.device(self.device) @@ -86,6 +91,13 @@ def __post_init__(self) -> None: if self.max_q_chunks is not None: object.__setattr__(self, "max_q_chunks", max(int(self.max_q_chunks), 1)) object.__setattr__(self, "page_size", max(int(self.page_size), 1)) + partial_dtype = self.dtype if self.partial_dtype is None else self.partial_dtype + if partial_dtype not in (self.dtype, torch.float32): + raise TypeError( + "partial_dtype must be the activation dtype or torch.float32, " + f"got {partial_dtype}" + ) + object.__setattr__(self, "partial_dtype", partial_dtype) @dataclass(kw_only=True) @@ -177,11 +189,32 @@ def _validate_device( ) +# Packed GLM query record: 512 E4M3 nope bytes, four fp32 pow2 tile scales +# (the ``q_sc`` values the in-kernel S0 stage would compute) and 64 bf16 rope +# values; the same framing as the 656-byte packed KV record. +PACKED_QUERY_RECORD_BYTES = 656 +PACKED_QUERY_HEAD_DIM = 576 + + +def is_packed_query(q: torch.Tensor, *, head_dim: int) -> bool: + """True for a uint8 ``(rows, heads, 656)`` packed query of a 576-dim head.""" + return ( + q.dtype == torch.uint8 + and q.ndim == 3 + and int(q.shape[-1]) == PACKED_QUERY_RECORD_BYTES + and int(head_dim) == PACKED_QUERY_HEAD_DIM + ) + + def _validate_q(q: torch.Tensor, *, scratch: object) -> torch.Tensor: if q.ndim != 3: raise ValueError(f"q must be rank-3, got {tuple(q.shape)}") - if q.dtype != scratch.dtype: - raise TypeError(f"q must have dtype {scratch.dtype}, got {q.dtype}") + packed = is_packed_query(q, head_dim=scratch.head_dim) + if q.dtype != scratch.dtype and not packed: + raise TypeError( + f"q must have dtype {scratch.dtype} (or a uint8 packed " + f"{PACKED_QUERY_RECORD_BYTES}-byte query record), got {q.dtype}" + ) _validate_device(q, scratch=scratch, name="q") if int(q.shape[0]) > int(scratch.max_total_q): raise ValueError( @@ -191,7 +224,7 @@ def _validate_q(q: torch.Tensor, *, scratch: object) -> torch.Tensor: raise ValueError( f"q heads {int(q.shape[1])} do not match scratch heads {scratch.num_q_heads}" ) - if int(q.shape[2]) != int(scratch.head_dim): + if not packed and int(q.shape[2]) != int(scratch.head_dim): raise ValueError( f"q head_dim {int(q.shape[2])} does not match scratch head_dim {scratch.head_dim}" ) @@ -325,16 +358,22 @@ def _sparse_mla_scratch_layout( if split: cursor = align_up(cursor, SCRATCH_ALIGN_BYTES) tmp_output_offset_bytes = cursor - # output_buffer aliases tmp_output[:, :, 0, :] (chunk-major stride), so no - # separate output allocation is needed for decode. + # With partials in the output dtype, output_buffer aliases + # tmp_output[:, :, 0, :] (chunk-major stride) and needs no separate + # allocation. fp32 partials cannot alias a bf16 output, so the output + # gets its own region after the partials. output_offset_bytes = cursor cursor += ( max_total_q * max_chunks_per_row * num_q_heads * v_head_dim - * dtype_nbytes(caps.dtype) + * dtype_nbytes(caps.partial_dtype) ) + if caps.partial_dtype != caps.dtype: + cursor = align_up(cursor, SCRATCH_ALIGN_BYTES) + output_offset_bytes = cursor + cursor += max_total_q * num_q_heads * v_head_dim * dtype_nbytes(caps.dtype) cursor = align_up(cursor, SCRATCH_ALIGN_BYTES) tmp_lse_offset_bytes = cursor cursor += ( @@ -400,12 +439,28 @@ def _materialize_sparse_mla_scratch( v_head_dim=v_head_dim, head_major_output=caps.head_major_output, ), - dtype=caps.dtype, - ) - output_buffer = _split_output_buffer_from_tmp( - tmp_output, - head_major_output=caps.head_major_output, + dtype=caps.partial_dtype, ) + if caps.partial_dtype == caps.dtype: + output_buffer = _split_output_buffer_from_tmp( + tmp_output, + head_major_output=caps.head_major_output, + ) + elif caps.head_major_output: + output_buffer, _ = materialize_scratch_strided_view( + scratch_storage, + offset_bytes=layout.output_offset_bytes, + shape=(max_total_q, num_q_heads, v_head_dim), + stride=(v_head_dim, max_total_q * v_head_dim, 1), + dtype=caps.dtype, + ) + else: + output_buffer, _ = materialize_scratch_view( + scratch_storage, + offset_bytes=layout.output_offset_bytes, + shape=(max_total_q, num_q_heads, v_head_dim), + dtype=caps.dtype, + ) tmp_lse, _ = materialize_scratch_view( scratch_storage, offset_bytes=layout.tmp_lse_offset_bytes, diff --git a/b12x/attention/sparse_mla/strided.py b/b12x/attention/sparse_mla/strided.py deleted file mode 100644 index a961dfa2d..000000000 --- a/b12x/attention/sparse_mla/strided.py +++ /dev/null @@ -1,503 +0,0 @@ -"""Planned strided-record sparse-MLA API.""" - -from __future__ import annotations - -import math -from dataclasses import dataclass - -import torch - -from ..._lib.gating import default_is_supported -from ..._lib.scratch import ScratchBufferSpec, scratch_buffer_spec, scratch_tensor -from ..._lib.scratch_layout import ( - SCRATCH_ALIGN_BYTES, - align_up, - materialize_scratch_view, -) -from .._shared import static_fp8_quant -from ..dense_mla._kernel import ( - clear_dense_mla_kernel_caches, - compile_dense_mla, - run_dense_mla, -) -from ..dense_mla._scratch import Binding as _NativeBinding -from ..dense_mla._scratch import Caps as _NativeCaps -from ..dense_mla._scratch import Plan as _NativePlan -from ..dense_mla._scratch import Scratch -from ..dense_mla._scratch import plan_dense_mla_scratch -from ..dense_mla.planner import Budget -from . import _paged_index_remap - -FP8 = torch.float8_e4m3fn -BLOCK_SIZE = 64 -TOTAL_HEADS = 128 -QK_DIM = 576 -VALUE_DIM = 512 -PHYSICAL_RECORD_WIDTH = 1088 -TOPK = 2048 -SM_SCALE = 1.0 / math.sqrt(192) - - -@dataclass(frozen=True, kw_only=True) -class Caps: - device: torch.device | str - num_q_heads: int - tp_size: int - max_q_rows: int - num_cache_blocks: int - max_physical_records: int | None = None - block_size: int = BLOCK_SIZE - topk: int = TOPK - kv_dtype: torch.dtype = FP8 - use_cuda_graph: bool = False - budget: Budget | None = None - - def __post_init__(self) -> None: - if int(self.tp_size) != 8: - raise ValueError("strided sparse MLA is qualified only for TP8") - if int(self.num_q_heads) != TOTAL_HEADS // int(self.tp_size): - raise ValueError("num_q_heads must equal total_heads / tp_size") - if int(self.block_size) != BLOCK_SIZE: - raise ValueError(f"strided sparse MLA block_size must be {BLOCK_SIZE}") - if int(self.topk) != TOPK: - raise ValueError(f"strided sparse MLA topk must be {TOPK}") - if self.kv_dtype != FP8: - raise TypeError("strided sparse MLA requires E4M3 KV cache") - if int(self.max_q_rows) <= 0: - raise ValueError("max_q_rows must be positive") - if int(self.num_cache_blocks) <= 0: - raise ValueError("num_cache_blocks must be positive") - if int(self.num_cache_blocks) > torch.iinfo(torch.int32).max // BLOCK_SIZE: - raise ValueError("num_cache_blocks exceeds the physical-slot int32 range") - max_physical_records = ( - int(self.num_cache_blocks) * BLOCK_SIZE - if self.max_physical_records is None - else int(self.max_physical_records) - ) - if not int(self.num_cache_blocks) * BLOCK_SIZE <= max_physical_records: - raise ValueError( - "max_physical_records must cover all contiguous physical slots" - ) - if max_physical_records > torch.iinfo(torch.int32).max: - raise ValueError("max_physical_records exceeds the index int32 range") - object.__setattr__(self, "max_physical_records", max_physical_records) - - -@dataclass(frozen=True) -class Binding: - native: _NativeBinding - query_quant: static_fp8_quant.Binding - index_remap: _paged_index_remap.Binding | None - kv_cache: torch.Tensor - selected_indices: torch.Tensor - selected_counts: torch.Tensor - - -@dataclass(frozen=True) -class Plan: - caps: Caps - native: _NativePlan - query_offset_bytes: int - physical_indices_offset_bytes: int - selected_counts_offset_bytes: int - _scratch_specs: tuple[ScratchBufferSpec, ...] - - def scratch_specs(self): - return self._scratch_specs - - def shapes_and_dtypes(self): - return tuple((spec.shape, spec.dtype) for spec in self._scratch_specs) - - @property - def num_splits(self) -> int: - return self.native.num_splits - - -def plan(caps: Caps) -> Plan: - native_caps = _NativeCaps( - device=caps.device, - mode="decode", - kv_dtype=FP8, - num_q_heads=int(caps.num_q_heads), - page_size=1, - max_total_q=int(caps.max_q_rows), - max_batch=int(caps.max_q_rows), - max_cache_tokens=TOPK, - max_page_table_width=TOPK, - num_cache_pages=int(caps.max_physical_records), - head_dim=QK_DIM, - v_head_dim=VALUE_DIM, - physical_record_width=PHYSICAL_RECORD_WIDTH, - use_cuda_graph=bool(caps.use_cuda_graph), - budget=caps.budget, - ) - native = plan_dense_mla_scratch(native_caps) - native_spec = native.scratch_specs()[0] - query_offset_bytes = align_up(native_spec.nbytes, SCRATCH_ALIGN_BYTES) - query_nbytes = int(caps.max_q_rows) * int(caps.num_q_heads) * QK_DIM - physical_indices_offset_bytes = align_up( - query_offset_bytes + query_nbytes, - SCRATCH_ALIGN_BYTES, - ) - physical_indices_nbytes = int(caps.max_q_rows) * TOPK * torch.int32.itemsize - selected_counts_offset_bytes = align_up( - physical_indices_offset_bytes + physical_indices_nbytes, - SCRATCH_ALIGN_BYTES, - ) - selected_counts_nbytes = int(caps.max_q_rows) * torch.int32.itemsize - total_nbytes = align_up( - selected_counts_offset_bytes + selected_counts_nbytes, - SCRATCH_ALIGN_BYTES, - ) - return Plan( - caps=caps, - native=native, - query_offset_bytes=query_offset_bytes, - physical_indices_offset_bytes=physical_indices_offset_bytes, - selected_counts_offset_bytes=selected_counts_offset_bytes, - _scratch_specs=( - scratch_buffer_spec( - "sparse_mla.strided.scratch", - nbytes=total_nbytes, - device=native_spec.device, - ), - ), - ) - - -def _physical_record_view( - plan: Plan, - kv_cache: torch.Tensor, -) -> tuple[torch.Tensor, int, int]: - expected_record_shape = (BLOCK_SIZE, PHYSICAL_RECORD_WIDTH) - if ( - kv_cache.dtype != FP8 - or kv_cache.ndim != 3 - or tuple(kv_cache.shape[1:]) != expected_record_shape - or not 0 < int(kv_cache.shape[0]) <= int(plan.caps.num_cache_blocks) - ): - raise ValueError( - "kv_cache must be E4M3 with shape " - f"[1..{plan.caps.num_cache_blocks}, {BLOCK_SIZE}, " - f"{PHYSICAL_RECORD_WIDTH}], got " - f"dtype={kv_cache.dtype}, shape={tuple(kv_cache.shape)}" - ) - block_stride, token_stride, element_stride = map(int, kv_cache.stride()) - if ( - element_stride != 1 - or token_stride < PHYSICAL_RECORD_WIDTH - or block_stride < BLOCK_SIZE * token_stride - or token_stride % PHYSICAL_RECORD_WIDTH - or block_stride % PHYSICAL_RECORD_WIDTH - ): - raise ValueError( - "kv_cache must use non-overlapping, whole 1088-element physical " - f"record strides, got stride={tuple(kv_cache.stride())}" - ) - block_stride_records = block_stride // PHYSICAL_RECORD_WIDTH - token_stride_records = token_stride // PHYSICAL_RECORD_WIDTH - physical_records = ( - (int(kv_cache.shape[0]) - 1) * block_stride_records - + (BLOCK_SIZE - 1) * token_stride_records - + 1 - ) - if physical_records > int(plan.caps.max_physical_records): - raise ValueError( - "kv_cache physical record span exceeds planned capacity: " - f"need {physical_records}, planned {plan.caps.max_physical_records}" - ) - flat_cache = torch.as_strided( - kv_cache, - size=(physical_records, 1, PHYSICAL_RECORD_WIDTH), - stride=(PHYSICAL_RECORD_WIDTH, PHYSICAL_RECORD_WIDTH, 1), - ) - return flat_cache, block_stride_records, token_stride_records - - -def _bind_native( - plan: Plan, - *, - scratch_storage: torch.Tensor, - q: torch.Tensor, - flat_cache: torch.Tensor, - original_cache: torch.Tensor, - output: torch.Tensor, - record_indices: torch.Tensor, - selected_counts: torch.Tensor, - cu_seqlens_q: torch.Tensor, - kv_scale: torch.Tensor, - q_scale: torch.Tensor, - active_splits: int | None = None, -) -> Binding: - if q.dtype != torch.bfloat16: - raise TypeError("strided sparse MLA absorbed query must be BF16") - rows = int(q.shape[0]) - if record_indices.dtype != torch.int32 or tuple(record_indices.shape) != ( - rows, - TOPK, - ): - raise ValueError(f"record_indices must be int32 with shape ({rows}, {TOPK})") - if selected_counts.dtype != torch.int32 or tuple(selected_counts.shape) != (rows,): - raise ValueError(f"selected_counts must be int32 with shape ({rows},)") - if tuple(cu_seqlens_q.shape) != (rows + 1,): - raise ValueError(f"cu_seqlens_q must have shape ({rows + 1},)") - q_fp8, _ = materialize_scratch_view( - scratch_storage, - offset_bytes=plan.query_offset_bytes, - shape=tuple(q.shape), - dtype=FP8, - ) - query_quant = static_fp8_quant.bind( - source=q, - output=q_fp8, - scale=q_scale, - max_numel=int(plan.caps.max_q_rows) * int(plan.caps.num_q_heads) * QK_DIM, - ) - native = plan.native.bind( - scratch=scratch_storage, - q=q_fp8, - kv_cache=flat_cache, - output=output, - page_table=record_indices, - cache_seqlens=selected_counts, - cu_seqlens_q=cu_seqlens_q, - kv_scale=kv_scale, - q_scale=q_scale, - sm_scale=SM_SCALE, - active_splits=active_splits, - ) - return Binding( - native=native, - query_quant=query_quant, - index_remap=None, - kv_cache=original_cache, - selected_indices=record_indices, - selected_counts=selected_counts, - ) - - -def _index_scratch( - plan: Plan, - scratch_storage: torch.Tensor, - rows: int, -) -> tuple[torch.Tensor, torch.Tensor]: - physical_indices, _ = materialize_scratch_view( - scratch_storage, - offset_bytes=plan.physical_indices_offset_bytes, - shape=(rows, TOPK), - dtype=torch.int32, - ) - selected_counts, _ = materialize_scratch_view( - scratch_storage, - offset_bytes=plan.selected_counts_offset_bytes, - shape=(rows,), - dtype=torch.int32, - ) - return physical_indices, selected_counts - - -def bind( - plan: Plan, - *, - scratch, - q: torch.Tensor, - kv_cache: torch.Tensor, - output: torch.Tensor, - selected_indices: torch.Tensor, - selected_counts: torch.Tensor, - cu_seqlens_q: torch.Tensor, - kv_scale: torch.Tensor, - q_scale: torch.Tensor, - active_splits: int | None = None, -) -> Binding: - scratch_storage = scratch_tensor( - scratch, - plan.scratch_specs(), - owner="strided sparse MLA", - ) - rows = int(q.shape[0]) - record_indices, remapped_counts = _index_scratch(plan, scratch_storage, rows) - flat_cache, block_stride_records, token_stride_records = _physical_record_view( - plan, kv_cache - ) - binding = _bind_native( - plan, - scratch_storage=scratch_storage, - q=q, - flat_cache=flat_cache, - original_cache=kv_cache, - output=output, - record_indices=record_indices, - selected_counts=remapped_counts, - cu_seqlens_q=cu_seqlens_q, - kv_scale=kv_scale, - q_scale=q_scale, - active_splits=active_splits, - ) - remap = _paged_index_remap.bind_physical_slots( - physical_slots=selected_indices, - input_counts=selected_counts, - physical_indices=record_indices, - selected_counts=remapped_counts, - max_q_rows=int(plan.caps.max_q_rows), - num_cache_blocks=int(kv_cache.shape[0]), - block_stride_records=block_stride_records, - token_stride_records=token_stride_records, - ) - return Binding( - native=binding.native, - query_quant=binding.query_quant, - index_remap=remap, - kv_cache=binding.kv_cache, - selected_indices=binding.selected_indices, - selected_counts=binding.selected_counts, - ) - - -def bind_indexed( - plan: Plan, - *, - scratch, - q: torch.Tensor, - kv_cache: torch.Tensor, - output: torch.Tensor, - logical_indices: torch.Tensor, - request_ids: torch.Tensor, - block_table: torch.Tensor, - cu_seqlens_q: torch.Tensor, - kv_scale: torch.Tensor, - q_scale: torch.Tensor, - active_splits: int | None = None, -) -> Binding: - scratch_storage = scratch_tensor( - scratch, - plan.scratch_specs(), - owner="strided sparse MLA", - ) - rows = int(q.shape[0]) - physical_indices, selected_counts = _index_scratch(plan, scratch_storage, rows) - flat_cache, block_stride_records, token_stride_records = _physical_record_view( - plan, kv_cache - ) - binding = _bind_native( - plan, - scratch_storage=scratch_storage, - q=q, - flat_cache=flat_cache, - original_cache=kv_cache, - output=output, - record_indices=physical_indices, - selected_counts=selected_counts, - cu_seqlens_q=cu_seqlens_q, - kv_scale=kv_scale, - q_scale=q_scale, - active_splits=active_splits, - ) - remap = _paged_index_remap.bind( - request_ids=request_ids, - block_table=block_table, - logical_indices=logical_indices, - physical_indices=physical_indices, - selected_counts=selected_counts, - max_q_rows=int(plan.caps.max_q_rows), - num_cache_blocks=int(kv_cache.shape[0]), - block_stride_records=block_stride_records, - token_stride_records=token_stride_records, - ) - return Binding( - native=binding.native, - query_quant=binding.query_quant, - index_remap=remap, - kv_cache=binding.kv_cache, - selected_indices=binding.selected_indices, - selected_counts=binding.selected_counts, - ) - - -def compile(*, binding: Binding) -> None: - if binding.index_remap is not None: - _paged_index_remap.compile(binding=binding.index_remap) - static_fp8_quant.compile(binding=binding.query_quant) - compile_dense_mla(binding=binding.native) - - -def run_decode(*, binding: Binding) -> tuple[torch.Tensor, torch.Tensor]: - if binding.index_remap is not None: - _paged_index_remap.run(binding=binding.index_remap) - static_fp8_quant.run(binding=binding.query_quant) - return run_dense_mla(binding=binding.native) - - -def run_extend(*, binding: Binding) -> tuple[torch.Tensor, torch.Tensor]: - if binding.index_remap is not None: - _paged_index_remap.run(binding=binding.index_remap) - static_fp8_quant.run(binding=binding.query_quant) - return run_dense_mla(binding=binding.native) - - -def reference( - q: torch.Tensor, - kv_cache: torch.Tensor, - selected_indices: torch.Tensor, - selected_counts: torch.Tensor, - *, - kv_scale: torch.Tensor | float, - q_scale: torch.Tensor | float, -) -> tuple[torch.Tensor, torch.Tensor]: - def scalar(value: torch.Tensor | float) -> float: - if isinstance(value, torch.Tensor): - return float(value.detach().cpu().item()) - return float(value) - - count_host = [int(value) for value in selected_counts.detach().cpu().tolist()] - q_mul = scalar(q_scale) - q_f32 = (q.float() / q_mul).to(FP8).float() * q_mul - kv_mul = scalar(kv_scale) - output = torch.empty( - q.shape[0], q.shape[1], VALUE_DIM, dtype=torch.float32, device=q.device - ) - lse = torch.empty(q.shape[:2], dtype=torch.float32, device=q.device) - for row, count in enumerate(count_host): - ids = selected_indices[row, :count].to(torch.long) - blocks = torch.div(ids, BLOCK_SIZE, rounding_mode="floor") - tokens = ids.remainder(BLOCK_SIZE) - records = kv_cache[blocks, tokens, :QK_DIM].float() * kv_mul - logits = torch.einsum("hd,kd->hk", q_f32[row], records) * SM_SCALE - probability = torch.softmax(logits, dim=-1) - output[row] = torch.einsum("hk,kd->hd", probability, records[:, :VALUE_DIM]) - lse[row] = torch.logsumexp(logits, dim=-1) - return output.to(torch.bfloat16), lse - - -def is_supported(device=None) -> bool: - if not default_is_supported(device, requires=()): - return False - if device is None: - device = torch.device("cuda", torch.cuda.current_device()) - else: - device = torch.device(device) - return tuple(torch.cuda.get_device_capability(device)) in ((12, 0), (12, 1)) - - -def clear_caches() -> None: - _paged_index_remap.clear_caches() - static_fp8_quant.clear_caches() - clear_dense_mla_kernel_caches() - - -__all__ = [ - "Binding", - "Budget", - "Caps", - "Plan", - "Scratch", - "bind", - "bind_indexed", - "clear_caches", - "compile", - "is_supported", - "plan", - "reference", - "run_decode", - "run_extend", -] diff --git a/b12x/attention/varlen/__init__.py b/b12x/attention/varlen/__init__.py index d5c078f90..ea66aa030 100644 --- a/b12x/attention/varlen/__init__.py +++ b/b12x/attention/varlen/__init__.py @@ -1,8 +1,7 @@ """Contiguous (non-paged) attention for SM12x: batched and varlen forward. -BF16/FP16 Q/K/V, including different Q/K and V head dimensions, causal + -sliding-window + attention-sink; tile shapes auto-selected by Q/K head -dimension and causality. ``run`` is the varlen (cu_seqlens) +BF16/FP16 Q/K/V, causal + sliding-window + attention-sink; tile shapes +auto-selected by head_dim/causality. ``run`` is the varlen (cu_seqlens) entry, ``run_batched`` the fixed-shape batched entry; each has its own plan/scratch/binding family (``create_plan*`` for the per-shape kernel plan, ``plan*`` for scratch sizing). diff --git a/b12x/gemm/mxfp8_linear/_kernel.py b/b12x/gemm/mxfp8_linear/_kernel.py index 7489229f5..985d7574b 100644 --- a/b12x/gemm/mxfp8_linear/_kernel.py +++ b/b12x/gemm/mxfp8_linear/_kernel.py @@ -86,6 +86,12 @@ def _pad_source_2d_k(source_2d: torch.Tensor, padded_k: int) -> torch.Tensor: return padded.contiguous() +def _dense_gemm_kwargs_for_n(out_features: int) -> dict[str, object]: + if int(out_features) < 64: + return {"mma_tiler_mn": (64, 32), "swap_ab": True} + return {} + + def is_mxfp8_linear_supported() -> tuple[bool, str | None]: if not hasattr(cute.nvgpu.warp, "MmaMXF8Op"): return False, "CUTLASS DSL does not expose cute.nvgpu.warp.MmaMXF8Op" @@ -191,6 +197,7 @@ def _mxfp8_linear_fused_op( sf_vec_size=MXFP8_SCALE_VEC_SIZE, expected_m=expected_m, stream=stream_int, + **_dense_gemm_kwargs_for_n(out_features), )[:, :, 0] diff --git a/b12x/gemm/tensor_fp8_linear/_kernel.py b/b12x/gemm/tensor_fp8_linear/_kernel.py index 6405a7821..1b95e550c 100644 --- a/b12x/gemm/tensor_fp8_linear/_kernel.py +++ b/b12x/gemm/tensor_fp8_linear/_kernel.py @@ -96,6 +96,12 @@ def _activation_scale_mma( ) +def _dense_gemm_kwargs_for_n(out_features: int) -> dict[str, object]: + if int(out_features) < 64: + return {"mma_tiler_mn": (64, 32), "swap_ab": True} + return {} + + def is_tensor_fp8_linear_supported() -> tuple[bool, str | None]: if not hasattr(cute.nvgpu.warp, "MmaMXF8Op"): return False, "CUTLASS DSL does not expose cute.nvgpu.warp.MmaMXF8Op" @@ -189,6 +195,7 @@ def _tensor_fp8_linear_fused_op( expected_m=expected_m, stream=stream_int, plain_fp8=True, + **_dense_gemm_kwargs_for_n(out_features), )[:, :, 0] diff --git a/b12x/gemm/trellis_linear/__init__.py b/b12x/gemm/trellis_linear/__init__.py index 50f6db801..4091faed7 100644 --- a/b12x/gemm/trellis_linear/__init__.py +++ b/b12x/gemm/trellis_linear/__init__.py @@ -1,7 +1,6 @@ -"""Trellis-coded dense linear for SM12x. +"""Native EXL3 Trellis dense linear for SM12x. -The operation consumes the checkpoint-native ``trellis_t256`` payload -(MCG or SQG-XOR-Cheb-T12 codebooks) and +The operation consumes the checkpoint-native ``trellis3_t256`` payload and its two Hadamard sign vectors. Preparation validates and records zero-copy views; execution performs input rotation, a W4A16 or direct E4M3-W4A8 GEMM, and output rotation. @@ -29,9 +28,11 @@ ), dtypes=("bf16", "fp16"), recipes=( - "w4a16/trellis_mcg", - "w4a16/trellis_sqg_e4m3", - "w4a8/trellis_sqg_e4m3", + "w4a16/exl3_trellis_mcg", + "w4a16/exl3_trellis_mul1_e4m3", + "w4a16/exl3_trellis_sqg_cheb_e4m3", + "w4a8/exl3_trellis_mul1_e4m3", + "w4a8/exl3_trellis_sqg_cheb_e4m3", ), provenance=Provenance( repo="https://github.com/local-inference-lab/b12x", @@ -43,7 +44,7 @@ ), test_path="tests/gemm/test_trellis_linear.py", since="1.0.1", - notes="Trellis-coded dense W4A16 and direct E4M3-W4A8 linear.", + notes="Native EXL3 Trellis dense W4A16 and direct E4M3-W4A8 linear.", ) if TYPE_CHECKING: diff --git a/b12x/gemm/trellis_linear/_small_m.py b/b12x/gemm/trellis_linear/_small_m.py new file mode 100644 index 000000000..4b3810396 --- /dev/null +++ b/b12x/gemm/trellis_linear/_small_m.py @@ -0,0 +1,112 @@ +"""JIT binding for the B12X-owned K6/MCG small-M CUDA kernel.""" + +from __future__ import annotations + +import hashlib +import os +from functools import lru_cache +from pathlib import Path + +import torch +from torch.utils.cpp_extension import load + + +_SOURCE_DIR = Path(__file__).resolve().parent / "csrc" +_SOURCE = _SOURCE_DIR / "trellis_k6_small.cu" +_VENDORED_FILES = tuple(sorted((_SOURCE_DIR / "vendor").rglob("*.[ch]*"))) + + +_GLM_K6_DECODE_SMS = { + # Q/indexer projection on the target stream. + (2048, 4096): 128, + # TP4 shared-expert FC1/FC2 run beside the target stream. These are the + # rank-local dimensions after column/row parallel slicing, not the full + # 4096/2048-wide shared MLP dimensions. The budgets match the E2E-optimal + # ExLlama autotuner result; using all 188 SMs serializes the graph branches. + (6144, 1024): 64, + (512, 6144): 96, +} + + +def _default_num_sms(size_k: int, size_n: int, available_sms: int) -> int: + """Select the measured GLM K6 decode overlap budget when applicable.""" + target = _GLM_K6_DECODE_SMS.get((size_k, size_n)) + return available_sms if target is None else min(available_sms, target) + + +@lru_cache(maxsize=None) +def _available_sms(device_index: int) -> int: + return int(torch.cuda.get_device_properties(device_index).multi_processor_count) + + +def _extension_name() -> str: + digest = hashlib.sha256() + for path in (_SOURCE, *_VENDORED_FILES): + if path.is_file(): + digest.update(path.relative_to(_SOURCE_DIR).as_posix().encode()) + digest.update(path.read_bytes()) + return f"b12x_trellis_k6_{digest.hexdigest()[:12]}" + + +@lru_cache(maxsize=1) +def _extension(): + build_directory = os.environ.get("B12X_TRELLIS_BUILD_DIR") + if build_directory: + Path(build_directory).mkdir(parents=True, exist_ok=True) + return load( + name=_extension_name(), + sources=[str(_SOURCE)], + extra_include_paths=[str(_SOURCE_DIR)], + extra_cuda_cflags=[ + "-O3", + "--use_fast_math", + "--expt-relaxed-constexpr", + "--expt-extended-lambda", + "-gencode=arch=compute_120,code=sm_120", + ], + extra_cflags=["-O3"], + build_directory=build_directory, + verbose=os.environ.get("B12X_JIT_VERBOSE", "0") == "1", + ) + + +def run_k6_mcg( + x: torch.Tensor, + trellis: torch.Tensor, + output: torch.Tensor, + suh: torch.Tensor, + rotated_input: torch.Tensor, + svh: torch.Tensor, + locks: torch.Tensor, + *, + num_sms: int = 0, +) -> None: + """Launch the capture-safe K6/MCG kernel on Torch's current stream.""" + capability = torch.cuda.get_device_capability(x.device) + if capability != (12, 0): + raise NotImplementedError( + "Trellis K6 small-M kernel is built for sm_120 only; " + f"device reports sm_{capability[0]}{capability[1]}" + ) + if num_sms <= 0: + device_index = x.device.index + if device_index is None: + device_index = torch.cuda.current_device() + num_sms = _default_num_sms( + int(x.shape[1]), + int(output.shape[1]), + _available_sms(int(device_index)), + ) + _extension().launch_k6_mcg( + x, + trellis, + output, + suh, + rotated_input, + svh, + locks, + int(num_sms), + ) + + +__all__ = ["run_k6_mcg"] diff --git a/b12x/gemm/trellis_linear/api.py b/b12x/gemm/trellis_linear/api.py index d9f029476..4714453cd 100644 --- a/b12x/gemm/trellis_linear/api.py +++ b/b12x/gemm/trellis_linear/api.py @@ -86,7 +86,7 @@ def run( gemm_output_f16: Optional[torch.Tensor] = None, output_f16: Optional[torch.Tensor] = None, hadamard_128=None, - _moe_block_size: int = 64, + _moe_block_size: int | None = None, _force_tile_config: tuple[int, int] | None = None, ) -> torch.Tensor: """Execute Trellis GEMM, optionally reusing all capture-time storage.""" diff --git a/b12x/gemm/trellis_linear/csrc/trellis_k6_small.cu b/b12x/gemm/trellis_linear/csrc/trellis_k6_small.cu new file mode 100644 index 000000000..bb9223521 --- /dev/null +++ b/b12x/gemm/trellis_linear/csrc/trellis_k6_small.cu @@ -0,0 +1,152 @@ +// SPDX-License-Identifier: MIT +// +// Narrow K6/MCG dense launcher adapted from ExLlamaV3. The vendored kernel +// headers retain their MIT license in vendor/LICENSE.exllamav3. B12X +// owns the tensor validation, workspace, dispatch, and CUDA-stream contract. + +#include +#include +#include +#include +#include +#include +#include + +#include "vendor/util.h" +#include "vendor/util.cuh" + +namespace cg = cooperative_groups; + +#include "vendor/quant/exl3_gemm_kernel.cuh" + +namespace { + +constexpr int kBits = 6; +constexpr int kCodebookMcg = 1; +constexpr int kTileM = 16; +constexpr int kTileK = 32; +constexpr int kTileN = 128; +constexpr int kSharedStages = 4; +constexpr int kFragmentStages = 3; +constexpr int kThreads = EXL3_GEMM_BASE_THREADS * (kTileK / 16); +constexpr int kFragmentsNPerWarp = + 2 * (kTileN / 16) / (EXL3_GEMM_BASE_THREADS / 32); +constexpr int kASharedElements = kTileM * kTileK; +constexpr int kBSharedElements = + (kTileK / 16) * (kTileN / 16) * 16 * kBits; +constexpr int kCSharedElements = + 4 * EXL3_GEMM_BASE_THREADS * kFragmentsNPerWarp > kTileN * kTileM + ? 4 * EXL3_GEMM_BASE_THREADS * kFragmentsNPerWarp + : kTileN * kTileM; +constexpr int kDynamicSmem = + kSharedStages * (2 * kASharedElements + 2 * kBSharedElements) + + 4 * kCSharedElements; +static_assert(kDynamicSmem <= EXL3_SMEM_MAX_BYTES, + "Trellis K6 kernel exceeds the vendored shared-memory limit"); + +void check_cuda(cudaError_t status, const char* operation) { + TORCH_CHECK(status == cudaSuccess, operation, ": ", cudaGetErrorString(status)); +} + +void launch_k6_mcg( + const torch::Tensor& input, + const torch::Tensor& trellis, + torch::Tensor& output, + const torch::Tensor& suh, + torch::Tensor& rotated_input, + const torch::Tensor& svh, + torch::Tensor& locks, + int64_t requested_sms) { + TORCH_CHECK(input.is_cuda() && trellis.is_cuda() && output.is_cuda() && + suh.is_cuda() && rotated_input.is_cuda() && svh.is_cuda() && + locks.is_cuda(), + "Trellis K6 tensors must be CUDA tensors"); + const int device = input.get_device(); + TORCH_CHECK(trellis.get_device() == device && output.get_device() == device && + suh.get_device() == device && + rotated_input.get_device() == device && + svh.get_device() == device && locks.get_device() == device, + "Trellis K6 tensors must share one CUDA device"); + const c10::cuda::CUDAGuard device_guard(input.device()); + TORCH_CHECK(input.scalar_type() == at::kHalf && output.scalar_type() == at::kHalf, + "Trellis K6 small-M path requires FP16 input and output"); + TORCH_CHECK(trellis.scalar_type() == at::kShort, + "Trellis K6 payload must be viewed as int16"); + TORCH_CHECK(suh.scalar_type() == at::kHalf && svh.scalar_type() == at::kHalf, + "Trellis K6 rotation scales must be FP16"); + TORCH_CHECK(rotated_input.scalar_type() == at::kHalf, + "Trellis K6 rotated input must be FP16"); + TORCH_CHECK(locks.scalar_type() == at::kInt, + "Trellis K6 lock workspace must be int32"); + TORCH_CHECK(input.is_contiguous() && trellis.is_contiguous() && + output.is_contiguous() && suh.is_contiguous() && + rotated_input.is_contiguous() && svh.is_contiguous() && + locks.is_contiguous(), + "Trellis K6 small-M tensors must be contiguous"); + TORCH_CHECK(input.dim() == 2 && output.dim() == 2 && trellis.dim() == 3, + "Trellis K6 expects A[M,K], B[K/16,N/16,96], C[M,N]"); + + int size_m = static_cast(input.size(0)); + int size_k = static_cast(input.size(1)); + int size_n = static_cast(output.size(1)); + TORCH_CHECK(size_m >= 1 && size_m <= 128, + "Trellis K6 small-M path supports 1..128 rows, got ", size_m); + TORCH_CHECK(output.size(0) == size_m && rotated_input.sizes() == input.sizes(), + "Trellis K6 output/scratch shapes do not match input"); + TORCH_CHECK(trellis.size(0) * 16 == size_k && + trellis.size(1) * 16 == size_n && trellis.size(2) == 96, + "Trellis K6 native payload shape does not match M/K/N"); + TORCH_CHECK(suh.numel() == size_k && svh.numel() == size_n, + "Trellis K6 rotation scale width mismatch"); + // The fused H128 preamble reads 128 contiguous scale elements per warp. + TORCH_CHECK(size_k % 128 == 0 && size_n % kTileN == 0, + "Trellis K6 K/N must be divisible by 128/128, got K=", size_k, + " N=", size_n); + TORCH_CHECK(locks.numel() >= size_n / 16, + "Trellis K6 lock workspace is too small"); + + int available_sms = 0; + check_cuda(cudaDeviceGetAttribute(&available_sms, cudaDevAttrMultiProcessorCount, + device), + "cudaDeviceGetAttribute"); + const int tiles = (size_k / kTileK) * (size_n / kTileN); + int num_sms = requested_sms > 0 ? static_cast(requested_sms) : available_sms; + num_sms = std::max(1, std::min(num_sms, std::min(available_sms, tiles))); + + auto kernel = exl3_gemm_kernel; + static std::once_flag smem_attribute_once; + std::call_once(smem_attribute_once, [&] { + check_cuda(cudaFuncSetAttribute(kernel, + cudaFuncAttributeMaxDynamicSharedMemorySize, + kDynamicSmem), + "cudaFuncSetAttribute"); + }); + + const half* a_ptr = reinterpret_cast(input.data_ptr()); + const uint16_t* b_ptr = + reinterpret_cast(trellis.data_ptr()); + half* c_ptr = reinterpret_cast(output.data_ptr()); + int* lock_ptr = locks.data_ptr(); + const half* suh_ptr = reinterpret_cast(suh.data_ptr()); + half* rotated_ptr = + reinterpret_cast(rotated_input.data_ptr()); + const half* svh_ptr = reinterpret_cast(svh.data_ptr()); + void* args[] = {&a_ptr, &b_ptr, &c_ptr, &size_m, &size_k, + &size_n, &lock_ptr, &suh_ptr, &rotated_ptr, &svh_ptr}; + + cudaStream_t stream = at::cuda::getCurrentCUDAStream(device).stream(); + // Keep this launch one-dimensional: vendored H128 scale indexing assumes + // gridDim.y == 1 because the caller already supplies its per-K scale offset. + check_cuda(cudaLaunchCooperativeKernel(reinterpret_cast(kernel), num_sms, + kThreads, args, kDynamicSmem, stream), + "cudaLaunchCooperativeKernel"); + check_cuda(cudaPeekAtLastError(), "Trellis K6 kernel launch"); +} + +} // namespace + +PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) { + module.def("launch_k6_mcg", &launch_k6_mcg, + "B12X K6/MCG small-M dense GEMM"); +} diff --git a/b12x/gemm/trellis_linear/csrc/vendor/LICENSE.exllamav3 b/b12x/gemm/trellis_linear/csrc/vendor/LICENSE.exllamav3 new file mode 100644 index 000000000..b40e3294c --- /dev/null +++ b/b12x/gemm/trellis_linear/csrc/vendor/LICENSE.exllamav3 @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2025 Turboderp + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. \ No newline at end of file diff --git a/b12x/gemm/trellis_linear/csrc/vendor/compat.cuh b/b12x/gemm/trellis_linear/csrc/vendor/compat.cuh new file mode 100644 index 000000000..42aa0a9e0 --- /dev/null +++ b/b12x/gemm/trellis_linear/csrc/vendor/compat.cuh @@ -0,0 +1,31 @@ +#pragma once + +// Approximate tanh + +__forceinline__ __device__ float copysignf_pos(float a, float b) +{ + float r; + r = __int_as_float(__float_as_int(a) | (__float_as_int(b) & 0x80000000)); + return r; +} + +#if defined(USE_ROCM) || (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ < 750 || CUDART_VERSION < 11000)) + +__inline__ __device__ float tanh_opt(float x) +{ + const float exp_val = -1.f * fabs(2 * x); + return copysignf_pos((1.0f - __expf(exp_val)) / (__expf(exp_val) + 1.0f), x); +} + +#else + +__inline__ __device__ float tanh_opt(float x) +{ + float r; + asm("tanh.approx.f32 %0,%1; \n\t" : "=f"(r) : "f"(x)); + return r; +} + +#endif +// Vendored from https://github.com/brandonmmusic-max/exllamav3 at 704aefd743b390af4bd0fb429d1906f9b964c7d8. +// License: b12x/gemm/trellis_linear/csrc/vendor/LICENSE.exllamav3 diff --git a/b12x/gemm/trellis_linear/csrc/vendor/ptx.cuh b/b12x/gemm/trellis_linear/csrc/vendor/ptx.cuh new file mode 100644 index 000000000..1d3095b13 --- /dev/null +++ b/b12x/gemm/trellis_linear/csrc/vendor/ptx.cuh @@ -0,0 +1,350 @@ +#pragma once +#include + +// Tensor core fragments + +template +struct Vec +{ + T elems[n]; + __device__ T& operator[](int i) { return elems[i]; } +}; + +using FragA = Vec; +using FragB = Vec; +using FragC = Vec; +using FragC_h = Vec; + +// m8n8k4 tensor core matmul (emulated on Ampere and later), don't use +// +// https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#matrix-fragments-for-mma-m8n8k4-with-f16-floating-point-type + +__device__ inline void ptx_mma_m8n8k4 +( + const Vec& frag_a, + const Vec& frag_b, + Vec& frag_c +) +{ + const uint32_t* a = reinterpret_cast(&frag_a); + const uint32_t* b = reinterpret_cast(&frag_b); + float* c = reinterpret_cast(&frag_c); + const float* d = reinterpret_cast(&frag_c); + + asm + ( + "mma.sync.aligned.m8n8k4.row.col.f32.f16.f16.f32 " + "{%0,%1,%2,%3,%4,%5,%6,%7}, {%8,%9}, {%10,%11}, {%12,%13,%14,%15,%16,%17,%18,%19};\n" + + : "=f"(c[0]), "=f"(c[1]), "=f"(c[2]), "=f"(c[3]),"=f"(c[4]), "=f"(c[5]), "=f"(c[6]), "=f"(c[7]) + + : "r"(a[0]), "r"(a[1]), + "r"(b[0]), "r"(b[1]), + "f"(d[0]), "f"(d[1]), "f"(d[2]), "f"(d[3]), "f"(d[4]), "f"(d[5]), "f"(d[6]), "f"(d[7]) + ); +} + +// m16n8k16 tensor core matmul +// +// https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#matrix-fragments-for-mma-m16n8k16-with-floating-point-type + +// FP16 @ FP16 + FP32 -> FP32 +__device__ inline void ptx_mma_m16n8k16 +( + const FragA& frag_a, + const FragB& frag_b, + FragC& frag_c +) +{ + const uint32_t* a = reinterpret_cast(&frag_a); + const uint32_t* b = reinterpret_cast(&frag_b); + float* c = reinterpret_cast(&frag_c); + const float* d = reinterpret_cast(&frag_c); + + asm + ( + "mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32 " + "{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%10,%11,%12,%13};\n" + + : "=f"(c[0]), "=f"(c[1]), "=f"(c[2]), "=f"(c[3]) + : "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), + "r"(b[0]), "r"(b[1]), + "f"(d[0]), "f"(d[1]), "f"(d[2]), "f"(d[3]) + ); +} + +// FP16 @ FP16 + FP16 -> FP16 +__device__ inline void ptx_mma_m16n8k16 +( + const FragA& frag_a, + const FragB& frag_b, + FragC_h& frag_c +) +{ + const uint32_t* a = reinterpret_cast(&frag_a); + const uint32_t* b = reinterpret_cast(&frag_b); + uint32_t* c = reinterpret_cast(&frag_c); + const uint32_t* d = reinterpret_cast(&frag_c); + + asm + ( + "mma.sync.aligned.m16n8k16.row.col.f16.f16.f16.f16 " + "{%0,%1}, {%2,%3,%4,%5}, {%6,%7}, {%8,%9};\n" + + : "=r"(c[0]), "=r"(c[1]) + : "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), + "r"(b[0]), "r"(b[1]), + "r"(d[0]), "r"(d[1]) + ); +} + +// Global barrier + +__device__ inline void barrier_acquire +( + int* lock, + int stage +) +{ + if (threadIdx.x == 0) + { + volatile int state = -1; + do + { + asm volatile ("ld.global.acquire.gpu.b32 %0, [%1];\n" : "=r"(state) : "l"(lock)); + } + while (state != stage); + } + __syncthreads(); +} + +__device__ inline void barrier_release +( + int* lock, + int val, + bool reset +) +{ + __syncthreads(); + if (threadIdx.x == 0) + { + if (reset) + { + *lock = 0; + return; + } + asm volatile ("fence.acq_rel.gpu;\n"); + asm volatile ("red.relaxed.gpu.global.add.s32 [%0], %1;\n" : : "l"(lock), "r"(val)); + } +} + +// Load global to shared memory, predicated. Seems to produce incorrect code when compiling for Blackwell, but +// `if (...) cp_async(...)` compiles to a predicated instruction anyway + +__device__ inline void cp_async_pred(void* smem_ptr, const void* glob_ptr, bool pred = true) +{ + const int bytes = 16; + uint32_t smem = static_cast(__cvta_generic_to_shared(smem_ptr)); + asm volatile( + "{\n" + " .reg .pred p;\n" + " setp.ne.b32 p, %0, 0;\n" + " @p cp.async.cg.shared.global [%1], [%2], %3;\n" + "}\n" :: "r"((int) pred), "r"(smem), "l"(glob_ptr), "n"(bytes) + ); +} + +// Load global to shared memory + +__device__ inline void cp_async(void* smem_ptr, const void* glob_ptr) +{ + const int bytes = 16; + uint32_t smem = static_cast(__cvta_generic_to_shared(smem_ptr)); + asm volatile( + "{\n" + " cp.async.cg.shared.global [%0], [%1], %2;\n" + "}\n" :: "r"(smem), "l"(glob_ptr), "n"(bytes) + ); +} + +// Load global to shared memory with cache hint to evict data from L2 ASAP + +__device__ inline void cp_async_stream(void* smem_ptr, const void* glob_ptr) +{ + uint32_t smem = static_cast(__cvta_generic_to_shared(smem_ptr)); + const int bytes = 16; + asm volatile + ( + "{\n" + " .reg .b64 p;\n" + " createpolicy.fractional.L2::evict_first.b64 p, 1.0;\n" + " cp.async.cg.shared.global.L2::cache_hint [%0], [%1], %2, p;\n" + "}\n" :: "r"(smem), "l"(glob_ptr), "n"(bytes) + ); +} + +// Async copy fence, commit all pending async copies + +__device__ inline void cp_async_fence() +{ + asm volatile("cp.async.commit_group;\n" ::); +} + +// Wait until at most n async groups are still pending. + +template +__device__ inline void cp_async_wait() +{ + asm volatile("cp.async.wait_group %0;\n" :: "n"(n)); +} + +// Load 16x16 matrix fragment from shared memory, directly in tensor core layout + +__device__ inline void ldsm4(FragA& frag_a, const void* smem_ptr) +{ + uint32_t* a = reinterpret_cast(&frag_a); + uint32_t smem = static_cast(__cvta_generic_to_shared(smem_ptr)); + asm volatile + ( + "ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];\n" + : "=r"(a[0]), "=r"(a[1]), "=r"(a[2]), "=r"(a[3]) : "r"(smem) + ); +} + +__device__ inline uint32_t mul_lo_u32(uint32_t x, uint32_t y) +{ + uint32_t w; + asm volatile + ( + "mul.lo.u32 %0, %1, %2;" + : "=r"(w) + : "r"(x), "r"(y) + ); + return w; +} + +__device__ inline uint32_t mul_hi_u32(uint32_t x, uint32_t y) +{ + uint32_t w; + asm volatile + ( + "mul.hi.u32 %0, %1, %2;" + : "=r"(w) + : "r"(x), "r"(y) + ); + return w; +} + +// Memory ops + +__device__ __forceinline__ void stg_wt_u32(uint32_t* p, uint32_t v) +{ + asm volatile("st.global.wt.u32 [%0], %1;" :: "l"(p), "r"(v)); +} + +__device__ __forceinline__ void stg_wt_u128(uint4* p, const uint4 v) +{ + asm volatile ("st.global.wt.v4.u32 [%0], {%1,%2,%3,%4};" + :: "l"(p), + "r"(v.x), "r"(v.y), "r"(v.z), "r"(v.w)); +} + +__device__ __forceinline__ uint32_t ldg_cv_u32(const uint32_t* p) +{ + uint32_t v; + asm volatile("ld.global.cv.u32 %0, [%1];" : "=r"(v) : "l"(p)); + return v; +} + +__device__ __forceinline__ uint4 ldg_cv_u128(const uint4* p) +{ + uint4 v; + asm volatile ("ld.global.cv.v4.u32 {%0,%1,%2,%3}, [%4];" + : "=r"(v.x), "=r"(v.y), "=r"(v.z), "=r"(v.w) + : "l"(p)); + return v; +} + +__device__ __forceinline__ uint32_t ldg_acquire_sys_u32(const uint32_t* p) +{ + uint32_t v; + asm volatile("ld.global.acquire.sys.u32 %0, [%1];" + : "=r"(v) : "l"(p)); + return v; +} + +__device__ __forceinline__ uint64_t ldg_acquire_sys_u64(const uint64_t* p) +{ + uint64_t v; + asm volatile("ld.global.acquire.sys.u64 %0, [%1];" : "=l"(v) : "l"(p) : "memory"); + return v; +} + +__device__ __forceinline__ void stg_release_sys_u32(uint32_t* p, uint32_t v) +{ + asm volatile("st.global.release.sys.u32 [%0], %1;" :: "l"(p), "r"(v) : "memory"); +} + +__device__ __forceinline__ void stg_release_sys_u64(uint64_t* p, uint64_t v) +{ + asm volatile("st.global.release.sys.u64 [%0], %1;" :: "l"(p), "l"(v) : "memory"); +} + +// Global time in nanoseconds + +__device__ __forceinline__ uint64_t globaltimer_ns() +{ + uint64_t t; + asm volatile("mov.u64 %0, %%globaltimer;" : "=l"(t)); + return t; +} + +// Bitfield stuff + +static __forceinline__ __device__ uint32_t bfe64(uint32_t lo, uint32_t hi, int offset, int length) +{ + uint64_t value = (static_cast(hi) << 32) | static_cast(lo); + uint64_t result64; + asm ("bfe.u64 %0, %1, %2, %3;" + : "=l"(result64) + : "l"(value), "r"(offset), "r"(length)); + return static_cast(result64); +} + +#define FSHF_IMM(dst, lo, hi, imm) asm("shf.r.wrap.b32 %0, %1, %2, " #imm ";" : "=r"(dst) : "r"(lo), "r"(hi)) +#define BFE16_IMM(dst, src, imm) asm("bfe.u32 %0, %1, " #imm ", 16;" : "=r"(dst) : "r"(src)) + +// Inter-block barrier + +__device__ inline void group_barrier +( + int group_id, + int group_size, + int* barrier_counters_sense // length 2*max(group_id). odd positions are flipped after sync (sense) +) +{ + __syncthreads(); + + if (threadIdx.x == 0) + { + cuda::atomic_ref counter(barrier_counters_sense[group_id * 2]); + cuda::atomic_ref sense(barrier_counters_sense[group_id * 2 + 1]); + + int old_sense = sense.load(cuda::memory_order_relaxed); + int old = counter.fetch_add(1, cuda::memory_order_acq_rel); + + if (old == group_size - 1) + { + counter.store(0, cuda::memory_order_relaxed); + sense.store(1 - old_sense, cuda::memory_order_release); + } + else + { + while (sense.load(cuda::memory_order_acquire) == old_sense) __nanosleep(32); + } + } + + __syncthreads(); +} +// Vendored from https://github.com/brandonmmusic-max/exllamav3 at 704aefd743b390af4bd0fb429d1906f9b964c7d8. +// License: b12x/gemm/trellis_linear/csrc/vendor/LICENSE.exllamav3 diff --git a/b12x/gemm/trellis_linear/csrc/vendor/quant/codebook.cuh b/b12x/gemm/trellis_linear/csrc/vendor/quant/codebook.cuh new file mode 100644 index 000000000..e3f6bac25 --- /dev/null +++ b/b12x/gemm/trellis_linear/csrc/vendor/quant/codebook.cuh @@ -0,0 +1,159 @@ +#pragma once + +// Force integer MAD on sm<=86. For some reason this performs better than letting the compiler emit IMUL +// TODO: Keep an eye on new behavior in future versions of nvcc. While this is faster on RTX 3090, it really shouldn't be. +template +__device__ __forceinline__ +uint32_t mul_const_u32(uint32_t x) +{ + #if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ == 860) + uint32_t r; + asm volatile ( + "{ .reg .u32 z,t;" + " mov.u32 t, %laneid;" // runtime SR + " sub.u32 z, t, t;" // z = 0 but data-dependent + " mad.lo.u32 %0, %1, %2, z;" + "}" + : "=r"(r) + : "r"(x), "n"(w)); + return r; + #else + return x * w; + #endif +} + +template +__device__ inline half decode_3inst(uint32_t x) +{ + if constexpr (cb == 0) + { + x *= 89226354u; + x += 64248484u; + asm ("lop3.b32 %0, %0, 0x8fff8fff, 0x3b603b60, 0x6a;" : "+r"(x)); + half2_uint32 xu(x); + return __hadd(__low2half(xu.as_half2), __high2half(xu.as_half2)); + } + if constexpr (cb == 1) + { +// x *= 0xCBAC1FEDu; + x = mul_const_u32<0xCBAC1FEDu>(x); + + asm ("lop3.b32 %0, %0, 0x8fff8fff, 0x3b603b60, 0x6a;" : "+r"(x)); + half2_uint32 xu(x); + return __hadd(__low2half(xu.as_half2), __high2half(xu.as_half2)); + } + if constexpr (cb == 2) + { + x *= 0x83DCD12Du; + uint32_t sum; + const uint32_t acc = 0x6400u; // 0x6400 -> 1024.0 .. 0x67FF -> 2047.0 + asm ("vabsdiff4.u32.u32.u32.add %0, %1, %2, %3;" : "=r"(sum) : "r"(x), "r"(0), "r"(acc) : ); + const __half k_inv_h = __ushort_as_half(0x1eee); // 0.00677 = 1/147.7 + const __half k_bias_h = __ushort_as_half(0xc931); // -10.39 = (-1024.0 - 510.0) * k_inv_h + half_uint16 h((uint16_t) sum); + return __hfma(h.as_half, k_inv_h, k_bias_h); + } +} + +template +__device__ inline half2 decode_3inst_2(uint32_t x0, uint32_t x1) +{ + if constexpr (cb == 0) + { + x0 *= 89226354u; + x1 *= 89226354u; + x0 += 64248484u; + x1 += 64248484u; + asm ("lop3.b32 %0, %0, 0x8fff8fff, 0x3b603b60, 0x6a;" : "+r"(x0)); + asm ("lop3.b32 %0, %0, 0x8fff8fff, 0x3b603b60, 0x6a;" : "+r"(x1)); + half2_uint32 xu0(x0); + half2_uint32 xu1(x1); + half2 d0 = __lows2half2(xu0.as_half2, xu1.as_half2); + half2 d1 = __highs2half2(xu0.as_half2, xu1.as_half2); + return __hadd2(d0, d1); + } + if constexpr (cb == 1) + { +// x0 *= 0xCBAC1FEDu; +// x1 *= 0xCBAC1FEDu; + x0 = mul_const_u32<0xCBAC1FEDu>(x0); + x1 = mul_const_u32<0xCBAC1FEDu>(x1); + asm ("lop3.b32 %0, %0, 0x8fff8fff, 0x3b603b60, 0x6a;" : "+r"(x0)); + asm ("lop3.b32 %0, %0, 0x8fff8fff, 0x3b603b60, 0x6a;" : "+r"(x1)); + half2_uint32 xu0(x0); + half2_uint32 xu1(x1); + half2 d0 = __lows2half2(xu0.as_half2, xu1.as_half2); + half2 d1 = __highs2half2(xu0.as_half2, xu1.as_half2); + return __hadd2(d0, d1); + } + if constexpr (cb == 2) + { + x0 *= 0x83DCD12Du; + x1 *= 0x83DCD12Du; + uint32_t sum0; + uint32_t sum1; + const uint32_t acc = 0x6400u; // 0x6400 -> 1024.0 .. 0x67FF -> 2047.0 + asm ("vabsdiff4.u32.u32.u32.add %0, %1, %2, %3;" : "=r"(sum0) : "r"(x0), "r"(0), "r"(acc) : ); + asm ("vabsdiff4.u32.u32.u32.add %0, %1, %2, %3;" : "=r"(sum1) : "r"(x1), "r"(0), "r"(acc) : ); + half2 k_inv_h2 = __half2half2(__ushort_as_half(0x1eee)); // 0.00677 = 1/147.7 + half2 k_bias_h2 = __half2half2(__ushort_as_half(0xc931)); // -10.39 = (-1024.0 - 510.0) * k_inv_h + half_uint16 h0((uint16_t) sum0); + half_uint16 h1((uint16_t) sum1); + return __hfma2(__halves2half2(h0.as_half, h1.as_half), k_inv_h2, k_bias_h2); + } +} + +__device__ inline half2 decode_mcg_product_2(uint32_t x0, uint32_t x1) +{ + asm ("lop3.b32 %0, %0, 0x8fff8fff, 0x3b603b60, 0x6a;" : "+r"(x0)); + asm ("lop3.b32 %0, %0, 0x8fff8fff, 0x3b603b60, 0x6a;" : "+r"(x1)); + half2_uint32 xu0(x0); + half2_uint32 xu1(x1); + half2 d0 = __lows2half2(xu0.as_half2, xu1.as_half2); + half2 d1 = __highs2half2(xu0.as_half2, xu1.as_half2); + return __hadd2(d0, d1); +} + +template +__device__ inline float decode_3inst_f(uint64_t x) +{ + return __half2float(decode_3inst(x)); +} + +template +__device__ inline float decode_3inst_f_diff(uint64_t x, float d) +{ + return __half2float(decode_3inst(x)) - d; +} + +// "2MAD" procedural codebook, much more overhead than 3INST, slightly better distribution at 2bpw +// Not used currently + +//__device__ inline half decode_2mad(uint64_t x) +//{ +// x = x * 264435761u + 1013904223u; +// x = ((x * 1664525u) >> 32) + x; +// int32_t c = (int32_t) __dp4a((uint32_t) x, 0x01010101u, 0xFFFFFE02u); +// half y = __hmul(__int2half_rn(c), __float2half_rn(0.008415)); +// return y; +//} +// +//__device__ inline float decode_2mad_f(uint64_t x) +//{ +// x = x * 264435761u + 1013904223u; +// x = ((x * 1664525u) >> 32) + x; +// int32_t c = (int32_t) __dp4a((uint32_t) x, 0x01010101u, 0xFFFFFE02u); +// float y = __int2float_rn(c) * 0.008415f; +// return y; +//} +// +//__device__ inline float decode_2mad_f_diff(uint64_t x, float d) +//{ +// x = x * 264435761u + 1013904223u; +// x = ((x * 1664525u) >> 32) + x; +// int32_t c = (int32_t) __dp4a((uint32_t) x, 0x01010101u, 0xFFFFFE02u); +// float y = fma(__int2float_rn(c), 0.008415f, -d); +// return y; +//} +// Vendored from https://github.com/brandonmmusic-max/exllamav3 at 704aefd743b390af4bd0fb429d1906f9b964c7d8. +// License: b12x/gemm/trellis_linear/csrc/vendor/LICENSE.exllamav3 diff --git a/b12x/gemm/trellis_linear/csrc/vendor/quant/exl3_devctx.cuh b/b12x/gemm/trellis_linear/csrc/vendor/quant/exl3_devctx.cuh new file mode 100644 index 000000000..281d9e694 --- /dev/null +++ b/b12x/gemm/trellis_linear/csrc/vendor/quant/exl3_devctx.cuh @@ -0,0 +1,57 @@ +#pragma once + +#include +#include + +// Max allowable output size, in tiles. Used to allocate global lock buffer per device for sync across threadblocks +#define MAX_TILES_C (1024 * 1024) +#define MAX_BARRIERS 1024 +#define BARRIER_LOCKS_OFFSET MAX_TILES_C + +// MoE expert scheduler state, after the barrier counters: [0] next ticket, +// [1] retired groups, [2 + group] ticket published to group. The state is +// self-resetting and is zero-initialized with the rest of the lock buffer. +#define MOE_MAX_GROUPS 64 +#define MOE_SCHED_OFFSET (MAX_TILES_C + 2 * MAX_BARRIERS) +#define MOE_SCHED_INTS (2 + MOE_MAX_GROUPS) + +// Workspace size +#define WORKSPACE_SIZE (16*1024*1024) + +// Treat hopper and blackwell as same arch for now +#define MAX_DEVICES 16 +#define CC_OLD 1 +#define CC_AMPERE 2 +#define CC_ADA 3 +#define CC_HOPPER 4 +#define CC_BLACKWELL 4 + +// Singleton to manage context for each device. Stores device attributes and a large-enough lock buffer per device +class DevCtx +{ +private: + int num_sms[MAX_DEVICES] = {}; + int cc[MAX_DEVICES] = {}; + void* locks[MAX_DEVICES] = {}; + void* ws[MAX_DEVICES] = {}; + std::mutex mtx; + +public: + static DevCtx& instance(); + int get_num_sms(int device); + int get_cc(int device); + void* get_ws(int device); + int* get_locks(int device); + +private: + DevCtx() = default; + DevCtx(const DevCtx&) = delete; + DevCtx& operator=(const DevCtx&) = delete; +}; + +int g_get_cc(int device); +int g_get_num_sms(int device); + +void prepare_ctx(int device); +// Vendored from https://github.com/brandonmmusic-max/exllamav3 at 704aefd743b390af4bd0fb429d1906f9b964c7d8. +// License: b12x/gemm/trellis_linear/csrc/vendor/LICENSE.exllamav3 diff --git a/b12x/gemm/trellis_linear/csrc/vendor/quant/exl3_dq.cuh b/b12x/gemm/trellis_linear/csrc/vendor/quant/exl3_dq.cuh new file mode 100644 index 000000000..caf4f60d8 --- /dev/null +++ b/b12x/gemm/trellis_linear/csrc/vendor/quant/exl3_dq.cuh @@ -0,0 +1,296 @@ +#pragma once + +#include "codebook.cuh" + +__device__ __forceinline__ uint32_t fshift(const uint32_t b, const uint32_t a, int shift) +{ + uint64_t merged = ((uint64_t)a << 32) | (uint64_t) b; + return (uint32_t)(merged >> shift); + + // Conditional funnel shift is somehow no longer faster + // if (shift < 32) return __funnelshift_r(b, a, shift); + // return a >> (shift - 32); +} + +template +__device__ __forceinline__ half dq(const uint32_t* ptr, int t_offset) +{ + int b0 = t_offset * bits + bits - 16 + 256 * bits; // bit index, start of word0 + int b1 = b0 + 16; // bit index, end of word0 + int i0 = b0 / 32; // uint32 containing first bit of word0 + int i1 = (b1 - 1) / 32; // uint32 containing last bit of word0, may be == i0 + int s0 = (i1 + 1) * 32 - b1; // shift value to align word1 to 32-bit boundary + + // Load 32 or 64 bits containing word0 + uint32_t a = ptr[i0 % (bits * 256 / 32)]; + uint32_t b = ptr[i1 % (bits * 256 / 32)]; + + // Shift into place + uint32_t w0 = __funnelshift_r(b, a, s0) & 0xffff; + return decode_3inst(w0); +} + +template +__device__ __forceinline__ half2 dq2(const uint32_t* ptr, int t_offset) +{ + int b0 = t_offset * bits + bits - 16 + 256 * bits; // bit index, start of word0 + int b1 = b0 + 16; // bit index, end of word0 + int i0 = b0 / 32; // uint32 containing first bit of word0 + int i1 = (b1 - 1) / 32; // uint32 containing last bit of word0, may be == i0 + int s0 = (i1 + 1) * 32 - b1; // shift value to align word1 to 32-bit boundary + + // Load 32 or 64 bits containing word0 + uint32_t a = ptr[i0 % (bits * 256 / 32)]; + uint32_t b = ptr[i1 % (bits * 256 / 32)]; + + // Shift into place + uint32_t w1 = __funnelshift_r(b, a, s0) & 0xffff; + uint32_t w0 = __funnelshift_r(b, a, s0 + bits) & 0xffff; + return decode_3inst_2(w0, w1); +} + +template +__device__ __forceinline__ void dq4(const uint32_t* ptr, int t_offset, FragB& frag) +{ + int b0 = (t_offset + 257) * bits - 16; // start of first word + int b1 = b0 + 3 * bits; // start of last word + int b2 = b1 + 16; // end of last word + int i0 = b0 / 32; // uint32 containing first bit of first word + int i2 = (b2 - 1) / 32; // uint32 containing last bit of last word, may be == i0 + int s2 = (i2 + 1) * 32 - b2; // shift value to align last word to 32-bit boundary + + uint32_t a = ptr[i0 % (bits * 256 / 32)]; + uint32_t b = ptr[i2 % (bits * 256 / 32)]; + uint32_t w3 = fshift(b, a, s2) & 0xffff; + uint32_t w2 = fshift(b, a, s2 + bits) & 0xffff; + uint32_t w1 = fshift(b, a, s2 + bits * 2) & 0xffff; + uint32_t w0 = fshift(b, a, s2 + bits * 3) & 0xffff; + half2 d0d1 = decode_3inst_2(w0, w1); + half2 d2d3 = decode_3inst_2(w2, w3); + frag[0] = d0d1; + frag[1] = d2d3; +} + +template +__device__ __forceinline__ void dq2x2(const uint32_t* ptr, int t_offset, FragB& frag) +{ + #pragma unroll + for (int i = 0; i < 2; ++i) + { + int b0 = (t_offset + 2 * i + 257) * bits - 16; // start of first word + int b1 = b0 + 1 * bits; // start of last word + int b2 = b1 + 16; // end of last word + int i0 = b0 / 32; // uint32 containing first bit of first word + int i2 = (b2 - 1) / 32; // uint32 containing last bit of last word, may be == i0 + int s2 = (i2 + 1) * 32 - b2; // shift value to align last word to 32-bit boundary + + uint32_t a = ptr[i0 % (bits * 256 / 32)]; + uint32_t b = ptr[i2 % (bits * 256 / 32)]; + uint32_t w1 = fshift(b, a, s2) & 0xffff; + uint32_t w0 = fshift(b, a, s2 + bits) & 0xffff; + half2 d0d1 = decode_3inst_2(w0, w1); + frag[i] = d0d1; + } +} + +template +__device__ __forceinline__ void dq8(const uint32_t* ptr, int t_offset, FragB& frag0, FragB& frag1) +{ + int b1 = (t_offset + 257) * bits; // end of first word + int b0 = b1 - 16; // start of first word + int b2 = b1 + bits * 7; + int i0 = b0 / 32; // uint32 containing first bit of word0 + int i2 = (b2 - 1) / 32; // uint32 containing last bit of word0, may be == i0 + int s2 = (i2 + 1) * 32 - b2; // shift value to align last word to 32-bit boundary + + uint32_t a = ptr[i0 % (bits * 256 / 32)]; + uint32_t b = ptr[i2 % (bits * 256 / 32)]; + uint32_t w0, w1, w2, w3, w4, w5, w6, w7; + if constexpr (align == 1) + { + w7 = fshift(b, a, s2); + w6 = fshift(b, a, s2 + bits); + w5 = fshift(b, a, s2 + bits * 2); + w4 = fshift(b, a, s2 + bits * 3); + w3 = fshift(b, a, s2 + bits * 4); + w2 = fshift(b, a, s2 + bits * 5); + w1 = fshift(b, a, s2 + bits * 6); + w0 = fshift(b, a, s2 + bits * 7); + } + if constexpr (align == 2) + { + w7 = fshift(b, a, s2); + w6 = w7 >> bits; + w5 = fshift(b, a, s2 + bits * 2); + w4 = w5 >> bits; + w3 = fshift(b, a, s2 + bits * 4); + w2 = w3 >> bits; + w1 = fshift(b, a, s2 + bits * 6); + w0 = w1 >> bits; + } + if constexpr (align == 4) + { + w7 = fshift(b, a, s2); + w6 = w7 >> bits; + w5 = w6 >> bits; + w4 = w5 >> bits; + w3 = fshift(b, a, s2 + bits * 4); + w2 = w3 >> bits; + w1 = w2 >> bits; + w0 = w1 >> bits; + } + if constexpr (align == 8) + { + w7 = fshift(b, a, s2); + w6 = w7 >> bits; + w5 = w6 >> bits; + w4 = w5 >> bits; + w3 = w4 >> bits; + w2 = w3 >> bits; + w1 = w2 >> bits; + w0 = w1 >> bits; + } + half2 d0d1 = decode_3inst_2(w0 & 0xffff, w1 & 0xffff); + half2 d2d3 = decode_3inst_2(w2 & 0xffff, w3 & 0xffff); + half2 d4d5 = decode_3inst_2(w4 & 0xffff, w5 & 0xffff); + half2 d6d7 = decode_3inst_2(w6 & 0xffff, w7 & 0xffff); + frag0[0] = d0d1; + frag0[1] = d2d3; + frag1[0] = d4d5; + frag1[1] = d6d7; +} + +template +__device__ __forceinline__ void dq8_aligned_4bits(const uint32_t* ptr, int t_offset, FragB& frag0, FragB& frag1) +{ + uint32_t i0, i1, a, b, s, w0, w1, w2, w3, w4, w5, w6, w7; + i1 = t_offset >> 3; + i0 = (i1 + 31) & 31; + a = ptr[i0]; + b = ptr[i1]; + FSHF_IMM(s, b, a, 20); + w7 = b & 0xffff; + BFE16_IMM(w6, b, 4); + BFE16_IMM(w5, b, 8); + BFE16_IMM(w4, b, 12); + BFE16_IMM(w3, b, 16); + w2 = s & 0xffff; + BFE16_IMM(w1, s, 4); + BFE16_IMM(w0, s, 8); + frag0[0] = decode_3inst_2(w0, w1); + frag0[1] = decode_3inst_2(w2, w3); + frag1[0] = decode_3inst_2(w4, w5); + frag1[1] = decode_3inst_2(w6, w7); +} + +template +__device__ __forceinline__ void dq8_aligned_2bits(const uint32_t* ptr, int t_offset, FragB& frag0, FragB& frag1) +{ + uint32_t i0, i1, a, b, w0, w1, w2, w3, w4, w5, w6, w7; + i1 = t_offset >> 4; + i0 = (i1 + 15) & 15; + a = ptr[i0]; + b = ptr[i1]; + b = fshift(b, a, ((~t_offset) & 8) << 1); + w7 = b & 0xffff; + BFE16_IMM(w6, b, 2); + BFE16_IMM(w5, b, 4); + BFE16_IMM(w4, b, 6); + BFE16_IMM(w3, b, 8); + BFE16_IMM(w2, b, 10); + BFE16_IMM(w1, b, 12); + BFE16_IMM(w0, b, 14); + frag0[0] = decode_3inst_2(w0, w1); + frag0[1] = decode_3inst_2(w2, w3); + frag1[0] = decode_3inst_2(w4, w5); + frag1[1] = decode_3inst_2(w6, w7); +} + +template +__device__ __forceinline__ void dq8_aligned_1bit(const uint32_t* ptr, int t_offset, FragB& frag0, FragB& frag1) +{ + uint32_t i0, i1, a, b, w0, w1, w2, w3, w4, w5, w6, w7; + i1 = t_offset >> 5; + i0 = (i1 + 7) & 7; + a = ptr[i0]; + b = ptr[i1]; + b = fshift(b, a, ((~t_offset) & 24)); + w7 = b & 0xffff; + BFE16_IMM(w6, b, 1); + BFE16_IMM(w5, b, 2); + BFE16_IMM(w4, b, 3); + BFE16_IMM(w3, b, 4); + BFE16_IMM(w2, b, 5); + BFE16_IMM(w1, b, 6); + BFE16_IMM(w0, b, 7); + frag0[0] = decode_3inst_2(w0, w1); + frag0[1] = decode_3inst_2(w2, w3); + frag1[0] = decode_3inst_2(w4, w5); + frag1[1] = decode_3inst_2(w6, w7); +} + + +template +__device__ __forceinline__ void dq8_aligned_4bits_bfe64(const uint32_t* ptr, int t_offset, FragB& frag0, FragB& frag1) +{ + int i1 = t_offset / 8; + int i0 = (i1 + 31) % 32; + uint32_t a = ptr[i0]; + uint32_t b = ptr[i1]; + uint32_t w7 = bfe64(b, a, 0, 16); + uint32_t w6 = bfe64(b, a, 4, 16); + uint32_t w5 = bfe64(b, a, 8, 16); + uint32_t w4 = bfe64(b, a, 12, 16); + uint32_t w3 = bfe64(b, a, 16, 16); + uint32_t w2 = bfe64(b, a, 20, 16); + uint32_t w1 = bfe64(b, a, 24, 16); + uint32_t w0 = bfe64(b, a, 28, 16); + frag0[0] = decode_3inst_2(w0, w1); + frag0[1] = decode_3inst_2(w2, w3); + frag1[0] = decode_3inst_2(w4, w5); + frag1[1] = decode_3inst_2(w6, w7); +} + +template +__device__ __forceinline__ void dq_dispatch(const uint32_t* ptr, int idx, FragB& frag0, FragB& frag1) +{ + static_assert(bits >= 1 && bits <= 8, "unsupported EXL3 bitrate"); + if constexpr (bits == 1) + { + dq8_aligned_1bit(ptr, idx, frag0, frag1); + } + else if constexpr (bits == 2) + { + dq8_aligned_2bits(ptr, idx, frag0, frag1); + } + else if constexpr (bits == 3) + { + dq8(ptr, idx, frag0, frag1); + } + else if constexpr (bits == 4) + { + dq8_aligned_4bits(ptr, idx, frag0, frag1); + } + else if constexpr (bits == 5) + { + dq4(ptr, idx, frag0); + dq4(ptr, idx + 4, frag1); + } + else if constexpr (bits == 6) + { + dq4(ptr, idx, frag0); + dq4(ptr, idx + 4, frag1); + } + else if constexpr (bits == 7) + { + dq2x2(ptr, idx, frag0); + dq2x2(ptr, idx + 4, frag1); + } + else if constexpr (bits == 8) + { + dq4(ptr, idx, frag0); + dq4(ptr, idx + 4, frag1); + } +} +// Vendored from https://github.com/brandonmmusic-max/exllamav3 at 704aefd743b390af4bd0fb429d1906f9b964c7d8. +// License: b12x/gemm/trellis_linear/csrc/vendor/LICENSE.exllamav3 diff --git a/b12x/gemm/trellis_linear/csrc/vendor/quant/exl3_gemm_inner.cuh b/b12x/gemm/trellis_linear/csrc/vendor/quant/exl3_gemm_inner.cuh new file mode 100644 index 000000000..7825f19c7 --- /dev/null +++ b/b12x/gemm/trellis_linear/csrc/vendor/quant/exl3_gemm_inner.cuh @@ -0,0 +1,782 @@ +#pragma once + +#include "../ptx.cuh" + +// Constants +#define EXL3_GEMM_BASE_THREADS 256 +#ifndef EXL3_SMEM_MAX_BYTES +#define EXL3_SMEM_MAX_BYTES (90 * 1024) +#endif +#ifndef SMEM_MAX +#define SMEM_MAX EXL3_SMEM_MAX_BYTES +#endif + +#include "exl3_dq.cuh" + +// On GA10x and sm_120 consumer silicon, fp32-accumulating HMMA runs at +// half rate. Accumulate each k-slice in fp16 and fold it into the persistent +// fp32 accumulators before reduction. This changes numerical accumulation and +// therefore remains an accuracy-gated optional variant. +#ifndef EXL3_GEMM_H_ACC + #if defined(__CUDA_ARCH__) && \ + ((__CUDA_ARCH__ == 860) || (__CUDA_ARCH__ == 1200)) + #define EXL3_GEMM_H_ACC 1 + #else + #define EXL3_GEMM_H_ACC 0 + #endif +#endif + +template +inline __device__ +void exl3_gemm_kernel_inner +( + const half* __restrict__ A, + const uint16_t* __restrict__ B, + void* __restrict__ C, + const int size_m, + const int size_k, + const int size_n, + int* __restrict__ locks, + const half* post_scale +) +{ + const int TILEBLOCKS_M = TILESIZE_M / 16; + const int TILEBLOCKS_K = TILESIZE_K / 16; + const int TILEBLOCKS_N = TILESIZE_N / 16; + // const int FRAGS_M = TILEBLOCKS_M; + const int FRAGS_N_PER_WARP = 2 * TILEBLOCKS_N / (EXL3_GEMM_BASE_THREADS / 32); + + const int sh_a_stage_size = TILESIZE_M * TILESIZE_K; // in halfs + const int sh_b_stage_size = TILEBLOCKS_K * TILEBLOCKS_N * 256 / 16 * bits; // in uint16s + const int sh_c_size = MAX // in floats + ( + 4 * EXL3_GEMM_BASE_THREADS * FRAGS_N_PER_WARP, + shmem_out_had ? TILESIZE_N * TILESIZE_M : 0 + ); + + // XOR-swizzle constants for bank-conflict-free A fragment loads + // col_swizzled = col ^ ((row >> SHIFT) & MASK) + const int A_COLS = TILESIZE_K / 8; // int4 columns per row + const int A_SWIZZLE_MASK = A_COLS - 1; + const int A_SWIZZLE_SHIFT = (A_COLS <= 2) ? 2 : 1; + + // Sanity checks + static_assert(EXL3_GEMM_BASE_THREADS == 256); + static_assert(TILESIZE_M % 16 == 0, "Invalid kernel params"); + static_assert(TILESIZE_K % 16 == 0, "Invalid kernel params"); + static_assert(TILESIZE_N % 128 == 0, "Invalid kernel params"); + static_assert + ( + SMEM_MAX >= SH_STAGES * (2 * sh_a_stage_size + 2 * sh_b_stage_size) + 4 * sh_c_size, + "Invalid kernel params (insufficient shared memory for shape)" + ); + + // Shared memory + extern __shared__ half shared[]; + half* sh_a = shared; + uint16_t* sh_b = (uint16_t*) (sh_a + SH_STAGES * sh_a_stage_size); + float* sh_c = (float*) (sh_b + sh_b_stage_size * SH_STAGES); + + // Thread index + int t = threadIdx.x % EXL3_GEMM_BASE_THREADS; + int sub_k = threadIdx.x / EXL3_GEMM_BASE_THREADS; + int warp_id = t / 32; + int lane_id = t % 32; + + // Dimensions + //int tiles_m = CEIL_DIVIDE(size_m, TILESIZE_M); + int tiles_k = size_k / TILESIZE_K; + int tiles_n = size_n / TILESIZE_N; + //int blocks_m = 1; + //int blocks_k = tiles_k * TILEBLOCKS_K; + int blocks_n = tiles_n * TILEBLOCKS_N; + + // Start and end index of current slice, must span at least one tile + int num_slices = gridDim.x; + int slice_beg = tiles_k * tiles_n * blockIdx.x / num_slices; + int slice_end = tiles_k * tiles_n * (blockIdx.x + 1) / num_slices; + int slice_len = slice_end - slice_beg; + if (slice_len < 1) return; + + auto index_m = [&] (int slice_i) { return 0; }; //blockIdx.y; }; + auto index_k = [&] (int slice_i) { return (slice_i % tiles_k); }; + auto index_n = [&] (int slice_i) { return (slice_i / tiles_k); }; + + // Batch dimension + // int slice_m = index_m(slice_beg); + // int max_m = MIN(size_m - slice_m * TILESIZE_M, TILESIZE_M); + const int slice_m = 0; + + // Pipe 0, global A, B tile and shared A, B tile + int slice0_k = index_k(slice_beg); + int slice0_n = index_n(slice_beg); + int slice0_iters = slice_len; + + int gl_a_stride_m = TILESIZE_M * size_k; + const int gl_a_stride_k = TILESIZE_K; + const int sh0_a_stride_m = TILESIZE_M * TILESIZE_K; + const half* gl_a_ptr = A + slice_m * gl_a_stride_m + slice0_k * gl_a_stride_k; + half* sh0_a_ptr = sh_a + (slice0_iters % SH_STAGES) * sh_a_stage_size; + + const int load_a_iters = CEIL_DIVIDE(sh0_a_stride_m / 8, EXL3_GEMM_BASE_THREADS); + bool pred_a_gl[load_a_iters]; + int load_a_gl[load_a_iters]; + int load_a_sh[load_a_iters]; + for (int i = 0; i < load_a_iters; ++i) + { + int k = (i * EXL3_GEMM_BASE_THREADS + t) % (gl_a_stride_k / 8); + int m = (i * EXL3_GEMM_BASE_THREADS + t) / (gl_a_stride_k / 8); + load_a_gl[i] = m * size_k / 8 + k; + load_a_sh[i] = m * A_COLS + (k ^ ((m >> A_SWIZZLE_SHIFT) & A_SWIZZLE_MASK)); + pred_a_gl[i] = m < size_m; + } + + int gl_b_stride_k = blocks_n * TILEBLOCKS_K * 256 / 16 * bits; + const int gl_b_stride_n = TILEBLOCKS_N * 256 / 16 * bits; + const int sh0_b_stride_k = TILEBLOCKS_K * TILEBLOCKS_N * 256 / 16 * bits; + const uint16_t* gl_b_ptr = B + slice0_k * gl_b_stride_k + slice0_n * gl_b_stride_n; + uint16_t* sh0_b_ptr = sh_b + (slice0_iters % SH_STAGES) * sh_b_stage_size; + + const int load_b_iters = CEIL_DIVIDE(sh0_b_stride_k / 8, EXL3_GEMM_BASE_THREADS); + bool pred_b_gl[load_b_iters]; + int load_b_gl[load_b_iters]; + for (int i = 0; i < load_b_iters; ++i) + { + int n = (i * EXL3_GEMM_BASE_THREADS + t) % (gl_b_stride_n / 8); + int k = (i * EXL3_GEMM_BASE_THREADS + t) / (gl_b_stride_n / 8); + load_b_gl[i] = k * (blocks_n * 256 / 16 * bits / 8) + n; + pred_b_gl[i] = i * EXL3_GEMM_BASE_THREADS + t < sh0_b_stride_k / 8; + } + + auto advance0 = [&] () + { + slice0_k++; + slice0_iters--; + + int stage = slice0_iters % SH_STAGES; + sh0_a_ptr = sh_a + stage * sh_a_stage_size; + sh0_b_ptr = sh_b + stage * sh_b_stage_size; + + if (slice0_k >= tiles_k) + { + slice0_k = 0; + slice0_n++; + gl_a_ptr = A + slice_m * gl_a_stride_m + slice0_k * gl_a_stride_k; + gl_b_ptr = B + slice0_k * gl_b_stride_k + slice0_n * gl_b_stride_n; + } + else + { + gl_a_ptr += gl_a_stride_k; + gl_b_ptr += gl_b_stride_k; + } + }; + + // Pipe 1, shared A, B tile and registers + int slice1_k = slice0_k; + int slice1_n = slice0_n; + int slice1_iters = slice0_iters; + + half* sh1_a_ptr = sh_a + (slice1_iters % SH_STAGES) * sh_a_stage_size; + uint16_t* sh1_b_ptr = sh_b + (slice1_iters % SH_STAGES) * sh_b_stage_size; + + auto advance1 = [&] () + { + slice1_k++; + slice1_iters--; + + int stage = slice1_iters % SH_STAGES; + sh1_a_ptr = sh_a + stage * sh_a_stage_size; + sh1_b_ptr = sh_b + stage * sh_b_stage_size; + + if (slice1_k >= tiles_k) + { + slice1_k = 0; + slice1_n++; + } + }; + + // Pipe 2 + int slice2_k = slice0_k; + int slice2_k0 = slice0_k; + int slice2_n = slice0_n; + int slice2_iters = slice0_iters; + + int gl_c_stride_n = TILESIZE_N; + int gl_c_stride_m = TILESIZE_M * size_n; + + half* gl_c_ptr_16 = ((half*) C) + slice_m * gl_c_stride_m + slice2_n * gl_c_stride_n; + float* gl_c_ptr_32 = ((float*) C) + slice_m * gl_c_stride_m + slice2_n * gl_c_stride_n; + + FragA frag_a[FRAG_STAGES][TILEBLOCKS_M]; + FragB frag_b[FRAG_STAGES][FRAGS_N_PER_WARP]; + FragC frag_c[TILEBLOCKS_M][FRAGS_N_PER_WARP]; + #if EXL3_GEMM_H_ACC + FragC_h frag_c_h[TILEBLOCKS_M][FRAGS_N_PER_WARP]; + #endif + + auto advance2 = [&] () + { + slice2_k++; + slice2_iters--; + + if (slice2_k >= tiles_k) + { + slice2_k = 0; + slice2_k0 = 0; + slice2_n++; + if constexpr (c_fp32) + gl_c_ptr_32 += gl_c_stride_n; + else + gl_c_ptr_16 += gl_c_stride_n; + } + }; + + // Schedule load of the next A, B tiles to shared memory and advance the pipeline + auto async_load_gl = [&] () + { + if (sub_k) + { + cp_async_fence(); + return; + } + + if (slice0_iters) + { + // Copy tile from row-major A matrix (XOR-swizzled for bank-conflict-free ldmatrix) + { + const int4* gl = (const int4*) gl_a_ptr; + int4* sh = (int4*) sh0_a_ptr; + #pragma unroll + for (int i = 0; i < load_a_iters; ++i) + { + if (pred_a_gl[i]) cp_async(sh + load_a_sh[i], gl + load_a_gl[i]); + } + } + + // Copy tile of 256-element blocks from quantized B matrix + { + const int4* gl = (const int4*) gl_b_ptr; + int4* sh = (int4*) sh0_b_ptr; + #pragma unroll + for (int i = 0; i < load_b_iters; ++i) + { + // cp_async_pred(sh + EXL3_GEMM_BASE_THREADS * i + t, gl + load_b_gl[i], pred_b_gl[i]); + if (pred_b_gl[i]) cp_async(sh + EXL3_GEMM_BASE_THREADS * i + t, gl + load_b_gl[i]); + } + } + advance0(); + } + + // Sync and advance + cp_async_fence(); + }; + + // Load fragments + // Ref. for fragment layout: + // https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#matrix-fragments-for-mma-m16n8k16-with-floating-point-type + auto load_frags = [&] (int buf) + { + if (!slice1_iters) return; + + // A fragments (XOR-swizzled shared memory layout) + { + int r = (lane_id % 8) + 8 * ((lane_id / 8) % 2); + int base_c = lane_id / 16 + sub_k * 2; + #pragma unroll + for (int m = 0; m < TILEBLOCKS_M; ++m) + { + int R = r + m * 16; + int c_swizzled = base_c ^ ((R >> A_SWIZZLE_SHIFT) & A_SWIZZLE_MASK); + ldsm4(frag_a[buf][m], (int4*) sh1_a_ptr + R * A_COLS + c_swizzled); + } + } + + // B fragments + #pragma unroll + for (int n2 = 0; n2 < FRAGS_N_PER_WARP; n2 += 2) + { + int sub_n2 = warp_id * FRAGS_N_PER_WARP / 2 + n2 / 2; + const uint32_t* shb = (const uint32_t*) (sh1_b_ptr + (sub_k * TILEBLOCKS_N + sub_n2) * 256 / 16 * bits); + + dq_dispatch(shb, lane_id << 3, frag_b[buf][n2], frag_b[buf][n2 + 1]); + } + + __syncthreads(); + advance1(); + }; + + // Clear C fragments + auto clear_frag_c = [&] () + { + #pragma unroll + for (int m = 0; m < TILEBLOCKS_M; ++m) + { + #pragma unroll + for (int n = 0; n < FRAGS_N_PER_WARP; ++n) + frag_c[m][n] = {}; + } + #if EXL3_GEMM_H_ACC + #pragma unroll + for (int m = 0; m < TILEBLOCKS_M; ++m) + { + #pragma unroll + for (int n = 0; n < FRAGS_N_PER_WARP; ++n) + frag_c_h[m][n] = {}; + } + #endif + }; + + // Threadblock reduction + auto threadblock_reduce = [&] () + { + auto store = [&] (int i, int m) + { + if (sub_k == i) + { + float* sh_red = sh_c + (FRAGS_N_PER_WARP * 4) * t; + #pragma unroll + for (int n = 0; n < FRAGS_N_PER_WARP; ++n) + { + #pragma unroll + for (int j = 0; j < 4; ++j) *sh_red++ = frag_c[m][n][j]; + } + } + __syncthreads(); + }; + + auto add = [&] (int i, int m) + { + if (sub_k == i) + { + float* sh_red = sh_c + (FRAGS_N_PER_WARP * 4) * t; + #pragma unroll + for (int n = 0; n < FRAGS_N_PER_WARP; ++n) + { + #pragma unroll + for (int j = 0; j < 4; ++j) frag_c[m][n][j] += *sh_red++; + } + } + }; + + auto store_small = [&] (int i, int m) + { + if (sub_k == i && m * 16 + lane_id / 4 < size_m) + { + float* sh_red = sh_c + (FRAGS_N_PER_WARP * 4) * t; + #pragma unroll + for (int n = 0; n < FRAGS_N_PER_WARP; ++n) + { + *sh_red++ = frag_c[m][n][0]; + *sh_red++ = frag_c[m][n][1]; + } + } + __syncthreads(); + }; + + auto add_small = [&] (int i, int m) + { + if (sub_k == i && m * 16 + lane_id / 4 < size_m) + { + float* sh_red = sh_c + (FRAGS_N_PER_WARP * 4) * t; + #pragma unroll + for (int n = 0; n < FRAGS_N_PER_WARP; ++n) + { + frag_c[m][n][0] += *sh_red++; + frag_c[m][n][1] += *sh_red++; + } + } + }; + + // Reuse the same reduction scratch for each 16-row M block. The + // barrier between blocks prevents the next store from racing the + // preceding add without increasing the dynamic-smem footprint. + #pragma unroll + for (int m = 0; m < TILEBLOCKS_M; ++m) + { + if (size_m <= 8) + { + if constexpr (TILEBLOCKS_K == 2) + { + store_small(1, m); + add_small(0, m); + } + if constexpr (TILEBLOCKS_K == 3) + { + store_small(1, m); + add_small(0, m); + store_small(2, m); + add_small(0, m); + } + if constexpr (TILEBLOCKS_K == 4) + { + store_small(3, m); + add_small(2, m); + store_small(1, m); + add_small(0, m); + store_small(2, m); + add_small(0, m); + } + } + else + { + if constexpr (TILEBLOCKS_K == 2) + { + store(1, m); + add(0, m); + } + if constexpr (TILEBLOCKS_K == 3) + { + store(1, m); + add(0, m); + store(2, m); + add(0, m); + } + if constexpr (TILEBLOCKS_K == 4) + { + store(3, m); + add(2, m); + store(1, m); + add(0, m); + store(2, m); + add(0, m); + } + } + + if constexpr (TILEBLOCKS_K > 1 && TILEBLOCKS_M > 1) + if (m + 1 < TILEBLOCKS_M) __syncthreads(); + } + }; + + // Pre-hadamard: Write final output tile to shmem + auto write_sum_tile_sh = [&]() + { + const int n0 = warp_id * FRAGS_N_PER_WARP; + #pragma unroll + for (int m = 0; m < TILEBLOCKS_M; ++m) + { + const int r0 = m * 16 + lane_id / 4; + const int r1 = r0 + 8; + if (r0 < size_m) + { + const int c = (lane_id % 4) * 2; + #pragma unroll + for (int n = 0; n < FRAGS_N_PER_WARP; ++n) + { + float* c_ptr = sh_c + r0 * TILESIZE_N + (n0 + n) * 8 + c; + *c_ptr++ = frag_c[m][n][0]; + *c_ptr++ = frag_c[m][n][1]; + } + } + if (r1 < size_m) + { + const int c = (lane_id % 4) * 2; + #pragma unroll + for (int n = 0; n < FRAGS_N_PER_WARP; ++n) + { + float* c_ptr = sh_c + r1 * TILESIZE_N + (n0 + n) * 8 + c; + *c_ptr++ = frag_c[m][n][2]; + *c_ptr++ = frag_c[m][n][3]; + } + } + } + }; + + // Copy output tile to global with hadamard transform and out scale + auto output_had_sh_gl = [&]() + { + int sh_warp = warp_id; + constexpr int active_warps = EXL3_GEMM_BASE_THREADS / 32; + for (;; sh_warp += active_warps) + { + int col = sh_warp % (TILESIZE_N / 128); + int row = sh_warp / (TILESIZE_N / 128); + if (row >= size_m) break; + + const float* had_in = sh_c + row * TILESIZE_N + col * 128; + const half* post_scale_c = post_scale + slice2_n * gl_c_stride_n + col * 128; + + if constexpr (c_fp32) + { + float* had_out = gl_c_ptr_32 + row * size_n + col * 128; + had_ff_r_128_inner(had_in, had_out, post_scale_c, 0.088388347648f); + } + else + { + half* had_out = gl_c_ptr_16 + row * size_n + col * 128; + had_fh_r_128_inner(had_in, had_out, post_scale_c, 0.088388347648f); + } + } + }; + + auto read_sum_gl = [&]() + { + int n0 = warp_id * FRAGS_N_PER_WARP; + #pragma unroll + for (int m = 0; m < TILEBLOCKS_M; ++m) + { + #pragma unroll + for (int n = 0; n < FRAGS_N_PER_WARP; ++n) + { + int r0 = m * 16 + lane_id / 4; + int r1 = r0 + 8; + int c = (lane_id % 4) * 2; + if (r0 < size_m) + { + if constexpr (c_fp32) + { + float* c_ptr = gl_c_ptr_32 + r0 * size_n + (n0 + n) * 8 + c; + frag_c[m][n][0] += *c_ptr++; + frag_c[m][n][1] += *c_ptr++; + } + else + { + half2* c_ptr = (half2*) (gl_c_ptr_16 + r0 * size_n + (n0 + n) * 8 + c); + float2 interm = __half22float2(*c_ptr); + frag_c[m][n][0] += interm.x; + frag_c[m][n][1] += interm.y; + } + } + if (r1 < size_m) + { + if constexpr (c_fp32) + { + float* c_ptr = gl_c_ptr_32 + r1 * size_n + (n0 + n) * 8 + c; + frag_c[m][n][2] += *c_ptr++; + frag_c[m][n][3] += *c_ptr++; + } + else + { + half2* c_ptr = (half2*) (gl_c_ptr_16 + r1 * size_n + (n0 + n) * 8 + c); + float2 interm = __half22float2(*c_ptr); + frag_c[m][n][2] += interm.x; + frag_c[m][n][3] += interm.y; + } + } + } + } + }; + + auto write_sum_gl = [&]() + { + int n0 = warp_id * FRAGS_N_PER_WARP; + #pragma unroll + for (int m = 0; m < TILEBLOCKS_M; ++m) + { + #pragma unroll + for (int n = 0; n < FRAGS_N_PER_WARP; ++n) + { + int r0 = m * 16 + lane_id / 4; + int r1 = r0 + 8; + int c = (lane_id % 4) * 2; + if (r0 < size_m) + { + if constexpr (c_fp32) + { + float* c_ptr = gl_c_ptr_32 + r0 * size_n + (n0 + n) * 8 + c; + *c_ptr++ = frag_c[m][n][0]; + *c_ptr++ = frag_c[m][n][1]; + } + else + { + half2* c_ptr = (half2*) (gl_c_ptr_16 + r0 * size_n + (n0 + n) * 8 + c); + half2 sum = __floats2half2_rn(frag_c[m][n][0], frag_c[m][n][1]); + *c_ptr = sum; + } + } + if (r1 < size_m) + { + if constexpr (c_fp32) + { + float* c_ptr = gl_c_ptr_32 + r1 * size_n + (n0 + n) * 8 + c; + *c_ptr++ = frag_c[m][n][2]; + *c_ptr++ = frag_c[m][n][3]; + } + else + { + half2* c_ptr = (half2*) (gl_c_ptr_16 + r1 * size_n + (n0 + n) * 8 + c); + half2 sum = __floats2half2_rn(frag_c[m][n][2], frag_c[m][n][3]); + *c_ptr = sum; + } + } + } + } + }; + + // Output reduction + auto reduce = [&] () + { + #if EXL3_GEMM_H_ACC + // Fold the fp16 MMA accumulators into fp32 once per k-slice. + #pragma unroll + for (int m = 0; m < TILEBLOCKS_M; ++m) + { + #pragma unroll + for (int n = 0; n < FRAGS_N_PER_WARP; ++n) + { + float2 f0 = __half22float2(frag_c_h[m][n][0]); + float2 f1 = __half22float2(frag_c_h[m][n][1]); + frag_c[m][n][0] += f0.x; + frag_c[m][n][1] += f0.y; + frag_c[m][n][2] += f1.x; + frag_c[m][n][3] += f1.y; + } + } + #endif + + // First reduce all partial sums along k for the current slice + threadblock_reduce(); + + // Process (partial) slices within column in reverse order so the threadblock doing the bottom slice is + // free to proceed to the next column right away + int lock_i = tiles_k - slice2_k - 1; + int lock_d = slice2_k - slice2_k0 + 1; + int* lock = &locks[slice_m * blocks_n + slice2_n]; + + barrier_acquire(lock, lock_i); + + bool first = lock_i == 0; + bool last = lock_i + lock_d == tiles_k; + + // Second and subsequent threadblocks in column read back the intermediate sum from global memory + if (!sub_k && !first) + { + read_sum_gl(); + } + + // All but last threadblock in column write the intermediate result to global memory + if (!sub_k && !last) + { + write_sum_gl(); + } + + // Last block writes in row-major format + if (!sub_k && last) + { + if constexpr (shmem_out_had) + write_sum_tile_sh(); + else + write_sum_gl(); + } + + if constexpr (shmem_out_had) + { + if (last) __syncthreads(); + if (!sub_k && last) + output_had_sh_gl(); + } + + barrier_release(lock, lock_d, last); + + clear_frag_c(); + }; + + // Wait until there are at most SH_STAGES - 2 async copies pending, i.e. at least one stage has finished loading + auto wait_stage = [&] () + { + cp_async_wait(); + __syncthreads(); + }; + + // Perform tensor core matmul on current tile + auto matmul = [&] (int buf) + { + #pragma unroll + for (int m = 0; m < TILEBLOCKS_M; ++m) + { + #pragma unroll + for (int n = 0; n < FRAGS_N_PER_WARP; ++n) + { + #if EXL3_GEMM_H_ACC + ptx_mma_m16n8k16(frag_a[buf][m], frag_b[buf][n], frag_c_h[m][n]); + #else + ptx_mma_m16n8k16(frag_a[buf][m], frag_b[buf][n], frag_c[m][n]); + #endif + } + } + }; + + // Start global to shared pipeline + #pragma unroll + for (int i = 0; i < SH_STAGES - 1; ++i) + async_load_gl(); + wait_stage(); + + // Start shared to register pipeline. + clear_frag_c(); + if constexpr (FRAG_STAGES > 1) + load_frags(0); + + // Main loop. Fragments are double buffered to allow more interleaving. This is especially important to hide the + // dequantization overhead, but we need two different iterations of the main loop to avoid confusing the compiler + // and making it (sometimes) place the fragment arrays in local memory + + #define FSTAGE_OLD(_load, _mul) \ + async_load_gl(); \ + wait_stage(); \ + load_frags(_load); \ + matmul(_mul); \ + if (slice2_k == tiles_k - 1 || slice2_iters == 1) { reduce(); slice2_k0 = slice2_k + 1; } \ + advance2(); \ + if (!slice2_iters) break; \ + + #define FSTAGE(_load, _mul) \ + async_load_gl(); \ + wait_stage(); \ + matmul(_mul); \ + if (slice2_k == tiles_k - 1 || slice2_iters == 1) { reduce(); slice2_k0 = slice2_k + 1; } \ + advance2(); \ + if (!slice2_iters) break; \ + load_frags(_load); \ + + if constexpr (FRAG_STAGES == 1) + { + while (true) + { + FSTAGE_OLD(0, 0); + } + } + + if constexpr (FRAG_STAGES == 2) + { + while (true) + { + FSTAGE(1, 0); + FSTAGE(0, 1); + } + } + + if constexpr (FRAG_STAGES == 3) + { + while (true) + { + FSTAGE(1, 0); + FSTAGE(2, 1); + FSTAGE(0, 2); + } + } + + if constexpr (FRAG_STAGES == 4) + { + while (true) + { + FSTAGE(1, 0); + FSTAGE(2, 1); + FSTAGE(3, 2); + FSTAGE(0, 3); + } + } + + if constexpr (FRAG_STAGES == 5) + { + while (true) + { + FSTAGE(1, 0); + FSTAGE(2, 1); + FSTAGE(3, 2); + FSTAGE(4, 3); + FSTAGE(0, 4); + } + } +} +// Vendored from https://github.com/brandonmmusic-max/exllamav3 at 704aefd743b390af4bd0fb429d1906f9b964c7d8. +// License: b12x/gemm/trellis_linear/csrc/vendor/LICENSE.exllamav3 diff --git a/b12x/gemm/trellis_linear/csrc/vendor/quant/exl3_gemm_kernel.cuh b/b12x/gemm/trellis_linear/csrc/vendor/quant/exl3_gemm_kernel.cuh new file mode 100644 index 000000000..001004510 --- /dev/null +++ b/b12x/gemm/trellis_linear/csrc/vendor/quant/exl3_gemm_kernel.cuh @@ -0,0 +1,281 @@ +#pragma once + +#include "exl3_kernel_map.cuh" +#include "hadamard_inner.cuh" +#include "exl3_gemm_inner.cuh" +#include "exl3_devctx.cuh" + +template +__global__ __launch_bounds__(EXL3_GEMM_BASE_THREADS * TILESIZE_K / 16) +void exl3_gemm_kernel(EXL3_GEMM_ARGS) +{ + auto grid = cg::this_grid(); + + // if (suh) + { + int total_warps = size_m * size_k / 128; + int warps_grid = gridDim.x * blockDim.x / 32; + int this_warp = threadIdx.x / 32 + blockDim.x / 32 * blockIdx.x; + + for(; this_warp < total_warps; this_warp += warps_grid) + had_hf_r_128_inner + ( + A + this_warp * 128, + A_had + this_warp * 128, + suh + (this_warp * 128) % size_k, + 0.088388347648f // 1/sqrt(128) + ); + + grid.sync(); + A = A_had; + } + + int size_m_ = size_m; + const half* A_ = A; + void* C_ = C; + + while (size_m_ > 0) + { + exl3_gemm_kernel_inner + + (A_, B, C_, MIN(size_m_, 16), size_k, size_n, locks, svh); + + A_ += 16 * size_k; + if constexpr (c_fp32) C_ = (void*) (((float*) C_) + 16 * size_n); + else C_ = (void*) (((half*) C_) + 16 * size_n); + size_m_ -= 16; + + if (size_m_ > 0 || svh) + grid.sync(); + } + + // if (svh) + /* + { + int total_warps = size_m * size_n / 128; + int warps_grid = gridDim.x * blockDim.x / 32; + int this_warp = threadIdx.x / 32 + blockDim.x / 32 * blockIdx.x; + + for(; this_warp < total_warps; this_warp += warps_grid) + { + if constexpr (c_fp32) + had_ff_r_128_inner + ( + ((const float*) C) + this_warp * 128, + ((float*) C) + this_warp * 128, + svh + (this_warp * 128) % size_n, + 0.088388347648f // 1/sqrt(128) + ); + else + had_hf_r_128_inner + ( + ((const half*) C) + this_warp * 128, + ((half*) C) + this_warp * 128, + svh + (this_warp * 128) % size_n, + 0.088388347648f // 1/sqrt(128) + ); + } + } + */ +} + +#define MAX_INDICES 128 + +__device__ int64_t v_indices[128]; +__device__ half v_weights[128]; +__device__ int bszm_sync; + +template +__global__ __launch_bounds__(EXL3_GEMM_BASE_THREADS * TILESIZE_K / 16) +void exl3_mgemm_kernel(EXL3_MGEMM_ARGS) +{ + int bszm = MAX(bszm_in, bszm_out); + auto grid = cg::this_grid(); + + #if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ > 890) + int* barrier_counters_sense = locks + BARRIER_LOCKS_OFFSET; + #endif + + // Pack indices within min_index <= idx < max_index + + if (min_index >= 0) + { + if (blockIdx.x == 0 && blockIdx.y == 0 && blockIdx.z == 0 && threadIdx.x == 0) + { + int j = 0; + for (int i = 0; i < bszm; ++i) + { + int idx = B_indices[i]; + if (idx >= min_index && idx < max_index) + { + v_indices[j] = idx - min_index; + if (B_weights) v_weights[j] = B_weights[i]; + j++; + } + } + bszm_sync = j; + for (; j < bszm; ++j) + { + v_indices[j] = -1; + } + } + __threadfence(); + grid.sync(); + B_indices = v_indices; + if (B_weights) B_weights = v_weights; + bszm = bszm_sync; + } + + for (int i = 0; i < bszm; i += gridDim.z) + { + int j = i + blockIdx.z; + int mat_index = -1; + const uint16_t* B = nullptr; + if (j >= bszm) j = -1; + else + { + mat_index = B_indices ? (int) B_indices[j] : j; + if (mat_index >= 0) + { + B = B_list[mat_index]; + } + } + + // Had and input scales + + if (B) + { + int total_warps = size_m * size_k / 128; + int warps_grid = gridDim.x * blockDim.x / 32; + int this_warp = threadIdx.x / 32 + blockDim.x / 32 * blockIdx.x; + + const half* suh = suh_list[mat_index]; + const half* A_ = bszm_in == 1 ? A : A + j * size_m * size_k; + half* A_had_ = A_had + j * size_m * size_k; + + for(; this_warp < total_warps; this_warp += warps_grid) + had_hf_r_128_inner + ( + A_ + this_warp * 128, + A_had_ + this_warp * 128, + suh + (this_warp * 128) % size_k, + 0.088388347648f // 1/sqrt(128) + ); + } + + #if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ > 890) + group_barrier(blockIdx.z, gridDim.x, barrier_counters_sense); + #else + grid.sync(); + #endif + + // Matmul + + int size_m_ = size_m; + half* A_ = A_had + j * size_m * size_k; + void* C_; + if constexpr (c_fp32) C_ = (void*) (((float*) C) + j * size_m * size_n); + else C_ = (void*) (((half*) C) + j * size_m * size_n); + + while (size_m_ > 0) + { + if (B) + { + int lock_offs = blockIdx.z * size_n / 128; + + exl3_gemm_kernel_inner + + (A_, B, C_, MIN(size_m_, 16), size_k, size_n, locks + lock_offs, nullptr); + } + + A_ += 16 * size_k; + if constexpr (c_fp32) C_ = (void*) (((float*) C_) + 16 * size_n); + else C_ = (void*) (((half*) C_) + 16 * size_n); + size_m_ -= 16; + + #if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ > 890) + group_barrier(blockIdx.z, gridDim.x, barrier_counters_sense); + #else + grid.sync(); + #endif + } + + // Had and output scales + + if (B) + { + int total_warps = size_m * size_n / 128; + int warps_grid = gridDim.x * blockDim.x / 32; + int this_warp = threadIdx.x / 32 + blockDim.x / 32 * blockIdx.x; + + const half* svh = svh_list[mat_index]; + float scale = 0.088388347648f; // 1/sqrt(128) + if (B_weights) scale *= __half2float(B_weights[j]); + + if constexpr (c_fp32) C_ = (void*) (((float*) C) + j * size_m * size_n); + else C_ = (void*) (((half*) C) + j * size_m * size_n); + + for(; this_warp < total_warps; this_warp += warps_grid) + { + if constexpr (c_fp32) + had_ff_r_128_inner + ( + ((const float*) C_) + this_warp * 128, + ((float*) C_) + this_warp * 128, + svh + (this_warp * 128) % size_n, + scale + ); + else + had_hf_r_128_inner + ( + ((const half*) C_) + this_warp * 128, + ((half*) C_) + this_warp * 128, + svh + (this_warp * 128) % size_n, + scale + ); + } + } + } + + if (B_weights) + grid.sync(); + + // Final reduction + if (B_weights && blockIdx.z == 0) + { + int total_warps = size_m * size_n / 32; + int warps_grid = gridDim.x * blockDim.x / 32; + int this_warp = threadIdx.x / 32 + blockDim.x / 32 * blockIdx.x; + int this_lane = threadIdx.x % 32; + + for(; this_warp < total_warps; this_warp += warps_grid) + { + if constexpr (c_fp32) + { + float* C__ = ((float*) C) + this_warp * 32 + this_lane; + float* C___ = C__; + float sum = 0.0f; + for (int j = 0; j < bszm; ++j) + { + sum += *C___; + C___ += size_m * size_n; + } + *C__ = sum; + } + else + { + half* C__ = ((half*) C) + this_warp * 32 + this_lane; + half* C___ = C__; + half sum = {}; + for (int j = 0; j < bszm; ++j) + { + sum = __hadd(sum, *C___); + C___ += size_m * size_n; + } + *C__ = sum; + } + } + } +} +// Vendored from https://github.com/brandonmmusic-max/exllamav3 at 704aefd743b390af4bd0fb429d1906f9b964c7d8. +// License: b12x/gemm/trellis_linear/csrc/vendor/LICENSE.exllamav3 diff --git a/b12x/gemm/trellis_linear/csrc/vendor/quant/exl3_kernel_map.cuh b/b12x/gemm/trellis_linear/csrc/vendor/quant/exl3_kernel_map.cuh new file mode 100644 index 000000000..585dd55e4 --- /dev/null +++ b/b12x/gemm/trellis_linear/csrc/vendor/quant/exl3_kernel_map.cuh @@ -0,0 +1,144 @@ +#pragma once + +int select_gemm_shape(int cc, int size_m, int size_k, int size_n, int bits, bool multi); +int exl3_gemm_num_kernel_shapes(); +bool exl3_gemm_shape_compat(int shape_idx, int size_m, int size_k, int size_n, int bits); + +#define EXL3_GEMM_T_ARGS \ + const int bits, \ + const bool c_fp32, \ + const int cb, \ + const int TILESIZE_M, \ + const int TILESIZE_K, \ + const int TILESIZE_N, \ + const int SH_STAGES, \ + const int FRAG_STAGES + +#define EXL3_GEMM_ARGS \ + const half* __restrict__ A, \ + const uint16_t* __restrict__ B, \ + void* __restrict__ C, \ + const int size_m, \ + const int size_k, \ + const int size_n, \ + int* __restrict__ locks, \ + const half* __restrict__ suh, \ + half* __restrict__ A_had, \ + const half* __restrict__ svh + +#define EXL3_MGEMM_ARGS \ + const half* __restrict__ A, \ + const uint16_t** __restrict__ B_list, \ + void* __restrict__ C, \ + const int size_m, \ + const int size_k, \ + const int size_n, \ + int* __restrict__ locks, \ + const half** __restrict__ suh_list, \ + half* __restrict__ A_had, \ + const half** __restrict__ svh_list, \ + int64_t* B_indices, \ + half* B_weights, \ + const int bszm_in, \ + const int bszm_out, \ + const int min_index, \ + const int max_index + +typedef void (*fp_exl3_gemm_kernel) (EXL3_GEMM_ARGS); +typedef void (*fp_exl3_mgemm_kernel) (EXL3_MGEMM_ARGS); + +#define EXL3_GEMM_SHAPE_1 16, 16, 128, 6, 5 +#define EXL3_GEMM_SHAPE_2 16, 32, 128, 4, 3 +#define EXL3_GEMM_SHAPE_3 16, 32, 256, 4, 3 +#define EXL3_GEMM_SHAPE_4 16, 16, 512, 4, 3 + +#define EXL3_GEMM_TILESIZE_K 0, 16, 32, 32, 16 +#define EXL3_GEMM_TILESIZE_N 0, 128, 128, 256, 512 +#define EXL3_GEMM_BLOCKDIM 0, 256, 512, 512, 256 + +#define EXL3_GEMM_NUM_SHAPES 4 + +// Shape 1 not currently used anywhere +#define EXL3_GEMM_KERNEL_INSTANCES(_bits, _c_fp32, cb) \ + nullptr, \ + exl3_gemm_kernel<_bits, _c_fp32, cb, EXL3_GEMM_SHAPE_1>, \ + exl3_gemm_kernel<_bits, _c_fp32, cb, EXL3_GEMM_SHAPE_2>, \ + exl3_gemm_kernel<_bits, _c_fp32, cb, EXL3_GEMM_SHAPE_3>, \ + exl3_gemm_kernel<_bits, _c_fp32, cb, EXL3_GEMM_SHAPE_4> + +#define EXL3_MGEMM_KERNEL_INSTANCES(_bits, _c_fp32, cb) \ + nullptr, \ + exl3_mgemm_kernel<_bits, _c_fp32, cb, EXL3_GEMM_SHAPE_1>, \ + exl3_mgemm_kernel<_bits, _c_fp32, cb, EXL3_GEMM_SHAPE_2>, \ + exl3_mgemm_kernel<_bits, _c_fp32, cb, EXL3_GEMM_SHAPE_3>, \ + exl3_mgemm_kernel<_bits, _c_fp32, cb, EXL3_GEMM_SHAPE_4> + +#define EXL3_GEMM_BASE_THREADS 256 + +#define ALL_EXL3_KERNEL_EXTERNS(K) \ + extern fp_exl3_gemm_kernel tfp_exl3_gemm_kernel_fp32_b##K[]; \ + extern fp_exl3_gemm_kernel tfp_exl3_gemm_kernel_fp16_b##K[]; \ + extern fp_exl3_mgemm_kernel tfp_exl3_mgemm_kernel_fp32_b##K[]; \ + extern fp_exl3_mgemm_kernel tfp_exl3_mgemm_kernel_fp16_b##K[]; \ + +#define ALL_EXL3_KERNEL_INSTANCES(K) \ + fp_exl3_gemm_kernel tfp_exl3_gemm_kernel_fp32_b##K[] = { \ + EXL3_GEMM_KERNEL_INSTANCES(K, true, 0), \ + EXL3_GEMM_KERNEL_INSTANCES(K, true, 1), \ + EXL3_GEMM_KERNEL_INSTANCES(K, true, 2) \ + }; \ + \ + fp_exl3_gemm_kernel tfp_exl3_gemm_kernel_fp16_b##K[] = { \ + EXL3_GEMM_KERNEL_INSTANCES(K, false, 0), \ + EXL3_GEMM_KERNEL_INSTANCES(K, false, 1), \ + EXL3_GEMM_KERNEL_INSTANCES(K, false, 2) \ + }; \ + \ + fp_exl3_mgemm_kernel tfp_exl3_mgemm_kernel_fp32_b##K[] = { \ + EXL3_MGEMM_KERNEL_INSTANCES(K, true, 0), \ + EXL3_MGEMM_KERNEL_INSTANCES(K, true, 1), \ + EXL3_MGEMM_KERNEL_INSTANCES(K, true, 2) \ + }; \ + \ + fp_exl3_mgemm_kernel tfp_exl3_mgemm_kernel_fp16_b##K[] = { \ + EXL3_MGEMM_KERNEL_INSTANCES(K, false, 0), \ + EXL3_MGEMM_KERNEL_INSTANCES(K, false, 1), \ + EXL3_MGEMM_KERNEL_INSTANCES(K, false, 2) \ + }; + +fp_exl3_gemm_kernel select_exl3_gemm_kernel +( + const int cc, + const int size_m, + const int size_k, + const int size_n, + const int bits, + const bool c_fp32, + const int force_shape_idx, + int* out_block_dim, + int* out_shape_idx, + int* out_num_sms, + const int cb +); + +fp_exl3_mgemm_kernel select_exl3_mgemm_kernel +( + const int cc, + const int size_m, + const int size_k, + const int size_n, + const int K, + const bool c_fp32, + const int force_shape_idx, + int* out_block_dim, + int* out_shape_idx, + int* out_num_sms, + const int cb, + const int bszm_in, + const int bszm_out +); + +fp_exl3_gemm_kernel get_gemm_kernel_ptr(int K, int shape_idx, bool c_fp32, int cb); +fp_exl3_mgemm_kernel get_mgemm_kernel_ptr(int K, int shape_idx, bool c_fp32, int cb); +// Vendored from https://github.com/brandonmmusic-max/exllamav3 at 704aefd743b390af4bd0fb429d1906f9b964c7d8. +// License: b12x/gemm/trellis_linear/csrc/vendor/LICENSE.exllamav3 diff --git a/b12x/gemm/trellis_linear/csrc/vendor/quant/hadamard_inner.cuh b/b12x/gemm/trellis_linear/csrc/vendor/quant/hadamard_inner.cuh new file mode 100644 index 000000000..f34ff1787 --- /dev/null +++ b/b12x/gemm/trellis_linear/csrc/vendor/quant/hadamard_inner.cuh @@ -0,0 +1,461 @@ +#pragma once +#include "../compat.cuh" + +#define ACT_SILU 0 +#define ACT_GELU 1 + +// Hadamard transform 128-element vector across one warp, with optional pre and post scales + +__device__ inline half hreduce(half2 x) +{ + return __hadd(__low2half(x), __high2half(x)); +} + +__device__ inline void shuffle_had_f4x32(float& h0, float& h1, float& h2, float& h3, const int lane_id) +{ + #pragma unroll + for (int i = 1; i < 32; i <<= 1) + { + uint32_t i0 = __float_as_uint(h0); + uint32_t i1 = __float_as_uint(h1); + uint32_t i2 = __float_as_uint(h2); + uint32_t i3 = __float_as_uint(h3); + uint64_t h01 = (uint64_t) i0 | (((uint64_t) i1) << 32); + uint64_t h23 = (uint64_t) i2 | (((uint64_t) i3) << 32); + uint64_t ph01 = __shfl_xor_sync(0xffffffff, h01, i); + uint64_t ph23 = __shfl_xor_sync(0xffffffff, h23, i); + float ph0 = __uint_as_float((uint32_t) (ph01 & 0xffffffff)); + float ph1 = __uint_as_float((uint32_t) (ph01 >> 32)); + float ph2 = __uint_as_float((uint32_t) (ph23 & 0xffffffff)); + float ph3 = __uint_as_float((uint32_t) (ph23 >> 32)); + int32_t sfm = -static_cast(lane_id & i) >> 31; + i0 ^= sfm & 0x80000000; + i1 ^= sfm & 0x80000000; + i2 ^= sfm & 0x80000000; + i3 ^= sfm & 0x80000000; + h0 = __uint_as_float(i0) + ph0; + h1 = __uint_as_float(i1) + ph1; + h2 = __uint_as_float(i2) + ph2; + h3 = __uint_as_float(i3) + ph3; + } +} + +__device__ inline void shuffle_had_f2x32(float& v, float& w, const int lane_id) +{ + #pragma unroll + for (int i = 1; i < 32; i <<= 1) + { + uint64_t vw = ((uint64_t) __float_as_uint(v)) | (((uint64_t) __float_as_uint(w)) << 32); + uint64_t pvw = __shfl_xor_sync(0xffffffff, vw, i); + float pv = __uint_as_float((uint32_t) (pvw & 0xffffffff)); + float pw = __uint_as_float((uint32_t) (pvw >> 32)); + uint32_t vi = __float_as_uint(v); + uint32_t wi = __float_as_uint(w); + int32_t sfm = -static_cast(lane_id & i) >> 31; + vi ^= (sfm & 0x80000000); + wi ^= (sfm & 0x80000000); + v = __uint_as_float(vi) + pv; + w = __uint_as_float(wi) + pw; + } +} + +__device__ inline float shuffle_had_fx32(float v, const int lane_id) +{ + for (int i = 1; i < 32; i <<= 1) + { + float pv = __shfl_xor_sync(0xffffffff, v, i); + uint32_t* vi = reinterpret_cast(&v); + int32_t sfm = -static_cast(lane_id & i) >> 31; + *vi ^= (sfm & 0x80000000); + v = v + pv; + } + return v; +} + +__device__ inline half2 shuffle_had_h2x32(half2 v, int lane_id) +{ + for (int i = 1; i < 32; i <<= 1) + { + half2 pv = __shfl_xor_sync(0xffffffff, v, i); + uint32_t* vi = reinterpret_cast(&v); + int32_t sfm = -static_cast(lane_id & i) >> 31; + *vi ^= (sfm & 0x80008000); + v = __hadd2(v, pv); + } + return v; +} + +// Half vector, half scales + +template +inline __device__ +void had_hf_r_128_inner +( + const half* __restrict__ input_ptr, + half* __restrict__ output_ptr, + const half* __restrict__ scale, + const float r_scale +) +{ + int t = threadIdx.x & 31; + + // Load + half4 v = ((half4*) input_ptr)[t]; + + // Pre scale + if constexpr (pre_scale) + { + int i = blockIdx.y * 32 + t; + half4 scales = ((half4*) scale)[i]; + v.x = __hmul2(v.x, scales.x); + v.y = __hmul2(v.y, scales.y); + } + + // 4 element had + float v0 = __half2float(__low2half(v.x)); + float v1 = __half2float(__high2half(v.x)); + float v2 = __half2float(__low2half(v.y)); + float v3 = __half2float(__high2half(v.y)); + float s0 = v0 + v1; + float d0 = v0 - v1; + float s1 = v2 + v3; + float d1 = v2 - v3; + float h0 = s0 + s1; + float h1 = d0 + d1; + float h2 = s0 - s1; + float h3 = d0 - d1; + + // 32 element had, warp shuffle + shuffle_had_f4x32(h0, h1, h2, h3, t); + v.x = __floats2half2_rn(h0 * r_scale, h1 * r_scale); + v.y = __floats2half2_rn(h2 * r_scale, h3 * r_scale); + + // Post scale + if constexpr (post_scale) + { + int i = blockIdx.y * 32 + t; + half4 scales = ((half4*) scale)[i]; + v.x = __hmul2(v.x, scales.x); + v.y = __hmul2(v.y, scales.y); + } + + // Store + ((half4*) output_ptr)[t] = v; +} + +// Float vector, half scales + +template +inline __device__ +void had_ff_r_128_inner +( + const float* __restrict__ input_ptr, + float* __restrict__ output_ptr, + const half* __restrict__ scale, + const float r_scale +) +{ + int t = threadIdx.x & 31; + + // Load + float4 v = ((float4*) input_ptr)[t]; + + // Pre scale + if constexpr (pre_scale) + { + int i = blockIdx.y * 32 + t; + half4 scales = ((half4*) scale)[i]; + v.x *= __low2float(scales.x); + v.y *= __high2float(scales.x); + v.z *= __low2float(scales.y); + v.w *= __high2float(scales.y); + } + + // 4 element had + float v0 = v.x; + float v1 = v.y; + float v2 = v.z; + float v3 = v.w; + float s0 = v0 + v1; + float d0 = v0 - v1; + float s1 = v2 + v3; + float d1 = v2 - v3; + v.x = s0 + s1; + v.y = d0 + d1; + v.z = s0 - s1; + v.w = d0 - d1; + + // 32 element had, warp shuffle + shuffle_had_f2x32(v.x, v.y, t); + shuffle_had_f2x32(v.z, v.w, t); + v.x *= r_scale; + v.y *= r_scale; + v.z *= r_scale; + v.w *= r_scale; + + // Post scale + if constexpr (post_scale) + { + int i = blockIdx.y * 32 + t; + half4 scales = ((half4*) scale)[i]; + v.x *= __low2float(scales.x); + v.y *= __high2float(scales.x); + v.z *= __low2float(scales.y); + v.w *= __high2float(scales.y); + } + + // Store + ((float4*) output_ptr)[t] = v; +} + +// Float vector, half scales, half output + +template +inline __device__ +void had_fh_r_128_inner +( + const float* __restrict__ input_ptr, + half* __restrict__ output_ptr, + const half* __restrict__ scale, + const float r_scale +) +{ + int t = threadIdx.x & 31; + + // Load + float4 v = ((float4*) input_ptr)[t]; + + // Pre scale + if constexpr (pre_scale) + { + int i = blockIdx.y * 32 + t; + half4 scales = ((half4*) scale)[i]; + v.x *= __low2float(scales.x); + v.y *= __high2float(scales.x); + v.z *= __low2float(scales.y); + v.w *= __high2float(scales.y); + } + + // 4 element had + float v0 = v.x; + float v1 = v.y; + float v2 = v.z; + float v3 = v.w; + float s0 = v0 + v1; + float d0 = v0 - v1; + float s1 = v2 + v3; + float d1 = v2 - v3; + v.x = s0 + s1; + v.y = d0 + d1; + v.z = s0 - s1; + v.w = d0 - d1; + + // 32 element had, warp shuffle + shuffle_had_f2x32(v.x, v.y, t); + shuffle_had_f2x32(v.z, v.w, t); + v.x *= r_scale; + v.y *= r_scale; + v.z *= r_scale; + v.w *= r_scale; + + half4 o; + o.x = __floats2half2_rn(v.x, v.y); + o.y = __floats2half2_rn(v.z, v.w); + + // Post scale + if constexpr (post_scale) + { + int i = blockIdx.y * 32 + t; + half4 scales = ((half4*) scale)[i]; + o.x = __hmul2(o.x, scales.x); + o.y = __hmul2(o.y, scales.y); + } + + // Store + ((half4*) output_ptr)[t] = o; +} + +// Fused op: o <- in_had(silu(out_had(g)) * out_had(u)) + +inline __device__ +void had_hf_r_128_guad_inner +( + const half* __restrict__ input_ptr_g, + const half* __restrict__ input_ptr_u, + half* __restrict__ output_ptr, + const half* __restrict__ post_scale_g, + const half* __restrict__ post_scale_u, + const half* __restrict__ pre_scale_d, + const float r_scale, + const float act_limit, + const int act_function +) +{ + int t = threadIdx.x & 31; + + auto had = [&](half4& v) + { + // 4 element had + float v0 = __half2float(__low2half(v.x)); + float v1 = __half2float(__high2half(v.x)); + float v2 = __half2float(__low2half(v.y)); + float v3 = __half2float(__high2half(v.y)); + float s0 = v0 + v1; + float d0 = v0 - v1; + float s1 = v2 + v3; + float d1 = v2 - v3; + float h0 = s0 + s1; + float h1 = d0 + d1; + float h2 = s0 - s1; + float h3 = d0 - d1; + + // 32 element had, warp shuffle + shuffle_had_f4x32(h0, h1, h2, h3, t); + v.x = __floats2half2_rn(h0 * r_scale, h1 * r_scale); + v.y = __floats2half2_rn(h2 * r_scale, h3 * r_scale); + + return v; + }; + + auto _silu = [&](const half2& x) + { + half2 one = __float2half2_rn(1.0f); + half2 neg_x = __hneg2(x); + half2 e = h2exp(neg_x); + half2 sum = __hadd2(one, e); + half2 r = h2rcp(sum); + half2 result = __hmul2(x, r); + return result; + }; + + auto _gelu = [&](const half2& x) + { + float2 xf = __half22float2(x); + const float c = 0.797884560803f; // sqrt(2/Pi) + xf.x = 0.5f * xf.x * (1.0f + tanh_opt(c * (xf.x + 0.044715f * xf.x * xf.x * xf.x))); + xf.y = 0.5f * xf.y * (1.0f + tanh_opt(c * (xf.y + 0.044715f * xf.y * xf.y * xf.y))); + return __float22half2_rn(xf); + }; + + // Load + half4 vg = ((half4*) input_ptr_g)[t]; + half4 vu = ((half4*) input_ptr_u)[t]; + + // Hadamard + vg = had(vg); + vu = had(vu); + + // Post scale TODO: should maybe do this in float32 + int i = blockIdx.y * 32 + t; + half4 scales_g = ((half4*) post_scale_g)[i]; + half4 scales_u = ((half4*) post_scale_u)[i]; + vg.x = __hmul2(vg.x, scales_g.x); + vg.y = __hmul2(vg.y, scales_g.y); + vu.x = __hmul2(vu.x, scales_u.x); + vu.y = __hmul2(vu.y, scales_u.y); + + // Activation + switch (act_function) + { + case ACT_SILU: + vg.x = _silu(vg.x); + vg.y = _silu(vg.y); + break; + + case ACT_GELU: + vg.x = _gelu(vg.x); + vg.y = _gelu(vg.y); + break; + + default: + break; + } + + // Optional activation limits + if (act_limit != 0.0f) + { + vu.x = __hmax2(vu.x, __float2half2_rn(-act_limit)); + vu.y = __hmax2(vu.y, __float2half2_rn(-act_limit)); + vu.x = __hmin2(vu.x, __float2half2_rn(act_limit)); + vu.y = __hmin2(vu.y, __float2half2_rn(act_limit)); + vg.x = __hmin2(vg.x, __float2half2_rn(act_limit)); + vg.y = __hmin2(vg.y, __float2half2_rn(act_limit)); + } + + // Gate + vg.x = __hmul2(vg.x, vu.x); + vg.y = __hmul2(vg.y, vu.y); + + // Pre scale (d) + half4 scales_d = ((half4*) pre_scale_d)[i]; + vg.x = __hmul2(vg.x, scales_d.x); + vg.y = __hmul2(vg.y, scales_d.y); + + // Hadamard + vg = had(vg); + + // Store + ((half4*) output_ptr)[t] = vg; +} + +// Fused op: o += float(out_had(i)), atomic + +inline __device__ +void had_hf_r_128_d_inner +( + const half* __restrict__ input_ptr, + float* __restrict__ output_ptr, + const half* __restrict__ post_scale, + const float r_scale +) +{ + int t = threadIdx.x & 31; + + // Load + half4 v = ((half4*) input_ptr)[t]; + + // 4 element had + float v0 = __half2float(__low2half(v.x)); + float v1 = __half2float(__high2half(v.x)); + float v2 = __half2float(__low2half(v.y)); + float v3 = __half2float(__high2half(v.y)); + float s0 = v0 + v1; + float d0 = v0 - v1; + float s1 = v2 + v3; + float d1 = v2 - v3; + float h0 = s0 + s1; + float h1 = d0 + d1; + float h2 = s0 - s1; + float h3 = d0 - d1; + + // 32 element had, warp shuffle + shuffle_had_f4x32(h0, h1, h2, h3, t); + h0 *= r_scale; + h1 *= r_scale; + h2 *= r_scale; + h3 *= r_scale; + + // Post scale + int i = blockIdx.y * 32 + t; + half4 scales = ((half4*) post_scale)[i]; + h0 *= __low2float(scales.x); + h1 *= __high2float(scales.x); + h2 *= __low2float(scales.y); + h3 *= __high2float(scales.y); + + // Reshuffle in shmem for coalesced store with atomicAdd + extern __shared__ float temp_shared[]; + int warp_id = threadIdx.x / 32; + float* sh = temp_shared + warp_id * 128; + sh[t * 4 + 0] = h0; + sh[t * 4 + 1] = h1; + sh[t * 4 + 2] = h2; + sh[t * 4 + 3] = h3; + __syncwarp(); + atomicAdd(output_ptr + 0 + t, sh[ 0 + t]); + atomicAdd(output_ptr + 32 + t, sh[32 + t]); + atomicAdd(output_ptr + 64 + t, sh[64 + t]); + atomicAdd(output_ptr + 96 + t, sh[96 + t]); +} +// Vendored from https://github.com/brandonmmusic-max/exllamav3 at 704aefd743b390af4bd0fb429d1906f9b964c7d8. +// License: b12x/gemm/trellis_linear/csrc/vendor/LICENSE.exllamav3 diff --git a/b12x/gemm/trellis_linear/csrc/vendor/util.cuh b/b12x/gemm/trellis_linear/csrc/vendor/util.cuh new file mode 100644 index 000000000..1b7347e71 --- /dev/null +++ b/b12x/gemm/trellis_linear/csrc/vendor/util.cuh @@ -0,0 +1,104 @@ +#pragma once + +typedef struct __align__(8) half4 +{ + half2 x; + half2 y; + __device__ half4() = default; + __device__ half4(half2 x_, half2 y_) : x(x_), y(y_) {} + __device__ half4(half h0, half h1, half h2, half h3) : + x(__halves2half2(h0, h1)), + y(__halves2half2(h2, h3)) {} +} +half4; + +typedef struct __align__(8) bfloat164 +{ + __nv_bfloat162 x; + __nv_bfloat162 y; + __device__ bfloat164() = default; + __device__ bfloat164(__nv_bfloat162 x_, __nv_bfloat162 y_): x(x_), y(y_) {} + __device__ bfloat164(__nv_bfloat16 b0, __nv_bfloat16 b1, __nv_bfloat16 b2, __nv_bfloat16 b3) : + x(__halves2bfloat162(b0, b1)), + y(__halves2bfloat162(b2, b3)) {} +} +bfloat164; + +typedef struct __align__(16) half8 +{ + half2 x; + half2 y; + half2 z; + half2 w; + __device__ half8() = default; + __device__ half8(half2 x_, half2 y_, half2 z_, half2 w_) : x(x_), y(y_), z(z_), w(w_) {} + __device__ half8(half h0, half h1, half h2, half h3, half h4, half h5, half h6, half h7) : + x(__halves2half2(h0, h1)), + y(__halves2half2(h2, h3)), + z(__halves2half2(h4, h5)), + w(__halves2half2(h6, h7)) {} +} +half8; + +struct Dim3 +{ + int m; + int k; + int n; + inline __device__ int numel_a() { return m * k; } + inline __device__ int numel_b() { return k * n; } + inline __device__ int numel_c() { return m * n; } +}; + +#define READ128(__x, __y) ((uint4*)&__x)[0] = ((uint4*)(__y))[0]; +#define WRITE128(__x, __y) ((uint4*)__x)[0] = ((uint4*)(&__y))[0]; +#define READ64(__x, __y) ((uint2*)&__x)[0] = ((uint2*)(__y))[0]; +#define WRITE64(__x, __y) ((uint2*)__x)[0] = ((uint2*)(&__y))[0]; + +#define LOW_TO_FLOAT(__x) __half2float(__low2half(__x)) +#define HIGH_TO_FLOAT(__x) __half2float(__high2half(__x)) + +#define LOW_TO_FLOAT(__x) __half2float(__low2half(__x)) +#define HIGH_TO_FLOAT(__x) __half2float(__high2half(__x)) + +#define CLAMP(__x, __min, __max) fmaxf(__min, fminf(__x, __max)) +#define CLAMP_FP16(__x) CLAMP(__x, -65504.0f, 65504.0f) + +#define SWAP16(__x) __byte_perm(__x, 0, 0x1032) + +union half2_uint32 +{ + uint32_t as_uint32; + half2 as_half2; + __device__ half2_uint32(uint32_t val) : as_uint32(val) {} + __device__ half2_uint32(half2 val) : as_half2(val) {} + __device__ half2_uint32() : as_uint32(0) {} +}; + +union half_uint16 +{ + uint16_t as_uint16; + half as_half; + __device__ half_uint16(uint16_t val) : as_uint16(val) {} + __device__ half_uint16(half val) : as_half(val) {} + __device__ half_uint16() : as_uint16(0) {} +}; + +__device__ inline float fxor(float v, uint32_t mask) +{ + uint32_t* vi = reinterpret_cast(&v); + *vi ^= mask; + return v; +} + +__device__ inline half2 h2xor(half2 v, uint32_t mask) +{ + uint32_t* vi = reinterpret_cast(&v); + *vi ^= mask; + return v; +} + +#define NEG_INF_F16 __ushort_as_half(0xFC00) +#define POS_INF_F16 __ushort_as_half(0x7C00) +// Vendored from https://github.com/brandonmmusic-max/exllamav3 at 704aefd743b390af4bd0fb429d1906f9b964c7d8. +// License: b12x/gemm/trellis_linear/csrc/vendor/LICENSE.exllamav3 diff --git a/b12x/gemm/trellis_linear/csrc/vendor/util.h b/b12x/gemm/trellis_linear/csrc/vendor/util.h new file mode 100644 index 000000000..a02b6305b --- /dev/null +++ b/b12x/gemm/trellis_linear/csrc/vendor/util.h @@ -0,0 +1,140 @@ +#pragma once + +#include + +#define CEIL_DIVIDE(x, size) (((x) + (size) - 1) / (size)) +#define MIN(x, y) ((x) < (y) ? (x) : (y)) +#define MAX(x, y) ((x) > (y) ? (x) : (y)) + +// Some decluttering macros +// +// TORCH_CHECK_DTYPE(x, T): assert x is dtype T +// TORCH_CHECK_DTYPE_OPT(x, T): assert x is dtype T, unless x is None +// TORCH_CHECK_FLOAT_HALF(x): assert x is either kFloat or kHalf +// TORCH_CHECK_SHAPES(x, i, y, j, scale): assert x.size(i) == y.size(j) * scale +// TORCH_CHECK_SHAPES_OPT(x, i, y, j, scale): assert x.size(i) == y.size(j) * scale, unless x is None +// TORCH_CHECK_SHAPES_FULL(x, y): assert x and y are same shape +// TORCH_CHECK_NUMEL(x, y): assert x and y have same number of elements +// TORCH_CHECK_DIV(x, i, divisor): assert x.size(i) is divisible by divisor +// TORCH_CHECK_DIM(x, D): assert x has D dimensions +// TORCH_CHECK_DIM_OPT(x, D): assert x has D dimensions, unless x is None +// TORCH_CHECK_SIZE(x, i, s): assert x.size(i) == s +// OPTPTR(x): x.data_ptr() or nullptr if x is None + +#define TORCH_CHECK_DTYPE(__x, __dtype) TORCH_CHECK((__x).dtype() == at::__dtype, #__x " is incorrect datatype, must be " #__dtype) +#define TORCH_CHECK_DTYPE_OPT(__x, __dtype) TORCH_CHECK((!__x.has_value()) || (__x).value().dtype() == at::__dtype, #__x " is incorrect datatype, must be " #__dtype) +#define TORCH_CHECK_FLOAT_HALF(__x) TORCH_CHECK((__x).dtype() == at::kHalf || (__x).dtype() == at::kFloat, #__x " is incorrect datatype, must be kHalf or kFloat") +#define TORCH_CHECK_SHAPES(__x, __dim_x, __y, __dim_y, __scale_y) TORCH_CHECK((__x).size(__dim_x) == (__y).size(__dim_y) * __scale_y, #__x " and " #__y " have incompatible shapes") +#define TORCH_CHECK_SHAPES_OPT(__x, __dim_x, __y, __dim_y, __scale_y) TORCH_CHECK((!(__x).has_value()) || (__x).value().size(__dim_x) == (__y).size(__dim_y) * __scale_y, #__x " and " #__y " have incompatible shapes") +#define TORCH_CHECK_SHAPES_FULL(__x, __y) TORCH_CHECK((__x).sizes() == (__y).sizes(), #__x " and " #__y " have incompatible shapes") +#define TORCH_CHECK_NUMEL(__x, __y) TORCH_CHECK((__x).numel() == (__y).numel(), #__x " and " #__y " have incompatible shapes") +#define TORCH_CHECK_DIV(__x, __dim_x, __div) TORCH_CHECK((__x).size(__dim_x) % __div == 0, #__x " dimension " #__dim_x " must be divisible by " #__div) +#define TORCH_CHECK_DIM(__x, __dims) TORCH_CHECK((__x).dim() == __dims, #__x " must have " #__dims " dimensions") +#define TORCH_CHECK_DIM_OPT(__x, __dims) TORCH_CHECK((!__x.has_value()) || (__x).value().dim() == __dims, #__x " must have " #__dims " dimensions") +#define TORCH_CHECK_SIZE(__x, __dim_x, __s) TORCH_CHECK((__x).size(__dim_x) == (__s), #__x " dimension " #__dim_x " is incorrect size") +#define OPTPTR(__x) (__x.has_value() ? __x.value().data_ptr() : nullptr) + +// Debug stuff + +#define DBGS(__x) printf("%s\n", __x) +#define DBGI(__x) \ + printf("%s: %i\n", #__x, __x) +#define DBGI2(__x, __y) \ + printf("%s, %s: %i, %i\n", #__x, #__y, __x, __y) +#define DBGI3(__x, __y, __z) \ + printf("%s, %s, %s: %i, %i, %i\n", #__x, #__y, #__z, __x, __y, __z) +#define DBGI4(__x, __y, __z, __w) \ + printf("%s, %s, %s, %s: %i, %i, %i, %i\n", #__x, #__y, #__z, #__w, __x, __y, __z, __w) +#define DBGI5(__x, __y, __z, __w, __v) \ + printf("%s, %s, %s, %s, %s: %i, %i, %i, %i, %i\n", #__x, #__y, #__z, #__w, #__v, __x, __y, __z, __w, __v) +#define DBGI6(__x, __y, __z, __w, __v, __u) \ + printf("%s, %s, %s, %s, %s, %s: %i, %i, %i, %i, %i, %i\n", #__x, #__y, #__z, #__w, #__v, #__u, __x, __y, __z, __w, __v, __u) +#define DBGI7(__x, __y, __z, __w, __v, __u, __t) \ + printf("%s, %s, %s, %s, %s, %s, %s: %i, %i, %i, %i, %i, %i, %i\n", #__x, #__y, #__z, #__w, #__v, #__u, #__t, __x, __y, __z, __w, __v, __u, __t) +#define DBGI8(__x, __y, __z, __w, __v, __u, __t, __s) \ + printf("%s, %s, %s, %s, %s, %s, %s, %s: %i, %i, %i, %i, %i, %i, %i, %i\n", #__x, #__y, #__z, #__w, #__v, #__u, #__t, #__s, __x, __y, __z, __w, __v, __u, __t, __s) +#define DBGI9(__x, __y, __z, __w, __v, __u, __t, __s, __r) \ + printf("%s, %s, %s, %s, %s, %s, %s, %s, %s: %i, %i, %i, %i, %i, %i, %i, %i, %i\n", #__x, #__y, #__z, #__w, #__v, #__u, #__t, #__s, #__r, __x, __y, __z, __w, __v, __u, __t, __s, __r) +#define DBGI10(__x, __y, __z, __w, __v, __u, __t, __s, __r, __q) \ + printf("%s, %s, %s, %s, %s, %s, %s, %s, %s, %s: %i, %i, %i, %i, %i, %i, %i, %i, %i, %i\n", #__x, #__y, #__z, #__w, #__v, #__u, #__t, #__s, #__r, #__q, __x, __y, __z, __w, __v, __u, __t, __s, __r, __q) +#define DBGX(__x) printf("%s: %x\n", #__x, __x) +#define DBGX2(__x, __y) printf("%s, %s: %x, %x\n", #__x, #__y, __x, __y) +#define DBGX3(__x, __y, __z) printf("%s, %s, %s: %x, %x, %x\n", #__x, #__y, #__z, __x, __y, __z) +#define DBGIX(__x, __y) printf("%s, %s: %i, %x\n", #__x, #__y, __x, __y) +#define DBGIX2(__x, __y, __z) printf("%s, %s, %s: %i, %x, %x\n", #__x, #__y, #__z, __x, __y, __z) +#define DBGIF(__x, __y) printf("%s, %s: %i, %f\n", #__x, #__y, __x, __y) +#define DBGIF2(__x, __y, __z) printf("%s, %s, %s: %i, %f, %f\n", #__x, #__y, #__z, __x, __y, __z) +#define DBGF(__x) printf("%s: %f\n", #__x, __x) +#define DBGF2(__x, __y) printf("%s, %s: %f, %f\n", #__x, #__y, __x, __y) +#define DBGF3(__x, __y, __z) printf("%s, %s, %s: %f, %f, %f\n", #__x, #__y, #__z, __x, __y, __z) +#define DBGF4(__x, __y, __z, __w) printf("%s, %s, %s, %s: %f, %f, %f, %f\n", #__x, #__y, #__z, #__w, __x, __y, __z, __w) +#define DBGH(__x) printf("%s: %f\n", #__x, __half2float(__x)) +#define DBGH2(__x, __y) printf("%s, %s: %f, %f\n", #__x, #__y, __half2float(__x), __half2float(__y)) +#define DBGH3(__x, __y, __z) printf("%s, %s, %s: %f, %f, %f\n", #__x, #__y, #__z, __half2float(__x), __half2float(__y), __half2float(__z)) +#define DBGIH(__x, __y) printf("%s, %s: %i, %f\n", #__x, #__y, __x, __half2float(__y)) +#define DBGIH2(__x, __y, __z) printf("%s, %s, %s: %i, %f, %f\n", #__x, #__y, #__z, __x, __half2float(__y), __half2float(__z)) +#define DBGI2H2(__x, __y, __z, __w) printf("%s, %s, %s, %s: %i, %i, %f, %f\n", #__x, #__y, #__z, #__w, __x, __y, __half2float(__z), __half2float(__w)) +#define DBGIH3(__x, __y, __z, __w) printf("%s, %s, %s, %s: %i, %f, %f, %f\n", #__x, #__y, #__z, #__w, __x, __half2float(__y), __half2float(__z), __half2float(__w)) +#define DBGIH4(__x, __y, __z, __w, __v) printf("%s, %s, %s, %s, %s: %i, %f, %f, %f, %f\n", #__x, #__y, #__z, #__w, #__v, __x, __half2float(__y), __half2float(__z), __half2float(__w), __half2float(__v)) +#define DBGA(__x) printf("%s: %016llx\n", #__x, __x) +#define DBGIA(__x, __y) printf("%s, %s: %i, %016llx\n", #__x, #__y, __x, __y) +#define DBGI2A(__x, __y, __z) printf("%s, %s, %s: %i, %i, %016llx\n", #__x, #__y, #__z, __x, __y, __z) + +#define TIME_START \ + auto start = std::chrono::high_resolution_clock::now() + +#define TIME_STOP \ + do { \ + auto stop = std::chrono::high_resolution_clock::now(); \ + auto duration_us = std::chrono::duration_cast(stop - start); \ + DBGI(duration_us); \ + } while (false) + +/* +Compile-time for loop. Supports template instancing. Example usage: + +int kernel_arg = select_kernel_somehow(); + +// Not nice +if (kernel_arg == 2) + launch_kernel_instance<2><<< ... >>>( ... ) +if (kernel_arg == 3) + launch_kernel_instance<3><<< ... >>>( ... ) +if (kernel_arg == 4) + launch_kernel_instance<4><<< ... >>>( ... ) +if (kernel_arg == 6) + launch_kernel_instance<6><<< ... >>>( ... ) +if (kernel_arg == 8) + launch_kernel_instance<8><<< ... >>>( ... ) + +// Nice? +static_for_pack<2, 3, 4, 6, 8>([&](auto ic) +{ + constexpr int i = decltype(ic)::value; + if (kernel_arg == i) + launch_kernel_instance<<< ... >>>( ... ) +}); + +// Ultimately much cleaner +#define __(i, j) quant_cache_paged_kernel +constexpr auto quant_cache_paged_kernel_instances = std::array +{ + std::array{ __(2, 2), __(2, 3), __(2, 4), __(2, 5), __(2, 6), __(2, 7), __(2, 8) }, + std::array{ __(3, 2), __(3, 3), __(3, 4), __(3, 5), __(3, 6), __(3, 7), __(3, 8) }, + std::array{ __(4, 2), __(4, 3), __(4, 4), __(4, 5), __(4, 6), __(4, 7), __(4, 8) }, + std::array{ __(5, 2), __(5, 3), __(5, 4), __(5, 5), __(5, 6), __(5, 7), __(5, 8) }, + std::array{ __(6, 2), __(6, 3), __(6, 4), __(6, 5), __(6, 6), __(6, 7), __(6, 8) }, + std::array{ __(7, 2), __(7, 3), __(7, 4), __(7, 5), __(7, 6), __(7, 7), __(7, 8) }, + std::array{ __(8, 2), __(8, 3), __(8, 4), __(8, 5), __(8, 6), __(8, 7), __(8, 8) } +}; +#undef __ +*/ + +// This breaks with nesting on VC++ older than 17.13 (late 2024 preview) +template +constexpr void static_for_pack(F&& f) +{ + (f(std::integral_constant{}), ...); +} +// Vendored from https://github.com/brandonmmusic-max/exllamav3 at 704aefd743b390af4bd0fb429d1906f9b964c7d8. +// License: b12x/gemm/trellis_linear/csrc/vendor/LICENSE.exllamav3 diff --git a/b12x/moe/__init__.py b/b12x/moe/__init__.py index 620697147..1ca2cbca0 100644 --- a/b12x/moe/__init__.py +++ b/b12x/moe/__init__.py @@ -1,8 +1,9 @@ """MoE ops for b12x. - ``fused_moe``: fused tensor-parallel routed-expert FFN (route -> FC1 -> - activation -> FC2 -> scatter); recipes nvfp4/w4a8_mx/w4a8_nvfp4/w6a8_mx/ - w4a16, including the TP-independent ``btx`` trellis checkpoint container. + activation -> FC2 -> scatter); recipes nvfp4/mxfp4/w4a8_mx/w4a8_nvfp4/w4a16, + including legacy ``exl3_trellis_mcg`` and TP-independent + ``qsrt_sqg_e4m3`` plus uniform-K5/K6 ``sqg_fp16_d3l`` W4A16 source formats. - ``ep_moe``: expert-parallel MoE (replicated input -> local partial; cross-rank reduction is the caller's job, typically ``comm.pcie``). """ diff --git a/b12x/moe/_shared/btx_schema.py b/b12x/moe/_shared/btx_schema.py deleted file mode 100644 index afa2ac16a..000000000 --- a/b12x/moe/_shared/btx_schema.py +++ /dev/null @@ -1,490 +0,0 @@ -"""BTX (b12x trellis exchange) checkpoint container schema. - -BTX is the TP-shard-independent checkpoint container for trellis-coded MoE -expert weights. Storage is organized around 32-channel *atom slots* on the -intermediate axis: every slot row holds all experts' code words for those -channels, so a tensor-parallel rank loads a contiguous slot range and -nothing else. All behavior is declared in a manifest — codebook, rate -structure, coupled-Hadamard transform, and geometry — and the reader -derives byte addressing purely from those declarations. - -Storage schema id: ``btx-atoms-v1``. A checkpoint directory contains -``btx-manifest.json`` plus one ``btx-layer-.safetensors`` file per -MoE layer. Per-layer tensors: - -- ``atoms``: u8 ``[atom_slots, row_stride]`` — trellis code words only, - expert-id-major bundles per row, zero padding to the row stride. -- ``rotations``: fp16 ``[atom_slots, num_experts, 3, atom_channels]`` — - per-channel intermediate-boundary values for gate/up/down in physical - channel order. -- ``rates_fc1``/``rates_fc2``: u8 ``[atom_slots/8, num_experts]`` — one - rate byte per (256-channel pair, expert); present iff the rate structure - is ``per_expert_pair``. -- ``gate_suh``/``up_suh``/``down_svh``: fp16 ``[hidden_size]`` or - ``[num_experts, hidden_size]`` — hidden-axis incoherence values. -- ``rotation_draws``: u8 ``[num_experts]`` in ``0..7`` — present iff the - coupled-Hadamard transform is declared. - -A rate byte is ``(low_bits << 4) | high_bits`` and is exactly the fused -kernel's pair-kind vocabulary expressed as data. Uniform checkpoints carry -no rate tables; their single bitrate is declared in the manifest. - -Within one ``atoms`` row, expert bundles are concatenated in expert-id -order; each bundle is gate ‖ up ‖ down and each matrix section stores its -low-record plane followed by its high-record plane (``[H/16][16*low]i16`` -‖ ``[H/16][16*high]i16``). Under a uniform rate structure the two planes -are the atom's two consecutive N16 (FC1) or K16 (FC2) tiles. - -This module is torch-free: manifest parsing, fail-closed validation, extent -legality, and byte arithmetic. Tensor I/O and preparation live with the -W4A16 kernel host code. -""" - -from __future__ import annotations - -from dataclasses import dataclass - -from .trellis_codebooks import ( - CODEBOOKS, - MCG, - MCG_MULTIPLIER, - validate_codebook_bits, -) - -BTX_SCHEMA = "btx-atoms-v1" -BTX_MANIFEST_KIND = "btx-manifest" -BTX_MANIFEST_FILENAME = "btx-manifest.json" - -ATOM_CHANNELS = 32 -ATOMS_PER_PAIR = 8 - -RATE_STRUCTURE_UNIFORM = "uniform" -RATE_STRUCTURE_PER_EXPERT_PAIR = "per_expert_pair" - -# The fused kernel's pair-kind vocabulary as rate bytes. -RATE_CODE_PAIR_KINDS: dict[int, str] = { - 0x22: "P22", - 0x33: "P33", - 0x24: "P24", - 0x43: "P43", - 0x44: "P44", -} -PAIR_KIND_RATE_CODES: dict[str, int] = { - kind: code for code, kind in RATE_CODE_PAIR_KINDS.items() -} - - -def layer_filename(layer_index: int) -> str: - return f"btx-layer-{int(layer_index):05d}.safetensors" - - -def rate_code(low_bits: int, high_bits: int) -> int: - return (int(low_bits) << 4) | int(high_bits) - - -def rate_code_bits(code: int) -> tuple[int, int]: - return (int(code) >> 4) & 0xF, int(code) & 0xF - - -def matrix_atom_bytes(hidden_size: int, low_bits: int, high_bits: int) -> int: - """Trellis bytes one atom contributes to one expert matrix. - - An atom holds two 16-channel record planes; each plane stores - ``hidden_size/16`` tiles of ``16*bits`` int16 words. - """ - - return (int(hidden_size) // 16) * 32 * (int(low_bits) + int(high_bits)) - - -def bundle_bytes( - hidden_size: int, fc1_code: int, fc2_code: int -) -> int: - """Per-(expert, atom) bundle size: gate ‖ up ‖ down trellis words.""" - - fc1_low, fc1_high = rate_code_bits(fc1_code) - fc2_low, fc2_high = rate_code_bits(fc2_code) - return 2 * matrix_atom_bytes(hidden_size, fc1_low, fc1_high) + ( - matrix_atom_bytes(hidden_size, fc2_low, fc2_high) - ) - - -def _require(condition: bool, message: str) -> None: - if not condition: - raise ValueError(message) - - -def _require_keys( - mapping: dict, *, required: set[str], optional: set[str], where: str -) -> None: - _require(isinstance(mapping, dict), f"{where} must be a JSON object") - keys = set(mapping.keys()) - unknown = keys - required - optional - _require(not unknown, f"{where} has unknown keys {sorted(unknown)}") - missing = required - keys - _require(not missing, f"{where} is missing keys {sorted(missing)}") - - -@dataclass(frozen=True) -class BtxGeometry: - num_experts: int - hidden_size: int - intermediate_size: int - atom_channels: int - atom_slots: int - moe_layer_indices: tuple[int, ...] - - -@dataclass(frozen=True) -class BtxRates: - structure: str - bits: int | None - pair_kinds: frozenset[str] | None - - def uniform_code(self) -> int | None: - if self.structure != RATE_STRUCTURE_UNIFORM: - return None - assert self.bits is not None - return rate_code(self.bits, self.bits) - - -@dataclass(frozen=True) -class BtxHadamard: - coupled: bool - pre_block: int | None - post_block: int | None - per_expert_input_rotations: bool - - -@dataclass(frozen=True) -class BtxLayout: - atom_row_alignment: int - extent_alignment_slots: int - extent_barriers: tuple[int, ...] - - -@dataclass(frozen=True) -class BtxLayerRef: - file: str - sha256: str - - -@dataclass(frozen=True) -class BtxManifest: - codebook: str - codebook_seed: int | None - geometry: BtxGeometry - rates: BtxRates - hadamard: BtxHadamard - layout: BtxLayout - layers: dict[int, BtxLayerRef] - - @staticmethod - def from_dict(data: dict) -> "BtxManifest": - _require_keys( - data, - required={ - "kind", - "schema", - "codebook", - "geometry", - "rates", - "hadamard", - "layout", - "layers", - }, - optional={"codebook_seed"}, - where="BTX manifest", - ) - _require( - data["kind"] == BTX_MANIFEST_KIND, - f"BTX manifest kind must be {BTX_MANIFEST_KIND!r}, " - f"got {data['kind']!r}", - ) - _require( - data["schema"] == BTX_SCHEMA, - f"BTX manifest schema must be {BTX_SCHEMA!r}, got {data['schema']!r}", - ) - - codebook = data["codebook"] - _require( - codebook in CODEBOOKS, - f"BTX codebook must be one of {sorted(CODEBOOKS)}, got {codebook!r}", - ) - seed = data.get("codebook_seed") - if codebook == MCG: - _require( - isinstance(seed, int) and seed == MCG_MULTIPLIER, - "BTX mcg checkpoints must declare codebook_seed " - f"{MCG_MULTIPLIER:#010x}", - ) - else: - _require( - seed is None, - f"BTX codebook_seed is valid only for mcg, not {codebook!r}", - ) - - geometry = _parse_geometry(data["geometry"]) - rates = _parse_rates(data["rates"], codebook=codebook) - hadamard = _parse_hadamard(data["hadamard"], geometry=geometry) - layout = _parse_layout(data["layout"], geometry=geometry) - layers = _parse_layers(data["layers"], geometry=geometry) - return BtxManifest( - codebook=codebook, - codebook_seed=seed, - geometry=geometry, - rates=rates, - hadamard=hadamard, - layout=layout, - layers=layers, - ) - - def validate_extent(self, first_slot: int, slot_count: int) -> None: - """Reject rank extents the layout declarations make illegal.""" - - alignment = self.layout.extent_alignment_slots - slots = self.geometry.atom_slots - _require( - slot_count > 0 and first_slot >= 0, - f"BTX extent [{first_slot}, {first_slot + slot_count}) is empty " - "or negative", - ) - _require( - first_slot + slot_count <= slots, - f"BTX extent [{first_slot}, {first_slot + slot_count}) exceeds " - f"{slots} atom slots", - ) - _require( - first_slot % alignment == 0 and slot_count % alignment == 0, - f"BTX extent [{first_slot}, {first_slot + slot_count}) must " - f"align to {alignment} slots", - ) - for barrier in self.layout.extent_barriers: - _require( - not (first_slot < barrier < first_slot + slot_count), - f"BTX extent [{first_slot}, {first_slot + slot_count}) " - f"crosses the declared barrier at slot {barrier}", - ) - - -def _parse_geometry(data: dict) -> BtxGeometry: - _require_keys( - data, - required={ - "num_experts", - "hidden_size", - "intermediate_size", - "atom_channels", - "atom_slots", - "moe_layer_indices", - }, - optional=set(), - where="BTX geometry", - ) - for name in ( - "num_experts", - "hidden_size", - "intermediate_size", - "atom_channels", - "atom_slots", - ): - _require( - isinstance(data[name], int) and data[name] > 0, - f"BTX geometry {name} must be a positive integer", - ) - _require( - data["atom_channels"] == ATOM_CHANNELS, - f"BTX atoms hold {ATOM_CHANNELS} channels; got {data['atom_channels']}", - ) - _require( - data["intermediate_size"] % data["atom_channels"] == 0, - "BTX intermediate_size must be a multiple of atom_channels", - ) - _require( - data["atom_slots"] * data["atom_channels"] == data["intermediate_size"], - "BTX atom_slots must equal intermediate_size / atom_channels", - ) - _require( - data["hidden_size"] % 16 == 0, - "BTX hidden_size must be a multiple of 16", - ) - indices = data["moe_layer_indices"] - _require( - isinstance(indices, list) - and len(indices) > 0 - and all(isinstance(i, int) and i >= 0 for i in indices) - and len(set(indices)) == len(indices), - "BTX moe_layer_indices must be distinct non-negative integers", - ) - return BtxGeometry( - num_experts=data["num_experts"], - hidden_size=data["hidden_size"], - intermediate_size=data["intermediate_size"], - atom_channels=data["atom_channels"], - atom_slots=data["atom_slots"], - moe_layer_indices=tuple(sorted(indices)), - ) - - -def _parse_rates(data: dict, *, codebook: str) -> BtxRates: - _require_keys( - data, - required={"structure"}, - optional={"bits", "pair_kinds"}, - where="BTX rates", - ) - structure = data["structure"] - if structure == RATE_STRUCTURE_UNIFORM: - _require( - "bits" in data and "pair_kinds" not in data, - "uniform BTX rates declare bits and no pair_kinds", - ) - bits = data["bits"] - _require( - isinstance(bits, int) and bits in (2, 3, 4, 5, 6), - f"BTX uniform bits must be one of 2..6, got {bits!r}", - ) - validate_codebook_bits(codebook, bits) - return BtxRates(structure=structure, bits=bits, pair_kinds=None) - if structure == RATE_STRUCTURE_PER_EXPERT_PAIR: - _require( - "pair_kinds" in data and "bits" not in data, - "per_expert_pair BTX rates declare pair_kinds and no bits", - ) - kinds = data["pair_kinds"] - _require( - isinstance(kinds, list) - and len(kinds) > 0 - and all(kind in PAIR_KIND_RATE_CODES for kind in kinds) - and len(set(kinds)) == len(kinds), - "BTX pair_kinds must be distinct members of " - f"{sorted(PAIR_KIND_RATE_CODES)}", - ) - for kind in kinds: - for bits in rate_code_bits(PAIR_KIND_RATE_CODES[kind]): - validate_codebook_bits(codebook, bits) - return BtxRates( - structure=structure, bits=None, pair_kinds=frozenset(kinds) - ) - raise ValueError( - "BTX rates structure must be 'uniform' or 'per_expert_pair', " - f"got {structure!r}" - ) - - -def _parse_hadamard(data: dict, *, geometry: BtxGeometry) -> BtxHadamard: - _require_keys( - data, - required={"coupled", "per_expert_input_rotations"}, - optional={"pre_block", "post_block"}, - where="BTX hadamard", - ) - coupled = data["coupled"] - _require( - isinstance(coupled, bool), "BTX hadamard coupled must be a boolean" - ) - per_expert = data["per_expert_input_rotations"] - _require( - isinstance(per_expert, bool), - "BTX per_expert_input_rotations must be a boolean", - ) - if not coupled: - _require( - "pre_block" not in data and "post_block" not in data, - "BTX hadamard blocks are valid only for coupled checkpoints", - ) - return BtxHadamard( - coupled=False, - pre_block=None, - post_block=None, - per_expert_input_rotations=per_expert, - ) - _require( - "pre_block" in data and "post_block" in data, - "coupled BTX checkpoints must declare pre_block and post_block", - ) - pre_block, post_block = data["pre_block"], data["post_block"] - for name, value in (("pre_block", pre_block), ("post_block", post_block)): - _require( - isinstance(value, int) and value > 0 and value % ATOM_CHANNELS == 0, - f"BTX hadamard {name} must be a positive multiple of " - f"{ATOM_CHANNELS}", - ) - _require( - geometry.intermediate_size % post_block == 0, - "BTX intermediate_size must be a multiple of post_block", - ) - _require( - geometry.hidden_size % pre_block == 0, - "BTX hidden_size must be a multiple of pre_block", - ) - return BtxHadamard( - coupled=True, - pre_block=pre_block, - post_block=post_block, - per_expert_input_rotations=per_expert, - ) - - -def _parse_layout(data: dict, *, geometry: BtxGeometry) -> BtxLayout: - _require_keys( - data, - required={"atom_row_alignment", "extent_alignment_slots"}, - optional={"extent_barriers"}, - where="BTX layout", - ) - alignment = data["atom_row_alignment"] - _require( - isinstance(alignment, int) and alignment > 0, - "BTX atom_row_alignment must be a positive integer", - ) - extent_alignment = data["extent_alignment_slots"] - _require( - isinstance(extent_alignment, int) - and extent_alignment > 0 - and geometry.atom_slots % extent_alignment == 0, - "BTX extent_alignment_slots must be a positive divisor of atom_slots", - ) - barriers = data.get("extent_barriers", []) - _require( - isinstance(barriers, list) - and all( - isinstance(b, int) and 0 < b < geometry.atom_slots - for b in barriers - ) - and len(set(barriers)) == len(barriers), - "BTX extent_barriers must be distinct interior slot indices", - ) - return BtxLayout( - atom_row_alignment=alignment, - extent_alignment_slots=extent_alignment, - extent_barriers=tuple(sorted(barriers)), - ) - - -def _parse_layers(data: dict, *, geometry: BtxGeometry) -> dict[int, BtxLayerRef]: - _require(isinstance(data, dict) and data, "BTX layers must be non-empty") - layers: dict[int, BtxLayerRef] = {} - for key, value in data.items(): - _require( - isinstance(key, str) and key.isdigit(), - f"BTX layer keys must be decimal strings, got {key!r}", - ) - index = int(key) - _require_keys( - value, - required={"file", "sha256"}, - optional=set(), - where=f"BTX layer {index}", - ) - _require( - isinstance(value["file"], str) - and isinstance(value["sha256"], str) - and len(value["sha256"]) == 64, - f"BTX layer {index} must declare file and hex sha256", - ) - layers[index] = BtxLayerRef(file=value["file"], sha256=value["sha256"]) - _require( - set(layers.keys()) == set(geometry.moe_layer_indices), - "BTX layers must cover exactly geometry.moe_layer_indices", - ) - return layers diff --git a/b12x/moe/_shared/execution.py b/b12x/moe/_shared/execution.py index d3dc65cd1..17b91c2ca 100644 --- a/b12x/moe/_shared/execution.py +++ b/b12x/moe/_shared/execution.py @@ -29,12 +29,6 @@ from dataclasses import dataclass from enum import Enum -from .trellis_codebooks import ( - CODEBOOKS as _TRELLIS_CODEBOOKS, - SQG_E4M3 as _TRELLIS_SQG_E4M3, - validate_codebook_bits as _validate_trellis_codebook_bits, -) - class _StringEnum(str, Enum): def __str__(self) -> str: @@ -160,20 +154,30 @@ class OutputReduction(_StringEnum): "fp4_e8m0_k32", "compressed_tensors", "mxfp6_e2m3", - "btx", + "exl3_trellis_mcg", + "qsrt_sqg_e4m3", + "sqg_fp16_d3l", } -_TRELLIS_SOURCE_FORMATS = frozenset({"btx"}) +_TRELLIS_SOURCE_FORMATS = frozenset( + {"exl3_trellis_mcg", "qsrt_sqg_e4m3", "sqg_fp16_d3l"} +) +_QSRT_ATOMS_V2_PROFILE_H308 = "k3x22_k4x2" +_QSRT_ATOMS_V2_PROFILE_COUPLED_K2 = "k2_coupled_h512_h128" +_QSRT_ATOMS_V2_PROFILE_COUPLED_H308 = "k3x22_k4x2_coupled_h512_h128" _SOURCES_BY_QUANT_MODE = { "nvfp4": frozenset({"modelopt_nvfp4"}), "w4a8_nvfp4": frozenset({"modelopt_nvfp4"}), - "w4a8_mx": frozenset({"fp4_e8m0_k32", "btx"}), - # The packed MX-FP6 source is exclusive to the w6a8_mx recipe. + "w4a8_mx": frozenset({"fp4_e8m0_k32", "qsrt_sqg_e4m3"}), + # W4A16 deliberately keeps its historical FP4 source trio; the packed + # MX-FP6 source is exclusive to the w6a8_mx recipe. "w4a16": frozenset( { "modelopt_nvfp4", "fp4_e8m0_k32", "compressed_tensors", - "btx", + "exl3_trellis_mcg", + "qsrt_sqg_e4m3", + "sqg_fp16_d3l", } ), "w6a8_mx": frozenset({"mxfp6_e2m3"}), @@ -278,13 +282,9 @@ class MoEWeightPreparationPlan: storage_policy: WeightStoragePolicy trellis_bits: int | None = None trellis_tile_config: tuple[int, int, int, int] | None = None + qsrt_storage_format: str | None = None + qsrt_profile: str | None = None coupled_hadamard: bool = False - # BTX declarations: the manifest-declared codebook, rate structure, - # pair-kind summary, and coupled-Hadamard block widths. - trellis_codebook: str | None = None - trellis_rate_structure: str | None = None - trellis_pair_kinds: frozenset[str] | None = None - coupled_hadamard_blocks: tuple[int, int] | None = None def __post_init__(self) -> None: specs = tuple(self.specs) @@ -315,21 +315,17 @@ def __post_init__(self) -> None: object.__setattr__(self, "coupled_hadamard", bool(self.coupled_hadamard)) if self.source_format in _TRELLIS_SOURCE_FORMATS: bits = 3 if self.trellis_bits is None else int(self.trellis_bits) - codebook = ( - None - if self.trellis_codebook is None - else str(self.trellis_codebook).lower() + valid_bits = ( + (2, 3, 4) + if self.source_format == "qsrt_sqg_e4m3" + else (5, 6) + if self.source_format == "sqg_fp16_d3l" + else (3, 4, 5, 6) ) - if codebook not in _TRELLIS_CODEBOOKS: + if bits not in valid_bits: raise ValueError( - "btx weights require trellis_codebook in " - f"{sorted(_TRELLIS_CODEBOOKS)}; got " - f"{self.trellis_codebook!r}" + f"trellis_bits must be one of {valid_bits}; got {bits}" ) - _validate_trellis_codebook_bits(codebook, bits) - if codebook == "mcg" and bits == 2: - raise ValueError("mcg btx weights require trellis_bits>=3") - object.__setattr__(self, "trellis_codebook", codebook) tile_config = self.trellis_tile_config or (64, 256, 64, 256) tile_config = tuple(int(value) for value in tile_config) if len(tile_config) != 4 or any( @@ -340,97 +336,93 @@ def __post_init__(self) -> None: ) object.__setattr__(self, "trellis_bits", bits) object.__setattr__(self, "trellis_tile_config", tile_config) - structure = ( - "uniform" - if self.trellis_rate_structure is None - else str(self.trellis_rate_structure).lower() - ) - pair_kinds = ( - None - if self.trellis_pair_kinds is None - else frozenset( - str(kind).upper() for kind in self.trellis_pair_kinds - ) - ) - if structure == "uniform": - if pair_kinds is not None: - raise ValueError( - "uniform btx rates declare no trellis_pair_kinds" - ) - elif structure == "per_expert_pair": - if self.coupled_hadamard: - raise ValueError( - "coupled-Hadamard btx execution is qualified " - "only for uniform rate structures" - ) - if bits != 3: - raise ValueError( - "per-expert-pair btx rates require the " - "trellis_bits=3 base specialization" - ) - if pair_kinds is None: - raise ValueError( - "per-expert-pair btx rates declare " - "trellis_pair_kinds" - ) - if "P44" in pair_kinds: - raise ValueError( - "btx pair-kind sets containing P44 (whole-expert" - " K4 tiers) have no fused execution arm; use" - " mixed-tier or multi-launch execution" - ) - if pair_kinds not in ( - frozenset({"P33"}), - frozenset({"P33", "P24"}), - frozenset({"P33", "P43"}), - ): - raise ValueError( - "btx pair-kind sets must be {P33}, {P33,P24}, " - f"or {{P33,P43}}; got {sorted(pair_kinds)}" - ) - else: - raise ValueError( - "trellis_rate_structure must be 'uniform' or " - f"'per_expert_pair'; got {structure!r}" - ) - blocks = ( + storage_format = ( None - if self.coupled_hadamard_blocks is None - else tuple( - int(value) for value in self.coupled_hadamard_blocks - ) + if self.qsrt_storage_format is None + else str(self.qsrt_storage_format).lower() ) - if self.coupled_hadamard: - if self.trellis_codebook != _TRELLIS_SQG_E4M3: + if self.source_format == "qsrt_sqg_e4m3": + if storage_format not in { + "qsrt_atoms_v1", + "qsrt_atoms_v2", + }: raise ValueError( - "coupled-Hadamard btx execution is qualified only" - " for the sqg_e4m3 codebook" + "QSRT MoE weights require " + "qsrt_storage_format='qsrt_atoms_v1', " + "or 'qsrt_atoms_v2'" ) - if blocks is None: - blocks = (512, 128) - if blocks != (512, 128): + profile = self.qsrt_profile + if storage_format == "qsrt_atoms_v2": + profile = profile or _QSRT_ATOMS_V2_PROFILE_H308 + if profile not in { + _QSRT_ATOMS_V2_PROFILE_H308, + _QSRT_ATOMS_V2_PROFILE_COUPLED_K2, + _QSRT_ATOMS_V2_PROFILE_COUPLED_H308, + }: + raise ValueError( + f"unsupported QSRT atoms-v2 profile {profile!r}" + ) + elif profile is not None: raise ValueError( - "coupled-Hadamard btx execution implements " - f"blocks (512, 128); got {blocks}" + "qsrt_profile is valid only for qsrt_atoms_v2 storage" ) - elif blocks is not None: + if profile == _QSRT_ATOMS_V2_PROFILE_COUPLED_K2: + if bits != 2: + raise ValueError( + "the coupled pure-K2 profile requires trellis_bits=2" + ) + if self.intermediate_size % 128: + raise ValueError( + "the coupled pure-K2 profile requires a local intermediate " + "size divisible by 128" + ) + if tile_config != (128, 128, 128, 128): + raise ValueError( + "the coupled pure-K2 profile requires tile_config=" + "(128, 128, 128, 128)" + ) + if not self.coupled_hadamard: + raise ValueError( + "the coupled pure-K2 profile requires " + "coupled_hadamard=True" + ) + else: + if bits != 3: + raise ValueError( + "the fixed high-rate QSRT profile requires trellis_bits=3" + ) + if self.intermediate_size != 256: + raise ValueError( + "the fixed high-rate QSRT pair kernel requires " + "intermediate_size=256" + ) + if tile_config[1] != 256: + raise ValueError( + "the fixed high-rate QSRT pair kernel requires " + "FC1 tile_n=256" + ) + if ( + profile == _QSRT_ATOMS_V2_PROFILE_COUPLED_H308 + ) != self.coupled_hadamard: + raise ValueError( + "the fixed high-rate coupled profile and " + "coupled_hadamard flag must agree" + ) + object.__setattr__(self, "qsrt_profile", profile) + elif storage_format is not None: raise ValueError( - "coupled_hadamard_blocks require coupled_hadamard" + "qsrt_storage_format is valid only for qsrt_sqg_e4m3" ) - object.__setattr__(self, "trellis_rate_structure", structure) - object.__setattr__(self, "trellis_pair_kinds", pair_kinds) - object.__setattr__(self, "coupled_hadamard_blocks", blocks) + object.__setattr__(self, "qsrt_storage_format", storage_format) elif ( self.trellis_bits is not None or self.trellis_tile_config is not None + or self.qsrt_storage_format is not None + or self.qsrt_profile is not None or self.coupled_hadamard - or self.trellis_codebook is not None - or self.trellis_rate_structure is not None - or self.trellis_pair_kinds is not None - or self.coupled_hadamard_blocks is not None ): raise ValueError( - "trellis storage settings require a trellis source format" + "QSRT storage settings require a QSRT source format" ) for name in ("num_experts", "hidden_size", "intermediate_size"): if getattr(self, name) <= 0: @@ -503,7 +495,7 @@ def reuses_source_storage(self) -> bool: @property def w4a16_weight_layout(self) -> str | None: if WeightPreparationTransform.W4A16_TRELLIS in self.transforms: - return "trellis_t256" + return "trellis3_t256" if WeightPreparationTransform.W4A16_NATIVE in self.transforms: return "modelopt" if WeightPreparationTransform.W4A16_PACKED in self.transforms: @@ -515,7 +507,7 @@ def w4a8_weight_layout(self) -> str | None: """Kernel weight-layout string for the w4a8_mx recipe, if planned.""" if WeightPreparationTransform.W4A8_TRELLIS in self.transforms: - return "trellis_t256" + return "trellis3_t256" if WeightPreparationTransform.W4A8_QMMA in self.transforms: return "qmma_repacked" return None @@ -724,11 +716,9 @@ def plan_moe_weight_preparation( w4a16_layout: PreparedWeightLayout | str | None = None, trellis_bits: int | None = None, trellis_tile_config: tuple[int, int, int, int] | None = None, + qsrt_storage_format: str | None = None, + qsrt_profile: str | None = None, coupled_hadamard: bool | None = None, - trellis_codebook: str | None = None, - trellis_rate_structure: str | None = None, - trellis_pair_kinds: Iterable[str] | None = None, - coupled_hadamard_blocks: tuple[int, int] | None = None, ) -> MoEWeightPreparationPlan: """Choose the minimal representation set for the requested recipes. @@ -775,7 +765,7 @@ def plan_moe_weight_preparation( and source_format not in _TRELLIS_SOURCE_FORMATS ): raise ValueError( - "trellis_native layout requires a trellis source format" + "trellis_native layout requires an EXL3 trellis source format" ) transforms: set[WeightPreparationTransform] = set() @@ -806,17 +796,15 @@ def plan_moe_weight_preparation( # after in-kernel decode; identity weight block scales). # The w4a8 trellis kernels window K and I in 16s and run # the 128-wide activation boundary per intermediate chunk. - if ( - trellis_codebook is not None - and str(trellis_codebook).lower() != "sqg_e4m3" - ): + if source_format != "qsrt_sqg_e4m3": raise ValueError( - "W4A8-MX btx execution requires the sqg_e4m3 codebook" + "W4A8-MX trellis execution requires the " + "qsrt_sqg_e4m3 source format" ) - if ( - trellis_rate_structure is not None - and str(trellis_rate_structure).lower() != "uniform" - ): + if qsrt_profile in { + _QSRT_ATOMS_V2_PROFILE_H308, + _QSRT_ATOMS_V2_PROFILE_COUPLED_H308, + }: # The W4A8 kernels decode one compile-time rate for the # whole payload; the mixed-rate pair machinery is # W4A16-only. @@ -956,7 +944,10 @@ def plan_moe_weight_preparation( storage_policy = WeightStoragePolicy.KEEP_SOURCE if coupled_hadamard is None: - coupled_hadamard = False + coupled_hadamard = qsrt_profile in { + _QSRT_ATOMS_V2_PROFILE_COUPLED_K2, + _QSRT_ATOMS_V2_PROFILE_COUPLED_H308, + } return MoEWeightPreparationPlan( specs=normalized_specs, num_experts=num_experts, @@ -968,15 +959,9 @@ def plan_moe_weight_preparation( storage_policy=storage_policy, trellis_bits=trellis_bits, trellis_tile_config=trellis_tile_config, + qsrt_storage_format=qsrt_storage_format, + qsrt_profile=qsrt_profile, coupled_hadamard=coupled_hadamard, - trellis_codebook=trellis_codebook, - trellis_rate_structure=trellis_rate_structure, - trellis_pair_kinds=( - None - if trellis_pair_kinds is None - else frozenset(str(kind) for kind in trellis_pair_kinds) - ), - coupled_hadamard_blocks=coupled_hadamard_blocks, ) diff --git a/b12x/moe/_shared/kernels/dynamic.py b/b12x/moe/_shared/kernels/dynamic.py index 457a4d71e..04224d70b 100644 --- a/b12x/moe/_shared/kernels/dynamic.py +++ b/b12x/moe/_shared/kernels/dynamic.py @@ -2944,66 +2944,43 @@ class Storage: # one thread per K32 quantization block, plus one thread per route. # Fold it into phase 0 so the existing resident barrier publishes # both the cleared output and the prepared input in one step. - m1_lane_id = Int32(tidx) & Int32(31) - m1_has_active_route = Int32(0) - if m1_lane_id == Int32(0): - m1_route_idx = Int32(0) - while m1_route_idx < total_pairs: - m1_route_expert_id = topk_ids[m1_route_idx].to(Int32) - if ( - m1_route_expert_id >= Int32(0) - and m1_route_expert_id < num_experts - ): - m1_has_active_route = Int32(1) - m1_route_idx = total_pairs - else: - m1_route_idx += Int32(1) - m1_has_active_route = cute.arch.shuffle_sync( - m1_has_active_route, Int32(0) - ) - - if m1_has_active_route > Int32(0): - m1_blk_idx = flat_tid - while m1_blk_idx < mx_blocks_per_row: - m1_block_start = m1_blk_idx * Int32(32) - m1_values, m1_block_max = _load_bf16x32_to_f32( - a_input, - m1_block_start, + m1_blk_idx = flat_tid + while m1_blk_idx < mx_blocks_per_row: + m1_block_start = m1_blk_idx * Int32(32) + m1_values, m1_block_max = _load_bf16x32_to_f32( + a_input, + m1_block_start, + ) + if cutlass.const_expr(self.w4a8_trellis): + m1_payload, m1_mx_scale_byte = quantize_block_fp8_mx( + _w4a8_trellis_permute_k32(m1_values), + m1_block_max, ) - if cutlass.const_expr(self.w4a8_trellis): - m1_payload, m1_mx_scale_byte = quantize_block_fp8_mx( - _w4a8_trellis_permute_k32(m1_values), - m1_block_max, - ) - else: - m1_payload, m1_mx_scale_byte = quantize_block_fp8_mx( - m1_values, - m1_block_max, - ) - for m1_payload_pair in cutlass.range_constexpr(4): - m1_packed64 = ( - Uint64(m1_payload[m1_payload_pair * 2 + 1]) << Uint64(32) - ) | Uint64(m1_payload[m1_payload_pair * 2]) - st_global_u64( - get_ptr_as_int64( - packed_a_storage, - m1_block_start + Int32(m1_payload_pair * 8), - ), - m1_packed64, - ) - scale_storage[m1_blk_idx] = Uint8( - m1_mx_scale_byte & Uint32(0xFF) + else: + m1_payload, m1_mx_scale_byte = quantize_block_fp8_mx( + m1_values, + m1_block_max, + ) + for m1_payload_pair in cutlass.range_constexpr(4): + m1_packed64 = ( + Uint64(m1_payload[m1_payload_pair * 2 + 1]) << Uint64(32) + ) | Uint64(m1_payload[m1_payload_pair * 2]) + st_global_u64( + get_ptr_as_int64( + packed_a_storage, + m1_block_start + Int32(m1_payload_pair * 8), + ), + m1_packed64, ) - m1_blk_idx += flat_stride + scale_storage[m1_blk_idx] = Uint8(m1_mx_scale_byte & Uint32(0xFF)) + m1_blk_idx += flat_stride if flat_tid < total_pairs: - m1_slot_expert_id = topk_ids[flat_tid].to(Int32) - if m1_slot_expert_id >= Int32(0) and m1_slot_expert_id < num_experts: - m1_physical_row = flat_tid * Int32(self.tile_shape_mnk[0]) - token_map[m1_physical_row] = Int32(0) - token_weights[m1_physical_row] = topk_weights[flat_tid].to( - cutlass.Float32 - ) + m1_physical_row = flat_tid * Int32(self.tile_shape_mnk[0]) + token_map[m1_physical_row] = Int32(0) + token_weights[m1_physical_row] = topk_weights[flat_tid].to( + cutlass.Float32 + ) cute.arch.sync_threads() self._resident_grid_barrier( @@ -3021,13 +2998,7 @@ class Storage: hist_idx = flat_tid while hist_idx < total_pairs: expert_id = topk_ids[hist_idx].to(Int32) - # A route outside the local expert domain is inactive. vLLM - # uses -1 for scheduler padding, and expert-parallel callers - # may use the same contract for non-local routes. - if expert_id >= Int32(0) and expert_id < num_experts: - atomic_add_global_i32( - get_ptr_as_int64(row_counts, expert_id), Int32(1) - ) + atomic_add_global_i32(get_ptr_as_int64(row_counts, expert_id), Int32(1)) hist_idx += flat_stride self._resident_grid_barrier( @@ -3131,33 +3102,29 @@ class Storage: pair_idx = token_idx * num_topk + topk_slot expert_id = topk_ids[pair_idx].to(Int32) weight = topk_weights[pair_idx].to(cutlass.Float32) - phys_row = Int32(-1) - if expert_id >= Int32(0) and expert_id < num_experts: - if cutlass.const_expr(self.direct_routing): - row = Int32(0) - phys_tile = pair_idx - else: - row = atomic_add_global_i32( - get_ptr_as_int64( - expert_write_rows, expert_id - ), - Int32(1), - ) - phys_tile = expert_tile_base[ - expert_id - ] + row // Int32(self.tile_shape_mnk[0]) - phys_row = phys_tile * Int32( - self.tile_shape_mnk[0] - ) + row % Int32(self.tile_shape_mnk[0]) - map_value = token_idx - if cutlass.const_expr(self.deterministic_output): - map_value = pair_idx - st_global_i32( - get_ptr_as_int64(token_map, phys_row), map_value - ) - st_global_f32( - get_ptr_as_int64(token_weights, phys_row), weight + if cutlass.const_expr(self.direct_routing): + row = Int32(0) + phys_tile = pair_idx + else: + row = atomic_add_global_i32( + get_ptr_as_int64(expert_write_rows, expert_id), + Int32(1), ) + phys_tile = expert_tile_base[ + expert_id + ] + row // Int32(self.tile_shape_mnk[0]) + phys_row = phys_tile * Int32( + self.tile_shape_mnk[0] + ) + row % Int32(self.tile_shape_mnk[0]) + map_value = token_idx + if cutlass.const_expr(self.deterministic_output): + map_value = pair_idx + st_global_i32( + get_ptr_as_int64(token_map, phys_row), map_value + ) + st_global_f32( + get_ptr_as_int64(token_weights, phys_row), weight + ) slot = route_slot_base + topk_slot _st_shared_i32( route_phys_rows_addr + slot * Int32(4), phys_row @@ -3171,157 +3138,72 @@ class Storage: else: cute.arch.sync_warp() - token_has_active_route = Int32(0) - if lane_id == Int32(0): - active_topk_slot = Int32(0) - while active_topk_slot < num_topk: - slot = route_slot_base + active_topk_slot - if _ld_shared_i32( - route_phys_rows_addr + slot * Int32(4) - ) >= Int32(0): - token_has_active_route = Int32(1) - active_topk_slot = num_topk - else: - active_topk_slot += Int32(1) - token_has_active_route = cute.arch.shuffle_sync( - token_has_active_route, Int32(0) - ) - - if token_has_active_route > Int32(0): - if cutlass.const_expr(self.is_w4a8): - # A token's BF16 row is identical for every routed - # expert. The materialized specialization stores one - # quantized row per token and gathers it in FC1; the - # route-expanded path fans it out to each route. - # That path keeps only physical rows in rmem here to - # stay below the two-CTA register-residency limit. - if num_topk == Int32(8): - for cache_slot in cutlass.range_constexpr(8): - slot = route_slot_base + Int32(cache_slot) - shared_route_phys_rows[cache_slot] = _ld_shared_i32( - route_phys_rows_addr + slot * Int32(4) - ) + if cutlass.const_expr(self.is_w4a8): + # A token's BF16 row is identical for every routed + # expert. The materialized specialization stores one + # quantized row per token and gathers it in FC1; the + # route-expanded path fans it out to each route. + # That path keeps only physical rows in rmem here to + # stay below the two-CTA register-residency limit. + if num_topk == Int32(8): + for cache_slot in cutlass.range_constexpr(8): + slot = route_slot_base + Int32(cache_slot) + shared_route_phys_rows[cache_slot] = _ld_shared_i32( + route_phys_rows_addr + slot * Int32(4) + ) - blk_idx = lane_id + token_partition * Int32(32) - while blk_idx < mx_blocks_per_row: - block_start = blk_idx * Int32(32) - values, block_max = _load_bf16x32_to_f32( - a_input, - token_idx * Int32(a_input.shape[1]) - + block_start, - ) - if cutlass.const_expr(self.w4a8_trellis): - payload, mx_scale_byte = ( - quantize_block_fp8_mx( - _w4a8_trellis_permute_k32( - values - ), - block_max, - ) + blk_idx = lane_id + token_partition * Int32(32) + while blk_idx < mx_blocks_per_row: + block_start = blk_idx * Int32(32) + values, block_max = _load_bf16x32_to_f32( + a_input, + token_idx * Int32(a_input.shape[1]) + + block_start, + ) + if cutlass.const_expr(self.w4a8_trellis): + payload, mx_scale_byte = ( + quantize_block_fp8_mx( + _w4a8_trellis_permute_k32( + values + ), + block_max, ) - else: - payload, mx_scale_byte = ( - quantize_block_fp8_mx( - values, block_max - ) + ) + else: + payload, mx_scale_byte = ( + quantize_block_fp8_mx( + values, block_max ) - for payload_pair in cutlass.range_constexpr(4): - packed64 = ( - Uint64(payload[payload_pair * 2 + 1]) - << Uint64(32) - ) | Uint64(payload[payload_pair * 2]) - if cutlass.const_expr( - self.materialize_intermediate - ): - output_offset = ( - token_idx * output_bytes_per_row - + block_start - + Int32(payload_pair * 8) - ) - st_global_u64( - get_ptr_as_int64( - packed_a_storage, - output_offset, - ), - packed64, - ) - else: - # The offset needs a concrete - # type on the inactive-route - # fall-through path of the - # dynamic guard below. - output_offset = Int32(0) - for cache_slot in cutlass.range_constexpr( - 8 - ): - phys_row = shared_route_phys_rows[ - cache_slot - ] - if phys_row >= Int32(0): - output_offset = ( - phys_row * output_bytes_per_row - + block_start - + Int32(payload_pair * 8) - ) - st_global_u64( - get_ptr_as_int64( - packed_a_storage, - output_offset, - ), - packed64, - ) + ) + for payload_pair in cutlass.range_constexpr(4): + packed64 = ( + Uint64(payload[payload_pair * 2 + 1]) + << Uint64(32) + ) | Uint64(payload[payload_pair * 2]) if cutlass.const_expr( self.materialize_intermediate ): - scale_storage[ - token_idx * mx_blocks_per_row + blk_idx - ] = Uint8(mx_scale_byte & Uint32(0xFF)) + output_offset = ( + token_idx * output_bytes_per_row + + block_start + + Int32(payload_pair * 8) + ) + st_global_u64( + get_ptr_as_int64( + packed_a_storage, + output_offset, + ), + packed64, + ) else: - for cache_slot in cutlass.range_constexpr(8): + for cache_slot in cutlass.range_constexpr( + 8 + ): phys_row = shared_route_phys_rows[ cache_slot ] - if phys_row >= Int32(0): - scale_storage[ - phys_row * mx_blocks_per_row + blk_idx - ] = Uint8( - mx_scale_byte & Uint32(0xFF) - ) - blk_idx += Int32(self.input_warps_per_token * 32) - else: - blk_idx = lane_id + token_partition * Int32(32) - while blk_idx < mx_blocks_per_row: - block_start = blk_idx * Int32(32) - values, block_max = _load_bf16x32_to_f32( - a_input, - token_idx * Int32(a_input.shape[1]) - + block_start, - ) - if cutlass.const_expr(self.w4a8_trellis): - payload, mx_scale_byte = ( - quantize_block_fp8_mx( - _w4a8_trellis_permute_k32( - values - ), - block_max, - ) - ) - else: - payload, mx_scale_byte = ( - quantize_block_fp8_mx( - values, block_max - ) - ) - for payload_pair in cutlass.range_constexpr(4): - packed64 = ( - Uint64(payload[payload_pair * 2 + 1]) - << Uint64(32) - ) | Uint64(payload[payload_pair * 2]) - if cutlass.const_expr( - self.materialize_intermediate - ): output_offset = ( - token_idx * output_bytes_per_row + phys_row * output_bytes_per_row + block_start + Int32(payload_pair * 8) ) @@ -3332,243 +3214,284 @@ class Storage: ), packed64, ) - else: - # CuTe loop-carried values must have - # a concrete type before entering a - # dynamic while. Keep the address - # offset in rmem and update it for - # each routed copy below. - output_offset = Int32(0) - topk_slot = Int32(0) - while topk_slot < num_topk: - slot = route_slot_base + topk_slot - phys_row = _ld_shared_i32( - route_phys_rows_addr - + slot * Int32(4) - ) - if phys_row >= Int32(0): - output_offset = ( - phys_row * output_bytes_per_row - + block_start - + Int32(payload_pair * 8) - ) - st_global_u64( - get_ptr_as_int64( - packed_a_storage, - output_offset, - ), - packed64, - ) - topk_slot += Int32(1) + if cutlass.const_expr( + self.materialize_intermediate + ): + scale_storage[ + token_idx * mx_blocks_per_row + blk_idx + ] = Uint8(mx_scale_byte & Uint32(0xFF)) + else: + for cache_slot in cutlass.range_constexpr(8): + phys_row = shared_route_phys_rows[ + cache_slot + ] + scale_storage[ + phys_row * mx_blocks_per_row + blk_idx + ] = Uint8(mx_scale_byte & Uint32(0xFF)) + blk_idx += Int32(self.input_warps_per_token * 32) + else: + blk_idx = lane_id + token_partition * Int32(32) + while blk_idx < mx_blocks_per_row: + block_start = blk_idx * Int32(32) + values, block_max = _load_bf16x32_to_f32( + a_input, + token_idx * Int32(a_input.shape[1]) + + block_start, + ) + if cutlass.const_expr(self.w4a8_trellis): + payload, mx_scale_byte = ( + quantize_block_fp8_mx( + _w4a8_trellis_permute_k32( + values + ), + block_max, + ) + ) + else: + payload, mx_scale_byte = ( + quantize_block_fp8_mx( + values, block_max + ) + ) + for payload_pair in cutlass.range_constexpr(4): + packed64 = ( + Uint64(payload[payload_pair * 2 + 1]) + << Uint64(32) + ) | Uint64(payload[payload_pair * 2]) if cutlass.const_expr( self.materialize_intermediate ): - scale_storage[ - token_idx * mx_blocks_per_row + blk_idx - ] = Uint8(mx_scale_byte & Uint32(0xFF)) + output_offset = ( + token_idx * output_bytes_per_row + + block_start + + Int32(payload_pair * 8) + ) + st_global_u64( + get_ptr_as_int64( + packed_a_storage, + output_offset, + ), + packed64, + ) else: + # CuTe loop-carried values must have + # a concrete type before entering a + # dynamic while. Keep the address + # offset in rmem and update it for + # each routed copy below. + output_offset = Int32(0) topk_slot = Int32(0) while topk_slot < num_topk: slot = route_slot_base + topk_slot phys_row = _ld_shared_i32( - route_phys_rows_addr + slot * Int32(4) + route_phys_rows_addr + + slot * Int32(4) ) - if phys_row >= Int32(0): - scale_storage[ - phys_row * mx_blocks_per_row + blk_idx - ] = Uint8( - mx_scale_byte & Uint32(0xFF) - ) - topk_slot += Int32(1) - blk_idx += Int32(self.input_warps_per_token * 32) - else: - gs_value = shared_input_gs_value - if num_topk == Int32(8): - for cache_slot in cutlass.range_constexpr(8): - slot = route_slot_base + Int32(cache_slot) - phys_row = _ld_shared_i32( - route_phys_rows_addr + slot * Int32(4) - ) - route_output_base[cache_slot] = Int32(-1) - route_scale_base[cache_slot] = Int32(-1) - if phys_row >= Int32(0): - route_output_base[cache_slot] = ( - phys_row * output_bytes_per_row - ) - # Scale storage is tiled in 128-row SF - # atoms, independently of the MMA tile. - sf_atom = phys_row >> Int32(7) - sf_row = phys_row & Int32(127) - route_scale_base[cache_slot] = ( - sf_atom - * num_k_tiles - * Int32(32 * 4 * 4) - + (sf_row % Int32(32)) * Int32(4 * 4) - + (sf_row // Int32(32)) * Int32(4) - ) - - sf_idx = lane_id - while sf_idx < sf_blocks_per_row: - block_start = sf_idx * Int32(16) - values = cute.make_rmem_tensor( - (16,), cutlass.Float32 - ) - block_max = cutlass.Float32(0.0) - for elem_idx in cutlass.range_constexpr(16): - value = cutlass.Float32( - a_input[ - token_idx, block_start + Int32(elem_idx) - ] - ) - values[elem_idx] = value - block_max = fmax_f32(block_max, fabs_f32(value)) - packed64 = Uint64(0) - scale_byte = Uint8(0) - if self.is_gated and self.fast_math: - packed64, scale_byte = quantize_block_fp4_fast( - values, block_max, gs_value - ) - else: - packed64, scale_byte = quantize_block_fp4( - values, block_max, gs_value - ) - - k_tile_idx = sf_idx // Int32(4) - inner_k_idx = sf_idx % Int32(4) - scale_k_base = ( - k_tile_idx * Int32(32 * 4 * 4) + inner_k_idx - ) - for cache_slot in cutlass.range_constexpr(8): - output_base = route_output_base[cache_slot] - if output_base >= Int32(0): - output_offset = output_base + sf_idx * Int32( - 8 + output_offset = ( + phys_row * output_bytes_per_row + + block_start + + Int32(payload_pair * 8) ) st_global_u64( get_ptr_as_int64( - packed_a_storage, output_offset + packed_a_storage, + output_offset, ), packed64, ) - scale_storage[ - route_scale_base[cache_slot] - + scale_k_base - ] = scale_byte - sf_idx += Int32(32) - else: - sf_idx = lane_id - while sf_idx < sf_blocks_per_row: - block_start = sf_idx * Int32(16) - values = cute.make_rmem_tensor( - (16,), cutlass.Float32 - ) - block_max = cutlass.Float32(0.0) - for elem_idx in cutlass.range_constexpr(16): - value = cutlass.Float32( - a_input[ - token_idx, block_start + Int32(elem_idx) - ] - ) - values[elem_idx] = value - block_max = fmax_f32(block_max, fabs_f32(value)) - packed64 = Uint64(0) - scale_byte = Uint8(0) - if self.is_gated and self.fast_math: - packed64, scale_byte = quantize_block_fp4_fast( - values, block_max, gs_value - ) - else: - packed64, scale_byte = quantize_block_fp4( - values, block_max, gs_value - ) - + topk_slot += Int32(1) + if cutlass.const_expr( + self.materialize_intermediate + ): + scale_storage[ + token_idx * mx_blocks_per_row + blk_idx + ] = Uint8(mx_scale_byte & Uint32(0xFF)) + else: topk_slot = Int32(0) while topk_slot < num_topk: slot = route_slot_base + topk_slot phys_row = _ld_shared_i32( route_phys_rows_addr + slot * Int32(4) ) - if phys_row >= Int32(0): - output_offset = ( - phys_row * output_bytes_per_row - + sf_idx * Int32(8) - ) - st_global_u64( - get_ptr_as_int64( - packed_a_storage, output_offset - ), - packed64, - ) - - # Scale storage uses 128-row SF atoms, - # independently of the MMA tile. - k_tile_idx = sf_idx // Int32(4) - sf_atom = phys_row >> Int32(7) - sf_row = phys_row & Int32(127) - outer_m_idx = sf_row % Int32(32) - inner_m_idx = sf_row // Int32(32) - inner_k_idx = sf_idx % Int32(4) - scale_offset = ( - sf_atom - * num_k_tiles - * Int32(32 * 4 * 4) - + k_tile_idx * Int32(32 * 4 * 4) - + outer_m_idx * Int32(4 * 4) - + inner_m_idx * Int32(4) - + inner_k_idx - ) - scale_storage[scale_offset] = scale_byte + scale_storage[ + phys_row * mx_blocks_per_row + blk_idx + ] = Uint8(mx_scale_byte & Uint32(0xFF)) topk_slot += Int32(1) - sf_idx += Int32(32) + blk_idx += Int32(self.input_warps_per_token * 32) + else: + gs_value = shared_input_gs_value + if num_topk == Int32(8): + for cache_slot in cutlass.range_constexpr(8): + slot = route_slot_base + Int32(cache_slot) + phys_row = _ld_shared_i32( + route_phys_rows_addr + slot * Int32(4) + ) + route_output_base[cache_slot] = ( + phys_row * output_bytes_per_row + ) + # Scale storage is tiled in 128-row SF + # atoms, independently of the MMA tile. + sf_atom = phys_row >> Int32(7) + sf_row = phys_row & Int32(127) + route_scale_base[cache_slot] = ( + sf_atom * num_k_tiles * Int32(32 * 4 * 4) + + (sf_row % Int32(32)) * Int32(4 * 4) + + (sf_row // Int32(32)) * Int32(4) + ) - if cutlass.const_expr(self.work_is_streaming): - cute.arch.sync_warp() - _threadfence() - cute.arch.sync_warp() + sf_idx = lane_id + while sf_idx < sf_blocks_per_row: + block_start = sf_idx * Int32(16) + values = cute.make_rmem_tensor( + (16,), cutlass.Float32 + ) + block_max = cutlass.Float32(0.0) + for elem_idx in cutlass.range_constexpr(16): + value = cutlass.Float32( + a_input[ + token_idx, block_start + Int32(elem_idx) + ] + ) + values[elem_idx] = value + block_max = fmax_f32(block_max, fabs_f32(value)) + packed64 = Uint64(0) + scale_byte = Uint8(0) + if self.is_gated and self.fast_math: + packed64, scale_byte = quantize_block_fp4_fast( + values, block_max, gs_value + ) + else: + packed64, scale_byte = quantize_block_fp4( + values, block_max, gs_value + ) - publish_routes = Int32(1) - if cutlass.const_expr(self.is_w4a8): - self._sync_input_warp_pair(token_owner_warp) - publish_routes = Int32(1) - token_partition + k_tile_idx = sf_idx // Int32(4) + inner_k_idx = sf_idx % Int32(4) + scale_k_base = ( + k_tile_idx * Int32(32 * 4 * 4) + inner_k_idx + ) + for cache_slot in cutlass.range_constexpr(8): + output_offset = route_output_base[ + cache_slot + ] + sf_idx * Int32(8) + st_global_u64( + get_ptr_as_int64( + packed_a_storage, output_offset + ), + packed64, + ) + scale_storage[ + route_scale_base[cache_slot] + scale_k_base + ] = scale_byte + sf_idx += Int32(32) + else: + sf_idx = lane_id + while sf_idx < sf_blocks_per_row: + block_start = sf_idx * Int32(16) + values = cute.make_rmem_tensor( + (16,), cutlass.Float32 + ) + block_max = cutlass.Float32(0.0) + for elem_idx in cutlass.range_constexpr(16): + value = cutlass.Float32( + a_input[ + token_idx, block_start + Int32(elem_idx) + ] + ) + values[elem_idx] = value + block_max = fmax_f32(block_max, fabs_f32(value)) + packed64 = Uint64(0) + scale_byte = Uint8(0) + if self.is_gated and self.fast_math: + packed64, scale_byte = quantize_block_fp4_fast( + values, block_max, gs_value + ) + else: + packed64, scale_byte = quantize_block_fp4( + values, block_max, gs_value + ) - if lane_id == Int32(0) and publish_routes > Int32(0): topk_slot = Int32(0) while topk_slot < num_topk: slot = route_slot_base + topk_slot phys_row = _ld_shared_i32( route_phys_rows_addr + slot * Int32(4) ) - expert_id = _ld_shared_i32( - route_expert_ids_addr + slot * Int32(4) + output_offset = ( + phys_row * output_bytes_per_row + + sf_idx * Int32(8) ) - if phys_row >= Int32(0): - phys_tile = phys_row // Int32( - self.tile_shape_mnk[0] - ) - completed = atomic_add_global_i32( - get_ptr_as_int64( - tile_write_count, phys_tile - ), - Int32(1), - ) + Int32(1) - if completed == Int32( - self.tile_shape_mnk[0] - ): - self._publish_ready_tasks( - task_tail, - task_ready, - task_expert, - task_m_tile, - task_slice_begin, - task_slice_count, - task_valid_rows, - route_gate_tile_cnt, - task_slice_chunk, - expert_id, - phys_tile, - Int32(self.tile_shape_mnk[0]), - ) + st_global_u64( + get_ptr_as_int64( + packed_a_storage, output_offset + ), + packed64, + ) + + # scale_storage uses 128-row SF atoms + # (tile_atom_to_shape_SF); index by the 128-atom + # (phys_row>>7) + row-within-atom, NOT the MMA + # tile (which may be 64). Identity at tile_m==128. + k_tile_idx = sf_idx // Int32(4) + sf_atom = phys_row >> Int32(7) + sf_row = phys_row & Int32(127) + outer_m_idx = sf_row % Int32(32) + inner_m_idx = sf_row // Int32(32) + inner_k_idx = sf_idx % Int32(4) + scale_offset = ( + sf_atom * num_k_tiles * Int32(32 * 4 * 4) + + k_tile_idx * Int32(32 * 4 * 4) + + outer_m_idx * Int32(4 * 4) + + inner_m_idx * Int32(4) + + inner_k_idx + ) + scale_storage[scale_offset] = scale_byte topk_slot += Int32(1) + sf_idx += Int32(32) + + if cutlass.const_expr(self.work_is_streaming): + cute.arch.sync_warp() + _threadfence() + cute.arch.sync_warp() + + publish_routes = Int32(1) + if cutlass.const_expr(self.is_w4a8): + self._sync_input_warp_pair(token_owner_warp) + publish_routes = Int32(1) - token_partition + + if lane_id == Int32(0) and publish_routes > Int32(0): + topk_slot = Int32(0) + while topk_slot < num_topk: + slot = route_slot_base + topk_slot + phys_row = _ld_shared_i32( + route_phys_rows_addr + slot * Int32(4) + ) + expert_id = _ld_shared_i32( + route_expert_ids_addr + slot * Int32(4) + ) + phys_tile = phys_row // Int32( + self.tile_shape_mnk[0] + ) + completed = atomic_add_global_i32( + get_ptr_as_int64(tile_write_count, phys_tile), + Int32(1), + ) + Int32(1) + if completed == Int32(self.tile_shape_mnk[0]): + self._publish_ready_tasks( + task_tail, + task_ready, + task_expert, + task_m_tile, + task_slice_begin, + task_slice_count, + task_valid_rows, + route_gate_tile_cnt, + task_slice_chunk, + expert_id, + phys_tile, + Int32(self.tile_shape_mnk[0]), + ) + topk_slot += Int32(1) else: warp_item = Int32(0) while warp_item < Int32(_PRODUCER_PAIRS_PER_WARP): @@ -3578,15 +3501,12 @@ class Storage: weight = cutlass.Float32(0.0) row = Int32(0) phys_tile = Int32(0) - route_is_valid = Int32(0) if pair_idx < total_pairs: expert_id = topk_ids[pair_idx].to(Int32) token_idx = pair_idx // num_topk weight = topk_weights[pair_idx].to(cutlass.Float32) - if expert_id >= Int32(0) and expert_id < num_experts: - route_is_valid = Int32(1) - if lane_id == Int32(0) and route_is_valid > Int32(0): + if lane_id == Int32(0): if cutlass.const_expr(self.direct_routing): row = Int32(0) phys_tile = pair_idx @@ -3615,222 +3535,212 @@ class Storage: phys_tile = cute.arch.shuffle_sync(phys_tile, Int32(0)) expert_id = cute.arch.shuffle_sync(expert_id, Int32(0)) token_idx = cute.arch.shuffle_sync(token_idx, Int32(0)) - route_is_valid = cute.arch.shuffle_sync( - route_is_valid, Int32(0) - ) - if route_is_valid > Int32(0): - gs_value = cutlass.Float32(0.0) - if cutlass.const_expr(not self.is_w4a8): - gs_value = input_global_scale[expert_id].to( - cutlass.Float32 + gs_value = input_global_scale[expert_id].to(cutlass.Float32) + if cutlass.const_expr(self.is_w4a8): + # w4a8: per-32 dynamic UE8M0 + E4M3 payload, no + # global scale. Payload stored plain row-major; + # scales stored plain [row, cols//32]. + phys_row = phys_tile * Int32( + self.tile_shape_mnk[0] + ) + row % Int32(self.tile_shape_mnk[0]) + blk_idx = lane_id + while blk_idx < mx_blocks_per_row: + block_start = blk_idx * Int32(32) + values = cute.make_rmem_tensor( + (32,), cutlass.Float32 ) - if cutlass.const_expr(self.is_w4a8): - # w4a8: per-32 dynamic UE8M0 + E4M3 payload, no - # global scale. Payload stored plain row-major; - # scales stored plain [row, cols//32]. - phys_row = phys_tile * Int32( - self.tile_shape_mnk[0] - ) + row % Int32(self.tile_shape_mnk[0]) - blk_idx = lane_id - while blk_idx < mx_blocks_per_row: - block_start = blk_idx * Int32(32) - values = cute.make_rmem_tensor( - (32,), cutlass.Float32 - ) - block_max = cutlass.Float32(0.0) - for elem_idx in cutlass.range_constexpr(32): - value = cutlass.Float32( - a_input[ - token_idx, block_start + Int32(elem_idx) - ] - ) - values[elem_idx] = value - block_max = fmax_f32(block_max, fabs_f32(value)) - if cutlass.const_expr(self.w4a8_trellis): - payload, mx_scale_byte = ( - quantize_block_fp8_mx( - _w4a8_trellis_permute_k32( - values - ), - block_max, - ) - ) - else: - payload, mx_scale_byte = ( - quantize_block_fp8_mx( - values, block_max - ) - ) - output_offset = ( - phys_row * output_bytes_per_row + block_start + block_max = cutlass.Float32(0.0) + for elem_idx in cutlass.range_constexpr(32): + value = cutlass.Float32( + a_input[ + token_idx, block_start + Int32(elem_idx) + ] ) - for pair_idx in cutlass.range_constexpr(4): - packed64 = ( - Uint64(payload[pair_idx * 2 + 1]) - << Uint64(32) - ) | Uint64(payload[pair_idx * 2]) - st_global_u64( - get_ptr_as_int64( - packed_a_storage, - output_offset + Int32(pair_idx * 8), + values[elem_idx] = value + block_max = fmax_f32(block_max, fabs_f32(value)) + if cutlass.const_expr(self.w4a8_trellis): + payload, mx_scale_byte = ( + quantize_block_fp8_mx( + _w4a8_trellis_permute_k32( + values ), - packed64, - ) - scale_storage[ - phys_row * mx_blocks_per_row + blk_idx - ] = Uint8(mx_scale_byte & Uint32(0xFF)) - blk_idx += Int32(32) - elif cutlass.const_expr(self.is_w6a8): - # w6a8_mx: MXFP8-E4M3 K/32 byte-container - # payload (same encoding as w4a8's activation - # side, but WITH the calibrated per-expert - # global scale folded in at quantize time), - # stored row-major for the A TMA; scale bytes - # go to the swizzled 128-row SF atoms exactly - # like nvfp4 (sf_idx now indexes K/32 blocks). - sf_idx = lane_id - while sf_idx < sf_blocks_per_row: - block_start = sf_idx * quant_block_elems - values = cute.make_rmem_tensor( - (32,), cutlass.Float32 - ) - block_max = cutlass.Float32(0.0) - for elem_idx in cutlass.range_constexpr(32): - value = cutlass.Float32( - a_input[ - token_idx, block_start + Int32(elem_idx) - ] - ) - values[elem_idx] = value - block_max = fmax_f32(block_max, fabs_f32(value)) - containers, scale_byte = ( - moe_mxfp6_quantize_input_block_containers( - values, block_max, - gs_value, - self.mxfp6_fmt_a, ) ) - output_offset = ( - phys_tile * Int32(self.tile_shape_mnk[0]) - + row % Int32(self.tile_shape_mnk[0]) - ) * output_bytes_per_row + ( - sf_idx * packed_bytes_per_sf_block - ) - moe_mxfp6_store_expanded_global( - packed_a_storage, - output_offset, - containers, - ) - - k_tile_idx = sf_idx // Int32(4) - inner_k_idx = sf_idx % Int32(4) - # Scale storage uses 128-row SF atoms, - # independently of the MMA tile. - phys_row = phys_tile * Int32( - self.tile_shape_mnk[0] - ) + row % Int32(self.tile_shape_mnk[0]) - sf_atom = phys_row >> Int32(7) - sf_row = phys_row & Int32(127) - scale_offset = ( - sf_atom - * num_k_tiles - * Int32(32 * 4 * 4) - + k_tile_idx * Int32(32 * 4 * 4) - + (sf_row % Int32(32)) * Int32(4 * 4) - + (sf_row // Int32(32)) * Int32(4) - + inner_k_idx - ) - scale_storage[scale_offset] = scale_byte - sf_idx += Int32(32) - else: - sf_idx = lane_id - while sf_idx < sf_blocks_per_row: - block_start = sf_idx * Int32(16) - values = cute.make_rmem_tensor( - (16,), cutlass.Float32 - ) - block_max = cutlass.Float32(0.0) - for elem_idx in cutlass.range_constexpr(16): - value = cutlass.Float32( - a_input[ - token_idx, block_start + Int32(elem_idx) - ] - ) - values[elem_idx] = value - block_max = fmax_f32(block_max, fabs_f32(value)) - packed64 = Uint64(0) - scale_byte = Uint8(0) - if self.is_gated and self.fast_math: - packed64, scale_byte = quantize_block_fp4_fast( - values, block_max, gs_value - ) - else: - packed64, scale_byte = quantize_block_fp4( - values, block_max, gs_value + else: + payload, mx_scale_byte = ( + quantize_block_fp8_mx( + values, block_max ) - - output_offset = ( - phys_tile * Int32(self.tile_shape_mnk[0]) - + row % Int32(self.tile_shape_mnk[0]) - ) * output_bytes_per_row + sf_idx * Int32(8) + ) + output_offset = ( + phys_row * output_bytes_per_row + block_start + ) + for pair_idx in cutlass.range_constexpr(4): + packed64 = ( + Uint64(payload[pair_idx * 2 + 1]) + << Uint64(32) + ) | Uint64(payload[pair_idx * 2]) st_global_u64( get_ptr_as_int64( - packed_a_storage, output_offset + packed_a_storage, + output_offset + Int32(pair_idx * 8), ), packed64, ) + scale_storage[ + phys_row * mx_blocks_per_row + blk_idx + ] = Uint8(mx_scale_byte & Uint32(0xFF)) + blk_idx += Int32(32) + elif cutlass.const_expr(self.is_w6a8): + # w6a8_mx: MXFP8-E4M3 K/32 byte-container + # payload (same encoding as w4a8's activation + # side, but WITH the calibrated per-expert + # global scale folded in at quantize time), + # stored row-major for the A TMA; scale bytes + # go to the swizzled 128-row SF atoms exactly + # like nvfp4 (sf_idx now indexes K/32 blocks). + sf_idx = lane_id + while sf_idx < sf_blocks_per_row: + block_start = sf_idx * quant_block_elems + values = cute.make_rmem_tensor( + (32,), cutlass.Float32 + ) + block_max = cutlass.Float32(0.0) + for elem_idx in cutlass.range_constexpr(32): + value = cutlass.Float32( + a_input[ + token_idx, block_start + Int32(elem_idx) + ] + ) + values[elem_idx] = value + block_max = fmax_f32(block_max, fabs_f32(value)) + containers, scale_byte = ( + moe_mxfp6_quantize_input_block_containers( + values, + block_max, + gs_value, + self.mxfp6_fmt_a, + ) + ) + output_offset = ( + phys_tile * Int32(self.tile_shape_mnk[0]) + + row % Int32(self.tile_shape_mnk[0]) + ) * output_bytes_per_row + ( + sf_idx * packed_bytes_per_sf_block + ) + moe_mxfp6_store_expanded_global( + packed_a_storage, + output_offset, + containers, + ) - k_tile_idx = sf_idx // Int32(4) - inner_k_idx = sf_idx % Int32(4) - # Scale storage uses 128-row SF atoms, - # independently of the MMA tile. - phys_row = phys_tile * Int32( - self.tile_shape_mnk[0] - ) + row % Int32(self.tile_shape_mnk[0]) - sf_atom = phys_row >> Int32(7) - sf_row = phys_row & Int32(127) - scale_offset = ( - sf_atom - * num_k_tiles - * Int32(32 * 4 * 4) - + k_tile_idx * Int32(32 * 4 * 4) - + (sf_row % Int32(32)) * Int32(4 * 4) - + (sf_row // Int32(32)) * Int32(4) - + inner_k_idx + k_tile_idx = sf_idx // Int32(4) + inner_k_idx = sf_idx % Int32(4) + # scale_storage uses 128-row SF atoms: index by + # the 128-atom + row-within-atom, not the MMA + # tile. Identity at tile_m==128. + phys_row = phys_tile * Int32( + self.tile_shape_mnk[0] + ) + row % Int32(self.tile_shape_mnk[0]) + sf_atom = phys_row >> Int32(7) + sf_row = phys_row & Int32(127) + scale_offset = ( + sf_atom * num_k_tiles * Int32(32 * 4 * 4) + + k_tile_idx * Int32(32 * 4 * 4) + + (sf_row % Int32(32)) * Int32(4 * 4) + + (sf_row // Int32(32)) * Int32(4) + + inner_k_idx + ) + scale_storage[scale_offset] = scale_byte + sf_idx += Int32(32) + else: + sf_idx = lane_id + while sf_idx < sf_blocks_per_row: + block_start = sf_idx * Int32(16) + values = cute.make_rmem_tensor( + (16,), cutlass.Float32 + ) + block_max = cutlass.Float32(0.0) + for elem_idx in cutlass.range_constexpr(16): + value = cutlass.Float32( + a_input[ + token_idx, block_start + Int32(elem_idx) + ] ) - scale_storage[scale_offset] = scale_byte - sf_idx += Int32(32) + values[elem_idx] = value + block_max = fmax_f32(block_max, fabs_f32(value)) + packed64 = Uint64(0) + scale_byte = Uint8(0) + if self.is_gated and self.fast_math: + packed64, scale_byte = quantize_block_fp4_fast( + values, block_max, gs_value + ) + else: + packed64, scale_byte = quantize_block_fp4( + values, block_max, gs_value + ) + + output_offset = ( + phys_tile * Int32(self.tile_shape_mnk[0]) + + row % Int32(self.tile_shape_mnk[0]) + ) * output_bytes_per_row + sf_idx * Int32(8) + st_global_u64( + get_ptr_as_int64( + packed_a_storage, output_offset + ), + packed64, + ) - if cutlass.const_expr(self.work_is_streaming): - cute.arch.sync_warp() - # When the whole launch has fewer than one M-tile of routed - # rows, only the final partial-tile flush can publish work. - # Skip the per-row fence/counter path in that common micro case. - _threadfence() - cute.arch.sync_warp() + k_tile_idx = sf_idx // Int32(4) + inner_k_idx = sf_idx % Int32(4) + # scale_storage uses 128-row SF atoms: index by the + # 128-atom + row-within-atom, not the MMA tile. + # Identity at tile_m==128. + phys_row = phys_tile * Int32( + self.tile_shape_mnk[0] + ) + row % Int32(self.tile_shape_mnk[0]) + sf_atom = phys_row >> Int32(7) + sf_row = phys_row & Int32(127) + scale_offset = ( + sf_atom * num_k_tiles * Int32(32 * 4 * 4) + + k_tile_idx * Int32(32 * 4 * 4) + + (sf_row % Int32(32)) * Int32(4 * 4) + + (sf_row // Int32(32)) * Int32(4) + + inner_k_idx + ) + scale_storage[scale_offset] = scale_byte + sf_idx += Int32(32) - if lane_id == Int32(0): - completed = atomic_add_global_i32( - get_ptr_as_int64(tile_write_count, phys_tile), - Int32(1), - ) + Int32(1) - if completed == Int32(self.tile_shape_mnk[0]): - self._publish_ready_tasks( - task_tail, - task_ready, - task_expert, - task_m_tile, - task_slice_begin, - task_slice_count, - task_valid_rows, - route_gate_tile_cnt, - task_slice_chunk, - expert_id, - phys_tile, - Int32(self.tile_shape_mnk[0]), - ) + if cutlass.const_expr(self.work_is_streaming): + cute.arch.sync_warp() + # When the whole launch has fewer than one M-tile of routed + # rows, only the final partial-tile flush can publish work. + # Skip the per-row fence/counter path in that common micro case. + _threadfence() + cute.arch.sync_warp() + + if lane_id == Int32(0): + completed = atomic_add_global_i32( + get_ptr_as_int64(tile_write_count, phys_tile), + Int32(1), + ) + Int32(1) + if completed == Int32(self.tile_shape_mnk[0]): + self._publish_ready_tasks( + task_tail, + task_ready, + task_expert, + task_m_tile, + task_slice_begin, + task_slice_count, + task_valid_rows, + route_gate_tile_cnt, + task_slice_chunk, + expert_id, + phys_tile, + Int32(self.tile_shape_mnk[0]), + ) warp_item += Int32(1) if cutlass.const_expr(not self.w4a8_m1_materialized): @@ -3859,17 +3769,6 @@ class Storage: pair_flush = Int32(bidz) while pair_flush < total_pairs: expert_flush = topk_ids[pair_flush].to(Int32) - valid_rows = Int32(0) - if ( - expert_flush >= Int32(0) - and expert_flush < num_experts - ): - valid_rows = Int32(1) - else: - # Deferred direct tasks use an arithmetic slot per - # route. Keep the slot addressable, but give an - # inactive route a safe expert index and no work. - expert_flush = Int32(0) self._publish_deferred_tasks( task_expert, task_valid_rows, @@ -3877,7 +3776,7 @@ class Storage: task_slice_chunk, expert_flush, pair_flush, - valid_rows, + Int32(1), ) pair_flush += Int32(gdim_z) else: @@ -4518,14 +4417,12 @@ class Storage: if materialized_slot < materialized_tail: route_idx = materialized_slot // route_gate_tile_cnt route_slice = materialized_slot - route_idx * route_gate_tile_cnt - route_expert = topk_ids[route_idx].to(Int32) - if route_expert >= Int32(0) and route_expert < num_experts: - work_item[_WORK_EXPERT] = route_expert - work_item[_WORK_M_TILE] = route_idx - work_item[_WORK_SLICE_BEGIN] = route_slice - work_item[_WORK_SLICE_COUNT] = Int32(1) - work_item[_WORK_VALID_ROWS] = Int32(1) - has_task = Int32(1) + work_item[_WORK_EXPERT] = topk_ids[route_idx].to(Int32) + work_item[_WORK_M_TILE] = route_idx + work_item[_WORK_SLICE_BEGIN] = route_slice + work_item[_WORK_SLICE_COUNT] = Int32(1) + work_item[_WORK_VALID_ROWS] = Int32(1) + has_task = Int32(1) else: is_done = Int32(1) materialized_slot += Int32(gdim_z) @@ -4644,12 +4541,6 @@ class Storage: not self.work_is_persistent_grid and not self.w4a8_m1_materialized ): self._load_shared_work_item(work_item, ctrl_base_addr) - # Direct routing retains arithmetic task slots for graph-stable - # scheduling. An inactive route carries zero valid rows and - # must not reach any expert weight or scale load. - if work_item[_WORK_VALID_ROWS] <= Int32(0): - has_task = Int32(0) - if has_task > Int32(0): task_m_tile_idx_cache = work_item[_WORK_M_TILE] task_valid_rows_cache = work_item[_WORK_VALID_ROWS] tile_m_base_cache = task_m_tile_idx_cache * Int32( @@ -8872,30 +8763,29 @@ class Storage: phase2_slot - phase2_m_tile * phase2_task_output_tiles ) phase2_expert = topk_ids[phase2_m_tile].to(Int32) - if phase2_expert >= Int32(0) and phase2_expert < num_experts: - self._run_w4a8_materialized_fc2( - intermediate_u32, - down_rp, - down_sfb_rp, - scatter_output, - token_map, - token_weights, - down_alpha, - global_scale, - sa_flat_addr, - w4a8_sb0, - w4a8_sb1, - w4a8_sfbb, - Int32(tidx), - warp_idx, - phase2_m_tile, - phase2_expert, - phase2_output_tile, - Int32(1), - rows_capacity, - gate_tile_cnt, - phase2_packed_output_tiles, - ) + self._run_w4a8_materialized_fc2( + intermediate_u32, + down_rp, + down_sfb_rp, + scatter_output, + token_map, + token_weights, + down_alpha, + global_scale, + sa_flat_addr, + w4a8_sb0, + w4a8_sb1, + w4a8_sfbb, + Int32(tidx), + warp_idx, + phase2_m_tile, + phase2_expert, + phase2_output_tile, + Int32(1), + rows_capacity, + gate_tile_cnt, + phase2_packed_output_tiles, + ) cute.arch.sync_threads() phase2_slot += Int32(gdim_z) else: diff --git a/b12x/moe/_shared/kernels/micro.py b/b12x/moe/_shared/kernels/micro.py index 3bec68a84..7e592bf62 100644 --- a/b12x/moe/_shared/kernels/micro.py +++ b/b12x/moe/_shared/kernels/micro.py @@ -406,6 +406,7 @@ def __init__( weight_layout: str = "modelopt", trellis_bits: int | None = None, trellis_coupled: bool = False, + stage_inactive_routes: bool = False, ): activation = normalize_moe_activation(activation) if int(compile_time_phase) not in {0, 1, 2}: @@ -425,22 +426,22 @@ def __init__( swiglu_beta = normalize_swiglu_beta_for_activation(activation, swiglu_beta) if w13_layout not in {"w13", "w31"}: raise ValueError(f"unsupported micro w13_layout {w13_layout!r}") - if weight_layout not in {"modelopt", "trellis_t256"}: + if weight_layout not in {"modelopt", "trellis3_t256"}: raise ValueError(f"unsupported micro weight_layout {weight_layout!r}") - if weight_layout == "trellis_t256": + if weight_layout == "trellis3_t256": if trellis_bits not in (2, 3, 4): raise ValueError( - "trellis_t256 micro weights require trellis_bits in " + "trellis3_t256 micro weights require trellis_bits in " f"{{2, 3, 4}}, got {trellis_bits!r}" ) if not is_gated_moe_activation(activation): raise ValueError( - "trellis_t256 micro weights require a gated activation" + "trellis3_t256 micro weights require a gated activation" ) elif trellis_bits is not None or trellis_coupled: raise ValueError( "trellis_bits/trellis_coupled require weight_layout " - "'trellis_t256'" + "'trellis3_t256'" ) self.scale_format = scale_format self.scale_format_e8m0_k32 = scale_format == "e8m0_k32" @@ -470,11 +471,13 @@ def __init__( # w4a8 prefill recipe. Same f16 dot-product math and weights. self.a8_mx_mode = a8_mx_mode self.weight_layout = weight_layout - self.weight_layout_trellis256 = weight_layout == "trellis_t256" + self.weight_layout_trellis256 = weight_layout == "trellis3_t256" self.trellis_bits = 0 if trellis_bits is None else int(trellis_bits) self.trellis_coupled = bool(trellis_coupled) + self.stage_inactive_routes = bool(stage_inactive_routes) self.trellis_ksplit = 1 self.trellis_scratch_u32 = 0 + self.route_slots = 0 self._cfg = None self.m_const = 0 self.m1_fc2_onepass = False @@ -497,6 +500,7 @@ def __cache_key__(self): self.weight_layout, self.trellis_bits, self.trellis_coupled, + self.stage_inactive_routes, self.trellis_ksplit, self.scale_format, self.e8m0_scale_layout, @@ -511,6 +515,7 @@ def __cache_key__(self): self.m1_fc2_onepass, self.m1_fc2_rows_per_cta, self.launch_block_dim, + self.route_slots, ) @cute.jit @@ -548,6 +553,32 @@ def _fp4_dot4_for_math( ) return fp4_dot4_sum_f32acc(u_packed, x0, x1, x2, x3) + @cute.jit + def _fc2_route( + self, + topk_ids: cute.Tensor, + topk_weights: cute.Tensor, + eid_addr: Int32, + route_expert_limit: Int32, + ) -> Tuple[Int32, Float32]: + """Return an addressable expert and its effective router weight. + + Fixed-M fused launches receive a route table sanitized in shared + memory before FC1 and FC2 execute. The FC2-only endpoint instead uses + runtime M with a compile-time expert-capacity bucket, so it must check + each route against the resident expert count supplied by the caller. + Invalid routes address expert zero with an exact-zero effective weight. + """ + if cutlass.const_expr(self.compile_time_phase == 2): + raw_expert = Int64(topk_ids[eid_addr]) + expert = Int32(0) + weight = Float32(0.0) + if raw_expert >= Int64(0) and raw_expert < Int64(route_expert_limit): + expert = Int32(raw_expert) + weight = Float32(topk_weights[eid_addr]) + return expert, weight + return Int32(topk_ids[eid_addr]), Float32(topk_weights[eid_addr]) + @cute.jit def _scale_byte_to_f32(self, byte: Uint32) -> Float32: """Decode one block-scale byte to f32 (E8M0 in MXFP4 mode, else E4M3).""" @@ -820,7 +851,7 @@ def configure( trellis_chunks = max(1, n // 128) if n % 128: raise ValueError( - "trellis_t256 micro weights require intermediate % 128 == 0" + "trellis3_t256 micro weights require intermediate % 128 == 0" ) num_fc1_chunks = trellis_chunks elif self.a8_mx_mode: @@ -894,6 +925,7 @@ def configure( self.m1_fc2_rows_per_cta = m1_fc2_rows self.launch_block_dim = _K_PER_CTA * 16 if m1_half_cta_fc2 else _BLOCK_DIM self.grid_x = grid_x + self.route_slots = m * cfg.num_topk # Required inter_fp32 element count: the per-token intermediate plus # the trellis K-split scratch tail. Callers must zero-initialize; the # kernel restores the tail to zero after each use. @@ -933,6 +965,7 @@ def _m1_fc2_rowpair_narrow( w2_alphas: cute.Tensor, topk_ids: cute.Tensor, topk_weights: cute.Tensor, + route_expert_limit: Int32, scatter_output: cute.Tensor, ): cfg = self._cfg @@ -965,8 +998,9 @@ def _m1_fc2_rowpair_narrow( for kk in cutlass.range_constexpr(cfg.num_topk): eid_addr = Int32(kk) - eid = Int32(topk_ids[eid_addr]) - router_w = topk_weights[eid_addr] + eid, router_w = self._fc2_route( + topk_ids, topk_weights, eid_addr, route_expert_limit + ) if cutlass.const_expr( self.w4a16_mode and (not self.is_gated) @@ -1165,6 +1199,7 @@ def _m1_fc2_rowpair_wide( w2_alphas: cute.Tensor, topk_ids: cute.Tensor, topk_weights: cute.Tensor, + route_expert_limit: Int32, scatter_output: cute.Tensor, ): cfg = self._cfg @@ -1190,8 +1225,9 @@ def _m1_fc2_rowpair_wide( for kk in cutlass.range_constexpr(cfg.num_topk): eid_addr = Int32(kk) - eid = Int32(topk_ids[eid_addr]) - router_w = topk_weights[eid_addr] + eid, router_w = self._fc2_route( + topk_ids, topk_weights, eid_addr, route_expert_limit + ) if cutlass.const_expr( self.w4a16_mode and (not self.is_gated) @@ -1405,6 +1441,7 @@ def _m2_fc2_rowquad_narrow( w2_alphas: cute.Tensor, topk_ids: cute.Tensor, topk_weights: cute.Tensor, + route_expert_limit: Int32, scatter_output: cute.Tensor, ): cfg = self._cfg @@ -1445,8 +1482,9 @@ def _m2_fc2_rowquad_narrow( for kk in cutlass.range_constexpr(cfg.num_topk): eid_addr = t * Int32(cfg.num_topk) + Int32(kk) - eid = Int32(topk_ids[eid_addr]) - router_w = topk_weights[eid_addr] + eid, router_w = self._fc2_route( + topk_ids, topk_weights, eid_addr, route_expert_limit + ) if cutlass.const_expr( self.w4a16_mode and (not self.is_gated) @@ -1636,6 +1674,7 @@ def _m2_fc2_rowpair_narrow( w2_alphas: cute.Tensor, topk_ids: cute.Tensor, topk_weights: cute.Tensor, + route_expert_limit: Int32, scatter_output: cute.Tensor, ): cfg = self._cfg @@ -1673,8 +1712,9 @@ def _m2_fc2_rowpair_narrow( for kk in cutlass.range_constexpr(cfg.num_topk): eid_addr = t * Int32(cfg.num_topk) + Int32(kk) - eid = Int32(topk_ids[eid_addr]) - router_w = topk_weights[eid_addr] + eid, router_w = self._fc2_route( + topk_ids, topk_weights, eid_addr, route_expert_limit + ) if cutlass.const_expr( self.w4a16_mode and (not self.is_gated) @@ -1870,6 +1910,7 @@ def _m2_fc2_rowquad_wide( w2_alphas: cute.Tensor, topk_ids: cute.Tensor, topk_weights: cute.Tensor, + route_expert_limit: Int32, scatter_output: cute.Tensor, ): cfg = self._cfg @@ -1919,8 +1960,9 @@ def _m2_fc2_rowquad_wide( for kk in cutlass.range_constexpr(cfg.num_topk): eid_addr = t * Int32(cfg.num_topk) + Int32(kk) - eid = Int32(topk_ids[eid_addr]) - router_w = topk_weights[eid_addr] + eid, router_w = self._fc2_route( + topk_ids, topk_weights, eid_addr, route_expert_limit + ) if cutlass.const_expr( self.w4a16_mode and (not self.is_gated) @@ -2393,12 +2435,44 @@ def kernel( barrier_epoch: cute.Tensor, trellis_lut: cute.Tensor, trellis_rotations: cute.Tensor, + route_expert_limit: Int32, m_val: Int32, ): cfg = self._cfg bidx_x, _, _ = cute.arch.block_idx() tidx, _, _ = cute.arch.thread_idx() gdim_x, _, _ = cute.arch.grid_dim() + if cutlass.const_expr(self.stage_inactive_routes): + # The native direct path bypasses dynamic route packing, so stage + # its bounded route table once per CTA. Invalid IDs receive an + # addressable placeholder and an exact-zero effective weight. + # Inner FC1/FC2 loops then retain their unchecked hot path while + # caller-owned route tensors remain unchanged across graph replay. + staged_ids_ptr = cute.arch.alloc_smem(Int32, self.route_slots) + staged_weights_ptr = cute.arch.alloc_smem(Float32, self.route_slots) + staged_ids = cute.make_tensor( + staged_ids_ptr, cute.make_layout(self.route_slots) + ) + staged_weights = cute.make_tensor( + staged_weights_ptr, cute.make_layout(self.route_slots) + ) + route_slot = Int32(tidx) + while route_slot < Int32(self.route_slots): + raw_expert = Int64(topk_ids[route_slot]) + expert = Int32(0) + weight = Float32(0.0) + if ( + raw_expert >= Int64(0) + and raw_expert < Int64(route_expert_limit) + ): + expert = Int32(raw_expert) + weight = Float32(topk_weights[route_slot]) + staged_ids[route_slot] = expert + staged_weights[route_slot] = weight + route_slot += Int32(self.launch_block_dim) + cute.arch.sync_threads() + topk_ids = staged_ids + topk_weights = staged_weights if cutlass.const_expr(self.compile_time_phase == 2): self._run_fc2( Int32(bidx_x), @@ -2412,6 +2486,7 @@ def kernel( intermediate, topk_ids, topk_weights, + route_expert_limit, scatter_output, trellis_lut, ) @@ -5500,6 +5575,7 @@ def kernel( intermediate, topk_ids, topk_weights, + route_expert_limit, scatter_output, trellis_lut, ) @@ -5631,6 +5707,7 @@ def _run_fc2( intermediate: cute.Tensor, topk_ids: cute.Tensor, topk_weights: cute.Tensor, + route_expert_limit: Int32, scatter_output: cute.Tensor, trellis_lut: cute.Tensor, ): @@ -5674,6 +5751,7 @@ def _run_fc2( w2_alphas, topk_ids, topk_weights, + route_expert_limit, scatter_output, ) else: @@ -5687,6 +5765,7 @@ def _run_fc2( w2_alphas, topk_ids, topk_weights, + route_expert_limit, scatter_output, ) else: @@ -5703,6 +5782,7 @@ def _run_fc2( w2_alphas, topk_ids, topk_weights, + route_expert_limit, scatter_output, ) else: @@ -5716,6 +5796,7 @@ def _run_fc2( w2_alphas, topk_ids, topk_weights, + route_expert_limit, scatter_output, ) fc2_task += Int32(gdim_x) @@ -5740,6 +5821,7 @@ def _run_fc2( w2_alphas, topk_ids, topk_weights, + route_expert_limit, scatter_output, ) elif cutlass.const_expr(cfg.fc2_n_chunks > 1): @@ -5753,6 +5835,7 @@ def _run_fc2( w2_alphas, topk_ids, topk_weights, + route_expert_limit, scatter_output, ) else: @@ -5766,6 +5849,7 @@ def _run_fc2( w2_alphas, topk_ids, topk_weights, + route_expert_limit, scatter_output, ) fc2_task += Int32(gdim_x) @@ -5788,6 +5872,7 @@ def __call__( out_ptr: cute.Pointer, barrier_count: cute.Tensor, barrier_epoch: cute.Tensor, + route_expert_limit: Int32, m_val: Int32, grid_x: Int32, stream, @@ -5888,6 +5973,7 @@ def __call__( if cutlass.const_expr(trellis_rot_ptr is not None) else barrier_count ), + route_expert_limit, m_val, ).launch( grid=(grid_x, Int32(1), Int32(1)), @@ -5950,6 +6036,7 @@ def ptr(dt, t): ptr(cutlass.BFloat16, out), barrier_count, barrier_epoch, + Int32(w1_alphas.numel()), Int32(m), Int32(grid_x), stream, diff --git a/b12x/moe/_shared/kernels/situ.py b/b12x/moe/_shared/kernels/situ.py index 991c37b47..e2c162cd3 100644 --- a/b12x/moe/_shared/kernels/situ.py +++ b/b12x/moe/_shared/kernels/situ.py @@ -18,10 +18,10 @@ class MoEMicroKernelSitu(MoEMicroKernelBackend): The compact NVFP4 intermediate quantization has not passed its correctness gate, so ``is_supported`` stays False and the nvfp4-family - dispatch never selects this class. The trellis_t256 arm quantizes its + dispatch never selects this class. The trellis3_t256 arm quantizes its intermediate through per-32 UE8M0 E4M3 (a8_mx) instead; the w4a8_mx trellis band dispatches it explicitly and constructs this class with - ``weight_layout="trellis_t256"``. + ``weight_layout="trellis3_t256"``. """ @classmethod @@ -37,9 +37,9 @@ def is_supported( return False def __init__(self, *args: object, **kwargs: object): - if kwargs.get("weight_layout") != "trellis_t256": + if kwargs.get("weight_layout") != "trellis3_t256": raise NotImplementedError( - "SiTU compact micro serves only the trellis_t256 weight " + "SiTU compact micro serves only the trellis3_t256 weight " "layout; select the fused dynamic kernel" ) kwargs["activation"] = SITU diff --git a/b12x/moe/_shared/kernels/tiny_decode.py b/b12x/moe/_shared/kernels/tiny_decode.py index 0cf0ca4aa..8d522a63e 100644 --- a/b12x/moe/_shared/kernels/tiny_decode.py +++ b/b12x/moe/_shared/kernels/tiny_decode.py @@ -231,14 +231,7 @@ def kernel( rem = bidx % fc1_per_rt nt = rem // Int32(c["fc1_ktg"]) ktg = rem % Int32(c["fc1_ktg"]) - route_eid = Int32(topk_ids[rt_idx]) - route_active = route_eid >= Int32(0) and route_eid < Int32(c["weight_E"]) - # vLLM represents CUDA-graph padding and non-local EP routes with - # an expert id outside [0, weight_E). Keep every address in range; - # inactive routes retain the pre-zeroed intermediate and output. - eid = Int64(0) - if route_active: - eid = Int64(route_eid) + eid = Int64(Int32(topk_ids[rt_idx])) tok = rt_idx // Int32(c["num_topk"]) we_base = w13_base + eid * Int64(c["w13_words"] * 4) se_base = sfb13_base + eid * Int64(c["sfb13_bytes"]) @@ -299,7 +292,7 @@ def kernel( acc1 = warp_reduce(acc1, lambda a, b: a + b, width=4) acc2 = warp_reduce(acc2, lambda a, b: a + b, width=4) acc3 = warp_reduce(acc3, lambda a, b: a + b, width=4) - if route_active and cgrp == Int32(0): + if cgrp == Int32(0): ibase_rt = inter_base + Int64(rt_idx) * Int64(c["two_n"] * 4) accs = (acc0, acc1, acc2, acc3) for v in cutlass.range_constexpr(4): @@ -333,11 +326,7 @@ def kernel( rem = bidx % fc2_per_rt nt = rem // Int32(c["fc2_ktg"]) ktg = rem % Int32(c["fc2_ktg"]) - route_eid = Int32(topk_ids[rt_idx]) - route_active = route_eid >= Int32(0) and route_eid < Int32(c["weight_E"]) - eid = Int64(0) - if route_active: - eid = Int64(route_eid) + eid = Int64(Int32(topk_ids[rt_idx])) tok = rt_idx // Int32(c["num_topk"]) rw = Float32(topk_weights[rt_idx]) we_base = w2_base + eid * Int64(c["w2_words"] * 4) @@ -444,7 +433,7 @@ def kernel( o1 = cute.arch.shuffle_sync_bfly(acc1, offset=4) o2 = cute.arch.shuffle_sync_bfly(acc2, offset=4) o3 = cute.arch.shuffle_sync_bfly(acc3, offset=4) - if route_active and cgrp == Int32(0): + if cgrp == Int32(0): if (r8 % Int32(2)) == Int32(0): ob = out_base + Int64(tok) * Int64(c["k"] * 2) accs = (acc0, acc1, acc2, acc3) diff --git a/b12x/moe/_shared/kernels/trellis_decode.py b/b12x/moe/_shared/kernels/trellis_decode.py new file mode 100644 index 000000000..6d0da0f05 --- /dev/null +++ b/b12x/moe/_shared/kernels/trellis_decode.py @@ -0,0 +1,212 @@ +"""Shared SQG-XOR-Cheb-T12 native-tile decode for W4A8 kernels. + +The mixin extracts overlapping L16 trellis windows from the packed native +payload with the t256 ring geometry, decodes them through the modal T12 +staircase, and maps the resulting E4M3 bytes into ``m16n8k32`` MMA B +fragments (including the butterfly pairing of adjacent N16 tiles). Kernel +classes that inherit it provide the payload addressing (``_tile_base``-style +helpers) against their own task coordinates. +""" + +from __future__ import annotations + +import cutlass +import cutlass.cute as cute +from cutlass.cutlass_dsl import Int32, Int64, Uint32 + +from b12x._lib.intrinsics import ( + packed_decode_sqg_xor_cheb_t12_to_e4m3x8, +) + +_PAIR_U32_PER_K16 = 384 # 8 N16 tiles * 8 u32/bit * (low_bits + high_bits=6) + + +class _NativeTrellisDecode: + """Shared native-tile window extraction and E4M3 fragment mapping.""" + + @cute.jit + def _decode_windows( + self, + win_a: Uint32, + win_b: Uint32, + bits: cutlass.Constexpr[int], + rank_lut_addr: Int64, + ): + return packed_decode_sqg_xor_cheb_t12_to_e4m3x8( + win_a, + win_b, + rank_lut_addr, + int(bits), + ) + + @cute.jit + def _lane_geom(self, lane: Int32, bits: cutlass.Constexpr[int]): + bits_i32 = Int32(int(bits)) + ring_u32 = Int32(8 * int(bits)) + t_offset = Int32(8) * lane + b1 = (t_offset + Int32(257)) * bits_i32 + b0 = b1 - Int32(16) + b2 = b1 + Int32(7 * int(bits)) + i0 = b0 >> Int32(5) + i2 = (b2 - Int32(1)) >> Int32(5) + ia = i0 - ring_u32 * (i0 >= ring_u32).to(Int32) + ib = i2 - ring_u32 * (i2 >= ring_u32).to(Int32) + s2 = (i2 + Int32(1)) * Int32(32) - b2 + return ia, ib, s2 + + @cute.jit + def _decode_tile_at( + self, + trellis: cute.Tensor, + lane: Int32, + base: Int64, + n_high: Int32, + bits: cutlass.Constexpr[int], + rank_lut_smem_addr: Int64, + ) -> Uint32: + ia, ib, s2 = self._lane_geom(lane, bits) + a = Uint32(trellis[base + Int64(ia)]) + b = Uint32(trellis[base + Int64(ib)]) + merged = (Int64(a) << Int64(32)) | Int64(b) + win_a = Uint32(merged >> Int64(s2)) + win_b = Uint32(merged >> Int64(s2 + Int32(4 * int(bits)))) + lo, hi = self._decode_windows( + win_a, + win_b, + bits, + rank_lut_smem_addr, + ) + value = lo + if n_high != Int32(0): + value = hi + return value + + @cute.jit + def _decode_tile_both( + self, + trellis: cute.Tensor, + lane: Int32, + base: Int64, + bits: cutlass.Constexpr[int], + rank_lut_smem_addr: Int64, + ): + """Decode both adjacent N8 halves of one native N16 tile once.""" + + ia, ib, s2 = self._lane_geom(lane, bits) + a = Uint32(trellis[base + Int64(ia)]) + b = Uint32(trellis[base + Int64(ib)]) + merged = (Int64(a) << Int64(32)) | Int64(b) + win_a = Uint32(merged >> Int64(s2)) + win_b = Uint32(merged >> Int64(s2 + Int32(4 * int(bits)))) + return self._decode_windows( + win_a, + win_b, + bits, + rank_lut_smem_addr, + ) + + @cute.jit + def _decode_tile_both_geom( + self, + trellis: cute.Tensor, + base: Int64, + ia: Int32, + ib: Int32, + s2: Int32, + bits: cutlass.Constexpr[int], + rank_lut_smem_addr: Int64, + ): + a = Uint32(trellis[base + Int64(ia)]) + b = Uint32(trellis[base + Int64(ib)]) + merged = (Int64(a) << Int64(32)) | Int64(b) + win_a = Uint32(merged >> Int64(s2)) + win_b = Uint32(merged >> Int64(s2 + Int32(4 * int(bits)))) + return self._decode_windows( + win_a, + win_b, + bits, + rank_lut_smem_addr, + ) + + @cute.jit + def _pair_native_words( + self, + trellis: cute.Tensor, + lane: Int32, + base0: Int64, + base1: Int64, + n_high: Int32, + bits: cutlass.Constexpr[int], + rank_lut_smem_addr: Int64, + ): + e0 = self._decode_tile_at( + trellis, lane, base0, n_high, bits, rank_lut_smem_addr + ) + e1 = self._decode_tile_at( + trellis, lane, base1, n_high, bits, rank_lut_smem_addr + ) + c = lane & Int32(3) + own = e0 + send = e1 + if c >= Int32(2): + own = e1 + send = e0 + peer = Uint32(cute.arch.shuffle_sync_bfly(send, offset=2)) + return own, peer + + @cute.jit + def _pair_native_words_both( + self, + trellis: cute.Tensor, + lane: Int32, + base0: Int64, + base1: Int64, + bits: cutlass.Constexpr[int], + rank_lut_smem_addr: Int64, + ): + ia, ib, s2 = self._lane_geom(lane, bits) + e0_lo, e0_hi = self._decode_tile_both_geom( + trellis, base0, ia, ib, s2, bits, rank_lut_smem_addr + ) + e1_lo, e1_hi = self._decode_tile_both_geom( + trellis, base1, ia, ib, s2, bits, rank_lut_smem_addr + ) + c = lane & Int32(3) + own_lo = e0_lo + own_hi = e0_hi + send_lo = e1_lo + send_hi = e1_hi + if c >= Int32(2): + own_lo = e1_lo + own_hi = e1_hi + send_lo = e0_lo + send_hi = e0_hi + peer_lo = Uint32(cute.arch.shuffle_sync_bfly(send_lo, offset=2)) + peer_hi = Uint32(cute.arch.shuffle_sync_bfly(send_hi, offset=2)) + return own_lo, peer_lo, own_hi, peer_hi + + @cute.jit + def _a_fragment( + self, + values: cute.Tensor, + scale_rows: cute.Tensor, + route: Int32, + k32: Int32, + lane: Int32, + ): + c = lane & Int32(3) + g = lane >> Int32(2) + a0 = Uint32(0) + a1 = Uint32(0) + a2 = Uint32(0) + a3 = Uint32(0) + if g == Int32(0): + word0 = k32 * Int32(8) + c * Int32(2) + a0 = Uint32(values[route, word0]) + a2 = Uint32(values[route, word0 + Int32(1)]) + + sf = Uint32(127) + if g == Int32(0): + if (lane & Int32(1)) == Int32(0): + sf = Uint32(scale_rows[route, k32]) + return a0, a1, a2, a3, sf * Uint32(0x01010101) diff --git a/b12x/moe/_shared/kernels/trellis_ring.py b/b12x/moe/_shared/kernels/trellis_ring.py deleted file mode 100644 index c9f63167e..000000000 --- a/b12x/moe/_shared/kernels/trellis_ring.py +++ /dev/null @@ -1,45 +0,0 @@ -"""t256 tail-biting ring window geometry shared by the W4A16 and W4A8 kernels. - -One ``trellis_t256`` tile packs 256 codes into a ring of ``8*bits`` uint32 -words (``256*bits`` bits total). Weight ``j`` of lane ``l`` occupies the -16-bit window ending at bit ``(8*l + j + 257) * bits``, taken modulo the -ring; consecutive weights overlap by ``16 - bits`` bits. The geometry -depends only on the lane id, the weight span, and the bitrate, so callers -hoist it out of their decode loops and feed the resulting word indices to a -64-bit funnel shift. -""" - -from __future__ import annotations - -import cutlass -import cutlass.cute as cute -from cutlass.cutlass_dsl import Int32 - - -@cute.jit -def trellis256_lane_geom_bits( - lane: Int32, - weight_offset: cutlass.Constexpr[int], - weight_count: cutlass.Constexpr[int], - bits: cutlass.Constexpr[int], -): - """Ring geometry for ``weight_count`` weights starting at ``weight_offset``. - - Returns ``(ia, ib, s2, span)``: the ring word indices of the first and - last 32-bit words covering the lane's windows, the funnel shift that - aligns the merged 64-bit read on the final window, and the word span - ``i2 - i0`` before ring wrap-around. - """ - - bits_i32 = Int32(int(bits)) - ring_u32 = Int32(8 * int(bits)) - t_offset = Int32(8) * lane + Int32(weight_offset) - b1 = (t_offset + Int32(257)) * bits_i32 - b0 = b1 - Int32(16) - b2 = b1 + Int32((int(weight_count) - 1) * int(bits)) - i0 = b0 >> Int32(5) - i2 = (b2 - Int32(1)) >> Int32(5) - ia = i0 - ring_u32 * (i0 >= ring_u32).to(Int32) - ib = i2 - ring_u32 * (i2 >= ring_u32).to(Int32) - s2 = (i2 + Int32(1)) * Int32(32) - b2 - return ia, ib, s2, i2 - i0 diff --git a/b12x/moe/_shared/kernels/w4a16/btx.py b/b12x/moe/_shared/kernels/w4a16/btx.py deleted file mode 100644 index f0ac004e3..000000000 --- a/b12x/moe/_shared/kernels/w4a16/btx.py +++ /dev/null @@ -1,640 +0,0 @@ -"""BTX container reading and W4A16 MoE weight preparation. - -One metadata-driven load path serves every declared configuration of the -`btx-atoms-v1` container: the manifest and rate tables locate every byte, -the extent rules come from the manifest's layout section, and preparation -reuses the shared trellis machinery (`prepare_trellis256_moe_weights` -for uniform rate structures, the pair finalizer for per-expert ones). -Nothing in this module depends on a specific model geometry, profile -name, or rate-placement convention. -""" - -from __future__ import annotations - -import hashlib -import json -import pathlib -from dataclasses import dataclass, replace - -import torch - -from b12x.moe._shared.btx_schema import ( - ATOMS_PER_PAIR, - BTX_MANIFEST_FILENAME, - BTX_SCHEMA, - BtxManifest, - RATE_CODE_PAIR_KINDS, - RATE_STRUCTURE_PER_EXPERT_PAIR, - RATE_STRUCTURE_UNIFORM, - matrix_atom_bytes, - rate_code, - rate_code_bits, -) -from b12x.moe._shared.kernels.w4a16.prepare import ( - PreparedW4A16MoeWeights, - _finalize_prepared_trellis_weights, - _coupled_rotation_signs, - _restore_plane_words, - prepare_trellis256_moe_weights, -) - -_EXPERT_CHUNK = 64 - - -@dataclass(frozen=True) -class BtxLayer: - """One layer's extent-sliced content, validated against the manifest.""" - - manifest: BtxManifest - layer_index: int - first_slot: int - slot_count: int - atoms: torch.Tensor - rotations: torch.Tensor - gate_suh: torch.Tensor - up_suh: torch.Tensor - down_svh: torch.Tensor - rates_fc1: torch.Tensor | None - rates_fc2: torch.Tensor | None - rotation_draws: torch.Tensor | None - - @property - def local_intermediate_size(self) -> int: - return self.slot_count * self.manifest.geometry.atom_channels - - -def read_btx_manifest(root: str | pathlib.Path) -> BtxManifest: - root = pathlib.Path(root) - data = json.loads((root / BTX_MANIFEST_FILENAME).read_text()) - return BtxManifest.from_dict(data) - - -def read_btx_layer( - root: str | pathlib.Path, - manifest: BtxManifest, - layer_index: int, - *, - first_slot: int, - slot_count: int, - verify_sha: bool = False, -) -> BtxLayer: - """Load one rank extent of one layer as CPU tensors. - - The extent is validated against the manifest's layout declarations and - the safetensors metadata is cross-checked against the manifest before - any tensor is interpreted. - """ - - from safetensors import safe_open - - root = pathlib.Path(root) - manifest.validate_extent(first_slot, slot_count) - if layer_index not in manifest.layers: - raise ValueError(f"BTX manifest does not declare layer {layer_index}") - ref = manifest.layers[layer_index] - path = root / ref.file - if verify_sha: - digest = hashlib.sha256(path.read_bytes()).hexdigest() - if digest != ref.sha256: - raise ValueError( - f"BTX layer {layer_index} sha256 mismatch: manifest " - f"{ref.sha256}, file {digest}" - ) - - geometry = manifest.geometry - per_expert = manifest.rates.structure == RATE_STRUCTURE_PER_EXPERT_PAIR - with safe_open(str(path), framework="pt") as handle: - metadata = handle.metadata() or {} - expected = { - "schema": BTX_SCHEMA, - "codebook": manifest.codebook, - "layer": str(int(layer_index)), - "num_experts": str(geometry.num_experts), - "hidden_size": str(geometry.hidden_size), - "intermediate_size": str(geometry.intermediate_size), - "atom_channels": str(geometry.atom_channels), - } - for key, value in expected.items(): - if metadata.get(key) != value: - raise ValueError( - f"BTX layer {layer_index} metadata {key!r} is " - f"{metadata.get(key)!r}; the manifest declares {value!r}" - ) - names = set(handle.keys()) - required = {"atoms", "rotations", "gate_suh", "up_suh", "down_svh"} - if per_expert: - required |= {"rates_fc1", "rates_fc2"} - if manifest.hadamard.coupled: - required |= {"rotation_draws"} - if names != required: - raise ValueError( - f"BTX layer {layer_index} tensors {sorted(names)} do not " - f"match the declared set {sorted(required)}" - ) - - atoms_slice = handle.get_slice("atoms") - atoms_shape = atoms_slice.get_shape() - if ( - len(atoms_shape) != 2 - or atoms_shape[0] != geometry.atom_slots - or atoms_shape[1] % manifest.layout.atom_row_alignment - ): - raise ValueError( - f"BTX layer {layer_index} atoms shape {atoms_shape} violates " - "the declared geometry or row alignment" - ) - atoms = atoms_slice[first_slot : first_slot + slot_count] - - rotations_slice = handle.get_slice("rotations") - if tuple(rotations_slice.get_shape()) != ( - geometry.atom_slots, - geometry.num_experts, - 3, - geometry.atom_channels, - ): - raise ValueError( - f"BTX layer {layer_index} rotations shape is not " - "[atom_slots, num_experts, 3, atom_channels]" - ) - rotations = rotations_slice[first_slot : first_slot + slot_count] - - h_shapes = ( - (geometry.num_experts, geometry.hidden_size) - if manifest.hadamard.per_expert_input_rotations - else (geometry.hidden_size,) - ) - sides = {} - for name in ("gate_suh", "up_suh", "down_svh"): - tensor = handle.get_tensor(name) - if tuple(tensor.shape) != h_shapes or tensor.dtype != torch.float16: - raise ValueError( - f"BTX layer {layer_index} {name} must be fp16 {h_shapes}" - ) - sides[name] = tensor - - rates_fc1 = rates_fc2 = None - if per_expert: - pairs = geometry.atom_slots // ATOMS_PER_PAIR - first_pair = first_slot // ATOMS_PER_PAIR - pair_count = slot_count // ATOMS_PER_PAIR - declared = manifest.rates.pair_kinds or frozenset() - observed: set[str] = set() - tables = {} - for name in ("rates_fc1", "rates_fc2"): - table_slice = handle.get_slice(name) - if tuple(table_slice.get_shape()) != ( - pairs, - geometry.num_experts, - ): - raise ValueError( - f"BTX layer {layer_index} {name} must be " - "[atom_slots/8, num_experts]" - ) - table = table_slice[first_pair : first_pair + pair_count] - for code in table.unique().tolist(): - kind = RATE_CODE_PAIR_KINDS.get(int(code)) - if kind is None: - raise ValueError( - f"BTX layer {layer_index} {name} contains " - f"unknown rate code {int(code):#x}" - ) - observed.add(kind) - tables[name] = table - if not observed <= set(declared): - raise ValueError( - f"BTX layer {layer_index} rate tables use kinds " - f"{sorted(observed)} outside the declared " - f"{sorted(declared)}" - ) - rates_fc1, rates_fc2 = tables["rates_fc1"], tables["rates_fc2"] - - rotation_draws = None - if manifest.hadamard.coupled: - rotation_draws = handle.get_tensor("rotation_draws") - if ( - tuple(rotation_draws.shape) != (geometry.num_experts,) - or rotation_draws.dtype != torch.uint8 - or bool(torch.any(rotation_draws > 7)) - ): - raise ValueError( - f"BTX layer {layer_index} rotation_draws must be " - "uint8[num_experts] in 0..7" - ) - - return BtxLayer( - manifest=manifest, - layer_index=layer_index, - first_slot=first_slot, - slot_count=slot_count, - atoms=atoms, - rotations=rotations, - gate_suh=sides["gate_suh"], - up_suh=sides["up_suh"], - down_svh=sides["down_svh"], - rates_fc1=rates_fc1, - rates_fc2=rates_fc2, - rotation_draws=rotation_draws, - ) - - -def _extent_rotation_tables( - layer: BtxLayer, device: torch.device -) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - """suh/svh device tensors plus [E, 3*I_local] boundary values.""" - - experts = layer.manifest.geometry.num_experts - local = layer.local_intermediate_size - values = layer.rotations.to(device=device) - columns = [ - values[:, :, matrix, :].permute(1, 0, 2).reshape(experts, local) - for matrix in range(3) - ] - intermediate = torch.cat(columns, dim=1).contiguous() - - def _side(tensor: torch.Tensor) -> torch.Tensor: - moved = tensor.to(device=device) - if moved.dim() == 1: - moved = moved.reshape(1, -1) - return moved.contiguous() - - return ( - _side(layer.gate_suh), - _side(layer.up_suh), - _side(layer.down_svh), - intermediate, - ) - - -def _coupled_rotation_rows( - layer: BtxLayer, intermediate: torch.Tensor, device: torch.device -) -> torch.Tensor: - """Append the frozen coupled sign rows: [values 3I | pre 2I | post I].""" - - manifest = layer.manifest - if (manifest.hadamard.pre_block, manifest.hadamard.post_block) != ( - 512, - 128, - ): - raise ValueError( - "coupled BTX preparation currently implements pre/post Hadamard " - "blocks (512, 128); the manifest declares " - f"({manifest.hadamard.pre_block}, {manifest.hadamard.post_block})" - ) - assert layer.rotation_draws is not None - experts = manifest.geometry.num_experts - global_i = manifest.geometry.intermediate_size - local = layer.local_intermediate_size - pre_begin = 2 * layer.first_slot * manifest.geometry.atom_channels - post_begin = layer.first_slot * manifest.geometry.atom_channels - signs = torch.empty((experts, 3 * local), dtype=torch.float16) - draws = layer.rotation_draws - for draw in sorted(set(int(value) for value in draws.tolist())): - rows = torch.nonzero(draws == draw, as_tuple=False).flatten() - pre = _coupled_rotation_signs(2 * global_i, draw=draw, axis=1)[ - pre_begin : pre_begin + 2 * local - ] - post = _coupled_rotation_signs(global_i, draw=draw, axis=2)[ - post_begin : post_begin + local - ] - signs.index_copy_( - 0, - rows, - torch.cat((pre, post)).to(torch.float16).expand(rows.numel(), -1), - ) - return torch.cat( - (intermediate, signs.to(device=device)), dim=1 - ).contiguous() - - -def _uniform_native_tensors( - layer: BtxLayer, device: torch.device -) -> tuple[torch.Tensor, torch.Tensor]: - """Assemble native trellis_t256 tensors from a uniform-rate extent. - - Returns projection-major FC1 ``[2, E, H/16, I_local/16, 16*bits]`` and - FC2 ``[E, I_local/16, H/16, 16*bits]`` int16 tensors. - """ - - manifest = layer.manifest - geometry = manifest.geometry - bits = manifest.rates.bits - assert bits is not None - experts = geometry.num_experts - hidden_tiles = geometry.hidden_size // 16 - slots = layer.slot_count - section = matrix_atom_bytes(geometry.hidden_size, bits, bits) - payload = experts * 3 * section - if layer.atoms.shape[1] < payload: - raise ValueError("BTX atoms rows are shorter than their expert bundles") - if bool(torch.any(layer.atoms[:, payload:] != 0)): - raise ValueError("BTX atoms row padding must be zero") - - w13 = torch.empty( - (2, experts, hidden_tiles, 2 * slots, 16 * bits), - dtype=torch.int16, - device=device, - ) - w2 = torch.empty( - (experts, 2 * slots, hidden_tiles, 16 * bits), - dtype=torch.int16, - device=device, - ) - bundles = layer.atoms[:, :payload].reshape(slots, experts, 3 * section) - for first in range(0, experts, _EXPERT_CHUNK): - count = min(_EXPERT_CHUNK, experts - first) - chunk = ( - bundles[:, first : first + count] - .contiguous() - .to(device=device) - .view(torch.int16) - .reshape(slots, count, 3, 2, hidden_tiles, 16 * bits) - ) - for matrix in range(2): - # FC1 planes are the atom's two consecutive N16 columns. - w13[matrix, first : first + count].copy_( - chunk[:, :, matrix].permute(1, 3, 0, 2, 4).reshape( - count, hidden_tiles, 2 * slots, 16 * bits - ) - ) - # FC2 planes are the atom's two consecutive K16 rows. - w2[first : first + count].copy_( - chunk[:, :, 2].permute(1, 0, 2, 3, 4).reshape( - count, 2 * slots, hidden_tiles, 16 * bits - ) - ) - return w13, w2 - - -def prepare_btx_moe_weights( - layer: BtxLayer, - *, - activation: str, - device: torch.device | str, - params_dtype: torch.dtype = torch.float16, - tile_config: tuple[int, int, int, int] | None = None, - dummy_scale: torch.Tensor | None = None, - workspace: torch.Tensor | None = None, -) -> PreparedW4A16MoeWeights: - """Prepare one BTX rank extent for the fused W4A16 serving path.""" - - manifest = layer.manifest - device = torch.device(device) - gate_suh, up_suh, down_svh, intermediate = _extent_rotation_tables( - layer, device - ) - rotations = intermediate - if manifest.hadamard.coupled: - rotations = _coupled_rotation_rows(layer, intermediate, device) - # The coupled residual transform interleaves gate/up into one - # length-2I axis whose two stored halves each carry one input-side - # table; both physical FC1 slots of a rank use the half its extent - # lies in. - pre_half_slots = manifest.geometry.atom_slots // 2 - source_suh = gate_suh if layer.first_slot < pre_half_slots else up_suh - gate_suh = source_suh - up_suh = source_suh - - if manifest.rates.structure == RATE_STRUCTURE_UNIFORM: - assert manifest.rates.bits is not None - w13, w2 = _uniform_native_tensors(layer, device) - if tile_config is None: - tile_config = ( - (128, 128, 128, 128) - if manifest.hadamard.coupled and manifest.rates.bits == 2 - else (64, 256, 64, 256) - ) - prepared = prepare_trellis256_moe_weights( - w13=w13, - w2=w2, - hidden_size=manifest.geometry.hidden_size, - intermediate_size=layer.local_intermediate_size, - num_experts=manifest.geometry.num_experts, - activation=activation, - fc1_tile_n=tile_config[1], - fc2_tile_n=tile_config[3], - device=device, - params_dtype=params_dtype, - w13_layout="trellis_t256_proj", - trellis_bits=manifest.rates.bits, - codebook=manifest.codebook, - gate_suh=gate_suh, - up_suh=up_suh, - intermediate_rotations=intermediate, - down_svh=down_svh, - tile_config=tile_config, - dummy_scale=dummy_scale, - workspace=workspace, - ) - if not manifest.hadamard.coupled: - return prepared - assert prepared.trellis is not None - return replace( - prepared, - trellis=replace( - prepared.trellis, - coupled_hadamard=True, - intermediate_rotations=rotations, - ), - ) - - return _prepare_btx_pair_extent( - layer, - device=device, - gate_suh=gate_suh, - up_suh=up_suh, - down_svh=down_svh, - rotations=rotations, - params_dtype=params_dtype, - tile_config=tile_config or (64, 256, 64, 256), - dummy_scale=dummy_scale, - workspace=workspace, - ) - - -def _prepare_btx_pair_extent( - layer: BtxLayer, - *, - device: torch.device, - gate_suh: torch.Tensor, - up_suh: torch.Tensor, - down_svh: torch.Tensor, - rotations: torch.Tensor, - params_dtype: torch.dtype, - tile_config: tuple[int, int, int, int], - dummy_scale: torch.Tensor | None, - workspace: torch.Tensor | None, -) -> PreparedW4A16MoeWeights: - """Prepare a per-expert-pair extent through the pair kernel machinery. - - The fused kernel's pair decode operates on one 256-channel pair per - rank (FC2 pairs lie on the local K axis), so a per-expert-rate extent - is exactly one pair of atom slots. - """ - - manifest = layer.manifest - geometry = manifest.geometry - if layer.slot_count != ATOMS_PER_PAIR: - raise ValueError( - "per-expert-pair BTX extents must cover exactly one " - f"256-channel pair ({ATOMS_PER_PAIR} atom slots); got " - f"{layer.slot_count}" - ) - if manifest.hadamard.coupled: - raise ValueError( - "coupled-Hadamard execution of per-expert-pair BTX extents has " - "no qualified kernel path" - ) - # The pair runtime orders each 256-channel pair record-major: every - # atom contributes its first 16 channels to the low record and its - # last 16 to the high record. Rotation rows must match that order. - experts_count = geometry.num_experts - per_matrix = [] - values = layer.rotations.to(rotations.device) - for matrix in range(3): - planes = values[:, :, matrix, :].reshape( - ATOMS_PER_PAIR, experts_count, 2, 16 - ) - low = planes[:, :, 0, :].permute(1, 0, 2).reshape(experts_count, -1) - high = planes[:, :, 1, :].permute(1, 0, 2).reshape(experts_count, -1) - per_matrix.append(torch.cat((low, high), dim=1)) - rotations = torch.cat(per_matrix, dim=1).contiguous() - assert layer.rates_fc1 is not None and layer.rates_fc2 is not None - fc1_codes = layer.rates_fc1[0].to(torch.int64) - fc2_codes = layer.rates_fc2[0].to(torch.int64) - experts = geometry.num_experts - hidden_tiles = geometry.hidden_size // 16 - - kinds = { - RATE_CODE_PAIR_KINDS[int(code)] - for code in torch.cat((fc1_codes, fc2_codes)).unique().tolist() - } - if kinds == {"P33"} or kinds == {"P33", "P24"}: - pair_kind = "PDYNAMIC" - high_code = rate_code(2, 4) - elif kinds == {"P33", "P43"}: - pair_kind = "P33_P43" - high_code = rate_code(4, 3) - else: - raise ValueError( - f"BTX per-expert extents with pair kinds {sorted(kinds)} have " - "no fused execution arm; whole-expert K4 tiers run through " - "mixed-tier or multi-launch execution" - ) - - def _restore(codes: torch.Tensor, matrix: int, *, fc1: bool): - sections = [] - for expert in range(experts): - low_bits, high_bits = rate_code_bits(int(codes[expert])) - begin = 0 - for m in range(matrix): - m_codes = fc1_codes if m < 2 else fc2_codes - lo, hi = rate_code_bits(int(m_codes[expert])) - begin += matrix_atom_bytes(geometry.hidden_size, lo, hi) - section = matrix_atom_bytes(geometry.hidden_size, low_bits, high_bits) - raw = layer.atoms[:, _bundle_offset(layer, expert) + begin :][ - :, :section - ] - words = ( - raw.contiguous() - .to(device=device) - .view(torch.int16) - .reshape(ATOMS_PER_PAIR, 1, -1) - .permute(1, 0, 2) - ) - low_words = hidden_tiles * 16 * low_bits - low = words[..., :low_words].reshape( - 1, ATOMS_PER_PAIR, hidden_tiles, 16 * low_bits - ) - high = words[..., low_words:].reshape( - 1, ATOMS_PER_PAIR, hidden_tiles, 16 * high_bits - ) - sections.append(_restore_plane_words(low, high, fc1=fc1)) - return sections - - if pair_kind == "PDYNAMIC": - fc1_modes = (fc1_codes == high_code).to(torch.int32).to(device) - fc2_modes = (fc2_codes == high_code).to(torch.int32).to(device) - gate = _restore(fc1_codes, 0, fc1=True) - up = _restore(fc1_codes, 1, fc1=True) - down = _restore(fc2_codes, 2, fc1=False) - w13 = torch.cat( - [torch.cat(gate, dim=0), torch.cat(up, dim=0)] - ).reshape(-1) - w2 = torch.cat(down, dim=0).reshape(-1) - fc1_pair_modes: torch.Tensor = fc1_modes.contiguous() - fc2_pair_modes: torch.Tensor = fc2_modes.contiguous() - else: - # Compact gap-free pools with per-expert descriptors, matching the - # fused kernel's P33_P43 addressing. - def _compact(codes, sections): - lengths = torch.tensor( - [section.numel() // 2 for section in sections], - dtype=torch.int64, - ) - offsets = torch.zeros_like(lengths) - offsets[1:] = torch.cumsum(lengths[:-1], dim=0) - modes = (codes == high_code).to(torch.int64) - descriptors = ((offsets << 1) | modes).to(device) - return offsets, descriptors - - gate = _restore(fc1_codes, 0, fc1=True) - up = _restore(fc1_codes, 1, fc1=True) - down = _restore(fc2_codes, 2, fc1=False) - _, fc1_descriptors = _compact(fc1_codes, gate) - _, fc2_descriptors = _compact(fc2_codes, down) - w13 = torch.cat( - [torch.cat(gate, dim=1), torch.cat(up, dim=1)] - ).reshape(-1) - w2 = torch.cat(down, dim=1).reshape(-1) - fc1_pair_modes = fc1_descriptors.contiguous() - fc2_pair_modes = fc2_descriptors.contiguous() - - return _finalize_prepared_trellis_weights( - context="BTX per-expert-pair preparation", - device=device, - hidden_size=geometry.hidden_size, - intermediate_size=layer.local_intermediate_size, - num_experts=experts, - params_dtype=params_dtype, - w13=w13, - w2=w2, - gate_suh=gate_suh, - up_suh=up_suh, - intermediate_rotations=rotations, - down_svh=down_svh, - rotation_columns=rotations.shape[1], - tile_config=tile_config, - required_fc1_tile_n=256, - dummy_scale=dummy_scale, - workspace=workspace, - codebook=manifest.codebook, - trellis_bits=3, - fc1_pair_kind=pair_kind, - fc2_pair_kind=pair_kind, - fc1_pair_modes=fc1_pair_modes, - fc2_pair_modes=fc2_pair_modes, - coupled_hadamard=manifest.hadamard.coupled, - ) - - -def _bundle_offset(layer: BtxLayer, expert: int) -> int: - """Byte offset of one expert's bundle within this extent's rows.""" - - manifest = layer.manifest - if manifest.rates.structure == RATE_STRUCTURE_UNIFORM: - assert manifest.rates.bits is not None - section = matrix_atom_bytes( - manifest.geometry.hidden_size, - manifest.rates.bits, - manifest.rates.bits, - ) - return expert * 3 * section - assert layer.rates_fc1 is not None and layer.rates_fc2 is not None - offset = 0 - for e in range(expert): - fc1_lo, fc1_hi = rate_code_bits(int(layer.rates_fc1[0, e])) - fc2_lo, fc2_hi = rate_code_bits(int(layer.rates_fc2[0, e])) - offset += 2 * matrix_atom_bytes( - manifest.geometry.hidden_size, fc1_lo, fc1_hi - ) + matrix_atom_bytes(manifest.geometry.hidden_size, fc2_lo, fc2_hi) - return offset diff --git a/b12x/moe/_shared/kernels/w4a16/btx_compat.py b/b12x/moe/_shared/kernels/w4a16/btx_compat.py deleted file mode 100644 index 9b9155f4b..000000000 --- a/b12x/moe/_shared/kernels/w4a16/btx_compat.py +++ /dev/null @@ -1,346 +0,0 @@ -"""Lift frozen QSRT atom containers into in-memory BTX extents. - -The `kquant_kimi_k3_qsrt_atoms_v1`/`_v2` containers derive rate placement -arithmetically and embed per-atom rotation spans inside expert bundles. -These lifts re-express one rank extent as a :class:`BtxLayer` — explicit -rate tables, a separate rotations tensor, pure-code-word bundles, and a -synthesized fail-closed manifest — so the declarative BTX preparation path -serves the frozen containers unchanged. - -This module is the only remaining holder of the frozen containers' layout -knowledge (the ``(5*expert + layer) % 12`` pair rotation, the embedded -64-byte span offsets, and the rate-class row grouping). Remove it when the -checkpoints it serves have been re-exported as `btx-atoms-v1` and the -re-exports validate. -""" - -from __future__ import annotations - -import torch - -from b12x.moe._shared.btx_schema import ( - ATOMS_PER_PAIR, - BTX_MANIFEST_KIND, - BTX_SCHEMA, - BtxManifest, - matrix_atom_bytes, - rate_code, -) -from b12x.moe._shared.kernels.w4a16.btx import BtxLayer -from b12x.moe._shared.trellis_codebooks import SQG_E4M3 - -# The frozen containers spread each expert's high-rate pairs across ranks -# with this multiplier; the lift materializes the result as tables. -_PAIR_ROTATION_MULTIPLIER = 5 -_SPAN_BYTES = 64 - - -def _lift_manifest( - *, - num_experts: int, - hidden_size: int, - global_intermediate_size: int, - layer_index: int, - rates: dict, - coupled: bool, -) -> BtxManifest: - hadamard: dict = { - "coupled": coupled, - "per_expert_input_rotations": False, - } - if coupled: - hadamard["pre_block"] = 512 - hadamard["post_block"] = 128 - atom_slots = global_intermediate_size // 32 - barriers = [atom_slots // 2] if coupled else [] - return BtxManifest.from_dict( - { - "kind": BTX_MANIFEST_KIND, - "schema": BTX_SCHEMA, - "codebook": SQG_E4M3, - "geometry": { - "num_experts": int(num_experts), - "hidden_size": int(hidden_size), - "intermediate_size": int(global_intermediate_size), - "atom_channels": 32, - "atom_slots": atom_slots, - "moe_layer_indices": [int(layer_index)], - }, - "rates": rates, - "hadamard": hadamard, - "layout": { - "atom_row_alignment": 1, - "extent_alignment_slots": 4, - "extent_barriers": barriers, - }, - "layers": { - str(int(layer_index)): { - "file": "lifted-in-memory", - "sha256": "0" * 64, - } - }, - } - ) - - -def _strip_spans( - bundles: torch.Tensor, *, trellis_bytes: int -) -> tuple[torch.Tensor, torch.Tensor]: - """Split embedded bundles into code words and per-atom rotation spans. - - ``bundles`` is ``[atoms, experts, trellis_bytes + 3*64]`` uint8. Returns - the pure code words ``[atoms, experts, trellis_bytes]`` and the rotation - values ``[atoms, experts, 3, 32]`` fp16 in physical channel order. - """ - - words = bundles[..., :trellis_bytes] - spans = ( - bundles[..., trellis_bytes:] - .contiguous() - .view(torch.float16) - .reshape(bundles.shape[0], bundles.shape[1], 3, 32) - ) - return words, spans - - -def _side_table(tensor: torch.Tensor) -> torch.Tensor: - return tensor.reshape(-1).contiguous() - - -def lift_qsrt_atoms_v1_extent( - atom_payload: torch.Tensor, - *, - first_atom_slot: int, - layer_index: int, - expert_ids: torch.Tensor, - format_codes: torch.Tensor, - hidden_size: int, - global_intermediate_size: int, - gate_suh: torch.Tensor, - up_suh: torch.Tensor, - down_svh: torch.Tensor, -) -> BtxLayer: - """Lift one fixed-payload atoms-v1 pair extent. - - ``atom_payload`` is ``[8, num_experts, bundle]`` uint8, one 256-channel - pair whose per-expert format byte packs the P24 pair counts - (``r13`` high nibble for FC1, ``r2`` low nibble for FC2). - """ - - if atom_payload.dim() != 3 or atom_payload.shape[0] != ATOMS_PER_PAIR: - raise ValueError( - "atoms-v1 extents are [8, num_experts, bundle] uint8 rows" - ) - num_experts = int(atom_payload.shape[1]) - matrix_bytes = matrix_atom_bytes(hidden_size, 3, 3) - if int(atom_payload.shape[2]) != 3 * matrix_bytes + 3 * _SPAN_BYTES: - raise ValueError( - "atoms-v1 bundle bytes disagree with the declared hidden size" - ) - expert_ids = expert_ids.to(dtype=torch.int64, device="cpu") - format_codes = format_codes.to(dtype=torch.int64, device="cpu") - if expert_ids.shape != (num_experts,) or format_codes.shape != ( - num_experts, - ): - raise ValueError( - "atoms-v1 expert_ids and format_codes must cover the extent's " - "experts" - ) - r13 = format_codes >> 4 - r2 = format_codes & 0xF - if bool(torch.any((r13 < 0) | (r13 > 2) | (r2 < 0) | (r2 > 2))): - raise ValueError("atoms-v1 format codes must encode R0/R1/R2") - physical_pair = first_atom_slot // ATOMS_PER_PAIR - rotation = (_PAIR_ROTATION_MULTIPLIER * expert_ids + layer_index) % 12 - logical_pair = (physical_pair - rotation) % 12 - p33 = rate_code(3, 3) - p24 = rate_code(2, 4) - fc1_codes = torch.where(logical_pair < r13, p24, p33).to(torch.uint8) - fc2_codes = torch.where(logical_pair < r2, p24, p33).to(torch.uint8) - - kinds = {"P33"} - if bool(torch.any(fc1_codes == p24)) or bool(torch.any(fc2_codes == p24)): - kinds.add("P24") - manifest = _lift_manifest( - num_experts=num_experts, - hidden_size=hidden_size, - global_intermediate_size=global_intermediate_size, - layer_index=layer_index, - rates={ - "structure": "per_expert_pair", - "pair_kinds": sorted(kinds), - }, - coupled=False, - ) - manifest.validate_extent(first_atom_slot, ATOMS_PER_PAIR) - - words, spans = _strip_spans(atom_payload, trellis_bytes=3 * matrix_bytes) - pairs = manifest.geometry.atom_slots // ATOMS_PER_PAIR - rates_fc1 = torch.zeros((pairs, num_experts), dtype=torch.uint8) - rates_fc2 = torch.zeros((pairs, num_experts), dtype=torch.uint8) - rates_fc1[physical_pair] = fc1_codes - rates_fc2[physical_pair] = fc2_codes - return BtxLayer( - manifest=manifest, - layer_index=layer_index, - first_slot=first_atom_slot, - slot_count=ATOMS_PER_PAIR, - atoms=words.reshape(ATOMS_PER_PAIR, -1), - rotations=spans, - gate_suh=_side_table(gate_suh), - up_suh=_side_table(up_suh), - down_svh=_side_table(down_svh), - rates_fc1=rates_fc1[physical_pair : physical_pair + 1], - rates_fc2=rates_fc2[physical_pair : physical_pair + 1], - rotation_draws=None, - ) - - -def lift_qsrt_atoms_v2_extent( - atom_payload: torch.Tensor, - *, - profile: str, - first_atom_slot: int, - layer_index: int, - hidden_size: int, - global_intermediate_size: int, - num_experts: int, - gate_suh: torch.Tensor, - up_suh: torch.Tensor, - down_svh: torch.Tensor, - rotation_draws: torch.Tensor | None = None, -) -> BtxLayer: - """Lift one atoms-v2 rank extent for a supported profile. - - ``atom_payload`` is ``[atom_count, row_bytes]`` uint8. The pure-K2 - coupled profile lifts to a uniform coupled extent of any legal length; - the fixed high-rate profile (``k3x22_k4x2``) lifts one pair with its - rate-class row grouping restored to expert-major order. The coupled - high-rate profile has no lift: its pair-kind mixes have no qualified - BTX execution path, so serving it requires a re-export decision. - """ - - if profile == "k2_coupled_h512_h128": - if rotation_draws is None: - raise ValueError("the coupled pure-K2 profile carries draws") - atom_count = int(atom_payload.shape[0]) - matrix_bytes = matrix_atom_bytes(hidden_size, 2, 2) - bundle = 3 * matrix_bytes + 3 * _SPAN_BYTES - payload = num_experts * bundle - if int(atom_payload.shape[1]) < payload: - raise ValueError("pure-K2 rows are shorter than their bundles") - if bool(torch.any(atom_payload[:, payload:] != 0)): - raise ValueError("pure-K2 row padding must be zero") - manifest = _lift_manifest( - num_experts=num_experts, - hidden_size=hidden_size, - global_intermediate_size=global_intermediate_size, - layer_index=layer_index, - rates={"structure": "uniform", "bits": 2}, - coupled=True, - ) - manifest.validate_extent(first_atom_slot, atom_count) - bundles = atom_payload[:, :payload].reshape( - atom_count, num_experts, bundle - ) - words, spans = _strip_spans(bundles, trellis_bytes=3 * matrix_bytes) - return BtxLayer( - manifest=manifest, - layer_index=layer_index, - first_slot=first_atom_slot, - slot_count=atom_count, - atoms=words.reshape(atom_count, -1), - rotations=spans, - gate_suh=_side_table(gate_suh), - up_suh=_side_table(up_suh), - down_svh=_side_table(down_svh), - rates_fc1=None, - rates_fc2=None, - rotation_draws=rotation_draws.to( - dtype=torch.uint8, device="cpu" - ).contiguous(), - ) - - if profile == "k3x22_k4x2": - if rotation_draws is not None: - raise ValueError("the fixed high-rate profile carries no draws") - if int(atom_payload.shape[0]) != ATOMS_PER_PAIR: - raise ValueError("fixed high-rate extents cover one pair") - physical_pair = first_atom_slot // ATOMS_PER_PAIR - expert_ids = torch.arange(num_experts, dtype=torch.int64) - rotation = ( - _PAIR_ROTATION_MULTIPLIER * expert_ids + layer_index - ) % 12 - base_pair = (physical_pair - rotation) % 12 - modes = (base_pair == 0) | (base_pair == 6) - p33_ids = torch.nonzero(~modes, as_tuple=False).flatten() - p43_ids = torch.nonzero(modes, as_tuple=False).flatten() - p33_bytes = 3 * matrix_atom_bytes(hidden_size, 3, 3) + 3 * _SPAN_BYTES - p43_bytes = 3 * matrix_atom_bytes(hidden_size, 4, 3) + 3 * _SPAN_BYTES - payload = ( - int(p33_ids.numel()) * p33_bytes + int(p43_ids.numel()) * p43_bytes - ) - if int(atom_payload.shape[1]) < payload: - raise ValueError( - "fixed high-rate rows are shorter than their compact groups" - ) - if bool(torch.any(atom_payload[:, payload:] != 0)): - raise ValueError("fixed high-rate row padding must be zero") - fc_codes = torch.where( - modes, rate_code(4, 3), rate_code(3, 3) - ).to(torch.uint8) - manifest = _lift_manifest( - num_experts=num_experts, - hidden_size=hidden_size, - global_intermediate_size=global_intermediate_size, - layer_index=layer_index, - rates={ - "structure": "per_expert_pair", - "pair_kinds": sorted({"P33", "P43"}), - }, - coupled=False, - ) - manifest.validate_extent(first_atom_slot, ATOMS_PER_PAIR) - - # Restore the rate-class row grouping to expert-major order. - span_rows = torch.empty( - (ATOMS_PER_PAIR, num_experts, 3, 32), dtype=torch.float16 - ) - expert_words: list[torch.Tensor | None] = [None] * num_experts - begin = 0 - for ids, bundle in ((p33_ids, p33_bytes), (p43_ids, p43_bytes)): - count = int(ids.numel()) - group = atom_payload[ - :, begin : begin + count * bundle - ].reshape(ATOMS_PER_PAIR, count, bundle) - words, spans = _strip_spans( - group, trellis_bytes=bundle - 3 * _SPAN_BYTES - ) - for position, expert in enumerate(ids.tolist()): - expert_words[expert] = words[:, position] - span_rows[:, expert] = spans[:, position] - begin += count * bundle - assert all(words is not None for words in expert_words) - atoms = torch.cat( - [words for words in expert_words if words is not None], dim=1 - ).reshape(ATOMS_PER_PAIR, -1) - return BtxLayer( - manifest=manifest, - layer_index=layer_index, - first_slot=first_atom_slot, - slot_count=ATOMS_PER_PAIR, - atoms=atoms, - rotations=span_rows, - gate_suh=_side_table(gate_suh), - up_suh=_side_table(up_suh), - down_svh=_side_table(down_svh), - rates_fc1=fc_codes.reshape(1, -1), - rates_fc2=fc_codes.reshape(1, -1), - rotation_draws=None, - ) - - raise ValueError( - f"QSRT atoms-v2 profile {profile!r} has no BTX lift; the coupled " - "high-rate profile's pair-kind mixes have no qualified fused " - "execution path" - ) diff --git a/b12x/moe/_shared/kernels/w4a16/btx_synth.py b/b12x/moe/_shared/kernels/w4a16/btx_synth.py deleted file mode 100644 index 1df619860..000000000 --- a/b12x/moe/_shared/kernels/w4a16/btx_synth.py +++ /dev/null @@ -1,317 +0,0 @@ -"""Synthetic BTX checkpoint writer for tests and benchmarks. - -Generates deterministic, schema-complete `btx-atoms-v1` checkpoints -(manifest plus per-layer safetensors) from random trellis words. This is a -fixture generator, not a converter: it exists so the reader, planner, and -serving paths can be exercised without a quantizer, and so one packer -implementation is shared by every test and benchmark. - -The per-(expert, slot) plane words are exposed as an intermediate -representation (`BtxLayerPayloads`) so equivalence tests can feed the same -logical weights to other packers and compare prepared tensors byte for -byte. -""" - -from __future__ import annotations - -import hashlib -import json -import pathlib -from dataclasses import dataclass, field - -import torch - -from b12x.moe._shared.btx_schema import ( - ATOM_CHANNELS, - ATOMS_PER_PAIR, - BTX_MANIFEST_FILENAME, - BTX_MANIFEST_KIND, - BTX_SCHEMA, - BtxManifest, - RATE_CODE_PAIR_KINDS, - RATE_STRUCTURE_PER_EXPERT_PAIR, - RATE_STRUCTURE_UNIFORM, - layer_filename, - matrix_atom_bytes, - rate_code, - rate_code_bits, -) -from b12x.moe._shared.trellis_codebooks import MCG, MCG_MULTIPLIER - - -@dataclass(frozen=True) -class BtxSynthConfig: - """Declarations for one synthetic checkpoint.""" - - codebook: str - num_experts: int - hidden_size: int - intermediate_size: int - moe_layer_indices: tuple[int, ...] - # Uniform structure declares bits; per-expert structure declares tables - # via ``rate_tables`` below. - bits: int | None = None - # Optional per-layer {layer: (rates_fc1, rates_fc2)} u8 tables of shape - # [atom_slots/8, num_experts]. Present iff bits is None. - rate_tables: dict[int, tuple[torch.Tensor, torch.Tensor]] | None = None - coupled: bool = False - pre_block: int | None = None - post_block: int | None = None - per_expert_input_rotations: bool = False - # Unit hidden-axis tables for routes that fold no input-side rotation. - unit_hidden_rotations: bool = False - atom_row_alignment: int = 4096 - extent_alignment_slots: int = 4 - extent_barriers: tuple[int, ...] = () - seed: int = 0 - - @property - def atom_slots(self) -> int: - return self.intermediate_size // ATOM_CHANNELS - - -@dataclass -class BtxLayerPayloads: - """One layer's logical content before row assembly. - - ``planes[(expert, slot, matrix)]`` is ``(low_plane, high_plane)``, - each an int16 tensor ``[hidden_size/16, 16*bits_plane]``. Matrices are - indexed 0=gate, 1=up, 2=down. - """ - - planes: dict[tuple[int, int, int], tuple[torch.Tensor, torch.Tensor]] - rotations: torch.Tensor - gate_suh: torch.Tensor - up_suh: torch.Tensor - down_svh: torch.Tensor - rotation_draws: torch.Tensor | None - rates_fc1: torch.Tensor | None - rates_fc2: torch.Tensor | None - metadata: dict[str, str] = field(default_factory=dict) - - -def _layer_rate_codes( - config: BtxSynthConfig, layer: int -) -> tuple[torch.Tensor, torch.Tensor]: - """Per-(pair, expert) rate-code tables, synthesized for uniform rates.""" - - pairs = config.atom_slots // ATOMS_PER_PAIR - if config.bits is not None: - code = rate_code(config.bits, config.bits) - table = torch.full( - (pairs, config.num_experts), code, dtype=torch.uint8 - ) - return table, table - assert config.rate_tables is not None - rates_fc1, rates_fc2 = config.rate_tables[layer] - expected = (pairs, config.num_experts) - for name, table in (("rates_fc1", rates_fc1), ("rates_fc2", rates_fc2)): - if tuple(table.shape) != expected or table.dtype != torch.uint8: - raise ValueError( - f"{name} must be uint8 {expected}, got " - f"{table.dtype} {tuple(table.shape)}" - ) - for code in table.unique().tolist(): - if int(code) not in RATE_CODE_PAIR_KINDS: - raise ValueError(f"{name} contains unknown rate code {code:#x}") - return rates_fc1, rates_fc2 - - -def synth_layer_payloads( - config: BtxSynthConfig, layer: int -) -> BtxLayerPayloads: - """Deterministic random content for one layer.""" - - generator = torch.Generator().manual_seed( - (config.seed << 20) ^ (layer * 2654435761 % (1 << 31)) - ) - hidden_tiles = config.hidden_size // 16 - experts = config.num_experts - slots = config.atom_slots - rates_fc1, rates_fc2 = _layer_rate_codes(config, layer) - - def _plane(bits: int) -> torch.Tensor: - return torch.randint( - -(1 << 15), - 1 << 15, - (hidden_tiles, 16 * bits), - dtype=torch.int16, - generator=generator, - ) - - planes: dict[tuple[int, int, int], tuple[torch.Tensor, torch.Tensor]] = {} - for expert in range(experts): - for slot in range(slots): - pair = slot // ATOMS_PER_PAIR - fc1_low, fc1_high = rate_code_bits(int(rates_fc1[pair, expert])) - fc2_low, fc2_high = rate_code_bits(int(rates_fc2[pair, expert])) - planes[(expert, slot, 0)] = (_plane(fc1_low), _plane(fc1_high)) - planes[(expert, slot, 1)] = (_plane(fc1_low), _plane(fc1_high)) - planes[(expert, slot, 2)] = (_plane(fc2_low), _plane(fc2_high)) - - def _values(shape: tuple[int, ...]) -> torch.Tensor: - raw = torch.rand(shape, generator=generator, dtype=torch.float32) - return (0.5 + raw).to(torch.float16) - - rotations = _values((slots, experts, 3, ATOM_CHANNELS)) - h_shape = ( - (experts, config.hidden_size) - if config.per_expert_input_rotations - else (config.hidden_size,) - ) - - def _h_values() -> torch.Tensor: - if config.unit_hidden_rotations: - return torch.ones(h_shape, dtype=torch.float16) - return _values(h_shape) - draws = None - if config.coupled: - draws = torch.randint( - 0, 8, (experts,), dtype=torch.uint8, generator=generator - ) - uniform = config.bits is not None - return BtxLayerPayloads( - planes=planes, - rotations=rotations, - gate_suh=_h_values(), - up_suh=_h_values(), - down_svh=_h_values(), - rotation_draws=draws, - rates_fc1=None if uniform else rates_fc1, - rates_fc2=None if uniform else rates_fc2, - ) - - -def assemble_atoms_rows( - config: BtxSynthConfig, layer: int, payloads: BtxLayerPayloads -) -> torch.Tensor: - """Pack plane words into the expert-major, zero-padded ``atoms`` tensor.""" - - rates_fc1, rates_fc2 = _layer_rate_codes(config, layer) - slots = config.atom_slots - row_bytes = [] - for slot in range(slots): - pair = slot // ATOMS_PER_PAIR - total = 0 - for expert in range(config.num_experts): - total += 2 * matrix_atom_bytes( - config.hidden_size, *rate_code_bits(int(rates_fc1[pair, expert])) - ) + matrix_atom_bytes( - config.hidden_size, *rate_code_bits(int(rates_fc2[pair, expert])) - ) - row_bytes.append(total) - alignment = config.atom_row_alignment - stride = (max(row_bytes) + alignment - 1) // alignment * alignment - atoms = torch.zeros((slots, stride), dtype=torch.uint8) - for slot in range(slots): - cursor = 0 - for expert in range(config.num_experts): - for matrix in range(3): - low, high = payloads.planes[(expert, slot, matrix)] - for plane in (low, high): - raw = plane.contiguous().view(torch.uint8).reshape(-1) - atoms[slot, cursor : cursor + raw.numel()] = raw - cursor += raw.numel() - return atoms - - -def _manifest_dict(config: BtxSynthConfig) -> dict: - rates: dict[str, object] = {} - if config.bits is not None: - rates = {"structure": RATE_STRUCTURE_UNIFORM, "bits": config.bits} - else: - kinds: set[str] = set() - assert config.rate_tables is not None - for rates_fc1, rates_fc2 in config.rate_tables.values(): - for table in (rates_fc1, rates_fc2): - kinds.update( - RATE_CODE_PAIR_KINDS[int(code)] - for code in table.unique().tolist() - ) - rates = { - "structure": RATE_STRUCTURE_PER_EXPERT_PAIR, - "pair_kinds": sorted(kinds), - } - hadamard: dict[str, object] = { - "coupled": config.coupled, - "per_expert_input_rotations": config.per_expert_input_rotations, - } - if config.coupled: - hadamard["pre_block"] = config.pre_block - hadamard["post_block"] = config.post_block - manifest: dict[str, object] = { - "kind": BTX_MANIFEST_KIND, - "schema": BTX_SCHEMA, - "codebook": config.codebook, - "geometry": { - "num_experts": config.num_experts, - "hidden_size": config.hidden_size, - "intermediate_size": config.intermediate_size, - "atom_channels": ATOM_CHANNELS, - "atom_slots": config.atom_slots, - "moe_layer_indices": list(config.moe_layer_indices), - }, - "rates": rates, - "hadamard": hadamard, - "layout": { - "atom_row_alignment": config.atom_row_alignment, - "extent_alignment_slots": config.extent_alignment_slots, - "extent_barriers": list(config.extent_barriers), - }, - "layers": {}, - } - if config.codebook == MCG: - manifest["codebook_seed"] = MCG_MULTIPLIER - return manifest - - -def layer_metadata(config: BtxSynthConfig, layer: int) -> dict[str, str]: - return { - "schema": BTX_SCHEMA, - "codebook": config.codebook, - "layer": str(int(layer)), - "num_experts": str(config.num_experts), - "hidden_size": str(config.hidden_size), - "intermediate_size": str(config.intermediate_size), - "atom_channels": str(ATOM_CHANNELS), - } - - -def write_btx_checkpoint( - root: str | pathlib.Path, config: BtxSynthConfig -) -> BtxManifest: - """Write a complete synthetic checkpoint and return its parsed manifest.""" - - from safetensors.torch import save_file - - root = pathlib.Path(root) - root.mkdir(parents=True, exist_ok=True) - manifest = _manifest_dict(config) - for layer in config.moe_layer_indices: - payloads = synth_layer_payloads(config, layer) - tensors: dict[str, torch.Tensor] = { - "atoms": assemble_atoms_rows(config, layer, payloads), - "rotations": payloads.rotations, - "gate_suh": payloads.gate_suh, - "up_suh": payloads.up_suh, - "down_svh": payloads.down_svh, - } - if payloads.rates_fc1 is not None: - assert payloads.rates_fc2 is not None - tensors["rates_fc1"] = payloads.rates_fc1 - tensors["rates_fc2"] = payloads.rates_fc2 - if payloads.rotation_draws is not None: - tensors["rotation_draws"] = payloads.rotation_draws - filename = layer_filename(layer) - save_file( - tensors, str(root / filename), metadata=layer_metadata(config, layer) - ) - digest = hashlib.sha256((root / filename).read_bytes()).hexdigest() - manifest["layers"][str(int(layer))] = { - "file": filename, - "sha256": digest, - } - (root / BTX_MANIFEST_FILENAME).write_text( - json.dumps(manifest, indent=2, sort_keys=True) + "\n" - ) - return BtxManifest.from_dict(manifest) diff --git a/b12x/moe/_shared/kernels/w4a16/host.py b/b12x/moe/_shared/kernels/w4a16/host.py index e2eb4d579..b095c1179 100644 --- a/b12x/moe/_shared/kernels/w4a16/host.py +++ b/b12x/moe/_shared/kernels/w4a16/host.py @@ -6,6 +6,7 @@ import torch +from b12x._lib.env import env_flag from b12x.moe._shared.kernels.activations import ( SUPPORTED_MOE_ACTIVATIONS, is_gated_moe_activation, @@ -18,6 +19,50 @@ _SUPPORTED_ACTIVATIONS = SUPPORTED_MOE_ACTIVATIONS +def prefill_fused_sum_enabled() -> bool: + """Enable direct FP32 route reduction for large-M W4A16 launches. + + The FC2 epilogue uses relaxed FP32 global reductions. Disable this path + when bitwise-identical route accumulation order is required. + """ + return env_flag("W4A16_PREFILL_FUSED_SUM") + + +def prefill_fused_sum_eligible( + *, + dtype: torch.dtype | str, + m: int, + full_rotation: bool, + weight_layout: str, + collect_activation_amax: bool, + enabled: bool | None = None, +) -> bool: + """Return whether one W4A16 launch may use direct route reduction. + + Args: + dtype: Activation element dtype as a PyTorch dtype or kernel dtype name. + m: Logical token rows represented by the launch. + full_rotation: Whether the launch uses full-rotation Trellis weights. + weight_layout: Prepared W4A16 weight layout. + collect_activation_amax: Whether the launch records activation maxima. + enabled: Frozen feature selection. ``None`` reads the process setting. + + Returns: + ``True`` when planning and execution may use the FP32 accumulator. + """ + if enabled is None: + enabled = prefill_fused_sum_enabled() + element_dtype = str(dtype).removeprefix("torch.") + return bool( + enabled + and element_dtype in {"bfloat16", "bf16"} + and int(m) > 8 + and not full_rotation + and weight_layout in {"packed", "modelopt"} + and not collect_activation_amax + ) + + @dataclass(frozen=True) class W4A16PackedShape: num_experts: int @@ -32,6 +77,7 @@ class W4A16PackedBuffers: intermediate_cache13: torch.Tensor intermediate_cache2: torch.Tensor output: torch.Tensor + prefill_sum_accum: torch.Tensor | None = None fc1_c_tmp: torch.Tensor | None = None fc2_c_tmp: torch.Tensor | None = None packed_route_indices: torch.Tensor | None = None @@ -55,6 +101,7 @@ class W4A16BufferPlan: intermediate_cache2_elements: int block_size_m: int rotation_a_elements: int = 0 + prefill_sum_accum_elements: int = 0 def validate_activation(activation: str) -> bool: @@ -190,6 +237,32 @@ def route_pack_token_capacity(tokens: int, topk: int) -> int: return 1 << (max(int(tokens), 1) - 1).bit_length() +def route_pack_warmup_token_counts(capacity: int) -> tuple[int, ...]: + """Return live row counts covering capacity and scalar specializations. + + Triton specializes the runtime ``live_numel`` scalar on alignment as well + as the constexpr route-capacity bucket. For a power-of-two bucket, the + bucket maximum and its first live member cover both alignment classes + reached by the GLM top-k route counts (for example 8 and 5 rows in the + 8-row bucket). Warming only bucket maxima leaves the less-aligned variant + to compile on the first odd-sized request. + """ + capacity = int(capacity) + if capacity < 1: + raise ValueError(f"route-pack warmup capacity must be positive, got {capacity}") + counts: list[int] = [] + bucket = 1 + while True: + first_in_bucket = 1 if bucket == 1 else bucket // 2 + 1 + if first_in_bucket > capacity: + break + for count in (first_in_bucket, min(bucket, capacity)): + if not counts or counts[-1] != count: + counts.append(count) + bucket *= 2 + return tuple(counts) + + def route_pack_capacity( numel: int, block_size: int, @@ -248,8 +321,12 @@ def plan_w4a16_buffers( topk: int, route_num_experts: int | None = None, sms: int, + dtype: torch.dtype | None = None, full_rotation: bool = False, block_size_m: int | None = None, + weight_layout: str | None = None, + collect_activation_amax: bool = False, + prefill_fused_sum: bool | None = None, ) -> W4A16BufferPlan: routed_rows = int(m) * int(topk) route_num_experts = ( @@ -280,6 +357,19 @@ def plan_w4a16_buffers( if int(m) <= 8 and bool(prepared.is_gated): gemm_route_slots = max(gemm_route_slots, routed_rows * block_size_m) scratch_sms = int(sms) + weight_layout = str( + weight_layout + if weight_layout is not None + else getattr(prepared, "weight_layout", "packed") + ) + use_prefill_fused_sum = prefill_fused_sum_eligible( + dtype=dtype if dtype is not None else "", + m=m, + full_rotation=full_rotation, + weight_layout=weight_layout, + collect_activation_amax=collect_activation_amax, + enabled=prefill_fused_sum, + ) return W4A16BufferPlan( routed_rows=routed_rows, fc1_cols=fc1_cols, @@ -297,10 +387,19 @@ def plan_w4a16_buffers( moe_block_size=block_size_m, sms=scratch_sms, ), - intermediate_cache13_elements=routed_rows * max(fc1_cols, hidden_size), + intermediate_cache13_elements=( + routed_rows * fc1_cols + if use_prefill_fused_sum + else routed_rows * max(fc1_cols, hidden_size) + ), intermediate_cache2_elements=routed_rows * intermediate_size, block_size_m=block_size_m, rotation_a_elements=(routed_rows * hidden_size if full_rotation else 0), + prefill_sum_accum_elements=( + int(m) * hidden_size + if use_prefill_fused_sum + else 0 + ), ) @@ -327,6 +426,7 @@ def make_w4a16_packed_buffers( topk=topk, route_num_experts=route_num_experts, sms=sms, + dtype=dtype, full_rotation=full_rotation, block_size_m=block_size_m, ) @@ -376,6 +476,15 @@ def make_w4a16_packed_buffers( dtype=torch.float32 if full_rotation else dtype, device=device, ), + prefill_sum_accum=( + torch.empty( + (plan.prefill_sum_accum_elements,), + dtype=torch.float32, + device=device, + ) + if plan.prefill_sum_accum_elements + else None + ), fc1_c_tmp=fc1_c_tmp, fc2_c_tmp=fc2_c_tmp, packed_route_indices=torch.empty( @@ -409,6 +518,7 @@ def make_w4a16_packed_buffers( "route_pack_numel_capacity", "route_pack_capacity", "route_pack_token_capacity", + "route_pack_warmup_token_counts", "select_route_block_size_m", "unswizzle_block_scale", "unswizzle_expert_scales", diff --git a/b12x/moe/_shared/kernels/w4a16/kernel.py b/b12x/moe/_shared/kernels/w4a16/kernel.py index 5b416153e..6a3ff164c 100644 --- a/b12x/moe/_shared/kernels/w4a16/kernel.py +++ b/b12x/moe/_shared/kernels/w4a16/kernel.py @@ -8,6 +8,7 @@ from typing import NamedTuple import cuda.bindings.driver as cuda +import cuda.bindings.runtime as cuda_runtime import cutlass import cutlass.cute as cute import torch @@ -67,6 +68,7 @@ pack_f32x2_to_bfloat2, pack_f32x2_to_f16x2, red_add_global_bf16x2, + red_add_global_v4_f32, red_add_global_release_i32, red_max_global_f32_nonnegative, shared_ptr_to_u32, @@ -85,16 +87,8 @@ warp_reduce, ) from b12x._lib.quant.sqg_e4m3 import sqg_xor_cheb_t12_lut -from b12x.moe._shared.kernels.trellis_ring import ( - trellis256_lane_geom_bits as _trellis_ring_lane_geom_bits, -) -from b12x.moe._shared.trellis_codebooks import ( - MCG, - SQG_E4M3, - SQG_FP16, - validate_codebook_bits, -) from b12x._lib.quant.sqg_fp16_d3l import ( + SQG_FP16_D3L, sqg_fp16_d3l_descriptors, ) from b12x._lib.utils import current_cuda_stream, make_ptr @@ -106,6 +100,7 @@ max_packed_route_slots, packed_gemm_scratch_elements, plan_w4a16_buffers, + prefill_fused_sum_eligible, select_route_block_size_m, validate_activation, ) @@ -171,13 +166,13 @@ def _sqg_xor_cheb_t12_smem_enabled() -> bool: _DEVICE_MAX_REG_BYTES = 255 * 1024 _DEFAULT_MAX_SHARED_MEM = 101_376 _SCALAR_ACC_FRAGMENT_WIDTH = 1 -_WEIGHT_LAYOUTS = {"packed", "modelopt", "trellis_t256"} +_WEIGHT_LAYOUTS = {"packed", "modelopt", "trellis3_t256"} _MODEL_OPT_W13_LAYOUTS = {"w13", "w31"} -_TRELLIS256_W13_LAYOUTS = {"packed", "trellis_t256_proj"} +_TRELLIS256_W13_LAYOUTS = {"packed", "trellis3_t256_proj"} # Native QSRT t256 tiles contain 256 tail-biting codes at one compile-time # bitrate. Their exact storage is [16*bits] int16 == [8*bits] uint32 per tile. _TRELLIS256_BITS = (2, 3, 4, 5, 6) -_TRELLIS256_CODEBOOKS = {MCG, SQG_E4M3, SQG_FP16} +_TRELLIS256_CODEBOOKS = {"mcg", "sqg_xor_cheb_t12", SQG_FP16_D3L} _SQG_XOR_CHEB_T12_LUT_ENTRIES = 1 << 12 _SQG_XOR_CHEB_T12_SMEM_REGION_BYTES = _SQG_XOR_CHEB_T12_LUT_ENTRIES _SCALE_FORMATS = { @@ -192,13 +187,21 @@ def _sqg_xor_cheb_t12_smem_enabled() -> bool: _FC2_DIRECT_MIN_EXPERT_CAPACITY = 1024 +def _validate_trellis256_codebook_bits(codebook: str, bits: int) -> None: + if codebook == "sqg_xor_cheb_t12" and bits not in (2, 3, 4): + raise ValueError("sqg_xor_cheb_t12 is defined only for K2/K3/K4") + if codebook == SQG_FP16_D3L and bits not in (5, 6): + raise ValueError("sqg_fp16_d3l is defined only for uniform K5/K6") + + def _trellis256_execution_lut( device: torch.device | str, codebook: str ) -> torch.Tensor: - if codebook == SQG_FP16: + if codebook == SQG_FP16_D3L: return sqg_fp16_d3l_descriptors(device) return sqg_xor_cheb_t12_lut(device) + # TC-decode: a small-M decode specialization that runs on the PACKED W4A16 # object (the same weights/scales the prefill GEMM uses). It reuses the packed # tensor-core MMA inner loop but folds the top-k sum into the FC2 store @@ -211,7 +214,6 @@ def _trellis256_execution_lut( _TC_DECODE_MAX_M = _W4A16_SMALL_M_DIRECT_MAX_M _TC_DECODE_M = tuple(range(1, _TC_DECODE_MAX_M + 1)) - @dsl_user_op def _materialize_w4a16_topk_route_f32(value, *, loc=None, ip=None): """Keep the BF16-to-F32 conversion outside the unrolled reduction add. @@ -580,7 +582,7 @@ class W4A16GemmCompileResult: w13_layout: str = "w13" dense_route_fast_path: bool = False trellis_bits: int = 3 - trellis_codebook: str = SQG_E4M3 + trellis_codebook: str = "sqg_xor_cheb_t12" trellis_pair_kind: str | None = None trellis_rate_axis: str | None = None @@ -639,17 +641,24 @@ class W4A16FusedMoeCompileResult: use_expert_map: bool = False scale_format: str = "e4m3_k16" tc_decode_fused_sum: bool = False + prefill_fused_sum_fp32: bool = False collect_activation_amax: bool = False schedule_whole_tiles: bool = False intermediate_rotation: bool = False dual_a: bool = False trellis_bits: int = 3 - trellis_codebook: str = SQG_E4M3 + trellis_codebook: str = "sqg_xor_cheb_t12" fc1_trellis_pair_kind: str | None = None fc2_trellis_pair_kind: str | None = None full_rotation: bool = False coupled_hadamard: bool = False rotation_input_dtype: str = "fp16" + # CUDA function attributes are available after a fresh compile. The + # on-disk object-cache loader does not currently expose the CUDA-dialect + # introspection surface, so cached entries retain the sentinel values. + kernel_symbol: str | None = None + registers_per_thread: int = -1 + local_memory_bytes: int = -1 cta_threads: int = -1 shared_memory_bytes: int = -1 @@ -764,6 +773,9 @@ def __init__( swiglu_beta=swiglu_beta, w13_layout=w13_layout, compile_time_phase=compile_time_phase, + # Fused direct launches have a compile-time route-table extent. + # FC2-only uses runtime M and sanitizes each route inside FC2. + stage_inactive_routes=int(compile_time_phase) != 2, ) @@ -787,7 +799,7 @@ def __init__( scale_format: str = "e4m3_k16", w13_layout: str = "w13", trellis_bits: int = 3, - trellis_codebook: str = SQG_E4M3, + trellis_codebook: str = "sqg_xor_cheb_t12", trellis_pair_kind: str | None = None, trellis_rate_axis: str | None = None, source_n_rotation: int = 0, @@ -797,10 +809,12 @@ def __init__( dual_a: bool = False, route_major_a: bool = False, fused_topk_sum: bool = False, + fused_sum_fp32: bool = False, fused_sum_topk: int = 1, schedule_whole_tiles: bool = False, dynamic_num_experts: bool = False, schedule_route_block_factor: int = 1, + paired_m8_routes: bool = False, ): if element_dtype not in {"bf16", "fp16"}: raise ValueError(f"unsupported element_dtype {element_dtype!r}") @@ -812,27 +826,27 @@ def __init__( if weight_layout == "modelopt": if w13_layout not in _MODEL_OPT_W13_LAYOUTS: raise ValueError(f"unsupported W4A16 w13_layout {w13_layout!r}") - elif weight_layout == "trellis_t256": + elif weight_layout == "trellis3_t256": if w13_layout not in _TRELLIS256_W13_LAYOUTS: - raise ValueError(f"unsupported trellis_t256 w13_layout {w13_layout!r}") + raise ValueError(f"unsupported trellis3_t256 w13_layout {w13_layout!r}") else: w13_layout = "packed" source_n_rotation = 0 - if weight_layout == "trellis_t256": + if weight_layout == "trellis3_t256": if trellis_codebook not in _TRELLIS256_CODEBOOKS: raise ValueError( - "trellis_t256 codebook must be one of " + "trellis3_t256 codebook must be one of " f"{sorted(_TRELLIS256_CODEBOOKS)}, got {trellis_codebook!r}" ) if trellis_bits not in _TRELLIS256_BITS: raise ValueError( - "trellis_t256 bits must be one of " + "trellis3_t256 bits must be one of " f"{_TRELLIS256_BITS}, got {trellis_bits}" ) - validate_codebook_bits(trellis_codebook, trellis_bits) + _validate_trellis256_codebook_bits(trellis_codebook, trellis_bits) if scale_format != "e4m3_k32": raise ValueError( - "trellis_t256 W4A16 weights require scale_format='e4m3_k32'" + "trellis3_t256 W4A16 weights require scale_format='e4m3_k32'" ) trellis_pair_kind = ( None if trellis_pair_kind is None else str(trellis_pair_kind).upper() @@ -845,8 +859,8 @@ def __init__( "trellis_pair_kind and trellis_rate_axis must be supplied together" ) if trellis_pair_kind is not None: - if weight_layout != "trellis_t256": - raise ValueError("trellis pairs require trellis_t256 weights") + if weight_layout != "trellis3_t256": + raise ValueError("trellis pairs require trellis3_t256 weights") if trellis_pair_kind not in { "P24", "P33", @@ -866,8 +880,7 @@ def __init__( ) if trellis_bits != 3: raise ValueError( - "QSRT pair decoding requires the trellis_bits=3 base " - "specialization" + "QSRT pair decoding requires the trellis_bits=3 base specialization" ) if epilogue_activation not in (None, "relu2"): raise ValueError( @@ -880,9 +893,7 @@ def __init__( "N-axis trellis pairs require size_n % 256 == 0 and " "tile_n=256 so every CTA consumes one complete fixed-rate pair" ) - if trellis_rate_axis == "k" and ( - size_k != 256 or tile_k > 128 or 128 % tile_k - ): + if trellis_rate_axis == "k" and (size_k != 256 or tile_k > 128 or 128 % tile_k): raise ValueError( "K-axis trellis pairs require size_k=256 and a tile_k that " "divides one 128-channel record" @@ -935,11 +946,11 @@ def __init__( self.dynamic_num_experts = bool(dynamic_num_experts) if self.dynamic_num_experts and weight_layout not in { "packed", - "trellis_t256", + "trellis3_t256", }: raise ValueError( "dynamic_num_experts is only supported for packed and " - "trellis_t256 weights" + "trellis3_t256 weights" ) self.top_k = int(top_k) self.mul_topk_weights = bool(mul_topk_weights) @@ -966,23 +977,23 @@ def __init__( "P43": (4, 3), "P44": (4, 4), } - self.trellis_pair_low_bits, self.trellis_pair_high_bits = ( - static_pair_rates.get(trellis_pair_kind, (3, 3)) + self.trellis_pair_low_bits, self.trellis_pair_high_bits = static_pair_rates.get( + trellis_pair_kind, (3, 3) ) self.sqg_xor_cheb_t12_smem = False # Small-M stripe split-K: opt out of the one-tile-per-CTA fast path # so decode-heavy small-M phases spread each mn-tile's K range across # multiple CTAs (existing tail scheduling plus cross-CTA finalize). self.small_m_splitk = _w4a16_small_m_splitk_enabled() - self.weight_layout_trellis256 = weight_layout == "trellis_t256" + self.weight_layout_trellis256 = weight_layout == "trellis3_t256" self.weight_layout_trellis256_proj = ( - self.weight_layout_trellis256 and w13_layout == "trellis_t256_proj" + self.weight_layout_trellis256 and w13_layout == "trellis3_t256_proj" ) if self.weight_layout_trellis256_proj and ( self.size_n % 2 != 0 or (self.size_n // 2) % self.tile_n != 0 ): raise ValueError( - "trellis_t256_proj requires each FC1 projection to contain " + "trellis3_t256_proj requires each FC1 projection to contain " "an integral number of CTA N tiles" ) self.b_region_variable = self.weight_layout_trellis256 @@ -1027,12 +1038,13 @@ def __init__( self.dual_a = bool(dual_a) if self.dual_a and not self.weight_layout_trellis256_proj: raise ValueError( - "dual_a is only valid for projection-major trellis_t256 FC1" + "dual_a is only valid for projection-major trellis3_t256 FC1" ) self.route_major_a = bool(route_major_a) if self.route_major_a and not self.dual_a: raise ValueError("route_major_a requires the exact dual-A FC1 path") self.fused_topk_sum = bool(fused_topk_sum) + self.fused_sum_fp32 = bool(fused_sum_fp32) self.fused_sum_topk = int(fused_sum_topk) # Whole-tile persistent scheduling: every mn-tile is computed by one # CTA over the full K (grid-strided waves, ragged last wave), skipping @@ -1057,14 +1069,24 @@ def __init__( and not self.weight_layout_trellis256 ): raise ValueError( - "schedule_whole_tiles requires direct_topk_routes or trellis_t256" + "schedule_whole_tiles requires direct_topk_routes or trellis3_t256" ) - if self.fused_topk_sum and not self.direct_topk_routes: - raise ValueError("fused_topk_sum requires direct_topk_routes") if self.fused_topk_sum and self.fused_sum_topk < 1: raise ValueError("fused_sum_topk must be >= 1") + if self.fused_sum_fp32 and not self.fused_topk_sum: + raise ValueError("fused_sum_fp32 requires fused_topk_sum") self.cta_m_blocks = int(_covering_count(moe_block_size, 16)) self.uses_m_block_8 = moe_block_size == 8 + self.paired_m8_routes = bool(paired_m8_routes) + if self.paired_m8_routes and ( + not self.uses_m_block_8 + or not self.schedule_whole_tiles + or self.schedule_route_block_factor != 2 + ): + raise ValueError( + "paired_m8_routes requires M8 whole-tile scheduling with " + "schedule_route_block_factor=2" + ) self.max_m_blocks = int(max_m_blocks) if torch.cuda.is_available(): props = torch.cuda.get_device_properties(torch.cuda.current_device()) @@ -1096,6 +1118,11 @@ def __init__( # W4A16 shared-memory geometry, in int4 units unless noted. self.a_sh_stride = 16 * self.cta_k_blocks // 8 self.a_sh_stage = self.a_sh_stride * (16 * self.cta_m_blocks) + if self.paired_m8_routes: + # The M8 ldmatrix mapping consumes a padded 16-row slab: rows 8-15 + # must remain zero. A paired tile therefore needs two independent + # 16-row slabs even though only eight rows in each slab are live. + self.a_sh_stage *= 2 self.a_gl_rd_delta_o = 16 * self.cta_k_blocks // 8 self.a_sh_wr_delta = self.a_sh_stride * ( self.cta_threads // self.a_gl_rd_delta_o @@ -1142,9 +1169,10 @@ def __init__( self.s_sh_stage = self.s_tb_groups * self.s_sh_stride self.tb_n_warps = self.cta_n_blocks // 4 - sh_block_route_indices = self.moe_block_size // 4 - sh_rd_block_route_indices = self.moe_block_size // 4 - sh_block_topk_weights = self.moe_block_size // 2 + route_metadata_rows = self.moe_block_size * (2 if self.paired_m8_routes else 1) + sh_block_route_indices = route_metadata_rows // 4 + sh_rd_block_route_indices = route_metadata_rows // 4 + sh_block_topk_weights = route_metadata_rows // 2 self.sh_valid_count_off = ( sh_block_route_indices + sh_rd_block_route_indices + sh_block_topk_weights ) @@ -1207,7 +1235,9 @@ def __cache_key__(self) -> tuple[object, ...]: self.dual_a, self.route_major_a, self.fused_topk_sum, + self.fused_sum_fp32, self.fused_sum_topk, + self.size_m if self.fused_sum_fp32 else None, self.cta_m_blocks, self.uses_m_block_8, self.shared_words, @@ -1217,6 +1247,7 @@ def __cache_key__(self) -> tuple[object, ...]: self.blocks_per_sm, self.schedule_whole_tiles, self.schedule_route_block_factor, + self.paired_m8_routes, self.sqg_xor_cheb_t12_smem, self.small_m_splitk, ) @@ -1232,6 +1263,13 @@ def _activation_smem_permuted_offset(self, i: Int32) -> Int32: def _int4_addr(self, smem_base: Int32, int4_off: Int32) -> Int32: return smem_base + int4_off * Int32(16) + @cute.jit + def _epilogue_sync(self, sync_barrier: cutlass.Constexpr = None): + if cutlass.const_expr(sync_barrier is None): + cute.arch.sync_threads() + else: + sync_barrier.arrive_and_wait() + @cute.jit def _dequant_e2m1x4_to_elem2x2(self, packed: Uint32): if cutlass.const_expr(self.is_fp16): @@ -1811,6 +1849,84 @@ def _read_moe_block_data( cute.arch.sync_threads() return valid_count + @cute.jit + def _read_moe_block_data_pair( + self, + packed_route_indices: cute.Tensor, + topk_weights_flat: cute.Tensor, + smem_base: Int32, + tid: Int32, + route_block_idx: Int32, + global_scale_f32: cutlass.Float32, + active_size_m: Int32, + ): + """Load two adjacent M8 route blocks into one 16-row metadata slab.""" + route_indices_int4_addr = self._int4_addr( + smem_base, Int32(self.sh_route_off) + tid + ) + route_indices_gmem = get_ptr_as_int64( + packed_route_indices, + route_block_idx * Int32(self.moe_block_size) + tid * Int32(4), + ) + cp_async4_shared_global_pred( + route_indices_int4_addr, + route_indices_gmem, + (tid < Int32(2 * self.moe_block_size // 4)).to(Int32), + ) + cute.arch.cp_async_commit_group() + cute.arch.cp_async_wait_group(0) + cute.arch.sync_threads() + + if tid >= Int32(self.cta_threads - 32): + lane = tid - Int32(self.cta_threads - 32) + valid0 = Int32(0) + valid1 = Int32(0) + if lane < Int32(2 * self.moe_block_size): + idx = ld_shared_i32_relaxed( + smem_base + Int32(self.sh_route_off * 16) + lane * Int32(4) + ) + valid = (idx < active_size_m * Int32(self.top_k)).to(Int32) + if lane < Int32(self.moe_block_size): + valid0 = valid + else: + valid1 = valid + valid0 = cute.arch.warp_redux_sync(valid0, "add") + valid1 = cute.arch.warp_redux_sync(valid1, "add") + if lane == Int32(0): + valid_addr = smem_base + Int32(self.sh_valid_count_off * 16) + st_shared_i32(valid_addr, valid0) + st_shared_i32(valid_addr + Int32(4), valid1) + + if tid < Int32(2 * self.moe_block_size): + idx = ld_shared_i32_relaxed( + smem_base + Int32(self.sh_route_off * 16) + tid * Int32(4) + ) + rd_row = idx // Int32(self.top_k) + if cutlass.const_expr(self.route_major_a): + rd_row = idx + st_shared_i32( + smem_base + Int32(self.sh_rd_route_off * 16) + tid * Int32(4), + rd_row, + ) + if cutlass.const_expr(self.mul_topk_weights): + safe_idx = idx + if idx >= active_size_m * Int32(self.top_k): + safe_idx = Int32(0) + topk = ( + topk_weights_flat[safe_idx].to(cutlass.Float32) * global_scale_f32 + ) + st_shared_u32( + smem_base + Int32(self.sh_topk_off * 16) + tid * Int32(4), + self._broadcast_f32_to_elem2(topk), + ) + + cute.arch.sync_threads() + valid_addr = smem_base + Int32(self.sh_valid_count_off * 16) + block_valid_rows0 = ld_shared_i32_relaxed(valid_addr) + block_valid_rows1 = ld_shared_i32_relaxed(valid_addr + Int32(4)) + cute.arch.sync_threads() + return block_valid_rows0, block_valid_rows1 + @cute.jit def _run_tile( self, @@ -2060,6 +2176,42 @@ def _tile_common_prologue( s_sh_rd, ) + @cute.jit + def _tile_common_prologue_pair( + self, + global_scale: cute.Tensor, + packed_route_indices: cute.Tensor, + topk_weights_flat: cute.Tensor, + smem_base: Int32, + tid: Int32, + route_block_idx: Int32, + expert_idx: Int32, + output_n_tile: Int32, + active_size_m: Int32, + ): + global_scale_f32 = global_scale[expert_idx].to(cutlass.Float32) + if cutlass.const_expr(self.scale_format_e8m0_k32): + if cutlass.const_expr(self.is_fp16): + global_scale_f32 *= cutlass.Float32(_E8M0_K32_FP16_GLOBAL_COMPENSATION) + else: + global_scale_f32 *= cutlass.Float32(_E8M0_K32_BF16_GLOBAL_COMPENSATION) + block_valid_rows0, block_valid_rows1 = self._read_moe_block_data_pair( + packed_route_indices, + topk_weights_flat, + smem_base, + tid, + route_block_idx, + global_scale_f32, + active_size_m, + ) + offsets = self._tile_stream_offsets(tid, expert_idx, output_n_tile) + return ( + global_scale_f32, + block_valid_rows0, + block_valid_rows1, + *offsets, + ) + @cute.jit def _tile_stream_offsets(self, tid: Int32, expert_idx: Int32, output_n_tile: Int32): a_gl_stride = Int32(self.size_k // 8) @@ -2200,6 +2352,8 @@ def _run_tile_m8( k_tiles, reduce_k_tile, block_valid_rows, + Int32(0), + False, a_gl_stride, b_gl_stride, s_gl_stride, @@ -2286,6 +2440,7 @@ def _run_tile_m8( tid, output_n_tile, block_valid_rows, + Int32(0), global_scale_f32, reduce_slice_count, reduce_slice_idx, @@ -2293,6 +2448,199 @@ def _run_tile_m8( True, ) + @cute.jit + def _run_tile_m8_pair( + self, + a_bf16_flat: cute.Tensor, + a_alt_bf16_flat: cute.Tensor, + b_i32_flat: cute.Tensor, + c_bf16_flat: cute.Tensor, + scales_i32_flat: cute.Tensor, + global_scale: cute.Tensor, + packed_route_indices: cute.Tensor, + topk_weights_flat: cute.Tensor, + c_tmp_f32_flat: cute.Tensor, + locks_i32_flat: cute.Tensor, + trellis_lut_addr: Int64, + smem_base: Int32, + tid: Int32, + route_block_idx: Int32, + expert_idx: Int32, + output_n_tile: Int32, + reduce_k_tile: Int32, + reduce_tile_count: Int32, + reduce_slice_count: Int32, + reduce_slice_idx: Int32, + lock_slot: Int32, + active_size_m: Int32, + ): + ( + global_scale_f32, + block_valid_rows0, + block_valid_rows1, + a_gl_stride, + b_gl_stride, + s_gl_stride, + scales_expert_off, + b_gl_rd_base, + a_gl_rd_row, + a_gl_rd_col0, + a_sh_wr, + a_rows_per_iter, + b_sh_rd, + s_sh_rd, + ) = self._tile_common_prologue_pair( + global_scale, + packed_route_indices, + topk_weights_flat, + smem_base, + tid, + route_block_idx, + expert_idx, + output_n_tile, + active_size_m, + ) + a0_sh_rd = self._a_shared_read_offset(tid, 8) + a1_sh_rd = a0_sh_rd + Int32(self.a_sh_rd_delta_i) + + acc0 = [ + cute.make_rmem_tensor((_SCALAR_ACC_FRAGMENT_WIDTH,), cutlass.Float32) + for _ in range(16 // _SCALAR_ACC_FRAGMENT_WIDTH) + ] + acc1 = [ + cute.make_rmem_tensor((_SCALAR_ACC_FRAGMENT_WIDTH,), cutlass.Float32) + for _ in range(16 // _SCALAR_ACC_FRAGMENT_WIDTH) + ] + for frag in cutlass.range_constexpr(16 // _SCALAR_ACC_FRAGMENT_WIDTH): + acc0[frag].fill(0.0) + acc1[frag].fill(0.0) + + k_tiles = reduce_tile_count + self._prefetch_initial_tiles( + a_bf16_flat, + a_alt_bf16_flat, + b_i32_flat, + scales_i32_flat, + smem_base, + tid, + k_tiles, + reduce_k_tile, + block_valid_rows0, + block_valid_rows1, + True, + a_gl_stride, + b_gl_stride, + s_gl_stride, + scales_expert_off, + b_gl_rd_base, + a_gl_rd_row, + a_gl_rd_col0, + a_sh_wr, + a_rows_per_iter, + output_n_tile, + expert_idx, + -1, + ) + + b_scale_cur = cute.make_rmem_tensor((2, 4), Uint32) + b_scale_next = cute.make_rmem_tensor((2, 4), Uint32) + self._load_b_scale_register_bundle( + b_scale_cur, + smem_base, + tid, + b_sh_rd, + s_sh_rd, + Int32(0), + Int32(0), + reduce_k_tile, + -1, + ) + a0_regs_cur = cute.make_rmem_tensor((2,), Uint32) + a0_regs_next = cute.make_rmem_tensor((2,), Uint32) + a1_regs_cur = cute.make_rmem_tensor((2,), Uint32) + a1_regs_next = cute.make_rmem_tensor((2,), Uint32) + self._load_a_registers_m8_bundle( + a0_regs_cur, smem_base, a0_sh_rd, Int32(0), Int32(0) + ) + self._load_a_registers_m8_bundle( + a1_regs_cur, smem_base, a1_sh_rd, Int32(0), Int32(0) + ) + self._run_mma_pipeline_m8_pair( + a_bf16_flat, + a_alt_bf16_flat, + b_i32_flat, + scales_i32_flat, + trellis_lut_addr, + smem_base, + tid, + acc0, + acc1, + b_scale_cur, + b_scale_next, + a0_regs_cur, + a0_regs_next, + a1_regs_cur, + a1_regs_next, + b_sh_rd, + s_sh_rd, + a0_sh_rd, + a1_sh_rd, + k_tiles, + reduce_k_tile, + block_valid_rows0, + block_valid_rows1, + a_gl_stride, + b_gl_stride, + s_gl_stride, + scales_expert_off, + b_gl_rd_base, + a_gl_rd_row, + a_gl_rd_col0, + a_sh_wr, + a_rows_per_iter, + output_n_tile, + expert_idx, + ) + + self._finish_tile( + acc0, + acc0, + acc0, + acc0, + c_bf16_flat, + c_tmp_f32_flat, + locks_i32_flat, + smem_base, + tid, + output_n_tile, + block_valid_rows0, + Int32(0), + global_scale_f32, + reduce_slice_count, + reduce_slice_idx, + lock_slot, + True, + ) + self._finish_tile( + acc1, + acc1, + acc1, + acc1, + c_bf16_flat, + c_tmp_f32_flat, + locks_i32_flat, + smem_base, + tid, + output_n_tile, + block_valid_rows1, + Int32(self.moe_block_size), + global_scale_f32, + reduce_slice_count, + reduce_slice_idx, + lock_slot + Int32(1), + True, + ) + @cute.jit def _run_tile_large_m( self, @@ -2391,6 +2739,8 @@ def _run_tile_large_m( k_tiles, reduce_k_tile, block_valid_rows, + Int32(0), + False, a_gl_stride, b_gl_stride, s_gl_stride, @@ -2477,6 +2827,7 @@ def _run_tile_large_m( tid, output_n_tile, block_valid_rows, + Int32(0), global_scale_f32, reduce_slice_count, reduce_slice_idx, @@ -2558,6 +2909,8 @@ def _run_mma_pipeline( k_tiles, reduce_k_tile, block_valid_rows, + Int32(0), + False, a_gl_stride, b_gl_stride, s_gl_stride, @@ -2573,9 +2926,7 @@ def _run_mma_pipeline( ) if cutlass.const_expr(self.trellis_pair_dynamic): - if cutlass.const_expr( - int(dynamic_pair_override) != 0 - ): + if cutlass.const_expr(int(dynamic_pair_override) != 0): self._dequant_and_accumulate_bundle( acc0, acc1, @@ -2665,6 +3016,152 @@ def _run_mma_pipeline( uses_m_block_8, ) + @cute.jit + def _run_mma_pipeline_m8_pair( + self, + a_bf16_flat: cute.Tensor, + a_alt_bf16_flat: cute.Tensor, + b_i32_flat: cute.Tensor, + scales_i32_flat: cute.Tensor, + trellis_lut_addr: Int64, + smem_base: Int32, + tid: Int32, + acc0, + acc1, + b_scale_cur: cute.Tensor, + b_scale_next: cute.Tensor, + a0_regs_cur: cute.Tensor, + a0_regs_next: cute.Tensor, + a1_regs_cur: cute.Tensor, + a1_regs_next: cute.Tensor, + b_sh_rd: Int32, + s_sh_rd: Int32, + a0_sh_rd: Int32, + a1_sh_rd: Int32, + k_tiles: Int32, + reduce_k_tile: Int32, + block_valid_rows0: Int32, + block_valid_rows1: Int32, + a_gl_stride: Int32, + b_gl_stride: Int32, + s_gl_stride: Int32, + scales_expert_off: Int32, + b_gl_rd_base: Int32, + a_gl_rd_row: Int32, + a_gl_rd_col0: Int32, + a_sh_wr: Int32, + a_rows_per_iter: Int32, + output_n_tile: Int32, + expert_idx: Int32, + ): + b_frag = cute.make_rmem_tensor((2, 2), Uint32) + tile_idx = Int32(0) + while tile_idx < k_tiles: + for pipe in cutlass.range_constexpr(_STAGES): + if tile_idx < k_tiles: + for kk in cutlass.range_constexpr(self.b_sh_wr_iters): + self._load_next_fragment_bundle_m8_pair( + b_scale_next, + a0_regs_next, + a1_regs_next, + smem_base, + tid, + b_sh_rd, + s_sh_rd, + a0_sh_rd, + a1_sh_rd, + pipe, + kk, + tile_idx, + k_tiles, + reduce_k_tile, + ) + + self._prefetch_pipeline_step( + a_bf16_flat, + a_alt_bf16_flat, + b_i32_flat, + scales_i32_flat, + smem_base, + tid, + pipe, + kk, + tile_idx, + k_tiles, + reduce_k_tile, + block_valid_rows0, + block_valid_rows1, + True, + a_gl_stride, + b_gl_stride, + s_gl_stride, + scales_expert_off, + b_gl_rd_base, + a_gl_rd_row, + a_gl_rd_col0, + a_sh_wr, + a_rows_per_iter, + output_n_tile, + expert_idx, + -1, + ) + + for jj in cutlass.range_constexpr(4): + if cutlass.const_expr(self.weight_layout_trellis256): + self._scaled_dequant_b_fragment_trellis256( + b_frag, + b_scale_cur[0, jj], + b_scale_cur[1, jj], + trellis_lut_addr, + ) + else: + q, s = self._select_b_scale_register(jj, b_scale_cur) + self._scaled_dequant_b_fragment(b_frag, q, s) + self._mma_accumulate_m8( + acc0, + jj, + a0_regs_cur, + b_frag, + ) + self._mma_accumulate_m8( + acc1, + jj, + a1_regs_cur, + b_frag, + ) + + self._copy_a_register_bundle_m8(a0_regs_cur, a0_regs_next) + self._copy_a_register_bundle_m8(a1_regs_cur, a1_regs_next) + self._copy_b_scale_register_bundle(b_scale_cur, b_scale_next) + tile_idx += Int32(1) + cute.arch.sync_threads() + if tile_idx < k_tiles: + self._load_b_scale_register_bundle( + b_scale_cur, + smem_base, + tid, + b_sh_rd, + s_sh_rd, + Int32(0), + Int32(0), + reduce_k_tile + tile_idx, + -1, + ) + self._load_a_registers_m8_bundle( + a0_regs_cur, + smem_base, + a0_sh_rd, + Int32(0), + Int32(0), + ) + self._load_a_registers_m8_bundle( + a1_regs_cur, + smem_base, + a1_sh_rd, + Int32(0), + Int32(0), + ) + @cute.jit def _dequant_and_accumulate_bundle( self, @@ -2688,10 +3185,7 @@ def _dequant_and_accumulate_bundle( and self.trellis_rate_axis == "k" and ( self.trellis_pair_kind in {"P24", "P43", "P44"} - or ( - self.trellis_pair_dynamic - and int(dynamic_pair_override) in (1, 2) - ) + or (self.trellis_pair_dynamic and int(dynamic_pair_override) in (1, 2)) ) ): # FC2 assigns an entire warp/kk fragment to one side of an @@ -2794,21 +3288,13 @@ def _dequant_and_accumulate_bundle( else: for mb in cutlass.range_constexpr(self.cta_m_blocks): if cutlass.const_expr(mb == 0): - self._mma_accumulate_large_m( - acc0, a_regs_cur, mb, jj, b_frag - ) + self._mma_accumulate_large_m(acc0, a_regs_cur, mb, jj, b_frag) elif cutlass.const_expr(mb == 1): - self._mma_accumulate_large_m( - acc1, a_regs_cur, mb, jj, b_frag - ) + self._mma_accumulate_large_m(acc1, a_regs_cur, mb, jj, b_frag) elif cutlass.const_expr(mb == 2): - self._mma_accumulate_large_m( - acc2, a_regs_cur, mb, jj, b_frag - ) + self._mma_accumulate_large_m(acc2, a_regs_cur, mb, jj, b_frag) else: - self._mma_accumulate_large_m( - acc3, a_regs_cur, mb, jj, b_frag - ) + self._mma_accumulate_large_m(acc3, a_regs_cur, mb, jj, b_frag) @cute.jit def _finish_tile( @@ -2824,11 +3310,13 @@ def _finish_tile( tid: Int32, output_n_tile: Int32, block_valid_rows: Int32, + metadata_row_base: Int32, global_scale_f32: cutlass.Float32, reduce_slice_count: Int32, reduce_slice_idx: Int32, lock_slot: Int32, uses_m_block_8: cutlass.Constexpr[bool], + sync_barrier: cutlass.Constexpr = None, ): if cutlass.const_expr(uses_m_block_8): self._fold_cta_partials_m8(acc0, smem_base, tid) @@ -2840,6 +3328,7 @@ def _finish_tile( acc3, smem_base, tid, + sync_barrier, ) if reduce_slice_count > Int32(1): @@ -2875,6 +3364,7 @@ def _finish_tile( tid, output_n_tile, block_valid_rows, + metadata_row_base, global_scale_f32, ) else: @@ -2889,6 +3379,7 @@ def _finish_tile( output_n_tile, block_valid_rows, global_scale_f32, + sync_barrier, ) @cute.jit @@ -3438,6 +3929,85 @@ def _load_next_fragment_bundle( uses_m_block_8, ) + @cute.jit + def _load_next_fragment_bundle_m8_pair( + self, + b_scale_next: cute.Tensor, + a0_regs_next: cute.Tensor, + a1_regs_next: cute.Tensor, + smem_base: Int32, + tid: Int32, + b_sh_rd: Int32, + s_sh_rd: Int32, + a0_sh_rd: Int32, + a1_sh_rd: Int32, + pipe: cutlass.Constexpr[int], + kk: cutlass.Constexpr[int], + tile_idx: Int32, + k_tiles: Int32, + reduce_k_tile: Int32, + ): + self._clear_b_scale_register_bundle(b_scale_next) + self._clear_a_register_bundle_m8(a0_regs_next) + self._clear_a_register_bundle_m8(a1_regs_next) + + if cutlass.const_expr(kk + 1 < self.b_sh_wr_iters): + if tile_idx < k_tiles: + self._load_b_scale_register_bundle( + b_scale_next, + smem_base, + tid, + b_sh_rd, + s_sh_rd, + Int32(pipe), + Int32(kk + 1), + reduce_k_tile + tile_idx, + -1, + ) + self._load_a_registers_m8_bundle( + a0_regs_next, + smem_base, + a0_sh_rd, + Int32(pipe), + Int32(kk + 1), + ) + self._load_a_registers_m8_bundle( + a1_regs_next, + smem_base, + a1_sh_rd, + Int32(pipe), + Int32(kk + 1), + ) + else: + next_tile = tile_idx + Int32(1) + if next_tile < k_tiles: + next_pipe = Int32((pipe + 1) % _STAGES) + self._load_b_scale_register_bundle( + b_scale_next, + smem_base, + tid, + b_sh_rd, + s_sh_rd, + next_pipe, + Int32(0), + reduce_k_tile + next_tile, + -1, + ) + self._load_a_registers_m8_bundle( + a0_regs_next, + smem_base, + a0_sh_rd, + next_pipe, + Int32(0), + ) + self._load_a_registers_m8_bundle( + a1_regs_next, + smem_base, + a1_sh_rd, + next_pipe, + Int32(0), + ) + @cute.jit def _scaled_dequant_b_fragment(self, frag: cute.Tensor, q: Uint32, s: Uint32): bq1 = q @@ -3455,7 +4025,6 @@ def _scaled_dequant_b_fragment(self, frag: cute.Tensor, q: Uint32, s: Uint32): frag[1, 0] = b1_0 frag[1, 1] = b1_1 - @cute.jit def _trellis256_lane_geom_bits( self, @@ -3464,9 +4033,18 @@ def _trellis256_lane_geom_bits( weight_count: cutlass.Constexpr[int], bits: cutlass.Constexpr[int], ): - return _trellis_ring_lane_geom_bits( - lane, weight_offset, weight_count, bits - ) + bits_i32 = Int32(int(bits)) + ring_u32 = Int32(8 * int(bits)) + t_offset = Int32(8) * lane + Int32(weight_offset) + b1 = (t_offset + Int32(257)) * bits_i32 + b0 = b1 - Int32(16) + b2 = b1 + Int32((int(weight_count) - 1) * int(bits)) + i0 = b0 >> Int32(5) + i2 = (b2 - Int32(1)) >> Int32(5) + ia = i0 - ring_u32 * (i0 >= ring_u32).to(Int32) + ib = i2 - ring_u32 * (i2 >= ring_u32).to(Int32) + s2 = (i2 + Int32(1)) * Int32(32) - b2 + return ia, ib, s2, i2 - i0 @cute.jit def _trellis256_lane_geom( @@ -3511,7 +4089,7 @@ def _scaled_dequant_b_fragment_trellis256_bits( o0, o1, o2, o3 = packed_dequant_trellis_to_bfloat2x4( win_a, win_b, int(bits) ) - elif cutlass.const_expr(self.trellis_codebook == SQG_FP16): + elif cutlass.const_expr(self.trellis_codebook == SQG_FP16_D3L): if cutlass.const_expr(self.is_fp16): o0, o1, o2, o3 = packed_decode_sqg_fp16_d3l_to_half2x4( win_a, win_b, trellis_lut_addr, int(bits) @@ -3685,10 +4263,7 @@ def _load_b_registers_trellis256_pair( regs = cute.make_rmem_tensor((2, 4), Uint32) if cutlass.const_expr( self.trellis_pair_kind == "P24" - or ( - self.trellis_pair_dynamic - and int(dynamic_pair_override) == 1 - ) + or (self.trellis_pair_dynamic and int(dynamic_pair_override) == 1) ): self._load_b_registers_trellis256_pair_bits( regs, smem_base, tid, pipe, kk, tile_idx, 2, 4 @@ -3750,9 +4325,8 @@ def _load_b_registers_trellis256_pair_bits( # original contiguous output-channel order. if cutlass.const_expr(jj < 2): record_n16 = Int32(2) * w_n + Int32(jj) - tile_base = ( - kt_local * Int32(pair_span_u32) - + record_n16 * Int32(8 * low_bits) + tile_base = kt_local * Int32(pair_span_u32) + record_n16 * Int32( + 8 * low_bits ) wa[jj], wb[jj] = self._load_trellis256_pair_tile_windows( b_region, tile_base, lane, low_bits @@ -3776,12 +4350,9 @@ def _load_b_registers_trellis256_pair_bits( for jj in cutlass.range_constexpr(4): local_n16 = Int32(4) * w_n + Int32(jj) tile_base = Int32(0) - if cutlass.const_expr(int(low_bits) == int(high_bits)): - tile_base = kt_base_u32 + local_n16 * Int32(8 * low_bits) - wa[jj], wb[jj] = self._load_trellis256_pair_tile_windows( - b_region, tile_base, lane, low_bits - ) - elif logical_k16 < Int32(8): + if cutlass.const_expr( + int(low_bits) == int(high_bits) + ) or logical_k16 < Int32(8): tile_base = kt_base_u32 + local_n16 * Int32(8 * low_bits) wa[jj], wb[jj] = self._load_trellis256_pair_tile_windows( b_region, tile_base, lane, low_bits @@ -4154,6 +4725,8 @@ def _stage_k_tile_async( pipe: Int32, tile_idx: Int32, block_valid_rows: Int32, + block_valid_rows1: Int32, + paired_m8: cutlass.Constexpr[bool], a_gl_stride: Int32, b_gl_stride: Int32, s_gl_stride: Int32, @@ -4169,12 +4742,23 @@ def _stage_k_tile_async( ): for i in cutlass.range_constexpr(self.a_sh_wr_iters): row = a_rows_per_iter * Int32(i) + a_gl_rd_row + metadata_row = row + route_rows = Int32(self.moe_block_size) + if cutlass.const_expr(paired_m8): + route_rows = Int32(2 * self.moe_block_size) + metadata_row = Int32(-1) + if row < Int32(self.moe_block_size): + metadata_row = row + elif row >= Int32(2 * self.moe_block_size) and row < Int32( + 3 * self.moe_block_size + ): + metadata_row = row - Int32(self.moe_block_size) route_index = Int32(0) - if row < Int32(self.moe_block_size): + if metadata_row >= Int32(0) and metadata_row < route_rows: route_index = ld_shared_i32_relaxed( smem_base + Int32(self.sh_rd_route_off * 16) - + row * Int32(4) + + metadata_row * Int32(4) ) a_int4 = ( Int64(route_index) * Int64(a_gl_stride) @@ -4198,9 +4782,19 @@ def _stage_k_tile_async( # stage, so this adds neither shared memory nor MMA work. if output_n_tile >= Int32(self.n_tiles // 2): a_src = get_ptr_as_int64(a_alt_bf16_flat, a_int4 * Int32(8)) + row_valid = row < block_valid_rows + if cutlass.const_expr(paired_m8): + if row < Int32(self.moe_block_size): + row_valid = row < block_valid_rows + elif row >= Int32(2 * self.moe_block_size) and row < Int32( + 3 * self.moe_block_size + ): + row_valid = row - Int32(2 * self.moe_block_size) < block_valid_rows1 + else: + row_valid = row < Int32(0) if cutlass.const_expr(self.has_k_tile_tail): a_k_int4 = tile_idx * Int32(self.a_gl_rd_delta_o) + a_gl_rd_col0 - if row < block_valid_rows and a_k_int4 < a_gl_stride: + if row_valid and a_k_int4 < a_gl_stride: cp_async4_shared_global( a_dst, a_src, @@ -4211,7 +4805,7 @@ def _stage_k_tile_async( cp_async4_shared_global_pred( a_dst, a_src, - (row < block_valid_rows).to(Int32), + row_valid.to(Int32), ) if cutlass.const_expr(self.weight_layout_trellis256): @@ -4221,19 +4815,13 @@ def _stage_k_tile_async( high_bits = 3 if cutlass.const_expr( self.trellis_pair_kind == "P24" - or ( - self.trellis_pair_dynamic - and int(dynamic_pair_override) == 1 - ) + or (self.trellis_pair_dynamic and int(dynamic_pair_override) == 1) ): low_bits = 2 high_bits = 4 elif cutlass.const_expr( self.trellis_pair_kind == "P43" - or ( - self.trellis_pair_dynamic - and int(dynamic_pair_override) == 2 - ) + or (self.trellis_pair_dynamic and int(dynamic_pair_override) == 2) ): low_bits = 4 high_bits = 3 @@ -4244,17 +4832,13 @@ def _stage_k_tile_async( # Preparation swizzles the reference record-major payload # into one fixed-size compact pair span per K16 row. pair_u32_per_k16 = Int32(8 * 8 * (low_bits + high_bits)) - t256_chunks_per_kt = self.cta_n_blocks * ( - low_bits + high_bits - ) + t256_chunks_per_kt = self.cta_n_blocks * (low_bits + high_bits) t256_total_chunks = self.cta_k_blocks * t256_chunks_per_kt t256_pair_u32 = (self.size_k // 16) * pair_u32_per_k16 for i in cutlass.range_constexpr(self.b_sh_wr_iters_var): t256_chunk = Int32(i * self.cta_threads) + tid t256_kt = t256_chunk // Int32(t256_chunks_per_kt) - t256_in_kt = ( - t256_chunk - t256_kt * Int32(t256_chunks_per_kt) - ) + t256_in_kt = t256_chunk - t256_kt * Int32(t256_chunks_per_kt) b_dst = ( smem_base + Int32(self.sh_b_off * 16) @@ -4269,23 +4853,23 @@ def _stage_k_tile_async( pair_plane_u32 = Int64(cute.size(b_i32_flat)) // Int64(2) if cutlass.const_expr(self.trellis_pair_compact_offsets): pair_descriptor = scales_i32_flat[expert_idx].to(Int64) - pair_base_i64 = ( - Int64(output_n_tile) * pair_plane_u32 - + (pair_descriptor >> Int64(1)) - ) + pair_base_i64 = Int64( + output_n_tile + ) * pair_plane_u32 + (pair_descriptor >> Int64(1)) else: - pair_base_i64 = ( - Int64(output_n_tile) * pair_plane_u32 - + Int64(expert_idx) * Int64(t256_pair_u32) + pair_base_i64 = Int64( + output_n_tile + ) * pair_plane_u32 + Int64(expert_idx) * Int64( + t256_pair_u32 ) else: # Dense/expert-major payload: [E, pair, K16, # compact-pair-row]. if cutlass.const_expr(self.trellis_pair_compact_offsets): pair_descriptor = scales_i32_flat[expert_idx].to(Int64) - pair_base_i64 = ( - pair_descriptor >> Int64(1) - ) + Int64(output_n_tile) * Int64(t256_pair_u32) + pair_base_i64 = (pair_descriptor >> Int64(1)) + Int64( + output_n_tile + ) * Int64(t256_pair_u32) else: pair_base_i64 = ( Int64(expert_idx) * Int64(self.size_n // 256) @@ -4312,21 +4896,14 @@ def _stage_k_tile_async( # prefix. max_chunks_per_kt = self.cta_n_blocks * 8 t256_expert_u32 = ( - (self.size_k // 16) - * t256_n16 - * 4 - * (low_bits + high_bits) + (self.size_k // 16) * t256_n16 * 4 * (low_bits + high_bits) ) low_record_u32 = Int32(8 * t256_n16 * 8 * low_bits) for i in cutlass.range_constexpr(self.b_sh_wr_iters_var): t256_chunk = Int32(i * self.cta_threads) + tid - t256_kt = t256_chunk // Int32(max_chunks_per_kt) - t256_in_kt = ( - t256_chunk - t256_kt * Int32(max_chunks_per_kt) - ) - logical_k16 = ( - tile_idx * Int32(self.cta_k_blocks) + t256_kt - ) + t256_kt = t256_chunk // Int32(max_chunks_per_kt) + t256_in_kt = t256_chunk - t256_kt * Int32(max_chunks_per_kt) + logical_k16 = tile_idx * Int32(self.cta_k_blocks) + t256_kt high_record = (logical_k16 >= Int32(8)).to(Int32) local_k16 = logical_k16 - high_record * Int32(8) record_bits = Int32(low_bits) @@ -4362,17 +4939,13 @@ def _stage_k_tile_async( ) else: t256_tile_u32 = 8 * self.trellis_bits - t256_expert_u32 = ( - (self.size_k // 16) * t256_n16 * t256_tile_u32 - ) + t256_expert_u32 = (self.size_k // 16) * t256_n16 * t256_tile_u32 t256_chunks_per_kt = self.cta_n_blocks * (2 * self.trellis_bits) t256_total_chunks = self.cta_k_blocks * t256_chunks_per_kt for i in cutlass.range_constexpr(self.b_sh_wr_iters_var): t256_chunk = Int32(i * self.cta_threads) + tid t256_kt = t256_chunk // Int32(t256_chunks_per_kt) - t256_in_kt = ( - t256_chunk - t256_kt * Int32(t256_chunks_per_kt) - ) + t256_in_kt = t256_chunk - t256_kt * Int32(t256_chunks_per_kt) b_dst = ( smem_base + Int32(self.sh_b_off * 16) @@ -4382,16 +4955,10 @@ def _stage_k_tile_async( if cutlass.const_expr(self.weight_layout_trellis256_proj): t256_half_n16 = t256_n16 // 2 t256_out_n16 = output_n_tile * Int32(self.cta_n_blocks) - t256_proj = ( - t256_out_n16 >= Int32(t256_half_n16) - ).to(Int32) - t256_local_n16 = ( - t256_out_n16 - t256_proj * Int32(t256_half_n16) - ) + t256_proj = (t256_out_n16 >= Int32(t256_half_n16)).to(Int32) + t256_local_n16 = t256_out_n16 - t256_proj * Int32(t256_half_n16) t256_proj_expert_u32 = ( - (self.size_k // 16) - * t256_half_n16 - * t256_tile_u32 + (self.size_k // 16) * t256_half_n16 * t256_tile_u32 ) # Projection-major W13 is physically [2, E, ...]. t256_plane_u32 = Int64(cute.size(b_i32_flat)) // Int64(2) @@ -4454,7 +5021,7 @@ def _stage_k_tile_async( Int32(i * self.cta_threads) + tid, ) - # trellis_t256 has no per-weight scale and its register-load arm returns + # trellis3_t256 has no per-weight scale and its register-load arm returns # before touching scale SMEM. Const-expr-elide the otherwise dead HBM # reads so every layer can share a four-byte aligned dummy scale tensor # instead of retaining 54 MiB of packed ones. @@ -4505,6 +5072,8 @@ def _prefetch_pipeline_step( k_tiles: Int32, reduce_k_tile: Int32, block_valid_rows: Int32, + block_valid_rows1: Int32, + paired_m8: cutlass.Constexpr[bool], a_gl_stride: Int32, b_gl_stride: Int32, s_gl_stride: Int32, @@ -4531,6 +5100,8 @@ def _prefetch_pipeline_step( k_tiles, reduce_k_tile, block_valid_rows, + block_valid_rows1, + paired_m8, a_gl_stride, b_gl_stride, s_gl_stride, @@ -4557,6 +5128,8 @@ def _prefetch_initial_tiles( k_tiles: Int32, reduce_k_tile: Int32, block_valid_rows: Int32, + block_valid_rows1: Int32, + paired_m8: cutlass.Constexpr[bool], a_gl_stride: Int32, b_gl_stride: Int32, s_gl_stride: Int32, @@ -4582,6 +5155,8 @@ def _prefetch_initial_tiles( Int32(pipe), reduce_k_tile + Int32(pipe), block_valid_rows, + block_valid_rows1, + paired_m8, a_gl_stride, b_gl_stride, s_gl_stride, @@ -4614,6 +5189,8 @@ def _prefetch_lookahead_tile( k_tiles: Int32, reduce_k_tile: Int32, block_valid_rows: Int32, + block_valid_rows1: Int32, + paired_m8: cutlass.Constexpr[bool], a_gl_stride: Int32, b_gl_stride: Int32, s_gl_stride: Int32, @@ -4639,6 +5216,8 @@ def _prefetch_lookahead_tile( Int32((pipe + _STAGES - 1) % _STAGES), reduce_k_tile + fetch_tile, block_valid_rows, + block_valid_rows1, + paired_m8, a_gl_stride, b_gl_stride, s_gl_stride, @@ -4866,13 +5445,16 @@ def _drain_output_smem( c_sh_rd: Int32, c_sh_rd_delta: Int32, block_valid_rows: Int32, + metadata_row_base: Int32, store_iters: cutlass.Constexpr[int], + sync_barrier: cutlass.Constexpr = None, ): for _ in cutlass.range_constexpr(store_iters): row = c_gl_wr // c_gl_stride if row < block_valid_rows: + metadata_row = metadata_row_base + row route_index = ld_shared_i32_relaxed( - smem_base + Int32(self.sh_route_off * 16) + row * Int32(4) + smem_base + Int32(self.sh_route_off * 16) + metadata_row * Int32(4) ) true_idx = Int64(route_index) * Int64(c_gl_stride) + Int64( c_gl_wr % c_gl_stride @@ -4884,7 +5466,7 @@ def _drain_output_smem( scale_bf2 = ld_shared_u32( smem_base + Int32(self.sh_topk_off * 16) - + row * Int32(4) + + metadata_row * Int32(4) ) q0 = self._elem2_mul(q0, scale_bf2) q1 = self._elem2_mul(q1, scale_bf2) @@ -4902,14 +5484,32 @@ def _drain_output_smem( # to token = route_index // top_k. bf16x2 add lands two # consecutive hidden lanes per word. token_idx = route_index // Int32(self.fused_sum_topk) - out_idx = Int64(token_idx) * Int64(c_gl_stride) + Int64( - c_gl_wr % c_gl_stride - ) - out_addr = get_ptr_as_int64(c_bf16_flat, out_idx * Int64(8)) - red_add_global_bf16x2(out_addr, q0) - red_add_global_bf16x2(out_addr + Int64(4), q1) - red_add_global_bf16x2(out_addr + Int64(8), q2) - red_add_global_bf16x2(out_addr + Int64(12), q3) + col_word = c_gl_wr % c_gl_stride + if cutlass.const_expr(self.fused_sum_fp32): + out_elem = ( + Int64(token_idx) * Int64(self.size_n) + + Int64(col_word) * Int64(8) + ) + out_addr = get_ptr_as_int64(c_bf16_flat, out_elem) + q00, q01 = self._elem2_to_f32x2(q0) + q10, q11 = self._elem2_to_f32x2(q1) + q20, q21 = self._elem2_to_f32x2(q2) + q30, q31 = self._elem2_to_f32x2(q3) + red_add_global_v4_f32(out_addr, q00, q01, q10, q11) + red_add_global_v4_f32( + out_addr + Int64(16), q20, q21, q30, q31 + ) + else: + out_idx = Int64(token_idx) * Int64(c_gl_stride) + Int64( + col_word + ) + out_addr = get_ptr_as_int64( + c_bf16_flat, out_idx * Int64(8) + ) + red_add_global_bf16x2(out_addr, q0) + red_add_global_bf16x2(out_addr + Int64(4), q1) + red_add_global_bf16x2(out_addr + Int64(8), q2) + red_add_global_bf16x2(out_addr + Int64(12), q3) else: st_global_v4_u32( get_ptr_as_int64(c_bf16_flat, true_idx * Int64(8)), @@ -4920,7 +5520,7 @@ def _drain_output_smem( ) c_gl_wr += c_gl_wr_delta c_sh_rd += c_sh_rd_delta - cute.arch.sync_threads() + self._epilogue_sync(sync_barrier) @cute.jit def _drain_output_smem_tail( @@ -4934,14 +5534,17 @@ def _drain_output_smem_tail( c_sh_rd: Int32, c_sh_rd_delta: Int32, block_valid_rows: Int32, + metadata_row_base: Int32, store_iters: cutlass.Constexpr[int], + sync_barrier: cutlass.Constexpr = None, ): for _ in cutlass.range_constexpr(store_iters): row = c_gl_wr // c_gl_stride_covered col_word = c_gl_wr - row * c_gl_stride_covered if row < block_valid_rows and col_word < c_gl_stride: + metadata_row = metadata_row_base + row route_index = ld_shared_i32_relaxed( - smem_base + Int32(self.sh_route_off * 16) + row * Int32(4) + smem_base + Int32(self.sh_route_off * 16) + metadata_row * Int32(4) ) true_idx = Int64(route_index) * Int64(c_gl_stride) + Int64(col_word) q0, q1, q2, q3 = ld_shared_v4_u32( @@ -4951,7 +5554,7 @@ def _drain_output_smem_tail( scale_bf2 = ld_shared_u32( smem_base + Int32(self.sh_topk_off * 16) - + row * Int32(4) + + metadata_row * Int32(4) ) q0 = self._elem2_mul(q0, scale_bf2) q1 = self._elem2_mul(q1, scale_bf2) @@ -4964,12 +5567,32 @@ def _drain_output_smem_tail( q3 = self._relu2_elem2(q3) if cutlass.const_expr(self.fused_topk_sum): token_idx = route_index // Int32(self.fused_sum_topk) - out_idx = Int64(token_idx) * Int64(c_gl_stride) + Int64(col_word) - out_addr = get_ptr_as_int64(c_bf16_flat, out_idx * Int64(8)) - red_add_global_bf16x2(out_addr, q0) - red_add_global_bf16x2(out_addr + Int64(4), q1) - red_add_global_bf16x2(out_addr + Int64(8), q2) - red_add_global_bf16x2(out_addr + Int64(12), q3) + if cutlass.const_expr(self.fused_sum_fp32): + out_elem = ( + Int64(token_idx) * Int64(self.size_n) + + Int64(col_word) * Int64(8) + ) + out_addr = get_ptr_as_int64(c_bf16_flat, out_elem) + q00, q01 = self._elem2_to_f32x2(q0) + q10, q11 = self._elem2_to_f32x2(q1) + q20, q21 = self._elem2_to_f32x2(q2) + q30, q31 = self._elem2_to_f32x2(q3) + red_add_global_v4_f32(out_addr, q00, q01, q10, q11) + red_add_global_v4_f32( + out_addr + Int64(16), q20, q21, q30, q31 + ) + else: + out_idx = ( + Int64(token_idx) * Int64(c_gl_stride) + + Int64(col_word) + ) + out_addr = get_ptr_as_int64( + c_bf16_flat, out_idx * Int64(8) + ) + red_add_global_bf16x2(out_addr, q0) + red_add_global_bf16x2(out_addr + Int64(4), q1) + red_add_global_bf16x2(out_addr + Int64(8), q2) + red_add_global_bf16x2(out_addr + Int64(12), q3) else: st_global_v4_u32( get_ptr_as_int64(c_bf16_flat, true_idx * Int64(8)), @@ -4980,7 +5603,7 @@ def _drain_output_smem_tail( ) c_gl_wr += c_gl_wr_delta c_sh_rd += c_sh_rd_delta - cute.arch.sync_threads() + self._epilogue_sync(sync_barrier) @cute.jit def _store_tile_m8( @@ -4991,6 +5614,7 @@ def _store_tile_m8( tid: Int32, output_n_tile: Int32, block_valid_rows: Int32, + metadata_row_base: Int32, global_scale_f32: cutlass.Float32, ): if cutlass.const_expr(self.has_n_tile_tail): @@ -5025,8 +5649,7 @@ def _store_tile_m8( for jj in cutlass.range_constexpr(4): wr = c_sh_wr + Int32(16 * jj) if cutlass.const_expr( - self.weight_layout_trellis256_pair - and self.trellis_rate_axis == "n" + self.weight_layout_trellis256_pair and self.trellis_rate_axis == "n" ): # MMA register assignment is [L0,L1,H0,H1] per warp for # balanced P24 work. Scatter the four N16 fragments back @@ -5037,11 +5660,7 @@ def _store_tile_m8( warp_n = tid // Int32(32) semantic_n16 = Int32(2) * warp_n + Int32(jj) if cutlass.const_expr(jj >= 2): - semantic_n16 = ( - Int32(8) - + Int32(2) * warp_n - + Int32(jj - 2) - ) + semantic_n16 = Int32(8) + Int32(2) * warp_n + Int32(jj - 2) compute_n16 = Int32(4) * warp_n + Int32(jj) wr += Int32(16) * (semantic_n16 - compute_n16) self._st_shared_elem_from_f32( @@ -5092,6 +5711,7 @@ def _store_tile_m8( c_sh_rd, c_sh_rd_delta, block_valid_rows, + metadata_row_base, store_iters, ) else: @@ -5104,6 +5724,7 @@ def _store_tile_m8( c_sh_rd, c_sh_rd_delta, block_valid_rows, + metadata_row_base, store_iters, ) @@ -5116,6 +5737,7 @@ def _fold_cta_partials_large_m( acc3, smem_base: Int32, tid: Int32, + sync_barrier: cutlass.Constexpr = None, ): red_off = self.cta_threads // self.b_sh_stride_threads // 2 if cutlass.const_expr(red_off >= 1): @@ -5133,6 +5755,7 @@ def _fold_cta_partials_large_m( red_sh_stride, red_sh_delta, red_sh_rd, + sync_barrier, ) elif cutlass.const_expr(mb == 1): self._fold_cta_partials_large_m_block( @@ -5143,6 +5766,7 @@ def _fold_cta_partials_large_m( red_sh_stride, red_sh_delta, red_sh_rd, + sync_barrier, ) elif cutlass.const_expr(mb == 2): self._fold_cta_partials_large_m_block( @@ -5153,6 +5777,7 @@ def _fold_cta_partials_large_m( red_sh_stride, red_sh_delta, red_sh_rd, + sync_barrier, ) else: self._fold_cta_partials_large_m_block( @@ -5163,6 +5788,7 @@ def _fold_cta_partials_large_m( red_sh_stride, red_sh_delta, red_sh_rd, + sync_barrier, ) @cute.jit @@ -5175,6 +5801,7 @@ def _fold_cta_partials_large_m_block( red_sh_stride: Int32, red_sh_delta: Int32, red_sh_rd: Int32, + sync_barrier: cutlass.Constexpr = None, ): if cutlass.const_expr(red_off == 2): if Int32(2) <= red_idx and red_idx < Int32(4): @@ -5197,7 +5824,7 @@ def _fold_cta_partials_large_m_block( (flat_j * 4 + 3) % _SCALAR_ACC_FRAGMENT_WIDTH ], ) - cute.arch.sync_threads() + self._epilogue_sync(sync_barrier) if Int32(1) <= red_idx and red_idx < Int32(2): for flat_j in cutlass.range_constexpr(8): @@ -5266,7 +5893,7 @@ def _fold_cta_partials_large_m_block( (flat_j * 4 + 3) % _SCALAR_ACC_FRAGMENT_WIDTH ], ) - cute.arch.sync_threads() + self._epilogue_sync(sync_barrier) if red_idx == Int32(0): for flat_j in cutlass.range_constexpr(8): @@ -5307,7 +5934,7 @@ def _fold_cta_partials_large_m_block( ] + r3 ) - cute.arch.sync_threads() + self._epilogue_sync(sync_barrier) @cute.jit def _write_bf16x2_shared( @@ -5337,6 +5964,7 @@ def _store_tile_large_m( output_n_tile: Int32, block_valid_rows: Int32, global_scale_f32: cutlass.Float32, + sync_barrier: cutlass.Constexpr = None, ): if cutlass.const_expr(self.has_n_tile_tail): ( @@ -5406,7 +6034,7 @@ def _store_tile_large_m( write_scale, ) c_sh_wr += Int32(16 * (4 * (2 * self.cta_n_blocks + 1))) - cute.arch.sync_threads() + self._epilogue_sync(sync_barrier) store_iters = _covering_count( 16 * self.cta_m_blocks, @@ -5423,7 +6051,9 @@ def _store_tile_large_m( c_sh_rd, c_sh_rd_delta, block_valid_rows, + Int32(0), store_iters, + sync_barrier, ) else: self._drain_output_smem( @@ -5435,7 +6065,9 @@ def _store_tile_large_m( c_sh_rd, c_sh_rd_delta, block_valid_rows, + Int32(0), store_iters, + sync_barrier, ) @cute.jit @@ -5451,18 +6083,13 @@ def _store_tile_large_m_block( for jj in cutlass.range_constexpr(4): wr = c_sh_wr + Int32(8 * jj) if cutlass.const_expr( - self.weight_layout_trellis256_pair - and self.trellis_rate_axis == "n" + self.weight_layout_trellis256_pair and self.trellis_rate_axis == "n" ): # Match the M<=8 epilogue: restore the reference record order # from the balanced LLHH per-warp MMA assignment before H128. semantic_n16 = Int32(2) * warp_n + Int32(jj) if cutlass.const_expr(jj >= 2): - semantic_n16 = ( - Int32(8) - + Int32(2) * warp_n - + Int32(jj - 2) - ) + semantic_n16 = Int32(8) + Int32(2) * warp_n + Int32(jj - 2) compute_n16 = Int32(4) * warp_n + Int32(jj) wr += Int32(8) * (semantic_n16 - compute_n16) self._write_bf16x2_shared( @@ -5540,12 +6167,13 @@ def __init__( scale_format: str = "e4m3_k16", w13_layout: str = "w13", trellis_bits: int = 3, - trellis_codebook: str = SQG_E4M3, + trellis_codebook: str = "sqg_xor_cheb_t12", fc1_trellis_pair_kind: str | None = None, fc2_trellis_pair_kind: str | None = None, direct_topk_routes: bool = False, use_expert_map: bool = False, tc_decode_fused_sum: bool = False, + prefill_fused_sum_fp32: bool = False, tc_zero_output: bool = True, collect_activation_amax: bool = False, schedule_whole_tiles: bool = False, @@ -5569,26 +6197,39 @@ def __init__( if weight_layout == "modelopt": if w13_layout not in _MODEL_OPT_W13_LAYOUTS: raise ValueError(f"unsupported W4A16 w13_layout {w13_layout!r}") - elif weight_layout == "trellis_t256": + elif weight_layout == "trellis3_t256": if w13_layout not in _TRELLIS256_W13_LAYOUTS: - raise ValueError(f"unsupported trellis_t256 w13_layout {w13_layout!r}") + raise ValueError(f"unsupported trellis3_t256 w13_layout {w13_layout!r}") else: w13_layout = "packed" self.tc_decode_fused_sum = bool(tc_decode_fused_sum) + self.prefill_fused_sum_fp32 = bool(prefill_fused_sum_fp32) + if self.tc_decode_fused_sum and self.prefill_fused_sum_fp32: + raise ValueError( + "TC-decode and large-M FP32 route reduction are mutually exclusive" + ) # When two TC-decode launches share one pre-zeroed output, only the # first must zero it. Default True preserves single-launch behavior. self.tc_zero_output = bool(tc_zero_output) self.collect_activation_amax = bool(collect_activation_amax) if self.collect_activation_amax and bool(direct_topk_routes): raise ValueError("activation amax collection requires route-packed W4A16") - if self.collect_activation_amax and self.tc_decode_fused_sum: + if self.collect_activation_amax and ( + self.tc_decode_fused_sum or self.prefill_fused_sum_fp32 + ): raise ValueError( - "activation amax collection is incompatible with TC-decode" + "activation amax collection is incompatible with fused route reduction" ) if self.tc_decode_fused_sum and not bool(direct_topk_routes): raise ValueError("tc_decode_fused_sum requires direct_topk_routes") if self.tc_decode_fused_sum and element_dtype != "bf16": raise ValueError("tc_decode_fused_sum currently requires bf16 activations") + if self.prefill_fused_sum_fp32 and element_dtype != "bf16": + raise ValueError("prefill_fused_sum_fp32 requires bf16 activations") + if self.prefill_fused_sum_fp32 and int(size_m) <= _TC_DECODE_MAX_M: + raise ValueError( + "prefill_fused_sum_fp32 requires a token capacity above the decode range" + ) fc1_cols = int(intermediate_size) * (2 if is_gated else 1) routed_rows = int(size_m) * int(top_k) self.size_m = int(size_m) @@ -5603,13 +6244,14 @@ def __init__( # the launch boundary. self.dynamic_num_experts = weight_layout in { "packed", - "trellis_t256", + "trellis3_t256", } self.top_k = int(top_k) self.moe_block_size = int(moe_block_size) - # Stripe split-K spreads each mn-tile's K range across many CTAs for - # decode-heavy small-M phases. It is incompatible with whole-tile - # scheduling and grouped FC2 route subtiles. + # Classic stripe split-K experiment: revert the whole-tile wave + # schedule (and the paired-route grouping that requires it) so + # decode-heavy small-M phases spread each mn-tile's K range across + # many CTAs instead of idling most of the grid. self.small_m_splitk = _w4a16_small_m_splitk_enabled() if self.small_m_splitk: schedule_whole_tiles = False @@ -5648,10 +6290,10 @@ def __init__( self.weight_layout = weight_layout self.trellis_bits = int(trellis_bits) self.trellis_codebook = str(trellis_codebook).lower() - if self.weight_layout == "trellis_t256": + if self.weight_layout == "trellis3_t256": if self.trellis_codebook not in _TRELLIS256_CODEBOOKS: raise ValueError( - "trellis_t256 codebook must be one of " + "trellis3_t256 codebook must be one of " f"{sorted(_TRELLIS256_CODEBOOKS)}, got {self.trellis_codebook!r}" ) self.fc1_trellis_pair_kind = ( @@ -5664,19 +6306,16 @@ def __init__( if fc2_trellis_pair_kind is None else str(fc2_trellis_pair_kind).upper() ) - if (self.fc1_trellis_pair_kind is None) != ( - self.fc2_trellis_pair_kind is None - ): + if (self.fc1_trellis_pair_kind is None) != (self.fc2_trellis_pair_kind is None): raise ValueError( "fused trellis pair weights require both FC1 and FC2 pair kinds" ) if self.fc1_trellis_pair_kind is not None: - if weight_layout != "trellis_t256": - raise ValueError("fused trellis pairs require trellis_t256 weights") + if weight_layout != "trellis3_t256": + raise ValueError("fused trellis pairs require trellis3_t256 weights") if self.trellis_bits != 3: raise ValueError( - "fused QSRT pairs require the trellis_bits=3 base " - "specialization" + "fused QSRT pairs require the trellis_bits=3 base specialization" ) dynamic_kinds = {"PDYNAMIC", "P33_P43"} static_kinds = {"P24", "P33", "P43", "P44"} @@ -5685,9 +6324,7 @@ def __init__( or self.fc2_trellis_pair_kind in dynamic_kinds ): if self.fc1_trellis_pair_kind != self.fc2_trellis_pair_kind: - raise ValueError( - "dynamic fused trellis pair kinds must match" - ) + raise ValueError("dynamic fused trellis pair kinds must match") elif ( self.fc1_trellis_pair_kind not in static_kinds or self.fc2_trellis_pair_kind not in static_kinds @@ -5705,14 +6342,14 @@ def __init__( if self.use_expert_map and not self.direct_topk_routes: raise ValueError("use_expert_map requires direct_topk_routes") self.schedule_whole_tiles = bool( - (schedule_whole_tiles or weight_layout == "trellis_t256") + (schedule_whole_tiles or weight_layout == "trellis3_t256") and not self.small_m_splitk ) self.intermediate_rotation = bool(intermediate_rotation) if self.intermediate_rotation: - if weight_layout != "trellis_t256": + if weight_layout != "trellis3_t256": raise ValueError( - "intermediate_rotation is only supported for trellis_t256" + "intermediate_rotation is only supported for trellis3_t256" ) if not is_gated or self.activation_is_swigluoai or self.has_swiglu_limit: raise ValueError( @@ -5737,8 +6374,8 @@ def __init__( raise ValueError("full_rotation requires fp16 GEMM operands") if self.rotation_input_dtype not in {"bf16", "fp16"}: raise ValueError("full_rotation input dtype must be 'bf16' or 'fp16'") - if self.tc_decode_fused_sum: - raise ValueError("full_rotation is incompatible with TC decode") + if self.tc_decode_fused_sum or self.prefill_fused_sum_fp32: + raise ValueError("full_rotation is incompatible with fused route reduction") if self.apply_router_weight_on_input: raise ValueError( "full_rotation applies router weights only in the fp32 top-k sum" @@ -5756,8 +6393,8 @@ def __init__( ) self.dual_a = bool( self.intermediate_rotation - and weight_layout == "trellis_t256" - and w13_layout == "trellis_t256_proj" + and weight_layout == "trellis3_t256" + and w13_layout == "trellis3_t256_proj" ) fc1_source_n_rotation = ( int(intermediate_size) @@ -5783,9 +6420,7 @@ def __init__( trellis_bits=self.trellis_bits, trellis_codebook=self.trellis_codebook, trellis_pair_kind=self.fc1_trellis_pair_kind, - trellis_rate_axis=( - "n" if self.fc1_trellis_pair_kind is not None else None - ), + trellis_rate_axis=("n" if self.fc1_trellis_pair_kind is not None else None), source_n_rotation=fc1_source_n_rotation, single_token_route_fast_path=size_m == 1 and not self.direct_topk_routes, direct_topk_routes=self.direct_topk_routes, @@ -5812,20 +6447,25 @@ def __init__( element_dtype=element_dtype, weight_layout=weight_layout, scale_format=scale_format, - w13_layout=("packed" if weight_layout == "trellis_t256" else w13_layout), + w13_layout=("packed" if weight_layout == "trellis3_t256" else w13_layout), trellis_bits=self.trellis_bits, trellis_codebook=self.trellis_codebook, trellis_pair_kind=self.fc2_trellis_pair_kind, - trellis_rate_axis=( - "k" if self.fc2_trellis_pair_kind is not None else None - ), + trellis_rate_axis=("k" if self.fc2_trellis_pair_kind is not None else None), single_token_route_fast_path=size_m == 1 and not self.direct_topk_routes, direct_topk_routes=self.direct_topk_routes, - fused_topk_sum=self.tc_decode_fused_sum, + fused_topk_sum=( + self.tc_decode_fused_sum or self.prefill_fused_sum_fp32 + ), + fused_sum_fp32=self.prefill_fused_sum_fp32, fused_sum_topk=int(top_k), schedule_whole_tiles=self.schedule_whole_tiles, dynamic_num_experts=self.dynamic_num_experts, schedule_route_block_factor=self.fc2_schedule_route_block_factor, + paired_m8_routes=( + self.fc2_moe_block_size == 8 + and self.fc2_schedule_route_block_factor == 2 + ), ) self.cta_threads = max(self.fc1.cta_threads, self.fc2.cta_threads) if self.fc1.cta_threads != self.fc2.cta_threads: @@ -5836,17 +6476,14 @@ def __init__( self.blocks_per_sm = min(self.fc1.blocks_per_sm, self.fc2.blocks_per_sm) self.shared_words = max(self.fc1.shared_words, self.fc2.shared_words) self.sqg_xor_cheb_t12_smem = ( - self.trellis_codebook == SQG_E4M3 + self.trellis_codebook == "sqg_xor_cheb_t12" and _sqg_xor_cheb_t12_smem_enabled() ) self.sqg_xor_cheb_t12_smem_off = 0 if self.sqg_xor_cheb_t12_smem: - self.sqg_xor_cheb_t12_smem_off = ( - self.shared_words * 4 + 15 - ) // 16 * 16 + self.sqg_xor_cheb_t12_smem_off = (self.shared_words * 4 + 15) // 16 * 16 self.shared_words = ( - self.sqg_xor_cheb_t12_smem_off - + _SQG_XOR_CHEB_T12_SMEM_REGION_BYTES + self.sqg_xor_cheb_t12_smem_off + _SQG_XOR_CHEB_T12_SMEM_REGION_BYTES ) // 4 self.fc1.sqg_xor_cheb_t12_smem = True self.fc2.sqg_xor_cheb_t12_smem = True @@ -6270,9 +6907,7 @@ class Storage: tid, ) cute.arch.sync_threads() - table_addr = Int64( - smem_base + Int32(self.sqg_xor_cheb_t12_smem_off) - ) + table_addr = Int64(smem_base + Int32(self.sqg_xor_cheb_t12_smem_off)) fc1_phase_lut_addr = table_addr fc2_phase_lut_addr = table_addr @@ -6410,9 +7045,7 @@ def _moe_body( fc1_phase_lut = fc1_trellis_lut_addr fc2_phase_lut = fc2_trellis_lut_addr if cutlass.const_expr(self.sqg_xor_cheb_t12_smem): - table_addr = Int64( - smem_base + Int32(self.sqg_xor_cheb_t12_smem_off) - ) + table_addr = Int64(smem_base + Int32(self.sqg_xor_cheb_t12_smem_off)) fc1_phase_lut = table_addr fc2_phase_lut = table_addr if cutlass.const_expr(self.full_rotation): @@ -6451,34 +7084,34 @@ def _moe_body( active_m, ) self._grid_barrier(locks_i32_flat, tid, grid_x) - if cutlass.const_expr(self.tc_decode_fused_sum): - # The TC-decode FC2 epilogue atomically accumulates per-route - # partials directly into the per-token output, so the output must be - # pre-zeroed. Previously this was a SEPARATE host-side output.zero_() - # kernel launch on the latency-bound decode critical path (an extra - # launch + its grid-fill memset before the fused kernel even starts). - # Fold it into the fused kernel prologue here: every CTA zeroes a - # grid-strided slice of the output BEFORE FC1, and the EXISTING - # post-FC1 grid barrier (already required to order FC1 writes before - # the activation/FC2 read) makes all zero stores globally visible - # before the first FC2 atomic -- so no extra barrier is added. The - # tiny m*hidden bf16 memset (decode: <=4*4096 elems) is dwarfed by - # FC1's whole-K FP4-weight stream, but we delete one whole kernel - # launch from the per-decode chain. The TC-decode output is per-token - # (top_k routes atomically summed into the SAME token row), so the - # zero span is active_m*hidden_size -- NOT the per-route - # active_m*top_k*hidden_size of _zero_fc2_output. - # tc_zero_output=False skips the zero (a paired earlier launch has - # already zeroed the shared output); the grid barrier below is - # unconditional so ordering is preserved either way. + if cutlass.const_expr( + self.tc_decode_fused_sum or self.prefill_fused_sum_fp32 + ): + # FC2 route reduction atomically accumulates into one row per token. + # Every CTA zeroes a grid-strided slice before FC1. The mandatory + # post-FC1 grid barrier orders these stores before every FC2 atomic. + # ``tc_zero_output=False`` is valid only when a paired launch has + # already zeroed the same output and participates in that barrier. if cutlass.const_expr(self.tc_zero_output): zidx = cta * Int32(self.cta_threads) + tid zstride = grid_x * Int32(self.cta_threads) - ztotal = active_m * Int32(self.hidden_size) - zzero = self._cast_elem(cutlass.Float32(0.0)) - while zidx < ztotal: - fc2_bf16_flat[zidx] = zzero - zidx += zstride + zzero = ( + cutlass.Float32(0.0) + if cutlass.const_expr(self.prefill_fused_sum_fp32) + else self._cast_elem(cutlass.Float32(0.0)) + ) + if cutlass.const_expr(self.prefill_fused_sum_fp32): + zidx_i64 = Int64(zidx) + zstride_i64 = Int64(zstride) + ztotal_i64 = Int64(active_m) * Int64(self.hidden_size) + while zidx_i64 < ztotal_i64: + fc2_bf16_flat[zidx_i64] = zzero + zidx_i64 += zstride_i64 + else: + ztotal = active_m * Int32(self.hidden_size) + while zidx < ztotal: + fc2_bf16_flat[zidx] = zzero + zidx += zstride if cutlass.const_expr(self.activation_is_gated): self.fc1._run_persistent_gemm( @@ -6608,7 +7241,6 @@ def _moe_body( active_m * Int32(self.top_k), fc2_emit_tile, ) - @cute.jit def _sqg_smem_copy( self, @@ -6790,9 +7422,7 @@ def _run_input_rotation_coupled( route_pos = unit // nblk blk = unit - route_pos * nblk route = packed_route_indices[route_pos].to(Int32) - expert = block_expert_ids[route_pos // Int32(self.moe_block_size)].to( - Int32 - ) + expert = block_expert_ids[route_pos // Int32(self.moe_block_size)].to(Int32) if cutlass.const_expr(self.direct_topk_routes): route = route_pos expert = packed_route_indices[route_pos].to(Int32) @@ -6860,18 +7490,10 @@ def _run_input_rotation_coupled( x_input_flat[x_base + Int32(387)].to(cutlass.Float32) ).to(cutlass.Float32) - h00, h01, h02, h03 = self._had128_quad( - x00, x01, x02, x03, lane - ) - h10, h11, h12, h13 = self._had128_quad( - x10, x11, x12, x13, lane - ) - h20, h21, h22, h23 = self._had128_quad( - x20, x21, x22, x23, lane - ) - h30, h31, h32, h33 = self._had128_quad( - x30, x31, x32, x33, lane - ) + h00, h01, h02, h03 = self._had128_quad(x00, x01, x02, x03, lane) + h10, h11, h12, h13 = self._had128_quad(x10, x11, x12, x13, lane) + h20, h21, h22, h23 = self._had128_quad(x20, x21, x22, x23, lane) + h30, h31, h32, h33 = self._had128_quad(x30, x31, x32, x33, lane) c00, c10, c20, c30 = self._had4_normalized(h00, h10, h20, h30) c01, c11, c21, c31 = self._had4_normalized(h01, h11, h21, h31) c02, c12, c22, c32 = self._had4_normalized(h02, h12, h22, h32) @@ -6997,9 +7619,7 @@ def _run_activation_coupled( route_pos = unit // nblk post_block = unit - route_pos * nblk row = packed_route_indices[route_pos].to(Int32) - expert = block_expert_ids[route_pos // Int32(self.moe_block_size)].to( - Int32 - ) + expert = block_expert_ids[route_pos // Int32(self.moe_block_size)].to(Int32) if cutlass.const_expr(self.direct_topk_routes): row = route_pos expert = packed_route_indices[route_pos].to(Int32) @@ -7448,7 +8068,7 @@ def _run_activation( active_m: cutlass.Int32, ): if cutlass.const_expr(self.intermediate_rotation): - # Rotation-aware epilogue (trellis_t256 tail). Warp-cooperative over + # Rotation-aware epilogue (trellis3_t256 tail). Warp-cooperative over # (routed-row, 128-block) units; each warp owns one 128-wide block of # a row's intermediate. Per row r the FC1 output is [gate(I) | up(I)] # (fc1_cols = 2I) and rot_scales_flat[r] = [svh_gate(I)|svh_up(I)| @@ -7574,8 +8194,6 @@ def _run_activation( idx += stride - - class W4A16ActivationKernel: def __init__( self, @@ -7881,9 +8499,7 @@ def kernel( ): tidx, _, _ = cute.arch.thread_idx() bidx, _, _ = cute.arch.block_idx() - if cutlass.const_expr( - self.coupled_hadamard and self.broadcast_svh - ): + if cutlass.const_expr(self.coupled_hadamard and self.broadcast_svh): # The output scale is shared by every expert. Linearity therefore # permits route reduction before the ordinary H128 cancellation: # @@ -7903,9 +8519,7 @@ def kernel( blk = unit - token * nblk block_col = blk * Int32(512) reduced_ptr = cute.arch.alloc_smem(cutlass.Float32, 512) - reduced = cute.make_tensor( - reduced_ptr, cute.make_layout(512) - ) + reduced = cute.make_tensor(reduced_ptr, cute.make_layout(512)) if warp < Int32(4): sub = warp @@ -7928,20 +8542,16 @@ def kernel( weight = topk_weights_flat[row].to(cutlass.Float32) base = row * Int32(self.hidden_size) + col0 acc0 += ( - fc2_flat[base + Int32(0)].to(cutlass.Float32) - * weight + fc2_flat[base + Int32(0)].to(cutlass.Float32) * weight ) acc1 += ( - fc2_flat[base + Int32(1)].to(cutlass.Float32) - * weight + fc2_flat[base + Int32(1)].to(cutlass.Float32) * weight ) acc2 += ( - fc2_flat[base + Int32(2)].to(cutlass.Float32) - * weight + fc2_flat[base + Int32(2)].to(cutlass.Float32) * weight ) acc3 += ( - fc2_flat[base + Int32(3)].to(cutlass.Float32) - * weight + fc2_flat[base + Int32(3)].to(cutlass.Float32) * weight ) acc0, acc1, acc2, acc3 = self._had128_quad( acc0, acc1, acc2, acc3, lane @@ -8065,7 +8675,9 @@ def kernel( acc2 = cutlass.Float32(0.0) acc3 = cutlass.Float32(0.0) for route in cutlass.range_constexpr(self.topk): - value_base = Int32(route * 512) + sub * Int32(128) + lane * Int32(4) + value_base = ( + Int32(route * 512) + sub * Int32(128) + lane * Int32(4) + ) weight = route_weights[Int32(route)] acc0 += route_values[value_base + Int32(0)] * weight acc1 += route_values[value_base + Int32(1)] * weight @@ -8271,11 +8883,156 @@ def _had4_normalized( ) +class W4A16DenseHadamard128Kernel: + """FP16 blockwise H128 used by native dense Trellis linears. + + EXL3 applies an incoherence scale before the input rotation and after the + output rotation. Keeping both forms in one B12X kernel removes the + runtime dependency on exllamav3_ext while preserving that ordering. + """ + + def __init__(self, *, width: int, scale_before: bool): + if width <= 0 or width % 128 != 0: + raise ValueError("dense H128 width must be a positive multiple of 128") + self.width = int(width) + self.scale_before = bool(scale_before) + self.cta_threads = 256 + + @property + def __cache_key__(self) -> tuple[object, ...]: + return (self.width, self.scale_before, self.cta_threads) + + @cute.jit + def __call__( + self, + input_ptr: cute.Pointer, + output_ptr: cute.Pointer, + scale_ptr: cute.Pointer, + active_m: cutlass.Int32, + stream: cuda.CUstream, + ): + input_flat = cute.make_tensor( + input_ptr, + layout=cute.make_layout((active_m * Int32(self.width),), stride=(1,)), + ) + output_flat = cute.make_tensor( + output_ptr, + layout=cute.make_layout((active_m * Int32(self.width),), stride=(1,)), + ) + scale_flat = cute.make_tensor( + scale_ptr, + layout=cute.make_layout((Int32(self.width),), stride=(1,)), + ) + total_units = active_m * Int32(self.width // 128) + grid = (_covering_count(total_units, self.cta_threads // 32), 1, 1) + self.kernel(input_flat, output_flat, scale_flat, active_m).launch( + grid=grid, + block=[self.cta_threads, 1, 1], + stream=stream, + ) + + @cute.kernel + def kernel( + self, + input_flat: cute.Tensor, + output_flat: cute.Tensor, + scale_flat: cute.Tensor, + active_m: cutlass.Int32, + ): + tidx, _, _ = cute.arch.thread_idx() + bidx, _, _ = cute.arch.block_idx() + tid = Int32(tidx) + lane = tid & Int32(31) + warp = tid >> Int32(5) + unit = Int32(bidx) * Int32(self.cta_threads // 32) + warp + nblocks = Int32(self.width // 128) + total_units = active_m * nblocks + if unit < total_units: + row = unit // nblocks + block = unit - row * nblocks + col0 = block * Int32(128) + lane * Int32(4) + base = row * Int32(self.width) + col0 + v0 = input_flat[base + Int32(0)].to(cutlass.Float32) + v1 = input_flat[base + Int32(1)].to(cutlass.Float32) + v2 = input_flat[base + Int32(2)].to(cutlass.Float32) + v3 = input_flat[base + Int32(3)].to(cutlass.Float32) + if cutlass.const_expr(self.scale_before): + # exllamav3 uses __hmul2 before the transform; preserve the + # intermediate fp16 rounding instead of promoting the product. + v0 = cutlass.Float16( + v0 * scale_flat[col0 + Int32(0)].to(cutlass.Float32) + ).to(cutlass.Float32) + v1 = cutlass.Float16( + v1 * scale_flat[col0 + Int32(1)].to(cutlass.Float32) + ).to(cutlass.Float32) + v2 = cutlass.Float16( + v2 * scale_flat[col0 + Int32(2)].to(cutlass.Float32) + ).to(cutlass.Float32) + v3 = cutlass.Float16( + v3 * scale_flat[col0 + Int32(3)].to(cutlass.Float32) + ).to(cutlass.Float32) + h0, h1, h2, h3 = self._had128_quad(v0, v1, v2, v3, lane) + if cutlass.const_expr(not self.scale_before): + # The reference rounds H128 to fp16 before the post-scale. + h0 = cutlass.Float16(h0).to(cutlass.Float32) * scale_flat[ + col0 + Int32(0) + ].to(cutlass.Float32) + h1 = cutlass.Float16(h1).to(cutlass.Float32) * scale_flat[ + col0 + Int32(1) + ].to(cutlass.Float32) + h2 = cutlass.Float16(h2).to(cutlass.Float32) * scale_flat[ + col0 + Int32(2) + ].to(cutlass.Float32) + h3 = cutlass.Float16(h3).to(cutlass.Float32) * scale_flat[ + col0 + Int32(3) + ].to(cutlass.Float32) + output_flat[base + Int32(0)] = cutlass.Float16(h0) + output_flat[base + Int32(1)] = cutlass.Float16(h1) + output_flat[base + Int32(2)] = cutlass.Float16(h2) + output_flat[base + Int32(3)] = cutlass.Float16(h3) + + @cute.jit + def _had128_quad( + self, + v0: cutlass.Float32, + v1: cutlass.Float32, + v2: cutlass.Float32, + v3: cutlass.Float32, + lane: Int32, + ): + s0 = v0 + v1 + d0 = v0 - v1 + s1 = v2 + v3 + d1 = v2 - v3 + h0 = s0 + s1 + h1 = d0 + d1 + h2 = s0 - s1 + h3 = d0 - d1 + for i in cutlass.range_constexpr(5): + step = 1 << i + p0 = cute.arch.shuffle_sync_bfly(h0, offset=step) + p1 = cute.arch.shuffle_sync_bfly(h1, offset=step) + p2 = cute.arch.shuffle_sync_bfly(h2, offset=step) + p3 = cute.arch.shuffle_sync_bfly(h3, offset=step) + if (lane & Int32(step)) != Int32(0): + h0 = p0 - h0 + h1 = p1 - h1 + h2 = p2 - h2 + h3 = p3 - h3 + else: + h0 = p0 + h0 + h1 = p1 + h1 + h2 = p2 + h2 + h3 = p3 + h3 + scale = cutlass.Float32(0.088388347648) + return h0 * scale, h1 * scale, h2 * scale, h3 * scale + _CACHE: dict[tuple, W4A16GemmCompileResult] = {} _FUSED_CACHE: dict[tuple, W4A16FusedMoeCompileResult] = {} _ACTIVATION_CACHE: dict[tuple, W4A16ActivationCompileResult] = {} _SUM_CACHE: dict[tuple, W4A16TopKSumCompileResult] = {} +_DENSE_HAD128_CACHE: dict[tuple, object] = {} _SMALL_M_DIRECT_CACHE: dict[tuple, _W4A16SmallMDirectLaunch] = {} _FC2_DIRECT_CACHE: dict[tuple, _W4A16FC2DirectLaunch] = {} @@ -8497,12 +9254,13 @@ def dummy(dt): dummy(cutlass.BFloat16), barrier_fake, barrier_fake, + Int32(num_experts), Int32(m), Int32(kernel.grid_x), current_cuda_stream(), compile_spec=KernelCompileSpec.from_facts( "moe.w4a16.small_m_direct", - 2, + 3, ("device_index", None if device is None else int(device.index or 0)), ("m", int(m)), ("hidden_size", int(hidden_size)), @@ -8634,12 +9392,13 @@ def dummy(dt): dummy(cutlass.BFloat16), barrier_fake, barrier_fake, + Int32(expert_capacity), Int32(2), Int32(kernel.grid_x), current_cuda_stream(), compile_spec=KernelCompileSpec.from_facts( "moe.w4a16.fc2_direct", - 4, + 5, ("device_index", int(device.index or 0)), ("hidden_size", int(hidden_size)), ("intermediate_size", int(intermediate_size)), @@ -8677,7 +9436,7 @@ def compile_w4a16_gemm( scale_format: str = "e4m3_k16", w13_layout: str = "packed", trellis_bits: int = 3, - trellis_codebook: str = SQG_E4M3, + trellis_codebook: str = "sqg_xor_cheb_t12", trellis_pair_kind: str | None = None, trellis_rate_axis: str | None = None, dense_route_fast_path: bool = False, @@ -8708,7 +9467,7 @@ def compile_w4a16_gemm( trellis_pair_kind=trellis_pair_kind, trellis_rate_axis=trellis_rate_axis, dense_route_fast_path=bool(dense_route_fast_path), - schedule_whole_tiles=weight_layout == "trellis_t256", + schedule_whole_tiles=weight_layout == "trellis3_t256", ) cache_key = ( "w4a16_gemm", @@ -8727,7 +9486,7 @@ def compile_w4a16_gemm( compile_route_blocks = 1 compile_route_slots = compile_route_blocks * int(moe_block_size) a_fake = make_ptr(cutlass_dtype, 16, cute.AddressSpace.gmem, assumed_align=16) - if weight_layout == "trellis_t256": + if weight_layout == "trellis3_t256": b_fake_elements = ( num_experts * (size_k // 16) * (size_n // 16) * (8 * int(trellis_bits)) ) @@ -8820,7 +9579,7 @@ def compile_w4a16_gemm( current_cuda_stream(), compile_spec=KernelCompileSpec.from_key( "moe.w4a16.gemm", - 4, + 3, cache_key, ), ) @@ -8867,12 +9626,13 @@ def compile_w4a16_fused_moe( scale_format: str = "e4m3_k16", w13_layout: str = "w13", trellis_bits: int = 3, - trellis_codebook: str = SQG_E4M3, + trellis_codebook: str = "sqg_xor_cheb_t12", fc1_trellis_pair_kind: str | None = None, fc2_trellis_pair_kind: str | None = None, direct_topk_routes: bool = False, use_expert_map: bool = False, tc_decode_fused_sum: bool = False, + prefill_fused_sum_fp32: bool = False, collect_activation_amax: bool = False, force_tile_config: tuple[int, int, int, int] | None = None, intermediate_rotation: bool = False, @@ -8902,16 +9662,16 @@ def compile_w4a16_fused_moe( if weight_layout not in _WEIGHT_LAYOUTS: raise ValueError(f"unsupported W4A16 weight_layout {weight_layout!r}") trellis_bits = int(trellis_bits) - if weight_layout == "trellis_t256": + if weight_layout == "trellis3_t256": if trellis_bits not in _TRELLIS256_BITS: raise ValueError( - f"trellis_t256 bits must be one of {_TRELLIS256_BITS}, got {trellis_bits}" + f"trellis3_t256 bits must be one of {_TRELLIS256_BITS}, got {trellis_bits}" ) elif trellis_bits != 3: - raise ValueError("trellis_bits is only valid for trellis_t256 weights") + raise ValueError("trellis_bits is only valid for trellis3_t256 weights") # Existing 3-bpw scheduling was conservatively planned as 4 bpw. Keep that # grid contract stable for D6; widen only the 5/6-bpw specializations. - weight_bits = max(4, trellis_bits) if weight_layout == "trellis_t256" else 4 + weight_bits = max(4, trellis_bits) if weight_layout == "trellis3_t256" else 4 # GATE 5: the PRODUCTION 256-weight-tile fused-megakernel B-staging is now # wired (per-warp native [K/16,N/16,8*bits u32] tile staging + the per-lane # bitrate-specialized read) and ADMITTED at 3 bpw against a full-GEMM @@ -8921,18 +9681,25 @@ def compile_w4a16_fused_moe( if weight_layout == "modelopt": if w13_layout not in _MODEL_OPT_W13_LAYOUTS: raise ValueError(f"unsupported W4A16 w13_layout {w13_layout!r}") - elif weight_layout == "trellis_t256": + elif weight_layout == "trellis3_t256": if w13_layout not in _TRELLIS256_W13_LAYOUTS: - raise ValueError(f"unsupported trellis_t256 w13_layout {w13_layout!r}") + raise ValueError(f"unsupported trellis3_t256 w13_layout {w13_layout!r}") else: w13_layout = "packed" direct_topk_routes = bool(direct_topk_routes) use_expert_map = bool(use_expert_map) tc_decode_fused_sum = bool(tc_decode_fused_sum) + prefill_fused_sum_fp32 = bool(prefill_fused_sum_fp32) + if tc_decode_fused_sum and prefill_fused_sum_fp32: + raise ValueError( + "TC-decode and large-M FP32 route reduction are mutually exclusive" + ) if use_expert_map and not direct_topk_routes: raise ValueError("use_expert_map requires direct_topk_routes") collect_activation_amax = bool(collect_activation_amax) - if collect_activation_amax and (direct_topk_routes or tc_decode_fused_sum): + if collect_activation_amax and ( + direct_topk_routes or tc_decode_fused_sum or prefill_fused_sum_fp32 + ): raise ValueError( "W4A16 activation amax collection requires the route-packed fused path" ) @@ -8943,25 +9710,25 @@ def compile_w4a16_fused_moe( if full_rotation: if not intermediate_rotation: raise ValueError("full_rotation requires intermediate_rotation") - if weight_layout != "trellis_t256": - raise ValueError("full_rotation is only supported for trellis_t256") + if weight_layout != "trellis3_t256": + raise ValueError("full_rotation is only supported for trellis3_t256") if element_dtype != "fp16": raise ValueError("full_rotation requires element_dtype='fp16'") if rotation_input_dtype not in {"bf16", "fp16"}: raise ValueError( "rotation_input_dtype must be 'bf16' or 'fp16' for full_rotation" ) - if tc_decode_fused_sum: - raise ValueError("full_rotation is incompatible with TC decode") + if tc_decode_fused_sum or prefill_fused_sum_fp32: + raise ValueError("full_rotation is incompatible with fused route reduction") if apply_router_weight_on_input: raise ValueError( "full_rotation requires apply_router_weight_on_input=False" ) if coupled_hadamard and not full_rotation: raise ValueError("coupled_hadamard requires full_rotation") - if collect_activation_amax and weight_layout == "trellis_t256": + if collect_activation_amax and weight_layout == "trellis3_t256": raise NotImplementedError( - "trellis_t256 activation-amax collection is not exposed through the " + "trellis3_t256 activation-amax collection is not exposed through the " "registered launch ABI; refusing to compile a bitrate-ambiguous kernel" ) # The TC-decode path validates M in {1,2,4,8} itself and uses direct-topk @@ -8972,7 +9739,7 @@ def compile_w4a16_fused_moe( else _MAX_DIRECT_TOPK_ROUTE_M ) direct_weight_layout_ok = weight_layout == "packed" or ( - full_rotation and weight_layout == "trellis_t256" + full_rotation and weight_layout == "trellis3_t256" ) if direct_topk_routes and ( int(size_m) > direct_topk_m_cap @@ -9254,6 +10021,7 @@ def compile_w4a16_fused_moe( direct_topk_routes=direct_topk_routes, use_expert_map=use_expert_map, tc_decode_fused_sum=tc_decode_fused_sum, + prefill_fused_sum_fp32=prefill_fused_sum_fp32, collect_activation_amax=collect_activation_amax, intermediate_rotation=intermediate_rotation, full_rotation=full_rotation, @@ -9355,14 +10123,16 @@ def compile_w4a16_fused_moe( assumed_align=16, ) fc2_fake = cute.runtime.make_fake_compact_tensor( - cutlass_dtype, - (compile_routed_rows * hidden_size,), + cutlass.Float32 if kernel.prefill_fused_sum_fp32 else cutlass_dtype, + ( + compile_size_m * hidden_size + if kernel.prefill_fused_sum_fp32 + else compile_routed_rows * hidden_size + ,), assumed_align=16, ) pair_metadata_cutlass_dtype = ( - cutlass.Int64 - if fc1_trellis_pair_kind == "P33_P43" - else cutlass.Int32 + cutlass.Int64 if fc1_trellis_pair_kind == "P33_P43" else cutlass.Int32 ) w13_scales_fake = make_ptr( pair_metadata_cutlass_dtype, 16, cute.AddressSpace.gmem, assumed_align=16 @@ -9479,11 +10249,17 @@ def compile_w4a16_fused_moe( current_cuda_stream(), compile_spec=KernelCompileSpec.from_key( "moe.w4a16.fused_moe", - 8, + 7, cache_key, ), dsl_compile_options=OptLevel(2), ) + resources = _query_w4a16_kernel_resources(compiled) + kernel_symbol = None + registers_per_thread = -1 + local_memory_bytes = -1 + if resources is not None: + kernel_symbol, registers_per_thread, local_memory_bytes = resources result = W4A16FusedMoeCompileResult( compiled=compiled, size_m=size_m, @@ -9512,6 +10288,7 @@ def compile_w4a16_fused_moe( use_expert_map=kernel.use_expert_map, scale_format=scale_format, tc_decode_fused_sum=bool(tc_decode_fused_sum), + prefill_fused_sum_fp32=bool(prefill_fused_sum_fp32), collect_activation_amax=collect_activation_amax, schedule_whole_tiles=kernel.schedule_whole_tiles, intermediate_rotation=intermediate_rotation, @@ -9523,12 +10300,54 @@ def compile_w4a16_fused_moe( full_rotation=full_rotation, coupled_hadamard=coupled_hadamard, rotation_input_dtype=rotation_input_dtype, + kernel_symbol=kernel_symbol, + registers_per_thread=registers_per_thread, + local_memory_bytes=local_memory_bytes, cta_threads=kernel.cta_threads, shared_memory_bytes=kernel.shared_words * 4, ) _FUSED_CACHE[cache_key] = result return result + +def _query_w4a16_kernel_resources(compiled: object) -> tuple[str, int, int] | None: + """Return (symbol, registers/thread, local bytes/thread) for a one-kernel + CUDA-dialect compile result, or None when the object does not expose the + introspection surface (e.g. an on-disk object-cache reload).""" + + kernel_info = getattr(compiled, "kernel_info", None) + to_executor = getattr(compiled, "to", None) + if not isinstance(kernel_info, dict) or not callable(to_executor): + return None + symbols = tuple(kernel_info) + if len(symbols) != 1 or not isinstance(symbols[0], str) or not symbols[0]: + return None + executor = to_executor(int(torch.cuda.current_device())) + libraries = tuple( + getattr(getattr(executor, "jit_module", None), "cuda_library", None) or () + ) + if len(libraries) != 1: + return None + success = cuda_runtime.cudaError_t(0) + kernel_status, kernel_handle = cuda_runtime.cudaLibraryGetKernel( + libraries[0], symbols[0].encode("utf-8") + ) + if kernel_status != success: + raise RuntimeError( + f"cudaLibraryGetKernel failed for {symbols[0]}: {kernel_status}" + ) + attributes_status, attributes = cuda_runtime.cudaFuncGetAttributes(kernel_handle) + if attributes_status != success: + raise RuntimeError( + f"cudaFuncGetAttributes failed for {symbols[0]}: {attributes_status}" + ) + registers_per_thread = int(getattr(attributes, "numRegs", -1)) + local_memory_bytes = int(getattr(attributes, "localSizeBytes", -1)) + if registers_per_thread < 0 or local_memory_bytes < 0: + raise RuntimeError(f"incomplete CUDA function attributes for {symbols[0]}") + return symbols[0], registers_per_thread, local_memory_bytes + + def _w4a16_weight_flat_elements( *, num_experts: int, @@ -9544,8 +10363,6 @@ def _w4a16_weight_flat_elements( return int(num_experts) * (int(size_k) // 16) * (int(size_n) // 16 * 32) - - def clear_w4a16_kernel_cache() -> None: _CACHE.clear() _FUSED_CACHE.clear() @@ -9806,6 +10623,7 @@ def ptr(dt, tensor: torch.Tensor): ptr(cutlass.BFloat16, output), barrier_count, barrier_epoch, + Int32(num_experts), Int32(m), Int32(direct_launch.grid_x), cuda.CUstream(stream_int), @@ -9953,6 +10771,7 @@ def ptr(dt, tensor: torch.Tensor): ptr(cutlass.BFloat16, output), barrier_count, barrier_epoch, + Int32(num_experts), Int32(m), Int32(launch.grid_x), cuda.CUstream(stream_int), @@ -10084,6 +10903,7 @@ def _w4a16_fused_moe_launch_flat( fc2_tile_n: int, direct_topk_routes: bool, tc_decode_fused_sum: bool, + prefill_fused_sum_fp32: bool, collect_activation_amax: bool, stream_int: int, expert_map: torch.Tensor | None = None, @@ -10091,7 +10911,7 @@ def _w4a16_fused_moe_launch_flat( intermediate_rotation: bool = False, a_input_up: torch.Tensor | None = None, trellis_bits: int = 3, - trellis_codebook: str = SQG_E4M3, + trellis_codebook: str = "sqg_xor_cheb_t12", fc1_trellis_pair_kind: str | None = None, fc2_trellis_pair_kind: str | None = None, full_rotation: bool = False, @@ -10150,11 +10970,26 @@ def _w4a16_fused_moe_launch_flat( ) suh_gate_arg = suh_gate_table.reshape(-1) suh_up_arg = suh_up_table.reshape(-1) - broadcast_suh = suh_gate_arg.numel() == hidden_size - if broadcast_suh != (suh_up_arg.numel() == hidden_size): + expanded_suh_elements = num_experts * hidden_size + valid_suh_elements = {hidden_size, expanded_suh_elements} + if ( + suh_gate_arg.numel() not in valid_suh_elements + or suh_up_arg.numel() not in valid_suh_elements + ): + raise ValueError( + "suh gate/up tables must contain either one broadcast row or " + f"one row per expert ({hidden_size} or {expanded_suh_elements} " + "elements)" + ) + gate_broadcast_suh = suh_gate_arg.numel() == hidden_size + up_broadcast_suh = suh_up_arg.numel() == hidden_size + if gate_broadcast_suh != up_broadcast_suh: raise ValueError( "suh gate/up tables must both be per-expert or both broadcast" ) + # For one expert, per-expert and broadcast storage are identical. Use + # the broadcast specialization so this valid tier remains unambiguous. + broadcast_suh = gate_broadcast_suh rotation_input_dtype = _normalize_element_dtype(rotation_input.dtype) else: suh_gate_arg = _rot_scales_dummy(w13_global_scale.device) @@ -10188,6 +11023,7 @@ def _w4a16_fused_moe_launch_flat( direct_topk_routes=bool(direct_topk_routes), use_expert_map=use_expert_map, tc_decode_fused_sum=bool(tc_decode_fused_sum), + prefill_fused_sum_fp32=bool(prefill_fused_sum_fp32), collect_activation_amax=collect_activation_amax, # The custom-op boundary cannot carry the compiled launch object. Re-pin # its selected geometry so tile-specific packs resolve the @@ -10206,14 +11042,12 @@ def _w4a16_fused_moe_launch_flat( packed_route_indices.data_ptr() if expert_map is None else expert_map.data_ptr() ) route_num_experts = 0 if expert_map is None else int(expert_map.numel()) - if weight_layout == "trellis_t256" and trellis_codebook != "mcg": - trellis_rank_lut = _trellis256_execution_lut( - a_input.device, trellis_codebook - ) + if weight_layout == "trellis3_t256": + trellis_rank_lut = _trellis256_execution_lut(a_input.device, trellis_codebook) fc1_trellis_lut_addr = trellis_rank_lut.data_ptr() fc2_trellis_lut_addr = trellis_rank_lut.data_ptr() else: - # Non-trellis and MCG kernels never dereference this ABI slot. + # Non-trellis kernels never dereference this ABI slot. fc1_trellis_lut_addr = w13_scale_i32.data_ptr() fc2_trellis_lut_addr = w13_scale_i32.data_ptr() fused.compiled( @@ -10253,17 +11087,13 @@ def _w4a16_fused_moe_launch_flat( activated, fc2_out, make_ptr( - cutlass.Int64 - if fc1_trellis_pair_kind == "P33_P43" - else cutlass.Int32, + cutlass.Int64 if fc1_trellis_pair_kind == "P33_P43" else cutlass.Int32, w13_scale_i32.data_ptr(), cute.AddressSpace.gmem, assumed_align=16, ), make_ptr( - cutlass.Int64 - if fc2_trellis_pair_kind == "P33_P43" - else cutlass.Int32, + cutlass.Int64 if fc2_trellis_pair_kind == "P33_P43" else cutlass.Int32, w2_scale_i32.data_ptr(), cute.AddressSpace.gmem, assumed_align=16, @@ -10520,6 +11350,7 @@ def _w4a16_fused_moe_launch_op( fc2_tile_n=fc2_tile_n, direct_topk_routes=direct_topk_routes, tc_decode_fused_sum=tc_decode_fused_sum, + prefill_fused_sum_fp32=False, collect_activation_amax=False, stream_int=stream_int, ) @@ -10677,6 +11508,7 @@ def _w4a16_fused_moe_calibrated_launch_op( fc2_tile_n=fc2_tile_n, direct_topk_routes=False, tc_decode_fused_sum=False, + prefill_fused_sum_fp32=False, collect_activation_amax=True, stream_int=stream_int, ) @@ -10758,9 +11590,7 @@ def _w4a16_topk_sum_launch_flat( ) route_num_experts = 0 if expert_map is None else int(expert_map.numel()) broadcast_svh = ( - full_rotation - and svh_table is not None - and svh_table.numel() == hidden_size + full_rotation and svh_table is not None and svh_table.numel() == hidden_size ) sum_kernel = compile_w4a16_topk_sum( m=m, @@ -10987,7 +11817,7 @@ def _compile_w4a16_gemm_launch( scale_format: str = "e4m3_k16", w13_layout: str = "packed", trellis_bits: int = 3, - trellis_codebook: str = SQG_E4M3, + trellis_codebook: str = "sqg_xor_cheb_t12", trellis_pair_kind: str | None = None, trellis_rate_axis: str | None = None, dense_route_fast_path: bool = False, @@ -10995,7 +11825,7 @@ def _compile_w4a16_gemm_launch( force_tile_config: tuple[int, int] | None = None, ) -> _W4A16GemmLaunch: planner_weight_bits = ( - max(4, int(trellis_bits)) if weight_layout == "trellis_t256" else 4 + max(4, int(trellis_bits)) if weight_layout == "trellis3_t256" else 4 ) if force_tile_config is None: tile_k, tile_n, _, _ = _select_tile_config( @@ -11106,7 +11936,7 @@ def pack_topk_routes_by_expert( def _trellis256_dense_tile_config(size_k: int, size_n: int) -> tuple[int, int]: """Return an m-invariant t256 tile so dense row bits cannot drift with m.""" if int(size_k) % 64 != 0: - raise ValueError(f"trellis_t256 dense K must be divisible by 64, got {size_k}") + raise ValueError(f"trellis3_t256 dense K must be divisible by 64, got {size_k}") if int(size_n) % 256 == 0: return (64, 256) if int(size_n) % 128 == 0: @@ -11116,7 +11946,7 @@ def _trellis256_dense_tile_config(size_k: int, size_n: int) -> tuple[int, int]: # 256-thread 4/8/8 register entry and fail before compilation. return (64, 128) raise ValueError( - "trellis_t256 dense N must be divisible by 128 (or 256 for the wide tile), " + "trellis3_t256 dense N must be divisible by 128 (or 256 for the wide tile), " f"got N={size_n}" ) @@ -11156,6 +11986,73 @@ def _trellis256_dense_launch_geometry( return default +def _compile_trellis_dense_hadamard128(*, width: int, scale_before: bool): + cache_key = ("trellis_dense_hadamard128", int(width), bool(scale_before)) + cached = _DENSE_HAD128_CACHE.get(cache_key) + if cached is not None: + return cached + kernel = W4A16DenseHadamard128Kernel( + width=int(width), + scale_before=bool(scale_before), + ) + fp16_fake = make_ptr(cutlass.Float16, 16, cute.AddressSpace.gmem, assumed_align=16) + raise_if_kernel_resolution_frozen( + "cute.compile", target=kernel, cache_key=cache_key + ) + compiled = b12x_compile( + kernel, + fp16_fake, + fp16_fake, + fp16_fake, + Int32(1), + current_cuda_stream(), + compile_spec=KernelCompileSpec.from_key( + "gemm.trellis_dense.hadamard128", + 1, + cache_key, + ), + ) + _DENSE_HAD128_CACHE[cache_key] = compiled + return compiled + + +def _run_trellis_dense_hadamard128( + x: torch.Tensor, + output: torch.Tensor, + scale: torch.Tensor, + *, + scale_before: bool, +) -> None: + if x.dtype != torch.float16 or output.dtype != torch.float16: + raise TypeError("native dense H128 requires fp16 input and output") + if x.ndim != 2 or output.shape != x.shape or not x.is_contiguous(): + raise ValueError("native dense H128 requires equal contiguous rank-2 tensors") + if not output.is_contiguous() or output.device != x.device: + raise ValueError( + "native dense H128 output must be contiguous on the input device" + ) + if ( + scale.dtype != torch.float16 + or scale.device != x.device + or not scale.is_contiguous() + or scale.numel() != x.shape[1] + ): + raise ValueError( + "native dense H128 scale must be contiguous fp16 with width elements" + ) + compiled = _compile_trellis_dense_hadamard128( + width=int(x.shape[1]), + scale_before=bool(scale_before), + ) + fp16 = cutlass.Float16 + compiled( + make_ptr(fp16, x.data_ptr(), cute.AddressSpace.gmem, assumed_align=16), + make_ptr(fp16, output.data_ptr(), cute.AddressSpace.gmem, assumed_align=16), + make_ptr(fp16, scale.data_ptr(), cute.AddressSpace.gmem, assumed_align=16), + int(x.shape[0]), + current_cuda_stream(), + ) + def _resolve_exl3_hadamard_128(hadamard_128): if hadamard_128 is None: @@ -11204,6 +12101,30 @@ def _trellis_dense_buffer( return buffer +def _use_k6_mcg_small( + *, + device: torch.device, + m: int, + trellis_bits: int, + trellis_codebook: str, + trellis_pair_kind, + compute_dtype: torch.dtype, + external_hadamard_128, + explicit_launch_config: bool, +) -> bool: + """Select the capture-safe K6/MCG kernel only on its compiled target.""" + return ( + not explicit_launch_config + and tuple(torch.cuda.get_device_capability(device)) == (12, 0) + and m <= 128 + and trellis_bits == 6 + and trellis_codebook == "mcg" + and trellis_pair_kind is None + and compute_dtype == torch.float16 + and external_hadamard_128 is None + ) + + def _run_trellis256_dense_current_device( x: torch.Tensor, prepared_dense, @@ -11218,7 +12139,7 @@ def _run_trellis256_dense_current_device( output_f16: torch.Tensor | None = None, hadamard_128=None, stream: cuda.CUstream | None = None, - _moe_block_size: int = 64, + _moe_block_size: int | None = None, _force_tile_config: tuple[int, int] | None = None, ) -> torch.Tensor: """Run one native EXL3 linear on the already-selected CUDA device. @@ -11228,17 +12149,14 @@ def _run_trellis256_dense_current_device( Outer rotations follow EXL3 order exactly: fp16 ``suh`` multiply before the input H128, and fp16 ``svh`` multiply after the output H128. """ - if getattr(prepared_dense, "weight_layout", None) != "trellis_t256": - raise ValueError("run_trellis256_dense requires prepared trellis_t256 weights") + if getattr(prepared_dense, "weight_layout", None) != "trellis3_t256": + raise ValueError("run_trellis256_dense requires prepared trellis3_t256 weights") if int(getattr(prepared_dense, "num_experts", 0)) != 1: raise ValueError("run_trellis256_dense requires an honest E=1 prepared weight") - trellis_codebook = str( - getattr(prepared_dense, "trellis_codebook", "") - ).lower() + trellis_codebook = str(getattr(prepared_dense, "trellis_codebook", "")).lower() if trellis_codebook not in _TRELLIS256_CODEBOOKS: raise NotImplementedError( - "run_trellis256_dense has no decoder for codebook " - f"{trellis_codebook!r}" + f"run_trellis256_dense has no decoder for codebook {trellis_codebook!r}" ) trellis_bits = int(getattr(prepared_dense, "trellis_bits", 0)) if trellis_bits not in _TRELLIS256_BITS: @@ -11283,7 +12201,76 @@ def _run_trellis256_dense_current_device( if c_tmp is not None and int(c_tmp.data_ptr()) % 16 != 0: raise ValueError("c_tmp must be at least 16-byte aligned") - hadamard_128 = _resolve_exl3_hadamard_128(hadamard_128) + external_hadamard_128 = ( + None if hadamard_128 is None else _resolve_exl3_hadamard_128(hadamard_128) + ) + + # Keep the established K6/MCG decode path independent from the generic + # Trellis scheduler. It owns both H128 rotations, needs no GEMM scratch, + # and is safe to capture with only caller-owned output/rotation storage. + # Compact pair payloads and the newer SQG codebooks use the generic path. + use_k6_mcg_small = _use_k6_mcg_small( + device=x.device, + m=m, + trellis_bits=trellis_bits, + trellis_codebook=trellis_codebook, + trellis_pair_kind=trellis_pair_kind, + compute_dtype=compute_dtype, + external_hadamard_128=external_hadamard_128, + explicit_launch_config=( + _moe_block_size is not None or _force_tile_config is not None + ), + ) + if use_k6_mcg_small: + if x.dtype == torch.float16: + x_f16 = x + else: + input_f16 = _trellis_dense_buffer( + "input_f16", + input_f16, + shape=(m, size_k), + dtype=torch.float16, + device=x.device, + ) + input_f16.copy_(x) + x_f16 = input_f16 + rotated_f16 = _trellis_dense_buffer( + "rotated_f16", + rotated_f16, + shape=(m, size_k), + dtype=torch.float16, + device=x.device, + ) + if output.dtype == torch.float16: + small_output = output + else: + output_f16 = _trellis_dense_buffer( + "output_f16", + output_f16, + shape=(m, size_n), + dtype=torch.float16, + device=x.device, + ) + small_output = output_f16 + from b12x.gemm.trellis_linear._small_m import run_k6_mcg + + trellis_i16 = prepared_dense.trellis.view(torch.int16).view( + size_k // 16, + size_n // 16, + trellis_bits * 16, + ) + run_k6_mcg( + x_f16, + trellis_i16, + small_output, + prepared_dense.suh, + rotated_f16, + prepared_dense.svh, + prepared_dense.workspace, + ) + if output.dtype != torch.float16: + output.copy_(small_output) + return output gemm_output = _trellis_dense_buffer( "gemm_output", @@ -11311,7 +12298,15 @@ def _run_trellis256_dense_current_device( dtype=torch.float16, device=x.device, ) - hadamard_128(x_f16, rotated_f16, prepared_dense.suh, None, 1.0) + if external_hadamard_128 is None: + _run_trellis_dense_hadamard128( + x_f16, + rotated_f16, + prepared_dense.suh, + scale_before=True, + ) + else: + external_hadamard_128(x_f16, rotated_f16, prepared_dense.suh, None, 1.0) if compute_dtype == torch.float16: rotated_compute = rotated_f16 else: @@ -11329,8 +12324,8 @@ def _run_trellis256_dense_current_device( max_shared_mem = int( getattr(props, "shared_memory_per_block_optin", _DEFAULT_MAX_SHARED_MEM) ) - moe_block_size = int(_moe_block_size) - if _force_tile_config is None and moe_block_size == 64: + moe_block_size = 64 if _moe_block_size is None else int(_moe_block_size) + if _force_tile_config is None and _moe_block_size is None: moe_block_size, (tile_k, tile_n) = _trellis256_dense_launch_geometry( size_m=m, size_k=size_k, @@ -11362,7 +12357,7 @@ def _run_trellis256_dense_current_device( max_shared_mem=max_shared_mem, device=x.device, c_tmp=c_tmp, - weight_layout="trellis_t256", + weight_layout="trellis3_t256", scale_format="e4m3_k32", w13_layout="packed", trellis_bits=trellis_bits, @@ -11408,10 +12403,7 @@ def _run_trellis256_dense_current_device( prepared_dense.global_scale, launch.c_tmp, prepared_dense.workspace, - # MCG kernels never dereference the LUT ABI slot. - dummy_i32 - if trellis_codebook == "mcg" - else _trellis256_execution_lut(x.device, trellis_codebook), + _trellis256_execution_lut(x.device, trellis_codebook), m, grid_x, stream, @@ -11430,7 +12422,15 @@ def _run_trellis256_dense_current_device( gemm_output_f16.copy_(gemm_output) gemm_f16 = gemm_output_f16 if output.dtype == torch.float16: - hadamard_128(gemm_f16, output, None, prepared_dense.svh, 1.0) + if external_hadamard_128 is None: + _run_trellis_dense_hadamard128( + gemm_f16, + output, + prepared_dense.svh, + scale_before=False, + ) + else: + external_hadamard_128(gemm_f16, output, None, prepared_dense.svh, 1.0) else: output_f16 = _trellis_dense_buffer( "output_f16", @@ -11439,7 +12439,15 @@ def _run_trellis256_dense_current_device( dtype=torch.float16, device=x.device, ) - hadamard_128(gemm_f16, output_f16, None, prepared_dense.svh, 1.0) + if external_hadamard_128 is None: + _run_trellis_dense_hadamard128( + gemm_f16, + output_f16, + prepared_dense.svh, + scale_before=False, + ) + else: + external_hadamard_128(gemm_f16, output_f16, None, prepared_dense.svh, 1.0) output.copy_(output_f16) return output @@ -11458,7 +12466,7 @@ def run_trellis256_dense( output_f16: torch.Tensor | None = None, hadamard_128=None, stream: cuda.CUstream | None = None, - _moe_block_size: int = 64, + _moe_block_size: int | None = None, _force_tile_config: tuple[int, int] | None = None, ) -> torch.Tensor: """Run one native or compact P24/P33 EXL3 linear through the t256 GEMM. @@ -11546,6 +12554,7 @@ def run_w4a16_moe( intermediate_cache13: torch.Tensor, intermediate_cache2: torch.Tensor, output: torch.Tensor, + prefill_sum_accum: torch.Tensor | None = None, fc1_c_tmp: torch.Tensor | None = None, fc2_c_tmp: torch.Tensor | None = None, packed_route_indices: torch.Tensor | None = None, @@ -11610,7 +12619,7 @@ def run_w4a16_moe( trellis_bits = int(getattr(prepared, "trellis_bits", 3)) coupled_hadamard = bool(getattr(prepared, "coupled_hadamard", False)) trellis_codebook = str( - getattr(prepared, "trellis_codebook", SQG_E4M3) + getattr(prepared, "trellis_codebook", "sqg_xor_cheb_t12") ).lower() fc1_trellis_pair_kind = getattr(prepared, "fc1_trellis_pair_kind", None) fc2_trellis_pair_kind = getattr(prepared, "fc2_trellis_pair_kind", None) @@ -11619,15 +12628,15 @@ def run_w4a16_moe( prepared_tile_config = getattr(prepared, "tile_config", None) if (fc1_trellis_pair_kind is None) != (fc2_trellis_pair_kind is None): raise ValueError("prepared trellis pair weights have incomplete pair metadata") - if weight_layout == "trellis_t256": + if weight_layout == "trellis3_t256": if trellis_bits not in _TRELLIS256_BITS: raise ValueError( - f"prepared trellis_t256 bitrate must be in {_TRELLIS256_BITS}, " + f"prepared trellis3_t256 bitrate must be in {_TRELLIS256_BITS}, " f"got {trellis_bits}" ) if trellis_codebook not in _TRELLIS256_CODEBOOKS: raise NotImplementedError( - "trellis_t256 execution has no decoder for codebook " + "trellis3_t256 execution has no decoder for codebook " f"{trellis_codebook!r}" ) if fc1_trellis_pair_kind is not None: @@ -11650,14 +12659,11 @@ def run_w4a16_moe( ): raise ValueError("unsupported prepared static trellis pair kind") pair_metadata_dtype = ( - torch.int64 - if fc1_trellis_pair_kind == "P33_P43" - else torch.int32 + torch.int64 if fc1_trellis_pair_kind == "P33_P43" else torch.int32 ) if trellis_bits != 3: raise ValueError( - "prepared QSRT pairs require the trellis_bits=3 base " - "specialization" + "prepared QSRT pairs require the trellis_bits=3 base specialization" ) if dynamic_pairs: for name, modes in ( @@ -11678,7 +12684,7 @@ def run_w4a16_moe( ) if activation_amax is not None: raise NotImplementedError( - "trellis_t256 activation-amax collection is not exposed through " + "trellis3_t256 activation-amax collection is not exposed through " "the registered launch ABI" ) if coupled_hadamard and not full_rotation: @@ -11713,15 +12719,15 @@ def run_w4a16_moe( if weight_layout == "modelopt": if w13_layout not in _MODEL_OPT_W13_LAYOUTS: raise ValueError(f"unsupported W4A16 w13_layout {w13_layout!r}") - elif weight_layout == "trellis_t256": + elif weight_layout == "trellis3_t256": if w13_layout not in _TRELLIS256_W13_LAYOUTS: - raise ValueError(f"unsupported trellis_t256 w13_layout {w13_layout!r}") + raise ValueError(f"unsupported trellis3_t256 w13_layout {w13_layout!r}") else: w13_layout = "packed" dual_a_required = bool( intermediate_rotation_scales is not None - and weight_layout == "trellis_t256" - and w13_layout == "trellis_t256_proj" + and weight_layout == "trellis3_t256" + and w13_layout == "trellis3_t256_proj" ) if full_rotation and not dual_a_required: raise ValueError( @@ -11729,11 +12735,11 @@ def run_w4a16_moe( ) if dual_a_required and a_input_up is None and not full_rotation: raise ValueError( - "exact projection-major trellis_t256 rotation requires a_input_up" + "exact projection-major trellis3_t256 rotation requires a_input_up" ) if a_input_up is not None and (not dual_a_required or full_rotation): raise ValueError( - "a_input_up is only valid for exact projection-major trellis_t256 rotation" + "a_input_up is only valid for exact projection-major trellis3_t256 rotation" ) if a_input_up is not None: if ( @@ -11975,10 +12981,15 @@ def run_w4a16_moe( # scheduling/epilogue changes. A global->local expert map is resolved by # the same direct-route FC1/FC2 emit hook, so compact hybrid tiers do not # need to materialize remapped ids or masked router weights first. - # A preplanned launch built with the TC-decode fused-sum epilogue carries - # ``tc_decode_fused_sum``; accept it through the binding path. A runtime - # ``fused_launch is None`` (e.g. the standalone benchmark) compiles its own. - preplanned_tc_decode = bool(getattr(fused_launch, "tc_decode_fused_sum", False)) + # A frozen workspace carries the route-reduction contract in its compiled + # launch metadata. A standalone call without a preplanned launch resolves + # the specialization from the runtime inputs and the explicit feature flag. + preplanned_tc_decode = bool( + getattr(fused_launch, "tc_decode_fused_sum", False) + ) + preplanned_prefill_fused_sum = bool( + getattr(fused_launch, "prefill_fused_sum_fp32", False) + ) use_tc_decode = bool( (not collect_activation_amax) and (fused_launch is None or preplanned_tc_decode) @@ -11998,7 +13009,7 @@ def run_w4a16_moe( _W4A16_SMALL_M_DIRECT_MAX_M if mapped_direct else _MAX_DIRECT_TOPK_ROUTE_M ) direct_layout_ok = weight_layout == "packed" or ( - mapped_direct and full_rotation and weight_layout == "trellis_t256" + mapped_direct and full_rotation and weight_layout == "trellis3_t256" ) direct_topk_eligible = ( (not collect_activation_amax) @@ -12030,6 +13041,36 @@ def run_w4a16_moe( # TC-decode requires the inline direct-topk route path (no route-pack). use_tc_decode = bool(use_tc_decode and use_direct_topk_routes) + prefill_fused_sum_requested = ( + preplanned_prefill_fused_sum + if fused_launch is not None + else prefill_sum_accum is not None + ) + use_prefill_fused_sum = prefill_fused_sum_eligible( + dtype=element_dtype, + m=m, + full_rotation=full_rotation, + weight_layout=weight_layout, + collect_activation_amax=collect_activation_amax, + enabled=prefill_fused_sum_requested, + ) + use_fused_topk_sum = bool(use_tc_decode or use_prefill_fused_sum) + + if use_prefill_fused_sum: + required_accum_elements = int(m) * hidden_size + if ( + prefill_sum_accum is None + or prefill_sum_accum.dtype != torch.float32 + or prefill_sum_accum.device != a_input.device + or not prefill_sum_accum.is_contiguous() + or prefill_sum_accum.numel() < required_accum_elements + ): + raise ValueError( + "W4A16 prefill fused sum requires a contiguous FP32 accumulator " + f"with at least {required_accum_elements} elements on " + f"{a_input.device}" + ) + # A preplanned TC-decode launch atomically accumulates FC2 partials into the # (pre-zeroed) output and emits no separate top-k sum. If it was selected but # the decode preconditions don't hold, running it would corrupt the output, @@ -12039,6 +13080,11 @@ def run_w4a16_moe( "preplanned TC-decode W4A16 launch requires small-M packed bf16 " f"decode (m <= {_TC_DECODE_MAX_M}, cuda int32/int64 topk_ids)" ) + if preplanned_prefill_fused_sum and not use_prefill_fused_sum: + raise RuntimeError( + "preplanned W4A16 prefill fused-sum launch requires the enabled " + "large-M packed or modelopt BF16 route-reduction contract" + ) route_slots_for_scratch = int(m) * int(topk) * int(block_size_m) required_m_blocks = int(m) * int(topk) if use_direct_topk_routes else 0 @@ -12107,9 +13153,7 @@ def run_w4a16_moe( prepared.w13_scale, prepared.w2_scale, expert_map=expert_map if use_direct_topk_routes else None, - w13_row_rotation=int( - getattr(prepared, "x4t_w13_row_rotation", 0) - ), + w13_row_rotation=int(getattr(prepared, "x4t_w13_row_rotation", 0)), expert_ids_unique=bool(use_direct_topk_routes and m == 1), stream=stream, ) @@ -12125,8 +13169,12 @@ def run_w4a16_moe( topk=topk, route_num_experts=route_num_experts, sms=sms, + dtype=(prepared_dtype if full_rotation else a_input.dtype), full_rotation=full_rotation, block_size_m=block_size_m, + weight_layout=weight_layout, + collect_activation_amax=collect_activation_amax, + prefill_fused_sum=use_prefill_fused_sum, ) intermediate_size = int(prepared.intermediate_size) fc1_cols = buffer_plan.fc1_cols @@ -12192,6 +13240,7 @@ def run_w4a16_moe( direct_topk_routes=use_direct_topk_routes, use_expert_map=mapped_direct and use_direct_topk_routes, tc_decode_fused_sum=use_tc_decode, + prefill_fused_sum_fp32=use_prefill_fused_sum, collect_activation_amax=collect_activation_amax, intermediate_rotation=intermediate_rotation_scales is not None, full_rotation=full_rotation, @@ -12235,6 +13284,8 @@ def run_w4a16_moe( fc2_trellis_pair_kind, bool(use_direct_topk_routes), mapped_direct and use_direct_topk_routes, + bool(use_tc_decode), + bool(use_prefill_fused_sum), bool(collect_activation_amax), block_size_m, bool(intermediate_rotation_scales is not None), @@ -12270,13 +13321,15 @@ def run_w4a16_moe( getattr( fused_launch, "trellis_codebook", - SQG_E4M3, + "sqg_xor_cheb_t12", ) ).lower(), getattr(fused_launch, "fc1_trellis_pair_kind", None), getattr(fused_launch, "fc2_trellis_pair_kind", None), bool(getattr(fused_launch, "direct_topk_routes", False)), bool(getattr(fused_launch, "use_expert_map", False)), + bool(getattr(fused_launch, "tc_decode_fused_sum", False)), + bool(getattr(fused_launch, "prefill_fused_sum_fp32", False)), bool(getattr(fused_launch, "collect_activation_amax", False)), int(fused_launch.moe_block_size), bool(getattr(fused_launch, "intermediate_rotation", False)), @@ -12296,12 +13349,22 @@ def run_w4a16_moe( fused = fused_launch capacity_m = int(fused.size_m) capacity_routed_rows = capacity_m * topk - if intermediate_cache13_flat.numel() < capacity_routed_rows * max( - fc1_cols, hidden_size - ): + required_cache13_elements = ( + capacity_routed_rows * fc1_cols + if use_prefill_fused_sum + else capacity_routed_rows * max(fc1_cols, hidden_size) + ) + if intermediate_cache13_flat.numel() < required_cache13_elements: raise ValueError( "intermediate_cache13 is smaller than the selected W4A16 launch capacity: " - f"capacity_rows={capacity_m}, topk={topk}" + f"capacity_rows={capacity_m}, topk={topk}, " + f"available_elements={intermediate_cache13_flat.numel()}, " + f"required_elements={required_cache13_elements}, " + f"fused_topk_sum={use_fused_topk_sum}, " + f"prefill_fused_sum={use_prefill_fused_sum}, " + f"collect_activation_amax={collect_activation_amax}, " + f"full_rotation={full_rotation}, weight_layout={weight_layout}, " + f"element_dtype={element_dtype}" ) if intermediate_cache2_flat.numel() < capacity_routed_rows * intermediate_size: raise ValueError( @@ -12310,14 +13373,13 @@ def run_w4a16_moe( ) fc1_out = intermediate_cache13_flat[: capacity_routed_rows * fc1_cols] activated = intermediate_cache2_flat[: capacity_routed_rows * intermediate_size] - if use_tc_decode: + if use_prefill_fused_sum: + assert prefill_sum_accum is not None + fc2_out = prefill_sum_accum[: capacity_m * hidden_size] + elif use_tc_decode: # FC2 atomically accumulates per-route partials directly into the - # per-token output, so the output is the FC2 store target and must be - # pre-zeroed. The fused tc_decode kernel now zeroes the output in its - # own prologue (before FC1, made visible by the existing post-FC1 grid - # barrier), so the separate host-side output.zero_() launch is removed - # from the decode critical path here. This drops the separate top-k-sum - # launch as well. + # per-token output. The fused kernel zeroes the output in its prologue + # before the mandatory post-FC1 grid barrier. fc2_out = output.view(-1) else: fc2_out = intermediate_cache13_flat[: capacity_routed_rows * hidden_size] @@ -12411,10 +13473,10 @@ def run_w4a16_moe( collect_activation_amax or use_tc_decode or (use_direct_topk_routes and not full_rotation) - or weight_layout != "trellis_t256" + or weight_layout != "trellis3_t256" ): raise ValueError( - "intermediate_rotation_scales requires the trellis_t256 fused path " + "intermediate_rotation_scales requires the trellis3_t256 fused path " "(no calibration / tc-decode; direct routing requires full_rotation)" ) if _intermediate_rotation: @@ -12448,8 +13510,9 @@ def run_w4a16_moe( ) elif ( _intermediate_rotation - or weight_layout == "trellis_t256" + or weight_layout == "trellis3_t256" or (mapped_direct and use_direct_topk_routes) + or use_prefill_fused_sum ): # Native t256 bypasses the registered torch op so its shape-derived # bitrate reaches compilation without widening the stable public op ABI. @@ -12556,6 +13619,7 @@ def run_w4a16_moe( fc2_tile_n=_lt_fc2tn, direct_topk_routes=bool(use_direct_topk_routes), tc_decode_fused_sum=bool(use_tc_decode), + prefill_fused_sum_fp32=bool(use_prefill_fused_sum), collect_activation_amax=False, stream_int=int(stream), expert_map=expert_map if use_direct_topk_routes else None, @@ -12581,6 +13645,10 @@ def run_w4a16_moe( int(stream), ) + if use_prefill_fused_sum: + assert prefill_sum_accum is not None + output.copy_(prefill_sum_accum[: m * hidden_size].view(m, hidden_size)) + return output if use_tc_decode: # FC2 already wrote the top-k-summed result into `output`. return output @@ -12686,18 +13754,6 @@ def build_w4a16_tier_local_map( return table.contiguous() - - - - - - - - - - - - __all__ = [ "W4A16ActivationCompileResult", "W4A16FusedMoeCompileResult", diff --git a/b12x/moe/_shared/kernels/w4a16/mixed_trellis.py b/b12x/moe/_shared/kernels/w4a16/mixed_trellis.py index 020508999..fc3c30cf2 100644 --- a/b12x/moe/_shared/kernels/w4a16/mixed_trellis.py +++ b/b12x/moe/_shared/kernels/w4a16/mixed_trellis.py @@ -1,11 +1,9 @@ -"""One-launch mixed-bitrate ``trellis_t256`` MoE execution. +"""One-launch mixed-bitrate EXL3 Trellis MoE execution. The route packer assigns every global expert to one combined expert namespace. Input/intermediate rotations therefore run once. Per-tile dispatch resolves the -combined expert to a bitrate-specialized K3 or K4 decoder while preserving the -single cooperative FC1/activation/FC2 grid used by homogeneous trellis -execution. The decoder codebook (MCG or SQG-XOR-Cheb-T12) is a compile-time -parameter shared with the fused W4A16 kernel ABI. +combined expert to a bitrate-specialized decoder while preserving the single +cooperative FC1/activation/FC2 grid used by homogeneous Trellis. The module stays internal because checkpoint interpretation and runtime planning belong to the serving framework; B12X owns only the prepared kernel path. @@ -26,17 +24,20 @@ from b12x._lib.compiler import KernelCompileSpec, compile as b12x_compile from b12x._lib.intrinsics import get_ptr_as_int64, shared_ptr_to_u32 +from b12x._lib.quant.sqg_e4m3 import sqg_xor_cheb_t12_lut from b12x._lib.runtime_control import raise_if_kernel_resolution_frozen from b12x._lib.utils import current_cuda_stream, make_ptr -from .host import max_packed_route_slots, packed_gemm_scratch_elements +from .host import ( + max_packed_route_slots, + packed_gemm_scratch_elements, + route_pack_warmup_token_counts, +) from .kernel import ( - _SQG_XOR_CHEB_T12_LUT_ENTRIES, - _SQG_XOR_CHEB_T12_SMEM_REGION_BYTES, W4A16FusedMoeKernel, _cutlass_element_dtype, _fake_m_for_specialization, - _trellis256_execution_lut, + _query_w4a16_kernel_resources, compile_w4a16_topk_sum, pack_topk_routes_by_expert, ) @@ -46,7 +47,6 @@ class MixedTrellisCompileResult: compiled: object topk_sum: object - trellis_lut: torch.Tensor size_m: int hidden_size: int intermediate_size: int @@ -63,16 +63,27 @@ class MixedTrellisCompileResult: moe_block_size: int fc2_moe_block_size: int fc2_schedule_route_block_factor: int + fc2_paired_m8_routes: bool max_m_blocks: int blocks_per_sm: int sms: int shared_memory_bytes: int + registers_per_thread: int + local_memory_bytes: int rotation_input_dtype: str route_ids_dtype: torch.dtype broadcast_suh: bool broadcast_svh: bool +@dataclass(frozen=True) +class MixedTrellis3CompileResult(MixedTrellisCompileResult): + """Launch metadata for the K3/K4/K5 three-tier specialization.""" + + tier2_num_experts: int + tier2_bits: int + + @dataclass(frozen=True) class MixedTrellisRotations: intermediate: torch.Tensor @@ -103,6 +114,7 @@ class MixedTrellisTier(Protocol): """Prepared Trellis tier fields consumed by the mixed launch.""" num_experts: int + trellis_codebook: str intermediate_rotations: torch.Tensor gate_suh: torch.Tensor up_suh: torch.Tensor @@ -118,7 +130,7 @@ class MixedTrellisTier(Protocol): class W4A16MixedTrellisKernel: """One cooperative grid over two native Trellis bitrates.""" - ABI_VERSION = 8 + ABI_VERSION = 12 def __init__( self, @@ -132,7 +144,7 @@ def __init__( raise ValueError(f"mixed Trellis {name} requires full rotation") if moe.direct_topk_routes or moe.tc_decode_fused_sum: raise ValueError(f"mixed Trellis {name} requires route packing") - if moe.weight_layout != "trellis_t256": + if moe.weight_layout != "trellis3_t256": raise ValueError(f"mixed Trellis {name} requires native t256 weights") if moe.element_dtype != "fp16": raise ValueError(f"mixed Trellis {name} requires fp16 GEMM operands") @@ -146,7 +158,6 @@ def __init__( "activation", "rotation_input_dtype", "broadcast_suh", - "trellis_codebook", "cta_threads", "sms", ): @@ -164,6 +175,7 @@ def __init__( g.cta_threads, g.moe_block_size, g.schedule_route_block_factor, + g.paired_m8_routes, ) for g in gemms ) @@ -178,6 +190,13 @@ def __init__( "mixed Trellis FC2 schedule factor must divide one packed " f"route block: factor={fc2_factor}, maximum={expected_factor}" ) + expected_pair = fc2_factor == 2 and driver.fc2.moe_block_size == 8 + if bool(driver.fc2.paired_m8_routes) != expected_pair: + raise ValueError( + "mixed Trellis FC2 pair contract mismatch: " + f"factor={fc2_factor}, m={driver.fc2.moe_block_size}, " + f"paired={driver.fc2.paired_m8_routes}" + ) if tier0.num_experts > 256 or tier1.num_experts > 256: raise ValueError("tier-local expert ids must fit in eight bits") if driver.num_experts != tier0.num_experts + tier1.num_experts: @@ -197,6 +216,14 @@ def __init__( self.shared_words = max( driver.shared_words, tier0.shared_words, tier1.shared_words ) + # Each bitrate has a different GEMM scratch footprint. The shared T12 + # table must follow the largest pre-LUT region, otherwise a wider tier + # can overwrite a table placed at the driver's (K3) offset. + self.sqg_xor_cheb_t12_smem_off = max( + driver.sqg_xor_cheb_t12_smem_off, + tier0.sqg_xor_cheb_t12_smem_off, + tier1.sqg_xor_cheb_t12_smem_off, + ) @property def __cache_key__(self) -> tuple[object, ...]: @@ -242,8 +269,12 @@ def _dispatch_tier_gemm( factor = gemm.schedule_route_block_factor first_route_block = route_block_idx * Int32(factor) first_lock_slot = lock_slot * Int32(factor) - for subtile in cutlass.range_constexpr(factor): - gemm._run_tile( + # MCG does not consume the generic Trellis LUT ABI slot. SQG must keep + # the caller-provided address so every bitrate can decode through T12. + if cutlass.const_expr(self.driver.trellis_codebook == "mcg"): + trellis_lut_addr = cutlass.Int64(0) + if cutlass.const_expr(gemm.paired_m8_routes): + gemm._run_tile_m8_pair( a_flat, a_alt_flat, b_flat, @@ -257,16 +288,42 @@ def _dispatch_tier_gemm( trellis_lut_addr, smem_base, tid, - first_route_block + Int32(subtile), + first_route_block, local_expert, output_n_tile, reduce_k_tile, reduce_tile_count, reduce_slice_count, reduce_slice_idx, - first_lock_slot + Int32(subtile), + first_lock_slot, active_size_m, ) + else: + for subtile in cutlass.range_constexpr(factor): + gemm._run_tile( + a_flat, + a_alt_flat, + b_flat, + c_flat, + scales_flat, + global_scale, + packed_route_indices, + topk_weights, + c_tmp, + locks, + trellis_lut_addr, + smem_base, + tid, + first_route_block + Int32(subtile), + local_expert, + output_n_tile, + reduce_k_tile, + reduce_tile_count, + reduce_slice_count, + reduce_slice_idx, + first_lock_slot + Int32(subtile), + active_size_m, + ) @cute.jit def _emit_tier_tile( @@ -293,6 +350,12 @@ def _emit_tier_tile( active_size_m: Int32, tier0_num_experts: Int32, tier1_num_experts: Int32, + tier0_fc2_experts: Int32, + tier1_fc2_experts: Int32, + tier0_gate_experts: Int32, + tier1_gate_experts: Int32, + tier0_up_experts: Int32, + tier1_up_experts: Int32, route_block_idx: Int32, output_n_tile: Int32, reduce_k_tile: Int32, @@ -312,12 +375,33 @@ def _emit_tier_tile( ) combined_expert = block_expert_ids[metadata_block_idx].to(Int32) total_experts = tier0_num_experts + tier1_num_experts + # glm52-r7-projtiers: gate and up may sit in different tiers, so the + # descriptor row is chosen per projection. FC2 resolves at compile time; + # FC1 splits on the N half, which trellis3_t256_proj keeps aligned to + # whole CTA N tiles. + descriptor_row = Int32(2) + if cutlass.const_expr(is_fc1): + fc1_half_tiles = Int32(self.driver.fc1.n_tiles // 2) + descriptor_row = Int32(0) + if output_n_tile >= fc1_half_tiles: + descriptor_row = Int32(1) if combined_expert >= Int32(0) and combined_expert < total_experts: - descriptor = descriptor_map[combined_expert].to(Int32) + descriptor = descriptor_map[ + descriptor_row * total_experts + combined_expert + ].to(Int32) if descriptor >= Int32(0): tier = descriptor >> Int32(8) local_expert = descriptor & Int32(0xFF) - if tier == Int32(0) and local_expert < tier0_num_experts: + # FC1 is bounded by the tier's FC1 slot count; FC2 by its own + # independent count, since per-projection membership lets the + # two differ. Both remain real bounds. + if cutlass.const_expr(is_fc1): + tier0_in_bounds = local_expert < tier0_gate_experts + if output_n_tile >= fc1_half_tiles: + tier0_in_bounds = local_expert < tier0_up_experts + else: + tier0_in_bounds = local_expert < tier0_fc2_experts + if tier == Int32(0) and tier0_in_bounds: if cutlass.const_expr(is_fc1): gemm = self.tier0.fc1 else: @@ -347,7 +431,13 @@ def _emit_tier_tile( lock_slot, active_size_m, ) - elif tier == Int32(1) and local_expert < tier1_num_experts: + if cutlass.const_expr(is_fc1): + tier1_in_bounds = local_expert < tier1_gate_experts + if output_n_tile >= fc1_half_tiles: + tier1_in_bounds = local_expert < tier1_up_experts + else: + tier1_in_bounds = local_expert < tier1_fc2_experts + if tier == Int32(1) and tier1_in_bounds: if cutlass.const_expr(is_fc1): gemm = self.tier1.fc1 else: @@ -413,11 +503,30 @@ def __call__( trellis_lut_ptr: cute.Pointer, tier0_num_experts: cutlass.Int32, tier1_num_experts: cutlass.Int32, + tier0_fc2_experts: cutlass.Int32, + tier1_fc2_experts: cutlass.Int32, active_m: cutlass.Int32, grid_x: cutlass.Int32, stream: cuda.CUstream, + # Appended LAST so every existing positional call keeps its slots; + # passed by keyword. No default: the CuTe DSL types every parameter + # when it traces, and a None default is untypeable. + tier0_gate_experts: cutlass.Int32, + tier1_gate_experts: cutlass.Int32, + tier0_up_experts: cutlass.Int32, + tier1_up_experts: cutlass.Int32, ): tier0_experts = cutlass.Int64(tier0_num_experts) + # FC2 extents are independent of the FC1 slot counts. + tier0_fc2 = cutlass.Int64(tier0_fc2_experts) + tier1_fc2 = cutlass.Int64(tier1_fc2_experts) + # The w13 descriptor is sized by the GATE count so the gemm's + # cute.size(w13)//2 up-block base lands at gate_count*proj_stride over + # a tight [gate|up] buffer. run_mixed_trellis defaults these to the + # tier expert counts when a caller supplies none, which reproduces the + # historical padded sizing exactly. + tier0_gate = cutlass.Int64(tier0_gate_experts) + tier1_gate = cutlass.Int64(tier1_gate_experts) tier1_experts = cutlass.Int64(tier1_num_experts) total_experts = tier0_experts + tier1_experts @@ -425,7 +534,7 @@ def __call__( t0_w13_ptr, layout=cute.make_layout( ( - tier0_experts + tier0_gate * cutlass.Int64(self.hidden_size // 16) * cutlass.Int64(self.driver.fc1_cols // 16) * cutlass.Int64(8 * self.tier0.trellis_bits), @@ -437,7 +546,7 @@ def __call__( t0_w2_ptr, layout=cute.make_layout( ( - tier0_experts + tier0_fc2 * cutlass.Int64(self.intermediate_size // 16) * cutlass.Int64(self.hidden_size // 16) * cutlass.Int64(8 * self.tier0.trellis_bits), @@ -449,7 +558,7 @@ def __call__( t1_w13_ptr, layout=cute.make_layout( ( - tier1_experts + tier1_gate * cutlass.Int64(self.hidden_size // 16) * cutlass.Int64(self.driver.fc1_cols // 16) * cutlass.Int64(8 * self.tier1.trellis_bits), @@ -461,7 +570,7 @@ def __call__( t1_w2_ptr, layout=cute.make_layout( ( - tier1_experts + tier1_fc2 * cutlass.Int64(self.intermediate_size // 16) * cutlass.Int64(self.hidden_size // 16) * cutlass.Int64(8 * self.tier1.trellis_bits), @@ -519,7 +628,7 @@ def __call__( ) t0_w2_global = cute.make_tensor( t0_w2_global_ptr, - layout=cute.make_layout((tier0_experts,), stride=(1,)), + layout=cute.make_layout((tier0_fc2,), stride=(1,)), ) t1_w13_global = cute.make_tensor( t1_w13_global_ptr, @@ -527,11 +636,13 @@ def __call__( ) t1_w2_global = cute.make_tensor( t1_w2_global_ptr, - layout=cute.make_layout((tier1_experts,), stride=(1,)), + layout=cute.make_layout((tier1_fc2,), stride=(1,)), ) + # glm52-r7-projtiers: rows are gate, up, down. Row 0 alone is the + # historical layout, so three identical rows reproduce it exactly. descriptor_map = cute.make_tensor( descriptor_map_ptr, - layout=cute.make_layout((total_experts,), stride=(1,)), + layout=cute.make_layout((cutlass.Int64(3) * total_experts,), stride=(1,)), ) intermediate_rotations = cute.make_tensor( intermediate_rotations_ptr, @@ -555,6 +666,11 @@ def __call__( (suh_rows * cutlass.Int64(self.hidden_size),), stride=(1,) ), ) + trellis_lut = cute.make_tensor( + trellis_lut_ptr, + layout=cute.make_layout((Int64(1 << 12),), stride=(1,)), + ) + trellis_lut_addr = get_ptr_as_int64(trellis_lut, Int32(0)) rotation_input = cute.make_tensor( rotation_input_ptr, layout=cute.make_layout( @@ -569,12 +685,6 @@ def __call__( stride=(1,), ), ) - trellis_lut = cute.make_tensor( - trellis_lut_ptr, - layout=cute.make_layout( - (cutlass.Int64(_SQG_XOR_CHEB_T12_LUT_ENTRIES),), stride=(1,) - ), - ) self.kernel( rotation_input, rotation_gate, @@ -605,9 +715,15 @@ def __call__( intermediate_rotations, gate_suh, up_suh, - trellis_lut, + trellis_lut_addr, tier0_num_experts, tier1_num_experts, + tier0_fc2_experts, + tier1_fc2_experts, + tier0_gate_experts, + tier1_gate_experts, + tier0_up_experts, + tier1_up_experts, active_m, ).launch( grid=(grid_x, 1, 1), @@ -649,9 +765,15 @@ def kernel( intermediate_rotations: cute.Tensor, gate_suh: cute.Tensor, up_suh: cute.Tensor, - trellis_lut: cute.Tensor, + trellis_lut_addr: Int64, tier0_num_experts: cutlass.Int32, tier1_num_experts: cutlass.Int32, + tier0_fc2_experts: cutlass.Int32, + tier1_fc2_experts: cutlass.Int32, + tier0_gate_experts: cutlass.Int32, + tier1_gate_experts: cutlass.Int32, + tier0_up_experts: cutlass.Int32, + tier1_up_experts: cutlass.Int32, active_m: cutlass.Int32, ): tidx, _, _ = cute.arch.thread_idx() @@ -671,19 +793,16 @@ class Storage: storage = smem.allocate(Storage) smem_base = shared_ptr_to_u32(storage.words.data_ptr()) - trellis_lut_addr = get_ptr_as_int64(trellis_lut, Int32(0)) - phase_lut_addr = trellis_lut_addr + decode_lut_addr = trellis_lut_addr if cutlass.const_expr(self.driver.sqg_xor_cheb_t12_smem): self.driver._sqg_smem_copy( trellis_lut_addr, - smem_base + Int32(self.driver.sqg_xor_cheb_t12_smem_off), - _SQG_XOR_CHEB_T12_SMEM_REGION_BYTES, + smem_base + Int32(self.sqg_xor_cheb_t12_smem_off), + 1 << 12, tid, ) cute.arch.sync_threads() - phase_lut_addr = Int64( - smem_base + Int32(self.driver.sqg_xor_cheb_t12_smem_off) - ) + decode_lut_addr = Int64(smem_base + Int32(self.sqg_xor_cheb_t12_smem_off)) fc1_emit = partial( self._emit_tier_tile, True, @@ -702,12 +821,18 @@ class Storage: topk_weights, fc1_scratch, workspace, - phase_lut_addr, + decode_lut_addr, smem_base, tid, active_m, tier0_num_experts, tier1_num_experts, + tier0_fc2_experts, + tier1_fc2_experts, + tier0_gate_experts, + tier1_gate_experts, + tier0_up_experts, + tier1_up_experts, ) fc2_emit = partial( self._emit_tier_tile, @@ -727,12 +852,18 @@ class Storage: topk_weights, fc2_scratch, workspace, - phase_lut_addr, + decode_lut_addr, smem_base, tid, active_m * Int32(self.top_k), tier0_num_experts, tier1_num_experts, + tier0_fc2_experts, + tier1_fc2_experts, + tier0_gate_experts, + tier1_gate_experts, + tier0_up_experts, + tier1_up_experts, ) total_experts = tier0_num_experts + tier1_num_experts self.driver._moe_body( @@ -761,8 +892,10 @@ class Storage: gate_suh, up_suh, descriptor_map, - trellis_lut_addr, - trellis_lut_addr, + # Mixed tier emit hooks own Trellis decoding, so the shared + # driver's LUT ABI slots are intentionally unused. + cutlass.Int64(0), + cutlass.Int64(0), total_experts, total_experts, smem_base, @@ -775,180 +908,1259 @@ class Storage: ) -_CACHE: dict[tuple[object, ...], MixedTrellisCompileResult] = {} - - -def compile_mixed_trellis( - *, - size_m: int, - hidden_size: int, - intermediate_size: int, - tier0_num_experts: int, - tier1_num_experts: int, - top_k: int, - max_m_blocks: int, - sms: int, - max_shared_mem: int, - force_tile_config: tuple[int, int, int, int], - tier0_bits: int = 3, - tier1_bits: int = 4, - trellis_codebook: str = "mcg", - moe_block_size: int = 8, - rotation_input_dtype: str = "bf16", - route_ids_dtype: torch.dtype = torch.int32, - broadcast_suh: bool = False, - broadcast_svh: bool = False, -) -> MixedTrellisCompileResult: - if route_ids_dtype not in (torch.int32, torch.int64): - raise TypeError("mixed Trellis route IDs must be int32 or int64") - if int(size_m) * int(top_k) > torch.iinfo(torch.int32).max: - raise ValueError("mixed Trellis routed-row count must fit in int32") - fc1_tile_k, fc1_tile_n, fc2_tile_k, fc2_tile_n = ( - int(value) for value in force_tile_config - ) - trellis_codebook = str(trellis_codebook).lower() - if fc1_tile_k < 128: - raise ValueError( - "mixed Trellis FC1 requires tile_k >= 128; narrower K tiles lose " - "large-M cross-tier partial reductions" - ) - total_experts = int(tier0_num_experts) + int(tier1_num_experts) - grouped_m8_fc2 = int(moe_block_size) in (32, 64) - - def make_kernel(num_experts: int, bits: int) -> W4A16FusedMoeKernel: - return W4A16FusedMoeKernel( - size_m=size_m, - hidden_size=hidden_size, - intermediate_size=intermediate_size, - num_experts=num_experts, - top_k=top_k, - activation="silu", - apply_router_weight_on_input=False, - zero_fc2_output=False, - fc1_tile_n=fc1_tile_n, - fc1_tile_k=fc1_tile_k, - fc2_tile_n=fc2_tile_n, - fc2_tile_k=fc2_tile_k, - moe_block_size=moe_block_size, - max_m_blocks=max_m_blocks, - fc2_moe_block_size=(8 if grouped_m8_fc2 else moe_block_size), - fc2_schedule_route_block_factor=(2 if grouped_m8_fc2 else 1), - element_dtype="fp16", - weight_layout="trellis_t256", - scale_format="e4m3_k32", - w13_layout="trellis_t256_proj", - trellis_bits=bits, - trellis_codebook=trellis_codebook, - intermediate_rotation=True, - full_rotation=True, - rotation_input_dtype=rotation_input_dtype, - broadcast_suh=broadcast_suh, - schedule_whole_tiles=True, - ) +class W4A16MixedTrellis3Kernel(W4A16MixedTrellisKernel): + """One cooperative grid over three native Trellis bitrates. - kernel = W4A16MixedTrellisKernel( - driver=make_kernel(total_experts, tier0_bits), - tier0=make_kernel(int(tier0_num_experts), int(tier0_bits)), - tier1=make_kernel(int(tier1_num_experts), int(tier1_bits)), - ) - # shared_words is the complete dynamically allocated MemRange used by the - # cooperative kernel. CUDA permits a launch exactly at the device's - # opt-in shared-memory limit; rejecting an additional 512 bytes here - # unnecessarily excludes the stock mixed-K tile geometry at block-64. - if kernel.shared_words * 4 > int(max_shared_mem): - raise ValueError( - "mixed Trellis shared-memory requirement exceeds the device limit: " - f"required={kernel.shared_words * 4} " - f"limit={int(max_shared_mem)}" - ) - device = int(torch.cuda.current_device()) - cache_key = ( - "mixed_trellis", - device, - kernel.__cache_key__, - str(route_ids_dtype), - int(size_m), - int(max_m_blocks), - ) - topk_sum = compile_w4a16_topk_sum( - m=size_m, - topk=top_k, - hidden_size=hidden_size, - element_dtype="fp16", - full_rotation=True, - num_experts=total_experts, - route_num_experts=total_experts, - route_ids_dtype=route_ids_dtype, - use_expert_map=True, - broadcast_svh=broadcast_svh, - ) - cached = _CACHE.get(cache_key) - if cached is not None: - # The compiled object is intentionally independent of the artifact's - # K3/K4 partition. Keep the current plan metadata and top-k launch, - # rather than leaking the first split that populated the cache. - return replace( - cached, - topk_sum=topk_sum, - tier0_num_experts=int(tier0_num_experts), - tier1_num_experts=int(tier1_num_experts), - sms=int(sms), - broadcast_suh=bool(broadcast_suh), - broadcast_svh=bool(broadcast_svh), - ) + This is a separate specialization so the established two-tier K3/K4 kernel + keeps its existing ABI and generated code. The GLM R7 checkpoint uses this + path only for layers that actually contain K5 payloads. + """ - compile_m = _fake_m_for_specialization(size_m) - compile_rows = compile_m * top_k - fc1_cols = 2 * intermediate_size - cutlass_dtype = cutlass.Float16 - rotation_dtype = _cutlass_element_dtype(rotation_input_dtype) + ABI_VERSION = 4 - def tensor(dtype, elements: int, *, align: int = 16): - return cute.runtime.make_fake_compact_tensor( - dtype, (max(int(elements), 1),), assumed_align=align + def __init__( + self, + *, + driver: W4A16FusedMoeKernel, + tier0: W4A16FusedMoeKernel, + tier1: W4A16FusedMoeKernel, + tier2: W4A16FusedMoeKernel, + ): + kernels = (driver, tier0, tier1, tier2) + for name, moe in zip( + ("driver", "tier0", "tier1", "tier2"), kernels, strict=True + ): + if not moe.full_rotation or not moe.intermediate_rotation: + raise ValueError(f"mixed Trellis {name} requires full rotation") + if moe.direct_topk_routes or moe.tc_decode_fused_sum: + raise ValueError(f"mixed Trellis {name} requires route packing") + if moe.weight_layout != "trellis3_t256": + raise ValueError(f"mixed Trellis {name} requires native t256 weights") + if moe.element_dtype != "fp16": + raise ValueError(f"mixed Trellis {name} requires fp16 GEMM operands") + for attr in ( + "size_m", + "hidden_size", + "intermediate_size", + "fc1_cols", + "top_k", + "moe_block_size", + "activation", + "rotation_input_dtype", + "broadcast_suh", + "cta_threads", + "sms", + ): + values = tuple(getattr(moe, attr) for moe in kernels) + if values[1:] != values[:-1]: + raise ValueError(f"mixed Trellis kernels disagree on {attr}: {values}") + for phase in ("fc1", "fc2"): + gemms = tuple(getattr(moe, phase) for moe in kernels) + geometry = tuple( + ( + gemm.n_tiles, + gemm.k_tiles, + gemm.tile_n, + gemm.tile_k, + gemm.cta_threads, + gemm.moe_block_size, + gemm.schedule_route_block_factor, + gemm.paired_m8_routes, + ) + for gemm in gemms + ) + if geometry[1:] != geometry[:-1]: + raise ValueError( + f"mixed Trellis kernels disagree on {phase} geometry: {geometry}" + ) + fc2_factor = int(driver.fc2.schedule_route_block_factor) + expected_factor = int(driver.moe_block_size // driver.fc2.moe_block_size) + if fc2_factor < 1 or expected_factor % fc2_factor != 0: + raise ValueError( + "mixed Trellis FC2 schedule factor must divide one packed " + f"route block: factor={fc2_factor}, maximum={expected_factor}" + ) + expected_pair = fc2_factor == 2 and driver.fc2.moe_block_size == 8 + if bool(driver.fc2.paired_m8_routes) != expected_pair: + raise ValueError( + "mixed Trellis FC2 pair contract mismatch: " + f"factor={fc2_factor}, m={driver.fc2.moe_block_size}, " + f"paired={driver.fc2.paired_m8_routes}" + ) + tiers = (tier0, tier1, tier2) + if any(tier.num_experts > _MAX_TIER_EXPERTS for tier in tiers): + raise ValueError("tier-local expert ids must fit in eight bits") + if driver.num_experts != sum(tier.num_experts for tier in tiers): + raise ValueError("driver expert count must equal the sum of all tiers") + self.driver = driver + self.tier0 = tier0 + self.tier1 = tier1 + self.tier2 = tier2 + self.size_m = driver.size_m + self.hidden_size = driver.hidden_size + self.intermediate_size = driver.intermediate_size + self.top_k = driver.top_k + self.cta_threads = driver.cta_threads + self.sms = driver.sms + self.blocks_per_sm = min(tier.blocks_per_sm for tier in kernels) + self.shared_words = max(tier.shared_words for tier in kernels) + self.sqg_xor_cheb_t12_smem_off = max( + tier.sqg_xor_cheb_t12_smem_off for tier in kernels ) - def tier_args(): + @property + def __cache_key__(self) -> tuple[object, ...]: return ( - make_ptr(cutlass.Int32, 16, cute.AddressSpace.gmem, assumed_align=16), - make_ptr(cutlass.Int32, 16, cute.AddressSpace.gmem, assumed_align=16), - make_ptr(cutlass.Int32, 16, cute.AddressSpace.gmem, assumed_align=16), - make_ptr(cutlass.Int32, 16, cute.AddressSpace.gmem, assumed_align=16), - make_ptr(cutlass.Float32, 16, cute.AddressSpace.gmem, assumed_align=16), - make_ptr(cutlass.Float32, 16, cute.AddressSpace.gmem, assumed_align=16), + "w4a16_mixed_trellis3", + self.ABI_VERSION, + self.driver.__cache_key__, + self.tier0.__cache_key__, + self.tier1.__cache_key__, + self.tier2.__cache_key__, + self.blocks_per_sm, + self.shared_words, ) - scratch_elements = max( - fc1_cols * compile_rows, - hidden_size * compile_rows, - 4 * 256 * moe_block_size * 256, - ) - compile_args = ( - make_ptr(rotation_dtype, 16, cute.AddressSpace.gmem, assumed_align=16), - tensor(cutlass_dtype, compile_rows * hidden_size), - tensor(cutlass_dtype, compile_rows * hidden_size), - *tier_args(), - *tier_args(), - tensor(cutlass_dtype, compile_rows * fc1_cols), - tensor(cutlass_dtype, compile_rows * intermediate_size), - tensor(cutlass_dtype, compile_rows * hidden_size), - tensor(cutlass.Int32, moe_block_size), - tensor(cutlass.Int32, 1), - tensor(cutlass.Int32, 1, align=4), - make_ptr(cutlass.Int32, 4, cute.AddressSpace.gmem, assumed_align=4), - make_ptr(cutlass.Float32, 4, cute.AddressSpace.gmem, assumed_align=4), - tensor(cutlass.Float32, scratch_elements), - tensor(cutlass.Float32, scratch_elements), - tensor(cutlass.Int32, 4 * 256 + 2), - make_ptr(cutlass.Float16, 16, cute.AddressSpace.gmem, assumed_align=16), - make_ptr(cutlass.Float16, 16, cute.AddressSpace.gmem, assumed_align=16), - make_ptr(cutlass.Float16, 16, cute.AddressSpace.gmem, assumed_align=16), - make_ptr(cutlass.Uint8, 16, cute.AddressSpace.gmem, assumed_align=16), - Int32(tier0_num_experts), - Int32(tier1_num_experts), - 1, - 1, + @cute.jit + def _emit_tier_tile3( + self, + is_fc1: cutlass.Constexpr, + a_flat: cute.Tensor, + a_alt_flat: cute.Tensor, + t0_b_flat: cute.Tensor, + t0_scales_flat: cute.Tensor, + t0_global_scale: cute.Tensor, + t1_b_flat: cute.Tensor, + t1_scales_flat: cute.Tensor, + t1_global_scale: cute.Tensor, + t2_b_flat: cute.Tensor, + t2_scales_flat: cute.Tensor, + t2_global_scale: cute.Tensor, + c_flat: cute.Tensor, + packed_route_indices: cute.Tensor, + block_expert_ids: cute.Tensor, + descriptor_map: cute.Tensor, + topk_weights: cute.Tensor, + c_tmp: cute.Tensor, + locks: cute.Tensor, + trellis_lut_addr: Int64, + smem_base: Int32, + tid: Int32, + active_size_m: Int32, + tier0_num_experts: Int32, + tier1_num_experts: Int32, + tier2_num_experts: Int32, + tier0_fc2_experts: Int32, + tier1_fc2_experts: Int32, + tier2_fc2_experts: Int32, + tier0_gate_experts: Int32, + tier1_gate_experts: Int32, + tier2_gate_experts: Int32, + tier0_up_experts: Int32, + tier1_up_experts: Int32, + tier2_up_experts: Int32, + route_block_idx: Int32, + output_n_tile: Int32, + reduce_k_tile: Int32, + reduce_tile_count: Int32, + reduce_slice_count: Int32, + reduce_slice_idx: Int32, + lock_slot: Int32, + ): + metadata_block_idx = route_block_idx + if cutlass.const_expr(not is_fc1): + metadata_block_idx = route_block_idx // Int32( + self.driver.moe_block_size + // ( + self.driver.fc2.moe_block_size + * self.driver.fc2.schedule_route_block_factor + ) + ) + combined_expert = block_expert_ids[metadata_block_idx].to(Int32) + total_experts = tier0_num_experts + tier1_num_experts + tier2_num_experts + descriptor_row = Int32(2) + if cutlass.const_expr(is_fc1): + fc1_half_tiles = Int32(self.driver.fc1.n_tiles // 2) + descriptor_row = Int32(0) + if output_n_tile >= fc1_half_tiles: + descriptor_row = Int32(1) + if combined_expert >= Int32(0) and combined_expert < total_experts: + descriptor = descriptor_map[ + descriptor_row * total_experts + combined_expert + ].to(Int32) + if descriptor >= Int32(0): + tier = descriptor >> Int32(8) + local_expert = descriptor & Int32(0xFF) + + if cutlass.const_expr(is_fc1): + t0_in_bounds = local_expert < tier0_gate_experts + if output_n_tile >= fc1_half_tiles: + t0_in_bounds = local_expert < tier0_up_experts + else: + t0_in_bounds = local_expert < tier0_fc2_experts + if tier == Int32(0) and t0_in_bounds: + if cutlass.const_expr(is_fc1): + gemm = self.tier0.fc1 + else: + gemm = self.tier0.fc2 + self._dispatch_tier_gemm( + gemm, + a_flat, + a_alt_flat, + t0_b_flat, + c_flat, + t0_scales_flat, + t0_global_scale, + packed_route_indices, + topk_weights, + c_tmp, + locks, + trellis_lut_addr, + smem_base, + tid, + route_block_idx, + local_expert, + output_n_tile, + reduce_k_tile, + reduce_tile_count, + reduce_slice_count, + reduce_slice_idx, + lock_slot, + active_size_m, + ) + + if cutlass.const_expr(is_fc1): + t1_in_bounds = local_expert < tier1_gate_experts + if output_n_tile >= fc1_half_tiles: + t1_in_bounds = local_expert < tier1_up_experts + else: + t1_in_bounds = local_expert < tier1_fc2_experts + if tier == Int32(1) and t1_in_bounds: + if cutlass.const_expr(is_fc1): + gemm = self.tier1.fc1 + else: + gemm = self.tier1.fc2 + self._dispatch_tier_gemm( + gemm, + a_flat, + a_alt_flat, + t1_b_flat, + c_flat, + t1_scales_flat, + t1_global_scale, + packed_route_indices, + topk_weights, + c_tmp, + locks, + trellis_lut_addr, + smem_base, + tid, + route_block_idx, + local_expert, + output_n_tile, + reduce_k_tile, + reduce_tile_count, + reduce_slice_count, + reduce_slice_idx, + lock_slot, + active_size_m, + ) + + if cutlass.const_expr(is_fc1): + t2_in_bounds = local_expert < tier2_gate_experts + if output_n_tile >= fc1_half_tiles: + t2_in_bounds = local_expert < tier2_up_experts + else: + t2_in_bounds = local_expert < tier2_fc2_experts + if tier == Int32(2) and t2_in_bounds: + if cutlass.const_expr(is_fc1): + gemm = self.tier2.fc1 + else: + gemm = self.tier2.fc2 + self._dispatch_tier_gemm( + gemm, + a_flat, + a_alt_flat, + t2_b_flat, + c_flat, + t2_scales_flat, + t2_global_scale, + packed_route_indices, + topk_weights, + c_tmp, + locks, + trellis_lut_addr, + smem_base, + tid, + route_block_idx, + local_expert, + output_n_tile, + reduce_k_tile, + reduce_tile_count, + reduce_slice_count, + reduce_slice_idx, + lock_slot, + active_size_m, + ) + + @cute.jit + def __call__( + self, + rotation_input_ptr: cute.Pointer, + rotation_gate: cute.Tensor, + rotation_up: cute.Tensor, + t0_w13_ptr: cute.Pointer, + t0_w2_ptr: cute.Pointer, + t0_w13_scales_ptr: cute.Pointer, + t0_w2_scales_ptr: cute.Pointer, + t0_w13_global_ptr: cute.Pointer, + t0_w2_global_ptr: cute.Pointer, + t1_w13_ptr: cute.Pointer, + t1_w2_ptr: cute.Pointer, + t1_w13_scales_ptr: cute.Pointer, + t1_w2_scales_ptr: cute.Pointer, + t1_w13_global_ptr: cute.Pointer, + t1_w2_global_ptr: cute.Pointer, + t2_w13_ptr: cute.Pointer, + t2_w2_ptr: cute.Pointer, + t2_w13_scales_ptr: cute.Pointer, + t2_w2_scales_ptr: cute.Pointer, + t2_w13_global_ptr: cute.Pointer, + t2_w2_global_ptr: cute.Pointer, + fc1: cute.Tensor, + activated: cute.Tensor, + fc2: cute.Tensor, + packed_route_indices: cute.Tensor, + block_expert_ids: cute.Tensor, + packed_route_count: cute.Tensor, + descriptor_map_ptr: cute.Pointer, + topk_weights_ptr: cute.Pointer, + fc1_scratch: cute.Tensor, + fc2_scratch: cute.Tensor, + workspace: cute.Tensor, + intermediate_rotations_ptr: cute.Pointer, + gate_suh_ptr: cute.Pointer, + up_suh_ptr: cute.Pointer, + trellis_lut_ptr: cute.Pointer, + tier0_num_experts: cutlass.Int32, + tier1_num_experts: cutlass.Int32, + tier2_num_experts: cutlass.Int32, + tier0_fc2_experts: cutlass.Int32, + tier1_fc2_experts: cutlass.Int32, + tier2_fc2_experts: cutlass.Int32, + active_m: cutlass.Int32, + grid_x: cutlass.Int32, + stream: cuda.CUstream, + tier0_gate_experts: cutlass.Int32, + tier1_gate_experts: cutlass.Int32, + tier2_gate_experts: cutlass.Int32, + tier0_up_experts: cutlass.Int32, + tier1_up_experts: cutlass.Int32, + tier2_up_experts: cutlass.Int32, + ): + tier0_experts = cutlass.Int64(tier0_num_experts) + tier1_experts = cutlass.Int64(tier1_num_experts) + tier2_experts = cutlass.Int64(tier2_num_experts) + tier0_fc2 = cutlass.Int64(tier0_fc2_experts) + tier1_fc2 = cutlass.Int64(tier1_fc2_experts) + tier2_fc2 = cutlass.Int64(tier2_fc2_experts) + tier0_gate = cutlass.Int64(tier0_gate_experts) + tier1_gate = cutlass.Int64(tier1_gate_experts) + tier2_gate = cutlass.Int64(tier2_gate_experts) + total_experts = tier0_experts + tier1_experts + tier2_experts + + def weight_tensor(ptr, elements): + return cute.make_tensor( + ptr, + layout=cute.make_layout((elements,), stride=(1,)), + ) + + t0_w13 = weight_tensor( + t0_w13_ptr, + tier0_gate + * cutlass.Int64(self.hidden_size // 16) + * cutlass.Int64(self.driver.fc1_cols // 16) + * cutlass.Int64(8 * self.tier0.trellis_bits), + ) + t0_w2 = weight_tensor( + t0_w2_ptr, + tier0_fc2 + * cutlass.Int64(self.intermediate_size // 16) + * cutlass.Int64(self.hidden_size // 16) + * cutlass.Int64(8 * self.tier0.trellis_bits), + ) + t1_w13 = weight_tensor( + t1_w13_ptr, + tier1_gate + * cutlass.Int64(self.hidden_size // 16) + * cutlass.Int64(self.driver.fc1_cols // 16) + * cutlass.Int64(8 * self.tier1.trellis_bits), + ) + t1_w2 = weight_tensor( + t1_w2_ptr, + tier1_fc2 + * cutlass.Int64(self.intermediate_size // 16) + * cutlass.Int64(self.hidden_size // 16) + * cutlass.Int64(8 * self.tier1.trellis_bits), + ) + t2_w13 = weight_tensor( + t2_w13_ptr, + tier2_gate + * cutlass.Int64(self.hidden_size // 16) + * cutlass.Int64(self.driver.fc1_cols // 16) + * cutlass.Int64(8 * self.tier2.trellis_bits), + ) + t2_w2 = weight_tensor( + t2_w2_ptr, + tier2_fc2 + * cutlass.Int64(self.intermediate_size // 16) + * cutlass.Int64(self.hidden_size // 16) + * cutlass.Int64(8 * self.tier2.trellis_bits), + ) + + t0_w13_scales = weight_tensor( + t0_w13_scales_ptr, + tier0_experts + * cutlass.Int64(self.tier0.fc1.scale_k_groups) + * cutlass.Int64(self.tier0.fc1.scale_size_n // 4), + ) + t0_w2_scales = weight_tensor( + t0_w2_scales_ptr, + tier0_experts + * cutlass.Int64(self.tier0.fc2.scale_k_groups) + * cutlass.Int64(self.tier0.fc2.scale_size_n // 4), + ) + t1_w13_scales = weight_tensor( + t1_w13_scales_ptr, + tier1_experts + * cutlass.Int64(self.tier1.fc1.scale_k_groups) + * cutlass.Int64(self.tier1.fc1.scale_size_n // 4), + ) + t1_w2_scales = weight_tensor( + t1_w2_scales_ptr, + tier1_experts + * cutlass.Int64(self.tier1.fc2.scale_k_groups) + * cutlass.Int64(self.tier1.fc2.scale_size_n // 4), + ) + t2_w13_scales = weight_tensor( + t2_w13_scales_ptr, + tier2_experts + * cutlass.Int64(self.tier2.fc1.scale_k_groups) + * cutlass.Int64(self.tier2.fc1.scale_size_n // 4), + ) + t2_w2_scales = weight_tensor( + t2_w2_scales_ptr, + tier2_experts + * cutlass.Int64(self.tier2.fc2.scale_k_groups) + * cutlass.Int64(self.tier2.fc2.scale_size_n // 4), + ) + t0_w13_global = weight_tensor(t0_w13_global_ptr, tier0_experts) + t0_w2_global = weight_tensor(t0_w2_global_ptr, tier0_fc2) + t1_w13_global = weight_tensor(t1_w13_global_ptr, tier1_experts) + t1_w2_global = weight_tensor(t1_w2_global_ptr, tier1_fc2) + t2_w13_global = weight_tensor(t2_w13_global_ptr, tier2_experts) + t2_w2_global = weight_tensor(t2_w2_global_ptr, tier2_fc2) + descriptor_map = weight_tensor( + descriptor_map_ptr, cutlass.Int64(3) * total_experts + ) + intermediate_rotations = weight_tensor( + intermediate_rotations_ptr, + total_experts * cutlass.Int64(3 * self.intermediate_size), + ) + suh_rows = total_experts + if cutlass.const_expr(self.driver.broadcast_suh): + suh_rows = cutlass.Int64(1) + gate_suh = weight_tensor( + gate_suh_ptr, suh_rows * cutlass.Int64(self.hidden_size) + ) + up_suh = weight_tensor(up_suh_ptr, suh_rows * cutlass.Int64(self.hidden_size)) + trellis_lut = weight_tensor(trellis_lut_ptr, Int64(1 << 12)) + trellis_lut_addr = get_ptr_as_int64(trellis_lut, Int32(0)) + rotation_input = weight_tensor( + rotation_input_ptr, + active_m.to(cutlass.Int64) * cutlass.Int64(self.hidden_size), + ) + topk_weights = weight_tensor( + topk_weights_ptr, + active_m.to(cutlass.Int64) * cutlass.Int64(self.top_k), + ) + self.kernel3( + rotation_input, + rotation_gate, + rotation_up, + t0_w13, + t0_w2, + t0_w13_scales, + t0_w2_scales, + t0_w13_global, + t0_w2_global, + t1_w13, + t1_w2, + t1_w13_scales, + t1_w2_scales, + t1_w13_global, + t1_w2_global, + t2_w13, + t2_w2, + t2_w13_scales, + t2_w2_scales, + t2_w13_global, + t2_w2_global, + fc1, + activated, + fc2, + packed_route_indices, + block_expert_ids, + packed_route_count, + descriptor_map, + topk_weights, + fc1_scratch, + fc2_scratch, + workspace, + intermediate_rotations, + gate_suh, + up_suh, + trellis_lut_addr, + tier0_num_experts, + tier1_num_experts, + tier2_num_experts, + tier0_fc2_experts, + tier1_fc2_experts, + tier2_fc2_experts, + tier0_gate_experts, + tier1_gate_experts, + tier2_gate_experts, + tier0_up_experts, + tier1_up_experts, + tier2_up_experts, + active_m, + ).launch( + grid=(grid_x, 1, 1), + block=[self.cta_threads, 1, 1], + min_blocks_per_mp=self.blocks_per_sm, + cooperative=True, + stream=stream, + ) + + @cute.kernel + def kernel3( + self, + rotation_input: cute.Tensor, + rotation_gate: cute.Tensor, + rotation_up: cute.Tensor, + t0_w13: cute.Tensor, + t0_w2: cute.Tensor, + t0_w13_scales: cute.Tensor, + t0_w2_scales: cute.Tensor, + t0_w13_global: cute.Tensor, + t0_w2_global: cute.Tensor, + t1_w13: cute.Tensor, + t1_w2: cute.Tensor, + t1_w13_scales: cute.Tensor, + t1_w2_scales: cute.Tensor, + t1_w13_global: cute.Tensor, + t1_w2_global: cute.Tensor, + t2_w13: cute.Tensor, + t2_w2: cute.Tensor, + t2_w13_scales: cute.Tensor, + t2_w2_scales: cute.Tensor, + t2_w13_global: cute.Tensor, + t2_w2_global: cute.Tensor, + fc1: cute.Tensor, + activated: cute.Tensor, + fc2: cute.Tensor, + packed_route_indices: cute.Tensor, + block_expert_ids: cute.Tensor, + packed_route_count: cute.Tensor, + descriptor_map: cute.Tensor, + topk_weights: cute.Tensor, + fc1_scratch: cute.Tensor, + fc2_scratch: cute.Tensor, + workspace: cute.Tensor, + intermediate_rotations: cute.Tensor, + gate_suh: cute.Tensor, + up_suh: cute.Tensor, + trellis_lut_addr: Int64, + tier0_num_experts: cutlass.Int32, + tier1_num_experts: cutlass.Int32, + tier2_num_experts: cutlass.Int32, + tier0_fc2_experts: cutlass.Int32, + tier1_fc2_experts: cutlass.Int32, + tier2_fc2_experts: cutlass.Int32, + tier0_gate_experts: cutlass.Int32, + tier1_gate_experts: cutlass.Int32, + tier2_gate_experts: cutlass.Int32, + tier0_up_experts: cutlass.Int32, + tier1_up_experts: cutlass.Int32, + tier2_up_experts: cutlass.Int32, + active_m: cutlass.Int32, + ): + tidx, _, _ = cute.arch.thread_idx() + bidx, _, _ = cute.arch.block_idx() + grid_x_raw, _, _ = cute.arch.grid_dim() + tid = Int32(tidx) + cta = Int32(bidx) + grid_x = Int32(grid_x_raw) + smem = cutlass.utils.SmemAllocator() + + @cute.struct + class Storage: + words: cute.struct.Align[ + cute.struct.MemRange[cutlass.Uint32, self.shared_words], 1024 + ] + + storage = smem.allocate(Storage) + smem_base = shared_ptr_to_u32(storage.words.data_ptr()) + decode_lut_addr = trellis_lut_addr + if cutlass.const_expr(self.driver.sqg_xor_cheb_t12_smem): + self.driver._sqg_smem_copy( + trellis_lut_addr, + smem_base + Int32(self.sqg_xor_cheb_t12_smem_off), + 1 << 12, + tid, + ) + cute.arch.sync_threads() + decode_lut_addr = Int64(smem_base + Int32(self.sqg_xor_cheb_t12_smem_off)) + common = ( + packed_route_indices, + block_expert_ids, + descriptor_map, + topk_weights, + ) + counts = ( + tier0_num_experts, + tier1_num_experts, + tier2_num_experts, + tier0_fc2_experts, + tier1_fc2_experts, + tier2_fc2_experts, + tier0_gate_experts, + tier1_gate_experts, + tier2_gate_experts, + tier0_up_experts, + tier1_up_experts, + tier2_up_experts, + ) + fc1_emit = partial( + self._emit_tier_tile3, + True, + rotation_gate, + rotation_up, + t0_w13, + t0_w13_scales, + t0_w13_global, + t1_w13, + t1_w13_scales, + t1_w13_global, + t2_w13, + t2_w13_scales, + t2_w13_global, + fc1, + *common, + fc1_scratch, + workspace, + decode_lut_addr, + smem_base, + tid, + active_m, + *counts, + ) + fc2_emit = partial( + self._emit_tier_tile3, + False, + activated, + activated, + t0_w2, + t0_w2_scales, + t0_w2_global, + t1_w2, + t1_w2_scales, + t1_w2_global, + t2_w2, + t2_w2_scales, + t2_w2_global, + fc2, + *common, + fc2_scratch, + workspace, + decode_lut_addr, + smem_base, + tid, + active_m * Int32(self.top_k), + *counts, + ) + total_experts = tier0_num_experts + tier1_num_experts + tier2_num_experts + self.driver._moe_body( + rotation_gate, + rotation_up, + rotation_input, + t0_w13, + t0_w2, + fc1, + activated, + fc2, + t0_w13_scales, + t0_w2_scales, + t0_w13_global, + t0_w2_global, + packed_route_indices, + block_expert_ids, + packed_route_count, + t0_w13_global, + Int32(0), + topk_weights, + fc1_scratch, + fc2_scratch, + workspace, + intermediate_rotations, + gate_suh, + up_suh, + descriptor_map, + # The tier emit hooks own Trellis decoding. The shared driver's + # generic LUT ABI slots must stay unused for every mixed bitrate. + cutlass.Int64(0), + cutlass.Int64(0), + total_experts, + total_experts, + smem_base, + tid, + cta, + grid_x, + active_m, + fc1_emit, + fc2_emit, + ) + + +_CACHE: dict[tuple[object, ...], MixedTrellisCompileResult] = {} +_CACHE3: dict[tuple[object, ...], MixedTrellis3CompileResult] = {} +_ROUTE_PACK_WARMED: set[tuple[object, ...]] = set() + + +def _mixed_route_num_experts( + expert_map: torch.Tensor, + expected_route_num_experts: int, +) -> int: + route_num_experts = int(expert_map.numel()) + if route_num_experts != int(expected_route_num_experts): + raise ValueError( + "mixed Trellis route map must match the compiled route namespace: " + f"map={route_num_experts}, compiled={int(expected_route_num_experts)}" + ) + return route_num_experts + + +def warmup_mixed_trellis_route_pack( + launch: MixedTrellisCompileResult, + buffers: MixedTrellisBuffers, + *, + expert_map: torch.Tensor, +) -> int: + """Materialize every route-pack specialization reachable by ``launch``. + + Route packing buckets token capacity to powers of two. A profile pass at + the maximum batch therefore does not cover smaller decode, speculative, + or final-prefill-chunk buckets. Load those CUDA modules eagerly while the + serving framework is still profiling persistent memory, so KV sizing sees + their real driver footprint instead of discovering it under live traffic. + """ + device = buffers.packed_route_indices.device + if device.type != "cuda": + raise RuntimeError("mixed Trellis route-pack warmup requires CUDA buffers") + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError("mixed Trellis route-pack warmup cannot run during capture") + + route_num_experts = _mixed_route_num_experts( + expert_map, int(launch.topk_sum.route_num_experts) + ) + warmed = 0 + pending_keys: list[tuple[object, ...]] = [] + with torch.cuda.device(device): + device_index = int(torch.cuda.current_device()) + for token_count in route_pack_warmup_token_counts(launch.size_m): + key = ( + device_index, + str(launch.route_ids_dtype), + int(token_count), + int(launch.top_k), + route_num_experts, + int(launch.moe_block_size), + True, + ) + if key in _ROUTE_PACK_WARMED: + continue + dummy_topk_ids = torch.zeros( + (token_count, launch.top_k), + dtype=launch.route_ids_dtype, + device=device, + ) + pack_topk_routes_by_expert( + dummy_topk_ids, + launch.moe_block_size, + route_num_experts, + expert_map=expert_map, + packed_route_indices=buffers.packed_route_indices, + block_expert_ids=buffers.block_expert_ids, + packed_route_count=buffers.packed_route_count, + expert_offsets=buffers.expert_offsets, + expert_counts=buffers.expert_counts, + ) + pending_keys.append(key) + warmed += 1 + torch.cuda.current_stream(device).synchronize() + _ROUTE_PACK_WARMED.update(pending_keys) + return warmed + + +def compile_mixed_trellis( + *, + size_m: int, + hidden_size: int, + intermediate_size: int, + tier0_num_experts: int, + tier1_num_experts: int, + top_k: int, + max_m_blocks: int, + sms: int, + max_shared_mem: int, + force_tile_config: tuple[int, int, int, int], + tier0_bits: int = 3, + tier1_bits: int = 4, + trellis_codebook: str = "mcg", + moe_block_size: int = 8, + rotation_input_dtype: str = "bf16", + route_ids_dtype: torch.dtype = torch.int32, + broadcast_suh: bool = False, + broadcast_svh: bool = False, + route_num_experts: int | None = None, +) -> MixedTrellisCompileResult: + if route_ids_dtype not in (torch.int32, torch.int64): + raise TypeError("mixed Trellis route IDs must be int32 or int64") + if int(size_m) * int(top_k) > torch.iinfo(torch.int32).max: + raise ValueError("mixed Trellis routed-row count must fit in int32") + fc1_tile_k, fc1_tile_n, fc2_tile_k, fc2_tile_n = ( + int(value) for value in force_tile_config + ) + if fc1_tile_k < 128: + raise ValueError( + "mixed Trellis FC1 requires tile_k >= 128; narrower K tiles lose " + "large-M cross-tier partial reductions" + ) + trellis_codebook = str(trellis_codebook).lower() + total_experts = int(tier0_num_experts) + int(tier1_num_experts) + if route_num_experts is None: + route_num_experts = total_experts + route_num_experts = int(route_num_experts) + if route_num_experts <= 0: + raise ValueError("mixed Trellis route_num_experts must be positive") + paired_m8_fc2 = int(moe_block_size) in (32, 64) + + def make_kernel(num_experts: int, bits: int) -> W4A16FusedMoeKernel: + return W4A16FusedMoeKernel( + size_m=size_m, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + num_experts=num_experts, + top_k=top_k, + activation="silu", + apply_router_weight_on_input=False, + zero_fc2_output=False, + fc1_tile_n=fc1_tile_n, + fc1_tile_k=fc1_tile_k, + fc2_tile_n=fc2_tile_n, + fc2_tile_k=fc2_tile_k, + moe_block_size=moe_block_size, + max_m_blocks=max_m_blocks, + fc2_moe_block_size=(8 if paired_m8_fc2 else moe_block_size), + fc2_schedule_route_block_factor=(2 if paired_m8_fc2 else 1), + element_dtype="fp16", + weight_layout="trellis3_t256", + scale_format="e4m3_k32", + w13_layout="trellis3_t256_proj", + trellis_bits=bits, + trellis_codebook=trellis_codebook, + intermediate_rotation=True, + full_rotation=True, + rotation_input_dtype=rotation_input_dtype, + broadcast_suh=broadcast_suh, + schedule_whole_tiles=True, + ) + + kernel = W4A16MixedTrellisKernel( + driver=make_kernel(total_experts, tier0_bits), + tier0=make_kernel(int(tier0_num_experts), int(tier0_bits)), + tier1=make_kernel(int(tier1_num_experts), int(tier1_bits)), + ) + # shared_words is the complete dynamically allocated MemRange used by the + # cooperative kernel. CUDA permits a launch exactly at the device's + # opt-in shared-memory limit; rejecting an additional 512 bytes here + # unnecessarily excludes the stock mixed-K tile geometry at block-64. + if kernel.shared_words * 4 > int(max_shared_mem): + raise ValueError( + "mixed Trellis shared-memory requirement exceeds the device limit: " + f"required={kernel.shared_words * 4} " + f"limit={int(max_shared_mem)}" + ) + device = int(torch.cuda.current_device()) + cache_key = ( + "mixed_trellis", + device, + kernel.__cache_key__, + str(route_ids_dtype), + int(size_m), + int(max_m_blocks), + ) + topk_sum = compile_w4a16_topk_sum( + m=size_m, + topk=top_k, + hidden_size=hidden_size, + element_dtype="fp16", + full_rotation=True, + num_experts=total_experts, + route_num_experts=route_num_experts, + route_ids_dtype=route_ids_dtype, + use_expert_map=True, + broadcast_svh=broadcast_svh, + ) + cached = _CACHE.get(cache_key) + if cached is not None: + # The compiled object is intentionally independent of the artifact's + # K3/K4 partition. Keep the current plan metadata and top-k launch, + # rather than leaking the first split that populated the cache. + return replace( + cached, + topk_sum=topk_sum, + tier0_num_experts=int(tier0_num_experts), + tier1_num_experts=int(tier1_num_experts), + sms=int(sms), + broadcast_suh=bool(broadcast_suh), + broadcast_svh=bool(broadcast_svh), + ) + + compile_m = _fake_m_for_specialization(size_m) + compile_rows = compile_m * top_k + fc1_cols = 2 * intermediate_size + cutlass_dtype = cutlass.Float16 + rotation_dtype = _cutlass_element_dtype(rotation_input_dtype) + + def tensor(dtype, elements: int, *, align: int = 16): + return cute.runtime.make_fake_compact_tensor( + dtype, (max(int(elements), 1),), assumed_align=align + ) + + def tier_args(): + return ( + make_ptr(cutlass.Int32, 16, cute.AddressSpace.gmem, assumed_align=16), + make_ptr(cutlass.Int32, 16, cute.AddressSpace.gmem, assumed_align=16), + make_ptr(cutlass.Int32, 16, cute.AddressSpace.gmem, assumed_align=16), + make_ptr(cutlass.Int32, 16, cute.AddressSpace.gmem, assumed_align=16), + make_ptr(cutlass.Float32, 16, cute.AddressSpace.gmem, assumed_align=16), + make_ptr(cutlass.Float32, 16, cute.AddressSpace.gmem, assumed_align=16), + ) + + scratch_elements = max( + fc1_cols * compile_rows, + hidden_size * compile_rows, + 4 * 256 * moe_block_size * 256, + ) + compile_args = ( + make_ptr(rotation_dtype, 16, cute.AddressSpace.gmem, assumed_align=16), + tensor(cutlass_dtype, compile_rows * hidden_size), + tensor(cutlass_dtype, compile_rows * hidden_size), + *tier_args(), + *tier_args(), + tensor(cutlass_dtype, compile_rows * fc1_cols), + tensor(cutlass_dtype, compile_rows * intermediate_size), + tensor(cutlass_dtype, compile_rows * hidden_size), + tensor(cutlass.Int32, moe_block_size), + tensor(cutlass.Int32, 1), + tensor(cutlass.Int32, 1, align=4), + make_ptr(cutlass.Int32, 4, cute.AddressSpace.gmem, assumed_align=4), + make_ptr(cutlass.Float32, 4, cute.AddressSpace.gmem, assumed_align=4), + tensor(cutlass.Float32, scratch_elements), + tensor(cutlass.Float32, scratch_elements), + tensor(cutlass.Int32, 4 * 256 + 2), + make_ptr(cutlass.Float16, 16, cute.AddressSpace.gmem, assumed_align=16), + make_ptr(cutlass.Float16, 16, cute.AddressSpace.gmem, assumed_align=16), + make_ptr(cutlass.Float16, 16, cute.AddressSpace.gmem, assumed_align=16), + make_ptr(cutlass.Uint8, 16, cute.AddressSpace.gmem, assumed_align=16), + Int32(tier0_num_experts), + Int32(tier1_num_experts), + # FC2 counts are independent artifact data; trace with the FC1 values. + Int32(tier0_num_experts), + Int32(tier1_num_experts), + 1, + 1, + current_cuda_stream(), + # Gate-count trace placeholders (keyword-only, last in the signature); + # real counts ride each launch. + Int32(tier0_num_experts), + Int32(tier1_num_experts), + # Up-count trace placeholders; real counts ride each launch too. + Int32(tier0_num_experts), + Int32(tier1_num_experts), + ) + raise_if_kernel_resolution_frozen( + "cute.compile", target=kernel, cache_key=cache_key + ) + compiled = b12x_compile( + kernel, + *compile_args, + compile_spec=KernelCompileSpec.from_key( + "moe.w4a16.mixed_trellis", W4A16MixedTrellisKernel.ABI_VERSION, cache_key + ), + dsl_compile_options=OptLevel(2), + ) + registers = -1 + local_bytes = -1 + resources = _query_w4a16_kernel_resources(compiled) + if resources is not None: + _, registers, local_bytes = resources + if local_bytes != 0: + raise RuntimeError( + "mixed Trellis codegen spills to local memory " + f"({local_bytes} bytes/thread)" + ) + result = MixedTrellisCompileResult( + compiled=compiled, + topk_sum=topk_sum, + size_m=int(size_m), + hidden_size=int(hidden_size), + intermediate_size=int(intermediate_size), + top_k=int(top_k), + tier0_num_experts=int(tier0_num_experts), + tier1_num_experts=int(tier1_num_experts), + tier0_bits=int(tier0_bits), + tier1_bits=int(tier1_bits), + trellis_codebook=trellis_codebook, + fc1_tile_k=fc1_tile_k, + fc1_tile_n=fc1_tile_n, + fc2_tile_k=fc2_tile_k, + fc2_tile_n=fc2_tile_n, + moe_block_size=int(moe_block_size), + fc2_moe_block_size=int(kernel.driver.fc2.moe_block_size), + fc2_schedule_route_block_factor=int( + kernel.driver.fc2.schedule_route_block_factor + ), + fc2_paired_m8_routes=bool(kernel.driver.fc2.paired_m8_routes), + max_m_blocks=int(max_m_blocks), + blocks_per_sm=int(kernel.blocks_per_sm), + sms=int(sms), + shared_memory_bytes=int(kernel.shared_words * 4), + registers_per_thread=registers, + local_memory_bytes=local_bytes, + rotation_input_dtype=str(rotation_input_dtype), + route_ids_dtype=route_ids_dtype, + broadcast_suh=bool(broadcast_suh), + broadcast_svh=bool(broadcast_svh), + ) + _CACHE[cache_key] = result + return result + + +def compile_mixed_trellis3( + *, + size_m: int, + hidden_size: int, + intermediate_size: int, + tier0_num_experts: int, + tier1_num_experts: int, + tier2_num_experts: int, + top_k: int, + max_m_blocks: int, + sms: int, + max_shared_mem: int, + force_tile_config: tuple[int, int, int, int], + tier0_bits: int = 3, + tier1_bits: int = 4, + tier2_bits: int = 5, + trellis_codebook: str = "mcg", + moe_block_size: int = 8, + rotation_input_dtype: str = "bf16", + route_ids_dtype: torch.dtype = torch.int32, + broadcast_suh: bool = False, + broadcast_svh: bool = False, + route_num_experts: int | None = None, +) -> MixedTrellis3CompileResult: + """Compile the dedicated three-bitrate cooperative Trellis grid.""" + + if route_ids_dtype not in (torch.int32, torch.int64): + raise TypeError("mixed Trellis route IDs must be int32 or int64") + if int(size_m) * int(top_k) > torch.iinfo(torch.int32).max: + raise ValueError("mixed Trellis routed-row count must fit in int32") + trellis_codebook = str(trellis_codebook).lower() + counts = tuple( + int(value) + for value in ( + tier0_num_experts, + tier1_num_experts, + tier2_num_experts, + ) + ) + if any(value <= 0 or value > _MAX_TIER_EXPERTS for value in counts): + raise ValueError( + "three-tier mixed Trellis requires each tier to contain 1..256 slots" + ) + bits = tuple(int(value) for value in (tier0_bits, tier1_bits, tier2_bits)) + if len(set(bits)) != 3 or any(value not in (3, 4, 5, 6) for value in bits): + raise ValueError( + "three-tier mixed Trellis requires three distinct bitrates in 3..6" + ) + fc1_tile_k, fc1_tile_n, fc2_tile_k, fc2_tile_n = ( + int(value) for value in force_tile_config + ) + if fc1_tile_k < 128: + raise ValueError( + "mixed Trellis FC1 requires tile_k >= 128; narrower K tiles lose " + "large-M cross-tier partial reductions" + ) + total_experts = sum(counts) + if route_num_experts is None: + route_num_experts = total_experts + route_num_experts = int(route_num_experts) + if route_num_experts <= 0: + raise ValueError("mixed Trellis route_num_experts must be positive") + paired_m8_fc2 = int(moe_block_size) in (32, 64) + + def make_kernel(num_experts: int, trellis_bits: int) -> W4A16FusedMoeKernel: + return W4A16FusedMoeKernel( + size_m=size_m, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + num_experts=num_experts, + top_k=top_k, + activation="silu", + apply_router_weight_on_input=False, + zero_fc2_output=False, + fc1_tile_n=fc1_tile_n, + fc1_tile_k=fc1_tile_k, + fc2_tile_n=fc2_tile_n, + fc2_tile_k=fc2_tile_k, + moe_block_size=moe_block_size, + max_m_blocks=max_m_blocks, + fc2_moe_block_size=(8 if paired_m8_fc2 else moe_block_size), + fc2_schedule_route_block_factor=(2 if paired_m8_fc2 else 1), + element_dtype="fp16", + weight_layout="trellis3_t256", + scale_format="e4m3_k32", + w13_layout="trellis3_t256_proj", + trellis_bits=trellis_bits, + trellis_codebook=trellis_codebook, + intermediate_rotation=True, + full_rotation=True, + rotation_input_dtype=rotation_input_dtype, + broadcast_suh=broadcast_suh, + schedule_whole_tiles=True, + ) + + kernel = W4A16MixedTrellis3Kernel( + driver=make_kernel(total_experts, bits[0]), + tier0=make_kernel(counts[0], bits[0]), + tier1=make_kernel(counts[1], bits[1]), + tier2=make_kernel(counts[2], bits[2]), + ) + if kernel.shared_words * 4 > int(max_shared_mem): + raise ValueError( + "mixed Trellis shared-memory requirement exceeds the device limit: " + f"required={kernel.shared_words * 4} limit={int(max_shared_mem)}" + ) + device = int(torch.cuda.current_device()) + cache_key = ( + "mixed_trellis3", + device, + kernel.__cache_key__, + str(route_ids_dtype), + int(size_m), + int(max_m_blocks), + ) + topk_sum = compile_w4a16_topk_sum( + m=size_m, + topk=top_k, + hidden_size=hidden_size, + element_dtype="fp16", + full_rotation=True, + num_experts=total_experts, + route_num_experts=route_num_experts, + route_ids_dtype=route_ids_dtype, + use_expert_map=True, + broadcast_svh=broadcast_svh, + ) + cached = _CACHE3.get(cache_key) + if cached is not None: + return replace( + cached, + topk_sum=topk_sum, + tier0_num_experts=counts[0], + tier1_num_experts=counts[1], + tier2_num_experts=counts[2], + sms=int(sms), + broadcast_suh=bool(broadcast_suh), + broadcast_svh=bool(broadcast_svh), + ) + + compile_m = _fake_m_for_specialization(size_m) + compile_rows = compile_m * top_k + fc1_cols = 2 * intermediate_size + cutlass_dtype = cutlass.Float16 + rotation_dtype = _cutlass_element_dtype(rotation_input_dtype) + + def tensor(dtype, elements: int, *, align: int = 16): + return cute.runtime.make_fake_compact_tensor( + dtype, (max(int(elements), 1),), assumed_align=align + ) + + def tier_args(): + return ( + make_ptr(cutlass.Int32, 16, cute.AddressSpace.gmem, assumed_align=16), + make_ptr(cutlass.Int32, 16, cute.AddressSpace.gmem, assumed_align=16), + make_ptr(cutlass.Int32, 16, cute.AddressSpace.gmem, assumed_align=16), + make_ptr(cutlass.Int32, 16, cute.AddressSpace.gmem, assumed_align=16), + make_ptr(cutlass.Float32, 16, cute.AddressSpace.gmem, assumed_align=16), + make_ptr(cutlass.Float32, 16, cute.AddressSpace.gmem, assumed_align=16), + ) + + scratch_elements = max( + fc1_cols * compile_rows, + hidden_size * compile_rows, + 4 * 256 * moe_block_size * 256, + ) + compile_args = ( + make_ptr(rotation_dtype, 16, cute.AddressSpace.gmem, assumed_align=16), + tensor(cutlass_dtype, compile_rows * hidden_size), + tensor(cutlass_dtype, compile_rows * hidden_size), + *tier_args(), + *tier_args(), + *tier_args(), + tensor(cutlass_dtype, compile_rows * fc1_cols), + tensor(cutlass_dtype, compile_rows * intermediate_size), + tensor(cutlass_dtype, compile_rows * hidden_size), + tensor(cutlass.Int32, moe_block_size), + tensor(cutlass.Int32, 1), + tensor(cutlass.Int32, 1, align=4), + make_ptr(cutlass.Int32, 4, cute.AddressSpace.gmem, assumed_align=4), + make_ptr(cutlass.Float32, 4, cute.AddressSpace.gmem, assumed_align=4), + tensor(cutlass.Float32, scratch_elements), + tensor(cutlass.Float32, scratch_elements), + tensor(cutlass.Int32, 4 * 256 + 2), + make_ptr(cutlass.Float16, 16, cute.AddressSpace.gmem, assumed_align=16), + make_ptr(cutlass.Float16, 16, cute.AddressSpace.gmem, assumed_align=16), + make_ptr(cutlass.Float16, 16, cute.AddressSpace.gmem, assumed_align=16), + make_ptr(cutlass.Uint8, 16, cute.AddressSpace.gmem, assumed_align=16), + Int32(counts[0]), + Int32(counts[1]), + Int32(counts[2]), + Int32(counts[0]), + Int32(counts[1]), + Int32(counts[2]), + 1, + 1, current_cuda_stream(), + Int32(counts[0]), + Int32(counts[1]), + Int32(counts[2]), + Int32(counts[0]), + Int32(counts[1]), + Int32(counts[2]), ) raise_if_kernel_resolution_frozen( "cute.compile", target=kernel, cache_key=cache_key @@ -957,24 +2169,35 @@ def tier_args(): kernel, *compile_args, compile_spec=KernelCompileSpec.from_key( - "moe.w4a16.mixed_trellis", W4A16MixedTrellisKernel.ABI_VERSION, cache_key + "moe.w4a16.mixed_trellis3", + W4A16MixedTrellis3Kernel.ABI_VERSION, + cache_key, ), dsl_compile_options=OptLevel(2), ) - result = MixedTrellisCompileResult( + registers = -1 + local_bytes = -1 + resources = _query_w4a16_kernel_resources(compiled) + if resources is not None: + _, registers, local_bytes = resources + if local_bytes != 0: + raise RuntimeError( + "three-tier mixed Trellis codegen spills to local memory " + f"({local_bytes} bytes/thread)" + ) + result = MixedTrellis3CompileResult( compiled=compiled, topk_sum=topk_sum, - trellis_lut=_trellis256_execution_lut( - torch.device("cuda", device), trellis_codebook - ), size_m=int(size_m), hidden_size=int(hidden_size), intermediate_size=int(intermediate_size), top_k=int(top_k), - tier0_num_experts=int(tier0_num_experts), - tier1_num_experts=int(tier1_num_experts), - tier0_bits=int(tier0_bits), - tier1_bits=int(tier1_bits), + tier0_num_experts=counts[0], + tier1_num_experts=counts[1], + tier2_num_experts=counts[2], + tier0_bits=bits[0], + tier1_bits=bits[1], + tier2_bits=bits[2], trellis_codebook=trellis_codebook, fc1_tile_k=fc1_tile_k, fc1_tile_n=fc1_tile_n, @@ -985,29 +2208,32 @@ def tier_args(): fc2_schedule_route_block_factor=int( kernel.driver.fc2.schedule_route_block_factor ), + fc2_paired_m8_routes=bool(kernel.driver.fc2.paired_m8_routes), max_m_blocks=int(max_m_blocks), blocks_per_sm=int(kernel.blocks_per_sm), sms=int(sms), shared_memory_bytes=int(kernel.shared_words * 4), + registers_per_thread=registers, + local_memory_bytes=local_bytes, rotation_input_dtype=str(rotation_input_dtype), route_ids_dtype=route_ids_dtype, broadcast_suh=bool(broadcast_suh), broadcast_svh=bool(broadcast_svh), ) - _CACHE[cache_key] = result + _CACHE3[cache_key] = result return result -def make_mixed_trellis_buffers( - launch: MixedTrellisCompileResult, +def _make_mixed_trellis_buffers( + launch: MixedTrellisCompileResult | MixedTrellis3CompileResult, *, device: torch.device, sms: int, + route_num_experts: int, ) -> MixedTrellisBuffers: capacity_rows = launch.size_m * launch.top_k - total_experts = launch.tier0_num_experts + launch.tier1_num_experts route_slots = max_packed_route_slots( - capacity_rows, launch.moe_block_size, total_experts + capacity_rows, launch.moe_block_size, route_num_experts ) route_blocks = (route_slots + launch.moe_block_size - 1) // launch.moe_block_size if route_blocks > launch.max_m_blocks: @@ -1041,8 +2267,10 @@ def make_mixed_trellis_buffers( packed_route_indices=torch.empty(route_slots, dtype=torch.int32, device=device), block_expert_ids=torch.empty(route_blocks, dtype=torch.int32, device=device), packed_route_count=torch.empty(1, dtype=torch.int32, device=device), - expert_offsets=torch.empty(total_experts + 1, dtype=torch.int32, device=device), - expert_counts=torch.empty(total_experts, dtype=torch.int32, device=device), + expert_offsets=torch.empty( + route_num_experts + 1, dtype=torch.int32, device=device + ), + expert_counts=torch.empty(route_num_experts, dtype=torch.int32, device=device), fc1_scratch=torch.empty( packed_gemm_scratch_elements( size_n=fc1_cols, @@ -1071,6 +2299,34 @@ def make_mixed_trellis_buffers( ) +def make_mixed_trellis_buffers( + launch: MixedTrellisCompileResult, + *, + device: torch.device, + sms: int, +) -> MixedTrellisBuffers: + return _make_mixed_trellis_buffers( + launch, + device=device, + sms=sms, + route_num_experts=int(launch.topk_sum.route_num_experts), + ) + + +def make_mixed_trellis3_buffers( + launch: MixedTrellis3CompileResult, + *, + device: torch.device, + sms: int, +) -> MixedTrellisBuffers: + return _make_mixed_trellis_buffers( + launch, + device=device, + sms=sms, + route_num_experts=int(launch.topk_sum.route_num_experts), + ) + + def build_ordered_maps( tier0_num_experts: int, tier1_num_experts: int, @@ -1086,6 +2342,10 @@ def build_ordered_maps( ) +# One tier-local expert index is encoded in the descriptor's low 8 bits. +_MAX_TIER_EXPERTS = 256 + + def build_tiered_maps( tier0_global_ids: Sequence[int], tier1_global_ids: Sequence[int], @@ -1111,25 +2371,178 @@ def build_tiered_maps( global_to_combined = torch.tensor( global_to_combined_host, dtype=torch.int32, device=device ) - descriptor = torch.tensor( + descriptor_row = torch.tensor( [*range(len(tier0_ids)), *((1 << 8) | i for i in range(len(tier1_ids)))], dtype=torch.int32, device=device, ) + # The descriptor table carries one row per projection (gate, up, down). + # Per-expert tiering is the degenerate case where all three rows are + # identical, reproducing single-row behaviour bit-for-bit. + descriptor = descriptor_row.repeat(3).contiguous() + # Publish the per-tier gate/up counts this descriptor encodes so + # run_mixed_trellis can fail closed on mismatched caller counts without a + # device sync. Per-expert tiering has gate == up == the tier partition. + counts = (len(tier0_ids), len(tier1_ids)) + descriptor._mt_projection_counts = (counts, counts) + return global_to_combined, descriptor + + +def build_projection_tiered_maps( + gate_tiers: Sequence[int], + up_tiers: Sequence[int], + down_tiers: Sequence[int], + *, + tier_slots: Sequence[int], + device: torch.device, +) -> tuple[torch.Tensor, torch.Tensor]: + """Build the route map and the three-row descriptor map for one R7 layer. + + Each argument is one tier id per global expert, for that + projection. Combined expert ids are the global ids -- with per-projection + tiering there is no single tier-contiguous ordering, so all tier knowledge + lives in the descriptor rows and global_to_combined is the identity. + + Returns (global_to_combined, descriptor_map) where descriptor_map is + int32[3 * sum(tier_slots)] laid out gate, up, down. Real entries are + (tier << 8) | tier_local_index and padding entries are -1. + glm52-r7-projtiers. + """ + + projections = ( + ("gate", tuple(int(t) for t in gate_tiers)), + ("up", tuple(int(t) for t in up_tiers)), + ("down", tuple(int(t) for t in down_tiers)), + ) + slots = tuple(int(value) for value in tier_slots) + if len(slots) not in (2, 3): + raise ValueError( + "mixed Trellis tier_slots must contain exactly two or three counts" + ) + if any(value < 0 or value > 256 for value in slots): + raise ValueError("mixed Trellis tier slots must be in [0, 256]") + num_experts = len(projections[0][1]) + rows: list[int] = [] + projection_counts: list[tuple[int, ...]] = [] + for name, tiers in projections: + if len(tiers) != num_experts: + raise ValueError( + "mixed Trellis projection tier lists must agree on expert count: " + f"{name} has {len(tiers)}, expected {num_experts}" + ) + if any(t < 0 or t >= len(slots) for t in tiers): + raise ValueError( + f"mixed Trellis {name} tier ids must be in [0, {len(slots)})" + ) + counters = [0] * len(slots) + row = [] + for tier in tiers: + local = counters[tier] + counters[tier] += 1 + if local > 0xFF: + raise ValueError( + f"mixed Trellis {name} tier {tier} exceeds 256 experts" + ) + row.append((tier << 8) | local) + projection_counts.append(tuple(counters)) + rows.extend(row) + # The launch sizes the descriptor namespace as the sum of the tier slot + # counts. A tier slot is max(gate_count, up_count), so with per-projection + # tiering that sum can exceed the real expert count. Routing remains a + # separate, exact-size namespace; only the descriptor rows are padded. + required_fc1 = tuple( + max(projection_counts[0][tier], projection_counts[1][tier]) + for tier in range(len(slots)) + ) + if any(slots[tier] < required_fc1[tier] for tier in range(len(slots))): + raise ValueError( + "mixed Trellis tier slots cannot address all gate/up locals: " + f"slots={slots}, required={required_fc1}" + ) + stride = sum(slots) + if stride < num_experts: + raise ValueError( + f"mixed Trellis tier slots ({stride}) cannot address {num_experts} experts" + ) + global_to_combined = torch.arange(num_experts, dtype=torch.int32, device=device) + descriptor = torch.full((3 * stride,), -1, dtype=torch.int32, device=device) + for row_index in range(3): + base = row_index * num_experts + descriptor[row_index * stride : row_index * stride + num_experts] = ( + torch.tensor( + rows[base : base + num_experts], dtype=torch.int32, device=device + ) + ) + # Publish the gate/up counts this descriptor encodes so run_mixed_trellis + # can fail closed on mismatched caller counts without a device sync. + descriptor._mt_projection_counts = ( + projection_counts[0], + projection_counts[1], + ) return global_to_combined, descriptor +def _check_descriptor_projection_counts( + descriptor_map: torch.Tensor, + total_experts: int, + *, + gate_counts: tuple[int, ...], + up_counts: tuple[int, ...], +) -> None: + """Fail closed when launch counts disagree with the descriptor map. + + The device dispatch bounds gate/up locals by the launch counts, so a + descriptor entry at or beyond its projection's count is silently skipped + and the corresponding FC1 half stays zero. The in-tree builders publish + the counts their descriptor encodes; descriptors from other producers pay + one host copy here, memoized on the tensor, so steady-state launches + never synchronize. + """ + + encoded = getattr(descriptor_map, "_mt_projection_counts", None) + if encoded is None: + rows = descriptor_map.detach().cpu().view(3, total_experts) + derived = [] + tier_count = len(gate_counts) + for row in rows[:2]: + live = row[row >= 0] + encoded_tiers = live >> 8 + if bool((encoded_tiers >= tier_count).any()): + raise ValueError( + "mixed Trellis descriptor contains a tier outside the " + f"launch range [0, {tier_count})" + ) + derived.append( + tuple(int((encoded_tiers == tier).sum()) for tier in range(tier_count)) + ) + encoded = (tuple(derived[0]), tuple(derived[1])) + descriptor_map._mt_projection_counts = encoded + expected_gate = tuple(int(value) for value in encoded[0]) + expected_up = tuple(int(value) for value in encoded[1]) + got_gate = tuple(int(value) for value in gate_counts) + got_up = tuple(int(value) for value in up_counts) + if got_gate != expected_gate or got_up != expected_up: + raise ValueError( + "mixed Trellis projection counts disagree with the descriptor " + f"map: gate {got_gate} vs encoded {expected_gate}, up {got_up} " + f"vs encoded {expected_up}" + ) + + def combine_trellis_rotations( - tier0: MixedTrellisTier, tier1: MixedTrellisTier + tier0: MixedTrellisTier, + tier1: MixedTrellisTier, + *additional_tiers: MixedTrellisTier, ) -> MixedTrellisRotations: """Materialize one tier-ordered table set once during model preparation.""" + tiers = (tier0, tier1, *additional_tiers) return MixedTrellisRotations( intermediate=torch.cat( - (tier0.intermediate_rotations, tier1.intermediate_rotations), dim=0 + tuple(tier.intermediate_rotations for tier in tiers), dim=0 ).contiguous(), - gate_suh=torch.cat((tier0.gate_suh, tier1.gate_suh), dim=0).contiguous(), - up_suh=torch.cat((tier0.up_suh, tier1.up_suh), dim=0).contiguous(), - down_svh=torch.cat((tier0.down_svh, tier1.down_svh), dim=0).contiguous(), + gate_suh=torch.cat(tuple(tier.gate_suh for tier in tiers), dim=0).contiguous(), + up_suh=torch.cat(tuple(tier.up_suh for tier in tiers), dim=0).contiguous(), + down_svh=torch.cat(tuple(tier.down_svh for tier in tiers), dim=0).contiguous(), ) @@ -1142,23 +2555,82 @@ def _validate_mixed_trellis_tier_storage( hidden_size: int, intermediate_size: int, device: torch.device, + gate_experts: int | None = None, + up_experts: int | None = None, ) -> None: """Fail closed before binding expert-sized storage as raw CuTe pointers.""" expected_experts = int(expected_experts) bits = int(bits) - fc1_cols = 2 * int(intermediate_size) + # Per-projection membership lets a tier hold a different number of FC2 + # (down) experts than FC1 (gate/up) slots, so the FC2 count cannot be + # assumed equal to expected_experts. Derive it from the W2 payload itself, + # which is the tensor that actually carries the data, then require the + # global-scale vector to agree. Deriving it from the scale vector instead + # would report a malformed scale as a confusing W2 extent error. + w2_expert_stride = ( + (int(intermediate_size) // 16) * (int(hidden_size) // 16) * (8 * bits) + ) + w2_elements = int(tier.w2.numel()) + # The FC2 count is NOT bounded by the FC1 slot count: per-projection + # membership routinely gives a tier more down experts than gate/up slots + # (measured 231 down vs 77 gate/up on a real R7 layer). The real ceiling is + # the descriptor's 8-bit tier-local index. + if ( + tier.w2.dtype != torch.int32 + or w2_expert_stride <= 0 + or w2_elements % w2_expert_stride != 0 + or not 1 <= w2_elements // w2_expert_stride <= _MAX_TIER_EXPERTS + ): + raise ValueError( + f"mixed Trellis {name}.w2 must be torch.int32 holding " + f"1..{_MAX_TIER_EXPERTS} whole FC2 experts of " + f"{w2_expert_stride} elements, got {w2_elements}" + ) + fc2_experts = w2_elements // w2_expert_stride + if gate_experts is None and up_experts is None: + gate_experts = expected_experts + up_experts = expected_experts + elif gate_experts is None or up_experts is None: + raise ValueError( + f"mixed Trellis {name} requires paired gate_experts/up_experts" + ) + gate_experts = int(gate_experts) + up_experts = int(up_experts) + if not ( + 0 <= gate_experts <= expected_experts and 0 <= up_experts <= expected_experts + ): + raise ValueError( + f"mixed Trellis {name} projection counts must both be in " + f"[0, {expected_experts}], got gate={gate_experts}, up={up_experts}" + ) + # Legacy callers have G=U=E and therefore require the historical 2E + # planes. Tight callers must couple the physical payload exactly to G+U. + # One dummy plane is permitted only for the empty/empty tier. + _proj_stride = ( + (int(hidden_size) // 16) * (int(intermediate_size) // 16) * (8 * bits) + ) + _w13_planes = int(tier.w13.numel()) // _proj_stride if _proj_stride else 0 + _expected_w13_planes = max(gate_experts + up_experts, 1) + if ( + tier.w13.dtype != torch.int32 + or tier.w13.device != device + or not tier.w13.is_contiguous() + or _proj_stride <= 0 + or int(tier.w13.numel()) % _proj_stride != 0 + or _w13_planes != _expected_w13_planes + or int(tier.w13.data_ptr()) % 16 != 0 + ): + raise ValueError( + f"mixed Trellis {name}.w13 must be contiguous int32 on {device} " + f"with exactly {_expected_w13_planes} projection planes, got " + f"{int(tier.w13.numel())} elements" + ) expected = ( - ( - "w13", - tier.w13, - torch.int32, - expected_experts * (int(hidden_size) // 16) * (fc1_cols // 16) * (8 * bits), - ), ( "w2", tier.w2, torch.int32, - expected_experts + fc2_experts * (int(intermediate_size) // 16) * (int(hidden_size) // 16) * (8 * bits), @@ -1178,7 +2650,7 @@ def _validate_mixed_trellis_tier_storage( "w2_global_scale", tier.w2_global_scale, torch.float32, - expected_experts, + fc2_experts, ), ) for field, tensor, expected_dtype, expected_elements in expected: @@ -1207,7 +2679,34 @@ def run_mixed_trellis( rotations: MixedTrellisRotations, launch: MixedTrellisCompileResult, buffers: MixedTrellisBuffers, + gate_experts: tuple[int, int] | None = None, + up_experts: tuple[int, int] | None = None, ) -> torch.Tensor: + def projection_counts(name, values, defaults): + if values is None: + return defaults + if ( + not isinstance(values, tuple) + or len(values) != 2 + or any( + not isinstance(value, int) or isinstance(value, bool) + for value in values + ) + ): + raise TypeError(f"mixed Trellis {name} must be a pair of integer counts") + return values + + if (gate_experts is None) != (up_experts is None): + raise ValueError( + "mixed Trellis projection-tight storage requires paired " + "gate_experts/up_experts" + ) + defaults = ( + int(launch.tier0_num_experts), + int(launch.tier1_num_experts), + ) + _gate0, _gate1 = projection_counts("gate_experts", gate_experts, defaults) + _up0, _up1 = projection_counts("up_experts", up_experts, defaults) m = int(x.shape[0]) if m <= 0: raise ValueError(f"mixed Trellis requires at least one active row, got {m}") @@ -1236,8 +2735,14 @@ def run_mixed_trellis( actual_experts = int(tier.num_experts) if actual_experts != int(expected_experts): raise ValueError( - f"mixed Trellis {name} has {actual_experts} experts, but the " - f"launch plan describes {int(expected_experts)}" + f"mixed Trellis {name} has {actual_experts} experts, expected " + f"the launch-plan count {int(expected_experts)}" + ) + tier_codebook = str(tier.trellis_codebook).lower() + if tier_codebook != launch.trellis_codebook: + raise ValueError( + f"mixed Trellis {name} uses codebook {tier_codebook!r}, expected " + f"the launch-plan codebook {launch.trellis_codebook!r}" ) _validate_mixed_trellis_tier_storage( name="tier0", @@ -1247,6 +2752,8 @@ def run_mixed_trellis( hidden_size=launch.hidden_size, intermediate_size=launch.intermediate_size, device=x.device, + gate_experts=_gate0, + up_experts=_up0, ) _validate_mixed_trellis_tier_storage( name="tier1", @@ -1256,22 +2763,35 @@ def run_mixed_trellis( hidden_size=launch.hidden_size, intermediate_size=launch.intermediate_size, device=x.device, + gate_experts=_gate1, + up_experts=_up1, ) total_experts = launch.tier0_num_experts + launch.tier1_num_experts - for name, mapping in ( - ("global_to_combined", global_to_combined), - ("descriptor_map", descriptor_map), + route_num_experts = _mixed_route_num_experts( + global_to_combined, int(launch.topk_sum.route_num_experts) + ) + for name, mapping, expected_entries in ( + ("global_to_combined", global_to_combined, route_num_experts), + ("descriptor_map", descriptor_map, 3 * total_experts), ): + # glm52-r7-projtiers: descriptor rows use the padded weight stride, + # while the route map covers only real global experts. if ( mapping.dtype != torch.int32 or mapping.device != x.device or not mapping.is_contiguous() - or int(mapping.numel()) != total_experts + or int(mapping.numel()) != expected_entries ): raise ValueError( f"mixed Trellis {name} must be contiguous int32 on {x.device} " - f"with {total_experts} elements" + f"with {expected_entries} elements" ) + _check_descriptor_projection_counts( + descriptor_map, + total_experts, + gate_counts=(_gate0, _gate1), + up_counts=(_up0, _up1), + ) for name, table, expected_elements in ( ( "intermediate rotations", @@ -1306,7 +2826,7 @@ def run_mixed_trellis( f"with {expected_elements} elements and at least 16-byte alignment" ) required_route_slots = max_packed_route_slots( - m * launch.top_k, launch.moe_block_size, total_experts + m * launch.top_k, launch.moe_block_size, route_num_experts ) required_route_blocks = ( required_route_slots + launch.moe_block_size - 1 @@ -1328,7 +2848,7 @@ def run_mixed_trellis( packed, block_experts, packed_count = pack_topk_routes_by_expert( topk_ids, launch.moe_block_size, - total_experts, + route_num_experts, expert_map=global_to_combined, packed_route_indices=buffers.packed_route_indices, block_expert_ids=buffers.block_expert_ids, @@ -1337,6 +2857,7 @@ def run_mixed_trellis( expert_counts=buffers.expert_counts, ) stream = current_cuda_stream() + trellis_rank_lut = sqg_xor_cheb_t12_lut(x.device) launch.compiled( make_ptr( _cutlass_element_dtype(launch.rotation_input_dtype), @@ -1459,15 +2980,368 @@ def run_mixed_trellis( ), make_ptr( cutlass.Uint8, - launch.trellis_lut.data_ptr(), + trellis_rank_lut.data_ptr(), cute.AddressSpace.gmem, assumed_align=16, ), Int32(launch.tier0_num_experts), Int32(launch.tier1_num_experts), + Int32(int(tier0.w2_global_scale.numel())), + Int32(int(tier1.w2_global_scale.numel())), + m, + max(int(launch.blocks_per_sm) * int(launch.sms), 1), + stream, + # Keyword, so the positional chain above stays intact. + tier0_gate_experts=Int32(_gate0), + tier1_gate_experts=Int32(_gate1), + tier0_up_experts=Int32(_up0), + tier1_up_experts=Int32(_up1), + ) + launch.topk_sum.compiled( + make_ptr( + cutlass.Float16, + buffers.fc2.data_ptr(), + cute.AddressSpace.gmem, + assumed_align=16, + ), + make_ptr( + cutlass.Float32, + buffers.output.data_ptr(), + cute.AddressSpace.gmem, + assumed_align=16, + ), + make_ptr( + cutlass.Float32, + topk_weights.data_ptr(), + cute.AddressSpace.gmem, + assumed_align=4, + ), + make_ptr( + cutlass.Int32 if topk_ids.dtype == torch.int32 else cutlass.Int64, + topk_ids.data_ptr(), + cute.AddressSpace.gmem, + assumed_align=4 if topk_ids.dtype == torch.int32 else 8, + ), + make_ptr( + cutlass.Int32, + global_to_combined.data_ptr(), + cute.AddressSpace.gmem, + assumed_align=4, + ), + make_ptr( + cutlass.Float16, + rotations.down_svh.data_ptr(), + cute.AddressSpace.gmem, + assumed_align=16, + ), + Int32(launch.topk_sum.num_experts), + Int32(launch.topk_sum.route_num_experts), + m, + stream, + ) + return buffers.output[:m] + + +def run_mixed_trellis3( + x: torch.Tensor, + tier0: MixedTrellisTier, + tier1: MixedTrellisTier, + tier2: MixedTrellisTier, + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, + global_to_combined: torch.Tensor, + descriptor_map: torch.Tensor, + rotations: MixedTrellisRotations, + launch: MixedTrellis3CompileResult, + buffers: MixedTrellisBuffers, + gate_experts: tuple[int, int, int] | None = None, + up_experts: tuple[int, int, int] | None = None, +) -> torch.Tensor: + """Run one graph-safe K3/K4/K5 cooperative MoE launch.""" + + def projection_counts(name, values, defaults): + if values is None: + return defaults + if ( + not isinstance(values, tuple) + or len(values) != 3 + or any( + not isinstance(value, int) or isinstance(value, bool) + for value in values + ) + ): + raise TypeError( + f"three-tier mixed Trellis {name} must contain three integer counts" + ) + return values + + if (gate_experts is None) != (up_experts is None): + raise ValueError( + "mixed Trellis projection-tight storage requires paired " + "gate_experts/up_experts" + ) + counts = ( + int(launch.tier0_num_experts), + int(launch.tier1_num_experts), + int(launch.tier2_num_experts), + ) + gate_counts = projection_counts("gate_experts", gate_experts, counts) + up_counts = projection_counts("up_experts", up_experts, counts) + m = int(x.shape[0]) + if m <= 0: + raise ValueError(f"mixed Trellis requires at least one active row, got {m}") + if m > launch.size_m: + raise ValueError(f"active rows {m} exceed launch capacity {launch.size_m}") + expected_input_dtype = ( + torch.bfloat16 if launch.rotation_input_dtype == "bf16" else torch.float16 + ) + for name, tensor, expected_dtype in ( + ("input", x, expected_input_dtype), + ("topk_ids", topk_ids, launch.route_ids_dtype), + ("topk_weights", topk_weights, torch.float32), + ): + if tensor.dtype != expected_dtype: + raise TypeError( + f"mixed Trellis {name} must be {expected_dtype}, got {tensor.dtype}" + ) + if not tensor.is_contiguous(): + raise ValueError(f"mixed Trellis {name} must be contiguous") + if int(x.data_ptr()) % 16 != 0: + raise ValueError("mixed Trellis input must have at least 16-byte alignment") + + tiers = (tier0, tier1, tier2) + bits = (launch.tier0_bits, launch.tier1_bits, launch.tier2_bits) + for tier_id, (tier, expected_experts, tier_bits) in enumerate( + zip(tiers, counts, bits, strict=True) + ): + actual_experts = int(tier.num_experts) + if actual_experts != expected_experts: + raise ValueError( + f"mixed Trellis tier{tier_id} has {actual_experts} experts, " + f"expected the launch-plan count {expected_experts}" + ) + tier_codebook = str(tier.trellis_codebook).lower() + if tier_codebook != launch.trellis_codebook: + raise ValueError( + f"mixed Trellis tier{tier_id} uses codebook {tier_codebook!r}, " + f"expected the launch-plan codebook {launch.trellis_codebook!r}" + ) + _validate_mixed_trellis_tier_storage( + name=f"tier{tier_id}", + tier=tier, + expected_experts=expected_experts, + bits=tier_bits, + hidden_size=launch.hidden_size, + intermediate_size=launch.intermediate_size, + device=x.device, + gate_experts=gate_counts[tier_id], + up_experts=up_counts[tier_id], + ) + + total_experts = sum(counts) + route_num_experts = _mixed_route_num_experts( + global_to_combined, int(launch.topk_sum.route_num_experts) + ) + for name, mapping, expected_entries in ( + ("global_to_combined", global_to_combined, route_num_experts), + ("descriptor_map", descriptor_map, 3 * total_experts), + ): + if ( + mapping.dtype != torch.int32 + or mapping.device != x.device + or not mapping.is_contiguous() + or int(mapping.numel()) != expected_entries + ): + raise ValueError( + f"mixed Trellis {name} must be contiguous int32 on {x.device} " + f"with {expected_entries} elements" + ) + _check_descriptor_projection_counts( + descriptor_map, + total_experts, + gate_counts=gate_counts, + up_counts=up_counts, + ) + for name, table, expected_elements in ( + ( + "intermediate rotations", + rotations.intermediate, + total_experts * 3 * launch.intermediate_size, + ), + ( + "gate SUH", + rotations.gate_suh, + (1 if launch.broadcast_suh else total_experts) * launch.hidden_size, + ), + ( + "up SUH", + rotations.up_suh, + (1 if launch.broadcast_suh else total_experts) * launch.hidden_size, + ), + ( + "down SVH", + rotations.down_svh, + (1 if launch.broadcast_svh else total_experts) * launch.hidden_size, + ), + ): + if ( + table.dtype != torch.float16 + or table.device != x.device + or not table.is_contiguous() + or int(table.numel()) != expected_elements + or int(table.data_ptr()) % 16 != 0 + ): + raise ValueError( + f"mixed Trellis {name} must be contiguous fp16 on {x.device} " + f"with {expected_elements} elements and at least 16-byte alignment" + ) + + required_route_slots = max_packed_route_slots( + m * launch.top_k, launch.moe_block_size, route_num_experts + ) + required_route_blocks = ( + required_route_slots + launch.moe_block_size - 1 + ) // launch.moe_block_size + if required_route_blocks > launch.max_m_blocks: + raise RuntimeError( + "mixed Trellis request requires " + f"{required_route_blocks} route blocks, but the launch supports " + f"{launch.max_m_blocks}" + ) + if buffers.packed_route_indices.numel() < required_route_slots: + raise RuntimeError( + "mixed Trellis packed-route buffer is below request capacity" + ) + if buffers.block_expert_ids.numel() < required_route_blocks: + raise RuntimeError( + "mixed Trellis block-expert buffer is below request capacity" + ) + packed, block_experts, packed_count = pack_topk_routes_by_expert( + topk_ids, + launch.moe_block_size, + route_num_experts, + expert_map=global_to_combined, + packed_route_indices=buffers.packed_route_indices, + block_expert_ids=buffers.block_expert_ids, + packed_route_count=buffers.packed_route_count, + expert_offsets=buffers.expert_offsets, + expert_counts=buffers.expert_counts, + ) + stream = current_cuda_stream() + trellis_rank_lut = sqg_xor_cheb_t12_lut(x.device) + + def tier_pointers(tier): + return ( + make_ptr( + cutlass.Int32, + tier.w13.data_ptr(), + cute.AddressSpace.gmem, + assumed_align=16, + ), + make_ptr( + cutlass.Int32, + tier.w2.data_ptr(), + cute.AddressSpace.gmem, + assumed_align=16, + ), + make_ptr( + cutlass.Int32, + tier.w13_scale.data_ptr(), + cute.AddressSpace.gmem, + assumed_align=16, + ), + make_ptr( + cutlass.Int32, + tier.w2_scale.data_ptr(), + cute.AddressSpace.gmem, + assumed_align=16, + ), + make_ptr( + cutlass.Float32, + tier.w13_global_scale.data_ptr(), + cute.AddressSpace.gmem, + assumed_align=16, + ), + make_ptr( + cutlass.Float32, + tier.w2_global_scale.data_ptr(), + cute.AddressSpace.gmem, + assumed_align=16, + ), + ) + + launch.compiled( + make_ptr( + _cutlass_element_dtype(launch.rotation_input_dtype), + x.data_ptr(), + cute.AddressSpace.gmem, + assumed_align=16, + ), + buffers.rotation_gate.view(-1), + buffers.rotation_up.view(-1), + *tier_pointers(tier0), + *tier_pointers(tier1), + *tier_pointers(tier2), + buffers.fc1.view(-1), + buffers.activated.view(-1), + buffers.fc2.view(-1), + packed, + block_experts, + packed_count, + make_ptr( + cutlass.Int32, + descriptor_map.data_ptr(), + cute.AddressSpace.gmem, + assumed_align=4, + ), + make_ptr( + cutlass.Float32, + topk_weights.data_ptr(), + cute.AddressSpace.gmem, + assumed_align=4, + ), + buffers.fc1_scratch, + buffers.fc2_scratch, + buffers.workspace, + make_ptr( + cutlass.Float16, + rotations.intermediate.data_ptr(), + cute.AddressSpace.gmem, + assumed_align=16, + ), + make_ptr( + cutlass.Float16, + rotations.gate_suh.data_ptr(), + cute.AddressSpace.gmem, + assumed_align=16, + ), + make_ptr( + cutlass.Float16, + rotations.up_suh.data_ptr(), + cute.AddressSpace.gmem, + assumed_align=16, + ), + make_ptr( + cutlass.Uint8, + trellis_rank_lut.data_ptr(), + cute.AddressSpace.gmem, + assumed_align=16, + ), + Int32(counts[0]), + Int32(counts[1]), + Int32(counts[2]), + Int32(int(tier0.w2_global_scale.numel())), + Int32(int(tier1.w2_global_scale.numel())), + Int32(int(tier2.w2_global_scale.numel())), m, max(int(launch.blocks_per_sm) * int(launch.sms), 1), stream, + tier0_gate_experts=Int32(gate_counts[0]), + tier1_gate_experts=Int32(gate_counts[1]), + tier2_gate_experts=Int32(gate_counts[2]), + tier0_up_experts=Int32(up_counts[0]), + tier1_up_experts=Int32(up_counts[1]), + tier2_up_experts=Int32(up_counts[2]), ) launch.topk_sum.compiled( make_ptr( @@ -1517,11 +3391,18 @@ def run_mixed_trellis( __all__ = [ "MixedTrellisBuffers", "MixedTrellisCompileResult", + "MixedTrellis3CompileResult", "MixedTrellisRotations", + "W4A16MixedTrellis3Kernel", "build_ordered_maps", + "build_projection_tiered_maps", "build_tiered_maps", "combine_trellis_rotations", "compile_mixed_trellis", + "compile_mixed_trellis3", "make_mixed_trellis_buffers", + "make_mixed_trellis3_buffers", "run_mixed_trellis", + "warmup_mixed_trellis_route_pack", + "run_mixed_trellis3", ] diff --git a/b12x/moe/_shared/kernels/w4a16/prepare.py b/b12x/moe/_shared/kernels/w4a16/prepare.py index f11afe9c4..8bd20ecf8 100644 --- a/b12x/moe/_shared/kernels/w4a16/prepare.py +++ b/b12x/moe/_shared/kernels/w4a16/prepare.py @@ -7,12 +7,7 @@ import torch -from b12x.moe._shared.trellis_codebooks import ( - MCG, - SQG_E4M3, - normalize_codebook, - validate_codebook_bits, -) +from b12x._lib.quant.sqg_fp16_d3l import SQG_FP16_D3L from b12x.moe._shared.kernels.w4a16.host import ( W4A16PackedBuffers, @@ -34,6 +29,28 @@ } _E8M0_K32_BF16_MAX_SCALE_BYTE = 247 _E8M0_LOGICAL_TAIL_SCALE_N_ALIGNMENT = 64 +_QSRT_ATOM_CHANNELS = 32 +_QSRT_ATOMS_PER_PAIR = 8 +_QSRT_ATOMS_PER_EXPERT = 96 +_QSRT_MATRIX_ATOM_TRELLIS_BYTES = 43_008 +_QSRT_MATRIX_ATOM_SCALE_BYTES = 64 +_QSRT_ATOM_TRELLIS_BYTES = 129_024 +_QSRT_ATOM_BUNDLE_BYTES = 129_216 +_QSRT_MATRIX_TRELLIS_OFFSETS = (0, 43_008, 86_016) +_QSRT_MATRIX_SCALE_OFFSETS = (129_024, 129_088, 129_152) +_QSRT_EXPERT_ROTATION_MULTIPLIER = 5 +_QSRT_V2_P33_MATRIX_TRELLIS_BYTES = 43_008 +_QSRT_V2_P43_MATRIX_TRELLIS_BYTES = 50_176 +_QSRT_V2_P44_MATRIX_TRELLIS_BYTES = 57_344 +_QSRT_V2_P33_ATOM_BUNDLE_BYTES = 129_216 +_QSRT_V2_P43_ATOM_BUNDLE_BYTES = 150_720 +_QSRT_V2_COUPLED_H308_P43_P33_ATOM_BUNDLE_BYTES = 143_552 +_QSRT_V2_COUPLED_H308_P43_P44_ATOM_BUNDLE_BYTES = 157_888 +_QSRT_V2_P22_MATRIX_TRELLIS_BYTES = 28_672 +_QSRT_V2_P22_ATOM_BUNDLE_BYTES = 86_208 +_QSRT_V2_PROFILE_H308 = "k3x22_k4x2" +_QSRT_V2_PROFILE_COUPLED_K2 = "k2_coupled_h512_h128" +_QSRT_V2_PROFILE_COUPLED_H308 = "k3x22_k4x2_coupled_h512_h128" # Canonical W13 layout names are "w13"/"w31"; accept the physical FC1-half # spellings as aliases. Logical checkpoint order "w13" arrives up/gate and # needs a swap before the kernel's SwiGLU; "w31" is already kernel-native @@ -117,50 +134,13 @@ class W4A16FC2Weights: scale_format: str = "e8m0_k32" -@dataclass(frozen=True) -class TrellisWeightState: - """Trellis-specific prepared-weight state. - - ``codebook``/``bits`` select the decode law of the native tiles. The - optional pair specialization keeps every matrix at a fixed payload - size: FC1's pair lies on N (one pair in each projection) and FC2's on - K, selected separately per matrix. The fixed-size K3/K3 versus K2/K4 - representation uses one int32 mode per local expert (zero selects - K3/K3); the K3/K3 versus K4/K3 representation uses one int64 - descriptor per local expert whose bit zero selects K4/K3 and whose - remaining bits store the exact uint32 offset into that matrix's - compact payload pool. - - ``gate_suh``/``up_suh``/``down_svh`` are the projection-specific input - incoherence values; ``intermediate_rotations`` holds the - intermediate-boundary values, and coupled transforms append the - activation-boundary sign rows so each rotation row reads - ``[gate I, up I, down I, preactivation 2I, postactivation I]``. - """ - - codebook: str - bits: int - fc1_pair_kind: str | None = None - fc2_pair_kind: str | None = None - fc1_pair_modes: torch.Tensor | None = None - fc2_pair_modes: torch.Tensor | None = None - gate_suh: torch.Tensor | None = None - up_suh: torch.Tensor | None = None - intermediate_rotations: torch.Tensor | None = None - down_svh: torch.Tensor | None = None - coupled_hadamard: bool = False - tile_config: tuple[int, int, int, int] | None = None - - @dataclass(frozen=True) class PreparedW4A16MoeWeights: """Runtime native-codebook W4A16 expert weights. - Trellis-coded sources (MCG, SQG-XOR-Cheb-T12, or FP16-D3L codebooks) keep - native codebook tiles and persistent full-rotation tables, sharing the - W4A16 host ABI and retaining the tile configuration used at preparation. - The trellis annex lives in one :class:`TrellisWeightState`; the accessor - properties expose its members under the host-ABI names. + EXL3 Trellis keeps native codebook tiles and persistent full-rotation + tables, sharing the W4A16 host ABI and retaining the tile configuration + used at preparation. """ w13: torch.Tensor @@ -177,63 +157,39 @@ class PreparedW4A16MoeWeights: params_dtype: torch.dtype fc1_tile_n: int fc2_tile_n: int - source_format: str = "btx" + source_format: str = "qsrt_sqg_e4m3" w13_layout: str = "packed" - weight_layout: str = "trellis_t256" + weight_layout: str = "trellis3_t256" scale_format: str = "e4m3_k32" - trellis: TrellisWeightState | None = None - - @property - def trellis_codebook(self) -> str | None: - return None if self.trellis is None else self.trellis.codebook - - @property - def trellis_bits(self) -> int: - return 3 if self.trellis is None else self.trellis.bits - - @property - def fc1_trellis_pair_kind(self) -> str | None: - return None if self.trellis is None else self.trellis.fc1_pair_kind - - @property - def fc2_trellis_pair_kind(self) -> str | None: - return None if self.trellis is None else self.trellis.fc2_pair_kind - - @property - def fc1_trellis_pair_modes(self) -> torch.Tensor | None: - return None if self.trellis is None else self.trellis.fc1_pair_modes - - @property - def fc2_trellis_pair_modes(self) -> torch.Tensor | None: - return None if self.trellis is None else self.trellis.fc2_pair_modes - - @property - def gate_suh(self) -> torch.Tensor | None: - return None if self.trellis is None else self.trellis.gate_suh - - @property - def up_suh(self) -> torch.Tensor | None: - return None if self.trellis is None else self.trellis.up_suh - - @property - def intermediate_rotations(self) -> torch.Tensor | None: - return ( - None - if self.trellis is None - else self.trellis.intermediate_rotations - ) - - @property - def down_svh(self) -> torch.Tensor | None: - return None if self.trellis is None else self.trellis.down_svh - - @property - def coupled_hadamard(self) -> bool: - return self.trellis is not None and self.trellis.coupled_hadamard - - @property - def tile_config(self) -> tuple[int, int, int, int] | None: - return None if self.trellis is None else self.trellis.tile_config + # Native tiles use MCG, K2/K3/K4 SQG-XOR-Cheb-T12, or K5/K6 FP16-D3L. + trellis_codebook: str | None = None + trellis_bits: int = 3 + # Optional compact fixed-payload pair specialization. FC1's pair lies on + # N (one pair in each projection); FC2's lies on K. The two kinds may be + # selected separately while every matrix remains exactly three bpw. + fc1_trellis_pair_kind: str | None = None + fc2_trellis_pair_kind: str | None = None + # The fixed-size K3/K3 versus K2/K4 representation uses one int32 mode per + # local expert: zero selects K3/K3 and one selects K2/K4. The K3/K3 versus + # K4/K3 representation uses one int64 descriptor per local expert. Bit + # zero selects K4/K3 and the remaining bits store the exact uint32 offset + # into that matrix's compact payload pool. + fc1_trellis_pair_modes: torch.Tensor | None = None + fc2_trellis_pair_modes: torch.Tensor | None = None + # Projection-specific EXL3 input incoherence scales. They remain optional + # for non-trellis and synthetic oracle preparation, but the coherent + # projection-major runtime binds both so gate/up can stage distinct rotated + # A operands without copying either table. + gate_suh: torch.Tensor | None = None + up_suh: torch.Tensor | None = None + intermediate_rotations: torch.Tensor | None = None + down_svh: torch.Tensor | None = None + # Coupled transforms add an exact activation-boundary reparameterization + # around the per-matrix incoherence transforms. The rotation rows contain + # [gate I, up I, down I, preactivation 2I, postactivation I]. The 512-wide + # residual transforms use a fixed all-positive draw and need no table. + coupled_hadamard: bool = False + tile_config: tuple[int, int, int, int] | None = None @dataclass(frozen=True) @@ -259,7 +215,7 @@ class PreparedTrellis256DenseWeight: mcg: torch.Tensor | None = None mul1_e4m3: torch.Tensor | None = None num_experts: int = 1 - weight_layout: str = "trellis_t256" + weight_layout: str = "trellis3_t256" scale_format: str = "e4m3_k32" w13_layout: str = "packed" # Optional fixed-size 256-channel pair container. The payload still averages @@ -1546,9 +1502,43 @@ def make_w4a16_packed_buffers( ) -_TRELLIS256_W13_LAYOUTS = {"packed", "trellis_t256_proj"} +_TRELLIS256_W13_LAYOUTS = {"packed", "trellis3_t256_proj"} +_TRELLIS256_CODEBOOKS = { + "mcg": "mcg", + "sqg_xor_cheb_t12": "sqg_xor_cheb_t12", + SQG_FP16_D3L: SQG_FP16_D3L, +} +_TRELLIS256_CODEBOOK_SENTINELS = { + 0xCBAC1FED: "mcg", +} + + def _normalize_trellis256_codebook(codebook: str | int) -> str: - return normalize_codebook(codebook) + if isinstance(codebook, int): + normalized = _TRELLIS256_CODEBOOK_SENTINELS.get( + int(codebook) & 0xFFFFFFFF + ) + if normalized is None: + raise ValueError( + "unsupported trellis256 codebook sentinel " + f"{int(codebook) & 0xFFFFFFFF:#010x}; expected MCG " + "0xcbac1fed" + ) + return normalized + normalized = _TRELLIS256_CODEBOOKS.get(str(codebook).strip().lower()) + if normalized is None: + raise ValueError( + f"unsupported trellis256 codebook {codebook!r}; expected " + "'mcg', 'sqg_xor_cheb_t12', or 'sqg_fp16_d3l'" + ) + return normalized + + +def _validate_trellis256_codebook_bits(codebook: str, bits: int) -> None: + if codebook == "sqg_xor_cheb_t12" and bits not in (2, 3, 4): + raise ValueError("sqg_xor_cheb_t12 is defined only for K2/K3/K4") + if codebook == SQG_FP16_D3L and bits not in (5, 6): + raise ValueError("sqg_fp16_d3l is defined only for uniform K5/K6") def _trellis256_random_native_tensor( @@ -1571,25 +1561,25 @@ def _trellis256_random_native_tensor( def _trellis256_bits_from_native_tensor(tensor: torch.Tensor, *, name: str) -> int: - """Recover the trellis bitrate from the native tile's final dimension.""" + """Recover the EXL3 bitrate from the native tile's final dimension.""" if tensor.ndim < 1: - raise ValueError(f"trellis_t256 {name} must have at least one dimension") + raise ValueError(f"trellis3_t256 {name} must have at least one dimension") words_per_bit = 16 if tensor.dtype == torch.int16 else 8 if tensor.dtype not in (torch.int16, torch.int32): raise TypeError( - f"trellis_t256 {name} must use native int16 or int32 storage, " + f"trellis3_t256 {name} must use native int16 or int32 storage, " f"got {tensor.dtype}" ) last = int(tensor.shape[-1]) if last % words_per_bit != 0: raise ValueError( - f"trellis_t256 {name} final dimension {last} is not a native " + f"trellis3_t256 {name} final dimension {last} is not a native " f"{tensor.dtype} tile width" ) bits = last // words_per_bit if bits not in (2, 3, 4, 5, 6): raise ValueError( - f"trellis_t256 {name} encodes unsupported {bits}-bpw storage; " + f"trellis3_t256 {name} encodes unsupported {bits}-bpw storage; " "expected 2, 3, 4, 5, or 6" ) return bits @@ -1608,7 +1598,7 @@ def _trellis256_flat_native_view( expected_i32_shape = (*expected_prefix_shape, 8 * trellis_bits) if tensor.device != device: raise ValueError( - f"trellis_t256 {name} must be on {device}, got {tensor.device}" + f"trellis3_t256 {name} must be on {device}, got {tensor.device}" ) if tensor.dtype == torch.int16: expected_shape = expected_i16_shape @@ -1616,19 +1606,19 @@ def _trellis256_flat_native_view( expected_shape = expected_i32_shape else: raise TypeError( - f"trellis_t256 {name} must be native int16 tiles ({16 * trellis_bits} words) " + f"trellis3_t256 {name} must be native int16 tiles ({16 * trellis_bits} words) " f"or the identical int32 view ({8 * trellis_bits} words), got {tensor.dtype}" ) if tuple(tensor.shape) != expected_shape: raise ValueError( - f"trellis_t256 {name} requires native {trellis_bits}-bit tile shape " + f"trellis3_t256 {name} requires native {trellis_bits}-bit EXL3 shape " f"{expected_shape} for dtype {tensor.dtype}, got {tuple(tensor.shape)}" ) if not tensor.is_contiguous(): - raise ValueError(f"trellis_t256 {name} must be contiguous") + raise ValueError(f"trellis3_t256 {name} must be contiguous") if int(tensor.data_ptr()) % 16 != 0: raise ValueError( - f"trellis_t256 {name} must be at least 16-byte aligned for cp.async" + f"trellis3_t256 {name} must be at least 16-byte aligned for cp.async" ) return tensor.view(torch.int32).reshape(-1) @@ -1657,18 +1647,18 @@ def prepare_trellis256_moe_weights( tile_config: tuple[int, int, int, int] | None = None, workspace: torch.Tensor | None = None, ) -> PreparedW4A16MoeWeights: - """Wrap or synthesize native ``trellis_t256`` tiles for any codebook. + """Wrap or synthesize native EXL3 tiles for ``trellis3_t256``. Supplying both ``w13`` and ``w2`` is the production path: no bytes are copied or permuted; each tensor is only viewed as contiguous int32 words and flattened. Omitting both tensors is the deterministic full-GEMM-oracle path selected by ``device`` and ``seed``. ``gate_suh`` and ``up_suh`` are - optional zero-copy bindings for the two projection-specific input + optional zero-copy bindings for the two projection-specific EXL3 input rotations; when supplied, both must be contiguous fp16 ``[E,H]`` tensors on the weight device. Plain FC1 storage is expert-major - ``[E,H/16,FC1_N/16,16*bits]i16``. ``trellis_t256_proj`` instead requires one + ``[E,H/16,FC1_N/16,16*bits]i16``. ``trellis3_t256_proj`` instead requires one projection-major backing ``[2,E,H/16,I/16,16*bits]i16`` so gate/up fallback views can continue to alias the same live storage. FC2 is always the plain native ``[E,I/16,H/16,16*bits]i16`` layout. @@ -1679,7 +1669,7 @@ def prepare_trellis256_moe_weights( fc1_tile_n = int(fc1_tile_n) fc2_tile_n = int(fc2_tile_n) if params_dtype not in (torch.bfloat16, torch.float16): - raise ValueError("trellis_t256 W4A16 weights require fp16 or bf16 activations") + raise ValueError("trellis3_t256 W4A16 weights require fp16 or bf16 activations") requested_trellis_bits = None if trellis_bits is None else int(trellis_bits) if requested_trellis_bits is not None and requested_trellis_bits not in ( 2, @@ -1689,31 +1679,31 @@ def prepare_trellis256_moe_weights( 6, ): raise ValueError( - "trellis_t256 bits must be one of 2, 3, 4, 5, or 6, " + "trellis3_t256 bits must be one of 2, 3, 4, 5, or 6, " f"got {requested_trellis_bits}" ) if num_experts <= 0: - raise ValueError(f"trellis_t256 requires num_experts > 0, got {num_experts}") + raise ValueError(f"trellis3_t256 requires num_experts > 0, got {num_experts}") if hidden_size <= 0 or intermediate_size <= 0: raise ValueError( - "trellis_t256 requires positive hidden_size and intermediate_size, " + "trellis3_t256 requires positive hidden_size and intermediate_size, " f"got H={hidden_size} I={intermediate_size}" ) if hidden_size % 16 != 0 or intermediate_size % 16 != 0: raise ValueError( - "native trellis tiles require hidden_size and intermediate_size to be " + "native EXL3 tiles require hidden_size and intermediate_size to be " f"multiples of 16, got H={hidden_size} I={intermediate_size}" ) if hidden_size % 32 != 0 or intermediate_size % 32 != 0: raise ValueError( - "trellis_t256 uses E4M3 K/32 kernel plumbing and therefore requires " + "trellis3_t256 uses E4M3 K/32 kernel plumbing and therefore requires " "hidden_size and intermediate_size to be multiples of 32; " f"got H={hidden_size} I={intermediate_size}" ) for name, tile_n in (("fc1_tile_n", fc1_tile_n), ("fc2_tile_n", fc2_tile_n)): if tile_n < 64 or tile_n % 16 != 0: raise ValueError( - f"trellis_t256 {name} must be a multiple of 16 and at least " + f"trellis3_t256 {name} must be a multiple of 16 and at least " f"64 for the current W4A16 kernel, got {tile_n}" ) @@ -1721,29 +1711,29 @@ def prepare_trellis256_moe_weights( w13_rows = intermediate_size * (2 if is_gated else 1) if w13_layout not in _TRELLIS256_W13_LAYOUTS: raise ValueError( - f"unsupported trellis_t256 w13_layout {w13_layout!r}; expected " - "'packed' or 'trellis_t256_proj'" + f"unsupported trellis3_t256 w13_layout {w13_layout!r}; expected " + "'packed' or 'trellis3_t256_proj'" ) - if w13_layout == "trellis_t256_proj": + if w13_layout == "trellis3_t256_proj": if not is_gated: raise ValueError( - "trellis_t256_proj requires a gated activation with separate " + "trellis3_t256_proj requires a gated activation with separate " "gate/up FC1 projections" ) if intermediate_size % fc1_tile_n != 0: raise ValueError( - "trellis_t256_proj requires each FC1 projection to contain an " + "trellis3_t256_proj requires each FC1 projection to contain an " f"integral number of CTA N tiles: I={intermediate_size}, " f"fc1_tile_n={fc1_tile_n}" ) elif w13_rows % fc1_tile_n != 0: raise ValueError( - "trellis_t256 has no FC1 logical-tail path: " + "trellis3_t256 has no FC1 logical-tail path: " f"FC1_N={w13_rows} must be divisible by fc1_tile_n={fc1_tile_n}" ) if hidden_size % fc2_tile_n != 0: raise ValueError( - "trellis_t256 has no FC2 logical-tail path: " + "trellis3_t256 has no FC2 logical-tail path: " f"FC2_N={hidden_size} must be divisible by fc2_tile_n={fc2_tile_n}" ) normalized_codebook = _normalize_trellis256_codebook(codebook) @@ -1751,13 +1741,13 @@ def prepare_trellis256_moe_weights( have_w13 = w13 is not None have_w2 = w2 is not None if have_w13 != have_w2: - raise ValueError("trellis_t256 requires both w13 and w2, or neither") + raise ValueError("trellis3_t256 requires both w13 and w2, or neither") synthetic = not have_w13 if synthetic: resolved_trellis_bits = requested_trellis_bits or 3 if device is None: raise ValueError( - "device is required when synthesizing trellis_t256 oracle weights" + "device is required when synthesizing trellis3_t256 oracle weights" ) if isinstance(device, int): resolved_device = torch.device("cuda", int(device)) @@ -1765,14 +1755,14 @@ def prepare_trellis256_moe_weights( resolved_device = torch.device(device) if resolved_device.type != "cuda": raise ValueError( - "trellis_t256 W4A16 weights require a CUDA device, got " + "trellis3_t256 W4A16 weights require a CUDA device, got " f"{resolved_device}" ) if resolved_device.index is None: resolved_device = torch.device("cuda", torch.cuda.current_device()) generator = torch.Generator(device=resolved_device) generator.manual_seed(int(seed)) - if w13_layout == "trellis_t256_proj": + if w13_layout == "trellis3_t256_proj": w13_i32_shape = ( 2, num_experts, @@ -1805,12 +1795,12 @@ def prepare_trellis256_moe_weights( w2_bits = _trellis256_bits_from_native_tensor(w2, name="w2") if w13_bits != w2_bits: raise ValueError( - f"trellis_t256 w13/w2 bitrate mismatch: {w13_bits} vs {w2_bits}" + f"trellis3_t256 w13/w2 bitrate mismatch: {w13_bits} vs {w2_bits}" ) resolved_trellis_bits = w13_bits if resolved_trellis_bits not in (2, 3, 4, 5, 6): raise ValueError( - "trellis_t256 tensors must encode 2, 3, 4, 5, or " + "trellis3_t256 tensors must encode 2, 3, 4, 5, or " f"6 bpw, got {resolved_trellis_bits}" ) if ( @@ -1824,7 +1814,7 @@ def prepare_trellis256_moe_weights( resolved_device = w13.device if resolved_device.type != "cuda": raise ValueError( - "trellis_t256 W4A16 weights require CUDA storage, got " + "trellis3_t256 W4A16 weights require CUDA storage, got " f"{resolved_device}" ) if device is not None: @@ -1839,10 +1829,10 @@ def prepare_trellis256_moe_weights( ) if not device_matches: raise ValueError( - "explicit trellis_t256 device does not match supplied weights: " + "explicit trellis3_t256 device does not match supplied weights: " f"device={requested_device}, weights={resolved_device}" ) - if w13_layout == "trellis_t256_proj": + if w13_layout == "trellis3_t256_proj": expected_w13_prefix = ( 2, num_experts, @@ -1879,23 +1869,23 @@ def prepare_trellis256_moe_weights( have_up_suh = up_suh is not None if have_gate_suh != have_up_suh: raise ValueError( - "trellis_t256 projection input scales require both gate_suh and up_suh" + "trellis3_t256 projection input scales require both gate_suh and up_suh" ) if have_gate_suh: - if w13_layout != "trellis_t256_proj": + if w13_layout != "trellis3_t256_proj": raise ValueError( - "trellis_t256 gate_suh/up_suh bindings require " - "w13_layout='trellis_t256_proj'" + "trellis3_t256 gate_suh/up_suh bindings require " + "w13_layout='trellis3_t256_proj'" ) assert gate_suh is not None and up_suh is not None for name, scale in (("gate_suh", gate_suh), ("up_suh", up_suh)): if scale.device != resolved_device: raise ValueError( - f"trellis_t256 {name} must be on {resolved_device}, got {scale.device}" + f"trellis3_t256 {name} must be on {resolved_device}, got {scale.device}" ) if scale.dtype != torch.float16: raise TypeError( - f"trellis_t256 {name} must be torch.float16, got {scale.dtype}" + f"trellis3_t256 {name} must be torch.float16, got {scale.dtype}" ) # (1, hidden_size) is a broadcast row shared by all experts # (kquant shared-su artifacts); kernels index it with expert @@ -1905,15 +1895,15 @@ def prepare_trellis256_moe_weights( (1, hidden_size), ): raise ValueError( - f"trellis_t256 {name} must have shape " + f"trellis3_t256 {name} must have shape " f"{(num_experts, hidden_size)} or {(1, hidden_size)}, " f"got {tuple(scale.shape)}" ) if not scale.is_contiguous(): - raise ValueError(f"trellis_t256 {name} must be contiguous") + raise ValueError(f"trellis3_t256 {name} must be contiguous") if (gate_suh.shape[0] == 1) != (up_suh.shape[0] == 1): raise ValueError( - "trellis_t256 gate_suh and up_suh must both be per-expert " + "trellis3_t256 gate_suh and up_suh must both be per-expert " "or both broadcast" ) @@ -1926,7 +1916,7 @@ def prepare_trellis256_moe_weights( for value in (gate_suh, up_suh, intermediate_rotations, down_svh) ): raise ValueError( - "full-rotation trellis_t256 requires gate_suh, up_suh, " + "full-rotation trellis3_t256 requires gate_suh, up_suh, " "intermediate_rotations, and down_svh" ) if have_full_rotation: @@ -1949,19 +1939,19 @@ def prepare_trellis256_moe_weights( ): if scale.device != resolved_device: raise ValueError( - f"trellis_t256 {name} must be on {resolved_device}, got {scale.device}" + f"trellis3_t256 {name} must be on {resolved_device}, got {scale.device}" ) if scale.dtype != torch.float16: raise TypeError( - f"trellis_t256 {name} must be torch.float16, got {scale.dtype}" + f"trellis3_t256 {name} must be torch.float16, got {scale.dtype}" ) if tuple(scale.shape) not in shapes: raise ValueError( - f"trellis_t256 {name} must have shape {shapes}, " + f"trellis3_t256 {name} must have shape {shapes}, " f"got {tuple(scale.shape)}" ) if not scale.is_contiguous(): - raise ValueError(f"trellis_t256 {name} must be contiguous") + raise ValueError(f"trellis3_t256 {name} must be contiguous") if tile_config is not None: tile_config = tuple(int(value) for value in tile_config) @@ -1969,41 +1959,43 @@ def prepare_trellis256_moe_weights( value <= 0 or value % 16 != 0 for value in tile_config ): raise ValueError( - "trellis_t256 tile_config must contain four positive multiples of 16" + "trellis3_t256 tile_config must contain four positive multiples of 16" ) fc1_tile_k, configured_fc1_n, fc2_tile_k, configured_fc2_n = tile_config if configured_fc1_n != fc1_tile_n or configured_fc2_n != fc2_tile_n: raise ValueError( - "trellis_t256 tile_config N dimensions disagree with preparation: " + "trellis3_t256 tile_config N dimensions disagree with preparation: " f"tile_config={tile_config}, fc1_tile_n={fc1_tile_n}, " f"fc2_tile_n={fc2_tile_n}" ) if hidden_size % fc1_tile_k != 0 or intermediate_size % fc2_tile_k != 0: raise ValueError( - "trellis_t256 tile_config K dimensions must divide the model geometry" + "trellis3_t256 tile_config K dimensions must divide the model geometry" ) - validate_codebook_bits(normalized_codebook, resolved_trellis_bits) + _validate_trellis256_codebook_bits( + normalized_codebook, resolved_trellis_bits + ) if dummy_scale is None: dummy_scale = torch.zeros(4, dtype=torch.uint8, device=resolved_device) else: if dummy_scale.device != resolved_device: raise ValueError( - "trellis_t256 dummy_scale must share the weight device, got " + "trellis3_t256 dummy_scale must share the weight device, got " f"{dummy_scale.device} and {resolved_device}" ) if dummy_scale.dtype != torch.uint8: raise TypeError( - f"trellis_t256 dummy_scale must be torch.uint8, got {dummy_scale.dtype}" + f"trellis3_t256 dummy_scale must be torch.uint8, got {dummy_scale.dtype}" ) if not dummy_scale.is_contiguous() or tuple(dummy_scale.shape) != (4,): raise ValueError( - "trellis_t256 dummy_scale must be a contiguous four-byte " + "trellis3_t256 dummy_scale must be a contiguous four-byte " f"tensor with shape (4,), got {tuple(dummy_scale.shape)}" ) if int(dummy_scale.data_ptr()) % 16 != 0: raise ValueError( - "trellis_t256 dummy_scale must be at least 16-byte aligned" + "trellis3_t256 dummy_scale must be at least 16-byte aligned" ) global_scale = torch.ones( @@ -2012,7 +2004,7 @@ def prepare_trellis256_moe_weights( if workspace is None: workspace = _make_workspace(resolved_device, max_blocks_per_sm=4) elif workspace.device != resolved_device or workspace.dtype != torch.int32: - raise ValueError("trellis_t256 workspace must be int32 on the weight device") + raise ValueError("trellis3_t256 workspace must be int32 on the weight device") return PreparedW4A16MoeWeights( w13=packed_w13, w13_scale=dummy_scale, @@ -2028,19 +2020,23 @@ def prepare_trellis256_moe_weights( params_dtype=params_dtype, fc1_tile_n=fc1_tile_n, fc2_tile_n=fc2_tile_n, - source_format="btx", + source_format=( + "exl3_trellis_mcg" + if normalized_codebook == "mcg" + else "qsrt_sqg_e4m3" + if normalized_codebook == "sqg_xor_cheb_t12" + else SQG_FP16_D3L + ), w13_layout=w13_layout, - weight_layout="trellis_t256", + weight_layout="trellis3_t256", scale_format="e4m3_k32", - trellis=TrellisWeightState( - codebook=normalized_codebook, - bits=resolved_trellis_bits, - gate_suh=gate_suh, - up_suh=up_suh, - intermediate_rotations=intermediate_rotations, - down_svh=down_svh, - tile_config=tile_config, - ), + trellis_codebook=normalized_codebook, + trellis_bits=resolved_trellis_bits, + gate_suh=gate_suh, + up_suh=up_suh, + intermediate_rotations=intermediate_rotations, + down_svh=down_svh, + tile_config=tile_config, ) @@ -2095,7 +2091,7 @@ def prepare_trellis256_dense_weight( params_dtype: torch.dtype = torch.float16, dummy_scale: torch.Tensor | None = None, ) -> PreparedTrellis256DenseWeight: - """Prepare one native trellis linear for the dense trellis256 entry point. + """Prepare one native EXL3 linear for the dense trellis256 entry point. The native payload is ``[K/16,N/16,16*bits]i16`` (or the byte-identical ``[...,8*bits]i32`` view), optionally with a leading singleton expert axis. @@ -2103,15 +2099,15 @@ def prepare_trellis256_dense_weight( bytes are copied, permuted, stacked, or concatenated. """ if params_dtype not in (torch.float16, torch.bfloat16): - raise ValueError("trellis_t256 dense compute requires fp16 or bf16 MMA inputs") + raise ValueError("trellis3_t256 dense compute requires fp16 or bf16 MMA inputs") if trellis.ndim not in (3, 4): raise ValueError( - "trellis_t256 dense payload must have rank 3 or a leading E=1 axis" + "trellis3_t256 dense payload must have rank 3 or a leading E=1 axis" ) if trellis.ndim == 4: if int(trellis.shape[0]) != 1: raise ValueError( - f"trellis_t256 dense payload requires E=1, got {int(trellis.shape[0])}" + f"trellis3_t256 dense payload requires E=1, got {int(trellis.shape[0])}" ) k16, n16 = int(trellis.shape[1]), int(trellis.shape[2]) else: @@ -2121,18 +2117,18 @@ def prepare_trellis256_dense_weight( out_features = n16 * 16 if in_features <= 0 or out_features <= 0: raise ValueError( - f"trellis_t256 dense dimensions must be positive, got {in_features}x{out_features}" + f"trellis3_t256 dense dimensions must be positive, got {in_features}x{out_features}" ) if in_features % 128 != 0 or out_features % 128 != 0: raise ValueError( - "trellis_t256 dense rotations require K and N divisible by 128; " + "trellis3_t256 dense rotations require K and N divisible by 128; " f"got K={in_features} N={out_features}" ) expected_prefix_shape = (1, k16, n16) if trellis.ndim == 4 else (k16, n16) device = trellis.device if device.type != "cuda": raise ValueError( - f"trellis_t256 dense weights require CUDA storage, got {device}" + f"trellis3_t256 dense weights require CUDA storage, got {device}" ) packed = _trellis256_flat_native_view( trellis, @@ -2146,24 +2142,24 @@ def prepare_trellis256_dense_weight( ("svh", svh, out_features), ): if scale.device != device: - raise ValueError(f"trellis_t256 dense {name} must be on {device}") + raise ValueError(f"trellis3_t256 dense {name} must be on {device}") if scale.dtype != torch.float16: raise TypeError( - f"trellis_t256 dense {name} must be torch.float16, got {scale.dtype}" + f"trellis3_t256 dense {name} must be torch.float16, got {scale.dtype}" ) if tuple(scale.shape) != (width,): raise ValueError( - f"trellis_t256 dense {name} must have shape {(width,)}, " + f"trellis3_t256 dense {name} must have shape {(width,)}, " f"got {tuple(scale.shape)}" ) if not scale.is_contiguous(): - raise ValueError(f"trellis_t256 dense {name} must be contiguous") + raise ValueError(f"trellis3_t256 dense {name} must be contiguous") normalized_codebook = _trellis256_marker_codebook( mcg=mcg, mul1_e4m3=mul1_e4m3, codebook=codebook, ) - validate_codebook_bits(normalized_codebook, trellis_bits) + _validate_trellis256_codebook_bits(normalized_codebook, trellis_bits) if dummy_scale is None: dummy_scale = torch.zeros(4, dtype=torch.uint8, device=device) else: @@ -2175,7 +2171,7 @@ def prepare_trellis256_dense_weight( or int(dummy_scale.data_ptr()) % 16 != 0 ): raise ValueError( - "trellis_t256 dense dummy_scale must be a contiguous, 16-byte-" + "trellis3_t256 dense dummy_scale must be a contiguous, 16-byte-" "aligned four-byte uint8 tensor on the weight device" ) return PreparedTrellis256DenseWeight( @@ -2208,7 +2204,7 @@ def prepare_trellis256_pair_dense_weight( rate_axis: str, mcg: torch.Tensor | int | None = None, mul1_e4m3: torch.Tensor | int | None = None, - codebook: str | None = SQG_E4M3, + codebook: str | None = "sqg_xor_cheb_t12", params_dtype: torch.dtype = torch.float16, dummy_scale: torch.Tensor | None = None, ) -> PreparedTrellis256DenseWeight: @@ -2230,7 +2226,7 @@ def prepare_trellis256_pair_dense_weight( if rate_axis not in {"k", "n"}: raise ValueError(f"trellis pair rate_axis must be 'k' or 'n', got {rate_axis!r}") if params_dtype not in (torch.float16, torch.bfloat16): - raise ValueError("trellis_t256 pair compute requires fp16 or bf16 MMA inputs") + raise ValueError("trellis3_t256 pair compute requires fp16 or bf16 MMA inputs") if payload.dtype != torch.int16: raise TypeError(f"trellis pair payload must use torch.int16, got {payload.dtype}") if payload.ndim != 1 or not payload.is_contiguous(): @@ -2351,54 +2347,331 @@ def prepare_trellis256_pair_dense_weight( ) -def _restore_plane_words( - low: torch.Tensor, high: torch.Tensor, *, fc1: bool -) -> torch.Tensor: - """Assemble low/high record planes into the runtime pair word order. +def prepare_qsrt_pair_moe_weights( + w13_payload: torch.Tensor, + w2_payload: torch.Tensor, + *, + hidden_size: int, + intermediate_size: int, + num_experts: int, + activation: str, + fc1_pair_kind: torch.Tensor, + fc2_pair_kind: torch.Tensor, + gate_suh: torch.Tensor, + up_suh: torch.Tensor, + intermediate_rotations: torch.Tensor, + down_svh: torch.Tensor, + params_dtype: torch.dtype = torch.float16, + codebook: str = "sqg_xor_cheb_t12", + tile_config: tuple[int, int, int, int] = (64, 256, 64, 256), + dummy_scale: torch.Tensor | None = None, + workspace: torch.Tensor | None = None, +) -> PreparedW4A16MoeWeights: + """Prepare the compact P24/P33 reference view used by QSRT tests. - ``low``/``high`` are int16 ``[count, atoms, hidden_tiles, 16*bits]``. - FC1 places both 128-channel records under each K16 tile; FC2 keeps its - K-major low-plane/high-plane ordering. Returns ``[count, words]``. + Checkpoint loading uses :func:`prepare_qsrt_atom_moe_weights`. This + internal adapter exists only to compare that canonical atom path with the + former pair-shaped representation. It accepts only the active + SQG-XOR-Cheb-T12 codebook and is independent of legacy EXL3. """ - count = low.shape[0] - if fc1: - low = low.permute(0, 2, 1, 3).reshape(count, low.shape[2], -1) - high = high.permute(0, 2, 1, 3).reshape(count, high.shape[2], -1) - return torch.cat((low, high), dim=2).reshape(count, -1) - return torch.cat( - (low.reshape(count, -1), high.reshape(count, -1)), dim=1 + hidden_size = int(hidden_size) + intermediate_size = int(intermediate_size) + num_experts = int(num_experts) + if hidden_size <= 0 or hidden_size % 128: + raise ValueError( + "trellis pairs require hidden_size to be a positive multiple of 128" + ) + if intermediate_size != 256: + raise ValueError("trellis pairs require intermediate_size=256") + if num_experts <= 0 or not validate_activation(activation): + raise ValueError("trellis pairs require gated experts") + if params_dtype != torch.float16: + raise ValueError("trellis pairs require fp16 operands") + if _normalize_trellis256_codebook(codebook) != "sqg_xor_cheb_t12": + raise ValueError("W4A8 trellis pairs support only sqg_xor_cheb_t12") + if w13_payload.dtype != torch.int16 or w2_payload.dtype != torch.int16: + raise TypeError("trellis-pair payloads must use torch.int16") + if not w13_payload.is_cuda or not w2_payload.is_cuda: + raise ValueError("trellis-pair payloads require CUDA storage") + if w13_payload.device != w2_payload.device: + raise ValueError("w13_payload and w2_payload must share one CUDA device") + if not w13_payload.is_contiguous() or not w2_payload.is_contiguous(): + raise ValueError("trellis-pair payloads must be contiguous") + device = w13_payload.device + + def _modes(name: str, value: torch.Tensor) -> torch.Tensor: + if not isinstance(value, torch.Tensor): + raise TypeError(f"{name} must be a per-expert mode table tensor") + if value.device != device or tuple(value.shape) != (num_experts,): + raise ValueError( + f"{name} mode table must have shape {(num_experts,)} on {device}" + ) + if value.dtype not in { + torch.bool, + torch.uint8, + torch.int8, + torch.int16, + torch.int32, + torch.int64, + }: + raise TypeError(f"{name} mode table must use an integer or bool dtype") + result = value.to(dtype=torch.int32).contiguous() + if not bool(torch.all((result == 0) | (result == 1))): + raise ValueError(f"{name} mode table values must be 0=P33 or 1=P24") + return result + + fc1_modes = _modes("fc1_pair_kind", fc1_pair_kind) + fc2_modes = _modes("fc2_pair_kind", fc2_pair_kind) + hidden_tiles = hidden_size // 16 + pair_words = hidden_tiles * 8 * 16 * 6 + expected_w13 = (2, num_experts, pair_words) + expected_w2 = (num_experts, pair_words) + if tuple(w13_payload.shape) != expected_w13: + raise ValueError( + f"w13_payload must have shape {expected_w13}, got {tuple(w13_payload.shape)}" + ) + if tuple(w2_payload.shape) != expected_w2: + raise ValueError( + f"w2_payload must have shape {expected_w2}, got {tuple(w2_payload.shape)}" + ) + if int(w13_payload.data_ptr()) % 16 or int(w2_payload.data_ptr()) % 16: + raise ValueError("trellis-pair payloads must be at least 16-byte aligned") + + prepared_w13 = torch.empty_like(w13_payload) + for mode, (low_bits, high_bits) in ((0, (3, 3)), (1, (2, 4))): + ids = torch.nonzero(fc1_modes == mode, as_tuple=False).flatten() + if int(ids.numel()) == 0: + continue + selected = w13_payload.index_select(1, ids) + low_words = hidden_tiles * 8 * 16 * low_bits + low = selected[..., :low_words].reshape( + 2, ids.numel(), hidden_tiles, 8 * 16 * low_bits + ) + high = selected[..., low_words:].reshape( + 2, ids.numel(), hidden_tiles, 8 * 16 * high_bits + ) + swizzled = torch.cat((low, high), dim=-1).reshape( + 2, ids.numel(), pair_words + ) + prepared_w13.index_copy_(1, ids, swizzled) + + for name, scale, shapes in ( + ("gate_suh", gate_suh, ((1, hidden_size), (num_experts, hidden_size))), + ("up_suh", up_suh, ((1, hidden_size), (num_experts, hidden_size))), + ( + "intermediate_rotations", + intermediate_rotations, + ((num_experts, 3 * intermediate_size),), + ), + ("down_svh", down_svh, ((1, hidden_size), (num_experts, hidden_size))), + ): + if ( + scale.device != device + or scale.dtype != torch.float16 + or tuple(scale.shape) not in shapes + or not scale.is_contiguous() + ): + raise ValueError( + f"{name} must be contiguous fp16 {shapes} on {device}; got " + f"{tuple(scale.shape)}/{scale.dtype}/{scale.device}" + ) + if not bool(torch.all(torch.isfinite(scale))): + raise ValueError(f"{name} contains non-finite values") + if (gate_suh.shape[0] == 1) != (up_suh.shape[0] == 1): + raise ValueError("gate_suh and up_suh must both be broadcast or per-expert") + + tile_config = tuple(int(value) for value in tile_config) + if len(tile_config) != 4: + raise ValueError("tile_config must contain fc1_k, fc1_n, fc2_k, fc2_n") + fc1_tile_k, fc1_tile_n, fc2_tile_k, fc2_tile_n = tile_config + if fc1_tile_n != 256: + raise ValueError("FC1 pair decode requires fc1_tile_n=256") + if fc1_tile_k <= 0 or hidden_size % fc1_tile_k: + raise ValueError("fc1_tile_k must divide hidden_size") + if fc2_tile_k <= 0 or fc2_tile_k > 128 or 128 % fc2_tile_k: + raise ValueError("fc2_tile_k must be a positive divisor of 128") + if fc2_tile_n <= 0 or hidden_size % fc2_tile_n: + raise ValueError("fc2_tile_n must divide hidden_size") + if (fc1_tile_k * fc1_tile_n) // 64 != (fc2_tile_k * fc2_tile_n) // 64: + raise ValueError("FC1 and FC2 pair tiles must use the same CTA thread count") + + if dummy_scale is None: + dummy_scale = torch.zeros(4, dtype=torch.uint8, device=device) + elif ( + dummy_scale.device != device + or dummy_scale.dtype != torch.uint8 + or tuple(dummy_scale.shape) != (4,) + or not dummy_scale.is_contiguous() + or int(dummy_scale.data_ptr()) % 16 + ): + raise ValueError( + "dummy_scale must be contiguous aligned uint8[4] on the weight device" + ) + if workspace is None: + workspace = _make_workspace(device, max_blocks_per_sm=4) + elif ( + workspace.device != device + or workspace.dtype != torch.int32 + or not workspace.is_contiguous() + ): + raise ValueError("workspace must be contiguous int32 on the weight device") + + global_scale = torch.ones((num_experts,), dtype=torch.float32, device=device) + return PreparedW4A16MoeWeights( + w13=prepared_w13.reshape(-1).view(torch.int32), + w13_scale=dummy_scale, + w13_global_scale=global_scale, + w2=w2_payload.reshape(-1).view(torch.int32), + w2_scale=dummy_scale, + w2_global_scale=global_scale, + workspace=workspace, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + num_experts=num_experts, + is_gated=True, + params_dtype=params_dtype, + fc1_tile_n=fc1_tile_n, + fc2_tile_n=fc2_tile_n, + source_format="qsrt_sqg_e4m3", + w13_layout="trellis3_t256_proj", + weight_layout="trellis3_t256", + scale_format="e4m3_k32", + trellis_codebook="sqg_xor_cheb_t12", + trellis_bits=3, + fc1_trellis_pair_kind="PDYNAMIC", + fc2_trellis_pair_kind="PDYNAMIC", + fc1_trellis_pair_modes=fc1_modes, + fc2_trellis_pair_modes=fc2_modes, + gate_suh=gate_suh, + up_suh=up_suh, + intermediate_rotations=intermediate_rotations, + down_svh=down_svh, + tile_config=tile_config, ) -def _finalize_prepared_trellis_weights( +def _prepare_qsrt_p33_p43_moe_weights( + w13_payload: torch.Tensor, + w2_payload: torch.Tensor, *, - context: str, - device: torch.device, + pair_offsets_u32: torch.Tensor, + pair_is_p43: torch.Tensor, hidden_size: int, intermediate_size: int, num_experts: int, - params_dtype: torch.dtype, - w13: torch.Tensor, - w2: torch.Tensor, + activation: str, gate_suh: torch.Tensor, up_suh: torch.Tensor, intermediate_rotations: torch.Tensor, down_svh: torch.Tensor, - rotation_columns: int, - tile_config: tuple[int, int, int, int], - required_fc1_tile_n: int, - dummy_scale: torch.Tensor | None, - workspace: torch.Tensor | None, - codebook: str, - trellis_bits: int, - fc1_pair_kind: str | None, - fc2_pair_kind: str | None, - fc1_pair_modes: torch.Tensor | None, - fc2_pair_modes: torch.Tensor | None, - coupled_hadamard: bool = False, + params_dtype: torch.dtype = torch.float16, + codebook: str = "sqg_xor_cheb_t12", + tile_config: tuple[int, int, int, int] = (64, 256, 64, 256), + dummy_scale: torch.Tensor | None = None, + workspace: torch.Tensor | None = None, ) -> PreparedW4A16MoeWeights: - """Shared validation and construction tail of the trellis preparers.""" + """Bind the P33/P43 pools decoded from one atoms-v2 extent. + + ``w13_payload`` is projection-major and ``w2_payload`` is expert-major. + Within every pool, each expert owns one complete 256-channel pair in + physical record order. P33 entries contain six bits per coefficient pair; + P43 entries contain seven. ``pair_offsets_u32`` gives the exact start of + each expert in one W13 projection plane and in W2. No expert is padded. + + The caller must already have applied one common whole-record placement to + gate, up, and down. This routine validates compact geometry; it never + changes logical funding or permutes channels. + """ + + hidden_size = int(hidden_size) + intermediate_size = int(intermediate_size) + num_experts = int(num_experts) + if hidden_size <= 0 or hidden_size % 128: + raise ValueError( + "QSRT P33/P43 preparation requires hidden_size to be a positive " + "multiple of 128" + ) + if intermediate_size != 256: + raise ValueError("one prepared QSRT P33/P43 extent must own 256 channels") + if num_experts <= 0 or not validate_activation(activation): + raise ValueError("QSRT P33/P43 preparation requires gated experts") + if params_dtype != torch.float16: + raise ValueError("QSRT P33/P43 preparation requires fp16 operands") + if _normalize_trellis256_codebook(codebook) != "sqg_xor_cheb_t12": + raise ValueError("QSRT P33/P43 preparation supports only sqg_xor_cheb_t12") + if w13_payload.dtype != torch.int16 or w2_payload.dtype != torch.int16: + raise TypeError("QSRT P33/P43 payloads must use torch.int16") + if not w13_payload.is_cuda or not w2_payload.is_cuda: + raise ValueError("QSRT P33/P43 preparation requires CUDA payloads") + if w13_payload.device != w2_payload.device: + raise ValueError("w13_payload and w2_payload must share one CUDA device") + if not w13_payload.is_contiguous() or not w2_payload.is_contiguous(): + raise ValueError("QSRT P33/P43 payloads must be contiguous") + if w13_payload.ndim != 1 or w2_payload.ndim != 1: + raise ValueError("QSRT P33/P43 payload pools must be flat") + if int(w13_payload.data_ptr()) % 16 or int(w2_payload.data_ptr()) % 16: + raise ValueError("QSRT P33/P43 payloads must be at least 16-byte aligned") + device = w13_payload.device + + for name, value in ( + ("pair_offsets_u32", pair_offsets_u32), + ("pair_is_p43", pair_is_p43), + ): + if not isinstance(value, torch.Tensor): + raise TypeError(f"{name} must be a tensor") + if value.device != device or tuple(value.shape) != (num_experts,): + raise ValueError(f"{name} must have shape {(num_experts,)} on {device}") + if not value.is_contiguous(): + raise ValueError(f"{name} must be contiguous") + if pair_offsets_u32.dtype != torch.int64: + raise TypeError("pair_offsets_u32 must use int64") + if pair_is_p43.dtype not in { + torch.bool, + torch.uint8, + torch.int8, + torch.int16, + torch.int32, + torch.int64, + }: + raise TypeError("pair_is_p43 must use an integer or bool dtype") + modes = pair_is_p43.to(dtype=torch.int64).contiguous() + if not bool(torch.all((modes == 0) | (modes == 1))): + raise ValueError("pair_is_p43 values must be 0=P33 or 1=P43") + + hidden_tiles = hidden_size // 16 + p33_u32 = hidden_tiles * 8 * 8 * 6 + p43_u32 = hidden_tiles * 8 * 8 * 7 + lengths_u32 = torch.where( + modes != 0, + torch.full_like(modes, p43_u32), + torch.full_like(modes, p33_u32), + ) + if not bool(torch.all(pair_offsets_u32 >= 0)): + raise ValueError("pair_offsets_u32 must be non-negative") + order = torch.argsort(pair_offsets_u32) + sorted_offsets = pair_offsets_u32.index_select(0, order) + sorted_lengths = lengths_u32.index_select(0, order) + expected_offsets = torch.empty_like(sorted_lengths) + expected_offsets[0] = 0 + if num_experts > 1: + expected_offsets[1:] = torch.cumsum(sorted_lengths[:-1], dim=0) + if not bool(torch.equal(sorted_offsets, expected_offsets)): + raise ValueError( + "pair_offsets_u32 must describe a gap-free, non-overlapping compact " + "packing of the exact P33/P43 payload lengths" + ) + total_u32 = int((sorted_offsets[-1] + sorted_lengths[-1]).item()) + if w13_payload.numel() != 4 * total_u32: + raise ValueError( + "projection-major w13_payload must contain two compact planes of " + f"{total_u32} uint32 values" + ) + if w2_payload.numel() != 2 * total_u32: + raise ValueError( + f"w2_payload must contain {total_u32} compact uint32 values" + ) + descriptors = ((pair_offsets_u32 << 1) | modes).contiguous() for name, scale, shapes in ( ("gate_suh", gate_suh, ((1, hidden_size), (num_experts, hidden_size))), @@ -2406,7 +2679,7 @@ def _finalize_prepared_trellis_weights( ( "intermediate_rotations", intermediate_rotations, - ((num_experts, rotation_columns),), + ((num_experts, 3 * intermediate_size),), ), ("down_svh", down_svh, ((1, hidden_size), (num_experts, hidden_size))), ): @@ -2429,8 +2702,8 @@ def _finalize_prepared_trellis_weights( if len(tile_config) != 4: raise ValueError("tile_config must contain fc1_k, fc1_n, fc2_k, fc2_n") fc1_tile_k, fc1_tile_n, fc2_tile_k, fc2_tile_n = tile_config - if fc1_tile_n != required_fc1_tile_n: - raise ValueError(f"{context} requires fc1_tile_n={required_fc1_tile_n}") + if fc1_tile_n != 256: + raise ValueError("QSRT P33/P43 preparation requires fc1_tile_n=256") if fc1_tile_k <= 0 or hidden_size % fc1_tile_k: raise ValueError("fc1_tile_k must divide hidden_size") if fc2_tile_k <= 0 or fc2_tile_k > 128 or 128 % fc2_tile_k: @@ -2463,10 +2736,10 @@ def _finalize_prepared_trellis_weights( global_scale = torch.ones((num_experts,), dtype=torch.float32, device=device) return PreparedW4A16MoeWeights( - w13=w13.view(torch.int32) if w13.dtype == torch.int16 else w13, + w13=w13_payload.view(torch.int32), w13_scale=dummy_scale, w13_global_scale=global_scale, - w2=w2.view(torch.int32) if w2.dtype == torch.int16 else w2, + w2=w2_payload.view(torch.int32), w2_scale=dummy_scale, w2_global_scale=global_scale, workspace=workspace, @@ -2477,28 +2750,25 @@ def _finalize_prepared_trellis_weights( params_dtype=params_dtype, fc1_tile_n=fc1_tile_n, fc2_tile_n=fc2_tile_n, - source_format="btx", - w13_layout="trellis_t256_proj", - weight_layout="trellis_t256", + source_format="qsrt_sqg_e4m3", + w13_layout="trellis3_t256_proj", + weight_layout="trellis3_t256", scale_format="e4m3_k32", - trellis=TrellisWeightState( - codebook=codebook, - bits=trellis_bits, - fc1_pair_kind=fc1_pair_kind, - fc2_pair_kind=fc2_pair_kind, - fc1_pair_modes=fc1_pair_modes, - fc2_pair_modes=fc2_pair_modes, - gate_suh=gate_suh, - up_suh=up_suh, - intermediate_rotations=intermediate_rotations, - down_svh=down_svh, - coupled_hadamard=coupled_hadamard, - tile_config=tile_config, - ), + trellis_codebook="sqg_xor_cheb_t12", + trellis_bits=3, + fc1_trellis_pair_kind="P33_P43", + fc2_trellis_pair_kind="P33_P43", + fc1_trellis_pair_modes=descriptors, + fc2_trellis_pair_modes=descriptors, + gate_suh=gate_suh, + up_suh=up_suh, + intermediate_rotations=intermediate_rotations, + down_svh=down_svh, + tile_config=tile_config, ) -def _coupled_rotation_signs( +def _qsrt_coupled_rotation_signs( length: int, *, draw: int, @@ -2526,9 +2796,1120 @@ def _coupled_rotation_signs( ) +def _prepare_qsrt_p22_atom_v2_moe_weights( + atom_payload: torch.Tensor, + *, + first_atom_slot: int, + rotation_draws: torch.Tensor, + hidden_size: int, + intermediate_size: int, + num_experts: int, + activation: str, + gate_suh: torch.Tensor, + up_suh: torch.Tensor, + down_svh: torch.Tensor, + params_dtype: torch.dtype, + codebook: str, + tile_config: tuple[int, int, int, int], + dummy_scale: torch.Tensor | None, + workspace: torch.Tensor | None, +) -> PreparedW4A16MoeWeights: + """Restore one balanced pure-K2 coupled-Hadamard atom extent.""" + + atom_count = int(atom_payload.shape[0]) + if atom_count <= 0 or first_atom_slot < 0 or first_atom_slot + atom_count > 96: + raise ValueError("pure-K2 atom extent lies outside the 96-slot axis") + if first_atom_slot % 4 or atom_count % 4: + raise ValueError( + "pure-K2 atom extents must close 128-channel postactivation blocks" + ) + pre_half_atoms = 48 + if first_atom_slot < pre_half_atoms < first_atom_slot + atom_count: + raise ValueError( + "pure-K2 atom extents must not cross the two transformed " + "preactivation halves" + ) + source_intermediate_size = atom_count * 32 + if ( + intermediate_size < source_intermediate_size + or intermediate_size % 128 + ): + raise ValueError( + "pure-K2 runtime intermediate size must contain its atom extent " + "and close 128-channel postactivation blocks" + ) + if atom_payload.shape[1] < num_experts * _QSRT_V2_P22_ATOM_BUNDLE_BYTES: + raise ValueError("pure-K2 atom row is shorter than its expert bundles") + payload_bytes = num_experts * _QSRT_V2_P22_ATOM_BUNDLE_BYTES + if bool(torch.any(atom_payload[:, payload_bytes:] != 0)): + raise ValueError("pure-K2 atom-row padding must be zero") + if ( + rotation_draws.dtype != torch.uint8 + or tuple(rotation_draws.shape) != (num_experts,) + or not rotation_draws.is_contiguous() + or rotation_draws.device.type != "cpu" + or bool(torch.any(rotation_draws > 7)) + ): + raise ValueError( + "pure-K2 rotation_draws must be contiguous CPU uint8[num_experts]" + ) + + device = gate_suh.device + source = atom_payload[:, :payload_bytes].reshape( + atom_count, num_experts, _QSRT_V2_P22_ATOM_BUNDLE_BYTES + ) + hidden_tiles = hidden_size // 16 + source_local_tiles = source_intermediate_size // 16 + local_tiles = intermediate_size // 16 + w13 = torch.zeros( + (2, num_experts, hidden_tiles, local_tiles, 32), + dtype=torch.int16, + device=device, + ) + w2 = torch.zeros( + (num_experts, local_tiles, hidden_tiles, 32), + dtype=torch.int16, + device=device, + ) + intermediate_rotations = torch.zeros( + (num_experts, 3 * intermediate_size), + dtype=torch.float16, + device=device, + ) + scale_base = 3 * _QSRT_V2_P22_MATRIX_TRELLIS_BYTES + + for first_expert in range(0, num_experts, 64): + count = min(64, num_experts - first_expert) + for matrix_index in range(3): + begin = matrix_index * _QSRT_V2_P22_MATRIX_TRELLIS_BYTES + raw = source[ + :, + first_expert : first_expert + count, + begin : begin + _QSRT_V2_P22_MATRIX_TRELLIS_BYTES, + ] + values = ( + raw.contiguous() + .to(device=device) + .view(torch.int16) + .reshape(atom_count, count, 2, hidden_tiles, 32) + ) + if matrix_index < 2: + restored = ( + values.permute(1, 3, 0, 2, 4) + .reshape(count, hidden_tiles, source_local_tiles, 32) + ) + w13[ + matrix_index, + first_expert : first_expert + count, + :, + :source_local_tiles, + ].copy_(restored) + else: + restored = ( + values.permute(1, 0, 2, 3, 4) + .reshape(count, source_local_tiles, hidden_tiles, 32) + ) + w2[ + first_expert : first_expert + count, + :source_local_tiles, + ].copy_(restored) + + scale_begin = scale_base + matrix_index * 64 + scales = ( + source[ + :, + first_expert : first_expert + count, + scale_begin : scale_begin + 64, + ] + .contiguous() + .to(device=device) + .view(torch.float16) + .permute(1, 0, 2) + .reshape(count, source_intermediate_size) + ) + intermediate_rotations[ + first_expert : first_expert + count, + matrix_index * intermediate_size : matrix_index + * intermediate_size + + source_intermediate_size, + ].copy_(scales) + + # The preactivation signs index the interleaved length-2I coordinate and + # the postactivation signs index the ordinary I coordinate. Atom-v2's + # coupled placement makes both slices contiguous for every balanced rank. + pre_begin = 2 * first_atom_slot * 32 + pre_count = 2 * source_intermediate_size + post_begin = first_atom_slot * 32 + signs = torch.ones( + (num_experts, 3 * intermediate_size), dtype=torch.float16 + ) + for draw in sorted(set(int(value) for value in rotation_draws.tolist())): + rows = torch.nonzero(rotation_draws == draw, as_tuple=False).flatten() + pre = _qsrt_coupled_rotation_signs(2 * 3072, draw=draw, axis=1)[ + pre_begin : pre_begin + pre_count + ] + post = _qsrt_coupled_rotation_signs(3072, draw=draw, axis=2)[ + post_begin : post_begin + source_intermediate_size + ] + draw_signs = torch.ones( + (rows.numel(), 3 * intermediate_size), dtype=torch.float16 + ) + draw_signs[:, :pre_count].copy_( + pre.to(torch.float16).expand(rows.numel(), -1) + ) + draw_signs[ + :, + 2 * intermediate_size : 2 * intermediate_size + + source_intermediate_size, + ].copy_(post.to(torch.float16).expand(rows.numel(), -1)) + signs.index_copy_(0, rows, draw_signs) + coupled_signs = signs.to(device=device, non_blocking=True).contiguous() + + # Each balanced rank owns a contiguous slice of the transformed length-2I + # preactivation axis. Both physical FC1 slots in that slice therefore use + # the same ordinary input-side scale table: the first 48 atoms derive from + # the first stored half, and the final 48 from the second. + source_suh = gate_suh if first_atom_slot < pre_half_atoms else up_suh + + prepared = prepare_trellis256_moe_weights( + w13, + w2, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + num_experts=num_experts, + activation=activation, + params_dtype=params_dtype, + fc1_tile_n=tile_config[1], + fc2_tile_n=tile_config[3], + w13_layout="trellis3_t256_proj", + trellis_bits=2, + dummy_scale=dummy_scale, + codebook=codebook, + gate_suh=source_suh, + up_suh=source_suh, + intermediate_rotations=intermediate_rotations, + down_svh=down_svh, + tile_config=tile_config, + workspace=workspace, + ) + return replace( + prepared, + coupled_hadamard=True, + intermediate_rotations=torch.cat( + (intermediate_rotations, coupled_signs), dim=1 + ).contiguous(), + ) + + +def _prepare_qsrt_coupled_h308_atom_v2_moe_weights( + atom_payload: torch.Tensor, + *, + first_atom_slot: int, + rotation_draws: torch.Tensor, + hidden_size: int, + intermediate_size: int, + num_experts: int, + activation: str, + gate_suh: torch.Tensor, + up_suh: torch.Tensor, + down_svh: torch.Tensor, + params_dtype: torch.dtype, + codebook: str, + tile_config: tuple[int, int, int, int], + dummy_scale: torch.Tensor | None, + workspace: torch.Tensor | None, +) -> PreparedW4A16MoeWeights: + """Restore one balanced coupled-Hadamard H308 atom extent.""" + + atom_count = int(atom_payload.shape[0]) + if atom_count != 8 or first_atom_slot % 8 or not 0 <= first_atom_slot <= 88: + raise ValueError( + "coupled H308 extents must contain one aligned eight-atom pair" + ) + if intermediate_size != 256: + raise ValueError("coupled H308 requires local intermediate_size=256") + physical_pair = first_atom_slot // 8 + if physical_pair == 5: + fc1_kind, fc2_kind = "P43", "P33" + fc1_matrix_bytes = _QSRT_V2_P43_MATRIX_TRELLIS_BYTES + fc2_matrix_bytes = _QSRT_V2_P33_MATRIX_TRELLIS_BYTES + bundle_bytes = _QSRT_V2_COUPLED_H308_P43_P33_ATOM_BUNDLE_BYTES + elif physical_pair == 11: + fc1_kind, fc2_kind = "P43", "P44" + fc1_matrix_bytes = _QSRT_V2_P43_MATRIX_TRELLIS_BYTES + fc2_matrix_bytes = _QSRT_V2_P44_MATRIX_TRELLIS_BYTES + bundle_bytes = _QSRT_V2_COUPLED_H308_P43_P44_ATOM_BUNDLE_BYTES + else: + fc1_kind = fc2_kind = "P33" + fc1_matrix_bytes = fc2_matrix_bytes = _QSRT_V2_P33_MATRIX_TRELLIS_BYTES + bundle_bytes = _QSRT_V2_P33_ATOM_BUNDLE_BYTES + if atom_payload.shape[1] < num_experts * bundle_bytes: + raise ValueError("coupled H308 atom row is shorter than its expert bundles") + payload_bytes = num_experts * bundle_bytes + if bool(torch.any(atom_payload[:, payload_bytes:] != 0)): + raise ValueError("coupled H308 atom-row padding must be zero") + if ( + rotation_draws.dtype != torch.uint8 + or tuple(rotation_draws.shape) != (num_experts,) + or not rotation_draws.is_contiguous() + or rotation_draws.device.type != "cpu" + or bool(torch.any(rotation_draws > 7)) + ): + raise ValueError( + "coupled H308 rotation_draws must be contiguous CPU uint8[num_experts]" + ) + if tuple(int(value) for value in tile_config) != (64, 256, 64, 256): + raise ValueError( + "coupled H308 currently requires tile_config=(64, 256, 64, 256)" + ) + + device = gate_suh.device + if device.type != "cuda": + raise ValueError("coupled H308 preparation requires CUDA output tensors") + for name, value, shapes in ( + ("gate_suh", gate_suh, ((1, hidden_size), (num_experts, hidden_size))), + ("up_suh", up_suh, ((1, hidden_size), (num_experts, hidden_size))), + ("down_svh", down_svh, ((1, hidden_size), (num_experts, hidden_size))), + ): + if ( + value.device != device + or value.dtype != torch.float16 + or tuple(value.shape) not in shapes + or not value.is_contiguous() + or not bool(torch.all(torch.isfinite(value))) + ): + raise ValueError( + f"{name} must be finite contiguous fp16 {shapes} on {device}" + ) + if (gate_suh.shape[0] == 1) != (up_suh.shape[0] == 1): + raise ValueError("gate_suh and up_suh must both be broadcast or per-expert") + + source = atom_payload[:, :payload_bytes].reshape( + atom_count, num_experts, bundle_bytes + ) + hidden_tiles = hidden_size // 16 + pair_bits = {"P33": (3, 3), "P43": (4, 3), "P44": (4, 4)} + + def restore_matrix( + begin: int, + matrix_bytes: int, + kind: str, + *, + fc1: bool, + out: torch.Tensor | None = None, + ) -> torch.Tensor: + low_bits, high_bits = pair_bits[kind] + low_bytes = hidden_tiles * 32 * low_bits + if low_bytes + hidden_tiles * 32 * high_bits != matrix_bytes: + raise AssertionError("coupled H308 matrix geometry drifted") + elements_per_expert = atom_count * matrix_bytes // 2 + if out is None: + out = torch.empty( + (num_experts, elements_per_expert), + dtype=torch.int16, + device=device, + ) + elif ( + out.dtype != torch.int16 + or out.device != device + or tuple(out.shape) != (num_experts, elements_per_expert) + or not out.is_contiguous() + ): + raise ValueError("coupled H308 restore destination has invalid geometry") + for first_expert in range(0, num_experts, 64): + count = min(64, num_experts - first_expert) + raw = source[ + :, + first_expert : first_expert + count, + begin : begin + matrix_bytes, + ] + low = ( + raw[..., :low_bytes] + .contiguous() + .to(device=device) + .view(torch.int16) + .reshape(atom_count, count, hidden_tiles, 16 * low_bits) + ) + high = ( + raw[..., low_bytes:] + .contiguous() + .to(device=device) + .view(torch.int16) + .reshape(atom_count, count, hidden_tiles, 16 * high_bits) + ) + if fc1: + low = low.permute(1, 2, 0, 3).reshape(count, hidden_tiles, -1) + high = high.permute(1, 2, 0, 3).reshape(count, hidden_tiles, -1) + restored = torch.cat((low, high), dim=2).reshape(count, -1) + else: + low = low.permute(1, 0, 2, 3).reshape(count, -1) + high = high.permute(1, 0, 2, 3).reshape(count, -1) + restored = torch.cat((low, high), dim=1) + out[first_expert : first_expert + count].copy_(restored) + return out + + matrix_offsets = (0, fc1_matrix_bytes, 2 * fc1_matrix_bytes) + fc1_elements_per_expert = atom_count * fc1_matrix_bytes // 2 + w13_storage = torch.empty( + (2, num_experts, fc1_elements_per_expert), + dtype=torch.int16, + device=device, + ) + restore_matrix( + matrix_offsets[0], + fc1_matrix_bytes, + fc1_kind, + fc1=True, + out=w13_storage[0], + ) + restore_matrix( + matrix_offsets[1], + fc1_matrix_bytes, + fc1_kind, + fc1=True, + out=w13_storage[1], + ) + w13 = w13_storage.reshape(-1) + w2 = restore_matrix( + matrix_offsets[2], fc2_matrix_bytes, fc2_kind, fc1=False + ).reshape(-1) + + scale_base = 2 * fc1_matrix_bytes + fc2_matrix_bytes + intermediate_scales = torch.empty( + (num_experts, 3 * intermediate_size), + dtype=torch.float16, + device=device, + ) + for matrix_index in range(3): + begin = scale_base + matrix_index * 64 + for first_expert in range(0, num_experts, 64): + count = min(64, num_experts - first_expert) + values = ( + source[ + :, + first_expert : first_expert + count, + begin : begin + 64, + ] + .contiguous() + .to(device=device) + .view(torch.float16) + .reshape(atom_count, count, 32) + ) + restored = torch.cat( + ( + values[:, :, :16].permute(1, 0, 2).reshape(count, 128), + values[:, :, 16:].permute(1, 0, 2).reshape(count, 128), + ), + dim=1, + ) + intermediate_scales[ + first_expert : first_expert + count, + matrix_index + * intermediate_size : (matrix_index + 1) + * intermediate_size, + ].copy_(restored) + + pre_begin = 2 * first_atom_slot * 32 + post_begin = first_atom_slot * 32 + signs = torch.empty((num_experts, 3 * intermediate_size), dtype=torch.float16) + for draw in sorted(set(int(value) for value in rotation_draws.tolist())): + rows = torch.nonzero(rotation_draws == draw, as_tuple=False).flatten() + pre = _qsrt_coupled_rotation_signs(6144, draw=draw, axis=1)[ + pre_begin : pre_begin + 2 * intermediate_size + ] + post = _qsrt_coupled_rotation_signs(3072, draw=draw, axis=2)[ + post_begin : post_begin + intermediate_size + ] + signs.index_copy_( + 0, + rows, + torch.cat((pre, post)).to(torch.float16).expand(rows.numel(), -1), + ) + rotations = torch.cat( + (intermediate_scales, signs.to(device=device, non_blocking=True)), dim=1 + ).contiguous() + + if dummy_scale is None: + dummy_scale = torch.zeros(4, dtype=torch.uint8, device=device) + if ( + dummy_scale.device != device + or dummy_scale.dtype != torch.uint8 + or tuple(dummy_scale.shape) != (4,) + or not dummy_scale.is_contiguous() + or dummy_scale.data_ptr() % 16 + ): + raise ValueError( + "dummy_scale must be contiguous aligned uint8[4] on the weight device" + ) + if workspace is None: + workspace = _make_workspace(device, max_blocks_per_sm=4) + if ( + workspace.device != device + or workspace.dtype != torch.int32 + or not workspace.is_contiguous() + ): + raise ValueError("workspace must be contiguous int32 on the weight device") + normalized_codebook = _trellis256_marker_codebook( + mcg=None, mul1_e4m3=None, codebook=codebook + ) + global_scale = torch.ones((num_experts,), dtype=torch.float32, device=device) + source_suh = gate_suh if physical_pair < 6 else up_suh + return PreparedW4A16MoeWeights( + w13=w13.view(torch.int32), + w13_scale=dummy_scale, + w13_global_scale=global_scale, + w2=w2.view(torch.int32), + w2_scale=dummy_scale, + w2_global_scale=global_scale, + workspace=workspace, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + num_experts=num_experts, + is_gated=True, + params_dtype=params_dtype, + fc1_tile_n=256, + fc2_tile_n=256, + source_format="qsrt_sqg_e4m3", + w13_layout="trellis3_t256_proj", + weight_layout="trellis3_t256", + scale_format="e4m3_k32", + trellis_codebook=normalized_codebook, + trellis_bits=3, + fc1_trellis_pair_kind=fc1_kind, + fc2_trellis_pair_kind=fc2_kind, + gate_suh=source_suh, + up_suh=source_suh, + intermediate_rotations=rotations, + down_svh=down_svh, + coupled_hadamard=True, + tile_config=(64, 256, 64, 256), + ) + + +def prepare_qsrt_atom_v2_moe_weights( + atom_payload: torch.Tensor, + *, + first_atom_slot: int, + layer_index: int, + profile: str = _QSRT_V2_PROFILE_H308, + rotation_draws: torch.Tensor | None = None, + hidden_size: int, + intermediate_size: int, + num_experts: int, + activation: str, + gate_suh: torch.Tensor, + up_suh: torch.Tensor, + down_svh: torch.Tensor, + params_dtype: torch.dtype = torch.float16, + codebook: str = "sqg_xor_cheb_t12", + tile_config: tuple[int, int, int, int] = (64, 256, 64, 256), + dummy_scale: torch.Tensor | None = None, + workspace: torch.Tensor | None = None, +) -> PreparedW4A16MoeWeights: + """Prepare one balanced rank extent from a QSRT atoms-v2 profile.""" + + hidden_size = int(hidden_size) + intermediate_size = int(intermediate_size) + num_experts = int(num_experts) + first_atom_slot = int(first_atom_slot) + layer_index = int(layer_index) + if num_experts != 896: + raise ValueError("QSRT atoms-v2 currently requires all 896 experts") + if hidden_size != 3584: + raise ValueError("QSRT atoms-v2 currently requires hidden_size=3584") + if profile not in { + _QSRT_V2_PROFILE_H308, + _QSRT_V2_PROFILE_COUPLED_K2, + _QSRT_V2_PROFILE_COUPLED_H308, + }: + raise ValueError(f"unsupported QSRT atoms-v2 profile {profile!r}") + if not validate_activation(activation) or params_dtype != torch.float16: + raise ValueError("QSRT atoms-v2 requires gated fp16 preparation") + if _normalize_trellis256_codebook(codebook) != "sqg_xor_cheb_t12": + raise ValueError("QSRT atoms-v2 supports only sqg_xor_cheb_t12") + if ( + atom_payload.dtype != torch.uint8 + or atom_payload.ndim != 2 + or atom_payload.shape[0] <= 0 + or not atom_payload.is_contiguous() + or atom_payload.device.type not in {"cpu", "cuda"} + ): + raise ValueError( + "atom_payload must be contiguous CPU/CUDA uint8 [atoms, row_stride]" + ) + if not 1 <= layer_index <= 92: + raise ValueError("layer_index must identify a Kimi-K3 MoE layer") + + if profile == _QSRT_V2_PROFILE_COUPLED_K2: + if rotation_draws is None: + raise ValueError("pure-K2 atoms-v2 requires coupled rotation draws") + return _prepare_qsrt_p22_atom_v2_moe_weights( + atom_payload, + first_atom_slot=first_atom_slot, + rotation_draws=rotation_draws, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + num_experts=num_experts, + activation=activation, + gate_suh=gate_suh, + up_suh=up_suh, + down_svh=down_svh, + params_dtype=params_dtype, + codebook=codebook, + tile_config=tile_config, + dummy_scale=dummy_scale, + workspace=workspace, + ) + if profile == _QSRT_V2_PROFILE_COUPLED_H308: + if rotation_draws is None: + raise ValueError("coupled H308 atoms-v2 requires rotation draws") + return _prepare_qsrt_coupled_h308_atom_v2_moe_weights( + atom_payload, + first_atom_slot=first_atom_slot, + rotation_draws=rotation_draws, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + num_experts=num_experts, + activation=activation, + gate_suh=gate_suh, + up_suh=up_suh, + down_svh=down_svh, + params_dtype=params_dtype, + codebook=codebook, + tile_config=tile_config, + dummy_scale=dummy_scale, + workspace=workspace, + ) + + if intermediate_size != 256: + raise ValueError("H308 atoms-v2 preparation requires 256 local channels") + if atom_payload.shape[0] != _QSRT_ATOMS_PER_PAIR: + raise ValueError("H308 atoms-v2 requires exactly eight atom rows") + if first_atom_slot % _QSRT_ATOMS_PER_PAIR or not 0 <= first_atom_slot < 96: + raise ValueError("H308 atoms-v2 requires a pair-aligned atom extent") + if rotation_draws is not None: + raise ValueError("H308 atoms-v2 must not carry coupled rotation draws") + + device = gate_suh.device + if device.type != "cuda": + raise ValueError("QSRT atoms-v2 preparation requires CUDA output tensors") + physical_pair = first_atom_slot // _QSRT_ATOMS_PER_PAIR + expert_ids = torch.arange(num_experts, dtype=torch.int64, device=device) + rotation = (_QSRT_EXPERT_ROTATION_MULTIPLIER * expert_ids + layer_index) % 12 + base_pair = (physical_pair - rotation) % 12 + modes = ((base_pair == 0) | (base_pair == 6)).to(torch.int64) + p33_ids = torch.nonzero(modes == 0, as_tuple=False).flatten() + p43_ids = torch.nonzero(modes == 1, as_tuple=False).flatten() + count33 = int(p33_ids.numel()) + count43 = int(p43_ids.numel()) + p33_bytes = count33 * _QSRT_V2_P33_ATOM_BUNDLE_BYTES + p43_bytes = count43 * _QSRT_V2_P43_ATOM_BUNDLE_BYTES + if atom_payload.shape[1] < p33_bytes + p43_bytes: + raise ValueError("QSRT atoms-v2 row is shorter than its compact groups") + if bool(torch.any(atom_payload[:, p33_bytes + p43_bytes :] != 0)): + raise ValueError("QSRT atoms-v2 row padding must be zero") + + groups: list[tuple[torch.Tensor, torch.Tensor, bool, int]] = [] + for ids, p43, begin, bundle in ( + (p33_ids, False, 0, _QSRT_V2_P33_ATOM_BUNDLE_BYTES), + (p43_ids, True, p33_bytes, _QSRT_V2_P43_ATOM_BUNDLE_BYTES), + ): + count = int(ids.numel()) + view = atom_payload[:, begin : begin + count * bundle].reshape( + _QSRT_ATOMS_PER_PAIR, count, bundle + ) + groups.append((ids, view, p43, bundle)) + + hidden_tiles = hidden_size // 16 + + def _restore_group_matrix_into( + source: torch.Tensor, + matrix_index: int, + destination: torch.Tensor, + *, + p43: bool, + fc1: bool, + ) -> None: + low_bits, high_bits = ((4, 3) if p43 else (3, 3)) + matrix_bytes = ( + _QSRT_V2_P43_MATRIX_TRELLIS_BYTES + if p43 + else _QSRT_V2_P33_MATRIX_TRELLIS_BYTES + ) + count = source.shape[1] + words_per_expert = destination.numel() // count + destination = destination.reshape(count, words_per_expert) + # The canonical slab is close to 1 GiB per rank. Stream bounded + # expert groups through the GPU so loading never needs the canonical + # slab and its equally-sized prepared view resident at once. + for first_expert in range(0, count, 64): + chunk_count = min(64, count - first_expert) + raw = source[ + :, + first_expert : first_expert + chunk_count, + matrix_index * matrix_bytes : (matrix_index + 1) * matrix_bytes, + ] + selected = ( + raw.contiguous() + .to(device=device) + .view(torch.int16) + .reshape(_QSRT_ATOMS_PER_PAIR, chunk_count, -1) + .permute(1, 0, 2) + ) + low_words = hidden_tiles * 16 * low_bits + low = selected[..., :low_words].reshape( + chunk_count, + _QSRT_ATOMS_PER_PAIR, + hidden_tiles, + 16 * low_bits, + ) + high = selected[..., low_words:].reshape( + chunk_count, + _QSRT_ATOMS_PER_PAIR, + hidden_tiles, + 16 * high_bits, + ) + target = destination[first_expert : first_expert + chunk_count] + if fc1: + target = target.reshape(chunk_count, hidden_tiles, -1) + target[..., : _QSRT_ATOMS_PER_PAIR * 16 * low_bits].reshape( + chunk_count, + hidden_tiles, + _QSRT_ATOMS_PER_PAIR, + 16 * low_bits, + ).copy_(low.permute(0, 2, 1, 3)) + target[..., _QSRT_ATOMS_PER_PAIR * 16 * low_bits :].reshape( + chunk_count, + hidden_tiles, + _QSRT_ATOMS_PER_PAIR, + 16 * high_bits, + ).copy_(high.permute(0, 2, 1, 3)) + else: + low_words_per_expert = low.numel() // chunk_count + target[:, :low_words_per_expert].copy_( + low.reshape(chunk_count, -1) + ) + target[:, low_words_per_expert:].copy_( + high.reshape(chunk_count, -1) + ) + + group_matrix_words = [] + for ids, _source, p43, _bundle in groups: + matrix_bytes = ( + _QSRT_V2_P43_MATRIX_TRELLIS_BYTES + if p43 + else _QSRT_V2_P33_MATRIX_TRELLIS_BYTES + ) + group_matrix_words.append( + int(ids.numel()) * _QSRT_ATOMS_PER_PAIR * matrix_bytes // 2 + ) + matrix_words = sum(group_matrix_words) + w13_payload = torch.empty(2 * matrix_words, dtype=torch.int16, device=device) + w2_payload = torch.empty(matrix_words, dtype=torch.int16, device=device) + + for matrix_index in range(2): + group_offset = matrix_index * matrix_words + for group_words, (_ids, source, p43, _bundle) in zip( + group_matrix_words, groups + ): + _restore_group_matrix_into( + source, + matrix_index, + w13_payload.narrow(0, group_offset, group_words), + p43=p43, + fc1=True, + ) + group_offset += group_words + group_offset = 0 + for group_words, (_ids, source, p43, _bundle) in zip( + group_matrix_words, groups + ): + _restore_group_matrix_into( + source, + 2, + w2_payload.narrow(0, group_offset, group_words), + p43=p43, + fc1=False, + ) + group_offset += group_words + + p33_u32 = hidden_tiles * 8 * 8 * 6 + p43_u32 = hidden_tiles * 8 * 8 * 7 + offsets = torch.empty(num_experts, dtype=torch.int64, device=device) + offsets.index_copy_(0, p33_ids, torch.arange(count33, device=device) * p33_u32) + offsets.index_copy_( + 0, + p43_ids, + count33 * p33_u32 + torch.arange(count43, device=device) * p43_u32, + ) + + intermediate_rotations = torch.empty( + (num_experts, 3 * intermediate_size), dtype=torch.float16, device=device + ) + for ids, source, p43, _bundle in groups: + matrix_bytes = ( + _QSRT_V2_P43_MATRIX_TRELLIS_BYTES + if p43 + else _QSRT_V2_P33_MATRIX_TRELLIS_BYTES + ) + count = int(ids.numel()) + matrices = [] + for matrix_index in range(3): + begin = 3 * matrix_bytes + matrix_index * _QSRT_MATRIX_ATOM_SCALE_BYTES + raw = source.narrow(2, begin, _QSRT_MATRIX_ATOM_SCALE_BYTES) + values = ( + raw.contiguous() + .to(device=device) + .view(torch.float16) + .reshape(_QSRT_ATOMS_PER_PAIR, count, _QSRT_ATOM_CHANNELS) + .permute(1, 0, 2) + ) + matrices.append( + torch.cat( + ( + values[..., :16].reshape(count, -1), + values[..., 16:].reshape(count, -1), + ), + dim=1, + ) + ) + intermediate_rotations.index_copy_(0, ids, torch.cat(tuple(matrices), dim=1)) + + return _prepare_qsrt_p33_p43_moe_weights( + w13_payload, + w2_payload, + pair_offsets_u32=offsets, + pair_is_p43=modes, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + num_experts=num_experts, + activation=activation, + gate_suh=gate_suh, + up_suh=up_suh, + intermediate_rotations=intermediate_rotations.contiguous(), + down_svh=down_svh, + params_dtype=params_dtype, + codebook=codebook, + tile_config=tile_config, + dummy_scale=dummy_scale, + workspace=workspace, + ) + + +def prepare_qsrt_atom_moe_weights( + atom_payload: torch.Tensor, + *, + first_atom_slot: int, + layer_index: int, + expert_ids: torch.Tensor, + format_codes: torch.Tensor, + hidden_size: int, + intermediate_size: int, + num_experts: int, + activation: str, + gate_suh: torch.Tensor, + up_suh: torch.Tensor, + down_svh: torch.Tensor, + params_dtype: torch.dtype = torch.float16, + codebook: str | None = "sqg_xor_cheb_t12", + tile_config: tuple[int, int, int, int] = (64, 256, 64, 256), + dummy_scale: torch.Tensor | None = None, + workspace: torch.Tensor | None = None, +) -> PreparedW4A16MoeWeights: + """Prepare one canonical QSRT atom extent for the fused decoder. + + ``atom_payload`` is the de-padded checkpoint tensor + ``[atom_slots, E, 129216]u8``. Storage is independent of a deployment's + shard count: every first-axis entry owns a complete 32-channel physical + atom. The currently qualified fused kernel consumes one aligned extent of + eight atoms (256 channels), so this load-time transform restores its local + P24/P33 view without decoding or re-encoding any trellis symbols. + + Expert format bytes encode ``r13`` in the high nibble and ``r2`` in the + low nibble. The physical atom rotation is inverted from ``layer_index`` + and the global ``expert_ids``; no rank-local pair-mode side table exists in + the checkpoint. + """ + + hidden_size = int(hidden_size) + intermediate_size = int(intermediate_size) + num_experts = int(num_experts) + if hidden_size <= 0 or hidden_size % 128: + raise ValueError( + "QSRT atoms require hidden_size to be a positive multiple " + f"of 128, got {hidden_size}" + ) + if intermediate_size != _QSRT_ATOMS_PER_PAIR * _QSRT_ATOM_CHANNELS: + raise ValueError( + "the current QSRT fused kernel requires one eight-atom extent " + "(intermediate_size=256), " + f"got {intermediate_size}" + ) + if num_experts <= 0: + raise ValueError("QSRT atoms require num_experts > 0") + if not validate_activation(activation): + raise ValueError("QSRT atoms require a gated expert activation") + if params_dtype != torch.float16: + raise ValueError("full-rotation QSRT requires fp16 MMA operands") + if atom_payload.dtype != torch.uint8: + raise TypeError("QSRT atom payloads must use torch.uint8") + if not atom_payload.is_cuda: + raise ValueError("QSRT atom preparation requires CUDA storage") + expected_atoms = ( + _QSRT_ATOMS_PER_PAIR, + num_experts, + _QSRT_ATOM_BUNDLE_BYTES, + ) + if tuple(atom_payload.shape) != expected_atoms: + raise ValueError( + f"atom_payload must have shape {expected_atoms}, got " + f"{tuple(atom_payload.shape)}" + ) + expected_inner_strides = (_QSRT_ATOM_BUNDLE_BYTES, 1) + if tuple(atom_payload.stride()[1:]) != expected_inner_strides or int( + atom_payload.stride(0) + ) < num_experts * _QSRT_ATOM_BUNDLE_BYTES: + raise ValueError( + "QSRT atom payloads must be expert-major within each atom row; " + "the row stride may include checkpoint alignment padding" + ) + device = atom_payload.device + first_atom_slot = int(first_atom_slot) + layer_index = int(layer_index) + if not 0 <= first_atom_slot < _QSRT_ATOMS_PER_EXPERT: + raise ValueError( + f"first_atom_slot must be in 0..{_QSRT_ATOMS_PER_EXPERT - 1}" + ) + if first_atom_slot % _QSRT_ATOMS_PER_PAIR: + raise ValueError("the current QSRT kernel requires a pair-aligned atom extent") + if not 1 <= layer_index <= 92: + raise ValueError("layer_index must identify a Kimi-K3 MoE layer in 1..92") + + def _normalize_vector(name: str, value: torch.Tensor) -> torch.Tensor: + if not isinstance(value, torch.Tensor): + raise TypeError(f"{name} must be a tensor") + if value.device != device or tuple(value.shape) != (num_experts,): + raise ValueError( + f"{name} must have shape {(num_experts,)} on {device}" + ) + if value.dtype not in { + torch.uint8, + torch.int8, + torch.int16, + torch.int32, + torch.int64, + }: + raise TypeError(f"{name} must use an integer dtype") + return value.to(dtype=torch.int32).contiguous() + + expert_ids_i32 = _normalize_vector("expert_ids", expert_ids) + format_codes_i32 = _normalize_vector("format_codes", format_codes) + if not bool(torch.all((expert_ids_i32 >= 0) & (expert_ids_i32 < 896))): + raise ValueError("expert_ids must lie in 0..895") + r13 = format_codes_i32 >> 4 + r2 = format_codes_i32 & 0xF + if not bool(torch.all((r13 >= 0) & (r13 <= 2) & (r2 >= 0) & (r2 <= 2))): + raise ValueError("compressed QSRT format codes must encode R0/R1/R2") + physical_pair = first_atom_slot // _QSRT_ATOMS_PER_PAIR + rotation = ( + _QSRT_EXPERT_ROTATION_MULTIPLIER * expert_ids_i32 + layer_index + ) % 12 + logical_pair = (physical_pair - rotation) % 12 + fc1_pair_modes = (logical_pair < r13).to(dtype=torch.int32).contiguous() + fc2_pair_modes = (logical_pair < r2).to(dtype=torch.int32).contiguous() + fc1_pair_kind = "PDYNAMIC" + fc2_pair_kind = "PDYNAMIC" + + hidden_tiles = hidden_size // 16 + pair_words = hidden_tiles * 8 * 16 * 6 + words_per_atom = _QSRT_MATRIX_ATOM_TRELLIS_BYTES // 2 + + def _matrix_words(matrix_index: int) -> torch.Tensor: + begin = _QSRT_MATRIX_TRELLIS_OFFSETS[matrix_index] + raw = atom_payload.narrow( + 2, begin, _QSRT_MATRIX_ATOM_TRELLIS_BYTES + ).contiguous() + return raw.view(torch.int16).reshape( + _QSRT_ATOMS_PER_PAIR, num_experts, words_per_atom + ).permute(1, 0, 2) + + def _restore_matrix( + matrix_index: int, modes: torch.Tensor, *, fc1: bool + ) -> torch.Tensor: + source = _matrix_words(matrix_index) + output = torch.empty( + (num_experts, pair_words), dtype=torch.int16, device=device + ) + for mode, (low_bits, high_bits) in ((0, (3, 3)), (1, (2, 4))): + ids = torch.nonzero(modes == mode, as_tuple=False).flatten() + if int(ids.numel()) == 0: + continue + selected = source.index_select(0, ids).narrow( + 2, 0, hidden_tiles * 16 * 6 + ) + low_words = hidden_tiles * 16 * low_bits + low = selected[..., :low_words].reshape( + -1, _QSRT_ATOMS_PER_PAIR, hidden_tiles, 16 * low_bits + ) + high = selected[..., low_words:].reshape( + -1, _QSRT_ATOMS_PER_PAIR, hidden_tiles, 16 * high_bits + ) + if fc1: + low = low.permute(0, 2, 1, 3).reshape( + ids.numel(), hidden_tiles, -1 + ) + high = high.permute(0, 2, 1, 3).reshape( + ids.numel(), hidden_tiles, -1 + ) + # FC1 places both 128-channel records under each K16 tile. + restored = torch.cat((low, high), dim=-1).reshape( + ids.numel(), -1 + ) + else: + # FC2 retains its K-major low-plane/high-plane ordering. + restored = torch.cat( + ( + low.reshape(ids.numel(), -1), + high.reshape(ids.numel(), -1), + ), + dim=1, + ) + output.index_copy_(0, ids, restored) + return output + + prepared_w13_i16 = torch.stack( + ( + _restore_matrix(0, fc1_pair_modes, fc1=True), + _restore_matrix(1, fc1_pair_modes, fc1=True), + ) + ).reshape(-1) + prepared_w2_i16 = _restore_matrix( + 2, fc2_pair_modes, fc1=False + ).reshape(-1) + + def _local_scale(matrix_index: int) -> torch.Tensor: + begin = _QSRT_MATRIX_SCALE_OFFSETS[matrix_index] + raw = atom_payload.narrow( + 2, begin, _QSRT_MATRIX_ATOM_SCALE_BYTES + ).contiguous() + values = raw.view(torch.float16).reshape( + _QSRT_ATOMS_PER_PAIR, num_experts, _QSRT_ATOM_CHANNELS + ).permute(1, 0, 2) + return torch.cat( + ( + values[..., :16].reshape(num_experts, -1), + values[..., 16:].reshape(num_experts, -1), + ), + dim=1, + ) + + intermediate_rotations = torch.cat( + (_local_scale(0), _local_scale(1), _local_scale(2)), dim=1 + ).contiguous() + + normalized_codebook = _trellis256_marker_codebook( + mcg=None, + mul1_e4m3=None, + codebook=codebook, + ) + + for name, scale, shapes in ( + ("gate_suh", gate_suh, ((1, hidden_size), (num_experts, hidden_size))), + ("up_suh", up_suh, ((1, hidden_size), (num_experts, hidden_size))), + ( + "intermediate_rotations", + intermediate_rotations, + ((num_experts, 3 * intermediate_size),), + ), + ("down_svh", down_svh, ((1, hidden_size), (num_experts, hidden_size))), + ): + if ( + scale.device != device + or scale.dtype != torch.float16 + or tuple(scale.shape) not in shapes + or not scale.is_contiguous() + ): + raise ValueError( + f"{name} must be contiguous fp16 {shapes} on {device}; got " + f"{tuple(scale.shape)}/{scale.dtype}/{scale.device}" + ) + if not bool(torch.all(torch.isfinite(scale))): + raise ValueError(f"{name} contains non-finite values") + if (gate_suh.shape[0] == 1) != (up_suh.shape[0] == 1): + raise ValueError("gate_suh and up_suh must both be broadcast or per-expert") + + tile_config = tuple(int(value) for value in tile_config) + if len(tile_config) != 4: + raise ValueError("tile_config must contain fc1_k, fc1_n, fc2_k, fc2_n") + fc1_tile_k, fc1_tile_n, fc2_tile_k, fc2_tile_n = tile_config + if fc1_tile_n != 256: + raise ValueError("QSRT pair decode requires fc1_tile_n=256") + if fc1_tile_k <= 0 or hidden_size % fc1_tile_k: + raise ValueError("fc1_tile_k must divide hidden_size") + if fc2_tile_k <= 0 or fc2_tile_k > 128 or 128 % fc2_tile_k: + raise ValueError("fc2_tile_k must be a positive divisor of 128") + if fc2_tile_n <= 0 or hidden_size % fc2_tile_n: + raise ValueError("fc2_tile_n must divide hidden_size") + if (fc1_tile_k * fc1_tile_n) // 64 != (fc2_tile_k * fc2_tile_n) // 64: + raise ValueError("FC1 and FC2 pair tiles must use the same CTA thread count") + + if dummy_scale is None: + dummy_scale = torch.zeros(4, dtype=torch.uint8, device=device) + elif ( + dummy_scale.device != device + or dummy_scale.dtype != torch.uint8 + or tuple(dummy_scale.shape) != (4,) + or not dummy_scale.is_contiguous() + or int(dummy_scale.data_ptr()) % 16 + ): + raise ValueError( + "dummy_scale must be contiguous aligned uint8[4] on the weight device" + ) + if workspace is None: + workspace = _make_workspace(device, max_blocks_per_sm=4) + elif ( + workspace.device != device + or workspace.dtype != torch.int32 + or not workspace.is_contiguous() + ): + raise ValueError("workspace must be contiguous int32 on the weight device") + + global_scale = torch.ones((num_experts,), dtype=torch.float32, device=device) + return PreparedW4A16MoeWeights( + w13=prepared_w13_i16.view(torch.int32), + w13_scale=dummy_scale, + w13_global_scale=global_scale, + w2=prepared_w2_i16.view(torch.int32), + w2_scale=dummy_scale, + w2_global_scale=global_scale, + workspace=workspace, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + num_experts=num_experts, + is_gated=True, + params_dtype=params_dtype, + fc1_tile_n=fc1_tile_n, + fc2_tile_n=fc2_tile_n, + source_format="qsrt_sqg_e4m3", + w13_layout="trellis3_t256_proj", + weight_layout="trellis3_t256", + scale_format="e4m3_k32", + trellis_codebook=normalized_codebook, + trellis_bits=3, + fc1_trellis_pair_kind=fc1_pair_kind, + fc2_trellis_pair_kind=fc2_pair_kind, + fc1_trellis_pair_modes=fc1_pair_modes, + fc2_trellis_pair_modes=fc2_pair_modes, + gate_suh=gate_suh, + up_suh=up_suh, + intermediate_rotations=intermediate_rotations, + down_svh=down_svh, + tile_config=tile_config, + ) + + __all__ = [ "PreparedW4A16MoeWeights", - "TrellisWeightState", "PreparedTrellis256DenseWeight", "W4A16PackedBuffers", "W4A16FC2Weights", @@ -2538,6 +3919,8 @@ def _coupled_rotation_signs( "prepare_trellis256_moe_weights", "prepare_trellis256_dense_weight", "prepare_trellis256_pair_dense_weight", + "prepare_qsrt_atom_moe_weights", + "prepare_qsrt_atom_v2_moe_weights", "prepare_w4a16_compressed_tensors_weights", "prepare_w4a16_e8m0_native_weights", "prepare_w4a16_fc2_e8m0_weights", diff --git a/b12x/moe/_shared/kernels/w4a16/route_pack.py b/b12x/moe/_shared/kernels/w4a16/route_pack.py index faaad8817..3156ff040 100644 --- a/b12x/moe/_shared/kernels/w4a16/route_pack.py +++ b/b12x/moe/_shared/kernels/w4a16/route_pack.py @@ -6,6 +6,7 @@ import triton.language as tl import torch +from b12x._lib.env import env_flag from b12x.moe._shared.kernels.w4a16.host import route_pack_capacity @@ -17,6 +18,9 @@ _FAST_COUNT_BLOCK_T = 1024 +_STABLE_SORT_MIN_ROUTES = 4096 +_STABLE_SORT_BLOCK_T = 4096 +_STABLE_SORT_EXPERTS_PER_PROGRAM = 1 @triton.jit @@ -285,6 +289,73 @@ def _pack_topk_routes_sort_kernel( tl.store(packed_route_indices + ranks, offsets, mask=valid) +@triton.jit +def _pack_topk_routes_stable_kernel( + topk_ids, + expert_map, + packed_route_indices, + expert_offsets, + live_numel, + NUMEL_CAPACITY: tl.constexpr, + NUM_EXPERTS: tl.constexpr, + HAS_EXPERT_MAP: tl.constexpr, + BLOCK_T: tl.constexpr, + EXPERTS_PER_PROGRAM: tl.constexpr, +): + """Pack each expert's routes in ascending token-major route order.""" + expert_ids = ( + tl.program_id(0) * EXPERTS_PER_PROGRAM + + tl.arange(0, EXPERTS_PER_PROGRAM) + ) + expert_mask = expert_ids < NUM_EXPERTS + output_starts = tl.load( + expert_offsets + expert_ids, + mask=expert_mask, + other=0, + ) + output_counts = tl.zeros((EXPERTS_PER_PROGRAM,), dtype=tl.int32) + lanes = tl.arange(0, BLOCK_T) + + for start in tl.range(0, NUMEL_CAPACITY, BLOCK_T): + offsets = start + lanes + raw_ids = tl.load( + topk_ids + offsets, + mask=offsets < live_numel, + other=-1, + ).to(tl.int32) + valid = ( + (offsets < live_numel) + & (raw_ids >= 0) + & (raw_ids < NUM_EXPERTS) + ) + ids = raw_ids + if HAS_EXPERT_MAP: + safe_ids = tl.minimum(tl.maximum(raw_ids, 0), NUM_EXPERTS - 1) + ids = tl.load(expert_map + safe_ids, mask=valid, other=-1).to( + tl.int32 + ) + valid = valid & (ids >= 0) & (ids < NUM_EXPERTS) + + matches = ( + expert_mask[:, None] + & valid[None, :] + & (expert_ids[:, None] == ids[None, :]) + ) + match_i32 = matches.to(tl.int32) + local_ranks = tl.cumsum(match_i32, axis=1) - 1 + output_indices = ( + output_starts[:, None] + + output_counts[:, None] + + local_ranks + ) + tl.store( + packed_route_indices + output_indices, + offsets[None, :], + mask=matches, + ) + output_counts += tl.sum(match_i32, axis=1) + + def pack_topk_routes_by_expert( topk_ids: torch.Tensor, block_size: int, @@ -509,17 +580,38 @@ def pack_topk_routes_by_expert( SEARCH_STEPS=block_e.bit_length(), num_warps=4, ) - _pack_topk_routes_sort_kernel[sort_grid]( - topk_ids, - expert_map_tensor, - packed_route_indices, - expert_offsets, - numel, - NUM_EXPERTS=int(num_experts), - HAS_EXPERT_MAP=expert_map is not None, - BLOCK_T=_SORT_BLOCK_T, - num_warps=4, - ) + if ( + env_flag("W4A16_STABLE_ROUTE_PACK") + and numel >= _STABLE_SORT_MIN_ROUTES + ): + stable_grid = ( + triton.cdiv(num_experts, _STABLE_SORT_EXPERTS_PER_PROGRAM), + ) + _pack_topk_routes_stable_kernel[stable_grid]( + topk_ids, + expert_map_tensor, + packed_route_indices, + expert_offsets, + numel, + NUMEL_CAPACITY=numel_capacity, + NUM_EXPERTS=int(num_experts), + HAS_EXPERT_MAP=expert_map is not None, + BLOCK_T=_STABLE_SORT_BLOCK_T, + EXPERTS_PER_PROGRAM=_STABLE_SORT_EXPERTS_PER_PROGRAM, + num_warps=8, + ) + else: + _pack_topk_routes_sort_kernel[sort_grid]( + topk_ids, + expert_map_tensor, + packed_route_indices, + expert_offsets, + numel, + NUM_EXPERTS=int(num_experts), + HAS_EXPERT_MAP=expert_map is not None, + BLOCK_T=_SORT_BLOCK_T, + num_warps=4, + ) return packed_route_indices, block_expert_ids, packed_route_count diff --git a/b12x/moe/_shared/kernels/w4a8_trellis_decode.py b/b12x/moe/_shared/kernels/w4a8_trellis_decode.py index aa300f299..fafb3dc24 100644 --- a/b12x/moe/_shared/kernels/w4a8_trellis_decode.py +++ b/b12x/moe/_shared/kernels/w4a8_trellis_decode.py @@ -21,7 +21,6 @@ packed_decode_sqg_xor_cheb_t12_to_e4m3x8, packed_decode_trellis_sqg_direct_lut_to_e4m3x8, ) -from b12x.moe._shared.kernels.trellis_ring import trellis256_lane_geom_bits @cute.jit @@ -134,11 +133,20 @@ def _w4a8_trellis_lane_geom(lane: Int32, bits: cutlass.Constexpr): ``ia``/``ib`` are the two ring word indices covering the lane's overlapping L16 windows and ``s2`` the merge shift; they depend only on - the lane id and the bitrate. The W4A8 kernels always decode the full - eight-weight span starting at the lane origin. + the lane id and the bitrate. """ - ia, ib, s2, _ = trellis256_lane_geom_bits(lane, 0, 8, bits) + bits_i32 = Int32(int(bits)) + ring_u32 = Int32(8 * int(bits)) + t_offset = Int32(8) * lane + b1 = (t_offset + Int32(257)) * bits_i32 + b0 = b1 - Int32(16) + b2 = b1 + Int32(7 * int(bits)) + i0 = b0 >> Int32(5) + i2 = (b2 - Int32(1)) >> Int32(5) + ia = i0 - ring_u32 * (i0 >= ring_u32).to(Int32) + ib = i2 - ring_u32 * (i2 >= ring_u32).to(Int32) + s2 = (i2 + Int32(1)) * Int32(32) - b2 return ia, ib, s2 diff --git a/b12x/moe/_shared/trellis_codebooks.py b/b12x/moe/_shared/trellis_codebooks.py deleted file mode 100644 index 82f0a61bf..000000000 --- a/b12x/moe/_shared/trellis_codebooks.py +++ /dev/null @@ -1,62 +0,0 @@ -"""Trellis codebook registry shared by preparation, planning, and kernels. - -A ``trellis_t256`` tile stores 256 tail-biting codes whose 16-bit decode -windows are interpreted by one codebook. The codebook is a model-level -setting; three are defined: - -- ``mcg``: multiplicative congruential decode (multiplier ``0xCBAC1FED``) - with a lop3 mask/or into two added fp16 halves. -- ``sqg_e4m3``: XOR-Cheb-T12 bijection over the retained L16 history with a - frozen E4M3 reconstruction staircase; defined for K2/K3/K4. -- ``sqg_fp16``: D3L descriptor decode to fp16; defined for uniform K5/K6. - -This module is torch-free. Kernel modules embed the ids as compile-time -constants, so the ids participate in kernel cache keys. -""" - -from __future__ import annotations - -MCG = "mcg" -SQG_E4M3 = "sqg_e4m3" -SQG_FP16 = "sqg_fp16" - -CODEBOOKS: tuple[str, ...] = (MCG, SQG_E4M3, SQG_FP16) - -MCG_MULTIPLIER = 0xCBAC1FED -CODEBOOK_SENTINELS: dict[int, str] = {MCG_MULTIPLIER: MCG} - -def normalize_codebook(codebook: str | int) -> str: - """Return the canonical codebook id for ``codebook``. - - Integers are checkpoint sentinels (the MCG multiplier); strings are - matched case-insensitively against the canonical ids. - """ - - if isinstance(codebook, int): - normalized = CODEBOOK_SENTINELS.get(int(codebook) & 0xFFFFFFFF) - if normalized is None: - raise ValueError( - "unsupported trellis codebook sentinel " - f"{int(codebook) & 0xFFFFFFFF:#010x}; expected MCG 0xcbac1fed" - ) - return normalized - text = str(codebook).strip().lower() - if text in CODEBOOKS: - return text - raise ValueError( - f"unsupported trellis codebook {codebook!r}; expected " - "'mcg', 'sqg_e4m3', or 'sqg_fp16'" - ) - - -def validate_codebook_bits(codebook: str, bits: int) -> None: - """Reject (codebook, bitrate) pairs the decoders do not define. - - MCG decodes any supported tile bitrate; the SQG codebooks are defined - only on their construction ranges. - """ - - if codebook == SQG_E4M3 and bits not in (2, 3, 4): - raise ValueError("sqg_e4m3 is defined only for K2/K3/K4") - if codebook == SQG_FP16 and bits not in (5, 6): - raise ValueError("sqg_fp16 is defined only for uniform K5/K6") diff --git a/b12x/moe/ep_moe/_impl.py b/b12x/moe/ep_moe/_impl.py index e982b8a57..f86c26538 100644 --- a/b12x/moe/ep_moe/_impl.py +++ b/b12x/moe/ep_moe/_impl.py @@ -356,7 +356,6 @@ def bind( block_expert_ids=views["block_expert_ids"], packed_route_count=views["packed_route_count"], expert_offsets=views["expert_offsets"], - expert_counts=views["expert_counts"], ) @@ -431,7 +430,6 @@ def plan_ep_moe_scratch(caps: EPMoEScratchCaps) -> EPMoEScratchPlan: ("block_expert_ids", route_blocks, torch.int32), ("packed_route_count", 1, torch.int32), ("expert_offsets", caps.global_num_experts + 1, torch.int32), - ("expert_counts", caps.global_num_experts, torch.int32), ) ) scratch_specs = ( @@ -474,7 +472,6 @@ class EPMoEFP4Binding: block_expert_ids: torch.Tensor packed_route_count: torch.Tensor expert_offsets: torch.Tensor - expert_counts: torch.Tensor def run(self) -> torch.Tensor: return b12x_ep_moe_fp4(binding=self) @@ -517,7 +514,6 @@ def b12x_ep_moe_fp4(*, binding: EPMoEFP4Binding) -> torch.Tensor: block_expert_ids=binding.block_expert_ids, packed_route_count=binding.packed_route_count, expert_offsets=binding.expert_offsets, - expert_counts=binding.expert_counts, expert_map=binding.expert_map.tensor, apply_router_weight_on_input=binding.apply_router_weight_on_input, fast_math=binding.fast_math, diff --git a/b12x/moe/fused_moe/__init__.py b/b12x/moe/fused_moe/__init__.py index 14988a459..23bb57dae 100644 --- a/b12x/moe/fused_moe/__init__.py +++ b/b12x/moe/fused_moe/__init__.py @@ -1,10 +1,10 @@ """Fused tensor-parallel MoE for SM12x: route -> FC1 -> activation -> FC2 -> scatter, in one launch family. -Recipes (``META.recipes``) are arguments, not separate ops: nvfp4, -w4a8_mx, w4a8_nvfp4, w6a8_mx, w4a16 (weight layouts packed/modelopt plus -the TP-independent ``btx`` trellis container across the mcg, sqg_e4m3, and -sqg_fp16 codebooks). +Recipes (``META.recipes``) are arguments, not separate ops: nvfp4, mxfp4, +w4a8_mx, w4a8_nvfp4, w4a16 (weight layouts packed/modelopt and the +legacy ``exl3_trellis_mcg`` or TP-independent ``qsrt_sqg_e4m3`` +full-rotation sources, plus uniform-K5/K6 ``sqg_fp16_d3l`` W4A16). Activations: silu, situ, relu2, swigluoai_uninterleave. Kernel regimes (micro / dynamic / tiny-decode / w4a16) are selected declaratively by @@ -76,12 +76,13 @@ dtypes=("bf16", "fp16"), recipes=( "nvfp4", + "mxfp4", "w4a8_mx", "w4a8_nvfp4", - "w6a8_mx", "w4a16", - "w4a16/btx", - "w4a8_mx/btx", + "w4a16/exl3_trellis_mcg", + "w4a16/qsrt_sqg_e4m3", + "w4a16/sqg_fp16_d3l", ), requires=("triton",), provenance=Provenance( diff --git a/b12x/moe/fused_moe/_impl.py b/b12x/moe/fused_moe/_impl.py index 343061ac1..261c142ae 100644 --- a/b12x/moe/fused_moe/_impl.py +++ b/b12x/moe/fused_moe/_impl.py @@ -80,11 +80,6 @@ make_moe_spec, plan_moe_weight_preparation, ) -from b12x.moe._shared.execution import ( - _SOURCE_FORMATS as _EXECUTION_SOURCE_FORMATS, - _TRELLIS_SOURCE_FORMATS as _EXECUTION_TRELLIS_SOURCE_FORMATS, -) -from b12x.moe._shared.trellis_codebooks import SQG_E4M3 from b12x.moe._shared.tuning import lookup_max_active_clusters from b12x._lib.runtime_control import ( raise_if_kernel_resolution_frozen, @@ -97,8 +92,7 @@ logger = logging.getLogger(__name__) _B12X_TIMING = ( - os.getenv("B12X_TIMING", "0") == "1" - or os.getenv("VLLM_B12X_TIMING", "0") == "1" + os.getenv("B12X_TIMING", "0") == "1" or os.getenv("VLLM_B12X_TIMING", "0") == "1" ) _B12X_TIMING_THRESHOLD_MS = float( os.getenv( @@ -128,10 +122,25 @@ _DYNAMIC_W4A8_MATERIALIZED_ENV = "B12X_DYNAMIC_W4A8_MATERIALIZED" _W4A8_CONVERT_SCRATCH_MB_ENV = "B12X_W4A8_CONVERT_SCRATCH_MB" _W4A8_CONVERT_SCRATCH_MB_DEFAULT = 64 -# The source-format vocabulary is owned by b12x.moe._shared.execution; -# these views keep the historical local names. -_FP4_SOURCE_FORMATS = {name: name for name in _EXECUTION_SOURCE_FORMATS} -_TRELLIS_SOURCE_FORMATS = frozenset(_EXECUTION_TRELLIS_SOURCE_FORMATS) +_FP4_SOURCE_FORMATS = { + "modelopt_nvfp4": "modelopt_nvfp4", + "fp4_e8m0_k32": "fp4_e8m0_k32", + "compressed_tensors": "compressed_tensors", + # Packed MX-FP6 E2M3 codes with UE8M0 K/32 grids; exclusive to the + # w6a8_mx recipe (see _validate_fp4_source_format_for_quant_mode). + "mxfp6_e2m3": "mxfp6_e2m3", + # Native fixed-payload QSRT with finite E4M3 reconstruction. + "qsrt_sqg_e4m3": "qsrt_sqg_e4m3", + # Existing general EXL3 full-rotation MCG source contract. + "exl3_trellis_mcg": "exl3_trellis_mcg", + # Uniform K5/K6 SQG with the frozen 416-byte FP16-D3L scalar law. + "sqg_fp16_d3l": "sqg_fp16_d3l", +} +_TRELLIS_SOURCE_FORMATS = { + "exl3_trellis_mcg", + "qsrt_sqg_e4m3", + "sqg_fp16_d3l", +} _W4A16_SCALE_FORMATS = { "e4m3_k16": "e4m3_k16", "e4m3_k32": "e4m3_k32", @@ -144,7 +153,7 @@ # in-place W13 row rotation ("w13"), "gate_up" is already kernel-native. "up_gate": "w13", "gate_up": "w31", - "trellis_t256_proj": "trellis_t256_proj", + "trellis3_t256_proj": "trellis3_t256_proj", } _DEVICE_CAPABILITY_CACHE: dict[int, tuple[int, int]] = {} @@ -256,6 +265,7 @@ class TPW4A16Workspace: routed_rows_capacity: int intermediate_cache13: torch.Tensor intermediate_cache2: torch.Tensor + prefill_sum_accum: torch.Tensor | None fc1_c_tmp: torch.Tensor fc2_c_tmp: torch.Tensor packed_route_indices: torch.Tensor @@ -270,8 +280,7 @@ class TPW4A16Workspace: full_rotation: bool = False trellis_bits: int = 3 trellis_tile_config: tuple[int, int, int, int] | None = None - trellis_pair_kinds: frozenset[str] | None = None - trellis_codebook: str | None = None + qsrt_storage_format: str | None = None coupled_hadamard: bool = False route_block_size_m: int | None = None planned_token_counts: frozenset[int] = field(default_factory=frozenset) @@ -281,6 +290,7 @@ class TPW4A16Workspace: planned_swiglu_beta: float = 0.0 planned_scale_format: str = "e4m3_k16" planned_collect_activation_amax: bool = False + planned_prefill_fused_sum_fp32: bool = False planned_fused_moe_launches: dict[object, object] = field(default_factory=dict) planned_topk_sum_launches: dict[object, object] = field(default_factory=dict) # Mapped direct-route launches, keyed by exact live token count. These @@ -692,10 +702,10 @@ class _TPCoreWorkspacePlan: full_rotation: bool = False trellis_bits: int = 3 trellis_tile_config: tuple[int, int, int, int] | None = None - trellis_pair_kinds: frozenset[str] | None = None - trellis_codebook: str | None = None + qsrt_storage_format: str | None = None coupled_hadamard: bool = False route_block_size_m: int | None = None + prefill_fused_sum_fp32: bool = False tensor_specs: Tuple[_TensorAllocSpec, ...] = () @@ -1033,6 +1043,7 @@ class TPMoEFP4Binding: materialized_intermediate: torch.Tensor | None = None intermediate_cache13: torch.Tensor | None = None intermediate_cache2: torch.Tensor | None = None + prefill_sum_accum: torch.Tensor | None = None fc1_c_tmp: torch.Tensor | None = None fc2_c_tmp: torch.Tensor | None = None packed_route_indices: torch.Tensor | None = None @@ -1319,8 +1330,10 @@ def _normalize_fp4_source_format(source_format: str) -> str: return _FP4_SOURCE_FORMATS[source_format.lower()] except KeyError as exc: raise ValueError( - "source_format must be one of " - f"{sorted(_FP4_SOURCE_FORMATS)}, got {source_format!r}" + "source_format must be one of 'modelopt_nvfp4', " + "'fp4_e8m0_k32', 'compressed_tensors', 'mxfp6_e2m3', or " + "'exl3_trellis_mcg', 'qsrt_sqg_e4m3', or 'sqg_fp16_d3l', " + f"got {source_format!r}" ) from exc @@ -1367,10 +1380,10 @@ def _w4a16_weight_layout_for_source(source_format: str) -> str: just not auto-routed here.) """ source_format = _normalize_fp4_source_format(source_format) - return "trellis_t256" if source_format in _TRELLIS_SOURCE_FORMATS else "packed" + return "trellis3_t256" if source_format in _TRELLIS_SOURCE_FORMATS else "packed" -_W4A16_WEIGHT_LAYOUTS = {"packed", "modelopt", "trellis_t256"} +_W4A16_WEIGHT_LAYOUTS = {"packed", "modelopt", "trellis3_t256"} def _normalize_w4a16_weight_layout(weight_layout: str) -> str: @@ -1415,28 +1428,6 @@ def _normalize_swiglu_params( ) -def _derive_trellis_codebook( - source_format: str, trellis_codebook: str | None -) -> str | None: - """Kernel decode codebook, declared by the btx plan.""" - - del source_format - return trellis_codebook - - -def _fc_trellis_pair_kind(workspace) -> str | None: - """Compile-time pair kind from the plan's declared rate summary. - - Compact P33/P43 descriptors serve {P33, P43}; PDYNAMIC mode tables - serve {P33} and {P33, P24}. - """ - - kinds = workspace.trellis_pair_kinds - if kinds: - return "P33_P43" if "P43" in kinds else "PDYNAMIC" - return None - - def _validate_fp4_source_format_for_quant_mode( *, source_format: str, quant_mode: str ) -> None: @@ -1460,7 +1451,7 @@ def _validate_fp4_source_format_for_quant_mode( return if source_format == "fp4_e8m0_k32" and quant_mode == "w4a8_mx": return - if source_format == "btx" and quant_mode == "w4a8_mx": + if source_format == "qsrt_sqg_e4m3" and quant_mode == "w4a8_mx": return raise ValueError( f"source_format={source_format!r} with quant_mode={quant_mode!r} is " @@ -2412,6 +2403,7 @@ def _build_tp_moe_fp4_binding_from_views( routed_rows_capacity=plan.routed_rows, intermediate_cache13=tensors["intermediate_cache13"], intermediate_cache2=tensors["intermediate_cache2"], + prefill_sum_accum=tensors.get("prefill_sum_accum"), fc1_c_tmp=tensors["fc1_c_tmp"], fc2_c_tmp=tensors["fc2_c_tmp"], packed_route_indices=tensors["packed_route_indices"], @@ -2546,10 +2538,10 @@ def _plan_core_workspace( w4a16_block_size_m: int | None = None, trellis_bits: int = 3, trellis_tile_config: tuple[int, int, int, int] | None = None, - trellis_pair_kinds: frozenset[str] | None = None, - trellis_codebook: str | None = None, + qsrt_storage_format: str | None = None, coupled_hadamard: bool = False, apply_router_weight_on_input: bool = False, + collect_activation_amax: bool = False, deterministic_output: bool = False, swiglu_limit: float | None = None, swiglu_alpha: float | None = None, @@ -2569,6 +2561,7 @@ def _plan_core_workspace( from b12x.moe._shared.kernels.w4a16.host import ( max_packed_route_slots, packed_gemm_scratch_elements, + prefill_fused_sum_eligible, route_block_sizes_for_capacity, select_route_block_size_m, ) @@ -2590,14 +2583,12 @@ def _plan_core_workspace( else _w4a16_weight_layout_for_source(source_format) ) requested_route_E = int(route_num_experts or weight_E) - full_rotation = weight_layout == "trellis_t256" - # Route metadata must address every physical weight and every global - # router ID accepted through route_expert_map. - route_E = max(int(weight_E), requested_route_E) + full_rotation = weight_layout == "trellis3_t256" + route_E = requested_route_E if full_rotation else int(weight_E) if full_rotation: if source_format not in _TRELLIS_SOURCE_FORMATS: raise ValueError( - "trellis_t256 workspace requires a trellis source format" + "trellis3_t256 workspace requires an EXL3 trellis source format" ) if int(k) % 128 != 0 or int(n) % 128 != 0: raise ValueError( @@ -2607,13 +2598,33 @@ def _plan_core_workspace( if int(trellis_bits) not in (2, 3, 4, 5, 6): raise ValueError("trellis_bits must be one of 2, 3, 4, 5, 6") trellis_tile_config = trellis_tile_config or (64, 256, 64, 256) - if trellis_pair_kinds and ( - int(trellis_bits) != 3 or int(n) != 256 - ): + if qsrt_storage_format not in { + None, + "qsrt_atoms_v1", + "qsrt_atoms_v2", + }: raise ValueError( - "per-expert-pair btx extents require n=256 and " - "trellis_bits=3" + "qsrt_storage_format must be None, 'qsrt_atoms_v1', " + "or 'qsrt_atoms_v2'" ) + if qsrt_storage_format == "qsrt_atoms_v1" and ( + int(trellis_bits) != 3 or int(n) != 256 + ): + raise ValueError("QSRT atoms-v1 requires n=256 and trellis_bits=3") + if qsrt_storage_format == "qsrt_atoms_v2": + if int(trellis_bits) == 3 and int(n) != 256: + raise ValueError( + "the fixed high-rate atoms-v2 profile requires n=256" + ) + if int(trellis_bits) == 2 and int(n) % 128: + raise ValueError( + "the coupled pure-K2 atoms-v2 profile requires n divisible " + "by 128" + ) + if int(trellis_bits) not in (2, 3): + raise ValueError("QSRT atoms-v2 supports only K2 or K3 bases") + elif qsrt_storage_format is not None: + raise ValueError("qsrt_storage_format requires QSRT full rotation") routed_capacity = max(int(routed_rows), 1) fc1_cols = _activation_w1_rows(activation, int(n)) route_slots_capacity = 1 @@ -2757,10 +2768,22 @@ def _plan_core_workspace( // _dtype_nbytes(dtype), ) cache_dtype = torch.float16 if full_rotation else dtype + use_prefill_fused_sum = prefill_fused_sum_eligible( + dtype=dtype, + m=token_capacity, + full_rotation=full_rotation, + weight_layout=weight_layout, + collect_activation_amax=collect_activation_amax, + ) + intermediate_cache13_elements = ( + routed_capacity * fc1_cols + if use_prefill_fused_sum + else routed_capacity * max(fc1_cols, int(k)) + ) tensor_specs = [ _TensorAllocSpec( "intermediate_cache13", - (routed_capacity * max(fc1_cols, int(k)),), + (intermediate_cache13_elements,), cache_dtype, ), _TensorAllocSpec( @@ -2778,6 +2801,14 @@ def _plan_core_workspace( _TensorAllocSpec("expert_offsets", (route_E + 1,), torch.int32), _TensorAllocSpec("expert_counts", (route_E,), torch.int32), ] + if use_prefill_fused_sum: + tensor_specs.append( + _TensorAllocSpec( + "prefill_sum_accum", + (token_capacity * int(k),), + torch.float32, + ) + ) if full_rotation: max_tokens = max(routed_capacity // max(int(num_topk), 1), 1) tensor_specs.extend( @@ -2824,10 +2855,10 @@ def _plan_core_workspace( full_rotation=full_rotation, trellis_bits=int(trellis_bits), trellis_tile_config=trellis_tile_config, - trellis_pair_kinds=trellis_pair_kinds, - trellis_codebook=trellis_codebook, + qsrt_storage_format=qsrt_storage_format, coupled_hadamard=bool(coupled_hadamard), route_block_size_m=w4a16_block_size_m, + prefill_fused_sum_fp32=bool(use_prefill_fused_sum), tensor_specs=tuple(tensor_specs), ) @@ -3245,6 +3276,7 @@ def _materialize_workspace_from_core_arena( routed_rows_capacity=plan.routed_rows, intermediate_cache13=tensors["intermediate_cache13"], intermediate_cache2=tensors["intermediate_cache2"], + prefill_sum_accum=tensors.get("prefill_sum_accum"), fc1_c_tmp=tensors["fc1_c_tmp"], fc2_c_tmp=tensors["fc2_c_tmp"], packed_route_indices=tensors["packed_route_indices"], @@ -3263,10 +3295,10 @@ def _materialize_workspace_from_core_arena( full_rotation=plan.full_rotation, trellis_bits=plan.trellis_bits, trellis_tile_config=plan.trellis_tile_config, - trellis_pair_kinds=plan.trellis_pair_kinds, - trellis_codebook=plan.trellis_codebook, + qsrt_storage_format=plan.qsrt_storage_format, coupled_hadamard=plan.coupled_hadamard, route_block_size_m=plan.route_block_size_m, + planned_prefill_fused_sum_fp32=plan.prefill_fused_sum_fp32, volatile_launch_state=bool(volatile_launch_state), ) if a1_gscale is None or a2_gscale is None: @@ -4711,7 +4743,7 @@ def _w4a8_prepared_dict(prepared: object) -> dict[str, torch.Tensor]: if isinstance(prepared, dict): values = prepared - elif getattr(prepared, "weight_layout", None) == "trellis_t256": + elif getattr(prepared, "weight_layout", None) == "trellis3_t256": # Trellis-native representation: the rp slots carry the # projection-major payloads and the sfb slots carry the fp16 # boundary rotations. The dynamic launch keys trellis mode off @@ -5002,11 +5034,9 @@ def plan_b12x_fp4_moe_weights( w4a16_layout: PreparedWeightLayout | str | None = None, trellis_bits: int | None = None, trellis_tile_config: tuple[int, int, int, int] | None = None, + qsrt_storage_format: str | None = None, + qsrt_profile: str | None = None, coupled_hadamard: bool | None = None, - trellis_codebook: str | None = None, - trellis_rate_structure: str | None = None, - trellis_pair_kinds: Sequence[str] | frozenset[str] | None = None, - coupled_hadamard_blocks: tuple[int, int] | None = None, ) -> MoEWeightPreparationPlan: """Plan the one canonical weight allocation used by selected recipes.""" @@ -5034,11 +5064,9 @@ def plan_b12x_fp4_moe_weights( w4a16_layout=w4a16_layout, trellis_bits=trellis_bits, trellis_tile_config=trellis_tile_config, + qsrt_storage_format=qsrt_storage_format, + qsrt_profile=qsrt_profile, coupled_hadamard=coupled_hadamard, - trellis_codebook=trellis_codebook, - trellis_rate_structure=trellis_rate_structure, - trellis_pair_kinds=trellis_pair_kinds, - coupled_hadamard_blocks=coupled_hadamard_blocks, ) @@ -5054,8 +5082,17 @@ def prepare_b12x_fp4_moe_weights( w2_blockscale: torch.Tensor | None = None, a1_gscale: torch.Tensor | None = None, a2_gscale: torch.Tensor | None = None, - btx_layer: object | None = None, - btx_device: torch.device | str | None = None, + gate_suh: torch.Tensor | None = None, + up_suh: torch.Tensor | None = None, + intermediate_rotations: torch.Tensor | None = None, + down_svh: torch.Tensor | None = None, + trellis_mcg: torch.Tensor | int | None = None, + qsrt_atom_payload: torch.Tensor | None = None, + qsrt_first_atom_slot: int | None = None, + qsrt_layer_index: int | None = None, + qsrt_expert_ids: torch.Tensor | None = None, + qsrt_format_codes: torch.Tensor | None = None, + qsrt_rotation_draws: torch.Tensor | None = None, dummy_scale: torch.Tensor | None = None, ) -> B12XFP4ExpertWeights: """Transfer source tensors into the planner-selected runtime owner.""" @@ -5067,91 +5104,278 @@ def prepare_b12x_fp4_moe_weights( raise ValueError( f"params_dtype={actual_dtype!r} does not match plan dtype={plan.io_dtype!r}" ) - if plan.source_format == "btx": - from b12x.moe._shared.kernels.w4a16.btx import ( - BtxLayer, - prepare_btx_moe_weights, - ) - - if not isinstance(btx_layer, BtxLayer): - raise ValueError( - "btx preparation requires btx_layer, a BtxLayer extent from " - "read_btx_layer" - ) - if btx_device is None: - raise ValueError( - "btx preparation requires btx_device, the CUDA device that " - "owns the prepared weights" - ) + if plan.source_format == "qsrt_sqg_e4m3": if len(plan.quant_modes) != 1 or not plan.quant_modes <= frozenset( {"w4a16", "w4a8_mx"} ): raise ValueError( - "btx weights support exactly one of quant_mode='w4a16' or " + "QSRT weights support exactly one of quant_mode='w4a16' or " "quant_mode='w4a8_mx'" ) - manifest = btx_layer.manifest - if manifest.codebook != plan.trellis_codebook: + qsrt_quant_mode = next(iter(plan.quant_modes)) + if plan.qsrt_storage_format not in {"qsrt_atoms_v1", "qsrt_atoms_v2"}: raise ValueError( - f"btx manifest codebook {manifest.codebook!r} does not match " - f"the plan's trellis_codebook {plan.trellis_codebook!r}" + "QSRT preparation requires qsrt_atoms_v1 or qsrt_atoms_v2 storage" + ) + if plan.qsrt_storage_format == "qsrt_atoms_v2": + if not all( + value is not None + for value in ( + qsrt_atom_payload, + qsrt_first_atom_slot, + qsrt_layer_index, + gate_suh, + up_suh, + down_svh, + ) + ): + raise ValueError( + "QSRT atoms-v2 preparation requires the atom extent, its " + "physical identity, and all shared rotations" + ) + assert qsrt_atom_payload is not None + assert qsrt_first_atom_slot is not None + assert qsrt_layer_index is not None + assert gate_suh is not None and up_suh is not None + assert down_svh is not None + from b12x.moe._shared.kernels.w4a16.prepare import ( + prepare_qsrt_atom_v2_moe_weights, + ) + + value = prepare_qsrt_atom_v2_moe_weights( + qsrt_atom_payload, + first_atom_slot=qsrt_first_atom_slot, + layer_index=qsrt_layer_index, + profile=plan.qsrt_profile or "k3x22_k4x2", + rotation_draws=qsrt_rotation_draws, + hidden_size=plan.hidden_size, + intermediate_size=plan.intermediate_size, + num_experts=plan.num_experts, + activation=plan.activation, + gate_suh=gate_suh, + up_suh=up_suh, + down_svh=down_svh, + params_dtype=torch.float16, + codebook="sqg_xor_cheb_t12", + tile_config=plan.trellis_tile_config or (64, 256, 64, 256), + dummy_scale=dummy_scale, + ) + representation = _PreparedWeightRepresentation( + quant_mode=qsrt_quant_mode, + layout=PreparedWeightLayout.TRELLIS_NATIVE, + value=value, + ) + # Atom-v2 may be streamed from a host-backed canonical extent, but + # the returned expert package is a CUDA runtime object. Keep even + # the unit activation-scale placeholders colocated with the + # prepared weights so later binding and validation never inherit + # the source extent's device. + input_scale = torch.ones((), dtype=torch.float32, device=value.w13.device) + return B12XFP4ExpertWeights( + plan=plan, + a1_gscale=a1_gscale if a1_gscale is not None else input_scale, + w1_fp4=value.w13, + w1_blockscale=value.w13_scale, + w1_alphas=value.w13_global_scale, + a2_gscale=a2_gscale if a2_gscale is not None else input_scale, + w2_fp4=value.w2, + w2_blockscale=value.w2_scale, + w2_alphas=value.w2_global_scale, + representation=representation, ) - if ( - manifest.rates.bits is not None - and manifest.rates.bits != plan.trellis_bits - ): - raise ValueError( - f"btx manifest declares bits={manifest.rates.bits}; the plan " - f"declares trellis_bits={plan.trellis_bits}" + if not all( + value is not None + for value in ( + qsrt_atom_payload, + qsrt_first_atom_slot, + qsrt_layer_index, + qsrt_expert_ids, + qsrt_format_codes, + gate_suh, + up_suh, + down_svh, ) - if manifest.rates.structure != ( - plan.trellis_rate_structure or "uniform" ): raise ValueError( - "btx manifest rate structure " - f"{manifest.rates.structure!r} does not match the plan's " - f"{plan.trellis_rate_structure!r}" - ) - if (manifest.rates.pair_kinds or None) != plan.trellis_pair_kinds: + "QSRT preparation requires the atom extent, its physical-slot " + "metadata, expert IDs, format codes, and shared rotations" + ) + assert qsrt_atom_payload is not None + assert qsrt_first_atom_slot is not None + assert qsrt_layer_index is not None + assert qsrt_expert_ids is not None and qsrt_format_codes is not None + assert gate_suh is not None and up_suh is not None + assert down_svh is not None + h_side_shapes = ( + (plan.num_experts, plan.hidden_size), + (1, plan.hidden_size), + ) + for name, tensor, shapes in ( + ("gate_suh", gate_suh, h_side_shapes), + ("up_suh", up_suh, h_side_shapes), + ("down_svh", down_svh, h_side_shapes), + ): + if tensor.dtype != torch.float16: + raise TypeError(f"{name} must be torch.float16, got {tensor.dtype}") + if tuple(tensor.shape) not in shapes: + raise ValueError( + f"{name} must have shape {shapes}, got {tuple(tensor.shape)}" + ) + if tensor.device != qsrt_atom_payload.device: + raise ValueError( + f"{name} must be on {qsrt_atom_payload.device}, got {tensor.device}" + ) + if not tensor.is_contiguous(): + raise ValueError(f"{name} must be contiguous") + if (gate_suh.shape[0] == 1) != (up_suh.shape[0] == 1): raise ValueError( - "btx manifest pair kinds do not match the plan's " - "trellis_pair_kinds" + "gate_suh and up_suh must both be per-expert or both broadcast" ) - if manifest.hadamard.coupled != plan.coupled_hadamard: + tile_config = plan.trellis_tile_config or (64, 256, 64, 256) + # The atom extent may retain an aligned row stride from a sliced + # safetensors tensor. Take the placeholder from one contiguous inner + # byte span instead of requiring the complete extent to be viewable. + workspace_placeholder = ( + qsrt_atom_payload[0, 0, :4].view(torch.int32).reshape(-1) + ) + from b12x.moe._shared.kernels.w4a16.prepare import ( + prepare_qsrt_atom_moe_weights, + ) + + value = prepare_qsrt_atom_moe_weights( + qsrt_atom_payload, + first_atom_slot=qsrt_first_atom_slot, + layer_index=qsrt_layer_index, + expert_ids=qsrt_expert_ids, + format_codes=qsrt_format_codes, + hidden_size=plan.hidden_size, + intermediate_size=plan.intermediate_size, + num_experts=plan.num_experts, + activation=plan.activation, + gate_suh=gate_suh, + up_suh=up_suh, + down_svh=down_svh, + params_dtype=torch.float16, + codebook="sqg_xor_cheb_t12", + tile_config=tile_config, + dummy_scale=dummy_scale, + workspace=workspace_placeholder, + ) + representation = _PreparedWeightRepresentation( + quant_mode=qsrt_quant_mode, + layout=PreparedWeightLayout.TRELLIS_NATIVE, + value=value, + ) + input_scale = torch.ones((), dtype=torch.float32, device=value.w13.device) + return B12XFP4ExpertWeights( + plan=plan, + a1_gscale=a1_gscale if a1_gscale is not None else input_scale, + w1_fp4=value.w13, + w1_blockscale=value.w13_scale, + w1_alphas=value.w13_global_scale, + a2_gscale=a2_gscale if a2_gscale is not None else input_scale, + w2_fp4=value.w2, + w2_blockscale=value.w2_scale, + w2_alphas=value.w2_global_scale, + representation=representation, + ) + + if plan.source_format in {"exl3_trellis_mcg", "sqg_fp16_d3l"}: + is_d3l = plan.source_format == "sqg_fp16_d3l" + if plan.quant_modes != frozenset({"w4a16"}): + raise ValueError("Trellis weights support only quant_mode='w4a16'") + if w1_fp4 is None or w2_fp4 is None: + raise ValueError("Trellis preparation requires both weight payloads") + if not is_d3l and trellis_mcg is None: raise ValueError( - "btx manifest coupled-Hadamard declaration does not match " - "the plan" - ) - if ( - manifest.geometry.num_experts != plan.num_experts - or manifest.geometry.hidden_size != plan.hidden_size + "EXL3 Trellis preparation requires the checkpoint's trellis_mcg " + "marker; the runtime decoder is MCG-only and must fail closed" + ) + if not is_d3l: + assert trellis_mcg is not None + marker = ( + int(trellis_mcg.item()) + if isinstance(trellis_mcg, torch.Tensor) + else int(trellis_mcg) + ) & 0xFFFFFFFF + if marker != 0xCBAC1FED: + raise ValueError( + f"unexpected EXL3 MCG marker {marker:#010x}; expected 0xcbac1fed" + ) + elif trellis_mcg is not None: + raise ValueError("sqg_fp16_d3l must not carry an MCG marker") + if not all( + value is not None + for value in (gate_suh, up_suh, intermediate_rotations, down_svh) ): raise ValueError( - "btx manifest geometry does not match the plan's expert " - "count or hidden size" - ) - if btx_layer.local_intermediate_size != plan.intermediate_size: + "EXL3 Trellis preparation requires gate_suh, up_suh, " + "intermediate_rotations, and down_svh" + ) + assert gate_suh is not None and up_suh is not None + assert intermediate_rotations is not None and down_svh is not None + h_side_shapes = ( + (plan.num_experts, plan.hidden_size), + (1, plan.hidden_size), + ) + for name, tensor, shapes in ( + ("gate_suh", gate_suh, h_side_shapes), + ("up_suh", up_suh, h_side_shapes), + ( + "intermediate_rotations", + intermediate_rotations, + ((plan.num_experts, 3 * plan.intermediate_size),), + ), + ("down_svh", down_svh, h_side_shapes), + ): + if tensor.dtype != torch.float16: + raise TypeError(f"{name} must be torch.float16, got {tensor.dtype}") + if tuple(tensor.shape) not in shapes: + raise ValueError( + f"{name} must have shape {shapes}, got {tuple(tensor.shape)}" + ) + if tensor.device != w1_fp4.device: + raise ValueError( + f"{name} must be on {w1_fp4.device}, got {tensor.device}" + ) + if not tensor.is_contiguous(): + raise ValueError(f"{name} must be contiguous") + if (gate_suh.shape[0] == 1) != (up_suh.shape[0] == 1): raise ValueError( - f"btx extent covers {btx_layer.local_intermediate_size} " - "intermediate channels; the plan declares " - f"{plan.intermediate_size}" + "gate_suh and up_suh must both be per-expert or both broadcast" ) - value = prepare_btx_moe_weights( - btx_layer, + from b12x.moe._shared.kernels.w4a16.prepare import ( + prepare_trellis256_moe_weights, + ) + + tile_config = plan.trellis_tile_config or (64, 256, 64, 256) + value = prepare_trellis256_moe_weights( + w1_fp4, + w2_fp4, + hidden_size=plan.hidden_size, + intermediate_size=plan.intermediate_size, + num_experts=plan.num_experts, activation=plan.activation, - device=btx_device, params_dtype=torch.float16, - tile_config=plan.trellis_tile_config, + fc1_tile_n=tile_config[1], + fc2_tile_n=tile_config[3], + w13_layout="trellis3_t256_proj", + trellis_bits=plan.trellis_bits, dummy_scale=dummy_scale, + codebook="sqg_fp16_d3l" if is_d3l else "mcg", + gate_suh=gate_suh, + up_suh=up_suh, + intermediate_rotations=intermediate_rotations, + down_svh=down_svh, + tile_config=tile_config, + workspace=w1_fp4.view(torch.int32).reshape(-1)[:1], ) representation = _PreparedWeightRepresentation( - quant_mode=next(iter(plan.quant_modes)), + quant_mode="w4a16", layout=PreparedWeightLayout.TRELLIS_NATIVE, value=value, ) - input_scale = torch.ones( - (), dtype=torch.float32, device=value.w13.device - ) + input_scale = torch.ones((), dtype=torch.float32, device=w1_fp4.device) return B12XFP4ExpertWeights( plan=plan, a1_gscale=a1_gscale if a1_gscale is not None else input_scale, @@ -5164,6 +5388,7 @@ def prepare_b12x_fp4_moe_weights( w2_alphas=value.w2_global_scale, representation=representation, ) + if w1_fp4 is None or w2_fp4 is None: raise ValueError("weight preparation requires both w1_fp4 and w2_fp4") @@ -5548,7 +5773,7 @@ def _resolve_workspace_layout( # W4A8-MX sizes so compacted checkpoints never require a second # source-native copy merely to enter the tiny-decode band. if normalized_quant_mode == "w4a8_mx": - if weight_layout == "trellis_t256": + if weight_layout == "trellis3_t256": # Trellis-native serving: the direct micro trellis arm owns # the decode band; the dynamic kernel's w4a8_trellis recipe # serves every larger batch. The QMMA tiny-decode kernel @@ -6099,13 +6324,17 @@ def _validate_frozen_w4a16_launch( f"scale_format={scale_format!r}, " f"collect_activation_amax={bool(collect_activation_amax)}" ) + fused_reduces_routes = bool( + getattr(fused, "tc_decode_fused_sum", False) + or getattr(fused, "prefill_fused_sum_fp32", False) + ) has_topk_sum = planned_capacity in workspace.planned_topk_sum_launches if workspace.full_rotation: has_topk_sum = any( isinstance(key, tuple) and key[0] == planned_capacity for key in workspace.planned_topk_sum_launches ) - if not has_topk_sum: + if not fused_reduces_routes and not has_topk_sum: raise RuntimeError( "frozen W4A16 MoE workspace is missing its preplanned top-k sum launch " f"for capacity={planned_capacity}" @@ -6173,6 +6402,11 @@ def _w4a16_preplanned_launches( fused = workspace.planned_fused_moe_launches.get( (weight_layout, scale_format, planned_capacity, collect_activation_amax) ) + if fused is not None and bool( + getattr(fused, "tc_decode_fused_sum", False) + or getattr(fused, "prefill_fused_sum_fp32", False) + ): + return fused, None topk_sum_key: object = planned_capacity if workspace.full_rotation: topk_sum_key = ( @@ -6488,12 +6722,10 @@ def plan_tp_moe_arena_layout( w4a16_block_size_m=w4a16_block_size_m, trellis_bits=weight_plan.trellis_bits or 3, trellis_tile_config=weight_plan.trellis_tile_config, - trellis_pair_kinds=weight_plan.trellis_pair_kinds, - trellis_codebook=_derive_trellis_codebook( - weight_plan.source_format, weight_plan.trellis_codebook - ), + qsrt_storage_format=weight_plan.qsrt_storage_format, coupled_hadamard=weight_plan.coupled_hadamard, apply_router_weight_on_input=apply_router_weight_on_input, + collect_activation_amax=collect_activation_amax, deterministic_output=plan.deterministic_output, swiglu_limit=plan.swiglu_limit, swiglu_alpha=plan.swiglu_alpha, @@ -6545,7 +6777,10 @@ def _plan_full_rotation_w4a16_launches( if torch.cuda.is_current_stream_capturing(): raise RuntimeError("Trellis launch planning cannot run during capture") - from b12x.moe._shared.kernels.w4a16.host import route_pack_capacity + from b12x.moe._shared.kernels.w4a16.host import ( + route_pack_capacity, + route_pack_warmup_token_counts, + ) from b12x.moe._shared.kernels.w4a16.kernel import ( _DEFAULT_MAX_SHARED_MEM, compile_w4a16_fused_moe, @@ -6565,13 +6800,13 @@ def _plan_full_rotation_w4a16_launches( scale_format = _normalize_w4a16_scale_format( caps.w4a16_scale_format or _w4a16_scale_format_for_source(caps.source_format) ) - if weight_layout != "trellis_t256" or scale_format != "e4m3_k32": + if weight_layout != "trellis3_t256" or scale_format != "e4m3_k32": raise RuntimeError( "full-rotation Trellis launch planning requires " - "weight_layout='trellis_t256' and scale_format='e4m3_k32'; " + "weight_layout='trellis3_t256' and scale_format='e4m3_k32'; " f"got weight_layout={weight_layout!r}, scale_format={scale_format!r}" ) - w13_layout = "trellis_t256_proj" + w13_layout = "trellis3_t256_proj" _, capacity_route_slots, capacity_m_blocks = route_pack_capacity( capacity_tokens * core_plan.num_topk, block_size_m, @@ -6626,7 +6861,6 @@ def compile_fused(token_count: int) -> object: scale_format=scale_format, w13_layout=w13_layout, trellis_bits=core_plan.trellis_bits, - trellis_codebook=core_plan.trellis_codebook or SQG_E4M3, force_tile_config=core_plan.trellis_tile_config, intermediate_rotation=True, full_rotation=True, @@ -6658,68 +6892,80 @@ def compile_fused(token_count: int) -> object: for ids_dtype in (torch.int32, torch.int64) for mapped in (False, True) ) - for mapped in (False, True): - route_pack_key = ( - core_plan.device.type, - int(torch.cuda.current_device()), - capacity_tokens * core_plan.num_topk, - int(block_size_m), - int(core_plan.route_E), - mapped, - ) - if route_pack_key in _W4A16_ROUTE_PACK_PREWARMED: - continue + packed_route_indices = torch.empty( + capacity_route_slots, + dtype=torch.int32, + device=core_plan.device, + ) + block_expert_ids = torch.empty( + capacity_m_blocks, + dtype=torch.int32, + device=core_plan.device, + ) + packed_route_count = torch.empty( + 1, + dtype=torch.int32, + device=core_plan.device, + ) + expert_offsets = torch.empty( + core_plan.route_E + 1, + dtype=torch.int32, + device=core_plan.device, + ) + expert_counts = torch.empty( + core_plan.route_E, + dtype=torch.int32, + device=core_plan.device, + ) + pending_route_pack_keys: list[tuple[object, ...]] = [] + for route_ids_dtype in (torch.int32, torch.int64): dummy_topk_ids = torch.zeros( capacity_tokens, core_plan.num_topk, - dtype=torch.int32, + dtype=route_ids_dtype, device=core_plan.device, ) - packed_route_indices = torch.empty( - capacity_route_slots, - dtype=torch.int32, - device=core_plan.device, - ) - block_expert_ids = torch.empty( - capacity_m_blocks, - dtype=torch.int32, - device=core_plan.device, - ) - packed_route_count = torch.empty( - 1, - dtype=torch.int32, - device=core_plan.device, - ) - expert_offsets = torch.empty( - core_plan.route_E + 1, - dtype=torch.int32, - device=core_plan.device, - ) - expert_counts = torch.empty( - core_plan.route_E, - dtype=torch.int32, - device=core_plan.device, + for mapped in (False, True): + expert_map = None + if mapped: + expert_map = torch.arange( + core_plan.route_E, + dtype=torch.int32, + device=core_plan.device, + ) + for token_count in route_pack_warmup_token_counts(capacity_tokens): + route_pack_key = ( + core_plan.device.type, + int(torch.cuda.current_device()), + str(route_ids_dtype), + token_count * core_plan.num_topk, + int(block_size_m), + int(core_plan.route_E), + mapped, + ) + if route_pack_key in _W4A16_ROUTE_PACK_PREWARMED: + continue + pack_topk_routes_by_expert( + dummy_topk_ids[:token_count], + block_size_m, + core_plan.route_E, + expert_map=expert_map, + packed_route_indices=packed_route_indices, + block_expert_ids=block_expert_ids, + packed_route_count=packed_route_count, + expert_offsets=expert_offsets, + expert_counts=expert_counts, + ) + pending_route_pack_keys.append(route_pack_key) + torch.cuda.current_stream(core_plan.device).synchronize() + _W4A16_ROUTE_PACK_PREWARMED.update(pending_route_pack_keys) + if pending_route_pack_keys: + logger.info( + "Prewarmed %d full-rotation W4A16 route-pack variant(s) " + "for token capacity %d.", + len(pending_route_pack_keys), + capacity_tokens, ) - expert_map = None - if mapped: - expert_map = torch.arange( - core_plan.route_E, - dtype=torch.int32, - device=core_plan.device, - ) - pack_topk_routes_by_expert( - dummy_topk_ids, - block_size_m, - core_plan.route_E, - expert_map=expert_map, - packed_route_indices=packed_route_indices, - block_expert_ids=block_expert_ids, - packed_route_count=packed_route_count, - expert_offsets=expert_offsets, - expert_counts=expert_counts, - ) - torch.cuda.current_stream(core_plan.device).synchronize() - _W4A16_ROUTE_PACK_PREWARMED.add(route_pack_key) return fused_launches, topk_sum_launches @@ -6826,12 +7072,10 @@ def plan_tp_moe_scratch(caps: TPMoEScratchCaps) -> TPMoEScratchPlan: w4a16_block_size_m=resolved_block_size_m, trellis_bits=caps.weight_plan.trellis_bits or 3, trellis_tile_config=caps.weight_plan.trellis_tile_config, - trellis_pair_kinds=caps.weight_plan.trellis_pair_kinds, - trellis_codebook=_derive_trellis_codebook( - caps.weight_plan.source_format, caps.weight_plan.trellis_codebook - ), + qsrt_storage_format=caps.weight_plan.qsrt_storage_format, coupled_hadamard=caps.weight_plan.coupled_hadamard, apply_router_weight_on_input=caps.apply_router_weight_on_input, + collect_activation_amax=caps.collect_activation_amax, deterministic_output=launch_plan.deterministic_output, swiglu_limit=launch_plan.swiglu_limit, swiglu_alpha=launch_plan.swiglu_alpha, @@ -6883,6 +7127,7 @@ def _prewarm_w4a16_planned_launches( weight_layout: str = "packed", w13_layout: str = "w13", collect_activation_amax: bool = False, + prefill_fused_sum: bool = False, ) -> None: """Resolve every W4A16 kernel shape owned by a frozen arena. @@ -6897,12 +7142,13 @@ def _prewarm_w4a16_planned_launches( weight_layout = _normalize_w4a16_weight_layout(weight_layout) full_rotation = bool(workspace.full_rotation) if full_rotation: - w13_layout = "trellis_t256_proj" + w13_layout = "trellis3_t256_proj" else: w13_layout = _normalize_w13_layout(w13_layout) collect_activation_amax = bool(collect_activation_amax) from b12x.moe._shared.kernels.w4a16.host import ( + prefill_fused_sum_eligible, route_pack_capacity, select_route_block_size_m, ) @@ -6965,6 +7211,14 @@ def _prewarm_w4a16_planned_launches( token_count, collect_activation_amax, ) + build_prefill_fused_sum = prefill_fused_sum_eligible( + dtype=element_dtype, + m=token_count, + full_rotation=full_rotation, + weight_layout=weight_layout, + collect_activation_amax=collect_activation_amax, + enabled=prefill_fused_sum, + ) # Rotation-table row count belongs to the prepared artifact, which # is not bound until after the workspace is frozen. Resolve both # full-rotation specializations now so either artifact contract is @@ -6992,10 +7246,24 @@ def _prewarm_w4a16_planned_launches( scale_format=scale_format, w13_layout=w13_layout, collect_activation_amax=collect_activation_amax, + prefill_fused_sum=bool(build_prefill_fused_sum), trellis_bits=workspace.trellis_bits, - trellis_codebook=workspace.trellis_codebook or SQG_E4M3, - fc1_trellis_pair_kind=_fc_trellis_pair_kind(workspace), - fc2_trellis_pair_kind=_fc_trellis_pair_kind(workspace), + fc1_trellis_pair_kind=( + "P33_P43" + if workspace.qsrt_storage_format == "qsrt_atoms_v2" + and workspace.trellis_bits == 3 + else "PDYNAMIC" + if workspace.qsrt_storage_format == "qsrt_atoms_v1" + else None + ), + fc2_trellis_pair_kind=( + "P33_P43" + if workspace.qsrt_storage_format == "qsrt_atoms_v2" + and workspace.trellis_bits == 3 + else "PDYNAMIC" + if workspace.qsrt_storage_format == "qsrt_atoms_v1" + else None + ), force_tile_config=workspace.trellis_tile_config, intermediate_rotation=full_rotation, full_rotation=full_rotation, @@ -7027,7 +7295,7 @@ def _prewarm_w4a16_planned_launches( topk_sum_launches[(token_count, ids_dtype, mapped)] = ( resolved_topk_sum ) - else: + elif not build_prefill_fused_sum: topk_sum_launches[token_count] = compile_w4a16_topk_sum( m=token_count, topk=workspace.num_topk, @@ -7059,6 +7327,7 @@ def _prewarm_w4a16_planned_launches( route_pack_key = ( workspace.device.type, int(torch.cuda.current_device()), + str(torch.int32), int(token_count) * int(workspace.num_topk), int(block_size_m), int(workspace.route_E), @@ -7157,11 +7426,21 @@ def _prewarm_w4a16_planned_launches( w13_layout=w13_layout, collect_activation_amax=False, trellis_bits=workspace.trellis_bits, - fc1_trellis_pair_kind=_fc_trellis_pair_kind( - workspace + fc1_trellis_pair_kind=( + "P33_P43" + if workspace.qsrt_storage_format == "qsrt_atoms_v2" + and workspace.trellis_bits == 3 + else "PDYNAMIC" + if workspace.qsrt_storage_format == "qsrt_atoms_v1" + else None ), - fc2_trellis_pair_kind=_fc_trellis_pair_kind( - workspace + fc2_trellis_pair_kind=( + "P33_P43" + if workspace.qsrt_storage_format == "qsrt_atoms_v2" + and workspace.trellis_bits == 3 + else "PDYNAMIC" + if workspace.qsrt_storage_format == "qsrt_atoms_v1" + else None ), force_tile_config=workspace.trellis_tile_config, direct_topk_routes=True, @@ -7327,12 +7606,10 @@ def materialize_tp_moe_arena_workspaces( w4a16_block_size_m=resolved_block_size_m, trellis_bits=weight_plan.trellis_bits or 3, trellis_tile_config=weight_plan.trellis_tile_config, - trellis_pair_kinds=weight_plan.trellis_pair_kinds, - trellis_codebook=_derive_trellis_codebook( - weight_plan.source_format, weight_plan.trellis_codebook - ), + qsrt_storage_format=weight_plan.qsrt_storage_format, coupled_hadamard=weight_plan.coupled_hadamard, apply_router_weight_on_input=apply_router_weight_on_input, + collect_activation_amax=collect_activation_amax, deterministic_output=plan.deterministic_output, swiglu_limit=plan.swiglu_limit, swiglu_alpha=plan.swiglu_alpha, @@ -7378,6 +7655,14 @@ def materialize_tp_moe_arena_workspaces( != w4a16_scale_format or bool(getattr(existing, "planned_collect_activation_amax", False)) != collect_activation_amax + or bool( + getattr( + existing, + "planned_prefill_fused_sum_fp32", + False, + ) + ) + != core_plan.prefill_fused_sum_fp32 ): pass else: @@ -7434,6 +7719,9 @@ def materialize_tp_moe_arena_workspaces( materialized.planned_swiglu_alpha = plan.swiglu_alpha materialized.planned_swiglu_beta = plan.swiglu_beta materialized.planned_collect_activation_amax = collect_activation_amax + materialized.planned_prefill_fused_sum_fp32 = ( + core_plan.prefill_fused_sum_fp32 + ) t_prewarm0 = time.perf_counter() if _B12X_TIMING else 0.0 _prewarm_w4a16_planned_launches( materialized, @@ -7446,6 +7734,7 @@ def materialize_tp_moe_arena_workspaces( weight_layout=w4a16_weight_layout, w13_layout=w13_layout, collect_activation_amax=collect_activation_amax, + prefill_fused_sum=core_plan.prefill_fused_sum_fp32, ) if _B12X_TIMING: prewarm_ms += (time.perf_counter() - t_prewarm0) * 1000.0 @@ -7727,6 +8016,7 @@ def build_tp_moe_fp4_binding( routed_rows_capacity=workspace.routed_rows_capacity, intermediate_cache13=workspace.intermediate_cache13, intermediate_cache2=workspace.intermediate_cache2, + prefill_sum_accum=workspace.prefill_sum_accum, fc1_c_tmp=workspace.fc1_c_tmp, fc2_c_tmp=workspace.fc2_c_tmp, packed_route_indices=workspace.packed_route_indices, @@ -8072,7 +8362,7 @@ def _get_micro_kernel( swiglu_alpha=swiglu_alpha, swiglu_beta=swiglu_beta, ) - micro_trellis = weight_layout == "trellis_t256" + micro_trellis = weight_layout == "trellis3_t256" if micro_trellis: micro_kwargs["weight_layout"] = weight_layout micro_kwargs["trellis_bits"] = int(trellis_bits) @@ -8137,6 +8427,7 @@ def dummy(dt): dummy(cutlass.BFloat16), # out_ptr barrier_fake, # barrier_count barrier_fake, # barrier_epoch + Int32(weight_E), # route_expert_limit Int32(compile_m), # m_val Int32(1), # grid_x current_cuda_stream(), # stream @@ -8552,21 +8843,15 @@ def __call__( # [E][K16(I)][N16(K)]. SFB tensors are compile-time dead # (identity UE8M0 word). _tr_bits = int(self._kernel.trellis_bits) - _tr_w13_u32 = ( - 2 * (self._k // 16) * ((self._w1_n // 2) // 16) * 8 * _tr_bits - ) + _tr_w13_u32 = 2 * (self._k // 16) * ((self._w1_n // 2) // 16) * 8 * _tr_bits _tr_down_u32 = (self._n // 16) * (self._k // 16) * 8 * _tr_bits w13_rp = cute.make_tensor( w13_rp_ptr, - layout=cute.make_layout( - (num_experts * _tr_w13_u32,), stride=(1,) - ), + layout=cute.make_layout((num_experts * _tr_w13_u32,), stride=(1,)), ) down_rp = cute.make_tensor( down_rp_ptr, - layout=cute.make_layout( - (num_experts * _tr_down_u32,), stride=(1,) - ), + layout=cute.make_layout((num_experts * _tr_down_u32,), stride=(1,)), ) _tr_sentinel = cute.make_layout((1,), stride=(1,)) w13_sfb_rp = cute.make_tensor(w13_sfb_rp_ptr, layout=_tr_sentinel) @@ -8700,9 +8985,7 @@ def __call__( row_counts.shape[0] * ( 6 - if getattr( - self._kernel, "trellis_coupled", False - ) + if getattr(self._kernel, "trellis_coupled", False) else 3 ) * self._n, @@ -9217,9 +9500,7 @@ def _launch_dynamic_flat( # A trellis-native w4a8 binding carries the fp16 boundary rotations in # the sfb slot (the QMMA repack stores int32 scale words there); the # trellis geometry is recovered from the operand extents. - w4a8_trellis = ( - _is_w4a8_quant_mode(quant_mode) and w13_sfb_rp.dtype == torch.float16 - ) + w4a8_trellis = _is_w4a8_quant_mode(quant_mode) and w13_sfb_rp.dtype == torch.float16 trellis_bits = 0 trellis_coupled = False trellis_lut_tensor = None @@ -9228,14 +9509,11 @@ def _launch_dynamic_flat( _trellis256_execution_lut, ) - trellis_lut_tensor = _trellis256_execution_lut(a.device, SQG_E4M3) + trellis_lut_tensor = _trellis256_execution_lut(a.device, "sqg_xor_cheb_t12") payload_u32 = w13_rp.numel() * w13_rp.element_size() // 4 window_words = 2 * E * (k // 16) * (n // 16) * 8 trellis_bits = payload_u32 // window_words - if ( - trellis_bits not in (2, 3, 4) - or trellis_bits * window_words != payload_u32 - ): + if trellis_bits not in (2, 3, 4) or trellis_bits * window_words != payload_u32: raise RuntimeError( "w4a8 trellis payload extent does not match the launch " f"shape (E={E}, k={k}, n={n}, payload_u32={payload_u32})" @@ -9896,7 +10174,7 @@ def _launch_compact_micro_flat( _trellis256_execution_lut, ) - trellis_lut = _trellis256_execution_lut(a.device, SQG_E4M3) + trellis_lut = _trellis256_execution_lut(a.device, "sqg_xor_cheb_t12") trellis_rotations = w1_scale_storage use_native_nvfp4_split = ( quant_mode == "w4a8_nvfp4" @@ -9975,7 +10253,7 @@ def _launch_compact_micro_flat( swiglu_limit=swiglu_limit, swiglu_alpha=swiglu_alpha, swiglu_beta=swiglu_beta, - weight_layout="trellis_t256" if w4a8_trellis else None, + weight_layout="trellis3_t256" if w4a8_trellis else None, trellis_bits=trellis_bits, trellis_coupled=trellis_coupled, ) @@ -10377,9 +10655,7 @@ def _launch_micro( # activation wrappers such as SiTU keep is_supported False to hold the # nvfp4-family compact path closed while serving the trellis arm. shape_supported = ( - MoEMicroKernelBackend.is_supported - if w4a8_trellis - else micro_cls.is_supported + MoEMicroKernelBackend.is_supported if w4a8_trellis else micro_cls.is_supported ) use_micro_direct = ( quant_mode == "nvfp4" or _is_w4a8_quant_mode(quant_mode) @@ -10653,7 +10929,7 @@ def b12x_moe_fp4(*, binding: TPMoEFP4Binding) -> torch.Tensor: raise RuntimeError( "the W4A16 weight plan did not materialize its required representation" ) - full_rotation = getattr(prepared, "weight_layout", "") == "trellis_t256" + full_rotation = getattr(prepared, "weight_layout", "") == "trellis3_t256" output_dtype = torch.float32 if full_rotation else a.dtype if output is None: if torch.cuda.is_current_stream_capturing(): @@ -10687,6 +10963,7 @@ def b12x_moe_fp4(*, binding: TPMoEFP4Binding) -> torch.Tensor: ) intermediate_cache13 = _require_binding_field(binding, "intermediate_cache13") intermediate_cache2 = _require_binding_field(binding, "intermediate_cache2") + prefill_sum_accum = binding.prefill_sum_accum fc1_c_tmp = _require_binding_field(binding, "fc1_c_tmp") fc2_c_tmp = _require_binding_field(binding, "fc2_c_tmp") packed_route_indices = _require_binding_field(binding, "packed_route_indices") @@ -10724,6 +11001,7 @@ def b12x_moe_fp4(*, binding: TPMoEFP4Binding) -> torch.Tensor: fast_math=fast_math, intermediate_cache13=intermediate_cache13, intermediate_cache2=intermediate_cache2, + prefill_sum_accum=prefill_sum_accum, output=scatter_output, fc1_c_tmp=fc1_c_tmp, fc2_c_tmp=fc2_c_tmp, @@ -10802,8 +11080,7 @@ def b12x_moe_fp4(*, binding: TPMoEFP4Binding) -> torch.Tensor: micro_w4a8_trellis = ( quant_mode == "w4a8_mx" - and getattr(prepared_payload, "weight_layout", None) - == "trellis_t256" + and getattr(prepared_payload, "weight_layout", None) == "trellis3_t256" ) if micro_w4a8_trellis: wv = _w4a8_trellis_weight_views( @@ -10993,8 +11270,7 @@ def b12x_moe_fp4(*, binding: TPMoEFP4Binding) -> torch.Tensor: activation in ("relu2", "silu") and m == 1 and a1_gscale.numel() == 1 - and os.environ.get("B12X_MICRO_SHARE_INPUT_ACROSS_EXPERTS", "1") - != "0" + and os.environ.get("B12X_MICRO_SHARE_INPUT_ACROSS_EXPERTS", "1") != "0" ), share_expert_scales=( activation in ("relu2", "silu") @@ -11353,9 +11629,7 @@ def _select_experts_reference( ) -def b12x_route_experts_fast( - *, binding: TPMoERouteBinding -) -> B12XTopKRouting: +def b12x_route_experts_fast(*, binding: TPMoERouteBinding) -> B12XTopKRouting: """Public sparse-routing entrypoint for higher-level integrations. This is the optimization seam for future fast routing work. The current diff --git a/benchmarks/benchmark_dense_mla_verify.py b/benchmarks/benchmark_dense_mla_verify.py new file mode 100644 index 000000000..b799b6c84 --- /dev/null +++ b/benchmarks/benchmark_dense_mla_verify.py @@ -0,0 +1,298 @@ +#!/usr/bin/env python3 +"""Compare deployed, row-specialized, and tiled K3 MLA verification plans.""" + +from __future__ import annotations + +import argparse +import json +import statistics +from dataclasses import dataclass + +import torch + +from b12x.attention import dense_mla + +FP8 = torch.float8_e4m3fn + + +@dataclass +class Arm: + name: str + plan: dense_mla.Plan + binding: dense_mla.Binding + graph: torch.cuda.CUDAGraph + output: torch.Tensor + lse: torch.Tensor + + +def _arguments() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--cache-tokens", type=int, default=131072) + parser.add_argument("--planned-cache-tokens", type=int, default=131072) + parser.add_argument("--page-size", type=int, default=1536) + parser.add_argument("--heads", type=int, default=96) + parser.add_argument("--query-len", type=int, default=4) + parser.add_argument("--samples", type=int, default=21) + parser.add_argument("--warmup", type=int, default=5) + parser.add_argument("--sparse-stride", type=int, default=1) + parser.add_argument("--sparse-min-tokens", type=int, default=32768) + parser.add_argument("--sparse-sink-chunks", type=int, default=8) + parser.add_argument("--sparse-recent-chunks", type=int, default=64) + parser.add_argument("--sparse-refresh-interval", type=int, default=0) + parser.add_argument("--split-candidates", default="") + parser.add_argument("--seed", type=int, default=20260901) + return parser.parse_args() + + +def _scratch(plan: dense_mla.Plan) -> torch.Tensor: + (spec,) = plan.scratch_specs() + return torch.empty(spec.shape, dtype=spec.dtype, device=spec.device) + + +def _make_plan( + *, + name: str, + device: torch.device, + heads: int, + query_len: int, + page_size: int, + planned_cache_tokens: int, + sparse_kwargs: dict[str, int], + max_splits: int | None, +) -> dense_mla.Plan: + planned_pages = (planned_cache_tokens + page_size - 1) // page_size + if name == "deployed": + mode = "decode" + max_total_q = 28 + max_batch = 28 + per_query_lens = False + elif name == "row_specialized": + mode = "decode" + max_total_q = query_len + max_batch = query_len + per_query_lens = False + elif name in ("tiled", "tiled_sparse") or name.startswith("tiled_s"): + mode = "verify" + max_total_q = query_len + max_batch = 1 + per_query_lens = True + else: + raise ValueError(f"Unknown benchmark arm: {name}") + return dense_mla.plan( + dense_mla.Caps( + device=device, + mode=mode, + kv_dtype=FP8, + num_q_heads=heads, + page_size=page_size, + max_total_q=max_total_q, + max_batch=max_batch, + max_cache_tokens=planned_cache_tokens, + max_page_table_width=planned_pages, + num_cache_pages=torch.iinfo(torch.int32).max, + use_cuda_graph=True, + uses_query_cache_seqlens=per_query_lens, + budget=( + dense_mla.Budget(max_splits=max_splits) + if max_splits is not None + else None + ), + **(sparse_kwargs if name == "tiled_sparse" else {}), + ) + ) + + +def _make_arm( + *, + name: str, + q: torch.Tensor, + cache: torch.Tensor, + page_table: torch.Tensor, + query_cache_lens: torch.Tensor, + q_scale: torch.Tensor, + kv_scale: torch.Tensor, + planned_cache_tokens: int, + sparse_kwargs: dict[str, int], + max_splits: int | None = None, +) -> Arm: + query_len, heads = int(q.shape[0]), int(q.shape[1]) + page_size = int(cache.shape[1]) + plan = _make_plan( + name=name, + device=q.device, + heads=heads, + query_len=query_len, + page_size=page_size, + planned_cache_tokens=planned_cache_tokens, + sparse_kwargs=sparse_kwargs, + max_splits=max_splits, + ) + output = torch.empty( + query_len, + heads, + 512, + dtype=torch.bfloat16, + device=q.device, + ) + if name in ("tiled", "tiled_sparse") or name.startswith("tiled_s"): + arm_table = page_table + cache_lens = query_cache_lens[-1:] + cu_seqlens_q = torch.tensor([0, query_len], dtype=torch.int32, device=q.device) + per_query_lens = query_cache_lens + else: + arm_table = page_table.expand(query_len, -1).contiguous() + cache_lens = query_cache_lens + cu_seqlens_q = torch.arange( + query_len + 1, + dtype=torch.int32, + device=q.device, + ) + per_query_lens = None + binding = dense_mla.bind( + plan, + scratch=_scratch(plan), + q=q, + kv_cache=cache, + output=output, + page_table=arm_table, + cache_seqlens=cache_lens, + query_cache_seqlens=per_query_lens, + cu_seqlens_q=cu_seqlens_q, + q_scale=q_scale, + kv_scale=kv_scale, + ) + dense_mla.compile(binding=binding) + actual, lse = dense_mla.run(binding=binding) + expected, expected_lse = dense_mla.reference( + q, + cache, + arm_table, + cache_lens, + cu_seqlens_q, + query_cache_seqlens=per_query_lens, + q_scale=q_scale, + kv_scale=kv_scale, + **(sparse_kwargs if name == "tiled_sparse" else {}), + ) + torch.cuda.synchronize(q.device) + cosine = torch.nn.functional.cosine_similarity( + actual.float().flatten(), + expected.float().flatten(), + dim=0, + ) + if float(cosine) <= 0.999 or not torch.isfinite(lse).all(): + raise RuntimeError(f"{name} correctness failed: cosine={float(cosine):.8f}") + torch.testing.assert_close(lse, expected_lse, rtol=2e-5, atol=2e-5) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + dense_mla.run(binding=binding) + return Arm(name, plan, binding, graph, actual, lse) + + +def main() -> None: + args = _arguments() + if args.query_len != 4: + raise SystemExit("This benchmark currently compares the fixed K3 q=4 path") + device = torch.device("cuda") + if torch.cuda.get_device_capability(device) not in ((12, 0), (12, 1)): + raise SystemExit("SM120/SM121 is required") + torch.manual_seed(args.seed) + pages = (args.cache_tokens + args.page_size - 1) // args.page_size + q_float = torch.randn(args.query_len, args.heads, 576, device=device) * 0.1 + cache_float = torch.randn(pages, args.page_size, 576, device=device) * 0.1 + q_scale = (q_float.abs().max() / 400).reshape(1).float() + kv_scale = (cache_float.abs().max() / 400).reshape(1).float() + q = (q_float / q_scale).to(FP8) + cache = (cache_float / kv_scale).to(FP8) + page_table = torch.arange(pages, dtype=torch.int32, device=device).view(1, -1) + query_cache_lens = torch.arange( + args.cache_tokens - args.query_len + 1, + args.cache_tokens + 1, + dtype=torch.int32, + device=device, + ) + sparse_kwargs = { + "sparse_stride": args.sparse_stride, + "sparse_min_tokens": args.sparse_min_tokens, + "sparse_sink_chunks": args.sparse_sink_chunks, + "sparse_recent_chunks": args.sparse_recent_chunks, + "sparse_refresh_interval": args.sparse_refresh_interval, + } + arm_names = ["deployed", "row_specialized", "tiled"] + split_candidates = [ + int(value) for value in args.split_candidates.split(",") if value.strip() + ] + arm_names.extend(f"tiled_s{value}" for value in split_candidates) + if args.sparse_stride > 1: + arm_names.append("tiled_sparse") + arms = [ + _make_arm( + name=name, + q=q, + cache=cache, + page_table=page_table, + query_cache_lens=query_cache_lens, + q_scale=q_scale, + kv_scale=kv_scale, + planned_cache_tokens=args.planned_cache_tokens, + sparse_kwargs=sparse_kwargs, + max_splits=( + int(name.removeprefix("tiled_s")) + if name.startswith("tiled_s") and name != "tiled_sparse" + else None + ), + ) + for name in arm_names + ] + for _ in range(args.warmup): + for arm in arms: + arm.graph.replay() + torch.cuda.synchronize(device) + + flush = torch.empty(256 * 1024 * 1024, dtype=torch.uint8, device=device) + timings: dict[str, list[float]] = {arm.name: [] for arm in arms} + for sample in range(args.samples): + ordered = arms[sample % len(arms) :] + arms[: sample % len(arms)] + for arm in ordered: + flush.zero_() + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + arm.graph.replay() + end.record() + end.synchronize() + timings[arm.name].append(float(start.elapsed_time(end))) + + medians = {name: statistics.median(values) for name, values in timings.items()} + baseline = medians["deployed"] + tiled_output = next(arm.output for arm in arms if arm.name == "tiled") + result = { + "device": torch.cuda.get_device_name(device), + "cache_tokens": args.cache_tokens, + "planned_cache_tokens": args.planned_cache_tokens, + "page_size": args.page_size, + "heads": args.heads, + "query_len": args.query_len, + "arms": { + arm.name: { + "query_tile": arm.plan.query_tile, + "num_splits": arm.plan.num_splits, + "median_ms": medians[arm.name], + "speedup_vs_deployed": baseline / medians[arm.name], + "cosine_vs_tiled": float( + torch.nn.functional.cosine_similarity( + arm.output.float().flatten(), + tiled_output.float().flatten(), + dim=0, + ) + ), + "raw_ms": timings[arm.name], + } + for arm in arms + }, + } + print(json.dumps(result, indent=2, sort_keys=True)) + + +if __name__ == "__main__": + main() diff --git a/benchmarks/benchmark_kimi_packed_mla.py b/benchmarks/benchmark_kimi_packed_mla.py new file mode 100644 index 000000000..9f4aef713 --- /dev/null +++ b/benchmarks/benchmark_kimi_packed_mla.py @@ -0,0 +1,164 @@ +"""Compare packed Kimi MLA decode at TP9 head counts and one-million-token capacity. + +Record CUDA-graph timings, output digests, and error against fp32 dense attention. +The cache keeps every visible token. Query-head padding changes only zero heads +that are removed before the DCP reduction. Run with an otherwise idle GPU. +""" + +from __future__ import annotations + +import argparse +import hashlib +import inspect +import json +import os +from pathlib import Path +import statistics +import sys + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +import torch + +from b12x.attention import sparse_mla +from b12x.attention._shared.mla.reference import ( + pack_mla_kv_cache_reference, + unpack_mla_kv_cache_reference, +) +from benchmarks.common import ( + bench_cuda_graph, + capture_cuda_graph, + make_l2_flush_fn, + nvidia_smi_gpu_mode_snapshot, +) + + +def digest(tensor): + return hashlib.sha256(tensor.contiguous().cpu().view(torch.uint8).numpy().tobytes()).hexdigest() + + +@torch.inference_mode() +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--output", type=Path, required=True) + parser.add_argument("--heads", type=int, choices=(104, 112), required=True) + parser.add_argument("--policy", choices=("static", "balanced"), required=True) + parser.add_argument("--partial", choices=("bf16", "fp32"), required=True) + parser.add_argument("--lengths", default="128,2048,8192,16384") + parser.add_argument("--seeds", default="42,43,44") + parser.add_argument("--amplitudes", default="0.25,1,4") + parser.add_argument("--capacity", type=int, default=116736) + parser.add_argument("--iters", type=int, default=40) + parser.add_argument("--flush-l2", action="store_true") + args = parser.parse_args() + if torch.cuda.get_device_capability()[0] != 12: + raise SystemExit("Requires SM12x") + torch.backends.cuda.matmul.allow_tf32 = False + rows, valid_heads, page = 4, 99, 1536 + sm_scale = 192 ** -0.5 + result = dict( + settings=vars(args) | {"output": str(args.output)}, + environment={key: os.environ.get(key) for key in ( + "B12X_MLA_SM120_GLM_FASTPATH", "B12X_MLA_SM120_GLM_W_HW_DEQUANT", + "B12X_MLA_SM120_BALANCED_WAVES", + )}, + gpu_before=nvidia_smi_gpu_mode_snapshot(), records=[], + ) + flush = make_l2_flush_fn(args.flush_l2) + for seed in map(int, args.seeds.split(",")): + for length in map(int, args.lengths.split(",")): + capacity = max(args.capacity, length) + physical = ((length + page - 1) // page) * page + gen = torch.Generator(device="cpu").manual_seed(seed + length) + # Generate only semantic heads before padding, so both geometries + # receive the same bytes under the same seed. + q_raw = torch.randn(rows, valid_heads, 576, generator=gen).to("cuda") + k_raw = torch.randn(physical, 512, generator=gen).to("cuda") + rope_raw = torch.randn(physical, 64, generator=gen).to("cuda") + lengths = torch.tensor([max(1, length - i) for i in range(rows)], device="cuda", dtype=torch.int32) + idx = torch.arange(capacity, device="cuda", dtype=torch.int32).repeat(rows, 1) + idx.masked_fill_(idx >= lengths[:, None], -1) + extra = {} + if "partial_dtype" in sparse_mla.Caps.__dataclass_fields__: + extra["partial_dtype"] = torch.float32 if args.partial == "fp32" else torch.bfloat16 + elif args.partial != "bf16": + raise RuntimeError("Reader does not support fp32 partials") + plan = sparse_mla.plan(sparse_mla.Caps( + device="cuda", num_q_heads=args.heads, max_q_rows=rows, + max_batch=rows, max_width=capacity, head_dim=576, v_head_dim=512, + dtype=torch.bfloat16, kv_dtype=torch.uint8, + max_chunks_per_row=64, page_size=page, head_major_output=False, + **extra, + )) + spec = plan.scratch_specs()[0] + scratch = torch.empty(spec.shape, dtype=spec.dtype, device="cuda") + q = torch.zeros(rows, args.heads, 576, dtype=torch.bfloat16, device="cuda") + binding = plan.bind( + scratch=scratch, q=q, selected_indices=idx, + cache_seqlens_int32=lengths, nsa_cache_seqlens_int32=lengths, + ) + kwargs = {} + if "split_policy" in inspect.signature(sparse_mla.run_decode).parameters: + kwargs["split_policy"] = args.policy + elif args.policy != "static": + raise RuntimeError("Reader does not support balanced splits") + for amplitude in map(float, args.amplitudes.split(",")): + q[:, :valid_heads].copy_((q_raw * amplitude).to(torch.bfloat16)) + packed = pack_mla_kv_cache_reference( + (k_raw * 0.25).to(torch.bfloat16), + (rope_raw * 0.25).to(torch.bfloat16), + ) + cache = packed.view(-1, page, 656) + last_outputs = {} + def run(bound=binding, kv=cache): + outputs = sparse_mla.run_decode( + binding=bound, kv_cache=kv, sm_scale=sm_scale, + v_head_dim=512, forced_num_splits=64, return_lse=True, + lse_scale="natural", **kwargs, + ) + last_outputs["value"] = outputs + return outputs + eager, eager_lse = run() + eager = eager[:, :valid_heads].clone() + eager_lse = eager_lse[:, :valid_heads].clone() + torch.cuda.synchronize() + graph = capture_cuda_graph(run, warmup=3) + samples = bench_cuda_graph(graph, replays=args.iters, l2_flush=flush)["replay_us"] + graph_out, graph_lse = last_outputs["value"] + assert torch.equal(eager, graph_out[:, :valid_heads]) + assert torch.equal(eager_lse, graph_lse[:, :valid_heads]) + again, again_lse = run() + assert torch.equal(eager, again[:, :valid_heads]) + assert torch.equal(eager_lse, again_lse[:, :valid_heads]) + assert torch.isfinite(eager).all() and torch.isfinite(eager_lse).all() + keys = unpack_mla_kv_cache_reference(packed).float().view(-1, 576) + expected = [] + expected_lse = [] + for row in range(rows): + visible = max(1, length - row) + scores = q[row, :valid_heads].float() @ keys[:visible].T * sm_scale + expected.append(torch.softmax(scores, dim=-1) @ keys[:visible, :512]) + expected_lse.append(torch.logsumexp(scores, dim=-1)) + expected = torch.stack(expected) + expected_lse = torch.stack(expected_lse) + error = eager.float() - expected + record = dict( + seed=seed, local_tokens=length, amplitude=amplitude, + output_sha256=digest(eager), lse_sha256=digest(eager_lse), + relative_l2=float(error.norm() / expected.norm()), + max_abs_error=float(error.abs().max()), + lse_max_abs_error=float((eager_lse - expected_lse).abs().max()), + graph_replay_us=samples, median_us=statistics.median(samples), + finite=True, eager_repeat_identical=True, + ) + result["records"].append(record) + print(json.dumps({k: v for k, v in record.items() if k != "graph_replay_us"}), flush=True) + del graph, eager, eager_lse, keys, expected, expected_lse, error, cache, packed + del scratch, binding, plan, q_raw, k_raw, rope_raw, q, idx + torch.cuda.empty_cache() + result["gpu_after"] = nvidia_smi_gpu_mode_snapshot() + args.output.write_text(json.dumps(result, indent=2) + "\n") + + +if __name__ == "__main__": + main() diff --git a/benchmarks/benchmark_moe.py b/benchmarks/benchmark_moe.py index 5c05c7bde..94f9528a0 100644 --- a/benchmarks/benchmark_moe.py +++ b/benchmarks/benchmark_moe.py @@ -342,6 +342,22 @@ class ModelProfile: top_k=8, ), ), + "kimi-k3-mxfp4-shape": ModelProfile( + label="Kimi-K3 MXFP4 TP16 (shape)", + checkpoint_family="kimi_k3_mxfp4_shape", + default_layer_idx=1, + tp_size=16, + hf_repo_id=None, + default_activation="situ", + default_quant_mode="w4a16", + default_validate="none", + shape=ShapeSpec( + hidden_size=7168, + intermediate_size=3072, + num_experts=896, + top_k=16, + ), + ), "dsv4f-nvfp4": ModelProfile( label="DSV4F NVFP4 (shape)", checkpoint_family="dsv4f_nvfp4_shape", @@ -803,11 +819,14 @@ def load_expert_weights( "nano35_w4a16_shape", "dsv4f_shape", "dsv4f_nvfp4_shape", + "kimi_k3_mxfp4_shape", "laguna_s21_shape", "minimax_m3_shape", }: shape_source_format = ( - "fp4_e8m0_k32" if checkpoint_family == "dsv4f_shape" else "modelopt_nvfp4" + "fp4_e8m0_k32" + if checkpoint_family in {"dsv4f_shape", "kimi_k3_mxfp4_shape"} + else "modelopt_nvfp4" ) return make_shape_only_expert_weights( spec, @@ -2818,6 +2837,15 @@ def bench_e2e() -> None: default=5, help="Repeat the timed measurement this many times per batch size and aggregate the results.", ) + parser.add_argument( + "--timing-json", + type=pathlib.Path, + default=None, + help=( + "write every CUDA-event timing sample and the benchmark command " + "to this JSON file" + ), + ) parser.add_argument("--batch-size-profile", choices=sorted(BATCH_SIZE_PROFILES), default="micro") parser.add_argument("--batch-sizes", type=int, nargs="+", default=None) parser.add_argument( @@ -2907,6 +2935,14 @@ def bench_e2e() -> None: ), ) parser.add_argument("--validate", choices=["none", "oracle"], default=None) + parser.add_argument( + "--compare-prefill-fused-sum", + action="store_true", + help=( + "compare W4A16 direct FP32 prefill route reduction with the " + "materialized-route reduction before timing" + ), + ) parser.add_argument( "--oracle-mode", choices=[ @@ -3028,6 +3064,10 @@ def bench_e2e() -> None: ) if args.force_mxfp4 and args.quant_mode != "w4a8_mx": raise ValueError("--force-mxfp4 requires --quant-mode w4a8_mx") + if args.compare_prefill_fused_sum and not use_w4a16: + raise ValueError( + "--compare-prefill-fused-sum requires --quant-mode w4a16" + ) if args.flashinfer_tune_max_num_tokens <= 0: raise ValueError("--flashinfer-tune-max-num-tokens must be positive") if ( @@ -3042,6 +3082,8 @@ def bench_e2e() -> None: raise ValueError("--quant-mode w4a16 currently supports --graph-mode single-op") if use_w4a16 and args.tp_parallel: raise ValueError("--quant-mode w4a16 currently does not support --tp-parallel") + if args.timing_json is not None and args.graph_mode != "single-op": + raise ValueError("--timing-json currently requires --graph-mode single-op") if args.graph_only and not args.cuda_graph: raise ValueError("--graph-only requires --cuda-graph") if args.routing_repeat_period < 0: @@ -3357,9 +3399,15 @@ def bench_e2e() -> None: print(f" {spec.tp_size} ranks done.") batch_results: dict[int, BatchResult] = {} + raw_timing_runs: dict[str, dict[str, dict[str, list[list[float]]]]] = {} accuracy_failures: list[str] = [] reference_warnings: list[str] = [] for batch_size in batch_sizes: + batch_timing_runs: dict[str, dict[str, list[list[float]]]] = { + "eager": {}, + "cuda_graph": {}, + } + raw_timing_runs[str(batch_size)] = batch_timing_runs print(f"\n{'=' * 70}") print(f" batch_size={batch_size} (tokens*top_k = {batch_size * spec.top_k} expert calls)") print(f"{'=' * 70}") @@ -3454,12 +3502,14 @@ def impl_launch(topk_ids_local: torch.Tensor, topk_weights_local: torch.Tensor) intermediate_cache13=backend_w4a16_buffers.intermediate_cache13, intermediate_cache2=backend_w4a16_buffers.intermediate_cache2, output=backend_output, + prefill_sum_accum=backend_w4a16_buffers.prefill_sum_accum, fc1_c_tmp=backend_w4a16_buffers.fc1_c_tmp, fc2_c_tmp=backend_w4a16_buffers.fc2_c_tmp, packed_route_indices=backend_w4a16_buffers.packed_route_indices, block_expert_ids=backend_w4a16_buffers.block_expert_ids, packed_route_count=backend_w4a16_buffers.packed_route_count, expert_offsets=backend_w4a16_buffers.expert_offsets, + expert_counts=backend_w4a16_buffers.expert_counts, **activation_params.kwargs(), ) assert backend_binding is not None @@ -3574,6 +3624,117 @@ def ref_launch() -> None: backend_out = backend_e2e().clone() torch.cuda.synchronize() + if args.compare_prefill_fused_sum: + env_name = "B12X_W4A16_PREFILL_FUSED_SUM" + original_value = os.environ.get(env_name) + try: + os.environ[env_name] = "0" + standard_buffers = make_backend_w4a16_buffers( + backend_w4a16_weights, + m=batch_size, + topk=spec.top_k, + dtype=torch.bfloat16, + device=device, + ) + standard_output_storage = torch.empty_like(x) + standard_output = w4a16_moe( + x, + backend_w4a16_weights, + topk_weights, + topk_ids, + activation=args.activation, + fast_math=args.fast_math, + intermediate_cache13=standard_buffers.intermediate_cache13, + intermediate_cache2=standard_buffers.intermediate_cache2, + output=standard_output_storage, + prefill_sum_accum=standard_buffers.prefill_sum_accum, + fc1_c_tmp=standard_buffers.fc1_c_tmp, + fc2_c_tmp=standard_buffers.fc2_c_tmp, + packed_route_indices=standard_buffers.packed_route_indices, + block_expert_ids=standard_buffers.block_expert_ids, + packed_route_count=standard_buffers.packed_route_count, + expert_offsets=standard_buffers.expert_offsets, + expert_counts=standard_buffers.expert_counts, + **activation_params.kwargs(), + ).clone() + os.environ[env_name] = "1" + fused_sum_buffers = make_backend_w4a16_buffers( + backend_w4a16_weights, + m=batch_size, + topk=spec.top_k, + dtype=torch.bfloat16, + device=device, + ) + fused_sum_output_storage = torch.empty_like(x) + fused_activation_kwargs = activation_params.kwargs() + + def run_fused_sum_comparison( + input_tensor=x, + prepared_weights=backend_w4a16_weights, + route_weights=topk_weights, + route_ids=topk_ids, + buffers=fused_sum_buffers, + output_storage=fused_sum_output_storage, + activation=args.activation, + fast_math=args.fast_math, + activation_kwargs=fused_activation_kwargs, + ) -> torch.Tensor: + return w4a16_moe( + input_tensor, + prepared_weights, + route_weights, + route_ids, + activation=activation, + fast_math=fast_math, + intermediate_cache13=buffers.intermediate_cache13, + intermediate_cache2=buffers.intermediate_cache2, + output=output_storage, + prefill_sum_accum=buffers.prefill_sum_accum, + fc1_c_tmp=buffers.fc1_c_tmp, + fc2_c_tmp=buffers.fc2_c_tmp, + packed_route_indices=buffers.packed_route_indices, + block_expert_ids=buffers.block_expert_ids, + packed_route_count=buffers.packed_route_count, + expert_offsets=buffers.expert_offsets, + expert_counts=buffers.expert_counts, + **activation_kwargs, + ) + + fused_sum_output = run_fused_sum_comparison().clone() + fused_sum_repeat = run_fused_sum_comparison().clone() + torch.cuda.synchronize() + finally: + if original_value is None: + os.environ.pop(env_name, None) + else: + os.environ[env_name] = original_value + fused_sum_metrics = compare_to_reference( + fused_sum_output, standard_output + ) + fused_sum_repeat_metrics = compare_to_reference( + fused_sum_repeat, fused_sum_output + ) + print( + " " + + format_oracle_metrics( + "prefill fused sum vs standard", fused_sum_metrics + ) + ) + print( + " " + + format_oracle_metrics( + "prefill fused sum repeat", fused_sum_repeat_metrics + ) + ) + if ( + not math.isfinite(fused_sum_metrics.cos) + or fused_sum_metrics.cos < 0.9999 + ): + accuracy_failures.append( + f" bs={batch_size} prefill fused sum vs standard: " + f"cos={fused_sum_metrics.cos:.6f} < 0.999900" + ) + if ref_output is not None: ref_compare_metrics = compare_to_reference(backend_out, ref_output) print(f" {format_oracle_metrics(f'{backend_label} vs {ref_name}', ref_compare_metrics)}") @@ -3671,6 +3832,11 @@ def ref_launch() -> None: ) ref_stats = summarize_timing_runs(ref_runs_ms) if ref_runs_ms else None backend_stats = summarize_timing_runs(backend_runs_ms) + if ref_kernel_name is not None and ref_kernel_runs_ms: + batch_timing_runs["eager"][ref_kernel_name] = ref_kernel_runs_ms + if ref_name is not None and ref_runs_ms: + batch_timing_runs["eager"][ref_name] = ref_runs_ms + batch_timing_runs["eager"][backend_label] = backend_runs_ms ratio_nograph = RatioStats(ratio_runs) if ratio_runs else None if ref_kernel_stats is not None and ref_kernel_name is not None: @@ -3721,6 +3887,7 @@ def replay(g: torch.cuda.CUDAGraph = graph) -> None: ] stats = summarize_timing_runs(graph_runs) graph_stats_by_name[name] = stats + batch_timing_runs["cuda_graph"][name] = graph_runs print(f" {fmt_timing_stats(stats)}") except Exception as exc: print(f" FAILED ({type(exc).__name__}: {exc})") @@ -3892,6 +4059,25 @@ def tp_replay( for f in reference_warnings: print(f) print(f"{'=' * 70}\033[0m") + if args.timing_json is not None: + args.timing_json.parent.mkdir(parents=True, exist_ok=True) + args.timing_json.write_text( + json.dumps( + { + "schema": "b12x.moe-benchmark-timing.v1", + "command": sys.argv, + "model_profile": args.model_profile, + "quant_mode": args.quant_mode, + "warmup_iterations": args.warmup, + "timed_iterations_per_repeat": args.iters, + "repeats": args.repeats, + "samples_ms": raw_timing_runs, + }, + indent=2, + sort_keys=True, + ) + + "\n" + ) if accuracy_failures: print(f"\n\033[1;31m{'=' * 70}") print(" ACCURACY CHECK FAILED") diff --git a/docs/evidence/kimi_packed_mla_tp9.json b/docs/evidence/kimi_packed_mla_tp9.json new file mode 100644 index 000000000..6ff29d802 --- /dev/null +++ b/docs/evidence/kimi_packed_mla_tp9.json @@ -0,0 +1,1579 @@ +{ + "schema_version": 1, + "status": "qualified for recorded kernel cases; not an end-to-end serving guarantee", + "measured_date": "2026-09-08", + "baseline_b12x_revision": "04f246d7e5e6a9384717e2d2d68ae3c53b83aba3", + "candidate_b12x_revision": "0bf9f177b237", + "equivalent_attention_source_revision": "0edbaef99ffa6f03588e0ca46b4bd65a143ca3fb", + "container_image_id": "sha256:bb9843ca63fe61b258077a3231a4136f143f942e259676225446df030afda767", + "microbenchmark_source_sha256": "2d51b4eda52006fd630a4d2e71d4e1b22274c1b51f67bcb7d79780fa7357628c", + "precision_contract": "Static BF16-partial variants preserve all 99 effective output and LSE bytes. Balanced FP32 changes association and is research-only for serving.", + "variants": { + "base104": { + "settings": { + "heads": 104, + "policy": "static", + "partial": "bf16", + "lengths": "2048,8192,16384", + "seeds": "42", + "amplitudes": "0.25,4", + "capacity": 116736, + "iters": 40, + "flush_l2": true + }, + "environment": { + "B12X_MLA_SM120_GLM_FASTPATH": "0", + "B12X_MLA_SM120_GLM_W_HW_DEQUANT": "0", + "B12X_MLA_SM120_BALANCED_WAVES": null + }, + "gpu_before": { + "command": [ + "nvidia-smi", + "--query-gpu=index,uuid,pstate,persistence_mode,compute_mode,clocks.current.sm,clocks.current.memory,clocks_throttle_reasons.active,power.draw,power.limit,temperature.gpu", + "--format=csv,noheader,nounits" + ], + "available": true, + "fields": { + "index": "0", + "pstate": "P1", + "persistence_mode": "Disabled", + "compute_mode": "Default", + "clocks.current.sm": "2580", + "clocks.current.memory": "16365", + "clocks_throttle_reasons.active": "0x0000000000000000", + "power.draw": "71.95", + "power.limit": "325.00", + "temperature.gpu": "51" + } + }, + "records": [ + { + "seed": 42, + "local_tokens": 2048, + "amplitude": 0.25, + "output_sha256": "3ebde4c28a7a1f25327d2af638091c77750fe3c1eba3301cac6b8ac05d045f4b", + "lse_sha256": "3d28e9a8f531fe04f51418cb2b20595eb39c70981aeaa737697481c5cf4676ae", + "relative_l2": 0.006217378191649914, + "max_abs_error": 0.00027348927687853575, + "lse_max_abs_error": 0.00019359588623046875, + "graph_replay_us": [ + 266.2400007247925, + 281.5040051937103, + 295.80798745155334, + 296.00000381469727, + 300.31999945640564, + 299.1679906845093, + 299.77598786354065, + 295.80798745155334, + 296.9599962234497, + 296.83199524879456, + 299.9359965324402, + 296.86400294303894, + 295.8720028400421, + 296.86400294303894, + 296.8960106372833, + 296.86400294303894, + 297.91998863220215, + 294.8159873485565, + 297.91998863220215, + 295.77600955963135, + 295.8720028400421, + 296.86400294303894, + 295.8720028400421, + 300.31999945640564, + 297.91998863220215, + 296.8960106372833, + 294.9120104312897, + 294.9120104312897, + 293.88800263404846, + 294.9120104312897, + 294.9120104312897, + 294.9120104312897, + 294.9120104312897, + 293.88800263404846, + 294.9120104312897, + 294.9120104312897, + 294.9120104312897, + 294.9120104312897, + 293.88800263404846, + 296.9599962234497 + ], + "median_us": 295.8720028400421, + "finite": true, + "eager_repeat_identical": true + }, + { + "seed": 42, + "local_tokens": 2048, + "amplitude": 4.0, + "output_sha256": "389871c78abd776013070a3eb009e51c919cc50a41a938ab7d82ed979c03e478", + "lse_sha256": "b2f0cf9b100c318cc0979c43915280a4bc844657b93290533f8ca7b21fddfb02", + "relative_l2": 0.03469689562916756, + "max_abs_error": 0.013672813773155212, + "lse_max_abs_error": 0.020244598388671875, + "graph_replay_us": [ + 272.38398790359497, + 268.22400093078613, + 297.85600304603577, + 301.0239899158478, + 300.6719946861267, + 300.79999566078186, + 295.80798745155334, + 296.9279885292053, + 295.83999514579773, + 294.624000787735, + 301.31199955940247, + 294.7840094566345, + 297.91998863220215, + 296.86400294303894, + 295.83999514579773, + 296.83199524879456, + 297.95199632644653, + 294.8159873485565, + 295.83999514579773, + 294.8159873485565, + 297.91998863220215, + 294.8479950428009, + 297.91998863220215, + 296.86400294303894, + 295.8720028400421, + 296.83199524879456, + 294.9120104312897, + 294.9120104312897, + 296.9599962234497, + 294.9120104312897, + 296.9599962234497, + 294.9120104312897, + 294.9120104312897, + 295.9359884262085, + 294.9120104312897, + 294.9120104312897, + 294.9120104312897, + 294.9120104312897, + 295.9359884262085, + 294.9120104312897 + ], + "median_us": 295.83999514579773, + "finite": true, + "eager_repeat_identical": true + }, + { + "seed": 42, + "local_tokens": 8192, + "amplitude": 0.25, + "output_sha256": "b19b7095b6b0a163b37878386840a4ae5ed5b5d5a582cfe69fb265e299f1da52", + "lse_sha256": "ae301f4e5bfbdb0574206ecc0d2333ea98e421cd12a217087427bbbd1c734c03", + "relative_l2": 0.00999539252370596, + "max_abs_error": 0.00023821555078029633, + "lse_max_abs_error": 0.00011444091796875, + "graph_replay_us": [ + 274.30400252342224, + 272.2879946231842, + 284.64001417160034, + 299.00801181793213, + 299.0399897098541, + 300.03198981285095, + 303.1040132045746, + 299.96800422668457, + 298.7520098686218, + 299.9359965324402, + 301.15199089050293, + 299.96800422668457, + 298.94399642944336, + 299.96800422668457, + 298.880010843277, + 299.96800422668457, + 298.40001463890076, + 303.0399978160858, + 298.880010843277, + 299.9359965324402, + 298.911988735199, + 299.96800422668457, + 298.911988735199, + 300.9600043296814, + 298.911988735199, + 299.96800422668457, + 299.00801181793213, + 296.9599962234497, + 297.9840040206909, + 299.00801181793213, + 296.9599962234497, + 299.00801181793213, + 299.00801181793213, + 297.9840040206909, + 299.00801181793213, + 296.9599962234497, + 299.00801181793213, + 296.9599962234497, + 296.9599962234497, + 296.9599962234497 + ], + "median_us": 299.00801181793213, + "finite": true, + "eager_repeat_identical": true + }, + { + "seed": 42, + "local_tokens": 8192, + "amplitude": 4.0, + "output_sha256": "ff9353dd21b3f6bc097de0a8160dee26b67d2681cda2b0d8fb206324c0685292", + "lse_sha256": "ee004d75ecc902576c0f1a6a7fff174c0c590e8a0185aaae7c12b333e4be5a03", + "relative_l2": 0.03264153003692627, + "max_abs_error": 0.005813680589199066, + "lse_max_abs_error": 0.018121719360351562, + "graph_replay_us": [ + 273.4079957008362, + 268.3520019054413, + 289.66400027275085, + 302.65599489212036, + 303.0399978160858, + 298.94399642944336, + 299.9039888381958, + 298.94399642944336, + 306.11199140548706, + 300.00001192092896, + 300.03198981285095, + 299.8400032520294, + 299.96800422668457, + 299.8720109462738, + 300.9920120239258, + 300.03198981285095, + 299.9359965324402, + 299.8400032520294, + 299.00801181793213, + 299.9039888381958, + 299.96800422668457, + 298.911988735199, + 299.96800422668457, + 299.8720109462738, + 299.96800422668457, + 298.880010843277, + 296.9599962234497, + 296.9599962234497, + 299.00801181793213, + 296.9599962234497, + 296.9599962234497, + 299.00801181793213, + 296.9599962234497, + 297.9840040206909, + 296.9599962234497, + 299.00801181793213, + 296.9599962234497, + 296.9599962234497, + 299.00801181793213, + 299.00801181793213 + ], + "median_us": 299.00801181793213, + "finite": true, + "eager_repeat_identical": true + }, + { + "seed": 42, + "local_tokens": 16384, + "amplitude": 0.25, + "output_sha256": "5680071f914a7175e44ccbe27db25df5cdf3065aae7c2fea19a638598dc04dcc", + "lse_sha256": "bf9faa0d5e643c7f2d32c021d3ac7bc7b81c44c4b87f0536f50d654e930256bb", + "relative_l2": 0.01325127761811018, + "max_abs_error": 0.0003041177988052368, + "lse_max_abs_error": 8.678436279296875e-05, + "graph_replay_us": [ + 405.5039882659912, + 405.2799940109253, + 407.9039990901947, + 425.05601048469543, + 426.04801058769226, + 428.0000030994415, + 428.0959963798523, + 426.88000202178955, + 427.93598771095276, + 425.9839951992035, + 427.90400981903076, + 426.91200971603394, + 425.85599422454834, + 426.94398760795593, + 426.88000202178955, + 426.94398760795593, + 427.93598771095276, + 428.99200320243835, + 426.91200971603394, + 428.99200320243835, + 428.1280040740967, + 428.99200320243835, + 425.85599422454834, + 426.91200971603394, + 425.9839951992035, + 425.9200096130371, + 423.93600940704346, + 425.9839951992035, + 425.9839951992035, + 424.9599874019623, + 425.9839951992035, + 425.9839951992035, + 425.9839951992035, + 425.9839951992035, + 425.9839951992035, + 425.9839951992035, + 424.9599874019623, + 423.93600940704346, + 423.93600940704346, + 423.93600940704346 + ], + "median_us": 425.9839951992035, + "finite": true, + "eager_repeat_identical": true + }, + { + "seed": 42, + "local_tokens": 16384, + "amplitude": 4.0, + "output_sha256": "f72b7343a706caaf930dfeb0a2abeb04dffcd8eaa1af5eb939b4e3ccb2302fb1", + "lse_sha256": "03a573887330b774ab3e156399601cf326049f1121aa9ecc3e0df94ea6abcd20", + "relative_l2": 0.029885800555348396, + "max_abs_error": 0.004447534680366516, + "lse_max_abs_error": 0.018520355224609375, + "graph_replay_us": [ + 405.5039882659912, + 402.27198600769043, + 402.3360013961792, + 418.8799858093262, + 429.05598878860474, + 425.9839951992035, + 432.2560131549835, + 427.90400981903076, + 426.94398760795593, + 427.90400981903076, + 426.94398760795593, + 425.7279932498932, + 431.0399889945984, + 427.90400981903076, + 427.96799540519714, + 427.93598771095276, + 426.94398760795593, + 425.8880019187927, + 429.05598878860474, + 423.8080084323883, + 427.0719885826111, + 425.9200096130371, + 427.0080029964447, + 425.82398653030396, + 426.94398760795593, + 429.3760061264038, + 423.93600940704346, + 425.9839951992035, + 425.9839951992035, + 424.9599874019623, + 425.9839951992035, + 425.9839951992035, + 425.9839951992035, + 425.9839951992035, + 423.93600940704346, + 423.93600940704346, + 422.91200160980225, + 425.9839951992035, + 423.93600940704346, + 425.9839951992035 + ], + "median_us": 425.9839951992035, + "finite": true, + "eager_repeat_identical": true + } + ], + "gpu_after": { + "command": [ + "nvidia-smi", + "--query-gpu=index,uuid,pstate,persistence_mode,compute_mode,clocks.current.sm,clocks.current.memory,clocks_throttle_reasons.active,power.draw,power.limit,temperature.gpu", + "--format=csv,noheader,nounits" + ], + "available": true, + "fields": { + "index": "0", + "pstate": "P1", + "persistence_mode": "Disabled", + "compute_mode": "Default", + "clocks.current.sm": "2572", + "clocks.current.memory": "16365", + "clocks_throttle_reasons.active": "0x0000000000000000", + "power.draw": "87.54", + "power.limit": "325.00", + "temperature.gpu": "54" + } + } + }, + "fast104": { + "settings": { + "heads": 104, + "policy": "static", + "partial": "bf16", + "lengths": "2048,8192,16384", + "seeds": "42", + "amplitudes": "0.25,4", + "capacity": 116736, + "iters": 40, + "flush_l2": true + }, + "environment": { + "B12X_MLA_SM120_GLM_FASTPATH": "1", + "B12X_MLA_SM120_GLM_W_HW_DEQUANT": "0", + "B12X_MLA_SM120_BALANCED_WAVES": null + }, + "gpu_before": { + "command": [ + "nvidia-smi", + "--query-gpu=index,uuid,pstate,persistence_mode,compute_mode,clocks.current.sm,clocks.current.memory,clocks_throttle_reasons.active,power.draw,power.limit,temperature.gpu", + "--format=csv,noheader,nounits" + ], + "available": true, + "fields": { + "index": "0", + "pstate": "P1", + "persistence_mode": "Disabled", + "compute_mode": "Default", + "clocks.current.sm": "2580", + "clocks.current.memory": "16365", + "clocks_throttle_reasons.active": "0x0000000000000000", + "power.draw": "72.46", + "power.limit": "325.00", + "temperature.gpu": "51" + } + }, + "records": [ + { + "seed": 42, + "local_tokens": 2048, + "amplitude": 0.25, + "output_sha256": "3ebde4c28a7a1f25327d2af638091c77750fe3c1eba3301cac6b8ac05d045f4b", + "lse_sha256": "3d28e9a8f531fe04f51418cb2b20595eb39c70981aeaa737697481c5cf4676ae", + "relative_l2": 0.006217378191649914, + "max_abs_error": 0.00027348927687853575, + "lse_max_abs_error": 0.00019359588623046875, + "graph_replay_us": [ + 180.2240014076233, + 187.3600035905838, + 192.44800508022308, + 195.45599818229675, + 198.62399995326996, + 198.91199469566345, + 199.74400103092194, + 195.5520063638687, + 195.39199769496918, + 194.5600062608719, + 195.6160068511963, + 194.43200528621674, + 195.39199769496918, + 194.5600062608719, + 195.39199769496918, + 194.5600062608719, + 195.5520063638687, + 198.62399995326996, + 196.4160054922104, + 194.5600062608719, + 195.42400538921356, + 195.51999866962433, + 195.39199769496918, + 198.65599274635315, + 195.5520063638687, + 194.5600062608719, + 194.5600062608719, + 194.5600062608719, + 194.5600062608719, + 192.51200556755066, + 192.51200556755066, + 194.5600062608719, + 192.51200556755066, + 192.51200556755066, + 194.5600062608719, + 192.51200556755066, + 192.51200556755066, + 193.53599846363068, + 192.51200556755066, + 192.51200556755066 + ], + "median_us": 194.5600062608719, + "finite": true, + "eager_repeat_identical": true + }, + { + "seed": 42, + "local_tokens": 2048, + "amplitude": 4.0, + "output_sha256": "389871c78abd776013070a3eb009e51c919cc50a41a938ab7d82ed979c03e478", + "lse_sha256": "b2f0cf9b100c318cc0979c43915280a4bc844657b93290533f8ca7b21fddfb02", + "relative_l2": 0.03469689562916756, + "max_abs_error": 0.013672813773155212, + "lse_max_abs_error": 0.020244598388671875, + "graph_replay_us": [ + 184.32000279426575, + 186.14399433135986, + 194.4960057735443, + 195.5839991569519, + 196.383997797966, + 195.45599818229675, + 198.46400618553162, + 195.5839991569519, + 194.43200528621674, + 194.5600062608719, + 196.48000597953796, + 194.5600062608719, + 198.88000190258026, + 194.5600062608719, + 194.5600062608719, + 194.5600062608719, + 194.5279985666275, + 194.5600062608719, + 197.56799936294556, + 194.5600062608719, + 194.43200528621674, + 194.5600062608719, + 194.43200528621674, + 194.5600062608719, + 198.46400618553162, + 194.5600062608719, + 194.5600062608719, + 194.5600062608719, + 194.5600062608719, + 192.51200556755066, + 194.5600062608719, + 194.5600062608719, + 194.5600062608719, + 192.51200556755066, + 194.5600062608719, + 192.51200556755066, + 192.51200556755066, + 193.53599846363068, + 192.51200556755066, + 192.51200556755066 + ], + "median_us": 194.5600062608719, + "finite": true, + "eager_repeat_identical": true + }, + { + "seed": 42, + "local_tokens": 8192, + "amplitude": 0.25, + "output_sha256": "b19b7095b6b0a163b37878386840a4ae5ed5b5d5a582cfe69fb265e299f1da52", + "lse_sha256": "ae301f4e5bfbdb0574206ecc0d2333ea98e421cd12a217087427bbbd1c734c03", + "relative_l2": 0.00999539252370596, + "max_abs_error": 0.00023821555078029633, + "lse_max_abs_error": 0.00011444091796875, + "graph_replay_us": [ + 184.32000279426575, + 187.391996383667, + 193.40799748897552, + 199.5519995689392, + 198.62399995326996, + 199.64799284934998, + 202.62399315834045, + 203.19999754428864, + 198.55999946594238, + 196.60800695419312, + 198.55999946594238, + 198.65599274635315, + 196.99199497699738, + 198.65599274635315, + 198.55999946594238, + 196.60800695419312, + 198.5280066728592, + 198.65599274635315, + 197.50399887561798, + 201.664000749588, + 198.62399995326996, + 198.65599274635315, + 198.5280066728592, + 198.65599274635315, + 198.71999323368073, + 201.664000749588, + 198.65599274635315, + 198.65599274635315, + 196.60800695419312, + 196.60800695419312, + 197.63199985027313, + 198.65599274635315, + 196.60800695419312, + 196.60800695419312, + 196.60800695419312, + 196.60800695419312, + 196.60800695419312, + 198.65599274635315, + 196.60800695419312, + 196.60800695419312 + ], + "median_us": 198.55999946594238, + "finite": true, + "eager_repeat_identical": true + }, + { + "seed": 42, + "local_tokens": 8192, + "amplitude": 4.0, + "output_sha256": "ff9353dd21b3f6bc097de0a8160dee26b67d2681cda2b0d8fb206324c0685292", + "lse_sha256": "ee004d75ecc902576c0f1a6a7fff174c0c590e8a0185aaae7c12b333e4be5a03", + "relative_l2": 0.03264153003692627, + "max_abs_error": 0.005813680589199066, + "lse_max_abs_error": 0.018121719360351562, + "graph_replay_us": [ + 188.4160041809082, + 184.28799510002136, + 191.3599967956543, + 199.5519995689392, + 203.64800095558167, + 201.92000269889832, + 201.79200172424316, + 197.4399983882904, + 196.60800695419312, + 199.48799908161163, + 198.65599274635315, + 201.79200172424316, + 198.65599274635315, + 198.4959989786148, + 199.61600005626678, + 197.4399983882904, + 198.65599274635315, + 203.5840004682541, + 197.63199985027313, + 196.6399997472763, + 198.65599274635315, + 199.52000677585602, + 198.65599274635315, + 198.7839937210083, + 197.63199985027313, + 197.4399983882904, + 196.60800695419312, + 196.60800695419312, + 198.65599274635315, + 196.60800695419312, + 198.65599274635315, + 196.60800695419312, + 196.60800695419312, + 198.65599274635315, + 198.65599274635315, + 198.65599274635315, + 196.60800695419312, + 196.60800695419312, + 196.60800695419312, + 198.65599274635315 + ], + "median_us": 198.65599274635315, + "finite": true, + "eager_repeat_identical": true + }, + { + "seed": 42, + "local_tokens": 16384, + "amplitude": 0.25, + "output_sha256": "5680071f914a7175e44ccbe27db25df5cdf3065aae7c2fea19a638598dc04dcc", + "lse_sha256": "bf9faa0d5e643c7f2d32c021d3ac7bc7b81c44c4b87f0536f50d654e930256bb", + "relative_l2": 0.01325127761811018, + "max_abs_error": 0.0003041177988052368, + "lse_max_abs_error": 8.678436279296875e-05, + "graph_replay_us": [ + 279.32798862457275, + 275.32801032066345, + 274.399995803833, + 287.1359884738922, + 291.26399755477905, + 289.0560030937195, + 290.49599170684814, + 286.72000765800476, + 286.624014377594, + 286.72000765800476, + 290.3999984264374, + 287.6479923725128, + 286.72000765800476, + 286.72000765800476, + 285.5679988861084, + 286.72000765800476, + 290.71998596191406, + 285.6000065803528, + 286.624014377594, + 288.7679934501648, + 286.5920066833496, + 287.6479923725128, + 285.8879864215851, + 284.67199206352234, + 286.5920066833496, + 286.72000765800476, + 286.72000765800476, + 286.72000765800476, + 286.72000765800476, + 286.72000765800476, + 286.72000765800476, + 286.72000765800476, + 286.72000765800476, + 286.72000765800476, + 286.72000765800476, + 286.72000765800476, + 286.72000765800476, + 284.67199206352234, + 286.72000765800476, + 286.72000765800476 + ], + "median_us": 286.72000765800476, + "finite": true, + "eager_repeat_identical": true + }, + { + "seed": 42, + "local_tokens": 16384, + "amplitude": 4.0, + "output_sha256": "f72b7343a706caaf930dfeb0a2abeb04dffcd8eaa1af5eb939b4e3ccb2302fb1", + "lse_sha256": "03a573887330b774ab3e156399601cf326049f1121aa9ecc3e0df94ea6abcd20", + "relative_l2": 0.029885800555348396, + "max_abs_error": 0.004447534680366516, + "lse_max_abs_error": 0.018520355224609375, + "graph_replay_us": [ + 278.52800488471985, + 276.4799892902374, + 273.98398518562317, + 287.6800000667572, + 291.3280129432678, + 290.68800806999207, + 287.6160144805908, + 286.655992269516, + 287.6800000667572, + 284.7679853439331, + 286.72000765800476, + 288.672000169754, + 285.72800755500793, + 286.75198554992676, + 288.7040078639984, + 290.68800806999207, + 288.7679934501648, + 286.624014377594, + 286.72000765800476, + 288.672000169754, + 286.72000765800476, + 287.9680097103119, + 290.5920147895813, + 288.672000169754, + 286.72000765800476, + 286.5920066833496, + 286.72000765800476, + 286.72000765800476, + 286.72000765800476, + 286.72000765800476, + 286.72000765800476, + 286.72000765800476, + 286.72000765800476, + 286.72000765800476, + 286.72000765800476, + 286.72000765800476, + 286.72000765800476, + 286.72000765800476, + 284.67199206352234, + 286.72000765800476 + ], + "median_us": 286.72000765800476, + "finite": true, + "eager_repeat_identical": true + } + ], + "gpu_after": { + "command": [ + "nvidia-smi", + "--query-gpu=index,uuid,pstate,persistence_mode,compute_mode,clocks.current.sm,clocks.current.memory,clocks_throttle_reasons.active,power.draw,power.limit,temperature.gpu", + "--format=csv,noheader,nounits" + ], + "available": true, + "fields": { + "index": "0", + "pstate": "P1", + "persistence_mode": "Disabled", + "compute_mode": "Default", + "clocks.current.sm": "2580", + "clocks.current.memory": "16365", + "clocks_throttle_reasons.active": "0x0000000000000000", + "power.draw": "88.46", + "power.limit": "325.00", + "temperature.gpu": "51" + } + } + }, + "fast112": { + "settings": { + "heads": 112, + "policy": "static", + "partial": "bf16", + "lengths": "2048,8192,16384", + "seeds": "42", + "amplitudes": "0.25,4", + "capacity": 116736, + "iters": 40, + "flush_l2": true + }, + "environment": { + "B12X_MLA_SM120_GLM_FASTPATH": "1", + "B12X_MLA_SM120_GLM_W_HW_DEQUANT": "0", + "B12X_MLA_SM120_BALANCED_WAVES": null + }, + "gpu_before": { + "command": [ + "nvidia-smi", + "--query-gpu=index,uuid,pstate,persistence_mode,compute_mode,clocks.current.sm,clocks.current.memory,clocks_throttle_reasons.active,power.draw,power.limit,temperature.gpu", + "--format=csv,noheader,nounits" + ], + "available": true, + "fields": { + "index": "0", + "pstate": "P1", + "persistence_mode": "Disabled", + "compute_mode": "Default", + "clocks.current.sm": "2580", + "clocks.current.memory": "16365", + "clocks_throttle_reasons.active": "0x0000000000000000", + "power.draw": "72.32", + "power.limit": "325.00", + "temperature.gpu": "51" + } + }, + "records": [ + { + "seed": 42, + "local_tokens": 2048, + "amplitude": 0.25, + "output_sha256": "3ebde4c28a7a1f25327d2af638091c77750fe3c1eba3301cac6b8ac05d045f4b", + "lse_sha256": "3d28e9a8f531fe04f51418cb2b20595eb39c70981aeaa737697481c5cf4676ae", + "relative_l2": 0.006217378191649914, + "max_abs_error": 0.00027348927687853575, + "lse_max_abs_error": 0.00019359588623046875, + "graph_replay_us": [ + 99.32799637317657, + 100.80000013113022, + 108.31999778747559, + 108.47999900579453, + 109.43999886512756, + 109.53599959611893, + 111.51999980211258, + 109.40799862146378, + 108.5439994931221, + 109.37599837779999, + 109.50399935245514, + 109.37599837779999, + 108.5439994931221, + 109.40799862146378, + 109.50399935245514, + 108.44799876213074, + 109.47199910879135, + 109.40799862146378, + 108.5439994931221, + 112.31999844312668, + 108.5439994931221, + 109.37599837779999, + 109.50399935245514, + 109.40799862146378, + 108.5439994931221, + 112.5440001487732, + 108.5439994931221, + 108.5439994931221, + 106.49599879980087, + 108.5439994931221, + 106.49599879980087, + 106.49599879980087, + 108.5439994931221, + 107.51999914646149, + 106.49599879980087, + 108.5439994931221, + 108.5439994931221, + 108.5439994931221, + 106.49599879980087, + 106.49599879980087 + ], + "median_us": 108.5439994931221, + "finite": true, + "eager_repeat_identical": true + }, + { + "seed": 42, + "local_tokens": 2048, + "amplitude": 4.0, + "output_sha256": "389871c78abd776013070a3eb009e51c919cc50a41a938ab7d82ed979c03e478", + "lse_sha256": "b2f0cf9b100c318cc0979c43915280a4bc844657b93290533f8ca7b21fddfb02", + "relative_l2": 0.03469689562916756, + "max_abs_error": 0.013672813773155212, + "lse_max_abs_error": 0.020244598388671875, + "graph_replay_us": [ + 102.4319976568222, + 100.73599964380264, + 109.40799862146378, + 111.48799955844879, + 113.3119985461235, + 111.77600175142288, + 109.43999886512756, + 108.51199924945831, + 108.5439994931221, + 110.46399921178818, + 108.5439994931221, + 108.41599851846695, + 108.5439994931221, + 108.60799998044968, + 109.47199910879135, + 107.10400342941284, + 112.35199868679047, + 108.41599851846695, + 108.5439994931221, + 108.41599851846695, + 108.5439994931221, + 108.44799876213074, + 111.55200004577637, + 108.44799876213074, + 109.47199910879135, + 108.44799876213074, + 108.5439994931221, + 106.49599879980087, + 106.49599879980087, + 106.49599879980087, + 106.49599879980087, + 106.49599879980087, + 106.49599879980087, + 107.51999914646149, + 106.49599879980087, + 108.5439994931221, + 108.5439994931221, + 106.49599879980087, + 108.5439994931221, + 106.49599879980087 + ], + "median_us": 108.5279993712902, + "finite": true, + "eager_repeat_identical": true + }, + { + "seed": 42, + "local_tokens": 8192, + "amplitude": 0.25, + "output_sha256": "b19b7095b6b0a163b37878386840a4ae5ed5b5d5a582cfe69fb265e299f1da52", + "lse_sha256": "ae301f4e5bfbdb0574206ecc0d2333ea98e421cd12a217087427bbbd1c734c03", + "relative_l2": 0.00999539252370596, + "max_abs_error": 0.00023821555078029633, + "lse_max_abs_error": 0.00011444091796875, + "graph_replay_us": [ + 101.59999877214432, + 100.8640006184578, + 105.31199723482132, + 110.52799969911575, + 113.53600025177002, + 112.60800063610077, + 113.63200098276138, + 112.5440001487732, + 116.67200177907944, + 112.5440001487732, + 113.6000007390976, + 112.5440001487732, + 111.51999980211258, + 112.60800063610077, + 115.61600118875504, + 112.5440001487732, + 113.72800171375275, + 112.44799941778183, + 113.66400122642517, + 112.5440001487732, + 117.69600212574005, + 112.47999966144562, + 113.56800049543381, + 112.5440001487732, + 113.6000007390976, + 112.57600039243698, + 112.64000087976456, + 110.59200018644333, + 110.59200018644333, + 112.64000087976456, + 112.64000087976456, + 110.59200018644333, + 110.59200018644333, + 110.59200018644333, + 110.59200018644333, + 112.64000087976456, + 112.64000087976456, + 112.64000087976456, + 110.59200018644333, + 112.64000087976456 + ], + "median_us": 112.56000027060509, + "finite": true, + "eager_repeat_identical": true + }, + { + "seed": 42, + "local_tokens": 8192, + "amplitude": 4.0, + "output_sha256": "ff9353dd21b3f6bc097de0a8160dee26b67d2681cda2b0d8fb206324c0685292", + "lse_sha256": "ee004d75ecc902576c0f1a6a7fff174c0c590e8a0185aaae7c12b333e4be5a03", + "relative_l2": 0.03264153003692627, + "max_abs_error": 0.005813680589199066, + "lse_max_abs_error": 0.018121719360351562, + "graph_replay_us": [ + 104.44799810647964, + 102.55999863147736, + 109.40799862146378, + 114.9120032787323, + 113.53600025177002, + 113.53600025177002, + 113.53600025177002, + 112.57600039243698, + 117.21599847078323, + 112.5119999051094, + 113.6000007390976, + 113.50400000810623, + 113.6000007390976, + 112.5440001487732, + 115.64800143241882, + 112.57600039243698, + 113.6000007390976, + 115.55200070142746, + 113.66400122642517, + 112.5440001487732, + 117.66400188207626, + 112.5119999051094, + 113.63200098276138, + 113.50400000810623, + 111.61600053310394, + 112.5440001487732, + 112.64000087976456, + 112.64000087976456, + 110.59200018644333, + 110.59200018644333, + 112.64000087976456, + 112.64000087976456, + 111.61600053310394, + 112.64000087976456, + 112.64000087976456, + 112.64000087976456, + 112.64000087976456, + 112.64000087976456, + 112.38399893045425, + 112.64000087976456 + ], + "median_us": 112.64000087976456, + "finite": true, + "eager_repeat_identical": true + }, + { + "seed": 42, + "local_tokens": 16384, + "amplitude": 0.25, + "output_sha256": "5680071f914a7175e44ccbe27db25df5cdf3065aae7c2fea19a638598dc04dcc", + "lse_sha256": "bf9faa0d5e643c7f2d32c021d3ac7bc7b81c44c4b87f0536f50d654e930256bb", + "relative_l2": 0.01325127761811018, + "max_abs_error": 0.0003041177988052368, + "lse_max_abs_error": 8.678436279296875e-05, + "graph_replay_us": [ + 192.47999787330627, + 190.43199717998505, + 198.71999323368073, + 217.98400580883026, + 216.2880003452301, + 216.35200083255768, + 215.00800549983978, + 213.8880044221878, + 210.84800362586975, + 213.919997215271, + 214.11199867725372, + 213.919997215271, + 212.89600431919098, + 212.99199759960175, + 211.04000508785248, + 210.78400313854218, + 214.9759978055954, + 211.87199652194977, + 212.89600431919098, + 212.99199759960175, + 210.84800362586975, + 212.09600567817688, + 211.67999505996704, + 216.63999557495117, + 210.81599593162537, + 213.919997215271, + 210.94399690628052, + 210.94399690628052, + 212.99199759960175, + 209.9200040102005, + 210.94399690628052, + 212.99199759960175, + 210.94399690628052, + 212.99199759960175, + 208.8959962129593, + 211.96800470352173, + 212.99199759960175, + 210.94399690628052, + 210.94399690628052, + 210.94399690628052 + ], + "median_us": 212.0320051908493, + "finite": true, + "eager_repeat_identical": true + }, + { + "seed": 42, + "local_tokens": 16384, + "amplitude": 4.0, + "output_sha256": "f72b7343a706caaf930dfeb0a2abeb04dffcd8eaa1af5eb939b4e3ccb2302fb1", + "lse_sha256": "03a573887330b774ab3e156399601cf326049f1121aa9ecc3e0df94ea6abcd20", + "relative_l2": 0.029885800555348396, + "max_abs_error": 0.004447534680366516, + "lse_max_abs_error": 0.018520355224609375, + "graph_replay_us": [ + 194.5600062608719, + 197.2160041332245, + 208.8319957256317, + 213.8880044221878, + 213.05599808692932, + 214.1440063714981, + 215.7759964466095, + 217.53600239753723, + 212.8639966249466, + 213.95200490951538, + 210.87999641895294, + 212.92799711227417, + 216.8000042438507, + 210.94399690628052, + 210.7519954442978, + 211.90400421619415, + 212.8639966249466, + 212.89600431919098, + 214.9440050125122, + 211.90400421619415, + 210.84800362586975, + 213.95200490951538, + 212.76800334453583, + 212.92799711227417, + 213.82400393486023, + 216.09599888324738, + 210.94399690628052, + 212.99199759960175, + 212.99199759960175, + 211.96800470352173, + 210.94399690628052, + 212.99199759960175, + 210.94399690628052, + 210.94399690628052, + 210.94399690628052, + 211.96800470352173, + 212.99199759960175, + 208.8959962129593, + 210.94399690628052, + 210.94399690628052 + ], + "median_us": 212.8159999847412, + "finite": true, + "eager_repeat_identical": true + } + ], + "gpu_after": { + "command": [ + "nvidia-smi", + "--query-gpu=index,uuid,pstate,persistence_mode,compute_mode,clocks.current.sm,clocks.current.memory,clocks_throttle_reasons.active,power.draw,power.limit,temperature.gpu", + "--format=csv,noheader,nounits" + ], + "available": true, + "fields": { + "index": "0", + "pstate": "P1", + "persistence_mode": "Disabled", + "compute_mode": "Default", + "clocks.current.sm": "2580", + "clocks.current.memory": "16365", + "clocks_throttle_reasons.active": "0x0000000000000000", + "power.draw": "88.44", + "power.limit": "325.00", + "temperature.gpu": "51" + } + } + }, + "balanced112": { + "settings": { + "heads": 112, + "policy": "balanced", + "partial": "fp32", + "lengths": "2048,8192,16384", + "seeds": "42", + "amplitudes": "0.25,4", + "capacity": 116736, + "iters": 40, + "flush_l2": true + }, + "environment": { + "B12X_MLA_SM120_GLM_FASTPATH": "1", + "B12X_MLA_SM120_GLM_W_HW_DEQUANT": "0", + "B12X_MLA_SM120_BALANCED_WAVES": null + }, + "gpu_before": { + "command": [ + "nvidia-smi", + "--query-gpu=index,uuid,pstate,persistence_mode,compute_mode,clocks.current.sm,clocks.current.memory,clocks_throttle_reasons.active,power.draw,power.limit,temperature.gpu", + "--format=csv,noheader,nounits" + ], + "available": true, + "fields": { + "index": "0", + "pstate": "P1", + "persistence_mode": "Disabled", + "compute_mode": "Default", + "clocks.current.sm": "2580", + "clocks.current.memory": "16365", + "clocks_throttle_reasons.active": "0x0000000000000000", + "power.draw": "72.05", + "power.limit": "325.00", + "temperature.gpu": "51" + } + }, + "records": [ + { + "seed": 42, + "local_tokens": 2048, + "amplitude": 0.25, + "output_sha256": "1b5493677d3c04ce2b9f4f239f597ba8d6880795b8b25c06c715e7bf5144dda1", + "lse_sha256": "ad6a9727edb378616b33cc243751a7653c25f617516adb42514beb2e0ce63727", + "relative_l2": 0.006005485542118549, + "max_abs_error": 0.00027348927687853575, + "lse_max_abs_error": 0.00019359588623046875, + "graph_replay_us": [ + 38.91199827194214, + 40.12800008058548, + 38.88000175356865, + 39.8080013692379, + 38.88000175356865, + 41.85599833726883, + 39.16800022125244, + 39.77600112557411, + 38.784001022577286, + 39.872001856565475, + 38.84800150990486, + 39.8080013692379, + 41.600000113248825, + 39.872001856565475, + 38.784001022577286, + 39.872001856565475, + 38.816001266241074, + 39.872001856565475, + 40.863998234272, + 38.88000175356865, + 38.816001266241074, + 39.872001856565475, + 38.784001022577286, + 39.872001856565475, + 37.63199970126152, + 42.49599948525429, + 36.86400130391121, + 36.86400130391121, + 36.86400130391121, + 36.86400130391121, + 36.86400130391121, + 36.86400130391121, + 38.91199827194214, + 36.86400130391121, + 36.86400130391121, + 36.86400130391121, + 36.86400130391121, + 36.86400130391121, + 36.86400130391121, + 36.86400130391121 + ], + "median_us": 38.864001631736755, + "finite": true, + "eager_repeat_identical": true + }, + { + "seed": 42, + "local_tokens": 2048, + "amplitude": 4.0, + "output_sha256": "8f1b42261a2c49632cd1012bafa7a1ebf702ce527cc171cdfee634a45057e392", + "lse_sha256": "d52bda0696e739fae67ca834317c97490f5d6819046aecb0f43d09a1f031142e", + "relative_l2": 0.03465093672275543, + "max_abs_error": 0.01341409981250763, + "lse_max_abs_error": 0.020244598388671875, + "graph_replay_us": [ + 38.91199827194214, + 39.68000039458275, + 39.8080013692379, + 42.27200150489807, + 42.43199899792671, + 41.24800115823746, + 37.856001406908035, + 37.79200091958046, + 38.91199827194214, + 42.17600077390671, + 37.82400116324425, + 38.816001266241074, + 37.856001406908035, + 37.696000188589096, + 39.872001856565475, + 41.79200157523155, + 39.872001856565475, + 38.816001266241074, + 37.79200091958046, + 38.784001022577286, + 39.872001856565475, + 37.31200098991394, + 42.33599826693535, + 38.784001022577286, + 39.872001856565475, + 36.80000081658363, + 36.86400130391121, + 36.86400130391121, + 36.86400130391121, + 36.86400130391121, + 36.86400130391121, + 36.86400130391121, + 36.86400130391121, + 36.86400130391121, + 36.86400130391121, + 36.86400130391121, + 36.86400130391121, + 36.86400130391121, + 36.86400130391121, + 36.86400130391121 + ], + "median_us": 37.84000128507614, + "finite": true, + "eager_repeat_identical": true + }, + { + "seed": 42, + "local_tokens": 8192, + "amplitude": 0.25, + "output_sha256": "504caaf53f372436cc5f28a59e14f6a1c12873d6e6684e09f0170c9486420d04", + "lse_sha256": "432d47ffb0b45f27451d9d19575b6c822f3f8b9322a0e01106f9affdd9aa34d2", + "relative_l2": 0.009880737401545048, + "max_abs_error": 0.00023821555078029633, + "lse_max_abs_error": 0.00011444091796875, + "graph_replay_us": [ + 81.11999928951263, + 82.17599987983704, + 88.3840024471283, + 88.95999938249588, + 91.00800007581711, + 91.00800007581711, + 90.08000046014786, + 95.10400146245956, + 90.01599997282028, + 90.11200070381165, + 90.94399958848953, + 91.07200056314468, + 89.9839997291565, + 95.16800194978714, + 90.01599997282028, + 91.0400003194809, + 90.94399958848953, + 91.07200056314468, + 90.01599997282028, + 91.0400003194809, + 93.91999989748001, + 91.0400003194809, + 90.94399958848953, + 90.11200070381165, + 90.94399958848953, + 90.11200070381165, + 89.08800035715103, + 90.11200070381165, + 90.11200070381165, + 90.11200070381165, + 90.11200070381165, + 90.11200070381165, + 90.11200070381165, + 90.11200070381165, + 88.06400001049042, + 89.08800035715103, + 90.11200070381165, + 88.06400001049042, + 90.11200070381165, + 90.11200070381165 + ], + "median_us": 90.11200070381165, + "finite": true, + "eager_repeat_identical": true + }, + { + "seed": 42, + "local_tokens": 8192, + "amplitude": 4.0, + "output_sha256": "7aa28303df3b9c9713f2d8df59ff2b280a9ad3e7144910e153b15eece7119799", + "lse_sha256": "42cdfb5dfe4ad08dc87332d4404d3f3b53fc58c280b8f0e90a4f559e8b321d65", + "relative_l2": 0.03263235092163086, + "max_abs_error": 0.005813680589199066, + "lse_max_abs_error": 0.018120765686035156, + "graph_replay_us": [ + 82.94399827718735, + 82.75199681520462, + 88.28800171613693, + 95.90400010347366, + 93.98400038480759, + 95.67999839782715, + 92.12800115346909, + 91.07200056314468, + 91.16800129413605, + 95.16800194978714, + 90.08000046014786, + 91.20000153779984, + 92.06400066614151, + 92.16000139713287, + 89.9839997291565, + 95.13600170612335, + 90.01599997282028, + 92.0960009098053, + 92.22400188446045, + 92.16000139713287, + 92.06400066614151, + 93.12000125646591, + 94.7519987821579, + 93.12000125646591, + 89.9839997291565, + 92.06400066614151, + 90.11200070381165, + 90.11200070381165, + 92.16000139713287, + 90.11200070381165, + 90.11200070381165, + 90.11200070381165, + 90.11200070381165, + 90.11200070381165, + 90.11200070381165, + 91.13600105047226, + 90.11200070381165, + 90.11200070381165, + 90.11200070381165, + 90.11200070381165 + ], + "median_us": 91.10400080680847, + "finite": true, + "eager_repeat_identical": true + }, + { + "seed": 42, + "local_tokens": 16384, + "amplitude": 0.25, + "output_sha256": "b492a5e520b5867e13341124a8d0a0ed273822b8a60126f9201cd82e0ea1a205", + "lse_sha256": "bf9faa0d5e643c7f2d32c021d3ac7bc7b81c44c4b87f0536f50d654e930256bb", + "relative_l2": 0.01316456776112318, + "max_abs_error": 0.0002736002206802368, + "lse_max_abs_error": 8.678436279296875e-05, + "graph_replay_us": [ + 192.51200556755066, + 202.04800367355347, + 209.82399582862854, + 211.0079973936081, + 211.84000372886658, + 214.20800685882568, + 211.8079960346222, + 211.776003241539, + 209.82399582862854, + 211.776003241539, + 208.8959962129593, + 214.84799683094025, + 211.87199652194977, + 213.82400393486023, + 211.90400421619415, + 213.82400393486023, + 211.87199652194977, + 210.24000644683838, + 215.87200462818146, + 210.84800362586975, + 211.90400421619415, + 211.8079960346222, + 211.90400421619415, + 208.80000293254852, + 215.96799790859222, + 210.84800362586975, + 210.9760046005249, + 210.94399690628052, + 208.8959962129593, + 209.9200040102005, + 208.92800390720367, + 208.8959962129593, + 210.94399690628052, + 208.8959962129593, + 210.94399690628052, + 207.87200331687927, + 210.94399690628052, + 208.8959962129593, + 208.8959962129593, + 210.94399690628052 + ], + "median_us": 210.94399690628052, + "finite": true, + "eager_repeat_identical": true + }, + { + "seed": 42, + "local_tokens": 16384, + "amplitude": 4.0, + "output_sha256": "cd78a701d75674529021d80b45d18a94b635f637f99a7638e130e4796ecd6ffe", + "lse_sha256": "03a573887330b774ab3e156399601cf326049f1121aa9ecc3e0df94ea6abcd20", + "relative_l2": 0.02987416461110115, + "max_abs_error": 0.004447534680366516, + "lse_max_abs_error": 0.018520355224609375, + "graph_replay_us": [ + 192.54399836063385, + 189.31199610233307, + 208.99200439453125, + 210.84800362586975, + 215.10399878025055, + 210.91200411319733, + 213.919997215271, + 210.87999641895294, + 213.95200490951538, + 208.80000293254852, + 211.90400421619415, + 210.7519954442978, + 211.90400421619415, + 210.87999641895294, + 209.85600352287292, + 214.91199731826782, + 211.93599700927734, + 209.82399582862854, + 209.85600352287292, + 210.81599593162537, + 211.90400421619415, + 214.91199731826782, + 208.8959962129593, + 209.88799631595612, + 211.87199652194977, + 208.80000293254852, + 208.8959962129593, + 208.8959962129593, + 210.94399690628052, + 210.94399690628052, + 210.94399690628052, + 208.92800390720367, + 210.94399690628052, + 210.94399690628052, + 208.8959962129593, + 207.87200331687927, + 208.8959962129593, + 208.92800390720367, + 208.8959962129593, + 208.8959962129593 + ], + "median_us": 210.78399568796158, + "finite": true, + "eager_repeat_identical": true + } + ], + "gpu_after": { + "command": [ + "nvidia-smi", + "--query-gpu=index,uuid,pstate,persistence_mode,compute_mode,clocks.current.sm,clocks.current.memory,clocks_throttle_reasons.active,power.draw,power.limit,temperature.gpu", + "--format=csv,noheader,nounits" + ], + "available": true, + "fields": { + "index": "0", + "pstate": "P1", + "persistence_mode": "Disabled", + "compute_mode": "Default", + "clocks.current.sm": "2580", + "clocks.current.memory": "16365", + "clocks_throttle_reasons.active": "0x0000000000000000", + "power.draw": "88.16", + "power.limit": "325.00", + "temperature.gpu": "51" + } + } + } + }, + "limitations": [ + "Synthetic query/KV values; not model output equivalence for every prompt.", + "One physical GPU and sequential variant order.", + "No source-level sparse policy or KV precision reduction is used.", + "128 Ki serving baseline was invalid and supplies no speedup evidence." + ] +} diff --git a/docs/evidence/kimi_packed_mla_tp9.md b/docs/evidence/kimi_packed_mla_tp9.md new file mode 100644 index 000000000..63061791d --- /dev/null +++ b/docs/evidence/kimi_packed_mla_tp9.md @@ -0,0 +1,92 @@ +# Packed Kimi MLA verification on TP9 + +Status: **qualified for the recorded kernel cases**. The measurement record +`kimi_packed_mla_tp9.json` contains source identities, GPU operating state, +individual replay samples, and output/LSE digests from 2026-09-08. + +Kimi-K3 under TP9 gathers 99 effective query heads. The dense adapter's +eight-head padding gives 104 heads, while the packed reader executes 16-head +tiles. Six full tiles and a remainder require two launches. Padding to 112 +heads closes the tile; the adapter removes zero heads before DCP reduction. + +The recorded workload uses four query rows, a 116,736-token local capacity, +64 split slots, 1,536-token pages, and 656-byte packed KV records. Each variant +uses the same 99 semantic heads and the same generated query/KV values. +All visible KV tokens remain selected. + +| Reader | Padded heads | Split policy / partials | Correctness | +| --- | ---: | --- | --- | +| Reference source `04f246d7e5e6` | 104 | Static / BF16 | Reference digest | +| Vector shared-memory loads | 104 | Static / BF16 | Output and LSE bytes equal | +| Vector loads with whole head tiles | 112 | Static / BF16 | Output and LSE bytes equal | +| Vector loads with balanced work | 112 | Balanced / FP32 | Different association; recorded FP32-reference errors | + +At 2,048/8,192/16,384 local tokens and two query amplitudes, the six static +cases preserve output and LSE digests for all effective heads. The raw timing +samples are in the JSON record; they are isolated kernel measurements with +sequential arm ordering, not a universal model-throughput claim. The balanced +FP32 arm is **research-only for serving** and is excluded from the bit-identity +claim. Hardware residual dequantization is disabled in these comparisons. + +The ranges below are the minimum and maximum per-case replay medians across +the same six cases; the complete per-case samples remain in the JSON record. + +| Load implementation and heads | JSON variant | Median range, microseconds | +| --- | --- | ---: | +| Scalar loads, 104 heads | `base104` | 295.8–426.0 | +| Vector loads, 104 heads | `fast104` | 194.6–286.7 | +| Vector loads, 112 heads | `fast112` | 108.5–212.8 | + +The vector-104 arm isolates the head-padding comparison from the load-path +change. The scalar-104 arm measures their combined effect. These arms must +not be described as the same reference implementation. + +The candidate's `b12x/attention/_shared/mla/` and +`b12x/attention/sparse_mla/` source trees are identical to PR #311 revision +`0edbaef99ffa6f03588e0ca46b4bd65a143ca3fb`. This includes the S4 +return-state fix; the intermediate fast-path revision `242d6ca` is not a valid +comparison or deployment candidate. + +## Reproduction + +Run on an idle SM12x GPU using the desired source checkout and record its +revision. The benchmark retains raw replay samples, GPU state, and digests. + +```bash +B12X_MLA_SM120_GLM_FASTPATH=1 \ +B12X_MLA_SM120_GLM_W_HW_DEQUANT=0 \ +python benchmarks/benchmark_kimi_packed_mla.py \ + --heads 112 --policy static --partial bf16 \ + --lengths 2048,8192,16384 --seeds 42 --amplitudes 0.25,4 \ + --flush-l2 --output packed-112.json +``` + +Use 104 heads for the tail-launch comparison. Disable the fast path on the +reference source for the scalar-load comparison. Do not interpret timings +taken concurrently with model serving as an isolated-kernel result. + +`validation/attention/check_kimi_packed_mla_high_pages.py` independently checks +physical byte addressing beyond signed 32-bit range. It compares low pages +with page 2,133 at byte offset 2,149,244,928 using four rows and 112 padded +heads. Output and LSE are bit-identical in the recorded frozen-source run. +The script needs more than 2 GiB of free device memory. + +```bash +B12X_MLA_SM120_GLM_FASTPATH=1 \ +B12X_MLA_SM120_GLM_W_HW_DEQUANT=0 \ +python validation/attention/check_kimi_packed_mla_high_pages.py +``` + +## Serving boundary + +The serving composition pairs vLLM `fa6ea71c01fd`, B12X `0bf9f177b237`, and +LMCache `9f8514c680e7`. At approximately 64 Ki context, rank-zero target graph +duration changes from 35.35 to 30.85 ms and packed MLA kernel sum from 7.50 to +2.83 ms per step. The query-head tail disappears: 192 to 96 reader launches +over four steps and 24 layers. Kernel sums across streams are not critical-path +durations, and these traces use different generated continuations. + +The 128 Ki serving baseline is **invalid**: it contains no client output and +borrowed another request's global server counters. No 128 Ki model speedup is +claimed. The benchmark ownership correction is +[llm-inference-bench #16](https://github.com/local-inference-lab/llm-inference-bench/pull/16). diff --git a/tests/attention/test_attention_mla_sm120.py b/tests/attention/test_attention_mla_sm120.py index d7c3948ed..374a7ea9d 100644 --- a/tests/attention/test_attention_mla_sm120.py +++ b/tests/attention/test_attention_mla_sm120.py @@ -1229,6 +1229,7 @@ def _run_unified_glm( seed, num_heads=_GLM_NUM_HEADS, use_length_tensor=True, + split_policy="static", ): """Build a glm_ref GLM decode case and run the real unified launcher.""" from b12x.attention._shared.mla.kernel import run_unified_decode @@ -1266,6 +1267,7 @@ def _run_unified_glm( sm_scale=sm_scale, swa_page_size=_GLM_PAGE, forced_num_splits=forced_num_splits, + split_policy=split_policy, ) torch.cuda.synchronize() return out[0].float(), exp_O, min(forced_num_splits, n_chunks) @@ -1987,3 +1989,529 @@ def test_unified_prefill_glm_mixed_per_token_length_with_zero_row( f"neg_pad={neg_pad_past_len}) O cos={cos}" ) assert (got[t] - exp_O[t]).abs().max().item() < 3e-2 + + +# ── Split policy and partial precision (Kimi-K3 packed dense decode) ───────── +# +# The vLLM K3 adapter drives the GLM_NSA decode kernel as an exact-dense reader +# with 64 planned splits over a capacity-sized slot table. These tests cover the +# two knobs it uses: ``split_policy="balanced"`` (runtime chunk ranges derived +# from each row's live chunk count) and ``partial_dtype=torch.float32`` (exact +# split partials merged in fp32). + + +def _run_glm_multitoken( + device, + *, + topk, + num_tokens, + forced_num_splits, + seed, + split_policy="static", + partial_dtype=None, + lengths=None, + graph_lengths=None, + num_heads=_GLM_NUM_HEADS, + cache_hook=None, +): + """Run the unified GLM decode with per-token lengths and return + ``(out, expected, lengths)``. With ``graph_lengths`` the launch is captured + into a CUDA graph at ``lengths`` and replayed after the length tensor is + overwritten with ``graph_lengths``; the returned output and expectation then + correspond to ``graph_lengths``. ``cache_hook(kv_cache_flat, q, idx)`` may + edit the token-major cache in place before the kernel and the reference + read it.""" + from b12x.attention._shared.mla.kernel import run_unified_decode + from b12x.attention.sparse_mla._scratch import ( + B12XSparseMLAScratchCaps, + plan_sparse_mla_scratch, + ) + + nblk = max(1, (topk + _GLM_PAGE - 1) // _GLM_PAGE) + case = glm_ref.make_glm_decode_case( + num_heads=num_heads, topk=topk, num_tokens=num_tokens, num_blocks=nblk, + page_block_size=_GLM_PAGE, invalidate_half=False, seed=seed, device=device, + ) + q = case["q"].contiguous() + kv_cache_flat = case["kv_cache"].contiguous() # (slots, 1, 656) token-major + idx = case["topk_indices"].contiguous() + sm_scale = case["sm_scale"] + s_kv = kv_cache_flat.shape[0] + if cache_hook is not None: + cache_hook(kv_cache_flat, q, idx) + # The launcher addresses pages as block * page_bytes + slot * 656, so the + # token-major reference cache is the same bytes as a (blocks, page, 656) view. + kv_cache = kv_cache_flat.view(nblk, _GLM_PAGE, kv_cache_flat.shape[-1]) + if lengths is None: + lengths = _mixed_lengths(num_tokens, topk, device) + lengths = lengths.to(device=device, dtype=torch.int32).contiguous() + + n_chunks = (topk + 64 - 1) // 64 + caps = B12XSparseMLAScratchCaps( + device=device, num_q_heads=num_heads, max_q_rows=num_tokens, + max_batch=num_tokens, max_width=topk, max_kv_rows=s_kv, + head_dim=glm_ref.GLM_Q_HEAD_DIM, v_head_dim=glm_ref.GLM_D_V, + max_chunks_per_row=max(8, n_chunks, forced_num_splits), page_size=_GLM_PAGE, + partial_dtype=partial_dtype, + ) + plan = plan_sparse_mla_scratch(caps) + (spec,) = plan.scratch_specs() + storage = torch.zeros(spec.shape, dtype=spec.dtype, device=device) + cache_seqlens = torch.full((num_tokens,), s_kv, dtype=torch.int32, device=device) + binding = plan.bind( + scratch=storage, q=q, selected_indices=idx, + cache_seqlens_int32=cache_seqlens, nsa_cache_seqlens_int32=lengths, + ) + + def launch(): + return run_unified_decode( + q_all=q, swa_k_cache=kv_cache, swa_indices=idx, swa_topk_lengths=lengths, + workspace=binding.scratch, sm_scale=sm_scale, swa_page_size=_GLM_PAGE, + forced_num_splits=forced_num_splits, split_policy=split_policy, + ) + + if graph_lengths is None: + out = launch() + torch.cuda.synchronize() + exp = glm_ref.glm_decode_reference( + q, kv_cache_flat, idx, sm_scale, active_token_counts=lengths, + ).float() + return out.float().clone(), exp, lengths + + stream = torch.cuda.Stream() + with torch.cuda.stream(stream): + launch() # warm up (compiles) outside capture + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + out = launch() + torch.cuda.synchronize() + graph_lengths = graph_lengths.to(device=device, dtype=torch.int32) + lengths.copy_(graph_lengths) + graph.replay() + torch.cuda.synchronize() + exp = glm_ref.glm_decode_reference( + q, kv_cache_flat, idx, sm_scale, active_token_counts=lengths, + ).float() + return out.float().clone(), exp, lengths.clone() + + +def _assert_glm_rows_match_reference(got, exp, lengths, *, label): + for t in range(int(got.shape[0])): + cos = _cosine(got[t], exp[t]) + assert cos > 0.995, f"{label} token {t} (len={int(lengths[t])}) O cos={cos}" + assert (got[t] - exp[t]).abs().max().item() < 3e-2, ( + f"{label} token {t} (len={int(lengths[t])}) O atol exceeded" + ) + + +@torch.inference_mode() +@pytest.mark.parametrize("num_tokens,topk", [(1, 512), (4, 2048), (16, 512)]) +def test_unified_decode_glm_balanced_split_policy(num_tokens, topk) -> None: + """``split_policy="balanced"`` reproduces the static-policy result up to the + merge rounding of bf16 partials, activates at most + ``min(num_splits, sm_count // (rows * head_blocks))`` splits, and matches the + reference for mixed per-token lengths.""" + device = require_b12x_sparse_mla() + import b12x.attention._shared.mla.kernel as launch + + n_chunks = (topk + 63) // 64 + forced = min(n_chunks, 64) + static_out, exp, lengths = _run_glm_multitoken( + device, topk=topk, num_tokens=num_tokens, forced_num_splits=forced, + seed=7100 + num_tokens, split_policy="static", + ) + assert launch.LAST_DECODE_PLAN.get("balanced_split_target") == 0 + assert launch.LAST_DECODE_PLAN.get("split_policy") == "static" + balanced_out, _, _ = _run_glm_multitoken( + device, topk=topk, num_tokens=num_tokens, forced_num_splits=forced, + seed=7100 + num_tokens, split_policy="balanced", + ) + plan = launch.LAST_DECODE_PLAN + assert plan.get("split_policy") == "balanced" + sm_count = torch.cuda.get_device_properties(device).multi_processor_count + h_blocks = (_GLM_NUM_HEADS + 15) // 16 + expected_target = max(1, min(plan["num_splits"], sm_count // (num_tokens * h_blocks))) + assert plan.get("balanced_split_target") == expected_target + + _assert_glm_rows_match_reference(balanced_out, exp, lengths, label="balanced") + _assert_glm_rows_match_reference(static_out, exp, lengths, label="static") + # Same attention math, different partial grouping: bf16 rounding only. + assert _cosine(balanced_out, static_out) > 0.99999 + assert (balanced_out - static_out).abs().max().item() < 1.5e-2 + + +@torch.inference_mode() +def test_unified_decode_glm_balanced_split_policy_active_splits() -> None: + """Under the balanced policy the partial LSE of split ``s`` is finite exactly + when ``s * ceil(n / target) < n`` for the row's live chunk count ``n``.""" + device = require_b12x_sparse_mla() + from b12x.attention._shared.mla.kernel import run_unified_decode + import b12x.attention._shared.mla.kernel as launch + from b12x.attention.sparse_mla._scratch import ( + B12XSparseMLAScratchCaps, + plan_sparse_mla_scratch, + ) + + topk, num_tokens, forced = 4096, 2, 8 + nblk = (topk + _GLM_PAGE - 1) // _GLM_PAGE + case = glm_ref.make_glm_decode_case( + num_heads=_GLM_NUM_HEADS, topk=topk, num_tokens=num_tokens, num_blocks=nblk, + page_block_size=_GLM_PAGE, invalidate_half=False, seed=7300, device=device, + ) + q = case["q"].contiguous() + kv_cache = case["kv_cache"].contiguous() + kv_cache = kv_cache.view(nblk, _GLM_PAGE, kv_cache.shape[-1]) + idx = case["topk_indices"].contiguous() + lengths = torch.tensor([64 * 3 + 5, 4096], dtype=torch.int32, device=device) + caps = B12XSparseMLAScratchCaps( + device=device, num_q_heads=_GLM_NUM_HEADS, max_q_rows=num_tokens, + max_batch=num_tokens, max_width=topk, max_kv_rows=nblk * _GLM_PAGE, + head_dim=glm_ref.GLM_Q_HEAD_DIM, v_head_dim=glm_ref.GLM_D_V, + max_chunks_per_row=64, page_size=_GLM_PAGE, + ) + plan = plan_sparse_mla_scratch(caps) + (spec,) = plan.scratch_specs() + storage = torch.zeros(spec.shape, dtype=spec.dtype, device=device) + cache_seqlens = torch.full((num_tokens,), nblk * _GLM_PAGE, dtype=torch.int32, device=device) + binding = plan.bind( + scratch=storage, q=q, selected_indices=idx, + cache_seqlens_int32=cache_seqlens, nsa_cache_seqlens_int32=lengths, + ) + binding.scratch.tmp_lse.fill_(float("nan")) + run_unified_decode( + q_all=q, swa_k_cache=kv_cache, swa_indices=idx, swa_topk_lengths=lengths, + workspace=binding.scratch, sm_scale=case["sm_scale"], swa_page_size=_GLM_PAGE, + forced_num_splits=forced, split_policy="balanced", + ) + torch.cuda.synchronize() + target = launch.LAST_DECODE_PLAN["balanced_split_target"] + num_splits = launch.LAST_DECODE_PLAN["num_splits"] + static_cps = launch.LAST_DECODE_PLAN["chunks_per_split"] + assert num_splits == forced and static_cps == 8 + assert target == forced # 188 // (2 rows * 8 head blocks) = 11 > 8 splits + lse = binding.scratch.tmp_lse[:num_tokens, :, :num_splits] + expected_active = [] + for t in range(num_tokens): + n = (int(lengths[t]) + 63) // 64 + cps = min(-(-n // target), static_cps) + active = -(-n // cps) + expected_active.append(active) + assert active <= max(target, num_splits) + finite = torch.isfinite(lse[t]).all(dim=0) + assert finite[:active].all(), f"row {t}: expected {active} active splits" + assert (lse[t, :, active:] == float("-inf")).all(), ( + f"row {t}: splits >= {active} must be neutral" + ) + # Row 0 (4 live chunks) spreads over 4 single-chunk splits where the static + # ranges would keep one 8-chunk split busy; row 1 (64 chunks) keeps the + # static 8-chunk ranges. + assert expected_active == [4, 8] + + +@torch.inference_mode() +def test_unified_decode_glm_balanced_split_policy_waves_env(monkeypatch) -> None: + """``B12X_MLA_SM120_BALANCED_WAVES`` scales the active-split bound.""" + from b12x.attention._shared.mla.kernel import balanced_split_target_for + + kwargs = dict(num_splits=64, rows=4, h_blocks=4, sm_count=188) + monkeypatch.delenv("B12X_MLA_SM120_BALANCED_WAVES", raising=False) + assert balanced_split_target_for(**kwargs) == 11 + monkeypatch.setenv("B12X_MLA_SM120_BALANCED_WAVES", "2") + assert balanced_split_target_for(**kwargs) == 23 + monkeypatch.setenv("B12X_MLA_SM120_BALANCED_WAVES", "8") + assert balanced_split_target_for(**kwargs) == 64 + monkeypatch.delenv("B12X_MLA_SM120_BALANCED_WAVES") + assert balanced_split_target_for(num_splits=64, rows=64, h_blocks=8, sm_count=188) == 1 + monkeypatch.setenv("B12X_MLA_SM120_BALANCED_WAVES", "0") + with pytest.raises(ValueError): + balanced_split_target_for(**kwargs) + + +@torch.inference_mode() +def test_unified_decode_glm_balanced_split_policy_requires_per_token_lengths() -> None: + """The balanced policy is limited to the per-token single-cache entry.""" + device = require_b12x_sparse_mla() + from b12x.attention._shared.mla.kernel import run_unified_decode + from b12x.attention.sparse_mla._scratch import ( + B12XSparseMLAScratchCaps, + plan_sparse_mla_scratch, + ) + + topk = 256 + nblk = (topk + _GLM_PAGE - 1) // _GLM_PAGE + case = glm_ref.make_glm_decode_case( + num_heads=_GLM_NUM_HEADS, topk=topk, num_blocks=nblk, + page_block_size=_GLM_PAGE, invalidate_half=False, seed=52_300, device=device, + ) + kv_cache = case["kv_cache"].contiguous() + kv_cache = kv_cache.view(nblk, _GLM_PAGE, kv_cache.shape[-1]) + caps = B12XSparseMLAScratchCaps( + device=device, num_q_heads=_GLM_NUM_HEADS, max_q_rows=1, max_batch=1, + max_width=topk, max_kv_rows=nblk * _GLM_PAGE, + head_dim=glm_ref.GLM_Q_HEAD_DIM, v_head_dim=glm_ref.GLM_D_V, + max_chunks_per_row=8, page_size=_GLM_PAGE, + ) + plan = plan_sparse_mla_scratch(caps) + (spec,) = plan.scratch_specs() + storage = torch.zeros(spec.shape, dtype=spec.dtype, device=device) + lens = torch.full((1,), topk, dtype=torch.int32, device=device) + binding = plan.bind( + scratch=storage, q=case["q"].contiguous(), selected_indices=case["topk_indices"].contiguous(), + cache_seqlens_int32=lens, nsa_cache_seqlens_int32=lens, + ) + with pytest.raises(ValueError, match="per-token lengths"): + run_unified_decode( + q_all=binding.q, swa_k_cache=kv_cache, swa_indices=binding.selected_indices, + swa_topk_lengths=None, workspace=binding.scratch, sm_scale=case["sm_scale"], + swa_page_size=_GLM_PAGE, forced_num_splits=4, split_policy="balanced", + ) + + +@torch.inference_mode() +def test_unified_decode_glm_balanced_split_policy_graph_replay() -> None: + """A CUDA graph captured under the balanced policy stays exact when the + per-token lengths change between replays: the grid is capacity-based and + every CTA derives its chunk range at replay time.""" + device = require_b12x_sparse_mla() + topk, num_tokens = 2048, 4 + capture = torch.tensor([2048, 700, 65, 1], dtype=torch.int32) + replay = torch.tensor([129, 2048, 1000, 64 * 7], dtype=torch.int32) + got, exp, lengths = _run_glm_multitoken( + device, topk=topk, num_tokens=num_tokens, forced_num_splits=32, + seed=7400, split_policy="balanced", lengths=capture, graph_lengths=replay, + ) + _assert_glm_rows_match_reference(got, exp, lengths, label="graph replay") + eager, _, _ = _run_glm_multitoken( + device, topk=topk, num_tokens=num_tokens, forced_num_splits=32, + seed=7400, split_policy="balanced", lengths=replay, + ) + assert torch.equal(got, eager), "graph replay must be bit-identical to eager" + + +@torch.inference_mode() +@pytest.mark.parametrize("num_tokens,topk,forced", [(1, 512, 8), (4, 2048, 32)]) +def test_unified_decode_glm_fp32_partials(num_tokens, topk, forced) -> None: + """``partial_dtype=torch.float32`` keeps split partials exact: the merged + result matches the fp32 reference at least as closely as bf16 partials, the + plan records the partial dtype, and the output buffer no longer aliases the + partial workspace.""" + device = require_b12x_sparse_mla() + import b12x.attention._shared.mla.kernel as launch + + bf16_out, exp, lengths = _run_glm_multitoken( + device, topk=topk, num_tokens=num_tokens, forced_num_splits=forced, + seed=7500 + num_tokens, split_policy="balanced", + ) + assert launch.LAST_DECODE_PLAN.get("partial_dtype") == "torch.bfloat16" + fp32_out, _, _ = _run_glm_multitoken( + device, topk=topk, num_tokens=num_tokens, forced_num_splits=forced, + seed=7500 + num_tokens, split_policy="balanced", partial_dtype=torch.float32, + ) + assert launch.LAST_DECODE_PLAN.get("partial_dtype") == "torch.float32" + _assert_glm_rows_match_reference(fp32_out, exp, lengths, label="fp32 partials") + err_bf16 = (bf16_out - exp).norm().item() + err_fp32 = (fp32_out - exp).norm().item() + assert err_fp32 <= err_bf16 * 1.05, (err_fp32, err_bf16) + assert _cosine(fp32_out, bf16_out) > 0.99999 + + +@torch.inference_mode() +def test_unified_decode_glm_fp32_partials_single_split_bit_identical() -> None: + """With one active split the fp32-partial merge reproduces the bf16-partial + result bit for bit: the merge scales the single partial by 1 and rounds + once, exactly like the direct bf16 store.""" + device = require_b12x_sparse_mla() + lengths = torch.tensor([64, 20, 64, 64], dtype=torch.int32) + bf16_out, _, _ = _run_glm_multitoken( + device, topk=64, num_tokens=4, forced_num_splits=1, seed=7600, + lengths=lengths, + ) + fp32_out, _, _ = _run_glm_multitoken( + device, topk=64, num_tokens=4, forced_num_splits=1, seed=7600, + lengths=lengths, partial_dtype=torch.float32, + ) + assert torch.equal(bf16_out, fp32_out) + + +def test_sparse_mla_scratch_fp32_partials_layout() -> None: + """fp32 partials get their own output region; bf16 partials keep aliasing.""" + from b12x.attention.sparse_mla._scratch import ( + B12XSparseMLAScratchCaps, + _sparse_mla_scratch_layout, + ) + + common = dict( + device="cpu", num_q_heads=64, max_q_rows=4, max_width=4096, + head_dim=576, v_head_dim=512, max_chunks_per_row=64, page_size=64, + ) + bf16 = _sparse_mla_scratch_layout(B12XSparseMLAScratchCaps(**common)) + fp32 = _sparse_mla_scratch_layout( + B12XSparseMLAScratchCaps(**common, partial_dtype=torch.float32) + ) + assert bf16.output_offset_bytes == bf16.tmp_output_offset_bytes + partial_bytes = 4 * 64 * 64 * 512 + assert fp32.output_offset_bytes >= fp32.tmp_output_offset_bytes + 4 * partial_bytes + assert fp32.nbytes > bf16.nbytes + with pytest.raises(TypeError): + B12XSparseMLAScratchCaps(**common, partial_dtype=torch.float16) + + +@torch.inference_mode() +@pytest.mark.parametrize("num_tokens,topk", [(1, 512), (4, 2048)]) +def test_unified_decode_glm_fastpath_bit_identical(monkeypatch, num_tokens, topk) -> None: + """``B12X_MLA_SM120_GLM_FASTPATH=1`` stages packed 656-byte records, loads + the PV B-fragments with the b8 transposed ldmatrix and keeps W hi/lo in + fixed slots; bytes and MMA order are unchanged, so the output is + bit-identical to the base path and the plan records the mode.""" + device = require_b12x_sparse_mla() + import b12x.attention._shared.mla.kernel as launch + + n_chunks = (topk + 63) // 64 + forced = min(n_chunks, 8) + monkeypatch.delenv("B12X_MLA_SM120_GLM_FASTPATH", raising=False) + base, exp, lengths = _run_glm_multitoken( + device, topk=topk, num_tokens=num_tokens, forced_num_splits=forced, + seed=7700 + num_tokens, split_policy="balanced", + ) + assert launch.LAST_DECODE_PLAN.get("glm_fastpath") is False + monkeypatch.setenv("B12X_MLA_SM120_GLM_FASTPATH", "1") + ldsm, _, _ = _run_glm_multitoken( + device, topk=topk, num_tokens=num_tokens, forced_num_splits=forced, + seed=7700 + num_tokens, split_policy="balanced", + ) + assert launch.LAST_DECODE_PLAN.get("glm_fastpath") is True + assert torch.equal(base, ldsm) + _assert_glm_rows_match_reference(ldsm, exp, lengths, label="glm_fastpath") + + +@torch.inference_mode() +@pytest.mark.parametrize("num_tokens", [1, 4]) +def test_unified_decode_glm_serial_chunks_rescale_late_maximum(num_tokens) -> None: + """A split that walks several chunks serially must rescale its accumulators + when the running maximum rises in a later chunk. Each of the 64 keys of the + last of ten chunks is the dominant key of one pair of row-0 heads (logit + about 7 against a background of about 0.15); the single-split serial walk + must match the reference and the one-chunk-per-split walk. Without the + rescale the earlier chunks keep their unscaled weight and the output error + is close to 100 %.""" + device = require_b12x_sparse_mla() + topk = 640 + n_chunks = topk // 64 + boost = 60.0 + + def dominant_keys_in_last_chunk(kv_cache_flat, q, idx): + rope = q[0, :, 512:].float() + rope = rope / rope.norm(dim=-1, keepdim=True) + # Heads h and h + 64 share key h; the key points along their mean. + pair_mean = rope.view(-1, 64, rope.shape[-1]).sum(dim=0) + pair_mean = boost * pair_mean / pair_mean.norm(dim=-1, keepdim=True) + for key in range(64): + slot = int(idx[0, topk - 64 + key]) + kv_cache_flat[slot, 0, 528:656] = pair_mean[key].to(torch.bfloat16).view(torch.uint8) + + lengths = torch.full((num_tokens,), topk, dtype=torch.int32, device=device) + serial, exp, _ = _run_glm_multitoken( + device, topk=topk, num_tokens=num_tokens, forced_num_splits=1, + seed=7900 + num_tokens, lengths=lengths, cache_hook=dominant_keys_in_last_chunk, + partial_dtype=torch.float32, + ) + per_chunk, _, _ = _run_glm_multitoken( + device, topk=topk, num_tokens=num_tokens, forced_num_splits=n_chunks, + seed=7900 + num_tokens, lengths=lengths, cache_hook=dominant_keys_in_last_chunk, + partial_dtype=torch.float32, + ) + for label, got in (("serial", serial), ("per_chunk", per_chunk)): + rel = ((got - exp).norm() / exp.norm()).item() + assert rel < 2e-2, f"{label} rel-L2 vs reference {rel}" + rel = ((serial - per_chunk).norm() / exp.norm()).item() + assert rel < 1e-2, f"serial vs per-chunk rel-L2 {rel}" + + +def pack_glm_query_reference(q: torch.Tensor) -> torch.Tensor: + """Torch replica of the in-kernel S0 query quantization as a packed record. + + ``q`` is ``(rows, heads, 576)`` bf16. Returns uint8 ``(rows, heads, 656)``: + 512 E4M3 nope bytes (per 128-dim tile: absmax -> ``max(amax, 1e-4) / + FP8_MAX`` rounded up to a power of two -> ``q * (1 / scale)`` clamped and + rounded to nearest even), the four fp32 pow2 scales, then the 64 bf16 rope + values. + """ + from b12x._lib.intrinsics import pow2_ceil_ue8m0_torch + + fp8_max = 448.0 + nope = q[..., :512].float().reshape(*q.shape[:2], 4, 128) + amax = nope.abs().amax(dim=-1, keepdim=True) + raw = torch.clamp_min(amax, 1e-4) * torch.tensor(1.0 / fp8_max, dtype=torch.float32, device=q.device) + rounded, _ = pow2_ceil_ue8m0_torch(raw) + scaled = (nope * (1.0 / rounded)).clamp(-fp8_max, fp8_max).to(torch.float8_e4m3fn) + nope_bytes = scaled.reshape(*q.shape[:2], 512).view(torch.uint8) + scale_bytes = rounded.reshape(*q.shape[:2], 4).contiguous().view(torch.uint8) + rope_bytes = q[..., 512:].contiguous().view(torch.uint8) + return torch.cat([nope_bytes, scale_bytes, rope_bytes], dim=-1).contiguous() + + +@torch.inference_mode() +@pytest.mark.parametrize("num_tokens,topk", [(1, 512), (4, 2048)]) +def test_unified_decode_glm_packed_query_bit_identical(monkeypatch, num_tokens, topk) -> None: + """A packed query record (host-side S0 quantization) is bit-identical to + the bf16 query on the fast path, and is rejected without the fast path.""" + device = require_b12x_sparse_mla() + from b12x.attention._shared.mla.kernel import run_unified_decode + from b12x.attention.sparse_mla._scratch import ( + B12XSparseMLAScratchCaps, + plan_sparse_mla_scratch, + ) + import b12x.attention._shared.mla.kernel as launch + + nblk = max(1, (topk + _GLM_PAGE - 1) // _GLM_PAGE) + case = glm_ref.make_glm_decode_case( + num_heads=_GLM_NUM_HEADS, topk=topk, num_tokens=num_tokens, num_blocks=nblk, + page_block_size=_GLM_PAGE, invalidate_half=False, seed=7800 + num_tokens, device=device, + ) + q = case["q"].contiguous() + kv_cache_flat = case["kv_cache"].contiguous() + kv_cache = kv_cache_flat.view(nblk, _GLM_PAGE, kv_cache_flat.shape[-1]) + idx = case["topk_indices"].contiguous() + lengths = _mixed_lengths(num_tokens, topk, device) + caps = B12XSparseMLAScratchCaps( + device=device, num_q_heads=_GLM_NUM_HEADS, max_q_rows=num_tokens, + max_batch=num_tokens, max_width=topk, max_kv_rows=nblk * _GLM_PAGE, + head_dim=glm_ref.GLM_Q_HEAD_DIM, v_head_dim=glm_ref.GLM_D_V, + max_chunks_per_row=max(8, (topk + 63) // 64), page_size=_GLM_PAGE, + partial_dtype=torch.float32, + ) + plan = plan_sparse_mla_scratch(caps) + (spec,) = plan.scratch_specs() + storage = torch.zeros(spec.shape, dtype=spec.dtype, device=device) + cache_seqlens = torch.full((num_tokens,), nblk * _GLM_PAGE, dtype=torch.int32, device=device) + packed = pack_glm_query_reference(q) + assert packed.shape == (num_tokens, _GLM_NUM_HEADS, 656) and packed.dtype == torch.uint8 + + def run(q_in): + binding = plan.bind( + scratch=storage, q=q_in, selected_indices=idx, + cache_seqlens_int32=cache_seqlens, nsa_cache_seqlens_int32=lengths, + ) + out = run_unified_decode( + q_all=binding.q, swa_k_cache=kv_cache, swa_indices=idx, swa_topk_lengths=lengths, + workspace=binding.scratch, sm_scale=case["sm_scale"], swa_page_size=_GLM_PAGE, + forced_num_splits=8, split_policy="balanced", + ) + torch.cuda.synchronize() + return out.float().clone() + + monkeypatch.setenv("B12X_MLA_SM120_GLM_FASTPATH", "1") + monkeypatch.setenv("B12X_MLA_SM120_GLM_W_HW_DEQUANT", "1") + bf16_out = run(q) + packed_out = run(packed) + assert launch.LAST_DECODE_PLAN.get("glm_fastpath") is True + assert torch.equal(bf16_out, packed_out) + exp = glm_ref.glm_decode_reference( + q, kv_cache_flat, idx, case["sm_scale"], active_token_counts=lengths, + ).float() + _assert_glm_rows_match_reference(packed_out, exp, lengths, label="packed query") + + monkeypatch.setenv("B12X_MLA_SM120_GLM_FASTPATH", "0") + with pytest.raises(ValueError, match="packed query"): + run(packed) diff --git a/tests/attention/test_dense_mla.py b/tests/attention/test_dense_mla.py index bc12770ef..bd00dc134 100644 --- a/tests/attention/test_dense_mla.py +++ b/tests/attention/test_dense_mla.py @@ -196,6 +196,70 @@ def test_partial_row_budget_changes_native_split_policy() -> None: assert plan.chunks_per_split == 2 +def test_fp8_multi_request_verify_plan_tiles_four_queries() -> None: + plan = dense_mla.plan( + dense_mla.Caps( + device="cpu", + mode="verify", + kv_dtype=FP8, + num_q_heads=HEADS, + page_size=16, + max_total_q=8, + max_batch=2, + max_cache_tokens=128, + max_page_table_width=8, + num_cache_pages=16, + uses_query_cache_seqlens=True, + ) + ) + + assert plan.query_tile == 4 + + +def test_dynamic_sparse_chunk_policy_preserves_sink_and_recent_chunks() -> None: + assert dense_mla.dynamic_sparse_chunk_indices( + 10, + stride=3, + sink_chunks=2, + recent_chunks=2, + ) == (0, 1, 2, 5, 8, 9) + + +def test_verify_plan_requires_query_cache_lengths() -> None: + plan = dense_mla.plan( + dense_mla.Caps( + device="cpu", + mode="verify", + kv_dtype=FP8, + num_q_heads=HEADS, + page_size=16, + max_total_q=4, + max_batch=1, + max_cache_tokens=64, + max_page_table_width=4, + num_cache_pages=4, + uses_query_cache_seqlens=True, + ) + ) + q = torch.empty(4, HEADS, QK_DIM, dtype=FP8) + cache = torch.empty(4, 16, QK_DIM, dtype=FP8) + output = torch.empty(4, HEADS, VALUE_DIM, dtype=torch.bfloat16) + + with pytest.raises(ValueError, match="requires per-query cache lengths"): + dense_mla.bind( + plan, + scratch=_scratch(plan), + q=q, + kv_cache=cache, + output=output, + page_table=torch.arange(4, dtype=torch.int32).view(1, 4), + cache_seqlens=torch.tensor([64], dtype=torch.int32), + cu_seqlens_q=torch.tensor([0, 4], dtype=torch.int32), + q_scale=torch.tensor(0.01), + kv_scale=torch.tensor(0.01), + ) + + @pytest.mark.parametrize("heads", [8, 12]) @torch.inference_mode() def test_bf16_multi_request_decode_matches_reference(heads: int) -> None: @@ -397,6 +461,179 @@ def test_fp8_query_tiled_causal_extend_matches_reference(heads: int) -> None: ) +@torch.inference_mode() +def test_fp8_tiled_dcp_visibility_matches_reference() -> None: + device = require_b12x() + torch.manual_seed(20260901) + batch = 2 + query_len = 4 + total_q = batch * query_len + heads = 8 + page_size = 16 + pages = 8 + plan = dense_mla.plan( + dense_mla.Caps( + device=device, + mode="verify", + kv_dtype=FP8, + num_q_heads=heads, + page_size=page_size, + max_total_q=total_q, + max_batch=batch, + max_cache_tokens=64, + max_page_table_width=4, + num_cache_pages=pages, + uses_query_cache_seqlens=True, + ) + ) + assert plan.query_tile == 4 + q_float = torch.randn(total_q, heads, QK_DIM, device=device) * 0.14 + cache_float = torch.randn(pages, page_size, QK_DIM, device=device) * 0.1 + q_scale = (q_float.abs().max() / 400).reshape(1).float() + kv_scale = (cache_float.abs().max() / 400).reshape(1).float() + q = (q_float / q_scale).to(FP8) + cache = (cache_float / kv_scale).to(FP8) + page_table = torch.tensor( + [[4, 7, 1, 5], [0, 3, 6, 2]], + dtype=torch.int32, + device=device, + ) + cache_seqlens = torch.tensor([31, 47], dtype=torch.int32, device=device) + query_cache_seqlens = torch.tensor( + [28, 29, 30, 31, 44, 45, 46, 47], + dtype=torch.int32, + device=device, + ) + cu_seqlens_q = torch.tensor([0, 4, 8], dtype=torch.int32, device=device) + output = torch.empty( + total_q, + heads, + VALUE_DIM, + dtype=torch.bfloat16, + device=device, + ) + binding = dense_mla.bind( + plan, + scratch=_scratch(plan), + q=q, + kv_cache=cache, + output=output, + page_table=page_table, + cache_seqlens=cache_seqlens, + query_cache_seqlens=query_cache_seqlens, + cu_seqlens_q=cu_seqlens_q, + q_scale=q_scale, + kv_scale=kv_scale, + ) + + actual_output, actual_lse = dense_mla.run(binding=binding) + expected_output, expected_lse = dense_mla.reference( + q, + cache, + page_table, + cache_seqlens, + cu_seqlens_q, + query_cache_seqlens=query_cache_seqlens, + q_scale=q_scale, + kv_scale=kv_scale, + ) + + _assert_matches( + actual_output, + actual_lse, + expected_output, + expected_lse, + ) + + +@torch.inference_mode() +def test_fp8_dynamic_sparse_verify_matches_sparse_reference() -> None: + device = require_b12x() + torch.manual_seed(20260902) + query_len = 4 + heads = 8 + page_size = 16 + pages = 16 + cache_len = 240 + sparse_kwargs = { + "sparse_stride": 3, + "sparse_min_tokens": 64, + "sparse_sink_chunks": 1, + "sparse_recent_chunks": 1, + } + plan = dense_mla.plan( + dense_mla.Caps( + device=device, + mode="verify", + kv_dtype=FP8, + num_q_heads=heads, + page_size=page_size, + max_total_q=query_len, + max_batch=1, + max_cache_tokens=256, + max_page_table_width=16, + num_cache_pages=pages, + uses_query_cache_seqlens=True, + **sparse_kwargs, + ) + ) + q_float = torch.randn(query_len, heads, QK_DIM, device=device) * 0.14 + cache_float = torch.randn(pages, page_size, QK_DIM, device=device) * 0.1 + q_scale = (q_float.abs().max() / 400).reshape(1).float() + kv_scale = (cache_float.abs().max() / 400).reshape(1).float() + q = (q_float / q_scale).to(FP8) + cache = (cache_float / kv_scale).to(FP8) + page_table = torch.arange(pages, dtype=torch.int32, device=device).view(1, -1) + cache_seqlens = torch.tensor([cache_len], dtype=torch.int32, device=device) + query_cache_seqlens = torch.arange( + cache_len - query_len + 1, + cache_len + 1, + dtype=torch.int32, + device=device, + ) + cu_seqlens_q = torch.tensor([0, query_len], dtype=torch.int32, device=device) + output = torch.empty( + query_len, + heads, + VALUE_DIM, + dtype=torch.bfloat16, + device=device, + ) + binding = dense_mla.bind( + plan, + scratch=_scratch(plan), + q=q, + kv_cache=cache, + output=output, + page_table=page_table, + cache_seqlens=cache_seqlens, + query_cache_seqlens=query_cache_seqlens, + cu_seqlens_q=cu_seqlens_q, + q_scale=q_scale, + kv_scale=kv_scale, + ) + + actual_output, actual_lse = dense_mla.run(binding=binding) + expected_output, expected_lse = dense_mla.reference( + q, + cache, + page_table, + cache_seqlens, + cu_seqlens_q, + query_cache_seqlens=query_cache_seqlens, + q_scale=q_scale, + kv_scale=kv_scale, + **sparse_kwargs, + ) + + _assert_matches( + actual_output, + actual_lse, + expected_output, + expected_lse, + ) + + @torch.inference_mode() def test_bf16_query_tiled_causal_extend_matches_reference() -> None: device = require_b12x() diff --git a/tests/moe/test_fused_moe_trellis.py b/tests/moe/test_fused_moe_trellis.py index d3e4749bb..875bf1251 100644 --- a/tests/moe/test_fused_moe_trellis.py +++ b/tests/moe/test_fused_moe_trellis.py @@ -267,6 +267,55 @@ def test_low_level_buffer_plan_honors_explicit_block_m() -> None: ) +def test_low_level_buffer_plan_bounds_large_m_route_reduction( + monkeypatch: pytest.MonkeyPatch, +) -> None: + prepared = SimpleNamespace( + num_experts=896, + hidden_size=7168, + intermediate_size=192, + is_gated=True, + ) + plan_kwargs = dict( + prepared=prepared, + m=4096, + topk=16, + route_num_experts=896, + sms=188, + dtype=torch.bfloat16, + ) + + monkeypatch.delenv("B12X_W4A16_PREFILL_FUSED_SUM", raising=False) + materialized = plan_w4a16_buffers(**plan_kwargs) + monkeypatch.setenv("B12X_W4A16_PREFILL_FUSED_SUM", "1") + fused = plan_w4a16_buffers(**plan_kwargs) + rotation = plan_w4a16_buffers(**plan_kwargs, full_rotation=True) + fp16 = plan_w4a16_buffers(**{**plan_kwargs, "dtype": torch.float16}) + trellis = plan_w4a16_buffers( + **plan_kwargs, + weight_layout="trellis_t256", + ) + activation_amax = plan_w4a16_buffers( + **plan_kwargs, + collect_activation_amax=True, + ) + + routed_rows = 4096 * 16 + fc1_cols = 2 * 192 + assert materialized.intermediate_cache13_elements == routed_rows * 7168 + assert materialized.prefill_sum_accum_elements == 0 + assert fused.intermediate_cache13_elements == routed_rows * fc1_cols + assert fused.prefill_sum_accum_elements == 4096 * 7168 + assert rotation.intermediate_cache13_elements == routed_rows * 7168 + assert rotation.prefill_sum_accum_elements == 0 + assert fp16.intermediate_cache13_elements == routed_rows * 7168 + assert fp16.prefill_sum_accum_elements == 0 + assert trellis.intermediate_cache13_elements == routed_rows * 7168 + assert trellis.prefill_sum_accum_elements == 0 + assert activation_amax.intermediate_cache13_elements == routed_rows * 7168 + assert activation_amax.prefill_sum_accum_elements == 0 + + def test_planned_route_block_overrides_live_batch_heuristic() -> None: from b12x.moe._shared.kernels.w4a16.kernel import ( _resolve_route_block_size_m, diff --git a/tests/moe/test_moe_launch_param_regression.py b/tests/moe/test_moe_launch_param_regression.py index e6d35b7d2..f427ef7b0 100644 --- a/tests/moe/test_moe_launch_param_regression.py +++ b/tests/moe/test_moe_launch_param_regression.py @@ -132,6 +132,7 @@ def _direct_micro_launchable( weight_E: int = 256, k: int = 4096, num_topk: int = 10, + compile_time_phase: int = 0, ) -> bool: from b12x.moe.fused_moe._impl import ( _DIRECT_MICRO_BLOCK_DIM, @@ -153,6 +154,7 @@ def _direct_micro_launchable( activation="silu", quant_mode=quant_mode, device=torch.device("cuda"), + compile_time_phase=compile_time_phase, ) return _compiled_direct_micro_accepts_block_dim(compiled, _DIRECT_MICRO_BLOCK_DIM) @@ -163,6 +165,21 @@ def test_nvfp4_direct_micro_launches_qwen_bs8_shape() -> None: assert _direct_micro_launchable("nvfp4", 8, 256, weight_E=512) +@pytest.mark.parametrize("compile_time_phase", [1, 2]) +def test_nvfp4_split_micro_compile_signature_matches_runtime_launch( + compile_time_phase: int, +) -> None: + _skip_if_no_sm120() + + assert _direct_micro_launchable( + "nvfp4", + 1, + 256, + weight_E=512, + compile_time_phase=compile_time_phase, + ) + + @pytest.mark.parametrize("case", ["alphas", "scales"]) def test_b12x_moe_accepts_parameter_backed_launch_args(case: str) -> None: """The static path should not segfault on Parameter-backed scale tensors.""" diff --git a/tests/moe/test_tp_moe_scratch_bindings.py b/tests/moe/test_tp_moe_scratch_bindings.py index 01c005bc2..707d8e1c2 100644 --- a/tests/moe/test_tp_moe_scratch_bindings.py +++ b/tests/moe/test_tp_moe_scratch_bindings.py @@ -525,6 +525,61 @@ def test_w4a16_scratch_plan_uses_route_pack_capacity_buckets( assert plan_4080.shapes_and_dtypes() == plan_4096.shapes_and_dtypes() +def test_w4a16_prefill_route_reduction_shrinks_caller_owned_scratch( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(tp_moe_impl, "get_num_sm", lambda _device: 188) + weight_plan = _weight_plan( + "w4a16", + source_format="fp4_e8m0_k32", + experts=896, + k=7168, + n=192, + activation="situ", + ) + base_caps = dict( + max_tokens=4096, + core_token_counts=(4096,), + num_topk=16, + route_num_experts=896, + device="cpu", + weight_plan=weight_plan, + quant_mode="w4a16", + ) + + monkeypatch.delenv("B12X_W4A16_PREFILL_FUSED_SUM", raising=False) + materialized = plan_tp_moe_scratch(TPMoEScratchCaps(**base_caps)) + monkeypatch.setenv("B12X_W4A16_PREFILL_FUSED_SUM", "1") + fused = plan_tp_moe_scratch(TPMoEScratchCaps(**base_caps)) + calibrated = plan_tp_moe_scratch( + TPMoEScratchCaps(**base_caps, collect_activation_amax=True) + ) + + def specs(plan): + return {spec.name: spec for spec in plan._core_workspace_plan.tensor_specs} + + materialized_specs = specs(materialized) + fused_specs = specs(fused) + calibrated_specs = specs(calibrated) + routed_rows = 4096 * 16 + fc1_cols = 2 * 192 + + assert materialized_specs["intermediate_cache13"].shape == ( + routed_rows * 7168, + ) + assert "prefill_sum_accum" not in materialized_specs + assert fused_specs["intermediate_cache13"].shape == ( + routed_rows * fc1_cols, + ) + assert fused_specs["prefill_sum_accum"].shape == (4096 * 7168,) + assert fused_specs["prefill_sum_accum"].dtype == torch.float32 + assert calibrated_specs["intermediate_cache13"].shape == ( + routed_rows * 7168, + ) + assert "prefill_sum_accum" not in calibrated_specs + assert fused.layout.core_workspace_nbytes < materialized.layout.core_workspace_nbytes + + def test_trellis_scratch_plan_preserves_exact_fixed_capacity( monkeypatch: pytest.MonkeyPatch, ) -> None: @@ -814,9 +869,15 @@ def test_w4a16_materialize_can_prewarm_activation_amax_variant( fused = object() def _fake_w4a16_prewarm( - workspace, *, token_counts, collect_activation_amax=False, **_kwargs + workspace, + *, + token_counts, + collect_activation_amax=False, + prefill_fused_sum=False, + **_kwargs, ) -> None: captured["collect_activation_amax"] = bool(collect_activation_amax) + captured["prefill_fused_sum"] = bool(prefill_fused_sum) workspace.planned_fused_moe_launches = { ( "packed", @@ -864,11 +925,57 @@ def _fake_w4a16_prewarm( ) assert captured["collect_activation_amax"] is True + assert captured["prefill_fused_sum"] is False assert workspace.planned_collect_activation_amax is True + assert workspace.planned_prefill_fused_sum_fp32 is False assert selected is fused assert topk_sum is not None +def test_w4a16_materialize_freezes_prefill_reduction_selection( + monkeypatch: pytest.MonkeyPatch, +) -> None: + captured = {} + + def _fake_w4a16_prewarm( + workspace, + *, + token_counts, + prefill_fused_sum=False, + **_kwargs, + ) -> None: + captured["prefill_fused_sum"] = bool(prefill_fused_sum) + monkeypatch.setenv("B12X_W4A16_PREFILL_FUSED_SUM", "0") + + monkeypatch.setenv("B12X_W4A16_PREFILL_FUSED_SUM", "1") + monkeypatch.setattr(tp_moe_impl, "get_num_sm", lambda _device: 120) + monkeypatch.setattr( + tp_moe_impl, + "_prewarm_w4a16_planned_launches", + _fake_w4a16_prewarm, + ) + pool = tp_moe_impl.allocate_tp_moe_workspace_pool(frozen=True) + weight_plan = _weight_plan( + "w4a16", + w4a16_layout=PreparedWeightLayout.MMA_PACKED, + ) + + tp_moe_impl.materialize_tp_moe_arena_workspaces( + pool, + caps=_caps( + max_tokens=16, + weight_plan=weight_plan, + core_token_counts=(16,), + route_num_experts=0, + ), + ) + + workspace = next(iter(pool.workspaces.values())) + assert captured["prefill_fused_sum"] is True + assert workspace.planned_prefill_fused_sum_fp32 is True + assert workspace.prefill_sum_accum is not None + + def test_w4a16_scratch_binding_carries_activation_amax_to_kernel( monkeypatch: pytest.MonkeyPatch, ) -> None: diff --git a/tests/moe/test_w4a16_e2e.py b/tests/moe/test_w4a16_e2e.py index 6d63a02e6..cb37851fa 100644 --- a/tests/moe/test_w4a16_e2e.py +++ b/tests/moe/test_w4a16_e2e.py @@ -62,6 +62,19 @@ def test_w4a16_small_m_host_barrier_reset_kill_switch( assert not _small_m_direct_host_barrier_reset_enabled() +def test_w4a16_fc2_runtime_m_does_not_use_fixed_route_staging() -> None: + kernel = MoEMicroKernelW4A16SmallMDirect( + activation="silu", + fast_math=False, + share_input_across_experts=False, + share_expert_scales=True, + single_token=False, + scale_format="e8m0_k32", + compile_time_phase=2, + ) + assert not kernel.stage_inactive_routes + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") @pytest.mark.parametrize("host_barrier_reset", [False, True]) def test_w4a16_small_m_direct_barrier_modes_eager_and_graph( @@ -1088,6 +1101,163 @@ def launch() -> torch.Tensor: _assert_matches_oracle(buffers.output, expected, activation=activation) +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +@pytest.mark.parametrize( + ("m", "route_ids_dtype"), + [ + (1, torch.int32), + (2, torch.int32), + (4, torch.int32), + (8, torch.int32), + (1, torch.int64), + ], +) +def test_w4a16_e8m0_native_micro_ignores_inactive_routes_during_graph_replay( + m: int, + route_ids_dtype: torch.dtype, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Native small-M execution must not address weights for inactive routes.""" + import b12x.moe._shared.kernels.w4a16.kernel as w4a16_kernel + + monkeypatch.setenv("B12X_W4A16_SMALL_M_DIRECT", "1") + direct_launches = 0 + real_direct_launch = w4a16_kernel._w4a16_small_m_direct_launch_flat + + def spy_direct_launch(*args, **kwargs) -> None: + nonlocal direct_launches + direct_launches += 1 + real_direct_launch(*args, **kwargs) + + monkeypatch.setattr( + w4a16_kernel, + "_w4a16_small_m_direct_launch_flat", + spy_direct_launch, + ) + experts, hidden_size, intermediate_size = 4, 128, 192 + topk, activation = 2, "situ" + rows = 2 * intermediate_size + torch.manual_seed(20260817 + m) + w13 = torch.randint( + 0, + 256, + (experts, rows, hidden_size // 2), + dtype=torch.uint8, + device="cuda", + ) + w2 = torch.randint( + 0, + 256, + (experts, hidden_size, intermediate_size // 2), + dtype=torch.uint8, + device="cuda", + ) + w13_scale = _pattern_e8m0((experts, rows, hidden_size // 32)) + w2_scale = _pattern_e8m0((experts, hidden_size, intermediate_size // 32), offset=1) + global_scale = torch.ones(experts, dtype=torch.float32, device="cuda") + prepared = prepare_w4a16_e8m0_native_weights( + w13, + w13_scale, + global_scale, + w2, + w2_scale, + global_scale, + activation=activation, + params_dtype=torch.bfloat16, + w13_layout="w31", + ) + buffers = make_w4a16_buffers( + prepared, + m=m, + topk=topk, + dtype=torch.bfloat16, + device=torch.device("cuda"), + ) + fc2_n_chunks = ((intermediate_size // 2) + 127) // 128 + intermediate_cache2 = torch.zeros( + 2 * m * fc2_n_chunks * 128 * topk, + dtype=torch.bfloat16, + device="cuda", + ) + inputs = torch.randn(m, hidden_size, dtype=torch.bfloat16, device="cuda") + topk_ids = torch.randint( + 0, experts, (m, topk), dtype=route_ids_dtype, device="cuda" + ) + topk_weights = torch.rand(m, topk, dtype=torch.float32, device="cuda") + + def launch() -> torch.Tensor: + return run_w4a16_moe( + inputs, + prepared, + topk_weights, + topk_ids, + activation=activation, + intermediate_cache13=buffers.intermediate_cache13, + intermediate_cache2=intermediate_cache2, + output=buffers.output, + fc1_c_tmp=buffers.fc1_c_tmp, + fc2_c_tmp=buffers.fc2_c_tmp, + packed_route_indices=buffers.packed_route_indices, + block_expert_ids=buffers.block_expert_ids, + packed_route_count=buffers.packed_route_count, + expert_offsets=buffers.expert_offsets, + ) + + valid_eager = launch().clone() + torch.cuda.synchronize() + assert direct_launches == 1 + assert bool(torch.isfinite(valid_eager).all().item()) + assert bool((valid_eager.abs().sum(dim=1) > 0).all().item()) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + captured = launch() + assert direct_launches == 2 + + inactive_ids = topk_ids.clone() + invalid_upper = 1 << 32 if route_ids_dtype == torch.int64 else experts + inactive_ids[-1] = torch.tensor( + [-1, invalid_upper], dtype=route_ids_dtype, device="cuda" + ) + if m > 1: + inactive_ids[0, 0] = -1 + topk_ids.copy_(inactive_ids) + original_weights = topk_weights.clone() + active = (inactive_ids >= 0) & (inactive_ids < experts) + reference_ids = torch.where(active, inactive_ids, torch.zeros_like(inactive_ids)) + reference_weights = torch.where( + active, topk_weights, torch.zeros_like(topk_weights) + ) + expected = moe_reference_w4a16_fp4_e8m0_k32( + inputs, + w13, + w13_scale, + global_scale, + w2, + w2_scale, + global_scale, + reference_ids, + reference_weights, + experts, + hidden_size, + intermediate_size, + activation=activation, + w13_layout="w31", + ) + + eager = launch().clone() + torch.cuda.synchronize() + _assert_matches_oracle(eager, expected, activation=activation) + torch.testing.assert_close(eager[-1], torch.zeros_like(eager[-1]), rtol=0, atol=0) + + captured.fill_(float("nan")) + graph.replay() + torch.cuda.synchronize() + assert bool(torch.isfinite(captured).all().item()) + torch.testing.assert_close(captured, eager, rtol=0, atol=0) + torch.testing.assert_close(topk_ids, inactive_ids, rtol=0, atol=0) + torch.testing.assert_close(topk_weights, original_weights, rtol=0, atol=0) + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") def test_w4a16_e8m0_native_aligned_generic_fc1_uses_k32_scale_stride() -> None: """K/32 scales keep their native expert stride in aligned generic FC1.""" @@ -1999,6 +2169,152 @@ def test_w4a16_tc_decode_preplanned_launch_matches_oracle(m: int) -> None: _assert_matches_oracle(actual, expected, activation=activation) +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_w4a16_prefill_fused_sum_is_graph_safe_and_matches_oracle( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Large-M route reduction uses fixed caller-owned scratch across replay. + + Relaxed FP32 atomics may change the final BF16 rounding bit when CTA order + changes. Every replay must remain finite, nonzero, and close to the same + FP32 oracle without allocating or changing scratch addresses. + """ + monkeypatch.setenv("B12X_W4A16_PREFILL_FUSED_SUM", "1") + torch.manual_seed(20260819) + m = 32 + experts, hidden_size, intermediate_size = 8, 128, 128 + topk, activation = 2, "silu" + weights = _make_weights( + experts=experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + activation=activation, + ) + x = (torch.randn(m, hidden_size, device="cuda") * 0.25).to(torch.bfloat16) + topk_ids = torch.randint( + 0, experts, (m, topk), device="cuda", dtype=torch.int32 + ) + topk_weights = torch.softmax(torch.randn(m, topk, device="cuda"), dim=-1) + prepared = prepare_w4a16_weights( + *weights, + activation=activation, + params_dtype=x.dtype, + ) + buffers = make_w4a16_buffers( + prepared, + m=m, + topk=topk, + dtype=x.dtype, + device=x.device, + ) + assert buffers.prefill_sum_accum is not None + assert buffers.prefill_sum_accum.dtype == torch.float32 + assert buffers.prefill_sum_accum.numel() == m * hidden_size + + props = torch.cuda.get_device_properties(x.device) + block_size_m = select_route_block_size_m(m, topk, experts) + _, _, max_m_blocks = route_pack_capacity( + m * topk, + block_size_m, + experts, + topk=topk, + ) + fused_launch = compile_w4a16_fused_moe( + size_m=m, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + num_experts=experts, + top_k=topk, + activation=activation, + apply_router_weight_on_input=False, + zero_fc2_output=False, + moe_block_size=block_size_m, + max_m_blocks=max_m_blocks, + element_dtype="bf16", + sms=int(props.multi_processor_count), + max_shared_mem=int( + getattr(props, "shared_memory_per_block_optin", _DEFAULT_MAX_SHARED_MEM) + ), + weight_layout="packed", + scale_format="e4m3_k16", + w13_layout="packed", + prefill_fused_sum_fp32=True, + ) + assert fused_launch.prefill_fused_sum_fp32 + assert not fused_launch.tc_decode_fused_sum + + def run() -> torch.Tensor: + return run_w4a16_moe( + x, + prepared, + topk_weights, + topk_ids, + activation=activation, + fast_math=True, + intermediate_cache13=buffers.intermediate_cache13, + intermediate_cache2=buffers.intermediate_cache2, + output=buffers.output, + prefill_sum_accum=buffers.prefill_sum_accum, + fc1_c_tmp=buffers.fc1_c_tmp, + fc2_c_tmp=buffers.fc2_c_tmp, + packed_route_indices=buffers.packed_route_indices, + block_expert_ids=buffers.block_expert_ids, + packed_route_count=buffers.packed_route_count, + expert_offsets=buffers.expert_offsets, + expert_counts=buffers.expert_counts, + fused_launch=fused_launch, + ) + + expected = _reference_w4a16( + x, + *weights, + topk_ids, + topk_weights, + activation=activation, + ) + eager = run().clone() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + captured = run() + pointer_contract = ( + buffers.intermediate_cache13.data_ptr(), + buffers.intermediate_cache2.data_ptr(), + buffers.prefill_sum_accum.data_ptr(), + buffers.output.data_ptr(), + ) + buffers.prefill_sum_accum.fill_(float("nan")) + allocated_before_replay = torch.cuda.memory_allocated() + graph.replay() + torch.cuda.synchronize() + allocated_after_replay = torch.cuda.memory_allocated() + first_replay = captured.clone() + graph.replay() + torch.cuda.synchronize() + second_replay = captured.clone() + + assert pointer_contract == ( + buffers.intermediate_cache13.data_ptr(), + buffers.intermediate_cache2.data_ptr(), + buffers.prefill_sum_accum.data_ptr(), + buffers.output.data_ptr(), + ) + assert allocated_after_replay == allocated_before_replay + for actual in (eager, first_replay, second_replay): + assert bool(torch.isfinite(actual).all().item()) + assert bool(torch.count_nonzero(actual).item()) + _assert_matches_oracle(actual, expected, activation=activation) + repeat_metrics = compare_to_reference(first_replay, second_replay) + assert repeat_metrics.cos >= 0.999999, repeat_metrics + torch.testing.assert_close( + second_replay, + buffers.prefill_sum_accum[: m * hidden_size] + .view(m, hidden_size) + .to(second_replay.dtype), + rtol=0, + atol=0, + ) + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") @pytest.mark.parametrize("m", [1, 4, 6]) def test_w4a16_small_m_packed_direct_topk_routes_matches_oracle(m: int) -> None: @@ -3256,6 +3572,78 @@ def test_w4a16_fc2_only_consumes_contiguous_bf16_and_native_mxfp4() -> None: torch.testing.assert_close(actual, expected, rtol=0, atol=0) +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +@pytest.mark.parametrize("intermediate_size", [256, 512]) +@pytest.mark.parametrize("route_ids_dtype", [torch.int32, torch.int64]) +def test_w4a16_fc2_only_zeroes_invalid_routes_at_runtime_m3( + intermediate_size: int, + route_ids_dtype: torch.dtype, +) -> None: + experts, hidden_size, routes = 2, 128, 3 + w2 = torch.empty( + (experts, hidden_size, intermediate_size // 2), + dtype=torch.uint8, + device="cuda", + ) + w2[0].fill_(0x11) + w2[1].fill_(0x22) + scales = torch.full( + (experts, hidden_size, intermediate_size // 32), + 127, + dtype=torch.uint8, + device="cuda", + ) + intermediate = torch.ones( + (routes, intermediate_size), dtype=torch.bfloat16, device="cuda" + ) + invalid_upper = 1 << 32 if route_ids_dtype == torch.int64 else experts + route_ids = torch.tensor( + [0, -1, invalid_upper], dtype=route_ids_dtype, device="cuda" + ) + route_weights = torch.tensor([0.25, 0.5, 1.0], dtype=torch.float32, device="cuda") + original_ids = route_ids.clone() + original_weights = route_weights.clone() + + prepared = prepare_w4a16_fc2_e8m0(w2, scales) + actual = run_w4a16_fc2_e8m0( + intermediate, + prepared, + route_ids, + route_weights, + ) + expected_values = torch.tensor( + [intermediate_size / 8.0, 0.0, 0.0], + dtype=torch.bfloat16, + device="cuda", + ) + expected = expected_values[:, None].expand_as(actual) + + assert bool(torch.isfinite(actual).all().item()) + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + torch.testing.assert_close(route_ids, original_ids, rtol=0, atol=0) + torch.testing.assert_close(route_weights, original_weights, rtol=0, atol=0) + + graph_output = torch.empty_like(actual) + prewarm_w4a16_fc2_e8m0(prepared, route_ids_dtype=route_ids_dtype) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + captured = run_w4a16_fc2_e8m0( + intermediate, + prepared, + route_ids, + route_weights, + output=graph_output, + ) + captured.fill_(float("nan")) + graph.replay() + torch.cuda.synchronize() + assert captured is graph_output + assert bool(torch.isfinite(captured).all().item()) + torch.testing.assert_close(captured, expected, rtol=0, atol=0) + torch.testing.assert_close(route_ids, original_ids, rtol=0, atol=0) + torch.testing.assert_close(route_weights, original_weights, rtol=0, atol=0) + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") def test_w4a16_fc2_only_is_cuda_graph_safe_with_preallocated_output() -> None: experts, hidden_size, intermediate_size, routes = 5, 128, 256, 7 @@ -3275,15 +3663,11 @@ def test_w4a16_fc2_only_is_cuda_graph_safe_with_preallocated_output() -> None: intermediate = torch.ones( (routes, intermediate_size), dtype=torch.bfloat16, device="cuda" ) - route_ids = torch.tensor( - [0, 1, 2, 3, 4, 0, 2], dtype=torch.int32, device="cuda" - ) + route_ids = torch.tensor([0, 1, 0, 1, 0, 1, 0], dtype=torch.int32, device="cuda") route_weights = torch.linspace( 0.125, 0.875, routes, dtype=torch.float32, device="cuda" ) - output = torch.empty( - (routes, hidden_size), dtype=torch.bfloat16, device="cuda" - ) + output = torch.empty((routes, hidden_size), dtype=torch.bfloat16, device="cuda") prepared = prepare_w4a16_fc2_e8m0(w2, scales) prewarm_w4a16_fc2_e8m0(prepared, route_ids_dtype=torch.int32) @@ -3301,14 +3685,46 @@ def test_w4a16_fc2_only_is_cuda_graph_safe_with_preallocated_output() -> None: graph.replay() torch.cuda.synchronize() - expected = torch.empty_like(output) + valid_values = ( + 256.0 + * route_weights + * torch.where( + route_ids == 0, + torch.tensor(0.5, device="cuda"), + torch.tensor(1.0, device="cuda"), + ) + ) + valid_expected = valid_values.to(torch.bfloat16)[:, None].expand_as(output) + torch.testing.assert_close(captured, valid_expected, rtol=0, atol=0) + + invalid_ids = torch.tensor( + [0, -1, 1, experts, 0, -9, 1], dtype=torch.int32, device="cuda" + ) + route_ids.copy_(invalid_ids) + original_weights = route_weights.clone() + invalid_values = torch.tensor( + [16.0, 0.0, 96.0, 0.0, 80.0, 0.0, 224.0], + dtype=torch.bfloat16, + device="cuda", + ) + invalid_expected = invalid_values[:, None].expand_as(output) + + eager = torch.empty_like(output) run_w4a16_fc2_e8m0( intermediate, prepared, route_ids, route_weights, - output=expected, + output=eager, ) torch.cuda.synchronize() + torch.testing.assert_close(eager, invalid_expected, rtol=0, atol=0) + + captured.fill_(float("nan")) + graph.replay() + torch.cuda.synchronize() assert captured is output - torch.testing.assert_close(captured, expected, rtol=0, atol=0) + assert bool(torch.isfinite(captured).all().item()) + torch.testing.assert_close(captured, invalid_expected, rtol=0, atol=0) + torch.testing.assert_close(route_ids, invalid_ids, rtol=0, atol=0) + torch.testing.assert_close(route_weights, original_weights, rtol=0, atol=0) diff --git a/validation/attention/check_kimi_packed_mla_high_pages.py b/validation/attention/check_kimi_packed_mla_high_pages.py new file mode 100644 index 000000000..2d005bc8c --- /dev/null +++ b/validation/attention/check_kimi_packed_mla_high_pages.py @@ -0,0 +1,64 @@ +"""Compare packed MLA low and high physical pages across a signed-32-bit byte boundary.""" + +from pathlib import Path +import sys + +sys.path.insert(0, str(Path(__file__).resolve().parents[2])) + +import torch +from b12x.attention import sparse_mla +from b12x.attention._shared.mla.reference import pack_mla_kv_cache_reference + + +def main(): + torch.manual_seed(20260908) + page, width, rows, heads = 1536, 2048, 4, 112 + base_page = (2**31 // (page * 656)) + 2 + base_slot = base_page * page + packed = pack_mla_kv_cache_reference( + torch.randn(3072, 512, device="cuda", dtype=torch.bfloat16) / 4, + torch.randn(3072, 64, device="cuda", dtype=torch.bfloat16) / 4, + ) + low = packed.view(2, page, 656) + high = torch.empty((base_page + 2, page, 656), device="cuda", dtype=torch.uint8) + high[0].zero_() + high[base_page:].copy_(low) + q = torch.randn(rows, heads, 576, device="cuda", dtype=torch.bfloat16) + q[:, 99:].zero_() + lens = torch.tensor([2048, 2047, 1920, 65], device="cuda", dtype=torch.int32) + indices = torch.arange(width, device="cuda", dtype=torch.int32).repeat(rows, 1) + indices.masked_fill_(indices >= lens[:, None], -1) + high_indices = torch.where(indices >= 0, indices + base_slot, -1) + plan = sparse_mla.plan(sparse_mla.Caps( + device="cuda", num_q_heads=heads, max_q_rows=rows, max_batch=rows, + max_width=116736, dtype=torch.bfloat16, kv_dtype=torch.uint8, + head_dim=576, v_head_dim=512, max_chunks_per_row=64, page_size=page, + partial_dtype=torch.bfloat16, + )) + spec = plan.scratch_specs()[0] + scratch = torch.empty(spec.shape, dtype=spec.dtype, device="cuda") + + def run(cache, slots): + binding = plan.bind( + scratch=scratch, q=q, selected_indices=slots, + cache_seqlens_int32=lens, nsa_cache_seqlens_int32=lens, + ) + out, lse = sparse_mla.run_decode( + binding=binding, kv_cache=cache, sm_scale=192**-0.5, + v_head_dim=512, forced_num_splits=64, split_policy="static", + return_lse=True, lse_scale="natural", + ) + torch.cuda.synchronize() + return out.clone(), lse.clone() + + expected, expected_lse = run(low, indices) + got, got_lse = run(high, high_indices) + assert torch.isfinite(got).all() and torch.isfinite(got_lse).all() + assert torch.equal(got, expected) + assert torch.equal(got_lse, expected_lse) + print({"status": "passed", "page_id": base_page, + "byte_offset": base_slot * 656, "output_and_lse_bit_identical": True}) + + +if __name__ == "__main__": + main() diff --git a/validation/performance/w4a16_inactive_routes_sm120.md b/validation/performance/w4a16_inactive_routes_sm120.md new file mode 100644 index 000000000..3fe4b6dd2 --- /dev/null +++ b/validation/performance/w4a16_inactive_routes_sm120.md @@ -0,0 +1,217 @@ +# Native W4A16 inactive-route validation on SM120 + +Status: **qualified** for native ModelOpt W4A16 fused small-M execution and +FC2-only runtime-M execution on NVIDIA RTX PRO 6000 Blackwell GPUs. + +## Operation contract + +An expert identifier outside `[0, resident_expert_count)` is an inactive route. +The operation must not read weights or scales through that identifier, and its +contribution must be exactly zero. Valid routes retain their identifiers and +weights. Caller-owned route tensors remain unchanged in eager execution and +CUDA Graph replay. + +Route validity is evaluated at the source tensor width. An `int64` identifier +is narrowed to the kernel's `int32` expert address only after the range check, +so values such as `2**32` cannot alias a resident expert. + +Fixed-M fused launches sanitize their compile-time-bounded route table in +shared memory. FC2-only launches accept runtime M, so they validate each route +inline against the runtime resident-expert count. The FC2 implementation must +not index shared route storage whose extent was compiled for a smaller M. + +## Source and runtime identity + +- Repository: `local-inference-lab/b12x` +- Runtime implementation revision: `af354efdbadb8b722da7c696e41d7e35b849b1ec` +- Runtime implementation tree: `dc1d9340849219f99b5ba93e915b0fbd10967028` +- Test-complete revision: `6a41770fb1514c4db03b4ff552380ec6821a3ae9` +- Test-complete tree: `5ea752169a9c4a2863f51f16ed77cc85371f1353` +- Compile-ABI qualification revision: + `97dfe7f837fa5383c5424cbdb2ccfabf57c4480c` +- Route-width qualification revision: + `0e167cdd7d32fbb5b6dcaed88c263ac1fa236c26` +- Production composition base revision: + `c25cdba2c1df7a69b2d7771e4243e12a8fbf19d5` +- Runtime-code qualification tree: + `9b916c2641c85c0ac78a2fa6c4bea0214a0d2cd8` +- Comparison revision for valid FC2 routes: + `debaafe156c9824396178d53e01e5f15d2a2a04a` +- Comparison tree: `ec6edd9da4687f83519fd37bd7322ea0800f0ace` +- Implementation worktree: + `/root/vllm/worktrees/b12x-ii-w4a16-inactive-routes-v2-20260817` +- Comparison worktree: + `/mnt/luke/kimi-k3-runs/pr227-fc2-review-20260817/b12x-debaafe` +- GPU: NVIDIA RTX PRO 6000 Blackwell Workstation Edition, SM120 +- Driver: 610.57.04 +- Runtime: CUDA 13.3, PyTorch 2.13.0, CUTLASS DSL 4.6.2 +- Test image: + `voipmonitor/vllm@sha256:ffb25774eaa90850b4cacfb88ed9e55072818e99bad977f1315c7118e7a730b2` + +## Correctness and memory safety + +The targeted test command selected the staging invariant, an established valid +FC2 case, both narrow and wide invalid-route M=3 cases, and the eager plus CUDA +Graph M=7 case: + +```bash +docker run --rm --gpus device=0 --ipc=host \ + --entrypoint /opt/venv/bin/python3 \ + -e PYTHONPATH=/workspace \ + -e PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True \ + -e CUDA_MODULE_LOADING=LAZY \ + -e TORCH_CUDA_ARCH_LIST=12.0a \ + -e B12X_COMPILE_CACHE_DIR=/cache/compile \ + -e B12X_CUTE_COMPILE_CACHE_DIR=/cache/cute \ + -e CUDA_CACHE_PATH=/cache/cuda \ + -e CUTE_DSL_CACHE_DIR=/cache/cute-dsl \ + -v /root/vllm/worktrees/b12x-ii-w4a16-inactive-routes-v2-20260817:/workspace:ro \ + -v /mnt/luke/kimi-k3-cache/pr227-fc2-review-20260817:/cache:rw \ + voipmonitor/vllm:kimi-k3-production-dspark-lmcache-vllmdf13924-b12xec6edd9-cu133-torch213-20260817-r6 \ + -m pytest -q -s \ + /workspace/tests/moe/test_w4a16_e2e.py::test_w4a16_fc2_runtime_m_does_not_use_fixed_route_staging \ + /workspace/tests/moe/test_w4a16_e2e.py::test_w4a16_fc2_only_consumes_contiguous_bf16_and_native_mxfp4 \ + /workspace/tests/moe/test_w4a16_e2e.py::test_w4a16_fc2_only_zeroes_invalid_routes_at_runtime_m3 \ + /workspace/tests/moe/test_w4a16_e2e.py::test_w4a16_fc2_only_is_cuda_graph_safe_with_preallocated_output +``` + +Result: five parametrized tests passed. The invalid identifiers included `-1`, +`-9`, and the first upper-bound identifier. The M=3 test exercised 256- and +512-element intermediate widths. Assertions covered finite output, nonzero +valid-route output, exact-zero inactive rows, an independent constant oracle, +stable graph replay, and immutable route inputs. + +Source-width validation used the same CUDA 13.3, PyTorch 2.13, and CUTLASS DSL +4.6.2 runtime: + +```bash +python -m pytest -q tests/moe/test_w4a16_e2e.py \ + -k 'test_w4a16_e8m0_native_micro_ignores_inactive_routes_during_graph_replay or test_w4a16_fc2_only_zeroes_invalid_routes_at_runtime_m3' \ + --maxfail=1 +``` + +Result: nine tests passed. Fused execution covered `int32` route tables at +M=1, 2, 4, and 8 plus an `int64` route table at M=1. FC2-only execution covered +`int32` and `int64` route tables with 256- and 512-element intermediate widths. +Both `int64` paths used `2**32` as an inactive identifier. The tests verify +eager output, CUDA Graph replay, exact-zero inactive contributions, valid-route +oracle parity, and immutable route tensors. + +The same invalid-route cases were executed under Compute Sanitizer: + +```bash +docker run --rm --gpus device=0 --ipc=host \ + --entrypoint /usr/local/cuda/bin/compute-sanitizer \ + -e PYTHONPATH=/workspace \ + -e PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True \ + -e CUDA_MODULE_LOADING=LAZY \ + -e TORCH_CUDA_ARCH_LIST=12.0a \ + -e B12X_COMPILE_CACHE_DIR=/cache/compile \ + -e B12X_CUTE_COMPILE_CACHE_DIR=/cache/cute \ + -e CUDA_CACHE_PATH=/cache/cuda \ + -e CUTE_DSL_CACHE_DIR=/cache/cute-dsl \ + -v /root/vllm/worktrees/b12x-ii-w4a16-inactive-routes-v2-20260817:/workspace:ro \ + -v /mnt/luke/kimi-k3-cache/pr227-fc2-review-20260817:/cache:rw \ + voipmonitor/vllm:kimi-k3-production-dspark-lmcache-vllmdf13924-b12xec6edd9-cu133-torch213-20260817-r6 \ + --tool memcheck --error-exitcode=99 \ + /opt/venv/bin/python3 -m pytest -q -s \ + /workspace/tests/moe/test_w4a16_e2e.py::test_w4a16_fc2_only_zeroes_invalid_routes_at_runtime_m3 \ + /workspace/tests/moe/test_w4a16_e2e.py::test_w4a16_fc2_only_is_cuda_graph_safe_with_preallocated_output +``` + +Result: three parametrized tests passed and Compute Sanitizer reported +`ERROR SUMMARY: 0 errors`. + +The generic direct-micro compiler receives the resident-expert limit before +the runtime M, grid extent, and CUDA stream arguments. GPU compile tests cover +the fused body and both split phases, preventing a positional ABI change from +binding the stream as an integer launch parameter. + +A composed source qualification applied this implementation to B12X master +`c25cdba2c1df7a69b2d7771e4243e12a8fbf19d5`. The following suites produced +329 passes and 18 skips: + +- `tests/moe/test_w4a16_e2e.py` +- `tests/moe/test_w4a16_mixed_trellis.py` +- `tests/moe/test_w4a16_route_pack.py` +- `tests/moe/test_moe_launch_param_regression.py` +- `tests/test_packaging.py` + +The result covers native and Trellis routes, fused and split NVFP4 compilation, +eager execution, CUDA Graph replay, mixed K3/K4/K5 dispatch, route packing, +and packaged source availability. + +## FC2 valid-route latency diagnostic + +The diagnostic used valid routes, a caller-owned output, 200 graph warmups, +nine samples, and 5,000 CUDA Graph replays per sample. Lower latency is better; +the reported ratio is implementation latency divided by comparison latency. + +| Source | M | Raw microseconds per replay | Median | +|---|---:|---|---:| +| comparison `debaafe1` | 2 | 2.226042, 2.201734, 2.150394, 2.139565, 2.154854, 2.202118, 2.181280, 2.181818, 2.180480 | 2.181280 | +| implementation `af354efd` | 2 | 2.170336, 2.137523, 2.150829, 2.141190, 2.157146, 2.165818, 2.146022, 2.139597, 2.221146 | 2.150829 | +| implementation `af354efd` | 7 | 4.098701, 4.098419, 4.098074, 4.097894, 4.098157, 4.097914, 4.098016, 4.097914, 4.098009 | 4.098016 | + +The M=2 median ratio is 0.9860. This diagnostic establishes that runtime route +validation did not produce a measurable valid-route FC2 regression. It is not +a speedup claim because the arms were measured sequentially without a locked +clock. + +## Full-model integration + +The candidate image copied the implementation-revision `b12x` package over the +published integration image without changing vLLM, LMCache, model weights, or +launch arguments: + +- Candidate image ID: + `sha256:d2060dab541504dfec43e352a817e353e9672eec34231fab734db340b69c3154` +- Base image digest: + `sha256:ffb25774eaa90850b4cacfb88ed9e55072818e99bad977f1315c7118e7a730b2` +- Candidate Dockerfile SHA-256: + `0707481e9ea583cf8ac13524c730e6d1913a60d53aa5409d5c7b0227fc2d3549` + +Conditions were the official Kimi-K3 MXFP4 target, the Inferact seven-token +DSpark draft, TP16/DCP16, FP8 target KV cache, a 1,000,000-token model limit, +native vision, LMCache, and CUDA Graph capture size eight. InstantTensor loaded +90.48 GiB per GPU. Physical target KV capacity was 1,033,126 tokens. CUDA Graph +capture completed and the API became healthy. + +The benchmark program is +[`benchmark-kimi-k3-dspark-decode.py`](https://github.com/local-inference-lab/rtx6kpro/blob/a82029c0ffa9c1cccfa9215e927c3db0ae2aeb57/models/kimi-k3/tools/benchmark-kimi-k3-dspark-decode.py), +SHA-256 `b465fb785fc11b5b2941510eca1cfbdd159a55d8258c827a3db110309768023b`. +The command used 256 stored prompt tokens, 1,024 generated tokens, temperature +zero, seed one, two warmups, and eight measured runs: + +```bash +python3 models/kimi-k3/tools/benchmark-kimi-k3-dspark-decode.py \ + --url http://127.0.0.1:8001 \ + --model Kimi-K3-MXFP4-DSpark7-DCP16-1M \ + --token-file models/kimi-k3/tools/decode-baseline-256-token-ids.json \ + --prompt-tokens 256 --max-tokens 1024 --warmups 2 --runs 8 \ + --output-dir full-model-dspark-normalized-256x1024 +``` + +Target-cycle rate is the acceptance-independent performance metric. Higher is +better; the reported ratio is candidate median divided by base-image median. + +| Arm | Raw target cycles/s | Median | +|---|---|---:| +| base image | 31.384470, 31.589480, 31.633320, 31.563746, 31.237360, 31.340825, 31.474961, 31.232879 | 31.429716 | +| candidate image | 31.391607, 31.526917, 31.527495, 31.454522, 31.589554, 31.332777, 31.679354, 31.512322 | 31.519620 | + +The candidate/base median ratio is 1.0029. Emitted throughput was 130.446 +tokens/s for the candidate and 118.773 tokens/s for the base image, but that +difference is not attributed to the kernel because median draft acceptance was +0.4476 and 0.3973 respectively. + +During a separate sustained 4,096-token candidate run, all 16 GPUs remained in +P1 at 100% utilization, memory clocks were 13,365 MHz, SM clocks ranged from +2,670 to 2,865 MHz, and the active clock-event mask was zero on every GPU. The +run measured 30.918 target cycles/s. GPU 0 was +`GPU-d8438b2d-f000-a617-5dcc-0197ce0365a3`. + +Conclusion: the implementation satisfies inactive-route address safety and +zero-contribution semantics for fused fixed-M and FC2 runtime-M execution. The +qualified full-model profile loads, captures, and decodes without a target-cycle +regression. diff --git a/validation/performance/w4a16_prefill_fused_sum_rtx_pro_6000_blackwell.json.gz b/validation/performance/w4a16_prefill_fused_sum_rtx_pro_6000_blackwell.json.gz new file mode 100644 index 000000000..0238b7a47 Binary files /dev/null and b/validation/performance/w4a16_prefill_fused_sum_rtx_pro_6000_blackwell.json.gz differ diff --git a/validation/performance/w4a16_prefill_fused_sum_rtx_pro_6000_blackwell.md b/validation/performance/w4a16_prefill_fused_sum_rtx_pro_6000_blackwell.md new file mode 100644 index 000000000..b8316e562 --- /dev/null +++ b/validation/performance/w4a16_prefill_fused_sum_rtx_pro_6000_blackwell.md @@ -0,0 +1,161 @@ +# W4A16 bounded prefill reduction qualification + +Status: **qualified** for the full-model serving hardware, source composition, +and launch contract recorded below. The isolated kernel latency comparison is +**diagnostic** because its alternating materialized/fused process sequence did +not retain a contemporaneous GPU-mode sample. + +## Purpose + +This record validates the opt-in W4A16 BF16 prefill path selected by +`B12X_W4A16_PREFILL_FUSED_SUM=1`. The path reduces routed FC2 results directly +into one FP32 row per input token. Its purpose is to bound scratch memory so a +4,096-token Kimi-K3 scheduler chunk fits beside the model and a physical +1,057,049-token KV cache. + +The machine-readable receipt +`w4a16_prefill_fused_sum_rtx_pro_6000_blackwell.json.gz` contains every +CUDA-event timing sample from the A-B-B-A kernel comparison. Its SHA-256 is +`65db1e3220c45246e0eede79da4d9ebb448f27c631bfcf23194a5248efdeb551`. + +## Source and hardware + +- Repository: `local-inference-lab/b12x` +- Implementation revision: `0c3be37138f74a6d0213c10202e0077c2d2a44da` +- Implementation tree: `4794737159345008fa897579d6d3b67ed671a151` +- Materialized-path base revision: + `c25cdba2c1df7a69b2d7771e4243e12a8fbf19d5` +- Measured worktree: + `/root/vllm/worktrees/b12x-k3-prefill-fused-reduce-20260819` +- Container: + `voipmonitor/vllm@sha256:bd8a4be5e87c89f37548ee0502c1a0dc186e9058d57f3278927c1ef5d01e65fa` +- CUDA runtime: 13.3 +- PyTorch: `2.13.0a0+9186a08` +- Driver: `610.57.04` +- Microbenchmark GPU: physical GPU 0, + `NVIDIA RTX PRO 6000 Blackwell Workstation Edition` +- Microbenchmark GPU UUID: + `GPU-d8438b2d-f000-a617-5dcc-0197ce0365a3` +- Full-model hardware: 16 GPUs of the same type + +The two persistent MoE artifacts use CUTLASS DSL 4.6.2 and PTXAS +`Cuda compilation tools, release 13.3, V13.3.27`: + +| Arm | Compile-cache key | Object SHA-256 | +|---|---|---| +| Materialized BF16 routes | `d44aed42740f285a69b9bd972756111de92eb4ec8f7a21e683be3750a58eeeb6` | `34c5396f3d9f140fad4fed016d11516f19e8e063bff3729764811abf85e0f96a` | +| FP32 bounded reduction | `19a9b91e405c43b3eea9751b6ac02bf85c94918e254bfe005a064cf37eb67579` | `fb313da8431a46cc86ab1a2e6972dcf0f10af76554ea7a71cc666a617aa494b3` | + +## Scratch contract + +The measured rank shape has 4,096 tokens, hidden width 7,168, tensor-parallel +intermediate width 192, 896 router experts, top-k 16, BF16 activations, and the +`situ` activation. + +| Caller-owned W4A16 scratch | Bytes per rank | MiB per rank | +|---|---:|---:| +| Materialized route output | 1,063,787,080 | 1,014.51 | +| FP32 fused reduction | 292,035,144 | 278.51 | +| Memory released | 771,751,936 | 736.00 | + +The fused total includes a 48 MiB FC1 cache, a 112 MiB FP32 per-token +accumulator, a 24 MiB activation cache, GEMM accumulation scratch, and route +metadata. The materialized total stores one 7,168-element BF16 output for each +of the 65,536 token-route pairs. + +## Kernel timing + +Arm A sets `B12X_W4A16_PREFILL_FUSED_SUM=0` and materializes BF16 route output. +Arm B sets the flag to `1` and performs direct FP32 per-token reduction. Both +arms use identical synthetic weights, activations, routing, compiled source, +L2 flushing, and CUDA-graph timing. + +Each A1-B1-B2-A2 process records ten repeats of 100 CUDA-event samples after +20 warmup iterations per repeat. The aggregate statistic is the median of the +20 per-repeat medians for each arm. + +| Arm | Aggregate median | Samples | +|---|---:|---:| +| Materialized BF16 routes | 5,214.17 us | 2,000 | +| FP32 fused reduction | 5,304.32 us | 2,000 | + +The ratio is fused latency divided by materialized latency: `1.01729`. The +bounded reduction is 1.73% slower as an isolated MoE launch. It is not claimed +as a kernel-latency optimization. Its serving gain comes from replacing four +1,024-token MoE launches with one 4,096-token launch under the same physical KV +allocation. + +The A-B-B-A processes did not retain P-state or throttle-mask snapshots. Their +latency ratio is diagnostic rather than hardware-qualified. The full-model +measurement below retained 7,072 GPU-mode samples and is the authoritative +serving-performance result. + +Reproduce either arm with this command and set the feature flag to `0` or `1`: + +```bash +B12X_W4A16_PREFILL_FUSED_SUM=1 \ +python benchmarks/benchmark_moe.py \ + --model-profile kimi-k3-mxfp4-shape \ + --batch-sizes 4096 \ + --quant-mode w4a16 \ + --validate none \ + --reference none \ + --compare-prefill-fused-sum \ + --graph-only \ + --warmup 20 \ + --iters 100 \ + --repeats 10 \ + --timing-json /tmp/w4a16-prefill-arm.json +``` + +## Full-model result + +The serving test uses the official `moonshotai/Kimi-K3` MXFP4 target, the +`Inferact/Kimi-K3-DSpark` draft, TP16/DCP16, one active sequence, Triton KDA, +B12X MLA, disabled prefix caching, and a 1,325,000,000-byte FP8 KV allocation +per rank. `MAX_NUM_BATCHED_TOKENS=4102` reserves six DSpark slots and leaves an +exact 4,096-token prefill scheduler limit. + +The complete launch command, source-composition procedure, benchmark command, +and machine-readable serving receipt are published in the +[Kimi-K3 full-MXFP4 4096-token prefill profile](https://github.com/local-inference-lab/rtx6kpro/blob/15d7347108673ad73aff51fdb40e74fc292401a3/models/kimi-k3/full-mxfp4-p4096-prefill.md). +The compressed receipt beside this document also contains every measured +request's time to first token and effective prefill throughput. + +| Prompt tokens | Four 1,024-token MoE launches | One 4,096-token MoE launch | Change | +|---:|---:|---:|---:| +| 8,192 | 2,723.1 tok/s | 3,861.7 tok/s | +41.81% | +| 32,768 | 2,897.0 tok/s | 3,732.5 tok/s | +28.84% | +| 65,535 | 2,839.7 tok/s | 3,554.4 tok/s | +25.17% | + +The 8,192-token result is the median of six requests. The other results are +medians of three requests. Every request supplies token IDs directly, emits +one token, measures streamed time to first token, and uses a unique cache salt. +The physical KV capacity reported by vLLM is 1,057,049 tokens. + +During the full-model measurements, every GPU reached 100% utilization, every +active sample was in P1, active SM clocks ranged from 2,595 to 2,880 MHz, and +the active clock median was 2,790 MHz. All 7,072 recorded GPU-mode samples had +an active throttle mask of zero. + +## Correctness and regression coverage + +- The default W4A16 GPU suite passed 292 tests and skipped 16 unsupported + cases. +- Five focused planner, arena, activation-amax, CUDA Graph, and fused-output + tests passed with the opt-in path selected. +- Fused output had cosine similarity 1.0 and maximum BF16 absolute difference + 0.015625 against materialized reduction. +- CUDA Graph replay retained scratch addresses, performed no replay-time + allocation, consumed the FP32 accumulator, and produced finite nonzero + output. +- A DSpark decode sanity test measured 31.486 normalized target cycles/s. + +## Limitations + +FP32 global additions use relaxed ordering. Bitwise equality with materialized +route reduction and bitwise repeat determinism are unsupported. The isolated +kernel measurement uses synthetic Kimi-K3-shaped inputs; the serving result is +the end-to-end evidence for scheduler throughput. Qualification does not cover +FP16, full-rotation Trellis, activation-amax capture, other GPU types, or other +topologies.