kernels-bot commited on
Commit
2c977e2
·
verified ·
1 Parent(s): 65478c1

Uploaded using `kernel-builder`.

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. build/torch212-cxx11-xpu20253-x86_64-linux/{_megablocks_xpu_c4d0dc5.abi3.so → _megablocks_xpu_addf474.abi3.so} +1 -1
  2. build/torch212-cxx11-xpu20253-x86_64-linux/_ops.py +3 -3
  3. build/torch212-cxx11-xpu20253-x86_64-linux/megablocks/__init__.py +0 -26
  4. build/torch212-cxx11-xpu20253-x86_64-linux/metadata.json +6 -6
  5. build/torch212-cxx11-xpu20253-x86_64-linux/metadata.json.sigstore +1 -1
  6. build/torch213-cxx11-xpu20260-x86_64-linux/{_megablocks_xpu_c4d0dc5.abi3.so → _megablocks_xpu_addf474.abi3.so} +2 -2
  7. build/torch213-cxx11-xpu20260-x86_64-linux/_ops.py +3 -3
  8. build/torch213-cxx11-xpu20260-x86_64-linux/megablocks/__init__.py +0 -26
  9. build/torch213-cxx11-xpu20260-x86_64-linux/metadata.json +6 -6
  10. build/torch213-cxx11-xpu20260-x86_64-linux/metadata.json.sigstore +1 -1
  11. build/torch214-cxx11-xpu20261-x86_64-linux/__init__.py +205 -0
  12. build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/__init__.py +0 -0
  13. build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/_triton_kernels/__init__.py +0 -0
  14. build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/_triton_kernels/gmm.py +574 -0
  15. build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/adapter.py +53 -0
  16. build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/configs.py +5 -0
  17. build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/gmm.py +567 -0
  18. build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/utils/__init__.py +0 -0
  19. build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/utils/_triton/__init__.py +0 -0
  20. build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/utils/_triton/arch_info.py +46 -0
  21. build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/utils/_triton/pid_preprocessing.py +100 -0
  22. build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/utils/gmm_common.py +752 -0
  23. build/torch214-cxx11-xpu20261-x86_64-linux/_grouped_gemm_triton/utils/logger.py +47 -0
  24. build/torch214-cxx11-xpu20261-x86_64-linux/_layers/__init__.py +10 -0
  25. build/torch214-cxx11-xpu20261-x86_64-linux/_layers/activation_fn.py +33 -0
  26. build/torch214-cxx11-xpu20261-x86_64-linux/_layers/all_to_all.py +54 -0
  27. build/torch214-cxx11-xpu20261-x86_64-linux/_layers/arguments.py +101 -0
  28. build/torch214-cxx11-xpu20261-x86_64-linux/_layers/common.py +26 -0
  29. build/torch214-cxx11-xpu20261-x86_64-linux/_layers/dmlp_registry.py +42 -0
  30. build/torch214-cxx11-xpu20261-x86_64-linux/_layers/dmoe.py +337 -0
  31. build/torch214-cxx11-xpu20261-x86_64-linux/_layers/gelu.py +52 -0
  32. build/torch214-cxx11-xpu20261-x86_64-linux/_layers/glu.py +244 -0
  33. build/torch214-cxx11-xpu20261-x86_64-linux/_layers/memory_test.py +103 -0
  34. build/torch214-cxx11-xpu20261-x86_64-linux/_layers/mlp.py +587 -0
  35. build/torch214-cxx11-xpu20261-x86_64-linux/_layers/moe.py +507 -0
  36. build/torch214-cxx11-xpu20261-x86_64-linux/_layers/mpu.py +94 -0
  37. build/torch214-cxx11-xpu20261-x86_64-linux/_layers/router.py +116 -0
  38. build/torch214-cxx11-xpu20261-x86_64-linux/_layers/sharedexpert_registry.py +32 -0
  39. build/torch214-cxx11-xpu20261-x86_64-linux/_megablocks_xpu_addf474.abi3.so +3 -0
  40. build/torch214-cxx11-xpu20261-x86_64-linux/_ops.py +9 -0
  41. build/torch214-cxx11-xpu20261-x86_64-linux/_version.py +6 -0
  42. build/torch214-cxx11-xpu20261-x86_64-linux/backend/__init__.py +2 -0
  43. build/torch214-cxx11-xpu20261-x86_64-linux/backend/kernels.py +557 -0
  44. build/torch214-cxx11-xpu20261-x86_64-linux/benchmark_util.py +35 -0
  45. build/torch214-cxx11-xpu20261-x86_64-linux/cpu_fused_moe.py +311 -0
  46. build/torch214-cxx11-xpu20261-x86_64-linux/cpu_moe_cpp.py +265 -0
  47. build/torch214-cxx11-xpu20261-x86_64-linux/grouped_gemm/__init__.py +2 -0
  48. build/torch214-cxx11-xpu20261-x86_64-linux/grouped_gemm/backend.py +33 -0
  49. build/torch214-cxx11-xpu20261-x86_64-linux/grouped_gemm/ops.py +33 -0
  50. 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:868eb695bd1e4c252332321da1ca3002e7322764498808ed957cf8e869b9ea5a
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 _megablocks_xpu_c4d0dc5
3
- ops = torch.ops._megablocks_xpu_c4d0dc5
4
 
5
  def add_op_namespace_prefix(op_name: str):
6
  """
7
  Prefix op by namespace.
8
  """
9
- return f"_megablocks_xpu_c4d0dc5::{op_name}"
 
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": "_megablocks_xpu_c4d0dc5",
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
- "_megablocks_xpu_c4d0dc5.abi3.so": "ho62lb0eTCUjMjIdocowAucyJ2RJiAjtlXz46Gm56lo=",
42
- "_ops.py": "KZ/xutl7d4r1cvlnq+lt8PGsR1a57OQznlwnP3eVkKI=",
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
- "sha": "81580bb92577f2f7228661ca2a221fc052375709",
100
  "dirty": false
101
  },
102
  "kernel": {
103
- "sha": "c4d0dc56661badb9cc47bde647e08fe3b4bf7458",
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:e9e67436714e7f1ac7513960d18a05eae815c650050754fe667d1ddef6cf2210
3
- size 3931504
 
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 _megablocks_xpu_c4d0dc5
3
- ops = torch.ops._megablocks_xpu_c4d0dc5
4
 
5
  def add_op_namespace_prefix(op_name: str):
6
  """
7
  Prefix op by namespace.
8
  """
9
- return f"_megablocks_xpu_c4d0dc5::{op_name}"
 
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": "_megablocks_xpu_c4d0dc5",
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
- "_megablocks_xpu_c4d0dc5.abi3.so": "6eZ0NnFOfxrHUTlg0YoF6ugVxlAFB1T+Zn0d3vbPIhA=",
42
- "_ops.py": "KZ/xutl7d4r1cvlnq+lt8PGsR1a57OQznlwnP3eVkKI=",
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
- "sha": "81580bb92577f2f7228661ca2a221fc052375709",
100
  "dirty": false
101
  },
102
  "kernel": {
103
- "sha": "c4d0dc56661badb9cc47bde647e08fe3b4bf7458",
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