"""ex-vllm-load — kv_cache_bytes. STARTER: the signature and docstring are the contract;
fill in the body.

Every decoding step needs one K and one V vector per layer, per KV head, kept for every
token in the sequence, for every request in the batch. Grouped-query attention (GQA) shares
K/V across groups of query heads, so what multiplies the cache is heads_kv (the NUMBER OF KV
HEADS) — not the number of query (attention) heads, which never appears in this formula.

    bytes = 2 (K and V) * layers * heads_kv * head_dim * seq * batch * dtype_bytes
"""

from __future__ import annotations

DTYPE_BYTES = {
    "fp16": 2,
    "bf16": 2,
    "fp32": 4,
    "int8": 1,
}


def kv_cache_bytes(
    layers: int, heads_kv: int, head_dim: int, seq: int, batch: int, dtype: str
) -> int:
    """Bytes of KV-cache memory for `batch` concurrent requests, each holding `seq` tokens
    of context, in a model with `layers` decoder layers and `heads_kv` KV heads (the GQA
    head count) of width `head_dim`, stored in `dtype` ("fp16" / "bf16" / "fp32" / "int8",
    looked up in DTYPE_BYTES).

    Raise ValueError if any of layers/heads_kv/head_dim/seq/batch is not a positive integer,
    or if dtype is not one of DTYPE_BYTES's keys.
    """
    ...
