tile-kernels 1.0.0__py3-none-any.whl

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 (86) hide show
  1. tile_kernels/__init__.py +15 -0
  2. tile_kernels/_version.py +34 -0
  3. tile_kernels/config.py +29 -0
  4. tile_kernels/engram/__init__.py +4 -0
  5. tile_kernels/engram/engram_fused_weight_kernel.py +59 -0
  6. tile_kernels/engram/engram_gate_kernel.py +570 -0
  7. tile_kernels/engram/engram_grad_w_reduce_kernel.py +89 -0
  8. tile_kernels/engram/engram_hash_kernel.py +97 -0
  9. tile_kernels/mhc/__init__.py +0 -0
  10. tile_kernels/mhc/expand_kernel.py +60 -0
  11. tile_kernels/mhc/head_compute_mix_kernel.py +85 -0
  12. tile_kernels/mhc/multilayer_recompute_kernel.py +196 -0
  13. tile_kernels/mhc/norm_fn_kernel.py +289 -0
  14. tile_kernels/mhc/post_kernel.py +221 -0
  15. tile_kernels/mhc/pre_apply_mix_kernel.py +110 -0
  16. tile_kernels/mhc/pre_big_fuse_kernel.py +131 -0
  17. tile_kernels/mhc/pre_split_mixes_kernel.py +163 -0
  18. tile_kernels/mhc/sinkhorn_kernel.py +165 -0
  19. tile_kernels/modeling/__init__.py +2 -0
  20. tile_kernels/modeling/engram/__init__.py +1 -0
  21. tile_kernels/modeling/engram/engram_gate.py +95 -0
  22. tile_kernels/modeling/mhc/__init__.py +1 -0
  23. tile_kernels/modeling/mhc/functional.py +160 -0
  24. tile_kernels/modeling/mhc/ops/__init__.py +9 -0
  25. tile_kernels/modeling/mhc/ops/expand.py +34 -0
  26. tile_kernels/modeling/mhc/ops/head_compute_mix.py +77 -0
  27. tile_kernels/modeling/mhc/ops/multilayer_recompute.py +3 -0
  28. tile_kernels/modeling/mhc/ops/norm_fn.py +189 -0
  29. tile_kernels/modeling/mhc/ops/post.py +35 -0
  30. tile_kernels/modeling/mhc/ops/pre_apply_mix.py +55 -0
  31. tile_kernels/modeling/mhc/ops/pre_big_fuse.py +91 -0
  32. tile_kernels/modeling/mhc/ops/pre_split_mixes.py +122 -0
  33. tile_kernels/modeling/mhc/ops/sinkhorn.py +32 -0
  34. tile_kernels/moe/__init__.py +11 -0
  35. tile_kernels/moe/aux_fi_kernel.py +72 -0
  36. tile_kernels/moe/common.py +52 -0
  37. tile_kernels/moe/expand_to_fused_kernel.py +200 -0
  38. tile_kernels/moe/get_fused_mapping_kernel.py +250 -0
  39. tile_kernels/moe/group_count_kernel.py +69 -0
  40. tile_kernels/moe/inplace_unique_group_indices_kernel.py +69 -0
  41. tile_kernels/moe/mask_indices_by_tp_kernel.py +73 -0
  42. tile_kernels/moe/normalize_weight_kernel.py +70 -0
  43. tile_kernels/moe/reduce_fused_kernel.py +136 -0
  44. tile_kernels/moe/scoring.py +26 -0
  45. tile_kernels/moe/top2_sum_gate_kernel.py +424 -0
  46. tile_kernels/moe/topk_gate_kernel.py +90 -0
  47. tile_kernels/moe/topk_sum_and_topk_group_idx_kernel.py +102 -0
  48. tile_kernels/quant/__init__.py +11 -0
  49. tile_kernels/quant/cast_back_e5m6_kernel.py +176 -0
  50. tile_kernels/quant/cast_back_kernel.py +143 -0
  51. tile_kernels/quant/common.py +295 -0
  52. tile_kernels/quant/per_block_cast_kernel.py +276 -0
  53. tile_kernels/quant/per_block_cast_lossless_kernel.py +209 -0
  54. tile_kernels/quant/per_channel_cast_and_transpose_kernel.py +123 -0
  55. tile_kernels/quant/per_channel_cast_fused_kernel.py +203 -0
  56. tile_kernels/quant/per_channel_cast_kernel.py +28 -0
  57. tile_kernels/quant/per_token_cast_kernel.py +302 -0
  58. tile_kernels/quant/per_token_cast_to_e5m6_kernel.py +215 -0
  59. tile_kernels/quant/swiglu_backward_and_per_token_cast_kernel.py +237 -0
  60. tile_kernels/quant/swiglu_forward_and_per_channel_cast_and_transpose_kernel.py +231 -0
  61. tile_kernels/quant/swiglu_forward_and_per_token_cast_kernel.py +260 -0
  62. tile_kernels/quant/types.py +3 -0
  63. tile_kernels/testing/__init__.py +2 -0
  64. tile_kernels/testing/bench.py +115 -0
  65. tile_kernels/testing/generator.py +105 -0
  66. tile_kernels/testing/numeric.py +65 -0
  67. tile_kernels/testing/quant.py +21 -0
  68. tile_kernels/torch/__init__.py +10 -0
  69. tile_kernels/torch/cast.py +252 -0
  70. tile_kernels/torch/cast_e5m6.py +258 -0
  71. tile_kernels/torch/engram.py +113 -0
  72. tile_kernels/torch/expand_to_fused.py +91 -0
  73. tile_kernels/torch/mhc.py +85 -0
  74. tile_kernels/torch/moe.py +120 -0
  75. tile_kernels/torch/per_channel_cast_fused.py +44 -0
  76. tile_kernels/torch/reduce_fused.py +74 -0
  77. tile_kernels/torch/swiglu.py +227 -0
  78. tile_kernels/torch/topk.py +206 -0
  79. tile_kernels/transpose/__init__.py +1 -0
  80. tile_kernels/transpose/batched_transpose_kernel.py +119 -0
  81. tile_kernels/utils.py +10 -0
  82. tile_kernels-1.0.0.dist-info/METADATA +116 -0
  83. tile_kernels-1.0.0.dist-info/RECORD +86 -0
  84. tile_kernels-1.0.0.dist-info/WHEEL +5 -0
  85. tile_kernels-1.0.0.dist-info/licenses/LICENSE +21 -0
  86. tile_kernels-1.0.0.dist-info/top_level.txt +1 -0
