--- 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,