diff --git a/docs_new/docs/advanced_features/hicache_design.mdx b/docs_new/docs/advanced_features/hicache_design.mdx index 15ab841af..24ca4dac1 100644 --- a/docs_new/docs/advanced_features/hicache_design.mdx +++ b/docs_new/docs/advanced_features/hicache_design.mdx @@ -56,15 +56,19 @@ After local matching, for the parts not found in L1 or L2, the system queries L3 After prefetching stops, the data already fetched is used together with the local data for the prefill computation. -For **timeout** strategy, HiCache introduces two configuration parameters to support fine-grained control over prefetch timeout conditions: +For **timeout** strategy, HiCache introduces three configuration parameters to support fine-grained control over prefetch timeout conditions: -* `prefetch_timeout_base`: the base timeout, representing overhead unrelated to the number of tokens (e.g., scheduling and synchronization). -* `prefetch_timeout_per_ki_token`: the incremental timeout per thousand tokens. +* `prefetch_timeout_base`: the base timeout, representing overhead unrelated to the number of tokens (e.g., scheduling and synchronization). Default: `2` seconds. +* `prefetch_timeout_per_ki_token`: the incremental timeout per thousand tokens. Default: `0.1` seconds per 1024 tokens. +* `prefetch_timeout_max`: the upper bound applied to the linear timeout, preventing very long prompts from waiting unboundedly. Default: `30` seconds. The timeout is computed as: ```python Example -timeout = prefetch_timeout_base + prefetch_timeout_per_ki_token * num_token_to_fetch / 1024 +timeout = min( + prefetch_timeout_max, + prefetch_timeout_base + prefetch_timeout_per_ki_token * num_token_to_fetch / 1024, +) ``` ### Data Write-back diff --git a/docs_new/docs/advanced_features/hicache_storage_runtime_attach_detach.mdx b/docs_new/docs/advanced_features/hicache_storage_runtime_attach_detach.mdx index b245bf520..1f36e0677 100644 --- a/docs_new/docs/advanced_features/hicache_storage_runtime_attach_detach.mdx +++ b/docs_new/docs/advanced_features/hicache_storage_runtime_attach_detach.mdx @@ -109,7 +109,7 @@ Notes: - `hicache_storage_backend_extra_config_json` can include both: - **Backend configuration** (e.g., Mooncake master/metadata/protocol, etc.) - - **Prefetch configuration** (`prefetch_threshold`, `prefetch_timeout_base`, `prefetch_timeout_per_ki_token`, `hicache_storage_pass_prefix_keys`) + - **Prefetch configuration** (`prefetch_threshold`, `prefetch_timeout_base`, `prefetch_timeout_per_ki_token`, `prefetch_timeout_max`, `hicache_storage_pass_prefix_keys`) ### 3.3 Detach (disable) the storage backend diff --git a/docs_new/docs/advanced_features/server_arguments.mdx b/docs_new/docs/advanced_features/server_arguments.mdx index 7a16ce1d8..994a42f5a 100644 --- a/docs_new/docs/advanced_features/server_arguments.mdx +++ b/docs_new/docs/advanced_features/server_arguments.mdx @@ -1683,7 +1683,7 @@ Please consult the documentation below and [server_args.py](https://github.com/s `--hicache-storage-prefetch-policy` Control when prefetching from the storage backend should stop. - `best_effort` + `timeout` `best_effort`, `wait_complete`, `timeout` diff --git a/docs_new/docs/hardware-platforms/ascend-npus/ascend_npu_support_features.mdx b/docs_new/docs/hardware-platforms/ascend-npus/ascend_npu_support_features.mdx index 884d8f1dc..e92622719 100644 --- a/docs_new/docs/hardware-platforms/ascend-npus/ascend_npu_support_features.mdx +++ b/docs_new/docs/hardware-platforms/ascend-npus/ascend_npu_support_features.mdx @@ -1712,7 +1712,7 @@ click [Server Arguments](../../advanced_features/server_arguments). --hicache-storage-<br/>prefetch-policy - best_effort + timeout best_effort,<br/> wait_complete,<br/> timeout Special for GPU diff --git a/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py b/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py index 5ec990081..d95cf3791 100644 --- a/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py +++ b/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py @@ -23,7 +23,12 @@ from sglang.srt.mem_cache.base_prefix_cache import ( MatchPrefixParams, MatchResult, ) -from sglang.srt.mem_cache.hicache_storage import PoolHitPolicy, PoolName, PoolTransfer +from sglang.srt.mem_cache.hicache_storage import ( + PoolHitPolicy, + PoolName, + PoolTransfer, + PrefetchTimeoutConfig, +) from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import ( PrefetchOperation, ) @@ -125,8 +130,7 @@ class HiMambaRadixCache(MambaRadixCache): ( extra_config, prefetch_threshold, - prefetch_timeout_base, - prefetch_timeout_per_ki_token, + prefetch_timeout_config, hicache_storage_pass_prefix_keys, ) = self._parse_storage_backend_extra_config( server_args.hicache_storage_backend_extra_config @@ -149,8 +153,7 @@ class HiMambaRadixCache(MambaRadixCache): self._apply_storage_runtime_config( storage_backend=server_args.hicache_storage_backend, prefetch_threshold=prefetch_threshold, - prefetch_timeout_base=prefetch_timeout_base, - prefetch_timeout_per_ki_token=prefetch_timeout_per_ki_token, + prefetch_timeout_config=prefetch_timeout_config, hicache_storage_pass_prefix_keys=hicache_storage_pass_prefix_keys, enable_storage=self.enable_storage, enable_storage_metrics=self.enable_storage_metrics, @@ -1233,17 +1236,12 @@ class HiMambaRadixCache(MambaRadixCache): *, storage_backend: Optional[str], prefetch_threshold: int, - prefetch_timeout_base: float, - prefetch_timeout_per_ki_token: float, + prefetch_timeout_config: PrefetchTimeoutConfig, hicache_storage_pass_prefix_keys: bool, enable_storage: bool, enable_storage_metrics: bool, extra_metric_labels: Optional[Dict[str, str]], ) -> None: - prefetch_timeout_per_page = ( - self.page_size / 1024 * prefetch_timeout_per_ki_token - ) - storage_metrics_collector = None if enable_storage_metrics: labels = { @@ -1259,8 +1257,7 @@ class HiMambaRadixCache(MambaRadixCache): self.enable_storage = enable_storage self.prefetch_threshold = prefetch_threshold - self.prefetch_timeout_base = prefetch_timeout_base - self.prefetch_timeout_per_page = prefetch_timeout_per_page + self.prefetch_timeout_config = prefetch_timeout_config self.hicache_storage_pass_prefix_keys = hicache_storage_pass_prefix_keys self.enable_storage_metrics = enable_storage_metrics if self.enable_storage_metrics: @@ -1327,8 +1324,7 @@ class HiMambaRadixCache(MambaRadixCache): ( extra_config, prefetch_threshold, - prefetch_timeout_base, - prefetch_timeout_per_ki_token, + prefetch_timeout_config, hicache_storage_pass_prefix_keys, ) = self._parse_storage_backend_extra_config( storage_backend_extra_config_json @@ -1358,8 +1354,7 @@ class HiMambaRadixCache(MambaRadixCache): self._apply_storage_runtime_config( storage_backend=storage_backend, prefetch_threshold=prefetch_threshold, - prefetch_timeout_base=prefetch_timeout_base, - prefetch_timeout_per_ki_token=prefetch_timeout_per_ki_token, + prefetch_timeout_config=prefetch_timeout_config, hicache_storage_pass_prefix_keys=hicache_storage_pass_prefix_keys, enable_storage=True, enable_storage_metrics=self._enable_metrics_flag, @@ -1536,11 +1531,13 @@ class HiMambaRadixCache(MambaRadixCache): logger.error(f"Invalid backend extra config JSON: {e}") raise e + defaults = PrefetchTimeoutConfig() prefetch_threshold = extra_config.pop("prefetch_threshold", 256) - prefetch_timeout_base = extra_config.pop("prefetch_timeout_base", 1) + prefetch_timeout_base = extra_config.pop("prefetch_timeout_base", defaults.base) prefetch_timeout_per_ki_token = extra_config.pop( - "prefetch_timeout_per_ki_token", 0.25 + "prefetch_timeout_per_ki_token", defaults.per_ki_token ) + prefetch_timeout_max = extra_config.pop("prefetch_timeout_max", defaults.max) hicache_storage_pass_prefix_keys = extra_config.pop( "hicache_storage_pass_prefix_keys", False ) @@ -1558,17 +1555,27 @@ class HiMambaRadixCache(MambaRadixCache): f"prefetch_timeout_per_ki_token must be number, got " f"{type(prefetch_timeout_per_ki_token).__name__}" ) + if not isinstance(prefetch_timeout_max, (int, float)): + raise ValueError( + f"prefetch_timeout_max must be number, got " + f"{type(prefetch_timeout_max).__name__}" + ) if not isinstance(hicache_storage_pass_prefix_keys, bool): raise ValueError( "hicache_storage_pass_prefix_keys must be bool, got " f"{type(hicache_storage_pass_prefix_keys).__name__}" ) + prefetch_timeout_config = PrefetchTimeoutConfig( + base=float(prefetch_timeout_base), + per_ki_token=float(prefetch_timeout_per_ki_token), + max=float(prefetch_timeout_max), + ) + return ( extra_config, prefetch_threshold, - float(prefetch_timeout_base), - float(prefetch_timeout_per_ki_token), + prefetch_timeout_config, hicache_storage_pass_prefix_keys, ) @@ -1620,11 +1627,10 @@ class HiMambaRadixCache(MambaRadixCache): ) def _prefetch_timeout_check_linear_func(self, operation: PrefetchOperation): - return ( - time.monotonic() - operation.start_time - > self.prefetch_timeout_base - + len(operation.hash_value) * self.prefetch_timeout_per_page - ) + cfg = self.prefetch_timeout_config + num_tokens = len(operation.hash_value) * self.page_size + timeout = min(cfg.max, cfg.base + cfg.per_ki_token * num_tokens / 1024) + return time.monotonic() - operation.start_time > timeout def can_terminate_prefetch(self, operation: PrefetchOperation): can_terminate = True diff --git a/python/sglang/srt/mem_cache/hicache_storage.py b/python/sglang/srt/mem_cache/hicache_storage.py index 203b0ec60..f301a2c94 100644 --- a/python/sglang/srt/mem_cache/hicache_storage.py +++ b/python/sglang/srt/mem_cache/hicache_storage.py @@ -36,6 +36,15 @@ class HiCacheStorageExtraInfo: extra_info: Optional[dict] = None +@dataclass(frozen=True) +class PrefetchTimeoutConfig: + """Knobs for the linear prefetch-timeout policy used by HiCache.""" + + base: float = 2.0 # seconds, fixed overhead unrelated to token count + per_ki_token: float = 0.1 # seconds per 1024 tokens + max: float = 30.0 # seconds, upper bound for the linear timeout + + class PoolName(str, Enum): """Well-known pool names used as PoolTransfer/PoolEntry identifiers.""" diff --git a/python/sglang/srt/mem_cache/hiradix_cache.py b/python/sglang/srt/mem_cache/hiradix_cache.py index febe60ae1..4a75e965a 100644 --- a/python/sglang/srt/mem_cache/hiradix_cache.py +++ b/python/sglang/srt/mem_cache/hiradix_cache.py @@ -30,6 +30,7 @@ from sglang.srt.mem_cache.hicache_storage import ( PoolHitPolicy, PoolName, PoolTransfer, + PrefetchTimeoutConfig, ) from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import ( HybridCacheController, @@ -111,8 +112,7 @@ class HiRadixCache(RadixCache): ( extra_config, prefetch_threshold, - prefetch_timeout_base, - prefetch_timeout_per_ki_token, + prefetch_timeout_config, hicache_storage_pass_prefix_keys, ) = self._parse_storage_backend_extra_config( server_args.hicache_storage_backend_extra_config @@ -156,8 +156,7 @@ class HiRadixCache(RadixCache): self._apply_storage_runtime_config( storage_backend=server_args.hicache_storage_backend, prefetch_threshold=prefetch_threshold, - prefetch_timeout_base=prefetch_timeout_base, - prefetch_timeout_per_ki_token=prefetch_timeout_per_ki_token, + prefetch_timeout_config=prefetch_timeout_config, hicache_storage_pass_prefix_keys=hicache_storage_pass_prefix_keys, enable_storage=self.enable_storage, enable_storage_metrics=self.enable_storage_metrics, @@ -222,21 +221,15 @@ class HiRadixCache(RadixCache): *, storage_backend: Optional[str], prefetch_threshold: int, - prefetch_timeout_base: float, - prefetch_timeout_per_ki_token: float, + prefetch_timeout_config: PrefetchTimeoutConfig, hicache_storage_pass_prefix_keys: bool, enable_storage: bool, enable_storage_metrics: bool, extra_metric_labels: Optional[Dict[str, str]], ) -> None: - prefetch_timeout_per_page = ( - self.page_size / 1024 * prefetch_timeout_per_ki_token - ) - self.enable_storage = enable_storage self.prefetch_threshold = prefetch_threshold - self.prefetch_timeout_base = prefetch_timeout_base - self.prefetch_timeout_per_page = prefetch_timeout_per_page + self.prefetch_timeout_config = prefetch_timeout_config self.hicache_storage_pass_prefix_keys = hicache_storage_pass_prefix_keys self.enable_storage_metrics = enable_storage_metrics @@ -349,8 +342,7 @@ class HiRadixCache(RadixCache): ( extra_config, prefetch_threshold, - prefetch_timeout_base, - prefetch_timeout_per_ki_token, + prefetch_timeout_config, hicache_storage_pass_prefix_keys, ) = self._parse_storage_backend_extra_config( storage_backend_extra_config_json @@ -379,8 +371,7 @@ class HiRadixCache(RadixCache): self._apply_storage_runtime_config( storage_backend=storage_backend, prefetch_threshold=prefetch_threshold, - prefetch_timeout_base=prefetch_timeout_base, - prefetch_timeout_per_ki_token=prefetch_timeout_per_ki_token, + prefetch_timeout_config=prefetch_timeout_config, hicache_storage_pass_prefix_keys=hicache_storage_pass_prefix_keys, enable_storage=True, enable_storage_metrics=self._enable_metrics_flag, @@ -551,7 +542,7 @@ class HiRadixCache(RadixCache): storage_backend_extra_config: JSON string containing extra configuration Returns: - tuple: (extra_config_dict, prefetch_threshold, prefetch_timeout_base, prefetch_timeout_per_ki_token, hicache_storage_pass_prefix_keys) + tuple: (extra_config_dict, prefetch_threshold, prefetch_timeout_config, hicache_storage_pass_prefix_keys) """ # Parse extra config if provided. Extra config can be a JSON string or a json/toml/yaml file path prefixed with "@". extra_config = {} @@ -583,11 +574,17 @@ class HiRadixCache(RadixCache): logger.error(f"Invalid backend extra config JSON: {e}") raise e + defaults = PrefetchTimeoutConfig() prefetch_threshold = extra_config.pop("prefetch_threshold", 256) # tokens - prefetch_timeout_base = extra_config.pop("prefetch_timeout_base", 1) # seconds + prefetch_timeout_base = extra_config.pop( + "prefetch_timeout_base", defaults.base + ) # seconds prefetch_timeout_per_ki_token = extra_config.pop( - "prefetch_timeout_per_ki_token", 0.25 + "prefetch_timeout_per_ki_token", defaults.per_ki_token ) # seconds per 1024 tokens + prefetch_timeout_max = extra_config.pop( + "prefetch_timeout_max", defaults.max + ) # seconds, upper bound for the linear timeout hicache_storage_pass_prefix_keys = extra_config.pop( "hicache_storage_pass_prefix_keys", False ) @@ -604,17 +601,26 @@ class HiRadixCache(RadixCache): raise ValueError( f"prefetch_timeout_per_ki_token must be number, got {type(prefetch_timeout_per_ki_token).__name__}" ) + if not isinstance(prefetch_timeout_max, (int, float)): + raise ValueError( + f"prefetch_timeout_max must be number, got {type(prefetch_timeout_max).__name__}" + ) if not isinstance(hicache_storage_pass_prefix_keys, bool): raise ValueError( "hicache_storage_pass_prefix_keys must be bool, got " f"{type(hicache_storage_pass_prefix_keys).__name__}" ) + prefetch_timeout_config = PrefetchTimeoutConfig( + base=float(prefetch_timeout_base), + per_ki_token=float(prefetch_timeout_per_ki_token), + max=float(prefetch_timeout_max), + ) + return ( extra_config, prefetch_threshold, - float(prefetch_timeout_base), - float(prefetch_timeout_per_ki_token), + prefetch_timeout_config, hicache_storage_pass_prefix_keys, ) @@ -1105,12 +1111,10 @@ class HiRadixCache(RadixCache): # Timeout is linearly increasing with the number of pages def _prefetch_timeout_check_linear_func(self, operation: PrefetchOperation): - # If hash_value has not been computed in timeout_base seconds, terminate it. - return ( - time.monotonic() - operation.start_time - > self.prefetch_timeout_base - + len(operation.hash_value) * self.prefetch_timeout_per_page - ) + cfg = self.prefetch_timeout_config + num_tokens = len(operation.hash_value) * self.page_size + timeout = min(cfg.max, cfg.base + cfg.per_ki_token * num_tokens / 1024) + return time.monotonic() - operation.start_time > timeout def can_terminate_prefetch(self, operation: PrefetchOperation): can_terminate = True diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 941ca78a1..f8ec870bd 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -642,7 +642,7 @@ class ServerArgs: hicache_io_backend: str = "kernel" hicache_mem_layout: str = "layer_first" hicache_storage_backend: Optional[str] = None - hicache_storage_prefetch_policy: str = "best_effort" + hicache_storage_prefetch_policy: str = "timeout" hicache_storage_backend_extra_config: Optional[str] = None # Hierarchical sparse attention