Skip to content
Closed
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
14 changes: 14 additions & 0 deletions astrai/inference/cache/pool.py
Original file line number Diff line number Diff line change
Expand Up @@ -104,6 +104,20 @@ def __init__(
page_size: int = 1,
n_tokens: Optional[int] = None,
):
if isinstance(page_size, bool) or not isinstance(page_size, int):
raise TypeError("page_size must be an integer")
if page_size <= 0:
raise ValueError("page_size must be positive")
if n_tokens is not None:
if isinstance(n_tokens, bool) or not isinstance(n_tokens, int):
raise TypeError("n_tokens must be an integer")
if n_tokens <= 0:
raise ValueError("n_tokens must be positive")
if n_tokens % page_size:
raise ValueError("n_tokens must be divisible by page_size")
elif page_size != 1:
raise ValueError("page_size requires n_tokens in paged mode")

self.page_size = page_size
self.max_batch_size = max_batch_size
self.max_seq_len = max_seq_len
Expand Down
4 changes: 4 additions & 0 deletions astrai/inference/engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -76,6 +76,8 @@ def __init__(
cache: Optional[PagePool] = None,
enable_cuda_graph: bool = True,
backend: Optional[Union[str, ATTN_BACKEND, AttentionBackend, type]] = None,
kv_cache_page_size: int = 1,
kv_cache_tokens: Optional[int] = None,
):
self.model = model
self.tokenizer = tokenizer
Expand All @@ -87,6 +89,8 @@ def __init__(
cache=cache,
enable_cuda_graph=enable_cuda_graph,
backend=backend,
kv_cache_page_size=kv_cache_page_size,
kv_cache_tokens=kv_cache_tokens,
)

self.scheduler.start()
Expand Down
18 changes: 17 additions & 1 deletion astrai/inference/network/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -111,6 +111,8 @@ def _create_engine(
dtype: torch.dtype = torch.bfloat16,
max_batch_size: int = 16,
max_seq_len: Optional[int] = None,
kv_cache_tokens: Optional[int] = None,
kv_cache_page_size: int = 1,
) -> InferenceEngine:
if not param_path.exists():
raise FileNotFoundError(f"Parameter directory not found: {param_path}")
Expand All @@ -125,8 +127,18 @@ def _create_engine(
tokenizer=tokenizer,
max_batch_size=max_batch_size,
max_seq_len=max_seq_len,
kv_cache_tokens=kv_cache_tokens,
kv_cache_page_size=kv_cache_page_size,
)
cache_mode = "paged" if kv_cache_tokens is not None else "contiguous"
logger.info(
"Inference engine initialized with max_batch_size=%s, "
"kv_cache_mode=%s, kv_cache_tokens=%s, kv_cache_page_size=%s",
max_batch_size,
cache_mode,
kv_cache_tokens,
kv_cache_page_size,
)
logger.info(f"Inference engine initialized with max_batch_size={max_batch_size}")
return engine


