perf: migrate Req token-id storage to array.array('q') in Scheduler (#25098)

Co-authored-by: jialino <jialino@fb.com>
This commit is contained in:
Jialin Ouyang
2026-05-22 10:51:07 -07:00
committed by GitHub
co-authored by jialino
parent 5e9bd21979
commit 06c23d55b5
34 changed files with 833 additions and 299 deletions
+334
View File
@@ -0,0 +1,334 @@
"""Benchmark `list[int]` vs `array.array('q')` storage for
`Req.origin_input_ids` / `Req.output_ids` over one request lifecycle.
Simulated steps (per batch):
1. ingest -- tokenizer list[int] -> storage container.
2. prefix_match -- scheduler radix-tree lookup; RadixKey.match()
zip+!= walk. Exposes the per-element PyLong-boxing
cost array.array introduces (list[int] iterates
existing PyLongs and pays nothing).
3. prefill -- (a) fill_ids = origin + output,
(b) per-req slice fill_ids[prefix_len:],
(c) cross-req flatten + pinned cuda tensor build.
4. decode -- per-step output.append(next_token) for n_decode steps.
5. finish -- cache_finished_req:
(a) concat (origin + output)[:kv_committed_len]
for the radix-tree insert.
(b) RadixKey.match() zip+!= walk during insert's
tree traversal — second PyLong-boxing hotspot
on the array.array path.
Usage:
python benchmark/scheduler/bench_token_storage.py
"""
from __future__ import annotations
import time
from array import array
from collections import defaultdict
from contextlib import contextmanager
from itertools import chain
from typing import Any, Callable, Iterator
import numpy as np
import torch
# Per-req stages accumulate across reqs in a batch; batch_torch_tensor
# is the single cross-req prepare_for_extend tensor build.
STAGES = (
"ingest",
"prefix_match",
"prefill_concat",
"prefill_perreq_slice",
"batch_torch_tensor",
"decode_append",
"finish_concat",
"cache_finished_req",
)
def _ingest_list(seed: list[int]) -> list[int]:
return seed
def _ingest_pyarray(seed: list[int]) -> array:
return array("q", seed)
def _empty_list() -> list[int]:
return []
def _empty_pyarray() -> array:
return array("q")
def _zip_iterate(t0: Any, t1: Any) -> int:
"""Simulate zip iteration which surface PyLong boxing cost in array scenario"""
i = 0
for a, b in zip(t0, t1):
if a != b:
break
i += 1
return i
def _batch_tensor_from_lists(parts: list[list[int]]) -> torch.Tensor:
flat = list(chain.from_iterable(parts))
return torch.tensor(flat, dtype=torch.int64, pin_memory=True).to(
"cuda", non_blocking=True
)
def _batch_tensor_from_pyarrays(parts: list[array]) -> torch.Tensor:
# np.frombuffer gives a zero-copy view; np.concatenate is one C-level
# memcpy. This bypasses the per-element PyLong->int64 walk that
# torch.tensor(array('q')) would otherwise do.
views = [np.frombuffer(p, dtype=np.int64) for p in parts]
combined = np.concatenate(views) if len(views) > 1 else views[0]
return torch.from_numpy(combined).pin_memory().to("cuda", non_blocking=True)
LIST_KIT = {
"ingest_fn": _ingest_list,
"empty_fn": _empty_list,
"batch_torch_fn": _batch_tensor_from_lists,
}
PYARRAY_KIT = {
"ingest_fn": _ingest_pyarray,
"empty_fn": _empty_pyarray,
"batch_torch_fn": _batch_tensor_from_pyarrays,
}
@contextmanager
def timed(timings: dict[str, float], stage: str) -> Iterator[None]:
t0 = time.monotonic_ns()
try:
yield
finally:
timings[stage] += time.monotonic_ns() - t0
def simulate(
seeds: list[list[int]],
n_decode: int,
*,
ingest_fn: Callable[[list[int]], Any],
empty_fn: Callable[[], Any],
batch_torch_fn: Callable[[list[Any]], torch.Tensor],
) -> dict[str, float]:
"""One scheduling-round lifecycle. Returns per-stage cumulative ns."""
timings: dict[str, float] = defaultdict(float)
n_reqs = len(seeds)
n_origins = [len(s) for s in seeds]
origins: list[Any] = [None] * n_reqs
outputs: list[Any] = [None] * n_reqs
# 1. ingest
for i, seed in enumerate(seeds):
with timed(timings, "ingest"):
origins[i] = ingest_fn(seed)
outputs[i] = empty_fn()
# 2. prefix_match: simulating the worse scenario of PyLong-boxing overhead during prefix_match
for i in range(n_reqs):
with timed(timings, "prefix_match"):
_ = _zip_iterate(origins[i], origins[i])
# 3. prefill
per_req_slices: list[Any] = [None] * n_reqs
for i in range(n_reqs):
# 3a. fill_ids = origin_input_ids + output_ids
with timed(timings, "prefill_concat"):
fill_ids = origins[i] + outputs[i]
# 3b. input_ids = fill_ids[len(prefix_indices):]; prefix_len=0 here.
with timed(timings, "prefill_perreq_slice"):
per_req_slices[i] = fill_ids[0:]
# 3c. prepare_for_extend tensor build: flatten per-req slices, then
# build the pinned GPU tensor (kit-specific path).
with timed(timings, "batch_torch_tensor"):
_ = batch_torch_fn(per_req_slices)
# 4. decode
for i in range(n_reqs):
with timed(timings, "decode_append"):
for j in range(n_decode):
outputs[i].append(j)
# 5. finish: cache_finished_req -> insert -> _insert_helper tree walk.
for i in range(n_reqs):
# 5a. (origin + output)[:kv_committed_len] for the radix-tree insert.
with timed(timings, "finish_concat"):
committed = (origins[i] + outputs[i])[: n_origins[i] + n_decode]
# 5b. simulating the worse scenario of PyLong-boxing overhead during cache_finished_req
with timed(timings, "cache_finished_req"):
_ = _zip_iterate(committed, committed)
return timings
def bench_lifecycle(
seeds: list[list[int]],
n_decode: int,
iterations: int,
*,
ingest_fn: Callable[[list[int]], Any],
empty_fn: Callable[[], Any],
batch_torch_fn: Callable[[list[Any]], torch.Tensor],
warmup: int = 5,
) -> dict[str, float]:
"""Run simulate() N times, return mean per-stage us per batch.
GPU sync is excluded from per-iteration timing: production issues
`to(device, non_blocking=True)` and continues, so we measure issue
cost rather than H2D completion.
"""
kit = {
"ingest_fn": ingest_fn,
"empty_fn": empty_fn,
"batch_torch_fn": batch_torch_fn,
}
torch.cuda.synchronize()
for _ in range(warmup):
simulate(seeds, n_decode, **kit)
torch.cuda.synchronize()
accum: dict[str, float] = defaultdict(float)
for _ in range(iterations):
t = simulate(seeds, n_decode, **kit)
for k, v in t.items():
accum[k] += v
torch.cuda.synchronize()
return {k: accum[k] / iterations / 1000.0 for k in STAGES} # ns -> us
def print_breakdown(title: str, results: dict[str, dict[str, float]]) -> None:
"""Print per-stage timings with delta us vs the first (baseline) column."""
labels = list(results.keys())
baseline_label = labels[0]
baseline = results[baseline_label]
width = max(len(s) for s in STAGES)
header_cells = [f"{baseline_label + ' us':>11s}"]
for lbl in labels[1:]:
header_cells.append(f"{lbl + ' us':>11s}")
header_cells.append(f"{'delta':>10s}")
print(f"=== {title} ===")
print(f"{'Stage':<{width}s} " + " ".join(header_cells))
print("-" * (width + 2 + sum(len(c) + 2 for c in header_cells)))
for s in STAGES:
cells = [f"{baseline[s]:>11.3f}"]
for lbl in labels[1:]:
v = results[lbl][s]
cells.append(f"{v:>11.3f}")
d = v - baseline[s]
cells.append(f"{d:>+10.3f}")
print(f"{s:<{width}s} " + " ".join(cells))
print("-" * (width + 2 + sum(len(c) + 2 for c in header_cells)))
base_total = sum(baseline.values())
total_cells = [f"{base_total:>11.3f}"]
for lbl in labels[1:]:
v = sum(results[lbl].values())
total_cells.append(f"{v:>11.3f}")
d = v - base_total
total_cells.append(f"{d:>+10.3f}")
print(f"{'TOTAL':<{width}s} " + " ".join(total_cells))
print()
for lbl in labels[1:]:
v = sum(results[lbl].values())
d = v - base_total
speedup = base_total / v if v > 0 else 0.0
verdict = "LOSES" if d > 0 else "WINS"
print(
f" {lbl:<14s} vs {baseline_label}: {verdict} by {abs(d):>8.2f} us ({speedup:.2f}x)"
)
print()
def microbench_torch_tensor_paths(
sizes: tuple[int, ...] = (1_000, 10_000, 100_000)
) -> None:
"""Compare three CPU-buffer -> pinned cuda tensor paths.
A. torch.tensor(list, pin) -> cuda
B. torch.tensor(array('q'), pin) -> cuda
C. torch.from_numpy(np.frombuffer(array('q'))).pin() -> cuda
"""
def t(fn, iterations: int) -> float:
for _ in range(20):
fn()
torch.cuda.synchronize()
t0 = time.monotonic_ns()
for _ in range(iterations):
fn()
torch.cuda.synchronize()
return (time.monotonic_ns() - t0) / iterations / 1000.0
print("=== microbench: CPU-buffer -> pinned cuda tensor (us/op) ===\n")
width = 56
print(f"{'Path':<{width}s} " + " ".join(f"{f'N={n}':>10s}" for n in sizes))
print("-" * (width + 2 + 12 * len(sizes)))
for label, build in [
(
"(A) torch.tensor(list, pin) -> cuda",
lambda x: torch.tensor(x, dtype=torch.int64, pin_memory=True).to(
"cuda", non_blocking=True
),
),
(
"(B) torch.tensor(array('q'), pin) -> cuda (naive)",
lambda x: torch.tensor(x, dtype=torch.int64, pin_memory=True).to(
"cuda", non_blocking=True
),
),
(
"(C) from_numpy(frombuf(array('q'))).pin() -> cuda",
lambda x: torch.from_numpy(np.frombuffer(x, dtype=np.int64))
.pin_memory()
.to("cuda", non_blocking=True),
),
]:
cells = []
for n in sizes:
iters = max(50, 200_000 // max(n, 1))
if "(A)" in label:
src = list(range(n))
else:
src = array("q", range(n))
us = t(lambda src=src, build=build: build(src), iters)
cells.append(f"{us:>10.2f}")
print(f"{label:<{width}s} " + " ".join(cells))
print()
def main() -> None:
microbench_torch_tensor_paths()
n_reqs = 2
cases = [
("short prompt N_origin=1K N_decode=1K", 1_000, 1_000, 1_000),
("medium prompt N_origin=10K N_decode=1K", 10_000, 1_000, 200),
("long prompt N_origin=100K N_decode=1K", 100_000, 1_000, 30),
]
print(f"Batch size = {n_reqs} reqs/batch (per-req stages accumulate)\n")
for label, n_origin, n_decode, iters in cases:
seeds = [list(range(n_origin)) for _ in range(n_reqs)]
results = {
"list": bench_lifecycle(seeds, n_decode, iters, **LIST_KIT),
"pyarray": bench_lifecycle(seeds, n_decode, iters, **PYARRAY_KIT),
}
print_breakdown(label, results)
if __name__ == "__main__":
main()
@@ -1,6 +1,7 @@
from __future__ import annotations
import logging
from array import array
from http import HTTPStatus
from typing import TYPE_CHECKING, List
@@ -71,7 +72,7 @@ class ScheduleBatchDisaggregationDecodeMixin:
# Set fields
self.input_ids = torch.tensor(
sum(input_ids, []), dtype=torch.int32, device=self.device
sum(input_ids, array("q")), dtype=torch.int32, device=self.device
)
self.req_pool_indices = torch.tensor(
req_pool_indices, dtype=torch.int64, device=self.device
@@ -7,6 +7,7 @@ import threading
import time
import uuid
from abc import ABC, abstractmethod
from array import array
from collections import OrderedDict, defaultdict
from enum import IntEnum
from http import HTTPStatus
@@ -588,7 +589,7 @@ class WaitingImageRequest:
**self.recv_embedding_data.get_mm_extra_meta(),
)
self.recv_req.mm_inputs = mm_inputs
self.recv_req.input_ids = mm_inputs.input_ids
self.recv_req.input_ids = array("q", mm_inputs.input_ids)
self.status = WaitingImageRequestStatus.SUCCESS
self.recv_socket.close()
+2 -1
View File
@@ -1,6 +1,7 @@
from __future__ import annotations
import enum
from array import array
from typing import TYPE_CHECKING, Optional
from sglang.srt.dllm.config import DllmConfig
@@ -62,7 +63,7 @@ class ReqDllmMixin:
self.fill_ids = (
self.origin_input_ids
+ self.output_ids
+ [self.dllm_config.mask_id] * self.dllm_config.block_size
+ array("q", [self.dllm_config.mask_id] * self.dllm_config.block_size)
)
def _update_block_offset_for_dllm(self):
+2 -1
View File
@@ -1,6 +1,7 @@
from __future__ import annotations
import logging
from array import array
from typing import TYPE_CHECKING, List, Optional, Set, Union
from sglang.srt.dllm.config import DllmConfig
@@ -79,7 +80,7 @@ class SchedulerDllmMixin:
if new_tokens == 0:
continue
req.fill_ids[-new_tokens:] = next_token_ids[:]
req.fill_ids[-new_tokens:] = array("q", next_token_ids)
self.metrics_reporter.num_generated_tokens += new_tokens
req.output_ids.extend(next_token_ids)
@@ -238,7 +238,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
if rid not in self.decode_status:
s = DecodeStatus(
decoded_text=recv_obj.decoded_texts[i],
decode_ids=recv_obj.decode_ids[i],
decode_ids=list(recv_obj.decode_ids[i]),
surr_offset=0,
read_offset=recv_obj.read_offsets[i],
)
+5 -4
View File
@@ -21,6 +21,7 @@ from __future__ import annotations
import copy
import uuid
from abc import ABC
from array import array
from collections import Counter
from dataclasses import dataclass, field
from enum import Enum
@@ -712,7 +713,7 @@ class TokenizedGenerateReqInput(BaseReq):
# The input text
input_text: str
# The input token ids
input_ids: List[int]
input_ids: Optional[array[int]]
# The multimodal inputs
mm_inputs: object
# The sampling parameters
@@ -1027,7 +1028,7 @@ class TokenizedEmbeddingReqInput(BaseReq):
# The input text
input_text: str
# The input token ids
input_ids: List[int]
input_ids: array[int]
# The image inputs
image_inputs: dict
# The token type ids
@@ -1075,10 +1076,10 @@ class BatchTokenIDOutput(BaseBatchReq, SpeculativeDecodingMetricsMixin):
finished_reasons: List[BaseFinishReason]
# For incremental decoding
decoded_texts: List[str]
decode_ids: List[int]
decode_ids: List[array[int]]
read_offsets: List[int]
# Only used when `--skip-tokenizer-init` is on
output_ids: Optional[List[int]]
output_ids: Optional[List[array[int]]]
# Detokenization configs
skip_special_tokens: List[bool]
spaces_between_special_tokens: List[bool]
+22 -17
View File
@@ -2,7 +2,11 @@ from __future__ import annotations
from sglang.srt.dllm.config import DllmConfig
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.utils.common import ceil_align, is_pin_memory_available
from sglang.srt.utils.common import (
ceil_align,
flatten_arrays_to_int64_tensor,
is_pin_memory_available,
)
# Copyright 2023-2024 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
@@ -35,11 +39,11 @@ import copy
import dataclasses
import logging
import re
from array import array
from concurrent.futures import Future
from enum import Enum, auto
from functools import lru_cache
from http import HTTPStatus
from itertools import chain
from typing import (
TYPE_CHECKING,
Any,
@@ -611,14 +615,14 @@ class Req(ReqDllmMixin):
self,
rid: str,
origin_input_text: str,
origin_input_ids: List[int],
origin_input_ids: array[int],
sampling_params: SamplingParams,
return_logprob: bool = False,
top_logprobs_num: int = 0,
dllm_config: Optional[DllmConfig] = None,
token_ids_logprob: List[int] = None,
stream: bool = False,
origin_input_ids_unpadded: Optional[Tuple[int]] = None,
origin_input_ids_unpadded: Optional[array[int]] = None,
lora_id: Optional[str] = None,
input_embeds: Optional[List[List[float]]] = None,
positional_embed_overrides: Optional[PositionalEmbeds] = None,
@@ -659,9 +663,10 @@ class Req(ReqDllmMixin):
)
self.origin_input_ids = origin_input_ids
# Each decode stage's output ids
self.output_ids = []
self.output_ids = array("q")
# fill_ids = origin_input_ids + output_ids. Updated if chunked.
self.fill_ids = []
self.fill_ids = array("q")
self.session = session
self.input_embeds = input_embeds
self.positional_embed_overrides = positional_embed_overrides
@@ -948,7 +953,7 @@ class Req(ReqDllmMixin):
return self.sampling_params.max_new_tokens == 0 and spec_alg is None
@property
def output_ids_through_stop(self) -> List[int]:
def output_ids_through_stop(self) -> array[int]:
"""Get the output ids through the stop condition. Stop position is included."""
if self.finished_len is not None:
return self.output_ids[: self.finished_len]
@@ -1037,7 +1042,7 @@ class Req(ReqDllmMixin):
# Disable prefix caching when embed overrides are present: same token IDs
# with different override vectors must not share cached KV values.
if self.positional_embed_overrides is not None:
token_ids_to_match = []
token_ids_to_match = array("q")
if tree_cache is not None:
if cow_mamba is None:
@@ -1305,7 +1310,7 @@ class Req(ReqDllmMixin):
# Therefore, we discard the generated output_ids and restart prefill and generation
# to ensure shape consistency in KV cache.
if self.input_embeds is not None:
self.output_ids = []
self.output_ids = array("q")
def offload_kv_cache(self, req_to_token_pool, token_to_kv_pool_allocator):
token_indices = req_to_token_pool.req_to_token[
@@ -1370,7 +1375,9 @@ class Req(ReqDllmMixin):
logger.error(f"{error_msg}, {self.rid=}")
self.multimodal_inputs = None
self.grammar = None
self.origin_input_ids = [0] # set it to one token to skip the long prefill
self.origin_input_ids = array(
"q", [0]
) # set it to one token to skip the long prefill
self.return_logprob = False
self.logprob_start_len = -1
self.to_finish = FINISH_ABORT(
@@ -1629,7 +1636,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
def is_dllm(self):
return self.dllm_config is not None
def prepare_encoder_info_extend(self, input_ids: List[int], seq_lens: List[int]):
def prepare_encoder_info_extend(
self, input_ids: List[array[int]], seq_lens: List[int]
):
_pin = is_pin_memory_available(self.device)
self.encoder_lens_cpu = []
self.encoder_cached = []
@@ -1678,9 +1687,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
pt += req.extend_input_len
# Reassign
self.input_ids = torch.tensor(
sum(input_ids, []), dtype=torch.int64, pin_memory=_pin
).to(self.device, non_blocking=True)
self.input_ids = flatten_arrays_to_int64_tensor(input_ids, self.device, _pin)
self.seq_lens = torch.tensor(seq_lens, dtype=torch.int64, pin_memory=_pin).to(
self.device, non_blocking=True
)
@@ -1783,9 +1790,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
]
_pin = is_pin_memory_available(self.device)
input_ids_tensor = torch.tensor(
list(chain.from_iterable(input_ids)), dtype=torch.int64, pin_memory=_pin
).to(self.device, non_blocking=True)
input_ids_tensor = flatten_arrays_to_int64_tensor(input_ids, self.device, _pin)
seq_lens_tensor = torch.tensor(seq_lens, dtype=torch.int64, pin_memory=_pin).to(
self.device, non_blocking=True
)
@@ -1,6 +1,7 @@
from __future__ import annotations
import logging
from array import array
from sglang.srt.environ import envs
from sglang.srt.managers.prefill_delayer import PrefillDelayerSinglePassExecutor
@@ -84,7 +85,7 @@ IGNORE_EOS_RESERVE_TOKENS = 1
def match_prefix_for_req(
tree_cache: BasePrefixCache,
req: Req,
token_ids: Optional[List[int]] = None,
token_ids: Optional[array[int]] = None,
*,
cow_mamba: bool = False,
include_req: bool = False,
+8 -6
View File
@@ -20,6 +20,7 @@ import os
import signal
import sys
import time
from array import array
from collections import deque
from contextlib import contextmanager, nullcontext
from functools import partial
@@ -959,6 +960,7 @@ class Scheduler(
)
self.dp_tp_cpu_group = self.dp_tp_group.cpu_group
# TODO(Jialin): Migrate pad_input_ids implementations to return array.
self.pad_input_ids_func = self.tp_worker.get_pad_input_ids_func()
set_random_seed(self.random_seed)
@@ -1782,8 +1784,7 @@ class Scheduler(
if recv_req.input_embeds is not None:
# Generate fake input_ids based on the length of input_embeds
seq_length = len(recv_req.input_embeds)
fake_input_ids = [1] * seq_length
recv_req.input_ids = fake_input_ids
recv_req.input_ids = array("q", [1]) * seq_length
if recv_req.bootstrap_port is None:
# Use default bootstrap port
@@ -1909,8 +1910,8 @@ class Scheduler(
# Expand a single image token into multiple dummy tokens for receiving image embeddings.
# The pad function is model-specific and can be None for some backends.
if self.pad_input_ids_func:
req.origin_input_ids = self.pad_input_ids_func(
req.origin_input_ids, image_inputs
req.origin_input_ids = array(
"q", self.pad_input_ids_func(req.origin_input_ids, image_inputs)
)
req.extend_image_inputs(image_inputs)
self._maybe_compute_mrope_positions(req)
@@ -2182,8 +2183,9 @@ class Scheduler(
# embedding models or models not requiring special padding.
# If None, `req.origin_input_ids` is expected to be correctly populated already.
if self.pad_input_ids_func:
req.origin_input_ids = self.pad_input_ids_func(
req.origin_input_ids, image_inputs
# See companion call site above for the array.array wrap rationale.
req.origin_input_ids = array(
"q", self.pad_input_ids_func(req.origin_input_ids, image_inputs)
)
req.extend_image_inputs(image_inputs)
@@ -3,6 +3,7 @@ from __future__ import annotations
import logging
import math
import time
from array import array
from collections import defaultdict, deque
from dataclasses import dataclass
from typing import TYPE_CHECKING, Dict, List, Optional, Tuple
@@ -556,7 +557,7 @@ class SchedulerPPMixin:
if self.pp_group.is_first_rank:
model_runner = self.tp_worker.model_runner
model_config = model_runner.model_config
input_ids_list = []
input_ids_list: List[array[int]] = []
for i in range(128):
chunk_size = int(
self.chunked_prefill_size * 1.25
@@ -564,9 +565,12 @@ class SchedulerPPMixin:
)
if chunk_size <= 0:
break
input_ids = np.random.randint(
0, 10000, size=chunk_size, dtype=np.int64
).tolist()
input_ids = array(
"q",
np.random.randint(
0, 10000, size=chunk_size, dtype=np.int64
).tobytes(),
)
input_ids_list.append(input_ids)
sampling_params = SamplingParams(
@@ -13,6 +13,8 @@
# ==============================================================================
"""TokenizerManager is a process that tokenizes the text."""
from __future__ import annotations
import asyncio
import copy
import dataclasses
@@ -24,6 +26,7 @@ import signal
import socket
import sys
import threading
from array import array
from collections import deque
from contextlib import nullcontext
from datetime import datetime
@@ -987,12 +990,15 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self,
obj: Union[GenerateReqInput, EmbeddingReqInput],
input_text: str,
input_ids: List[int],
input_ids: Optional[List[int]],
input_embeds: Optional[Union[List[float], None]] = None,
mm_inputs=None,
token_type_ids: Optional[List[int]] = None,
) -> Union[TokenizedGenerateReqInput, TokenizedEmbeddingReqInput]:
"""Create a tokenized request object from common parameters."""
input_ids_arr: Optional[array[int]] = (
array("q", input_ids) if input_ids is not None else None
)
# Parse sampling parameters
# Note: if there are preferred sampling params, we use them if they are not
# explicitly passed in sampling_params
@@ -1020,7 +1026,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
tokenized_obj = TokenizedGenerateReqInput(
input_text,
input_ids,
input_ids_arr,
mm_inputs,
sampling_params,
obj.return_logprob,
@@ -1062,12 +1068,12 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
and obj.embed_override_token_id is not None
):
positional_embed_overrides = self._resolve_embed_overrides(
input_ids, obj.embed_override_token_id, obj.embed_overrides
input_ids_arr, obj.embed_override_token_id, obj.embed_overrides
)
tokenized_obj = TokenizedEmbeddingReqInput(
input_text,
input_ids,
input_ids_arr,
mm_inputs,
token_type_ids,
sampling_params,
@@ -1088,7 +1094,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
@staticmethod
def _resolve_embed_overrides(
input_ids: List[int],
input_ids: array[int],
token_id: int,
embeds: List[torch.Tensor],
) -> PositionalEmbeds:
@@ -1787,7 +1793,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self.server_args.incremental_streaming_output and is_stream
)
delta_text = recv_obj.output_strs[i]
delta_output_ids = recv_obj.output_ids[i]
delta_output_ids = list(recv_obj.output_ids[i])
output_offset = state.last_output_offset
state.append_text(delta_text)
state.output_ids.extend(delta_output_ids)
@@ -1830,7 +1836,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
incremental = (
self.server_args.incremental_streaming_output and is_stream
)
delta_output_ids = recv_obj.output_ids[i]
delta_output_ids = list(recv_obj.output_ids[i])
output_offset = state.last_output_offset
state.output_ids.extend(delta_output_ids)
+1 -1
View File
@@ -65,7 +65,7 @@ class KVCacheEventMixin:
if is_bigram:
page_tokens = [(raw[j], raw[j + 1]) for j in range(start, end)]
else:
page_tokens = raw[start:end]
page_tokens = list(raw[start:end])
block_hash = hash_str_to_int64(node.hash_value[page_index])
@@ -804,10 +804,6 @@ class HiRadixCache(RadixCache):
def evictable_size(self):
return self.evictable_size_
def _to_radix_key(self, token_ids: List[int]) -> RadixKey:
"""Convert raw token_ids to a RadixKey; must be list (not tuple) for paged match."""
return RadixKey(token_ids=list(token_ids))
def inc_lock_ref(self, node: TreeNode) -> IncLockRefResult:
if self.disable:
return IncLockRefResult(delta=0)
@@ -20,6 +20,7 @@ The radix tree data structure for managing the hybrid (full and Mamba) KV cache.
"""
import heapq
from array import array
from collections import defaultdict
from functools import lru_cache
from typing import TYPE_CHECKING, List, Optional, Tuple
@@ -456,7 +457,7 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache):
def reset(self) -> None:
self.root_node = TreeNode()
self.root_node.key = RadixKey([], None)
self.root_node.key = RadixKey(array("q"), None)
self.root_node.value = []
self.root_node.hash_value = []
self.root_node.full_lock_ref = 1
+12 -14
View File
@@ -26,6 +26,7 @@ import heapq
import logging
import sys
import time
from array import array
from collections import defaultdict
from functools import lru_cache
from typing import TYPE_CHECKING, Any, Iterator, List, Optional, Tuple, Union
@@ -70,7 +71,7 @@ class RadixKey:
def __init__(
self,
token_ids: List[int],
token_ids: array[int],
extra_key: Optional[str] = None,
is_bigram: bool = False,
):
@@ -87,6 +88,7 @@ class RadixKey:
return n - 1 if n > 0 else 0
return len(self.token_ids)
# TODO(Jialin): vectorize with numpy without PyLong boxing
def __iter__(self) -> Iterator:
if self.is_bigram:
t = self.token_ids
@@ -110,7 +112,7 @@ class RadixKey:
if self.is_bigram:
# bigrams [start, stop) span raw tokens [start, stop + 1);
# empty slice -> empty raw tokens (not a dangling boundary token).
raw = self.token_ids[start : stop + 1] if stop > start else []
raw = self.token_ids[start : stop + 1] if stop > start else array("q")
return RadixKey(raw, self.extra_key, is_bigram=True)
return RadixKey(self.token_ids[start:stop], self.extra_key)
@@ -144,6 +146,7 @@ class RadixKey:
f"{self.extra_key=} != {other.extra_key=}"
)
# TODO(Jialin): replace zip with numpy to skip per-element PyLong boxing
def match(self, other: "RadixKey", page_size: int = 1) -> int:
"""Logical-unit prefix length shared with ``other``. Result is rounded down to ``page_size``."""
self._check_compatible(other)
@@ -337,7 +340,7 @@ class RadixCache(KVCacheEventMixin, BasePrefixCache):
def reset(self):
# Initialize root with minimum priority so any real priority overrides it
self.root_node = TreeNode(priority=-sys.maxsize)
self.root_node.key = RadixKey(token_ids=[], extra_key=None)
self.root_node.key = RadixKey(token_ids=array("q"), extra_key=None)
self.root_node.value = []
self.root_node.host_value = []
self.root_node.lock_ref = 1
@@ -811,20 +814,15 @@ class RadixCache(KVCacheEventMixin, BasePrefixCache):
if __name__ == "__main__":
tree = RadixCache.create_simulated()
# Example token id sequences (as lists of ints)
tree.insert(InsertParams(key=RadixKey(token_ids=[1, 2, 3], extra_key=None)))
tree.insert(InsertParams(key=RadixKey(token_ids=[1, 2, 3], extra_key=None)))
tree.insert(InsertParams(key=RadixKey(token_ids=[1, 2, 4, 5], extra_key=None)))
tree.insert(
InsertParams(key=RadixKey(token_ids=[1, 2, 4, 5, 6, 7], extra_key=None))
)
tree.insert(
InsertParams(key=RadixKey(token_ids=[8, 9, 10, 11, 12], extra_key=None))
)
tree.insert(InsertParams(key=RadixKey(token_ids=array("q", [1, 2, 3]))))
tree.insert(InsertParams(key=RadixKey(token_ids=array("q", [1, 2, 3]))))
tree.insert(InsertParams(key=RadixKey(token_ids=array("q", [1, 2, 4, 5]))))
tree.insert(InsertParams(key=RadixKey(token_ids=array("q", [1, 2, 4, 5, 6, 7]))))
tree.insert(InsertParams(key=RadixKey(token_ids=array("q", [8, 9, 10, 11, 12]))))
tree.pretty_print()
print(
tree.match_prefix(
MatchPrefixParams(key=RadixKey(token_ids=[1, 2, 3, 13, 14], extra_key=None))
MatchPrefixParams(key=RadixKey(token_ids=array("q", [1, 2, 3, 13, 14])))
)
)
@@ -3,6 +3,7 @@ from __future__ import annotations
import logging
import threading
import time
from array import array
from collections import defaultdict
from functools import partial
from typing import TYPE_CHECKING, Any, Optional
@@ -257,7 +258,7 @@ class UnifiedRadixCache(BasePrefixCache):
def _reset_full(self) -> None:
"""Full reset: destroy entire tree and all state."""
self.root_node = UnifiedTreeNode(self.tree_components)
self.root_node.key = RadixKey([], None)
self.root_node.key = RadixKey(array("q"), None)
self.root_node.component_data[BASE_COMPONENT_TYPE].value = []
for ct in self.tree_components:
self.root_node.component_data[ct].lock_ref = 1
+8 -2
View File
@@ -13,8 +13,11 @@
# ==============================================================================
"""Inference-only LLaVa model compatible with HuggingFace weights."""
from __future__ import annotations
import math
import re
from array import array
from functools import lru_cache
from typing import Dict, Iterable, List, Optional, Tuple, Type, Union
@@ -73,7 +76,9 @@ class LlavaBaseForCausalLM(nn.Module):
return "pad"
return "anyres"
def pad_input_ids(self, input_ids: List[int], image_inputs: MultimodalInputs):
def pad_input_ids(
self, input_ids: array[int], image_inputs: MultimodalInputs
) -> array[int]:
image_sizes = flatten_nested_list(
[item.image_sizes for item in image_inputs.mm_items]
)
@@ -125,9 +130,10 @@ class LlavaBaseForCausalLM(nn.Module):
except ValueError:
offset = 0
# old_len + pad_len - 1, because we need to remove image_token_id
pad_token = pad_values[image_idx % len(pad_values)]
input_ids = (
input_ids[:offset]
+ [pad_values[image_idx % len(pad_values)]] * new_image_feature_len
+ array("q", [pad_token]) * new_image_feature_len
+ input_ids[offset + 1 :]
)
offset_list.append(offset)
+8 -3
View File
@@ -13,7 +13,10 @@
# ==============================================================================
"""Inference-only LLaVa video model compatible with HuggingFace weights."""
from typing import Iterable, List, Optional, Tuple
from __future__ import annotations
from array import array
from typing import Iterable, Optional, Tuple
import numpy as np
import torch
@@ -57,8 +60,10 @@ class LlavaVidForCausalLM(nn.Module):
torch.empty(config.text_config.hidden_size, dtype=torch.float16)
)
def pad_input_ids(self, input_ids: List[int], image_inputs: MultimodalInputs):
pad_values = [item.pad_value for item in image_inputs.mm_items]
def pad_input_ids(
self, input_ids: array[int], image_inputs: MultimodalInputs
) -> array[int]:
pad_values = array("q", (item.pad_value for item in image_inputs.mm_items))
new_image_feature_len = self.image_feature_len
pad_ids = pad_values * (
+7 -2
View File
@@ -4,7 +4,10 @@
# https://github.com/vllm-project/vllm/blob/7193774b1ff8603ad5bf4598e5efba0d9a39b436/vllm/model_executor/models/mllama.py
"""PyTorch Mllama model."""
from __future__ import annotations
import math
from array import array
from typing import Iterable, List, Optional, Tuple, Union
import torch
@@ -823,9 +826,11 @@ class MllamaForConditionalGeneration(nn.Module):
)
self.logits_processor = LogitsProcessor(config.text_config)
def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs):
def pad_input_ids(
self, input_ids: array[int], mm_inputs: MultimodalInputs
) -> array[int]:
pixel_values = torch.cat([item.feature for item in mm_inputs.mm_items], dim=0)
pad_values = [item.pad_value for item in mm_inputs.mm_items]
pad_values = array("q", (item.pad_value for item in mm_inputs.mm_items))
num_concurrent_media, num_tiles = pixel_values.shape[1:3]
num_patches = self.vision_model.num_patches
+7 -4
View File
@@ -3,6 +3,7 @@
from __future__ import annotations
import logging
from array import array
from functools import partial
from typing import Iterable, List, Optional, Tuple
@@ -1122,15 +1123,17 @@ class MossVLForConditionalGeneration(nn.Module):
return total_len
def _build_encoder_prefix_pad_ids(self, mm_inputs: MultimodalInputs) -> List[int]:
def _build_encoder_prefix_pad_ids(self, mm_inputs: MultimodalInputs) -> array[int]:
encoder_len = self._get_encoder_len(mm_inputs)
if encoder_len == 0 or not mm_inputs.mm_items:
return []
return array("q")
pad_value = mm_inputs.mm_items[0].pad_value
return [pad_value] * encoder_len
return array("q", [pad_value]) * encoder_len
def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs):
def pad_input_ids(
self, input_ids: array[int], mm_inputs: MultimodalInputs
) -> array[int]:
encoder_len = self._get_encoder_len(mm_inputs)
mm_inputs.num_image_tokens = encoder_len
if encoder_len == 0:
+8 -4
View File
@@ -1,4 +1,7 @@
from typing import Any, Iterable, List, Optional, Tuple
from __future__ import annotations
from array import array
from typing import Any, Iterable, Optional, Tuple
import torch
from transformers import WhisperConfig
@@ -418,14 +421,15 @@ class WhisperForConditionalGeneration(torch.nn.Module):
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight)
def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs):
def pad_input_ids(
self, input_ids: array[int], mm_inputs: MultimodalInputs
) -> array[int]:
# Prepend dummy encoder tokens so that prepare_encoder_info_extend
# correctly allocates encoder KV cache locations in the KV pool.
# These dummy tokens are stripped before the model forward receives input_ids.
encoder_len = self.config.max_source_positions
mm_inputs.num_image_tokens = encoder_len
pad_ids = [0] * encoder_len
return pad_ids + input_ids
return array("q", [0]) * encoder_len + input_ids
def forward(
self,
+16
View File
@@ -45,6 +45,7 @@ import traceback
import types
import uuid
import warnings
from array import array
from collections import OrderedDict, defaultdict
from contextlib import contextmanager
from dataclasses import dataclass
@@ -101,6 +102,21 @@ logger = logging.getLogger(__name__)
torch_release = pkg_version.parse(torch.__version__).release
def flatten_arrays_to_int64_tensor(
parts: List[array[int]], device, pin: bool
) -> torch.Tensor:
"""Flatten a list of array.array('q') buffers into one int64 tensor.
Uses NumPy here to speed up the conversion by using memcpy
instead of a per-element PyLong-to-int64 walk.
"""
combined = np.concatenate([np.frombuffer(p, dtype=np.int64) for p in parts])
cpu_t = torch.from_numpy(combined)
if pin:
cpu_t = cpu_t.pin_memory()
return cpu_t.to(device, non_blocking=True)
# https://pytorch.org/docs/stable/notes/hip.html#checking-for-hip
@lru_cache(maxsize=1)
def is_hip() -> bool:
@@ -1,6 +1,7 @@
"""Regression tests for the SWA chunked-req stash gate (#24252)."""
import unittest
from array import array
from types import SimpleNamespace
from unittest.mock import MagicMock
@@ -27,9 +28,9 @@ def _make_req(
) -> Req:
req = Req.__new__(Req)
req.rid = "test-req"
req.origin_input_ids = list(fill_ids)
req.output_ids = []
req.fill_ids = list(fill_ids)
req.origin_input_ids = array("q", fill_ids)
req.output_ids = array("q")
req.fill_ids = array("q", fill_ids)
req.prefix_indices = prefix_indices
req.req_pool_idx = req_pool_idx
req.extend_input_len = extend_input_len
@@ -26,6 +26,7 @@ register_cuda_ci(est_time=10, stage="base-b", runner_config="1-gpu-small")
register_amd_ci(est_time=10, suite="stage-b-test-1-gpu-small-amd")
import unittest
from array import array
from unittest.mock import MagicMock
import torch
@@ -65,11 +66,11 @@ class MockReq:
"""Minimal mock Req with fields needed by cache_unfinished/finished_req."""
def __init__(self, fill_ids, req_pool_idx=0, cache_protected_len=0, last_node=None):
self.fill_ids = list(fill_ids)
self.origin_input_ids = (
list(fill_ids[:-1]) if len(fill_ids) > 1 else list(fill_ids)
self.fill_ids = array("q", fill_ids)
self.origin_input_ids = array(
"q", fill_ids[:-1] if len(fill_ids) > 1 else fill_ids
)
self.output_ids = [fill_ids[-1]] if len(fill_ids) > 1 else []
self.output_ids = array("q", [fill_ids[-1]] if len(fill_ids) > 1 else [])
self.req_pool_idx = req_pool_idx
self.cache_protected_len = cache_protected_len
self.last_node = last_node
@@ -99,7 +100,7 @@ class TestDecodeLockRefScenarios(unittest.TestCase):
"""Insert a prefix into the tree so future requests can match it."""
cache.insert(
InsertParams(
key=RadixKey(prefix_ids),
key=RadixKey(array("q", prefix_ids)),
value=torch.tensor(prefix_values, dtype=torch.int64),
)
)
@@ -119,7 +120,7 @@ class TestDecodeLockRefScenarios(unittest.TestCase):
self._populate_prefix(cache, prefix, prefix_vals)
# Match prefix (simulates _match_prefix_and_lock in pop_preallocated)
result = cache.match_prefix(MatchPrefixParams(key=RadixKey(prefix)))
result = cache.match_prefix(MatchPrefixParams(key=RadixKey(array("q", prefix))))
matched_node = result.last_device_node
prefix_len = len(result.device_indices)
self.assertEqual(prefix_len, 3)
@@ -164,7 +165,9 @@ class TestDecodeLockRefScenarios(unittest.TestCase):
# No prefix in tree -- match returns root
full_ids = [10, 20, 30]
result = cache.match_prefix(MatchPrefixParams(key=RadixKey(full_ids)))
result = cache.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", full_ids)))
)
matched_node = result.last_device_node
self.assertEqual(len(result.device_indices), 0) # no match
# matched_node is root
@@ -212,7 +215,7 @@ class TestDecodeLockRefScenarios(unittest.TestCase):
self._populate_prefix(cache, prefix, prefix_vals)
# Match and lock
result = cache.match_prefix(MatchPrefixParams(key=RadixKey(prefix)))
result = cache.match_prefix(MatchPrefixParams(key=RadixKey(array("q", prefix))))
matched_node = result.last_device_node
prefix_len = len(result.device_indices)
@@ -256,7 +259,9 @@ class TestDecodeLockRefScenarios(unittest.TestCase):
# No prefix in tree -- match returns root (simulates _match_prefix_and_lock)
full_ids = [10, 20, 30]
result = cache.match_prefix(MatchPrefixParams(key=RadixKey(full_ids)))
result = cache.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", full_ids)))
)
matched_node = result.last_device_node
self.assertIs(matched_node, cache.root_node)
@@ -356,7 +361,9 @@ class TestDecodeLockRefScenarios(unittest.TestCase):
self._populate_prefix(cache, prefix, prefix_vals)
for iteration in range(5):
result = cache.match_prefix(MatchPrefixParams(key=RadixKey(prefix)))
result = cache.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", prefix)))
)
matched_node = result.last_device_node
prefix_len = len(result.device_indices)
@@ -1,4 +1,5 @@
import unittest
from array import array
import torch
@@ -116,7 +117,7 @@ class TestMamba(unittest.TestCase):
req = Req(
rid=0,
origin_input_text="",
origin_input_ids=[],
origin_input_ids=array("q"),
sampling_params=sampling_params,
)
@@ -158,7 +159,7 @@ class TestMamba(unittest.TestCase):
print(
f"req1: inserting, req1_token_ids: {req1_token_ids}, req1_kv_indices: {req1_kv_indices}"
)
key = RadixKey(req1_token_ids)
key = RadixKey(array("q", req1_token_ids))
result = tree.insert(
InsertParams(
key=key,
@@ -176,7 +177,7 @@ class TestMamba(unittest.TestCase):
print(
f"req2: inserting, req2_token_ids: {req2_token_ids}, req2_kv_indices: {req2_kv_indices}"
)
key = RadixKey(req2_token_ids)
key = RadixKey(array("q", req2_token_ids))
result = tree.insert(
InsertParams(
key=key,
@@ -195,7 +196,7 @@ class TestMamba(unittest.TestCase):
print(
f"req3: inserting, req3_token_ids: {req3_token_ids}, req3_kv_indices: {req3_kv_indices}"
)
key = RadixKey(req3_token_ids)
key = RadixKey(array("q", req3_token_ids))
result = tree.insert(
InsertParams(
key=key,
@@ -213,7 +214,7 @@ class TestMamba(unittest.TestCase):
print(
f"req4: inserting, req4_token_ids: {req4_token_ids}, req4_kv_indices: {req4_kv_indices}"
)
key = RadixKey(req4_token_ids)
key = RadixKey(array("q", req4_token_ids))
result = tree.insert(
InsertParams(
key=key,
@@ -244,7 +245,9 @@ class TestMamba(unittest.TestCase):
tree.pretty_print()
req5_token_ids = [1, 2, 3, 4, 5]
result = tree.match_prefix(MatchPrefixParams(key=RadixKey(req5_token_ids)))
result = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", req5_token_ids)))
)
kv_indices, last_node = result.device_indices, result.last_device_node
print(
f"req5: token_ids: {req5_token_ids}, matched kv_indices: {kv_indices}, last_node.key: {last_node.key}"
@@ -252,7 +255,9 @@ class TestMamba(unittest.TestCase):
assert len(kv_indices) == 0
req6_token_ids = [1, 2, 3, 4, 5, 60, 70]
result = tree.match_prefix(MatchPrefixParams(key=RadixKey(req6_token_ids)))
result = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", req6_token_ids)))
)
kv_indices, last_node = result.device_indices, result.last_device_node
print(
f"req6: token_ids: {req6_token_ids}, matched kv_indices: {kv_indices}, last_node.key: {last_node.key}"
@@ -261,7 +266,9 @@ class TestMamba(unittest.TestCase):
assert len(last_node.key) == 2
req7_token_ids = [1, 2, 3, 4, 5, 6, 7]
result = tree.match_prefix(MatchPrefixParams(key=RadixKey(req7_token_ids)))
result = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", req7_token_ids)))
)
kv_indices, last_node = result.device_indices, result.last_device_node
print(
f"req7: token_ids: {req7_token_ids}, matched kv_indices: {kv_indices}, last_node.key: {last_node.key}"
@@ -278,7 +285,9 @@ class TestMamba(unittest.TestCase):
tree.pretty_print()
req8_token_ids = [1, 2, 3, 4, 5, 60, 70]
result = tree.match_prefix(MatchPrefixParams(key=RadixKey(req8_token_ids)))
result = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", req8_token_ids)))
)
kv_indices, last_node = result.device_indices, result.last_device_node
print(
f"req8: token_ids: {req8_token_ids}, matched kv_indices: {kv_indices}, last_node.key: {last_node.key}"
@@ -289,7 +298,9 @@ class TestMamba(unittest.TestCase):
req9_token_ids = [1, 2, 3, 4, 5, 6, 7]
req9 = make_dummy_req()
result = tree.match_prefix(
MatchPrefixParams(key=RadixKey(req9_token_ids), req=req9, cow_mamba=True)
MatchPrefixParams(
key=RadixKey(array("q", req9_token_ids)), req=req9, cow_mamba=True
)
)
kv_indices, last_node = result.device_indices, result.last_device_node
assert req9.mamba_pool_idx is not None
@@ -315,7 +326,7 @@ class TestMamba(unittest.TestCase):
stored_hashes = []
req1 = make_dummy_req()
key1 = RadixKey([1, 2, 3])
key1 = RadixKey(array("q", [1, 2, 3]))
tree.insert(
InsertParams(
key=key1,
@@ -330,7 +341,7 @@ class TestMamba(unittest.TestCase):
stored_hashes.extend(e.block_hashes[0] for e in stored_events)
req2 = make_dummy_req()
key2 = RadixKey([1, 2, 3, 4, 5])
key2 = RadixKey(array("q", [1, 2, 3, 4, 5]))
tree.insert(
InsertParams(
key=key2,
@@ -367,7 +378,7 @@ class TestMamba(unittest.TestCase):
tree.take_events() # Clear the reset event.
req1 = make_dummy_req()
key1 = RadixKey([1, 2, 3, 4])
key1 = RadixKey(array("q", [1, 2, 3, 4]))
tree.insert(
InsertParams(
key=key1,
@@ -382,7 +393,7 @@ class TestMamba(unittest.TestCase):
split_parent_hash = first_insert_events[1].block_hashes[0]
req2 = make_dummy_req()
key2 = RadixKey([1, 2, 5, 6])
key2 = RadixKey(array("q", [1, 2, 5, 6]))
tree.insert(
InsertParams(
key=key2,
@@ -394,7 +405,7 @@ class TestMamba(unittest.TestCase):
e for e in tree.take_events() if isinstance(e, BlockStored)
]
self.assertEqual(len(second_insert_events), 2)
self.assertEqual(second_insert_events[0].token_ids, [5])
self.assertEqual(list(second_insert_events[0].token_ids), [5])
self.assertEqual(second_insert_events[0].parent_block_hash, split_parent_hash)
def _setup_tree_and_allocator(self, enable_kv_cache_events=False):
@@ -478,7 +489,7 @@ class TestMamba(unittest.TestCase):
req = Req(
rid=0,
origin_input_text="",
origin_input_ids=[],
origin_input_ids=array("q"),
sampling_params=sampling_params,
)
req_to_token_pool.alloc([req])
@@ -492,9 +503,9 @@ class TestMamba(unittest.TestCase):
parent = TreeNode()
deleted = TreeNode()
root.key = RadixKey([])
parent.key = RadixKey([1])
deleted.key = RadixKey([2])
root.key = RadixKey(array("q", []))
parent.key = RadixKey(array("q", [1]))
deleted.key = RadixKey(array("q", [2]))
parent.parent = root
deleted.parent = parent
parent.value = torch.tensor([1], dtype=torch.int64)
@@ -668,7 +679,7 @@ class TestMamba(unittest.TestCase):
# Step 1: Insert [1,2,3] to create first node
req1 = make_dummy_req()
key1 = RadixKey([1, 2, 3])
key1 = RadixKey(array("q", [1, 2, 3]))
tree.insert(
InsertParams(
key=key1,
@@ -681,7 +692,7 @@ class TestMamba(unittest.TestCase):
# Step 2: Insert [1,2,3,4,5,6,7] with prev_prefix_len=0 (free all matched)
# Creates tree: [1,2,3] -> [4,5,6,7]
req2 = make_dummy_req()
key2 = RadixKey([1, 2, 3, 4, 5, 6, 7])
key2 = RadixKey(array("q", [1, 2, 3, 4, 5, 6, 7]))
result = tree.insert(
InsertParams(
key=key2,
@@ -699,7 +710,7 @@ class TestMamba(unittest.TestCase):
# Matched prefix = 7 (across two nodes: [1,2,3] len=3, [4,5,6,7] len=4)
# Protected [0..1], freed [2..6] = 5 slots, new [7] = 1 slot stored
req3 = make_dummy_req()
key3 = RadixKey([1, 2, 3, 4, 5, 6, 7, 8])
key3 = RadixKey(array("q", [1, 2, 3, 4, 5, 6, 7, 8]))
result = tree.insert(
InsertParams(
key=key3,
@@ -716,7 +727,7 @@ class TestMamba(unittest.TestCase):
# Step 4: Insert [1,2,3,4,5,6,7,8,9] with prev_prefix_len=8 (covers all matched)
# Matched prefix = 8, prev_prefix_len=8 => nothing freed
req4 = make_dummy_req()
key4 = RadixKey([1, 2, 3, 4, 5, 6, 7, 8, 9])
key4 = RadixKey(array("q", [1, 2, 3, 4, 5, 6, 7, 8, 9]))
result = tree.insert(
InsertParams(
key=key4,
@@ -1,4 +1,5 @@
import unittest
from array import array
import torch
@@ -63,7 +64,7 @@ class TestSLRUAccuracy(unittest.TestCase):
"""Test that SLRU eviction mechanism works correctly"""
# Insert one key-value three times (high frequency access)
frequent_key = RadixKey([1, 2]) # High hit rate, should be retained
frequent_key = RadixKey(array("q", [1, 2])) # High hit rate, should be retained
frequent_val = torch.tensor([10, 20], dtype=torch.int64)
# Insert the frequent key multiple times to increase its hit count
@@ -71,7 +72,9 @@ class TestSLRUAccuracy(unittest.TestCase):
self.cache.insert(InsertParams(key=frequent_key, value=frequent_val))
# Insert first low-frequency key-value pair that should be evicted
first_low_freq_key = RadixKey([5, 6]) # Low hit rate, should be evicted
first_low_freq_key = RadixKey(
array("q", [5, 6])
) # Low hit rate, should be evicted
first_low_freq_val = torch.tensor([50, 60], dtype=torch.int64)
self.cache.insert(
@@ -81,14 +84,14 @@ class TestSLRUAccuracy(unittest.TestCase):
# Insert other key-values once each (low frequency access) - fill up the cache
other_keys = []
for i in range(4): # Reduce the number to fit in our smaller cache
key = RadixKey([i + 10]) # Unique keys for low-frequency items
key = RadixKey(array("q", [i + 10])) # Unique keys for low-frequency items
val = torch.tensor([i + 100], dtype=torch.int64)
self.cache.insert(InsertParams(key=key, value=val))
other_keys.append(key)
# Now insert more items to trigger evictions
for i in range(6, 10): # Add more items to definitely exceed capacity
key = RadixKey([i * 2]) # Different pattern to avoid conflicts
key = RadixKey(array("q", [i * 2])) # Different pattern to avoid conflicts
val = torch.tensor([i * 200], dtype=torch.int64)
self.cache.insert(InsertParams(key=key, value=val))
@@ -28,6 +28,7 @@ import random
import time
import unittest
import unittest.mock
from array import array
import torch
@@ -50,30 +51,30 @@ class TestRadixKey(unittest.TestCase):
def test_init_basic(self):
"""Test basic initialization of RadixKey."""
token_ids = [1, 2, 3, 4]
key = RadixKey(token_ids)
self.assertEqual(key.token_ids, token_ids)
key = RadixKey(array("q", token_ids))
self.assertEqual(list(key.token_ids), token_ids)
self.assertIsNone(key.extra_key)
def test_init_with_extra_key(self):
"""Test initialization with extra_key."""
token_ids = [1, 2, 3]
extra_key = "test_key"
key = RadixKey(token_ids, extra_key)
self.assertEqual(key.token_ids, token_ids)
key = RadixKey(array("q", token_ids), extra_key)
self.assertEqual(list(key.token_ids), token_ids)
self.assertEqual(key.extra_key, extra_key)
def test_len(self):
"""Test __len__ method."""
key = RadixKey([1, 2, 3])
key = RadixKey(array("q", [1, 2, 3]))
self.assertEqual(len(key), 3)
empty_key = RadixKey([])
empty_key = RadixKey(array("q", []))
self.assertEqual(len(empty_key), 0)
def test_iter(self):
"""Test __iter__ method."""
token_ids = [1, 2, 3, 4]
key = RadixKey(token_ids)
key = RadixKey(array("q", token_ids))
self.assertEqual(list(key), token_ids)
def test_len_and_iter(self):
@@ -86,7 +87,7 @@ class TestRadixKey(unittest.TestCase):
for tokens, expected in test_cases:
with self.subTest(tokens=tokens):
key = RadixKey(tokens)
key = RadixKey(array("q", tokens))
self.assertEqual(len(key), expected)
self.assertEqual(list(key), tokens)
@@ -100,34 +101,34 @@ class TestRadixKey(unittest.TestCase):
for tokens, index, expected in test_cases:
with self.subTest(tokens=tokens, index=index):
key = RadixKey(tokens)
key = RadixKey(array("q", tokens))
result = key[index]
self.assertIsInstance(result, RadixKey)
self.assertEqual(result.token_ids, expected)
self.assertEqual(list(result.token_ids), expected)
def test_getitem_slice(self):
"""Test __getitem__ with slice and edge cases."""
key = RadixKey([1, 2, 3, 4, 5], "extra")
key = RadixKey(array("q", [1, 2, 3, 4, 5]), "extra")
# Basic slice
sliced = key[1:4]
self.assertIsInstance(sliced, RadixKey)
self.assertEqual(sliced.token_ids, [2, 3, 4])
self.assertEqual(list(sliced.token_ids), [2, 3, 4])
self.assertEqual(sliced.extra_key, "extra")
# Edge cases
self.assertEqual(key[2:2].token_ids, []) # Empty slice
self.assertEqual(key[:].token_ids, [1, 2, 3, 4, 5]) # Full slice
self.assertEqual(list(key[2:2].token_ids), []) # Empty slice
self.assertEqual(list(key[:].token_ids), [1, 2, 3, 4, 5]) # Full slice
def test_getitem_invalid_index(self):
"""Test __getitem__ with invalid indices."""
key = RadixKey([1, 2, 3])
key = RadixKey(array("q", [1, 2, 3]))
with self.assertRaises(IndexError):
_ = key[10] # Out of bounds
def test_repr(self):
"""Test __repr__ method."""
key = RadixKey([1, 2, 3], "test")
key = RadixKey(array("q", [1, 2, 3]), "test")
repr_str = repr(key)
self.assertIn("RadixKey", repr_str)
self.assertIn("extra_key='test'", repr_str)
@@ -136,7 +137,7 @@ class TestRadixKey(unittest.TestCase):
def test_repr_long_token_ids(self):
"""Test __repr__ with long token_ids."""
long_tokens = list(range(15))
key = RadixKey(long_tokens)
key = RadixKey(array("q", long_tokens))
repr_str = repr(key)
self.assertIn("...", repr_str) # Should be truncated
@@ -274,7 +275,7 @@ class TestRadixCache(unittest.TestCase):
# Insert some data
cache.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
key=RadixKey(array("q", [1, 2, 3])),
value=torch.tensor([10, 20, 30], dtype=torch.int64),
)
)
@@ -292,7 +293,7 @@ class TestRadixCache(unittest.TestCase):
with self.subTest(disable_cache=disable_cache):
cache = RadixCache.create_simulated(disable=disable_cache)
key = RadixKey([1, 2, 3])
key = RadixKey(array("q", [1, 2, 3]))
value = torch.tensor([10, 20, 30], dtype=torch.int64)
result = cache.insert(InsertParams(key=key, value=value))
prefix_len = result.prefix_len
@@ -307,12 +308,16 @@ class TestRadixCache(unittest.TestCase):
self.assertEqual(cache.evictable_size(), 3)
# Test match_prefix
result = cache.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3])))
result = cache.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", [1, 2, 3])))
)
self.assertEqual(len(result.device_indices), 3)
torch.testing.assert_close(result.device_indices, value)
# Test partial match
result = cache.match_prefix(MatchPrefixParams(key=RadixKey([1, 2])))
result = cache.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", [1, 2])))
)
self.assertEqual(len(result.device_indices), 2)
torch.testing.assert_close(
result.device_indices, torch.tensor([10, 20], dtype=torch.int64)
@@ -322,7 +327,7 @@ class TestRadixCache(unittest.TestCase):
"""Test insert with None value (should use token_ids as list)."""
cache = RadixCache.create_simulated()
key = RadixKey([1, 2, 3])
key = RadixKey(array("q", [1, 2, 3]))
result = cache.insert(InsertParams(key=key, value=None))
prefix_len = result.prefix_len
@@ -338,7 +343,7 @@ class TestRadixCache(unittest.TestCase):
cache.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
key=RadixKey(array("q", [1, 2, 3])),
value=torch.tensor([10, 20, 30], dtype=torch.int64),
)
)
@@ -346,7 +351,8 @@ class TestRadixCache(unittest.TestCase):
cache.insert(
InsertParams(
key=RadixKey([4, 5]), value=torch.tensor([40, 50], dtype=torch.int64)
key=RadixKey(array("q", [4, 5])),
value=torch.tensor([40, 50], dtype=torch.int64),
)
)
self.assertEqual(cache.total_size(), 5)
@@ -366,7 +372,9 @@ class TestRadixCache(unittest.TestCase):
)
# Insert data
cache.insert(InsertParams(key=RadixKey([1, 2, 3, 4, 5]), value=None))
cache.insert(
InsertParams(key=RadixKey(array("q", [1, 2, 3, 4, 5])), value=None)
)
# Take events
events = cache.take_events()
@@ -395,7 +403,7 @@ class TestRadixCache(unittest.TestCase):
# Insert and then evict data
cache.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
key=RadixKey(array("q", [1, 2, 3])),
value=torch.tensor([10, 20, 30], dtype=torch.int64),
)
)
@@ -427,29 +435,35 @@ class TestRadixCache(unittest.TestCase):
# Insert same token sequence with different extra keys
cache.insert(
InsertParams(
key=RadixKey([1, 2, 3], "key1"),
key=RadixKey(array("q", [1, 2, 3]), "key1"),
value=torch.tensor([10, 20, 30], dtype=torch.int64),
)
)
cache.insert(
InsertParams(
key=RadixKey([1, 2, 3], "key2"),
key=RadixKey(array("q", [1, 2, 3]), "key2"),
value=torch.tensor([40, 50, 60], dtype=torch.int64),
)
)
cache.insert(
InsertParams(
key=RadixKey([1, 2, 3], None),
key=RadixKey(array("q", [1, 2, 3]), None),
value=torch.tensor([70, 80, 90], dtype=torch.int64),
)
)
# Keys with different extra_key should not match each other
result1 = cache.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3], "key1")))
result2 = cache.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3], "key2")))
result3 = cache.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3], None)))
result1 = cache.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", [1, 2, 3]), "key1"))
)
result2 = cache.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", [1, 2, 3]), "key2"))
)
result3 = cache.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", [1, 2, 3]), None))
)
result4 = cache.match_prefix(
MatchPrefixParams(key=RadixKey([1, 2, 3], "nonexistent"))
MatchPrefixParams(key=RadixKey(array("q", [1, 2, 3]), "nonexistent"))
)
# Each should match only its own data
@@ -478,13 +492,15 @@ class TestRadixCache(unittest.TestCase):
# Insert sequence
cache.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
key=RadixKey(array("q", [1, 2, 3])),
value=torch.tensor([10, 20, 30], dtype=torch.int64),
)
)
# Get node
result = cache.match_prefix(MatchPrefixParams(key=RadixKey([1, 2, 3])))
result = cache.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", [1, 2, 3])))
)
node = result.last_device_node
initial_evictable = cache.evictable_size()
@@ -510,12 +526,14 @@ class TestRadixCache(unittest.TestCase):
# Insert sequences
cache.insert(
InsertParams(
key=RadixKey([1, 2]), value=torch.tensor([10, 20], dtype=torch.int64)
key=RadixKey(array("q", [1, 2])),
value=torch.tensor([10, 20], dtype=torch.int64),
)
)
cache.insert(
InsertParams(
key=RadixKey([3, 4]), value=torch.tensor([30, 40], dtype=torch.int64)
key=RadixKey(array("q", [3, 4])),
value=torch.tensor([30, 40], dtype=torch.int64),
)
)
@@ -547,7 +565,7 @@ class TestRadixCache(unittest.TestCase):
cache = RadixCache.create_simulated(page_size=page_size)
tokens = list(range(sequence_length))
key = RadixKey(tokens)
key = RadixKey(array("q", tokens))
cache.insert(
InsertParams(
key=key,
@@ -555,7 +573,9 @@ class TestRadixCache(unittest.TestCase):
)
)
result = cache.match_prefix(MatchPrefixParams(key=RadixKey(tokens)))
result = cache.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", tokens)))
)
self.assertGreater(len(result.device_indices), 0)
# Match length should be page-aligned
@@ -568,7 +588,7 @@ class TestRadixCache(unittest.TestCase):
cache.insert(
InsertParams(
key=RadixKey([1, 2, 3]),
key=RadixKey(array("q", [1, 2, 3])),
value=torch.tensor([10, 20, 30], dtype=torch.int64),
)
)
@@ -585,12 +605,14 @@ class TestRadixCache(unittest.TestCase):
cache.insert(
InsertParams(
key=RadixKey([1, 2]), value=torch.tensor([10, 20], dtype=torch.int64)
key=RadixKey(array("q", [1, 2])),
value=torch.tensor([10, 20], dtype=torch.int64),
)
)
cache.insert(
InsertParams(
key=RadixKey([3, 4]), value=torch.tensor([30, 40], dtype=torch.int64)
key=RadixKey(array("q", [3, 4])),
value=torch.tensor([30, 40], dtype=torch.int64),
)
)
@@ -609,12 +631,12 @@ class TestRadixCache(unittest.TestCase):
# Insert a long sequence that will be split later.
seq1 = [1, 2, 3, 4, 5, 6, 7, 8]
val1 = torch.tensor([x * 10 for x in seq1], dtype=torch.int64)
cache.insert(InsertParams(key=RadixKey(seq1), value=val1))
cache.insert(InsertParams(key=RadixKey(array("q", seq1)), value=val1))
# Insert a diverging branch to create an internal node on the path.
seq2 = [1, 2, 9, 10]
val2 = torch.tensor([x * 10 for x in seq2], dtype=torch.int64)
cache.insert(InsertParams(key=RadixKey(seq2), value=val2))
cache.insert(InsertParams(key=RadixKey(array("q", seq2)), value=val2))
print(cache.pretty_print())
baseline_total = cache.total_size()
@@ -624,24 +646,30 @@ class TestRadixCache(unittest.TestCase):
# Match that causes a split inside an existing node:
# take first 4 tokens of seq1, then diverge.
query1 = [1, 2, 3, 4, 999, 1000]
result1 = cache.match_prefix(MatchPrefixParams(key=RadixKey(query1)))
result1 = cache.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", query1)))
)
torch.testing.assert_close(result1.device_indices, val1[:4])
# No data change after structural split during matching.
self.assertEqual(cache.total_size(), baseline_total)
# Full match of the long sequence still returns the full indices.
result_full = cache.match_prefix(MatchPrefixParams(key=RadixKey(seq1)))
result_full = cache.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", seq1)))
)
torch.testing.assert_close(result_full.device_indices, val1)
# Another split deeper on the path (after matching 6 tokens, then diverge).
query2 = [1, 2, 3, 4, 5, 6, 777, 888]
result2 = cache.match_prefix(MatchPrefixParams(key=RadixKey(query2)))
result2 = cache.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", query2)))
)
torch.testing.assert_close(result2.device_indices, val1[:6])
self.assertEqual(cache.total_size(), baseline_total)
# Matching the short diverging branch should return exactly its indices.
result_branch = cache.match_prefix(
MatchPrefixParams(key=RadixKey(seq2))
MatchPrefixParams(key=RadixKey(array("q", seq2)))
)
torch.testing.assert_close(result_branch.device_indices, val2)
@@ -653,7 +681,9 @@ class TestRadixCache(unittest.TestCase):
)
# Insert a sequence
cache.insert(InsertParams(key=RadixKey([1, 2, 3, 4, 5, 6, 7, 8]), value=None))
cache.insert(
InsertParams(key=RadixKey(array("q", [1, 2, 3, 4, 5, 6, 7, 8])), value=None)
)
# Trigger event emission to compute hash_value lazily
cache.take_events()
@@ -679,7 +709,9 @@ class TestRadixCache(unittest.TestCase):
)
# Insert a sequence with repeating token pattern: [1,2,3,4, 1,2,3,4]
cache.insert(InsertParams(key=RadixKey([1, 2, 3, 4, 1, 2, 3, 4]), value=None))
cache.insert(
InsertParams(key=RadixKey(array("q", [1, 2, 3, 4, 1, 2, 3, 4])), value=None)
)
events = cache.take_events()
block_stored_events = [e for e in events if isinstance(e, BlockStored)]
@@ -713,11 +745,11 @@ class TestRadixCache(unittest.TestCase):
)
# Insert a sequence that will cause a split
cache.insert(InsertParams(key=RadixKey([1, 2, 3, 4]), value=None))
cache.insert(InsertParams(key=RadixKey(array("q", [1, 2, 3, 4])), value=None))
cache.take_events() # Clear events and compute hash_value for first node
# Insert a diverging sequence that will cause a split at page boundary
cache.insert(InsertParams(key=RadixKey([1, 2, 5, 6]), value=None))
cache.insert(InsertParams(key=RadixKey(array("q", [1, 2, 5, 6])), value=None))
cache.take_events() # Trigger event emission to compute hash_value
# Find the split node
@@ -754,7 +786,7 @@ class TestRadixCache(unittest.TestCase):
cache: RadixCache = RadixCache.create_simulated()
for key, value in zip(keys, values):
cache.insert(InsertParams(key=RadixKey(key), value=value))
cache.insert(InsertParams(key=RadixKey(array("q", key)), value=value))
del values
@@ -11,6 +11,7 @@ register_cpu_ci(est_time=5, suite="base-a-test-cpu")
import unittest
import unittest.mock
from array import array
import torch
@@ -27,8 +28,8 @@ from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey
class _StubReq:
def __init__(self, token_ids):
self.origin_input_ids = list(token_ids)
self.output_ids = []
self.origin_input_ids = array("q", token_ids)
self.output_ids = array("q")
self.extra_key = None
self.prefix_indices = None
self.last_node = None
@@ -42,9 +43,9 @@ class _StubReq:
class TestZeroMatchResult(unittest.TestCase):
def test_zero_replaces_indices_and_nodes(self):
tree = RadixCache.create_simulated()
tree.insert(InsertParams(key=RadixKey(token_ids=[1, 2, 3, 4, 5])))
tree.insert(InsertParams(key=RadixKey(token_ids=array("q", [1, 2, 3, 4, 5]))))
match = tree.match_prefix(
MatchPrefixParams(key=RadixKey(token_ids=[1, 2, 3, 9]))
MatchPrefixParams(key=RadixKey(token_ids=array("q", [1, 2, 3, 9])))
)
self.assertGreater(len(match.device_indices), 0)
zeroed = zero_match_result(tree, match)
@@ -76,7 +77,9 @@ class TestMatchPrefixForReqForceMiss(unittest.TestCase):
def test_force_miss_zeros_req_prefix(self):
tree = RadixCache.create_simulated()
tree.insert(
InsertParams(key=RadixKey(token_ids=[10, 11, 12, 13, 14, 15, 16, 17]))
InsertParams(
key=RadixKey(token_ids=array("q", [10, 11, 12, 13, 14, 15, 16, 17]))
)
)
# Sanity: without the flag, the same lookup hits.
@@ -12,6 +12,7 @@ Covers:
"""
import unittest
from array import array
import torch
@@ -110,6 +111,7 @@ def _swa_alloc(allocator, need_size):
def _insert_chain(tree, allocator, token_ids):
token_ids = array("q", token_ids)
indices = _swa_alloc(allocator, len(token_ids))
assert indices is not None
tree.insert(InsertParams(key=RadixKey(token_ids), value=indices))
@@ -1,4 +1,5 @@
import unittest
from array import array
import torch
@@ -113,12 +114,12 @@ def _swa_alloc(allocator, need_size):
def _insert(tree, allocator, token_ids):
indices = _swa_alloc(allocator, len(token_ids))
assert indices is not None
tree.insert(InsertParams(key=RadixKey(token_ids), value=indices))
tree.insert(InsertParams(key=RadixKey(array("q", token_ids)), value=indices))
def _insert_chain(tree, allocator, token_ids):
_insert(tree, allocator, token_ids)
match = tree.match_prefix(MatchPrefixParams(key=RadixKey(token_ids)))
match = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", token_ids))))
return match.last_device_node
@@ -193,7 +194,7 @@ class TestSWA(unittest.TestCase):
e for e in tree.take_events() if isinstance(e, BlockStored)
]
self.assertEqual(len(second_insert_events), 2)
self.assertEqual(second_insert_events[0].token_ids, [5])
self.assertEqual(list(second_insert_events[0].token_ids), [5])
self.assertEqual(second_insert_events[0].parent_block_hash, split_parent_hash)
def test_swa_memory_pool(self):
@@ -313,7 +314,7 @@ class TestSWA(unittest.TestCase):
print(
f"req1: inserting, req1_token_ids: {req1_token_ids}, req1_kv_indices: {req1_kv_indices}"
)
key = RadixKey(req1_token_ids)
key = RadixKey(array("q", req1_token_ids))
result = tree.insert(InsertParams(key=key, value=req1_kv_indices[: len(key)]))
prefix_len = result.prefix_len
print(
@@ -324,7 +325,7 @@ class TestSWA(unittest.TestCase):
print(
f"req2: inserting, req2_token_ids: {req2_token_ids}, req2_kv_indices: {req2_kv_indices}"
)
key = RadixKey(req2_token_ids)
key = RadixKey(array("q", req2_token_ids))
result = tree.insert(InsertParams(key=key, value=req2_kv_indices[: len(key)]))
prefix_len = result.prefix_len
print(
@@ -335,7 +336,7 @@ class TestSWA(unittest.TestCase):
print(
f"req3: inserting, req3_token_ids: {req3_token_ids}, req3_kv_indices: {req3_kv_indices}"
)
key = RadixKey(req3_token_ids)
key = RadixKey(array("q", req3_token_ids))
result = tree.insert(InsertParams(key=key, value=req3_kv_indices[: len(key)]))
prefix_len = result.prefix_len
print(
@@ -346,7 +347,7 @@ class TestSWA(unittest.TestCase):
print(
f"req4: inserting, req4_token_ids: {req4_token_ids}, req4_kv_indices: {req4_kv_indices}"
)
key = RadixKey(req4_token_ids)
key = RadixKey(array("q", req4_token_ids))
result = tree.insert(InsertParams(key=key, value=req4_kv_indices[: len(key)]))
prefix_len = result.prefix_len
print(
@@ -376,7 +377,9 @@ class TestSWA(unittest.TestCase):
tree.pretty_print()
req5_token_ids = [1, 2, 3, 4, 5]
result = tree.match_prefix(MatchPrefixParams(key=RadixKey(req5_token_ids)))
result = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", req5_token_ids)))
)
kv_indices, last_node = result.device_indices, result.last_device_node
print(
f"req5: token_ids: {req5_token_ids}, matched kv_indices: {kv_indices}, last_node.key: {last_node.key}"
@@ -384,7 +387,9 @@ class TestSWA(unittest.TestCase):
self.assertEqual(len(kv_indices), 0)
req6_token_ids = [1, 2, 3, 4, 5, 60, 70]
result = tree.match_prefix(MatchPrefixParams(key=RadixKey(req6_token_ids)))
result = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", req6_token_ids)))
)
kv_indices, last_node = result.device_indices, result.last_device_node
print(
f"req6: token_ids: {req6_token_ids}, matched kv_indices: {kv_indices}, last_node.key: {last_node.key}"
@@ -468,7 +473,7 @@ class TestSWA(unittest.TestCase):
print(
f"req1: inserting, req1_token_ids: {req1_token_ids}, req1_kv_indices: {req1_kv_indices}"
)
key = RadixKey(req1_token_ids)
key = RadixKey(array("q", req1_token_ids))
result = tree.insert(InsertParams(key=key, value=req1_kv_indices[: len(key)]))
prefix_len = result.prefix_len
self.assertEqual(prefix_len, 0)
@@ -480,7 +485,7 @@ class TestSWA(unittest.TestCase):
print(
f"req2: inserting, req2_token_ids: {req2_token_ids}, req2_kv_indices: {req2_kv_indices}"
)
key = RadixKey(req2_token_ids)
key = RadixKey(array("q", req2_token_ids))
result = tree.insert(InsertParams(key=key, value=req2_kv_indices[: len(key)]))
prefix_len = result.prefix_len
self.assertEqual(prefix_len, 2)
@@ -492,7 +497,7 @@ class TestSWA(unittest.TestCase):
print(
f"req3: inserting, req3_token_ids: {req3_token_ids}, req3_kv_indices: {req3_kv_indices}"
)
key = RadixKey(req3_token_ids)
key = RadixKey(array("q", req3_token_ids))
result = tree.insert(InsertParams(key=key, value=req3_kv_indices[: len(key)]))
prefix_len = result.prefix_len
self.assertEqual(prefix_len, 0)
@@ -504,7 +509,7 @@ class TestSWA(unittest.TestCase):
print(
f"req4: inserting, req4_token_ids: {req4_token_ids}, req4_kv_indices: {req4_kv_indices}"
)
key = RadixKey(req4_token_ids)
key = RadixKey(array("q", req4_token_ids))
result = tree.insert(InsertParams(key=key, value=req4_kv_indices[: len(key)]))
prefix_len = result.prefix_len
self.assertEqual(prefix_len, 4)
@@ -553,7 +558,9 @@ class TestSWA(unittest.TestCase):
tree.pretty_print()
req5_token_ids = [1, 2, 3, 4, 5]
result = tree.match_prefix(MatchPrefixParams(key=RadixKey(req5_token_ids)))
result = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", req5_token_ids)))
)
kv_indices, last_node = result.device_indices, result.last_device_node
print(
f"req5: token_ids: {req5_token_ids}, matched kv_indices: {kv_indices}, last_node.key: {last_node.key}"
@@ -561,7 +568,9 @@ class TestSWA(unittest.TestCase):
self.assertEqual(len(kv_indices), 0) # no swa prefix matched
req6_token_ids = [1, 2, 3, 4, 5, 60, 70]
result = tree.match_prefix(MatchPrefixParams(key=RadixKey(req6_token_ids)))
result = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", req6_token_ids)))
)
kv_indices, last_node = result.device_indices, result.last_device_node
print(
f"req6: token_ids: {req6_token_ids}, matched kv_indices: {kv_indices}, last_node.key: {last_node.key}"
@@ -578,8 +587,8 @@ class TestSWA(unittest.TestCase):
# Case 1: is_insert=True should pass bigram key and use cache_protected_len.
req = _DummyReq()
req.req_pool_idx = 0
req.origin_input_ids = [1, 2, 3, 4, 5, 6]
req.output_ids = []
req.origin_input_ids = array("q", [1, 2, 3, 4, 5, 6])
req.output_ids = array("q")
req._kv_committed_len = len(req.origin_input_ids)
kv_indices = allocator.alloc(req._kv_committed_len)
req_to_token_pool.write(
@@ -613,8 +622,8 @@ class TestSWA(unittest.TestCase):
# even when len(prefix_indices) is intentionally larger.
req2 = _DummyReq()
req2.req_pool_idx = 1
req2.origin_input_ids = [11, 12, 13, 14, 15, 16]
req2.output_ids = []
req2.origin_input_ids = array("q", [11, 12, 13, 14, 15, 16])
req2.output_ids = array("q")
req2._kv_committed_len = len(req2.origin_input_ids)
kv_indices2 = allocator.alloc(req2._kv_committed_len)
req_to_token_pool.write(
@@ -730,7 +739,9 @@ class TestSWASplitLeafOnInsert(CustomTestCase):
with envs.SGLANG_OPT_SWA_SPLIT_LEAF_ON_INSERT.override(True):
inserted_leaf = _insert_chain(tree, allocator, token_ids)
self.assertEqual(len(inserted_leaf.value), 4)
match = tree.match_prefix(MatchPrefixParams(key=RadixKey(token_ids)))
match = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", token_ids)))
)
self.assertEqual(match.device_indices.shape[0], 12)
self.assertIs(match.last_device_node, inserted_leaf)
@@ -12,6 +12,7 @@ import random
import statistics
import time
import unittest
from array import array
from contextlib import contextmanager
from dataclasses import dataclass
from typing import Callable
@@ -335,7 +336,7 @@ def _insert_seq(env, seq):
if env.has_mamba:
req = env.make_req()
mamba_val = req.mamba_pool_idx.unsqueeze(0)
key = RadixKey(seq)
key = RadixKey(array("q", seq))
env.tree.insert(InsertParams(key=key, value=v[: len(key)], mamba_value=mamba_val))
return True
@@ -357,7 +358,7 @@ def _fill_no_evict(env):
if env.has_mamba:
req = env.make_req()
mamba_val = req.mamba_pool_idx.unsqueeze(0)
key = RadixKey(seq)
key = RadixKey(array("q", seq))
env.tree.insert(
InsertParams(key=key, value=v[: len(key)], mamba_value=mamba_val)
)
@@ -505,7 +506,7 @@ def bench_match_prefix(
queries.append([rng.randint(1, 32000)] * rng.randint(50, 300))
def verify_fn(q):
k = RadixKey(q)
k = RadixKey(array("q", q))
r1 = env.tree.match_prefix(MatchPrefixParams(key=k))
r2 = env.tree.match_prefix(MatchPrefixParams(key=k))
assert len(r1.device_indices) == len(r2.device_indices), "match not idempotent"
@@ -514,7 +515,7 @@ def bench_match_prefix(
return bench_api(
"match_prefix",
lambda: queries,
lambda q: env.tree.match_prefix(MatchPrefixParams(key=RadixKey(q))),
lambda q: env.tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", q)))),
min(len(queries) - warmup, num_seqs),
env.avg_tokens,
warmup,
@@ -566,7 +567,7 @@ def bench_lock_unlock(
nodes = []
for seq in env.seqs[: num_seqs // 2]:
r = env.tree.match_prefix(MatchPrefixParams(key=RadixKey(seq)))
r = env.tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
if r.last_device_node != env.tree.root_node:
nodes.append(r.last_device_node)
if not nodes:
@@ -613,7 +614,7 @@ def bench_cache_finished(
# Pre-build Req objects with token IDs filled into req_to_token
req_items: list = []
for seq in env.seqs:
key = RadixKey(seq)
key = RadixKey(array("q", seq))
mr = env.tree.match_prefix(MatchPrefixParams(key=key))
matched_len = len(mr.device_indices)
node = mr.last_device_node
@@ -635,9 +636,9 @@ def bench_cache_finished(
kv_indices = mr.device_indices
req = env.make_req()
req.origin_input_ids = list(seq)
req.output_ids = []
req.fill_ids = list(seq)
req.origin_input_ids = array("q", seq)
req.output_ids = array("q")
req.fill_ids = array("q", seq)
req.last_node = node
req.cache_protected_len = matched_len
req.kv_committed_len = len(seq)
@@ -1,6 +1,7 @@
"""Unit tests for UnifiedRadixCache"""
import unittest
from array import array
from dataclasses import dataclass
from typing import Optional
from unittest import mock
@@ -272,7 +273,7 @@ class UnifiedRadixCacheSuite:
def _insert(self, tree, allocator, req_to_token_pool, tokens):
"""Insert tokens, attaching mamba data when the config has mamba."""
key = RadixKey(tokens)
key = RadixKey(array("q", tokens))
value = self._alloc(allocator, len(tokens))
params = InsertParams(key=key, value=value[: len(key)])
if self.cfg.has_mamba:
@@ -290,15 +291,17 @@ class UnifiedRadixCacheSuite:
result = self._insert(tree, allocator, req_to_token_pool, seq_b)
self.assertEqual(result.prefix_len, len(seq_a))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq_b)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq_b))))
self.assertEqual(len(m.device_indices), len(seq_b))
m = tree.match_prefix(
MatchPrefixParams(key=RadixKey(seq_a + self._make_seq(9000, 1)))
MatchPrefixParams(key=RadixKey(array("q", seq_a + self._make_seq(9000, 1))))
)
self.assertEqual(len(m.device_indices), len(seq_a))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(self._make_seq(5000, 2))))
m = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", self._make_seq(5000, 2))))
)
self.assertEqual(len(m.device_indices), 0)
tree.sanity_check()
@@ -317,11 +320,11 @@ class UnifiedRadixCacheSuite:
self.assertEqual(result_b.prefix_len, len(base))
for seq in (branch_a, branch_b):
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
self.assertEqual(len(m.device_indices), len(seq))
m = tree.match_prefix(
MatchPrefixParams(key=RadixKey(base + self._make_seq(999, 1)))
MatchPrefixParams(key=RadixKey(array("q", base + self._make_seq(999, 1))))
)
self.assertEqual(len(m.device_indices), len(base))
tree.sanity_check()
@@ -350,13 +353,13 @@ class UnifiedRadixCacheSuite:
self._insert(tree, allocator, req_to_token_pool, seq_a)
self._insert(tree, allocator, req_to_token_pool, seq_b)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq_a)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq_a))))
lock_result = tree.inc_lock_ref(m.last_device_node)
result = tree.evict(EvictParams(num_tokens=len(seq_a) + len(seq_b)))
self.assertGreaterEqual(result.num_tokens_evicted, len(seq_b))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq_a)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq_a))))
self.assertEqual(len(m.device_indices), len(seq_a))
# Unlock -> should now be evictable
@@ -395,7 +398,7 @@ class UnifiedRadixCacheSuite:
if self.cfg.has_mamba:
self.assertEqual(tree.mamba_evictable_size(), 0)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seqs[0])))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seqs[0]))))
self.assertEqual(len(m.device_indices), 0)
tree.sanity_check()
@@ -413,7 +416,7 @@ class UnifiedRadixCacheSuite:
self.assertEqual(allocator.available_size(), initial_avail - len(seq_1p))
# Step 2: insert 2 pages with prev_prefix_len=0 → frees overlap of 1 page
key_2p = RadixKey(seq_2p)
key_2p = RadixKey(array("q", seq_2p))
value_2p = self._alloc(allocator, len(seq_2p))
params = InsertParams(
key=key_2p,
@@ -432,7 +435,7 @@ class UnifiedRadixCacheSuite:
# Step 3: insert 3 pages with prev_prefix_len=len(seq_2p) → nothing freed
avail_before = allocator.available_size()
key_3p = RadixKey(seq_3p)
key_3p = RadixKey(array("q", seq_3p))
value_3p = self._alloc(allocator, len(seq_3p))
params = InsertParams(
key=key_3p,
@@ -461,11 +464,11 @@ class UnifiedRadixCacheSuite:
self.assertEqual(result.prefix_len, len(base))
for seq in (fork_a, fork_b):
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
self.assertEqual(len(m.device_indices), len(seq))
m = tree.match_prefix(
MatchPrefixParams(key=RadixKey(base + self._make_seq(999, 1)))
MatchPrefixParams(key=RadixKey(array("q", base + self._make_seq(999, 1))))
)
self.assertEqual(len(m.device_indices), len(base))
tree.sanity_check()
@@ -477,8 +480,8 @@ class UnifiedRadixCacheSuite:
req = self._make_req(req_to_token_pool)
input_ids = self._make_seq(1, 3)
output_ids = self._make_seq(2000, 1)
req.origin_input_ids = input_ids
req.output_ids = output_ids
req.origin_input_ids = array("q", input_ids)
req.output_ids = array("q", output_ids)
kv_len = len(input_ids) + len(output_ids)
kv_indices = self._alloc(allocator, kv_len)
req_to_token_pool.write((req.req_pool_idx, slice(0, kv_len)), kv_indices)
@@ -487,7 +490,7 @@ class UnifiedRadixCacheSuite:
req.cache_protected_len = 0
req.swa_uuid_for_lock = None
req.extra_key = None
req.fill_ids = input_ids + output_ids
req.fill_ids = array("q", input_ids + output_ids)
if self.cfg.has_mamba:
req.mamba_last_track_seqlen = kv_len
@@ -495,7 +498,9 @@ class UnifiedRadixCacheSuite:
all_ids = input_ids + output_ids
aligned_len = (len(all_ids) // ps) * ps
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(all_ids[:aligned_len])))
m = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", all_ids[:aligned_len])))
)
self.assertEqual(len(m.device_indices), aligned_len)
tree.sanity_check()
@@ -506,9 +511,9 @@ class UnifiedRadixCacheSuite:
req = self._make_req(req_to_token_pool)
prompt_ids = self._make_seq(1, 3)
output_ids = self._make_seq(2000, 7)
req.origin_input_ids = prompt_ids
req.output_ids = output_ids
req.fill_ids = prompt_ids + output_ids
req.origin_input_ids = array("q", prompt_ids)
req.output_ids = array("q", output_ids)
req.fill_ids = array("q", prompt_ids + output_ids)
kv_len = len(req.fill_ids)
kv_indices = self._alloc(allocator, kv_len)
req_to_token_pool.write((req.req_pool_idx, slice(0, kv_len)), kv_indices)
@@ -538,7 +543,9 @@ class UnifiedRadixCacheSuite:
prompt_aligned = (len(prompt_ids) // ps) * ps
# Thinking+answer must not be reachable past the prompt.
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(prompt_ids + output_ids)))
m = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", prompt_ids + output_ids)))
)
self.assertEqual(len(m.device_indices), prompt_aligned)
# Only prompt-aligned pages remain owned by the tree.
self.assertEqual(
@@ -550,8 +557,8 @@ class UnifiedRadixCacheSuite:
tree, allocator, req_to_token_pool = build_fixture(self.cfg)
req = self._make_req(req_to_token_pool)
tokens = self._make_seq(1, 2)
req.origin_input_ids = tokens
req.output_ids = []
req.origin_input_ids = array("q", tokens)
req.output_ids = array("q")
kv_len = len(tokens)
kv_indices = self._alloc(allocator, kv_len)
req_to_token_pool.write((req.req_pool_idx, slice(0, kv_len)), kv_indices)
@@ -560,13 +567,13 @@ class UnifiedRadixCacheSuite:
req.cache_protected_len = 0
req.swa_uuid_for_lock = None
req.extra_key = None
req.fill_ids = tokens
req.fill_ids = array("q", tokens)
avail_before = allocator.available_size()
tree.cache_finished_req(req, is_insert=False)
self.assertEqual(allocator.available_size(), avail_before + kv_len)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(tokens)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", tokens))))
self.assertEqual(len(m.device_indices), 0)
tree.sanity_check()
@@ -575,9 +582,9 @@ class UnifiedRadixCacheSuite:
req = self._make_req(req_to_token_pool)
tokens = self._make_seq(1, 3)
req.origin_input_ids = tokens
req.output_ids = []
req.fill_ids = tokens[:]
req.origin_input_ids = array("q", tokens)
req.output_ids = array("q")
req.fill_ids = array("q", tokens)
kv_len = len(tokens)
kv_indices = self._alloc(allocator, kv_len)
req_to_token_pool.write((req.req_pool_idx, slice(0, kv_len)), kv_indices)
@@ -628,11 +635,11 @@ class UnifiedRadixCacheSuite:
for suffix_start in [100, 200, 300]:
seq = base + self._make_seq(suffix_start, 2)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
self.assertEqual(len(m.device_indices), len(seq))
m = tree.match_prefix(
MatchPrefixParams(key=RadixKey(base + self._make_seq(999, 1)))
MatchPrefixParams(key=RadixKey(array("q", base + self._make_seq(999, 1))))
)
self.assertEqual(len(m.device_indices), len(base))
tree.sanity_check()
@@ -641,7 +648,7 @@ class UnifiedRadixCacheSuite:
if self.cfg.page_size == 1:
self.skipTest("page_size > 1 only")
tree, _, _ = build_fixture(self.cfg)
key = RadixKey(self._make_seq(1, 1))
key = RadixKey(array("q", self._make_seq(1, 1)))
child_key = key.child_key(tree.page_size)
self.assertIsInstance(child_key, tuple)
@@ -656,11 +663,13 @@ class UnifiedRadixCacheSuite:
# Tree truncates unaligned tail internally, so it matches the seq prefix.
unaligned = seq + list(range(9000, 9000 + ps - 1))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(unaligned)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", unaligned))))
self.assertEqual(len(m.device_indices), len(seq))
# Below-page-size key aligns to 0 -> no match.
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq[: ps - 1])))
m = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", seq[: ps - 1])))
)
self.assertEqual(len(m.device_indices), 0)
tree.sanity_check()
@@ -679,12 +688,12 @@ class UnifiedRadixCacheSuite:
# Mismatch in second page → only first page matches
bad_page2 = seq[:ps] + [9999] * ps
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(bad_page2)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", bad_page2))))
self.assertEqual(len(m.device_indices), ps)
# Mismatch in first page → 0 match
bad_page1 = [9999] + seq[1:]
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(bad_page1)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", bad_page1))))
self.assertEqual(len(m.device_indices), 0)
tree.sanity_check()
@@ -699,8 +708,8 @@ class UnifiedRadixCacheSuite:
tail_extra = ps // 2
input_ids = self._make_seq(1, 1) + list(range(8000, 8000 + tail_extra))
req = self._make_req(req_to_token_pool)
req.origin_input_ids = input_ids
req.output_ids = []
req.origin_input_ids = array("q", input_ids)
req.output_ids = array("q")
kv_len = len(input_ids)
kv_indices = self._alloc(allocator, kv_len)
req_to_token_pool.write((req.req_pool_idx, slice(0, kv_len)), kv_indices)
@@ -709,7 +718,7 @@ class UnifiedRadixCacheSuite:
req.cache_protected_len = 0
req.swa_uuid_for_lock = None
req.extra_key = None
req.fill_ids = input_ids
req.fill_ids = array("q", input_ids)
if self.cfg.has_mamba:
req.mamba_last_track_seqlen = kv_len
@@ -718,7 +727,7 @@ class UnifiedRadixCacheSuite:
self.assertEqual(allocator.available_size(), avail_before + tail_extra)
aligned = input_ids[: (len(input_ids) // ps) * ps]
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(aligned)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", aligned))))
self.assertEqual(len(m.device_indices), len(aligned))
tree.sanity_check()
@@ -749,7 +758,7 @@ class UnifiedRadixCacheSuite:
tree.evict(EvictParams(num_tokens=0, mamba_num=10))
self.assertEqual(tree.mamba_evictable_size(), 0)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq_long)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq_long))))
self.assertEqual(len(m.device_indices), 0)
tree.sanity_check()
@@ -788,7 +797,7 @@ class UnifiedRadixCacheSuite:
req2 = self._make_req(req_to_token_pool)
m = tree.match_prefix(
MatchPrefixParams(key=RadixKey(seq), cow_mamba=True, req=req2)
MatchPrefixParams(key=RadixKey(array("q", seq)), cow_mamba=True, req=req2)
)
self.assertEqual(len(m.device_indices), len(seq))
self.assertIsNotNone(req2.mamba_pool_idx)
@@ -809,7 +818,7 @@ class UnifiedRadixCacheSuite:
seq = self._make_seq(1, 3)
self._insert(tree, allocator, req_to_token_pool, seq)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
self.assertEqual(len(m.device_indices), len(seq))
tree.sanity_check()
@@ -866,13 +875,13 @@ class UnifiedRadixCacheSuite:
self._insert(tree, allocator, req_to_token_pool, seq_a)
self._insert(tree, allocator, req_to_token_pool, seq_b)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq_a)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq_a))))
lock_result = tree.inc_lock_ref(m.last_device_node)
result = tree.evict(EvictParams(num_tokens=len(seq_a) + len(seq_b)))
self.assertGreaterEqual(result.num_tokens_evicted, len(seq_b))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq_a)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq_a))))
self.assertEqual(len(m.device_indices), len(seq_a))
tree.dec_lock_ref(
@@ -886,8 +895,8 @@ class UnifiedRadixCacheSuite:
parent = UnifiedTreeNode(self.cfg.components)
deleted = UnifiedTreeNode(self.cfg.components)
parent.key = RadixKey(self._make_seq(1, 1))
deleted.key = RadixKey(self._make_seq(1000, 1))
parent.key = RadixKey(array("q", self._make_seq(1, 1)))
deleted.key = RadixKey(array("q", self._make_seq(1000, 1)))
parent.parent = tree.root_node
deleted.parent = parent
parent.component_data[ComponentType.FULL].value = torch.arange(
@@ -924,15 +933,15 @@ class UnifiedRadixCacheSuite:
node_count_before = count_nodes(tree.root_node)
self.assertEqual(node_count_before, 2)
tree._match_prefix_helper(RadixKey([1, 2]))
tree._match_prefix_helper(RadixKey(array("q", [1, 2])))
(
value,
best_match_node,
best_match_device_node,
best_value_len,
) = tree._match_prefix_helper(RadixKey([1, 2, 3, 4]))
) = tree._match_prefix_helper(RadixKey(array("q", [1, 2, 3, 4])))
self.assertEqual(best_value_len, 2)
self.assertEqual(best_match_node.key.token_ids, [3, 4])
self.assertEqual(list(best_match_node.key.token_ids), [3, 4])
self.assertIs(best_match_device_node, best_match_node)
node_count_after_regular = count_nodes(tree.root_node)
self.assertEqual(node_count_after_regular, node_count_before + 2)
@@ -942,9 +951,9 @@ class UnifiedRadixCacheSuite:
best_match_node,
best_match_device_node,
best_value_len,
) = tree._match_prefix_helper_readonly(RadixKey([1, 2, 3]))
) = tree._match_prefix_helper_readonly(RadixKey(array("q", [1, 2, 3])))
self.assertEqual(best_value_len, 1)
self.assertEqual(best_match_node.key.token_ids, [1, 2])
self.assertEqual(list(best_match_node.key.token_ids), [1, 2])
self.assertIs(best_match_device_node, best_match_node)
node_count_after_readonly = count_nodes(tree.root_node)
self.assertEqual(node_count_after_readonly, node_count_after_regular)
@@ -971,7 +980,7 @@ class UnifiedRadixCacheSuite:
seq = self._make_seq(1, 2)
self._insert(tree, allocator, req_to_token_pool, seq)
match = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq)))
match = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
node = match.last_device_node
full_cd = node.component_data[ComponentType.FULL]
aux_cd = node.component_data[aux]
@@ -1040,7 +1049,7 @@ class UnifiedRadixCacheSuite:
self._insert(tree, allocator, req_to_token_pool, leaf)
# Lock the base node to prevent it from being evicted
m_base = tree.match_prefix(MatchPrefixParams(key=RadixKey(base)))
m_base = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", base))))
lock_result = tree.inc_lock_ref(m_base.last_device_node)
# Evict the leaf — parent (base) should become D-leaf after unlock
@@ -1089,14 +1098,14 @@ class UnifiedRadixCacheSuite:
self._insert(tree, allocator, req_to_token_pool, seq_new)
# Touch seq_new to make it MRU
tree.match_prefix(MatchPrefixParams(key=RadixKey(seq_new)))
tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq_new))))
# Evict just enough for one sequence
tree.evict(EvictParams(num_tokens=len(seq_old)))
# seq_old should be gone (LRU), seq_new should remain
m_old = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq_old)))
m_new = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq_new)))
m_old = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq_old))))
m_new = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq_new))))
self.assertEqual(len(m_old.device_indices), 0)
self.assertEqual(len(m_new.device_indices), len(seq_new))
tree.sanity_check()
@@ -1136,13 +1145,13 @@ class UnifiedRadixCacheSuite:
self._insert(tree, allocator, req_to_token_pool, branch_b)
# Lock branch_b
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(branch_b)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", branch_b))))
lr = tree.inc_lock_ref(m.last_device_node)
# Evict — branch_a should go, base + branch_b stay
tree.evict(EvictParams(num_tokens=len(branch_a)))
m_b = tree.match_prefix(MatchPrefixParams(key=RadixKey(branch_b)))
m_b = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", branch_b))))
self.assertEqual(len(m_b.device_indices), len(branch_b))
tree.dec_lock_ref(
@@ -1173,7 +1182,7 @@ class UnifiedRadixCacheSuite:
self._insert(tree, allocator, req_to_token_pool, seq_b)
# Lock seq_a
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq_a)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq_a))))
lr = tree.inc_lock_ref(m.last_device_node)
# Try to evict everything
@@ -1181,7 +1190,7 @@ class UnifiedRadixCacheSuite:
result = tree.evict(EvictParams(num_tokens=total))
# seq_a should still be matchable (protected)
m2 = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq_a)))
m2 = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq_a))))
self.assertEqual(len(m2.device_indices), len(seq_a))
tree.dec_lock_ref(
@@ -1220,7 +1229,7 @@ class UnifiedRadixCacheSuite:
# Re-insert
seq_b = self._make_seq(500, 2)
self._insert(tree, allocator, req_to_token_pool, seq_b)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq_b)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq_b))))
self.assertEqual(len(m.device_indices), len(seq_b))
tree.sanity_check()
@@ -1247,7 +1256,7 @@ class UnifiedRadixCacheSuite:
self._insert(tree, allocator, req_to_token_pool, s)
# Lock some, evict some, unlock
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seqs[0])))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seqs[0]))))
lr = tree.inc_lock_ref(m.last_device_node)
tree.evict(EvictParams(num_tokens=len(seqs[1])))
@@ -1409,7 +1418,7 @@ class UnifiedRadixCacheSuite:
self._insert(tree, allocator, req_to_token_pool, seq)
# Find the leaf node
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
node = m.last_device_node
self.assertIsNot(node, tree.root_node)
@@ -1435,7 +1444,7 @@ class UnifiedRadixCacheSuite:
seq = self._make_seq(1, 2)
self._insert(tree, allocator, req_to_token_pool, seq)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
node = m.last_device_node
self._backup_node(tree, node)
@@ -1469,7 +1478,7 @@ class UnifiedRadixCacheSuite:
self._backup_tree(tree)
# Lock leaf so only base can be evicted
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(leaf)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", leaf))))
lr = tree.inc_lock_ref(m.last_device_node)
# Evict base (inner node won't be evicted while child is locked)
@@ -1479,7 +1488,7 @@ class UnifiedRadixCacheSuite:
m.last_device_node,
DecLockRefParams(swa_uuid_for_lock=getattr(lr, "swa_uuid_for_lock", None)),
)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(leaf)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", leaf))))
self.assertGreaterEqual(len(m.device_indices), len(base))
tree.sanity_check()
@@ -1495,7 +1504,7 @@ class UnifiedRadixCacheSuite:
query = expected_prefix + self._make_seq(9000, 1)
self._insert(tree, allocator, req_to_token_pool, seq)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
node = m.last_device_node
self._backup_node(tree, node)
@@ -1503,7 +1512,7 @@ class UnifiedRadixCacheSuite:
self.assertTrue(node.evicted)
self.assertTrue(node.backuped)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(query)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", query))))
self.assertEqual(len(m.device_indices), 0)
self.assertIs(m.last_device_node, tree.root_node)
@@ -1512,8 +1521,8 @@ class UnifiedRadixCacheSuite:
self.assertIsNot(split_parent, tree.root_node)
self.assertTrue(split_parent.evicted)
self.assertTrue(split_parent.backuped)
self.assertEqual(split_parent.key.token_ids, expected_prefix)
self.assertEqual(node.key.token_ids, expected_suffix)
self.assertEqual(list(split_parent.key.token_ids), expected_prefix)
self.assertEqual(list(node.key.token_ids), expected_suffix)
if self.cfg.has_mamba:
self.assertEqual(m.host_hit_length, 0)
@@ -1536,7 +1545,7 @@ class UnifiedRadixCacheSuite:
self._insert(tree, allocator, req_to_token_pool, s)
for i in range(2):
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seqs[i])))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seqs[i]))))
self._backup_node(tree, m.last_device_node)
# Evict one backed-up node
@@ -1555,7 +1564,7 @@ class UnifiedRadixCacheSuite:
seq = self._make_seq(1, 2)
self._insert(tree, allocator, req_to_token_pool, seq)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
node = m.last_device_node
self._backup_node(tree, node)
@@ -1580,7 +1589,7 @@ class UnifiedRadixCacheSuite:
base = self._make_seq(1, 2)
self._insert(tree, allocator, req_to_token_pool, base)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(base)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", base))))
node = m.last_device_node
original_device_indices = m.device_indices.clone()
self._fill_full_kv(allocator, original_device_indices, marker=3)
@@ -1662,7 +1671,7 @@ class UnifiedRadixCacheSuite:
seq = self._make_seq(1, 2)
self._insert(tree, allocator, req_to_token_pool, seq)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
node = m.last_device_node
for aux in aux_types:
@@ -1688,7 +1697,7 @@ class UnifiedRadixCacheSuite:
for i in range(num_pages):
seq = seq + self._make_seq(1000 * (i + 1), 1)
self._insert(tree, allocator, req_to_token_pool, seq)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
chain: list = []
cur = m.last_device_node
while cur is not tree.root_node:
@@ -1738,7 +1747,7 @@ class UnifiedRadixCacheSuite:
seq = self._make_seq(1, (min_tokens + ps - 1) // ps)
self._insert(tree, allocator, req_to_token_pool, seq)
result = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq)))
result = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
self.assertEqual(len(result.device_indices), len(seq))
self.assertIs(result.best_match_node, result.last_device_node)
@@ -1760,7 +1769,7 @@ class UnifiedRadixCacheSuite:
tree.evict(EvictParams(num_tokens=len(leaf.key)))
self.assertTrue(leaf.evicted)
result = tree.match_prefix(MatchPrefixParams(key=RadixKey(tokens)))
result = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", tokens))))
self.assertIs(result.best_match_node, leaf)
self.assertIs(result.last_device_node, parent)
@@ -1782,7 +1791,7 @@ class UnifiedRadixCacheSuite:
tree.evict(EvictParams(num_tokens=len(leaf.key)))
self.assertTrue(leaf.evicted)
result = tree.match_prefix(MatchPrefixParams(key=RadixKey(tokens)))
result = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", tokens))))
self.assertIs(result.best_match_node, leaf)
self.assertIs(result.last_device_node, parent)
@@ -1797,12 +1806,14 @@ class UnifiedRadixCacheSuite:
tokens = self._make_seq(1, chunk_size + 1)
self._insert(tree, allocator, req_to_token_pool, tokens)
leaf = tree.match_prefix(
MatchPrefixParams(key=RadixKey(tokens))
MatchPrefixParams(key=RadixKey(array("q", tokens)))
).last_device_node
mamba_cd = leaf.component_data[ComponentType.MAMBA]
mamba_cd.value = None
no_hicache = tree.match_prefix(MatchPrefixParams(key=RadixKey(tokens)))
no_hicache = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", tokens)))
)
self.assertIs(no_hicache.best_match_node, tree.root_node)
self.assertIs(no_hicache.last_device_node, tree.root_node)
self.assertEqual(no_hicache.mamba_branching_seqlen, chunk_size)
@@ -1810,11 +1821,13 @@ class UnifiedRadixCacheSuite:
tree_h, allocator_h, req_to_token_pool_h = self._build_hicache_fixture()
self._insert(tree_h, allocator_h, req_to_token_pool_h, tokens)
leaf_h = tree_h.match_prefix(
MatchPrefixParams(key=RadixKey(tokens))
MatchPrefixParams(key=RadixKey(array("q", tokens)))
).last_device_node
self._backup_node(tree_h, leaf_h)
tree_h.evict(EvictParams(num_tokens=len(tokens)))
with_hicache = tree_h.match_prefix(MatchPrefixParams(key=RadixKey(tokens)))
with_hicache = tree_h.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", tokens)))
)
self.assertIs(with_hicache.best_match_node, leaf_h)
self.assertIs(with_hicache.last_device_node, tree_h.root_node)
self.assertIsNone(with_hicache.mamba_branching_seqlen)
@@ -1834,7 +1847,9 @@ class UnifiedRadixCacheSuite:
self.assertTrue(leaf.evicted)
req = self._make_req(req_to_token_pool)
match = tree.match_prefix(MatchPrefixParams(key=RadixKey(tokens), req=req))
match = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", tokens)), req=req)
)
req.prefix_indices = match.device_indices
req.last_node = match.last_device_node
req.best_match_node = match.best_match_node
@@ -1876,7 +1891,9 @@ class UnifiedRadixCacheSuite:
self._set_aux_host_tombstone(tree, leaf, aux)
req = self._make_req(req_to_token_pool)
match = tree.match_prefix(MatchPrefixParams(key=RadixKey(tokens), req=req))
match = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", tokens)), req=req)
)
req.prefix_indices = match.device_indices
req.last_node = match.last_device_node
req.best_match_node = match.best_match_node
@@ -1915,7 +1932,9 @@ class UnifiedRadixCacheSuite:
tree.evict(EvictParams(num_tokens=len(leaf.key)))
req = self._make_req(req_to_token_pool)
match = tree.match_prefix(MatchPrefixParams(key=RadixKey(tokens), req=req))
match = tree.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", tokens)), req=req)
)
req.prefix_indices = match.device_indices
req.last_node = match.last_device_node
req.best_match_node = match.best_match_node
@@ -1986,7 +2005,7 @@ class UnifiedRadixCacheSuite:
tree, allocator, req_to_token_pool = build_fixture(self.cfg)
seq = self._make_seq(1, 2)
self._insert(tree, allocator, req_to_token_pool, seq)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
node = m.last_device_node
self._simulate_backup(tree, node)
@@ -2077,7 +2096,9 @@ class UnifiedRadixCacheSuite:
)
result = swa_comp.finalize_match_result(
result=result,
params=MatchPrefixParams(key=RadixKey(self._make_seq(1, 1))),
params=MatchPrefixParams(
key=RadixKey(array("q", self._make_seq(1, 1)))
),
value_chunks=[],
best_value_len=0,
)
@@ -2214,7 +2235,7 @@ class UnifiedRadixCacheSuite:
def test_hicache_swa_match_prefix_picks_best_match_node_above_last_host(self):
tree, _, n, y, x, tokens = self._swa_anchor_setup()
result = tree.match_prefix(MatchPrefixParams(key=RadixKey(tokens)))
result = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", tokens))))
self.assertIs(result.best_match_node, x)
self.assertIs(result.last_device_node, n.parent)
self.assertIs(result.last_host_node, y)
@@ -2247,7 +2268,7 @@ class UnifiedRadixCacheSuite:
)
result = swa_comp.finalize_match_result(
result=base,
params=MatchPrefixParams(key=RadixKey(self._make_seq(1, 1))),
params=MatchPrefixParams(key=RadixKey(array("q", self._make_seq(1, 1)))),
value_chunks=[],
best_value_len=0,
)
@@ -2365,7 +2386,7 @@ class UnifiedRadixCacheSuite:
tree, allocator, req_to_token_pool = build_fixture(self.cfg)
seq = self._make_seq(1, 2)
self._insert(tree, allocator, req_to_token_pool, seq)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
node = m.last_device_node
cd = node.component_data[ComponentType.MAMBA]
old_mamba = cd.value
@@ -2414,7 +2435,7 @@ class UnifiedRadixCacheSuite:
tree.sanity_check()
for i in range(3):
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seqs[i])))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seqs[i]))))
self._backup_node(tree, m.last_device_node)
# Evict to free some tokens
@@ -2443,7 +2464,7 @@ class UnifiedRadixCacheSuite:
self._insert(tree, allocator, req_to_token_pool, base)
self._insert(tree, allocator, req_to_token_pool, leaf_seq)
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(leaf_seq)))
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", leaf_seq))))
leaf = m.last_device_node
parent = leaf.parent
self.assertIsNot(parent, tree.root_node)
+50
View File
@@ -0,0 +1,50 @@
import unittest
from array import array
import torch
from sglang.srt.utils.common import flatten_arrays_to_int64_tensor
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=5, stage="base-b", runner_config="1-gpu-small")
@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA")
class TestFlattenArraysToInt64Tensor(CustomTestCase):
"""`flatten_arrays_to_int64_tensor` is invoked by `prepare_for_extend`
to build the per-batch input_ids tensor (pinned, async H2D) from a
list of array.array('q') per-req fill_ids slices. Tests the full
matrix of (device, pin) the production code paths through.
"""
DEVICES = ("cpu", "cuda")
PIN_OPTIONS = (False, True)
def _check(self, parts: list, expected: list[int]) -> None:
for device in self.DEVICES:
for pin in self.PIN_OPTIONS:
with self.subTest(device=device, pin=pin):
out = flatten_arrays_to_int64_tensor(parts, device, pin)
if device == "cuda":
torch.cuda.synchronize()
self.assertEqual(out.dtype, torch.int64)
self.assertEqual(out.device.type, device)
self.assertEqual(out.shape, (len(expected),))
self.assertEqual(out.cpu().tolist(), expected)
def test_single_part(self):
parts = [array("q", [1, 2, 3, 4, 5])]
self._check(parts, [1, 2, 3, 4, 5])
def test_multiple_parts(self):
parts = [
array("q", [10, 20, 30]),
array("q", [100, 200]),
array("q", [1000]),
]
self._check(parts, [10, 20, 30, 100, 200, 1000])
if __name__ == "__main__":
unittest.main()