diff --git a/benchmark/hicache/README-proactive.md b/benchmark/hicache/README-proactive.md new file mode 100644 index 000000000000..e6c80967bc37 --- /dev/null +++ b/benchmark/hicache/README-proactive.md @@ -0,0 +1,61 @@ +# Proactive resident HiCache restore (experimental) + +Restore an **exact token prefix** before an ordinary continuation reaches `/generate`. +Requires a fixed text-generation model, TP1/PP1/DP1, FULL-only resident (`cache`) +HiCache, and the file backend. LoRA, speculation, disaggregation, sidecars, SWA, +embedding models, and multimodal models are rejected. This control does not send a +model request or move KV to the GPU. + +```sh +curl -s http://localhost:30000/hicache/prefetch -H 'Content-Type: application/json' \ + -d '{"operation_id":"tool-1","input_ids":[1,2,3,4],"ttl_ms":10000}' +curl -s http://localhost:30000/hicache/prefetch -H 'Content-Type: application/json' \ + -d '{"operation_id":"tool-1","action":"status"}' +curl -s http://localhost:30000/hicache/prefetch -H 'Content-Type: application/json' \ + -d '{"operation_id":"tool-1","action":"cancel"}' +``` + +Use a prefix at least as long as the storage prefetch threshold (the four-token +example only illustrates the request shape). Prefixes are aligned down to complete +cache pages. For an exact continuation prompt, normally omit the final token and +align the remainder, matching ordinary prompt-cache lookup. Include the same +`cache_salt` in both requests when used; there is no implicit session lookup. +Existing admin API authentication applies. + +Only one restore is active; 32 recent outcomes are retained. Repeating an ID with +the same normalized prefix/salt/TTL is idempotent; a different payload is rejected. +Accepted operations begin `RUNNING` and finish `SUCCESS`, `MISS`, `FAILURE`, +`CANCELLED`, or `EXPIRED`. `CACHED` is a no-I/O result, `DECLINED` means existing +controller admission did not start the restore. A matching early continuation +joins the in-flight restore instead of submitting duplicate storage reads. It +then uses ordinary prefix matching and H2D. Other requests remain ordinary. + +TTL cancels pending work; it does **not** create a resident cache lease. Published +pages use normal eviction. Cancelling a completed restore does not invalidate +shared KV. Finite failures use the separate lifecycle fix; a backend call that +never returns cannot be reclaimed safely. Cancelled allocated work retains its +existing ownership until terminal ACK. Engine pause cancels this control and +continues draining its ACKs. + +## Benchmark + +```sh +python benchmark/hicache/bench_proactive_prefetch.py \ + --model-path /path/to/model --work-dir /tmp/hicache-bench \ + --results-dir /tmp/hicache-results --repetitions 3 +``` + +A: recompute with empty L3; B: request-time file L3 restore; C: control restore +before arrival. Each trial has a fresh process/cache, identical model/server/ +sampling configuration, and no injected I/O delay. Five tool gaps: 0/100/500/1000/ +3000 ms. Client TTFT comes from the first SSE token event. Existing plugin hooks +record actual storage reads, publication, host occupancy/evictions, and H2D; +all three modes have identical instrumentation. Source file KV is copied into +B/C trial directories; this measures a local file backend with potentially warm +OS page cache. It is not a remote-storage latency claim. + +Zero-gap includes control RPC latency; `actual_signal_to_arrival_ms` records that +overhead. Compare observed TTFT savings with measured hidden restore work rather +than assuming an exact nominal zero. Results include raw per-trial timings, +identical output-token verification, duplicate reads, and a no-continuation +experiment followed by ordinary cache flush to verify reclaimability. diff --git a/benchmark/hicache/bench_proactive_prefetch.py b/benchmark/hicache/bench_proactive_prefetch.py new file mode 100644 index 000000000000..756c895c9d47 --- /dev/null +++ b/benchmark/hicache/bench_proactive_prefetch.py @@ -0,0 +1,448 @@ +"""Real-model A/B/C tool-gap benchmark, fresh cache per trial, no delay injection. + +Example: python bench_proactive_prefetch.py --model-path /opt/model --work-dir /work/bench --results-dir /work/results +Requires one FULL-attention model/worker, CUDA, and the file backend. +""" + +import argparse +import collections +import json +import os +import random +import shutil +import signal +import statistics +import subprocess +import sys +import time +import urllib.error +import urllib.request +from pathlib import Path + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--model-path", required=True) + parser.add_argument("--work-dir", type=Path, required=True) + parser.add_argument("--results-dir", type=Path, required=True) + parser.add_argument("--port", type=int, default=30000) + parser.add_argument("--repetitions", type=int, default=3) + parser.add_argument( + "--gaps-ms", type=int, nargs="+", default=[0, 100, 500, 1000, 3000] + ) + args = parser.parse_args() + work, out = args.work_dir, args.results_dir + work.mkdir(parents=True, exist_ok=True) + out.mkdir(parents=True, exist_ok=True) + source = work / "source-l3" + source.mkdir() + trace = out / "storage-trace.jsonl" + url = f"http://127.0.0.1:{args.port}" + flags = [ + "--model-path", + args.model_path, + "--host", + "127.0.0.1", + "--port", + str(args.port), + "--tp-size", + "1", + "--pp-size", + "1", + "--dtype", + "bfloat16", + "--attention-backend", + "triton", + "--sampling-backend", + "pytorch", + "--enable-deterministic-inference", + "--random-seed", + "42", + "--max-running-requests", + "1", + "--context-length", + "8192", + "--max-total-tokens", + "8192", + "--page-size", + "16", + "--enable-hierarchical-cache", + "--hicache-host-memory-mode", + "cache", + "--hicache-size", + "4", + "--hicache-io-backend", + "kernel", + "--hicache-mem-layout", + "page_first", + "--hicache-write-policy", + "write_through", + "--hicache-storage-backend", + "file", + "--hicache-storage-prefetch-policy", + "wait_complete", + "--hicache-storage-backend-extra-config", + '{"prefetch_threshold":64}', + "--enable-metrics", + "--stream-interval", + "1", + "--cuda-graph-backend-decode", + "disabled", + "--cuda-graph-backend-prefill", + "disabled", + ] + (out / "server-command.json").write_text( + json.dumps([sys.executable, "-m", "sglang.launch_server"] + flags, indent=2) + ) + from transformers import AutoTokenizer + + tokenizer = AutoTokenizer.from_pretrained(args.model_path, local_files_only=True) + ids = tokenizer.encode( + "The validation notebook records observations about the lifecycle of cached tokens. " + * 400, + add_special_tokens=False, + )[:4096] + assert len(ids) == 4096 + prefix = ids[:4080] + payload = dict( + input_ids=ids, + sampling_params=dict( + temperature=0, max_new_tokens=32, ignore_eos=True, sampling_seed=42 + ), + return_logprob=True, + stream=True, + ) + (out / "request.json").write_text(json.dumps(payload)) + + def request(path, data=None, timeout=180): + req = urllib.request.Request( + url + path, + data=json.dumps(data).encode() if data is not None else None, + headers={"Content-Type": "application/json"}, + ) + with urllib.request.urlopen(req, timeout=timeout) as r: + return r.read().decode() + + def control(action, operation_id, **kwargs): + result = json.loads( + request( + "/hicache/prefetch", + dict(action=action, operation_id=operation_id, **kwargs), + ) + ) + assert result["success"], result + return result["result"] + + def generate(): + start = time.monotonic_ns() + first = None + last = None + req = urllib.request.Request( + url + "/generate", + data=json.dumps(payload).encode(), + headers={"Content-Type": "application/json"}, + ) + with urllib.request.urlopen(req, timeout=180) as response: + for line in response: + if not line.startswith(b"data: "): + continue + data = line[6:].strip() + if data == b"[DONE]": + break + last = json.loads(data) + assert "error" not in last, last + if last.get("output_ids") and first is None: + first = time.monotonic_ns() + assert first is not None and len(last["output_ids"]) == 32, last + return last, dict( + arrival_ns=start, ttft_ms=(first - start) / 1e6, end_ns=time.monotonic_ns() + ) + + def stop(proc): + if proc.poll() is None: + os.killpg(proc.pid, signal.SIGTERM) + try: + proc.wait(timeout=15) + except subprocess.TimeoutExpired: + os.killpg(proc.pid, signal.SIGKILL) + proc.wait(timeout=10) + + def start(label, directory): + env = os.environ.copy() + env.update( + SGLANG_HICACHE_FILE_BACKEND_STORAGE_DIR=str(directory), + HICACHE_BENCH_TRACE=str(trace), + HICACHE_BENCH_LABEL=label, + SGLANG_PLUGINS="hicache_trace", + ) + env["PYTHONPATH"] = ( + str(Path(__file__).parent / "trace") + ":" + env.get("PYTHONPATH", "") + ) + with (out / f"server-{label}.log").open("w") as log: + proc = subprocess.Popen( + [sys.executable, "-m", "sglang.launch_server"] + flags, + env=env, + stdout=log, + stderr=subprocess.STDOUT, + start_new_session=True, + ) + try: + deadline = time.monotonic() + 180 + while time.monotonic() < deadline: + if proc.poll() is not None: + raise RuntimeError(f"Server {label} exited {proc.returncode}") + try: + request("/health", timeout=2) + return proc + except (urllib.error.URLError, TimeoutError): + time.sleep(0.5) + raise TimeoutError(f"Server {label} did not become ready in 3 minutes") + except BaseException: + stop(proc) + raise + + def events(label): + result = [] + if trace.exists(): + for line in trace.read_text().splitlines(): + try: + item = json.loads(line) + except json.JSONDecodeError: + continue + if item["label"] == label: + result.append(item) + return result + + def upload(): + # Optional adapter in the validated container; never includes KV/model files. + script = os.environ.get("HICACHE_RESULTS_UPLOADER") + if script: + subprocess.run([sys.executable, script], check=True, timeout=120) + + producer = start("producer", source) + try: + baseline, _ = generate() + expected = baseline["output_ids"] + previous = None + stable = 0 + deadline = time.monotonic() + 60 + while time.monotonic() < deadline: + current = sorted((x.name, x.stat().st_size) for x in source.glob("*.bin")) + stable = stable + 1 if current and current == previous else 0 + previous = current + if stable >= 3: + break + time.sleep(1) + else: + raise TimeoutError("L3 writes did not settle") + finally: + stop(producer) + (out / "l3-source-manifest.json").write_text(json.dumps(previous)) + expected_file = out / "expected-output-ids.json" + if expected_file.exists(): + assert json.loads(expected_file.read_text()) == expected, ( + "Resume changed baseline output" + ) + expected_file.write_text(json.dumps(expected)) + upload() + rows_file = out / "trials.json" + rows = json.loads(rows_file.read_text()) if rows_file.exists() else [] + completed = {(r["repetition"], r["gap_ms"], r["mode"]) for r in rows} + plans = [ + (rep, gap, mode) + for rep in range(args.repetitions) + for gap in args.gaps_ms + for mode in "ABC" + ] + random.Random(42).shuffle(plans) + for rep, gap, mode in plans: + if (rep, gap, mode) in completed: + continue + label = f"{mode}-{gap}-{rep}" + directory = work / label + directory.mkdir() + if mode != "A": + for page in source.glob("*.bin"): + shutil.copyfile(page, directory / page.name) + proc = start(label, directory) + try: + tool_start = time.monotonic_ns() + accepted = None + if mode == "C": + accepted = control("submit", label, input_ids=prefix, ttl_ms=60000) + time.sleep( + max(0, (tool_start + gap * 1_000_000 - time.monotonic_ns()) / 1e9) + ) + response, timing = generate() + assert response["output_ids"] == expected, ( + label, + response["output_ids"], + expected, + ) + status = control("status", label) if mode == "C" else None + records = events(label) + arrival = timing["arrival_ns"] + io = [x for x in records if x["kind"] == "restore_io"] + reads = [x for x in records if x["kind"] == "read"] + keys = collections.Counter(k for r in reads for k in r["keys"]) + published = [x for x in records if x["kind"] == "publish"] + stats = dict( + mode=mode, + gap_ms=gap, + repetition=rep, + **timing, + actual_signal_to_arrival_ms=(arrival - tool_start) / 1e6, + restore_io_ms=sum((x["end_ns"] - x["start_ns"]) / 1e6 for x in io), + restore_hidden_ms=sum( + max(0, min(x["end_ns"], arrival) - x["start_ns"]) / 1e6 for x in io + ), + restored_bytes=sum(x["bytes"] for x in reads), + backend_read_pages=sum(keys.values()), + duplicate_backend_read_pages=sum(max(0, n - 1) for n in keys.values()), + read_pages_completed_before_arrival=sum( + len(x["keys"]) for x in reads if x["end_ns"] <= arrival + ), + published_before_arrival=any( + x["at_ns"] <= arrival and x["restored_tokens"] > 0 + for x in published + ), + restored_tokens=sum(x["restored_tokens"] for x in published), + host_evicted_tokens=sum( + x["evicted_tokens"] for x in records if x["kind"] == "evict_host" + ), + host_occupancy_snapshots=[ + { + k: x[k] + for k in [ + "kind", + "at_ns", + "host_used_tokens", + "inflight_tokens", + ] + } + for x in records + if "host_used_tokens" in x + ], + status=status, + cached_tokens_details=response["meta_info"].get( + "cached_tokens_details" + ), + server_first_token_latency_ms=response["meta_info"].get( + "first_token_latency", 0 + ) + * 1000, + output_correct=True, + ) + if mode == "A": + assert not reads, stats + if mode == "B": + assert io and all(x["start_ns"] >= arrival for x in io), stats + if mode == "C": + assert status["state"] in ("SUCCESS", "CACHED"), status + assert io, stats + # Query/allocation can overlap a short gap even if file reads + # start later. Validate submission of the entire restore, then + # measure actual I/O overlap independently (which may be zero). + assert all(x["operation_start_ns"] < arrival for x in io), stats + assert not any( + x["kind"] == "h2d_submit" and tool_start <= x["at_ns"] < arrival + for x in records + ), "Proactive H2D is out of scope" + rows.append(stats) + (out / "trials.json").write_text(json.dumps(rows, indent=2)) + print( + "TRIAL", + json.dumps( + { + k: stats[k] + for k in [ + "mode", + "gap_ms", + "repetition", + "ttft_ms", + "restore_io_ms", + "restore_hidden_ms", + "duplicate_backend_read_pages", + "published_before_arrival", + ] + } + ), + flush=True, + ) + upload() + finally: + stop(proc) + shutil.rmtree(directory) + directory = work / "wasted" + shutil.copytree(source, directory) + proc = start("wasted", directory) + try: + initial = control("submit", "no-continuation", input_ids=prefix, ttl_ms=10000) + deadline = time.monotonic() + 15 + while time.monotonic() < deadline: + status = control("status", "no-continuation") + if status["state"] != "RUNNING": + break + time.sleep(0.01) + assert status["state"] == "SUCCESS", status + assert status["inflight_tokens"] == 0 and not status["cleanup_pending"], status + after_cancel = control("cancel", "no-continuation") + waste = dict( + restored_tokens=status["restored_tokens"], + wasted_bytes=status["restored_bytes"], + host_available_token_delta=initial["host_available_tokens"] + - status["host_available_tokens"], + resident_after_cancel=after_cancel["state"], + policy="Completed restore is ordinary evictable L2, no pin or post-publication lease; entire staged prefix is wasted if continuation never arrives", + ) + request("/flush_cache", {}) + after_flush = control("status", "no-continuation") + waste["host_available_after_flush"] = after_flush["host_available_tokens"] + assert ( + after_flush["host_available_tokens"] >= initial["host_available_tokens"] + ), waste + (out / "wasted-prefetch.json").write_text(json.dumps(waste, indent=2)) + finally: + stop(proc) + summary = [] + for gap in args.gaps_ms: + bymode = { + mode: [x for x in rows if x["mode"] == mode and x["gap_ms"] == gap] + for mode in "ABC" + } + item = dict( + gap_ms=gap, + **{ + f"{mode}_ttft_median_ms": statistics.median( + x["ttft_ms"] for x in bymode[mode] + ) + for mode in "ABC" + }, + C_hidden_median_ms=statistics.median( + x["restore_hidden_ms"] for x in bymode["C"] + ), + B_restore_median_ms=statistics.median( + x["restore_io_ms"] for x in bymode["B"] + ), + ) + item["B_minus_C_ms"] = item["B_ttft_median_ms"] - item["C_ttft_median_ms"] + summary.append(item) + result = dict( + trials=len(rows), + repetitions=args.repetitions, + summary=summary, + all_outputs_correct=all(x["output_correct"] for x in rows), + duplicate_backend_read_pages=sum( + x["duplicate_backend_read_pages"] for x in rows + ), + host_evicted_tokens=sum(x["host_evicted_tokens"] for x in rows), + wasted_prefetch=waste, + scope="one L4, FULL resident file, TP1 PP1, no injected delay", + ) + (out / "benchmark.json").write_text(json.dumps(result, indent=2)) + print("BENCHMARK_RESULT", json.dumps(result), flush=True) + upload() + + +if __name__ == "__main__": + main() diff --git a/benchmark/hicache/trace/hicache_trace.py b/benchmark/hicache/trace/hicache_trace.py new file mode 100644 index 000000000000..fcb6d3219244 --- /dev/null +++ b/benchmark/hicache/trace/hicache_trace.py @@ -0,0 +1,127 @@ +"""Benchmark-only instrumentation, enabled through the existing plugin loader. + +No artificial I/O latency. All A/B/C runs receive the same hooks. +""" + +import functools +import json +import os +import time + + +def event(kind, **data): + path = os.environ.get("HICACHE_BENCH_TRACE") + if not path: + return + data.update( + kind=kind, + label=os.environ.get("HICACHE_BENCH_LABEL"), + pid=os.getpid(), + at_ns=time.monotonic_ns(), + ) + fd = os.open(path, os.O_WRONLY | os.O_CREAT | os.O_APPEND, 0o600) + try: + os.write(fd, (json.dumps(data, separators=(",", ":")) + "\n").encode()) + finally: + os.close(fd) + + +def occupancy(cache): + pool = cache.cache_controller.mem_pool_host + return dict( + host_used_tokens=pool.anchor_entry.host_pool.size - pool.available_size(), + inflight_tokens=cache.cache_controller.prefetch_tokens_occupied, + ) + + +def install(): + from sglang.srt.managers.scheduler import Scheduler + from sglang.srt.mem_cache.hicache_storage import HiCacheFile + from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import ( + HybridCacheController, + ) + from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache + + if getattr(HiCacheFile, "_benchmark_trace_installed", False): + return + HiCacheFile._benchmark_trace_installed = True + + original_get = HiCacheFile.batch_get + + @functools.wraps(original_get) + def get(self, keys, *args, **kwargs): + start = time.monotonic_ns() + result = original_get(self, keys, *args, **kwargs) + event( + "read", + start_ns=start, + end_ns=time.monotonic_ns(), + keys=list(keys), + bytes=sum(x.numel() * x.element_size() for x in result if x is not None), + ) + return result + + HiCacheFile.batch_get = get + + original_transfer = HybridCacheController._page_transfer + + @functools.wraps(original_transfer) + def transfer(self, operation): + start = time.monotonic_ns() + result = original_transfer(self, operation) + event( + "restore_io", + operation_start_ns=int(operation.start_time * 1e9), + start_ns=start, + end_ns=time.monotonic_ns(), + rid=operation.handle.rid, + pages=len(operation.hash_value), + ) + return result + + HybridCacheController._page_transfer = transfer + + original_publish = UnifiedRadixCache._handle_prefetch_result + + @functools.wraps(original_publish) + def publish(self, operation): + result = original_publish(self, operation) + event( + "publish", + rid=operation.handle.rid, + restored_tokens=self.prefetch_loaded_tokens_by_reqid.get( + operation.handle, 0 + ), + **occupancy(self), + ) + return result + + UnifiedRadixCache._handle_prefetch_result = publish + + original_evict = UnifiedRadixCache.evict_host + + @functools.wraps(original_evict) + def evict(self, *args, **kwargs): + result = original_evict(self, *args, **kwargs) + event("evict_host", evicted_tokens=result, **occupancy(self)) + return result + + UnifiedRadixCache.evict_host = evict + + original_load = UnifiedRadixCache.load_back + + @functools.wraps(original_load) + def load(self, *args, **kwargs): + event("h2d_submit", **occupancy(self)) + return original_load(self, *args, **kwargs) + + UnifiedRadixCache.load_back = load + + original_request = Scheduler.handle_generate_request + + @functools.wraps(original_request) + def request(self, obj): + event("generation_received", rid=obj.rid, **occupancy(self.tree_cache)) + return original_request(self, obj) + + Scheduler.handle_generate_request = request diff --git a/benchmark/hicache/trace/sglang_hicache_trace-0.0.dist-info/METADATA b/benchmark/hicache/trace/sglang_hicache_trace-0.0.dist-info/METADATA new file mode 100644 index 000000000000..f557002e0829 --- /dev/null +++ b/benchmark/hicache/trace/sglang_hicache_trace-0.0.dist-info/METADATA @@ -0,0 +1,3 @@ +Metadata-Version: 2.1 +Name: sglang-hicache-trace +Version: 0.0 diff --git a/benchmark/hicache/trace/sglang_hicache_trace-0.0.dist-info/entry_points.txt b/benchmark/hicache/trace/sglang_hicache_trace-0.0.dist-info/entry_points.txt new file mode 100644 index 000000000000..720f0b9fbf50 --- /dev/null +++ b/benchmark/hicache/trace/sglang_hicache_trace-0.0.dist-info/entry_points.txt @@ -0,0 +1,2 @@ +[sglang.srt.plugins] +hicache_trace = hicache_trace:install diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index 587b3e782f18..4f1cb6ad6211 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -141,6 +141,7 @@ ParseFunctionCallReq, PauseGenerationReqInput, PdRoleSwitchReqInput, + ProactivePrefetchReqInput, ProfileReq, ReleaseMemoryOccupationReqInput, ResumeMemoryOccupationReqInput, @@ -1062,6 +1063,21 @@ async def list_external_corpora(): # example usage: # curl -s -X POST http://127.0.0.1:30000/hicache/storage-backend/clear +@app.post("/hicache/prefetch") +@auth_level(AuthLevel.ADMIN_OPTIONAL) +async def proactive_prefetch(obj: Annotated[ProactivePrefetchReqInput, Body()]): + """Restore exact token-prefix KV into resident host cache before generation. + + Experimental single-worker FULL/file scope. Actions: submit, status, cancel. + This route allocates no generation request and never transfers KV to HBM. + """ + ret = await _global_state.tokenizer_manager.proactive_prefetch(obj) + return ORJSONResponse( + {"success": ret.success, "message": ret.message, "result": ret.result}, + status_code=200 if ret.success else HTTPStatus.BAD_REQUEST, + ) + + @app.api_route("/hicache/storage-backend/clear", methods=["POST"]) @auth_level(AuthLevel.ADMIN_OPTIONAL) async def clear_hicache_storage_backend(): diff --git a/python/sglang/srt/managers/cache_controller.py b/python/sglang/srt/managers/cache_controller.py index a05c7623dd5e..c6990cf73318 100644 --- a/python/sglang/srt/managers/cache_controller.py +++ b/python/sglang/srt/managers/cache_controller.py @@ -222,6 +222,8 @@ class PrefetchAck: # Number of hits in extra pools. pool_hits: Optional[dict[str, int]] = None completed_req: Optional[bool] = None + # Only emitted for a single-worker backend exception after I/O returns. + failed: bool = False class StorageOperation: @@ -249,6 +251,9 @@ def __init__( self.stats_requested_tokens = 0 # Absolute token offset at which this storage-prefetched span starts. self.storage_start = 0 + # Terminal consumption is scheduler-owned, including cancelled/stale ops. + self.terminal_ack_consumed = False + self.terminal_outcome: Optional[str] = None self.id = StorageOperation.counter StorageOperation.counter += 1 @@ -1125,7 +1130,12 @@ def _page_transfer(self, operation: PrefetchOperation) -> int: kv_derived_transfers, ) except Exception: - if not get_memory().enable_unified_memory: + # The resident single-worker drain distinguishes failure + # from a miss and reclaims acknowledged progress plus tail. + if ( + self._supports_local_prefetch_failure() + or not get_memory().enable_unified_memory + ): raise logger.exception( "HiCache prefetch transfer failed for request %s", @@ -1196,17 +1206,36 @@ def prefetch_io_aux_func(self): continue if operation is None: continue + failed = False try: self._page_transfer(operation) + except Exception: + # Local recovery does not change multi-rank collective ordering. + if not self._supports_local_prefetch_failure(): + raise + logger.exception( + "HiCache prefetch read failed: %s", operation.request_id + ) + failed = True finally: self.prefetch_sync_queue.put( PrefetchAck( rid=operation.request_id, completed_req=True, operation=operation, + failed=failed, ) ) + def _supports_local_prefetch_failure(self) -> bool: + # _create_sync_groups excludes single-rank groups. PP tickets also + # carry their own completion protocol, even before local allocation. + return getattr(self, "_supports_prefetch_failure_ack", False) and not ( + self.prefetch_hits_sync_groups + or self.prefetch_completion_sync_groups + or getattr(self, "pp_prefetch_command_group", None) is not None + ) + def prefetch_rate_limited(self) -> bool: """ Rate limit the prefetching operations to avoid overwhelming the storage backend. @@ -1287,6 +1316,19 @@ def prefetch_thread_func(self): operation ) except Exception: + if self._supports_local_prefetch_failure(): + logger.exception( + "HiCache prefetch query failed: %s", operation.request_id + ) + self.prefetch_sync_queue.put( + PrefetchAck( + rid=operation.request_id, + operation=operation, + completed_req=True, + failed=True, + ) + ) + continue if not get_memory().enable_unified_memory: raise logger.exception( diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index 9629b7aab621..b2b7866f4bd5 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -1743,6 +1743,20 @@ class BatchEmbeddingOutput(BaseBatchReq, kw_only=True): pooled_hidden_states: Optional[List[Optional[torch.Tensor]]] = None +class ProactivePrefetchReqInput(BaseReq, kw_only=True): + operation_id: str + action: str = "submit" + input_ids: Optional[List[int]] = None + cache_salt: Optional[str] = None + ttl_ms: int = 10000 + + +class ProactivePrefetchReqOutput(BaseReq, kw_only=True): + success: bool + message: str = "" + result: Optional[dict] = None + + class ClearHiCacheReqInput(BaseReq, kw_only=True): pass diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index b6a3aa45c95f..7e055a168806 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -162,6 +162,8 @@ OpenSessionReqInput, PauseGenerationReqInput, PdRoleSwitchReqInput, + ProactivePrefetchReqInput, + ProactivePrefetchReqOutput, ProfileReq, ReleaseMemoryOccupationReqInput, RemoveExternalCorpusReqInput, @@ -1712,6 +1714,7 @@ def init_request_dispatcher(self): (BatchTokenizedGenerateReqInput, self.handle_batch_generate_request), (BatchTokenizedEmbeddingReqInput, self.handle_batch_embedding_request), (FlushCacheReqInput, self.flush_wrapper.handle), + (ProactivePrefetchReqInput, self.handle_proactive_prefetch), (ClearHiCacheReqInput, self.clear_hicache_storage_wrapped), (AttachHiCacheStorageReqInput, self.attach_hicache_storage_wrapped), (DetachHiCacheStorageReqInput, self.detach_hicache_storage_wrapped), @@ -3129,6 +3132,9 @@ def handle_batch_generate_request( self.handle_generate_request(tokenized_req) def _prefetch_kvcache(self, req: Req, storage_hit_end: Optional[int] = None): + proactive = getattr(self, "proactive_prefetch", None) + if proactive is not None and proactive.blocks(req): + return # The existing control operation is already reading this prefix. if self.enable_hicache_storage or self.enable_lmcache: req.init_next_round_input(self.tree_cache, cow_mamba=False) tree_cache = self.tree_cache @@ -3633,6 +3639,9 @@ def _build_hisparse_decode_batch(self, reqs): def _process_hicache_events( self, should_retry_storage_prefetch: bool = True ) -> None: + proactive = getattr(self, "proactive_prefetch", None) + if proactive is not None: + proactive.tick() # Expire before terminal ACK can publish. # The HiCache drain is TP-wide consensus; run it before rank-local # decisions (_should_defer_prefill) or ranks enter different collectives. if ( @@ -3644,6 +3653,8 @@ def _process_hicache_events( self.tree_cache.check_hicache_events() if self.enable_hicache_storage and should_retry_storage_prefetch: self._process_storage_prefetch_retries() + if proactive is not None: + proactive.tick() # Consume terminal outcomes and control accounting. @scheduler_stage_method(SCHEDULER_STAGE_GET_NEXT_BATCH) def get_next_batch_to_run( @@ -3994,6 +4005,9 @@ def _get_new_batch_prefill_raw( ): break + proactive = getattr(self, "proactive_prefetch", None) + if proactive is not None and proactive.blocks(req): + continue # Rematch normally after resident L2 publication. if self.enable_hicache_storage or self.enable_lmcache: prefetch_done = self.tree_cache.check_prefetch_progress( req.cache_request_handle @@ -4908,6 +4922,69 @@ def list_external_corpora( ) return self.external_corpus_manager.list(recv_req) + def handle_proactive_prefetch(self, obj: ProactivePrefetchReqInput): + from sglang.srt.mem_cache.proactive_prefetch import ProactivePrefetch + from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache + + try: + parallel = get_parallel() + if ( + any( + getattr(parallel, name) != 1 + for name in ( + "tp_size", + "pp_size", + "dp_size", + "nnodes", + "attn_cp_size", + "attn_dp_size", + ) + ) + or self.disaggregation_mode != DisaggregationMode.NULL + or get_spec().speculative_algorithm is not None + or get_lora().enable_lora + or self.model_config.is_multimodal + or not self.model_config.is_generation + or not isinstance(self.tree_cache, UnifiedRadixCache) + ): + raise ValueError( + "Requires single-worker ordinary full-attention HiCache" + ) + if obj.action == "submit": + if self._engine_paused: + raise ValueError("Engine is paused") + if ( + not obj.input_ids + or len(obj.input_ids) > self.max_req_input_len + or any( + type(t) is not int or not 0 <= t < self.model_config.vocab_size + for t in obj.input_ids + ) + ): + raise ValueError( + "Provide a bounded exact token prefix within the vocabulary" + ) + if getattr(self, "proactive_prefetch", None) is None: + self.proactive_prefetch = ProactivePrefetch(self.tree_cache) + result = self.proactive_prefetch.submit( + obj.operation_id, obj.input_ids, obj.cache_salt, obj.ttl_ms + ) + elif obj.action in ("status", "cancel"): + manager = getattr(self, "proactive_prefetch", None) + if manager is None: + raise KeyError(obj.operation_id) + manager.tick() + result = ( + manager.status(obj.operation_id) + if obj.action == "status" + else manager.cancel(obj.operation_id) + ) + else: + raise ValueError("action must be submit, status or cancel") + return ProactivePrefetchReqOutput(success=True, result=result) + except (ValueError, KeyError) as exc: + return ProactivePrefetchReqOutput(success=False, message=str(exc)) + def clear_hicache_storage_wrapped(self, recv_req: ClearHiCacheReqInput): if self.enable_hierarchical_cache or self.enable_lmcache: self.tree_cache.clear_storage_backend() @@ -5009,6 +5086,9 @@ def on_idle(self): self.metrics_reporter.record_scheduler_idle() def _record_scheduler_state_for_paused_engine(self) -> None: + proactive = getattr(self, "proactive_prefetch", None) + if proactive is not None and proactive.active is not None: + self._process_hicache_events() # Drain cancelled control I/O while paused. if self.is_fully_idle(): self.metrics_reporter.record_scheduler_idle() else: @@ -5043,6 +5123,8 @@ def is_fully_idle(self, for_health_check=False, ignore_waiting=False) -> bool: idle &= len(self.disagg_decode_prealloc_queue.retracted_queue) == 0 if not for_health_check: + proactive = getattr(self, "proactive_prefetch", None) + idle &= proactive is None or proactive.active is None # Grammar queue and prefill inflight queue may not produce batch # results instantly, but they still indicate the server is not idle. idle &= len(self.grammar_manager.grammar_queue) == 0 @@ -5447,6 +5529,9 @@ def _pause_engine(self) -> Tuple[List[Req], int]: raise NotImplementedError() def pause_generation(self, recv_req: PauseGenerationReqInput): + proactive = getattr(self, "proactive_prefetch", None) + if proactive is not None and proactive.active is not None: + proactive.cancel(proactive.active.operation_id) assert recv_req.mode in ("in_place", "retract") self._engine_paused = True diff --git a/python/sglang/srt/managers/tokenizer_control_mixin.py b/python/sglang/srt/managers/tokenizer_control_mixin.py index f62afbae8ed6..a8cd4c089860 100644 --- a/python/sglang/srt/managers/tokenizer_control_mixin.py +++ b/python/sglang/srt/managers/tokenizer_control_mixin.py @@ -55,6 +55,8 @@ OpenSessionReqInput, PdRoleSwitchReqInput, PdRoleSwitchReqOutput, + ProactivePrefetchReqInput, + ProactivePrefetchReqOutput, ProfileReq, ProfileReqOutput, ProfileReqType, @@ -128,6 +130,7 @@ ("add_external_corpus", AddExternalCorpusReqOutput), ("remove_external_corpus", RemoveExternalCorpusReqOutput), ("list_external_corpora", ListExternalCorporaReqOutput), + ("proactive_prefetch", ProactivePrefetchReqOutput), ("clear_hicache_storage", ClearHiCacheReqOutput), ("attach_hicache_storage", AttachHiCacheStorageReqOutput), ("detach_hicache_storage", DetachHiCacheStorageReqOutput), @@ -329,6 +332,27 @@ async def flush_cache( self.mm_processor.clear_preprocess_cache() return result + async def proactive_prefetch( + self: TokenizerManager, obj: ProactivePrefetchReqInput + ): + parallel = get_parallel() + if any( + getattr(parallel, name) != 1 + for name in ( + "tp_size", + "pp_size", + "dp_size", + "nnodes", + "attn_cp_size", + "attn_dp_size", + ) + ): + return ProactivePrefetchReqOutput( + success=False, message="Single worker required" + ) + self.auto_create_handle_loop() + return (await self.proactive_prefetch_communicator(obj))[0] + async def clear_hicache_storage(self: TokenizerManager) -> ClearHiCacheReqOutput: """Clear the hierarchical cache storage.""" self.auto_create_handle_loop() diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py index ec4d7390dede..04c997544255 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py @@ -211,6 +211,13 @@ def __init__( host_memory_mode: str = "cache", ): startup_storage_backend = storage_backend + # Only the resident Unified FULL drain currently consumes failure ACKs. + # Keep legacy and multi-pool consumers on their existing protocol. + self._supports_prefetch_failure_ack = ( + host_memory_mode == "cache" + and len(mem_pool_host.entries) == 1 + and mem_pool_host.anchor_entry.name == PoolName.KV + ) self.extra_host_mem_release_queues: dict[PoolName, Queue[torch.Tensor]] = {} self.pp_prefetch_command_group = None self.pp_prefetch_command_thread = None diff --git a/python/sglang/srt/mem_cache/proactive_prefetch.py b/python/sglang/srt/mem_cache/proactive_prefetch.py new file mode 100644 index 000000000000..1a8de14e3609 --- /dev/null +++ b/python/sglang/srt/mem_cache/proactive_prefetch.py @@ -0,0 +1,202 @@ +"""Bounded requestless L3-to-resident-L2 restoration; scheduler-thread only.""" + +import time +import uuid +from array import array +from collections import OrderedDict +from dataclasses import dataclass +from typing import Optional + +from sglang.srt.mem_cache.base_prefix_cache import CacheRequestHandle, MatchPrefixParams +from sglang.srt.mem_cache.radix_cache import RadixKey +from sglang.srt.mem_cache.unified_cache.components import ComponentType + + +@dataclass +class _Restore: + operation_id: str + tokens: tuple[int, ...] + cache_salt: Optional[str] + ttl_ms: int + handle: CacheRequestHandle + started: float + deadline: float + operation: object = None + state: str = "RUNNING" + finished: Optional[float] = None + restored_tokens: int = 0 + + +class ProactivePrefetch: + """One active restore, with bounded recent outcomes and ordinary ownership. + + TTL limits in-flight work only. Published host pages are normally evictable; + cancellation after publication does not evict or pin shared cache state. + """ + + def __init__(self, cache): + if ( + not cache.enable_storage + or cache.host_memory_mode != "cache" + or set(cache.tree_components) != {ComponentType.FULL} + or cache.sidecar_pool_specs + or cache.tree_core.is_eagle + or cache.cache_controller.storage_backend.__class__.__name__ + != "HiCacheFile" + ): + raise ValueError("Requires resident FULL KV-only HiCache with file storage") + self.cache = cache + self._backend = cache.cache_controller.storage_backend + self.active: Optional[_Restore] = None + self._waiting_req = None + self.records: OrderedDict[str, _Restore] = OrderedDict() + + def tick(self): + record = self.active + if record is None: + return + if record.state == "RUNNING" and time.monotonic() >= record.deadline: + self.cancel(record.operation_id, state="EXPIRED") + if record.handle in self.cache.ongoing_prefetch: + return + if record.state == "RUNNING": + record.restored_tokens, _ = self.cache.pop_prefetch_loaded_span( + record.handle + ) + record.state = record.operation.terminal_outcome or "MISS" + record.finished = time.monotonic() + self.cache.discard_storage_prefetch_accounting(record.handle) + self.cache.storage_prefetch_retries.cancel(record.handle.rid) + # An allocated cancelled read still owns its tail until terminal ACK. + # Keep admission of another control operation bounded until it drains. + if ( + record.operation.host_indices is None + or record.operation.terminal_ack_consumed + ): + self.active = None + self._waiting_req = None + + def submit(self, operation_id, input_ids, cache_salt=None, ttl_ms=10000): + if ( + not self.cache.enable_storage + or self.cache.cache_controller.storage_backend is not self._backend + ): + raise ValueError("Storage backend changed; restart the fixed-model server") + self.tick() + if not operation_id or len(operation_id) > 128: + raise ValueError("operation_id must contain 1..128 characters") + if cache_salt is not None and len(cache_salt) > 256: + raise ValueError("cache_salt must contain at most 256 characters") + if not 1 <= ttl_ms <= 60000: + raise ValueError("ttl_ms must be in 1..60000") + key = RadixKey(array("q", input_ids), cache_salt=cache_salt).page_aligned( + self.cache.page_size + ) + if len(key) < self.cache.prefetch_threshold: + raise ValueError("Prefix is shorter than the storage prefetch threshold") + tokens = tuple(key.token_ids) + if operation_id in self.records: + record = self.records[operation_id] + if (record.tokens, record.cache_salt, record.ttl_ms) != ( + tokens, + key.cache_salt, + ttl_ms, + ): + raise ValueError("operation_id already identifies a different restore") + return self.status(operation_id) + if self.active is not None: + raise ValueError("One proactive restore is already active") + match = self.cache.match_prefix(MatchPrefixParams(key=key)) + matched = len(match.device_indices) + match.host_hit_length + anchor = match.last_host_node + if matched < len(key) and not ( + self.cache.is_root(anchor) or self.cache.is_backuped(anchor) + ): + raise ValueError("The matched anchor is not backed by host KV") + now = time.monotonic() + record = _Restore( + operation_id, + tokens, + key.cache_salt, + ttl_ms, + CacheRequestHandle("__proactive__" + uuid.uuid4().hex, 0), + now, + now + ttl_ms / 1000, + ) + if matched >= len(key): + record.state, record.finished = "CACHED", now + else: + if len(key) - matched < self.cache.prefetch_threshold: + raise ValueError("Uncached suffix is below the prefetch threshold") + self.cache.prefetch_from_storage( + record.handle, + anchor, + array("q", tokens[matched:]), + self.cache.get_last_hash_value(anchor), + ( + self.cache.get_prefix_hash_values(anchor) + if self.cache.hicache_storage_pass_prefix_keys + else None + ), + matched_prefix_tokens=array("q", tokens[:matched]), + cache_salt=key.cache_salt, + ) + info = self.cache.ongoing_prefetch.get(record.handle) + if info is None: + self.cache.storage_prefetch_retries.cancel(record.handle.rid) + record.state, record.finished = "DECLINED", now + else: + record.operation = info.operation + self.active = record + self.records[operation_id] = record + while len(self.records) > 32: + self.records.popitem(last=False) + return self.status(operation_id) + + def cancel(self, operation_id, state="CANCELLED"): + record = self.records[operation_id] + if record.state == "RUNNING": + # Ownership/ACK cleanup is entirely the existing controller contract. + self._waiting_req = None + self.cache.release_aborted_request(record.handle) + record.state, record.finished = state, time.monotonic() + return self.status(operation_id) + + def blocks(self, req): + # A queued Req keeps its original token prefix. Compare it once, not on + # every idle scheduling poll (which otherwise contends with file I/O). + if self.active is None or self.active.state != "RUNNING": + return False + if req is self._waiting_req: + return True + if self.waits_for(req): + self._waiting_req = req + return True + return False + + def waits_for(self, req): + record = self.active + if record is None or record.state != "RUNNING": + return False + return ( + req.extra_key is None + and (req.cache_salt or None) == record.cache_salt + and tuple(req.origin_input_ids[: len(record.tokens)]) == record.tokens + ) + + def status(self, operation_id): + record = self.records[operation_id] + pool = self.cache.cache_controller.mem_pool_host + end = record.finished or time.monotonic() + return { + "operation_id": operation_id, + "state": record.state, + "requested_tokens": len(record.tokens), + "restored_tokens": record.restored_tokens, + "restored_bytes": record.restored_tokens + * pool.anchor_entry.host_pool.get_size_per_token(), + "elapsed_ms": (end - record.started) * 1000, + "cleanup_pending": record is self.active and record.state != "RUNNING", + "host_available_tokens": pool.available_size(), + "inflight_tokens": self.cache.cache_controller.prefetch_tokens_occupied, + } diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index 5b1475d90d68..cc1d7f39bb18 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -2578,6 +2578,8 @@ def release_aborted_request(self, request: CacheRequestHandle) -> None: return info = self.ongoing_prefetch[request] + if info.operation.terminal_outcome is None: + info.operation.terminal_outcome = "CANCELLED" if info.operation.host_indices is None: self.cache_controller.terminate_prefetch(info.operation) self.revoke_pending_prefetch(request) @@ -2754,6 +2756,8 @@ def revoke_pending_prefetch(self, request: CacheRequestHandle) -> None: info = self.ongoing_prefetch.get(request) if info is None: return + if info.operation.terminal_outcome is None: + info.operation.terminal_outcome = "MISS" self._invalidate_absent_from_hit_query(info.operation) # Every revoke path runs before the bounce alloc, so buffer mode # holds no occupancy here; post-alloc aborts go through @@ -2985,6 +2989,8 @@ def _drain_and_alloc_storage_hit(): def _drain_ack_prefetch(): for ack in _drain_queue(cc.ack_prefetch_queue, n_ack_prefetch): operation = ack.operation + if operation.terminal_ack_consumed: + continue info = self.ongoing_prefetch.get(operation.handle) is_current = info is not None and info.operation is operation if ack.completed_tokens is not None: @@ -2998,11 +3004,27 @@ def _drain_ack_prefetch(): ) operation.pool_transfers_done = True if ack.completed_req: - if is_current: + operation.terminal_ack_consumed = True + if operation.terminal_outcome is None: + operation.terminal_outcome = ( + "FAILURE" + if ack.failed + else "SUCCESS" + if operation.completed_tokens > 0 + else "MISS" + ) + if is_current and ack.failed: + # The backend has returned/raised: abort frees only + # acknowledged pages; the terminal drain owns the tail. + self.release_aborted_request(operation.handle) + elif is_current: # check_prefetch_progress() is not called for this rid yet. # Let us insert the prefetch result into the radix tree. self._handle_prefetch_result(operation) - if operation.ack_releases_incomplete_host_indices: + if ( + operation.host_indices is not None + and operation.ack_releases_incomplete_host_indices + ): cc.append_host_mem_release( operation.host_indices[operation.completed_tokens :], ( diff --git a/test/registered/unit/mem_cache/test_prefetch_finite_io.py b/test/registered/unit/mem_cache/test_prefetch_finite_io.py new file mode 100644 index 000000000000..11e27a00b63f --- /dev/null +++ b/test/registered/unit/mem_cache/test_prefetch_finite_io.py @@ -0,0 +1,370 @@ +"""Finite-I/O regressions with real Unified cache, host pool and file storage. + +Run on CPU with SGLANG_USE_CPU_ENGINE=1 and Python TreeCore. CPU runs use +unpinned host tensors; no GPU transfer or model inference is claimed. +""" + +import tempfile +import threading +import time +import unittest +from array import array +from pathlib import Path +from unittest import mock + +import test_unified_radix_cache_unittest as fixtures +import torch + +from sglang.srt.managers.cache_controller import PrefetchAck +from sglang.srt.mem_cache.base_prefix_cache import CacheRequestHandle, MatchPrefixParams +from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost +from sglang.srt.mem_cache.radix_cache import RadixKey +from sglang.srt.mem_cache.utils import get_storage_hash_str +from sglang.srt.runtime_context import get_memory +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=15, suite="base-a-test-cpu") + + +class TestFiniteIO(fixtures.CustomTestCase): + cfg = fixtures.CacheConfig( + page_size=4, + num_layers=2, + head_num=1, + head_dim=16, + kv_size=64, + max_context_len=64, + ) + + def setUp(self): + self.directory = tempfile.TemporaryDirectory() + self.addCleanup(self.directory.cleanup) + self.cache, self.allocator, _ = fixtures.build_fixture(self.cfg) + original = MHATokenToKVPoolHost.__init__ + + def cpu_host(pool, *args, **kwargs): + if not torch.cuda.is_available(): + kwargs["pin_memory"] = False + original(pool, *args, **kwargs) + + # Production reserves 10 GiB; this fixture allocates less than 64 KiB. + # Ignore that deployment admission floor on constrained CPU runners. + with ( + mock.patch.object(MHATokenToKVPoolHost, "__init__", cpu_host), + mock.patch( + "sglang.srt.mem_cache.pool_host.base.host_memory_budget_bytes", + side_effect=lambda requested: requested, + ), + ): + fixtures.UnifiedRadixCacheSuite._init_hicache( + self, + self.cache, + storage_backend="file", + storage_dir=self.directory.name, + prefetch_threshold=4, + ) + # Match the running Unified cache flag: upstream catches some batch + # exceptions only when this flag is enabled. + unified_memory = get_memory().override(enable_unified_memory=True) + unified_memory.__enter__() + self.addCleanup(unified_memory.__exit__, None, None, None) + self.cc = self.cache.cache_controller + self.pool = self.cc.mem_pool_host + self.initial_slots = self.pool.available_size() + self.cache.enable_storage_metrics = True + self.cache.storage_metrics_collector = mock.Mock() + self.tokens = array("q", range(1, 13)) + self.key = RadixKey(self.tokens) + self.hashes = get_storage_hash_str(self.key, page_size=4) + self.backend = self.cc.storage_backend + host = self.pool.anchor_entry.host_pool + page = host.get_dummy_flat_data_page() + page.fill_(7) + for key in self.hashes: + self.assertTrue(self.backend.set(key, page)) + self.handle_counter = 0 + + def submit(self, handle=None): + self.handle_counter += 1 + handle = handle or CacheRequestHandle(f"finite-{self.handle_counter}", 0) + self.cache.prefetch_from_storage( + handle, self.cache.root_node_handle(), self.tokens, None, None + ) + return handle + + def pump_until(self, predicate, timeout=5): + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + self.cache.drain_storage_control_queues() + if predicate(): + return + time.sleep(0.005) + self.fail("prefetch lifecycle did not settle") + + def settle(self, handle): + self.pump_until(lambda: handle not in self.cache.ongoing_prefetch) + self.cache.drain_storage_control_queues() + + def conservation(self, handle, resident=0): + if resident: + self.cache.finish_storage_prefetch_admission(handle, resident, reason=None) + self.cache.pop_prefetch_loaded_span(handle) + self.cc.mem_pool_host.anchor_entry.host_pool._merge_release_slots() + self.assertEqual(self.pool.available_size(), self.initial_slots - resident) + self.assertEqual( + int(self.pool.anchor_entry.host_pool.slot_used.sum()), resident + ) + self.assertEqual(self.cache.ongoing_prefetch, {}) + self.assertEqual(self.cc.prefetch_tokens_occupied, 0) + self.assertFalse(self.cache._storage_prefetch_hit_remaining_by_reqid) + self.assertFalse(self.cache.prefetch_loaded_tokens_by_reqid) + self.assertFalse(self.cache.prefetch_loaded_storage_start_by_reqid) + self.cache.sanity_check() + + def next_prefetch(self): + handle = self.submit() + self.settle(handle) + match = self.cache.match_prefix(MatchPrefixParams(key=self.key)) + self.assertEqual(match.host_hit_length, len(self.tokens)) + self.assertTrue(self.cc.prefetch_io_aux_thread.is_alive()) + node = self.cache.tree_core.node_by_id(match.last_host_node) + slots = node.component_data[fixtures.ComponentType.FULL].host_value + page = self.pool.anchor_entry.host_pool.get_data_page(int(slots[0]), flat=True) + torch.testing.assert_close(page, torch.full_like(page, 7)) + self.conservation(handle, resident=len(self.tokens)) + + def test_finite_read_exception_worker_survives(self): + with mock.patch.object( + self.backend, "batch_get", side_effect=RuntimeError("read failed") + ): + handle = self.submit() + self.settle(handle) + self.conservation(handle) + self.next_prefetch() + + def test_finite_short_file_worker_survives(self): + path = Path(self.directory.name) / ( + self.backend._get_suffixed_key(self.hashes[0]) + ".bin" + ) + original = path.read_bytes() + path.write_bytes(b"x") + handle = self.submit() + self.settle(handle) + self.conservation(handle) + path.write_bytes(original) + self.next_prefetch() + + def test_finite_query_exception_worker_survives(self): + with mock.patch.object( + self.backend, "batch_exists", side_effect=OSError("query failed") + ): + handle = self.submit() + self.settle(handle) + self.conservation(handle) + self.assertTrue(self.cc.prefetch_thread.is_alive()) + self.next_prefetch() + + def test_finite_file_disappears_after_query(self): + original = self.backend.batch_get + + def disappeared(*args, **kwargs): + for p in Path(self.directory.name).glob("*.bin"): + p.unlink() + return original(*args, **kwargs) + + with mock.patch.object(self.backend, "batch_get", side_effect=disappeared): + handle = self.submit() + self.settle(handle) + self.conservation(handle) + + def test_finite_cancel_running_late_success_and_duplicate_terminal(self): + entered, resume = threading.Event(), threading.Event() + original = self.backend.batch_get + + def blocked(*args, **kwargs): + entered.set() + if not resume.wait(10): + raise RuntimeError("test gate not released") + return original(*args, **kwargs) + + try: + with mock.patch.object(self.backend, "batch_get", side_effect=blocked): + handle = self.submit() + self.pump_until(entered.is_set) + operation = self.cache.ongoing_prefetch[handle].operation + self.cache.release_aborted_request(handle) + self.cache.drain_storage_control_queues() + self.assertLess(self.pool.available_size(), self.initial_slots) + self.assertEqual( + self.cache.match_prefix( + MatchPrefixParams(key=self.key) + ).host_hit_length, + 0, + ) + resume.set() + self.pump_until( + lambda: self.pool.available_size() == self.initial_slots + ) + self.conservation(handle) + terminal = PrefetchAck(operation.request_id, operation, completed_req=True) + self.cc.ack_prefetch_queue.put(terminal) + self.cache.drain_storage_control_queues() + self.conservation(handle) + self.assertEqual( + self.cache.match_prefix( + MatchPrefixParams(key=self.key) + ).host_hit_length, + 0, + ) + self.next_prefetch() + finally: + resume.set() + + def test_finite_cancel_before_allocation(self): + entered, resume = threading.Event(), threading.Event() + original = self.backend.batch_exists + + def blocked(*args, **kwargs): + entered.set() + if not resume.wait(10): + raise RuntimeError("test gate not released") + return original(*args, **kwargs) + + try: + with mock.patch.object(self.backend, "batch_exists", side_effect=blocked): + handle = self.submit() + self.pump_until(entered.is_set) + operation = self.cache.ongoing_prefetch[handle].operation + self.assertIsNone(operation.host_indices) + self.cache.release_aborted_request(handle) + self.conservation(handle) + resume.set() + self.pump_until(lambda: operation.storage_hit_count > 0) + self.conservation(handle) + self.next_prefetch() + finally: + resume.set() + + def test_finite_cancel_allocated_before_read(self): + entered, resume = threading.Event(), threading.Event() + original = self.cc._page_transfer + + def blocked(operation): + entered.set() + if not resume.wait(10): + raise RuntimeError("test gate not released") + return original(operation) + + try: + with ( + mock.patch.object(self.cc, "_page_transfer", side_effect=blocked), + mock.patch.object( + self.backend, "batch_get", wraps=self.backend.batch_get + ) as read, + ): + handle = self.submit() + self.pump_until(entered.is_set) + operation = self.cache.ongoing_prefetch[handle].operation + self.assertIsNotNone(operation.host_indices) + self.cache.release_aborted_request(handle) + self.cache.drain_storage_control_queues() + self.assertLess(self.pool.available_size(), self.initial_slots) + resume.set() + self.pump_until( + lambda: self.pool.available_size() == self.initial_slots + ) + read.assert_not_called() + self.conservation(handle) + self.next_prefetch() + finally: + resume.set() + + def test_finite_partial_progress_then_failure(self): + original = self.backend.batch_get + calls = 0 + + def partial(*args, **kwargs): + nonlocal calls + calls += 1 + if calls == 2: + raise RuntimeError("second batch failed") + return original(*args, **kwargs) + + with ( + mock.patch("sglang.srt.managers.cache_controller.STORAGE_BATCH_SIZE", 1), + mock.patch.object(self.backend, "batch_get", side_effect=partial), + ): + handle = self.submit() + operation = self.cache.ongoing_prefetch[handle].operation + self.settle(handle) + self.assertEqual(operation.completed_tokens, 4) + self.conservation(handle) + self.assertEqual(operation.terminal_outcome, "FAILURE") + self.assertTrue(operation.terminal_ack_consumed) + self.next_prefetch() + + def test_finite_failure_releases_nonroot_anchor_lock(self): + full_tokens, full_key = self.tokens, self.key + self.tokens = self.tokens[:4] + handle = self.submit() + self.settle(handle) + self.conservation(handle, resident=4) + self.tokens, self.key = full_tokens, full_key + anchor = self.cache.match_prefix(MatchPrefixParams(key=self.key)).last_host_node + node = self.cache.tree_core.node_by_id(anchor) + full = node.component_data[fixtures.ComponentType.FULL] + baseline = full.host_lock_ref + with mock.patch.object( + self.backend, "batch_get", side_effect=RuntimeError("anchor read failed") + ): + handle = CacheRequestHandle("anchor-failure", 0) + self.cache.prefetch_from_storage( + handle, + anchor, + self.tokens[4:], + self.cache.get_last_hash_value(anchor), + None, + matched_prefix_tokens=list(self.tokens[:4]), + ) + self.settle(handle) + self.assertEqual(full.host_lock_ref, baseline) + self.conservation(handle, resident=4) + + def test_finite_late_completion_keeps_replacement_operation(self): + entered, resume = threading.Event(), threading.Event() + original = self.backend.batch_get + calls = 0 + + def blocked(*args, **kwargs): + nonlocal calls + calls += 1 + if calls == 1: + entered.set() + if not resume.wait(10): + raise RuntimeError("test gate not released") + return original(*args, **kwargs) + + try: + with mock.patch.object(self.backend, "batch_get", side_effect=blocked): + handle = self.submit() + self.pump_until(entered.is_set) + old = self.cache.ongoing_prefetch[handle].operation + self.cache.release_aborted_request(handle) + self.submit(handle) + replacement = self.cache.ongoing_prefetch[handle].operation + self.assertIsNot(old, replacement) + resume.set() + self.settle(handle) + self.assertEqual( + self.cache.match_prefix( + MatchPrefixParams(key=self.key) + ).host_hit_length, + len(self.tokens), + ) + self.conservation(handle, resident=len(self.tokens)) + finally: + resume.set() + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/mem_cache/test_proactive_prefetch.py b/test/registered/unit/mem_cache/test_proactive_prefetch.py new file mode 100644 index 000000000000..97dc741a246b --- /dev/null +++ b/test/registered/unit/mem_cache/test_proactive_prefetch.py @@ -0,0 +1,208 @@ +"""Real file/controller fixtures: requestless publication and bounded cleanup.""" + +import threading +import unittest +from types import SimpleNamespace +from unittest import mock + +import test_prefetch_finite_io as finite + +from sglang.srt.mem_cache.base_prefix_cache import MatchPrefixParams +from sglang.srt.mem_cache.proactive_prefetch import ProactivePrefetch +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=20, suite="base-a-test-cpu") + + +class TestProactivePrefetch(finite.TestFiniteIO): + def setUp(self): + super().setUp() + self.manager = ProactivePrefetch(self.cache) + + def advance(self, predicate): + def done(): + self.manager.tick() + return predicate() + + self.pump_until(done) + self.cache.drain_storage_control_queues() + self.manager.tick() + + def restore(self, operation_id="p", tokens=None, **kwargs): + self.manager.submit(operation_id, tokens or list(self.tokens), **kwargs) + self.advance(lambda: self.manager.status(operation_id)["state"] != "RUNNING") + return self.manager.status(operation_id) + + def test_requestless_publication_and_accounting(self): + before = self.allocator.available_size() + result = self.restore() + self.assertEqual(result["state"], "SUCCESS") + self.assertEqual(result["restored_tokens"], 12) + self.assertGreater(result["restored_bytes"], 0) + match = self.cache.match_prefix(MatchPrefixParams(key=self.key)) + self.assertEqual(match.host_hit_length, 12) + self.assertEqual(len(match.device_indices), 0) + self.assertEqual(self.allocator.available_size(), before) + self.conservation(self.manager.records["p"].handle, resident=12) + + def test_idempotency_and_conflicting_identity(self): + result = self.restore() + with mock.patch.object( + self.backend, "batch_get", side_effect=AssertionError("duplicate read") + ): + self.assertEqual( + self.manager.submit("p", list(self.tokens))["state"], "SUCCESS" + ) + self.assertEqual(self.restore("cached")["state"], "CACHED") + with self.assertRaisesRegex(ValueError, "different restore"): + self.manager.submit("p", [9] * 12) + self.assertEqual(result["restored_tokens"], 12) + + def test_nonroot_restore_extends_existing_host_prefix(self): + self.assertEqual( + self.restore("first", list(self.tokens[:4]))["restored_tokens"], 4 + ) + self.assertEqual(self.restore("rest")["restored_tokens"], 8) + self.assertEqual( + self.cache.match_prefix(MatchPrefixParams(key=self.key)).host_hit_length, 12 + ) + self.conservation(self.manager.records["rest"].handle, resident=12) + + def test_namespace_miss_does_not_publish_or_leak(self): + result = self.restore(cache_salt="other-session") + self.assertEqual(result["state"], "MISS") + self.conservation(self.manager.records["p"].handle) + self.assertEqual( + self.cache.match_prefix(MatchPrefixParams(key=self.key)).host_hit_length, 0 + ) + + def test_failure_terminal_does_not_leave_control_accounting(self): + with mock.patch.object( + self.backend, "batch_get", side_effect=OSError("read failed") + ): + self.assertEqual(self.restore()["state"], "FAILURE") + self.conservation(self.manager.records["p"].handle) + self.assertEqual(self.restore("next")["state"], "SUCCESS") + + def test_cancel_running_read_keeps_tail_until_ack(self): + entered, resume = threading.Event(), threading.Event() + original = self.backend.batch_get + + def read(*args, **kwargs): + entered.set() + self.assertTrue(resume.wait(5)) + return original(*args, **kwargs) + + try: + with mock.patch.object(self.backend, "batch_get", side_effect=read): + self.manager.submit("p", list(self.tokens)) + self.advance(entered.is_set) + self.assertTrue(self.manager.cancel("p")["cleanup_pending"]) + with self.assertRaisesRegex(ValueError, "already active"): + self.manager.submit("new", list(self.tokens)) + resume.set() + self.advance(lambda: self.manager.active is None) + finally: + resume.set() + self.assertEqual(self.manager.status("p")["state"], "CANCELLED") + self.assertEqual( + self.cache.match_prefix(MatchPrefixParams(key=self.key)).host_hit_length, 0 + ) + self.conservation(self.manager.records["p"].handle) + + def test_expiry_precedes_late_success_publication(self): + entered, resume = threading.Event(), threading.Event() + original = self.backend.batch_get + + def read(*args, **kwargs): + entered.set() + self.assertTrue(resume.wait(5)) + return original(*args, **kwargs) + + try: + with mock.patch.object(self.backend, "batch_get", side_effect=read): + self.manager.submit("p", list(self.tokens)) + self.advance(entered.is_set) + self.manager.active.deadline = 0 + self.manager.tick() + self.assertEqual(self.manager.status("p")["state"], "EXPIRED") + resume.set() + self.advance(lambda: self.manager.active is None) + finally: + resume.set() + self.conservation(self.manager.records["p"].handle) + self.assertEqual( + self.cache.match_prefix(MatchPrefixParams(key=self.key)).host_hit_length, 0 + ) + + def test_early_continuation_joins_only_matching_namespace_prefix(self): + self.manager.submit("p", list(self.tokens)) + req = SimpleNamespace( + origin_input_ids=list(self.tokens) + [13], extra_key=None, cache_salt=None + ) + self.assertTrue(self.manager.waits_for(req)) + req.cache_salt = "unrelated" + self.assertFalse(self.manager.waits_for(req)) + req.cache_salt = None + req.origin_input_ids[0] = 999 + self.assertFalse(self.manager.waits_for(req)) + self.advance(lambda: self.manager.active is None) + self.assertFalse(self.manager.waits_for(req)) + + def test_join_polling_does_not_rescan_or_join_replacement_request(self): + self.manager.submit("p", list(self.tokens)) + req = SimpleNamespace( + origin_input_ids=list(self.tokens) + [13], extra_key=None, cache_salt=None + ) + with mock.patch.object( + self.manager, "waits_for", wraps=self.manager.waits_for + ) as match: + self.assertTrue(self.manager.blocks(req)) + self.assertTrue(self.manager.blocks(req)) + self.assertEqual(match.call_count, 1) + replacement = SimpleNamespace( + origin_input_ids=[999], extra_key=None, cache_salt=None + ) + self.assertFalse(self.manager.blocks(replacement)) + self.manager.cancel("p") + self.assertFalse(self.manager.blocks(req)) + self.assertIsNone(self.manager._waiting_req) + self.advance(lambda: self.manager.active is None) + + def test_control_registry_is_bounded(self): + self.restore() + for i in range(50): + self.assertEqual(self.restore(str(i))["state"], "CACHED") + self.assertEqual(len(self.manager.records), 32) + + def test_backend_replacement_rejects_new_submit(self): + self.restore() + with mock.patch.object( + self.cache.cache_controller, "storage_backend", object() + ): + with self.assertRaisesRegex(ValueError, "backend changed"): + self.manager.submit("next", list(self.tokens)) + self.assertEqual(self.manager.status("p")["state"], "SUCCESS") + + def test_terminal_join_releases_request_reference(self): + self.manager.submit("p", list(self.tokens)) + req = SimpleNamespace( + origin_input_ids=list(self.tokens) + [13], extra_key=None, cache_salt=None + ) + self.assertTrue(self.manager.blocks(req)) + self.advance(lambda: self.manager.active is None) + self.assertIsNone(self.manager._waiting_req) + + def test_rejects_buffer_mode_and_non_file_backend(self): + with mock.patch.object(self.cache, "host_memory_mode", "buffer_only"): + with self.assertRaisesRegex(ValueError, "resident FULL"): + ProactivePrefetch(self.cache) + with mock.patch.object( + self.cache.cache_controller, "storage_backend", object() + ): + with self.assertRaisesRegex(ValueError, "resident FULL"): + ProactivePrefetch(self.cache) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/mem_cache/test_proactive_prefetch_api.py b/test/registered/unit/mem_cache/test_proactive_prefetch_api.py new file mode 100644 index 000000000000..68b76a17ac4a --- /dev/null +++ b/test/registered/unit/mem_cache/test_proactive_prefetch_api.py @@ -0,0 +1,151 @@ +"""HTTP control transport is independent of generation admission.""" + +import unittest +from types import SimpleNamespace +from unittest import mock + +from fastapi.testclient import TestClient + +from sglang.srt.entrypoints import http_server +from sglang.srt.managers.io_struct import ProactivePrefetchReqOutput +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +class TestProactiveAPI(unittest.TestCase): + def test_exact_prefix_and_cancel_use_only_control_transport(self): + received = [] + + async def control(obj): + received.append(obj) + return ProactivePrefetchReqOutput(success=True, result={"state": "RUNNING"}) + + state = SimpleNamespace( + tokenizer_manager=SimpleNamespace(proactive_prefetch=control) + ) + with mock.patch.object(http_server, "_global_state", state): + client = TestClient(http_server.app) + response = client.post( + "/hicache/prefetch", + json={ + "operation_id": "tool-gap", + "input_ids": [1, 2, 3, 4], + "cache_salt": "session", + }, + ) + self.assertEqual(response.status_code, 200, response.text) + self.assertEqual(received[0].input_ids, [1, 2, 3, 4]) + self.assertEqual(received[0].cache_salt, "session") + response = client.post( + "/hicache/prefetch", + json={"operation_id": "tool-gap", "action": "cancel"}, + ) + self.assertEqual(response.status_code, 200, response.text) + self.assertEqual(received[1].action, "cancel") + self.assertIsNone(received[1].input_ids) + + def test_unsupported_control_is_reported_as_bad_request(self): + async def control(obj): + return ProactivePrefetchReqOutput( + success=False, message="Single worker required" + ) + + state = SimpleNamespace( + tokenizer_manager=SimpleNamespace(proactive_prefetch=control) + ) + with mock.patch.object(http_server, "_global_state", state): + response = TestClient(http_server.app).post( + "/hicache/prefetch", json={"operation_id": "p", "input_ids": [1]} + ) + self.assertEqual(response.status_code, 400) + self.assertFalse(response.json()["success"]) + + def test_multiworker_rejects_before_control_fanout(self): + import asyncio + + from sglang.srt.managers.io_struct import ProactivePrefetchReqInput + from sglang.srt.managers.tokenizer_control_mixin import TokenizerControlMixin + + parallel = SimpleNamespace( + tp_size=2, pp_size=1, dp_size=1, nnodes=1, attn_cp_size=1, attn_dp_size=1 + ) + with mock.patch( + "sglang.srt.managers.tokenizer_control_mixin.get_parallel", + return_value=parallel, + ): + result = asyncio.run( + TokenizerControlMixin.proactive_prefetch( + SimpleNamespace(), + ProactivePrefetchReqInput(operation_id="p", input_ids=[1]), + ) + ) + self.assertFalse(result.success) + + def test_scheduler_control_creates_no_generation_request(self): + import test_prefetch_finite_io as finite + + from sglang.srt.disaggregation.utils import DisaggregationMode + from sglang.srt.managers.io_struct import ProactivePrefetchReqInput + from sglang.srt.managers.scheduler import Scheduler + + fixture = finite.TestFiniteIO( + methodName="test_finite_read_exception_worker_survives" + ) + fixture.setUp() + try: + scheduler = SimpleNamespace( + tree_cache=fixture.cache, + _engine_paused=False, + max_req_input_len=64, + model_config=SimpleNamespace( + vocab_size=128, is_multimodal=False, is_generation=True + ), + disaggregation_mode=DisaggregationMode.NULL, + ) + parallel = SimpleNamespace( + tp_size=1, + pp_size=1, + dp_size=1, + nnodes=1, + attn_cp_size=1, + attn_dp_size=1, + ) + with ( + mock.patch( + "sglang.srt.managers.scheduler.get_parallel", return_value=parallel + ), + mock.patch( + "sglang.srt.managers.scheduler.get_spec", + return_value=SimpleNamespace(speculative_algorithm=None), + ), + mock.patch( + "sglang.srt.managers.scheduler.get_lora", + return_value=SimpleNamespace(enable_lora=False), + ), + mock.patch( + "sglang.srt.managers.scheduler.Req", + side_effect=AssertionError("control must not construct Req"), + ), + ): + result = Scheduler.handle_proactive_prefetch( + scheduler, + ProactivePrefetchReqInput( + operation_id="p", input_ids=list(fixture.tokens) + ), + ) + self.assertTrue(result.success, result.message) + fixture.pump_until(lambda: not fixture.cache.ongoing_prefetch) + scheduler.proactive_prefetch.tick() + self.assertEqual( + scheduler.proactive_prefetch.status("p")["restored_tokens"], 12 + ) + fixture.conservation( + scheduler.proactive_prefetch.records["p"].handle, resident=12 + ) + finally: + fixture.doCleanups() + + +if __name__ == "__main__": + unittest.main()