Clean up wrapper in flashinfer backend (#2638)

This commit is contained in:
Lianmin Zheng
2024-12-29 00:45:57 -08:00
committed by GitHub
parent fd34f2da35
commit 3815b23ccb
12 changed files with 197 additions and 94 deletions
@@ -2,7 +2,6 @@ from abc import ABC, abstractmethod
from typing import Optional
import torch
from torch import nn
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
@@ -8,8 +8,9 @@ Each backend supports two operators: extend (i.e. prefill with cached prefix) an
"""
import os
from dataclasses import dataclass
from enum import Enum, auto
from typing import TYPE_CHECKING, List
from typing import TYPE_CHECKING, List, Union
import torch
import triton
@@ -38,12 +39,25 @@ class WrapperDispatch(Enum):
CROSS_ATTENTION = auto()
@dataclass
class DecodeMetadata:
decode_wrappers: List[BatchDecodeWithPagedKVCacheWrapper]
@dataclass
class PrefillMetadata:
prefill_wrappers: List[BatchPrefillWithPagedKVCacheWrapper]
use_ragged: bool
extend_no_prefix: bool
class FlashInferAttnBackend(AttentionBackend):
"""Flashinfer attention kernels."""
def __init__(self, model_runner: ModelRunner):
super().__init__()
# Parse constants
self.decode_use_tensor_cores = should_use_tensor_core(
kv_cache_dtype=model_runner.kv_cache_dtype,
num_attention_heads=model_runner.model_config.num_attention_heads
@@ -52,7 +66,6 @@ class FlashInferAttnBackend(AttentionBackend):
model_runner.tp_size
),
)
self.max_context_len = model_runner.model_config.context_len
assert not (
@@ -120,8 +133,8 @@ class FlashInferAttnBackend(AttentionBackend):
)
# Other metadata
self.forward_metadata = None
self.cuda_graph_metadata = {}
self.forward_metadata: Union[PrefillMetadata, DecodeMetadata] = None
self.decode_cuda_graph_metadata = {}
def init_forward_metadata(self, forward_batch: ForwardBatch):
if forward_batch.forward_mode.is_decode():
@@ -129,10 +142,10 @@ class FlashInferAttnBackend(AttentionBackend):
forward_batch.req_pool_indices,
forward_batch.seq_lens,
forward_batch.seq_lens_sum,
decode_wrappers=None,
decode_wrappers=self.decode_wrappers,
encoder_lens=forward_batch.encoder_lens,
)
self.forward_metadata = (self.decode_wrappers,)
self.forward_metadata = DecodeMetadata(self.decode_wrappers)
else:
prefix_lens = forward_batch.extend_prefix_lens
@@ -149,11 +162,13 @@ class FlashInferAttnBackend(AttentionBackend):
forward_batch.seq_lens,
forward_batch.seq_lens_sum,
prefix_lens,
prefill_wrappers=self.prefill_wrappers_paged,
use_ragged=use_ragged,
encoder_lens=forward_batch.encoder_lens,
)
self.forward_metadata = (use_ragged, extend_no_prefix)
self.forward_metadata = PrefillMetadata(
self.prefill_wrappers_paged, use_ragged, extend_no_prefix
)
def init_cuda_graph_state(self, max_bs: int):
cuda_graph_kv_indices = torch.zeros(
@@ -194,8 +209,8 @@ class FlashInferAttnBackend(AttentionBackend):
decode_wrappers=decode_wrappers,
encoder_lens=encoder_lens,
)
self.cuda_graph_metadata[bs] = decode_wrappers
self.forward_metadata = (decode_wrappers,)
self.decode_cuda_graph_metadata[bs] = decode_wrappers
self.forward_metadata = DecodeMetadata(decode_wrappers)
def init_forward_metadata_replay_cuda_graph(
self,
@@ -209,7 +224,7 @@ class FlashInferAttnBackend(AttentionBackend):
req_pool_indices[:bs],
seq_lens[:bs],
seq_lens_sum,
decode_wrappers=self.cuda_graph_metadata[bs],
decode_wrappers=self.decode_cuda_graph_metadata[bs],
encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None,
)
@@ -225,18 +240,16 @@ class FlashInferAttnBackend(AttentionBackend):
forward_batch: ForwardBatch,
save_kv_cache=True,
):
prefill_wrapper_paged = self.prefill_wrappers_paged[
prefill_wrapper_paged = self.forward_metadata.prefill_wrappers[
self._get_wrapper_idx(layer)
]
use_ragged, extend_no_prefix = self.forward_metadata
cache_loc = (
forward_batch.out_cache_loc
if not layer.is_cross_attention
else forward_batch.encoder_out_cache_loc
)
if not use_ragged:
if not self.forward_metadata.use_ragged:
if k is not None:
assert v is not None
if save_kv_cache:
@@ -260,7 +273,7 @@ class FlashInferAttnBackend(AttentionBackend):
logits_soft_cap=layer.logit_cap,
)
if extend_no_prefix:
if self.forward_metadata.extend_no_prefix:
o = o1
else:
o2, s2 = prefill_wrapper_paged.forward_return_lse(
@@ -287,7 +300,9 @@ class FlashInferAttnBackend(AttentionBackend):
forward_batch: ForwardBatch,
save_kv_cache=True,
):
decode_wrapper = self.forward_metadata[0][self._get_wrapper_idx(layer)]
decode_wrapper = self.forward_metadata.decode_wrappers[
self._get_wrapper_idx(layer)
]
cache_loc = (
forward_batch.out_cache_loc
if not layer.is_cross_attention
@@ -322,7 +337,7 @@ class FlashInferAttnBackend(AttentionBackend):
class FlashInferIndicesUpdaterDecode:
def __init__(self, model_runner: ModelRunner, attn_backend: AttentionBackend):
# Constants
# Parse Constants
self.num_qo_heads = (
model_runner.model_config.num_attention_heads // model_runner.tp_size
)
@@ -340,9 +355,8 @@ class FlashInferIndicesUpdaterDecode:
self.kv_indptr = attn_backend.kv_indptr
self.kv_last_page_len = attn_backend.kv_last_page_len
self.req_to_token = model_runner.req_to_token_pool.req_to_token
self.decode_wrappers = attn_backend.decode_wrappers
# Dispatch
# Dispatch the update function
if self.attn_backend.dispatch_reason == WrapperDispatch.SLIDING_WINDOW:
self.update = self.update_sliding_window
elif self.attn_backend.dispatch_reason == WrapperDispatch.CROSS_ATTENTION:
@@ -356,7 +370,7 @@ class FlashInferIndicesUpdaterDecode:
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
decode_wrappers: List,
decode_wrappers: List[BatchDecodeWithPagedKVCacheWrapper],
encoder_lens: torch.Tensor,
):
# Keep the signature for type checking. It will be assigned during runtime.
@@ -367,7 +381,7 @@ class FlashInferIndicesUpdaterDecode:
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
decode_wrappers: List,
decode_wrappers: List[BatchDecodeWithPagedKVCacheWrapper],
encoder_lens: torch.Tensor,
):
decode_wrappers = decode_wrappers or self.decode_wrappers
@@ -385,11 +399,9 @@ class FlashInferIndicesUpdaterDecode:
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
decode_wrappers: List,
decode_wrappers: List[BatchDecodeWithPagedKVCacheWrapper],
encoder_lens: torch.Tensor,
):
decode_wrappers = decode_wrappers or self.decode_wrappers
for wrapper_id in range(2):
if wrapper_id == 0:
# Sliding window attention
@@ -419,11 +431,9 @@ class FlashInferIndicesUpdaterDecode:
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
decode_wrappers: List,
decode_wrappers: List[BatchDecodeWithPagedKVCacheWrapper],
encoder_lens: torch.Tensor,
):
decode_wrappers = decode_wrappers or self.decode_wrappers
for wrapper_id in range(2):
if wrapper_id == 0:
# Normal attention
@@ -446,7 +456,7 @@ class FlashInferIndicesUpdaterDecode:
def call_begin_forward(
self,
wrapper,
wrapper: BatchDecodeWithPagedKVCacheWrapper,
req_pool_indices: torch.Tensor,
paged_kernel_lens: torch.Tensor,
paged_kernel_lens_sum: int,
@@ -486,7 +496,7 @@ class FlashInferIndicesUpdaterDecode:
class FlashInferIndicesUpdaterPrefill:
def __init__(self, model_runner: ModelRunner, attn_backend: AttentionBackend):
# Constants
# Parse Constants
self.num_qo_heads = (
model_runner.model_config.num_attention_heads // model_runner.tp_size
)
@@ -505,10 +515,9 @@ class FlashInferIndicesUpdaterPrefill:
self.kv_last_page_len = attn_backend.kv_last_page_len
self.qo_indptr = attn_backend.qo_indptr
self.req_to_token = model_runner.req_to_token_pool.req_to_token
self.wrapper_ragged = attn_backend.prefill_wrapper_ragged
self.wrappers_paged = attn_backend.prefill_wrappers_paged
self.prefill_wrapper_ragged = attn_backend.prefill_wrapper_ragged
# Dispatch
# Dispatch the update function
if self.attn_backend.dispatch_reason == WrapperDispatch.SLIDING_WINDOW:
self.update = self.update_sliding_window
elif self.attn_backend.dispatch_reason == WrapperDispatch.CROSS_ATTENTION:
@@ -523,6 +532,7 @@ class FlashInferIndicesUpdaterPrefill:
seq_lens: torch.Tensor,
seq_lens_sum: int,
prefix_lens: torch.Tensor,
prefill_wrappers: List[BatchPrefillWithPagedKVCacheWrapper],
use_ragged: bool,
encoder_lens: torch.Tensor,
):
@@ -535,6 +545,7 @@ class FlashInferIndicesUpdaterPrefill:
seq_lens: torch.Tensor,
seq_lens_sum: int,
prefix_lens: torch.Tensor,
prefill_wrappers: List[BatchPrefillWithPagedKVCacheWrapper],
use_ragged: bool,
encoder_lens: torch.Tensor,
):
@@ -546,8 +557,8 @@ class FlashInferIndicesUpdaterPrefill:
paged_kernel_lens_sum = seq_lens_sum
self.call_begin_forward(
self.wrapper_ragged,
self.wrappers_paged[0],
self.prefill_wrapper_ragged,
prefill_wrappers[0],
req_pool_indices,
paged_kernel_lens,
paged_kernel_lens_sum,
@@ -565,6 +576,7 @@ class FlashInferIndicesUpdaterPrefill:
seq_lens: torch.Tensor,
seq_lens_sum: int,
prefix_lens: torch.Tensor,
prefill_wrappers: List[BatchPrefillWithPagedKVCacheWrapper],
use_ragged: bool,
encoder_lens: torch.Tensor,
):
@@ -584,8 +596,8 @@ class FlashInferIndicesUpdaterPrefill:
kv_start_idx = seq_lens - paged_kernel_lens
self.call_begin_forward(
self.wrapper_ragged,
self.wrappers_paged[wrapper_id],
self.prefill_wrapper_ragged,
prefill_wrappers[wrapper_id],
req_pool_indices,
paged_kernel_lens,
paged_kernel_lens_sum,
@@ -603,6 +615,7 @@ class FlashInferIndicesUpdaterPrefill:
seq_lens: torch.Tensor,
seq_lens_sum: int,
prefix_lens: torch.Tensor,
prefill_wrappers: List[BatchPrefillWithPagedKVCacheWrapper],
use_ragged: bool,
encoder_lens: torch.Tensor,
):
@@ -619,8 +632,8 @@ class FlashInferIndicesUpdaterPrefill:
paged_kernel_lens_sum = paged_kernel_lens.sum().item()
self.call_begin_forward(
self.wrapper_ragged,
self.wrappers_paged[wrapper_id],
self.prefill_wrapper_ragged,
prefill_wrappers[wrapper_id],
req_pool_indices,
paged_kernel_lens,
paged_kernel_lens_sum,
@@ -634,8 +647,8 @@ class FlashInferIndicesUpdaterPrefill:
def call_begin_forward(
self,
wrapper_ragged,
wrapper_paged,
wrapper_ragged: BatchPrefillWithRaggedKVCacheWrapper,
wrapper_paged: BatchPrefillWithPagedKVCacheWrapper,
req_pool_indices: torch.Tensor,
paged_kernel_lens: torch.Tensor,
paged_kernel_lens_sum: int,