Clean up wrapper in flashinfer backend (#2638)
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user