Skip to content
Open
Show file tree
Hide file tree
Changes from 10 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 34 additions & 0 deletions include/sgl_flash_kernel_ops.h
Original file line number Diff line number Diff line change
Expand Up @@ -76,6 +76,40 @@ std::tuple<at::Tensor, at::Tensor, at::Tensor, at::Tensor> mha_fwd(
int const sm_margin,
std::optional<at::Tensor>& out_);

std::tuple<at::Tensor, at::Tensor, at::Tensor, at::Tensor> mha_fwd_appendkv(
const at::Tensor& q,
const at::Tensor& k,
const at::Tensor& v,
std::optional<const at::Tensor>& q_v_,
const at::Tensor& cu_seqlens_q,
const at::Tensor& cu_seqlens_k,
int max_seqlen_q,
int max_seqlen_k,
std::optional<const at::Tensor>& page_table,
std::optional<const at::Tensor>& kv_batch_idx_,
std::optional<const at::Tensor>& leftpad_k_,
std::optional<const at::Tensor>& rotary_cos_,
std::optional<const at::Tensor>& rotary_sin_,
std::optional<const at::Tensor>& seqlens_rotary_,
std::optional<at::Tensor>& q_descale_,
std::optional<at::Tensor>& k_descale_,
std::optional<at::Tensor>& v_descale_,
float const softmax_scale,
std::optional<const at::Tensor>& sinks,
bool is_causal,
int window_size_left,
int window_size_right,
float const softcap,
bool const is_rotary_interleaved,
std::optional<at::Tensor>& scheduler_metadata_,
int num_kv_splits,
std::optional<bool> pack_gqa_,
int const sm_margin,
std::optional<at::Tensor>& out_,
std::optional<const at::Tensor>& k_new_,
Comment thread
yuankuns marked this conversation as resolved.
std::optional<const at::Tensor>& v_new_,
std::optional<const at::Tensor>& cu_seqlens_k_new_);

