perf: replace paged KV cache with contiguous ContiguousCache, decode all groups
- Add KVCache/CacheView abstract base classes in cache.py - Add ContiguousCache (contiguous per-slot buffer, default) alongside PageCache (paged, renamed from old KVCache) - Merge make_table_tensor + bind into bind_tasks on KVCache interface - Remove task_cached/task_record_hashes from base class (PageCache-only) - Scheduler: decode all position groups instead of just the largest (eliminates 63% group skip rate) - Scheduler: accept optional cache param for swapping implementations - Model layer type hints use CacheView base class - Batch 1-32: 1-7% speedup from eliminating Storage.gather overhead - All 183 inference tests pass
This commit is contained in:
@@ -5,7 +5,7 @@ import torch.nn as nn
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.inference.core.cache import KvcacheView
|
||||
from astrai.inference.core.cache import CacheView
|
||||
from astrai.model.automodel import AutoModel
|
||||
from astrai.model.components.decoder_block import DecoderBlock
|
||||
from astrai.model.components.embedding import Embedding
|
||||
@@ -112,7 +112,7 @@ class AutoRegressiveLM(AutoModel):
|
||||
self,
|
||||
input_ids: Tensor,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
paged_cache: Optional[KvcacheView] = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
position_ids: Optional[Tensor] = None,
|
||||
) -> Dict[str, Tensor]:
|
||||
assert input_ids.ndim == 2
|
||||
|
||||
Reference in New Issue
Block a user