File size: 2,112 Bytes
26de23c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
--- a/python/sglang/srt/layers/attention/tilelang_fa_v100/_kernels_paged.py
+++ b/python/sglang/srt/layers/attention/tilelang_fa_v100/_kernels_paged.py
@@ -10,6 +10,7 @@
    which doubles warp count from 2 to 4, improving V100 occupancy.
 """
 import math
+import os
 import torch
 import tilelang
 import tilelang.language as T
@@ -51,9 +52,13 @@
                        max_blocks_per_seq, num_pages, is_causal,
                        sliding_window_size=-1,
                        block_M=32, block_N=128, num_stages=0, threads=256,
-                       num_splits=1):
+                       num_splits=1, dynamic_shapes=False):
     scale = (1.0 / dim) ** 0.5
     nt = T.dynamic("nt")
+    if dynamic_shapes:
+        batch = T.dynamic("batch")
+        max_blocks_per_seq = T.dynamic("max_blocks_per_seq")
+        num_pages = T.dynamic("num_pages")
     use_kv_union = dim > _USE_KV_UNION_FOR_DIM
 
     if num_splits > 1:
@@ -377,15 +382,18 @@
                      max_blocks, causal, sliding_window_size=-1):
     """Return compiled kernel."""
     cfg = _BEST_CONFIGS.get(dim, dict(block_M=32, block_N=128, threads=256, num_stages=0, num_splits=1))
+    dynamic = os.environ.get("SGLANG_V100_DYNAMIC_PAGED", "0") == "1"
+    shape_key = (0, 0, 0) if dynamic else (batch, max_blocks, num_pages)
     key = (heads, heads_kv, dim, block_size, causal, sliding_window_size,
            cfg["block_M"], cfg["block_N"], cfg["threads"], cfg["num_stages"], cfg["num_splits"],
-           batch, max_blocks, num_pages)
+           dynamic, *shape_key)
     if key not in _KERNEL_CACHE:
         kt = _paged_kernel_func(
-            batch=batch, heads=heads, heads_kv=heads_kv, dim=dim,
+            batch=shape_key[0], heads=heads, heads_kv=heads_kv, dim=dim,
             page_block_size=block_size,
-            max_blocks_per_seq=max_blocks,
-            num_pages=num_pages,
+            max_blocks_per_seq=shape_key[1],
+            num_pages=shape_key[2],
+            dynamic_shapes=dynamic,
             is_causal=causal,
             sliding_window_size=sliding_window_size,
             **cfg,