From 72d22c8f39269a4029307ad22794b7835cc82b79 Mon Sep 17 00:00:00 2001 From: BioMikeUkr Date: Fri, 10 Jul 2026 12:52:25 +0000 Subject: [PATCH 01/25] caching utils added --- gliclass/streaming/cache.py | 351 ++++++++++++++++++++++++++++++++++++ 1 file changed, 351 insertions(+) create mode 100644 gliclass/streaming/cache.py diff --git a/gliclass/streaming/cache.py b/gliclass/streaming/cache.py new file mode 100644 index 0000000..a7941ba --- /dev/null +++ b/gliclass/streaming/cache.py @@ -0,0 +1,351 @@ +""" +KV cache management for GLiClass streaming classification. + +Includes: + - CacheState / DynamicKVCacheManager (low-level per-session ops) + - BatchedKVHelper (stack / unstack heterogeneous KV caches for batched inference) +""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any + +import torch + + +# --------------------------------------------------------------------------- +# CacheState +# --------------------------------------------------------------------------- + +@dataclass +class CacheState: + past_key_values: Any + input_ids: torch.Tensor + attention_mask: torch.Tensor + position_ids: torch.Tensor | None = None + cached_length: int = 0 + session_id: str | None = None + metadata: dict = field(default_factory=dict) + + def to(self, device: torch.device) -> "CacheState": + return CacheState( + past_key_values=_move_past_kv(self.past_key_values, device), + input_ids=self.input_ids.to(device), + attention_mask=self.attention_mask.to(device), + position_ids=self.position_ids.to(device) if self.position_ids is not None else None, + cached_length=self.cached_length, + session_id=self.session_id, + metadata=self.metadata.copy(), + ) + + +def _move_past_kv(past_kv, device): + if past_kv is None: + return None + if hasattr(past_kv, "to"): + return past_kv.to(device) + if isinstance(past_kv, (tuple, list)): + return type(past_kv)(_move_past_kv(item, device) for item in past_kv) + return past_kv + + +# --------------------------------------------------------------------------- +def create_empty_cache(session_id: str | None = None, device=None) -> CacheState: + kw = {"device": device} if device is not None else {} + return CacheState( + past_key_values=None, + input_ids=torch.empty(0, dtype=torch.long, **kw), + attention_mask=torch.empty(0, dtype=torch.long, **kw), + position_ids=None, + cached_length=0, + session_id=session_id, + metadata={}, + ) + + +def truncate_cache(cache_state: CacheState, max_length: int) -> CacheState: + if cache_state.cached_length <= max_length or cache_state.past_key_values is None: + return cache_state + truncated = _crop_past_kv(cache_state.past_key_values, max_length) + if truncated is None: + return cache_state + return CacheState( + past_key_values=truncated, + input_ids=cache_state.input_ids[-max_length:], + attention_mask=cache_state.attention_mask[-max_length:], + position_ids=cache_state.position_ids[-max_length:] if cache_state.position_ids is not None else None, + cached_length=max_length, + session_id=cache_state.session_id, + metadata=cache_state.metadata.copy(), + ) + + +def _deep_copy_past_kv(past_kv): + if past_kv is None: + return None + from transformers.cache_utils import DynamicCache + if isinstance(past_kv, DynamicCache): + new = DynamicCache() + for i, layer in enumerate(past_kv.layers): + if layer.is_initialized: + new.update(layer.keys.clone(), layer.values.clone(), i) + return new + if hasattr(past_kv, "clone"): + return past_kv.clone() + if isinstance(past_kv, (tuple, list)): + return type(past_kv)(_deep_copy_past_kv(x) for x in past_kv) + return past_kv + + +def _crop_past_kv(past_kv, max_length: int): + """Crop past_key_values to last max_length tokens. Returns None if not supported.""" + if past_kv is None: + return None + from transformers.cache_utils import DynamicCache + if isinstance(past_kv, DynamicCache): + new = DynamicCache() + for i, layer in enumerate(past_kv.layers): + if layer.is_initialized: + new.update(layer.keys[:, :, -max_length:, :].clone(), + layer.values[:, :, -max_length:, :].clone(), i) + return new + return None + + +# --------------------------------------------------------------------------- +# BatchedKVHelper – stack / unstack heterogeneous KV caches +# --------------------------------------------------------------------------- + +class BatchedKVHelper: + """ + Utilities for batching decoder forwards across sessions with different + KV cache lengths. + + Stacking strategy: + - KV caches are prepend-padded with zeros to max_cached_len so that + position embeddings (RoPE) baked into each K/V tensor remain correct. + - The attention_mask covers [prepend_zeros | real_cached | new_tokens | new_pad]. + - position_ids for new tokens are set per-session: [cached_len_i .. cached_len_i + new_len_i). + """ + + @staticmethod + def stack_for_update( + caches: list[CacheState], + new_input_ids: list[torch.Tensor], # one per session, already on device + new_attention_masks: list[torch.Tensor], + device: torch.device, + ) -> dict: + """ + Prepare a batched decoder forward for the cache-update stage. + + Returns a dict with keys: + input_ids, attention_mask, position_ids, past_key_values, + cached_lengths, new_lengths, max_cached_len + """ + batch_size = len(caches) + cached_lengths = [c.cached_length for c in caches] + new_lengths = [ids.shape[-1] for ids in new_input_ids] + max_cached = max(cached_lengths) + max_new = max(new_lengths) + + # --- pad new tokens (right-pad) --- + padded_ids = new_input_ids[0].new_zeros(batch_size, max_new) + padded_new_mask = new_input_ids[0].new_zeros(batch_size, max_new) + position_ids = new_input_ids[0].new_zeros(batch_size, max_new) + + for i, (ids, mask, clen, nlen) in enumerate( + zip(new_input_ids, new_attention_masks, cached_lengths, new_lengths) + ): + padded_ids[i, :nlen] = ids[0] if ids.dim() == 2 else ids + padded_new_mask[i, :nlen] = mask[0] if mask.dim() == 2 else mask + position_ids[i, :nlen] = torch.arange(clen, clen + nlen, device=device) + + # --- build full attention mask: [prepend_zeros | cached | new | new_pad] --- + full_mask = torch.zeros(batch_size, max_cached + max_new, dtype=torch.long, device=device) + for i, (clen, nlen) in enumerate(zip(cached_lengths, new_lengths)): + pad = max_cached - clen + full_mask[i, pad : pad + clen] = 1 # real cached tokens + full_mask[i, max_cached : max_cached + nlen] = 1 # real new tokens + + # --- stack past_key_values (prepend-pad to max_cached) --- + stacked_past_kv = None + if max_cached > 0: + stacked_past_kv = BatchedKVHelper._stack_past_kv(caches, max_cached, device) + + return { + "input_ids": padded_ids, + "attention_mask": full_mask, + "position_ids": position_ids, + "past_key_values": stacked_past_kv, + "cached_lengths": cached_lengths, + "new_lengths": new_lengths, + "max_cached_len": max_cached, + } + + @staticmethod + def unstack_after_update( + new_past_kv, + stacked_info: dict, + old_caches: list[CacheState], + new_input_ids: list[torch.Tensor], + new_attention_masks: list[torch.Tensor], + ) -> list[CacheState]: + """ + Unstack per-session CacheState from a batched decoder output. + + new_past_kv has shape [..., max_cached + max_new, ...] along seq dim. + We slice [max_cached - clen_i : max_cached + new_len_i] to recover + only the real tokens for each session. + """ + cached_lengths = stacked_info["cached_lengths"] + new_lengths = stacked_info["new_lengths"] + max_cached = stacked_info["max_cached_len"] + + results = [] + for i, (old, clen, nlen) in enumerate(zip(old_caches, cached_lengths, new_lengths)): + real_start = max_cached - clen + real_end = max_cached + nlen + sliced_kv = BatchedKVHelper._slice_past_kv(new_past_kv, i, real_start, real_end) + new_clen = clen + nlen + + # rebuild input_ids / attention_mask for CacheState + new_ids = new_input_ids[i] + if new_ids.dim() == 2: + new_ids = new_ids[0] + new_mask = new_attention_masks[i] + if new_mask.dim() == 2: + new_mask = new_mask[0] + + if clen > 0: + full_ids = torch.cat([old.input_ids, new_ids[:nlen]], dim=0) + full_mask = torch.cat([old.attention_mask, new_mask[:nlen]], dim=0) + else: + full_ids = new_ids[:nlen] + full_mask = new_mask[:nlen] + + results.append(CacheState( + past_key_values=sliced_kv, + input_ids=full_ids, + attention_mask=full_mask, + position_ids=None, + cached_length=new_clen, + session_id=old.session_id, + metadata=old.metadata.copy(), + )) + + return results + + @staticmethod + def stack_for_classify( + caches: list[CacheState], + label_ids: list[torch.Tensor], + label_masks: list[torch.Tensor], + device: torch.device, + ) -> dict: + """ + Prepare a batched decoder forward for the classification stage. + + Label tokens are right-padded. KV caches are prepend-padded. + Returns input_ids, attention_mask, position_ids, past_key_values, + plus metadata needed for slicing scorer inputs afterwards. + """ + batch_size = len(caches) + cached_lengths = [c.cached_length for c in caches] + label_lengths = [ids.shape[-1] for ids in label_ids] + max_cached = max(cached_lengths) + max_label = max(label_lengths) + + padded_label_ids = label_ids[0].new_zeros(batch_size, max_label) + padded_label_mask = label_ids[0].new_zeros(batch_size, max_label) + position_ids = label_ids[0].new_zeros(batch_size, max_label) + + for i, (ids, mask, clen, llen) in enumerate( + zip(label_ids, label_masks, cached_lengths, label_lengths) + ): + flat_ids = ids[0] if ids.dim() == 2 else ids + flat_mask = mask[0] if mask.dim() == 2 else mask + padded_label_ids[i, :llen] = flat_ids + padded_label_mask[i, :llen] = flat_mask + position_ids[i, :llen] = torch.arange(clen, clen + llen, device=device) + + full_mask = torch.zeros(batch_size, max_cached + max_label, dtype=torch.long, device=device) + for i, (clen, llen) in enumerate(zip(cached_lengths, label_lengths)): + pad = max_cached - clen + full_mask[i, pad : pad + clen] = 1 + full_mask[i, max_cached : max_cached + llen] = 1 + + stacked_past_kv = None + if max_cached > 0: + stacked_past_kv = BatchedKVHelper._stack_past_kv(caches, max_cached, device) + + return { + "input_ids": padded_label_ids, + "attention_mask": full_mask, # full: cache + labels (for decoder) + "label_mask": padded_label_mask, # labels only (for scorer) + "position_ids": position_ids, + "past_key_values": stacked_past_kv, + "label_lengths": label_lengths, + "max_label_len": max_label, + } + + # ------------------------------------------------------------------ + # Internal helpers + # ------------------------------------------------------------------ + + @staticmethod + def _stack_past_kv(caches: list[CacheState], max_cached: int, device: torch.device): + """Prepend-pad each session's KV to max_cached and stack along batch dim.""" + from transformers.cache_utils import DynamicCache + + non_empty = [c for c in caches if c.past_key_values is not None] + num_layers = len(non_empty[0].past_key_values.layers) + + # Get reference shape from first non-empty layer + ref_layer = non_empty[0].past_key_values.layers[0] + ref_k = ref_layer.keys # [1, heads, seq, head_dim] + + stacked = DynamicCache() + for layer_idx in range(num_layers): + keys, values = [], [] + for c in caches: + kv = c.past_key_values + if kv is None or not kv.layers[layer_idx].is_initialized: + k = torch.zeros(1, ref_k.shape[1], max_cached, ref_k.shape[3], + device=device, dtype=ref_k.dtype) + v = torch.zeros_like(k) + else: + layer = kv.layers[layer_idx] + k = layer.keys.to(device) # [1, heads, clen, head_dim] + v = layer.values.to(device) + clen = k.shape[2] + pad = max_cached - clen + if pad > 0: + zeros_k = torch.zeros(1, k.shape[1], pad, k.shape[3], device=device, dtype=k.dtype) + zeros_v = torch.zeros_like(zeros_k) + k = torch.cat([zeros_k, k], dim=2) + v = torch.cat([zeros_v, v], dim=2) + keys.append(k) + values.append(v) + + stacked.update(torch.cat(keys, dim=0), torch.cat(values, dim=0), layer_idx) + + return stacked + + @staticmethod + def _slice_past_kv(past_kv, batch_idx: int, start: int, end: int): + """Extract single-session KV from a batched past_key_values.""" + if past_kv is None: + return None + + from transformers.cache_utils import DynamicCache + sliced = DynamicCache() + for i, layer in enumerate(past_kv.layers): + if not layer.is_initialized: + continue + sliced.update( + layer.keys[batch_idx : batch_idx + 1, :, start:end, :], + layer.values[batch_idx : batch_idx + 1, :, start:end, :], + i, + ) + return sliced From 2937d33d94a3af5ad48ccba43b87ec97d499b109 Mon Sep 17 00:00:00 2001 From: BioMikeUkr Date: Fri, 10 Jul 2026 12:52:47 +0000 Subject: [PATCH 02/25] streaming pipeline added --- gliclass/streaming/pipeline.py | 443 +++++++++++++++++++++++++++++++++ 1 file changed, 443 insertions(+) create mode 100644 gliclass/streaming/pipeline.py diff --git a/gliclass/streaming/pipeline.py b/gliclass/streaming/pipeline.py new file mode 100644 index 0000000..22eaa83 --- /dev/null +++ b/gliclass/streaming/pipeline.py @@ -0,0 +1,443 @@ +""" +StreamingPipeline – batched multi-session streaming classification. + +Flow per __call__: + 1. Cleanup sessions absent from current batch. + 2. Create CacheState for new session_ids. + 3. Stage 1 – batch update KV caches with new text (skip empty text). + 4. Stage 2 – resolve strategies; batch classify sessions that triggered. +""" + +from __future__ import annotations + +import torch + +from .cache import BatchedKVHelper, CacheState, create_empty_cache, truncate_cache +from .types import SessionInput, SessionOutput + + +class StreamingPipeline: + def __init__( + self, + model, + tokenizer, + device: str | torch.device = "cpu", + max_cache_len: int | None = None, + batch_size: int | None = None, + offload_to_cpu: bool = False, + use_pinned_memory: bool = False, + score_on_cpu: bool = False, + ): + self.model = model + self.tokenizer = tokenizer + self.device = torch.device(device) if isinstance(device, str) else device + self.max_cache_len = max_cache_len + self.batch_size = batch_size + self.offload_to_cpu = offload_to_cpu + self.use_pinned_memory = use_pinned_memory and offload_to_cpu and self.device.type == "cuda" + self.score_on_cpu = score_on_cpu + self._caches: dict[str, CacheState] = {} + + scorer_target = "cpu" if score_on_cpu else self.device + self.model.model.scorer.to(scorer_target) + + # ------------------------------------------------------------------ + # Public API + # ------------------------------------------------------------------ + + def __call__(self, inputs: list[SessionInput]) -> list[SessionOutput]: + self._cleanup_dead_sessions(inputs) + self._ensure_sessions(inputs) + + if self.batch_size is None or len(inputs) <= self.batch_size: + return self._process_batch(inputs) + + outputs = [None] * len(inputs) + for start in range(0, len(inputs), self.batch_size): + chunk = inputs[start : start + self.batch_size] + chunk_ids = {inp.session_id for inp in chunk} + + if self.offload_to_cpu: + self._load_to_device(chunk_ids, self.device) + + for i, out in enumerate(self._process_batch(chunk)): + outputs[start + i] = out + + if self.offload_to_cpu: + self._load_to_device(chunk_ids, torch.device("cpu")) + + return outputs + + def _process_batch(self, inputs: list[SessionInput]) -> list[SessionOutput]: + tokens_added = self._stage_update(inputs) + if self.max_cache_len is not None: + self._enforce_max_cache_len() + return self._stage_classify(inputs, tokens_added) + + def _load_to_device(self, session_ids: set[str], device: torch.device) -> None: + to_cpu = device.type == "cpu" + for sid in session_ids: + cache = self._caches.get(sid) + if cache is None or cache.past_key_values is None: + continue + if to_cpu and self.use_pinned_memory: + self._caches[sid] = _move_cache_pinned(cache) + else: + self._caches[sid] = _move_cache_nonblocking(cache, device) + + # synchronize before using on GPU so non_blocking transfers complete + if not to_cpu and self.use_pinned_memory and self.device.type == "cuda": + torch.cuda.synchronize(self.device) + + # ------------------------------------------------------------------ + # Session management + # ------------------------------------------------------------------ + + def _cleanup_dead_sessions(self, inputs: list[SessionInput]) -> None: + active = {inp.session_id for inp in inputs} + for sid in list(self._caches): + if sid not in active: + del self._caches[sid] + + def _enforce_max_cache_len(self) -> None: + for sid, cache in self._caches.items(): + if cache.cached_length > self.max_cache_len: + self._caches[sid] = truncate_cache(cache, self.max_cache_len) + + def _ensure_sessions(self, inputs: list[SessionInput]) -> None: + for inp in inputs: + if inp.session_id not in self._caches: + self._caches[inp.session_id] = create_empty_cache( + session_id=inp.session_id, device=self.device + ) + + # ------------------------------------------------------------------ + # Stage 1: update KV caches + # ------------------------------------------------------------------ + + def _stage_update(self, inputs: list[SessionInput]) -> list[int]: + """ + Update KV cache for each session. Returns tokens_added per input. + Sessions with empty text are skipped (tokens_added = 0). + """ + tokens_added = [0] * len(inputs) + to_update = [(i, inp) for i, inp in enumerate(inputs) if inp.text] + + if not to_update: + return tokens_added + + indices = [i for i, _ in to_update] + active_inputs = [inp for _, inp in to_update] + + # Tokenize new text per session + tokenized = [ + self.tokenizer(inp.text, return_tensors="pt", add_special_tokens=False) + for inp in active_inputs + ] + new_ids = [t["input_ids"].to(self.device) for t in tokenized] + new_masks = [t["attention_mask"].to(self.device) for t in tokenized] + + caches = [self._caches[inp.session_id] for inp in active_inputs] + + # Check if any session has a non-empty cache (need KV stacking) + has_cache = any(c.cached_length > 0 for c in caches) + + if has_cache: + stacked = BatchedKVHelper.stack_for_update(caches, new_ids, new_masks, self.device) + decoder = self.model.model.decoder_model + + with torch.no_grad(): + out = decoder( + input_ids=stacked["input_ids"], + attention_mask=stacked["attention_mask"], + position_ids=stacked["position_ids"], + past_key_values=stacked["past_key_values"], + use_cache=True, + return_dict=True, + ) + + updated = BatchedKVHelper.unstack_after_update( + out.past_key_values, stacked, caches, new_ids, new_masks + ) + else: + # All caches empty – simple batched forward without stacking + updated = self._update_from_scratch(caches, new_ids, new_masks) + + for session_idx, cache, new_len in zip( + indices, updated, stacked["new_lengths"] if has_cache else [ids.shape[-1] for ids in new_ids] + ): + inp = inputs[session_idx] + self._caches[inp.session_id] = cache + tokens_added[session_idx] = new_len + + return tokens_added + + def _update_from_scratch( + self, + caches: list[CacheState], + new_ids: list[torch.Tensor], + new_masks: list[torch.Tensor], + ) -> list[CacheState]: + """Batched update when all caches are empty (no KV stacking needed).""" + max_len = max(ids.shape[-1] for ids in new_ids) + batch_size = len(caches) + + padded_ids = new_ids[0].new_zeros(batch_size, max_len) + padded_masks = new_ids[0].new_zeros(batch_size, max_len) + position_ids = new_ids[0].new_zeros(batch_size, max_len) + + new_lengths = [] + for i, (ids, mask) in enumerate(zip(new_ids, new_masks)): + flat = ids[0] if ids.dim() == 2 else ids + flat_m = mask[0] if mask.dim() == 2 else mask + nlen = flat.shape[0] + padded_ids[i, :nlen] = flat + padded_masks[i, :nlen] = flat_m + position_ids[i, :nlen] = torch.arange(nlen, device=self.device) + new_lengths.append(nlen) + + decoder = self.model.model.decoder_model + with torch.no_grad(): + out = decoder( + input_ids=padded_ids, + attention_mask=padded_masks, + position_ids=position_ids, + use_cache=True, + return_dict=True, + ) + + results = [] + for i, (cache, nlen) in enumerate(zip(caches, new_lengths)): + sliced_kv = BatchedKVHelper._slice_past_kv(out.past_key_values, i, 0, nlen) + flat_ids = new_ids[i][0] if new_ids[i].dim() == 2 else new_ids[i] + flat_mask = new_masks[i][0] if new_masks[i].dim() == 2 else new_masks[i] + results.append(CacheState( + past_key_values=sliced_kv, + input_ids=flat_ids[:nlen], + attention_mask=flat_mask[:nlen], + cached_length=nlen, + session_id=cache.session_id, + metadata=cache.metadata.copy(), + )) + + return results + + # ------------------------------------------------------------------ + # Stage 2: classify triggered sessions + # ------------------------------------------------------------------ + + def _stage_classify( + self, inputs: list[SessionInput], tokens_added: list[int] + ) -> list[SessionOutput]: + # Resolve which sessions trigger classification + triggered_indices = [] + for i, (inp, n_added) in enumerate(zip(inputs, tokens_added)): + if n_added == 0: + continue # no new tokens → cache unchanged, skip classification + cache = self._caches[inp.session_id] + if inp.strategy.should_classify(n_added, cache.cached_length, inp.text): + triggered_indices.append(i) + + outputs: list[SessionOutput | None] = [None] * len(inputs) + + # Fill non-triggered outputs + for i, inp in enumerate(inputs): + if i not in triggered_indices: + cache = self._caches[inp.session_id] + outputs[i] = SessionOutput( + session_id=inp.session_id, + triggered=False, + predictions=None, + cached_length=cache.cached_length, + tokens_added=tokens_added[i], + ) + + if not triggered_indices: + return outputs + + # Prepare label sequences for triggered sessions + triggered_inputs = [inputs[i] for i in triggered_indices] + triggered_caches = [ + inputs[i].strategy.get_window(self._caches[inputs[i].session_id]) + for i in triggered_indices + ] + + label_seqs = [ + self._prepare_label_seq(inp.labels) for inp in triggered_inputs + ] + tokenized_labels = [ + self.tokenizer(seq, return_tensors="pt", add_special_tokens=False) + for seq in label_seqs + ] + label_ids = [t["input_ids"].to(self.device) for t in tokenized_labels] + label_masks = [t["attention_mask"].to(self.device) for t in tokenized_labels] + + stacked = BatchedKVHelper.stack_for_classify( + triggered_caches, label_ids, label_masks, self.device + ) + + decoder = self.model.model.decoder_model + with torch.no_grad(): + out = decoder( + input_ids=stacked["input_ids"], + attention_mask=stacked["attention_mask"], + position_ids=stacked["position_ids"], + past_key_values=stacked["past_key_values"], + use_cache=False, + return_dict=True, + ) + + # Slice label hidden states and pass to scorer + hidden = out.last_hidden_state # [batch, max_cached + max_label, hidden] + max_cached = stacked["past_key_values"] and max(c.cached_length for c in triggered_caches) or 0 + label_lengths = stacked["label_lengths"] + max_label = stacked["max_label_len"] + + # hidden for labels starts at index max_cached (prepend-padded KV) + # but we need the actual label slice per session + scorer = self.model.model.scorer + batch_logits = self._run_scorer_batched( + scorer, hidden, stacked["input_ids"], stacked["label_mask"], + label_lengths, max_label, + ) + + for local_idx, global_idx in enumerate(triggered_indices): + inp = inputs[global_idx] + cache = self._caches[inp.session_id] + logits = batch_logits[local_idx] # [num_labels] + preds = self._decode_predictions(logits, inp.labels, inp.classification_type) + outputs[global_idx] = SessionOutput( + session_id=inp.session_id, + triggered=True, + predictions=preds, + cached_length=cache.cached_length, + tokens_added=tokens_added[global_idx], + ) + + return outputs + + # ------------------------------------------------------------------ + # Helpers + # ------------------------------------------------------------------ + + def _prepare_label_seq(self, labels: list[str]) -> str: + label_token = "<