Expand Down Expand Up @@ -189,6 +201,8 @@ def run_server(
dtype: torch.dtype = torch.bfloat16,
max_batch_size: int = 16,
max_seq_len: Optional[int] = None,
kv_cache_tokens: Optional[int] = None,
kv_cache_page_size: int = 1,
):
app = get_app()
app.state.server_config = {
Expand All @@ -197,6 +211,8 @@ def run_server(
"param_path": param_path,
"max_batch_size": max_batch_size,
"max_seq_len": max_seq_len,
"kv_cache_tokens": kv_cache_tokens,
"kv_cache_page_size": kv_cache_page_size,
}
uvicorn.run(
app,
Expand Down
15 changes: 15 additions & 0 deletions astrai/inference/scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,8 @@ def __init__(
cache: Optional[PagePool] = None,
enable_cuda_graph: bool = True,
backend: Optional[Union[str, ATTN_BACKEND, AttentionBackend, type]] = None,
kv_cache_page_size: int = 1,
kv_cache_tokens: Optional[int] = None,
):
config = model.config

Expand All @@ -60,6 +62,11 @@ def __init__(
head_dim = config.hidden_size // config.num_attention_heads

if cache is not None:
if kv_cache_tokens is not None or kv_cache_page_size != 1:
raise ValueError(
"cache cannot be combined with kv_cache_tokens or "
"kv_cache_page_size"
)
self._cache = cache
else:
self._cache = PagePool(
Expand All @@ -70,6 +77,8 @@ def __init__(
max_seq_len=self.max_seq_len,
device=self.device,
dtype=self.dtype,
page_size=kv_cache_page_size,
n_tokens=kv_cache_tokens,
)

self._metrics = MetricsCollector()
Expand Down Expand Up @@ -125,6 +134,12 @@ def remove_task(self, task_id: str) -> bool:
def get_stats(self) -> Dict[str, Any]:
stats = self._task_mgr.get_stats()
stats["kv_cache_tasks"] = self._task_cache.task_count
stats["kv_cache_mode"] = "contiguous" if self._cache.contiguous else "paged"
stats["kv_cache_tokens"] = self._cache.n_tokens
stats["kv_cache_page_size"] = self._cache.page_size
stats["kv_cache_prefix_caching"] = (
not self._cache.contiguous and self._cache.page_size > 1
)
return stats

@property
Expand Down
2 changes: 2 additions & 0 deletions docs/developer/docker-serving.md
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,8 @@ server:
dtype: bfloat16 # bfloat16 | float16 | float32
max_batch_size: 16
max_seq_len: null # falls back to model config
kv_cache_tokens: null # set to enable paged allocation
kv_cache_page_size: 1 # values above 1 enable prefix caching
```

- Relative paths resolve from the YAML file's directory, not the current shell.
Expand Down
10 changes: 10 additions & 0 deletions docs/guides/inference.md
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,11 @@ PagePool (top-level manager, orchestrates all layers)
- **Contiguous (default)**: pre-allocates `max_batch_size * max_seq_len` token slots. `req_to_token` is a trivial linear mapping (`slot = req_idx * max_seq_len + pos`). No dynamic allocation.
- **Paged** (`page_size=1` or `>1` with `n_tokens` set): shared token pool with on-demand allocation. `Allocator` provides ref-counted allocation and LRU eviction. When `page_size > 1`, `RadixCache` also enables prefix sharing.

The server exposes the same choice as `kv_cache_tokens` and
`kv_cache_page_size`. Leaving `kv_cache_tokens` unset preserves contiguous
allocation. Setting it enables paged allocation; a page size above 1 also
enables prefix caching. The token capacity must be divisible by the page size.

`RadixCache` indexes complete token pages as parent-linked radix edges. Lookup walks from the root and compares each page's exact token tuple, so an identical page can only be reused under the same parent prefix. Hash values are retained for introspection, but never determine a match.

Only fully materialized KV pages enter the radix. A partial final page remains private to its request and is released when the request ends. On completion, the scheduler records the prompt plus generated tokens already decoded into KV; it excludes the final sampled token because that token has not yet passed through the model. A later request resumes prefill immediately after the longest complete-page hit.
Expand Down Expand Up @@ -192,13 +197,18 @@ server:
dtype: bfloat16
max_batch_size: 16
max_seq_len: null
kv_cache_tokens: 131072 # enables paged allocation
kv_cache_page_size: 64 # values above 1 enable prefix caching
```

```bash
python scripts/tools/server.py --config serve.yaml
python scripts/tools/server.py --config serve.yaml --port 9000 # CLI wins
```

`GET /stats` reports the effective `kv_cache_mode`, token capacity, page size,
and whether prefix caching is active.

In Docker, `scripts/serve.sh` drives the same YAML (a `runtime:` section
controls ports/GPU/mounts); see
[Docker Serving](../developer/docker-serving.md).
Expand Down
6 changes: 6 additions & 0 deletions docs/guides/params.md
Original file line number Diff line number Diff line change
Expand Up @@ -211,6 +211,8 @@ nohup python scripts/tools/train.py \
| `--dtype` | str | `bfloat16` | Model weights dtype (`bfloat16`, `float16`, `float32`) |
| `--max_batch_size` | int | `16` | Maximum batch size for continuous batching |
| `--max_seq_len` | int | model config `max_position_embeddings` | Maximum sequence length (KV cache size + prompt truncation) |
| `--kv_cache_tokens` | int | `None` | Shared KV token capacity; setting it enables paged allocation |
| `--kv_cache_page_size` | int | `1` | Paged allocation size; values above 1 enable prefix caching |
| `--reload` | flag | `False` | Enable auto-reload for development |

Usage:
Expand All @@ -230,7 +232,11 @@ server:
dtype: bfloat16
max_batch_size: 16
max_seq_len: null
kv_cache_tokens: 131072
kv_cache_page_size: 64
```
`kv_cache_tokens` must be positive and divisible by `kv_cache_page_size`.
Leave it unset to retain the contiguous cache default.
`serve.yaml` also carries a `runtime:` section for the Docker wrapper; see
[Docker Serving](../developer/docker-serving.md).

Expand Down
53 changes: 53 additions & 0 deletions scripts/tools/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,8 @@
"dtype",
"max_batch_size",
"max_seq_len",
"kv_cache_tokens",
"kv_cache_page_size",
)


Expand Down Expand Up @@ -62,6 +64,25 @@ def _as_int(value, name: str) -> int | None:
raise click.UsageError(f"{name} must be an integer, got {value!r}") from None


def _validate_kv_cache_config(
kv_cache_tokens: int | None, kv_cache_page_size: int
) -> None:
if kv_cache_page_size <= 0:
raise click.UsageError("server.kv_cache_page_size must be positive")
if kv_cache_tokens is None:
if kv_cache_page_size != 1:
raise click.UsageError(
"server.kv_cache_page_size requires server.kv_cache_tokens"
)
return
if kv_cache_tokens <= 0:
raise click.UsageError("server.kv_cache_tokens must be positive")
if kv_cache_tokens % kv_cache_page_size:
raise click.UsageError(
"server.kv_cache_tokens must be divisible by server.kv_cache_page_size"
)


def _resolve_server_config(
config_path: str,
passed_kwargs: dict,
Expand All @@ -79,11 +100,21 @@ def _resolve_server_config(
_as_int(resolved["max_batch_size"], "server.max_batch_size") or 16
)
resolved["max_seq_len"] = _as_int(resolved["max_seq_len"], "server.max_seq_len")
resolved["kv_cache_tokens"] = _as_int(
resolved.get("kv_cache_tokens"), "server.kv_cache_tokens"
)
page_size = _as_int(
resolved.get("kv_cache_page_size", 1), "server.kv_cache_page_size"
)
resolved["kv_cache_page_size"] = 1 if page_size is None else page_size
resolved["reload"] = bool(resolved["reload"])
if resolved["dtype"] not in _DTYPES:
raise click.UsageError(
f"server.dtype must be one of {', '.join(_DTYPES)}, got {resolved['dtype']!r}"
)
_validate_kv_cache_config(
resolved["kv_cache_tokens"], resolved["kv_cache_page_size"]
)
return resolved


Expand Down Expand Up @@ -126,6 +157,18 @@ def _resolve_server_config(
default=None,
help="Maximum sequence length (KV cache size + prompt truncation). Uses model config if not set.",
)
@click.option(
"--kv_cache_tokens",
type=int,
default=None,
help="Shared KV token capacity. Setting it enables paged allocation.",
)
@click.option(
"--kv_cache_page_size",
type=int,
default=1,
help="Paged KV allocation size. Values above 1 enable prefix caching.",
)
@click.pass_context
def server_command(
ctx,
Expand All @@ -138,6 +181,8 @@ def server_command(
dtype,
max_batch_size,
max_seq_len,
kv_cache_tokens,
kv_cache_page_size,
):
"""Launch inference server (OpenAI-compatible API)."""
if config_path:
Expand All @@ -150,6 +195,8 @@ def server_command(
"dtype": dtype,
"max_batch_size": max_batch_size,
"max_seq_len": max_seq_len,
"kv_cache_tokens": kv_cache_tokens,
"kv_cache_page_size": kv_cache_page_size,
}
explicit_keys = {
key
Expand All @@ -165,8 +212,12 @@ def server_command(
dtype = resolved["dtype"]
max_batch_size = resolved["max_batch_size"]
max_seq_len = resolved["max_seq_len"]
kv_cache_tokens = resolved["kv_cache_tokens"]
kv_cache_page_size = resolved["kv_cache_page_size"]
click.echo(f"Config: {config_path}")

_validate_kv_cache_config(kv_cache_tokens, kv_cache_page_size)

dtype_map = {
"bfloat16": torch.bfloat16,
"float16": torch.float16,
Expand All @@ -186,6 +237,8 @@ def server_command(
param_path=Path(param_path),
max_batch_size=max_batch_size,
max_seq_len=max_seq_len,
kv_cache_tokens=kv_cache_tokens,
kv_cache_page_size=kv_cache_page_size,
)


Expand Down
25 changes: 25 additions & 0 deletions tests/inference/test_cache.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
"""Unit tests for inference cache components."""

import pytest
import torch

from astrai.inference.cache import (
Expand Down Expand Up @@ -231,6 +232,30 @@ def _make_contiguous_pool(**kwargs):
return PagePool(**defaults)


@pytest.mark.parametrize(
("kwargs", "error", "message"),
[
({"page_size": 0}, ValueError, "page_size must be positive"),
({"page_size": True}, TypeError, "page_size must be an integer"),
({"n_tokens": 0}, ValueError, "n_tokens must be positive"),
({"n_tokens": True}, TypeError, "n_tokens must be an integer"),
(
{"page_size": 8, "n_tokens": 10},
ValueError,
"n_tokens must be divisible by page_size",
),
(
{"page_size": 8},
ValueError,
"page_size requires n_tokens in paged mode",
),
],
)
def test_page_pool_rejects_invalid_capacity_settings(kwargs, error, message):
with pytest.raises(error, match=message):
_make_contiguous_pool(**kwargs)


def test_page_pool_contiguous_task_alloc_free():
pool = _make_contiguous_pool()
task_cache = _make_task_cache(pool)
Expand Down
15 changes: 15 additions & 0 deletions tests/inference/test_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -288,6 +288,21 @@ def test_engine_passes_backend_to_scheduler():
assert MockSched.call_args.kwargs["backend"] == "torch_native"


def test_engine_passes_paged_cache_settings_to_scheduler():
mock_model, mock_tokenizer = _make_engine_mocks()

with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
InferenceEngine(
mock_model,
mock_tokenizer,
kv_cache_tokens=4096,
kv_cache_page_size=64,
)

assert MockSched.call_args.kwargs["kv_cache_tokens"] == 4096
assert MockSched.call_args.kwargs["kv_cache_page_size"] == 64


def test_generate_captures_calling_backend_context():
mock_model, mock_tokenizer = _make_engine_mocks()
captured = []
Expand Down
22 changes: 22 additions & 0 deletions tests/inference/test_scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -246,6 +246,28 @@ def get_stats():
assert stats["total_tasks"] >= 0


def test_scheduler_exposes_paged_cache_settings(mock_model_and_tokenizer):
mock_model, mock_tokenizer = mock_model_and_tokenizer

scheduler = InferenceScheduler(
model=mock_model,
tokenizer=mock_tokenizer,
max_batch_size=4,
max_seq_len=64,
device="cpu",
kv_cache_tokens=512,
kv_cache_page_size=8,
)
try:
stats = scheduler.get_stats()
assert stats["kv_cache_mode"] == "paged"
assert stats["kv_cache_tokens"] == 512
assert stats["kv_cache_page_size"] == 8
assert stats["kv_cache_prefix_caching"] is True
finally:
scheduler.stop()


def _make_real_scheduler(device):
"""Build a scheduler backed by a tiny real model for run_batch tests."""
cfg = make_rollout_config(max_position_embeddings=64)
Expand Down
Loading
Loading