From 283f505f9cbd8fe83b36aa089a513c01dad39037 Mon Sep 17 00:00:00 2001 From: Kangyan-Zhou Date: Wed, 5 Aug 2026 00:26:19 -0700 Subject: [PATCH 01/47] [Cherry-pick to release/v0.5.17] feat(grpc): add generation request semantics (#32588) (#33668) Signed-off-by: Connor Carpenter Co-authored-by: Connor Carpenter Co-authored-by: ishandhanani <82981111+ishandhanani@users.noreply.github.com> Co-authored-by: Alex Nails --- proto/sglang/runtime/v1/sglang.proto | 26 +- python/sglang/srt/entrypoints/grpc_bridge.py | 30 +- python/sglang/srt/entrypoints/http_server.py | 5 +- python/sglang/srt/managers/io_struct.py | 40 ++- .../sglang/srt/managers/tokenizer_manager.py | 303 +++++++++++++++--- rust/sglang-grpc/src/server.rs | 4 +- rust/sglang-grpc/src/utils/request_utils.rs | 252 ++++++++++++++- .../unit/entrypoints/test_grpc_bridge.py | 115 +++++++ .../unit/managers/test_io_struct.py | 41 +++ .../test_tokenizer_manager_rid_cleanup.py | 252 ++++++++++++++- 10 files changed, 982 insertions(+), 86 deletions(-) create mode 100644 test/registered/unit/entrypoints/test_grpc_bridge.py diff --git a/proto/sglang/runtime/v1/sglang.proto b/proto/sglang/runtime/v1/sglang.proto index 980506957aab..fb8e4f7a1750 100644 --- a/proto/sglang/runtime/v1/sglang.proto +++ b/proto/sglang/runtime/v1/sglang.proto @@ -63,8 +63,24 @@ message SamplingParams { repeated int32 stop_token_ids = 11; optional bool ignore_eos = 12; optional int32 n = 13; - optional string json_schema = 14; - optional string regex = 15; + optional string json_schema = 14 [deprecated = true]; + optional string regex = 15 [deprecated = true]; + optional int64 seed = 16; + optional GuidedDecoding guided_decoding = 17; +} + +message GuidedDecoding { + oneof constraint { + string json_schema = 1; + string regex = 2; + string ebnf = 3; + ChoiceConstraint choice = 4; + string structural_tag = 5; + } +} + +message ChoiceConstraint { + repeated string values = 1; } // ---- Text-based generate (text in, text out) ---- @@ -84,6 +100,9 @@ message TextGenerateRequest { map trace_headers = 12; optional string session_id = 13; optional DisaggregatedParams disaggregated_params = 14; + optional int32 priority = 15; + optional bool require_reasoning = 16; + optional uint32 max_thinking_tokens = 17; } message TextGenerateResponse { @@ -108,6 +127,9 @@ message GenerateRequest { map trace_headers = 11; optional string session_id = 12; optional DisaggregatedParams disaggregated_params = 13; + optional int32 priority = 14; + optional bool require_reasoning = 15; + optional uint32 max_thinking_tokens = 16; } message GenerateResponse { diff --git a/python/sglang/srt/entrypoints/grpc_bridge.py b/python/sglang/srt/entrypoints/grpc_bridge.py index dfaf2840fd76..176503a50760 100644 --- a/python/sglang/srt/entrypoints/grpc_bridge.py +++ b/python/sglang/srt/entrypoints/grpc_bridge.py @@ -293,14 +293,23 @@ def submit_request( async def _run_generate(self, obj, chunk_callback, stream: bool, request): ready_event = None + gen = None try: - ready_event = self._install_on_ready(chunk_callback) if stream else None + ready_event = self._install_on_ready(chunk_callback) gen = self.tokenizer_manager.generate_request(obj, request=request) if stream: + completed_choices = set() + expected_choices = obj.batch_size * obj.parallel_sample_num async for chunk in gen: - finished = ( + choice_finished = ( chunk.get("meta_info", {}).get("finish_reason") is not None ) + if choice_finished: + choice_id = chunk.get( + "index", chunk.get("meta_info", {}).get("id") + ) + completed_choices.add(choice_id) + finished = len(completed_choices) >= expected_choices keep_going = await self._send_with_backpressure( chunk_callback, ready_event, @@ -314,15 +323,26 @@ async def _run_generate(self, obj, chunk_callback, stream: bool, request): self._safe_callback(chunk_callback, {}, finished=True) else: result = await gen.__anext__() - self._safe_callback(chunk_callback, result, finished=True) + chunks = result if isinstance(result, list) else [result] + for index, chunk in enumerate(chunks): + keep_going = await self._send_with_backpressure( + chunk_callback, + ready_event, + chunk, + finished=index == len(chunks) - 1, + timeout_abort_rid=obj.rid, + ) + if not keep_going: + return except StopAsyncIteration: self._safe_callback(chunk_callback, {}, finished=True) except Exception as e: logger.error("gRPC generate error for rid=%s: %s", obj.rid, e) self._send_native_error(chunk_callback, str(e)) finally: - if stream: - self._uninstall_on_ready(chunk_callback) + if gen is not None: + await gen.aclose() + self._uninstall_on_ready(chunk_callback) async def _run_embed(self, obj, chunk_callback, request): try: diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index 289d00d6bd29..ac38b4668c81 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -895,7 +895,7 @@ async def stream_results() -> AsyncIterator[bytes]: "error": { "message": str(e), "type": "invalid_request_error", - "code": 400, + "code": getattr(e, "status_code", 400), "retryable": False, } } @@ -2048,7 +2048,8 @@ async def vertex_generate( def _create_error_response(e): return ORJSONResponse( - {"error": {"message": str(e)}}, status_code=HTTPStatus.BAD_REQUEST + {"error": {"message": str(e)}}, + status_code=getattr(e, "status_code", HTTPStatus.BAD_REQUEST), ) diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index 4e7a54b39c98..928348eb52ad 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -158,8 +158,9 @@ class SessionParams(msgspec.Struct, kw_only=True, array_like=True): @dataclass class GenerateReqInput: - # Request ID(s). If omitted, generated during normalization. For batch - # requests, a string is expanded to per-item IDs using it as a prefix. + # Logical request ID(s). If omitted, generated during normalization. For + # batch requests, a string is expanded to one ID per original batch item. + # Parallel-sampling child IDs are internal to TokenizerManager. rid: Optional[Union[str, List[str]]] = field(default=None, kw_only=True) # Stable identity shared by requests in the same session. Unlike # session_params, this does not alter or reconstruct the prompt. @@ -276,6 +277,9 @@ class GenerateReqInput: background: bool = False # Require reasoning for the request (hybrid reasoning model only) require_reasoning: bool = False + # Per-request thinking budget. Requires strict thinking so the runtime can + # enforce the limit rather than silently treating it as metadata. + max_thinking_tokens: Optional[int] = None # Priority for the request priority: Optional[int] = None @@ -319,12 +323,17 @@ class GenerateReqInput: # Batch-level: List[List[int]] (one per request). After __getitem__: List[int]. multi_item_delimiter_indices: Optional[Union[List[List[int]], List[int]]] = None - def regenerate_rid(self): + def regenerate_rid(self, prefix: Optional[str] = None): """Generate a new request ID and return it.""" + + def new_rid() -> str: + suffix = uuid.uuid4().hex + return f"{prefix}_{suffix}" if prefix is not None else suffix + if isinstance(self.rid, list): - self.rid = [uuid.uuid4().hex for _ in range(len(self.rid))] + self.rid = [new_rid() for _ in range(len(self.rid))] else: - self.rid = uuid.uuid4().hex + self.rid = new_rid() return self.rid def _validate_rid_uniqueness(self): @@ -480,7 +489,7 @@ def _normalize_batch_inputs(self): # Expand input based on type self._expand_inputs(num) - self._normalize_rid(num) + self._normalize_rid() self._normalize_lora_paths(num) self._normalize_image_data(num) self._normalize_video_data(num) @@ -590,16 +599,16 @@ def _normalize_sampling_params(self, num): else: # Already a list self.sampling_params = self.sampling_params * self.parallel_sample_num - def _normalize_rid(self, num): - """Normalize request IDs for batch processing.""" + def _normalize_rid(self): + """Normalize one logical request ID per original batch item.""" if self.rid is None: - self.rid = [uuid.uuid4().hex for _ in range(num)] + self.rid = [uuid.uuid4().hex for _ in range(self.batch_size)] elif isinstance(self.rid, str): - new_rids = [f"{self.rid}_{i}" for i in range(num)] - self.rid = new_rids + if self.batch_size == 1: + self.rid = [self.rid] + else: + self.rid = [f"{self.rid}_{i}" for i in range(self.batch_size)] elif isinstance(self.rid, list): - # Note: the length of rid shall be the same as the batch_size, - # as the rid would be expanded for parallel sampling in tokenizer_manager if len(self.rid) != self.batch_size: raise ValueError( "The specified rids length mismatch with the batch_size for batch processing." @@ -751,8 +760,9 @@ def __getitem__(self, i): cache = self.__dict__.setdefault("_sub_obj_cache", {}) if i in cache: return cache[i] + logical_index = i % self.batch_size sub = GenerateReqInput( - rid=self.rid[i], + rid=self.rid[logical_index], session_id=self.session_id, text=self.text[i] if self.text is not None else None, input_ids=self.input_ids[i] if self.input_ids is not None else None, @@ -813,6 +823,8 @@ def __getitem__(self, i): disagg_prefill_dp_rank=self.disagg_prefill_dp_rank, conversation_id=self.conversation_id, http_worker_ipc=self.http_worker_ipc, + require_reasoning=self.require_reasoning, + max_thinking_tokens=self.max_thinking_tokens, priority=self.priority, extra_key=self.extra_key[i] if self.extra_key is not None else None, no_logs=self.no_logs, diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 3658aa209f30..80351f330430 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -195,6 +195,10 @@ def _ragged_verify_cap_accept() -> bool: ) +class RequestAbortedError(ValueError): + status_code = 499 + + @dataclasses.dataclass class ReqState: """Store the state a request.""" @@ -206,6 +210,9 @@ class ReqState: # For performance metrics time_stats: APIServerReqTimeStats + abort_requested: bool = False + lifecycle_id: object = dataclasses.field(default_factory=object) + dispatched: bool = False last_completion_tokens: int = 1 ttft_observed: bool = False @@ -545,6 +552,10 @@ async def _async_dispatch_to_scheduler(self, obj: Any) -> None: def init_running_status(self): # Request states self.rid_to_state: Dict[str, ReqState] = {} + # Parallel sampling keeps one caller-visible logical RID per original + # prompt while the scheduler operates on separate prefix/sample RIDs. + self.logical_rid_to_child_rids: Dict[str, set[str]] = {} + self.child_rid_to_logical_rid: Dict[str, str] = {} self.event_loop = None self.asyncio_tasks = set() @@ -740,6 +751,15 @@ async def generate_request( # Normalize the request obj.normalize_batch_and_arguments() self._set_default_priority(obj) + if ( + isinstance(obj, GenerateReqInput) + and obj.max_thinking_tokens is not None + and not self.server_args.enable_strict_thinking + ): + raise ValueError( + "max_thinking_tokens requires the server to be launched with " + "--enable-strict-thinking" + ) if isinstance(obj, GenerateReqInput) and obj.routed_dp_rank is not None: dp_size = self.elastic_worker_count @@ -752,7 +772,7 @@ async def generate_request( f"routed_dp_rank={obj.routed_dp_rank} out of range [0, {dp_size})" ) - self._init_req_state(obj, request) + request_lifecycles = self._init_req_state(obj, request) try: if self.server_args.language_only: self._handle_epd_disaggregation_encode_request(obj) @@ -762,13 +782,16 @@ async def generate_request( async with self.is_pause_cond: await self.is_pause_cond.wait_for(lambda: not self.is_pause) + self._raise_if_logical_request_aborted(obj) async with self.model_update_lock.reader_lock: await self._validate_and_resolve_lora(obj) + self._raise_if_logical_request_aborted(obj) # Tokenize the request and send it to the scheduler if obj.is_single: tokenized_obj = await self._tokenize_one_request(obj) + self._raise_if_logical_rid_aborted(obj.rid) state = self.rid_to_state[obj.rid] if obj.return_prompt_token_ids: state.prompt_token_ids = list(tokenized_obj.input_ids) @@ -778,7 +801,7 @@ async def generate_request( else: async for response in self._handle_batch_request(obj, request): yield response - except Exception: + except BaseException: # _init_req_state created a rid_to_state entry per (sub-)request up # front. The normal remover is the scheduler-response path # (_handle_batch_output), so a failure *before* a request reaches the @@ -786,7 +809,7 @@ async def generate_request( # request -- would otherwise leak those entries forever. Drop any that # are still pending; entries already removed on the normal completion # path are left untouched (pop is a no-op). - self._discard_pending_req_states(obj) + self._discard_pending_req_states(obj, request_lifecycles) raise def _detect_input_format( @@ -1308,6 +1331,11 @@ def _create_tokenized_object( sampling_kwargs = {**self.preferred_sampling_params, **obj.sampling_params} else: sampling_kwargs = obj.sampling_params + if isinstance(obj, GenerateReqInput) and obj.max_thinking_tokens is not None: + sampling_kwargs = dict(sampling_kwargs) + custom_params = dict(sampling_kwargs.get("custom_params") or {}) + custom_params["thinking_budget"] = obj.max_thinking_tokens + sampling_kwargs["custom_params"] = custom_params sampling_params = self.sampling_params_class(**sampling_kwargs) sampling_params.normalize(self.tokenizer) sampling_params.verify(self.model_config.vocab_size) @@ -1518,6 +1546,9 @@ def _send_one_request( time_stats = tokenized_obj.time_stats tokenized_obj.wrap_pickle_fields() self._dispatch_to_scheduler(tokenized_obj) + state = self.rid_to_state.get(tokenized_obj.rid) + if state is not None: + state.dispatched = True tokenized_obj.time_stats = time_stats tokenized_obj.time_stats.set_api_server_dispatch_finish_time() @@ -1539,6 +1570,10 @@ def _send_batch_request( batch_req = BatchTokenizedEmbeddingReqInput(batch=tokenized_objs) self._dispatch_to_scheduler(batch_req) + for tokenized_obj in tokenized_objs: + state = self.rid_to_state.get(tokenized_obj.rid) + if state is not None: + state.dispatched = True for tokenized_obj, time_stat in zip(tokenized_objs, time_stats): tokenized_obj.time_stats = time_stat set_time_batch(tokenized_objs, "set_api_server_dispatch_finish_time") @@ -1610,7 +1645,7 @@ async def _handle_abort_finish_reason( # Delete the key to prevent resending abort request to the scheduler and # to ensure aborted request state is cleaned up. if state.obj.rid in self.rid_to_state: - del self.rid_to_state[state.obj.rid] + self._remove_req_state(state.obj.rid) # Mark ongoing LoRA request as finished. if self.enable_lora and state.obj.lora_path: @@ -1744,6 +1779,7 @@ async def _handle_batch_request( if getattr(obj, "parallel_sample_num", 1) == 1: if self._should_use_batch_tokenization(batch_size, obj): tokenized_objs = await self._batch_tokenize_and_process(batch_size, obj) + self._raise_if_logical_request_aborted(obj) self._send_batch_request(tokenized_objs) # Set up generators for each request in the batch @@ -1766,6 +1802,7 @@ async def _handle_batch_request( for i in range(batch_size): tmp_obj = obj[i] tokenized_obj = await self._tokenize_one_request(tmp_obj) + self._raise_if_logical_rid_aborted(tmp_obj.rid) state = self.rid_to_state[tmp_obj.rid] if tmp_obj.return_prompt_token_ids: state.prompt_token_ids = list(tokenized_obj.input_ids) @@ -1786,9 +1823,12 @@ async def _handle_batch_request( tokenized_objs = await asyncio.gather( *(self._tokenize_one_request(obj) for obj in objs) ) + self._raise_if_logical_request_aborted(obj) # Cache the common prefix for parallel sampling for i in range(batch_size): + logical_rid = objs[i].rid + self._raise_if_logical_rid_aborted(logical_rid) tmp_obj = copy.copy(objs[i]) tokenized_obj = copy.copy(tokenized_objs[i]) # Ensure independent mm_items so wrap_shm_features won't mutate the original @@ -1797,17 +1837,20 @@ async def _handle_batch_request( tokenized_obj.mm_inputs.mm_items = [ copy.copy(item) for item in tokenized_obj.mm_inputs.mm_items ] - tokenized_obj.rid = tmp_obj.regenerate_rid() + tokenized_obj.rid = tmp_obj.regenerate_rid(prefix=logical_rid) tokenized_obj.sampling_params = copy.copy(tokenized_obj.sampling_params) tokenized_obj.sampling_params.max_new_tokens = 0 tokenized_obj.stream = False - self._init_req_state(tmp_obj) + self._init_child_req_state(logical_rid, tmp_obj) self._send_one_request(tokenized_obj) await self._wait_one_response(tmp_obj, request).__anext__() + self._raise_if_logical_rid_aborted(logical_rid) # Expand requests, assign new rids for them, and send them for i in range(batch_size): + logical_rid = objs[i].rid for _ in range(obj.parallel_sample_num): + self._raise_if_logical_rid_aborted(logical_rid) tmp_obj = copy.copy(objs[i]) tokenized_obj = copy.copy(tokenized_objs[i]) # Ensure independent mm_items so wrap_shm_features won't mutate the original @@ -1816,8 +1859,8 @@ async def _handle_batch_request( tokenized_obj.mm_inputs.mm_items = [ copy.copy(item) for item in tokenized_obj.mm_inputs.mm_items ] - tokenized_obj.rid = tmp_obj.regenerate_rid() - self._init_req_state(tmp_obj) + tokenized_obj.rid = tmp_obj.regenerate_rid(prefix=logical_rid) + self._init_child_req_state(logical_rid, tmp_obj) state = self.rid_to_state[tmp_obj.rid] tokenized_obj.time_stats = state.time_stats if tmp_obj.return_prompt_token_ids: @@ -1826,17 +1869,38 @@ async def _handle_batch_request( generators.append(self._wait_one_response(tmp_obj, request)) rids.append(tmp_obj.rid) - self.rid_to_state[objs[i].rid].time_stats.set_finished_time() - del self.rid_to_state[objs[i].rid] + parent_state = self.rid_to_state.get(logical_rid) + if parent_state is not None: + parent_state.time_stats.set_finished_time() + self._remove_req_state(logical_rid) # Wait for all requests is_stream = hasattr(obj, "stream") and obj.stream if not is_stream: - outputs = await asyncio.gather(*(gen.__anext__() for gen in generators)) + outputs = await self._collect_batch_responses(generators) yield outputs else: - rid_to_index = {rid: i for i, rid in enumerate(rids)} - task_map = {asyncio.create_task(gen.__anext__()): gen for gen in generators} + async for response in self._stream_batch_responses(generators, rids): + yield response + + async def _collect_batch_responses(self, generators): + tasks = [asyncio.create_task(gen.__anext__()) for gen in generators] + try: + return await asyncio.gather(*tasks) + finally: + for task in tasks: + if not task.done(): + task.cancel() + await asyncio.gather(*tasks, return_exceptions=True) + await asyncio.gather( + *(gen.aclose() for gen in generators), + return_exceptions=True, + ) + + async def _stream_batch_responses(self, generators, rids): + rid_to_index = {rid: i for i, rid in enumerate(rids)} + task_map = {asyncio.create_task(gen.__anext__()): gen for gen in generators} + try: while task_map: done, _ = await asyncio.wait( task_map.keys(), return_when=asyncio.FIRST_COMPLETED @@ -1852,20 +1916,55 @@ async def _handle_batch_request( task_map[new_task] = gen except StopAsyncIteration: pass + finally: + pending_tasks = list(task_map) + for task in pending_tasks: + task.cancel() + if pending_tasks: + await asyncio.gather(*pending_tasks, return_exceptions=True) + await asyncio.gather( + *(gen.aclose() for gen in generators), + return_exceptions=True, + ) def abort_request(self, rid: str = "", abort_all: bool = False): # Empty rid would startswith-match every request on the scheduler. if not abort_all and not rid: logger.warning("Ignore abort_request with empty rid and abort_all=False") return - if ( - not abort_all - and self.server_args.tokenizer_worker_num == 1 - and rid not in self.rid_to_state - ): + if abort_all: + for state_rid, state in self.rid_to_state.items(): + if state_rid not in self.child_rid_to_logical_rid: + state.abort_requested = True + target_rids = (rid,) + elif rid in self.child_rid_to_logical_rid: + # Preserve direct child aborts for internal callers. + target_rids = (rid,) + elif rid in self.rid_to_state: + state = self.rid_to_state[rid] + state.abort_requested = True + parallel_sample_num = getattr(state.obj, "parallel_sample_num", None) + if parallel_sample_num is None: + sampling_params = getattr(state.obj, "sampling_params", None) + parallel_sample_num = ( + sampling_params.get("n", 1) + if isinstance(sampling_params, dict) + else 1 + ) + if parallel_sample_num > 1: + # Snapshot because scheduler abort echoes remove child ownership. + target_rids = tuple(sorted(self.logical_rid_to_child_rids.get(rid, ()))) + else: + target_rids = (rid,) + elif child_rids := self.logical_rid_to_child_rids.get(rid): + target_rids = tuple(sorted(child_rids)) + elif self.server_args.tokenizer_worker_num == 1: return - req = AbortReq(rid=rid, abort_all=abort_all) - self._dispatch_to_scheduler(req) + else: + target_rids = (rid,) + + for target_rid in target_rids: + self._dispatch_to_scheduler(AbortReq(rid=target_rid, abort_all=abort_all)) if self.enable_metrics: # TODO: also use custom_labels from the request self.metrics_collector.observe_one_aborted_request( @@ -2352,7 +2451,7 @@ async def _handle_batch_output( ) ) - del self.rid_to_state[rid] + self._remove_req_state(rid) # Mark ongoing LoRA request as finished. if self.enable_lora and state.obj.lora_path: @@ -3088,7 +3187,7 @@ def _handle_abort_req(self, recv_obj: AbortReq): "output_ids": output_ids, "meta_info": meta_info, } - del self.rid_to_state[recv_obj.rid] + self._remove_req_state(recv_obj.rid) state.out_list.append(out) state.event.set() @@ -3244,11 +3343,80 @@ async def _resolve_lora_path(self, obj: Union[GenerateReqInput, EmbeddingReqInpu obj.lora_id[i] if isinstance(obj.lora_id, list) else obj.lora_id ) + @staticmethod + def _logical_rids(obj) -> List[str]: + if not hasattr(obj, "is_single") or obj.is_single: + return [obj.rid] + return list(obj.rid) + + def _register_child_rid(self, logical_rid: str, child_rid: str) -> None: + if child_rid == logical_rid: + raise ValueError( + "Parallel-sampling child RID must differ from its logical RID" + ) + owner = self.child_rid_to_logical_rid.get(child_rid) + if owner is not None and owner != logical_rid: + raise ValueError( + f"Request ID {child_rid} is already owned by logical request {owner}" + ) + self.child_rid_to_logical_rid[child_rid] = logical_rid + self.logical_rid_to_child_rids.setdefault(logical_rid, set()).add(child_rid) + + def _init_child_req_state( + self, + logical_rid: str, + obj: Union[GenerateReqInput, EmbeddingReqInput], + request: Optional[fastapi.Request] = None, + ) -> None: + self._raise_if_logical_rid_aborted(logical_rid) + logical_state = self.rid_to_state[logical_rid] + self._init_req_state( + obj, + request, + lifecycle_id=logical_state.lifecycle_id, + ) + try: + self._register_child_rid(logical_rid, obj.rid) + except BaseException: + self._remove_req_state(obj.rid) + raise + + def _remove_req_state( + self, + rid: str, + lifecycle_id: Optional[object] = None, + ) -> Optional[ReqState]: + """Remove a request state and its parallel-sampling ownership.""" + state = self.rid_to_state.get(rid) + if state is None or ( + lifecycle_id is not None and state.lifecycle_id is not lifecycle_id + ): + return None + self.rid_to_state.pop(rid) + logical_rid = self.child_rid_to_logical_rid.pop(rid, None) + if logical_rid is not None: + children = self.logical_rid_to_child_rids.get(logical_rid) + if children is not None: + children.discard(rid) + if not children: + self.logical_rid_to_child_rids.pop(logical_rid, None) + return state + + def _raise_if_logical_rid_aborted(self, logical_rid: str) -> None: + state = self.rid_to_state.get(logical_rid) + if state is None or state.abort_requested: + raise RequestAbortedError(f"Request {logical_rid} was aborted") + + def _raise_if_logical_request_aborted(self, obj) -> None: + for logical_rid in self._logical_rids(obj): + self._raise_if_logical_rid_aborted(logical_rid) + def _init_req_state( self, obj: Union[GenerateReqInput, EmbeddingReqInput], request: Optional[fastapi.Request] = None, - ): + lifecycle_id: Optional[object] = None, + ) -> Dict[str, object]: created_time = obj.received_time external_trace_header = None @@ -3279,29 +3447,90 @@ def _init_req_state( for i in range(len(obj.rid)) ] - for rid, sub_obj, bootstrap_room in items: - if rid in self.rid_to_state: + rids = [rid for rid, _, _ in items] + seen_rids = set() + for rid in rids: + if rid in seen_rids: + raise ValueError(f"Duplicate request ID detected: {rid}") + seen_rids.add(rid) + if ( + rid in self.rid_to_state + or rid in self.logical_rid_to_child_rids + or rid in self.child_rid_to_logical_rid + ): raise ValueError(f"Duplicate request ID detected: {rid}") + + # Mutate only after every RID passes duplicate validation so a rejected + # batch cannot leave a partial rid_to_state insertion behind. + lifecycle_ids = {} + for rid, sub_obj, bootstrap_room in items: time_stats = APIServerReqTimeStats(disagg_mode=self.disaggregation_mode) - state = ReqState([], False, asyncio.Event(), sub_obj, time_stats) + state = ReqState( + [], + False, + asyncio.Event(), + sub_obj, + time_stats, + lifecycle_id=lifecycle_id if lifecycle_id is not None else object(), + ) self.rid_to_state[rid] = state + lifecycle_ids[rid] = state.lifecycle_id if self.enable_trace: time_stats.init_trace_ctx(rid, bootstrap_room, external_trace_header) time_stats.set_created_time(created_time) + return lifecycle_ids - def _discard_pending_req_states(self, obj): - """Drop rid_to_state entries created by _init_req_state for *obj*. + def _discard_pending_req_states( + self, + obj, + lifecycle_ids: Optional[Dict[str, object]] = None, + ): + """Drop all logical and child state owned by *obj*. - Safe to call after a partial/failed dispatch: only entries still present - are removed, and the scheduler-response path looks up state with - ``.get(...)`` so a later output for a discarded rid is ignored, not fatal. + Safe to call after a partial/failed dispatch: only requests known to have + reached the scheduler are aborted, all owned state is removed, and a later + output for a discarded RID is ignored by the scheduler-response path. """ - if not hasattr(obj, "is_single") or obj.is_single: - rids = [obj.rid] - else: - rids = obj.rid - for rid in rids: - self.rid_to_state.pop(rid, None) + if lifecycle_ids is None: + lifecycle_ids = { + logical_rid: state.lifecycle_id + for logical_rid in self._logical_rids(obj) + if (state := self.rid_to_state.get(logical_rid)) is not None + } + for logical_rid in self._logical_rids(obj): + lifecycle_id = lifecycle_ids.get(logical_rid) + if lifecycle_id is None: + continue + child_rids = tuple( + child_rid + for child_rid in self.logical_rid_to_child_rids.get(logical_rid, ()) + if ( + (state := self.rid_to_state.get(child_rid)) is not None + and state.lifecycle_id is lifecycle_id + ) + ) + logical_state = self.rid_to_state.get(logical_rid) + owns_logical_state = ( + logical_state is not None and logical_state.lifecycle_id is lifecycle_id + ) + target_rids = tuple( + rid for rid in child_rids if self.rid_to_state[rid].dispatched + ) + if not child_rids and owns_logical_state and logical_state.dispatched: + target_rids = (logical_rid,) + for target_rid in target_rids: + try: + self._dispatch_to_scheduler( + AbortReq(rid=target_rid, abort_all=False) + ) + except Exception: + logger.exception( + "Failed to abort request rid=%s", + target_rid, + ) + for child_rid in child_rids: + self._remove_req_state(child_rid, lifecycle_id) + self._remove_req_state(logical_rid, lifecycle_id) def _should_dispatch_to_encoder( self, obj: Union[GenerateReqInput, EmbeddingReqInput] diff --git a/rust/sglang-grpc/src/server.rs b/rust/sglang-grpc/src/server.rs index 88ce594e963d..c05736ce4bed 100644 --- a/rust/sglang-grpc/src/server.rs +++ b/rust/sglang-grpc/src/server.rs @@ -229,7 +229,7 @@ impl proto::sglang_service_server::SglangService for SglangServiceImpl { .rid .clone() .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()); - let req_dict = build_text_generate_dict(&rid, &req); + let req_dict = build_text_generate_dict(&rid, &req).map_err(Status::invalid_argument)?; let mut receiver = self .bridge @@ -298,7 +298,7 @@ impl proto::sglang_service_server::SglangService for SglangServiceImpl { .rid .clone() .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()); - let req_dict = build_generate_dict(&rid, &req); + let req_dict = build_generate_dict(&rid, &req).map_err(Status::invalid_argument)?; let mut receiver = self .bridge diff --git a/rust/sglang-grpc/src/utils/request_utils.rs b/rust/sglang-grpc/src/utils/request_utils.rs index 4685a82093bc..20a192f14151 100644 --- a/rust/sglang-grpc/src/utils/request_utils.rs +++ b/rust/sglang-grpc/src/utils/request_utils.rs @@ -2,8 +2,25 @@ use std::collections::HashMap; use crate::proto; +fn regex_escape_literal(value: &str) -> String { + let mut escaped = String::with_capacity(value.len()); + for character in value.chars() { + if matches!( + character, + '.' | '+' | '*' | '?' | '^' | '$' | '(' | ')' | '[' | ']' | '{' | '}' | '|' | '\\' + ) { + escaped.push('\\'); + } + escaped.push(character); + } + escaped +} + /// Convert proto SamplingParams to a serde_json map (used as Python dict via PyO3). -fn sampling_params_to_map(params: &Option) -> serde_json::Value { +#[allow(deprecated)] +fn sampling_params_to_map( + params: &Option, +) -> Result { match params { Some(p) => { let mut map = serde_json::Map::new(); @@ -46,15 +63,90 @@ fn sampling_params_to_map(params: &Option) -> serde_json: if let Some(v) = p.n { map.insert("n".into(), serde_json::json!(v)); } - if let Some(ref v) = p.json_schema { - map.insert("json_schema".into(), serde_json::json!(v)); + if let Some(v) = p.seed { + map.insert("sampling_seed".into(), serde_json::json!(v)); } - if let Some(ref v) = p.regex { - map.insert("regex".into(), serde_json::json!(v)); + if p.guided_decoding.is_some() && (p.json_schema.is_some() || p.regex.is_some()) { + return Err( + "legacy json_schema/regex cannot be combined with guided_decoding".into(), + ); } - serde_json::Value::Object(map) + if let Some(guided) = p.guided_decoding.as_ref() { + use proto::guided_decoding::Constraint; + match guided.constraint.as_ref() { + Some(Constraint::JsonSchema(value)) if !value.is_empty() => { + map.insert("json_schema".into(), serde_json::json!(value)); + } + Some(Constraint::Regex(value)) if !value.is_empty() => { + map.insert("regex".into(), serde_json::json!(value)); + } + Some(Constraint::Ebnf(value)) if !value.is_empty() => { + map.insert("ebnf".into(), serde_json::json!(value)); + } + Some(Constraint::Choice(choice)) + if !choice.values.is_empty() + && choice.values.iter().all(|value| !value.is_empty()) => + { + let alternatives = choice + .values + .iter() + .map(|value| regex_escape_literal(value)) + .collect::>() + .join("|"); + map.insert( + "regex".into(), + serde_json::json!(format!("(?:{alternatives})")), + ); + } + Some(Constraint::StructuralTag(value)) if !value.is_empty() => { + map.insert("structural_tag".into(), serde_json::json!(value)); + } + Some(Constraint::Choice(_)) => { + return Err("guided choice must contain only non-empty values".into()); + } + Some(_) => return Err("guided decoding constraint must not be empty".into()), + None => return Err("guided decoding constraint must be specified".into()), + } + } else { + if let Some(value) = p.json_schema.as_ref() { + if value.is_empty() { + return Err("legacy json_schema must not be empty".into()); + } + map.insert("json_schema".into(), serde_json::json!(value)); + } + if let Some(value) = p.regex.as_ref() { + if value.is_empty() { + return Err("legacy regex must not be empty".into()); + } + map.insert("regex".into(), serde_json::json!(value)); + } + } + Ok(serde_json::Value::Object(map)) } - None => serde_json::Value::Object(serde_json::Map::new()), + None => Ok(serde_json::Value::Object(serde_json::Map::new())), + } +} + +fn insert_generation_controls( + d: &mut HashMap, + priority: Option, + require_reasoning: Option, + max_thinking_tokens: Option, +) { + if let Some(priority) = priority { + d.insert("priority".into(), serde_json::json!(priority)); + } + if let Some(require_reasoning) = require_reasoning { + d.insert( + "require_reasoning".into(), + serde_json::json!(require_reasoning), + ); + } + if let Some(max_thinking_tokens) = max_thinking_tokens { + d.insert( + "max_thinking_tokens".into(), + serde_json::json!(max_thinking_tokens), + ); } } @@ -111,13 +203,13 @@ pub(crate) fn extract_model_path(json_info: &str) -> String { pub(crate) fn build_text_generate_dict( rid: &str, req: &proto::TextGenerateRequest, -) -> HashMap { +) -> Result, String> { let mut d = HashMap::new(); d.insert("rid".into(), serde_json::json!(rid)); d.insert("text".into(), serde_json::json!(req.text)); d.insert( "sampling_params".into(), - sampling_params_to_map(&req.sampling_params), + sampling_params_to_map(&req.sampling_params)?, ); d.insert( "stream".into(), @@ -151,25 +243,31 @@ pub(crate) fn build_text_generate_dict( if let Some(ref session_id) = req.session_id { d.insert("session_id".into(), serde_json::json!(session_id)); } + insert_generation_controls( + &mut d, + req.priority, + req.require_reasoning, + req.max_thinking_tokens, + ); insert_disaggregated_params(&mut d, &req.disaggregated_params); if let Some(trace) = trace_headers_to_json(&req.trace_headers) { d.insert("external_trace_header".into(), trace); } d.insert("received_time".into(), serde_json::json!(now_timestamp())); - d + Ok(d) } /// Build a request dict for GenerateReqInput from proto GenerateRequest (tokenized). pub(crate) fn build_generate_dict( rid: &str, req: &proto::GenerateRequest, -) -> HashMap { +) -> Result, String> { let mut d = HashMap::new(); d.insert("rid".into(), serde_json::json!(rid)); d.insert("input_ids".into(), serde_json::json!(req.input_ids)); d.insert( "sampling_params".into(), - sampling_params_to_map(&req.sampling_params), + sampling_params_to_map(&req.sampling_params)?, ); d.insert( "stream".into(), @@ -199,12 +297,18 @@ pub(crate) fn build_generate_dict( if let Some(ref session_id) = req.session_id { d.insert("session_id".into(), serde_json::json!(session_id)); } + insert_generation_controls( + &mut d, + req.priority, + req.require_reasoning, + req.max_thinking_tokens, + ); insert_disaggregated_params(&mut d, &req.disaggregated_params); if let Some(trace) = trace_headers_to_json(&req.trace_headers) { d.insert("external_trace_header".into(), trace); } d.insert("received_time".into(), serde_json::json!(now_timestamp())); - d + Ok(d) } /// Build a request dict for EmbeddingReqInput from proto TextEmbedRequest. @@ -267,6 +371,7 @@ pub(crate) fn build_classify_dict( } #[cfg(test)] +#[allow(deprecated)] mod tests { use super::*; @@ -283,11 +388,15 @@ mod tests { }; assert_eq!( - build_text_generate_dict("request-1", &text_req).get("session_id"), + build_text_generate_dict("request-1", &text_req) + .unwrap() + .get("session_id"), Some(&serde_json::json!("session-1")) ); assert_eq!( - build_generate_dict("request-2", &token_req).get("session_id"), + build_generate_dict("request-2", &token_req) + .unwrap() + .get("session_id"), Some(&serde_json::json!("session-1")) ); } @@ -312,6 +421,7 @@ mod tests { build_text_generate_dict("request-1", &text_req), build_generate_dict("request-2", &token_req), ] { + let request = request.unwrap(); assert_eq!( request.get("bootstrap_host"), Some(&serde_json::json!("10.0.0.1")) @@ -330,8 +440,9 @@ mod tests { #[test] fn generate_dicts_omit_disaggregated_params_when_absent() { let text_request = - build_text_generate_dict("request-1", &proto::TextGenerateRequest::default()); - let token_request = build_generate_dict("request-2", &proto::GenerateRequest::default()); + build_text_generate_dict("request-1", &proto::TextGenerateRequest::default()).unwrap(); + let token_request = + build_generate_dict("request-2", &proto::GenerateRequest::default()).unwrap(); for request in [text_request, token_request] { assert!(!request.contains_key("bootstrap_host")); @@ -339,4 +450,111 @@ mod tests { assert!(!request.contains_key("bootstrap_room")); } } + + #[test] + fn generate_dicts_preserve_optional_generation_controls() { + let sampling_params = proto::SamplingParams { + seed: Some(42), + ..Default::default() + }; + let text_request = proto::TextGenerateRequest { + sampling_params: Some(sampling_params.clone()), + priority: Some(3), + require_reasoning: Some(false), + max_thinking_tokens: Some(128), + ..Default::default() + }; + let token_request = proto::GenerateRequest { + sampling_params: Some(proto::SamplingParams { + seed: Some(42), + ..Default::default() + }), + priority: Some(3), + require_reasoning: Some(false), + max_thinking_tokens: Some(128), + ..Default::default() + }; + + for mapped in [ + build_text_generate_dict("text-request", &text_request).unwrap(), + build_generate_dict("token-request", &token_request).unwrap(), + ] { + assert_eq!(mapped["priority"], serde_json::json!(3)); + assert_eq!(mapped["require_reasoning"], serde_json::json!(false)); + assert_eq!(mapped["max_thinking_tokens"], serde_json::json!(128)); + assert_eq!( + mapped["sampling_params"]["sampling_seed"], + serde_json::json!(42) + ); + } + + for mapped in [ + build_text_generate_dict("text-request", &Default::default()).unwrap(), + build_generate_dict("token-request", &Default::default()).unwrap(), + ] { + assert!(!mapped.contains_key("priority")); + assert!(!mapped.contains_key("require_reasoning")); + assert!(!mapped.contains_key("max_thinking_tokens")); + } + } + + #[test] + fn guided_choice_maps_to_escaped_regex() { + let request = proto::GenerateRequest { + sampling_params: Some(proto::SamplingParams { + guided_decoding: Some(proto::GuidedDecoding { + constraint: Some(proto::guided_decoding::Constraint::Choice( + proto::ChoiceConstraint { + values: vec!["a+b".into(), "x.y".into()], + }, + )), + }), + ..Default::default() + }), + ..Default::default() + }; + let mapped = build_generate_dict("request", &request).unwrap(); + assert_eq!( + mapped["sampling_params"]["regex"], + serde_json::json!("(?:a\\+b|x\\.y)") + ); + } + + #[test] + fn invalid_guidance_combinations_are_rejected() { + let conflicting = proto::GenerateRequest { + sampling_params: Some(proto::SamplingParams { + regex: Some("[a-z]+".into()), + guided_decoding: Some(proto::GuidedDecoding { + constraint: Some(proto::guided_decoding::Constraint::Regex("[0-9]+".into())), + }), + ..Default::default() + }), + ..Default::default() + }; + + let empty_choice = proto::GenerateRequest { + sampling_params: Some(proto::SamplingParams { + guided_decoding: Some(proto::GuidedDecoding { + constraint: Some(proto::guided_decoding::Constraint::Choice( + proto::ChoiceConstraint { values: vec![] }, + )), + }), + ..Default::default() + }), + ..Default::default() + }; + + let empty_legacy_regex = proto::GenerateRequest { + sampling_params: Some(proto::SamplingParams { + regex: Some(String::new()), + ..Default::default() + }), + ..Default::default() + }; + + for request in [conflicting, empty_choice, empty_legacy_regex] { + assert!(build_generate_dict("request", &request).is_err()); + } + } } diff --git a/test/registered/unit/entrypoints/test_grpc_bridge.py b/test/registered/unit/entrypoints/test_grpc_bridge.py new file mode 100644 index 000000000000..dccc7b24bf62 --- /dev/null +++ b/test/registered/unit/entrypoints/test_grpc_bridge.py @@ -0,0 +1,115 @@ +import asyncio +import enum +import unittest +from types import SimpleNamespace + +from sglang.srt.entrypoints.grpc_bridge import RuntimeHandle +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=1, suite="base-a-test-cpu") + + +class _ChunkStatus(enum.Enum): + Ready = 1 + Pending = 2 + Closed = 3 + + +class _RecordingCallback: + def __init__(self): + self.calls = [] + + def __call__(self, payload, *, finished=False, error=None): + self.calls.append((payload, finished, error)) + return _ChunkStatus.Ready + + +class _FakeTokenizerManager: + def __init__(self, responses): + self.responses = responses + + def generate_request(self, obj, request=None): + async def generate(): + for response in self.responses: + yield response + + return generate() + + +def _make_runtime_handle(responses): + handle = RuntimeHandle.__new__(RuntimeHandle) + handle.tokenizer_manager = _FakeTokenizerManager(responses) + return handle + + +class TestNativeGrpcParallelResponses(CustomTestCase): + def test_non_streaming_returns_every_choice_before_finishing(self): + callback = _RecordingCallback() + responses = [ + [ + {"output_ids": [1], "meta_info": {"id": "choice-0"}}, + {"output_ids": [2], "meta_info": {"id": "choice-1"}}, + ] + ] + handle = _make_runtime_handle(responses) + obj = SimpleNamespace(rid="logical", batch_size=1, parallel_sample_num=2) + + asyncio.run( + handle._run_generate( + obj, + callback, + stream=False, + request=None, + ) + ) + + self.assertEqual([call[0]["output_ids"] for call in callback.calls], [[1], [2]]) + self.assertEqual([call[1] for call in callback.calls], [False, True]) + + def test_streaming_first_finished_choice_is_not_batch_terminal(self): + callback = _RecordingCallback() + responses = [ + { + "index": 0, + "output_ids": [1], + "meta_info": {"id": "choice-0", "finish_reason": None}, + }, + { + "index": 0, + "output_ids": [2], + "meta_info": { + "id": "choice-0", + "finish_reason": {"type": "stop"}, + }, + }, + { + "index": 1, + "output_ids": [3], + "meta_info": { + "id": "choice-1", + "finish_reason": {"type": "stop"}, + }, + }, + ] + handle = _make_runtime_handle(responses) + obj = SimpleNamespace(rid="logical", batch_size=1, parallel_sample_num=2) + + asyncio.run( + handle._run_generate( + obj, + callback, + stream=True, + request=None, + ) + ) + + self.assertEqual( + [call[0]["output_ids"] for call in callback.calls], + [[1], [2], [3]], + ) + self.assertEqual([call[1] for call in callback.calls], [False, False, True]) + + +if __name__ == "__main__": + unittest.main(verbosity=2) diff --git a/test/registered/unit/managers/test_io_struct.py b/test/registered/unit/managers/test_io_struct.py index eb80428551a4..835c846c5dc2 100644 --- a/test/registered/unit/managers/test_io_struct.py +++ b/test/registered/unit/managers/test_io_struct.py @@ -305,6 +305,38 @@ def test_single_to_batch_with_parallel_sampling(self): # Modalities should be set for all 3 examples self.assertEqual(req.modalities, ["image", "image", "image"]) + def test_parallel_sampling_keeps_one_logical_rid_per_prompt(self): + """Test logical RID and reasoning control preservation across parallel samples.""" + single = GenerateReqInput( + text="Hello", + rid="single", + sampling_params={"n": 3}, + require_reasoning=True, + max_thinking_tokens=128, + ) + single.normalize_batch_and_arguments() + + self.assertEqual(single.rid, ["single"]) + self.assertEqual([single[i].rid for i in range(3)], ["single"] * 3) + self.assertTrue(all(single[i].require_reasoning for i in range(3))) + self.assertEqual( + [single[i].max_thinking_tokens for i in range(3)], + [128] * 3, + ) + + batch = GenerateReqInput( + text=["Hello", "World"], + rid="batch", + sampling_params={"n": 2}, + ) + batch.normalize_batch_and_arguments() + + self.assertEqual(batch.rid, ["batch_0", "batch_1"]) + self.assertEqual( + [batch[i].rid for i in range(4)], + ["batch_0", "batch_1", "batch_0", "batch_1"], + ) + def test_audio_data_handling(self): """Test handling of audio_data.""" req = copy.deepcopy(self.base_req) @@ -648,6 +680,15 @@ def test_regenerate_rid(self): self.assertNotEqual(original_rid, new_rid) self.assertEqual(req.rid, new_rid) + def test_regenerate_rid_with_parent_prefix(self): + """Test RID regeneration with a logical parent prefix.""" + req = GenerateReqInput(text="Hello", rid="logical") + req.normalize_batch_and_arguments() + + new_rid = req.regenerate_rid(prefix="logical") + + self.assertTrue(new_rid.startswith("logical_")) + def test_error_cases(self): """Test various error cases.""" # Test when neither text, input_ids, nor input_embeds is provided diff --git a/test/registered/unit/managers/test_tokenizer_manager_rid_cleanup.py b/test/registered/unit/managers/test_tokenizer_manager_rid_cleanup.py index 2926a22a5e04..ff93d8ec4712 100644 --- a/test/registered/unit/managers/test_tokenizer_manager_rid_cleanup.py +++ b/test/registered/unit/managers/test_tokenizer_manager_rid_cleanup.py @@ -23,9 +23,19 @@ maybe_stub_sgl_kernel() -from sglang.srt.managers.io_struct import AbortReq, BatchStrOutput, GenerateReqInput -from sglang.srt.managers.tokenizer_manager import ReqState, TokenizerManager -from sglang.srt.observability.req_time_stats import APIServerReqTimeStats +from sglang.srt.managers.io_struct import ( # noqa: E402 + AbortReq, + BatchStrOutput, + GenerateReqInput, +) +from sglang.srt.managers.tokenizer_manager import ( # noqa: E402 + ReqState, + RequestAbortedError, + TokenizerManager, +) +from sglang.srt.observability.req_time_stats import ( # noqa: E402 + APIServerReqTimeStats, +) register_cpu_ci(est_time=15, suite="base-a-test-cpu") @@ -111,6 +121,8 @@ def _make_tokenizer_manager() -> TokenizerManager: tm.server_args.dp_size = 1 tm.disaggregation_mode = "none" tm.rid_to_state = {} + tm.logical_rid_to_child_rids = {} + tm.child_rid_to_logical_rid = {} tm.enable_metrics = False tm.enable_trace = False tm.enable_lora = False @@ -120,10 +132,11 @@ def _make_tokenizer_manager() -> TokenizerManager: tm.dump_requests_folder = "" tm.crash_dump_folder = "" tm.send_to_scheduler = MagicMock() + tm._dispatch_to_scheduler = Mock() return tm -def _make_req_state(rid: str = "test_rid") -> ReqState: +def _make_req_state(rid: str = "test_rid", *, dispatched: bool = False) -> ReqState: """Create a minimal ReqState for testing.""" obj = Mock(spec=GenerateReqInput) obj.rid = rid @@ -137,6 +150,7 @@ def _make_req_state(rid: str = "test_rid") -> ReqState: event=asyncio.Event(), obj=obj, time_stats=APIServerReqTimeStats(), + dispatched=dispatched, ) @@ -338,6 +352,19 @@ def test_unique_rid_succeeds(self): tm._init_req_state(obj) self.assertIn(rid, tm.rid_to_state) + def test_batch_duplicate_preflight_does_not_insert_partial_state(self): + tm = _make_tokenizer_manager() + existing_rid = "existing" + existing_state = _make_req_state(existing_rid) + tm.rid_to_state[existing_rid] = existing_state + obj = _make_generate_obj(["new", existing_rid], is_single=False) + + with self.assertRaisesRegex(ValueError, "Duplicate request ID"): + tm._init_req_state(obj) + + self.assertNotIn("new", tm.rid_to_state) + self.assertIs(tm.rid_to_state[existing_rid], existing_state) + class TestResubmitAfterCompletion(CustomTestCase): """End-to-end test: complete a request, then resubmit with the same rid.""" @@ -409,6 +436,7 @@ def _make_tm_for_generate() -> TokenizerManager: tm = _make_tokenizer_manager() tm.server_args.language_only = False tm.server_args.tokenizer_worker_num = 1 + tm.server_args.enable_strict_thinking = False tm.auto_create_handle_loop = Mock() tm._set_default_priority = Mock() tm.request_logger = Mock() @@ -429,6 +457,7 @@ def _make_generate_obj(rid, is_single): obj.received_time = 0.0 obj.external_trace_header = None obj.bootstrap_room = None + obj.max_thinking_tokens = None obj.normalize_batch_and_arguments = Mock() if not is_single: obj.__getitem__.side_effect = lambda i: Mock() @@ -438,17 +467,20 @@ def _make_generate_obj(rid, is_single): class TestDiscardPendingReqStates(CustomTestCase): """Direct tests for _discard_pending_req_states.""" - def test_discard_single(self): + def test_discard_single_aborts_scheduler_before_cleanup(self): tm = _make_tokenizer_manager() rid = "d_single" - tm.rid_to_state[rid] = _make_req_state(rid) + tm.rid_to_state[rid] = _make_req_state(rid, dispatched=True) obj = Mock(spec=GenerateReqInput) obj.is_single = True obj.rid = rid tm._discard_pending_req_states(obj) self.assertNotIn(rid, tm.rid_to_state) + abort_req = tm._dispatch_to_scheduler.call_args.args[0] + self.assertEqual(abort_req.rid, rid) + self.assertFalse(abort_req.abort_all) - def test_discard_batch_removes_all(self): + def test_discard_unsent_batch_without_scheduler_abort(self): tm = _make_tokenizer_manager() rids = ["d0", "d1", "d2"] for r in rids: @@ -459,6 +491,7 @@ def test_discard_batch_removes_all(self): tm._discard_pending_req_states(obj) for r in rids: self.assertNotIn(r, tm.rid_to_state) + tm._dispatch_to_scheduler.assert_not_called() def test_discard_ignores_already_removed(self): """Popping a rid that is no longer present must not raise.""" @@ -470,6 +503,150 @@ def test_discard_ignores_already_removed(self): tm._discard_pending_req_states(obj) # must not raise self.assertNotIn("p1", tm.rid_to_state) + def test_parallel_cleanup_aborts_children_and_allows_parent_reuse(self): + tm = _make_tokenizer_manager() + parent = _make_generate_obj("parent", is_single=True) + lifecycle_ids = tm._init_req_state(parent) + + child_rids = {"prefix", "choice_0", "choice_1"} + for child_rid in child_rids: + child = _make_generate_obj(child_rid, is_single=True) + tm._init_child_req_state("parent", child) + tm.rid_to_state[child_rid].dispatched = True + tm._remove_req_state("parent") + + tm._discard_pending_req_states(parent, lifecycle_ids) + + aborted_rids = { + call.args[0].rid for call in tm._dispatch_to_scheduler.call_args_list + } + self.assertEqual(aborted_rids, child_rids) + self.assertFalse(tm.rid_to_state) + self.assertFalse(tm.logical_rid_to_child_rids) + self.assertFalse(tm.child_rid_to_logical_rid) + + tm._init_req_state(_make_generate_obj("parent", is_single=True)) + self.assertIn("parent", tm.rid_to_state) + + def test_stale_cleanup_does_not_remove_reused_rid(self): + tm = _make_tokenizer_manager() + old_obj = _make_generate_obj("reused", is_single=True) + old_lifecycle_ids = tm._init_req_state(old_obj) + tm._remove_req_state("reused") + + replacement = _make_generate_obj("reused", is_single=True) + tm._init_req_state(replacement) + replacement_state = tm.rid_to_state["reused"] + + tm._discard_pending_req_states(old_obj, old_lifecycle_ids) + + self.assertIs(tm.rid_to_state["reused"], replacement_state) + tm._dispatch_to_scheduler.assert_not_called() + + +class TestParallelAbortRouting(CustomTestCase): + def test_parent_abort_fans_out_to_children(self): + tm = _make_tokenizer_manager() + tm.server_args.tokenizer_worker_num = 1 + tm._register_child_rid("parent", "choice_0") + tm._register_child_rid("parent", "choice_1") + + tm.abort_request("parent") + + requests = [call.args[0] for call in tm._dispatch_to_scheduler.call_args_list] + self.assertEqual( + {request.rid for request in requests}, {"choice_0", "choice_1"} + ) + self.assertTrue(all(not request.abort_all for request in requests)) + + +class TestParallelStreamTaskCleanup(CustomTestCase): + def test_failing_choice_cancels_and_closes_sibling_waiters(self): + tm = _make_tokenizer_manager() + + async def drive(): + sibling_closed = asyncio.Event() + + async def failing_choice(): + await asyncio.sleep(0) + raise RuntimeError("choice failed") + yield # pragma: no cover + + async def blocked_choice(): + try: + await asyncio.Event().wait() + yield # pragma: no cover + finally: + sibling_closed.set() + + stream = tm._stream_batch_responses( + [failing_choice(), blocked_choice()], + ["choice-0", "choice-1"], + ) + with self.assertRaisesRegex(RuntimeError, "choice failed"): + await stream.__anext__() + self.assertTrue(sibling_closed.is_set()) + + asyncio.run(drive()) + + def test_failing_non_stream_choice_cancels_and_closes_sibling_waiters(self): + tm = _make_tokenizer_manager() + + async def drive(): + sibling_closed = asyncio.Event() + + async def failing_choice(): + await asyncio.sleep(0) + raise RuntimeError("choice failed") + yield # pragma: no cover + + async def blocked_choice(): + try: + await asyncio.Event().wait() + yield # pragma: no cover + finally: + sibling_closed.set() + + with self.assertRaisesRegex(RuntimeError, "choice failed"): + await tm._collect_batch_responses([failing_choice(), blocked_choice()]) + self.assertTrue(sibling_closed.is_set()) + + asyncio.run(drive()) + + +class TestParallelRidReuse(CustomTestCase): + def test_completed_n2_request_can_repeat_the_same_logical_rid(self): + tm = _make_tokenizer_manager() + + async def complete_child(rid): + await tm._handle_batch_output(_make_batch_str_output(rid)) + + for _ in range(2): + logical = GenerateReqInput( + text="hello", + rid="repeat-n2", + sampling_params={"n": 2}, + ) + logical.normalize_batch_and_arguments() + tm._init_req_state(logical) + + prefix = GenerateReqInput(text="hello", rid="prefix") + prefix.normalize_batch_and_arguments() + tm._init_child_req_state("repeat-n2", prefix) + asyncio.run(complete_child("prefix")) + + for child_rid in ("choice-0", "choice-1"): + child = GenerateReqInput(text="hello", rid=child_rid) + child.normalize_batch_and_arguments() + tm._init_child_req_state("repeat-n2", child) + tm._remove_req_state("repeat-n2") + asyncio.run(complete_child("choice-0")) + asyncio.run(complete_child("choice-1")) + + self.assertFalse(tm.rid_to_state) + self.assertFalse(tm.logical_rid_to_child_rids) + self.assertFalse(tm.child_rid_to_logical_rid) + class TestGenerateRequestCleanupOnDispatchFailure(CustomTestCase): """generate_request must not leak rid_to_state when dispatch fails. @@ -497,6 +674,7 @@ async def drive(): # Got past _init_req_state (which created the entry) ... tm._tokenize_one_request.assert_awaited_once() tm._send_one_request.assert_not_called() + tm._dispatch_to_scheduler.assert_not_called() # ... and the entry was cleaned up rather than leaked. self.assertNotIn(rid, tm.rid_to_state) @@ -521,6 +699,66 @@ async def drive(): # All sub-request entries created by _init_req_state are cleaned up. for r in rids: self.assertNotIn(r, tm.rid_to_state) + tm._dispatch_to_scheduler.assert_not_called() + + def test_interrupted_parallel_tokenization_prevents_child_dispatch(self): + for remove_state in (False, True): + with self.subTest(remove_state=remove_state): + tm = _make_tm_for_generate() + tm._send_one_request = Mock() + obj = GenerateReqInput( + text="hello", + rid="interrupted-during-tokenization", + sampling_params={"n": 2}, + ) + + async def drive(): + tokenization_started = asyncio.Event() + allow_tokenization = asyncio.Event() + + async def blocked_tokenization(_obj): + tokenization_started.set() + await allow_tokenization.wait() + return MagicMock() + + tm._tokenize_one_request = blocked_tokenization + response = tm.generate_request(obj) + task = asyncio.create_task(response.__anext__()) + await tokenization_started.wait() + if remove_state: + tm._remove_req_state("interrupted-during-tokenization") + else: + tm.abort_request("interrupted-during-tokenization") + allow_tokenization.set() + with self.assertRaisesRegex( + RequestAbortedError, "interrupted-during-tokenization" + ): + await task + + asyncio.run(drive()) + + tm._send_one_request.assert_not_called() + tm._dispatch_to_scheduler.assert_not_called() + self.assertFalse(tm.rid_to_state) + self.assertFalse(tm.logical_rid_to_child_rids) + self.assertFalse(tm.child_rid_to_logical_rid) + + def test_thinking_budget_rejects_runtime_without_strict_thinking(self): + tm = _make_tm_for_generate() + obj = GenerateReqInput( + text="hello", + rid="thinking-budget", + sampling_params={}, + max_thinking_tokens=32, + ) + + async def drive(): + await tm.generate_request(obj).__anext__() + + with self.assertRaisesRegex(ValueError, "--enable-strict-thinking"): + asyncio.run(drive()) + + self.assertFalse(tm.rid_to_state) if __name__ == "__main__": From bc2fc41fae79f9651fdb06a7da47d902bc5ca6c3 Mon Sep 17 00:00:00 2001 From: Brayden Zhong Date: Wed, 5 Aug 2026 14:27:05 -0700 Subject: [PATCH 02/47] Fix broken Nemotron DP attention (#33123) Co-authored-by: Brayden Zhong --- python/sglang/srt/environ.py | 3 ++ .../layers/attention/mamba/mamba2_metadata.py | 2 +- .../layers/moe/token_dispatcher/flashinfer.py | 11 ++-- .../model_runner_components/layer_setup.py | 6 ++- python/sglang/srt/models/nemotron_h_mtp.py | 7 ++- python/sglang/srt/server_args.py | 5 +- .../test_nvidia_nemotron_3_super_nvfp4.py | 54 ++++++++++++++++++- 7 files changed, 74 insertions(+), 14 deletions(-) diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 820d0b3626cb..81cbfc18ddf2 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -675,6 +675,9 @@ class Envs: SGLANG_FLASHINFER_USE_PAGED = EnvBool(False) # Default to the pick from flashinfer SGLANG_FLASHINFER_WORKSPACE_SIZE = EnvInt(384 * 1024 * 1024) + # Per-rank dispatch capacity of the FlashInfer MoE A2A dispatcher. Unset + # means each call site keeps its own default. + SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK = EnvInt(None) # Enable NVFP4 per-token activation scaling path for FlashInfer TRT-LLM MoE. SGLANG_FLASHINFER_NVFP4_PER_TOKEN_ACTIVATION = EnvBool(False) # Launch the TRT-LLM MoE grouped GEMMs with PDL only at or below this diff --git a/python/sglang/srt/layers/attention/mamba/mamba2_metadata.py b/python/sglang/srt/layers/attention/mamba/mamba2_metadata.py index 42d8a0c979d4..2cb3dd4f403e 100644 --- a/python/sglang/srt/layers/attention/mamba/mamba2_metadata.py +++ b/python/sglang/srt/layers/attention/mamba/mamba2_metadata.py @@ -240,7 +240,7 @@ def prepare_mixed( batch_size = getattr(forward_batch, "_original_batch_size", None) if batch_size is None: batch_size = len(forward_batch.seq_lens) - num_decodes = batch_size - num_prefills + num_decodes = max(0, batch_size - num_prefills) context_lens_tensor = forward_batch.extend_prefix_lens assert context_lens_tensor is not None has_initial_states = context_lens_tensor > 0 diff --git a/python/sglang/srt/layers/moe/token_dispatcher/flashinfer.py b/python/sglang/srt/layers/moe/token_dispatcher/flashinfer.py index 670fb428ce07..35759b9199b3 100644 --- a/python/sglang/srt/layers/moe/token_dispatcher/flashinfer.py +++ b/python/sglang/srt/layers/moe/token_dispatcher/flashinfer.py @@ -29,7 +29,6 @@ from sglang.srt.layers.moe.utils import get_moe_runner_backend from sglang.srt.runtime_context import get_schedule, get_spec from sglang.srt.speculative.spec_info import SpeculativeAlgorithm -from sglang.srt.utils import get_int_env_var try: from flashinfer import nvfp4_block_scale_interleave @@ -125,9 +124,13 @@ def __init__( # (which warms up at batch_size = req_to_token_pool.size). cps = get_schedule().chunked_prefill_size default_max_tokens = max(cps if cps and cps > 0 else 4096, 4096) - self.max_num_tokens = get_int_env_var( - "SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK", - default_max_tokens, + configured_max_tokens = ( + envs.SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK.get() + ) + self.max_num_tokens = ( + configured_max_tokens + if configured_max_tokens is not None + else default_max_tokens ) # Calculate workspace size. For eagle mode, use the larger workspace size since nextn layer will be unquantized. diff --git a/python/sglang/srt/model_executor/model_runner_components/layer_setup.py b/python/sglang/srt/model_executor/model_runner_components/layer_setup.py index c6485abca01b..83a94a63b993 100644 --- a/python/sglang/srt/model_executor/model_runner_components/layer_setup.py +++ b/python/sglang/srt/model_executor/model_runner_components/layer_setup.py @@ -3,6 +3,7 @@ from typing import TYPE_CHECKING, Any, NamedTuple import msgspec +from torch import nn if TYPE_CHECKING: from sglang.srt.configs.model_config import ModelConfig @@ -23,7 +24,10 @@ def compute_attention_and_moe_layers(layer_model: Any) -> AttentionAndMoeLayers: moe_fusions: list[Any] = [] dsa_indexers: list[Any] = [] mha_companion_layers: list[Any] = [] - for layer in layer_model.layers: + layers = layer_model.layers + if isinstance(layers, nn.ModuleDict): + layers = layers.values() + for layer in layers: attn_layer = None mha_companion_layer = None if hasattr(layer, "self_attn"): diff --git a/python/sglang/srt/models/nemotron_h_mtp.py b/python/sglang/srt/models/nemotron_h_mtp.py index b3cd79fadd36..3c375e90c84d 100644 --- a/python/sglang/srt/models/nemotron_h_mtp.py +++ b/python/sglang/srt/models/nemotron_h_mtp.py @@ -288,13 +288,14 @@ def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor: def forward( self, input_ids: torch.Tensor, - hidden_states: torch.Tensor, + positions: torch.Tensor, forward_batch: ForwardBatch, inputs_embeds: torch.Tensor | None = None, ) -> torch.Tensor: if inputs_embeds is None: inputs_embeds = self.get_input_embeddings(input_ids) + hidden_states = forward_batch.spec_info.hidden_states residual = None for i in range(self.pattern_len): @@ -352,11 +353,9 @@ def forward( input_embeds: torch.Tensor | None = None, **kwargs, ) -> torch.Tensor: - hidden_states = forward_batch.spec_info.hidden_states - hidden_states = self.model( input_ids, - hidden_states, + positions, forward_batch, input_embeds, ) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index d93983ee020a..18664959e959 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -68,7 +68,6 @@ get_device, get_device_memory_capacity, get_device_sm, - get_int_env_var, get_quantization_config, human_readable_int, is_blackwell_supported, @@ -6632,8 +6631,8 @@ def _validate_cutedsl_a2a_token_budget(self): ): return required_tokens = self.cutedsl_moe_max_num_tokens() - max_dispatch_tokens_per_rank = get_int_env_var( - "SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK", 1024 + max_dispatch_tokens_per_rank = ( + envs.SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK.get() or 1024 ) max_cutedsl_tokens = max_dispatch_tokens_per_rank * view.ep_size if max_cutedsl_tokens < required_tokens: diff --git a/test/registered/4-gpu-models/test_nvidia_nemotron_3_super_nvfp4.py b/test/registered/4-gpu-models/test_nvidia_nemotron_3_super_nvfp4.py index 4737d0461d3b..ba37a7ae9889 100644 --- a/test/registered/4-gpu-models/test_nvidia_nemotron_3_super_nvfp4.py +++ b/test/registered/4-gpu-models/test_nvidia_nemotron_3_super_nvfp4.py @@ -12,7 +12,7 @@ popen_launch_server, ) -register_cuda_ci(est_time=540, suite="nightly-4-gpu-b200", nightly=True) +register_cuda_ci(est_time=810, suite="nightly-4-gpu-b200", nightly=True) NEMOTRON_3_SUPER_NVFP4_MODEL = "nvidia/NVIDIA-Nemotron-3-Super-120B-A12B-NVFP4" @@ -29,6 +29,31 @@ '{"enable_multithread_load": true, "num_threads": 17}', ] +DP_ATTENTION_EP_ARGS = [ + "--dp-size", + "4", + "--enable-dp-attention", + "--enable-dp-lm-head", + "--ep-size", + "4", + "--moe-a2a-backend", + "flashinfer", + "--moe-runner-backend", + "flashinfer_cutedsl", + "--mamba-full-memory-ratio", + "5.0", + "--mamba-radix-cache-strategy", + "extra_buffer", + "--attention-backend", + "trtllm_mha", + "--max-running-requests", + "1024", + "--mem-fraction-static", + "0.93", + "--max-prefill-tokens", + "8192", +] + MTP_ARGS = [ "--speculative-algorithm", "EAGLE", @@ -107,5 +132,32 @@ def test_gsm8k(self): _run_gsm8k(self) +class TestNvidiaNemotron3SuperNVFP4DPAttentionEP(CustomTestCase): + """DP attention + EP with the FlashInfer one-sided A2A and CuteDSL MoE runner.""" + + @classmethod + def setUpClass(cls): + cls.model = NEMOTRON_3_SUPER_NVFP4_MODEL + cls.base_url = DEFAULT_URL_FOR_TEST + with ( + envs.SGLANG_ENABLE_ASYNC_ASSERT.override(0), + envs.SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(4096), + envs.SGLANG_FLASHINFER_WORKSPACE_SIZE.override(1024 * 1024 * 1024), + ): + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=NEMOTRON_3_SUPER_NVFP4_ARGS + DP_ATTENTION_EP_ARGS, + ) + + @classmethod + def tearDownClass(cls): + kill_process_tree(cls.process.pid) + + def test_gsm8k(self): + _run_gsm8k(self) + + if __name__ == "__main__": unittest.main() From 4fbc160c4a8652aa5260505ba14de75ff28b7827 Mon Sep 17 00:00:00 2001 From: Kangyan-Zhou Date: Wed, 5 Aug 2026 15:49:11 -0700 Subject: [PATCH 03/47] [Cherry-pick to release/v0.5.17] [DCP] Match the replicated draft KV pool's page granularity to its allocator (#33348) (#33762) Co-authored-by: Khoa Pham --- .../srt/mem_cache/kv_cache_configurator.py | 56 +++++++++++-------- 1 file changed, 34 insertions(+), 22 deletions(-) diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index acd0f413742c..221c73dccdfb 100644 --- a/python/sglang/srt/mem_cache/kv_cache_configurator.py +++ b/python/sglang/srt/mem_cache/kv_cache_configurator.py @@ -270,6 +270,19 @@ def configure(self, *, pre_model_load_memory: int) -> KVCacheConfigResult: unified_memory_pool=pools.unified_memory_pool, ) + # Note(kpham-sgl): + # 1. A replicated draft indexes the allocator's virtual locs raw, so its pools + # span and page that space; the sharded target translates and stays per-rank. + # 2. A pool must page as its allocator does, or its last page falls short. + @property + def loc_space_scale(self) -> int: + dcp_size = self.server_args.dcp_size + return dcp_size if (self.is_draft_worker and dcp_size > 1) else 1 + + @property + def pool_page_size(self) -> int: + return get_schedule().page_size * self.loc_space_scale + def _derive_pool_sizes(self, *, config: MemoryPoolConfig) -> _PoolSizes: max_total_num_tokens = config.max_total_num_tokens max_running_requests = config.max_running_requests @@ -281,13 +294,12 @@ def _derive_pool_sizes(self, *, config: MemoryPoolConfig) -> _PoolSizes: # Draft pools are replicated, not DCP-sharded, yet consume the shared # allocator's virtual locs in [0, max_total * dcp_size) untranslated. - dcp_size = self.server_args.dcp_size - if self.is_draft_worker and dcp_size > 1: - max_total_num_tokens *= dcp_size - if full_max_total_num_tokens is not None: - full_max_total_num_tokens *= dcp_size - if swa_max_total_num_tokens is not None: - swa_max_total_num_tokens *= dcp_size + loc_scale = self.loc_space_scale + max_total_num_tokens *= loc_scale + if full_max_total_num_tokens is not None: + full_max_total_num_tokens *= loc_scale + if swa_max_total_num_tokens is not None: + swa_max_total_num_tokens *= loc_scale # DSV4 compressed-attention pool sizes. Draft worker reuses target's # full/swa sizes but does NOT own c4/c128/state pools (those live on @@ -1000,7 +1012,7 @@ def _build_dsv4_kv_pool( c128_size=c128_max_total_num_tokens, c4_state_pool_size=c4_state_pool_size, c128_state_pool_size=c128_state_pool_size, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, swa_page_size=swa_page_size, sliding_window=self.model_config.window_size, dtype=self.kv_cache_dtype, @@ -1026,7 +1038,7 @@ def _build_oot_dsa_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: PoolCls = current_platform.get_dsa_kv_pool_cls() token_to_kv_pool = PoolCls( max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, kv_lora_rank=self.model_config.kv_lora_rank, qk_rope_head_dim=self.model_config.qk_rope_head_dim, @@ -1050,7 +1062,7 @@ def _build_oot_mla_kv_pool( PoolCls = current_platform.get_mla_kv_pool_cls() token_to_kv_pool = PoolCls( max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, kv_lora_rank=self.model_config.kv_lora_rank, qk_rope_head_dim=self.model_config.qk_rope_head_dim, @@ -1067,7 +1079,7 @@ def _build_oot_mha_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: PoolCls = current_platform.get_mha_kv_pool_cls() token_to_kv_pool = PoolCls( max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), head_dim=self.model_config.head_dim, @@ -1104,7 +1116,7 @@ def _build_ascend_swa_kv_pool( token_to_kv_pool = SWAKVPool( size=full_max_total_num_tokens, size_swa=swa_max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, post_capture_active=self.post_capture_kv_active, head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), @@ -1126,7 +1138,7 @@ def _build_ascend_mla_kv_pool( token_to_kv_pool = NPUMLATokenToKVPool( max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, kv_lora_rank=self.model_config.kv_lora_rank, qk_rope_head_dim=self.model_config.qk_rope_head_dim, @@ -1146,7 +1158,7 @@ def _build_ascend_mha_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: token_to_kv_pool = NPUMHATokenToKVPool( max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), head_dim=self.model_config.head_dim, @@ -1186,7 +1198,7 @@ def _build_dsa_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: PoolCls = DSATokenToKVPool token_to_kv_pool = PoolCls( max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, kv_lora_rank=self.model_config.kv_lora_rank, qk_rope_head_dim=self.model_config.qk_rope_head_dim, @@ -1208,7 +1220,7 @@ def _build_dsa_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: def _build_mla_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: token_to_kv_pool = MLATokenToKVPoolFP4( max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, kv_lora_rank=self.model_config.kv_lora_rank, qk_rope_head_dim=self.model_config.qk_rope_head_dim, @@ -1223,7 +1235,7 @@ def _build_mla_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: def _build_mla_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: token_to_kv_pool = MLATokenToKVPool( max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, kv_lora_rank=self.model_config.kv_lora_rank, qk_rope_head_dim=self.model_config.qk_rope_head_dim, @@ -1286,7 +1298,7 @@ def _build_hybrid_swa_kv_pool( token_to_kv_pool = SWAKVPool( size=full_max_total_num_tokens, size_swa=size_swa, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, post_capture_active=self.post_capture_kv_active, head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), @@ -1311,7 +1323,7 @@ def _build_minimax_sparse_kv_pool(self, *, max_total_num_tokens: int) -> KVCache ) token_to_kv_pool = MiniMaxSparseKVPool( size=max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, # fp8 attn-GEMM mode opts the lightning-indexer cache into # fp8 too (fp8 indexer GEMMs); fp8 KV without the mode @@ -1368,7 +1380,7 @@ def _build_hybrid_linear_kv_pool( else mha_pool_class ) token_to_kv_pool = HybridLinearKVPool( - page_size=get_schedule().page_size, + page_size=self.pool_page_size, size=max_total_num_tokens, dtype=self.kv_cache_dtype, head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), @@ -1391,7 +1403,7 @@ def _build_hybrid_linear_kv_pool( def _build_mha_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: token_to_kv_pool = MHATokenToKVPoolFP4( max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), head_dim=self.model_config.head_dim, @@ -1424,7 +1436,7 @@ def _build_mha_kv_pool( pool_kwargs["post_capture_active"] = self.post_capture_kv_active token_to_kv_pool = pool_cls( max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), head_dim=self.model_config.head_dim, From f2890ada2f7148af5e066af6ecb6d3f83a485285 Mon Sep 17 00:00:00 2001 From: Kangyan-Zhou Date: Wed, 5 Aug 2026 18:10:03 -0700 Subject: [PATCH 04/47] [Cherry-pick to release/v0.5.17] Fix Nightly NV CI (#33564) (#33779) Co-authored-by: Brayden Zhong --- python/sglang/srt/arg_groups/overrides.py | 3 ++ python/sglang/srt/layers/n_gram_embedding.py | 2 +- .../registered/8-gpu-models/test_glm52_fp8.py | 2 +- .../test_longcat_flash_lite_fp8.py | 50 +++++++++++-------- 4 files changed, 34 insertions(+), 23 deletions(-) diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index 41f3550be935..245f5943ca5f 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -1725,6 +1725,9 @@ def _deepseek_moe_quant_resolution(view: Any) -> dict: if ( view.moe_a2a_backend == "none" and view.moe_runner_backend == "auto" + # LongCat top-k spans the zero-expert logits, which trtllm-gen's + # fused routing cannot see. + and not model_arch.startswith("LongcatFlash") and ( quantization in ["fp8", "modelopt_fp8", "modelopt_fp4", "modelopt_mixed"] diff --git a/python/sglang/srt/layers/n_gram_embedding.py b/python/sglang/srt/layers/n_gram_embedding.py index 6bdb16d62260..699e7c326a52 100644 --- a/python/sglang/srt/layers/n_gram_embedding.py +++ b/python/sglang/srt/layers/n_gram_embedding.py @@ -49,7 +49,7 @@ def __init__( + int(over_embedding_m + i * 2 + 1) ) self.oe_embeder = VocabParallelEmbedding( - num_embeddings=self.exclusive_oe_embedder_size_sums[-1], + num_embeddings=int(self.exclusive_oe_embedder_size_sums[-1]), embedding_dim=oe_hidden_dim, use_attn_tp_group=use_attn_tp_group, ) diff --git a/test/registered/8-gpu-models/test_glm52_fp8.py b/test/registered/8-gpu-models/test_glm52_fp8.py index 8995298aa731..dca35ed6d50d 100644 --- a/test/registered/8-gpu-models/test_glm52_fp8.py +++ b/test/registered/8-gpu-models/test_glm52_fp8.py @@ -15,7 +15,7 @@ "--trust-remote-code", "--reasoning-parser=glm45", "--tool-call-parser=glm47", - "--mem-fraction-static=0.85", + "--mem-fraction-static=0.8", "--enable-metrics", ] diff --git a/test/registered/8-gpu-models/test_longcat_flash_lite_fp8.py b/test/registered/8-gpu-models/test_longcat_flash_lite_fp8.py index b48825e10dbc..feec8d0e7fd3 100644 --- a/test/registered/8-gpu-models/test_longcat_flash_lite_fp8.py +++ b/test/registered/8-gpu-models/test_longcat_flash_lite_fp8.py @@ -6,8 +6,7 @@ from sglang.test.run_combined_tests import run_combined_tests from sglang.test.test_utils import ModelLaunchSettings -# Runs on both H200 and B200 via the nightly-8-gpu-common suite. -register_cuda_ci(est_time=1200, suite="nightly-8-gpu-common", nightly=True) +register_cuda_ci(est_time=1200, suite="nightly-8-gpu-h200", nightly=True) # LongCat-Flash-Lite-FP8 is the smallest member of the LongCat family # (~138 GB FP8 weights, hidden=3072, 14 layers, 256 routed + 128 zero @@ -20,7 +19,6 @@ # LongcatFlashNgramForCausalLM is remapped to model_type "longcat_flash") # with LongCat-Flash-Chat-FP8 and LongCat-2.0-FP8, so it guards the shared # code paths that recent LongCat EP fixes touched: -# - the scheduler moe-topk gate for --moe-a2a-backend (PR #30975) # - ScMoE dense-branch gather (RoPE) + the MoE-vs-DeepEPMoE double # all_reduce fix (PR #31311) # - the zero-expert (identity) compute path (zero_expert_num=128) and @@ -41,28 +39,14 @@ class TestLongCatFlashLiteFp8(unittest.TestCase): Two variants exercise the two MoE all-to-all backends that the LongCat EP fixes gate on: - - EP8 + deepep : real expert parallelism (the path #30975/#31311 fix) - EP8 + none : EP-over-TP baseline (all_reduce / gather correctness) + - EP8 + deepep : real expert parallelism (the path #30975/#31311 fix), + skipped until DeepEP raises its low-latency top-k cap to 12 """ - def test_longcat_flash_lite_fp8(self): - variants = [ - ModelLaunchSettings( - LONGCAT_FLASH_LITE_FP8_MODEL_PATH, - tp_size=8, - extra_args=COMMON_ARGS + ["--ep=8"], - variant="TP8+EP8+none", - ), - ModelLaunchSettings( - LONGCAT_FLASH_LITE_FP8_MODEL_PATH, - tp_size=8, - extra_args=COMMON_ARGS + ["--ep=8", "--moe-a2a-backend=deepep"], - variant="TP8+EP8+deepep", - ), - ] - + def _run(self, variant: ModelLaunchSettings): run_combined_tests( - models=variants, + models=[variant], test_name="LongCat-Flash-Lite-FP8", # Measured 2026-07-22 on 8xH100-80GB, gsm8k 200q, 5-shot, greedy: # TP8+EP8+none -> 0.840 @@ -78,6 +62,30 @@ def test_longcat_flash_lite_fp8(self): ), ) + def test_longcat_flash_lite_fp8(self): + self._run( + ModelLaunchSettings( + LONGCAT_FLASH_LITE_FP8_MODEL_PATH, + tp_size=8, + extra_args=COMMON_ARGS + ["--ep=8"], + variant="TP8+EP8+none", + ) + ) + + @unittest.skip( + "Blocked: DeepEP low-latency dispatch asserts num_topk <= kNumMaxTopK " + "(11 in internode_ll.cu), LongCat moe_topk is 12." + ) + def test_longcat_flash_lite_fp8_deepep(self): + self._run( + ModelLaunchSettings( + LONGCAT_FLASH_LITE_FP8_MODEL_PATH, + tp_size=8, + extra_args=COMMON_ARGS + ["--ep=8", "--moe-a2a-backend=deepep"], + variant="TP8+EP8+deepep", + ) + ) + if __name__ == "__main__": unittest.main() From 7163b4ae58c1ac8c965ff453f724c8410a1206bc Mon Sep 17 00:00:00 2001 From: Kangyan-Zhou Date: Wed, 5 Aug 2026 18:54:46 -0700 Subject: [PATCH 05/47] [Cherry-pick to release/v0.5.17] [CI] Temporarily disable prefill cuda graph for qwen3.5 nightly test (#33772) (#33786) Co-authored-by: Baizhou Zhang --- test/registered/8-gpu-models/test_qwen35.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/test/registered/8-gpu-models/test_qwen35.py b/test/registered/8-gpu-models/test_qwen35.py index bf7fb2d01e12..f8e48e75727d 100644 --- a/test/registered/8-gpu-models/test_qwen35.py +++ b/test/registered/8-gpu-models/test_qwen35.py @@ -30,7 +30,7 @@ def test_qwen35(self): "--tool-call-parser=qwen3_coder", "--mem-fraction-static=0.8", ] - dp_args = ["--dp=8", "--enable-dp-attention"] + dp_args = ["--dp=8", "--enable-dp-attention", "--disable-prefill-cuda-graph"] mtp_args = [ "--speculative-algorithm=EAGLE", "--speculative-num-steps=3", From 20cef1816b68c8575b61c156723efff7a7b3890d Mon Sep 17 00:00:00 2001 From: Brayden Zhong Date: Thu, 6 Aug 2026 13:59:28 -0700 Subject: [PATCH 06/47] Fix Mistral-Large-3 EAGLE draft skipping DeepseekV2Model.__init__ --- .../srt/models/mistral_large_3_eagle.py | 52 +++---------------- .../8-gpu-models/test_mistral_large3.py | 20 ++++--- 2 files changed, 18 insertions(+), 54 deletions(-) diff --git a/python/sglang/srt/models/mistral_large_3_eagle.py b/python/sglang/srt/models/mistral_large_3_eagle.py index ae487c52c86e..b42351e9cf5b 100644 --- a/python/sglang/srt/models/mistral_large_3_eagle.py +++ b/python/sglang/srt/models/mistral_large_3_eagle.py @@ -4,19 +4,12 @@ from typing import Optional import torch -from torch import nn from transformers import PretrainedConfig -from sglang.srt.configs.model_config import is_deepseek_dsa -from sglang.srt.distributed import get_pp_group -from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp -from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import RowParallelLinear from sglang.srt.layers.quantization.base_config import QuantizationConfig -from sglang.srt.layers.utils.cp_utils import is_prefill_context_parallel_enabled -from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors -from sglang.srt.models.deepseek_v2 import DeepseekV2DecoderLayer, DeepseekV2Model +from sglang.srt.models.deepseek_v2 import DeepseekV2Model from sglang.srt.models.mistral_large_3 import MistralLarge3ForCausalLM from sglang.srt.utils import add_prefix @@ -31,50 +24,17 @@ def __init__( quant_config: Optional[QuantizationConfig] = None, prefix: str = "", ): - nn.Module.__init__(self) - - self.config = config - self.vocab_size = config.vocab_size - assert get_pp_group().world_size == 1 - self.pp_group = get_pp_group() - self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() - self.mla_enable_prefill_cp = ( - is_prefill_context_parallel_enabled() and not is_deepseek_dsa(config) - ) - - self.embed_tokens = VocabParallelEmbedding( - config.vocab_size, - config.hidden_size, - prefix=add_prefix("embed_tokens", prefix), - ) - - self.layers = nn.ModuleList( - [ - DeepseekV2DecoderLayer( - config=config, - prefix=add_prefix(prefix, f"layers.{i}"), - quant_config=quant_config, - layer_id=i, - dsa_enable_prefill_cp=self.dsa_enable_prefill_cp, - mla_enable_prefill_cp=self.mla_enable_prefill_cp, - ) - for i in range(self.config.num_hidden_layers) - ] - ) - self.start_layer = 0 - self.end_layer = self.config.num_hidden_layers + super().__init__(config, quant_config, prefix=prefix) + assert self.pp_group.world_size == 1 self.fc = RowParallelLinear( - self.config.hidden_size * 2, - self.config.hidden_size, + config.hidden_size * 2, + config.hidden_size, bias=False, quant_config=quant_config, - prefix=add_prefix(prefix, "fc"), + prefix=add_prefix("fc", prefix), input_is_parallel=False, ) - self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) - self.layers_to_capture = [] - self.llama_4_scaling_config = getattr(config, "llama_4_scaling", None) def forward( self, diff --git a/test/registered/8-gpu-models/test_mistral_large3.py b/test/registered/8-gpu-models/test_mistral_large3.py index 58587d45e3e2..9dc9c221110d 100644 --- a/test/registered/8-gpu-models/test_mistral_large3.py +++ b/test/registered/8-gpu-models/test_mistral_large3.py @@ -1,6 +1,7 @@ import os import unittest +from sglang.srt.environ import envs from sglang.test.accuracy_test_runner import AccuracyTestParams from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.performance_test_runner import PerformanceTestParams @@ -84,14 +85,17 @@ def test_mistral_large3_all_variants(self): ), ] - run_combined_tests( - models=variants, - test_name="Mistral-Large-3", - accuracy_params=AccuracyTestParams(dataset="gsm8k", baseline_accuracy=0.85), - performance_params=PerformanceTestParams( - profile_dir="performance_profiles_mistral_large3", - ), - ) + with envs.SGLANG_ENABLE_ASYNC_ASSERT.override(0): + run_combined_tests( + models=variants, + test_name="Mistral-Large-3", + accuracy_params=AccuracyTestParams( + dataset="gsm8k", baseline_accuracy=0.85 + ), + performance_params=PerformanceTestParams( + profile_dir="performance_profiles_mistral_large3", + ), + ) if __name__ == "__main__": From d3cba9109d66d320f6caaae7e57f81b8945bd923 Mon Sep 17 00:00:00 2001 From: Yuwei An Date: Thu, 6 Aug 2026 18:48:47 -0700 Subject: [PATCH 07/47] fix(gdn): skip the -1 padding sentinel in the chunked extend kernel (#33810) --- python/sglang/kernels/ops/attention/fla/chunk_delta_h.py | 7 +++++-- test/registered/8-gpu-models/test_qwen35.py | 2 +- 2 files changed, 6 insertions(+), 3 deletions(-) diff --git a/python/sglang/kernels/ops/attention/fla/chunk_delta_h.py b/python/sglang/kernels/ops/attention/fla/chunk_delta_h.py index 7f79c6be6d15..2fe5e623d4ec 100644 --- a/python/sglang/kernels/ops/attention/fla/chunk_delta_h.py +++ b/python/sglang/kernels/ops/attention/fla/chunk_delta_h.py @@ -119,6 +119,9 @@ def chunk_gated_delta_rule_fwd_kernel_h_blockdim64( # per-slot pitch spans ALL layers' state, not H*V*K. int64: envelope pitches # overflow an int32 index product. index = tl.load(initial_state_indices + i_n).to(tl.int64) + # Padded rows carry the -1 sentinel; the decode kernel guards on it + # (fused_recurrent.py), the chunked extend path did not. + valid_state = index >= 0 h0 = initial_state + index * stride_init_state ht = initial_state + index * stride_init_state if USE_INITIAL_STATE: @@ -127,7 +130,7 @@ def chunk_gated_delta_rule_fwd_kernel_h_blockdim64( ht = ht + i_h * V * K # load initial state - if USE_INITIAL_STATE: + if USE_INITIAL_STATE and valid_state: p_h0_1 = tl.make_block_ptr(h0, (V, K), (K, 1), (i_v * BV, 0), (BV, 64), (1, 0)) b_h1 += tl.load(p_h0_1, boundary_check=(0, 1)).to(tl.float32) if K > 64: @@ -290,7 +293,7 @@ def chunk_gated_delta_rule_fwd_kernel_h_blockdim64( b_h4 += tl.trans(tl.dot(b_k, b_v)) # epilogue - if INPLACE_UPDATE: + if INPLACE_UPDATE and valid_state: p_ht = tl.make_block_ptr(ht, (V, K), (K, 1), (i_v * BV, 0), (BV, 64), (1, 0)) tl.store(p_ht, b_h1.to(p_ht.dtype.element_ty), boundary_check=(0, 1)) if K > 64: diff --git a/test/registered/8-gpu-models/test_qwen35.py b/test/registered/8-gpu-models/test_qwen35.py index f8e48e75727d..bf7fb2d01e12 100644 --- a/test/registered/8-gpu-models/test_qwen35.py +++ b/test/registered/8-gpu-models/test_qwen35.py @@ -30,7 +30,7 @@ def test_qwen35(self): "--tool-call-parser=qwen3_coder", "--mem-fraction-static=0.8", ] - dp_args = ["--dp=8", "--enable-dp-attention", "--disable-prefill-cuda-graph"] + dp_args = ["--dp=8", "--enable-dp-attention"] mtp_args = [ "--speculative-algorithm=EAGLE", "--speculative-num-steps=3", From 9e3287ea0b630421d182ba60d160f6fdee8b15cf Mon Sep 17 00:00:00 2001 From: Kangyan-Zhou Date: Fri, 7 Aug 2026 13:54:58 -0700 Subject: [PATCH 08/47] [Cherry-pick to release/v0.5.17] docker: add Kimi K3 artifacts and build hpc-ops with C++20 (#33956) (#34028) Co-authored-by: Baizhou Zhang --- docker/Dockerfile | 52 ++++++++++++++++++++++++++++++++++++++++++++--- 1 file changed, 49 insertions(+), 3 deletions(-) diff --git a/docker/Dockerfile b/docker/Dockerfile index 710e65570cd4..fb70db6d8c0d 100644 --- a/docker/Dockerfile +++ b/docker/Dockerfile @@ -20,6 +20,9 @@ ARG UBUNTU_MIRROR ARG GITHUB_ARTIFACTORY=github.com ARG INSTALL_FLASHINFER_JIT_CACHE=0 ARG FLASHINFER_VERSION=0.6.15.post1 +ARG TRTLLM_GEN_MOE_CUBIN_URL="https://github.com/sgl-project/whl/releases/download/trtllm_gen_moe_cubin_20260617/trtllm_gen_moe_cubin_pool_20260617_v0613rc1.zip" +ARG TRTLLM_GEN_MOE_CUBIN_SHA256="4900501cbe782a76b08a5858f9f07152287b97cb68114466dac286366b66c192" +ARG TRTLLM_GEN_MOE_CUBIN_ARCHIVE_ROOT="trtllm_gen_moe_cubin_pool_20260617_v0613rc1" ARG MOONCAKE_VERSION=0.3.12.post1 ARG MSCCLPP_VERSION=sglang-v0.9.1 #if need other arg please add in MOONCAKE_COMPILE_ARG @@ -28,7 +31,8 @@ ARG MOONCAKE_COMPILE_ARG="-DUSE_HTTP=ON -DUSE_MNNVL=ON -DUSE_CUDA=ON -DWITH_EP=O ENV DEBIAN_FRONTEND=noninteractive \ CUDA_HOME=/usr/local/cuda \ GDRCOPY_HOME=/usr/src/gdrdrv-${GDRCOPY_VERSION}/ \ - FLASHINFER_VERSION=${FLASHINFER_VERSION} + FLASHINFER_VERSION=${FLASHINFER_VERSION} \ + SGLANG_TRTLLM_GEN_MOE_CUBIN_POOL=/opt/trtllm_gen_moe_cubin_pool # Add GKE default lib and bin locations ENV PATH="${PATH}:/usr/local/nvidia/bin" \ @@ -75,6 +79,7 @@ RUN --mount=type=cache,target=/var/cache/apt,id=base-apt \ build-essential \ cmake \ perl \ + patch \ patchelf \ ccache \ git-lfs \ @@ -166,6 +171,7 @@ ENV LANG=en_US.UTF-8 \ # | # +-- devtools_builder (independent) # +-- gateway_builder (independent, only needs gateway source) +# +-- trtllm_cubin_builder (independent) # | # v # framework (combines all artifacts) @@ -342,7 +348,7 @@ FROM torch_deps AS hpc_ops_builder # HPC-Ops (https://github.com/Tencent/hpc-ops, MIT): fused attention / MoE / # RoPE kernels from the Tencent Hunyuan AI Infra team, consumed by the opt-in # hpc_ops attention and MoE runner backends. -ARG HPC_OPS_COMMIT=6e2ecede9d6d47b2e680e839cc7ad7422bc8d88b +ARG HPC_OPS_COMMIT=ab1a402724635507037426068f6cddc3d30dc0a8 WORKDIR /build @@ -466,6 +472,27 @@ RUN --mount=type=cache,target=/root/.cache/pip \ && cp target/release/sgl-model-gateway /build/sgl-model-gateway-bin \ && rm -rf /root/.cargo /root/.rustup /build/sgl-model-gateway/target /build/sgl-model-gateway/bindings/python/target +######################################################## +# PARALLEL STAGE 6: TRT-LLM Generated-MoE Cubin Pool +######################################################## +FROM base AS trtllm_cubin_builder + +RUN cubin_archive="/tmp/trtllm_gen_moe_cubin_pool.zip" && \ + cubin_extract_dir="/tmp/trtllm_gen_moe_cubin_extract" && \ + wget --no-verbose --output-document="${cubin_archive}" \ + "${TRTLLM_GEN_MOE_CUBIN_URL}" && \ + echo "${TRTLLM_GEN_MOE_CUBIN_SHA256} ${cubin_archive}" | \ + sha256sum --check --strict - && \ + mkdir -p "${cubin_extract_dir}" && \ + unzip -q "${cubin_archive}" -d "${cubin_extract_dir}" && \ + test ! -e "${SGLANG_TRTLLM_GEN_MOE_CUBIN_POOL}" && \ + mv "${cubin_extract_dir}/${TRTLLM_GEN_MOE_CUBIN_ARCHIVE_ROOT}" \ + "${SGLANG_TRTLLM_GEN_MOE_CUBIN_POOL}" && \ + test "$(find "${SGLANG_TRTLLM_GEN_MOE_CUBIN_POOL}" \ + -type f -name '*.cubin' | wc -l)" -eq 1696 && \ + rm -f "${cubin_archive}" && \ + rm -rf "${cubin_extract_dir}" + ######################################################## ########## Final Framework Image ###################### ######################################################## @@ -507,6 +534,21 @@ RUN --mount=type=cache,target=/root/.cache/pip \ # Copy flashinfer cubin (always) and jit-cache (if installed) packages COPY --from=flashinfer_cache /flashinfer_jit_output/ /usr/local/lib/python3.12/dist-packages/ +# Apply the FlashInfer CuTeDSL MLA decode-context-parallel runtime patch. +# Exclude tests because they are not included in the installed wheel. +COPY docker/kimi_k3/flashinfer-perkz-dcp-0.6.15.txt /tmp/flashinfer-perkz-dcp-0.6.15.txt +RUN FLASHINFER_DCP_PATCH=/tmp/flashinfer-perkz-dcp-0.6.15.txt && \ + FLASHINFER_SITE_PACKAGES="$(python3 -c 'from pathlib import Path; import flashinfer; print(Path(flashinfer.__file__).resolve().parent.parent)')" && \ + sed '/^diff --git a\/tests\//,$d' "${FLASHINFER_DCP_PATCH}" | \ + patch --dry-run --batch --forward --strip=1 --directory="${FLASHINFER_SITE_PACKAGES}" && \ + sed '/^diff --git a\/tests\//,$d' "${FLASHINFER_DCP_PATCH}" | \ + patch --batch --forward --strip=1 --directory="${FLASHINFER_SITE_PACKAGES}" && \ + rm -f "${FLASHINFER_DCP_PATCH}" && \ + rm -rf /root/.cache/flashinfer /root/.cache/pip + +# Copy the pinned FlashInfer MXFP4 MoE runner cubin pool +COPY --from=trtllm_cubin_builder /opt/trtllm_gen_moe_cubin_pool /opt/trtllm_gen_moe_cubin_pool + # Copy dev tools COPY --from=devtools_builder /tools/diff-so-fancy /usr/local/bin/ COPY --from=devtools_builder /tools/clang-format /usr/local/bin/ @@ -771,7 +813,8 @@ ARG GDRCOPY_VERSION=2.5.1 ENV DEBIAN_FRONTEND=noninteractive \ CUDA_HOME=/usr/local/cuda \ - GDRCOPY_HOME=/usr/src/gdrdrv-${GDRCOPY_VERSION}/ + GDRCOPY_HOME=/usr/src/gdrdrv-${GDRCOPY_VERSION}/ \ + SGLANG_TRTLLM_GEN_MOE_CUBIN_POOL=/opt/trtllm_gen_moe_cubin_pool # Add GKE default lib and bin locations + CUDA compiler paths for FlashInfer JIT ENV PATH="${PATH}:/usr/local/nvidia/bin:/usr/local/cuda/bin:/usr/local/cuda/nvvm/bin" \ @@ -857,6 +900,9 @@ RUN --mount=type=cache,target=/var/cache/apt,id=runtime-apt \ # Copy Python site-packages from framework (already cleaned of __pycache__/tests/pyc files) COPY --from=framework_final /usr/local/lib/python3.12/dist-packages /usr/local/lib/python3.12/dist-packages +# Copy the pinned FlashInfer MXFP4 MoE runner cubin pool +COPY --from=framework_final /opt/trtllm_gen_moe_cubin_pool /opt/trtllm_gen_moe_cubin_pool + # Copy SGLang workspace COPY --from=framework_final /sgl-workspace /sgl-workspace From 8af8fbaf895d408d34589b0b0ee9be816e613d4a Mon Sep 17 00:00:00 2001 From: Kangyan-Zhou Date: Fri, 7 Aug 2026 14:40:24 -0700 Subject: [PATCH 09/47] [Cherry-pick to release/v0.5.17] Fix MXFP4 scale placeholder initialization (#33500) (#34032) Co-authored-by: weireweire Co-authored-by: weireweire <20922698+weireweire@users.noreply.github.com> --- .../sglang/srt/layers/quantization/mxfp4.py | 26 +++++++++++++------ 1 file changed, 18 insertions(+), 8 deletions(-) diff --git a/python/sglang/srt/layers/quantization/mxfp4.py b/python/sglang/srt/layers/quantization/mxfp4.py index c9c8b6ae57db..131bec476da4 100644 --- a/python/sglang/srt/layers/quantization/mxfp4.py +++ b/python/sglang/srt/layers/quantization/mxfp4.py @@ -70,6 +70,10 @@ has_triton_kernels = is_triton_kernels_available() +# Serialized MXFP4 scales use raw UE8M0 bytes. Keep fresh parameters valid for +# post-load transforms and dummy initialization; 127 is the neutral scale (1.0). +_UE8M0_ONE = 127 + if is_flashinfer_available(): from flashinfer import ( @@ -473,10 +477,13 @@ def create_weights( set_weight_attrs(w13_weight, extra_weight_attrs) w13_weight_scale = torch.nn.Parameter( - torch.zeros( - layer.num_local_experts, - 2 * intermediate_size_per_partition_after_pad, - hidden_size // mxfp4_block, + torch.full( + ( + layer.num_local_experts, + 2 * intermediate_size_per_partition_after_pad, + hidden_size // mxfp4_block, + ), + fill_value=_UE8M0_ONE, dtype=scale_dtype, ), requires_grad=False, @@ -512,10 +519,13 @@ def create_weights( set_weight_attrs(w2_weight, extra_weight_attrs) w2_weight_scale = torch.nn.Parameter( - torch.zeros( - layer.num_local_experts, - hidden_size, - intermediate_size_per_partition_after_pad // mxfp4_block, + torch.full( + ( + layer.num_local_experts, + hidden_size, + intermediate_size_per_partition_after_pad // mxfp4_block, + ), + fill_value=_UE8M0_ONE, dtype=scale_dtype, ), requires_grad=False, From 4c06614809d63404c47f7315cd258686b4b02df4 Mon Sep 17 00:00:00 2001 From: Kangyan-Zhou Date: Fri, 7 Aug 2026 14:41:05 -0700 Subject: [PATCH 10/47] [Cherry-pick to release/v0.5.17] [AMD] Fuse Kimi-K3 attn-residual aggregation (#33599) (#34033) Co-authored-by: Xinyi Song --- .../kernels/ops/kimi_k3/attn_res_hip.py | 212 ++++++++++++++++++ python/sglang/srt/layers/attn_residual.py | 84 ++++++- 2 files changed, 295 insertions(+), 1 deletion(-) create mode 100644 python/sglang/kernels/ops/kimi_k3/attn_res_hip.py diff --git a/python/sglang/kernels/ops/kimi_k3/attn_res_hip.py b/python/sglang/kernels/ops/kimi_k3/attn_res_hip.py new file mode 100644 index 000000000000..7b6c8d470cd3 --- /dev/null +++ b/python/sglang/kernels/ops/kimi_k3/attn_res_hip.py @@ -0,0 +1,212 @@ +"""Triton attention-residual aggregation for Kimi-K3 on ROCm. + +The HIP counterpart of attn_res.py: same aggregation point (score the bank rows +against the current prefix, softmax, weighted sum, output RMSNorm), one launch, +but built for a GPU with no TMA and no tcgen05. See _agg_kernel for why the +shape differs so much from the SM100 kernel's. +""" + +from __future__ import annotations + +from functools import cache +from typing import Optional + +import torch +import triton +import triton.language as tl + +# _agg_kernel keeps a next_pow2(nvb) x next_pow2(H) fp32 tile in registers, and +# that is the whole basis of its speed. Past this budget it spills and loses to +# the 2-kernel Triton pipeline it replaces: measured at T=4 on MI355X, +# next_pow2(nvb)=8 is 2.0x faster, 16 is 1.2x, 32 is 0.2x. K3 sits right at the +# limit with H=7168 (one masked tile of 8192) and nvb <= 8. +MAX_REGISTER_TILE: int = 8 * 8192 + + +@cache +def supports_attn_res_hip(hidden_size: int, nvb: int) -> bool: + """Whether this shape fits the register budget. Callers must additionally + check that they are on ROCm; this is only the shape constraint.""" + return _tile_size(hidden_size, nvb) <= MAX_REGISTER_TILE + + +def _tile_size(hidden_size: int, nvb: int) -> int: + return triton.next_power_of_2(max(nvb, 1)) * triton.next_power_of_2(hidden_size) + + +@triton.jit +def _agg_kernel( + prefix_ptr, # [T, H] + addend_ptr, # [T, H]; the pending residual, or prefix_ptr when not HAS_ADD + prefix_out_ptr, # [T, H]; materialized prefix, written when HAS_ADD + bank_ptr, # [T, NB, H] + cw_ptr, # [H] fp32; score_norm weight * score_proj weight + ow_ptr, # [H]; out RMSNorm weight, unread when not APPLY_OUT_NORM + out_ptr, # [T, H] + score_eps, + out_eps, + stride_pm, + stride_am, + stride_om, + stride_bm, + stride_bb, + stride_o, + H: tl.constexpr, + BLOCK_H: tl.constexpr, + NVB: tl.constexpr, + R_PAD: tl.constexpr, + HAS_ADD: tl.constexpr, + WRITE_BANK: tl.constexpr, + APPLY_OUT_NORM: tl.constexpr, +): + """One CTA per token: score the NVB+1 rows, softmax, mix, apply the output + RMSNorm, all in one launch. + + T is the decode batch size, so this runs a handful of CTAs on a 256-CU GPU + and what binds is the load latency and bandwidth of a *single* CU. That is + what dictates the shape, and why it is not a port of the SM100 kernel. That + one's online softmax reads each row once but chains one block-wide reduction + per row; here each link in that chain costs a full HBM round-trip — + measured 1.1us/row, dead linear in NVB, because there is no TMA pipeline to + hide it behind. Scoring the rows as one [R_PAD, BLOCK_H] tile instead puts + every row's load in flight at once, and keeping that tile in registers lets + the mix reuse it rather than re-reading the bank, which at one active CU is + the difference between ~9us and ~12.5us at NVB=8. + + Holding the tile is why NVB is a constexpr and why MAX_REGISTER_TILE caps + the shape: the register budget is what this trades away. + + The prefix row streams through anyway, so the pending residual add and the + bank snapshot ride along for free — no other program reads bank row NVB, so + those stores need no synchronization. + + Taking the global max before exponentiating makes this bit-comparable to + the 2-kernel pipeline rather than to the SM100 kernel. + """ + t = tl.program_id(0) + offs = tl.arange(0, BLOCK_H) + mask = offs < H + + # The prefix row is score row NVB, and the only row that needs writing back. + row = tl.load(prefix_ptr + t * stride_pm + offs, mask=mask, other=0.0) + if HAS_ADD: + # Round to the storage dtype before scoring: downstream readers see + # these bits, so the score has to as well. + row = ( + row.to(tl.float32) + + tl.load(addend_ptr + t * stride_am + offs, mask=mask, other=0.0).to( + tl.float32 + ) + ).to(prefix_out_ptr.dtype.element_ty) + tl.store(prefix_out_ptr + t * stride_om + offs, row, mask=mask) + if WRITE_BANK: + tl.store(bank_ptr + t * stride_bm + NVB * stride_bb + offs, row, mask=mask) + pv = row.to(tl.float32) + + cw = tl.load(cw_ptr + offs, mask=mask, other=0.0) + p_score = tl.sum(pv * cw) / tl.sqrt(tl.sum(pv * pv) / H + score_eps) + + # The whole bank, in registers: every row's load is in flight at once, and + # the mix below reuses it instead of going back to HBM. + offs_r = tl.arange(0, R_PAD) + mask_r = offs_r < NVB + tile = tl.load( + bank_ptr + t * stride_bm + offs_r[:, None] * stride_bb + offs[None, :], + mask=mask_r[:, None] & mask[None, :], + other=0.0, + ).to(tl.float32) + b_score = tl.sum(tile * cw[None, :], axis=1) / tl.sqrt( + tl.sum(tile * tile, axis=1) / H + score_eps + ) + + m = tl.maximum(tl.max(tl.where(mask_r, b_score, -float("inf"))), p_score) + b_w = tl.where(mask_r, tl.exp(b_score - m), 0.0) + p_w = tl.exp(p_score - m) + inv = 1.0 / (tl.sum(b_w) + p_w) + + acc = pv * (p_w * inv) + tl.sum(b_w[:, None] * inv * tile, axis=0) + + if APPLY_OUT_NORM: + scale = 1.0 / tl.sqrt(tl.sum(acc * acc) / H + out_eps) + ow = tl.load(ow_ptr + offs, mask=mask, other=0.0).to(tl.float32) + acc = acc * scale * ow + tl.store(out_ptr + t * stride_o + offs, acc.to(out_ptr.dtype.element_ty), mask=mask) + + +def attn_res_hip( + prefix_sum: torch.Tensor, + bank: torch.Tensor, + cw: torch.Tensor, + ow: Optional[torch.Tensor], + out: torch.Tensor, + nvb: int, + score_eps: float, + out_eps: float, + *, + addend: Optional[torch.Tensor] = None, + prefix_out: Optional[torch.Tensor] = None, + write_prefix: bool = False, +) -> None: + """Single-kernel attention-residual aggregation for ROCm. + + Restrictions: nvb >= 1, and the shape must pass supports_attn_res_hip(). + + Parameters + ---------- + prefix_sum : [T, H] bf16 — the running prefix, or its first term if addend + is given + bank : [T, NB, H] bf16 (rows 0..nvb-1 are aggregated) + cw : [H] fp32 — precomputed score_norm weight * proj weight + ow : [H] output RMSNorm weight, or None to return the pre-norm + softmax mixture (the aggregate-stream value) + out : [T, H] bf16 output buffer + nvb : number of valid bank rows (>= 1) + score_eps, out_eps : RMSNorm epsilons; unlike the SM100 kernel these need + not be equal + addend : fold a pending residual add in, so the aggregated prefix is + prefix_sum + addend; requires prefix_out + prefix_out : [T, H] bf16 buffer receiving that materialized prefix + write_prefix : also snapshot the prefix row into bank[:, nvb, :] (bit-exact + copy, fused into the score pass which already has the row in + registers); requires NB > nvb + """ + T, H = prefix_sum.shape + assert nvb >= 1, "nvb == 0 has nothing to aggregate; the caller must handle it" + assert supports_attn_res_hip(H, nvb), ( + f"attn_res_hip: register tile {_tile_size(H, nvb)} exceeds " + f"{MAX_REGISTER_TILE} (H={H}, nvb={nvb})" + ) + has_add = addend is not None + assert not has_add or prefix_out is not None, "addend requires prefix_out" + + # Triton needs a real pointer for every argument; the flags decide whether + # these are ever dereferenced. + addend_arg = addend if has_add else prefix_sum + prefix_out_arg = prefix_out if has_add else prefix_sum + ow_arg = ow if ow is not None else cw + + _agg_kernel[(T,)]( + prefix_sum, + addend_arg, + prefix_out_arg, + bank, + cw, + ow_arg, + out, + score_eps, + out_eps, + prefix_sum.stride(0), + addend_arg.stride(0), + prefix_out_arg.stride(0), + bank.stride(0), + bank.stride(1), + out.stride(0), + H=H, + BLOCK_H=triton.next_power_of_2(H), + NVB=nvb, + R_PAD=triton.next_power_of_2(nvb), + HAS_ADD=has_add, + WRITE_BANK=write_prefix, + APPLY_OUT_NORM=ow is not None, + num_warps=4, + ) diff --git a/python/sglang/srt/layers/attn_residual.py b/python/sglang/srt/layers/attn_residual.py index 9db79ca8bb51..491a7b4f8a88 100644 --- a/python/sglang/srt/layers/attn_residual.py +++ b/python/sglang/srt/layers/attn_residual.py @@ -9,6 +9,8 @@ # online-softmax consumers over a double-buffered chunk ring, out # norm fused, per-nvb tuned launch config, one persistent CTA per # SM. Taken on SM100+ with H=7168. +# hip — single Triton kernel, everything in one launch; taken on ROCm +# within its register budget. # fused — Triton 2-kernel pipeline with full H-parallelism; the fallback # everywhere the fast kernel does not apply. # aggregate_stream_torch is the eager reference (tests and the @@ -22,11 +24,13 @@ from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ReplicatedLinear +from sglang.srt.utils import is_hip _BLOCK_H: int = 1024 # H = 7168 = 7 x 1024 _MAX_ROWS: int = 16 # next_pow2(8 + 1), K3 has <= 8 snapshots _FAST_SUPPORTED = None +_HIP_SHAPE_GATE = None def _use_fast(hidden_size: int) -> bool: @@ -39,6 +43,19 @@ def _use_fast(hidden_size: int) -> bool: return _FAST_SUPPORTED and hidden_size == 7168 +def _use_hip_fused(hidden_size: int, nvb: int) -> bool: + """This gate picks the single-kernel ROCm Triton kernel instead of the + 2-kernel pipeline.""" + if not is_hip(): + return False + global _HIP_SHAPE_GATE + if _HIP_SHAPE_GATE is None: + from sglang.kernels.ops.kimi_k3.attn_res_hip import supports_attn_res_hip + + _HIP_SHAPE_GATE = supports_attn_res_hip + return _HIP_SHAPE_GATE(hidden_size, nvb) + + def get_cw( proj: ReplicatedLinear, norm: RMSNorm, @@ -244,6 +261,41 @@ def _aggregate_fused( return out_norm(_mix_fused(prefix_sum, bank, nvb, score_proj, score_norm)) +def _aggregate_hip( + prefix_sum: torch.Tensor, + addend: Optional[torch.Tensor], + bank: torch.Tensor, + nvb: int, + score_proj: ReplicatedLinear, + score_norm: RMSNorm, + out_norm: Optional[RMSNorm], + write_bank_row: bool = False, +) -> tuple[torch.Tensor, torch.Tensor]: + """Single ROCm Triton kernel: the bank stays in registers so scoring and + mixing share one read, and the pending residual add, the bank snapshot and + the output RMSNorm all fold into the same launch. out_norm None returns the + pre-norm mixture instead. Returns (result, prefix).""" + from sglang.kernels.ops.kimi_k3.attn_res_hip import attn_res_hip + + cw = get_cw(score_proj, score_norm) + prefix = prefix_sum if addend is None else torch.empty_like(prefix_sum) + out = torch.empty_like(prefix_sum) + attn_res_hip( + prefix_sum, + bank, + cw, + out_norm.weight if out_norm is not None else None, + out, + nvb, + score_norm.variance_epsilon, + out_norm.variance_epsilon if out_norm is not None else 0.0, + addend=addend, + prefix_out=prefix, + write_prefix=write_bank_row, + ) + return out, prefix + + def aggregate_stream_torch( prefix_sum: torch.Tensor, bank: torch.Tensor, @@ -277,6 +329,10 @@ def aggregate_stream( raw wire only carries the current block's running prefix.""" if nvb == 0: return prefix_sum + if _use_hip_fused(prefix_sum.shape[1], nvb): + return _aggregate_hip( + prefix_sum, None, bank, nvb, score_proj, score_norm, None + )[0] if prefix_sum.shape[1] % _BLOCK_H != 0: return aggregate_stream_torch(prefix_sum, bank, nvb, score_proj, score_norm) return _mix_fused(prefix_sum, bank, nvb, score_proj, score_norm) @@ -295,6 +351,18 @@ def _aggregate_fused_add( """Aggregation point with a pending upstream residual add: materialize prefix = prefix_a + prefix_b, then aggregate. Returns (normed, prefix). write_bank_row rides _aggregate (fast path only).""" + if _use_hip_fused(prefix_a.shape[1], nvb): + # The hip kernel reads the prefix row anyway, so the add folds into it. + return _aggregate_hip( + prefix_a, + prefix_b, + bank, + nvb, + score_proj, + score_norm, + out_norm, + write_bank_row=write_bank_row, + ) prefix = prefix_a + prefix_b return ( _aggregate( @@ -336,6 +404,17 @@ def _aggregate( out_norm, write_bank_row=write_bank_row, ) + if _use_hip_fused(prefix_sum.shape[1], nvb): + return _aggregate_hip( + prefix_sum, + None, + bank, + nvb, + score_proj, + score_norm, + out_norm, + write_bank_row=write_bank_row, + )[0] assert not write_bank_row, "fused bank write is fast-path only" return _aggregate_fused(prefix_sum, bank, nvb, score_proj, score_norm, out_norm) @@ -403,7 +482,10 @@ def forward( self.block_residual if rows is None else self.block_residual[rows] ) - fused_write = write and _use_fast(hidden_states.shape[1]) + fused_write = write and ( + _use_fast(hidden_states.shape[1]) + or _use_hip_fused(hidden_states.shape[1], nvb) + ) if prefix_sum is None: # hidden_states already is the whole head (PP entry or a # block-boundary restart). From 58be7dbdb635487c859729e230cc9749be0f467e Mon Sep 17 00:00:00 2001 From: Kangyan-Zhou Date: Fri, 7 Aug 2026 14:50:34 -0700 Subject: [PATCH 11/47] [Cherry-pick to release/v0.5.17] [Kimi-K3] Allow DSPARK verify on cutedsl_mla (fold_sq) (#33650) (#34034) Co-authored-by: Yuhao Yang <47235274+yhyang201@users.noreply.github.com> --- python/sglang/srt/arg_groups/overrides.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index 245f5943ca5f..5a56b19c078d 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -327,8 +327,10 @@ def _dspark_verify_on_decode_backend( if backend == "tokenspeed_mla": return kv_cache_dtype == "fp8_e4m3" and q_len <= 8 if backend == "cutedsl_mla": - # The cute-dsl kernel rejects q_len >= 5 with no fallback. - return q_len <= 4 + # cute-dsl monolithic MLA decode folds the verify tokens into the head + # dim (fold_sq), so it serves any DSPARK verify width. Needs flashinfer + # >= 0.6.15 (older builds reject q_len >= 5). + return True return False From 29481685462732237d80d86076d6563e1f658102 Mon Sep 17 00:00:00 2001 From: Kangyan-Zhou Date: Fri, 7 Aug 2026 14:50:47 -0700 Subject: [PATCH 12/47] [Cherry-pick to release/v0.5.17] fix(PP): size the mamba pool per pipeline stage, not per whole model (#33666) (#34035) Co-authored-by: YAMY <74099316+YAMY1234@users.noreply.github.com> --- .../srt/mem_cache/kv_cache_configurator.py | 35 ++++++-- .../test_mamba_donated_alloc_ratio.py | 79 +++++++++++++++++++ 2 files changed, 108 insertions(+), 6 deletions(-) diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index 221c73dccdfb..de9efbf20227 100644 --- a/python/sglang/srt/mem_cache/kv_cache_configurator.py +++ b/python/sglang/srt/mem_cache/kv_cache_configurator.py @@ -24,6 +24,7 @@ is_minimax_sparse, ) from sglang.srt.distributed.parallel_state import get_world_group +from sglang.srt.distributed.utils import get_pp_indices from sglang.srt.environ import envs from sglang.srt.layers.quantization.fp4_kv_cache_quant_method import ( get_kv_cache_quant_method, @@ -1817,6 +1818,27 @@ def _handle_max_mamba_cache(self, total_rest_memory): server_args = self.server_args assert config is not None + # mamba_cache_per_req covers every mamba layer, but under PP a rank only + # allocates its own [start_layer, end_layer) slice. Charge the largest + # per-stage share so every rank derives the same pool without a collective. + all_mamba_layers = config.mamba2_cache_params.layers + if self.ps.pp_size > 1 and all_mamba_layers: + max_stage_mamba_layers = max( + sum(1 for i in all_mamba_layers if start <= i < end) + for start, end in ( + get_pp_indices( + self.model_config.num_hidden_layers, rank, self.ps.pp_size + ) + for rank in range(self.ps.pp_size) + ) + ) + else: + max_stage_mamba_layers = len(all_mamba_layers) + pp_layer_scale = max_stage_mamba_layers / max(len(all_mamba_layers), 1) + stage_per_req = int( + config.mamba2_cache_params.mamba_cache_per_req * pp_layer_scale + ) + has_spec_dec = not self.spec_algorithm.is_none() # ReplaySSM drops the per-step intermediate_ssm scratch, so its mamba budget # no longer reserves the (1 + D/ratio) intermediate factor -- the whole @@ -1844,6 +1866,7 @@ def _handle_max_mamba_cache(self, total_rest_memory): ) else: replayssm_ring_per_req = 0 + replayssm_ring_per_req = int(replayssm_ring_per_req * pp_layer_scale) if has_spec_dec: assert get_spec().speculative_num_draft_tokens is not None assert get_schedule().max_running_requests is not None @@ -1865,7 +1888,7 @@ def _handle_max_mamba_cache(self, total_rest_memory): get_schedule().max_mamba_cache_size // ratio, ) intermediate_size = ( - config.mamba2_cache_params.mamba_cache_per_req + stage_per_req * (capped_reqs + 1) * get_spec().speculative_num_draft_tokens ) @@ -1884,15 +1907,15 @@ def _handle_max_mamba_cache(self, total_rest_memory): # pool's padding slot). Skipped under replayssm. if has_spec_dec and not replayssm_active: intermediate_size = ( - config.mamba2_cache_params.mamba_cache_per_req + stage_per_req * (get_schedule().max_mamba_cache_size + 1) * get_spec().speculative_num_draft_tokens ) total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30)) else: # Use ratio-based calculation to auto-fit available memory - assert config.mamba2_cache_params.mamba_cache_per_req > 0 - per_req = config.mamba2_cache_params.mamba_cache_per_req + assert stage_per_req > 0 + per_req = stage_per_req # Solve jointly for max_mamba_cache_size (K), including the pool's # +1 padding slot on both buffers (see memory_pool.py): @@ -1941,7 +1964,7 @@ def _handle_max_mamba_cache(self, total_rest_memory): f"Not enough GPU memory for hybrid (mamba/linear-attention) state cache. " f"Computed max_mamba_cache_size={get_schedule().max_mamba_cache_size} " f"(total_rest_memory={total_rest_memory:.2f} GB, " - f"mamba_cache_per_req={config.mamba2_cache_params.mamba_cache_per_req / (1 << 20):.2f} MB). " + f"mamba_cache_per_req={stage_per_req / (1 << 20):.2f} MB). " f"Try: (1) reduce --max-running-requests, " f"(2) increase --mem-fraction-static, " f"(3) reduce --speculative-num-draft-tokens, or " @@ -1953,7 +1976,7 @@ def _handle_max_mamba_cache(self, total_rest_memory): # the ring is not allocated). mamba_state_memory = ( (get_schedule().max_mamba_cache_size + 1) - * (config.mamba2_cache_params.mamba_cache_per_req + replayssm_ring_per_req) + * (stage_per_req + replayssm_ring_per_req) / (1 << 30) ) return total_rest_memory - mamba_state_memory diff --git a/test/registered/unit/mem_cache/test_mamba_donated_alloc_ratio.py b/test/registered/unit/mem_cache/test_mamba_donated_alloc_ratio.py index 4f39375627ab..9b108af7a918 100644 --- a/test/registered/unit/mem_cache/test_mamba_donated_alloc_ratio.py +++ b/test/registered/unit/mem_cache/test_mamba_donated_alloc_ratio.py @@ -232,5 +232,84 @@ def test_decode_steady_evictable_prefix_ratio2_ok(self): self.assertEqual(len(cache.prefix_nodes), N - 1) +class TestPPMambaPoolSizing(unittest.TestCase): + """A PP rank only allocates mamba state for its own [start_layer, end_layer) + slice, so charging it for the whole model's layers starves the pool. Sizing + uses the largest per-stage share, which also keeps every rank on the same + pool size (and hence the same max_running_requests / pp_max_micro_batch_size) + without a collective.""" + + # Kimi-K3 shaped: 93 layers, linear attention everywhere except every 4th and + # the last, so the 69 mamba layers split unevenly over 8 stages (9 or 8 each). + TOTAL_LAYERS = 93 + MAMBA_LAYERS = [i for i in range(93) if (i + 1) % 4 != 0 and i <= 90] + BUDGET_GB = 8.0 + + @classmethod + def _pool_size(cls, pp_rank, pp_size): + from sglang.srt import runtime_context as rc + from sglang.srt.configs.mamba_utils import ( + Mamba2CacheParams, + Mamba2StateDType, + Mamba2StateShape, + ) + from sglang.srt.distributed.utils import get_pp_indices + from sglang.srt.mem_cache.kv_cache_configurator import KVCacheConfigurator + from sglang.srt.runtime_context import get_schedule + + shape = Mamba2StateShape( + conv=[(4096, 3)], + temporal=(64, 128, 128), + intermediate_size=0, + conv_dim=0, + ssm_state_size=0, + num_heads=0, + head_dim=0, + state_size=0, + conv_kernel=0, + num_k_heads_per_tp=8, + ) + params = Mamba2CacheParams( + shape=shape, + dtype=Mamba2StateDType(conv=torch.bfloat16, temporal=torch.float32), + layers=list(cls.MAMBA_LAYERS), + ) + start, end = get_pp_indices(cls.TOTAL_LAYERS, pp_rank, pp_size) + fake = SimpleNamespace( + mambaish_config=SimpleNamespace(mamba2_cache_params=params), + server_args=SimpleNamespace(), + spec_algorithm=SimpleNamespace(is_none=lambda: True), + layer_info=SimpleNamespace(start_layer=start, end_layer=end), + ps=SimpleNamespace(attn_dp_size=1, pp_size=pp_size), + hybrid_gdn_config=None, + model_config=SimpleNamespace( + hf_config=SimpleNamespace(), num_hidden_layers=cls.TOTAL_LAYERS + ), + ) + with rc.get_context().override_server_args( + disable_radix_cache=False, + max_mamba_cache_size=None, + max_running_requests=None, + mamba_full_memory_ratio=0.5, + enable_linear_replayssm_spec=False, + ): + KVCacheConfigurator._handle_max_mamba_cache(fake, cls.BUDGET_GB) + return get_schedule().max_mamba_cache_size + + def test_stage_is_not_charged_for_the_whole_model(self): + solo = self._pool_size(0, 1) + staged = self._pool_size(0, 8) + # The busiest stage holds 9 of the 69 mamba layers, so it should fit + # roughly 69/9 more slots than a rank holding all of them. pp_size=1 is + # unchanged: that rank does hold every layer. + self.assertGreater(staged, solo * 5) + + def test_every_stage_agrees_on_the_pool_size(self): + sizes = {self._pool_size(r, 8) for r in range(8)} + self.assertEqual( + len(sizes), 1, f"per-rank pool sizes diverged: {sorted(sizes)}" + ) + + if __name__ == "__main__": unittest.main() From f704fe29f9fa68ca8785607729545f5ab60c2566 Mon Sep 17 00:00:00 2001 From: Hank Han Date: Wed, 12 Aug 2026 15:27:10 +0800 Subject: [PATCH 13/47] feat(onion): install Onion for private delivery Keep the Onion runtime dependency independently reviewable and revertible. Signed-off-by: Hank Han --- docker/Dockerfile | 26 ++++++++++++++++++++++++++ 1 file changed, 26 insertions(+) diff --git a/docker/Dockerfile b/docker/Dockerfile index fb70db6d8c0d..adf8ea461b51 100644 --- a/docker/Dockerfile +++ b/docker/Dockerfile @@ -600,6 +600,19 @@ RUN --mount=type=cache,target=/var/cache/apt,id=framework-apt \ && apt install -y --no-install-recommends nsight-systems-cli \ && rm -rf /var/lib/apt/lists/* +# Install Onion model/data tooling. The onion-ai-data package provides oniond. +RUN --mount=type=cache,target=/var/cache/apt,id=framework-apt \ + OS="$(. /etc/os-release && echo "${ID}")" \ + && VERSION_CODENAME="$(. /etc/os-release && echo "${VERSION_CODENAME}")" \ + && curl --fail --show-error --location --retry 5 --retry-all-errors --retry-delay 2 "https://mirrors.ivolces.com/extra-tools/${OS}/GPG-KEY-system" \ + | gpg --dearmor -o /etc/apt/trusted.gpg.d/volc-extra-tools.gpg \ + && echo "deb [arch=$(dpkg --print-architecture) signed-by=/etc/apt/trusted.gpg.d/volc-extra-tools.gpg] http://mirrors.ivolces.com/extra-tools/${OS} ${VERSION_CODENAME} main" \ + > /etc/apt/sources.list.d/volctools.list \ + && apt-get update \ + && apt-get install -y --no-install-recommends onion-ai-data \ + && rm -rf /var/lib/apt/lists/* \ + && apt-get clean + # ============================================================================= # Python packages and tools (before source copy for better caching) # ============================================================================= @@ -877,6 +890,19 @@ RUN --mount=type=cache,target=/var/cache/apt,id=runtime-apt \ && rm -rf /var/lib/apt/lists/* \ && apt-get clean +# Install Onion model/data tooling. The onion-ai-data package provides oniond. +RUN --mount=type=cache,target=/var/cache/apt,id=runtime-apt \ + OS="$(. /etc/os-release && echo "${ID}")" \ + && VERSION_CODENAME="$(. /etc/os-release && echo "${VERSION_CODENAME}")" \ + && curl --fail --show-error --location --retry 5 --retry-all-errors --retry-delay 2 "https://mirrors.ivolces.com/extra-tools/${OS}/GPG-KEY-system" \ + | gpg --dearmor -o /etc/apt/trusted.gpg.d/volc-extra-tools.gpg \ + && echo "deb [arch=$(dpkg --print-architecture) signed-by=/etc/apt/trusted.gpg.d/volc-extra-tools.gpg] http://mirrors.ivolces.com/extra-tools/${OS} ${VERSION_CODENAME} main" \ + > /etc/apt/sources.list.d/volctools.list \ + && apt-get update \ + && apt-get install -y --no-install-recommends onion-ai-data \ + && rm -rf /var/lib/apt/lists/* \ + && apt-get clean + # Set up locale RUN apt-get update && apt-get install -y --no-install-recommends locales \ && locale-gen en_US.UTF-8 \ From 0188325e68b590001d5d64d9d7eeb9d8fec1f2b7 Mon Sep 17 00:00:00 2001 From: Hank Han Date: Wed, 12 Aug 2026 15:27:10 +0800 Subject: [PATCH 14/47] feat(eic): add private delivery integration Install the EIC SDK in the private runtime image and provide the v0.5.17-compatible deployment integration check as one feature. Signed-off-by: Hank Han --- docker/Dockerfile | 5 + scripts/eic_integration_check.py | 293 +++++++++++++++++++++++++++++++ 2 files changed, 298 insertions(+) create mode 100644 scripts/eic_integration_check.py diff --git a/docker/Dockerfile b/docker/Dockerfile index adf8ea461b51..d05af8f88bc8 100644 --- a/docker/Dockerfile +++ b/docker/Dockerfile @@ -665,6 +665,11 @@ RUN --mount=type=cache,target=/root/.cache/pip \ termplotlib \ "runai-model-streamer[s3,gcs,azure]>=0.15.7" +# Install the EIC SDK used by the EIC HiCache backend. +RUN python3 -m pip install --force-reinstall \ + https://eic-sdk-release.tos-cn-beijing.volces.com/python/eic-1.5.2-py3-none-any.whl \ + --no-cache-dir + # Per-CUDA-major package installs. The `nixl` stub package is needed (it owns # the `nixl` import path) but unconditionally requires nixl-cu12, so we install # it with --no-deps and pair it with the matching nixl-cu12 / nixl-cu13 binary diff --git a/scripts/eic_integration_check.py b/scripts/eic_integration_check.py new file mode 100644 index 000000000000..0fb5b3af7ca6 --- /dev/null +++ b/scripts/eic_integration_check.py @@ -0,0 +1,293 @@ +"""Run a post-deployment check against SGLang's v0.5.17 EIC backend. + +The script uses the same SDK calls and ``remote-eic.yaml`` fields as +``EICStorage``. It writes only uniquely prefixed temporary keys and removes +them on every normal or exceptional exit. + + python scripts/eic_integration_check.py + python scripts/eic_integration_check.py --page-bytes $((512*1024)) --flood-gib 8 +""" + +import argparse +import os +import time +import traceback +import uuid + +import eic +import torch +import yaml + +FAILURES = [] +PROBE = None + + +def check(name, ok, detail=""): + print(f"{'PASS' if ok else 'FAIL'} {name}{' ' + detail if detail else ''}") + if not ok: + FAILURES.append(name) + return ok + + +def _status_mask(outcome, count): + codes = list(getattr(outcome, "status_codes", ())) + if not codes: + return [False] * count + mask = [code == eic.StatusCode.SUCCESS for code in codes[:count]] + return mask + [False] * (count - len(mask)) + + +class Probe: + def __init__(self, config, page_bytes, dtype=torch.bfloat16): + numel = page_bytes // dtype.itemsize + if numel <= 0: + raise ValueError("--page-bytes is smaller than one tensor element") + self.shape = (numel,) + self.dtype = dtype + self.page_bytes = numel * dtype.itemsize + self.namespace = config.get("eic_namespace", "") + self.run_id = uuid.uuid4().hex[:12] + self.written = set() + self.connection = self._connect(config) + + @staticmethod + def _connect(config): + remote_url = config.get("remote_url") + if not isinstance(remote_url, str) or not remote_url.startswith("eic://"): + raise ValueError("remote_url must be an eic:// URL") + + log_dir = config.get("eic_log_dir") + if not log_dir: + raise ValueError("eic_log_dir is required") + os.makedirs(log_dir, exist_ok=True) + + init_option = eic.InitOption() + init_option.log_dir = log_dir + init_option.log_level = eic.LogLevel(config.get("eic_log_level", 2)) + init_option.transport_type = eic.TransportType(config.get("eic_trans_type", 3)) + init_option.flag_file = config.get("eic_flag_file") + + connection = eic.Client() + ret = connection.init( + config.get("eic_instance_id"), remote_url[len("eic://") :], init_option + ) + if ret != 0: + raise RuntimeError(f"EIC client initialization failed with code {ret}") + return connection + + def key(self, index): + return f"eic_integration_check/{self.run_id}/{index}" + + def page(self, seed): + generator = torch.Generator().manual_seed(seed) + return torch.randint( + -128, 127, self.shape, generator=generator, dtype=torch.int16 + ).to(self.dtype) + + @staticmethod + def _keys(keys): + result = eic.StringVector() + for key in keys: + result.append(key) + return result + + @staticmethod + def _buffers(pages): + result = eic.IOBuffers() + for page in pages: + result.append(page.data_ptr(), page.numel() * page.element_size(), False) + return result + + def write(self, keys, pages): + option = eic.SetOption() + option.ns = self.namespace + option.ttl_second = -1 + status, outcome = self.connection.mset( + self._keys(keys), self._buffers(pages), option + ) + mask = _status_mask(outcome, len(keys)) + if status == eic.StatusCode.SUCCESS and not mask: + mask = [True] * len(keys) + self.written.update(key for key, ok in zip(keys, mask) if ok) + return mask + + def read(self, keys): + pages = [torch.zeros(self.shape, dtype=self.dtype) for _ in keys] + option = eic.GetOption() + option.ns = self.namespace + status, _, outcome = self.connection.mget( + self._keys(keys), option, self._buffers(pages) + ) + mask = _status_mask(outcome, len(keys)) + if status == eic.StatusCode.SUCCESS and not mask: + mask = [True] * len(keys) + return pages, mask + + def exists(self, keys): + option = eic.ExistOption() + option.ns = self.namespace + status, outcome = self.connection.mexist(self._keys(keys), option) + mask = _status_mask(outcome, len(keys)) + if status == eic.StatusCode.SUCCESS and not mask: + mask = [True] * len(keys) + return mask + + def delete(self, keys): + option = eic.DelOption() + option.ns = self.namespace + status, outcome = self.connection.mdel(self._keys(keys), option) + mask = _status_mask(outcome, len(keys)) + if status == eic.StatusCode.SUCCESS and not mask: + mask = [True] * len(keys) + for key, ok in zip(keys, mask): + if ok: + self.written.discard(key) + return mask + + def cleanup(self): + keys = sorted(self.written) + for start in range(0, len(keys), 256): + self.delete(keys[start : start + 256]) + print(f"cleaned up {len(keys)} temporary keys") + + +def parse_args(): + parser = argparse.ArgumentParser() + parser.add_argument( + "--config", + default=os.environ.get( + "REMOTE_EIC_YAML", "/sgl-workspace/config/remote-eic.yaml" + ), + help="EIC YAML config; defaults to REMOTE_EIC_YAML or SGLang's standard path", + ) + parser.add_argument("--page-bytes", type=int, default=1 << 20) + parser.add_argument("--pages", type=int, default=32, help="pages per probe batch") + parser.add_argument( + "--flood-gib", + type=float, + default=0, + help="write this much temporary data, then check exists/get consistency", + ) + args = parser.parse_args() + if args.pages < 2: + parser.error("--pages must be at least 2") + if args.flood_gib < 0: + parser.error("--flood-gib must be non-negative") + return args + + +def main(): + args = parse_args() + config_path = os.path.abspath(args.config) + if not check("config file present", os.path.isfile(config_path), config_path): + return 1 + with open(config_path, encoding="utf-8") as config_file: + config = yaml.safe_load(config_file) or {} + + global PROBE + start = time.perf_counter() + probe = PROBE = Probe(config, args.page_bytes) + check( + "client init", + True, + f"ns={probe.namespace or ''} page={probe.page_bytes}B " + f"{time.perf_counter() - start:.1f}s", + ) + + count = args.pages + keys = [probe.key(index) for index in range(count)] + check("absent key reports miss", not any(probe.exists(keys))) + + pages = [probe.page(index) for index in range(count)] + start = time.perf_counter() + write_mask = probe.write(keys, pages) + elapsed = time.perf_counter() - start + check( + "write", + all(write_mask), + f"{count} pages {count * probe.page_bytes / elapsed / 2**30:.2f} GiB/s", + ) + check("exists after write", all(probe.exists(keys))) + + start = time.perf_counter() + values, read_mask = probe.read(keys) + elapsed = time.perf_counter() - start + check( + "read", + all(read_mask), + f"{count} pages {count * probe.page_bytes / elapsed / 2**30:.2f} GiB/s", + ) + mismatches = [ + index + for index, ok in enumerate(read_mask) + if ok and not torch.equal(values[index], pages[index]) + ] + check("read-back is byte-identical", not mismatches, f"mismatched={mismatches[:4]}") + + half = count // 2 + ghosts = [probe.key(f"ghost{index}") for index in range(half)] + _, mixed = probe.read(keys[:half] + ghosts) + check( + "mixed batch reports per-key hits", + mixed[:half] == [True] * half and mixed[half:] == [False] * len(ghosts), + f"hits={sum(mixed)}/{len(mixed)}", + ) + + replacement = probe.page(10**6) + probe.write(keys[:1], [replacement]) + values, mask = probe.read(keys[:1]) + check("overwrite wins", mask[0] and torch.equal(values[0], replacement)) + + victims = keys[:4] + delete_mask = probe.delete(victims) + _, mask = probe.read(victims) + check("delete removes the value", all(delete_mask) and not any(mask)) + + if args.flood_gib > 0: + flood_pages = int(args.flood_gib * 2**30 / probe.page_bytes) + start = time.perf_counter() + for offset in range(0, flood_pages, count): + batch = min(count, flood_pages - offset) + flood_keys = [probe.key(f"flood{offset + index}") for index in range(batch)] + flood_values = [ + probe.page(10**7 + offset + index) for index in range(batch) + ] + probe.write(flood_keys, flood_values) + elapsed = time.perf_counter() - start + print( + f"flooded {flood_pages} pages ({args.flood_gib} GiB) in {elapsed:.1f}s " + f"{flood_pages * probe.page_bytes / elapsed / 2**30:.2f} GiB/s" + ) + + survivors = keys[4:] + exists = probe.exists(survivors) + _, mask = probe.read(survivors) + phantoms = [ + key + for key, present, read in zip(survivors, exists, mask) + if present and not read + ] + check( + "eviction keeps exists and get consistent", + not phantoms, + f"survived={sum(mask)}/{len(survivors)} phantom={len(phantoms)}", + ) + + return 1 if FAILURES else 0 + + +if __name__ == "__main__": + exit_code = 1 + try: + exit_code = main() + except BaseException: + traceback.print_exc() + finally: + if PROBE is not None: + PROBE.cleanup() + print( + "\nall checks passed" + if exit_code == 0 + else f"\n{len(FAILURES) or 'aborted'} failed" + ) + raise SystemExit(exit_code) From 9e214b6704716ca53624a6fdeec775105a408a1a Mon Sep 17 00:00:00 2001 From: Hank Han Date: Wed, 12 Aug 2026 15:27:23 +0800 Subject: [PATCH 15/47] ci: restore private delivery pipelines Consolidate the Volcengine image build and sync paths, CUDA 13 variants, DeepSeek V4 nightly, reusable kernel build, ep_main PR suites, runner hardening, private schedule policy, gateway build metadata, bounded image provenance, and immutable-image runtime verification into one CI feature. Signed-off-by: Hank Han --- .../workflows/_docker-build-and-publish.yml | 507 ++++++++++-------- .github/workflows/nightly-release-gateway.yml | 3 - .github/workflows/pr-test-extra.yml | 1 + .github/workflows/pr-test-rust.yml | 4 +- .github/workflows/pr-test.yml | 1 + .../release-docker-deepseek-v4-nightly.yml | 166 ++++++ .github/workflows/release-docker-dev.yml | 191 ++++++- .github/workflows/release-docker-runtime.yml | 4 +- .github/workflows/release-docker.yml | 123 ++++- .github/workflows/release-whl-kernel.yml | 89 ++- .../sync-docker-images-to-volcengine.yml | 125 +++++ .github/workflows/sync-lmsys-sglang-blogs.yml | 2 - docker/Dockerfile | 56 +- scripts/ci/amd/check_vram_clear.sh | 5 +- scripts/ci/get_volcengine_image_tag.py | 83 +++ .../ci/sync_docker_images_to_volcengine.py | 437 +++++++++++++++ .../test_sync_docker_images_to_volcengine.py | 249 +++++++++ .../ci/utils/docker_build_metadata_args.py | 22 + scripts/ci/verify_private_image_runtime.py | 29 + scripts/code_sync/install_github_cli.sh | 39 +- .../bindings/python/pyproject.toml | 4 +- .../tools/test_docker_build_metadata_args.py | 22 +- 22 files changed, 1873 insertions(+), 289 deletions(-) create mode 100644 .github/workflows/release-docker-deepseek-v4-nightly.yml create mode 100644 .github/workflows/sync-docker-images-to-volcengine.yml create mode 100755 scripts/ci/get_volcengine_image_tag.py create mode 100755 scripts/ci/sync_docker_images_to_volcengine.py create mode 100755 scripts/ci/test_sync_docker_images_to_volcengine.py create mode 100644 scripts/ci/verify_private_image_runtime.py diff --git a/.github/workflows/_docker-build-and-publish.yml b/.github/workflows/_docker-build-and-publish.yml index ba55dc939f94..ec6155bf0f45 100644 --- a/.github/workflows/_docker-build-and-publish.yml +++ b/.github/workflows/_docker-build-and-publish.yml @@ -1,17 +1,28 @@ -name: Build and Publish Multi-Arch Docker Images +name: Build and Publish Volcengine Docker Image -# Reusable workflow: builds CUDA 12 + CUDA 13 images for amd64 and arm64, -# then creates multi-arch manifests with caller-specified tags. +# Reusable workflow: builds one amd64 CUDA image variant and publishes the +# matching Volcengine CR tag. Callers use a matrix when multiple variants are +# needed, so the Docker build command stays in one file. on: workflow_call: inputs: docker_target: - description: "Dockerfile target stage (framework or runtime)" + description: "Dockerfile target stage" required: true type: string + cuda_key: + description: "Logical CUDA tag key, e.g. cu130" + required: false + type: string + default: "cu130" + cuda_version: + description: "CUDA_VERSION build arg" + required: false + type: string + default: "13.0.1" sgl_version: - description: "Version string passed as SGL_VERSION build arg (empty to skip)" + description: "Version string passed as SGL_VERSION build arg" required: false type: string default: "" @@ -21,37 +32,87 @@ on: type: string default: "" checkout_ref: - description: "Git ref to checkout (empty for default)" + description: "Git ref to checkout" required: false type: string default: "" tag_config: - description: 'JSON array of {"cuda":"cu129|cu130","tags":["tag1","tag2"]}. Tags support {version} substitution.' - required: true + description: "Deprecated compatibility input; Volcengine tags are generated from tag_mode." + required: false + type: string + default: "" + tag_mode: + description: "Tag mode passed to scripts/ci/get_volcengine_image_tag.py" + required: false + type: string + default: "version" + tag_value: + description: "Tag value passed for version-mode Volcengine tags" + required: false + type: string + default: "" + image_tag_override: + description: "Optional fully resolved image tag to publish instead of generating one" + required: false + type: string + default: "" + variant_suffix: + description: "Optional variant suffix appended before CUDA suffix" + required: false + type: string + default: "" + cuda_suffix: + description: "Optional CUDA suffix appended to Volcengine tags" + required: false type: string + default: "" + publish_default_cuda_alias: + description: "Also publish an unsuffixed default tag for this CUDA variant" + required: false + type: boolean + default: false use_environment: description: "GitHub environment name (e.g. prod) or empty for none" required: false type: string default: "" + registry_host: + description: "Volcengine CR registry host" + required: false + type: string + default: "" image_repo: - description: "Docker Hub repo to push to (e.g. lmsysorg/sglang-staging for testing)" + description: "Full Volcengine CR image repository" required: false type: string - default: "lmsysorg/sglang" + default: "" + artifact_name: + description: "Optional local sgl-kernel wheel artifact name" + required: false + type: string + default: "" jobs: - build-x86: - if: github.repository == 'sgl-project/sglang' + build-and-publish: + if: github.repository == 'bytedance-iaas/sglang' environment: ${{ inputs.use_environment || null }} runs-on: x64-docker-build-node env: - TAG_CONFIG: ${{ inputs.tag_config }} - SGL_VERSION: ${{ inputs.sgl_version }} + CUDA_KEY: ${{ inputs.cuda_key }} + CUDA_SUFFIX: ${{ inputs.cuda_suffix }} IMAGE_REPO: ${{ inputs.image_repo }} - outputs: - digest-cu129: ${{ steps.build-cu129.outputs.digest }} - digest-cu130: ${{ steps.build-cu130.outputs.digest }} + PUBLISH_DEFAULT_CUDA_ALIAS: ${{ inputs.publish_default_cuda_alias }} + REGISTRY_HOST: ${{ inputs.registry_host }} + SGL_VERSION: ${{ inputs.sgl_version }} + TAG_MODE: ${{ inputs.tag_mode }} + TAG_VALUE: ${{ inputs.tag_value }} + IMAGE_TAG_OVERRIDE: ${{ inputs.image_tag_override }} + VARIANT_SUFFIX: ${{ inputs.variant_suffix }} + VOLCENGINE_CR_REGISTRY: ${{ vars.VOLCENGINE_CR_REGISTRY }} + VOLCENGINE_CR_NAMESPACE: ${{ vars.VOLCENGINE_CR_NAMESPACE }} + VOLCENGINE_CR_REPOSITORY: ${{ vars.VOLCENGINE_CR_REPOSITORY || 'sglang' }} + VOLCENGINE_CR_USERNAME: ${{ secrets.VOLCENGINE_CR_USERNAME }} + VOLCENGINE_CR_PASSWORD: ${{ secrets.VOLCENGINE_CR_PASSWORD }} steps: - name: Delete huge unnecessary tools folder run: rm -rf /opt/hostedtoolcache @@ -63,33 +124,124 @@ jobs: uses: actions/checkout@v4 with: ref: ${{ inputs.checkout_ref || github.ref }} + fetch-depth: 0 + fetch-tags: true + + - name: Resolve Volcengine publish target + run: | + set -euo pipefail + + if [ -z "${IMAGE_REPO}" ]; then + if [ -z "${VOLCENGINE_CR_REGISTRY}" ] || [ -z "${VOLCENGINE_CR_NAMESPACE}" ]; then + echo "::error::Volcengine CR variables are required: VOLCENGINE_CR_REGISTRY and VOLCENGINE_CR_NAMESPACE" + exit 1 + fi + IMAGE_REPO="${VOLCENGINE_CR_REGISTRY}/${VOLCENGINE_CR_NAMESPACE}/${VOLCENGINE_CR_REPOSITORY:-sglang}" + fi + + if [ -z "${REGISTRY_HOST}" ] && [ -n "${VOLCENGINE_CR_REGISTRY}" ] && [[ "${IMAGE_REPO}" == "${VOLCENGINE_CR_REGISTRY}/"* ]]; then + REGISTRY_HOST="${VOLCENGINE_CR_REGISTRY}" + fi + + if [ -z "${IMAGE_REPO}" ] || [[ "${IMAGE_REPO}" == /* ]] || [[ "${IMAGE_REPO}" == *'//'* ]]; then + echo "::error::Invalid image repository: '${IMAGE_REPO}'" + exit 1 + fi + + echo "IMAGE_REPO=${IMAGE_REPO}" >> "${GITHUB_ENV}" + echo "REGISTRY_HOST=${REGISTRY_HOST}" >> "${GITHUB_ENV}" + echo "Resolved image repository ${IMAGE_REPO}" + + - name: Resolve Volcengine image tag + id: image-tag + run: | + set -euo pipefail + resolve_tag() { + local cuda_suffix="$1" + local tag_cmd=(python3 scripts/ci/get_volcengine_image_tag.py --mode "${TAG_MODE}") + if [ -n "${TAG_VALUE}" ]; then + tag_cmd+=(--tag-value "${TAG_VALUE}") + fi + if [ -n "${VARIANT_SUFFIX}" ]; then + tag_cmd+=(--variant-suffix "${VARIANT_SUFFIX}") + fi + if [ -n "${cuda_suffix}" ]; then + tag_cmd+=(--cuda-suffix "${cuda_suffix}") + fi + "${tag_cmd[@]}" + } + + IMAGE_TAGS=() + if [ -n "${IMAGE_TAG_OVERRIDE}" ]; then + IMAGE_TAGS+=("${IMAGE_TAG_OVERRIDE}") + elif [ "${PUBLISH_DEFAULT_CUDA_ALIAS}" = "true" ]; then + IMAGE_TAGS+=("$(resolve_tag "")") + IMAGE_TAGS+=("$(resolve_tag "${CUDA_SUFFIX}")") + else + IMAGE_TAGS+=("$(resolve_tag "${CUDA_SUFFIX}")") + fi + + { + echo "image-tags<> "$GITHUB_OUTPUT" - name: Compute Docker build metadata args + id: build-metadata + env: + IMAGE_TAGS: ${{ steps.image-tag.outputs.image-tags }} run: | set -euo pipefail BUILD_COMMIT="$(git rev-parse HEAD)" + BUILD_TREE="$(git rev-parse HEAD^{tree})" + PYTHON_MANIFEST_SHA256="$(python3 - <<'PY' + import hashlib + import subprocess + + paths = subprocess.check_output( + ["git", "ls-files", "-z", "--", "python"] + ).split(b"\0") + digest = hashlib.sha256() + for raw_path in sorted(path for path in paths if path): + digest.update(raw_path) + digest.update(b"\0") + with open(raw_path, "rb") as source_file: + digest.update(hashlib.sha256(source_file.read()).digest()) + print(digest.hexdigest()) + PY + )" + BUILD_SOURCE="${GITHUB_SERVER_URL}/${GITHUB_REPOSITORY}" BUILD_URL="${GITHUB_SERVER_URL}/${GITHUB_REPOSITORY}/actions/runs/${GITHUB_RUN_ID}" - for CUDA_VARIANT in cu129 cu130; do - python3 scripts/ci/utils/docker_build_metadata_args.py \ - --cuda "${CUDA_VARIANT}" \ - --tag-config "${TAG_CONFIG}" \ - --image-repo "${IMAGE_REPO}" \ - --sgl-version "${SGL_VERSION}" \ - --build-commit "${BUILD_COMMIT}" \ - --build-url "${BUILD_URL}" \ - > "/tmp/docker-metadata-${CUDA_VARIANT}.args" - done + TAG_CONFIG="$(CUDA_KEY="${CUDA_KEY}" IMAGE_TAGS="${IMAGE_TAGS}" python3 -c 'import json, os; tags = [tag for tag in os.environ["IMAGE_TAGS"].splitlines() if tag]; print(json.dumps([{"cuda": os.environ["CUDA_KEY"], "tags": tags}], separators=(",", ":")))')" + python3 scripts/ci/utils/docker_build_metadata_args.py \ + --cuda "${CUDA_KEY}" \ + --tag-config "${TAG_CONFIG}" \ + --image-repo "${IMAGE_REPO}" \ + --sgl-version "${SGL_VERSION}" \ + --build-commit "${BUILD_COMMIT}" \ + --build-tree "${BUILD_TREE}" \ + --python-manifest-sha256 "${PYTHON_MANIFEST_SHA256}" \ + --build-source "${BUILD_SOURCE}" \ + --build-url "${BUILD_URL}" \ + > /tmp/docker-metadata.args + { + echo "build-commit=${BUILD_COMMIT}" + echo "build-tree=${BUILD_TREE}" + echo "python-manifest-sha256=${PYTHON_MANIFEST_SHA256}" + echo "build-source=${BUILD_SOURCE}" + } >> "$GITHUB_OUTPUT" - name: Free disk space uses: jlumbroso/free-disk-space@main with: - tool-cache: true - docker-images: true + tool-cache: false + docker-images: false android: true dotnet: true haskell: true large-packages: true - swap-storage: true + swap-storage: false - name: Prune Docker to reclaim disk space run: | @@ -97,235 +249,132 @@ jobs: docker system prune -af --filter "until=72h" docker volume prune -af - - name: Set up Docker Buildx - uses: docker/setup-buildx-action@v3 + - name: Download local sgl-kernel wheel + if: ${{ inputs.artifact_name != '' }} + uses: actions/download-artifact@v4 + with: + name: ${{ inputs.artifact_name }} + path: .ci-artifacts/sgl-kernel - - name: Login to Docker Hub - uses: docker/login-action@v2 + - name: Set up Docker Buildx with local Docker Hub mirror + uses: docker/setup-buildx-action@v3 with: - username: ${{ secrets.DOCKERHUB_USERNAME }} - password: ${{ secrets.DOCKERHUB_TOKEN }} + driver-opts: | + network=host + buildkitd-config-inline: | + [registry."docker.io"] + mirrors = ["127.0.0.1:5000"] - - name: Build and push AMD64 image (CUDA 12) - id: build-cu129 + - name: Login to Volcengine CR run: | - VERSION_ARG="" - if [ -n "${SGL_VERSION}" ]; then - VERSION_ARG="--build-arg SGL_VERSION=${SGL_VERSION}" - fi - mapfile -t METADATA_ARGS < /tmp/docker-metadata-cu129.args - - docker buildx build \ - --target ${{ inputs.docker_target }} \ - --platform linux/amd64 \ - --output type=image,name=${{ inputs.image_repo }},push-by-digest=true,name-canonical=true,push=true \ - -f docker/Dockerfile \ - --build-arg CUDA_VERSION=12.9.1 \ - --build-arg BUILD_TYPE=all \ - --build-arg GRACE_BLACKWELL=0 \ - --build-arg INSTALL_FLASHINFER_JIT_CACHE=1 \ - "${METADATA_ARGS[@]}" \ - ${VERSION_ARG} \ - ${{ inputs.extra_build_args }} \ - --metadata-file /tmp/metadata-cu129.json \ - --no-cache \ - . - - DIGEST=$(python3 -c "import json; print(json.load(open('/tmp/metadata-cu129.json'))['containerimage.digest'])") - echo "Pushed digest: ${DIGEST}" - echo "digest=${DIGEST}" >> $GITHUB_OUTPUT + set -euo pipefail + echo "${VOLCENGINE_CR_PASSWORD}" | docker login "${REGISTRY_HOST}" --username "${VOLCENGINE_CR_USERNAME}" --password-stdin - - name: Build and push AMD64 image (CUDA 13) - id: build-cu130 + - name: Build and push AMD64 image + id: build run: | + set -euo pipefail VERSION_ARG="" if [ -n "${SGL_VERSION}" ]; then VERSION_ARG="--build-arg SGL_VERSION=${SGL_VERSION}" fi - mapfile -t METADATA_ARGS < /tmp/docker-metadata-cu130.args + mapfile -t METADATA_ARGS < /tmp/docker-metadata.args docker buildx build \ --target ${{ inputs.docker_target }} \ --platform linux/amd64 \ - --output type=image,name=${{ inputs.image_repo }},push-by-digest=true,name-canonical=true,push=true \ + --output type=image,name=${IMAGE_REPO},push-by-digest=true,name-canonical=true,push=true \ -f docker/Dockerfile \ - --build-arg CUDA_VERSION=13.0.1 \ + --build-arg CUDA_VERSION=${{ inputs.cuda_version }} \ --build-arg BUILD_TYPE=all \ --build-arg GRACE_BLACKWELL=0 \ --build-arg INSTALL_FLASHINFER_JIT_CACHE=1 \ "${METADATA_ARGS[@]}" \ ${VERSION_ARG} \ ${{ inputs.extra_build_args }} \ - --metadata-file /tmp/metadata-cu130.json \ + --metadata-file /tmp/metadata.json \ --no-cache \ . - DIGEST=$(python3 -c "import json; print(json.load(open('/tmp/metadata-cu130.json'))['containerimage.digest'])") + DIGEST=$(python3 -c "import json; print(json.load(open('/tmp/metadata.json'))['containerimage.digest'])") echo "Pushed digest: ${DIGEST}" - echo "digest=${DIGEST}" >> $GITHUB_OUTPUT + echo "digest=${DIGEST}" >> "$GITHUB_OUTPUT" - build-arm64: - if: github.repository == 'sgl-project/sglang' - environment: ${{ inputs.use_environment || null }} - runs-on: arm-docker-build-node - env: - TAG_CONFIG: ${{ inputs.tag_config }} - SGL_VERSION: ${{ inputs.sgl_version }} - IMAGE_REPO: ${{ inputs.image_repo }} - outputs: - digest-cu129: ${{ steps.build-cu129.outputs.digest }} - digest-cu130: ${{ steps.build-cu130.outputs.digest }} - steps: - - name: Delete huge unnecessary tools folder - run: rm -rf /opt/hostedtoolcache - - - name: Cleanup workspace (remove root-owned files from prior runs) - run: sudo rm -rf "$GITHUB_WORKSPACE"/* || true - - - name: Checkout repository - uses: actions/checkout@v4 - with: - ref: ${{ inputs.checkout_ref || github.ref }} - - - name: Compute Docker build metadata args + - name: Create Volcengine image tag + env: + IMAGE_TAGS: ${{ steps.image-tag.outputs.image-tags }} run: | set -euo pipefail - BUILD_COMMIT="$(git rev-parse HEAD)" - BUILD_URL="${GITHUB_SERVER_URL}/${GITHUB_REPOSITORY}/actions/runs/${GITHUB_RUN_ID}" - for CUDA_VARIANT in cu129 cu130; do - python3 scripts/ci/utils/docker_build_metadata_args.py \ - --cuda "${CUDA_VARIANT}" \ - --tag-config "${TAG_CONFIG}" \ - --image-repo "${IMAGE_REPO}" \ - --sgl-version "${SGL_VERSION}" \ - --build-commit "${BUILD_COMMIT}" \ - --build-url "${BUILD_URL}" \ - > "/tmp/docker-metadata-${CUDA_VARIANT}.args" + mapfile -t image_tags <<< "${IMAGE_TAGS}" + tag_args=() + for image_tag in "${image_tags[@]}"; do + [ -n "${image_tag}" ] || continue + tag_args+=(-t "${IMAGE_REPO}:${image_tag}") done - - name: Prune Docker to reclaim disk space - run: | - docker buildx prune --filter "until=72h" -f - docker system prune -af --filter "until=72h" - docker volume prune -af - - - name: Set up Docker Buildx - uses: docker/setup-buildx-action@v3 - - - name: Login to Docker Hub - uses: docker/login-action@v2 - with: - username: ${{ secrets.DOCKERHUB_USERNAME }} - password: ${{ secrets.DOCKERHUB_TOKEN }} - - - name: Build and push ARM64 image (CUDA 12) - id: build-cu129 - run: | - VERSION_ARG="" - if [ -n "${SGL_VERSION}" ]; then - VERSION_ARG="--build-arg SGL_VERSION=${SGL_VERSION}" - fi - mapfile -t METADATA_ARGS < /tmp/docker-metadata-cu129.args - - docker buildx build \ - --target ${{ inputs.docker_target }} \ - --platform linux/arm64 \ - --output type=image,name=${{ inputs.image_repo }},push-by-digest=true,name-canonical=true,push=true \ - -f docker/Dockerfile \ - --build-arg CUDA_VERSION=12.9.1 \ - --build-arg BUILD_TYPE=all \ - --build-arg GRACE_BLACKWELL=1 \ - --build-arg INSTALL_FLASHINFER_JIT_CACHE=1 \ - "${METADATA_ARGS[@]}" \ - ${VERSION_ARG} \ - ${{ inputs.extra_build_args }} \ - --metadata-file /tmp/metadata-cu129.json \ - --no-cache \ - . - - DIGEST=$(python3 -c "import json; print(json.load(open('/tmp/metadata-cu129.json'))['containerimage.digest'])") - echo "Pushed digest: ${DIGEST}" - echo "digest=${DIGEST}" >> $GITHUB_OUTPUT - - - name: Build and push ARM64 image (CUDA 13) - id: build-cu130 - run: | - VERSION_ARG="" - if [ -n "${SGL_VERSION}" ]; then - VERSION_ARG="--build-arg SGL_VERSION=${SGL_VERSION}" - fi - mapfile -t METADATA_ARGS < /tmp/docker-metadata-cu130.args - - docker buildx build \ - --target ${{ inputs.docker_target }} \ - --platform linux/arm64 \ - --output type=image,name=${{ inputs.image_repo }},push-by-digest=true,name-canonical=true,push=true \ - -f docker/Dockerfile \ - --build-arg CUDA_VERSION=13.0.1 \ - --build-arg BUILD_TYPE=all \ - --build-arg GRACE_BLACKWELL=1 \ - --build-arg INSTALL_FLASHINFER_JIT_CACHE=1 \ - "${METADATA_ARGS[@]}" \ - ${VERSION_ARG} \ - ${{ inputs.extra_build_args }} \ - --metadata-file /tmp/metadata-cu130.json \ - --no-cache \ - . - - DIGEST=$(python3 -c "import json; print(json.load(open('/tmp/metadata-cu130.json'))['containerimage.digest'])") - echo "Pushed digest: ${DIGEST}" - echo "digest=${DIGEST}" >> $GITHUB_OUTPUT - - create-manifests: - runs-on: ubuntu-latest - needs: [build-x86, build-arm64] - if: github.repository == 'sgl-project/sglang' - environment: ${{ inputs.use_environment || null }} - steps: - - name: Set up Docker Buildx - uses: docker/setup-buildx-action@v3 - - - name: Login to Docker Hub - uses: docker/login-action@v2 - with: - username: ${{ secrets.DOCKERHUB_USERNAME }} - password: ${{ secrets.DOCKERHUB_TOKEN }} + docker buildx imagetools create \ + "${tag_args[@]}" \ + "${IMAGE_REPO}@${{ steps.build.outputs.digest }}" + for image_tag in "${image_tags[@]}"; do + [ -n "${image_tag}" ] || continue + echo "Published ${IMAGE_REPO}:${image_tag}" + done - - name: Create multi-arch manifests + - name: Verify pushed AMD64 image and source provenance env: - TAG_CONFIG: ${{ inputs.tag_config }} - SGL_VERSION: ${{ inputs.sgl_version }} - IMAGE_REPO: ${{ inputs.image_repo }} - X86_CU129: ${{ needs.build-x86.outputs.digest-cu129 }} - X86_CU130: ${{ needs.build-x86.outputs.digest-cu130 }} - ARM64_CU129: ${{ needs.build-arm64.outputs.digest-cu129 }} - ARM64_CU130: ${{ needs.build-arm64.outputs.digest-cu130 }} - SHORT_SHA: ${{ github.sha }} + EXPECTED_BUILD_COMMIT: ${{ steps.build-metadata.outputs.build-commit }} + EXPECTED_BUILD_SOURCE: ${{ steps.build-metadata.outputs.build-source }} + EXPECTED_BUILD_TREE: ${{ steps.build-metadata.outputs.build-tree }} + EXPECTED_PYTHON_MANIFEST_SHA256: ${{ steps.build-metadata.outputs.python-manifest-sha256 }} + IMAGE_DIGEST: ${{ steps.build.outputs.digest }} run: | - echo "${TAG_CONFIG}" | jq -c '.[]' | while read -r entry; do - CUDA=$(echo "${entry}" | jq -r '.cuda') - - if [ "${CUDA}" = "cu129" ]; then - X86_DIGEST="${X86_CU129}" - ARM64_DIGEST="${ARM64_CU129}" - else - X86_DIGEST="${X86_CU130}" - ARM64_DIGEST="${ARM64_CU130}" - fi - - TAG_ARGS="" - for tag in $(echo "${entry}" | jq -r '.tags[]'); do - # Substitute template variables - tag=$(echo "${tag}" | sed "s/{version}/${SGL_VERSION}/g") - tag=$(echo "${tag}" | sed "s/{date}/$(date +%Y%m%d)/g") - tag=$(echo "${tag}" | sed "s/{short_sha}/${SHORT_SHA:0:8}/g") - TAG_ARGS="${TAG_ARGS} -t ${IMAGE_REPO}:${tag}" - done - - docker buildx imagetools create \ - ${TAG_ARGS} \ - ${IMAGE_REPO}@${X86_DIGEST} \ - ${IMAGE_REPO}@${ARM64_DIGEST} - - echo "Published:${TAG_ARGS}" - done + set -euo pipefail + IMAGE_REF="${IMAGE_REPO}@${IMAGE_DIGEST}" + docker buildx imagetools inspect "${IMAGE_REF}" + docker pull "${IMAGE_REF}" + + python3 - "${IMAGE_REF}" <<'PY' + import json + import os + import subprocess + import sys + + inspect = json.loads( + subprocess.check_output(["docker", "image", "inspect", sys.argv[1]]) + )[0] + if inspect["Architecture"] != "amd64": + raise SystemExit(f"unexpected architecture: {inspect['Architecture']}") + labels = inspect["Config"].get("Labels") or {} + expected = { + "ai.sglang.build.commit": os.environ["EXPECTED_BUILD_COMMIT"], + "ai.sglang.build.tree": os.environ["EXPECTED_BUILD_TREE"], + "ai.sglang.build.python-manifest-sha256": os.environ[ + "EXPECTED_PYTHON_MANIFEST_SHA256" + ], + "org.opencontainers.image.source": os.environ["EXPECTED_BUILD_SOURCE"], + } + mismatches = { + key: (labels.get(key), value) + for key, value in expected.items() + if labels.get(key) != value + } + if mismatches: + raise SystemExit(f"image provenance mismatch: {mismatches}") + print(json.dumps({"architecture": inspect["Architecture"], **expected})) + PY + + docker run --rm \ + -e EXPECTED_BUILD_COMMIT \ + -e EXPECTED_BUILD_TREE \ + -e EXPECTED_PYTHON_MANIFEST_SHA256 \ + --entrypoint bash "${IMAGE_REF}" -lc ' + set -euo pipefail + test "$SGLANG_BUILD_COMMIT" = "$EXPECTED_BUILD_COMMIT" + test "$SGLANG_BUILD_TREE" = "$EXPECTED_BUILD_TREE" + test "$SGLANG_PYTHON_MANIFEST_SHA256" = "$EXPECTED_PYTHON_MANIFEST_SHA256" + command -v oniond + python3 scripts/ci/verify_private_image_runtime.py + python3 -m py_compile scripts/eic_integration_check.py + python3 -m pytest -q test/registered/unit/test_runtime_context.py -k TestMoeFlagsGroup + python3 -m pytest -q test/registered/unit/test_model_overrides.py -k deepseek_spec_moe_resolution + ' diff --git a/.github/workflows/nightly-release-gateway.yml b/.github/workflows/nightly-release-gateway.yml index 19c952103c40..de6dc0aca1dd 100644 --- a/.github/workflows/nightly-release-gateway.yml +++ b/.github/workflows/nightly-release-gateway.yml @@ -3,9 +3,6 @@ name: Nightly Release SGLang Model Gateway to PyPI on: - schedule: - # Run at 2 AM UTC every day - - cron: '0 2 * * *' workflow_dispatch: # Allow manual trigger jobs: diff --git a/.github/workflows/pr-test-extra.yml b/.github/workflows/pr-test-extra.yml index 3155c77e3528..e4c68a78714c 100644 --- a/.github/workflows/pr-test-extra.yml +++ b/.github/workflows/pr-test-extra.yml @@ -15,6 +15,7 @@ name: PR Test Extra on: pull_request: + branches: [main, ep_main] # `labeled` lets the workflow re-fire when `run-ci-extra` (or `run-ci`) # is added after the latest push — otherwise GHA leaves the stale # skipped run in place and there's no way to enable extra without diff --git a/.github/workflows/pr-test-rust.yml b/.github/workflows/pr-test-rust.yml index 5c8f7fbf4f04..0a60efae243e 100644 --- a/.github/workflows/pr-test-rust.yml +++ b/.github/workflows/pr-test-rust.yml @@ -2,14 +2,14 @@ name: PR Test (SMG) on: push: - branches: [ main ] + branches: [ main, ep_main ] paths: - "sgl-model-gateway/**" - ".github/workflows/pr-test-rust.yml" - "scripts/ci/cuda/ci_install_dependency.sh" - "scripts/ci/cuda/ci_install_gateway_dependencies.sh" pull_request: - branches: [ main ] + branches: [ main, ep_main ] types: [opened, synchronize, reopened, labeled] paths: - "sgl-model-gateway/**" diff --git a/.github/workflows/pr-test.yml b/.github/workflows/pr-test.yml index 6d39630ef624..29a308f6bf57 100644 --- a/.github/workflows/pr-test.yml +++ b/.github/workflows/pr-test.yml @@ -11,6 +11,7 @@ on: schedule: - cron: '0 11,23 * * *' # Run 2x daily, 12h apart, off PR-push peaks pull_request: + branches: [main, ep_main] workflow_dispatch: inputs: force_continue_on_error: diff --git a/.github/workflows/release-docker-deepseek-v4-nightly.yml b/.github/workflows/release-docker-deepseek-v4-nightly.yml new file mode 100644 index 000000000000..e96308e2bc8c --- /dev/null +++ b/.github/workflows/release-docker-deepseek-v4-nightly.yml @@ -0,0 +1,166 @@ +name: Build DeepSeek V4 Nightly CUDA 13 Docker Image + +# NOTE: +# - GitHub Actions schedule runs from the default branch workflow definition. +# - This workflow should eventually live on ep_main/default-branch, but always +# checks out bytedance/deepseek_v4 (or the manually selected ref) explicitly. + +on: + workflow_dispatch: + inputs: + target_ref: + description: "Git ref to build. Leave empty to use the default bytedance/deepseek_v4 branch." + required: false + default: "" + compile_kernel: + description: "Build sgl-kernel from source before building the image" + required: false + type: boolean + default: true + image_repo: + description: "Optional full Volcengine CR image repo override. Leave empty to use repo variables like release-docker-dev.yml." + required: false + default: "" + schedule: + # UTC+8 02:00 + - cron: "0 18 * * *" + +concurrency: + group: release-docker-deepseek-v4-nightly-${{ github.event_name == 'workflow_dispatch' && (inputs.target_ref || 'default-target-ref') || 'default-target-ref' }} + cancel-in-progress: true + +jobs: + resolve-target-ref: + if: ${{ github.repository == 'bytedance-iaas/sglang' }} + runs-on: ubuntu-22.04 + outputs: + target_ref: ${{ steps.resolve.outputs.target_ref }} + steps: + - name: Resolve target ref + id: resolve + run: | + set -euo pipefail + target_ref="${{ inputs.target_ref || 'bytedance/deepseek_v4' }}" + echo "target_ref=${target_ref}" >> "${GITHUB_OUTPUT}" + + resolve-image-metadata: + needs: [resolve-target-ref] + if: ${{ github.repository == 'bytedance-iaas/sglang' }} + environment: prod + runs-on: ubuntu-22.04 + outputs: + image_repo: ${{ steps.resolve.outputs.image_repo }} + image_tag: ${{ steps.resolve.outputs.image_tag }} + image_ref: ${{ steps.resolve.outputs.image_ref }} + target_ref: ${{ steps.resolve.outputs.target_ref }} + env: + IMAGE_REPO: ${{ inputs.image_repo || '' }} + VOLCENGINE_CR_REGISTRY: ${{ vars.VOLCENGINE_CR_REGISTRY }} + VOLCENGINE_CR_NAMESPACE: ${{ vars.VOLCENGINE_CR_NAMESPACE }} + VOLCENGINE_CR_REPOSITORY: ${{ vars.VOLCENGINE_CR_REPOSITORY || 'sglang' }} + steps: + - name: Checkout target ref + uses: actions/checkout@v4 + with: + ref: ${{ needs.resolve-target-ref.outputs.target_ref }} + fetch-depth: 0 + fetch-tags: true + + - name: Resolve image metadata + id: resolve + run: | + set -euo pipefail + + target_ref="${{ needs.resolve-target-ref.outputs.target_ref }}" + image_repo="${IMAGE_REPO}" + if [ -z "${image_repo}" ]; then + if [ -z "${VOLCENGINE_CR_REGISTRY}" ] || [ -z "${VOLCENGINE_CR_NAMESPACE}" ]; then + echo "::error::Volcengine CR variables are required: VOLCENGINE_CR_REGISTRY and VOLCENGINE_CR_NAMESPACE" + exit 1 + fi + image_repo="${VOLCENGINE_CR_REGISTRY}/${VOLCENGINE_CR_NAMESPACE}/${VOLCENGINE_CR_REPOSITORY:-sglang}" + fi + + if [ -z "${image_repo}" ] || [[ "${image_repo}" == /* ]] || [[ "${image_repo}" == *'//'* ]]; then + echo "::error::Invalid image repository: '${image_repo}'" + exit 1 + fi + + tag_mode="manual" + if [ "${{ github.event_name }}" = "schedule" ]; then + tag_mode="nightly" + fi + version="$(python3 python/tools/get_version_tag.py --tag-only)" + version="${version#v}" + timestamp="$(TZ=Asia/Shanghai date +%Y%m%d%H%M)" + if [ "${tag_mode}" = "nightly" ]; then + image_tag="v${version}.iaas.nightly.${timestamp}" + else + image_tag="v${version}.iaas.dev.${timestamp}" + fi + image_tag="${image_tag}-deepseek-v4-cu130" + + echo "image_repo=${image_repo}" >> "${GITHUB_OUTPUT}" + echo "image_tag=${image_tag}" >> "${GITHUB_OUTPUT}" + echo "image_ref=${image_repo}:${image_tag}" >> "${GITHUB_OUTPUT}" + echo "target_ref=${target_ref}" >> "${GITHUB_OUTPUT}" + + - name: Write image metadata artifact + run: | + set -euo pipefail + mkdir -p /tmp/deepseek-v4-image-metadata + jq -n \ + --arg image_repo "${{ steps.resolve.outputs.image_repo }}" \ + --arg image_tag "${{ steps.resolve.outputs.image_tag }}" \ + --arg image_ref "${{ steps.resolve.outputs.image_ref }}" \ + --arg target_ref "${{ steps.resolve.outputs.target_ref }}" \ + --arg sglang_commit "$(git rev-parse HEAD)" \ + --arg workflow_run_id "${GITHUB_RUN_ID}" \ + --arg workflow_run_url "${GITHUB_SERVER_URL}/${GITHUB_REPOSITORY}/actions/runs/${GITHUB_RUN_ID}" \ + '{image_repo:$image_repo,image_tag:$image_tag,image_ref:$image_ref,target_ref:$target_ref,sglang_commit:$sglang_commit,workflow_run_id:$workflow_run_id,workflow_run_url:$workflow_run_url}' \ + > /tmp/deepseek-v4-image-metadata/image-metadata.json + + - name: Upload image metadata artifact + uses: actions/upload-artifact@v4 + with: + name: deepseek-v4-image-metadata + path: /tmp/deepseek-v4-image-metadata/image-metadata.json + retention-days: 30 + + build-kernel-wheel: + needs: [resolve-target-ref] + if: ${{ github.repository == 'bytedance-iaas/sglang' && (github.event_name == 'schedule' || inputs.compile_kernel) }} + uses: ./.github/workflows/release-whl-kernel.yml + with: + checkout_ref: ${{ needs.resolve-target-ref.outputs.target_ref }} + cuda_version: "13.0" + artifact_name: sgl-kernel-deepseek-v4-cu13 + + build-nightly: + needs: [resolve-image-metadata, build-kernel-wheel] + if: ${{ github.repository == 'bytedance-iaas/sglang' && always() && needs.resolve-image-metadata.result == 'success' && (needs.build-kernel-wheel.result == 'success' || needs.build-kernel-wheel.result == 'skipped') }} + uses: ./.github/workflows/_docker-build-and-publish.yml + with: + docker_target: framework_final + checkout_ref: ${{ needs.resolve-image-metadata.outputs.target_ref }} + cuda_key: cu130 + cuda_version: 13.0.1 + cuda_suffix: cu130 + publish_default_cuda_alias: false + variant_suffix: deepseek-v4 + tag_mode: ${{ github.event_name == 'schedule' && 'nightly' || 'manual' }} + image_tag_override: ${{ needs.resolve-image-metadata.outputs.image_tag }} + use_environment: prod + image_repo: ${{ inputs.image_repo || '' }} + artifact_name: ${{ (github.event_name == 'schedule' || inputs.compile_kernel) && 'sgl-kernel-deepseek-v4-cu13' || '' }} + extra_build_args: >- + --build-arg BRANCH_TYPE=local + --build-arg CMAKE_BUILD_PARALLEL_LEVEL=$(nproc) + --build-arg INSTALL_LOCAL_SGL_KERNEL_WHEEL=${{ (github.event_name == 'schedule' || inputs.compile_kernel) && '1' || '0' }} + --build-arg HTTP_PROXY=http://100.68.162.211:3128 + --build-arg HTTPS_PROXY=http://100.68.162.211:3128 + --build-arg ALL_PROXY=http://100.68.162.211:3128 + --build-arg NO_PROXY=localhost,127.0.0.1,::1,mirrors.ivolces.com,.ivolces.com,.volceapi.com,.byted.org,eic-sdk-release.tos-cn-beijing.ivolces.com,scqq9isgq31i0fb8nt4eg.apigateway-cn-beijing.volceapi.com,bytedpypi.byted.org,mirrors.byted.org,iaas-gpu-cn-beijing.cr.volces.com + --build-arg BYTED_INTERNAL_EXTRA_TOOLS_APT_URL=${{ vars.BYTED_INTERNAL_EXTRA_TOOLS_APT_URL || 'http://mirrors.ivolces.com/extra-tools/ubuntu' }} + --build-arg BYTED_INTERNAL_EXTRA_TOOLS_GPG_URL=${{ vars.BYTED_INTERNAL_EXTRA_TOOLS_GPG_URL || 'https://mirrors.ivolces.com/extra-tools/ubuntu/GPG-KEY-system' }} + secrets: inherit diff --git a/.github/workflows/release-docker-dev.yml b/.github/workflows/release-docker-dev.yml index 855903a63d49..7c1d8a388c9f 100644 --- a/.github/workflows/release-docker-dev.yml +++ b/.github/workflows/release-docker-dev.yml @@ -8,7 +8,7 @@ on: required: false default: "" tag: - description: "Custom tag suffix (overrides pr_number in tag). E.g. 'my-test' → dev-my-test, dev-cu13-my-test, etc." + description: "Custom tag suffix" required: false default: "" image_repo: @@ -32,11 +32,25 @@ on: description: "Overlay output suffix appended to the base tag. EMPTY = overwrite the base tag(s); non-empty = append (e.g. 'msa')." required: false default: "" + compile_kernel: + description: "Build sgl-kernel from source (Volcengine fork only)" + required: false + type: boolean + default: true + private_debug_base_only: + description: "Volcengine fork only: build and publish exactly one AMD64 CUDA 13 base image, without compiling a kernel wheel or building SBO/W4A8 variants." + required: false + type: boolean + default: false + private_debug_verify_image: + description: "Volcengine fork only: verify an existing private debug image by immutable digest (repo@sha256:...), without rebuilding it." + required: false + default: "" schedule: - cron: "0 0 * * *" concurrency: - group: release-docker-dev-${{ inputs.tag || inputs.pr_number || 'nightly' }} + group: release-docker-dev-${{ inputs.private_debug_verify_image != '' && format('verify-{0}', github.run_id) || inputs.tag || inputs.pr_number || 'nightly' }} cancel-in-progress: true jobs: @@ -81,14 +95,12 @@ jobs: SUFFIX="-pr-${{ inputs.pr_number }}" fi - # Build tag config. dev-cu13 / nightly-dev-cu13 are published as - # aliases on the cu130 image for backwards compatibility with - # consumers pinned to the pre-flip names. + # Build tag config. Development images are CUDA 13 only. if [ -z "${SUFFIX}" ]; then # Nightly: include dated tags - TAG_CONFIG='[{"cuda":"cu129","tags":["dev-cu12","nightly-dev-cu12-{date}-{short_sha}"]},{"cuda":"cu130","tags":["dev","dev-cu13","nightly-dev-{date}-{short_sha}","nightly-dev-cu13-{date}-{short_sha}"]}]' + TAG_CONFIG='[{"cuda":"cu130","tags":["dev","dev-cu13","nightly-dev-{date}-{short_sha}","nightly-dev-cu13-{date}-{short_sha}"]}]' else - TAG_CONFIG="[{\"cuda\":\"cu129\",\"tags\":[\"dev-cu12${SUFFIX}\"]},{\"cuda\":\"cu130\",\"tags\":[\"dev${SUFFIX}\",\"dev-cu13${SUFFIX}\"]}]" + TAG_CONFIG="[{\"cuda\":\"cu130\",\"tags\":[\"dev${SUFFIX}\",\"dev-cu13${SUFFIX}\"]}]" fi echo "tag_config=${TAG_CONFIG}" >> $GITHUB_OUTPUT @@ -140,7 +152,7 @@ jobs: build-and-publish: needs: prepare - if: ${{ !inputs.build_only }} + if: ${{ github.repository == 'sgl-project/sglang' && !inputs.build_only }} uses: ./.github/workflows/_docker-build-and-publish.yml with: docker_target: framework_final @@ -211,10 +223,10 @@ jobs: cleanup-nightly: needs: build-and-publish - if: ${{ !inputs.build_only && !inputs.tag && !inputs.pr_number }} + if: ${{ github.repository == 'sgl-project/sglang' && !inputs.build_only && !inputs.tag && !inputs.pr_number }} uses: ./.github/workflows/_docker-cleanup-nightly.yml with: - tag_prefixes: '["nightly-dev", "nightly-dev-cu12", "nightly-dev-cu13"]' + tag_prefixes: '["nightly-dev", "nightly-dev-cu13"]' image_repo: ${{ inputs.image_repo || 'lmsysorg/sglang' }} secrets: inherit @@ -261,3 +273,162 @@ jobs: -f Dockerfile.overlay \ $OUT_ARGS \ --push . + + build-dev-kernel-wheel: + if: ${{ github.repository == 'bytedance-iaas/sglang' && !inputs.private_debug_base_only && inputs.private_debug_verify_image == '' && (github.event_name == 'schedule' || inputs.compile_kernel) }} + strategy: + fail-fast: false + matrix: + include: + - arch_tag: x86-cu13 + cuda_version: "13.0" + uses: ./.github/workflows/release-whl-kernel.yml + with: + checkout_ref: ${{ inputs.pr_number && format('refs/pull/{0}/head', inputs.pr_number) || github.ref }} + cuda_version: ${{ matrix.cuda_version }} + artifact_name: sgl-kernel-${{ matrix.arch_tag }} + + build-dev: + needs: [build-dev-kernel-wheel] + if: ${{ github.repository == 'bytedance-iaas/sglang' && !inputs.private_debug_base_only && inputs.private_debug_verify_image == '' && always() && (needs.build-dev-kernel-wheel.result == 'success' || needs.build-dev-kernel-wheel.result == 'skipped') }} + strategy: + fail-fast: false + matrix: + include: + - cuda_key: cu130 + cuda_version: 13.0.1 + cuda_suffix: cu130 + publish_default_cuda_alias: true + hopper_sbo: 0 + w4a8_per_tensor_transfer: 0 + variant_suffix: "" + artifact_name: sgl-kernel-x86-cu13 + - cuda_key: cu130 + cuda_version: 13.0.1 + cuda_suffix: cu130 + publish_default_cuda_alias: true + hopper_sbo: 1 + w4a8_per_tensor_transfer: 0 + variant_suffix: sbo + artifact_name: sgl-kernel-x86-cu13 + - cuda_key: cu130 + cuda_version: 13.0.1 + cuda_suffix: cu130 + publish_default_cuda_alias: true + hopper_sbo: 0 + w4a8_per_tensor_transfer: 1 + variant_suffix: w4a8 + artifact_name: sgl-kernel-x86-cu13 + uses: ./.github/workflows/_docker-build-and-publish.yml + with: + docker_target: framework_final + checkout_ref: ${{ inputs.pr_number && format('refs/pull/{0}/head', inputs.pr_number) || github.ref }} + cuda_key: ${{ matrix.cuda_key }} + cuda_version: ${{ matrix.cuda_version }} + cuda_suffix: ${{ matrix.cuda_suffix }} + publish_default_cuda_alias: ${{ matrix.publish_default_cuda_alias }} + variant_suffix: ${{ matrix.variant_suffix }} + tag_mode: ${{ github.event_name == 'schedule' && 'nightly' || 'manual' }} + use_environment: prod + artifact_name: ${{ (github.event_name == 'schedule' || inputs.compile_kernel) && matrix.artifact_name || '' }} + extra_build_args: >- + --build-arg BRANCH_TYPE=local + --build-arg CMAKE_BUILD_PARALLEL_LEVEL=$(nproc) + --build-arg INSTALL_LOCAL_SGL_KERNEL_WHEEL=${{ (github.event_name == 'schedule' || inputs.compile_kernel) && '1' || '0' }} + --build-arg HOPPER_SBO=${{ matrix.hopper_sbo }} + --build-arg W4A8_PER_TENSOR_TRANSFER=${{ matrix.w4a8_per_tensor_transfer }} + --build-arg HTTP_PROXY=http://100.68.162.211:3128 + --build-arg HTTPS_PROXY=http://100.68.162.211:3128 + --build-arg ALL_PROXY=http://100.68.162.211:3128 + --build-arg NO_PROXY=localhost,127.0.0.1,::1,mirrors.ivolces.com,.ivolces.com,.volceapi.com,.byted.org,eic-sdk-release.tos-cn-beijing.ivolces.com,scqq9isgq31i0fb8nt4eg.apigateway-cn-beijing.volceapi.com,bytedpypi.byted.org,mirrors.byted.org,iaas-gpu-cn-beijing.cr.volces.com + --build-arg BYTED_INTERNAL_EXTRA_TOOLS_APT_URL=${{ vars.BYTED_INTERNAL_EXTRA_TOOLS_APT_URL || 'http://mirrors.ivolces.com/extra-tools/ubuntu' }} + --build-arg BYTED_INTERNAL_EXTRA_TOOLS_GPG_URL=${{ vars.BYTED_INTERNAL_EXTRA_TOOLS_GPG_URL || 'https://mirrors.ivolces.com/extra-tools/ubuntu/GPG-KEY-system' }} + secrets: inherit + + build-dev-debug-base: + if: ${{ github.repository == 'bytedance-iaas/sglang' && inputs.private_debug_base_only && inputs.private_debug_verify_image == '' }} + uses: ./.github/workflows/_docker-build-and-publish.yml + with: + docker_target: framework_final + checkout_ref: ${{ inputs.pr_number && format('refs/pull/{0}/head', inputs.pr_number) || github.ref }} + cuda_key: cu130 + cuda_version: 13.0.1 + cuda_suffix: cu130 + publish_default_cuda_alias: false + image_tag_override: ${{ format('debug-{0}-{1}-{2}', inputs.tag || 'servingkit', github.run_id, github.run_attempt) }} + use_environment: prod + extra_build_args: >- + --build-arg BRANCH_TYPE=local + --build-arg CMAKE_BUILD_PARALLEL_LEVEL=$(nproc) + --build-arg INSTALL_LOCAL_SGL_KERNEL_WHEEL=0 + --build-arg HOPPER_SBO=0 + --build-arg W4A8_PER_TENSOR_TRANSFER=0 + --build-arg HTTP_PROXY=http://100.68.162.211:3128 + --build-arg HTTPS_PROXY=http://100.68.162.211:3128 + --build-arg ALL_PROXY=http://100.68.162.211:3128 + --build-arg NO_PROXY=localhost,127.0.0.1,::1,mirrors.ivolces.com,.ivolces.com,.volceapi.com,.byted.org,eic-sdk-release.tos-cn-beijing.ivolces.com,scqq9isgq31i0fb8nt4eg.apigateway-cn-beijing.volceapi.com,bytedpypi.byted.org,mirrors.byted.org,iaas-gpu-cn-beijing.cr.volces.com + --build-arg BYTED_INTERNAL_EXTRA_TOOLS_APT_URL=${{ vars.BYTED_INTERNAL_EXTRA_TOOLS_APT_URL || 'http://mirrors.ivolces.com/extra-tools/ubuntu' }} + --build-arg BYTED_INTERNAL_EXTRA_TOOLS_GPG_URL=${{ vars.BYTED_INTERNAL_EXTRA_TOOLS_GPG_URL || 'https://mirrors.ivolces.com/extra-tools/ubuntu/GPG-KEY-system' }} + secrets: inherit + + verify-private-debug-image: + if: ${{ github.repository == 'bytedance-iaas/sglang' && inputs.private_debug_verify_image != '' }} + environment: prod + runs-on: x64-docker-build-node + env: + IMAGE_REF: ${{ inputs.private_debug_verify_image }} + VOLCENGINE_CR_USERNAME: ${{ secrets.VOLCENGINE_CR_USERNAME }} + VOLCENGINE_CR_PASSWORD: ${{ secrets.VOLCENGINE_CR_PASSWORD }} + steps: + - name: Checkout verification script + uses: actions/checkout@v4 + + - name: Login to Volcengine CR + run: | + set -euo pipefail + REGISTRY_HOST="${IMAGE_REF%%/*}" + echo "${VOLCENGINE_CR_PASSWORD}" | docker login "${REGISTRY_HOST}" --username "${VOLCENGINE_CR_USERNAME}" --password-stdin + + - name: Verify existing AMD64 image and source provenance + run: | + set -euo pipefail + case "${IMAGE_REF}" in + *@sha256:*) ;; + *) echo "private_debug_verify_image must be an immutable repo@sha256:... reference" >&2; exit 1 ;; + esac + docker buildx imagetools inspect "${IMAGE_REF}" + docker pull "${IMAGE_REF}" + + python3 - "${IMAGE_REF}" <<'PY' + import json + import subprocess + import sys + + inspect = json.loads( + subprocess.check_output(["docker", "image", "inspect", sys.argv[1]]) + )[0] + if inspect["Architecture"] != "amd64": + raise SystemExit(f"unexpected architecture: {inspect['Architecture']}") + labels = inspect["Config"].get("Labels") or {} + required = ( + "ai.sglang.build.commit", + "ai.sglang.build.tree", + "ai.sglang.build.python-manifest-sha256", + "org.opencontainers.image.source", + ) + missing = [key for key in required if not labels.get(key)] + if missing: + raise SystemExit(f"missing image provenance labels: {missing}") + if labels["org.opencontainers.image.source"] != "https://github.com/bytedance-iaas/sglang": + raise SystemExit("unexpected image source label") + print(json.dumps({"architecture": inspect["Architecture"], **{key: labels[key] for key in required}})) + PY + + docker run --rm -i --entrypoint python3 "${IMAGE_REF}" - < scripts/ci/verify_private_image_runtime.py + docker run --rm --entrypoint bash "${IMAGE_REF}" -lc ' + set -euo pipefail + command -v oniond + python3 -m py_compile scripts/eic_integration_check.py + python3 -m pytest -q test/registered/unit/test_runtime_context.py -k TestMoeFlagsGroup + python3 -m pytest -q test/registered/unit/test_model_overrides.py -k deepseek_spec_moe_resolution + ' diff --git a/.github/workflows/release-docker-runtime.yml b/.github/workflows/release-docker-runtime.yml index 0e224bf91e33..1f5e4e5e5e2a 100644 --- a/.github/workflows/release-docker-runtime.yml +++ b/.github/workflows/release-docker-runtime.yml @@ -2,7 +2,6 @@ name: Release Docker Runtime Images # # Builds and publishes runtime Docker images (production-optimized, ~50% smaller): # - lmsysorg/sglang:v{version}-runtime, lmsysorg/sglang:latest-runtime -# - lmsysorg/sglang:v{version}-cu129-runtime, lmsysorg/sglang:latest-cu129-runtime # on: push: @@ -49,7 +48,6 @@ jobs: image_repo: ${{ inputs.image_repo || 'lmsysorg/sglang' }} tag_config: | [ - {"cuda": "cu130", "tags": ["v${{ needs.resolve-version.outputs.version }}-runtime", "latest-runtime", "v${{ needs.resolve-version.outputs.version }}-cu130-runtime", "latest-cu130-runtime"]}, - {"cuda": "cu129", "tags": ["v${{ needs.resolve-version.outputs.version }}-cu129-runtime", "latest-cu129-runtime"]} + {"cuda": "cu130", "tags": ["v${{ needs.resolve-version.outputs.version }}-runtime", "latest-runtime", "v${{ needs.resolve-version.outputs.version }}-cu130-runtime", "latest-cu130-runtime"]} ] secrets: inherit diff --git a/.github/workflows/release-docker.yml b/.github/workflows/release-docker.yml index edf21469e134..88009560c92c 100644 --- a/.github/workflows/release-docker.yml +++ b/.github/workflows/release-docker.yml @@ -1,22 +1,23 @@ name: Release Docker Images # -# Builds and publishes framework Docker images (full development environment): -# - lmsysorg/sglang:v{version}, lmsysorg/sglang:latest (cuda 13) -# - lmsysorg/sglang:v{version}-cu129, lmsysorg/sglang:latest-cu129 +# Builds and publishes framework Docker images to Volcengine CR: +# - v{sglang}.byted.{internal_release}.{timestamp} (CUDA 13) +# - v{sglang}.byted.{internal_release}.{timestamp}-cu130 +# - v{sglang}.byted.{internal_release}.{timestamp}-{sbo,w4a8}[-cu130] # on: push: tags: - "v[0-9]+.*" + - "[0-9]+.*" workflow_dispatch: inputs: version: - description: "Version to build (without v prefix, e.g., 0.5.7)" - required: true - image_repo: - description: "Docker Hub repo to push to. Use lmsysorg/sglang-staging for testing." + description: "SGLang version for sgl-project builds; fallback internal version for Byted builds" + required: false + internal_version: + description: "Byted internal version for Volcengine tags, e.g. 0.0.11" required: false - default: "lmsysorg/sglang" jobs: resolve-version: @@ -40,6 +41,7 @@ jobs: echo "version=${VERSION}" >> $GITHUB_OUTPUT build-and-publish: + if: github.repository == 'sgl-project/sglang' needs: resolve-version uses: ./.github/workflows/_docker-build-and-publish.yml with: @@ -49,7 +51,108 @@ jobs: image_repo: ${{ inputs.image_repo || 'lmsysorg/sglang' }} tag_config: | [ - {"cuda": "cu130", "tags": ["v${{ needs.resolve-version.outputs.version }}", "latest", "v${{ needs.resolve-version.outputs.version }}-cu130", "latest-cu130"]}, - {"cuda": "cu129", "tags": ["v${{ needs.resolve-version.outputs.version }}-cu129", "latest-cu129"]} + {"cuda": "cu130", "tags": ["v${{ needs.resolve-version.outputs.version }}", "latest", "v${{ needs.resolve-version.outputs.version }}-cu130", "latest-cu130"]} ] secrets: inherit + + resolve-volcengine-release: + if: github.repository == 'bytedance-iaas/sglang' + runs-on: ubuntu-22.04 + outputs: + community-version: ${{ steps.release.outputs.community-version }} + internal-version: ${{ steps.release.outputs.internal-version }} + steps: + - name: Checkout repository + uses: actions/checkout@v4 + with: + fetch-depth: 0 + + - name: Fetch version tags + run: git fetch --force --tags origin + + - name: Resolve Volcengine release versions + id: release + run: | + set -euo pipefail + + if [ "${{ github.event_name }}" = "workflow_dispatch" ]; then + INTERNAL_VERSION="${{ github.event.inputs.internal_version || github.event.inputs.version }}" + RAW_TAG="v${INTERNAL_VERSION}" + if ! git rev-parse -q --verify "refs/tags/${RAW_TAG}" >/dev/null; then + RAW_TAG="${INTERNAL_VERSION}" + fi + else + RAW_TAG="${GITHUB_REF_NAME}" + INTERNAL_VERSION="${RAW_TAG#v}" + fi + + if [ -z "${RAW_TAG}" ] || ! git rev-parse -q --verify "refs/tags/${RAW_TAG}" >/dev/null; then + echo "::error::Volcengine image publishing requires an existing git tag for the internal version" + exit 1 + fi + + INTERNAL_VERSION="${INTERNAL_VERSION#v}" + if [ -z "${INTERNAL_VERSION}" ] || ! echo "${INTERNAL_VERSION}" | grep -qE '^[0-9]+\.[0-9]+\.[0-9]+([._-][0-9A-Za-z]+)*$'; then + echo "::error::Invalid internal version: ${INTERNAL_VERSION} (expected: X.Y.Z, for example 0.0.11)" + exit 1 + fi + + COMMUNITY_VERSION="$(python3 python/tools/get_version_tag.py --tag-only | sed 's/^v//')" + if [ -z "${COMMUNITY_VERSION}" ] || ! echo "${COMMUNITY_VERSION}" | grep -qE '^[0-9]+\.[0-9]+\.[0-9]+'; then + echo "::error::Invalid community version resolved from checkout: ${COMMUNITY_VERSION}" + exit 1 + fi + + echo "community-version=${COMMUNITY_VERSION}" >> "$GITHUB_OUTPUT" + echo "internal-version=${INTERNAL_VERSION}" >> "$GITHUB_OUTPUT" + + build-and-publish-volcengine: + if: github.repository == 'bytedance-iaas/sglang' + needs: [resolve-volcengine-release] + strategy: + fail-fast: false + matrix: + include: + - cuda_key: cu130 + cuda_version: 13.0.1 + cuda_suffix: cu130 + publish_default_cuda_alias: true + hopper_sbo: 0 + w4a8_per_tensor_transfer: 0 + variant_suffix: "" + - cuda_key: cu130 + cuda_version: 13.0.1 + cuda_suffix: cu130 + publish_default_cuda_alias: true + hopper_sbo: 1 + w4a8_per_tensor_transfer: 0 + variant_suffix: sbo + - cuda_key: cu130 + cuda_version: 13.0.1 + cuda_suffix: cu130 + publish_default_cuda_alias: true + hopper_sbo: 0 + w4a8_per_tensor_transfer: 1 + variant_suffix: w4a8 + uses: ./.github/workflows/_docker-build-and-publish.yml + with: + docker_target: framework_final + cuda_key: ${{ matrix.cuda_key }} + cuda_version: ${{ matrix.cuda_version }} + cuda_suffix: ${{ matrix.cuda_suffix }} + publish_default_cuda_alias: ${{ matrix.publish_default_cuda_alias }} + variant_suffix: ${{ matrix.variant_suffix }} + sgl_version: ${{ needs.resolve-volcengine-release.outputs.community-version }} + tag_mode: version + tag_value: ${{ needs.resolve-volcengine-release.outputs.internal-version }} + use_environment: prod + extra_build_args: >- + --build-arg HTTP_PROXY=http://100.68.162.211:3128 + --build-arg HTTPS_PROXY=http://100.68.162.211:3128 + --build-arg ALL_PROXY=http://100.68.162.211:3128 + --build-arg HOPPER_SBO=${{ matrix.hopper_sbo }} + --build-arg W4A8_PER_TENSOR_TRANSFER=${{ matrix.w4a8_per_tensor_transfer }} + --build-arg NO_PROXY=localhost,127.0.0.1,::1,mirrors.ivolces.com,.ivolces.com,.volceapi.com,.byted.org,eic-sdk-release.tos-cn-beijing.ivolces.com,scqq9isgq31i0fb8nt4eg.apigateway-cn-beijing.volceapi.com,bytedpypi.byted.org,mirrors.byted.org,iaas-gpu-cn-beijing.cr.volces.com + --build-arg BYTED_INTERNAL_EXTRA_TOOLS_APT_URL=${{ vars.BYTED_INTERNAL_EXTRA_TOOLS_APT_URL || 'http://mirrors.ivolces.com/extra-tools/ubuntu' }} + --build-arg BYTED_INTERNAL_EXTRA_TOOLS_GPG_URL=${{ vars.BYTED_INTERNAL_EXTRA_TOOLS_GPG_URL || 'https://mirrors.ivolces.com/extra-tools/ubuntu/GPG-KEY-system' }} + secrets: inherit diff --git a/.github/workflows/release-whl-kernel.yml b/.github/workflows/release-whl-kernel.yml index 50adb14f37f7..b4fe36ab9bf1 100644 --- a/.github/workflows/release-whl-kernel.yml +++ b/.github/workflows/release-whl-kernel.yml @@ -6,6 +6,50 @@ on: - main paths: - python/sglang/kernels/aot/python/sgl_kernel/version.py + workflow_call: + inputs: + checkout_ref: + description: "Git ref to build from" + required: true + type: string + python_version: + description: "Python version passed to sgl-kernel/build.sh" + required: false + type: string + default: "3.10" + cuda_version: + description: "CUDA version passed to sgl-kernel/build.sh" + required: true + type: string + arch: + description: "Optional architecture override passed to sgl-kernel/build.sh" + required: false + type: string + default: "" + artifact_name: + description: "Uploaded artifact name" + required: true + type: string + runner: + description: "Runner label for the wheel build" + required: false + type: string + default: "x64-kernel-build-node" + build_jobs: + description: "BUILD_JOBS passed to sgl-kernel/build.sh" + required: false + type: string + default: "64" + nvcc_threads: + description: "NVCC_THREADS passed to sgl-kernel/build.sh" + required: false + type: string + default: "8" + use_ccache: + description: "USE_CCACHE passed to sgl-kernel/build.sh" + required: false + type: string + default: "1" workflow_dispatch: inputs: target: @@ -30,10 +74,53 @@ on: required: false concurrency: - group: release-sglang-kernels-${{ github.ref }} + group: release-sglang-kernels-${{ github.ref }}-${{ inputs.artifact_name || github.event.inputs.target || 'standalone' }} cancel-in-progress: true jobs: + build-wheel: + if: ${{ inputs.artifact_name != '' }} + runs-on: ${{ inputs.runner }} + steps: + # Self-hosted build nodes retain the workspace across jobs. Prior builds + # leave root-owned artifacts under python/sglang/kernels/aot/build/ that actions/checkout + # cannot remove, causing EACCES on rmdir. Wipe them via a throwaway root + # container before checkout recreates the workspace. + - name: Clean workspace (remove root-owned files from prior runs) + run: | + docker run --rm -v "${{ github.workspace }}:/workspace" alpine:3 \ + sh -c 'rm -rf /workspace/..?* /workspace/.[!.]* /workspace/*' || true + + - uses: actions/checkout@v4 + with: + ref: ${{ inputs.checkout_ref || github.ref }} + + - name: Set up Python ${{ inputs.python_version }} + uses: actions/setup-python@v5 + with: + python-version: ${{ inputs.python_version }} + + - name: Build wheels + run: | + cd python/sglang/kernels/aot + chmod +x ./build.sh + if [ -n "${{ inputs.arch }}" ]; then + ./build.sh "${{ inputs.python_version }}" "${{ inputs.cuda_version }}" "${{ inputs.arch }}" + else + ./build.sh "${{ inputs.python_version }}" "${{ inputs.cuda_version }}" + fi + env: + BUILD_JOBS: ${{ inputs.build_jobs }} + NVCC_THREADS: ${{ inputs.nvcc_threads }} + USE_CCACHE: ${{ inputs.use_ccache }} + + - name: Upload artifacts + uses: actions/upload-artifact@v4 + with: + name: ${{ inputs.artifact_name }} + path: python/sglang/kernels/aot/dist/* + if-no-files-found: error + # cu130 is the PyPI-released variant; cu129 wheels are published only to the # sgl-project/whl index (consumed via `pip install ...+cu129` for the legacy # cuda 12.9 path), not to PyPI. diff --git a/.github/workflows/sync-docker-images-to-volcengine.yml b/.github/workflows/sync-docker-images-to-volcengine.yml new file mode 100644 index 000000000000..80c7ed8e1c61 --- /dev/null +++ b/.github/workflows/sync-docker-images-to-volcengine.yml @@ -0,0 +1,125 @@ +name: Sync Docker Images to Volcengine CR + +on: + workflow_dispatch: + inputs: + sglang_source: + description: "Source SGLang image repository without tag" + required: false + default: "docker.io/lmsysorg/sglang" + sglang_tags: + description: "Comma or newline separated SGLang tags to sync. Special values: version, today-nightly." + required: false + default: "version,today-nightly" + sglang_repository: + description: "Destination SGLang repository inside the Volcengine CR namespace" + required: false + default: "" + vllm_source: + description: "Source vLLM image repository without tag" + required: false + default: "docker.io/vllm/vllm-openai" + vllm_tags: + description: "Comma or newline separated vLLM tags to sync. Automatic selectors prefer Ubuntu 24.04 tags and fall back to unsuffixed tags when unavailable." + required: false + default: "version,today-nightly" + vllm_repository: + description: "Destination vLLM repository inside the Volcengine CR namespace" + required: false + default: "" + platform: + description: "Platform to sync from multi-arch source images" + required: false + default: "linux/amd64" + schedule: + # UTC+8 17:00, after the usual upstream SGLang and vLLM nightly image builds. + - cron: "0 9 * * *" + +concurrency: + group: sync-docker-images-to-volcengine-${{ github.event_name == 'workflow_dispatch' && github.run_id || 'nightly' }} + cancel-in-progress: true + +permissions: + contents: read + +jobs: + sync: + if: github.repository == 'bytedance-iaas/sglang' + runs-on: x64-docker-build-node + environment: prod + env: + VOLCENGINE_CR_REGISTRY: ${{ vars.VOLCENGINE_CR_REGISTRY }} + VOLCENGINE_CR_NAMESPACE: ${{ vars.VOLCENGINE_CR_NAMESPACE }} + VOLCENGINE_CR_USERNAME: ${{ secrets.VOLCENGINE_CR_USERNAME }} + VOLCENGINE_CR_PASSWORD: ${{ secrets.VOLCENGINE_CR_PASSWORD }} + SGLANG_SOURCE: ${{ inputs.sglang_source || 'docker.io/lmsysorg/sglang' }} + SGLANG_TAGS: ${{ inputs.sglang_tags || 'version,today-nightly' }} + SGLANG_REPOSITORY: ${{ inputs.sglang_repository || vars.VOLCENGINE_CR_REPOSITORY || 'sglang' }} + VLLM_SOURCE: ${{ inputs.vllm_source || 'docker.io/vllm/vllm-openai' }} + VLLM_TAGS: ${{ inputs.vllm_tags || 'version,today-nightly' }} + VLLM_REPOSITORY: ${{ inputs.vllm_repository || vars.VOLCENGINE_CR_VLLM_REPOSITORY || 'vllm' }} + SYNC_PLATFORM: ${{ inputs.platform || 'linux/amd64' }} + steps: + - name: Delete huge unnecessary tools folder + run: rm -rf /opt/hostedtoolcache + + - name: Cleanup workspace + run: sudo rm -rf "$GITHUB_WORKSPACE"/* || true + + - name: Checkout repository + uses: actions/checkout@v4 + + - name: Validate Volcengine CR configuration + run: | + set -euo pipefail + if [ -z "${VOLCENGINE_CR_REGISTRY}" ] || [ -z "${VOLCENGINE_CR_NAMESPACE}" ]; then + echo "::error::Volcengine CR variables are required: VOLCENGINE_CR_REGISTRY and VOLCENGINE_CR_NAMESPACE" + exit 1 + fi + if [ -z "${VOLCENGINE_CR_USERNAME}" ] || [ -z "${VOLCENGINE_CR_PASSWORD}" ]; then + echo "::error::Volcengine CR secrets are required: VOLCENGINE_CR_USERNAME and VOLCENGINE_CR_PASSWORD" + exit 1 + fi + + - name: Set up Docker Buildx with local Docker Hub mirror + uses: docker/setup-buildx-action@v3 + with: + driver-opts: | + network=host + buildkitd-config-inline: | + [registry."docker.io"] + mirrors = ["127.0.0.1:5000"] + + - name: Login to Volcengine CR + run: | + set -euo pipefail + echo "${VOLCENGINE_CR_PASSWORD}" | docker login "${VOLCENGINE_CR_REGISTRY}" --username "${VOLCENGINE_CR_USERNAME}" --password-stdin + + - name: Show sync plan + run: | + set -euo pipefail + python3 scripts/ci/sync_docker_images_to_volcengine.py \ + --registry "${VOLCENGINE_CR_REGISTRY}" \ + --namespace "${VOLCENGINE_CR_NAMESPACE}" \ + --sglang-source "${SGLANG_SOURCE}" \ + --sglang-repository "${SGLANG_REPOSITORY}" \ + --sglang-tags "${SGLANG_TAGS}" \ + --vllm-source "${VLLM_SOURCE}" \ + --vllm-repository "${VLLM_REPOSITORY}" \ + --vllm-tags "${VLLM_TAGS}" \ + --platform "${SYNC_PLATFORM}" + + - name: Sync images + run: | + set -euo pipefail + python3 scripts/ci/sync_docker_images_to_volcengine.py \ + --registry "${VOLCENGINE_CR_REGISTRY}" \ + --namespace "${VOLCENGINE_CR_NAMESPACE}" \ + --sglang-source "${SGLANG_SOURCE}" \ + --sglang-repository "${SGLANG_REPOSITORY}" \ + --sglang-tags "${SGLANG_TAGS}" \ + --vllm-source "${VLLM_SOURCE}" \ + --vllm-repository "${VLLM_REPOSITORY}" \ + --vllm-tags "${VLLM_TAGS}" \ + --platform "${SYNC_PLATFORM}" \ + --execute diff --git a/.github/workflows/sync-lmsys-sglang-blogs.yml b/.github/workflows/sync-lmsys-sglang-blogs.yml index 03d7f5a9da92..f6c9870c64d1 100644 --- a/.github/workflows/sync-lmsys-sglang-blogs.yml +++ b/.github/workflows/sync-lmsys-sglang-blogs.yml @@ -2,8 +2,6 @@ name: Sync LMSYS SGLang blogs on: workflow_dispatch: - schedule: - - cron: "0 */12 * * *" permissions: contents: write diff --git a/docker/Dockerfile b/docker/Dockerfile index d05af8f88bc8..89e2d319d2ab 100644 --- a/docker/Dockerfile +++ b/docker/Dockerfile @@ -48,11 +48,12 @@ fi # Ubuntu 24.04 ships Python 3.12 in main, so we no longer need the deadsnakes # PPA. Dropping it avoids transient Launchpad 504s in `add-apt-repository`. RUN --mount=type=cache,target=/var/cache/apt,id=base-apt \ - apt update && apt install -y --no-install-recommends wget software-properties-common \ + apt update && apt install -y --no-install-recommends wget curl software-properties-common \ && apt install -y --no-install-recommends python3.12-full python3.12-dev \ && update-alternatives --install /usr/bin/python3 python3 /usr/bin/python3.12 2 \ && update-alternatives --set python3 /usr/bin/python3.12 \ - && wget -q https://bootstrap.pypa.io/get-pip.py \ + && curl --fail --show-error --location --retry 5 --retry-all-errors --retry-delay 2 \ + -o get-pip.py https://bootstrap.pypa.io/get-pip.py \ && python3 get-pip.py --break-system-packages \ && rm get-pip.py \ # Allow pip to install packages globally (PEP 668 workaround for Ubuntu 24.04) @@ -192,8 +193,11 @@ WORKDIR /sgl-workspace # Rust toolchain for setuptools-rust extensions (e.g. sglang-grpc). # Requires >= 1.85 (edition 2024). Inherited by framework via FROM torch_deps. ENV PATH="/root/.cargo/bin:${PATH}" -RUN curl --proto '=https' --tlsv1.2 --retry 3 --retry-delay 2 -sSf https://sh.rustup.rs \ - | sh -s -- -y --no-modify-path --profile minimal \ +RUN curl --proto '=https' --tlsv1.2 --fail --show-error --location \ + --retry 5 --retry-all-errors --retry-delay 2 \ + -o /tmp/rustup-init.sh https://sh.rustup.rs \ + && sh /tmp/rustup-init.sh -y --no-modify-path --profile minimal \ + && rm /tmp/rustup-init.sh \ && rustc --version && cargo --version # Install sgl-kernel (from pre-built wheel) @@ -434,12 +438,19 @@ RUN CMAKE_VERSION=3.31.1 \ && mkdir -p /tools/share && cp -r "/tmp/${CMAKE_INSTALLER}/share/"* /tools/share/ \ && rm -rf "/tmp/${CMAKE_INSTALLER}" "/tmp/${CMAKE_INSTALLER}.tar.gz" -RUN curl --proto '=https' --tlsv1.2 --retry 3 --retry-delay 2 -sSf https://just.systems/install.sh | \ - sed "s|https://github.com|https://${GITHUB_ARTIFACTORY}|g" | \ - bash -s -- --tag 1.42.4 --to /tools +RUN curl --proto '=https' --tlsv1.2 --fail --show-error --location \ + --retry 5 --retry-all-errors --retry-delay 2 \ + -o /tmp/just-install.sh https://just.systems/install.sh \ + && sed -i "s|https://github.com|https://${GITHUB_ARTIFACTORY}|g" /tmp/just-install.sh \ + && bash /tmp/just-install.sh --tag 1.42.4 --to /tools \ + && rm /tmp/just-install.sh \ + && test -x /tools/just # Install oh-my-zsh and plugins -RUN sh -c "$(curl --retry 3 --retry-delay 2 -fsSL https://raw.githubusercontent.com/ohmyzsh/ohmyzsh/master/tools/install.sh)" "" --unattended \ +RUN curl --fail --show-error --location --retry 5 --retry-all-errors --retry-delay 2 \ + -o /tmp/oh-my-zsh-install.sh https://raw.githubusercontent.com/ohmyzsh/ohmyzsh/master/tools/install.sh \ + && sh /tmp/oh-my-zsh-install.sh "" --unattended \ + && rm /tmp/oh-my-zsh-install.sh \ && git clone --depth 1 https://github.com/zsh-users/zsh-autosuggestions ${ZSH_CUSTOM:-/root/.oh-my-zsh/custom}/plugins/zsh-autosuggestions \ && git clone --depth 1 https://github.com/zsh-users/zsh-syntax-highlighting.git ${ZSH_CUSTOM:-/root/.oh-my-zsh/custom}/plugins/zsh-syntax-highlighting @@ -462,7 +473,11 @@ COPY sgl-model-gateway /build/sgl-model-gateway # Install Rust, build gateway binary and Python bindings, then clean up Rust toolchain RUN --mount=type=cache,target=/root/.cache/pip \ - curl --proto '=https' --tlsv1.2 --retry 3 --retry-delay 2 -sSf https://sh.rustup.rs | sh -s -- -y \ + curl --proto '=https' --tlsv1.2 --fail --show-error --location \ + --retry 5 --retry-all-errors --retry-delay 2 \ + -o /tmp/rustup-init.sh https://sh.rustup.rs \ + && sh /tmp/rustup-init.sh -y \ + && rm /tmp/rustup-init.sh \ && export PATH="/root/.cargo/bin:${PATH}" \ && python3 -m pip install maturin \ && cd /build/sgl-model-gateway/bindings/python \ @@ -792,16 +807,24 @@ WORKDIR /sgl-workspace/sglang # Keep build provenance at the end so metadata changes do not invalidate build layers. ARG SGLANG_BUILD_COMMIT=unknown +ARG SGLANG_BUILD_TREE=unknown +ARG SGLANG_PYTHON_MANIFEST_SHA256=unknown +ARG SGLANG_BUILD_SOURCE=https://github.com/sgl-project/sglang ARG SGLANG_BUILD_URL= ARG SGLANG_IMAGE_TAG=local/sglang:dev ENV SGLANG_BUILD_COMMIT=${SGLANG_BUILD_COMMIT:-unknown} \ + SGLANG_BUILD_TREE=${SGLANG_BUILD_TREE:-unknown} \ + SGLANG_PYTHON_MANIFEST_SHA256=${SGLANG_PYTHON_MANIFEST_SHA256:-unknown} \ + SGLANG_BUILD_SOURCE=${SGLANG_BUILD_SOURCE:-https://github.com/sgl-project/sglang} \ SGLANG_BUILD_URL=${SGLANG_BUILD_URL:-} \ SGLANG_IMAGE_TAG=${SGLANG_IMAGE_TAG:-local/sglang:dev} -LABEL org.opencontainers.image.source="https://github.com/sgl-project/sglang" \ +LABEL org.opencontainers.image.source="${SGLANG_BUILD_SOURCE}" \ org.opencontainers.image.revision="${SGLANG_BUILD_COMMIT}" \ org.opencontainers.image.version="${SGLANG_IMAGE_TAG}" \ org.opencontainers.image.url="${SGLANG_BUILD_URL}" \ ai.sglang.build.commit="${SGLANG_BUILD_COMMIT}" \ + ai.sglang.build.tree="${SGLANG_BUILD_TREE}" \ + ai.sglang.build.python-manifest-sha256="${SGLANG_PYTHON_MANIFEST_SHA256}" \ ai.sglang.build.url="${SGLANG_BUILD_URL}" \ ai.sglang.image.tag="${SGLANG_IMAGE_TAG}" @@ -887,7 +910,8 @@ RUN --mount=type=cache,target=/var/cache/apt,id=runtime-apt \ && update-alternatives --install /usr/bin/python3 python3 /usr/bin/python3.12 2 \ && update-alternatives --set python3 /usr/bin/python3.12 \ && ln -sf /usr/bin/python3.12 /usr/bin/python \ - && wget -q https://bootstrap.pypa.io/get-pip.py \ + && curl --fail --show-error --location --retry 5 --retry-all-errors --retry-delay 2 \ + -o get-pip.py https://bootstrap.pypa.io/get-pip.py \ && python3 get-pip.py --break-system-packages \ && rm get-pip.py \ # Allow pip to install packages globally (PEP 668 workaround for Ubuntu 24.04) @@ -962,16 +986,24 @@ WORKDIR /sgl-workspace/sglang # Keep build provenance at the end so metadata changes do not invalidate build layers. ARG SGLANG_BUILD_COMMIT=unknown +ARG SGLANG_BUILD_TREE=unknown +ARG SGLANG_PYTHON_MANIFEST_SHA256=unknown +ARG SGLANG_BUILD_SOURCE=https://github.com/sgl-project/sglang ARG SGLANG_BUILD_URL= ARG SGLANG_IMAGE_TAG=local/sglang:dev ENV SGLANG_BUILD_COMMIT=${SGLANG_BUILD_COMMIT:-unknown} \ + SGLANG_BUILD_TREE=${SGLANG_BUILD_TREE:-unknown} \ + SGLANG_PYTHON_MANIFEST_SHA256=${SGLANG_PYTHON_MANIFEST_SHA256:-unknown} \ + SGLANG_BUILD_SOURCE=${SGLANG_BUILD_SOURCE:-https://github.com/sgl-project/sglang} \ SGLANG_BUILD_URL=${SGLANG_BUILD_URL:-} \ SGLANG_IMAGE_TAG=${SGLANG_IMAGE_TAG:-local/sglang:dev} -LABEL org.opencontainers.image.source="https://github.com/sgl-project/sglang" \ +LABEL org.opencontainers.image.source="${SGLANG_BUILD_SOURCE}" \ org.opencontainers.image.revision="${SGLANG_BUILD_COMMIT}" \ org.opencontainers.image.version="${SGLANG_IMAGE_TAG}" \ org.opencontainers.image.url="${SGLANG_BUILD_URL}" \ ai.sglang.build.commit="${SGLANG_BUILD_COMMIT}" \ + ai.sglang.build.tree="${SGLANG_BUILD_TREE}" \ + ai.sglang.build.python-manifest-sha256="${SGLANG_PYTHON_MANIFEST_SHA256}" \ ai.sglang.build.url="${SGLANG_BUILD_URL}" \ ai.sglang.image.tag="${SGLANG_IMAGE_TAG}" diff --git a/scripts/ci/amd/check_vram_clear.sh b/scripts/ci/amd/check_vram_clear.sh index 51e5a915fad3..43acd7037bde 100755 --- a/scripts/ci/amd/check_vram_clear.sh +++ b/scripts/ci/amd/check_vram_clear.sh @@ -17,7 +17,10 @@ check_vram_clear() { echo "✓ VRAM usage is within acceptable limits on all GPUs" return 0 fi - fi + fi + + echo "ERROR: rocm-smi is not available; cannot verify that VRAM is clear." >&2 + return 1 } # If this script is run directly (not sourced), run the check diff --git a/scripts/ci/get_volcengine_image_tag.py b/scripts/ci/get_volcengine_image_tag.py new file mode 100755 index 000000000000..b1537a340ce9 --- /dev/null +++ b/scripts/ci/get_volcengine_image_tag.py @@ -0,0 +1,83 @@ +#!/usr/bin/env python3 + +from __future__ import annotations + +import argparse +import re +import subprocess +from datetime import datetime +from pathlib import Path +from zoneinfo import ZoneInfo + + +def get_sglang_version() -> str: + repo_root = Path(__file__).resolve().parents[2] + version_file = repo_root / "python/sglang/_version.py" + if version_file.exists(): + content = version_file.read_text() + match = re.search(r"__version__\s*=\s*version\s*=\s*'([^']+)'", content) + if match: + return match.group(1) + + result = subprocess.run( + ["python3", "python/tools/get_version_tag.py", "--tag-only"], + cwd=repo_root, + capture_output=True, + text=True, + ) + if result.returncode == 0 and result.stdout.strip(): + return result.stdout.strip().lstrip("v") + + raise SystemExit( + "failed to extract sglang version from python/sglang/_version.py or git tags" + ) + + +def main() -> None: + parser = argparse.ArgumentParser( + description="Generate Volcengine CR image tags for fork workflows." + ) + parser.add_argument( + "--mode", choices=["manual", "nightly", "version"], required=True + ) + parser.add_argument( + "--tag-value", + default="", + help="Required for version mode; inserted after .byted.", + ) + parser.add_argument("--cuda-suffix", choices=["", "cu129", "cu130"], default="") + parser.add_argument( + "--variant-suffix", + default="", + help="Optional build variant suffix appended before the CUDA suffix.", + ) + args = parser.parse_args() + + if args.variant_suffix and not re.fullmatch( + r"[0-9A-Za-z][0-9A-Za-z_.-]*", args.variant_suffix + ): + raise SystemExit("--variant-suffix must be a Docker tag-safe suffix") + + version = get_sglang_version() + timestamp = datetime.now(ZoneInfo("Asia/Shanghai")).strftime("%Y%m%d%H%M") + + if args.mode == "manual": + tag = f"v{version}.iaas.dev.{timestamp}" + elif args.mode == "nightly": + tag = f"v{version}.iaas.nightly.{timestamp}" + else: + if not args.tag_value: + raise SystemExit("--tag-value is required when --mode=version") + tag = f"v{version}.byted.{args.tag_value}.{timestamp}" + + if args.variant_suffix: + tag = f"{tag}-{args.variant_suffix}" + + if args.cuda_suffix: + tag = f"{tag}-{args.cuda_suffix}" + + print(tag) + + +if __name__ == "__main__": + main() diff --git a/scripts/ci/sync_docker_images_to_volcengine.py b/scripts/ci/sync_docker_images_to_volcengine.py new file mode 100755 index 000000000000..f5bd5e6e0a35 --- /dev/null +++ b/scripts/ci/sync_docker_images_to_volcengine.py @@ -0,0 +1,437 @@ +#!/usr/bin/env python3 + +from __future__ import annotations + +import argparse +import json +import re +import subprocess +from dataclasses import dataclass +from datetime import date, datetime +from typing import Iterable +from urllib.parse import quote +from urllib.request import urlopen +from zoneinfo import ZoneInfo + +TAG_RE = re.compile(r"^[A-Za-z0-9_][A-Za-z0-9_.-]{0,127}$") +VERSION_TAG_RE = re.compile(r"^v(\d+)\.(\d+)\.(\d+)(?:\.post(\d+))?([A-Za-z0-9_.-]*)?$") +VLLM_UBUNTU2404_VERSION_TAG_RE = re.compile(r"^v(\d+)\.(\d+)\.(\d+)-ubuntu2404$") +VLLM_UBUNTU2404_NIGHTLY_TAG_RE = re.compile( + r"^(nightly|nightly-[0-9a-f]{7,64})-ubuntu2404$" +) +AUTO_TAG_SPECS = {"version", "today-nightly"} +DEFAULT_VARIANT_RE = re.compile( + r"^(latest|dev|nightly|nightly-[0-9a-f]{7,64}|nightly-dev-[0-9]{8}-[0-9a-f]{7,64}|v\d+\.\d+\.\d+(?:\.post\d+)?)$" +) + + +@dataclass(frozen=True) +class SyncItem: + source: str + destination: str + + +@dataclass(frozen=True) +class DockerTag: + name: str + last_updated: str + + +def parse_tags(value: str) -> list[str]: + tags = [tag.strip() for tag in re.split(r"[,\n]+", value) if tag.strip()] + if not tags: + raise SystemExit("at least one image tag is required") + for tag in tags: + if tag not in AUTO_TAG_SPECS and not TAG_RE.fullmatch(tag): + raise SystemExit(f"invalid Docker tag: {tag}") + return tags + + +def docker_hub_repository(image_name: str) -> str: + image_name = normalize_image_name(image_name, "source image") + if image_name.startswith("docker.io/"): + image_name = image_name.removeprefix("docker.io/") + parts = image_name.split("/") + if len(parts) == 1: + return f"library/{parts[0]}" + return "/".join(parts) + + +def fetch_docker_hub_tags(image_name: str, *, pages: int = 5) -> list[DockerTag]: + repository = docker_hub_repository(image_name) + page_url = ( + "https://hub.docker.com/v2/repositories/" + f"{quote(repository, safe='/')}/tags?page_size=100&ordering=last_updated" + ) + tags: list[DockerTag] = [] + for _ in range(pages): + with urlopen(page_url, timeout=30) as response: + payload = json.load(response) + for item in payload.get("results", []): + name = item.get("name") + last_updated = item.get("last_updated") + if name and last_updated: + tags.append(DockerTag(name=name, last_updated=last_updated)) + page_url = payload.get("next") + if not page_url: + break + return tags + + +def version_key(tag_name: str) -> tuple[int, int, int, int] | None: + match = VERSION_TAG_RE.fullmatch(tag_name) + if not match: + return None + major, minor, patch, post, suffix = match.groups() + if suffix and not suffix.startswith("-"): + return None + return (int(major), int(minor), int(patch), int(post or 0)) + + +def latest_version_tags(tags: Iterable[DockerTag]) -> list[str]: + names = {tag.name for tag in tags} + keyed_versions = [ + (key, name) + for name in names + if DEFAULT_VARIANT_RE.fullmatch(name) + and (key := version_key(name.removesuffix(""))) is not None + ] + if not keyed_versions: + raise SystemExit("no version image tags found in source repository") + + latest_key = max(key for key, _ in keyed_versions) + version_names = sorted(name for key, name in keyed_versions if key == latest_key) + latest_aliases = ["latest"] if "latest" in names else [] + return sorted(set(latest_aliases + version_names)) + + +def tag_updated_date(tag: DockerTag, timezone: ZoneInfo) -> date: + value = tag.last_updated + if value.endswith("Z"): + value = value[:-1] + "+00:00" + if "." in value: + prefix, suffix = value.split(".", 1) + match = re.fullmatch(r"(\d+)(.*)", suffix) + if not match: + raise SystemExit( + f"invalid Docker Hub last_updated timestamp: {tag.last_updated}" + ) + fraction, offset = match.groups() + fraction = fraction[:6] + fraction = fraction.ljust(6, "0") + value = f"{prefix}.{fraction}{offset}" + updated = datetime.fromisoformat(value) + return updated.astimezone(timezone).date() + + +def today_nightly_tags( + tags: Iterable[DockerTag], + *, + today: date, + daily_aliases: set[str] | None = None, + timezone: ZoneInfo | None = None, +) -> list[str]: + timezone = timezone or ZoneInfo("Asia/Shanghai") + daily_aliases = daily_aliases or set() + selected = { + tag.name + for tag in tags + if tag_updated_date(tag, timezone) == today + and ("nightly" in tag.name or tag.name in daily_aliases) + and DEFAULT_VARIANT_RE.fullmatch(tag.name) + } + if not selected: + raise SystemExit(f"no today-nightly image tags found for {today.isoformat()}") + return sorted(selected) + + +def resolve_tag_specs( + specs: Iterable[str], + docker_tags: Iterable[DockerTag], + *, + today: date, + daily_aliases: set[str] | None = None, + timezone: ZoneInfo | None = None, +) -> list[str]: + docker_tags = list(docker_tags) + resolved: list[str] = [] + for spec in specs: + if spec == "version": + resolved.extend(latest_version_tags(docker_tags)) + elif spec == "today-nightly": + resolved.extend( + today_nightly_tags( + docker_tags, + today=today, + daily_aliases=daily_aliases, + timezone=timezone, + ) + ) + else: + resolved.append(spec) + + deduped: list[str] = [] + seen: set[str] = set() + for tag in resolved: + if tag not in seen: + deduped.append(tag) + seen.add(tag) + return deduped + + +def vllm_ubuntu2404_version_key(tag_name: str) -> tuple[int, int, int] | None: + match = VLLM_UBUNTU2404_VERSION_TAG_RE.fullmatch(tag_name) + if not match: + return None + major, minor, patch = match.groups() + return (int(major), int(minor), int(patch)) + + +def latest_vllm_version_tags(tags: Iterable[DockerTag]) -> list[str]: + names = {tag.name for tag in tags} + keyed_versions = [ + (key, name) + for name in names + if (key := vllm_ubuntu2404_version_key(name)) is not None + ] + if not keyed_versions: + return latest_version_tags(tags) + + latest_key = max(key for key, _ in keyed_versions) + version_names = sorted(name for key, name in keyed_versions if key == latest_key) + latest_aliases = ["latest-ubuntu2404"] if "latest-ubuntu2404" in names else [] + return sorted(set(latest_aliases + version_names)) + + +def today_vllm_nightly_tags( + tags: Iterable[DockerTag], + *, + today: date, + timezone: ZoneInfo | None = None, +) -> list[str]: + tags = list(tags) + timezone = timezone or ZoneInfo("Asia/Shanghai") + ubuntu2404_tags = sorted( + { + tag.name + for tag in tags + if tag_updated_date(tag, timezone) == today + and VLLM_UBUNTU2404_NIGHTLY_TAG_RE.fullmatch(tag.name) + } + ) + if ubuntu2404_tags: + return ubuntu2404_tags + return today_nightly_tags(tags, today=today, timezone=timezone) + + +def resolve_vllm_tag_specs( + specs: Iterable[str], + docker_tags: Iterable[DockerTag], + *, + today: date, + timezone: ZoneInfo | None = None, +) -> list[str]: + docker_tags = list(docker_tags) + resolved: list[str] = [] + for spec in specs: + if spec == "version": + resolved.extend(latest_vllm_version_tags(docker_tags)) + elif spec == "today-nightly": + resolved.extend( + today_vllm_nightly_tags( + docker_tags, + today=today, + timezone=timezone, + ) + ) + else: + resolved.append(spec) + + deduped: list[str] = [] + seen: set[str] = set() + for tag in resolved: + if tag not in seen: + deduped.append(tag) + seen.add(tag) + return deduped + + +def normalize_image_name(value: str, field_name: str) -> str: + value = value.strip().rstrip("/") + if not value: + raise SystemExit(f"{field_name} is required") + if ":" in value.rsplit("/", 1)[-1]: + raise SystemExit(f"{field_name} must not include a tag: {value}") + if value.startswith("/") or "//" in value: + raise SystemExit(f"invalid {field_name}: {value}") + return value + + +def normalize_repository(value: str, field_name: str) -> str: + value = value.strip().strip("/") + if not value: + raise SystemExit(f"{field_name} is required") + if ":" in value or "//" in value: + raise SystemExit(f"invalid {field_name}: {value}") + return value + + +def build_sync_plan( + *, + registry: str, + namespace: str, + sglang_source: str, + sglang_repository: str, + sglang_tags: Iterable[str], + vllm_source: str, + vllm_repository: str, + vllm_tags: Iterable[str], +) -> list[SyncItem]: + registry = normalize_repository(registry, "registry") + namespace = normalize_repository(namespace, "namespace") + sglang_source = normalize_image_name(sglang_source, "sglang source image") + sglang_repository = normalize_repository( + sglang_repository, "sglang destination repository" + ) + vllm_source = normalize_image_name(vllm_source, "vllm source image") + vllm_repository = normalize_repository( + vllm_repository, "vllm destination repository" + ) + + plan: list[SyncItem] = [] + for source_image, destination_repo, tags in ( + (sglang_source, sglang_repository, sglang_tags), + (vllm_source, vllm_repository, vllm_tags), + ): + for tag in tags: + if not TAG_RE.fullmatch(tag): + raise SystemExit(f"invalid Docker tag: {tag}") + plan.append( + SyncItem( + source=f"{source_image}:{tag}", + destination=f"{registry}/{namespace}/{destination_repo}:{tag}", + ) + ) + return plan + + +def build_imagetools_command(item: SyncItem, *, platform: str = "") -> list[str]: + command = ["docker", "buildx", "imagetools", "create"] + if platform: + command.extend(["--platform", platform]) + command.extend(["-t", item.destination, item.source]) + return command + + +def sync_images(plan: Iterable[SyncItem], *, execute: bool, platform: str = "") -> None: + for item in plan: + command = build_imagetools_command(item, platform=platform) + print(f"{item.source} -> {item.destination}", flush=True) + if execute: + subprocess.run(command, check=True) + else: + print("+ " + " ".join(command), flush=True) + + +def main() -> None: + parser = argparse.ArgumentParser( + description="Sync public SGLang and vLLM Docker image tags to Volcengine CR." + ) + parser.add_argument("--registry", required=True, help="Volcengine CR registry host") + parser.add_argument("--namespace", required=True, help="Volcengine CR namespace") + parser.add_argument( + "--sglang-source", + default="docker.io/lmsysorg/sglang", + help="Source SGLang image repository without tag", + ) + parser.add_argument( + "--sglang-repository", + default="sglang", + help="Destination SGLang repository inside the Volcengine CR namespace", + ) + parser.add_argument( + "--sglang-tags", + default="version,today-nightly", + help=( + "Comma or newline separated SGLang tags to sync. Special values: " + "version, today-nightly." + ), + ) + parser.add_argument( + "--vllm-source", + default="docker.io/vllm/vllm-openai", + help="Source vLLM image repository without tag", + ) + parser.add_argument( + "--vllm-repository", + default="vllm", + help="Destination vLLM repository inside the Volcengine CR namespace", + ) + parser.add_argument( + "--vllm-tags", + default="version,today-nightly", + help=( + "Comma or newline separated vLLM tags to sync. Special values: " + "version, today-nightly. The automatic vLLM selector prefers " + "ubuntu2404 tags and falls back to unsuffixed tags when no " + "ubuntu2404 tag exists for that image family." + ), + ) + parser.add_argument( + "--execute", + action="store_true", + help="Run docker buildx imagetools create. Without this, print a dry run.", + ) + parser.add_argument( + "--timezone", + default="Asia/Shanghai", + help="Timezone used to decide which Docker Hub tags were updated today.", + ) + parser.add_argument( + "--platform", + default="linux/amd64", + help="Optional platform filter passed to docker buildx imagetools create.", + ) + args = parser.parse_args() + + timezone = ZoneInfo(args.timezone) + today = datetime.now(timezone).date() + sglang_tag_specs = parse_tags(args.sglang_tags) + vllm_tag_specs = parse_tags(args.vllm_tags) + sglang_docker_tags = ( + fetch_docker_hub_tags(args.sglang_source) + if AUTO_TAG_SPECS.intersection(sglang_tag_specs) + else [] + ) + vllm_docker_tags = ( + fetch_docker_hub_tags(args.vllm_source) + if AUTO_TAG_SPECS.intersection(vllm_tag_specs) + else [] + ) + + plan = build_sync_plan( + registry=args.registry, + namespace=args.namespace, + sglang_source=args.sglang_source, + sglang_repository=args.sglang_repository, + sglang_tags=resolve_tag_specs( + sglang_tag_specs, + sglang_docker_tags, + today=today, + daily_aliases={"dev", "dev-cu12", "dev-cu13"}, + timezone=timezone, + ), + vllm_source=args.vllm_source, + vllm_repository=args.vllm_repository, + vllm_tags=resolve_vllm_tag_specs( + vllm_tag_specs, + vllm_docker_tags, + today=today, + timezone=timezone, + ), + ) + sync_images(plan, execute=args.execute, platform=args.platform) + + +if __name__ == "__main__": + main() diff --git a/scripts/ci/test_sync_docker_images_to_volcengine.py b/scripts/ci/test_sync_docker_images_to_volcengine.py new file mode 100755 index 000000000000..6c0156bc81e4 --- /dev/null +++ b/scripts/ci/test_sync_docker_images_to_volcengine.py @@ -0,0 +1,249 @@ +#!/usr/bin/env python3 + +from __future__ import annotations + +import sys +import unittest +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parent)) + +from datetime import date +from zoneinfo import ZoneInfo + +from sync_docker_images_to_volcengine import ( + DockerTag, + SyncItem, + build_imagetools_command, + build_sync_plan, + parse_tags, + resolve_tag_specs, + resolve_vllm_tag_specs, + tag_updated_date, +) + + +class SyncDockerImagesToVolcengineTest(unittest.TestCase): + def test_builds_default_sglang_and_vllm_latest_plan(self) -> None: + plan = build_sync_plan( + registry="iaas-gpu-cn-beijing.cr.volces.com", + namespace="serving", + sglang_source="docker.io/lmsysorg/sglang", + sglang_repository="sglang", + sglang_tags=["latest"], + vllm_source="docker.io/vllm/vllm-openai", + vllm_repository="vllm", + vllm_tags=["latest"], + ) + + self.assertEqual( + [(item.source, item.destination) for item in plan], + [ + ( + "docker.io/lmsysorg/sglang:latest", + "iaas-gpu-cn-beijing.cr.volces.com/serving/sglang:latest", + ), + ( + "docker.io/vllm/vllm-openai:latest", + "iaas-gpu-cn-beijing.cr.volces.com/serving/vllm:latest", + ), + ], + ) + + def test_parse_tags_accepts_commas_and_newlines(self) -> None: + self.assertEqual( + parse_tags("latest, nightly\nlatest-cu130"), + ["latest", "nightly", "latest-cu130"], + ) + + def test_resolves_default_latest_version_tags_only(self) -> None: + tags = [ + DockerTag("nightly-dev-20260616-abcdef1", "2026-06-16T01:53:18Z"), + DockerTag("latest", "2026-06-15T08:39:02Z"), + DockerTag("v0.5.13.post1", "2026-06-15T08:39:01Z"), + DockerTag("latest-cu130", "2026-06-15T08:39:06Z"), + DockerTag("v0.5.13.post1-cu130", "2026-06-15T08:39:04Z"), + DockerTag("v0.5.12", "2026-05-10T08:00:00Z"), + DockerTag("latest-runtime", "2026-06-15T09:25:15Z"), + DockerTag("v0.5.13.post1-runtime", "2026-06-15T09:25:13Z"), + ] + + self.assertEqual( + resolve_tag_specs(["version"], tags, today=date(2026, 6, 16)), + [ + "latest", + "v0.5.13.post1", + ], + ) + + def test_resolves_default_today_nightly_tags_only(self) -> None: + tags = [ + DockerTag("nightly-dev-cu13-20260616-abcdef1", "2026-06-16T01:53:20Z"), + DockerTag("nightly-dev-20260616-abcdef1", "2026-06-16T01:53:18Z"), + DockerTag("dev-cu13", "2026-06-16T01:53:17Z"), + DockerTag("dev", "2026-06-16T01:53:15Z"), + DockerTag("nightly-dev-cu12-20260615-old", "2026-06-15T01:53:12Z"), + DockerTag("latest", "2026-06-15T08:39:02Z"), + ] + + self.assertEqual( + resolve_tag_specs( + ["today-nightly"], + tags, + today=date(2026, 6, 16), + daily_aliases={"dev"}, + ), + [ + "dev", + "nightly-dev-20260616-abcdef1", + ], + ) + + def test_resolves_default_vllm_today_nightly_tags_only(self) -> None: + tags = [ + DockerTag("cu129-nightly-abcdef1", "2026-06-16T06:28:48Z"), + DockerTag("cu129-nightly", "2026-06-16T06:28:46Z"), + DockerTag("nightly-abcdef1", "2026-06-16T06:15:25Z"), + DockerTag("nightly", "2026-06-16T06:15:24Z"), + DockerTag("nightly-aarch64", "2026-06-16T06:15:21Z"), + DockerTag("nightly-x86_64", "2026-06-16T06:05:25Z"), + ] + + self.assertEqual( + resolve_tag_specs(["today-nightly"], tags, today=date(2026, 6, 16)), + [ + "nightly", + "nightly-abcdef1", + ], + ) + + def test_resolves_vllm_ubuntu2404_version_tags_only(self) -> None: + tags = [ + DockerTag("latest-ubuntu2404", "2026-06-13T01:49:50Z"), + DockerTag("v0.23.0-ubuntu2404", "2026-06-13T01:49:52Z"), + DockerTag("latest", "2026-06-13T00:36:44Z"), + DockerTag("v0.23.0", "2026-06-13T00:36:45Z"), + DockerTag("latest-x86_64-ubuntu2404", "2026-06-13T01:39:39Z"), + DockerTag("v0.23.0-x86_64-ubuntu2404", "2026-06-13T01:39:41Z"), + DockerTag("latest-cu129-ubuntu2404", "2026-06-13T02:30:42Z"), + DockerTag("v0.23.0-cu129-ubuntu2404", "2026-06-13T02:30:43Z"), + DockerTag("v0.22.1-ubuntu2404", "2026-06-05T08:34:55Z"), + DockerTag("nightly", "2026-06-17T06:16:46Z"), + ] + + self.assertEqual( + resolve_vllm_tag_specs(["version"], tags, today=date(2026, 6, 17)), + [ + "latest-ubuntu2404", + "v0.23.0-ubuntu2404", + ], + ) + + def test_resolves_vllm_default_version_tags_when_ubuntu2404_is_unavailable( + self, + ) -> None: + tags = [ + DockerTag("latest", "2026-06-13T00:36:44Z"), + DockerTag("v0.23.0", "2026-06-13T00:36:45Z"), + DockerTag("latest-x86_64", "2026-06-13T01:39:39Z"), + DockerTag("v0.23.0-x86_64", "2026-06-13T01:39:41Z"), + DockerTag("latest-cu129", "2026-06-13T02:30:42Z"), + DockerTag("v0.23.0-cu129", "2026-06-13T02:30:43Z"), + DockerTag("v0.22.1", "2026-06-05T08:34:55Z"), + ] + + self.assertEqual( + resolve_vllm_tag_specs(["version"], tags, today=date(2026, 6, 17)), + [ + "latest", + "v0.23.0", + ], + ) + + def test_resolves_vllm_ubuntu2404_version_and_default_today_nightly_tags( + self, + ) -> None: + tags = [ + DockerTag("latest-ubuntu2404", "2026-06-13T01:49:50Z"), + DockerTag("v0.23.0-ubuntu2404", "2026-06-13T01:49:52Z"), + DockerTag("latest", "2026-06-13T00:36:44Z"), + DockerTag("v0.23.0", "2026-06-13T00:36:45Z"), + DockerTag("cu129-nightly-abcdef1", "2026-06-17T06:28:48Z"), + DockerTag("cu129-nightly", "2026-06-17T06:28:46Z"), + DockerTag("nightly-abcdef1", "2026-06-17T06:15:25Z"), + DockerTag("nightly", "2026-06-17T06:15:24Z"), + DockerTag("nightly-aarch64", "2026-06-17T06:15:21Z"), + DockerTag("nightly-x86_64", "2026-06-17T06:05:25Z"), + ] + + self.assertEqual( + resolve_vllm_tag_specs( + ["version", "today-nightly"], tags, today=date(2026, 6, 17) + ), + [ + "latest-ubuntu2404", + "v0.23.0-ubuntu2404", + "nightly", + "nightly-abcdef1", + ], + ) + + def test_tag_updated_date_accepts_non_six_digit_fraction(self) -> None: + self.assertEqual( + tag_updated_date( + DockerTag("dev-cu13", "2026-06-16T01:53:17.12614Z"), + timezone=ZoneInfo("Asia/Shanghai"), + ), + date(2026, 6, 16), + ) + + def test_imagetools_command_filters_platform_when_requested(self) -> None: + self.assertEqual( + build_imagetools_command( + SyncItem( + source="docker.io/lmsysorg/sglang:latest", + destination="iaas-gpu-cn-beijing.cr.volces.com/serving/sglang:latest", + ), + platform="linux/amd64", + ), + [ + "docker", + "buildx", + "imagetools", + "create", + "--platform", + "linux/amd64", + "-t", + "iaas-gpu-cn-beijing.cr.volces.com/serving/sglang:latest", + "docker.io/lmsysorg/sglang:latest", + ], + ) + + def test_rejects_missing_registry_or_namespace(self) -> None: + with self.assertRaisesRegex(SystemExit, "registry is required"): + build_sync_plan( + registry="", + namespace="serving", + sglang_source="docker.io/lmsysorg/sglang", + sglang_repository="sglang", + sglang_tags=["latest"], + vllm_source="docker.io/vllm/vllm-openai", + vllm_repository="vllm", + vllm_tags=["latest"], + ) + + with self.assertRaisesRegex(SystemExit, "namespace is required"): + build_sync_plan( + registry="iaas-gpu-cn-beijing.cr.volces.com", + namespace="", + sglang_source="docker.io/lmsysorg/sglang", + sglang_repository="sglang", + sglang_tags=["latest"], + vllm_source="docker.io/vllm/vllm-openai", + vllm_repository="vllm", + vllm_tags=["latest"], + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/scripts/ci/utils/docker_build_metadata_args.py b/scripts/ci/utils/docker_build_metadata_args.py index 79a41a656d22..e5e7bb04942d 100644 --- a/scripts/ci/utils/docker_build_metadata_args.py +++ b/scripts/ci/utils/docker_build_metadata_args.py @@ -48,6 +48,9 @@ def build_arg_tokens( image_repo: str, version: str, build_commit: str, + build_tree: str, + python_manifest_sha256: str, + build_source: str, build_url: str, date: str, short_sha: str, @@ -55,6 +58,9 @@ def build_arg_tokens( image_tag = select_tag(tag_config, cuda, version, date, short_sha) build_args = { "SGLANG_BUILD_COMMIT": build_commit, + "SGLANG_BUILD_TREE": build_tree, + "SGLANG_PYTHON_MANIFEST_SHA256": python_manifest_sha256, + "SGLANG_BUILD_SOURCE": build_source, "SGLANG_BUILD_URL": build_url, "SGLANG_IMAGE_TAG": f"{image_repo}:{image_tag}", } @@ -78,6 +84,19 @@ def parse_args() -> argparse.Namespace: required=True, help="Commit checked out for the Docker build.", ) + parser.add_argument( + "--build-tree", required=True, help="Git tree checked out for the Docker build." + ) + parser.add_argument( + "--python-manifest-sha256", + required=True, + help="Deterministic digest of the tracked Python source manifest.", + ) + parser.add_argument( + "--build-source", + required=True, + help="Repository URL for the checked-out source.", + ) parser.add_argument("--build-url", default="", help="CI run URL.") parser.add_argument( "--date", @@ -103,6 +122,9 @@ def main() -> int: image_repo=args.image_repo, version=args.sgl_version, build_commit=args.build_commit, + build_tree=args.build_tree, + python_manifest_sha256=args.python_manifest_sha256, + build_source=args.build_source, build_url=args.build_url, date=args.date, short_sha=short_sha, diff --git a/scripts/ci/verify_private_image_runtime.py b/scripts/ci/verify_private_image_runtime.py new file mode 100644 index 000000000000..6d3b7a4f0612 --- /dev/null +++ b/scripts/ci/verify_private_image_runtime.py @@ -0,0 +1,29 @@ +"""Verify private-delivery packages without initializing the EIC client. + +The EIC SDK aborts during import on a build runner without its runtime +environment. Real client initialization and read/write checks belong in a +configured task Pod; image CI only proves that the SDK module is installed. +""" + +import importlib.metadata +import importlib.util + +for distribution in ("sglang", "sgl-kernel", "eic", "onion-ai-data", "deep-ep"): + try: + version = importlib.metadata.version(distribution) + except importlib.metadata.PackageNotFoundError: + version = "not-installed-as-distribution" + print(f"{distribution}={version}") + +eic_spec = importlib.util.find_spec("eic") +if eic_spec is None: + raise SystemExit("EIC SDK module is not installed") +print(f"eic-module={eic_spec.origin}") + +import torch + +print(f"torch={torch.__version__} cuda={torch.version.cuda}") + +import sglang + +print(f"sglang-module={sglang.__file__}") diff --git a/scripts/code_sync/install_github_cli.sh b/scripts/code_sync/install_github_cli.sh index 2ef1db023952..6f8512408b7d 100755 --- a/scripts/code_sync/install_github_cli.sh +++ b/scripts/code_sync/install_github_cli.sh @@ -1,18 +1,31 @@ #!/bin/bash +set -euo pipefail # Check if gh is installed before attempting to install it -if ! command -v gh &> /dev/null -then -echo "GitHub CLI not found. Installing now..." -(type -p wget >/dev/null || ( apt update && apt install wget -y)) \ -&& mkdir -p -m 755 /etc/apt/keyrings \ -&& out=$(mktemp) && wget -nv -O$out https://cli.github.com/packages/githubcli-archive-keyring.gpg \ -&& cat $out | tee /etc/apt/keyrings/githubcli-archive-keyring.gpg > /dev/null \ -&& chmod go+r /etc/apt/keyrings/githubcli-archive-keyring.gpg \ -&& mkdir -p -m 755 /etc/apt/sources.list.d \ -&& echo "deb [arch=$(dpkg --print-architecture) signed-by=/etc/apt/keyrings/githubcli-archive-keyring.gpg] https://cli.github.com/packages stable main" | tee /etc/apt/sources.list.d/github-cli.list > /dev/null \ -&& apt update \ -&& apt install gh -y +if ! command -v gh >/dev/null 2>&1; then + echo "GitHub CLI not found. Installing now..." + + if ! command -v wget >/dev/null 2>&1; then + apt-get update + apt-get install -y wget + fi + + install -d -m 755 /etc/apt/keyrings /etc/apt/sources.list.d + keyring_tmp=$(mktemp) + trap 'rm -f "$keyring_tmp"' EXIT + + wget -nv -O "$keyring_tmp" \ + https://cli.github.com/packages/githubcli-archive-keyring.gpg + install -m 0644 "$keyring_tmp" \ + /etc/apt/keyrings/githubcli-archive-keyring.gpg + chmod go+r /etc/apt/keyrings/githubcli-archive-keyring.gpg + + printf '%s\n' \ + "deb [arch=$(dpkg --print-architecture) signed-by=/etc/apt/keyrings/githubcli-archive-keyring.gpg] https://cli.github.com/packages stable main" \ + >/etc/apt/sources.list.d/github-cli.list + + apt-get update + apt-get install -y gh else -echo "GitHub CLI is already installed. Skipping installation." + echo "GitHub CLI is already installed. Skipping installation." fi diff --git a/sgl-model-gateway/bindings/python/pyproject.toml b/sgl-model-gateway/bindings/python/pyproject.toml index c44b1b96abb4..a4628fdd5c54 100644 --- a/sgl-model-gateway/bindings/python/pyproject.toml +++ b/sgl-model-gateway/bindings/python/pyproject.toml @@ -13,7 +13,7 @@ authors = [ {name = "Byron Hsu", email = "byronhsu1230@gmail.com"} ] requires-python = ">=3.8" -readme = "../../README.md" +readme = "README.md" license = { text = "Apache-2.0" } classifiers = [ "Programming Language :: Python :: Implementation :: CPython", @@ -51,7 +51,7 @@ sglang-router = "sglang_router.cli:main" [tool.maturin] python-source = "src" module-name = "sglang_router.sglang_router_rs" -# Exclude bindings/python/README.md to use root README only +# Keep the README as package metadata, but do not include it in the wheel payload. exclude = ["README.md"] [tool.pytest.ini_options] diff --git a/test/registered/unit/tools/test_docker_build_metadata_args.py b/test/registered/unit/tools/test_docker_build_metadata_args.py index 7aca2444bda9..5c54063de39b 100644 --- a/test/registered/unit/tools/test_docker_build_metadata_args.py +++ b/test/registered/unit/tools/test_docker_build_metadata_args.py @@ -35,6 +35,9 @@ def run_helper( image_repo: str = "lmsysorg/sglang", version: str = "0.6.0", build_commit: str = "abcdef1234567890", + build_tree: str = "tree1234567890abcdef", + python_manifest_sha256: str = "manifest1234567890abcdef", + build_source: str = "https://github.com/bytedance-iaas/sglang", build_url: str = "https://github.com/sgl-project/sglang/actions/runs/1", date: str = "20260429", ) -> list[str]: @@ -52,6 +55,12 @@ def run_helper( version, "--build-commit", build_commit, + "--build-tree", + build_tree, + "--python-manifest-sha256", + python_manifest_sha256, + "--build-source", + build_source, "--build-url", build_url, "--date", @@ -87,6 +96,9 @@ def test_release_metadata_prefers_versioned_tag(self): self.build_args(args), { "SGLANG_BUILD_COMMIT": "abcdef1234567890", + "SGLANG_BUILD_TREE": "tree1234567890abcdef", + "SGLANG_PYTHON_MANIFEST_SHA256": "manifest1234567890abcdef", + "SGLANG_BUILD_SOURCE": "https://github.com/bytedance-iaas/sglang", "SGLANG_BUILD_URL": ( "https://github.com/sgl-project/sglang/actions/runs/1" ), @@ -174,16 +186,24 @@ def test_final_dockerfile_stages_embed_metadata_contract(self): for stage in (framework_stage, runtime_stage): for expected in ( "ARG SGLANG_BUILD_COMMIT=unknown", + "ARG SGLANG_BUILD_TREE=unknown", + "ARG SGLANG_PYTHON_MANIFEST_SHA256=unknown", + "ARG SGLANG_BUILD_SOURCE=https://github.com/sgl-project/sglang", "ARG SGLANG_BUILD_URL=", "ARG SGLANG_IMAGE_TAG=local/sglang:dev", "SGLANG_BUILD_COMMIT=${SGLANG_BUILD_COMMIT:-unknown}", + "SGLANG_BUILD_TREE=${SGLANG_BUILD_TREE:-unknown}", + "SGLANG_PYTHON_MANIFEST_SHA256=${SGLANG_PYTHON_MANIFEST_SHA256:-unknown}", + "SGLANG_BUILD_SOURCE=${SGLANG_BUILD_SOURCE:-https://github.com/sgl-project/sglang}", "SGLANG_BUILD_URL=${SGLANG_BUILD_URL:-}", "SGLANG_IMAGE_TAG=${SGLANG_IMAGE_TAG:-local/sglang:dev}", - 'org.opencontainers.image.source="https://github.com/sgl-project/sglang"', + 'org.opencontainers.image.source="${SGLANG_BUILD_SOURCE}"', 'org.opencontainers.image.revision="${SGLANG_BUILD_COMMIT}"', 'org.opencontainers.image.version="${SGLANG_IMAGE_TAG}"', 'org.opencontainers.image.url="${SGLANG_BUILD_URL}"', 'ai.sglang.build.commit="${SGLANG_BUILD_COMMIT}"', + 'ai.sglang.build.tree="${SGLANG_BUILD_TREE}"', + 'ai.sglang.build.python-manifest-sha256="${SGLANG_PYTHON_MANIFEST_SHA256}"', 'ai.sglang.build.url="${SGLANG_BUILD_URL}"', 'ai.sglang.image.tag="${SGLANG_IMAGE_TAG}"', ): From a312850ecb2c5d33af0394a780cd24362dece0e7 Mon Sep 17 00:00:00 2001 From: Hank Han Date: Thu, 13 Aug 2026 15:06:52 +0800 Subject: [PATCH 16/47] ci: add zstd and nydus image formats to private delivery build Add opt-in zstd (layer compression) and nydus (lazy-loading) image formats to the SGLang private delivery build, on top of the existing gzip OCI output. Formats are selected via a CSV `image_formats` input (default `oci`, so production callers are byte-for-byte unchanged) and tagged with `-zstd` / `-nydus` suffixes after the cuda suffix. - get_volcengine_image_tag.py: extract pure build_tag()/validate_suffix() helpers and add `--format-suffix` (appended after the cuda suffix); covered by scripts/ci/test_get_volcengine_image_tag.py. - _docker-build-and-publish.yml: add `image_formats` CSV input; build once then derive zstd via a cache-hit buildx re-export (compression=zstd,force-compression=true,oci-mediatypes=true) and nydus via `nydusify convert` from the pushed digest. Each format is gated by contains(inputs.image_formats, ...). base/zstd reuse the docker pull+run provenance check (falling back to a manifest media-type check when the daemon lacks zstd); nydus uses `nydusify check` + manifest media-type/annotation assertions (no run). - release-docker-dev.yml: wire `private_debug_image_formats` into build-dev-debug-base so one bounded debug job exercises all three formats; production callers keep the default oci. Also harden all framework_final egress fetches in docker/Dockerfile and the nydus tooling/convert steps behind an outer retry() loop (github tarballs, flashinfer pip index, nydus static tarball, just/oh-my-zsh installers, git clones, and nydusify convert), so intermittent internal proxy failures (curl 56 / 504 / closed pipe) are retried instead of hard-failing. These robustness changes are orthogonal to the format feature and also benefit the default oci production build. Co-authored-by: TRAE CLI --- .../workflows/_docker-build-and-publish.yml | 254 ++++++++++++++++++ .github/workflows/release-docker-dev.yml | 5 + docker/Dockerfile | 132 +++++---- scripts/ci/get_volcengine_image_tag.py | 78 ++++-- scripts/ci/test_get_volcengine_image_tag.py | 101 +++++++ 5 files changed, 506 insertions(+), 64 deletions(-) create mode 100644 scripts/ci/test_get_volcengine_image_tag.py diff --git a/.github/workflows/_docker-build-and-publish.yml b/.github/workflows/_docker-build-and-publish.yml index ec6155bf0f45..986a065271cc 100644 --- a/.github/workflows/_docker-build-and-publish.yml +++ b/.github/workflows/_docker-build-and-publish.yml @@ -91,6 +91,11 @@ on: required: false type: string default: "" + image_formats: + description: "CSV of image formats to build and push: oci (default), zstd, nydus" + required: false + type: string + default: "oci" jobs: build-and-publish: @@ -100,6 +105,7 @@ jobs: env: CUDA_KEY: ${{ inputs.cuda_key }} CUDA_SUFFIX: ${{ inputs.cuda_suffix }} + IMAGE_FORMATS: ${{ inputs.image_formats }} IMAGE_REPO: ${{ inputs.image_repo }} PUBLISH_DEFAULT_CUDA_ALIAS: ${{ inputs.publish_default_cuda_alias }} REGISTRY_HOST: ${{ inputs.registry_host }} @@ -187,6 +193,28 @@ jobs: echo "EOF" } >> "$GITHUB_OUTPUT" + - name: Validate requested image formats + run: | + set -euo pipefail + python3 - <<'PY' + import os + import sys + + raw = os.environ.get("IMAGE_FORMATS", "oci") + tokens = [token.strip() for token in raw.split(",") if token.strip()] + if not tokens: + tokens = ["oci"] + allowed = {"oci", "zstd", "nydus"} + unknown = [token for token in tokens if token not in allowed] + if unknown: + print( + f"::error::unknown image_formats tokens: {unknown}; " + f"allowed: {sorted(allowed)}" + ) + sys.exit(1) + print(f"Requested image formats: {tokens}") + PY + - name: Compute Docker build metadata args id: build-metadata env: @@ -378,3 +406,229 @@ jobs: python3 -m pytest -q test/registered/unit/test_runtime_context.py -k TestMoeFlagsGroup python3 -m pytest -q test/registered/unit/test_model_overrides.py -k deepseek_spec_moe_resolution ' + + # --- zstd format: re-export the just-built (cached) layers with zstd + # compression, then reuse the same provenance/runtime verification. This + # is a cache hit (no --no-cache), so the config is identical to base and + # only layer compression changes. + - name: Build and push AMD64 zstd image + id: build-zstd + if: ${{ contains(inputs.image_formats, 'zstd') }} + run: | + set -euo pipefail + VERSION_ARG="" + if [ -n "${SGL_VERSION}" ]; then + VERSION_ARG="--build-arg SGL_VERSION=${SGL_VERSION}" + fi + mapfile -t METADATA_ARGS < /tmp/docker-metadata.args + + docker buildx build \ + --target ${{ inputs.docker_target }} \ + --platform linux/amd64 \ + --output type=image,name=${IMAGE_REPO},push-by-digest=true,name-canonical=true,push=true,compression=zstd,compression-level=3,force-compression=true,oci-mediatypes=true \ + -f docker/Dockerfile \ + --build-arg CUDA_VERSION=${{ inputs.cuda_version }} \ + --build-arg BUILD_TYPE=all \ + --build-arg GRACE_BLACKWELL=0 \ + --build-arg INSTALL_FLASHINFER_JIT_CACHE=1 \ + "${METADATA_ARGS[@]}" \ + ${VERSION_ARG} \ + ${{ inputs.extra_build_args }} \ + --metadata-file /tmp/metadata-zstd.json \ + . + + DIGEST=$(python3 -c "import json; print(json.load(open('/tmp/metadata-zstd.json'))['containerimage.digest'])") + echo "Pushed zstd digest: ${DIGEST}" + echo "digest=${DIGEST}" >> "$GITHUB_OUTPUT" + + - name: Create Volcengine zstd image tag + if: ${{ contains(inputs.image_formats, 'zstd') }} + env: + IMAGE_TAGS: ${{ steps.image-tag.outputs.image-tags }} + run: | + set -euo pipefail + mapfile -t image_tags <<< "${IMAGE_TAGS}" + tag_args=() + for image_tag in "${image_tags[@]}"; do + [ -n "${image_tag}" ] || continue + tag_args+=(-t "${IMAGE_REPO}:${image_tag}-zstd") + done + + docker buildx imagetools create \ + "${tag_args[@]}" \ + "${IMAGE_REPO}@${{ steps.build-zstd.outputs.digest }}" + for image_tag in "${image_tags[@]}"; do + [ -n "${image_tag}" ] || continue + echo "Published ${IMAGE_REPO}:${image_tag}-zstd" + done + + - name: Verify pushed AMD64 zstd image and source provenance + if: ${{ contains(inputs.image_formats, 'zstd') }} + env: + EXPECTED_BUILD_COMMIT: ${{ steps.build-metadata.outputs.build-commit }} + EXPECTED_BUILD_SOURCE: ${{ steps.build-metadata.outputs.build-source }} + EXPECTED_BUILD_TREE: ${{ steps.build-metadata.outputs.build-tree }} + EXPECTED_PYTHON_MANIFEST_SHA256: ${{ steps.build-metadata.outputs.python-manifest-sha256 }} + IMAGE_DIGEST: ${{ steps.build-zstd.outputs.digest }} + run: | + set -euo pipefail + IMAGE_REF="${IMAGE_REPO}@${IMAGE_DIGEST}" + docker buildx imagetools inspect "${IMAGE_REF}" + + # zstd-compressed layers require Docker Engine >= 23 / containerd >= 1.5 + # to pull and run. If this daemon cannot, degrade to manifest-only + # inspection (media-type + provenance labels) instead of failing. + if docker pull "${IMAGE_REF}"; then + python3 - "${IMAGE_REF}" <<'PY' + import json + import os + import subprocess + import sys + + inspect = json.loads( + subprocess.check_output(["docker", "image", "inspect", sys.argv[1]]) + )[0] + if inspect["Architecture"] != "amd64": + raise SystemExit(f"unexpected architecture: {inspect['Architecture']}") + labels = inspect["Config"].get("Labels") or {} + expected = { + "ai.sglang.build.commit": os.environ["EXPECTED_BUILD_COMMIT"], + "ai.sglang.build.tree": os.environ["EXPECTED_BUILD_TREE"], + "ai.sglang.build.python-manifest-sha256": os.environ[ + "EXPECTED_PYTHON_MANIFEST_SHA256" + ], + "org.opencontainers.image.source": os.environ["EXPECTED_BUILD_SOURCE"], + } + mismatches = { + key: (labels.get(key), value) + for key, value in expected.items() + if labels.get(key) != value + } + if mismatches: + raise SystemExit(f"image provenance mismatch: {mismatches}") + print(json.dumps({"architecture": inspect["Architecture"], **expected})) + PY + + docker run --rm \ + -e EXPECTED_BUILD_COMMIT \ + -e EXPECTED_BUILD_TREE \ + -e EXPECTED_PYTHON_MANIFEST_SHA256 \ + --entrypoint bash "${IMAGE_REF}" -lc ' + set -euo pipefail + test "$SGLANG_BUILD_COMMIT" = "$EXPECTED_BUILD_COMMIT" + test "$SGLANG_BUILD_TREE" = "$EXPECTED_BUILD_TREE" + test "$SGLANG_PYTHON_MANIFEST_SHA256" = "$EXPECTED_PYTHON_MANIFEST_SHA256" + command -v oniond + python3 scripts/ci/verify_private_image_runtime.py + ' + else + echo "::warning::daemon cannot pull zstd image; degrading to manifest-only check" + python3 - "${IMAGE_REF}" <<'PY' + import json + import subprocess + import sys + + raw = subprocess.check_output( + ["docker", "buildx", "imagetools", "inspect", "--raw", sys.argv[1]] + ) + manifest = json.loads(raw) + layers = manifest.get("layers", []) + media_types = {layer.get("mediaType", "") for layer in layers} + if not any("zstd" in media_type for media_type in media_types): + raise SystemExit(f"expected zstd layer media types, got: {media_types}") + print(json.dumps({"zstd_layer_media_types": sorted(media_types)})) + PY + fi + + # --- nydus format: convert the already-pushed base OCI image into a + # nydus (lazy-loading) image with nydusify. A plain `docker run` cannot + # run nydus images (needs nydus-snapshotter), so CI validates with + # `nydusify check` + manifest/label inspection only. Real snapshotter + # mount + GPU serving is deferred to a ServingKit node. + - name: Install nydus tooling + if: ${{ contains(inputs.image_formats, 'nydus') }} + env: + NYDUS_VERSION: "2.3.0" + run: | + set -euo pipefail + tarball="nydus-static-v${NYDUS_VERSION}-linux-amd64.tgz" + url="https://github.com/dragonflyoss/nydus/releases/download/v${NYDUS_VERSION}/${tarball}" + echo "Downloading ${url}" + # The CI proxy is intermittently unstable, so wrap the download in an + # outer retry loop and let curl retry every error class (including + # curl 56 "Connection died") instead of only its default retryable set. + retry() { for i in 1 2 3 4 5; do "$@" && return 0; echo "attempt $i failed; sleeping 15s"; sleep 15; done; return 1; } + retry curl -fL --retry 5 --retry-all-errors --retry-delay 2 -o "/tmp/${tarball}" "${url}" + tar -xzf "/tmp/${tarball}" -C /tmp + sudo install -m 0755 /tmp/nydus-static/nydusify /usr/local/bin/nydusify + sudo install -m 0755 /tmp/nydus-static/nydus-image /usr/local/bin/nydus-image + nydusify --version + nydus-image --version + + - name: Convert and push nydus image + if: ${{ contains(inputs.image_formats, 'nydus') }} + env: + IMAGE_TAGS: ${{ steps.image-tag.outputs.image-tags }} + BASE_DIGEST: ${{ steps.build.outputs.digest }} + run: | + set -euo pipefail + mapfile -t image_tags <<< "${IMAGE_TAGS}" + # nydusify convert is pull(base OCI)->convert->push(nydus). The pull and + # convert are deterministic and the push overwrites the same tag, so the + # whole convert is idempotent and safe to re-run. The internal registry + # push occasionally drops the connection mid-copy ("io: read/write on + # closed pipe"), so wrap the full convert in an outer retry loop. + retry() { for i in 1 2 3 4 5; do "$@" && return 0; echo "attempt $i failed; sleeping 20s"; sleep 20; done; return 1; } + for image_tag in "${image_tags[@]}"; do + [ -n "${image_tag}" ] || continue + target="${IMAGE_REPO}:${image_tag}-nydus" + echo "Converting ${IMAGE_REPO}@${BASE_DIGEST} -> ${target}" + retry nydusify convert \ + --nydus-image /usr/local/bin/nydus-image \ + --source "${IMAGE_REPO}@${BASE_DIGEST}" \ + --target "${target}" + echo "Published ${target}" + done + + - name: Verify pushed nydus image + if: ${{ contains(inputs.image_formats, 'nydus') }} + env: + IMAGE_TAGS: ${{ steps.image-tag.outputs.image-tags }} + EXPECTED_BUILD_COMMIT: ${{ steps.build-metadata.outputs.build-commit }} + EXPECTED_BUILD_SOURCE: ${{ steps.build-metadata.outputs.build-source }} + run: | + set -euo pipefail + mapfile -t image_tags <<< "${IMAGE_TAGS}" + for image_tag in "${image_tags[@]}"; do + [ -n "${image_tag}" ] || continue + target="${IMAGE_REPO}:${image_tag}-nydus" + echo "Checking ${target}" + nydusify check --target "${target}" + docker buildx imagetools inspect "${target}" + python3 - "${target}" <<'PY' + import json + import os + import subprocess + import sys + + target = sys.argv[1] + raw = subprocess.check_output( + ["docker", "buildx", "imagetools", "inspect", "--raw", target] + ) + manifest = json.loads(raw) + layers = manifest.get("layers", []) + media_types = [layer.get("mediaType", "") for layer in layers] + annotations = {} + for layer in layers: + annotations.update(layer.get("annotations", {}) or {}) + has_nydus = any( + "nydus" in media_type for media_type in media_types + ) or any("nydus" in key for key in annotations) + if not has_nydus: + raise SystemExit( + f"expected nydus markers in manifest; layers={media_types} " + f"annotations={sorted(annotations)}" + ) + print(json.dumps({"nydus_layer_media_types": media_types})) + PY + done diff --git a/.github/workflows/release-docker-dev.yml b/.github/workflows/release-docker-dev.yml index 7c1d8a388c9f..65ca8aced1ca 100644 --- a/.github/workflows/release-docker-dev.yml +++ b/.github/workflows/release-docker-dev.yml @@ -46,6 +46,10 @@ on: description: "Volcengine fork only: verify an existing private debug image by immutable digest (repo@sha256:...), without rebuilding it." required: false default: "" + private_debug_image_formats: + description: "Volcengine fork only: CSV of image formats for the debug base build (oci, zstd, nydus). Applies only when private_debug_base_only is set." + required: false + default: "oci,zstd,nydus" schedule: - cron: "0 0 * * *" @@ -355,6 +359,7 @@ jobs: cuda_version: 13.0.1 cuda_suffix: cu130 publish_default_cuda_alias: false + image_formats: ${{ inputs.private_debug_image_formats || 'oci,zstd,nydus' }} image_tag_override: ${{ format('debug-{0}-{1}-{2}', inputs.tag || 'servingkit', github.run_id, github.run_attempt) }} use_environment: prod extra_build_args: >- diff --git a/docker/Dockerfile b/docker/Dockerfile index 89e2d319d2ab..ce9c18205281 100644 --- a/docker/Dockerfile +++ b/docker/Dockerfile @@ -52,7 +52,8 @@ RUN --mount=type=cache,target=/var/cache/apt,id=base-apt \ && apt install -y --no-install-recommends python3.12-full python3.12-dev \ && update-alternatives --install /usr/bin/python3 python3 /usr/bin/python3.12 2 \ && update-alternatives --set python3 /usr/bin/python3.12 \ - && curl --fail --show-error --location --retry 5 --retry-all-errors --retry-delay 2 \ + && retry() { for i in 1 2 3 4 5; do "$@" && return 0; echo "attempt $i failed; sleeping 15s"; sleep 15; done; return 1; } \ + && retry curl --fail --show-error --location --retry 5 --retry-all-errors --retry-delay 2 \ -o get-pip.py https://bootstrap.pypa.io/get-pip.py \ && python3 get-pip.py --break-system-packages \ && rm get-pip.py \ @@ -140,8 +141,12 @@ RUN if [ -n "${PIP_DEFAULT_INDEX}" ]; then \ fi # GDRCopy installation +# Wrap the tarball download in a whole-command retry loop: the internal egress +# proxy is intermittently unstable and can stay down longer than curl's own +# --retry window, so an outer loop with sleep survives multi-minute outages. RUN mkdir -p /tmp/gdrcopy && cd /tmp \ - && curl --retry 3 --retry-delay 2 -fsSL -o v${GDRCOPY_VERSION}.tar.gz \ + && retry() { for i in 1 2 3 4 5; do "$@" && return 0; echo "attempt $i failed; sleeping 15s"; sleep 15; done; return 1; } \ + && retry curl --retry 5 --retry-all-errors --retry-delay 2 -fsSL -o v${GDRCOPY_VERSION}.tar.gz \ https://${GITHUB_ARTIFACTORY}/NVIDIA/gdrcopy/archive/refs/tags/v${GDRCOPY_VERSION}.tar.gz \ && tar -xzf v${GDRCOPY_VERSION}.tar.gz && rm v${GDRCOPY_VERSION}.tar.gz \ && cd gdrcopy-${GDRCOPY_VERSION}/packages \ @@ -193,7 +198,8 @@ WORKDIR /sgl-workspace # Rust toolchain for setuptools-rust extensions (e.g. sglang-grpc). # Requires >= 1.85 (edition 2024). Inherited by framework via FROM torch_deps. ENV PATH="/root/.cargo/bin:${PATH}" -RUN curl --proto '=https' --tlsv1.2 --fail --show-error --location \ +RUN retry() { for i in 1 2 3 4 5; do "$@" && return 0; echo "attempt $i failed; sleeping 15s"; sleep 15; done; return 1; } \ + && retry curl --proto '=https' --tlsv1.2 --fail --show-error --location \ --retry 5 --retry-all-errors --retry-delay 2 \ -o /tmp/rustup-init.sh https://sh.rustup.rs \ && sh /tmp/rustup-init.sh -y --no-modify-path --profile minimal \ @@ -202,7 +208,8 @@ RUN curl --proto '=https' --tlsv1.2 --fail --show-error --location \ # Install sgl-kernel (from pre-built wheel) RUN --mount=type=cache,target=/root/.cache/pip \ - python3 -m pip install --upgrade pip setuptools wheel html5lib six \ + retry() { for i in 1 2 3 4 5; do "$@" && return 0; echo "attempt $i failed; sleeping 15s"; sleep 15; done; return 1; } \ + && retry python3 -m pip install --upgrade pip setuptools wheel html5lib six \ && case "$CUDA_VERSION" in \ 12.6.1) CUINDEX=126 ;; \ 12.8.1) CUINDEX=128 ;; \ @@ -211,14 +218,14 @@ RUN --mount=type=cache,target=/root/.cache/pip \ *) echo "Unsupported CUDA version: $CUDA_VERSION" && exit 1 ;; \ esac \ && if [ "$CUDA_VERSION" = "12.6.1" ]; then \ - python3 -m pip install https://${GITHUB_ARTIFACTORY}/sgl-project/whl/releases/download/v${SGL_KERNEL_VERSION}/sglang_kernel-${SGL_KERNEL_VERSION}+cu124-cp310-abi3-manylinux2014_$(uname -m).whl --force-reinstall --no-deps \ + retry python3 -m pip install https://${GITHUB_ARTIFACTORY}/sgl-project/whl/releases/download/v${SGL_KERNEL_VERSION}/sglang_kernel-${SGL_KERNEL_VERSION}+cu124-cp310-abi3-manylinux2014_$(uname -m).whl --force-reinstall --no-deps \ ; \ elif [ "$CUDA_VERSION" = "12.8.1" ] || [ "$CUDA_VERSION" = "12.9.1" ]; then \ - python3 -m pip install https://github.com/sgl-project/whl/releases/download/v${SGL_KERNEL_VERSION}/sglang_kernel-${SGL_KERNEL_VERSION}+cu129-cp310-abi3-manylinux2014_$(uname -m).whl --force-reinstall --no-deps \ + retry python3 -m pip install https://github.com/sgl-project/whl/releases/download/v${SGL_KERNEL_VERSION}/sglang_kernel-${SGL_KERNEL_VERSION}+cu129-cp310-abi3-manylinux2014_$(uname -m).whl --force-reinstall --no-deps \ ; \ elif [ "$CUDA_VERSION" = "13.0.1" ]; then \ # --no-deps prevents pip from pulling torch from default PyPI - python3 -m pip install sglang-kernel==${SGL_KERNEL_VERSION} --force-reinstall --no-deps \ + retry python3 -m pip install sglang-kernel==${SGL_KERNEL_VERSION} --force-reinstall --no-deps \ ; \ else \ echo "Unsupported CUDA version: $CUDA_VERSION" && exit 1 \ @@ -238,7 +245,8 @@ COPY proto /tmp/sglang_deps/proto # Generate constraints.txt to prevent reinstalling these deps in later stages RUN --mount=type=cache,target=/root/.cache/pip \ --mount=type=cache,target=/root/.cargo/registry \ - case "$CUDA_VERSION" in \ + retry() { for i in 1 2 3 4 5; do "$@" && return 0; echo "attempt $i failed; sleeping 15s"; sleep 15; done; return 1; } \ + && case "$CUDA_VERSION" in \ 12.6.1) CUINDEX=126 ;; \ 12.8.1) CUINDEX=128 ;; \ 12.9.1) CUINDEX=129 ;; \ @@ -256,13 +264,13 @@ RUN --mount=type=cache,target=/root/.cache/pip \ sed -i 's/flashinfer_python\[cu13\]/flashinfer_python[cu12]/' pyproject.toml && \ sed -i 's/nvidia-cutlass-dsl\[cu13\]/nvidia-cutlass-dsl/' pyproject.toml; \ fi \ - && python3 -m pip install --extra-index-url https://download.pytorch.org/whl/cu${CUINDEX} ".[${BUILD_TYPE}]" \ + && retry python3 -m pip install --extra-index-url https://download.pytorch.org/whl/cu${CUINDEX} ".[${BUILD_TYPE}]" \ && if [ "${CUDA_VERSION%%.*}" = "12" ]; then \ pip list --format=freeze | awk -F'==' '/-cu13(==|$)/ {print $1}' \ | xargs -r python3 -m pip uninstall -y && \ - python3 -m pip install --index-url https://download.pytorch.org/whl/cu${CUINDEX} \ + retry python3 -m pip install --index-url https://download.pytorch.org/whl/cu${CUINDEX} \ torch==2.11.0 torchvision==0.26.0 torchaudio==2.11.0 --force-reinstall; \ - python3 -m pip install https://github.com/sgl-project/whl/releases/download/v${SGL_DEEP_GEMM_VERSION}/sgl_deep_gemm-${SGL_DEEP_GEMM_VERSION}+cu129-py3-none-manylinux2014_$(uname -m).whl --force-reinstall; \ + retry python3 -m pip install https://github.com/sgl-project/whl/releases/download/v${SGL_DEEP_GEMM_VERSION}/sgl_deep_gemm-${SGL_DEEP_GEMM_VERSION}+cu129-py3-none-manylinux2014_$(uname -m).whl --force-reinstall; \ fi \ && cd /sgl-workspace \ && rm -rf /tmp/sglang_deps \ @@ -313,7 +321,8 @@ RUN set -eux; \ sed -i 's/#define NUM_TIMEOUT_CYCLES 200000000000ull/#define NUM_TIMEOUT_CYCLES 2000000000000ull/' csrc/kernels/configs.cuh && \ cd .. ; \ else \ - curl --retry 3 --retry-delay 2 -fsSL -o ${DEEPEP_COMMIT}.zip \ + retry() { for i in 1 2 3 4 5; do "$@" && return 0; echo "attempt $i failed; sleeping 15s"; sleep 15; done; return 1; }; \ + retry curl --retry 5 --retry-all-errors --retry-delay 2 -fsSL -o ${DEEPEP_COMMIT}.zip \ https://${GITHUB_ARTIFACTORY}/deepseek-ai/DeepEP/archive/${DEEPEP_COMMIT}.zip && \ unzip -q ${DEEPEP_COMMIT}.zip && rm ${DEEPEP_COMMIT}.zip && mv DeepEP-${DEEPEP_COMMIT} DeepEP && cd DeepEP && \ sed -i 's/#define NUM_CPU_TIMEOUT_SECS 100/#define NUM_CPU_TIMEOUT_SECS 1000/' csrc/kernels/configs.cuh && \ @@ -360,9 +369,11 @@ WORKDIR /build # setup.py derives the version from `git rev-parse`, so keep the .git dir # (a source zip archive would not build). RUN --mount=type=cache,target=/root/.cache/pip \ - mkdir -p /wheels && \ + retry() { for i in 1 2 3 4 5; do "$@" && return 0; echo "attempt $i failed; sleeping 15s"; sleep 15; done; return 1; } \ + && clone_fresh() { dir="$1"; shift; rm -rf "$dir" && git clone "$@" "$dir" ; } \ + && mkdir -p /wheels && \ if [ "$(uname -m)" = "x86_64" ]; then \ - git clone https://github.com/Tencent/hpc-ops.git && \ + retry clone_fresh hpc-ops https://github.com/Tencent/hpc-ops.git && \ cd hpc-ops && \ git checkout ${HPC_OPS_COMMIT} && \ python3 setup.py bdist_wheel -d /wheels; \ @@ -387,12 +398,16 @@ RUN --mount=type=cache,target=/root/.cache/pip \ *) echo "Unsupported CUDA version: $CUDA_VERSION" && exit 1 ;; \ esac \ && mkdir -p /flashinfer_jit_output \ + # Retry external-index pip installs: the internal egress proxy is briefly + # unstable (proxy-connection-failed / 504) but recovers within a minute, so + # a whole-command retry with sleep is more robust than pip's inner --retries. + && retry() { for i in 1 2 3 4 5; do "$@" && return 0; echo "attempt $i failed; sleeping 15s"; sleep 15; done; return 1; } \ # flashinfer-cubin is CUDA-version-agnostic, unlike jit-cache, so its index-url has no cu${CUINDEX} suffix - && python3 -m pip install flashinfer-cubin==${FLASHINFER_VERSION} --index-url https://flashinfer.ai/whl \ + && retry python3 -m pip install flashinfer-cubin==${FLASHINFER_VERSION} --index-url https://flashinfer.ai/whl \ && cp -r /usr/local/lib/python3.12/dist-packages/flashinfer_cubin /flashinfer_jit_output/ \ && cp -r /usr/local/lib/python3.12/dist-packages/flashinfer_cubin-*.dist-info /flashinfer_jit_output/ \ && if [ "$INSTALL_FLASHINFER_JIT_CACHE" = "1" ]; then \ - python3 -m pip install flashinfer-jit-cache==${FLASHINFER_VERSION} --index-url https://flashinfer.ai/whl/cu${CUINDEX} \ + retry python3 -m pip install flashinfer-jit-cache==${FLASHINFER_VERSION} --index-url https://flashinfer.ai/whl/cu${CUINDEX} \ && cp -r /usr/local/lib/python3.12/dist-packages/flashinfer_jit_cache /flashinfer_jit_output/ \ && cp -r /usr/local/lib/python3.12/dist-packages/flashinfer_jit_cache-*.dist-info /flashinfer_jit_output/ ; \ fi @@ -413,15 +428,20 @@ RUN --mount=type=cache,target=/var/cache/apt,id=devtools-apt \ && rm -rf /var/lib/apt/lists/* # Download CLI tools (each in its own layer for parallel downloads) -RUN curl --retry 3 --retry-delay 2 -LSso /tools/diff-so-fancy \ +# Each download is wrapped in an outer retry loop because the internal egress +# proxy can stay down longer than curl's own --retry window. +RUN retry() { for i in 1 2 3 4 5; do "$@" && return 0; echo "attempt $i failed; sleeping 15s"; sleep 15; done; return 1; } \ + && retry curl --retry 5 --retry-all-errors --retry-delay 2 -LSso /tools/diff-so-fancy \ https://${GITHUB_ARTIFACTORY}/so-fancy/diff-so-fancy/releases/download/v1.4.4/diff-so-fancy \ && chmod +x /tools/diff-so-fancy -RUN curl --retry 3 --retry-delay 2 -LSso /tools/clang-format \ +RUN retry() { for i in 1 2 3 4 5; do "$@" && return 0; echo "attempt $i failed; sleeping 15s"; sleep 15; done; return 1; } \ + && retry curl --retry 5 --retry-all-errors --retry-delay 2 -LSso /tools/clang-format \ https://${GITHUB_ARTIFACTORY}/muttleyxd/clang-tools-static-binaries/releases/download/master-32d3ac78/clang-format-16_linux-amd64 \ && chmod +x /tools/clang-format -RUN curl --retry 3 --retry-delay 2 -fsSL -o /tmp/clangd.zip \ +RUN retry() { for i in 1 2 3 4 5; do "$@" && return 0; echo "attempt $i failed; sleeping 15s"; sleep 15; done; return 1; } \ + && retry curl --retry 5 --retry-all-errors --retry-delay 2 -fsSL -o /tmp/clangd.zip \ https://${GITHUB_ARTIFACTORY}/clangd/clangd/releases/download/18.1.3/clangd-linux-18.1.3.zip \ && unzip -q /tmp/clangd.zip -d /tmp \ && cp /tmp/clangd_18.1.3/bin/* /tools/ \ @@ -431,28 +451,39 @@ RUN curl --retry 3 --retry-delay 2 -fsSL -o /tmp/clangd.zip \ RUN CMAKE_VERSION=3.31.1 \ && ARCH=$(uname -m) \ && CMAKE_INSTALLER="cmake-${CMAKE_VERSION}-linux-${ARCH}" \ - && curl --retry 3 --retry-delay 2 -fsSL -o "/tmp/${CMAKE_INSTALLER}.tar.gz" \ + && retry() { for i in 1 2 3 4 5; do "$@" && return 0; echo "attempt $i failed; sleeping 15s"; sleep 15; done; return 1; } \ + && retry curl --retry 5 --retry-all-errors --retry-delay 2 -fsSL -o "/tmp/${CMAKE_INSTALLER}.tar.gz" \ "https://${GITHUB_ARTIFACTORY}/Kitware/CMake/releases/download/v${CMAKE_VERSION}/${CMAKE_INSTALLER}.tar.gz" \ && tar -xzf "/tmp/${CMAKE_INSTALLER}.tar.gz" -C /tmp \ && cp -r "/tmp/${CMAKE_INSTALLER}/bin/"* /tools/ \ && mkdir -p /tools/share && cp -r "/tmp/${CMAKE_INSTALLER}/share/"* /tools/share/ \ && rm -rf "/tmp/${CMAKE_INSTALLER}" "/tmp/${CMAKE_INSTALLER}.tar.gz" -RUN curl --proto '=https' --tlsv1.2 --fail --show-error --location \ - --retry 5 --retry-all-errors --retry-delay 2 \ - -o /tmp/just-install.sh https://just.systems/install.sh \ - && sed -i "s|https://github.com|https://${GITHUB_ARTIFACTORY}|g" /tmp/just-install.sh \ - && bash /tmp/just-install.sh --tag 1.42.4 --to /tools \ - && rm /tmp/just-install.sh \ - && test -x /tools/just +RUN retry() { for i in 1 2 3 4 5 6 7 8; do "$@" && return 0; echo "attempt $i failed; sleeping 20s"; sleep 20; done; return 1; } \ + && install_just() { \ + curl --proto '=https' --tlsv1.2 --fail --show-error --location \ + --retry 5 --retry-all-errors --retry-delay 2 \ + -o /tmp/just-install.sh https://just.systems/install.sh \ + && sed -i "s|https://github.com|https://${GITHUB_ARTIFACTORY}|g" /tmp/just-install.sh \ + && bash /tmp/just-install.sh --tag 1.42.4 --to /tools \ + && rm -f /tmp/just-install.sh \ + && test -x /tools/just ; \ + } \ + && retry install_just # Install oh-my-zsh and plugins -RUN curl --fail --show-error --location --retry 5 --retry-all-errors --retry-delay 2 \ - -o /tmp/oh-my-zsh-install.sh https://raw.githubusercontent.com/ohmyzsh/ohmyzsh/master/tools/install.sh \ - && sh /tmp/oh-my-zsh-install.sh "" --unattended \ - && rm /tmp/oh-my-zsh-install.sh \ - && git clone --depth 1 https://github.com/zsh-users/zsh-autosuggestions ${ZSH_CUSTOM:-/root/.oh-my-zsh/custom}/plugins/zsh-autosuggestions \ - && git clone --depth 1 https://github.com/zsh-users/zsh-syntax-highlighting.git ${ZSH_CUSTOM:-/root/.oh-my-zsh/custom}/plugins/zsh-syntax-highlighting +RUN retry() { for i in 1 2 3 4 5 6 7 8; do "$@" && return 0; echo "attempt $i failed; sleeping 20s"; sleep 20; done; return 1; } \ + && install_omz() { \ + rm -rf /root/.oh-my-zsh \ + && curl --fail --show-error --location --retry 5 --retry-all-errors --retry-delay 2 \ + -o /tmp/oh-my-zsh-install.sh https://raw.githubusercontent.com/ohmyzsh/ohmyzsh/master/tools/install.sh \ + && sh /tmp/oh-my-zsh-install.sh "" --unattended \ + && rm -f /tmp/oh-my-zsh-install.sh ; \ + } \ + && clone_fresh() { dir="$1"; shift; rm -rf "$dir" && git clone --depth 1 "$@" "$dir" ; } \ + && retry install_omz \ + && retry clone_fresh ${ZSH_CUSTOM:-/root/.oh-my-zsh/custom}/plugins/zsh-autosuggestions https://github.com/zsh-users/zsh-autosuggestions \ + && retry clone_fresh ${ZSH_CUSTOM:-/root/.oh-my-zsh/custom}/plugins/zsh-syntax-highlighting https://github.com/zsh-users/zsh-syntax-highlighting.git ######################################################## # PARALLEL STAGE 5: Gateway Builder (starts from base) @@ -473,7 +504,8 @@ COPY sgl-model-gateway /build/sgl-model-gateway # Install Rust, build gateway binary and Python bindings, then clean up Rust toolchain RUN --mount=type=cache,target=/root/.cache/pip \ - curl --proto '=https' --tlsv1.2 --fail --show-error --location \ + retry() { for i in 1 2 3 4 5; do "$@" && return 0; echo "attempt $i failed; sleeping 15s"; sleep 15; done; return 1; } \ + && retry curl --proto '=https' --tlsv1.2 --fail --show-error --location \ --retry 5 --retry-all-errors --retry-delay 2 \ -o /tmp/rustup-init.sh https://sh.rustup.rs \ && sh /tmp/rustup-init.sh -y \ @@ -634,16 +666,19 @@ RUN --mount=type=cache,target=/var/cache/apt,id=framework-apt \ # Install Mooncake RUN --mount=type=cache,target=/root/.cache/pip \ - CUDA_MAJOR="${CUDA_VERSION%%.*}" && \ + retry() { for i in 1 2 3 4 5; do "$@" && return 0; echo "attempt $i failed; sleeping 15s"; sleep 15; done; return 1; } \ + && CUDA_MAJOR="${CUDA_VERSION%%.*}" && \ if [ "$CUDA_MAJOR" -ge 13 ]; then \ - python3 -m pip install mooncake-transfer-engine-cuda13==${MOONCAKE_VERSION}; \ + retry python3 -m pip install mooncake-transfer-engine-cuda13==${MOONCAKE_VERSION}; \ else \ - python3 -m pip install mooncake-transfer-engine==${MOONCAKE_VERSION}; \ + retry python3 -m pip install mooncake-transfer-engine==${MOONCAKE_VERSION}; \ fi # Install MSCCL++ Python dependencies and package (builds extension via CMake through pip) RUN --mount=type=cache,target=/root/.cache/pip \ - git clone --depth=1 --branch ${MSCCLPP_VERSION} https://${GITHUB_ARTIFACTORY}/microsoft/mscclpp.git /tmp/mscclpp \ + retry() { for i in 1 2 3 4 5; do "$@" && return 0; echo "attempt $i failed; sleeping 15s"; sleep 15; done; return 1; } \ + && clone_fresh() { dir="$1"; shift; rm -rf "$dir" && git clone "$@" "$dir" ; } \ + && retry clone_fresh /tmp/mscclpp --depth=1 --branch ${MSCCLPP_VERSION} https://${GITHUB_ARTIFACTORY}/microsoft/mscclpp.git \ && case "${CUDA_VERSION}" in \ 12.*) \ CMAKE_ARGS="-DMSCCLPP_BYPASS_GPU_CHECK=ON -DMSCCLPP_USE_CUDA=ON -DMSCCLPP_GPU_ARCHS=80,90,100,100a,103,103a" \ @@ -661,7 +696,8 @@ RUN --mount=type=cache,target=/root/.cache/pip \ # Install essential Python packages (use constraints to prevent conflicts) RUN --mount=type=cache,target=/root/.cache/pip \ - python3 -m pip install -c /sgl-workspace/constraints.txt \ + retry() { for i in 1 2 3 4 5; do "$@" && return 0; echo "attempt $i failed; sleeping 15s"; sleep 15; done; return 1; } \ + && retry python3 -m pip install -c /sgl-workspace/constraints.txt \ datamodel_code_generator \ pre-commit \ pytest \ @@ -681,7 +717,8 @@ RUN --mount=type=cache,target=/root/.cache/pip \ "runai-model-streamer[s3,gcs,azure]>=0.15.7" # Install the EIC SDK used by the EIC HiCache backend. -RUN python3 -m pip install --force-reinstall \ +RUN retry() { for i in 1 2 3 4 5; do "$@" && return 0; echo "attempt $i failed; sleeping 15s"; sleep 15; done; return 1; } \ + && retry python3 -m pip install --force-reinstall \ https://eic-sdk-release.tos-cn-beijing.volces.com/python/eic-1.5.2-py3-none-any.whl \ --no-cache-dir @@ -689,12 +726,14 @@ RUN python3 -m pip install --force-reinstall \ # the `nixl` import path) but unconditionally requires nixl-cu12, so we install # it with --no-deps and pair it with the matching nixl-cu12 / nixl-cu13 binary # to avoid shipping wrong-CUDA libs on cu13 images. -RUN --mount=type=cache,target=/root/.cache/pip if [ "${CUDA_VERSION%%.*}" = "12" ]; then \ - python3 -m pip install nixl nixl-cu12 --no-deps ; \ - python3 -m pip install cuda-python==12.9 ; \ +RUN --mount=type=cache,target=/root/.cache/pip \ + retry() { for i in 1 2 3 4 5; do "$@" && return 0; echo "attempt $i failed; sleeping 15s"; sleep 15; done; return 1; } \ + && if [ "${CUDA_VERSION%%.*}" = "12" ]; then \ + retry python3 -m pip install nixl nixl-cu12 --no-deps ; \ + retry python3 -m pip install cuda-python==12.9 ; \ elif [ "${CUDA_VERSION%%.*}" = "13" ]; then \ - python3 -m pip install nixl nixl-cu13 --no-deps ; \ - python3 -m pip install cuda-python==13.2.0 ; \ + retry python3 -m pip install nixl nixl-cu13 --no-deps ; \ + retry python3 -m pip install cuda-python==13.2.0 ; \ fi # Add yank script @@ -716,7 +755,8 @@ COPY docker/configs/.zshrc /root/.zshrc # libsqlite3-0: CVE-2025-{6965,7709} # libtasn1-6: CVE-2025-13151 # dpkg: CVE-2025-6297 -RUN python3 -m pip install --upgrade "urllib3>=2.6.3" "pillow>=12.1.1" +RUN retry() { for i in 1 2 3 4 5; do "$@" && return 0; echo "attempt $i failed; sleeping 15s"; sleep 15; done; return 1; } \ + && retry python3 -m pip install --upgrade "urllib3>=2.6.3" "pillow>=12.1.1" RUN --mount=type=cache,target=/var/cache/apt,id=framework-apt \ apt-get update && apt-get install -y --only-upgrade \ binutils binutils-common binutils-x86-64-linux-gnu libbinutils \ diff --git a/scripts/ci/get_volcengine_image_tag.py b/scripts/ci/get_volcengine_image_tag.py index b1537a340ce9..3cef2a6c2b58 100755 --- a/scripts/ci/get_volcengine_image_tag.py +++ b/scripts/ci/get_volcengine_image_tag.py @@ -9,6 +9,49 @@ from pathlib import Path from zoneinfo import ZoneInfo +TAG_SUFFIX_RE = re.compile(r"[0-9A-Za-z][0-9A-Za-z_.-]*") + + +def validate_suffix(flag: str, value: str) -> None: + if value and not TAG_SUFFIX_RE.fullmatch(value): + raise SystemExit(f"--{flag} must be a Docker tag-safe suffix") + + +def build_tag( + mode: str, + version: str, + timestamp: str, + tag_value: str = "", + variant_suffix: str = "", + cuda_suffix: str = "", + format_suffix: str = "", +) -> str: + """Compose the final image tag. + + Suffix order is fixed as ``variant`` -> ``cuda`` -> ``format`` so that the + image format marker (e.g. ``zstd`` / ``nydus``) always trails the CUDA + marker: ``v.byted..[-][-cu130][-zstd]``. + """ + if mode == "manual": + tag = f"v{version}.iaas.dev.{timestamp}" + elif mode == "nightly": + tag = f"v{version}.iaas.nightly.{timestamp}" + else: + if not tag_value: + raise SystemExit("--tag-value is required when --mode=version") + tag = f"v{version}.byted.{tag_value}.{timestamp}" + + if variant_suffix: + tag = f"{tag}-{variant_suffix}" + + if cuda_suffix: + tag = f"{tag}-{cuda_suffix}" + + if format_suffix: + tag = f"{tag}-{format_suffix}" + + return tag + def get_sglang_version() -> str: repo_root = Path(__file__).resolve().parents[2] @@ -51,30 +94,29 @@ def main() -> None: default="", help="Optional build variant suffix appended before the CUDA suffix.", ) + parser.add_argument( + "--format-suffix", + default="", + help="Optional image format suffix (e.g. zstd, nydus) appended after " + "the CUDA suffix.", + ) args = parser.parse_args() - if args.variant_suffix and not re.fullmatch( - r"[0-9A-Za-z][0-9A-Za-z_.-]*", args.variant_suffix - ): - raise SystemExit("--variant-suffix must be a Docker tag-safe suffix") + validate_suffix("variant-suffix", args.variant_suffix) + validate_suffix("format-suffix", args.format_suffix) version = get_sglang_version() timestamp = datetime.now(ZoneInfo("Asia/Shanghai")).strftime("%Y%m%d%H%M") - if args.mode == "manual": - tag = f"v{version}.iaas.dev.{timestamp}" - elif args.mode == "nightly": - tag = f"v{version}.iaas.nightly.{timestamp}" - else: - if not args.tag_value: - raise SystemExit("--tag-value is required when --mode=version") - tag = f"v{version}.byted.{args.tag_value}.{timestamp}" - - if args.variant_suffix: - tag = f"{tag}-{args.variant_suffix}" - - if args.cuda_suffix: - tag = f"{tag}-{args.cuda_suffix}" + tag = build_tag( + mode=args.mode, + version=version, + timestamp=timestamp, + tag_value=args.tag_value, + variant_suffix=args.variant_suffix, + cuda_suffix=args.cuda_suffix, + format_suffix=args.format_suffix, + ) print(tag) diff --git a/scripts/ci/test_get_volcengine_image_tag.py b/scripts/ci/test_get_volcengine_image_tag.py new file mode 100644 index 000000000000..6d8548fc4ec0 --- /dev/null +++ b/scripts/ci/test_get_volcengine_image_tag.py @@ -0,0 +1,101 @@ +#!/usr/bin/env python3 + +from __future__ import annotations + +import sys +import unittest +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parent)) + +from get_volcengine_image_tag import build_tag, validate_suffix + + +class GetVolcengineImageTagTest(unittest.TestCase): + def test_version_tag_without_suffixes(self) -> None: + self.assertEqual( + build_tag( + mode="version", + version="0.5.17", + timestamp="202608121200", + tag_value="glm5-2", + ), + "v0.5.17.byted.glm5-2.202608121200", + ) + + def test_manual_and_nightly_prefixes(self) -> None: + self.assertEqual( + build_tag(mode="manual", version="0.5.17", timestamp="202608121200"), + "v0.5.17.iaas.dev.202608121200", + ) + self.assertEqual( + build_tag(mode="nightly", version="0.5.17", timestamp="202608121200"), + "v0.5.17.iaas.nightly.202608121200", + ) + + def test_format_suffix_trails_cuda_suffix(self) -> None: + self.assertEqual( + build_tag( + mode="version", + version="0.5.17", + timestamp="202608121200", + tag_value="glm5-2", + cuda_suffix="cu130", + format_suffix="zstd", + ), + "v0.5.17.byted.glm5-2.202608121200-cu130-zstd", + ) + self.assertEqual( + build_tag( + mode="version", + version="0.5.17", + timestamp="202608121200", + tag_value="glm5-2", + cuda_suffix="cu130", + format_suffix="nydus", + ), + "v0.5.17.byted.glm5-2.202608121200-cu130-nydus", + ) + + def test_full_suffix_order_variant_cuda_format(self) -> None: + self.assertEqual( + build_tag( + mode="version", + version="0.5.17", + timestamp="202608121200", + tag_value="glm5-2", + variant_suffix="w4a8", + cuda_suffix="cu130", + format_suffix="zstd", + ), + "v0.5.17.byted.glm5-2.202608121200-w4a8-cu130-zstd", + ) + + def test_format_suffix_without_cuda_suffix(self) -> None: + self.assertEqual( + build_tag( + mode="version", + version="0.5.17", + timestamp="202608121200", + tag_value="glm5-2", + format_suffix="zstd", + ), + "v0.5.17.byted.glm5-2.202608121200-zstd", + ) + + def test_version_mode_requires_tag_value(self) -> None: + with self.assertRaisesRegex(SystemExit, "--tag-value is required"): + build_tag(mode="version", version="0.5.17", timestamp="202608121200") + + def test_validate_suffix_accepts_safe_values(self) -> None: + for value in ("zstd", "nydus", "cu130", "w4a8", "deepseek-v4", ""): + validate_suffix("format-suffix", value) + + def test_validate_suffix_rejects_unsafe_values(self) -> None: + for value in ("-zstd", "zstd/", "zs td", "-", "z$td"): + with self.assertRaisesRegex(SystemExit, "must be a Docker tag-safe suffix"): + validate_suffix("format-suffix", value) + + +if __name__ == "__main__": + unittest.main() From 8dba973152df95bf7ecbf5a94634f4fed6fe7ddb Mon Sep 17 00:00:00 2001 From: Hank Han Date: Thu, 13 Aug 2026 20:08:18 +0800 Subject: [PATCH 17/47] ci: default private delivery build to oci+zstd+nydus for all callers Flip the reusable _docker-build-and-publish.yml image_formats default from "oci" to "oci,zstd,nydus" so every private caller (dev regular matrix, release-docker, deepseek-v4 nightly, and the runtime target) produces and pushes all three formats by default instead of only the debug base path. Callers can still pass a subset to narrow the set. The zstd/nydus steps already iterate every resolved tag, derive all parameters from inputs (docker_target/cuda/extra_build_args), and reuse the same provenance verification, so generalizing to the production multi-tag / multi-variant callers is safe. The runtime target is fully verifiable too: it ships oniond, the full /sgl-workspace source tree (test/ + scripts/ci/verify_private_image_runtime.py), pytest (installed in the framework stage), and identical provenance env/labels, so the existing target-agnostic verify step applies without degradation. Also add a private_debug_docker_target input to release-docker-dev.yml (default framework_final) so a bounded debug dispatch can build the runtime target with all three formats to empirically confirm the runtime verify path. Co-authored-by: TRAE CLI --- .github/workflows/_docker-build-and-publish.yml | 8 ++++---- .github/workflows/release-docker-dev.yml | 6 +++++- scripts/ci/test_get_volcengine_image_tag.py | 0 3 files changed, 9 insertions(+), 5 deletions(-) mode change 100644 => 100755 scripts/ci/test_get_volcengine_image_tag.py diff --git a/.github/workflows/_docker-build-and-publish.yml b/.github/workflows/_docker-build-and-publish.yml index 986a065271cc..72164d072032 100644 --- a/.github/workflows/_docker-build-and-publish.yml +++ b/.github/workflows/_docker-build-and-publish.yml @@ -92,10 +92,10 @@ on: type: string default: "" image_formats: - description: "CSV of image formats to build and push: oci (default), zstd, nydus" + description: "CSV of image formats to build and push. Default builds all three (oci,zstd,nydus); pass a subset to narrow." required: false type: string - default: "oci" + default: "oci,zstd,nydus" jobs: build-and-publish: @@ -200,10 +200,10 @@ jobs: import os import sys - raw = os.environ.get("IMAGE_FORMATS", "oci") + raw = os.environ.get("IMAGE_FORMATS", "oci,zstd,nydus") tokens = [token.strip() for token in raw.split(",") if token.strip()] if not tokens: - tokens = ["oci"] + tokens = ["oci", "zstd", "nydus"] allowed = {"oci", "zstd", "nydus"} unknown = [token for token in tokens if token not in allowed] if unknown: diff --git a/.github/workflows/release-docker-dev.yml b/.github/workflows/release-docker-dev.yml index 65ca8aced1ca..c019b9164b7c 100644 --- a/.github/workflows/release-docker-dev.yml +++ b/.github/workflows/release-docker-dev.yml @@ -50,6 +50,10 @@ on: description: "Volcengine fork only: CSV of image formats for the debug base build (oci, zstd, nydus). Applies only when private_debug_base_only is set." required: false default: "oci,zstd,nydus" + private_debug_docker_target: + description: "Volcengine fork only: Dockerfile target stage for the debug base build (framework_final or runtime). Applies only when private_debug_base_only is set." + required: false + default: "framework_final" schedule: - cron: "0 0 * * *" @@ -353,7 +357,7 @@ jobs: if: ${{ github.repository == 'bytedance-iaas/sglang' && inputs.private_debug_base_only && inputs.private_debug_verify_image == '' }} uses: ./.github/workflows/_docker-build-and-publish.yml with: - docker_target: framework_final + docker_target: ${{ inputs.private_debug_docker_target || 'framework_final' }} checkout_ref: ${{ inputs.pr_number && format('refs/pull/{0}/head', inputs.pr_number) || github.ref }} cuda_key: cu130 cuda_version: 13.0.1 diff --git a/scripts/ci/test_get_volcengine_image_tag.py b/scripts/ci/test_get_volcengine_image_tag.py old mode 100644 new mode 100755 From 8f64ffd7a6091725512dba48bf606a13c7f92882 Mon Sep 17 00:00:00 2001 From: Hank Han Date: Fri, 14 Aug 2026 00:49:34 +0800 Subject: [PATCH 18/47] ci: reuse base build cache for zstd re-export via registry mode=max The zstd re-export step was recompiling the entire image from scratch (~124min, longer than the base build) because BuildKit's local layer cache was evicted between the base build and the zstd step by the large image plus the verify `docker pull`, so every intermediate stage (torch_deps `.[all]`, framework, deepep) re-ran with zero cache hits. Export a full mode=max registry build cache from the base build (only when a zstd format is requested) keyed to the resolved primary tag, and import it in the zstd step via --cache-from. mode=max is required because the heavy stages reach the final image via COPY --from and mode=min/inline would not capture them. The zstd step now hits cache for all layers and only redoes zstd compression. Co-authored-by: TRAE CLI --- .../workflows/_docker-build-and-publish.yml | 35 ++++++++++++++++--- 1 file changed, 31 insertions(+), 4 deletions(-) diff --git a/.github/workflows/_docker-build-and-publish.yml b/.github/workflows/_docker-build-and-publish.yml index 72164d072032..3613a7a80663 100644 --- a/.github/workflows/_docker-build-and-publish.yml +++ b/.github/workflows/_docker-build-and-publish.yml @@ -300,6 +300,8 @@ jobs: - name: Build and push AMD64 image id: build + env: + IMAGE_TAGS: ${{ steps.image-tag.outputs.image-tags }} run: | set -euo pipefail VERSION_ARG="" @@ -308,6 +310,21 @@ jobs: fi mapfile -t METADATA_ARGS < /tmp/docker-metadata.args + # When a zstd re-export will follow, export a full (mode=max) registry + # build cache so the zstd step can re-use these freshly built layers + # instead of recompiling from scratch. The heavy intermediate stages + # (torch_deps/framework/deepep) reach the final image via COPY --from, + # so mode=max is required to cache them; mode=min/inline would not. + # The cache ref is keyed to the resolved primary tag (unique per + # variant+cuda+run) to avoid collisions across concurrent matrix jobs. + CACHE_ARGS=() + if printf '%s' "${IMAGE_FORMATS}" | tr ',' '\n' | grep -qx "zstd"; then + PRIMARY_TAG=$(printf '%s\n' "${IMAGE_TAGS}" | grep -v '^[[:space:]]*$' | head -1 | tr -d '[:space:]') + CACHE_REF="${IMAGE_REPO}:buildcache-${PRIMARY_TAG}" + echo "cache-ref=${CACHE_REF}" >> "$GITHUB_OUTPUT" + CACHE_ARGS+=(--cache-to "type=registry,ref=${CACHE_REF},mode=max,image-manifest=true,oci-mediatypes=true") + fi + docker buildx build \ --target ${{ inputs.docker_target }} \ --platform linux/amd64 \ @@ -320,6 +337,7 @@ jobs: "${METADATA_ARGS[@]}" \ ${VERSION_ARG} \ ${{ inputs.extra_build_args }} \ + "${CACHE_ARGS[@]}" \ --metadata-file /tmp/metadata.json \ --no-cache \ . @@ -407,10 +425,13 @@ jobs: python3 -m pytest -q test/registered/unit/test_model_overrides.py -k deepseek_spec_moe_resolution ' - # --- zstd format: re-export the just-built (cached) layers with zstd - # compression, then reuse the same provenance/runtime verification. This - # is a cache hit (no --no-cache), so the config is identical to base and - # only layer compression changes. + # --- zstd format: re-export the just-built layers with zstd compression, + # then reuse the same provenance/runtime verification. It pulls the + # mode=max registry build cache written by the base step (--cache-from), + # so all intermediate stages are cache hits and only layer compression is + # redone. The config stays identical to base; only layer compression + # changes. (Without --cache-from the local BuildKit cache can be evicted + # between steps by the large image + verify pull, forcing a full rebuild.) - name: Build and push AMD64 zstd image id: build-zstd if: ${{ contains(inputs.image_formats, 'zstd') }} @@ -422,6 +443,11 @@ jobs: fi mapfile -t METADATA_ARGS < /tmp/docker-metadata.args + CACHE_FROM_ARGS=() + if [ -n "${{ steps.build.outputs.cache-ref }}" ]; then + CACHE_FROM_ARGS+=(--cache-from "type=registry,ref=${{ steps.build.outputs.cache-ref }}") + fi + docker buildx build \ --target ${{ inputs.docker_target }} \ --platform linux/amd64 \ @@ -434,6 +460,7 @@ jobs: "${METADATA_ARGS[@]}" \ ${VERSION_ARG} \ ${{ inputs.extra_build_args }} \ + "${CACHE_FROM_ARGS[@]}" \ --metadata-file /tmp/metadata-zstd.json \ . From 31bf24496a6f03c6414cc38ebbd195dc3207899d Mon Sep 17 00:00:00 2001 From: Hank Han Date: Fri, 14 Aug 2026 11:30:20 +0800 Subject: [PATCH 19/47] ci: fix kernel wheel path after deepseek_v4 moved AOT kernels to sgl-kernel/ The bytedance/deepseek_v4 branch relocated its kernel build tree from python/sglang/kernels/aot/ to sgl-kernel/. The nightly daily build (release-docker-deepseek-v4-nightly.yml) checks out deepseek_v4 for the kernel wheel but ran the reusable release-whl-kernel.yml from ep_main, which still cd'd into the removed python/sglang/kernels/aot/ path, failing build-kernel-wheel in <1s with "No such file or directory" and cascading to skip build-nightly. The workflow inputs were already described as sgl-kernel/build.sh, so the job body was an incomplete migration. Align all 35 path references (cd, artifact paths, version.py, working-directory) with sgl-kernel/, as already done in deepseek_v4's own release-whl-kernel.yml. Co-authored-by: TRAE CLI --- .github/workflows/release-whl-kernel.yml | 70 ++++++++++++------------ 1 file changed, 35 insertions(+), 35 deletions(-) diff --git a/.github/workflows/release-whl-kernel.yml b/.github/workflows/release-whl-kernel.yml index b4fe36ab9bf1..9e686d8ebe28 100644 --- a/.github/workflows/release-whl-kernel.yml +++ b/.github/workflows/release-whl-kernel.yml @@ -5,7 +5,7 @@ on: branches: - main paths: - - python/sglang/kernels/aot/python/sgl_kernel/version.py + - sgl-kernel/python/sgl_kernel/version.py workflow_call: inputs: checkout_ref: @@ -83,7 +83,7 @@ jobs: runs-on: ${{ inputs.runner }} steps: # Self-hosted build nodes retain the workspace across jobs. Prior builds - # leave root-owned artifacts under python/sglang/kernels/aot/build/ that actions/checkout + # leave root-owned artifacts under sgl-kernel/build/ that actions/checkout # cannot remove, causing EACCES on rmdir. Wipe them via a throwaway root # container before checkout recreates the workspace. - name: Clean workspace (remove root-owned files from prior runs) @@ -102,7 +102,7 @@ jobs: - name: Build wheels run: | - cd python/sglang/kernels/aot + cd sgl-kernel chmod +x ./build.sh if [ -n "${{ inputs.arch }}" ]; then ./build.sh "${{ inputs.python_version }}" "${{ inputs.cuda_version }}" "${{ inputs.arch }}" @@ -118,7 +118,7 @@ jobs: uses: actions/upload-artifact@v4 with: name: ${{ inputs.artifact_name }} - path: python/sglang/kernels/aot/dist/* + path: sgl-kernel/dist/* if-no-files-found: error # cu130 is the PyPI-released variant; cu129 wheels are published only to the @@ -141,7 +141,7 @@ jobs: runs-on: ${{ matrix.runner }} steps: # Self-hosted build nodes retain the workspace across jobs. Prior builds - # leave root-owned artifacts under python/sglang/kernels/aot/build/ that actions/checkout + # leave root-owned artifacts under sgl-kernel/build/ that actions/checkout # cannot remove, causing EACCES on rmdir. Wipe them via a throwaway root # container before checkout recreates the workspace. - name: Clean workspace (remove root-owned files from prior runs) @@ -161,7 +161,7 @@ jobs: - name: Build wheels run: | - cd python/sglang/kernels/aot + cd sgl-kernel chmod +x ./build.sh ./build.sh "${{ matrix.python-version }}" "${{ matrix.cuda-version }}" ${{ matrix.arch == 'aarch64' && 'aarch64' || '' }} env: @@ -172,7 +172,7 @@ jobs: uses: actions/upload-artifact@v4 with: name: wheel-python${{ matrix.python-version }}-cuda${{ matrix.cuda-version }}${{ matrix.arch == 'aarch64' && '-aarch64' || '' }} - path: python/sglang/kernels/aot/dist/* + path: sgl-kernel/dist/* release-cu129: needs: build-cu129-matrix @@ -185,7 +185,7 @@ jobs: - name: Download artifacts uses: actions/download-artifact@v4 with: - path: python/sglang/kernels/aot/dist/ + path: sgl-kernel/dist/ merge-multiple: true pattern: wheel-* @@ -193,7 +193,7 @@ jobs: id: set_tag_name run: | if [ -z "${{ inputs.tag_name }}" ]; then - TAG_NAME="v$(cat python/sglang/kernels/aot/python/sgl_kernel/version.py | cut -d'"' -f2)" + TAG_NAME="v$(cat sgl-kernel/python/sgl_kernel/version.py | cut -d'"' -f2)" echo "tag_name=$TAG_NAME" >> $GITHUB_OUTPUT else echo "tag_name=${{ inputs.tag_name }}" >> $GITHUB_OUTPUT @@ -206,7 +206,7 @@ jobs: repository: sgl-project/whl token: ${{ secrets.GH_PAT_FOR_WHL_RELEASE }} files: | - python/sglang/kernels/aot/dist/* + sgl-kernel/dist/* - name: Clone wheel index run: git clone https://oauth2:${WHL_TOKEN}@github.com/sgl-project/whl.git sgl-whl @@ -242,7 +242,7 @@ jobs: runs-on: ${{ matrix.runner }} steps: # Self-hosted build nodes retain the workspace across jobs. Prior builds - # leave root-owned artifacts under python/sglang/kernels/aot/build/ that actions/checkout + # leave root-owned artifacts under sgl-kernel/build/ that actions/checkout # cannot remove, causing EACCES on rmdir. Wipe them via a throwaway root # container before checkout recreates the workspace. - name: Clean workspace (remove root-owned files from prior runs) @@ -262,7 +262,7 @@ jobs: - name: Build wheels run: | - cd python/sglang/kernels/aot + cd sgl-kernel chmod +x ./build.sh ./build.sh "${{ matrix.python-version }}" "${{ matrix.cuda-version }}" ${{ matrix.arch == 'aarch64' && 'aarch64' || '' }} env: @@ -270,7 +270,7 @@ jobs: NVCC_THREADS: 8 - name: Strip +cu130 local version for PyPI upload - working-directory: python/sglang/kernels/aot + working-directory: sgl-kernel run: | set -eux pip install wheel @@ -295,7 +295,7 @@ jobs: ls -lh dist-pypi/ - name: Upload to PyPI - working-directory: python/sglang/kernels/aot + working-directory: sgl-kernel run: | pip install twine python3 -m twine upload --skip-existing dist-pypi/* -u __token__ -p ${{ secrets.PYPI_TOKEN_SGLANG_KERNEL }} @@ -304,7 +304,7 @@ jobs: uses: actions/upload-artifact@v4 with: name: wheel-python${{ matrix.python-version }}-cuda${{ matrix.cuda-version }}${{ matrix.arch == 'aarch64' && '-aarch64' || '' }} - path: python/sglang/kernels/aot/dist/* + path: sgl-kernel/dist/* release-cu130: needs: build-cu130-matrix @@ -317,7 +317,7 @@ jobs: - name: Download artifacts uses: actions/download-artifact@v4 with: - path: python/sglang/kernels/aot/dist/ + path: sgl-kernel/dist/ merge-multiple: true pattern: wheel-* @@ -325,7 +325,7 @@ jobs: id: set_tag_name run: | if [ -z "${{ inputs.tag_name }}" ]; then - TAG_NAME="v$(cat python/sglang/kernels/aot/python/sgl_kernel/version.py | cut -d'"' -f2)" + TAG_NAME="v$(cat sgl-kernel/python/sgl_kernel/version.py | cut -d'"' -f2)" echo "tag_name=$TAG_NAME" >> $GITHUB_OUTPUT else echo "tag_name=${{ inputs.tag_name }}" >> $GITHUB_OUTPUT @@ -338,7 +338,7 @@ jobs: repository: sgl-project/whl token: ${{ secrets.GH_PAT_FOR_WHL_RELEASE }} files: | - python/sglang/kernels/aot/dist/* + sgl-kernel/dist/* - name: Clone wheel index run: git clone https://oauth2:${WHL_TOKEN}@github.com/sgl-project/whl.git sgl-whl @@ -368,7 +368,7 @@ jobs: rocm-version: ["700", "720"] steps: # Self-hosted build nodes retain the workspace across jobs. Prior builds - # leave root-owned artifacts under python/sglang/kernels/aot/build/ that actions/checkout + # leave root-owned artifacts under sgl-kernel/build/ that actions/checkout # cannot remove, causing EACCES on rmdir. Wipe them via a throwaway root # container before checkout recreates the workspace. - name: Clean workspace (remove root-owned files from prior runs) @@ -388,8 +388,8 @@ jobs: - name: Build wheels run: | - cp 3rdparty/amd/wheel/sgl-kernel/* python/sglang/kernels/aot/ - cd python/sglang/kernels/aot + cp 3rdparty/amd/wheel/sgl-kernel/* sgl-kernel/ + cd sgl-kernel chmod +x ./build_rocm.sh ./build_rocm.sh "${{ matrix.rocm-version }}" @@ -397,7 +397,7 @@ jobs: uses: actions/upload-artifact@v4 with: name: wheel-python${{ matrix.python-version }}-rocm${{ matrix.rocm-version }} - path: python/sglang/kernels/aot/dist/* + path: sgl-kernel/dist/* release-rocm700: needs: build-rocm-matrix @@ -410,7 +410,7 @@ jobs: - name: Download artifacts uses: actions/download-artifact@v4 with: - path: python/sglang/kernels/aot/dist/ + path: sgl-kernel/dist/ merge-multiple: true pattern: wheel-*-rocm700 @@ -418,7 +418,7 @@ jobs: id: set_tag_name run: | if [ -z "${{ inputs.tag_name }}" ]; then - TAG_NAME="v$(cat python/sglang/kernels/aot/python/sgl_kernel/version.py | cut -d'"' -f2)" + TAG_NAME="v$(cat sgl-kernel/python/sgl_kernel/version.py | cut -d'"' -f2)" echo "tag_name=$TAG_NAME" >> $GITHUB_OUTPUT else echo "tag_name=${{ inputs.tag_name }}" >> $GITHUB_OUTPUT @@ -431,7 +431,7 @@ jobs: repository: sgl-project/whl token: ${{ secrets.GH_PAT_FOR_WHL_RELEASE }} files: | - python/sglang/kernels/aot/dist/* + sgl-kernel/dist/* - name: Clone wheel index run: git clone https://oauth2:${WHL_TOKEN}@github.com/sgl-project/whl.git sgl-whl @@ -461,7 +461,7 @@ jobs: - name: Download artifacts uses: actions/download-artifact@v4 with: - path: python/sglang/kernels/aot/dist/ + path: sgl-kernel/dist/ merge-multiple: true pattern: wheel-*-rocm720 @@ -469,7 +469,7 @@ jobs: id: set_tag_name run: | if [ -z "${{ inputs.tag_name }}" ]; then - TAG_NAME="v$(cat python/sglang/kernels/aot/python/sgl_kernel/version.py | cut -d'"' -f2)" + TAG_NAME="v$(cat sgl-kernel/python/sgl_kernel/version.py | cut -d'"' -f2)" echo "tag_name=$TAG_NAME" >> $GITHUB_OUTPUT else echo "tag_name=${{ inputs.tag_name }}" >> $GITHUB_OUTPUT @@ -482,7 +482,7 @@ jobs: repository: sgl-project/whl token: ${{ secrets.GH_PAT_FOR_WHL_RELEASE }} files: | - python/sglang/kernels/aot/dist/* + sgl-kernel/dist/* - name: Clone wheel index run: git clone https://oauth2:${WHL_TOKEN}@github.com/sgl-project/whl.git sgl-whl @@ -512,7 +512,7 @@ jobs: musa-version: ["43"] steps: # Self-hosted build nodes retain the workspace across jobs. Prior builds - # leave root-owned artifacts under python/sglang/kernels/aot/build/ that actions/checkout + # leave root-owned artifacts under sgl-kernel/build/ that actions/checkout # cannot remove, causing EACCES on rmdir. Wipe them via a throwaway root # container before checkout recreates the workspace. - name: Clean workspace (remove root-owned files from prior runs) @@ -538,19 +538,19 @@ jobs: - name: Build wheels run: | - cd python/sglang/kernels/aot + cd sgl-kernel mv pyproject_musa.toml pyproject.toml python setup_musa.py sdist bdist_wheel - name: Rename MUSA wheels run: | - bash scripts/ci/musa/rename_wheels_musa.sh ${{ matrix.musa-version }} python/sglang/kernels/aot/dist + bash scripts/ci/musa/rename_wheels_musa.sh ${{ matrix.musa-version }} sgl-kernel/dist - name: Upload artifacts uses: actions/upload-artifact@v4 with: name: wheel-python${{ matrix.python-version }}-musa${{ matrix.musa-version }} - path: python/sglang/kernels/aot/dist/* + path: sgl-kernel/dist/* release-musa43: needs: build-musa43 @@ -561,7 +561,7 @@ jobs: - name: Download artifacts uses: actions/download-artifact@v4 with: - path: python/sglang/kernels/aot/dist/ + path: sgl-kernel/dist/ merge-multiple: true pattern: wheel-* @@ -569,7 +569,7 @@ jobs: id: set_tag_name run: | if [ -z "${{ inputs.tag_name }}" ]; then - TAG_NAME="v$(cat python/sglang/kernels/aot/python/sgl_kernel/version.py | cut -d'"' -f2)" + TAG_NAME="v$(cat sgl-kernel/python/sgl_kernel/version.py | cut -d'"' -f2)" echo "tag_name=$TAG_NAME" >> $GITHUB_OUTPUT else echo "tag_name=${{ inputs.tag_name }}" >> $GITHUB_OUTPUT @@ -582,7 +582,7 @@ jobs: repository: sgl-project/whl token: ${{ secrets.GH_PAT_FOR_WHL_RELEASE }} files: | - python/sglang/kernels/aot/dist/* + sgl-kernel/dist/* - name: Clone wheel index run: git clone https://oauth2:${WHL_TOKEN}@github.com/sgl-project/whl.git sgl-whl From b806888d4319163d1fd8ea036828d0f3455c81eb Mon Sep 17 00:00:00 2001 From: Hank Han Date: Fri, 14 Aug 2026 15:03:33 +0800 Subject: [PATCH 20/47] ci: split nightly kernel build to deepseek_v4's own workflow The nightly builds deepseek_v4 source but reused ep_main's release-whl-kernel.yml, coupling that shared file to deepseek_v4's directory layout. ep_main follows upstream (kernels under python/sglang/kernels/aot/, per #32648) while deepseek_v4 keeps them under sgl-kernel/, so one shared file cannot serve both. Split the two consumers: - Nightly now calls the kernel workflow from the branch it builds: uses: .../release-whl-kernel.yml@bytedance/deepseek_v4 (which now exposes a workflow_call build-wheel job using sgl-kernel/). - ep_main's release-whl-kernel.yml is reverted to python/sglang/kernels/aot, undoing the earlier sgl-kernel path change (31bf24496a). Its only local consumer, release-docker-dev.yml, builds ep_main source, which uses the upstream python/sglang/kernels/aot/ layout. Each branch's kernel build now tracks its own source structure. Co-authored-by: TRAE CLI --- .../release-docker-deepseek-v4-nightly.yml | 5 +- .github/workflows/release-whl-kernel.yml | 70 +++++++++---------- 2 files changed, 39 insertions(+), 36 deletions(-) diff --git a/.github/workflows/release-docker-deepseek-v4-nightly.yml b/.github/workflows/release-docker-deepseek-v4-nightly.yml index e96308e2bc8c..df3b5c08c40b 100644 --- a/.github/workflows/release-docker-deepseek-v4-nightly.yml +++ b/.github/workflows/release-docker-deepseek-v4-nightly.yml @@ -130,7 +130,10 @@ jobs: build-kernel-wheel: needs: [resolve-target-ref] if: ${{ github.repository == 'bytedance-iaas/sglang' && (github.event_name == 'schedule' || inputs.compile_kernel) }} - uses: ./.github/workflows/release-whl-kernel.yml + # Reference the kernel build workflow from the deepseek_v4 branch itself, so + # the wheel build logic tracks that branch's source layout (sgl-kernel/) + # instead of forcing ep_main's shared release-whl-kernel.yml to chase it. + uses: bytedance-iaas/sglang/.github/workflows/release-whl-kernel.yml@bytedance/deepseek_v4 with: checkout_ref: ${{ needs.resolve-target-ref.outputs.target_ref }} cuda_version: "13.0" diff --git a/.github/workflows/release-whl-kernel.yml b/.github/workflows/release-whl-kernel.yml index 9e686d8ebe28..b4fe36ab9bf1 100644 --- a/.github/workflows/release-whl-kernel.yml +++ b/.github/workflows/release-whl-kernel.yml @@ -5,7 +5,7 @@ on: branches: - main paths: - - sgl-kernel/python/sgl_kernel/version.py + - python/sglang/kernels/aot/python/sgl_kernel/version.py workflow_call: inputs: checkout_ref: @@ -83,7 +83,7 @@ jobs: runs-on: ${{ inputs.runner }} steps: # Self-hosted build nodes retain the workspace across jobs. Prior builds - # leave root-owned artifacts under sgl-kernel/build/ that actions/checkout + # leave root-owned artifacts under python/sglang/kernels/aot/build/ that actions/checkout # cannot remove, causing EACCES on rmdir. Wipe them via a throwaway root # container before checkout recreates the workspace. - name: Clean workspace (remove root-owned files from prior runs) @@ -102,7 +102,7 @@ jobs: - name: Build wheels run: | - cd sgl-kernel + cd python/sglang/kernels/aot chmod +x ./build.sh if [ -n "${{ inputs.arch }}" ]; then ./build.sh "${{ inputs.python_version }}" "${{ inputs.cuda_version }}" "${{ inputs.arch }}" @@ -118,7 +118,7 @@ jobs: uses: actions/upload-artifact@v4 with: name: ${{ inputs.artifact_name }} - path: sgl-kernel/dist/* + path: python/sglang/kernels/aot/dist/* if-no-files-found: error # cu130 is the PyPI-released variant; cu129 wheels are published only to the @@ -141,7 +141,7 @@ jobs: runs-on: ${{ matrix.runner }} steps: # Self-hosted build nodes retain the workspace across jobs. Prior builds - # leave root-owned artifacts under sgl-kernel/build/ that actions/checkout + # leave root-owned artifacts under python/sglang/kernels/aot/build/ that actions/checkout # cannot remove, causing EACCES on rmdir. Wipe them via a throwaway root # container before checkout recreates the workspace. - name: Clean workspace (remove root-owned files from prior runs) @@ -161,7 +161,7 @@ jobs: - name: Build wheels run: | - cd sgl-kernel + cd python/sglang/kernels/aot chmod +x ./build.sh ./build.sh "${{ matrix.python-version }}" "${{ matrix.cuda-version }}" ${{ matrix.arch == 'aarch64' && 'aarch64' || '' }} env: @@ -172,7 +172,7 @@ jobs: uses: actions/upload-artifact@v4 with: name: wheel-python${{ matrix.python-version }}-cuda${{ matrix.cuda-version }}${{ matrix.arch == 'aarch64' && '-aarch64' || '' }} - path: sgl-kernel/dist/* + path: python/sglang/kernels/aot/dist/* release-cu129: needs: build-cu129-matrix @@ -185,7 +185,7 @@ jobs: - name: Download artifacts uses: actions/download-artifact@v4 with: - path: sgl-kernel/dist/ + path: python/sglang/kernels/aot/dist/ merge-multiple: true pattern: wheel-* @@ -193,7 +193,7 @@ jobs: id: set_tag_name run: | if [ -z "${{ inputs.tag_name }}" ]; then - TAG_NAME="v$(cat sgl-kernel/python/sgl_kernel/version.py | cut -d'"' -f2)" + TAG_NAME="v$(cat python/sglang/kernels/aot/python/sgl_kernel/version.py | cut -d'"' -f2)" echo "tag_name=$TAG_NAME" >> $GITHUB_OUTPUT else echo "tag_name=${{ inputs.tag_name }}" >> $GITHUB_OUTPUT @@ -206,7 +206,7 @@ jobs: repository: sgl-project/whl token: ${{ secrets.GH_PAT_FOR_WHL_RELEASE }} files: | - sgl-kernel/dist/* + python/sglang/kernels/aot/dist/* - name: Clone wheel index run: git clone https://oauth2:${WHL_TOKEN}@github.com/sgl-project/whl.git sgl-whl @@ -242,7 +242,7 @@ jobs: runs-on: ${{ matrix.runner }} steps: # Self-hosted build nodes retain the workspace across jobs. Prior builds - # leave root-owned artifacts under sgl-kernel/build/ that actions/checkout + # leave root-owned artifacts under python/sglang/kernels/aot/build/ that actions/checkout # cannot remove, causing EACCES on rmdir. Wipe them via a throwaway root # container before checkout recreates the workspace. - name: Clean workspace (remove root-owned files from prior runs) @@ -262,7 +262,7 @@ jobs: - name: Build wheels run: | - cd sgl-kernel + cd python/sglang/kernels/aot chmod +x ./build.sh ./build.sh "${{ matrix.python-version }}" "${{ matrix.cuda-version }}" ${{ matrix.arch == 'aarch64' && 'aarch64' || '' }} env: @@ -270,7 +270,7 @@ jobs: NVCC_THREADS: 8 - name: Strip +cu130 local version for PyPI upload - working-directory: sgl-kernel + working-directory: python/sglang/kernels/aot run: | set -eux pip install wheel @@ -295,7 +295,7 @@ jobs: ls -lh dist-pypi/ - name: Upload to PyPI - working-directory: sgl-kernel + working-directory: python/sglang/kernels/aot run: | pip install twine python3 -m twine upload --skip-existing dist-pypi/* -u __token__ -p ${{ secrets.PYPI_TOKEN_SGLANG_KERNEL }} @@ -304,7 +304,7 @@ jobs: uses: actions/upload-artifact@v4 with: name: wheel-python${{ matrix.python-version }}-cuda${{ matrix.cuda-version }}${{ matrix.arch == 'aarch64' && '-aarch64' || '' }} - path: sgl-kernel/dist/* + path: python/sglang/kernels/aot/dist/* release-cu130: needs: build-cu130-matrix @@ -317,7 +317,7 @@ jobs: - name: Download artifacts uses: actions/download-artifact@v4 with: - path: sgl-kernel/dist/ + path: python/sglang/kernels/aot/dist/ merge-multiple: true pattern: wheel-* @@ -325,7 +325,7 @@ jobs: id: set_tag_name run: | if [ -z "${{ inputs.tag_name }}" ]; then - TAG_NAME="v$(cat sgl-kernel/python/sgl_kernel/version.py | cut -d'"' -f2)" + TAG_NAME="v$(cat python/sglang/kernels/aot/python/sgl_kernel/version.py | cut -d'"' -f2)" echo "tag_name=$TAG_NAME" >> $GITHUB_OUTPUT else echo "tag_name=${{ inputs.tag_name }}" >> $GITHUB_OUTPUT @@ -338,7 +338,7 @@ jobs: repository: sgl-project/whl token: ${{ secrets.GH_PAT_FOR_WHL_RELEASE }} files: | - sgl-kernel/dist/* + python/sglang/kernels/aot/dist/* - name: Clone wheel index run: git clone https://oauth2:${WHL_TOKEN}@github.com/sgl-project/whl.git sgl-whl @@ -368,7 +368,7 @@ jobs: rocm-version: ["700", "720"] steps: # Self-hosted build nodes retain the workspace across jobs. Prior builds - # leave root-owned artifacts under sgl-kernel/build/ that actions/checkout + # leave root-owned artifacts under python/sglang/kernels/aot/build/ that actions/checkout # cannot remove, causing EACCES on rmdir. Wipe them via a throwaway root # container before checkout recreates the workspace. - name: Clean workspace (remove root-owned files from prior runs) @@ -388,8 +388,8 @@ jobs: - name: Build wheels run: | - cp 3rdparty/amd/wheel/sgl-kernel/* sgl-kernel/ - cd sgl-kernel + cp 3rdparty/amd/wheel/sgl-kernel/* python/sglang/kernels/aot/ + cd python/sglang/kernels/aot chmod +x ./build_rocm.sh ./build_rocm.sh "${{ matrix.rocm-version }}" @@ -397,7 +397,7 @@ jobs: uses: actions/upload-artifact@v4 with: name: wheel-python${{ matrix.python-version }}-rocm${{ matrix.rocm-version }} - path: sgl-kernel/dist/* + path: python/sglang/kernels/aot/dist/* release-rocm700: needs: build-rocm-matrix @@ -410,7 +410,7 @@ jobs: - name: Download artifacts uses: actions/download-artifact@v4 with: - path: sgl-kernel/dist/ + path: python/sglang/kernels/aot/dist/ merge-multiple: true pattern: wheel-*-rocm700 @@ -418,7 +418,7 @@ jobs: id: set_tag_name run: | if [ -z "${{ inputs.tag_name }}" ]; then - TAG_NAME="v$(cat sgl-kernel/python/sgl_kernel/version.py | cut -d'"' -f2)" + TAG_NAME="v$(cat python/sglang/kernels/aot/python/sgl_kernel/version.py | cut -d'"' -f2)" echo "tag_name=$TAG_NAME" >> $GITHUB_OUTPUT else echo "tag_name=${{ inputs.tag_name }}" >> $GITHUB_OUTPUT @@ -431,7 +431,7 @@ jobs: repository: sgl-project/whl token: ${{ secrets.GH_PAT_FOR_WHL_RELEASE }} files: | - sgl-kernel/dist/* + python/sglang/kernels/aot/dist/* - name: Clone wheel index run: git clone https://oauth2:${WHL_TOKEN}@github.com/sgl-project/whl.git sgl-whl @@ -461,7 +461,7 @@ jobs: - name: Download artifacts uses: actions/download-artifact@v4 with: - path: sgl-kernel/dist/ + path: python/sglang/kernels/aot/dist/ merge-multiple: true pattern: wheel-*-rocm720 @@ -469,7 +469,7 @@ jobs: id: set_tag_name run: | if [ -z "${{ inputs.tag_name }}" ]; then - TAG_NAME="v$(cat sgl-kernel/python/sgl_kernel/version.py | cut -d'"' -f2)" + TAG_NAME="v$(cat python/sglang/kernels/aot/python/sgl_kernel/version.py | cut -d'"' -f2)" echo "tag_name=$TAG_NAME" >> $GITHUB_OUTPUT else echo "tag_name=${{ inputs.tag_name }}" >> $GITHUB_OUTPUT @@ -482,7 +482,7 @@ jobs: repository: sgl-project/whl token: ${{ secrets.GH_PAT_FOR_WHL_RELEASE }} files: | - sgl-kernel/dist/* + python/sglang/kernels/aot/dist/* - name: Clone wheel index run: git clone https://oauth2:${WHL_TOKEN}@github.com/sgl-project/whl.git sgl-whl @@ -512,7 +512,7 @@ jobs: musa-version: ["43"] steps: # Self-hosted build nodes retain the workspace across jobs. Prior builds - # leave root-owned artifacts under sgl-kernel/build/ that actions/checkout + # leave root-owned artifacts under python/sglang/kernels/aot/build/ that actions/checkout # cannot remove, causing EACCES on rmdir. Wipe them via a throwaway root # container before checkout recreates the workspace. - name: Clean workspace (remove root-owned files from prior runs) @@ -538,19 +538,19 @@ jobs: - name: Build wheels run: | - cd sgl-kernel + cd python/sglang/kernels/aot mv pyproject_musa.toml pyproject.toml python setup_musa.py sdist bdist_wheel - name: Rename MUSA wheels run: | - bash scripts/ci/musa/rename_wheels_musa.sh ${{ matrix.musa-version }} sgl-kernel/dist + bash scripts/ci/musa/rename_wheels_musa.sh ${{ matrix.musa-version }} python/sglang/kernels/aot/dist - name: Upload artifacts uses: actions/upload-artifact@v4 with: name: wheel-python${{ matrix.python-version }}-musa${{ matrix.musa-version }} - path: sgl-kernel/dist/* + path: python/sglang/kernels/aot/dist/* release-musa43: needs: build-musa43 @@ -561,7 +561,7 @@ jobs: - name: Download artifacts uses: actions/download-artifact@v4 with: - path: sgl-kernel/dist/ + path: python/sglang/kernels/aot/dist/ merge-multiple: true pattern: wheel-* @@ -569,7 +569,7 @@ jobs: id: set_tag_name run: | if [ -z "${{ inputs.tag_name }}" ]; then - TAG_NAME="v$(cat sgl-kernel/python/sgl_kernel/version.py | cut -d'"' -f2)" + TAG_NAME="v$(cat python/sglang/kernels/aot/python/sgl_kernel/version.py | cut -d'"' -f2)" echo "tag_name=$TAG_NAME" >> $GITHUB_OUTPUT else echo "tag_name=${{ inputs.tag_name }}" >> $GITHUB_OUTPUT @@ -582,7 +582,7 @@ jobs: repository: sgl-project/whl token: ${{ secrets.GH_PAT_FOR_WHL_RELEASE }} files: | - sgl-kernel/dist/* + python/sglang/kernels/aot/dist/* - name: Clone wheel index run: git clone https://oauth2:${WHL_TOKEN}@github.com/sgl-project/whl.git sgl-whl From 7ae8a83cdd7bb096da088415af79c9ccba74525d Mon Sep 17 00:00:00 2001 From: Hank Han Date: Mon, 17 Aug 2026 17:58:38 +0800 Subject: [PATCH 21/47] ci: add backfill-image-formats workflow for zstd/nydus format backfill Add a workflow_dispatch workflow on x64-docker-build-node that backfills zstd and nydus image formats for existing OCI images in the Volcengine serving registry. Uses the same buildx and nydusify commands as the private delivery build pipeline. Co-authored-by: TRAE CLI --- .github/workflows/backfill-image-formats.yml | 192 +++++++++++++++++++ 1 file changed, 192 insertions(+) create mode 100644 .github/workflows/backfill-image-formats.yml diff --git a/.github/workflows/backfill-image-formats.yml b/.github/workflows/backfill-image-formats.yml new file mode 100644 index 000000000000..daf7e266ee11 --- /dev/null +++ b/.github/workflows/backfill-image-formats.yml @@ -0,0 +1,192 @@ +name: Backfill Image Formats (zstd + nydus) + +on: + workflow_dispatch: + inputs: + image_refs: + description: | + 镜像引用列表,每行一个,不含格式后缀。 + 格式: iaas-gpu-cn-beijing.cr.volces.com/serving/: + 示例: iaas-gpu-cn-beijing.cr.volces.com/serving/sglang:v0.5.16.iaas.202607280000-kimi-k3-cu130 + required: true + type: string + sync_to_customer: + description: "补全格式后触发 Bits 流水线同步到客户仓库" + required: false + type: boolean + default: true + +concurrency: + group: backfill-image-formats-${{ github.run_id }} + cancel-in-progress: false + +jobs: + backfill: + if: github.repository == 'bytedance-iaas/sglang' + runs-on: x64-docker-build-node + environment: prod + env: + VOLCENGINE_CR_REGISTRY: ${{ vars.VOLCENGINE_CR_REGISTRY }} + VOLCENGINE_CR_NAMESPACE: ${{ vars.VOLCENGINE_CR_NAMESPACE }} + VOLCENGINE_CR_REPOSITORY: ${{ vars.VOLCENGINE_CR_REPOSITORY || 'sglang' }} + VOLCENGINE_CR_USERNAME: ${{ secrets.VOLCENGINE_CR_USERNAME }} + VOLCENGINE_CR_PASSWORD: ${{ secrets.VOLCENGINE_CR_PASSWORD }} + NYDUS_VERSION: "2.3.0" + steps: + - name: Delete huge unnecessary tools folder + run: rm -rf /opt/hostedtoolcache + + - name: Cleanup workspace + run: sudo rm -rf "$GITHUB_WORKSPACE"/* || true + + - name: Checkout repository + uses: actions/checkout@v4 + + - name: Set up Docker Buildx + uses: docker/setup-buildx-action@v3 + + - name: Login to Volcengine CR + run: | + set -euo pipefail + echo "${VOLCENGINE_CR_PASSWORD}" | docker login "${VOLCENGINE_CR_REGISTRY}" --username "${VOLCENGINE_CR_USERNAME}" --password-stdin + + - name: Install nydus tooling + run: | + set -euo pipefail + if command -v nydusify &>/dev/null && command -v nydus-image &>/dev/null; then + echo "[nydus] tools already installed: $(nydusify --version 2>&1 | head -1)" + exit 0 + fi + tarball="nydus-static-v${NYDUS_VERSION}-linux-amd64.tgz" + url="https://github.com/dragonflyoss/nydus/releases/download/v${NYDUS_VERSION}/${tarball}" + retry() { for i in $(seq 1 5); do "$@" && return 0; echo "attempt $i failed; sleeping 15s"; sleep 15; done; return 1; } + retry curl -fL --retry 5 --retry-all-errors --retry-delay 2 -o "/tmp/${tarball}" "${url}" + tar -xzf "/tmp/${tarball}" -C /tmp + sudo install -m 0755 /tmp/nydus-static/nydusify /usr/local/bin/nydusify + sudo install -m 0755 /tmp/nydus-static/nydus-image /usr/local/bin/nydus-image + nydusify --version + nydus-image --version + + - name: Backfill image formats + env: + IMAGE_REFS: ${{ inputs.image_refs }} + run: | + set -euo pipefail + + retry() { + for i in $(seq 1 5); do + "$@" && return 0 + echo " [retry] attempt $i/5 failed; sleeping 15s" + sleep 15 + done + return 1 + } + + backfill_one() { + local image_ref="$1" + + echo "" + echo "=== $(date '+%H:%M:%S') ${image_ref} ===" + + # Verify OCI base exists + if ! skopeo inspect "docker://${image_ref}" >/dev/null 2>&1; then + echo " ERROR: base image not found: ${image_ref}" + return 1 + fi + local oci_digest + oci_digest=$(skopeo inspect "docker://${image_ref}" | jq -r '.Digest') + echo " OCI digest: ${oci_digest}" + + # zstd: pull then re-export with zstd compression (same as CI zstd step) + if skopeo inspect "docker://${image_ref}-zstd" >/dev/null 2>&1; then + echo " [zstd] already exists" + else + echo " [zstd] backfilling..." + docker pull "${image_ref}" + tmp_dockerfile=$(mktemp /tmp/Dockerfile.zstd.XXXXXX) + echo "FROM ${image_ref}" > "${tmp_dockerfile}" + docker buildx build \ + --platform linux/amd64 \ + --output "type=image,name=${image_ref}-zstd,push=true,compression=zstd,compression-level=3,force-compression=true,oci-mediatypes=true" \ + -f "${tmp_dockerfile}" \ + /tmp + rm -f "${tmp_dockerfile}" + echo " [zstd] published ${image_ref}-zstd" + fi + + # nydus: nydusify convert (same as CI nydus step) + if skopeo inspect "docker://${image_ref}-nydus" >/dev/null 2>&1; then + echo " [nydus] already exists" + else + echo " [nydus] backfilling..." + retry nydusify convert \ + --nydus-image /usr/local/bin/nydus-image \ + --source "${image_ref}" \ + --target "${image_ref}-nydus" + echo " [nydus] published ${image_ref}-nydus" + fi + + echo " done: ${image_ref}" + echo "" + } + + while IFS= read -r line; do + line="${line//[$'\r\n']/}" + [[ -z "$line" || "$line" == \#* ]] && continue + backfill_one "$line" + done <<< "${IMAGE_REFS}" + + echo "All backfills complete." + + - name: Trigger Bits sync pipeline + if: ${{ inputs.sync_to_customer }} + env: + IMAGE_REFS: ${{ inputs.image_refs }} + run: | + set -euo pipefail + + BITS_PIPELINE_ID="1008460003842" + BITS_SPACE_ID="779378790402" + + for image_ref in ${IMAGE_REFS}; do + image_ref="${image_ref//[$'\r\n']/}" + [[ -z "$image_ref" ]] && continue + + # Parse framework and tag from image_ref + # iaas-gpu-cn-beijing.cr.volces.com/serving/sglang:v0.5.16.iaas.202607280000-kimi-k3-cu130 + framework="${image_ref##*/}" # sglang:v0.5.16... + framework="${framework%%:*}" # sglang + tag="${image_ref##*:}" # v0.5.16.iaas.202607280000-kimi-k3-cu130 + + echo "=== Triggering Bits sync for ${framework}:${tag} ===" + + # Try to trigger via bytedcli + bits_pipeline_cli + if command -v bytedcli &>/dev/null && command -v bits_pipeline_cli &>/dev/null; then + JWT=$(bytedcli auth get-bytecloud-jwt-token --json 2>/dev/null | python3 -c "import sys,json; print(json.load(sys.stdin)['data']['jwt'])" 2>/dev/null || echo "") + if [ -n "$JWT" ]; then + bits_pipeline_cli call \ + --env cn \ + --username hanhan.hank \ + --rpc RunPipeline \ + --path-param "pipeline_id=${BITS_PIPELINE_ID}" \ + --jwt-token "$JWT" \ + --body-json "$(python3 -c " +import json +print(json.dumps({ + 'run_by': 'hanhan.hank', + 'custom_vars': [ + {'name': 'custom.framework', 'value': {'text': '${framework}'}}, + {'name': 'custom.tag', 'value': {'text': '${tag}'}}, + ] +})) + ")" + echo " Bits pipeline triggered: ${framework}:${tag}" + else + echo " ::warning:: could not get JWT; trigger manually" + fi + else + echo " ::notice:: bytedcli/bits_pipeline_cli not available on this runner" + echo " Manual trigger: https://bits.bytedance.net/devops/${BITS_SPACE_ID}/pipeline/detail/${BITS_PIPELINE_ID}" + echo " Variables: framework=${framework} tag=${tag}" + fi + done \ No newline at end of file From 9fd6ff5ea9e1efbd4c609835ebf63ede6270ae7a Mon Sep 17 00:00:00 2001 From: Hank Han Date: Mon, 17 Aug 2026 18:06:31 +0800 Subject: [PATCH 22/47] ci: fix backfill workflow sync_to_customer default and simplify trigger step Co-authored-by: TRAE CLI --- .github/workflows/backfill-image-formats.yml | 50 ++++---------------- 1 file changed, 10 insertions(+), 40 deletions(-) diff --git a/.github/workflows/backfill-image-formats.yml b/.github/workflows/backfill-image-formats.yml index daf7e266ee11..23a3603a7667 100644 --- a/.github/workflows/backfill-image-formats.yml +++ b/.github/workflows/backfill-image-formats.yml @@ -14,7 +14,7 @@ on: description: "补全格式后触发 Bits 流水线同步到客户仓库" required: false type: boolean - default: true + default: false concurrency: group: backfill-image-formats-${{ github.run_id }} @@ -148,45 +148,15 @@ jobs: BITS_PIPELINE_ID="1008460003842" BITS_SPACE_ID="779378790402" + echo "Images to sync:" for image_ref in ${IMAGE_REFS}; do image_ref="${image_ref//[$'\r\n']/}" [[ -z "$image_ref" ]] && continue - - # Parse framework and tag from image_ref - # iaas-gpu-cn-beijing.cr.volces.com/serving/sglang:v0.5.16.iaas.202607280000-kimi-k3-cu130 - framework="${image_ref##*/}" # sglang:v0.5.16... - framework="${framework%%:*}" # sglang - tag="${image_ref##*:}" # v0.5.16.iaas.202607280000-kimi-k3-cu130 - - echo "=== Triggering Bits sync for ${framework}:${tag} ===" - - # Try to trigger via bytedcli + bits_pipeline_cli - if command -v bytedcli &>/dev/null && command -v bits_pipeline_cli &>/dev/null; then - JWT=$(bytedcli auth get-bytecloud-jwt-token --json 2>/dev/null | python3 -c "import sys,json; print(json.load(sys.stdin)['data']['jwt'])" 2>/dev/null || echo "") - if [ -n "$JWT" ]; then - bits_pipeline_cli call \ - --env cn \ - --username hanhan.hank \ - --rpc RunPipeline \ - --path-param "pipeline_id=${BITS_PIPELINE_ID}" \ - --jwt-token "$JWT" \ - --body-json "$(python3 -c " -import json -print(json.dumps({ - 'run_by': 'hanhan.hank', - 'custom_vars': [ - {'name': 'custom.framework', 'value': {'text': '${framework}'}}, - {'name': 'custom.tag', 'value': {'text': '${tag}'}}, - ] -})) - ")" - echo " Bits pipeline triggered: ${framework}:${tag}" - else - echo " ::warning:: could not get JWT; trigger manually" - fi - else - echo " ::notice:: bytedcli/bits_pipeline_cli not available on this runner" - echo " Manual trigger: https://bits.bytedance.net/devops/${BITS_SPACE_ID}/pipeline/detail/${BITS_PIPELINE_ID}" - echo " Variables: framework=${framework} tag=${tag}" - fi - done \ No newline at end of file + framework="${image_ref##*/}"; framework="${framework%%:*}" + tag="${image_ref##*:}" + echo " framework=${framework} tag=${tag}" + done + + echo "" + echo "Bits 流水线: https://bits.bytedance.net/devops/${BITS_SPACE_ID}/pipeline/detail/${BITS_PIPELINE_ID}" + echo "请手动触发流水线,填入对应的 framework 和 tag。" \ No newline at end of file From a3a67ff05b39e30433ff2e6962dd9391b1559d58 Mon Sep 17 00:00:00 2001 From: Hank Han Date: Mon, 17 Aug 2026 21:35:37 +0800 Subject: [PATCH 23/47] ci: make in-container image verify opt-out for deepseek_v4 nightly The shared _docker-build-and-publish.yml verify step runs ep_main-specific runtime smoke tests inside the freshly built image: `scripts/eic_integration_check.py` (py_compile) plus pytest for `test_runtime_context.py::TestMoeFlagsGroup` and `test_model_overrides.py::deepseek_spec_moe_resolution`. Those tests import `sglang.srt.runtime_context` and `sglang.srt.arg_groups.arg_utils`, which exist only on ep_main; the deepseek_v4 source tree has neither the modules nor the test files. The nightly builds v4 source through ep_main's reusable workflow, so the verify step failed after the image built and pushed successfully, skipping the zstd/nydus stages. Add a `verify_image` boolean input (default true, preserving the ep_main dev/runtime/release callers) that gates all three verify steps, and pass `verify_image: false` from the deepseek_v4 nightly. The provenance labels/env are still baked into every image regardless; only the in-container smoke tests are skipped for v4. Co-authored-by: TRAE CLI --- .github/workflows/_docker-build-and-publish.yml | 10 ++++++++-- .../workflows/release-docker-deepseek-v4-nightly.yml | 5 +++++ 2 files changed, 13 insertions(+), 2 deletions(-) diff --git a/.github/workflows/_docker-build-and-publish.yml b/.github/workflows/_docker-build-and-publish.yml index 3613a7a80663..04fb6a91a928 100644 --- a/.github/workflows/_docker-build-and-publish.yml +++ b/.github/workflows/_docker-build-and-publish.yml @@ -96,6 +96,11 @@ on: required: false type: string default: "oci,zstd,nydus" + verify_image: + description: "Run in-container provenance/runtime verification after each pushed image. Default true. Callers that build a source tree without the ep_main runtime smoke tests (e.g. the deepseek_v4 nightly) pass false; provenance labels/env are still baked into the image regardless." + required: false + type: boolean + default: true jobs: build-and-publish: @@ -367,6 +372,7 @@ jobs: done - name: Verify pushed AMD64 image and source provenance + if: ${{ inputs.verify_image }} env: EXPECTED_BUILD_COMMIT: ${{ steps.build-metadata.outputs.build-commit }} EXPECTED_BUILD_SOURCE: ${{ steps.build-metadata.outputs.build-source }} @@ -490,7 +496,7 @@ jobs: done - name: Verify pushed AMD64 zstd image and source provenance - if: ${{ contains(inputs.image_formats, 'zstd') }} + if: ${{ inputs.verify_image && contains(inputs.image_formats, 'zstd') }} env: EXPECTED_BUILD_COMMIT: ${{ steps.build-metadata.outputs.build-commit }} EXPECTED_BUILD_SOURCE: ${{ steps.build-metadata.outputs.build-source }} @@ -618,7 +624,7 @@ jobs: done - name: Verify pushed nydus image - if: ${{ contains(inputs.image_formats, 'nydus') }} + if: ${{ inputs.verify_image && contains(inputs.image_formats, 'nydus') }} env: IMAGE_TAGS: ${{ steps.image-tag.outputs.image-tags }} EXPECTED_BUILD_COMMIT: ${{ steps.build-metadata.outputs.build-commit }} diff --git a/.github/workflows/release-docker-deepseek-v4-nightly.yml b/.github/workflows/release-docker-deepseek-v4-nightly.yml index df3b5c08c40b..7ff03b25765b 100644 --- a/.github/workflows/release-docker-deepseek-v4-nightly.yml +++ b/.github/workflows/release-docker-deepseek-v4-nightly.yml @@ -149,6 +149,11 @@ jobs: cuda_key: cu130 cuda_version: 13.0.1 cuda_suffix: cu130 + # deepseek_v4 source lacks the ep_main runtime smoke tests baked into the + # shared verify step (sglang.srt.runtime_context / arg_groups.arg_utils and + # scripts/eic_integration_check.py exist only on ep_main). Provenance + # labels/env are still baked into the image; skip the in-container verify. + verify_image: false publish_default_cuda_alias: false variant_suffix: deepseek-v4 tag_mode: ${{ github.event_name == 'schedule' && 'nightly' || 'manual' }} From 2e13118775f6bd2637bd2088f21b0d8f640d6e45 Mon Sep 17 00:00:00 2001 From: Hank Han Date: Tue, 18 Aug 2026 14:37:23 +0800 Subject: [PATCH 24/47] ci: use docker manifest inspect instead of skopeo for backfill workflow skopeo does not share docker's credential store, so `skopeo inspect` against a private registry fails even after `docker login`. Switch to `docker manifest inspect` which uses the docker daemon's auth. Co-authored-by: TRAE CLI --- .github/workflows/backfill-image-formats.yml | 11 ++++++----- 1 file changed, 6 insertions(+), 5 deletions(-) diff --git a/.github/workflows/backfill-image-formats.yml b/.github/workflows/backfill-image-formats.yml index 23a3603a7667..39f8802546cd 100644 --- a/.github/workflows/backfill-image-formats.yml +++ b/.github/workflows/backfill-image-formats.yml @@ -88,17 +88,18 @@ jobs: echo "" echo "=== $(date '+%H:%M:%S') ${image_ref} ===" - # Verify OCI base exists - if ! skopeo inspect "docker://${image_ref}" >/dev/null 2>&1; then + # Verify OCI base exists. docker manifest inspect uses docker's + # credential store (docker login) and does not require skopeo. + if ! docker manifest inspect "${image_ref}" >/dev/null 2>&1; then echo " ERROR: base image not found: ${image_ref}" return 1 fi local oci_digest - oci_digest=$(skopeo inspect "docker://${image_ref}" | jq -r '.Digest') + oci_digest=$(docker manifest inspect "${image_ref}" | jq -r '.config.digest // empty') echo " OCI digest: ${oci_digest}" # zstd: pull then re-export with zstd compression (same as CI zstd step) - if skopeo inspect "docker://${image_ref}-zstd" >/dev/null 2>&1; then + if docker manifest inspect "${image_ref}-zstd" >/dev/null 2>&1; then echo " [zstd] already exists" else echo " [zstd] backfilling..." @@ -115,7 +116,7 @@ jobs: fi # nydus: nydusify convert (same as CI nydus step) - if skopeo inspect "docker://${image_ref}-nydus" >/dev/null 2>&1; then + if docker manifest inspect "${image_ref}-nydus" >/dev/null 2>&1; then echo " [nydus] already exists" else echo " [nydus] backfilling..." From c435f334b03c2553d639c86f054ce09c03039153 Mon Sep 17 00:00:00 2001 From: Hank Han Date: Thu, 20 Aug 2026 02:21:42 +0800 Subject: [PATCH 25/47] ci: add rename step to backfill workflow for historical image migration Add source_image_ref input to allow renaming old-format images to standard daily-build tags before backfilling zstd/nydus formats. Co-authored-by: TRAE CLI --- .github/workflows/backfill-image-formats.yml | 56 +++++++++++++++----- 1 file changed, 42 insertions(+), 14 deletions(-) diff --git a/.github/workflows/backfill-image-formats.yml b/.github/workflows/backfill-image-formats.yml index 39f8802546cd..aaa202ff0302 100644 --- a/.github/workflows/backfill-image-formats.yml +++ b/.github/workflows/backfill-image-formats.yml @@ -5,11 +5,19 @@ on: inputs: image_refs: description: | - 镜像引用列表,每行一个,不含格式后缀。 + 目标镜像引用列表,每行一个,不含格式后缀。 格式: iaas-gpu-cn-beijing.cr.volces.com/serving/: 示例: iaas-gpu-cn-beijing.cr.volces.com/serving/sglang:v0.5.16.iaas.202607280000-kimi-k3-cu130 required: true type: string + source_image_ref: + description: | + 可选:源镜像引用,用于将旧命名格式的镜像重命名为新 tag(docker pull + tag + push)。 + 留空则跳过重命名,直接对 image_refs 中的镜像补全格式。 + 格式: iaas-gpu-cn-beijing.cr.volces.com/serving/: + required: false + type: string + default: "" sync_to_customer: description: "补全格式后触发 Bits 流水线同步到客户仓库" required: false @@ -67,6 +75,30 @@ jobs: nydusify --version nydus-image --version + - name: Rename image (old tag -> new tag) + if: ${{ inputs.source_image_ref != '' }} + env: + TARGET_IMAGE_REF: ${{ inputs.image_refs }} + SOURCE_IMAGE_REF: ${{ inputs.source_image_ref }} + run: | + set -euo pipefail + + TARGET="$(echo "${TARGET_IMAGE_REF}" | head -1 | tr -d '\r\n' | sed 's/^[[:space:]]*//;s/[[:space:]]*$//')" + + echo "=== $(date '+%H:%M:%S') Rename ===" + echo " Source: ${SOURCE_IMAGE_REF}" + echo " Target: ${TARGET}" + + if docker manifest inspect "${TARGET}" >/dev/null 2>&1; then + echo " Target already exists, skipping rename" + exit 0 + fi + + docker pull "${SOURCE_IMAGE_REF}" + docker tag "${SOURCE_IMAGE_REF}" "${TARGET}" + docker push "${TARGET}" + echo " Renamed: ${TARGET}" + - name: Backfill image formats env: IMAGE_REFS: ${{ inputs.image_refs }} @@ -88,17 +120,13 @@ jobs: echo "" echo "=== $(date '+%H:%M:%S') ${image_ref} ===" - # Verify OCI base exists. docker manifest inspect uses docker's - # credential store (docker login) and does not require skopeo. if ! docker manifest inspect "${image_ref}" >/dev/null 2>&1; then echo " ERROR: base image not found: ${image_ref}" return 1 fi - local oci_digest - oci_digest=$(docker manifest inspect "${image_ref}" | jq -r '.config.digest // empty') - echo " OCI digest: ${oci_digest}" + echo " OCI base verified" - # zstd: pull then re-export with zstd compression (same as CI zstd step) + # zstd if docker manifest inspect "${image_ref}-zstd" >/dev/null 2>&1; then echo " [zstd] already exists" else @@ -115,7 +143,7 @@ jobs: echo " [zstd] published ${image_ref}-zstd" fi - # nydus: nydusify convert (same as CI nydus step) + # nydus if docker manifest inspect "${image_ref}-nydus" >/dev/null 2>&1; then echo " [nydus] already exists" else @@ -150,13 +178,13 @@ jobs: BITS_SPACE_ID="779378790402" echo "Images to sync:" - for image_ref in ${IMAGE_REFS}; do - image_ref="${image_ref//[$'\r\n']/}" - [[ -z "$image_ref" ]] && continue - framework="${image_ref##*/}"; framework="${framework%%:*}" - tag="${image_ref##*:}" + while IFS= read -r line; do + line="${line//[$'\r\n']/}" + [[ -z "$line" || "$line" == \#* ]] && continue + framework="${line##*/}"; framework="${framework%%:*}" + tag="${line##*:}" echo " framework=${framework} tag=${tag}" - done + done <<< "${IMAGE_REFS}" echo "" echo "Bits 流水线: https://bits.bytedance.net/devops/${BITS_SPACE_ID}/pipeline/detail/${BITS_PIPELINE_ID}" From bfa0d96bd10b8b1c81e0e92bbb961a7c4509bc13 Mon Sep 17 00:00:00 2001 From: Qi Sun <65223714+shiyu7@users.noreply.github.com> Date: Fri, 21 Aug 2026 16:01:38 +0800 Subject: [PATCH 26/47] [DSV4][DSpark] Support PP/CP/DP/DeepEP (#727) Co-authored-by: sunqi.7 --- .../sglang/srt/arg_groups/speculative_hook.py | 28 +- python/sglang/srt/disaggregation/decode.py | 12 +- .../srt/disaggregation/mooncake/conn.py | 45 +- python/sglang/srt/disaggregation/nixl/conn.py | 44 +- python/sglang/srt/disaggregation/prefill.py | 106 ++++- python/sglang/srt/disaggregation/utils.py | 91 ++++ .../sglang/srt/distributed/parallel_state.py | 2 +- .../srt/layers/moe/token_dispatcher/deepep.py | 6 +- .../sglang/srt/managers/scheduler_pp_mixin.py | 132 ++++-- python/sglang/srt/managers/utils.py | 11 +- .../srt/mem_cache/deepseek_v4_memory_pool.py | 16 + .../sglang/srt/mem_cache/kv_cache_builder.py | 2 + .../model_runner_components/layer_setup.py | 2 + .../srt/model_executor/runner/base_runner.py | 3 +- .../srt/model_executor/runner/eager_runner.py | 22 +- .../model_executor/runner_utils/buffers.py | 3 +- python/sglang/srt/models/deepseek_v4.py | 40 +- .../sglang/srt/models/deepseek_v4_dspark.py | 283 ++++++++++++- python/sglang/srt/models/dflash.py | 63 +++ python/sglang/srt/models/dspark.py | 38 ++ python/sglang/srt/models/kimi_k3.py | 66 ++- python/sglang/srt/models/kimi_linear.py | 33 +- python/sglang/srt/server_args.py | 16 +- .../dspark_components/dspark_config.py | 39 ++ .../dspark_components/dspark_kv_inject.py | 154 ++++++- .../dspark_components/dspark_worker_v2.py | 390 ++++++++++++++++-- .../cp/test_deepseek_v4_flash_fp4_b200_cp.py | 44 ++ .../test_disaggregation_wire.py | 21 + .../disaggregation/test_pp_pd_consensus.py | 296 +++++++++++++ .../unit/spec/test_dspark_pp_context.py | 387 +++++++++++++++++ 30 files changed, 2202 insertions(+), 193 deletions(-) create mode 100644 test/registered/unit/disaggregation/test_pp_pd_consensus.py create mode 100644 test/registered/unit/spec/test_dspark_pp_context.py diff --git a/python/sglang/srt/arg_groups/speculative_hook.py b/python/sglang/srt/arg_groups/speculative_hook.py index 427acdd3817f..8ffbeec2bf38 100644 --- a/python/sglang/srt/arg_groups/speculative_hook.py +++ b/python/sglang/srt/arg_groups/speculative_hook.py @@ -279,18 +279,20 @@ def _handle_dspark(server_args: ServerArgs) -> None: if not server_args.device.startswith("cuda"): raise ValueError("DSpark speculative decoding only supports CUDA device.") - if server_args.enable_dp_attention: + # dp_size==1 with dp_attention is a degenerate flag under DSV4 CP; skip DP-only checks. + if server_args.enable_dp_attention and server_args.dp_size > 1: if not server_args.enable_dp_lm_head: raise ValueError("DSpark with dp attention requires --enable-dp-lm-head.") - if server_args.moe_a2a_backend != "none": - raise ValueError( - "DSpark with dp attention only supports the built-in TP MoE " - f"(moe_a2a_backend='none'), got {server_args.moe_a2a_backend!r}." - ) - if server_args.attn_cp_size > 1: + supports_dspark_dp_moe = server_args.moe_a2a_backend == "none" or ( + server_args.moe_a2a_backend == "deepep" + and server_args.moe_runner_backend == "deep_gemm" + ) + if not supports_dspark_dp_moe: raise ValueError( - "DSpark with dp attention does not support context parallel " - f"(attn_cp_size={server_args.attn_cp_size})." + "DSpark with dp attention only supports moe_a2a_backend='none' " + "or moe_a2a_backend='deepep' with moe_runner_backend='deep_gemm'; " + f"got moe_a2a_backend={server_args.moe_a2a_backend!r}, " + f"moe_runner_backend={server_args.moe_runner_backend!r}." ) if ( server_args.speculative_moe_a2a_backend is not None @@ -302,9 +304,13 @@ def _handle_dspark(server_args: ServerArgs) -> None: f"(got {server_args.speculative_moe_a2a_backend!r})." ) - if server_args.pp_size != 1: + if server_args.pp_size != 1 and server_args.disaggregation_mode not in ( + "prefill", + "decode", + ): raise ValueError( - "Currently DSpark speculative decoding only supports pp_size == 1." + "Currently DSpark speculative decoding with pp_size > 1 is only " + "supported under PD disaggregation." ) if server_args.speculative_draft_model_path is None: diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 240e20deffb7..a1137da8cc33 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -52,6 +52,8 @@ _is_fake_transfer, get_dsv4_c128_state_indices, get_kv_class, + get_transfer_draft_kv_layer_ids, + get_transfer_kv_layer_ids, is_dsv4_c128_online_enabled, is_mla_backend, poll_and_all_reduce, @@ -434,6 +436,9 @@ def _init_kv_manager(self) -> CommonKVManager: kv_data_ptrs, kv_data_lens, kv_item_lens = ( transfer_kv_pool.get_contiguous_buf_infos() ) + kv_layer_ids = get_transfer_kv_layer_ids( + self.token_to_kv_pool, len(kv_data_ptrs) + ) kv_data_mem_kinds = ( ["DRAM"] * len(kv_data_ptrs) if self.scheduler.enable_hisparse @@ -449,6 +454,7 @@ def _init_kv_manager(self) -> CommonKVManager: kv_data_ptrs += device_kv_data_ptrs[c4_layer_num:] kv_data_lens += device_kv_data_lens[c4_layer_num:] kv_item_lens += device_kv_item_lens[c4_layer_num:] + kv_layer_ids = [] kv_data_mem_kinds += ["VRAM"] * len(device_kv_data_ptrs[c4_layer_num:]) if self.draft_token_to_kv_pool is not None: # We should also transfer draft model kv cache. The indices are @@ -456,6 +462,7 @@ def _init_kv_manager(self) -> CommonKVManager: draft_kv_data_ptrs, draft_kv_data_lens, draft_kv_item_lens = ( self.draft_token_to_kv_pool.get_contiguous_buf_infos() ) + kv_layer_ids += get_transfer_draft_kv_layer_ids(len(draft_kv_data_ptrs)) kv_data_ptrs += draft_kv_data_ptrs kv_data_lens += draft_kv_data_lens kv_item_lens += draft_kv_item_lens @@ -465,10 +472,7 @@ def _init_kv_manager(self) -> CommonKVManager: kv_args.kv_data_lens = kv_data_lens kv_args.kv_item_lens = kv_item_lens kv_args.kv_layer_ids = ( - self.token_to_kv_pool.get_kv_layer_ids() - if self.draft_token_to_kv_pool is None - and hasattr(self.token_to_kv_pool, "get_kv_layer_ids") - else [] + kv_layer_ids if len(kv_layer_ids) == len(kv_data_ptrs) else [] ) if self.transfer_backend == TransferBackend.NIXL: kv_args.kv_data_mem_kinds = kv_data_mem_kinds diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 1907e0ee69f8..c0328f0eb8d2 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -46,7 +46,10 @@ DisaggregationMode, build_transfer_entry_pairs, compute_mamba_state_slice_byte_blocks, + pack_state_types, resolve_dcp_dst_entry_indices, + resolve_state_component_dst_index, + unpack_state_types, ) from sglang.srt.distributed.parallel_state import get_mooncake_transfer_engine from sglang.srt.environ import envs @@ -140,6 +143,7 @@ class KVArgsRegisterInfo: dst_state_dim_per_tensor: List[List[int]] dst_kv_layer_ids: List[int] dst_state_layer_ids: List[List[int]] + dst_state_types: List[StateType] = dataclasses.field(default_factory=list) dst_dcp_size: int = 1 dst_dcp_rank: int = 0 requires_dcp_relayout: bool = False @@ -176,6 +180,7 @@ def from_zmq(cls, msg: List[bytes]): if len(msg) > 13 and msg[13] != b"" else [] ), + dst_state_types=unpack_state_types(msg[18]) if len(msg) > 18 else [], # msg[14:16] belong to the staging field below; DCP trails it. dst_dcp_size=( int(msg[16].decode("ascii")) if len(msg) > 16 and msg[16] != b"" else 1 @@ -1207,32 +1212,52 @@ def maybe_send_extra( src_state_layer_ids = ( src_state_layer_ids[i] if i < len(src_state_layer_ids) else [] ) + dst_component_index = i if target_rank_registration_info is not None: + dst_component_index = resolve_state_component_dst_index( + state_types, + target_rank_registration_info.dst_state_types, + i, + ) dst_data_ptrs = ( - target_rank_registration_info.dst_state_data_ptrs[i] - if i < len(target_rank_registration_info.dst_state_data_ptrs) + target_rank_registration_info.dst_state_data_ptrs[ + dst_component_index + ] + if dst_component_index + < len(target_rank_registration_info.dst_state_data_ptrs) else [] ) dst_item_lens = ( - target_rank_registration_info.dst_state_item_lens[i] - if i < len(target_rank_registration_info.dst_state_item_lens) + target_rank_registration_info.dst_state_item_lens[ + dst_component_index + ] + if dst_component_index + < len(target_rank_registration_info.dst_state_item_lens) else [] ) dst_dim_per_tensor = ( - target_rank_registration_info.dst_state_dim_per_tensor[i] - if i < len(target_rank_registration_info.dst_state_dim_per_tensor) + target_rank_registration_info.dst_state_dim_per_tensor[ + dst_component_index + ] + if dst_component_index + < len(target_rank_registration_info.dst_state_dim_per_tensor) else [] ) dst_state_layer_ids = ( - target_rank_registration_info.dst_state_layer_ids[i] - if i < len(target_rank_registration_info.dst_state_layer_ids) + target_rank_registration_info.dst_state_layer_ids[ + dst_component_index + ] + if dst_component_index + < len(target_rank_registration_info.dst_state_layer_ids) else [] ) else: dst_data_ptrs, dst_item_lens, dst_dim_per_tensor = [], [], [] dst_state_layer_ids = [] dst_indices = ( - req.dst_state_indices[i] if i < len(req.dst_state_indices) else [] + req.dst_state_indices[dst_component_index] + if dst_component_index < len(req.dst_state_indices) + else [] ) if st == StateType.MAMBA: @@ -2257,6 +2282,7 @@ def _register_kv_args(self) -> bool: packed_state_layer_ids = pack_int_lists( self.kv_mgr.kv_args.state_layer_ids, "I" ) + packed_state_types = pack_state_types(self.kv_mgr.kv_args.state_types) packed_kv_layer_ids = b"".join( struct.pack("I", layer_id) for layer_id in self.kv_mgr.kv_args.kv_layer_ids @@ -2309,6 +2335,7 @@ def _register_kv_args(self) -> bool: staging_total_size_str, dst_dcp_size, dst_dcp_rank, + packed_state_types, ] ) except zmq.ZMQError: diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index 86a5b290453a..0183de236ff5 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -41,7 +41,10 @@ DisaggregationMode, build_transfer_entry_pairs, compute_mamba_state_slice_byte_blocks, + pack_state_types, resolve_dcp_dst_entry_indices, + resolve_state_component_dst_index, + unpack_state_types, ) from sglang.srt.environ import envs from sglang.srt.runtime_context import get_schedule @@ -231,6 +234,7 @@ class KVArgsRegisterInfo: dst_state_item_lens: List[List[int]] = dataclasses.field(default_factory=list) dst_state_dim_per_tensor: List[List[int]] = dataclasses.field(default_factory=list) dst_state_layer_ids: List[List[int]] = dataclasses.field(default_factory=list) + dst_state_types: List[StateType] = dataclasses.field(default_factory=list) dst_homogeneous_mem_kind: Optional[str] = None kv_xfer_segments: Optional[List[_KVXferPreparedSegment]] = None # Keep last: optional, parsed from a variable-length tail of the ZMQ @@ -276,6 +280,7 @@ def from_zmq(cls, msg: List[bytes]): if len(msg) > 20 and msg[20] != b"" else [] ) + dst_state_types = unpack_state_types(msg[23]) if len(msg) > 23 else [] return cls( room=str(msg[0].decode("ascii")), @@ -304,6 +309,7 @@ def from_zmq(cls, msg: List[bytes]): dst_state_item_lens=dst_state_item_lens, dst_state_dim_per_tensor=dst_state_dim_per_tensor, dst_state_layer_ids=dst_state_layer_ids, + dst_state_types=dst_state_types, staging=StagingRegisterInfo.from_zmq_fields(msg, 14), ) @@ -1283,6 +1289,7 @@ def transfer_worker(self, queue: FastQueue, staging_buffer=None): dst_state_item_lens=dst_info.dst_state_item_lens, dst_state_dim_per_tensor=dst_info.dst_state_dim_per_tensor, dst_state_layer_ids=dst_info.dst_state_layer_ids, + dst_state_types=dst_info.dst_state_types, ) handles.extend( h for h in state_xfer_handles if h is not None @@ -2208,6 +2215,7 @@ def maybe_send_extra( dst_state_item_lens: List[List[int]] | None = None, dst_state_dim_per_tensor: List[List[int]] | None = None, dst_state_layer_ids: List[List[int]] | None = None, + dst_state_types: List[StateType] | None = None, ): """Send state per hybrid component, dispatching by state_type[i].""" state_types = getattr(self.kv_args, "state_types", []) or [] @@ -2226,9 +2234,15 @@ def maybe_send_extra( dst_state_item_lens = dst_state_item_lens or [] dst_state_dim_per_tensor = dst_state_dim_per_tensor or [] dst_state_layer_ids = dst_state_layer_ids or [] + dst_state_types = dst_state_types or [] handles = [] for i, st in enumerate(state_types): + dst_component_index = resolve_state_component_dst_index( + state_types, + dst_state_types, + i, + ) src_indices = ( prefill_state_indices[i] if i < len(prefill_state_indices) else None ) @@ -2250,13 +2264,31 @@ def maybe_send_extra( else [] ) src_lids = src_state_layer_ids[i] if i < len(src_state_layer_ids) else [] - dst_ptrs = dst_state_data_ptrs[i] if i < len(dst_state_data_ptrs) else [] - dst_indices = dst_state_indices[i] if i < len(dst_state_indices) else [] - dst_lens = dst_state_item_lens[i] if i < len(dst_state_item_lens) else [] + dst_ptrs = ( + dst_state_data_ptrs[dst_component_index] + if dst_component_index < len(dst_state_data_ptrs) + else [] + ) + dst_indices = ( + dst_state_indices[dst_component_index] + if dst_component_index < len(dst_state_indices) + else [] + ) + dst_lens = ( + dst_state_item_lens[dst_component_index] + if dst_component_index < len(dst_state_item_lens) + else [] + ) dst_dims = ( - dst_state_dim_per_tensor[i] if i < len(dst_state_dim_per_tensor) else [] + dst_state_dim_per_tensor[dst_component_index] + if dst_component_index < len(dst_state_dim_per_tensor) + else [] + ) + dst_lids = ( + dst_state_layer_ids[dst_component_index] + if dst_component_index < len(dst_state_layer_ids) + else [] ) - dst_lids = dst_state_layer_ids[i] if i < len(dst_state_layer_ids) else [] comp_notif = f"{notif}_{i}" if st == StateType.MAMBA: @@ -2952,6 +2984,7 @@ def _register_kv_args(self) -> bool: packed_state_layer_ids = pack_int_lists( self.kv_mgr.kv_args.state_layer_ids, "I" ) + packed_state_types = pack_state_types(self.kv_mgr.kv_args.state_types) # Include staging allocator metadata if available if ( @@ -3000,6 +3033,7 @@ def _register_kv_args(self) -> bool: packed_kv_layer_ids, str(self.kv_mgr.dcp_size).encode("ascii"), str(self.kv_mgr.dcp_rank).encode("ascii"), + packed_state_types, ] ) except zmq.ZMQError: diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 5642c64efaa4..8af6655e0fea 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -24,7 +24,7 @@ from array import array from collections import deque from http import HTTPStatus -from typing import TYPE_CHECKING, List, Optional +from typing import TYPE_CHECKING, List, Optional, Tuple import numpy as np import torch @@ -41,6 +41,8 @@ TransferBackend, get_dsv4_c128_state_indices, get_kv_class, + get_transfer_draft_kv_layer_ids, + get_transfer_kv_layer_ids, is_aborted, is_dsv4_c128_online_enabled, is_mla_backend, @@ -193,8 +195,8 @@ def _init_kv_manager(self) -> CommonKVManager: layer_shard_rank = getattr(self.token_to_kv_pool, "layer_shard_rank", None) layer_shard_size = getattr(self.token_to_kv_pool, "layer_shard_size", 1) transfer_draft_cache = ( - not layer_shard_enabled or layer_shard_rank == layer_shard_size - 1 - ) + self.pp_size <= 1 or self.pp_rank == self.pp_size - 1 + ) and (not layer_shard_enabled or layer_shard_rank == layer_shard_size - 1) kv_args.prefill_start_layer = ( getattr( self.token_to_kv_pool, @@ -208,6 +210,9 @@ def _init_kv_manager(self) -> CommonKVManager: kv_data_ptrs, kv_data_lens, kv_item_lens = ( self.token_to_kv_pool.get_contiguous_buf_infos() ) + kv_layer_ids = get_transfer_kv_layer_ids( + self.token_to_kv_pool, len(kv_data_ptrs) + ) kv_args.prefill_end_layer = ( kv_args.prefill_start_layer + len(kv_data_ptrs) if layer_shard_enabled @@ -220,6 +225,7 @@ def _init_kv_manager(self) -> CommonKVManager: draft_kv_data_ptrs, draft_kv_data_lens, draft_kv_item_lens = ( self.draft_token_to_kv_pool.get_contiguous_buf_infos() ) + kv_layer_ids += get_transfer_draft_kv_layer_ids(len(draft_kv_data_ptrs)) kv_data_ptrs += draft_kv_data_ptrs kv_data_lens += draft_kv_data_lens kv_item_lens += draft_kv_item_lens @@ -228,10 +234,7 @@ def _init_kv_manager(self) -> CommonKVManager: kv_args.kv_data_lens = kv_data_lens kv_args.kv_item_lens = kv_item_lens kv_args.kv_layer_ids = ( - self.token_to_kv_pool.get_kv_layer_ids() - if self.draft_token_to_kv_pool is None - and hasattr(self.token_to_kv_pool, "get_kv_layer_ids") - else [] + kv_layer_ids if len(kv_layer_ids) == len(kv_data_ptrs) else [] ) if not self.is_mla_backend: kv_args.kv_head_num = self.token_to_kv_pool.head_num @@ -460,6 +463,43 @@ def pop_bootstrapped( else: return bootstrapped_reqs, failed_reqs + def get_ready_bootstrapped_rids_for_pp(self) -> Tuple[List[str], List[str]]: + """Return ordered PP candidates without reserving local resources.""" + good_rids: List[str] = [] + failed_rids: List[str] = [] + if len(self.queue) == 0: + return good_rids, failed_rids + + polls = poll_and_all_reduce_attn_cp_tp_group( + [req.disagg_kv_sender for req in self.queue], + self.scheduler.attn_cp_cpu_group, + self.scheduler.attn_tp_cpu_group, + ) + metadata_credits = self.req_to_metadata_buffer_idx_allocator.available_size() + admission_blocked = False + + for req, poll in zip(self.queue, polls): + if poll == KVPoll.Failed: + failed_rids.append(req.rid) + elif poll == KVPoll.WaitingForInput: + if admission_blocked: + continue + metadata_cost = 1 if req.metadata_buffer_index < 0 else 0 + if metadata_cost > metadata_credits: + admission_blocked = True + continue + metadata_credits -= metadata_cost + good_rids.append(req.rid) + elif poll == KVPoll.Bootstrapping: + continue + else: + raise RuntimeError( + f"Unexpected poll state {poll} for req {req.rid} " + "in get_ready_bootstrapped_rids_for_pp" + ) + + return good_rids, failed_rids + def release_memory_occupation(self): self.queue.clear() if hasattr(self.kv_manager, "deregister_buffer_to_engine"): @@ -815,16 +855,19 @@ def advance_logprob_pt(i: int, req: Req) -> None: ) def process_disagg_prefill_inflight_queue( - self: Scheduler, rids_to_check: Optional[List[str]] = None + self: Scheduler, + transfer_status: Optional[Tuple[List[str], List[str]]] = None, ) -> List[Req]: """ Poll the requests in the middle of transfer. If done, return the request. - rids_to_check: For PP, on rank > 0, check the rids from the previous rank has consensus with the current rank. + transfer_status: For PP, the consensus success and failure request ids. """ if len(self.disagg_prefill_inflight_queue) == 0: return [] done_reqs = [] + success_rids = set(transfer_status[0]) if transfer_status is not None else set() + failed_rids = set(transfer_status[1]) if transfer_status is not None else set() polls = poll_and_all_reduce_attn_cp_tp_group( [req.disagg_kv_sender for req in self.disagg_prefill_inflight_queue], @@ -835,8 +878,32 @@ def process_disagg_prefill_inflight_queue( undone_reqs: List[Req] = [] # Check .poll() for the reqs in disagg_prefill_inflight_queue. If Success, respond to the client and remove it from the queue for req, poll in zip(self.disagg_prefill_inflight_queue, polls): - if rids_to_check is not None: - if req.rid not in rids_to_check: + if transfer_status is not None: + consensus_failed = req.rid in failed_rids + failure_pending = isinstance(req.finished_reason, FINISH_ABORT) + if consensus_failed or failure_pending: + if consensus_failed and not failure_pending: + prepare_abort( + req, + ( + "Prefill transfer failed on another PP rank; " + "waiting for the local transfer to stop" + ), + status_code=HTTPStatus.INTERNAL_SERVER_ERROR, + ) + local_transfer_stopped = req.pending_bootstrap or poll in ( + KVPoll.Success, + KVPoll.Failed, + ) + if not local_transfer_stopped: + undone_reqs.append(req) + continue + + self.handle_inflight_transfer_failure(req) + done_reqs.append(req) + continue + + if req.rid not in success_rids: undone_reqs.append(req) continue @@ -950,9 +1017,11 @@ def handle_inflight_transfer_failure( self.metrics_collector.increment_transfer_failed_reqs() return exc - def get_transferred_rids(self: Scheduler) -> List[str]: + def get_transferred_rids( + self: Scheduler, + ) -> Tuple[List[str], List[str]]: """ - Used by PP, get the transferred rids but **do not pop** + Used by PP, inspect terminal transfer states without popping requests. """ polls = poll_and_all_reduce_attn_cp_tp_group( [req.disagg_kv_sender for req in self.disagg_prefill_inflight_queue], @@ -960,13 +1029,16 @@ def get_transferred_rids(self: Scheduler) -> List[str]: self.attn_tp_cpu_group, ) - transferred_rids: List[str] = [] + success_rids: List[str] = [] + failed_rids: List[str] = [] for req, poll in zip(self.disagg_prefill_inflight_queue, polls): - if poll == KVPoll.Success or poll == KVPoll.Failed: - transferred_rids.append(req.rid) + if poll == KVPoll.Success: + success_rids.append(req.rid) + elif poll == KVPoll.Failed: + failed_rids.append(req.rid) - return transferred_rids + return success_rids, failed_rids def handle_bootstrap_failure(self: Scheduler, req: Req) -> None: error_message = ( diff --git a/python/sglang/srt/disaggregation/utils.py b/python/sglang/srt/disaggregation/utils.py index 237f293ec503..dbba648afc2b 100644 --- a/python/sglang/srt/disaggregation/utils.py +++ b/python/sglang/srt/disaggregation/utils.py @@ -959,6 +959,97 @@ def resolve_dcp_dst_entry_indices( ] +_DRAFT_KV_LAYER_ID_BASE = 1_000_000 + + +def get_transfer_kv_layer_ids(kv_pool, num_entries: int) -> List[int]: + """Return global layer ids aligned with ``kv_pool.get_contiguous_buf_infos``. + + Pools with a custom sparse/MLA layout expose ``get_kv_layer_ids`` directly. + Plain MHA-like draft pools usually expose only ``start_layer``/``end_layer``; + infer either one entry per layer or K/V tensor-major entries. + """ + if kv_pool is None or num_entries <= 0: + return [] + + if hasattr(kv_pool, "get_kv_layer_ids"): + layer_ids = list(kv_pool.get_kv_layer_ids()) + if len(layer_ids) == num_entries: + return layer_ids + + start_layer = int(getattr(kv_pool, "start_layer", 0) or 0) + end_layer = getattr(kv_pool, "end_layer", None) + if end_layer is not None: + layer_ids = list(range(start_layer, int(end_layer))) + if len(layer_ids) == num_entries: + return layer_ids + if len(layer_ids) * 2 == num_entries: + return layer_ids * 2 + + return [] + + +def get_transfer_draft_kv_layer_ids(num_entries: int) -> List[int]: + """Return layer-id metadata for draft KV entries. + + Draft KV is not target-model KV and should not share the target layer-id + namespace, especially when PP prefill sends target KV by layer subset but a + single rank also sends replicated draft KV. + """ + if num_entries <= 0: + return [] + return [_DRAFT_KV_LAYER_ID_BASE + i for i in range(num_entries)] + + +def pack_state_types(state_types) -> bytes: + return ",".join( + state_type.value if hasattr(state_type, "value") else str(state_type) + for state_type in (state_types or []) + ).encode("ascii") + + +def unpack_state_types(data: bytes): + from sglang.srt.disaggregation.base.conn import StateType + + if not data: + return [] + return [StateType(value) for value in data.decode("ascii").split(",") if value] + + +def resolve_state_component_dst_index(src_state_types, dst_state_types, src_index: int): + """Map a source state component to the matching destination component. + + Older registrations did not carry state_types; keep positional behavior in + that case. When available, match by StateType occurrence so optional target + components (for example C128_STATE) do not shift draft state components. + """ + if not dst_state_types: + return src_index + if not src_state_types: + raise RuntimeError( + "Destination state_types are present but source state_types are empty." + ) + if src_index >= len(src_state_types): + raise RuntimeError( + f"Source state component index {src_index} exceeds " + f"state_types length {len(src_state_types)}." + ) + state_type = src_state_types[src_index] + occurrence = sum( + 1 for item in src_state_types[: src_index + 1] if item == state_type + ) + seen = 0 + for dst_index, dst_state_type in enumerate(dst_state_types): + if dst_state_type == state_type: + seen += 1 + if seen == occurrence: + return dst_index + raise RuntimeError( + f"Decode peer is missing state component {state_type!s} " + f"occurrence {occurrence}." + ) + + def append_state_component( kv_args: KVArgs, state_type: StateType, diff --git a/python/sglang/srt/distributed/parallel_state.py b/python/sglang/srt/distributed/parallel_state.py index 6fb820082c94..74dc252b8ec7 100644 --- a/python/sglang/srt/distributed/parallel_state.py +++ b/python/sglang/srt/distributed/parallel_state.py @@ -1277,7 +1277,7 @@ def all_gather( return torch.ops.sgl_kernel.shm_allgather(input_, dim) else: torch.distributed.all_gather_into_tensor( - output_tensor, input_, group=self.device_group + output_tensor, input_, group=self.cpu_group ) else: self.all_gather_into_tensor(output_tensor, input_) diff --git a/python/sglang/srt/layers/moe/token_dispatcher/deepep.py b/python/sglang/srt/layers/moe/token_dispatcher/deepep.py index f33c74f77edb..2c0cef024710 100644 --- a/python/sglang/srt/layers/moe/token_dispatcher/deepep.py +++ b/python/sglang/srt/layers/moe/token_dispatcher/deepep.py @@ -1,5 +1,6 @@ from __future__ import annotations +import inspect import logging from contextlib import nullcontext from dataclasses import dataclass @@ -278,7 +279,10 @@ def get_deepep_buffer( # auto-enables fabric in C++ when supported, so we skip it: # https://github.com/fzyzcjy/DeepEP/blob/814e508537c6ffc775d59f6f1b9ba43f3a65968c/csrc/deep_ep.cpp#L52 is_cu12 = get_cuda_version()[0] == 12 - if not is_cu12 and use_mnnvl_fabric: + supports_use_fabric = ( + "use_fabric" in inspect.signature(Buffer.__init__).parameters + ) + if not is_cu12 and use_mnnvl_fabric and supports_use_fabric: buffer_kwargs["use_fabric"] = True state.buffer = Buffer(group, num_nvl_bytes, num_rdma_bytes, **buffer_kwargs) diff --git a/python/sglang/srt/managers/scheduler_pp_mixin.py b/python/sglang/srt/managers/scheduler_pp_mixin.py index 715e20ebf34f..4b12fffa82c3 100644 --- a/python/sglang/srt/managers/scheduler_pp_mixin.py +++ b/python/sglang/srt/managers/scheduler_pp_mixin.py @@ -6,7 +6,7 @@ from array import array from collections import defaultdict, deque from dataclasses import dataclass -from typing import TYPE_CHECKING, Dict, List, Optional, Tuple +from typing import TYPE_CHECKING, Dict, List, Optional, Tuple, Union import numpy as np import torch @@ -46,6 +46,9 @@ if TYPE_CHECKING: from sglang.srt.managers.scheduler import Scheduler +PPTransferStatus = Tuple[List[str], List[str]] +PPReleasePayload = Union[List[str], PPTransferStatus] + def _pp_can_skip_output_comm(batch: ScheduleBatch) -> bool: """Check if output send/recv can be skipped for this batch.""" @@ -59,6 +62,35 @@ def _pp_can_skip_output_comm(batch: ScheduleBatch) -> bool: ) +def _pp_ordered_intersection(left: List[str], right: List[str]) -> List[str]: + right_set = set(right) + return [rid for rid in left if rid in right_set] + + +def _pp_ordered_union(left: List[str], right: List[str]) -> List[str]: + seen = set(left) + merged = list(left) + for rid in right: + if rid not in seen: + seen.add(rid) + merged.append(rid) + return merged + + +def _pp_merge_transfer_status( + previous: PPTransferStatus, + current: PPTransferStatus, +) -> PPTransferStatus: + previous_success, previous_failed = previous + current_success, current_failed = current + failed = _pp_ordered_union(previous_failed, current_failed) + success = _pp_ordered_intersection(previous_success, current_success) + if failed: + failed_set = set(failed) + success = [rid for rid in success if rid not in failed_set] + return success, failed + + @dataclass class PPBatchMetadata: can_run_cuda_graph: bool @@ -218,8 +250,8 @@ def event_loop_pp_disagg_prefill(self: Scheduler): bmbs = [None] * self.pp_loop_size tmbs = [None] * self.pp_loop_size consensus_bootstrapped_rids: Optional[List[str]] = None - transferred_rids: List[str] = [] - release_rids: Optional[List[str]] = None + transferred_rids: PPTransferStatus = ([], []) + release_rids: Optional[PPTransferStatus] = None send_bootstrapped_work = [] send_transfer_work = [] send_consensus_bootstrapped_work = [] @@ -817,18 +849,20 @@ def process_bootstrapped_queue( ) ) self.waiting_queue.extend(good_reqs) - return [[req.rid for req in good_reqs], [req.rid for req in failed_reqs]] + return [ + [req.rid for req in good_reqs], + _pp_ordered_union( + bad_consensus_bootstrapped_rids, + [req.rid for req in failed_reqs], + ), + ] return None def _pp_pd_get_bootstrapped_ids(self: Scheduler): # communicate pre-consensus bootstrapp reqs if self.pp_group.is_first_rank: - # First rank, pop the bootstrap reqs from the bootstrap queue - good_bootstrapped_rids, bad_bootstrapped_rids = self.get_rids( - self.disagg_prefill_bootstrap_queue.queue, - True, - [KVPoll.WaitingForInput], - [KVPoll.Failed], + good_bootstrapped_rids, bad_bootstrapped_rids = ( + self.disagg_prefill_bootstrap_queue.get_ready_bootstrapped_rids_for_pp() ) else: # Other ranks, receive the bootstrap reqs info from the previous rank and ensure the consensus @@ -836,17 +870,16 @@ def _pp_pd_get_bootstrapped_ids(self: Scheduler): prev_good_bootstrapped_rids, prev_bad_bootstrapped_rids = ( prev_bootstrapped_rids ) - curr_good_bootstrapped_rids, curr_bad_bootstrapped_rids = self.get_rids( - self.disagg_prefill_bootstrap_queue.queue, - True, - [KVPoll.WaitingForInput], - [KVPoll.Failed], + curr_good_bootstrapped_rids, curr_bad_bootstrapped_rids = ( + self.disagg_prefill_bootstrap_queue.get_ready_bootstrapped_rids_for_pp() ) - good_bootstrapped_rids = list( - set(prev_good_bootstrapped_rids) & set(curr_good_bootstrapped_rids) + good_bootstrapped_rids = _pp_ordered_intersection( + prev_good_bootstrapped_rids, + curr_good_bootstrapped_rids, ) - bad_bootstrapped_rids = list( - set(prev_bad_bootstrapped_rids) | set(curr_bad_bootstrapped_rids) + bad_bootstrapped_rids = _pp_ordered_union( + prev_bad_bootstrapped_rids, + curr_bad_bootstrapped_rids, ) # Route locally-aborted reqs through the bad-union consensus so every PP # rank flushes them in the same consensus round, regardless of when the @@ -862,28 +895,21 @@ def _pp_pd_get_bootstrapped_ids(self: Scheduler): ) return [good_bootstrapped_rids, bad_bootstrapped_rids] - def _pp_pd_get_prefill_transferred_ids(self: Scheduler): + def _pp_pd_get_prefill_transferred_ids( + self: Scheduler, + ) -> PPTransferStatus: # get the current stage transfer success + current_status = self.get_transferred_rids() if self.pp_group.is_first_rank: - transferred_rids = self.get_rids( - self.disagg_prefill_inflight_queue, - True, - [KVPoll.Success, KVPoll.Failed], - ) + transferred_rids = current_status # if other ranks, do intersection with the previous rank's transferred rids else: # 2 (Release): Receive the transferred rids from the previous rank # 1. recv previous stage's transferred reqs info - prev_transferred_rids = self._pp_recv_pyobj_from_prev_stage() - # 2. get the current stage's transferred reqs info - curr_transferred_rids = self.get_rids( - self.disagg_prefill_inflight_queue, - True, - [KVPoll.Success, KVPoll.Failed], - ) - # 3. new consensus rids = intersection(previous consensus rids, transfer finished rids) - transferred_rids = list( - set(prev_transferred_rids) & set(curr_transferred_rids) + previous_status = self._pp_recv_pyobj_from_prev_stage() + transferred_rids = _pp_merge_transfer_status( + previous=previous_status, + current=current_status, ) return transferred_rids @@ -912,10 +938,10 @@ def _pp_pd_send_consensus_bootstrapped_ids( def _pp_pd_send_consensus_release_ids( self: Scheduler, - tmbs: List[List[str]], + tmbs: List[Optional[PPReleasePayload]], next_first_rank_mb_id: int, - release_rids: List[str], - transferred_rids: List[str], + release_rids: Optional[PPReleasePayload], + transferred_rids: PPReleasePayload, ): send_release_work = [] if self.pp_group.is_last_rank: @@ -1148,13 +1174,30 @@ def _pp_prep_batch_result( extend_input_len_per_req, extend_logprob_start_len_per_req, ) = get_logprob_from_pp_outputs(pp_outputs) - next_token_ids = pp_outputs["next_token_ids"].to(torch.int64) + # PP outputs may be on CPU after result processing, while draft state is + # filtered with device-resident batch indices. + next_token_ids = pp_outputs["next_token_ids"].to( + device=batch.device, + dtype=torch.int64, + non_blocking=True, + ) # PP rank 0 also relays into output_tokens_buf so the next iter's # resolve_forward_inputs finds these tokens for the decode portion # of mixed-chunk batches (which gather via mix_running_indices). self.future_map.stash( batch.req_pool_indices, RelayPayload(bonus_tokens=next_token_ids) ) + next_draft_input = None + if batch.spec_algorithm.is_dspark(): + from sglang.srt.speculative.dspark_components.dspark_draft import ( + make_next_draft_input, + ) + + next_draft_input = make_next_draft_input( + bonus_tokens=next_token_ids, + new_seq_lens=batch.seq_lens, + ) + batch.spec_info = next_draft_input batch.input_ids = None output_result = GenerationBatchResult( logits_output=logits_output, @@ -1163,6 +1206,7 @@ def _pp_prep_batch_result( extend_input_len_per_req=extend_input_len_per_req, extend_logprob_start_len_per_req=extend_logprob_start_len_per_req, can_run_cuda_graph=mb_metadata.can_run_cuda_graph, + next_draft_input=next_draft_input, ) return output_result @@ -1291,7 +1335,15 @@ def _pp_launch_batch( "set_run_batch_cpu_start_time", trace_only=True, ) - result = self.run_batch(cur_batch, pp_proxy_tensors) + if cur_batch.spec_algorithm.is_dspark(): + self.model_worker.set_pp_proxy_tensors_for_next_forward( + pp_proxy_tensors + ) + try: + result = self.run_batch(cur_batch, pp_proxy_tensors) + finally: + if cur_batch.spec_algorithm.is_dspark(): + self.model_worker.set_pp_proxy_tensors_for_next_forward(None) set_time_batch( cur_batch.reqs, "set_run_batch_cpu_end_time", diff --git a/python/sglang/srt/managers/utils.py b/python/sglang/srt/managers/utils.py index fe883c264d28..2b4dfca70167 100644 --- a/python/sglang/srt/managers/utils.py +++ b/python/sglang/srt/managers/utils.py @@ -120,7 +120,7 @@ def copy_to_cpu(self, return_logprob: bool, return_hidden_states: bool = True): Only the tensors which are needed for processing results are copied, e.g., next_token_ids, logits outputs """ - if return_logprob: + if self.logits_output is not None and return_logprob: if self.logits_output.next_token_logprobs is not None: self.logits_output.next_token_logprobs = _async_d2h( self.logits_output.next_token_logprobs @@ -144,11 +144,16 @@ def copy_to_cpu(self, return_logprob: bool, return_hidden_states: bool = True): _async_d2h(v) if torch.is_tensor(v) else v for v in self.logits_output.next_token_token_ids_logprobs_val ] - if return_hidden_states and self.logits_output.hidden_states is not None: + if ( + self.logits_output is not None + and return_hidden_states + and self.logits_output.hidden_states is not None + ): self.logits_output.hidden_states = _async_d2h( self.logits_output.hidden_states ) - self.next_token_ids = _async_d2h(self.next_token_ids) + if self.next_token_ids is not None: + self.next_token_ids = _async_d2h(self.next_token_ids) if self.accept_lens is not None: self.accept_lens = _async_d2h(self.accept_lens) diff --git a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py index 062e60bc9c05..75ad574f7406 100644 --- a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py +++ b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py @@ -666,6 +666,22 @@ def translate_loc_from_full_to_swa(self, kv_indices: torch.Tensor): assert self.full_to_swa_index_mapping is not None return self.full_to_swa_index_mapping[kv_indices] + def get_kv_layer_ids(self) -> List[int]: + """Global layer IDs aligned with the compressed KV entry layout.""" + stage_ratios = self.compression_ratios[self._stage_start : self._stage_end] + c4_layer_ids = [ + self._stage_start + local_layer_id + for local_layer_id, ratio in enumerate(stage_ratios) + if ratio == 4 + ] + c128_layer_ids = [ + self._stage_start + local_layer_id + for local_layer_id, ratio in enumerate(stage_ratios) + if ratio == 128 + ] + # get_contiguous_buf_infos orders entries as C4, C4 indexer, C128. + return c4_layer_ids + c4_layer_ids + c128_layer_ids + def get_contiguous_buf_infos(self) -> Tuple[List[int], List[int], List[int]]: data_ptrs: List[int] = [] data_lens: List[int] = [] diff --git a/python/sglang/srt/mem_cache/kv_cache_builder.py b/python/sglang/srt/mem_cache/kv_cache_builder.py index 50f797101046..5cf652e50870 100644 --- a/python/sglang/srt/mem_cache/kv_cache_builder.py +++ b/python/sglang/srt/mem_cache/kv_cache_builder.py @@ -61,6 +61,8 @@ def get_draft_kv_pool( or None when no draft KV pool is available.""" if draft_worker is None or spec_algorithm.is_ngram(): return None + if spec_algorithm.is_dspark() and draft_worker.is_lifecycle_only_pp_prefill_rank: + return None # V2 workers nest the draft runner under `.draft_worker`. if server_args.enable_multi_layer_eagle: diff --git a/python/sglang/srt/model_executor/model_runner_components/layer_setup.py b/python/sglang/srt/model_executor/model_runner_components/layer_setup.py index 83a94a63b993..2323ce1d20e2 100644 --- a/python/sglang/srt/model_executor/model_runner_components/layer_setup.py +++ b/python/sglang/srt/model_executor/model_runner_components/layer_setup.py @@ -197,9 +197,11 @@ def _assert_pp_mtp_compat( num_effective_layers: int, model_num_layers: int, ) -> None: + # DSPARK uses a separate draft worker even when the target config bundles MTP layers. assert ( (not model_has_mtp_layers) or (spec_algorithm.is_none()) + or spec_algorithm.is_dspark() or ( (not spec_algorithm.is_none()) and (num_effective_layers == model_num_layers) diff --git a/python/sglang/srt/model_executor/runner/base_runner.py b/python/sglang/srt/model_executor/runner/base_runner.py index 656b070d562e..a5f63141cae7 100644 --- a/python/sglang/srt/model_executor/runner/base_runner.py +++ b/python/sglang/srt/model_executor/runner/base_runner.py @@ -112,8 +112,9 @@ def _allocate_decode_buffers( # mHC (e.g. DSV4) flattens residual into hidden_states (size = hc_hidden_size). is_mhc = hc_hidden_size is not None hs = hc_hidden_size if is_mhc else hidden_size + # Target verify expands each request into num_tokens_per_req rows. pp_proxy_tensors = { - "hidden_states": torch.zeros((max_bs, hs), dtype=dtype), + "hidden_states": torch.zeros((max_num_token, hs), dtype=dtype), } if not is_mhc: # Only Kimi K3 supplies num_blocks: its PP bank is token-major diff --git a/python/sglang/srt/model_executor/runner/eager_runner.py b/python/sglang/srt/model_executor/runner/eager_runner.py index 3c716c840cbd..f80343af32f3 100644 --- a/python/sglang/srt/model_executor/runner/eager_runner.py +++ b/python/sglang/srt/model_executor/runner/eager_runner.py @@ -371,15 +371,31 @@ def _execute_extend_cp_v2( else hidden_states ) - hidden_states = cp_gather_after_forward( - hidden_states, forward_batch, torch.cuda.current_stream() - ) + stream = torch.cuda.current_stream() + hidden_states = cp_gather_after_forward(hidden_states, forward_batch, stream) + # DSpark aux tensors ride the same CP token split; gather them the same way. + if aux_hidden_states is not None: + if isinstance(aux_hidden_states, torch.Tensor): + aux_hidden_states = cp_gather_after_forward( + aux_hidden_states, forward_batch, stream + ) + else: + aux_hidden_states = [ + cp_gather_after_forward(aux, forward_batch, stream) + for aux in aux_hidden_states + ] + logits_kwargs = {} + # DSV4 returns (hidden_states, hidden_states_before_norm) from its model body. + if isinstance(hidden_states, tuple): + hidden_states, hidden_states_before_norm = hidden_states + logits_kwargs["hidden_states_before_norm"] = hidden_states_before_norm return model.logits_processor( forward_batch.input_ids, hidden_states, model.lm_head, forward_batch, aux_hidden_states, + **logits_kwargs, ) def _execute_idle( diff --git a/python/sglang/srt/model_executor/runner_utils/buffers.py b/python/sglang/srt/model_executor/runner_utils/buffers.py index 97c7702c7288..572e8345aee5 100644 --- a/python/sglang/srt/model_executor/runner_utils/buffers.py +++ b/python/sglang/srt/model_executor/runner_utils/buffers.py @@ -132,8 +132,9 @@ def create( if pp_size > 1: is_mhc = hc_hidden_size is not None hs = hc_hidden_size if is_mhc else hidden_size + # Target verify expands each request into num_tokens_per_req rows. pp_proxy_tensors = { - "hidden_states": torch.zeros((max_bs, hs), dtype=dtype), + "hidden_states": torch.zeros((max_num_token, hs), dtype=dtype), } if not is_mhc: # Only Kimi K3 supplies num_blocks: its PP bank is token-major diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 65dd55549924..194003a394ca 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -2389,12 +2389,6 @@ def forward( if hasattr(forward_batch, _attr): delattr(forward_batch, _attr) capture_dspark = self.dspark_layers_to_capture is not None - if capture_dspark and dsa_use_prefill_cp(forward_batch): - raise NotImplementedError( - "DSpark aux hidden-state capture is not supported together with " - "DeepSeek-V4 prefill context parallelism (attn_cp_size > 1). Disable one " - "of them: DSpark static-verify is CP-off for v1." - ) dspark_aux_hidden_states: List[torch.Tensor] = [] # DSpark aux capture needs the per-layer eager loop (TBO's overlapped # execution cannot expose per-layer completed hidden states), so skip @@ -2445,16 +2439,35 @@ def forward( # CP all-gather only on the last PP rank; PP IPC carries CP-split tensors. if self.pp_group.is_last_rank and dsa_use_prefill_cp(forward_batch): + stream = torch.cuda.current_stream() hidden_states = cp_all_gather_rerange_output( hidden_states, self.cp_size, forward_batch, - torch.cuda.current_stream(), + stream, ) + # Gather DSpark aux tensors on the same CP token split. + if capture_dspark: + dspark_aux_hidden_states = [ + cp_all_gather_rerange_output( + aux, self.cp_size, forward_batch, stream + ) + for aux in dspark_aux_hidden_states + ] if not self.pp_group.is_last_rank: # Flatten 3D mHC tensor for PP IPC. - return PPProxyTensors({"hidden_states": hidden_states.flatten(1)}) + proxy_tensors = {"hidden_states": hidden_states.flatten(1)} + if capture_dspark: + if dspark_aux_hidden_states: + proxy_tensors["dspark_aux_hidden_states"] = torch.cat( + dspark_aux_hidden_states, dim=-1 + ) + else: + proxy_tensors["dspark_aux_hidden_states"] = hidden_states.new_empty( + hidden_states.shape[0], 0 + ) + return PPProxyTensors(proxy_tensors) pre_hc_head = hidden_states.flatten(1) @@ -2543,14 +2556,17 @@ def get_input_embeddings(self) -> nn.Module: return self.model.get_input_embeddings() def set_dspark_layers_to_capture(self, layer_ids: List[int]) -> None: - if not self.pp_group.is_last_rank: - return if layer_ids is None: raise ValueError( "DSPARK requires explicit layer_ids for aux hidden capture." ) - self.capture_aux_hidden_states = True - self.model.dspark_layers_to_capture = list(layer_ids) + local_layer_ids = [ + int(layer_id) + for layer_id in layer_ids + if self.model.start_layer <= int(layer_id) < self.model.end_layer + ] + self.capture_aux_hidden_states = bool(local_layer_ids) + self.model.dspark_layers_to_capture = local_layer_ids or None def determine_num_fused_shared_experts(self): self.num_fused_shared_experts = 0 diff --git a/python/sglang/srt/models/deepseek_v4_dspark.py b/python/sglang/srt/models/deepseek_v4_dspark.py index 2b01ec2e2a08..0eb5e3d50b1e 100644 --- a/python/sglang/srt/models/deepseek_v4_dspark.py +++ b/python/sglang/srt/models/deepseek_v4_dspark.py @@ -18,6 +18,11 @@ from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.layers.quantization.base_config import QuantizationConfig +from sglang.srt.layers.quantization.fp8 import Fp8LinearMethod +from sglang.srt.layers.quantization.fp8_utils import ( + inverse_transform_scale_ue8m0, + transform_scale_ue8m0, +) from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool @@ -40,9 +45,10 @@ gather_and_crop_vocab, run_markov_block, ) -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_disagg, get_parallel from sglang.srt.speculative.dspark_components.dspark_config import ( parse_dspark_draft_config, + use_lifecycle_only_draft_model, ) from sglang.srt.speculative.ragged_verify import ( RaggedVerifyMode, @@ -61,6 +67,83 @@ ) +class _BlockFp8LinearSlice(nn.Module): + """Persistent K-block slice of a loaded block-FP8 ReplicatedLinear.""" + + def __init__( + self, + *, + source: ReplicatedLinear, + feature_indices: List[int], + feature_width: int, + ) -> None: + super().__init__() + quant_method = source.quant_method + if not ( + isinstance(quant_method, Fp8LinearMethod) + and quant_method.block_quant + and not quant_method.use_mxfp8 + and not quant_method.use_marlin + ): + raise ValueError( + "DSpark block-FP8 projection slice requires a non-MXFP8 " + "block-quantized Fp8LinearMethod." + ) + + block_k = int(quant_method.weight_block_size[1]) + if feature_width % block_k != 0: + raise ValueError( + f"DSpark feature width {feature_width} must align to FP8 " + f"block_k={block_k}." + ) + blocks_per_feature = feature_width // block_k + device = source.weight.device + weight_columns = torch.cat( + [ + torch.arange( + feature_index * feature_width, + (feature_index + 1) * feature_width, + device=device, + ) + for feature_index in feature_indices + ] + ) + scale_columns = torch.cat( + [ + torch.arange( + feature_index * blocks_per_feature, + (feature_index + 1) * blocks_per_feature, + device=device, + ) + for feature_index in feature_indices + ] + ) + + self.quant_method = quant_method + self.weight = nn.Parameter( + source.weight.detach().index_select(1, weight_columns).contiguous(), + requires_grad=False, + ) + source_scale = source.weight_scale_inv.detach() + scale_is_ue8m0 = source.weight_scale_inv.format_ue8m0 + if scale_is_ue8m0: + source_scale = inverse_transform_scale_ue8m0( + source_scale, mn=source.weight.shape[0] + ) + local_scale = source_scale.index_select(1, scale_columns).contiguous() + if scale_is_ue8m0: + local_scale = transform_scale_ue8m0(local_scale, mn=source.weight.shape[0]) + self.weight_scale_inv = nn.Parameter( + local_scale, + requires_grad=False, + ) + self.weight_scale_inv.format_ue8m0 = scale_is_ue8m0 + self.register_parameter("bias", None) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + return self.quant_method.apply(self, hidden_states, bias=None) + + def apply_rotary_emb( x: torch.Tensor, freqs_cis: torch.Tensor, inverse: bool = False ) -> torch.Tensor: @@ -300,6 +383,7 @@ def __init__(self, *, vocab_size: int, markov_rank: int) -> None: self.markov_rank, self.vocab_size, bias=False, dtype=markov_w2_dtype ) self._tp_shard: Optional[MarkovW2ShardGeometry] = None + self._shard_group = None def configure_tp_shard(self, *, lm_head: nn.Module) -> None: if not self._opt_markov_w2_tp_shard: @@ -319,16 +403,23 @@ def configure_tp_shard(self, *, lm_head: nn.Module) -> None: f"num_embeddings_per_partition({per_partition}) * tp_size({tp_size}) != " f"num_embeddings_padded({num_padded})." ) - attn_tp_size = get_parallel().attn_tp_group.world_size - if attn_tp_size != tp_size: + # Follow lm_head's group choice; attn_tp_group degenerates to size 1 + # under prefill CP while lm_head still shards over the full TP group. + parallel = get_parallel() + shard_group = ( + parallel.attn_tp_group + if getattr(lm_head, "use_attn_tp_group", False) + else parallel.tp_group + ) + shard_group_size = shard_group.world_size + if shard_group_size != tp_size: raise ValueError( - "DSpark markov_w2 TP-shard needs the attn-TP group (used for the per-step " - f"all-gather) to equal the lm_head shard group, got attn_tp_size=" - f"{attn_tp_size} vs lm_head tp_size={tp_size}. This config (e.g. DP " - "attention without --enable-dp-lm-head, where lm_head shards over the " - "global TP group) is unsupported; disable " - "SGLANG_DSPARK_OPT_MARKOV_W2_TP_SHARD." + "DSpark markov_w2 TP-shard needs the per-step all-gather group to " + f"equal the lm_head shard group, got shard_group_size=" + f"{shard_group_size} vs lm_head tp_size={tp_size}. " + "Disable SGLANG_DSPARK_OPT_MARKOV_W2_TP_SHARD." ) + self._shard_group = shard_group self._tp_shard = MarkovW2ShardGeometry( tp_size=tp_size, org_vocab_start=int(lm_head.shard_indices.org_vocab_start_index), @@ -381,7 +472,8 @@ def _apply_step_logits_sharded( bias = F.linear(latent.float(), weight_local) step_local = BuildStepLocal.execute(bias=bias, base_local=base_local) if shard.tp_size > 1: - full = get_parallel().attn_tp_group.all_gather(step_local, dim=-1) + assert self._shard_group is not None + full = self._shard_group.all_gather(step_local, dim=-1) else: full = step_local return full[..., : self.vocab_size] @@ -591,6 +683,39 @@ def __init__( self.start_layer = 0 self.end_layer = self.num_stages + parallel = get_parallel() + self.is_lifecycle_only = use_lifecycle_only_draft_model( + disaggregation_mode=get_disagg().disaggregation_mode, + pp_rank=parallel.pp_rank, + pp_size=parallel.pp_size, + target_layer_ids=[ + int(layer_id) for layer_id in (dspark_config.target_layer_ids or []) + ], + num_hidden_layers=target_num_layers, + ) + self.hc_mult = int(config.hc_mult) + self.norm_eps = float(config.rms_norm_eps) + self.hc_eps = float(config.hc_eps) + self.embed_tokens: Optional[nn.Module] = None + self.lm_head: Optional[nn.Module] = None + self._partial_feature_indices: tuple[int, ...] = () + self._partial_main_proj: Optional[_BlockFp8LinearSlice] = None + self._use_fp32_lm_head = envs.SGLANG_DSPARK_FP32_LM_HEAD.get() + self._opt_markov_w2_tp_shard = envs.SGLANG_DSPARK_OPT_MARKOV_W2_TP_SHARD.get() + if self.is_lifecycle_only: + # Keep ModelRunner's distributed lifecycle aligned without building + # draft compute modules on PP ranks that cannot contribute context. + self.stages = nn.ModuleList() + self.markov_head = None + self.confidence_head = None + logger.info( + "DSpark PP rank %s uses a lifecycle-only draft model; all target " + "features are owned by final PP rank %s.", + parallel.pp_rank, + parallel.pp_size - 1, + ) + return + use_multi_stream = ( envs.SGLANG_OPT_USE_MULTI_STREAM_OVERLAP.get() and envs.SGLANG_DSPARK_ENABLE_MULTI_STREAM.get() @@ -621,14 +746,6 @@ def __init__( self.confidence_head = build_dspark_v4_confidence_head( config=config, markov_rank=int(dspark_config.markov_rank) ) - self.hc_mult = int(config.hc_mult) - self.norm_eps = float(config.rms_norm_eps) - self.hc_eps = float(config.hc_eps) - - self.embed_tokens: Optional[nn.Module] = None - self.lm_head: Optional[nn.Module] = None - self._use_fp32_lm_head = envs.SGLANG_DSPARK_FP32_LM_HEAD.get() - self._opt_markov_w2_tp_shard = envs.SGLANG_DSPARK_OPT_MARKOV_W2_TP_SHARD.get() @property def enable_confidence_head(self) -> bool: @@ -641,10 +758,101 @@ def attach_shared_modules( self.lm_head = lm_head self.markov_head.configure_tp_shard(lm_head=lm_head) - def project_target_hidden(self, main_hidden: torch.Tensor) -> torch.Tensor: + def prune_to_ctx_projection(self) -> None: + if self.is_lifecycle_only: + return stage0 = self.stages[0] - projected, _ = stage0.main_proj(main_hidden) - return stage0.main_norm(projected) + projection_stage = nn.Module() + projection_stage.main_proj = stage0.main_proj + projection_stage.main_norm = stage0.main_norm + # Preserve stages.0.main_proj parameter names for online weight updates. + self.stages = nn.ModuleList([projection_stage]) + self.markov_head = None + self.confidence_head = None + self.embed_tokens = None + self.lm_head = None + del stage0 + torch.cuda.empty_cache() + + def project_target_hidden(self, main_hidden: torch.Tensor) -> torch.Tensor: + projected, _ = self.stages[0].main_proj(main_hidden) + return self.stages[0].main_norm(projected) + + def prepare_target_hidden_partial(self, feature_indices: List[int]) -> None: + feature_indices = [int(index) for index in feature_indices] + main_proj = self.stages[0].main_proj + quant_method = main_proj.quant_method + self._partial_feature_indices = tuple(feature_indices) + if not ( + isinstance(quant_method, Fp8LinearMethod) + and quant_method.block_quant + and not quant_method.use_mxfp8 + and not quant_method.use_marlin + ): + self._partial_main_proj = None + logger.warning( + "DSpark partial projection cannot slice quant method %s; " + "falling back to the full-K projection.", + type(quant_method).__name__, + ) + return + self._partial_main_proj = _BlockFp8LinearSlice( + source=main_proj, + feature_indices=feature_indices, + feature_width=int(self.config.hidden_size), + ) + logger.info( + "DSpark block-FP8 partial projection uses feature columns %s " + "(local K=%s, full K=%s).", + feature_indices, + len(feature_indices) * int(self.config.hidden_size), + int(main_proj.weight.shape[1]), + ) + + def project_target_hidden_partial( + self, main_hidden: torch.Tensor, feature_indices: list[int] + ) -> torch.Tensor: + """Project PP-local target features into an additive pre-norm context.""" + if not feature_indices: + raise ValueError("feature_indices must be non-empty.") + feature_indices = [int(index) for index in feature_indices] + if min(feature_indices) < 0 or max(feature_indices) >= self.num_target_features: + raise ValueError( + "DeepSeek-V4 DSpark feature_indices out of range: " + f"{feature_indices=} {self.num_target_features=}." + ) + + hidden_size = int(self.config.hidden_size) + expected = len(feature_indices) * hidden_size + if main_hidden.ndim != 2 or int(main_hidden.shape[-1]) != expected: + raise ValueError( + "DeepSeek-V4 DSpark partial main_hidden feature dim mismatch. " + f"Expected shape [N, {expected}] for {feature_indices=}, " + f"but got shape={tuple(main_hidden.shape)}." + ) + + if ( + self._partial_main_proj is not None + and tuple(feature_indices) == self._partial_feature_indices + ): + return self._partial_main_proj(main_hidden) + + local_features = main_hidden.view( + main_hidden.shape[0], len(feature_indices), hidden_size + ) + full_features = main_hidden.new_zeros( + main_hidden.shape[0], self.num_target_features, hidden_size + ) + # Reuse ReplicatedLinear.forward so FP8 weights and scales follow the + # same quantization path as the full, non-PP projection. + feature_index = torch.tensor( + feature_indices, dtype=torch.long, device=main_hidden.device + ) + full_features.index_copy_(1, feature_index, local_features) + projected, _ = self.stages[0].main_proj( + full_features.view(main_hidden.shape[0], -1) + ) + return projected def write_target_hidden_kv( self, @@ -655,6 +863,37 @@ def write_target_hidden_kv( pool: DeepSeekV4TokenToKVPool, ) -> None: main_x = self.project_target_hidden(main_hidden) + self._write_context_hidden_kv( + main_x=main_x, + swa_loc=swa_loc, + positions=positions, + pool=pool, + ) + + def write_projected_context_kv( + self, + *, + projected_context: torch.Tensor, + swa_loc: torch.Tensor, + positions: torch.Tensor, + pool: DeepSeekV4TokenToKVPool, + ) -> None: + main_x = self.stages[0].main_norm(projected_context) + self._write_context_hidden_kv( + main_x=main_x, + swa_loc=swa_loc, + positions=positions, + pool=pool, + ) + + def _write_context_hidden_kv( + self, + *, + main_x: torch.Tensor, + swa_loc: torch.Tensor, + positions: torch.Tensor, + pool: DeepSeekV4TokenToKVPool, + ) -> None: swa_loc = swa_loc.to(torch.int32) kvs = CommitKvProj.execute( main_x=main_x, @@ -758,6 +997,8 @@ def compute_confidence( return confidence def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> None: + if self.is_lifecycle_only: + return params_dict = dict(self.named_parameters()) loaded_params = set() diff --git a/python/sglang/srt/models/dflash.py b/python/sglang/srt/models/dflash.py index 69c01c7adc9d..60f471127200 100644 --- a/python/sglang/srt/models/dflash.py +++ b/python/sglang/srt/models/dflash.py @@ -399,6 +399,40 @@ def project_target_hidden(self, target_hidden: torch.Tensor) -> torch.Tensor: ) return self.hidden_norm(self.fc(target_hidden)) + def project_target_hidden_partial( + self, target_hidden: torch.Tensor, feature_indices: list[int] + ) -> torch.Tensor: + """Project PP-local target features into an additive pre-norm context.""" + if not feature_indices: + raise ValueError("feature_indices must be non-empty.") + feature_indices = [int(i) for i in feature_indices] + if ( + min(feature_indices) < 0 + or max(feature_indices) >= self.num_context_features + ): + raise ValueError( + "feature_indices out of range for DFLASH context projection: " + f"{feature_indices=} {self.num_context_features=}." + ) + hidden_size = int(self.config.hidden_size) + expected = len(feature_indices) * hidden_size + if target_hidden.ndim != 2 or int(target_hidden.shape[-1]) != expected: + raise ValueError( + "DFLASH partial target_hidden feature dim mismatch. " + f"Expected shape [N, {expected}] for {feature_indices=}, " + f"but got shape={tuple(target_hidden.shape)}." + ) + + cols = [] + for idx in feature_indices: + start = idx * hidden_size + cols.extend(range(start, start + hidden_size)) + index = torch.tensor(cols, dtype=torch.long, device=self.fc.weight.device) + weight = self.fc.weight.index_select(1, index) + if target_hidden.dtype != weight.dtype: + target_hidden = target_hidden.to(weight.dtype) + return F.linear(target_hidden, weight) + @torch.no_grad() def forward( self, @@ -585,5 +619,34 @@ def project_target_hidden(self, target_hidden: torch.Tensor) -> torch.Tensor: fused = normed.reshape(target_hidden.shape[0], -1) return self.hidden_norm(self.fc(fused)) + def project_target_hidden_partial( + self, target_hidden: torch.Tensor, feature_indices: list[int] + ) -> torch.Tensor: + if not feature_indices: + raise ValueError("feature_indices must be non-empty.") + feature_indices = [int(i) for i in feature_indices] + hidden_size = int(self.config.hidden_size) + expected = len(feature_indices) * hidden_size + if target_hidden.ndim != 2 or int(target_hidden.shape[-1]) != expected: + raise ValueError( + "Laguna DFLASH partial target_hidden feature dim mismatch. " + f"Expected shape [N, {expected}] for {feature_indices=}, " + f"but got shape={tuple(target_hidden.shape)}." + ) + slices = target_hidden.view( + target_hidden.shape[0], len(feature_indices), hidden_size + ) + compute_dtype = self.fc.weight.dtype + if slices.dtype != compute_dtype: + slices = slices.to(compute_dtype) + normed = torch.empty_like(slices) + for out_idx, feature_idx in enumerate(feature_indices): + normed[:, out_idx, :] = self.aux_hidden_norms[feature_idx]( + slices[:, out_idx, :] + ) + return super().project_target_hidden_partial( + normed.reshape(target_hidden.shape[0], -1), feature_indices + ) + EntryClass = [DFlashDraftModel, DFlashLagunaForCausalLM] diff --git a/python/sglang/srt/models/dspark.py b/python/sglang/srt/models/dspark.py index f9ff64733eec..75ee16595276 100644 --- a/python/sglang/srt/models/dspark.py +++ b/python/sglang/srt/models/dspark.py @@ -513,6 +513,44 @@ def write_target_hidden_kv( commit_lens: Optional[torch.Tensor] = None, ) -> None: ctx_hidden = self.project_target_hidden(target_hidden) + self.write_context_hidden_kv( + ctx_hidden=ctx_hidden, + pool=pool, + positions=positions, + cache_loc=cache_loc, + cache_loc_2d=cache_loc_2d, + commit_lens=commit_lens, + ) + + def write_projected_context_kv( + self, + *, + projected_context: torch.Tensor, + pool, + positions: torch.Tensor, + cache_loc: torch.Tensor, + cache_loc_2d: Optional[torch.Tensor] = None, + commit_lens: Optional[torch.Tensor] = None, + ) -> None: + self.write_context_hidden_kv( + ctx_hidden=self.hidden_norm(projected_context), + pool=pool, + positions=positions, + cache_loc=cache_loc, + cache_loc_2d=cache_loc_2d, + commit_lens=commit_lens, + ) + + def write_context_hidden_kv( + self, + *, + ctx_hidden: torch.Tensor, + pool, + positions: torch.Tensor, + cache_loc: torch.Tensor, + cache_loc_2d: Optional[torch.Tensor] = None, + commit_lens: Optional[torch.Tensor] = None, + ) -> None: stacked = self._stacked_ctx_kv_params() if stacked is not None: k_all, v_all = self._project_ctx_kv_stacked( diff --git a/python/sglang/srt/models/kimi_k3.py b/python/sglang/srt/models/kimi_k3.py index 1d07b16a0e16..78929c7f03a9 100644 --- a/python/sglang/srt/models/kimi_k3.py +++ b/python/sglang/srt/models/kimi_k3.py @@ -2500,16 +2500,29 @@ def forward( ) if not self.pp_group.is_last_rank: + proxy_tensors = { + "hidden_states": hidden_states, + "residual": residual, + } + if self.dspark_layers_to_capture is not None: + if aux_hidden_states: + proxy_tensors["dspark_aux_hidden_states"] = torch.cat( + aux_hidden_states, dim=-1 + ) + else: + proxy_tensors["dspark_aux_hidden_states"] = hidden_states.new_empty( + hidden_states.shape[0], 0 + ) assert not sp_sharded if attn_res is not None: if residual is not None: # Materialize the delayed MLP add: the wire carries the # full stream head (bit-identical to the fused fold). hidden_states = residual + hidden_states + proxy_tensors["hidden_states"] = hidden_states residual = attn_res.block_residual # raw bank across ranks - return PPProxyTensors( - {"hidden_states": hidden_states, "residual": residual} - ) + proxy_tensors["residual"] = residual + return PPProxyTensors(proxy_tensors) if hidden_states.shape[0] != 0: if attn_res is not None: @@ -2618,18 +2631,17 @@ def get_input_embeddings(self): return self.model.embed_tokens def set_dspark_layers_to_capture(self, layer_ids: list[int]) -> None: - if self.pp_group.world_size > 1: - # Capture layers living on non-last PP ranks would be silently - # skipped (the flag is only set on the last rank). - raise NotImplementedError("DSPARK aux hidden capture requires PP=1.") - if not self.pp_group.is_last_rank: - return if layer_ids is None: raise ValueError( "DSPARK requires explicit layer_ids for aux hidden capture." ) - self.capture_aux_hidden_states = True - self.model.dspark_layers_to_capture = list(layer_ids) + local_layer_ids = [ + int(layer_id) + for layer_id in layer_ids + if self.model.start_layer <= int(layer_id) < self.model.end_layer + ] + self.capture_aux_hidden_states = bool(local_layer_ids) + self.model.dspark_layers_to_capture = local_layer_ids or None @torch.no_grad() def forward( @@ -2740,6 +2752,35 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]): loaded_params: set[str] = set() num_hidden_layers = self.config.num_hidden_layers + + def maybe_load_truncated_moe_gate( + param_name: str, param: torch.Tensor, loaded_weight: torch.Tensor + ) -> bool: + if self.config.num_experts is None: + return False + if not ( + param_name.endswith(".mlp.gate.weight") + or param_name.endswith(".mlp.gate.e_score_correction_bias") + ): + return False + if loaded_weight.shape == param.data.shape: + return False + keep_num_experts = int(self.config.num_experts) + if loaded_weight.shape[0] < keep_num_experts: + raise ValueError( + f"Cannot truncate Kimi-K3 MoE gate weight {param_name}: " + f"checkpoint has {loaded_weight.shape[0]} experts, " + f"requested {keep_num_experts}." + ) + truncated_weight = loaded_weight.narrow(0, 0, keep_num_experts) + if truncated_weight.shape != param.data.shape: + raise ValueError( + f"Unexpected truncated Kimi-K3 MoE gate shape for {param_name}: " + f"{truncated_weight.shape=} {param.data.shape=}." + ) + param.data.copy_(truncated_weight) + return True + for args in weights: name, loaded_weight = args[:2] kwargs = args[2] if len(args) > 2 else {} @@ -2846,6 +2887,9 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]): if name not in params_dict: continue param = params_dict[name] + if maybe_load_truncated_moe_gate(name, param, loaded_weight): + loaded_params.add(name) + continue weight_loader = getattr( param, "weight_loader", default_weight_loader ) diff --git a/python/sglang/srt/models/kimi_linear.py b/python/sglang/srt/models/kimi_linear.py index b8fc10636d6a..90498683c468 100644 --- a/python/sglang/srt/models/kimi_linear.py +++ b/python/sglang/srt/models/kimi_linear.py @@ -623,12 +623,20 @@ def forward( ) if not self.pp_group.is_last_rank: - return PPProxyTensors( - { - "hidden_states": hidden_states, - "residual": residual, - } - ) + proxy_tensors = { + "hidden_states": hidden_states, + "residual": residual, + } + if self.dspark_layers_to_capture is not None: + if aux_hidden_states: + proxy_tensors["dspark_aux_hidden_states"] = torch.cat( + aux_hidden_states, dim=-1 + ) + else: + proxy_tensors["dspark_aux_hidden_states"] = hidden_states.new_empty( + hidden_states.shape[0], 0 + ) + return PPProxyTensors(proxy_tensors) else: if hidden_states.shape[0] != 0: if residual is None: @@ -674,16 +682,17 @@ def get_input_embeddings(self): return self.model.embed_tokens def set_dspark_layers_to_capture(self, layer_ids: list[int]) -> None: - if self.pp_group.world_size > 1: - raise NotImplementedError("DSPARK aux hidden capture requires PP=1.") - if not self.pp_group.is_last_rank: - return if layer_ids is None: raise ValueError( "DSPARK requires explicit layer_ids for aux hidden capture." ) - self.capture_aux_hidden_states = True - self.model.dspark_layers_to_capture = list(layer_ids) + local_layer_ids = [ + int(layer_id) + for layer_id in layer_ids + if self.model.start_layer <= int(layer_id) < self.model.end_layer + ] + self.capture_aux_hidden_states = bool(local_layer_ids) + self.model.dspark_layers_to_capture = local_layer_ids or None @torch.no_grad() def forward( diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 18664959e959..1452fb1e8ec7 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -8720,8 +8720,20 @@ def check_server_args(self): if self.pp_size > 1: assert ( - self.disable_overlap_schedule and self.speculative_algorithm is None - ), "Pipeline parallelism is not compatible with overlap schedule, speculative decoding" + self.disable_overlap_schedule + ), "Pipeline parallelism is not compatible with overlap schedule" + pp_dspark_prefill = ( + self.speculative_algorithm or "" + ).upper() == "DSPARK" and self.disaggregation_mode == "prefill" + assert self.speculative_algorithm is None or pp_dspark_prefill, ( + "Pipeline parallelism with speculative decoding is only supported " + "for DSPARK on a PD prefill server" + ) + assert self.min_free_slots_delay is None, ( + "--min-free-slots-delay is not supported with pipeline " + "parallelism: allocatable slots per microbatch are bounded by " + "pp-max-micro-batch-size, so the threshold may never be reached" + ) assert not ( self.dp_size > 1 and self.nnodes != 1 and not self.enable_dp_attention diff --git a/python/sglang/srt/speculative/dspark_components/dspark_config.py b/python/sglang/srt/speculative/dspark_components/dspark_config.py index 0cac50977eee..f93becc43e7e 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_config.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_config.py @@ -38,6 +38,45 @@ def draft_is_deepseek_v4(*, server_args: ServerArgs) -> bool: return draft_hf_config is not None and is_deepseek_v4(draft_hf_config) +def resolve_single_owner_pp_rank( + *, target_layer_ids: List[int], num_hidden_layers: int, pp_size: int +) -> Optional[int]: + from sglang.srt.distributed.utils import get_pp_indices + + if not target_layer_ids: + return None + for pp_rank in range(pp_size): + start_layer, end_layer = get_pp_indices( + num_hidden_layers=num_hidden_layers, + pp_rank=pp_rank, + pp_size=pp_size, + ) + if all(start_layer <= layer_id < end_layer for layer_id in target_layer_ids): + return pp_rank + return None + + +def use_lifecycle_only_draft_model( + *, + disaggregation_mode: str, + pp_rank: int, + pp_size: int, + target_layer_ids: List[int], + num_hidden_layers: int, +) -> bool: + owner_pp_rank = resolve_single_owner_pp_rank( + target_layer_ids=target_layer_ids, + num_hidden_layers=num_hidden_layers, + pp_size=pp_size, + ) + return ( + disaggregation_mode == "prefill" + and pp_size > 1 + and owner_pp_rank == pp_size - 1 + and pp_rank != owner_pp_rank + ) + + def dspark_gamma_from_num_draft_tokens(num_draft_tokens: int) -> int: gamma = int(num_draft_tokens) - 1 if gamma < 1: diff --git a/python/sglang/srt/speculative/dspark_components/dspark_kv_inject.py b/python/sglang/srt/speculative/dspark_components/dspark_kv_inject.py index d1e44145a360..4e9b39a510cc 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_kv_inject.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_kv_inject.py @@ -2,11 +2,15 @@ import torch +from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import ( + is_unified_kv_triton, +) from sglang.kernels.ops.speculative.cache_locs import assign_extend_cache_locs_func from sglang.kernels.ops.speculative.dspark.dspark_verify_window import ( BuildCommitInjectLayout, ) from sglang.srt.managers.schedule_batch import ScheduleBatch +from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.speculative.ragged_verify import RaggedVerifyLayout @@ -36,6 +40,8 @@ def inject_target_hidden( positions: torch.Tensor, cache_loc_2d: Optional[torch.Tensor] = None, commit_lens: Optional[torch.Tensor] = None, + state_slot: Optional[torch.Tensor] = None, + final_pos: Optional[torch.Tensor] = None, ) -> None: if target_hidden is None or target_hidden.numel() == 0: return @@ -54,6 +60,14 @@ def inject_target_hidden( commit_lens = commit_lens.to( device=device, dtype=torch.int32, non_blocking=True ) + if state_slot is not None: + state_slot = state_slot.to( + device=device, dtype=torch.int64, non_blocking=True + ) + if final_pos is not None: + final_pos = final_pos.to( + device=device, dtype=torch.int64, non_blocking=True + ) pool = self.draft_model_runner.token_to_kv_pool if hasattr(pool, "set_swa_key_buffer_radix_fused_norm_rope"): @@ -64,6 +78,8 @@ def inject_target_hidden( positions=positions, cache_loc_2d=cache_loc_2d, commit_lens=commit_lens, + state_slot=state_slot, + final_pos=final_pos, ) return @@ -77,6 +93,84 @@ def inject_target_hidden( commit_lens=commit_lens, ) + def inject_projected_context( + self, + *, + projected_context: torch.Tensor, + cache_loc: torch.Tensor, + positions: torch.Tensor, + cache_loc_2d: Optional[torch.Tensor] = None, + commit_lens: Optional[torch.Tensor] = None, + state_slot: Optional[torch.Tensor] = None, + final_pos: Optional[torch.Tensor] = None, + ) -> None: + if projected_context is None or projected_context.numel() == 0: + return + device = self.model_runner.device + cache_loc = cache_loc.to(device=device, dtype=torch.int64, non_blocking=True) + positions = positions.to(device=device, dtype=torch.int64, non_blocking=True) + projected_context = projected_context.to(device=device, non_blocking=True) + n_real = positions.shape[0] + if projected_context.shape[0] > n_real: + projected_context = projected_context[:n_real] + if cache_loc_2d is not None: + cache_loc_2d = cache_loc_2d.to( + device=device, dtype=torch.int64, non_blocking=True + ) + if commit_lens is not None: + commit_lens = commit_lens.to( + device=device, dtype=torch.int32, non_blocking=True + ) + if state_slot is not None: + state_slot = state_slot.to( + device=device, dtype=torch.int64, non_blocking=True + ) + if final_pos is not None: + final_pos = final_pos.to( + device=device, dtype=torch.int64, non_blocking=True + ) + + pool = self.draft_model_runner.token_to_kv_pool + if isinstance(pool, DeepSeekV4TokenToKVPool): + if is_unified_kv_triton(): + swa_loc = self._unified_inject_loc( + pool=pool, + positions=positions, + cache_loc_2d=cache_loc_2d, + commit_lens=commit_lens, + state_slot=state_slot, + final_pos=final_pos, + ) + else: + swa_loc = pool.translate_loc_from_full_to_swa(cache_loc).to(torch.int32) + if commit_lens is not None and cache_loc_2d is not None: + bs, verify_len = cache_loc_2d.shape + col = torch.arange(verify_len, device=cache_loc.device).view(1, -1) + committed_mask = ( + col < commit_lens.to(torch.long).view(-1, 1) + ).reshape(-1) + swa_loc = torch.where( + committed_mask, swa_loc, torch.full_like(swa_loc, -1) + ) + with torch.inference_mode(): + self.draft_model.write_projected_context_kv( + projected_context=projected_context, + swa_loc=swa_loc, + positions=positions, + pool=pool, + ) + return + + with torch.inference_mode(): + self.draft_model.write_projected_context_kv( + projected_context=projected_context, + pool=pool, + positions=positions, + cache_loc=cache_loc, + cache_loc_2d=cache_loc_2d, + commit_lens=commit_lens, + ) + def _inject_mla( self, *, @@ -86,13 +180,29 @@ def _inject_mla( positions: torch.Tensor, cache_loc_2d: Optional[torch.Tensor], commit_lens: Optional[torch.Tensor], + state_slot: Optional[torch.Tensor] = None, + final_pos: Optional[torch.Tensor] = None, ) -> None: - swa_loc = pool.translate_loc_from_full_to_swa(cache_loc).to(torch.int32) - if commit_lens is not None and cache_loc_2d is not None: - bs, verify_len = cache_loc_2d.shape - col = torch.arange(verify_len, device=cache_loc.device).view(1, -1) - committed_mask = (col < commit_lens.to(torch.long).view(-1, 1)).reshape(-1) - swa_loc = torch.where(committed_mask, swa_loc, torch.full_like(swa_loc, -1)) + if is_unified_kv_triton(): + swa_loc = self._unified_inject_loc( + pool=pool, + positions=positions, + cache_loc_2d=cache_loc_2d, + commit_lens=commit_lens, + state_slot=state_slot, + final_pos=final_pos, + ) + else: + swa_loc = pool.translate_loc_from_full_to_swa(cache_loc).to(torch.int32) + if commit_lens is not None and cache_loc_2d is not None: + bs, verify_len = cache_loc_2d.shape + col = torch.arange(verify_len, device=cache_loc.device).view(1, -1) + committed_mask = ( + col < commit_lens.to(torch.long).view(-1, 1) + ).reshape(-1) + swa_loc = torch.where( + committed_mask, swa_loc, torch.full_like(swa_loc, -1) + ) with torch.inference_mode(): self.draft_model.write_target_hidden_kv( @@ -102,6 +212,38 @@ def _inject_mla( pool=pool, ) + def _unified_inject_loc( + self, + *, + pool, + positions: torch.Tensor, + cache_loc_2d: Optional[torch.Tensor], + commit_lens: Optional[torch.Tensor], + state_slot: Optional[torch.Tensor], + final_pos: Optional[torch.Tensor], + ) -> torch.Tensor: + """Build unified-KV SWA ring locations for target-hidden injection.""" + if state_slot is None: + raise RuntimeError( + "unified_kv target-hidden injection requires state_slot " + "(per-token draft req_pool_indices)." + ) + ring = pool.unified_swa_ring_size + win = pool.unified_swa_window + pos = positions.to(torch.int64) + loc = state_slot.to(torch.int64) * ring + pos % ring + if final_pos is not None: + keep = pos > (final_pos.to(torch.int64) - win) + loc = torch.where(keep, loc, torch.full_like(loc, -1)) + if commit_lens is not None and cache_loc_2d is not None: + _, verify_len = cache_loc_2d.shape + col = torch.arange(verify_len, device=positions.device).view(1, -1) + committed = ( + col < commit_lens.to(torch.long).view(-1, 1) + ).reshape(-1) + loc = torch.where(committed, loc, torch.full_like(loc, -1)) + return loc.to(torch.int32) + def inject_ragged( self, *, diff --git a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py index d5b1ec03b6d1..bfc5a1f5415d 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py @@ -1,13 +1,20 @@ import logging -from contextlib import nullcontext +from contextlib import ExitStack, contextmanager from dataclasses import replace from typing import Optional import torch +from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import ( + is_unified_kv_triton, +) from sglang.srt.configs.hybrid_arch import mambaish_config from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.environ import envs +from sglang.srt.layers.moe.utils import ( + speculative_moe_a2a_backend_context, + speculative_moe_backend_context, +) from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.scheduler import GenerationBatchResult from sglang.srt.managers.tp_worker import TpModelWorker @@ -29,6 +36,7 @@ DSV4_DRAFT_ATTENTION_BACKEND, draft_is_deepseek_v4, resolve_runtime_config, + resolve_single_owner_pp_rank, ) from sglang.srt.speculative.dspark_components.dspark_draft import ( DraftBlockProposer, @@ -67,6 +75,12 @@ logger = logging.getLogger(__name__) +def _is_context_only_pp_prefill_rank( + *, disaggregation_mode: str, pp_rank: int, pp_size: int +) -> bool: + return disaggregation_mode == "prefill" and pp_size > 1 and pp_rank < pp_size - 1 + + class DSparkWorkerV2(BaseSpecWorker): def __init__( @@ -93,6 +107,27 @@ def __init__( server_args.enable_dp_attention and not self._draft_is_moe ) self._is_pd_prefill = server_args.disaggregation_mode == "prefill" + self._is_context_only_pp_prefill_rank = _is_context_only_pp_prefill_rank( + disaggregation_mode=server_args.disaggregation_mode, + pp_rank=ps.pp_rank, + pp_size=ps.pp_size, + ) + self._use_full_projection_prefill = False + if self._is_pd_prefill and self._draft_is_moe and ps.pp_size > 1: + target_layer_ids = [ + int(layer_id) + for layer_id in ( + self.model_runner.spec_aux_config.dflash_target_layer_ids or [] + ) + ] + owner_pp_rank = resolve_single_owner_pp_rank( + target_layer_ids=target_layer_ids, + num_hidden_layers=self.model_runner.model_config.num_hidden_layers, + pp_size=ps.pp_size, + ) + self._use_full_projection_prefill = ( + owner_pp_rank == ps.pp_size - 1 and ps.pp_rank == owner_pp_rank + ) self._decode_graph_allowed = ( not server_args.disable_cuda_graph and not self._is_pd_prefill ) @@ -111,7 +146,7 @@ def __init__( bundle = build_draft_tp_worker( server_args=server_args, gpu_id=gpu_id, - ps=replace(ps, pp_rank=0), + ps=ps, nccl_port=nccl_port, target_model_config=target_worker.model_runner.model_config, algo_label="DSPARK", @@ -123,6 +158,9 @@ def __init__( self.draft_model_runner = bundle.draft_model_runner self.draft_model = bundle.draft_model self._draft_sampler = None + self._is_lifecycle_only_pp_prefill_rank = ( + self._draft_is_moe and self.draft_model.is_lifecycle_only + ) runtime_config = resolve_runtime_config( draft_hf_config=self.draft_model_runner.model_config.hf_config, @@ -135,6 +173,12 @@ def __init__( self.verify_num_draft_tokens = runtime_config.verify_num_draft_tokens self.speculative_num_draft_tokens = self.verify_num_draft_tokens self._mask_token_id = runtime_config.mask_token_id + self._pp_context_feature_indices: Optional[list[int]] = None + self._next_pp_proxy_tensors = None + + if self._is_lifecycle_only_pp_prefill_rank: + self._init_lifecycle_only_prefill() + return if self.ps.tp_rank == 0: logger.info( @@ -158,14 +202,16 @@ def __init__( target_model = self.target_worker.model_runner.model lm_head = getattr(target_model, "lm_head", None) - if lm_head is None or not hasattr(lm_head, "weight"): + needs_lm_head = not self._is_context_only_pp_prefill_rank + if needs_lm_head and (lm_head is None or not hasattr(lm_head, "weight")): raise RuntimeError( "DSpark requires the target model to expose `lm_head` with `weight`." ) - self.draft_model.attach_shared_modules( - embed_tokens=self._resolve_target_embed_tokens(target_model), - lm_head=lm_head, - ) + if needs_lm_head: + self.draft_model.attach_shared_modules( + embed_tokens=self._resolve_target_embed_tokens(target_model), + lm_head=lm_head, + ) self._verify_planner = DSparkVerifyPlanner( draft_model=self.draft_model, @@ -269,20 +315,88 @@ def __init__( simulate_acc_len=self._simulate_acc_len, ) - if self._is_pd_prefill and not self._draft_is_moe: - self.draft_model.prune_to_ctx_kv_injection() + if self._is_pd_prefill: + if self._draft_is_moe and self._is_context_only_pp_prefill_rank: + self.draft_model.prune_to_ctx_projection() + elif not self._draft_is_moe: + self.draft_model.prune_to_ctx_kv_injection() + self._init_pp_context_feature_indices() + + def _init_lifecycle_only_prefill(self) -> None: + self._verify_planner = None + self._kv_injector = None + self._proposer = None + self._verify_epilogue = None + self._verify_executor = None + self._simulate_acc_len = float(envs.SGLANG_SIMULATE_ACC_LEN.get()) + self._forced_budget_frac = None + self._need_mamba_verify_commit = False + self._observers = None def _resolve_target_embed_tokens(self, target_model): if hasattr(target_model, "get_input_embeddings"): return target_model.get_input_embeddings() return target_model.model.get_input_embeddings() + def _init_pp_context_feature_indices(self) -> None: + if self.ps.pp_size <= 1: + return + if not hasattr(self.draft_model, "project_target_hidden_partial"): + return + + target_layer_ids = getattr( + self.model_runner.spec_aux_config, "dflash_target_layer_ids", None + ) + if not target_layer_ids: + return + + target_model = self.target_worker.model_runner.model + start_layer = int(getattr(target_model, "start_layer", 0)) + end_layer = int(getattr(target_model, "end_layer", 0)) + target_layer_to_feature = { + int(layer_id): idx for idx, layer_id in enumerate(target_layer_ids) + } + local_feature_indices = [ + target_layer_to_feature[layer_id] + for layer_id in range(start_layer, end_layer) + if layer_id in target_layer_to_feature + ] + if not local_feature_indices: + return + + self._pp_context_feature_indices = local_feature_indices + if self._draft_is_moe and not self._use_full_projection_prefill: + self.draft_model.prepare_target_hidden_partial(local_feature_indices) + if self.ps.tp_rank == 0: + logger.info( + "DSpark PP-local context projection will accumulate feature indices %s " + "from target layers %s for PP rank %s local layers [%s, %s).", + local_feature_indices, + list(target_layer_ids), + self.ps.pp_rank, + start_layer, + end_layer, + ) + @property def carries_confidence(self) -> bool: + if self._is_lifecycle_only_pp_prefill_rank: + return False return self._verify_planner.carries_confidence + @property + def is_lifecycle_only_pp_prefill_rank(self) -> bool: + return self._is_lifecycle_only_pp_prefill_rank + + def _draft_model_runners(self) -> tuple: + if self._is_lifecycle_only_pp_prefill_rank: + return () + return super()._draft_model_runners() + @property def spec_v2_attn_backends(self) -> tuple: + if self._is_context_only_pp_prefill_rank: + return (self._target_worker.model_runner.attn_backend,) return ( self._target_worker.model_runner.attn_backend, self.draft_model_runner.attn_backend, @@ -293,10 +407,14 @@ def __getattr__(self, name): raise AttributeError(name) return getattr(self.target_worker, name) + @contextmanager def _draft_context(self): - if self._draft_dp_context_enabled: - return draft_tp_context(get_parallel().attn_tp_group) - return nullcontext() + with ExitStack() as stack: + if self._draft_dp_context_enabled: + stack.enter_context(draft_tp_context(get_parallel().attn_tp_group)) + stack.enter_context(speculative_moe_backend_context()) + stack.enter_context(speculative_moe_a2a_backend_context()) + yield def alloc_memory_pool( self, @@ -304,6 +422,26 @@ def alloc_memory_pool( req_to_token_pool=None, token_to_kv_pool_allocator=None, ): + if self._is_lifecycle_only_pp_prefill_rank: + return + if memory_pool_config is not None and self._is_context_only_pp_prefill_rank: + page_size = int(self.page_size) + + def _minimal_capacity(capacity): + if capacity is None or capacity == 0: + return capacity + return page_size + + memory_pool_config = replace( + memory_pool_config, + max_total_num_tokens=page_size, + full_max_total_num_tokens=_minimal_capacity( + memory_pool_config.full_max_total_num_tokens + ), + swa_max_total_num_tokens=_minimal_capacity( + memory_pool_config.swa_max_total_num_tokens + ), + ) self._draft_worker.alloc_memory_pool( memory_pool_config=memory_pool_config, req_to_token_pool=req_to_token_pool, @@ -311,6 +449,9 @@ def alloc_memory_pool( ) def init_attention_backends(self): + if self._is_context_only_pp_prefill_rank: + self._need_mamba_verify_commit = False + return with self._draft_context(): self._draft_worker.init_attention_backends() self._need_mamba_verify_commit = mambaish_config( @@ -321,6 +462,8 @@ def init_attention_backends(self): ) def init_cuda_graphs(self): + if self._is_context_only_pp_prefill_rank: + return capture_decode_cuda_graph = self._decode_graph_allowed if is_cuda() and capture_decode_cuda_graph: available_mem = get_available_gpu_memory(self.device, self.gpu_id) @@ -365,64 +508,153 @@ def _maybe_build_draft_sampler(self): def clear_cache_pool(self): pass + def _refresh_partial_projection(self) -> None: + if ( + self._draft_is_moe + and not self._is_lifecycle_only_pp_prefill_rank + and not self._use_full_projection_prefill + and self._pp_context_feature_indices is not None + ): + self.draft_model.prepare_target_hidden_partial( + self._pp_context_feature_indices + ) + + def update_weights_from_disk(self, recv_req): + success, message = super().update_weights_from_disk(recv_req) + if success: + self._refresh_partial_projection() + return success, message + + def update_weights_from_ipc(self, recv_req): + success, message = super().update_weights_from_ipc(recv_req) + if success: + self._refresh_partial_projection() + return success, message + def set_dspark_forced_budget_frac(self, frac: Optional[float]) -> None: self._forced_budget_frac = frac + if self._is_lifecycle_only_pp_prefill_rank: + return self._verify_planner.set_forced_budget_frac(frac) def dump_info_records(self) -> Optional[dict]: + if self._is_lifecycle_only_pp_prefill_rank: + return None return self._observers.dump_info_records() def clear_info_records(self) -> None: + if self._is_lifecycle_only_pp_prefill_rank: + return self._observers.clear_info_records() def block_accept_estimate_log_suffix(self) -> Optional[str]: + if self._is_lifecycle_only_pp_prefill_rank: + return None return self._observers.block_accept_estimate_log_suffix() def note_request_finished(self, *, rid: str, natural_stop: bool) -> None: + if self._is_lifecycle_only_pp_prefill_rank: + return self._observers.note_request_finished(rid=rid, natural_stop=natural_stop) + def set_pp_proxy_tensors_for_next_forward(self, pp_proxy_tensors) -> None: + self._next_pp_proxy_tensors = pp_proxy_tensors + def forward_batch_generation( self, batch: ScheduleBatch, on_publish=None, grammar_barrier=None, ) -> GenerationBatchResult: + pp_proxy_tensors = self._next_pp_proxy_tensors + self._next_pp_proxy_tensors = None + if getattr(batch, "return_logprob", False): raise ValueError( "DSpark speculative decoding does not support return_logprob yet." ) + if self._is_lifecycle_only_pp_prefill_rank: + return self._forward_lifecycle_only_prefill( + batch=batch, + on_publish=on_publish, + pp_proxy_tensors=pp_proxy_tensors, + ) + if batch.forward_mode.is_extend() or batch.is_extend_in_batch: self._verify_planner.note_non_decode_step() self._observers.note_prefill_step() - return self._forward_prefill(batch, on_publish) + return self._forward_prefill(batch, on_publish, pp_proxy_tensors) return self._forward_decode(batch, on_publish, grammar_barrier) + def _forward_lifecycle_only_prefill( + self, *, batch: ScheduleBatch, on_publish, pp_proxy_tensors + ) -> GenerationBatchResult: + if batch.forward_mode.is_idle(): + return self._forward_idle_prefill( + batch=batch, + on_publish=on_publish, + pp_proxy_tensors=pp_proxy_tensors, + capture_hidden_mode=CaptureHiddenMode.NULL, + ) + if not (batch.forward_mode.is_extend() or batch.is_extend_in_batch): + raise RuntimeError( + "Lifecycle-only DSpark worker only supports prefill batches." + ) + + batch_output = self.target_worker.forward_batch_generation( + batch, + pp_proxy_tensors=pp_proxy_tensors, + capture_hidden_mode=CaptureHiddenMode.NULL, + ) + batch_output.new_seq_lens = batch.seq_lens + if on_publish is not None: + on_publish(batch_output.new_seq_lens) + if batch_output.logits_output is not None: + batch_output.logits_output.hidden_states = None + return batch_output + def _forward_prefill( - self, batch: ScheduleBatch, on_publish + self, batch: ScheduleBatch, on_publish, pp_proxy_tensors=None ) -> GenerationBatchResult: if batch.forward_mode.is_idle(): - if get_parallel().enable_dp_attention: - self.target_worker.forward_batch_generation( - batch, capture_hidden_mode=CaptureHiddenMode.FULL - ) - return self._decode_idle_result(on_publish=on_publish) + return self._forward_idle_prefill( + batch=batch, + on_publish=on_publish, + pp_proxy_tensors=pp_proxy_tensors, + capture_hidden_mode=CaptureHiddenMode.FULL, + ) batch_output = self.target_worker.forward_batch_generation( - batch, capture_hidden_mode=CaptureHiddenMode.FULL + batch, + pp_proxy_tensors=pp_proxy_tensors, + capture_hidden_mode=CaptureHiddenMode.FULL, ) logits_output = batch_output.logits_output + output_pp_proxy_tensors = batch_output.pp_hidden_states_proxy_tensors + target_hidden = ( + logits_output.hidden_states + if logits_output is not None + else ( + output_pp_proxy_tensors.tensors.get("dspark_aux_hidden_states") + if output_pp_proxy_tensors is not None + else None + ) + ) next_token_ids = batch_output.next_token_ids batch_output.new_seq_lens = batch.seq_lens if on_publish is not None: on_publish(batch_output.new_seq_lens) - if logits_output.hidden_states is None: + if target_hidden is None and self.ps.pp_size <= 1: raise RuntimeError( "DSpark requires target aux hidden capture for prefill, but got None. " "Make sure the target model has DFlash layers-to-capture configured." ) + has_local_target_hidden = ( + target_hidden is not None and target_hidden.numel() > 0 + ) if batch.extend_lens is None or batch.prefix_lens is None: raise RuntimeError( "DSpark expected extend_lens / prefix_lens in extend mode, got None." @@ -432,7 +664,10 @@ def _forward_prefill( # Must inject before prefill returns: the scheduler may update radix # afterward, invalidating out_cache_loc. - device = next_token_ids.device + # Non-last PP ranks return PPProxyTensors and do not sample tokens, so + # next_token_ids is None there. Positions are still needed for local + # draft-KV injection, and should live on this worker's device. + device = self.device ctx_lens = torch.tensor(batch.extend_lens, dtype=torch.int32, device=device) draft_seq_lens = torch.tensor( batch.prefix_lens, dtype=torch.int32, device=device @@ -443,20 +678,106 @@ def _forward_prefill( ctx_lens, int(sum(batch.extend_lens)), ) - self._kv_injector.inject_target_hidden( - target_hidden=logits_output.hidden_states, - cache_loc=batch.out_cache_loc, - positions=positions, - ) + state_slot = final_pos = None + if is_unified_kv_triton(): + repeats = ctx_lens.to(torch.int64) + state_slot = torch.repeat_interleave( + batch.req_pool_indices.to(device=device, dtype=torch.int64), repeats + ) + final_pos = torch.repeat_interleave( + (draft_seq_lens + ctx_lens - 1).to(torch.int64), repeats + ) + if self._use_full_projection_prefill: + if not has_local_target_hidden or output_pp_proxy_tensors is not None: + raise RuntimeError( + "Single-owner DSpark prefill requires target hidden states " + "on the final PP rank." + ) + self._kv_injector.inject_target_hidden( + target_hidden=target_hidden, + cache_loc=batch.out_cache_loc, + positions=positions, + state_slot=state_slot, + final_pos=final_pos, + ) + else: + incoming_ctx = ( + pp_proxy_tensors.tensors.get("dspark_ctx_acc") + if pp_proxy_tensors is not None + else None + ) + local_ctx = None + if ( + self.ps.pp_size > 1 + and has_local_target_hidden + and self._pp_context_feature_indices is not None + ): + local_ctx = self.draft_model.project_target_hidden_partial( + target_hidden, self._pp_context_feature_indices + ) + if output_pp_proxy_tensors is not None: + output_pp_proxy_tensors.tensors.pop("dspark_aux_hidden_states", None) + ctx_acc = None + if incoming_ctx is not None and local_ctx is not None: + ctx_acc = ( + incoming_ctx.to(device=local_ctx.device, dtype=local_ctx.dtype) + + local_ctx + ) + elif incoming_ctx is not None: + ctx_acc = incoming_ctx.to(device=self.device, non_blocking=True) + elif local_ctx is not None: + ctx_acc = local_ctx + + if output_pp_proxy_tensors is not None: + if ctx_acc is not None: + output_pp_proxy_tensors.tensors["dspark_ctx_acc"] = ctx_acc + elif ctx_acc is not None: + self._kv_injector.inject_projected_context( + projected_context=ctx_acc, + cache_loc=batch.out_cache_loc, + positions=positions, + state_slot=state_slot, + final_pos=final_pos, + ) + elif has_local_target_hidden and not ( + self.ps.pp_size > 1 and not self._draft_is_moe + ): + self._kv_injector.inject_target_hidden( + target_hidden=target_hidden, + cache_loc=batch.out_cache_loc, + positions=positions, + state_slot=state_slot, + final_pos=final_pos, + ) # Avoid copying large hidden-state buffers to CPU in overlap scheduling. - logits_output.hidden_states = None + if logits_output is not None: + logits_output.hidden_states = None - batch_output.next_draft_input = make_next_draft_input( - bonus_tokens=next_token_ids, - new_seq_lens=batch.seq_lens, - ) + if next_token_ids is not None: + batch_output.next_draft_input = make_next_draft_input( + bonus_tokens=next_token_ids, + new_seq_lens=batch.seq_lens, + ) return batch_output + def _forward_idle_prefill( + self, + *, + batch: ScheduleBatch, + on_publish, + pp_proxy_tensors, + capture_hidden_mode: CaptureHiddenMode, + ) -> GenerationBatchResult: + if get_parallel().enable_dp_attention: + batch_output = self.target_worker.forward_batch_generation( + batch, + pp_proxy_tensors=pp_proxy_tensors, + capture_hidden_mode=capture_hidden_mode, + ) + if self._is_context_only_pp_prefill_rank: + return batch_output + return self._decode_idle_result(on_publish=on_publish) + def _idle_verify_ragged_layout(self, batch: ScheduleBatch): if batch.global_num_tokens is None or not self._verify_planner.is_compact_mode: return None @@ -520,7 +841,8 @@ def _forward_decode( self._observers.note_idle_decode_step() if get_parallel().enable_dp_attention: if self._draft_is_moe: - self._proposer.run_idle_participation(batch) + with self._draft_context(): + self._proposer.run_idle_participation(batch) self._verify_executor.run_idle_participation( batch=batch, idle_layout=self._idle_verify_ragged_layout(batch) ) @@ -785,4 +1107,6 @@ def _commit_target_mamba_states_after_verify( ) def get_confidence_budget_prepare(self): + if self._is_lifecycle_only_pp_prefill_rank: + return None return self._verify_planner.confidence_budget_prepare() diff --git a/test/registered/cp/test_deepseek_v4_flash_fp4_b200_cp.py b/test/registered/cp/test_deepseek_v4_flash_fp4_b200_cp.py index cfac77f383fd..748eb75ae0ec 100644 --- a/test/registered/cp/test_deepseek_v4_flash_fp4_b200_cp.py +++ b/test/registered/cp/test_deepseek_v4_flash_fp4_b200_cp.py @@ -195,5 +195,49 @@ def tearDownClass(cls): kill_process_tree(cls.process.pid) +# DSPARK draft is bundled with the -DSpark checkpoint. +DSPARK_MODEL = "deepseek-ai/DeepSeek-V4-Flash-DSpark" + + +class TestDSV4FlashFP4B200_CP_DSpark( + BasicDecodeCorrectnessMixin, + GSM8KMixin, + CustomTestCase, +): + """DSPARK speculation + prefill CP (interleave, CP_V2, attn_cp=tp).""" + + gsm8k_accuracy_thres = 0.90 + + @classmethod + def setUpClass(cls): + cls.model = try_cached_model(DSPARK_MODEL) + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=SERVER_LAUNCH_TIMEOUT, + other_args=[ + "--trust-remote-code", + "--tp", + "4", + "--attn-cp-size", + "4", + "--speculative-algorithm", + "DSPARK", + "--enable-prefill-cp", + "--cp-strategy", + "interleave", + "--moe-runner-backend", # for fp4 checkpoint + "flashinfer_mxfp4", + ], + env={"SGLANG_ENABLE_CP_V2": "1"}, + ) + + @classmethod + def tearDownClass(cls): + if hasattr(cls, "process") and cls.process: + kill_process_tree(cls.process.pid) + + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/disaggregation/test_disaggregation_wire.py b/test/registered/unit/disaggregation/test_disaggregation_wire.py index b96982b0c8b0..ce33d9ce15ed 100644 --- a/test/registered/unit/disaggregation/test_disaggregation_wire.py +++ b/test/registered/unit/disaggregation/test_disaggregation_wire.py @@ -15,7 +15,10 @@ from sglang.srt.disaggregation.utils import ( MetadataBuffers, get_dsv4_c128_state_indices, + pack_state_types, + resolve_state_component_dst_index, setup_state_kv_args, + unpack_state_types, ) from sglang.srt.managers.overlap_utils import FutureMap, RelayPayload from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool @@ -58,6 +61,24 @@ def test_list_of_buffers_roundtrip(self): bufs = [b"abc", b"", b"de", b"x" * 17] self.assertEqual(unpack_list_of_buffers(pack_list_of_buffers(bufs)), bufs) + def test_state_component_matching_uses_type_occurrence(self): + src_state_types = [StateType.SWA, StateType.SWA] + dst_state_types = [StateType.SWA, StateType.C128_STATE, StateType.SWA] + + self.assertEqual( + resolve_state_component_dst_index(src_state_types, dst_state_types, 0), + 0, + ) + self.assertEqual( + resolve_state_component_dst_index(src_state_types, dst_state_types, 1), + 2, + ) + + def test_state_types_roundtrip(self): + state_types = [StateType.SWA, StateType.C128_STATE, StateType.SWA_RING] + + self.assertEqual(unpack_state_types(pack_state_types(state_types)), state_types) + class TestGroupConcurrentContiguous(unittest.TestCase): @staticmethod diff --git a/test/registered/unit/disaggregation/test_pp_pd_consensus.py b/test/registered/unit/disaggregation/test_pp_pd_consensus.py new file mode 100644 index 000000000000..226127b31d9d --- /dev/null +++ b/test/registered/unit/disaggregation/test_pp_pd_consensus.py @@ -0,0 +1,296 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import Mock, patch + +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel + +maybe_stub_sgl_kernel() + +from sglang.srt.disaggregation.base import KVPoll # noqa: E402 +from sglang.srt.disaggregation.prefill import ( # noqa: E402 + PrefillBootstrapQueue, + SchedulerDisaggregationPrefillMixin, +) +from sglang.srt.disaggregation.utils import ( # noqa: E402 + _DRAFT_KV_LAYER_ID_BASE, + build_transfer_entry_pairs, +) +from sglang.srt.managers.schedule_batch import FINISH_ABORT # noqa: E402 +from sglang.srt.managers.scheduler_pp_mixin import ( # noqa: E402 + _pp_merge_transfer_status, +) + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +class TestPPPDConsensus(CustomTestCase): + @staticmethod + def _make_prefill_queue(pp_rank): + class FakeKVArgs: + pass + + class FakeManager: + def __init__(self, kv_args, *_args): + self.kv_args = kv_args + + class FakePool: + start_layer = 4 + end_layer = 6 + head_num = 1 + page_size = 64 + + @staticmethod + def get_contiguous_buf_infos(): + return [10, 11], [100, 100], [10, 10] + + class FakeDraftPool: + start_layer = 0 + end_layer = 1 + + @staticmethod + def get_contiguous_buf_infos(): + return [20], [100], [10] + + queue = PrefillBootstrapQueue.__new__(PrefillBootstrapQueue) + queue.transfer_backend = "mooncake" + queue.tp_rank = 0 + queue.pp_rank = pp_rank + queue.pp_size = 2 + queue.scheduler = SimpleNamespace( + ps=SimpleNamespace(dp_rank=0, gpu_id=0), + server_args=SimpleNamespace(disaggregation_ib_device=None), + tp_worker=SimpleNamespace( + model_runner=SimpleNamespace(kv_cache_dtype_str="auto") + ), + model_config=SimpleNamespace( + num_hidden_layers=8, + get_total_num_kv_heads=lambda: 1, + ), + req_to_token_pool=None, + ) + queue.token_to_kv_pool = FakePool() + queue.draft_token_to_kv_pool = FakeDraftPool() + queue.metadata_buffers = SimpleNamespace(get_buf_infos=lambda: ([], [], [])) + queue.is_mla_backend = False + return queue, FakeKVArgs, FakeManager + + def test_transfer_failure_overrides_ordered_success_intersection(self): + """A failure on one PP rank must terminate an otherwise successful rid.""" + status = _pp_merge_transfer_status( + previous=(["req-a", "req-b", "req-c"], ["req-x"]), + current=(["req-c", "req-a", "req-b"], ["req-b", "req-y"]), + ) + + self.assertEqual( + status, + (["req-a", "req-c"], ["req-x", "req-b", "req-y"]), + ) + + def test_bootstrap_probe_respects_local_metadata_credit_prefix(self): + """A slower PP rank must not advertise requests it cannot admit.""" + queue = PrefillBootstrapQueue.__new__(PrefillBootstrapQueue) + queue.queue = [ + SimpleNamespace( + rid="req-failed", + metadata_buffer_index=-1, + disagg_kv_sender=object(), + ), + SimpleNamespace( + rid="req-ready", + metadata_buffer_index=-1, + disagg_kv_sender=object(), + ), + SimpleNamespace( + rid="req-blocked", + metadata_buffer_index=-1, + disagg_kv_sender=object(), + ), + ] + queue.scheduler = SimpleNamespace( + attn_cp_cpu_group=object(), + attn_tp_cpu_group=object(), + ) + queue.req_to_metadata_buffer_idx_allocator = SimpleNamespace( + available_size=lambda: 1 + ) + + with patch( + "sglang.srt.disaggregation.prefill." "poll_and_all_reduce_attn_cp_tp_group", + return_value=[ + KVPoll.Failed, + KVPoll.WaitingForInput, + KVPoll.WaitingForInput, + ], + ): + good_rids, failed_rids = queue.get_ready_bootstrapped_rids_for_pp() + + self.assertEqual(good_rids, ["req-ready"]) + self.assertEqual(failed_rids, ["req-failed"]) + self.assertEqual( + [req.metadata_buffer_index for req in queue.queue], + [-1, -1, -1], + ) + + def test_bootstrap_probe_reports_failures_after_metadata_backpressure(self): + """Admission backpressure must not hide terminal failures later in FIFO.""" + queue = PrefillBootstrapQueue.__new__(PrefillBootstrapQueue) + queue.queue = [ + SimpleNamespace( + rid="req-blocked", + metadata_buffer_index=-1, + disagg_kv_sender=object(), + ), + SimpleNamespace( + rid="req-failed", + metadata_buffer_index=-1, + disagg_kv_sender=object(), + ), + SimpleNamespace( + rid="req-ready-after-block", + metadata_buffer_index=0, + disagg_kv_sender=object(), + ), + ] + queue.scheduler = SimpleNamespace( + attn_cp_cpu_group=object(), + attn_tp_cpu_group=object(), + ) + queue.req_to_metadata_buffer_idx_allocator = SimpleNamespace( + available_size=lambda: 0 + ) + + with patch( + "sglang.srt.disaggregation.prefill." "poll_and_all_reduce_attn_cp_tp_group", + return_value=[ + KVPoll.WaitingForInput, + KVPoll.Failed, + KVPoll.WaitingForInput, + ], + ): + good_rids, failed_rids = queue.get_ready_bootstrapped_rids_for_pp() + + self.assertEqual(good_rids, []) + self.assertEqual(failed_rids, ["req-failed"]) + + def test_remote_failure_waits_for_local_transfer_terminal_state(self): + """Do not release source KV while the local Mooncake worker is reading it.""" + sender = SimpleNamespace() + req = SimpleNamespace( + rid="req-race", + disagg_kv_sender=sender, + finished_reason=None, + pending_bootstrap=False, + return_logprob=False, + time_stats=SimpleNamespace(set_completion_time=Mock()), + ) + handle_failure = Mock() + scheduler = SimpleNamespace( + disagg_prefill_inflight_queue=[req], + attn_cp_cpu_group=object(), + attn_tp_cpu_group=object(), + ps=SimpleNamespace(pp_rank=1), + handle_inflight_transfer_failure=handle_failure, + output_streamer=SimpleNamespace(stream_output=Mock()), + req_to_metadata_buffer_idx_allocator=object(), + ) + + def mark_abort(target_req, message, status_code): + del message, status_code + target_req.finished_reason = FINISH_ABORT("remote PP failure") + + with ( + patch( + "sglang.srt.disaggregation.prefill." + "poll_and_all_reduce_attn_cp_tp_group", + return_value=[KVPoll.Transferring], + ), + patch( + "sglang.srt.disaggregation.prefill.prepare_abort", + side_effect=mark_abort, + ), + ): + done_reqs = SchedulerDisaggregationPrefillMixin.process_disagg_prefill_inflight_queue( + scheduler, + transfer_status=([], ["req-race"]), + ) + + self.assertEqual(done_reqs, []) + self.assertEqual(scheduler.disagg_prefill_inflight_queue, [req]) + handle_failure.assert_not_called() + + with ( + patch( + "sglang.srt.disaggregation.prefill." + "poll_and_all_reduce_attn_cp_tp_group", + return_value=[KVPoll.Success], + ), + patch("sglang.srt.disaggregation.prefill.maybe_release_metadata_buffer"), + ): + done_reqs = SchedulerDisaggregationPrefillMixin.process_disagg_prefill_inflight_queue( + scheduler, + transfer_status=([], []), + ) + + self.assertEqual(done_reqs, [req]) + self.assertEqual(scheduler.disagg_prefill_inflight_queue, []) + handle_failure.assert_called_once_with(req) + + def test_only_last_pp_registers_draft_kv_for_transfer(self): + """Draft KV has one prefill owner even though every PP rank has a worker.""" + layer_ids_by_rank = [] + for pp_rank in (0, 1): + queue, fake_args, fake_manager = self._make_prefill_queue(pp_rank) + + def get_kv_class(_backend, class_type): + return fake_args if class_type.value == "kvargs" else fake_manager + + with ( + patch( + "sglang.srt.disaggregation.prefill.get_kv_class", + side_effect=get_kv_class, + ), + patch( + "sglang.srt.disaggregation.prefill.setup_state_kv_args", + ), + ): + manager = queue._init_kv_manager() + layer_ids_by_rank.append(manager.kv_args.kv_layer_ids) + + self.assertEqual(layer_ids_by_rank[0], [4, 5]) + self.assertEqual( + layer_ids_by_rank[1], + [4, 5, _DRAFT_KV_LAYER_ID_BASE], + ) + + def test_pp_prefill_entries_pair_with_pp1_decode_entries(self): + """Each PP source maps target layers plus the unique draft entry by id.""" + decode_layer_ids = [ + 0, + 1, + 2, + 3, + 4, + 5, + _DRAFT_KV_LAYER_ID_BASE, + ] + + rank_0_pairs = build_transfer_entry_pairs( + src_layer_ids=[0, 1, 2, 3], + dst_layer_ids=decode_layer_ids, + n_src=4, + n_dst=7, + ) + rank_1_pairs = build_transfer_entry_pairs( + src_layer_ids=[4, 5, _DRAFT_KV_LAYER_ID_BASE], + dst_layer_ids=decode_layer_ids, + n_src=3, + n_dst=7, + ) + + self.assertEqual(rank_0_pairs, [(0, 0), (1, 1), (2, 2), (3, 3)]) + self.assertEqual(rank_1_pairs, [(0, 4), (1, 5), (2, 6)]) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/spec/test_dspark_pp_context.py b/test/registered/unit/spec/test_dspark_pp_context.py new file mode 100644 index 000000000000..b8d62f653ee3 --- /dev/null +++ b/test/registered/unit/spec/test_dspark_pp_context.py @@ -0,0 +1,387 @@ +import os +import unittest +from types import SimpleNamespace +from unittest.mock import Mock, patch + +import torch + +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel + +maybe_stub_sgl_kernel() + +from sglang.srt.layers.quantization.fp8 import Fp8Config, Fp8LinearMethod # noqa: E402 +from sglang.srt.mem_cache.kv_cache_builder import get_draft_kv_pool # noqa: E402 +from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig # noqa: E402 +from sglang.srt.model_executor.runner.base_runner import ( # noqa: E402 + _allocate_decode_buffers, +) +from sglang.srt.model_executor.runner_utils.buffers import ( # noqa: E402 + DecodeInputBuffers, +) +from sglang.srt.models.deepseek_v4 import DeepseekV4ForCausalLM # noqa: E402 +from sglang.srt.models.deepseek_v4_dspark import ( # noqa: E402 + DeepseekV4ForCausalLMDSpark, + _BlockFp8LinearSlice, +) +from sglang.srt.models.dflash import DFlashDraftModel # noqa: E402 +from sglang.srt.speculative.dspark_components.dspark_config import ( # noqa: E402 + resolve_single_owner_pp_rank, + use_lifecycle_only_draft_model, +) +from sglang.srt.speculative.dspark_components.dspark_worker_v2 import ( # noqa: E402 + DSparkWorkerV2, + _is_context_only_pp_prefill_rank, +) + +register_cpu_ci(est_time=3, suite="base-a-test-cpu") + + +class _TupleLinear(torch.nn.Module): + def __init__(self, input_size: int, output_size: int): + super().__init__() + self.linear = torch.nn.Linear(input_size, output_size, bias=False) + + def forward(self, hidden_states: torch.Tensor): + return self.linear(hidden_states), None + + +def _make_deepseek_v4_dspark_projection_model( + *, hidden_size: int, num_target_features: int +) -> DeepseekV4ForCausalLMDSpark: + model = DeepseekV4ForCausalLMDSpark.__new__(DeepseekV4ForCausalLMDSpark) + torch.nn.Module.__init__(model) + model.config = SimpleNamespace(hidden_size=hidden_size) + model.num_target_features = num_target_features + stage = torch.nn.Module() + stage.main_proj = _TupleLinear(hidden_size * num_target_features, hidden_size) + stage.main_norm = torch.nn.RMSNorm(hidden_size, eps=1e-6) + model.stages = torch.nn.ModuleList([stage]) + model.markov_head = torch.nn.Identity() + model.confidence_head = torch.nn.Identity() + model.embed_tokens = None + model.lm_head = None + model.is_lifecycle_only = False + model._partial_feature_indices = () + model._partial_main_proj = None + return model + + +class TestDSparkPPContext(CustomTestCase): + def test_full_projection_fast_path_requires_final_pp_owner(self): + """Only final-rank ownership can bypass the ctx_acc handoff.""" + with patch.dict( + os.environ, + {"SGLANG_PP_LAYER_PARTITION": "6,5,6,5,6,5,5,5"}, + ): + self.assertEqual( + resolve_single_owner_pp_rank( + target_layer_ids=[40, 41, 42], + num_hidden_layers=43, + pp_size=8, + ), + 7, + ) + self.assertIsNone( + resolve_single_owner_pp_rank( + target_layer_ids=[35, 40], + num_hidden_layers=43, + pp_size=8, + ) + ) + + def test_lifecycle_only_draft_model_is_limited_to_non_owner_ranks(self): + common = dict( + disaggregation_mode="prefill", + pp_size=8, + target_layer_ids=[40, 41, 42], + num_hidden_layers=43, + ) + for pp_rank in range(7): + self.assertTrue(use_lifecycle_only_draft_model(pp_rank=pp_rank, **common)) + self.assertFalse(use_lifecycle_only_draft_model(pp_rank=7, **common)) + self.assertFalse( + use_lifecycle_only_draft_model( + pp_rank=0, + **{**common, "target_layer_ids": [35, 40]}, + ) + ) + + def test_lifecycle_only_model_does_not_consume_checkpoint_weights(self): + model = DeepseekV4ForCausalLMDSpark.__new__(DeepseekV4ForCausalLMDSpark) + torch.nn.Module.__init__(model) + model.is_lifecycle_only = True + + def weights(): + raise AssertionError("lifecycle-only model consumed checkpoint weights") + yield + + model.load_weights(weights()) + + def test_block_fp8_projection_slice_selects_matching_weight_and_scale_blocks(self): + feature_width = 128 + output_size = 2 + quant_method = Fp8LinearMethod( + Fp8Config( + is_checkpoint_fp8_serialized=True, + activation_scheme="dynamic", + weight_block_size=[128, 128], + ) + ) + weight = torch.arange( + output_size * feature_width * 3, dtype=torch.float32 + ).reshape(output_size, feature_width * 3) + weight_scale = torch.tensor([[11.0, 22.0, 33.0]]) + source = SimpleNamespace( + quant_method=quant_method, + weight=torch.nn.Parameter(weight, requires_grad=False), + weight_scale_inv=torch.nn.Parameter(weight_scale, requires_grad=False), + ) + source.weight_scale_inv.format_ue8m0 = False + + projection_slice = _BlockFp8LinearSlice( + source=source, + feature_indices=[0, 2], + feature_width=feature_width, + ) + + expected_weight = torch.cat( + [weight[:, :feature_width], weight[:, 2 * feature_width :]], dim=1 + ) + self.assertTrue(torch.equal(projection_slice.weight, expected_weight)) + self.assertTrue( + torch.equal( + projection_slice.weight_scale_inv, + torch.tensor([[11.0, 33.0]]), + ) + ) + + def test_partial_projection_uses_prepared_quantized_slice(self): + model = _make_deepseek_v4_dspark_projection_model( + hidden_size=4, num_target_features=3 + ) + expected = torch.randn(2, 4) + projection_slice = Mock(return_value=expected) + model._partial_feature_indices = (1,) + model._partial_main_proj = projection_slice + local_hidden = torch.randn(2, 4) + + actual = model.project_target_hidden_partial(local_hidden, [1]) + + projection_slice.assert_called_once_with(local_hidden) + self.assertIs(actual, expected) + + def test_pp_spec_verify_buffers_use_token_axis(self): + """PP verify buffers must cover bs times speculative token width.""" + max_bs = 64 + num_tokens_per_req = 6 + max_num_token = max_bs * num_tokens_per_req + common_kwargs = dict( + device=torch.device("cpu"), + max_bs=max_bs, + max_num_token=max_num_token, + hidden_size=4, + dtype=torch.float32, + dp_size=1, + pp_size=8, + is_encoder_decoder=False, + require_mlp_tp_gather=False, + seq_len_fill_value=1, + encoder_len_fill_value=0, + num_tokens_per_req=num_tokens_per_req, + cache_loc_dtype=torch.int64, + enable_mamba_track=False, + hc_hidden_size=16, + ) + eager_buffers = _allocate_decode_buffers(vocab_size=8, **common_kwargs) + graph_buffers = DecodeInputBuffers.create( + next_token_logits_buffer=torch.zeros((max_num_token, 8)), + **common_kwargs, + ) + + self.assertEqual( + eager_buffers.pp_proxy_tensors["hidden_states"].shape, + (max_num_token, 16), + ) + self.assertEqual( + graph_buffers.pp_proxy_tensors["hidden_states"].shape, + (max_num_token, 16), + ) + + def test_partial_projection_sum_matches_full_projection(self): + """PP partial projections must preserve the full pre-norm FC result.""" + torch.manual_seed(0) + hidden_size = 4 + model = DFlashDraftModel.__new__(DFlashDraftModel) + torch.nn.Module.__init__(model) + model.config = SimpleNamespace(hidden_size=hidden_size) + model.num_context_features = 3 + model.fc = torch.nn.Linear(3 * hidden_size, hidden_size, bias=False) + model.hidden_norm = torch.nn.RMSNorm(hidden_size, eps=1e-6) + + feature_hidden = [ + torch.randn(5, hidden_size, dtype=torch.float32) for _ in range(3) + ] + full_hidden = torch.cat(feature_hidden, dim=-1) + full_projected = model.project_target_hidden(full_hidden) + + stage_0 = model.project_target_hidden_partial( + torch.cat([feature_hidden[0], feature_hidden[2]], dim=-1), + [0, 2], + ) + stage_1 = model.project_target_hidden_partial(feature_hidden[1], [1]) + pp_projected = model.hidden_norm(stage_0 + stage_1) + + torch.testing.assert_close(pp_projected, full_projected) + + def test_deepseek_v4_partial_projection_survives_context_only_pruning(self): + """Cross-rank contributions must retain full projection math after pruning.""" + torch.manual_seed(1) + hidden_size = 4 + model = _make_deepseek_v4_dspark_projection_model( + hidden_size=hidden_size, num_target_features=3 + ) + features = [torch.randn(5, hidden_size, dtype=torch.float32) for _ in range(3)] + full_projected = model.project_target_hidden(torch.cat(features, dim=-1)) + + stage_0 = model.project_target_hidden_partial( + torch.cat([features[0], features[2]], dim=-1), + [0, 2], + ) + model.prune_to_ctx_projection() + stage_1 = model.project_target_hidden_partial(features[1], [1]) + pp_projected = model.stages[0].main_norm(stage_0 + stage_1) + write_context_hidden_kv = Mock() + model._write_context_hidden_kv = write_context_hidden_kv + model.write_projected_context_kv( + projected_context=stage_0 + stage_1, + swa_loc=torch.arange(5), + positions=torch.arange(5), + pool=object(), + ) + + self.assertEqual(list(model.stages[0]._modules), ["main_proj", "main_norm"]) + torch.testing.assert_close(pp_projected, full_projected) + torch.testing.assert_close( + write_context_hidden_kv.call_args.kwargs["main_x"], + full_projected, + ) + + def test_deepseek_v4_capture_is_local_to_each_pp_rank(self): + """A target feature on a non-last rank must not disappear from ctx_acc.""" + model = DeepseekV4ForCausalLM.__new__(DeepseekV4ForCausalLM) + torch.nn.Module.__init__(model) + model.pp_group = SimpleNamespace(is_last_rank=False) + model.model = SimpleNamespace( + start_layer=10, + end_layer=20, + dspark_layers_to_capture=None, + ) + model.capture_aux_hidden_states = False + + model.set_dspark_layers_to_capture([5, 12, 18, 25]) + + self.assertTrue(model.capture_aux_hidden_states) + self.assertEqual(model.model.dspark_layers_to_capture, [12, 18]) + + def test_non_last_pp_prefill_does_not_require_target_lm_head(self): + """PP0-PP(N-2) must initialize without the last rank's lm_head.""" + self.assertTrue( + _is_context_only_pp_prefill_rank( + disaggregation_mode="prefill", + pp_rank=0, + pp_size=8, + ) + ) + self.assertFalse( + _is_context_only_pp_prefill_rank( + disaggregation_mode="prefill", + pp_rank=7, + pp_size=8, + ) + ) + + def test_context_only_rank_does_not_require_draft_attention_backend(self): + """Projection-only ranks must initialize overlap state without draft attention.""" + target_backend = object() + worker = DSparkWorkerV2.__new__(DSparkWorkerV2) + worker._is_context_only_pp_prefill_rank = True + worker._target_worker = SimpleNamespace( + model_runner=SimpleNamespace(attn_backend=target_backend) + ) + worker.draft_model_runner = SimpleNamespace() + + self.assertEqual(worker.spec_v2_attn_backends, (target_backend,)) + + def test_lifecycle_only_rank_does_not_allocate_draft_pool(self): + worker = DSparkWorkerV2.__new__(DSparkWorkerV2) + worker._draft_worker = Mock() + worker._is_lifecycle_only_pp_prefill_rank = True + + worker.alloc_memory_pool(memory_pool_config=Mock()) + + worker._draft_worker.alloc_memory_pool.assert_not_called() + + def test_lifecycle_only_rank_does_not_publish_draft_pool(self): + worker = SimpleNamespace(is_lifecycle_only_pp_prefill_rank=True) + spec_algorithm = SimpleNamespace( + is_ngram=lambda: False, + is_dspark=lambda: True, + ) + + self.assertIsNone( + get_draft_kv_pool( + draft_worker=worker, + spec_algorithm=spec_algorithm, + server_args=SimpleNamespace(enable_multi_layer_eagle=False), + ) + ) + + def test_non_last_pp_prefill_uses_minimal_draft_kv_pool(self): + """A context-only PP rank must not reserve the full draft KV capacity.""" + worker = DSparkWorkerV2.__new__(DSparkWorkerV2) + worker._draft_worker = Mock() + worker._is_pd_prefill = True + worker._draft_is_moe = True + worker._is_context_only_pp_prefill_rank = True + worker._is_lifecycle_only_pp_prefill_rank = False + worker.ps = SimpleNamespace(pp_rank=0, pp_size=2) + worker.page_size = 64 + full_config = MemoryPoolConfig( + max_total_num_tokens=4096, + max_running_requests=32, + ) + + worker.alloc_memory_pool(memory_pool_config=full_config) + + passed_config = worker._draft_worker.alloc_memory_pool.call_args.kwargs[ + "memory_pool_config" + ] + self.assertEqual(passed_config.max_total_num_tokens, 64) + self.assertEqual(passed_config.max_running_requests, 32) + self.assertEqual(full_config.max_total_num_tokens, 4096) + + def test_last_pp_prefill_keeps_full_draft_kv_pool(self): + worker = DSparkWorkerV2.__new__(DSparkWorkerV2) + worker._draft_worker = Mock() + worker._is_pd_prefill = True + worker._draft_is_moe = True + worker._is_context_only_pp_prefill_rank = False + worker._is_lifecycle_only_pp_prefill_rank = False + worker.ps = SimpleNamespace(pp_rank=1, pp_size=2) + worker.page_size = 64 + full_config = MemoryPoolConfig( + max_total_num_tokens=4096, + max_running_requests=32, + ) + + worker.alloc_memory_pool(memory_pool_config=full_config) + + passed_config = worker._draft_worker.alloc_memory_pool.call_args.kwargs[ + "memory_pool_config" + ] + self.assertIs(passed_config, full_config) + + +if __name__ == "__main__": + unittest.main() From 73ed7e6bf74e9dc4207381068a5478785e49207e Mon Sep 17 00:00:00 2001 From: Qi Sun <65223714+shiyu7@users.noreply.github.com> Date: Fri, 21 Aug 2026 17:38:39 +0800 Subject: [PATCH 27/47] [DSV4][DSpark] Budget draft KV only on its PP owner (#730) --- .../srt/model_executor/pool_configurator.py | 22 ++++++++- .../test_dsv4_pool_configurator.py | 47 +++++++++++++++++++ 2 files changed, 68 insertions(+), 1 deletion(-) create mode 100644 test/registered/unit/model_executor/test_dsv4_pool_configurator.py diff --git a/python/sglang/srt/model_executor/pool_configurator.py b/python/sglang/srt/model_executor/pool_configurator.py index 1fc00f5403e1..6b95f535478f 100644 --- a/python/sglang/srt/model_executor/pool_configurator.py +++ b/python/sglang/srt/model_executor/pool_configurator.py @@ -667,7 +667,7 @@ def __init__(self, kvc: KVCacheConfigurator): self.num_layers_ca128 = sum(1 for r in self.compression_ratios if r == 128) self.bytes_per_full_token = self._get_bytes_per_full_token() - if self.is_speculative: + if self.is_speculative and self._should_budget_draft_pool(kvc): # Reserve memory for the speculative draft worker by inflating # per-token bytes by (target+draft)/target. Equivalent to dflash's # scale_kv_cell_size_per_token_for_dflash but applied to @@ -706,6 +706,26 @@ def __init__(self, kvc: KVCacheConfigurator): "DSV4 compressed attention: online c128 enabled (ring_size=1)" ) + @staticmethod + def _should_budget_draft_pool(kvc: KVCacheConfigurator) -> bool: + """Return whether this PP stage owns the speculative draft KV pool. + + DSV4 DSpark PD prefill keeps the full draft SWA pool on the final PP + stage, matching ``prepare_dspark_hicache_draft_plan``. Earlier stages + either have no draft pool or retain only a one-page lifecycle pool, so + charging the full draft geometry there understates the shared PP token + capacity. Other speculative layouts keep the existing budgeting + behavior. + """ + + if not kvc.spec_algorithm.is_dspark(): + return True + return not ( + kvc.server_args.disaggregation_mode == "prefill" + and kvc.ps.pp_size > 1 + and kvc.ps.pp_rank != kvc.ps.pp_size - 1 + ) + def _get_bytes_per_full_token(self) -> float: kv_bytes = self.qk_nope_head_dim + self.qk_rope_head_dim * 2 + 8 diff --git a/test/registered/unit/model_executor/test_dsv4_pool_configurator.py b/test/registered/unit/model_executor/test_dsv4_pool_configurator.py new file mode 100644 index 000000000000..b036b0c48f18 --- /dev/null +++ b/test/registered/unit/model_executor/test_dsv4_pool_configurator.py @@ -0,0 +1,47 @@ +"""CPU-only tests for the DSV4 memory-pool configurator.""" + +import unittest +from types import SimpleNamespace + +from sglang.srt.model_executor.pool_configurator import DSV4PoolConfigurator +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +class TestDSV4DSparkDraftPoolBudget(CustomTestCase): + @staticmethod + def _make_kvc(*, pp_rank, pp_size=4, is_dspark=True, mode="prefill"): + return SimpleNamespace( + spec_algorithm=SimpleNamespace(is_dspark=lambda: is_dspark), + server_args=SimpleNamespace(disaggregation_mode=mode), + ps=SimpleNamespace(pp_rank=pp_rank, pp_size=pp_size), + ) + + def test_pd_prefill_budgets_draft_pool_only_on_final_pp_stage(self): + should_budget = [ + DSV4PoolConfigurator._should_budget_draft_pool( + self._make_kvc(pp_rank=rank) + ) + for rank in range(4) + ] + + self.assertEqual(should_budget, [False, False, False, True]) + + def test_other_speculative_layouts_keep_existing_budgeting(self): + cases = ( + self._make_kvc(pp_rank=0, is_dspark=False), + self._make_kvc(pp_rank=0, mode="decode"), + self._make_kvc(pp_rank=0, pp_size=1), + ) + + for kvc in cases: + with self.subTest(kvc=kvc): + self.assertTrue( + DSV4PoolConfigurator._should_budget_draft_pool(kvc) + ) + + +if __name__ == "__main__": + unittest.main() From dab2bccd1c45dc6fe0e41b06f783bdea0134e068 Mon Sep 17 00:00:00 2001 From: Qi Sun <65223714+shiyu7@users.noreply.github.com> Date: Fri, 21 Aug 2026 20:13:10 +0800 Subject: [PATCH 28/47] [DSpark][MoE] Support MegaMoE under DP attention (#728) --- .../sglang/srt/arg_groups/speculative_hook.py | 22 ++++++-- .../dspark/test_dspark_draft_path_default.py | 54 +++++++++++++++++++ 2 files changed, 73 insertions(+), 3 deletions(-) diff --git a/python/sglang/srt/arg_groups/speculative_hook.py b/python/sglang/srt/arg_groups/speculative_hook.py index 8ffbeec2bf38..d0314d90a09e 100644 --- a/python/sglang/srt/arg_groups/speculative_hook.py +++ b/python/sglang/srt/arg_groups/speculative_hook.py @@ -283,17 +283,33 @@ def _handle_dspark(server_args: ServerArgs) -> None: if server_args.enable_dp_attention and server_args.dp_size > 1: if not server_args.enable_dp_lm_head: raise ValueError("DSpark with dp attention requires --enable-dp-lm-head.") - supports_dspark_dp_moe = server_args.moe_a2a_backend == "none" or ( + supports_dspark_dp_moe = server_args.moe_a2a_backend in ( + "none", + "megamoe", + ) or ( server_args.moe_a2a_backend == "deepep" and server_args.moe_runner_backend == "deep_gemm" ) if not supports_dspark_dp_moe: raise ValueError( - "DSpark with dp attention only supports moe_a2a_backend='none' " - "or moe_a2a_backend='deepep' with moe_runner_backend='deep_gemm'; " + "DSpark with dp attention supports moe_a2a_backend='none', " + "moe_a2a_backend='megamoe', or moe_a2a_backend='deepep' with " + "moe_runner_backend='deep_gemm'; " f"got moe_a2a_backend={server_args.moe_a2a_backend!r}, " f"moe_runner_backend={server_args.moe_runner_backend!r}." ) + if server_args.moe_a2a_backend != "none": + from sglang.srt.speculative.ragged_verify import ( + RaggedVerifyMode, + read_ragged_verify_mode, + ) + + if read_ragged_verify_mode() is not RaggedVerifyMode.STATIC: + raise ValueError( + "DSpark with dp attention + " + f"moe_a2a_backend={server_args.moe_a2a_backend!r} requires " + "SGLANG_RAGGED_VERIFY_MODE=static." + ) if ( server_args.speculative_moe_a2a_backend is not None and server_args.speculative_moe_a2a_backend != server_args.moe_a2a_backend diff --git a/test/registered/spec/dspark/test_dspark_draft_path_default.py b/test/registered/spec/dspark/test_dspark_draft_path_default.py index b09ae0651435..24ec524c70b0 100644 --- a/test/registered/spec/dspark/test_dspark_draft_path_default.py +++ b/test/registered/spec/dspark/test_dspark_draft_path_default.py @@ -5,6 +5,7 @@ _handle_dspark, _target_checkpoint_bundles_dspark_draft, ) +from sglang.srt.environ import envs from sglang.srt.server_args import ServerArgs from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -84,5 +85,58 @@ def test_explicit_draft_path_is_not_overwritten(self): ) +class TestDsparkDpAttentionMoeA2aGate(CustomTestCase): + """Gate contract for DSpark + DP attention + MoE A2A backends.""" + + def _dp_server_args( + self, *, moe_a2a_backend: str, moe_runner_backend: str = "auto" + ) -> ServerArgs: + server_args = _make_dspark_server_args( + model_path=_BUNDLED_MODEL_PATH, hf_config=_bundled_hf_config() + ) + server_args.enable_dp_attention = True + server_args.enable_dp_lm_head = True + server_args.dp_size = 2 + server_args.tp_size = 2 + server_args.moe_a2a_backend = moe_a2a_backend + server_args.moe_runner_backend = moe_runner_backend + return server_args + + def test_supported_a2a_backends(self): + with envs.SGLANG_RAGGED_VERIFY_MODE.override("static"): + for backend, runner in ( + ("none", "auto"), + ("megamoe", "deep_gemm"), + ("deepep", "deep_gemm"), + ): + _handle_dspark( + self._dp_server_args( + moe_a2a_backend=backend, moe_runner_backend=runner + ) + ) + + for backend, runner in ( + ("deepep", "marlin"), + ("pplx", "deep_gemm"), + ): + with self.assertRaisesRegex(ValueError, backend): + _handle_dspark( + self._dp_server_args( + moe_a2a_backend=backend, moe_runner_backend=runner + ) + ) + + def test_a2a_backends_require_static_verify_mode(self): + with envs.SGLANG_RAGGED_VERIFY_MODE.override("compact"): + for backend in ("megamoe", "deepep"): + with self.assertRaisesRegex(ValueError, "static"): + _handle_dspark( + self._dp_server_args( + moe_a2a_backend=backend, + moe_runner_backend="deep_gemm", + ) + ) + + if __name__ == "__main__": unittest.main() From 53a782bd862c34e00b4070429c0f012d670039b1 Mon Sep 17 00:00:00 2001 From: Qi Sun <65223714+shiyu7@users.noreply.github.com> Date: Fri, 21 Aug 2026 20:27:31 +0800 Subject: [PATCH 29/47] [MoE][MXFP4] Shard weights before Marlin padding (#731) Co-authored-by: sunqi.7 --- python/sglang/srt/layers/quantization/mxfp4_marlin_moe.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/python/sglang/srt/layers/quantization/mxfp4_marlin_moe.py b/python/sglang/srt/layers/quantization/mxfp4_marlin_moe.py index 7826032efdd7..94212af7c4c9 100644 --- a/python/sglang/srt/layers/quantization/mxfp4_marlin_moe.py +++ b/python/sglang/srt/layers/quantization/mxfp4_marlin_moe.py @@ -69,7 +69,10 @@ def create_weights( layer._dsv4_mxfp4_backend = None # set in process_weights_after_loading fp4_block_k = 32 - intermediate_size_per_partition = round_up(intermediate_size_per_partition, 128) + # Keep the loader tensors at the logical TP shard size. Marlin tile + # padding is applied by prepare_moe_mxfp4_layer_for_marlin() after the + # checkpoint has been sharded. Padding here makes the generic loader + # use a padded TP stride and can index past the full checkpoint tensor. hidden_size = round_up(hidden_size, 256) self.hidden_pad = hidden_size - layer.hidden_size From 6ed5780c4fd3075edd8a76fe6f09af51e4e900a6 Mon Sep 17 00:00:00 2001 From: Qi Sun <65223714+shiyu7@users.noreply.github.com> Date: Sat, 22 Aug 2026 10:07:15 +0800 Subject: [PATCH 30/47] perf(moe): add Triton 3.6 H20 configs for DSV4 (#733) --- ...dtype=fp8_w8a8,block_shape=[128, 128].json | 146 ++++++++++++++++ ...=fp8_w8a8,block_shape=[128, 128]_down.json | 164 ++++++++++++++++++ 2 files changed, 310 insertions(+) create mode 100644 python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=256,N=512,device_name=NVIDIA_H20,dtype=fp8_w8a8,block_shape=[128, 128].json create mode 100644 python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=256,N=512,device_name=NVIDIA_H20,dtype=fp8_w8a8,block_shape=[128, 128]_down.json diff --git a/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=256,N=512,device_name=NVIDIA_H20,dtype=fp8_w8a8,block_shape=[128, 128].json b/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=256,N=512,device_name=NVIDIA_H20,dtype=fp8_w8a8,block_shape=[128, 128].json new file mode 100644 index 000000000000..77f260b3ddc7 --- /dev/null +++ b/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=256,N=512,device_name=NVIDIA_H20,dtype=fp8_w8a8,block_shape=[128, 128].json @@ -0,0 +1,146 @@ +{ + "1": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 5 + }, + "2": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 5 + }, + "4": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 16, + "num_warps": 4, + "num_stages": 4 + }, + "8": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 16, + "num_warps": 4, + "num_stages": 2 + }, + "16": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 4 + }, + "24": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 16, + "num_warps": 4, + "num_stages": 3 + }, + "32": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 3 + }, + "48": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 3 + }, + "64": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 4 + }, + "96": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 16, + "num_warps": 4, + "num_stages": 3 + }, + "128": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 16, + "num_warps": 4, + "num_stages": 3 + }, + "256": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 4 + }, + "512": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 16, + "num_warps": 4, + "num_stages": 3 + }, + "1024": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 3 + }, + "1536": { + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 64, + "num_warps": 4, + "num_stages": 3 + }, + "2048": { + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 3 + }, + "3072": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 4 + }, + "4096": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 16, + "num_warps": 4, + "num_stages": 4 + } +} diff --git a/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=256,N=512,device_name=NVIDIA_H20,dtype=fp8_w8a8,block_shape=[128, 128]_down.json b/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=256,N=512,device_name=NVIDIA_H20,dtype=fp8_w8a8,block_shape=[128, 128]_down.json new file mode 100644 index 000000000000..957c67d159f4 --- /dev/null +++ b/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=256,N=512,device_name=NVIDIA_H20,dtype=fp8_w8a8,block_shape=[128, 128]_down.json @@ -0,0 +1,164 @@ +{ + "1": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + }, + "2": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 2, + "USE_TMA": true + }, + "4": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 2, + "USE_TMA": true + }, + "8": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + }, + "16": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 2, + "USE_TMA": true + }, + "24": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 2, + "USE_TMA": true + }, + "32": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + }, + "48": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + }, + "64": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + }, + "96": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 2, + "USE_TMA": true + }, + "128": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 2, + "USE_TMA": true + }, + "256": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 16, + "num_warps": 4, + "num_stages": 2, + "USE_TMA": true + }, + "512": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 16, + "num_warps": 4, + "num_stages": 2, + "USE_TMA": true + }, + "1024": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 16, + "num_warps": 4, + "num_stages": 2, + "USE_TMA": true + }, + "1536": { + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 16, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + }, + "2048": { + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 2, + "USE_TMA": true + }, + "3072": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + }, + "4096": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + } +} From 3c78ea61878e249c5224c509a352aa1e0b3cd60d Mon Sep 17 00:00:00 2001 From: luoroger37 Date: Mon, 24 Aug 2026 11:29:05 +0800 Subject: [PATCH 31/47] [HiCache] Split the host-memory budget across co-located ranks (#732) --- .../sglang/srt/mem_cache/memory_pool_host.py | 15 +++---- python/sglang/srt/mem_cache/pool_host/base.py | 35 +++++++++++++++- python/sglang/srt/mem_cache/pool_host/mha.py | 6 +-- .../unit/mem_cache/test_mem_pool_host.py | 40 +++++++++++++++++++ .../test_minimax_sparse_pool_host_unit.py | 2 +- 5 files changed, 81 insertions(+), 17 deletions(-) diff --git a/python/sglang/srt/mem_cache/memory_pool_host.py b/python/sglang/srt/mem_cache/memory_pool_host.py index 2670de52c526..fbf4c2e1e460 100644 --- a/python/sglang/srt/mem_cache/memory_pool_host.py +++ b/python/sglang/srt/mem_cache/memory_pool_host.py @@ -10,7 +10,6 @@ from sglang.srt.mem_cache.pool_host.mla import MLATokenToKVPoolHost import numpy as np -import psutil import torch from sglang.kernels.ops.kvcache.hicache import ( @@ -52,7 +51,7 @@ from sglang.srt.mem_cache.pool_host import HostKVCache from sglang.srt.mem_cache.pool_host.base import ( _WRITE_BACK_STAGING_PAGE_CHUNK, - HICACHE_HOST_MEMORY_RESERVE_BYTES, + host_memory_budget_bytes, sync_fixed_hicache_size, synchronized, ) @@ -121,9 +120,8 @@ def __init__( device_pool.size, ) - host_mem = psutil.virtual_memory() requested_bytes = self.size * self.size_per_token - available_bytes = host_mem.available - HICACHE_HOST_MEMORY_RESERVE_BYTES + available_bytes = host_memory_budget_bytes() if requested_bytes > available_bytes: raise ValueError( f"Not enough host memory available. Requesting " @@ -760,8 +758,7 @@ def __init__( self.gpu_device = device_buffers[0].device if device_buffers else device requested_bytes = self.layer_num * num_host_pages * self.item_bytes - host_mem = psutil.virtual_memory() - available_bytes = host_mem.available - HICACHE_HOST_MEMORY_RESERVE_BYTES + available_bytes = host_memory_budget_bytes() if requested_bytes > available_bytes: raise ValueError( f"Not enough host memory for V4 paged pool {pool_name}. " @@ -1157,8 +1154,7 @@ def __init__( self.size_per_token = self.state_page_bytes requested_bytes = self.layer_num * num_host_pages * self.state_page_bytes - host_mem = psutil.virtual_memory() - available_bytes = host_mem.available - HICACHE_HOST_MEMORY_RESERVE_BYTES + available_bytes = host_memory_budget_bytes() if requested_bytes > available_bytes: raise ValueError( f"Not enough host memory for V4 state pool {pool_name}. " @@ -1755,8 +1751,7 @@ def __init__( buf_elem_size = self.page_num * self.layer_num * self.indexer_page_stride_size requested_bytes = buf_elem_size * self.indexer_dtype.itemsize - host_mem = psutil.virtual_memory() - available_bytes = host_mem.available - HICACHE_HOST_MEMORY_RESERVE_BYTES + available_bytes = host_memory_budget_bytes() if requested_bytes > available_bytes: raise ValueError( f"Not enough host memory for DSA indexer hierarchical cache. " diff --git a/python/sglang/srt/mem_cache/pool_host/base.py b/python/sglang/srt/mem_cache/pool_host/base.py index 57dc8c1a2a20..5e12f5f6510b 100644 --- a/python/sglang/srt/mem_cache/pool_host/base.py +++ b/python/sglang/srt/mem_cache/pool_host/base.py @@ -9,11 +9,13 @@ import psutil import torch +from sglang.srt.distributed.parallel_state import get_world_group from sglang.srt.mem_cache.memory_pool import KVCache from sglang.srt.mem_cache.pool_host.common import ( _cuda_host_unregister, get_allocator_from_storage, ) +from sglang.srt.runtime_context import get_parallel from sglang.srt.utils import is_cuda, is_hip logger = logging.getLogger(__name__) @@ -27,6 +29,36 @@ _WRITE_BACK_STAGING_PAGE_CHUNK = 64 +def ranks_per_host() -> int: + """Number of ranks of this job running on the same machine as this one. + + Derived as world_size // nnodes: the launcher slices ranks uniformly + across nodes (resolution asserts divisibility), so no hostname collective + is needed — a collective here would have to be issued the same number of + times on every rank, and ranks build different numbers of host pools. + """ + if not (torch.distributed.is_available() and torch.distributed.is_initialized()): + return 1 + try: + world_group = get_world_group() + except AssertionError: + return 1 + if world_group.world_size == 1: + return 1 + return max(world_group.world_size // get_parallel().nnodes, 1) + + +def host_memory_budget_bytes() -> int: + """Host RAM this rank may claim for a HiCache pool. + + psutil reports the whole machine, so co-located ranks each see the same free + memory; without the split every rank sizes its pool against all of it and + the host is oversubscribed by the number of ranks it holds. + """ + free = psutil.virtual_memory().available - HICACHE_HOST_MEMORY_RESERVE_BYTES + return free // ranks_per_host() + + def sync_fixed_hicache_size(size: int, host_size: int) -> int: """Sync fixed-size HiCache token capacity across PP ranks. @@ -139,9 +171,8 @@ def __init__( ) # Verify there is enough available host memory. - host_mem = psutil.virtual_memory() requested_bytes = self.size * self.size_per_token - available_bytes = host_mem.available - HICACHE_HOST_MEMORY_RESERVE_BYTES + available_bytes = host_memory_budget_bytes() if requested_bytes > available_bytes: raise ValueError( f"Not enough host memory available. Requesting " diff --git a/python/sglang/srt/mem_cache/pool_host/mha.py b/python/sglang/srt/mem_cache/pool_host/mha.py index 150c37b72628..3637fdb53ba4 100644 --- a/python/sglang/srt/mem_cache/pool_host/mha.py +++ b/python/sglang/srt/mem_cache/pool_host/mha.py @@ -3,7 +3,6 @@ import logging import threading -import psutil import torch from sglang.kernels.ops.kvcache.hicache import ( @@ -31,8 +30,8 @@ from sglang.srt.mem_cache.memory_pool import MHATokenToKOnlyPool, MHATokenToKVPool from sglang.srt.mem_cache.pool_host.base import ( _WRITE_BACK_STAGING_PAGE_CHUNK, - HICACHE_HOST_MEMORY_RESERVE_BYTES, HostKVCache, + host_memory_budget_bytes, ) from sglang.srt.mem_cache.pool_host.common import ( ALLOC_MEMORY_FUNCS, @@ -660,9 +659,8 @@ def __init__( self.page_num = anchor_host.page_num self.size_per_token = self.get_size_per_token() - host_mem = psutil.virtual_memory() requested_bytes = self.size * self.size_per_token - available_bytes = host_mem.available - HICACHE_HOST_MEMORY_RESERVE_BYTES + available_bytes = host_memory_budget_bytes() if requested_bytes > available_bytes: raise ValueError( f"Not enough host memory for MiniMax index-K hierarchical cache. " diff --git a/test/registered/unit/mem_cache/test_mem_pool_host.py b/test/registered/unit/mem_cache/test_mem_pool_host.py index ce94365d41bd..f038f8cd14cd 100644 --- a/test/registered/unit/mem_cache/test_mem_pool_host.py +++ b/test/registered/unit/mem_cache/test_mem_pool_host.py @@ -2,6 +2,7 @@ import threading import unittest +import unittest.mock import torch @@ -10,7 +11,9 @@ DeepSeekV4PagedHostPool, MambaPoolHost, ) +from sglang.srt.mem_cache.pool_host import base from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost +from sglang.srt.runtime_context import get_context from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -182,5 +185,42 @@ def test_deepseek_v4_pool_lazy_release(self): self.assertEqual(len(pool.alloc(1)), 2) +class TestHostMemoryBudget(CustomTestCase): + # Pinned so the two budget reads below see identical free memory; the real + # psutil value drifts between calls and would flake the equality checks. + _AVAILABLE = base.HICACHE_HOST_MEMORY_RESERVE_BYTES + 64 * (1024**3) + + def _budget_with_ranks(self, ranks): + # Deliberate single-accessor stub: isolates the budget math from the + # topology derivation, which the ranks_per_host case below covers. + fake_mem = unittest.mock.Mock(available=self._AVAILABLE) + with unittest.mock.patch.object( + base, "ranks_per_host", return_value=ranks + ), unittest.mock.patch.object( + base.psutil, "virtual_memory", return_value=fake_mem + ): + return base.host_memory_budget_bytes() + + def test_budget_is_split_across_co_located_ranks(self): + solo = self._budget_with_ranks(1) + self.assertEqual(self._budget_with_ranks(4), solo // 4) + + def test_reserve_is_taken_before_the_split(self): + # Each rank must not get its own copy of the reserve. + budget = self._budget_with_ranks(8) + self.assertLessEqual( + budget * 8, self._AVAILABLE - base.HICACHE_HOST_MEMORY_RESERVE_BYTES + ) + + def test_ranks_per_host_divides_world_size_by_nodes(self): + # The launcher slices ranks uniformly across nodes, so the co-located + # rank count is world_size // nnodes — no hostname collective. + fake_group = unittest.mock.Mock(world_size=16) + with get_context().override_server_args(nnodes=2), unittest.mock.patch.object( + torch.distributed, "is_initialized", return_value=True + ), unittest.mock.patch.object(base, "get_world_group", return_value=fake_group): + self.assertEqual(base.ranks_per_host(), 8) + + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/mem_cache/test_minimax_sparse_pool_host_unit.py b/test/registered/unit/mem_cache/test_minimax_sparse_pool_host_unit.py index 3bdf9df0faea..d0acd9b543e4 100644 --- a/test/registered/unit/mem_cache/test_minimax_sparse_pool_host_unit.py +++ b/test/registered/unit/mem_cache/test_minimax_sparse_pool_host_unit.py @@ -9,7 +9,7 @@ HybridCacheController, ) from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool -from sglang.srt.mem_cache.memory_pool_host import ( +from sglang.srt.mem_cache.pool_host.base import ( HICACHE_HOST_MEMORY_RESERVE_BYTES, ) from sglang.srt.mem_cache.pool_host.common import ( From 9b358684811a947f0a10dc8653edeef818ba512d Mon Sep 17 00:00:00 2001 From: Qi Sun <65223714+shiyu7@users.noreply.github.com> Date: Mon, 24 Aug 2026 11:30:22 +0800 Subject: [PATCH 32/47] [HiCache] Batch PP write/load completion synchronization (#729) Co-authored-by: luoroger37 --- .../srt/mem_cache/unified_radix_cache.py | 14 ++++- .../mem_cache/test_hiradix_pp_sync_drain.py | 57 ++++++++++++++++++- 2 files changed, 68 insertions(+), 3 deletions(-) diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index c11a02e30f68..e013fd6901a1 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -1937,8 +1937,18 @@ def check_hicache_events(self) -> None: self._drain_async_work() if self.pp_size != 1: - self.writing_check() - self.loading_check() + finish_counts = torch.zeros(2, dtype=torch.int, device="cpu") + if self.pp_rank == 0 and self.cache_controller is not None: + finish_counts[0] = self._count_ready_acks( + self.cache_controller.ack_write_queue + ) + finish_counts[1] = self._count_ready_acks( + self.cache_controller.ack_load_queue + ) + self._all_reduce(finish_counts, torch.distributed.ReduceOp.MIN) + write_finish_count, load_finish_count = map(int, finish_counts.tolist()) + self.writing_check(finish_count=write_finish_count) + self.loading_check(finish_count=load_finish_count) if self.enable_storage: self.drain_storage_control_queues() else: diff --git a/test/registered/unit/mem_cache/test_hiradix_pp_sync_drain.py b/test/registered/unit/mem_cache/test_hiradix_pp_sync_drain.py index 6e1ba43fed21..868aa03af89c 100644 --- a/test/registered/unit/mem_cache/test_hiradix_pp_sync_drain.py +++ b/test/registered/unit/mem_cache/test_hiradix_pp_sync_drain.py @@ -1,10 +1,13 @@ -"""Unit test for HiRadixCache._drain_async_work PP-sync backpressure.""" +"""Unit tests for HiCache PP synchronization.""" import unittest +from types import SimpleNamespace +from unittest.mock import MagicMock from sglang.srt.mem_cache.hiradix_cache import HiRadixCache from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase register_cpu_ci(est_time=1, suite="base-a-test-cpu") @@ -46,5 +49,57 @@ def test_drain_empty_is_noop(self): self.assertEqual(holder.work_list, []) +class TestUnifiedPPSyncBatching(CustomTestCase): + def _make_cache(self, pp_rank, write_ready, load_ready): + cache = object.__new__(UnifiedRadixCache) + cache.tree_core = SimpleNamespace(enable_storage=False) + cache.pp_rank = pp_rank + cache.pp_size = 2 + cache.enable_storage_metrics = False + cache.storage_metrics_collector = None + cache._drain_async_work = MagicMock() + cache._all_reduce = MagicMock() + cache.writing_check = MagicMock() + cache.loading_check = MagicMock() + cache.drain_storage_control_queues = MagicMock() + cache.cache_controller = SimpleNamespace( + ack_write_queue=[ + SimpleNamespace( + finish_event=SimpleNamespace(query=MagicMock(return_value=ready)) + ) + for ready in write_ready + ], + ack_load_queue=[ + SimpleNamespace( + finish_event=SimpleNamespace(query=MagicMock(return_value=ready)) + ) + for ready in load_ready + ], + ) + return cache + + def test_pp_batches_write_and_load_counts_once(self): + leader = self._make_cache(0, [True, False], [True, True]) + leader.check_hicache_events() + + leader._all_reduce.assert_called_once() + self.assertEqual(leader._all_reduce.call_args.args[0].tolist(), [1, 2]) + leader.writing_check.assert_called_once_with(finish_count=1) + leader.loading_check.assert_called_once_with(finish_count=2) + + follower = self._make_cache(1, [True], [True]) + follower._all_reduce.side_effect = lambda counts, _: counts.fill_(1) + follower.check_hicache_events() + + for queue in ( + follower.cache_controller.ack_write_queue, + follower.cache_controller.ack_load_queue, + ): + queue[0].finish_event.query.assert_not_called() + follower._all_reduce.assert_called_once() + follower.writing_check.assert_called_once_with(finish_count=1) + follower.loading_check.assert_called_once_with(finish_count=1) + + if __name__ == "__main__": unittest.main() From cbdec8111ea045e726c4f8e9d214174a7ebff40d Mon Sep 17 00:00:00 2001 From: luoroger37 Date: Mon, 24 Aug 2026 15:08:43 +0800 Subject: [PATCH 33/47] [HiCache] Merge upstream write-back duplicate reclamation fixes (#735) --- .../unified_cache/unified_tree_core.py | 186 +++++++++++++++++- .../unified_tree_core_interface.py | 9 + .../srt/mem_cache/unified_radix_cache.py | 28 ++- .../test_unified_radix_cache_unittest.py | 79 ++++++++ 4 files changed, 293 insertions(+), 9 deletions(-) diff --git a/python/sglang/srt/mem_cache/unified_cache/unified_tree_core.py b/python/sglang/srt/mem_cache/unified_cache/unified_tree_core.py index e0110067cb0e..6a819bf52c58 100644 --- a/python/sglang/srt/mem_cache/unified_cache/unified_tree_core.py +++ b/python/sglang/srt/mem_cache/unified_cache/unified_tree_core.py @@ -84,6 +84,11 @@ logger = logging.getLogger(__name__) +# 42 bits: digest * 1000003 (< 2^20) stays under 2^62, so the update never +# overflows int64 with plain (non-wrapping) arithmetic in the Rust port, and +# the TP consistency check can still all_reduce [digest, -digest] in int64. +_RECLAIM_DIGEST_MASK = (1 << 42) - 1 + class StorageBackupSpec(NamedTuple): """A node's device->storage backup spec, gathered tree-side.""" @@ -123,6 +128,9 @@ def __init__(self, tree_components: tuple[ComponentType, ...], priority: int = 0 self.id = UnifiedTreeNode.counter UnifiedTreeNode.counter += 1 self.write_through_pending_id: Optional[int] = None + # Anchor NodeId of an in-flight H->D load-back reading this node's + # host slots; such host copies must not be reclaimed until the ack. + self.load_back_pending_id: Optional[int] = None def component(self, component_type: ComponentType) -> ComponentData: return self.component_data[component_type] @@ -460,6 +468,11 @@ def reset(self) -> None: ) for ct in self.component_types } + # Full KV on both tiers -> redundant host copy, reclaimed first by + # write_back; insertion-ordered dict keeps victims TP-deterministic. + self.full_host_duplicates: dict[NodeId, UnifiedTreeNode] = {} + # Rolling digest of reclaim victim ids, cross-checked across TP ranks. + self.write_back_duplicate_reclaim_digest: int = 0 self._empty_match_result = MatchResult( device_indices=torch.empty( @@ -1016,6 +1029,8 @@ def _split_node( new_node.key = child.key[:split_len] new_node.hit_count = child.hit_count new_node.creation_time = child.creation_time + # Split fragments stay on the anchor's root path for the ack's walk. + new_node.load_back_pending_id = child.load_back_pending_id self._for_each_component_lru(child, UnifiedLRUList.remove_node) @@ -1051,6 +1066,8 @@ def _split_node( self._update_evictable_leaf_sets(new_node) self._update_evictable_leaf_sets(child) + # Only the new fragment needs qualifying; the child keeps its id. + self._update_duplicate_tracking(new_node) return new_node, action def _add_new_node( @@ -1086,6 +1103,8 @@ def _unevict_node_on_insert( cd.value = fresh_value.clone() self.component_evictable_size_[ct] += n self._update_evictable_leaf_sets(node) + # A backuped node restored from fresh KV is a duplicate right away. + self._update_duplicate_tracking(node) if node.parent is not None: self._update_evictable_leaf_sets(node.parent) self._record_store_event(node, medium=StorageMedium.GPU) @@ -1102,6 +1121,26 @@ def _update_evictable_leaf_sets(self, node: UnifiedTreeNode) -> None: else: self.evictable_host_leaves.discard(node) + def _update_duplicate_tracking(self, node: UnifiedTreeNode) -> None: + """Register where duplicates are born (acks, split, unevict); + deregistration is lazy, so entries may be stale and re-checked live.""" + if self._is_settled_full_host_duplicate(node): + self.full_host_duplicates.setdefault(node.id, node) + else: + self.full_host_duplicates.pop(node.id, None) + + def _is_settled_full_host_duplicate(self, node: UnifiedTreeNode) -> bool: + """Full KV present on both tiers with no in-flight DMA on the node's + host slots; mid-transfer nodes join the tracking at their ack.""" + cd = node.component_data[BASE_COMPONENT_TYPE] + return ( + node is not self.root_node + and cd.value is not None + and cd.host_value is not None + and node.write_through_pending_id is None + and node.load_back_pending_id is None + ) + def _for_each_component_lru( self, node: UnifiedTreeNode, @@ -1267,10 +1306,18 @@ def _delete_unbacked_device_leaf( def drive_host_eviction( self, component_type: ComponentType, num_tokens: int ) -> DriveHostEvictionResult: - """Evict a component's host-side resources; no-op if the component is absent.""" + """Evict a component's host-side resources; no-op if absent. Under + write_back, FULL pressure reclaims redundant Full host copies first.""" result = DriveHostEvictionResult() comp = self.components_by_type.get(component_type) if comp is not None: + if self.is_write_back and component_type == BASE_COMPONENT_TYPE: + self._reclaim_full_host_duplicates( + num_tokens, + result.tracker, + result.device_frees, + result.host_frees, + ) comp.drive_host_eviction( num_tokens, result.tracker, @@ -1279,6 +1326,74 @@ def drive_host_eviction( ) return result + def _reclaim_full_host_duplicates( + self, + num_tokens: int, + tracker: dict[ComponentType, int], + device_frees: dict[ComponentType, list[torch.Tensor]], + host_frees: dict[ComponentType, list[torch.Tensor]], + ) -> None: + """Reclaim Full host duplicates until num_tokens are freed; pass 1 + spares evictable D-leaves (imminent free demotes), pass 2 takes them.""" + swept_ids: list[NodeId] = [] + for spare_imminent_demotes in (True, False): + if tracker[BASE_COMPONENT_TYPE] >= num_tokens: + break + for node in self.full_host_duplicates.values(): + if tracker[BASE_COMPONENT_TYPE] >= num_tokens: + break + cd = node.component_data[BASE_COMPONENT_TYPE] + if cd.value is None or cd.host_value is None: + swept_ids.append(node.id) # stale entry + continue + if spare_imminent_demotes and node in self.evictable_device_leaves: + continue + if not self._can_reclaim_full_host_duplicate(node): + continue + self._release_full_host_duplicate( + node, tracker, device_frees, host_frees + ) + swept_ids.append(node.id) # released -> no longer a duplicate + # Sweep after the walk: the dict must not be mutated mid-iteration. + for nid in swept_ids: + self.full_host_duplicates.pop(nid, None) + + def _can_reclaim_full_host_duplicate(self, node: UnifiedTreeNode) -> bool: + """Full on both tiers, no in-flight DMA, no Full host lock; checked + live because tracking may be stale.""" + cd = node.component_data[BASE_COMPONENT_TYPE] + if node is self.root_node or cd.value is None or cd.host_value is None: + return False + if ( + node.write_through_pending_id is not None + or node.load_back_pending_id is not None + ): + return False + return cd.host_lock_ref == 0 + + def _release_full_host_duplicate( + self, + node: UnifiedTreeNode, + tracker: dict[ComponentType, int], + device_frees: dict[ComponentType, list[torch.Tensor]], + host_frees: dict[ComponentType, list[torch.Tensor]], + ) -> None: + """Free only the Full host layer; aux host slices stay under their own + pools' LRU (a host-only aux slice may be a sole copy).""" + assert self._can_reclaim_full_host_duplicate(node) + self._record_remove_event(node, medium=StorageMedium.CPU) + self._evict_component_and_detach_lru( + node, + self.components_by_type[BASE_COMPONENT_TYPE], + target=EvictLayer.HOST, + tracker=tracker, + device_frees=device_frees, + host_frees=host_frees, + ) + self.write_back_duplicate_reclaim_digest = ( + self.write_back_duplicate_reclaim_digest * 1000003 + node.id + 1 + ) & _RECLAIM_DIGEST_MASK + def _evict_host_leaf( self, node: UnifiedTreeNode, @@ -1420,6 +1535,8 @@ def _remove_leaf_from_parent(self, node: UnifiedTreeNode): key = node.key.child_key(self.page_size) v = node.parent.children.pop(key, None) assert v == node + # Deleted nodes must not linger in duplicate tracking as ghosts. + self.full_host_duplicates.pop(node.id, None) self._unregister_node(node) def _evict_component_and_detach_lru( @@ -1537,7 +1654,8 @@ def _is_host_leaf(self, node: UnifiedTreeNode) -> bool: """H-leaf: evicted, Full host value present, no children, unlocked, not root. Only the Full (base) component host_value is required; auxiliary - components are not mandatory for H-leaf membership.""" + components are not mandatory for H-leaf membership. In-flight DMA + marks need no check: marked nodes are never ``evicted``.""" if node is self.root_node or not node.evicted: return False if not node.backuped: @@ -1788,6 +1906,20 @@ def commit_load_back( rebuild is deferred to the orchestration layer.""" node = self.node_by_id(node_id) cache_actions: list[CacheAction | ComponentAction] = [] + if self.is_write_back: + # Write-back may reclaim a duplicate host copy while H->D DMA is + # still reading it, so pin every source node until the ack. + for xfers in ([kv_xfer], *comp_xfers.values()): + for xfer in xfers: + for nid in xfer.nodes_to_load or (): + pinned = self.node_by_id(nid) + # One live load-back per node; only the same anchor may + # re-pin (a node can sit in Full and aux transfer lists). + assert pinned.load_back_pending_id in (None, node_id), ( + f"node {nid} pinned by load-back " + f"{pinned.load_back_pending_id}, new anchor {node_id}" + ) + pinned.load_back_pending_id = node_id kv_xfer.device_indices = device_indices self.components_by_type[BASE_COMPONENT_TYPE].commit_hicache_transfer( node, @@ -1808,6 +1940,24 @@ def commit_load_back( self._update_evictable_leaf_sets(node) return cache_actions + def finish_load_back(self, anchor_node_id: NodeId) -> None: + """Finalize H->D load-back state along the anchor's root path. + + Write-back clears source-node pins at ack time. Write-through does not + use those pins, but still refreshes duplicate tracking after the device + copies become visible. Split fragments stay on the path, so the walk + covers them. + """ + node = self.node_by_id(anchor_node_id) + while node is not None and node is not self.root_node: + if self.is_write_back: + if node.load_back_pending_id != anchor_node_id: + node = node.parent + continue + node.load_back_pending_id = None + self._update_duplicate_tracking(node) + node = node.parent + def mark_write_through_pending(self, node_id: NodeId) -> None: """Mark a node as having an in-flight write-through backup.""" node = self.node_by_id(node_id) @@ -1820,6 +1970,8 @@ def finish_write_through(self, node_ids: list[NodeId], ack_id: int) -> None: node = self.node_by_id(node_id) if node.write_through_pending_id == ack_id: node.write_through_pending_id = None + # The backed-up copy becomes a tracked duplicate only now. + self._update_duplicate_tracking(node) self._record_store_event(node, medium=StorageMedium.CPU) def set_component_device_value( @@ -1887,6 +2039,7 @@ def sanity_check( # ── PART 2: Per-node state machine and leaf qualification ── expected_dev_leaves: set[UnifiedTreeNode] = set() expected_hst_leaves: set[UnifiedTreeNode] = set() + expected_duplicates: set[UnifiedTreeNode] = set() for node in all_nodes: if node is self.root_node: @@ -1903,7 +2056,10 @@ def sanity_check( if cd.value is not None and not full_dev: E(f"node {nid} {ct} device present but Full.value=None") if cd.host_value is not None and not full_hst: - E(f"node {nid} {ct} host present but Full.host_value=None") + # write_back reclaim takes only the Full host layer; an + # aux host slice may outlive it while Full device is live. + if not (self.is_write_back and full_dev): + E(f"node {nid} {ct} host present but Full.host_value=None") # Every node must keep Full data on at least one layer. if not full_dev and not full_hst: @@ -1936,6 +2092,8 @@ def sanity_check( expected_dev_leaves.add(node) if self._is_host_leaf(node): expected_hst_leaves.add(node) + if self._is_settled_full_host_duplicate(node): + expected_duplicates.add(node) # ── PART 3: Tracking structures ── @@ -1957,6 +2115,16 @@ def sanity_check( if missing: E(f"H-leaf missing: {[n.id for n in list(missing)[:5]]}") + # Lazy deregistration: stale extras are legal; settled duplicates must + # be tracked and entries must not outlive their node. + expected_ids = {n.id for n in expected_duplicates} + dup_ids = set(self.full_host_duplicates.keys()) + if expected_ids - dup_ids: + E(f"Duplicate missing: {list(expected_ids - dup_ids)[:5]}") + ghost_ids = dup_ids - {n.id for n in all_nodes} + if ghost_ids: + E(f"Duplicate ghosts: {list(ghost_ids)[:5]}") + # D-leaf ∩ H-leaf = ∅ overlap = self.evictable_device_leaves & self.evictable_host_leaves if overlap: @@ -2068,6 +2236,18 @@ def sanity_check( E( f"[Ongoing] load_back node {nid} lock_ref={n.component_data[FCT].lock_ref}" ) + # Every in-flight H->D mark must belong to a live load-back; a leaked + # mark would pin the node's host copy against reclaim forever. + ongoing_load_ids = {node_id for _, node_id in ongoing_load_back} + for node in all_nodes: + if ( + node.load_back_pending_id is not None + and node.load_back_pending_id not in ongoing_load_ids + ): + E( + f"[Ongoing] node {node.id} load_back_pending_id=" + f"{node.load_back_pending_id} has no live load-back" + ) if errors: msg = ( diff --git a/python/sglang/srt/mem_cache/unified_cache/unified_tree_core_interface.py b/python/sglang/srt/mem_cache/unified_cache/unified_tree_core_interface.py index cbd9910ffbed..4edb198be097 100644 --- a/python/sglang/srt/mem_cache/unified_cache/unified_tree_core_interface.py +++ b/python/sglang/srt/mem_cache/unified_cache/unified_tree_core_interface.py @@ -439,6 +439,15 @@ def commit_load_back( """Commit a successful H->D load-back onto the node; returns any cache actions.""" ... + @abstractmethod + def finish_load_back(self, anchor_node_id: NodeId) -> None: + """Clear the in-flight H->D marks on the anchor's root path at ack time.""" + ... + + # Order-sensitive digest of write_back duplicate-reclaim victim ids, + # cross-checked across TP ranks; cores that never reclaim keep 0. + write_back_duplicate_reclaim_digest: int = 0 + @abstractmethod def mark_write_through_pending(self, node_id: NodeId) -> None: """Mark a node as having an in-flight write-through backup.""" diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index e013fd6901a1..1120b1b74b2e 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -1794,22 +1794,30 @@ def _sync_hicache_ready_counts( else () ) + # Piggybacked TP check: [digest, -digest] MIN-reduces to [min, -max], + # equal iff reclaim victim order matched on every rank. + digest = self.tree_core.write_back_duplicate_reclaim_digest ready_counts = torch.tensor( [ write_acks, load_acks, *storage_queue_sizes, + digest, + -digest, ], - dtype=torch.int, + dtype=torch.int64, device="cpu", ) self._all_reduce(ready_counts, torch.distributed.ReduceOp.MIN) count_values = list(map(int, ready_counts.tolist())) + assert ( + count_values[-2] == -count_values[-1] + ), "write_back duplicate-reclaim victims diverged across TP ranks" return ( count_values[0], count_values[1], - tuple(count_values[2:]), + tuple(count_values[2:-2]), extra_pool_names, ) @@ -1864,11 +1872,17 @@ def loading_check(self, finish_count: Optional[int] = None) -> None: finish_count = 0 if self.pp_rank == 0: finish_count = self._count_ready_acks(cc.ack_load_queue) - finish_count_tensor = torch.tensor( - finish_count, dtype=torch.int, device="cpu" + # Piggybacked TP check: [digest, -digest] MIN-reduces to [min, -max], + # equal iff reclaim victim order matched on every rank. + digest = self.tree_core.write_back_duplicate_reclaim_digest + sync_tensor = torch.tensor( + [finish_count, digest, -digest], dtype=torch.int64, device="cpu" ) - self._all_reduce(finish_count_tensor, torch.distributed.ReduceOp.MIN) - finish_count = finish_count_tensor.item() + self._all_reduce(sync_tensor, torch.distributed.ReduceOp.MIN) + finish_count = int(sync_tensor[0].item()) + assert ( + sync_tensor[1].item() == -sync_tensor[2].item() + ), "write_back duplicate-reclaim victims diverged across TP ranks" while finish_count > 0: ack = cc.ack_load_queue.pop(0) @@ -1877,6 +1891,8 @@ def loading_check(self, finish_count: Optional[int] = None) -> None: node, lock_params, host_lock_params = self.ongoing_load_back.pop(ack_id) self.dec_lock_ref(node, lock_params) self.dec_host_lock_ref(node, host_lock_params) + # Unpin the loaded nodes; host copies stay as reclaimable duplicates. + self.tree_core.finish_load_back(node) if self.metrics_collector is not None: self.metrics_collector.increment_load_back_num_tokens(ack.num_tokens) diff --git a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py index b875e4b9b74d..a810800225fa 100644 --- a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py +++ b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py @@ -65,6 +65,7 @@ EvictLayer, TreeComponent, ) +from sglang.srt.mem_cache.unified_cache.unified_tree_core import UnifiedTreeCore from sglang.srt.mem_cache.unified_cache.unified_tree_core_interface import ( DecSwaLockOnlyResult, DemoteResult, @@ -254,6 +255,82 @@ def make_node(): self.assertEqual(n4.get_prefix_hash_values(n3), ["h1", "h2", "h3"]) +class TestUnifiedTreeCoreLoadBackPending(CustomTestCase): + def _build_core(self, *, is_write_back: bool): + component_types = (ComponentType.FULL,) + root = UnifiedTreeNode(component_types) + shared = UnifiedTreeNode(component_types) + anchor_a = UnifiedTreeNode(component_types) + anchor_b = UnifiedTreeNode(component_types) + shared.parent = root + anchor_a.parent = shared + anchor_b.parent = shared + + nodes = {node.id: node for node in (root, shared, anchor_a, anchor_b)} + core = mock.Mock() + core.is_write_back = is_write_back + core.root_node = root + core.node_by_id.side_effect = nodes.__getitem__ + core.components_by_type = {ComponentType.FULL: mock.Mock()} + core.full_host_duplicates = {} + core._is_settled_full_host_duplicate.side_effect = ( + lambda node: UnifiedTreeCore._is_settled_full_host_duplicate(core, node) + ) + core._update_duplicate_tracking.side_effect = ( + lambda node: UnifiedTreeCore._update_duplicate_tracking(core, node) + ) + return core, shared, anchor_a, anchor_b + + def _commit_load_back(self, core, anchor, source): + transfer = PoolTransfer( + name=PoolName.KV, + host_indices=torch.tensor([1], dtype=torch.int64), + nodes_to_load=[source.id], + ) + return UnifiedTreeCore.commit_load_back( + core, + anchor.id, + torch.tensor([2], dtype=torch.int64), + transfer, + {}, + ) + + def test_write_through_different_anchors_track_duplicate_without_pending(self): + core, shared, anchor_a, anchor_b = self._build_core(is_write_back=False) + full = shared.component_data[ComponentType.FULL] + full.value = torch.tensor([1], dtype=torch.int64) + full.host_value = torch.tensor([2], dtype=torch.int64) + + self._commit_load_back(core, anchor_a, shared) + self._commit_load_back(core, anchor_b, shared) + UnifiedTreeCore.finish_load_back(core, anchor_a.id) + + self.assertIsNone(shared.load_back_pending_id) + self.assertIn(shared.id, core.full_host_duplicates) + core._update_duplicate_tracking.assert_has_calls( + [mock.call(anchor_a), mock.call(shared)] + ) + + def test_write_back_pending_blocks_reclaim_until_ack(self): + core, shared, anchor_a, anchor_b = self._build_core(is_write_back=True) + full = shared.component_data[ComponentType.FULL] + full.value = torch.tensor([1], dtype=torch.int64) + full.host_value = torch.tensor([2], dtype=torch.int64) + + self._commit_load_back(core, anchor_a, shared) + + self.assertEqual(shared.load_back_pending_id, anchor_a.id) + self.assertFalse(UnifiedTreeCore._can_reclaim_full_host_duplicate(core, shared)) + with self.assertRaisesRegex(AssertionError, "new anchor"): + self._commit_load_back(core, anchor_b, shared) + + UnifiedTreeCore.finish_load_back(core, anchor_a.id) + + self.assertIsNone(shared.load_back_pending_id) + self.assertTrue(UnifiedTreeCore._can_reclaim_full_host_duplicate(core, shared)) + core._update_duplicate_tracking.assert_called_once_with(shared) + + def _write_backup(cache, node, write_back: bool = False) -> int: """Back up one node's KV D->H via the tree's build+execute primitives.""" return cache._execute_and_commit_kv_backup( @@ -2786,6 +2863,8 @@ def _simulate_backup(self, cache, node): cd = ancestor.component_data[ct] if cd.value is not None and cd.host_value is None: cd.host_value = cd.value.clone() + # A real backup registers duplicate tracking at its ack. + cache.tree_core._update_duplicate_tracking(ancestor) def _simulate_backup_tree(self, cache): """Backup all non-root nodes (simulates write-through).""" From 12cbb0d5b3d568121b3eb29a94cc5cff28c78b1e Mon Sep 17 00:00:00 2001 From: luoroger37 Date: Mon, 24 Aug 2026 16:41:01 +0800 Subject: [PATCH 34/47] [HiCache] Clamp tombstoned SWA locations in unified translation (#736) --- .../sglang/srt/mem_cache/unified_memory_pool.py | 16 +++++++++++----- 1 file changed, 11 insertions(+), 5 deletions(-) diff --git a/python/sglang/srt/mem_cache/unified_memory_pool.py b/python/sglang/srt/mem_cache/unified_memory_pool.py index 1575616ac053..76b2d8eb0e52 100644 --- a/python/sglang/srt/mem_cache/unified_memory_pool.py +++ b/python/sglang/srt/mem_cache/unified_memory_pool.py @@ -1345,12 +1345,18 @@ def translate_loc_from_full_to_swa(self, kv_indices: torch.Tensor): "attach_allocators" ) ps = self._swa_allocator.page_size + # Tombstone-safety clamp, matching MultiEndedAllocator.translate_kv_loc: + # a tombstoned v2p entry (-1) must not reach the caller as a negative + # loc. Clamp to 0 routes it to the reserved padding sink instead. if ps == 1: - return self._swa_allocator.virtual_to_physical[kv_indices].to(torch.int32) - virt_pages = kv_indices // ps - offsets = kv_indices % ps - swa_phys_pages = self._swa_allocator.virtual_to_physical[virt_pages] - return (swa_phys_pages * ps + offsets).to(torch.int32) + swa_locs = self._swa_allocator.virtual_to_physical[kv_indices] + else: + virt_pages = kv_indices // ps + offsets = kv_indices % ps + swa_phys_pages = self._swa_allocator.virtual_to_physical[virt_pages] + # Tombstoned page: -1 * ps + offset lands in [-ps, -1]. + swa_locs = swa_phys_pages * ps + offsets + return swa_locs.clamp(min=0).to(torch.int32) def get_state_buf_infos(self): return self.swa_kv_pool.get_contiguous_buf_infos() From 9f883bbc702983d52c60d2ecc4faa67fe7296a69 Mon Sep 17 00:00:00 2001 From: Hank Han Date: Mon, 24 Aug 2026 18:13:46 +0800 Subject: [PATCH 35/47] ci: support customer registry source for zstd-to-gzip conversion Add source_image_is_customer input and customer registry login for converting customer zstd images back to gzip format in dev registry. Co-authored-by: TRAE CLI Co-authored-by: TRAE CLI --- .github/workflows/backfill-image-formats.yml | 15 +++++++++++++-- 1 file changed, 13 insertions(+), 2 deletions(-) diff --git a/.github/workflows/backfill-image-formats.yml b/.github/workflows/backfill-image-formats.yml index aaa202ff0302..c1b765513945 100644 --- a/.github/workflows/backfill-image-formats.yml +++ b/.github/workflows/backfill-image-formats.yml @@ -18,6 +18,11 @@ on: required: false type: string default: "" + source_image_is_customer: + description: "源镜像在客户仓库(ai-containers-cn-beijing)而非研发仓库" + required: false + type: boolean + default: false sync_to_customer: description: "补全格式后触发 Bits 流水线同步到客户仓库" required: false @@ -54,9 +59,12 @@ jobs: uses: docker/setup-buildx-action@v3 - name: Login to Volcengine CR + env: + CUSTOMER_CR_REGISTRY: ${{ vars.CUSTOMER_CR_REGISTRY || 'ai-containers-cn-beijing.cr.volces.com' }} run: | set -euo pipefail echo "${VOLCENGINE_CR_PASSWORD}" | docker login "${VOLCENGINE_CR_REGISTRY}" --username "${VOLCENGINE_CR_USERNAME}" --password-stdin + echo "${VOLCENGINE_CR_PASSWORD}" | docker login "${CUSTOMER_CR_REGISTRY}" --username "${VOLCENGINE_CR_USERNAME}" --password-stdin || echo "[login] customer registry login failed (non-fatal)" - name: Install nydus tooling run: | @@ -80,12 +88,13 @@ jobs: env: TARGET_IMAGE_REF: ${{ inputs.image_refs }} SOURCE_IMAGE_REF: ${{ inputs.source_image_ref }} + SOURCE_IS_CUSTOMER: ${{ inputs.source_image_is_customer }} run: | set -euo pipefail TARGET="$(echo "${TARGET_IMAGE_REF}" | head -1 | tr -d '\r\n' | sed 's/^[[:space:]]*//;s/[[:space:]]*$//')" - echo "=== $(date '+%H:%M:%S') Rename ===" + echo "=== $(date '+%H:%M:%S') Rename (zstd-customer -> gzip-dev) ===" echo " Source: ${SOURCE_IMAGE_REF}" echo " Target: ${TARGET}" @@ -94,10 +103,12 @@ jobs: exit 0 fi + # Pull from customer registry (zstd compressed), docker auto-decompresses. + # Push to dev registry uses default gzip compression, completing zstd->gzip conversion. docker pull "${SOURCE_IMAGE_REF}" docker tag "${SOURCE_IMAGE_REF}" "${TARGET}" docker push "${TARGET}" - echo " Renamed: ${TARGET}" + echo " Converted & renamed: ${SOURCE_IMAGE_REF} -> ${TARGET}" - name: Backfill image formats env: From 919e675a1e0a3c7b939879907aed8f3a7e4befa4 Mon Sep 17 00:00:00 2001 From: Hank Han Date: Tue, 25 Aug 2026 00:40:51 +0800 Subject: [PATCH 36/47] ci: simplify backfill workflow, merge rename+backfill, remove customer sync Co-authored-by: TRAE CLI --- .github/workflows/backfill-image-formats.yml | 201 +++++++++++++------ 1 file changed, 136 insertions(+), 65 deletions(-) diff --git a/.github/workflows/backfill-image-formats.yml b/.github/workflows/backfill-image-formats.yml index c1b765513945..21efdcf0b4c4 100644 --- a/.github/workflows/backfill-image-formats.yml +++ b/.github/workflows/backfill-image-formats.yml @@ -6,28 +6,48 @@ on: image_refs: description: | 目标镜像引用列表,每行一个,不含格式后缀。 + auto_rename=true 且 source_image_ref 为空时,以此为源镜像列表进行重命名+补全。 格式: iaas-gpu-cn-beijing.cr.volces.com/serving/: 示例: iaas-gpu-cn-beijing.cr.volces.com/serving/sglang:v0.5.16.iaas.202607280000-kimi-k3-cu130 - required: true + required: false type: string + default: "" source_image_ref: description: | - 可选:源镜像引用,用于将旧命名格式的镜像重命名为新 tag(docker pull + tag + push)。 - 留空则跳过重命名,直接对 image_refs 中的镜像补全格式。 + 可选:当 auto_rename=true 时,以此为源镜像(而非 image_refs)进行重命名。 + 留空则使用 image_refs 中的镜像作为源。 格式: iaas-gpu-cn-beijing.cr.volces.com/serving/: required: false type: string default: "" - source_image_is_customer: - description: "源镜像在客户仓库(ai-containers-cn-beijing)而非研发仓库" + auto_rename: + description: | + 启用自动重命名:从源镜像 tag 提取版本号,按 ServingKit 日构建规则自动生成标准 tag。 + 有 source_image_ref 时以它为准;否则以 image_refs 中的镜像为准。 + 重命名后自动补全 zstd + nydus 格式。 required: false type: boolean default: false - sync_to_customer: - description: "补全格式后触发 Bits 流水线同步到客户仓库" + auto_rename_mode: + description: "自动重命名模式: manual→iaas.dev, nightly→iaas.nightly, version→byted.{tag_value}" required: false - type: boolean - default: false + type: string + default: "manual" + auto_rename_tag_value: + description: "version 模式下的 tag_value(仅 auto_rename_mode=version 时需要)" + required: false + type: string + default: "" + auto_rename_variant_suffix: + description: "变体后缀(如 deepseek-v4, kimi-k3, glm52-pd4)" + required: false + type: string + default: "" + auto_rename_cuda_suffix: + description: "CUDA 后缀(如 cu130)" + required: false + type: string + default: "" concurrency: group: backfill-image-formats-${{ github.run_id }} @@ -59,12 +79,9 @@ jobs: uses: docker/setup-buildx-action@v3 - name: Login to Volcengine CR - env: - CUSTOMER_CR_REGISTRY: ${{ vars.CUSTOMER_CR_REGISTRY || 'ai-containers-cn-beijing.cr.volces.com' }} run: | set -euo pipefail echo "${VOLCENGINE_CR_PASSWORD}" | docker login "${VOLCENGINE_CR_REGISTRY}" --username "${VOLCENGINE_CR_USERNAME}" --password-stdin - echo "${VOLCENGINE_CR_PASSWORD}" | docker login "${CUSTOMER_CR_REGISTRY}" --username "${VOLCENGINE_CR_USERNAME}" --password-stdin || echo "[login] customer registry login failed (non-fatal)" - name: Install nydus tooling run: | @@ -83,36 +100,24 @@ jobs: nydusify --version nydus-image --version - - name: Rename image (old tag -> new tag) - if: ${{ inputs.source_image_ref != '' }} - env: - TARGET_IMAGE_REF: ${{ inputs.image_refs }} - SOURCE_IMAGE_REF: ${{ inputs.source_image_ref }} - SOURCE_IS_CUSTOMER: ${{ inputs.source_image_is_customer }} + - name: Validate inputs run: | set -euo pipefail - - TARGET="$(echo "${TARGET_IMAGE_REF}" | head -1 | tr -d '\r\n' | sed 's/^[[:space:]]*//;s/[[:space:]]*$//')" - - echo "=== $(date '+%H:%M:%S') Rename (zstd-customer -> gzip-dev) ===" - echo " Source: ${SOURCE_IMAGE_REF}" - echo " Target: ${TARGET}" - - if docker manifest inspect "${TARGET}" >/dev/null 2>&1; then - echo " Target already exists, skipping rename" - exit 0 + if [ "${{ inputs.auto_rename }}" = "false" ] && [ -z "${{ inputs.image_refs }}" ]; then + echo "ERROR: image_refs is required when auto_rename is disabled" + exit 1 fi + echo "Input validation passed." - # Pull from customer registry (zstd compressed), docker auto-decompresses. - # Push to dev registry uses default gzip compression, completing zstd->gzip conversion. - docker pull "${SOURCE_IMAGE_REF}" - docker tag "${SOURCE_IMAGE_REF}" "${TARGET}" - docker push "${TARGET}" - echo " Converted & renamed: ${SOURCE_IMAGE_REF} -> ${TARGET}" - - - name: Backfill image formats + - name: Auto-rename and backfill env: IMAGE_REFS: ${{ inputs.image_refs }} + SOURCE_IMAGE_REF: ${{ inputs.source_image_ref }} + AUTO_RENAME: ${{ inputs.auto_rename }} + AUTO_RENAME_MODE: ${{ inputs.auto_rename_mode }} + AUTO_RENAME_TAG_VALUE: ${{ inputs.auto_rename_tag_value }} + AUTO_RENAME_VARIANT_SUFFIX: ${{ inputs.auto_rename_variant_suffix }} + AUTO_RENAME_CUDA_SUFFIX: ${{ inputs.auto_rename_cuda_suffix }} run: | set -euo pipefail @@ -125,11 +130,80 @@ jobs: return 1 } + generate_tag() { + local source_tag="$1" + local version + version="$(echo "${source_tag}" | sed -n 's/^v\([0-9][0-9.]*\)\..*/\1/p')" + if [ -z "${version}" ]; then + echo " ERROR: cannot extract version from tag: ${source_tag}" + echo " Expected format: v., e.g. v0.5.16.iaas.202607280000-kimi-k3-cu130" + return 1 + fi + echo " Extracted version: ${version}" + + local timestamp="$(TZ=Asia/Shanghai date +'%Y%m%d%H%M')" + local new_tag + case "${AUTO_RENAME_MODE}" in + nightly) new_tag="v${version}.iaas.nightly.${timestamp}" ;; + manual) new_tag="v${version}.iaas.dev.${timestamp}" ;; + version) + if [ -z "${AUTO_RENAME_TAG_VALUE}" ]; then + echo " ERROR: auto_rename_tag_value is required for version mode" + return 1 + fi + new_tag="v${version}.byted.${AUTO_RENAME_TAG_VALUE}.${timestamp}" + ;; + *) + echo " ERROR: unknown auto_rename_mode: ${AUTO_RENAME_MODE}" + return 1 + ;; + esac + + if [ -n "${AUTO_RENAME_VARIANT_SUFFIX}" ]; then + new_tag="${new_tag}-${AUTO_RENAME_VARIANT_SUFFIX}" + fi + if [ -n "${AUTO_RENAME_CUDA_SUFFIX}" ]; then + new_tag="${new_tag}-${AUTO_RENAME_CUDA_SUFFIX}" + fi + + echo "${new_tag}" + } + + rename_and_backfill() { + local source_ref="$1" + + echo "" + echo "=== $(date '+%H:%M:%S') Processing: ${source_ref} ===" + + local source_tag="${source_ref##*:}" + local source_registry_ns="${source_ref%%:*}" + + if [ "${AUTO_RENAME}" = "true" ]; then + local new_tag + new_tag="$(generate_tag "${source_tag}")" || return 1 + local target_ref="${source_registry_ns}:${new_tag}" + echo " Generated target: ${target_ref}" + + if docker manifest inspect "${target_ref}" >/dev/null 2>&1; then + echo " Target already exists, using existing: ${target_ref}" + else + echo " Renaming: ${source_ref} -> ${target_ref}" + docker pull "${source_ref}" + docker tag "${source_ref}" "${target_ref}" + docker push "${target_ref}" + echo " Renamed: ${target_ref}" + fi + backfill_one "${target_ref}" + else + backfill_one "${source_ref}" + fi + } + backfill_one() { local image_ref="$1" echo "" - echo "=== $(date '+%H:%M:%S') ${image_ref} ===" + echo "=== $(date '+%H:%M:%S') Backfill: ${image_ref} ===" if ! docker manifest inspect "${image_ref}" >/dev/null 2>&1; then echo " ERROR: base image not found: ${image_ref}" @@ -167,36 +241,33 @@ jobs: fi echo " done: ${image_ref}" - echo "" } - while IFS= read -r line; do - line="${line//[$'\r\n']/}" - [[ -z "$line" || "$line" == \#* ]] && continue - backfill_one "$line" - done <<< "${IMAGE_REFS}" - - echo "All backfills complete." - - - name: Trigger Bits sync pipeline - if: ${{ inputs.sync_to_customer }} - env: - IMAGE_REFS: ${{ inputs.image_refs }} - run: | - set -euo pipefail - - BITS_PIPELINE_ID="1008460003842" - BITS_SPACE_ID="779378790402" + echo "=== $(date '+%H:%M:%S') Backfill Image Formats ===" + echo " auto_rename: ${AUTO_RENAME}" + echo " auto_rename_mode: ${AUTO_RENAME_MODE}" + echo " variant_suffix: ${AUTO_RENAME_VARIANT_SUFFIX:-}" + echo " cuda_suffix: ${AUTO_RENAME_CUDA_SUFFIX:-}" - echo "Images to sync:" - while IFS= read -r line; do - line="${line//[$'\r\n']/}" - [[ -z "$line" || "$line" == \#* ]] && continue - framework="${line##*/}"; framework="${framework%%:*}" - tag="${line##*:}" - echo " framework=${framework} tag=${tag}" - done <<< "${IMAGE_REFS}" + # Determine source images + if [ "${AUTO_RENAME}" = "true" ] && [ -n "${SOURCE_IMAGE_REF}" ]; then + echo " Using source_image_ref as rename source" + rename_and_backfill "${SOURCE_IMAGE_REF}" + elif [ "${AUTO_RENAME}" = "true" ]; then + echo " Using image_refs as rename source" + while IFS= read -r line; do + line="${line//[$'\r\n']/}" + [[ -z "$line" || "$line" == \#* ]] && continue + rename_and_backfill "$line" + done <<< "${IMAGE_REFS}" + else + echo " Direct backfill (no rename)" + while IFS= read -r line; do + line="${line//[$'\r\n']/}" + [[ -z "$line" || "$line" == \#* ]] && continue + backfill_one "$line" + done <<< "${IMAGE_REFS}" + fi echo "" - echo "Bits 流水线: https://bits.bytedance.net/devops/${BITS_SPACE_ID}/pipeline/detail/${BITS_PIPELINE_ID}" - echo "请手动触发流水线,填入对应的 framework 和 tag。" \ No newline at end of file + echo "All backfills complete." \ No newline at end of file From bd15d4f1d908c4d30de23eaa0471d7738cf23d0e Mon Sep 17 00:00:00 2001 From: Hank Han Date: Tue, 25 Aug 2026 00:54:49 +0800 Subject: [PATCH 37/47] ci: detect sglang version from historical images Co-authored-by: TRAE CLI --- .github/workflows/backfill-image-formats.yml | 41 ++++++++++++++++---- 1 file changed, 33 insertions(+), 8 deletions(-) diff --git a/.github/workflows/backfill-image-formats.yml b/.github/workflows/backfill-image-formats.yml index 21efdcf0b4c4..9b7e0022aafa 100644 --- a/.github/workflows/backfill-image-formats.yml +++ b/.github/workflows/backfill-image-formats.yml @@ -132,14 +132,39 @@ jobs: generate_tag() { local source_tag="$1" + local source_ref="$2" local version version="$(echo "${source_tag}" | sed -n 's/^v\([0-9][0-9.]*\)\..*/\1/p')" if [ -z "${version}" ]; then - echo " ERROR: cannot extract version from tag: ${source_tag}" - echo " Expected format: v., e.g. v0.5.16.iaas.202607280000-kimi-k3-cu130" - return 1 + echo " Version is not present in tag '${source_tag}'; inspecting the image..." >&2 + if ! docker pull "${source_ref}" >&2; then + echo " ERROR: failed to pull source image for version detection: ${source_ref}" >&2 + return 1 + fi + + local detected_version="" + local python_cmd + for python_cmd in python3 python; do + if detected_version="$( + docker run --rm --entrypoint "${python_cmd}" "${source_ref}" \ + -c 'import importlib.metadata as m; print(m.version("sglang"))' \ + 2>/dev/null + )" && [ -n "${detected_version}" ]; then + echo " Detected installed sglang version with ${python_cmd}: ${detected_version}" >&2 + break + fi + detected_version="" + done + + version="$(echo "${detected_version}" | grep -oE '[0-9]+\.[0-9]+\.[0-9]+' | head -1 || true)" + if [ -z "${version}" ]; then + echo " ERROR: the tag has no version and the container does not expose an installed sglang package version" >&2 + echo " Source image: ${source_ref}" >&2 + echo " Tried: python3/python + importlib.metadata.version(\"sglang\")" >&2 + return 1 + fi fi - echo " Extracted version: ${version}" + echo " Using sglang version: ${version}" >&2 local timestamp="$(TZ=Asia/Shanghai date +'%Y%m%d%H%M')" local new_tag @@ -148,13 +173,13 @@ jobs: manual) new_tag="v${version}.iaas.dev.${timestamp}" ;; version) if [ -z "${AUTO_RENAME_TAG_VALUE}" ]; then - echo " ERROR: auto_rename_tag_value is required for version mode" + echo " ERROR: auto_rename_tag_value is required for version mode" >&2 return 1 fi new_tag="v${version}.byted.${AUTO_RENAME_TAG_VALUE}.${timestamp}" ;; *) - echo " ERROR: unknown auto_rename_mode: ${AUTO_RENAME_MODE}" + echo " ERROR: unknown auto_rename_mode: ${AUTO_RENAME_MODE}" >&2 return 1 ;; esac @@ -180,7 +205,7 @@ jobs: if [ "${AUTO_RENAME}" = "true" ]; then local new_tag - new_tag="$(generate_tag "${source_tag}")" || return 1 + new_tag="$(generate_tag "${source_tag}" "${source_ref}")" || return 1 local target_ref="${source_registry_ns}:${new_tag}" echo " Generated target: ${target_ref}" @@ -270,4 +295,4 @@ jobs: fi echo "" - echo "All backfills complete." \ No newline at end of file + echo "All backfills complete." From 307f2b945c25127c4f4390a6806cd4838ba86dc3 Mon Sep 17 00:00:00 2001 From: Hank Han Date: Tue, 25 Aug 2026 01:10:25 +0800 Subject: [PATCH 38/47] ci: reject placeholder versions in image backfill Co-authored-by: TRAE CLI --- .github/workflows/backfill-image-formats.yml | 63 ++++++++++++++++++-- 1 file changed, 59 insertions(+), 4 deletions(-) diff --git a/.github/workflows/backfill-image-formats.yml b/.github/workflows/backfill-image-formats.yml index 9b7e0022aafa..2701296e9f76 100644 --- a/.github/workflows/backfill-image-formats.yml +++ b/.github/workflows/backfill-image-formats.yml @@ -134,7 +134,7 @@ jobs: local source_tag="$1" local source_ref="$2" local version - version="$(echo "${source_tag}" | sed -n 's/^v\([0-9][0-9.]*\)\..*/\1/p')" + version="$(echo "${source_tag}" | grep -oE '^v[0-9]+\.[0-9]+\.[0-9]+(\.post[0-9]+)?' | sed 's/^v//' || true)" if [ -z "${version}" ]; then echo " Version is not present in tag '${source_tag}'; inspecting the image..." >&2 if ! docker pull "${source_ref}" >&2; then @@ -156,11 +156,66 @@ jobs: detected_version="" done - version="$(echo "${detected_version}" | grep -oE '[0-9]+\.[0-9]+\.[0-9]+' | head -1 || true)" + version="$(echo "${detected_version}" | grep -oE '[0-9]+\.[0-9]+\.[0-9]+(\.post[0-9]+)?' | head -1 || true)" + if [ "${version}" = "0.0.0" ]; then + echo " Ignoring placeholder package version: ${detected_version}" >&2 + version="" + fi + + # Source-overlay images can intentionally retain a 0.0.0 wheel + # version. Resolve their recorded build commit to the nearest + # release tag in this repository. + if [ -z "${version}" ]; then + local build_commit + build_commit="$( + docker image inspect "${source_ref}" \ + --format '{{ index .Config.Labels "ai.sglang.build.commit" }}' \ + 2>/dev/null || true + )" + if [[ "${build_commit}" =~ ^[0-9a-fA-F]{7,40}$ ]]; then + echo " Resolving recorded build commit: ${build_commit}" >&2 + if ! git cat-file -e "${build_commit}^{commit}" 2>/dev/null; then + git fetch --quiet --no-tags origin "${build_commit}" --depth=1000 || true + fi + git fetch --quiet --tags origin || true + local release_tag + release_tag="$( + git describe --tags --match 'v[0-9]*' --abbrev=0 "${build_commit}" \ + 2>/dev/null || true + )" + version="$( + echo "${release_tag}" | \ + grep -oE '^v[0-9]+\.[0-9]+\.[0-9]+(\.post[0-9]+)?' | \ + sed 's/^v//' || true + )" + if [ -n "${version}" ]; then + echo " Resolved ${build_commit} to release baseline ${release_tag}" >&2 + fi + fi + fi + + # Older images without source provenance may still record the + # versioned ServingKit image they were built from. + if [ -z "${version}" ]; then + local base_image + base_image="$( + docker image inspect "${source_ref}" \ + --format '{{ index .Config.Labels "ai.sglang.build.base-image" }}' \ + 2>/dev/null || true + )" + version="$( + echo "${base_image}" | \ + grep -oE 'v[0-9]+\.[0-9]+\.[0-9]+(\.post[0-9]+)?' | \ + head -1 | sed 's/^v//' || true + )" + if [ -n "${version}" ]; then + echo " Detected version from base-image metadata: ${version}" >&2 + fi + fi + if [ -z "${version}" ]; then - echo " ERROR: the tag has no version and the container does not expose an installed sglang package version" >&2 + echo " ERROR: unable to determine a non-placeholder sglang version from the tag, container, build commit, or image metadata" >&2 echo " Source image: ${source_ref}" >&2 - echo " Tried: python3/python + importlib.metadata.version(\"sglang\")" >&2 return 1 fi fi From 0faf49f28802363c6f58ff888e42151491079fb0 Mon Sep 17 00:00:00 2001 From: Hank Han Date: Tue, 25 Aug 2026 12:21:44 +0800 Subject: [PATCH 39/47] ci: simplify historical image auto-rename inputs Co-authored-by: TRAE CLI --- .github/workflows/backfill-image-formats.yml | 23 -------------------- 1 file changed, 23 deletions(-) diff --git a/.github/workflows/backfill-image-formats.yml b/.github/workflows/backfill-image-formats.yml index 2701296e9f76..393a980812d3 100644 --- a/.github/workflows/backfill-image-formats.yml +++ b/.github/workflows/backfill-image-formats.yml @@ -38,17 +38,6 @@ on: required: false type: string default: "" - auto_rename_variant_suffix: - description: "变体后缀(如 deepseek-v4, kimi-k3, glm52-pd4)" - required: false - type: string - default: "" - auto_rename_cuda_suffix: - description: "CUDA 后缀(如 cu130)" - required: false - type: string - default: "" - concurrency: group: backfill-image-formats-${{ github.run_id }} cancel-in-progress: false @@ -116,8 +105,6 @@ jobs: AUTO_RENAME: ${{ inputs.auto_rename }} AUTO_RENAME_MODE: ${{ inputs.auto_rename_mode }} AUTO_RENAME_TAG_VALUE: ${{ inputs.auto_rename_tag_value }} - AUTO_RENAME_VARIANT_SUFFIX: ${{ inputs.auto_rename_variant_suffix }} - AUTO_RENAME_CUDA_SUFFIX: ${{ inputs.auto_rename_cuda_suffix }} run: | set -euo pipefail @@ -239,13 +226,6 @@ jobs: ;; esac - if [ -n "${AUTO_RENAME_VARIANT_SUFFIX}" ]; then - new_tag="${new_tag}-${AUTO_RENAME_VARIANT_SUFFIX}" - fi - if [ -n "${AUTO_RENAME_CUDA_SUFFIX}" ]; then - new_tag="${new_tag}-${AUTO_RENAME_CUDA_SUFFIX}" - fi - echo "${new_tag}" } @@ -326,9 +306,6 @@ jobs: echo "=== $(date '+%H:%M:%S') Backfill Image Formats ===" echo " auto_rename: ${AUTO_RENAME}" echo " auto_rename_mode: ${AUTO_RENAME_MODE}" - echo " variant_suffix: ${AUTO_RENAME_VARIANT_SUFFIX:-}" - echo " cuda_suffix: ${AUTO_RENAME_CUDA_SUFFIX:-}" - # Determine source images if [ "${AUTO_RENAME}" = "true" ] && [ -n "${SOURCE_IMAGE_REF}" ]; then echo " Using source_image_ref as rename source" From d10114945bfea39cafa3281ce7d463bc16fc2881 Mon Sep 17 00:00:00 2001 From: Qi Sun <65223714+shiyu7@users.noreply.github.com> Date: Tue, 25 Aug 2026 14:50:04 +0800 Subject: [PATCH 40/47] fix: prevent false decode retraction with SWA (#739) --- python/sglang/srt/managers/schedule_batch.py | 73 +++++++-- python/sglang/srt/speculative/eagle_utils.py | 4 - python/sglang/srt/speculative/spec_utils.py | 9 +- .../test_schedule_batch_swa_retraction.py | 151 ++++++++++++++++++ .../spec/test_decode_bookkeeping_ownership.py | 17 +- .../unit/spec/test_spec_prepare_for_decode.py | 69 ++++++++ 6 files changed, 299 insertions(+), 24 deletions(-) create mode 100644 test/registered/unit/managers/test_schedule_batch_swa_retraction.py create mode 100644 test/registered/unit/spec/test_spec_prepare_for_decode.py diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index a45d122f0ff5..0af6a67cf249 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -1841,7 +1841,8 @@ def release_req( tree_cache: BasePrefixCache, hisparse_coordinator: Optional[HiSparseCoordinator], offload_kv: bool = True, -) -> None: + abort_on_unsupported_backup: bool = False, +) -> bool: if hisparse_coordinator is not None and not req.finished(): hisparse_coordinator.retract_req(req) @@ -1849,8 +1850,20 @@ def release_req( # restored later without recompute (see resume_retracted_reqs/load_kv_cache). # Callers that will recompute the KV instead (PD true-retraction rebootstrap) # pass offload_kv=False to skip the wasteful device->host copy. + backup_succeeded = True if server_args.disaggregation_mode == "decode" and offload_kv: - req.offload_kv_cache(req_to_token_pool, token_to_kv_pool_allocator) + try: + req.offload_kv_cache(req_to_token_pool, token_to_kv_pool_allocator) + except NotImplementedError: + if not abort_on_unsupported_backup: + raise + backup_succeeded = False + req.kv_cache_cpu = None + logger.error( + "CPU-tensor retraction backup is unsupported for request %s; " + "aborting the request instead of crashing the scheduler.", + req.rid, + ) # TODO (csy): for preempted requests, we may want to insert into the tree release_kv_cache(req, tree_cache, is_insert=False) # NOTE(lsyin): we should use the newly evictable memory instantly. @@ -1858,6 +1871,7 @@ def release_req( evict_from_tree_cache(tree_cache, num_tokens) req.reset_for_retract() + return backup_succeeded def retract_all( @@ -2729,6 +2743,20 @@ def check_decode_mem(self, selected_indices: Optional[List[int]] = None): whether the next decode step fits in the KV pool.""" num_tokens = self.new_tokens_required_next_decode(selected_indices) evict_from_tree_cache(self.tree_cache, num_tokens) + if self.token_to_kv_pool_allocator.available_size() >= num_tokens: + return True + + # SWA eviction is normally done while preparing the next decode batch. + # When SWA is the first limiting pool, the scheduler can reach this OOM + # check before prepare_for_decode() releases out-of-window SWA pages. + if ( + self.forward_mode is not None + and self.forward_mode.is_decode() + and self.tree_cache.supports_swa() + ): + self.maybe_evict_swa(force=True) + evict_from_tree_cache(self.tree_cache, num_tokens) + return self.token_to_kv_pool_allocator.available_size() >= num_tokens def retract_decode( @@ -2738,6 +2766,7 @@ def retract_decode( sorted_indices = self._get_decode_retraction_order(self.reqs, server_args) retracted_reqs = [] + reqs_to_abort: List[Req] = [] first_iter = True while first_iter or ( not self.check_decode_mem(selected_indices=sorted_indices) @@ -2749,11 +2778,26 @@ def retract_decode( first_iter = False idx = sorted_indices.pop() req = self.reqs[idx] - retracted_reqs.append(req) # release memory and don't insert into the tree because we need the space instantly - self.release_req(idx, len(sorted_indices), server_args) + if self.release_req( + idx, + len(sorted_indices), + server_args, + abort_on_unsupported_backup=True, + ): + retracted_reqs.append(req) + else: + req.to_finish = FINISH_ABORT( + "Retraction KV backup is unsupported. Aborting the request.", + status_code=HTTPStatus.INTERNAL_SERVER_ERROR, + ) + reqs_to_abort.append(req) + logger.warning( + "retract_decode: aborted request %s because KV backup is " + "unsupported", + req.rid, + ) - reqs_to_abort: List[Req] = [] if len(sorted_indices) <= 1 and not self.check_decode_mem( selected_indices=sorted_indices ): @@ -2767,7 +2811,7 @@ def retract_decode( status_code=HTTPStatus.INTERNAL_SERVER_ERROR, ) reqs_to_abort.append(last_req) - self.release_req(last_idx, 0, server_args) + self.release_req(last_idx, 0, server_args, offload_kv=False) logger.warning( "retract_decode: aborted last request %s due to OOM", last_req.rid ) @@ -2822,8 +2866,15 @@ def retraction_key(req: Req) -> Tuple[int, int, int]: ) return sorted_indices - def release_req(self, idx: int, remaing_req_count: int, server_args: ServerArgs): - release_req( + def release_req( + self, + idx: int, + remaing_req_count: int, + server_args: ServerArgs, + offload_kv: bool = True, + abort_on_unsupported_backup: bool = False, + ) -> bool: + return release_req( req=self.reqs[idx], remaing_req_count=remaing_req_count, server_args=server_args, @@ -2831,6 +2882,8 @@ def release_req(self, idx: int, remaing_req_count: int, server_args: ServerArgs) token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, tree_cache=self.tree_cache, hisparse_coordinator=self.hisparse_coordinator, + offload_kv=offload_kv, + abort_on_unsupported_backup=abort_on_unsupported_backup, ) def prepare_encoder_info_decode(self): @@ -3204,7 +3257,7 @@ def copy(self): extend_num_tokens=self.extend_num_tokens, ) - def maybe_evict_swa(self): + def maybe_evict_swa(self, force: bool = False): if self.tree_cache.supports_swa(): sliding_window_size = self.tree_cache.sliding_window_size server_args = get_server_args() @@ -3221,7 +3274,7 @@ def maybe_evict_swa(self): # We set evict_swa condition here with two reasons: # 1. In overlap scheduler, we cannot evict swa when req.decode_batch_idx == 0 since the prev extend batch is still running. # 2. Evict swa every eviction_interval iterations to reduce the overhead. - if swa_maintenance_step and req.decode_batch_idx >= 1: + if req.decode_batch_idx >= 1 and (force or swa_maintenance_step): self._evict_swa(req, req.seqlen - 1) # DSV4-NPU only (no-op elsewhere): the small paged compress-state diff --git a/python/sglang/srt/speculative/eagle_utils.py b/python/sglang/srt/speculative/eagle_utils.py index e87ecd91ef9e..6c8f302d5bb7 100644 --- a/python/sglang/srt/speculative/eagle_utils.py +++ b/python/sglang/srt/speculative/eagle_utils.py @@ -877,8 +877,6 @@ def eagle_sample( def eagle_prepare_for_decode(batch: ScheduleBatch): - batch.maybe_evict_swa() - bs = batch.batch_size() # Accumulate penalty @@ -909,8 +907,6 @@ def eagle_prepare_for_decode(batch: ScheduleBatch): cur_kv_lens[i] = cur nxt_kv_lens[i] = nxt num_needed_tokens += nxt - cur - r.decode_batch_idx += 1 - cur_kv_lens_cpu = torch.tensor(cur_kv_lens, dtype=torch.int32, device="cpu") nxt_kv_lens_cpu = torch.tensor(nxt_kv_lens, dtype=torch.int32, device="cpu") diff --git a/python/sglang/srt/speculative/spec_utils.py b/python/sglang/srt/speculative/spec_utils.py index 2845a2d6c5a2..82cbf8eb532c 100644 --- a/python/sglang/srt/speculative/spec_utils.py +++ b/python/sglang/srt/speculative/spec_utils.py @@ -1024,9 +1024,7 @@ def commit_mamba_states_after_verify( def spec_prepare_for_decode(batch: ScheduleBatch) -> None: - """eagle/ngram share a stateless free function; dflash keeps stateful - prep on its draft input -- the dispatcher routes. - """ + """Run common spec-v2 bookkeeping, then dispatch algorithm-specific prep.""" server_args = get_server_args() if server_args.enable_mamba_extra_buffer_lazy(): # Scheduler phase (outside forward isolation). @@ -1034,6 +1032,11 @@ def spec_prepare_for_decode(batch: ScheduleBatch) -> None: get_exec().mamba.mamba_track_interval, server_args.max_speculative_num_draft_tokens, ) + + batch.maybe_evict_swa() + for req in batch.reqs: + req.decode_batch_idx += 1 + if batch.spec_algorithm.is_dflash_family(): batch.spec_info.prepare_for_decode(batch) else: diff --git a/test/registered/unit/managers/test_schedule_batch_swa_retraction.py b/test/registered/unit/managers/test_schedule_batch_swa_retraction.py new file mode 100644 index 000000000000..834fb3d73723 --- /dev/null +++ b/test/registered/unit/managers/test_schedule_batch_swa_retraction.py @@ -0,0 +1,151 @@ +"""Unit tests for SWA reclamation and decode retraction fail-safes.""" + +import unittest +from http import HTTPStatus +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel + +maybe_stub_sgl_kernel() + +from sglang.srt.managers.schedule_batch import ( # noqa: E402 + FINISH_ABORT, + NewTokenRatioTracker, + ScheduleBatch, + release_req, +) + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +class TestScheduleBatchSwaRetraction(CustomTestCase): + def test_check_decode_mem_forces_swa_reclamation_before_retraction(self): + batch = ScheduleBatch(reqs=[]) + batch.forward_mode = MagicMock() + batch.forward_mode.is_decode.return_value = True + batch.tree_cache = MagicMock() + batch.tree_cache.supports_swa.return_value = True + batch.token_to_kv_pool_allocator = MagicMock() + batch.token_to_kv_pool_allocator.available_size.side_effect = [3, 8] + batch.new_tokens_required_next_decode = MagicMock(return_value=8) + batch.maybe_evict_swa = MagicMock() + + with patch("sglang.srt.managers.schedule_batch.evict_from_tree_cache") as evict: + self.assertTrue(batch.check_decode_mem()) + + batch.maybe_evict_swa.assert_called_once_with(force=True) + self.assertEqual(evict.call_count, 2) + + def test_force_swa_reclamation_ignores_periodic_interval(self): + req = SimpleNamespace(decode_batch_idx=1, seqlen=32) + batch = ScheduleBatch(reqs=[req]) + batch.forward_mode = MagicMock() + batch.forward_mode.is_decode.return_value = True + batch.tree_cache = MagicMock(sliding_window_size=16) + batch.tree_cache.supports_swa.return_value = True + batch.forward_iter = 1 + batch._evict_swa = MagicMock() + + with ( + patch( + "sglang.srt.managers.schedule_batch.envs.SGLANG_SWA_EVICTION_INTERVAL.get", + return_value=8, + ), + patch( + "sglang.srt.managers.schedule_batch.envs." + "SGLANG_OPT_SWA_RELEASE_LEAF_LOCK_AFTER_WINDOW.get", + return_value=False, + ), + patch("sglang.srt.managers.schedule_batch.get_server_args"), + patch("sglang.srt.managers.schedule_batch.maybe_evict_dsv4_state"), + ): + batch.maybe_evict_swa() + batch._evict_swa.assert_not_called() + + batch.maybe_evict_swa(force=True) + + batch._evict_swa.assert_called_once_with(req, req.seqlen - 1) + + def test_unsupported_backup_aborts_only_retracted_request(self): + kept_req = MagicMock(rid="kept") + aborted_req = MagicMock(rid="aborted") + batch = ScheduleBatch(reqs=[kept_req, aborted_req]) + batch._get_decode_retraction_order = MagicMock(return_value=[0, 1]) + batch.check_decode_mem = MagicMock(return_value=True) + batch.release_req = MagicMock(return_value=False) + batch.filter_batch = MagicMock() + server_args = MagicMock() + + with patch.object( + NewTokenRatioTracker, + "estimate_new_token_ratio_after_retract", + return_value=0.25, + ): + retracted, ratio, reqs_to_abort = batch.retract_decode(server_args) + + self.assertEqual(retracted, []) + self.assertEqual(ratio, 0.25) + self.assertEqual(reqs_to_abort, [aborted_req]) + self.assertIsInstance(aborted_req.to_finish, FINISH_ABORT) + self.assertEqual( + aborted_req.to_finish.status_code, HTTPStatus.INTERNAL_SERVER_ERROR + ) + batch.release_req.assert_called_once_with( + 1, + 1, + server_args, + abort_on_unsupported_backup=True, + ) + batch.filter_batch.assert_called_once_with(keep_indices=[0]) + + def test_last_request_oom_skips_unusable_backup(self): + req = MagicMock(rid="last") + batch = ScheduleBatch(reqs=[req]) + batch._get_decode_retraction_order = MagicMock(return_value=[0]) + batch.check_decode_mem = MagicMock(return_value=False) + batch.release_req = MagicMock(return_value=True) + batch.filter_batch = MagicMock() + server_args = MagicMock() + + with patch.object( + NewTokenRatioTracker, + "estimate_new_token_ratio_after_retract", + return_value=0.0, + ): + retracted, _, reqs_to_abort = batch.retract_decode(server_args) + + self.assertEqual(retracted, []) + self.assertEqual(reqs_to_abort, [req]) + batch.release_req.assert_called_once_with(0, 0, server_args, offload_kv=False) + + def test_release_req_handles_unsupported_backup_when_requested(self): + req = MagicMock(rid="unsupported") + req.finished.return_value = False + req.offload_kv_cache.side_effect = NotImplementedError + server_args = SimpleNamespace(disaggregation_mode="decode") + + with ( + patch("sglang.srt.managers.schedule_batch.release_kv_cache") as release, + patch("sglang.srt.managers.schedule_batch.evict_from_tree_cache"), + ): + backup_succeeded = release_req( + req=req, + remaing_req_count=1, + server_args=server_args, + req_to_token_pool=MagicMock(), + token_to_kv_pool_allocator=MagicMock(), + tree_cache=MagicMock(), + hisparse_coordinator=None, + abort_on_unsupported_backup=True, + ) + + self.assertFalse(backup_succeeded) + self.assertIsNone(req.kv_cache_cpu) + release.assert_called_once() + req.reset_for_retract.assert_called_once_with() + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/spec/test_decode_bookkeeping_ownership.py b/test/registered/unit/spec/test_decode_bookkeeping_ownership.py index 90a52cea5d8e..520d83b12d6b 100644 --- a/test/registered/unit/spec/test_decode_bookkeeping_ownership.py +++ b/test/registered/unit/spec/test_decode_bookkeeping_ownership.py @@ -3,9 +3,9 @@ Per-request accounting state (`decode_batch_idx` / `extend_batch_idx` iter clocks, `kv_committed_len` / `kv_allocated_len` KV watermarks, `spec_verify_ct`, and the `maybe_evict_swa()` call) must only be advanced by -the reviewed owner sites in _OWNER_SITES; spec-v2 draft workers must not -repeat any of them (the scheduler-driven free function / resolve path already -does). +the reviewed owner sites in _OWNER_SITES; spec-v2 algorithm-specific prep and +draft workers must not repeat any of them (the common scheduler-driven free +function / resolve path already does). A clock that runs fast fires SWA eviction in the overlap race window and releases the SWA prefix lock early; neither shows up in e2e CI or the idle leak checker, hence this AST-level guard. @@ -40,7 +40,7 @@ # attribute (`= 0` resets exempt) or "evict" for a `maybe_evict_swa()` call. # Any added/removed/recounted site fails until reviewed here. _SB = "managers/schedule_batch.py" -_EAGLE_DECODE = ("speculative/eagle_utils.py", "eagle_prepare_for_decode") +_SPEC_DECODE = ("speculative/spec_utils.py", "spec_prepare_for_decode") _RESOLVE = ( "managers/scheduler_components/batch_result_processor.py", "SchedulerBatchResultProcessor._resolve_spec_v2_tokens", @@ -52,16 +52,19 @@ (_SB, "ScheduleBatch.prepare_for_decode", "kv_committed_len"): 1, (_SB, "ScheduleBatch.prepare_for_extend", "extend_batch_idx"): 1, (_SB, "ScheduleBatch.prepare_for_extend", "kv_committed_len"): 1, + # Decode OOM fallback may force SWA maintenance before retraction. + (_SB, "ScheduleBatch.check_decode_mem", "evict"): 1, # kv_allocated_len is settled inside the owned-kv alloc functions (op28). ("mem_cache/allocation.py", "alloc_for_extend", "evict"): 1, ("mem_cache/allocation.py", "alloc_for_extend", "kv_allocated_len"): 1, ("mem_cache/allocation.py", "alloc_for_decode", "evict"): 1, ("mem_cache/allocation.py", "alloc_for_decode", "kv_allocated_len"): 1, - # spec v2: no pre-claim; resolve commits the full accepted run uniformly. + # spec v2: common preparation owns the decode clock and SWA maintenance for + # every algorithm; resolve commits the full accepted run uniformly. # kv_allocated_len for spec v2 draft decode (eagle + dflash) is settled # inside the owned-kv alloc_for_spec_decode function (op42). - (*_EAGLE_DECODE, "decode_batch_idx"): 1, - (*_EAGLE_DECODE, "evict"): 1, + (*_SPEC_DECODE, "decode_batch_idx"): 1, + (*_SPEC_DECODE, "evict"): 1, ( "mem_cache/allocation.py", "alloc_for_spec_decode", diff --git a/test/registered/unit/spec/test_spec_prepare_for_decode.py b/test/registered/unit/spec/test_spec_prepare_for_decode.py new file mode 100644 index 000000000000..abf740815cd7 --- /dev/null +++ b/test/registered/unit/spec/test_spec_prepare_for_decode.py @@ -0,0 +1,69 @@ +"""Unit tests for common speculative decode preparation.""" + +import unittest +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel + +maybe_stub_sgl_kernel() + +from sglang.srt.speculative.spec_utils import spec_prepare_for_decode # noqa: E402 + +register_cpu_ci(est_time=3, suite="base-a-test-cpu") + + +class TestSpecPrepareForDecode(CustomTestCase): + def _run_prepare(self, *, is_dflash_family: bool): + events = [] + req = SimpleNamespace(decode_batch_idx=4) + batch = SimpleNamespace( + reqs=[req], + spec_algorithm=SimpleNamespace(is_dflash_family=lambda: is_dflash_family), + spec_info=SimpleNamespace(), + ) + batch.maybe_evict_swa = MagicMock( + side_effect=lambda: events.append(("evict", req.decode_batch_idx)) + ) + batch.spec_info.prepare_for_decode = MagicMock( + side_effect=lambda _: events.append(("dflash", req.decode_batch_idx)) + ) + eagle_prepare = MagicMock( + side_effect=lambda _: events.append(("eagle", req.decode_batch_idx)) + ) + server_args = SimpleNamespace(enable_mamba_extra_buffer_lazy=lambda: False) + + with ( + patch( + "sglang.srt.speculative.spec_utils.get_server_args", + return_value=server_args, + ), + patch( + "sglang.srt.speculative.eagle_utils.eagle_prepare_for_decode", + eagle_prepare, + ), + ): + spec_prepare_for_decode(batch) + + return batch, eagle_prepare, events, req + + def test_dflash_gets_common_swa_bookkeeping_before_prepare(self): + batch, eagle_prepare, events, req = self._run_prepare(is_dflash_family=True) + + self.assertEqual(req.decode_batch_idx, 5) + self.assertEqual(events, [("evict", 4), ("dflash", 5)]) + batch.spec_info.prepare_for_decode.assert_called_once_with(batch) + eagle_prepare.assert_not_called() + + def test_eagle_keeps_single_common_bookkeeping_tick(self): + batch, eagle_prepare, events, req = self._run_prepare(is_dflash_family=False) + + self.assertEqual(req.decode_batch_idx, 5) + self.assertEqual(events, [("evict", 4), ("eagle", 5)]) + batch.spec_info.prepare_for_decode.assert_not_called() + eagle_prepare.assert_called_once_with(batch) + + +if __name__ == "__main__": + unittest.main() From 7da44242e2edeaba95d6165cff0060296654d69b Mon Sep 17 00:00:00 2001 From: Qi Sun <65223714+shiyu7@users.noreply.github.com> Date: Tue, 25 Aug 2026 15:35:27 +0800 Subject: [PATCH 41/47] ci: support injected values in manual image tags (#740) --- .../workflows/_docker-build-and-publish.yml | 2 +- .github/workflows/release-docker-dev.yml | 1 + scripts/ci/get_volcengine_image_tag.py | 10 ++++-- scripts/ci/test_get_volcengine_image_tag.py | 35 ++++++++++++++++++- 4 files changed, 43 insertions(+), 5 deletions(-) diff --git a/.github/workflows/_docker-build-and-publish.yml b/.github/workflows/_docker-build-and-publish.yml index 04fb6a91a928..f53b89019d35 100644 --- a/.github/workflows/_docker-build-and-publish.yml +++ b/.github/workflows/_docker-build-and-publish.yml @@ -47,7 +47,7 @@ on: type: string default: "version" tag_value: - description: "Tag value passed for version-mode Volcengine tags" + description: "Optional manual tag override; required as the value in version-mode tags" required: false type: string default: "" diff --git a/.github/workflows/release-docker-dev.yml b/.github/workflows/release-docker-dev.yml index c019b9164b7c..605ac70f841b 100644 --- a/.github/workflows/release-docker-dev.yml +++ b/.github/workflows/release-docker-dev.yml @@ -337,6 +337,7 @@ jobs: publish_default_cuda_alias: ${{ matrix.publish_default_cuda_alias }} variant_suffix: ${{ matrix.variant_suffix }} tag_mode: ${{ github.event_name == 'schedule' && 'nightly' || 'manual' }} + tag_value: ${{ inputs.tag }} use_environment: prod artifact_name: ${{ (github.event_name == 'schedule' || inputs.compile_kernel) && matrix.artifact_name || '' }} extra_build_args: >- diff --git a/scripts/ci/get_volcengine_image_tag.py b/scripts/ci/get_volcengine_image_tag.py index 3cef2a6c2b58..a794ce0fe05a 100755 --- a/scripts/ci/get_volcengine_image_tag.py +++ b/scripts/ci/get_volcengine_image_tag.py @@ -28,12 +28,14 @@ def build_tag( ) -> str: """Compose the final image tag. + Manual tags use ``tag_value`` verbatim when it is provided, otherwise they + fall back to the generated ``v.iaas.dev.`` name. Suffix order is fixed as ``variant`` -> ``cuda`` -> ``format`` so that the image format marker (e.g. ``zstd`` / ``nydus``) always trails the CUDA - marker: ``v.byted..[-][-cu130][-zstd]``. + marker: ``[-][-cu130][-zstd]``. """ if mode == "manual": - tag = f"v{version}.iaas.dev.{timestamp}" + tag = tag_value or f"v{version}.iaas.dev.{timestamp}" elif mode == "nightly": tag = f"v{version}.iaas.nightly.{timestamp}" else: @@ -86,7 +88,8 @@ def main() -> None: parser.add_argument( "--tag-value", default="", - help="Required for version mode; inserted after .byted.", + help="Used verbatim as the base tag in manual mode; required for " + "version mode and inserted after .byted.", ) parser.add_argument("--cuda-suffix", choices=["", "cu129", "cu130"], default="") parser.add_argument( @@ -102,6 +105,7 @@ def main() -> None: ) args = parser.parse_args() + validate_suffix("tag-value", args.tag_value) validate_suffix("variant-suffix", args.variant_suffix) validate_suffix("format-suffix", args.format_suffix) diff --git a/scripts/ci/test_get_volcengine_image_tag.py b/scripts/ci/test_get_volcengine_image_tag.py index 6d8548fc4ec0..1ae6e7604abf 100755 --- a/scripts/ci/test_get_volcengine_image_tag.py +++ b/scripts/ci/test_get_volcengine_image_tag.py @@ -33,6 +33,31 @@ def test_manual_and_nightly_prefixes(self) -> None: "v0.5.17.iaas.nightly.202608121200", ) + def test_manual_tag_with_injected_value(self) -> None: + self.assertEqual( + build_tag( + mode="manual", + version="0.5.17", + timestamp="202608121200", + tag_value="day0_deepseek_v4_official", + ), + "day0_deepseek_v4_official", + ) + + def test_manual_tag_with_injected_value_and_suffixes(self) -> None: + self.assertEqual( + build_tag( + mode="manual", + version="0.5.17", + timestamp="202608121200", + tag_value="day0_deepseek_v4_official", + variant_suffix="w4a8", + cuda_suffix="cu130", + format_suffix="nydus", + ), + "day0_deepseek_v4_official-w4a8-cu130-nydus", + ) + def test_format_suffix_trails_cuda_suffix(self) -> None: self.assertEqual( build_tag( @@ -88,7 +113,15 @@ def test_version_mode_requires_tag_value(self) -> None: build_tag(mode="version", version="0.5.17", timestamp="202608121200") def test_validate_suffix_accepts_safe_values(self) -> None: - for value in ("zstd", "nydus", "cu130", "w4a8", "deepseek-v4", ""): + for value in ( + "zstd", + "nydus", + "cu130", + "w4a8", + "deepseek-v4", + "day0_deepseek_v4_official", + "", + ): validate_suffix("format-suffix", value) def test_validate_suffix_rejects_unsafe_values(self) -> None: From 8644ab11dc58d9054b95fcc49b7397e6abaa3e81 Mon Sep 17 00:00:00 2001 From: Hank Han Date: Thu, 27 Aug 2026 10:59:31 +0800 Subject: [PATCH 42/47] ci: default unknown historical image versions Co-authored-by: TRAE CLI --- .github/workflows/backfill-image-formats.yml | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/.github/workflows/backfill-image-formats.yml b/.github/workflows/backfill-image-formats.yml index 393a980812d3..c9664f6eade9 100644 --- a/.github/workflows/backfill-image-formats.yml +++ b/.github/workflows/backfill-image-formats.yml @@ -54,6 +54,7 @@ jobs: VOLCENGINE_CR_USERNAME: ${{ secrets.VOLCENGINE_CR_USERNAME }} VOLCENGINE_CR_PASSWORD: ${{ secrets.VOLCENGINE_CR_PASSWORD }} NYDUS_VERSION: "2.3.0" + DEFAULT_SGLANG_VERSION: "0.5.18" steps: - name: Delete huge unnecessary tools folder run: rm -rf /opt/hostedtoolcache @@ -201,9 +202,10 @@ jobs: fi if [ -z "${version}" ]; then - echo " ERROR: unable to determine a non-placeholder sglang version from the tag, container, build commit, or image metadata" >&2 + version="${DEFAULT_SGLANG_VERSION}" + echo " WARNING: unable to determine a non-placeholder sglang version from the tag, container, build commit, or image metadata" >&2 echo " Source image: ${source_ref}" >&2 - return 1 + echo " Falling back to default sglang version: ${version}" >&2 fi fi echo " Using sglang version: ${version}" >&2 From e2910440ec55bd26e11d032e44be8dcbd9529ad1 Mon Sep 17 00:00:00 2001 From: Hank Han Date: Thu, 27 Aug 2026 11:00:10 +0800 Subject: [PATCH 43/47] ci: use dev sentinel for unknown image versions Co-authored-by: TRAE CLI --- .github/workflows/backfill-image-formats.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/backfill-image-formats.yml b/.github/workflows/backfill-image-formats.yml index c9664f6eade9..164d5248c323 100644 --- a/.github/workflows/backfill-image-formats.yml +++ b/.github/workflows/backfill-image-formats.yml @@ -54,7 +54,7 @@ jobs: VOLCENGINE_CR_USERNAME: ${{ secrets.VOLCENGINE_CR_USERNAME }} VOLCENGINE_CR_PASSWORD: ${{ secrets.VOLCENGINE_CR_PASSWORD }} NYDUS_VERSION: "2.3.0" - DEFAULT_SGLANG_VERSION: "0.5.18" + DEFAULT_SGLANG_VERSION: "0.0.0dev" steps: - name: Delete huge unnecessary tools folder run: rm -rf /opt/hostedtoolcache From 0455119f0a32092b6c06abb5f8f32ec611fdbf75 Mon Sep 17 00:00:00 2001 From: Hank Han Date: Mon, 31 Aug 2026 16:51:22 +0800 Subject: [PATCH 44/47] ci: support exact image normalization before backfill Co-authored-by: TRAE CLI --- .github/workflows/backfill-image-formats.yml | 51 ++++++++++++++++++-- 1 file changed, 48 insertions(+), 3 deletions(-) diff --git a/.github/workflows/backfill-image-formats.yml b/.github/workflows/backfill-image-formats.yml index 164d5248c323..812cf56ecc04 100644 --- a/.github/workflows/backfill-image-formats.yml +++ b/.github/workflows/backfill-image-formats.yml @@ -14,9 +14,9 @@ on: default: "" source_image_ref: description: | - 可选:当 auto_rename=true 时,以此为源镜像(而非 image_refs)进行重命名。 - 留空则使用 image_refs 中的镜像作为源。 - 格式: iaas-gpu-cn-beijing.cr.volces.com/serving/: + 可选源镜像。auto_rename=true 时从该镜像生成标准 tag; + auto_rename=false 时将它复制到 image_refs 指定的唯一目标,再补全格式。 + 可用于把旧 repo 中的历史镜像规范化到 serving/:。 required: false type: string default: "" @@ -68,6 +68,9 @@ jobs: - name: Set up Docker Buildx uses: docker/setup-buildx-action@v3 + - name: Set up crane + uses: imjasonh/setup-crane@v0.4 + - name: Login to Volcengine CR run: | set -euo pipefail @@ -97,6 +100,18 @@ jobs: echo "ERROR: image_refs is required when auto_rename is disabled" exit 1 fi + if [ "${{ inputs.auto_rename }}" = "false" ] && [ -n "${{ inputs.source_image_ref }}" ]; then + target_count=0 + while IFS= read -r line; do + line="${line//[$'\r\n']/}" + [[ -z "$line" || "$line" == \#* ]] && continue + target_count=$((target_count + 1)) + done <<< "${{ inputs.image_refs }}" + if [ "${target_count}" -ne 1 ]; then + echo "ERROR: exact-copy mode requires exactly one image_refs target" + exit 1 + fi + fi echo "Input validation passed." - name: Auto-rename and backfill @@ -261,6 +276,27 @@ jobs: fi } + copy_and_backfill() { + local source_ref="$1" + local target_ref="$2" + local allowed_prefix="${VOLCENGINE_CR_REGISTRY}/${VOLCENGINE_CR_NAMESPACE}/" + + if [[ "${target_ref}" != "${allowed_prefix}"* ]]; then + echo "ERROR: exact-copy target must start with ${allowed_prefix}" + return 1 + fi + + echo "" + echo "=== $(date '+%H:%M:%S') Normalize: ${source_ref} -> ${target_ref} ===" + if docker manifest inspect "${target_ref}" >/dev/null 2>&1; then + echo " Target already exists, using existing: ${target_ref}" + else + retry crane copy "${source_ref}" "${target_ref}" + echo " Copied: ${target_ref}" + fi + backfill_one "${target_ref}" + } + backfill_one() { local image_ref="$1" @@ -319,6 +355,15 @@ jobs: [[ -z "$line" || "$line" == \#* ]] && continue rename_and_backfill "$line" done <<< "${IMAGE_REFS}" + elif [ -n "${SOURCE_IMAGE_REF}" ]; then + echo " Exact copy to image_refs target, then direct backfill" + target_ref="" + while IFS= read -r line; do + line="${line//[$'\r\n']/}" + [[ -z "$line" || "$line" == \#* ]] && continue + target_ref="$line" + done <<< "${IMAGE_REFS}" + copy_and_backfill "${SOURCE_IMAGE_REF}" "${target_ref}" else echo " Direct backfill (no rename)" while IFS= read -r line; do From d5349c46692a0822169579ac55c9f431011abbe2 Mon Sep 17 00:00:00 2001 From: Hank Han Date: Mon, 31 Aug 2026 16:54:21 +0800 Subject: [PATCH 45/47] ci: bypass proxy for internal registry copy Co-authored-by: TRAE CLI --- .github/workflows/backfill-image-formats.yml | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/.github/workflows/backfill-image-formats.yml b/.github/workflows/backfill-image-formats.yml index 812cf56ecc04..7701fd718e54 100644 --- a/.github/workflows/backfill-image-formats.yml +++ b/.github/workflows/backfill-image-formats.yml @@ -291,7 +291,14 @@ jobs: if docker manifest inspect "${target_ref}" >/dev/null 2>&1; then echo " Target already exists, using existing: ${target_ref}" else - retry crane copy "${source_ref}" "${target_ref}" + # Both registries are reachable directly from the self-hosted + # runner. Its inherited proxy returns a gateway timeout for the + # production CR endpoint, so keep this registry-to-registry copy + # on the internal route. + retry env \ + HTTPS_PROXY= HTTP_PROXY= ALL_PROXY= \ + https_proxy= http_proxy= all_proxy= \ + crane copy "${source_ref}" "${target_ref}" echo " Copied: ${target_ref}" fi backfill_one "${target_ref}" From b709718c6a8d341526313765949ddafff860983c Mon Sep 17 00:00:00 2001 From: TobyMint Date: Tue, 1 Sep 2026 10:37:27 +0800 Subject: [PATCH 46/47] [qwen3_5] P0 correctness fixes for MTP verify paths (from sgl-project) - qwen3_5_mtp: captured prefill pads embeddings while target hidden states keep real height; graft the real rows into an equal-height slot before cat+fc instead of letting cat broadcast-fail or misalign (upstream qwen3_5_mtp forward padding fix). - gdn_backend: honor --linear-attn-verify-backend. The dispatcher re-derived the verify kernel with the auto rule and ignored the stored choice, so an explicit triton selection (required when --mamba-ssm-dtype bfloat16 meets FlashInfer SM90 verify's fp32-state requirement, upstream #36611) had no effect. - linear/utils: raise when --enable-deterministic-inference combines with a FlashInfer GDN prefill (upstream _validate_gdn_linear_attn_backends). Validated on H20 TP1 with NEXTN (steps 3 / topk 1 / draft 4): boots, gsm8k 200q = 0.975 (bf16 baseline 0.980), accept length 3.55-3.60. --- .../layers/attention/linear/gdn_backend.py | 20 ++++++++++++++----- .../srt/layers/attention/linear/utils.py | 7 +++++++ python/sglang/srt/models/qwen3_5_mtp.py | 10 ++++++++++ 3 files changed, 32 insertions(+), 5 deletions(-) diff --git a/python/sglang/srt/layers/attention/linear/gdn_backend.py b/python/sglang/srt/layers/attention/linear/gdn_backend.py index 7539cc1cdb67..c9dba2d43506 100644 --- a/python/sglang/srt/layers/attention/linear/gdn_backend.py +++ b/python/sglang/srt/layers/attention/linear/gdn_backend.py @@ -15,6 +15,7 @@ build_verify_intermediate_state_indices, get_linear_attn_decode_backend, get_linear_attn_prefill_backend, + get_linear_attn_verify_backend, ) from sglang.srt.layers.radix_linear_attention import RadixLinearAttention from sglang.srt.mem_cache.memory_pool import MambaPool @@ -110,6 +111,7 @@ def __init__( self, decode_backend: LinearAttnKernelBackend, prefill_backend: LinearAttnKernelBackend, + verify_backend: Optional[LinearAttnKernelBackend] = None, ): triton_kernel = TritonGDNKernel() self.tree_verify_kernel = triton_kernel @@ -177,10 +179,16 @@ def __init__( else: raise ValueError(f"Unsupported GDN prefill backend: {prefill_backend}") - # Verify kernel: use FlashInfer when the selected FlashInfer kernel - # supports MTP verify. SM90 uses the fp32-state path; SM100 uses the - # bf16-state adapter in FlashInferGDNKernel. - if ( + # Verify kernel. An explicitly configured --linear-attn-verify-backend + # (triton) wins; the historical auto rule (FlashInfer when the selected + # FlashInfer kernel supports MTP verify) only applies when that choice + # was not made. SM90 FlashInfer verify requires a fp32 SSM state, so + # e.g. --mamba-ssm-dtype bfloat16 setups must be able to force Triton + # here (same class of fix as upstream #36611 for H200). + if verify_backend is not None and verify_backend.is_triton(): + self.verify_kernel = triton_kernel + self.verify_kernel_is_flashinfer = False + elif ( decode_backend.is_flashinfer() or prefill_backend.is_flashinfer() ) and flashinfer_kernel.supports_target_verify: self.verify_kernel = flashinfer_kernel @@ -346,7 +354,9 @@ def __init__(self, model_runner: ModelRunner): decode_backend = get_linear_attn_decode_backend() prefill_backend = get_linear_attn_prefill_backend() - self.kernel_dispatcher = GDNKernelDispatcher(decode_backend, prefill_backend) + self.kernel_dispatcher = GDNKernelDispatcher( + decode_backend, prefill_backend, get_linear_attn_verify_backend() + ) # Sized past the pool for attn_tp-padded warmup/MLP-sync batches (see helper). self.verify_intermediate_state_indices = ( build_verify_intermediate_state_indices( diff --git a/python/sglang/srt/layers/attention/linear/utils.py b/python/sglang/srt/layers/attention/linear/utils.py index 99bdc07f1a70..f74f0c1bcb48 100644 --- a/python/sglang/srt/layers/attention/linear/utils.py +++ b/python/sglang/srt/layers/attention/linear/utils.py @@ -68,6 +68,13 @@ def initialize_linear_attn_config( _BACKENDS["decode"] = LinearAttnKernelBackend(decode) _BACKENDS["prefill"] = LinearAttnKernelBackend(prefill) + if server_args.enable_deterministic_inference and _BACKENDS["prefill"].is_flashinfer(): + raise ValueError( + "FlashInfer GDN prefill is not supported with " + "--enable-deterministic-inference. Use " + "--linear-attn-prefill-backend triton." + ) + # Verify backend. Unset -> follow decode (flashinfer -> its recurrent kernel, # else triton), preserving historical behavior. verify = server_args.linear_attn_verify_backend diff --git a/python/sglang/srt/models/qwen3_5_mtp.py b/python/sglang/srt/models/qwen3_5_mtp.py index ded5ef3b8835..d2c66884ebc5 100644 --- a/python/sglang/srt/models/qwen3_5_mtp.py +++ b/python/sglang/srt/models/qwen3_5_mtp.py @@ -189,6 +189,16 @@ def forward( if not forward_batch.forward_mode.is_idle(): input_embeds = self.pre_fc_norm_embedding(input_embeds) hidden_states = self.pre_fc_norm_hidden(hidden_states) + # Captured prefill gives padded embeddings but real-height target states; + # place the real rows in an equal-height slot whose padding stays unread. + if hidden_states.shape[0] != input_embeds.shape[0]: + rows = min(hidden_states.shape[0], input_embeds.shape[0]) + slot = hidden_states.new_zeros( + (input_embeds.shape[0], hidden_states.shape[1]) + ) + slot[:rows] = hidden_states[:rows] + hidden_states = slot + hidden_states = torch.cat([input_embeds, hidden_states], dim=-1) hidden_states = self.fc(hidden_states) From 2711db463082fabba89120bfd055906c9e8d0785 Mon Sep 17 00:00:00 2001 From: TobyMint Date: Tue, 1 Sep 2026 10:49:29 +0800 Subject: [PATCH 47/47] [qwen3_5] Fix mrope axis handling in fused QK-norm+RoPE+gate kernel (#34446) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The fused kernel loaded a single position per token, so with image inputs ([3, T] temporal/height/width positions) every rotary lane silently read the temporal row — wrong RoPE on image tokens in all full-attention layers of Qwen3.5/3.8 hybrids. Text was unaffected (the three rows coincide), so this never shows in text-only smoke. Port: the kernel takes an mrope_axis_map ([rotary_dim//2] lane->axis) and reads positions[axis[lane], t] when positions is 2-D; MRotaryEmbedding now builds the axis map for every mrope_section style (contiguous, interleaved, GLM round-robin) instead of GLM only, while the legacy sgl_kernel call sites keep the GLM-only map via _legacy_axis_map. Unit test mirrors the kernel math bitwise for 1-D and mrope positions and checks both axis-map styles. --- .../attention/fused_qk_rmsnorm_rope_gate.py | 35 +++- .../srt/layers/rotary_embedding/mrope.py | 75 ++++--- python/sglang/srt/models/qwen3_5.py | 1 + .../test_fused_qk_rmsnorm_rope_gate_mrope.py | 197 ++++++++++++++++++ 4 files changed, 271 insertions(+), 37 deletions(-) create mode 100644 test/registered/unit/layers/attention/test_fused_qk_rmsnorm_rope_gate_mrope.py diff --git a/python/sglang/kernels/ops/attention/fused_qk_rmsnorm_rope_gate.py b/python/sglang/kernels/ops/attention/fused_qk_rmsnorm_rope_gate.py index 56d442098c02..f84a07eaa6e9 100644 --- a/python/sglang/kernels/ops/attention/fused_qk_rmsnorm_rope_gate.py +++ b/python/sglang/kernels/ops/attention/fused_qk_rmsnorm_rope_gate.py @@ -1,7 +1,7 @@ """Fused Q/K GemmaRMSNorm + NeoX RoPE + gate deinterleave (Triton). -Single kernel launch fusing per-head GemmaRMSNorm, partial NeoX RoPE, -and gate deinterleave for Qwen3.5's interleaved Q+Gate layout. +Single kernel launch fusing per-head GemmaRMSNorm, partial NeoX RoPE over 1-D or +mrope positions, and gate deinterleave for Qwen3.5's interleaved Q+Gate layout. 2D grid (T, num_q_heads + num_kv_heads) — each program handles one (token, head) pair. Q programs also copy the gate slice. @@ -39,12 +39,14 @@ def _fused_qk_rmsnorm_rope_gate_kernel( k_weight_ptr, cos_sin_cache_ptr, positions_ptr, + mrope_axis_map_ptr, stride_qg_t, stride_k_t, stride_qo_t, stride_ko_t, stride_gate_t, stride_cos_t, + stride_pos_axis, NUM_Q_HEADS: tl.constexpr, NUM_KV_HEADS: tl.constexpr, HEAD_DIM: tl.constexpr, @@ -56,6 +58,7 @@ def _fused_qk_rmsnorm_rope_gate_kernel( FP16: tl.constexpr, HAS_PASS: tl.constexpr, HAS_GATE: tl.constexpr, + MROPE: tl.constexpr, ENABLE_PDL: tl.constexpr, ): token = tl.program_id(0) @@ -104,8 +107,14 @@ def _fused_qk_rmsnorm_rope_gate_kernel( xr1 = (xr1 * inv_rms * (wr1 + 1.0)).to(out_dtype).to(tl.float32) xr2 = (xr2 * inv_rms * (wr2 + 1.0)).to(out_dtype).to(tl.float32) - pos = tl.load(positions_ptr + token).to(tl.int64) - cache_off = pos * stride_cos_t + if MROPE: + axis = tl.load(mrope_axis_map_ptr + rot_offs, mask=rot_mask, other=0) + pos = tl.load( + positions_ptr + axis * stride_pos_axis + token, mask=rot_mask, other=0 + ) + else: + pos = tl.load(positions_ptr + token) + cache_off = pos.to(tl.int64) * stride_cos_t cos = tl.load( cos_sin_cache_ptr + cache_off + rot_offs, mask=rot_mask, other=0.0 ).to(tl.float32) @@ -141,6 +150,7 @@ def fused_qk_gemma_rmsnorm_rope_gate( head_dim: int, rotary_dim: int, has_gate: bool = True, + mrope_axis_map: Optional[torch.Tensor] = None, ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: """Fused QK GemmaRMSNorm + NeoX RoPE + gate deinterleave. @@ -149,8 +159,20 @@ def fused_qk_gemma_rmsnorm_rope_gate( k: [T, num_kv_heads * head_dim] q_weight, k_weight: [head_dim] — raw GemmaRMSNorm weights (kernel adds +1.0) cos_sin_cache: [max_seq_len, rotary_dim] — [cos..., sin...] - positions: [T] — token positions + positions: [T] token positions, or [3, T] mrope rows (temporal, height, width) + mrope_axis_map: [rotary_dim // 2] — the axis owning each rotary lane, from + MRotaryEmbedding """ + assert positions.dim() in (1, 2), f"want [T] or [3, T], got {positions.shape}" + mrope = positions.dim() == 2 + assert mrope == (mrope_axis_map is not None), "mrope_axis_map needs [3, T]" + if mrope: + assert positions.shape[0] == 3 and positions.stride(1) == 1, ( + f"want [3, T] contiguous over T, got {positions.shape} " + f"stride {positions.stride()}" + ) + lanes = rotary_dim // 2 + assert mrope_axis_map.shape == (lanes,), f"want one axis per lane ({lanes})" T = q_gate.shape[0] q_size = num_q_heads * head_dim kv_size = num_kv_heads * head_dim @@ -178,12 +200,14 @@ def fused_qk_gemma_rmsnorm_rope_gate( k_weight, cos_sin_cache, positions, + mrope_axis_map, q_gate.stride(0), k.stride(0), q_out.stride(0), k_out.stride(0), gate_out.stride(0), cos_sin_cache.stride(0), + positions.stride(0), NUM_Q_HEADS=num_q_heads, NUM_KV_HEADS=num_kv_heads, HEAD_DIM=head_dim, @@ -195,6 +219,7 @@ def fused_qk_gemma_rmsnorm_rope_gate( FP16=q_gate.dtype == torch.float16, HAS_PASS=rotary_dim < head_dim, HAS_GATE=has_gate, + MROPE=mrope, ENABLE_PDL=_ENABLE_PDL, ) diff --git a/python/sglang/srt/layers/rotary_embedding/mrope.py b/python/sglang/srt/layers/rotary_embedding/mrope.py index a49246802566..631922724d77 100644 --- a/python/sglang/srt/layers/rotary_embedding/mrope.py +++ b/python/sglang/srt/layers/rotary_embedding/mrope.py @@ -100,39 +100,50 @@ def __init__( f"Corrected mrope_section: {self.mrope_section} (sum={sum(self.mrope_section)})" ) - # MRoPE axis_map interleaving pattern depends on mrope_section sizes. - # The algorithm cycles through axes [0(T), 1(H), 2(W)] round-robin, - # skipping any axis that has exhausted its allocated pairs. - # - # For GLM-V (mrope_section=[8,12,12]): - # T(8) < H(12) = W(12), so T exhausts first at pair 24. - # Result: [0,1,2, 0,1,2, 0,1,2, 0,1,2, 0,1,2, 0,1,2, 0,1,2, 0,1,2, 1,1,2, 1,1,2, 2,2] - # After T runs out, only H and W fill the remaining slots. - # - # For Qwen3-VL (mrope_section=[24,20,20]): - # T(24) > H(20) = W(20), so H and W exhaust first near the tail. - # Result: [0,1,2, 0,1,2, ...repeated evenly..., 0,1, 0,1, 0,0] - # After H/W run out, T fills the remaining slots. + self.register_buffer("axis_map", self._build_axis_map(), persistent=False) + if get_exec().deterministic.rl_on_policy_target is not None: + self._forward_method = self.forward_native + def _build_axis_map(self) -> Optional[torch.Tensor]: + """Which of the temporal, height and width axes owns each rotary lane. + + The interleaving pattern depends on mrope_section sizes: + - GLM-V interleaved-glm (mrope_section=[8,12,12]): round-robin over + axes, skipping exhausted ones. + - Interleaved: lanes 0..2 cycle T/H/W, each axis owning every third + lane up to its section. + - Standard contiguous (e.g. Qwen3-VL mrope_section=[24,20,20]): + [0]*s0 + [1]*s1 + [2]*s2. + """ + if not self.mrope_section: + return None + section = self.mrope_section + num_pairs = self.rotary_dim // 2 + assert ( + len(section) == 3 and sum(section) == num_pairs + ), f"mrope_section {section} must be three axes summing to {num_pairs}" if self.mrope_interleaved_glm: - num_pairs = rotary_dim // 2 - axis_map = torch.empty(num_pairs, dtype=torch.long) - assert sum(self.mrope_section) == num_pairs - counts = [0, 0, 0] - current_ax = 0 - - for i in range(num_pairs): - current_ax = i % 3 - while counts[current_ax] >= self.mrope_section[current_ax]: - current_ax = (current_ax + 1) % 3 - - axis_map[i] = current_ax - counts[current_ax] += 1 - self.register_buffer("axis_map", axis_map, persistent=False) + axes = [] + spent = [0, 0, 0] + for lane in range(num_pairs): + axis = lane % 3 + while spent[axis] >= section[axis]: + axis = (axis + 1) % 3 + spent[axis] += 1 + axes.append(axis) + elif self.mrope_interleaved: + axes = [0] * num_pairs + for axis in (1, 2): + for lane in range(axis, min(3 * section[axis], num_pairs), 3): + axes[lane] = axis else: - self.axis_map = None - if get_exec().deterministic.rl_on_policy_target is not None: - self._forward_method = self.forward_native + axes = [axis for axis, size in enumerate(section) for _ in range(size)] + return torch.tensor(axes, dtype=torch.long, device=self.cos_sin_cache.device) + + @property + def _legacy_axis_map(self) -> Optional[torch.Tensor]: + """The map only where the older rope kernels read it; one is out of tree.""" + return self.axis_map if self.mrope_interleaved_glm else None def get_cos_sin_with_position(self, positions): if positions.ndim == 1: @@ -266,7 +277,7 @@ def forward_triton( self.mrope_interleaved, self.mrope_interleaved_glm, self.is_neox_style, - self.axis_map, + self._legacy_axis_map, ) return query, key @@ -316,7 +327,7 @@ def forward_xpu( self.mrope_interleaved, self.mrope_interleaved_glm, self.is_neox_style, - self.axis_map, + self._legacy_axis_map, ) return query, key return self.forward_native(positions, query, key, fused_set_kv_buffer_arg) diff --git a/python/sglang/srt/models/qwen3_5.py b/python/sglang/srt/models/qwen3_5.py index 08c47b9942b5..2650dacec1e2 100644 --- a/python/sglang/srt/models/qwen3_5.py +++ b/python/sglang/srt/models/qwen3_5.py @@ -1037,6 +1037,7 @@ def forward_prepare_cuda_fused(self, positions, hidden_states): self.head_dim, self.rotary_emb.rotary_dim, has_gate=self.attn_output_gate, + mrope_axis_map=(self.rotary_emb.axis_map if positions.dim() == 2 else None), ) seq_len = hidden_states.shape[0] q = q_out.view(seq_len, -1) diff --git a/test/registered/unit/layers/attention/test_fused_qk_rmsnorm_rope_gate_mrope.py b/test/registered/unit/layers/attention/test_fused_qk_rmsnorm_rope_gate_mrope.py new file mode 100644 index 000000000000..ce755b8c4592 --- /dev/null +++ b/test/registered/unit/layers/attention/test_fused_qk_rmsnorm_rope_gate_mrope.py @@ -0,0 +1,197 @@ +"""Fused QK GemmaRMSNorm+RoPE+gate kernel: 1-D and mrope [3, T] positions. + +The mrope branch (ported from sgl-project #34446) indexes cos/sin per rotary +lane through the MRotaryEmbedding axis map; before it, every lane silently +read the temporal row, corrupting RoPE on image tokens in every full-attention +layer of Qwen3.5/3.8 hybrids (text was unaffected since all three rows match). +""" + +import unittest + +import torch + +from sglang.kernels.ops.attention.fused_qk_rmsnorm_rope_gate import ( + fused_qk_gemma_rmsnorm_rope_gate, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=5, stage="base-b", runner_config="1-gpu-large") + + +def _reference( + q_gate, + k, + q_weight, + k_weight, + cos_sin_cache, + positions, + eps, + num_q_heads, + num_kv_heads, + head_dim, + rotary_dim, + axis_map, +): + """Torch mirror of the kernel math, including its bf16 round-trips.""" + out_dtype = q_gate.dtype + + def norm(x, w): + x = x.to(torch.float32) + w = w.to(torch.float32) + var = (x * x).sum(-1, keepdim=True) / x.shape[-1] + inv_rms = torch.rsqrt(var + eps) + return (x * inv_rms * (w + 1.0)).to(out_dtype).to(torch.float32) + + T = q_gate.shape[0] + q = q_gate.view(T, num_q_heads, 2 * head_dim)[..., :head_dim] + k = k.view(T, num_kv_heads, head_dim) + qn = norm(q, q_weight) + kn = norm(k, k_weight) + + half = rotary_dim // 2 + if positions.dim() == 1: + # One position per token: every lane reads the same cache row. + pos = positions.to(torch.long).view(T) + cos = cos_sin_cache[pos, :half].to(torch.float32).unsqueeze(1) + sin = cos_sin_cache[pos, half : 2 * half].to(torch.float32).unsqueeze(1) + else: + # Per-lane positions: lane l of token t reads cache[pos[t, l], l]. + pos = positions.index_select(0, axis_map).t().to(torch.long) + lanes = torch.arange(half, device=cos_sin_cache.device) + cos = cos_sin_cache[pos, lanes].to(torch.float32).unsqueeze(1) + sin = cos_sin_cache[pos, half + lanes].to(torch.float32).unsqueeze(1) + + def rope(xn): + x1, x2 = xn[..., :half], xn[..., half:rotary_dim] + return torch.cat( + [x1 * cos - x2 * sin, x2 * cos + x1 * sin, xn[..., rotary_dim:]], + dim=-1, + ).to(out_dtype) + + q_out = rope(qn) + k_out = rope(kn) + return q_out, k_out + + +class TestFusedQKRmsnormRopeGateMrope(unittest.TestCase): + """The fused kernel must match the torch reference for 1-D and mrope + positions, and MRotaryEmbedding must build the right axis map per style.""" + + HEAD_DIM = 256 + ROTARY_DIM = 64 + NUM_Q_HEADS = 8 + NUM_KV_HEADS = 2 + EPS = 1e-6 + + def setUp(self): + if not torch.cuda.is_available(): + self.skipTest("Test requires CUDA") + from sglang.srt.runtime_context import publish + from sglang.srt.server_args import ServerArgs + + # A real local model dir keeps ServerArgs resolution offline and fast; + # publish() is required because MRotaryEmbedding reads the exec bag. + publish(ServerArgs(model_path="/data02/models/Qwen3.8-27B"), role="test") + torch.manual_seed(0) + self.device = "cuda" + + def _run_case(self, positions, axis_map): + T = positions.shape[-1] + q_size = self.NUM_Q_HEADS * self.HEAD_DIM + kv_size = self.NUM_KV_HEADS * self.HEAD_DIM + q_gate = torch.randn( + T, q_size * 2, dtype=torch.bfloat16, device=self.device + ) + k = torch.randn(T, kv_size, dtype=torch.bfloat16, device=self.device) + q_weight = torch.randn(self.HEAD_DIM, dtype=torch.bfloat16, device=self.device) + k_weight = torch.randn(self.HEAD_DIM, dtype=torch.bfloat16, device=self.device) + cos_sin_cache = torch.randn( + 4096, self.ROTARY_DIM, dtype=torch.float32, device=self.device + ) + + q_out, k_out, gate_out = fused_qk_gemma_rmsnorm_rope_gate( + q_gate, + k, + q_weight, + k_weight, + cos_sin_cache, + positions, + self.EPS, + self.NUM_Q_HEADS, + self.NUM_KV_HEADS, + self.HEAD_DIM, + self.ROTARY_DIM, + has_gate=True, + mrope_axis_map=axis_map, + ) + + ref_q, ref_k = _reference( + q_gate, + k, + q_weight, + k_weight, + cos_sin_cache, + positions, + self.EPS, + self.NUM_Q_HEADS, + self.NUM_KV_HEADS, + self.HEAD_DIM, + self.ROTARY_DIM, + axis_map, + ) + torch.testing.assert_close(q_out, ref_q.view_as(q_out), rtol=0, atol=0) + torch.testing.assert_close(k_out, ref_k.view_as(k_out), rtol=0, atol=0) + # Gate is a pure copy of the interleaved Q+Gate tail. + gate_ref = q_gate.view(T, self.NUM_Q_HEADS, 2 * self.HEAD_DIM)[ + ..., self.HEAD_DIM : + ] + torch.testing.assert_close(gate_out, gate_ref, rtol=0, atol=0) + + def test_1d_positions(self): + positions = torch.randint(0, 4096, (33,), device=self.device) + self._run_case(positions, None) + + def test_mrope_positions(self): + # Contiguous-style map: lanes 0..1..2 own T/H/W in section order. + half = self.ROTARY_DIM // 2 + axis_map = torch.tensor( + [0] * 12 + [1] * 10 + [2] * 10, dtype=torch.long, device=self.device + ) + assert axis_map.numel() == half + positions = torch.stack( + [ + torch.randint(0, 4096, (33,), device=self.device), + torch.randint(0, 512, (33,), device=self.device), + torch.randint(0, 512, (33,), device=self.device), + ] + ) + self._run_case(positions, axis_map) + + def test_axis_map_styles(self): + from sglang.srt.layers.rotary_embedding.mrope import MRotaryEmbedding + + def build(rotary_dim, **kw): + return MRotaryEmbedding( + head_size=64, + rotary_dim=rotary_dim, + max_position_embeddings=4096, + base=1000000, + is_neox_style=True, + dtype=torch.bfloat16, + **kw, + ) + + # Standard contiguous (Qwen3-VL style): sections must sum to + # rotary_dim // 2 or the constructor rescales them. + r = build(64, mrope_section=[12, 10, 10]) + self.assertEqual(r.axis_map.tolist(), [0] * 12 + [1] * 10 + [2] * 10) + self.assertIsNone(r._legacy_axis_map) + + # GLM interleaved: round-robin skipping exhausted axes. + glm = build(16, mrope_section=[2, 3, 3], mrope_interleaved_glm=True) + self.assertEqual(glm._legacy_axis_map.tolist(), glm.axis_map.tolist()) + self.assertEqual(len(glm.axis_map), 8) + + +if __name__ == "__main__": + unittest.main()