Kernels:
Trusted publisher
Uploaded using `kernel-builder`.
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- build/torch212-cxx11-xpu20253-x86_64-linux/{_megablocks_xpu_c4d0dc5.abi3.so → _megablocks_xpu_addf474.abi3.so} +1 -1
- build/torch212-cxx11-xpu20253-x86_64-linux/_ops.py +3 -3
- build/torch212-cxx11-xpu20253-x86_64-linux/megablocks/__init__.py +0 -26
- build/torch212-cxx11-xpu20253-x86_64-linux/metadata.json +6 -6
- build/torch212-cxx11-xpu20253-x86_64-linux/metadata.json.sigstore +1 -1
- build/torch213-cxx11-xpu20260-x86_64-linux/{_megablocks_xpu_c4d0dc5.abi3.so → _megablocks_xpu_addf474.abi3.so} +2 -2
- build/torch213-cxx11-xpu20260-x86_64-linux/_ops.py +3 -3
- build/torch213-cxx11-xpu20260-x86_64-linux/megablocks/__init__.py +0 -26
- build/torch213-cxx11-xpu20260-x86_64-linux/metadata.json +6 -6
- build/torch213-cxx11-xpu20260-x86_64-linux/metadata.json.sigstore +1 -1
- build/torch214-cxx11-xpu20261-x86_64-linux/__init__.py +205 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/__init__.py +0 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/_triton_kernels/__init__.py +0 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/_triton_kernels/gmm.py +574 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/adapter.py +53 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/configs.py +5 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/gmm.py +567 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/utils/__init__.py +0 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/utils/_triton/__init__.py +0 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/utils/_triton/arch_info.py +46 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/utils/_triton/pid_preprocessing.py +100 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/utils/gmm_common.py +752 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/utils/logger.py +47 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_layers/__init__.py +10 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_layers/activation_fn.py +33 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_layers/all_to_all.py +54 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_layers/arguments.py +101 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_layers/common.py +26 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_layers/dmlp_registry.py +42 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_layers/dmoe.py +337 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_layers/gelu.py +52 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_layers/glu.py +244 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_layers/memory_test.py +103 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_layers/mlp.py +587 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_layers/moe.py +507 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_layers/mpu.py +94 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_layers/router.py +116 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_layers/sharedexpert_registry.py +32 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_megablocks_xpu_addf474.abi3.so +3 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_ops.py +9 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/_version.py +6 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/backend/__init__.py +2 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/backend/kernels.py +557 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/benchmark_util.py +35 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/cpu_fused_moe.py +311 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/cpu_moe_cpp.py +265 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/grouped_gemm/__init__.py +2 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/grouped_gemm/backend.py +33 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/grouped_gemm/ops.py +33 -0
- build/torch214-cxx11-xpu20261-x86_64-linux/grouped_gemm_util.py +31 -0
build/torch212-cxx11-xpu20253-x86_64-linux/{_megablocks_xpu_c4d0dc5.abi3.so → _megablocks_xpu_addf474.abi3.so}
RENAMED
|
@@ -1,3 +1,3 @@
|
|
| 1 |
version https://git-lfs.github.com/spec/v1
|
| 2 |
-
oid sha256:
|
| 3 |
size 4321128
|
|
|
|
| 1 |
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:69446c0ed6ada667cbc27e2cf08f257dd2d3e9955657d5ac87eff09d3350ed3d
|
| 3 |
size 4321128
|
build/torch212-cxx11-xpu20253-x86_64-linux/_ops.py
CHANGED
|
@@ -1,9 +1,9 @@
|
|
| 1 |
import torch
|
| 2 |
-
from . import
|
| 3 |
-
ops = torch.ops.
|
| 4 |
|
| 5 |
def add_op_namespace_prefix(op_name: str):
|
| 6 |
"""
|
| 7 |
Prefix op by namespace.
|
| 8 |
"""
|
| 9 |
-
return f"
|
|
|
|
| 1 |
import torch
|
| 2 |
+
from . import _megablocks_xpu_addf474
|
| 3 |
+
ops = torch.ops._megablocks_xpu_addf474
|
| 4 |
|
| 5 |
def add_op_namespace_prefix(op_name: str):
|
| 6 |
"""
|
| 7 |
Prefix op by namespace.
|
| 8 |
"""
|
| 9 |
+
return f"_megablocks_xpu_addf474::{op_name}"
|
build/torch212-cxx11-xpu20253-x86_64-linux/megablocks/__init__.py
DELETED
|
@@ -1,26 +0,0 @@
|
|
| 1 |
-
import ctypes
|
| 2 |
-
import importlib.util
|
| 3 |
-
import sys
|
| 4 |
-
from pathlib import Path
|
| 5 |
-
from types import ModuleType
|
| 6 |
-
|
| 7 |
-
|
| 8 |
-
def _import_from_path(file_path: Path) -> ModuleType:
|
| 9 |
-
# We cannot use the module name as-is, after adding it to `sys.modules`,
|
| 10 |
-
# it would also be used for other imports. So, we make a module name that
|
| 11 |
-
# depends on the path for it to be unique using the hex-encoded hash of
|
| 12 |
-
# the path.
|
| 13 |
-
path_hash = "{:x}".format(ctypes.c_size_t(hash(file_path.absolute())).value)
|
| 14 |
-
module_name = path_hash
|
| 15 |
-
spec = importlib.util.spec_from_file_location(module_name, file_path)
|
| 16 |
-
if spec is None:
|
| 17 |
-
raise ImportError(f"Cannot load spec for {module_name} from {file_path}")
|
| 18 |
-
module = importlib.util.module_from_spec(spec)
|
| 19 |
-
if module is None:
|
| 20 |
-
raise ImportError(f"Cannot load module {module_name} from spec")
|
| 21 |
-
sys.modules[module_name] = module
|
| 22 |
-
spec.loader.exec_module(module) # type: ignore
|
| 23 |
-
return module
|
| 24 |
-
|
| 25 |
-
|
| 26 |
-
globals().update(vars(_import_from_path(Path(__file__).parent.parent / "__init__.py")))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
build/torch212-cxx11-xpu20253-x86_64-linux/metadata.json
CHANGED
|
@@ -1,9 +1,10 @@
|
|
| 1 |
{
|
| 2 |
"name": "megablocks",
|
| 3 |
-
"id": "
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"python-depends": [],
|
|
|
|
| 7 |
"backend": {
|
| 8 |
"type": "xpu"
|
| 9 |
},
|
|
@@ -38,8 +39,8 @@
|
|
| 38 |
"_layers/mpu.py": "BX61/pmqRCVeEMnzpuK8x5NpITBJyZfOawGnDegW1O8=",
|
| 39 |
"_layers/router.py": "GKC/5//Uf43HT6zXlBTStCMTml2YXn7Jbgq9EgTgDC0=",
|
| 40 |
"_layers/sharedexpert_registry.py": "GfwHKjbC7lBXYq9kNifp3HudTdZKSqwsFxGxFLc2krw=",
|
| 41 |
-
"
|
| 42 |
-
"_ops.py": "
|
| 43 |
"_version.py": "W3l9oUrnmvfBoIy06Lv0y2ezTQhNU04ug/txHcdWJD0=",
|
| 44 |
"backend/__init__.py": "4bunAqqjwE93nym8CigHbqVAal7UaYGdONb8lGRfrek=",
|
| 45 |
"backend/kernels.py": "2UFNTlmSDk3S5y6UW14MJMPVVsVhmozxQ5wFDxtZHhQ=",
|
|
@@ -51,7 +52,6 @@
|
|
| 51 |
"grouped_gemm/ops.py": "Kj7yRgv7afB5SVwXIkN9RNIF2kxD2qrZ0dVsKX5ph2s=",
|
| 52 |
"grouped_gemm_util.py": "bVPrdsEtBVJJ+v/bNrGNf/QHNlVL4VxHg92HvaL3oq4=",
|
| 53 |
"layers.py": "mNTlX2oKJyHBJllRd24lPNdGLa2H9cEUgVLvkZJpoNo=",
|
| 54 |
-
"megablocks/__init__.py": "DFYPlrhXwYjEqCl/8n0SmWGZV8NFml5DPhMjKfv98GY=",
|
| 55 |
"ops/__init__.py": "MHQsiMEgkKeZCa8TAYFMGDFnwFIgijfAN7DhMUt/48U=",
|
| 56 |
"ops/all_to_all_benchmark.py": "sxXSCSDdsh+62IWpKNTyBphLTngY1jlgJ1WgNooSEO8=",
|
| 57 |
"ops/binned_gather.py": "pkWne+ThfiTR/bMbDs2QDimKxAI4klqD1FdvhytE7lU=",
|
|
@@ -96,11 +96,11 @@
|
|
| 96 |
"provenance": {
|
| 97 |
"kernel-builder": {
|
| 98 |
"version": "0.17.0-dev0",
|
| 99 |
-
"
|
| 100 |
"dirty": false
|
| 101 |
},
|
| 102 |
"kernel": {
|
| 103 |
-
"
|
| 104 |
"dirty": false
|
| 105 |
}
|
| 106 |
}
|
|
|
|
| 1 |
{
|
| 2 |
"name": "megablocks",
|
| 3 |
+
"id": "_megablocks_xpu_addf474",
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"python-depends": [],
|
| 7 |
+
"kernel-depends": [],
|
| 8 |
"backend": {
|
| 9 |
"type": "xpu"
|
| 10 |
},
|
|
|
|
| 39 |
"_layers/mpu.py": "BX61/pmqRCVeEMnzpuK8x5NpITBJyZfOawGnDegW1O8=",
|
| 40 |
"_layers/router.py": "GKC/5//Uf43HT6zXlBTStCMTml2YXn7Jbgq9EgTgDC0=",
|
| 41 |
"_layers/sharedexpert_registry.py": "GfwHKjbC7lBXYq9kNifp3HudTdZKSqwsFxGxFLc2krw=",
|
| 42 |
+
"_megablocks_xpu_addf474.abi3.so": "aURsDtatpmfLwn4s8I8lfdLT6ZVWV9Wsh+/wnTNQ7T0=",
|
| 43 |
+
"_ops.py": "QBFLHcpmqB9+JkTkhdpkiEE9yER4FUDrsZgGzybNoVc=",
|
| 44 |
"_version.py": "W3l9oUrnmvfBoIy06Lv0y2ezTQhNU04ug/txHcdWJD0=",
|
| 45 |
"backend/__init__.py": "4bunAqqjwE93nym8CigHbqVAal7UaYGdONb8lGRfrek=",
|
| 46 |
"backend/kernels.py": "2UFNTlmSDk3S5y6UW14MJMPVVsVhmozxQ5wFDxtZHhQ=",
|
|
|
|
| 52 |
"grouped_gemm/ops.py": "Kj7yRgv7afB5SVwXIkN9RNIF2kxD2qrZ0dVsKX5ph2s=",
|
| 53 |
"grouped_gemm_util.py": "bVPrdsEtBVJJ+v/bNrGNf/QHNlVL4VxHg92HvaL3oq4=",
|
| 54 |
"layers.py": "mNTlX2oKJyHBJllRd24lPNdGLa2H9cEUgVLvkZJpoNo=",
|
|
|
|
| 55 |
"ops/__init__.py": "MHQsiMEgkKeZCa8TAYFMGDFnwFIgijfAN7DhMUt/48U=",
|
| 56 |
"ops/all_to_all_benchmark.py": "sxXSCSDdsh+62IWpKNTyBphLTngY1jlgJ1WgNooSEO8=",
|
| 57 |
"ops/binned_gather.py": "pkWne+ThfiTR/bMbDs2QDimKxAI4klqD1FdvhytE7lU=",
|
|
|
|
| 96 |
"provenance": {
|
| 97 |
"kernel-builder": {
|
| 98 |
"version": "0.17.0-dev0",
|
| 99 |
+
"commit": "7d2828e6592de641b44fb8eb896719101a2ab2fb",
|
| 100 |
"dirty": false
|
| 101 |
},
|
| 102 |
"kernel": {
|
| 103 |
+
"commit": "addf4741fafb55b4c101b15985dfaeea7c7c2f3a",
|
| 104 |
"dirty": false
|
| 105 |
}
|
| 106 |
}
|
build/torch212-cxx11-xpu20253-x86_64-linux/metadata.json.sigstore
CHANGED
|
@@ -1 +1 @@
|
|
| 1 |
-
{"mediaType":"application/vnd.dev.sigstore.bundle.v0.3+json", "verificationMaterial":{"certificate":{"rawBytes":"MIIHTDCCBtKgAwIBAgIUXl5Lj0ZZxmJSBAElnUur5bPOFSUwCgYIKoZIzj0EAwMwNzEVMBMGA1UEChMMc2lnc3RvcmUuZGV2MR4wHAYDVQQDExVzaWdzdG9yZS1pbnRlcm1lZGlhdGUwHhcNMjYwODI1MTIwNzQ3WhcNMjYwODI1MTIxNzQ3WjAAMFkwEwYHKoZIzj0CAQYIKoZIzj0DAQcDQgAErc2vjnOd2F7gDDhqBurUUxT/ujdwHbHjt8JcM57UgNubELzPeSCPV+EjHNgld7qmRfUjLa21P6DWD7BypZNRNqOCBfEwggXtMA4GA1UdDwEB/wQEAwIHgDATBgNVHSUEDDAKBggrBgEFBQcDAzAdBgNVHQ4EFgQU+Uxv/tFJgmXCXe5vuYfg2ML3TocwHwYDVR0jBBgwFoAU39Ppz1YkEZb5qNjpKFWixi4YZD8wawYDVR0RAQH/BGEwX4ZdaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5Ly5naXRodWIvd29ya2Zsb3dzL2J1aWxkLnlhbWxAcmVmcy9oZWFkcy9tYWluMDkGCisGAQQBg78wAQEEK2h0dHBzOi8vdG9rZW4uYWN0aW9ucy5naXRodWJ1c2VyY29udGVudC5jb20wHwYKKwYBBAGDvzABAgQRd29ya2Zsb3dfZGlzcGF0Y2gwNgYKKwYBBAGDvzABAwQoYzRkMGRjNTY2NjFiYWRiOWNjNDdiZGU2NDdlMDhmZTNiNGJmNzQ1ODATBgorBgEEAYO/MAEEBAVCdWlsZDArBgorBgEEAYO/MAEFBB1odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eTAdBgorBgEEAYO/MAEGBA9yZWZzL2hlYWRzL21haW4wOwYKKwYBBAGDvzABCAQtDCtodHRwczovL3Rva2VuLmFjdGlvbnMuZ2l0aHVidXNlcmNvbnRlbnQuY29tMG0GCisGAQQBg78wAQkEXwxdaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5Ly5naXRodWIvd29ya2Zsb3dzL2J1aWxkLnlhbWxAcmVmcy9oZWFkcy9tYWluMDgGCisGAQQBg78wAQoEKgwoYzRkMGRjNTY2NjFiYWRiOWNjNDdiZGU2NDdlMDhmZTNiNGJmNzQ1ODAbBgorBgEEAYO/MAELBA0MC3NlbGYtaG9zdGVkMEAGCisGAQQBg78wAQwEMgwwaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5MDgGCisGAQQBg78wAQ0EKgwoYzRkMGRjNTY2NjFiYWRiOWNjNDdiZGU2NDdlMDhmZTNiNGJmNzQ1ODAfBgorBgEEAYO/MAEOBBEMD3JlZnMvaGVhZHMvbWFpbjAaBgorBgEEAYO/MAEPBAwMCjEwNzE0NzU1MjkwLgYKKwYBBAGDvzABEAQgDB5odHRwczovL2dpdGh1Yi5jb20vaHVnZ2luZ2ZhY2UwGAYKKwYBBAGDvzABEQQKDAgyNTcyMDc0MzBtBgorBgEEAYO/MAESBF8MXWh0dHBzOi8vZ2l0aHViLmNvbS9odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eS8uZ2l0aHViL3dvcmtmbG93cy9idWlsZC55YW1sQHJlZnMvaGVhZHMvbWFpbjA4BgorBgEEAYO/MAETBCoMKGM0ZDBkYzU2NjYxYmFkYjljYzQ3YmRlNjQ3ZTA4ZmUzYjRiZjc0NTgwIQYKKwYBBAGDvzABFAQTDBF3b3JrZmxvd19kaXNwYXRjaDBkBgorBgEEAYO/MAEVBFYMVGh0dHBzOi8vZ2l0aHViLmNvbS9odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eS9hY3Rpb25zL3J1bnMvMzI4NDQ2ODk4OTQvYXR0ZW1wdHMvMTAWBgorBgEEAYO/MAEWBAgMBnB1YmxpYzBGBgorBgEEAYO/MAEYBDgMNnJlcG86aHVnZ2luZ2ZhY2Uva2VybmVscy1jb21tdW5pdHk6cmVmOnJlZnMvaGVhZHMvbWFpbjCBiwYKKwYBBAHWeQIEAgR9BHsAeQB3AN09MGrGxxEyYxkeHJlnNwKiSl643jyt/4eKcoAvKe6OAAABoDjRxswAAAQDAEgwRgIhAN+W8o+X7UgM1dt/DI3YoGQwFZtyRfMAGOEyZnvXz5DPAiEA0zBGzcriJLXqkz+htQgXqRdIzS3LI70Vj61poOh2B+MwCgYIKoZIzj0EAwMDaAAwZQIxANO7WbSfiTPHzqSEAfB+JpfkSWRdxfujCm2jE0ncT/4jIC6v14qyZVBlHp7+teGYYAIwOx4kJzvP6rP/tlp5LPhe7rF5kD+Vja6AtjwfMNXLBgTcU4Aew0pof22Kdj5DsGO9"}, "tlogEntries":[{"logIndex":"2583274380", "logId":{"keyId":"wNI9atQGlz+VWfO6LRygH4QUfY/8W4RFwiT5i5WRgB0="}, "kindVersion":{"kind":"hashedrekord", "version":"0.0.1"}, "integratedTime":"1787659667", "inclusionPromise":{"signedEntryTimestamp":"MEYCIQCpPbM0AS/szLFi2f2RNpUPV4OvvmCRi5a6WAzN9iI8/AIhAJA2d6b7uUNN4ebditgvyIz/6tvKd7yv4hi/AP2Mwxti"}, "inclusionProof":{"logIndex":"2461370118", "rootHash":"cQ7wW1xW/L20rkOvTt9r9iCKPByWXNI6FCkRaju37Ts=", "treeSize":"2461370120", "hashes":["y7dnw8OA1nxLA6kcdGFSV0dFZoGA1uZ1xOITOKTf2Js=", "uJ0tsePLf6AS/EahjuJN0+b+EnlFud7NirVAZMI0B4I=", "jjdaExfTfO3Htve6MMlB8UBrmhQIgQW9lSQ4kVJYuG4=", "BJlxWyo+0U2jxgcimrAeK1pey6FSpHWLOetp0UXRWI4=", "JsX9xGIEwFv7MFelTobmTLz4j1Jl78aSs2G7dkybw0w=", "nwobU/9yRWhsbw80xq3x/uNV1nniMWdUKAqc3UWYsJM=", "EXKrSbeG2qm5APD5QSDN66X7UHz9OSofiW9PyNU7fPM=", "cjIeeutzThxIPTlFFHJqa7bAzn3k2FzK8Nrdd56OWvE=", "+pNXWGIXzU1dqSy/aePHKYtSkDnFinI+/bL5HnurJ6w=", "G1R9F5B2KizFA2NyhDznjNKXPOF1SDmbGHz+omLG7TY=", "RZ7XwGuqMr5gDUh2HpUcaXp+AmIiEcPJiMui4GLJE7o=", "mBm+vQtn0C4thMnlTnxfxqowq1dXsPBCKaJ85da3JeU=", "b80/J/68RC8/tx5xRdlKk6MmTpZiVRotquWBrE/z1pI=", "SndbMKVtcTenAkwi2JBfGzD+mhexp1qJbRIY+A1JRIU=", "xH/DCseLHr9eKoYT8qsORZK7zVdEGYWHuVtsVrD95wY="], "checkpoint":{"envelope":"rekor.sigstore.dev - 1193050959916656506\n2461370120\ncQ7wW1xW/L20rkOvTt9r9iCKPByWXNI6FCkRaju37Ts=\n\n— rekor.sigstore.dev wNI9ajBFAiA9mr274M9R9qIT6J8dTEeSqj88FaQxjItF4bRLmWCsmwIhAMsWTJ1A8HcsmgMoEJo+YpOQ01H1eQMrBHPvtx3pT2mT\n"}}, "canonicalizedBody":"eyJhcGlWZXJzaW9uIjoiMC4wLjEiLCJraW5kIjoiaGFzaGVkcmVrb3JkIiwic3BlYyI6eyJkYXRhIjp7Imhhc2giOnsiYWxnb3JpdGhtIjoic2hhMjU2IiwidmFsdWUiOiJlMjU4NGVmMDU2OGZkMmVkMTM4MmE2ODZjODNhMGIyNzNhMmZjYzE4MWI1YzZkMmQ2ODRjNzgwNDc5MzNjM2JkIn19LCJzaWduYXR1cmUiOnsiY29udGVudCI6Ik1FVUNJUUNwRlg3N3NjWTF1VUZwOWdNQWIrUWJ3VVluYlNFdXNjYURBYlJCbFVQNHJnSWdCT1E3T3U4QnNaTGo1RWNtNDArOTViOXZoUFhGVmZZNmxxV3kvUHpBbWE0PSIsInB1YmxpY0tleSI6eyJjb250ZW50IjoiTFMwdExTMUNSVWRKVGlCRFJWSlVTVVpKUTBGVVJTMHRMUzB0Q2sxSlNVaFVSRU5EUW5STFowRjNTVUpCWjBsVldHdzFUR293V2xwNGJVcFRRa0ZGYkc1VmRYSTFZbEJQUmxOVmQwTm5XVWxMYjFwSmVtb3dSVUYzVFhjS1RucEZWazFDVFVkQk1WVkZRMmhOVFdNeWJHNWpNMUoyWTIxVmRWcEhWakpOVWpSM1NFRlpSRlpSVVVSRmVGWjZZVmRrZW1SSE9YbGFVekZ3WW01U2JBcGpiVEZzV2tkc2FHUkhWWGRJYUdOT1RXcFpkMDlFU1RGTlZFbDNUbnBSTTFkb1kwNU5hbGwzVDBSSk1VMVVTWGhPZWxFelYycEJRVTFHYTNkRmQxbElDa3R2V2tsNmFqQkRRVkZaU1V0dldrbDZhakJFUVZGalJGRm5RVVZ5WXpKMmFtNVBaREpHTjJkRVJHaHhRblZ5VlZWNFZDOTFhbVIzU0dKSWFuUTRTbU1LVFRVM1ZXZE9kV0pGVEhwUVpWTkRVRllyUldwSVRtZHNaRGR4YlZKbVZXcE1ZVEl4VURaRVYwUTNRbmx3V2s1U1RuRlBRMEptUlhkbloxaDBUVUUwUndwQk1WVmtSSGRGUWk5M1VVVkJkMGxJWjBSQlZFSm5UbFpJVTFWRlJFUkJTMEpuWjNKQ1owVkdRbEZqUkVGNlFXUkNaMDVXU0ZFMFJVWm5VVlVyVlhoMkNpOTBSa3BuYlZoRFdHVTFkblZaWm1jeVRVd3pWRzlqZDBoM1dVUldVakJxUWtKbmQwWnZRVlV6T1ZCd2VqRlphMFZhWWpWeFRtcHdTMFpYYVhocE5Ga0tXa1E0ZDJGM1dVUldVakJTUVZGSUwwSkhSWGRZTkZwa1lVaFNNR05JVFRaTWVUbHVZVmhTYjJSWFNYVlpNamwwVERKb01Wb3laSEJpYldSdFdWZE9iQXBNTW5Sc1kyMDFiR0pJVFhSWk1qbDBZbGhXZFdGWVVqVk1lVFZ1WVZoU2IyUlhTWFprTWpsNVlUSmFjMkl6WkhwTU1rb3hZVmQ0YTB4dWJHaGlWM2hCQ21OdFZtMWplVGx2V2xkR2EyTjVPWFJaVjJ4MVRVUnJSME5wYzBkQlVWRkNaemM0ZDBGUlJVVkxNbWd3WkVoQ2VrOXBPSFprUnpseVdsYzBkVmxYVGpBS1lWYzVkV041Tlc1aFdGSnZaRmRLTVdNeVZubFpNamwxWkVkV2RXUkROV3BpTWpCM1NIZFpTMHQzV1VKQ1FVZEVkbnBCUWtGblVWSmtNamw1WVRKYWN3cGlNMlJtV2tkc2VtTkhSakJaTW1kM1RtZFpTMHQzV1VKQ1FVZEVkbnBCUWtGM1VXOVplbEpyVFVkU2FrNVVXVEpPYWtacFdWZFNhVTlYVG1wT1JHUnBDbHBIVlRKT1JHUnNUVVJvYlZwVVRtbE9SMHB0VG5wUk1VOUVRVlJDWjI5eVFtZEZSVUZaVHk5TlFVVkZRa0ZXUTJSWGJITmFSRUZ5UW1kdmNrSm5SVVVLUVZsUEwwMUJSVVpDUWpGdlpGZGtibUZYTlc1YWJVWnFXbE01Y2xwWVNuVmFWM2g2VEZkT2RtSlhNVEZpYld3d1pWUkJaRUpuYjNKQ1owVkZRVmxQTHdwTlFVVkhRa0U1ZVZwWFducE1NbWhzV1ZkU2Vrd3lNV2hoVnpSM1QzZFpTMHQzV1VKQ1FVZEVkbnBCUWtOQlVYUkVRM1J2WkVoU2QyTjZiM1pNTTFKMkNtRXlWblZNYlVacVpFZHNkbUp1VFhWYU1td3dZVWhXYVdSWVRteGpiVTUyWW01U2JHSnVVWFZaTWpsMFRVY3dSME5wYzBkQlVWRkNaemM0ZDBGUmEwVUtXSGQ0WkdGSVVqQmpTRTAyVEhrNWJtRllVbTlrVjBsMVdUSTVkRXd5YURGYU1tUndZbTFrYlZsWFRteE1NblJzWTIwMWJHSklUWFJaTWpsMFlsaFdkUXBoV0ZJMVRIazFibUZZVW05a1YwbDJaREk1ZVdFeVduTmlNMlI2VERKS01XRlhlR3RNYm14b1lsZDRRV050Vm0xamVUbHZXbGRHYTJONU9YUlpWMngxQ2sxRVowZERhWE5IUVZGUlFtYzNPSGRCVVc5RlMyZDNiMWw2VW10TlIxSnFUbFJaTWs1cVJtbFpWMUpwVDFkT2FrNUVaR2xhUjFVeVRrUmtiRTFFYUcwS1dsUk9hVTVIU20xT2VsRXhUMFJCWWtKbmIzSkNaMFZGUVZsUEwwMUJSVXhDUVRCTlF6Tk9iR0pIV1hSaFJ6bDZaRWRXYTAxRlFVZERhWE5IUVZGUlFncG5OemgzUVZGM1JVMW5kM2RoU0ZJd1kwaE5Oa3g1T1c1aFdGSnZaRmRKZFZreU9YUk1NbWd4V2pKa2NHSnRaRzFaVjA1c1RESjBiR050Tld4aVNFMTBDbGt5T1hSaVdGWjFZVmhTTlUxRVowZERhWE5IUVZGUlFtYzNPSGRCVVRCRlMyZDNiMWw2VW10TlIxSnFUbFJaTWs1cVJtbFpWMUpwVDFkT2FrNUVaR2tLV2tkVk1rNUVaR3hOUkdodFdsUk9hVTVIU20xT2VsRXhUMFJCWmtKbmIzSkNaMFZGUVZsUEwwMUJSVTlDUWtWTlJETktiRnB1VFhaaFIxWm9Xa2hOZGdwaVYwWndZbXBCWVVKbmIzSkNaMFZGUVZsUEwwMUJSVkJDUVhkTlEycEZkMDU2UlRCT2VsVXhUV3ByZDB4bldVdExkMWxDUWtGSFJIWjZRVUpGUVZGbkNrUkNOVzlrU0ZKM1kzcHZka3d5WkhCa1IyZ3hXV2sxYW1JeU1IWmhTRlp1V2pKc2RWb3lXbWhaTWxWM1IwRlpTMHQzV1VKQ1FVZEVkbnBCUWtWUlVVc0tSRUZuZVU1VVkzbE5SR013VFhwQ2RFSm5iM0pDWjBWRlFWbFBMMDFCUlZOQ1JqaE5XRmRvTUdSSVFucFBhVGgyV2pKc01HRklWbWxNYlU1MllsTTVid3BrVjJSdVlWYzFibHB0Um1wYVV6bHlXbGhLZFZwWGVIcE1WMDUyWWxjeE1XSnRiREJsVXpoMVdqSnNNR0ZJVm1sTU0yUjJZMjEwYldKSE9UTmplVGxwQ21SWGJITmFRelUxV1ZjeGMxRklTbXhhYmsxMllVZFdhRnBJVFhaaVYwWndZbXBCTkVKbmIzSkNaMFZGUVZsUEwwMUJSVlJDUTI5TlMwZE5NRnBFUW1zS1dYcFZNazVxV1hoWmJVWnJXV3BzYWxsNlVUTlpiVkpzVG1wUk0xcFVRVFJhYlZWNldXcFNhVnBxWXpCT1ZHZDNTVkZaUzB0M1dVSkNRVWRFZG5wQlFncEdRVkZVUkVKR00ySXpTbkphYlhoMlpERTVhMkZZVG5kWldGSnFZVVJDYTBKbmIzSkNaMFZGUVZsUEwwMUJSVlpDUmxsTlZrZG9NR1JJUW5wUGFUaDJDbG95YkRCaFNGWnBURzFPZG1KVE9XOWtWMlJ1WVZjMWJscHRSbXBhVXpseVdsaEtkVnBYZUhwTVYwNTJZbGN4TVdKdGJEQmxVemxvV1ROU2NHSXlOWG9LVEROS01XSnVUWFpOZWtrMFRrUlJNazlFYXpSUFZGRjJXVmhTTUZwWE1YZGtTRTEyVFZSQlYwSm5iM0pDWjBWRlFWbFBMMDFCUlZkQ1FXZE5RbTVDTVFwWmJYaHdXWHBDUjBKbmIzSkNaMFZGUVZsUEwwMUJSVmxDUkdkTlRtNUtiR05IT0RaaFNGWnVXakpzZFZveVdtaFpNbFYyWVRKV2VXSnRWbk5qZVRGcUNtSXlNWFJrVnpWd1pFaHJObU50Vm0xUGJrcHNXbTVOZG1GSFZtaGFTRTEyWWxkR2NHSnFRMEpwZDFsTFMzZFpRa0pCU0ZkbFVVbEZRV2RTT1VKSWMwRUtaVkZDTTBGT01EbE5SM0pIZUhoRmVWbDRhMlZJU214dVRuZExhVk5zTmpRemFubDBMelJsUzJOdlFYWkxaVFpQUVVGQlFtOUVhbEo0YzNkQlFVRlJSQXBCUldkM1VtZEphRUZPSzFjNGJ5dFlOMVZuVFRGa2RDOUVTVE5aYjBkUmQwWmFkSGxTWmsxQlIwOUZlVnB1ZGxoNk5VUlFRV2xGUVRCNlFrZDZZM0pwQ2twTVdIRnJlaXRvZEZGbldIRlNaRWw2VXpOTVNUY3dWbW8yTVhCdlQyZ3lRaXROZDBObldVbExiMXBKZW1vd1JVRjNUVVJoUVVGM1dsRkplRUZPVHpjS1YySlRabWxVVUVoNmNWTkZRV1pDSzBwd1ptdFRWMUprZUdaMWFrTnRNbXBGTUc1alZDODBha2xETm5ZeE5IRjVXbFpDYkVod055dDBaVWRaV1VGSmR3cFBlRFJyU25wMlVEWnlVQzkwYkhBMVRGQm9aVGR5UmpWclJDdFdhbUUyUVhScWQyWk5UbGhNUW1kVVkxVTBRV1YzTUhCdlpqSXlTMlJxTlVSelIwODVDaTB0TFMwdFJVNUVJRU5GVWxSSlJrbERRVlJGTFMwdExTMEsifX19fQ=="}], "timestampVerificationData":{"rfc3161Timestamps":[{"signedTimestamp":"MIICyjADAgEAMIICwQYJKoZIhvcNAQcCoIICsjCCAq4CAQMxDTALBglghkgBZQMEAgEwgbcGCyqGSIb3DQEJEAEEoIGnBIGkMIGhAgEBBgkrBgEEAYO/MAIwMTANBglghkgBZQMEAgEFAAQgAF4NzZ4QDxxyny7GD5G9gutsEM6gS48lZIjdv0q0L2oCFCqbR84O99XA9ORjAz6eaf5ry4FfGA8yMDI2MDgyNTEyMDc0N1owAwIBAaAypDAwLjEVMBMGA1UEChMMc2lnc3RvcmUuZGV2MRUwEwYDVQQDEwxzaWdzdG9yZS10c2GgADGCAdwwggHYAgEBMFEwOTEVMBMGA1UEChMMc2lnc3RvcmUuZGV2MSAwHgYDVQQDExdzaWdzdG9yZS10c2Etc2VsZnNpZ25lZAIUOhNULwyQYe68wUMvy4qOiyojiwwwCwYJYIZIAWUDBAIBoIH8MBoGCSqGSIb3DQEJAzENBgsqhkiG9w0BCRABBDAcBgkqhkiG9w0BCQUxDxcNMjYwODI1MTIwNzQ3WjAvBgkqhkiG9w0BCQQxIgQg0xV8kuoZmTLPREuzDn9A2nFk6uc+jwxngJv06AITOqwwgY4GCyqGSIb3DQEJEAIvMX8wfTB7MHkEIIX5J7wHq2LKw7RDVsEO/IGyxog/2nq55thw2dE6zQW3MFUwPaQ7MDkxFTATBgNVBAoTDHNpZ3N0b3JlLmRldjEgMB4GA1UEAxMXc2lnc3RvcmUtdHNhLXNlbGZzaWduZWQCFDoTVC8MkGHuvMFDL8uKjosqI4sMMAoGCCqGSM49BAMCBGgwZgIxAMORtmStWCwKwKsaiLGpF8sh0OG1iC4tOrFRmmNktIbLmI9CIQUsE5ApgIm8LWbMwAIxAOEkQlSqUSkaNMPcj+l3vymorX1kyGFtf1kzOX789cV5zvO8oJs+Z4Ob8ZAs+jdMhQ=="}]}}, "messageSignature":{"messageDigest":{"algorithm":"SHA2_256", "digest":"4lhO8FaP0u0TgqaGyDoLJzovzBgbXG0taEx4BHkzw70="}, "signature":"MEUCIQCpFX77scY1uUFp9gMAb+QbwUYnbSEuscaDAbRBlUP4rgIgBOQ7Ou8BsZLj5Ecm40+95b9vhPXFVfY6lqWy/PzAma4="}}
|
|
|
|
| 1 |
+
{"mediaType":"application/vnd.dev.sigstore.bundle.v0.3+json", "verificationMaterial":{"certificate":{"rawBytes":"MIIHTTCCBtKgAwIBAgIUeFaEsDg9TNOkT8jPh1eH7XW1ZG0wCgYIKoZIzj0EAwMwNzEVMBMGA1UEChMMc2lnc3RvcmUuZGV2MR4wHAYDVQQDExVzaWdzdG9yZS1pbnRlcm1lZGlhdGUwHhcNMjYwODI3MDcwMTMzWhcNMjYwODI3MDcxMTMzWjAAMFkwEwYHKoZIzj0CAQYIKoZIzj0DAQcDQgAE3Mql367h5c67VBxvbhUlurjO4WCr8b8hEpLd/9c1GxY5TKhxm/QZb3/NAioZJru7lljD/Lq62har3sLL0gz2J6OCBfEwggXtMA4GA1UdDwEB/wQEAwIHgDATBgNVHSUEDDAKBggrBgEFBQcDAzAdBgNVHQ4EFgQUdcb1KKYaYn60xVmuNMNRUFJtLsQwHwYDVR0jBBgwFoAU39Ppz1YkEZb5qNjpKFWixi4YZD8wawYDVR0RAQH/BGEwX4ZdaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5Ly5naXRodWIvd29ya2Zsb3dzL2J1aWxkLnlhbWxAcmVmcy9oZWFkcy9tYWluMDkGCisGAQQBg78wAQEEK2h0dHBzOi8vdG9rZW4uYWN0aW9ucy5naXRodWJ1c2VyY29udGVudC5jb20wHwYKKwYBBAGDvzABAgQRd29ya2Zsb3dfZGlzcGF0Y2gwNgYKKwYBBAGDvzABAwQoYWRkZjQ3NDFmYWZiNTViNGMxMDFiMTU5ODVkZmFlZWE3YzdjMmYzYTATBgorBgEEAYO/MAEEBAVCdWlsZDArBgorBgEEAYO/MAEFBB1odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eTAdBgorBgEEAYO/MAEGBA9yZWZzL2hlYWRzL21haW4wOwYKKwYBBAGDvzABCAQtDCtodHRwczovL3Rva2VuLmFjdGlvbnMuZ2l0aHVidXNlcmNvbnRlbnQuY29tMG0GCisGAQQBg78wAQkEXwxdaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5Ly5naXRodWIvd29ya2Zsb3dzL2J1aWxkLnlhbWxAcmVmcy9oZWFkcy9tYWluMDgGCisGAQQBg78wAQoEKgwoYWRkZjQ3NDFmYWZiNTViNGMxMDFiMTU5ODVkZmFlZWE3YzdjMmYzYTAbBgorBgEEAYO/MAELBA0MC3NlbGYtaG9zdGVkMEAGCisGAQQBg78wAQwEMgwwaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5MDgGCisGAQQBg78wAQ0EKgwoYWRkZjQ3NDFmYWZiNTViNGMxMDFiMTU5ODVkZmFlZWE3YzdjMmYzYTAfBgorBgEEAYO/MAEOBBEMD3JlZnMvaGVhZHMvbWFpbjAaBgorBgEEAYO/MAEPBAwMCjEwNzE0NzU1MjkwLgYKKwYBBAGDvzABEAQgDB5odHRwczovL2dpdGh1Yi5jb20vaHVnZ2luZ2ZhY2UwGAYKKwYBBAGDvzABEQQKDAgyNTcyMDc0MzBtBgorBgEEAYO/MAESBF8MXWh0dHBzOi8vZ2l0aHViLmNvbS9odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eS8uZ2l0aHViL3dvcmtmbG93cy9idWlsZC55YW1sQHJlZnMvaGVhZHMvbWFpbjA4BgorBgEEAYO/MAETBCoMKGFkZGY0NzQxZmFmYjU1YjRjMTAxYjE1OTg1ZGZhZWVhN2M3YzJmM2EwIQYKKwYBBAGDvzABFAQTDBF3b3JrZmxvd19kaXNwYXRjaDBkBgorBgEEAYO/MAEVBFYMVGh0dHBzOi8vZ2l0aHViLmNvbS9odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eS9hY3Rpb25zL3J1bnMvMzMwNDUzOTU5NTIvYXR0ZW1wdHMvMTAWBgorBgEEAYO/MAEWBAgMBnB1YmxpYzBGBgorBgEEAYO/MAEYBDgMNnJlcG86aHVnZ2luZ2ZhY2Uva2VybmVscy1jb21tdW5pdHk6cmVmOnJlZnMvaGVhZHMvbWFpbjCBiwYKKwYBBAHWeQIEAgR9BHsAeQB3AN09MGrGxxEyYxkeHJlnNwKiSl643jyt/4eKcoAvKe6OAAABoEIGI2UAAAQDAEgwRgIhAMEl1OtXrc+rgzO8F3FC5n353ulp7ViQeVkv/zbKyx+zAiEAw3YdtCwMf1v8wCiqHIkKwR8cwG1hKVQ6OU2C8Pb1gwYwCgYIKoZIzj0EAwMDaQAwZgIxANkQhJ0dt8SOSna6XMugwcm27C39vXpGxMZpkhkGz9jtX32uMYxS7PU7bOQ8UK/J8gIxAJOXk26SNI4T6amV9Qajdke3UVMOwIWFLxjo1eq6gai1C+I5ew1b3oJ9p5MB56BxMA=="}, "tlogEntries":[{"logIndex":"2613969849", "logId":{"keyId":"wNI9atQGlz+VWfO6LRygH4QUfY/8W4RFwiT5i5WRgB0="}, "kindVersion":{"kind":"hashedrekord", "version":"0.0.1"}, "integratedTime":"1787814093", "inclusionPromise":{"signedEntryTimestamp":"MEQCIEw/H/N+EZIwnyYubbQjbe4l1vEThipNr1UfJoGz9FDkAiBDxRJPxNCN7YSdiMTbOeTcChXz/3rXbPSz4Ggs0cq0xg=="}, "inclusionProof":{"logIndex":"2492065587", "rootHash":"HMabR7yd/EVq4WbRn7sHFJlrkYwtTQFLVw8bfn+V/3U=", "treeSize":"2492065618", "hashes":["/47HCMUHWevEtJadH20kmmpM1uwTvrnQPGl7XIxbn0g=", "S6DPOYJB7pC6yD+CDg/76n9DvCEnYP1bWwSyBOxtyBQ=", "08UjMwpJoruGrE0eG1Oa0j+eZiI3bGuJgz49Ck04JEE=", "y1zjtlmIKKMItKXBRoSrfvYznV5kx1gp1TH/JZnMvDQ=", "3QWj8I1L/XMP6bvhV6UfDbCFQtLroZAXR8elY5sR1f8=", "eXbrdCAOkgErRva+XLB62kDcJwlWngQ4NqKGch5Wz40=", "YpZEITJioo+i82eneJRh/rPljqYzpU6JkSijrmcgZZs=", "7f1DiyYwq4ZLq0MB+i2yDsDcVgCFJbIE4HvxSM1mYAo=", "+1z062t3T904zuuUCVWmzD3bmZKoTqs9bZ1VmfwnByQ=", "JZTdykpRua3w/oyS32G9vqOIoVajdstcZWLNMlPQfac=", "UbnDpZmg3ac4T8gCmplTGV9yU5j3fWOpeZi+i3apFqM=", "ygNYNZXNiY/+rjtW6HI0ngBM75qjGCTD6YstKiLiOZY=", "RvtgIftU0w4e16LAcGD7A1ErF9x2p/QLAyDLR+PaDMU=", "kW2yseLcah3rCzJJ1VS5Uhr9NCuHTJMDHsUkYy/0dV0=", "VGrFt5JvlIUQvoGEdBBxrCe3lYbu581OhgLq/Xdyat4=", "7TfDCaMWgEwSJEZd/e33b24919rF4IDJLmbqNNRNhZs=", "0HfdxDt/zGugZLIuHasrdLEW9s8OaTykSSqBEDFQXdw=", "SndbMKVtcTenAkwi2JBfGzD+mhexp1qJbRIY+A1JRIU=", "xH/DCseLHr9eKoYT8qsORZK7zVdEGYWHuVtsVrD95wY="], "checkpoint":{"envelope":"rekor.sigstore.dev - 1193050959916656506\n2492065618\nHMabR7yd/EVq4WbRn7sHFJlrkYwtTQFLVw8bfn+V/3U=\n\n— rekor.sigstore.dev wNI9ajBFAiEAndoa5ADD2zMzk/PVdEol6LUHALqSJ1aHkW2ZwNrNI5wCIGOxmepO1nyBeZT4ffX+mio5WU1k0ls6sfCVXThd1FB2\n"}}, "canonicalizedBody":"eyJhcGlWZXJzaW9uIjoiMC4wLjEiLCJraW5kIjoiaGFzaGVkcmVrb3JkIiwic3BlYyI6eyJkYXRhIjp7Imhhc2giOnsiYWxnb3JpdGhtIjoic2hhMjU2IiwidmFsdWUiOiJjYjNiM2YwMWMxMzA1OTQ2ZTBmMWRkMTQ3MjA5NDE5ODVlNTEyMWMxYzY1MTMzNDA2ZmIxNjdiNDdiOGRlYTIxIn19LCJzaWduYXR1cmUiOnsiY29udGVudCI6Ik1FVUNJUURZQkV4dHlHcm1nQXVMVEE2SVVtRmtHY0F1K2ZOd3VyTmN5SXZyYnNTc2xRSWdkRUxKcUNqQkowWGJ1UnpsQnQ0QVVNMU84NGhkTnNEWHNQa2t1MGlkZGlNPSIsInB1YmxpY0tleSI6eyJjb250ZW50IjoiTFMwdExTMUNSVWRKVGlCRFJWSlVTVVpKUTBGVVJTMHRMUzB0Q2sxSlNVaFVWRU5EUW5STFowRjNTVUpCWjBsVlpVWmhSWE5FWnpsVVRrOXJWRGhxVUdneFpVZzNXRmN4V2tjd2QwTm5XVWxMYjFwSmVtb3dSVUYzVFhjS1RucEZWazFDVFVkQk1WVkZRMmhOVFdNeWJHNWpNMUoyWTIxVmRWcEhWakpOVWpSM1NFRlpSRlpSVVVSRmVGWjZZVmRrZW1SSE9YbGFVekZ3WW01U2JBcGpiVEZzV2tkc2FHUkhWWGRJYUdOT1RXcFpkMDlFU1ROTlJHTjNUVlJOZWxkb1kwNU5hbGwzVDBSSk0wMUVZM2hOVkUxNlYycEJRVTFHYTNkRmQxbElDa3R2V2tsNmFqQkRRVkZaU1V0dldrbDZhakJFUVZGalJGRm5RVVV6VFhGc016WTNhRFZqTmpkV1FuaDJZbWhWYkhWeWFrODBWME55T0dJNGFFVndUR1FLTHpsak1VZDRXVFZVUzJoNGJTOVJXbUl6TDA1QmFXOWFTbkoxTjJ4c2FrUXZUSEUyTW1oaGNqTnpURXd3WjNveVNqWlBRMEptUlhkbloxaDBUVUUwUndwQk1WVmtSSGRGUWk5M1VVVkJkMGxJWjBSQlZFSm5UbFpJVTFWRlJFUkJTMEpuWjNKQ1owVkdRbEZqUkVGNlFXUkNaMDVXU0ZFMFJVWm5VVlZrWTJJeENrdExXV0ZaYmpZd2VGWnRkVTVOVGxKVlJrcDBUSE5SZDBoM1dVUldVakJxUWtKbmQwWnZRVlV6T1ZCd2VqRlphMFZhWWpWeFRtcHdTMFpYYVhocE5Ga0tXa1E0ZDJGM1dVUldVakJTUVZGSUwwSkhSWGRZTkZwa1lVaFNNR05JVFRaTWVUbHVZVmhTYjJSWFNYVlpNamwwVERKb01Wb3laSEJpYldSdFdWZE9iQXBNTW5Sc1kyMDFiR0pJVFhSWk1qbDBZbGhXZFdGWVVqVk1lVFZ1WVZoU2IyUlhTWFprTWpsNVlUSmFjMkl6WkhwTU1rb3hZVmQ0YTB4dWJHaGlWM2hCQ21OdFZtMWplVGx2V2xkR2EyTjVPWFJaVjJ4MVRVUnJSME5wYzBkQlVWRkNaemM0ZDBGUlJVVkxNbWd3WkVoQ2VrOXBPSFprUnpseVdsYzBkVmxYVGpBS1lWYzVkV041Tlc1aFdGSnZaRmRLTVdNeVZubFpNamwxWkVkV2RXUkROV3BpTWpCM1NIZFpTMHQzV1VKQ1FVZEVkbnBCUWtGblVWSmtNamw1WVRKYWN3cGlNMlJtV2tkc2VtTkhSakJaTW1kM1RtZFpTMHQzV1VKQ1FVZEVkbnBCUWtGM1VXOVpWMUpyV21wUk0wNUVSbTFaVjFwcFRsUldhVTVIVFhoTlJFWnBDazFVVlRWUFJGWnJXbTFHYkZwWFJUTlplbVJxVFcxWmVsbFVRVlJDWjI5eVFtZEZSVUZaVHk5TlFVVkZRa0ZXUTJSWGJITmFSRUZ5UW1kdmNrSm5SVVVLUVZsUEwwMUJSVVpDUWpGdlpGZGtibUZYTlc1YWJVWnFXbE01Y2xwWVNuVmFWM2g2VEZkT2RtSlhNVEZpYld3d1pWUkJaRUpuYjNKQ1owVkZRVmxQTHdwTlFVVkhRa0U1ZVZwWFducE1NbWhzV1ZkU2Vrd3lNV2hoVnpSM1QzZFpTMHQzV1VKQ1FVZEVkbnBCUWtOQlVYUkVRM1J2WkVoU2QyTjZiM1pNTTFKMkNtRXlWblZNYlVacVpFZHNkbUp1VFhWYU1td3dZVWhXYVdSWVRteGpiVTUyWW01U2JHSnVVWFZaTWpsMFRVY3dSME5wYzBkQlVWRkNaemM0ZDBGUmEwVUtXSGQ0WkdGSVVqQmpTRTAyVEhrNWJtRllVbTlrVjBsMVdUSTVkRXd5YURGYU1tUndZbTFrYlZsWFRteE1NblJzWTIwMWJHSklUWFJaTWpsMFlsaFdkUXBoV0ZJMVRIazFibUZZVW05a1YwbDJaREk1ZVdFeVduTmlNMlI2VERKS01XRlhlR3RNYm14b1lsZDRRV050Vm0xamVUbHZXbGRHYTJONU9YUlpWMngxQ2sxRVowZERhWE5IUVZGUlFtYzNPSGRCVVc5RlMyZDNiMWxYVW10YWFsRXpUa1JHYlZsWFdtbE9WRlpwVGtkTmVFMUVSbWxOVkZVMVQwUldhMXB0Um13S1dsZEZNMWw2WkdwTmJWbDZXVlJCWWtKbmIzSkNaMFZGUVZsUEwwMUJSVXhDUVRCTlF6Tk9iR0pIV1hSaFJ6bDZaRWRXYTAxRlFVZERhWE5IUVZGUlFncG5OemgzUVZGM1JVMW5kM2RoU0ZJd1kwaE5Oa3g1T1c1aFdGSnZaRmRKZFZreU9YUk1NbWd4V2pKa2NHSnRaRzFaVjA1c1RESjBiR050Tld4aVNFMTBDbGt5T1hSaVdGWjFZVmhTTlUxRVowZERhWE5IUVZGUlFtYzNPSGRCVVRCRlMyZDNiMWxYVW10YWFsRXpUa1JHYlZsWFdtbE9WRlpwVGtkTmVFMUVSbWtLVFZSVk5VOUVWbXRhYlVac1dsZEZNMWw2WkdwTmJWbDZXVlJCWmtKbmIzSkNaMFZGUVZsUEwwMUJSVTlDUWtWTlJETktiRnB1VFhaaFIxWm9Xa2hOZGdwaVYwWndZbXBCWVVKbmIzSkNaMFZGUVZsUEwwMUJSVkJDUVhkTlEycEZkMDU2UlRCT2VsVXhUV3ByZDB4bldVdExkMWxDUWtGSFJIWjZRVUpGUVZGbkNrUkNOVzlrU0ZKM1kzcHZka3d5WkhCa1IyZ3hXV2sxYW1JeU1IWmhTRlp1V2pKc2RWb3lXbWhaTWxWM1IwRlpTMHQzV1VKQ1FVZEVkbnBCUWtWUlVVc0tSRUZuZVU1VVkzbE5SR013VFhwQ2RFSm5iM0pDWjBWRlFWbFBMMDFCUlZOQ1JqaE5XRmRvTUdSSVFucFBhVGgyV2pKc01HRklWbWxNYlU1MllsTTVid3BrVjJSdVlWYzFibHB0Um1wYVV6bHlXbGhLZFZwWGVIcE1WMDUyWWxjeE1XSnRiREJsVXpoMVdqSnNNR0ZJVm1sTU0yUjJZMjEwYldKSE9UTmplVGxwQ21SWGJITmFRelUxV1ZjeGMxRklTbXhhYmsxMllVZFdhRnBJVFhaaVYwWndZbXBCTkVKbmIzSkNaMFZGUVZsUEwwMUJSVlJDUTI5TlMwZEdhMXBIV1RBS1RucFJlRnB0Um0xWmFsVXhXV3BTYWsxVVFYaFpha1V4VDFSbk1WcEhXbWhhVjFab1RqSk5NMWw2U20xTk1rVjNTVkZaUzB0M1dVSkNRVWRFZG5wQlFncEdRVkZVUkVKR00ySXpTbkphYlhoMlpERTVhMkZZVG5kWldGSnFZVVJDYTBKbmIzSkNaMFZGUVZsUEwwMUJSVlpDUmxsTlZrZG9NR1JJUW5wUGFUaDJDbG95YkRCaFNGWnBURzFPZG1KVE9XOWtWMlJ1WVZjMWJscHRSbXBhVXpseVdsaEtkVnBYZUhwTVYwNTJZbGN4TVdKdGJEQmxVemxvV1ROU2NHSXlOWG9LVEROS01XSnVUWFpOZWsxM1RrUlZlazlVVlRWT1ZFbDJXVmhTTUZwWE1YZGtTRTEyVFZSQlYwSm5iM0pDWjBWRlFWbFBMMDFCUlZkQ1FXZE5RbTVDTVFwWmJYaHdXWHBDUjBKbmIzSkNaMFZGUVZsUEwwMUJSVmxDUkdkTlRtNUtiR05IT0RaaFNGWnVXakpzZFZveVdtaFpNbFYyWVRKV2VXSnRWbk5qZVRGcUNtSXlNWFJrVnpWd1pFaHJObU50Vm0xUGJrcHNXbTVOZG1GSFZtaGFTRTEyWWxkR2NHSnFRMEpwZDFsTFMzZFpRa0pCU0ZkbFVVbEZRV2RTT1VKSWMwRUtaVkZDTTBGT01EbE5SM0pIZUhoRmVWbDRhMlZJU214dVRuZExhVk5zTmpRemFubDBMelJsUzJOdlFYWkxaVFpQUVVGQlFtOUZTVWRKTWxWQlFVRlJSQXBCUldkM1VtZEphRUZOUld3eFQzUlljbU1yY21kNlR6aEdNMFpETlc0ek5UTjFiSEEzVm1sUlpWWnJkaTk2WWt0NWVDdDZRV2xGUVhjeldXUjBRM2ROQ21ZeGRqaDNRMmx4U0VsclMzZFNPR04zUnpGb1MxWlJOazlWTWtNNFVHSXhaM2RaZDBObldVbExiMXBKZW1vd1JVRjNUVVJoVVVGM1dtZEplRUZPYTFFS2FFb3daSFE0VTA5VGJtRTJXRTExWjNkamJUSTNRek01ZGxod1IzaE5XbkJyYUd0SGVqbHFkRmd6TW5WTldYaFROMUJWTjJKUFVUaFZTeTlLT0dkSmVBcEJTazlZYXpJMlUwNUpORlEyWVcxV09WRmhhbVJyWlROVlZrMVBkMGxYUmt4NGFtOHhaWEUyWjJGcE1VTXJTVFZsZHpGaU0yOUtPWEExVFVJMU5rSjRDazFCUFQwS0xTMHRMUzFGVGtRZ1EwVlNWRWxHU1VOQlZFVXRMUzB0TFFvPSJ9fX19"}], "timestampVerificationData":{"rfc3161Timestamps":[{"signedTimestamp":"MIICyTADAgEAMIICwAYJKoZIhvcNAQcCoIICsTCCAq0CAQMxDTALBglghkgBZQMEAgEwgbgGCyqGSIb3DQEJEAEEoIGoBIGlMIGiAgEBBgkrBgEEAYO/MAIwMTANBglghkgBZQMEAgEFAAQgu68d7Mo1lxzqTLEcIIIYY2FmvZOIIiqmW+wOyc6iEOACFQDFjYG2FqqSK3niCsbatiEsgmH+QRgPMjAyNjA4MjcwNzAxMzNaMAMCAQGgMqQwMC4xFTATBgNVBAoTDHNpZ3N0b3JlLmRldjEVMBMGA1UEAxMMc2lnc3RvcmUtdHNhoAAxggHaMIIB1gIBATBRMDkxFTATBgNVBAoTDHNpZ3N0b3JlLmRldjEgMB4GA1UEAxMXc2lnc3RvcmUtdHNhLXNlbGZzaWduZWQCFDoTVC8MkGHuvMFDL8uKjosqI4sMMAsGCWCGSAFlAwQCAaCB/DAaBgkqhkiG9w0BCQMxDQYLKoZIhvcNAQkQAQQwHAYJKoZIhvcNAQkFMQ8XDTI2MDgyNzA3MDEzM1owLwYJKoZIhvcNAQkEMSIEINIMPF837ttfhQyR/UXYW4dIULHr9R9e+KH3gLevt0F8MIGOBgsqhkiG9w0BCRACLzF/MH0wezB5BCCF+Se8B6tiysO0Q1bBDvyBssaIP9p6uebYcNnROs0FtzBVMD2kOzA5MRUwEwYDVQQKEwxzaWdzdG9yZS5kZXYxIDAeBgNVBAMTF3NpZ3N0b3JlLXRzYS1zZWxmc2lnbmVkAhQ6E1QvDJBh7rzBQy/Lio6LKiOLDDAKBggqhkjOPQQDAgRmMGQCMGP+aJ7k1spyl0u7RBSD1p+hHTP2zGhDLvU3rSRMMvfeuXeAHLq07N1/ENr6rb8JCgIwYYxKHPJ+IzAw+GokkyFSJzNhjziLIL2bw8VZnZPFzgUgyaVTUZAnjp5UXpsOBaQX"}]}}, "messageSignature":{"messageDigest":{"algorithm":"SHA2_256", "digest":"yzs/AcEwWUbg8d0UcglBmF5RIcHGUTNAb7FntHuN6iE="}, "signature":"MEUCIQDYBExtyGrmgAuLTA6IUmFkGcAu+fNwurNcyIvrbsSslQIgdELJqCjBJ0XbuRzlBt4AUM1O84hdNsDXsPkku0iddiM="}}
|
build/torch213-cxx11-xpu20260-x86_64-linux/{_megablocks_xpu_c4d0dc5.abi3.so → _megablocks_xpu_addf474.abi3.so}
RENAMED
|
@@ -1,3 +1,3 @@
|
|
| 1 |
version https://git-lfs.github.com/spec/v1
|
| 2 |
-
oid sha256:
|
| 3 |
-
size
|
|
|
|
| 1 |
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:40d4c6b0feb85bd2e83ca7d619f3c6a0fb79f4da6f8b44f50f70da570de5b247
|
| 3 |
+
size 8288824
|
build/torch213-cxx11-xpu20260-x86_64-linux/_ops.py
CHANGED
|
@@ -1,9 +1,9 @@
|
|
| 1 |
import torch
|
| 2 |
-
from . import
|
| 3 |
-
ops = torch.ops.
|
| 4 |
|
| 5 |
def add_op_namespace_prefix(op_name: str):
|
| 6 |
"""
|
| 7 |
Prefix op by namespace.
|
| 8 |
"""
|
| 9 |
-
return f"
|
|
|
|
| 1 |
import torch
|
| 2 |
+
from . import _megablocks_xpu_addf474
|
| 3 |
+
ops = torch.ops._megablocks_xpu_addf474
|
| 4 |
|
| 5 |
def add_op_namespace_prefix(op_name: str):
|
| 6 |
"""
|
| 7 |
Prefix op by namespace.
|
| 8 |
"""
|
| 9 |
+
return f"_megablocks_xpu_addf474::{op_name}"
|
build/torch213-cxx11-xpu20260-x86_64-linux/megablocks/__init__.py
DELETED
|
@@ -1,26 +0,0 @@
|
|
| 1 |
-
import ctypes
|
| 2 |
-
import importlib.util
|
| 3 |
-
import sys
|
| 4 |
-
from pathlib import Path
|
| 5 |
-
from types import ModuleType
|
| 6 |
-
|
| 7 |
-
|
| 8 |
-
def _import_from_path(file_path: Path) -> ModuleType:
|
| 9 |
-
# We cannot use the module name as-is, after adding it to `sys.modules`,
|
| 10 |
-
# it would also be used for other imports. So, we make a module name that
|
| 11 |
-
# depends on the path for it to be unique using the hex-encoded hash of
|
| 12 |
-
# the path.
|
| 13 |
-
path_hash = "{:x}".format(ctypes.c_size_t(hash(file_path.absolute())).value)
|
| 14 |
-
module_name = path_hash
|
| 15 |
-
spec = importlib.util.spec_from_file_location(module_name, file_path)
|
| 16 |
-
if spec is None:
|
| 17 |
-
raise ImportError(f"Cannot load spec for {module_name} from {file_path}")
|
| 18 |
-
module = importlib.util.module_from_spec(spec)
|
| 19 |
-
if module is None:
|
| 20 |
-
raise ImportError(f"Cannot load module {module_name} from spec")
|
| 21 |
-
sys.modules[module_name] = module
|
| 22 |
-
spec.loader.exec_module(module) # type: ignore
|
| 23 |
-
return module
|
| 24 |
-
|
| 25 |
-
|
| 26 |
-
globals().update(vars(_import_from_path(Path(__file__).parent.parent / "__init__.py")))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
build/torch213-cxx11-xpu20260-x86_64-linux/metadata.json
CHANGED
|
@@ -1,9 +1,10 @@
|
|
| 1 |
{
|
| 2 |
"name": "megablocks",
|
| 3 |
-
"id": "
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"python-depends": [],
|
|
|
|
| 7 |
"backend": {
|
| 8 |
"type": "xpu"
|
| 9 |
},
|
|
@@ -38,8 +39,8 @@
|
|
| 38 |
"_layers/mpu.py": "BX61/pmqRCVeEMnzpuK8x5NpITBJyZfOawGnDegW1O8=",
|
| 39 |
"_layers/router.py": "GKC/5//Uf43HT6zXlBTStCMTml2YXn7Jbgq9EgTgDC0=",
|
| 40 |
"_layers/sharedexpert_registry.py": "GfwHKjbC7lBXYq9kNifp3HudTdZKSqwsFxGxFLc2krw=",
|
| 41 |
-
"
|
| 42 |
-
"_ops.py": "
|
| 43 |
"_version.py": "W3l9oUrnmvfBoIy06Lv0y2ezTQhNU04ug/txHcdWJD0=",
|
| 44 |
"backend/__init__.py": "4bunAqqjwE93nym8CigHbqVAal7UaYGdONb8lGRfrek=",
|
| 45 |
"backend/kernels.py": "2UFNTlmSDk3S5y6UW14MJMPVVsVhmozxQ5wFDxtZHhQ=",
|
|
@@ -51,7 +52,6 @@
|
|
| 51 |
"grouped_gemm/ops.py": "Kj7yRgv7afB5SVwXIkN9RNIF2kxD2qrZ0dVsKX5ph2s=",
|
| 52 |
"grouped_gemm_util.py": "bVPrdsEtBVJJ+v/bNrGNf/QHNlVL4VxHg92HvaL3oq4=",
|
| 53 |
"layers.py": "mNTlX2oKJyHBJllRd24lPNdGLa2H9cEUgVLvkZJpoNo=",
|
| 54 |
-
"megablocks/__init__.py": "DFYPlrhXwYjEqCl/8n0SmWGZV8NFml5DPhMjKfv98GY=",
|
| 55 |
"ops/__init__.py": "MHQsiMEgkKeZCa8TAYFMGDFnwFIgijfAN7DhMUt/48U=",
|
| 56 |
"ops/all_to_all_benchmark.py": "sxXSCSDdsh+62IWpKNTyBphLTngY1jlgJ1WgNooSEO8=",
|
| 57 |
"ops/binned_gather.py": "pkWne+ThfiTR/bMbDs2QDimKxAI4klqD1FdvhytE7lU=",
|
|
@@ -96,11 +96,11 @@
|
|
| 96 |
"provenance": {
|
| 97 |
"kernel-builder": {
|
| 98 |
"version": "0.17.0-dev0",
|
| 99 |
-
"
|
| 100 |
"dirty": false
|
| 101 |
},
|
| 102 |
"kernel": {
|
| 103 |
-
"
|
| 104 |
"dirty": false
|
| 105 |
}
|
| 106 |
}
|
|
|
|
| 1 |
{
|
| 2 |
"name": "megablocks",
|
| 3 |
+
"id": "_megablocks_xpu_addf474",
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"python-depends": [],
|
| 7 |
+
"kernel-depends": [],
|
| 8 |
"backend": {
|
| 9 |
"type": "xpu"
|
| 10 |
},
|
|
|
|
| 39 |
"_layers/mpu.py": "BX61/pmqRCVeEMnzpuK8x5NpITBJyZfOawGnDegW1O8=",
|
| 40 |
"_layers/router.py": "GKC/5//Uf43HT6zXlBTStCMTml2YXn7Jbgq9EgTgDC0=",
|
| 41 |
"_layers/sharedexpert_registry.py": "GfwHKjbC7lBXYq9kNifp3HudTdZKSqwsFxGxFLc2krw=",
|
| 42 |
+
"_megablocks_xpu_addf474.abi3.so": "QNTGsP64W9LoPKfWGfPGoPt59Npvi0T1D3DaVw3lskc=",
|
| 43 |
+
"_ops.py": "QBFLHcpmqB9+JkTkhdpkiEE9yER4FUDrsZgGzybNoVc=",
|
| 44 |
"_version.py": "W3l9oUrnmvfBoIy06Lv0y2ezTQhNU04ug/txHcdWJD0=",
|
| 45 |
"backend/__init__.py": "4bunAqqjwE93nym8CigHbqVAal7UaYGdONb8lGRfrek=",
|
| 46 |
"backend/kernels.py": "2UFNTlmSDk3S5y6UW14MJMPVVsVhmozxQ5wFDxtZHhQ=",
|
|
|
|
| 52 |
"grouped_gemm/ops.py": "Kj7yRgv7afB5SVwXIkN9RNIF2kxD2qrZ0dVsKX5ph2s=",
|
| 53 |
"grouped_gemm_util.py": "bVPrdsEtBVJJ+v/bNrGNf/QHNlVL4VxHg92HvaL3oq4=",
|
| 54 |
"layers.py": "mNTlX2oKJyHBJllRd24lPNdGLa2H9cEUgVLvkZJpoNo=",
|
|
|
|
| 55 |
"ops/__init__.py": "MHQsiMEgkKeZCa8TAYFMGDFnwFIgijfAN7DhMUt/48U=",
|
| 56 |
"ops/all_to_all_benchmark.py": "sxXSCSDdsh+62IWpKNTyBphLTngY1jlgJ1WgNooSEO8=",
|
| 57 |
"ops/binned_gather.py": "pkWne+ThfiTR/bMbDs2QDimKxAI4klqD1FdvhytE7lU=",
|
|
|
|
| 96 |
"provenance": {
|
| 97 |
"kernel-builder": {
|
| 98 |
"version": "0.17.0-dev0",
|
| 99 |
+
"commit": "7d2828e6592de641b44fb8eb896719101a2ab2fb",
|
| 100 |
"dirty": false
|
| 101 |
},
|
| 102 |
"kernel": {
|
| 103 |
+
"commit": "addf4741fafb55b4c101b15985dfaeea7c7c2f3a",
|
| 104 |
"dirty": false
|
| 105 |
}
|
| 106 |
}
|
build/torch213-cxx11-xpu20260-x86_64-linux/metadata.json.sigstore
CHANGED
|
@@ -1 +1 @@
|
|
| 1 |
-
{"mediaType":"application/vnd.dev.sigstore.bundle.v0.3+json", "verificationMaterial":{"certificate":{"rawBytes":"MIIHSTCCBtCgAwIBAgIUEsfCDaGiIjMTh2fWbFNDO7MtkCowCgYIKoZIzj0EAwMwNzEVMBMGA1UEChMMc2lnc3RvcmUuZGV2MR4wHAYDVQQDExVzaWdzdG9yZS1pbnRlcm1lZGlhdGUwHhcNMjYwODI1MTIwNzQ3WhcNMjYwODI1MTIxNzQ3WjAAMFkwEwYHKoZIzj0CAQYIKoZIzj0DAQcDQgAEmPdHGofC4sgawEo7z3SJT/Ayt9ERh4bgxvOL7buMV+eCrAw2gmH3lYDRyAHRdCkSP2KQuzjQAsjJ19LZCqzw6KOCBe8wggXrMA4GA1UdDwEB/wQEAwIHgDATBgNVHSUEDDAKBggrBgEFBQcDAzAdBgNVHQ4EFgQUPsC5HZTI3gn0jq+BQjq4r5S2OgswHwYDVR0jBBgwFoAU39Ppz1YkEZb5qNjpKFWixi4YZD8wawYDVR0RAQH/BGEwX4ZdaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5Ly5naXRodWIvd29ya2Zsb3dzL2J1aWxkLnlhbWxAcmVmcy9oZWFkcy9tYWluMDkGCisGAQQBg78wAQEEK2h0dHBzOi8vdG9rZW4uYWN0aW9ucy5naXRodWJ1c2VyY29udGVudC5jb20wHwYKKwYBBAGDvzABAgQRd29ya2Zsb3dfZGlzcGF0Y2gwNgYKKwYBBAGDvzABAwQoYzRkMGRjNTY2NjFiYWRiOWNjNDdiZGU2NDdlMDhmZTNiNGJmNzQ1ODATBgorBgEEAYO/MAEEBAVCdWlsZDArBgorBgEEAYO/MAEFBB1odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eTAdBgorBgEEAYO/MAEGBA9yZWZzL2hlYWRzL21haW4wOwYKKwYBBAGDvzABCAQtDCtodHRwczovL3Rva2VuLmFjdGlvbnMuZ2l0aHVidXNlcmNvbnRlbnQuY29tMG0GCisGAQQBg78wAQkEXwxdaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5Ly5naXRodWIvd29ya2Zsb3dzL2J1aWxkLnlhbWxAcmVmcy9oZWFkcy9tYWluMDgGCisGAQQBg78wAQoEKgwoYzRkMGRjNTY2NjFiYWRiOWNjNDdiZGU2NDdlMDhmZTNiNGJmNzQ1ODAbBgorBgEEAYO/MAELBA0MC3NlbGYtaG9zdGVkMEAGCisGAQQBg78wAQwEMgwwaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5MDgGCisGAQQBg78wAQ0EKgwoYzRkMGRjNTY2NjFiYWRiOWNjNDdiZGU2NDdlMDhmZTNiNGJmNzQ1ODAfBgorBgEEAYO/MAEOBBEMD3JlZnMvaGVhZHMvbWFpbjAaBgorBgEEAYO/MAEPBAwMCjEwNzE0NzU1MjkwLgYKKwYBBAGDvzABEAQgDB5odHRwczovL2dpdGh1Yi5jb20vaHVnZ2luZ2ZhY2UwGAYKKwYBBAGDvzABEQQKDAgyNTcyMDc0MzBtBgorBgEEAYO/MAESBF8MXWh0dHBzOi8vZ2l0aHViLmNvbS9odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eS8uZ2l0aHViL3dvcmtmbG93cy9idWlsZC55YW1sQHJlZnMvaGVhZHMvbWFpbjA4BgorBgEEAYO/MAETBCoMKGM0ZDBkYzU2NjYxYmFkYjljYzQ3YmRlNjQ3ZTA4ZmUzYjRiZjc0NTgwIQYKKwYBBAGDvzABFAQTDBF3b3JrZmxvd19kaXNwYXRjaDBkBgorBgEEAYO/MAEVBFYMVGh0dHBzOi8vZ2l0aHViLmNvbS9odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eS9hY3Rpb25zL3J1bnMvMzI4NDQ2ODk4OTQvYXR0ZW1wdHMvMTAWBgorBgEEAYO/MAEWBAgMBnB1YmxpYzBGBgorBgEEAYO/MAEYBDgMNnJlcG86aHVnZ2luZ2ZhY2Uva2VybmVscy1jb21tdW5pdHk6cmVmOnJlZnMvaGVhZHMvbWFpbjCBiQYKKwYBBAHWeQIEAgR7BHkAdwB1AN09MGrGxxEyYxkeHJlnNwKiSl643jyt/4eKcoAvKe6OAAABoDjRyaUAAAQDAEYwRAIgM0GM4ReR0pFWw7Wx/buzmiWJqvaZoMT/xfGlpLZRm1cCIGUMb+0iRQuknEKoQbK+5DveMo3AC6jEpy4nLUXlBadNMAoGCCqGSM49BAMDA2cAMGQCMFodrf0jX0YBZhtrHolQlOygwpaTKfIqEd1KsAEBGzEeyglfDdeCtnlJqyjC2QabTgIwMSJ9rAuw/7JMGuRwKX7SN6y10dj6bcgrWH1h1pMOpRkDviO3UcOW6D7YK0QJh6sI"}, "tlogEntries":[{"logIndex":"2583274400", "logId":{"keyId":"wNI9atQGlz+VWfO6LRygH4QUfY/8W4RFwiT5i5WRgB0="}, "kindVersion":{"kind":"hashedrekord", "version":"0.0.1"}, "integratedTime":"1787659668", "inclusionPromise":{"signedEntryTimestamp":"MEYCIQCqfQPWi1/v5rUMFeKZU9k687Y42OyxhJoWYlADuFTFUwIhAIRq8YTYK2Y2WqGWhf4mVDyufmMqpoud0k4fLVFBgMVq"}, "inclusionProof":{"logIndex":"2461370138", "rootHash":"sKa21DCpYU/7LZ9ORZD9C6zjkbDLdRZrxb0LEXD2oF8=", "treeSize":"2461370140", "hashes":["56CyIHBhqSmA3l4hOHXc32F/TItpTDPvieoQQ97VSao=", "y6SajW2RSxSC+qZEPkVTOda5uUFbFye//C1W3E4fOD0=", "cTavF/Ocd3ZeX1e1jnEV1te21vU3vvj+JGzwC/gStgk=", "yW+Ys91jOHBbqiKaKMiTKMYmn8iwJam2ZsB6CNDrzzs=", "BJlxWyo+0U2jxgcimrAeK1pey6FSpHWLOetp0UXRWI4=", "JsX9xGIEwFv7MFelTobmTLz4j1Jl78aSs2G7dkybw0w=", "nwobU/9yRWhsbw80xq3x/uNV1nniMWdUKAqc3UWYsJM=", "EXKrSbeG2qm5APD5QSDN66X7UHz9OSofiW9PyNU7fPM=", "cjIeeutzThxIPTlFFHJqa7bAzn3k2FzK8Nrdd56OWvE=", "+pNXWGIXzU1dqSy/aePHKYtSkDnFinI+/bL5HnurJ6w=", "G1R9F5B2KizFA2NyhDznjNKXPOF1SDmbGHz+omLG7TY=", "RZ7XwGuqMr5gDUh2HpUcaXp+AmIiEcPJiMui4GLJE7o=", "mBm+vQtn0C4thMnlTnxfxqowq1dXsPBCKaJ85da3JeU=", "b80/J/68RC8/tx5xRdlKk6MmTpZiVRotquWBrE/z1pI=", "SndbMKVtcTenAkwi2JBfGzD+mhexp1qJbRIY+A1JRIU=", "xH/DCseLHr9eKoYT8qsORZK7zVdEGYWHuVtsVrD95wY="], "checkpoint":{"envelope":"rekor.sigstore.dev - 1193050959916656506\n2461370140\nsKa21DCpYU/7LZ9ORZD9C6zjkbDLdRZrxb0LEXD2oF8=\n\n— rekor.sigstore.dev wNI9ajBFAiBoocRiTqMBmVriZO1sD/o5GwCtMgDY+WHYB3j+0fWVLAIhAPgdJas1OQKZUc4PR5zJpIrE5p3Emu3DDxDwPi16ZReL\n"}}, "canonicalizedBody":"eyJhcGlWZXJzaW9uIjoiMC4wLjEiLCJraW5kIjoiaGFzaGVkcmVrb3JkIiwic3BlYyI6eyJkYXRhIjp7Imhhc2giOnsiYWxnb3JpdGhtIjoic2hhMjU2IiwidmFsdWUiOiI3MzUwOTg3YTRmYjA2ODdiNDYxY2EzMjJlM2NiYzVhZmY3Y2VhZGQ4ZjBkYzcyZDQ0MmU3OTQ0MzRkNzIxYjc2In19LCJzaWduYXR1cmUiOnsiY29udGVudCI6Ik1FVUNJR09BTGdFYjdMaEJsZ09CSmpZVkhwMUk2bGZIVEZ6alZjVzdJTG9mL2U0cEFpRUFueks1RGZibGF2KzVuYVdseHV1U2lxUWJ0QmN5aHE5cFhvTE8wb0t0eHVRPSIsInB1YmxpY0tleSI6eyJjb250ZW50IjoiTFMwdExTMUNSVWRKVGlCRFJWSlVTVVpKUTBGVVJTMHRMUzB0Q2sxSlNVaFRWRU5EUW5SRFowRjNTVUpCWjBsVlJYTm1RMFJoUjJsSmFrMVVhREptVjJKR1RrUlBOMDEwYTBOdmQwTm5XVWxMYjFwSmVtb3dSVUYzVFhjS1RucEZWazFDVFVkQk1WVkZRMmhOVFdNeWJHNWpNMUoyWTIxVmRWcEhWakpOVWpSM1NFRlpSRlpSVVVSRmVGWjZZVmRrZW1SSE9YbGFVekZ3WW01U2JBcGpiVEZzV2tkc2FHUkhWWGRJYUdOT1RXcFpkMDlFU1RGTlZFbDNUbnBSTTFkb1kwNU5hbGwzVDBSSk1VMVVTWGhPZWxFelYycEJRVTFHYTNkRmQxbElDa3R2V2tsNmFqQkRRVkZaU1V0dldrbDZhakJFUVZGalJGRm5RVVZ0VUdSSVIyOW1RelJ6WjJGM1JXODNlak5UU2xRdlFYbDBPVVZTYURSaVozaDJUMHdLTjJKMVRWWXJaVU55UVhjeVoyMUlNMnhaUkZKNVFVaFNaRU5yVTFBeVMxRjFlbXBSUVhOcVNqRTVURnBEY1hwM05rdFBRMEpsT0hkbloxaHlUVUUwUndwQk1WVmtSSGRGUWk5M1VVVkJkMGxJWjBSQlZFSm5UbFpJVTFWRlJFUkJTMEpuWjNKQ1owVkdRbEZqUkVGNlFXUkNaMDVXU0ZFMFJVWm5VVlZRYzBNMUNraGFWRWt6WjI0d2FuRXJRbEZxY1RSeU5WTXlUMmR6ZDBoM1dVUldVakJxUWtKbmQwWnZRVlV6T1ZCd2VqRlphMFZhWWpWeFRtcHdTMFpYYVhocE5Ga0tXa1E0ZDJGM1dVUldVakJTUVZGSUwwSkhSWGRZTkZwa1lVaFNNR05JVFRaTWVUbHVZVmhTYjJSWFNYVlpNamwwVERKb01Wb3laSEJpYldSdFdWZE9iQXBNTW5Sc1kyMDFiR0pJVFhSWk1qbDBZbGhXZFdGWVVqVk1lVFZ1WVZoU2IyUlhTWFprTWpsNVlUSmFjMkl6WkhwTU1rb3hZVmQ0YTB4dWJHaGlWM2hCQ21OdFZtMWplVGx2V2xkR2EyTjVPWFJaVjJ4MVRVUnJSME5wYzBkQlVWRkNaemM0ZDBGUlJVVkxNbWd3WkVoQ2VrOXBPSFprUnpseVdsYzBkVmxYVGpBS1lWYzVkV041Tlc1aFdGSnZaRmRLTVdNeVZubFpNamwxWkVkV2RXUkROV3BpTWpCM1NIZFpTMHQzV1VKQ1FVZEVkbnBCUWtGblVWSmtNamw1WVRKYWN3cGlNMlJtV2tkc2VtTkhSakJaTW1kM1RtZFpTMHQzV1VKQ1FVZEVkbnBCUWtGM1VXOVplbEpyVFVkU2FrNVVXVEpPYWtacFdWZFNhVTlYVG1wT1JHUnBDbHBIVlRKT1JHUnNUVVJvYlZwVVRtbE9SMHB0VG5wUk1VOUVRVlJDWjI5eVFtZEZSVUZaVHk5TlFVVkZRa0ZXUTJSWGJITmFSRUZ5UW1kdmNrSm5SVVVLUVZsUEwwMUJSVVpDUWpGdlpGZGtibUZYTlc1YWJVWnFXbE01Y2xwWVNuVmFWM2g2VEZkT2RtSlhNVEZpYld3d1pWUkJaRUpuYjNKQ1owVkZRVmxQTHdwTlFVVkhRa0U1ZVZwWFducE1NbWhzV1ZkU2Vrd3lNV2hoVnpSM1QzZFpTMHQzV1VKQ1FVZEVkbnBCUWtOQlVYUkVRM1J2WkVoU2QyTjZiM1pNTTFKMkNtRXlWblZNYlVacVpFZHNkbUp1VFhWYU1td3dZVWhXYVdSWVRteGpiVTUyWW01U2JHSnVVWFZaTWpsMFRVY3dSME5wYzBkQlVWRkNaemM0ZDBGUmEwVUtXSGQ0WkdGSVVqQmpTRTAyVEhrNWJtRllVbTlrVjBsMVdUSTVkRXd5YURGYU1tUndZbTFrYlZsWFRteE1NblJzWTIwMWJHSklUWFJaTWpsMFlsaFdkUXBoV0ZJMVRIazFibUZZVW05a1YwbDJaREk1ZVdFeVduTmlNMlI2VERKS01XRlhlR3RNYm14b1lsZDRRV050Vm0xamVUbHZXbGRHYTJONU9YUlpWMngxQ2sxRVowZERhWE5IUVZGUlFtYzNPSGRCVVc5RlMyZDNiMWw2VW10TlIxSnFUbFJaTWs1cVJtbFpWMUpwVDFkT2FrNUVaR2xhUjFVeVRrUmtiRTFFYUcwS1dsUk9hVTVIU20xT2VsRXhUMFJCWWtKbmIzSkNaMFZGUVZsUEwwMUJSVXhDUVRCTlF6Tk9iR0pIV1hSaFJ6bDZaRWRXYTAxRlFVZERhWE5IUVZGUlFncG5OemgzUVZGM1JVMW5kM2RoU0ZJd1kwaE5Oa3g1T1c1aFdGSnZaRmRKZFZreU9YUk1NbWd4V2pKa2NHSnRaRzFaVjA1c1RESjBiR050Tld4aVNFMTBDbGt5T1hSaVdGWjFZVmhTTlUxRVowZERhWE5IUVZGUlFtYzNPSGRCVVRCRlMyZDNiMWw2VW10TlIxSnFUbFJaTWs1cVJtbFpWMUpwVDFkT2FrNUVaR2tLV2tkVk1rNUVaR3hOUkdodFdsUk9hVTVIU20xT2VsRXhUMFJCWmtKbmIzSkNaMFZGUVZsUEwwMUJSVTlDUWtWTlJETktiRnB1VFhaaFIxWm9Xa2hOZGdwaVYwWndZbXBCWVVKbmIzSkNaMFZGUVZsUEwwMUJSVkJDUVhkTlEycEZkMDU2UlRCT2VsVXhUV3ByZDB4bldVdExkMWxDUWtGSFJIWjZRVUpGUVZGbkNrUkNOVzlrU0ZKM1kzcHZka3d5WkhCa1IyZ3hXV2sxYW1JeU1IWmhTRlp1V2pKc2RWb3lXbWhaTWxWM1IwRlpTMHQzV1VKQ1FVZEVkbnBCUWtWUlVVc0tSRUZuZVU1VVkzbE5SR013VFhwQ2RFSm5iM0pDWjBWRlFWbFBMMDFCUlZOQ1JqaE5XRmRvTUdSSVFucFBhVGgyV2pKc01HRklWbWxNYlU1MllsTTVid3BrVjJSdVlWYzFibHB0Um1wYVV6bHlXbGhLZFZwWGVIcE1WMDUyWWxjeE1XSnRiREJsVXpoMVdqSnNNR0ZJVm1sTU0yUjJZMjEwYldKSE9UTmplVGxwQ21SWGJITmFRelUxV1ZjeGMxRklTbXhhYmsxMllVZFdhRnBJVFhaaVYwWndZbXBCTkVKbmIzSkNaMFZGUVZsUEwwMUJSVlJDUTI5TlMwZE5NRnBFUW1zS1dYcFZNazVxV1hoWmJVWnJXV3BzYWxsNlVUTlpiVkpzVG1wUk0xcFVRVFJhYlZWNldXcFNhVnBxWXpCT1ZHZDNTVkZaUzB0M1dVSkNRVWRFZG5wQlFncEdRVkZVUkVKR00ySXpTbkphYlhoMlpERTVhMkZZVG5kWldGSnFZVVJDYTBKbmIzSkNaMFZGUVZsUEwwMUJSVlpDUmxsTlZrZG9NR1JJUW5wUGFUaDJDbG95YkRCaFNGWnBURzFPZG1KVE9XOWtWMlJ1WVZjMWJscHRSbXBhVXpseVdsaEtkVnBYZUhwTVYwNTJZbGN4TVdKdGJEQmxVemxvV1ROU2NHSXlOWG9LVEROS01XSnVUWFpOZWtrMFRrUlJNazlFYXpSUFZGRjJXVmhTTUZwWE1YZGtTRTEyVFZSQlYwSm5iM0pDWjBWRlFWbFBMMDFCUlZkQ1FXZE5RbTVDTVFwWmJYaHdXWHBDUjBKbmIzSkNaMFZGUVZsUEwwMUJSVmxDUkdkTlRtNUtiR05IT0RaaFNGWnVXakpzZFZveVdtaFpNbFYyWVRKV2VXSnRWbk5qZVRGcUNtSXlNWFJrVnpWd1pFaHJObU50Vm0xUGJrcHNXbTVOZG1GSFZtaGFTRTEyWWxkR2NHSnFRMEpwVVZsTFMzZFpRa0pCU0ZkbFVVbEZRV2RTTjBKSWEwRUtaSGRDTVVGT01EbE5SM0pIZUhoRmVWbDRhMlZJU214dVRuZExhVk5zTmpRemFubDBMelJsUzJOdlFYWkxaVFpQUVVGQlFtOUVhbEo1WVZWQlFVRlJSQXBCUlZsM1VrRkpaMDB3UjAwMFVtVlNNSEJHVjNjM1YzZ3ZZblY2YldsWFNuRjJZVnB2VFZRdmVHWkhiSEJNV2xKdE1XTkRTVWRWVFdJck1HbFNVWFZyQ201RlMyOVJZa3NyTlVSMlpVMXZNMEZETm1wRmNIazBia3hWV0d4Q1lXUk9UVUZ2UjBORGNVZFRUVFE1UWtGTlJFRXlZMEZOUjFGRFRVWnZaSEptTUdvS1dEQlpRbHBvZEhKSWIyeFJiRTk1WjNkd1lWUkxaa2x4UldReFMzTkJSVUpIZWtWbGVXZHNaa1JrWlVOMGJteEtjWGxxUXpKUllXSlVaMGwzVFZOS09RcHlRWFYzTHpkS1RVZDFVbmRMV0RkVFRqWjVNVEJrYWpaaVkyZHlWMGd4YURGd1RVOXdVbXRFZG1sUE0xVmpUMWMyUkRkWlN6QlJTbWcyYzBrS0xTMHRMUzFGVGtRZ1EwVlNWRWxHU1VOQlZFVXRMUzB0TFFvPSJ9fX19"}], "timestampVerificationData":{"rfc3161Timestamps":[{"signedTimestamp":"MIICyDADAgEAMIICvwYJKoZIhvcNAQcCoIICsDCCAqwCAQMxDTALBglghkgBZQMEAgEwgbcGCyqGSIb3DQEJEAEEoIGnBIGkMIGhAgEBBgkrBgEEAYO/MAIwMTANBglghkgBZQMEAgEFAAQgRIyiiGiz024oUd3PbWjmpU4Y6+JvU70/lSZrsvLelY0CFF9e0UyujyckpCkQ8/3pRK36GnSgGA8yMDI2MDgyNTEyMDc0N1owAwIBAaAypDAwLjEVMBMGA1UEChMMc2lnc3RvcmUuZGV2MRUwEwYDVQQDEwxzaWdzdG9yZS10c2GgADGCAdowggHWAgEBMFEwOTEVMBMGA1UEChMMc2lnc3RvcmUuZGV2MSAwHgYDVQQDExdzaWdzdG9yZS10c2Etc2VsZnNpZ25lZAIUOhNULwyQYe68wUMvy4qOiyojiwwwCwYJYIZIAWUDBAIBoIH8MBoGCSqGSIb3DQEJAzENBgsqhkiG9w0BCRABBDAcBgkqhkiG9w0BCQUxDxcNMjYwODI1MTIwNzQ3WjAvBgkqhkiG9w0BCQQxIgQgakL2LLBSBy+oQUb+Pe25PcygEP3fQJZCV+Gtb6iAf3IwgY4GCyqGSIb3DQEJEAIvMX8wfTB7MHkEIIX5J7wHq2LKw7RDVsEO/IGyxog/2nq55thw2dE6zQW3MFUwPaQ7MDkxFTATBgNVBAoTDHNpZ3N0b3JlLmRldjEgMB4GA1UEAxMXc2lnc3RvcmUtdHNhLXNlbGZzaWduZWQCFDoTVC8MkGHuvMFDL8uKjosqI4sMMAoGCCqGSM49BAMCBGYwZAIwMkmYobhjvzAoBXodVkDPISEvVJD2D8VPMKig4/BXuTpd0ZUS22yF79X5tHAMOkNmAjB/7sdz/RW1HEIQ3AN1fMQwpJY7iKj5JDvrcHaXrQJZczm6BnMfClNitGhQzvaOzDY="}]}}, "messageSignature":{"messageDigest":{"algorithm":"SHA2_256", "digest":"c1CYek+waHtGHKMi48vFr/fOrdjw3HLUQueUQ01yG3Y="}, "signature":"MEUCIGOALgEb7LhBlgOBJjYVHp1I6lfHTFzjVcW7ILof/e4pAiEAnzK5Dfblav+5naWlxuuSiqQbtBcyhq9pXoLO0oKtxuQ="}}
|
|
|
|
| 1 |
+
{"mediaType":"application/vnd.dev.sigstore.bundle.v0.3+json", "verificationMaterial":{"certificate":{"rawBytes":"MIIHSzCCBtGgAwIBAgIUS6NAYvKLaslno4afvbnIQguXVIwwCgYIKoZIzj0EAwMwNzEVMBMGA1UEChMMc2lnc3RvcmUuZGV2MR4wHAYDVQQDExVzaWdzdG9yZS1pbnRlcm1lZGlhdGUwHhcNMjYwODI3MDcwMTMwWhcNMjYwODI3MDcxMTMwWjAAMFkwEwYHKoZIzj0CAQYIKoZIzj0DAQcDQgAEkkjUAxP5KcArfTKbAVpacWLLJex7TpKZBpVwxwfTsp7a5297XvEVrmqpJL9dUaB0UbqikK8EdFFibFlh0pa8ZqOCBfAwggXsMA4GA1UdDwEB/wQEAwIHgDATBgNVHSUEDDAKBggrBgEFBQcDAzAdBgNVHQ4EFgQUc9wV0QOoF2EImIrrC3wic3v6GkQwHwYDVR0jBBgwFoAU39Ppz1YkEZb5qNjpKFWixi4YZD8wawYDVR0RAQH/BGEwX4ZdaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5Ly5naXRodWIvd29ya2Zsb3dzL2J1aWxkLnlhbWxAcmVmcy9oZWFkcy9tYWluMDkGCisGAQQBg78wAQEEK2h0dHBzOi8vdG9rZW4uYWN0aW9ucy5naXRodWJ1c2VyY29udGVudC5jb20wHwYKKwYBBAGDvzABAgQRd29ya2Zsb3dfZGlzcGF0Y2gwNgYKKwYBBAGDvzABAwQoYWRkZjQ3NDFmYWZiNTViNGMxMDFiMTU5ODVkZmFlZWE3YzdjMmYzYTATBgorBgEEAYO/MAEEBAVCdWlsZDArBgorBgEEAYO/MAEFBB1odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eTAdBgorBgEEAYO/MAEGBA9yZWZzL2hlYWRzL21haW4wOwYKKwYBBAGDvzABCAQtDCtodHRwczovL3Rva2VuLmFjdGlvbnMuZ2l0aHVidXNlcmNvbnRlbnQuY29tMG0GCisGAQQBg78wAQkEXwxdaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5Ly5naXRodWIvd29ya2Zsb3dzL2J1aWxkLnlhbWxAcmVmcy9oZWFkcy9tYWluMDgGCisGAQQBg78wAQoEKgwoYWRkZjQ3NDFmYWZiNTViNGMxMDFiMTU5ODVkZmFlZWE3YzdjMmYzYTAbBgorBgEEAYO/MAELBA0MC3NlbGYtaG9zdGVkMEAGCisGAQQBg78wAQwEMgwwaHR0cHM6Ly9naXRodWIuY29tL2h1Z2dpbmdmYWNlL2tlcm5lbHMtY29tbXVuaXR5MDgGCisGAQQBg78wAQ0EKgwoYWRkZjQ3NDFmYWZiNTViNGMxMDFiMTU5ODVkZmFlZWE3YzdjMmYzYTAfBgorBgEEAYO/MAEOBBEMD3JlZnMvaGVhZHMvbWFpbjAaBgorBgEEAYO/MAEPBAwMCjEwNzE0NzU1MjkwLgYKKwYBBAGDvzABEAQgDB5odHRwczovL2dpdGh1Yi5jb20vaHVnZ2luZ2ZhY2UwGAYKKwYBBAGDvzABEQQKDAgyNTcyMDc0MzBtBgorBgEEAYO/MAESBF8MXWh0dHBzOi8vZ2l0aHViLmNvbS9odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eS8uZ2l0aHViL3dvcmtmbG93cy9idWlsZC55YW1sQHJlZnMvaGVhZHMvbWFpbjA4BgorBgEEAYO/MAETBCoMKGFkZGY0NzQxZmFmYjU1YjRjMTAxYjE1OTg1ZGZhZWVhN2M3YzJmM2EwIQYKKwYBBAGDvzABFAQTDBF3b3JrZmxvd19kaXNwYXRjaDBkBgorBgEEAYO/MAEVBFYMVGh0dHBzOi8vZ2l0aHViLmNvbS9odWdnaW5nZmFjZS9rZXJuZWxzLWNvbW11bml0eS9hY3Rpb25zL3J1bnMvMzMwNDUzOTU5NTIvYXR0ZW1wdHMvMTAWBgorBgEEAYO/MAEWBAgMBnB1YmxpYzBGBgorBgEEAYO/MAEYBDgMNnJlcG86aHVnZ2luZ2ZhY2Uva2VybmVscy1jb21tdW5pdHk6cmVmOnJlZnMvaGVhZHMvbWFpbjCBigYKKwYBBAHWeQIEAgR8BHoAeAB2AN09MGrGxxEyYxkeHJlnNwKiSl643jyt/4eKcoAvKe6OAAABoEIGFjUAAAQDAEcwRQIhANTRF20NOmZePyHcaLZoB+Ynh9gmeXuMkalMRZCeOecnAiBxjNS1Z2+jOeWRa6XIk3XrqnK/Iq41yIUIcglNc0MpEjAKBggqhkjOPQQDAwNoADBlAjAgugwsZ4aytQwxr/7pho7GVZocMS9UWugVQfh7RRrrQVwoNVxR0XBw0DD+SihdQKcCMQDVq/kD8ugfS/ELUyJ5sRiCGAQBNNocsFu4xUsJoY0p2XigO38hOTbnf8urjL0MOas="}, "tlogEntries":[{"logIndex":"2613968932", "logId":{"keyId":"wNI9atQGlz+VWfO6LRygH4QUfY/8W4RFwiT5i5WRgB0="}, "kindVersion":{"kind":"hashedrekord", "version":"0.0.1"}, "integratedTime":"1787814090", "inclusionPromise":{"signedEntryTimestamp":"MEQCIGkMZUpX0r1RdUT/qCpwNzsPBREIVwYwNNL0jyh3nV4qAiBfpetY+w32PkqAfhxX3vTpFJ6TAceJVyhQfX+QPt1uoQ=="}, "inclusionProof":{"logIndex":"2492064670", "rootHash":"1h92M2dTLnvbT7MAA2OgFilyI0+Wvly08dno+84I5lE=", "treeSize":"2492064684", "hashes":["CWmlwIKdRbEbSDeboUjozDIsx0bqOB/WEfIayAqx5YU=", "6LLc+NRuoFNMAhwGV746EP4qP91BNNMTnykCearhUc0=", "H5GPRpgXkD94jzZLD2lG1o2dNa+Dp4wZ3x3quVORDxk=", "vY0nAvZzZfKX0Pm0YTeJ2Nr7dUioCo09gqKzHqFm9O0=", "Lq+5cItalyvqW82mU5fvHUlud2o6lVXhO9TYKPjr+4A=", "2o4EZw1FLXLVQFKDPjPjvYsApHjoG1wq9ybfzyk7Kkw=", "+12tCov0Gm6R26NJWDl9P7e7QCoq/wf0/Szk2i4m4ZA=", "8den7EkYRWGnAu0x2oUyxgALn1Q7Qk2YTFObIa7zPEI=", "w0py6ki7Xp/BdPISxRvuHaSG06mdZqtk6VM+YYIIbEA=", "UbnDpZmg3ac4T8gCmplTGV9yU5j3fWOpeZi+i3apFqM=", "ygNYNZXNiY/+rjtW6HI0ngBM75qjGCTD6YstKiLiOZY=", "RvtgIftU0w4e16LAcGD7A1ErF9x2p/QLAyDLR+PaDMU=", "kW2yseLcah3rCzJJ1VS5Uhr9NCuHTJMDHsUkYy/0dV0=", "VGrFt5JvlIUQvoGEdBBxrCe3lYbu581OhgLq/Xdyat4=", "7TfDCaMWgEwSJEZd/e33b24919rF4IDJLmbqNNRNhZs=", "0HfdxDt/zGugZLIuHasrdLEW9s8OaTykSSqBEDFQXdw=", "SndbMKVtcTenAkwi2JBfGzD+mhexp1qJbRIY+A1JRIU=", "xH/DCseLHr9eKoYT8qsORZK7zVdEGYWHuVtsVrD95wY="], "checkpoint":{"envelope":"rekor.sigstore.dev - 1193050959916656506\n2492064684\n1h92M2dTLnvbT7MAA2OgFilyI0+Wvly08dno+84I5lE=\n\n— rekor.sigstore.dev wNI9ajBFAiEA8LYuNM1a+J0KgCyc6r2RNOn5HNRKCJeErkpmjrlweCgCICrwiO3l3JOK78OZTo/eGOzQwczFgUgoKhqut6MKY+a6\n"}}, "canonicalizedBody":"eyJhcGlWZXJzaW9uIjoiMC4wLjEiLCJraW5kIjoiaGFzaGVkcmVrb3JkIiwic3BlYyI6eyJkYXRhIjp7Imhhc2giOnsiYWxnb3JpdGhtIjoic2hhMjU2IiwidmFsdWUiOiJlY2UwNTcxYTZmMmYyNzQxOTUyNmQ2YjAwMTBmOGM1OGMxNTQ3NDYyNWVmZTQwZGUyNzI4M2JhMTE5ZjQ0ODZmIn19LCJzaWduYXR1cmUiOnsiY29udGVudCI6Ik1FVUNJUUN2WW5OR2J5TndpZFNiSzd0R1BiNG1VaVNIUjR5R1FJMEtBMGZUSU1RaEJBSWdNZnhEcWVqaTMwcG1IWEJCZHk3Q2IvYlpFZXdwNkxyV1VVR0dFZ2hJQ0xvPSIsInB1YmxpY0tleSI6eyJjb250ZW50IjoiTFMwdExTMUNSVWRKVGlCRFJWSlVTVVpKUTBGVVJTMHRMUzB0Q2sxSlNVaFRla05EUW5SSFowRjNTVUpCWjBsVlV6Wk9RVmwyUzB4aGMyeHVielJoWm5aaWJrbFJaM1ZZVmtsM2QwTm5XVWxMYjFwSmVtb3dSVUYzVFhjS1RucEZWazFDVFVkQk1WVkZRMmhOVFdNeWJHNWpNMUoyWTIxVmRWcEhWakpOVWpSM1NFRlpSRlpSVVVSRmVGWjZZVmRrZW1SSE9YbGFVekZ3WW01U2JBcGpiVEZzV2tkc2FHUkhWWGRJYUdOT1RXcFpkMDlFU1ROTlJHTjNUVlJOZDFkb1kwNU5hbGwzVDBSSk0wMUVZM2hOVkUxM1YycEJRVTFHYTNkRmQxbElDa3R2V2tsNmFqQkRRVkZaU1V0dldrbDZhakJFUVZGalJGRm5RVVZyYTJwVlFYaFFOVXRqUVhKbVZFdGlRVlp3WVdOWFRFeEtaWGczVkhCTFdrSndWbmNLZUhkbVZITndOMkUxTWprM1dIWkZWbkp0Y1hCS1REbGtWV0ZDTUZWaWNXbHJTemhGWkVaR2FXSkdiR2d3Y0dFNFduRlBRMEptUVhkbloxaHpUVUUwUndwQk1WVmtSSGRGUWk5M1VVVkJkMGxJWjBSQlZFSm5UbFpJVTFWRlJFUkJTMEpuWjNKQ1owVkdRbEZqUkVGNlFXUkNaMDVXU0ZFMFJVWm5VVlZqT1hkV0NqQlJUMjlHTWtWSmJVbHlja016ZDJsak0zWTJSMnRSZDBoM1dVUldVakJxUWtKbmQwWnZRVlV6T1ZCd2VqRlphMFZhWWpWeFRtcHdTMFpYYVhocE5Ga0tXa1E0ZDJGM1dVUldVakJTUVZGSUwwSkhSWGRZTkZwa1lVaFNNR05JVFRaTWVUbHVZVmhTYjJSWFNYVlpNamwwVERKb01Wb3laSEJpYldSdFdWZE9iQXBNTW5Sc1kyMDFiR0pJVFhSWk1qbDBZbGhXZFdGWVVqVk1lVFZ1WVZoU2IyUlhTWFprTWpsNVlUSmFjMkl6WkhwTU1rb3hZVmQ0YTB4dWJHaGlWM2hCQ21OdFZtMWplVGx2V2xkR2EyTjVPWFJaVjJ4MVRVUnJSME5wYzBkQlVWRkNaemM0ZDBGUlJVVkxNbWd3WkVoQ2VrOXBPSFprUnpseVdsYzBkVmxYVGpBS1lWYzVkV041Tlc1aFdGSnZaRmRLTVdNeVZubFpNamwxWkVkV2RXUkROV3BpTWpCM1NIZFpTMHQzV1VKQ1FVZEVkbnBCUWtGblVWSmtNamw1WVRKYWN3cGlNMlJtV2tkc2VtTkhSakJaTW1kM1RtZFpTMHQzV1VKQ1FVZEVkbnBCUWtGM1VXOVpWMUpyV21wUk0wNUVSbTFaVjFwcFRsUldhVTVIVFhoTlJFWnBDazFVVlRWUFJGWnJXbTFHYkZwWFJUTlplbVJxVFcxWmVsbFVRVlJDWjI5eVFtZEZSVUZaVHk5TlFVVkZRa0ZXUTJSWGJITmFSRUZ5UW1kdmNrSm5SVVVLUVZsUEwwMUJSVVpDUWpGdlpGZGtibUZYTlc1YWJVWnFXbE01Y2xwWVNuVmFWM2g2VEZkT2RtSlhNVEZpYld3d1pWUkJaRUpuYjNKQ1owVkZRVmxQTHdwTlFVVkhRa0U1ZVZwWFducE1NbWhzV1ZkU2Vrd3lNV2hoVnpSM1QzZFpTMHQzV1VKQ1FVZEVkbnBCUWtOQlVYUkVRM1J2WkVoU2QyTjZiM1pNTTFKMkNtRXlWblZNYlVacVpFZHNkbUp1VFhWYU1td3dZVWhXYVdSWVRteGpiVTUyWW01U2JHSnVVWFZaTWpsMFRVY3dSME5wYzBkQlVWRkNaemM0ZDBGUmEwVUtXSGQ0WkdGSVVqQmpTRTAyVEhrNWJtRllVbTlrVjBsMVdUSTVkRXd5YURGYU1tUndZbTFrYlZsWFRteE1NblJzWTIwMWJHSklUWFJaTWpsMFlsaFdkUXBoV0ZJMVRIazFibUZZVW05a1YwbDJaREk1ZVdFeVduTmlNMlI2VERKS01XRlhlR3RNYm14b1lsZDRRV050Vm0xamVUbHZXbGRHYTJONU9YUlpWMngxQ2sxRVowZERhWE5IUVZGUlFtYzNPSGRCVVc5RlMyZDNiMWxYVW10YWFsRXpUa1JHYlZsWFdtbE9WRlpwVGtkTmVFMUVSbWxOVkZVMVQwUldhMXB0Um13S1dsZEZNMWw2WkdwTmJWbDZXVlJCWWtKbmIzSkNaMFZGUVZsUEwwMUJSVXhDUVRCTlF6Tk9iR0pIV1hSaFJ6bDZaRWRXYTAxRlFVZERhWE5IUVZGUlFncG5OemgzUVZGM1JVMW5kM2RoU0ZJd1kwaE5Oa3g1T1c1aFdGSnZaRmRKZFZreU9YUk1NbWd4V2pKa2NHSnRaRzFaVjA1c1RESjBiR050Tld4aVNFMTBDbGt5T1hSaVdGWjFZVmhTTlUxRVowZERhWE5IUVZGUlFtYzNPSGRCVVRCRlMyZDNiMWxYVW10YWFsRXpUa1JHYlZsWFdtbE9WRlpwVGtkTmVFMUVSbWtLVFZSVk5VOUVWbXRhYlVac1dsZEZNMWw2WkdwTmJWbDZXVlJCWmtKbmIzSkNaMFZGUVZsUEwwMUJSVTlDUWtWTlJETktiRnB1VFhaaFIxWm9Xa2hOZGdwaVYwWndZbXBCWVVKbmIzSkNaMFZGUVZsUEwwMUJSVkJDUVhkTlEycEZkMDU2UlRCT2VsVXhUV3ByZDB4bldVdExkMWxDUWtGSFJIWjZRVUpGUVZGbkNrUkNOVzlrU0ZKM1kzcHZka3d5WkhCa1IyZ3hXV2sxYW1JeU1IWmhTRlp1V2pKc2RWb3lXbWhaTWxWM1IwRlpTMHQzV1VKQ1FVZEVkbnBCUWtWUlVVc0tSRUZuZVU1VVkzbE5SR013VFhwQ2RFSm5iM0pDWjBWRlFWbFBMMDFCUlZOQ1JqaE5XRmRvTUdSSVFucFBhVGgyV2pKc01HRklWbWxNYlU1MllsTTVid3BrVjJSdVlWYzFibHB0Um1wYVV6bHlXbGhLZFZwWGVIcE1WMDUyWWxjeE1XSnRiREJsVXpoMVdqSnNNR0ZJVm1sTU0yUjJZMjEwYldKSE9UTmplVGxwQ21SWGJITmFRelUxV1ZjeGMxRklTbXhhYmsxMllVZFdhRnBJVFhaaVYwWndZbXBCTkVKbmIzSkNaMFZGUVZsUEwwMUJSVlJDUTI5TlMwZEdhMXBIV1RBS1RucFJlRnB0Um0xWmFsVXhXV3BTYWsxVVFYaFpha1V4VDFSbk1WcEhXbWhhVjFab1RqSk5NMWw2U20xTk1rVjNTVkZaUzB0M1dVSkNRVWRFZG5wQlFncEdRVkZVUkVKR00ySXpTbkphYlhoMlpERTVhMkZZVG5kWldGSnFZVVJDYTBKbmIzSkNaMFZGUVZsUEwwMUJSVlpDUmxsTlZrZG9NR1JJUW5wUGFUaDJDbG95YkRCaFNGWnBURzFPZG1KVE9XOWtWMlJ1WVZjMWJscHRSbXBhVXpseVdsaEtkVnBYZUhwTVYwNTJZbGN4TVdKdGJEQmxVemxvV1ROU2NHSXlOWG9LVEROS01XSnVUWFpOZWsxM1RrUlZlazlVVlRWT1ZFbDJXVmhTTUZwWE1YZGtTRTEyVFZSQlYwSm5iM0pDWjBWRlFWbFBMMDFCUlZkQ1FXZE5RbTVDTVFwWmJYaHdXWHBDUjBKbmIzSkNaMFZGUVZsUEwwMUJSVmxDUkdkTlRtNUtiR05IT0RaaFNGWnVXakpzZFZveVdtaFpNbFYyWVRKV2VXSnRWbk5qZVRGcUNtSXlNWFJrVnpWd1pFaHJObU50Vm0xUGJrcHNXbTVOZG1GSFZtaGFTRTEyWWxkR2NHSnFRMEpwWjFsTFMzZFpRa0pCU0ZkbFVVbEZRV2RTT0VKSWIwRUtaVUZDTWtGT01EbE5SM0pIZUhoRmVWbDRhMlZJU214dVRuZExhVk5zTmpRemFubDBMelJsUzJOdlFYWkxaVFpQUVVGQlFtOUZTVWRHYWxWQlFVRlJSQXBCUldOM1VsRkphRUZPVkZKR01qQk9UMjFhWlZCNVNHTmhURnB2UWl0WmJtZzVaMjFsV0hWTmEyRnNUVkphUTJWUFpXTnVRV2xDZUdwT1V6RmFNaXRxQ2s5bFYxSmhObGhKYXpOWWNuRnVTeTlKY1RReGVVbFZTV05uYkU1ak1FMXdSV3BCUzBKblozRm9hMnBQVUZGUlJFRjNUbTlCUkVKc1FXcEJaM1ZuZDNNS1dqUmhlWFJSZDNoeUx6ZHdhRzgzUjFaYWIyTk5VemxWVjNWblZsRm1hRGRTVW5KeVVWWjNiMDVXZUZJd1dFSjNNRVJFSzFOcGFHUlJTMk5EVFZGRVZncHhMMnRFT0hWblpsTXZSVXhWZVVvMWMxSnBRMGRCVVVKT1RtOWpjMFoxTkhoVmMwcHZXVEJ3TWxocFowOHpPR2hQVkdKdVpqaDFjbXBNTUUxUFlYTTlDaTB0TFMwdFJVNUVJRU5GVWxSSlJrbERRVlJGTFMwdExTMEsifX19fQ=="}], "timestampVerificationData":{"rfc3161Timestamps":[{"signedTimestamp":"MIICyTADAgEAMIICwAYJKoZIhvcNAQcCoIICsTCCAq0CAQMxDTALBglghkgBZQMEAgEwgbgGCyqGSIb3DQEJEAEEoIGoBIGlMIGiAgEBBgkrBgEEAYO/MAIwMTANBglghkgBZQMEAgEFAAQgW8NRjp7oXtREaXDIukSp7psVkQj2lN7XSpes9QY1uoICFQDK2i1Rrxz3BNQNjaKBZFvL/WPHCxgPMjAyNjA4MjcwNzAxMzBaMAMCAQGgMqQwMC4xFTATBgNVBAoTDHNpZ3N0b3JlLmRldjEVMBMGA1UEAxMMc2lnc3RvcmUtdHNhoAAxggHaMIIB1gIBATBRMDkxFTATBgNVBAoTDHNpZ3N0b3JlLmRldjEgMB4GA1UEAxMXc2lnc3RvcmUtdHNhLXNlbGZzaWduZWQCFDoTVC8MkGHuvMFDL8uKjosqI4sMMAsGCWCGSAFlAwQCAaCB/DAaBgkqhkiG9w0BCQMxDQYLKoZIhvcNAQkQAQQwHAYJKoZIhvcNAQkFMQ8XDTI2MDgyNzA3MDEzMFowLwYJKoZIhvcNAQkEMSIEIC8EDpIkUJBXewTGsG4JIfOZtYX8DWyTISXasNinPG+PMIGOBgsqhkiG9w0BCRACLzF/MH0wezB5BCCF+Se8B6tiysO0Q1bBDvyBssaIP9p6uebYcNnROs0FtzBVMD2kOzA5MRUwEwYDVQQKEwxzaWdzdG9yZS5kZXYxIDAeBgNVBAMTF3NpZ3N0b3JlLXRzYS1zZWxmc2lnbmVkAhQ6E1QvDJBh7rzBQy/Lio6LKiOLDDAKBggqhkjOPQQDAgRmMGQCMGDN0sYVL3TD1HrB6Wh/Hqd43ZA7KZfW2UdfAe1NSvZkbyZ/FMz43fIxpBTZA21YpQIwDNOsqeMnIf36VRzdvs0+a5qhkSa7whBgvg5SROyG67Ql1V04UjSlV31aIJhTQW/a"}]}}, "messageSignature":{"messageDigest":{"algorithm":"SHA2_256", "digest":"7OBXGm8vJ0GVJtawAQ+MWMFUdGJe/kDeJyg7oRn0SG8="}, "signature":"MEUCIQCvYnNGbyNwidSbK7tGPb4mUiSHR4yGQI0KA0fTIMQhBAIgMfxDqeji30pmHXBBdy7Cb/bZEewp6LrWUUGGEghICLo="}}
|
build/torch214-cxx11-xpu20261-x86_64-linux/__init__.py
ADDED
|
@@ -0,0 +1,205 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2024 Databricks
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
import torch
|
| 5 |
+
|
| 6 |
+
# Stable alias: bare `ops` is shadowed by `from . import layers` below.
|
| 7 |
+
from ._ops import ops as _compiled_ops
|
| 8 |
+
from . import ops
|
| 9 |
+
|
| 10 |
+
from .grouped_gemm import backend as gg_backend
|
| 11 |
+
from .grouped_gemm import ops as gg_ops
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
from ._layers.arguments import Arguments
|
| 15 |
+
from ._layers.dmoe import ParallelDroplessMLP, dMoE
|
| 16 |
+
from ._layers.glu import SparseGLU
|
| 17 |
+
from ._layers.mlp import MLP, SparseMLP
|
| 18 |
+
from ._layers.moe import MoE, ParallelMLP, get_load_balancing_loss
|
| 19 |
+
|
| 20 |
+
from . import layers
|
| 21 |
+
|
| 22 |
+
# This section contains the direct kernel exports (not inlcuded in the original code)
|
| 23 |
+
def exclusive_cumsum(x: torch.Tensor, dim: int, out: torch.Tensor) -> torch.Tensor:
|
| 24 |
+
"""
|
| 25 |
+
Compute exclusive cumulative sum along the specified dimension.
|
| 26 |
+
|
| 27 |
+
Args:
|
| 28 |
+
x: Input tensor
|
| 29 |
+
dim: Dimension along which to compute cumsum
|
| 30 |
+
out: Output tensor (modified in-place)
|
| 31 |
+
|
| 32 |
+
Returns:
|
| 33 |
+
The output tensor
|
| 34 |
+
"""
|
| 35 |
+
result = ops.exclusive_cumsum(x, dim)
|
| 36 |
+
out.copy_(result)
|
| 37 |
+
return out
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
def inclusive_cumsum(x: torch.Tensor, dim: int, out: torch.Tensor) -> torch.Tensor:
|
| 41 |
+
"""
|
| 42 |
+
Compute inclusive cumulative sum along the specified dimension.
|
| 43 |
+
|
| 44 |
+
Args:
|
| 45 |
+
x: Input tensor
|
| 46 |
+
dim: Dimension along which to compute cumsum
|
| 47 |
+
out: Output tensor (modified in-place)
|
| 48 |
+
|
| 49 |
+
Returns:
|
| 50 |
+
The output tensor
|
| 51 |
+
"""
|
| 52 |
+
result = ops.inclusive_cumsum(x, dim)
|
| 53 |
+
out.copy_(result)
|
| 54 |
+
return out
|
| 55 |
+
|
| 56 |
+
|
| 57 |
+
def histogram(x: torch.Tensor, num_bins: int) -> torch.Tensor:
|
| 58 |
+
"""
|
| 59 |
+
Compute histogram of input tensor values.
|
| 60 |
+
|
| 61 |
+
Args:
|
| 62 |
+
x: Input tensor
|
| 63 |
+
num_bins: Number of histogram bins
|
| 64 |
+
|
| 65 |
+
Returns:
|
| 66 |
+
Histogram tensor with counts for each bin
|
| 67 |
+
"""
|
| 68 |
+
return ops.histogram(x, num_bins)
|
| 69 |
+
|
| 70 |
+
|
| 71 |
+
def indices(
|
| 72 |
+
padded_bins: torch.Tensor,
|
| 73 |
+
block_size: int,
|
| 74 |
+
output_block_rows: int,
|
| 75 |
+
output_block_columns: int,
|
| 76 |
+
) -> torch.Tensor:
|
| 77 |
+
"""
|
| 78 |
+
Construct indices from padded bins for sparse operations.
|
| 79 |
+
|
| 80 |
+
Args:
|
| 81 |
+
padded_bins: Tensor containing bin boundaries
|
| 82 |
+
block_size: Size of each block
|
| 83 |
+
output_block_rows: Number of rows in output blocks
|
| 84 |
+
output_block_columns: Number of columns in output blocks
|
| 85 |
+
|
| 86 |
+
Returns:
|
| 87 |
+
Tensor containing constructed indices
|
| 88 |
+
"""
|
| 89 |
+
return ops.indices(padded_bins, block_size, output_block_rows, output_block_columns)
|
| 90 |
+
|
| 91 |
+
|
| 92 |
+
def replicate_forward(
|
| 93 |
+
x: torch.Tensor, bins: torch.Tensor, out: torch.Tensor
|
| 94 |
+
) -> torch.Tensor:
|
| 95 |
+
"""
|
| 96 |
+
Forward pass of replicate operation - replicate values according to bin sizes.
|
| 97 |
+
|
| 98 |
+
Args:
|
| 99 |
+
x: Input tensor with values to replicate
|
| 100 |
+
bins: Tensor containing bin sizes
|
| 101 |
+
out: Output tensor (modified in-place)
|
| 102 |
+
|
| 103 |
+
Returns:
|
| 104 |
+
The output tensor
|
| 105 |
+
"""
|
| 106 |
+
return ops.replicate_forward(x, bins, out)
|
| 107 |
+
|
| 108 |
+
|
| 109 |
+
def replicate_backward(
|
| 110 |
+
grad: torch.Tensor, bins: torch.Tensor, out: torch.Tensor
|
| 111 |
+
) -> torch.Tensor:
|
| 112 |
+
"""
|
| 113 |
+
Backward pass of replicate operation - reduce gradients back to bins.
|
| 114 |
+
|
| 115 |
+
Args:
|
| 116 |
+
grad: Gradient tensor to reduce
|
| 117 |
+
bins: Tensor containing bin sizes
|
| 118 |
+
out: Output tensor (modified in-place)
|
| 119 |
+
|
| 120 |
+
Returns:
|
| 121 |
+
The output tensor
|
| 122 |
+
"""
|
| 123 |
+
return ops.replicate_backward(grad, bins, out)
|
| 124 |
+
|
| 125 |
+
|
| 126 |
+
def sort(
|
| 127 |
+
x: torch.Tensor, end_bit: int, x_out: torch.Tensor, iota_out: torch.Tensor
|
| 128 |
+
) -> torch.Tensor:
|
| 129 |
+
"""
|
| 130 |
+
Radix sort with index tracking.
|
| 131 |
+
|
| 132 |
+
Args:
|
| 133 |
+
x: Input tensor to sort
|
| 134 |
+
end_bit: Number of bits to consider in sorting
|
| 135 |
+
x_out: Output tensor for sorted values
|
| 136 |
+
iota_out: Output tensor for sorted indices
|
| 137 |
+
|
| 138 |
+
Returns:
|
| 139 |
+
The sorted values tensor
|
| 140 |
+
"""
|
| 141 |
+
_compiled_ops.sort(x, end_bit, x_out, iota_out)
|
| 142 |
+
return x_out
|
| 143 |
+
|
| 144 |
+
|
| 145 |
+
# Convenience functions for common use cases
|
| 146 |
+
def cumsum(x: torch.Tensor, dim: int = -1, exclusive: bool = False) -> torch.Tensor:
|
| 147 |
+
"""
|
| 148 |
+
Compute cumulative sum with automatic output allocation.
|
| 149 |
+
|
| 150 |
+
Args:
|
| 151 |
+
x: Input tensor
|
| 152 |
+
dim: Dimension along which to compute cumsum (default: last dimension)
|
| 153 |
+
exclusive: Whether to compute exclusive (True) or inclusive (False) cumsum
|
| 154 |
+
|
| 155 |
+
Returns:
|
| 156 |
+
New tensor containing the cumulative sum
|
| 157 |
+
"""
|
| 158 |
+
out = torch.empty_like(x)
|
| 159 |
+
if exclusive:
|
| 160 |
+
return exclusive_cumsum(x, dim, out)
|
| 161 |
+
else:
|
| 162 |
+
return inclusive_cumsum(x, dim, out)
|
| 163 |
+
|
| 164 |
+
|
| 165 |
+
def argsort(x: torch.Tensor, end_bit: int = 32) -> tuple[torch.Tensor, torch.Tensor]:
|
| 166 |
+
"""
|
| 167 |
+
Sort tensor and return both sorted values and indices.
|
| 168 |
+
|
| 169 |
+
Args:
|
| 170 |
+
x: Input tensor to sort
|
| 171 |
+
end_bit: Number of bits to consider in sorting
|
| 172 |
+
|
| 173 |
+
Returns:
|
| 174 |
+
Tuple of (sorted_values, sorted_indices)
|
| 175 |
+
"""
|
| 176 |
+
x_out = torch.empty_like(x)
|
| 177 |
+
iota_out = torch.empty_like(x)
|
| 178 |
+
sort(x, end_bit, x_out, iota_out)
|
| 179 |
+
return x_out, iota_out
|
| 180 |
+
|
| 181 |
+
|
| 182 |
+
# Export public API
|
| 183 |
+
__all__ = [
|
| 184 |
+
"MyReplacementLayer",
|
| 185 |
+
# Direct kernel exports
|
| 186 |
+
"exclusive_cumsum",
|
| 187 |
+
"inclusive_cumsum",
|
| 188 |
+
"histogram",
|
| 189 |
+
"indices",
|
| 190 |
+
"replicate_forward",
|
| 191 |
+
"replicate_backward",
|
| 192 |
+
"sort",
|
| 193 |
+
"cumsum",
|
| 194 |
+
"argsort",
|
| 195 |
+
# Original exports
|
| 196 |
+
"Arguments",
|
| 197 |
+
"ParallelDroplessMLP",
|
| 198 |
+
"dMoE",
|
| 199 |
+
"SparseGLU",
|
| 200 |
+
"MLP",
|
| 201 |
+
"SparseMLP",
|
| 202 |
+
"MoE",
|
| 203 |
+
"ParallelMLP",
|
| 204 |
+
"get_load_balancing_loss",
|
| 205 |
+
]
|
build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/__init__.py
ADDED
|
File without changes
|
build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/_triton_kernels/__init__.py
ADDED
|
File without changes
|
build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/_triton_kernels/gmm.py
ADDED
|
@@ -0,0 +1,574 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: MIT
|
| 2 |
+
# Copyright (C) 2025-2026, Advanced Micro Devices, Inc. All rights reserved.
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
# Imports.
|
| 6 |
+
# ------------------------------------------------------------------------------
|
| 7 |
+
|
| 8 |
+
# Python standard library
|
| 9 |
+
import functools
|
| 10 |
+
|
| 11 |
+
# Triton
|
| 12 |
+
import triton
|
| 13 |
+
import triton.language as tl
|
| 14 |
+
|
| 15 |
+
# AITER
|
| 16 |
+
from ..configs import CONFIGS as _CONFIGS
|
| 17 |
+
from ..utils._triton import arch_info
|
| 18 |
+
from ..utils._triton.pid_preprocessing import pid_grid, remap_xcd
|
| 19 |
+
|
| 20 |
+
# Kernel config.
|
| 21 |
+
# ------------------------------------------------------------------------------
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
@functools.lru_cache()
|
| 25 |
+
def get_config(
|
| 26 |
+
gmm_type: str, M: int, K: int, N: int, G: int, accumulate: bool = False
|
| 27 |
+
) -> dict[str, int]:
|
| 28 |
+
assert gmm_type in {
|
| 29 |
+
"gmm",
|
| 30 |
+
"ptgmm",
|
| 31 |
+
"nptgmm",
|
| 32 |
+
}, f"'{gmm_type}' is an invalid GMM variant."
|
| 33 |
+
dev = arch_info.get_arch()
|
| 34 |
+
assert (
|
| 35 |
+
dev in _CONFIGS
|
| 36 |
+
), f"No GMM configuration tuned for arch '{dev}'. Supported: {sorted(_CONFIGS)}."
|
| 37 |
+
arch_configs = _CONFIGS[dev]
|
| 38 |
+
assert (
|
| 39 |
+
"default" in arch_configs[gmm_type]
|
| 40 |
+
), "Default configuration is absent."
|
| 41 |
+
key = "accumulate" if accumulate else "default"
|
| 42 |
+
return arch_configs[gmm_type][key]
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
# Common code shared by GMM and TGMM kernels.
|
| 46 |
+
# ------------------------------------------------------------------------------
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
# XCD remapping followed by 1D PID to 2D grid mapping.
|
| 50 |
+
@triton.jit
|
| 51 |
+
def _remap_xcd_tile_grid(
|
| 52 |
+
tile_in_mm,
|
| 53 |
+
num_row_tiles,
|
| 54 |
+
num_col_tiles,
|
| 55 |
+
GROUP_SIZE: tl.constexpr = 1,
|
| 56 |
+
NUM_XCDS: tl.constexpr = 8,
|
| 57 |
+
):
|
| 58 |
+
return pid_grid(
|
| 59 |
+
remap_xcd(tile_in_mm, num_row_tiles * num_col_tiles, NUM_XCDS=NUM_XCDS),
|
| 60 |
+
num_row_tiles,
|
| 61 |
+
num_col_tiles,
|
| 62 |
+
GROUP_SIZE_M=GROUP_SIZE,
|
| 63 |
+
)
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
# GMM kernel.
|
| 67 |
+
# ------------------------------------------------------------------------------
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
@triton.heuristics(
|
| 71 |
+
{
|
| 72 |
+
"K_DIVISIBLE_BY_BLOCK_SIZE_K": lambda META: META["K"] % META["BLOCK_SIZE_K"]
|
| 73 |
+
== 0,
|
| 74 |
+
}
|
| 75 |
+
)
|
| 76 |
+
@triton.jit
|
| 77 |
+
def gmm_kernel(
|
| 78 |
+
# Tensor pointers:
|
| 79 |
+
lhs_ptr,
|
| 80 |
+
rhs_ptr,
|
| 81 |
+
group_sizes_ptr,
|
| 82 |
+
out_ptr,
|
| 83 |
+
bias_ptr,
|
| 84 |
+
# Tensor shapes:
|
| 85 |
+
M: int,
|
| 86 |
+
K: int,
|
| 87 |
+
N: int,
|
| 88 |
+
G: int,
|
| 89 |
+
# Meta-parameters:
|
| 90 |
+
TRANS_RHS: tl.constexpr,
|
| 91 |
+
BLOCK_SIZE_M: tl.constexpr,
|
| 92 |
+
BLOCK_SIZE_K: tl.constexpr,
|
| 93 |
+
BLOCK_SIZE_N: tl.constexpr,
|
| 94 |
+
K_DIVISIBLE_BY_BLOCK_SIZE_K: tl.constexpr,
|
| 95 |
+
GROUP_SIZE: tl.constexpr,
|
| 96 |
+
GRID_DIM: tl.constexpr,
|
| 97 |
+
USE_BIAS: tl.constexpr,
|
| 98 |
+
):
|
| 99 |
+
tl.assume(M > 0)
|
| 100 |
+
tl.assume(K > 0)
|
| 101 |
+
tl.assume(N > 0)
|
| 102 |
+
tl.assume(G > 0)
|
| 103 |
+
|
| 104 |
+
num_n_tiles = tl.cdiv(N, BLOCK_SIZE_N)
|
| 105 |
+
tl.device_assert(num_n_tiles > 0, "num_n_tiles <= 0")
|
| 106 |
+
|
| 107 |
+
# Current tile. Each program computes multiple tiles of each group.
|
| 108 |
+
tile = tl.program_id(0)
|
| 109 |
+
tl.device_assert(tile >= 0, "tile < 0 (at initialization)")
|
| 110 |
+
|
| 111 |
+
# Tile limit of last MM problem (inclusive).
|
| 112 |
+
last_mm_tile = 0
|
| 113 |
+
|
| 114 |
+
# Last input row of lhs and output row of out. Each group reads some rows of
|
| 115 |
+
# lhs and writes some rows to out.
|
| 116 |
+
last_m = 0
|
| 117 |
+
|
| 118 |
+
# Loop through all (m, K, N) MM problems:
|
| 119 |
+
# (m, K) x (K, N) = (m, N)
|
| 120 |
+
# sum(m) = M
|
| 121 |
+
for g in range(G):
|
| 122 |
+
# Get m dimension of current MM problem.
|
| 123 |
+
m = tl.load(group_sizes_ptr + g)
|
| 124 |
+
# m can be zero if group is empty
|
| 125 |
+
tl.device_assert(m >= 0, "m < 0")
|
| 126 |
+
|
| 127 |
+
num_m_tiles = tl.cdiv(m, BLOCK_SIZE_M)
|
| 128 |
+
# num_m_tiles can be zero if group is empty
|
| 129 |
+
tl.device_assert(num_m_tiles >= 0, "num_m_tiles < 0")
|
| 130 |
+
|
| 131 |
+
num_tiles = num_m_tiles * num_n_tiles
|
| 132 |
+
# num_tiles can be zero if group is empty
|
| 133 |
+
tl.device_assert(num_tiles >= 0, "num_tiles < 0")
|
| 134 |
+
|
| 135 |
+
# Loop through tiles of current MM problem.
|
| 136 |
+
while tile >= last_mm_tile and tile < last_mm_tile + num_tiles:
|
| 137 |
+
# Figure out tile coordinates in current MM problem.
|
| 138 |
+
tile_in_mm = tile - last_mm_tile
|
| 139 |
+
tl.device_assert(tile_in_mm >= 0, "tile_in_mm < 0")
|
| 140 |
+
|
| 141 |
+
tile_m, tile_n = _remap_xcd_tile_grid(
|
| 142 |
+
tile_in_mm, num_m_tiles, num_n_tiles, GROUP_SIZE=GROUP_SIZE
|
| 143 |
+
)
|
| 144 |
+
|
| 145 |
+
# Do regular MM:
|
| 146 |
+
|
| 147 |
+
tl.device_assert(tile_m * BLOCK_SIZE_M >= 0, "tile_m * BLOCK_SIZE_M < 0")
|
| 148 |
+
tl.device_assert(tile_n * BLOCK_SIZE_N >= 0, "tile_n * BLOCK_SIZE_N < 0")
|
| 149 |
+
|
| 150 |
+
offs_lhs_m = (
|
| 151 |
+
tile_m.to(tl.int64) * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
|
| 152 |
+
) % m
|
| 153 |
+
offs_rhs_n = (
|
| 154 |
+
tile_n.to(tl.int64) * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
|
| 155 |
+
) % N
|
| 156 |
+
offs_k = tl.arange(0, BLOCK_SIZE_K).to(tl.int64)
|
| 157 |
+
|
| 158 |
+
lhs_ptrs = lhs_ptr + (last_m + offs_lhs_m[:, None]) * K + offs_k[None, :]
|
| 159 |
+
|
| 160 |
+
if TRANS_RHS:
|
| 161 |
+
rhs_ptrs = (
|
| 162 |
+
rhs_ptr
|
| 163 |
+
+ g.to(tl.int64) * K * N
|
| 164 |
+
+ offs_k[:, None]
|
| 165 |
+
+ offs_rhs_n[None, :] * K
|
| 166 |
+
)
|
| 167 |
+
else:
|
| 168 |
+
rhs_ptrs = (
|
| 169 |
+
rhs_ptr
|
| 170 |
+
+ g.to(tl.int64) * K * N
|
| 171 |
+
+ offs_k[:, None] * N
|
| 172 |
+
+ offs_rhs_n[None, :]
|
| 173 |
+
)
|
| 174 |
+
|
| 175 |
+
acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
|
| 176 |
+
|
| 177 |
+
for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)):
|
| 178 |
+
if K_DIVISIBLE_BY_BLOCK_SIZE_K:
|
| 179 |
+
lhs = tl.load(lhs_ptrs)
|
| 180 |
+
rhs = tl.load(rhs_ptrs)
|
| 181 |
+
else:
|
| 182 |
+
k_mask_limit = K - k * BLOCK_SIZE_K
|
| 183 |
+
lhs = tl.load(
|
| 184 |
+
lhs_ptrs, mask=offs_k[None, :] < k_mask_limit, other=0
|
| 185 |
+
)
|
| 186 |
+
rhs = tl.load(
|
| 187 |
+
rhs_ptrs, mask=offs_k[:, None] < k_mask_limit, other=0
|
| 188 |
+
)
|
| 189 |
+
|
| 190 |
+
acc = tl.dot(lhs, rhs, acc=acc)
|
| 191 |
+
|
| 192 |
+
lhs_ptrs += BLOCK_SIZE_K
|
| 193 |
+
|
| 194 |
+
if TRANS_RHS:
|
| 195 |
+
rhs_ptrs += BLOCK_SIZE_K
|
| 196 |
+
else:
|
| 197 |
+
rhs_ptrs += BLOCK_SIZE_K * N
|
| 198 |
+
|
| 199 |
+
# Add bias if enabled
|
| 200 |
+
if USE_BIAS:
|
| 201 |
+
offs_bias_n = tile_n.to(tl.int64) * BLOCK_SIZE_N + tl.arange(
|
| 202 |
+
0, BLOCK_SIZE_N
|
| 203 |
+
)
|
| 204 |
+
bias_ptrs = bias_ptr + g.to(tl.int64) * N + offs_bias_n
|
| 205 |
+
bias = tl.load(bias_ptrs, mask=offs_bias_n < N, other=0.0)
|
| 206 |
+
# Convert bias to float32 to match accumulator precision
|
| 207 |
+
bias = bias.to(tl.float32)
|
| 208 |
+
# Broadcast bias across M dimension and add in float32
|
| 209 |
+
acc += bias[None, :]
|
| 210 |
+
|
| 211 |
+
# Convert to output dtype after all computations
|
| 212 |
+
acc = acc.to(out_ptr.type.element_ty)
|
| 213 |
+
|
| 214 |
+
offs_out_m = tile_m.to(tl.int64) * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
|
| 215 |
+
offs_out_n = tile_n.to(tl.int64) * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
|
| 216 |
+
|
| 217 |
+
out_ptrs = (
|
| 218 |
+
out_ptr + (last_m + offs_out_m[:, None]) * N + offs_out_n[None, :]
|
| 219 |
+
)
|
| 220 |
+
|
| 221 |
+
tl.store(
|
| 222 |
+
out_ptrs,
|
| 223 |
+
acc,
|
| 224 |
+
mask=(offs_out_m[:, None] < m) & (offs_out_n[None, :] < N),
|
| 225 |
+
)
|
| 226 |
+
|
| 227 |
+
# Go to the next tile by advancing number of programs.
|
| 228 |
+
tile += GRID_DIM
|
| 229 |
+
tl.device_assert(tile > 0, "tile <= 0 (at update)")
|
| 230 |
+
|
| 231 |
+
# Get ready to go to the next MM problem.
|
| 232 |
+
|
| 233 |
+
last_mm_tile += num_tiles
|
| 234 |
+
# last_mm_tile can be zero if group 0 is skipped
|
| 235 |
+
tl.device_assert(last_mm_tile >= 0, "last_mm_tile < 0 (at update)")
|
| 236 |
+
|
| 237 |
+
last_m += m
|
| 238 |
+
# last_m can be zero if group 0 is skipped
|
| 239 |
+
tl.device_assert(last_m >= 0, "last_m < 0 (at update)")
|
| 240 |
+
tl.device_assert(last_m <= M, "last_m > M (at update)")
|
| 241 |
+
|
| 242 |
+
|
| 243 |
+
# Persistent TGMM kernel.
|
| 244 |
+
# ------------------------------------------------------------------------------
|
| 245 |
+
|
| 246 |
+
|
| 247 |
+
@triton.jit
|
| 248 |
+
def tgmm_persistent_kernel(
|
| 249 |
+
# Tensor pointers:
|
| 250 |
+
lhs_ptr,
|
| 251 |
+
rhs_ptr,
|
| 252 |
+
group_sizes_ptr,
|
| 253 |
+
out_ptr,
|
| 254 |
+
bias_grad_ptr,
|
| 255 |
+
# Tensor shapes:
|
| 256 |
+
M: int,
|
| 257 |
+
K: int,
|
| 258 |
+
N: int,
|
| 259 |
+
G: int,
|
| 260 |
+
# Meta-parameters:
|
| 261 |
+
TRANS_LHS: tl.constexpr,
|
| 262 |
+
BLOCK_SIZE_M: tl.constexpr,
|
| 263 |
+
BLOCK_SIZE_K: tl.constexpr,
|
| 264 |
+
BLOCK_SIZE_N: tl.constexpr,
|
| 265 |
+
GROUP_SIZE: tl.constexpr,
|
| 266 |
+
GRID_DIM: tl.constexpr,
|
| 267 |
+
COMPUTE_BIAS_GRAD: tl.constexpr,
|
| 268 |
+
ACCUMULATE: tl.constexpr,
|
| 269 |
+
):
|
| 270 |
+
tl.assume(M > 0)
|
| 271 |
+
tl.assume(K > 0)
|
| 272 |
+
tl.assume(N > 0)
|
| 273 |
+
tl.assume(G > 0)
|
| 274 |
+
|
| 275 |
+
num_k_tiles = tl.cdiv(K, BLOCK_SIZE_K)
|
| 276 |
+
tl.device_assert(num_k_tiles > 0, "num_k_tiles <= 0")
|
| 277 |
+
|
| 278 |
+
num_n_tiles = tl.cdiv(N, BLOCK_SIZE_N)
|
| 279 |
+
tl.device_assert(num_n_tiles > 0, "num_n_tiles <= 0")
|
| 280 |
+
|
| 281 |
+
num_tiles = num_k_tiles * num_n_tiles
|
| 282 |
+
tl.device_assert(num_tiles > 0, "num_tiles <= 0")
|
| 283 |
+
|
| 284 |
+
# Current tile. Each program computes multiple tiles of each group.
|
| 285 |
+
tile = tl.program_id(0)
|
| 286 |
+
tl.device_assert(tile >= 0, "tile < 0 (at initialization)")
|
| 287 |
+
|
| 288 |
+
# Tile limit of last MM problem (inclusive).
|
| 289 |
+
last_mm_tile = 0
|
| 290 |
+
|
| 291 |
+
# Last input column of lhs and input row of rhs. Each group reads some
|
| 292 |
+
# columns of lhs and some rows of rhs.
|
| 293 |
+
last_m = 0
|
| 294 |
+
|
| 295 |
+
# Loop through all (K, m, N) MM problems:
|
| 296 |
+
# (K, m) x (m, N) = (K, N)
|
| 297 |
+
# sum(m) = M
|
| 298 |
+
for g in range(G):
|
| 299 |
+
# Get m dimension of current MM problem.
|
| 300 |
+
m = tl.load(group_sizes_ptr + g)
|
| 301 |
+
# m can be zero if group is empty
|
| 302 |
+
tl.device_assert(m >= 0, "m < 0")
|
| 303 |
+
|
| 304 |
+
# Loop through tiles of current MM problem.
|
| 305 |
+
while tile >= last_mm_tile and tile < last_mm_tile + num_tiles:
|
| 306 |
+
# Figure out tile coordinates in current MM problem.
|
| 307 |
+
tile_in_mm = tile - last_mm_tile
|
| 308 |
+
tl.device_assert(tile_in_mm >= 0, "tile_in_mm < 0")
|
| 309 |
+
|
| 310 |
+
tile_k, tile_n = _remap_xcd_tile_grid(
|
| 311 |
+
tile_in_mm, num_k_tiles, num_n_tiles, GROUP_SIZE=GROUP_SIZE
|
| 312 |
+
)
|
| 313 |
+
|
| 314 |
+
# Do regular MM:
|
| 315 |
+
|
| 316 |
+
tl.device_assert(tile_k * BLOCK_SIZE_K >= 0, "tile_k * BLOCK_SIZE_K < 0")
|
| 317 |
+
tl.device_assert(tile_n * BLOCK_SIZE_N >= 0, "tile_n * BLOCK_SIZE_N < 0")
|
| 318 |
+
|
| 319 |
+
offs_lhs_k = (
|
| 320 |
+
tile_k.to(tl.int64) * BLOCK_SIZE_K + tl.arange(0, BLOCK_SIZE_K)
|
| 321 |
+
) % K
|
| 322 |
+
offs_rhs_n = (
|
| 323 |
+
tile_n.to(tl.int64) * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
|
| 324 |
+
) % N
|
| 325 |
+
offs_m = tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
|
| 326 |
+
|
| 327 |
+
if TRANS_LHS:
|
| 328 |
+
lhs_ptrs = (
|
| 329 |
+
lhs_ptr + offs_lhs_k[:, None] + (last_m + offs_m[None, :]) * K
|
| 330 |
+
)
|
| 331 |
+
else:
|
| 332 |
+
lhs_ptrs = (
|
| 333 |
+
lhs_ptr + offs_lhs_k[:, None] * M + (last_m + offs_m[None, :])
|
| 334 |
+
)
|
| 335 |
+
|
| 336 |
+
rhs_ptrs = rhs_ptr + (last_m + offs_m[:, None]) * N + offs_rhs_n[None, :]
|
| 337 |
+
|
| 338 |
+
loop_m = tl.cdiv(m, BLOCK_SIZE_M)
|
| 339 |
+
m_divisible_by_block_m = m % BLOCK_SIZE_M == 0
|
| 340 |
+
if not m_divisible_by_block_m:
|
| 341 |
+
loop_m -= 1
|
| 342 |
+
|
| 343 |
+
acc = tl.zeros((BLOCK_SIZE_K, BLOCK_SIZE_N), dtype=tl.float32)
|
| 344 |
+
|
| 345 |
+
# Initialize bias accumulator
|
| 346 |
+
bias_acc = tl.zeros((BLOCK_SIZE_K,), dtype=tl.float32)
|
| 347 |
+
|
| 348 |
+
for _ in range(0, loop_m):
|
| 349 |
+
lhs = tl.load(lhs_ptrs)
|
| 350 |
+
rhs = tl.load(rhs_ptrs)
|
| 351 |
+
|
| 352 |
+
acc = tl.dot(lhs, rhs, acc=acc)
|
| 353 |
+
|
| 354 |
+
# Accumulate for bias gradient: sum lhs across M dimension
|
| 355 |
+
if COMPUTE_BIAS_GRAD and tile_n == 0:
|
| 356 |
+
bias_acc += tl.sum(
|
| 357 |
+
lhs, axis=1
|
| 358 |
+
) # Sum across M dimension [K, M] -> [K]
|
| 359 |
+
|
| 360 |
+
if TRANS_LHS:
|
| 361 |
+
lhs_ptrs += BLOCK_SIZE_M * K
|
| 362 |
+
else:
|
| 363 |
+
lhs_ptrs += BLOCK_SIZE_M
|
| 364 |
+
|
| 365 |
+
rhs_ptrs += BLOCK_SIZE_M * N
|
| 366 |
+
|
| 367 |
+
if not m_divisible_by_block_m:
|
| 368 |
+
offs_lhs_k = (
|
| 369 |
+
tile_k.to(tl.int64) * BLOCK_SIZE_K + tl.arange(0, BLOCK_SIZE_K)
|
| 370 |
+
) % K
|
| 371 |
+
offs_rhs_n = (
|
| 372 |
+
tile_n.to(tl.int64) * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
|
| 373 |
+
) % N
|
| 374 |
+
offs_m = loop_m.to(tl.int64) * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
|
| 375 |
+
lhs = tl.load(lhs_ptrs, mask=offs_m[None, :] < m, other=0)
|
| 376 |
+
rhs = tl.load(rhs_ptrs, mask=offs_m[:, None] < m, other=0)
|
| 377 |
+
acc = tl.dot(lhs, rhs, acc=acc)
|
| 378 |
+
|
| 379 |
+
# Accumulate last chunk for bias gradient
|
| 380 |
+
if COMPUTE_BIAS_GRAD and tile_n == 0:
|
| 381 |
+
bias_acc += tl.sum(lhs, axis=1)
|
| 382 |
+
|
| 383 |
+
acc = acc.to(out_ptr.type.element_ty)
|
| 384 |
+
|
| 385 |
+
offs_out_k = tile_k.to(tl.int64) * BLOCK_SIZE_K + tl.arange(0, BLOCK_SIZE_K)
|
| 386 |
+
offs_out_n = tile_n.to(tl.int64) * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
|
| 387 |
+
|
| 388 |
+
out_ptrs = (
|
| 389 |
+
out_ptr
|
| 390 |
+
+ g.to(tl.int64) * K * N
|
| 391 |
+
+ offs_out_k[:, None] * N
|
| 392 |
+
+ offs_out_n[None, :]
|
| 393 |
+
)
|
| 394 |
+
|
| 395 |
+
mask = (offs_out_k[:, None] < K) & (offs_out_n[None, :] < N)
|
| 396 |
+
if ACCUMULATE:
|
| 397 |
+
# Load existing values and add to them (like beta=1 in BLAS)
|
| 398 |
+
old_vals = tl.load(out_ptrs, mask=mask, other=0.0)
|
| 399 |
+
tl.store(out_ptrs, acc + old_vals, mask=mask)
|
| 400 |
+
else:
|
| 401 |
+
# Overwrite output (like beta=0 in BLAS)
|
| 402 |
+
tl.store(out_ptrs, acc, mask=mask)
|
| 403 |
+
|
| 404 |
+
# Store bias gradient (only for first N tile, sum across all M)
|
| 405 |
+
if COMPUTE_BIAS_GRAD and tile_n == 0:
|
| 406 |
+
# Keep as float32 for atomic_add (bf16 not supported for atomics)
|
| 407 |
+
bias_grad_ptrs = bias_grad_ptr + g.to(tl.int64) * K + offs_out_k
|
| 408 |
+
# Use atomic add since multiple K-tiles may write to same expert's bias
|
| 409 |
+
tl.atomic_add(
|
| 410 |
+
bias_grad_ptrs, bias_acc, mask=offs_out_k < K, sem="relaxed"
|
| 411 |
+
)
|
| 412 |
+
|
| 413 |
+
# Go to the next tile by advancing number of programs.
|
| 414 |
+
tile += GRID_DIM
|
| 415 |
+
tl.device_assert(tile > 0, "tile <= 0 (at update)")
|
| 416 |
+
|
| 417 |
+
# Get ready to go to the next MM problem.
|
| 418 |
+
|
| 419 |
+
last_mm_tile += num_tiles
|
| 420 |
+
# last_mm_tile can be zero if group 0 is skipped
|
| 421 |
+
tl.device_assert(last_mm_tile >= 0, "last_mm_tile < 0 (at update)")
|
| 422 |
+
|
| 423 |
+
last_m += m
|
| 424 |
+
# last_m can be zero if group 0 is skipped
|
| 425 |
+
tl.device_assert(last_m >= 0, "last_m < 0 (at update)")
|
| 426 |
+
tl.device_assert(last_m <= M, "last_m > M (at update)")
|
| 427 |
+
|
| 428 |
+
|
| 429 |
+
# Regular non-persistent TGMM kernel.
|
| 430 |
+
# ------------------------------------------------------------------------------
|
| 431 |
+
|
| 432 |
+
|
| 433 |
+
@triton.heuristics({"BLOCK_SIZE_G": lambda META: triton.next_power_of_2(META["G"])})
|
| 434 |
+
@triton.jit
|
| 435 |
+
def tgmm_non_persistent_kernel(
|
| 436 |
+
# Tensor pointers:
|
| 437 |
+
lhs_ptr,
|
| 438 |
+
rhs_ptr,
|
| 439 |
+
group_sizes_ptr,
|
| 440 |
+
out_ptr,
|
| 441 |
+
bias_grad_ptr,
|
| 442 |
+
# Tensor shapes:
|
| 443 |
+
M: int,
|
| 444 |
+
K: int,
|
| 445 |
+
N: int,
|
| 446 |
+
G: int,
|
| 447 |
+
# Meta-parameters:
|
| 448 |
+
TRANS_LHS: tl.constexpr,
|
| 449 |
+
BLOCK_SIZE_G: tl.constexpr,
|
| 450 |
+
BLOCK_SIZE_M: tl.constexpr,
|
| 451 |
+
BLOCK_SIZE_K: tl.constexpr,
|
| 452 |
+
BLOCK_SIZE_N: tl.constexpr,
|
| 453 |
+
GROUP_SIZE: tl.constexpr,
|
| 454 |
+
COMPUTE_BIAS_GRAD: tl.constexpr,
|
| 455 |
+
ACCUMULATE: tl.constexpr,
|
| 456 |
+
):
|
| 457 |
+
tl.assume(M > 0)
|
| 458 |
+
tl.assume(K > 0)
|
| 459 |
+
tl.assume(N > 0)
|
| 460 |
+
tl.assume(G > 0)
|
| 461 |
+
|
| 462 |
+
# Get group ID from grid.
|
| 463 |
+
g = tl.program_id(0)
|
| 464 |
+
tl.device_assert(g >= 0, "g < 0")
|
| 465 |
+
tl.device_assert(g < G, "g >= G")
|
| 466 |
+
|
| 467 |
+
# Get m dimension of current MM group.
|
| 468 |
+
m = tl.load(group_sizes_ptr + g)
|
| 469 |
+
# m can be zero if group is empty.
|
| 470 |
+
tl.device_assert(m >= 0, "m < 0")
|
| 471 |
+
|
| 472 |
+
# Skip empty groups.
|
| 473 |
+
if m == 0:
|
| 474 |
+
return
|
| 475 |
+
|
| 476 |
+
# Compute sum(group_sizes) until current group g.
|
| 477 |
+
# It's the starting column of lhs and starting row of rhs.
|
| 478 |
+
offs_g = tl.arange(0, BLOCK_SIZE_G)
|
| 479 |
+
group_sizes = tl.load(group_sizes_ptr + offs_g, mask=offs_g < g, other=0)
|
| 480 |
+
start_m = tl.sum(group_sizes)
|
| 481 |
+
|
| 482 |
+
num_k_tiles = tl.cdiv(K, BLOCK_SIZE_K)
|
| 483 |
+
tl.device_assert(num_k_tiles > 0, "num_k_tiles <= 0")
|
| 484 |
+
|
| 485 |
+
num_n_tiles = tl.cdiv(N, BLOCK_SIZE_N)
|
| 486 |
+
tl.device_assert(num_n_tiles > 0, "num_n_tiles <= 0")
|
| 487 |
+
|
| 488 |
+
# Get MM tile from grid.
|
| 489 |
+
tile_in_mm = tl.program_id(1)
|
| 490 |
+
tl.device_assert(tile_in_mm >= 0, "tile_in_mm < 0")
|
| 491 |
+
|
| 492 |
+
tile_k, tile_n = _remap_xcd_tile_grid(
|
| 493 |
+
tile_in_mm, num_k_tiles, num_n_tiles, GROUP_SIZE=GROUP_SIZE
|
| 494 |
+
)
|
| 495 |
+
|
| 496 |
+
tl.device_assert(tile_k * BLOCK_SIZE_K >= 0, "tile_k * BLOCK_SIZE_K < 0")
|
| 497 |
+
tl.device_assert(tile_n * BLOCK_SIZE_N >= 0, "tile_n * BLOCK_SIZE_N < 0")
|
| 498 |
+
|
| 499 |
+
offs_lhs_k = (tile_k.to(tl.int64) * BLOCK_SIZE_K + tl.arange(0, BLOCK_SIZE_K)) % K
|
| 500 |
+
offs_rhs_n = (tile_n.to(tl.int64) * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
|
| 501 |
+
offs_m = tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
|
| 502 |
+
|
| 503 |
+
if TRANS_LHS:
|
| 504 |
+
lhs_ptrs = lhs_ptr + offs_lhs_k[:, None] + (start_m + offs_m[None, :]) * K
|
| 505 |
+
else:
|
| 506 |
+
lhs_ptrs = lhs_ptr + offs_lhs_k[:, None] * M + (start_m + offs_m[None, :])
|
| 507 |
+
|
| 508 |
+
rhs_ptrs = rhs_ptr + (start_m + offs_m[:, None]) * N + offs_rhs_n[None, :]
|
| 509 |
+
|
| 510 |
+
loop_m = tl.cdiv(m, BLOCK_SIZE_M)
|
| 511 |
+
m_divisible_by_block_m = m % BLOCK_SIZE_M == 0
|
| 512 |
+
if not m_divisible_by_block_m:
|
| 513 |
+
loop_m -= 1
|
| 514 |
+
|
| 515 |
+
acc = tl.zeros((BLOCK_SIZE_K, BLOCK_SIZE_N), dtype=tl.float32)
|
| 516 |
+
# Initialize bias accumulator
|
| 517 |
+
bias_acc = tl.zeros((BLOCK_SIZE_K,), dtype=tl.float32)
|
| 518 |
+
|
| 519 |
+
for _ in range(0, loop_m):
|
| 520 |
+
lhs = tl.load(lhs_ptrs)
|
| 521 |
+
rhs = tl.load(rhs_ptrs)
|
| 522 |
+
|
| 523 |
+
acc = tl.dot(lhs, rhs, acc=acc)
|
| 524 |
+
|
| 525 |
+
# Accumulate for bias gradient: sum lhs across M dimension
|
| 526 |
+
if COMPUTE_BIAS_GRAD and tile_n == 0:
|
| 527 |
+
bias_acc += tl.sum(lhs, axis=1) # [K, M] -> [K]
|
| 528 |
+
|
| 529 |
+
if TRANS_LHS:
|
| 530 |
+
lhs_ptrs += BLOCK_SIZE_M * K
|
| 531 |
+
else:
|
| 532 |
+
lhs_ptrs += BLOCK_SIZE_M
|
| 533 |
+
|
| 534 |
+
rhs_ptrs += BLOCK_SIZE_M * N
|
| 535 |
+
|
| 536 |
+
if not m_divisible_by_block_m:
|
| 537 |
+
offs_lhs_k = (
|
| 538 |
+
tile_k.to(tl.int64) * BLOCK_SIZE_K + tl.arange(0, BLOCK_SIZE_K)
|
| 539 |
+
) % K
|
| 540 |
+
offs_rhs_n = (
|
| 541 |
+
tile_n.to(tl.int64) * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
|
| 542 |
+
) % N
|
| 543 |
+
offs_m = loop_m.to(tl.int64) * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
|
| 544 |
+
lhs = tl.load(lhs_ptrs, mask=offs_m[None, :] < m, other=0)
|
| 545 |
+
rhs = tl.load(rhs_ptrs, mask=offs_m[:, None] < m, other=0)
|
| 546 |
+
acc = tl.dot(lhs, rhs, acc=acc)
|
| 547 |
+
# Accumulate last chunk for bias gradient
|
| 548 |
+
if COMPUTE_BIAS_GRAD and tile_n == 0:
|
| 549 |
+
bias_acc += tl.sum(lhs, axis=1)
|
| 550 |
+
|
| 551 |
+
acc = acc.to(out_ptr.type.element_ty)
|
| 552 |
+
|
| 553 |
+
offs_out_k = tile_k.to(tl.int64) * BLOCK_SIZE_K + tl.arange(0, BLOCK_SIZE_K)
|
| 554 |
+
offs_out_n = tile_n.to(tl.int64) * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
|
| 555 |
+
|
| 556 |
+
out_ptrs = (
|
| 557 |
+
out_ptr + g.to(tl.int64) * K * N + offs_out_k[:, None] * N + offs_out_n[None, :]
|
| 558 |
+
)
|
| 559 |
+
|
| 560 |
+
mask = (offs_out_k[:, None] < K) & (offs_out_n[None, :] < N)
|
| 561 |
+
if ACCUMULATE:
|
| 562 |
+
# Load existing values and add to them (like beta=1 in BLAS)
|
| 563 |
+
old_vals = tl.load(out_ptrs, mask=mask, other=0.0)
|
| 564 |
+
tl.store(out_ptrs, acc + old_vals, mask=mask)
|
| 565 |
+
else:
|
| 566 |
+
# Overwrite output (like beta=0 in BLAS)
|
| 567 |
+
tl.store(out_ptrs, acc, mask=mask)
|
| 568 |
+
|
| 569 |
+
# Store bias gradient (only for first N tile, sum across all M)
|
| 570 |
+
if COMPUTE_BIAS_GRAD and tile_n == 0:
|
| 571 |
+
# Keep as float32 for atomic_add (bf16/fp16 not supported for atomics)
|
| 572 |
+
bias_grad_ptrs = bias_grad_ptr + g.to(tl.int64) * K + offs_out_k
|
| 573 |
+
# Use atomic add since multiple K-tiles may write to same expert's bias
|
| 574 |
+
tl.atomic_add(bias_grad_ptrs, bias_acc, mask=offs_out_k < K, sem="relaxed")
|
build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/adapter.py
ADDED
|
@@ -0,0 +1,53 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
"""Adapt AITER's Triton grouped GEMM to MegaBlocks' ``gmm`` calling convention.
|
| 3 |
+
|
| 4 |
+
MegaBlocks (following tgale96/grouped_gemm) uses a single ``gmm`` entry point
|
| 5 |
+
with ``trans_a`` / ``trans_b`` flags:
|
| 6 |
+
|
| 7 |
+
* ``trans_a=False, trans_b=False``: a(M,K) @ b(G,K,N) -> c(M,N)
|
| 8 |
+
* ``trans_a=False, trans_b=True`` : a(M,K) @ b(G,N,K)^T -> c(M,N) (dgrad)
|
| 9 |
+
* ``trans_a=True`` : a(M,K)^T @ b(M,N) per group -> c(G,K,N) (wgrad)
|
| 10 |
+
|
| 11 |
+
AITER exposes these as two kernels: ``gmm`` ((M,K)@(G,K,N)->(M,N), transposition
|
| 12 |
+
of the 3D operand inferred from strides) and ``ptgmm`` ((K,M)@(M,N)->(G,K,N),
|
| 13 |
+
transposition of the 2D operand inferred from strides).
|
| 14 |
+
"""
|
| 15 |
+
|
| 16 |
+
import torch
|
| 17 |
+
|
| 18 |
+
from .gmm import gmm as _aiter_gmm
|
| 19 |
+
from .gmm import ptgmm as _aiter_ptgmm
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
def gmm(a, b, c, batch_sizes, trans_a=False, trans_b=False):
|
| 23 |
+
# AITER requires group sizes to be int32 and to live on the compute device.
|
| 24 |
+
group_sizes = batch_sizes.to(device=a.device, dtype=torch.int32)
|
| 25 |
+
|
| 26 |
+
# AITER asserts exact strides: gmm wants lhs/rhs row-major (a transposed
|
| 27 |
+
# 3D operand must be exactly column-major), tgmm wants rhs row-major and
|
| 28 |
+
# lhs row/column-major. Make operands contiguous first so the transposed
|
| 29 |
+
# views have the precise strides the kernels expect. `.contiguous()` is a
|
| 30 |
+
# no-op when the tensor is already contiguous.
|
| 31 |
+
if trans_a:
|
| 32 |
+
# Weight gradient: a(M,K), b(M,N) -> c(G,K,N).
|
| 33 |
+
# Pass a transposed so AITER sees lhs(K,M) column-major (TRANS_LHS).
|
| 34 |
+
_aiter_ptgmm(
|
| 35 |
+
a.contiguous().transpose(0, 1),
|
| 36 |
+
b.contiguous(),
|
| 37 |
+
group_sizes,
|
| 38 |
+
preferred_element_type=c.dtype,
|
| 39 |
+
existing_out=c,
|
| 40 |
+
)
|
| 41 |
+
else:
|
| 42 |
+
# trans_b contracts b's last dim: pass a column-major (G,K,N) view.
|
| 43 |
+
rhs = b.contiguous()
|
| 44 |
+
if trans_b:
|
| 45 |
+
rhs = rhs.transpose(1, 2)
|
| 46 |
+
_aiter_gmm(
|
| 47 |
+
a.contiguous(),
|
| 48 |
+
rhs,
|
| 49 |
+
group_sizes,
|
| 50 |
+
preferred_element_type=c.dtype,
|
| 51 |
+
existing_out=c,
|
| 52 |
+
)
|
| 53 |
+
return c
|
build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/configs.py
ADDED
|
@@ -0,0 +1,5 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: MIT
|
| 2 |
+
# Tuned GMM configs vendored from ROCm/aiter (aiter/ops/triton/configs/).
|
| 3 |
+
# Inlined as a Python module so packaging always includes them.
|
| 4 |
+
|
| 5 |
+
CONFIGS = {'gfx1250': {'gmm': {'default': {'BLOCK_SIZE_M': 256, 'BLOCK_SIZE_K': 64, 'BLOCK_SIZE_N': 256, 'GROUP_SIZE': 1, 'GRID_DIM': 256, 'num_warps': 8, 'num_stages': 1}}, 'ptgmm': {'default': {'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_K': 256, 'BLOCK_SIZE_N': 256, 'GROUP_SIZE': 1, 'GRID_DIM': 256, 'num_warps': 8, 'num_stages': 1}, 'accumulate': {'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_K': 128, 'BLOCK_SIZE_N': 128, 'GROUP_SIZE': 1, 'GRID_DIM': 256, 'num_warps': 8, 'num_stages': 1}}, 'nptgmm': {'default': {'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_K': 256, 'BLOCK_SIZE_N': 256, 'GROUP_SIZE': 1, 'num_warps': 8, 'num_stages': 1}, 'accumulate': {'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_K': 128, 'BLOCK_SIZE_N': 128, 'GROUP_SIZE': 1, 'num_warps': 8, 'num_stages': 1}}}, 'gfx942': {'gmm': {'default': {'BLOCK_SIZE_M': 256, 'BLOCK_SIZE_K': 64, 'BLOCK_SIZE_N': 256, 'GROUP_SIZE': 1, 'GRID_DIM': 304, 'num_warps': 8, 'num_stages': 1}}, 'ptgmm': {'default': {'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_K': 256, 'BLOCK_SIZE_N': 256, 'GROUP_SIZE': 1, 'GRID_DIM': 304, 'num_warps': 8, 'num_stages': 1}, 'accumulate': {'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_K': 128, 'BLOCK_SIZE_N': 128, 'GROUP_SIZE': 1, 'GRID_DIM': 304, 'num_warps': 8, 'num_stages': 1}}, 'nptgmm': {'default': {'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_K': 256, 'BLOCK_SIZE_N': 256, 'GROUP_SIZE': 1, 'num_warps': 8, 'num_stages': 1}, 'accumulate': {'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_K': 128, 'BLOCK_SIZE_N': 128, 'GROUP_SIZE': 1, 'num_warps': 8, 'num_stages': 1}}}, 'gfx950': {'gmm': {'default': {'BLOCK_SIZE_M': 256, 'BLOCK_SIZE_K': 64, 'BLOCK_SIZE_N': 256, 'GROUP_SIZE': 1, 'GRID_DIM': 256, 'num_warps': 8, 'num_stages': 1}}, 'ptgmm': {'default': {'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_K': 256, 'BLOCK_SIZE_N': 256, 'GROUP_SIZE': 1, 'GRID_DIM': 256, 'num_warps': 8, 'num_stages': 1}, 'accumulate': {'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_K': 128, 'BLOCK_SIZE_N': 128, 'GROUP_SIZE': 1, 'GRID_DIM': 256, 'num_warps': 8, 'num_stages': 1}}, 'nptgmm': {'default': {'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_K': 256, 'BLOCK_SIZE_N': 256, 'GROUP_SIZE': 1, 'num_warps': 8, 'num_stages': 1}, 'accumulate': {'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_K': 128, 'BLOCK_SIZE_N': 128, 'GROUP_SIZE': 1, 'num_warps': 8, 'num_stages': 1}}}}
|
build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/gmm.py
ADDED
|
@@ -0,0 +1,567 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: MIT
|
| 2 |
+
# Copyright (C) 2025, Advanced Micro Devices, Inc. All rights reserved.
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
# Imports.
|
| 6 |
+
# ------------------------------------------------------------------------------
|
| 7 |
+
|
| 8 |
+
# PyTorch
|
| 9 |
+
import torch
|
| 10 |
+
from torch import Tensor
|
| 11 |
+
|
| 12 |
+
# Triton
|
| 13 |
+
import triton
|
| 14 |
+
|
| 15 |
+
# AITER: GMM utility functions
|
| 16 |
+
from .utils.gmm_common import (
|
| 17 |
+
DTYPE,
|
| 18 |
+
is_power_of_2,
|
| 19 |
+
check_input_device_dtype,
|
| 20 |
+
check_bias_shape_stride,
|
| 21 |
+
get_gmm_shape,
|
| 22 |
+
get_gmm_output,
|
| 23 |
+
get_gmm_transposition,
|
| 24 |
+
get_tgmm_shape,
|
| 25 |
+
get_tgmm_output,
|
| 26 |
+
get_tgmm_bias_grad,
|
| 27 |
+
get_tgmm_transposition,
|
| 28 |
+
)
|
| 29 |
+
|
| 30 |
+
# AITER: GMM Triton kernels
|
| 31 |
+
from ._triton_kernels.gmm import (
|
| 32 |
+
gmm_kernel,
|
| 33 |
+
tgmm_persistent_kernel,
|
| 34 |
+
tgmm_non_persistent_kernel,
|
| 35 |
+
get_config,
|
| 36 |
+
)
|
| 37 |
+
|
| 38 |
+
# GMM PyTorch wrapper.
|
| 39 |
+
# ------------------------------------------------------------------------------
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
def _gmm_grid(
|
| 43 |
+
N: int,
|
| 44 |
+
block_size_m: int,
|
| 45 |
+
block_size_n: int,
|
| 46 |
+
group_sizes: Tensor,
|
| 47 |
+
grid_dim: int,
|
| 48 |
+
) -> tuple[int]:
|
| 49 |
+
assert N > 0, f"N must be positive, it's {N}."
|
| 50 |
+
assert is_power_of_2(
|
| 51 |
+
block_size_m
|
| 52 |
+
), f"M-dimension tile size must be a power of 2 (it's {block_size_m})."
|
| 53 |
+
assert is_power_of_2(
|
| 54 |
+
block_size_n
|
| 55 |
+
), f"N-dimension tile size must be a power of 2 (it's {block_size_n})."
|
| 56 |
+
assert torch.all(group_sizes >= 0).item(), "All group_sizes must be non-negative."
|
| 57 |
+
assert grid_dim > 0, f"Grid dimension must be positive (it's {grid_dim})."
|
| 58 |
+
num_m_tiles = (group_sizes + block_size_m - 1) // block_size_m
|
| 59 |
+
assert torch.all(num_m_tiles >= 0).item(), "All num_m_tiles must be non-negative."
|
| 60 |
+
num_n_tiles = triton.cdiv(N, block_size_n)
|
| 61 |
+
assert num_n_tiles > 0, f"num_n_tiles must be positive, it's {num_n_tiles}."
|
| 62 |
+
num_tiles = torch.sum(num_m_tiles * num_n_tiles).item()
|
| 63 |
+
assert num_tiles > 0, f"num_tiles must be positive, it's {num_tiles}."
|
| 64 |
+
num_programs = int(min(grid_dim, num_tiles))
|
| 65 |
+
assert num_programs > 0, f"num_programs must be positive, it's {num_programs}."
|
| 66 |
+
return (num_programs,)
|
| 67 |
+
|
| 68 |
+
|
| 69 |
+
def gmm(
|
| 70 |
+
lhs: Tensor,
|
| 71 |
+
rhs: Tensor,
|
| 72 |
+
group_sizes: Tensor,
|
| 73 |
+
preferred_element_type: torch.dtype = DTYPE,
|
| 74 |
+
existing_out: Tensor | None = None,
|
| 75 |
+
config: dict[str, int] | None = None,
|
| 76 |
+
bias: Tensor | None = None,
|
| 77 |
+
) -> Tensor:
|
| 78 |
+
"""
|
| 79 |
+
Perform Group Matrix Multiplication (GMM): out = lhs @ rhs + bias
|
| 80 |
+
|
| 81 |
+
lhs rows are divided into G groups. Each group of lhs rows is matrix multiplied with a plane of
|
| 82 |
+
rhs 3D tensor and then stored in a slice of out. In PyTorch parlance, it can be implemented as
|
| 83 |
+
follows for a given group g:
|
| 84 |
+
out[group_start:group_end, :] = lhs[group_start:group_end, :] @ rhs[g] + bias[g]
|
| 85 |
+
|
| 86 |
+
The size of each group, and their respective start and end positions are specified by
|
| 87 |
+
group_sizes tensor. For instance, suppose that group_sizes = [3, 2, 4, 1]. In this particular
|
| 88 |
+
case we have 4 groups. The 1st group starts at 0 and ends at 2, the second group starts at 3 and
|
| 89 |
+
ends at 4, the third group starts at 5 and ends at 8, and the fourth and final group consists of
|
| 90 |
+
just the 10th (last) row of lhs.
|
| 91 |
+
|
| 92 |
+
Parameters
|
| 93 |
+
----------
|
| 94 |
+
lhs : torch.Tensor
|
| 95 |
+
Left-hand side 2D input tensor. Shape: (M, K).
|
| 96 |
+
lhs data type must be torch.float16 or torch.bfloat16, and must match rhs data type.
|
| 97 |
+
lhs must be on the same device of rhs and group_sizes.
|
| 98 |
+
rhs : torch.Tensor
|
| 99 |
+
Right-hand side 3D input tensor. Shape: (G, K, N).
|
| 100 |
+
rhs data type must be torch.float16 or torch.bfloat16, and must match lhs data type.
|
| 101 |
+
rhs must be on the same device of lhs and group_sizes.
|
| 102 |
+
group_sizes : torch.Tensor
|
| 103 |
+
1D input tensor describing group sizes. Shape: (G,).
|
| 104 |
+
group_sizes data type must be torch.int32 and all its elements must be non-negative.
|
| 105 |
+
group_sizes must be on the same device of lhs and rhs.
|
| 106 |
+
preferred_element_type : torch.dtype, optional
|
| 107 |
+
Desired data type for output tensor. Default is torch.bfloat16.
|
| 108 |
+
Supported output types are torch.float16 and torch.bfloat16.
|
| 109 |
+
existing_out : torch.Tensor or None, optional
|
| 110 |
+
Preallocated output tensor. Default is None.
|
| 111 |
+
If provided, results are written into this tensor. Otherwise, a new output tensor is
|
| 112 |
+
allocated.
|
| 113 |
+
If provided then it must have shape (M, N), its data type must match preferred_element_type
|
| 114 |
+
and it must be on the same device of other input tensors.
|
| 115 |
+
config : dict[str, int] or None, optional
|
| 116 |
+
Optional dictionary with kernel metaparameters. If absent, config will be queried from
|
| 117 |
+
internal tuning database.
|
| 118 |
+
bias : torch.Tensor or None, optional
|
| 119 |
+
Optional bias tensor. Shape: (G, N).
|
| 120 |
+
If provided, bias data type must match lhs and rhs data type, and bias must be on the same
|
| 121 |
+
device as other input tensors. Each group g adds bias[g] to the output.
|
| 122 |
+
|
| 123 |
+
Returns
|
| 124 |
+
-------
|
| 125 |
+
torch.Tensor
|
| 126 |
+
The computed output 2D tensor. Shape: (M, N).
|
| 127 |
+
Output tensor data type is given by preferred_element_type.
|
| 128 |
+
If existing_out is provided then existing_out is also returned.
|
| 129 |
+
|
| 130 |
+
Implementation Notes
|
| 131 |
+
--------------------
|
| 132 |
+
- GMM is implemented with a persistent Triton kernel.
|
| 133 |
+
- lhs must be row-major (lhs.stride() == (K, 1)).
|
| 134 |
+
- rhs can be row-major (rhs.stride() == (K * N, N, 1)) or column-major (rhs.stride() ==
|
| 135 |
+
(K * N, 1, K)). If rhs is row-major then kernel parameter TRANS_RHS == False, this is useful
|
| 136 |
+
for implementing forward pass. If rhs is column-major then kernel parameter TRANS_RHS == True,
|
| 137 |
+
this is useful for computing the lhs derivative in the backward pass, while fusing the
|
| 138 |
+
transposition.
|
| 139 |
+
- out must be row-major (out.stride() == (N, 1)).
|
| 140 |
+
- bias must be row-major (bias.stride() == (N, 1)) if provided.
|
| 141 |
+
"""
|
| 142 |
+
use_bias = bias is not None
|
| 143 |
+
check_input_device_dtype(lhs, rhs, group_sizes, bias)
|
| 144 |
+
|
| 145 |
+
M, K, N, G = get_gmm_shape(lhs, rhs, group_sizes)
|
| 146 |
+
|
| 147 |
+
if use_bias:
|
| 148 |
+
check_bias_shape_stride(bias, G, N)
|
| 149 |
+
|
| 150 |
+
out = get_gmm_output(
|
| 151 |
+
M,
|
| 152 |
+
N,
|
| 153 |
+
device=lhs.device,
|
| 154 |
+
preferred_element_type=preferred_element_type,
|
| 155 |
+
existing_out=existing_out,
|
| 156 |
+
)
|
| 157 |
+
|
| 158 |
+
trans_rhs, _ = get_gmm_transposition(lhs, rhs, out)
|
| 159 |
+
|
| 160 |
+
if config is None:
|
| 161 |
+
config = get_config("gmm", M, K, N, G)
|
| 162 |
+
|
| 163 |
+
assert all(
|
| 164 |
+
key in config
|
| 165 |
+
and isinstance(config[key], int)
|
| 166 |
+
and (
|
| 167 |
+
is_power_of_2(config[key])
|
| 168 |
+
if key.startswith("BLOCK_SIZE_")
|
| 169 |
+
else config[key] > 0
|
| 170 |
+
)
|
| 171 |
+
for key in {
|
| 172 |
+
"BLOCK_SIZE_M",
|
| 173 |
+
"BLOCK_SIZE_K",
|
| 174 |
+
"BLOCK_SIZE_N",
|
| 175 |
+
"GROUP_SIZE",
|
| 176 |
+
"GRID_DIM",
|
| 177 |
+
}
|
| 178 |
+
), "Invalid GMM kernel config."
|
| 179 |
+
|
| 180 |
+
grid = _gmm_grid(
|
| 181 |
+
N,
|
| 182 |
+
config["BLOCK_SIZE_M"],
|
| 183 |
+
config["BLOCK_SIZE_N"],
|
| 184 |
+
group_sizes,
|
| 185 |
+
config["GRID_DIM"],
|
| 186 |
+
)
|
| 187 |
+
|
| 188 |
+
# fmt: off
|
| 189 |
+
gmm_kernel[grid](
|
| 190 |
+
# Tensor pointers:
|
| 191 |
+
lhs, rhs, group_sizes, out, bias,
|
| 192 |
+
# Tensor shapes:
|
| 193 |
+
M, K, N, G,
|
| 194 |
+
# Meta-parameters:
|
| 195 |
+
TRANS_RHS=trans_rhs,
|
| 196 |
+
USE_BIAS=use_bias,
|
| 197 |
+
**config,
|
| 198 |
+
)
|
| 199 |
+
# fmt: on
|
| 200 |
+
|
| 201 |
+
return out
|
| 202 |
+
|
| 203 |
+
|
| 204 |
+
# Persistent TGMM PyTorch wrapper.
|
| 205 |
+
# ------------------------------------------------------------------------------
|
| 206 |
+
|
| 207 |
+
|
| 208 |
+
def _ptgmm_grid(
|
| 209 |
+
K: int,
|
| 210 |
+
N: int,
|
| 211 |
+
G: int,
|
| 212 |
+
block_size_k: int,
|
| 213 |
+
block_size_n: int,
|
| 214 |
+
grid_dim: int,
|
| 215 |
+
) -> tuple[int]:
|
| 216 |
+
assert K > 0, f"K must be positive, it's {K}."
|
| 217 |
+
assert N > 0, f"N must be positive, it's {N}."
|
| 218 |
+
assert G > 0, f"G must be positive, it's {G}."
|
| 219 |
+
assert is_power_of_2(
|
| 220 |
+
block_size_k
|
| 221 |
+
), f"K-dimension tile size must be a power of 2 (it's {block_size_k})."
|
| 222 |
+
assert is_power_of_2(
|
| 223 |
+
block_size_n
|
| 224 |
+
), f"N-dimension tile size must be a power of 2 (it's {block_size_n})."
|
| 225 |
+
assert grid_dim > 0, f"Grid dimension must be positive (it's {grid_dim})."
|
| 226 |
+
num_k_tiles = triton.cdiv(K, block_size_k)
|
| 227 |
+
assert num_k_tiles > 0, f"num_k_tiles must be positive, it's {num_k_tiles}."
|
| 228 |
+
num_n_tiles = triton.cdiv(N, block_size_n)
|
| 229 |
+
assert num_n_tiles > 0, f"num_n_tiles must be positive, it's {num_n_tiles}."
|
| 230 |
+
num_tiles = G * num_k_tiles * num_n_tiles
|
| 231 |
+
assert num_tiles > 0, f"num_tiles must be positive, it's {num_tiles}."
|
| 232 |
+
num_programs = min(grid_dim, num_tiles)
|
| 233 |
+
assert num_programs > 0, f"num_programs must be positive, it's {num_programs}."
|
| 234 |
+
return (num_programs,)
|
| 235 |
+
|
| 236 |
+
|
| 237 |
+
def ptgmm(
|
| 238 |
+
lhs: Tensor,
|
| 239 |
+
rhs: Tensor,
|
| 240 |
+
group_sizes: Tensor,
|
| 241 |
+
preferred_element_type: torch.dtype = DTYPE,
|
| 242 |
+
existing_out: Tensor | None = None,
|
| 243 |
+
config: dict[str, int] | None = None,
|
| 244 |
+
bias_grad: Tensor | None = None,
|
| 245 |
+
accumulate: bool = False,
|
| 246 |
+
) -> Tensor:
|
| 247 |
+
"""
|
| 248 |
+
Perform a Group Matrix Multiplication (GMM) variant: out = lhs @ rhs
|
| 249 |
+
|
| 250 |
+
lhs columns and rhs rows are divided into G groups. Each group of lhs is matrix multiplied with
|
| 251 |
+
the respective group of rhs and then stored in a plane of the output 3D tensor. In PyTorch
|
| 252 |
+
parlance, it can be implemented as follows for a given group g:
|
| 253 |
+
out[g] = lhs[:, group_start:group_end] @ rhs[group_start:group_end, :]
|
| 254 |
+
|
| 255 |
+
The 't' in the operator name derives from MaxText implementation
|
| 256 |
+
(https://github.com/AI-Hypercomputer/maxtext/blob/main/src/MaxText/kernels/megablox/gmm.py),
|
| 257 |
+
which served as the initial inspiration for this one. TGMM differs from GMM in terms of tensor
|
| 258 |
+
shapes. GMM does (M, K) @ (G, K, N) = (M, N) while TGMM does (K, M) @ (M, N) = (G, K, N).
|
| 259 |
+
|
| 260 |
+
The 'p' in the operator name means that it is implemented with a persistent kernel. There is
|
| 261 |
+
also the non-persistent variation, which is implemented with a regular kernel. Please take a
|
| 262 |
+
look at nptgmm operator. Both ptgmm and nptgmm implement the same computation, choosing one or
|
| 263 |
+
the other is a matter of performance for the target workload.
|
| 264 |
+
|
| 265 |
+
Parameters
|
| 266 |
+
----------
|
| 267 |
+
lhs : torch.Tensor
|
| 268 |
+
Left-hand side 2D input tensor. Shape: (K, M).
|
| 269 |
+
lhs data type must be torch.float16 or torch.bfloat16, and must match rhs data type.
|
| 270 |
+
lhs must be on the same device of rhs and group_sizes.
|
| 271 |
+
rhs : torch.Tensor
|
| 272 |
+
Right-hand side 2D input tensor. Shape: (M, N).
|
| 273 |
+
rhs data type must be torch.float16 or torch.bfloat16, and must match lhs data type.
|
| 274 |
+
rhs must be on the same device of lhs and group_sizes.
|
| 275 |
+
group_sizes : torch.Tensor
|
| 276 |
+
1D input tensor describing group sizes. Shape: (G,).
|
| 277 |
+
group_sizes data type must be torch.int32 and all its elements must be non-negative.
|
| 278 |
+
group_sizes must be on the same device of lhs and rhs.
|
| 279 |
+
preferred_element_type : torch.dtype, optional
|
| 280 |
+
Desired data type for output tensor. Default is torch.bfloat16.
|
| 281 |
+
Supported output types are torch.float16 and torch.bfloat16.
|
| 282 |
+
existing_out : torch.Tensor or None, optional
|
| 283 |
+
Preallocated output tensor. Default is None.
|
| 284 |
+
If provided, results are written into this tensor. Otherwise, a new output tensor is
|
| 285 |
+
allocated.
|
| 286 |
+
If provided then it must have shape (G, K, N), its data type must match
|
| 287 |
+
preferred_element_type and it must be on the same device of other input tensors.
|
| 288 |
+
config : dict[str, int] or None, optional
|
| 289 |
+
Optional dictionary with kernel metaparameters. If absent, config will be queried from
|
| 290 |
+
internal tuning database.
|
| 291 |
+
bias_grad : torch.Tensor or None, optional
|
| 292 |
+
Optional bias gradient output tensor. Shape: (G, K).
|
| 293 |
+
If provided, the kernel will compute the bias gradient and write it to this tensor.
|
| 294 |
+
bias_grad must be torch.float32 (kernel uses atomic_add which requires float32),
|
| 295 |
+
accumulate : bool, optional
|
| 296 |
+
Whether to accumulate into existing output tensor values. Default is False.
|
| 297 |
+
If False, output will be overwritten with fresh computation.
|
| 298 |
+
If True, results will be added to existing output tensor values.
|
| 299 |
+
|
| 300 |
+
Returns
|
| 301 |
+
-------
|
| 302 |
+
torch.Tensor
|
| 303 |
+
The computed output 3D tensor. Shape: (G, K, N).
|
| 304 |
+
Output tensor data type is given by preferred_element_type.
|
| 305 |
+
If existing_out is provided then existing_out is also returned.
|
| 306 |
+
|
| 307 |
+
Implementation Notes
|
| 308 |
+
--------------------
|
| 309 |
+
- PTGMM is implemented with a persistent Triton kernel.
|
| 310 |
+
- lhs can be row-major (lhs.stride() == (M, 1)) or column-major (lhs.stride() == (1, K)). If lhs
|
| 311 |
+
is row-major then kernel parameter TRANS_LHS == False. If lhs is column-major then kernel
|
| 312 |
+
parameter TRANS_LHS == True, this is useful for computing the rhs derivative in the backward
|
| 313 |
+
pass, while fusing the transposition.
|
| 314 |
+
- rhs must be row-major (rhs.stride() == (N, 1)).
|
| 315 |
+
- out must be row-major (out.stride() == (K * N, N, 1)).
|
| 316 |
+
"""
|
| 317 |
+
check_input_device_dtype(lhs, rhs, group_sizes)
|
| 318 |
+
|
| 319 |
+
M, K, N, G = get_tgmm_shape(lhs, rhs, group_sizes)
|
| 320 |
+
|
| 321 |
+
out = get_tgmm_output(
|
| 322 |
+
K,
|
| 323 |
+
N,
|
| 324 |
+
G,
|
| 325 |
+
device=lhs.device,
|
| 326 |
+
preferred_element_type=preferred_element_type,
|
| 327 |
+
existing_out=existing_out,
|
| 328 |
+
)
|
| 329 |
+
|
| 330 |
+
trans_lhs, _ = get_tgmm_transposition(lhs, rhs, out)
|
| 331 |
+
|
| 332 |
+
if config is None:
|
| 333 |
+
config = get_config("ptgmm", M, K, N, G, accumulate)
|
| 334 |
+
|
| 335 |
+
assert all(
|
| 336 |
+
key in config
|
| 337 |
+
and isinstance(config[key], int)
|
| 338 |
+
and (
|
| 339 |
+
is_power_of_2(config[key])
|
| 340 |
+
if key.startswith("BLOCK_SIZE_")
|
| 341 |
+
else config[key] > 0
|
| 342 |
+
)
|
| 343 |
+
for key in {
|
| 344 |
+
"BLOCK_SIZE_M",
|
| 345 |
+
"BLOCK_SIZE_K",
|
| 346 |
+
"BLOCK_SIZE_N",
|
| 347 |
+
"GROUP_SIZE",
|
| 348 |
+
"GRID_DIM",
|
| 349 |
+
}
|
| 350 |
+
), "Invalid PTGMM kernel config."
|
| 351 |
+
|
| 352 |
+
# Bias gradient handling.
|
| 353 |
+
# -----------------------
|
| 354 |
+
# Get or validate bias gradient tensor.
|
| 355 |
+
compute_bias_grad = bias_grad is not None
|
| 356 |
+
bias_grad_ptr = get_tgmm_bias_grad(
|
| 357 |
+
K,
|
| 358 |
+
G,
|
| 359 |
+
device=lhs.device,
|
| 360 |
+
existing_bias_grad=bias_grad,
|
| 361 |
+
)
|
| 362 |
+
|
| 363 |
+
grid = _ptgmm_grid(
|
| 364 |
+
K,
|
| 365 |
+
N,
|
| 366 |
+
G,
|
| 367 |
+
config["BLOCK_SIZE_K"],
|
| 368 |
+
config["BLOCK_SIZE_N"],
|
| 369 |
+
config["GRID_DIM"],
|
| 370 |
+
)
|
| 371 |
+
|
| 372 |
+
# fmt: off
|
| 373 |
+
tgmm_persistent_kernel[grid](
|
| 374 |
+
# Tensor pointers:
|
| 375 |
+
lhs, rhs, group_sizes, out, bias_grad_ptr,
|
| 376 |
+
# Tensor shapes:
|
| 377 |
+
M, K, N, G,
|
| 378 |
+
# Meta-parameters:
|
| 379 |
+
TRANS_LHS=trans_lhs,
|
| 380 |
+
COMPUTE_BIAS_GRAD=compute_bias_grad,
|
| 381 |
+
ACCUMULATE=accumulate,
|
| 382 |
+
**config,
|
| 383 |
+
)
|
| 384 |
+
# fmt: on
|
| 385 |
+
|
| 386 |
+
return out
|
| 387 |
+
|
| 388 |
+
|
| 389 |
+
# Regular non-persistent TGMM PyTorch wrapper.
|
| 390 |
+
# ------------------------------------------------------------------------------
|
| 391 |
+
|
| 392 |
+
|
| 393 |
+
def _nptgmm_grid(
|
| 394 |
+
K: int,
|
| 395 |
+
N: int,
|
| 396 |
+
G: int,
|
| 397 |
+
block_size_k: int,
|
| 398 |
+
block_size_n: int,
|
| 399 |
+
) -> tuple[int, int]:
|
| 400 |
+
assert K > 0, f"K must be positive, it's {K}."
|
| 401 |
+
assert N > 0, f"N must be positive, it's {N}."
|
| 402 |
+
assert G > 0, f"G must be positive, it's {G}."
|
| 403 |
+
assert is_power_of_2(
|
| 404 |
+
block_size_k
|
| 405 |
+
), f"K-dimension tile size must be a power of 2 (it's {block_size_k})."
|
| 406 |
+
assert is_power_of_2(
|
| 407 |
+
block_size_n
|
| 408 |
+
), f"N-dimension tile size must be a power of 2 (it's {block_size_n})."
|
| 409 |
+
num_k_tiles = triton.cdiv(K, block_size_k)
|
| 410 |
+
assert num_k_tiles > 0, f"num_k_tiles must be positive, it's {num_k_tiles}."
|
| 411 |
+
num_n_tiles = triton.cdiv(N, block_size_n)
|
| 412 |
+
assert num_n_tiles > 0, f"num_n_tiles must be positive, it's {num_n_tiles}."
|
| 413 |
+
num_tiles_per_mm = num_k_tiles * num_n_tiles
|
| 414 |
+
assert (
|
| 415 |
+
num_tiles_per_mm > 0
|
| 416 |
+
), f"num_tiles_per_mm must be positive, it's {num_tiles_per_mm}."
|
| 417 |
+
return (G, num_tiles_per_mm)
|
| 418 |
+
|
| 419 |
+
|
| 420 |
+
def nptgmm(
|
| 421 |
+
lhs: Tensor,
|
| 422 |
+
rhs: Tensor,
|
| 423 |
+
group_sizes: Tensor,
|
| 424 |
+
preferred_element_type: torch.dtype = DTYPE,
|
| 425 |
+
existing_out: Tensor | None = None,
|
| 426 |
+
config: dict[str, int] | None = None,
|
| 427 |
+
bias_grad: Tensor | None = None,
|
| 428 |
+
accumulate: bool = False,
|
| 429 |
+
) -> Tensor:
|
| 430 |
+
"""
|
| 431 |
+
Perform a Group Matrix Multiplication (GMM) variant: out = lhs @ rhs
|
| 432 |
+
|
| 433 |
+
lhs columns and rhs rows are divided into G groups. Each group of lhs is matrix multiplied with
|
| 434 |
+
the respective group of rhs and then stored in a plane of the output 3D tensor. In PyTorch
|
| 435 |
+
parlance, it can be implemented as follows for a given group g:
|
| 436 |
+
out[g] = lhs[:, group_start:group_end] @ rhs[group_start:group_end, :]
|
| 437 |
+
|
| 438 |
+
The 't' in the operator name derives from MaxText implementation
|
| 439 |
+
(https://github.com/AI-Hypercomputer/maxtext/blob/main/src/MaxText/kernels/megablox/gmm.py),
|
| 440 |
+
which served as the initial inspiration for this one. TGMM differs from GMM in terms of tensor
|
| 441 |
+
shapes. GMM does (M, K) @ (G, K, N) = (M, N) while TGMM does (K, M) @ (M, N) = (G, K, N).
|
| 442 |
+
|
| 443 |
+
The 'np' in the operator name means that it is implemented with a non-persistent, i.e. regular
|
| 444 |
+
kernel. There is also the persistent variation, which is implemented with a persistent kernel.
|
| 445 |
+
Please take a look at ptgmm operator. Both nptgmm and ptgmm implement the same computation,
|
| 446 |
+
choosing one or the other is a matter of performance for the target workload.
|
| 447 |
+
|
| 448 |
+
Parameters
|
| 449 |
+
----------
|
| 450 |
+
lhs : torch.Tensor
|
| 451 |
+
Left-hand side 2D input tensor. Shape: (K, M).
|
| 452 |
+
lhs data type must be torch.float16 or torch.bfloat16, and must match rhs data type.
|
| 453 |
+
lhs must be on the same device of rhs and group_sizes.
|
| 454 |
+
rhs : torch.Tensor
|
| 455 |
+
Right-hand side 2D input tensor. Shape: (M, N).
|
| 456 |
+
rhs data type must be torch.float16 or torch.bfloat16, and must match lhs data type.
|
| 457 |
+
rhs must be on the same device of lhs and group_sizes.
|
| 458 |
+
group_sizes : torch.Tensor
|
| 459 |
+
1D input tensor describing group sizes. Shape: (G,).
|
| 460 |
+
group_sizes data type must be torch.int32 and all its elements must be non-negative.
|
| 461 |
+
group_sizes must be on the same device of lhs and rhs.
|
| 462 |
+
preferred_element_type : torch.dtype, optional
|
| 463 |
+
Desired data type for output tensor. Default is torch.bfloat16.
|
| 464 |
+
Supported output types are torch.float16 and torch.bfloat16.
|
| 465 |
+
existing_out : torch.Tensor or None, optional
|
| 466 |
+
Preallocated output tensor. Default is None.
|
| 467 |
+
If provided, results are written into this tensor. Otherwise, a new output tensor is
|
| 468 |
+
allocated.
|
| 469 |
+
If provided then it must have shape (G, K, N), its data type must match
|
| 470 |
+
preferred_element_type and it must be on the same device of other input tensors.
|
| 471 |
+
config : dict[str, int] or None, optional
|
| 472 |
+
Optional dictionary with kernel metaparameters. If absent, config will be queried from
|
| 473 |
+
internal tuning database.
|
| 474 |
+
bias_grad : torch.Tensor or None, optional
|
| 475 |
+
Optional bias gradient output tensor. Shape: (G, K).
|
| 476 |
+
If provided, the kernel will compute the bias gradient and write it to this tensor.
|
| 477 |
+
bias_grad must be torch.float32 (kernel uses atomic_add which requires float32),
|
| 478 |
+
accumulate : bool, optional
|
| 479 |
+
Whether to accumulate into existing output tensor values. Default is False.
|
| 480 |
+
If False, output will be overwritten with fresh computation.
|
| 481 |
+
If True, results will be added to existing output tensor values.
|
| 482 |
+
|
| 483 |
+
Returns
|
| 484 |
+
-------
|
| 485 |
+
torch.Tensor
|
| 486 |
+
The computed output 3D tensor. Shape: (G, K, N).
|
| 487 |
+
Output tensor data type is given by preferred_element_type.
|
| 488 |
+
If existing_out is provided then existing_out is also returned.
|
| 489 |
+
|
| 490 |
+
Implementation Notes
|
| 491 |
+
--------------------
|
| 492 |
+
- NPTGMM is implemented with a non-persistent regular Triton kernel.
|
| 493 |
+
- lhs can be row-major (lhs.stride() == (M, 1)) or column-major (lhs.stride() == (1, K)). If lhs
|
| 494 |
+
is row-major then kernel parameter TRANS_LHS == False. If lhs is column-major then kernel
|
| 495 |
+
parameter TRANS_LHS == True, this is useful for computing the rhs derivative in the backward
|
| 496 |
+
pass, while fusing the transposition.
|
| 497 |
+
- rhs must be row-major (rhs.stride() == (N, 1)).
|
| 498 |
+
- out must be row-major (out.stride() == (K * N, N, 1)).
|
| 499 |
+
"""
|
| 500 |
+
check_input_device_dtype(lhs, rhs, group_sizes)
|
| 501 |
+
|
| 502 |
+
M, K, N, G = get_tgmm_shape(lhs, rhs, group_sizes)
|
| 503 |
+
|
| 504 |
+
out = get_tgmm_output(
|
| 505 |
+
K,
|
| 506 |
+
N,
|
| 507 |
+
G,
|
| 508 |
+
device=lhs.device,
|
| 509 |
+
preferred_element_type=preferred_element_type,
|
| 510 |
+
existing_out=existing_out,
|
| 511 |
+
)
|
| 512 |
+
|
| 513 |
+
trans_lhs, _ = get_tgmm_transposition(lhs, rhs, out)
|
| 514 |
+
|
| 515 |
+
# Bias gradient handling.
|
| 516 |
+
# -----------------------
|
| 517 |
+
# Get or validate bias gradient tensor.
|
| 518 |
+
compute_bias_grad = bias_grad is not None
|
| 519 |
+
bias_grad_ptr = get_tgmm_bias_grad(
|
| 520 |
+
K,
|
| 521 |
+
G,
|
| 522 |
+
device=lhs.device,
|
| 523 |
+
existing_bias_grad=bias_grad,
|
| 524 |
+
)
|
| 525 |
+
|
| 526 |
+
if config is None:
|
| 527 |
+
config = get_config("nptgmm", M, K, N, G, accumulate)
|
| 528 |
+
|
| 529 |
+
assert all(
|
| 530 |
+
key in config
|
| 531 |
+
and isinstance(config[key], int)
|
| 532 |
+
and (
|
| 533 |
+
is_power_of_2(config[key])
|
| 534 |
+
if key.startswith("BLOCK_SIZE_")
|
| 535 |
+
else config[key] > 0
|
| 536 |
+
)
|
| 537 |
+
for key in {
|
| 538 |
+
"BLOCK_SIZE_M",
|
| 539 |
+
"BLOCK_SIZE_K",
|
| 540 |
+
"BLOCK_SIZE_N",
|
| 541 |
+
"GROUP_SIZE",
|
| 542 |
+
}
|
| 543 |
+
), "Invalid NPTGMM kernel config."
|
| 544 |
+
|
| 545 |
+
grid = _nptgmm_grid(
|
| 546 |
+
K,
|
| 547 |
+
N,
|
| 548 |
+
G,
|
| 549 |
+
config["BLOCK_SIZE_K"],
|
| 550 |
+
config["BLOCK_SIZE_N"],
|
| 551 |
+
)
|
| 552 |
+
|
| 553 |
+
# fmt: off
|
| 554 |
+
tgmm_non_persistent_kernel[grid](
|
| 555 |
+
# Tensor pointers:
|
| 556 |
+
lhs, rhs, group_sizes, out, bias_grad_ptr,
|
| 557 |
+
# Tensor shapes:
|
| 558 |
+
M, K, N, G,
|
| 559 |
+
# Meta-parameters:
|
| 560 |
+
TRANS_LHS=trans_lhs,
|
| 561 |
+
COMPUTE_BIAS_GRAD=compute_bias_grad,
|
| 562 |
+
ACCUMULATE=accumulate,
|
| 563 |
+
**config,
|
| 564 |
+
)
|
| 565 |
+
# fmt: on
|
| 566 |
+
|
| 567 |
+
return out
|
build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/utils/__init__.py
ADDED
|
File without changes
|
build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/utils/_triton/__init__.py
ADDED
|
File without changes
|
build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/utils/_triton/arch_info.py
ADDED
|
@@ -0,0 +1,46 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import triton
|
| 2 |
+
|
| 3 |
+
# Detect the GPU arch lazily: querying the triton driver at import time fails
|
| 4 |
+
# in headless environments (e.g. the kernel-builder ABI check sandbox has no
|
| 5 |
+
# GPU), and the original JAX fallback pulled in an unrelated runtime dep. The
|
| 6 |
+
# arch is only actually needed when a GMM kernel is dispatched, so resolve and
|
| 7 |
+
# cache on first call.
|
| 8 |
+
_CACHED_ARCH = None
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
def get_arch():
|
| 12 |
+
global _CACHED_ARCH
|
| 13 |
+
if _CACHED_ARCH is not None:
|
| 14 |
+
return _CACHED_ARCH
|
| 15 |
+
try:
|
| 16 |
+
_CACHED_ARCH = triton.runtime.driver.active.get_current_target().arch
|
| 17 |
+
except RuntimeError:
|
| 18 |
+
try:
|
| 19 |
+
from jax._src.lib import gpu_triton as triton_kernel_call_lib
|
| 20 |
+
_CACHED_ARCH = triton_kernel_call_lib.get_arch_details("0").split(":")[0]
|
| 21 |
+
except ImportError as e:
|
| 22 |
+
raise RuntimeError(
|
| 23 |
+
"Cannot determine GPU arch: triton driver is inactive and "
|
| 24 |
+
"JAX is not available. A GPU is required for grouped GEMM."
|
| 25 |
+
) from e
|
| 26 |
+
return _CACHED_ARCH
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
def is_gluon_avail():
|
| 30 |
+
return get_arch() in ("gfx950", "gfx1250")
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
def is_fp4_avail():
|
| 34 |
+
return get_arch() in ("gfx950", "gfx1250")
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
def is_fp8_avail():
|
| 38 |
+
return get_arch() in ("gfx942", "gfx950", "gfx1250", "gfx1200", "gfx1201")
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
def is_mx_scale_preshuffling_avail():
|
| 42 |
+
return get_arch() in ("gfx950", "gfx1250")
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
def is_tdm_avail():
|
| 46 |
+
return get_arch() in ("gfx1250",)
|
build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/utils/_triton/pid_preprocessing.py
ADDED
|
@@ -0,0 +1,100 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: MIT
|
| 2 |
+
|
| 3 |
+
# Copyright (C) 2024-2026, Advanced Micro Devices, Inc. All rights reserved.
|
| 4 |
+
|
| 5 |
+
import triton
|
| 6 |
+
import triton.language as tl
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
@triton.jit
|
| 10 |
+
def remap_xcd_chunked(
|
| 11 |
+
pid, GRID_MN, NUM_XCDS: tl.constexpr = 8, CHUNK_SIZE: tl.constexpr = 2
|
| 12 |
+
):
|
| 13 |
+
# Compute current XCD and local PID
|
| 14 |
+
xcd = pid % NUM_XCDS
|
| 15 |
+
# distribute the modulo pids in round robin
|
| 16 |
+
if pid > (GRID_MN // (NUM_XCDS * CHUNK_SIZE)) * (NUM_XCDS * CHUNK_SIZE):
|
| 17 |
+
return pid
|
| 18 |
+
local_pid = pid // NUM_XCDS
|
| 19 |
+
# Calculate chunk index and position within chunk
|
| 20 |
+
chunk_idx = local_pid // CHUNK_SIZE
|
| 21 |
+
pos_in_chunk = local_pid % CHUNK_SIZE
|
| 22 |
+
# Calculate new PID
|
| 23 |
+
new_pid = chunk_idx * NUM_XCDS * CHUNK_SIZE + xcd * CHUNK_SIZE + pos_in_chunk
|
| 24 |
+
return new_pid
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
@triton.jit
|
| 28 |
+
def remap_xcd(pid, GRID_MN, NUM_XCDS: tl.constexpr = 8):
|
| 29 |
+
## pid remapping on xcds
|
| 30 |
+
# Number of pids per XCD in the new arrangement
|
| 31 |
+
pids_per_xcd = (GRID_MN + NUM_XCDS - 1) // NUM_XCDS
|
| 32 |
+
# When GRID_MN cannot divide NUM_XCDS, some xcds will have
|
| 33 |
+
# pids_per_xcd pids, the other will have pids_per_xcd - 1 pids.
|
| 34 |
+
# We calculate the number of xcds that have pids_per_xcd pids as
|
| 35 |
+
# tall_xcds
|
| 36 |
+
tall_xcds = GRID_MN % NUM_XCDS
|
| 37 |
+
tall_xcds = NUM_XCDS if tall_xcds == 0 else tall_xcds
|
| 38 |
+
# Compute current XCD and local pid within the XCD
|
| 39 |
+
xcd = pid % NUM_XCDS
|
| 40 |
+
local_pid = pid // NUM_XCDS
|
| 41 |
+
# Calculate new pid based on the new grouping
|
| 42 |
+
# Note that we need to consider the following two cases:
|
| 43 |
+
# 1. the current pid is on a tall xcd
|
| 44 |
+
# 2. the current pid is on a short xcd
|
| 45 |
+
if xcd < tall_xcds:
|
| 46 |
+
pid = xcd * pids_per_xcd + local_pid
|
| 47 |
+
else:
|
| 48 |
+
pid = (
|
| 49 |
+
tall_xcds * pids_per_xcd
|
| 50 |
+
+ (xcd - tall_xcds) * (pids_per_xcd - 1)
|
| 51 |
+
+ local_pid
|
| 52 |
+
)
|
| 53 |
+
|
| 54 |
+
return pid
|
| 55 |
+
|
| 56 |
+
|
| 57 |
+
@triton.jit
|
| 58 |
+
def pid_grid(pid: int, num_pid_m: int, num_pid_n: int, GROUP_SIZE_M: tl.constexpr = 1):
|
| 59 |
+
"""
|
| 60 |
+
Maps 1D pid to 2D grid coords (pid_m, pid_n).
|
| 61 |
+
|
| 62 |
+
Args:
|
| 63 |
+
- pid: 1D pid
|
| 64 |
+
- num_pid_m: grid m size
|
| 65 |
+
- num_pid_n: grid n size
|
| 66 |
+
- GROUP_SIZE_M: tl.constexpr: default is 1
|
| 67 |
+
"""
|
| 68 |
+
if GROUP_SIZE_M == 1:
|
| 69 |
+
pid_m = pid // num_pid_n
|
| 70 |
+
pid_n = pid % num_pid_n
|
| 71 |
+
else:
|
| 72 |
+
num_pid_in_group = GROUP_SIZE_M * num_pid_n
|
| 73 |
+
group_id = pid // num_pid_in_group
|
| 74 |
+
first_pid_m = group_id * GROUP_SIZE_M
|
| 75 |
+
group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
|
| 76 |
+
tl.assume(group_size_m >= 0)
|
| 77 |
+
pid_m = first_pid_m + (pid % group_size_m)
|
| 78 |
+
pid_n = (pid % num_pid_in_group) // group_size_m
|
| 79 |
+
|
| 80 |
+
return pid_m, pid_n
|
| 81 |
+
|
| 82 |
+
|
| 83 |
+
@triton.jit
|
| 84 |
+
def pid_grid_3d(pid: int, num_pid_m: int, num_pid_n: int, num_pid_k):
|
| 85 |
+
"""
|
| 86 |
+
Maps 1D pid to 3D grid coords (pid_m, pid_n, pid_k).
|
| 87 |
+
Args:
|
| 88 |
+
- pid: 1D pid
|
| 89 |
+
- num_pid_m: grid m size
|
| 90 |
+
- num_pid_n: grid n size
|
| 91 |
+
- num_pid_k: grid k size
|
| 92 |
+
|
| 93 |
+
Returns:
|
| 94 |
+
- pid_m, pid_n, pid_k: 3D grid coordinates
|
| 95 |
+
"""
|
| 96 |
+
pid_m = pid % num_pid_m
|
| 97 |
+
pid_n = (pid // num_pid_m) % num_pid_n
|
| 98 |
+
pid_k = pid // (num_pid_m * num_pid_n) % num_pid_k
|
| 99 |
+
|
| 100 |
+
return pid_m, pid_n, pid_k
|
build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/utils/gmm_common.py
ADDED
|
@@ -0,0 +1,752 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: MIT
|
| 2 |
+
# Copyright (C) 2025, Advanced Micro Devices, Inc. All rights reserved.
|
| 3 |
+
|
| 4 |
+
# Imports.
|
| 5 |
+
# ------------------------------------------------------------------------------
|
| 6 |
+
|
| 7 |
+
# PyTorch
|
| 8 |
+
import torch
|
| 9 |
+
from torch import Tensor
|
| 10 |
+
|
| 11 |
+
# AITER: logging
|
| 12 |
+
from .logger import AiterTritonLogger
|
| 13 |
+
|
| 14 |
+
_LOGGER: AiterTritonLogger = AiterTritonLogger()
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
# Supported data types.
|
| 18 |
+
# ------------------------------------------------------------------------------
|
| 19 |
+
|
| 20 |
+
# Supported data types, as strings.
|
| 21 |
+
SUPPORTED_DTYPES_STR: set[str] = {"fp16", "bf16"}
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
# Convert string data type to PyTorch data type.
|
| 25 |
+
def dtype_from_str(dtype_str: str) -> torch.dtype:
|
| 26 |
+
dtype_str = dtype_str.strip().lower()
|
| 27 |
+
dtype_str = dtype_str[1:] if dtype_str[0] in {"i", "o"} else dtype_str
|
| 28 |
+
assert (
|
| 29 |
+
dtype_str in SUPPORTED_DTYPES_STR
|
| 30 |
+
), "String data type isn't in set of supported string data types."
|
| 31 |
+
return {"fp16": torch.float16, "bf16": torch.bfloat16}[dtype_str]
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
# Supported data types, as PyTorch types.
|
| 35 |
+
SUPPORTED_DTYPES: set[torch.dtype] = {
|
| 36 |
+
dtype_from_str(dtype_str) for dtype_str in SUPPORTED_DTYPES_STR
|
| 37 |
+
}
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
# Convert PyTorch data type to string data type.
|
| 41 |
+
def str_from_dtype(dtype: torch.dtype) -> str:
|
| 42 |
+
assert (
|
| 43 |
+
dtype in SUPPORTED_DTYPES
|
| 44 |
+
), "PyTorch data type isn't in set of supported PyTorch data types."
|
| 45 |
+
return {torch.float16: "fp16", torch.bfloat16: "bf16"}[dtype]
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
# Default data type, as string.
|
| 49 |
+
DTYPE_STR: str = "bf16"
|
| 50 |
+
assert (
|
| 51 |
+
DTYPE_STR in SUPPORTED_DTYPES_STR
|
| 52 |
+
), "Default string data type isn't in set of supported string data types."
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
# Default data type, as PyTorch type.
|
| 56 |
+
DTYPE: torch.dtype = dtype_from_str(DTYPE_STR)
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
# Other defaults.
|
| 60 |
+
# ------------------------------------------------------------------------------
|
| 61 |
+
|
| 62 |
+
# Default device.
|
| 63 |
+
DEVICE: torch.device | str = "cuda"
|
| 64 |
+
|
| 65 |
+
# Default RNG seed for input generation.
|
| 66 |
+
RNG_SEED: int = 0
|
| 67 |
+
|
| 68 |
+
# Default number of group sizes.
|
| 69 |
+
NUM_GROUP_SIZES: int = 1
|
| 70 |
+
|
| 71 |
+
# Default transposition (NN).
|
| 72 |
+
TRANS_LHS: bool = False
|
| 73 |
+
TRANS_RHS: bool = False
|
| 74 |
+
|
| 75 |
+
|
| 76 |
+
# Parameter checking functions.
|
| 77 |
+
# ------------------------------------------------------------------------------
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
def is_power_of_2(x: int) -> bool:
|
| 81 |
+
return (x > 0) and (x & (x - 1) == 0)
|
| 82 |
+
|
| 83 |
+
|
| 84 |
+
def check_input_device_dtype(
|
| 85 |
+
lhs: Tensor, rhs: Tensor, group_sizes: Tensor, bias: Tensor | None = None
|
| 86 |
+
) -> None:
|
| 87 |
+
assert (
|
| 88 |
+
lhs.device == rhs.device == group_sizes.device
|
| 89 |
+
), f"All input tensors must be in the same device (lhs = {lhs.device}, rhs = {rhs.device}, group_sizes = {group_sizes.device})."
|
| 90 |
+
assert (
|
| 91 |
+
lhs.dtype == rhs.dtype
|
| 92 |
+
), f"lhs and rhs types must match (lhs = {lhs.dtype}, rhs = {rhs.dtype})."
|
| 93 |
+
assert group_sizes.dtype == torch.int32, "group_sizes type must be int32."
|
| 94 |
+
|
| 95 |
+
if bias is not None:
|
| 96 |
+
assert (
|
| 97 |
+
bias.device == lhs.device
|
| 98 |
+
), f"bias must be on the same device as lhs (bias = {bias.device}, lhs = {lhs.device})."
|
| 99 |
+
assert (
|
| 100 |
+
bias.dtype == lhs.dtype
|
| 101 |
+
), f"bias dtype must match lhs dtype (bias = {bias.dtype}, lhs = {lhs.dtype})."
|
| 102 |
+
|
| 103 |
+
|
| 104 |
+
def check_bias_shape_stride(bias: Tensor, G: int, N: int) -> None:
|
| 105 |
+
assert bias.shape == (
|
| 106 |
+
G,
|
| 107 |
+
N,
|
| 108 |
+
), f"bias must have shape (G, N) = ({G}, {N}), got {bias.shape}."
|
| 109 |
+
assert bias.stride() == (N, 1), "bias must be row-major (bias.stride() == (N, 1))."
|
| 110 |
+
|
| 111 |
+
|
| 112 |
+
# Generation of group sizes.
|
| 113 |
+
# ------------------------------------------------------------------------------
|
| 114 |
+
|
| 115 |
+
|
| 116 |
+
# Probabilities for generating random group sizes.
|
| 117 |
+
UNUSED_TOKENS_PROB: float = 0.0
|
| 118 |
+
UNUSED_EXPERTS_PROB: float = 0.1
|
| 119 |
+
|
| 120 |
+
|
| 121 |
+
def gen_uniform_group_sizes(
|
| 122 |
+
M: int,
|
| 123 |
+
G: int,
|
| 124 |
+
device: torch.device | str = DEVICE,
|
| 125 |
+
) -> Tensor:
|
| 126 |
+
assert M >= 0, f"Number of tokens M must be non-negative (it's {M})."
|
| 127 |
+
assert G > 0, f"Number of experts G must be positive (it's {G})."
|
| 128 |
+
|
| 129 |
+
base = M // G
|
| 130 |
+
remainder = M % G
|
| 131 |
+
group_sizes = torch.full((G,), base, dtype=torch.int32, device=device)
|
| 132 |
+
if remainder > 0:
|
| 133 |
+
group_sizes[:remainder] += 1
|
| 134 |
+
|
| 135 |
+
assert (
|
| 136 |
+
len(group_sizes) == G
|
| 137 |
+
), f"Group sizes don't have {G} elements (it's {len(group_sizes)})."
|
| 138 |
+
assert torch.all(group_sizes >= 0).item(), "All group sizes must be non-negative."
|
| 139 |
+
assert (
|
| 140 |
+
torch.sum(group_sizes).item() == M
|
| 141 |
+
), f"Group sizes don't add up to total tokens {M}."
|
| 142 |
+
assert group_sizes.dtype == torch.int32, "Group sizes must be int32."
|
| 143 |
+
|
| 144 |
+
return group_sizes
|
| 145 |
+
|
| 146 |
+
|
| 147 |
+
def gen_group_sizes(
|
| 148 |
+
M: int,
|
| 149 |
+
G: int,
|
| 150 |
+
device: torch.device | str = DEVICE,
|
| 151 |
+
rng_seed: int | None = RNG_SEED,
|
| 152 |
+
unused_tokens_prob: float = UNUSED_TOKENS_PROB,
|
| 153 |
+
unused_experts_prob: float = UNUSED_EXPERTS_PROB,
|
| 154 |
+
) -> Tensor:
|
| 155 |
+
assert M >= 0, f"Number of tokens M must be non-negative (it's {M})."
|
| 156 |
+
assert G > 0, f"Number of experts G must be positive (it's {G})."
|
| 157 |
+
assert (
|
| 158 |
+
0 <= unused_tokens_prob <= 1
|
| 159 |
+
), f"Probability of unused tokens must be in [0, 1] interval (it's {unused_tokens_prob})."
|
| 160 |
+
assert (
|
| 161 |
+
0 <= unused_experts_prob <= 1
|
| 162 |
+
), f"Probability of unused experts must be in [0, 1] interval (it's {unused_experts_prob})."
|
| 163 |
+
|
| 164 |
+
if rng_seed is not None:
|
| 165 |
+
torch.manual_seed(rng_seed)
|
| 166 |
+
|
| 167 |
+
if unused_tokens_prob > 0:
|
| 168 |
+
# Optionally drop tokens to simulate routing sparsity, some tokens may not be routed.
|
| 169 |
+
num_unused_tokens = M
|
| 170 |
+
while num_unused_tokens == M:
|
| 171 |
+
num_unused_tokens = int(
|
| 172 |
+
torch.binomial(
|
| 173 |
+
torch.tensor(float(M), device=device),
|
| 174 |
+
torch.tensor(unused_tokens_prob, device=device),
|
| 175 |
+
).item()
|
| 176 |
+
)
|
| 177 |
+
else:
|
| 178 |
+
num_unused_tokens = 0
|
| 179 |
+
num_used_tokens = M - num_unused_tokens
|
| 180 |
+
assert (
|
| 181 |
+
num_unused_tokens >= 0
|
| 182 |
+
), f"Number of unused tokens must be non-negative (it's {num_unused_tokens})."
|
| 183 |
+
assert (
|
| 184 |
+
num_used_tokens > 0
|
| 185 |
+
), f"Number of used tokens must be positive (it's {num_used_tokens})."
|
| 186 |
+
assert (
|
| 187 |
+
num_used_tokens + num_unused_tokens == M
|
| 188 |
+
), f"Unused + used tokens don't add up total tokens ({num_used_tokens} + {num_unused_tokens} != {M})."
|
| 189 |
+
|
| 190 |
+
if num_unused_tokens > 0:
|
| 191 |
+
_LOGGER.debug(
|
| 192 |
+
f"Group sizes generation: dropped {num_unused_tokens} token{'s' if num_unused_tokens > 1 else ''}.",
|
| 193 |
+
)
|
| 194 |
+
|
| 195 |
+
if unused_experts_prob > 0:
|
| 196 |
+
# Some experts may have zero tokens assigned to them.
|
| 197 |
+
num_used_experts = 0
|
| 198 |
+
while num_used_experts == 0:
|
| 199 |
+
used_experts = torch.nonzero(
|
| 200 |
+
torch.rand((G,), device=device) >= unused_experts_prob
|
| 201 |
+
).squeeze()
|
| 202 |
+
num_used_experts = used_experts.numel()
|
| 203 |
+
else:
|
| 204 |
+
used_experts = torch.arange(0, G, device=device)
|
| 205 |
+
num_used_experts = G
|
| 206 |
+
num_unused_experts = G - num_used_experts
|
| 207 |
+
assert (
|
| 208 |
+
num_unused_experts >= 0
|
| 209 |
+
), f"Number of unused experts must be non-negative (it's {num_unused_experts})."
|
| 210 |
+
assert (
|
| 211 |
+
num_used_experts >= 1
|
| 212 |
+
), f"At least one expert must be used (it's {num_used_experts})."
|
| 213 |
+
assert (
|
| 214 |
+
num_unused_experts + num_used_experts == G
|
| 215 |
+
), f"Unused + used experts don't add up total experts ({num_unused_experts} + {num_used_experts} != {G})."
|
| 216 |
+
|
| 217 |
+
if num_unused_experts > 0:
|
| 218 |
+
_LOGGER.debug(
|
| 219 |
+
f"Group sizes generation: dropped {num_unused_experts} expert{'s' if num_unused_experts > 1 else ''}.",
|
| 220 |
+
)
|
| 221 |
+
|
| 222 |
+
group_sizes = torch.bincount(
|
| 223 |
+
used_experts[
|
| 224 |
+
torch.randint(low=0, high=num_used_experts, size=(num_used_tokens,))
|
| 225 |
+
],
|
| 226 |
+
minlength=G,
|
| 227 |
+
).to(torch.int32)
|
| 228 |
+
|
| 229 |
+
assert (
|
| 230 |
+
len(group_sizes) == G
|
| 231 |
+
), f"Group sizes don't have {G} elements (it's {len(group_sizes)})."
|
| 232 |
+
assert torch.all(group_sizes >= 0).item(), "All group sizes must be non-negative."
|
| 233 |
+
assert (
|
| 234 |
+
torch.sum(group_sizes).item() == num_used_tokens
|
| 235 |
+
), f"Group sizes don't add up to used tokens {num_used_tokens}."
|
| 236 |
+
assert group_sizes.dtype == torch.int32, "Group sizes must be int32."
|
| 237 |
+
|
| 238 |
+
return group_sizes
|
| 239 |
+
|
| 240 |
+
|
| 241 |
+
def gen_multiple_group_sizes(
|
| 242 |
+
num_group_sizes: int,
|
| 243 |
+
M: int,
|
| 244 |
+
G: int,
|
| 245 |
+
device: torch.device | str = DEVICE,
|
| 246 |
+
rng_seed: int | None = RNG_SEED,
|
| 247 |
+
unused_tokens_prob: float = UNUSED_TOKENS_PROB,
|
| 248 |
+
unused_experts_prob: float = UNUSED_EXPERTS_PROB,
|
| 249 |
+
group_sizes_0: Tensor | None = None,
|
| 250 |
+
) -> list[Tensor]:
|
| 251 |
+
assert (
|
| 252 |
+
num_group_sizes > 0
|
| 253 |
+
), f"Number of group sizes to be generated must be positive, it's {num_group_sizes}."
|
| 254 |
+
multiple_group_sizes = [
|
| 255 |
+
gen_group_sizes(
|
| 256 |
+
M,
|
| 257 |
+
G,
|
| 258 |
+
device=device,
|
| 259 |
+
rng_seed=rng_seed if g == 0 else None,
|
| 260 |
+
unused_tokens_prob=unused_tokens_prob,
|
| 261 |
+
unused_experts_prob=unused_experts_prob,
|
| 262 |
+
)
|
| 263 |
+
for g in range(
|
| 264 |
+
num_group_sizes if group_sizes_0 is None else num_group_sizes - 1
|
| 265 |
+
)
|
| 266 |
+
]
|
| 267 |
+
if group_sizes_0 is not None:
|
| 268 |
+
multiple_group_sizes.insert(0, group_sizes_0)
|
| 269 |
+
assert (
|
| 270 |
+
len(multiple_group_sizes) == num_group_sizes
|
| 271 |
+
), f"Expecting {num_group_sizes} distinct group sizes (it's {len(multiple_group_sizes)})."
|
| 272 |
+
return multiple_group_sizes
|
| 273 |
+
|
| 274 |
+
|
| 275 |
+
# GMM helpers: tensor generation.
|
| 276 |
+
# ------------------------------------------------------------------------------
|
| 277 |
+
|
| 278 |
+
|
| 279 |
+
def gen_gmm_input(
|
| 280 |
+
M: int,
|
| 281 |
+
K: int,
|
| 282 |
+
N: int,
|
| 283 |
+
G: int,
|
| 284 |
+
device: torch.device | str = DEVICE,
|
| 285 |
+
preferred_element_type: torch.dtype = DTYPE,
|
| 286 |
+
trans_rhs: bool = TRANS_RHS,
|
| 287 |
+
rng_seed: int | None = RNG_SEED,
|
| 288 |
+
unif_group_sizes: bool = False,
|
| 289 |
+
) -> tuple[Tensor, Tensor, Tensor]:
|
| 290 |
+
assert M > 0, f"Number of lhs rows M must be positive (M = {M})."
|
| 291 |
+
assert K > 0, f"Number of lhs columns / rhs rows K must be positive (K = {K})."
|
| 292 |
+
assert N > 0, f"Number of rhs columns N must be positive (N = {N})."
|
| 293 |
+
assert G > 0, f"Number of groups G must be positive (G = {G})."
|
| 294 |
+
|
| 295 |
+
if rng_seed is not None:
|
| 296 |
+
torch.manual_seed(rng_seed)
|
| 297 |
+
|
| 298 |
+
lhs = torch.randn((M, K), dtype=torch.float32, device=device)
|
| 299 |
+
lhs = lhs.to(preferred_element_type)
|
| 300 |
+
|
| 301 |
+
if trans_rhs:
|
| 302 |
+
rhs = torch.randn((G, N, K), dtype=torch.float32, device=device).permute(
|
| 303 |
+
0, 2, 1
|
| 304 |
+
)
|
| 305 |
+
else:
|
| 306 |
+
rhs = torch.randn((G, K, N), dtype=torch.float32, device=device)
|
| 307 |
+
rhs = rhs.to(preferred_element_type)
|
| 308 |
+
|
| 309 |
+
group_sizes = (
|
| 310 |
+
gen_uniform_group_sizes(M, G, device=device)
|
| 311 |
+
if unif_group_sizes
|
| 312 |
+
else gen_group_sizes(M, G, device=device, rng_seed=None)
|
| 313 |
+
)
|
| 314 |
+
|
| 315 |
+
return lhs, rhs, group_sizes
|
| 316 |
+
|
| 317 |
+
|
| 318 |
+
def gen_gmm_output(
|
| 319 |
+
M: int,
|
| 320 |
+
N: int,
|
| 321 |
+
device: torch.device | str = DEVICE,
|
| 322 |
+
preferred_element_type: torch.dtype = DTYPE,
|
| 323 |
+
) -> Tensor:
|
| 324 |
+
assert M > 0, f"Number of out rows M must be positive (M = {M})."
|
| 325 |
+
assert N > 0, f"Number of out columns N must be positive (N = {N})."
|
| 326 |
+
|
| 327 |
+
out = torch.empty((M, N), dtype=preferred_element_type, device=device)
|
| 328 |
+
|
| 329 |
+
return out
|
| 330 |
+
|
| 331 |
+
|
| 332 |
+
def gen_gmm_tensors(
|
| 333 |
+
M: int,
|
| 334 |
+
K: int,
|
| 335 |
+
N: int,
|
| 336 |
+
G: int,
|
| 337 |
+
num_group_sizes: int,
|
| 338 |
+
device: torch.device | str = DEVICE,
|
| 339 |
+
input_type: torch.dtype = DTYPE,
|
| 340 |
+
output_type: torch.dtype = DTYPE,
|
| 341 |
+
trans_lhs: bool = False,
|
| 342 |
+
trans_rhs: bool = TRANS_RHS,
|
| 343 |
+
rng_seed: int | None = RNG_SEED,
|
| 344 |
+
unif_group_sizes: bool = False,
|
| 345 |
+
use_bias: bool = False,
|
| 346 |
+
) -> tuple[Tensor, Tensor, list[Tensor], Tensor, Tensor | None]:
|
| 347 |
+
lhs, rhs, group_sizes_0 = gen_gmm_input(
|
| 348 |
+
M,
|
| 349 |
+
K,
|
| 350 |
+
N,
|
| 351 |
+
G,
|
| 352 |
+
device=device,
|
| 353 |
+
preferred_element_type=input_type,
|
| 354 |
+
trans_rhs=trans_rhs,
|
| 355 |
+
rng_seed=rng_seed,
|
| 356 |
+
unif_group_sizes=unif_group_sizes,
|
| 357 |
+
)
|
| 358 |
+
multiple_group_sizes = gen_multiple_group_sizes(
|
| 359 |
+
num_group_sizes, M, G, device=device, rng_seed=None, group_sizes_0=group_sizes_0
|
| 360 |
+
)
|
| 361 |
+
out = gen_gmm_output(M, N, device=device, preferred_element_type=output_type)
|
| 362 |
+
bias = None
|
| 363 |
+
if use_bias:
|
| 364 |
+
torch.manual_seed(rng_seed + 1000) # Different seed for bias
|
| 365 |
+
bias = torch.randn(G, N, dtype=input_type, device=device)
|
| 366 |
+
|
| 367 |
+
return lhs, rhs, multiple_group_sizes, out, bias
|
| 368 |
+
|
| 369 |
+
|
| 370 |
+
# GMM helpers: get information from tensors.
|
| 371 |
+
# ------------------------------------------------------------------------------
|
| 372 |
+
|
| 373 |
+
|
| 374 |
+
def get_gmm_shape(
|
| 375 |
+
lhs: Tensor, rhs: Tensor, group_sizes: Tensor
|
| 376 |
+
) -> tuple[int, int, int, int]:
|
| 377 |
+
assert lhs.dim() == 2, f"lhs must have 2 dimensions (it's {lhs.dim()})."
|
| 378 |
+
assert rhs.dim() == 3, f"rhs must have 3 dimensions (it's {rhs.dim()})."
|
| 379 |
+
assert (
|
| 380 |
+
group_sizes.dim() == 1
|
| 381 |
+
), f"group_sizes must have 1 dimension (it's {group_sizes.dim()})."
|
| 382 |
+
|
| 383 |
+
M, lhs_k = lhs.shape
|
| 384 |
+
rhs_g, rhs_k, N = rhs.shape
|
| 385 |
+
group_sizes_g = group_sizes.shape[0]
|
| 386 |
+
|
| 387 |
+
assert (
|
| 388 |
+
lhs_k == rhs_k
|
| 389 |
+
), f"K dimension of lhs and rhs don't match (lhs = {lhs_k}, rhs = {rhs_k})."
|
| 390 |
+
K = lhs_k
|
| 391 |
+
assert (
|
| 392 |
+
rhs_g == group_sizes_g
|
| 393 |
+
), f"G dimension of rhs and group_sizes don't match (rhs = {rhs_g}, group_sizes = {group_sizes_g})."
|
| 394 |
+
G = rhs_g
|
| 395 |
+
|
| 396 |
+
assert M > 0, f"M must be positive, it's {M}."
|
| 397 |
+
assert K > 0, f"K must be positive, it's {K}."
|
| 398 |
+
assert N > 0, f"N must be positive, it's {N}"
|
| 399 |
+
assert G > 0, f"G must be positive, it's {G}"
|
| 400 |
+
|
| 401 |
+
return M, K, N, G
|
| 402 |
+
|
| 403 |
+
|
| 404 |
+
def get_gmm_output(
|
| 405 |
+
M: int,
|
| 406 |
+
N: int,
|
| 407 |
+
device: torch.device | str = DEVICE,
|
| 408 |
+
preferred_element_type: torch.dtype = DTYPE,
|
| 409 |
+
existing_out: Tensor | None = None,
|
| 410 |
+
) -> Tensor:
|
| 411 |
+
assert M > 0, f"Number of out rows M must be positive (M = {M})."
|
| 412 |
+
assert N > 0, f"Number of out columns N must be positive (N = {N})."
|
| 413 |
+
|
| 414 |
+
if existing_out is not None:
|
| 415 |
+
assert (
|
| 416 |
+
existing_out.device == device
|
| 417 |
+
), f"Existing output device and provided device don't match (existing = {existing_out.device}, provided = {device})."
|
| 418 |
+
assert (
|
| 419 |
+
existing_out.dtype == preferred_element_type
|
| 420 |
+
), f"Existing output type and preferred output type don't match (existing = {existing_out.dtype}, preferred = {preferred_element_type})."
|
| 421 |
+
assert existing_out.shape == (
|
| 422 |
+
M,
|
| 423 |
+
N,
|
| 424 |
+
), f"Existing output shape and GMM shape don't match (existing = {tuple(existing_out.shape)}, provided = {(M, N)})."
|
| 425 |
+
return existing_out
|
| 426 |
+
|
| 427 |
+
return gen_gmm_output(
|
| 428 |
+
M,
|
| 429 |
+
N,
|
| 430 |
+
device=device,
|
| 431 |
+
preferred_element_type=preferred_element_type,
|
| 432 |
+
)
|
| 433 |
+
|
| 434 |
+
|
| 435 |
+
def get_gmm_transposition(lhs: Tensor, rhs: Tensor, out: Tensor) -> tuple[bool, int]:
|
| 436 |
+
assert lhs.dim() == 2, f"lhs must have 2 dimensions (it's {lhs.dim()})."
|
| 437 |
+
assert rhs.dim() == 3, f"rhs must have 3 dimensions (it's {rhs.dim()})."
|
| 438 |
+
assert out.dim() == 2, f"out must have 2 dimensions (it's {out.dim()})."
|
| 439 |
+
|
| 440 |
+
lhs_m, lhs_k = lhs.shape
|
| 441 |
+
G, rhs_k, rhs_n = rhs.shape
|
| 442 |
+
out_m, out_n = out.shape
|
| 443 |
+
|
| 444 |
+
assert (
|
| 445 |
+
lhs_m == out_m
|
| 446 |
+
), f"M dimension of lhs and out don't match (lhs = {lhs_m}, rhs = {out_m})."
|
| 447 |
+
M = lhs_m
|
| 448 |
+
assert (
|
| 449 |
+
lhs_k == rhs_k
|
| 450 |
+
), f"K dimension of lhs and rhs don't match (lhs = {lhs_k}, rhs = {rhs_k})."
|
| 451 |
+
K = lhs_k
|
| 452 |
+
assert (
|
| 453 |
+
rhs_n == out_n
|
| 454 |
+
), f"N dimension of rhs and out don't match (lhs = {rhs_n}, rhs = {out_n})."
|
| 455 |
+
N = rhs_n
|
| 456 |
+
|
| 457 |
+
assert M > 0, f"M must be positive, it's {M}."
|
| 458 |
+
assert K > 0, f"K must be positive, it's {K}."
|
| 459 |
+
assert N > 0, f"N must be positive, it's {N}"
|
| 460 |
+
assert G > 0, f"G must be positive, it's {G}"
|
| 461 |
+
|
| 462 |
+
is_lhs_row_major = lhs.stride() == (K, 1)
|
| 463 |
+
assert is_lhs_row_major, "lhs must be row-major."
|
| 464 |
+
is_rhs_row_major = rhs.stride() == (K * N, N, 1)
|
| 465 |
+
is_rhs_col_major = rhs.stride() == (K * N, 1, K)
|
| 466 |
+
assert (
|
| 467 |
+
is_rhs_row_major != is_rhs_col_major
|
| 468 |
+
), "rhs must be row-major or column-major."
|
| 469 |
+
is_out_row_major = out.stride() == (N, 1)
|
| 470 |
+
assert is_out_row_major, "out must be row-major."
|
| 471 |
+
|
| 472 |
+
# Get rhs leading dimension according to transposition configuration.
|
| 473 |
+
ld_rhs = N if is_rhs_row_major else K
|
| 474 |
+
|
| 475 |
+
return is_rhs_col_major, ld_rhs
|
| 476 |
+
|
| 477 |
+
|
| 478 |
+
# TGMM helpers: tensor generation.
|
| 479 |
+
# ------------------------------------------------------------------------------
|
| 480 |
+
|
| 481 |
+
|
| 482 |
+
def gen_tgmm_input(
|
| 483 |
+
M: int,
|
| 484 |
+
K: int,
|
| 485 |
+
N: int,
|
| 486 |
+
G: int,
|
| 487 |
+
device: torch.device | str = DEVICE,
|
| 488 |
+
preferred_element_type: torch.dtype = DTYPE,
|
| 489 |
+
trans_lhs: bool = TRANS_LHS,
|
| 490 |
+
rng_seed: int | None = RNG_SEED,
|
| 491 |
+
unif_group_sizes: bool = False,
|
| 492 |
+
) -> tuple[Tensor, Tensor, Tensor]:
|
| 493 |
+
assert K > 0, f"Number of lhs rows K must be positive (M = {K})."
|
| 494 |
+
assert M > 0, f"Number of lhs columns / rhs rows M must be positive (K = {M})."
|
| 495 |
+
assert N > 0, f"Number of rhs columns N must be positive (N = {N})."
|
| 496 |
+
assert G > 0, f"Number of groups G must be positive (G = {G})."
|
| 497 |
+
|
| 498 |
+
if rng_seed is not None:
|
| 499 |
+
torch.manual_seed(rng_seed)
|
| 500 |
+
|
| 501 |
+
if trans_lhs:
|
| 502 |
+
lhs = torch.randn((M, K), dtype=torch.float32, device=device).T
|
| 503 |
+
else:
|
| 504 |
+
lhs = torch.randn((K, M), dtype=torch.float32, device=device)
|
| 505 |
+
lhs = lhs.to(preferred_element_type)
|
| 506 |
+
|
| 507 |
+
rhs = torch.randn((M, N), dtype=torch.float32, device=device)
|
| 508 |
+
rhs = rhs.to(preferred_element_type)
|
| 509 |
+
|
| 510 |
+
group_sizes = (
|
| 511 |
+
gen_uniform_group_sizes(M, G, device=device)
|
| 512 |
+
if unif_group_sizes
|
| 513 |
+
else gen_group_sizes(M, G, device=device, rng_seed=None)
|
| 514 |
+
)
|
| 515 |
+
|
| 516 |
+
return lhs, rhs, group_sizes
|
| 517 |
+
|
| 518 |
+
|
| 519 |
+
def gen_tgmm_output(
|
| 520 |
+
K: int,
|
| 521 |
+
N: int,
|
| 522 |
+
G: int,
|
| 523 |
+
device: torch.device | str = DEVICE,
|
| 524 |
+
preferred_element_type: torch.dtype = DTYPE,
|
| 525 |
+
) -> Tensor:
|
| 526 |
+
assert K > 0, f"Number of out rows K must be positive (K = {K})."
|
| 527 |
+
assert N > 0, f"Number of out columns N must be positive (N = {N})."
|
| 528 |
+
assert G > 0, f"Number of groups G must be positive (G = {G})."
|
| 529 |
+
|
| 530 |
+
out = torch.empty((G, K, N), dtype=preferred_element_type, device=device)
|
| 531 |
+
|
| 532 |
+
return out
|
| 533 |
+
|
| 534 |
+
|
| 535 |
+
def gen_tgmm_bias_grad(
|
| 536 |
+
K: int,
|
| 537 |
+
G: int,
|
| 538 |
+
device: torch.device | str = DEVICE,
|
| 539 |
+
with_bias_grad: bool = False,
|
| 540 |
+
) -> Tensor:
|
| 541 |
+
if with_bias_grad:
|
| 542 |
+
assert K > 0, f"Number of bias_grad rows K must be positive (K = {K})."
|
| 543 |
+
assert G > 0, f"Number of groups G must be positive (G = {G})."
|
| 544 |
+
return torch.empty((G, K), device=device, dtype=torch.float32)
|
| 545 |
+
else:
|
| 546 |
+
# Return dummy pointer when bias_grad is not needed.
|
| 547 |
+
# Must be float32 because atomic_add does not support bf16/fp16,
|
| 548 |
+
# and Triton validates the pointer dtype even in dead branches.
|
| 549 |
+
return torch.tensor([], device=device, dtype=torch.float32)
|
| 550 |
+
|
| 551 |
+
|
| 552 |
+
def gen_tgmm_tensors(
|
| 553 |
+
M: int,
|
| 554 |
+
K: int,
|
| 555 |
+
N: int,
|
| 556 |
+
G: int,
|
| 557 |
+
num_group_sizes: int,
|
| 558 |
+
device: torch.device | str = DEVICE,
|
| 559 |
+
input_type: torch.dtype = DTYPE,
|
| 560 |
+
output_type: torch.dtype = DTYPE,
|
| 561 |
+
trans_lhs: bool = TRANS_LHS,
|
| 562 |
+
trans_rhs: bool = False,
|
| 563 |
+
rng_seed: int | None = RNG_SEED,
|
| 564 |
+
unif_group_sizes: bool = False,
|
| 565 |
+
use_bias: bool = False,
|
| 566 |
+
) -> tuple[Tensor, Tensor, list[Tensor], Tensor, Tensor | None]:
|
| 567 |
+
lhs, rhs, group_sizes_0 = gen_tgmm_input(
|
| 568 |
+
M,
|
| 569 |
+
K,
|
| 570 |
+
N,
|
| 571 |
+
G,
|
| 572 |
+
device=device,
|
| 573 |
+
preferred_element_type=input_type,
|
| 574 |
+
trans_lhs=trans_lhs,
|
| 575 |
+
rng_seed=rng_seed,
|
| 576 |
+
unif_group_sizes=unif_group_sizes,
|
| 577 |
+
)
|
| 578 |
+
multiple_group_sizes = gen_multiple_group_sizes(
|
| 579 |
+
num_group_sizes, M, G, device=device, rng_seed=None, group_sizes_0=group_sizes_0
|
| 580 |
+
)
|
| 581 |
+
out = gen_tgmm_output(K, N, G, device=device, preferred_element_type=output_type)
|
| 582 |
+
if use_bias:
|
| 583 |
+
bias_grad = gen_tgmm_bias_grad(K, G, device=device, with_bias_grad=True)
|
| 584 |
+
else:
|
| 585 |
+
bias_grad = None
|
| 586 |
+
return lhs, rhs, multiple_group_sizes, out, bias_grad
|
| 587 |
+
|
| 588 |
+
|
| 589 |
+
# TGMM helpers: get information from tensors.
|
| 590 |
+
# ------------------------------------------------------------------------------
|
| 591 |
+
|
| 592 |
+
|
| 593 |
+
def get_tgmm_shape(
|
| 594 |
+
lhs: Tensor, rhs: Tensor, group_sizes: Tensor
|
| 595 |
+
) -> tuple[int, int, int, int]:
|
| 596 |
+
assert lhs.dim() == 2, f"lhs must have 2 dimensions (it's {lhs.dim()})."
|
| 597 |
+
assert rhs.dim() == 2, f"rhs must have 2 dimensions (it's {rhs.dim()})."
|
| 598 |
+
assert (
|
| 599 |
+
group_sizes.dim() == 1
|
| 600 |
+
), f"group_sizes must have 1 dimension (it's {group_sizes.dim()})."
|
| 601 |
+
|
| 602 |
+
K, lhs_m = lhs.shape
|
| 603 |
+
rhs_m, N = rhs.shape
|
| 604 |
+
G = group_sizes.shape[0]
|
| 605 |
+
|
| 606 |
+
assert (
|
| 607 |
+
lhs_m == rhs_m
|
| 608 |
+
), f"M dimension of lhs and rhs don't match (lhs = {lhs_m}, rhs = {rhs_m})."
|
| 609 |
+
M = lhs_m
|
| 610 |
+
|
| 611 |
+
assert M > 0, f"M must be positive, it's {M}."
|
| 612 |
+
assert K > 0, f"K must be positive, it's {K}."
|
| 613 |
+
assert N > 0, f"N must be positive, it's {N}"
|
| 614 |
+
assert G > 0, f"G must be positive, it's {G}"
|
| 615 |
+
|
| 616 |
+
return M, K, N, G
|
| 617 |
+
|
| 618 |
+
|
| 619 |
+
def get_tgmm_output(
|
| 620 |
+
K: int,
|
| 621 |
+
N: int,
|
| 622 |
+
G: int,
|
| 623 |
+
device: torch.device | str = DEVICE,
|
| 624 |
+
preferred_element_type: torch.dtype = DTYPE,
|
| 625 |
+
existing_out: Tensor | None = None,
|
| 626 |
+
) -> Tensor:
|
| 627 |
+
assert K > 0, f"Number of out rows K must be positive (K = {K})."
|
| 628 |
+
assert N > 0, f"Number of out columns N must be positive (N = {N})."
|
| 629 |
+
assert G > 0, f"Number of groups G must be positive (G = {G})."
|
| 630 |
+
|
| 631 |
+
if existing_out is not None:
|
| 632 |
+
assert (
|
| 633 |
+
existing_out.device == device
|
| 634 |
+
), f"Existing output device and provided device don't match (existing = {existing_out.device}, provided = {device})."
|
| 635 |
+
assert (
|
| 636 |
+
existing_out.dtype == preferred_element_type
|
| 637 |
+
), f"Existing output type and preferred output type don't match (existing = {existing_out.dtype}, preferred = {preferred_element_type})."
|
| 638 |
+
assert existing_out.shape == (
|
| 639 |
+
G,
|
| 640 |
+
K,
|
| 641 |
+
N,
|
| 642 |
+
), f"Existing output shape and GMM shape don't match (existing = {tuple(existing_out.shape)}, provided = {(G, K, N)})."
|
| 643 |
+
return existing_out
|
| 644 |
+
|
| 645 |
+
return gen_tgmm_output(
|
| 646 |
+
K,
|
| 647 |
+
N,
|
| 648 |
+
G,
|
| 649 |
+
device=device,
|
| 650 |
+
preferred_element_type=preferred_element_type,
|
| 651 |
+
)
|
| 652 |
+
|
| 653 |
+
|
| 654 |
+
def get_tgmm_bias_grad(
|
| 655 |
+
K: int,
|
| 656 |
+
G: int,
|
| 657 |
+
device: torch.device | str = DEVICE,
|
| 658 |
+
existing_bias_grad: Tensor | None = None,
|
| 659 |
+
) -> Tensor:
|
| 660 |
+
"""
|
| 661 |
+
Get or validate bias gradient tensor for TGMM.
|
| 662 |
+
|
| 663 |
+
If existing_bias_grad is provided, validates its shape, device, dtype, and stride,
|
| 664 |
+
and always zeros it before returning (since the kernel uses atomic_add).
|
| 665 |
+
If existing_bias_grad is None, returns a dummy tensor (for use when COMPUTE_BIAS_GRAD=False).
|
| 666 |
+
Parameters
|
| 667 |
+
----------
|
| 668 |
+
K : int
|
| 669 |
+
Number of rows in the bias gradient tensor.
|
| 670 |
+
G : int
|
| 671 |
+
Number of groups.
|
| 672 |
+
device : torch.device or str
|
| 673 |
+
Device for the tensor.
|
| 674 |
+
existing_bias_grad : torch.Tensor or None
|
| 675 |
+
Existing bias gradient tensor to validate and use.
|
| 676 |
+
Returns
|
| 677 |
+
-------
|
| 678 |
+
torch.Tensor
|
| 679 |
+
Valid bias gradient tensor or dummy tensor.
|
| 680 |
+
"""
|
| 681 |
+
assert K > 0, f"Number of bias_grad rows K must be positive (K = {K})."
|
| 682 |
+
assert G > 0, f"Number of groups G must be positive (G = {G})."
|
| 683 |
+
|
| 684 |
+
if existing_bias_grad is not None:
|
| 685 |
+
# Validate existing bias_grad tensor.
|
| 686 |
+
expected_shape = (G, K)
|
| 687 |
+
assert (
|
| 688 |
+
tuple(existing_bias_grad.shape) == expected_shape
|
| 689 |
+
), f"bias_grad must have shape {expected_shape}, got {tuple(existing_bias_grad.shape)}."
|
| 690 |
+
assert (
|
| 691 |
+
existing_bias_grad.device == device
|
| 692 |
+
), f"bias_grad must be on the same device (bias_grad = {existing_bias_grad.device}, device = {device})."
|
| 693 |
+
assert (
|
| 694 |
+
existing_bias_grad.dtype == torch.float32
|
| 695 |
+
), f"bias_grad must be torch.float32 (kernel uses atomic_add which requires float32), got {existing_bias_grad.dtype}."
|
| 696 |
+
assert existing_bias_grad.stride() == (
|
| 697 |
+
K,
|
| 698 |
+
1,
|
| 699 |
+
), f"bias_grad must be row-major with stride (K, 1) = ({K}, 1), got {existing_bias_grad.stride()}."
|
| 700 |
+
|
| 701 |
+
# Always zero the tensor since bias_grad represents gradients for the current
|
| 702 |
+
# computation and should start fresh. The kernel uses atomic_add which adds to
|
| 703 |
+
# existing values, so we must zero before the kernel runs.
|
| 704 |
+
existing_bias_grad.zero_()
|
| 705 |
+
|
| 706 |
+
return existing_bias_grad
|
| 707 |
+
|
| 708 |
+
else:
|
| 709 |
+
return gen_tgmm_bias_grad(K, G, device=device, with_bias_grad=False)
|
| 710 |
+
|
| 711 |
+
|
| 712 |
+
def get_tgmm_transposition(lhs: Tensor, rhs: Tensor, out: Tensor) -> tuple[bool, int]:
|
| 713 |
+
assert lhs.dim() == 2, f"lhs must have 2 dimensions (it's {lhs.dim()})."
|
| 714 |
+
assert rhs.dim() == 2, f"rhs must have 2 dimensions (it's {rhs.dim()})."
|
| 715 |
+
assert out.dim() == 3, f"out must have 3 dimensions (it's {out.dim()})."
|
| 716 |
+
|
| 717 |
+
lhs_k, lhs_m = lhs.shape
|
| 718 |
+
rhs_m, rhs_n = rhs.shape
|
| 719 |
+
G, out_k, out_n = out.shape
|
| 720 |
+
|
| 721 |
+
assert (
|
| 722 |
+
lhs_m == rhs_m
|
| 723 |
+
), f"M dimension of lhs and rhs don't match (lhs = {lhs_m}, rhs = {rhs_m})."
|
| 724 |
+
M = lhs_m
|
| 725 |
+
assert (
|
| 726 |
+
lhs_k == out_k
|
| 727 |
+
), f"K dimension of lhs and out don't match (lhs = {lhs_k}, rhs = {out_k})."
|
| 728 |
+
K = lhs_k
|
| 729 |
+
assert (
|
| 730 |
+
rhs_n == out_n
|
| 731 |
+
), f"N dimension of rhs and out don't match (lhs = {rhs_n}, rhs = {out_n})."
|
| 732 |
+
N = rhs_n
|
| 733 |
+
|
| 734 |
+
assert M > 0, f"M must be positive, it's {M}."
|
| 735 |
+
assert K > 0, f"K must be positive, it's {K}."
|
| 736 |
+
assert N > 0, f"N must be positive, it's {N}"
|
| 737 |
+
assert G > 0, f"G must be positive, it's {G}"
|
| 738 |
+
|
| 739 |
+
is_lhs_row_major = lhs.stride() == (M, 1)
|
| 740 |
+
is_lhs_col_major = lhs.stride() == (1, K)
|
| 741 |
+
assert (
|
| 742 |
+
is_lhs_row_major != is_lhs_col_major
|
| 743 |
+
), "lhs must be row-major or column-major."
|
| 744 |
+
is_rhs_row_major = rhs.stride() == (N, 1)
|
| 745 |
+
assert is_rhs_row_major, "rhs must be row-major."
|
| 746 |
+
is_out_row_major = out.stride() == (K * N, N, 1)
|
| 747 |
+
assert is_out_row_major, "out must be row-major."
|
| 748 |
+
|
| 749 |
+
# Get lhs leading dimension according to transposition configuration.
|
| 750 |
+
ld_lhs = M if is_lhs_row_major else K
|
| 751 |
+
|
| 752 |
+
return is_lhs_col_major, ld_lhs
|
build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/utils/logger.py
ADDED
|
@@ -0,0 +1,47 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
import logging
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
# AITER Triton Logger which is singleton object around python logging.
|
| 6 |
+
# Note: Python logging is also a singleton object, but we want to read the
|
| 7 |
+
# env var AITER_LOG_LEVEL once at the beginning. Another alternative is to do
|
| 8 |
+
# this in __init__.py. In fact, that's how CK logger is setup. We can look at
|
| 9 |
+
# switching to that at some point
|
| 10 |
+
#
|
| 11 |
+
# AITER_LOG_LEVEL follows python logging levels
|
| 12 |
+
# DEBUG
|
| 13 |
+
# INFO
|
| 14 |
+
# WARNING
|
| 15 |
+
# ERROR
|
| 16 |
+
# CRITICAL
|
| 17 |
+
#
|
| 18 |
+
class AiterTritonLogger(object):
|
| 19 |
+
_instance = None
|
| 20 |
+
|
| 21 |
+
def __new__(cls):
|
| 22 |
+
if cls._instance is None:
|
| 23 |
+
cls._instance = super(AiterTritonLogger, cls).__new__(cls)
|
| 24 |
+
log_level_str = os.getenv("AITER_TRITON_LOG_LEVEL", "WARNING").upper()
|
| 25 |
+
numeric_level = getattr(logging, log_level_str, logging.WARNING)
|
| 26 |
+
cls._instance._logger = logging.getLogger("AITER_TRITON")
|
| 27 |
+
cls._instance._logger.setLevel(numeric_level)
|
| 28 |
+
|
| 29 |
+
return cls._instance
|
| 30 |
+
|
| 31 |
+
def get_logger(self):
|
| 32 |
+
return self._logger
|
| 33 |
+
|
| 34 |
+
def debug(self, msg):
|
| 35 |
+
self._logger.debug(msg)
|
| 36 |
+
|
| 37 |
+
def info(self, msg):
|
| 38 |
+
self._logger.info(msg)
|
| 39 |
+
|
| 40 |
+
def warning(self, msg):
|
| 41 |
+
self._logger.warning(msg)
|
| 42 |
+
|
| 43 |
+
def error(self, msg):
|
| 44 |
+
self._logger.error(msg)
|
| 45 |
+
|
| 46 |
+
def critical(self, msg):
|
| 47 |
+
self._logger.critical(msg)
|
build/torch214-cxx11-xpu20261-x86_64-linux/_layers/__init__.py
ADDED
|
@@ -0,0 +1,10 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2024 Databricks
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
# from megablocks.layers.dmoe import dMoE
|
| 5 |
+
from .moe import MoE
|
| 6 |
+
|
| 7 |
+
__all__ = [
|
| 8 |
+
'MoE',
|
| 9 |
+
# 'dMoE',
|
| 10 |
+
]
|
build/torch214-cxx11-xpu20261-x86_64-linux/_layers/activation_fn.py
ADDED
|
@@ -0,0 +1,33 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2024 Databricks
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
from typing import Any, Callable, Union
|
| 5 |
+
|
| 6 |
+
import torch
|
| 7 |
+
from ..stk import Matrix
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
def act_fn(
|
| 11 |
+
x: Matrix,
|
| 12 |
+
function: Callable,
|
| 13 |
+
return_grad_fn: bool = False,
|
| 14 |
+
**kwargs,
|
| 15 |
+
) -> Union[tuple[Matrix, Any] | Matrix]:
|
| 16 |
+
assert isinstance(x, Matrix)
|
| 17 |
+
with torch.set_grad_enabled(torch.is_grad_enabled() or return_grad_fn):
|
| 18 |
+
if return_grad_fn:
|
| 19 |
+
x.data.requires_grad = True
|
| 20 |
+
out = function(x.data, **kwargs)
|
| 21 |
+
y = Matrix(
|
| 22 |
+
x.size(),
|
| 23 |
+
out,
|
| 24 |
+
x.row_indices,
|
| 25 |
+
x.column_indices,
|
| 26 |
+
x.offsets,
|
| 27 |
+
x.column_indices_t,
|
| 28 |
+
x.offsets_t,
|
| 29 |
+
x.block_offsets_t,
|
| 30 |
+
)
|
| 31 |
+
if return_grad_fn:
|
| 32 |
+
return y, out.backward
|
| 33 |
+
return y
|
build/torch214-cxx11-xpu20261-x86_64-linux/_layers/all_to_all.py
ADDED
|
@@ -0,0 +1,54 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2024 Databricks
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
import torch
|
| 5 |
+
import torch.distributed as dist
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
class AllToAllOp(torch.autograd.Function):
|
| 9 |
+
|
| 10 |
+
@staticmethod
|
| 11 |
+
def forward(ctx, x, output_split_sizes, input_split_sizes, group, async_op):
|
| 12 |
+
out = torch.empty((sum(output_split_sizes),) + x.shape[1:], device=x.device, dtype=x.dtype)
|
| 13 |
+
|
| 14 |
+
ctx.input_shape = x.shape
|
| 15 |
+
ctx.output_split_sizes = output_split_sizes
|
| 16 |
+
ctx.input_split_sizes = input_split_sizes
|
| 17 |
+
ctx.group = group
|
| 18 |
+
handle = dist.all_to_all_single(
|
| 19 |
+
out,
|
| 20 |
+
x,
|
| 21 |
+
output_split_sizes=output_split_sizes,
|
| 22 |
+
input_split_sizes=input_split_sizes,
|
| 23 |
+
group=group,
|
| 24 |
+
async_op=async_op,
|
| 25 |
+
)
|
| 26 |
+
return out, handle
|
| 27 |
+
|
| 28 |
+
@staticmethod
|
| 29 |
+
def backward(ctx, grad, _):
|
| 30 |
+
if ctx.needs_input_grad[0]:
|
| 31 |
+
out = torch.empty(
|
| 32 |
+
ctx.input_shape,
|
| 33 |
+
device=grad.device,
|
| 34 |
+
dtype=grad.dtype,
|
| 35 |
+
)
|
| 36 |
+
dist.all_to_all_single(
|
| 37 |
+
out,
|
| 38 |
+
grad,
|
| 39 |
+
output_split_sizes=ctx.input_split_sizes,
|
| 40 |
+
input_split_sizes=ctx.output_split_sizes,
|
| 41 |
+
group=ctx.group,
|
| 42 |
+
)
|
| 43 |
+
return out, None, None, None, None
|
| 44 |
+
return None, None, None, None, None
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
def all_to_all(x, output_split_sizes, input_split_sizes, group, async_op=False):
|
| 48 |
+
return AllToAllOp.apply(
|
| 49 |
+
x,
|
| 50 |
+
output_split_sizes,
|
| 51 |
+
input_split_sizes,
|
| 52 |
+
group,
|
| 53 |
+
async_op,
|
| 54 |
+
)
|
build/torch214-cxx11-xpu20261-x86_64-linux/_layers/arguments.py
ADDED
|
@@ -0,0 +1,101 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2024 Databricks
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
import dataclasses
|
| 5 |
+
from functools import partial
|
| 6 |
+
from typing import Any, Callable, Optional, Union
|
| 7 |
+
|
| 8 |
+
import torch
|
| 9 |
+
import torch.distributed as dist
|
| 10 |
+
import torch.nn.functional as F
|
| 11 |
+
|
| 12 |
+
# import megablocks.grouped_gemm_util as grouped_gemm
|
| 13 |
+
from .. import grouped_gemm_util as grouped_gemm
|
| 14 |
+
|
| 15 |
+
# Type annotation for in-place Tensor initialization function.
|
| 16 |
+
InitFn = Union[Callable[[torch.Tensor], None], partial[torch.Tensor]]
|
| 17 |
+
|
| 18 |
+
_ALLOWED_BITWIDTHS = (-1, 4, 8)
|
| 19 |
+
|
| 20 |
+
DEFAULT_ACTIVATION_FN = partial(F.gelu, approximate='tanh')
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
@dataclasses.dataclass
|
| 24 |
+
class Arguments:
|
| 25 |
+
# Model arguments.
|
| 26 |
+
hidden_size: int = 1024
|
| 27 |
+
ffn_hidden_size: int = 4096
|
| 28 |
+
num_layers: int = 1
|
| 29 |
+
bias: bool = True
|
| 30 |
+
return_bias: bool = True
|
| 31 |
+
activation_fn: Optional[Callable] = DEFAULT_ACTIVATION_FN
|
| 32 |
+
|
| 33 |
+
# MoE arguments.
|
| 34 |
+
moe_num_experts: int = 1
|
| 35 |
+
moe_top_k: int = 1
|
| 36 |
+
moe_capacity_factor: int = 1
|
| 37 |
+
moe_normalize_expert_weights: Optional[Union[int, float]] = None
|
| 38 |
+
moe_loss_weight: float = 0.1
|
| 39 |
+
moe_jitter_eps: Optional[float] = None
|
| 40 |
+
moe_lbl_in_fp32: bool = False
|
| 41 |
+
|
| 42 |
+
# Parallelism arguments.
|
| 43 |
+
moe_expert_model_parallelism: bool = False
|
| 44 |
+
expert_parallel_group: Optional[dist.ProcessGroup] = None
|
| 45 |
+
pipeline_model_parallel_size: int = 1
|
| 46 |
+
num_layers_per_virtual_pipeline_stage: Optional[int] = None
|
| 47 |
+
|
| 48 |
+
# Compute arguments.
|
| 49 |
+
memory_optimized_mlp: bool = False
|
| 50 |
+
mlp_type: str = 'mlp'
|
| 51 |
+
mlp_impl: str = 'sparse'
|
| 52 |
+
|
| 53 |
+
# Initialization arguments.
|
| 54 |
+
fp16: bool = True
|
| 55 |
+
bf16: bool = False
|
| 56 |
+
device: Union[int, torch.device] = dataclasses.field(default_factory=torch.cuda.current_device)
|
| 57 |
+
init_method: InitFn = partial(torch.nn.init.normal_, mean=0.0, std=0.02)
|
| 58 |
+
output_layer_init_method: InitFn = init_method
|
| 59 |
+
|
| 60 |
+
# Benchmarking arguments.
|
| 61 |
+
uniform_expert_assignment: bool = False
|
| 62 |
+
|
| 63 |
+
# shared expert arguments
|
| 64 |
+
shared_expert: bool = False # enable using shared expert
|
| 65 |
+
fc_cls: Any = torch.nn.Linear # class of the fully connected layer in shared expert (purpose: to allow using custom FC layer eg te.Linear (for FP8))
|
| 66 |
+
fc_kwargs: dict[str, Any] = dataclasses.field(default_factory=dict,) # kwargs for custom fc layers
|
| 67 |
+
remat_act_fn: bool = True # enable act fn to be rematerialized instead of stored
|
| 68 |
+
shared_expert_hidden_size: Optional[
|
| 69 |
+
int] = None # hidden size of the shared expert IF we want to set it to something different from hidden_size
|
| 70 |
+
shared_expert_weighted_sum: bool = False # enable using weighted sum for shared expert output (wieghted by number of experts used)
|
| 71 |
+
|
| 72 |
+
# Router Z-loss arguments
|
| 73 |
+
moe_zloss_weight: float = 0 # 1e-3 is a reasonable value
|
| 74 |
+
moe_zloss_in_fp32: bool = False
|
| 75 |
+
|
| 76 |
+
def __post_init__(self):
|
| 77 |
+
# Sparse MLP is not supported with triton >=3.2.0
|
| 78 |
+
# TODO: Remove this once sparse is supported with triton >=3.2.0
|
| 79 |
+
if self.__getattribute__('mlp_impl') == 'sparse':
|
| 80 |
+
try:
|
| 81 |
+
import triton
|
| 82 |
+
if triton.__version__ >= '3.2.0':
|
| 83 |
+
raise ValueError(
|
| 84 |
+
'Sparse MLP is not supported with triton >=3.2.0. Please use mlp_impl="grouped" instead.',
|
| 85 |
+
)
|
| 86 |
+
except ImportError:
|
| 87 |
+
raise ImportError('Triton is required for sparse MLP implementation')
|
| 88 |
+
|
| 89 |
+
if self.__getattribute__('mlp_impl') == 'grouped':
|
| 90 |
+
grouped_gemm.assert_grouped_gemm_is_available()
|
| 91 |
+
|
| 92 |
+
if self.shared_expert_hidden_size is None:
|
| 93 |
+
self.shared_expert_hidden_size = self.ffn_hidden_size
|
| 94 |
+
|
| 95 |
+
|
| 96 |
+
def from_megatron(megatron_args: Any):
|
| 97 |
+
args = Arguments()
|
| 98 |
+
for field in dataclasses.fields(args):
|
| 99 |
+
if hasattr(megatron_args, field.name):
|
| 100 |
+
setattr(args, field.name, getattr(megatron_args, field.name))
|
| 101 |
+
return args
|
build/torch214-cxx11-xpu20261-x86_64-linux/_layers/common.py
ADDED
|
@@ -0,0 +1,26 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2024 Databricks
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
import torch
|
| 5 |
+
|
| 6 |
+
from .arguments import Arguments
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
def dtype(args: Arguments):
|
| 10 |
+
if args.fp16:
|
| 11 |
+
return torch.float16
|
| 12 |
+
elif args.bf16:
|
| 13 |
+
return torch.bfloat16
|
| 14 |
+
return None
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
def cast_if_autocast_enabled(tensor):
|
| 18 |
+
if torch.is_autocast_enabled():
|
| 19 |
+
if tensor.device.type == 'cuda':
|
| 20 |
+
dtype = torch.get_autocast_gpu_dtype()
|
| 21 |
+
elif tensor.device.type == 'cpu':
|
| 22 |
+
dtype = torch.get_autocast_cpu_dtype()
|
| 23 |
+
else:
|
| 24 |
+
raise NotImplementedError()
|
| 25 |
+
return tensor.to(dtype=dtype)
|
| 26 |
+
return tensor
|
build/torch214-cxx11-xpu20261-x86_64-linux/_layers/dmlp_registry.py
ADDED
|
@@ -0,0 +1,42 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2024 Databricks
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
from typing import Union
|
| 5 |
+
|
| 6 |
+
from . import glu, mlp
|
| 7 |
+
from .arguments import Arguments
|
| 8 |
+
|
| 9 |
+
MlpType = Union[mlp.SparseMLP, glu.SparseGLU]
|
| 10 |
+
|
| 11 |
+
_REGISTRY = {
|
| 12 |
+
'mlp': {
|
| 13 |
+
'grouped': mlp.GroupedMLP,
|
| 14 |
+
'sparse': mlp.SparseMLP,
|
| 15 |
+
},
|
| 16 |
+
'glu': {
|
| 17 |
+
'grouped': glu.GroupedGLU,
|
| 18 |
+
'sparse': glu.SparseGLU,
|
| 19 |
+
},
|
| 20 |
+
}
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def get(args: Arguments) -> MlpType:
|
| 24 |
+
"""Returns an MLP for use in a dMoE instance.
|
| 25 |
+
|
| 26 |
+
Uses the provided arguments to instantiate the appropriate
|
| 27 |
+
MLP instance. This only contains MLPs for use in dMoEs
|
| 28 |
+
(ie. only for the dropless versions of MoEs).
|
| 29 |
+
|
| 30 |
+
Args:
|
| 31 |
+
args: propagated Arguments dataclass.
|
| 32 |
+
|
| 33 |
+
Returns:
|
| 34 |
+
An instantiated MLP constructed using the input args.
|
| 35 |
+
"""
|
| 36 |
+
if args.mlp_type not in _REGISTRY:
|
| 37 |
+
raise ValueError(f'Unsupported mlp type: {args.mlp_type}')
|
| 38 |
+
|
| 39 |
+
if args.mlp_impl not in _REGISTRY[args.mlp_type]:
|
| 40 |
+
raise ValueError(f'{args.mlp_type} does not support {args.mlp_impl} backend.',)
|
| 41 |
+
|
| 42 |
+
return _REGISTRY[args.mlp_type][args.mlp_impl](args)
|
build/torch214-cxx11-xpu20261-x86_64-linux/_layers/dmoe.py
ADDED
|
@@ -0,0 +1,337 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2024 Databricks
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
import numpy as np
|
| 5 |
+
import torch
|
| 6 |
+
|
| 7 |
+
# try:
|
| 8 |
+
# import stk.ops
|
| 9 |
+
# except ImportError:
|
| 10 |
+
# import warnings
|
| 11 |
+
# warnings.warn(
|
| 12 |
+
# 'Please add `stanford-stk` if megablocks/_layers/dmoe.py is needed.',
|
| 13 |
+
# )
|
| 14 |
+
|
| 15 |
+
# import megablocks.ops as ops
|
| 16 |
+
# # from megablocks.ops import ops
|
| 17 |
+
# from megablocks.layers import common, dmlp_registry, moe, mpu
|
| 18 |
+
# from megablocks.layers.arguments import Arguments
|
| 19 |
+
|
| 20 |
+
from .. import stk
|
| 21 |
+
from .. import ops
|
| 22 |
+
from . import common, dmlp_registry, moe, mpu
|
| 23 |
+
from .arguments import Arguments
|
| 24 |
+
|
| 25 |
+
def promote_scalar(x):
|
| 26 |
+
return x.view(1) if not len(x.size()) else x
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
class ParallelDroplessMLP(moe.ParallelMLP):
|
| 30 |
+
|
| 31 |
+
def __init__(self, args: Arguments):
|
| 32 |
+
super(ParallelDroplessMLP, self).__init__(args)
|
| 33 |
+
self.hidden_size = args.hidden_size
|
| 34 |
+
self.ffn_hidden_size = mpu.features_per_rank(args)
|
| 35 |
+
self.blocking = 128
|
| 36 |
+
self.mlp = dmlp_registry.get(args)
|
| 37 |
+
|
| 38 |
+
# Calculate the number of bits needed to represent the column indices
|
| 39 |
+
# in the intermediate sparse matrix.
|
| 40 |
+
max_column_index = ((self.ffn_hidden_size * self.num_experts) // self.blocking)
|
| 41 |
+
self.transpose_sort_end_bit = max(
|
| 42 |
+
int(np.ceil(np.log2(max_column_index))),
|
| 43 |
+
1,
|
| 44 |
+
)
|
| 45 |
+
|
| 46 |
+
def sparse_transpose(self, size, row_indices, column_indices, offsets):
|
| 47 |
+
block_columns = size[1] // self.blocking
|
| 48 |
+
|
| 49 |
+
# Sort row indices by column indices to get the transposed matrix's
|
| 50 |
+
# column indices.
|
| 51 |
+
#
|
| 52 |
+
# NOTE: Our sort operation uses the same width indices as the input values.
|
| 53 |
+
# To avoid overflow when we have large activation matrices we cast to
|
| 54 |
+
# 32-bit before sorting.
|
| 55 |
+
_, gather_indices = ops.sort(
|
| 56 |
+
column_indices.int(),
|
| 57 |
+
self.transpose_sort_end_bit,
|
| 58 |
+
)
|
| 59 |
+
|
| 60 |
+
# There are a constant number of blocks in every row of the sparse matrix.
|
| 61 |
+
# A blocks offset is:
|
| 62 |
+
#
|
| 63 |
+
# row_index * blocks_per_row + column_index % blocks_per_row
|
| 64 |
+
#
|
| 65 |
+
# Once we have the block offsets ordered for transposition we can divide
|
| 66 |
+
# by blocks_per_row to get the transposed column indices.
|
| 67 |
+
column_indices_t = row_indices.gather(0, gather_indices.long())
|
| 68 |
+
block_offsets_t = gather_indices.int()
|
| 69 |
+
|
| 70 |
+
zero = torch.zeros((1,), dtype=torch.int32, device=row_indices.device)
|
| 71 |
+
nnz_per_column = ops.histogram(column_indices, block_columns)
|
| 72 |
+
nnz_per_column = ops.inclusive_cumsum(nnz_per_column, 0)
|
| 73 |
+
if nnz_per_column.dim() == 0:
|
| 74 |
+
# This addresses an edge case when ffn_hidden_size is equal to self.blocking.
|
| 75 |
+
nnz_per_column = nnz_per_column.unsqueeze(0)
|
| 76 |
+
offsets_t = torch.cat([zero, nnz_per_column])
|
| 77 |
+
return column_indices_t, offsets_t, block_offsets_t
|
| 78 |
+
|
| 79 |
+
def topology(self, x, padded_bins):
|
| 80 |
+
padded_tokens, _ = x.size()
|
| 81 |
+
assert padded_tokens % self.blocking == 0
|
| 82 |
+
if self.ffn_hidden_size % self.blocking != 0:
|
| 83 |
+
raise ValueError(
|
| 84 |
+
f'The ffn_hidden_size {self.ffn_hidden_size} must be divisible by ' +
|
| 85 |
+
f'the block size {self.blocking}. Please update your configuration.',
|
| 86 |
+
)
|
| 87 |
+
|
| 88 |
+
# Offsets for the sparse matrix. All rows have the
|
| 89 |
+
# same number of nonzero blocks dictated by the
|
| 90 |
+
# dimensionality of a single expert.
|
| 91 |
+
block_rows = padded_tokens // self.blocking
|
| 92 |
+
blocks_per_row = self.ffn_hidden_size // self.blocking
|
| 93 |
+
offsets = torch.arange(
|
| 94 |
+
0,
|
| 95 |
+
block_rows * blocks_per_row + 1,
|
| 96 |
+
blocks_per_row,
|
| 97 |
+
dtype=torch.int32,
|
| 98 |
+
device=x.device,
|
| 99 |
+
)
|
| 100 |
+
|
| 101 |
+
# Indices for the sparse matrix. The indices for
|
| 102 |
+
# the intermediate matrix are dynamic depending
|
| 103 |
+
# on the mapping of tokens to experts.
|
| 104 |
+
column_indices = ops.topology(
|
| 105 |
+
padded_bins,
|
| 106 |
+
self.blocking,
|
| 107 |
+
block_rows,
|
| 108 |
+
blocks_per_row,
|
| 109 |
+
)
|
| 110 |
+
|
| 111 |
+
# TODO(tgale): This is unused. Remove the need for this in stk.
|
| 112 |
+
# For now, use meta init to save the device memory.
|
| 113 |
+
data = torch.empty(
|
| 114 |
+
column_indices.numel(),
|
| 115 |
+
self.blocking,
|
| 116 |
+
self.blocking,
|
| 117 |
+
dtype=common.dtype(self.args),
|
| 118 |
+
device='meta',
|
| 119 |
+
)
|
| 120 |
+
shape = (
|
| 121 |
+
padded_tokens,
|
| 122 |
+
self.ffn_hidden_size * mpu.experts_per_rank(self.args),
|
| 123 |
+
)
|
| 124 |
+
row_indices = stk.ops.row_indices(shape, data, offsets, column_indices)
|
| 125 |
+
column_indices_t, offsets_t, block_offsets_t = self.sparse_transpose(
|
| 126 |
+
shape,
|
| 127 |
+
row_indices,
|
| 128 |
+
column_indices,
|
| 129 |
+
offsets,
|
| 130 |
+
)
|
| 131 |
+
return stk.Matrix(
|
| 132 |
+
shape,
|
| 133 |
+
data,
|
| 134 |
+
row_indices,
|
| 135 |
+
column_indices,
|
| 136 |
+
offsets,
|
| 137 |
+
column_indices_t,
|
| 138 |
+
offsets_t,
|
| 139 |
+
block_offsets_t,
|
| 140 |
+
)
|
| 141 |
+
|
| 142 |
+
def indices_and_padded_bins(self, top_experts):
|
| 143 |
+
# Sort the expert ids to produce the scatter/gather
|
| 144 |
+
# indices for the permutation.
|
| 145 |
+
top_experts = top_experts.int()
|
| 146 |
+
bin_ids, indices = ops.sort(top_experts, self.sort_end_bit)
|
| 147 |
+
|
| 148 |
+
# Histogram the expert ids to identify the number of
|
| 149 |
+
# tokens routed to each expert.
|
| 150 |
+
tokens_per_expert = ops.histogram(top_experts, self.num_experts)
|
| 151 |
+
|
| 152 |
+
# Round the token counts up to the block size used in
|
| 153 |
+
# the matrix muliplications. Caculate the starting
|
| 154 |
+
# position of each bin.
|
| 155 |
+
padded_tokens_per_expert = ops.round_up(
|
| 156 |
+
tokens_per_expert,
|
| 157 |
+
self.blocking,
|
| 158 |
+
)
|
| 159 |
+
padded_bins = ops.inclusive_cumsum(padded_tokens_per_expert, 0)
|
| 160 |
+
padded_bins = promote_scalar(padded_bins)
|
| 161 |
+
|
| 162 |
+
# Calculate the bin bounds for the sorted tokens.
|
| 163 |
+
bins = ops.inclusive_cumsum(tokens_per_expert, 0)
|
| 164 |
+
bins = promote_scalar(bins)
|
| 165 |
+
return indices, bin_ids, bins, padded_bins, tokens_per_expert
|
| 166 |
+
|
| 167 |
+
def sparse_forward_once(self, x, expert_weights, top_experts):
|
| 168 |
+
# x: [sl, bs, hs]
|
| 169 |
+
# expert_weights: [sl * bs, top-k]
|
| 170 |
+
# top_experts: [sl * bs, top-k]
|
| 171 |
+
expert_weights = expert_weights.flatten()
|
| 172 |
+
top_experts = top_experts.flatten()
|
| 173 |
+
with torch.no_grad():
|
| 174 |
+
indices, bin_ids, bins, padded_bins, tokens_per_expert = (self.indices_and_padded_bins(top_experts))
|
| 175 |
+
|
| 176 |
+
# Route the tokens for MoE computation.
|
| 177 |
+
x = x.view(-1, x.shape[-1])
|
| 178 |
+
x = ops.padded_gather(
|
| 179 |
+
x,
|
| 180 |
+
indices,
|
| 181 |
+
bin_ids,
|
| 182 |
+
bins,
|
| 183 |
+
padded_bins,
|
| 184 |
+
self.top_k,
|
| 185 |
+
)
|
| 186 |
+
|
| 187 |
+
# Create the sparse matrix topology.
|
| 188 |
+
with torch.no_grad():
|
| 189 |
+
topo = self.topology(x, padded_bins)
|
| 190 |
+
|
| 191 |
+
# Perform the expert computation.
|
| 192 |
+
x = self.mlp(x, topo)
|
| 193 |
+
|
| 194 |
+
# Un-route the data for the MoE output.
|
| 195 |
+
x = ops.padded_scatter(
|
| 196 |
+
x,
|
| 197 |
+
indices,
|
| 198 |
+
bin_ids,
|
| 199 |
+
expert_weights,
|
| 200 |
+
bins,
|
| 201 |
+
padded_bins,
|
| 202 |
+
self.top_k,
|
| 203 |
+
)
|
| 204 |
+
return x, tokens_per_expert
|
| 205 |
+
|
| 206 |
+
# For use in the base-class parallel_forward_once.
|
| 207 |
+
def sparse_permute_and_compute(
|
| 208 |
+
self,
|
| 209 |
+
x,
|
| 210 |
+
tokens_per_expert,
|
| 211 |
+
indices,
|
| 212 |
+
bin_ids,
|
| 213 |
+
expert_weights,
|
| 214 |
+
bins,
|
| 215 |
+
expert_capactiy, # unused
|
| 216 |
+
top_k,
|
| 217 |
+
):
|
| 218 |
+
|
| 219 |
+
# Round the token counts up to the block size used in the matrix
|
| 220 |
+
# multiplication. Calculate the starting position of each bin.
|
| 221 |
+
padded_tokens_per_expert = ops.round_up(
|
| 222 |
+
tokens_per_expert,
|
| 223 |
+
self.blocking,
|
| 224 |
+
)
|
| 225 |
+
padded_bins = ops.inclusive_cumsum(padded_tokens_per_expert, 0)
|
| 226 |
+
padded_bins = promote_scalar(padded_bins)
|
| 227 |
+
|
| 228 |
+
# Route the tokens for MoE computation.
|
| 229 |
+
x = x.view(-1, x.shape[-1])
|
| 230 |
+
x = ops.padded_gather(x, indices, bin_ids, bins, padded_bins, top_k)
|
| 231 |
+
|
| 232 |
+
# Create the sparse matrix topology.
|
| 233 |
+
with torch.no_grad():
|
| 234 |
+
topo = self.topology(x, padded_bins)
|
| 235 |
+
|
| 236 |
+
# Perform the expert computation.
|
| 237 |
+
x = self.mlp(x, topo)
|
| 238 |
+
|
| 239 |
+
# Un-route the data for the MoE output.
|
| 240 |
+
return ops.padded_scatter(
|
| 241 |
+
x,
|
| 242 |
+
indices,
|
| 243 |
+
bin_ids,
|
| 244 |
+
expert_weights,
|
| 245 |
+
bins,
|
| 246 |
+
padded_bins,
|
| 247 |
+
top_k,
|
| 248 |
+
)
|
| 249 |
+
|
| 250 |
+
def grouped_forward_once(self, x, expert_weights, top_experts):
|
| 251 |
+
# x: [sl, bs, hs]
|
| 252 |
+
# expert_weights: [sl * bs, top-k]
|
| 253 |
+
# top_experts: [sl * bs, top-k]
|
| 254 |
+
expert_weights = expert_weights.flatten()
|
| 255 |
+
top_experts = top_experts.flatten()
|
| 256 |
+
with torch.no_grad():
|
| 257 |
+
indices, bin_ids, bins, tokens_per_expert = (self.indices_and_bins(top_experts))
|
| 258 |
+
|
| 259 |
+
out = self.grouped_permute_and_compute(
|
| 260 |
+
x,
|
| 261 |
+
tokens_per_expert,
|
| 262 |
+
indices,
|
| 263 |
+
bin_ids,
|
| 264 |
+
expert_weights,
|
| 265 |
+
bins,
|
| 266 |
+
-1, # unused
|
| 267 |
+
self.args.moe_top_k,
|
| 268 |
+
)
|
| 269 |
+
return out, tokens_per_expert
|
| 270 |
+
|
| 271 |
+
def grouped_permute_and_compute(
|
| 272 |
+
self,
|
| 273 |
+
x,
|
| 274 |
+
tokens_per_expert,
|
| 275 |
+
indices,
|
| 276 |
+
bin_ids,
|
| 277 |
+
expert_weights,
|
| 278 |
+
bins,
|
| 279 |
+
expert_capactiy, # unused
|
| 280 |
+
top_k,
|
| 281 |
+
):
|
| 282 |
+
|
| 283 |
+
# Route the tokens for MoE computation.
|
| 284 |
+
x = x.view(-1, x.shape[-1])
|
| 285 |
+
x = ops.gather(x, indices, bin_ids, bins, top_k)
|
| 286 |
+
|
| 287 |
+
# Perform the expert computation.
|
| 288 |
+
x = self.mlp(x, tokens_per_expert)
|
| 289 |
+
|
| 290 |
+
# Un-route the data for the MoE output.
|
| 291 |
+
return ops.scatter(x, indices, bin_ids, expert_weights, bins, top_k)
|
| 292 |
+
|
| 293 |
+
def forward_once(self, x, expert_weights, top_experts):
|
| 294 |
+
if self.args.mlp_impl == 'sparse':
|
| 295 |
+
return self.sparse_forward_once(x, expert_weights, top_experts)
|
| 296 |
+
else:
|
| 297 |
+
return self.grouped_forward_once(x, expert_weights, top_experts)
|
| 298 |
+
|
| 299 |
+
def permute_and_compute(
|
| 300 |
+
self,
|
| 301 |
+
x,
|
| 302 |
+
tokens_per_expert,
|
| 303 |
+
indices,
|
| 304 |
+
bin_ids,
|
| 305 |
+
expert_weights,
|
| 306 |
+
bins,
|
| 307 |
+
expert_capactiy,
|
| 308 |
+
top_k,
|
| 309 |
+
):
|
| 310 |
+
if self.args.mlp_impl == 'sparse':
|
| 311 |
+
return self.sparse_permute_and_compute(
|
| 312 |
+
x,
|
| 313 |
+
tokens_per_expert,
|
| 314 |
+
indices,
|
| 315 |
+
bin_ids,
|
| 316 |
+
expert_weights,
|
| 317 |
+
bins,
|
| 318 |
+
expert_capactiy,
|
| 319 |
+
top_k,
|
| 320 |
+
)
|
| 321 |
+
else:
|
| 322 |
+
return self.grouped_permute_and_compute(
|
| 323 |
+
x,
|
| 324 |
+
tokens_per_expert,
|
| 325 |
+
indices,
|
| 326 |
+
bin_ids,
|
| 327 |
+
expert_weights,
|
| 328 |
+
bins,
|
| 329 |
+
expert_capactiy,
|
| 330 |
+
top_k,
|
| 331 |
+
)
|
| 332 |
+
|
| 333 |
+
|
| 334 |
+
class dMoE(moe.MoE):
|
| 335 |
+
|
| 336 |
+
def _init_experts_mlp(self, args: Arguments):
|
| 337 |
+
return ParallelDroplessMLP(args)
|
build/torch214-cxx11-xpu20261-x86_64-linux/_layers/gelu.py
ADDED
|
@@ -0,0 +1,52 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2024 Databricks
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
# try:
|
| 5 |
+
# import stk
|
| 6 |
+
# except ImportError:
|
| 7 |
+
# import warnings
|
| 8 |
+
# warnings.warn(
|
| 9 |
+
# 'Please add `stanford-stk` if megablocks/_layers/gelu.py is needed.',
|
| 10 |
+
# )
|
| 11 |
+
|
| 12 |
+
from .. import stk
|
| 13 |
+
|
| 14 |
+
import torch
|
| 15 |
+
import torch.nn.functional as F
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
@torch.jit.script
|
| 19 |
+
def _gelu_backward_inplace(g, x):
|
| 20 |
+
tanh_out = torch.tanh(0.79788456 * x * (1 + 0.044715 * x * x))
|
| 21 |
+
ff = (0.5 * x * ((1 - tanh_out * tanh_out) * (0.79788456 + 0.1070322243 * x * x)) + 0.5 * (1 + tanh_out))
|
| 22 |
+
return g.mul_(ff)
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
def gelu_backward_(grad: stk.Matrix, x: stk.Matrix):
|
| 26 |
+
# NOTE: The two sparse matrices must have the same topology.
|
| 27 |
+
if isinstance(grad, stk.Matrix) and isinstance(x, stk.Matrix):
|
| 28 |
+
return stk.Matrix(
|
| 29 |
+
x.size(),
|
| 30 |
+
_gelu_backward_inplace(grad.data, x.data),
|
| 31 |
+
x.row_indices,
|
| 32 |
+
x.column_indices,
|
| 33 |
+
x.offsets,
|
| 34 |
+
x.column_indices_t,
|
| 35 |
+
x.offsets_t,
|
| 36 |
+
x.block_offsets_t,
|
| 37 |
+
)
|
| 38 |
+
return _gelu_backward_inplace(grad, x)
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
def gelu(x: stk.Matrix):
|
| 42 |
+
assert isinstance(x, stk.Matrix)
|
| 43 |
+
return stk.Matrix(
|
| 44 |
+
x.size(),
|
| 45 |
+
F.gelu(x.data, approximate='tanh'),
|
| 46 |
+
x.row_indices,
|
| 47 |
+
x.column_indices,
|
| 48 |
+
x.offsets,
|
| 49 |
+
x.column_indices_t,
|
| 50 |
+
x.offsets_t,
|
| 51 |
+
x.block_offsets_t,
|
| 52 |
+
)
|
build/torch214-cxx11-xpu20261-x86_64-linux/_layers/glu.py
ADDED
|
@@ -0,0 +1,244 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2024 Databricks
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
# import stk.ops
|
| 5 |
+
# try:
|
| 6 |
+
# import stk.ops
|
| 7 |
+
# except ImportError:
|
| 8 |
+
# import warnings
|
| 9 |
+
# warnings.warn(
|
| 10 |
+
# 'Please add `stanford-stk` if megablocks/_layers/glu.py is needed.',
|
| 11 |
+
# )
|
| 12 |
+
|
| 13 |
+
from .. import stk
|
| 14 |
+
|
| 15 |
+
import torch
|
| 16 |
+
|
| 17 |
+
# from megablocks import grouped_gemm_util as gg
|
| 18 |
+
# from megablocks.layers import common, mpu
|
| 19 |
+
# from megablocks.layers.activation_fn import act_fn
|
| 20 |
+
# from megablocks.layers.arguments import Arguments
|
| 21 |
+
# from megablocks.layers.mlp import (
|
| 22 |
+
# SharedMLP,
|
| 23 |
+
# SparseMLP,
|
| 24 |
+
# create_dmoe_expert_weights,
|
| 25 |
+
# resolve_dtensor,
|
| 26 |
+
# )
|
| 27 |
+
|
| 28 |
+
from .. import grouped_gemm_util as gg
|
| 29 |
+
from . import common, mpu
|
| 30 |
+
from .activation_fn import act_fn
|
| 31 |
+
from .arguments import Arguments
|
| 32 |
+
from .mlp import (
|
| 33 |
+
SharedMLP,
|
| 34 |
+
SparseMLP,
|
| 35 |
+
create_dmoe_expert_weights,
|
| 36 |
+
resolve_dtensor,
|
| 37 |
+
)
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
class SparseGLU(SparseMLP):
|
| 41 |
+
|
| 42 |
+
def __init__(self, args: Arguments):
|
| 43 |
+
super().__init__(args)
|
| 44 |
+
self.v1 = torch.nn.Parameter(
|
| 45 |
+
torch.empty(
|
| 46 |
+
self._num_rows_per_rank,
|
| 47 |
+
args.hidden_size,
|
| 48 |
+
device=args.device,
|
| 49 |
+
dtype=common.dtype(args),
|
| 50 |
+
),
|
| 51 |
+
)
|
| 52 |
+
with torch.no_grad():
|
| 53 |
+
self.v1.copy_(
|
| 54 |
+
create_dmoe_expert_weights(
|
| 55 |
+
args,
|
| 56 |
+
args.moe_num_experts,
|
| 57 |
+
args.ffn_hidden_size,
|
| 58 |
+
args.hidden_size,
|
| 59 |
+
args.init_method,
|
| 60 |
+
),
|
| 61 |
+
)
|
| 62 |
+
|
| 63 |
+
mpu.set_expert_model_parallel_attributes(
|
| 64 |
+
self.v1,
|
| 65 |
+
self._should_set_parallelism_attribute,
|
| 66 |
+
)
|
| 67 |
+
|
| 68 |
+
def forward(self, x, topo):
|
| 69 |
+
if self.args.memory_optimized_mlp:
|
| 70 |
+
raise NotImplementedError(
|
| 71 |
+
'Memory optimized implementation not yet supported with GLU with sparse kernels.',
|
| 72 |
+
)
|
| 73 |
+
|
| 74 |
+
w1, v1, w2 = self.scale_grad(self.w1), self.scale_grad(self.v1,), self.scale_grad(self.w2)
|
| 75 |
+
w1, v1, w2 = resolve_dtensor(w1), resolve_dtensor(v1,), resolve_dtensor(w2)
|
| 76 |
+
|
| 77 |
+
# Compute the GLU.
|
| 78 |
+
x1 = stk.ops.sdd(x, w1.t(), topo)
|
| 79 |
+
x2 = stk.ops.sdd(x, v1.t(), topo)
|
| 80 |
+
|
| 81 |
+
activation_fn_out = act_fn(x1, self.args.activation_fn)
|
| 82 |
+
x1 = stk.ops.mul(activation_fn_out, x2)
|
| 83 |
+
|
| 84 |
+
return stk.ops.dsd(x1, w2)
|
| 85 |
+
|
| 86 |
+
|
| 87 |
+
class MemoryOptimizedGroupedGLU(torch.autograd.Function):
|
| 88 |
+
"""GroupedMLP with manually scheduled memory reuse."""
|
| 89 |
+
|
| 90 |
+
@staticmethod
|
| 91 |
+
@torch.amp.autocast_mode.custom_fwd(device_type='cuda')
|
| 92 |
+
def forward(ctx, x, w1, v1, w2, batch_sizes, activation_fn):
|
| 93 |
+
# Cast inputs using ctx dtype from AMP
|
| 94 |
+
if ctx._fwd_used_autocast:
|
| 95 |
+
x = x.to(ctx._dtype)
|
| 96 |
+
w1 = w1.to(ctx._dtype)
|
| 97 |
+
v1 = v1.to(ctx._dtype)
|
| 98 |
+
w2 = w2.to(ctx._dtype)
|
| 99 |
+
# x: [m, k], w1: [n, k], v1: [n, k], w2: [n, k]
|
| 100 |
+
if (not x.is_contiguous() or not w1.is_contiguous() or not v1.is_contiguous() or not w2.is_contiguous()):
|
| 101 |
+
raise ValueError("Expected contiguous 'x', 'w1', 'v1' and 'w2'.")
|
| 102 |
+
|
| 103 |
+
# Layer 0: x @ w1.t().
|
| 104 |
+
assert gg.backend is not None
|
| 105 |
+
sdd_out = gg.backend.gmm(x, w1, batch_sizes, trans_b=True)
|
| 106 |
+
v1_out = gg.backend.gmm(x, v1, batch_sizes, trans_b=True)
|
| 107 |
+
|
| 108 |
+
# GeLU.
|
| 109 |
+
activation_fn_out = activation_fn(sdd_out) * v1_out
|
| 110 |
+
|
| 111 |
+
# Layer 1: x @ w2.
|
| 112 |
+
dsd_out = gg.backend.gmm(activation_fn_out, w2, batch_sizes)
|
| 113 |
+
|
| 114 |
+
# NOTE: Save the input to the layer and the activation_fn input for
|
| 115 |
+
# gradient computation. We'll re-compute the activation_fn forward
|
| 116 |
+
# pass in the backward pass to avoid materializing another
|
| 117 |
+
# intermediate.
|
| 118 |
+
ctx.x_shape = x.shape
|
| 119 |
+
ctx.sdd_out_shape = sdd_out.shape
|
| 120 |
+
ctx.dtype = x.dtype
|
| 121 |
+
ctx.activation_fn = activation_fn
|
| 122 |
+
ctx.save_for_backward(w1, v1, w2, batch_sizes, x, sdd_out, v1_out)
|
| 123 |
+
return dsd_out
|
| 124 |
+
|
| 125 |
+
@staticmethod
|
| 126 |
+
@torch.amp.autocast_mode.custom_bwd(device_type='cuda')
|
| 127 |
+
def backward(ctx, ddsd_out):
|
| 128 |
+
if (not ctx.needs_input_grad[0] or not ctx.needs_input_grad[1] or not ctx.needs_input_grad[2]):
|
| 129 |
+
raise ValueError('Expected all MLP inputs to need grad.')
|
| 130 |
+
|
| 131 |
+
# Unpack saved tensors
|
| 132 |
+
# dtype = ctx.dtype
|
| 133 |
+
saved_tensors = ctx.saved_tensors
|
| 134 |
+
w1, v1, w2 = saved_tensors[:3]
|
| 135 |
+
batch_sizes = saved_tensors[3]
|
| 136 |
+
x = saved_tensors[4]
|
| 137 |
+
sdd_out, v1_out = saved_tensors[5:7]
|
| 138 |
+
|
| 139 |
+
# Rematerialize activation_fn output.
|
| 140 |
+
activation_fn = ctx.activation_fn
|
| 141 |
+
with torch.set_grad_enabled(True):
|
| 142 |
+
sdd_out.requires_grad = True
|
| 143 |
+
v1_out.requires_grad = True
|
| 144 |
+
activation_fn_out = activation_fn(sdd_out) * v1_out
|
| 145 |
+
activation_grad_fn = activation_fn_out.backward
|
| 146 |
+
|
| 147 |
+
# Compute dw2 with recomputed activation_fn output.
|
| 148 |
+
assert gg.backend is not None
|
| 149 |
+
dw2 = gg.backend.gmm(
|
| 150 |
+
activation_fn_out,
|
| 151 |
+
ddsd_out,
|
| 152 |
+
batch_sizes,
|
| 153 |
+
trans_a=True,
|
| 154 |
+
)
|
| 155 |
+
|
| 156 |
+
# Compute dactivation_fn_out.
|
| 157 |
+
#
|
| 158 |
+
# NOTE: We reuse the activation_fn_out allocation.
|
| 159 |
+
dactivation_fn_out = activation_fn_out
|
| 160 |
+
gg.backend.gmm(
|
| 161 |
+
ddsd_out,
|
| 162 |
+
w2,
|
| 163 |
+
batch_sizes,
|
| 164 |
+
trans_b=True,
|
| 165 |
+
c=dactivation_fn_out,
|
| 166 |
+
)
|
| 167 |
+
|
| 168 |
+
# Compute dsdd_out.
|
| 169 |
+
#
|
| 170 |
+
# NOTE: This reuses the dactivation_fn_out allocation.
|
| 171 |
+
assert activation_grad_fn is not None
|
| 172 |
+
activation_grad_fn(dactivation_fn_out)
|
| 173 |
+
dsdd_out = sdd_out.grad
|
| 174 |
+
dv1_out = v1_out.grad
|
| 175 |
+
|
| 176 |
+
# Compute dw1.
|
| 177 |
+
dw1 = gg.backend.gmm(dsdd_out, x, batch_sizes, trans_a=True)
|
| 178 |
+
|
| 179 |
+
# Compute dv1.
|
| 180 |
+
dv1 = gg.backend.gmm(dv1_out, x, batch_sizes, trans_a=True)
|
| 181 |
+
|
| 182 |
+
# Compute dx.
|
| 183 |
+
#
|
| 184 |
+
# NOTE: This reuses the ddsd_out allocation.
|
| 185 |
+
dx = ddsd_out
|
| 186 |
+
gg.backend.gmm(dsdd_out, w1, batch_sizes, c=dx)
|
| 187 |
+
dx += gg.backend.gmm(dv1_out, v1, batch_sizes)
|
| 188 |
+
return dx, dw1, dv1, dw2, None, None
|
| 189 |
+
|
| 190 |
+
|
| 191 |
+
memory_optimized_grouped_glu = MemoryOptimizedGroupedGLU.apply
|
| 192 |
+
|
| 193 |
+
|
| 194 |
+
class GroupedGLU(SparseGLU):
|
| 195 |
+
|
| 196 |
+
def forward(self, x, tokens_per_expert):
|
| 197 |
+
batch_sizes = tokens_per_expert.cpu().to(torch.long)
|
| 198 |
+
w1, v1, w2 = (
|
| 199 |
+
self.scale_grad(self.w1),
|
| 200 |
+
self.scale_grad(self.v1),
|
| 201 |
+
self.scale_grad(self.w2),
|
| 202 |
+
)
|
| 203 |
+
w1, v1, w2 = resolve_dtensor(w1), resolve_dtensor(v1,), resolve_dtensor(w2)
|
| 204 |
+
|
| 205 |
+
# Re-shape the weights for the grouped GEMMs.
|
| 206 |
+
ne = mpu.experts_per_rank(self.args)
|
| 207 |
+
w1 = w1.view(ne, -1, self.args.hidden_size)
|
| 208 |
+
v1 = v1.view(ne, -1, self.args.hidden_size)
|
| 209 |
+
w2 = w2.view(ne, -1, self.args.hidden_size)
|
| 210 |
+
|
| 211 |
+
if self.args.memory_optimized_mlp:
|
| 212 |
+
return memory_optimized_grouped_glu(
|
| 213 |
+
x,
|
| 214 |
+
w1,
|
| 215 |
+
v1,
|
| 216 |
+
w2,
|
| 217 |
+
batch_sizes,
|
| 218 |
+
self.args.activation_fn,
|
| 219 |
+
)
|
| 220 |
+
|
| 221 |
+
# Compute the MLP.
|
| 222 |
+
assert gg.ops is not None
|
| 223 |
+
x1 = gg.ops.gmm(x, w1, batch_sizes, trans_b=True)
|
| 224 |
+
x2 = gg.ops.gmm(x, v1, batch_sizes, trans_b=True)
|
| 225 |
+
x1 = self.args.activation_fn(x1) * x2
|
| 226 |
+
return gg.ops.gmm(x1, w2, batch_sizes)
|
| 227 |
+
|
| 228 |
+
|
| 229 |
+
class SharedGLU(SharedMLP):
|
| 230 |
+
"""GPU for shared expert.
|
| 231 |
+
|
| 232 |
+
Note: this is a copy -> pasta -> modify of the LLM-Foundry MPTGLU class
|
| 233 |
+
"""
|
| 234 |
+
|
| 235 |
+
def __init__(self, args: Arguments):
|
| 236 |
+
super().__init__(args)
|
| 237 |
+
self.gate_proj = args.fc_cls(
|
| 238 |
+
args.hidden_size,
|
| 239 |
+
self.args.shared_expert_hidden_size,
|
| 240 |
+
**self.fc_kwargs,
|
| 241 |
+
)
|
| 242 |
+
|
| 243 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 244 |
+
return self.down_proj(self.act(self.gate_proj(x)) * self.up_proj(x))
|
build/torch214-cxx11-xpu20261-x86_64-linux/_layers/memory_test.py
ADDED
|
@@ -0,0 +1,103 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2024 Databricks
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
import gc
|
| 5 |
+
|
| 6 |
+
import torch
|
| 7 |
+
import torch.distributed as dist
|
| 8 |
+
|
| 9 |
+
# from megablocks.layers import arguments, dmoe
|
| 10 |
+
from . import arguments, dmoe
|
| 11 |
+
|
| 12 |
+
_TESTS = ((8, 2048, 4096, 4096, 32, 4),)
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
def get_tensors():
|
| 16 |
+
ptrs = set()
|
| 17 |
+
out = []
|
| 18 |
+
for obj in gc.get_objects():
|
| 19 |
+
if torch.is_tensor(obj):
|
| 20 |
+
if not obj.is_contiguous() or obj.data_ptr() in ptrs:
|
| 21 |
+
continue
|
| 22 |
+
out.append(obj)
|
| 23 |
+
ptrs.add(obj.data_ptr())
|
| 24 |
+
return out
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
def test_memory(
|
| 28 |
+
group,
|
| 29 |
+
batch_size,
|
| 30 |
+
sequence_length,
|
| 31 |
+
hidden_size,
|
| 32 |
+
ffn_hidden_size,
|
| 33 |
+
num_experts,
|
| 34 |
+
top_k,
|
| 35 |
+
):
|
| 36 |
+
args = arguments.Arguments(
|
| 37 |
+
hidden_size=hidden_size,
|
| 38 |
+
ffn_hidden_size=ffn_hidden_size,
|
| 39 |
+
moe_num_experts=num_experts,
|
| 40 |
+
moe_top_k=top_k,
|
| 41 |
+
moe_expert_model_parallelism=True,
|
| 42 |
+
expert_parallel_group=group,
|
| 43 |
+
fp16=False,
|
| 44 |
+
bf16=True,
|
| 45 |
+
device=torch.cuda.current_device(),
|
| 46 |
+
)
|
| 47 |
+
layer = dmoe.dMoE(args).cuda()
|
| 48 |
+
|
| 49 |
+
x = torch.randn((batch_size, sequence_length, hidden_size),
|
| 50 |
+
device=torch.cuda.current_device(),
|
| 51 |
+
dtype=torch.bfloat16).requires_grad_(True)
|
| 52 |
+
torch.cuda.empty_cache()
|
| 53 |
+
|
| 54 |
+
# Run forward + backward.
|
| 55 |
+
# with torch.autograd.detect_anomaly():
|
| 56 |
+
out, _ = layer(x)
|
| 57 |
+
out.mean().backward()
|
| 58 |
+
|
| 59 |
+
# Report peak memory.
|
| 60 |
+
mem = torch.cuda.max_memory_allocated()
|
| 61 |
+
print('Max Memory Allocated = {:0.0f}MiB'.format(mem / 1e6))
|
| 62 |
+
print('Max Memory Reserved = {:0.0f}MiB'.format(torch.cuda.max_memory_reserved() / 1e6,),)
|
| 63 |
+
|
| 64 |
+
# Calculate weight and gradient memory usage.
|
| 65 |
+
weight_memory = 2 * (
|
| 66 |
+
layer.router.layer.weight.numel() + layer.experts.mlp.w1.numel() + layer.experts.mlp.w2.numel()
|
| 67 |
+
)
|
| 68 |
+
|
| 69 |
+
def grad_numel(x):
|
| 70 |
+
if x.grad is not None:
|
| 71 |
+
return x.grad.numel()
|
| 72 |
+
return 0
|
| 73 |
+
|
| 74 |
+
grad_memory = 2 * (
|
| 75 |
+
grad_numel(layer.router.layer.weight) + grad_numel(layer.experts.mlp.w1) + grad_numel(layer.experts.mlp.w2)
|
| 76 |
+
)
|
| 77 |
+
weight_memory += grad_memory
|
| 78 |
+
|
| 79 |
+
print('Weight Memory Allocated = {:0.0f}MiB'.format(weight_memory / 1e6))
|
| 80 |
+
print('Activation Memory Allocated = {:0.0f}MiB'.format((mem - weight_memory) / 1e6,),)
|
| 81 |
+
|
| 82 |
+
# Manually calculate GPU memory usage from the garbage
|
| 83 |
+
# collector.
|
| 84 |
+
gc.collect()
|
| 85 |
+
total = 0
|
| 86 |
+
tensors = get_tensors()
|
| 87 |
+
tensors = sorted(tensors, key=lambda x: -x.numel())
|
| 88 |
+
for i, t in enumerate(tensors):
|
| 89 |
+
total += t.numel()
|
| 90 |
+
print(f'{i}: {t.shape}, {t.numel() * 2}')
|
| 91 |
+
del tensors
|
| 92 |
+
|
| 93 |
+
print('Total Bytes Found = {:0.0f}MiB'.format(total * 2 / 1e6))
|
| 94 |
+
|
| 95 |
+
|
| 96 |
+
if __name__ == '__main__':
|
| 97 |
+
assert dist.is_available()
|
| 98 |
+
group = dist.init_process_group(backend='nccl')
|
| 99 |
+
local_rank = dist.get_rank(group)
|
| 100 |
+
torch.cuda.set_device(local_rank)
|
| 101 |
+
|
| 102 |
+
for args in _TESTS:
|
| 103 |
+
test_memory(group, *args)
|
build/torch214-cxx11-xpu20261-x86_64-linux/_layers/mlp.py
ADDED
|
@@ -0,0 +1,587 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2024 Databricks
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
from typing import Any
|
| 5 |
+
|
| 6 |
+
# try:
|
| 7 |
+
# import stk
|
| 8 |
+
# import stk.backend.triton_kernels
|
| 9 |
+
# import stk.ops
|
| 10 |
+
# except ImportError:
|
| 11 |
+
# import warnings
|
| 12 |
+
# warnings.warn(
|
| 13 |
+
# 'Please add `stanford-stk` if megablocks/_layers/mlp.py is needed.',
|
| 14 |
+
# )
|
| 15 |
+
|
| 16 |
+
from .. import stk
|
| 17 |
+
|
| 18 |
+
import torch
|
| 19 |
+
from packaging import version
|
| 20 |
+
|
| 21 |
+
# from megablocks import grouped_gemm_util as gg
|
| 22 |
+
# from megablocks.layers import common, gelu, mpu
|
| 23 |
+
# from megablocks.layers.activation_fn import act_fn
|
| 24 |
+
# from megablocks.layers.arguments import DEFAULT_ACTIVATION_FN, Arguments, InitFn
|
| 25 |
+
|
| 26 |
+
from .. import grouped_gemm_util as gg
|
| 27 |
+
from . import common, gelu, mpu
|
| 28 |
+
from .activation_fn import act_fn
|
| 29 |
+
from .arguments import DEFAULT_ACTIVATION_FN, Arguments, InitFn
|
| 30 |
+
|
| 31 |
+
class ScaleGradient(torch.autograd.Function):
|
| 32 |
+
|
| 33 |
+
@staticmethod
|
| 34 |
+
@torch.amp.autocast_mode.custom_fwd(device_type='cuda')
|
| 35 |
+
def forward(ctx: Any, x: torch.Tensor, scale: float):
|
| 36 |
+
ctx.scale = scale
|
| 37 |
+
return x
|
| 38 |
+
|
| 39 |
+
@staticmethod
|
| 40 |
+
@torch.amp.autocast_mode.custom_bwd(device_type='cuda')
|
| 41 |
+
def backward(ctx: torch.Tensor, grad: torch.Tensor):
|
| 42 |
+
return grad * ctx.scale, None
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
scale_gradient = ScaleGradient.apply
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
def resolve_dtensor(weight: torch.Tensor):
|
| 49 |
+
if version.parse(torch.__version__) >= version.parse('2.0.0'):
|
| 50 |
+
from torch.distributed._tensor import DTensor
|
| 51 |
+
if isinstance(weight, DTensor):
|
| 52 |
+
return weight.to_local()
|
| 53 |
+
return weight
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
def create_moe_expert_weights(
|
| 57 |
+
args: Arguments,
|
| 58 |
+
num_experts: int,
|
| 59 |
+
ffn_hidden_size: int,
|
| 60 |
+
hidden_size: int,
|
| 61 |
+
init_method: InitFn,
|
| 62 |
+
):
|
| 63 |
+
# Create the entire weight matrix such that the sampled weights will
|
| 64 |
+
# not vary between data parallelism and expert model parallelism for
|
| 65 |
+
# the same random seed.
|
| 66 |
+
master_weights = torch.empty(
|
| 67 |
+
num_experts,
|
| 68 |
+
ffn_hidden_size,
|
| 69 |
+
hidden_size,
|
| 70 |
+
device=args.device,
|
| 71 |
+
dtype=common.dtype(args),
|
| 72 |
+
)
|
| 73 |
+
init_method(master_weights)
|
| 74 |
+
|
| 75 |
+
if not args.moe_expert_model_parallelism:
|
| 76 |
+
return master_weights
|
| 77 |
+
|
| 78 |
+
# Calculate the amount of sharding in each dimension.
|
| 79 |
+
expert_sharding_degree = mpu.expert_sharding_degree(args)
|
| 80 |
+
hidden_sharding_degree = mpu.hidden_sharding_degree(args)
|
| 81 |
+
|
| 82 |
+
# Calculate the experts per rank.
|
| 83 |
+
#
|
| 84 |
+
# NOTE: We assign ranks to be expert parallel before going
|
| 85 |
+
# tensor parallel.
|
| 86 |
+
rank = mpu.get_expert_parallel_rank(args)
|
| 87 |
+
expert_rank = rank % expert_sharding_degree
|
| 88 |
+
num_experts_per_rank = num_experts // expert_sharding_degree
|
| 89 |
+
start_expert = expert_rank * num_experts_per_rank
|
| 90 |
+
end_expert = (expert_rank + 1) * num_experts_per_rank
|
| 91 |
+
|
| 92 |
+
# Calculate the rows per rank.
|
| 93 |
+
row_rank = rank // expert_sharding_degree
|
| 94 |
+
num_rows_per_rank = ffn_hidden_size // hidden_sharding_degree
|
| 95 |
+
start_row = row_rank * num_rows_per_rank
|
| 96 |
+
end_row = (row_rank + 1) * num_rows_per_rank
|
| 97 |
+
|
| 98 |
+
# Slice the weight matrix to get the chunk for this rank.
|
| 99 |
+
with torch.no_grad():
|
| 100 |
+
weights = master_weights[start_expert:end_expert, start_row:end_row]
|
| 101 |
+
return weights
|
| 102 |
+
|
| 103 |
+
|
| 104 |
+
class MLP(torch.nn.Module):
|
| 105 |
+
|
| 106 |
+
def __init__(self, args: Arguments):
|
| 107 |
+
super().__init__()
|
| 108 |
+
self.args = args
|
| 109 |
+
# expert_parallel_world_size = mpu.get_expert_parallel_world_size(args)
|
| 110 |
+
experts_per_rank = mpu.experts_per_rank(args)
|
| 111 |
+
|
| 112 |
+
self.w1 = torch.nn.Parameter(
|
| 113 |
+
torch.empty(
|
| 114 |
+
experts_per_rank,
|
| 115 |
+
args.hidden_size,
|
| 116 |
+
mpu.features_per_rank(args),
|
| 117 |
+
device=args.device,
|
| 118 |
+
dtype=common.dtype(args),
|
| 119 |
+
),
|
| 120 |
+
)
|
| 121 |
+
self.w2 = torch.nn.Parameter(
|
| 122 |
+
torch.empty(
|
| 123 |
+
experts_per_rank,
|
| 124 |
+
mpu.features_per_rank(args),
|
| 125 |
+
args.hidden_size,
|
| 126 |
+
device=args.device,
|
| 127 |
+
dtype=common.dtype(args),
|
| 128 |
+
),
|
| 129 |
+
)
|
| 130 |
+
mpu.set_expert_model_parallel_attributes(
|
| 131 |
+
self.w1,
|
| 132 |
+
args.moe_expert_model_parallelism,
|
| 133 |
+
)
|
| 134 |
+
mpu.set_expert_model_parallel_attributes(
|
| 135 |
+
self.w2,
|
| 136 |
+
args.moe_expert_model_parallelism,
|
| 137 |
+
)
|
| 138 |
+
|
| 139 |
+
# Initialize the parameters for the MLP.
|
| 140 |
+
#
|
| 141 |
+
# NOTE: It is important that we create the weight tensors prior
|
| 142 |
+
# to creating the master weights and slicing our the piece for
|
| 143 |
+
# this rank. If the master weights are created first the PyTorch
|
| 144 |
+
# caching allocator appears to use the same memory block for these
|
| 145 |
+
# and the slice which causes large increases in our peak memory
|
| 146 |
+
# usage.
|
| 147 |
+
with torch.no_grad():
|
| 148 |
+
w1 = create_moe_expert_weights(
|
| 149 |
+
args,
|
| 150 |
+
args.moe_num_experts,
|
| 151 |
+
args.ffn_hidden_size,
|
| 152 |
+
args.hidden_size,
|
| 153 |
+
args.init_method,
|
| 154 |
+
)
|
| 155 |
+
self.w1.copy_(w1.transpose(1, 2).contiguous())
|
| 156 |
+
self.w2.copy_(
|
| 157 |
+
create_moe_expert_weights(
|
| 158 |
+
args,
|
| 159 |
+
args.moe_num_experts,
|
| 160 |
+
args.ffn_hidden_size,
|
| 161 |
+
args.hidden_size,
|
| 162 |
+
args.output_layer_init_method,
|
| 163 |
+
),
|
| 164 |
+
)
|
| 165 |
+
|
| 166 |
+
self.gradient_scale = None
|
| 167 |
+
if self.args.moe_expert_model_parallelism:
|
| 168 |
+
self.gradient_scale = 1 / mpu.get_expert_parallel_world_size(self.args,)
|
| 169 |
+
|
| 170 |
+
def scale_grad(self, w):
|
| 171 |
+
if self.gradient_scale is None:
|
| 172 |
+
return w
|
| 173 |
+
return scale_gradient(w, self.gradient_scale)
|
| 174 |
+
|
| 175 |
+
def forward(self, x):
|
| 176 |
+
w1, w2 = self.scale_grad(self.w1), self.scale_grad(self.w2)
|
| 177 |
+
w1, w2 = resolve_dtensor(w1), resolve_dtensor(w2)
|
| 178 |
+
x = torch.bmm(x, w1)
|
| 179 |
+
x = self.args.activation_fn(x)
|
| 180 |
+
return torch.bmm(x, w2)
|
| 181 |
+
|
| 182 |
+
|
| 183 |
+
def create_dmoe_expert_weights(
|
| 184 |
+
args: Arguments,
|
| 185 |
+
num_experts: int,
|
| 186 |
+
rows: int,
|
| 187 |
+
columns: int,
|
| 188 |
+
init_method: InitFn,
|
| 189 |
+
):
|
| 190 |
+
weights = create_moe_expert_weights(
|
| 191 |
+
args,
|
| 192 |
+
num_experts,
|
| 193 |
+
rows,
|
| 194 |
+
columns,
|
| 195 |
+
init_method,
|
| 196 |
+
)
|
| 197 |
+
return weights.view([-1, columns])
|
| 198 |
+
|
| 199 |
+
|
| 200 |
+
class MemoryOptimizedMLP(torch.autograd.Function):
|
| 201 |
+
"""Sparse MLP with manually scheduled memory reuse."""
|
| 202 |
+
|
| 203 |
+
@staticmethod
|
| 204 |
+
@torch.amp.autocast_mode.custom_fwd(device_type='cuda')
|
| 205 |
+
def forward(ctx, x, w1, w2, topo, activation_fn):
|
| 206 |
+
# Cast inputs using ctx dtype from AMP
|
| 207 |
+
if ctx._fwd_used_autocast:
|
| 208 |
+
x = x.to(ctx._dtype)
|
| 209 |
+
w1 = w1.to(ctx._dtype)
|
| 210 |
+
w2 = w2.to(ctx._dtype)
|
| 211 |
+
# x: [m, k], w1: [n, k], w2: [n, k]
|
| 212 |
+
if (not x.is_contiguous() or not w1.is_contiguous() or not w2.is_contiguous()):
|
| 213 |
+
raise ValueError("Expected contiguous 'x', 'w1' and 'w2'.")
|
| 214 |
+
|
| 215 |
+
topo_tensors = (
|
| 216 |
+
topo.row_indices,
|
| 217 |
+
topo.column_indices,
|
| 218 |
+
topo.offsets,
|
| 219 |
+
topo.column_indices_t,
|
| 220 |
+
topo.offsets_t,
|
| 221 |
+
topo.block_offsets_t,
|
| 222 |
+
)
|
| 223 |
+
|
| 224 |
+
# Layer 0: x @ w1.t().
|
| 225 |
+
sdd_out = stk.ops.sdd(x, w1.t(), topo)
|
| 226 |
+
|
| 227 |
+
# GeLU.
|
| 228 |
+
activation_fn_out = act_fn(sdd_out, activation_fn)
|
| 229 |
+
|
| 230 |
+
# Layer 1: x @ w2.
|
| 231 |
+
dsd_out = stk.ops.dsd(activation_fn_out, w2)
|
| 232 |
+
|
| 233 |
+
# NOTE: Save the input to the layer and the activation_fn input for
|
| 234 |
+
# gradient computation. We'll re-compute the activation_fn forward
|
| 235 |
+
# pass in the backward pass to avoid materializing another
|
| 236 |
+
# intermediate.
|
| 237 |
+
ctx.shape = topo.shape
|
| 238 |
+
ctx.x_shape = x.shape
|
| 239 |
+
ctx.sdd_out_shape = sdd_out.data.shape
|
| 240 |
+
ctx.dtype = x.dtype
|
| 241 |
+
ctx.activation_fn = activation_fn
|
| 242 |
+
ctx.save_for_backward(w1, w2, *topo_tensors, x, sdd_out.data)
|
| 243 |
+
return dsd_out
|
| 244 |
+
|
| 245 |
+
@staticmethod
|
| 246 |
+
@torch.amp.autocast_mode.custom_bwd(device_type='cuda')
|
| 247 |
+
def backward(ctx, ddsd_out):
|
| 248 |
+
if (not ctx.needs_input_grad[0] or not ctx.needs_input_grad[1] or not ctx.needs_input_grad[2]):
|
| 249 |
+
raise ValueError('Expected all MLP inputs to need grad.')
|
| 250 |
+
|
| 251 |
+
# unpack saved tensors
|
| 252 |
+
# dtype = ctx.dtype
|
| 253 |
+
saved_tensors = ctx.saved_tensors
|
| 254 |
+
w1, w2 = saved_tensors[:2]
|
| 255 |
+
topo_tensors = saved_tensors[2:8]
|
| 256 |
+
x = saved_tensors[8]
|
| 257 |
+
sdd_out_data = saved_tensors[9]
|
| 258 |
+
|
| 259 |
+
# rematerialize activation function output
|
| 260 |
+
activation_fn = ctx.activation_fn
|
| 261 |
+
sdd_out = stk.Matrix(ctx.shape, sdd_out_data, *topo_tensors)
|
| 262 |
+
activation_fn_out, activation_grad_fn = act_fn(
|
| 263 |
+
sdd_out,
|
| 264 |
+
activation_fn,
|
| 265 |
+
return_grad_fn=True,
|
| 266 |
+
)
|
| 267 |
+
|
| 268 |
+
# Compute dw2 with recomputed activation_fn output.
|
| 269 |
+
dw2 = stk.ops.dsd(activation_fn_out.t(), ddsd_out)
|
| 270 |
+
|
| 271 |
+
# Compute dactivation_fn_out.
|
| 272 |
+
#
|
| 273 |
+
# NOTE: We reuse the activation_fn_out allocation.
|
| 274 |
+
dactivation_fn_out = activation_fn_out
|
| 275 |
+
stk.backend.triton_kernels.sdd(
|
| 276 |
+
ddsd_out,
|
| 277 |
+
w2.t(),
|
| 278 |
+
dactivation_fn_out.shape,
|
| 279 |
+
dactivation_fn_out.data,
|
| 280 |
+
dactivation_fn_out.offsets,
|
| 281 |
+
dactivation_fn_out.row_indices,
|
| 282 |
+
dactivation_fn_out.column_indices,
|
| 283 |
+
)
|
| 284 |
+
|
| 285 |
+
# Compute dsdd_out.
|
| 286 |
+
#
|
| 287 |
+
# NOTE: This reuses the dactivation_fn_out allocation.
|
| 288 |
+
if activation_fn is DEFAULT_ACTIVATION_FN:
|
| 289 |
+
dsdd_out = gelu.gelu_backward_(dactivation_fn_out, sdd_out)
|
| 290 |
+
else:
|
| 291 |
+
assert activation_grad_fn is not None
|
| 292 |
+
activation_grad_fn(dactivation_fn_out.data)
|
| 293 |
+
dsdd_out = stk.Matrix(ctx.shape, sdd_out.data.grad, *topo_tensors)
|
| 294 |
+
|
| 295 |
+
# Compute dw1.
|
| 296 |
+
dw1 = stk.ops.dsd(dsdd_out.t(), x)
|
| 297 |
+
|
| 298 |
+
# Compute dx.
|
| 299 |
+
#
|
| 300 |
+
# NOTE: This reuses the ddsd_out allocation.
|
| 301 |
+
stk.backend.triton_kernels.dsd(
|
| 302 |
+
dsdd_out.shape,
|
| 303 |
+
dsdd_out.data,
|
| 304 |
+
dsdd_out.offsets,
|
| 305 |
+
dsdd_out.row_indices,
|
| 306 |
+
dsdd_out.column_indices,
|
| 307 |
+
dsdd_out.offsets_t,
|
| 308 |
+
dsdd_out.column_indices_t,
|
| 309 |
+
dsdd_out.block_offsets_t,
|
| 310 |
+
False,
|
| 311 |
+
w1,
|
| 312 |
+
ddsd_out,
|
| 313 |
+
)
|
| 314 |
+
dx = ddsd_out
|
| 315 |
+
return dx, dw1, dw2, None, None
|
| 316 |
+
|
| 317 |
+
|
| 318 |
+
memory_optimized_mlp = MemoryOptimizedMLP.apply
|
| 319 |
+
|
| 320 |
+
|
| 321 |
+
class SparseMLP(torch.nn.Module):
|
| 322 |
+
|
| 323 |
+
def __init__(self, args: Arguments):
|
| 324 |
+
super().__init__()
|
| 325 |
+
self.args = args
|
| 326 |
+
self._num_rows_per_rank = mpu.experts_per_rank(args) * mpu.features_per_rank(args)
|
| 327 |
+
|
| 328 |
+
self.w1 = torch.nn.Parameter(
|
| 329 |
+
torch.empty(
|
| 330 |
+
self._num_rows_per_rank,
|
| 331 |
+
args.hidden_size,
|
| 332 |
+
device=args.device,
|
| 333 |
+
dtype=common.dtype(args),
|
| 334 |
+
),
|
| 335 |
+
)
|
| 336 |
+
self.w2 = torch.nn.Parameter(
|
| 337 |
+
torch.empty(
|
| 338 |
+
self._num_rows_per_rank,
|
| 339 |
+
args.hidden_size,
|
| 340 |
+
device=args.device,
|
| 341 |
+
dtype=common.dtype(args),
|
| 342 |
+
),
|
| 343 |
+
)
|
| 344 |
+
|
| 345 |
+
# Initialize the parameters for the MLP.
|
| 346 |
+
#
|
| 347 |
+
# NOTE: It is important that we create the weight tensors prior
|
| 348 |
+
# to creating the master weights and slicing our the piece for
|
| 349 |
+
# this rank. If the master weights are created first the PyTorch
|
| 350 |
+
# caching allocator appears to use the same memory block for these
|
| 351 |
+
# and the slice which causes large increases in our peak memory
|
| 352 |
+
# usage.
|
| 353 |
+
with torch.no_grad():
|
| 354 |
+
self.w1.copy_(
|
| 355 |
+
create_dmoe_expert_weights(
|
| 356 |
+
args,
|
| 357 |
+
args.moe_num_experts,
|
| 358 |
+
args.ffn_hidden_size,
|
| 359 |
+
args.hidden_size,
|
| 360 |
+
args.init_method,
|
| 361 |
+
),
|
| 362 |
+
)
|
| 363 |
+
self.w2.copy_(
|
| 364 |
+
create_dmoe_expert_weights(
|
| 365 |
+
args,
|
| 366 |
+
args.moe_num_experts,
|
| 367 |
+
args.ffn_hidden_size,
|
| 368 |
+
args.hidden_size,
|
| 369 |
+
args.output_layer_init_method,
|
| 370 |
+
),
|
| 371 |
+
)
|
| 372 |
+
|
| 373 |
+
self._should_set_parallelism_attribute = args.moe_expert_model_parallelism
|
| 374 |
+
mpu.set_expert_model_parallel_attributes(
|
| 375 |
+
self.w1,
|
| 376 |
+
self._should_set_parallelism_attribute,
|
| 377 |
+
)
|
| 378 |
+
mpu.set_expert_model_parallel_attributes(
|
| 379 |
+
self.w2,
|
| 380 |
+
self._should_set_parallelism_attribute,
|
| 381 |
+
)
|
| 382 |
+
|
| 383 |
+
self.gradient_scale = None
|
| 384 |
+
if self.args.moe_expert_model_parallelism:
|
| 385 |
+
self.gradient_scale = 1 / mpu.get_expert_parallel_world_size(self.args,)
|
| 386 |
+
|
| 387 |
+
def scale_grad(self, w):
|
| 388 |
+
if self.gradient_scale is None:
|
| 389 |
+
return w
|
| 390 |
+
return scale_gradient(w, self.gradient_scale)
|
| 391 |
+
|
| 392 |
+
def forward(self, x, topo):
|
| 393 |
+
w1, w2 = self.scale_grad(self.w1), self.scale_grad(self.w2)
|
| 394 |
+
w1, w2 = resolve_dtensor(w1), resolve_dtensor(w2)
|
| 395 |
+
if self.args.memory_optimized_mlp:
|
| 396 |
+
return memory_optimized_mlp(
|
| 397 |
+
x,
|
| 398 |
+
w1,
|
| 399 |
+
w2,
|
| 400 |
+
topo,
|
| 401 |
+
self.args.activation_fn,
|
| 402 |
+
)
|
| 403 |
+
|
| 404 |
+
# Compute the MLP.
|
| 405 |
+
x = stk.ops.sdd(x, w1.t(), topo)
|
| 406 |
+
activation_fn_out = act_fn(x, self.args.activation_fn)
|
| 407 |
+
return stk.ops.dsd(activation_fn_out, w2)
|
| 408 |
+
|
| 409 |
+
|
| 410 |
+
class MemoryOptimizedGroupedMLP(torch.autograd.Function):
|
| 411 |
+
"""GroupedMLP with manually scheduled memory reuse."""
|
| 412 |
+
|
| 413 |
+
@staticmethod
|
| 414 |
+
@torch.amp.autocast_mode.custom_fwd(device_type='cuda')
|
| 415 |
+
def forward(ctx, x, w1, w2, batch_sizes, activation_fn):
|
| 416 |
+
# Cast inputs using ctx dtype from AMP
|
| 417 |
+
if ctx._fwd_used_autocast:
|
| 418 |
+
x = x.to(ctx._dtype)
|
| 419 |
+
w1 = w1.to(ctx._dtype)
|
| 420 |
+
w2 = w2.to(ctx._dtype)
|
| 421 |
+
# x: [m, k], w1: [n, k], w2: [n, k]
|
| 422 |
+
if (not x.is_contiguous() or not w1.is_contiguous() or not w2.is_contiguous()):
|
| 423 |
+
raise ValueError("Expected contiguous 'x', 'w1' and 'w2'.")
|
| 424 |
+
|
| 425 |
+
# Layer 0: x @ w1.t().
|
| 426 |
+
assert gg.backend is not None
|
| 427 |
+
sdd_out = gg.backend.gmm(x, w1, batch_sizes, trans_b=True)
|
| 428 |
+
|
| 429 |
+
# activation_fn
|
| 430 |
+
activation_fn_out = activation_fn(sdd_out)
|
| 431 |
+
|
| 432 |
+
# Layer 1: x @ w2.
|
| 433 |
+
dsd_out = gg.backend.gmm(activation_fn_out, w2, batch_sizes)
|
| 434 |
+
|
| 435 |
+
# NOTE: Save the input to the layer and the activation_fn input for
|
| 436 |
+
# gradient computation. We'll re-compute the activation_fn forward
|
| 437 |
+
# pass in the backward pass to avoid materializing another
|
| 438 |
+
# intermediate.
|
| 439 |
+
ctx.x_shape = x.shape
|
| 440 |
+
ctx.sdd_out_shape = sdd_out.shape
|
| 441 |
+
ctx.dtype = x.dtype
|
| 442 |
+
ctx.activation_fn = activation_fn
|
| 443 |
+
ctx.save_for_backward(w1, w2, batch_sizes, x, sdd_out)
|
| 444 |
+
return dsd_out
|
| 445 |
+
|
| 446 |
+
@staticmethod
|
| 447 |
+
@torch.amp.autocast_mode.custom_bwd(device_type='cuda')
|
| 448 |
+
def backward(ctx: Any, ddsd_out: torch.Tensor):
|
| 449 |
+
if (not ctx.needs_input_grad[0] or not ctx.needs_input_grad[1] or not ctx.needs_input_grad[2]):
|
| 450 |
+
raise ValueError('Expected all MLP inputs to need grad.')
|
| 451 |
+
|
| 452 |
+
# Unpack saved tensors
|
| 453 |
+
# dtype = ctx.dtype
|
| 454 |
+
saved_tensors = ctx.saved_tensors
|
| 455 |
+
w1, w2 = saved_tensors[:2]
|
| 456 |
+
batch_sizes = saved_tensors[2]
|
| 457 |
+
x = saved_tensors[3]
|
| 458 |
+
sdd_out = saved_tensors[4]
|
| 459 |
+
|
| 460 |
+
# Rematerialize activation_fn output.
|
| 461 |
+
activation_fn = ctx.activation_fn
|
| 462 |
+
with torch.set_grad_enabled(True):
|
| 463 |
+
sdd_out.requires_grad = True
|
| 464 |
+
activation_fn_out = activation_fn(sdd_out)
|
| 465 |
+
activation_grad_fn = activation_fn_out.backward
|
| 466 |
+
|
| 467 |
+
# Compute dw2 with recomputed activation_fn output.
|
| 468 |
+
assert gg.backend is not None
|
| 469 |
+
dw2 = gg.backend.gmm(
|
| 470 |
+
activation_fn_out,
|
| 471 |
+
ddsd_out,
|
| 472 |
+
batch_sizes,
|
| 473 |
+
trans_a=True,
|
| 474 |
+
)
|
| 475 |
+
|
| 476 |
+
# Compute dactivation_fn_out.
|
| 477 |
+
#
|
| 478 |
+
# NOTE: We reuse the activation_fn_out allocation.
|
| 479 |
+
dactivation_fn_out = activation_fn_out
|
| 480 |
+
gg.backend.gmm(
|
| 481 |
+
ddsd_out,
|
| 482 |
+
w2,
|
| 483 |
+
batch_sizes,
|
| 484 |
+
trans_b=True,
|
| 485 |
+
c=dactivation_fn_out,
|
| 486 |
+
)
|
| 487 |
+
|
| 488 |
+
# Compute dsdd_out.
|
| 489 |
+
#
|
| 490 |
+
# NOTE: This reuses the dactivation_fn_out allocation.
|
| 491 |
+
if activation_fn is DEFAULT_ACTIVATION_FN:
|
| 492 |
+
dsdd_out = gelu.gelu_backward_(dactivation_fn_out, sdd_out)
|
| 493 |
+
else:
|
| 494 |
+
assert activation_grad_fn is not None
|
| 495 |
+
activation_grad_fn(dactivation_fn_out)
|
| 496 |
+
dsdd_out = sdd_out.grad
|
| 497 |
+
|
| 498 |
+
# Compute dw1.
|
| 499 |
+
dw1 = gg.backend.gmm(dsdd_out, x, batch_sizes, trans_a=True)
|
| 500 |
+
|
| 501 |
+
# Compute dx.
|
| 502 |
+
#
|
| 503 |
+
# NOTE: This reuses the ddsd_out allocation.
|
| 504 |
+
gg.backend.gmm(dsdd_out, w1, batch_sizes, c=ddsd_out)
|
| 505 |
+
dx = ddsd_out
|
| 506 |
+
return dx, dw1, dw2, None, None
|
| 507 |
+
|
| 508 |
+
|
| 509 |
+
memory_optimized_grouped_mlp = MemoryOptimizedGroupedMLP.apply
|
| 510 |
+
|
| 511 |
+
|
| 512 |
+
class GroupedMLP(SparseMLP):
|
| 513 |
+
|
| 514 |
+
def forward(self, x, tokens_per_expert):
|
| 515 |
+
batch_sizes = tokens_per_expert.cpu().to(torch.long)
|
| 516 |
+
w1, w2 = (self.scale_grad(self.w1), self.scale_grad(self.w2))
|
| 517 |
+
|
| 518 |
+
# Re-shape the weights for the grouped GEMMs.
|
| 519 |
+
ne = mpu.experts_per_rank(self.args)
|
| 520 |
+
w1 = resolve_dtensor(w1).view(ne, -1, self.args.hidden_size)
|
| 521 |
+
w2 = resolve_dtensor(w2).view(ne, -1, self.args.hidden_size)
|
| 522 |
+
|
| 523 |
+
if self.args.memory_optimized_mlp:
|
| 524 |
+
return memory_optimized_grouped_mlp(
|
| 525 |
+
x,
|
| 526 |
+
w1,
|
| 527 |
+
w2,
|
| 528 |
+
batch_sizes,
|
| 529 |
+
self.args.activation_fn,
|
| 530 |
+
)
|
| 531 |
+
|
| 532 |
+
# Compute the MLP.
|
| 533 |
+
assert gg.ops is not None
|
| 534 |
+
x = gg.ops.gmm(x, w1, batch_sizes, trans_b=True)
|
| 535 |
+
x = self.args.activation_fn(x)
|
| 536 |
+
return gg.ops.gmm(x, w2, batch_sizes)
|
| 537 |
+
|
| 538 |
+
|
| 539 |
+
class SharedMLP(torch.nn.Module):
|
| 540 |
+
"""MLP for shared expert.
|
| 541 |
+
|
| 542 |
+
Note: this is a copy -> pasta -> modify of the LLM-Foundry MPTMLP class
|
| 543 |
+
"""
|
| 544 |
+
|
| 545 |
+
def __init__(self, args: Arguments):
|
| 546 |
+
super().__init__()
|
| 547 |
+
self.args = args
|
| 548 |
+
self.fc_kwargs: dict[str, Any] = {
|
| 549 |
+
'bias': args.bias,
|
| 550 |
+
'device': args.device,
|
| 551 |
+
}
|
| 552 |
+
self.fc_kwargs.update(args.fc_kwargs)
|
| 553 |
+
|
| 554 |
+
self.up_proj = args.fc_cls(
|
| 555 |
+
args.hidden_size,
|
| 556 |
+
args.shared_expert_hidden_size,
|
| 557 |
+
**self.fc_kwargs,
|
| 558 |
+
)
|
| 559 |
+
self.act = args.activation_fn
|
| 560 |
+
self.down_proj = args.fc_cls(
|
| 561 |
+
args.shared_expert_hidden_size,
|
| 562 |
+
args.hidden_size,
|
| 563 |
+
**self.fc_kwargs,
|
| 564 |
+
)
|
| 565 |
+
self.down_proj._is_residual = True # a flag for llm-foundry init
|
| 566 |
+
|
| 567 |
+
def add_experts_sharedexpert(
|
| 568 |
+
self,
|
| 569 |
+
shared_expert_out: torch.Tensor,
|
| 570 |
+
expert_out: torch.Tensor,
|
| 571 |
+
) -> torch.Tensor:
|
| 572 |
+
# Helper function to add expert output to shared expert output
|
| 573 |
+
# with optional weighted sum.
|
| 574 |
+
if self.args.shared_expert_weighted_sum:
|
| 575 |
+
# enable using weighted sum for shared expert output
|
| 576 |
+
# wieghted by number of experts used
|
| 577 |
+
t_experts = self.args.moe_top_k + 1
|
| 578 |
+
sh_mlp_out = shared_expert_out / t_experts
|
| 579 |
+
return sh_mlp_out.add(
|
| 580 |
+
expert_out,
|
| 581 |
+
alpha=(self.args.moe_top_k / t_experts),
|
| 582 |
+
)
|
| 583 |
+
|
| 584 |
+
return shared_expert_out + expert_out
|
| 585 |
+
|
| 586 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 587 |
+
return self.down_proj(self.act(self.up_proj(x)))
|
build/torch214-cxx11-xpu20261-x86_64-linux/_layers/moe.py
ADDED
|
@@ -0,0 +1,507 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2024 Databricks
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
from typing import Optional, Tuple
|
| 4 |
+
|
| 5 |
+
import numpy as np
|
| 6 |
+
import torch
|
| 7 |
+
import torch.distributed as dist
|
| 8 |
+
|
| 9 |
+
# import megablocks.ops as ops
|
| 10 |
+
# from megablocks.layers import common, mlp, mpu, router, sharedexpert_registry
|
| 11 |
+
# from megablocks.layers.all_to_all import all_to_all
|
| 12 |
+
# from megablocks.layers.arguments import Arguments
|
| 13 |
+
|
| 14 |
+
from ..ops import (
|
| 15 |
+
sort,
|
| 16 |
+
histogram,
|
| 17 |
+
inclusive_cumsum,
|
| 18 |
+
exclusive_cumsum,
|
| 19 |
+
binned_gather,
|
| 20 |
+
binned_scatter,
|
| 21 |
+
gather,
|
| 22 |
+
scatter,
|
| 23 |
+
repeat,
|
| 24 |
+
replicate,
|
| 25 |
+
)
|
| 26 |
+
|
| 27 |
+
from . import common, mlp, mpu, router, sharedexpert_registry
|
| 28 |
+
from .arguments import Arguments
|
| 29 |
+
from .all_to_all import all_to_all
|
| 30 |
+
|
| 31 |
+
_LOAD_BALANCING_LOSS = []
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
def save_load_balancing_loss(loss):
|
| 35 |
+
global _LOAD_BALANCING_LOSS
|
| 36 |
+
_LOAD_BALANCING_LOSS.append(loss)
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def get_load_balancing_loss():
|
| 40 |
+
global _LOAD_BALANCING_LOSS
|
| 41 |
+
return _LOAD_BALANCING_LOSS
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
def clear_load_balancing_loss():
|
| 45 |
+
global _LOAD_BALANCING_LOSS
|
| 46 |
+
_LOAD_BALANCING_LOSS.clear()
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
def batched_load_balancing_loss(args: Arguments):
|
| 50 |
+
if args.moe_loss_weight == 0:
|
| 51 |
+
return 0.0
|
| 52 |
+
|
| 53 |
+
# tokens_per_expert[i].shape = (num_experts)
|
| 54 |
+
# expert_scores[i].shape = (tokens, num_experts)
|
| 55 |
+
tokens_per_expert, expert_scores = zip(*get_load_balancing_loss())
|
| 56 |
+
num_layers_per_pipeline_stage = (args.num_layers // args.pipeline_model_parallel_size)
|
| 57 |
+
if args.num_layers_per_virtual_pipeline_stage is not None:
|
| 58 |
+
num_layers_per_pipeline_stage = args.num_layers_per_virtual_pipeline_stage
|
| 59 |
+
|
| 60 |
+
if len(tokens_per_expert) != num_layers_per_pipeline_stage:
|
| 61 |
+
raise ValueError(
|
| 62 |
+
f'Expected {num_layers_per_pipeline_stage} token_per_experts '
|
| 63 |
+
f'but found {len(tokens_per_expert)}.\nnum_layers = '
|
| 64 |
+
f'{args.num_layers}\npipeline_model_parallel_size = '
|
| 65 |
+
f'{args.pipeline_model_parallel_size}\n'
|
| 66 |
+
'num_layers_per_virtual_pipeline_stage'
|
| 67 |
+
f' = {args.num_layers_per_virtual_pipeline_stage}',
|
| 68 |
+
)
|
| 69 |
+
if len(expert_scores) != num_layers_per_pipeline_stage:
|
| 70 |
+
raise ValueError(
|
| 71 |
+
f'Expected {num_layers_per_pipeline_stage} expert_scores '
|
| 72 |
+
f'but found {len(tokens_per_expert)}.\nnum_layers = '
|
| 73 |
+
f'{args.num_layers}\npipeline_model_parallel_size = '
|
| 74 |
+
f'{args.pipeline_model_parallel_size}\n'
|
| 75 |
+
'num_layers_per_virtual_pipeline_stage'
|
| 76 |
+
f' = {args.num_layers_per_virtual_pipeline_stage}',
|
| 77 |
+
)
|
| 78 |
+
|
| 79 |
+
# Verify the shape of the tokens_per_expert and expert_scores tensors.
|
| 80 |
+
assert all((x.ndim == 1 and x.numel() == args.moe_num_experts for x in tokens_per_expert))
|
| 81 |
+
|
| 82 |
+
tokens = expert_scores[0].shape[0]
|
| 83 |
+
assert all(((x.ndim == 2 and x.shape[1] == args.moe_num_experts and x.shape[0] == tokens) for x in expert_scores))
|
| 84 |
+
|
| 85 |
+
# Concatenate the contributions of each layer and convert to
|
| 86 |
+
# the correct types and formats for the dot product.
|
| 87 |
+
expert_scores = torch.cat(expert_scores, dim=1)
|
| 88 |
+
if args.moe_lbl_in_fp32:
|
| 89 |
+
expert_scores = expert_scores.float()
|
| 90 |
+
if tokens != 0:
|
| 91 |
+
expert_scores = expert_scores.mean(dim=0)
|
| 92 |
+
else:
|
| 93 |
+
expert_scores = expert_scores.sum(dim=0)
|
| 94 |
+
tokens_per_expert = torch.cat(tokens_per_expert).to(expert_scores.dtype)
|
| 95 |
+
|
| 96 |
+
expected_values = num_layers_per_pipeline_stage * args.moe_num_experts
|
| 97 |
+
assert tokens_per_expert.numel() == expected_values
|
| 98 |
+
assert expert_scores.numel() == expected_values
|
| 99 |
+
|
| 100 |
+
# Calculate the total scale across all factors.
|
| 101 |
+
#
|
| 102 |
+
# loss_weight * num_experts / (num_layers * tokens * top_k)
|
| 103 |
+
scale_numerator = (args.moe_num_experts * args.moe_loss_weight)
|
| 104 |
+
scale_denominator = (args.num_layers * tokens * args.moe_top_k)
|
| 105 |
+
scale = scale_numerator / scale_denominator
|
| 106 |
+
return scale * torch.dot(tokens_per_expert, expert_scores)
|
| 107 |
+
|
| 108 |
+
|
| 109 |
+
# NOTE: This class defines MoE expert computation, including expert model parallel
|
| 110 |
+
# communication. When using FSDP on top of MegaBlocks this is the module that should
|
| 111 |
+
# be wrapped s.t. the weight all-gathers can be scheduled *before* the expert model
|
| 112 |
+
# parallel all2all.
|
| 113 |
+
class ParallelMLP(torch.nn.Module):
|
| 114 |
+
|
| 115 |
+
def __init__(self, args: Arguments):
|
| 116 |
+
super(ParallelMLP, self).__init__()
|
| 117 |
+
self.args = args
|
| 118 |
+
|
| 119 |
+
# Calculate the number of experts in total and the number of experts
|
| 120 |
+
# owned by this rank.
|
| 121 |
+
# world_size = mpu.get_expert_parallel_world_size(args)
|
| 122 |
+
self.num_experts = args.moe_num_experts
|
| 123 |
+
self.top_k = self.args.moe_top_k
|
| 124 |
+
|
| 125 |
+
# Calculate the number of bits needed to represent the expert indices
|
| 126 |
+
# so that we can pass it to radix sort.
|
| 127 |
+
self.sort_end_bit = max(int(np.ceil(np.log2(self.num_experts))), 1)
|
| 128 |
+
|
| 129 |
+
# Expert MLP.
|
| 130 |
+
self.mlp = mlp.MLP(args)
|
| 131 |
+
|
| 132 |
+
self.bias: Optional[torch.Tensor]
|
| 133 |
+
if self.args.bias:
|
| 134 |
+
# Note that the output bias is not parallelized with expert
|
| 135 |
+
# model parallelism.
|
| 136 |
+
self.bias = torch.nn.Parameter(
|
| 137 |
+
torch.empty(
|
| 138 |
+
args.hidden_size,
|
| 139 |
+
device=args.device,
|
| 140 |
+
dtype=common.dtype(args),
|
| 141 |
+
),
|
| 142 |
+
)
|
| 143 |
+
torch.nn.init.zeros_(self.bias)
|
| 144 |
+
else:
|
| 145 |
+
self.register_parameter('bias', None)
|
| 146 |
+
|
| 147 |
+
# Select the forward function for the operating mode.
|
| 148 |
+
self.forward_fn = (self.parallel_forward_once if args.moe_expert_model_parallelism else self.forward_once)
|
| 149 |
+
|
| 150 |
+
def expert_capacity(self, tokens: int) -> int:
|
| 151 |
+
world_size = mpu.get_expert_parallel_world_size(self.args)
|
| 152 |
+
tokens_per_expert = (self.top_k * tokens * world_size / self.num_experts)
|
| 153 |
+
return int(self.args.moe_capacity_factor * tokens_per_expert)
|
| 154 |
+
|
| 155 |
+
def load_balancing_loss(self, tokens_per_expert: torch.Tensor, expert_scores: torch.Tensor):
|
| 156 |
+
"""Calculate the load balancing loss contribution."""
|
| 157 |
+
assert len(expert_scores.size()) == 2
|
| 158 |
+
tokens, num_experts = expert_scores.size()
|
| 159 |
+
assert num_experts == self.num_experts
|
| 160 |
+
assert len(tokens_per_expert.size()) == 1
|
| 161 |
+
num_experts, = tokens_per_expert.size()
|
| 162 |
+
assert num_experts == self.num_experts
|
| 163 |
+
scale = self.num_experts / (tokens * self.top_k)
|
| 164 |
+
return scale * torch.dot(
|
| 165 |
+
tokens_per_expert.to(expert_scores.dtype),
|
| 166 |
+
expert_scores.mean(dim=0),
|
| 167 |
+
)
|
| 168 |
+
|
| 169 |
+
def indices_and_bins(self,
|
| 170 |
+
top_expert: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
| 171 |
+
# Sort the expert ids to produce the scatter/gather
|
| 172 |
+
# indices for the permutation.
|
| 173 |
+
#
|
| 174 |
+
# TODO(tgale): Is it worth doing this conversion to 32-bit
|
| 175 |
+
# prior? Could we place the `torch.max` operation to return
|
| 176 |
+
# 32-bit expert indices?
|
| 177 |
+
top_expert = top_expert.int()
|
| 178 |
+
# output = ops.sort(top_expert, self.sort_end_bit)
|
| 179 |
+
output = sort(top_expert, self.sort_end_bit)
|
| 180 |
+
assert output is not None
|
| 181 |
+
bin_ids, indices = output
|
| 182 |
+
|
| 183 |
+
# Histogram the expert ids to identify the number of
|
| 184 |
+
# tokens routed to each expert.
|
| 185 |
+
#
|
| 186 |
+
# TODO(tgale): Does the sorted data produce a more favorable
|
| 187 |
+
# data distribution for histogram? Or is the op parallelism
|
| 188 |
+
# worth more?
|
| 189 |
+
# tokens_per_expert = ops.histogram(top_expert, self.num_experts)
|
| 190 |
+
tokens_per_expert = histogram(top_expert, self.num_experts)
|
| 191 |
+
|
| 192 |
+
# Calculate the bin bounds for the sorted tokens.
|
| 193 |
+
# bins = ops.inclusive_cumsum(tokens_per_expert, 0)
|
| 194 |
+
bins = inclusive_cumsum(tokens_per_expert, 0)
|
| 195 |
+
assert bins is not None
|
| 196 |
+
bins = bins.view(1) if not len(bins.size()) else bins
|
| 197 |
+
|
| 198 |
+
assert isinstance(indices, torch.Tensor)
|
| 199 |
+
assert isinstance(bin_ids, torch.Tensor)
|
| 200 |
+
assert isinstance(bins, torch.Tensor)
|
| 201 |
+
assert isinstance(tokens_per_expert, torch.Tensor)
|
| 202 |
+
|
| 203 |
+
return indices, bin_ids, bins, tokens_per_expert
|
| 204 |
+
|
| 205 |
+
def permute_and_compute(
|
| 206 |
+
self,
|
| 207 |
+
x: torch.Tensor,
|
| 208 |
+
tokens_per_expert: int, # unused
|
| 209 |
+
indices: torch.Tensor,
|
| 210 |
+
bin_ids: torch.Tensor, # unused
|
| 211 |
+
expert_weights: torch.Tensor,
|
| 212 |
+
bins: torch.Tensor,
|
| 213 |
+
expert_capacity: int,
|
| 214 |
+
top_k: int,
|
| 215 |
+
):
|
| 216 |
+
# Route the tokens for MoE computation.
|
| 217 |
+
x = x.view(-1, x.shape[-1])
|
| 218 |
+
# output = ops.binned_gather(x, indices, bins, expert_capacity, top_k)
|
| 219 |
+
output = binned_gather(x, indices, bins, expert_capacity, top_k)
|
| 220 |
+
assert output is not None
|
| 221 |
+
x = output
|
| 222 |
+
|
| 223 |
+
# Perform the expert computation. Note that we don't
|
| 224 |
+
# use biases for these linear operations.
|
| 225 |
+
x = self.mlp(x)
|
| 226 |
+
|
| 227 |
+
# Un-route the data for the MoE output.
|
| 228 |
+
# return ops.binned_scatter(x, indices, expert_weights, bins, top_k)
|
| 229 |
+
return binned_scatter(x, indices, expert_weights, bins, top_k)
|
| 230 |
+
|
| 231 |
+
|
| 232 |
+
def forward_once(self, x: torch.Tensor, expert_weights: torch.Tensor, top_experts: torch.Tensor):
|
| 233 |
+
# x: [sl, bs, hs]
|
| 234 |
+
# expert_weights: [sl * bs, top-k]
|
| 235 |
+
# top_experts: [sl * bs, top-k]
|
| 236 |
+
expert_weights = expert_weights.flatten()
|
| 237 |
+
top_experts = top_experts.flatten()
|
| 238 |
+
with torch.no_grad():
|
| 239 |
+
indices, bin_ids, bins, tokens_per_expert = (self.indices_and_bins(top_experts))
|
| 240 |
+
|
| 241 |
+
# If expert_capacity is set to zero, set the number of tokens
|
| 242 |
+
# per expert to the maximum we need to avoid dropping tokens.
|
| 243 |
+
sl, bs, _ = x.size()
|
| 244 |
+
expert_capacity = self.expert_capacity(sl * bs)
|
| 245 |
+
if expert_capacity == 0:
|
| 246 |
+
expert_capacity = torch.max(tokens_per_expert).item()
|
| 247 |
+
|
| 248 |
+
x = self.permute_and_compute(
|
| 249 |
+
x,
|
| 250 |
+
tokens_per_expert,
|
| 251 |
+
indices,
|
| 252 |
+
bin_ids,
|
| 253 |
+
expert_weights,
|
| 254 |
+
bins,
|
| 255 |
+
expert_capacity,
|
| 256 |
+
self.top_k,
|
| 257 |
+
)
|
| 258 |
+
return x, tokens_per_expert
|
| 259 |
+
|
| 260 |
+
def parallel_forward_once(self, x: torch.Tensor, expert_weights: torch.Tensor, top_experts: torch.Tensor):
|
| 261 |
+
# NOTE: This function implements the same computation as forward_once
|
| 262 |
+
# but with expert model parallelism.
|
| 263 |
+
#
|
| 264 |
+
# 1. Permute the tokens locally so that they are grouped by their
|
| 265 |
+
# expert assignments. This allows us to transfer all of the tokens
|
| 266 |
+
# for a remote device in one communication primitive.
|
| 267 |
+
#
|
| 268 |
+
# 2. Permute the tokens across the expert parallel devices. After
|
| 269 |
+
# this is completed each device has all of the tokens assigned to
|
| 270 |
+
# its set of experts in its local HBM.
|
| 271 |
+
#
|
| 272 |
+
# 3. Permute the tokens locally so that they are grouped by their
|
| 273 |
+
# expert assignement. After the distributed permutation the tokens
|
| 274 |
+
# are grouped by which device they came from. We re-order them
|
| 275 |
+
# locally to allow for efficient computation.
|
| 276 |
+
#
|
| 277 |
+
# After this series of permutations we compute the linear layers
|
| 278 |
+
# and then repeat these three steps in reverse to produce the final
|
| 279 |
+
# output.
|
| 280 |
+
#
|
| 281 |
+
# Compute the mapping of local tokens to experts.
|
| 282 |
+
expert_weights = expert_weights.flatten()
|
| 283 |
+
top_experts = top_experts.flatten()
|
| 284 |
+
with torch.no_grad():
|
| 285 |
+
indices, bin_ids, bins, tokens_per_expert = (self.indices_and_bins(top_experts))
|
| 286 |
+
|
| 287 |
+
# If we're sharding the experts along the hidden dimension
|
| 288 |
+
# multiple devices own parts of the same sets of experts.
|
| 289 |
+
# Replicate the token counts so every device gets the counts.
|
| 290 |
+
# repeated_tokens_per_expert = ops.repeat(
|
| 291 |
+
repeated_tokens_per_expert = repeat(
|
| 292 |
+
tokens_per_expert,
|
| 293 |
+
(mpu.hidden_sharding_degree(self.args),),
|
| 294 |
+
)
|
| 295 |
+
|
| 296 |
+
# Pass token count information to the device on which the
|
| 297 |
+
# target expert resides.
|
| 298 |
+
parallel_tokens_per_expert = torch.empty_like(repeated_tokens_per_expert,)
|
| 299 |
+
tpe_handle = dist.all_to_all_single(
|
| 300 |
+
parallel_tokens_per_expert,
|
| 301 |
+
repeated_tokens_per_expert,
|
| 302 |
+
group=self.args.expert_parallel_group,
|
| 303 |
+
async_op=True,
|
| 304 |
+
)
|
| 305 |
+
|
| 306 |
+
# Permute locally and without any padding so that tokens for each
|
| 307 |
+
# parallel device are stored contiguously.
|
| 308 |
+
#
|
| 309 |
+
# This view updates the shape of the tensor from [sl, bs, hs] to
|
| 310 |
+
# [sl * bs, hs] prior to the permutation.
|
| 311 |
+
x = x.view(-1, x.shape[-1])
|
| 312 |
+
# output = ops.gather(x, indices, bin_ids, bins, self.top_k)
|
| 313 |
+
output = gather(x, indices, bin_ids, bins, self.top_k)
|
| 314 |
+
assert output is not None
|
| 315 |
+
x = output
|
| 316 |
+
|
| 317 |
+
# Compute the number of tokens that will be received from each
|
| 318 |
+
# device and permute the input data across the devices.
|
| 319 |
+
with torch.no_grad():
|
| 320 |
+
tpe_handle.wait()
|
| 321 |
+
experts_per_rank = mpu.experts_per_rank(self.args)
|
| 322 |
+
|
| 323 |
+
# Reshape to [world_size, num_experts_per_rank].
|
| 324 |
+
world_size = mpu.get_expert_parallel_world_size(self.args)
|
| 325 |
+
repeated_tokens_per_expert = (repeated_tokens_per_expert.view(world_size, experts_per_rank))
|
| 326 |
+
parallel_tokens_per_expert = (parallel_tokens_per_expert.view(world_size, experts_per_rank))
|
| 327 |
+
|
| 328 |
+
# TODO(tgale): It might be faster to do this on the GPU and
|
| 329 |
+
# then communicate the results back to the host.
|
| 330 |
+
send_counts = repeated_tokens_per_expert.cpu().sum(dim=-1)
|
| 331 |
+
parallel_tokens_per_expert_cpu = parallel_tokens_per_expert.cpu()
|
| 332 |
+
recv_counts = parallel_tokens_per_expert_cpu.sum(dim=-1)
|
| 333 |
+
|
| 334 |
+
# Convert the send/recv counts to lists.
|
| 335 |
+
send_counts = send_counts.tolist()
|
| 336 |
+
recv_counts = recv_counts.tolist()
|
| 337 |
+
tokens_received = sum(recv_counts)
|
| 338 |
+
|
| 339 |
+
# If we're sharding the experts along the hidden dimension
|
| 340 |
+
# multiple devices own parts of the same sets of experts.
|
| 341 |
+
# Replicate the token counts so devices that share experts
|
| 342 |
+
# get all of the tokens assigned to them.
|
| 343 |
+
#
|
| 344 |
+
# TODO(tgale): Fuse this into the prior, local permutation.
|
| 345 |
+
# x = ops.repeat(x, (mpu.hidden_sharding_degree(self.args), 1))
|
| 346 |
+
x = repeat(x, (mpu.hidden_sharding_degree(self.args), 1))
|
| 347 |
+
|
| 348 |
+
# Start the cross-device permutation asynchronously so we can
|
| 349 |
+
# overlap communication with computation.
|
| 350 |
+
parallel_x, parallel_x_handle = all_to_all(
|
| 351 |
+
x,
|
| 352 |
+
recv_counts,
|
| 353 |
+
send_counts,
|
| 354 |
+
self.args.expert_parallel_group,
|
| 355 |
+
async_op=True,
|
| 356 |
+
)
|
| 357 |
+
|
| 358 |
+
with torch.no_grad():
|
| 359 |
+
# After we do the cross-device permutation we have the tokens on the
|
| 360 |
+
# correct device but not yet grouped by expert because we received
|
| 361 |
+
# tokens from each device as contiguous chunks. To group the tokens
|
| 362 |
+
# for expert computation we'll do one more local permutation. The
|
| 363 |
+
# rest of this torch.no_grad() scope sets up the indices and bins
|
| 364 |
+
# for this permutation.
|
| 365 |
+
# replicate_bins = ops.inclusive_cumsum(
|
| 366 |
+
replicate_bins = inclusive_cumsum(
|
| 367 |
+
parallel_tokens_per_expert.flatten(),
|
| 368 |
+
0,
|
| 369 |
+
)
|
| 370 |
+
replicate_bins = (replicate_bins.view(1) if not len(replicate_bins.size()) else replicate_bins)
|
| 371 |
+
|
| 372 |
+
# Construct the expert indices for the permuted tokens.
|
| 373 |
+
parallel_top_expert = torch.remainder(
|
| 374 |
+
torch.arange(
|
| 375 |
+
self.num_experts * mpu.hidden_sharding_degree(self.args),
|
| 376 |
+
dtype=torch.int32,
|
| 377 |
+
device=indices.device,
|
| 378 |
+
),
|
| 379 |
+
mpu.experts_per_rank(self.args),
|
| 380 |
+
)
|
| 381 |
+
# parallel_top_expert = ops.replicate(
|
| 382 |
+
parallel_top_expert = replicate(
|
| 383 |
+
parallel_top_expert.unsqueeze(dim=0),
|
| 384 |
+
replicate_bins,
|
| 385 |
+
tokens_received,
|
| 386 |
+
).flatten()
|
| 387 |
+
|
| 388 |
+
# TODO(tgale): The sort_end_bit here can be reduced.
|
| 389 |
+
# parallel_bin_ids, parallel_indices = ops.sort(
|
| 390 |
+
parallel_bin_ids, parallel_indices = sort(
|
| 391 |
+
parallel_top_expert,
|
| 392 |
+
self.sort_end_bit,
|
| 393 |
+
)
|
| 394 |
+
|
| 395 |
+
# Calculate the bins boundaries from the token counts.
|
| 396 |
+
parallel_tokens_per_expert = parallel_tokens_per_expert.sum(
|
| 397 |
+
dim=0,
|
| 398 |
+
dtype=torch.int,
|
| 399 |
+
)
|
| 400 |
+
# parallel_bins = ops.inclusive_cumsum(parallel_tokens_per_expert, 0)
|
| 401 |
+
parallel_bins = inclusive_cumsum(parallel_tokens_per_expert, 0)
|
| 402 |
+
parallel_bins = (parallel_bins.view(1) if not len(parallel_bins.size()) else parallel_bins)
|
| 403 |
+
|
| 404 |
+
# If expert_capacity is set to zero, set the number of tokens
|
| 405 |
+
# per expert to the maximum we need to avoid dropping tokens.
|
| 406 |
+
tokens, _ = x.size()
|
| 407 |
+
expert_capacity = self.expert_capacity(tokens)
|
| 408 |
+
if expert_capacity == 0:
|
| 409 |
+
expert_capacity = torch.max(parallel_tokens_per_expert).item()
|
| 410 |
+
|
| 411 |
+
# Locally permute the tokens and perform the expert computation.
|
| 412 |
+
# Block to make sure that the cross-device permutation is complete.
|
| 413 |
+
if self.args.mlp_impl == 'grouped':
|
| 414 |
+
# GroupedMLP requires counts on CPU. We can use the tensor already
|
| 415 |
+
# moved to CPU for the prior all_to_all, which avoids an extra
|
| 416 |
+
# device synchronization.
|
| 417 |
+
parallel_tokens_per_expert = parallel_tokens_per_expert_cpu.sum(
|
| 418 |
+
dim=0,
|
| 419 |
+
dtype=torch.int,
|
| 420 |
+
)
|
| 421 |
+
parallel_x_handle.wait()
|
| 422 |
+
parallel_x = self.permute_and_compute(
|
| 423 |
+
parallel_x,
|
| 424 |
+
parallel_tokens_per_expert,
|
| 425 |
+
parallel_indices,
|
| 426 |
+
parallel_bin_ids,
|
| 427 |
+
None, # expert_weights
|
| 428 |
+
parallel_bins,
|
| 429 |
+
expert_capacity,
|
| 430 |
+
top_k=1,
|
| 431 |
+
)
|
| 432 |
+
|
| 433 |
+
# Un-permute the tokens across the devices.
|
| 434 |
+
x, _ = all_to_all(
|
| 435 |
+
parallel_x,
|
| 436 |
+
send_counts,
|
| 437 |
+
recv_counts,
|
| 438 |
+
self.args.expert_parallel_group,
|
| 439 |
+
)
|
| 440 |
+
|
| 441 |
+
# Reduce along the hidden sharding to get the final outputs.
|
| 442 |
+
#
|
| 443 |
+
# TODO(tgale): Fuse this into the following local permutation.
|
| 444 |
+
shape = (
|
| 445 |
+
mpu.hidden_sharding_degree(self.args),
|
| 446 |
+
-1,
|
| 447 |
+
self.args.hidden_size,
|
| 448 |
+
)
|
| 449 |
+
# x = ops.sum(x.view(shape), dim=0)
|
| 450 |
+
x = x.view(shape).sum(dim=0)
|
| 451 |
+
|
| 452 |
+
# Un-permute locally to setup for the next series of operations.
|
| 453 |
+
# x = ops.scatter(x, indices, bin_ids, expert_weights, bins, self.top_k)
|
| 454 |
+
x = scatter(x, indices, bin_ids, expert_weights, bins, self.top_k)
|
| 455 |
+
return x, tokens_per_expert.flatten()
|
| 456 |
+
|
| 457 |
+
def forward(self, x: torch.Tensor, scores: torch.Tensor, expert_weights: torch.Tensor, top_experts: torch.Tensor):
|
| 458 |
+
in_shape = x.size()
|
| 459 |
+
|
| 460 |
+
# Compute the experts.
|
| 461 |
+
x, tokens_per_expert = self.forward_fn(x, expert_weights, top_experts)
|
| 462 |
+
if self.training and self.args.moe_loss_weight > 0:
|
| 463 |
+
save_load_balancing_loss((tokens_per_expert, scores))
|
| 464 |
+
x = x.view(in_shape)
|
| 465 |
+
if self.bias is not None:
|
| 466 |
+
if self.args.return_bias:
|
| 467 |
+
return x, self.bias
|
| 468 |
+
return x + self.bias
|
| 469 |
+
return x
|
| 470 |
+
|
| 471 |
+
|
| 472 |
+
class MoE(torch.nn.Module):
|
| 473 |
+
|
| 474 |
+
def __init__(self, args: Arguments):
|
| 475 |
+
super(MoE, self).__init__()
|
| 476 |
+
|
| 477 |
+
# Token router.
|
| 478 |
+
self.router = router.LearnedRouter(args)
|
| 479 |
+
|
| 480 |
+
# Expert computation helper.
|
| 481 |
+
self.experts = self._init_experts_mlp(args)
|
| 482 |
+
|
| 483 |
+
self.shared_expert = None
|
| 484 |
+
if args.shared_expert:
|
| 485 |
+
# SharedExpert computation helper.
|
| 486 |
+
self.shared_expert = sharedexpert_registry.get(args)
|
| 487 |
+
|
| 488 |
+
def _init_experts_mlp(self, args: Arguments):
|
| 489 |
+
return ParallelMLP(args)
|
| 490 |
+
|
| 491 |
+
def forward(self, x: torch.Tensor):
|
| 492 |
+
# NOTE: If we're going to cast the activations to lower precision
|
| 493 |
+
# do it before we permute the tokens to save bandwidth.
|
| 494 |
+
x = common.cast_if_autocast_enabled(x)
|
| 495 |
+
|
| 496 |
+
# Compute the expert scores and assignments.
|
| 497 |
+
scores, expert_weights, top_experts = self.router(x)
|
| 498 |
+
|
| 499 |
+
# Compute the experts.
|
| 500 |
+
out = self.experts(x, scores, expert_weights, top_experts)
|
| 501 |
+
if self.shared_expert is not None:
|
| 502 |
+
shared_expert_out = self.shared_expert(x)
|
| 503 |
+
out = self.shared_expert.add_experts_sharedexpert(
|
| 504 |
+
shared_expert_out,
|
| 505 |
+
out,
|
| 506 |
+
)
|
| 507 |
+
return out
|
build/torch214-cxx11-xpu20261-x86_64-linux/_layers/mpu.py
ADDED
|
@@ -0,0 +1,94 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2024 Databricks
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
from typing import Optional
|
| 5 |
+
|
| 6 |
+
import torch
|
| 7 |
+
import torch.distributed as dist
|
| 8 |
+
|
| 9 |
+
# from megablocks.layers.arguments import Arguments
|
| 10 |
+
from .arguments import Arguments
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
class MoeParam(torch.Tensor):
|
| 14 |
+
|
| 15 |
+
def __init__(self):
|
| 16 |
+
super().__init__(self)
|
| 17 |
+
self.expert_model_parallel: bool
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def is_moe_param(tensor: torch.Tensor) -> bool:
|
| 21 |
+
return hasattr(tensor, 'expert_model_parallel')
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
def get_expert_parallel_world_size(args: Arguments) -> int:
|
| 25 |
+
return (dist.get_world_size(args.expert_parallel_group) if args.moe_expert_model_parallelism else 1)
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def get_expert_parallel_rank(args: Arguments) -> int:
|
| 29 |
+
return (dist.get_rank(args.expert_parallel_group) if args.moe_expert_model_parallelism else 0)
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def set_expert_model_parallel_attributes(
|
| 33 |
+
tensor: torch.Tensor,
|
| 34 |
+
is_parallel: bool,
|
| 35 |
+
):
|
| 36 |
+
assert not hasattr(tensor, 'expert_model_parallel')
|
| 37 |
+
setattr(tensor, 'expert_model_parallel', is_parallel)
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
def param_is_expert_model_parallel(param: MoeParam) -> bool:
|
| 41 |
+
return (hasattr(param, 'expert_model_parallel') and param.expert_model_parallel)
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
def copy_expert_model_parallel_attributes(
|
| 45 |
+
destination_tensor: torch.Tensor,
|
| 46 |
+
source_tensor: torch.Tensor,
|
| 47 |
+
):
|
| 48 |
+
if hasattr(source_tensor, 'expert_model_parallel'):
|
| 49 |
+
setattr(
|
| 50 |
+
destination_tensor,
|
| 51 |
+
'expert_model_parallel',
|
| 52 |
+
getattr(source_tensor, 'expert_model_parallel'),
|
| 53 |
+
)
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
def synchronized_print(group: Optional[dist.ProcessGroup], *x: torch.Tensor):
|
| 57 |
+
world_size = dist.get_world_size(group)
|
| 58 |
+
rank = dist.get_rank(group)
|
| 59 |
+
for i in range(world_size):
|
| 60 |
+
dist.barrier(group)
|
| 61 |
+
if i == rank:
|
| 62 |
+
print(f'rank = {rank}', *x)
|
| 63 |
+
|
| 64 |
+
|
| 65 |
+
# Helpers for expert/tensor sharding.
|
| 66 |
+
def expert_sharding_degree(args: Arguments) -> int:
|
| 67 |
+
world_size = get_expert_parallel_world_size(args)
|
| 68 |
+
esd = min(world_size, args.moe_num_experts)
|
| 69 |
+
|
| 70 |
+
if (args.moe_num_experts % esd) != 0:
|
| 71 |
+
raise ValueError(f'Cannot shard {args.moe_num_experts} experts {esd} ways.',)
|
| 72 |
+
return esd
|
| 73 |
+
|
| 74 |
+
|
| 75 |
+
def hidden_sharding_degree(args: Arguments) -> int:
|
| 76 |
+
world_size = get_expert_parallel_world_size(args)
|
| 77 |
+
esd = expert_sharding_degree(args)
|
| 78 |
+
hsd = world_size // esd
|
| 79 |
+
|
| 80 |
+
if (args.ffn_hidden_size % hsd) != 0:
|
| 81 |
+
raise ValueError(f'Cannot shard {args.ffn_hidden_size} features {hsd} ways.',)
|
| 82 |
+
if (esd * hsd) != world_size:
|
| 83 |
+
raise ValueError(
|
| 84 |
+
f"Invalid sharding. 'expert_sharding_degree' ({esd}) * hidden_sharding_degree ({hsd}) != world_size ({world_size}).",
|
| 85 |
+
)
|
| 86 |
+
return hsd
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
def experts_per_rank(args: Arguments) -> int:
|
| 90 |
+
return args.moe_num_experts // expert_sharding_degree(args)
|
| 91 |
+
|
| 92 |
+
|
| 93 |
+
def features_per_rank(args: Arguments) -> int:
|
| 94 |
+
return args.ffn_hidden_size // hidden_sharding_degree(args)
|
build/torch214-cxx11-xpu20261-x86_64-linux/_layers/router.py
ADDED
|
@@ -0,0 +1,116 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2024 Databricks
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
from typing import Any
|
| 4 |
+
|
| 5 |
+
import torch
|
| 6 |
+
|
| 7 |
+
# from megablocks.layers import common
|
| 8 |
+
# from megablocks.layers.arguments import Arguments
|
| 9 |
+
from . import common
|
| 10 |
+
from .arguments import Arguments
|
| 11 |
+
|
| 12 |
+
_ROUTER_LOGITS = []
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
def _save_router_logits(logits: torch.Tensor, args: Arguments):
|
| 16 |
+
if args.moe_zloss_weight == 0:
|
| 17 |
+
return
|
| 18 |
+
global _ROUTER_LOGITS
|
| 19 |
+
_ROUTER_LOGITS.append(logits)
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
def clear_router_zloss():
|
| 23 |
+
global _ROUTER_LOGITS
|
| 24 |
+
_ROUTER_LOGITS.clear()
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
def batched_router_zloss(args: Arguments):
|
| 28 |
+
global _ROUTER_LOGITS
|
| 29 |
+
|
| 30 |
+
if args.moe_zloss_weight == 0:
|
| 31 |
+
import warnings
|
| 32 |
+
warnings.warn('Call to batched_router_zloss, but moe_zloss_weight=0')
|
| 33 |
+
return 0
|
| 34 |
+
|
| 35 |
+
logits_per_router = _ROUTER_LOGITS
|
| 36 |
+
|
| 37 |
+
if args.moe_zloss_in_fp32:
|
| 38 |
+
logits_per_router = [logits.float() for logits in logits_per_router]
|
| 39 |
+
|
| 40 |
+
unscaled_zloss_per_router = torch.stack([
|
| 41 |
+
torch.logsumexp(logits, dim=1).square().mean() for logits in logits_per_router
|
| 42 |
+
])
|
| 43 |
+
|
| 44 |
+
return args.moe_zloss_weight * unscaled_zloss_per_router
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
# NOTE: To enable end-to-end benchmarking without convergence we
|
| 48 |
+
# support a flag to force the router to assign tokens uniformly
|
| 49 |
+
# across the experts. We do this with a custom autograd operation
|
| 50 |
+
# so that PyTorch still executes the full set of router operation.
|
| 51 |
+
class _UniformExpertAssignment(torch.autograd.Function):
|
| 52 |
+
|
| 53 |
+
@staticmethod
|
| 54 |
+
def forward(ctx: Any, x: torch.Tensor, num_experts: int):
|
| 55 |
+
out = torch.arange(x.numel(), dtype=x.dtype, device=x.device)
|
| 56 |
+
out = torch.remainder(out, num_experts)
|
| 57 |
+
return out.view(x.shape)
|
| 58 |
+
|
| 59 |
+
|
| 60 |
+
_uniform_expert_assignment = _UniformExpertAssignment.apply
|
| 61 |
+
|
| 62 |
+
|
| 63 |
+
class LearnedRouter(torch.nn.Module):
|
| 64 |
+
|
| 65 |
+
def __init__(self, args: Arguments):
|
| 66 |
+
super().__init__()
|
| 67 |
+
self.args = args
|
| 68 |
+
|
| 69 |
+
# Learned router parameters.
|
| 70 |
+
#
|
| 71 |
+
# NOTE: This weight matrix is not parallelized with expert model
|
| 72 |
+
# parallelism. Each device needs the entire router weight matrix
|
| 73 |
+
# so that it can route its batch of data correctly.
|
| 74 |
+
self.layer = torch.nn.Linear(
|
| 75 |
+
args.hidden_size,
|
| 76 |
+
args.moe_num_experts,
|
| 77 |
+
bias=False,
|
| 78 |
+
dtype=common.dtype(args),
|
| 79 |
+
device=args.device,
|
| 80 |
+
)
|
| 81 |
+
args.init_method(self.layer.weight)
|
| 82 |
+
|
| 83 |
+
def jitter(self, x: torch.Tensor):
|
| 84 |
+
low: float = 1.0 - self.args.moe_jitter_eps
|
| 85 |
+
high: float = 1.0 + self.args.moe_jitter_eps
|
| 86 |
+
noise = torch.rand(x.size(), dtype=x.dtype, device=x.device)
|
| 87 |
+
return low + noise * (high - low)
|
| 88 |
+
|
| 89 |
+
def _top_k(self, scores: torch.Tensor):
|
| 90 |
+
if self.args.moe_top_k == 1:
|
| 91 |
+
return scores.max(dim=-1, keepdim=True)
|
| 92 |
+
return torch.topk(scores, self.args.moe_top_k, dim=-1)
|
| 93 |
+
|
| 94 |
+
def forward(self, x: torch.Tensor):
|
| 95 |
+
if self.training and self.args.moe_jitter_eps is not None:
|
| 96 |
+
x = x * self.jitter(x)
|
| 97 |
+
|
| 98 |
+
logits = self.layer(x.view(-1, x.shape[-1]))
|
| 99 |
+
_save_router_logits(logits, self.args)
|
| 100 |
+
scores = logits.softmax(dim=-1)
|
| 101 |
+
expert_weights, expert_indices = self._top_k(scores)
|
| 102 |
+
if self.args.moe_normalize_expert_weights:
|
| 103 |
+
expert_weights = expert_weights / torch.norm(
|
| 104 |
+
expert_weights,
|
| 105 |
+
p=self.args.moe_normalize_expert_weights,
|
| 106 |
+
dim=-1,
|
| 107 |
+
keepdim=True,
|
| 108 |
+
)
|
| 109 |
+
|
| 110 |
+
expert_indices = (
|
| 111 |
+
_uniform_expert_assignment(
|
| 112 |
+
expert_indices,
|
| 113 |
+
self.args.moe_num_experts,
|
| 114 |
+
) if self.args.uniform_expert_assignment else expert_indices
|
| 115 |
+
)
|
| 116 |
+
return scores, expert_weights, expert_indices
|
build/torch214-cxx11-xpu20261-x86_64-linux/_layers/sharedexpert_registry.py
ADDED
|
@@ -0,0 +1,32 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2024 Databricks
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
from typing import Union
|
| 5 |
+
|
| 6 |
+
# from megablocks.layers import glu, mlp
|
| 7 |
+
# from megablocks.layers.arguments import Arguments
|
| 8 |
+
from . import glu, mlp
|
| 9 |
+
from .arguments import Arguments
|
| 10 |
+
|
| 11 |
+
_REGISTRY = {
|
| 12 |
+
'mlp': mlp.SharedMLP,
|
| 13 |
+
'glu': glu.SharedGLU,
|
| 14 |
+
}
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
def get(args: Arguments) -> Union[mlp.SharedMLP, glu.SharedGLU]:
|
| 18 |
+
"""Returns an SharedMLP for use in a dMoE instance.
|
| 19 |
+
|
| 20 |
+
Uses the provided arguments to instantiate the appropriate
|
| 21 |
+
SharedMLP instance.
|
| 22 |
+
|
| 23 |
+
Args:
|
| 24 |
+
args: propagated Arguments dataclass.
|
| 25 |
+
|
| 26 |
+
Returns:
|
| 27 |
+
An instantiated SharedMLP constructed using the input args.
|
| 28 |
+
"""
|
| 29 |
+
if args.mlp_type not in _REGISTRY:
|
| 30 |
+
raise ValueError(f'Unsupported mlp type: {args.mlp_type}')
|
| 31 |
+
|
| 32 |
+
return _REGISTRY[args.mlp_type](args)
|
build/torch214-cxx11-xpu20261-x86_64-linux/_megablocks_xpu_addf474.abi3.so
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:6c8d62889d098114f0801bbfcf41f6212fa9d4efc9c53de340a7a2922f76aad7
|
| 3 |
+
size 3939216
|
build/torch214-cxx11-xpu20261-x86_64-linux/_ops.py
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from . import _megablocks_xpu_addf474
|
| 3 |
+
ops = torch.ops._megablocks_xpu_addf474
|
| 4 |
+
|
| 5 |
+
def add_op_namespace_prefix(op_name: str):
|
| 6 |
+
"""
|
| 7 |
+
Prefix op by namespace.
|
| 8 |
+
"""
|
| 9 |
+
return f"_megablocks_xpu_addf474::{op_name}"
|
build/torch214-cxx11-xpu20261-x86_64-linux/_version.py
ADDED
|
@@ -0,0 +1,6 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2024 Databricks
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
"""The MegaBlocks Version."""
|
| 5 |
+
|
| 6 |
+
__version__ = '0.11.0.dev0'
|
build/torch214-cxx11-xpu20261-x86_64-linux/backend/__init__.py
ADDED
|
@@ -0,0 +1,2 @@
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2024 Databricks
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
build/torch214-cxx11-xpu20261-x86_64-linux/backend/kernels.py
ADDED
|
@@ -0,0 +1,557 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2024 Databricks
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
import torch
|
| 5 |
+
import triton
|
| 6 |
+
import triton.language as tl
|
| 7 |
+
|
| 8 |
+
# Stub triton autotune when testing in a env that does not have CUDA
|
| 9 |
+
# this approach preserves the original code but enables testing without a GPU
|
| 10 |
+
if torch.cuda.is_available() is False:
|
| 11 |
+
import warnings
|
| 12 |
+
|
| 13 |
+
warnings.warn("CUDA is not available. Triton autotuning is disabled.")
|
| 14 |
+
|
| 15 |
+
def _no_autotune(*args, **kwargs):
|
| 16 |
+
def deco(fn):
|
| 17 |
+
return fn
|
| 18 |
+
return deco
|
| 19 |
+
|
| 20 |
+
triton.autotune = _no_autotune
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def assert_is_tensor(x, ndim):
|
| 24 |
+
if x.ndim != ndim:
|
| 25 |
+
raise ValueError(f'Expected {ndim}-tensor but got {x.ndim}-tensor')
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def assert_is_matrix(x):
|
| 29 |
+
assert_is_tensor(x, 2)
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def assert_is_vector(x):
|
| 33 |
+
if x.ndim != 1:
|
| 34 |
+
raise ValueError(f'Expected 1-tensor but got {x.ndim}-tensor')
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
def assert_equal(a, b):
|
| 38 |
+
if a != b:
|
| 39 |
+
raise ValueError(f'Expected dimensions to be equal but got {a} and {b}.',)
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
# a: (tokens, hidden_size), real.
|
| 43 |
+
# indices: (tokens * top_k), integer.
|
| 44 |
+
# bin_ids: (tokens * top_k), integer.
|
| 45 |
+
# weights: (tokens * top_k), real.
|
| 46 |
+
# bins: (num_experts), integer.
|
| 47 |
+
# padded_bins: (num_experts), integer.
|
| 48 |
+
@triton.autotune(
|
| 49 |
+
configs=[
|
| 50 |
+
triton.Config({'BLOCK_X': 64}, num_warps=2),
|
| 51 |
+
triton.Config({'BLOCK_X': 128}, num_warps=2),
|
| 52 |
+
triton.Config({'BLOCK_X': 256}, num_warps=2),
|
| 53 |
+
triton.Config({'BLOCK_X': 128}, num_warps=4),
|
| 54 |
+
triton.Config({'BLOCK_X': 256}, num_warps=4),
|
| 55 |
+
],
|
| 56 |
+
key=['NUM_COLUMNS'],
|
| 57 |
+
)
|
| 58 |
+
@triton.jit
|
| 59 |
+
def _padded_copy(
|
| 60 |
+
a,
|
| 61 |
+
b,
|
| 62 |
+
indices,
|
| 63 |
+
bin_ids,
|
| 64 |
+
weights,
|
| 65 |
+
bins,
|
| 66 |
+
padded_bins,
|
| 67 |
+
NUM_COLUMNS: tl.constexpr,
|
| 68 |
+
TOP_K: tl.constexpr,
|
| 69 |
+
BLOCK_X: tl.constexpr,
|
| 70 |
+
A_TO_B: tl.constexpr,
|
| 71 |
+
SCALE: tl.constexpr,
|
| 72 |
+
):
|
| 73 |
+
# Our index into array 'a'.
|
| 74 |
+
index_a = tl.load(indices + tl.program_id(0))
|
| 75 |
+
|
| 76 |
+
# One threadblock per row in 'a'. Array 'b' has greater or equal
|
| 77 |
+
# number of rows since they could be padded.
|
| 78 |
+
bin_idx = tl.load(bin_ids + tl.program_id(0))
|
| 79 |
+
|
| 80 |
+
# Now we know what bin we're assigned to, but we need to know how
|
| 81 |
+
# many threadblocks were assigned to earlier bins so we can offset
|
| 82 |
+
# in our bin properly.
|
| 83 |
+
offset_in_bin = tl.program_id(0)
|
| 84 |
+
if bin_idx > 0:
|
| 85 |
+
offset_in_bin -= tl.load(bins + bin_idx - 1)
|
| 86 |
+
|
| 87 |
+
# Load the starting index of our bin in array 'b'.
|
| 88 |
+
index_b = offset_in_bin
|
| 89 |
+
if bin_idx > 0:
|
| 90 |
+
index_b += tl.load(padded_bins + bin_idx - 1)
|
| 91 |
+
|
| 92 |
+
# Offset the input and output pointers.
|
| 93 |
+
#
|
| 94 |
+
# If we're going from A to B, divide the input index to copy
|
| 95 |
+
# the same input repeatedly. If we're going from B to A we
|
| 96 |
+
# need to reduce the result. Using atomics is slow, so we
|
| 97 |
+
# do the reduce step in a second kernel.
|
| 98 |
+
offset = index_a // TOP_K if A_TO_B else index_a
|
| 99 |
+
a += tl.multiple_of(offset * NUM_COLUMNS, NUM_COLUMNS)
|
| 100 |
+
b += tl.multiple_of(index_b * NUM_COLUMNS, NUM_COLUMNS)
|
| 101 |
+
offsets = tl.max_contiguous(tl.arange(0, BLOCK_X), BLOCK_X)
|
| 102 |
+
|
| 103 |
+
# Load the scale, if requested.
|
| 104 |
+
scale = tl.load(weights + index_a) if SCALE else 1
|
| 105 |
+
|
| 106 |
+
# Swap the pointers depending on the direction.
|
| 107 |
+
iptr = a if A_TO_B else b
|
| 108 |
+
optr = b if A_TO_B else a
|
| 109 |
+
|
| 110 |
+
iterations = tl.cdiv(NUM_COLUMNS, BLOCK_X)
|
| 111 |
+
for _ in range(iterations):
|
| 112 |
+
mask = offsets < NUM_COLUMNS
|
| 113 |
+
x = tl.load(iptr + offsets, mask=mask)
|
| 114 |
+
x = x.to(tl.float32) * scale.to(tl.float32)
|
| 115 |
+
|
| 116 |
+
tl.store(optr + offsets, x.to(optr.dtype.element_ty), mask=mask)
|
| 117 |
+
|
| 118 |
+
offsets += BLOCK_X
|
| 119 |
+
|
| 120 |
+
|
| 121 |
+
def padded_gather(x, indices, bin_ids, weights, bins, padded_bins, top_k):
|
| 122 |
+
# Validate the input shapes.
|
| 123 |
+
assert_is_matrix(x)
|
| 124 |
+
assert_is_vector(indices)
|
| 125 |
+
assert_is_vector(bin_ids)
|
| 126 |
+
assert_is_vector(bins)
|
| 127 |
+
assert_is_vector(padded_bins)
|
| 128 |
+
assert_equal(indices.shape[0], x.shape[0] * top_k)
|
| 129 |
+
assert_equal(bin_ids.shape[0], x.shape[0] * top_k)
|
| 130 |
+
assert_equal(bins.size(), padded_bins.size())
|
| 131 |
+
|
| 132 |
+
if weights is not None:
|
| 133 |
+
assert_equal(weights.shape[0], x.shape[0] * top_k)
|
| 134 |
+
|
| 135 |
+
# NOTE: Because of the padding, the output size is dynamic.
|
| 136 |
+
# We load the final padded bin bound to get the output rows.
|
| 137 |
+
output_rows = padded_bins[-1].cpu().item()
|
| 138 |
+
out = torch.zeros((output_rows, x.shape[1]), dtype=x.dtype, device=x.device)
|
| 139 |
+
_padded_copy[(indices.shape[0],)](
|
| 140 |
+
x,
|
| 141 |
+
out,
|
| 142 |
+
indices,
|
| 143 |
+
bin_ids,
|
| 144 |
+
weights,
|
| 145 |
+
bins,
|
| 146 |
+
padded_bins,
|
| 147 |
+
NUM_COLUMNS=x.shape[1],
|
| 148 |
+
A_TO_B=True,
|
| 149 |
+
TOP_K=top_k,
|
| 150 |
+
SCALE=weights is not None,
|
| 151 |
+
)
|
| 152 |
+
return out
|
| 153 |
+
|
| 154 |
+
|
| 155 |
+
def gather(x, indices, bin_ids, weights, bins, top_k):
|
| 156 |
+
# Validate the input shapes.
|
| 157 |
+
assert_is_matrix(x)
|
| 158 |
+
assert_is_vector(indices)
|
| 159 |
+
assert_is_vector(bin_ids)
|
| 160 |
+
assert_is_vector(bins)
|
| 161 |
+
assert_equal(indices.shape[0], x.shape[0] * top_k)
|
| 162 |
+
assert_equal(bin_ids.shape[0], x.shape[0] * top_k)
|
| 163 |
+
|
| 164 |
+
if weights is not None:
|
| 165 |
+
assert_equal(weights.shape[0], x.shape[0] * top_k)
|
| 166 |
+
|
| 167 |
+
# NOTE: There is no padding so the output rows equals the
|
| 168 |
+
# input rows multiplied by top_k.
|
| 169 |
+
output_rows = x.shape[0] * top_k
|
| 170 |
+
out = torch.empty((output_rows, x.shape[1]), dtype=x.dtype, device=x.device)
|
| 171 |
+
_padded_copy[(indices.shape[0],)](
|
| 172 |
+
x,
|
| 173 |
+
out,
|
| 174 |
+
indices,
|
| 175 |
+
bin_ids,
|
| 176 |
+
weights,
|
| 177 |
+
bins,
|
| 178 |
+
bins,
|
| 179 |
+
NUM_COLUMNS=x.shape[1],
|
| 180 |
+
A_TO_B=True,
|
| 181 |
+
TOP_K=top_k,
|
| 182 |
+
SCALE=weights is not None,
|
| 183 |
+
)
|
| 184 |
+
return out
|
| 185 |
+
|
| 186 |
+
|
| 187 |
+
def padded_scatter(x, indices, bin_ids, weights, bins, padded_bins, top_k):
|
| 188 |
+
# Validate the input shapes.
|
| 189 |
+
assert_is_matrix(x)
|
| 190 |
+
assert_is_vector(indices)
|
| 191 |
+
assert_is_vector(bin_ids)
|
| 192 |
+
assert_is_vector(bins)
|
| 193 |
+
assert_is_vector(padded_bins)
|
| 194 |
+
assert_equal(indices.shape[0], bin_ids.shape[0])
|
| 195 |
+
assert_equal(bins.size(), padded_bins.size())
|
| 196 |
+
|
| 197 |
+
if weights is not None:
|
| 198 |
+
assert_equal(indices.shape[0], weights.shape[0])
|
| 199 |
+
|
| 200 |
+
tokens = indices.shape[0] // top_k
|
| 201 |
+
out = torch.empty((tokens, top_k, x.shape[1]), dtype=x.dtype, device=x.device)
|
| 202 |
+
_padded_copy[(indices.shape[0],)](
|
| 203 |
+
out,
|
| 204 |
+
x,
|
| 205 |
+
indices,
|
| 206 |
+
bin_ids,
|
| 207 |
+
weights,
|
| 208 |
+
bins,
|
| 209 |
+
padded_bins,
|
| 210 |
+
NUM_COLUMNS=x.shape[1],
|
| 211 |
+
A_TO_B=False,
|
| 212 |
+
TOP_K=top_k,
|
| 213 |
+
SCALE=weights is not None,
|
| 214 |
+
)
|
| 215 |
+
|
| 216 |
+
# Reduce along the top-k dimension, if needed.
|
| 217 |
+
return out.sum(dim=1) if top_k > 1 else out.view(tokens, x.shape[1])
|
| 218 |
+
|
| 219 |
+
|
| 220 |
+
def scatter(x, indices, bin_ids, weights, bins, top_k):
|
| 221 |
+
return padded_scatter(x, indices, bin_ids, weights, bins, bins, top_k)
|
| 222 |
+
|
| 223 |
+
|
| 224 |
+
# x: (tokens, top_k, hidden_size), real
|
| 225 |
+
# grad: (tokens, hidden_size), real.
|
| 226 |
+
# wgrad: (tokens, top_k), real.
|
| 227 |
+
# indices: (tokens * top_k), integer.
|
| 228 |
+
# bin_ids: (tokens * top_k), integer.
|
| 229 |
+
# bins: (num_experts), integer.
|
| 230 |
+
# padded_bins: (num_experts), integer.
|
| 231 |
+
@triton.autotune(
|
| 232 |
+
configs=[
|
| 233 |
+
triton.Config({'BLOCK_X': 64}, num_warps=2),
|
| 234 |
+
triton.Config({'BLOCK_X': 128}, num_warps=2),
|
| 235 |
+
triton.Config({'BLOCK_X': 256}, num_warps=2),
|
| 236 |
+
triton.Config({'BLOCK_X': 128}, num_warps=4),
|
| 237 |
+
triton.Config({'BLOCK_X': 256}, num_warps=4),
|
| 238 |
+
],
|
| 239 |
+
key=['NUM_COLUMNS'],
|
| 240 |
+
)
|
| 241 |
+
@triton.jit
|
| 242 |
+
def _padded_copy_wgrad(
|
| 243 |
+
x,
|
| 244 |
+
grad,
|
| 245 |
+
wgrad,
|
| 246 |
+
indices,
|
| 247 |
+
bin_ids,
|
| 248 |
+
bins,
|
| 249 |
+
padded_bins,
|
| 250 |
+
NUM_COLUMNS: tl.constexpr,
|
| 251 |
+
TOP_K: tl.constexpr,
|
| 252 |
+
BLOCK_X: tl.constexpr,
|
| 253 |
+
):
|
| 254 |
+
# Our index into 'tokens * top_k'.
|
| 255 |
+
index_out = tl.load(indices + tl.program_id(0))
|
| 256 |
+
|
| 257 |
+
# One threadblock per row in 'a'. Array 'b' has greater or equal
|
| 258 |
+
# number of rows since they could be padded.
|
| 259 |
+
bin_idx = tl.load(bin_ids + tl.program_id(0))
|
| 260 |
+
|
| 261 |
+
# Now we know what bin we're assigned to, but we need to know how
|
| 262 |
+
# many threadblocks were assigned to earlier bins so we can offset
|
| 263 |
+
# in our bin properly.
|
| 264 |
+
offset_in_bin = tl.program_id(0)
|
| 265 |
+
if bin_idx > 0:
|
| 266 |
+
offset_in_bin -= tl.load(bins + bin_idx - 1)
|
| 267 |
+
|
| 268 |
+
# Load the starting index of our bin in array 'x'.
|
| 269 |
+
index_x = offset_in_bin
|
| 270 |
+
if bin_idx > 0:
|
| 271 |
+
index_x += tl.load(padded_bins + bin_idx - 1)
|
| 272 |
+
|
| 273 |
+
# Offset the input and output pointers.
|
| 274 |
+
wgrad += index_out
|
| 275 |
+
grad += tl.multiple_of((index_out // TOP_K) * NUM_COLUMNS, NUM_COLUMNS)
|
| 276 |
+
x += tl.multiple_of(index_x * NUM_COLUMNS, NUM_COLUMNS)
|
| 277 |
+
offsets = tl.max_contiguous(tl.arange(0, BLOCK_X), BLOCK_X)
|
| 278 |
+
|
| 279 |
+
acc = tl.zeros((BLOCK_X,), dtype=tl.float32)
|
| 280 |
+
iterations = tl.cdiv(NUM_COLUMNS, BLOCK_X)
|
| 281 |
+
for _ in range(iterations):
|
| 282 |
+
mask = offsets < NUM_COLUMNS
|
| 283 |
+
data = tl.load(x + offsets, mask=mask).to(tl.float32)
|
| 284 |
+
scale = tl.load(grad + offsets, mask=mask).to(tl.float32)
|
| 285 |
+
acc += data * scale
|
| 286 |
+
offsets += BLOCK_X
|
| 287 |
+
|
| 288 |
+
# Reduce to get the final result and store.
|
| 289 |
+
out = tl.sum(acc).to(wgrad.dtype.element_ty)
|
| 290 |
+
tl.store(wgrad, out)
|
| 291 |
+
|
| 292 |
+
|
| 293 |
+
def padded_scatter_wgrad(x, grad, indices, bin_ids, bins, padded_bins, top_k):
|
| 294 |
+
# Validate the input shapes.
|
| 295 |
+
assert_is_matrix(x)
|
| 296 |
+
assert_is_matrix(grad)
|
| 297 |
+
assert_is_vector(indices)
|
| 298 |
+
assert_is_vector(bin_ids)
|
| 299 |
+
assert_is_vector(bins)
|
| 300 |
+
assert_is_vector(padded_bins)
|
| 301 |
+
assert_equal(indices.shape[0], bin_ids.shape[0])
|
| 302 |
+
assert_equal(bins.size(), padded_bins.size())
|
| 303 |
+
|
| 304 |
+
tokens = indices.shape[0] // top_k
|
| 305 |
+
out = torch.empty((tokens * top_k), dtype=x.dtype, device=x.device)
|
| 306 |
+
_padded_copy_wgrad[(indices.shape[0],)](
|
| 307 |
+
x,
|
| 308 |
+
grad,
|
| 309 |
+
out,
|
| 310 |
+
indices,
|
| 311 |
+
bin_ids,
|
| 312 |
+
bins,
|
| 313 |
+
padded_bins,
|
| 314 |
+
NUM_COLUMNS=x.shape[1],
|
| 315 |
+
TOP_K=top_k,
|
| 316 |
+
)
|
| 317 |
+
return out
|
| 318 |
+
|
| 319 |
+
|
| 320 |
+
def scatter_wgrad(x, grad, indices, bin_ids, bins, top_k):
|
| 321 |
+
return padded_scatter_wgrad(x, grad, indices, bin_ids, bins, bins, top_k)
|
| 322 |
+
|
| 323 |
+
|
| 324 |
+
# a: (tokens, hidden_size), real.
|
| 325 |
+
# b: (num_experts, expert_capacity, num_columns), real.
|
| 326 |
+
# indices: (tokens * top_k), integer.
|
| 327 |
+
# weights: (tokens * top_k), real.
|
| 328 |
+
# bins: (num_experts), integer.
|
| 329 |
+
@triton.autotune(
|
| 330 |
+
configs=[
|
| 331 |
+
triton.Config({'BLOCK_X': 64}, num_warps=2),
|
| 332 |
+
triton.Config({'BLOCK_X': 128}, num_warps=2),
|
| 333 |
+
triton.Config({'BLOCK_X': 256}, num_warps=2),
|
| 334 |
+
triton.Config({'BLOCK_X': 128}, num_warps=4),
|
| 335 |
+
triton.Config({'BLOCK_X': 256}, num_warps=4),
|
| 336 |
+
],
|
| 337 |
+
key=['NUM_COLUMNS'],
|
| 338 |
+
)
|
| 339 |
+
@triton.jit
|
| 340 |
+
def _binned_copy(
|
| 341 |
+
a,
|
| 342 |
+
b,
|
| 343 |
+
num_experts,
|
| 344 |
+
expert_capacity,
|
| 345 |
+
indices,
|
| 346 |
+
weights,
|
| 347 |
+
bins,
|
| 348 |
+
NUM_COLUMNS: tl.constexpr,
|
| 349 |
+
TOP_K: tl.constexpr,
|
| 350 |
+
BLOCK_X: tl.constexpr,
|
| 351 |
+
A_TO_B: tl.constexpr,
|
| 352 |
+
SCALE: tl.constexpr,
|
| 353 |
+
):
|
| 354 |
+
# Load our indices into the output.
|
| 355 |
+
expert_idx = tl.program_id(0)
|
| 356 |
+
entry_idx = tl.program_id(1)
|
| 357 |
+
|
| 358 |
+
# Calculate our offset into the output.
|
| 359 |
+
index_b = expert_idx * expert_capacity + entry_idx
|
| 360 |
+
|
| 361 |
+
# Load the index bounds for our bin and calculate
|
| 362 |
+
# the number of tokens assigned to our expert.
|
| 363 |
+
start = 0
|
| 364 |
+
if expert_idx > 0:
|
| 365 |
+
start = tl.load(bins + expert_idx - 1)
|
| 366 |
+
end = tl.load(bins + expert_idx)
|
| 367 |
+
num_tokens = end - start
|
| 368 |
+
|
| 369 |
+
# Calculate our offset into the input. If we don't
|
| 370 |
+
# have an input exit early.
|
| 371 |
+
if entry_idx >= num_tokens:
|
| 372 |
+
return
|
| 373 |
+
index_a = tl.load(indices + start + entry_idx)
|
| 374 |
+
|
| 375 |
+
# Offset the input and output pointers.
|
| 376 |
+
#
|
| 377 |
+
# If we're going from A to B, divide the input index to copy
|
| 378 |
+
# the same input repeatedly. If we're going from B to A we
|
| 379 |
+
# need to reduce the result. Using atomics is slow, so we
|
| 380 |
+
# do the reduce step in a second kernel.
|
| 381 |
+
offset = index_a // TOP_K if A_TO_B else index_a
|
| 382 |
+
a += tl.multiple_of(offset * NUM_COLUMNS, NUM_COLUMNS)
|
| 383 |
+
b += tl.multiple_of(index_b * NUM_COLUMNS, NUM_COLUMNS)
|
| 384 |
+
offsets = tl.max_contiguous(tl.arange(0, BLOCK_X), BLOCK_X)
|
| 385 |
+
|
| 386 |
+
# Load the scale, if requested.
|
| 387 |
+
scale = tl.load(weights + index_a) if SCALE else 1
|
| 388 |
+
|
| 389 |
+
# Swap the pointers depending on the direction.
|
| 390 |
+
#
|
| 391 |
+
# NOTE: We need to zero the output in both directions.
|
| 392 |
+
iptr = a if A_TO_B else b
|
| 393 |
+
optr = b if A_TO_B else a
|
| 394 |
+
|
| 395 |
+
iterations = tl.cdiv(NUM_COLUMNS, BLOCK_X)
|
| 396 |
+
for _ in range(iterations):
|
| 397 |
+
mask = offsets < NUM_COLUMNS
|
| 398 |
+
x = tl.load(iptr + offsets, mask=mask)
|
| 399 |
+
x = x.to(tl.float32) * scale.to(tl.float32)
|
| 400 |
+
|
| 401 |
+
tl.store(optr + offsets, x.to(optr.dtype.element_ty), mask=mask)
|
| 402 |
+
|
| 403 |
+
offsets += BLOCK_X
|
| 404 |
+
|
| 405 |
+
|
| 406 |
+
def binned_gather(x, indices, weights, bins, expert_capacity, top_k):
|
| 407 |
+
# Validate the input shapes.
|
| 408 |
+
assert_is_matrix(x)
|
| 409 |
+
assert_is_vector(indices)
|
| 410 |
+
assert_is_vector(bins)
|
| 411 |
+
assert_equal(indices.shape[0], x.shape[0] * top_k)
|
| 412 |
+
|
| 413 |
+
if weights is not None:
|
| 414 |
+
assert_equal(weights.shape[0], x.shape[0] * top_k)
|
| 415 |
+
|
| 416 |
+
num_experts = bins.shape[0]
|
| 417 |
+
out = torch.zeros((num_experts, expert_capacity, x.shape[1]), dtype=x.dtype, device=x.device)
|
| 418 |
+
|
| 419 |
+
_binned_copy[(num_experts, expert_capacity)](
|
| 420 |
+
x,
|
| 421 |
+
out,
|
| 422 |
+
num_experts,
|
| 423 |
+
expert_capacity,
|
| 424 |
+
indices,
|
| 425 |
+
weights,
|
| 426 |
+
bins,
|
| 427 |
+
NUM_COLUMNS=x.shape[1],
|
| 428 |
+
A_TO_B=True,
|
| 429 |
+
TOP_K=top_k,
|
| 430 |
+
SCALE=weights is not None,
|
| 431 |
+
)
|
| 432 |
+
return out
|
| 433 |
+
|
| 434 |
+
|
| 435 |
+
def binned_scatter(x, indices, weights, bins, top_k):
|
| 436 |
+
# Validate the input shapes.
|
| 437 |
+
assert_is_tensor(x, 3)
|
| 438 |
+
assert_is_vector(indices)
|
| 439 |
+
assert_is_vector(bins)
|
| 440 |
+
assert_equal(bins.shape[0], x.shape[0])
|
| 441 |
+
|
| 442 |
+
if weights is not None:
|
| 443 |
+
assert_equal(indices.shape[0], weights.shape[0])
|
| 444 |
+
|
| 445 |
+
num_experts, expert_capacity, hidden_size = x.shape
|
| 446 |
+
tokens = indices.shape[0] // top_k
|
| 447 |
+
out = torch.zeros((tokens, top_k, hidden_size), dtype=x.dtype, device=x.device)
|
| 448 |
+
_binned_copy[(num_experts, expert_capacity)](
|
| 449 |
+
out,
|
| 450 |
+
x,
|
| 451 |
+
num_experts,
|
| 452 |
+
expert_capacity,
|
| 453 |
+
indices,
|
| 454 |
+
weights,
|
| 455 |
+
bins,
|
| 456 |
+
NUM_COLUMNS=hidden_size,
|
| 457 |
+
A_TO_B=False,
|
| 458 |
+
TOP_K=top_k,
|
| 459 |
+
SCALE=weights is not None,
|
| 460 |
+
)
|
| 461 |
+
|
| 462 |
+
# Reduce along the top-k dimension, if needed.
|
| 463 |
+
return out.sum(dim=1) if top_k > 1 else out.view(tokens, hidden_size)
|
| 464 |
+
|
| 465 |
+
|
| 466 |
+
# a: (tokens, hidden_size), real.
|
| 467 |
+
# b: (num_experts, expert_capacity, num_columns), real.
|
| 468 |
+
# indices: (tokens * top_k), integer.
|
| 469 |
+
# weights: (tokens * top_k), real.
|
| 470 |
+
# bins: (num_experts), integer.
|
| 471 |
+
@triton.autotune(
|
| 472 |
+
configs=[
|
| 473 |
+
triton.Config({'BLOCK_X': 64}, num_warps=2),
|
| 474 |
+
triton.Config({'BLOCK_X': 128}, num_warps=2),
|
| 475 |
+
triton.Config({'BLOCK_X': 256}, num_warps=2),
|
| 476 |
+
triton.Config({'BLOCK_X': 128}, num_warps=4),
|
| 477 |
+
triton.Config({'BLOCK_X': 256}, num_warps=4),
|
| 478 |
+
],
|
| 479 |
+
key=['NUM_COLUMNS'],
|
| 480 |
+
)
|
| 481 |
+
@triton.jit
|
| 482 |
+
def _binned_copy_wgrad(
|
| 483 |
+
x,
|
| 484 |
+
grad,
|
| 485 |
+
wgrad,
|
| 486 |
+
num_experts,
|
| 487 |
+
expert_capacity,
|
| 488 |
+
indices,
|
| 489 |
+
bins,
|
| 490 |
+
NUM_COLUMNS: tl.constexpr,
|
| 491 |
+
TOP_K: tl.constexpr,
|
| 492 |
+
BLOCK_X: tl.constexpr,
|
| 493 |
+
):
|
| 494 |
+
# Load our indices into the output.
|
| 495 |
+
expert_idx = tl.program_id(0)
|
| 496 |
+
entry_idx = tl.program_id(1)
|
| 497 |
+
|
| 498 |
+
# Calculate our offset into the output.
|
| 499 |
+
index_x = expert_idx * expert_capacity + entry_idx
|
| 500 |
+
|
| 501 |
+
# Load the index bounds for our bin and calculate
|
| 502 |
+
# the number of tokens assigned to our expert.
|
| 503 |
+
start = 0
|
| 504 |
+
if expert_idx > 0:
|
| 505 |
+
start = tl.load(bins + expert_idx - 1)
|
| 506 |
+
end = tl.load(bins + expert_idx)
|
| 507 |
+
num_tokens = end - start
|
| 508 |
+
|
| 509 |
+
# Calculate our offset into the input. If we don't
|
| 510 |
+
# have an input exit early.
|
| 511 |
+
if entry_idx >= num_tokens:
|
| 512 |
+
return
|
| 513 |
+
index_out = tl.load(indices + start + entry_idx)
|
| 514 |
+
|
| 515 |
+
# Offset the input and output pointers.
|
| 516 |
+
wgrad += index_out
|
| 517 |
+
grad += tl.multiple_of((index_out // TOP_K) * NUM_COLUMNS, NUM_COLUMNS)
|
| 518 |
+
x += tl.multiple_of(index_x * NUM_COLUMNS, NUM_COLUMNS)
|
| 519 |
+
offsets = tl.max_contiguous(tl.arange(0, BLOCK_X), BLOCK_X)
|
| 520 |
+
|
| 521 |
+
acc = tl.zeros((BLOCK_X,), dtype=tl.float32)
|
| 522 |
+
iterations = tl.cdiv(NUM_COLUMNS, BLOCK_X)
|
| 523 |
+
for _ in range(iterations):
|
| 524 |
+
mask = offsets < NUM_COLUMNS
|
| 525 |
+
data = tl.load(x + offsets, mask=mask).to(tl.float32)
|
| 526 |
+
scale = tl.load(grad + offsets, mask=mask).to(tl.float32)
|
| 527 |
+
acc += data * scale
|
| 528 |
+
offsets += BLOCK_X
|
| 529 |
+
|
| 530 |
+
# Reduce to get the final result and store.
|
| 531 |
+
out = tl.sum(acc).to(wgrad.dtype.element_ty)
|
| 532 |
+
tl.store(wgrad, out)
|
| 533 |
+
|
| 534 |
+
|
| 535 |
+
def binned_scatter_wgrad(x, grad, indices, bins, top_k):
|
| 536 |
+
# Validate the input shapes.
|
| 537 |
+
assert_is_tensor(x, 3)
|
| 538 |
+
assert_is_matrix(grad)
|
| 539 |
+
assert_is_vector(indices)
|
| 540 |
+
assert_is_vector(bins)
|
| 541 |
+
assert_equal(bins.shape[0], x.shape[0])
|
| 542 |
+
|
| 543 |
+
num_experts, expert_capacity, hidden_size = x.shape
|
| 544 |
+
tokens = indices.shape[0] // top_k
|
| 545 |
+
out = torch.zeros((tokens * top_k), dtype=x.dtype, device=x.device)
|
| 546 |
+
_binned_copy_wgrad[(num_experts, expert_capacity)](
|
| 547 |
+
x,
|
| 548 |
+
grad,
|
| 549 |
+
out,
|
| 550 |
+
num_experts,
|
| 551 |
+
expert_capacity,
|
| 552 |
+
indices,
|
| 553 |
+
bins,
|
| 554 |
+
NUM_COLUMNS=hidden_size,
|
| 555 |
+
TOP_K=top_k,
|
| 556 |
+
)
|
| 557 |
+
return out
|
build/torch214-cxx11-xpu20261-x86_64-linux/benchmark_util.py
ADDED
|
@@ -0,0 +1,35 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2024 Databricks
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
import numpy as np
|
| 5 |
+
import torch
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
def log_benchmark(name, arguments, time, std):
|
| 9 |
+
print('=' * 60)
|
| 10 |
+
print(f'{name} Benchmark')
|
| 11 |
+
print('Benchmark Parameters:')
|
| 12 |
+
for (key, value) in arguments.items():
|
| 13 |
+
print(f'{key} = {value}')
|
| 14 |
+
print('Results:')
|
| 15 |
+
print('mean time = {:.3f}ms, std time = {:.3f}ms'.format(time, std))
|
| 16 |
+
print('=' * 60)
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
def benchmark_function(fn, iterations=100, warmup=10):
|
| 20 |
+
# Warmup iterations.
|
| 21 |
+
for _ in range(warmup):
|
| 22 |
+
fn()
|
| 23 |
+
|
| 24 |
+
times = []
|
| 25 |
+
for i in range(iterations):
|
| 26 |
+
start = torch.cuda.Event(enable_timing=True)
|
| 27 |
+
end = torch.cuda.Event(enable_timing=True)
|
| 28 |
+
|
| 29 |
+
start.record()
|
| 30 |
+
fn()
|
| 31 |
+
end.record()
|
| 32 |
+
|
| 33 |
+
torch.cuda.synchronize()
|
| 34 |
+
times.append(start.elapsed_time(end))
|
| 35 |
+
return np.mean(times), np.std(times)
|
build/torch214-cxx11-xpu20261-x86_64-linux/cpu_fused_moe.py
ADDED
|
@@ -0,0 +1,311 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
# MegaBlocks CPU Fused MoE Implementation
|
| 3 |
+
#
|
| 4 |
+
# This is a pure Python/PyTorch implementation for CPU.
|
| 5 |
+
# For better performance, consider using the C++ kernel implementation.
|
| 6 |
+
#
|
| 7 |
+
import torch
|
| 8 |
+
import torch.nn.functional as F
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
def swigluoai_activation(gate: torch.Tensor, up: torch.Tensor,
|
| 12 |
+
alpha: float = 1.702, limit: float = 7.0) -> torch.Tensor:
|
| 13 |
+
"""
|
| 14 |
+
SwigluOAI activation function used in GptOss models.
|
| 15 |
+
|
| 16 |
+
Formula:
|
| 17 |
+
gate = clamp(gate, max=limit)
|
| 18 |
+
up = clamp(up, -limit, limit)
|
| 19 |
+
glu = gate * sigmoid(gate * alpha)
|
| 20 |
+
output = (up + 1) * glu
|
| 21 |
+
|
| 22 |
+
Args:
|
| 23 |
+
gate: Gate tensor from gate projection
|
| 24 |
+
up: Up tensor from up projection
|
| 25 |
+
alpha: Scaling factor for sigmoid (default: 1.702)
|
| 26 |
+
limit: Clamp limit (default: 7.0)
|
| 27 |
+
|
| 28 |
+
Returns:
|
| 29 |
+
Activated tensor
|
| 30 |
+
"""
|
| 31 |
+
gate = gate.clamp(max=limit)
|
| 32 |
+
up = up.clamp(min=-limit, max=limit)
|
| 33 |
+
glu = gate * torch.sigmoid(gate * alpha)
|
| 34 |
+
return (up + 1) * glu
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
def silu_and_mul_activation(gate: torch.Tensor, up: torch.Tensor) -> torch.Tensor:
|
| 38 |
+
"""
|
| 39 |
+
SiLU (Swish) activation with element-wise multiplication.
|
| 40 |
+
|
| 41 |
+
Formula:
|
| 42 |
+
output = silu(gate) * up
|
| 43 |
+
|
| 44 |
+
Args:
|
| 45 |
+
gate: Gate tensor
|
| 46 |
+
up: Up tensor
|
| 47 |
+
|
| 48 |
+
Returns:
|
| 49 |
+
Activated tensor
|
| 50 |
+
"""
|
| 51 |
+
return F.silu(gate) * up
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
def route_tokens_cpu(
|
| 55 |
+
x: torch.Tensor,
|
| 56 |
+
router_weight: torch.Tensor,
|
| 57 |
+
router_bias: torch.Tensor | None,
|
| 58 |
+
moe_top_k: int,
|
| 59 |
+
moe_num_experts: int,
|
| 60 |
+
moe_normalize_expert_weights: int | None = None,
|
| 61 |
+
) -> tuple:
|
| 62 |
+
"""
|
| 63 |
+
Route tokens to experts and compute expert weights and indices (CPU version).
|
| 64 |
+
|
| 65 |
+
Args:
|
| 66 |
+
x: Input tensor [batch, seq, hidden] or [tokens, hidden]
|
| 67 |
+
router_weight: Router weight [num_experts, hidden]
|
| 68 |
+
router_bias: Router bias [num_experts] or None
|
| 69 |
+
moe_top_k: Number of experts per token
|
| 70 |
+
moe_num_experts: Total number of experts
|
| 71 |
+
moe_normalize_expert_weights: Normalization order or None
|
| 72 |
+
|
| 73 |
+
Returns:
|
| 74 |
+
Tuple of (logits, expert_weights, expert_indices)
|
| 75 |
+
"""
|
| 76 |
+
x_flat = x.view(-1, x.shape[-1])
|
| 77 |
+
logits = F.linear(x_flat, router_weight, router_bias)
|
| 78 |
+
|
| 79 |
+
if moe_top_k == 1:
|
| 80 |
+
expert_weights, expert_indices = logits.max(dim=-1, keepdim=True)
|
| 81 |
+
else:
|
| 82 |
+
expert_weights, expert_indices = torch.topk(logits, moe_top_k, dim=-1)
|
| 83 |
+
|
| 84 |
+
expert_weights = expert_weights.softmax(dim=-1)
|
| 85 |
+
|
| 86 |
+
if moe_normalize_expert_weights is not None:
|
| 87 |
+
expert_weights = expert_weights / torch.norm(
|
| 88 |
+
expert_weights,
|
| 89 |
+
p=moe_normalize_expert_weights,
|
| 90 |
+
dim=-1,
|
| 91 |
+
keepdim=True,
|
| 92 |
+
)
|
| 93 |
+
|
| 94 |
+
return logits, expert_weights, expert_indices
|
| 95 |
+
|
| 96 |
+
|
| 97 |
+
def cpu_fused_moe(
|
| 98 |
+
hidden_states: torch.Tensor,
|
| 99 |
+
w1: torch.Tensor,
|
| 100 |
+
w2: torch.Tensor,
|
| 101 |
+
topk_weights: torch.Tensor,
|
| 102 |
+
topk_ids: torch.Tensor,
|
| 103 |
+
w1_bias: torch.Tensor | None = None,
|
| 104 |
+
w2_bias: torch.Tensor | None = None,
|
| 105 |
+
activation: str = "silu",
|
| 106 |
+
alpha: float = 1.702,
|
| 107 |
+
limit: float = 7.0,
|
| 108 |
+
is_interleaved: bool = True,
|
| 109 |
+
) -> torch.Tensor:
|
| 110 |
+
"""
|
| 111 |
+
CPU Fused MoE using PyTorch operations.
|
| 112 |
+
|
| 113 |
+
This implementation processes all experts in parallel using batched operations
|
| 114 |
+
instead of sequential for loops, which is more efficient on CPU.
|
| 115 |
+
|
| 116 |
+
Args:
|
| 117 |
+
hidden_states: [num_tokens, hidden_size]
|
| 118 |
+
w1: [num_experts, hidden_size, 2*inter_size] - gate_up_proj weights
|
| 119 |
+
w2: [num_experts, inter_size, hidden_size] - down_proj weights
|
| 120 |
+
topk_weights: [num_tokens, topk] - routing weights
|
| 121 |
+
topk_ids: [num_tokens, topk] - expert indices
|
| 122 |
+
w1_bias: [num_experts, 2*inter_size] or None
|
| 123 |
+
w2_bias: [num_experts, hidden_size] or None
|
| 124 |
+
activation: "silu" or "swigluoai"
|
| 125 |
+
alpha: swigluoai alpha parameter
|
| 126 |
+
limit: swigluoai limit parameter
|
| 127 |
+
is_interleaved: whether gate_up is interleaved [g0,u0,g1,u1,...] (True for GptOss)
|
| 128 |
+
|
| 129 |
+
Returns:
|
| 130 |
+
output: [num_tokens, hidden_size]
|
| 131 |
+
"""
|
| 132 |
+
num_tokens, hidden_size = hidden_states.shape
|
| 133 |
+
num_experts = w1.shape[0]
|
| 134 |
+
inter_size = w2.shape[1]
|
| 135 |
+
topk = topk_weights.shape[1]
|
| 136 |
+
|
| 137 |
+
# Initialize output
|
| 138 |
+
output = torch.zeros_like(hidden_states)
|
| 139 |
+
|
| 140 |
+
# Build expert mask: which tokens go to which expert
|
| 141 |
+
# expert_mask[expert_id] contains indices of (token_idx, topk_pos) pairs
|
| 142 |
+
for expert_idx in range(num_experts):
|
| 143 |
+
# Find tokens assigned to this expert
|
| 144 |
+
# mask shape: [num_tokens, topk], True where topk_ids == expert_idx
|
| 145 |
+
mask = (topk_ids == expert_idx)
|
| 146 |
+
|
| 147 |
+
if not mask.any():
|
| 148 |
+
continue
|
| 149 |
+
|
| 150 |
+
# Get token indices and topk positions
|
| 151 |
+
token_indices, topk_positions = torch.where(mask)
|
| 152 |
+
|
| 153 |
+
if len(token_indices) == 0:
|
| 154 |
+
continue
|
| 155 |
+
|
| 156 |
+
# Gather input tokens for this expert
|
| 157 |
+
# current_hidden: [num_selected_tokens, hidden_size]
|
| 158 |
+
current_hidden = hidden_states[token_indices]
|
| 159 |
+
|
| 160 |
+
# Get weights for this expert
|
| 161 |
+
# w1[expert_idx]: [hidden_size, 2*inter_size]
|
| 162 |
+
# w2[expert_idx]: [inter_size, hidden_size]
|
| 163 |
+
expert_w1 = w1[expert_idx] # [hidden_size, 2*inter_size]
|
| 164 |
+
expert_w2 = w2[expert_idx] # [inter_size, hidden_size]
|
| 165 |
+
|
| 166 |
+
# First projection: hidden @ w1 -> [num_selected, 2*inter_size]
|
| 167 |
+
gate_up = current_hidden @ expert_w1
|
| 168 |
+
|
| 169 |
+
# Add bias if present
|
| 170 |
+
if w1_bias is not None:
|
| 171 |
+
gate_up = gate_up + w1_bias[expert_idx]
|
| 172 |
+
|
| 173 |
+
# Split gate and up projections
|
| 174 |
+
if is_interleaved:
|
| 175 |
+
# GptOss uses interleaved layout: [g0, u0, g1, u1, ...]
|
| 176 |
+
gate = gate_up[..., ::2] # [num_selected, inter_size]
|
| 177 |
+
up = gate_up[..., 1::2] # [num_selected, inter_size]
|
| 178 |
+
else:
|
| 179 |
+
# Standard layout: [gate_all, up_all]
|
| 180 |
+
gate = gate_up[..., :inter_size]
|
| 181 |
+
up = gate_up[..., inter_size:]
|
| 182 |
+
|
| 183 |
+
# Apply activation
|
| 184 |
+
if activation == "swigluoai":
|
| 185 |
+
activated = swigluoai_activation(gate, up, alpha, limit)
|
| 186 |
+
else: # silu
|
| 187 |
+
activated = silu_and_mul_activation(gate, up)
|
| 188 |
+
|
| 189 |
+
# Second projection: activated @ w2 -> [num_selected, hidden_size]
|
| 190 |
+
expert_out = activated @ expert_w2
|
| 191 |
+
|
| 192 |
+
# Add bias if present
|
| 193 |
+
if w2_bias is not None:
|
| 194 |
+
expert_out = expert_out + w2_bias[expert_idx]
|
| 195 |
+
|
| 196 |
+
# Apply routing weights and accumulate
|
| 197 |
+
# weights shape: [num_selected]
|
| 198 |
+
weights = topk_weights[token_indices, topk_positions].unsqueeze(-1)
|
| 199 |
+
weighted_out = expert_out * weights
|
| 200 |
+
|
| 201 |
+
# Accumulate to output
|
| 202 |
+
output.index_add_(0, token_indices, weighted_out.to(output.dtype))
|
| 203 |
+
|
| 204 |
+
return output
|
| 205 |
+
|
| 206 |
+
|
| 207 |
+
class MegaBlocksMoeMLP(torch.nn.Module):
|
| 208 |
+
"""
|
| 209 |
+
CPU MoE MLP module that can be used as a drop-in replacement for
|
| 210 |
+
the transformers GptOssMLP when using @use_kernel_forward_from_hub.
|
| 211 |
+
"""
|
| 212 |
+
can_torch_compile: bool = True
|
| 213 |
+
|
| 214 |
+
def forward(self, x: torch.Tensor) -> tuple:
|
| 215 |
+
"""
|
| 216 |
+
Forward pass through the MoE layer.
|
| 217 |
+
|
| 218 |
+
Args:
|
| 219 |
+
x: Input tensor of shape [batch_size, seq_len, hidden_size] or [tokens, hidden_size]
|
| 220 |
+
|
| 221 |
+
Returns:
|
| 222 |
+
Tuple of (output, expert_weights) where:
|
| 223 |
+
- output: Tensor of same shape as input
|
| 224 |
+
- expert_weights: Expert weights for each token [tokens, top_k]
|
| 225 |
+
"""
|
| 226 |
+
# Get MoE parameters from the wrapped modules
|
| 227 |
+
moe_top_k = getattr(self.router, "top_k", 4)
|
| 228 |
+
moe_num_experts = getattr(self.experts, "num_experts", 128)
|
| 229 |
+
moe_normalize_expert_weights = getattr(
|
| 230 |
+
self.experts, "normalize_expert_weights", None
|
| 231 |
+
)
|
| 232 |
+
|
| 233 |
+
# Detect activation type
|
| 234 |
+
if hasattr(self.experts, "alpha") and hasattr(self.experts, "limit"):
|
| 235 |
+
activation = "swigluoai"
|
| 236 |
+
alpha = self.experts.alpha
|
| 237 |
+
limit = self.experts.limit
|
| 238 |
+
else:
|
| 239 |
+
activation = getattr(self.experts, "activation", "silu")
|
| 240 |
+
alpha = 1.702
|
| 241 |
+
limit = 7.0
|
| 242 |
+
|
| 243 |
+
# Get weight tensors
|
| 244 |
+
if hasattr(self.experts, "gate_up_proj"):
|
| 245 |
+
w1 = self.experts.gate_up_proj
|
| 246 |
+
is_interleaved = True # GptOss uses interleaved layout
|
| 247 |
+
elif hasattr(self.experts, "w1"):
|
| 248 |
+
w1 = self.experts.w1
|
| 249 |
+
w3 = getattr(self.experts, "w3", None)
|
| 250 |
+
if w3 is not None:
|
| 251 |
+
w1 = torch.cat([w1, w3], dim=-1)
|
| 252 |
+
is_interleaved = False
|
| 253 |
+
else:
|
| 254 |
+
raise AttributeError("experts module must have 'gate_up_proj' or 'w1' attribute")
|
| 255 |
+
|
| 256 |
+
if hasattr(self.experts, "down_proj"):
|
| 257 |
+
w2 = self.experts.down_proj
|
| 258 |
+
elif hasattr(self.experts, "w2"):
|
| 259 |
+
w2 = self.experts.w2
|
| 260 |
+
else:
|
| 261 |
+
raise AttributeError("experts module must have 'down_proj' or 'w2' attribute")
|
| 262 |
+
|
| 263 |
+
# Get optional bias tensors
|
| 264 |
+
w1_bias = getattr(self.experts, "gate_up_proj_bias", None)
|
| 265 |
+
w2_bias = getattr(self.experts, "down_proj_bias", None)
|
| 266 |
+
|
| 267 |
+
# Store original shape
|
| 268 |
+
in_shape = x.size()
|
| 269 |
+
|
| 270 |
+
# Route tokens to experts
|
| 271 |
+
logits, expert_weights, expert_indices = route_tokens_cpu(
|
| 272 |
+
x,
|
| 273 |
+
self.router.weight,
|
| 274 |
+
getattr(self.router, "bias", None),
|
| 275 |
+
moe_top_k,
|
| 276 |
+
moe_num_experts,
|
| 277 |
+
moe_normalize_expert_weights,
|
| 278 |
+
)
|
| 279 |
+
|
| 280 |
+
# Reshape input for fused MoE
|
| 281 |
+
x_flat = x.view(-1, x.shape[-1])
|
| 282 |
+
|
| 283 |
+
# Call CPU fused MoE
|
| 284 |
+
output = cpu_fused_moe(
|
| 285 |
+
hidden_states=x_flat,
|
| 286 |
+
w1=w1,
|
| 287 |
+
w2=w2,
|
| 288 |
+
topk_weights=expert_weights,
|
| 289 |
+
topk_ids=expert_indices,
|
| 290 |
+
w1_bias=w1_bias,
|
| 291 |
+
w2_bias=w2_bias,
|
| 292 |
+
activation=activation,
|
| 293 |
+
alpha=alpha,
|
| 294 |
+
limit=limit,
|
| 295 |
+
is_interleaved=is_interleaved,
|
| 296 |
+
)
|
| 297 |
+
|
| 298 |
+
# Restore original shape
|
| 299 |
+
output = output.view(in_shape)
|
| 300 |
+
|
| 301 |
+
return output, expert_weights
|
| 302 |
+
|
| 303 |
+
|
| 304 |
+
# Export classes and functions
|
| 305 |
+
__all__ = [
|
| 306 |
+
"MegaBlocksMoeMLP",
|
| 307 |
+
"cpu_fused_moe",
|
| 308 |
+
"route_tokens_cpu",
|
| 309 |
+
"swigluoai_activation",
|
| 310 |
+
"silu_and_mul_activation",
|
| 311 |
+
]
|
build/torch214-cxx11-xpu20261-x86_64-linux/cpu_moe_cpp.py
ADDED
|
@@ -0,0 +1,265 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
# MegaBlocks C++ Optimized CPU MoE
|
| 3 |
+
|
| 4 |
+
"""
|
| 5 |
+
C++ accelerated MoE with brgemm optimization for Intel AMX.
|
| 6 |
+
Direct replacement for cpu_fused_moe.MegaBlocksMoeMLP with better performance.
|
| 7 |
+
"""
|
| 8 |
+
|
| 9 |
+
import torch
|
| 10 |
+
from typing import Optional
|
| 11 |
+
from .cpu_fused_moe import route_tokens_cpu
|
| 12 |
+
from ._ops import ops
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
def _to_local_tensor(tensor: Optional[torch.Tensor]) -> Optional[torch.Tensor]:
|
| 16 |
+
"""Convert DTensor to local torch.Tensor if needed for custom ops compatibility."""
|
| 17 |
+
if tensor is None:
|
| 18 |
+
return None
|
| 19 |
+
# Check if it's a DTensor by looking for the to_local() method
|
| 20 |
+
if hasattr(tensor, "to_local"):
|
| 21 |
+
return tensor.to_local()
|
| 22 |
+
return tensor
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
def fused_moe_cpp(
|
| 26 |
+
hidden_states: torch.Tensor,
|
| 27 |
+
w1: torch.Tensor,
|
| 28 |
+
w2: torch.Tensor,
|
| 29 |
+
topk_weights: torch.Tensor,
|
| 30 |
+
topk_ids: torch.Tensor,
|
| 31 |
+
inplace: bool = False,
|
| 32 |
+
use_int8_w8a8: bool = False,
|
| 33 |
+
use_fp8_w8a16: bool = False,
|
| 34 |
+
use_mxfp4: bool = False,
|
| 35 |
+
w1_scale: Optional[torch.Tensor] = None,
|
| 36 |
+
w2_scale: Optional[torch.Tensor] = None,
|
| 37 |
+
block_size: Optional[list] = None,
|
| 38 |
+
a1_scale: Optional[torch.Tensor] = None,
|
| 39 |
+
a2_scale: Optional[torch.Tensor] = None,
|
| 40 |
+
w1_bias: Optional[torch.Tensor] = None,
|
| 41 |
+
w2_bias: Optional[torch.Tensor] = None,
|
| 42 |
+
alpha: Optional[float] = None,
|
| 43 |
+
limit: Optional[float] = None,
|
| 44 |
+
is_vnni: bool = False,
|
| 45 |
+
) -> torch.Tensor:
|
| 46 |
+
"""
|
| 47 |
+
C++ Fused MoE with brgemm optimization (sglang compatible interface).
|
| 48 |
+
|
| 49 |
+
Uses at::native::cpublas::brgemm for efficient batch GEMM on Intel CPUs.
|
| 50 |
+
Supports both silu_and_mul (standard SwiGLU) and swigluoai (GptOss) activations.
|
| 51 |
+
|
| 52 |
+
Args:
|
| 53 |
+
hidden_states: Input tensor [M, K]
|
| 54 |
+
w1: Gate and up projections [E, 2N, K]
|
| 55 |
+
w2: Down projection [E, K, N]
|
| 56 |
+
topk_weights: Expert weights [M, topk]
|
| 57 |
+
topk_ids: Expert indices [M, topk]
|
| 58 |
+
inplace: Whether to use hidden_states as output
|
| 59 |
+
use_int8_w8a8: Use int8 quantization
|
| 60 |
+
use_fp8_w8a16: Use fp8 quantization
|
| 61 |
+
use_mxfp4: Use mxfp4 quantization
|
| 62 |
+
w1_scale, w2_scale: Quantization scales
|
| 63 |
+
block_size: Block size for fp8
|
| 64 |
+
a1_scale, a2_scale: Activation scales
|
| 65 |
+
w1_bias, w2_bias: Optional biases
|
| 66 |
+
alpha: swigluoai alpha parameter (set to enable swiglu)
|
| 67 |
+
limit: swigluoai limit parameter (set to enable swiglu)
|
| 68 |
+
is_vnni: Whether w1/w2 are already in VNNI packed format
|
| 69 |
+
"""
|
| 70 |
+
# MXFP4/FP8 kernels only support bf16, convert if needed
|
| 71 |
+
orig_dtype = hidden_states.dtype
|
| 72 |
+
need_convert = ((use_mxfp4 or use_fp8_w8a16) and orig_dtype != torch.bfloat16) or orig_dtype == torch.float32
|
| 73 |
+
if need_convert:
|
| 74 |
+
hidden_states = hidden_states.to(torch.bfloat16)
|
| 75 |
+
|
| 76 |
+
# bias must match hidden_states dtype
|
| 77 |
+
if w1_bias is not None:
|
| 78 |
+
w1_bias = w1_bias.to(hidden_states.dtype)
|
| 79 |
+
if w2_bias is not None:
|
| 80 |
+
w2_bias = w2_bias.to(hidden_states.dtype)
|
| 81 |
+
|
| 82 |
+
# Convert DTensor to local tensor for custom ops compatibility (TP mode)
|
| 83 |
+
hidden_states = _to_local_tensor(hidden_states)
|
| 84 |
+
w1 = _to_local_tensor(w1)
|
| 85 |
+
w2 = _to_local_tensor(w2)
|
| 86 |
+
topk_weights = _to_local_tensor(topk_weights)
|
| 87 |
+
topk_ids = _to_local_tensor(topk_ids)
|
| 88 |
+
w1_scale = _to_local_tensor(w1_scale)
|
| 89 |
+
w2_scale = _to_local_tensor(w2_scale)
|
| 90 |
+
a1_scale = _to_local_tensor(a1_scale)
|
| 91 |
+
a2_scale = _to_local_tensor(a2_scale)
|
| 92 |
+
w1_bias = _to_local_tensor(w1_bias)
|
| 93 |
+
w2_bias = _to_local_tensor(w2_bias)
|
| 94 |
+
|
| 95 |
+
output = ops.fused_experts(
|
| 96 |
+
hidden_states, w1, w2, topk_weights, topk_ids,
|
| 97 |
+
inplace, use_int8_w8a8, use_fp8_w8a16, use_mxfp4,
|
| 98 |
+
w1_scale, w2_scale, block_size, a1_scale, a2_scale,
|
| 99 |
+
w1_bias, w2_bias, alpha, limit, is_vnni
|
| 100 |
+
)
|
| 101 |
+
|
| 102 |
+
# Convert back to original dtype if needed
|
| 103 |
+
if need_convert:
|
| 104 |
+
output = output.to(orig_dtype)
|
| 105 |
+
return output
|
| 106 |
+
|
| 107 |
+
|
| 108 |
+
class CPUMegaBlocksMoeMLP(torch.nn.Module):
|
| 109 |
+
"""
|
| 110 |
+
C++ optimized MoE MLP using brgemm.
|
| 111 |
+
Drop-in replacement for cpu_fused_moe.MegaBlocksMoeMLP with better performance.
|
| 112 |
+
|
| 113 |
+
Usage in transformers:
|
| 114 |
+
# Will be used via @use_kernel_forward_from_hub decorator
|
| 115 |
+
"""
|
| 116 |
+
can_torch_compile: bool = True
|
| 117 |
+
|
| 118 |
+
def forward(self, x: torch.Tensor) -> tuple:
|
| 119 |
+
"""
|
| 120 |
+
Forward pass through the MoE layer using C++ kernel.
|
| 121 |
+
|
| 122 |
+
Args:
|
| 123 |
+
x: Input tensor [batch_size, seq_len, hidden_size]
|
| 124 |
+
|
| 125 |
+
Returns:
|
| 126 |
+
Tuple of (output, expert_weights)
|
| 127 |
+
"""
|
| 128 |
+
# Optimization for GPT-OSS model
|
| 129 |
+
if getattr(self, "use_mxfp4", None) is None:
|
| 130 |
+
self.use_mxfp4 = False
|
| 131 |
+
|
| 132 |
+
w1_scale = None
|
| 133 |
+
w2_scale = None
|
| 134 |
+
|
| 135 |
+
if (
|
| 136 |
+
not getattr(self, "packed_scales", False)
|
| 137 |
+
and hasattr(self.experts, "gate_up_proj")
|
| 138 |
+
and getattr(self.experts, "gate_up_proj_precision_config", None) is not None
|
| 139 |
+
):
|
| 140 |
+
# convert scales
|
| 141 |
+
data_1 = ops.convert_scale_packed(self.experts.gate_up_proj_precision_config.weight_scale.data.transpose(-1, -2).contiguous())
|
| 142 |
+
data_2 = ops.convert_scale_packed(self.experts.down_proj_precision_config.weight_scale.data.transpose(-1, -2).contiguous())
|
| 143 |
+
self.experts.gate_up_proj_precision_config.weight_scale.storage.data = data_1
|
| 144 |
+
self.experts.down_proj_precision_config.weight_scale.storage.data = data_2
|
| 145 |
+
self.packed_scales = True
|
| 146 |
+
self.use_mxfp4 = True
|
| 147 |
+
|
| 148 |
+
if not getattr(self, "packed_weight", False) and hasattr(
|
| 149 |
+
self.experts, "gate_up_proj"
|
| 150 |
+
):
|
| 151 |
+
# convert weights
|
| 152 |
+
data_1 = self.experts.gate_up_proj.data.transpose(-1, -2).contiguous()
|
| 153 |
+
data_2 = self.experts.down_proj.data.transpose(-1, -2).contiguous()
|
| 154 |
+
if self.use_mxfp4:
|
| 155 |
+
self.experts.gate_up_proj.storage.data = ops.convert_weight_packed(data_1)
|
| 156 |
+
self.experts.down_proj.storage.data = ops.convert_weight_packed(data_2)
|
| 157 |
+
else:
|
| 158 |
+
# convert_weight_packed only supports bfloat16, float16, int8, fp8_e4m3 or uint8(mxfp4 or int4).
|
| 159 |
+
data_1 = data_1.to(torch.bfloat16) if data_1.dtype == torch.float32 else data_1
|
| 160 |
+
data_2 = data_2.to(torch.bfloat16) if data_2.dtype == torch.float32 else data_2
|
| 161 |
+
self.experts.gate_up_proj.data = ops.convert_weight_packed(data_1)
|
| 162 |
+
self.experts.down_proj.data = ops.convert_weight_packed(data_2)
|
| 163 |
+
|
| 164 |
+
# C++ kernel does not support float32.
|
| 165 |
+
dtype = torch.bfloat16 if x.dtype == torch.float32 else x.dtype
|
| 166 |
+
if getattr(self.experts, "gate_up_proj_bias", None) is not None:
|
| 167 |
+
self.experts.gate_up_proj_bias.data = self.experts.gate_up_proj_bias.data.to(dtype)
|
| 168 |
+
if getattr(self.experts, "down_proj_bias", None) is not None:
|
| 169 |
+
self.experts.down_proj_bias.data = self.experts.down_proj_bias.data.to(dtype)
|
| 170 |
+
|
| 171 |
+
self.packed_weight = True
|
| 172 |
+
|
| 173 |
+
# Get MoE parameters
|
| 174 |
+
moe_top_k = getattr(self.router, "top_k", 4)
|
| 175 |
+
moe_num_experts = getattr(self.experts, "num_experts", 128)
|
| 176 |
+
moe_normalize_expert_weights = getattr(self.experts, "normalize_expert_weights", None)
|
| 177 |
+
|
| 178 |
+
# Detect activation type
|
| 179 |
+
if hasattr(self.experts, "alpha") and hasattr(self.experts, "limit"):
|
| 180 |
+
activation = "swigluoai"
|
| 181 |
+
alpha = self.experts.alpha
|
| 182 |
+
limit = self.experts.limit
|
| 183 |
+
else:
|
| 184 |
+
activation = getattr(self.experts, "activation", "silu")
|
| 185 |
+
alpha = 1.702
|
| 186 |
+
limit = 7.0
|
| 187 |
+
|
| 188 |
+
# Get weight tensors
|
| 189 |
+
if hasattr(self.experts, "gate_up_proj"):
|
| 190 |
+
w1 = self.experts.gate_up_proj
|
| 191 |
+
elif hasattr(self.experts, "w1"):
|
| 192 |
+
w1 = self.experts.w1
|
| 193 |
+
w3 = getattr(self.experts, "w3", None)
|
| 194 |
+
if w3 is not None:
|
| 195 |
+
w1 = torch.cat([w1, w3], dim=-1)
|
| 196 |
+
else:
|
| 197 |
+
raise AttributeError("experts module must have 'gate_up_proj' or 'w1' attribute")
|
| 198 |
+
|
| 199 |
+
if hasattr(self.experts, "down_proj"):
|
| 200 |
+
w2 = self.experts.down_proj
|
| 201 |
+
elif hasattr(self.experts, "w2"):
|
| 202 |
+
w2 = self.experts.w2
|
| 203 |
+
else:
|
| 204 |
+
raise AttributeError("experts module must have 'down_proj' or 'w2' attribute")
|
| 205 |
+
|
| 206 |
+
# Get optional bias tensors
|
| 207 |
+
w1_bias = getattr(self.experts, "gate_up_proj_bias", None)
|
| 208 |
+
w2_bias = getattr(self.experts, "down_proj_bias", None)
|
| 209 |
+
w1_bias = w1_bias if w1_bias is None else w1_bias.data
|
| 210 |
+
w2_bias = w2_bias if w2_bias is None else w2_bias.data
|
| 211 |
+
|
| 212 |
+
if self.use_mxfp4:
|
| 213 |
+
w1_scale = self.experts.gate_up_proj_precision_config.weight_scale.data
|
| 214 |
+
w2_scale = self.experts.down_proj_precision_config.weight_scale.data
|
| 215 |
+
|
| 216 |
+
# Store original shape
|
| 217 |
+
in_shape = x.size()
|
| 218 |
+
|
| 219 |
+
# Route tokens to experts (Python implementation is fast enough)
|
| 220 |
+
logits, expert_weights, expert_indices = route_tokens_cpu(
|
| 221 |
+
x,
|
| 222 |
+
self.router.weight,
|
| 223 |
+
getattr(self.router, "bias", None),
|
| 224 |
+
moe_top_k,
|
| 225 |
+
moe_num_experts,
|
| 226 |
+
moe_normalize_expert_weights,
|
| 227 |
+
)
|
| 228 |
+
|
| 229 |
+
# Flatten input
|
| 230 |
+
x_flat = x.view(-1, x.shape[-1])
|
| 231 |
+
|
| 232 |
+
# Determine alpha/limit for swiglu activation
|
| 233 |
+
use_alpha = alpha if activation == "swigluoai" else None
|
| 234 |
+
use_limit = limit if activation == "swigluoai" else None
|
| 235 |
+
|
| 236 |
+
# Call C++ optimized kernel
|
| 237 |
+
output = fused_moe_cpp(
|
| 238 |
+
hidden_states=x_flat,
|
| 239 |
+
w1=w1.data,
|
| 240 |
+
w2=w2.data,
|
| 241 |
+
topk_weights=expert_weights,
|
| 242 |
+
topk_ids=expert_indices.to(torch.int32),
|
| 243 |
+
inplace=False,
|
| 244 |
+
use_int8_w8a8=False,
|
| 245 |
+
use_fp8_w8a16=False,
|
| 246 |
+
use_mxfp4=self.use_mxfp4,
|
| 247 |
+
w1_scale=w1_scale,
|
| 248 |
+
w2_scale=w2_scale,
|
| 249 |
+
block_size=None,
|
| 250 |
+
a1_scale=None,
|
| 251 |
+
a2_scale=None,
|
| 252 |
+
w1_bias=w1_bias,
|
| 253 |
+
w2_bias=w2_bias,
|
| 254 |
+
alpha=use_alpha,
|
| 255 |
+
limit=use_limit,
|
| 256 |
+
is_vnni=getattr(self, "packed_weight", False),
|
| 257 |
+
)
|
| 258 |
+
|
| 259 |
+
# Restore original shape
|
| 260 |
+
output = output.view(in_shape)
|
| 261 |
+
|
| 262 |
+
return output, expert_weights
|
| 263 |
+
|
| 264 |
+
|
| 265 |
+
__all__ = ["fused_moe_cpp", "MegaBlocksMoeMLP"]
|
build/torch214-cxx11-xpu20261-x86_64-linux/grouped_gemm/__init__.py
ADDED
|
@@ -0,0 +1,2 @@
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from . import ops
|
| 2 |
+
from . import backend
|
build/torch214-cxx11-xpu20261-x86_64-linux/grouped_gemm/backend.py
ADDED
|
@@ -0,0 +1,33 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# NOTE: Torch needs to be imported before the custom
|
| 2 |
+
# extensions. Otherwise libc10.so cannot be found.
|
| 3 |
+
import torch
|
| 4 |
+
|
| 5 |
+
# On ROCm there is no CUTLASS grouped GEMM; dispatch to the vendored AITER
|
| 6 |
+
# Triton kernels instead. On CUDA we use the compiled CUTLASS `gmm` op.
|
| 7 |
+
_IS_ROCM = torch.version.hip is not None
|
| 8 |
+
|
| 9 |
+
if _IS_ROCM:
|
| 10 |
+
from .._grouped_gemm_triton import adapter as backend
|
| 11 |
+
else:
|
| 12 |
+
# We import the backend operations from the megablocks package as
|
| 13 |
+
# grouped_gemm is vendored in megablocks in this repository.
|
| 14 |
+
from .._ops import ops as backend # type: ignore
|
| 15 |
+
|
| 16 |
+
def _allocate_output(a, b, batch_sizes, trans_a, trans_b):
|
| 17 |
+
assert not (trans_a and trans_b)
|
| 18 |
+
assert batch_sizes.ndim == 1, "Expected 1d tensor for batch_sizes"
|
| 19 |
+
assert a.ndim == 2, "Expected 2d tensor for 'a'"
|
| 20 |
+
assert b.ndim == (2 if trans_a else 3)
|
| 21 |
+
|
| 22 |
+
shape = (
|
| 23 |
+
(batch_sizes.shape[0], a.shape[1], b.shape[1])
|
| 24 |
+
if trans_a else
|
| 25 |
+
(a.shape[0], (b.shape[1] if trans_b else b.shape[2]))
|
| 26 |
+
)
|
| 27 |
+
return torch.empty(*shape, device=a.device, dtype=a.dtype)
|
| 28 |
+
|
| 29 |
+
def gmm(a, b, batch_sizes, trans_a=False, trans_b=False, c=None):
|
| 30 |
+
if c is None:
|
| 31 |
+
c = _allocate_output(a, b, batch_sizes, trans_a, trans_b)
|
| 32 |
+
backend.gmm(a, b, c, batch_sizes, trans_a, trans_b)
|
| 33 |
+
return c
|
build/torch214-cxx11-xpu20261-x86_64-linux/grouped_gemm/ops.py
ADDED
|
@@ -0,0 +1,33 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from . import backend
|
| 2 |
+
import torch
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
class GroupedGemm(torch.autograd.Function):
|
| 6 |
+
|
| 7 |
+
@staticmethod
|
| 8 |
+
def forward(ctx, a, b, batch_sizes, trans_b):
|
| 9 |
+
ctx.save_for_backward(a, b, batch_sizes)
|
| 10 |
+
ctx.trans_b = trans_b
|
| 11 |
+
return backend.gmm(a, b, batch_sizes, trans_a=False, trans_b=trans_b)
|
| 12 |
+
|
| 13 |
+
@staticmethod
|
| 14 |
+
def backward(ctx, grad):
|
| 15 |
+
grad = grad.contiguous()
|
| 16 |
+
a, b, batch_sizes = ctx.saved_tensors
|
| 17 |
+
trans_b = ctx.trans_b
|
| 18 |
+
|
| 19 |
+
agrad = None
|
| 20 |
+
if ctx.needs_input_grad[0]:
|
| 21 |
+
agrad = backend.gmm(
|
| 22 |
+
grad, b, batch_sizes, trans_a=False, trans_b=not trans_b)
|
| 23 |
+
|
| 24 |
+
bgrad = None
|
| 25 |
+
if ctx.needs_input_grad[1]:
|
| 26 |
+
lhs, rhs = (grad, a) if trans_b else (a, grad)
|
| 27 |
+
bgrad = backend.gmm(
|
| 28 |
+
lhs, rhs, batch_sizes, trans_a=True, trans_b=False)
|
| 29 |
+
return agrad, bgrad, None, None
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def gmm(a, b, batch_sizes, trans_b=False):
|
| 33 |
+
return GroupedGemm.apply(a, b, batch_sizes, trans_b)
|
build/torch214-cxx11-xpu20261-x86_64-linux/grouped_gemm_util.py
ADDED
|
@@ -0,0 +1,31 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2024 Databricks
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
import warnings
|
| 4 |
+
|
| 5 |
+
_grouped_gemm_is_available: bool = False
|
| 6 |
+
try:
|
| 7 |
+
# import grouped_gemm
|
| 8 |
+
pass
|
| 9 |
+
_grouped_gemm_is_available = True
|
| 10 |
+
except ImportError as error:
|
| 11 |
+
warnings.warn('Grouped GEMM not available.')
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def grouped_gemm_is_available():
|
| 15 |
+
return _grouped_gemm_is_available
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
def assert_grouped_gemm_is_available():
|
| 19 |
+
msg = (
|
| 20 |
+
'Grouped GEMM not available. Please run '
|
| 21 |
+
'`pip install git+https://github.com/tgale96/grouped_gemm@main`.',
|
| 22 |
+
)
|
| 23 |
+
assert _grouped_gemm_is_available, msg
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
# backend = grouped_gemm.backend if grouped_gemm_is_available() else None
|
| 27 |
+
# ops = grouped_gemm.ops if grouped_gemm_is_available() else None
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
from .grouped_gemm import backend as ops
|
| 31 |
+
from .grouped_gemm import ops as backend
|