From d2146df1c243a5a5098bed3fc3c37eafc7562824 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Dani=C3=ABl=20de=20Kok?= Date: Sat, 4 Jul 2026 06:06:20 +0000 Subject: [PATCH 1/6] Fix XPU and ROCm issues for Torch 2.13 --- activation/flake.lock | 6 +++--- aiter-flash-attn/flake.lock | 6 +++--- aiter-kernels/flake.lock | 6 +++--- aiter-rope/flake.lock | 6 +++--- bitsandbytes-mps/flake.lock | 6 +++--- causal-conv1d/flake.lock | 6 +++--- cv-utils/flake.lock | 6 +++--- deep-gemm/flake.lock | 6 +++--- deformable-detr/flake.lock | 6 +++--- finegrained-fp8/flake.lock | 6 +++--- flash-attn-ops/flake.lock | 6 +++--- flash-attn2/flake.lock | 6 +++--- flash-attn3/flake.lock | 6 +++--- flash-attn4/flake.lock | 6 +++--- flash-mla/flake.lock | 6 +++--- fp8-fbgemm/flake.lock | 6 +++--- gpt-oss-metal-kernels/flake.lock | 6 +++--- gpt-oss-triton-kernels/flake.lock | 6 +++--- layer-norm/flake.lock | 6 +++--- liger-kernels/flake.lock | 6 +++--- mamba-ssm/flake.lock | 6 +++--- megablocks/flake.lock | 6 +++--- metal-flash-sdpa/flake.lock | 6 +++--- mlx-quantization-metal-kernels/flake.lock | 6 +++--- mlx-rmsnorm/flake.lock | 6 +++--- mra/flake.lock | 6 +++--- msa/flake.lock | 6 +++--- paged-attention/flake.lock | 6 +++--- punica-sgmv/flake.lock | 6 +++--- quantization-bitsandbytes/flake.lock | 6 +++--- quantization-eetq/flake.lock | 6 +++--- quantization-gptq/flake.lock | 6 +++--- relu/flake.lock | 6 +++--- rmsnorm/flake.lock | 6 +++--- rotary/flake.lock | 6 +++--- rwkv/flake.lock | 6 +++--- sage-attention/flake.lock | 6 +++--- scattermoe/flake.lock | 6 +++--- sgl-flash-attn3/flake.lock | 6 +++--- sonic-moe/flake.lock | 6 +++--- tinygrad-rms/flake.lock | 6 +++--- trimul-gpumode/flake.lock | 6 +++--- triton-kernels/flake.lock | 6 +++--- vllm-flash-attn3/flake.lock | 6 +++--- vllm-moe/flake.lock | 6 +++--- yoso/flake.lock | 6 +++--- 46 files changed, 138 insertions(+), 138 deletions(-) diff --git a/activation/flake.lock b/activation/flake.lock index 829e5c03..5e13de28 100644 --- a/activation/flake.lock +++ b/activation/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/aiter-flash-attn/flake.lock b/aiter-flash-attn/flake.lock index 829e5c03..5e13de28 100644 --- a/aiter-flash-attn/flake.lock +++ b/aiter-flash-attn/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/aiter-kernels/flake.lock b/aiter-kernels/flake.lock index 829e5c03..5e13de28 100644 --- a/aiter-kernels/flake.lock +++ b/aiter-kernels/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/aiter-rope/flake.lock b/aiter-rope/flake.lock index 829e5c03..5e13de28 100644 --- a/aiter-rope/flake.lock +++ b/aiter-rope/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/bitsandbytes-mps/flake.lock b/bitsandbytes-mps/flake.lock index 829e5c03..5e13de28 100644 --- a/bitsandbytes-mps/flake.lock +++ b/bitsandbytes-mps/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/causal-conv1d/flake.lock b/causal-conv1d/flake.lock index 829e5c03..5e13de28 100644 --- a/causal-conv1d/flake.lock +++ b/causal-conv1d/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/cv-utils/flake.lock b/cv-utils/flake.lock index 829e5c03..5e13de28 100644 --- a/cv-utils/flake.lock +++ b/cv-utils/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/deep-gemm/flake.lock b/deep-gemm/flake.lock index 829e5c03..5e13de28 100644 --- a/deep-gemm/flake.lock +++ b/deep-gemm/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/deformable-detr/flake.lock b/deformable-detr/flake.lock index 829e5c03..5e13de28 100644 --- a/deformable-detr/flake.lock +++ b/deformable-detr/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/finegrained-fp8/flake.lock b/finegrained-fp8/flake.lock index 829e5c03..5e13de28 100644 --- a/finegrained-fp8/flake.lock +++ b/finegrained-fp8/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/flash-attn-ops/flake.lock b/flash-attn-ops/flake.lock index 829e5c03..5e13de28 100644 --- a/flash-attn-ops/flake.lock +++ b/flash-attn-ops/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/flash-attn2/flake.lock b/flash-attn2/flake.lock index 829e5c03..5e13de28 100644 --- a/flash-attn2/flake.lock +++ b/flash-attn2/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/flash-attn3/flake.lock b/flash-attn3/flake.lock index 829e5c03..5e13de28 100644 --- a/flash-attn3/flake.lock +++ b/flash-attn3/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/flash-attn4/flake.lock b/flash-attn4/flake.lock index 829e5c03..5e13de28 100644 --- a/flash-attn4/flake.lock +++ b/flash-attn4/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/flash-mla/flake.lock b/flash-mla/flake.lock index 829e5c03..5e13de28 100644 --- a/flash-mla/flake.lock +++ b/flash-mla/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/fp8-fbgemm/flake.lock b/fp8-fbgemm/flake.lock index 829e5c03..5e13de28 100644 --- a/fp8-fbgemm/flake.lock +++ b/fp8-fbgemm/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/gpt-oss-metal-kernels/flake.lock b/gpt-oss-metal-kernels/flake.lock index 829e5c03..5e13de28 100644 --- a/gpt-oss-metal-kernels/flake.lock +++ b/gpt-oss-metal-kernels/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/gpt-oss-triton-kernels/flake.lock b/gpt-oss-triton-kernels/flake.lock index 829e5c03..5e13de28 100644 --- a/gpt-oss-triton-kernels/flake.lock +++ b/gpt-oss-triton-kernels/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/layer-norm/flake.lock b/layer-norm/flake.lock index 829e5c03..5e13de28 100644 --- a/layer-norm/flake.lock +++ b/layer-norm/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/liger-kernels/flake.lock b/liger-kernels/flake.lock index 829e5c03..5e13de28 100644 --- a/liger-kernels/flake.lock +++ b/liger-kernels/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/mamba-ssm/flake.lock b/mamba-ssm/flake.lock index 829e5c03..5e13de28 100644 --- a/mamba-ssm/flake.lock +++ b/mamba-ssm/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/megablocks/flake.lock b/megablocks/flake.lock index 829e5c03..5e13de28 100644 --- a/megablocks/flake.lock +++ b/megablocks/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/metal-flash-sdpa/flake.lock b/metal-flash-sdpa/flake.lock index 829e5c03..5e13de28 100644 --- a/metal-flash-sdpa/flake.lock +++ b/metal-flash-sdpa/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/mlx-quantization-metal-kernels/flake.lock b/mlx-quantization-metal-kernels/flake.lock index 829e5c03..5e13de28 100644 --- a/mlx-quantization-metal-kernels/flake.lock +++ b/mlx-quantization-metal-kernels/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/mlx-rmsnorm/flake.lock b/mlx-rmsnorm/flake.lock index 829e5c03..5e13de28 100644 --- a/mlx-rmsnorm/flake.lock +++ b/mlx-rmsnorm/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/mra/flake.lock b/mra/flake.lock index 829e5c03..5e13de28 100644 --- a/mra/flake.lock +++ b/mra/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/msa/flake.lock b/msa/flake.lock index 829e5c03..5e13de28 100644 --- a/msa/flake.lock +++ b/msa/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/paged-attention/flake.lock b/paged-attention/flake.lock index 829e5c03..5e13de28 100644 --- a/paged-attention/flake.lock +++ b/paged-attention/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/punica-sgmv/flake.lock b/punica-sgmv/flake.lock index 829e5c03..5e13de28 100644 --- a/punica-sgmv/flake.lock +++ b/punica-sgmv/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/quantization-bitsandbytes/flake.lock b/quantization-bitsandbytes/flake.lock index 829e5c03..5e13de28 100644 --- a/quantization-bitsandbytes/flake.lock +++ b/quantization-bitsandbytes/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/quantization-eetq/flake.lock b/quantization-eetq/flake.lock index 829e5c03..5e13de28 100644 --- a/quantization-eetq/flake.lock +++ b/quantization-eetq/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/quantization-gptq/flake.lock b/quantization-gptq/flake.lock index 829e5c03..5e13de28 100644 --- a/quantization-gptq/flake.lock +++ b/quantization-gptq/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/relu/flake.lock b/relu/flake.lock index 829e5c03..5e13de28 100644 --- a/relu/flake.lock +++ b/relu/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/rmsnorm/flake.lock b/rmsnorm/flake.lock index 829e5c03..5e13de28 100644 --- a/rmsnorm/flake.lock +++ b/rmsnorm/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/rotary/flake.lock b/rotary/flake.lock index 829e5c03..5e13de28 100644 --- a/rotary/flake.lock +++ b/rotary/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/rwkv/flake.lock b/rwkv/flake.lock index 829e5c03..5e13de28 100644 --- a/rwkv/flake.lock +++ b/rwkv/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/sage-attention/flake.lock b/sage-attention/flake.lock index 829e5c03..5e13de28 100644 --- a/sage-attention/flake.lock +++ b/sage-attention/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/scattermoe/flake.lock b/scattermoe/flake.lock index 829e5c03..5e13de28 100644 --- a/scattermoe/flake.lock +++ b/scattermoe/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/sgl-flash-attn3/flake.lock b/sgl-flash-attn3/flake.lock index 829e5c03..5e13de28 100644 --- a/sgl-flash-attn3/flake.lock +++ b/sgl-flash-attn3/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/sonic-moe/flake.lock b/sonic-moe/flake.lock index 829e5c03..5e13de28 100644 --- a/sonic-moe/flake.lock +++ b/sonic-moe/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/tinygrad-rms/flake.lock b/tinygrad-rms/flake.lock index 829e5c03..5e13de28 100644 --- a/tinygrad-rms/flake.lock +++ b/tinygrad-rms/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/trimul-gpumode/flake.lock b/trimul-gpumode/flake.lock index 829e5c03..5e13de28 100644 --- a/trimul-gpumode/flake.lock +++ b/trimul-gpumode/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/triton-kernels/flake.lock b/triton-kernels/flake.lock index 829e5c03..5e13de28 100644 --- a/triton-kernels/flake.lock +++ b/triton-kernels/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/vllm-flash-attn3/flake.lock b/vllm-flash-attn3/flake.lock index 829e5c03..5e13de28 100644 --- a/vllm-flash-attn3/flake.lock +++ b/vllm-flash-attn3/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/vllm-moe/flake.lock b/vllm-moe/flake.lock index 829e5c03..5e13de28 100644 --- a/vllm-moe/flake.lock +++ b/vllm-moe/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { diff --git a/yoso/flake.lock b/yoso/flake.lock index 829e5c03..5e13de28 100644 --- a/yoso/flake.lock +++ b/yoso/flake.lock @@ -41,11 +41,11 @@ "rust-overlay": "rust-overlay" }, "locked": { - "lastModified": 1782921863, - "narHash": "sha256-xROT3h5roOQ0oMQtAw7tVqALFGQIPqGWLwwvNDsZQjU=", + "lastModified": 1783096525, + "narHash": "sha256-MVB8qI/KDrCBQeJiBb32GdQBV7BmsI/MF0Jq+2rbSwo=", "owner": "huggingface", "repo": "kernels", - "rev": "b79face32a072b46e4e95c37869b3d1cc2e7735f", + "rev": "99cf6beb01f360e0d49ea630781481686d258783", "type": "github" }, "original": { From e775b0a5597d63cabf7150f916913891e99a30e7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Dani=C3=ABl=20de=20Kok?= Date: Sun, 5 Jul 2026 14:16:21 +0000 Subject: [PATCH 2/6] scattermoe: fix incorrect op registrations --- .../torch-ext/scattermoe/kernels/ops.py | 460 ++++++++++++------ .../torch-ext/scattermoe/parallel_experts.py | 177 +++++-- 2 files changed, 447 insertions(+), 190 deletions(-) diff --git a/scattermoe/torch-ext/scattermoe/kernels/ops.py b/scattermoe/torch-ext/scattermoe/kernels/ops.py index b47534c2..cb590a24 100644 --- a/scattermoe/torch-ext/scattermoe/kernels/ops.py +++ b/scattermoe/torch-ext/scattermoe/kernels/ops.py @@ -1,20 +1,28 @@ +from typing import Optional + import torch import triton import triton.language as tl -from typing import Optional + +from .._ops import add_op_namespace_prefix BLOCK_M = 128 ALLOW_TF32 = True - - @triton.jit def _compute_expert_block( - E_idx, E_mask, + E_idx, + E_mask, M_in_idx, - N_block, N_mask, - X_ptr, stride_xm, stride_xk, - W_ptr, stride_we, stride_wk, stride_wn, + N_block, + N_mask, + X_ptr, + stride_xm, + stride_xk, + W_ptr, + stride_we, + stride_wk, + stride_wn, K, acc, no_k_mask, @@ -24,7 +32,12 @@ def _compute_expert_block( K_block = tl.arange(0, BLOCK_K) X_blk_ptrs = X_ptr + M_in_idx[:, None] * stride_xm + K_block[None, :] * stride_xk - W_blk_ptrs = W_ptr + K_block[:, None] * stride_wk + N_block[None, :] * stride_wn + E_idx * stride_we + W_blk_ptrs = ( + W_ptr + + K_block[:, None] * stride_wk + + N_block[None, :] * stride_wn + + E_idx * stride_we + ) iters = tl.cdiv(K, BLOCK_K) for K_block_id in range(iters): @@ -44,30 +57,53 @@ def _compute_expert_block( def _scatter2scatter_configs(): return [ - triton.Config({'BLOCK_N': 128, 'BLOCK_K': 32}, num_stages=4, num_warps=4), + triton.Config({"BLOCK_N": 128, "BLOCK_K": 32}, num_stages=4, num_warps=4), ] -@triton.autotune(configs=_scatter2scatter_configs(), key=['M', 'N', 'K'], ) -@triton.heuristics({ - "NO_K_MASK": lambda args: (args['K'] % args['BLOCK_K']) == 0, - "NO_N_MASK": lambda args: (args['N'] % args['BLOCK_N']) == 0, -}) + +@triton.autotune( + configs=_scatter2scatter_configs(), + key=["M", "N", "K"], +) +@triton.heuristics( + { + "NO_K_MASK": lambda args: (args["K"] % args["BLOCK_K"]) == 0, + "NO_N_MASK": lambda args: (args["N"] % args["BLOCK_N"]) == 0, + } +) @triton.jit def _scatter2scatter( - X_ptr, stride_xm: tl.constexpr, stride_xk: tl.constexpr, - W_ptr, stride_we, stride_wk: tl.constexpr, stride_wn: tl.constexpr, - Y_ptr, stride_ym: tl.constexpr, stride_yn: tl.constexpr, - B_ptr, stride_be: tl.constexpr, stride_bn: tl.constexpr, - grouped_idx_ptr, expert_idxs_ptr, + X_ptr, + stride_xm: tl.constexpr, + stride_xk: tl.constexpr, + W_ptr, + stride_we, + stride_wk: tl.constexpr, + stride_wn: tl.constexpr, + Y_ptr, + stride_ym: tl.constexpr, + stride_yn: tl.constexpr, + B_ptr, + stride_be: tl.constexpr, + stride_bn: tl.constexpr, + grouped_idx_ptr, + expert_idxs_ptr, # block_start_idx_ptr, FAN_OUT: tl.constexpr, - M, K: tl.constexpr, N: tl.constexpr, E: tl.constexpr, - BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, + M, + K: tl.constexpr, + N: tl.constexpr, + E: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, + BLOCK_K: tl.constexpr, ACC_TYPE: tl.constexpr, # OUT_M, allow_tf32: tl.constexpr, - x_grouped: tl.constexpr, y_grouped: tl.constexpr, - NO_K_MASK: tl.constexpr, NO_N_MASK: tl.constexpr + x_grouped: tl.constexpr, + y_grouped: tl.constexpr, + NO_K_MASK: tl.constexpr, + NO_N_MASK: tl.constexpr, ): pid = tl.program_id(axis=0) @@ -95,10 +131,18 @@ def _scatter2scatter( else: M_in_idx = E_M_idx // FAN_OUT acc = _compute_expert_block( - E_idx, E_mask, - M_in_idx, N_block, N_mask, - X_ptr, stride_xm, stride_xk, - W_ptr, stride_we, stride_wk, stride_wn, + E_idx, + E_mask, + M_in_idx, + N_block, + N_mask, + X_ptr, + stride_xm, + stride_xk, + W_ptr, + stride_we, + stride_wk, + stride_wn, K, acc, no_k_mask, @@ -117,10 +161,18 @@ def _scatter2scatter( Y_blk_ptrs = Y_ptr + (M_out_idx[:, None] * stride_ym + N_block[None, :] * stride_yn) tl.store(Y_blk_ptrs, acc, mask=M_boundary_mask[:, None] & N_mask[None, :]) -def scatter2scatter(X, W, sorted_expert_idxs, sorted_scattered_idxs, k, - b=None, - x_grouped=False, y_grouped=False, - out=None): + +def scatter2scatter( + X, + W, + sorted_expert_idxs, + sorted_scattered_idxs, + k, + b=None, + x_grouped=False, + y_grouped=False, + out=None, +): assert sorted_scattered_idxs.size(0) == sorted_expert_idxs.size(0) assert sorted_scattered_idxs.size(0) == X.size(0) * k # Pre-kernel setup @@ -132,25 +184,38 @@ def scatter2scatter(X, W, sorted_expert_idxs, sorted_scattered_idxs, k, assert out.size(0) == L_scattered and out.size(1) == y_dim output = out - scatter2scatter_compileable(output, W, X, k, sorted_expert_idxs, sorted_scattered_idxs, - b, x_grouped, y_grouped) + scatter2scatter_compileable( + output, + W, + X, + k, + sorted_expert_idxs, + sorted_scattered_idxs, + b, + x_grouped, + y_grouped, + ) return output -@torch.library.custom_op("scattermoe::scatter2scatter", mutates_args={"output"}) +@torch.library.custom_op( + add_op_namespace_prefix("scatter2scatter"), mutates_args={"output"} +) def scatter2scatter_compileable( - output: torch.Tensor, - W: torch.Tensor, - X: torch.Tensor, - k: int, - sorted_expert_idxs: torch.Tensor, - sorted_scattered_idxs: torch.Tensor, - b: Optional[torch.Tensor], - x_grouped: bool, y_grouped: bool) -> None: + output: torch.Tensor, + W: torch.Tensor, + X: torch.Tensor, + k: int, + sorted_expert_idxs: torch.Tensor, + sorted_scattered_idxs: torch.Tensor, + b: Optional[torch.Tensor], + x_grouped: bool, + y_grouped: bool, +) -> None: def grid(META): grid_num = ( - triton.cdiv(sorted_expert_idxs.size(0), META["BLOCK_M"]) * - triton.cdiv(META['N'], META['BLOCK_N']), + triton.cdiv(sorted_expert_idxs.size(0), META["BLOCK_M"]) + * triton.cdiv(META["N"], META["BLOCK_N"]), ) return grid_num @@ -162,32 +227,46 @@ def grid(META): _scatter2scatter[grid]( # X_ptr, stride_xm, stride_xk, - X, X.stride(0), X.stride(1), + X, + X.stride(0), + X.stride(1), # W_ptr, stride_we, stride_wk, stride_wn, - W, W.stride(0), W.stride(1), W.stride(2), + W, + W.stride(0), + W.stride(1), + W.stride(2), # Y_ptr, stride_ym, stride_yn, - output, output.stride(0), output.stride(1), + output, + output.stride(0), + output.stride(1), # B_ptr, stride_be, stride_bk - b, stride_be, stride_bk, + b, + stride_be, + stride_bk, grouped_idx_ptr=sorted_scattered_idxs, expert_idxs_ptr=sorted_expert_idxs, # block_start_idx_ptr=padded_block_idxs, FAN_OUT=k, M=X.size(0), K=X.size(1), - N=output.size(1), E=W.size(0), + N=output.size(1), + E=W.size(0), BLOCK_M=BLOCK_M, ACC_TYPE=tl.float32, allow_tf32=ALLOW_TF32, - x_grouped=x_grouped, y_grouped=y_grouped, + x_grouped=x_grouped, + y_grouped=y_grouped, ) def _config_XtY(): return [ - triton.Config({'BLOCK_N': 128, 'BLOCK_K': 128, 'BLOCK_M': 32}, num_stages=4, num_warps=4), + triton.Config( + {"BLOCK_N": 128, "BLOCK_K": 128, "BLOCK_M": 32}, num_stages=4, num_warps=4 + ), ] + def group_bwd_W(DY, X, expert_offsets, E, has_bias=False): DWt = torch.zeros((E, DY.size(-1), X.size(-1)), device=DY.device, dtype=DY.dtype) DW = DWt.permute(0, 2, 1) @@ -199,21 +278,22 @@ def group_bwd_W(DY, X, expert_offsets, E, has_bias=False): return DW, Db -@torch.library.custom_op("scattermoe::groupXtY", mutates_args={"DW"}) +@torch.library.custom_op(add_op_namespace_prefix("groupXtY"), mutates_args={"DW"}) def groupXtY_compileable( - E: int, - DW: torch.Tensor, - Db: Optional[torch.Tensor], - DY: torch.Tensor, - X: torch.Tensor, - expert_offsets: torch.Tensor) -> None: + E: int, + DW: torch.Tensor, + Db: Optional[torch.Tensor], + DY: torch.Tensor, + X: torch.Tensor, + expert_offsets: torch.Tensor, +) -> None: def grid(META): grid = ( - E * triton.cdiv(META['K'], META['BLOCK_K']), - triton.cdiv(META['N'], META['BLOCK_N']), + E * triton.cdiv(META["K"], META["BLOCK_K"]), + triton.cdiv(META["N"], META["BLOCK_N"]), ) return grid - + if Db is None: stride_dbe = 0 stride_dbn = 0 @@ -222,40 +302,70 @@ def grid(META): _groupXtY[grid]( # DY_ptr, stride_dym, stride_dyk, - DY, DY.stride(0), DY.stride(1), + DY, + DY.stride(0), + DY.stride(1), # X_ptr, stride_xm, stride_xn, - X, X.stride(0), X.stride(1), + X, + X.stride(0), + X.stride(1), # DW_ptr, stride_dwe, stride_dwk, stride_dwn, - DW, DW.stride(0), DW.stride(1), DW.stride(2), + DW, + DW.stride(0), + DW.stride(1), + DW.stride(2), # Db_ptr, stride_dwe, stride_dbn, - Db, stride_dbe, stride_dbn, + Db, + stride_dbe, + stride_dbn, # expert_offsets_ptr, expert_offsets, # K: tl.constexpr, N: tl.constexpr, - M=DY.size(0), N=DY.size(-1), K=X.size(-1), + M=DY.size(0), + N=DY.size(-1), + K=X.size(-1), # ACC_TYPE: tl.constexpr, ACC_TYPE=tl.float32, - allow_tf32=ALLOW_TF32 + allow_tf32=ALLOW_TF32, ) -@triton.autotune(configs=_config_XtY(), key=['M', 'N', 'K'], ) -@triton.heuristics({ - "NO_K_MASK": lambda args: (args['K'] % args['BLOCK_K']) == 0, - "NO_N_MASK": lambda args: (args['N'] % args['BLOCK_N']) == 0, -}) +@triton.autotune( + configs=_config_XtY(), + key=["M", "N", "K"], +) +@triton.heuristics( + { + "NO_K_MASK": lambda args: (args["K"] % args["BLOCK_K"]) == 0, + "NO_N_MASK": lambda args: (args["N"] % args["BLOCK_N"]) == 0, + } +) @triton.jit def _groupXtY( - DY_ptr, stride_dym, stride_dyk, - X_ptr, stride_xm, stride_xn, - DW_ptr, stride_dwe, stride_dwk, stride_dwn, - Db_ptr, stride_dbe, stride_dbn, + DY_ptr, + stride_dym, + stride_dyk, + X_ptr, + stride_xm, + stride_xn, + DW_ptr, + stride_dwe, + stride_dwk, + stride_dwn, + Db_ptr, + stride_dbe, + stride_dbn, expert_offsets_ptr, - M, K: tl.constexpr, N: tl.constexpr, - BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, + M, + K: tl.constexpr, + N: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, + BLOCK_K: tl.constexpr, ACC_TYPE: tl.constexpr, allow_tf32: tl.constexpr, - NO_K_MASK: tl.constexpr, NO_N_MASK: tl.constexpr + NO_K_MASK: tl.constexpr, + NO_N_MASK: tl.constexpr, ): pid0 = tl.program_id(axis=0) pid1 = tl.program_id(axis=1) @@ -275,7 +385,6 @@ def _groupXtY( start_idx = tl.load(expert_offsets_ptr + E_idx - 1).to(tl.int32) end_idx = tl.load(expert_offsets_ptr + E_idx).to(tl.int32) - if end_idx > start_idx: M_block = tl.max_contiguous(start_idx + tl.arange(0, BLOCK_M), BLOCK_M) @@ -289,48 +398,101 @@ def _groupXtY( M_idxs = M_block xt_blk_ptrs = X_ptr + K_block[:, None] * stride_xn + M_idxs[None, :] * stride_xm - dy_blk_ptrs = DY_ptr + M_idxs[:, None] * stride_dym + N_block[None, :] * stride_dyk + dy_blk_ptrs = ( + DY_ptr + M_idxs[:, None] * stride_dym + N_block[None, :] * stride_dyk + ) if (Db_ptr is not None) and (K_block_id == 0): _xty_and_bias( - E_idx, start_idx, end_idx, + E_idx, + start_idx, + end_idx, M_block, - K_block, K_mask, N_block, N_mask, - dy_blk_ptrs, stride_dym, - xt_blk_ptrs, stride_xm, - DW_ptr, stride_dwe, stride_dwk, stride_dwn, - Db_ptr, stride_dbe, stride_dbn, - BLOCK_M, BLOCK_N, BLOCK_K, ACC_TYPE, - allow_tf32, NO_K_MASK, NO_N_MASK, - compute_bias=True + K_block, + K_mask, + N_block, + N_mask, + dy_blk_ptrs, + stride_dym, + xt_blk_ptrs, + stride_xm, + DW_ptr, + stride_dwe, + stride_dwk, + stride_dwn, + Db_ptr, + stride_dbe, + stride_dbn, + BLOCK_M, + BLOCK_N, + BLOCK_K, + ACC_TYPE, + allow_tf32, + NO_K_MASK, + NO_N_MASK, + compute_bias=True, ) else: _xty_and_bias( - E_idx, start_idx, end_idx, + E_idx, + start_idx, + end_idx, M_block, - K_block, K_mask, N_block, N_mask, - dy_blk_ptrs, stride_dym, - xt_blk_ptrs, stride_xm, - DW_ptr, stride_dwe, stride_dwk, stride_dwn, - Db_ptr, stride_dbe, stride_dbn, - BLOCK_M, BLOCK_N, BLOCK_K, ACC_TYPE, - allow_tf32, NO_K_MASK, NO_N_MASK, - compute_bias=False + K_block, + K_mask, + N_block, + N_mask, + dy_blk_ptrs, + stride_dym, + xt_blk_ptrs, + stride_xm, + DW_ptr, + stride_dwe, + stride_dwk, + stride_dwn, + Db_ptr, + stride_dbe, + stride_dbn, + BLOCK_M, + BLOCK_N, + BLOCK_K, + ACC_TYPE, + allow_tf32, + NO_K_MASK, + NO_N_MASK, + compute_bias=False, ) @triton.jit def _xty_and_bias( - E_idx, start_idx, end_idx, - M_block, - K_block, K_mask, N_block, N_mask, - dy_blk_ptrs, stride_dym, - xt_blk_ptrs, stride_xm, - DW_ptr, stride_dwe, stride_dwk, stride_dwn, - Db_ptr, stride_dbe, stride_dbn, - BLOCK_M, BLOCK_N, BLOCK_K, ACC_TYPE, - allow_tf32, NO_K_MASK, NO_N_MASK, - compute_bias: tl.constexpr - ): + E_idx, + start_idx, + end_idx, + M_block, + K_block, + K_mask, + N_block, + N_mask, + dy_blk_ptrs, + stride_dym, + xt_blk_ptrs, + stride_xm, + DW_ptr, + stride_dwe, + stride_dwk, + stride_dwn, + Db_ptr, + stride_dbe, + stride_dbn, + BLOCK_M, + BLOCK_N, + BLOCK_K, + ACC_TYPE, + allow_tf32, + NO_K_MASK, + NO_N_MASK, + compute_bias: tl.constexpr, +): if compute_bias: db_acc = tl.zeros((BLOCK_N,), dtype=ACC_TYPE) @@ -349,7 +511,7 @@ def _xty_and_bias( dy = tl.load(dy_blk_ptrs, mask=M_mask[:, None]) else: dy = tl.load(dy_blk_ptrs, mask=M_mask[:, None] & N_mask[None, :]) - + acc += tl.dot(xt, dy, out_dtype=ACC_TYPE, allow_tf32=allow_tf32) xt_blk_ptrs += BLOCK_M * stride_xm @@ -358,22 +520,27 @@ def _xty_and_bias( if compute_bias: db_acc += tl.sum(dy, axis=0) - DW_blk_ptrs = DW_ptr + E_idx * stride_dwe + K_block[:, None] * stride_dwk + N_block[None, :] * stride_dwn + DW_blk_ptrs = ( + DW_ptr + + E_idx * stride_dwe + + K_block[:, None] * stride_dwk + + N_block[None, :] * stride_dwn + ) acc = acc.to(DW_blk_ptrs.dtype.element_ty) tl.store(DW_blk_ptrs, acc, mask=K_mask[:, None] & N_mask[None, :]) if compute_bias: - Db_blk_ptrs = Db_ptr + E_idx * stride_dbe + N_block * stride_dbn + Db_blk_ptrs = Db_ptr + E_idx * stride_dbe + N_block * stride_dbn tl.store(Db_blk_ptrs, db_acc, mask=N_mask) - def _config_grouping(): return [ - triton.Config({'BLOCK_N': 256, 'BLOCK_K': 128}, num_stages=4, num_warps=4), + triton.Config({"BLOCK_N": 256, "BLOCK_K": 128}, num_stages=4, num_warps=4), # triton.Config({'BLOCK_N': 128, 'BLOCK_K': 64}, num_stages=4, num_warps=4), # triton.Config({'BLOCK_N': 64, 'BLOCK_K': 32}, num_stages=4, num_warps=4), ] + def group(A, sorted_expert_idxs, coeff=None, fan_out=1, out=None): N = sorted_expert_idxs.size(0) K = A.size(1) @@ -386,42 +553,60 @@ def group(A, sorted_expert_idxs, coeff=None, fan_out=1, out=None): return Y -@torch.library.custom_op("scattermoe::group", mutates_args={"Y"}) +@torch.library.custom_op(add_op_namespace_prefix("group"), mutates_args={"Y"}) def group_compileable( - A: torch.Tensor, - K: int, - N: int, - Y: torch.Tensor, - coeff: torch.Tensor, has_coeff: bool, - fan_out: int, - sorted_expert_idxs: torch.Tensor) -> None: + A: torch.Tensor, + K: int, + N: int, + Y: torch.Tensor, + coeff: torch.Tensor, + has_coeff: bool, + fan_out: int, + sorted_expert_idxs: torch.Tensor, +) -> None: def grid(META): - grid_num = (triton.cdiv(META['N'], META['BLOCK_N']),) + grid_num = (triton.cdiv(META["N"], META["BLOCK_N"]),) return grid_num + _group[grid]( # A_ptr, stride_an, stride_ai, - A, A.stride(0), A.stride(1), has_coeff, coeff, fan_out, + A, + A.stride(0), + A.stride(1), + has_coeff, + coeff, + fan_out, # Y_ptr, stride_yn, stride_yk, - Y, Y.stride(0), Y.stride(1), + Y, + Y.stride(0), + Y.stride(1), # grouped_idx_ptr, sorted_expert_idxs, # N: tl.constexpr, K: tl.constexpr, - N, K + N, + K, ) -@triton.autotune(configs=_config_grouping(), key=['K']) -@triton.heuristics({ - "NO_K_MASK": lambda args: (args['K'] % args['BLOCK_K']) == 0 -}) +@triton.autotune(configs=_config_grouping(), key=["K"]) +@triton.heuristics({"NO_K_MASK": lambda args: (args["K"] % args["BLOCK_K"]) == 0}) @triton.jit def _group( - src_ptr, stride_sn, stride_sk, has_coeff: tl.constexpr, coeff_ptr, FAN_OUT: tl.constexpr, - tgt_ptr, stride_tn, stride_ti, + src_ptr, + stride_sn, + stride_sk, + has_coeff: tl.constexpr, + coeff_ptr, + FAN_OUT: tl.constexpr, + tgt_ptr, + stride_tn, + stride_ti, grouped_idx_ptr, - N, K: tl.constexpr, - BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, - NO_K_MASK: tl.constexpr + N, + K: tl.constexpr, + BLOCK_N: tl.constexpr, + BLOCK_K: tl.constexpr, + NO_K_MASK: tl.constexpr, ): pid = tl.program_id(axis=0) @@ -432,6 +617,9 @@ def _group( N_idx = tl.load(grouped_idx_ptr + N_blk, mask=N_mask, other=0) K_blk = tl.arange(0, BLOCK_K) + src_blk_ptrs = ( + src_ptr + (N_idx // FAN_OUT)[:, None] * stride_sn + K_blk[None, :] * stride_sk + ) src_blk_ptrs = src_ptr + (N_idx // FAN_OUT)[:, None] * stride_sn + K_blk[None, :] * stride_sk tgt_blk_ptrs = tgt_ptr + N_blk[:, None] * stride_tn + K_blk[None, :] * stride_ti diff --git a/scattermoe/torch-ext/scattermoe/parallel_experts.py b/scattermoe/torch-ext/scattermoe/parallel_experts.py index fe67f253..11d7c933 100644 --- a/scattermoe/torch-ext/scattermoe/parallel_experts.py +++ b/scattermoe/torch-ext/scattermoe/parallel_experts.py @@ -1,73 +1,96 @@ +from typing import Optional + import torch import torch.nn as nn + from . import kernels -from typing import Optional +from ._ops import add_op_namespace_prefix + -@torch.library.custom_op("scattermoe::bincount", mutates_args={}) +@torch.library.custom_op(add_op_namespace_prefix("bincount"), mutates_args={}) def compileable_bincount(x: torch.Tensor, minlength: int) -> torch.Tensor: - return x.bincount(minlength=minlength) + return x.bincount(minlength=minlength) + @compileable_bincount.register_fake def _(x: torch.Tensor, minlength: int) -> torch.Tensor: return torch.empty(minlength, dtype=torch.long, device=x.device) + @torch.compile def flatten_sort_count(expert_idxs: torch.Tensor, num_experts: int): with torch.no_grad(): flattened_expert_idxs = expert_idxs.flatten() sorted_expert_idxs, sorted_scattered_idxs = torch.sort(flattened_expert_idxs) - expert_counts = compileable_bincount(flattened_expert_idxs, minlength=num_experts) + expert_counts = compileable_bincount( + flattened_expert_idxs, minlength=num_experts + ) expert_offsets = expert_counts.cumsum(-1) return sorted_expert_idxs, sorted_scattered_idxs, expert_offsets - class ParallelLinear(torch.autograd.Function): @staticmethod def forward( - ctx, - x: torch.Tensor, expert_weights: torch.Tensor, k: int, - sorted_expert_idxs: torch.Tensor, sorted_scattered_idxs: torch.Tensor, + ctx, + x: torch.Tensor, + expert_weights: torch.Tensor, + k: int, + sorted_expert_idxs: torch.Tensor, + sorted_scattered_idxs: torch.Tensor, expert_offsets: torch.Tensor, - expert_biases: Optional[torch.Tensor]=None, - gates: Optional[torch.Tensor]=None, - grouped_in: bool =False, grouped_out: bool=False, + expert_biases: Optional[torch.Tensor] = None, + gates: Optional[torch.Tensor] = None, + grouped_in: bool = False, + grouped_out: bool = False, ): with torch.device(x.device): output = kernels.ops.scatter2scatter( - X=x, W=expert_weights, - b=expert_biases, k=k, + X=x, + W=expert_weights, + b=expert_biases, + k=k, sorted_expert_idxs=sorted_expert_idxs, sorted_scattered_idxs=sorted_scattered_idxs, - x_grouped=grouped_in, y_grouped=grouped_out + x_grouped=grouped_in, + y_grouped=grouped_out, ) if gates is not None: - output_expanded = output.view(gates.size(0), gates.size(1), output.size(-1)) + output_expanded = output.view( + gates.size(0), gates.size(1), output.size(-1) + ) output = (gates.unsqueeze(1) @ output_expanded).squeeze(1) else: output_expanded = None ctx.save_for_backward( - x, expert_weights, + x, + expert_weights, expert_biases, sorted_expert_idxs, sorted_scattered_idxs, expert_offsets, gates, - output_expanded + output_expanded, ) ctx.grouped_in = grouped_in ctx.grouped_out = grouped_out ctx.k = k return output + @staticmethod def backward(ctx, grad_out: torch.Tensor): with torch.device(grad_out.device): - (x, expert_weights, expert_biases, - sorted_expert_idxs, - sorted_scattered_idxs, - expert_offsets, - gates, output_expanded) = ctx.saved_tensors + ( + x, + expert_weights, + expert_biases, + sorted_expert_idxs, + sorted_scattered_idxs, + expert_offsets, + gates, + output_expanded, + ) = ctx.saved_tensors k = ctx.k grouped_in = ctx.grouped_in grouped_out = ctx.grouped_out @@ -79,7 +102,9 @@ def backward(ctx, grad_out: torch.Tensor): d_gates = (output_expanded @ grad_out.unsqueeze(-1)).squeeze(-1) gates_flat = gates.flatten() gate_fan = gates.size(1) - grouped_grad_out = output_expanded.flatten(0, 1) # reuse expanded buffer later + grouped_grad_out = output_expanded.flatten( + 0, 1 + ) # reuse expanded buffer later else: d_gates = None gates_flat = None @@ -89,9 +114,13 @@ def backward(ctx, grad_out: torch.Tensor): if grouped_out: grouped_grad_out = grad_out else: - grouped_grad_out = kernels.ops.group(grad_out, sorted_scattered_idxs, - fan_out=gate_fan, coeff=gates_flat, - out=grouped_grad_out) + grouped_grad_out = kernels.ops.group( + grad_out, + sorted_scattered_idxs, + fan_out=gate_fan, + coeff=gates_flat, + out=grouped_grad_out, + ) if grouped_in: grouped_x = x d_expanded_input = None @@ -100,51 +129,76 @@ def backward(ctx, grad_out: torch.Tensor): d_expanded_input = grouped_x d_weights, d_biases = kernels.ops.group_bwd_W( - DY=grouped_grad_out, X=grouped_x, + DY=grouped_grad_out, + X=grouped_x, expert_offsets=expert_offsets, E=expert_weights.size(0), - has_bias=expert_biases is not None + has_bias=expert_biases is not None, ) - d_expanded_input = kernels.ops.scatter2scatter( - X=grouped_grad_out, x_grouped=True, + X=grouped_grad_out, + x_grouped=True, W=expert_weights.permute(0, 2, 1), sorted_expert_idxs=sorted_expert_idxs, sorted_scattered_idxs=sorted_scattered_idxs, k=1, y_grouped=grouped_in, - out=d_expanded_input # Reuse grouped_x buffer + out=d_expanded_input, # Reuse grouped_x buffer ) if k == 1: d_input = d_expanded_input else: - d_input = d_expanded_input.view(x.size(0), k, d_expanded_input.size(-1)).sum(-2) + d_input = d_expanded_input.view( + x.size(0), k, d_expanded_input.size(-1) + ).sum(-2) # print("backward end.") return ( # x, expert_weights, - d_input, d_weights, + d_input, + d_weights, # k, sorted_expert_idxs, sorted_scattered_idxs, expert_offsets, - None, None, None, None, + None, + None, + None, + None, # bias, gates - d_biases, d_gates, + d_biases, + d_gates, # grouped_in, grouped_out, - None, None + None, + None, ) -def parallel_linear(inputs, expert_weights, k, - sorted_expert_idxs, sorted_scattered_idxs, - expert_offsets, - expert_biases=None, - gates=None, grouped_in=False, grouped_out=False): - results = ParallelLinear.apply(inputs, expert_weights, k, - sorted_expert_idxs, sorted_scattered_idxs, - expert_offsets, - expert_biases, - gates, grouped_in, grouped_out) + +def parallel_linear( + inputs, + expert_weights, + k, + sorted_expert_idxs, + sorted_scattered_idxs, + expert_offsets, + expert_biases=None, + gates=None, + grouped_in=False, + grouped_out=False, +): + results = ParallelLinear.apply( + inputs, + expert_weights, + k, + sorted_expert_idxs, + sorted_scattered_idxs, + expert_offsets, + expert_biases, + gates, + grouped_in, + grouped_out, + ) return results + class ParallelExperts(nn.Module): def __init__(self, num_experts, input_size, output_size, bias=False) -> None: super().__init__() @@ -161,22 +215,37 @@ def __init__(self, num_experts, input_size, output_size, bias=False) -> None: self.reset_parameters() def extra_repr(self): - return 'num_experts={}, input_size={}, output_size={}'.format( - self.num_experts, self.input_size, self.output_size) + return "num_experts={}, input_size={}, output_size={}".format( + self.num_experts, self.input_size, self.output_size + ) def reset_parameters(self) -> None: nn.init.normal_(self.weight, std=0.02) if self.bias is not None: nn.init.zeros_(self.bias) - def forward(self, inputs, k, sorted_expert_idxs, sorted_scattered_idxs, - expert_offsets, - gates=None, grouped_in=False, grouped_out=False): + def forward( + self, + inputs, + k, + sorted_expert_idxs, + sorted_scattered_idxs, + expert_offsets, + gates=None, + grouped_in=False, + grouped_out=False, + ): results = parallel_linear( - inputs, self.weight.permute(0, 2, 1), k, - sorted_expert_idxs, sorted_scattered_idxs, expert_offsets, + inputs, + self.weight.permute(0, 2, 1), + k, + sorted_expert_idxs, + sorted_scattered_idxs, + expert_offsets, expert_biases=self.bias, - gates=gates, grouped_in=grouped_in, grouped_out=grouped_out + gates=gates, + grouped_in=grouped_in, + grouped_out=grouped_out, ) return results From ad6d554acf104b4eb8504e5c7b7ae1403a645414 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Dani=C3=ABl=20de=20Kok?= Date: Sun, 5 Jul 2026 14:46:55 +0000 Subject: [PATCH 3/6] Add some general instructions for agents --- AGENTS.md | 61 +++++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 61 insertions(+) diff --git a/AGENTS.md b/AGENTS.md index 620385e7..4d69abb0 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -1,3 +1,64 @@ +# General instructions that apply to all kernels + +## Imports + +All imports in a kernel (files inside `torch-ext/`) of other +modules within the same kernel must be relative. For example: + +```python +# Incorrect: +from activation.activations import silu_and_mul + +# Correct: +from .activations import silu_and_mul +``` + +## `_ops` module + +The `_ops` module (`_ops.py`) is generated by the build system in +`torch-ext/`. This module contains the following: + +- An `ops` variable that is assigned the Torch ops for the kernel + (`ops = torch.ops.`). +- A function `def add_op_namespace_prefix(op_name: str) -> str`, + which returns `f"::{op_name}"`. + +## Registering Torch operators + +All Torch operators that are registered must have a prefix that is unique +to the kernel and is generated by the build system. Under no circumstances, +the prefix should be fixed. The aforementioned `add_op_namespace_prefix` +is used to prepend this prefix. For instance, consider this registration +with `custom_op`: + +```python +# Incorrect: +@torch.library.custom_op("scattermoe::scatter2scatter", mutates_args={"output"}) + +# Correct: +@torch.library.custom_op(add_op_namespace_prefix("scatter2scatter"), mutates_args={"output"}) +``` + +Similarly, an op must always be prefixed using the kernel-unique prefix: + +```python +# Incorrect: +@torch.library.custom_op("_flash_attn_forward", mutates_args=(), device_types="cuda") + +# Correct: +@torch.library.custom_op(add_op_namespace_prefix("_flash_attn_forward"), mutates_args=(), device_types="cuda") +``` + +The same applies for other registrations of ops, such as fake ops: + +```python +# Incorrect: +@register_fake("moe::single_marlin_gemm_moe") + +# Correct: +@register_fake(add_op_namespace_prefix("single_marlin_gemm_moe")) +``` + # Kernel-specific instructions ## flash-attn3 From f082668d2ba9d30aeb94970d767efb7ced711ccc Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Dani=C3=ABl=20de=20Kok?= Date: Sun, 5 Jul 2026 15:25:54 +0000 Subject: [PATCH 4/6] sonic-moe: remove incorrect namespace prefix wrapper --- sonic-moe/README.md | 6 ++-- sonic-moe/torch-ext/sonic_moe/_ops_compat.py | 10 ------ .../sonic_moe/functional/backward.py | 2 +- .../torch-ext/sonic_moe/functional/forward.py | 2 +- .../functional/triton_kernels/__init__.py | 2 +- .../torch-ext/sonic_moe/quack/_ops_compat.py | 12 +++++-- .../sonic_moe/quack/cross_entropy.py | 6 ++-- .../sonic_moe/quack/gemm_interface.py | 32 +++++++++---------- .../sonic_moe/quack/rms_final_reduce.py | 4 +-- .../torch-ext/sonic_moe/quack/rmsnorm.py | 6 ++-- .../torch-ext/sonic_moe/quack/softmax.py | 6 ++-- sonic-moe/torch-ext/sonic_moe/quack/topk.py | 6 ++-- 12 files changed, 45 insertions(+), 49 deletions(-) delete mode 100644 sonic-moe/torch-ext/sonic_moe/_ops_compat.py diff --git a/sonic-moe/README.md b/sonic-moe/README.md index fc8dabcb..94235a38 100644 --- a/sonic-moe/README.md +++ b/sonic-moe/README.md @@ -82,9 +82,9 @@ concatenated layout directly without a pre-pass permutation. This kernel vendors [QuACK](https://github.com/Dao-AILab/quack) v0.3.11 for CuTe-DSL grouped GEMM infrastructure (Hopper + Blackwell). The vendored copy is -under `torch-ext/sonic_moe/quack/`. All `quack::` torch operator names are -rewritten to the `sonicmoe::quack__*` namespace so builds don't collide with a -user-installed `quack-kernels`. +under `torch-ext/sonic_moe/quack/`. Torch operators are registered through +`add_op_namespace_prefix`, which the build system prefixes with a kernel-unique +namespace, so they cannot collide with a user-installed `quack-kernels`. ## License diff --git a/sonic-moe/torch-ext/sonic_moe/_ops_compat.py b/sonic-moe/torch-ext/sonic_moe/_ops_compat.py deleted file mode 100644 index f4d00b10..00000000 --- a/sonic-moe/torch-ext/sonic_moe/_ops_compat.py +++ /dev/null @@ -1,10 +0,0 @@ -"""Compatibility helpers for op namespacing in source and built layouts.""" - -try: - from ._ops import add_op_namespace_prefix as _generated_add_op_namespace_prefix -except ImportError: - def _generated_add_op_namespace_prefix(name: str) -> str: - return name if "::" in name else f"sonicmoe::{name}" - -def add_op_namespace_prefix(name: str) -> str: - return _generated_add_op_namespace_prefix(name) diff --git a/sonic-moe/torch-ext/sonic_moe/functional/backward.py b/sonic-moe/torch-ext/sonic_moe/functional/backward.py index 01a7ac03..2ad6f478 100644 --- a/sonic-moe/torch-ext/sonic_moe/functional/backward.py +++ b/sonic-moe/torch-ext/sonic_moe/functional/backward.py @@ -11,7 +11,7 @@ import triton.language as tl from ..quack.gemm_interface import gemm, gemm_dgated -from .._ops_compat import add_op_namespace_prefix +from .._ops import add_op_namespace_prefix from ..utils import get_powers_of_2 from .reduction_over_k_gather import token_gather_and_sum_varlen_K_triton diff --git a/sonic-moe/torch-ext/sonic_moe/functional/forward.py b/sonic-moe/torch-ext/sonic_moe/functional/forward.py index 2cc4e780..ec75347f 100644 --- a/sonic-moe/torch-ext/sonic_moe/functional/forward.py +++ b/sonic-moe/torch-ext/sonic_moe/functional/forward.py @@ -11,7 +11,7 @@ from ..quack.cute_dsl_utils import torch2cute_dtype_map from ..quack.gemm_interface import gemm, gemm_gated -from .._ops_compat import add_op_namespace_prefix +from .._ops import add_op_namespace_prefix from .reduction_over_k_gather import token_gather_and_sum_varlen_K_triton from .topk import Softmax_Over_TopK, TopK_Over_Softmax diff --git a/sonic-moe/torch-ext/sonic_moe/functional/triton_kernels/__init__.py b/sonic-moe/torch-ext/sonic_moe/functional/triton_kernels/__init__.py index 90c121ae..f47a57bc 100644 --- a/sonic-moe/torch-ext/sonic_moe/functional/triton_kernels/__init__.py +++ b/sonic-moe/torch-ext/sonic_moe/functional/triton_kernels/__init__.py @@ -4,7 +4,7 @@ import triton import triton.language as tl -from ..._ops_compat import add_op_namespace_prefix +from ..._ops import add_op_namespace_prefix from .bitmatrix import _bitmatrix_metadata_compute_stage1, _bitmatrix_metadata_compute_stage2, _keyed_add diff --git a/sonic-moe/torch-ext/sonic_moe/quack/_ops_compat.py b/sonic-moe/torch-ext/sonic_moe/quack/_ops_compat.py index 9ce465bd..788ee62c 100644 --- a/sonic-moe/torch-ext/sonic_moe/quack/_ops_compat.py +++ b/sonic-moe/torch-ext/sonic_moe/quack/_ops_compat.py @@ -1,4 +1,10 @@ -from .._ops_compat import add_op_namespace_prefix +from .._ops import add_op_namespace_prefix as _add_op_namespace_prefix -def add_quack_op_namespace_prefix(name: str) -> str: - return add_op_namespace_prefix(f"quack__{name}") + + +# For quack we need to prefix the function name because some names +# overlap between quack and sonic-moe itself. Name the function the +# same as the function it is wrapping for the prefix check to be +# happy. +def add_op_namespace_prefix(name: str) -> str: + return _add_op_namespace_prefix(f"quack__{name}") diff --git a/sonic-moe/torch-ext/sonic_moe/quack/cross_entropy.py b/sonic-moe/torch-ext/sonic_moe/quack/cross_entropy.py index d3057bc4..81eb84bf 100644 --- a/sonic-moe/torch-ext/sonic_moe/quack/cross_entropy.py +++ b/sonic-moe/torch-ext/sonic_moe/quack/cross_entropy.py @@ -5,7 +5,7 @@ from typing import Optional, Type, Literal import torch -from ._ops_compat import add_quack_op_namespace_prefix +from ._ops_compat import add_op_namespace_prefix from torch import Tensor import cuda.bindings.driver as cuda @@ -287,7 +287,7 @@ def _compile_cross_entropy_fwd( ) -@torch.library.custom_op(add_quack_op_namespace_prefix("cross_entropy_fwd_out"), mutates_args={"loss", "lse", "dx"}) +@torch.library.custom_op(add_op_namespace_prefix("cross_entropy_fwd_out"), mutates_args={"loss", "lse", "dx"}) def cross_entropy_fwd_out( x: Tensor, target: Tensor, @@ -599,7 +599,7 @@ def _cross_entropy_backward( ) -@torch.library.custom_op(add_quack_op_namespace_prefix("cross_entropy_bwd_out"), mutates_args={"dx"}) +@torch.library.custom_op(add_op_namespace_prefix("cross_entropy_bwd_out"), mutates_args={"dx"}) def cross_entropy_bwd_out( x: torch.Tensor, target: torch.Tensor, diff --git a/sonic-moe/torch-ext/sonic_moe/quack/gemm_interface.py b/sonic-moe/torch-ext/sonic_moe/quack/gemm_interface.py index ce710b47..e8d28e4e 100644 --- a/sonic-moe/torch-ext/sonic_moe/quack/gemm_interface.py +++ b/sonic-moe/torch-ext/sonic_moe/quack/gemm_interface.py @@ -3,7 +3,7 @@ from functools import partial import torch -from ._ops_compat import add_quack_op_namespace_prefix +from ._ops_compat import add_op_namespace_prefix import torch.nn.functional as F from torch import Tensor @@ -455,7 +455,7 @@ def gemm( @torch.library.custom_op( - add_quack_op_namespace_prefix("gemm_out"), + add_op_namespace_prefix("gemm_out"), mutates_args=("out",), device_types="cuda", # We have to split out alpha and alpha_tensor since torch.library requires @@ -653,7 +653,7 @@ def gemm_add( @torch.library.custom_op( - add_quack_op_namespace_prefix("gemm_add_out"), + add_op_namespace_prefix("gemm_add_out"), mutates_args=("out",), device_types="cuda", # We have to split out alpha and alpha_tensor since torch.library requires @@ -833,7 +833,7 @@ def gemm_add_inplace( @torch.library.custom_op( - add_quack_op_namespace_prefix("gemm_add_inplace"), + add_op_namespace_prefix("gemm_add_inplace"), mutates_args=("out",), device_types="cuda", # We have to split out alpha and alpha_tensor since torch.library requires @@ -951,7 +951,7 @@ def gemm_act( @torch.library.custom_op( - add_quack_op_namespace_prefix("gemm_act_out"), + add_op_namespace_prefix("gemm_act_out"), mutates_args=("preact_out", "postact_out"), device_types="cuda", schema="(Tensor A, Tensor B, Tensor(a2!)? preact_out, Tensor(a3!) postact_out, Tensor? C=None, Tensor? bias=None, str? activation=None, Tensor? cu_seqlens_m=None, Tensor? A_idx=None, bool dynamic_scheduler=False, bool tuned=True) -> ()", @@ -1087,7 +1087,7 @@ def gemm_dact( @torch.library.custom_op( - add_quack_op_namespace_prefix("gemm_dact_out"), + add_op_namespace_prefix("gemm_dact_out"), mutates_args=("dx_out", "postact_out"), device_types="cuda", schema="(Tensor A, Tensor B, Tensor PreAct, Tensor(a3!) dx_out, Tensor(a4!) postact_out, str? activation=None, Tensor? cu_seqlens_m=None, Tensor? A_idx=None, bool dynamic_scheduler=True, bool tuned=True) -> ()", @@ -1153,7 +1153,7 @@ def gemm_dact_ref( @torch.library.custom_op( - add_quack_op_namespace_prefix("gemm_symmetric_out"), + add_op_namespace_prefix("gemm_symmetric_out"), mutates_args=("out",), device_types="cuda", schema="(Tensor A, Tensor B, Tensor(a2!) out, Tensor? C=None, bool dynamic_scheduler=False, float alpha=1.0, float beta=1.0) -> ()", @@ -1423,7 +1423,7 @@ def gemm_dgated_tuned( @torch.library.custom_op( - add_quack_op_namespace_prefix("gemm_gated_out"), + add_op_namespace_prefix("gemm_gated_out"), mutates_args=("preact_out", "postact_out"), device_types="cuda", schema="(Tensor A, Tensor B, Tensor(a2!)? preact_out, Tensor(a3!) postact_out, Tensor? C=None, Tensor? bias=None, str activation='swiglu', Tensor? cu_seqlens_m=None, Tensor? A_idx=None, bool dynamic_scheduler=False, bool tuned=True, str? concat_layout=None) -> ()", @@ -1460,7 +1460,7 @@ def gemm_gated_out( @torch.library.custom_op( - add_quack_op_namespace_prefix("gemm_dgated_out"), + add_op_namespace_prefix("gemm_dgated_out"), mutates_args=("dx_out", "postact_out"), device_types="cuda", schema="(Tensor A, Tensor B, Tensor PreAct, Tensor(a!) dx_out, Tensor(b!) postact_out, Tensor? colvec_scale=None, str activation='swiglu', bool colvec_reduce=False, Tensor? cu_seqlens_m=None, Tensor? A_idx=None, bool dynamic_scheduler=True, bool tuned=True) -> Tensor", @@ -1499,7 +1499,7 @@ def gemm_dgated_out( return result -@torch.library.register_fake(add_quack_op_namespace_prefix("gemm_dgated_out")) +@torch.library.register_fake(add_op_namespace_prefix("gemm_dgated_out")) def gemm_dgated_out_fake( A: Tensor, B: Tensor, @@ -1763,7 +1763,7 @@ def _gemm_rms_tuned( @torch.library.custom_op( - add_quack_op_namespace_prefix("gemm_rms_out"), + add_op_namespace_prefix("gemm_rms_out"), mutates_args=("out",), device_types="cuda", schema="(Tensor A, Tensor B, Tensor(a!) out, Tensor? C=None, Tensor? norm_weight=None, float eps=1e-6, bool dynamic_scheduler=False, bool tuned=True) -> Tensor", @@ -1794,7 +1794,7 @@ def _gemm_rms_out( ) -@torch.library.register_fake(add_quack_op_namespace_prefix("gemm_rms_out")) +@torch.library.register_fake(add_op_namespace_prefix("gemm_rms_out")) def _gemm_rms_out_fake( A: Tensor, B: Tensor, @@ -1999,7 +1999,7 @@ def gemm_norm_gated_tuned( @torch.library.custom_op( - add_quack_op_namespace_prefix("gemm_norm_act_out"), + add_op_namespace_prefix("gemm_norm_act_out"), mutates_args=("preact_out", "postact_out"), device_types="cuda", schema="(Tensor A, Tensor B, Tensor(a2!)? preact_out, Tensor(a3!) postact_out, Tensor? C=None, Tensor? rstd=None, str? activation=None, bool dynamic_scheduler=False, bool tuned=True) -> ()", @@ -2019,7 +2019,7 @@ def gemm_norm_act_out( fn(A, B, preact_out, postact_out, C, rstd, activation, dynamic_scheduler) -@torch.library.register_fake(add_quack_op_namespace_prefix("gemm_norm_act_out")) +@torch.library.register_fake(add_op_namespace_prefix("gemm_norm_act_out")) def _gemm_norm_act_out_fake( A, B, @@ -2035,7 +2035,7 @@ def _gemm_norm_act_out_fake( @torch.library.custom_op( - add_quack_op_namespace_prefix("gemm_norm_gated_out"), + add_op_namespace_prefix("gemm_norm_gated_out"), mutates_args=("preact_out", "postact_out"), device_types="cuda", schema="(Tensor A, Tensor B, Tensor(a2!)? preact_out, Tensor(a3!) postact_out, Tensor? C=None, Tensor? rstd=None, str activation='swiglu', bool dynamic_scheduler=False, bool tuned=True) -> ()", @@ -2055,7 +2055,7 @@ def gemm_norm_gated_out( fn(A, B, preact_out, postact_out, C, rstd, activation, dynamic_scheduler) -@torch.library.register_fake(add_quack_op_namespace_prefix("gemm_norm_gated_out")) +@torch.library.register_fake(add_op_namespace_prefix("gemm_norm_gated_out")) def _gemm_norm_gated_out_fake( A, B, diff --git a/sonic-moe/torch-ext/sonic_moe/quack/rms_final_reduce.py b/sonic-moe/torch-ext/sonic_moe/quack/rms_final_reduce.py index 1b65d95e..d5dd7641 100644 --- a/sonic-moe/torch-ext/sonic_moe/quack/rms_final_reduce.py +++ b/sonic-moe/torch-ext/sonic_moe/quack/rms_final_reduce.py @@ -13,7 +13,7 @@ from cutlass import Float32, const_expr import torch -from ._ops_compat import add_quack_op_namespace_prefix +from ._ops_compat import add_op_namespace_prefix from torch import Tensor from . import copy_utils as copy_utils @@ -136,7 +136,7 @@ def _compile_rms_final_reduce(dtype, N): @torch.library.custom_op( - add_quack_op_namespace_prefix("rms_final_reduce_out"), + add_op_namespace_prefix("rms_final_reduce_out"), mutates_args=("rstd",), device_types="cuda", ) diff --git a/sonic-moe/torch-ext/sonic_moe/quack/rmsnorm.py b/sonic-moe/torch-ext/sonic_moe/quack/rmsnorm.py index b3bcfc36..82b7d692 100644 --- a/sonic-moe/torch-ext/sonic_moe/quack/rmsnorm.py +++ b/sonic-moe/torch-ext/sonic_moe/quack/rmsnorm.py @@ -11,7 +11,7 @@ from cutlass import Float32, Int32, const_expr import torch -from ._ops_compat import add_quack_op_namespace_prefix +from ._ops_compat import add_op_namespace_prefix from torch import Tensor from . import utils as utils @@ -318,7 +318,7 @@ def kernel( @torch.library.custom_op( - add_quack_op_namespace_prefix("_rmsnorm_fwd"), + add_op_namespace_prefix("_rmsnorm_fwd"), mutates_args=("out", "rstd", "mean", "residual_out"), device_types="cuda", # We need to specify the schema manually since we're mutating an optional tensor @@ -921,7 +921,7 @@ def _get_sm_count(N: int, device: torch.device) -> int: @torch.library.custom_op( - add_quack_op_namespace_prefix("_rmsnorm_bwd"), + add_op_namespace_prefix("_rmsnorm_bwd"), mutates_args={"dx", "dw_partial", "db_partial", "dresidual"}, device_types="cuda", # We need to specify the schema manually since we're mutating an optional tensor diff --git a/sonic-moe/torch-ext/sonic_moe/quack/softmax.py b/sonic-moe/torch-ext/sonic_moe/quack/softmax.py index c32e6e4c..aaf14f51 100644 --- a/sonic-moe/torch-ext/sonic_moe/quack/softmax.py +++ b/sonic-moe/torch-ext/sonic_moe/quack/softmax.py @@ -6,7 +6,7 @@ import torch -from ._ops_compat import add_quack_op_namespace_prefix +from ._ops_compat import add_op_namespace_prefix import cuda.bindings.driver as cuda import cutlass @@ -191,7 +191,7 @@ def _compile_softmax_fwd(dtype, out_dtype, N): ) -@torch.library.custom_op(add_quack_op_namespace_prefix("_softmax_fwd"), mutates_args={"out"}) +@torch.library.custom_op(add_op_namespace_prefix("_softmax_fwd"), mutates_args={"out"}) def _softmax_fwd(x: torch.Tensor, out: torch.Tensor) -> None: """Softmax forward pass. Args: @@ -388,7 +388,7 @@ def _compile_softmax_backward(dtype, y_dtype, dx_dtype, N): ) -@torch.library.custom_op(add_quack_op_namespace_prefix("_softmax_backward"), mutates_args={"dx"}) +@torch.library.custom_op(add_op_namespace_prefix("_softmax_backward"), mutates_args={"dx"}) def _softmax_backward(dy: torch.Tensor, y: torch.Tensor, dx: torch.Tensor) -> None: """Softmax backward pass. Args: diff --git a/sonic-moe/torch-ext/sonic_moe/quack/topk.py b/sonic-moe/torch-ext/sonic_moe/quack/topk.py index 7cd089b3..30d45a0f 100644 --- a/sonic-moe/torch-ext/sonic_moe/quack/topk.py +++ b/sonic-moe/torch-ext/sonic_moe/quack/topk.py @@ -6,7 +6,7 @@ import torch -from ._ops_compat import add_quack_op_namespace_prefix +from ._ops_compat import add_op_namespace_prefix import cuda.bindings.driver as cuda import cutlass @@ -216,7 +216,7 @@ def kernel( cute.autovec_copy(topk_indices[None, i], mIndices_store[None, col]) -@torch.library.custom_op(add_quack_op_namespace_prefix("_topk_fwd"), mutates_args={"values", "indices"}) +@torch.library.custom_op(add_op_namespace_prefix("_topk_fwd"), mutates_args={"values", "indices"}) def _topk_fwd( x: torch.Tensor, k: int, softmax: bool, values: torch.Tensor, indices: torch.Tensor ) -> None: @@ -457,7 +457,7 @@ def kernel( copy_dx(tXrdX, tXgdX) -@torch.library.custom_op(add_quack_op_namespace_prefix("_topk_bwd"), mutates_args={"dx"}) +@torch.library.custom_op(add_op_namespace_prefix("_topk_bwd"), mutates_args={"dx"}) def _topk_bwd( dvalues: torch.Tensor, values: Optional[torch.Tensor], From 6bf44a249c003886bd3b4027c67a0f1375b6c9b2 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Dani=C3=ABl=20de=20Kok?= Date: Sun, 5 Jul 2026 15:30:50 +0000 Subject: [PATCH 5/6] flash-attn-ops: remove add_op_namespace_prefix wrapper See Slack for discussion. These wrappers can hide import errors (leading to incorrect/conflicting op registrations) and impede static analysis. --- .../torch-ext/flash_attn_ops/_ops_compat.py | 18 -- .../torch-ext/flash_attn_ops/layer_norm.py | 209 ++++++++++++------ 2 files changed, 139 insertions(+), 88 deletions(-) delete mode 100644 flash-attn-ops/torch-ext/flash_attn_ops/_ops_compat.py diff --git a/flash-attn-ops/torch-ext/flash_attn_ops/_ops_compat.py b/flash-attn-ops/torch-ext/flash_attn_ops/_ops_compat.py deleted file mode 100644 index 08984f62..00000000 --- a/flash-attn-ops/torch-ext/flash_attn_ops/_ops_compat.py +++ /dev/null @@ -1,18 +0,0 @@ -"""Compatibility helpers for op namespacing in source and built layouts. - -In the built (Hub) layout, kernel-builder generates an `_ops` module that -exposes `add_op_namespace_prefix`, which prefixes op names with a unique, -build-hashed namespace so custom ops never collide across kernels/versions. When -running directly from source there is no generated `_ops`, so we fall back to a -fixed namespace. -""" - -try: - from ._ops import add_op_namespace_prefix as _generated_add_op_namespace_prefix -except ImportError: - def _generated_add_op_namespace_prefix(name: str) -> str: - return name if "::" in name else f"flash_attn_ops::{name}" - - -def add_op_namespace_prefix(name: str) -> str: - return _generated_add_op_namespace_prefix(name) diff --git a/flash-attn-ops/torch-ext/flash_attn_ops/layer_norm.py b/flash-attn-ops/torch-ext/flash_attn_ops/layer_norm.py index 660f1021..f20803f2 100644 --- a/flash-attn-ops/torch-ext/flash_attn_ops/layer_norm.py +++ b/flash-attn-ops/torch-ext/flash_attn_ops/layer_norm.py @@ -7,18 +7,17 @@ # The models we train have hidden dim up to 8k anyway (e.g. Llama 70B), so this is fine. import math -from typing import Optional, List +from typing import List, Optional import torch import torch.nn.functional as F -from torch import Tensor - import triton import triton.language as tl +from torch import Tensor -from .utils.torch import custom_fwd, custom_bwd +from ._ops import add_op_namespace_prefix from .utils.library import triton_op -from ._ops_compat import add_op_namespace_prefix +from .utils.torch import custom_bwd, custom_fwd def maybe_contiguous_lastdim(x): @@ -40,9 +39,17 @@ def triton_autotune_configs(): warp_size = 32 if torch.cuda.is_available(): warp_size = getattr(torch.cuda.get_device_properties(torch.cuda.current_device()), "warp_size", 32) + warp_size = getattr( + torch.cuda.get_device_properties(torch.cuda.current_device()), + "warp_size", + 32, + ) # Autotune for warp counts which are powers of 2 and do not exceed thread per block limit - return [triton.Config({}, num_warps=warp_count) for warp_count in [1, 2, 4, 8, 16, 32] - if warp_count * warp_size <= max_threads_per_block] + return [ + triton.Config({}, num_warps=warp_count) + for warp_count in [1, 2, 4, 8, 16, 32] + if warp_count * warp_size <= max_threads_per_block + ] # return [triton.Config({}, num_warps=8)] @@ -94,9 +101,9 @@ def layer_norm_ref( x = x + x1 if residual is not None: x = (x + residual).to(x.dtype) - out = F.layer_norm(x.to(weight.dtype), x.shape[-1:], weight=weight, bias=bias, eps=eps).to( - dtype - ) + out = F.layer_norm( + x.to(weight.dtype), x.shape[-1:], weight=weight, bias=bias, eps=eps + ).to(dtype) if weight1 is None: return out if not prenorm else (out, x) else: @@ -155,19 +162,30 @@ def rms_norm_ref( if residual is not None: x = (x + residual).to(x.dtype) rstd = 1 / torch.sqrt((x.square()).mean(dim=-1, keepdim=True) + eps) - out = ((x * rstd * weight) + bias if bias is not None else (x * rstd * weight)).to(dtype) + out = ((x * rstd * weight) + bias if bias is not None else (x * rstd * weight)).to( + dtype + ) if weight1 is None: return out if not prenorm else (out, x) else: - out1 = ((x * rstd * weight1) + bias1 if bias1 is not None else (x * rstd * weight1)).to( - dtype - ) + out1 = ( + (x * rstd * weight1) + bias1 if bias1 is not None else (x * rstd * weight1) + ).to(dtype) return (out, out1) if not prenorm else (out, out1, x) @triton.autotune( configs=triton_autotune_configs(), - key=["N", "HAS_RESIDUAL", "STORE_RESIDUAL_OUT", "IS_RMS_NORM", "HAS_BIAS", "HAS_X1", "HAS_W1", "HAS_B1"], + key=[ + "N", + "HAS_RESIDUAL", + "STORE_RESIDUAL_OUT", + "IS_RMS_NORM", + "HAS_BIAS", + "HAS_X1", + "HAS_W1", + "HAS_B1", + ], ) # torch compile doesn't like triton.heuristics, so we set these manually when calling the kernel # @triton.heuristics({"HAS_BIAS": lambda args: args["B"] is not None}) @@ -237,7 +255,9 @@ def _layer_norm_fwd_1pass_kernel( if HAS_DROPOUT: # Compute dropout mask # 7 rounds is good enough, and reduces register pressure - keep_mask = tl.rand(tl.load(SEEDS + row).to(tl.uint32), cols, n_rounds=7) > dropout_p + keep_mask = ( + tl.rand(tl.load(SEEDS + row).to(tl.uint32), cols, n_rounds=7) > dropout_p + ) x = tl.where(keep_mask, x / (1.0 - dropout_p), 0.0) if STORE_DROPOUT_MASK: tl.store(DROPOUT_MASK + row * N + cols, keep_mask, mask=cols < N) @@ -250,7 +270,8 @@ def _layer_norm_fwd_1pass_kernel( # Compute dropout mask # 7 rounds is good enough, and reduces register pressure keep_mask = ( - tl.rand(tl.load(SEEDS + M + row).to(tl.uint32), cols, n_rounds=7) > dropout_p + tl.rand(tl.load(SEEDS + M + row).to(tl.uint32), cols, n_rounds=7) + > dropout_p ) x1 = tl.where(keep_mask, x1 / (1.0 - dropout_p), 0.0) if STORE_DROPOUT_MASK: @@ -309,7 +330,7 @@ def _layer_norm_fwd( is_rms_norm: bool = False, return_dropout_mask: bool = False, out: Optional[Tensor] = None, - residual_out: Optional[Tensor] = None + residual_out: Optional[Tensor] = None, ) -> (Tensor, Tensor, Tensor, Tensor, Tensor, Tensor, Tensor, Tensor): # Need to wrap to handle the case where residual_out is a alias of x, which makes torch.library # and torch.compile unhappy. Also allocate memory for out and residual_out if they are None @@ -355,8 +376,11 @@ def _layer_norm_fwd( # [2025-04-28] torch.library.triton_op ignores the schema argument, but here we need the schema # since we're returning a tuple of tensors -@triton_op(add_op_namespace_prefix("layer_norm_fwd_impl"), mutates_args={"out", "residual_out"}, - schema="(Tensor x, Tensor weight, Tensor bias, float eps, Tensor(a!) out, Tensor? residual, Tensor? x1, Tensor? weight1, Tensor? bias1, float dropout_p, Tensor? rowscale, bool zero_centered_weight, bool is_rms_norm, bool return_dropout_mask, Tensor(a!)? residual_out) -> (Tensor y1, Tensor mean, Tensor rstd, Tensor seeds, Tensor dropout_mask, Tensor dropout_mask1)") +@triton_op( + add_op_namespace_prefix("layer_norm_fwd_impl"), + mutates_args={"out", "residual_out"}, + schema="(Tensor x, Tensor weight, Tensor bias, float eps, Tensor(a!) out, Tensor? residual, Tensor? x1, Tensor? weight1, Tensor? bias1, float dropout_p, Tensor? rowscale, bool zero_centered_weight, bool is_rms_norm, bool return_dropout_mask, Tensor(a!)? residual_out) -> (Tensor y1, Tensor mean, Tensor rstd, Tensor seeds, Tensor dropout_mask, Tensor dropout_mask1)", +) def _layer_norm_fwd_impl( x: Tensor, weight: Tensor, @@ -372,7 +396,7 @@ def _layer_norm_fwd_impl( zero_centered_weight: bool = False, is_rms_norm: bool = False, return_dropout_mask: bool = False, - residual_out: Optional[Tensor] = None + residual_out: Optional[Tensor] = None, ) -> (Tensor, Tensor, Tensor, Tensor, Tensor, Tensor): M, N = x.shape assert x.stride(-1) == 1 @@ -407,7 +431,11 @@ def _layer_norm_fwd_impl( assert y1.stride(-1) == 1 else: y1 = None - mean = torch.empty((M,), dtype=torch.float32, device=x.device) if not is_rms_norm else None + mean = ( + torch.empty((M,), dtype=torch.float32, device=x.device) + if not is_rms_norm + else None + ) rstd = torch.empty((M,), dtype=torch.float32, device=x.device) if dropout_p > 0.0: seeds = torch.randint( @@ -475,7 +503,14 @@ def _layer_norm_fwd_impl( @triton.autotune( configs=triton_autotune_configs(), - key=["N", "HAS_DRESIDUAL", "STORE_DRESIDUAL", "IS_RMS_NORM", "HAS_BIAS", "HAS_DROPOUT"], + key=[ + "N", + "HAS_DRESIDUAL", + "STORE_DRESIDUAL", + "IS_RMS_NORM", + "HAS_BIAS", + "HAS_DROPOUT", + ], ) # torch compile doesn't like triton.heuristics, so we set these manually when calling the kernel # @triton.heuristics({"HAS_BIAS": lambda args: args["B"] is not None}) @@ -609,14 +644,18 @@ def _layer_norm_bwd_kernel( if HAS_DX1: if HAS_DROPOUT: keep_mask = ( - tl.rand(tl.load(SEEDS + M + row).to(tl.uint32), cols, n_rounds=7) > dropout_p + tl.rand(tl.load(SEEDS + M + row).to(tl.uint32), cols, n_rounds=7) + > dropout_p ) dx1 = tl.where(keep_mask, dx / (1.0 - dropout_p), 0.0) else: dx1 = dx tl.store(DX1 + cols, dx1, mask=mask) if HAS_DROPOUT: - keep_mask = tl.rand(tl.load(SEEDS + row).to(tl.uint32), cols, n_rounds=7) > dropout_p + keep_mask = ( + tl.rand(tl.load(SEEDS + row).to(tl.uint32), cols, n_rounds=7) + > dropout_p + ) dx = tl.where(keep_mask, dx / (1.0 - dropout_p), 0.0) if HAS_ROWSCALE: rowscale = tl.load(ROWSCALE + row).to(tl.float32) @@ -699,11 +738,12 @@ def _layer_norm_bwd( return dx, dw, db, dresidual_in, dx1, dw1, db1, y - -@triton_op(add_op_namespace_prefix("layer_norm_bwd_impl"), mutates_args={}, - schema="(Tensor dy, Tensor x, Tensor weight, Tensor bias, float eps, Tensor mean, Tensor rstd, Tensor? dresidual, Tensor? dy1, Tensor? weight1, Tensor? bias1, Tensor? seeds, float dropout_p, Tensor? rowscale, bool has_residual, bool has_x1, bool zero_centered_weight, bool is_rms_norm, ScalarType? x_dtype, bool recompute_output) -> (Tensor dx, Tensor dw, Tensor db, Tensor dresidual_in, Tensor dx1, Tensor dw1, Tensor db1, Tensor y)", - allow_decomposition=False, # Don't let torch.compile trace inside - ) +@triton_op( + add_op_namespace_prefix("layer_norm_bwd_impl"), + mutates_args={}, + schema="(Tensor dy, Tensor x, Tensor weight, Tensor bias, float eps, Tensor mean, Tensor rstd, Tensor? dresidual, Tensor? dy1, Tensor? weight1, Tensor? bias1, Tensor? seeds, float dropout_p, Tensor? rowscale, bool has_residual, bool has_x1, bool zero_centered_weight, bool is_rms_norm, ScalarType? x_dtype, bool recompute_output) -> (Tensor dx, Tensor dw, Tensor db, Tensor dresidual_in, Tensor dx1, Tensor dw1, Tensor db1, Tensor y)", + allow_decomposition=False, # Don't let torch.compile trace inside +) def _layer_norm_bwd_impl( dy: Tensor, x: Tensor, @@ -770,9 +810,15 @@ def _layer_norm_bwd_impl( else None ) dx1 = torch.empty_like(dx) if (has_x1 and dropout_p > 0.0) else None - y = torch.empty(M, N, dtype=dy.dtype, device=dy.device) if recompute_output else None + y = ( + torch.empty(M, N, dtype=dy.dtype, device=dy.device) + if recompute_output + else None + ) if recompute_output: - assert weight1 is None, "recompute_output is not supported with parallel LayerNorm" + assert weight1 is None, ( + "recompute_output is not supported with parallel LayerNorm" + ) # Less than 64KB per feature: enqueue fused kernel MAX_FUSED_SIZE = 65536 // x.element_size() @@ -849,7 +895,6 @@ def _layer_norm_bwd_impl( class LayerNormFn(torch.autograd.Function): - @staticmethod def forward( ctx, @@ -870,14 +915,16 @@ def forward( return_dropout_mask=False, out_dtype=None, out=None, - residual_out=None + residual_out=None, ): x_shape_og = x.shape # reshape input data into 2D tensor x = maybe_contiguous_lastdim(x.reshape(-1, x.shape[-1])) if residual is not None: assert residual.shape == x_shape_og - residual = maybe_contiguous_lastdim(residual.reshape(-1, residual.shape[-1])) + residual = maybe_contiguous_lastdim( + residual.reshape(-1, residual.shape[-1]) + ) if x1 is not None: assert x1.shape == x_shape_og assert rowscale is None, "rowscale is not supported with parallel LayerNorm" @@ -897,24 +944,26 @@ def forward( out = out.reshape(-1, out.shape[-1]) if residual_out is not None: residual_out = residual_out.reshape(-1, residual_out.shape[-1]) - y, y1, mean, rstd, residual_out, seeds, dropout_mask, dropout_mask1 = _layer_norm_fwd( - x, - weight, - bias, - eps, - residual, - x1, - weight1, - bias1, - dropout_p=dropout_p, - rowscale=rowscale, - out_dtype=out_dtype, - residual_dtype=residual_dtype, - zero_centered_weight=zero_centered_weight, - is_rms_norm=is_rms_norm, - return_dropout_mask=return_dropout_mask, - out=out, - residual_out=residual_out, + y, y1, mean, rstd, residual_out, seeds, dropout_mask, dropout_mask1 = ( + _layer_norm_fwd( + x, + weight, + bias, + eps, + residual, + x1, + weight1, + bias1, + dropout_p=dropout_p, + rowscale=rowscale, + out_dtype=out_dtype, + residual_dtype=residual_dtype, + zero_centered_weight=zero_centered_weight, + is_rms_norm=is_rms_norm, + return_dropout_mask=return_dropout_mask, + out=out, + residual_out=residual_out, + ) ) ctx.save_for_backward( residual_out, weight, bias, weight1, bias1, rowscale, seeds, mean, rstd @@ -930,9 +979,15 @@ def forward( ctx.zero_centered_weight = zero_centered_weight y = y.reshape(x_shape_og) y1 = y1.reshape(x_shape_og) if y1 is not None else None - residual_out = residual_out.reshape(x_shape_og) if residual_out is not None else None - dropout_mask = dropout_mask.reshape(x_shape_og) if dropout_mask is not None else None - dropout_mask1 = dropout_mask1.reshape(x_shape_og) if dropout_mask1 is not None else None + residual_out = ( + residual_out.reshape(x_shape_og) if residual_out is not None else None + ) + dropout_mask = ( + dropout_mask.reshape(x_shape_og) if dropout_mask is not None else None + ) + dropout_mask1 = ( + dropout_mask1.reshape(x_shape_og) if dropout_mask1 is not None else None + ) if not return_dropout_mask: if weight1 is None: return y if not prenorm else (y, residual_out) @@ -1030,7 +1085,7 @@ def layer_norm_fn( return_dropout_mask=False, out_dtype=None, out=None, - residual_out=None + residual_out=None, ): return LayerNormFn.apply( x, @@ -1050,7 +1105,7 @@ def layer_norm_fn( return_dropout_mask, out_dtype, out, - residual_out + residual_out, ) @@ -1071,7 +1126,7 @@ def rms_norm_fn( return_dropout_mask=False, out_dtype=None, out=None, - residual_out=None + residual_out=None, ): return LayerNormFn.apply( x, @@ -1091,14 +1146,20 @@ def rms_norm_fn( return_dropout_mask, out_dtype, out, - residual_out + residual_out, ) class RMSNorm(torch.nn.Module): - - def __init__(self, hidden_size, eps=1e-5, dropout_p=0.0, zero_centered_weight=False, - device=None, dtype=None): + def __init__( + self, + hidden_size, + eps=1e-5, + dropout_p=0.0, + zero_centered_weight=False, + device=None, + dtype=None, + ): factory_kwargs = {"device": device, "dtype": dtype} super().__init__() self.eps = eps @@ -1132,7 +1193,6 @@ def forward(self, x, residual=None, prenorm=False, residual_in_fp32=False): class LayerNormLinearFn(torch.autograd.Function): - @staticmethod @custom_fwd def forward( @@ -1153,7 +1213,9 @@ def forward( x = maybe_contiguous_lastdim(x.reshape(-1, x.shape[-1])) if residual is not None: assert residual.shape == x_shape_og - residual = maybe_contiguous_lastdim(residual.reshape(-1, residual.shape[-1])) + residual = maybe_contiguous_lastdim( + residual.reshape(-1, residual.shape[-1]) + ) norm_weight = norm_weight.contiguous() norm_bias = maybe_contiguous(norm_bias) residual_dtype = ( @@ -1167,17 +1229,23 @@ def forward( norm_bias, eps, residual, - out_dtype=None if not torch.is_autocast_enabled() else torch.get_autocast_dtype("cuda"), + out_dtype=None + if not torch.is_autocast_enabled() + else torch.get_autocast_dtype("cuda"), residual_dtype=residual_dtype, is_rms_norm=is_rms_norm, ) y = y.reshape(x_shape_og) - dtype = torch.get_autocast_dtype("cuda") if torch.is_autocast_enabled() else y.dtype + dtype = ( + torch.get_autocast_dtype("cuda") if torch.is_autocast_enabled() else y.dtype + ) linear_weight = linear_weight.to(dtype) linear_bias = linear_bias.to(dtype) if linear_bias is not None else None out = F.linear(y.to(linear_weight.dtype), linear_weight, linear_bias) # We don't store y, will be recomputed in the backward pass to save memory - ctx.save_for_backward(residual_out, norm_weight, norm_bias, linear_weight, mean, rstd) + ctx.save_for_backward( + residual_out, norm_weight, norm_bias, linear_weight, mean, rstd + ) ctx.x_shape_og = x_shape_og ctx.eps = eps ctx.is_rms_norm = is_rms_norm @@ -1198,8 +1266,9 @@ def backward(ctx, dout, *args): assert dy.shape == x.shape if ctx.prenorm: dresidual = args[0] - dresidual = maybe_contiguous_lastdim(dresidual.reshape(-1, dresidual.shape[-1])) - assert dresidual.shape == x.shape + dresidual = maybe_contiguous_lastdim( + dresidual.reshape(-1, dresidual.shape[-1]) + ) else: dresidual = None dx, dnorm_weight, dnorm_bias, dresidual_in, _, _, _, y = _layer_norm_bwd( From 852690778c761658eff9a407ddacfd18698b3281 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Dani=C3=ABl=20de=20Kok?= Date: Sun, 5 Jul 2026 15:33:49 +0000 Subject: [PATCH 6/6] vllm-moe: disable CUDA 13.x until we sync with upstream --- vllm-moe/build.toml | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/vllm-moe/build.toml b/vllm-moe/build.toml index c610cbad..775b2a6d 100644 --- a/vllm-moe/build.toml +++ b/vllm-moe/build.toml @@ -8,6 +8,11 @@ backends = ["cuda"] [general.hub] repo-id = "kernels-community/vllm-moe" +[general.cuda] +# Not compatible with cub from CUDA 13.0, should work again after sync of +# kernel with upstream. +maxver = "12.9" + [torch] include = ["."] pyext = [