Sliding-window flash attention kernels split from attention_kernels.c. More...
Go to the source code of this file.
Functions | |
| static void | attention_flash_query_sliding (const float *q_vec, const float *k_head, const float *v_head, int query_pos, int kv_tokens, int head_dim, int aligned_head_dim, float scale, float *out_vec, int sliding_window) |
| void | attention_forward_causal_head_major_gqa_flash_strided_sliding (const float *q, const float *k, const float *v, float *output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, int sliding_window) |
| void | attention_forward_causal_head_major_gqa_flash_strided_sliding_gemma4 (const float *q, const float *k, const float *v, float *output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, int sliding_window) |
| static void | attention_forward_causal_head_major_gqa_flash_strided_sliding_gemma4_impl (const float *q, const float *k, const float *v, float *output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, int sliding_window, int output_token_major) |
| void | attention_forward_causal_head_major_gqa_flash_strided_sliding_gemma4_token_output (const float *q, const float *k, const float *v, float *output, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, int sliding_window) |
| void | attention_forward_causal_head_major_shared_kv_sliding_gemma4 (const float *q, float *output, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int kv_stride_tokens, int sliding_window) |
| void | attention_forward_decode_head_major_gqa_flash_sliding (const float *q_token, const float *k_cache, const float *v_cache, float *out_token, int num_heads, int num_kv_heads, int kv_tokens, int cache_capacity, int head_dim, int aligned_head_dim, int sliding_window) |
| void | attention_forward_decode_head_major_gqa_flash_sliding_gemma4 (const float *q_token, const float *k_cache, const float *v_cache, float *out_token, int num_heads, int num_kv_heads, int kv_tokens, int cache_capacity, int head_dim, int aligned_head_dim, int sliding_window) |
| void | attention_forward_decode_head_major_shared_kv_sliding_gemma4 (const float *q_token, const float *k_cache, const float *v_cache, float *out_token, int num_heads, int kv_tokens, int cache_capacity, int head_dim, int aligned_head_dim, int sliding_window) |
| static size_t | attention_output_index (int h, int t, int num_heads, int num_tokens, int aligned_head_dim, int output_token_major) |
| static int | ck_env_int_default (const char *name, int fallback) |
| static void | ck_sliding_attention_compute_one (const ck_sliding_attention_args_t *a, int job) |
| static int | ck_sliding_attention_parallel_disabled (void) |
| static int | ck_sliding_attention_pick_threads (ck_threadpool_t *pool, int total_jobs, int num_tokens, int head_dim) |
| static void | ck_sliding_attention_work_fn (int ith, int nth, void *args) |
| static size_t | qkv_index (int h, int t, int d, int num_tokens, int aligned_head_dim) |
Sliding-window flash attention kernels split from attention_kernels.c.
Definition in file attention_kernels_sliding.c.
| #define CK_SLIDING_FLASH_IMPL attention_flash_query_sliding |
| #define SLIDING_DECODE_IMPL attention_flash_query_sliding |
| #define SLIDING_DECODE_IMPL_GEMMA4 attention_flash_query_sliding |
| #define SLIDING_FLASH_IMPL attention_flash_query_sliding |
| #define SLIDING_FLASH_IMPL_GEMMA4 attention_flash_query_sliding |
|
static |
Definition at line 392 of file attention_kernels_sliding.c.
References score.
| void attention_forward_causal_head_major_gqa_flash_strided_sliding | ( | const float * | q, |
| const float * | k, | ||
| const float * | v, | ||
| float * | output, | ||
| int | num_heads, | ||
| int | num_kv_heads, | ||
| int | num_tokens, | ||
| int | head_dim, | ||
| int | aligned_head_dim, | ||
| int | kv_stride_tokens, | ||
| int | sliding_window | ||
| ) |
Flash attention forward with sliding window (prefill)
Sliding-window attention for prefill: each token attends to the last W tokens. When sliding_window <= 0, behaves like regular causal attention.
After changes: make test
Definition at line 566 of file attention_kernels_sliding.c.
References attention_forward_causal_head_major_gqa_flash_strided(), attention_output_index(), ck_sliding_attention_pick_threads(), ck_sliding_attention_work_fn(), ck_threadpool_dispatch_n(), ck_threadpool_global(), qkv_index(), and SLIDING_FLASH_IMPL.
| void attention_forward_causal_head_major_gqa_flash_strided_sliding_gemma4 | ( | const float * | q, |
| const float * | k, | ||
| const float * | v, | ||
| float * | output, | ||
| int | num_heads, | ||
| int | num_kv_heads, | ||
| int | num_tokens, | ||
| int | head_dim, | ||
| int | aligned_head_dim, | ||
| int | kv_stride_tokens, | ||
| int | sliding_window | ||
| ) |
Definition at line 763 of file attention_kernels_sliding.c.
References attention_forward_causal_head_major_gqa_flash_strided_sliding_gemma4_impl().
Referenced by attention_forward_causal_head_major_shared_kv_sliding_gemma4().
|
static |
Flash attention decode with sliding window
Single query token attends to the last W tokens in the KV cache. For decode: effective_kv_tokens = min(kv_tokens, sliding_window)
After changes: make test
Definition at line 668 of file attention_kernels_sliding.c.
References attention_forward_causal_head_major_gqa_flash_strided_gemma4(), attention_forward_causal_head_major_gqa_flash_strided_gemma4_token_output(), attention_output_index(), ck_sliding_attention_pick_threads(), ck_sliding_attention_work_fn(), ck_threadpool_dispatch_n(), ck_threadpool_global(), qkv_index(), and SLIDING_FLASH_IMPL_GEMMA4.
Referenced by attention_forward_causal_head_major_gqa_flash_strided_sliding_gemma4(), and attention_forward_causal_head_major_gqa_flash_strided_sliding_gemma4_token_output().
| void attention_forward_causal_head_major_gqa_flash_strided_sliding_gemma4_token_output | ( | const float * | q, |
| const float * | k, | ||
| const float * | v, | ||
| float * | output, | ||
| int | num_heads, | ||
| int | num_kv_heads, | ||
| int | num_tokens, | ||
| int | head_dim, | ||
| int | aligned_head_dim, | ||
| int | kv_stride_tokens, | ||
| int | sliding_window | ||
| ) |
Definition at line 782 of file attention_kernels_sliding.c.
References attention_forward_causal_head_major_gqa_flash_strided_sliding_gemma4_impl().
| void attention_forward_causal_head_major_shared_kv_sliding_gemma4 | ( | const float * | q, |
| float * | output, | ||
| int | num_heads, | ||
| int | num_tokens, | ||
| int | head_dim, | ||
| int | aligned_head_dim, | ||
| int | kv_stride_tokens, | ||
| int | sliding_window | ||
| ) |
Definition at line 801 of file attention_kernels_sliding.c.
References attention_forward_causal_head_major_gqa_flash_strided_sliding_gemma4().
| void attention_forward_decode_head_major_gqa_flash_sliding | ( | const float * | q_token, |
| const float * | k_cache, | ||
| const float * | v_cache, | ||
| float * | out_token, | ||
| int | num_heads, | ||
| int | num_kv_heads, | ||
| int | kv_tokens, | ||
| int | cache_capacity, | ||
| int | head_dim, | ||
| int | aligned_head_dim, | ||
| int | sliding_window | ||
| ) |
Definition at line 817 of file attention_kernels_sliding.c.
References attention_forward_decode_head_major_gqa_flash(), and SLIDING_DECODE_IMPL.
| void attention_forward_decode_head_major_gqa_flash_sliding_gemma4 | ( | const float * | q_token, |
| const float * | k_cache, | ||
| const float * | v_cache, | ||
| float * | out_token, | ||
| int | num_heads, | ||
| int | num_kv_heads, | ||
| int | kv_tokens, | ||
| int | cache_capacity, | ||
| int | head_dim, | ||
| int | aligned_head_dim, | ||
| int | sliding_window | ||
| ) |
Definition at line 903 of file attention_kernels_sliding.c.
References attention_forward_decode_head_major_gqa_flash_gemma4(), and SLIDING_DECODE_IMPL_GEMMA4.
Referenced by attention_forward_decode_head_major_shared_kv_sliding_gemma4().
| void attention_forward_decode_head_major_shared_kv_sliding_gemma4 | ( | const float * | q_token, |
| const float * | k_cache, | ||
| const float * | v_cache, | ||
| float * | out_token, | ||
| int | num_heads, | ||
| int | kv_tokens, | ||
| int | cache_capacity, | ||
| int | head_dim, | ||
| int | aligned_head_dim, | ||
| int | sliding_window | ||
| ) |
Definition at line 976 of file attention_kernels_sliding.c.
References attention_forward_decode_head_major_gqa_flash_sliding_gemma4().
|
inlinestatic |
Definition at line 26 of file attention_kernels_sliding.c.
References qkv_index().
Referenced by attention_forward_causal_head_major_gqa_flash_strided_sliding(), attention_forward_causal_head_major_gqa_flash_strided_sliding_gemma4_impl(), and ck_sliding_attention_compute_one().
|
static |
Definition at line 472 of file attention_kernels_sliding.c.
References end.
Referenced by ck_sliding_attention_pick_threads().
|
static |
Definition at line 509 of file attention_kernels_sliding.c.
References attention_output_index(), CK_SLIDING_FLASH_IMPL, and qkv_index().
Referenced by ck_sliding_attention_work_fn().
|
static |
Definition at line 484 of file attention_kernels_sliding.c.
Referenced by ck_sliding_attention_pick_threads().
|
static |
Definition at line 490 of file attention_kernels_sliding.c.
References ck_env_int_default(), ck_sliding_attention_parallel_disabled(), ck_threadpool_n_threads(), and ck_threadpool_thread_id().
Referenced by attention_forward_causal_head_major_gqa_flash_strided_sliding(), and attention_forward_causal_head_major_gqa_flash_strided_sliding_gemma4_impl().
|
static |
Definition at line 547 of file attention_kernels_sliding.c.
References ck_sliding_attention_compute_one().
Referenced by attention_forward_causal_head_major_gqa_flash_strided_sliding(), and attention_forward_causal_head_major_gqa_flash_strided_sliding_gemma4_impl().
|
inlinestatic |
Definition at line 16 of file attention_kernels_sliding.c.
Referenced by attention_forward_causal_head_major_gqa_flash_strided_sliding(), attention_forward_causal_head_major_gqa_flash_strided_sliding_gemma4_impl(), attention_output_index(), and ck_sliding_attention_compute_one().