@@ -0,0 +1,15 @@
1
+ import tilelang
2
+
3
+ from . import (
4
+ config,
5
+ engram,
6
+ mhc,
7
+ modeling,
8
+ moe,
9
+ quant,
10
+ transpose,
11
+ torch,
12
+ testing,
13
+ )
14
+
15
+ from .config import get_num_sms, get_device_num_sms, set_num_sms
@@ -0,0 +1,34 @@
1
+ # file generated by setuptools-scm
2
+ # don't change, don't track in version control
3
+
4
+ __all__ = [
5
+ "__version__",
6
+ "__version_tuple__",
7
+ "version",
8
+ "version_tuple",
9
+ "__commit_id__",
10
+ "commit_id",
11
+ ]
12
+
13
+ TYPE_CHECKING = False
14
+ if TYPE_CHECKING:
15
+ from typing import Tuple
16
+ from typing import Union
17
+
18
+ VERSION_TUPLE = Tuple[Union[int, str], ...]
19
+ COMMIT_ID = Union[str, None]
20
+ else:
21
+ VERSION_TUPLE = object
22
+ COMMIT_ID = object
23
+
24
+ version: str
25
+ __version__: str
26
+ __version_tuple__: VERSION_TUPLE
27
+ version_tuple: VERSION_TUPLE
28
+ commit_id: COMMIT_ID
29
+ __commit_id__: COMMIT_ID
30
+
31
+ __version__ = version = '1.0.0'
32
+ __version_tuple__ = version_tuple = (1, 0, 0)
33
+
34
+ __commit_id__ = commit_id = None
tile_kernels/config.py ADDED
@@ -0,0 +1,29 @@
1
+ import functools
2
+ import torch
3
+
4
+ _num_sms = 0
5
+
6
+
7
+ @functools.lru_cache(maxsize=None)
8
+ def get_device_num_sms() -> int:
9
+ prop = torch.cuda.get_device_properties(torch.cuda.current_device())
10
+ return prop.multi_processor_count
11
+
12
+
13
+ def set_num_sms(num_sms: int) -> None:
14
+ global _num_sms
15
+ assert 0 < num_sms <= get_device_num_sms()
16
+ _num_sms = num_sms
17
+
18
+
19
+ def get_num_sms() -> int:
20
+ global _num_sms
21
+ if _num_sms == 0:
22
+ return get_device_num_sms()
23
+ return _num_sms
24
+
25
+
26
+ @functools.lru_cache(maxsize=None)
27
+ def get_max_smem_per_sm() -> int:
28
+ prop = torch.cuda.get_device_properties(torch.cuda.current_device())
29
+ return prop.shared_memory_per_multiprocessor
@@ -0,0 +1,4 @@
1
+ from .engram_fused_weight_kernel import fused_weight
2
+ from .engram_gate_kernel import engram_gate_fwd, engram_gate_bwd
3
+ from .engram_grad_w_reduce_kernel import grad_w_reduce
4
+ from .engram_hash_kernel import engram_hash
@@ -0,0 +1,59 @@
1
+ import os
2
+
3
+ import torch
4
+ import tilelang
5
+ from tilelang import language as T
6
+
7
+
8
+ @tilelang.jit(
9
+ pass_configs={
10
+ tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True,
11
+ },
12
+ )
13
+ def get_engram_fused_weight_kernel(hidden_size: int, hc_mult: int):
14
+ """Elementwise bf16 x bf16 -> fp32 for weight_hidden * weight_embed."""
15
+ threads = 32
16
+ vec_size = 8
17
+ blk_d = threads * vec_size
18
+ assert hidden_size % blk_d == 0
19
+ num_blk = hidden_size // blk_d
20
+
21
+ @T.prim_func
22
+ def engram_fused_weight_kernel(
23
+ weight_hidden: T.Tensor[(hc_mult, hidden_size), T.bfloat16],
24
+ weight_embed: T.Tensor[(hc_mult, hidden_size), T.bfloat16],
25
+ weight_fused: T.Tensor[(hc_mult, hidden_size), T.float],
26
+ ):
27
+ with T.Kernel(hc_mult, num_blk, threads=threads) as (pid_h, pid_b):
28
+ tid = T.get_thread_binding()
29
+ a_local = T.alloc_local((vec_size,), T.float)
30
+ b_local = T.alloc_local((vec_size,), T.float)
31
+ for i_k in T.vectorized(vec_size):
32
+ a_local[i_k] = weight_hidden[pid_h, pid_b * blk_d + tid * vec_size + i_k]
33
+ b_local[i_k] = weight_embed[pid_h, pid_b * blk_d + tid * vec_size + i_k]
34
+ for i_k in T.vectorized(vec_size):
35
+ weight_fused[pid_h, pid_b * blk_d + tid * vec_size + i_k] = a_local[i_k] * b_local[i_k]
36
+
37
+ return engram_fused_weight_kernel
38
+
39
+
40
+ def fused_weight(weight_hidden: torch.Tensor, weight_embed: torch.Tensor) -> torch.Tensor:
41
+ """Compute weight_hidden * weight_embed in fp32.
42
+
43
+ Args:
44
+ weight_hidden: Shape (hc_mult, hidden_size), bfloat16.
45
+ weight_embed: Shape (hc_mult, hidden_size), bfloat16.
46
+
47
+ Returns:
48
+ weight_fused: Shape (hc_mult, hidden_size), float32.
49
+ """
50
+ hc_mult, hidden_size = weight_hidden.shape
51
+
52
+ kernel = get_engram_fused_weight_kernel(hidden_size, hc_mult)
53
+ if int(os.getenv('TK_PRINT_KERNEL_SOURCE', 0)):
54
+ print(kernel.get_kernel_source())
55
+
56
+ weight_fused = torch.empty(hc_mult, hidden_size, dtype=torch.float32, device=weight_hidden.device)
57
+ kernel(weight_hidden, weight_embed, weight_fused)
58
+
59
+ return weight_fused