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.
- tile_kernels/__init__.py +15 -0
- tile_kernels/_version.py +34 -0
- tile_kernels/config.py +29 -0
- tile_kernels/engram/__init__.py +4 -0
- tile_kernels/engram/engram_fused_weight_kernel.py +59 -0
- tile_kernels/engram/engram_gate_kernel.py +570 -0
- tile_kernels/engram/engram_grad_w_reduce_kernel.py +89 -0
- tile_kernels/engram/engram_hash_kernel.py +97 -0
- tile_kernels/mhc/__init__.py +0 -0
- tile_kernels/mhc/expand_kernel.py +60 -0
- tile_kernels/mhc/head_compute_mix_kernel.py +85 -0
- tile_kernels/mhc/multilayer_recompute_kernel.py +196 -0
- tile_kernels/mhc/norm_fn_kernel.py +289 -0
- tile_kernels/mhc/post_kernel.py +221 -0
- tile_kernels/mhc/pre_apply_mix_kernel.py +110 -0
- tile_kernels/mhc/pre_big_fuse_kernel.py +131 -0
- tile_kernels/mhc/pre_split_mixes_kernel.py +163 -0
- tile_kernels/mhc/sinkhorn_kernel.py +165 -0
- tile_kernels/modeling/__init__.py +2 -0
- tile_kernels/modeling/engram/__init__.py +1 -0
- tile_kernels/modeling/engram/engram_gate.py +95 -0
- tile_kernels/modeling/mhc/__init__.py +1 -0
- tile_kernels/modeling/mhc/functional.py +160 -0
- tile_kernels/modeling/mhc/ops/__init__.py +9 -0
- tile_kernels/modeling/mhc/ops/expand.py +34 -0
- tile_kernels/modeling/mhc/ops/head_compute_mix.py +77 -0
- tile_kernels/modeling/mhc/ops/multilayer_recompute.py +3 -0
- tile_kernels/modeling/mhc/ops/norm_fn.py +189 -0
- tile_kernels/modeling/mhc/ops/post.py +35 -0
- tile_kernels/modeling/mhc/ops/pre_apply_mix.py +55 -0
- tile_kernels/modeling/mhc/ops/pre_big_fuse.py +91 -0
- tile_kernels/modeling/mhc/ops/pre_split_mixes.py +122 -0
- tile_kernels/modeling/mhc/ops/sinkhorn.py +32 -0
- tile_kernels/moe/__init__.py +11 -0
- tile_kernels/moe/aux_fi_kernel.py +72 -0
- tile_kernels/moe/common.py +52 -0
- tile_kernels/moe/expand_to_fused_kernel.py +200 -0
- tile_kernels/moe/get_fused_mapping_kernel.py +250 -0
- tile_kernels/moe/group_count_kernel.py +69 -0
- tile_kernels/moe/inplace_unique_group_indices_kernel.py +69 -0
- tile_kernels/moe/mask_indices_by_tp_kernel.py +73 -0
- tile_kernels/moe/normalize_weight_kernel.py +70 -0
- tile_kernels/moe/reduce_fused_kernel.py +136 -0
- tile_kernels/moe/scoring.py +26 -0
- tile_kernels/moe/top2_sum_gate_kernel.py +424 -0
- tile_kernels/moe/topk_gate_kernel.py +90 -0
- tile_kernels/moe/topk_sum_and_topk_group_idx_kernel.py +102 -0
- tile_kernels/quant/__init__.py +11 -0
- tile_kernels/quant/cast_back_e5m6_kernel.py +176 -0
- tile_kernels/quant/cast_back_kernel.py +143 -0
- tile_kernels/quant/common.py +295 -0
- tile_kernels/quant/per_block_cast_kernel.py +276 -0
- tile_kernels/quant/per_block_cast_lossless_kernel.py +209 -0
- tile_kernels/quant/per_channel_cast_and_transpose_kernel.py +123 -0
- tile_kernels/quant/per_channel_cast_fused_kernel.py +203 -0
- tile_kernels/quant/per_channel_cast_kernel.py +28 -0
- tile_kernels/quant/per_token_cast_kernel.py +302 -0
- tile_kernels/quant/per_token_cast_to_e5m6_kernel.py +215 -0
- tile_kernels/quant/swiglu_backward_and_per_token_cast_kernel.py +237 -0
- tile_kernels/quant/swiglu_forward_and_per_channel_cast_and_transpose_kernel.py +231 -0
- tile_kernels/quant/swiglu_forward_and_per_token_cast_kernel.py +260 -0
- tile_kernels/quant/types.py +3 -0
- tile_kernels/testing/__init__.py +2 -0
- tile_kernels/testing/bench.py +115 -0
- tile_kernels/testing/generator.py +105 -0
- tile_kernels/testing/numeric.py +65 -0
- tile_kernels/testing/quant.py +21 -0
- tile_kernels/torch/__init__.py +10 -0
- tile_kernels/torch/cast.py +252 -0
- tile_kernels/torch/cast_e5m6.py +258 -0
- tile_kernels/torch/engram.py +113 -0
- tile_kernels/torch/expand_to_fused.py +91 -0
- tile_kernels/torch/mhc.py +85 -0
- tile_kernels/torch/moe.py +120 -0
- tile_kernels/torch/per_channel_cast_fused.py +44 -0
- tile_kernels/torch/reduce_fused.py +74 -0
- tile_kernels/torch/swiglu.py +227 -0
- tile_kernels/torch/topk.py +206 -0
- tile_kernels/transpose/__init__.py +1 -0
- tile_kernels/transpose/batched_transpose_kernel.py +119 -0
- tile_kernels/utils.py +10 -0
- tile_kernels-1.0.0.dist-info/METADATA +116 -0
- tile_kernels-1.0.0.dist-info/RECORD +86 -0
- tile_kernels-1.0.0.dist-info/WHEEL +5 -0
- tile_kernels-1.0.0.dist-info/licenses/LICENSE +21 -0
- tile_kernels-1.0.0.dist-info/top_level.txt +1 -0
tile_kernels/__init__.py
ADDED
tile_kernels/_version.py
ADDED
|
@@ -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,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
|