openjev / code /serving /patches /dynamic-paged.patch
AlexWortega's picture
Add image Decisions serving and explicit option probabilities
26de23c verified
Raw History Blame Contribute Delete
2.11 kB
--- 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,