[EAGLE] Fix slow Triton compilation in EAGLE KV cache copy by chunking large num_locs_upper (#15111)
This commit is contained in:
@@ -673,6 +673,9 @@ class MHATokenToKVPool(KVCache):
|
|||||||
else:
|
else:
|
||||||
bytes_per_tile = _KV_COPY_TILE_SIZE_SMALL
|
bytes_per_tile = _KV_COPY_TILE_SIZE_SMALL
|
||||||
|
|
||||||
|
# Calculate num_locs_upper to avoid large Triton specialization (e.g. 8192)
|
||||||
|
chunk_upper = 128 if bytes_per_tile >= _KV_COPY_TILE_SIZE_LARGE else 256
|
||||||
|
|
||||||
self._kv_copy_config = {
|
self._kv_copy_config = {
|
||||||
"bytes_per_tile": bytes_per_tile,
|
"bytes_per_tile": bytes_per_tile,
|
||||||
"byte_tiles": (stride_bytes + bytes_per_tile - 1) // bytes_per_tile,
|
"byte_tiles": (stride_bytes + bytes_per_tile - 1) // bytes_per_tile,
|
||||||
@@ -681,9 +684,10 @@ class MHATokenToKVPool(KVCache):
|
|||||||
if bytes_per_tile <= _KV_COPY_TILE_SIZE_MEDIUM
|
if bytes_per_tile <= _KV_COPY_TILE_SIZE_MEDIUM
|
||||||
else _KV_COPY_NUM_WARPS_LARGE_TILE
|
else _KV_COPY_NUM_WARPS_LARGE_TILE
|
||||||
),
|
),
|
||||||
|
"num_locs_upper": chunk_upper,
|
||||||
}
|
}
|
||||||
|
|
||||||
dummy_loc = torch.zeros(1, dtype=torch.int32, device=self.device)
|
dummy_loc = torch.zeros(chunk_upper, dtype=torch.int64, device=self.device)
|
||||||
grid = (self.data_ptrs.numel(), self._kv_copy_config["byte_tiles"])
|
grid = (self.data_ptrs.numel(), self._kv_copy_config["byte_tiles"])
|
||||||
|
|
||||||
copy_all_layer_kv_cache_tiled[grid](
|
copy_all_layer_kv_cache_tiled[grid](
|
||||||
@@ -692,7 +696,7 @@ class MHATokenToKVPool(KVCache):
|
|||||||
dummy_loc,
|
dummy_loc,
|
||||||
dummy_loc,
|
dummy_loc,
|
||||||
1,
|
1,
|
||||||
1,
|
chunk_upper,
|
||||||
BYTES_PER_TILE=self._kv_copy_config["bytes_per_tile"],
|
BYTES_PER_TILE=self._kv_copy_config["bytes_per_tile"],
|
||||||
num_warps=self._kv_copy_config["num_warps"],
|
num_warps=self._kv_copy_config["num_warps"],
|
||||||
num_stages=2,
|
num_stages=2,
|
||||||
@@ -902,20 +906,40 @@ class MHATokenToKVPool(KVCache):
|
|||||||
), "KV copy not initialized. Set enable_kv_cache_copy=True in __init__"
|
), "KV copy not initialized. Set enable_kv_cache_copy=True in __init__"
|
||||||
|
|
||||||
cfg = self._kv_copy_config
|
cfg = self._kv_copy_config
|
||||||
N_upper = next_power_of_2(N)
|
cap = int(cfg.get("num_locs_upper", 256))
|
||||||
grid = (self.data_ptrs.numel(), cfg["byte_tiles"])
|
grid = (self.data_ptrs.numel(), cfg["byte_tiles"])
|
||||||
|
|
||||||
copy_all_layer_kv_cache_tiled[grid](
|
if N <= cap:
|
||||||
self.data_ptrs,
|
upper = next_power_of_2(N)
|
||||||
self.data_strides,
|
copy_all_layer_kv_cache_tiled[grid](
|
||||||
tgt_loc,
|
self.data_ptrs,
|
||||||
src_loc,
|
self.data_strides,
|
||||||
N,
|
tgt_loc,
|
||||||
N_upper,
|
src_loc,
|
||||||
BYTES_PER_TILE=cfg["bytes_per_tile"],
|
N,
|
||||||
num_warps=cfg["num_warps"],
|
upper,
|
||||||
num_stages=2,
|
BYTES_PER_TILE=cfg["bytes_per_tile"],
|
||||||
)
|
num_warps=cfg["num_warps"],
|
||||||
|
num_stages=2,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
# Huge N: chunk, but each chunk's upper is still pow2(<= cap)
|
||||||
|
for start in range(0, N, cap):
|
||||||
|
end = min(start + cap, N)
|
||||||
|
chunk_len = end - start
|
||||||
|
upper = next_power_of_2(chunk_len)
|
||||||
|
copy_all_layer_kv_cache_tiled[grid](
|
||||||
|
self.data_ptrs,
|
||||||
|
self.data_strides,
|
||||||
|
tgt_loc[start:end],
|
||||||
|
src_loc[start:end],
|
||||||
|
chunk_len,
|
||||||
|
upper,
|
||||||
|
BYTES_PER_TILE=cfg["bytes_per_tile"],
|
||||||
|
num_warps=cfg["num_warps"],
|
||||||
|
num_stages=2,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class MHATokenToKVPoolFP4(MHATokenToKVPool):
|
class MHATokenToKVPoolFP4(MHATokenToKVPool):
|
||||||
|
|||||||
Reference in New Issue
Block a user