| import torch |
| from torch import nn |
| import triton |
| import triton.language as tl |
|
|
| from jetengine_ext.utils.context import get_context |
| from jetengine_ext.engine.sequence import RunType |
| from jetengine_ext.kernels.triton.attention import sparse_attn_varlen |
| |
| |
| from flash_attn import flash_attn_with_kvcache |
|
|
|
|
| @triton.jit |
| def store_kvcache_kernel( |
| key_ptr, |
| key_stride, |
| value_ptr, |
| value_stride, |
| k_cache_ptr, |
| v_cache_ptr, |
| slot_mapping_ptr, |
| D: tl.constexpr, |
| ): |
| idx = tl.program_id(0) |
| key_offsets = idx * key_stride + tl.arange(0, D) |
| value_offsets = idx * value_stride + tl.arange(0, D) |
| key = tl.load(key_ptr + key_offsets) |
| value = tl.load(value_ptr + value_offsets) |
| slot = tl.load(slot_mapping_ptr + idx) |
| cache_offsets = slot * D + tl.arange(0, D) |
| tl.store(k_cache_ptr + cache_offsets, key) |
| tl.store(v_cache_ptr + cache_offsets, value) |
|
|
|
|
| def store_kvcache(key: torch.Tensor, value: torch.Tensor, k_cache: torch.Tensor, v_cache: torch.Tensor, slot_mapping: torch.Tensor): |
| N, num_heads, head_dim = key.shape |
| D = num_heads * head_dim |
| assert key.stride(-1) == 1 and value.stride(-1) == 1 |
| assert key.stride(1) == head_dim and value.stride(1) == head_dim |
| assert k_cache.stride(1) == D and v_cache.stride(1) == D |
| assert slot_mapping.numel() == N |
| store_kvcache_kernel[(N,)](key, key.stride(0), value, value.stride(0), k_cache, v_cache, slot_mapping, D) |
|
|
|
|
| class Attention(nn.Module): |
|
|
| def __init__( |
| self, |
| num_heads, |
| head_dim, |
| scale, |
| num_kv_heads, |
| ): |
| super().__init__() |
| self.num_heads = num_heads |
| self.head_dim = head_dim |
| self.scale = scale |
| self.num_kv_heads = num_kv_heads |
| self.k_cache = self.v_cache = torch.tensor([]) |
|
|
| def forward(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor): |
| pass |
|
|
| class BlockAttention(Attention): |
| def __init__( |
| self, |
| num_heads, |
| head_dim, |
| scale, |
| num_kv_heads, |
| ): |
| super().__init__(num_heads, head_dim, scale, num_kv_heads) |
| |
| def forward(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor): |
| o: torch.Tensor |
| q = q.view(-1, self.num_heads, self.head_dim) |
| k = k.view(-1, self.num_kv_heads, self.head_dim) |
| v = v.view(-1, self.num_kv_heads, self.head_dim) |
| context = get_context() |
| k_cache, v_cache = self.k_cache, self.v_cache |
|
|
| should_store_whole = (context.run_type == RunType.PREFILL) |
| if should_store_whole and k_cache.numel() and v_cache.numel(): |
| store_kvcache(k, v, k_cache, v_cache, context.slot_mapping) |
| |
| if context.run_type == RunType.PREFILL: |
| o = sparse_attn_varlen(q, k, v, |
| cu_seqlens_q=context.cu_seqlens_q, |
| cu_seqlens_k=context.cu_seqlens_k, |
| staircase_size=context.block_length) |
| else: |
| q = q.view(-1, context.block_length, self.num_heads, self.head_dim) |
| k = k.view(-1, context.block_length, self.num_kv_heads, self.head_dim) |
| v = v.view(-1, context.block_length, self.num_kv_heads, self.head_dim) |
| o = flash_attn_with_kvcache(q, k_cache=k_cache, v_cache=v_cache, k=k, v=v, |
| cache_seqlens=context.context_lens, |
| block_table=context.block_tables, |
| causal=False) |
| o = o.view(-1, self.num_heads * self.head_dim) |
| return o |
|
|
| |
|
|