void flash_mla_decode(
torch::Tensor& out,
const torch::Tensor& q_nope,
Expand Down
104 changes: 73 additions & 31 deletions python/sgl_kernel/flash_attn.py
Original file line number Diff line number Diff line change
Expand Up @@ -264,37 +264,79 @@ def flash_attn_with_kvcache(
if cache_seqlens is not None:
assert cache_seqlens.size(0) + 1 == cu_seqlens_q.size(0)
cu_seqlens_k = cache_seqlens
out, softmax_lse, *rest = torch.ops.sgl_kernel.fwd.default(
q,
k_cache,
v_cache,
qv,
cu_seqlens_q,
cu_seqlens_k,
max_seqlen_q,
1,
page_table,
cache_batch_idx,
cache_leftpad,
rotary_cos,
rotary_sin,
rotary_seqlens,
q_descale,
k_descale,
v_descale,
softmax_scale,
sinks,
causal,
window_size[0],
window_size[1],
softcap,
rotary_interleaved,
scheduler_metadata,
num_splits,
pack_gqa,
sm_margin,
out,
)
has_new_kv = k is not None or v is not None or cu_seqlens_k_new is not None
if has_new_kv:
native_max_seqlen_k = max_seqlen_k
if native_max_seqlen_k is None or native_max_seqlen_k == 0:
native_max_seqlen_k = (
k.shape[1] if k is not None and k.dim() == 4 else (max_seqlen_q or 1)
)
out, softmax_lse, *rest = torch.ops.sgl_kernel.fwd_appendkv.default(
q,
k_cache,
v_cache,
qv,
cu_seqlens_q,
cu_seqlens_k,
max_seqlen_q,
native_max_seqlen_k,
page_table,
cache_batch_idx,
cache_leftpad,
rotary_cos,
rotary_sin,
rotary_seqlens,
q_descale,
k_descale,
v_descale,
softmax_scale,
sinks,
causal,
window_size[0],
window_size[1],
softcap,
rotary_interleaved,
scheduler_metadata,
num_splits,
pack_gqa,
sm_margin,
out,
k,
v,
cu_seqlens_k_new,
)
else:
out, softmax_lse, *rest = torch.ops.sgl_kernel.fwd.default(
q,
k_cache,
v_cache,
qv,
cu_seqlens_q,
cu_seqlens_k,
max_seqlen_q,
1,
page_table,
cache_batch_idx,
cache_leftpad,
rotary_cos,
rotary_sin,
rotary_seqlens,
q_descale,
k_descale,
v_descale,
softmax_scale,
sinks,
causal,
window_size[0],
window_size[1],
softcap,
rotary_interleaved,
scheduler_metadata,
num_splits,
pack_gqa,
sm_margin,
out,
)
return (out, softmax_lse, *rest) if return_softmax_lse else out


Expand Down
87 changes: 79 additions & 8 deletions src/FMHAPrefillXe20.cmake
Original file line number Diff line number Diff line change
Expand Up @@ -31,19 +31,37 @@ set(FMHA_PREFILL_NUM_SG_64 8)

set(FMHA_PREFILL_TILED_Q_96 128)
set(FMHA_PREFILL_TILED_KV_96 64)
set(FMHA_PREFILL_NUM_SG_96 8)
set(FMHA_PREFILL_NUM_SG_96 16)

set(FMHA_PREFILL_TILED_Q_128 256)
set(FMHA_PREFILL_TILED_KV_128 32)
set(FMHA_PREFILL_TILED_Q_128 128)
set(FMHA_PREFILL_TILED_KV_128 64)
set(FMHA_PREFILL_NUM_SG_128 16)
option(FMHA_PREFILL_HD128_LARGE_TILE "Enable q256/k32 paged head_dim=128 path for model-sized Q lengths" ON)
set(FMHA_PREFILL_HD128_LARGE_TILE_MIN_Q 256)
set(FMHA_PREFILL_HD128_LARGE_TILE_Q 256)
set(FMHA_PREFILL_HD128_LARGE_TILE_KV 32)
set(FMHA_PREFILL_HD128_LARGE_NUM_SG 16)

set(FMHA_PREFILL_TILED_Q_192 256)
set(FMHA_PREFILL_TILED_Q_192 128)
set(FMHA_PREFILL_TILED_KV_192 64)
set(FMHA_PREFILL_NUM_SG_192 32)
set(FMHA_PREFILL_NUM_SG_192 16)

set(FMHA_PREFILL_TILED_Q_256 256)
set(FMHA_PREFILL_TILED_Q_256 128)
set(FMHA_PREFILL_TILED_KV_256 64)
set(FMHA_PREFILL_NUM_SG_256 32)
set(FMHA_PREFILL_NUM_SG_256 16)

# A q256 tile reduces duplicate fused KV-cache stores for uniform 256-token
# appends while retaining the tuned default tiles for shorter and ragged Q.
set(FMHA_PREFILL_APPENDKV_Q256_HEAD_DIMS 64 192 256)
set(FMHA_PREFILL_APPENDKV_Q256_Q_64 256)
set(FMHA_PREFILL_APPENDKV_Q256_KV_64 64)
set(FMHA_PREFILL_APPENDKV_Q256_NUM_SG_64 16)
set(FMHA_PREFILL_APPENDKV_Q256_Q_192 256)
set(FMHA_PREFILL_APPENDKV_Q256_KV_192 64)
set(FMHA_PREFILL_APPENDKV_Q256_NUM_SG_192 32)
set(FMHA_PREFILL_APPENDKV_Q256_Q_256 256)
set(FMHA_PREFILL_APPENDKV_Q256_KV_256 64)
set(FMHA_PREFILL_APPENDKV_Q256_NUM_SG_256 32)

set(FMHA_PREFILL_TILED_Q_512 256)
set(FMHA_PREFILL_TILED_KV_512 64)
Expand All @@ -66,9 +84,19 @@ set(FMHA_PREFILL_TILED_Q_NP_80 256)
set(FMHA_PREFILL_TILED_KV_NP_80 64)
set(FMHA_PREFILL_NUM_SG_NP_80 16)

set(FMHA_PREFILL_TILED_Q_NP_96 256)
set(FMHA_PREFILL_TILED_Q_NP_96 128)
set(FMHA_PREFILL_TILED_KV_NP_96 64)
set(FMHA_PREFILL_NUM_SG_NP_96 16)
option(FMHA_PREFILL_HD96_NP_SMALL_TILE "Enable q32 non-paged head_dim=96 path for short queries" ON)
set(FMHA_PREFILL_HD96_NP_SMALL_TILE_MAX_Q 32)
set(FMHA_PREFILL_HD96_NP_SMALL_TILE_Q 32)
set(FMHA_PREFILL_HD96_NP_SMALL_TILE_KV 64)
set(FMHA_PREFILL_HD96_NP_SMALL_NUM_SG 4)
option(FMHA_PREFILL_HD96_NP_LARGE_TILE "Enable q256 non-paged head_dim=96 path for long queries" ON)
set(FMHA_PREFILL_HD96_NP_LARGE_TILE_MIN_Q 257)
set(FMHA_PREFILL_HD96_NP_LARGE_TILE_Q 256)
set(FMHA_PREFILL_HD96_NP_LARGE_TILE_KV 64)
set(FMHA_PREFILL_HD96_NP_LARGE_NUM_SG 16)

set(FMHA_PREFILL_TILED_Q_NP_128 256)
set(FMHA_PREFILL_TILED_KV_NP_128 32)
Expand Down Expand Up @@ -96,6 +124,29 @@ foreach(HEAD_DIM ${FMHA_PREFILL_PAGED_HEAD_DIMS})
set(TILED_OUT ${FMHA_PREFILL_TILED_OUT_${HEAD_DIM}})
endif()

if(HEAD_DIM STREQUAL "128" AND FMHA_PREFILL_HD128_LARGE_TILE)
set(HD128_PAGED_LARGE_TILE 1)
else()
set(HD128_PAGED_LARGE_TILE 0)
endif()
set(HD128_PAGED_LARGE_TILE_MIN_Q ${FMHA_PREFILL_HD128_LARGE_TILE_MIN_Q})
set(HD128_PAGED_LARGE_TILE_Q ${FMHA_PREFILL_HD128_LARGE_TILE_Q})
set(HD128_PAGED_LARGE_TILE_KV ${FMHA_PREFILL_HD128_LARGE_TILE_KV})
set(HD128_PAGED_LARGE_NUM_SG ${FMHA_PREFILL_HD128_LARGE_NUM_SG})

list(FIND FMHA_PREFILL_APPENDKV_Q256_HEAD_DIMS ${HEAD_DIM} APPENDKV_Q256_INDEX)
if(APPENDKV_Q256_INDEX GREATER -1)
set(APPENDKV_Q256_TILE 1)
set(APPENDKV_Q256_Q ${FMHA_PREFILL_APPENDKV_Q256_Q_${HEAD_DIM}})
set(APPENDKV_Q256_KV ${FMHA_PREFILL_APPENDKV_Q256_KV_${HEAD_DIM}})
set(APPENDKV_Q256_NUM_SG ${FMHA_PREFILL_APPENDKV_Q256_NUM_SG_${HEAD_DIM}})
else()
set(APPENDKV_Q256_TILE 0)
set(APPENDKV_Q256_Q 256)
set(APPENDKV_Q256_KV 64)
set(APPENDKV_Q256_NUM_SG 16)
endif()

set(GENERATED_FILE
"${CMAKE_CURRENT_BINARY_DIR}/sycl/xe_fmha_fwd_prefill_paged_kernel_${HEAD_DIM}.cpp")
configure_file(${FMHA_PREFILL_TEMPLATE} ${GENERATED_FILE} @ONLY)
Expand All @@ -116,6 +167,26 @@ foreach(HEAD_DIM ${FMHA_PREFILL_NP_HEAD_DIMS})
message(FATAL_ERROR "Missing non-paged tile params for prefill HEAD_DIM=${HEAD_DIM}")
endif()

if(HEAD_DIM STREQUAL "96" AND FMHA_PREFILL_HD96_NP_SMALL_TILE)
set(HD96_NP_SMALL_TILE 1)
else()
set(HD96_NP_SMALL_TILE 0)
endif()
set(HD96_NP_SMALL_TILE_MAX_Q ${FMHA_PREFILL_HD96_NP_SMALL_TILE_MAX_Q})
set(HD96_NP_SMALL_TILE_Q ${FMHA_PREFILL_HD96_NP_SMALL_TILE_Q})
set(HD96_NP_SMALL_TILE_KV ${FMHA_PREFILL_HD96_NP_SMALL_TILE_KV})
set(HD96_NP_SMALL_NUM_SG ${FMHA_PREFILL_HD96_NP_SMALL_NUM_SG})

if(HEAD_DIM STREQUAL "96" AND FMHA_PREFILL_HD96_NP_LARGE_TILE)
set(HD96_NP_LARGE_TILE 1)
else()
set(HD96_NP_LARGE_TILE 0)
endif()
set(HD96_NP_LARGE_TILE_MIN_Q ${FMHA_PREFILL_HD96_NP_LARGE_TILE_MIN_Q})
set(HD96_NP_LARGE_TILE_Q ${FMHA_PREFILL_HD96_NP_LARGE_TILE_Q})
set(HD96_NP_LARGE_TILE_KV ${FMHA_PREFILL_HD96_NP_LARGE_TILE_KV})
set(HD96_NP_LARGE_NUM_SG ${FMHA_PREFILL_HD96_NP_LARGE_NUM_SG})

set(GENERATED_NP_FILE
"${CMAKE_CURRENT_BINARY_DIR}/sycl/xe_fmha_fwd_prefill_nopage_kernel_${HEAD_DIM}.cpp")
configure_file(${FMHA_PREFILL_NOPAGE_TEMPLATE} ${GENERATED_NP_FILE} @ONLY)
Expand Down
Loading
Loading