Skip to content
Open
Show file tree
Hide file tree
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
8 changes: 6 additions & 2 deletions mistralrs-core/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -155,7 +155,11 @@ pub use mistralrs_mcp::{
};
pub use mistralrs_quant::{IsqBits, IsqType};
pub use mistralrs_sandbox::{NetworkMode, SandboxPolicy};
pub use paged_attention::{MemoryGpuConfig, PagedAttentionConfig, PagedCacheType};
pub use paged_attention::{
compute_block_hashes, BlockHash, BlockHashWithGroupId, InMemoryKvCacheConnector,
KVCacheManager, KvCacheConnector, MemoryGpuConfig, NoopKvCacheConnector, PagedAttentionConfig,
PagedCacheType,
};
pub use pipeline::hf::{
get_model_file, hf_home_dir, hf_hub_cache_dir, hf_token_path, is_hf_hub_offline,
list_model_files, probe_hf_repo_files, read_model_file_range, try_get_model_file,
Expand Down Expand Up @@ -2689,7 +2693,7 @@ impl MistralRs {
loader_config.silent,
loader_config.device_map_setting.clone(),
loader_config.isq,
loader_config.paged_attn_config,
loader_config.paged_attn_config.clone(),
)
.map_err(|e| MistralRsError::ReloadFailed(format!("Failed to load model: {e}")))?;

Expand Down
26 changes: 25 additions & 1 deletion mistralrs-core/src/paged_attention/block_pool.rs
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,10 @@

use std::collections::HashMap;

use std::sync::Arc;

use super::block_hash::{BlockHash, BlockHashWithGroupId};
use super::kv_cache_connector::{KvCacheConnector, NoopKvCacheConnector};

/// Sentinel value for "no link" in the doubly-linked free list.
const NO_LINK: usize = usize::MAX;
Expand Down Expand Up @@ -279,15 +282,32 @@ pub struct BlockPool {
null_block_id: usize,
/// The block size (number of tokens per block) for hash computation.
hash_block_size: usize,
/// Optional external KV connector (no-op by default).
connector: Arc<dyn KvCacheConnector>,
}

impl BlockPool {
/// Create a new block pool.
/// Create a new block pool with the default no-op connector.
///
/// `num_gpu_blocks`: Number of physical GPU blocks.
/// `enable_caching`: Whether to enable prefix caching.
/// `hash_block_size`: Block size used for hash computation.
pub fn new(num_gpu_blocks: usize, enable_caching: bool, hash_block_size: usize) -> Self {
Self::with_connector(
num_gpu_blocks,
enable_caching,
hash_block_size,
Arc::new(NoopKvCacheConnector),
)
}

/// Create a new block pool with a custom KV cache connector.
pub fn with_connector(
num_gpu_blocks: usize,
enable_caching: bool,
hash_block_size: usize,
connector: Arc<dyn KvCacheConnector>,
) -> Self {
assert!(num_gpu_blocks > 0, "Must have at least 1 GPU block");

// Allocate blocks: [0..num_gpu_blocks) are real blocks,
Expand All @@ -310,6 +330,7 @@ impl BlockPool {
num_gpu_blocks,
null_block_id: 0, // Will be set below
hash_block_size,
connector,
};

// Pop the first block as the null block (placeholder, never freed)
Expand Down Expand Up @@ -496,6 +517,9 @@ impl BlockPool {
/// Evict a cached block's hash from the cache map and reset its hash.
fn maybe_evict_cached_block(&mut self, block_id: usize) {
let block_hashes = std::mem::take(&mut self.blocks[block_id].block_hashes);
if !block_hashes.is_empty() {
self.connector.observe_evict(&block_hashes, block_id);
}
for hash in &block_hashes {
self.cached_block_hash_to_block.pop(hash, block_id);
}
Expand Down
15 changes: 14 additions & 1 deletion mistralrs-core/src/paged_attention/cache_engine.rs
Original file line number Diff line number Diff line change
Expand Up @@ -38,12 +38,25 @@ impl FromStr for PagedCacheType {
}
}

#[derive(Clone, Debug)]
#[derive(Clone)]
pub struct CacheConfig {
pub block_size: usize,
pub num_gpu_blocks: usize,
pub cache_type: PagedCacheType,
pub kv_cache_group_ids: Vec<u32>,
pub kv_cache_connector: std::sync::Arc<dyn super::KvCacheConnector>,
}

impl std::fmt::Debug for CacheConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CacheConfig")
.field("block_size", &self.block_size)
.field("num_gpu_blocks", &self.num_gpu_blocks)
.field("cache_type", &self.cache_type)
.field("kv_cache_group_ids", &self.kv_cache_group_ids)
.field("kv_cache_connector", &"Arc<dyn KvCacheConnector>")
.finish()
}
}

pub type KVCache = (Tensor, Tensor);
Expand Down
107 changes: 107 additions & 0 deletions mistralrs-core/src/paged_attention/in_memory_kv_cache_connector.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,107 @@
//! In-memory reference connector for the KV cache seam.
//!
//! This is an example "external" tier: it keeps a hash -> local block-id index
//! outside the block pool's own prefix map. It does not offload tensor bytes
//! (disk/S3 would still need a hydrate path before returning block IDs).

use std::collections::HashMap;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Mutex;

use super::block_hash::{BlockHash, BlockHashWithGroupId};
use super::kv_cache_connector::KvCacheConnector;

/// Process-local external KV index used as a reference `KvCacheConnector`.
#[derive(Debug, Default)]
pub struct InMemoryKvCacheConnector {
blocks: Mutex<HashMap<BlockHash, usize>>,
lookups: AtomicUsize,
hits: AtomicUsize,
stores: AtomicUsize,
evicts: AtomicUsize,
}

impl InMemoryKvCacheConnector {
pub fn new() -> Self {
Self::default()
}

pub fn len(&self) -> usize {
self.blocks.lock().expect("connector lock").len()
}

pub fn is_empty(&self) -> bool {
self.len() == 0
}

pub fn lookup_count(&self) -> usize {
self.lookups.load(Ordering::SeqCst)
}

pub fn hit_count(&self) -> usize {
self.hits.load(Ordering::SeqCst)
}

pub fn store_count(&self) -> usize {
self.stores.load(Ordering::SeqCst)
}

pub fn evict_count(&self) -> usize {
self.evicts.load(Ordering::SeqCst)
}
}

impl KvCacheConnector for InMemoryKvCacheConnector {
fn lookup_blocks(
&self,
block_hashes: &[BlockHash],
_group_ids: &[u32],
max_blocks: usize,
) -> Option<Vec<usize>> {
self.lookups.fetch_add(1, Ordering::SeqCst);
let guard = self.blocks.lock().expect("connector lock");
let mut ids = Vec::new();
for hash in block_hashes.iter().take(max_blocks) {
match guard.get(hash) {
Some(&id) => ids.push(id),
None => break,
}
}
if ids.is_empty() {
None
} else {
Some(ids)
}
}

fn observe_hit(
&self,
_block_hashes: &[BlockHash],
_block_ids: &[usize],
_num_computed_tokens: usize,
) {
self.hits.fetch_add(1, Ordering::SeqCst);
}

fn observe_store(
&self,
block_hashes: &[BlockHash],
block_ids: &[usize],
_num_full_blocks: usize,
) {
self.stores.fetch_add(1, Ordering::SeqCst);
let mut guard = self.blocks.lock().expect("connector lock");
for (hash, &block_id) in block_hashes.iter().zip(block_ids.iter()) {
guard.insert(*hash, block_id);
}
}

fn observe_evict(&self, block_hashes: &[BlockHashWithGroupId], block_id: usize) {
self.evicts.fetch_add(1, Ordering::SeqCst);
let mut guard = self.blocks.lock().expect("connector lock");
for entry in block_hashes {
guard.remove(&entry.block_hash);
}
guard.retain(|_, id| *id != block_id);
}
}
47 changes: 47 additions & 0 deletions mistralrs-core/src/paged_attention/kv_cache_connector.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,47 @@
//! Optional KV cache connector seam for external cache tiers.
//!
//! Default behavior is local-only prefix caching. A connector may observe
//! lookup/store/evict traffic or supply already-materialized local block IDs.

use super::block_hash::{BlockHash, BlockHashWithGroupId};

/// Hook into paged-attention prefix-cache block flow.
///
/// Implementations must keep default local behavior intact when they miss:
/// `lookup_blocks` returning `None` means "fall back to the in-process pool."
/// Returned block IDs, when `Some`, must already be valid IDs in the local
/// `BlockPool` (e.g. after the connector has hydrated external KV into them).
pub trait KvCacheConnector: Send + Sync {
fn lookup_blocks(
&self,
_block_hashes: &[BlockHash],
_group_ids: &[u32],
_max_blocks: usize,
) -> Option<Vec<usize>> {
None
}

fn observe_hit(
&self,
_block_hashes: &[BlockHash],
_block_ids: &[usize],
_num_computed_tokens: usize,
) {
}

fn observe_store(
&self,
_block_hashes: &[BlockHash],
_block_ids: &[usize],
_num_full_blocks: usize,
) {
}

fn observe_evict(&self, _block_hashes: &[BlockHashWithGroupId], _block_id: usize) {}
}

/// Default connector: always miss, never observe.
#[derive(Debug, Default, Clone, Copy)]
pub struct NoopKvCacheConnector;

impl KvCacheConnector for NoopKvCacheConnector {}
Loading
Loading