"""Deterministic prompts with controlled prefix sharing.
Serving engines cache the key/value state of prompt prefixes they have
seen. Whether two requests share a prefix therefore changes how much work
the second one needs, so a benchmark has to choose it on purpose:
- ``repeat`` sends one prompt for every request of a case, warmup included,
as Stormlog always has. After the first request, most of each prompt can
come from the cache.
- ``unique`` starts every request with its own nonce, so no two requests
share a prefix beyond whatever the server's chat template adds.
- ``shared-prefix`` gives each request one of N seeded group prefixes that
covers a set share of its tokens, then a request nonce and filler.
Nonces depend on the seed, the case and the phase, so a run can be repeated
exactly, and neither another case nor the warmup shares a prefix with the
measured requests.
"""
from __future__ import annotations
import hashlib
from collections.abc import Iterable
from dataclasses import dataclass, field
from functools import cached_property
from typing import Any
from .tokens import TokenCount, TokenCounter, generate_prompt
REPEAT = "repeat"
UNIQUE = "unique"
SHARED_PREFIX = "shared-prefix"
PROMPT_MODES = (REPEAT, UNIQUE, SHARED_PREFIX)
# A nonce takes about eight subword tokens, so shorter prompts overshoot
# their target and cannot hold a set prefix share.
MIN_CONTROLLED_TOKENS = 32
# Bump when the generated text changes, so digests from older runs are not
# mistaken for the same prompts.
GENERATOR_VERSION = 2
_FILLER = (
"profile",
"inference",
"latency",
"throughput",
"memory",
"tokens",
"scheduler",
"request",
"streaming",
"capacity",
)
[docs]
@dataclass(frozen=True)
class PromptSpec:
"""How the prompts of a run share their prefixes."""
mode: str = REPEAT
shared_prefix_ratio: float | None = None
prefix_groups: int | None = None
def __post_init__(self) -> None:
if self.mode not in PROMPT_MODES:
raise ValueError(f"prompt mode must be one of {', '.join(PROMPT_MODES)}")
shared = self.mode == SHARED_PREFIX
if shared != (self.shared_prefix_ratio is not None):
raise ValueError("a shared-prefix ratio goes with shared-prefix prompts")
if self.prefix_groups is not None and not shared:
raise ValueError("prefix groups go with shared-prefix prompts")
ratio = self.shared_prefix_ratio
if ratio is not None and not 0 < ratio < 1:
raise ValueError("shared-prefix ratio must be between 0 and 1")
if self.prefix_groups is not None and self.prefix_groups < 1:
raise ValueError("prefix groups must be >= 1")
[docs]
def to_record(self) -> dict[str, Any]:
record: dict[str, Any] = {"mode": self.mode}
if self.mode == SHARED_PREFIX:
record["shared_prefix_ratio"] = self.shared_prefix_ratio
record["prefix_groups"] = self.groups
return record
@property
def groups(self) -> int:
return self.prefix_groups or 1
[docs]
@dataclass(frozen=True)
class Prompt:
"""One request's prompt and what it shares with others.
``planned`` is the size the generator aimed for, known without
tokenizing the prompt. The exact ``count`` is only worked out when
something reads it: most servers report the prompt's tokens themselves.
"""
text: str
prompt_id: str
counter: TokenCounter = field(repr=False, compare=False)
planned: TokenCount | None = None
prefix_group: int | None = None
shared_prefix_tokens: int | None = None
@cached_property
def count(self) -> TokenCount:
return self.counter.count_text(self.text)
@property
def planned_count(self) -> TokenCount:
return self.planned or self.count
@cached_property
def digest(self) -> str:
return _digest(self.text)
[docs]
class PromptSource:
"""The prompts of one case phase, generated on demand and cached."""
def __init__(
self,
spec: PromptSpec,
*,
counter: TokenCounter,
seed: int,
case_id: str,
phase: str,
input_tokens: int,
repeated: dict[tuple[int, int], Prompt] | None = None,
) -> None:
self.spec = spec
self.counter = counter
self.seed = seed
self.input_tokens = input_tokens
self._namespace = f"{seed}:{case_id}:{phase}"
self._prompts: dict[int, Prompt] = {}
self._digests: dict[int, str] = {}
# Each group's prefix text and its token count.
self._prefixes: dict[int, tuple[str, int]] = {}
# Filler text and its token count, by the count it was built to reach.
self._fillers: dict[int, tuple[str, int]] = {}
self._word_tokens: list[int] | None = None
# The repeated prompt by (seed, length), shared by phases given one dict.
self._repeated = {} if repeated is None else repeated
[docs]
def prompt(self, index: int) -> Prompt:
"""Build, or return the built, prompt for one request."""
prompt = self._prompts.get(index)
if prompt is None:
prompt = self._build(index)
self._prompts[index] = prompt
self._digests[index] = prompt.digest
return prompt
[docs]
def take(self, index: int) -> Prompt:
"""The prompt a request is about to use; forget it once it is done."""
return self.prompt(index)
[docs]
def forget(self, index: int) -> None:
"""Drop a used prompt's text; its digest stays for the phase digest."""
self._prompts.pop(index, None)
[docs]
def warm(self, sample: int = 64) -> None:
"""Do a phase's one-off prompt work before its clock starts.
That is the repeated prompt, every group prefix, and the filler for
each length the first ``sample`` nonces leave room for. Each later
prompt then only tokenizes its short nonce. The sample prompts are
not handed out, so they are not part of the phase digest.
"""
for index in range(sample):
self._build(index)
if self.spec.mode == SHARED_PREFIX:
for group in range(self.spec.groups):
self._prefix(group)
[docs]
def prepare(self, indices: Iterable[int]) -> None:
"""Build prompts ahead of use."""
for index in indices:
self.prompt(index)
[docs]
def digest(self) -> str | None:
"""One digest over every prompt handed out, in index order."""
if not self._digests:
return None
digests = [self._digests[index] for index in sorted(self._digests)]
return _digest("\n".join(digests))
def _build(self, index: int) -> Prompt:
if self.spec.mode == REPEAT:
return self._repeat()
if self.spec.mode == UNIQUE:
nonce = self._nonce(f"request:{index}")
text, planned = self._extend(f"[{nonce}]", self.input_tokens)
return Prompt(text, f"r-{nonce}", self.counter, self._planned(planned))
return self._shared(index)
def _repeat(self) -> Prompt:
# The same text Stormlog has always generated for this length.
key = (self.seed, self.input_tokens)
if key not in self._repeated:
text = generate_prompt(
self.input_tokens, self.counter, seed=self.seed + self.input_tokens
)
exact = self.counter.count_text(text)
self._repeated[key] = Prompt(text, REPEAT, self.counter, planned=exact)
return self._repeated[key]
def _shared(self, index: int) -> Prompt:
group = self._group(index)
prefix, prefix_tokens = self._prefix(group)
nonce = self._nonce(f"request:{index}")
marker = f"[{nonce}]"
used = prefix_tokens + self._count(marker)
text, planned = self._extend(f"{prefix} {marker}", self.input_tokens, used)
return Prompt(
text,
f"g{group}-{nonce}",
self.counter,
self._planned(planned),
prefix_group=group,
shared_prefix_tokens=prefix_tokens,
)
def _prefix(self, group: int) -> tuple[str, int]:
if group not in self._prefixes:
ratio = self.spec.shared_prefix_ratio or 0.0
tokens = max(1, round(ratio * self.input_tokens))
text, _ = self._extend(f"[{self._nonce(f'prefix:{group}')}]", tokens)
self._prefixes[group] = (text, self._count(text))
return self._prefixes[group]
def _extend(
self, head: str, target_tokens: int, used: int | None = None
) -> tuple[str, int]:
"""``head`` and then filler sized to reach the target, and the size.
Only the head is tokenized for each prompt: the filler for each size
is built once per phase. Token counts are treated as adding up across
the space between them, which holds for whitespace estimates and,
within a token, for subword tokenizers.
"""
if used is None:
used = self._count(head)
filler, filler_tokens = self._filler(target_tokens - used)
return (f"{head} {filler}" if filler else head), used + filler_tokens
def _planned(self, tokens: int) -> TokenCount:
return TokenCount(value=tokens, source=self.counter.source, exact=False)
def _filler(self, tokens: int) -> tuple[str, int]:
"""The fewest filler words that reach ``tokens``, built once per length.
The word count is estimated from each filler word's own token count,
which subword tokenizers add up across spaces, then checked against
the counter: usually two counts, a few more for a counter whose
counts do not add up.
"""
if tokens <= 0:
return "", 0
if tokens not in self._fillers:
if self._word_tokens is None:
self._word_tokens = [
max(1, self._count(f" {word}")) for word in _FILLER
]
words = _estimated_words(tokens, self._word_tokens)
# Counted with the space before it, as it follows the head.
while words > 1 and self._count(" " + _filler_words(words - 1)) >= tokens:
words -= 1
while self._count(" " + _filler_words(words)) < tokens:
words += 1
text = _filler_words(words)
self._fillers[tokens] = (text, self._count(" " + text))
return self._fillers[tokens]
def _count(self, text: str) -> int:
return self.counter.count_text(text).value
def _group(self, index: int) -> int:
draw = hashlib.sha256(f"{self._namespace}:group:{index}".encode()).digest()
return int.from_bytes(draw[:8], "big") % self.spec.groups
def _nonce(self, label: str) -> str:
return _digest(f"{self._namespace}:{label}")[:12]
def _estimated_words(tokens: int, word_tokens: list[int]) -> int:
words = total = 0
while total < tokens:
total += word_tokens[words % len(word_tokens)]
words += 1
return words
def _filler_words(count: int) -> str:
return " ".join(_FILLER[index % len(_FILLER)] for index in range(count))
def _digest(text: str) -> str:
return hashlib.sha256(text.encode("utf-8")).hexdigest()[:16]