tile-kernels 1.0.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (135) hide show
  1. tile_kernels-1.0.0/.editorconfig +58 -0
  2. tile_kernels-1.0.0/.gitignore +20 -0
  3. tile_kernels-1.0.0/LICENSE +21 -0
  4. tile_kernels-1.0.0/PKG-INFO +116 -0
  5. tile_kernels-1.0.0/README.md +88 -0
  6. tile_kernels-1.0.0/pyproject.toml +58 -0
  7. tile_kernels-1.0.0/setup.cfg +4 -0
  8. tile_kernels-1.0.0/tests/__init__.py +0 -0
  9. tile_kernels-1.0.0/tests/conftest.py +10 -0
  10. tile_kernels-1.0.0/tests/engram/test_engram_fused_weight.py +56 -0
  11. tile_kernels-1.0.0/tests/engram/test_engram_gate_bwd.py +112 -0
  12. tile_kernels-1.0.0/tests/engram/test_engram_gate_fwd.py +108 -0
  13. tile_kernels-1.0.0/tests/engram/test_engram_grad_w_reduce.py +76 -0
  14. tile_kernels-1.0.0/tests/engram/test_engram_hash.py +71 -0
  15. tile_kernels-1.0.0/tests/mhc/test_expand.py +45 -0
  16. tile_kernels-1.0.0/tests/mhc/test_head_compute_mix.py +54 -0
  17. tile_kernels-1.0.0/tests/mhc/test_multilayer_recompute.py +157 -0
  18. tile_kernels-1.0.0/tests/mhc/test_norm_fn.py +130 -0
  19. tile_kernels-1.0.0/tests/mhc/test_post.py +72 -0
  20. tile_kernels-1.0.0/tests/mhc/test_pre_apply_mix.py +52 -0
  21. tile_kernels-1.0.0/tests/mhc/test_pre_big_fuse.py +138 -0
  22. tile_kernels-1.0.0/tests/mhc/test_pre_split_mixes.py +103 -0
  23. tile_kernels-1.0.0/tests/mhc/test_sinkhorn.py +43 -0
  24. tile_kernels-1.0.0/tests/moe/test_aux_fi.py +65 -0
  25. tile_kernels-1.0.0/tests/moe/test_expand_to_fused.py +150 -0
  26. tile_kernels-1.0.0/tests/moe/test_get_fused_mapping.py +103 -0
  27. tile_kernels-1.0.0/tests/moe/test_group_count.py +55 -0
  28. tile_kernels-1.0.0/tests/moe/test_inplace_unique_group_indices.py +79 -0
  29. tile_kernels-1.0.0/tests/moe/test_mask_indices_by_tp.py +70 -0
  30. tile_kernels-1.0.0/tests/moe/test_normalize_weight.py +57 -0
  31. tile_kernels-1.0.0/tests/moe/test_reduce_fused.py +112 -0
  32. tile_kernels-1.0.0/tests/moe/test_top2_sum_gate.py +359 -0
  33. tile_kernels-1.0.0/tests/moe/test_topk_gate.py +75 -0
  34. tile_kernels-1.0.0/tests/moe/test_topk_sum_and_topk_idx.py +86 -0
  35. tile_kernels-1.0.0/tests/pytest_benchmark_plugin.py +477 -0
  36. tile_kernels-1.0.0/tests/pytest_random_plugin.py +18 -0
  37. tile_kernels-1.0.0/tests/quant/test_cast_back.py +157 -0
  38. tile_kernels-1.0.0/tests/quant/test_cast_back_e5m6.py +120 -0
  39. tile_kernels-1.0.0/tests/quant/test_per_block_cast.py +110 -0
  40. tile_kernels-1.0.0/tests/quant/test_per_block_cast_lossless.py +114 -0
  41. tile_kernels-1.0.0/tests/quant/test_per_channel_cast.py +75 -0
  42. tile_kernels-1.0.0/tests/quant/test_per_channel_cast_and_transpose.py +77 -0
  43. tile_kernels-1.0.0/tests/quant/test_per_channel_cast_fused.py +104 -0
  44. tile_kernels-1.0.0/tests/quant/test_per_token_cast.py +166 -0
  45. tile_kernels-1.0.0/tests/quant/test_per_token_cast_to_e5m6.py +118 -0
  46. tile_kernels-1.0.0/tests/quant/test_swiglu_backward_and_per_token_cast.py +149 -0
  47. tile_kernels-1.0.0/tests/quant/test_swiglu_forward_and_per_channel_cast_and_transpose.py +114 -0
  48. tile_kernels-1.0.0/tests/quant/test_swiglu_forward_and_per_token_cast.py +179 -0
  49. tile_kernels-1.0.0/tests/transpose/test_transpose.py +117 -0
  50. tile_kernels-1.0.0/tile_kernels/__init__.py +15 -0
  51. tile_kernels-1.0.0/tile_kernels/_version.py +34 -0
  52. tile_kernels-1.0.0/tile_kernels/config.py +29 -0
  53. tile_kernels-1.0.0/tile_kernels/engram/__init__.py +4 -0
  54. tile_kernels-1.0.0/tile_kernels/engram/engram_fused_weight_kernel.py +59 -0
  55. tile_kernels-1.0.0/tile_kernels/engram/engram_gate_kernel.py +570 -0
  56. tile_kernels-1.0.0/tile_kernels/engram/engram_grad_w_reduce_kernel.py +89 -0
  57. tile_kernels-1.0.0/tile_kernels/engram/engram_hash_kernel.py +97 -0
  58. tile_kernels-1.0.0/tile_kernels/mhc/__init__.py +0 -0
  59. tile_kernels-1.0.0/tile_kernels/mhc/expand_kernel.py +60 -0
  60. tile_kernels-1.0.0/tile_kernels/mhc/head_compute_mix_kernel.py +85 -0
  61. tile_kernels-1.0.0/tile_kernels/mhc/multilayer_recompute_kernel.py +196 -0
  62. tile_kernels-1.0.0/tile_kernels/mhc/norm_fn_kernel.py +289 -0
  63. tile_kernels-1.0.0/tile_kernels/mhc/post_kernel.py +221 -0
  64. tile_kernels-1.0.0/tile_kernels/mhc/pre_apply_mix_kernel.py +110 -0
  65. tile_kernels-1.0.0/tile_kernels/mhc/pre_big_fuse_kernel.py +131 -0
  66. tile_kernels-1.0.0/tile_kernels/mhc/pre_split_mixes_kernel.py +163 -0
  67. tile_kernels-1.0.0/tile_kernels/mhc/sinkhorn_kernel.py +165 -0
  68. tile_kernels-1.0.0/tile_kernels/modeling/__init__.py +2 -0
  69. tile_kernels-1.0.0/tile_kernels/modeling/engram/__init__.py +1 -0
  70. tile_kernels-1.0.0/tile_kernels/modeling/engram/engram_gate.py +95 -0
  71. tile_kernels-1.0.0/tile_kernels/modeling/mhc/__init__.py +1 -0
  72. tile_kernels-1.0.0/tile_kernels/modeling/mhc/functional.py +160 -0
  73. tile_kernels-1.0.0/tile_kernels/modeling/mhc/ops/__init__.py +9 -0
  74. tile_kernels-1.0.0/tile_kernels/modeling/mhc/ops/expand.py +34 -0
  75. tile_kernels-1.0.0/tile_kernels/modeling/mhc/ops/head_compute_mix.py +77 -0
  76. tile_kernels-1.0.0/tile_kernels/modeling/mhc/ops/multilayer_recompute.py +3 -0
  77. tile_kernels-1.0.0/tile_kernels/modeling/mhc/ops/norm_fn.py +189 -0
  78. tile_kernels-1.0.0/tile_kernels/modeling/mhc/ops/post.py +35 -0
  79. tile_kernels-1.0.0/tile_kernels/modeling/mhc/ops/pre_apply_mix.py +55 -0
  80. tile_kernels-1.0.0/tile_kernels/modeling/mhc/ops/pre_big_fuse.py +91 -0
  81. tile_kernels-1.0.0/tile_kernels/modeling/mhc/ops/pre_split_mixes.py +122 -0
  82. tile_kernels-1.0.0/tile_kernels/modeling/mhc/ops/sinkhorn.py +32 -0
  83. tile_kernels-1.0.0/tile_kernels/moe/__init__.py +11 -0
  84. tile_kernels-1.0.0/tile_kernels/moe/aux_fi_kernel.py +72 -0
  85. tile_kernels-1.0.0/tile_kernels/moe/common.py +52 -0
  86. tile_kernels-1.0.0/tile_kernels/moe/expand_to_fused_kernel.py +200 -0
  87. tile_kernels-1.0.0/tile_kernels/moe/get_fused_mapping_kernel.py +250 -0
  88. tile_kernels-1.0.0/tile_kernels/moe/group_count_kernel.py +69 -0
  89. tile_kernels-1.0.0/tile_kernels/moe/inplace_unique_group_indices_kernel.py +69 -0
  90. tile_kernels-1.0.0/tile_kernels/moe/mask_indices_by_tp_kernel.py +73 -0
  91. tile_kernels-1.0.0/tile_kernels/moe/normalize_weight_kernel.py +70 -0
  92. tile_kernels-1.0.0/tile_kernels/moe/reduce_fused_kernel.py +136 -0
  93. tile_kernels-1.0.0/tile_kernels/moe/scoring.py +26 -0
  94. tile_kernels-1.0.0/tile_kernels/moe/top2_sum_gate_kernel.py +424 -0
  95. tile_kernels-1.0.0/tile_kernels/moe/topk_gate_kernel.py +90 -0
  96. tile_kernels-1.0.0/tile_kernels/moe/topk_sum_and_topk_group_idx_kernel.py +102 -0
  97. tile_kernels-1.0.0/tile_kernels/quant/__init__.py +11 -0
  98. tile_kernels-1.0.0/tile_kernels/quant/cast_back_e5m6_kernel.py +176 -0
  99. tile_kernels-1.0.0/tile_kernels/quant/cast_back_kernel.py +143 -0
  100. tile_kernels-1.0.0/tile_kernels/quant/common.py +295 -0
  101. tile_kernels-1.0.0/tile_kernels/quant/per_block_cast_kernel.py +276 -0
  102. tile_kernels-1.0.0/tile_kernels/quant/per_block_cast_lossless_kernel.py +209 -0
  103. tile_kernels-1.0.0/tile_kernels/quant/per_channel_cast_and_transpose_kernel.py +123 -0
  104. tile_kernels-1.0.0/tile_kernels/quant/per_channel_cast_fused_kernel.py +203 -0
  105. tile_kernels-1.0.0/tile_kernels/quant/per_channel_cast_kernel.py +28 -0
  106. tile_kernels-1.0.0/tile_kernels/quant/per_token_cast_kernel.py +302 -0
  107. tile_kernels-1.0.0/tile_kernels/quant/per_token_cast_to_e5m6_kernel.py +215 -0
  108. tile_kernels-1.0.0/tile_kernels/quant/swiglu_backward_and_per_token_cast_kernel.py +237 -0
  109. tile_kernels-1.0.0/tile_kernels/quant/swiglu_forward_and_per_channel_cast_and_transpose_kernel.py +231 -0
  110. tile_kernels-1.0.0/tile_kernels/quant/swiglu_forward_and_per_token_cast_kernel.py +260 -0
  111. tile_kernels-1.0.0/tile_kernels/quant/types.py +3 -0
  112. tile_kernels-1.0.0/tile_kernels/testing/__init__.py +2 -0
  113. tile_kernels-1.0.0/tile_kernels/testing/bench.py +115 -0
  114. tile_kernels-1.0.0/tile_kernels/testing/generator.py +105 -0
  115. tile_kernels-1.0.0/tile_kernels/testing/numeric.py +65 -0
  116. tile_kernels-1.0.0/tile_kernels/testing/quant.py +21 -0
  117. tile_kernels-1.0.0/tile_kernels/torch/__init__.py +10 -0
  118. tile_kernels-1.0.0/tile_kernels/torch/cast.py +252 -0
  119. tile_kernels-1.0.0/tile_kernels/torch/cast_e5m6.py +258 -0
  120. tile_kernels-1.0.0/tile_kernels/torch/engram.py +113 -0
  121. tile_kernels-1.0.0/tile_kernels/torch/expand_to_fused.py +91 -0
  122. tile_kernels-1.0.0/tile_kernels/torch/mhc.py +85 -0
  123. tile_kernels-1.0.0/tile_kernels/torch/moe.py +120 -0
  124. tile_kernels-1.0.0/tile_kernels/torch/per_channel_cast_fused.py +44 -0
  125. tile_kernels-1.0.0/tile_kernels/torch/reduce_fused.py +74 -0
  126. tile_kernels-1.0.0/tile_kernels/torch/swiglu.py +227 -0
  127. tile_kernels-1.0.0/tile_kernels/torch/topk.py +206 -0
  128. tile_kernels-1.0.0/tile_kernels/transpose/__init__.py +1 -0
  129. tile_kernels-1.0.0/tile_kernels/transpose/batched_transpose_kernel.py +119 -0
  130. tile_kernels-1.0.0/tile_kernels/utils.py +10 -0
  131. tile_kernels-1.0.0/tile_kernels.egg-info/PKG-INFO +116 -0
  132. tile_kernels-1.0.0/tile_kernels.egg-info/SOURCES.txt +133 -0
  133. tile_kernels-1.0.0/tile_kernels.egg-info/dependency_links.txt +1 -0
  134. tile_kernels-1.0.0/tile_kernels.egg-info/requires.txt +10 -0
  135. tile_kernels-1.0.0/tile_kernels.egg-info/top_level.txt +1 -0
