Skip to content
Open
Changes from all 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
48 changes: 22 additions & 26 deletions python/minisgl/attention/fi.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@
from minisgl.core import Batch, get_global_ctx
from minisgl.distributed import get_tp_info
from minisgl.env import ENV
from minisgl.utils import div_even, init_logger
from minisgl.utils import div_even, init_logger, div_ceil

from .base import BaseAttnBackend, BaseAttnMetadata
from .utils import BaseCaptureData
Expand Down Expand Up @@ -54,7 +54,7 @@ class FIMetadata(BaseAttnMetadata):
num_qo_heads: int
num_kv_heads: int
head_dim: int
page_size: Literal[1] # currently only support page_size=1
page_size: int
pos_encoding_mode: str
seq_lens_cpu: torch.Tensor # on cpu
dtype: torch.dtype
Expand All @@ -63,7 +63,6 @@ class FIMetadata(BaseAttnMetadata):
# fmt: on

def __post_init__(self) -> None:
assert self.page_size == 1, "Currently only page_size=1 is supported."
assert (
self.cu_seqlens_k_cpu.is_cpu
and self.cu_seqlens_q_cpu.is_cpu
Expand All @@ -85,8 +84,10 @@ def __init__(self, config: ModelConfig) -> None:
)

self.config = config
self.kvcache = get_global_ctx().kv_cache
ctx = get_global_ctx()
self.kvcache = ctx.kv_cache
self.device = self.kvcache.device
self.page_size = ctx.page_size
self.float_workspace_buffer = torch.empty(
128 * 1024 * 1024, dtype=torch.uint8, device=self.device
)
Expand All @@ -111,7 +112,6 @@ def __init__(self, config: ModelConfig) -> None:
self.qo_head_local = div_even(self.config.num_qo_heads, tp_size)
self.kv_head_local = div_even(self.config.num_kv_heads, tp_size, allow_replicate=True)

self.cached_ones_cpu: torch.Tensor = torch.tensor([], dtype=torch.int32, pin_memory=True)
# for cuda graph
self.capture_bs: List[int] = []
self.max_graph_bs = 0
Expand Down Expand Up @@ -165,26 +165,14 @@ def _initialize_metadata_once(self, metadata: FIMetadata) -> None:
)
self.last_event.record()

def _get_ones_cpu(self, bs: int) -> torch.Tensor:
if bs <= len(self.cached_ones_cpu):
return self.cached_ones_cpu[:bs]
# padding to next pow of 2
next_len = _next_power_of_2(bs)
self.cached_ones_cpu = torch.ones(next_len, dtype=torch.int32, pin_memory=True)
return self.cached_ones_cpu[:bs]

def forward(
self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, layer_id: int, batch: Batch
) -> torch.Tensor:
def _flatten_cache(cache: torch.Tensor) -> torch.Tensor: # treat page = 1
return cache.view(-1, 1, cache.shape[2], cache.shape[3])

metadata = batch.attn_metadata
assert isinstance(metadata, FIMetadata)
self._initialize_metadata_once(metadata)
self.kvcache.store_kv(k, v, batch.out_loc, layer_id)
kv_cache = (self.kvcache.k_cache(layer_id), self.kvcache.v_cache(layer_id))
kv_cache = (_flatten_cache(kv_cache[0]), _flatten_cache(kv_cache[1]))
return metadata.wrapper.run(q=q, paged_kv_cache=kv_cache)

def prepare_metadata(self, batch: Batch) -> None:
Expand All @@ -193,31 +181,39 @@ def prepare_metadata(self, batch: Batch) -> None:
padded_size = len(reqs)
seqlens_q = [req.extend_len for req in reqs]
seqlens_k = [req.device_len for req in reqs]
cached_lens = [req.cached_len for req in reqs]
max_seqlen_q = max(seqlens_q)
CPU_KWARGS = {"device": "cpu", "dtype": torch.int32, "pin_memory": True}

device = self.device
seq_len_cpu = torch.tensor(seqlens_k, **CPU_KWARGS)
cu_seqlens_k_cpu = torch.tensor([0] + seqlens_k, **CPU_KWARGS).cumsum_(dim=0)
num_pages_per_seq = [div_ceil(l, self.page_size) for l in seqlens_k]
cu_seqlens_k_cpu = torch.tensor([0] + num_pages_per_seq, **CPU_KWARGS).cumsum_(dim=0)

if max_seqlen_q == 1: # decode with all extend_len = 1
cu_seqlens_q_cpu = torch.arange(0, padded_size + 1, **CPU_KWARGS)
elif all(l == 0 for l in cached_lens): # prefill with no cache hit
cu_seqlens_q_cpu = cu_seqlens_k_cpu
else: # normal extend prefill, with partial cache hit
else: # prefill, extend prefill, or chunk prefill
cu_seqlens_q_cpu = torch.tensor([0] + seqlens_q, **CPU_KWARGS).cumsum_(dim=0)

page_table = get_global_ctx().page_table

indices = torch.cat([page_table[req.table_idx, : req.device_len : self.page_size] for req in reqs])
if self.page_size > 1:
indices.div_(self.page_size, rounding_mode="floor")

last_page_len_cpu = torch.tensor([
((l - 1) % self.page_size) + 1 for l in seqlens_k
], **CPU_KWARGS)

batch.attn_metadata = FIMetadata(
cu_seqlens_q_cpu=cu_seqlens_q_cpu,
cu_seqlens_k_cpu=cu_seqlens_k_cpu,
cu_seqlens_q_gpu=cu_seqlens_q_cpu.to(device, non_blocking=True),
indices=torch.cat([page_table[req.table_idx, : req.device_len] for req in reqs]),
last_page_len_cpu=self._get_ones_cpu(padded_size),
indices=indices,
last_page_len_cpu=last_page_len_cpu,
num_qo_heads=self.qo_head_local,
num_kv_heads=self.kv_head_local,
head_dim=self.config.head_dim,
page_size=1,
page_size=self.page_size,
pos_encoding_mode="NONE",
seq_lens_cpu=seq_len_cpu,
dtype=self.kvcache.dtype,
Expand All @@ -227,7 +223,7 @@ def prepare_metadata(self, batch: Batch) -> None:
def init_capture_graph(self, max_seq_len: int, bs_list: List[int]) -> None:
assert self.capture is None, "Capture already initialized."
max_bs = max(bs_list)
capture = FICaptureData.create(max_bs, max_seq_len, self.kvcache.device)
capture = FICaptureData.create(max_bs, max_seq_len // self.page_size, self.kvcache.device)
capture.page_table = capture.page_table.view(-1) # use 1D as ragged indices
self.max_graph_bs = max_bs
self.capture = capture
Expand Down