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.
- tile_kernels-1.0.0/.editorconfig +58 -0
- tile_kernels-1.0.0/.gitignore +20 -0
- tile_kernels-1.0.0/LICENSE +21 -0
- tile_kernels-1.0.0/PKG-INFO +116 -0
- tile_kernels-1.0.0/README.md +88 -0
- tile_kernels-1.0.0/pyproject.toml +58 -0
- tile_kernels-1.0.0/setup.cfg +4 -0
- tile_kernels-1.0.0/tests/__init__.py +0 -0
- tile_kernels-1.0.0/tests/conftest.py +10 -0
- tile_kernels-1.0.0/tests/engram/test_engram_fused_weight.py +56 -0
- tile_kernels-1.0.0/tests/engram/test_engram_gate_bwd.py +112 -0
- tile_kernels-1.0.0/tests/engram/test_engram_gate_fwd.py +108 -0
- tile_kernels-1.0.0/tests/engram/test_engram_grad_w_reduce.py +76 -0
- tile_kernels-1.0.0/tests/engram/test_engram_hash.py +71 -0
- tile_kernels-1.0.0/tests/mhc/test_expand.py +45 -0
- tile_kernels-1.0.0/tests/mhc/test_head_compute_mix.py +54 -0
- tile_kernels-1.0.0/tests/mhc/test_multilayer_recompute.py +157 -0
- tile_kernels-1.0.0/tests/mhc/test_norm_fn.py +130 -0
- tile_kernels-1.0.0/tests/mhc/test_post.py +72 -0
- tile_kernels-1.0.0/tests/mhc/test_pre_apply_mix.py +52 -0
- tile_kernels-1.0.0/tests/mhc/test_pre_big_fuse.py +138 -0
- tile_kernels-1.0.0/tests/mhc/test_pre_split_mixes.py +103 -0
- tile_kernels-1.0.0/tests/mhc/test_sinkhorn.py +43 -0
- tile_kernels-1.0.0/tests/moe/test_aux_fi.py +65 -0
- tile_kernels-1.0.0/tests/moe/test_expand_to_fused.py +150 -0
- tile_kernels-1.0.0/tests/moe/test_get_fused_mapping.py +103 -0
- tile_kernels-1.0.0/tests/moe/test_group_count.py +55 -0
- tile_kernels-1.0.0/tests/moe/test_inplace_unique_group_indices.py +79 -0
- tile_kernels-1.0.0/tests/moe/test_mask_indices_by_tp.py +70 -0
- tile_kernels-1.0.0/tests/moe/test_normalize_weight.py +57 -0
- tile_kernels-1.0.0/tests/moe/test_reduce_fused.py +112 -0
- tile_kernels-1.0.0/tests/moe/test_top2_sum_gate.py +359 -0
- tile_kernels-1.0.0/tests/moe/test_topk_gate.py +75 -0
- tile_kernels-1.0.0/tests/moe/test_topk_sum_and_topk_idx.py +86 -0
- tile_kernels-1.0.0/tests/pytest_benchmark_plugin.py +477 -0
- tile_kernels-1.0.0/tests/pytest_random_plugin.py +18 -0
- tile_kernels-1.0.0/tests/quant/test_cast_back.py +157 -0
- tile_kernels-1.0.0/tests/quant/test_cast_back_e5m6.py +120 -0
- tile_kernels-1.0.0/tests/quant/test_per_block_cast.py +110 -0
- tile_kernels-1.0.0/tests/quant/test_per_block_cast_lossless.py +114 -0
- tile_kernels-1.0.0/tests/quant/test_per_channel_cast.py +75 -0
- tile_kernels-1.0.0/tests/quant/test_per_channel_cast_and_transpose.py +77 -0
- tile_kernels-1.0.0/tests/quant/test_per_channel_cast_fused.py +104 -0
- tile_kernels-1.0.0/tests/quant/test_per_token_cast.py +166 -0
- tile_kernels-1.0.0/tests/quant/test_per_token_cast_to_e5m6.py +118 -0
- tile_kernels-1.0.0/tests/quant/test_swiglu_backward_and_per_token_cast.py +149 -0
- tile_kernels-1.0.0/tests/quant/test_swiglu_forward_and_per_channel_cast_and_transpose.py +114 -0
- tile_kernels-1.0.0/tests/quant/test_swiglu_forward_and_per_token_cast.py +179 -0
- tile_kernels-1.0.0/tests/transpose/test_transpose.py +117 -0
- tile_kernels-1.0.0/tile_kernels/__init__.py +15 -0
- tile_kernels-1.0.0/tile_kernels/_version.py +34 -0
- tile_kernels-1.0.0/tile_kernels/config.py +29 -0
- tile_kernels-1.0.0/tile_kernels/engram/__init__.py +4 -0
- tile_kernels-1.0.0/tile_kernels/engram/engram_fused_weight_kernel.py +59 -0
- tile_kernels-1.0.0/tile_kernels/engram/engram_gate_kernel.py +570 -0
- tile_kernels-1.0.0/tile_kernels/engram/engram_grad_w_reduce_kernel.py +89 -0
- tile_kernels-1.0.0/tile_kernels/engram/engram_hash_kernel.py +97 -0
- tile_kernels-1.0.0/tile_kernels/mhc/__init__.py +0 -0
- tile_kernels-1.0.0/tile_kernels/mhc/expand_kernel.py +60 -0
- tile_kernels-1.0.0/tile_kernels/mhc/head_compute_mix_kernel.py +85 -0
- tile_kernels-1.0.0/tile_kernels/mhc/multilayer_recompute_kernel.py +196 -0
- tile_kernels-1.0.0/tile_kernels/mhc/norm_fn_kernel.py +289 -0
- tile_kernels-1.0.0/tile_kernels/mhc/post_kernel.py +221 -0
- tile_kernels-1.0.0/tile_kernels/mhc/pre_apply_mix_kernel.py +110 -0
- tile_kernels-1.0.0/tile_kernels/mhc/pre_big_fuse_kernel.py +131 -0
- tile_kernels-1.0.0/tile_kernels/mhc/pre_split_mixes_kernel.py +163 -0
- tile_kernels-1.0.0/tile_kernels/mhc/sinkhorn_kernel.py +165 -0
- tile_kernels-1.0.0/tile_kernels/modeling/__init__.py +2 -0
- tile_kernels-1.0.0/tile_kernels/modeling/engram/__init__.py +1 -0
- tile_kernels-1.0.0/tile_kernels/modeling/engram/engram_gate.py +95 -0
- tile_kernels-1.0.0/tile_kernels/modeling/mhc/__init__.py +1 -0
- tile_kernels-1.0.0/tile_kernels/modeling/mhc/functional.py +160 -0
- tile_kernels-1.0.0/tile_kernels/modeling/mhc/ops/__init__.py +9 -0
- tile_kernels-1.0.0/tile_kernels/modeling/mhc/ops/expand.py +34 -0
- tile_kernels-1.0.0/tile_kernels/modeling/mhc/ops/head_compute_mix.py +77 -0
- tile_kernels-1.0.0/tile_kernels/modeling/mhc/ops/multilayer_recompute.py +3 -0
- tile_kernels-1.0.0/tile_kernels/modeling/mhc/ops/norm_fn.py +189 -0
- tile_kernels-1.0.0/tile_kernels/modeling/mhc/ops/post.py +35 -0
- tile_kernels-1.0.0/tile_kernels/modeling/mhc/ops/pre_apply_mix.py +55 -0
- tile_kernels-1.0.0/tile_kernels/modeling/mhc/ops/pre_big_fuse.py +91 -0
- tile_kernels-1.0.0/tile_kernels/modeling/mhc/ops/pre_split_mixes.py +122 -0
- tile_kernels-1.0.0/tile_kernels/modeling/mhc/ops/sinkhorn.py +32 -0
- tile_kernels-1.0.0/tile_kernels/moe/__init__.py +11 -0
- tile_kernels-1.0.0/tile_kernels/moe/aux_fi_kernel.py +72 -0
- tile_kernels-1.0.0/tile_kernels/moe/common.py +52 -0
- tile_kernels-1.0.0/tile_kernels/moe/expand_to_fused_kernel.py +200 -0
- tile_kernels-1.0.0/tile_kernels/moe/get_fused_mapping_kernel.py +250 -0
- tile_kernels-1.0.0/tile_kernels/moe/group_count_kernel.py +69 -0
- tile_kernels-1.0.0/tile_kernels/moe/inplace_unique_group_indices_kernel.py +69 -0
- tile_kernels-1.0.0/tile_kernels/moe/mask_indices_by_tp_kernel.py +73 -0
- tile_kernels-1.0.0/tile_kernels/moe/normalize_weight_kernel.py +70 -0
- tile_kernels-1.0.0/tile_kernels/moe/reduce_fused_kernel.py +136 -0
- tile_kernels-1.0.0/tile_kernels/moe/scoring.py +26 -0
- tile_kernels-1.0.0/tile_kernels/moe/top2_sum_gate_kernel.py +424 -0
- tile_kernels-1.0.0/tile_kernels/moe/topk_gate_kernel.py +90 -0
- tile_kernels-1.0.0/tile_kernels/moe/topk_sum_and_topk_group_idx_kernel.py +102 -0
- tile_kernels-1.0.0/tile_kernels/quant/__init__.py +11 -0
- tile_kernels-1.0.0/tile_kernels/quant/cast_back_e5m6_kernel.py +176 -0
- tile_kernels-1.0.0/tile_kernels/quant/cast_back_kernel.py +143 -0
- tile_kernels-1.0.0/tile_kernels/quant/common.py +295 -0
- tile_kernels-1.0.0/tile_kernels/quant/per_block_cast_kernel.py +276 -0
- tile_kernels-1.0.0/tile_kernels/quant/per_block_cast_lossless_kernel.py +209 -0
- tile_kernels-1.0.0/tile_kernels/quant/per_channel_cast_and_transpose_kernel.py +123 -0
- tile_kernels-1.0.0/tile_kernels/quant/per_channel_cast_fused_kernel.py +203 -0
- tile_kernels-1.0.0/tile_kernels/quant/per_channel_cast_kernel.py +28 -0
- tile_kernels-1.0.0/tile_kernels/quant/per_token_cast_kernel.py +302 -0
- tile_kernels-1.0.0/tile_kernels/quant/per_token_cast_to_e5m6_kernel.py +215 -0
- tile_kernels-1.0.0/tile_kernels/quant/swiglu_backward_and_per_token_cast_kernel.py +237 -0
- tile_kernels-1.0.0/tile_kernels/quant/swiglu_forward_and_per_channel_cast_and_transpose_kernel.py +231 -0
- tile_kernels-1.0.0/tile_kernels/quant/swiglu_forward_and_per_token_cast_kernel.py +260 -0
- tile_kernels-1.0.0/tile_kernels/quant/types.py +3 -0
- tile_kernels-1.0.0/tile_kernels/testing/__init__.py +2 -0
- tile_kernels-1.0.0/tile_kernels/testing/bench.py +115 -0
- tile_kernels-1.0.0/tile_kernels/testing/generator.py +105 -0
- tile_kernels-1.0.0/tile_kernels/testing/numeric.py +65 -0
- tile_kernels-1.0.0/tile_kernels/testing/quant.py +21 -0
- tile_kernels-1.0.0/tile_kernels/torch/__init__.py +10 -0
- tile_kernels-1.0.0/tile_kernels/torch/cast.py +252 -0
- tile_kernels-1.0.0/tile_kernels/torch/cast_e5m6.py +258 -0
- tile_kernels-1.0.0/tile_kernels/torch/engram.py +113 -0
- tile_kernels-1.0.0/tile_kernels/torch/expand_to_fused.py +91 -0
- tile_kernels-1.0.0/tile_kernels/torch/mhc.py +85 -0
- tile_kernels-1.0.0/tile_kernels/torch/moe.py +120 -0
- tile_kernels-1.0.0/tile_kernels/torch/per_channel_cast_fused.py +44 -0
- tile_kernels-1.0.0/tile_kernels/torch/reduce_fused.py +74 -0
- tile_kernels-1.0.0/tile_kernels/torch/swiglu.py +227 -0
- tile_kernels-1.0.0/tile_kernels/torch/topk.py +206 -0
- tile_kernels-1.0.0/tile_kernels/transpose/__init__.py +1 -0
- tile_kernels-1.0.0/tile_kernels/transpose/batched_transpose_kernel.py +119 -0
- tile_kernels-1.0.0/tile_kernels/utils.py +10 -0
- tile_kernels-1.0.0/tile_kernels.egg-info/PKG-INFO +116 -0
- tile_kernels-1.0.0/tile_kernels.egg-info/SOURCES.txt +133 -0
- tile_kernels-1.0.0/tile_kernels.egg-info/dependency_links.txt +1 -0
- tile_kernels-1.0.0/tile_kernels.egg-info/requires.txt +10 -0
- 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,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"
|
|
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
|
+
)
|