@@ -0,0 +1,58 @@
1
+ # https://editorconfig.org/
2
+
3
+ root = true
4
+
5
+ [*]
6
+ charset = utf-8
7
+ end_of_line = lf
8
+ indent_style = space
9
+ indent_size = 4
10
+ trim_trailing_whitespace = true
11
+ insert_final_newline = true
12
+
13
+ [*.{py,pyi}]
14
+ indent_size = 4
15
+
16
+ [*.{cpp,hpp,cxx,cc,c,h,cu,cuh}]
17
+ indent_size = 4
18
+
19
+ [*.rs]
20
+ indent_size = 4
21
+
22
+ [*.go]
23
+ indent_style = tab
24
+
25
+ [*.{yaml,yml}]
26
+ indent_size = 2
27
+
28
+ [.clang-{format,tidy}]
29
+ indent_size = 2
30
+
31
+ [Makefile]
32
+ indent_style = tab
33
+
34
+ [*.sh]
35
+ indent_size = 4
36
+
37
+ [*.bat]
38
+ indent_size = 4
39
+ end_of_line = crlf
40
+
41
+ [*.md]
42
+ indent_size = 2
43
+ x-soft-wrap-text = true
44
+
45
+ [*.rst]
46
+ indent_size = 4
47
+ x-soft-wrap-text = true
48
+
49
+ [*.{html,xml,css,scss,js,jsx,ts,tsx,vue}]
50
+ indent_size = 2
51
+
52
+ [**/test{,s,ing}/**/*.txt]
53
+ trim_trailing_whitespace = false
54
+ insert_final_newline = false
55
+
56
+ [**/example{,s}/**/*.txt]
57
+ trim_trailing_whitespace = false
58
+ insert_final_newline = false
@@ -0,0 +1,20 @@
1
+ .DS_Store
2
+ build
3
+ dist
4
+ *.egg-info
5
+ *.pyc
6
+ __pycache__/
7
+
8
+ # PyCharm Settings
9
+ .idea/
10
+ *.iml
11
+ *.iws
12
+
13
+ # Python packaging version file
14
+ _version.py
15
+
16
+ # VSCode Settings
17
+ /.vscode
18
+
19
+ # workspace
20
+ workspace
@@ -0,0 +1,21 @@
1
+ MIT License
2
+
3
+ Copyright (c) 2026 DeepSeek
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
@@ -0,0 +1,116 @@
1
+ Metadata-Version: 2.4
2
+ Name: tile_kernels
3
+ Version: 1.0.0
4
+ Summary: Tilelang-based kernels.
5
+ Author-email: Chenhao Xu <xch@deepseek.com>, Xiangwen Wang <xiangwen.wang@deepseek.com>, Huanqi Cao <caohuanqi@deepseek.com>, Rui Tian <tianr22@deepseek.com>, Weilin Zhao <zhaoweilin@deepseek.com>, Kuai Yu <yukuai@deepseek.com>, Chenggang Zhao <chenggangz@deepseek.com>
6
+ License-Expression: MIT
7
+ Project-URL: Homepage, https://github.com/deepseek-ai/TileKernels
8
+ Classifier: Development Status :: 3 - Alpha
9
+ Classifier: Intended Audience :: Developers
10
+ Classifier: Programming Language :: Python :: 3
11
+ Classifier: Programming Language :: Python :: 3.10
12
+ Classifier: Programming Language :: Python :: 3.11
13
+ Classifier: Programming Language :: Python :: 3.12
14
+ Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
15
+ Requires-Python: >=3.10
16
+ Description-Content-Type: text/markdown
17
+ License-File: LICENSE
18
+ Requires-Dist: torch>=2.10
19
+ Requires-Dist: tilelang>=0.1.9
20
+ Provides-Extra: dev
21
+ Requires-Dist: setuptools; extra == "dev"
22
+ Requires-Dist: wheel; extra == "dev"
23
+ Requires-Dist: setuptools-scm>=8; extra == "dev"
24
+ Requires-Dist: pytest; extra == "dev"
25
+ Requires-Dist: pytest-xdist; extra == "dev"
26
+ Requires-Dist: pytest-repeat; extra == "dev"
27
+ Dynamic: license-file
28
+
29
+ # Tile Kernels
30
+
31
+ Optimized GPU kernels for LLM operations, built with [TileLang](https://github.com/tile-ai/tilelang). TileLang is a domain-specific language for expressing high-performance GPU kernels in Python, featuring easy migration, agile development, and automatic optimization.
32
+
33
+ Most kernels in this project approach the limit of hardware performance regarding the compute intensity and memory bandwidth. Some of them have already been used in internal training and inference scenarios. However, they do not represent best practices and we are actively working on improving the code quality and documentation.
34
+
35
+ ## Features
36
+
37
+ - **Gating** — Top-k expert selection and scoring for Mixture of Experts routing
38
+ - **MoE Routing** — Token-to-expert mapping, fused expansion/reduction and weight normalization
39
+ - **Quantization** — Per-token, per-block, and per-channel FP8/FP4/E5M6 casting with fused SwiGLU+quantization ops
40
+ - **Transpose** — Batched transpose operations
41
+ - **Engram** — Engram gating kernels with fused RMSNorm, forward/backward passes and weight gradient reduction
42
+ - **Manifold HyperConnection** — Hyper-connection kernels including Sinkhorn normalization and mix splitting/application
43
+ - **Modeling** — High-level `torch.autograd.Function` wrappers composing low-level kernels into trainable layers (engram gate, mHC pipeline)
44
+
45
+ ## Requirements
46
+
47
+ - Python 3.10 or higher
48
+ - PyTorch 2.10 or higher
49
+ - TileLang 0.1.9 or higher
50
+ - NVIDIA SM90 or SM100 architecture GPU
51
+ - CUDA Toolkit 13.1 or higher
52
+
53
+ ## Installation
54
+
55
+ ### Install a local development version
56
+
57
+ ```bash
58
+ pip install -e ".[dev]"
59
+ ```
60
+
61
+ ### Install a release version
62
+
63
+ ```bash
64
+ pip install tile-kernels
65
+ ```
66
+
67
+ ## Testing
68
+
69
+ Tests using pytest:
70
+
71
+ ### Test single test file
72
+
73
+ ```bash
74
+ pytest tests/transpose/test_transpose.py -n 4 # Correctness only with 4 workers
75
+ pytest tests/transpose/test_transpose.py --run-benchmark # Correctness + Benchmarking
76
+ ```
77
+
78
+ ### Pressure test
79
+
80
+ ```bash
81
+ TK_FULL_TEST=1 pytest -n 4 --count 2
82
+ ```
83
+
84
+ ## Project Structure
85
+
86
+ ```txt
87
+ tile_kernels/
88
+ ├── moe/ # Mixture of Experts routing related kernels
89
+ ├── quant/ # FP8/FP4/E5M6 quantization
90
+ ├── transpose/ # Batched transpose
91
+ ├── engram/ # Engram gating kernels
92
+ ├── mhc/ # Manifold HyperConnection kernels
93
+ ├── modeling/ # High-level autograd modeling layers (engram, mHC)
94
+ ├── torch/ # PyTorch reference implementations
95
+ └── testing/ # Test and benchmark utilities
96
+ ```
97
+
98
+ ## Acknowledgement
99
+
100
+ This project is built on [TileLang](https://github.com/tile-ai/tilelang). Thanks and respect to the developers!
101
+
102
+ ## License
103
+
104
+ This code repository is released under [the MIT License](LICENSE).
105
+
106
+ ## Citation
107
+
108
+ ```bibtex
109
+ @misc{tilekernels,
110
+ title={TileKernels},
111
+ author={Xiangwen Wang, Chenhao Xu, Huanqi Cao, Rui Tian, Weilin Zhao, Kuai Yu and Chenggang Zhao},
112
+ year={2026},
113
+ publisher = {GitHub},
114
+ howpublished = {\url{https://github.com/deepseek-ai/TileKernels}},
115
+ }
116
+ ```
@@ -0,0 +1,88 @@
1
+ # Tile Kernels
2
+
3
+ Optimized GPU kernels for LLM operations, built with [TileLang](https://github.com/tile-ai/tilelang). TileLang is a domain-specific language for expressing high-performance GPU kernels in Python, featuring easy migration, agile development, and automatic optimization.
4
+
5
+ Most kernels in this project approach the limit of hardware performance regarding the compute intensity and memory bandwidth. Some of them have already been used in internal training and inference scenarios. However, they do not represent best practices and we are actively working on improving the code quality and documentation.
6
+
7
+ ## Features
8
+
9
+ - **Gating** — Top-k expert selection and scoring for Mixture of Experts routing
10
+ - **MoE Routing** — Token-to-expert mapping, fused expansion/reduction and weight normalization
11
+ - **Quantization** — Per-token, per-block, and per-channel FP8/FP4/E5M6 casting with fused SwiGLU+quantization ops
12
+ - **Transpose** — Batched transpose operations
13
+ - **Engram** — Engram gating kernels with fused RMSNorm, forward/backward passes and weight gradient reduction
14
+ - **Manifold HyperConnection** — Hyper-connection kernels including Sinkhorn normalization and mix splitting/application
15
+ - **Modeling** — High-level `torch.autograd.Function` wrappers composing low-level kernels into trainable layers (engram gate, mHC pipeline)
16
+
17
+ ## Requirements
18
+
19
+ - Python 3.10 or higher
20
+ - PyTorch 2.10 or higher
21
+ - TileLang 0.1.9 or higher
22
+ - NVIDIA SM90 or SM100 architecture GPU
23
+ - CUDA Toolkit 13.1 or higher
24
+
25
+ ## Installation
26
+
27
+ ### Install a local development version
28
+
29
+ ```bash
30
+ pip install -e ".[dev]"
31
+ ```
32
+
33
+ ### Install a release version
34
+
35
+ ```bash
36
+ pip install tile-kernels
37
+ ```
38
+
39
+ ## Testing
40
+
41
+ Tests using pytest:
42
+
43
+ ### Test single test file
44
+
45
+ ```bash
46
+ pytest tests/transpose/test_transpose.py -n 4 # Correctness only with 4 workers
47
+ pytest tests/transpose/test_transpose.py --run-benchmark # Correctness + Benchmarking
48
+ ```
49
+
50
+ ### Pressure test
51
+
52
+ ```bash
53
+ TK_FULL_TEST=1 pytest -n 4 --count 2
54
+ ```
55
+
56
+ ## Project Structure
57
+
58
+ ```txt
59
+ tile_kernels/
60
+ ├── moe/ # Mixture of Experts routing related kernels
61
+ ├── quant/ # FP8/FP4/E5M6 quantization
62
+ ├── transpose/ # Batched transpose
63
+ ├── engram/ # Engram gating kernels
64
+ ├── mhc/ # Manifold HyperConnection kernels
65
+ ├── modeling/ # High-level autograd modeling layers (engram, mHC)
66
+ ├── torch/ # PyTorch reference implementations
67
+ └── testing/ # Test and benchmark utilities
68
+ ```
69
+
70
+ ## Acknowledgement
71
+
72
+ This project is built on [TileLang](https://github.com/tile-ai/tilelang). Thanks and respect to the developers!
73
+
74
+ ## License
75
+
76
+ This code repository is released under [the MIT License](LICENSE).
77
+
78
+ ## Citation
79
+
80
+ ```bibtex
81
+ @misc{tilekernels,
82
+ title={TileKernels},
83
+ author={Xiangwen Wang, Chenhao Xu, Huanqi Cao, Rui Tian, Weilin Zhao, Kuai Yu and Chenggang Zhao},
84
+ year={2026},
85
+ publisher = {GitHub},
86
+ howpublished = {\url{https://github.com/deepseek-ai/TileKernels}},
87
+ }
88
+ ```
@@ -0,0 +1,58 @@
1
+ [build-system]
2
+ requires = ["setuptools", "wheel", "setuptools-scm>=8"]
3
+ build-backend = "setuptools.build_meta"
4
+
5
+ [tool.setuptools_scm]
6
+ version_file = "tile_kernels/_version.py"
7
+
8
+ [project]
9
+ name = "tile_kernels"
10
+ dynamic = ["version"]
11
+ description = "Tilelang-based kernels."
12
+ readme = "README.md"
13
+ license = "MIT"
14
+ authors = [
15
+ { name = "Chenhao Xu", email = "xch@deepseek.com" },
16
+ { name = "Xiangwen Wang", email = "xiangwen.wang@deepseek.com" },
17
+ { name = "Huanqi Cao", email = "caohuanqi@deepseek.com" },
18
+ { name = "Rui Tian", email = "tianr22@deepseek.com"},
19
+ { name = "Weilin Zhao", email = "zhaoweilin@deepseek.com" },
20
+ { name = "Kuai Yu", email = "yukuai@deepseek.com" },
21
+ { name = "Chenggang Zhao", email = "chenggangz@deepseek.com"}
22
+ ]
23
+ dependencies = [
24
+ "torch>=2.10",
25
+ "tilelang>=0.1.9"
26
+ ]
27
+ requires-python = ">=3.10"
28
+ classifiers = [
29
+ "Development Status :: 3 - Alpha",
30
+ "Intended Audience :: Developers",
31
+ "Programming Language :: Python :: 3",
32
+ "Programming Language :: Python :: 3.10",
33
+ "Programming Language :: Python :: 3.11",
34
+ "Programming Language :: Python :: 3.12",
35
+ "Topic :: Scientific/Engineering :: Artificial Intelligence",
36
+ ]
37
+
38
+ [project.optional-dependencies]
39
+ dev = ["setuptools", "wheel", "setuptools-scm>=8", "pytest", "pytest-xdist", "pytest-repeat"]
40
+
41
+ [tool.setuptools.packages.find]
42
+ where = ["."]
43
+ include = ["tile_kernels*"]
44
+
45
+ [project.urls]
46
+ Homepage = "https://github.com/deepseek-ai/TileKernels"
47
+
48
+ # Linter tools
49
+
50
+ [tool.ruff]
51
+ line-length = 150
52
+
53
+ [tool.ruff.lint]
54
+ select = ["Q000"]
55
+ fixable = ["Q000"]
56
+
57
+ [tool.ruff.lint.flake8-quotes]
58
+ inline-quotes = "single"
@@ -0,0 +1,4 @@
1
+ [egg_info]
2
+ tag_build =
3
+ tag_date = 0
4
+
File without changes
@@ -0,0 +1,10 @@
1
+ # Root-level conftest
2
+ #
3
+ # Loads the benchmark plugin (CLI options, markers, fixtures).
4
+ # The plugin lives in a file deliberately NOT named conftest.py to
5
+ # avoid pluggy's duplicate-registration error.
6
+
7
+ pytest_plugins = [
8
+ 'tests.pytest_random_plugin',
9
+ 'tests.pytest_benchmark_plugin',
10
+ ]
@@ -0,0 +1,56 @@
1
+ import os
2
+ import pytest
3
+ import torch
4
+
5
+ from tile_kernels.engram import fused_weight
6
+ from tile_kernels.testing.numeric import assert_equal, count_bytes
7
+ from tile_kernels.testing.generator import generate_hidden_sizes
8
+ from tile_kernels.testing.bench import make_param_id
9
+
10
+ # Disable TileLang prints
11
+ os.environ['TILELANG_PRINT_ON_COMPILATION'] = '0'
12
+
13
+
14
+ def generate_test_data(params):
15
+ hc_mult = params['hc']
16
+ hidden_size = params['hidden']
17
+ wh_data = torch.randn(hc_mult, hidden_size, dtype=torch.bfloat16, device='cuda')
18
+ we_data = torch.randn(hc_mult, hidden_size, dtype=torch.bfloat16, device='cuda')
19
+ return (wh_data, we_data)
20
+
21
+
22
+ def generate_test_params(is_benchmark: bool) -> list[dict]:
23
+ return [
24
+ {'hc': hc, 'hidden': hidden_size}
25
+ for hc in (4,)
26
+ for hidden_size in generate_hidden_sizes(128)
27
+ ]
28
+
29
+
30
+ @pytest.mark.parametrize('params', generate_test_params(is_benchmark=False), ids=make_param_id)
31
+ def test_engram_fused_weight(params):
32
+ wh_data, we_data = generate_test_data(params)
33
+
34
+ ref = wh_data.float() * we_data.float()
35
+ out = fused_weight(wh_data, we_data)
36
+
37
+ assert_equal(out, ref)
38
+
39
+
40
+ @pytest.mark.benchmark
41
+ @pytest.mark.parametrize('params', generate_test_params(is_benchmark=True), ids=make_param_id)
42
+ def test_engram_fused_weight_benchmark(benchmark_timer, benchmark_record, params):
43
+ wh_data, we_data = generate_test_data(params)
44
+ out = fused_weight(wh_data, we_data)
45
+
46
+ t_us = benchmark_timer(lambda: fused_weight(wh_data, we_data))
47
+
48
+ num_bytes = count_bytes(wh_data, we_data, out)
49
+ bandwidth_gbs = num_bytes / t_us / 1e3
50
+ benchmark_record(
51
+ kernel='fused_weight',
52
+ operation='fwd',
53
+ params=params,
54
+ time_us=t_us,
55
+ bandwidth_gbs=bandwidth_gbs,
56
+ )
@@ -0,0 +1,112 @@
1
+ import os
2
+ import pytest
3
+ import torch
4
+
5
+ from tile_kernels.engram import engram_gate_bwd
6
+ from tile_kernels.torch.engram import engram_gate_ref
7
+ from tile_kernels.testing.numeric import calc_diff, count_bytes
8
+ from tile_kernels.testing.generator import generate_hidden_sizes, generate_num_tokens
9
+ from tile_kernels.testing.bench import make_param_id
10
+
11
+ # Disable TileLang prints
12
+ os.environ['TILELANG_PRINT_ON_COMPILATION'] = '0'
13
+
14
+
15
+ def generate_test_data(params):
16
+ num_tokens = params['num_tokens']
17
+ hc_mult = params['hc']
18
+ hidden_size = params['hidden']
19
+ eps = 1e-20
20
+ clamp_value = 1e-6
21
+ x_data = torch.randn(num_tokens, hc_mult, hidden_size, dtype=torch.bfloat16, device='cuda')
22
+ k_data = torch.randn(num_tokens, hc_mult, hidden_size, dtype=torch.bfloat16, device='cuda')
23
+ v_data = torch.randn(num_tokens, hidden_size, dtype=torch.bfloat16, device='cuda')
24
+ wh_data = torch.randn(hc_mult, hidden_size, dtype=torch.bfloat16, device='cuda')
25
+ we_data = torch.randn(hc_mult, hidden_size, dtype=torch.bfloat16, device='cuda')
26
+ weight_fused = wh_data.float() * we_data.float()
27
+ grad_out = torch.randn(num_tokens, hc_mult, hidden_size, dtype=torch.bfloat16, device='cuda')
28
+ return (x_data, k_data, v_data, wh_data, we_data, weight_fused, grad_out, eps, clamp_value)
29
+
30
+
31
+ def generate_test_params(is_benchmark: bool) -> list[dict]:
32
+ return [
33
+ {'num_tokens': t, 'hc': hc, 'hidden': hidden_size}
34
+ for t in generate_num_tokens(is_benchmark=is_benchmark)
35
+ for hc in (4,)
36
+ for hidden_size in generate_hidden_sizes(128)
37
+ ]
38
+
39
+
40
+ @pytest.mark.parametrize('params', generate_test_params(is_benchmark=False), ids=make_param_id)
41
+ def test_engram_gate_bwd(params):
42
+ (x_data, k_data, v_data, wh_data, we_data, weight_fused, grad_out, eps, clamp_value) = generate_test_data(params)
43
+
44
+ # Reference: forward with intermediates + autograd backward
45
+ x_ref = x_data.clone().requires_grad_(True)
46
+ k_ref = k_data.clone().requires_grad_(True)
47
+ v_ref = v_data.clone().requires_grad_(True)
48
+ # Cast to float32 so autograd produces fp32 gradients matching the kernel
49
+ wh_ref = wh_data.float().requires_grad_(True)
50
+ we_ref = we_data.float().requires_grad_(True)
51
+ o_ref, dot_ref, gate_score_ref, rstd_x_ref, rstd_k_ref = engram_gate_ref(
52
+ x_ref, k_ref, v_ref, wh_ref, we_ref, clamp_value, eps, save_for_backward=True,
53
+ )
54
+ o_ref.backward(grad_out)
55
+
56
+ # Kernel backward using ref intermediates
57
+ grad_x, grad_k, grad_v, grad_w_partial = engram_gate_bwd(
58
+ grad_out, x_data, k_data, v_data, weight_fused,
59
+ dot_ref, gate_score_ref, rstd_x_ref, rstd_k_ref, clamp_value,
60
+ )
61
+ grad_w_fused = grad_w_partial.sum(0)
62
+ grad_wh = grad_w_fused * we_data.float()
63
+ grad_we = grad_w_fused * wh_data.float()
64
+
65
+ # Correctness
66
+ diff_x = calc_diff(grad_x, x_ref.grad)
67
+ assert diff_x < 1e-8, f'grad_x mismatch: {diff_x:.6e}'
68
+ diff_k = calc_diff(grad_k, k_ref.grad)
69
+ assert diff_k < 1e-8, f'grad_k mismatch: {diff_k:.6e}'
70
+ diff_v = calc_diff(grad_v, v_ref.grad)
71
+ assert diff_v < 1e-8, f'grad_v mismatch: {diff_v:.6e}'
72
+ diff_wh = calc_diff(grad_wh, wh_ref.grad)
73
+ assert diff_wh < 1e-8, f'grad_wh mismatch: {diff_wh:.6e}'
74
+ diff_we = calc_diff(grad_we, we_ref.grad)
75
+ assert diff_we < 1e-8, f'grad_we mismatch: {diff_we:.6e}'
76
+
77
+
78
+ @pytest.mark.benchmark
79
+ @pytest.mark.parametrize('params', generate_test_params(is_benchmark=True), ids=make_param_id)
80
+ def test_engram_gate_bwd_benchmark(benchmark_timer, benchmark_record, params):
81
+ (x_data, k_data, v_data, wh_data, we_data, weight_fused, grad_out, eps, clamp_value) = generate_test_data(params)
82
+
83
+ # Forward to get intermediates
84
+ o_ref, dot_ref, gate_score_ref, rstd_x_ref, rstd_k_ref = engram_gate_ref(
85
+ x_data, k_data, v_data, wh_data, we_data, clamp_value, eps, save_for_backward=True,
86
+ )
87
+
88
+ grad_x, grad_k, grad_v, grad_w_partial = engram_gate_bwd(
89
+ grad_out, x_data, k_data, v_data, weight_fused,
90
+ dot_ref, gate_score_ref, rstd_x_ref, rstd_k_ref, clamp_value,
91
+ )
92
+ grad_w_fused = grad_w_partial.sum(0)
93
+ grad_wh = grad_w_fused * we_data.float()
94
+ grad_we = grad_w_fused * wh_data.float()
95
+
96
+ func_bwd = lambda: engram_gate_bwd(
97
+ grad_out, x_data, k_data, v_data, weight_fused,
98
+ dot_ref, gate_score_ref, rstd_x_ref, rstd_k_ref, clamp_value,
99
+ )
100
+ t_us = benchmark_timer(func_bwd)
101
+ num_bytes = count_bytes(
102
+ grad_out, x_data, k_data, v_data, weight_fused,
103
+ dot_ref, gate_score_ref, rstd_x_ref, rstd_k_ref,
104
+ grad_x, grad_k, grad_v, grad_wh, grad_we,
105
+ )
106
+ benchmark_record(
107
+ kernel='engram_gate_bwd',
108
+ operation='bwd',
109
+ params=params,
110
+ time_us=t_us,
111
+ bandwidth_gbs=num_bytes / t_us / 1e3,
112
+ )
@@ -0,0 +1,108 @@
1
+ import os
2
+ import pytest
3
+ import torch
4
+
5
+ from tile_kernels.engram import engram_gate_fwd
6
+ from tile_kernels.torch.engram import engram_gate_ref
7
+ from tile_kernels.testing.numeric import assert_equal, calc_diff, count_bytes
8
+ from tile_kernels.testing.generator import generate_hidden_sizes, generate_num_tokens
9
+ from tile_kernels.testing.bench import make_param_id
10
+
11
+ # Disable TileLang prints
12
+ os.environ['TILELANG_PRINT_ON_COMPILATION'] = '0'
13
+
14
+
15
+ def generate_test_data(params):
16
+ num_tokens = params['num_tokens']
17
+ hc_mult = params['hc']
18
+ hidden_size = params['hidden']
19
+ eps = 1e-20
20
+ clamp_value = 1e-6
21
+ x_data = torch.randn(num_tokens, hc_mult, hidden_size, dtype=torch.bfloat16, device='cuda')
22
+ k_data = torch.randn(num_tokens, hc_mult, hidden_size, dtype=torch.bfloat16, device='cuda')
23
+ v_data = torch.randn(num_tokens, hidden_size, dtype=torch.bfloat16, device='cuda')
24
+ wh_data = torch.randn(hc_mult, hidden_size, dtype=torch.bfloat16, device='cuda')
25
+ we_data = torch.randn(hc_mult, hidden_size, dtype=torch.bfloat16, device='cuda')
26
+ weight_fused = wh_data.float() * we_data.float()
27
+ return (x_data, k_data, v_data, wh_data, we_data, weight_fused, eps, clamp_value)
28
+
29
+
30
+ def generate_test_params(is_benchmark: bool) -> list[dict]:
31
+ return [
32
+ {'num_tokens': t, 'hc': hc, 'hidden': hidden_size}
33
+ for t in generate_num_tokens(is_benchmark=is_benchmark)
34
+ for hc in (4,)
35
+ for hidden_size in generate_hidden_sizes(128)
36
+ ]
37
+
38
+
39
+ @pytest.mark.parametrize('params', generate_test_params(is_benchmark=False), ids=make_param_id)
40
+ def test_engram_gate_fwd(params):
41
+ (x_data, k_data, v_data, wh_data, we_data, weight_fused, eps, clamp_value) = generate_test_data(params)
42
+
43
+ out_ref, dot_ref, gate_score_ref, rstd_x_ref, rstd_k_ref = engram_gate_ref(
44
+ x_data, k_data, v_data, wh_data, we_data, clamp_value, eps, save_for_backward=True,
45
+ )
46
+
47
+ # Correctness: save_for_backward=True
48
+ out_save, dot, gate_score, rstd_x, rstd_k = engram_gate_fwd(
49
+ x_data, k_data, v_data, weight_fused, eps, clamp_value, save_for_backward=True,
50
+ )
51
+ assert dot is not None and gate_score is not None and rstd_x is not None and rstd_k is not None
52
+ diff_out = calc_diff(out_save, out_ref)
53
+ assert diff_out < 2e-10, f'out_save mismatch: {diff_out:.6e}'
54
+ diff_dot = calc_diff(dot, dot_ref)
55
+ assert diff_dot < 2e-10, f'dot mismatch: {diff_dot:.6e}'
56
+ diff_gate = calc_diff(gate_score, gate_score_ref)
57
+ assert diff_gate < 2e-10, f'gate_score mismatch: {diff_gate:.6e}'
58
+ diff_rstd_x = calc_diff(rstd_x, rstd_x_ref)
59
+ assert diff_rstd_x < 2e-10, f'rstd_x mismatch: {diff_rstd_x:.6e}'
60
+ diff_rstd_k = calc_diff(rstd_k, rstd_k_ref)
61
+ assert diff_rstd_k < 2e-10, f'rstd_k mismatch: {diff_rstd_k:.6e}'
62
+
63
+ # Correctness: save_for_backward=False
64
+ out_no_save, dot_n, gate_score_n, rstd_x_n, rstd_k_n = engram_gate_fwd(
65
+ x_data, k_data, v_data, weight_fused, eps, clamp_value, save_for_backward=False,
66
+ )
67
+ assert dot_n is None and gate_score_n is None and rstd_x_n is None and rstd_k_n is None
68
+ diff_out = calc_diff(out_no_save, out_ref)
69
+ assert diff_out < 2e-10, f'out_no_save mismatch: {diff_out:.6e}'
70
+ assert_equal(out_no_save, out_save)
71
+
72
+
73
+ @pytest.mark.benchmark
74
+ @pytest.mark.parametrize('params', generate_test_params(is_benchmark=True), ids=make_param_id)
75
+ def test_engram_gate_fwd_benchmark(benchmark_timer, benchmark_record, params):
76
+ (x_data, k_data, v_data, _, _, weight_fused, eps, clamp_value) = generate_test_data(params)
77
+
78
+ # Benchmark save_for_backward=True
79
+ out_save, dot, gate_score, rstd_x, rstd_k = engram_gate_fwd(
80
+ x_data, k_data, v_data, weight_fused, eps, clamp_value, save_for_backward=True,
81
+ )
82
+ t_save_us = benchmark_timer(lambda: engram_gate_fwd(
83
+ x_data, k_data, v_data, weight_fused, eps, clamp_value, save_for_backward=True,
84
+ ))
85
+ num_bytes_save = count_bytes(x_data, k_data, v_data, weight_fused, out_save, dot, gate_score, rstd_x, rstd_k)
86
+ benchmark_record(
87
+ kernel='engram_gate_fwd',
88
+ operation='fwd',
89
+ params={**params, 'save': True},
90
+ time_us=t_save_us,
91
+ bandwidth_gbs=num_bytes_save / t_save_us / 1e3,
92
+ )
93
+
94
+ # Benchmark save_for_backward=False
95
+ out_no_save = engram_gate_fwd(
96
+ x_data, k_data, v_data, weight_fused, eps, clamp_value, save_for_backward=False,
97
+ )[0]
98
+ t_no_save_us = benchmark_timer(lambda: engram_gate_fwd(
99
+ x_data, k_data, v_data, weight_fused, eps, clamp_value, save_for_backward=False,
100
+ ))
101
+ num_bytes_no_save = count_bytes(x_data, k_data, v_data, weight_fused, out_no_save)
102
+ benchmark_record(
103
+ kernel='engram_gate_fwd',
104
+ operation='fwd',
105
+ params={**params, 'save': False},
106
+ time_us=t_no_save_us,
107
+ bandwidth_gbs=num_bytes_no_save / t_no_save_us / 1e3,
108
+ )