use std::sync::Arc;
use std::sync::atomic::AtomicU32;
use rmp_serde as rmps;
use rustc_hash::FxHashMap;
use crate::protocols::{DpRank, PlacementEvent, WorkerWithDpRank};
mod convert;
mod deserialize;
mod extra_keys;
mod filter;
#[cfg(test)]
mod tests;
mod types;
pub use convert::{convert_event, create_stored_block_from_parts, create_stored_blocks};
pub use extra_keys::{extra_keys_to_block_mm_infos, parse_mm_hash_from_extra_key};
pub use filter::KvCacheSpecKind;
pub use types::{BlockHashValue, ExtraKeyItem, KvEventBatch, KvTokenIds, RawKvEvent};
use filter::KvCacheEventMetadata;
pub fn decode_event_batch(payload: &[u8]) -> Result<KvEventBatch, rmps::decode::Error> {
rmps::from_slice(payload)
}
#[derive(Debug, Clone)]
pub struct ZmqEventNormalizer {
kv_block_size: u32,
image_token_id: Option<u32>,
warning_count: Arc<AtomicU32>,
group_metadata: FxHashMap<(DpRank, u32), KvCacheGroupMetadata>,
}
#[derive(Debug, Clone, Copy)]
struct KvCacheGroupMetadata {
kind: KvCacheSpecKind,
sliding_window: Option<u32>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ZmqEventFilterReason {
IgnoredEvent,
NonMainAttentionKind,
UnknownKind,
NonMainAttentionGroup,
UnlearnedGroupIdx,
}
impl ZmqEventFilterReason {
pub fn as_label(self) -> &'static str {
match self {
Self::IgnoredEvent => "ignored_event",
Self::NonMainAttentionKind => "non_main_attention_kind",
Self::UnknownKind => "unknown_kind",
Self::NonMainAttentionGroup => "non_main_attention_group",
Self::UnlearnedGroupIdx => "unlearned_group_idx",
}
}
}
impl ZmqEventNormalizer {
pub fn new(kv_block_size: u32) -> Self {
Self {
kv_block_size,
image_token_id: None,
warning_count: Arc::new(AtomicU32::new(0)),
group_metadata: FxHashMap::default(),
}
}
pub fn with_warning_count(kv_block_size: u32, warning_count: Arc<AtomicU32>) -> Self {
Self {
kv_block_size,
image_token_id: None,
warning_count,
group_metadata: FxHashMap::default(),
}
}
pub fn with_image_token_id(mut self, image_token_id: Option<u32>) -> Self {
self.image_token_id = image_token_id;
self
}
pub fn preprocess(&mut self, raw: RawKvEvent, worker: WorkerWithDpRank) -> Option<RawKvEvent> {
self.preprocess_with_reason(raw, worker).ok()
}
pub fn preprocess_with_reason(
&mut self,
raw: RawKvEvent,
worker: WorkerWithDpRank,
) -> Result<RawKvEvent, ZmqEventFilterReason> {
if raw.is_ignored() {
return Err(ZmqEventFilterReason::IgnoredEvent);
}
let metadata = raw.metadata();
if matches!(raw, RawKvEvent::BlockStored { .. }) {
self.learn_metadata(metadata, worker.dp_rank);
}
if let Some(reason) = self.filter_reason(metadata, worker.dp_rank) {
return Err(reason);
}
Ok(raw)
}
pub fn normalize_preprocessed(
&self,
raw: RawKvEvent,
event_id: u64,
worker: WorkerWithDpRank,
) -> Option<PlacementEvent> {
convert_event(
raw,
event_id,
self.kv_block_size,
worker,
&self.warning_count,
self.image_token_id,
)
}
pub fn normalize(
&mut self,
raw: RawKvEvent,
event_id: u64,
worker: WorkerWithDpRank,
) -> Option<PlacementEvent> {
let raw = self.preprocess(raw, worker)?;
self.normalize_preprocessed(raw, event_id, worker)
}
fn learn_metadata(&mut self, metadata: KvCacheEventMetadata, dp_rank: DpRank) {
let (Some(group_idx), Some(kind)) = (metadata.group_idx, metadata.kv_cache_spec_kind)
else {
return;
};
self.group_metadata.insert(
(dp_rank, group_idx),
KvCacheGroupMetadata {
kind,
sliding_window: metadata.kv_cache_spec_sliding_window,
},
);
}
fn filter_reason(
&self,
metadata: KvCacheEventMetadata,
dp_rank: DpRank,
) -> Option<ZmqEventFilterReason> {
if let Some(kind) = metadata.kv_cache_spec_kind {
if kind.is_main_attention() {
return None;
}
if kind == KvCacheSpecKind::Unknown {
return Some(ZmqEventFilterReason::UnknownKind);
}
return Some(ZmqEventFilterReason::NonMainAttentionKind);
}
let group_idx = metadata.group_idx?;
if let Some(metadata) = self.group_metadata.get(&(dp_rank, group_idx)) {
let _sliding_window = metadata.sliding_window;
if metadata.kind.is_main_attention() {
return None;
}
if metadata.kind == KvCacheSpecKind::Unknown {
return Some(ZmqEventFilterReason::UnknownKind);
}
return Some(ZmqEventFilterReason::NonMainAttentionGroup);
}
if group_idx == 0 {
None
} else {
Some(ZmqEventFilterReason::UnlearnedGroupIdx)
}
}
}