diff --git a/README.md b/README.md index 1a829cb..c4847b1 100644 --- a/README.md +++ b/README.md @@ -341,6 +341,38 @@ config = GLiClassModelConfig( --- +### Streaming Classification + +GLiClass supports incremental, multi-session text classification over a decoder-KV model. Instead of re-encoding the full document on every update, it maintains a persistent KV cache per session and runs the scorer only when a pluggable strategy decides classification should fire. + +```bash +pip install gliclass[streaming] +``` + +```python +from gliclass.streaming import StreamingPipeline, SessionInput, EveryNTokensStrategy + +pipeline = StreamingPipeline(model, tokenizer, device="cuda", max_cache_len=1024) +strategy = EveryNTokensStrategy(n=50) + +for chunk in text_chunks: + outputs = pipeline([SessionInput( + session_id="doc_001", + text=chunk, + labels=["science", "politics", "finance"], + strategy=strategy, + classification_type="multi-label", + )]) + if outputs[0].triggered: + print(outputs[0].predictions) +``` + +Built-in strategies: `EveryChunkStrategy`, `EveryNTokensStrategy`, `OnDelimiterStrategy`, `SlidingWindowStrategy`, `ComposedStrategy`, `NeverStrategy`. + +For full documentation on session management, KV cache internals, batching, CPU offloading, and custom strategies, see [docs/streaming.md](docs/streaming.md). + +--- + ### Flash Attention Backends GLiClass supports optional flash attention backends for faster inference. diff --git a/docs/streaming.md b/docs/streaming.md new file mode 100644 index 0000000..529c08f --- /dev/null +++ b/docs/streaming.md @@ -0,0 +1,274 @@ +# Streaming classification + +`StreamingZeroShotClassificationPipeline` incrementally classifies text with a decoder-KV model. It keeps one text-only KV cache per active session and evaluates label sequences without adding those labels to the persistent cache. + +## Installation + +```bash +pip install gliclass[streaming] +``` + +## Basic usage + +```python +from gliclass.streaming import EveryNTokensStrategy, StreamingZeroShotClassificationPipeline + +pipeline = StreamingZeroShotClassificationPipeline( + model, + tokenizer, + device="cuda", + max_cache_len=1024, + default_strategy=EveryNTokensStrategy(50), +) + +for chunk in text_chunks: + output = pipeline( + chunk, + ["science", "politics", "finance"], + session_ids="document-1", + threshold=0.5, + )[0] + if output["triggered"]: + print(output["predictions"]) +``` + +The call returns one dictionary per input text: + +```python +{ + "session_id": "document-1", + "triggered": True, + "predictions": [{"label": "science", "score": 0.91}], + "cached_length": 128, + "tokens_added": 17, +} +``` + +## Batched sessions + +```python +outputs = pipeline( + texts=[chunk_a, chunk_b], + labels=["positive", "negative"], + session_ids=["session-a", "session-b"], + strategies=[strategy_a, strategy_b], + batch_size=8, +) +``` + +`texts`, `session_ids`, and a per-session `strategies` list must have matching lengths. Labels may be shared, specified per input, or provided as a hierarchical dictionary like the regular zero-shot pipeline. A single strategy is copied independently for every new session; a strategy list assigns one strategy to each new session. + +Session IDs must be unique within a call. + +## Session lifecycle + +The session IDs in each top-level call are the complete active set. A cached session is deleted automatically when it is absent from the next call. + +```python +pipeline([chunk_a, chunk_b], labels, session_ids=["a", "b"]) +pipeline([next_b, chunk_c], labels, session_ids=["b", "c"]) + +# Session "a" has been deleted; "b" was extended; "c" was created. +assert pipeline.active_sessions == ["b", "c"] +``` + +Cleanup is performed once for the complete call, before internal `batch_size` splitting. Therefore, sessions are not accidentally removed merely because they belong to different internal sub-batches. + +To retain a session without adding text, include it with an empty string: + +```python +pipeline( + texts=["", next_chunk_b], + labels=labels, + session_ids=["a", "b"], +) +``` + +Session `a` remains active, reports `tokens_added=0`, and is not classified. Omitting `a` would delete it. + +Manual lifecycle methods are also available: + +```python +pipeline.delete_session("b") +pipeline.clear_sessions() +print(pipeline.active_sessions) +``` + +Calling the pipeline with empty `texts` and `session_ids` clears all sessions: + +```python +pipeline([], [], session_ids=[]) +``` + +## Prompt and examples + +Prompt and few-shot examples become a one-time prefix when a session is created: + +```python +pipeline( + first_chunk, + labels, + session_ids="document-1", + prompt="Classify this document: ", + examples=examples, +) +``` + +The prefix is not appended again on later calls. Repeating the same prompt/examples is allowed. Once a session contains cached tokens, changing its prompt or examples raises an error; delete the session first when a different prefix is required. + +If a new session has an empty text but a non-empty prompt or examples, the prefix is still cached and counted in `tokens_added`. Prefix-only updates never trigger classification. + +## Classification strategies + +Available strategies: + +- `EveryChunkStrategy`: classify every non-empty update. +- `EveryNTokensStrategy(n)`: classify after accumulating `n` tokens. +- `OnDelimiterStrategy(delimiter)`: classify when the incoming text contains the delimiter. +- `NeverStrategy`: update the cache without classifying. +- `SlidingWindowStrategy(window_size)`: classify every chunk using only the newest cached tokens. +- `ComposedStrategy(trigger, window)`: combine a trigger with a classification window. + +Each session owns an independent copy of its strategy. Stateful counters are never shared across sessions, even when one strategy is supplied as the pipeline default. `default_strategy=None` selects `EveryChunkStrategy`. + +The `strategies` call argument initializes new sessions. Passing another strategy for an already-active session does not reset its state. Replace an active strategy explicitly with: + +```python +pipeline.set_session_strategy("document-1", EveryChunkStrategy()) +``` + +## Two-stage execution + +Each call has two stages: + +1. New text is passed to `model.update_decoder_cache(...)` with `use_cache=True`. The resulting text-only cache replaces the cache for that session. +2. Triggered sessions pass `<>label<