From 700336bfe685c84cc9c6cdbca3876a74127089d4 Mon Sep 17 00:00:00 2001 From: Zhangheng Date: Sat, 16 May 2026 00:40:55 +0800 Subject: [PATCH 01/50] [UnifiedTree]: Fix UnifiedRadixCache device match semantics with HiCache (#25277) --- .../full_component.py | 11 +- .../mamba_component.py | 12 +- .../unified_cache_components/swa_component.py | 6 +- .../tree_component.py | 8 +- .../srt/mem_cache/unified_radix_cache.py | 184 +++++------ .../test_unified_radix_hicache_kl.py | 2 +- .../test_unified_radix_cache_unittest.py | 285 ++++++++++++++++-- 7 files changed, 398 insertions(+), 110 deletions(-) diff --git a/python/sglang/srt/mem_cache/unified_cache_components/full_component.py b/python/sglang/srt/mem_cache/unified_cache_components/full_component.py index 3263cfbd1092..bccb866a7c07 100644 --- a/python/sglang/srt/mem_cache/unified_cache_components/full_component.py +++ b/python/sglang/srt/mem_cache/unified_cache_components/full_component.py @@ -41,8 +41,15 @@ def __init__(self, cache, params): # HiCache state: set to host KV pool when HiCache enabled self._full_kv_pool_host = None - def create_match_validator(self) -> Callable[[UnifiedTreeNode], bool]: - # HiCache: evicted + backuped nodes are valid match boundaries + def create_match_validator( + self, match_device_only: bool = False + ) -> Callable[[UnifiedTreeNode], bool]: + if match_device_only: + return ( + lambda node: node.component_data[self.component_type].value is not None + ) + + # HiCache: evicted + backuped nodes are valid match boundaries. return lambda node: ( node.component_data[self.component_type].value is not None or node.backuped ) diff --git a/python/sglang/srt/mem_cache/unified_cache_components/mamba_component.py b/python/sglang/srt/mem_cache/unified_cache_components/mamba_component.py index c1ef99b88406..7142bd9b3846 100644 --- a/python/sglang/srt/mem_cache/unified_cache_components/mamba_component.py +++ b/python/sglang/srt/mem_cache/unified_cache_components/mamba_component.py @@ -50,8 +50,13 @@ def __init__(self, cache: UnifiedRadixCache, params: CacheInitParams): # HiCache state self._mamba_pool_host = None # set to host mamba pool when HiCache enabled - def create_match_validator(self) -> Callable[[UnifiedTreeNode], bool]: + def create_match_validator( + self, match_device_only: bool = False + ) -> Callable[[UnifiedTreeNode], bool]: ct = self.component_type + if match_device_only: + return lambda node: node.component_data[ct].value is not None + # HiCache: evicted + backuped (host_value present) is also a valid match return lambda node: ( node.component_data[ct].value is not None @@ -69,7 +74,10 @@ def finalize_match_result( req = params.req last_node = result.best_match_node - if len(value_chunks) > best_value_len: + # HiCache can still use prefix matches and load back host-backed Mamba + # states. We temporarily skip branching-state fill in that mode and can + # add a HiCache-aware branching policy later. + if self.cache.cache_controller is None and len(value_chunks) > best_value_len: chunk_size = get_global_server_args().mamba_cache_chunk_size aligned_seqlen = ( sum(len(v) for v in value_chunks) // chunk_size diff --git a/python/sglang/srt/mem_cache/unified_cache_components/swa_component.py b/python/sglang/srt/mem_cache/unified_cache_components/swa_component.py index 8a0d6bac22b7..63223625e906 100644 --- a/python/sglang/srt/mem_cache/unified_cache_components/swa_component.py +++ b/python/sglang/srt/mem_cache/unified_cache_components/swa_component.py @@ -69,7 +69,9 @@ def _restore_device_value(self, node: UnifiedTreeNode, value: torch.Tensor) -> N self.cache.lru_lists[ct].insert_mru(node) self.cache.component_evictable_size_[ct] += len(value) - def create_match_validator(self) -> Callable[[UnifiedTreeNode], bool]: + def create_match_validator( + self, match_device_only: bool = False + ) -> Callable[[UnifiedTreeNode], bool]: sliding_window_size = self.sliding_window_size ct = self.component_type state = {"len": float("inf")} @@ -78,7 +80,7 @@ def validator(node: UnifiedTreeNode) -> bool: cd = node.component_data[ct] # HiCache: a host-only tombstone is a valid match boundary too # — load_back will restore SWA from host before use. - if cd.value is None and cd.host_value is None: + if cd.value is None and (match_device_only or cd.host_value is None): state["len"] = 0 return False state["len"] += len(node.key) diff --git a/python/sglang/srt/mem_cache/unified_cache_components/tree_component.py b/python/sglang/srt/mem_cache/unified_cache_components/tree_component.py index 317ee3012787..ae6f71167a7b 100644 --- a/python/sglang/srt/mem_cache/unified_cache_components/tree_component.py +++ b/python/sglang/srt/mem_cache/unified_cache_components/tree_component.py @@ -116,11 +116,15 @@ def value_len(self, node: UnifiedTreeNode) -> int: return len(value) if value is not None else 0 @abstractmethod - def create_match_validator(self) -> Callable[[UnifiedTreeNode], bool]: + def create_match_validator( + self, match_device_only: bool = False + ) -> Callable[[UnifiedTreeNode], bool]: """Return a per-match stateful predicate that decides whether a node is a valid match boundary for this component. Called once per match_prefix; the returned closure may carry state. - - Full: always True (every node is valid). + When match_device_only is true, host-backed nodes must not be accepted + as valid match boundaries. + - Full: returns True if the node has full component data. - SWA: tracks accumulated length since last gap; returns True only when the contiguous window reaches swa_sliding_window_size. - Mamba: returns True iff the node has mamba component data.""" diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index f275c992532f..8c587028121a 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -351,9 +351,18 @@ def match_prefix(self, params: MatchPrefixParams) -> MatchResult: if len(key) == 0: return self._empty_match_result - value, best_match_node, best_value_len = self._match_prefix_helper(key) + ( + value, + best_match_node, + best_match_device_node, + best_match_device_value_len, + ) = self._match_prefix_helper(key) return self._match_post_processor( - params, value, best_match_node, best_value_len + params, + value, + best_match_node, + best_match_device_node, + best_match_device_value_len, ) def insert(self, params: InsertParams) -> InsertResult: @@ -585,67 +594,53 @@ def cache_unfinished_req(self, req: Req, chunked=False, **kwargs) -> None: # ---- Internal Helpers ---- - def _match_prefix_helper_readonly( - self, key: RadixKey - ) -> tuple[list[torch.Tensor], UnifiedTreeNode, int]: - """Read-only version of _match_prefix_helper that does not split nodes. - Only considers fully matched nodes, ignores partial matches. - - Not used yet; reserved for future read-only match operations.""" - node = self.root_node - child_key = key.child_key(self.page_size) - value: list[torch.Tensor] = [] - best_value_len = 0 - best_match_node = node - validators = tuple( - comp.create_match_validator() for comp in self._components_tuple - ) - - def _update_best_if_valid(node): - nonlocal best_value_len, best_match_node - if all(v(node) for v in validators): - best_value_len = len(value) - best_match_node = node - - while len(key) > 0 and child_key in node.children: - child = node.children[child_key] - - # HiCache: dead node (evicted + not backuped) — stop traversal - if child.evicted and not child.backuped: - break - - prefix_len = child.key.match(key, page_size=self.page_size) - if prefix_len < len(child.key): - # Read-only: do not split, ignore partial match and stop - break - - if not child.evicted: - value.append(child.component_data[BASE_COMPONENT_TYPE].value) - node = child - _update_best_if_valid(node) - key = key[prefix_len:] - if len(key): - child_key = key.child_key(self.page_size) - return value, best_match_node, best_value_len - def _match_prefix_helper( self, key: RadixKey - ) -> tuple[list[torch.Tensor], UnifiedTreeNode, int]: + ) -> tuple[list[torch.Tensor], UnifiedTreeNode, UnifiedTreeNode, int]: + # Non-HiCache mode has only device-resident matches, so the scheduler + # device anchor follows the best match. In HiCache mode, host-backed + # nodes can also match, so we separately track the best device-resident + # match for scheduler prefix indices and locking. node = self.root_node child_key = key.child_key(self.page_size) value: list[torch.Tensor] = [] - best_value_len = 0 best_match_node = node - validators = tuple( - comp.create_match_validator() for comp in self._components_tuple - ) + best_match_device_node = node + best_match_device_value_len = 0 + separate_device_match = self.cache_controller is not None + if separate_device_match: + validators = tuple( + comp.create_match_validator() for comp in self._components_tuple + ) + device_validators = tuple( + comp.create_match_validator(match_device_only=True) + for comp in self._components_tuple + ) + else: + validators = tuple( + comp.create_match_validator(match_device_only=True) + for comp in self._components_tuple + ) + + def _all_valid(validators, node): + return all([v(node) for v in validators]) def _update_best_if_valid(node): - nonlocal best_value_len, best_match_node - if all(v(node) for v in validators): - best_value_len = len(value) + nonlocal best_match_node + nonlocal best_match_device_value_len, best_match_device_node + matched = _all_valid(validators, node) + if matched: best_match_node = node + if not separate_device_match: + if matched: + best_match_device_value_len = len(value) + best_match_device_node = node + return + if _all_valid(device_validators, node): + best_match_device_value_len = len(value) + best_match_device_node = node + while len(key) > 0 and child_key in node.children: child = node.children[child_key] @@ -668,14 +663,21 @@ def _update_best_if_valid(node): key = key[prefix_len:] if len(key): child_key = key.child_key(self.page_size) - return value, best_match_node, best_value_len + + return ( + value, + best_match_node, + best_match_device_node, + best_match_device_value_len, + ) def _match_post_processor( self, params: MatchPrefixParams, value: list[torch.Tensor], best_match_node: UnifiedTreeNode, - best_value_len: int, + best_match_device_node: UnifiedTreeNode, + best_match_device_value_len: int, ) -> MatchResult: node_update = best_match_node for comp in self._components_tuple: @@ -691,23 +693,21 @@ def _match_post_processor( cur_time -= 0.00001 node_update = node_update.parent - # Walk up to find last_device_node - last_device_node = best_match_node - while last_device_node is not self.root_node and last_device_node.evicted: - last_device_node = last_device_node.parent - - # Walk up to find last_host_node - last_host_node = best_match_node - while last_host_node is not self.root_node and not last_host_node.backuped: - last_host_node = last_host_node.parent + # Walk up to find last_host_node for full component. + if self.cache_controller is None: + last_host_node = best_match_device_node + else: + last_host_node = best_match_node + while last_host_node is not self.root_node and not last_host_node.backuped: + last_host_node = last_host_node.parent - if best_value_len > 0: - device_indices = torch.cat(value[:best_value_len]) + if best_match_device_value_len > 0: + device_indices = torch.cat(value[:best_match_device_value_len]) else: device_indices = self._empty_match_result.device_indices result = MatchResult( device_indices=device_indices, - last_device_node=last_device_node, + last_device_node=best_match_device_node, last_host_node=last_host_node, best_match_node=best_match_node, host_hit_length=0, @@ -718,7 +718,7 @@ def _match_post_processor( result=result, params=params, value_chunks=value, - best_value_len=best_value_len, + best_value_len=best_match_device_value_len, ) return result @@ -1219,10 +1219,10 @@ def load_back( best_match_node: UnifiedTreeNode, mem_quota: Optional[int] = None, req=None, - ) -> Optional[torch.Tensor]: + ) -> bool: """Load evicted KV data from host back to device (H→D).""" if self.cache_controller is None: - return None + return False # Build KV transfer kv_xfer = self.components[BASE_COMPONENT_TYPE].build_hicache_transfers( @@ -1255,7 +1255,7 @@ def load_back( mem_quota is not None and kv_tokens > mem_quota + result.delta ): self.dec_lock_ref(best_match_node, ancestor_lock_params) - return None + return False avail = self.token_to_kv_pool_allocator.available_size() if avail < kv_tokens: @@ -1263,7 +1263,7 @@ def load_back( result = self.evict(EvictParams(num_tokens=needed)) if result.num_tokens_evicted < needed: self.dec_lock_ref(best_match_node, ancestor_lock_params) - return None + return False # Load H→D aux_xfers = [x for xfers in comp_xfers.values() for x in xfers] @@ -1276,7 +1276,7 @@ def load_back( self.dec_lock_ref(best_match_node, ancestor_lock_params) if device_indices is None: - return None + return False # Commit: each component gets only its own transfers kv_xfer.device_indices = device_indices @@ -1297,7 +1297,7 @@ def load_back( best_match_node, self.inc_lock_ref(best_match_node).to_dec_params(), ) - return device_indices + return True def _build_sidecar_transfers( self, @@ -1432,25 +1432,41 @@ def init_load_back( best_match_node = params.best_match_node mem_quota = params.mem_quota req = params.req + assert req is not None + last_best_match_device_node = req.last_node + + def _collect_new_prefix_indices() -> torch.Tensor: + prefix_chunks: list[torch.Tensor] = [] + node = best_match_node + while node is not last_best_match_device_node: + value = node.component_data[BASE_COMPONENT_TYPE].value + assert value is not None + prefix_chunks.append(value) + node = node.parent + if not prefix_chunks: + return self._empty_match_result.device_indices + prefix_chunks.reverse() + return torch.cat(prefix_chunks) if best_match_node.evicted or params.host_hit_length > 0: - loading_values = self.load_back(best_match_node, mem_quota, req=req) - if loading_values is not None: + if self.load_back(best_match_node, mem_quota, req=req): + new_indices = _collect_new_prefix_indices() + if new_indices.numel() == 0: + return ( + self._empty_match_result.device_indices, + last_best_match_device_node, + ) + logger.debug( "init_load_back success: loaded %d tokens for node %d", - len(loading_values), + len(new_indices), best_match_node.id, ) - return loading_values, best_match_node - - # Fallback: walk up to non-evicted ancestor - # TODO(ispobock): The fallback path is not correct. The last_device_node should consider all the components. - while best_match_node is not self.root_node and best_match_node.evicted: - best_match_node = best_match_node.parent + return new_indices, best_match_node return ( self._empty_match_result.device_indices, - best_match_node, + last_best_match_device_node, ) def check_hicache_events(self) -> None: diff --git a/test/registered/radix_cache/test_unified_radix_hicache_kl.py b/test/registered/radix_cache/test_unified_radix_hicache_kl.py index c6e5cdb16dbf..841af8733fe2 100644 --- a/test/registered/radix_cache/test_unified_radix_hicache_kl.py +++ b/test/registered/radix_cache/test_unified_radix_hicache_kl.py @@ -155,7 +155,7 @@ def setUpClass(cls): "--max-total-tokens", "20000", "--max-running-requests", - "4", + "2", ], env={ "SGLANG_DSV4_FP4_EXPERTS": "0", 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 a96c8a82acd6..7074084d92fd 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 @@ -15,6 +15,7 @@ DecLockRefParams, EvictParams, EvictResult, + InitLoadBackParams, InsertParams, MatchPrefixParams, MatchResult, @@ -908,6 +909,8 @@ def test_internal_readonly_does_not_modify_tree(self): """Verify readonly match does not modify tree structure (no split).""" if self.cfg.page_size > 1 or self.cfg.has_mamba or self.cfg.has_swa: self.skipTest("Full-only page_size=1 only") + if not hasattr(UnifiedRadixCache, "_match_prefix_helper_readonly"): + self.skipTest("_match_prefix_helper_readonly is not available") tree, allocator, req_to_token_pool = build_fixture(self.cfg) self._insert(tree, allocator, req_to_token_pool, [1, 2, 3, 4, 5]) @@ -922,19 +925,27 @@ def count_nodes(node): self.assertEqual(node_count_before, 2) tree._match_prefix_helper(RadixKey([1, 2])) - value, best_match_node, best_value_len = tree._match_prefix_helper( - RadixKey([1, 2, 3, 4]) - ) + ( + value, + best_match_node, + best_match_device_node, + best_value_len, + ) = tree._match_prefix_helper(RadixKey([1, 2, 3, 4])) self.assertEqual(best_value_len, 2) self.assertEqual(best_match_node.key.token_ids, [3, 4]) + self.assertIs(best_match_device_node, best_match_node) node_count_after_regular = count_nodes(tree.root_node) self.assertEqual(node_count_after_regular, node_count_before + 2) - value, best_match_node, best_value_len = tree._match_prefix_helper_readonly( - RadixKey([1, 2, 3]) - ) + ( + value, + best_match_node, + best_match_device_node, + best_value_len, + ) = tree._match_prefix_helper_readonly(RadixKey([1, 2, 3])) self.assertEqual(best_value_len, 1) self.assertEqual(best_match_node.key.token_ids, [1, 2]) + self.assertIs(best_match_device_node, best_match_node) node_count_after_readonly = count_nodes(tree.root_node) self.assertEqual(node_count_after_readonly, node_count_after_regular) @@ -1258,8 +1269,8 @@ def test_evict_d_leaf_set_consistency(self): # ================================================================ def _skip_unsupported_hicache_test(self): - if self.cfg.has_swa: - self.skipTest("HiCache tests do not run on SWA stacks") + if self.cfg.has_swa and self.cfg.has_mamba: + self.skipTest("HiCache unit fixture does not support SWA + Mamba stacks") return False def _simulate_backup(self, tree, node): @@ -1343,14 +1354,14 @@ def _backup_tree(self, tree): self._backup_node(tree, node) def _load_back_node(self, tree, node): - device_indices = tree.load_back(node) - self.assertIsNotNone(device_indices) + loaded = tree.load_back(node) + self.assertTrue(loaded) producer_id = tree.ready_to_load_host_cache() self.assertNotEqual(producer_id, -1) for _, finish_event, _ in list(tree.cache_controller.ack_load_queue): finish_event.synchronize() tree.loading_check() - return device_indices + return node.component_data[ComponentType.FULL].value def _get_full_kv_pool(self, allocator): kv_pool = allocator.get_kvcache() @@ -1474,7 +1485,9 @@ def test_hicache_match_through_evicted_node(self): def test_hicache_partial_match_splits_evicted_backed_up_node(self): """Partial matches on host-only nodes must keep the host prefix usable.""" - tree, allocator, req_to_token_pool = build_fixture(self.cfg) + if self._skip_unsupported_hicache_test(): + return + tree, allocator, req_to_token_pool = self._build_hicache_fixture() ps = self.cfg.page_size seq = self._make_seq(1, 4) expected_prefix = seq[: 2 * ps] @@ -1484,7 +1497,7 @@ def test_hicache_partial_match_splits_evicted_backed_up_node(self): self._insert(tree, allocator, req_to_token_pool, seq) m = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq))) node = m.last_device_node - self._simulate_backup(tree, node) + self._backup_node(tree, node) tree.evict(EvictParams(num_tokens=len(seq))) self.assertTrue(node.evicted) @@ -1684,6 +1697,244 @@ def _build_chain_pages(self, tree, allocator, req_to_token_pool, num_pages): chain.reverse() return chain + def _release_ongoing_load_back_locks(self, tree): + for node, lock_params in list(tree.ongoing_load_back.values()): + tree.dec_lock_ref(node, lock_params) + tree.ongoing_load_back.clear() + + def _finish_pending_loads(self, tree): + producer_id = tree.ready_to_load_host_cache() + self.assertNotEqual(producer_id, -1) + for _, finish_event, _ in list(tree.cache_controller.ack_load_queue): + finish_event.synchronize() + tree.loading_check() + + def _match_tokens_for_chain(self, chain): + tokens: list[int] = [] + for node in chain: + tokens.extend(node.key.token_ids) + return tokens + + def _set_aux_host_tombstone(self, tree, node, component_type): + cd = node.component_data[component_type] + self.assertIsNotNone(cd.value) + if cd.host_value is None: + cd.host_value = cd.value.clone() + old_value = cd.value + cd.value = None + if component_type in tree.lru_lists and tree.lru_lists[component_type].in_list( + node + ): + tree.lru_lists[component_type].remove_node(node) + tree.host_lru_lists[component_type].insert_mru(node) + tree.component_evictable_size_[component_type] -= len(old_value) + + def test_match_prefix_best_and_device_node_without_hicache(self): + tree, allocator, req_to_token_pool = build_fixture(self.cfg) + ps = self.cfg.page_size + min_tokens = 2 * ps + if self.cfg.has_swa: + min_tokens = max(min_tokens, self.cfg.sliding_window_size + ps) + seq = self._make_seq(1, (min_tokens + ps - 1) // ps) + self._insert(tree, allocator, req_to_token_pool, seq) + + result = tree.match_prefix(MatchPrefixParams(key=RadixKey(seq))) + + self.assertEqual(len(result.device_indices), len(seq)) + self.assertIs(result.best_match_node, result.last_device_node) + self.assertIs(result.last_host_node, result.last_device_node) + self.assertEqual(result.host_hit_length, 0) + + def test_hicache_mamba_host_best_match_keeps_device_anchor(self): + if not self.cfg.has_mamba or self.cfg.has_swa or self.cfg.page_size != 1: + self.skipTest("requires page_size=1 Full+Mamba") + tree, allocator, req_to_token_pool = self._build_hicache_fixture() + chain = self._build_chain_pages(tree, allocator, req_to_token_pool, 3) + if len(chain) < 3: + self.skipTest("chain too short") + leaf = chain[-1] + parent = chain[-2] + tokens = self._match_tokens_for_chain(chain) + + self._backup_node(tree, leaf) + tree.evict(EvictParams(num_tokens=len(leaf.key))) + self.assertTrue(leaf.evicted) + + result = tree.match_prefix(MatchPrefixParams(key=RadixKey(tokens))) + + self.assertIs(result.best_match_node, leaf) + self.assertIs(result.last_device_node, parent) + self.assertEqual(len(result.device_indices), len(tokens) - len(leaf.key)) + self.assertEqual(result.host_hit_length, len(leaf.key)) + + def test_hicache_swa_host_best_match_keeps_device_anchor(self): + if not self.cfg.has_swa or self.cfg.has_mamba or self.cfg.page_size != 1: + self.skipTest("requires page_size=1 Full+SWA") + tree, allocator, req_to_token_pool = self._build_hicache_fixture() + chain = self._build_chain_pages(tree, allocator, req_to_token_pool, 3) + if len(chain) < 3: + self.skipTest("chain too short") + leaf = chain[-1] + parent = chain[-2] + tokens = self._match_tokens_for_chain(chain) + + self._backup_node(tree, leaf) + tree.evict(EvictParams(num_tokens=len(leaf.key))) + self.assertTrue(leaf.evicted) + + result = tree.match_prefix(MatchPrefixParams(key=RadixKey(tokens))) + + self.assertIs(result.best_match_node, leaf) + self.assertIs(result.last_device_node, parent) + self.assertEqual(len(result.device_indices), len(tokens) - len(leaf.key)) + self.assertEqual(result.host_hit_length, 1) + + def test_mamba_branching_seqlen_disabled_under_hicache(self): + if not self.cfg.has_mamba or self.cfg.has_swa or self.cfg.page_size != 1: + self.skipTest("requires page_size=1 Full+Mamba") + tree, allocator, req_to_token_pool = build_fixture(self.cfg) + chunk_size = get_global_server_args().mamba_cache_chunk_size + tokens = self._make_seq(1, chunk_size + 1) + self._insert(tree, allocator, req_to_token_pool, tokens) + leaf = tree.match_prefix( + MatchPrefixParams(key=RadixKey(tokens)) + ).last_device_node + + mamba_cd = leaf.component_data[ComponentType.MAMBA] + mamba_cd.value = None + no_hicache = tree.match_prefix(MatchPrefixParams(key=RadixKey(tokens))) + self.assertIs(no_hicache.best_match_node, tree.root_node) + self.assertIs(no_hicache.last_device_node, tree.root_node) + self.assertEqual(no_hicache.mamba_branching_seqlen, chunk_size) + + tree_h, allocator_h, req_to_token_pool_h = self._build_hicache_fixture() + self._insert(tree_h, allocator_h, req_to_token_pool_h, tokens) + leaf_h = tree_h.match_prefix( + MatchPrefixParams(key=RadixKey(tokens)) + ).last_device_node + self._backup_node(tree_h, leaf_h) + tree_h.evict(EvictParams(num_tokens=len(tokens))) + with_hicache = tree_h.match_prefix(MatchPrefixParams(key=RadixKey(tokens))) + self.assertIs(with_hicache.best_match_node, leaf_h) + self.assertIs(with_hicache.last_device_node, tree_h.root_node) + self.assertIsNone(with_hicache.mamba_branching_seqlen) + + def test_scheduler_hicache_full_mamba_init_load_back_appends_new_indices(self): + if not self.cfg.has_mamba or self.cfg.has_swa or self.cfg.page_size != 1: + self.skipTest("requires page_size=1 Full+Mamba") + tree, allocator, req_to_token_pool = self._build_hicache_fixture() + chain = self._build_chain_pages(tree, allocator, req_to_token_pool, 3) + if len(chain) < 3: + self.skipTest("chain too short") + leaf = chain[-1] + tokens = self._match_tokens_for_chain(chain) + + self._backup_node(tree, leaf) + tree.evict(EvictParams(num_tokens=len(leaf.key))) + self.assertTrue(leaf.evicted) + + req = self._make_req(req_to_token_pool) + match = tree.match_prefix(MatchPrefixParams(key=RadixKey(tokens), req=req)) + req.prefix_indices = match.device_indices + req.last_node = match.last_device_node + req.best_match_node = match.best_match_node + req.host_hit_length = match.host_hit_length + + new_indices, new_node = tree.init_load_back( + InitLoadBackParams( + best_match_node=req.best_match_node, + host_hit_length=req.host_hit_length, + req=req, + ) + ) + + self.assertIs(new_node, leaf) + self.assertEqual(len(torch.cat([req.prefix_indices, new_indices])), len(tokens)) + self.assertIsNotNone(leaf.component_data[ComponentType.MAMBA].value) + self._finish_pending_loads(tree) + self._release_ongoing_load_back_locks(tree) + + def test_scheduler_hicache_aux_only_load_back_appends_full_device_indices(self): + if self.cfg.page_size != 1: + self.skipTest("page_size=1 keeps the expected suffix precise") + aux = None + if self.cfg.has_swa and not self.cfg.has_mamba: + aux = ComponentType.SWA + elif self.cfg.has_mamba and not self.cfg.has_swa: + aux = ComponentType.MAMBA + if aux is None: + self.skipTest("requires exactly one aux component") + + tree, allocator, req_to_token_pool = self._build_hicache_fixture() + chain = self._build_chain_pages(tree, allocator, req_to_token_pool, 3) + if len(chain) < 3: + self.skipTest("chain too short") + leaf = chain[-1] + tokens = self._match_tokens_for_chain(chain) + leaf_full = leaf.component_data[ComponentType.FULL].value.clone() + self._backup_node(tree, leaf) + self._set_aux_host_tombstone(tree, leaf, aux) + + req = self._make_req(req_to_token_pool) + match = tree.match_prefix(MatchPrefixParams(key=RadixKey(tokens), req=req)) + req.prefix_indices = match.device_indices + req.last_node = match.last_device_node + req.best_match_node = match.best_match_node + req.host_hit_length = match.host_hit_length + + new_indices, new_node = tree.init_load_back( + InitLoadBackParams( + best_match_node=req.best_match_node, + host_hit_length=req.host_hit_length, + req=req, + ) + ) + + self.assertIs(new_node, leaf) + self.assertEqual(new_indices.tolist(), leaf_full.tolist()) + self.assertEqual(len(torch.cat([req.prefix_indices, new_indices])), len(tokens)) + self.assertEqual( + leaf.component_data[ComponentType.FULL].value.tolist(), + leaf_full.tolist(), + ) + self.assertIsNotNone(leaf.component_data[aux].value) + self._finish_pending_loads(tree) + self._release_ongoing_load_back_locks(tree) + + def test_scheduler_hicache_load_back_fallback_keeps_old_anchor(self): + if not self.cfg.has_mamba or self.cfg.has_swa or self.cfg.page_size != 1: + self.skipTest("requires page_size=1 Full+Mamba") + tree, allocator, req_to_token_pool = self._build_hicache_fixture() + chain = self._build_chain_pages(tree, allocator, req_to_token_pool, 3) + if len(chain) < 3: + self.skipTest("chain too short") + leaf = chain[-1] + tokens = self._match_tokens_for_chain(chain) + + self._backup_node(tree, leaf) + tree.evict(EvictParams(num_tokens=len(leaf.key))) + + req = self._make_req(req_to_token_pool) + match = tree.match_prefix(MatchPrefixParams(key=RadixKey(tokens), req=req)) + req.prefix_indices = match.device_indices + req.last_node = match.last_device_node + req.best_match_node = match.best_match_node + req.host_hit_length = match.host_hit_length + + new_indices, new_node = tree.init_load_back( + InitLoadBackParams( + best_match_node=req.best_match_node, + host_hit_length=req.host_hit_length, + req=req, + mem_quota=-1_000_000, + ) + ) + + self.assertEqual(len(new_indices), 0) + self.assertIs(new_node, match.last_device_node) + self.assertIsNone(leaf.component_data[ComponentType.FULL].value) + self.assertIsNone(leaf.component_data[ComponentType.MAMBA].value) + def test_hicache_swa_load_back_min_suffix(self): """LOAD_BACK collects only the suffix nodes needed to cover sliding_window_size.""" if not self.cfg.has_swa: @@ -1940,11 +2191,11 @@ def _swa_anchor_setup(self): if chain_pages * ps > self.cfg.kv_size // 2: self.skipTest("kv_size too small for the desired chain") - tree, allocator, req_to_token_pool = build_fixture(self.cfg) + tree, allocator, req_to_token_pool = self._build_hicache_fixture() chain = self._build_chain_pages(tree, allocator, req_to_token_pool, chain_pages) if len(chain) < chain_pages: self.skipTest("chain too short") - self._simulate_backup_tree(tree) + self._backup_tree(tree) x = chain[-1] y = chain[-window_pages] @@ -1962,10 +2213,10 @@ def _swa_anchor_setup(self): return tree, chain, n, y, x, tokens def test_hicache_swa_match_prefix_picks_best_match_node_above_last_host(self): - tree, _, _, y, x, tokens = self._swa_anchor_setup() + tree, _, n, y, x, tokens = self._swa_anchor_setup() result = tree.match_prefix(MatchPrefixParams(key=RadixKey(tokens))) self.assertIs(result.best_match_node, x) - self.assertIs(result.last_device_node, x) + self.assertIs(result.last_device_node, n.parent) self.assertIs(result.last_host_node, y) def test_hicache_swa_load_back_anchored_on_best_match_node(self): From 6bacd0c5123bd1ae3df0097e0ec9fdbec022f2f5 Mon Sep 17 00:00:00 2001 From: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com> Date: Fri, 15 May 2026 22:32:50 +0100 Subject: [PATCH 02/50] [Fix] DeepSeek-V3.2: build structural tag locally to encode both wrapper and invoke layers (#25233) --- .../srt/function_call/deepseekv32_detector.py | 118 +++++++++++++++++- 1 file changed, 115 insertions(+), 3 deletions(-) diff --git a/python/sglang/srt/function_call/deepseekv32_detector.py b/python/sglang/srt/function_call/deepseekv32_detector.py index 4febe5458a40..5d4391e913dc 100644 --- a/python/sglang/srt/function_call/deepseekv32_detector.py +++ b/python/sglang/srt/function_call/deepseekv32_detector.py @@ -1,10 +1,11 @@ import json import logging import re +from typing import List, Literal, Optional, Union from partial_json_parser.core.options import Allow -from sglang.srt.entrypoints.openai.protocol import Tool +from sglang.srt.entrypoints.openai.protocol import Tool, ToolChoice from sglang.srt.function_call.base_format_detector import BaseFormatDetector from sglang.srt.function_call.core_types import ( StreamingParseResult, @@ -14,8 +15,30 @@ ) from sglang.srt.function_call.utils import _find_common_prefix, _partial_json_loads +try: + from xgrammar import StructuralTag + from xgrammar.structural_tag import ( + AnyTextFormat, + ConstStringFormat, + JSONSchemaFormat, + SequenceFormat, + TagFormat, + TagsWithSeparatorFormat, + TriggeredTagsFormat, + ) +except ImportError: + StructuralTag = None # type: ignore + logger = logging.getLogger(__name__) +# Names mirror the DeepSeek-V3.2 official chat template tokens +# (see encoding_dsv32.TOOLS_SYSTEM_TEMPLATE). +_INVOKE_BEGIN_PREFIX = '<|DSML|invoke name="' +_INVOKE_BEGIN_SUFFIX = '">\n' +_THINK_TAG_END = "" +_THINK_EXCLUDE_TOKENS = ["", ""] +_XML_STYLE = "deepseek_xml" + class DeepSeekV32Detector(BaseFormatDetector): """ @@ -368,5 +391,94 @@ def structure_info(self) -> _GetInfoFunc: trigger="<|DSML|invoke", ) - def get_structural_tag_name(self) -> str: - return "deepseek_v3_2" + def get_structural_tag( + self, + tools: Union[List[Tool], None] = None, + tool_choice: Union[ToolChoice, Literal["auto", "required"]] = "auto", + thinking_mode: bool = False, + ) -> Optional["StructuralTag"]: + """ + Build an xgrammar StructuralTag locally for DeepSeek-V3.2. + + Both layers — the outer `<|DSML|function_calls|>...` wrapper and + the inner `<|DSML|invoke>...` blocks — are encoded + directly in the grammar with a single-newline join between + consecutive invokes, matching DeepSeek-V3.2's official chat + template. This avoids two layered defects that surfaced with the + prior `xgrammar.get_model_structural_tag("deepseek_v3_2")` path: + + - the xgrammar builtin template (pre mlc-ai/xgrammar#638) forced + a double-newline join, which deterministically collapsed + parallel tool calls to one at greedy decoding. + - falling back to the legacy structural tag (built from + `structure_info()`) only constrains the inner invoke block; + the outer wrapper is off-grammar and the model can skip it + under `at_least_one=True`, leaving `detect_and_parse` with no + `<|DSML|function_calls>` marker to anchor on. + + Returning a fully-formed StructuralTag from the detector keeps + both fixes local to sglang and decoupled from the xgrammar + release cadence. + """ + if not tools or StructuralTag is None: + return None + + # `INVOKE_END` and the empty separator together yield a single `\n` + # between consecutive invokes — matching DeepSeek-V3.2's chat template + # `"\n".join(invoke_blocks)`. + function_calls_begin = self.bot_token + "\n" + invoke_end = self.invoke_end_token + "\n" + + def _invoke_tag(tool: Tool) -> TagFormat: + return TagFormat( + begin=_INVOKE_BEGIN_PREFIX + tool.function.name + _INVOKE_BEGIN_SUFFIX, + content=JSONSchemaFormat( + json_schema=tool.function.parameters or {}, + style=_XML_STYLE, + ), + end=invoke_end, + ) + + if isinstance(tool_choice, ToolChoice): + target = next( + (t for t in tools if t.function.name == tool_choice.function.name), + None, + ) + if target is None: + return None + invoke_tags = [_invoke_tag(target)] + is_required = True + else: + invoke_tags = [_invoke_tag(t) for t in tools] + is_required = tool_choice == "required" + + inner_tool_calls = TagsWithSeparatorFormat( + tags=invoke_tags, separator="", at_least_one=True + ) + + if is_required: + suffix_tag = SequenceFormat( + elements=[ + ConstStringFormat(value=function_calls_begin), + inner_tool_calls, + ConstStringFormat(value=self.eot_token), + ] + ) + else: + suffix_tag = TriggeredTagsFormat( + triggers=[self.bot_token], + tags=[ + TagFormat( + begin=function_calls_begin, + content=inner_tool_calls, + end=self.eot_token, + ) + ], + excludes=_THINK_EXCLUDE_TOKENS, + ) + + if not thinking_mode: + return StructuralTag(format=suffix_tag) + + prefix_tag = TagFormat(begin="", content=AnyTextFormat(), end=_THINK_TAG_END) + return StructuralTag(format=SequenceFormat(elements=[prefix_tag, suffix_tag])) From e028556db0e62ced06be42c15c0afa490c01f002 Mon Sep 17 00:00:00 2001 From: Yuhao Yang <47235274+yhyang201@users.noreply.github.com> Date: Sat, 16 May 2026 08:26:26 +0800 Subject: [PATCH 03/50] Add multi-detokenizer support (#24944) Co-authored-by: ybyang <10629930+whybeyoung@users.noreply.github.com> Co-authored-by: Shangming Cai --- python/sglang/srt/entrypoints/engine.py | 84 ++++++++-- .../srt/managers/detokenizer_manager.py | 14 +- .../srt/managers/multi_tokenizer_mixin.py | 148 ++++++++++++++++-- python/sglang/srt/server_args.py | 14 ++ .../tokenizer/test_multi_detokenizer.py | 80 ++++++++++ 5 files changed, 308 insertions(+), 32 deletions(-) create mode 100644 test/registered/tokenizer/test_multi_detokenizer.py diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index f96445c31055..1dde8bed80e9 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -27,6 +27,7 @@ import os import random import signal +import tempfile import threading import time from typing import ( @@ -80,7 +81,10 @@ UpdateWeightsFromIPCReqInput, UpdateWeightsFromTensorReqInput, ) -from sglang.srt.managers.multi_tokenizer_mixin import MultiTokenizerRouter +from sglang.srt.managers.multi_tokenizer_mixin import ( + MultiTokenizerRouter, + run_multi_detokenizer_router_process, +) from sglang.srt.managers.scheduler import run_scheduler_process from sglang.srt.managers.template_detection import resolve_auto_parsers from sglang.srt.managers.template_manager import TemplateManager @@ -675,6 +679,63 @@ def wait_for_completion(): scheduler_procs, ) + @classmethod + def _launch_detokenizer_subprocesses( + cls, + server_args: ServerArgs, + port_args: PortArgs, + run_detokenizer_process_func: Callable, + ) -> Tuple[List[mp.Process], List[str]]: + """Launch detokenizer worker(s). + + - When ``detokenizer_worker_num == 1``: a single detokenizer process listens on + ``port_args.detokenizer_ipc_name`` (the original behavior). + - When ``detokenizer_worker_num > 1``: each detokenizer worker gets its own + private IPC socket, and a ``MultiDetokenizerRouter`` process owns the + original ``port_args.detokenizer_ipc_name`` and fans out to them. + + Returns (processes, names) for SubprocessWatchdog. + """ + processes: List[mp.Process] = [] + names: List[str] = [] + + if server_args.detokenizer_worker_num <= 1: + proc = mp.Process( + target=run_detokenizer_process_func, + args=(server_args, port_args), + ) + proc.start() + processes.append(proc) + names.append("detokenizer") + return processes, names + + router_ipc_name = port_args.detokenizer_ipc_name + worker_ipc_names: List[str] = [] + try: + for i in range(server_args.detokenizer_worker_num): + worker_ipc = f"ipc://{tempfile.NamedTemporaryFile(delete=False).name}" + port_args.detokenizer_ipc_name = worker_ipc + proc = mp.Process( + target=run_detokenizer_process_func, + args=(server_args, port_args), + ) + proc.start() + processes.append(proc) + names.append(f"detokenizer_{i}") + worker_ipc_names.append(worker_ipc) + finally: + port_args.detokenizer_ipc_name = router_ipc_name + + router_proc = mp.Process( + target=run_multi_detokenizer_router_process, + args=(worker_ipc_names, server_args, port_args), + ) + router_proc.start() + processes.append(router_proc) + names.append("detokenizer_router") + + return processes, names + @classmethod def _launch_subprocesses( cls, @@ -776,16 +837,15 @@ def _launch_subprocesses( None, ) - # Launch detokenizer process - detoken_proc = mp.Process( - target=run_detokenizer_process_func, - args=( - server_args, - port_args, - ), + # Launch detokenizer process(es) — optionally fronted by a router when + # detokenizer_worker_num > 1. + detoken_procs, detoken_names = cls._launch_detokenizer_subprocesses( + server_args=server_args, + port_args=port_args, + run_detokenizer_process_func=run_detokenizer_process_func, ) - detoken_proc.start() - scheduler_init_result.all_child_pids.append(detoken_proc.pid) + for p in detoken_procs: + scheduler_init_result.all_child_pids.append(p.pid) # Init tokenizer manager first, as the bootstrap server is initialized here if server_args.tokenizer_worker_num == 1: @@ -809,8 +869,8 @@ def _launch_subprocesses( # Note: RayEngine returns scheduler_procs=None as it uses Ray actors instead of mp.Process processes = list(scheduler_procs or []) names = [f"scheduler_{i}" for i in range(len(processes))] - processes.append(detoken_proc) - names.append("detokenizer") + processes.extend(detoken_procs) + names.extend(detoken_names) subprocess_watchdog = SubprocessWatchdog( processes=processes, process_names=names ) diff --git a/python/sglang/srt/managers/detokenizer_manager.py b/python/sglang/srt/managers/detokenizer_manager.py index a4547bf36c24..a35e98167b90 100644 --- a/python/sglang/srt/managers/detokenizer_manager.py +++ b/python/sglang/srt/managers/detokenizer_manager.py @@ -81,7 +81,7 @@ def __init__( port_args: PortArgs, ): # Init inter-process communication - self.init_ipc_channels(port_args) + self.init_ipc_channels(port_args, server_args) # Init tokenizer self.init_tokenizer(server_args) @@ -92,14 +92,18 @@ def __init__( # Init dispatcher self.init_request_dispatcher() - def init_ipc_channels(self, port_args: PortArgs): + def init_ipc_channels(self, port_args: PortArgs, server_args: ServerArgs): context = zmq.Context(2) self.recv_from_scheduler = get_zmq_socket( context, zmq.PULL, port_args.detokenizer_ipc_name, True ) - self.send_to_tokenizer = get_zmq_socket( - context, zmq.PUSH, port_args.tokenizer_ipc_name, False - ) + # In multi-tokenizer mode, results are pushed back to each TokenizerWorker + # directly via SocketMapping inside multi_http_worker_event_loop, so the + # single send_to_tokenizer socket is unused. + if server_args.tokenizer_worker_num == 1: + self.send_to_tokenizer = get_zmq_socket( + context, zmq.PUSH, port_args.tokenizer_ipc_name, False + ) def init_tokenizer(self, server_args: ServerArgs): if server_args.skip_tokenizer_init: diff --git a/python/sglang/srt/managers/multi_tokenizer_mixin.py b/python/sglang/srt/managers/multi_tokenizer_mixin.py index 658acc01c9aa..baf25d332e2e 100644 --- a/python/sglang/srt/managers/multi_tokenizer_mixin.py +++ b/python/sglang/srt/managers/multi_tokenizer_mixin.py @@ -24,11 +24,14 @@ import multiprocessing as multiprocessing import os import pickle +import signal import sys import threading +import zlib from multiprocessing import shared_memory -from typing import TYPE_CHECKING, Any, Dict, Optional, Union +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +import psutil import setproctitle import zmq import zmq.asyncio @@ -43,13 +46,18 @@ BatchStrOutput, BatchTokenIDOutput, ContinueGenerationReqInput, + FreezeGCReq, PauseContinueBroadcast, PauseGenerationReqInput, TokenizerWorkerRegistration, ) from sglang.srt.managers.tokenizer_manager import TokenizerManager from sglang.srt.server_args import PortArgs, ServerArgs -from sglang.srt.utils import kill_process_tree +from sglang.srt.utils import ( + configure_logger, + kill_itself_when_parent_died, + kill_process_tree, +) from sglang.srt.utils.network import get_zmq_socket from sglang.utils import get_exception_traceback @@ -78,14 +86,14 @@ def _register_ipc_mapping(self, ipc_name: str, is_tokenizer: bool): socket = get_zmq_socket(self._zmq_context, zmq.PUSH, ipc_name, False) self._mapping[ipc_name] = socket - def send_output(self, ipc_name: str, output: Any): + def send_output(self, ipc_name: str, output: Any, is_tokenizer: bool = False): if ipc_name is None: # Some unhandled cases logger.warning(f"IPC name is None, output type={type(output)}, skipping...") return if ipc_name not in self._mapping: - self._register_ipc_mapping(ipc_name, is_tokenizer=False) + self._register_ipc_mapping(ipc_name, is_tokenizer=is_tokenizer) self._mapping[ipc_name].send_pyobj(output) @@ -110,9 +118,7 @@ def _extract_field_by_index( if isinstance(field, dict): new_field = {} for k, v in field.items(): - if len(v) <= index: - new_field[k] = None - new_field[k] = v[index] + new_field[k] = v[index] if len(v) > index else None return new_field if check_length: @@ -196,11 +202,22 @@ def _handle_output_by_index(output, i): output_hidden_states=_extract_field_by_index( output, "output_hidden_states", i, check_length=False ), + routed_experts=_extract_field_by_index( + output, "routed_experts", i, check_length=False + ), + indexer_topk=_extract_field_by_index( + output, "indexer_topk", i, check_length=False + ), + retraction_counts=_extract_field_by_index(output, "retraction_counts", i), placeholder_tokens_idx=None, placeholder_tokens_val=None, token_steps=_extract_field_by_index( output, "token_steps", i, check_length=False ), + customized_info=_extract_field_by_index( + output, "customized_info", i, check_length=False + ), + dp_ranks=_extract_field_by_index(output, "dp_ranks", i, check_length=False), ) elif isinstance(output, BatchEmbeddingOutput): new_output = BatchEmbeddingOutput( @@ -310,14 +327,23 @@ def multi_http_worker_event_loop(self: DetokenizerManager): if output is None: continue - assert isinstance( - recv_obj, BaseBatchReq - ), "for multi-http-worker, recv_obj must be BaseBatchReq" - - # Send data using the corresponding socket - for i, ipc_name in enumerate(recv_obj.http_worker_ipcs): - new_output = _handle_output_by_index(output, i) - self.socket_mapping.send_output(ipc_name, new_output) + # Fan out the output back to the originating tokenizer worker(s). + # In multi-detokenizer mode the upstream MultiDetokenizerRouter may + # forward either batched or single requests, so handle both shapes. + if isinstance(recv_obj, BaseBatchReq): + for i, ipc_name in enumerate(recv_obj.http_worker_ipcs): + new_output = _handle_output_by_index(output, i) + self.socket_mapping.send_output( + ipc_name, new_output, is_tokenizer=True + ) + elif isinstance(recv_obj, BaseReq): + self.socket_mapping.send_output( + recv_obj.http_worker_ipc, output, is_tokenizer=True + ) + else: + raise ValueError( + f"multi_http_worker_event_loop got unexpected req type {type(recv_obj)}" + ) class MultiTokenizerRouter: @@ -415,6 +441,98 @@ async def _distribute_result_to_workers(self, recv_obj): self.socket_mapping.send_output(ipc_name, new_recv_obj) +class MultiDetokenizerRouter: + """Route scheduler outputs to one of N DetokenizerManager workers. + + Each request is pinned to a worker by hashing its ``http_worker_ipc`` with + ``zlib.crc32`` (deterministic across runs), so all outputs of the same rid + always land on the same detokenizer and ``decode_status`` stays consistent. + """ + + def __init__(self, ipc_name_list: List[str], port_args: PortArgs): + self.ipc_name_list = ipc_name_list + self.num_workers = len(ipc_name_list) + self.socket_mapping = SocketMapping() + context = zmq.Context(2) + self.recv_from_scheduler = get_zmq_socket( + context, zmq.PULL, port_args.detokenizer_ipc_name, True + ) + + def _pick(self, key: str) -> str: + return self.ipc_name_list[zlib.crc32(key.encode()) % self.num_workers] + + def _send(self, ipc_name: str, obj: Any) -> None: + self.socket_mapping.send_output(ipc_name, obj, is_tokenizer=False) + + def event_loop(self): + while True: + recv_obj = self.recv_from_scheduler.recv_pyobj() + + # FreezeGCReq must freeze every detokenizer process. + if isinstance(recv_obj, FreezeGCReq): + for ipc in self.ipc_name_list: + self._send(ipc, recv_obj) + continue + + # Single request: route by its own http_worker_ipc. + if isinstance(recv_obj, BaseReq): + assert ( + recv_obj.http_worker_ipc is not None + ), f"Single req {recv_obj.rid=} missing http_worker_ipc" + self._send(self._pick(recv_obj.http_worker_ipc), recv_obj) + continue + + # Batch request. + if isinstance(recv_obj, BaseBatchReq): + # Idle/no-op batch (rids=[]): broadcast to all detokenizers + if not recv_obj.rids: + for ipc in self.ipc_name_list: + self._send(ipc, recv_obj) + continue + + ipcs = recv_obj.http_worker_ipcs + assert ( + ipcs is not None + and len(ipcs) == len(recv_obj.rids) + and all(x is not None for x in ipcs) + ), f"Batch req {recv_obj.rids=} has invalid http_worker_ipcs" + + # Split per-item and route each by its own ipc. + for i, ipc_key in enumerate(ipcs): + one = _handle_output_by_index(recv_obj, i) + if one is recv_obj: + raise TypeError(f"Cannot split {type(recv_obj)}") + one.http_worker_ipcs = [ipc_key] + self._send(self._pick(ipc_key), one) + continue + + raise ValueError( + f"MultiDetokenizerRouter got unsupported type {type(recv_obj)}" + ) + + +def run_multi_detokenizer_router_process( + ipc_name_list: List[str], + server_args: ServerArgs, + port_args: PortArgs, +): + kill_itself_when_parent_died() + setproctitle.setproctitle("sglang::detokenizer_router") + configure_logger(server_args) + parent_process = psutil.Process().parent() + + router = None + try: + router = MultiDetokenizerRouter(ipc_name_list, port_args) + router.event_loop() + except Exception: + traceback = get_exception_traceback() + logger.error(f"MultiDetokenizerRouter hit an exception: {traceback}") + if router is not None: + router.socket_mapping.clear_all_sockets() + parent_process.send_signal(signal.SIGQUIT) + + class TokenizerWorker(TokenizerManager): """Tokenizer Worker in multi-http-worker mode""" diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index a1dba0f7290c..5e48ad36d1be 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -369,6 +369,7 @@ class ServerArgs: tokenizer_mode: str = "auto" tokenizer_backend: str = "huggingface" tokenizer_worker_num: int = 1 + detokenizer_worker_num: int = 1 skip_tokenizer_init: bool = False load_format: str = "auto" model_loader_extra_config: str = "{}" @@ -4186,6 +4187,12 @@ def _handle_tokenizer_batching(self): f"(requested {self.tokenizer_worker_num})." ) self.tokenizer_worker_num = 1 + if self.detokenizer_worker_num != 1: + logger.warning( + "skip_tokenizer_init=True disables detokenizer workers; forcing detokenizer_worker_num=1 " + f"(requested {self.detokenizer_worker_num})." + ) + self.detokenizer_worker_num = 1 if self.enable_tokenizer_batch_encode: logger.warning( @@ -4529,6 +4536,12 @@ def add_cli_args(parser: argparse.ArgumentParser): default=ServerArgs.tokenizer_worker_num, help="The worker num of the tokenizer manager.", ) + parser.add_argument( + "--detokenizer-worker-num", + type=int, + default=ServerArgs.detokenizer_worker_num, + help="The worker num of the detokenizer manager.", + ) parser.add_argument( "--skip-tokenizer-init", action="store_true", @@ -7275,6 +7288,7 @@ def check_server_args(self): ) assert self.tokenizer_worker_num > 0, "Tokenizer worker num must >= 1" + assert self.detokenizer_worker_num > 0, "Detokenizer worker num must >= 1" self.validate_buckets_rule( "--prompt-tokens-buckets", self.prompt_tokens_buckets ) diff --git a/test/registered/tokenizer/test_multi_detokenizer.py b/test/registered/tokenizer/test_multi_detokenizer.py new file mode 100644 index 000000000000..09ca5e9b101e --- /dev/null +++ b/test/registered/tokenizer/test_multi_detokenizer.py @@ -0,0 +1,80 @@ +import unittest + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci +from sglang.test.kits.eval_accuracy_kit import MMLUMixin +from sglang.test.test_utils import ( + DEFAULT_MODEL_NAME_FOR_TEST, + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + CustomTestCase, + auto_config_device, + get_benchmark_args, + is_in_amd_ci, + is_in_ci, + popen_launch_server, + run_benchmark, + write_github_step_summary, +) + +register_cuda_ci(est_time=211, suite="stage-b-test-1-gpu-large") +register_amd_ci(est_time=345, suite="stage-b-test-1-gpu-small-amd") + + +class TestMultiDetokenizer(CustomTestCase, MMLUMixin): + mmlu_score_threshold = 0.65 + mmlu_num_examples = 64 + mmlu_num_threads = 32 + + @classmethod + def setUpClass(cls): + cls.model = DEFAULT_MODEL_NAME_FOR_TEST + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=[ + "--tokenizer-worker-num", + 8, + "--detokenizer-worker-num", + 4, + "--mem-fraction-static", + 0.7, + ], + ) + + @classmethod + def tearDownClass(cls): + kill_process_tree(cls.process.pid) + + def test_multi_detokenizer_ttft(self): + args = get_benchmark_args( + base_url=self.base_url, + dataset_name="random", + dataset_path="", + tokenizer=None, + num_prompts=100, + random_input_len=4096, + random_output_len=2048, + sharegpt_context_len=None, + request_rate=1, + disable_stream=False, + disable_ignore_eos=False, + seed=0, + device=auto_config_device(), + lora_name=None, + ) + res = run_benchmark(args) + if is_in_ci(): + write_github_step_summary( + f"### test_multi_detokenizer_ttft\n" + f"median_e2e_latency_ms: {res['median_e2e_latency_ms']:.2f} ms\n" + ) + self.assertLess(res["median_e2e_latency_ms"], 11000) + self.assertLess(res["median_ttft_ms"], 130 if is_in_amd_ci() else 86) + self.assertLess(res["median_itl_ms"], 10) + + +if __name__ == "__main__": + unittest.main() From a004d0a49b4dedbcdabc854b2b6ce51980ef1bea Mon Sep 17 00:00:00 2001 From: Brayden Zhong Date: Fri, 15 May 2026 20:28:35 -0400 Subject: [PATCH 04/50] Fix Mistral Large 3 nightly test (#25407) Co-authored-by: b8zhong --- .../schemes/compressed_tensors_w4a4_nvfp4_moe.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a4_nvfp4_moe.py b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a4_nvfp4_moe.py index 69e572498c66..0a992187478e 100644 --- a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a4_nvfp4_moe.py +++ b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a4_nvfp4_moe.py @@ -311,10 +311,10 @@ def apply_weights( router_logits = topk_output.router_logits topk_config = topk_output.topk_config - # Quantize input hidden states using fp4_quantize + # global_scale must be shape [1] (strict in cute-dsl backend). hs_fp4_bytes, hs_sf_bytes = fp4_quantize( x, - layer.w13_input_scale_quant, + layer.w13_input_scale_quant[:1], self.group_size, # sf_vec_size False, # use_ue8m0 False, # is_sf_swizzled_layout From 2e5b65b46a3edde965704872734ff7d33639b49b Mon Sep 17 00:00:00 2001 From: Jimmy Shong <69131491+Jiminator@users.noreply.github.com> Date: Fri, 15 May 2026 20:14:47 -0700 Subject: [PATCH 05/50] [CI] Lower mem-fraction-static for GLM-5.1 FP8 8-GPU test to 0.85 (#25453) --- test/registered/8-gpu-models/test_glm_51_fp8.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/test/registered/8-gpu-models/test_glm_51_fp8.py b/test/registered/8-gpu-models/test_glm_51_fp8.py index ace31a06f070..f96cc16b7da2 100644 --- a/test/registered/8-gpu-models/test_glm_51_fp8.py +++ b/test/registered/8-gpu-models/test_glm_51_fp8.py @@ -15,7 +15,7 @@ "--trust-remote-code", "--reasoning-parser=glm45", "--tool-call-parser=glm47", - "--mem-fraction-static=0.9", + "--mem-fraction-static=0.85", "--enable-metrics", ] From 6977e95d4b00718f640856e5e76462cc99e47635 Mon Sep 17 00:00:00 2001 From: Jimmy Shong <69131491+Jiminator@users.noreply.github.com> Date: Fri, 15 May 2026 22:01:44 -0700 Subject: [PATCH 06/50] [Fix] Probe speculative draft config via sglang get_config (#25428) --- python/sglang/srt/server_args.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 5e48ad36d1be..8ebdf088c5e7 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -324,9 +324,9 @@ def _resolve_speculative_algorithm_alias( is_gemma4_draft = False if speculative_draft_model_path: - from transformers import AutoConfig + from sglang.srt.utils.hf_transformers_utils import get_config - cfg = AutoConfig.from_pretrained( + cfg = get_config( speculative_draft_model_path, trust_remote_code=trust_remote_code ) is_gemma4_draft = "Gemma4AssistantForCausalLM" in ( From 251e9c9636bde93bb3b48df23872d446b5aa9327 Mon Sep 17 00:00:00 2001 From: JoeLee314 Date: Sat, 16 May 2026 13:46:44 +0800 Subject: [PATCH 07/50] [Disagg] Fix MegaMoE topk_ids dtype mismatch and FakeKVManager missing kv_args (#25380) Co-authored-by: JoeLee314 --- python/sglang/srt/disaggregation/fake/conn.py | 1 + python/sglang/srt/layers/moe/mega_moe.py | 4 ++-- 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/disaggregation/fake/conn.py b/python/sglang/srt/disaggregation/fake/conn.py index d59641c3c428..e44cbb7b3e48 100644 --- a/python/sglang/srt/disaggregation/fake/conn.py +++ b/python/sglang/srt/disaggregation/fake/conn.py @@ -28,6 +28,7 @@ def __init__( is_mla_backend: Optional[bool] = False, ): super().__init__(args, disaggregation_mode, server_args, is_mla_backend) + self.kv_args = args self.req_to_decode_prefix_len = {} def register_to_bootstrap(self): diff --git a/python/sglang/srt/layers/moe/mega_moe.py b/python/sglang/srt/layers/moe/mega_moe.py index 93f4d9a155a6..1c13f7be9885 100644 --- a/python/sglang/srt/layers/moe/mega_moe.py +++ b/python/sglang/srt/layers/moe/mega_moe.py @@ -206,8 +206,8 @@ def _run_mega_routed( ) if num_tokens > 0: - topk_ids_in = topk_ids - topk_weights_in = topk_weights + topk_ids_in = topk_ids.to(torch.int32) + topk_weights_in = topk_weights.to(torch.float32) else: topk_ids_in = hidden_states.new_empty((0, top_k), dtype=torch.int32) topk_weights_in = hidden_states.new_empty((0, top_k), dtype=torch.float32) From 9d2f0e438d10a6a00c67835a8eb9647a33ab64cc Mon Sep 17 00:00:00 2001 From: ybyang <10629930+whybeyoung@users.noreply.github.com> Date: Sat, 16 May 2026 13:54:32 +0800 Subject: [PATCH 08/50] feat: add Pipeline Parallelism (PP) and PD support for DeepSeek-V4 (#24704) Co-authored-by: Shangming Cai Co-authored-by: xuyongfei --- python/sglang/srt/configs/model_config.py | 3 + python/sglang/srt/disaggregation/base/conn.py | 14 +- .../sglang/srt/disaggregation/common/conn.py | 132 ++++++++++++++++- python/sglang/srt/disaggregation/prefill.py | 10 ++ python/sglang/srt/managers/schedule_batch.py | 5 +- .../sglang/srt/managers/scheduler_pp_mixin.py | 6 +- .../srt/mem_cache/deepseek_v4_memory_pool.py | 122 +++++++++------- .../srt/model_executor/cuda_graph_runner.py | 14 +- python/sglang/srt/models/deepseek_v4.py | 138 +++++++++++++----- 9 files changed, 341 insertions(+), 103 deletions(-) diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index 5a9a931017db..111145ef6d2f 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -766,6 +766,9 @@ def _derive_model_shapes(self): self.spec_hidden_size = ( self.hidden_size * hc_mult if hc_mult > 1 else self.hidden_size ) + # mHC-flattened hidden size; None when not running an mHC model + # (e.g. non-DeepSeek-V4 configs without ``hc_mult``). + self.hc_hidden_size = self.spec_hidden_size if hc_mult > 1 else None self.num_hidden_layers = self.hf_text_config.num_hidden_layers self.num_attention_layers = self.num_hidden_layers if "LongcatFlashForCausalLM" in self.hf_config.architectures: diff --git a/python/sglang/srt/disaggregation/base/conn.py b/python/sglang/srt/disaggregation/base/conn.py index 264f3ca0d054..1e1b4b4f50d6 100644 --- a/python/sglang/srt/disaggregation/base/conn.py +++ b/python/sglang/srt/disaggregation/base/conn.py @@ -47,11 +47,21 @@ class KVArgs: kv_head_num: int total_kv_head_num: int page_size: int + # for system dp + system_dp_rank: int # for pp prefill pp_rank: int prefill_start_layer: int - # for system dp - system_dp_rank: int + # Absolute end layer (exclusive) for this prefill PP stage. Needed to + # reconstruct PP sub-ranges when kv_data_ptrs does not use a flat + # layer-indexed layout (e.g. DeepSeek V4's buffer-type-organized flat + # list). + prefill_end_layer: int + # For DeepSeek V4 (and other compressed-MLA) memory pools only. + # Full-model compression ratio per layer (entries are 0/4/128). Used by + # the connection layer to slice the buffer-type-organized flat list in a + # PP-aware manner. + mla_compression_ratios: Optional[List[int]] # Only used of npu, for kv buf groups kv_buf_groups: int # Only used of npu, for decode total kv layers diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index 596b58303bc6..555ef5215449 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -474,15 +474,133 @@ def get_mha_kv_ptrs_with_pp( def get_mla_kv_ptrs_with_pp( self, src_kv_ptrs: List[int], dst_kv_ptrs: List[int] ) -> Tuple[List[int], List[int], int]: + # Fast path: both sides use exactly the same PP layout + if len(src_kv_ptrs) == len(dst_kv_ptrs): + return src_kv_ptrs, dst_kv_ptrs, len(src_kv_ptrs) + + mla_ratios = getattr(self.kv_args, "mla_compression_ratios", None) + if mla_ratios: + # Compressed-MLA (e.g. DeepSeek V4): the flat list is organized + # by buffer type (compression-ratio bucket) rather than by + # layer, so we locate the sub-range for this PP stage inside each + # section of the dst flat list. + sliced_src_kv_ptrs, sliced_dst_kv_ptrs = self._mla_slice_ptrs_for_pp( + src_kv_ptrs, dst_kv_ptrs, mla_ratios + ) + return ( + sliced_src_kv_ptrs, + sliced_dst_kv_ptrs, + len(sliced_src_kv_ptrs), + ) + + # Regular MLA PP slicing start_layer = self.kv_args.prefill_start_layer end_layer = start_layer + len(src_kv_ptrs) - if len(src_kv_ptrs) == len(dst_kv_ptrs): - sliced_dst_kv_ptrs = dst_kv_ptrs - else: - # Decode pp size should be equal to prefill pp size or 1 - sliced_dst_kv_ptrs = dst_kv_ptrs[start_layer:end_layer] - layers_current_pp_stage = len(src_kv_ptrs) - return src_kv_ptrs, sliced_dst_kv_ptrs, layers_current_pp_stage + # Decode pp size should be equal to prefill pp size or 1 + sliced_dst_kv_ptrs = dst_kv_ptrs[start_layer:end_layer] + return src_kv_ptrs, sliced_dst_kv_ptrs, len(src_kv_ptrs) + + def _mla_slice_ptrs_for_pp( + self, + src_kv_ptrs: List[int], + dst_kv_ptrs: List[int], + mla_ratios: List[int], + ) -> Tuple[List[int], List[int]]: + """Produce aligned (src, dst) pointer lists for compressed-MLA + pools (e.g. DeepSeek V4) under PP. + + The pool produces two possible flat-list layouts (selected via dst + length): + + - kv_data layout, length = 2 * c4_L + c128_L: + [c4_layer_{0..c4_L-1}, + c4_indexer_layer_{0..c4_L-1}, + c128_layer_{0..c128_L-1}] + Each section is indexed by compressed-layer id within that + compression bucket. + + - state_data layout, length = swa_L + 2 * c4_L + c128_L: + [swa_layer_{0..swa_L-1}, + compress_state_{non-None, c4_L + c128_L}, + indexer_compress_state_{non-None, c4_L}] + ``swa_L`` is the SWA pool's actual buffer count + (``num_effective_layers``), which can be smaller than + ``len(mla_ratios)`` when the HF config's ``compress_ratios`` + list contains entries for layers not materialized into the SWA + pool (e.g. an MTP/nextn slot at the tail). + + src is already PP-filtered on the prefill side. dst is the + decode-side full-model list (when decode is PP=1). We slice dst to + match src's PP stage. If src itself is also full-model, it is + returned unchanged. + """ + start_layer = self.kv_args.prefill_start_layer + end_layer = getattr(self.kv_args, "prefill_end_layer", None) + assert end_layer is not None, ( + "KVArgs.prefill_end_layer must be set when using " + "compressed-MLA PD with PP" + ) + + c4_full = sum(1 for r in mla_ratios if r == 4) + c128_full = sum(1 for r in mla_ratios if r == 128) + kv_layout_len = 2 * c4_full + c128_full + + c4_off_s = sum(1 for r in mla_ratios[:start_layer] if r == 4) + c4_off_e = sum(1 for r in mla_ratios[:end_layer] if r == 4) + c128_off_s = sum(1 for r in mla_ratios[:start_layer] if r == 128) + c128_off_e = sum(1 for r in mla_ratios[:end_layer] if r == 128) + + if len(dst_kv_ptrs) == kv_layout_len: + sliced_dst = ( + list(dst_kv_ptrs[c4_off_s:c4_off_e]) + + list(dst_kv_ptrs[c4_full + c4_off_s : c4_full + c4_off_e]) + + list(dst_kv_ptrs[2 * c4_full + c128_off_s : 2 * c4_full + c128_off_e]) + ) + return src_kv_ptrs, sliced_dst + + # State-data layout. ``swa_L`` is derived from the actual dst + # length so we tolerate cases where the SWA pool has fewer + # buffers than ``len(mla_ratios)`` (e.g. nextn padding). + swa_L = len(dst_kv_ptrs) - 2 * c4_full - c128_full + if swa_L < 0 or swa_L > len(mla_ratios): + raise ValueError( + f"Unexpected compressed-MLA dst_kv_ptrs length " + f"{len(dst_kv_ptrs)}; expected either {kv_layout_len} " + f"(kv_data) or swa_L + {2 * c4_full + c128_full} " + f"(state_data) given compression_ratios " + f"(c4={c4_full}, c128={c128_full}, " + f"total={len(mla_ratios)})." + ) + # Guard against asking the prefill side to read past the SWA + # pool boundary. + assert end_layer <= swa_L, ( + f"prefill_end_layer ({end_layer}) exceeds dst SWA pool " + f"buffer count ({swa_L}); compression_ratios may include " + f"layers (e.g. nextn) that the SWA pool does not cover." + ) + + # compress_state non-None count up to L = count(r != 0). + c_non_zero_s = sum(1 for r in mla_ratios[:start_layer] if r != 0) + c_non_zero_e = sum(1 for r in mla_ratios[:end_layer] if r != 0) + compress_section_start = swa_L + indexer_section_start = swa_L + (c4_full + c128_full) + sliced_dst = ( + list(dst_kv_ptrs[start_layer:end_layer]) + + list( + dst_kv_ptrs[ + compress_section_start + + c_non_zero_s : compress_section_start + + c_non_zero_e + ] + ) + + list( + dst_kv_ptrs[ + indexer_section_start + c4_off_s : indexer_section_start + c4_off_e + ] + ) + ) + + return src_kv_ptrs, sliced_dst class CommonKVSender(BaseKVSender): diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 7ddcbe169d7d..867824c50ede 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -55,6 +55,7 @@ maybe_cache_unfinished_req, release_kv_cache, ) +from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.observability.req_time_stats import set_schedule_time_batch if TYPE_CHECKING: @@ -146,6 +147,8 @@ def _init_kv_manager(self) -> CommonKVManager: kv_args.pp_rank = self.pp_rank kv_args.system_dp_rank = self.scheduler.dp_rank kv_args.prefill_start_layer = self.token_to_kv_pool.start_layer + kv_args.prefill_end_layer = self.token_to_kv_pool.end_layer + kv_args.mla_compression_ratios = None kv_data_ptrs, kv_data_lens, kv_item_lens = ( self.token_to_kv_pool.get_contiguous_buf_infos() ) @@ -185,6 +188,13 @@ def _init_kv_manager(self) -> CommonKVManager: req_to_token_pool=req_to_token_pool, ) + if isinstance(self.token_to_kv_pool, DeepSeekV4TokenToKVPool): + # V4's KVCache is organized by compression-ratio + # buckets rather than by layer. + kv_args.mla_compression_ratios = list( + self.token_to_kv_pool.compression_ratios + ) + kv_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER) kv_manager = kv_manager_class( kv_args, diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index ab8cbc78fca2..feecc544160b 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -2328,7 +2328,10 @@ def prepare_for_decode(self): ) # Update fields - self.input_ids = self.output_ids + # Coerce to int64: torch sampling helpers (sampling_from_probs_torch / + # top_k_top_p_min_p_sampling_from_probs_torch) return int32 token ids, + # but downstream kernels enforce int64 (e.g. DeepSeek-V4 hash_topk). + self.input_ids = self.output_ids.to(torch.int64) self.output_ids = None if self.model_config.is_encoder_decoder: diff --git a/python/sglang/srt/managers/scheduler_pp_mixin.py b/python/sglang/srt/managers/scheduler_pp_mixin.py index 939c83f3c2f6..9a0aecb41d10 100644 --- a/python/sglang/srt/managers/scheduler_pp_mixin.py +++ b/python/sglang/srt/managers/scheduler_pp_mixin.py @@ -613,9 +613,13 @@ def profile_and_init_predictor(self: Scheduler): batch.global_num_tokens = global_num_tokens batch.global_num_tokens_for_logprob = global_num_tokens + hs = ( + getattr(model_config, "hc_hidden_size", None) + or model_config.hidden_size + ) proxy_tensors = { "hidden_states": torch.zeros( - (current_seq_len, model_config.hidden_size), + (current_seq_len, hs), dtype=model_config.dtype, device=self.device, ), 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 94b930b890c8..a185b02c5729 100644 --- a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py +++ b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py @@ -401,6 +401,19 @@ def __init__( self.state_dtype = state_dtype self.compression_ratios = compression_ratios + # Determine this PP stage's absolute layer range + if ( + start_layer is not None + and end_layer is not None + and len(compression_ratios) >= end_layer + ): + self._stage_start = start_layer + self._stage_end = end_layer + else: + self._stage_start = 0 + self._stage_end = len(compression_ratios) + stage_ratios = compression_ratios[self._stage_start : self._stage_end] + assert page_size % swa_page_size == 0 self.swa_size = swa_size @@ -412,8 +425,8 @@ def __init__( self.qk_rope_head_dim = qk_rope_head_dim self.indexer_head_dim = indexer_head_dim - c4_layer_num = sum(1 for r in compression_ratios if r == 4) - c128_layer_num = sum(1 for r in compression_ratios if r == 128) + c4_layer_num = sum(1 for r in stage_ratios if r == 4) + c128_layer_num = sum(1 for r in stage_ratios if r == 128) c4_page_size = page_size // 4 c128_page_size = page_size // 128 self.swa_kv_pool = DeepSeekV4SingleKVPool( @@ -467,6 +480,7 @@ def __init__( self._init_paged_compress_states(enable_memory_saver) self._should_cache_swa = envs.SGLANG_OPT_CACHE_SWA_TRANSLATION.get() + self.cached_loc = None def register_mapping(self, full_to_swa_index_mapping: torch.Tensor): self.full_to_swa_index_mapping = full_to_swa_index_mapping @@ -535,29 +549,34 @@ def get_state_buf_infos(self) -> Tuple[List[int], List[int], List[int]]: def _init_paged_compress_states(self, enable_memory_saver: bool): c4_state_pool_size = self.c4_state_pool_size c128_state_pool_size = self.c128_state_pool_size - self.compress_state_pools: List[CompressStatePool] = [] - self.indexer_compress_state_pools: List[CompressStatePool] = [] - - for ratio in self.compression_ratios: + total_L = len(self.compression_ratios) + self.compress_state_pools: List[Optional[CompressStatePool]] = [None] * total_L + self.indexer_compress_state_pools: List[Optional[CompressStatePool]] = [ + None + ] * total_L + + for idx in range(self._stage_start, self._stage_end): + ratio = self.compression_ratios[idx] + if ratio == 0: + continue overlap = ratio == 4 - compress_state_pool = indexer_compress_state_pool = None size = c4_state_pool_size if ratio == 4 else c128_state_pool_size - ring_size = self.get_ring_size(ratio) if ratio != 0 else 0 - if ratio != 0: - compress_state_pool = CompressStatePool( - size=size, - ring_size=ring_size, - overlap=overlap, - head_dim=self.qk_nope_head_dim + self.qk_rope_head_dim, - dtype=self.state_dtype, - device=self.device, - enable_memory_saver=enable_memory_saver, - ratio=ratio, - online=(ratio == 128 and ONLINE_C128), - ) + ring_size = self.get_ring_size(ratio) + + self.compress_state_pools[idx] = CompressStatePool( + size=size, + ring_size=ring_size, + overlap=overlap, + head_dim=self.qk_nope_head_dim + self.qk_rope_head_dim, + dtype=self.state_dtype, + device=self.device, + enable_memory_saver=enable_memory_saver, + ratio=ratio, + online=(ratio == 128 and ONLINE_C128), + ) if ratio == 4: - indexer_compress_state_pool = CompressStatePool( + self.indexer_compress_state_pools[idx] = CompressStatePool( size=size, ring_size=ring_size, overlap=overlap, @@ -568,38 +587,31 @@ def _init_paged_compress_states(self, enable_memory_saver: bool): ratio=ratio, ) - self.compress_state_pools.append(compress_state_pool) - self.indexer_compress_state_pools.append(indexer_compress_state_pool) - def _init_compressed_layer_mapping(self): - c1_cnt, c4_cnt, c128_cnt = 0, 0, 0 - self.layer_mapping: List[DeepSeekV4LayerItem] = [] + c1_cnt = c4_cnt = c128_cnt = 0 + total_L = len(self.compression_ratios) + self.layer_mapping: List[Optional[DeepSeekV4LayerItem]] = [None] * total_L - for ratio in self.compression_ratios: + for idx in range(self._stage_start, self._stage_end): + ratio = self.compression_ratios[idx] if ratio == 0: - self.layer_mapping.append( - DeepSeekV4LayerItem( - compress_ratio=0, - compress_layer_id=c1_cnt, - ) + self.layer_mapping[idx] = DeepSeekV4LayerItem( + compress_ratio=0, + compress_layer_id=c1_cnt, ) c1_cnt += 1 elif ratio == 4: - self.layer_mapping.append( - DeepSeekV4LayerItem( - compress_ratio=4, - compress_layer_id=c4_cnt, - compress_kv_pool=self.c4_kv_pool, - ) + self.layer_mapping[idx] = DeepSeekV4LayerItem( + compress_ratio=4, + compress_layer_id=c4_cnt, + compress_kv_pool=self.c4_kv_pool, ) c4_cnt += 1 elif ratio == 128: - self.layer_mapping.append( - DeepSeekV4LayerItem( - compress_ratio=128, - compress_layer_id=c128_cnt, - compress_kv_pool=self.c128_kv_pool, - ) + self.layer_mapping[idx] = DeepSeekV4LayerItem( + compress_ratio=128, + compress_layer_id=c128_cnt, + compress_kv_pool=self.c128_kv_pool, ) c128_cnt += 1 else: @@ -625,9 +637,13 @@ def get_indexer_compress_states(self, layer_id: int) -> CompressStatePool: ), "Only c4 layers have indexer states." return indexer_compress_state_pool + def _swa_local_layer_id(self, layer_id: int) -> int: + """Convert absolute model layer_id to SWA-pool-local (PP-stage-local) index.""" + return layer_id - self._stage_start + def get_swa_key_buffer(self, layer_id: int) -> torch.Tensor: self.wait_layer_transfer(layer_id) - return self.swa_kv_pool.get_key_buffer(layer_id) + return self.swa_kv_pool.get_key_buffer(self._swa_local_layer_id(layer_id)) def set_swa_key_buffer( self, @@ -635,7 +651,9 @@ def set_swa_key_buffer( loc: torch.Tensor, cache_nope_fp8_rope_bf16_pack: NopeFp8RopeBf16Pack, ) -> None: - self.swa_kv_pool.set_key_buffer(layer_id, loc, cache_nope_fp8_rope_bf16_pack) + self.swa_kv_pool.set_key_buffer( + self._swa_local_layer_id(layer_id), loc, cache_nope_fp8_rope_bf16_pack + ) def get_extra_key_page_size(self, layer_id: int) -> int: _, _, compress_kv_pool = self.layer_mapping[layer_id] @@ -715,12 +733,12 @@ def set_swa_key_buffer_radix( ) -> None: swa_loc = self.translate_loc_from_full_to_swa(raw_loc) self.swa_kv_pool.set_key_buffer( - layer_id, swa_loc, cache_nope_fp8_rope_bf16_pack + self._swa_local_layer_id(layer_id), swa_loc, cache_nope_fp8_rope_bf16_pack ) def get_swa_key_buffer_radix(self, layer_id: int) -> torch.Tensor: self.wait_layer_transfer(layer_id) - return self.swa_kv_pool.get_key_buffer(layer_id) + return self.swa_kv_pool.get_key_buffer(self._swa_local_layer_id(layer_id)) def set_swa_key_buffer_radix_fused( self, @@ -729,12 +747,14 @@ def set_swa_key_buffer_radix_fused( cache_k: torch.Tensor, ) -> None: if self._should_cache_swa: - if layer_id == 0: + if layer_id == self.start_layer or self.cached_loc is None: self.cached_loc = self.translate_loc_from_full_to_swa(raw_loc) swa_loc = self.cached_loc else: swa_loc = self.translate_loc_from_full_to_swa(raw_loc) - return self.swa_kv_pool.set_key_buffer_fused(layer_id, swa_loc, cache_k) + return self.swa_kv_pool.set_key_buffer_fused( + self._swa_local_layer_id(layer_id), swa_loc, cache_k + ) def set_swa_key_buffer_radix_fused_norm_rope( self, @@ -759,7 +779,7 @@ def set_swa_key_buffer_radix_fused_norm_rope( freqs_cis=freqs_cis, positions=positions, out_loc=swa_loc, - kvcache=self.swa_kv_pool.kv_buffer[layer_id], + kvcache=self.swa_kv_pool.kv_buffer[self._swa_local_layer_id(layer_id)], page_size=self.swa_kv_pool.page_size, ) diff --git a/python/sglang/srt/model_executor/cuda_graph_runner.py b/python/sglang/srt/model_executor/cuda_graph_runner.py index 55118fa17e17..c037b20dd9af 100644 --- a/python/sglang/srt/model_executor/cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/cuda_graph_runner.py @@ -174,6 +174,7 @@ def create( enable_mamba_track: bool, ne_token_table: Optional[torch.Tensor] = None, is_hybrid_swa: bool = False, + hc_hidden_size: Optional[int] = None, ) -> "DecodeInputBuffers": with torch.device(device): input_ids = torch.zeros((max_num_token,), dtype=torch.int64) @@ -207,10 +208,16 @@ def create( ) if pp_size > 1: + # 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 pp_proxy_tensors = { - "hidden_states": torch.zeros((max_bs, hidden_size), dtype=dtype), - "residual": torch.zeros((max_bs, hidden_size), dtype=dtype), + "hidden_states": torch.zeros((max_bs, hs), dtype=dtype), } + if not is_mhc: + pp_proxy_tensors["residual"] = torch.zeros( + (max_bs, hidden_size), dtype=dtype + ) else: pp_proxy_tensors = None @@ -715,6 +722,9 @@ def __init__( model_runner.token_table if self.use_ngram_embedding else None ), is_hybrid_swa=model_runner.is_hybrid_swa, + hc_hidden_size=getattr( + self.model_runner.model_config, "hc_hidden_size", None + ), ) self.buffers.share_buffers() diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index f4128884ad58..b8df0c3135e2 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -2,7 +2,16 @@ import concurrent.futures import logging -from typing import TYPE_CHECKING, Iterable, List, Literal, Optional, Set, Tuple +from typing import ( + TYPE_CHECKING, + Iterable, + List, + Literal, + Optional, + Set, + Tuple, + Union, +) import torch import torch.nn as nn @@ -49,7 +58,7 @@ from sglang.srt.layers.moe import get_moe_a2a_backend from sglang.srt.layers.moe.fused_moe_triton import FusedMoE from sglang.srt.layers.quantization.fp8_kernel import sglang_per_token_group_quant_fp8 -from sglang.srt.layers.utils import get_layer_id +from sglang.srt.layers.utils import PPMissingLayer, get_layer_id from sglang.srt.layers.utils.cp_utils import ( cp_all_gather_rerange_output, cp_split_and_rebuild_data, @@ -62,6 +71,7 @@ compile_in_capture_mode, get_is_capture_mode, ) +from sglang.srt.model_executor.forward_batch_info import PPProxyTensors from sglang.srt.model_loader.utils import maybe_executor_submit, should_async_load from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.dbrx import ReplicatedLinear @@ -86,10 +96,7 @@ ) from sglang.srt.layers.quantization import QuantizationConfig from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool - from sglang.srt.model_executor.forward_batch_info import ( - ForwardBatch, - PPProxyTensors, - ) + from sglang.srt.model_executor.forward_batch_info import ForwardBatch @triton.jit @@ -870,11 +877,15 @@ def __init__( ) -> None: super().__init__() self.pp_group = get_pp_group() - self.embed_tokens = VocabParallelEmbedding( - config.vocab_size, - config.hidden_size, - enable_tp=not is_dp_attention_enabled(), - ) + self.hidden_size = config.hidden_size + if self.pp_group.is_first_rank: + self.embed_tokens = VocabParallelEmbedding( + config.vocab_size, + config.hidden_size, + enable_tp=not is_dp_attention_enabled(), + ) + else: + self.embed_tokens = PPMissingLayer() self.rms_norm_eps = config.rms_norm_eps self.alt_streams = ( [torch.cuda.Stream() for _ in range(5)] if (_is_cuda or _is_hip) else None @@ -892,17 +903,21 @@ def __init__( pp_size=self.pp_group.world_size, prefix=add_prefix("layers", prefix), ) - self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + if self.pp_group.is_last_rank: + self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + else: + self.norm = PPMissingLayer() self.gemm_output_zero_allocator_size = 0 self.hc_eps = config.hc_eps self.hc_mult = hc_mult = config.hc_mult self.norm_eps = config.rms_norm_eps - hc_dim = hc_mult * config.hidden_size - self.hc_head_fn = nn.Parameter( - torch.empty(hc_mult, hc_dim, dtype=torch.float32) - ) - self.hc_head_base = nn.Parameter(torch.empty(hc_mult, dtype=torch.float32)) - self.hc_head_scale = nn.Parameter(torch.empty(1, dtype=torch.float32)) + if self.pp_group.is_last_rank: + hc_dim = hc_mult * config.hidden_size + self.hc_head_fn = nn.Parameter( + torch.empty(hc_mult, hc_dim, dtype=torch.float32) + ) + self.hc_head_base = nn.Parameter(torch.empty(hc_mult, dtype=torch.float32)) + self.hc_head_scale = nn.Parameter(torch.empty(1, dtype=torch.float32)) self.nsa_enable_prefill_cp = is_nsa_enable_prefill_cp() if self.nsa_enable_prefill_cp: @@ -940,9 +955,19 @@ def forward( positions: torch.Tensor, forward_batch: ForwardBatch, input_embeds: Optional[torch.Tensor], - ) -> torch.Tensor: - hidden_states = self.embed_tokens(input_ids) - hidden_states = hidden_states.unsqueeze(1).repeat(1, self.hc_mult, 1) + pp_proxy_tensors: Optional[PPProxyTensors] = None, + ) -> Union[torch.Tensor, PPProxyTensors]: + if self.pp_group.is_first_rank: + hidden_states = self.embed_tokens(input_ids) + hidden_states = hidden_states.unsqueeze(1).repeat(1, self.hc_mult, 1) + else: + assert pp_proxy_tensors is not None + hidden_states = pp_proxy_tensors["hidden_states"] + # Unflatten 2D PP IPC tensor back to 3D mHC shape. + if hidden_states.ndim == 2: + hidden_states = hidden_states.view( + hidden_states.shape[0], self.hc_mult, self.hidden_size + ) if get_attention_dp_size() > 1 and get_moe_a2a_backend().is_none(): input_ids_global = torch.empty( @@ -956,7 +981,8 @@ def forward( input_ids_global = input_ids if nsa_use_prefill_cp(forward_batch): - hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states) + if self.pp_group.is_first_rank: + hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states) positions = cp_split_and_rebuild_position(forward_batch, positions) for i in range(self.start_layer, self.end_layer): @@ -969,7 +995,8 @@ def forward( input_ids_global=input_ids_global, ) - if nsa_use_prefill_cp(forward_batch): + # CP all-gather only on the last PP rank; PP IPC carries CP-split tensors. + if self.pp_group.is_last_rank and nsa_use_prefill_cp(forward_batch): hidden_states = cp_all_gather_rerange_output( hidden_states, self.cp_size, @@ -977,6 +1004,10 @@ def forward( torch.cuda.current_stream(), ) + if not self.pp_group.is_last_rank: + # Flatten 3D mHC tensor for PP IPC. + return PPProxyTensors({"hidden_states": hidden_states.flatten(1)}) + pre_hc_head = hidden_states.flatten(1) hidden_states = self.hc_head( @@ -1003,28 +1034,37 @@ def __init__( config, quant_config, prefix=add_prefix("model", prefix) ) self.pp_group = get_pp_group() - if config.tie_word_embeddings: - self.lm_head = self.model.embed_tokens + if self.pp_group.is_last_rank: + if self.pp_group.world_size == 1 and config.tie_word_embeddings: + self.lm_head = self.model.embed_tokens + else: + self.lm_head = ParallelLMHead( + config.vocab_size, + config.hidden_size, + quant_config=quant_config, + prefix=add_prefix("lm_head", prefix), + use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + ) else: - self.lm_head = ParallelLMHead( - config.vocab_size, - config.hidden_size, - quant_config=quant_config, - prefix=add_prefix("lm_head", prefix), - use_attn_tp_group=get_global_server_args().enable_dp_lm_head, - ) + self.lm_head = PPMissingLayer() self.logits_processor = LogitsProcessor(config) self.capture_aux_hidden_states = False get_attn_tp_context().init_context(config.q_lora_rank, is_nsa=True) self._routed_experts_weights_of_layer = LazyValue( lambda: { - layer_id: layer.mlp.get_moe_weights() - for layer_id, layer in enumerate(self.model.layers) - if isinstance(layer.mlp, deepseek_v2.DeepseekV2MoE) + layer_id: self.model.layers[layer_id].mlp.get_moe_weights() + for layer_id in range(self.model.start_layer, self.model.end_layer) + if isinstance( + self.model.layers[layer_id].mlp, deepseek_v2.DeepseekV2MoE + ) } ) + # Expose start_layer/end_layer for model_runner PP support + self.start_layer = self.model.start_layer + self.end_layer = self.model.end_layer + self.nsa_enable_prefill_cp = is_nsa_enable_prefill_cp() if self.nsa_enable_prefill_cp: self.cp_rank = get_attention_cp_rank() @@ -1077,8 +1117,11 @@ def forward( with get_attn_tp_context().maybe_input_scattered(forward_batch): hidden_states = self.model.forward( - input_ids, positions, forward_batch, input_embeds + input_ids, positions, forward_batch, input_embeds, pp_proxy_tensors ) + if not self.pp_group.is_last_rank: + return hidden_states + aux_hidden_states = None if self.capture_aux_hidden_states: hidden_states, aux_hidden_states = hidden_states @@ -1098,7 +1141,10 @@ def _setup_fp8_wo_a_scales(self, is_nextn: bool) -> None: if is_nextn: layers = [self.model.decoder] else: - layers = self.model.layers + layers = [ + self.model.layers[layer_id] + for layer_id in range(self.model.start_layer, self.model.end_layer) + ] for layer in layers: attn = layer.self_attn G = attn.n_local_groups @@ -1121,7 +1167,8 @@ def post_load_weights(self, is_nextn=False, weight_names=None): if is_nextn: return - for layer in self.model.layers: + for layer_id in range(self.model.start_layer, self.model.end_layer): + layer = self.model.layers[layer_id] self_attn = layer.self_attn if self_attn.compress_ratio != 0 and not self_attn.compressor.ape_converted: self_attn.compressor.apply_ape_hotfix() @@ -1389,7 +1436,15 @@ def auto_weight_loader(module): and not self.pp_group.is_first_rank ): continue - if ".norm." in name and not self.pp_group.is_last_rank: + if ( + name == "model.norm.weight" + and not self.pp_group.is_last_rank + ): + continue + if ( + name.startswith("model.hc_head_") + or name == "lm_head.weight" + ) and not self.pp_group.is_last_rank: continue elif COMPRESSOR_PART in name: is_kv = name.endswith(".wkv.weight") @@ -1493,6 +1548,11 @@ def auto_weight_loader(module): unloaded_params = params_dict.keys() - loaded_params skipped_checking_patterns = ["attn_mqa.k_scale", "attn_mqa.v_scale"] + if not self.pp_group.is_first_rank: + skipped_checking_patterns.append("embed_tokens") + if not self.pp_group.is_last_rank: + skipped_checking_patterns.append("model.norm.") + skipped_checking_patterns.extend(["lm_head", "hc_head_"]) if is_nextn: skipped_checking_patterns.extend(["lm_head", "embed_tokens"]) unloaded_params = { From a637481f7b290bc1ca13da7abae860c895621f45 Mon Sep 17 00:00:00 2001 From: Yuhao Yang <47235274+yhyang201@users.noreply.github.com> Date: Sat, 16 May 2026 15:18:43 +0800 Subject: [PATCH 09/50] [MoE] Decouple Mega MoE from DeepEP backend (#25406) --- python/sglang/srt/environ.py | 5 ++-- .../srt/layers/moe/fused_moe_triton/layer.py | 2 +- python/sglang/srt/layers/moe/mega_moe.py | 3 +- .../srt/layers/moe/moe_runner/deep_gemm.py | 1 - python/sglang/srt/layers/moe/utils.py | 4 +++ python/sglang/srt/layers/quantization/fp8.py | 2 +- python/sglang/srt/models/deepseek_v2.py | 1 + python/sglang/srt/server_args.py | 29 ++++++++++++++++++- test/manual/dsv4/test_b200_flash.py | 1 - test/manual/dsv4/test_b200_pro.py | 1 - test/manual/dsv4/test_b300_flash.py | 1 - test/manual/dsv4/test_b300_pro.py | 1 - test/manual/dsv4/test_dsv4_flash_mtp_tp8.py | 2 -- test/manual/dsv4/test_dsv4_pro_mtp.py | 2 -- test/manual/dsv4/test_gb300_flash.py | 1 - test/manual/dsv4/test_gb300_pro.py | 1 - test/manual/dsv4/test_h200_fp8_flash.py | 1 - ...test_deepseek_v4_flash_fp4_megamoe_b200.py | 12 ++------ 18 files changed, 41 insertions(+), 29 deletions(-) diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index a194ac05ec15..c78ad659826d 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -584,7 +584,7 @@ class Envs: SGLANG_OPT_USE_TILELANG_MHC_PRE = EnvBool(True) SGLANG_OPT_USE_TILELANG_MHC_POST = EnvBool(True) SGLANG_OPT_USE_TILELANG_INDEXER = EnvBool(False) - SGLANG_OPT_USE_JIT_INDEXER_METADATA = EnvBool(False) + SGLANG_OPT_USE_JIT_INDEXER_METADATA = EnvBool(True) SGLANG_OPT_USE_ONLINE_COMPRESS = EnvBool(False) SGLANG_OPT_USE_COMPRESSOR_V2 = EnvBool(True) SGLANG_FP8_PAGED_MQA_LOGITS_TORCH = EnvBool(False) @@ -616,13 +616,12 @@ class Envs: # TopK SGLANG_OPT_USE_FUSED_HASH_TOPK = EnvBool(True) SGLANG_OPT_USE_JIT_KERNEL_FUSED_TOPK = EnvBool(True) - SGLANG_OPT_USE_TOPK_V2 = EnvBool(False) + SGLANG_OPT_USE_TOPK_V2 = EnvBool(True) # GEMM / kernel fusion SGLANG_OPT_FP8_WO_A_GEMM = EnvBool(True) SGLANG_OPT_BF16_FP32_GEMM_ALGO = EnvStr("cublas") SGLANG_OPT_USE_JIT_EP_ACTIVATION = EnvBool(True) - SGLANG_OPT_USE_JIT_NORM = EnvBool(False) SGLANG_OPT_FUSE_WQA_WKV = EnvBool(True) SGLANG_OPT_SWIGLU_CLAMP_FUSION = EnvBool(True) diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py index 56eaf8a1e8bf..97a5bfc3d2d4 100644 --- a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py +++ b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py @@ -82,7 +82,7 @@ def create_moe_dispatcher(moe_runner_config: MoeRunnerConfig) -> BaseDispatcher: a2a_backend = get_moe_a2a_backend() - if a2a_backend.is_none(): + if a2a_backend.is_none() or a2a_backend.is_megamoe(): return StandardDispatcher(moe_runner_config) elif ( a2a_backend.is_deepep() diff --git a/python/sglang/srt/layers/moe/mega_moe.py b/python/sglang/srt/layers/moe/mega_moe.py index 1c13f7be9885..9574bd2da68d 100644 --- a/python/sglang/srt/layers/moe/mega_moe.py +++ b/python/sglang/srt/layers/moe/mega_moe.py @@ -25,6 +25,7 @@ from sglang.srt.environ import envs from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo from sglang.srt.layers.dp_attention import get_dp_global_num_tokens +from sglang.srt.layers.moe.utils import get_moe_a2a_backend from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode if TYPE_CHECKING: @@ -94,7 +95,7 @@ def _get_mega_moe_symm_buffer( def should_use_mega_moe(moe: "DeepseekV2MoE", hidden_states: torch.Tensor) -> bool: - if not envs.SGLANG_OPT_USE_DEEPGEMM_MEGA_MOE.get(): + if not get_moe_a2a_backend().is_megamoe(): return False if not getattr(moe.experts, "_mega_moe_weights_built", False): return False diff --git a/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py b/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py index da6f13fcd5e1..61af5533f5b3 100644 --- a/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py +++ b/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py @@ -131,7 +131,6 @@ def __init__(self, config: MoeRunnerConfig): if envs.SGLANG_OPT_FIX_MEGA_MOE_MEMORY.get(): assert envs.SGLANG_OPT_SWIGLU_CLAMP_FUSION.get() assert envs.SGLANG_OPT_USE_JIT_EP_ACTIVATION.get() - assert envs.SGLANG_OPT_USE_DEEPGEMM_MEGA_MOE.get() self.use_swizzle = True def run( diff --git a/python/sglang/srt/layers/moe/utils.py b/python/sglang/srt/layers/moe/utils.py index e05167da972a..fbca714d4130 100644 --- a/python/sglang/srt/layers/moe/utils.py +++ b/python/sglang/srt/layers/moe/utils.py @@ -29,6 +29,7 @@ class MoeA2ABackend(Enum): MORI = "mori" ASCEND_FUSEEP = "ascend_fuseep" FLASHINFER = "flashinfer" + MEGAMOE = "megamoe" CUSTOMIZED = "customized" @classmethod @@ -61,6 +62,9 @@ def is_ascend_fuseep(self): def is_mori(self): return self == MoeA2ABackend.MORI + def is_megamoe(self): + return self == MoeA2ABackend.MEGAMOE + def is_customized(self): return self == MoeA2ABackend.CUSTOMIZED diff --git a/python/sglang/srt/layers/quantization/fp8.py b/python/sglang/srt/layers/quantization/fp8.py index 99d86a56ea32..f1d5cc4a9396 100644 --- a/python/sglang/srt/layers/quantization/fp8.py +++ b/python/sglang/srt/layers/quantization/fp8.py @@ -1193,7 +1193,7 @@ def process_weights_after_loading_block_quant(self, layer: Module) -> None: layer.w13_weight.data = layer.w13_weight.data.view(torch.int8) layer.w2_weight.data = layer.w2_weight.data.view(torch.int8) - if envs.SGLANG_OPT_USE_DEEPGEMM_MEGA_MOE.get(): + if get_moe_a2a_backend().is_megamoe(): from sglang.srt.layers.moe.mega_moe import ( build_mega_moe_experts_weights, ) diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index e1b77562ce47..adf440672488 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -601,6 +601,7 @@ def __init__( or get_moe_a2a_backend().is_mori() or get_moe_a2a_backend().is_ascend_fuseep() or get_moe_a2a_backend().is_flashinfer() + or get_moe_a2a_backend().is_megamoe() or should_use_flashinfer_cutlass_moe_fp4_allgather() or envs.SGLANG_SHARED_EXPERT_TP1.get() ) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 8ebdf088c5e7..173bdbbf4c55 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -209,6 +209,7 @@ "mori", "ascend_fuseep", "flashinfer", + "megamoe", ] FP8_GEMM_RUNNER_BACKEND_CHOICES = [ @@ -610,7 +611,14 @@ class ServerArgs: # Expert parallelism ep_size: int = 1 moe_a2a_backend: Literal[ - "none", "deepep", "mooncake", "nixl", "mori", "ascend_fuseep", "flashinfer" + "none", + "deepep", + "mooncake", + "nixl", + "mori", + "ascend_fuseep", + "flashinfer", + "megamoe", ] = "none" moe_runner_backend: str = "auto" record_nolora_graph: bool = True @@ -3185,6 +3193,25 @@ def _handle_a2a_moe(self): ) self.moe_a2a_backend = "deepep" + if ( + envs.SGLANG_OPT_USE_DEEPGEMM_MEGA_MOE.get() + and self.moe_a2a_backend != "megamoe" + ): + self.moe_a2a_backend = "megamoe" + logger.info( + "SGLANG_OPT_USE_DEEPGEMM_MEGA_MOE is set, " + "auto-configuring --moe-a2a-backend megamoe." + ) + + if self.moe_a2a_backend == "megamoe": + self.ep_size = self.tp_size + if not envs.SGLANG_OPT_FIX_MEGA_MOE_MEMORY.is_set(): + envs.SGLANG_OPT_FIX_MEGA_MOE_MEMORY.set(True) + logger.info( + f"Mega MoE is enabled. The expert parallel size is adjusted " + f"to be the same as the tensor parallel size[{self.tp_size}]." + ) + if self.moe_a2a_backend == "deepep": if self.deepep_mode == "normal": logger.warning("Cuda graph is disabled because deepep_mode=`normal`") diff --git a/test/manual/dsv4/test_b200_flash.py b/test/manual/dsv4/test_b200_flash.py index 3c828116190b..05d738bcf13e 100644 --- a/test/manual/dsv4/test_b200_flash.py +++ b/test/manual/dsv4/test_b200_flash.py @@ -100,7 +100,6 @@ class TestB200FlashCP(DSV4FlashAime25TestBase): DEEPEP_LARGE_SMS_CONFIG, ] EXTRA_ENV = { - "SGLANG_OPT_USE_JIT_INDEXER_METADATA": "1", "SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "1024", } diff --git a/test/manual/dsv4/test_b200_pro.py b/test/manual/dsv4/test_b200_pro.py index 67eaadd21c74..eba9b65460af 100644 --- a/test/manual/dsv4/test_b200_pro.py +++ b/test/manual/dsv4/test_b200_pro.py @@ -116,7 +116,6 @@ class TestB200ProCP(DSV4ProAime25TestBase): DEEPEP_LARGE_SMS_CONFIG, ] EXTRA_ENV = { - "SGLANG_OPT_USE_JIT_INDEXER_METADATA": "1", "SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "256", } diff --git a/test/manual/dsv4/test_b300_flash.py b/test/manual/dsv4/test_b300_flash.py index 261bac4764db..4e800526d67d 100644 --- a/test/manual/dsv4/test_b300_flash.py +++ b/test/manual/dsv4/test_b300_flash.py @@ -102,7 +102,6 @@ class TestB300FlashCP(DSV4FlashAime25TestBase): DEEPEP_LARGE_SMS_CONFIG, ] EXTRA_ENV = { - "SGLANG_OPT_USE_JIT_INDEXER_METADATA": "1", "SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "1024", } diff --git a/test/manual/dsv4/test_b300_pro.py b/test/manual/dsv4/test_b300_pro.py index 1ed254e5253d..701a42a39fb8 100644 --- a/test/manual/dsv4/test_b300_pro.py +++ b/test/manual/dsv4/test_b300_pro.py @@ -118,7 +118,6 @@ class TestB300ProCP(DSV4ProAime25TestBase): DEEPEP_LARGE_SMS_CONFIG, ] EXTRA_ENV = { - "SGLANG_OPT_USE_JIT_INDEXER_METADATA": "1", "SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "256", } diff --git a/test/manual/dsv4/test_dsv4_flash_mtp_tp8.py b/test/manual/dsv4/test_dsv4_flash_mtp_tp8.py index 4913f0bb96bd..87c63f9c3351 100644 --- a/test/manual/dsv4/test_dsv4_flash_mtp_tp8.py +++ b/test/manual/dsv4/test_dsv4_flash_mtp_tp8.py @@ -24,9 +24,7 @@ DSV4_FLASH_BASE_ENV = { "SGLANG_ENABLE_SPEC_V2": "1", - "SGLANG_OPT_USE_TOPK_V2": "1", "SGLANG_DSV4_FP4_EXPERTS": "0", - "SGLANG_JIT_DEEPGEMM_PRECOMPILE": "0", } DSV4_FLASH_SERVER_ARGS = [ diff --git a/test/manual/dsv4/test_dsv4_pro_mtp.py b/test/manual/dsv4/test_dsv4_pro_mtp.py index d989f22db4e0..7e5cf62ae726 100644 --- a/test/manual/dsv4/test_dsv4_pro_mtp.py +++ b/test/manual/dsv4/test_dsv4_pro_mtp.py @@ -41,9 +41,7 @@ DSV4_PRO_BASE_ENV = { "SGLANG_ENABLE_SPEC_V2": "1", - "SGLANG_OPT_USE_TOPK_V2": "1", "SGLANG_OPT_USE_CUSTOM_ALL_REDUCE_V2": "1", - "SGLANG_JIT_DEEPGEMM_PRECOMPILE": "0", } DSV4_PRO_SERVER_ARGS = [ diff --git a/test/manual/dsv4/test_gb300_flash.py b/test/manual/dsv4/test_gb300_flash.py index 4e7a7ec5e582..ef997f50ec32 100644 --- a/test/manual/dsv4/test_gb300_flash.py +++ b/test/manual/dsv4/test_gb300_flash.py @@ -100,7 +100,6 @@ class TestGB300FlashCP(DSV4FlashAime25TestBase): DEEPEP_LARGE_SMS_CONFIG, ] EXTRA_ENV = { - "SGLANG_OPT_USE_JIT_INDEXER_METADATA": "1", "SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "1024", } diff --git a/test/manual/dsv4/test_gb300_pro.py b/test/manual/dsv4/test_gb300_pro.py index 2e186591b3a6..1c53129b2923 100644 --- a/test/manual/dsv4/test_gb300_pro.py +++ b/test/manual/dsv4/test_gb300_pro.py @@ -118,7 +118,6 @@ class TestGB300ProCP(DSV4ProAime25TestBase): DEEPEP_LARGE_SMS_CONFIG, ] EXTRA_ENV = { - "SGLANG_OPT_USE_JIT_INDEXER_METADATA": "1", "SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "256", } diff --git a/test/manual/dsv4/test_h200_fp8_flash.py b/test/manual/dsv4/test_h200_fp8_flash.py index 2aacca9ae0b0..fe53ded83f93 100644 --- a/test/manual/dsv4/test_h200_fp8_flash.py +++ b/test/manual/dsv4/test_h200_fp8_flash.py @@ -112,7 +112,6 @@ class TestH200Fp8FlashCP(DSV4FlashAime25TestBase): ] EXTRA_ENV = { **H200_FP8_ENV, - "SGLANG_OPT_USE_JIT_INDEXER_METADATA": "1", "SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "1024", } diff --git a/test/registered/dsv4/test_deepseek_v4_flash_fp4_megamoe_b200.py b/test/registered/dsv4/test_deepseek_v4_flash_fp4_megamoe_b200.py index 1c4df29c483c..7ef83c3027bf 100644 --- a/test/registered/dsv4/test_deepseek_v4_flash_fp4_megamoe_b200.py +++ b/test/registered/dsv4/test_deepseek_v4_flash_fp4_megamoe_b200.py @@ -28,20 +28,12 @@ _W4A8_MEGAMOE_ENV = { - "SGLANG_OPT_USE_DEEPGEMM_MEGA_MOE": "1", - "SGLANG_OPT_FIX_MEGA_MOE_MEMORY": "1", - "SGLANG_OPT_FIX_NEXTN_MEGA_MOE": "1", "SGLANG_OPT_DEEPGEMM_MEGA_MOE_NUM_MAX_TOKENS_PER_RANK": "4096", - "SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "0", } _W4A4_MEGAMOE_ENV = { - "SGLANG_OPT_USE_DEEPGEMM_MEGA_MOE": "1", - "SGLANG_OPT_FIX_MEGA_MOE_MEMORY": "1", - "SGLANG_OPT_FIX_NEXTN_MEGA_MOE": "1", "SGLANG_OPT_DEEPGEMM_MEGA_MOE_NUM_MAX_TOKENS_PER_RANK": "4096", - "SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "0", "SGLANG_OPT_DEEPGEMM_MEGA_MOE_USE_FP4_ACTS": "1", "SGLANG_OPT_DEEPGEMM_MEGA_MOE_USE_MXF4_KIND": "1", } @@ -81,7 +73,7 @@ def setUpClass(cls): "4", "--enable-dp-attention", "--moe-a2a-backend", - "deepep", + "megamoe", "--speculative-algorithm", "EAGLE", "--speculative-num-steps", @@ -122,7 +114,7 @@ def setUpClass(cls): "4", "--enable-dp-attention", "--moe-a2a-backend", - "deepep", + "megamoe", "--speculative-algorithm", "EAGLE", "--speculative-num-steps", From d52c6436ead438ded56581da9efed7385ad5de92 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Fri, 15 May 2026 04:09:15 -0700 Subject: [PATCH 10/50] pr-states: fix fork-PR token + add run-ci label awareness (#25392) --- .github/workflows/pr-states.yml | 31 ++++++++++++++++------- scripts/ci/utils/slash_command_handler.py | 4 --- 2 files changed, 22 insertions(+), 13 deletions(-) diff --git a/.github/workflows/pr-states.yml b/.github/workflows/pr-states.yml index 6c6b72328c2e..794d82379b39 100644 --- a/.github/workflows/pr-states.yml +++ b/.github/workflows/pr-states.yml @@ -1,8 +1,11 @@ name: PR States -# Maintains a CI-states block at the bottom of the PR body. +# Maintains a CI-states block at the bottom of the PR body. Triggered by +# `pull_request_target` (not `pull_request`) so fork PRs get a write-enabled +# GITHUB_TOKEN. Safe because the workflow never checks out PR head code — +# only reads metadata via API and PATCHes the PR body. on: - pull_request: + pull_request_target: types: [opened, synchronize, reopened, labeled, unlabeled] permissions: @@ -23,6 +26,7 @@ jobs: script: | const sha = context.payload.pull_request.head.sha; const labels = context.payload.pull_request.labels.map(l => l.name); + const hasCI = labels.includes('run-ci'); const hasExtra = labels.includes('run-ci-extra'); // Retry briefly: pr-test* may not be API-visible yet when we start. @@ -45,20 +49,29 @@ jobs: return null; } - const ptRun = await findRunWithRetry('pr-test.yml'); - const peRun = hasExtra ? await findRunWithRetry('pr-test-extra.yml') : null; + const ptRun = hasCI ? await findRunWithRetry('pr-test.yml') : null; + // pr-test-extra's gate requires BOTH run-ci AND run-ci-extra + // (see pr-test-extra.yml check-changes if-condition), so without + // run-ci the extra workflow doesn't run either. + const peRun = (hasCI && hasExtra) ? await findRunWithRetry('pr-test-extra.yml') : null; // Treat a fully-skipped run as "no real run" — happens when the PR // was opened without the label and label was added later (GHA does // not retrigger on `labeled`). const isReal = (run) => run && run.conclusion !== 'skipped'; - const notEnabledText = ':warning: **Not enabled** — add `run-ci-extra` label to opt in.'; + const missingCIText = ':x: **Missing `run-ci` label** — add it to run CI tests.'; + const peBlockedByCIText = ':x: **Blocked** — `run-ci` is required first.'; + const notExtraEnabledText = ':warning: **Not enabled** — add `run-ci-extra` label to opt in.'; const stalePushText = ':warning: **Not run on latest push** — push again or use `/rerun-failed-ci` to dispatch.'; - const ptText = isReal(ptRun) ? `[Run #${ptRun.id}](${ptRun.html_url})` : '_Not run yet_'; - const peText = !hasExtra - ? notEnabledText - : (isReal(peRun) ? `[Run #${peRun.id}](${peRun.html_url})` : stalePushText); + const ptText = !hasCI + ? missingCIText + : (isReal(ptRun) ? `[Run #${ptRun.id}](${ptRun.html_url})` : '_Not run yet_'); + const peText = !hasCI + ? peBlockedByCIText + : !hasExtra + ? notExtraEnabledText + : (isReal(peRun) ? `[Run #${peRun.id}](${peRun.html_url})` : stalePushText); const outerStart = ''; const outerEnd = ''; diff --git a/scripts/ci/utils/slash_command_handler.py b/scripts/ci/utils/slash_command_handler.py index 9beb2f8424af..ad58d78ec7c5 100644 --- a/scripts/ci/utils/slash_command_handler.py +++ b/scripts/ci/utils/slash_command_handler.py @@ -1098,10 +1098,6 @@ def main(): print("Combined command finished, but no actions were taken.") elif first_line.startswith("/rerun-stage"): - # /rerun-stage is deprecated. Stage-level granularity is too coarse to map to - # a specific feature, and a stage rerun re-pays the cost of all unrelated tests - # in that stage. Use /rerun-test for selective UT runs, or /rerun-failed-ci / - # `run-ci` / `run-ci-extra` labels for a full rerun. print("/rerun-stage is deprecated; posting deprecation notice.") comment.create_reaction("-1") pr.create_issue_comment( From b2e9661776dfcf9fe1cb8895e80815d3527d1ae4 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Fri, 15 May 2026 04:19:13 -0700 Subject: [PATCH 11/50] move runs_on + rdma into runner_configs.yml (#25264) --- .github/workflows/_pr-test-check-changes.yml | 13 +++-- .github/workflows/_pr-test-stage.yml | 19 ++++--- .github/workflows/pr-test-extra.yml | 7 --- .github/workflows/pr-test-jit-kernel.yml | 7 ++- .github/workflows/pr-test-multimodal-gen.yml | 7 ++- .github/workflows/pr-test-sgl-kernel.yml | 7 ++- .github/workflows/pr-test.yml | 21 +++----- scripts/ci/runner_configs.py | 53 ++++++++++++++++---- scripts/ci/runner_configs.yml | 35 +++++++------ 9 files changed, 104 insertions(+), 65 deletions(-) diff --git a/.github/workflows/_pr-test-check-changes.yml b/.github/workflows/_pr-test-check-changes.yml index 61cccb5991dc..366a55286c7e 100644 --- a/.github/workflows/_pr-test-check-changes.yml +++ b/.github/workflows/_pr-test-check-changes.yml @@ -29,8 +29,8 @@ on: value: ${{ jobs.run.outputs.partitions }} partition_model_sha: value: ${{ jobs.run.outputs.partition_model_sha }} - b200_runner: - value: ${{ jobs.run.outputs.b200_runner }} + runs_on_map: + value: ${{ jobs.run.outputs.runs_on_map }} enable_retry: value: ${{ jobs.run.outputs.enable_retry }} continue_on_error: @@ -47,7 +47,7 @@ jobs: multimodal_gen: ${{ steps.filter.outputs.multimodal_gen || steps.run-mode.outputs.run_all_tests }} partitions: ${{ steps.partitions.outputs.partitions }} partition_model_sha: ${{ steps.partition-model-sha.outputs.sha }} - b200_runner: ${{ steps.set-runner.outputs.b200_runner }} + runs_on_map: ${{ steps.runner-map.outputs.runs_on_map }} enable_retry: ${{ steps.set-retry.outputs.enable_retry }} continue_on_error: ${{ steps.set-continue-on-error.outputs.continue_on_error }} steps: @@ -174,6 +174,13 @@ jobs: echo "b200_runner=4-gpu-b200" >> $GITHUB_OUTPUT fi + - name: Build runs_on_map (resolves $b200_runner sentinel) + id: runner-map + run: | + python3 scripts/ci/runner_configs.py --map \ + '${{ steps.set-runner.outputs.b200_runner }}' \ + >> "$GITHUB_OUTPUT" + - name: Enable retry for CI id: set-retry run: | diff --git a/.github/workflows/_pr-test-stage.yml b/.github/workflows/_pr-test-stage.yml index 2c6e248295b8..90f60665804a 100644 --- a/.github/workflows/_pr-test-stage.yml +++ b/.github/workflows/_pr-test-stage.yml @@ -15,15 +15,11 @@ on: type: string required: true runner_config: - description: 'Key in scripts/ci/runner_configs.yml (install script / artifact version / install timeout).' - type: string - required: true - runs_on: - description: 'GHA runner label. B200 stages pass needs.check-changes.outputs.b200_runner for dynamic selection.' + description: 'Key in scripts/ci/runner_configs.yml. Resolves install script, artifact version, install timeout, runs-on label, and rdma_devices.' type: string required: true check_changes: - description: 'toJson(needs.check-changes.outputs). Read via fromJson(...).main_package / sgl_kernel / continue_on_error etc.' + description: 'toJson(needs.check-changes.outputs). Read via fromJson(...).main_package / sgl_kernel / continue_on_error / runs_on_map etc.' type: string required: true caller_inputs: @@ -79,11 +75,11 @@ jobs: always() && ((github.event_name == 'schedule' || fromJson(inputs.caller_inputs).test_parallel_dispatch == true) || (!failure() && !cancelled())) && (fromJson(inputs.check_changes).main_package == 'true' || fromJson(inputs.check_changes).sgl_kernel == 'true') - runs-on: ${{ inputs.runs_on }} + # runs-on resolved from runs_on_map; check-changes already substituted + # $b200_runner (see runner_configs.py --map). rdma_devices is exported + # below in a setup step via $GITHUB_ENV. + runs-on: ${{ fromJson(fromJson(inputs.check_changes).runs_on_map)[inputs.runner_config] }} timeout-minutes: 240 - env: - # Only stage-c-test-8-gpu-h20 needs the RDMA device list. - SGLANG_CI_RDMA_ALL_DEVICES: ${{ inputs.runner_config == '8-gpu-h20' && 'mlx5_1,mlx5_2,mlx5_3,mlx5_4' || '' }} strategy: fail-fast: false max-parallel: ${{ fromJson(inputs.partitions)[inputs.self_name].max_parallel }} @@ -98,6 +94,9 @@ jobs: id: rc run: python3 scripts/ci/runner_configs.py '${{ inputs.runner_config }}' >> "$GITHUB_OUTPUT" + - name: Export rdma_devices to job env + run: echo "SGLANG_CI_RDMA_ALL_DEVICES=${{ steps.rc.outputs.rdma_devices || '' }}" >> "$GITHUB_ENV" + - uses: ./.github/actions/check-stage-health - uses: ./.github/actions/check-maintenance diff --git a/.github/workflows/pr-test-extra.yml b/.github/workflows/pr-test-extra.yml index 4039c7e6cad6..12212f8945fc 100644 --- a/.github/workflows/pr-test-extra.yml +++ b/.github/workflows/pr-test-extra.yml @@ -107,7 +107,6 @@ jobs: with: self_name: extra-a-test-1-gpu-small runner_config: 1-gpu-small - runs_on: 1-gpu-5090 check_changes: ${{ toJson(needs.check-changes.outputs) }} caller_inputs: ${{ toJson(inputs) }} partitions: ${{ needs.check-changes.outputs.partitions }} @@ -121,7 +120,6 @@ jobs: with: self_name: extra-a-test-1-gpu-large runner_config: 1-gpu-large - runs_on: 1-gpu-h100 check_changes: ${{ toJson(needs.check-changes.outputs) }} caller_inputs: ${{ toJson(inputs) }} partitions: ${{ needs.check-changes.outputs.partitions }} @@ -136,7 +134,6 @@ jobs: with: self_name: extra-a-test-2-gpu-large runner_config: 2-gpu-large - runs_on: 2-gpu-h100 check_changes: ${{ toJson(needs.check-changes.outputs) }} caller_inputs: ${{ toJson(inputs) }} partitions: ${{ needs.check-changes.outputs.partitions }} @@ -151,7 +148,6 @@ jobs: with: self_name: extra-b-test-4-gpu-h100 runner_config: 4-gpu-h100 - runs_on: 4-gpu-h100 check_changes: ${{ toJson(needs.check-changes.outputs) }} caller_inputs: ${{ toJson(inputs) }} partitions: ${{ needs.check-changes.outputs.partitions }} @@ -165,7 +161,6 @@ jobs: with: self_name: extra-b-test-4-gpu-b200 runner_config: 4-gpu-b200 - runs_on: ${{ needs.check-changes.outputs.b200_runner }} check_changes: ${{ toJson(needs.check-changes.outputs) }} caller_inputs: ${{ toJson(inputs) }} partitions: ${{ needs.check-changes.outputs.partitions }} @@ -180,7 +175,6 @@ jobs: with: self_name: extra-b-test-8-gpu-h200 runner_config: 8-gpu-h200 - runs_on: 8-gpu-h200 check_changes: ${{ toJson(needs.check-changes.outputs) }} caller_inputs: ${{ toJson(inputs) }} partitions: ${{ needs.check-changes.outputs.partitions }} @@ -194,7 +188,6 @@ jobs: with: self_name: extra-b-test-deepep-8-gpu-h200 runner_config: deepep-8-gpu-h200 - runs_on: 8-gpu-h200-deepep check_changes: ${{ toJson(needs.check-changes.outputs) }} caller_inputs: ${{ toJson(inputs) }} partitions: ${{ needs.check-changes.outputs.partitions }} diff --git a/.github/workflows/pr-test-jit-kernel.yml b/.github/workflows/pr-test-jit-kernel.yml index 74c0b90e46a7..9e3e597aa498 100644 --- a/.github/workflows/pr-test-jit-kernel.yml +++ b/.github/workflows/pr-test-jit-kernel.yml @@ -9,7 +9,10 @@ on: sgl_kernel: required: true type: string - b200_runner: + runner_config: + required: true + type: string + runs_on_map: required: true type: string git_ref: @@ -157,7 +160,7 @@ jobs: if: | github.event_name != 'schedule' && inputs.test_parallel_dispatch != 'true' - runs-on: ${{ inputs.b200_runner }} + runs-on: ${{ fromJson(inputs.runs_on_map)[inputs.runner_config] }} timeout-minutes: 240 steps: - uses: actions/checkout@v4 diff --git a/.github/workflows/pr-test-multimodal-gen.yml b/.github/workflows/pr-test-multimodal-gen.yml index 18cf49387284..4ebdb2cd4f81 100644 --- a/.github/workflows/pr-test-multimodal-gen.yml +++ b/.github/workflows/pr-test-multimodal-gen.yml @@ -9,7 +9,10 @@ on: sgl_kernel: required: true type: string - b200_runner: + runner_config: + required: true + type: string + runs_on_map: required: true type: string continue_on_error: @@ -264,7 +267,7 @@ jobs: if: | ((github.event_name == 'schedule' || inputs.test_parallel_dispatch == 'true') || (inputs.caller_needs_failure != 'true' && !cancelled())) && inputs.multimodal_gen == 'true' - runs-on: ${{ inputs.b200_runner }} + runs-on: ${{ fromJson(inputs.runs_on_map)[inputs.runner_config] }} timeout-minutes: 240 steps: - name: Checkout code diff --git a/.github/workflows/pr-test-sgl-kernel.yml b/.github/workflows/pr-test-sgl-kernel.yml index 54cfd3736221..af1552954910 100644 --- a/.github/workflows/pr-test-sgl-kernel.yml +++ b/.github/workflows/pr-test-sgl-kernel.yml @@ -6,7 +6,10 @@ on: sgl_kernel: required: true type: string - b200_runner: + runner_config: + required: true + type: string + runs_on_map: required: true type: string git_ref: @@ -110,7 +113,7 @@ jobs: echo "All benchmark tests completed!" sgl-kernel-b200-test: - runs-on: ${{ inputs.b200_runner }} + runs-on: ${{ fromJson(inputs.runs_on_map)[inputs.runner_config] }} timeout-minutes: 240 steps: - uses: actions/checkout@v4 diff --git a/.github/workflows/pr-test.yml b/.github/workflows/pr-test.yml index a11db0c3a0c1..d5f4d824cb68 100644 --- a/.github/workflows/pr-test.yml +++ b/.github/workflows/pr-test.yml @@ -209,8 +209,9 @@ jobs: needs.check-changes.outputs.sgl_kernel == 'true' uses: ./.github/workflows/pr-test-sgl-kernel.yml with: + runner_config: 4-gpu-b200 + runs_on_map: ${{ needs.check-changes.outputs.runs_on_map }} sgl_kernel: ${{ needs.check-changes.outputs.sgl_kernel }} - b200_runner: ${{ needs.check-changes.outputs.b200_runner }} git_ref: ${{ inputs.git_ref || '' }} skip_stage_health_check: ${{ inputs.skip_stage_health_check == true }} secrets: inherit @@ -227,9 +228,10 @@ jobs: needs.check-changes.outputs.jit_kernel == 'true' uses: ./.github/workflows/pr-test-jit-kernel.yml with: + runner_config: 4-gpu-b200 + runs_on_map: ${{ needs.check-changes.outputs.runs_on_map }} jit_kernel: ${{ needs.check-changes.outputs.jit_kernel }} sgl_kernel: ${{ needs.check-changes.outputs.sgl_kernel }} - b200_runner: ${{ needs.check-changes.outputs.b200_runner }} git_ref: ${{ inputs.git_ref || '' }} test_parallel_dispatch: ${{ inputs.test_parallel_dispatch == true && 'true' || 'false' }} skip_stage_health_check: ${{ inputs.skip_stage_health_check == true }} @@ -245,7 +247,6 @@ jobs: with: self_name: stage-a-test-1-gpu-small runner_config: 1-gpu-small - runs_on: 1-gpu-5090 check_changes: ${{ toJson(needs.check-changes.outputs) }} caller_inputs: ${{ toJson(inputs) }} partitions: ${{ needs.check-changes.outputs.partitions }} @@ -325,7 +326,6 @@ jobs: with: self_name: stage-b-test-1-gpu-small runner_config: 1-gpu-small - runs_on: 1-gpu-5090 check_changes: ${{ toJson(needs.check-changes.outputs) }} caller_inputs: ${{ toJson(inputs) }} partitions: ${{ needs.check-changes.outputs.partitions }} @@ -340,7 +340,6 @@ jobs: with: self_name: stage-b-test-1-gpu-large runner_config: 1-gpu-large - runs_on: 1-gpu-h100 check_changes: ${{ toJson(needs.check-changes.outputs) }} caller_inputs: ${{ toJson(inputs) }} partitions: ${{ needs.check-changes.outputs.partitions }} @@ -355,7 +354,6 @@ jobs: with: self_name: stage-b-test-2-gpu-large runner_config: 2-gpu-large - runs_on: 2-gpu-h100 check_changes: ${{ toJson(needs.check-changes.outputs) }} caller_inputs: ${{ toJson(inputs) }} partitions: ${{ needs.check-changes.outputs.partitions }} @@ -369,7 +367,6 @@ jobs: with: self_name: stage-b-test-4-gpu-b200 runner_config: 4-gpu-b200 - runs_on: ${{ needs.check-changes.outputs.b200_runner }} check_changes: ${{ toJson(needs.check-changes.outputs) }} caller_inputs: ${{ toJson(inputs) }} partitions: ${{ needs.check-changes.outputs.partitions }} @@ -386,9 +383,10 @@ jobs: needs.check-changes.outputs.multimodal_gen == 'true' uses: ./.github/workflows/pr-test-multimodal-gen.yml with: + runner_config: 4-gpu-b200 + runs_on_map: ${{ needs.check-changes.outputs.runs_on_map }} multimodal_gen: ${{ needs.check-changes.outputs.multimodal_gen }} sgl_kernel: ${{ needs.check-changes.outputs.sgl_kernel }} - b200_runner: ${{ needs.check-changes.outputs.b200_runner }} continue_on_error: ${{ needs.check-changes.outputs.continue_on_error }} git_ref: ${{ inputs.git_ref || '' }} test_parallel_dispatch: ${{ inputs.test_parallel_dispatch == true && 'true' || 'false' }} @@ -403,7 +401,6 @@ jobs: with: self_name: stage-c-test-4-gpu-h100 runner_config: 4-gpu-h100 - runs_on: 4-gpu-h100 check_changes: ${{ toJson(needs.check-changes.outputs) }} caller_inputs: ${{ toJson(inputs) }} partitions: ${{ needs.check-changes.outputs.partitions }} @@ -417,7 +414,6 @@ jobs: with: self_name: stage-c-test-8-gpu-h200 runner_config: 8-gpu-h200 - runs_on: 8-gpu-h200 check_changes: ${{ toJson(needs.check-changes.outputs) }} caller_inputs: ${{ toJson(inputs) }} partitions: ${{ needs.check-changes.outputs.partitions }} @@ -438,7 +434,6 @@ jobs: with: self_name: stage-c-test-8-gpu-h20 runner_config: 8-gpu-h20 - runs_on: 8-gpu-h20 check_changes: ${{ toJson(needs.check-changes.outputs) }} caller_inputs: ${{ toJson(inputs) }} partitions: ${{ needs.check-changes.outputs.partitions }} @@ -452,7 +447,6 @@ jobs: with: self_name: stage-c-test-deepep-4-gpu-h100 runner_config: deepep-4-gpu-h100 - runs_on: 4-gpu-h100 check_changes: ${{ toJson(needs.check-changes.outputs) }} caller_inputs: ${{ toJson(inputs) }} partitions: ${{ needs.check-changes.outputs.partitions }} @@ -468,7 +462,6 @@ jobs: with: self_name: stage-c-test-4-gpu-b200 runner_config: 4-gpu-b200 - runs_on: ${{ needs.check-changes.outputs.b200_runner }} check_changes: ${{ toJson(needs.check-changes.outputs) }} caller_inputs: ${{ toJson(inputs) }} partitions: ${{ needs.check-changes.outputs.partitions }} @@ -483,7 +476,6 @@ jobs: with: self_name: stage-c-test-dsv4-4-gpu-b200 runner_config: dsv4-4-gpu-b200 - runs_on: ${{ needs.check-changes.outputs.b200_runner }} check_changes: ${{ toJson(needs.check-changes.outputs) }} caller_inputs: ${{ toJson(inputs) }} partitions: ${{ needs.check-changes.outputs.partitions }} @@ -498,7 +490,6 @@ jobs: with: self_name: stage-c-test-dsv4-8-gpu-h200 runner_config: dsv4-8-gpu-h200 - runs_on: 8-gpu-h200 check_changes: ${{ toJson(needs.check-changes.outputs) }} caller_inputs: ${{ toJson(inputs) }} partitions: ${{ needs.check-changes.outputs.partitions }} diff --git a/scripts/ci/runner_configs.py b/scripts/ci/runner_configs.py index 227dcf73b685..26d9b115d516 100644 --- a/scripts/ci/runner_configs.py +++ b/scripts/ci/runner_configs.py @@ -1,14 +1,24 @@ -"""Emit a runner_config's setup details (install / artifact_version / -install_timeout) in $GITHUB_OUTPUT format. Reads scripts/ci/runner_configs.yml. -Called by .github/workflows/_pr-test-stage.yml. +"""Emit runner_config setup for GitHub Actions $GITHUB_OUTPUT. + +runner_configs.py + Per-field `key=value` lines (install / artifact_version / + install_timeout / rdma_devices). `runs_on` is intentionally omitted — + it carries the `$b200_runner` sentinel and is resolved via --map. + Called per stage by _pr-test-stage.yml. + +runner_configs.py --map + `runs_on_map={json}` — flat dict {runner_config: runs_on}, with + `$b200_runner` substituted. Called once by _pr-test-check-changes.yml. """ +import json import os import sys import yaml _YAML_PATH = os.path.join(os.path.dirname(__file__), "runner_configs.yml") +_B200_SENTINEL = "$b200_runner" def load() -> dict: @@ -16,12 +26,35 @@ def load() -> dict: return yaml.safe_load(f)["runner_configs"] -if __name__ == "__main__": - if len(sys.argv) != 2: - sys.exit("usage: runner_configs.py ") - rc = sys.argv[1] - config = load().get(rc) - if config is None: +def _emit_single(rc: str) -> None: + # runs_on goes through --map (resolves $b200_runner). Suppress it here so a + # consumer can't accidentally read the raw sentinel value. + cfg = load().get(rc) + if cfg is None: sys.exit(f"unknown runner_config: {rc!r}") - for key, value in config.items(): + for key, value in cfg.items(): + if key == "runs_on": + continue print(f"{key}={value}") + + +def _emit_map(b200_runner: str) -> None: + runs_on = { + name: (b200_runner if cfg.get("runs_on") == _B200_SENTINEL else cfg["runs_on"]) + for name, cfg in load().items() + } + print(f"runs_on_map={json.dumps(runs_on, separators=(',', ':'))}") + + +if __name__ == "__main__": + args = sys.argv[1:] + if len(args) == 1: + _emit_single(args[0]) + elif len(args) == 2 and args[0] == "--map": + _emit_map(args[1]) + else: + sys.exit( + "usage:\n" + " runner_configs.py \n" + " runner_configs.py --map " + ) diff --git a/scripts/ci/runner_configs.yml b/scripts/ci/runner_configs.yml index fde798da3c7e..2cccf432d805 100644 --- a/scripts/ci/runner_configs.yml +++ b/scripts/ci/runner_configs.yml @@ -3,9 +3,16 @@ # scripts/ci/runner_configs.py (CLI wrapper), which is in turn called by # .github/workflows/_pr-test-stage.yml. # -# Each runner_config carries install script, actions/download-artifact major -# version, and install-step wall-clock cap (minutes, enforced via -# `timeout-minutes:` on the install step in _pr-test-stage.yml). +# Each runner_config carries: +# - install: install script path +# - artifact_version: actions/download-artifact major version +# - install_timeout: install-step wall-clock cap (minutes), enforced via +# `timeout-minutes:` on the install step in _pr-test-stage.yml +# - runs_on: GHA runner label for the stage's `runs-on:`. The literal +# `$b200_runner` is substituted at workflow-load time with the dynamic +# b200 runner tag from check-changes (see runner_configs.py --map). +# - rdma_devices (optional): exported as SGLANG_CI_RDMA_ALL_DEVICES env +# to the stage job; absent means unset. _anchors: default_install: &default scripts/ci/cuda/ci_install_dependency.sh @@ -13,14 +20,14 @@ _anchors: dsv4_install: &dsv4 scripts/ci/cuda/ci_install_dsv4_dep.sh runner_configs: - 1-gpu-small: { install: *default, artifact_version: v4, install_timeout: "20" } - 1-gpu-large: { install: *default, artifact_version: v4, install_timeout: "20" } - 2-gpu-large: { install: *default, artifact_version: v4, install_timeout: "20" } - 4-gpu-b200: { install: *default, artifact_version: v6, install_timeout: "20" } - 4-gpu-h100: { install: *default, artifact_version: v4, install_timeout: "20" } - 8-gpu-h200: { install: *default, artifact_version: v4, install_timeout: "20" } - 8-gpu-h20: { install: *deepep, artifact_version: v4, install_timeout: "20" } - deepep-4-gpu-h100: { install: *deepep, artifact_version: v4, install_timeout: "20" } - deepep-8-gpu-h200: { install: *deepep, artifact_version: v4, install_timeout: "20" } - dsv4-4-gpu-b200: { install: *dsv4, artifact_version: v6, install_timeout: "30" } - dsv4-8-gpu-h200: { install: *dsv4, artifact_version: v4, install_timeout: "30" } + 1-gpu-small: { install: *default, artifact_version: v4, install_timeout: "20", runs_on: 1-gpu-5090 } + 1-gpu-large: { install: *default, artifact_version: v4, install_timeout: "20", runs_on: 1-gpu-h100 } + 2-gpu-large: { install: *default, artifact_version: v4, install_timeout: "20", runs_on: 2-gpu-h100 } + 4-gpu-b200: { install: *default, artifact_version: v6, install_timeout: "20", runs_on: $b200_runner } + 4-gpu-h100: { install: *default, artifact_version: v4, install_timeout: "20", runs_on: 4-gpu-h100 } + 8-gpu-h200: { install: *default, artifact_version: v4, install_timeout: "20", runs_on: 8-gpu-h200 } + 8-gpu-h20: { install: *deepep, artifact_version: v4, install_timeout: "20", runs_on: 8-gpu-h20, rdma_devices: "mlx5_1,mlx5_2,mlx5_3,mlx5_4" } + deepep-4-gpu-h100: { install: *deepep, artifact_version: v4, install_timeout: "20", runs_on: 4-gpu-h100 } + deepep-8-gpu-h200: { install: *deepep, artifact_version: v4, install_timeout: "20", runs_on: 8-gpu-h200-deepep } + dsv4-4-gpu-b200: { install: *dsv4, artifact_version: v6, install_timeout: "30", runs_on: $b200_runner } + dsv4-8-gpu-h200: { install: *dsv4, artifact_version: v4, install_timeout: "30", runs_on: 8-gpu-h200 } From 99fc29d256a900156f099606056b8b0c13161774 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Sat, 16 May 2026 00:51:48 -0700 Subject: [PATCH 12/50] [Misc] Update release branch cut script (#25468) --- .github/workflows/release-branch-cut.yml | 32 ++++++++++++++++++------ 1 file changed, 25 insertions(+), 7 deletions(-) diff --git a/.github/workflows/release-branch-cut.yml b/.github/workflows/release-branch-cut.yml index a4ed645d5131..d510b1ec3df4 100644 --- a/.github/workflows/release-branch-cut.yml +++ b/.github/workflows/release-branch-cut.yml @@ -26,6 +26,7 @@ jobs: environment: 'prod' outputs: branch_name: ${{ steps.set_output.outputs.branch_name }} + branch_exists: ${{ steps.check_exists.outputs.branch_exists }} steps: - name: Checkout repository uses: actions/checkout@v4 @@ -77,18 +78,20 @@ jobs: echo "Validated commit SHA: $COMMIT_SHA" - name: Check if branch already exists + id: check_exists run: | BRANCH_NAME="${{ github.event.inputs.branch_name }}" if git ls-remote --heads origin "$BRANCH_NAME" | grep -q "$BRANCH_NAME"; then - echo "::error::Branch '$BRANCH_NAME' already exists" - exit 1 + echo "Branch '$BRANCH_NAME' already exists, skipping creation and proceeding to downstream tests" + echo "branch_exists=true" >> $GITHUB_OUTPUT + else + echo "Branch '$BRANCH_NAME' does not exist, proceeding with creation" + echo "branch_exists=false" >> $GITHUB_OUTPUT fi - echo "Branch '$BRANCH_NAME' does not exist, proceeding with creation" - - name: Create release branch - id: set_output + if: steps.check_exists.outputs.branch_exists != 'true' run: | COMMIT_SHA="${{ steps.validate.outputs.COMMIT_SHA }}" BRANCH_NAME="${{ github.event.inputs.branch_name }}" @@ -99,10 +102,10 @@ jobs: # Create branch from the specified commit git checkout -b "$BRANCH_NAME" "$COMMIT_SHA" - echo "branch_name=$BRANCH_NAME" >> $GITHUB_OUTPUT echo "Successfully created branch '$BRANCH_NAME' from commit '$COMMIT_SHA'" - name: Update version references in documentation + if: steps.check_exists.outputs.branch_exists != 'true' run: | BRANCH_NAME="${{ github.event.inputs.branch_name }}" # Extract version from branch name (e.g., release/v0.5.8 -> v0.5.8) @@ -122,15 +125,29 @@ jobs: fi - name: Push release branch + if: steps.check_exists.outputs.branch_exists != 'true' run: | - BRANCH_NAME="${{ steps.set_output.outputs.branch_name }}" + BRANCH_NAME="${{ github.event.inputs.branch_name }}" git push origin "$BRANCH_NAME" echo "Successfully pushed branch '$BRANCH_NAME'" + - name: Emit branch_name output + id: set_output + run: | + BRANCH_NAME="${{ github.event.inputs.branch_name }}" + echo "branch_name=$BRANCH_NAME" >> $GITHUB_OUTPUT + - name: Summary run: | COMMIT_SHA="${{ steps.validate.outputs.COMMIT_SHA }}" BRANCH_NAME="${{ github.event.inputs.branch_name }}" + BRANCH_EXISTS="${{ steps.check_exists.outputs.branch_exists }}" + + if [ "$BRANCH_EXISTS" = "true" ]; then + STATUS="Branch already existed — creation skipped, downstream tests will run" + else + STATUS="Newly created" + fi echo "## Release Branch Cut Summary" >> $GITHUB_STEP_SUMMARY echo "" >> $GITHUB_STEP_SUMMARY @@ -138,6 +155,7 @@ jobs: echo "|----------|-------|" >> $GITHUB_STEP_SUMMARY echo "| Branch | \`$BRANCH_NAME\` |" >> $GITHUB_STEP_SUMMARY echo "| Commit | \`$COMMIT_SHA\` |" >> $GITHUB_STEP_SUMMARY + echo "| Status | $STATUS |" >> $GITHUB_STEP_SUMMARY echo "| Triggered by | @${{ github.actor }} |" >> $GITHUB_STEP_SUMMARY echo "" >> $GITHUB_STEP_SUMMARY echo "### Next Steps" >> $GITHUB_STEP_SUMMARY From cd09fb97f8dc2a1441423d69088842c23e659a27 Mon Sep 17 00:00:00 2001 From: Yuhao Yang Date: Sat, 16 May 2026 22:14:11 +0800 Subject: [PATCH 13/50] Fix HiCache crash on sparse layer_mapping with None entries layer_mapping was changed from a compact list to a sparse array containing None for out-of-stage layers by PP support (#24704), but the HiCache consumer was not updated to handle None entries. --- .../sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py index 3de18edde6d7..217982ebde28 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py @@ -294,6 +294,8 @@ def build_deepseek_v4_hicache_stack( c4_state_global_layers = [] c128_state_global_layers = [] for layer_id, layer_item in enumerate(kvcache.layer_mapping): + if layer_item is None: + continue if layer_item.compress_ratio == 4: c4_layer_mapping[layer_id] = layer_item.compress_layer_id c4_state_global_layers.append(layer_id) From a4e4c45f4e1dce66db8f193bf417765983037554 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Sat, 16 May 2026 10:03:54 -0700 Subject: [PATCH 14/50] Revert "Fix HiCache crash on sparse layer_mapping with None entries" This reverts commit cd09fb97f8dc2a1441423d69088842c23e659a27. --- .../sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py index 217982ebde28..3de18edde6d7 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py @@ -294,8 +294,6 @@ def build_deepseek_v4_hicache_stack( c4_state_global_layers = [] c128_state_global_layers = [] for layer_id, layer_item in enumerate(kvcache.layer_mapping): - if layer_item is None: - continue if layer_item.compress_ratio == 4: c4_layer_mapping[layer_id] = layer_item.compress_layer_id c4_state_global_layers.append(layer_id) From ce43e86265855af95ac3942df3ebf3d8836efb8a Mon Sep 17 00:00:00 2001 From: ybyang <10629930+whybeyoung@users.noreply.github.com> Date: Sat, 16 May 2026 19:49:26 +0800 Subject: [PATCH 15/50] fix(pd): fix kv pools without end_layer (#25476) --- python/sglang/srt/disaggregation/base/conn.py | 2 +- python/sglang/srt/disaggregation/prefill.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/disaggregation/base/conn.py b/python/sglang/srt/disaggregation/base/conn.py index 1e1b4b4f50d6..097a8415914c 100644 --- a/python/sglang/srt/disaggregation/base/conn.py +++ b/python/sglang/srt/disaggregation/base/conn.py @@ -56,7 +56,7 @@ class KVArgs: # reconstruct PP sub-ranges when kv_data_ptrs does not use a flat # layer-indexed layout (e.g. DeepSeek V4's buffer-type-organized flat # list). - prefill_end_layer: int + prefill_end_layer: Optional[int] # For DeepSeek V4 (and other compressed-MLA) memory pools only. # Full-model compression ratio per layer (entries are 0/4/128). Used by # the connection layer to slice the buffer-type-organized flat list in a diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 867824c50ede..0e2ed6a1904e 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -147,7 +147,7 @@ def _init_kv_manager(self) -> CommonKVManager: kv_args.pp_rank = self.pp_rank kv_args.system_dp_rank = self.scheduler.dp_rank kv_args.prefill_start_layer = self.token_to_kv_pool.start_layer - kv_args.prefill_end_layer = self.token_to_kv_pool.end_layer + kv_args.prefill_end_layer = getattr(self.token_to_kv_pool, "end_layer", None) kv_args.mla_compression_ratios = None kv_data_ptrs, kv_data_lens, kv_item_lens = ( self.token_to_kv_pool.get_contiguous_buf_infos() From 127b9e3283f7c2a43234b852ff5c9f1796d53624 Mon Sep 17 00:00:00 2001 From: Zhangheng Date: Sat, 16 May 2026 23:50:01 +0800 Subject: [PATCH 16/50] [BugFix]: Fix DeepSeek V4 HiCache layer count logic (#25477) --- .../hybrid_cache/hybrid_pool_assembler.py | 9 +- .../test_unified_radix_cache_kl_hicache.py | 155 ++++++++++++++++++ ...unified_radix_cache_kl_hicache_nightly.py} | 141 ---------------- 3 files changed, 161 insertions(+), 144 deletions(-) create mode 100644 test/registered/radix_cache/test_unified_radix_cache_kl_hicache.py rename test/registered/radix_cache/{test_unified_radix_hicache_kl.py => test_unified_radix_cache_kl_hicache_nightly.py} (51%) diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py index 3de18edde6d7..1d4f48b13280 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py @@ -283,7 +283,8 @@ def build_deepseek_v4_hicache_stack( pp_size: int = 1, enable_storage_metrics: bool = False, ) -> tuple[HostPoolGroup, HybridCacheController]: - transfer_layer_num = len(kvcache.compression_ratios) + # TODO(hzh0425): Support PP for deepseek v4 with hicache + transfer_layer_num = kvcache.end_layer - kvcache.start_layer full_layer_mapping = {layer_id: layer_id for layer_id in range(transfer_layer_num)} swa_layer_mapping = { layer_id: layer_id for layer_id in range(len(kvcache.swa_kv_pool.kv_buffer)) @@ -293,7 +294,9 @@ def build_deepseek_v4_hicache_stack( c128_layer_mapping = {} c4_state_global_layers = [] c128_state_global_layers = [] - for layer_id, layer_item in enumerate(kvcache.layer_mapping): + for layer_id, layer_item in enumerate( + kvcache.layer_mapping[kvcache.start_layer : kvcache.end_layer] + ): if layer_item.compress_ratio == 4: c4_layer_mapping[layer_id] = layer_item.compress_layer_id c4_state_global_layers.append(layer_id) @@ -730,7 +733,7 @@ def attach_hybrid_pool_to_unified_cache( indices_from_pool=indices_from_pool, ) ) - transfer_layer_num = len(kvcache.compression_ratios) + transfer_layer_num = kvcache.end_layer - kvcache.start_layer elif mamba_stack: full_layer_mapping = dict(kvcache.full_attention_layer_id_mapping) mamba_layer_mapping = dict(params.req_to_token_pool.mamba_map) diff --git a/test/registered/radix_cache/test_unified_radix_cache_kl_hicache.py b/test/registered/radix_cache/test_unified_radix_cache_kl_hicache.py new file mode 100644 index 000000000000..2f3fc32ba071 --- /dev/null +++ b/test/registered/radix_cache/test_unified_radix_cache_kl_hicache.py @@ -0,0 +1,155 @@ +import unittest + +from test_unified_radix_cache_kl import UnifiedRadixTreeTestMixin + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kl_multiturn_utils import ( + get_input_ids, + make_mamba_decode_assert, + make_mamba_prefill_assert, +) +from sglang.test.test_utils import ( + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + CustomTestCase, + is_in_ci, + popen_launch_server, +) + +MAMBA_MODEL = "Qwen/Qwen3-Next-80B-A3B-Instruct" +MAMBA_CHUNK_SIZE = 64 +MAMBA_TRACK_INTERVAL = 128 + +DSV4_FLASH_MODEL = "sgl-project/DeepSeek-V4-Flash-FP8" +DSV4_FLASH_LAUNCH_TIMEOUT = 3600 + +register_cuda_ci(est_time=768, stage="base-c", runner_config="8-gpu-h200") + + +class TestUnifiedMambaHiCache(UnifiedRadixTreeTestMixin, CustomTestCase): + """Mamba hybrid + HiCache + UnifiedRadixCache.""" + + kl_threshold = 0.005 + prefill_cache_assert = staticmethod( + make_mamba_prefill_assert(chunk_size=MAMBA_CHUNK_SIZE) + ) + decode_cache_assert = staticmethod( + make_mamba_decode_assert(track_interval=MAMBA_TRACK_INTERVAL) + ) + + @classmethod + def setUpClass(cls): + cls.model = MAMBA_MODEL + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=[ + "--tp-size", + "4", + "--chunked-prefill-size", + "2048", + "--mem-fraction-static", + "0.85", + "--mamba-scheduler-strategy", + "extra_buffer", + "--mamba-track-interval", + str(MAMBA_TRACK_INTERVAL), + "--enable-hierarchical-cache", + "--hicache-ratio", + "4", + "--hicache-write-policy", + "write_through", + "--hicache-io-backend", + "direct", + "--hicache-mem-layout", + "page_first_direct", + "--max-total-tokens", + "12000", + "--max-mamba-cache-size", + "500", + "--max-running-requests", + "4", + ], + env={"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1"}, + ) + cls.input_ids = get_input_ids(cls.model, num_samples=18) + + @classmethod + def tearDownClass(cls): + kill_process_tree(cls.process.pid) + + +def _assert_dsv4_decode_cached_tokens(result, history_len, output_len, label): + expected = history_len + output_len + actual = result["meta_info"]["cached_tokens"] + lower = max(0, expected - 256) + assert actual >= lower, f"{label}: expected cached_tokens>={lower}, got {actual}" + + +class TestUnifiedDeepSeekV4FlashHiCache(UnifiedRadixTreeTestMixin, CustomTestCase): + """DeepSeek V4 Flash FP8 + HiCache + UnifiedRadixCache.""" + + kl_threshold = 0.005 + sampling_temperature = 0 + decode_cache_assert = staticmethod(_assert_dsv4_decode_cached_tokens) + gsm8k_threshold = 0.90 + num_gsm8k_questions = 100 + + @unittest.skipIf(is_in_ci(), "To reduce the CI execution time.") + def test_multiturn_logprobs_match(self): + pass + + @classmethod + def setUpClass(cls): + cls.model = DSV4_FLASH_MODEL + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DSV4_FLASH_LAUNCH_TIMEOUT, + other_args=[ + "--trust-remote-code", + "--tp-size", + "4", + "--attention-backend", + "compressed", + "--page-size", + "256", + "--chunked-prefill-size", + "8192", + "--mem-fraction-static", + "0.9", + "--disable-shared-experts-fusion", + "--enable-hierarchical-cache", + "--hicache-ratio", + "4", + "--hicache-write-policy", + "write_through", + "--hicache-io-backend", + "direct", + "--hicache-mem-layout", + "page_first_direct", + "--swa-full-tokens-ratio", + "0.25", + "--max-total-tokens", + "20000", + "--max-running-requests", + "2", + ], + env={ + "SGLANG_DSV4_FP4_EXPERTS": "0", + "SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1", + }, + ) + cls.input_ids = get_input_ids(cls.model, num_samples=18) + + @classmethod + def tearDownClass(cls): + kill_process_tree(cls.process.pid) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/radix_cache/test_unified_radix_hicache_kl.py b/test/registered/radix_cache/test_unified_radix_cache_kl_hicache_nightly.py similarity index 51% rename from test/registered/radix_cache/test_unified_radix_hicache_kl.py rename to test/registered/radix_cache/test_unified_radix_cache_kl_hicache_nightly.py index 841af8733fe2..799e94a311f9 100644 --- a/test/registered/radix_cache/test_unified_radix_hicache_kl.py +++ b/test/registered/radix_cache/test_unified_radix_cache_kl_hicache_nightly.py @@ -13,162 +13,21 @@ from urllib.parse import urlparse import requests -from test_unified_radix_cache_kl import UnifiedRadixTreeTestMixin from sglang.srt.utils import kill_process_tree from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.kl_multiturn_utils import ( - get_input_ids, - make_mamba_decode_assert, - make_mamba_prefill_assert, -) from sglang.test.test_utils import ( - DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, DEFAULT_URL_FOR_TEST, CustomTestCase, popen_launch_server, ) -MAMBA_MODEL = "Qwen/Qwen3-Next-80B-A3B-Instruct-FP8" -MAMBA_CHUNK_SIZE = 64 -MAMBA_TRACK_INTERVAL = 128 - -DSV4_FLASH_MODEL = "sgl-project/DeepSeek-V4-Flash-FP8" -DSV4_FLASH_LAUNCH_TIMEOUT = 3600 - -DSV32_MODEL = "deepseek-ai/DeepSeek-V3.2" -DSV32_LAUNCH_TIMEOUT = 3600 - GLM5_MODEL = "zai-org/GLM-5.1-FP8" GLM5_LAUNCH_TIMEOUT = 3600 register_cuda_ci(est_time=900, suite="nightly-8-gpu-h200", nightly=True) -class TestUnifiedMambaHiCache(UnifiedRadixTreeTestMixin, CustomTestCase): - """Mamba hybrid + HiCache + UnifiedRadixCache.""" - - kl_threshold = 0.003 - prefill_cache_assert = staticmethod( - make_mamba_prefill_assert(chunk_size=MAMBA_CHUNK_SIZE) - ) - decode_cache_assert = staticmethod( - make_mamba_decode_assert(track_interval=MAMBA_TRACK_INTERVAL) - ) - - @classmethod - def setUpClass(cls): - cls.model = MAMBA_MODEL - cls.base_url = DEFAULT_URL_FOR_TEST - cls.process = popen_launch_server( - cls.model, - cls.base_url, - timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, - other_args=[ - "--tp-size", - "4", - "--chunked-prefill-size", - "2048", - "--mem-fraction-static", - "0.85", - "--mamba-scheduler-strategy", - "extra_buffer", - "--mamba-track-interval", - str(MAMBA_TRACK_INTERVAL), - "--enable-hierarchical-cache", - "--hicache-ratio", - "4", - "--hicache-write-policy", - "write_through", - "--hicache-io-backend", - "direct", - "--hicache-mem-layout", - "page_first_direct", - "--max-total-tokens", - "12000", - "--max-mamba-cache-size", - "500", - "--max-running-requests", - "4", - ], - env={"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1"}, - ) - cls.input_ids = get_input_ids(cls.model, num_samples=18) - - @classmethod - def tearDownClass(cls): - kill_process_tree(cls.process.pid) - - -def _assert_dsv4_decode_cached_tokens(result, history_len, output_len, label): - expected = history_len + output_len - actual = result["meta_info"]["cached_tokens"] - lower = max(0, expected - 256) - assert actual >= lower, f"{label}: expected cached_tokens>={lower}, got {actual}" - - -class TestUnifiedDeepSeekV4FlashHiCache(UnifiedRadixTreeTestMixin, CustomTestCase): - """DeepSeek V4 Flash FP8 + HiCache + UnifiedRadixCache.""" - - kl_threshold = 0.0035 - sampling_temperature = 0 - decode_cache_assert = staticmethod(_assert_dsv4_decode_cached_tokens) - gsm8k_threshold = 0.90 - num_gsm8k_questions = 100 - - @unittest.skip("no stable.") - def test_multiturn_logprobs_match(self): - pass - - @classmethod - def setUpClass(cls): - cls.model = DSV4_FLASH_MODEL - cls.base_url = DEFAULT_URL_FOR_TEST - cls.process = popen_launch_server( - cls.model, - cls.base_url, - timeout=DSV4_FLASH_LAUNCH_TIMEOUT, - other_args=[ - "--trust-remote-code", - "--tp-size", - "4", - "--attention-backend", - "compressed", - "--page-size", - "256", - "--chunked-prefill-size", - "8192", - "--mem-fraction-static", - "0.9", - "--disable-shared-experts-fusion", - "--enable-hierarchical-cache", - "--hicache-ratio", - "4", - "--hicache-write-policy", - "write_through", - "--hicache-io-backend", - "direct", - "--hicache-mem-layout", - "page_first_direct", - "--swa-full-tokens-ratio", - "0.25", - "--max-total-tokens", - "20000", - "--max-running-requests", - "2", - ], - env={ - "SGLANG_DSV4_FP4_EXPERTS": "0", - "SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1", - }, - ) - cls.input_ids = get_input_ids(cls.model, num_samples=18) - - @classmethod - def tearDownClass(cls): - kill_process_tree(cls.process.pid) - - class GSM8KTwoPassMixin: """Mixin: run GSM8K twice with flush in between, verify accuracy diff. From f5b3fd2bcca0c3378768e5e2757c08f752ea7ffc Mon Sep 17 00:00:00 2001 From: Yueming Yuan Date: Wed, 20 May 2026 19:36:37 -0700 Subject: [PATCH 17/50] [1/14] [sglang-miles] True on-policy training support for FSDP2 (#18639) --- .../sglang/srt/distributed/parallel_state.py | 5 +- python/sglang/srt/layers/layernorm.py | 36 ++-- python/sglang/srt/layers/logits_processor.py | 5 - .../moe/moe_runner/triton_utils/fused_moe.py | 5 +- .../srt/layers/rotary_embedding/base.py | 3 - python/sglang/srt/layers/sampler.py | 14 +- python/sglang/srt/models/qwen2_moe.py | 12 +- python/sglang/srt/models/qwen3.py | 6 +- python/sglang/srt/models/qwen3_moe.py | 47 ++++- python/sglang/srt/models/qwen3_vl.py | 172 ++++++------------ 10 files changed, 143 insertions(+), 162 deletions(-) diff --git a/python/sglang/srt/distributed/parallel_state.py b/python/sglang/srt/distributed/parallel_state.py index 863f5f1a7349..33408f39e592 100644 --- a/python/sglang/srt/distributed/parallel_state.py +++ b/python/sglang/srt/distributed/parallel_state.py @@ -2172,7 +2172,10 @@ def get_tensor_model_parallel_world_size(): def get_tensor_model_parallel_rank(): """Return my rank for the tensor model parallel group.""" - return get_tp_group().rank_in_group + try: + return get_tp_group().rank_in_group + except Exception: + return 0 # ATTN_TP diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index e9c9e7be8c71..1079434581f2 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -144,7 +144,7 @@ def _forward_with_allreduce_fusion( if world_size > 1: if post_residual_addition is not None: - residual = residual + post_residual_addition + x = x + post_residual_addition # Prefer AITER fused AR+RMSNorm when enabled on AMD. if _use_aiter: @@ -179,20 +179,17 @@ def __init__( eps: float = 1e-6, var_hidden_size: Optional[int] = None, cast_x_before_out_mul: bool = False, - fp32_residual: bool = False, + fp32_residual: bool = True, has_weight: bool = True, - weight_dtype: Optional = None, - override_orig_dtype: Optional = None, ) -> None: super().__init__() self.has_weight = has_weight self.cast_x_before_out_mul = cast_x_before_out_mul self.fp32_residual = fp32_residual - self.override_orig_dtype = override_orig_dtype if self.has_weight: - self.weight = nn.Parameter(torch.ones(hidden_size, dtype=weight_dtype)) + self.weight = nn.Parameter(torch.ones(hidden_size)) else: - self.weight = torch.ones(hidden_size, dtype=weight_dtype) + self.weight = torch.ones(hidden_size) self.variance_epsilon = eps self.hidden_size = hidden_size self.variance_size_override = ( @@ -294,10 +291,10 @@ def forward_aiter( elif not x.is_contiguous(): x = x.contiguous() if residual is not None: - residual_out = torch.empty_like(x) - output = torch.empty_like(x) if post_residual_addition is not None: residual = residual + post_residual_addition + residual_out = torch.empty_like(x) + output = torch.empty_like(x) fused_add_rms_norm( output, x, @@ -326,10 +323,10 @@ def forward_hip( # NOTE: Remove this if aiter kernel supports discontinuous input x = x.contiguous() if residual is not None: - out = torch.empty_like(x) - residual_out = torch.empty_like(x) if post_residual_addition is not None: residual = residual + post_residual_addition + out = torch.empty_like(x) + residual_out = torch.empty_like(x) fused_add_rms_norm( out, x, residual_out, residual, self.weight.data, self.variance_epsilon ) @@ -366,16 +363,19 @@ def forward_native( ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: if not x.is_contiguous(): x = x.contiguous() - orig_dtype = self.override_orig_dtype or x.dtype + orig_dtype = x.dtype + + if residual is not None and not self.fp32_residual: + x = x + residual + if post_residual_addition is not None: + x = x + post_residual_addition + residual = x.clone() x = x.to(torch.float32) - if residual is not None: + if residual is not None and self.fp32_residual: x = x + residual.to(torch.float32) if post_residual_addition is not None: x = x + post_residual_addition.to(torch.float32) - if self.fp32_residual: - residual = x.clone() - else: - residual = x.to(orig_dtype) + residual = x.to(orig_dtype) hidden_size = x.shape[-1] if hidden_size != self.hidden_size: @@ -603,7 +603,7 @@ def forward_native( orig_dtype = x.dtype if residual is not None: if post_residual_addition is not None: - residual = residual + post_residual_addition + x = x + post_residual_addition x = x + residual residual = x diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index e14895e979dc..881fe11ee61c 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -906,11 +906,6 @@ def _compute_lm_head( None, # bias True, # is_vnni ) - elif get_global_server_args().rl_on_policy_target is not None: - # Due to tie-weight, we may not be able to change lm_head's weight dtype - logits = torch.matmul( - hidden_states.bfloat16(), lm_head.weight.T.bfloat16() - ) else: logits = torch.matmul( hidden_states.to(lm_head.weight.dtype), lm_head.weight.T diff --git a/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py b/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py index 8062e93c4e9b..e905b08069d4 100644 --- a/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py +++ b/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py @@ -728,7 +728,10 @@ def _fused_moe_kernel_sequence( ).squeeze(dim=1) else: # According to micro benchmark results, torch.compile can get better performance for small token. - if num_tokens <= 32: + if ( + not get_global_server_args().enable_deterministic_inference + and num_tokens <= 32 + ): moe_sum_reduce_torch_compile( intermediate_cache3.view(*intermediate_cache3.shape), out_hidden_states, diff --git a/python/sglang/srt/layers/rotary_embedding/base.py b/python/sglang/srt/layers/rotary_embedding/base.py index 2b13c1594d82..ac5d59d4ac79 100644 --- a/python/sglang/srt/layers/rotary_embedding/base.py +++ b/python/sglang/srt/layers/rotary_embedding/base.py @@ -102,9 +102,6 @@ def __init__( # XXX (MUSA): Implement sgl_kernel.rotary_embedding support for MUSA backend if get_global_server_args().rl_on_policy_target is not None or _is_musa: self._forward_method = self.forward_native - self._apply_rotary_emb_wrapped = torch.compile(dynamic=True)( - apply_rotary_emb - ) self.position_cos, self.position_sin = None, None def _match_cos_sin_cache_dtype(self, query: torch.Tensor) -> None: diff --git a/python/sglang/srt/layers/sampler.py b/python/sglang/srt/layers/sampler.py index 9181fbac5e4e..1660bcb93a7a 100644 --- a/python/sglang/srt/layers/sampler.py +++ b/python/sglang/srt/layers/sampler.py @@ -136,17 +136,13 @@ def forward( if return_logprob and SGLANG_RETURN_ORIGINAL_LOGPROB: original_logprobs = torch.log_softmax(logits, dim=-1) + # Post process logits + logits.div_(sampling_info.temperatures) + # In RL on-policy mode, we use log_softmax to compute logprobs to match the trainer. logprobs_via_logsoftmax_kernel = None if self.rl_on_policy_target is not None: - # TODO: use more inplace ops to save memory - logits_div_temperature = ( - logits.bfloat16().div(sampling_info.temperatures).bfloat16() - ) - logprobs_via_logsoftmax_kernel = torch.log_softmax( - logits_div_temperature, dim=-1 - ) - del logits_div_temperature + logprobs_via_logsoftmax_kernel = torch.log_softmax(logits, dim=-1) if self.use_ascend_backend: # Ascend backend: sample from logits directly. @@ -168,8 +164,6 @@ def forward( logprobs = logprobs_via_logsoftmax_kernel else: # Standard path: do softmax and sample from probs. - logits.div_(sampling_info.temperatures) - # In-place op to save memory logits[:] = torch.softmax(logits, dim=-1) probs = logits diff --git a/python/sglang/srt/models/qwen2_moe.py b/python/sglang/srt/models/qwen2_moe.py index 8bdc0598e12a..415c6044ff60 100644 --- a/python/sglang/srt/models/qwen2_moe.py +++ b/python/sglang/srt/models/qwen2_moe.py @@ -761,7 +761,17 @@ def __init__( prefix=add_prefix("layers", prefix), ) if self.pp_group.is_last_rank: - self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + norm_kwargs = ( + dict( + cast_x_before_out_mul=True, + fp32_residual=False, + ) + if get_global_server_args().rl_on_policy_target is not None + else {} + ) + self.norm = RMSNorm( + config.hidden_size, eps=config.rms_norm_eps, **norm_kwargs + ) else: self.norm = PPMissingLayer(return_tuple=True) diff --git a/python/sglang/srt/models/qwen3.py b/python/sglang/srt/models/qwen3.py index 30333f9998c9..be8f747caf25 100644 --- a/python/sglang/srt/models/qwen3.py +++ b/python/sglang/srt/models/qwen3.py @@ -103,8 +103,8 @@ def __init__( norm_kwargs = ( dict( - weight_dtype=torch.float32, cast_x_before_out_mul=True, + fp32_residual=False, ) if get_global_server_args().rl_on_policy_target is not None else {} @@ -353,10 +353,8 @@ def __init__( norm_kwargs = ( dict( - weight_dtype=torch.float32, cast_x_before_out_mul=True, - override_orig_dtype=torch.float32, - fp32_residual=True, + fp32_residual=False, ) if get_global_server_args().rl_on_policy_target is not None else {} diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py index f255b90fde99..9fb6808678d9 100644 --- a/python/sglang/srt/models/qwen3_moe.py +++ b/python/sglang/srt/models/qwen3_moe.py @@ -22,6 +22,7 @@ from typing import Any, Dict, Iterable, List, Optional, Tuple, TypeVar import torch +import torch.nn.functional as F from torch import nn from transformers import PretrainedConfig @@ -54,7 +55,7 @@ ) from sglang.srt.layers.moe.ep_moe.layer import get_moe_impl_class from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE -from sglang.srt.layers.moe.topk import TopK +from sglang.srt.layers.moe.topk import StandardTopKOutput, TopK from sglang.srt.layers.moe.utils import ( RoutingMethodType, filter_moe_weight_param_global_expert, @@ -328,7 +329,20 @@ def forward_normal( # router_logits: (num_tokens, n_experts) router_logits, _ = self.gate(hidden_states) - topk_output = self.topk(hidden_states, router_logits) + if get_global_server_args().rl_on_policy_target is not None: + routing_weights = F.softmax(router_logits, dim=1, dtype=torch.float) + routing_weights, selected_experts = torch.topk( + routing_weights, self.top_k, dim=-1 + ) + routing_weights /= routing_weights.sum(dim=-1, keepdim=True) + routing_weights = routing_weights.to(hidden_states.dtype) + topk_output = StandardTopKOutput( + topk_weights=routing_weights, + topk_ids=selected_experts, + router_logits=router_logits, + ) + else: + topk_output = self.topk(hidden_states, router_logits) final_hidden_states = self.experts(hidden_states, topk_output) if self.ep_size > 1 and not should_skip_post_experts_all_reduce( @@ -518,7 +532,7 @@ def __init__( ) self.compatible_with_fused_kv_buffer = ( False if isinstance(self.rotary_emb, MRotaryEmbedding) else True - ) + ) and (get_global_server_args().rl_on_policy_target is None) self.compatible_with_fused_qk_norm_rope = not isinstance( self.rotary_emb, MRotaryEmbedding ) and self.head_dim in (64, 128, 256) @@ -533,6 +547,7 @@ def __init__( torch.bfloat16, _yarn_factor != 1.0, ) + and (get_global_server_args().rl_on_policy_target is None) ) self._used_fused_qk_norm_rope_last_call = False @@ -545,8 +560,16 @@ def __init__( prefix=add_prefix("attn", prefix), ) - self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps) - self.k_norm = RMSNorm(self.head_dim, eps=rms_norm_eps) + norm_kwargs = ( + dict( + cast_x_before_out_mul=True, + fp32_residual=False, + ) + if get_global_server_args().rl_on_policy_target is not None + else {} + ) + self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps, **norm_kwargs) + self.k_norm = RMSNorm(self.head_dim, eps=rms_norm_eps, **norm_kwargs) self.alt_stream = alt_stream def op_prepare(self, state): @@ -790,9 +813,19 @@ def __init__( quant_config=quant_config, prefix=add_prefix("mlp", prefix), ) - self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + norm_kwargs = ( + dict( + cast_x_before_out_mul=True, + fp32_residual=False, + ) + if get_global_server_args().rl_on_policy_target is not None + else {} + ) + self.input_layernorm = RMSNorm( + config.hidden_size, eps=config.rms_norm_eps, **norm_kwargs + ) self.post_attention_layernorm = RMSNorm( - config.hidden_size, eps=config.rms_norm_eps + config.hidden_size, eps=config.rms_norm_eps, **norm_kwargs ) self.layer_communicator = LayerCommunicator( diff --git a/python/sglang/srt/models/qwen3_vl.py b/python/sglang/srt/models/qwen3_vl.py index 1b6c185bcbda..9db813fdeaca 100644 --- a/python/sglang/srt/models/qwen3_vl.py +++ b/python/sglang/srt/models/qwen3_vl.py @@ -450,75 +450,73 @@ def rot_pos_emb( return cos_combined, sin_combined - def _get_interpolation_indices(self, dim_size: int) -> torch.Tensor: - """ - Compute continuous interpolation indices for a single dimension. + def fast_pos_embed_interpolate(self, grid_thw): + grid_ts, grid_hs, grid_ws = grid_thw[:, 0], grid_thw[:, 1], grid_thw[:, 2] + num_grid_per_side = int(self.num_position_embeddings**0.5) + device = self.pos_embed.weight.device - Returns continuous indices. - """ - if self.align_corners: - indices = np.linspace( - 0, self.num_grid_per_side - 1, dim_size, dtype=np.float32 - ) - else: - indices = (np.arange(dim_size, dtype=np.float32) + 0.5) * ( - self.num_grid_per_side / dim_size - ) - 0.5 - indices = np.clip(indices, 0, self.num_grid_per_side - 1) - return indices + idx_list = [[] for _ in range(4)] + weight_list = [[] for _ in range(4)] - def _calculate_indices_and_weights(self, h_idxs, w_idxs): - """ - Compute bilinear interpolation indices and weights. + for t, h, w in zip(grid_ts, grid_hs, grid_ws): + h_idxs = torch.linspace(0, num_grid_per_side - 1, h) + w_idxs = torch.linspace(0, num_grid_per_side - 1, w) - Returns tuple of (indices, weights), each as 4 numpy arrays for the 4 corner points. - """ - h_f = np.floor(h_idxs).astype(np.int64) - h_c = np.clip(h_f + 1, 0, self.num_grid_per_side - 1) - dh = h_idxs - h_f + h_idxs_floor = h_idxs.int() + w_idxs_floor = w_idxs.int() + h_idxs_ceil = (h_idxs.int() + 1).clip(max=num_grid_per_side - 1) + w_idxs_ceil = (w_idxs.int() + 1).clip(max=num_grid_per_side - 1) - w_f = np.floor(w_idxs).astype(np.int64) - w_c = np.clip(w_f + 1, 0, self.num_grid_per_side - 1) - dw = w_idxs - w_f + dh = h_idxs - h_idxs_floor + dw = w_idxs - w_idxs_floor - side = self.num_grid_per_side + base_h = h_idxs_floor * num_grid_per_side + base_h_ceil = h_idxs_ceil * num_grid_per_side - indices = [ - (h_f[:, None] * side + w_f).flatten(), - (h_f[:, None] * side + w_c).flatten(), - (h_c[:, None] * side + w_f).flatten(), - (h_c[:, None] * side + w_c).flatten(), - ] - weights = [ - ((1 - dh)[:, None] * (1 - dw)).flatten(), - ((1 - dh)[:, None] * dw).flatten(), - (dh[:, None] * (1 - dw)).flatten(), - (dh[:, None] * dw).flatten(), - ] - return indices, weights + indices = [ + (base_h[None].T + w_idxs_floor[None]).flatten(), + (base_h[None].T + w_idxs_ceil[None]).flatten(), + (base_h_ceil[None].T + w_idxs_floor[None]).flatten(), + (base_h_ceil[None].T + w_idxs_ceil[None]).flatten(), + ] - def _get_position_embedding(self, patch_pos_embeds, grid_ts, grid_hs, grid_ws): - """ - Tile and reorganize position embeddings to align with the token sequence. - """ - result_parts = [] - merge_size = self.spatial_merge_size + weights = [ + ((1 - dh)[None].T * (1 - dw)[None]).flatten(), + ((1 - dh)[None].T * dw[None]).flatten(), + (dh[None].T * (1 - dw)[None]).flatten(), + (dh[None].T * dw[None]).flatten(), + ] - for pos_embed, t, h, w in zip(patch_pos_embeds, grid_ts, grid_hs, grid_ws): - pos_embed = pos_embed.repeat(t, 1) + for i in range(4): + idx_list[i].extend(indices[i].tolist()) + weight_list[i].extend(weights[i].tolist()) - h_merge = h // merge_size - w_merge = w // merge_size + idx_tensor = torch.tensor(idx_list, dtype=torch.long, device=device) + weight_tensor = torch.tensor( + weight_list, dtype=self.pos_embed.weight.dtype, device=device + ) + pos_embeds = self.pos_embed(idx_tensor).to(device) * weight_tensor[:, :, None] + patch_pos_embeds = pos_embeds[0] + pos_embeds[1] + pos_embeds[2] + pos_embeds[3] + + patch_pos_embeds = patch_pos_embeds.split( + [h * w for h, w in zip(grid_hs, grid_ws)] + ) + patch_pos_embeds_permute = [] + merge_size = self.spatial_merge_size + for pos_embed, t, h, w in zip(patch_pos_embeds, grid_ts, grid_hs, grid_ws): + pos_embed = pos_embed.repeat(t, 1) pos_embed = ( - pos_embed.view(t, h_merge, merge_size, w_merge, merge_size, -1) + pos_embed.view( + t, h // merge_size, merge_size, w // merge_size, merge_size, -1 + ) .permute(0, 1, 3, 2, 4, 5) .flatten(0, 4) ) - result_parts.append(pos_embed) + patch_pos_embeds_permute.append(pos_embed) - return torch.cat(result_parts, dim=0) + return torch.cat(patch_pos_embeds_permute) def _torch_interp_indices( self, dim_size: int, device: torch.device @@ -626,61 +624,6 @@ def bucket_flashinfer_max_seqlen(self, real_max_seqlen: int) -> int: round_up(real_max_seqlen, FLASHINFER_MAX_SEQLEN_BUCKETS[-1]), ) - def fast_pos_embed_interpolate(self, grid_thw): - """Interpolate position embeddings for (batch, 3) size input dimensions. - - Performs bilinear interpolation on spatial dimensions (height, width) and replicates - along temporal dimension. The result is reorganized according to spatial_merge_size. - - Args: - grid_thw: Tensor of shape [batch_size, 3] with (temporal, height, width) dimensions - in patches for each sample. - - Returns: - Interpolated position embeddings tensor. - """ - grid_thw_cpu = grid_thw.cpu().numpy() - - # transfer data to CPU before loop - temporal_dims = grid_thw_cpu[:, 0].tolist() - height_dims = grid_thw_cpu[:, 1].tolist() - width_dims = grid_thw_cpu[:, 2].tolist() - - device = self.pos_embed.weight.device - dtype = self.pos_embed.weight.dtype - - patches_size = [h * w for h, w in zip(height_dims, width_dims)] - total_patches = sum(patches_size) - all_indices_np = np.zeros((4, total_patches), dtype=np.int64) - all_weights_np = np.zeros((4, total_patches), dtype=np.float32) - - current_idx = 0 - - # calculate indices and weights on CPU - for t, h, w in zip(temporal_dims, height_dims, width_dims): - h_idxs = self._get_interpolation_indices(h) - w_idxs = self._get_interpolation_indices(w) - - indices, weights = self._calculate_indices_and_weights(h_idxs, w_idxs) - - end_idx = current_idx + h * w - for i in range(4): - all_indices_np[i, current_idx:end_idx] = indices[i] - all_weights_np[i, current_idx:end_idx] = weights[i] - current_idx = end_idx - - idx_tensor = torch.from_numpy(all_indices_np).to(device) - weight_tensor = torch.from_numpy(all_weights_np).to(dtype=dtype, device=device) - - # calculate interpolation - pos_embeds = self.pos_embed(idx_tensor.view(-1)) - pos_embeds = pos_embeds.view(4, total_patches, -1) - patch_pos_embeds = (pos_embeds * weight_tensor.unsqueeze(-1)).sum(dim=0) - patch_pos_embeds = patch_pos_embeds.split(patches_size) - return self._get_position_embedding( - patch_pos_embeds, temporal_dims, height_dims, width_dims - ) - def compute_flashinfer_batch_offsets_packed( self, token_cu_seqlens: np.ndarray, @@ -1024,14 +967,19 @@ def forward( hidden_states + residual if residual is not None else hidden_states ) + deepstack_embeds = None + if input_deepstack_embeds is not None: + prev_layer_idx = layer_idx - 1 + if prev_layer_idx in self.deepstack_embed_to_decoder_layer: + sep = self.hidden_size * prev_layer_idx + deepstack_embeds = input_deepstack_embeds[ + :, sep : sep + self.hidden_size + ] + # SGLang applies residual at the START of the next layer, not at the END like HuggingFace. # See: https://github.com/huggingface/transformers/blob/v5.0.0rc0/src/transformers/models/qwen3_vl/modeling_qwen3_vl.py#L549 # To match HF behavior, deepstack must be added AFTER residual: (hidden_states + residual) + deepstack # The order matters because addition with different tensors is not associative in practice. - # Deepstack for prev_layer is applied at the start of current layer via post_residual_addition. - deepstack_embeds = self.get_deepstack_embeds( - layer_idx - 1, input_deepstack_embeds - ) hidden_states, residual = layer( positions, hidden_states, From 1a83197865bbd69b18e79ae78df6713e429bd8f4 Mon Sep 17 00:00:00 2001 From: Yueming Yuan Date: Wed, 20 May 2026 19:37:55 -0700 Subject: [PATCH 18/50] [2/14] [sglang-miles] R3 (Rollout Routing Replay) DeepEP and MTP support (#18642) --- .../scheduler_output_processor_mixin.py | 2 ++ .../sglang/srt/model_executor/model_runner.py | 20 +++++++++++++++---- 2 files changed, 18 insertions(+), 4 deletions(-) diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py index ae6f732fe934..e02f21163b62 100644 --- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py +++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py @@ -1262,6 +1262,8 @@ def stream_output_generation( # Send to detokenizer if reqs or is_idle_batch: + if self.model_config.is_multimodal_gen: + return self.send_to_detokenizer.send_output( BatchTokenIDOutput( rids=rids, diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index ff937dd9ca4e..c32fbacbcaf0 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -752,7 +752,8 @@ def initialize(self, pre_model_load_memory: float): self.maybe_init_ngram_embedding() # Init routed experts capturer - self.init_routed_experts_capturer() + if not self.is_draft_worker: + self.init_routed_experts_capturer() self.init_indexer_capturer() @@ -3339,11 +3340,22 @@ def forward( output.expert_distribution_metrics = recorder_outputs.get("metrics") no_copy_to_cpu = not self.server_args.disable_overlap_schedule - if (experts_capturer := get_global_experts_capturer()) is not None: + # In speculative decoding, num_tokens_per_bs > 1, so pass the actual + # number of tokens per DP rank in CUDA graph, not the batch size. + cuda_graph_num_tokens = None + if getattr(self.graph_runner, "bs", None): + cuda_graph_num_tokens = ( + self.graph_runner.bs * self.graph_runner.num_tokens_per_bs + ) + + if ( + not self.is_draft_worker + and (experts_capturer := get_global_experts_capturer()) is not None + ): output.routed_experts_output = experts_capturer.on_forward_end( forward_batch=forward_batch, can_run_graph=output.can_run_graph, - cuda_graph_batch=getattr(self.graph_runner, "bs", None), + cuda_graph_batch=cuda_graph_num_tokens, no_copy_to_cpu=no_copy_to_cpu, ) @@ -3351,7 +3363,7 @@ def forward( output.indexer_topk_output = indexer_capturer.on_forward_end( forward_batch=forward_batch, can_run_graph=output.can_run_graph, - cuda_graph_batch=getattr(self.graph_runner, "bs", None), + cuda_graph_batch=cuda_graph_num_tokens, no_copy_to_cpu=no_copy_to_cpu, ) From 5bd6c6bc45de1e3466e2e3901b4ed00ce1f48c91 Mon Sep 17 00:00:00 2001 From: Yueming Yuan Date: Wed, 20 May 2026 19:40:11 -0700 Subject: [PATCH 19/50] [3/14] [sglang-miles] Support INT4 QAT for RL (#18565) --- python/sglang/srt/entrypoints/engine.py | 19 ++++++ python/sglang/srt/entrypoints/http_server.py | 17 +++++ .../srt/layers/moe/fused_moe_triton/layer.py | 1 + .../compressed_tensors/compressed_tensors.py | 6 +- .../schemes/compressed_tensors_wNa16_moe.py | 66 ++++++++++++++++--- python/sglang/srt/managers/io_struct.py | 14 ++++ python/sglang/srt/managers/scheduler.py | 2 + .../scheduler_update_weights_mixin.py | 22 +++++++ .../srt/managers/tokenizer_control_mixin.py | 14 ++++ python/sglang/srt/managers/tp_worker.py | 6 ++ .../sglang/srt/model_executor/model_runner.py | 26 ++++++++ 11 files changed, 184 insertions(+), 9 deletions(-) diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index 1dde8bed80e9..729a5e6fe3c6 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -71,6 +71,7 @@ LoadLoRAAdapterReqInput, MultimodalDataInputFormat, OpenSessionReqInput, + PostProcessWeightsReqInput, ReleaseMemoryOccupationReqInput, ResumeMemoryOccupationReqInput, RpcReqInput, @@ -1090,6 +1091,24 @@ def update_weights_from_ipc( self.tokenizer_manager.update_weights_from_ipc(obj, None) ) + def post_process_weights( + self, + restore_weights_before_load: bool = False, + post_process_quantization: bool = False, + ): + """ + Optional post-processing for updated weights (e.g., Marlin conversion). + Should be called after weight update is finished. + """ + obj = PostProcessWeightsReqInput( + restore_weights_before_load=restore_weights_before_load, + post_process_quantization=post_process_quantization, + ) + + return self.loop.run_until_complete( + self.tokenizer_manager.post_process_weights(obj, None) + ) + def get_weights_by_name(self, name: str, truncate_size: int = 100): """Get weights by parameter name.""" obj = GetWeightsByNameReqInput(name=name, truncate_size=truncate_size) diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index 80081fc5790b..b18af2f30504 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -127,6 +127,7 @@ OpenSessionReqInput, ParseFunctionCallReq, PauseGenerationReqInput, + PostProcessWeightsReqInput, ProfileReqInput, ReleaseMemoryOccupationReqInput, ResumeMemoryOccupationReqInput, @@ -1222,6 +1223,22 @@ async def update_weights_from_ipc(obj: UpdateWeightsFromIPCReqInput, request: Re return ORJSONResponse(content, status_code=HTTPStatus.BAD_REQUEST) +@app.post("/post_process_weights") +async def post_process_weights(req: PostProcessWeightsReqInput, request: Request): + """ + Optional post-processing for updated weights (e.g., Marlin conversion). + This should be called selectively after `update_weights_from_distributed/update_weights_from_tensor`. + """ + success, message = await _global_state.tokenizer_manager.post_process_weights( + req, request + ) + + content = {"success": success, "message": message} + return ORJSONResponse( + content, status_code=200 if success else HTTPStatus.BAD_REQUEST + ) + + @app.post("/update_weight_version") @auth_level(AuthLevel.ADMIN_OPTIONAL) async def update_weight_version(obj: UpdateWeightVersionReqInput, request: Request): diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py index 97a5bfc3d2d4..2e2aae7f532c 100644 --- a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py +++ b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py @@ -771,6 +771,7 @@ def _weight_loader_impl( "CompressedTensorsWNA16TritonMoE", ] ) + and "zero" not in weight_name else loaded_weight ) diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py b/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py index 81056a17e03d..d7a013d19ebc 100644 --- a/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py +++ b/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py @@ -502,7 +502,7 @@ def _is_wNa16_group_channel( ) is_static = not weight_quant.dynamic - return is_channel_group and input_quant_none and is_symmetric and is_static + return is_channel_group and input_quant_none and is_static def _is_mxint4a16(self, weight_quant: BaseModel, input_quant: BaseModel) -> bool: input_quant_none = input_quant is None @@ -978,6 +978,10 @@ def __init__(self, quantization_config: CompressedTensorsConfig): def process_weights_after_loading(self, layer: torch.nn.Module) -> None: layer.scheme.process_weights_after_loading(layer) + def restore_weights_before_loading(self, layer: torch.nn.Module) -> None: + if hasattr(layer.scheme, "restore_weights_before_loading"): + layer.scheme.restore_weights_before_loading(layer) + def create_weights( self, layer: torch.nn.Module, diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py index 0ac18784c316..82f421ac9ebe 100644 --- a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py +++ b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py @@ -17,7 +17,10 @@ CompressedTensorsMoEScheme, ) from sglang.srt.layers.quantization.gptq import gptq_marlin_moe_repack -from sglang.srt.layers.quantization.marlin_utils import marlin_moe_permute_scales +from sglang.srt.layers.quantization.marlin_utils import ( + marlin_moe_permute_scales, + moe_awq_to_marlin_zero_points, +) from sglang.srt.layers.quantization.utils import replace_parameter from sglang.srt.utils import get_bool_env_var, is_cuda, is_hip, set_weight_attrs @@ -64,7 +67,7 @@ def __init__(self, quant_config: CompressedTensorsConfig, num_gpu_experts=-1): self.strategy = config.strategy self.group_size = config.group_size self.actorder = config.actorder - assert config.symmetric, "Only symmetric quantization is supported for MoE" + self.sym = config.symmetric if not ( self.quant_config.quant_format == CompressionFormat.pack_quantized.value @@ -124,7 +127,7 @@ def create_weights( # In the case where we have actorder/g_idx, # we do not partition the w2 scales - load_full_w2 = self.actorder and self.group_size != -1 + load_full_w2 = (self.actorder != "static") and self.group_size != -1 if load_full_w2: w2_scales_size = intermediate_size_per_partition * layer.moe_tp_size @@ -172,6 +175,32 @@ def create_weights( layer.register_parameter("w13_weight_shape", w13_weight_shape) set_weight_attrs(w13_weight_shape, extra_weight_attrs) + # add zero param + if not self.sym: + w13_qzeros = torch.nn.Parameter( + torch.empty( + num_experts, + num_groups_w13, + 2 * intermediate_size_per_partition // self.packed_factor, + dtype=torch.int32, + ), + requires_grad=False, + ) + layer.register_parameter("w13_weight_zero_point", w13_qzeros) + set_weight_attrs(w13_qzeros, extra_weight_attrs) + + w2_qzeros = torch.nn.Parameter( + torch.empty( + num_experts, + num_groups_w2, + hidden_size // self.packed_factor, + dtype=torch.int32, + ), + requires_grad=False, + ) + layer.register_parameter("w2_weight_zero_point", w2_qzeros) + set_weight_attrs(w2_qzeros, extra_weight_attrs) + w13_g_idx = torch.nn.Parameter( torch.empty( num_experts, @@ -225,14 +254,16 @@ def create_weights( # Force record: these are the target GPTQ shapes for rollback. layer._original_shapes["w13_weight_packed"] = tuple(w13_weight.shape) - layer._original_shapes["w2_weight_packed"] = tuple(w2_weight.shape) + layer._original_shapes["w13_weight_scale"] = tuple(w13_scale.shape) + if not self.sym: + layer._original_shapes["w13_weight_zero_point"] = w13_qzeros.shape - # Also record the shapes of the scales. + layer._original_shapes["w2_weight_packed"] = tuple(w2_weight.shape) layer._original_shapes["w2_weight_scale"] = tuple(w2_scale.shape) - layer._original_shapes["w13_weight_scale"] = tuple(w13_scale.shape) + if not self.sym: + layer._original_shapes["w2_weight_zero_point"] = tuple(w2_qzeros.shape) def process_weights_after_loading(self, layer: torch.nn.Module) -> None: - # Skip if the layer is already converted to Marlin format to prevent double-packing. if getattr(layer, "is_marlin_converted", False): return @@ -334,11 +365,28 @@ def replace_tensor(name, new_t): ) replace_tensor("w2_weight_scale", marlin_w2_scales) + # Repack zero + if not self.sym: + marlin_w13_zp = moe_awq_to_marlin_zero_points( + layer.w13_weight_zero_point, + size_k=layer.w13_weight_zero_point.shape[1], + size_n=layer.w13_weight_zero_point.shape[2] * self.packed_factor, + num_bits=self.num_bits, + ) + replace_tensor("w13_weight_zero_point", marlin_w13_zp) + + marlin_w2_zp = moe_awq_to_marlin_zero_points( + layer.w2_weight_zero_point, + size_k=layer.w2_weight_zero_point.shape[1], + size_n=layer.w2_weight_zero_point.shape[2] * self.packed_factor, + num_bits=self.num_bits, + ) + replace_tensor("w2_weight_zero_point", marlin_w2_zp) + layer.is_marlin_converted = True def restore_weights_before_loading(self, layer: torch.nn.Module): """Forcibly resize parameters back to their original shapes (e.g., GPTQ format) before loading weights.""" - if not hasattr(layer, "_original_shapes"): return @@ -416,6 +464,8 @@ def apply_weights( g_idx2=layer.w2_weight_g_idx, sort_indices1=layer.w13_g_idx_sort_indices, sort_indices2=layer.w2_g_idx_sort_indices, + w1_zeros=layer.w13_weight_zero_point if not self.sym else None, + w2_zeros=layer.w2_weight_zero_point if not self.sym else None, num_bits=self.num_bits, is_k_full=self.is_k_full, routed_scaling_factor=self.moe_runner_config.routed_scaling_factor, diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index 293335f646d1..6db8f0684ff8 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -1533,6 +1533,20 @@ class InitWeightsSendGroupForRemoteInstanceReqOutput(BaseReq): message: str +@dataclass +class PostProcessWeightsReqInput(BaseReq): + # Whether to restore weights before loading new weights + restore_weights_before_load: bool = False + # Whether to enable quantization post-processing + post_process_quantization: bool = False + + +@dataclass +class PostProcessWeightsReqOutput(BaseReq): + success: bool + message: str + + @dataclass class SendWeightsToRemoteInstanceReqInput(BaseReq): # The master address diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 143054cd6e57..48a4d1b02cdd 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -124,6 +124,7 @@ LoadLoRAAdapterReqOutput, OpenSessionReqInput, PauseGenerationReqInput, + PostProcessWeightsReqInput, ProfileReq, ReleaseMemoryOccupationReqInput, RemoveExternalCorpusReqInput, @@ -1452,6 +1453,7 @@ def init_request_dispatcher(self): ), (UpdateWeightsFromTensorReqInput, self.update_weights_from_tensor), (UpdateWeightsFromIPCReqInput, self.update_weights_from_ipc), + (PostProcessWeightsReqInput, self.post_process_weights), (GetWeightsByNameReqInput, self.get_weights_by_name), (ReleaseMemoryOccupationReqInput, self.release_memory_occupation), (ResumeMemoryOccupationReqInput, self.resume_memory_occupation), diff --git a/python/sglang/srt/managers/scheduler_update_weights_mixin.py b/python/sglang/srt/managers/scheduler_update_weights_mixin.py index f2daf644def1..792fd6194db7 100644 --- a/python/sglang/srt/managers/scheduler_update_weights_mixin.py +++ b/python/sglang/srt/managers/scheduler_update_weights_mixin.py @@ -12,6 +12,7 @@ GPU_MEMORY_TYPE_KV_CACHE, GPU_MEMORY_TYPE_WEIGHTS, ) +from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.managers.io_struct import ( CheckWeightsReqInput, CheckWeightsReqOutput, @@ -21,6 +22,8 @@ GetWeightsByNameReqOutput, InitWeightsUpdateGroupReqInput, InitWeightsUpdateGroupReqOutput, + PostProcessWeightsReqInput, + PostProcessWeightsReqOutput, ReleaseMemoryOccupationReqInput, ReleaseMemoryOccupationReqOutput, ResumeMemoryOccupationReqInput, @@ -120,6 +123,11 @@ def update_weights_from_ipc( torch.distributed.barrier(group=self.tp_cpu_group) return UpdateWeightsFromIPCReqOutput(success, message) + def post_process_weights(self, recv_req: PostProcessWeightsReqInput): + """Optional post-processing for updated weights (e.g., Marlin conversion).""" + success, message = self.tp_worker.post_process_weights(recv_req) + return PostProcessWeightsReqOutput(success, message) + def get_weights_by_name(self: Scheduler, recv_req: GetWeightsByNameReqInput): parameter = self.tp_worker.get_weights_by_name(recv_req) return GetWeightsByNameReqOutput(parameter) @@ -143,6 +151,13 @@ def release_memory_occupation( self.memory_saver_adapter.pause(GPU_MEMORY_TYPE_KV_CACHE) self.flush_cache() + if self.disaggregation_mode == DisaggregationMode.DECODE: + if hasattr(self, "disagg_decode_prealloc_queue"): + self.disagg_decode_prealloc_queue.release_memory_occupation() + elif self.disaggregation_mode == DisaggregationMode.PREFILL: + if hasattr(self, "disagg_prefill_bootstrap_queue"): + self.disagg_prefill_bootstrap_queue.release_memory_occupation() + if GPU_MEMORY_TYPE_WEIGHTS in tags: self.stashed_model_static_state = _export_static_state( self.tp_worker.model_runner.model @@ -183,6 +198,13 @@ def resume_memory_occupation( if GPU_MEMORY_TYPE_KV_CACHE in tags: self.memory_saver_adapter.resume(GPU_MEMORY_TYPE_KV_CACHE) + if self.disaggregation_mode == DisaggregationMode.DECODE: + if hasattr(self, "disagg_decode_prealloc_queue"): + self.disagg_decode_prealloc_queue.resume_memory_occupation() + elif self.disaggregation_mode == DisaggregationMode.PREFILL: + if hasattr(self, "disagg_prefill_bootstrap_queue"): + self.disagg_prefill_bootstrap_queue.resume_memory_occupation() + return ResumeMemoryOccupationReqOutput() def check_weights(self: Scheduler, recv_req: CheckWeightsReqInput): diff --git a/python/sglang/srt/managers/tokenizer_control_mixin.py b/python/sglang/srt/managers/tokenizer_control_mixin.py index 05382e073eda..edb685dff9a5 100644 --- a/python/sglang/srt/managers/tokenizer_control_mixin.py +++ b/python/sglang/srt/managers/tokenizer_control_mixin.py @@ -55,6 +55,8 @@ LoadLoRAAdapterReqOutput, LoRAUpdateOutput, OpenSessionReqInput, + PostProcessWeightsReqInput, + PostProcessWeightsReqOutput, ProfileReq, ProfileReqOutput, ProfileReqType, @@ -102,6 +104,7 @@ ("send_weights_to_remote_instance", SendWeightsToRemoteInstanceReqOutput), ("update_weights_from_tensor", UpdateWeightsFromTensorReqOutput), ("update_weights_from_ipc", UpdateWeightsFromIPCReqOutput), + ("post_process_weights", PostProcessWeightsReqOutput), ("get_weights_by_name", GetWeightsByNameReqOutput), ("release_memory_occupation", ReleaseMemoryOccupationReqOutput), ("resume_memory_occupation", ResumeMemoryOccupationReqOutput), @@ -531,6 +534,17 @@ async def update_weights_from_ipc( return success, message + async def post_process_weights( + self: TokenizerManager, + obj: PostProcessWeightsReqInput, + request: Optional[fastapi.Request] = None, + ) -> Tuple[bool, str]: + """Trigger post-processing hooks for weights after loading.""" + self.auto_create_handle_loop() + async with self.model_update_lock.writer_lock: + results = await self.post_process_weights_communicator(obj) + return FanOutCommunicator.merge_results(results) + async def _unload_lora_adapter_locked( self: TokenizerManager, obj: UnloadLoRAAdapterReqInput, diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py index 60e105d93963..1687de74be49 100644 --- a/python/sglang/srt/managers/tp_worker.py +++ b/python/sglang/srt/managers/tp_worker.py @@ -29,6 +29,7 @@ InitWeightsUpdateGroupReqInput, LoadLoRAAdapterFromTensorsReqInput, LoadLoRAAdapterReqInput, + PostProcessWeightsReqInput, SendWeightsToRemoteInstanceReqInput, UnloadLoRAAdapterReqInput, UpdateWeightFromDiskReqInput, @@ -171,6 +172,11 @@ def update_weights_from_ipc(self, recv_req: UpdateWeightsFromIPCReqInput): success, message = self.model_runner.update_weights_from_ipc(recv_req) return success, message + def post_process_weights(self, recv_req: PostProcessWeightsReqInput): + """Perform optional post-processing on the updated model weights (e.g., Marlin conversion).""" + success, message = self.model_runner.post_process_weights(recv_req) + return success, message + def get_weights_by_name(self, recv_req: GetWeightsByNameReqInput): parameter = self.model_runner.get_weights_by_name( recv_req.name, recv_req.truncate_size diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index c32fbacbcaf0..fce3f9f28cb8 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -3646,6 +3646,32 @@ def _maybe_rebalance_after_rank_fault( ) return output + def post_process_weights(self, recv_req): + """Run quantization-specific post-processing hooks after loading weights.""" + from sglang.srt.model_loader.loader import device_loading_context + + target_device = torch.device("cuda", torch.cuda.current_device()) + + if recv_req.restore_weights_before_load: + for _, module in self.model.named_modules(): + quant_method = getattr(module, "quant_method", None) + if quant_method is not None and hasattr( + quant_method, "restore_weights_before_loading" + ): + with device_loading_context(module, target_device): + quant_method.restore_weights_before_loading(module) + + if recv_req.post_process_quantization: + for _, module in self.model.named_modules(): + quant_method = getattr(module, "quant_method", None) + if quant_method is not None and hasattr( + quant_method, "process_weights_after_loading" + ): + with device_loading_context(module, target_device): + quant_method.process_weights_after_loading(module) + + return True, "Success" + def _model_load_weights_direct(model, named_tensors: List[Tuple[str, torch.Tensor]]): params_dict = dict(model.named_parameters()) From ded3a7835c45b12c051c42c86f34d8e1b637b887 Mon Sep 17 00:00:00 2001 From: Yueming Yuan Date: Wed, 20 May 2026 19:40:16 -0700 Subject: [PATCH 20/50] [4/14] [sglang-miles] PD disaggregation for RL (#18646) --- python/sglang/srt/disaggregation/decode.py | 10 ++++++++++ python/sglang/srt/disaggregation/mooncake/conn.py | 13 +++++++++++++ python/sglang/srt/disaggregation/prefill.py | 9 +++++++++ python/sglang/srt/managers/schedule_batch.py | 6 ++++-- 4 files changed, 36 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 3cdf2af17a05..0afce703adb2 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -445,6 +445,16 @@ def _init_kv_manager(self) -> CommonKVManager: ) return kv_manager + def release_memory_occupation(self): + self.queue.clear() + self.retracted_queue.clear() + if hasattr(self.kv_manager, "deregister_buffer_to_engine"): + self.kv_manager.deregister_buffer_to_engine() + + def resume_memory_occupation(self): + if hasattr(self.kv_manager, "register_buffer_to_engine"): + self.kv_manager.register_buffer_to_engine() + def add(self, req: Req, is_retracted: bool = False) -> None: """Add a request to the pending queue.""" if self._check_if_req_exceed_kv_capacity(req): diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 634f2eae5b79..83ff375d977e 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -569,6 +569,19 @@ def send_kvcache_staged( ) return ret + def deregister_buffer_to_engine(self): + # Batch deregister KV data buffers + if self.kv_args.kv_data_ptrs: + self.engine.batch_deregister(self.kv_args.kv_data_ptrs) + + # Batch deregister auxiliary data buffers + if self.kv_args.aux_data_ptrs: + self.engine.batch_deregister(self.kv_args.aux_data_ptrs) + + # Batch deregister state/extra pool data buffers + if self.kv_args.state_data_ptrs: + self.engine.batch_deregister(self.kv_args.state_data_ptrs) + def _transfer_data(self, mooncake_session_id, transfer_blocks): if not transfer_blocks: return 0 diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 0e2ed6a1904e..c193250d11e9 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -354,6 +354,15 @@ def pop_bootstrapped( else: return bootstrapped_reqs, failed_reqs + def release_memory_occupation(self): + self.queue.clear() + if hasattr(self.kv_manager, "deregister_buffer_to_engine"): + self.kv_manager.deregister_buffer_to_engine() + + def resume_memory_occupation(self): + if hasattr(self.kv_manager, "register_buffer_to_engine"): + self.kv_manager.register_buffer_to_engine() + class SchedulerDisaggregationPrefillMixin: """ diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index feecc544160b..609681ae67bc 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -2193,8 +2193,10 @@ def retract_decode( while first_iter or ( not self.check_decode_mem(selected_indices=sorted_indices) ): - if len(sorted_indices) == 1: - # Always keep at least one request + # We should allow all requests to be retracted in decode disaggregation mode + # because there can be prealloc prefill requests. + num_minimum_reqs = 0 if server_args.disaggregation_mode == "decode" else 1 + if len(sorted_indices) == num_minimum_reqs: break first_iter = False From b5697ff88eb8d52367ca33a291993444f5c3ee8b Mon Sep 17 00:00:00 2001 From: Yueming Yuan Date: Wed, 20 May 2026 19:41:58 -0700 Subject: [PATCH 21/50] [5/14] [sglang-miles] MTP related fix (#18647) --- python/sglang/srt/server_args.py | 6 ++++++ .../srt/speculative/eagle_draft_cuda_graph_runner.py | 8 ++++++-- python/sglang/srt/speculative/eagle_worker.py | 5 ++++- 3 files changed, 16 insertions(+), 3 deletions(-) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 173bdbbf4c55..ca19dc1553e6 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -707,6 +707,7 @@ class ServerArgs: cuda_graph_max_bs: Optional[int] = None cuda_graph_bs: Optional[List[int]] = None disable_cuda_graph: bool = False + disable_draft_cuda_graph: bool = False disable_cuda_graph_padding: bool = False enable_breakable_cuda_graph: bool = False enable_profile_cuda_graph: bool = False @@ -6441,6 +6442,11 @@ def add_cli_args(parser: argparse.ArgumentParser): action="store_true", help="Disable cuda graph.", ) + parser.add_argument( + "--disable-draft-cuda-graph", + action="store_true", + help="Disable cuda graph for draft model in speculative decoding.", + ) parser.add_argument( "--disable-cuda-graph-padding", action="store_true", diff --git a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py index 1dd4bb03940a..4be3a8c31be4 100644 --- a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py @@ -442,8 +442,12 @@ def replay(self, forward_batch: ForwardBatch): "EagleDraftCudaGraphRunner.replay: topk_index vs vocab_size=" f"{self.model_runner.model_config.vocab_size}", ) - buffers.topk_p[:raw_bs].copy_(forward_batch.spec_info.topk_p) - buffers.topk_index[:raw_bs].copy_(forward_batch.spec_info.topk_index) + buffers.topk_p[:raw_bs].copy_(forward_batch.spec_info.topk_p.clamp(0, 1)) + buffers.topk_index[:raw_bs].copy_( + forward_batch.spec_info.topk_index.clamp( + 0, self.model_runner.model_config.vocab_size - 1 + ) + ) if ( buffers.hidden_states is not None and forward_batch.spec_info.hidden_states is not None diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py index 0a50b9182b72..a81e20add0ca 100644 --- a/python/sglang/srt/speculative/eagle_worker.py +++ b/python/sglang/srt/speculative/eagle_worker.py @@ -270,7 +270,10 @@ def init_cuda_graphs(self): self.cuda_graph_runner = None self.cuda_graph_runner_for_draft_extend = None - if self.server_args.disable_cuda_graph: + if ( + self.server_args.disable_cuda_graph + or self.server_args.disable_draft_cuda_graph + ): return Device2DraftCudaGraphRunner = { From 6efbb87ddbcb085f8c383b3abe5cfe23236d4d57 Mon Sep 17 00:00:00 2001 From: Yueming Yuan Date: Wed, 20 May 2026 19:42:06 -0700 Subject: [PATCH 22/50] [6/14] [sglang-miles] VLM training multimodal fallback fixes (#18781) --- python/sglang/srt/managers/scheduler_output_processor_mixin.py | 2 +- python/sglang/srt/multimodal/processors/qwen_vl.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py index e02f21163b62..0dbc50b1deb1 100644 --- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py +++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py @@ -1262,7 +1262,7 @@ def stream_output_generation( # Send to detokenizer if reqs or is_idle_batch: - if self.model_config.is_multimodal_gen: + if getattr(self.model_config, "is_multimodal_gen", False): return self.send_to_detokenizer.send_output( BatchTokenIDOutput( diff --git a/python/sglang/srt/multimodal/processors/qwen_vl.py b/python/sglang/srt/multimodal/processors/qwen_vl.py index fb9fd856be0a..b1cbc9b480eb 100644 --- a/python/sglang/srt/multimodal/processors/qwen_vl.py +++ b/python/sglang/srt/multimodal/processors/qwen_vl.py @@ -505,7 +505,7 @@ async def process_mm_data_async( **kwargs, ): entry_time = time.perf_counter() - base_output = self.load_mm_data( + base_output = self.legacy_load_mm_data( prompt=input_text, image_data=image_data, video_data=request_obj.video_data, From 0b891eb2ed7b637906cc3ad58dc9dc6cf5421519 Mon Sep 17 00:00:00 2001 From: Jiajun Li Date: Wed, 20 May 2026 19:45:36 -0700 Subject: [PATCH 23/50] [7/14] [sglang-miles] Support better token id return for TITO (#19731) --- .../sglang/srt/entrypoints/openai/protocol.py | 6 + .../srt/entrypoints/openai/serving_chat.py | 18 +- python/sglang/srt/entrypoints/openai/utils.py | 3 +- python/sglang/srt/managers/io_struct.py | 3 + .../sglang/srt/managers/tokenizer_manager.py | 20 + .../basic/test_return_token_ids.py | 491 ++++++++++++++++++ 6 files changed, 538 insertions(+), 3 deletions(-) create mode 100644 test/registered/openai_server/basic/test_return_token_ids.py diff --git a/python/sglang/srt/entrypoints/openai/protocol.py b/python/sglang/srt/entrypoints/openai/protocol.py index 58a137b19e5e..c550c46187ef 100644 --- a/python/sglang/srt/entrypoints/openai/protocol.py +++ b/python/sglang/srt/entrypoints/openai/protocol.py @@ -92,6 +92,7 @@ class LogProbs(BaseModel): text_offset: List[int] = Field(default_factory=list) token_logprobs: List[Optional[float]] = Field(default_factory=list) tokens: List[str] = Field(default_factory=list) + token_ids: List[int] = Field(default_factory=list) top_logprobs: List[Optional[Dict[str, float]]] = Field(default_factory=list) @@ -103,6 +104,7 @@ class TopLogprob(BaseModel): class ChatCompletionTokenLogprob(BaseModel): token: str + token_id: int bytes: List[int] logprob: float top_logprobs: List[TopLogprob] @@ -635,6 +637,7 @@ class ChatCompletionRequest(BaseModel): return_routed_experts: bool = False routed_experts_start_len: int = 0 return_cached_tokens_details: bool = False + return_prompt_token_ids: bool = False reasoning_effort: Optional[Literal["none", "low", "medium", "high", "max"]] = Field( default=None, description="Constrains effort on reasoning for reasoning models. " @@ -908,12 +911,15 @@ class ChatCompletionResponseChoice(BaseModel): ] = None matched_stop: Union[None, int, str] = None hidden_states: Optional[object] = None + prompt_token_ids: Optional[List[int]] = None @model_serializer(mode="wrap") def _serialize(self, handler): data = handler(self) if self.hidden_states is None: data.pop("hidden_states", None) + if self.prompt_token_ids is None: + data.pop("prompt_token_ids", None) return data diff --git a/python/sglang/srt/entrypoints/openai/serving_chat.py b/python/sglang/srt/entrypoints/openai/serving_chat.py index 26293fc27dcc..a077cf00952f 100644 --- a/python/sglang/srt/entrypoints/openai/serving_chat.py +++ b/python/sglang/srt/entrypoints/openai/serving_chat.py @@ -383,6 +383,12 @@ def _convert_to_internal_request( request.reasoning_effort = reasoning_effort """Convert OpenAI chat completion request to internal format""" + if request.return_prompt_token_ids and request.stream: + raise ValueError( + "return_prompt_token_ids is not supported with streaming. " + "Please set stream=false when using return_prompt_token_ids=true." + ) + is_multimodal = self.tokenizer_manager.model_config.is_multimodal # Process messages and apply chat template @@ -449,6 +455,7 @@ def _convert_to_internal_request( video_max_dynamic_patch=vid_max_dynamic_patch, max_dynamic_patch=getattr(request, "max_dynamic_patch", None), use_audio_in_video=getattr(request, "use_audio_in_video", False), + return_prompt_token_ids=request.return_prompt_token_ids, ) return adapted_request, request @@ -1139,6 +1146,11 @@ def _build_chat_response( # Handle hidden states hidden_states = process_hidden_states_from_ret(ret_item, request) + choice_prompt_token_ids = ( + ret_item.get("prompt_token_ids") + if request.return_prompt_token_ids + else None + ) finish_reason = ret_item["meta_info"]["finish_reason"] text = ret_item["text"] @@ -1199,6 +1211,7 @@ def _build_chat_response( else None ), hidden_states=hidden_states, + prompt_token_ids=choice_prompt_token_ids, ) choices.append(choice_data) @@ -1230,8 +1243,8 @@ def _process_logprobs_tokens( """ token_logprobs = [] - for token_idx, (token, logprob) in enumerate( - zip(logprobs.tokens, logprobs.token_logprobs) + for token_idx, (token, token_id, logprob) in enumerate( + zip(logprobs.tokens, logprobs.token_ids, logprobs.token_logprobs) ): token_bytes = list(token.encode("utf-8")) top_logprobs = [] @@ -1253,6 +1266,7 @@ def _process_logprobs_tokens( token_logprobs.append( ChatCompletionTokenLogprob( token=token, + token_id=token_id, bytes=token_bytes, logprob=logprob, top_logprobs=top_logprobs, diff --git a/python/sglang/srt/entrypoints/openai/utils.py b/python/sglang/srt/entrypoints/openai/utils.py index 7586f62f6c9b..38a31c29d9de 100644 --- a/python/sglang/srt/entrypoints/openai/utils.py +++ b/python/sglang/srt/entrypoints/openai/utils.py @@ -23,8 +23,9 @@ def to_openai_style_logprobs( ret_logprobs = LogProbs() def append_token_logprobs(token_logprobs): - for logprob, _, token_text in token_logprobs: + for logprob, token_id, token_text in token_logprobs: ret_logprobs.tokens.append(token_text) + ret_logprobs.token_ids.append(token_id) ret_logprobs.token_logprobs.append(logprob) # Not supported yet diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index 6db8f0684ff8..ca46b82108c5 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -242,6 +242,9 @@ class GenerateReqInput(BaseReq): # Whether to return entropy return_entropy: bool = False + # Whether to return prompt token IDs without computing logprobs + return_prompt_token_ids: bool = False + # Propagates trace context via Engine.generate/async_generate external_trace_header: Optional[Dict] = None received_time: Optional[float] = None diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 6375a1d8b7e3..e6d337de43c3 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -194,6 +194,7 @@ def get_crash_dump_output(self) -> Dict[Any, Any]: output_top_logprobs: List[Any] = dataclasses.field(default_factory=list) input_token_ids_logprobs: List[Any] = dataclasses.field(default_factory=list) output_token_ids_logprobs: List[Any] = dataclasses.field(default_factory=list) + prompt_token_ids: Optional[List[int]] = None def _slice_streaming_output_meta_info( @@ -554,6 +555,10 @@ async def generate_request( if obj.is_single: tokenized_obj = await self._tokenize_one_request(obj) self._send_one_request(tokenized_obj) + if getattr(obj, "return_prompt_token_ids", False): + self.rid_to_state[obj.rid].prompt_token_ids = list( + tokenized_obj.input_ids + ) async for response in self._wait_one_response(obj, request): yield response else: @@ -1405,6 +1410,10 @@ async def _handle_batch_request( # Set up generators for each request in the batch for i in range(batch_size): tmp_obj = obj[i] + if getattr(tmp_obj, "return_prompt_token_ids", False): + self.rid_to_state[tmp_obj.rid].prompt_token_ids = list( + tokenized_objs[i].input_ids + ) generators.append(self._wait_one_response(tmp_obj, request)) rids.append(tmp_obj.rid) else: @@ -1418,6 +1427,10 @@ async def _handle_batch_request( tmp_obj = obj[i] tokenized_obj = await self._tokenize_one_request(tmp_obj) self._send_one_request(tokenized_obj) + if getattr(tmp_obj, "return_prompt_token_ids", False): + self.rid_to_state[tmp_obj.rid].prompt_token_ids = list( + tokenized_obj.input_ids + ) generators.append(self._wait_one_response(tmp_obj, request)) rids.append(tmp_obj.rid) else: @@ -1456,6 +1469,10 @@ async def _handle_batch_request( self._init_req_state(tmp_obj) tokenized_obj.time_stats = self.rid_to_state[tmp_obj.rid].time_stats self._send_one_request(tokenized_obj) + if getattr(tmp_obj, "return_prompt_token_ids", False): + self.rid_to_state[tmp_obj.rid].prompt_token_ids = list( + tokenized_objs[i].input_ids + ) generators.append(self._wait_one_response(tmp_obj, request)) rids.append(tmp_obj.rid) @@ -1836,6 +1853,9 @@ async def _handle_batch_output( ): out_dict["pooled_hidden_state"] = recv_obj.pooled_hidden_states[i] + if out_dict is not None and state.prompt_token_ids is not None: + out_dict["prompt_token_ids"] = state.prompt_token_ids + # Set first_token_time on the first output batch. # This is the single write point for first_token_time. if state.time_stats.first_token_time == 0.0: diff --git a/test/registered/openai_server/basic/test_return_token_ids.py b/test/registered/openai_server/basic/test_return_token_ids.py new file mode 100644 index 000000000000..de4bb503a4e3 --- /dev/null +++ b/test/registered/openai_server/basic/test_return_token_ids.py @@ -0,0 +1,491 @@ +""" +Unit tests for the return_prompt_token_ids feature in ChatCompletion endpoint. + +Tests that: +1. Protocol models correctly handle return_prompt_token_ids / prompt_token_ids fields +2. ChatCompletionTokenLogprob includes token_id field +3. Request conversion passes return_prompt_token_ids flag through +4. Non-streaming response includes prompt_token_ids +5. Fields are omitted from JSON when return_prompt_token_ids is False (default) + +Run with: + python -m pytest test/registered/openai_server/basic/test_return_prompt_token_ids.py -v +or: + python test/registered/openai_server/basic/test_return_prompt_token_ids.py -v +""" + +import json +import sys +import unittest +from unittest.mock import MagicMock + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=10, suite="stage-a-test-cpu") + +# --------------------------------------------------------------------------- +# Mock out heavy GPU dependencies so tests run on CPU-only machines. +# We install a MagicMock for every missing module in the import chain. +# --------------------------------------------------------------------------- + +_GPU_MODULES = [ + # PyTorch + "torch", + "torch.nn", + "torch.nn.functional", + "torch.nn.parameter", + "torch.cuda", + "torch.distributed", + "torch.library", + "torch.utils", + "torch.utils.checkpoint", + "torch.fx", + "torch.profiler", + "torch.autograd", + "torch.amp", + "torch.optim", + # Triton + "triton", + "triton.language", + "triton.runtime", + # SGLang kernel / vLLM / transformers + "sgl_kernel", + "vllm", + "vllm.config", + "vllm.model_executor", + "transformers", + "transformers.models", + "outlines", + # CUDA-specific + "cuda", + "cupy", + "numba", + # Packaging + "packaging", + "packaging.version", +] + +_mock_cache = {} + +for mod_name in _GPU_MODULES: + if mod_name not in sys.modules: + mock = MagicMock() + sys.modules[mod_name] = mock + _mock_cache[mod_name] = mock + +# --------------------------------------------------------------------------- +# Now safe to import sglang modules +# --------------------------------------------------------------------------- + +import asyncio + +from sglang.srt.entrypoints.openai.protocol import ( + ChatCompletionRequest, + ChatCompletionResponse, + ChatCompletionResponseChoice, + ChatCompletionTokenLogprob, + ChatMessage, + ChoiceLogprobs, + UsageInfo, +) + +# These may fail on CPU-only if the import chain hits something we missed. +# We protect with try/except and skip the dependent tests. +_HAS_IO_STRUCT = False +_HAS_SERVING_CHAT = False +_HAS_TOKENIZER_MANAGER = False + +try: + from sglang.srt.managers.io_struct import GenerateReqInput + + _HAS_IO_STRUCT = True +except Exception: + pass + +try: + from sglang.srt.entrypoints.openai.protocol import MessageProcessingResult + from sglang.srt.entrypoints.openai.serving_chat import OpenAIServingChat + + _HAS_SERVING_CHAT = True +except Exception: + pass + +try: + from sglang.srt.managers.tokenizer_manager import ReqState + + _HAS_TOKENIZER_MANAGER = True +except Exception: + pass + +# --------------------------------------------------------------------------- +# Constants +# --------------------------------------------------------------------------- + +MOCK_PROMPT_TOKEN_IDS = [128000, 882, 1234, 5678, 9012] +MOCK_OUTPUT_TOKEN_IDS = [100, 200, 300] + + +# =========================================================================== +# 1. Protocol Tests — pure Pydantic model serialization (always runnable) +# =========================================================================== + + +class TestReturnTokenIdsProtocol(unittest.TestCase): + """Test protocol model fields for return_prompt_token_ids.""" + + # --- Request --- + + def test_request_default_false(self): + req = ChatCompletionRequest( + model="test", + messages=[{"role": "user", "content": "Hi"}], + ) + self.assertFalse(req.return_prompt_token_ids) + + def test_request_explicit_true(self): + req = ChatCompletionRequest( + model="test", + messages=[{"role": "user", "content": "Hi"}], + return_prompt_token_ids=True, + ) + self.assertTrue(req.return_prompt_token_ids) + + # --- Response (non-streaming) --- + + def test_choice_omits_prompt_token_ids_when_none(self): + choice = ChatCompletionResponseChoice( + index=0, + message=ChatMessage(role="assistant", content="hi"), + finish_reason="stop", + ) + data = choice.model_dump() + self.assertNotIn("prompt_token_ids", data) + + def test_choice_includes_prompt_token_ids_when_set(self): + choice = ChatCompletionResponseChoice( + index=0, + message=ChatMessage(role="assistant", content="hi"), + finish_reason="stop", + prompt_token_ids=[1, 2, 3], + ) + data = choice.model_dump() + self.assertIn("prompt_token_ids", data) + self.assertEqual(data["prompt_token_ids"], [1, 2, 3]) + + # --- ChatCompletionTokenLogprob --- + + def test_token_logprob_includes_token_id(self): + logprob = ChatCompletionTokenLogprob( + token="hello", + token_id=12345, + bytes=list(b"hello"), + logprob=-0.5, + top_logprobs=[], + ) + data = logprob.model_dump() + self.assertEqual(data["token_id"], 12345) + + def test_token_logprob_in_choice_logprobs(self): + """token_id should appear in serialized logprobs.content entries.""" + logprob_entry = ChatCompletionTokenLogprob( + token="hi", + token_id=100, + bytes=list(b"hi"), + logprob=-1.0, + top_logprobs=[], + ) + choice = ChatCompletionResponseChoice( + index=0, + message=ChatMessage(role="assistant", content="hi"), + logprobs=ChoiceLogprobs(content=[logprob_entry]), + finish_reason="stop", + ) + data = choice.model_dump() + self.assertEqual(data["logprobs"]["content"][0]["token_id"], 100) + + # --- Full JSON round-trip --- + + def test_full_response_json_with_prompt_token_ids(self): + choice = ChatCompletionResponseChoice( + index=0, + message=ChatMessage(role="assistant", content="hello"), + finish_reason="stop", + prompt_token_ids=MOCK_PROMPT_TOKEN_IDS, + ) + resp = ChatCompletionResponse( + id="test-id", + model="test", + choices=[choice], + usage=UsageInfo(prompt_tokens=5, completion_tokens=3, total_tokens=8), + ) + data = json.loads(resp.model_dump_json()) + self.assertEqual(data["choices"][0]["prompt_token_ids"], MOCK_PROMPT_TOKEN_IDS) + + def test_full_response_json_without_token_ids(self): + choice = ChatCompletionResponseChoice( + index=0, + message=ChatMessage(role="assistant", content="hello"), + finish_reason="stop", + ) + resp = ChatCompletionResponse( + id="test-id", + model="test", + choices=[choice], + usage=UsageInfo(prompt_tokens=5, completion_tokens=3, total_tokens=8), + ) + data = json.loads(resp.model_dump_json()) + self.assertNotIn("prompt_token_ids", data["choices"][0]) + + +# =========================================================================== +# 2. GenerateReqInput Tests +# =========================================================================== + + +@unittest.skipUnless(_HAS_IO_STRUCT, "io_struct import requires GPU deps") +class TestReturnTokenIdsIOStruct(unittest.TestCase): + """Test GenerateReqInput return_prompt_token_ids field.""" + + def test_default_false(self): + req = GenerateReqInput(text="hello") + self.assertFalse(req.return_prompt_token_ids) + + def test_explicit_true(self): + req = GenerateReqInput(text="hello", return_prompt_token_ids=True) + self.assertTrue(req.return_prompt_token_ids) + + def test_does_not_affect_logprob_fields(self): + req = GenerateReqInput( + text="hello", + return_prompt_token_ids=True, + return_logprob=False, + logprob_start_len=-1, + ) + self.assertTrue(req.return_prompt_token_ids) + self.assertFalse(req.return_logprob) + self.assertEqual(req.logprob_start_len, -1) + + +# =========================================================================== +# 3. Request Conversion Tests +# =========================================================================== + + +@unittest.skipUnless(_HAS_SERVING_CHAT, "OpenAIServingChat import requires GPU deps") +class TestReturnTokenIdsRequestConversion(unittest.TestCase): + """Test that return_prompt_token_ids flows through _convert_to_internal_request.""" + + def setUp(self): + from unittest.mock import Mock + + tm = Mock() + tm.model_config = Mock(is_multimodal=False) + tm.server_args = Mock( + enable_cache_report=False, + tool_call_parser=None, + reasoning_parser=None, + ) + mock_hf_config = Mock() + mock_hf_config.architectures = ["LlamaForCausalLM"] + tm.model_config.hf_config = mock_hf_config + tm.chat_template_name = "llama-3" + tm.tokenizer = Mock() + tm.tokenizer.encode.return_value = [1, 2, 3] + tm.tokenizer.chat_template = None + tm.tokenizer.bos_token_id = 1 + + template_mgr = Mock() + template_mgr.chat_template_name = "llama-3" + template_mgr.jinja_template_content_format = None + template_mgr.completion_template_name = None + template_mgr.force_reasoning = False + + self.chat = OpenAIServingChat(tm, template_mgr) + + def _convert(self, return_prompt_token_ids: bool): + from unittest.mock import patch + + req = ChatCompletionRequest( + model="x", + messages=[{"role": "user", "content": "Hi"}], + return_prompt_token_ids=return_prompt_token_ids, + ) + with patch.object(self.chat, "_process_messages") as proc_mock: + proc_mock.return_value = MessageProcessingResult( + "Test prompt", + [1, 2, 3], + None, + None, + [], + [""], + None, + ) + adapted, _ = self.chat._convert_to_internal_request(req) + return adapted + + def test_flag_passed_when_true(self): + adapted = self._convert(return_prompt_token_ids=True) + self.assertIsInstance(adapted, GenerateReqInput) + self.assertTrue(adapted.return_prompt_token_ids) + + def test_flag_passed_when_false(self): + adapted = self._convert(return_prompt_token_ids=False) + self.assertIsInstance(adapted, GenerateReqInput) + self.assertFalse(adapted.return_prompt_token_ids) + + def test_logprob_not_affected(self): + adapted = self._convert(return_prompt_token_ids=True) + self.assertFalse(adapted.return_logprob) + self.assertEqual(adapted.logprob_start_len, -1) + + def test_stream_with_return_prompt_token_ids_raises(self): + """return_prompt_token_ids=True + stream=True should raise ValueError.""" + from unittest.mock import patch + + req = ChatCompletionRequest( + model="x", + messages=[{"role": "user", "content": "Hi"}], + return_prompt_token_ids=True, + stream=True, + ) + with patch.object(self.chat, "_process_messages") as proc_mock: + proc_mock.return_value = MessageProcessingResult( + "Test prompt", + [1, 2, 3], + None, + None, + [], + [""], + None, + ) + with self.assertRaises(ValueError): + self.chat._convert_to_internal_request(req) + + +# =========================================================================== +# 4. Response Building Tests +# =========================================================================== + + +@unittest.skipUnless(_HAS_SERVING_CHAT, "OpenAIServingChat import requires GPU deps") +class TestReturnTokenIdsResponseBuilding(unittest.TestCase): + """Test _build_chat_response includes prompt_token_ids when requested.""" + + def setUp(self): + from unittest.mock import Mock + + tm = Mock() + tm.model_config = Mock(is_multimodal=False) + tm.server_args = Mock( + enable_cache_report=False, + tool_call_parser=None, + reasoning_parser=None, + ) + mock_hf_config = Mock() + mock_hf_config.architectures = ["LlamaForCausalLM"] + tm.model_config.hf_config = mock_hf_config + tm.chat_template_name = "llama-3" + tm.tokenizer = Mock() + tm.tokenizer.chat_template = None + tm.tokenizer.bos_token_id = 1 + + template_mgr = Mock() + template_mgr.chat_template_name = "llama-3" + template_mgr.jinja_template_content_format = None + template_mgr.completion_template_name = None + template_mgr.force_reasoning = False + + self.chat = OpenAIServingChat(tm, template_mgr) + + def _make_ret(self, include_prompt_token_ids: bool = False): + ret = { + "text": "Test response", + "output_ids": MOCK_OUTPUT_TOKEN_IDS, + "meta_info": { + "id": "chatcmpl-test", + "prompt_tokens": 5, + "completion_tokens": 3, + "cached_tokens": 0, + "finish_reason": {"type": "stop", "matched": None}, + "output_token_logprobs": [], + "output_top_logprobs": None, + "weight_version": "default", + }, + } + if include_prompt_token_ids: + ret["prompt_token_ids"] = MOCK_PROMPT_TOKEN_IDS + return ret + + def test_with_return_prompt_token_ids(self): + req = ChatCompletionRequest( + model="x", + messages=[{"role": "user", "content": "Hi"}], + return_prompt_token_ids=True, + ) + ret = [self._make_ret(include_prompt_token_ids=True)] + response = self.chat._build_chat_response(req, ret, created=0) + + self.assertIsInstance(response, ChatCompletionResponse) + self.assertEqual(response.choices[0].prompt_token_ids, MOCK_PROMPT_TOKEN_IDS) + + def test_without_return_prompt_token_ids(self): + req = ChatCompletionRequest( + model="x", + messages=[{"role": "user", "content": "Hi"}], + ) + ret = [self._make_ret(include_prompt_token_ids=False)] + response = self.chat._build_chat_response(req, ret, created=0) + + self.assertIsNone(response.choices[0].prompt_token_ids) + + data = json.loads(response.model_dump_json()) + self.assertNotIn("prompt_token_ids", data["choices"][0]) + + def test_json_round_trip(self): + req = ChatCompletionRequest( + model="x", + messages=[{"role": "user", "content": "Hi"}], + return_prompt_token_ids=True, + ) + ret = [self._make_ret(include_prompt_token_ids=True)] + response = self.chat._build_chat_response(req, ret, created=0) + + data = json.loads(response.model_dump_json()) + self.assertEqual(data["choices"][0]["prompt_token_ids"], MOCK_PROMPT_TOKEN_IDS) + + +# =========================================================================== +# 5. ReqState Tests +# =========================================================================== + + +@unittest.skipUnless( + _HAS_TOKENIZER_MANAGER, "tokenizer_manager import requires GPU deps" +) +class TestReturnTokenIdsReqState(unittest.TestCase): + """Test that ReqState stores prompt_token_ids correctly.""" + + def test_reqstate_default_none(self): + state = ReqState( + out_list=[], + finished=False, + event=asyncio.Event(), + obj=MagicMock(), + time_stats=MagicMock(), + ) + self.assertIsNone(state.prompt_token_ids) + + def test_reqstate_stores_prompt_token_ids(self): + state = ReqState( + out_list=[], + finished=False, + event=asyncio.Event(), + obj=MagicMock(), + time_stats=MagicMock(), + ) + state.prompt_token_ids = MOCK_PROMPT_TOKEN_IDS + self.assertEqual(state.prompt_token_ids, MOCK_PROMPT_TOKEN_IDS) + + +if __name__ == "__main__": + unittest.main() From 1ca4d3046fe9c6c20da8f1b721e9f39951c408ec Mon Sep 17 00:00:00 2001 From: Jiajun Li Date: Wed, 20 May 2026 19:48:02 -0700 Subject: [PATCH 24/50] [8/14] [sglang-miles] Support cross turn token in after last user message (#20066) --- .../sglang/srt/entrypoints/openai/protocol.py | 11 +- .../srt/entrypoints/openai/serving_chat.py | 35 +- python/sglang/srt/entrypoints/openai/utils.py | 3 +- .../basic/test_return_token_ids.py | 491 ------------------ 4 files changed, 38 insertions(+), 502 deletions(-) delete mode 100644 test/registered/openai_server/basic/test_return_token_ids.py diff --git a/python/sglang/srt/entrypoints/openai/protocol.py b/python/sglang/srt/entrypoints/openai/protocol.py index c550c46187ef..a5b9f82dd890 100644 --- a/python/sglang/srt/entrypoints/openai/protocol.py +++ b/python/sglang/srt/entrypoints/openai/protocol.py @@ -92,7 +92,6 @@ class LogProbs(BaseModel): text_offset: List[int] = Field(default_factory=list) token_logprobs: List[Optional[float]] = Field(default_factory=list) tokens: List[str] = Field(default_factory=list) - token_ids: List[int] = Field(default_factory=list) top_logprobs: List[Optional[Dict[str, float]]] = Field(default_factory=list) @@ -104,7 +103,6 @@ class TopLogprob(BaseModel): class ChatCompletionTokenLogprob(BaseModel): token: str - token_id: int bytes: List[int] logprob: float top_logprobs: List[TopLogprob] @@ -638,6 +636,7 @@ class ChatCompletionRequest(BaseModel): routed_experts_start_len: int = 0 return_cached_tokens_details: bool = False return_prompt_token_ids: bool = False + return_meta_info: bool = False reasoning_effort: Optional[Literal["none", "low", "medium", "high", "max"]] = Field( default=None, description="Constrains effort on reasoning for reasoning models. " @@ -687,6 +686,11 @@ class ChatCompletionRequest(BaseModel): custom_logit_processor: Optional[Union[List[Optional[str]], str]] = None custom_params: Optional[Dict] = None + # Pre-computed prompt token IDs: when provided, bypasses chat template + # tokenization entirely. Messages are still used to derive stop tokens + # and tool_call_constraint. + input_ids: Optional[List[int]] = None + # For request id rid: Optional[Union[List[str], str]] = None # Extra key for classifying the request (e.g. cache_salt) @@ -912,6 +916,7 @@ class ChatCompletionResponseChoice(BaseModel): matched_stop: Union[None, int, str] = None hidden_states: Optional[object] = None prompt_token_ids: Optional[List[int]] = None + meta_info: Optional[Dict[str, Any]] = None @model_serializer(mode="wrap") def _serialize(self, handler): @@ -920,6 +925,8 @@ def _serialize(self, handler): data.pop("hidden_states", None) if self.prompt_token_ids is None: data.pop("prompt_token_ids", None) + if self.meta_info is None: + data.pop("meta_info", None) return data diff --git a/python/sglang/srt/entrypoints/openai/serving_chat.py b/python/sglang/srt/entrypoints/openai/serving_chat.py index a077cf00952f..5e2a70376541 100644 --- a/python/sglang/srt/entrypoints/openai/serving_chat.py +++ b/python/sglang/srt/entrypoints/openai/serving_chat.py @@ -393,7 +393,6 @@ def _convert_to_internal_request( # Process messages and apply chat template processed_messages = self._process_messages(request, is_multimodal) - # Build sampling parameters sampling_params = request.to_sampling_params( stop=processed_messages.stop, @@ -511,8 +510,19 @@ def _process_messages( ) tool_call_constraint = ("json_schema", json_schema) - # Use chat template - if self.template_manager.chat_template_name is None: + # When input_ids are provided, skip template tokenization entirely; + # only stop tokens and tool_call_constraint are needed. + if request.input_ids is not None: + result = MessageProcessingResult( + prompt=self.tokenizer_manager.tokenizer.decode(request.input_ids), + prompt_ids=request.input_ids, + image_data=None, + audio_data=None, + video_data=None, + modalities=[], + stop=request.stop or [], + ) + elif self.template_manager.chat_template_name is None: result = self._apply_jinja_template(request, tools, is_multimodal) else: result = self._apply_conversation_template(request, is_multimodal) @@ -1195,11 +1205,22 @@ def _build_chat_response( history_tool_calls_cnt, ) + # Extract prompt_token_ids if requested + choice_prompt_token_ids = ( + ret_item.get("prompt_token_ids") + if request.return_prompt_token_ids + else None + ) + + choice_meta_info = ( + ret_item["meta_info"] if request.return_meta_info else None + ) + # NOTE: content should not be None but empty string to make sure retokenize consistency. choice_data = ChatCompletionResponseChoice( index=idx, message=ChatMessage( role="assistant", - content=text if text else None, + content=text if text else "", tool_calls=tool_calls, reasoning_content=reasoning_text if reasoning_text else None, ), @@ -1212,6 +1233,7 @@ def _build_chat_response( ), hidden_states=hidden_states, prompt_token_ids=choice_prompt_token_ids, + meta_info=choice_meta_info, ) choices.append(choice_data) @@ -1243,8 +1265,8 @@ def _process_logprobs_tokens( """ token_logprobs = [] - for token_idx, (token, token_id, logprob) in enumerate( - zip(logprobs.tokens, logprobs.token_ids, logprobs.token_logprobs) + for token_idx, (token, logprob) in enumerate( + zip(logprobs.tokens, logprobs.token_logprobs) ): token_bytes = list(token.encode("utf-8")) top_logprobs = [] @@ -1266,7 +1288,6 @@ def _process_logprobs_tokens( token_logprobs.append( ChatCompletionTokenLogprob( token=token, - token_id=token_id, bytes=token_bytes, logprob=logprob, top_logprobs=top_logprobs, diff --git a/python/sglang/srt/entrypoints/openai/utils.py b/python/sglang/srt/entrypoints/openai/utils.py index 38a31c29d9de..7586f62f6c9b 100644 --- a/python/sglang/srt/entrypoints/openai/utils.py +++ b/python/sglang/srt/entrypoints/openai/utils.py @@ -23,9 +23,8 @@ def to_openai_style_logprobs( ret_logprobs = LogProbs() def append_token_logprobs(token_logprobs): - for logprob, token_id, token_text in token_logprobs: + for logprob, _, token_text in token_logprobs: ret_logprobs.tokens.append(token_text) - ret_logprobs.token_ids.append(token_id) ret_logprobs.token_logprobs.append(logprob) # Not supported yet diff --git a/test/registered/openai_server/basic/test_return_token_ids.py b/test/registered/openai_server/basic/test_return_token_ids.py deleted file mode 100644 index de4bb503a4e3..000000000000 --- a/test/registered/openai_server/basic/test_return_token_ids.py +++ /dev/null @@ -1,491 +0,0 @@ -""" -Unit tests for the return_prompt_token_ids feature in ChatCompletion endpoint. - -Tests that: -1. Protocol models correctly handle return_prompt_token_ids / prompt_token_ids fields -2. ChatCompletionTokenLogprob includes token_id field -3. Request conversion passes return_prompt_token_ids flag through -4. Non-streaming response includes prompt_token_ids -5. Fields are omitted from JSON when return_prompt_token_ids is False (default) - -Run with: - python -m pytest test/registered/openai_server/basic/test_return_prompt_token_ids.py -v -or: - python test/registered/openai_server/basic/test_return_prompt_token_ids.py -v -""" - -import json -import sys -import unittest -from unittest.mock import MagicMock - -from sglang.test.ci.ci_register import register_cpu_ci - -register_cpu_ci(est_time=10, suite="stage-a-test-cpu") - -# --------------------------------------------------------------------------- -# Mock out heavy GPU dependencies so tests run on CPU-only machines. -# We install a MagicMock for every missing module in the import chain. -# --------------------------------------------------------------------------- - -_GPU_MODULES = [ - # PyTorch - "torch", - "torch.nn", - "torch.nn.functional", - "torch.nn.parameter", - "torch.cuda", - "torch.distributed", - "torch.library", - "torch.utils", - "torch.utils.checkpoint", - "torch.fx", - "torch.profiler", - "torch.autograd", - "torch.amp", - "torch.optim", - # Triton - "triton", - "triton.language", - "triton.runtime", - # SGLang kernel / vLLM / transformers - "sgl_kernel", - "vllm", - "vllm.config", - "vllm.model_executor", - "transformers", - "transformers.models", - "outlines", - # CUDA-specific - "cuda", - "cupy", - "numba", - # Packaging - "packaging", - "packaging.version", -] - -_mock_cache = {} - -for mod_name in _GPU_MODULES: - if mod_name not in sys.modules: - mock = MagicMock() - sys.modules[mod_name] = mock - _mock_cache[mod_name] = mock - -# --------------------------------------------------------------------------- -# Now safe to import sglang modules -# --------------------------------------------------------------------------- - -import asyncio - -from sglang.srt.entrypoints.openai.protocol import ( - ChatCompletionRequest, - ChatCompletionResponse, - ChatCompletionResponseChoice, - ChatCompletionTokenLogprob, - ChatMessage, - ChoiceLogprobs, - UsageInfo, -) - -# These may fail on CPU-only if the import chain hits something we missed. -# We protect with try/except and skip the dependent tests. -_HAS_IO_STRUCT = False -_HAS_SERVING_CHAT = False -_HAS_TOKENIZER_MANAGER = False - -try: - from sglang.srt.managers.io_struct import GenerateReqInput - - _HAS_IO_STRUCT = True -except Exception: - pass - -try: - from sglang.srt.entrypoints.openai.protocol import MessageProcessingResult - from sglang.srt.entrypoints.openai.serving_chat import OpenAIServingChat - - _HAS_SERVING_CHAT = True -except Exception: - pass - -try: - from sglang.srt.managers.tokenizer_manager import ReqState - - _HAS_TOKENIZER_MANAGER = True -except Exception: - pass - -# --------------------------------------------------------------------------- -# Constants -# --------------------------------------------------------------------------- - -MOCK_PROMPT_TOKEN_IDS = [128000, 882, 1234, 5678, 9012] -MOCK_OUTPUT_TOKEN_IDS = [100, 200, 300] - - -# =========================================================================== -# 1. Protocol Tests — pure Pydantic model serialization (always runnable) -# =========================================================================== - - -class TestReturnTokenIdsProtocol(unittest.TestCase): - """Test protocol model fields for return_prompt_token_ids.""" - - # --- Request --- - - def test_request_default_false(self): - req = ChatCompletionRequest( - model="test", - messages=[{"role": "user", "content": "Hi"}], - ) - self.assertFalse(req.return_prompt_token_ids) - - def test_request_explicit_true(self): - req = ChatCompletionRequest( - model="test", - messages=[{"role": "user", "content": "Hi"}], - return_prompt_token_ids=True, - ) - self.assertTrue(req.return_prompt_token_ids) - - # --- Response (non-streaming) --- - - def test_choice_omits_prompt_token_ids_when_none(self): - choice = ChatCompletionResponseChoice( - index=0, - message=ChatMessage(role="assistant", content="hi"), - finish_reason="stop", - ) - data = choice.model_dump() - self.assertNotIn("prompt_token_ids", data) - - def test_choice_includes_prompt_token_ids_when_set(self): - choice = ChatCompletionResponseChoice( - index=0, - message=ChatMessage(role="assistant", content="hi"), - finish_reason="stop", - prompt_token_ids=[1, 2, 3], - ) - data = choice.model_dump() - self.assertIn("prompt_token_ids", data) - self.assertEqual(data["prompt_token_ids"], [1, 2, 3]) - - # --- ChatCompletionTokenLogprob --- - - def test_token_logprob_includes_token_id(self): - logprob = ChatCompletionTokenLogprob( - token="hello", - token_id=12345, - bytes=list(b"hello"), - logprob=-0.5, - top_logprobs=[], - ) - data = logprob.model_dump() - self.assertEqual(data["token_id"], 12345) - - def test_token_logprob_in_choice_logprobs(self): - """token_id should appear in serialized logprobs.content entries.""" - logprob_entry = ChatCompletionTokenLogprob( - token="hi", - token_id=100, - bytes=list(b"hi"), - logprob=-1.0, - top_logprobs=[], - ) - choice = ChatCompletionResponseChoice( - index=0, - message=ChatMessage(role="assistant", content="hi"), - logprobs=ChoiceLogprobs(content=[logprob_entry]), - finish_reason="stop", - ) - data = choice.model_dump() - self.assertEqual(data["logprobs"]["content"][0]["token_id"], 100) - - # --- Full JSON round-trip --- - - def test_full_response_json_with_prompt_token_ids(self): - choice = ChatCompletionResponseChoice( - index=0, - message=ChatMessage(role="assistant", content="hello"), - finish_reason="stop", - prompt_token_ids=MOCK_PROMPT_TOKEN_IDS, - ) - resp = ChatCompletionResponse( - id="test-id", - model="test", - choices=[choice], - usage=UsageInfo(prompt_tokens=5, completion_tokens=3, total_tokens=8), - ) - data = json.loads(resp.model_dump_json()) - self.assertEqual(data["choices"][0]["prompt_token_ids"], MOCK_PROMPT_TOKEN_IDS) - - def test_full_response_json_without_token_ids(self): - choice = ChatCompletionResponseChoice( - index=0, - message=ChatMessage(role="assistant", content="hello"), - finish_reason="stop", - ) - resp = ChatCompletionResponse( - id="test-id", - model="test", - choices=[choice], - usage=UsageInfo(prompt_tokens=5, completion_tokens=3, total_tokens=8), - ) - data = json.loads(resp.model_dump_json()) - self.assertNotIn("prompt_token_ids", data["choices"][0]) - - -# =========================================================================== -# 2. GenerateReqInput Tests -# =========================================================================== - - -@unittest.skipUnless(_HAS_IO_STRUCT, "io_struct import requires GPU deps") -class TestReturnTokenIdsIOStruct(unittest.TestCase): - """Test GenerateReqInput return_prompt_token_ids field.""" - - def test_default_false(self): - req = GenerateReqInput(text="hello") - self.assertFalse(req.return_prompt_token_ids) - - def test_explicit_true(self): - req = GenerateReqInput(text="hello", return_prompt_token_ids=True) - self.assertTrue(req.return_prompt_token_ids) - - def test_does_not_affect_logprob_fields(self): - req = GenerateReqInput( - text="hello", - return_prompt_token_ids=True, - return_logprob=False, - logprob_start_len=-1, - ) - self.assertTrue(req.return_prompt_token_ids) - self.assertFalse(req.return_logprob) - self.assertEqual(req.logprob_start_len, -1) - - -# =========================================================================== -# 3. Request Conversion Tests -# =========================================================================== - - -@unittest.skipUnless(_HAS_SERVING_CHAT, "OpenAIServingChat import requires GPU deps") -class TestReturnTokenIdsRequestConversion(unittest.TestCase): - """Test that return_prompt_token_ids flows through _convert_to_internal_request.""" - - def setUp(self): - from unittest.mock import Mock - - tm = Mock() - tm.model_config = Mock(is_multimodal=False) - tm.server_args = Mock( - enable_cache_report=False, - tool_call_parser=None, - reasoning_parser=None, - ) - mock_hf_config = Mock() - mock_hf_config.architectures = ["LlamaForCausalLM"] - tm.model_config.hf_config = mock_hf_config - tm.chat_template_name = "llama-3" - tm.tokenizer = Mock() - tm.tokenizer.encode.return_value = [1, 2, 3] - tm.tokenizer.chat_template = None - tm.tokenizer.bos_token_id = 1 - - template_mgr = Mock() - template_mgr.chat_template_name = "llama-3" - template_mgr.jinja_template_content_format = None - template_mgr.completion_template_name = None - template_mgr.force_reasoning = False - - self.chat = OpenAIServingChat(tm, template_mgr) - - def _convert(self, return_prompt_token_ids: bool): - from unittest.mock import patch - - req = ChatCompletionRequest( - model="x", - messages=[{"role": "user", "content": "Hi"}], - return_prompt_token_ids=return_prompt_token_ids, - ) - with patch.object(self.chat, "_process_messages") as proc_mock: - proc_mock.return_value = MessageProcessingResult( - "Test prompt", - [1, 2, 3], - None, - None, - [], - [""], - None, - ) - adapted, _ = self.chat._convert_to_internal_request(req) - return adapted - - def test_flag_passed_when_true(self): - adapted = self._convert(return_prompt_token_ids=True) - self.assertIsInstance(adapted, GenerateReqInput) - self.assertTrue(adapted.return_prompt_token_ids) - - def test_flag_passed_when_false(self): - adapted = self._convert(return_prompt_token_ids=False) - self.assertIsInstance(adapted, GenerateReqInput) - self.assertFalse(adapted.return_prompt_token_ids) - - def test_logprob_not_affected(self): - adapted = self._convert(return_prompt_token_ids=True) - self.assertFalse(adapted.return_logprob) - self.assertEqual(adapted.logprob_start_len, -1) - - def test_stream_with_return_prompt_token_ids_raises(self): - """return_prompt_token_ids=True + stream=True should raise ValueError.""" - from unittest.mock import patch - - req = ChatCompletionRequest( - model="x", - messages=[{"role": "user", "content": "Hi"}], - return_prompt_token_ids=True, - stream=True, - ) - with patch.object(self.chat, "_process_messages") as proc_mock: - proc_mock.return_value = MessageProcessingResult( - "Test prompt", - [1, 2, 3], - None, - None, - [], - [""], - None, - ) - with self.assertRaises(ValueError): - self.chat._convert_to_internal_request(req) - - -# =========================================================================== -# 4. Response Building Tests -# =========================================================================== - - -@unittest.skipUnless(_HAS_SERVING_CHAT, "OpenAIServingChat import requires GPU deps") -class TestReturnTokenIdsResponseBuilding(unittest.TestCase): - """Test _build_chat_response includes prompt_token_ids when requested.""" - - def setUp(self): - from unittest.mock import Mock - - tm = Mock() - tm.model_config = Mock(is_multimodal=False) - tm.server_args = Mock( - enable_cache_report=False, - tool_call_parser=None, - reasoning_parser=None, - ) - mock_hf_config = Mock() - mock_hf_config.architectures = ["LlamaForCausalLM"] - tm.model_config.hf_config = mock_hf_config - tm.chat_template_name = "llama-3" - tm.tokenizer = Mock() - tm.tokenizer.chat_template = None - tm.tokenizer.bos_token_id = 1 - - template_mgr = Mock() - template_mgr.chat_template_name = "llama-3" - template_mgr.jinja_template_content_format = None - template_mgr.completion_template_name = None - template_mgr.force_reasoning = False - - self.chat = OpenAIServingChat(tm, template_mgr) - - def _make_ret(self, include_prompt_token_ids: bool = False): - ret = { - "text": "Test response", - "output_ids": MOCK_OUTPUT_TOKEN_IDS, - "meta_info": { - "id": "chatcmpl-test", - "prompt_tokens": 5, - "completion_tokens": 3, - "cached_tokens": 0, - "finish_reason": {"type": "stop", "matched": None}, - "output_token_logprobs": [], - "output_top_logprobs": None, - "weight_version": "default", - }, - } - if include_prompt_token_ids: - ret["prompt_token_ids"] = MOCK_PROMPT_TOKEN_IDS - return ret - - def test_with_return_prompt_token_ids(self): - req = ChatCompletionRequest( - model="x", - messages=[{"role": "user", "content": "Hi"}], - return_prompt_token_ids=True, - ) - ret = [self._make_ret(include_prompt_token_ids=True)] - response = self.chat._build_chat_response(req, ret, created=0) - - self.assertIsInstance(response, ChatCompletionResponse) - self.assertEqual(response.choices[0].prompt_token_ids, MOCK_PROMPT_TOKEN_IDS) - - def test_without_return_prompt_token_ids(self): - req = ChatCompletionRequest( - model="x", - messages=[{"role": "user", "content": "Hi"}], - ) - ret = [self._make_ret(include_prompt_token_ids=False)] - response = self.chat._build_chat_response(req, ret, created=0) - - self.assertIsNone(response.choices[0].prompt_token_ids) - - data = json.loads(response.model_dump_json()) - self.assertNotIn("prompt_token_ids", data["choices"][0]) - - def test_json_round_trip(self): - req = ChatCompletionRequest( - model="x", - messages=[{"role": "user", "content": "Hi"}], - return_prompt_token_ids=True, - ) - ret = [self._make_ret(include_prompt_token_ids=True)] - response = self.chat._build_chat_response(req, ret, created=0) - - data = json.loads(response.model_dump_json()) - self.assertEqual(data["choices"][0]["prompt_token_ids"], MOCK_PROMPT_TOKEN_IDS) - - -# =========================================================================== -# 5. ReqState Tests -# =========================================================================== - - -@unittest.skipUnless( - _HAS_TOKENIZER_MANAGER, "tokenizer_manager import requires GPU deps" -) -class TestReturnTokenIdsReqState(unittest.TestCase): - """Test that ReqState stores prompt_token_ids correctly.""" - - def test_reqstate_default_none(self): - state = ReqState( - out_list=[], - finished=False, - event=asyncio.Event(), - obj=MagicMock(), - time_stats=MagicMock(), - ) - self.assertIsNone(state.prompt_token_ids) - - def test_reqstate_stores_prompt_token_ids(self): - state = ReqState( - out_list=[], - finished=False, - event=asyncio.Event(), - obj=MagicMock(), - time_stats=MagicMock(), - ) - state.prompt_token_ids = MOCK_PROMPT_TOKEN_IDS - self.assertEqual(state.prompt_token_ids, MOCK_PROMPT_TOKEN_IDS) - - -if __name__ == "__main__": - unittest.main() From 3f41db87272f3ccb1fc4909c1fdfbd08a0cd75ef Mon Sep 17 00:00:00 2001 From: JD-ETH Date: Wed, 20 May 2026 19:51:59 -0700 Subject: [PATCH 25/50] [9/14] [sglang-miles] P2P weight update support and fixes (#21278, #22663) --- .../sglang/srt/distributed/parallel_state.py | 180 ++++++++++- python/sglang/srt/entrypoints/engine.py | 9 + .../engine_info_bootstrap_server.py | 41 ++- python/sglang/srt/entrypoints/http_server.py | 26 ++ python/sglang/srt/managers/io_struct.py | 3 + .../sglang/srt/model_executor/model_runner.py | 70 ++++- .../srt/model_loader/parameter_mapper.py | 261 ++++++++++++++++ .../deepseek_common/deepseek_weight_loader.py | 65 ++-- python/sglang/srt/models/deepseek_v2.py | 47 +++ python/sglang/srt/models/glm4.py | 20 +- python/sglang/srt/models/glm4_moe.py | 70 ++--- python/sglang/srt/models/glm4_moe_lite.py | 43 +++ python/sglang/srt/models/glm4_moe_nextn.py | 17 + python/sglang/srt/models/llama.py | 28 +- python/sglang/srt/models/qwen2.py | 20 +- python/sglang/srt/models/qwen3.py | 32 +- python/sglang/srt/models/qwen3_moe.py | 34 +- .../test_parallelism_context_integration.py | 275 +++++++++++++++++ test/srt/models/test_params_mapping.py | 292 ++++++++++++++++++ 19 files changed, 1383 insertions(+), 150 deletions(-) create mode 100644 python/sglang/srt/model_loader/parameter_mapper.py create mode 100644 test/registered/distributed/test_parallelism_context_integration.py create mode 100644 test/srt/models/test_params_mapping.py diff --git a/python/sglang/srt/distributed/parallel_state.py b/python/sglang/srt/distributed/parallel_state.py index 33408f39e592..6675af6658e6 100644 --- a/python/sglang/srt/distributed/parallel_state.py +++ b/python/sglang/srt/distributed/parallel_state.py @@ -32,11 +32,11 @@ import weakref from collections import namedtuple from contextlib import contextmanager, nullcontext -from dataclasses import dataclass +from dataclasses import asdict, dataclass from datetime import timedelta from multiprocessing import shared_memory from typing import Any, Callable, Dict, List, Optional, Tuple, Union -from unittest.mock import patch +from unittest.mock import MagicMock, patch import torch import torch.distributed @@ -2414,3 +2414,179 @@ def monkey_patch_vllm_parallel_state(reverse: bool = False): setattr(vllm_parallel_state, "get_pp_group", get_pp_group) setattr(vllm_parallel_state, "get_tp_group", get_tp_group) setattr(vllm_parallel_state, "get_world_group", get_world_group) + + +@dataclass +class RankParallelismConfig: + """ + Complete parallelism configuration for a single inference rank. + + This configuration captures all the parallelism settings needed to recreate + a model shard outside of sglang. It supports: + - TP/PP/EP for model parallelism + - MoE-TP/Attn-TP/Attn-DP for MoE and DP attention. + """ + + tp_size: int = 1 + tp_rank: int = 0 + pp_size: int = 1 + pp_rank: int = 0 + ep_size: int = 1 + ep_rank: int = 0 + moe_tp_size: int = 1 + moe_tp_rank: int = 0 + attn_tp_size: int = 1 + attn_tp_rank: int = 0 + attn_dp_size: int = 1 + attn_dp_rank: int = 0 + attn_cp_size: int = 1 + attn_cp_rank: int = 0 + moe_dp_size: int = 1 + moe_dp_rank: int = 0 + + world_size: int = 1 + global_rank: int = 0 + local_rank: int = 0 + + def to_dict(self) -> Dict[str, Any]: + """Convert to dictionary for serialization.""" + return asdict(self) + + @classmethod + def from_dict(cls, data: Dict[str, Any]) -> "RankParallelismConfig": + """Create from dictionary, filtering unknown fields.""" + import dataclasses + + valid_fields = {f.name for f in dataclasses.fields(cls)} + filtered_data = {k: v for k, v in data.items() if k in valid_fields} + return cls(**filtered_data) + + @classmethod + def from_parallel_state(cls, local_rank: int = 0) -> "RankParallelismConfig": + """Extract current parallelism settings from the global parallel state.""" + tp_size = get_tensor_model_parallel_world_size() + tp_rank = get_tensor_model_parallel_rank() + + # Import dp_attention lazily to avoid circular imports + from sglang.srt.layers.dp_attention import ( + get_attention_cp_rank, + get_attention_cp_size, + get_attention_dp_rank, + get_attention_dp_size, + get_attention_tp_rank, + get_attention_tp_size, + ) + + return cls( + tp_size=tp_size, + tp_rank=tp_rank, + pp_size=get_pipeline_model_parallel_world_size(), + pp_rank=get_pipeline_model_parallel_rank(), + ep_size=get_moe_expert_parallel_world_size(), + ep_rank=get_moe_expert_parallel_rank(), + moe_tp_size=get_moe_tensor_parallel_world_size(), + moe_tp_rank=get_moe_tensor_parallel_rank(), + attn_tp_size=get_attention_tp_size(), + attn_tp_rank=get_attention_tp_rank(), + attn_dp_size=get_attention_dp_size(), + attn_dp_rank=get_attention_dp_rank(), + attn_cp_size=get_attention_cp_size(), + attn_cp_rank=get_attention_cp_rank(), + moe_dp_size=get_moe_data_parallel_world_size(), + moe_dp_rank=get_moe_data_parallel_rank(), + world_size=( + torch.distributed.get_world_size() + if torch.distributed.is_initialized() + else 1 + ), + global_rank=( + torch.distributed.get_rank() + if torch.distributed.is_initialized() + else 0 + ), + local_rank=local_rank, + ) + + +# Globals on parallel_state module to save/restore +_PS_GLOBALS = ("_TP", "_PP", "_MOE_EP", "_MOE_TP", "_ATTN_TP", "_ATTN_CP", "_MOE_DP") +# Globals on dp_attention module to save/restore +_DA_GLOBALS = ("_ATTN_DP_RANK", "_ATTN_DP_SIZE", "_ENABLE_DP_ATTENTION_FLAG") + + +class ParallelismContext: + """ + Context manager for creating model replicas with specific parallelism settings. + + Temporarily sets global variables to allow creating model shards outside of a + real distributed environment. + Usage: + with ParallelismContext(RankParallelismConfig.from_dict(parallelism_info)): + model = get_model(...) + """ + + def __init__(self, parallelism_config: RankParallelismConfig): + self.config = parallelism_config + self._original_globals: Dict[str, Any] = {} + + def _create_mock_group(self, world_size: int, rank_in_group: int): + """Create a mock group coordinator with all necessary properties.""" + mock_group = MagicMock() + mock_group.world_size = world_size + mock_group.rank_in_group = rank_in_group + mock_group.rank = rank_in_group + mock_group.local_rank = rank_in_group + mock_group.ranks = list(range(world_size)) + mock_group.first_rank = 0 + mock_group.last_rank = world_size - 1 + mock_group.is_first_rank = rank_in_group == 0 + mock_group.is_last_rank = rank_in_group == world_size - 1 + mock_group.next_rank = mock_group.ranks[(rank_in_group + 1) % world_size] + mock_group.prev_rank = mock_group.ranks[(rank_in_group - 1) % world_size] + return mock_group + + def __enter__(self): + conf = self.config + + from sglang.srt.distributed import parallel_state + from sglang.srt.layers import dp_attention + + # Save original globals + for name in _PS_GLOBALS: + self._original_globals[name] = getattr(parallel_state, name, None) + for name in _DA_GLOBALS: + self._original_globals[name] = getattr(dp_attention, name, None) + + # Build and set mock group objects on parallel_state + _ps_new_values = { + "_TP": self._create_mock_group(conf.tp_size, conf.tp_rank), + "_PP": self._create_mock_group(conf.pp_size, conf.pp_rank), + "_MOE_EP": self._create_mock_group(conf.ep_size, conf.ep_rank), + "_MOE_TP": self._create_mock_group(conf.moe_tp_size, conf.moe_tp_rank), + "_ATTN_TP": self._create_mock_group(conf.attn_tp_size, conf.attn_tp_rank), + "_ATTN_CP": self._create_mock_group(conf.attn_cp_size, conf.attn_cp_rank), + "_MOE_DP": self._create_mock_group(conf.moe_dp_size, conf.moe_dp_rank), + } + for name, value in _ps_new_values.items(): + setattr(parallel_state, name, value) + + # Set dp_attention scalar globals + dp_attention._ATTN_DP_RANK = conf.attn_dp_rank + dp_attention._ATTN_DP_SIZE = conf.attn_dp_size + dp_attention._ENABLE_DP_ATTENTION_FLAG = conf.attn_dp_size > 1 + + logger.info(f"[ParallelismContext] Activated: {conf}") + return self + + def __exit__(self, *args): + from sglang.srt.distributed import parallel_state + from sglang.srt.layers import dp_attention + + # Restore original globals + for name in _PS_GLOBALS: + setattr(parallel_state, name, self._original_globals.get(name)) + for name in _DA_GLOBALS: + setattr(dp_attention, name, self._original_globals.get(name)) + + logger.info("[ParallelismContext] Deactivated") + return False diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index 729a5e6fe3c6..5b7ba8fc7506 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -1095,14 +1095,23 @@ def post_process_weights( self, restore_weights_before_load: bool = False, post_process_quantization: bool = False, + post_load_weights: bool = False, ): """ Optional post-processing for updated weights (e.g., Marlin conversion). Should be called after weight update is finished. + + Args: + restore_weights_before_load: Restore weights to pre-quantization state. + post_process_quantization: Re-apply quantization post-processing. + post_load_weights: Call model.post_load_weights() for models that + need post-load decomposition (e.g., DeepSeek MLA kv_b_proj + decomposition into w_kc/w_vc tensors after RDMA weight transfer). """ obj = PostProcessWeightsReqInput( restore_weights_before_load=restore_weights_before_load, post_process_quantization=post_process_quantization, + post_load_weights=post_load_weights, ) return self.loop.run_until_complete( diff --git a/python/sglang/srt/entrypoints/engine_info_bootstrap_server.py b/python/sglang/srt/entrypoints/engine_info_bootstrap_server.py index 77de7fc7d030..88e48075a644 100644 --- a/python/sglang/srt/entrypoints/engine_info_bootstrap_server.py +++ b/python/sglang/srt/entrypoints/engine_info_bootstrap_server.py @@ -31,7 +31,8 @@ class EngineInfoBootstrapServer: accesses the collected info directly in-process; external consumers can query via HTTP GET. - Currently supports transfer engine memory registration info. + Currently supports transfer engine memory registration info and + per-rank parallelism configuration. """ def __init__(self, host: str, port: int): @@ -40,6 +41,8 @@ def __init__(self, host: str, port: int): # Storage: {tp_rank: (session_id, weights_info_dict)} self.transfer_engine_info: Dict[int, Tuple] = {} + # Storage: {tp_rank: parallelism_config_dict} + self.parallelism_config: Dict[int, dict] = {} self.lock = threading.Lock() app = FastAPI() @@ -89,6 +92,38 @@ def get_transfer_engine_info(rank: int): config = uvicorn.Config(app, host=host, port=port, log_level="warning") self._server = uvicorn.Server(config) + + @app.put("/register_parallelism_config") + def register_parallelism_config(data: dict): + try: + tp_rank = data["tp_rank"] + config = data["parallelism_config"] + + with self.lock: + self.parallelism_config[tp_rank] = config + + logger.info(f"Registered parallelism config for tp_rank={tp_rank}") + return PlainTextResponse("OK") + except Exception as e: + logger.error(f"Failed to register parallelism config: {e}") + raise HTTPException(status_code=400, detail=str(e)) + + @app.get("/get_parallelism_config") + def get_parallelism_config(rank: int): + if rank < 0: + raise HTTPException(status_code=400, detail="Invalid rank parameter") + + with self.lock: + config = self.parallelism_config.get(rank) + + if config is None: + raise HTTPException( + status_code=404, + detail=f"No parallelism config for rank {rank}", + ) + + return config + self._thread = threading.Thread( target=self._server.run, daemon=True, @@ -103,3 +138,7 @@ def close(self): def get_transfer_engine_info(self, rank: int) -> Optional[Tuple]: """Direct in-process access for co-located HTTP server (no HTTP round-trip).""" return self.transfer_engine_info.get(rank) + + def get_parallelism_config_info(self, rank: int) -> Optional[dict]: + """Direct in-process access for parallelism config (no HTTP round-trip).""" + return self.parallelism_config.get(rank) diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index b18af2f30504..6b895c223f97 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -1134,6 +1134,32 @@ async def remote_instance_transfer_engine_info(rank: int = None): ) +@app.get("/parallelism_config") +@auth_level(AuthLevel.ADMIN_OPTIONAL) +async def parallelism_config(rank: int = None): + """Get per-rank parallelism config from the bootstrap server.""" + if rank is None or rank < 0: + return ORJSONResponse( + {"error": {"message": "Missing or invalid rank parameter"}}, + status_code=HTTPStatus.BAD_REQUEST, + ) + + server_args = _global_state.tokenizer_manager.server_args + try: + + resp = requests.get( + f"{server_args.engine_info_bootstrap_url}/get_parallelism_config", + params={"rank": rank}, + timeout=5, + ) + if resp.status_code == 200: + return resp.json() + except Exception: + pass + + return Response(status_code=HTTPStatus.BAD_REQUEST) + + @app.post("/init_weights_update_group") @auth_level(AuthLevel.ADMIN_OPTIONAL) async def init_weights_update_group( diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index ca46b82108c5..de740d898751 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -1542,6 +1542,9 @@ class PostProcessWeightsReqInput(BaseReq): restore_weights_before_load: bool = False # Whether to enable quantization post-processing post_process_quantization: bool = False + # Whether to call model.post_load_weights() after weight update + # (e.g., DeepSeek MLA kv_b_proj decomposition into w_kc/w_vc tensors) + post_load_weights: bool = False @dataclass diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index fce3f9f28cb8..ee7ec4c81645 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -82,7 +82,10 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) -from sglang.srt.distributed.parallel_state import monkey_patch_vllm_parallel_state +from sglang.srt.distributed.parallel_state import ( + RankParallelismConfig, + monkey_patch_vllm_parallel_state, +) from sglang.srt.elastic_ep.elastic_ep import ( ElasticEPStateManager, join_process_groups, @@ -405,6 +408,7 @@ def __init__( self.remote_instance_transfer_engine = None self.remote_instance_transfer_engine_session_id = "" self.remote_instance_transfer_engine_weight_info = None + self.parallelism_config = None self.msprobe_debugger = None if server_args.msprobe_dump_config is not None: @@ -615,6 +619,9 @@ def initialize(self, pre_model_load_memory: float): if self.server_args.remote_instance_weight_loader_use_transfer_engine(): self.remote_instance_init_transfer_engine() + self.parallelism_config = RankParallelismConfig.from_parallel_state( + self.tp_rank + ) if not self.is_draft_worker: set_global_expert_location_metadata( @@ -676,6 +683,13 @@ def initialize(self, pre_model_load_memory: float): ) self._register_to_engine_info_bootstrap() + # Register parallelism config with the bootstrap server + if ( + self.server_args.remote_instance_weight_loader_use_transfer_engine() + and self.parallelism_config is not None + ): + self._register_parallelism_config_to_bootstrap() + # For MTP models like DeepSeek-V3 or GLM-4.5, the MTP layer(s) are used separately as draft # models for speculative decoding. In those cases, `num_nextn_predict_layers` is used to # determine the number of layers. @@ -921,6 +935,7 @@ def remote_instance_init_transfer_engine(self): "Please install mooncake for using remote instance transfer engine: pip install mooncake" ) return + self.remote_instance_transfer_engine = TransferEngine() local_ip = get_local_ip_auto() self.remote_instance_transfer_engine.initialize( @@ -1125,6 +1140,47 @@ def _build_nixl_worker_metadata(self, p2p_pb2): return worker, len(tensors) + def _register_parallelism_config_to_bootstrap(self): + """Register parallelism config with the EngineInfoBootstrapServer via HTTP PUT.""" + import requests as http_requests + + bootstrap_url = self._get_bootstrap_url() + url = f"{bootstrap_url}/register_parallelism_config" + + payload = { + "tp_rank": self.tp_rank, + "parallelism_config": self.parallelism_config.to_dict(), + } + + try: + resp = http_requests.put(url, json=payload, timeout=5) + if resp.status_code == 200: + logger.info( + f"Registered parallelism config for tp_rank={self.tp_rank} " + f"with bootstrap server at {bootstrap_url}" + ) + else: + logger.error( + f"Failed to register parallelism config for tp_rank={self.tp_rank}: " + f"{resp.status_code}, {resp.text}" + ) + except Exception as e: + logger.error( + f"Failed to register parallelism config for tp_rank={self.tp_rank}: {e}" + ) + + def _get_bootstrap_url(self): + """Get the base URL for the EngineInfoBootstrapServer.""" + if self.server_args.dist_init_addr: + bootstrap_host = ( + NetworkAddress.parse(self.server_args.dist_init_addr).resolved().host + ) + else: + bootstrap_host = "127.0.0.1" + + bootstrap_port = self.server_args.engine_info_bootstrap_port + return NetworkAddress(bootstrap_host, bootstrap_port).to_url() + def model_specific_adjustment(self): server_args = self.server_args @@ -3647,11 +3703,21 @@ def _maybe_rebalance_after_rank_fault( return output def post_process_weights(self, recv_req): - """Run quantization-specific post-processing hooks after loading weights.""" + """ + Execute post-processing logic for model weights, such as Marlin quantization format conversion + and model-specific post_load_weights hooks (e.g., DeepSeek MLA kv_b_proj decomposition). + """ from sglang.srt.model_loader.loader import device_loading_context target_device = torch.device("cuda", torch.cuda.current_device()) + if recv_req.post_load_weights: + # Call model.post_load_weights() if available (e.g., for DeepSeek MLA + # models that need to decompose kv_b_proj.weight into w_kc/w_vc tensors + # after RDMA weight transfer) + if hasattr(self.model, "post_load_weights"): + self.model.post_load_weights() + if recv_req.restore_weights_before_load: for _, module in self.model.named_modules(): quant_method = getattr(module, "quant_method", None) diff --git a/python/sglang/srt/model_loader/parameter_mapper.py b/python/sglang/srt/model_loader/parameter_mapper.py new file mode 100644 index 000000000000..56056d570843 --- /dev/null +++ b/python/sglang/srt/model_loader/parameter_mapper.py @@ -0,0 +1,261 @@ +# Copyright 2023-2024 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Parameter mapping from HuggingFace checkpoint names to SGLang model parameters. + +This module provides utilities for translating weight names between HuggingFace +checkpoint format and SGLang's internal parameter naming, handling: + +1. Stacked Parameter Fusion + - gate_proj + up_proj → gate_up_proj (num_shards=2) + - q_proj + k_proj + v_proj → qkv_proj (num_shards=3) + - q_a_proj + kv_a_proj_with_mqa → fused_qkv_a_proj_with_mqa (DeepSeek MLA) + +2. Expert Parameter Sharding (MoE models) + - experts.{id}.gate_proj + experts.{id}.up_proj → experts.w13_weight (num_shards=2) + - experts.{id}.down_proj → experts.w2_weight (num_shards=1) + - Handles expert parallelism: num_local_experts = n_routed // ep_size + shared + +3. Scale Remapping (Quantized models) + - k_proj.k_scale → attn.k_scale + - v_proj.v_scale → attn.v_scale + - Quark-specific: output_scale → per-component scales + +Supported Models: + Dense: Llama, Qwen2, Qwen3, GLM4 + MoE: DeepSeekV2/V3/R1, Qwen3-MoE, GLM4-MoE, GLM4-MoE-Lite (GLM-4.7) + +Example: + >>> mapper = ParameterMapper.from_model(model) + >>> result = mapper.map("model.layers.0.mlp.gate_proj.weight") + >>> result.sglang_name # "model.layers.0.mlp.gate_up_proj.weight" + >>> result.shard_id # 0 + >>> result.num_shards # 2 +""" + +from dataclasses import dataclass +from typing import Callable, Dict, List, Optional, Tuple, Union + +StackedParamsEntry = Tuple[str, str, Union[int, str]] +ExpertParamsEntry = Tuple[str, str, int, Union[int, str]] + + +@dataclass +class MappingResult: + """Result of mapping a HuggingFace checkpoint weight name to SGLang parameter.""" + + sglang_name: str + shard_id: Optional[Union[int, str]] + num_shards: int + expert_id: Optional[int] + num_local_experts: Optional[int] + + +# Standard FP8 scale remapping patterns +_SCALE_REMAP_PATTERNS: List[Tuple[str, str, str]] = [ + (".k_scale", ".self_attn.k_proj.k_scale", ".self_attn.attn.k_scale"), + (".v_scale", ".self_attn.v_proj.v_scale", ".self_attn.attn.v_scale"), + (".k_scale", ".k_scale", ".attn.k_scale"), + (".v_scale", ".v_scale", ".attn.v_scale"), +] + +# Quark quantization scale remapping +_QUARK_SCALE_REMAP: Dict[str, str] = { + ".q_proj.output_scale": ".attn.q_scale", + ".k_proj.output_scale": ".attn.k_scale", + ".v_proj.output_scale": ".attn.v_scale", + "self_attn.prob_output_scale": ".attn.prob_scale", +} + + +class ParameterMapper: + """Maps HuggingFace checkpoint weight names to SGLang model parameters. + + This class pre-computes lookup tables at initialization for efficient + repeated mapping. It handles: + - Stacked/fused parameter mapping (gate_up_proj, qkv_proj, etc.) + - Expert parameter mapping with shard information + - Scale remapping for quantized models + - Model-specific weight name mutations + """ + + def __init__( + self, + stacked_params_mapping: List[StackedParamsEntry], + expert_params_mapping: List[ExpertParamsEntry], + num_local_experts: int = 0, + mutate_weight_preload: Optional[Callable[[str], str]] = None, + custom_scale_remap: Optional[Callable[[str], str]] = None, + ): + """Initialize the parameter mapper with model-specific configuration. + + Args: + stacked_params_mapping: List of (sglang_name, hf_name, shard_id) tuples. + Example: [("gate_up_proj", "gate_proj", 0), ("gate_up_proj", "up_proj", 1)] + expert_params_mapping: List of (sglang_name, hf_name, expert_id, shard_id) tuples. + Example: [("w13_weight", "experts.0.gate_proj.weight", 0, 0), ...] + num_local_experts: Number of experts in the current model rank. + For EP=1: num_local_experts = n_routed_experts + num_fused_shared_experts + For EP>1: num_local_experts = n_routed_experts // ep_size + num_fused_shared_experts + mutate_weight_preload: Optional function to transform weight names before mapping. + Used for shared expert fusion in DeepSeek (shared_experts → experts.{n_routed}). + custom_scale_remap: Optional function for model-specific scale remapping. + Used for DeepSeek k_proj/v_proj → attn_mqa scale mapping. + """ + self.num_local_experts = num_local_experts + self._mutate_weight_preload = mutate_weight_preload + self._custom_scale_remap = custom_scale_remap + + self._stacked_lookup, self._stacked_num_shards = self._build_stacked_lookup( + stacked_params_mapping + ) + self._expert_lookup, self._expert_num_shards = self._build_expert_lookup( + expert_params_mapping + ) + + @staticmethod + def _build_stacked_lookup( + mapping: List[StackedParamsEntry], + ) -> Tuple[Dict[str, Tuple[str, Union[int, str]]], Dict[str, int]]: + """Build lookup table and num_shards from stacked params mapping.""" + lookup: Dict[str, Tuple[str, Union[int, str]]] = {} + shard_counts: Dict[str, int] = {} + + for sglang_name, hf_name, shard_id in mapping: + lookup[hf_name] = (sglang_name, shard_id) + shard_counts[sglang_name] = shard_counts.get(sglang_name, 0) + 1 + + return lookup, shard_counts + + @staticmethod + def _build_expert_lookup( + mapping: List[ExpertParamsEntry], + ) -> Tuple[Dict[str, Tuple[str, int, Union[int, str]]], Dict[str, int]]: + """Build lookup table and num_shards from expert params mapping.""" + lookup: Dict[str, Tuple[str, int, Union[int, str]]] = {} + shard_counts: Dict[str, int] = {} + + for sglang_name, hf_name, expert_id, shard_id in mapping: + lookup[hf_name] = (sglang_name, expert_id, shard_id) + + for sglang_name, _, _, shard_id in mapping: + key = sglang_name + if key not in shard_counts: + unique_shards = set( + s_id for s_name, _, _, s_id in mapping if s_name == sglang_name + ) + shard_counts[key] = len(unique_shards) + + return lookup, shard_counts + + def _apply_scale_remap(self, name: str) -> str: + """Apply standard and Quark scale remapping patterns.""" + for suffix, pattern, replacement in _SCALE_REMAP_PATTERNS: + if name.endswith(suffix) and pattern in name: + return name.replace(pattern, replacement) + + for quark_suffix, replacement in _QUARK_SCALE_REMAP.items(): + if name.endswith(quark_suffix): + return name.replace(quark_suffix, replacement) + + return name + + def map(self, hf_weight_name: str) -> MappingResult: + """Map a HuggingFace checkpoint weight name to SGLang parameter info. + + Args: + hf_weight_name: The weight name from HuggingFace checkpoint. + + Returns: + MappingResult with mapped name and sharding information. + """ + name = hf_weight_name + + if self._mutate_weight_preload is not None: + name = self._mutate_weight_preload(name) + + if "scale" in name: + if self._custom_scale_remap is not None: + remapped = self._custom_scale_remap(name) + if remapped != name: + name = remapped + else: + name = self._apply_scale_remap(name) + else: + name = self._apply_scale_remap(name) + + for hf_pattern, ( + sglang_name, + expert_id, + shard_id, + ) in self._expert_lookup.items(): + if hf_pattern in name: + mapped_name = name.replace(hf_pattern, sglang_name) + return MappingResult( + sglang_name=mapped_name, + shard_id=shard_id, + num_shards=self._expert_num_shards.get(sglang_name, 1), + expert_id=expert_id, + num_local_experts=self.num_local_experts, + ) + + for hf_pattern, (sglang_name, shard_id) in self._stacked_lookup.items(): + if hf_pattern in name: + mapped_name = name.replace(hf_pattern, sglang_name) + return MappingResult( + sglang_name=mapped_name, + shard_id=shard_id, + num_shards=self._stacked_num_shards.get(sglang_name, 1), + expert_id=None, + num_local_experts=None, + ) + + return MappingResult( + sglang_name=name, + shard_id=None, + num_shards=1, + expert_id=None, + num_local_experts=None, + ) + + @classmethod + def from_model(cls, model) -> "ParameterMapper": + """Create a ParameterMapper from a model instance; currently supports + DeepseekV2ForCausalLM, Glm4ForCausalLM, Glm4MoeForCausalLM, + Glm4MoeLiteForCausalLM, LlamaForCausalLM, Qwen2ForCausalLM, + Qwen3ForCausalLM, Qwen3MoeForCausalLM.""" + stacked_mapping = list(getattr(model, "stacked_params_mapping", []) or []) + expert_mapping = list(getattr(model, "expert_params_mapping", []) or []) + + num_local_experts = 0 + if hasattr(model, "num_local_experts"): + num_local_experts = model.num_local_experts + elif expert_mapping: + expert_ids = set(entry[2] for entry in expert_mapping) + num_local_experts = len(expert_ids) + + mutate_fn = None + if hasattr(model, "mutate_weight_preload"): + mutate_fn = model.mutate_weight_preload + + scale_fn = None + if hasattr(model, "custom_scale_remap"): + scale_fn = model.custom_scale_remap + + return cls( + stacked_params_mapping=stacked_mapping, + expert_params_mapping=expert_mapping, + num_local_experts=num_local_experts, + mutate_weight_preload=mutate_fn, + custom_scale_remap=scale_fn, + ) diff --git a/python/sglang/srt/models/deepseek_common/deepseek_weight_loader.py b/python/sglang/srt/models/deepseek_common/deepseek_weight_loader.py index 754dc1bbacd0..0f9fbc2590c8 100644 --- a/python/sglang/srt/models/deepseek_common/deepseek_weight_loader.py +++ b/python/sglang/srt/models/deepseek_common/deepseek_weight_loader.py @@ -25,7 +25,6 @@ from sglang.srt.distributed.parallel_state import GroupCoordinator from sglang.srt.environ import envs from sglang.srt.layers import deep_gemm_wrapper -from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.fp8_utils import ( block_quant_dequant, @@ -101,6 +100,18 @@ class DeepseekV2WeightLoaderMixin: quant_config: Optional[QuantizationConfig] pp_group: GroupCoordinator num_fused_shared_experts: int + # Weight mapping relationships determined at model initialization time. + fuse_qkv_a_proj: bool + stacked_params_mapping: List[Tuple[str, str, int]] + expert_params_mapping: List[Tuple[str, str, int, int]] + + def mutate_weight_preload(self, name: str) -> str: + """Override in subclass for model-specific weight name mutations.""" + return name + + def custom_scale_remap(self, name: str) -> str: + """Override in subclass for model-specific scale remapping.""" + return name def do_load_weights( self, @@ -119,33 +130,7 @@ def do_load_weights( weights, NVFP4_CKPT_FP8_ATTN_QUANT_MODULES, nextn_conf ) - stacked_params_mapping = [ - # (param_name, shard_name, shard_id) - ("gate_up_proj", "gate_proj", 0), - ("gate_up_proj", "up_proj", 1), - ] - - # Params for weights, fp8 weight scales, fp8 activation scales - # (param_name, weight_name, expert_id, shard_id) - expert_params_mapping = FusedMoE.make_expert_params_mapping( - ckpt_gate_proj_name="gate_proj", - ckpt_down_proj_name="down_proj", - ckpt_up_proj_name="up_proj", - num_experts=self.config.n_routed_experts + self.num_fused_shared_experts, - ) - # Params for special naming rules in mixed-precision models, for example: - # model.layers.xx.mlp.experts.xx.w1.input_scale. For details, - # see https://huggingface.co/Barrrrry/DeepSeek-R1-W4AFP8/blob/main. - if self.quant_config and self.quant_config.get_name() == "w4afp8": - expert_params_mapping += FusedMoE.make_expert_input_scale_params_mapping( - num_experts=self.config.n_routed_experts - ) - - # Fuse q_a_proj and kv_a_proj_with_mqa along output dimension when q_lora_rank is not None - fuse_qkv_a_proj = hasattr(self.config, "q_lora_rank") and ( - self.config.q_lora_rank is not None - ) - cached_a_proj = {} if fuse_qkv_a_proj else None + cached_a_proj = {} if self.fuse_qkv_a_proj else None if self.num_fused_shared_experts > 0: assert self.num_fused_shared_experts == 1 @@ -167,11 +152,7 @@ def do_load_weights( ) ): continue - if self.num_fused_shared_experts > 0 and "mlp.shared_experts" in name: - name = name.replace( - "mlp.shared_experts", - f"mlp.experts.{self.config.n_routed_experts}", - ) + name = self.mutate_weight_preload(name) weight_names.append(name) @@ -207,7 +188,7 @@ def do_load_weights( if "rotary_emb.inv_freq" in name: continue - for param_name, weight_name, shard_id in stacked_params_mapping: + for param_name, weight_name, shard_id in self.stacked_params_mapping: # Skip non-stacked layers and experts (experts handled below). if weight_name not in name: continue @@ -221,6 +202,10 @@ def do_load_weights( # for mlp.experts[0].gate_gate_up_proj, which breaks load. if ("mlp.experts." in name) and name not in params_dict: continue + # q_a_proj / kv_a_proj_with_mqa must bypass stacked loading + # and use the cache+concat fused A-proj path below. + if param_name == "fused_qkv_a_proj_with_mqa": + continue name = name.replace(weight_name, param_name) # Skip loading extra bias for GPTQ models. if name.endswith(".bias") and name not in params_dict: @@ -236,7 +221,7 @@ def do_load_weights( ) break else: - for mapping in expert_params_mapping: + for mapping in self.expert_params_mapping: param_name, weight_name, expert_id, shard_id = mapping if weight_name not in name: continue @@ -273,7 +258,7 @@ def do_load_weights( # Skip loading norm if not last rank in pipeline parallelism if ".norm." in name and not self.pp_group.is_last_rank: continue - if fuse_qkv_a_proj and ( + if self.fuse_qkv_a_proj and ( "q_a_proj" in name or "kv_a_proj_with_mqa" in name ): cached_a_proj[name] = _clone_if_runai_streamed_tensor( @@ -343,13 +328,7 @@ def do_load_weights( if ( "k_scale" in name or "v_scale" in name ) and name not in params_dict: - # modelopt attn kv scale is named differently - for scale in ["k_scale", "v_scale"]: - if scale in name: - name = name.replace( - f"{scale[0]}_proj", "attn_mqa" - ) - break + name = self.custom_scale_remap(name) if name not in params_dict: # modelopt ckpt contains not needed weights for MTP module: # model.decoder.self_attn.attn_mqa.v_scale and diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index adf440672488..cca41abb0db9 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -2382,6 +2382,37 @@ def __init__( self.lm_head = PPMissingLayer() self.logits_processor = LogitsProcessor(config) + # Stacked params mapping for unified weight loading API + self.stacked_params_mapping = [ + # (param_name, shard_name, shard_id) + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), + ] + # Add A-proj fusion mapping when q_lora_rank is enabled + # q_a_proj + kv_a_proj_with_mqa -> fused_qkv_a_proj_with_mqa + if self.fuse_qkv_a_proj: + self.stacked_params_mapping.extend( + [ + ("fused_qkv_a_proj_with_mqa", "q_a_proj", 0), + ("fused_qkv_a_proj_with_mqa", "kv_a_proj_with_mqa", 1), + ] + ) + self.expert_params_mapping = FusedMoE.make_expert_params_mapping( + ckpt_gate_proj_name="gate_proj", + ckpt_down_proj_name="down_proj", + ckpt_up_proj_name="up_proj", + num_experts=self.config.n_routed_experts + self.num_fused_shared_experts, + ) + # Params for special naming rules in mixed-precision models, for example: + # model.layers.xx.mlp.experts.xx.w1.input_scale. For details, + # see https://huggingface.co/Barrrrry/DeepSeek-R1-W4AFP8/blob/main. + if self.quant_config and self.quant_config.get_name() == "w4afp8": + self.expert_params_mapping += ( + FusedMoE.make_expert_input_scale_params_mapping( + num_experts=self.config.n_routed_experts + ) + ) + self._routed_experts_weights_of_layer = LazyValue( lambda: { layer_id: layer.mlp.get_moe_weights() @@ -2401,6 +2432,22 @@ def __init__( q_lora_rank = config.q_lora_rank if hasattr(config, "q_lora_rank") else None get_attn_tp_context().init_context(q_lora_rank, is_deepseek_nsa(config)) + def mutate_weight_preload(self, name: str) -> str: + """DeepSeek V2: shared expert fusion.""" + if self.num_fused_shared_experts > 0 and "mlp.shared_experts" in name: + return name.replace( + "mlp.shared_experts", + f"mlp.experts.{self.config.n_routed_experts}", + ) + return name + + def custom_scale_remap(self, name: str) -> str: + """DeepSeek V2: k_proj -> attn_mqa when k_scale in name, v_proj -> attn_mqa when v_scale in name.""" + for scale in ["k_scale", "v_scale"]: + if scale in name: + return name.replace(f"{scale[0]}_proj", "attn_mqa") + return name + @property def routed_experts_weights_of_layer(self): return self._routed_experts_weights_of_layer.value diff --git a/python/sglang/srt/models/glm4.py b/python/sglang/srt/models/glm4.py index 016941b4b6c6..c9dd59f9cec1 100644 --- a/python/sglang/srt/models/glm4.py +++ b/python/sglang/srt/models/glm4.py @@ -463,6 +463,17 @@ def __init__( self.logits_processor = LogitsProcessor(config) self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True) + + # Stacked params mapping for unified weight loading API + self.stacked_params_mapping = [ + # (param_name, shard_name, shard_id) + (".qkv_proj", ".q_proj", "q"), + (".qkv_proj", ".k_proj", "k"), + (".qkv_proj", ".v_proj", "v"), + (".gate_up_proj", ".gate_proj", 0), + (".gate_up_proj", ".up_proj", 1), + ] + # For EAGLE3 support self.capture_aux_hidden_states = False @@ -557,14 +568,7 @@ def end_layer(self): return self.model.end_layer def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): - stacked_params_mapping = [ - # (param_name, shard_name, shard_id) - (".qkv_proj", ".q_proj", "q"), - (".qkv_proj", ".k_proj", "k"), - (".qkv_proj", ".v_proj", "v"), - (".gate_up_proj", ".up_proj", 1), - (".gate_up_proj", ".gate_proj", 0), - ] + stacked_params_mapping = self.stacked_params_mapping params_dict = dict(self.named_parameters()) for name, loaded_weight in weights: diff --git a/python/sglang/srt/models/glm4_moe.py b/python/sglang/srt/models/glm4_moe.py index 155173731a2b..0bc11ed3d53d 100644 --- a/python/sglang/srt/models/glm4_moe.py +++ b/python/sglang/srt/models/glm4_moe.py @@ -15,7 +15,6 @@ """Inference-only GLM-4.5, GLM-4.6 and GLM-4.7 model compatible with HuggingFace weights""" import logging -import re from typing import Any, Dict, Iterable, List, Optional, Tuple, Union import torch @@ -1196,9 +1195,37 @@ def __init__( ) self.logits_processor = LogitsProcessor(config) + # Stacked params mapping for unified weight loading API + self.stacked_params_mapping = [ + # (param_name, shard_name, shard_id) + ("qkv_proj", "q_proj", "q"), + ("qkv_proj", "k_proj", "k"), + ("qkv_proj", "v_proj", "v"), + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), + ] + self.expert_params_mapping = FusedMoE.make_expert_params_mapping( + ckpt_gate_proj_name="gate_proj", + ckpt_down_proj_name="down_proj", + ckpt_up_proj_name="up_proj", + num_experts=self.config.n_routed_experts + self.num_fused_shared_experts, + ) + # For EAGLE3 support self.capture_aux_hidden_states = False + def mutate_weight_preload(self, name: str) -> str: + """GLM4-MoE: shared expert fusion.""" + if self.num_fused_shared_experts > 0 and "mlp.shared_experts" in name: + return name.replace( + "mlp.shared_experts", + f"mlp.experts.{self.config.n_routed_experts}", + ) + return name + + def get_input_embeddings(self) -> nn.Embedding: + return self.model.embed_tokens + def determine_num_fused_shared_experts(self): if get_global_server_args().disable_shared_experts_fusion: return @@ -1286,43 +1313,8 @@ def load_weights( else: raise ValueError("num_nextn_predict_layers is not in the config") - stacked_params_mapping = [ - # (param_name, shard_name, shard_id) - ("qkv_proj", "q_proj", "q"), - ("qkv_proj", "k_proj", "k"), - ("qkv_proj", "v_proj", "v"), - ("gate_up_proj", "gate_proj", 0), - ("gate_up_proj", "up_proj", 1), - ] - - if self.num_fused_shared_experts > 0: - assert self.num_fused_shared_experts == 1 - - def iter_weights_with_fused_shared_experts( - weights: Iterable[Tuple[str, torch.Tensor]], - ) -> Iterable[Tuple[str, torch.Tensor]]: - - pattern = re.compile( - r"^model\.layers\.(\d+)\.mlp\.shared_experts\.(.+)$" - ) - for name, weight in weights: - match = pattern.match(name) - if match: - layer_id = int(match.group(1)) - suffix = match.group(2) - name = f"model.layers.{layer_id}.mlp.experts.{self.config.n_routed_experts}.{suffix}" - yield name, weight - - weights = iter_weights_with_fused_shared_experts(weights) - - # Params for weights, fp8 weight scales, fp8 activation scales - # (param_name, weight_name, expert_id, shard_id) - expert_params_mapping = FusedMoE.make_expert_params_mapping( - ckpt_gate_proj_name="gate_proj", - ckpt_down_proj_name="down_proj", - ckpt_up_proj_name="up_proj", - num_experts=self.config.n_routed_experts + self.num_fused_shared_experts, - ) + stacked_params_mapping = self.stacked_params_mapping + expert_params_mapping = self.expert_params_mapping if is_nextn: nextn_layer_prefix = f"model.layers.{nextn_layer_id}" @@ -1343,6 +1335,8 @@ def iter_weights_with_fused_shared_experts( for name, loaded_weight in weights: weight_names.append(name) + name = self.mutate_weight_preload(name) + if not is_nextn: if hasattr(self.config, "num_nextn_predict_layers"): num_nextn_layers = self.config.num_nextn_predict_layers diff --git a/python/sglang/srt/models/glm4_moe_lite.py b/python/sglang/srt/models/glm4_moe_lite.py index 80a0351628ab..e0b06f6d78bd 100644 --- a/python/sglang/srt/models/glm4_moe_lite.py +++ b/python/sglang/srt/models/glm4_moe_lite.py @@ -503,6 +503,33 @@ def __init__( ) self.capture_aux_hidden_states = False + # Weight loading mappings for ParameterMapper compatibility + self.stacked_params_mapping = [ + # (param_name, shard_name, shard_id) + ("qkv_proj", "q_proj", "q"), + ("qkv_proj", "k_proj", "k"), + ("qkv_proj", "v_proj", "v"), + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), + ] + # Add A-proj fusion mapping when q_lora_rank is enabled (MLA) + self.fuse_qkv_a_proj = hasattr(config, "q_lora_rank") and ( + config.q_lora_rank is not None + ) + if self.fuse_qkv_a_proj: + self.stacked_params_mapping.extend( + [ + ("fused_qkv_a_proj_with_mqa", "q_a_proj", 0), + ("fused_qkv_a_proj_with_mqa", "kv_a_proj_with_mqa", 1), + ] + ) + self.expert_params_mapping = FusedMoE.make_expert_params_mapping( + ckpt_gate_proj_name="gate_proj", + ckpt_down_proj_name="down_proj", + ckpt_up_proj_name="up_proj", + num_experts=config.n_routed_experts + self.num_fused_shared_experts, + ) + self.nsa_enable_prefill_cp = is_nsa_enable_prefill_cp() if self.nsa_enable_prefill_cp: self.cp_rank = get_attention_tp_rank() @@ -539,6 +566,22 @@ def determine_num_fused_shared_experts( self.num_fused_shared_experts = self.config.n_shared_experts + def mutate_weight_preload(self, name: str) -> str: + """GLM4-MoE-Lite: shared expert fusion.""" + if self.num_fused_shared_experts > 0 and "mlp.shared_experts" in name: + return name.replace( + "mlp.shared_experts", + f"mlp.experts.{self.config.n_routed_experts}", + ) + return name + + def custom_scale_remap(self, name: str) -> str: + """GLM4-MoE-Lite: k_proj/v_proj -> attn_mqa for MLA kv scale.""" + for s in ["k_scale", "v_scale"]: + if s in name: + return name.replace(f"{s[0]}_proj", "attn_mqa") + return name + def load_weights( self, weights: Iterable[Tuple[str, torch.Tensor]], diff --git a/python/sglang/srt/models/glm4_moe_nextn.py b/python/sglang/srt/models/glm4_moe_nextn.py index dfbd4583dbd2..f57c90d66410 100644 --- a/python/sglang/srt/models/glm4_moe_nextn.py +++ b/python/sglang/srt/models/glm4_moe_nextn.py @@ -28,6 +28,7 @@ from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.logits_processor import LogitsProcessor +from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.vocab_parallel_embedding import ( ParallelLMHead, @@ -149,6 +150,22 @@ def __init__( 0 if get_global_server_args().disable_shared_experts_fusion else 1 ) + # Weight loading mappings (must match parent for load_weights to work) + self.stacked_params_mapping = [ + # (param_name, shard_name, shard_id) + ("qkv_proj", "q_proj", "q"), + ("qkv_proj", "k_proj", "k"), + ("qkv_proj", "v_proj", "v"), + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), + ] + self.expert_params_mapping = FusedMoE.make_expert_params_mapping( + ckpt_gate_proj_name="gate_proj", + ckpt_down_proj_name="down_proj", + ckpt_up_proj_name="up_proj", + num_experts=self.config.n_routed_experts + self.num_fused_shared_experts, + ) + @torch.no_grad() def forward( self, diff --git a/python/sglang/srt/models/llama.py b/python/sglang/srt/models/llama.py index 9fc16874d082..a369c39d4f0d 100644 --- a/python/sglang/srt/models/llama.py +++ b/python/sglang/srt/models/llama.py @@ -500,9 +500,21 @@ def __init__( (".gate_up_proj", ".gate_proj", 0), (".gate_up_proj", ".up_proj", 1), ] + # Llama-specific scale remapping patterns (suffix, pattern, replacement) + self._llama_scale_remap_patterns = [ + (".activation_scale", ".activation_scale", ".input_scale"), + (".weight_scale_inv", ".weight_scale_inv", ".weight_scale"), + ] self.capture_aux_hidden_states = False + def custom_scale_remap(self, name: str) -> str: + """Llama: activation_scale->input_scale, weight_scale_inv->weight_scale.""" + for suffix, pattern, replacement in self._llama_scale_remap_patterns: + if name.endswith(suffix) and pattern in name: + return name.replace(pattern, replacement) + return name + def _init_model( self, config: LlamaConfig, @@ -613,23 +625,11 @@ def get_num_params(self): return len(params_dict) def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): - stacked_params_mapping = [ - # (param_name, shard_name, shard_id) - (".qkv_proj", ".q_proj", "q"), - (".qkv_proj", ".k_proj", "k"), - (".qkv_proj", ".v_proj", "v"), - (".gate_up_proj", ".gate_proj", 0), - (".gate_up_proj", ".up_proj", 1), - ] params_dict = dict(self.named_parameters()) for name, loaded_weight in weights: - if name.endswith(".activation_scale"): - name = name.replace(".activation_scale", ".input_scale") - if name.endswith(".weight_scale_inv"): - name = name.replace(".weight_scale_inv", ".weight_scale") - + name = self.custom_scale_remap(name) layer_id = get_layer_id(name) if ( layer_id is not None @@ -656,7 +656,7 @@ def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): if name is None: continue - for param_name, weight_name, shard_id in stacked_params_mapping: + for param_name, weight_name, shard_id in self.stacked_params_mapping: if weight_name not in name: continue name = name.replace(weight_name, param_name) diff --git a/python/sglang/srt/models/qwen2.py b/python/sglang/srt/models/qwen2.py index 39e404884d55..fe1ec95522c5 100644 --- a/python/sglang/srt/models/qwen2.py +++ b/python/sglang/srt/models/qwen2.py @@ -464,6 +464,17 @@ def __init__( self.logits_processor = LogitsProcessor(config) self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True) + + # Stacked params mapping for unified weight loading API + self.stacked_params_mapping = [ + # (param_name, shard_name, shard_id) + ("qkv_proj", "q_proj", "q"), + ("qkv_proj", "k_proj", "k"), + ("qkv_proj", "v_proj", "v"), + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), + ] + # For EAGLE3 support self.capture_aux_hidden_states = False @@ -558,14 +569,7 @@ def end_layer(self): return self.model.end_layer def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): - stacked_params_mapping = [ - # (param_name, shard_name, shard_id) - ("qkv_proj", "q_proj", "q"), - ("qkv_proj", "k_proj", "k"), - ("qkv_proj", "v_proj", "v"), - ("gate_up_proj", "gate_proj", 0), - ("gate_up_proj", "up_proj", 1), - ] + stacked_params_mapping = self.stacked_params_mapping params_dict = dict(self.named_parameters()) for name, loaded_weight in weights: diff --git a/python/sglang/srt/models/qwen3.py b/python/sglang/srt/models/qwen3.py index be8f747caf25..64c18883c445 100644 --- a/python/sglang/srt/models/qwen3.py +++ b/python/sglang/srt/models/qwen3.py @@ -32,6 +32,7 @@ from sglang.srt.models.utils import apply_qk_norm from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import add_prefix, get_bool_env_var, is_cuda, is_hip, is_npu +from sglang.srt.utils.hf_transformers_utils import get_rope_config Qwen3Config = None @@ -316,16 +317,7 @@ def __init__( ) -> None: super().__init__() self.hidden_size = config.hidden_size - if ( - hasattr(config, "rope_parameters") - and config.rope_parameters - and "rope_theta" in config.rope_parameters - ): - rope_theta = config.rope_parameters["rope_theta"] - rope_scaling = config.rope_parameters - else: - rope_theta = getattr(config, "rope_theta", 1000000) - rope_scaling = getattr(config, "rope_scaling", None) + rope_theta, rope_scaling = get_rope_config(config) max_position_embeddings = getattr(config, "max_position_embeddings", 32768) head_dim = getattr(config, "head_dim", None) self.self_attn = Qwen3Attention( @@ -477,6 +469,16 @@ def __init__( config, quant_config=quant_config, prefix=add_prefix("model", prefix) ) + # Stacked params mapping for unified weight loading API + self.stacked_params_mapping = [ + # (param_name, shard_name, shard_id) + ("qkv_proj", "q_proj", "q"), + ("qkv_proj", "k_proj", "k"), + ("qkv_proj", "v_proj", "v"), + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), + ] + # handle the lm head on different pp ranks if self.pp_group.is_last_rank: if self.pp_group.world_size == 1 and config.tie_word_embeddings: @@ -588,15 +590,7 @@ def end_layer(self): return self.model.end_layer def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): - stacked_params_mapping = [ - # (param_name, shard_name, shard_id) - ("qkv_proj", "q_proj", "q"), - ("qkv_proj", "k_proj", "k"), - ("qkv_proj", "v_proj", "v"), - ("gate_up_proj", "gate_proj", 0), - ("gate_up_proj", "up_proj", 1), - ] - + stacked_params_mapping = self.stacked_params_mapping params_dict = dict(self.named_parameters()) for name, loaded_weight in weights: if not name.startswith("model.") and ( diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py index 9fb6808678d9..66c8016fa361 100644 --- a/python/sglang/srt/models/qwen3_moe.py +++ b/python/sglang/srt/models/qwen3_moe.py @@ -1004,6 +1004,23 @@ def __init__( use_attn_tp_group=get_global_server_args().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) + + # Stacked params mapping for unified weight loading API + self.stacked_params_mapping = [ + # (param_name, shard_name, shard_id) + ("qkv_proj", "q_proj", "q"), + ("qkv_proj", "k_proj", "k"), + ("qkv_proj", "v_proj", "v"), + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), + ] + self.expert_params_mapping = FusedMoE.make_expert_params_mapping( + ckpt_gate_proj_name="gate_proj", + ckpt_down_proj_name="down_proj", + ckpt_up_proj_name="up_proj", + num_experts=self.config.num_experts, + ) + self.capture_aux_hidden_states = False self.attn_cp_size = get_attn_context_model_parallel_world_size() @@ -1140,21 +1157,8 @@ def set_dflash_layers_to_capture(self, layer_ids: List[int]): self.model.set_dflash_layers_to_capture([val + 1 for val in layer_ids]) def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): - stacked_params_mapping = [ - # (param_name, shard_name, shard_id) - ("qkv_proj", "q_proj", "q"), - ("qkv_proj", "k_proj", "k"), - ("qkv_proj", "v_proj", "v"), - ("gate_up_proj", "gate_proj", 0), - ("gate_up_proj", "up_proj", 1), - ] - - expert_params_mapping = FusedMoE.make_expert_params_mapping( - ckpt_gate_proj_name="gate_proj", - ckpt_down_proj_name="down_proj", - ckpt_up_proj_name="up_proj", - num_experts=self.config.num_experts, - ) + stacked_params_mapping = self.stacked_params_mapping + expert_params_mapping = self.expert_params_mapping # Pre-define `params_dict` to avoid repeated expensive traversal of model parameters. params_dict = dict(self.named_parameters()) diff --git a/test/registered/distributed/test_parallelism_context_integration.py b/test/registered/distributed/test_parallelism_context_integration.py new file mode 100644 index 000000000000..e4dde076fc78 --- /dev/null +++ b/test/registered/distributed/test_parallelism_context_integration.py @@ -0,0 +1,275 @@ +""" +Integration tests for ParallelismContext with real sglang servers. + +Tests that ParallelismContext can instantiate models with correct tensor parallel +sharding by comparing parameter names and sizes against a running sglang server. + +Run with: + pytest test/registered/distributed/test_parallelism_context_integration.py -v + +Full test suite (non-CI): + - TP=2 small model (Qwem2.5-1.5B-Instruct) + - EP=2 small MOE model (DeepSeek-Coder-V2-Lite-Instruct) + - MLA model with hybrid dp attention (DeepSeek-Coder-V2-Lite-Instruct) + +CI test (reduced): + - TP=2 small model only +""" + +import dataclasses +import gc +from typing import Dict, List, Tuple + +import pytest +import requests +import torch + +from sglang.srt.distributed.parallel_state import RankParallelismConfig +from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci +from sglang.test.test_utils import ( + DEFAULT_MLA_MODEL_NAME_FOR_TEST, + DEFAULT_SMALL_MODEL_NAME_FOR_TEST_QWEN, + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + is_in_ci, + popen_launch_server, +) +from sglang.utils import terminate_process + +register_cuda_ci(est_time=145, stage="extra-a", runner_config="2-gpu-large") +register_amd_ci(est_time=72, suite="stage-b-test-2-gpu-large-amd") + + +def get_transfer_engine_info(url: str, rank: int) -> Dict: + """Get transfer engine info (parameter names and sizes) for a rank.""" + response = requests.get( + f"{url}/remote_instance_transfer_engine_info", + params={"rank": rank}, + ) + response.raise_for_status() + return response.json() + + +def get_parallelism_config(url: str, rank: int) -> Dict: + """Get parallelism config for a rank.""" + response = requests.get(f"{url}/parallelism_config", params={"rank": rank}) + response.raise_for_status() + return response.json() + + +def get_server_info(url: str) -> Dict: + """Get server info.""" + response = requests.get(f"{url}/server_info") + response.raise_for_status() + return response.json() + + +def verify_model_params_match_for_rank( + url: str, + rank: int, + server_info: Dict, + test_gpu_id: int, +): + """Verify model parameters match for a specific rank by recreating a model shard.""" + transfer_info = get_transfer_engine_info(url, rank) + server_weights_info = transfer_info["remote_instance_transfer_engine_info"][1] + + # Get parallelism config from running server + parallelism_config_data = get_parallelism_config(url, rank) + parallelism_config = RankParallelismConfig.from_dict(parallelism_config_data) + # Get server args from server info + from sglang.srt.server_args import ServerArgs + + valid_fields = {f.name for f in dataclasses.fields(ServerArgs)} + filtered_info = {k: v for k, v in server_info.items() if k in valid_fields} + filtered_info.pop("model_config", None) + server_args = ServerArgs(**filtered_info) + + from sglang.srt import server_args as server_args_module + from sglang.srt.distributed.parallel_state import ParallelismContext + + original_global_server_args = server_args_module._global_server_args + + try: + # In a Mock ParallelismContext, instantiate the model for this rank. + # Use a separate GPU (test_gpu_id) to avoid memory conflicts with the running server. + server_args_module._global_server_args = server_args + with ParallelismContext(parallelism_config): + from sglang.srt.configs.device_config import DeviceConfig + from sglang.srt.configs.load_config import LoadConfig + from sglang.srt.configs.model_config import ModelConfig + from sglang.srt.model_loader import get_model + + model_config = ModelConfig.from_server_args(server_args) + load_config = LoadConfig(load_format="dummy") + device_config = DeviceConfig(device="cuda", gpu_id=test_gpu_id) + + torch.cuda.set_device(test_gpu_id) + model = get_model( + model_config=model_config, + load_config=load_config, + device_config=device_config, + ) + model_params = {} + for name, param in model.named_parameters(): + model_params[name] = param.numel() * param.element_size() + + # Verify all server parameters exist in model with same size + mismatches = [] + missing = [] + for param_name, (ptr, numel, elem_size) in server_weights_info.items(): + expected_size = numel * elem_size + if param_name not in model_params: + missing.append(param_name) + elif model_params[param_name] != expected_size: + mismatches.append( + f"{param_name}: model={model_params[param_name]}, server={expected_size}" + ) + + assert not missing, f"Rank {rank}: Missing parameters: {missing}" + assert not mismatches, f"Rank {rank}: Size mismatches: {mismatches}" + del model + torch.cuda.empty_cache() + + finally: + server_args_module._global_server_args = original_global_server_args + + +TEST_CONFIGS: List[Tuple[str, str, int, List[str], int]] = [ + # Basic TP=2 test (CI only) + ( + "tp2_small", + DEFAULT_SMALL_MODEL_NAME_FOR_TEST_QWEN, + 2, + [], + 2, + ), + # EP=2: MoE experts split across 2 groups, moe_tp=1 per group + ( + "mla_ep2", + DEFAULT_MLA_MODEL_NAME_FOR_TEST, + 2, + ["--ep-size", "2"], + 2, + ), + ( + "mla_dp2_tp4", + DEFAULT_MLA_MODEL_NAME_FOR_TEST, + 4, + ["--enable-dp-attention", "--dp", "2"], + 4, + ), + ( + "mla_dp2_ep2_tp4", + DEFAULT_MLA_MODEL_NAME_FOR_TEST, + 4, + ["--enable-dp-attention", "--dp", "2", "--ep-size", "2"], + 4, + ), + ( + "mla_dp2_ep4_tp4", + DEFAULT_MLA_MODEL_NAME_FOR_TEST, + 4, + ["--enable-dp-attention", "--dp", "2", "--ep-size", "4"], + 4, + ), + ( + "mla_dp4_ep2_tp4", + DEFAULT_MLA_MODEL_NAME_FOR_TEST, + 4, + ["--enable-dp-attention", "--dp", "4", "--ep-size", "2"], + 4, + ), +] + + +def get_test_configs(): + if is_in_ci(): + return [TEST_CONFIGS[0]] + else: + return TEST_CONFIGS + + +def _get_test_params(): + """Generate pytest parameters based on test configs.""" + configs = get_test_configs() + params = [] + ids = [] + for ( + test_id, + model_name, + tp_size, + extra_args, + min_gpus, + ) in configs: + params.append( + pytest.param( + (model_name, tp_size, extra_args, min_gpus), + id=test_id, + ) + ) + return params + + +class TestParallelismContextIntegration: + """ + Test that ParallelismContext can instantiate models with the same + parameter names and sizes as the sglang server engine. + """ + + @pytest.mark.parametrize("config", _get_test_params()) + def test_model_instantiation_matches_server(self, config): + """ + Test that a model instantiated with ParallelismContext has the same + parameter names and sizes as the model in the sglang server. + + This test: + 1. Starts a server with specified parallelism config + 2. Gets transfer_engine_info for all ranks (contains param names and sizes) + 3. Gets parallelism_config and server_info + 4. Uses ParallelismContext to instantiate a model for each rank + 5. Compares the parameter names and sizes + """ + model_name, tp_size, extra_args, min_gpus = config + url = DEFAULT_URL_FOR_TEST + + # Need min_gpus for server + 1 extra GPU for test model instantiation + required_gpus = min_gpus + 1 + if torch.cuda.device_count() < required_gpus: + pytest.skip( + f"Need at least {required_gpus} GPUs (server={min_gpus} + test=1), have {torch.cuda.device_count()}" + ) + test_gpu_id = min_gpus # e.g., if server uses 0-1, test uses 2 + + # Build server args + other_args = [ + "--tp-size", + str(tp_size), + "--remote-instance-weight-loader-start-seed-via-transfer-engine", + "--trust-remote-code", + ] + other_args.extend(extra_args) + + process = None + try: + process = popen_launch_server( + model_name, + url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=other_args, + ) + server_info = get_server_info(url) + + for rank in range(tp_size): + verify_model_params_match_for_rank(url, rank, server_info, test_gpu_id) + + finally: + if process is not None: + terminate_process(process) + gc.collect() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + + +if __name__ == "__main__": + pytest.main([__file__, "-v"]) diff --git a/test/srt/models/test_params_mapping.py b/test/srt/models/test_params_mapping.py new file mode 100644 index 000000000000..e064266cab6a --- /dev/null +++ b/test/srt/models/test_params_mapping.py @@ -0,0 +1,292 @@ +"""Unit tests for ParameterMapper.""" + +from types import SimpleNamespace + +import pytest + +from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE +from sglang.srt.model_loader.parameter_mapper import ParameterMapper + +_DEEPSEEK_N_ROUTED = 4 +_DEEPSEEK_N_LOCAL = _DEEPSEEK_N_ROUTED + 1 # +1 fused shared expert +_QWEN3MOE_N = 4 +_GLM4LITE_N_ROUTED = 4 +_GLM4LITE_N_LOCAL = _GLM4LITE_N_ROUTED + 1 # +1 fused shared expert + + +def _make_model(**kwargs): + """Create a stub model object for ParameterMapper.from_model().""" + return SimpleNamespace(**kwargs) + + +def _deepseek_mutate(name): + if "mlp.shared_experts" in name: + return name.replace("mlp.shared_experts", f"mlp.experts.{_DEEPSEEK_N_ROUTED}") + return name + + +def _deepseek_scale_remap(name): + for s in ["k_scale", "v_scale"]: + if s in name: + return name.replace(f"{s[0]}_proj", "attn_mqa") + return name + + +def _glm4lite_mutate(name): + if "mlp.shared_experts" in name: + return name.replace("mlp.shared_experts", f"mlp.experts.{_GLM4LITE_N_ROUTED}") + return name + + +_LLAMA_SCALE_PATTERNS = [ + (".activation_scale", ".activation_scale", ".input_scale"), + (".weight_scale_inv", ".weight_scale_inv", ".weight_scale"), +] + + +def _llama_scale_remap(name): + for suffix, pattern, replacement in _LLAMA_SCALE_PATTERNS: + if name.endswith(suffix) and pattern in name: + return name.replace(pattern, replacement) + return name + + +@pytest.fixture +def qwen_mapper(): + """Qwen2/Qwen3 (dense): QKV fusion, gate/up fusion, no experts.""" + return ParameterMapper.from_model( + _make_model( + stacked_params_mapping=[ + ("qkv_proj", "q_proj", "q"), + ("qkv_proj", "k_proj", "k"), + ("qkv_proj", "v_proj", "v"), + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), + ], + ) + ) + + +@pytest.fixture +def llama_mapper(): + """Llama/GLM4 (dense): dot-prefixed stacked params, custom scale remap.""" + return ParameterMapper.from_model( + _make_model( + stacked_params_mapping=[ + (".qkv_proj", ".q_proj", "q"), + (".qkv_proj", ".k_proj", "k"), + (".qkv_proj", ".v_proj", "v"), + (".gate_up_proj", ".gate_proj", 0), + (".gate_up_proj", ".up_proj", 1), + ], + custom_scale_remap=_llama_scale_remap, + ) + ) + + +@pytest.fixture +def qwen3moe_mapper(): + """Qwen3-MoE: QKV fusion + experts, no shared expert fusion.""" + return ParameterMapper.from_model( + _make_model( + stacked_params_mapping=[ + ("qkv_proj", "q_proj", "q"), + ("qkv_proj", "k_proj", "k"), + ("qkv_proj", "v_proj", "v"), + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), + ], + expert_params_mapping=FusedMoE.make_expert_params_mapping( + ckpt_gate_proj_name="gate_proj", + ckpt_down_proj_name="down_proj", + ckpt_up_proj_name="up_proj", + num_experts=_QWEN3MOE_N, + ), + ) + ) + + +@pytest.fixture +def deepseek_mapper(): + """DeepSeek V2/V3: MLA A-proj fusion, shared expert fusion, custom scale remap.""" + return ParameterMapper.from_model( + _make_model( + stacked_params_mapping=[ + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), + ("fused_qkv_a_proj_with_mqa", "q_a_proj", 0), + ("fused_qkv_a_proj_with_mqa", "kv_a_proj_with_mqa", 1), + ], + expert_params_mapping=FusedMoE.make_expert_params_mapping( + ckpt_gate_proj_name="gate_proj", + ckpt_down_proj_name="down_proj", + ckpt_up_proj_name="up_proj", + num_experts=_DEEPSEEK_N_LOCAL, + ), + mutate_weight_preload=_deepseek_mutate, + custom_scale_remap=_deepseek_scale_remap, + ) + ) + + +@pytest.fixture +def glm4lite_mapper(): + """GLM4-MoE-Lite (GLM-4.7): QKV fusion, MLA A-proj fusion, shared expert fusion, custom scale remap.""" + return ParameterMapper.from_model( + _make_model( + stacked_params_mapping=[ + ("qkv_proj", "q_proj", "q"), + ("qkv_proj", "k_proj", "k"), + ("qkv_proj", "v_proj", "v"), + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), + ("fused_qkv_a_proj_with_mqa", "q_a_proj", 0), + ("fused_qkv_a_proj_with_mqa", "kv_a_proj_with_mqa", 1), + ], + expert_params_mapping=FusedMoE.make_expert_params_mapping( + ckpt_gate_proj_name="gate_proj", + ckpt_down_proj_name="down_proj", + ckpt_up_proj_name="up_proj", + num_experts=_GLM4LITE_N_LOCAL, + ), + mutate_weight_preload=_glm4lite_mutate, + custom_scale_remap=_deepseek_scale_remap, + ) + ) + + +# ── Helpers ────────────────────────────────────────────────────────────────── + + +def to_expect(name, shard=None, n=1, expert=None, n_exp=None): + """Shorthand for expected MappingResult fields.""" + return (name, shard, n, expert, n_exp) + + +def _assert(mapper, ckpt, expected): + r = mapper.map(ckpt) + name, shard, n, expert, n_exp = expected + assert ( + r.sglang_name, + r.shard_id, + r.num_shards, + r.expert_id, + r.num_local_experts, + ) == (name, shard, n, expert, n_exp), f"map({ckpt!r}) = {r}" + + +# ── Tests ──────────────────────────────────────────────────────────────────── + +# fmt: off +_QWEN_CASES = [ + # QKV fusion (Qwen2, Qwen3, GLM4-MoE) + ("layers.0.attn.q_proj.weight", to_expect("layers.0.attn.qkv_proj.weight", "q", 3)), + ("layers.0.attn.k_proj.weight", to_expect("layers.0.attn.qkv_proj.weight", "k", 3)), + ("layers.0.attn.v_proj.weight", to_expect("layers.0.attn.qkv_proj.weight", "v", 3)), + # Gate/Up fusion + ("layers.0.mlp.gate_proj.weight", to_expect("layers.0.mlp.gate_up_proj.weight", 0, 2)), + ("layers.0.mlp.up_proj.weight", to_expect("layers.0.mlp.gate_up_proj.weight", 1, 2)), + # Pass-through + ("layers.0.mlp.down_proj.weight", to_expect("layers.0.mlp.down_proj.weight")), + ("embed_tokens.weight", to_expect("embed_tokens.weight")), + # Standard scale remap (no custom_scale_remap) + ("model.layers.0.self_attn.k_scale", to_expect("model.layers.0.self_attn.attn.k_scale")), + ("model.layers.0.self_attn.v_scale", to_expect("model.layers.0.self_attn.attn.v_scale")), +] + +_LLAMA_CASES = [ + # Dot-prefixed QKV fusion (Llama, GLM4) + ("model.layers.0.self_attn.q_proj.weight", to_expect("model.layers.0.self_attn.qkv_proj.weight", "q", 3)), + ("model.layers.0.self_attn.k_proj.weight", to_expect("model.layers.0.self_attn.qkv_proj.weight", "k", 3)), + # Dot-prefixed gate/up + ("model.layers.0.mlp.gate_proj.weight", to_expect("model.layers.0.mlp.gate_up_proj.weight", 0, 2)), + # Llama-specific scale remap + stacked (scales follow their weights) + ("model.layers.0.mlp.gate_proj.activation_scale", to_expect("model.layers.0.mlp.gate_up_proj.input_scale", 0, 2)), + ("model.layers.0.mlp.gate_proj.weight_scale_inv", to_expect("model.layers.0.mlp.gate_up_proj.weight_scale", 0, 2)), + # Pass-through + ("model.layers.0.mlp.down_proj.weight", to_expect("model.layers.0.mlp.down_proj.weight")), +] + +_QWEN3MOE_CASES = [ + # QKV fusion + ("model.layers.0.self_attn.q_proj.weight", to_expect("model.layers.0.self_attn.qkv_proj.weight", "q", 3)), + # Expert mapping (no shared expert fusion) + ("model.layers.0.mlp.experts.0.gate_proj.weight", to_expect("model.layers.0.mlp.experts.w13_weight", "w1", 2, 0, _QWEN3MOE_N)), + ("model.layers.0.mlp.experts.3.down_proj.weight", to_expect("model.layers.0.mlp.experts.w2_weight", "w2", 1, 3, _QWEN3MOE_N)), + # shared_experts falls through to stacked mapping (no mutate_weight_preload) + ("model.layers.0.mlp.shared_experts.gate_proj.weight", to_expect("model.layers.0.mlp.shared_experts.gate_up_proj.weight", 0, 2)), +] + +_DEEPSEEK_CASES = [ + # MLA A-proj fusion + ("model.layers.0.self_attn.q_a_proj.weight", to_expect("model.layers.0.self_attn.fused_qkv_a_proj_with_mqa.weight", 0, 2)), + ("model.layers.0.self_attn.kv_a_proj_with_mqa.weight", to_expect("model.layers.0.self_attn.fused_qkv_a_proj_with_mqa.weight", 1, 2)), + # Shared expert fusion via mutate_weight_preload + ("model.layers.0.mlp.shared_experts.gate_proj.weight", to_expect("model.layers.0.mlp.experts.w13_weight", "w1", 2, _DEEPSEEK_N_ROUTED, _DEEPSEEK_N_LOCAL)), + ("model.layers.0.mlp.shared_experts.down_proj.weight", to_expect("model.layers.0.mlp.experts.w2_weight", "w2", 1, _DEEPSEEK_N_ROUTED, _DEEPSEEK_N_LOCAL)), + # Custom scale remap (k_proj/v_proj -> attn_mqa, NOT double-remapped) + ("model.layers.0.self_attn.k_proj.k_scale", to_expect("model.layers.0.self_attn.attn_mqa.k_scale")), + ("model.layers.0.self_attn.v_proj.v_scale", to_expect("model.layers.0.self_attn.attn_mqa.v_scale")), + # kv_b_proj pass-through (decomposed in post_load_weights) + ("model.layers.0.self_attn.kv_b_proj.weight", to_expect("model.layers.0.self_attn.kv_b_proj.weight")), +] + +_GLM4LITE_CASES = [ + # QKV fusion (GLM-4.7 uses standard QKV unlike DeepSeek which uses MLA-only) + ("model.layers.0.self_attn.q_proj.weight", to_expect("model.layers.0.self_attn.qkv_proj.weight", "q", 3)), + ("model.layers.0.self_attn.k_proj.weight", to_expect("model.layers.0.self_attn.qkv_proj.weight", "k", 3)), + ("model.layers.0.self_attn.v_proj.weight", to_expect("model.layers.0.self_attn.qkv_proj.weight", "v", 3)), + # MLA A-proj fusion (GLM-4.7 also uses MLA with q_lora_rank) + ("model.layers.0.self_attn.q_a_proj.weight", to_expect("model.layers.0.self_attn.fused_qkv_a_proj_with_mqa.weight", 0, 2)), + ("model.layers.0.self_attn.kv_a_proj_with_mqa.weight", to_expect("model.layers.0.self_attn.fused_qkv_a_proj_with_mqa.weight", 1, 2)), + # Gate/Up fusion (non-expert layers) + ("model.layers.0.mlp.gate_proj.weight", to_expect("model.layers.0.mlp.gate_up_proj.weight", 0, 2)), + ("model.layers.0.mlp.up_proj.weight", to_expect("model.layers.0.mlp.gate_up_proj.weight", 1, 2)), + # Expert mapping + ("model.layers.0.mlp.experts.0.gate_proj.weight", to_expect("model.layers.0.mlp.experts.w13_weight", "w1", 2, 0, _GLM4LITE_N_LOCAL)), + ("model.layers.0.mlp.experts.0.up_proj.weight", to_expect("model.layers.0.mlp.experts.w13_weight", "w3", 2, 0, _GLM4LITE_N_LOCAL)), + ("model.layers.0.mlp.experts.3.down_proj.weight", to_expect("model.layers.0.mlp.experts.w2_weight", "w2", 1, 3, _GLM4LITE_N_LOCAL)), + # Shared expert fusion via mutate_weight_preload + ("model.layers.0.mlp.shared_experts.gate_proj.weight", to_expect("model.layers.0.mlp.experts.w13_weight", "w1", 2, _GLM4LITE_N_ROUTED, _GLM4LITE_N_LOCAL)), + ("model.layers.0.mlp.shared_experts.down_proj.weight", to_expect("model.layers.0.mlp.experts.w2_weight", "w2", 1, _GLM4LITE_N_ROUTED, _GLM4LITE_N_LOCAL)), + # Custom scale remap (same as DeepSeek: k_proj/v_proj -> attn_mqa) + ("model.layers.0.self_attn.k_proj.k_scale", to_expect("model.layers.0.self_attn.attn_mqa.k_scale")), + ("model.layers.0.self_attn.v_proj.v_scale", to_expect("model.layers.0.self_attn.attn_mqa.v_scale")), + # Pass-through + ("model.layers.0.mlp.down_proj.weight", to_expect("model.layers.0.mlp.down_proj.weight")), + ("model.layers.0.self_attn.kv_b_proj.weight", to_expect("model.layers.0.self_attn.kv_b_proj.weight")), +] +# fmt: on + + +@pytest.mark.parametrize("ckpt,expected", _QWEN_CASES, ids=[c[0] for c in _QWEN_CASES]) +def test_qwen(qwen_mapper, ckpt, expected): + _assert(qwen_mapper, ckpt, expected) + + +@pytest.mark.parametrize( + "ckpt,expected", _LLAMA_CASES, ids=[c[0] for c in _LLAMA_CASES] +) +def test_llama(llama_mapper, ckpt, expected): + _assert(llama_mapper, ckpt, expected) + + +@pytest.mark.parametrize( + "ckpt,expected", _QWEN3MOE_CASES, ids=[c[0] for c in _QWEN3MOE_CASES] +) +def test_qwen3moe(qwen3moe_mapper, ckpt, expected): + _assert(qwen3moe_mapper, ckpt, expected) + + +@pytest.mark.parametrize( + "ckpt,expected", _DEEPSEEK_CASES, ids=[c[0] for c in _DEEPSEEK_CASES] +) +def test_deepseek(deepseek_mapper, ckpt, expected): + _assert(deepseek_mapper, ckpt, expected) + + +@pytest.mark.parametrize( + "ckpt,expected", _GLM4LITE_CASES, ids=[c[0] for c in _GLM4LITE_CASES] +) +def test_glm4lite(glm4lite_mapper, ckpt, expected): + _assert(glm4lite_mapper, ckpt, expected) From 024496193938eb4a824b98a44c8961ff958bb10f Mon Sep 17 00:00:00 2001 From: maocheng23 Date: Mon, 13 Apr 2026 21:42:32 -0700 Subject: [PATCH 26/50] [10/14] [sglang-miles] Fix pause-aware weight update deadlocks (#22754, #22623) --- python/sglang/srt/layers/quantization/fp8.py | 16 +++-- python/sglang/srt/managers/scheduler.py | 67 ++++++++++++++++--- .../scheduler_update_weights_mixin.py | 17 +++++ .../srt/managers/tokenizer_control_mixin.py | 12 +++- 4 files changed, 97 insertions(+), 15 deletions(-) diff --git a/python/sglang/srt/layers/quantization/fp8.py b/python/sglang/srt/layers/quantization/fp8.py index f1d5cc4a9396..9c488539c35c 100644 --- a/python/sglang/srt/layers/quantization/fp8.py +++ b/python/sglang/srt/layers/quantization/fp8.py @@ -369,11 +369,19 @@ def validate_block_quant_shapes( f"{input_size_per_partition} is not divisible by " f"weight quantization block_k = {block_k}." ) - # Required by column parallel or enabling merged weights - if ( + # Required by column parallel or enabling merged weights. + is_tp_split = ( tp_size > 1 and output_size // output_size_per_partition == tp_size - ) or len(output_partition_sizes) > 1: - for output_partition_size in output_partition_sizes: + ) + is_merged_gemm = len(output_partition_sizes) > 1 + if is_tp_split or is_merged_gemm: + sizes_to_check = output_partition_sizes + if not is_tp_split and is_merged_gemm: + # Match validate_fp8_block_shape: merged weights may have a + # ragged final logical matrix, and scale tensors are already + # allocated with ceil-divided block counts. + sizes_to_check = output_partition_sizes[:-1] + for output_partition_size in sizes_to_check: if output_partition_size % block_n != 0: raise ValueError( f"Weight output_partition_size = " diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 48a4d1b02cdd..94a95dc7a24e 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -565,7 +565,8 @@ def init_ipc_channels(self, port_args: PortArgs): [ self.recv_from_tokenizer, self.recv_from_rpc, - ] + ], + can_empty_cache=lambda: not self._engine_paused, ) else: self.recv_from_tokenizer = None @@ -1688,6 +1689,39 @@ def recv_requests( except zmq.ZMQError: break recv_reqs.append(recv_rpc) + + should_block_for_paused_rpc = True + if self.server_args.enable_dp_attention: + # In DP-attention mode, control requests are broadcast via + # tp_group below. Only the tp_group source should wait on + # the RPC socket while paused; the other ranks must enter + # the broadcast as receivers. If they block on ZMQ here, + # pause/flush can deadlock the control broadcast. + should_block_for_paused_rpc = self.tp_group.is_first_rank + + if ( + self._engine_paused + and len(recv_reqs) == 0 + and should_block_for_paused_rpc + ): + poller = zmq.Poller() + poller.register(self.recv_from_tokenizer, zmq.POLLIN) + poller.register(self.recv_from_rpc, zmq.POLLIN) + + while len(recv_reqs) == 0: + events = dict(poller.poll()) + for socket in (self.recv_from_tokenizer, self.recv_from_rpc): + if socket not in events: + continue + while True: + try: + if self.recv_limit_reached(len(recv_reqs)): + break + recv_reqs.append(socket.recv_pyobj(zmq.NOBLOCK)) + except zmq.ZMQError: + break + if self.recv_limit_reached(len(recv_reqs)): + break else: recv_reqs = None else: @@ -3215,7 +3249,9 @@ def _check_pending_flush(self): pending_req, deadline = self._pending_flush - if self.is_fully_idle(): + if self.is_fully_idle() or ( + self._engine_paused and self.running_batch.is_empty() + ): success = self.flush_cache() self._pending_flush = None self.send_to_tokenizer.send_output( @@ -3276,10 +3312,18 @@ def flush_cache_wrapped( timeout_s = float(recv_req.timeout_s or 0.0) if timeout_s <= 0.0: - return FlushCacheReqOutput(success=self.flush_cache()) + success = self.flush_cache() + if self.tp_cpu_group is not None: + barrier(group=self.tp_cpu_group) + return FlushCacheReqOutput(success=success) - if self.is_fully_idle(): - return FlushCacheReqOutput(success=self.flush_cache()) + if self.is_fully_idle() or ( + self._engine_paused and self.running_batch.is_empty() + ): + success = self.flush_cache() + if self.tp_cpu_group is not None: + barrier(group=self.tp_cpu_group) + return FlushCacheReqOutput(success=success) self._pending_flush = (recv_req, time.monotonic() + timeout_s) return None @@ -3443,7 +3487,10 @@ def detach_hicache_storage_wrapped( def flush_cache(self, empty_cache: bool = True): """Flush memory pools (e.g., KV cache, Mamba cache) and optionally empty device allocator cache.""" - if self.is_fully_idle(): + can_flush = self.is_fully_idle() or ( + self._engine_paused and self.running_batch.is_empty() + ) + if can_flush: self.cur_batch = None self.last_batch = None self.tree_cache.reset() @@ -3455,7 +3502,7 @@ def flush_cache(self, empty_cache: bool = True): if self.draft_worker: self.draft_worker.clear_cache_pool() - if empty_cache: + if empty_cache and not self._engine_paused: empty_device_cache(self.device_module) logger.info("Cache flushed successfully!") success = True @@ -3839,9 +3886,10 @@ class IdleSleeper: data that needs handling immediately. """ - def __init__(self, sockets): + def __init__(self, sockets, can_empty_cache=None): self.poller = zmq.Poller() self.last_empty_time = real_time() + self.can_empty_cache = can_empty_cache for s in sockets: self.poller.register(s, zmq.POLLIN) @@ -3854,7 +3902,8 @@ def maybe_sleep(self): and real_time() - self.last_empty_time > self.empty_cache_interval ): self.last_empty_time = real_time() - empty_device_cache() + if self.can_empty_cache is None or self.can_empty_cache(): + empty_device_cache() def is_health_check_generate_req(recv_req): diff --git a/python/sglang/srt/managers/scheduler_update_weights_mixin.py b/python/sglang/srt/managers/scheduler_update_weights_mixin.py index 792fd6194db7..e4704204ca45 100644 --- a/python/sglang/srt/managers/scheduler_update_weights_mixin.py +++ b/python/sglang/srt/managers/scheduler_update_weights_mixin.py @@ -52,6 +52,16 @@ def flush_cache_after_weight_update(self: Scheduler, recv_req) -> None: ) assert flush_cache_success, "Cache flush failed after updating weights" + def _quiesce_for_weight_update(self: Scheduler): + """Drain in-flight forward work before any NCCL weight mutation. + Synchronize forward_stream and schedule_stream to ensure all ranks are quiescent. + """ + if self.enable_overlap: + self.forward_stream.synchronize() + self.schedule_stream.synchronize() + if self.tp_cpu_group is not None: + torch.distributed.barrier(group=self.tp_cpu_group) + def update_weights_from_disk( self: Scheduler, recv_req: UpdateWeightFromDiskReqInput ): @@ -85,17 +95,20 @@ def update_weights_from_distributed( recv_req: UpdateWeightsFromDistributedReqInput, ) -> Tuple[bool, str]: """Update the online model parameter.""" + self._quiesce_for_weight_update() success, message = self.tp_worker.update_weights_from_distributed(recv_req) if success: self.flush_cache_after_weight_update(recv_req) else: logger.error(message) + torch.distributed.barrier(group=self.tp_cpu_group) return UpdateWeightsFromDistributedReqOutput(success, message) def update_weights_from_tensor( self: Scheduler, recv_req: UpdateWeightsFromTensorReqInput ): """Update the online model parameter from tensors.""" + self._quiesce_for_weight_update() if recv_req.disable_draft_model: worker = self.tp_worker else: @@ -112,6 +125,7 @@ def update_weights_from_ipc( self: Scheduler, recv_req: UpdateWeightsFromIPCReqInput ): """Update the online model parameter from IPC for checkpoint-engine integration.""" + self._quiesce_for_weight_update() success, message = self.tp_worker.update_weights_from_ipc(recv_req) tp_success = success if success and self.draft_worker is not None: @@ -125,7 +139,10 @@ def update_weights_from_ipc( def post_process_weights(self, recv_req: PostProcessWeightsReqInput): """Optional post-processing for updated weights (e.g., Marlin conversion).""" + self._quiesce_for_weight_update() success, message = self.tp_worker.post_process_weights(recv_req) + if self.tp_cpu_group is not None: + torch.distributed.barrier(group=self.tp_cpu_group) return PostProcessWeightsReqOutput(success, message) def get_weights_by_name(self: Scheduler, recv_req: GetWeightsByNameReqInput): diff --git a/python/sglang/srt/managers/tokenizer_control_mixin.py b/python/sglang/srt/managers/tokenizer_control_mixin.py index edb685dff9a5..a5ae7d4829e0 100644 --- a/python/sglang/srt/managers/tokenizer_control_mixin.py +++ b/python/sglang/srt/managers/tokenizer_control_mixin.py @@ -4,6 +4,7 @@ import logging import time import uuid +from contextlib import nullcontext from typing import ( TYPE_CHECKING, Any, @@ -541,9 +542,16 @@ async def post_process_weights( ) -> Tuple[bool, str]: """Trigger post-processing hooks for weights after loading.""" self.auto_create_handle_loop() - async with self.model_update_lock.writer_lock: + + async with self.is_pause_cond: + is_paused = self.is_pause + + lock_context = ( + self.model_update_lock.writer_lock if not is_paused else nullcontext() + ) + async with lock_context: results = await self.post_process_weights_communicator(obj) - return FanOutCommunicator.merge_results(results) + return FanOutCommunicator.merge_results(results) async def _unload_lora_adapter_locked( self: TokenizerManager, From 329828baeb2ebe955b387f4779d9bbecf8b70f82 Mon Sep 17 00:00:00 2001 From: zyzshishui Date: Thu, 16 Apr 2026 14:31:10 -0700 Subject: [PATCH 27/50] [11/14] [sglang-miles] R3 support on PD disaggregation mini_lb (#22916) --- python/sglang/srt/disaggregation/prefill.py | 2 + .../bindings/python/pyproject.toml | 1 + .../python/src/sglang_router/mini_lb.py | 54 +++++++++++++++---- 3 files changed, 47 insertions(+), 10 deletions(-) diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index c193250d11e9..672d08a6df7f 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -648,6 +648,8 @@ def process_disagg_prefill_inflight_queue( if poll in [KVPoll.WaitingForInput, KVPoll.Transferring]: undone_reqs.append(req) elif poll == KVPoll.Success: # transfer done + if req.return_routed_experts: + self.maybe_collect_routed_experts(req) release_kv_cache(req, self.tree_cache) # unlock the tree req.finished_reason = FINISH_LENGTH(length=0) # FIXME: clean up req's data in transfer engine diff --git a/sgl-model-gateway/bindings/python/pyproject.toml b/sgl-model-gateway/bindings/python/pyproject.toml index c44b1b96abb4..98a9fbe3d9de 100644 --- a/sgl-model-gateway/bindings/python/pyproject.toml +++ b/sgl-model-gateway/bindings/python/pyproject.toml @@ -32,6 +32,7 @@ dependencies = [ "setproctitle", "aiohttp", "orjson", + "pybase64", "uvicorn", "fastapi", ] diff --git a/sgl-model-gateway/bindings/python/src/sglang_router/mini_lb.py b/sgl-model-gateway/bindings/python/src/sglang_router/mini_lb.py index 11ef2f4c6b64..924edf393067 100644 --- a/sgl-model-gateway/bindings/python/src/sglang_router/mini_lb.py +++ b/sgl-model-gateway/bindings/python/src/sglang_router/mini_lb.py @@ -14,6 +14,7 @@ import aiohttp import orjson +import pybase64 import uvicorn from fastapi import FastAPI, HTTPException from fastapi.responses import ORJSONResponse, Response, StreamingResponse @@ -34,6 +35,44 @@ def maybe_wrap_ipv6_address(address: str) -> str: return address +def _merge_routed_experts(prefill: dict, decode: dict): + if "routed_experts" not in prefill or "routed_experts" not in decode: + return False + + prefill_bytes = pybase64.b64decode(prefill["routed_experts"], validate=True) + decode_bytes = pybase64.b64decode(decode["routed_experts"], validate=True) + decode["routed_experts"] = pybase64.b64encode( + prefill_bytes + decode_bytes[len(prefill_bytes) :] + ).decode("utf-8") + return True + + +def _merge_input_token_logprobs(prefill_meta: dict, decode_meta: dict): + if ( + "input_token_logprobs" not in prefill_meta + or "input_token_logprobs" not in decode_meta + ): + return + + decode_meta["input_token_logprobs"] = ( + prefill_meta["input_token_logprobs"] + decode_meta["input_token_logprobs"] + ) + + +def _merge_prefill_json(prefill_json, decode_json): + if "meta_info" in prefill_json and "meta_info" in decode_json: + prefill_meta = prefill_json["meta_info"] + decode_meta = decode_json["meta_info"] + _merge_input_token_logprobs(prefill_meta, decode_meta) + _merge_routed_experts(prefill_meta, decode_meta) + + if "sglext" not in prefill_json: + return + + if "sglext" in decode_json: + _merge_routed_experts(prefill_json["sglext"], decode_json["sglext"]) + + class MiniLoadBalancer: def __init__( self, @@ -140,18 +179,13 @@ async def generate( # Wait for both responses to complete. Prefill should end first. prefill_response, decode_response = await asyncio.gather(*tasks) - if "return_logprob" in modified_request: - + if ( + "return_logprob" in modified_request + or "return_routed_experts" in modified_request + ): prefill_json = await prefill_response.json() ret_json = await decode_response.json() - - # merge `meta_info.input_token_logprobs` from prefill to decode - if "meta_info" in ret_json: - if "input_token_logprobs" in ret_json["meta_info"]: - ret_json["meta_info"]["input_token_logprobs"] = ( - prefill_json["meta_info"]["input_token_logprobs"] - + ret_json["meta_info"]["input_token_logprobs"] - ) + _merge_prefill_json(prefill_json, ret_json) else: ret_json = await decode_response.json() From 6ce6b5f2a4e31b8465e0d97334aee5456728969c Mon Sep 17 00:00:00 2001 From: Byron Hsu Date: Sun, 26 Apr 2026 21:48:50 -0700 Subject: [PATCH 28/50] [12/14] [sglang-miles] Improve PD pause handling (#23672, #23887) --- python/sglang/srt/disaggregation/decode.py | 3 + python/sglang/srt/managers/scheduler.py | 14 +++ .../test_disaggregation_basic.py | 97 +++++++++++++++++++ .../test_scheduler_pause_generation.py | 2 + 4 files changed, 116 insertions(+) diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 0afce703adb2..0e2f402b9def 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -1585,7 +1585,9 @@ def event_loop_normal_disagg_decode(self: Scheduler): # Receive requests recv_reqs = self.recv_requests() self.process_input_requests(recv_reqs) + self.process_decode_queue() + if self._engine_paused: continue @@ -1613,6 +1615,7 @@ def event_loop_overlap_disagg_decode(self: Scheduler): # Receive requests recv_reqs = self.recv_requests() self.process_input_requests(recv_reqs) + self.process_decode_queue() if self._engine_paused: continue diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 94a95dc7a24e..a64ede6e7fdd 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -3707,6 +3707,20 @@ def _pause_engine(self) -> Tuple[List[Req], int]: raise NotImplementedError() def pause_generation(self, recv_req: PauseGenerationReqInput): + # PD disaggregation: `retract` has no decode-to-prefill rebootstrap path, + # so retracted decode-side requests cannot be re-prefilled under new + # weights. Fail fast before mutating any scheduler state to avoid + # leaving the engine in a half-paused / inconsistent state that later + # crashes inside radix-cache cleanup on flush_cache. + assert not ( + recv_req.mode == "retract" + and self.disaggregation_mode != DisaggregationMode.NULL + ), ( + "pause_generation(mode='retract') is not supported in PD " + "disaggregation mode yet. Decode-side retracted requests need " + "a rebootstrap path back to prefill." + ) + self._engine_paused = True if recv_req.mode == "in_place": diff --git a/test/registered/disaggregation/test_disaggregation_basic.py b/test/registered/disaggregation/test_disaggregation_basic.py index f14cc17561fd..c729bd5eb5c8 100644 --- a/test/registered/disaggregation/test_disaggregation_basic.py +++ b/test/registered/disaggregation/test_disaggregation_basic.py @@ -1,7 +1,9 @@ import asyncio import json import os +import time import unittest +from concurrent.futures import ThreadPoolExecutor, as_completed from types import SimpleNamespace import aiohttp @@ -100,6 +102,101 @@ def test_structured_output(self): # ensure the output is a valid JSON json.loads(output) + def test_pause_resume_in_place(self): + """Send requests, pause mid-generation, verify no progress during pause, resume.""" + NUM_REQUESTS = 32 + MAX_NEW_TOKENS = 512 + REQUEST_TIMEOUT = 180 + PAUSE_DURATION = 5 + + def _generate(prompt_id): + return requests.post( + self.lb_url + "/generate", + json={ + "text": f"Question {prompt_id}: Write a short essay about the number {prompt_id}.", + "sampling_params": { + "temperature": 0.8, + "max_new_tokens": MAX_NEW_TOKENS, + }, + }, + timeout=REQUEST_TIMEOUT, + ) + + with ThreadPoolExecutor(max_workers=NUM_REQUESTS) as executor: + futures = {executor.submit(_generate, i): i for i in range(NUM_REQUESTS)} + + time.sleep(1) + + requests.post( + self.prefill_url + "/pause_generation", + json={"mode": "in_place"}, + timeout=30, + ).raise_for_status() + requests.post( + self.decode_url + "/pause_generation", + json={"mode": "in_place"}, + timeout=30, + ).raise_for_status() + + time.sleep(0.5) + done_before = sum(1 for f in futures if f.done()) + + time.sleep(PAUSE_DURATION) + done_after = sum(1 for f in futures if f.done()) + + self.assertLess( + done_before, + NUM_REQUESTS, + "All requests completed before pause took effect — " + "increase MAX_NEW_TOKENS to make the test meaningful.", + ) + + self.assertEqual( + done_after - done_before, + 0, + f"{done_after - done_before} requests completed during pause " + f"({done_before} before, {done_after} after) — " + f"pause_generation was not respected by the disagg scheduler.", + ) + + requests.post( + self.decode_url + "/continue_generation", + json={}, + timeout=30, + ).raise_for_status() + requests.post( + self.prefill_url + "/continue_generation", + json={}, + timeout=30, + ).raise_for_status() + + completed = 0 + errors = [] + for future in as_completed(futures, timeout=REQUEST_TIMEOUT): + prompt_id = futures[future] + try: + resp = future.result() + if resp.status_code == 200: + body = resp.json() + self.assertIn("text", body) + self.assertGreater(len(body["text"]), 0) + completed += 1 + else: + errors.append(f"Request {prompt_id}: status={resp.status_code}") + except Exception as e: + errors.append(f"Request {prompt_id}: exception={e}") + + self.assertEqual( + completed + len(errors), + NUM_REQUESTS, + "Some requests did not resolve within the timeout — likely hung during pause.", + ) + self.assertEqual( + completed, + NUM_REQUESTS, + f"Some requests failed: {completed}/{NUM_REQUESTS} succeeded. Errors: {errors}", + ) + def test_first_token_finish(self): client = openai.Client(api_key="empty", base_url=f"{self.lb_url}/v1") tokenizer = AutoTokenizer.from_pretrained(self.model) diff --git a/test/registered/unit/managers/test_scheduler_pause_generation.py b/test/registered/unit/managers/test_scheduler_pause_generation.py index b1059a33e15d..816c3545fe40 100644 --- a/test/registered/unit/managers/test_scheduler_pause_generation.py +++ b/test/registered/unit/managers/test_scheduler_pause_generation.py @@ -7,6 +7,7 @@ maybe_stub_sgl_kernel() +from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.managers.io_struct import PauseGenerationReqInput from sglang.srt.managers.scheduler import Scheduler from sglang.srt.managers.scheduler_runtime_checker_mixin import PoolStats @@ -22,6 +23,7 @@ def _new_scheduler(self) -> Scheduler: scheduler.last_batch = None scheduler.cur_batch = None scheduler.chunked_req = None + scheduler.disaggregation_mode = DisaggregationMode.NULL scheduler.running_batch = MagicMock() scheduler.running_batch.reqs = [] scheduler.running_batch.is_empty.return_value = True From 44f350d10538eaf7c2643d043ff2fb227d74fe2e Mon Sep 17 00:00:00 2001 From: Jiajun Li <48857426+guapisolo@users.noreply.github.com> Date: Wed, 13 May 2026 18:15:05 -0700 Subject: [PATCH 29/50] [13/14] [sglang-miles] Add KimiK2 raw tool call id parser (#25196) --- .../srt/entrypoints/openai/serving_chat.py | 17 +++++++++----- python/sglang/srt/function_call/core_types.py | 6 +++++ .../srt/function_call/function_call_parser.py | 6 ++++- .../srt/function_call/kimik2_detector.py | 22 +++++++++++++++++++ 4 files changed, 45 insertions(+), 6 deletions(-) diff --git a/python/sglang/srt/entrypoints/openai/serving_chat.py b/python/sglang/srt/entrypoints/openai/serving_chat.py index 5e2a70376541..e8e4a6f412ec 100644 --- a/python/sglang/srt/entrypoints/openai/serving_chat.py +++ b/python/sglang/srt/entrypoints/openai/serving_chat.py @@ -1312,11 +1312,7 @@ def _process_tool_call_id( history_tool_calls_cnt: int, ) -> str: """Process for generating a new and unique `tool_call_id`""" - if self.tool_call_parser != "kimi_k2": - # A simple uuid is sufficient for all models except for Kimi-K2. - tool_call_id = f"call_{uuid.uuid4().hex[:24]}" - return tool_call_id - else: + if self.tool_call_parser == "kimi_k2": # Align with Kimi-K2 format: functions.{name}:{index} # Kimi-K2 allows multiple tool_calls in one message; SGLang sets call_item.tool_index to the *local* position inside that message. # Therefore, the index must be corrected by using `history_tool_calls_cnt + call_item.tool_index` to ensure globally unique and properly ordered. @@ -1325,6 +1321,17 @@ def _process_tool_call_id( f"Process tool call idx, parser: {self.tool_call_parser}, tool_call_id: {tool_call_id}, history_cnt: {history_tool_calls_cnt}" ) return tool_call_id + if self.tool_call_parser == "kimi_k2_raw_id": + # RL training needs the model-emitted tool_call_id round-tripped verbatim, + # so we skip the history-based renumbering above and return whatever the + # detector captured. Fall back to the canonical Kimi-K2 reconstruction + # (without history offset) if for any reason the detector did not record + # a raw id — the raw id field is best-effort but the format is stable. + if call_item.tool_call_id: + return call_item.tool_call_id + return f"functions.{call_item.name}:{call_item.tool_index}" + # A simple uuid is sufficient for all other models. + return f"call_{uuid.uuid4().hex[:24]}" def _process_tool_calls( self, diff --git a/python/sglang/srt/function_call/core_types.py b/python/sglang/srt/function_call/core_types.py index 1ea87df798c8..297dd2712ef3 100644 --- a/python/sglang/srt/function_call/core_types.py +++ b/python/sglang/srt/function_call/core_types.py @@ -10,6 +10,12 @@ class ToolCallItem(BaseModel): tool_index: int name: Optional[str] = None parameters: str # JSON string + # The tool_call_id string emitted by the model, captured verbatim. + # Only populated by detectors whose downstream consumers need the exact + # model-emitted id (e.g. RL training trajectories). Existing detectors + # leave this as None and the serving layer falls back to its usual id + # generation strategy. + tool_call_id: Optional[str] = None class StreamingParseResult(BaseModel): diff --git a/python/sglang/srt/function_call/function_call_parser.py b/python/sglang/srt/function_call/function_call_parser.py index 602e93fdd67c..a0ad790156a5 100644 --- a/python/sglang/srt/function_call/function_call_parser.py +++ b/python/sglang/srt/function_call/function_call_parser.py @@ -24,7 +24,10 @@ from sglang.srt.function_call.hermes_detector import HermesDetector from sglang.srt.function_call.hunyuan_detector import HunyuanDetector from sglang.srt.function_call.internlm_detector import InternlmDetector -from sglang.srt.function_call.kimik2_detector import KimiK2Detector +from sglang.srt.function_call.kimik2_detector import ( + KimiK2Detector, + KimiK2RawIdDetector, +) from sglang.srt.function_call.lfm2_detector import Lfm2Detector from sglang.srt.function_call.llama32_detector import Llama32Detector from sglang.srt.function_call.mimo_detector import MiMoDetector @@ -63,6 +66,7 @@ class FunctionCallParser: "glm47": Glm47MoeDetector, "gpt-oss": GptOssDetector, "kimi_k2": KimiK2Detector, + "kimi_k2_raw_id": KimiK2RawIdDetector, "lfm2": Lfm2Detector, "llama3": Llama32Detector, "mimo": MiMoDetector, diff --git a/python/sglang/srt/function_call/kimik2_detector.py b/python/sglang/srt/function_call/kimik2_detector.py index da2c76fd09ff..1e54c8c4bfc5 100644 --- a/python/sglang/srt/function_call/kimik2_detector.py +++ b/python/sglang/srt/function_call/kimik2_detector.py @@ -186,6 +186,7 @@ def detect_and_parse(self, text: str, tools: List[Tool]) -> StreamingParseResult tool_index=function_idx, name=function_name, parameters=function_args, + tool_call_id=function_id, ) ) @@ -254,6 +255,7 @@ def parse_streaming_increment( tool_index=self.current_tool_id, name=function_name, parameters="", + tool_call_id=function_id, ) ) self.current_tool_name_sent = True @@ -342,3 +344,23 @@ def get_info(name: str) -> StructureInfo: # the base get_structural_tag_name (returns None) keeps FunctionCallParser # on the legacy path, whose structure_info bakes the section markers in. # TODO: re-enable the builtin once https://github.com/mlc-ai/xgrammar/issues/622 is fixed. + + +class KimiK2RawIdDetector(KimiK2Detector): + """ + Variant of KimiK2Detector that preserves the model-emitted tool_call_id verbatim. + + The default kimi_k2 path renumbers ids via `history_tool_calls_cnt + tool_index` + in the serving layer so that multi-turn conversations get globally unique, + monotonically increasing ids (see PR #10600). That is the right behavior for + chat use cases. + + RL training has the opposite requirement: the trajectory must round-trip the + exact tool_call_id the model produced (e.g. `functions.foo:5`), so that the + follow-up tool result turn references the same id the policy emitted. This + subclass exists purely as a marker so the serving layer can branch on the + parser name and use `ToolCallItem.tool_call_id` directly. Parsing logic + is identical to KimiK2Detector. + """ + + pass From e687c743b568e07750124d6441dbe639684a2a76 Mon Sep 17 00:00:00 2001 From: maocheng23 <35615230+maocheng23@users.noreply.github.com> Date: Sun, 17 May 2026 23:20:41 -0700 Subject: [PATCH 30/50] [14/14] [sglang-miles] Add true on-policy qwen_dense support --- python/sglang/srt/debug_utils/dumper.py | 15 +- .../srt/distributed/communication_op.py | 9 + python/sglang/srt/layers/activation.py | 4 +- python/sglang/srt/layers/attention/vision.py | 8 +- python/sglang/srt/layers/communicator.py | 10 + python/sglang/srt/layers/layernorm.py | 38 +- python/sglang/srt/layers/linear.py | 13 +- python/sglang/srt/layers/logits_processor.py | 8 + python/sglang/srt/layers/on_policy_utils.py | 6 + .../srt/layers/rotary_embedding/base.py | 10 +- .../srt/layers/rotary_embedding/mrope.py | 4 +- python/sglang/srt/layers/sampler.py | 7 +- .../srt/model_executor/cuda_graph_runner.py | 33 +- .../srt/model_executor/forward_batch_info.py | 11 +- .../sglang/srt/model_executor/model_runner.py | 5 + python/sglang/srt/models/qwen2.py | 31 +- python/sglang/srt/models/qwen2_moe.py | 10 +- python/sglang/srt/models/qwen3.py | 60 +- python/sglang/srt/models/qwen3_moe.py | 29 +- python/sglang/srt/models/sdar.py | 34 +- python/sglang/srt/models/sdar_moe.py | 34 +- python/sglang/srt/models/step3p5.py | 7 +- python/sglang/srt/models/utils.py | 2 + .../multimodal/processors/base_processor.py | 4 +- python/sglang/srt/server_args.py | 64 +- .../sglang/srt/tp_invariant_ops/__init__.py | 23 + .../srt/tp_invariant_ops/tp_invariant_ops.py | 1941 +++++++++++++++++ python/sglang/srt/true_on_policy/__init__.py | 49 + python/sglang/srt/true_on_policy/config.py | 150 ++ python/sglang/srt/true_on_policy/contracts.py | 111 + python/sglang/srt/true_on_policy/schema.py | 33 + test/manual/layers/test_layernorm.py | 20 + .../core/test_dense_deterministic_math.py | 338 +++ test/registered/core/test_on_policy_wiring.py | 606 +++++ test/registered/core/test_tp_invariant_ops.py | 866 ++++++++ 35 files changed, 4423 insertions(+), 170 deletions(-) create mode 100644 python/sglang/srt/layers/on_policy_utils.py create mode 100644 python/sglang/srt/tp_invariant_ops/__init__.py create mode 100644 python/sglang/srt/tp_invariant_ops/tp_invariant_ops.py create mode 100644 python/sglang/srt/true_on_policy/__init__.py create mode 100644 python/sglang/srt/true_on_policy/config.py create mode 100644 python/sglang/srt/true_on_policy/contracts.py create mode 100644 python/sglang/srt/true_on_policy/schema.py create mode 100644 test/registered/core/test_dense_deterministic_math.py create mode 100644 test/registered/core/test_on_policy_wiring.py create mode 100644 test/registered/core/test_tp_invariant_ops.py diff --git a/python/sglang/srt/debug_utils/dumper.py b/python/sglang/srt/debug_utils/dumper.py index 6c7d3f6e570a..413e4f45b88c 100644 --- a/python/sglang/srt/debug_utils/dumper.py +++ b/python/sglang/srt/debug_utils/dumper.py @@ -492,11 +492,16 @@ def _dump_inner( meta_only_fields={**(value_meta_only_fields or {}), **recompute_meta}, ) - if ( - enable_curr_grad - and isinstance(value, torch.Tensor) - and (g := value.grad) is not None - ): + if enable_curr_grad and isinstance(value, torch.Tensor): + g = ( + value.grad + if value.grad is not None + else getattr(value, "main_grad", None) + ) + else: + g = None + + if g is not None: self._dump_single( tag=grad_tag, tags={**tags, "name": f"grad__{name}"}, diff --git a/python/sglang/srt/distributed/communication_op.py b/python/sglang/srt/distributed/communication_op.py index 89f9986e4f68..674300d3b718 100644 --- a/python/sglang/srt/distributed/communication_op.py +++ b/python/sglang/srt/distributed/communication_op.py @@ -7,6 +7,9 @@ import torch import torch.distributed +from sglang.srt.tp_invariant_ops import tree_all_reduce_sum +from sglang.srt.true_on_policy import should_use_tp_invariant_tree_all_reduce + from .parallel_state import ( get_attn_tp_group, get_moe_ep_group, @@ -17,6 +20,8 @@ def tensor_model_parallel_all_reduce(input_: torch.Tensor) -> torch.Tensor: """All-reduce the input tensor across model parallel group.""" + if should_use_tp_invariant_tree_all_reduce(): + return tree_all_reduce_sum(input_, device_group=get_tp_group().device_group) return get_tp_group().all_reduce(input_) @@ -64,6 +69,10 @@ def broadcast_tensor_dict( def attention_tensor_model_parallel_all_reduce(input_: torch.Tensor) -> torch.Tensor: """All-reduce the input tensor across attention parallel group.""" + if should_use_tp_invariant_tree_all_reduce(): + return tree_all_reduce_sum( + input_, device_group=get_attn_tp_group().device_group + ) return get_attn_tp_group().all_reduce(input_) diff --git a/python/sglang/srt/layers/activation.py b/python/sglang/srt/layers/activation.py index 216e37a234ae..64bd7b033e50 100644 --- a/python/sglang/srt/layers/activation.py +++ b/python/sglang/srt/layers/activation.py @@ -30,7 +30,7 @@ from sglang.srt.environ import envs from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.utils import MultiPlatformOp -from sglang.srt.server_args import get_global_server_args +from sglang.srt.true_on_policy import is_true_on_policy_enabled from sglang.srt.utils import ( cpu_has_amx_support, is_cpu, @@ -80,7 +80,7 @@ def _(x): class SiluAndMul(MultiPlatformOp): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) - if get_global_server_args().rl_on_policy_target is not None: + if is_true_on_policy_enabled(): self._forward_method = self.forward_native def forward_native(self, x: torch.Tensor) -> torch.Tensor: diff --git a/python/sglang/srt/layers/attention/vision.py b/python/sglang/srt/layers/attention/vision.py index 6791d0c79078..30a4f02da40f 100644 --- a/python/sglang/srt/layers/attention/vision.py +++ b/python/sglang/srt/layers/attention/vision.py @@ -68,6 +68,7 @@ from sglang.srt.layers.quantization import QuantizationConfig from sglang.srt.layers.rotary_embedding import apply_rotary_pos_emb from sglang.srt.server_args import get_global_server_args +from sglang.srt.true_on_policy import is_true_on_policy_enabled from sglang.srt.utils import add_prefix, get_bool_env_var _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip @@ -914,7 +915,7 @@ def _init_qk_norm( weight_dtype=torch.float32, cast_x_before_out_mul=True, ) - if get_global_server_args().rl_on_policy_target is not None + if is_true_on_policy_enabled() else {} ) q_norm = RMSNorm( @@ -1043,10 +1044,7 @@ def forward( if x.dim() == 2: x = x.unsqueeze(0) assert x.dim() == 3, x.shape - if ( - get_global_server_args().rl_on_policy_target is not None - and position_embeddings is not None - ): + if is_true_on_policy_enabled() and position_embeddings is not None: assert isinstance(position_embeddings, tuple), ( "expected position_embeddings to be a tuple of two tensors,\n" f"but got {type(position_embeddings)}, change if needed" diff --git a/python/sglang/srt/layers/communicator.py b/python/sglang/srt/layers/communicator.py index 30f068595188..48de54b1a09b 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -68,6 +68,10 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.server_args import get_global_server_args from sglang.srt.speculative.spec_info import SpeculativeAlgorithm +from sglang.srt.true_on_policy import ( + should_disable_mlp_allreduce_fusion_for_on_policy, + should_disable_reduce_scatter_for_on_policy, +) from sglang.srt.utils import ( get_bool_env_var, is_cuda, @@ -696,6 +700,9 @@ def postprocess_layer( ) def should_use_reduce_scatter(self, forward_batch: ForwardBatch): + if should_disable_reduce_scatter_for_on_policy(): + return False + if not self.allow_reduce_scatter: return False if ( @@ -716,6 +723,9 @@ def should_use_reduce_scatter(self, forward_batch: ForwardBatch): def should_fuse_mlp_allreduce_with_next_layer( self, forward_batch: ForwardBatch ) -> bool: + if should_disable_mlp_allreduce_fusion_for_on_policy(): + return False + # When MOE_FULL is active (moe_cp allgather), fusion must be disabled because # the fusion path skips postprocess_layer which contains the moe_cp scatter. # Without scatter, hidden_states remain at MOE_FULL size while residual is at diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index 1079434581f2..7158aabedff5 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -27,6 +27,10 @@ from sglang.srt.environ import envs from sglang.srt.layers.utils import MultiPlatformOp from sglang.srt.server_args import get_global_server_args +from sglang.srt.true_on_policy import ( + get_on_policy_rms_norm_kwargs, + is_true_on_policy_enabled, +) from sglang.srt.utils import ( cpu_has_amx_support, get_bool_env_var, @@ -181,13 +185,35 @@ def __init__( cast_x_before_out_mul: bool = False, fp32_residual: bool = True, has_weight: bool = True, + weight_dtype: Optional[torch.dtype] = None, + override_orig_dtype: Optional[torch.dtype] = None, + true_on_policy_weight_dtype: Optional[torch.dtype] = None, + true_on_policy_override_orig_dtype: Optional[torch.dtype] = None, + true_on_policy_fp32_residual: bool = False, ) -> None: super().__init__() + true_on_policy_kwargs = get_on_policy_rms_norm_kwargs( + weight_dtype=true_on_policy_weight_dtype, + override_orig_dtype=true_on_policy_override_orig_dtype, + fp32_residual=true_on_policy_fp32_residual, + ) + if not cast_x_before_out_mul: + cast_x_before_out_mul = true_on_policy_kwargs.get( + "cast_x_before_out_mul", cast_x_before_out_mul + ) + fp32_residual = true_on_policy_kwargs.get("fp32_residual", fp32_residual) + if weight_dtype is None: + weight_dtype = true_on_policy_kwargs.get("weight_dtype", weight_dtype) + if override_orig_dtype is None: + override_orig_dtype = true_on_policy_kwargs.get( + "override_orig_dtype", override_orig_dtype + ) self.has_weight = has_weight self.cast_x_before_out_mul = cast_x_before_out_mul self.fp32_residual = fp32_residual + self.override_orig_dtype = override_orig_dtype if self.has_weight: - self.weight = nn.Parameter(torch.ones(hidden_size)) + self.weight = nn.Parameter(torch.ones(hidden_size, dtype=weight_dtype)) else: self.weight = torch.ones(hidden_size) self.variance_epsilon = eps @@ -217,11 +243,17 @@ def forward_cuda( x = x.contiguous().reshape(-1, original_shape[-1]) if self.variance_size_override is not None: return self.forward_native(x, residual, post_residual_addition) + if ( + self.weight.dtype != x.dtype + or self.cast_x_before_out_mul + or self.override_orig_dtype is not None + ): + return self.forward_native(x, residual, post_residual_addition) if is_batch_invariant_mode_enabled(): if ( residual is not None or self.cast_x_before_out_mul - or get_global_server_args().rl_on_policy_target == "fsdp" + or is_true_on_policy_enabled() ): return self.forward_native(x, residual, post_residual_addition) return rms_norm_batch_invariant( @@ -363,7 +395,7 @@ def forward_native( ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: if not x.is_contiguous(): x = x.contiguous() - orig_dtype = x.dtype + orig_dtype = self.override_orig_dtype or x.dtype if residual is not None and not self.fp32_residual: x = x + residual diff --git a/python/sglang/srt/layers/linear.py b/python/sglang/srt/layers/linear.py index fc197860da30..b8fd06334fab 100644 --- a/python/sglang/srt/layers/linear.py +++ b/python/sglang/srt/layers/linear.py @@ -14,6 +14,7 @@ from sglang.kernel_api_logging import wrap_method_with_debug_kernel_once from sglang.srt.distributed import ( + attention_tensor_model_parallel_all_reduce, divide, get_tensor_model_parallel_rank, get_tensor_model_parallel_world_size, @@ -41,6 +42,7 @@ ) from sglang.srt.layers.utils import pad_or_narrow_weight from sglang.srt.server_args import get_global_server_args +from sglang.srt.true_on_policy import should_use_tp_invariant_row_linear from sglang.srt.utils import get_bool_env_var, is_cpu, is_hip, is_npu, set_weight_attrs if TYPE_CHECKING: @@ -1538,11 +1540,18 @@ def forward(self, input_, skip_all_reduce=False, forward_batch=None): get_tp_group(), disabled=not is_allocation_symmetric() ) with symm_ctx: - output_parallel = self.quant_method.apply(self, input_parallel, bias=bias_) + if should_use_tp_invariant_row_linear(input_parallel.shape[-1]): + output_parallel = torch.ops.tp_inv_ops.matmul_tp_inv( + input_parallel, self.weight.t(), bias_ + ) + else: + output_parallel = self.quant_method.apply( + self, input_parallel, bias=bias_ + ) if self.reduce_results and self.tp_size > 1 and not skip_all_reduce: if self.use_dp_attention_reduce: - output = get_attention_tp_group().all_reduce(output_parallel) + output = attention_tensor_model_parallel_all_reduce(output_parallel) else: quantize_communications = ( ( diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index 881fe11ee61c..3d9b2b839153 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -56,6 +56,7 @@ ForwardMode, ) from sglang.srt.server_args import get_global_server_args +from sglang.srt.true_on_policy import should_force_bfloat16_lm_head from sglang.srt.utils.common import is_npu, use_intel_amx_backend logger = logging.getLogger(__name__) @@ -899,6 +900,11 @@ def _compute_lm_head( logits = torch.matmul( hidden_states.to(torch.float32), lm_head.weight.to(torch.float32).T ) + elif should_force_bfloat16_lm_head(use_fp32_lm_head=self.use_fp32_lm_head): + logits = torch.matmul( + hidden_states.to(torch.bfloat16), + lm_head.weight.to(torch.bfloat16).T, + ).to(torch.bfloat16) elif use_intel_amx_backend(lm_head): logits = torch.ops.sgl_kernel.weight_packed_linear( hidden_states.to(lm_head.weight.dtype), @@ -987,6 +993,8 @@ def _copy_logits_to_buffer( assert logits_buffer.dtype == torch.float logits_buffer.copy_(logits[:, : self.vocab_size]) logits = logits_buffer + elif should_force_bfloat16_lm_head(use_fp32_lm_head=self.use_fp32_lm_head): + logits = logits[:, : self.vocab_size].to(torch.bfloat16) else: logits = logits[:, : self.vocab_size].float() return logits diff --git a/python/sglang/srt/layers/on_policy_utils.py b/python/sglang/srt/layers/on_policy_utils.py new file mode 100644 index 000000000000..234397b9b33c --- /dev/null +++ b/python/sglang/srt/layers/on_policy_utils.py @@ -0,0 +1,6 @@ +"""Compatibility imports for SGLang true-on-policy helpers. + +New code should import from :mod:`sglang.srt.true_on_policy`. +""" + +from sglang.srt.true_on_policy import * # noqa: F403 diff --git a/python/sglang/srt/layers/rotary_embedding/base.py b/python/sglang/srt/layers/rotary_embedding/base.py index ac5d59d4ac79..ba0170485043 100644 --- a/python/sglang/srt/layers/rotary_embedding/base.py +++ b/python/sglang/srt/layers/rotary_embedding/base.py @@ -8,7 +8,7 @@ from sglang.srt.layers.rotary_embedding.utils import apply_rotary_emb from sglang.srt.layers.utils import MultiPlatformOp -from sglang.srt.server_args import get_global_server_args +from sglang.srt.true_on_policy import is_true_on_policy_enabled from sglang.srt.utils import ( cpu_has_amx_support, get_bool_env_var, @@ -100,7 +100,7 @@ def __init__( self._apply_rotary_emb_wrapped = apply_rotary_emb # XXX (MUSA): Implement sgl_kernel.rotary_embedding support for MUSA backend - if get_global_server_args().rl_on_policy_target is not None or _is_musa: + if is_true_on_policy_enabled() or _is_musa: self._forward_method = self.forward_native self.position_cos, self.position_sin = None, None @@ -119,9 +119,7 @@ def _compute_inv_freq(self, base: Union[int, float]) -> torch.Tensor: # use CPU to compute the cache and then move it to GPU. However, we # create the cache on GPU for faster initialization. This may cause # a slight numerical difference between the HF implementation and ours. - init_device = ( - "cpu" if get_global_server_args().rl_on_policy_target is not None else None - ) + init_device = "cpu" if is_true_on_policy_enabled() else None inv_freq = 1.0 / ( base ** ( @@ -131,7 +129,7 @@ def _compute_inv_freq(self, base: Union[int, float]) -> torch.Tensor: / self.rotary_dim ) ) - if get_global_server_args().rl_on_policy_target is not None: + if is_true_on_policy_enabled(): inv_freq = inv_freq.cuda() return inv_freq diff --git a/python/sglang/srt/layers/rotary_embedding/mrope.py b/python/sglang/srt/layers/rotary_embedding/mrope.py index 9c93ad1ffd21..c896cf1d6b72 100644 --- a/python/sglang/srt/layers/rotary_embedding/mrope.py +++ b/python/sglang/srt/layers/rotary_embedding/mrope.py @@ -18,7 +18,7 @@ yarn_get_mscale_simple, yarn_linear_ramp_mask, ) -from sglang.srt.server_args import get_global_server_args +from sglang.srt.true_on_policy import is_true_on_policy_enabled from sglang.srt.utils import cpu_has_amx_support, is_cuda, is_npu _is_cuda = is_cuda() @@ -119,7 +119,7 @@ def __init__( self.register_buffer("axis_map", axis_map, persistent=False) else: self.axis_map = None - if get_global_server_args().rl_on_policy_target is not None: + if is_true_on_policy_enabled(): self._forward_method = self.forward_native def get_cos_sin_with_position(self, positions): diff --git a/python/sglang/srt/layers/sampler.py b/python/sglang/srt/layers/sampler.py index 1660bcb93a7a..8ad1ac6b788d 100644 --- a/python/sglang/srt/layers/sampler.py +++ b/python/sglang/srt/layers/sampler.py @@ -16,6 +16,7 @@ from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.sampling.sampling_params import TOP_K_ALL from sglang.srt.server_args import get_global_server_args +from sglang.srt.true_on_policy import resolve_true_on_policy_runtime_policy from sglang.srt.utils.common import ( crash_on_warnings, get_bool_env_var, @@ -62,13 +63,13 @@ def __init__(self): if is_dp_attention_enabled(): self.tp_sync_group = get_attention_tp_group().device_group - self.rl_on_policy_target = get_global_server_args().rl_on_policy_target + true_on_policy = resolve_true_on_policy_runtime_policy(get_global_server_args()) # In RL on-policy mode, deterministic inference is automatically enabled. self.enable_deterministic = ( get_global_server_args().enable_deterministic_inference ) # In RL on-policy mode, we use log_softmax to compute logprobs to match the trainer. - self.use_log_softmax_logprob = self.rl_on_policy_target is not None + self.use_log_softmax_logprob = true_on_policy.enabled self.use_ascend_backend = get_global_server_args().sampling_backend == "ascend" def _preprocess_logits( @@ -141,7 +142,7 @@ def forward( # In RL on-policy mode, we use log_softmax to compute logprobs to match the trainer. logprobs_via_logsoftmax_kernel = None - if self.rl_on_policy_target is not None: + if self.use_log_softmax_logprob: logprobs_via_logsoftmax_kernel = torch.log_softmax(logits, dim=-1) if self.use_ascend_backend: diff --git a/python/sglang/srt/model_executor/cuda_graph_runner.py b/python/sglang/srt/model_executor/cuda_graph_runner.py index c037b20dd9af..74c0bc11e37f 100644 --- a/python/sglang/srt/model_executor/cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/cuda_graph_runner.py @@ -71,6 +71,9 @@ ) from sglang.srt.model_executor.input_buffers import ForwardInputBuffers from sglang.srt.multiplex.pdmux_context import get_current_stream_idx, get_stream_groups +from sglang.srt.true_on_policy import ( + patch_prefill_only_deterministic_inference_for_cuda_graph, +) from sglang.srt.utils import ( empty_context, get_available_gpu_memory, @@ -911,20 +914,30 @@ def _capture_one_stream(stream_idx: Optional[int] = None): # Trigger CUDA graph capture for specific shapes. # Capture the large shapes first so that the smaller shapes # can reuse the memory pool allocated for the large shapes. - with freeze_gc(self.model_runner.server_args.enable_cudagraph_gc): - if not self.enable_pdmux: - with graph_capture() as graph_capture_context, profile_context as prof: - self.stream = graph_capture_context.stream - _capture_one_stream() - else: - set_pdmux_status(False) - for i, sg in enumerate(self.stream_groups): + with patch_prefill_only_deterministic_inference_for_cuda_graph( + self.model_runner.server_args, + attn_backend=getattr(self.model_runner, "attn_backend", None), + dvr_target_verify_cuda_graph=getattr( + self.model_runner, "enable_dvr_target_verify_cuda_graph", False + ), + ): + with freeze_gc(self.model_runner.server_args.enable_cudagraph_gc): + if not self.enable_pdmux: with ( - graph_capture(stream=sg[1]) as graph_capture_context, + graph_capture() as graph_capture_context, profile_context as prof, ): self.stream = graph_capture_context.stream - _capture_one_stream(i) + _capture_one_stream() + else: + set_pdmux_status(False) + for i, sg in enumerate(self.stream_groups): + with ( + graph_capture(stream=sg[1]) as graph_capture_context, + profile_context as prof, + ): + self.stream = graph_capture_context.stream + _capture_one_stream(i) _set_capture_lora_variant(None) diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 1671c5e42846..2557c8073c9e 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -56,6 +56,7 @@ ForwardBatchDeepSeekMHAMixin, ) from sglang.srt.server_args import get_global_server_args +from sglang.srt.true_on_policy import is_true_on_policy_enabled from sglang.srt.utils import ( is_cuda, is_hip, @@ -771,9 +772,8 @@ def _compute_mrope_positions( mm_input = batch.multimodal_inputs[batch_idx] if self.forward_mode.is_decode(): # 3 * N - if ( - mm_input is None - or get_global_server_args().rl_on_policy_target is not None + if mm_input is None or is_true_on_policy_enabled( + get_global_server_args() ): mrope_positions_list[batch_idx] = torch.full( (3, 1), @@ -790,9 +790,8 @@ def _compute_mrope_positions( batch.extend_seq_lens[batch_idx], batch.extend_prefix_lens[batch_idx], ) - if ( - mm_input is None - or get_global_server_args().rl_on_policy_target is not None + if mm_input is None or is_true_on_policy_enabled( + get_global_server_args() ): # text only mrope_positions = torch.tensor( diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index ee7ec4c81645..48c4f701e1e0 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -185,6 +185,7 @@ get_global_experts_capturer, set_global_experts_capturer, ) +from sglang.srt.true_on_policy import is_tp_invariant_target from sglang.srt.utils import ( MultiprocessingSerializer, broadcast_pyobj, @@ -755,6 +756,10 @@ def initialize(self, pre_model_load_memory: float): from sglang.srt.batch_invariant_ops import enable_batch_invariant_mode enable_batch_invariant_mode() + if is_tp_invariant_target(): + from sglang.srt.tp_invariant_ops import enable_tp_invariant_mode + + enable_tp_invariant_mode() # Deduce KV cache dtype self.configure_kv_cache_dtype() diff --git a/python/sglang/srt/models/qwen2.py b/python/sglang/srt/models/qwen2.py index fe1ec95522c5..3500acf6b53a 100644 --- a/python/sglang/srt/models/qwen2.py +++ b/python/sglang/srt/models/qwen2.py @@ -50,7 +50,11 @@ default_weight_loader, kv_cache_scales_loader, ) -from sglang.srt.server_args import get_global_server_args +from sglang.srt.true_on_policy import ( + get_on_policy_rms_norm_kwargs, + is_true_on_policy_enabled, + should_force_bfloat16_dense_tensor_math, +) from sglang.srt.utils import add_prefix, make_layers from sglang.srt.utils.hf_transformers_utils import get_rope_config @@ -96,9 +100,11 @@ def forward( x: torch.Tensor, forward_batch: ForwardBatch = None, ) -> torch.Tensor: - if get_global_server_args().rl_on_policy_target is not None: - x = x.bfloat16() - + if ( + should_force_bfloat16_dense_tensor_math() + or x.dtype != self.gate_up_proj.weight.dtype + ): + x = x.to(self.gate_up_proj.weight.dtype) gate_up, _ = self.gate_up_proj(x) x = self.act_fn(gate_up) x, _ = self.down_proj(x, forward_batch=forward_batch) @@ -284,11 +290,7 @@ def __init__( quant_config=quant_config, use_attn_tp_group=is_dp_attention_enabled(), prefix=add_prefix("embed_tokens", prefix), - params_dtype=( - torch.float32 - if get_global_server_args().rl_on_policy_target is not None - else None - ), + params_dtype=(torch.float32 if is_true_on_policy_enabled() else None), ) else: self.embed_tokens = PPMissingLayer() @@ -309,16 +311,7 @@ def __init__( prefix=add_prefix("layers", prefix), ) if self.pp_group.is_last_rank: - norm_kwargs = ( - dict( - weight_dtype=torch.float32, - cast_x_before_out_mul=True, - override_orig_dtype=torch.float32, - fp32_residual=True, - ) - if get_global_server_args().rl_on_policy_target is not None - else {} - ) + norm_kwargs = get_on_policy_rms_norm_kwargs() self.norm = RMSNorm( config.hidden_size, eps=config.rms_norm_eps, **norm_kwargs ) diff --git a/python/sglang/srt/models/qwen2_moe.py b/python/sglang/srt/models/qwen2_moe.py index 415c6044ff60..98dfd2cf11c1 100644 --- a/python/sglang/srt/models/qwen2_moe.py +++ b/python/sglang/srt/models/qwen2_moe.py @@ -89,6 +89,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.server_args import get_global_server_args +from sglang.srt.true_on_policy import get_on_policy_rms_norm_kwargs from sglang.srt.utils import ( add_prefix, cpu_has_amx_support, @@ -761,14 +762,7 @@ def __init__( prefix=add_prefix("layers", prefix), ) if self.pp_group.is_last_rank: - norm_kwargs = ( - dict( - cast_x_before_out_mul=True, - fp32_residual=False, - ) - if get_global_server_args().rl_on_policy_target is not None - else {} - ) + norm_kwargs = get_on_policy_rms_norm_kwargs(fp32_residual=False) self.norm = RMSNorm( config.hidden_size, eps=config.rms_norm_eps, **norm_kwargs ) diff --git a/python/sglang/srt/models/qwen3.py b/python/sglang/srt/models/qwen3.py index 64c18883c445..2527c8f73f91 100644 --- a/python/sglang/srt/models/qwen3.py +++ b/python/sglang/srt/models/qwen3.py @@ -31,6 +31,10 @@ from sglang.srt.models.qwen2 import Qwen2Model from sglang.srt.models.utils import apply_qk_norm from sglang.srt.server_args import get_global_server_args +from sglang.srt.true_on_policy import ( + should_disable_fused_qk_norm_mrope, + should_force_bfloat16_dense_tensor_math, +) from sglang.srt.utils import add_prefix, get_bool_env_var, is_cuda, is_hip, is_npu from sglang.srt.utils.hf_transformers_utils import get_rope_config @@ -102,16 +106,16 @@ def __init__( self.max_position_embeddings = max_position_embeddings self.tp_rank = get_tensor_model_parallel_rank() - norm_kwargs = ( - dict( - cast_x_before_out_mul=True, - fp32_residual=False, - ) - if get_global_server_args().rl_on_policy_target is not None - else {} + self.q_norm = RMSNorm( + self.head_dim, + eps=rms_norm_eps, + true_on_policy_weight_dtype=torch.float32, + ) + self.k_norm = RMSNorm( + self.head_dim, + eps=rms_norm_eps, + true_on_policy_weight_dtype=torch.float32, ) - self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps, **norm_kwargs) - self.k_norm = RMSNorm(self.head_dim, eps=rms_norm_eps, **norm_kwargs) self.qkv_proj = QKVParallelLinear( hidden_size, @@ -267,14 +271,20 @@ def forward( hidden_states: torch.Tensor, forward_batch: ForwardBatch, ) -> torch.Tensor: - if get_global_server_args().rl_on_policy_target is not None: - hidden_states = hidden_states.bfloat16() + if ( + should_force_bfloat16_dense_tensor_math() + or hidden_states.dtype != self.qkv_proj.weight.dtype + ): + # True-on-policy RMSNorm can produce fp32 activations while dense + # projections remain bf16, including during cuda-graph capture when + # the global on-policy flag is temporarily cleared. + hidden_states = hidden_states.to(self.qkv_proj.weight.dtype) save_kv_cache = True use_aiter_fused = ( self.use_fused_qk_norm_mrope and forward_batch.forward_mode.is_decode() - and get_global_server_args().rl_on_policy_target is None + and not should_disable_fused_qk_norm_mrope() ) if use_aiter_fused: @@ -297,9 +307,9 @@ def forward( forward_batch=forward_batch, ) - if get_global_server_args().rl_on_policy_target is not None: - q = q.to(torch.bfloat16) - k = k.to(torch.bfloat16) + if should_force_bfloat16_dense_tensor_math() or q.dtype != v.dtype: + q = q.to(v.dtype) + k = k.to(v.dtype) attn_output = self.attn(q, k, v, forward_batch, save_kv_cache=save_kv_cache) output, _ = self.o_proj(attn_output) @@ -343,19 +353,19 @@ def __init__( prefix=add_prefix("mlp", prefix), ) - norm_kwargs = ( - dict( - cast_x_before_out_mul=True, - fp32_residual=False, - ) - if get_global_server_args().rl_on_policy_target is not None - else {} - ) self.input_layernorm = RMSNorm( - config.hidden_size, eps=config.rms_norm_eps, **norm_kwargs + config.hidden_size, + eps=config.rms_norm_eps, + true_on_policy_weight_dtype=torch.float32, + true_on_policy_override_orig_dtype=torch.float32, + true_on_policy_fp32_residual=True, ) self.post_attention_layernorm = RMSNorm( - config.hidden_size, eps=config.rms_norm_eps, **norm_kwargs + config.hidden_size, + eps=config.rms_norm_eps, + true_on_policy_weight_dtype=torch.float32, + true_on_policy_override_orig_dtype=torch.float32, + true_on_policy_fp32_residual=True, ) self.layer_scatter_modes = LayerScatterModes.init_new( diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py index 66c8016fa361..fe67985371d0 100644 --- a/python/sglang/srt/models/qwen3_moe.py +++ b/python/sglang/srt/models/qwen3_moe.py @@ -80,6 +80,11 @@ enable_fused_set_kv_buffer, ) from sglang.srt.server_args import get_global_server_args +from sglang.srt.true_on_policy import ( + get_on_policy_rms_norm_kwargs, + is_true_on_policy_enabled, + should_disable_fused_qk_norm_mrope, +) from sglang.srt.utils import ( LazyValue, add_prefix, @@ -329,7 +334,7 @@ def forward_normal( # router_logits: (num_tokens, n_experts) router_logits, _ = self.gate(hidden_states) - if get_global_server_args().rl_on_policy_target is not None: + if is_true_on_policy_enabled(): routing_weights = F.softmax(router_logits, dim=1, dtype=torch.float) routing_weights, selected_experts = torch.topk( routing_weights, self.top_k, dim=-1 @@ -532,7 +537,7 @@ def __init__( ) self.compatible_with_fused_kv_buffer = ( False if isinstance(self.rotary_emb, MRotaryEmbedding) else True - ) and (get_global_server_args().rl_on_policy_target is None) + ) and not is_true_on_policy_enabled() self.compatible_with_fused_qk_norm_rope = not isinstance( self.rotary_emb, MRotaryEmbedding ) and self.head_dim in (64, 128, 256) @@ -547,7 +552,7 @@ def __init__( torch.bfloat16, _yarn_factor != 1.0, ) - and (get_global_server_args().rl_on_policy_target is None) + and not should_disable_fused_qk_norm_mrope() ) self._used_fused_qk_norm_rope_last_call = False @@ -560,14 +565,7 @@ def __init__( prefix=add_prefix("attn", prefix), ) - norm_kwargs = ( - dict( - cast_x_before_out_mul=True, - fp32_residual=False, - ) - if get_global_server_args().rl_on_policy_target is not None - else {} - ) + norm_kwargs = get_on_policy_rms_norm_kwargs(fp32_residual=False) self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps, **norm_kwargs) self.k_norm = RMSNorm(self.head_dim, eps=rms_norm_eps, **norm_kwargs) self.alt_stream = alt_stream @@ -813,14 +811,7 @@ def __init__( quant_config=quant_config, prefix=add_prefix("mlp", prefix), ) - norm_kwargs = ( - dict( - cast_x_before_out_mul=True, - fp32_residual=False, - ) - if get_global_server_args().rl_on_policy_target is not None - else {} - ) + norm_kwargs = get_on_policy_rms_norm_kwargs(fp32_residual=False) self.input_layernorm = RMSNorm( config.hidden_size, eps=config.rms_norm_eps, **norm_kwargs ) diff --git a/python/sglang/srt/models/sdar.py b/python/sglang/srt/models/sdar.py index 70ab59a48980..933098c59e25 100644 --- a/python/sglang/srt/models/sdar.py +++ b/python/sglang/srt/models/sdar.py @@ -44,6 +44,10 @@ enable_fused_set_kv_buffer, ) from sglang.srt.server_args import get_global_server_args +from sglang.srt.true_on_policy import ( + get_on_policy_rms_norm_kwargs, + should_force_bfloat16_dense_tensor_math, +) from sglang.srt.utils import add_prefix, is_cuda, make_layers logger = logging.getLogger(__name__) @@ -206,7 +210,7 @@ def forward( hidden_states: torch.Tensor, forward_batch: ForwardBatch, ): - if get_global_server_args().rl_on_policy_target is not None: + if should_force_bfloat16_dense_tensor_math(): hidden_states = hidden_states.bfloat16() qkv, _ = self.qkv_proj(hidden_states) @@ -234,7 +238,7 @@ def forward( ), ) - if get_global_server_args().rl_on_policy_target is not None: + if should_force_bfloat16_dense_tensor_math(): q = q.to(torch.bfloat16) k = k.to(torch.bfloat16) @@ -262,15 +266,10 @@ def __init__( self.hidden_size = config.hidden_size self.layer_id = layer_id - norm_kwargs = ( - dict( - weight_dtype=torch.float32, - cast_x_before_out_mul=True, - override_orig_dtype=torch.float32, - fp32_residual=True, - ) - if get_global_server_args().rl_on_policy_target is not None - else {} + norm_kwargs = get_on_policy_rms_norm_kwargs( + weight_dtype=torch.float32, + override_orig_dtype=torch.float32, + fp32_residual=True, ) self.input_layernorm = RMSNorm( self.hidden_size, eps=config.rms_norm_eps, **norm_kwargs @@ -386,15 +385,10 @@ def __init__( prefix=add_prefix("layers", prefix), ) if self.pp_group.is_last_rank: - norm_kwargs = ( - dict( - weight_dtype=torch.float32, - cast_x_before_out_mul=True, - override_orig_dtype=torch.float32, - fp32_residual=True, - ) - if get_global_server_args().rl_on_policy_target is not None - else {} + norm_kwargs = get_on_policy_rms_norm_kwargs( + weight_dtype=torch.float32, + override_orig_dtype=torch.float32, + fp32_residual=True, ) self.norm = RMSNorm(self.embed_dim, eps=config.rms_norm_eps, **norm_kwargs) else: diff --git a/python/sglang/srt/models/sdar_moe.py b/python/sglang/srt/models/sdar_moe.py index c09bfeb17dfd..f671c3eda365 100644 --- a/python/sglang/srt/models/sdar_moe.py +++ b/python/sglang/srt/models/sdar_moe.py @@ -62,6 +62,10 @@ enable_fused_set_kv_buffer, ) from sglang.srt.server_args import get_global_server_args +from sglang.srt.true_on_policy import ( + get_on_policy_rms_norm_kwargs, + should_force_bfloat16_dense_tensor_math, +) from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers logger = logging.getLogger(__name__) @@ -284,7 +288,7 @@ def forward( hidden_states: torch.Tensor, forward_batch: ForwardBatch, ) -> torch.Tensor: - if get_global_server_args().rl_on_policy_target is not None: + if should_force_bfloat16_dense_tensor_math(): hidden_states = hidden_states.bfloat16() qkv, _ = self.qkv_proj(hidden_states) @@ -312,7 +316,7 @@ def forward( ), ) - if get_global_server_args().rl_on_policy_target is not None: + if should_force_bfloat16_dense_tensor_math(): q = q.to(torch.bfloat16) k = k.to(torch.bfloat16) @@ -341,15 +345,10 @@ def __init__( self.hidden_size = config.hidden_size self.layer_id = layer_id - norm_kwargs = ( - dict( - weight_dtype=torch.float32, - cast_x_before_out_mul=True, - override_orig_dtype=torch.float32, - fp32_residual=True, - ) - if get_global_server_args().rl_on_policy_target is not None - else {} + norm_kwargs = get_on_policy_rms_norm_kwargs( + weight_dtype=torch.float32, + override_orig_dtype=torch.float32, + fp32_residual=True, ) self.input_layernorm = RMSNorm( self.hidden_size, eps=config.rms_norm_eps, **norm_kwargs @@ -479,15 +478,10 @@ def __init__( ) if self.pp_group.is_last_rank: - norm_kwargs = ( - dict( - weight_dtype=torch.float32, - cast_x_before_out_mul=True, - override_orig_dtype=torch.float32, - fp32_residual=True, - ) - if get_global_server_args().rl_on_policy_target is not None - else {} + norm_kwargs = get_on_policy_rms_norm_kwargs( + weight_dtype=torch.float32, + override_orig_dtype=torch.float32, + fp32_residual=True, ) self.norm = RMSNorm(self.embed_dim, eps=config.rms_norm_eps, **norm_kwargs) else: diff --git a/python/sglang/srt/models/step3p5.py b/python/sglang/srt/models/step3p5.py index 1f3a4d221758..f1a24ee224f9 100644 --- a/python/sglang/srt/models/step3p5.py +++ b/python/sglang/srt/models/step3p5.py @@ -51,6 +51,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.server_args import get_global_server_args +from sglang.srt.true_on_policy import is_true_on_policy_enabled from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers Step3p5Config = None @@ -680,11 +681,7 @@ def __init__( quant_config=quant_config, enable_tp=not is_dp_attention_enabled(), prefix=add_prefix("embed_tokens", prefix), - params_dtype=( - torch.float32 - if get_global_server_args().rl_on_policy_target is not None - else None - ), + params_dtype=torch.float32 if is_true_on_policy_enabled() else None, ) else: self.embed_tokens = PPMissingLayer() diff --git a/python/sglang/srt/models/utils.py b/python/sglang/srt/models/utils.py index 92588e1775e6..a88dc3650c13 100644 --- a/python/sglang/srt/models/utils.py +++ b/python/sglang/srt/models/utils.py @@ -446,6 +446,8 @@ def apply_qk_norm( _is_cuda # TODO(dark): have not tested on ROCm or other backends and allow_inplace # TODO(dark): this can be relaxed if needed and (q_eps == k_eps) # TODO(dark): this can also be relaxed + and q_norm.weight.dtype == q.dtype + and k_norm.weight.dtype == k.dtype and not envs.SGLANG_ENABLE_DETERMINISTIC_INFERENCE.get() and get_global_server_args().piecewise_cuda_graph_compiler != "inductor" # let inductor fuse QK norm diff --git a/python/sglang/srt/multimodal/processors/base_processor.py b/python/sglang/srt/multimodal/processors/base_processor.py index 5b886a8791cb..5f207f9f2bbf 100644 --- a/python/sglang/srt/multimodal/processors/base_processor.py +++ b/python/sglang/srt/multimodal/processors/base_processor.py @@ -18,7 +18,7 @@ MultimodalInputFormat, MultimodalProcessorOutput, ) -from sglang.srt.server_args import get_global_server_args +from sglang.srt.true_on_policy import is_true_on_policy_enabled from sglang.srt.utils import ( envs, is_cpu, @@ -431,7 +431,7 @@ def process_mm_data( and isinstance(processor.image_processor, BaseImageProcessor) and not self.server_args.disable_fast_image_processor ): - if _is_cpu or get_global_server_args().rl_on_policy_target is not None: + if _is_cpu or is_true_on_policy_enabled(): kwargs["device"] = "cpu" elif _is_xpu: kwargs["device"] = "xpu" diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index ca19dc1553e6..4b8bec3fc8d4 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -34,6 +34,10 @@ from sglang.srt.layers.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE from sglang.srt.lora.lora_registry import LoRARef from sglang.srt.parser.reasoning_parser import ReasoningParser +from sglang.srt.true_on_policy.contracts import ( + resolve_true_on_policy_runtime_policy, + validate_true_on_policy_contract, +) from sglang.srt.utils.common import ( LORA_TARGET_ALL_MODULES, SUPPORTED_LORA_TARGET_MODULES, @@ -234,7 +238,7 @@ RADIX_EVICTION_POLICY_CHOICES = ["lru", "lfu", "slru", "priority"] -RL_ON_POLICY_TARGET_CHOICES = ["fsdp"] +RL_ON_POLICY_TARGET_CHOICES = ["fsdp", "fsdp_tp"] LORA_BACKEND_CHOICES = ["triton", "csgmv", "ascend", "torch_native"] @@ -767,7 +771,9 @@ class ServerArgs: scheduler_recv_interval: int = 1 numa_node: Optional[List[int]] = None enable_deterministic_inference: bool = False + enable_prefill_only_deterministic_inference: bool = False rl_on_policy_target: Optional[str] = None + true_on_policy_contract: Optional[str] = None enable_attn_tp_input_scattered: bool = False gc_threshold: Optional[List[int]] = None # Context parallelism used in the long sequence prefill phase of DeepSeek v3.2 @@ -3061,9 +3067,12 @@ def _handle_context_parallelism(self): ), "attn_cp_size != moe_dp_size is only supported when moe_dp_size == 1" def _handle_data_parallelism(self): + use_context_parallel_lm_head = self.enable_dp_lm_head and self.attn_cp_size > 1 + if self.dp_size == 1: self.enable_dp_attention = False - self.enable_dp_lm_head = False + if not use_context_parallel_lm_head: + self.enable_dp_lm_head = False if self.enable_dp_attention: self.schedule_conservativeness = self.schedule_conservativeness * 0.3 @@ -3074,9 +3083,10 @@ def _handle_data_parallelism(self): ) if self.enable_dp_lm_head: - assert ( - self.enable_dp_attention - ), "Please enable dp attention when setting enable_dp_lm_head. " + assert use_context_parallel_lm_head or self.enable_dp_attention, ( + "Please enable dp attention when setting enable_dp_lm_head, " + "unless attention context parallelism is enabled." + ) def _handle_moe_kernel_config(self): if self.quantization == "mxfp8": @@ -4298,16 +4308,40 @@ def _handle_cache_compatibility(self): raise ValueError("--swa-full-tokens-ratio should be in range (0, 1.0].") def _handle_deterministic_inference(self): + validate_true_on_policy_contract(self) + + if self.enable_prefill_only_deterministic_inference: + self.enable_deterministic_inference = True + if self.rl_on_policy_target is not None: logger.warning( - "Enable deterministic inference because of rl_on_policy_target." + "Enable deterministic inference because of legacy rl_on_policy_target." + ) + self.enable_deterministic_inference = True + + if self.true_on_policy_contract is not None: + logger.warning( + "Enable deterministic inference because of true_on_policy_contract." ) self.enable_deterministic_inference = True # For VLM envs.SGLANG_VLM_CACHE_SIZE_MB.set(0) - # TODO remove this environment variable as a whole - envs.SGLANG_ENABLE_DETERMINISTIC_INFERENCE.set(True) + + if ( + resolve_true_on_policy_runtime_policy( + self + ).disable_flashinfer_allreduce_fusion + and self.enable_flashinfer_allreduce_fusion + ): + self.enable_flashinfer_allreduce_fusion = False + logger.warning( + "Disable flashinfer allreduce fusion because of " + "true_on_policy_contract with TP rollout." + ) + + if self.enable_deterministic_inference: + envs.SGLANG_ENABLE_DETERMINISTIC_INFERENCE.set("1") if self.enable_deterministic_inference: if self.enable_aiter_allreduce_fusion: @@ -6763,6 +6797,11 @@ def add_cli_args(parser: argparse.ArgumentParser): action="store_true", help="Enable deterministic inference mode with batch invariant ops.", ) + parser.add_argument( + "--enable-prefill-only-deterministic-inference", + action="store_true", + help="Enable prefill-only deterministic inference mode with batch invariant ops.", + ) parser.add_argument( "--rl-on-policy-target", type=str, @@ -6770,6 +6809,15 @@ def add_cli_args(parser: argparse.ArgumentParser): choices=RL_ON_POLICY_TARGET_CHOICES, help="The training system that SGLang needs to match for true on-policy.", ) + parser.add_argument( + "--true-on-policy-contract", + type=str, + default=ServerArgs.true_on_policy_contract, + help=( + "Internal true-on-policy parity contract selected by the launcher. " + "Normal users should prefer the Miles true_on_policy switch." + ), + ) parser.add_argument( "--enable-attn-tp-input-scattered", action="store_true", diff --git a/python/sglang/srt/tp_invariant_ops/__init__.py b/python/sglang/srt/tp_invariant_ops/__init__.py new file mode 100644 index 000000000000..2a3f7a07f4a7 --- /dev/null +++ b/python/sglang/srt/tp_invariant_ops/__init__.py @@ -0,0 +1,23 @@ +from .tp_invariant_ops import ( + disable_tp_invariant_mode, + enable_tp_invariant_mode, + is_tp_invariant_mode_enabled, + matmul_tp_inv, + matmul_tp_persistent, + moe_sum_tree_reduce, + set_tp_invariant_mode, + tree_all_reduce_sum, +) + +__version__ = "0.1.0" + +__all__ = [ + "matmul_tp_persistent", + "matmul_tp_inv", + "moe_sum_tree_reduce", + "tree_all_reduce_sum", + "set_tp_invariant_mode", + "is_tp_invariant_mode_enabled", + "disable_tp_invariant_mode", + "enable_tp_invariant_mode", +] diff --git a/python/sglang/srt/tp_invariant_ops/tp_invariant_ops.py b/python/sglang/srt/tp_invariant_ops/tp_invariant_ops.py new file mode 100644 index 000000000000..b465ff4a1977 --- /dev/null +++ b/python/sglang/srt/tp_invariant_ops/tp_invariant_ops.py @@ -0,0 +1,1941 @@ +import contextlib +import math +import os +import sys +from typing import Any, Callable, Dict + +import torch +import torch.distributed as dist +import triton +import triton.language as tl + +# Triton's constexpr tree unrolling causes deep AST recursion in the JIT +# compiler. The two-level tree (v2) bounds compilation depth to +# max(log2(SUBTREE), log2(E/SUBTREE)) instead of log2(E), but we still +# need headroom for the per-level AST visitor overhead. +if sys.getrecursionlimit() < 16384: + sys.setrecursionlimit(16384) + + +def _matmul_launch_metadata( + grid: Callable[..., Any], kernel: Any, args: Dict[str, Any] +) -> Dict[str, Any]: + ret = {} + m, n, k = args["M"], args["N"], args["K"] + ret["name"] = f"{kernel.name} [M={m}, N={n}, K={k}]" + if "tiles_per_update" in args: + ret["name"] = ( + f"{kernel.name} [M={m}, N={n}, K={k}, tiles_per_update={args['tiles_per_update']:02}]" + ) + if "c_ptr" in args: + bytes_per_elem = args["c_ptr"].element_size() + else: + bytes_per_elem = 1 if args["FP8_OUTPUT"] else 2 + ret[f"flops{bytes_per_elem * 8}"] = 2.0 * m * n * k + ret["bytes"] = bytes_per_elem * (m * k + n * k + m * n) + return ret + + +@triton.jit +def _compute_pid(tile_id, num_pid_in_group, num_pid_m, GROUP_SIZE_M, NUM_SMS): + group_id = tile_id // num_pid_in_group + first_pid_m = group_id * GROUP_SIZE_M + group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M) + pid_m = first_pid_m + (tile_id % group_size_m) + pid_n = (tile_id % num_pid_in_group) // group_size_m + return pid_m, pid_n + + +def _get_tl_dtype(dtype): + if dtype == torch.float32: + return tl.float32 + elif dtype == torch.float16: + return tl.float16 + elif dtype == torch.bfloat16: + return tl.bfloat16 + + +# ---- kernel ---- +@triton.jit(launch_metadata=_matmul_launch_metadata) +def matmul_kernel_tp_persistent( + A_ptr, + B_ptr, + C_ptr, + M, + N, + K, + stride_am, + stride_ak, + stride_bk, + stride_bn, + stride_cm, + stride_cn, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, + BLOCK_K: tl.constexpr, + GROUP_SIZE_M: tl.constexpr, + NUM_SMS: tl.constexpr, + LEVEL_K: tl.constexpr, + TILE_K: tl.constexpr, + FIRST_LEVEL_BLOCK: tl.constexpr, + NEXT_POWER_OF_LEVEL: tl.constexpr, + NEXT_POWER_OF_REMAIN_LEVEL: tl.constexpr, + ACC_DTYPE: tl.constexpr, + OUT_DTYPE: tl.constexpr, + A_LARGE: tl.constexpr, + B_LARGE: tl.constexpr, + C_LARGE: tl.constexpr, +): + pid = tl.program_id(axis=0) + num_pid_m = tl.cdiv(M, BLOCK_M) + num_pid_n = tl.cdiv(N, BLOCK_N) + num_tiles = num_pid_m * num_pid_n + + num_pid_in_group = GROUP_SIZE_M * num_pid_n + + manual_acc = 3 + acc1 = tl.zeros((BLOCK_M, BLOCK_N), dtype=ACC_DTYPE) + acc2 = tl.zeros((BLOCK_M, BLOCK_N), dtype=ACC_DTYPE) + acc3 = tl.zeros((BLOCK_M, BLOCK_N), dtype=ACC_DTYPE) + + S = tl.zeros((NEXT_POWER_OF_REMAIN_LEVEL, BLOCK_M, BLOCK_N), dtype=ACC_DTYPE) + + S_mask = tl.arange(0, NEXT_POWER_OF_REMAIN_LEVEL)[:, None, None] + level_ids = tl.arange(0, NEXT_POWER_OF_LEVEL) + + base_offs_m = tl.arange(0, BLOCK_M) + base_offs_n = tl.arange(0, BLOCK_N) + offs_k = tl.arange(0, BLOCK_K) + offs_k = tl.max_contiguous(tl.multiple_of(offs_k, BLOCK_K), BLOCK_K) + + for tile_id in tl.range(pid, num_tiles, NUM_SMS, flatten=False): + pid_m, pid_n = _compute_pid( + tile_id, num_pid_in_group, num_pid_m, GROUP_SIZE_M, NUM_SMS + ) + + start_m = pid_m * BLOCK_M + start_n = pid_n * BLOCK_N + + offs_am = start_m + base_offs_m + if A_LARGE: + offs_am = offs_am.to(tl.int64) + offs_am = tl.where(offs_am < M, offs_am, 0) + + offs_bn = start_n + base_offs_n + if B_LARGE: + offs_bn = offs_bn.to(tl.int64) + offs_bn = tl.where(offs_bn < N, offs_bn, 0) + + offs_am = tl.max_contiguous(tl.multiple_of(offs_am, BLOCK_M), BLOCK_M) + offs_bn = tl.max_contiguous(tl.multiple_of(offs_bn, BLOCK_N), BLOCK_N) + + a_ptrs = A_ptr + offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak + b_ptrs = B_ptr + offs_bn[None, :] * stride_bn + offs_k[:, None] * stride_bk + + count = tl.zeros((NEXT_POWER_OF_LEVEL,), dtype=tl.int32) + + acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=ACC_DTYPE) + for s_tile_idx in range(0, TILE_K): + k0 = s_tile_idx * BLOCK_K + a = tl.load( + a_ptrs, + mask=(offs_am[:, None] < M) & ((k0 + offs_k)[None, :] < K), + other=0.0, + ) + b = tl.load( + b_ptrs, + mask=((k0 + offs_k)[:, None] < K) & (offs_bn[None, :] < N), + other=0.0, + ) + a_ptrs += BLOCK_K * stride_ak + b_ptrs += BLOCK_K * stride_bk + + acc = tl.dot(a, b).to(ACC_DTYPE) + + break_flag = 0 + for level in range(LEVEL_K): + if break_flag == 0: + idx_mask = level_ids == level + + count_value_added = tl.sum(count * idx_mask) + 1 + + table_value = FIRST_LEVEL_BLOCK if level == 0 else 2 + + carry_over = (table_value == count_value_added).to(tl.int1) + + if count_value_added > 1: + if level == 0: + acc = acc1 + acc + elif level == 1: + acc = acc2 + acc + elif level == 2: + acc = acc3 + acc + else: + tmp_acc_mask = S_mask == (level - manual_acc) + acc = ( + tl.sum(S * tmp_acc_mask, axis=0, dtype=ACC_DTYPE) + acc + ) + + count = tl.where( + idx_mask, count_value_added * (1 - carry_over), count + ) + if not carry_over: + break_flag = 1 + if level == 0: + acc1 = acc + elif level == 1: + acc2 = acc + elif level == 2: + acc3 = acc + else: + tmp_acc_mask = S_mask == (level - manual_acc) + S = tl.where(tmp_acc_mask, acc[None, :, :], S) + + c_ptr = C_ptr + (offs_am[:, None] * stride_cm + offs_bn[None, :] * stride_cn) + offs_cm = start_m + base_offs_m + offs_cn = start_n + base_offs_n + if C_LARGE: + offs_cm = offs_cm.to(tl.int64) + offs_cn = offs_cn.to(tl.int64) + offs_cm = tl.where(offs_cm < M, offs_cm, 0) + offs_cn = tl.where(offs_cn < N, offs_cn, 0) + offs_cm = tl.max_contiguous(tl.multiple_of(offs_cm, BLOCK_M), BLOCK_M) + offs_cn = tl.max_contiguous(tl.multiple_of(offs_cn, BLOCK_N), BLOCK_N) + mask_c = (offs_cm[:, None] < M) & (offs_cn[None, :] < N) + tl.store(c_ptr, acc.to(OUT_DTYPE), mask=mask_c) + + +@triton.jit(launch_metadata=_matmul_launch_metadata) +def matmul_kernel_tp_persistent_optim( + A_ptr, + B_ptr, + C_ptr, + M, + N, + K, + stride_am, + stride_ak, + stride_bk, + stride_bn, + stride_cm, + stride_cn, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, + BLOCK_K: tl.constexpr, + GROUP_SIZE_M: tl.constexpr, + NUM_SMS: tl.constexpr, + LEVEL_K: tl.constexpr, + TILE_K: tl.constexpr, + FIRST_LEVEL_BLOCK: tl.constexpr, + NEXT_POWER_OF_LEVEL: tl.constexpr, + NEXT_POWER_OF_REMAIN_LEVEL: tl.constexpr, + ACC_DTYPE: tl.constexpr, + OUT_DTYPE: tl.constexpr, + A_LARGE: tl.constexpr, + B_LARGE: tl.constexpr, + C_LARGE: tl.constexpr, +): + pid = tl.program_id(axis=0) + num_pid_m = tl.cdiv(M, BLOCK_M) + num_pid_n = tl.cdiv(N, BLOCK_N) + num_tiles = num_pid_m * num_pid_n + + num_pid_in_group = GROUP_SIZE_M * num_pid_n + + # Flat loop (same structure as original) with inlined level-0 handling. + # Level-0 uses scalar counter; tree merge (levels 1+) only on carry. + # Accumulators are conditionally allocated based on LEVEL_K (constexpr) + # to minimize register pressure and maximize occupancy. + acc1 = tl.zeros((BLOCK_M, BLOCK_N), dtype=ACC_DTYPE) + acc2 = tl.zeros((BLOCK_M, BLOCK_N), dtype=ACC_DTYPE) + acc3 = tl.zeros((BLOCK_M, BLOCK_N), dtype=ACC_DTYPE) + acc4 = tl.zeros((BLOCK_M, BLOCK_N), dtype=ACC_DTYPE) + acc5 = tl.zeros((BLOCK_M, BLOCK_N), dtype=ACC_DTYPE) + + S = tl.zeros((NEXT_POWER_OF_REMAIN_LEVEL, BLOCK_M, BLOCK_N), dtype=ACC_DTYPE) + S_mask = tl.arange(0, NEXT_POWER_OF_REMAIN_LEVEL)[:, None, None] + level_ids = tl.arange(0, NEXT_POWER_OF_LEVEL) + + offs_k = tl.arange(0, BLOCK_K) + offs_k = tl.max_contiguous(tl.multiple_of(offs_k, BLOCK_K), BLOCK_K) + + for tile_id in tl.range(pid, num_tiles, NUM_SMS, flatten=False): + pid_m, pid_n = _compute_pid( + tile_id, num_pid_in_group, num_pid_m, GROUP_SIZE_M, NUM_SMS + ) + + start_m = pid_m * BLOCK_M + start_n = pid_n * BLOCK_N + + offs_am = start_m + tl.arange(0, BLOCK_M) + mask_m = offs_am < M + if A_LARGE: + offs_am = offs_am.to(tl.int64) + offs_am = tl.where(mask_m, offs_am, 0) + + offs_bn = start_n + tl.arange(0, BLOCK_N) + mask_n = offs_bn < N + if B_LARGE: + offs_bn = offs_bn.to(tl.int64) + offs_bn = tl.where(mask_n, offs_bn, 0) + + offs_am = tl.max_contiguous(tl.multiple_of(offs_am, BLOCK_M), BLOCK_M) + offs_bn = tl.max_contiguous(tl.multiple_of(offs_bn, BLOCK_N), BLOCK_N) + mask_m_bc = mask_m[:, None] + mask_n_bc = mask_n[None, :] + + a_ptrs = A_ptr + offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak + b_ptrs = B_ptr + offs_bn[None, :] * stride_bn + offs_k[:, None] * stride_bk + + c0 = 0 + c1 = 0 + c2 = 0 + c3 = 0 + c4 = 0 + count = tl.zeros((NEXT_POWER_OF_LEVEL,), dtype=tl.int32) + + acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=ACC_DTYPE) + for _ in range(0, TILE_K): + a = tl.load(a_ptrs, mask=mask_m_bc, other=0.0) + b = tl.load(b_ptrs, mask=mask_n_bc, other=0.0) + a_ptrs += BLOCK_K * stride_ak + b_ptrs += BLOCK_K * stride_bk + + acc = tl.dot(a, b).to(ACC_DTYPE) + + c0 += 1 + if c0 > 1: + acc = acc1 + acc + if c0 < FIRST_LEVEL_BLOCK: + acc1 = acc + else: + c0 = 0 + break_flag = 0 + for level in range(1, LEVEL_K): + if break_flag == 0: + if level == 1: + c1 += 1 + if c1 > 1: + acc = acc2 + acc + if c1 == 2: + c1 = 0 + else: + acc2 = acc + break_flag = 1 + elif level == 2: + c2 += 1 + if c2 > 1: + acc = acc3 + acc + if c2 == 2: + c2 = 0 + else: + acc3 = acc + break_flag = 1 + elif level == 3: + c3 += 1 + if c3 > 1: + acc = acc4 + acc + if c3 == 2: + c3 = 0 + else: + acc4 = acc + break_flag = 1 + elif level == 4: + c4 += 1 + if c4 > 1: + acc = acc5 + acc + if c4 == 2: + c4 = 0 + else: + acc5 = acc + break_flag = 1 + else: + idx_mask = level_ids == level + count_value_added = tl.sum(count * idx_mask) + 1 + carry_over = (2 == count_value_added).to(tl.int1) + if count_value_added > 1: + tmp_acc_mask = S_mask == (level - 5) + acc = ( + tl.sum(S * tmp_acc_mask, axis=0, dtype=ACC_DTYPE) + + acc + ) + count = tl.where( + idx_mask, count_value_added * (1 - carry_over), count + ) + if not carry_over: + break_flag = 1 + tmp_acc_mask = S_mask == (level - 5) + S = tl.where(tmp_acc_mask, acc[None, :, :], S) + + offs_cm = start_m + tl.arange(0, BLOCK_M) + offs_cn = start_n + tl.arange(0, BLOCK_N) + if C_LARGE: + offs_cm = offs_cm.to(tl.int64) + offs_cn = offs_cn.to(tl.int64) + offs_cm = tl.where(mask_m, offs_cm, 0) + offs_cn = tl.where(mask_n, offs_cn, 0) + offs_cm = tl.max_contiguous(tl.multiple_of(offs_cm, BLOCK_M), BLOCK_M) + offs_cn = tl.max_contiguous(tl.multiple_of(offs_cn, BLOCK_N), BLOCK_N) + c_ptr = C_ptr + (offs_cm[:, None] * stride_cm + offs_cn[None, :] * stride_cn) + mask_c = mask_m_bc & mask_n_bc + tl.store(c_ptr, acc.to(OUT_DTYPE), mask=mask_c) + + +def _matmul_tp_persistent_impl( + A: torch.Tensor, + B: torch.Tensor, + bias: torch.Tensor = None, + fp32_accum: bool = False, + use_optim_kernel: bool = False, +): + assert A.shape[-1] == B.shape[-2], "Dim doesn't match" + + out_dtype = A.dtype + acc_dtype = torch.float32 if fp32_accum else A.dtype + + NUM_SMS = torch.cuda.get_device_properties(A.device).multi_processor_count + + # 1D launch kernel where each block gets its own program. + def grid(META): + return ( + min( + NUM_SMS, + triton.cdiv(M, META["BLOCK_M"]) * triton.cdiv(N, META["BLOCK_N"]), + ), + ) + + base_configs = { + torch.bfloat16: { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 8, + "num_stages": 2, + "num_warps": 8, + }, + torch.float16: { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 8, + "num_stages": 2, + "num_warps": 8, + }, + torch.float32: { + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 8, + "num_stages": 2, + "num_warps": 8, + }, + } + + configs = base_configs + BLOCK_M = configs[out_dtype]["BLOCK_SIZE_M"] + BLOCK_N = configs[out_dtype]["BLOCK_SIZE_N"] + BLOCK_K = configs[out_dtype]["BLOCK_SIZE_K"] + GROUP_SIZE_M = configs[out_dtype]["GROUP_SIZE_M"] + num_stages = configs[out_dtype]["num_stages"] + num_warps = configs[out_dtype]["num_warps"] + + M, K = A.shape + _, N = B.shape + assert ( + K % BLOCK_K == 0 + ), f"Dimension K should be divisible by BLOCK_K. Got K={K}, BLOCK_K={BLOCK_K}." + T = K // BLOCK_K + FIRST_LEVEL_BLOCK = T + + if use_optim_kernel: + num_n_tiles = triton.cdiv(N, BLOCK_N) + total_tiles = triton.cdiv(M, BLOCK_M) * num_n_tiles + + if total_tiles * 4 <= NUM_SMS: + while BLOCK_M > 16 and triton.cdiv(M, BLOCK_M) * num_n_tiles < NUM_SMS: + BLOCK_M //= 2 + elif total_tiles * 2 <= NUM_SMS: + while BLOCK_M > 32 and triton.cdiv(M, BLOCK_M) * num_n_tiles < NUM_SMS: + BLOCK_M //= 2 + + if out_dtype in (torch.bfloat16, torch.float16): + num_warps = 4 if BLOCK_M <= 16 else 8 + num_stages = 3 if K >= 1024 else 2 + + LEVEL_K = 1 + while FIRST_LEVEL_BLOCK > 2 and FIRST_LEVEL_BLOCK % 2 == 0: + FIRST_LEVEL_BLOCK //= 2 + LEVEL_K += 1 + + C = torch.empty((M, N), device=A.device, dtype=out_dtype) + + # Original kernel manually handles levels 0-2 (3 accumulators); + # optim kernel handles levels 0-4 (5 register accumulators), + # so S tensor is smaller / unused for common LEVEL_K <= 5. + manual_acc = 5 if use_optim_kernel else 3 + + NEXT_POWER_OF_LEVEL = 2 ** math.ceil(math.log2(LEVEL_K)) + NEXT_POWER_OF_REMAIN_LEVEL = ( + 2 ** math.ceil(math.log2(LEVEL_K - manual_acc)) if LEVEL_K > manual_acc else 1 + ) + + kernel = ( + matmul_kernel_tp_persistent_optim + if use_optim_kernel + else matmul_kernel_tp_persistent + ) + kernel[grid]( + A, + B, + C, + M, + N, + K, + *A.stride(), + *B.stride(), + *C.stride(), + BLOCK_M=BLOCK_M, + BLOCK_N=BLOCK_N, + BLOCK_K=BLOCK_K, + GROUP_SIZE_M=GROUP_SIZE_M, + NUM_SMS=NUM_SMS, + NEXT_POWER_OF_REMAIN_LEVEL=NEXT_POWER_OF_REMAIN_LEVEL, + LEVEL_K=LEVEL_K, + TILE_K=T, + FIRST_LEVEL_BLOCK=FIRST_LEVEL_BLOCK, + NEXT_POWER_OF_LEVEL=NEXT_POWER_OF_LEVEL, + ACC_DTYPE=_get_tl_dtype(acc_dtype), + OUT_DTYPE=_get_tl_dtype(out_dtype), + A_LARGE=A.numel() > 2**31, + B_LARGE=B.numel() > 2**31, + C_LARGE=C.numel() > 2**31, + num_warps=num_warps, + num_stages=num_stages, + ) + if bias is not None: + C += bias + return C + + +def matmul_tp_persistent( + A: torch.Tensor, + B: torch.Tensor, + bias: torch.Tensor = None, + fp32_accum: bool = False, +): + return _matmul_tp_persistent_impl( + A=A, + B=B, + bias=bias, + fp32_accum=fp32_accum, + use_optim_kernel=False, + ) + + +def matmul_tp_inv( + A: torch.Tensor, + B: torch.Tensor, + bias: torch.Tensor = None, + fp32_accum: bool = False, +): + return matmul_tp_persistent(A, B, bias=bias, fp32_accum=fp32_accum) + + +def matmul_tp_persistent_optim( + A: torch.Tensor, + B: torch.Tensor, + bias: torch.Tensor = None, + fp32_accum: bool = False, +): + return _matmul_tp_persistent_impl( + A=A, + B=B, + bias=bias, + fp32_accum=fp32_accum, + use_optim_kernel=True, + ) + + +def tree_all_reduce_sum(x: torch.Tensor, device_group=None) -> torch.Tensor: + rank = dist.get_rank(device_group) + world_size = dist.get_world_size(device_group) + + if world_size & (world_size - 1) != 0: + raise ValueError( + "world_size must be a power of 2 in order to use all_reduce_sum." + ) + + result = [torch.zeros_like(x) for _ in range(world_size)] + dist.all_gather(result, x, group=device_group) + + for level in range(1, world_size.bit_length()): + for left in range(0, world_size, 1 << level): + right = left + (1 << (level - 1)) + result[left] += result[right] + + return result[0] + + +def tree_all_reduce_sum_optim(x: torch.Tensor, device_group=None) -> torch.Tensor: + if not x.is_cuda: + raise ValueError("x must be a CUDA tensor.") + if not x.is_contiguous(): + raise ValueError( + "x must be contiguous. Call x = x.contiguous() OUTSIDE graph capture." + ) + + world_size = dist.get_world_size(device_group) + if world_size & (world_size - 1) != 0: + raise ValueError("world_size must be a power of 2.") + + # cache + if not hasattr(tree_all_reduce_sum, "_cache"): + tree_all_reduce_sum._cache = {} + + key = (id(device_group), x.device.index, tuple(x.shape), x.dtype, world_size) + st = tree_all_reduce_sum._cache.get(key) + if st is None: + gather = torch.empty( + (world_size,) + tuple(x.shape), device=x.device, dtype=x.dtype + ) + out = torch.empty_like(x) + st = tree_all_reduce_sum._cache[key] = (gather, out) + + gather, out = st + + # 1) all_gather into one contiguous buffer + dist.all_gather_into_tensor(gather, x, group=device_group) + + # 2) deterministic tree pairing EXACTLY like your original: + # for level in range(1, bit_length): + # for left in range(0, world_size, 1<> 1 + # Views only (no alloc); one add_ kernel per level + gather[0:world_size:step].add_(gather[half:world_size:step]) + + out.copy_(gather[0]) + tree_all_reduce_sum._cache.clear() + return out + + +_tp_inv_MODE = False + +try: + def_lib = torch.library.Library("tp_inv_ops", "DEF") + def_lib.define("matmul_tp_inv(Tensor a, Tensor b, Tensor? bias=None) -> Tensor") +except RuntimeError: + pass + +try: + impl = torch.library.Library("tp_inv_ops", "IMPL") + + impl.impl("matmul_tp_inv", matmul_tp_persistent, "CUDA") +except RuntimeError: + pass + + +def is_tp_invariant_mode_enabled(): + return _tp_inv_MODE + + +def enable_tp_invariant_mode(): + global _tp_inv_MODE + + if _tp_inv_MODE: + return + + _tp_inv_MODE = True + + +def disable_tp_invariant_mode(): + global _tp_inv_MODE + + _tp_inv_MODE = False + + +@contextlib.contextmanager +def set_tp_invariant_mode(enabled=True): + global _tp_inv_MODE + + old_state = _tp_inv_MODE + + if enabled: + enable_tp_invariant_mode() + else: + disable_tp_invariant_mode() + + try: + yield + finally: + _tp_inv_MODE = old_state + + +def scatter_input_by_local_expert( + topk: torch.Tensor, input: torch.Tensor, E: int +) -> torch.Tensor: + """ + Args: + topk: [M, topk], long, -1 means remote expert + input: [M, topk, hidden_size], float + E: int, number of local experts (expert ids in [0, E)) + Returns: + output: [M, E, hidden_size] + """ + M, _, hidden_size = input.shape + + # Mask out remote experts in output + valid = (topk != -1).unsqueeze(-1) # [M, topk, 1] + output_masked = input * valid.to(input.dtype) # [M, topk, hidden_size] + + # Replace -1 with 0 for safe indexing (value doesn't matter because output is zero) + topk_index = topk.clamp(min=0) # turns -1 into 0, leaves others unchanged + + # Expand index to match output + index = topk_index.unsqueeze(-1).expand( + -1, -1, hidden_size + ) # [M, topk, hidden_size] + + # Initialize result + output = torch.zeros(M, E, hidden_size, device=input.device, dtype=input.dtype) + + # Scatter add + output.scatter_add_(1, index, output_masked) + + return output + + +@triton.jit +def _load_expert_tile( + input_ptr, + input_base, + input_stride_1, + topk_ids_base, + topk_ids_stride_1, + zero_ptrs, + offs_dim, + mask, + mask_token, + e: tl.constexpr, + TOPK: tl.constexpr, + BLOCK_M: tl.constexpr, +): + """Load tile for expert e: search topk_ids for slot, redirect to zero_buf if remote.""" + found_slot = tl.full([BLOCK_M], -1, dtype=tl.int32) + for k in range(TOPK): + kid = tl.load( + topk_ids_base + k * topk_ids_stride_1, + mask=mask_token, + other=-1, + ).to(tl.int32) + match = (kid == e) & (found_slot == -1) + found_slot = tl.where(match, k, found_slot) + + is_valid = found_slot != -1 + slot_safe = tl.maximum(found_slot, 0) + input_ptrs = input_base + slot_safe[:, None] * input_stride_1 + offs_dim[None, :] + load_ptrs = tl.where(is_valid[:, None], input_ptrs, zero_ptrs) + return tl.load(load_ptrs, mask=mask, other=0.0).to(tl.float32) + + +@triton.jit +def _tree_reduce_pair( + input_ptr, + input_base, + input_stride_1, + topk_ids_base, + topk_ids_stride_1, + zero_ptrs, + offs_dim, + mask, + mask_token, + start: tl.constexpr, + size: tl.constexpr, + TOPK: tl.constexpr, + BLOCK_M: tl.constexpr, +): + """ + Recursively compute binary-tree sum over experts [start, start+size). + Returns a fp32 tile [BLOCK_M, BLOCK_DIM]. + + Tree structure (e.g. size=4, start=0): + _tree_reduce_pair(0, 4) + = _tree_reduce_pair(0, 2) + _tree_reduce_pair(2, 2) + = (_tree_reduce_pair(0,1) + _tree_reduce_pair(1,1)) + + (_tree_reduce_pair(2,1) + _tree_reduce_pair(3,1)) + = (load(e0) + load(e1)) + (load(e2) + load(e3)) + + Since size is constexpr and always a power of 2, Triton fully unrolls this + into a fixed sequence of loads and adds with no dynamic branching. + Max register depth = log2(E) tiles, e.g. E=64 -> 6 tiles. + """ + if size == 1: + return _load_expert_tile( + input_ptr, + input_base, + input_stride_1, + topk_ids_base, + topk_ids_stride_1, + zero_ptrs, + offs_dim, + mask, + mask_token, + start, + TOPK, + BLOCK_M, + ) + else: + half: tl.constexpr = size // 2 + left = _tree_reduce_pair( + input_ptr, + input_base, + input_stride_1, + topk_ids_base, + topk_ids_stride_1, + zero_ptrs, + offs_dim, + mask, + mask_token, + start, + half, + TOPK, + BLOCK_M, + ) + right = _tree_reduce_pair( + input_ptr, + input_base, + input_stride_1, + topk_ids_base, + topk_ids_stride_1, + zero_ptrs, + offs_dim, + mask, + mask_token, + start + half, + half, + TOPK, + BLOCK_M, + ) + return left + right + + +@triton.jit +def _fused_tree_reduce_kernel( + # input: [M, topk, hidden_dim], contiguous + input_ptr, + input_stride_0, + input_stride_1, + # topk_ids: [M, topk], -1 means remote + topk_ids_ptr, + topk_ids_stride_0, + topk_ids_stride_1, + # zero_buf: [hidden_dim], all zeros + zero_buf_ptr, + # output: [M, hidden_dim] + output_ptr, + output_stride_0, + # scalars + token_num, + hidden_dim, + routed_scaling_factor, + # constexpr + E: tl.constexpr, + TOPK: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_DIM: tl.constexpr, +): + """ + Fused scatter + binary-tree reduce in a single kernel. + Zero extra memory: all intermediate results live in registers. + + For each token block, loads expert tiles on-the-fly (with zero-buf redirect + for remote experts), and reduces them in binary-tree order using recursive + constexpr unrolling. The tree structure is: + result = tree_sum(0, E) + tree_sum(s, n) = tree_sum(s, n/2) + tree_sum(s+n/2, n/2) if n > 1 + tree_sum(s, 1) = load_expert(s) base case + + Register pressure: log2(E) tiles of [BLOCK_M, BLOCK_DIM] fp32. + E.g. E=64, BLOCK_M=1, BLOCK_DIM=2048 -> 6 * 8KB = 48KB, well within limits. + """ + input_stride_0 = tl.cast(input_stride_0, dtype=tl.int64) + input_stride_1 = tl.cast(input_stride_1, dtype=tl.int64) + topk_ids_stride_0 = tl.cast(topk_ids_stride_0, dtype=tl.int64) + topk_ids_stride_1 = tl.cast(topk_ids_stride_1, dtype=tl.int64) + output_stride_0 = tl.cast(output_stride_0, dtype=tl.int64) + + token_block_id = tl.program_id(0) + dim_block_id = tl.program_id(1) + + offs_token = token_block_id * BLOCK_M + tl.arange(0, BLOCK_M) + offs_dim = dim_block_id * BLOCK_DIM + tl.arange(0, BLOCK_DIM) + + mask_token = offs_token < token_num + mask_dim = offs_dim < hidden_dim + mask = mask_token[:, None] & mask_dim[None, :] + + zero_ptrs = zero_buf_ptr + offs_dim[None, :] + input_base = input_ptr + offs_token[:, None] * input_stride_0 + topk_ids_base = topk_ids_ptr + offs_token * topk_ids_stride_0 + + # Binary tree reduce over all E experts, entirely in registers. + result = _tree_reduce_pair( + input_ptr, + input_base, + input_stride_1, + topk_ids_base, + topk_ids_stride_1, + zero_ptrs, + offs_dim, + mask, + mask_token, + 0, + E, + TOPK, + BLOCK_M, + ) + + result *= routed_scaling_factor + + store_ptrs = output_ptr + offs_token[:, None] * output_stride_0 + offs_dim[None, :] + tl.store(store_ptrs, result.to(input_ptr.dtype.element_ty), mask=mask) + + +# Persistent zero buffer: only [hidden_dim] elements, negligible memory. +# Allocated once on first call, never freed. +_zero_buf_cache: torch.Tensor | None = None + + +def moe_sum_tree_reduce_v1( + input: torch.Tensor, # [M, topk, hidden_dim] + output: torch.Tensor, # [M, hidden_dim] + curr_topk_ids: torch.Tensor, # [M, topk], -1 means remote + routed_scaling_factor: float, + E: int, +): + """ + Fused MoE tree reduce: zero extra memory, CUDA Graph safe. + + Single kernel: loads expert tiles on-the-fly with zero-buf pointer redirect + for remote experts (L1 cache hit), reduces in binary-tree order entirely + in registers. No scratch buffer needed. + + Invariant guarantee: binary tree reduce order is fixed by expert id, + identical across all EP ranks regardless of which experts are local. + + Memory overhead: only a single [hidden_dim] zero buffer (~14KB for H=7168 bf16). + """ + assert input.is_contiguous() + assert output.is_contiguous() + + token_num, topk, hidden_dim = input.shape + assert output.shape[0] == token_num and output.shape[1] == hidden_dim + assert (E & (E - 1)) == 0, f"E must be power of 2, got {E}" + + # Fast path for the dominant K=8 case: avoids generic expert-tree loads/scans. + if topk == 8: + _moe_sum_tree_reduce_k8_fast_path( + input=input, + output=output, + curr_topk_ids=curr_topk_ids, + routed_scaling_factor=routed_scaling_factor, + E=E, + ) + return + + # Zero buffer: [hidden_dim], allocated once, reused forever + global _zero_buf_cache + if ( + _zero_buf_cache is None + or _zero_buf_cache.device != input.device + or _zero_buf_cache.dtype != input.dtype + or _zero_buf_cache.numel() < hidden_dim + ): + _zero_buf_cache = torch.zeros( + hidden_dim, device=input.device, dtype=input.dtype + ) + zero_buf = _zero_buf_cache + + BLOCK_M = 1 + BLOCK_DIM = 2048 + num_warps = 16 + + grid = ( + triton.cdiv(token_num, BLOCK_M), + triton.cdiv(hidden_dim, BLOCK_DIM), + ) + + _fused_tree_reduce_kernel[grid]( + input, + input.stride(0), + input.stride(1), + curr_topk_ids, + curr_topk_ids.stride(0), + curr_topk_ids.stride(1), + zero_buf, + output, + output.stride(0), + token_num=token_num, + hidden_dim=hidden_dim, + routed_scaling_factor=routed_scaling_factor, + E=E, + TOPK=topk, + BLOCK_M=BLOCK_M, + BLOCK_DIM=BLOCK_DIM, + num_warps=num_warps, + ) + return + + +import math + +import torch +import triton +import triton.language as tl + + +@triton.jit +def _moe_sum_tree_reduce_k8_fused_kernel_opt2d( + x_ptr, + ids_ptr, + out_ptr, + sx_m: tl.constexpr, + sx_k: tl.constexpr, + sx_h: tl.constexpr, # x: [M,8,H] + sid_m: tl.constexpr, + sid_k: tl.constexpr, # ids: [M,8] + so_m: tl.constexpr, + so_h: tl.constexpr, # out: [M,H] + M, + H, # runtime OK + E_LEVEL, # runtime int (log2(E)) + routed_scaling_factor, # runtime scalar + BLOCK_M: tl.constexpr, + BLOCK_H: tl.constexpr, +): + pid_m = tl.program_id(0) + pid_h = tl.program_id(1) + + m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) + h = pid_h * BLOCK_H + tl.arange(0, BLOCK_H) + + mask_m = m < M + mask_h = h < H + mask_mh = mask_m[:, None] & mask_h[None, :] + + # ---- load ids (int32), -1 means remote ---- + ids0 = tl.load(ids_ptr + m * sid_m + 0 * sid_k, mask=mask_m, other=-1).to(tl.int32) + ids1 = tl.load(ids_ptr + m * sid_m + 1 * sid_k, mask=mask_m, other=-1).to(tl.int32) + ids2 = tl.load(ids_ptr + m * sid_m + 2 * sid_k, mask=mask_m, other=-1).to(tl.int32) + ids3 = tl.load(ids_ptr + m * sid_m + 3 * sid_k, mask=mask_m, other=-1).to(tl.int32) + ids4 = tl.load(ids_ptr + m * sid_m + 4 * sid_k, mask=mask_m, other=-1).to(tl.int32) + ids5 = tl.load(ids_ptr + m * sid_m + 5 * sid_k, mask=mask_m, other=-1).to(tl.int32) + ids6 = tl.load(ids_ptr + m * sid_m + 6 * sid_k, mask=mask_m, other=-1).to(tl.int32) + ids7 = tl.load(ids_ptr + m * sid_m + 7 * sid_k, mask=mask_m, other=-1).to(tl.int32) + + # h comes from tl.arange and can be annotated for alignment. + tl.multiple_of(h, 8) + tl.max_contiguous(h, BLOCK_H) + + # m is also a tensor. + tl.multiple_of(m, BLOCK_M) + + # ---- load values: remote handled by mask->other=0 (NO zero_tile / NO tl.where big tiles) ---- + m0 = mask_mh & (ids0 != -1)[:, None] + m1 = mask_mh & (ids1 != -1)[:, None] + m2 = mask_mh & (ids2 != -1)[:, None] + m3 = mask_mh & (ids3 != -1)[:, None] + m4 = mask_mh & (ids4 != -1)[:, None] + m5 = mask_mh & (ids5 != -1)[:, None] + m6 = mask_mh & (ids6 != -1)[:, None] + m7 = mask_mh & (ids7 != -1)[:, None] + + v0 = tl.load( + x_ptr + m[:, None] * sx_m + 0 * sx_k + h[None, :] * sx_h, mask=m0, other=0.0 + ) + v1 = tl.load( + x_ptr + m[:, None] * sx_m + 1 * sx_k + h[None, :] * sx_h, mask=m1, other=0.0 + ) + v2 = tl.load( + x_ptr + m[:, None] * sx_m + 2 * sx_k + h[None, :] * sx_h, mask=m2, other=0.0 + ) + v3 = tl.load( + x_ptr + m[:, None] * sx_m + 3 * sx_k + h[None, :] * sx_h, mask=m3, other=0.0 + ) + v4 = tl.load( + x_ptr + m[:, None] * sx_m + 4 * sx_k + h[None, :] * sx_h, mask=m4, other=0.0 + ) + v5 = tl.load( + x_ptr + m[:, None] * sx_m + 5 * sx_k + h[None, :] * sx_h, mask=m5, other=0.0 + ) + v6 = tl.load( + x_ptr + m[:, None] * sx_m + 6 * sx_k + h[None, :] * sx_h, mask=m6, other=0.0 + ) + v7 = tl.load( + x_ptr + m[:, None] * sx_m + 7 * sx_k + h[None, :] * sx_h, mask=m7, other=0.0 + ) + + x_dtype = x_ptr.dtype.element_ty + + # ---- deterministic dense-tree-equivalent reduce (same order concept as baseline) ---- + # Remote entries have already been masked to 0.0 at load time (m0..m7). + + for bit in tl.range(0, E_LEVEL): + bitmask = 1 << bit + + # ========== lane0 as source ========== + cond = (ids0 != -1) & ((ids0 & bitmask) != 0) + target = ids0 ^ bitmask + mm1 = cond & (ids1 == target) + mm2 = cond & (ids2 == target) + mm3 = cond & (ids3 == target) + mm4 = cond & (ids4 == target) + mm5 = cond & (ids5 == target) + mm6 = cond & (ids6 == target) + mm7 = cond & (ids7 == target) + hit = mm1 | mm2 | mm3 | mm4 | mm5 | mm6 | mm7 + + src = v0.to(tl.float32) + v1 = tl.where(mm1[:, None], (v1.to(tl.float32) + src).to(x_dtype), v1) + v2 = tl.where(mm2[:, None], (v2.to(tl.float32) + src).to(x_dtype), v2) + v3 = tl.where(mm3[:, None], (v3.to(tl.float32) + src).to(x_dtype), v3) + v4 = tl.where(mm4[:, None], (v4.to(tl.float32) + src).to(x_dtype), v4) + v5 = tl.where(mm5[:, None], (v5.to(tl.float32) + src).to(x_dtype), v5) + v6 = tl.where(mm6[:, None], (v6.to(tl.float32) + src).to(x_dtype), v6) + v7 = tl.where(mm7[:, None], (v7.to(tl.float32) + src).to(x_dtype), v7) + + v0 = tl.where(hit[:, None], 0.0, v0) + ids0 = tl.where(hit, -1, ids0) + ids0 = tl.where(cond & (~hit), target, ids0) + + # ========== lane1 as source ========== + cond = (ids1 != -1) & ((ids1 & bitmask) != 0) + target = ids1 ^ bitmask + mm0 = cond & (ids0 == target) + mm2 = cond & (ids2 == target) + mm3 = cond & (ids3 == target) + mm4 = cond & (ids4 == target) + mm5 = cond & (ids5 == target) + mm6 = cond & (ids6 == target) + mm7 = cond & (ids7 == target) + hit = mm0 | mm2 | mm3 | mm4 | mm5 | mm6 | mm7 + + src = v1.to(tl.float32) + v0 = tl.where(mm0[:, None], (v0.to(tl.float32) + src).to(x_dtype), v0) + v2 = tl.where(mm2[:, None], (v2.to(tl.float32) + src).to(x_dtype), v2) + v3 = tl.where(mm3[:, None], (v3.to(tl.float32) + src).to(x_dtype), v3) + v4 = tl.where(mm4[:, None], (v4.to(tl.float32) + src).to(x_dtype), v4) + v5 = tl.where(mm5[:, None], (v5.to(tl.float32) + src).to(x_dtype), v5) + v6 = tl.where(mm6[:, None], (v6.to(tl.float32) + src).to(x_dtype), v6) + v7 = tl.where(mm7[:, None], (v7.to(tl.float32) + src).to(x_dtype), v7) + + v1 = tl.where(hit[:, None], 0.0, v1) + ids1 = tl.where(hit, -1, ids1) + ids1 = tl.where(cond & (~hit), target, ids1) + + # ========== lane2 as source ========== + cond = (ids2 != -1) & ((ids2 & bitmask) != 0) + target = ids2 ^ bitmask + mm0 = cond & (ids0 == target) + mm1 = cond & (ids1 == target) + mm3 = cond & (ids3 == target) + mm4 = cond & (ids4 == target) + mm5 = cond & (ids5 == target) + mm6 = cond & (ids6 == target) + mm7 = cond & (ids7 == target) + hit = mm0 | mm1 | mm3 | mm4 | mm5 | mm6 | mm7 + + src = v2.to(tl.float32) + v0 = tl.where(mm0[:, None], (v0.to(tl.float32) + src).to(x_dtype), v0) + v1 = tl.where(mm1[:, None], (v1.to(tl.float32) + src).to(x_dtype), v1) + v3 = tl.where(mm3[:, None], (v3.to(tl.float32) + src).to(x_dtype), v3) + v4 = tl.where(mm4[:, None], (v4.to(tl.float32) + src).to(x_dtype), v4) + v5 = tl.where(mm5[:, None], (v5.to(tl.float32) + src).to(x_dtype), v5) + v6 = tl.where(mm6[:, None], (v6.to(tl.float32) + src).to(x_dtype), v6) + v7 = tl.where(mm7[:, None], (v7.to(tl.float32) + src).to(x_dtype), v7) + + v2 = tl.where(hit[:, None], 0.0, v2) + ids2 = tl.where(hit, -1, ids2) + ids2 = tl.where(cond & (~hit), target, ids2) + + # ========== lane3 as source ========== + cond = (ids3 != -1) & ((ids3 & bitmask) != 0) + target = ids3 ^ bitmask + mm0 = cond & (ids0 == target) + mm1 = cond & (ids1 == target) + mm2 = cond & (ids2 == target) + mm4 = cond & (ids4 == target) + mm5 = cond & (ids5 == target) + mm6 = cond & (ids6 == target) + mm7 = cond & (ids7 == target) + hit = mm0 | mm1 | mm2 | mm4 | mm5 | mm6 | mm7 + + src = v3.to(tl.float32) + v0 = tl.where(mm0[:, None], (v0.to(tl.float32) + src).to(x_dtype), v0) + v1 = tl.where(mm1[:, None], (v1.to(tl.float32) + src).to(x_dtype), v1) + v2 = tl.where(mm2[:, None], (v2.to(tl.float32) + src).to(x_dtype), v2) + v4 = tl.where(mm4[:, None], (v4.to(tl.float32) + src).to(x_dtype), v4) + v5 = tl.where(mm5[:, None], (v5.to(tl.float32) + src).to(x_dtype), v5) + v6 = tl.where(mm6[:, None], (v6.to(tl.float32) + src).to(x_dtype), v6) + v7 = tl.where(mm7[:, None], (v7.to(tl.float32) + src).to(x_dtype), v7) + + v3 = tl.where(hit[:, None], 0.0, v3) + ids3 = tl.where(hit, -1, ids3) + ids3 = tl.where(cond & (~hit), target, ids3) + + # ========== lane4 as source ========== + cond = (ids4 != -1) & ((ids4 & bitmask) != 0) + target = ids4 ^ bitmask + mm0 = cond & (ids0 == target) + mm1 = cond & (ids1 == target) + mm2 = cond & (ids2 == target) + mm3 = cond & (ids3 == target) + mm5 = cond & (ids5 == target) + mm6 = cond & (ids6 == target) + mm7 = cond & (ids7 == target) + hit = mm0 | mm1 | mm2 | mm3 | mm5 | mm6 | mm7 + + src = v4.to(tl.float32) + v0 = tl.where(mm0[:, None], (v0.to(tl.float32) + src).to(x_dtype), v0) + v1 = tl.where(mm1[:, None], (v1.to(tl.float32) + src).to(x_dtype), v1) + v2 = tl.where(mm2[:, None], (v2.to(tl.float32) + src).to(x_dtype), v2) + v3 = tl.where(mm3[:, None], (v3.to(tl.float32) + src).to(x_dtype), v3) + v5 = tl.where(mm5[:, None], (v5.to(tl.float32) + src).to(x_dtype), v5) + v6 = tl.where(mm6[:, None], (v6.to(tl.float32) + src).to(x_dtype), v6) + v7 = tl.where(mm7[:, None], (v7.to(tl.float32) + src).to(x_dtype), v7) + + v4 = tl.where(hit[:, None], 0.0, v4) + ids4 = tl.where(hit, -1, ids4) + ids4 = tl.where(cond & (~hit), target, ids4) + + # ========== lane5 as source ========== + cond = (ids5 != -1) & ((ids5 & bitmask) != 0) + target = ids5 ^ bitmask + mm0 = cond & (ids0 == target) + mm1 = cond & (ids1 == target) + mm2 = cond & (ids2 == target) + mm3 = cond & (ids3 == target) + mm4 = cond & (ids4 == target) + mm6 = cond & (ids6 == target) + mm7 = cond & (ids7 == target) + hit = mm0 | mm1 | mm2 | mm3 | mm4 | mm6 | mm7 + + src = v5.to(tl.float32) + v0 = tl.where(mm0[:, None], (v0.to(tl.float32) + src).to(x_dtype), v0) + v1 = tl.where(mm1[:, None], (v1.to(tl.float32) + src).to(x_dtype), v1) + v2 = tl.where(mm2[:, None], (v2.to(tl.float32) + src).to(x_dtype), v2) + v3 = tl.where(mm3[:, None], (v3.to(tl.float32) + src).to(x_dtype), v3) + v4 = tl.where(mm4[:, None], (v4.to(tl.float32) + src).to(x_dtype), v4) + v6 = tl.where(mm6[:, None], (v6.to(tl.float32) + src).to(x_dtype), v6) + v7 = tl.where(mm7[:, None], (v7.to(tl.float32) + src).to(x_dtype), v7) + + v5 = tl.where(hit[:, None], 0.0, v5) + ids5 = tl.where(hit, -1, ids5) + ids5 = tl.where(cond & (~hit), target, ids5) + + # ========== lane6 as source ========== + cond = (ids6 != -1) & ((ids6 & bitmask) != 0) + target = ids6 ^ bitmask + mm0 = cond & (ids0 == target) + mm1 = cond & (ids1 == target) + mm2 = cond & (ids2 == target) + mm3 = cond & (ids3 == target) + mm4 = cond & (ids4 == target) + mm5 = cond & (ids5 == target) + mm7 = cond & (ids7 == target) + hit = mm0 | mm1 | mm2 | mm3 | mm4 | mm5 | mm7 + + src = v6.to(tl.float32) + v0 = tl.where(mm0[:, None], (v0.to(tl.float32) + src).to(x_dtype), v0) + v1 = tl.where(mm1[:, None], (v1.to(tl.float32) + src).to(x_dtype), v1) + v2 = tl.where(mm2[:, None], (v2.to(tl.float32) + src).to(x_dtype), v2) + v3 = tl.where(mm3[:, None], (v3.to(tl.float32) + src).to(x_dtype), v3) + v4 = tl.where(mm4[:, None], (v4.to(tl.float32) + src).to(x_dtype), v4) + v5 = tl.where(mm5[:, None], (v5.to(tl.float32) + src).to(x_dtype), v5) + v7 = tl.where(mm7[:, None], (v7.to(tl.float32) + src).to(x_dtype), v7) + + v6 = tl.where(hit[:, None], 0.0, v6) + ids6 = tl.where(hit, -1, ids6) + ids6 = tl.where(cond & (~hit), target, ids6) + + # ========== lane7 as source ========== + cond = (ids7 != -1) & ((ids7 & bitmask) != 0) + target = ids7 ^ bitmask + mm0 = cond & (ids0 == target) + mm1 = cond & (ids1 == target) + mm2 = cond & (ids2 == target) + mm3 = cond & (ids3 == target) + mm4 = cond & (ids4 == target) + mm5 = cond & (ids5 == target) + mm6 = cond & (ids6 == target) + hit = mm0 | mm1 | mm2 | mm3 | mm4 | mm5 | mm6 + + src = v7.to(tl.float32) + v0 = tl.where(mm0[:, None], (v0.to(tl.float32) + src).to(x_dtype), v0) + v1 = tl.where(mm1[:, None], (v1.to(tl.float32) + src).to(x_dtype), v1) + v2 = tl.where(mm2[:, None], (v2.to(tl.float32) + src).to(x_dtype), v2) + v3 = tl.where(mm3[:, None], (v3.to(tl.float32) + src).to(x_dtype), v3) + v4 = tl.where(mm4[:, None], (v4.to(tl.float32) + src).to(x_dtype), v4) + v5 = tl.where(mm5[:, None], (v5.to(tl.float32) + src).to(x_dtype), v5) + v6 = tl.where(mm6[:, None], (v6.to(tl.float32) + src).to(x_dtype), v6) + + v7 = tl.where(hit[:, None], 0.0, v7) + ids7 = tl.where(hit, -1, ids7) + ids7 = tl.where(cond & (~hit), target, ids7) + + # ---- final: since ids unique per token, bucket0 is unique; compute directly in fp32 (no out_tile chain) ---- + acc = tl.zeros((BLOCK_M, BLOCK_H), dtype=tl.float32) + acc += tl.where((ids0 == 0)[:, None], v0.to(tl.float32), 0.0) + acc += tl.where((ids1 == 0)[:, None], v1.to(tl.float32), 0.0) + acc += tl.where((ids2 == 0)[:, None], v2.to(tl.float32), 0.0) + acc += tl.where((ids3 == 0)[:, None], v3.to(tl.float32), 0.0) + acc += tl.where((ids4 == 0)[:, None], v4.to(tl.float32), 0.0) + acc += tl.where((ids5 == 0)[:, None], v5.to(tl.float32), 0.0) + acc += tl.where((ids6 == 0)[:, None], v6.to(tl.float32), 0.0) + acc += tl.where((ids7 == 0)[:, None], v7.to(tl.float32), 0.0) + + acc *= routed_scaling_factor + + out_ptrs = out_ptr + m[:, None] * so_m + h[None, :] * so_h + tl.store(out_ptrs, acc.to(out_ptr.dtype.element_ty), mask=mask_mh) + + +def _moe_sum_tree_reduce_k8_fast_path( + input: torch.Tensor, # [M, 8, H] + output: torch.Tensor, # [M, H] + curr_topk_ids: torch.Tensor, # [M, 8], -1 means remote + routed_scaling_factor: float, + E: int, +): + assert input.is_contiguous() + assert output.is_contiguous() + assert curr_topk_ids.is_contiguous() + M, K, H = input.shape + assert K == 8 + assert output.shape == (M, H) + assert (E & (E - 1)) == 0 + + E_LEVEL = int(math.log2(E)) + + # K=8 specialization: use wider H tiles for better memory throughput. + if H >= 4096: + BLOCK_M = 8 + BLOCK_H = 128 + num_warps = 4 + elif H >= 2048: + BLOCK_M = 8 + BLOCK_H = 256 + num_warps = 8 + else: + BLOCK_M = 8 + BLOCK_H = 128 + num_warps = 4 + + grid = (triton.cdiv(M, BLOCK_M), triton.cdiv(H, BLOCK_H)) + _moe_sum_tree_reduce_k8_fused_kernel_opt2d[grid]( + input, + curr_topk_ids, + output, + input.stride(0), + input.stride(1), + input.stride(2), + curr_topk_ids.stride(0), + curr_topk_ids.stride(1), + output.stride(0), + output.stride(1), + M, + H, + E_LEVEL, + routed_scaling_factor, + BLOCK_M=BLOCK_M, + BLOCK_H=BLOCK_H, + num_warps=num_warps, + ) + return output + + +@triton.jit +def _moe_sum_tree_reduce_topk16_sparse_kernel( + x_ptr, + ids_ptr, + out_ptr, + sx_m, + sx_k, + sx_h, # x: [M, K, H] + sid_m, + sid_k, # ids: [M, K] + so_m, + so_h, # out: [M, H] + M, + K, + H, # runtime + E_LEVEL, # log2(E) + routed_scaling_factor, + BLOCK_H: tl.constexpr, + MAX_TOPK: tl.constexpr, # fixed compile-time capacity, e.g. 16 +): + # One program handles one token + one hidden tile. + pid_m = tl.program_id(0) + pid_h = tl.program_id(1) + + m = pid_m + h = pid_h * BLOCK_H + tl.arange(0, BLOCK_H) + + mask_m = m < M + mask_h = h < H + mask_mh = mask_m & mask_h + + slot = tl.arange(0, MAX_TOPK) + in_k = slot < K + + ids = tl.load( + ids_ptr + m * sid_m + slot * sid_k, mask=(mask_m & in_k), other=-1 + ).to(tl.int32) + valid = (ids != -1) & in_k & mask_m + + vals = tl.zeros((MAX_TOPK, BLOCK_H), dtype=tl.float32) + for s in range(MAX_TOPK): + vmask = mask_h & valid[s] + v = tl.load( + x_ptr + m * sx_m + s * sx_k + h * sx_h, + mask=vmask, + other=0.0, + ).to(tl.float32) + vals = tl.where((slot == s)[:, None], v[None, :], vals) + + idx = tl.arange(0, MAX_TOPK) + for bit in tl.range(0, E_LEVEL): + bitmask = 1 << bit + + for s in range(MAX_TOPK): + id_s = ids[s] + cond = (id_s != -1) & ((id_s & bitmask) != 0) + target = id_s ^ bitmask + + match = (ids == target) & (idx != s) + # 1-based index; 0 means no match. + match_pos1 = tl.max(tl.where(match, idx + 1, 0)) + has = match_pos1 > 0 + j = match_pos1 - 1 + + src = vals[s, :] + hit_mask = idx == j + vals = vals + tl.where(hit_mask[:, None] & has, src[None, :], 0.0) + vals = tl.where((idx == s)[:, None] & has, 0.0, vals) + ids = tl.where((idx == s) & has, -1, ids) + + # No partner found at this level: update bucket id only. + ids = tl.where((idx == s) & cond & (~has), target, ids) + + acc = tl.zeros((BLOCK_H,), dtype=tl.float32) + for s in range(MAX_TOPK): + acc += tl.where(ids[s] == 0, vals[s, :], 0.0) + + acc *= routed_scaling_factor + tl.store( + out_ptr + m * so_m + h * so_h, acc.to(out_ptr.dtype.element_ty), mask=mask_mh + ) + + +def moe_sum_tree_reduce_v1_topk_sparse16( + input: torch.Tensor, # [M, K, H] + output: torch.Tensor, # [M, H] + curr_topk_ids: torch.Tensor, # [M, K], -1 means remote + routed_scaling_factor: float, + E: int, + max_topk: int = 16, +): + """ + Supplemental kernel path for topk != 8: + - optimized for small topk (<=16) + - keeps deterministic tree-equivalent reduction semantics + - does not replace existing v1 entrypoint automatically + """ + assert input.is_contiguous() + assert output.is_contiguous() + assert curr_topk_ids.is_contiguous() + M, K, H = input.shape + assert output.shape == (M, H) + assert K <= max_topk, f"K={K} > max_topk={max_topk}" + assert (E & (E - 1)) == 0, "E must be power of 2" + + E_LEVEL = int(math.log2(E)) + BLOCK_H = 128 if H >= 4096 else 256 + num_warps = 4 if BLOCK_H == 128 else 8 + + grid = (M, triton.cdiv(H, BLOCK_H)) + _moe_sum_tree_reduce_topk16_sparse_kernel[grid]( + input, + curr_topk_ids, + output, + input.stride(0), + input.stride(1), + input.stride(2), + curr_topk_ids.stride(0), + curr_topk_ids.stride(1), + output.stride(0), + output.stride(1), + M, + K, + H, + E_LEVEL, + routed_scaling_factor, + BLOCK_H=BLOCK_H, + MAX_TOPK=max_topk, + num_warps=num_warps, + ) + return output + + +def moe_sum_tree_reduce_v0( + input: torch.Tensor, # [M, 8, H] + output: torch.Tensor, # [M, H] + curr_topk_ids: torch.Tensor, # [M, 8], -1 means remote + routed_scaling_factor: float, + E: int, +): + # Backward-compatible alias. + return _moe_sum_tree_reduce_k8_fast_path( + input=input, + output=output, + curr_topk_ids=curr_topk_ids, + routed_scaling_factor=routed_scaling_factor, + E=E, + ) + + +@triton.jit +def _moe_tree_reduce_sparse_k8_kernel( + x_ptr, + ids_ptr, + out_ptr, + sx_m: tl.constexpr, + sx_k: tl.constexpr, + sx_h: tl.constexpr, + sid_m: tl.constexpr, + sid_k: tl.constexpr, + so_m: tl.constexpr, + so_h: tl.constexpr, + M, + H, + routed_scaling_factor, + LOGE, # runtime: tl.range does not unroll + CAST_MODE: tl.constexpr, # 0: per-level bf16 round, 1: no intermediate cast + BLOCK_M: tl.constexpr, + BLOCK_H: tl.constexpr, +): + pid_m = tl.program_id(0) + pid_h = tl.program_id(1) + + m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) + h = pid_h * BLOCK_H + tl.arange(0, BLOCK_H) + + tl.multiple_of(m, BLOCK_M) + tl.multiple_of(h, 8) + tl.max_contiguous(h, BLOCK_H) + + mask_m = m < M + mask_h = h < H + mask_mh = mask_m[:, None] & mask_h[None, :] + + x_ty = x_ptr.dtype.element_ty + NEG2 = -2 # sentinel: topk_ids cannot be -2 + + # ---- ids ---- + ids0 = tl.load(ids_ptr + m * sid_m + 0 * sid_k, mask=mask_m, other=-1).to(tl.int32) + ids1 = tl.load(ids_ptr + m * sid_m + 1 * sid_k, mask=mask_m, other=-1).to(tl.int32) + ids2 = tl.load(ids_ptr + m * sid_m + 2 * sid_k, mask=mask_m, other=-1).to(tl.int32) + ids3 = tl.load(ids_ptr + m * sid_m + 3 * sid_k, mask=mask_m, other=-1).to(tl.int32) + ids4 = tl.load(ids_ptr + m * sid_m + 4 * sid_k, mask=mask_m, other=-1).to(tl.int32) + ids5 = tl.load(ids_ptr + m * sid_m + 5 * sid_k, mask=mask_m, other=-1).to(tl.int32) + ids6 = tl.load(ids_ptr + m * sid_m + 6 * sid_k, mask=mask_m, other=-1).to(tl.int32) + ids7 = tl.load(ids_ptr + m * sid_m + 7 * sid_k, mask=mask_m, other=-1).to(tl.int32) + + # ---- vals: load -> fp32 ---- + f0 = tl.load( + x_ptr + m[:, None] * sx_m + 0 * sx_k + h[None, :] * sx_h, + mask=mask_mh & (ids0 != -1)[:, None], + other=0.0, + ).to(tl.float32) + f1 = tl.load( + x_ptr + m[:, None] * sx_m + 1 * sx_k + h[None, :] * sx_h, + mask=mask_mh & (ids1 != -1)[:, None], + other=0.0, + ).to(tl.float32) + f2 = tl.load( + x_ptr + m[:, None] * sx_m + 2 * sx_k + h[None, :] * sx_h, + mask=mask_mh & (ids2 != -1)[:, None], + other=0.0, + ).to(tl.float32) + f3 = tl.load( + x_ptr + m[:, None] * sx_m + 3 * sx_k + h[None, :] * sx_h, + mask=mask_mh & (ids3 != -1)[:, None], + other=0.0, + ).to(tl.float32) + f4 = tl.load( + x_ptr + m[:, None] * sx_m + 4 * sx_k + h[None, :] * sx_h, + mask=mask_mh & (ids4 != -1)[:, None], + other=0.0, + ).to(tl.float32) + f5 = tl.load( + x_ptr + m[:, None] * sx_m + 5 * sx_k + h[None, :] * sx_h, + mask=mask_mh & (ids5 != -1)[:, None], + other=0.0, + ).to(tl.float32) + f6 = tl.load( + x_ptr + m[:, None] * sx_m + 6 * sx_k + h[None, :] * sx_h, + mask=mask_mh & (ids6 != -1)[:, None], + other=0.0, + ).to(tl.float32) + f7 = tl.load( + x_ptr + m[:, None] * sx_m + 7 * sx_k + h[None, :] * sx_h, + mask=mask_mh & (ids7 != -1)[:, None], + other=0.0, + ).to(tl.float32) + + # ---- main: runtime loop ---- + for bit in tl.range(0, LOGE): + bitmask = 1 << bit + + # snapshot ids for matching (critical!) + a0, a1, a2, a3, a4, a5, a6, a7 = ids0, ids1, ids2, ids3, ids4, ids5, ids6, ids7 + + # ---------------- s0 ---------------- + c0 = (a0 != -1) & ((a0 & bitmask) != 0) + t0 = tl.where(c0, a0 ^ bitmask, NEG2) + m01 = a1 == t0 + m02 = a2 == t0 + m03 = a3 == t0 + m04 = a4 == t0 + m05 = a5 == t0 + m06 = a6 == t0 + m07 = a7 == t0 + hit0 = m01 | m02 | m03 | m04 | m05 | m06 | m07 + src = f0 + f1 = tl.where(m01[:, None], f1 + src, f1) + f2 = tl.where(m02[:, None], f2 + src, f2) + f3 = tl.where(m03[:, None], f3 + src, f3) + f4 = tl.where(m04[:, None], f4 + src, f4) + f5 = tl.where(m05[:, None], f5 + src, f5) + f6 = tl.where(m06[:, None], f6 + src, f6) + f7 = tl.where(m07[:, None], f7 + src, f7) + f0 = tl.where(hit0[:, None], 0.0, f0) + ids0 = tl.where(c0, tl.where(hit0, -1, t0), ids0) + + # ---------------- s1 ---------------- + c1 = (a1 != -1) & ((a1 & bitmask) != 0) + t1 = tl.where(c1, a1 ^ bitmask, NEG2) + m10 = a0 == t1 + m12 = a2 == t1 + m13 = a3 == t1 + m14 = a4 == t1 + m15 = a5 == t1 + m16 = a6 == t1 + m17 = a7 == t1 + hit1 = m10 | m12 | m13 | m14 | m15 | m16 | m17 + src = f1 + f0 = tl.where(m10[:, None], f0 + src, f0) + f2 = tl.where(m12[:, None], f2 + src, f2) + f3 = tl.where(m13[:, None], f3 + src, f3) + f4 = tl.where(m14[:, None], f4 + src, f4) + f5 = tl.where(m15[:, None], f5 + src, f5) + f6 = tl.where(m16[:, None], f6 + src, f6) + f7 = tl.where(m17[:, None], f7 + src, f7) + f1 = tl.where(hit1[:, None], 0.0, f1) + ids1 = tl.where(c1, tl.where(hit1, -1, t1), ids1) + + # ---------------- s2 ---------------- + c2 = (a2 != -1) & ((a2 & bitmask) != 0) + t2 = tl.where(c2, a2 ^ bitmask, NEG2) + m20 = a0 == t2 + m21 = a1 == t2 + m23 = a3 == t2 + m24 = a4 == t2 + m25 = a5 == t2 + m26 = a6 == t2 + m27 = a7 == t2 + hit2 = m20 | m21 | m23 | m24 | m25 | m26 | m27 + src = f2 + f0 = tl.where(m20[:, None], f0 + src, f0) + f1 = tl.where(m21[:, None], f1 + src, f1) + f3 = tl.where(m23[:, None], f3 + src, f3) + f4 = tl.where(m24[:, None], f4 + src, f4) + f5 = tl.where(m25[:, None], f5 + src, f5) + f6 = tl.where(m26[:, None], f6 + src, f6) + f7 = tl.where(m27[:, None], f7 + src, f7) + f2 = tl.where(hit2[:, None], 0.0, f2) + ids2 = tl.where(c2, tl.where(hit2, -1, t2), ids2) + + # ---------------- s3 ---------------- + c3 = (a3 != -1) & ((a3 & bitmask) != 0) + t3 = tl.where(c3, a3 ^ bitmask, NEG2) + m30 = a0 == t3 + m31 = a1 == t3 + m32 = a2 == t3 + m34 = a4 == t3 + m35 = a5 == t3 + m36 = a6 == t3 + m37 = a7 == t3 + hit3 = m30 | m31 | m32 | m34 | m35 | m36 | m37 + src = f3 + f0 = tl.where(m30[:, None], f0 + src, f0) + f1 = tl.where(m31[:, None], f1 + src, f1) + f2 = tl.where(m32[:, None], f2 + src, f2) + f4 = tl.where(m34[:, None], f4 + src, f4) + f5 = tl.where(m35[:, None], f5 + src, f5) + f6 = tl.where(m36[:, None], f6 + src, f6) + f7 = tl.where(m37[:, None], f7 + src, f7) + f3 = tl.where(hit3[:, None], 0.0, f3) + ids3 = tl.where(c3, tl.where(hit3, -1, t3), ids3) + + # ---------------- s4 ---------------- + c4 = (a4 != -1) & ((a4 & bitmask) != 0) + t4 = tl.where(c4, a4 ^ bitmask, NEG2) + m40 = a0 == t4 + m41 = a1 == t4 + m42 = a2 == t4 + m43 = a3 == t4 + m45 = a5 == t4 + m46 = a6 == t4 + m47 = a7 == t4 + hit4 = m40 | m41 | m42 | m43 | m45 | m46 | m47 + src = f4 + f0 = tl.where(m40[:, None], f0 + src, f0) + f1 = tl.where(m41[:, None], f1 + src, f1) + f2 = tl.where(m42[:, None], f2 + src, f2) + f3 = tl.where(m43[:, None], f3 + src, f3) + f5 = tl.where(m45[:, None], f5 + src, f5) + f6 = tl.where(m46[:, None], f6 + src, f6) + f7 = tl.where(m47[:, None], f7 + src, f7) + f4 = tl.where(hit4[:, None], 0.0, f4) + ids4 = tl.where(c4, tl.where(hit4, -1, t4), ids4) + + # ---------------- s5 ---------------- + c5 = (a5 != -1) & ((a5 & bitmask) != 0) + t5 = tl.where(c5, a5 ^ bitmask, NEG2) + m50 = a0 == t5 + m51 = a1 == t5 + m52 = a2 == t5 + m53 = a3 == t5 + m54 = a4 == t5 + m56 = a6 == t5 + m57 = a7 == t5 + hit5 = m50 | m51 | m52 | m53 | m54 | m56 | m57 + src = f5 + f0 = tl.where(m50[:, None], f0 + src, f0) + f1 = tl.where(m51[:, None], f1 + src, f1) + f2 = tl.where(m52[:, None], f2 + src, f2) + f3 = tl.where(m53[:, None], f3 + src, f3) + f4 = tl.where(m54[:, None], f4 + src, f4) + f6 = tl.where(m56[:, None], f6 + src, f6) + f7 = tl.where(m57[:, None], f7 + src, f7) + f5 = tl.where(hit5[:, None], 0.0, f5) + ids5 = tl.where(c5, tl.where(hit5, -1, t5), ids5) + + # ---------------- s6 ---------------- + c6 = (a6 != -1) & ((a6 & bitmask) != 0) + t6 = tl.where(c6, a6 ^ bitmask, NEG2) + m60 = a0 == t6 + m61 = a1 == t6 + m62 = a2 == t6 + m63 = a3 == t6 + m64 = a4 == t6 + m65 = a5 == t6 + m67 = a7 == t6 + hit6 = m60 | m61 | m62 | m63 | m64 | m65 | m67 + src = f6 + f0 = tl.where(m60[:, None], f0 + src, f0) + f1 = tl.where(m61[:, None], f1 + src, f1) + f2 = tl.where(m62[:, None], f2 + src, f2) + f3 = tl.where(m63[:, None], f3 + src, f3) + f4 = tl.where(m64[:, None], f4 + src, f4) + f5 = tl.where(m65[:, None], f5 + src, f5) + f7 = tl.where(m67[:, None], f7 + src, f7) + f6 = tl.where(hit6[:, None], 0.0, f6) + ids6 = tl.where(c6, tl.where(hit6, -1, t6), ids6) + + # ---------------- s7 ---------------- + c7 = (a7 != -1) & ((a7 & bitmask) != 0) + t7 = tl.where(c7, a7 ^ bitmask, NEG2) + m70 = a0 == t7 + m71 = a1 == t7 + m72 = a2 == t7 + m73 = a3 == t7 + m74 = a4 == t7 + m75 = a5 == t7 + m76 = a6 == t7 + hit7 = m70 | m71 | m72 | m73 | m74 | m75 | m76 + src = f7 + f0 = tl.where(m70[:, None], f0 + src, f0) + f1 = tl.where(m71[:, None], f1 + src, f1) + f2 = tl.where(m72[:, None], f2 + src, f2) + f3 = tl.where(m73[:, None], f3 + src, f3) + f4 = tl.where(m74[:, None], f4 + src, f4) + f5 = tl.where(m75[:, None], f5 + src, f5) + f6 = tl.where(m76[:, None], f6 + src, f6) + f7 = tl.where(hit7[:, None], 0.0, f7) + ids7 = tl.where(c7, tl.where(hit7, -1, t7), ids7) + + # ---- per-level bf16 rounding (optional) ---- + if CAST_MODE == 0: + # emulate "store bf16 per level then reload": round once per level + f0 = f0.to(x_ty).to(tl.float32) + f1 = f1.to(x_ty).to(tl.float32) + f2 = f2.to(x_ty).to(tl.float32) + f3 = f3.to(x_ty).to(tl.float32) + f4 = f4.to(x_ty).to(tl.float32) + f5 = f5.to(x_ty).to(tl.float32) + f6 = f6.to(x_ty).to(tl.float32) + f7 = f7.to(x_ty).to(tl.float32) + + # ---- gather root (id==0) ---- + acc = tl.zeros((BLOCK_M, BLOCK_H), dtype=tl.float32) + acc += tl.where((ids0 == 0)[:, None], f0, 0.0) + acc += tl.where((ids1 == 0)[:, None], f1, 0.0) + acc += tl.where((ids2 == 0)[:, None], f2, 0.0) + acc += tl.where((ids3 == 0)[:, None], f3, 0.0) + acc += tl.where((ids4 == 0)[:, None], f4, 0.0) + acc += tl.where((ids5 == 0)[:, None], f5, 0.0) + acc += tl.where((ids6 == 0)[:, None], f6, 0.0) + acc += tl.where((ids7 == 0)[:, None], f7, 0.0) + + acc *= routed_scaling_factor + + out_ptrs = out_ptr + m[:, None] * so_m + h[None, :] * so_h + tl.store(out_ptrs, acc.to(out_ptr.dtype.element_ty), mask=mask_mh) + + +def _launch_sparse_tree_k8_k10( + input: torch.Tensor, # [M, K, H], bf16 contiguous + output: torch.Tensor, # [M, H], bf16 contiguous + curr_topk_ids: torch.Tensor, # [M, K], int32/int64 contiguous, -1 remote + routed_scaling_factor: float, + E: int, + cast_mode: int = 1, # <-- you asked for "last cast" version, so default = 1 +): + assert ( + input.is_contiguous() + and output.is_contiguous() + and curr_topk_ids.is_contiguous() + ) + M, K, H = input.shape + assert output.shape == (M, H) + assert K in (8, 10) + assert (E & (E - 1)) == 0 + LOGE = int(math.log2(E)) + + # stable params for CUDA Graph + BLOCK_M = 8 + BLOCK_H = 256 + num_warps = 8 + grid = (triton.cdiv(M, BLOCK_M), triton.cdiv(H, BLOCK_H)) + + # NOTE: for strict CUDA Graph safety, avoid dtype conversion here. + # Please ensure curr_topk_ids is int32 upstream. + assert ( + curr_topk_ids.dtype == torch.int32 + ), "make ids int32 upstream for CUDA Graph stability" + + if K == 8: + _moe_tree_reduce_sparse_k8_kernel[grid]( + input, + curr_topk_ids, + output, + input.stride(0), + input.stride(1), + input.stride(2), + curr_topk_ids.stride(0), + curr_topk_ids.stride(1), + output.stride(0), + output.stride(1), + M, + H, + routed_scaling_factor, + LOGE=LOGE, + CAST_MODE=cast_mode, + BLOCK_M=BLOCK_M, + BLOCK_H=BLOCK_H, + num_warps=num_warps, + ) + + +def moe_sum_tree_reduce_v2( + input: torch.Tensor, + output: torch.Tensor, + curr_topk_ids: torch.Tensor, + routed_scaling_factor: float, + E: int, + *, + cast_mode: int = 0, # default to "last cast" as requested +): + assert ( + input.is_contiguous() + and output.is_contiguous() + and curr_topk_ids.is_contiguous() + ) + M, K, H = input.shape + if M == 0: + return output + if K in (8, 10): + _launch_sparse_tree_k8_k10( + input=input, + output=output, + curr_topk_ids=curr_topk_ids, + routed_scaling_factor=routed_scaling_factor, + E=E, + cast_mode=cast_mode, + ) + return output + + # fallback: keep your existing generic path here + return output + + +def moe_sum_tree_reduce( + input: torch.Tensor, + output: torch.Tensor, + curr_topk_ids: torch.Tensor, + routed_scaling_factor: float, + E: int, +): + curr_topk_ids = curr_topk_ids.to(torch.int32) + if os.environ.get("SGLANG_MOE_TREE_REDUCE_USE_V2", "0") == "1": + return moe_sum_tree_reduce_v2( + input=input, + output=output, + curr_topk_ids=curr_topk_ids, + routed_scaling_factor=routed_scaling_factor, + E=E, + ) + + if input.shape[1] == 8: + return moe_sum_tree_reduce_v0( + input=input, + output=output, + curr_topk_ids=curr_topk_ids, + routed_scaling_factor=routed_scaling_factor, + E=E, + ) + return moe_sum_tree_reduce_v1( + input=input, + output=output, + curr_topk_ids=curr_topk_ids, + routed_scaling_factor=routed_scaling_factor, + E=E, + ) + + +# if os.getenv("MOE_SUM_OPTIM", "0") == "1": +# moe_sum_tree_reduce = moe_sum_tree_reduce_optim +# print("using optimized moe_sum_tree_reduce_optim") +# else: +# print("using original moe_sum_tree_reduce_original") + +if os.getenv("TREE_ALL_REDUCE_OPTIM", "0") == "1": + tree_all_reduce_sum = tree_all_reduce_sum_optim diff --git a/python/sglang/srt/true_on_policy/__init__.py b/python/sglang/srt/true_on_policy/__init__.py new file mode 100644 index 000000000000..9fbaaeee48f3 --- /dev/null +++ b/python/sglang/srt/true_on_policy/__init__.py @@ -0,0 +1,49 @@ +"""True-on-policy runtime contract helpers.""" + +from .config import ( + ROW_LINEAR_INV_BLOCK_K, + get_on_policy_rms_norm_kwargs, + get_rl_on_policy_target, + is_tp_invariant_target, + is_true_on_policy_enabled, + patch_prefill_only_deterministic_inference_for_cuda_graph, + should_disable_flashinfer_allreduce_fusion, + should_disable_fused_qk_norm_mrope, + should_disable_mlp_allreduce_fusion_for_on_policy, + should_disable_reduce_scatter_for_on_policy, + should_force_bfloat16_dense_tensor_math, + should_force_bfloat16_lm_head, + should_use_tp_invariant_row_linear, + should_use_tp_invariant_tree_all_reduce, +) +from .contracts import ( + QWEN3_DENSE_TRUE_ON_POLICY_V1, + SGLangTrueOnPolicyContract, + SGLangTrueOnPolicyRuntimePolicy, + get_true_on_policy_contract, + resolve_true_on_policy_runtime_policy, + validate_true_on_policy_contract, +) + +__all__ = [ + "QWEN3_DENSE_TRUE_ON_POLICY_V1", + "ROW_LINEAR_INV_BLOCK_K", + "SGLangTrueOnPolicyContract", + "SGLangTrueOnPolicyRuntimePolicy", + "get_true_on_policy_contract", + "get_on_policy_rms_norm_kwargs", + "get_rl_on_policy_target", + "is_tp_invariant_target", + "is_true_on_policy_enabled", + "patch_prefill_only_deterministic_inference_for_cuda_graph", + "resolve_true_on_policy_runtime_policy", + "validate_true_on_policy_contract", + "should_disable_flashinfer_allreduce_fusion", + "should_disable_fused_qk_norm_mrope", + "should_disable_mlp_allreduce_fusion_for_on_policy", + "should_disable_reduce_scatter_for_on_policy", + "should_force_bfloat16_dense_tensor_math", + "should_force_bfloat16_lm_head", + "should_use_tp_invariant_row_linear", + "should_use_tp_invariant_tree_all_reduce", +] diff --git a/python/sglang/srt/true_on_policy/config.py b/python/sglang/srt/true_on_policy/config.py new file mode 100644 index 000000000000..67c054955bc6 --- /dev/null +++ b/python/sglang/srt/true_on_policy/config.py @@ -0,0 +1,150 @@ +from __future__ import annotations + +import contextlib +from typing import Any, Iterator, Optional + +import torch + +from sglang.srt.true_on_policy.contracts import resolve_true_on_policy_runtime_policy + +ROW_LINEAR_INV_BLOCK_K = 128 + + +def _get_global_server_args() -> Any: + from sglang.srt.server_args import get_global_server_args + + return get_global_server_args() + + +def get_rl_on_policy_target() -> Optional[str]: + return getattr(_get_global_server_args(), "rl_on_policy_target", None) + + +def is_true_on_policy_enabled() -> bool: + return resolve_true_on_policy_runtime_policy(_get_global_server_args()).enabled + + +def is_tp_invariant_target() -> bool: + return resolve_true_on_policy_runtime_policy( + _get_global_server_args() + ).tp_invariant_row_linear + + +def should_disable_reduce_scatter_for_on_policy() -> bool: + return resolve_true_on_policy_runtime_policy( + _get_global_server_args() + ).disable_reduce_scatter + + +def should_disable_mlp_allreduce_fusion_for_on_policy() -> bool: + return resolve_true_on_policy_runtime_policy( + _get_global_server_args() + ).disable_mlp_allreduce_fusion + + +def should_disable_flashinfer_allreduce_fusion() -> bool: + return resolve_true_on_policy_runtime_policy( + _get_global_server_args() + ).disable_flashinfer_allreduce_fusion + + +def should_force_bfloat16_dense_tensor_math() -> bool: + return resolve_true_on_policy_runtime_policy( + _get_global_server_args() + ).force_bfloat16_dense_tensor_math + + +def should_force_bfloat16_lm_head( + *, + use_fp32_lm_head: bool = False, +) -> bool: + return ( + resolve_true_on_policy_runtime_policy( + _get_global_server_args() + ).force_bfloat16_lm_head + and not use_fp32_lm_head + ) + + +def should_disable_fused_qk_norm_mrope() -> bool: + return resolve_true_on_policy_runtime_policy( + _get_global_server_args() + ).disable_fused_qk_norm_mrope + + +def get_on_policy_rms_norm_kwargs( + *, + weight_dtype: Optional[torch.dtype] = None, + override_orig_dtype: Optional[torch.dtype] = None, + fp32_residual: bool = False, +) -> dict[str, Any]: + if not is_true_on_policy_enabled(): + return {} + + kwargs: dict[str, Any] = { + "cast_x_before_out_mul": True, + "fp32_residual": fp32_residual, + } + if weight_dtype is not None: + kwargs["weight_dtype"] = weight_dtype + if override_orig_dtype is not None: + kwargs["override_orig_dtype"] = override_orig_dtype + return kwargs + + +def should_use_tp_invariant_row_linear( + k_size: int, + row_linear_enable_inv: Optional[bool] = None, +) -> bool: + policy_enabled = resolve_true_on_policy_runtime_policy( + _get_global_server_args() + ).tp_invariant_row_linear + if row_linear_enable_inv is not None: + policy_enabled = policy_enabled and row_linear_enable_inv + + return ( + policy_enabled + and k_size >= ROW_LINEAR_INV_BLOCK_K + and k_size % ROW_LINEAR_INV_BLOCK_K == 0 + ) + + +def should_use_tp_invariant_tree_all_reduce( + accl_binary_tree_enabled: Optional[bool] = None, +) -> bool: + policy_enabled = resolve_true_on_policy_runtime_policy( + _get_global_server_args() + ).deterministic_tree_all_reduce + if accl_binary_tree_enabled is not None: + policy_enabled = policy_enabled and not accl_binary_tree_enabled + + return policy_enabled + + +@contextlib.contextmanager +def patch_prefill_only_deterministic_inference_for_cuda_graph( + server_args: Any, + *, + attn_backend: Optional[Any] = None, + dvr_target_verify_cuda_graph: bool = False, +) -> Iterator[bool]: + enabled = ( + getattr(server_args, "enable_prefill_only_deterministic_inference", False) + and not dvr_target_verify_cuda_graph + ) + if not enabled: + yield False + return + + saved_num_splits = None + if attn_backend is not None and hasattr(attn_backend, "num_splits"): + saved_num_splits = attn_backend.num_splits + + try: + if attn_backend is not None and hasattr(attn_backend, "num_splits"): + attn_backend.num_splits = 0 + + yield True + finally: + if attn_backend is not None and hasattr(attn_backend, "num_splits"): + attn_backend.num_splits = saved_num_splits diff --git a/python/sglang/srt/true_on_policy/contracts.py b/python/sglang/srt/true_on_policy/contracts.py new file mode 100644 index 000000000000..310778c842ce --- /dev/null +++ b/python/sglang/srt/true_on_policy/contracts.py @@ -0,0 +1,111 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Optional + +from sglang.srt.true_on_policy.schema import ( + QWEN3_DENSE_TRUE_ON_POLICY_V1_SCHEMA, + TrueOnPolicyContractName, + TrueOnPolicyContractSchema, +) + +QWEN3_DENSE_TRUE_ON_POLICY_V1 = QWEN3_DENSE_TRUE_ON_POLICY_V1_SCHEMA.name + + +@dataclass(frozen=True) +class SGLangTrueOnPolicyRuntimePolicy: + """SGLang-local behavior implied by a true-on-policy parity contract.""" + + contract_name: Optional[str] + enabled: bool + force_bfloat16_dense_tensor_math: bool + force_bfloat16_lm_head: bool + disable_reduce_scatter: bool + disable_mlp_allreduce_fusion: bool + disable_flashinfer_allreduce_fusion: bool + tp_invariant_row_linear: bool + deterministic_tree_all_reduce: bool + disable_fused_qk_norm_mrope: bool + + +DEFAULT_RUNTIME_POLICY = SGLangTrueOnPolicyRuntimePolicy( + contract_name=None, + enabled=False, + force_bfloat16_dense_tensor_math=False, + force_bfloat16_lm_head=False, + disable_reduce_scatter=False, + disable_mlp_allreduce_fusion=False, + disable_flashinfer_allreduce_fusion=False, + tp_invariant_row_linear=False, + deterministic_tree_all_reduce=False, + disable_fused_qk_norm_mrope=False, +) + + +@dataclass(frozen=True) +class SGLangTrueOnPolicyContract: + """SGLang-local adapter from a shared contract schema to runtime policy.""" + + schema: TrueOnPolicyContractSchema + + @property + def name(self) -> TrueOnPolicyContractName: + return self.schema.name + + def policy_for(self, server_args: Any) -> SGLangTrueOnPolicyRuntimePolicy: + uses_tp_invariant_rollout = getattr(server_args, "tp_size", 1) > 1 + return SGLangTrueOnPolicyRuntimePolicy( + contract_name=self.name, + enabled=True, + force_bfloat16_dense_tensor_math=True, + force_bfloat16_lm_head=True, + disable_reduce_scatter=True, + disable_mlp_allreduce_fusion=True, + disable_flashinfer_allreduce_fusion=uses_tp_invariant_rollout, + tp_invariant_row_linear=uses_tp_invariant_rollout, + deterministic_tree_all_reduce=uses_tp_invariant_rollout, + disable_fused_qk_norm_mrope=True, + ) + + +QWEN3_DENSE_TRUE_ON_POLICY_CONTRACT = SGLangTrueOnPolicyContract( + schema=QWEN3_DENSE_TRUE_ON_POLICY_V1_SCHEMA, +) + + +_CONTRACT_BY_NAME = { + QWEN3_DENSE_TRUE_ON_POLICY_CONTRACT.name: QWEN3_DENSE_TRUE_ON_POLICY_CONTRACT, +} + + +def get_true_on_policy_contract(contract_name: str) -> SGLangTrueOnPolicyContract: + try: + return _CONTRACT_BY_NAME[contract_name] + except KeyError as exc: + supported = ", ".join(sorted(_CONTRACT_BY_NAME)) + raise ValueError( + f"Unsupported SGLang true-on-policy contract {contract_name!r}. " + f"Supported contracts: {supported}" + ) from exc + + +def _contract_name_for(server_args: Any) -> Optional[str]: + return getattr(server_args, "true_on_policy_contract", None) + + +def validate_true_on_policy_contract(server_args: Any) -> None: + contract_name = getattr(server_args, "true_on_policy_contract", None) + if contract_name is None: + return + get_true_on_policy_contract(contract_name) + + +def resolve_true_on_policy_runtime_policy( + server_args: Any, +) -> SGLangTrueOnPolicyRuntimePolicy: + contract_name = _contract_name_for(server_args) + if contract_name is None: + return DEFAULT_RUNTIME_POLICY + + validate_true_on_policy_contract(server_args) + return get_true_on_policy_contract(contract_name).policy_for(server_args) diff --git a/python/sglang/srt/true_on_policy/schema.py b/python/sglang/srt/true_on_policy/schema.py new file mode 100644 index 000000000000..0628573c6880 --- /dev/null +++ b/python/sglang/srt/true_on_policy/schema.py @@ -0,0 +1,33 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Literal + +TrueOnPolicyContractName = Literal["qwen3_dense_true_on_policy_v1"] +ModelFamily = Literal["qwen3_dense", "qwen3_moe", "qwen3_next"] +KernelContract = Literal["qwen3_dense_sglang_math"] +LogprobContract = Literal["sglang_prefill"] + + +@dataclass(frozen=True) +class TrueOnPolicyContractSchema: + """Declarative cross-repo identity for a true-on-policy parity contract.""" + + name: TrueOnPolicyContractName + model_family: ModelFamily + required_kernel_contracts: tuple[KernelContract, ...] + logprob_contract: LogprobContract + sglang_attention_backend: str + fsdp_attention_implementation: str + disable_megatron_sequence_parallel: bool + + +QWEN3_DENSE_TRUE_ON_POLICY_V1_SCHEMA = TrueOnPolicyContractSchema( + name="qwen3_dense_true_on_policy_v1", + model_family="qwen3_dense", + required_kernel_contracts=("qwen3_dense_sglang_math",), + logprob_contract="sglang_prefill", + sglang_attention_backend="fa3", + fsdp_attention_implementation="flash_attention_3", + disable_megatron_sequence_parallel=True, +) diff --git a/test/manual/layers/test_layernorm.py b/test/manual/layers/test_layernorm.py index 299e5dcffaf4..8a0e68601ba6 100644 --- a/test/manual/layers/test_layernorm.py +++ b/test/manual/layers/test_layernorm.py @@ -56,6 +56,26 @@ def test_rms_norm(self): ): self._run_rms_norm_test(*params) + def test_rms_norm_cuda_uses_native_for_fp32_weight(self): + torch.manual_seed(0) + + hidden_size = 256 + layer = RMSNorm( + hidden_size, + cast_x_before_out_mul=True, + weight_dtype=torch.float32, + override_orig_dtype=torch.float32, + ) + layer.weight.data.normal_(mean=1.0, std=0.1) + x = torch.randn(17, hidden_size, dtype=torch.bfloat16) / hidden_size + + with torch.inference_mode(): + ref_out = layer.forward_native(x) + out = layer(x) + + self.assertEqual(out.dtype, torch.float32) + self.assertTrue(torch.equal(out, ref_out)) + class TestGemmaRMSNorm(CustomTestCase): DTYPES = [torch.half, torch.bfloat16] diff --git a/test/registered/core/test_dense_deterministic_math.py b/test/registered/core/test_dense_deterministic_math.py new file mode 100644 index 000000000000..7c5c79ee52bd --- /dev/null +++ b/test/registered/core/test_dense_deterministic_math.py @@ -0,0 +1,338 @@ +import json +import os +import subprocess +import textwrap +import unittest +from types import SimpleNamespace + +import torch + +from sglang.srt.true_on_policy import ( + QWEN3_DENSE_TRUE_ON_POLICY_V1, + get_on_policy_rms_norm_kwargs, + should_force_bfloat16_dense_tensor_math, + should_force_bfloat16_lm_head, +) +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=12, suite="stage-a-test-cpu") + + +def _run_dense_math_script(script_body: str) -> dict[str, object]: + stubbed_imports = textwrap.dedent(""" + import importlib.machinery + import json + import sys + import types + from pydantic import BaseModel + + def install_openai_stubs(): + openai_mod = types.ModuleType("openai") + openai_types_mod = types.ModuleType("openai.types") + openai_responses_mod = types.ModuleType("openai.types.responses") + openai_response_mod = types.ModuleType("openai.types.responses.response") + openai_tool_mod = types.ModuleType("openai.types.responses.tool") + + openai_mod.__spec__ = importlib.machinery.ModuleSpec("openai", loader=None) + openai_types_mod.__spec__ = importlib.machinery.ModuleSpec("openai.types", loader=None) + openai_responses_mod.__spec__ = importlib.machinery.ModuleSpec( + "openai.types.responses", loader=None + ) + openai_response_mod.__spec__ = importlib.machinery.ModuleSpec( + "openai.types.responses.response", loader=None + ) + openai_tool_mod.__spec__ = importlib.machinery.ModuleSpec( + "openai.types.responses.tool", loader=None + ) + + for name in [ + "ResponseFunctionToolCall", + "ResponseInputItemParam", + "ResponseOutputItem", + "ResponseOutputMessage", + "ResponseOutputText", + "ResponseReasoningItem", + ]: + setattr(openai_responses_mod, name, type(name, (BaseModel,), {})) + + openai_response_mod.ToolChoice = type("ToolChoice", (BaseModel,), {}) + openai_tool_mod.Tool = type("Tool", (BaseModel,), {}) + + sys.modules.setdefault("openai", openai_mod) + sys.modules.setdefault("openai.types", openai_types_mod) + sys.modules.setdefault("openai.types.responses", openai_responses_mod) + sys.modules.setdefault("openai.types.responses.response", openai_response_mod) + sys.modules.setdefault("openai.types.responses.tool", openai_tool_mod) + + install_openai_stubs() + + hf_utils_mod = types.ModuleType("sglang.srt.utils.hf_transformers_utils") + hf_utils_mod.__spec__ = importlib.machinery.ModuleSpec( + "sglang.srt.utils.hf_transformers_utils", loader=None + ) + hf_utils_mod.check_gguf_file = lambda *args, **kwargs: False + hf_utils_mod.get_rope_config = lambda config: ( + getattr(config, "rope_theta", 1000000), + getattr(config, "rope_scaling", None), + ) + sys.modules.setdefault("sglang.srt.utils.hf_transformers_utils", hf_utils_mod) + + gguf_mod = types.ModuleType("gguf") + gguf_mod.__spec__ = importlib.machinery.ModuleSpec("gguf", loader=None) + gguf_mod.GGMLQuantizationType = type( + "GGMLQuantizationType", + (), + { + "F32": 0, + "F16": 1, + "BF16": 2, + "Q4_0": 3, + "Q4_1": 4, + "Q5_0": 5, + "Q5_1": 6, + "Q8_0": 7, + "Q8_1": 8, + "Q2_K": 9, + "Q3_K": 10, + "Q4_K": 11, + "Q5_K": 12, + "Q6_K": 13, + "IQ1_S": 14, + "IQ1_M": 15, + "IQ2_XXS": 16, + "IQ2_XS": 17, + "IQ2_S": 18, + "IQ3_XXS": 19, + "IQ3_S": 20, + "IQ4_NL": 21, + "IQ4_XS": 22, + }, + ) + sys.modules.setdefault("gguf", gguf_mod) + """) + + env = dict(os.environ) + pythonpath = env.get("PYTHONPATH") + repo_python = "python" + env["PYTHONPATH"] = ( + f"{repo_python}{os.pathsep}{pythonpath}" if pythonpath else repo_python + ) + script = f"{stubbed_imports}\n{script_body}" + completed = subprocess.run( + ["python", "-c", script], + check=True, + capture_output=True, + text=True, + env=env, + ) + return json.loads(completed.stdout) + + +class TestDenseOnPolicyHelpers(unittest.TestCase): + def test_default_dense_math_helpers_are_inactive(self): + server_args = SimpleNamespace( + true_on_policy_contract=None, + tp_size=1, + ) + + self.assertFalse(should_force_bfloat16_dense_tensor_math(server_args)) + self.assertFalse( + should_force_bfloat16_lm_head( + server_args=server_args, + use_fp32_lm_head=False, + ) + ) + self.assertEqual(get_on_policy_rms_norm_kwargs(server_args), {}) + + def test_on_policy_dense_math_helpers_enable_bfloat16_and_rms_norm_kwargs(self): + server_args = SimpleNamespace( + true_on_policy_contract=QWEN3_DENSE_TRUE_ON_POLICY_V1, + tp_size=1, + ) + + kwargs = get_on_policy_rms_norm_kwargs( + server_args, + weight_dtype=torch.float32, + override_orig_dtype=torch.float32, + fp32_residual=True, + ) + + self.assertTrue(should_force_bfloat16_dense_tensor_math(server_args)) + self.assertTrue( + should_force_bfloat16_lm_head( + server_args=server_args, + use_fp32_lm_head=False, + ) + ) + self.assertFalse( + should_force_bfloat16_lm_head( + server_args=server_args, + use_fp32_lm_head=True, + ) + ) + self.assertEqual(kwargs["weight_dtype"], torch.float32) + self.assertEqual(kwargs["override_orig_dtype"], torch.float32) + self.assertTrue(kwargs["cast_x_before_out_mul"]) + self.assertTrue(kwargs["fp32_residual"]) + + +class TestDenseOnPolicyContracts(unittest.TestCase): + def test_qwen3_style_rms_norm_keeps_fp32_weight_output_and_residual(self): + result = _run_dense_math_script(textwrap.dedent(""" + import json + from types import SimpleNamespace + + import torch + + from sglang.srt.layers.layernorm import RMSNorm + from sglang.srt.true_on_policy import get_on_policy_rms_norm_kwargs + + from sglang.srt.true_on_policy import QWEN3_DENSE_TRUE_ON_POLICY_V1 + + server_args = SimpleNamespace( + true_on_policy_contract=QWEN3_DENSE_TRUE_ON_POLICY_V1, + tp_size=1, + ) + norm = RMSNorm( + 4, + eps=1e-6, + **get_on_policy_rms_norm_kwargs( + server_args, + weight_dtype=torch.float32, + override_orig_dtype=torch.float32, + fp32_residual=True, + ), + ) + x = torch.randn(2, 4, dtype=torch.bfloat16) + residual = torch.randn(2, 4, dtype=torch.bfloat16) + out, residual_out = norm.forward_native(x, residual) + print( + json.dumps( + { + "weight_dtype": str(norm.weight.dtype), + "out_dtype": str(out.dtype), + "residual_dtype": str(residual_out.dtype), + } + ) + ) + """)) + + self.assertEqual(result["weight_dtype"], "torch.float32") + self.assertEqual(result["out_dtype"], "torch.float32") + self.assertEqual(result["residual_dtype"], "torch.float32") + + def test_rms_norm_can_self_configure_from_true_on_policy_role_hints(self): + result = _run_dense_math_script(textwrap.dedent(""" + import json + + import torch + + from sglang.srt.layers.layernorm import RMSNorm + from sglang.srt.server_args import ( + ServerArgs, + get_global_server_args, + set_global_server_args_for_scheduler, + ) + from sglang.srt.true_on_policy import QWEN3_DENSE_TRUE_ON_POLICY_V1 + + set_global_server_args_for_scheduler(ServerArgs(model_path="dummy")) + server_args = get_global_server_args() + server_args.true_on_policy_contract = QWEN3_DENSE_TRUE_ON_POLICY_V1 + server_args.tp_size = 1 + norm = RMSNorm( + 4, + eps=1e-6, + true_on_policy_weight_dtype=torch.float32, + true_on_policy_override_orig_dtype=torch.float32, + true_on_policy_fp32_residual=True, + ) + print( + json.dumps( + { + "weight_dtype": str(norm.weight.dtype), + "cast_x_before_out_mul": norm.cast_x_before_out_mul, + "fp32_residual": norm.fp32_residual, + "override_orig_dtype": str(norm.override_orig_dtype), + } + ) + ) + """)) + + self.assertEqual(result["weight_dtype"], "torch.float32") + self.assertTrue(result["cast_x_before_out_mul"]) + self.assertTrue(result["fp32_residual"]) + self.assertEqual(result["override_orig_dtype"], "torch.float32") + + def test_on_policy_lm_head_forces_bfloat16_matmul_inputs(self): + result = _run_dense_math_script(textwrap.dedent(""" + import json + from types import SimpleNamespace + from unittest.mock import patch + + import torch + import torch.nn as nn + + from sglang.srt.layers.logits_processor import LogitsProcessor + from sglang.srt.server_args import ( + ServerArgs, + get_global_server_args, + set_global_server_args_for_scheduler, + ) + from sglang.srt.true_on_policy import QWEN3_DENSE_TRUE_ON_POLICY_V1 + + class DummyMeta: + gathered_buffer = None + next_token_logits_buffer = None + + def compute_dp_attention_metadata(self): + return None + + class LMHeadStub(nn.Module): + def __init__(self): + super().__init__() + self.weight = nn.Parameter(torch.randn(8, 4, dtype=torch.float32)) + + set_global_server_args_for_scheduler(ServerArgs(model_path="dummy")) + get_global_server_args().enable_dp_lm_head = False + get_global_server_args().enable_fp32_lm_head = False + get_global_server_args().true_on_policy_contract = QWEN3_DENSE_TRUE_ON_POLICY_V1 + get_global_server_args().tp_size = 1 + + processor = LogitsProcessor( + SimpleNamespace(vocab_size=8, final_logit_softcapping=None), + skip_all_gather=True, + logit_scale=None, + ) + hidden_states = torch.randn(2, 4, dtype=torch.float32) + head = LMHeadStub() + captured = {} + + original_matmul = torch.matmul + + def probe_matmul(a, b, *args, **kwargs): + if not captured: + captured["a_dtype"] = str(a.dtype) + captured["b_dtype"] = str(b.dtype) + return original_matmul(a, b, *args, **kwargs) + + with patch("torch.matmul", new=probe_matmul): + logits = processor._get_logits(hidden_states, head, DummyMeta()) + + print( + json.dumps( + { + "a_dtype": captured["a_dtype"], + "b_dtype": captured["b_dtype"], + "logits_dtype": str(logits.dtype), + } + ) + ) + """)) + + self.assertEqual(result["a_dtype"], "torch.bfloat16") + self.assertEqual(result["b_dtype"], "torch.bfloat16") + self.assertEqual(result["logits_dtype"], "torch.bfloat16") + + +if __name__ == "__main__": + unittest.main(verbosity=2) diff --git a/test/registered/core/test_on_policy_wiring.py b/test/registered/core/test_on_policy_wiring.py new file mode 100644 index 000000000000..946f1922a9e0 --- /dev/null +++ b/test/registered/core/test_on_policy_wiring.py @@ -0,0 +1,606 @@ +import json +import os +import subprocess +import textwrap +import unittest +from types import SimpleNamespace +from unittest.mock import patch + +import torch + +from sglang.srt.true_on_policy import ( + QWEN3_DENSE_TRUE_ON_POLICY_V1, + get_rl_on_policy_target, + get_true_on_policy_contract, + is_tp_invariant_target, + is_true_on_policy_enabled, + patch_prefill_only_deterministic_inference_for_cuda_graph, + resolve_true_on_policy_runtime_policy, + should_disable_flashinfer_allreduce_fusion, + should_disable_fused_qk_norm_mrope, + should_disable_mlp_allreduce_fusion_for_on_policy, + should_disable_reduce_scatter_for_on_policy, + should_use_tp_invariant_row_linear, + should_use_tp_invariant_tree_all_reduce, +) +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=12, suite="stage-a-test-cpu") + +_PATCH_TARGET = "sglang.srt.server_args.get_global_server_args" + + +def _run_server_args_script(argv: list[str]) -> dict[str, object]: + stubbed_imports = textwrap.dedent(""" + import argparse + import importlib.machinery + import json + import sys + import types + from types import SimpleNamespace + from unittest.mock import patch + + from pydantic import BaseModel + + def install_openai_stubs(): + openai_mod = types.ModuleType("openai") + openai_types_mod = types.ModuleType("openai.types") + openai_responses_mod = types.ModuleType("openai.types.responses") + openai_response_mod = types.ModuleType("openai.types.responses.response") + openai_tool_mod = types.ModuleType("openai.types.responses.tool") + + openai_mod.__spec__ = importlib.machinery.ModuleSpec("openai", loader=None) + openai_types_mod.__spec__ = importlib.machinery.ModuleSpec("openai.types", loader=None) + openai_responses_mod.__spec__ = importlib.machinery.ModuleSpec( + "openai.types.responses", loader=None + ) + openai_response_mod.__spec__ = importlib.machinery.ModuleSpec( + "openai.types.responses.response", loader=None + ) + openai_tool_mod.__spec__ = importlib.machinery.ModuleSpec( + "openai.types.responses.tool", loader=None + ) + + for name in [ + "ResponseFunctionToolCall", + "ResponseInputItemParam", + "ResponseOutputItem", + "ResponseOutputMessage", + "ResponseOutputText", + "ResponseReasoningItem", + ]: + setattr(openai_responses_mod, name, type(name, (BaseModel,), {})) + + openai_response_mod.ToolChoice = type("ToolChoice", (BaseModel,), {}) + openai_tool_mod.Tool = type("Tool", (BaseModel,), {}) + + sys.modules.setdefault("openai", openai_mod) + sys.modules.setdefault("openai.types", openai_types_mod) + sys.modules.setdefault("openai.types.responses", openai_responses_mod) + sys.modules.setdefault("openai.types.responses.response", openai_response_mod) + sys.modules.setdefault("openai.types.responses.tool", openai_tool_mod) + + install_openai_stubs() + + hf_utils_mod = types.ModuleType("sglang.srt.utils.hf_transformers_utils") + hf_utils_mod.__spec__ = importlib.machinery.ModuleSpec( + "sglang.srt.utils.hf_transformers_utils", loader=None + ) + hf_utils_mod.check_gguf_file = lambda *args, **kwargs: False + sys.modules.setdefault("sglang.srt.utils.hf_transformers_utils", hf_utils_mod) + + from sglang.srt.server_args import ServerArgs + + def _mock_model_config(): + return SimpleNamespace( + hf_config=SimpleNamespace(architectures=["Qwen2ForCausalLM"]) + ) + + parser = argparse.ArgumentParser() + ServerArgs.add_cli_args(parser) + cli_args = parser.parse_args(ARGV) + + with patch("sglang.srt.server_args.get_device", return_value="cuda"), patch.object( + ServerArgs, "get_model_config", return_value=_mock_model_config() + ): + server_args = ServerArgs.from_cli_args(cli_args) + server_args._handle_deterministic_inference() + + print( + json.dumps( + { + "enable_deterministic_inference": server_args.enable_deterministic_inference, + "enable_prefill_only_deterministic_inference": server_args.enable_prefill_only_deterministic_inference, + "enable_flashinfer_allreduce_fusion": server_args.enable_flashinfer_allreduce_fusion, + "rl_on_policy_target": server_args.rl_on_policy_target, + "true_on_policy_contract": server_args.true_on_policy_contract, + "sampling_backend": server_args.sampling_backend, + } + ) + ) + """) + + env = dict(os.environ) + pythonpath = env.get("PYTHONPATH") + repo_python = "python" + env["PYTHONPATH"] = ( + f"{repo_python}{os.pathsep}{pythonpath}" if pythonpath else repo_python + ) + script = f"ARGV = {argv!r}\n{stubbed_imports}" + completed = subprocess.run( + ["python", "-c", script], + check=True, + capture_output=True, + text=True, + env=env, + ) + return json.loads(completed.stdout) + + +class TestOnPolicyServerArgs(unittest.TestCase): + def test_cli_parses_prefill_only_deterministic_flag(self): + result = _run_server_args_script( + [ + "--model-path", + "dummy", + "--attention-backend", + "triton", + "--enable-prefill-only-deterministic-inference", + ] + ) + + self.assertTrue(result["enable_prefill_only_deterministic_inference"]) + self.assertTrue(result["enable_deterministic_inference"]) + self.assertIsNone(result["rl_on_policy_target"]) + self.assertEqual(result["sampling_backend"], "pytorch") + + def test_cli_accepts_fsdp_and_fsdp_tp_targets(self): + fsdp_tp_result = _run_server_args_script( + [ + "--model-path", + "dummy", + "--attention-backend", + "triton", + "--rl-on-policy-target", + "fsdp_tp", + ] + ) + self.assertEqual(fsdp_tp_result["rl_on_policy_target"], "fsdp_tp") + self.assertIsNone(fsdp_tp_result["true_on_policy_contract"]) + self.assertTrue(fsdp_tp_result["enable_deterministic_inference"]) + + fsdp_result = _run_server_args_script( + [ + "--model-path", + "dummy", + "--attention-backend", + "triton", + "--rl-on-policy-target", + "fsdp", + ] + ) + self.assertEqual(fsdp_result["rl_on_policy_target"], "fsdp") + self.assertIsNone(fsdp_result["true_on_policy_contract"]) + self.assertTrue(fsdp_result["enable_deterministic_inference"]) + + def test_cli_accepts_explicit_true_on_policy_contract(self): + result = _run_server_args_script( + [ + "--model-path", + "dummy", + "--attention-backend", + "triton", + "--true-on-policy-contract", + QWEN3_DENSE_TRUE_ON_POLICY_V1, + ] + ) + + self.assertIsNone(result["rl_on_policy_target"]) + self.assertEqual( + result["true_on_policy_contract"], QWEN3_DENSE_TRUE_ON_POLICY_V1 + ) + self.assertTrue(result["enable_deterministic_inference"]) + + def test_contract_tp_rollout_disables_flashinfer_allreduce_fusion(self): + result = _run_server_args_script( + [ + "--model-path", + "dummy", + "--attention-backend", + "triton", + "--tensor-parallel-size", + "2", + "--true-on-policy-contract", + QWEN3_DENSE_TRUE_ON_POLICY_V1, + "--enable-flashinfer-allreduce-fusion", + ] + ) + self.assertFalse(result["enable_flashinfer_allreduce_fusion"]) + + def test_legacy_target_keeps_flashinfer_allreduce_fusion_available(self): + result = _run_server_args_script( + [ + "--model-path", + "dummy", + "--attention-backend", + "triton", + "--rl-on-policy-target", + "fsdp_tp", + "--enable-flashinfer-allreduce-fusion", + ] + ) + self.assertTrue(result["enable_flashinfer_allreduce_fusion"]) + + +def _mock_args(**kwargs): + defaults = dict( + rl_on_policy_target=None, + true_on_policy_contract=None, + tp_size=1, + ) + defaults.update(kwargs) + return SimpleNamespace(**defaults) + + +def _contract_args(*, tp_size: int = 1): + return _mock_args( + true_on_policy_contract=QWEN3_DENSE_TRUE_ON_POLICY_V1, + tp_size=tp_size, + ) + + +class TestDefaultPathUnchanged(unittest.TestCase): + """Default serving must not enter true-on-policy policy paths.""" + + def setUp(self): + self.default_args = _mock_args() + + def test_default_args_no_on_policy(self): + with patch(_PATCH_TARGET, return_value=self.default_args): + self.assertIsNone(get_rl_on_policy_target()) + self.assertFalse(is_true_on_policy_enabled()) + self.assertFalse(is_tp_invariant_target()) + + def test_default_args_row_linear_uses_quant_method(self): + with patch(_PATCH_TARGET, return_value=self.default_args): + self.assertFalse( + should_use_tp_invariant_row_linear( + 256, + row_linear_enable_inv=True, + ) + ) + + def test_default_args_tree_allreduce_not_selected(self): + with patch(_PATCH_TARGET, return_value=self.default_args): + self.assertFalse( + should_use_tp_invariant_tree_all_reduce( + accl_binary_tree_enabled=False, + ) + ) + + def test_default_args_reduce_scatter_available(self): + with patch(_PATCH_TARGET, return_value=self.default_args): + self.assertFalse(should_disable_reduce_scatter_for_on_policy()) + + def test_default_args_mlp_fusion_available(self): + with patch(_PATCH_TARGET, return_value=self.default_args): + self.assertFalse(should_disable_mlp_allreduce_fusion_for_on_policy()) + + def test_default_args_flashinfer_fusion_available(self): + with patch(_PATCH_TARGET, return_value=self.default_args): + self.assertFalse(should_disable_flashinfer_allreduce_fusion()) + + def test_default_server_args_cli_no_on_policy_flags(self): + result = _run_server_args_script( + ["--model-path", "dummy", "--attention-backend", "triton"] + ) + self.assertIsNone(result["rl_on_policy_target"]) + self.assertFalse(result["enable_deterministic_inference"]) + self.assertFalse(result["enable_prefill_only_deterministic_inference"]) + + +class TestOnPolicyHelpers(unittest.TestCase): + def test_tp_invariant_row_linear_selection(self): + with patch(_PATCH_TARGET, return_value=_contract_args(tp_size=2)): + self.assertTrue( + should_use_tp_invariant_row_linear( + 256, + row_linear_enable_inv=True, + ) + ) + + def test_tp_invariant_row_linear_selection_is_contract_owned(self): + with patch(_PATCH_TARGET, return_value=_contract_args(tp_size=2)): + with patch.dict(os.environ, {"ROW_LINEAR_ENABLE_INV": "0"}): + self.assertTrue(should_use_tp_invariant_row_linear(256)) + + def test_contract_resolver_ignores_legacy_target_without_contract(self): + policy = resolve_true_on_policy_runtime_policy( + _mock_args(rl_on_policy_target="fsdp_tp", tp_size=2) + ) + + self.assertIsNone(policy.contract_name) + self.assertFalse(policy.enabled) + self.assertFalse(policy.tp_invariant_row_linear) + self.assertFalse(policy.deterministic_tree_all_reduce) + + def test_contract_resolver_accepts_explicit_qwen3_dense_contract(self): + args_tp1 = _contract_args(tp_size=1) + policy = resolve_true_on_policy_runtime_policy(args_tp1) + + self.assertTrue(policy.enabled) + self.assertTrue(policy.force_bfloat16_dense_tensor_math) + self.assertFalse(policy.tp_invariant_row_linear) + with patch(_PATCH_TARGET, return_value=args_tp1): + self.assertFalse( + should_use_tp_invariant_row_linear( + 96, + row_linear_enable_inv=True, + ) + ) + self.assertFalse( + should_use_tp_invariant_row_linear( + 256, + row_linear_enable_inv=True, + ) + ) + + def test_contract_object_owns_sglang_runtime_policy_values(self): + contract = get_true_on_policy_contract(QWEN3_DENSE_TRUE_ON_POLICY_V1) + + policy = contract.policy_for(_contract_args(tp_size=2)) + + self.assertEqual(contract.schema.name, QWEN3_DENSE_TRUE_ON_POLICY_V1) + self.assertEqual(contract.schema.model_family, "qwen3_dense") + self.assertEqual(policy.contract_name, QWEN3_DENSE_TRUE_ON_POLICY_V1) + self.assertTrue(policy.enabled) + self.assertTrue(policy.force_bfloat16_dense_tensor_math) + self.assertTrue(policy.force_bfloat16_lm_head) + self.assertTrue(policy.disable_reduce_scatter) + self.assertTrue(policy.disable_mlp_allreduce_fusion) + self.assertTrue(policy.disable_flashinfer_allreduce_fusion) + self.assertTrue(policy.tp_invariant_row_linear) + self.assertTrue(policy.deterministic_tree_all_reduce) + self.assertTrue(policy.disable_fused_qk_norm_mrope) + + def test_reduce_scatter_and_fusion_are_disabled_for_contract(self): + with patch(_PATCH_TARGET, return_value=_contract_args(tp_size=1)): + self.assertTrue(should_disable_reduce_scatter_for_on_policy()) + with patch(_PATCH_TARGET, return_value=_contract_args(tp_size=2)): + self.assertTrue(should_disable_mlp_allreduce_fusion_for_on_policy()) + with patch(_PATCH_TARGET, return_value=_mock_args()): + self.assertFalse(should_disable_reduce_scatter_for_on_policy()) + + def test_tree_all_reduce_selection_requires_tp_rollout_and_no_accl(self): + with patch(_PATCH_TARGET, return_value=_contract_args(tp_size=2)): + self.assertTrue( + should_use_tp_invariant_tree_all_reduce( + accl_binary_tree_enabled=False, + ) + ) + self.assertFalse( + should_use_tp_invariant_tree_all_reduce( + accl_binary_tree_enabled=True, + ) + ) + with patch(_PATCH_TARGET, return_value=_contract_args(tp_size=1)): + self.assertFalse( + should_use_tp_invariant_tree_all_reduce( + accl_binary_tree_enabled=False, + ) + ) + + def test_tree_all_reduce_selection_is_contract_owned(self): + with patch(_PATCH_TARGET, return_value=_contract_args(tp_size=2)): + with patch.dict(os.environ, {"ACCL_BINARY_TREE_ENABLE": "1"}): + self.assertTrue(should_use_tp_invariant_tree_all_reduce()) + + def test_attention_handoff_tree_reduce_uses_attention_tp_group(self): + from sglang.srt.layers.communicator import ( + CommunicateWithAllReduceAndLayerNormFn, + ) + + hidden_states = torch.ones(2, 4) + residual = torch.full((2, 4), 3.0) + + class FakeNorm: + def __call__(self, x, residual): + return x + residual, residual + + with ( + patch( + "sglang.srt.layers.communicator.get_attn_tp_context", + return_value=SimpleNamespace(input_scattered=False), + ), + patch( + "sglang.srt.layers.communicator.apply_aiter_all_reduce_fusion", + return_value=False, + ), + patch( + "sglang.srt.layers.communicator.apply_flashinfer_allreduce_fusion", + return_value=False, + ), + patch( + "sglang.srt.layers.communicator.attention_tensor_model_parallel_all_reduce", + side_effect=lambda x: x + 10.0, + ) as attn_tree_reduce, + ): + output, output_residual = ( + CommunicateWithAllReduceAndLayerNormFn._gather_hidden_states_and_residual( + hidden_states, + residual, + forward_batch=None, + layernorm=FakeNorm(), + context=SimpleNamespace(attn_dp_size=1, cache=None), + residual_input_mode=None, + ) + ) + + attn_tree_reduce.assert_called_once() + torch.testing.assert_close(output, hidden_states + 10.0 + residual) + torch.testing.assert_close(output_residual, residual) + + def test_prefill_only_cuda_graph_patch_only_scopes_attention_splits(self): + server_args = SimpleNamespace( + enable_prefill_only_deterministic_inference=True, + enable_deterministic_inference=True, + enable_flashinfer_allreduce_fusion=False, + rl_on_policy_target="fsdp_tp", + true_on_policy_contract=QWEN3_DENSE_TRUE_ON_POLICY_V1, + disable_custom_all_reduce=True, + ) + attn_backend = SimpleNamespace(num_splits=1) + + with patch.dict( + os.environ, + { + "SGLANG_ENABLE_DETERMINISTIC_INFERENCE": "1", + "SGLANG_DISABLE_CUSTOM_ALL_REDUCE": "1", + "NCCL_ALGO": "allreduce:tree", + }, + clear=False, + ): + with patch_prefill_only_deterministic_inference_for_cuda_graph( + server_args, + attn_backend=attn_backend, + ) as patched: + self.assertTrue(patched) + self.assertTrue(server_args.enable_deterministic_inference) + self.assertFalse(server_args.enable_flashinfer_allreduce_fusion) + self.assertEqual(server_args.rl_on_policy_target, "fsdp_tp") + self.assertEqual( + server_args.true_on_policy_contract, + QWEN3_DENSE_TRUE_ON_POLICY_V1, + ) + self.assertEqual(attn_backend.num_splits, 0) + self.assertEqual( + os.environ["SGLANG_ENABLE_DETERMINISTIC_INFERENCE"], "1" + ) + self.assertEqual(os.environ["SGLANG_DISABLE_CUSTOM_ALL_REDUCE"], "1") + self.assertEqual(os.environ["NCCL_ALGO"], "allreduce:tree") + + self.assertTrue(server_args.enable_deterministic_inference) + self.assertFalse(server_args.enable_flashinfer_allreduce_fusion) + self.assertEqual(server_args.rl_on_policy_target, "fsdp_tp") + self.assertEqual( + server_args.true_on_policy_contract, + QWEN3_DENSE_TRUE_ON_POLICY_V1, + ) + self.assertTrue(server_args.disable_custom_all_reduce) + self.assertEqual(attn_backend.num_splits, 1) + self.assertEqual(os.environ["SGLANG_ENABLE_DETERMINISTIC_INFERENCE"], "1") + self.assertEqual(os.environ["SGLANG_DISABLE_CUSTOM_ALL_REDUCE"], "1") + self.assertEqual(os.environ["NCCL_ALGO"], "allreduce:tree") + + def test_row_linear_k_alignment_edge_cases(self): + with patch(_PATCH_TARGET, return_value=_contract_args(tp_size=2)): + self.assertFalse( + should_use_tp_invariant_row_linear(64, row_linear_enable_inv=True), + ) + self.assertTrue( + should_use_tp_invariant_row_linear(128, row_linear_enable_inv=True), + ) + self.assertFalse( + should_use_tp_invariant_row_linear(300, row_linear_enable_inv=True), + ) + self.assertTrue( + should_use_tp_invariant_row_linear(3584, row_linear_enable_inv=True), + ) + + def test_row_linear_explicit_override_can_disable(self): + with patch(_PATCH_TARGET, return_value=_contract_args(tp_size=2)): + self.assertFalse( + should_use_tp_invariant_row_linear(256, row_linear_enable_inv=False) + ) + + def test_flashinfer_allreduce_fusion_helpers(self): + with patch(_PATCH_TARGET, return_value=_contract_args(tp_size=2)): + self.assertTrue(should_disable_flashinfer_allreduce_fusion()) + with patch(_PATCH_TARGET, return_value=_contract_args(tp_size=1)): + self.assertFalse(should_disable_flashinfer_allreduce_fusion()) + with patch(_PATCH_TARGET, return_value=_mock_args()): + self.assertFalse(should_disable_flashinfer_allreduce_fusion()) + + def test_fused_qk_norm_mrope_helper_follows_true_on_policy_contract(self): + with patch(_PATCH_TARGET, return_value=_contract_args(tp_size=1)): + self.assertTrue(should_disable_fused_qk_norm_mrope()) + with patch(_PATCH_TARGET, return_value=_mock_args()): + self.assertFalse(should_disable_fused_qk_norm_mrope()) + + def test_get_rl_on_policy_target_returns_correct_value(self): + with patch( + _PATCH_TARGET, return_value=_mock_args(rl_on_policy_target="fsdp_tp") + ): + self.assertEqual(get_rl_on_policy_target(), "fsdp_tp") + with patch(_PATCH_TARGET, return_value=_mock_args(rl_on_policy_target="fsdp")): + self.assertEqual(get_rl_on_policy_target(), "fsdp") + with patch(_PATCH_TARGET, return_value=_mock_args()): + self.assertIsNone(get_rl_on_policy_target()) + + def test_is_true_on_policy_enabled_for_both_targets(self): + with patch(_PATCH_TARGET, return_value=_contract_args(tp_size=1)): + self.assertTrue(is_true_on_policy_enabled()) + with patch(_PATCH_TARGET, return_value=_contract_args(tp_size=2)): + self.assertTrue(is_true_on_policy_enabled()) + with patch(_PATCH_TARGET, return_value=_mock_args()): + self.assertFalse(is_true_on_policy_enabled()) + + def test_is_tp_invariant_target_only_fsdp_tp(self): + with patch(_PATCH_TARGET, return_value=_contract_args(tp_size=2)): + self.assertTrue(is_tp_invariant_target()) + with patch(_PATCH_TARGET, return_value=_contract_args(tp_size=1)): + self.assertFalse(is_tp_invariant_target()) + with patch(_PATCH_TARGET, return_value=_mock_args()): + self.assertFalse(is_tp_invariant_target()) + + def test_cuda_graph_patch_noop_when_disabled(self): + server_args = SimpleNamespace( + enable_prefill_only_deterministic_inference=False, + enable_deterministic_inference=True, + rl_on_policy_target="fsdp_tp", + ) + with patch_prefill_only_deterministic_inference_for_cuda_graph( + server_args, + ) as patched: + self.assertFalse(patched) + self.assertTrue(server_args.enable_deterministic_inference) + self.assertEqual(server_args.rl_on_policy_target, "fsdp_tp") + + def test_cuda_graph_patch_noop_when_dvr_verify(self): + server_args = SimpleNamespace( + enable_prefill_only_deterministic_inference=True, + enable_deterministic_inference=True, + enable_flashinfer_allreduce_fusion=False, + rl_on_policy_target="fsdp_tp", + disable_custom_all_reduce=True, + ) + with patch_prefill_only_deterministic_inference_for_cuda_graph( + server_args, + dvr_target_verify_cuda_graph=True, + ) as patched: + self.assertFalse(patched) + self.assertTrue(server_args.enable_deterministic_inference) + self.assertEqual(server_args.rl_on_policy_target, "fsdp_tp") + + def test_tp_invariant_ops_import_is_available(self): + import sglang.srt.tp_invariant_ops as tp_invariant_ops + + self.assertTrue(hasattr(tp_invariant_ops, "matmul_tp_inv")) + + def test_legacy_on_policy_utils_import_matches_true_on_policy_namespace(self): + from sglang.srt import true_on_policy + from sglang.srt.layers import on_policy_utils as legacy + + self.assertIs( + legacy.should_use_tp_invariant_row_linear, + true_on_policy.should_use_tp_invariant_row_linear, + ) + self.assertIs( + legacy.patch_prefill_only_deterministic_inference_for_cuda_graph, + true_on_policy.patch_prefill_only_deterministic_inference_for_cuda_graph, + ) + self.assertTrue(hasattr(torch.ops, "tp_inv_ops")) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/core/test_tp_invariant_ops.py b/test/registered/core/test_tp_invariant_ops.py new file mode 100644 index 000000000000..c20f853940d8 --- /dev/null +++ b/test/registered/core/test_tp_invariant_ops.py @@ -0,0 +1,866 @@ +"""Tests for TP-invariant kernels (PR1). + +TP-invariance property: + Given the same TP degree and the same input data, matmul_tp_persistent + plus tree_all_reduce_sum produces bitwise identical results across runs. + When K/BLOCK_K is divisible by tp_size AND each shard yields a power-of-two + block count, TP=1 and TP=N also agree (isomorphic tree structure). + + For production K values (e.g. 3584, 5120) where block counts per shard are + not power-of-two, the tree structure differs between TP degrees. The + invariance guarantee is *determinism for a fixed TP degree*. + +All bitwise assertions use torch.equal, never approximate tolerances. +""" + +import os +import random +import unittest + +import torch +import torch.distributed as dist + +from sglang.srt.tp_invariant_ops import ( + disable_tp_invariant_mode, + enable_tp_invariant_mode, + is_tp_invariant_mode_enabled, + matmul_tp_inv, + matmul_tp_persistent, + moe_sum_tree_reduce, + set_tp_invariant_mode, + tree_all_reduce_sum, +) +from sglang.srt.tp_invariant_ops.tp_invariant_ops import ( + _MATMUL_K_BLOCK, + _fixed_tree_sum_tensors, + _is_power_of_two, +) +from sglang.test.ci.ci_register import register_cpu_ci, register_cuda_ci + +register_cpu_ci(est_time=12, suite="stage-a-test-cpu") +register_cuda_ci(est_time=30, suite="stage-b-kernel-unit-8-gpu-h200") + +BLOCK_K = _MATMUL_K_BLOCK # 128 + + +def _simulate_tp_matmul(A, B, tp_size, **kwargs): + """Simulate what production TP does: each rank runs matmul_tp_persistent + on its K-shard, then tree_all_reduce_sum gathers and tree-sums.""" + K = A.shape[1] + shard = K // tp_size + partials = [] + for r in range(tp_size): + start = r * shard + end = start + shard + partials.append( + matmul_tp_persistent(A[:, start:end], B[start:end, :], **kwargs) + ) + return _fixed_tree_sum_tensors(partials) + + +# --------------------------------------------------------------------------- +# Mode flag +# --------------------------------------------------------------------------- +class TestTPInvariantMode(unittest.TestCase): + def tearDown(self): + disable_tp_invariant_mode() + + def test_mode_context_restores_previous_state(self): + disable_tp_invariant_mode() + self.assertFalse(is_tp_invariant_mode_enabled()) + + with set_tp_invariant_mode(True): + self.assertTrue(is_tp_invariant_mode_enabled()) + self.assertFalse(is_tp_invariant_mode_enabled()) + + enable_tp_invariant_mode() + with set_tp_invariant_mode(False): + self.assertFalse(is_tp_invariant_mode_enabled()) + self.assertTrue(is_tp_invariant_mode_enabled()) + + def test_enable_is_idempotent(self): + enable_tp_invariant_mode() + enable_tp_invariant_mode() + self.assertTrue(is_tp_invariant_mode_enabled()) + disable_tp_invariant_mode() + self.assertFalse(is_tp_invariant_mode_enabled()) + + def test_context_restores_after_exception(self): + disable_tp_invariant_mode() + try: + with set_tp_invariant_mode(True): + raise RuntimeError("deliberate") + except RuntimeError: + pass + self.assertFalse(is_tp_invariant_mode_enabled()) + + +# --------------------------------------------------------------------------- +# Reference ops: correctness +# --------------------------------------------------------------------------- +class TestTPInvariantReferenceOps(unittest.TestCase): + def test_fixed_tree_sum_order_is_stable(self): + values = [ + torch.tensor([1.0e20], dtype=torch.float32), + torch.tensor([1.0], dtype=torch.float32), + torch.tensor([-1.0e20], dtype=torch.float32), + torch.tensor([3.0], dtype=torch.float32), + ] + tree_result = _fixed_tree_sum_tensors(values) + sequential_result = values[0] + values[1] + values[2] + values[3] + + self.assertEqual(tree_result.item(), 0.0) + self.assertEqual(sequential_result.item(), 3.0) + + def test_matmul_tp_persistent_matches_torch_matmul_fp32(self): + torch.manual_seed(0) + A = torch.randn(5, 257, dtype=torch.float32) + B = torch.randn(257, 7, dtype=torch.float32) + bias = torch.randn(7, dtype=torch.float32) + + actual = matmul_tp_persistent(A, B, bias=bias) + expected = A @ B + bias + torch.testing.assert_close(actual, expected, rtol=1e-5, atol=1e-5) + + def test_matmul_tp_persistent_bf16_approximate_to_torch(self): + """BF16 block-tree matmul vs native torch matmul. Different accumulation + order means results differ by up to a few BF16 ULPs; this test only + checks that the magnitude is in the same ballpark, not bitwise equality.""" + torch.manual_seed(0) + A = torch.randn(6, 256, dtype=torch.bfloat16) + B = torch.randn(256, 10, dtype=torch.bfloat16) + bias = torch.randn(10, dtype=torch.bfloat16) + + actual = matmul_tp_persistent(A, B, bias=bias) + expected = A @ B + bias + torch.testing.assert_close(actual, expected, rtol=1e-1, atol=1e-1) + + def test_torch_custom_op_dispatches_to_matmul(self): + A = torch.arange(6, dtype=torch.float32).reshape(2, 3) + B = torch.arange(12, dtype=torch.float32).reshape(3, 4) + + actual = torch.ops.tp_inv_ops.matmul_tp_inv(A, B) + expected = A @ B + torch.testing.assert_close(actual, expected) + + def test_torch_custom_op_dispatches_with_bias(self): + torch.manual_seed(99) + A = torch.randn(4, 128, dtype=torch.float32) + B = torch.randn(128, 8, dtype=torch.float32) + bias = torch.randn(8, dtype=torch.float32) + + actual = torch.ops.tp_inv_ops.matmul_tp_inv(A, B, bias) + expected = matmul_tp_persistent(A, B, bias=bias) + self.assertTrue(torch.equal(actual, expected)) + + def test_matmul_tp_inv_matches_persistent(self): + torch.manual_seed(11) + A = torch.randn(4, 256, dtype=torch.bfloat16) + B = torch.randn(256, 16, dtype=torch.bfloat16) + self.assertTrue(torch.equal(matmul_tp_inv(A, B), matmul_tp_persistent(A, B))) + + def test_moe_sum_tree_reduce_matches_expert_order_reference(self): + input_tensor = torch.tensor( + [ + [ + [1.0e20, 1.0], + [1.0, 2.0], + [-1.0e20, 4.0], + [3.0, 8.0], + ], + [ + [5.0, 7.0], + [11.0, 13.0], + [17.0, 19.0], + [23.0, 29.0], + ], + ], + dtype=torch.float32, + ) + curr_topk_ids = torch.tensor( + [[0, 1, 2, 3], [3, -1, 1, 0]], + dtype=torch.int64, + ) + output = torch.empty(2, 2, dtype=torch.float32) + + moe_sum_tree_reduce( + input=input_tensor, + output=output, + curr_topk_ids=curr_topk_ids, + routed_scaling_factor=0.5, + E=4, + ) + + expected = torch.tensor([[0.0, 7.5], [22.5, 27.5]], dtype=torch.float32) + torch.testing.assert_close(output, expected) + + def test_moe_sum_tree_reduce_rejects_non_power_of_two_expert_count(self): + with self.assertRaisesRegex(ValueError, "power of two"): + moe_sum_tree_reduce( + input=torch.zeros(1, 1, 2), + output=torch.zeros(1, 2), + curr_topk_ids=torch.zeros(1, 1, dtype=torch.int64), + routed_scaling_factor=1.0, + E=3, + ) + + +# --------------------------------------------------------------------------- +# Matmul TP invariance: the core contract +# +# Two classes of tests: +# 1. Cross-TP: TP=1 == TP=N, requires isomorphic tree (K/BLOCK_K power-of-two +# multiple of tp_size). Uses torch.equal. +# 2. Determinism: same TP degree, two runs == bitwise identical. Works for +# any valid K alignment. +# --------------------------------------------------------------------------- +class TestTPInvarianceCrossTP(unittest.TestCase): + """TP=1 == TP=N when the binary tree structure is isomorphic.""" + + def test_fp32_cross_tp_k512(self): + torch.manual_seed(42) + A = torch.randn(8, 512, dtype=torch.float32) + B = torch.randn(512, 16, dtype=torch.float32) + + result_tp1 = matmul_tp_persistent(A, B) + result_tp2 = _simulate_tp_matmul(A, B, tp_size=2) + result_tp4 = _simulate_tp_matmul(A, B, tp_size=4) + + self.assertTrue(torch.equal(result_tp1, result_tp2)) + self.assertTrue(torch.equal(result_tp1, result_tp4)) + + def test_bf16_cross_tp_k512(self): + torch.manual_seed(42) + A = torch.randn(8, 512, dtype=torch.bfloat16) + B = torch.randn(512, 16, dtype=torch.bfloat16) + + result_tp1 = matmul_tp_persistent(A, B) + result_tp2 = _simulate_tp_matmul(A, B, tp_size=2) + result_tp4 = _simulate_tp_matmul(A, B, tp_size=4) + + self.assertTrue(torch.equal(result_tp1, result_tp2)) + self.assertTrue(torch.equal(result_tp1, result_tp4)) + + def test_bf16_cross_tp_k1024(self): + torch.manual_seed(7) + A = torch.randn(4, 1024, dtype=torch.bfloat16) + B = torch.randn(1024, 32, dtype=torch.bfloat16) + + result_tp1 = matmul_tp_persistent(A, B) + result_tp2 = _simulate_tp_matmul(A, B, tp_size=2) + result_tp4 = _simulate_tp_matmul(A, B, tp_size=4) + result_tp8 = _simulate_tp_matmul(A, B, tp_size=8) + + self.assertTrue(torch.equal(result_tp1, result_tp2)) + self.assertTrue(torch.equal(result_tp1, result_tp4)) + self.assertTrue(torch.equal(result_tp1, result_tp8)) + + def test_bf16_cross_tp_k2048_all_sizes(self): + """K=2048: 16 blocks. TP={1,2,4,8,16} all produce isomorphic trees.""" + torch.manual_seed(102) + A = torch.randn(4, 2048, dtype=torch.bfloat16) + B = torch.randn(2048, 16, dtype=torch.bfloat16) + + result_tp1 = matmul_tp_persistent(A, B) + for tp_size in [2, 4, 8, 16]: + result_tpN = _simulate_tp_matmul(A, B, tp_size=tp_size) + self.assertTrue( + torch.equal(result_tp1, result_tpN), + f"TP=1 != TP={tp_size} for K=2048", + ) + + def test_bf16_cross_tp_k4096(self): + """K=4096: 32 blocks.""" + torch.manual_seed(103) + A = torch.randn(2, 4096, dtype=torch.bfloat16) + B = torch.randn(4096, 8, dtype=torch.bfloat16) + + result_tp1 = matmul_tp_persistent(A, B) + for tp_size in [2, 4, 8]: + result_tpN = _simulate_tp_matmul(A, B, tp_size=tp_size) + self.assertTrue( + torch.equal(result_tp1, result_tpN), + f"TP=1 != TP={tp_size} for K=4096", + ) + + def test_fp16_cross_tp_k512(self): + torch.manual_seed(500) + A = torch.randn(4, 512, dtype=torch.float16) + B = torch.randn(512, 16, dtype=torch.float16) + + result_tp1 = matmul_tp_persistent(A, B) + result_tp4 = _simulate_tp_matmul(A, B, tp_size=4) + + self.assertTrue(torch.equal(result_tp1, result_tp4)) + + def test_torch_ops_dispatch_bf16_cross_tp(self): + """torch.ops.tp_inv_ops.matmul_tp_inv dispatch preserves cross-TP invariance.""" + torch.manual_seed(400) + A = torch.randn(4, 512, dtype=torch.bfloat16) + B = torch.randn(512, 16, dtype=torch.bfloat16) + + result_full = torch.ops.tp_inv_ops.matmul_tp_inv(A, B) + + K = A.shape[1] + shard = K // 4 + partials = [] + for r in range(4): + s, e = r * shard, (r + 1) * shard + partials.append(torch.ops.tp_inv_ops.matmul_tp_inv(A[:, s:e], B[s:e, :])) + result_tp4 = _fixed_tree_sum_tensors(partials) + + self.assertTrue(torch.equal(result_full, result_tp4)) + + +class TestTPInvarianceDeterminism(unittest.TestCase): + """Same TP degree, two runs -> bitwise identical. Works for all production K.""" + + def _assert_deterministic(self, K, tp_size, dtype=torch.bfloat16): + torch.manual_seed(42) + A = torch.randn(4, K, dtype=dtype) + B = torch.randn(K, 16, dtype=dtype) + + result_a = _simulate_tp_matmul(A, B, tp_size=tp_size) + result_b = _simulate_tp_matmul(A, B, tp_size=tp_size) + self.assertTrue( + torch.equal(result_a, result_b), + f"K={K} TP={tp_size} dtype={dtype} not deterministic", + ) + + def test_bf16_k3584_tp2(self): + """Qwen3-4B hidden_size. 3584/2=1792 -> 14 blocks per shard.""" + self._assert_deterministic(3584, tp_size=2) + + def test_bf16_k3584_tp4(self): + self._assert_deterministic(3584, tp_size=4) + + def test_bf16_k5120_tp4(self): + """Qwen3-30B hidden_size. 5120/4=1280 -> 10 blocks per shard.""" + self._assert_deterministic(5120, tp_size=4) + + def test_bf16_k5120_tp8(self): + self._assert_deterministic(5120, tp_size=8) + + def test_bf16_k4096_tp8(self): + self._assert_deterministic(4096, tp_size=8) + + def test_fp32_k3584_tp2(self): + self._assert_deterministic(3584, tp_size=2, dtype=torch.float32) + + def test_fp32_accum_deterministic(self): + """fp32_accum=True is deterministic for a fixed TP degree.""" + torch.manual_seed(200) + A = torch.randn(4, 512, dtype=torch.bfloat16) + B = torch.randn(512, 16, dtype=torch.bfloat16) + + result_a = _simulate_tp_matmul(A, B, tp_size=2, fp32_accum=True) + result_b = _simulate_tp_matmul(A, B, tp_size=2, fp32_accum=True) + self.assertTrue(torch.equal(result_a, result_b)) + + def test_fp32_accum_output_dtype_is_input_dtype(self): + torch.manual_seed(300) + A = torch.randn(4, 512, dtype=torch.bfloat16) + B = torch.randn(512, 16, dtype=torch.bfloat16) + + result = matmul_tp_persistent(A, B, fp32_accum=True) + self.assertEqual(result.dtype, torch.bfloat16) + + +# --------------------------------------------------------------------------- +# BFloat16 ops: approximate correctness vs torch +# --------------------------------------------------------------------------- +class TestBFloat16Ops(unittest.TestCase): + def test_moe_sum_tree_reduce_bf16(self): + input_tensor = torch.tensor( + [[[1.0, 2.0], [3.0, 4.0]]], + dtype=torch.bfloat16, + ) + curr_topk_ids = torch.tensor([[0, 1]], dtype=torch.int64) + output = torch.empty(1, 2, dtype=torch.bfloat16) + + moe_sum_tree_reduce( + input=input_tensor, + output=output, + curr_topk_ids=curr_topk_ids, + routed_scaling_factor=1.0, + E=2, + ) + + expected = torch.tensor([[4.0, 6.0]], dtype=torch.bfloat16) + torch.testing.assert_close(output, expected) + + +# --------------------------------------------------------------------------- +# MoE tree reduce: EP/slot invariance +# --------------------------------------------------------------------------- +class TestMoeReduceSlotInvariance(unittest.TestCase): + """moe_sum_tree_reduce must produce bitwise identical results regardless + of which topk slot an expert appears in. This is the EP invariance + property: different EP ranks may route the same experts to different + slot positions, but the tree-reduce result must be bitwise identical.""" + + def _run_moe_reduce(self, input_tensor, topk_ids, E, scaling=1.0): + output = torch.zeros( + input_tensor.shape[0], input_tensor.shape[2], dtype=input_tensor.dtype + ) + moe_sum_tree_reduce( + input=input_tensor, + output=output, + curr_topk_ids=topk_ids, + routed_scaling_factor=scaling, + E=E, + ) + return output + + def _permute_slots(self, values, ids_src, ids_dst, topk): + """Rearrange values so that each expert's data moves to its new slot.""" + tokens = values.shape[0] + result = torch.zeros_like(values) + for t in range(tokens): + for slot_dst in range(topk): + eid = ids_dst[t, slot_dst].item() + if eid == -1: + continue + slot_src = (ids_src[t] == eid).nonzero(as_tuple=True)[0].item() + result[t, slot_dst] = values[t, slot_src] + return result + + def test_slot_permutation_gives_identical_result_fp32(self): + torch.manual_seed(10) + H = 64 + values = torch.randn(1, 4, H, dtype=torch.float32) + + ids_a = torch.tensor([[0, 1, 2, 3]], dtype=torch.int64) + ids_b = torch.tensor([[2, 0, 3, 1]], dtype=torch.int64) + + input_b = self._permute_slots(values, ids_a, ids_b, topk=4) + + result_a = self._run_moe_reduce(values, ids_a, E=4) + result_b = self._run_moe_reduce(input_b, ids_b, E=4) + self.assertTrue(torch.equal(result_a, result_b)) + + def test_slot_permutation_gives_identical_result_bf16(self): + torch.manual_seed(20) + H = 128 + values = torch.randn(2, 8, H, dtype=torch.bfloat16) + + expert_ids = list(range(8)) + rng = random.Random(42) + shuffled = expert_ids.copy() + rng.shuffle(shuffled) + + ids_a = torch.tensor([expert_ids, expert_ids], dtype=torch.int64) + ids_b = torch.tensor([shuffled, shuffled], dtype=torch.int64) + + input_b = self._permute_slots(values, ids_a, ids_b, topk=8) + + result_a = self._run_moe_reduce(values, ids_a, E=8) + result_b = self._run_moe_reduce(input_b, ids_b, E=8) + self.assertTrue(torch.equal(result_a, result_b)) + + def test_slot_invariance_with_remote_experts(self): + """Remote experts (-1) in different positions must not change the result.""" + torch.manual_seed(30) + H = 32 + values_a = torch.randn(1, 4, H, dtype=torch.float32) + + ids_a = torch.tensor([[0, -1, 2, -1]], dtype=torch.int64) + ids_b = torch.tensor([[-1, 2, -1, 0]], dtype=torch.int64) + + values_b = torch.randn(1, 4, H, dtype=torch.float32) + for slot_b in range(4): + eid = ids_b[0, slot_b].item() + if eid == -1: + continue + slot_a = (ids_a[0] == eid).nonzero(as_tuple=True)[0].item() + values_b[0, slot_b] = values_a[0, slot_a] + + result_a = self._run_moe_reduce(values_a, ids_a, E=4) + result_b = self._run_moe_reduce(values_b, ids_b, E=4) + self.assertTrue(torch.equal(result_a, result_b)) + + def test_moe_bf16_large_hidden_invariance(self): + """Production-scale hidden dim (H=7168) with E=64 in BF16.""" + torch.manual_seed(40) + H = 7168 + topk = 8 + E = 64 + tokens = 2 + + expert_ids = torch.zeros(tokens, topk, dtype=torch.int64) + for t in range(tokens): + chosen = torch.randperm(E)[:topk] + expert_ids[t] = chosen + + values = torch.randn(tokens, topk, H, dtype=torch.bfloat16) + + output_a = torch.zeros(tokens, H, dtype=torch.bfloat16) + moe_sum_tree_reduce( + input=values, + output=output_a, + curr_topk_ids=expert_ids, + routed_scaling_factor=0.25, + E=E, + ) + + perm = torch.randperm(topk) + values_b = values[:, perm, :] + ids_b = expert_ids[:, perm] + + output_b = torch.zeros(tokens, H, dtype=torch.bfloat16) + moe_sum_tree_reduce( + input=values_b, + output=output_b, + curr_topk_ids=ids_b, + routed_scaling_factor=0.25, + E=E, + ) + + self.assertTrue(torch.equal(output_a, output_b)) + + def test_moe_deterministic_two_runs(self): + """Same input, two calls -> bitwise identical.""" + torch.manual_seed(50) + values = torch.randn(4, 4, 256, dtype=torch.bfloat16) + ids = torch.tensor( + [[0, 1, 2, 3], [3, 2, 1, 0], [0, 0, 1, 1], [2, 3, 0, 1]], + dtype=torch.int64, + ) + + out_a = torch.zeros(4, 256, dtype=torch.bfloat16) + moe_sum_tree_reduce( + input=values, + output=out_a, + curr_topk_ids=ids, + routed_scaling_factor=0.5, + E=4, + ) + + out_b = torch.zeros(4, 256, dtype=torch.bfloat16) + moe_sum_tree_reduce( + input=values, + output=out_b, + curr_topk_ids=ids, + routed_scaling_factor=0.5, + E=4, + ) + + self.assertTrue(torch.equal(out_a, out_b)) + + +# --------------------------------------------------------------------------- +# tree_all_reduce_sum: non-distributed +# --------------------------------------------------------------------------- +class TestTreeAllReduceNonDistributed(unittest.TestCase): + def test_returns_clone_when_dist_not_initialized(self): + if dist.is_initialized(): + self.skipTest("dist already initialized") + x = torch.tensor([1.0, 2.0, 3.0]) + result = tree_all_reduce_sum(x) + self.assertTrue(torch.equal(x, result)) + self.assertFalse(x.data_ptr() == result.data_ptr()) + + def test_fixed_tree_sum_is_order_deterministic(self): + torch.manual_seed(50) + world_size = 8 + shards = [torch.randn(16, dtype=torch.bfloat16) for _ in range(world_size)] + + result_a = _fixed_tree_sum_tensors(shards) + result_b = _fixed_tree_sum_tensors(list(shards)) + self.assertTrue(torch.equal(result_a, result_b)) + + def test_tree_sum_power_of_two_sizes(self): + for n in [1, 2, 4, 8, 16]: + shards = [torch.tensor([float(i + 1)]) for i in range(n)] + result = _fixed_tree_sum_tensors(shards) + self.assertAlmostEqual(result.item(), n * (n + 1) / 2, places=5) + + +# --------------------------------------------------------------------------- +# MoE reduce: order-matters proof +# --------------------------------------------------------------------------- +class TestMoeReduceOrderMatters(unittest.TestCase): + def test_expert_order_differs_from_slot_order(self): + input_tensor = torch.tensor( + [[[1.0e16, 0.0], [1.0, 0.0], [-1.0e16, 0.0], [2.0, 0.0]]], + dtype=torch.float32, + ) + curr_topk_ids = torch.tensor([[2, 0, 3, 1]], dtype=torch.int64) + output = torch.empty(1, 2, dtype=torch.float32) + + moe_sum_tree_reduce( + input=input_tensor, + output=output, + curr_topk_ids=curr_topk_ids, + routed_scaling_factor=1.0, + E=4, + ) + + slot_order_sum = ( + input_tensor[0, 0] + + input_tensor[0, 1] + + input_tensor[0, 2] + + input_tensor[0, 3] + ) + self.assertNotEqual(output[0, 0].item(), slot_order_sum[0].item()) + + +# --------------------------------------------------------------------------- +# Edge cases +# --------------------------------------------------------------------------- +class TestEdgeCases(unittest.TestCase): + def test_fixed_tree_sum_single_tensor(self): + t = torch.tensor([5.0]) + result = _fixed_tree_sum_tensors([t]) + self.assertEqual(result.item(), 5.0) + + def test_fixed_tree_sum_odd_count(self): + values = [torch.tensor([1.0]), torch.tensor([2.0]), torch.tensor([3.0])] + result = _fixed_tree_sum_tensors(values) + self.assertEqual(result.item(), 6.0) + + def test_fixed_tree_sum_empty_raises(self): + with self.assertRaises(ValueError): + _fixed_tree_sum_tensors([]) + + def test_matmul_tp_persistent_k_less_than_block(self): + torch.manual_seed(0) + A = torch.randn(3, 64, dtype=torch.float32) + B = torch.randn(64, 5, dtype=torch.float32) + actual = matmul_tp_persistent(A, B) + expected = A @ B + torch.testing.assert_close(actual, expected) + + def test_matmul_tp_persistent_k_equals_block(self): + torch.manual_seed(0) + A = torch.randn(3, 128, dtype=torch.float32) + B = torch.randn(128, 5, dtype=torch.float32) + actual = matmul_tp_persistent(A, B) + expected = A @ B + torch.testing.assert_close(actual, expected) + + def test_matmul_tp_persistent_k_multi_block_non_aligned(self): + """K not divisible by BLOCK_K: reference handles remainder gracefully.""" + torch.manual_seed(0) + A = torch.randn(3, 300, dtype=torch.float32) + B = torch.randn(300, 5, dtype=torch.float32) + actual = matmul_tp_persistent(A, B) + expected = A @ B + torch.testing.assert_close(actual, expected, rtol=1e-5, atol=1e-5) + + def test_moe_sum_tree_reduce_single_expert(self): + input_tensor = torch.tensor([[[1.0, 2.0]]], dtype=torch.float32) + curr_topk_ids = torch.tensor([[0]], dtype=torch.int64) + output = torch.empty(1, 2, dtype=torch.float32) + + moe_sum_tree_reduce( + input=input_tensor, + output=output, + curr_topk_ids=curr_topk_ids, + routed_scaling_factor=1.0, + E=1, + ) + expected = torch.tensor([[1.0, 2.0]], dtype=torch.float32) + torch.testing.assert_close(output, expected) + + def test_moe_sum_tree_reduce_single_topk(self): + input_tensor = torch.tensor( + [[[10.0, 20.0]], [[30.0, 40.0]]], + dtype=torch.float32, + ) + curr_topk_ids = torch.tensor([[1], [0]], dtype=torch.int64) + output = torch.empty(2, 2, dtype=torch.float32) + + moe_sum_tree_reduce( + input=input_tensor, + output=output, + curr_topk_ids=curr_topk_ids, + routed_scaling_factor=0.5, + E=2, + ) + expected = torch.tensor([[5.0, 10.0], [15.0, 20.0]], dtype=torch.float32) + torch.testing.assert_close(output, expected) + + def test_moe_sum_tree_reduce_all_remote(self): + input_tensor = torch.tensor( + [[[99.0, 99.0], [99.0, 99.0]]], + dtype=torch.float32, + ) + curr_topk_ids = torch.tensor([[-1, -1]], dtype=torch.int64) + output = torch.empty(1, 2, dtype=torch.float32) + + moe_sum_tree_reduce( + input=input_tensor, + output=output, + curr_topk_ids=curr_topk_ids, + routed_scaling_factor=1.0, + E=2, + ) + expected = torch.zeros(1, 2, dtype=torch.float32) + torch.testing.assert_close(output, expected) + + def test_moe_sum_tree_reduce_duplicate_expert_ids(self): + """Same expert in multiple slots: both contributions must be summed.""" + input_tensor = torch.tensor( + [[[10.0, 20.0], [30.0, 40.0]]], + dtype=torch.float32, + ) + curr_topk_ids = torch.tensor([[0, 0]], dtype=torch.int64) + output = torch.empty(1, 2, dtype=torch.float32) + + moe_sum_tree_reduce( + input=input_tensor, + output=output, + curr_topk_ids=curr_topk_ids, + routed_scaling_factor=1.0, + E=2, + ) + expected = torch.tensor([[40.0, 60.0]], dtype=torch.float32) + torch.testing.assert_close(output, expected) + + +# --------------------------------------------------------------------------- +# Input validation +# --------------------------------------------------------------------------- +class TestInputValidation(unittest.TestCase): + def test_matmul_rejects_non_2d_inputs(self): + with self.assertRaisesRegex(ValueError, "expected 2D"): + matmul_tp_persistent(torch.randn(2, 3, 4), torch.randn(4, 5)) + with self.assertRaisesRegex(ValueError, "expected 2D"): + matmul_tp_persistent(torch.randn(2, 3), torch.randn(3)) + + def test_matmul_rejects_dimension_mismatch(self): + with self.assertRaisesRegex(ValueError, "dimension mismatch"): + matmul_tp_persistent(torch.randn(2, 3), torch.randn(4, 5)) + + def test_moe_rejects_wrong_input_dims(self): + with self.assertRaisesRegex(ValueError, "tokens, topk, hidden"): + moe_sum_tree_reduce( + input=torch.zeros(4, 8), + output=torch.zeros(4, 8), + curr_topk_ids=torch.zeros(4, 2, dtype=torch.int64), + routed_scaling_factor=1.0, + E=2, + ) + + def test_moe_rejects_wrong_topk_ids_dims(self): + with self.assertRaisesRegex(ValueError, "tokens, topk"): + moe_sum_tree_reduce( + input=torch.zeros(4, 2, 8), + output=torch.zeros(4, 8), + curr_topk_ids=torch.zeros(4, dtype=torch.int64), + routed_scaling_factor=1.0, + E=2, + ) + + def test_moe_rejects_shape_mismatch_between_input_and_ids(self): + with self.assertRaisesRegex(ValueError, "must match"): + moe_sum_tree_reduce( + input=torch.zeros(4, 2, 8), + output=torch.zeros(4, 8), + curr_topk_ids=torch.zeros(4, 3, dtype=torch.int64), + routed_scaling_factor=1.0, + E=2, + ) + + def test_moe_rejects_wrong_output_shape(self): + with self.assertRaisesRegex(ValueError, "output must have shape"): + moe_sum_tree_reduce( + input=torch.zeros(4, 2, 8), + output=torch.zeros(4, 4), + curr_topk_ids=torch.zeros(4, 2, dtype=torch.int64), + routed_scaling_factor=1.0, + E=2, + ) + + def test_is_power_of_two_helper(self): + self.assertTrue(_is_power_of_two(1)) + self.assertTrue(_is_power_of_two(2)) + self.assertTrue(_is_power_of_two(64)) + self.assertFalse(_is_power_of_two(0)) + self.assertFalse(_is_power_of_two(3)) + self.assertFalse(_is_power_of_two(6)) + + +# --------------------------------------------------------------------------- +# Distributed tree all-reduce (multi-GPU only) +# --------------------------------------------------------------------------- +class TestDistributedTreeAllReduce(unittest.TestCase): + @unittest.skipUnless( + int(os.environ.get("WORLD_SIZE", "1")) > 1, + "requires torchrun with WORLD_SIZE > 1", + ) + def test_tree_all_reduce_sum_distributed(self): + own_pg = False + if not dist.is_initialized(): + backend = "nccl" if torch.cuda.is_available() else "gloo" + dist.init_process_group(backend=backend) + own_pg = True + + try: + world_size = dist.get_world_size() + if world_size & (world_size - 1) != 0: + self.skipTest("requires power-of-two world size") + + local_rank = int(os.environ.get("LOCAL_RANK", "0")) + if torch.cuda.is_available(): + torch.cuda.set_device(local_rank) + device = torch.device("cuda", local_rank) + else: + device = torch.device("cpu") + + rank = dist.get_rank() + value = torch.full((4,), float(rank + 1), device=device) + actual = tree_all_reduce_sum(value) + expected = torch.full( + (4,), + float(world_size * (world_size + 1) // 2), + device=device, + ) + + torch.testing.assert_close(actual, expected) + dist.barrier() + finally: + if own_pg: + dist.destroy_process_group() + + @unittest.skipUnless( + int(os.environ.get("WORLD_SIZE", "1")) > 1, + "requires torchrun with WORLD_SIZE > 1", + ) + def test_tree_all_reduce_bf16_bitwise_deterministic(self): + """Run tree all-reduce twice with same inputs, verify bitwise identical.""" + own_pg = False + if not dist.is_initialized(): + backend = "nccl" if torch.cuda.is_available() else "gloo" + dist.init_process_group(backend=backend) + own_pg = True + + try: + world_size = dist.get_world_size() + if world_size & (world_size - 1) != 0: + self.skipTest("requires power-of-two world size") + + local_rank = int(os.environ.get("LOCAL_RANK", "0")) + if torch.cuda.is_available(): + torch.cuda.set_device(local_rank) + device = torch.device("cuda", local_rank) + else: + device = torch.device("cpu") + + rank = dist.get_rank() + torch.manual_seed(rank * 1000 + 777) + value = torch.randn(256, device=device, dtype=torch.bfloat16) + + result_a = tree_all_reduce_sum(value) + result_b = tree_all_reduce_sum(value) + + self.assertTrue(torch.equal(result_a, result_b)) + dist.barrier() + finally: + if own_pg: + dist.destroy_process_group() + + +if __name__ == "__main__": + unittest.main() From 0be4d1ae4ac847e3ac023edff8a3b5a1f1b1d5c8 Mon Sep 17 00:00:00 2001 From: yueming-yuan Date: Thu, 21 May 2026 17:37:17 -0700 Subject: [PATCH 31/50] Fallback DeepGEMM activation for unsupported shapes --- python/sglang/srt/layers/moe/moe_runner/deep_gemm.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py b/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py index 61af5533f5b3..b6700ff57512 100644 --- a/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py +++ b/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py @@ -863,8 +863,11 @@ def _varlen_deep_gemm_silu_mul_quant( dtype=torch.float8_e4m3fn, ) - if envs.SGLANG_OPT_USE_JIT_EP_ACTIVATION.get(): - assert N % 4 == 0 and G % 4 == 0 + use_jit_ep_activation = envs.SGLANG_OPT_USE_JIT_EP_ACTIVATION.get() + if N % 4 != 0 or G % 4 != 0: + use_jit_ep_activation = False + + if use_jit_ep_activation: packed_ue8m0 = deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0 down_input_scale = torch.empty( (E, G // 4, N) if packed_ue8m0 else (E, N, G), From caed371b0da1d1a9c7d14e89344356192395549f Mon Sep 17 00:00:00 2001 From: yueming-yuan Date: Thu, 21 May 2026 20:14:38 -0700 Subject: [PATCH 32/50] Fix GLM4 MoE Lite shared expert TP flag --- python/sglang/srt/models/glm4_moe_lite.py | 15 ++++++++------- 1 file changed, 8 insertions(+), 7 deletions(-) diff --git a/python/sglang/srt/models/glm4_moe_lite.py b/python/sglang/srt/models/glm4_moe_lite.py index e0b06f6d78bd..823c8ec1cf48 100644 --- a/python/sglang/srt/models/glm4_moe_lite.py +++ b/python/sglang/srt/models/glm4_moe_lite.py @@ -273,10 +273,16 @@ def __init__( self.shared_experts_is_int8 = False self.shared_experts_is_fp8 = False + self._shared_expert_tp1 = False # self.shared_experts_weight_block_size = None if config.n_shared_experts is not None and self.num_fused_shared_experts == 0: intermediate_size = config.moe_intermediate_size * config.n_shared_experts # disable tp for shared experts when enable deepep moe, or with fp4 allgather + _shared_expert_use_tp1 = ( + get_moe_a2a_backend().is_deepep() + or get_moe_a2a_backend().is_mooncake() + or should_use_flashinfer_cutlass_moe_fp4_allgather() + ) self.shared_experts = Glm4MoeLiteMLP( hidden_size=config.hidden_size, intermediate_size=intermediate_size, @@ -284,14 +290,9 @@ def __init__( quant_config=quant_config, reduce_results=False, prefix=add_prefix("shared_experts", prefix), - **( - dict(tp_rank=0, tp_size=1) - if get_moe_a2a_backend().is_deepep() - or get_moe_a2a_backend().is_mooncake() - or should_use_flashinfer_cutlass_moe_fp4_allgather() - else {} - ), + **(dict(tp_rank=0, tp_size=1) if _shared_expert_use_tp1 else {}), ) + self._shared_expert_tp1 = _shared_expert_use_tp1 is_packed_weight = hasattr( self.shared_experts.gate_up_proj.quant_method, "quant_config" ) From 244ff926fafe9cfeb491bbb535c79e12e6678793 Mon Sep 17 00:00:00 2001 From: yueming-yuan Date: Thu, 21 May 2026 21:18:00 -0700 Subject: [PATCH 33/50] Fix GLM NextN draft KV cache v head dim --- python/sglang/srt/configs/model_config.py | 13 +++++++++++++ 1 file changed, 13 insertions(+) diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index 111145ef6d2f..02afafa0a664 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -658,6 +658,19 @@ def _derive_model_shapes(self): self.scaling = compute_mla_mscale_scaling( self.hf_config.rope_scaling, self.scaling ) + elif "Glm4MoeForCausalLMNextN" in self.hf_config.architectures: + if self.head_dim is None: + self.head_dim = getattr( + self.hf_text_config, + "qk_rope_head_dim", + self.hf_text_config.hidden_size + // self.hf_text_config.num_attention_heads, + ) + if self.swa_head_dim is None: + self.swa_head_dim = self.head_dim + self.v_head_dim = self.head_dim + self.swa_v_head_dim = self.swa_head_dim + self.attention_arch = AttentionArch.MHA elif "MiniCPM3ForCausalLM" in self.hf_config.architectures: self.head_dim = 128 self.attention_arch = AttentionArch.MLA From 71bda1af3b01f2821b077cd91d25da2a15573e4d Mon Sep 17 00:00:00 2001 From: yueming-yuan Date: Fri, 22 May 2026 13:21:11 -0700 Subject: [PATCH 34/50] Use logical seqlen for routed topk returns --- .../managers/scheduler_output_processor_mixin.py | 13 ++++++++----- 1 file changed, 8 insertions(+), 5 deletions(-) diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py index 0dbc50b1deb1..d95435dec962 100644 --- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py +++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py @@ -126,25 +126,27 @@ def maybe_collect_routed_experts(self: Scheduler, req: Req): if capturer is None: return start_len = req.routed_experts_start_len + seqlen = len(req.origin_input_ids) + len(req.output_ids_through_stop) req.routed_experts = capturer.get_topk( req_pool_idx=req.req_pool_idx, - seqlen=req.seqlen, + seqlen=seqlen, req_to_token_pool=self.req_to_token_pool, start_len=start_len, ) - expected_rows = max(0, req.seqlen - 1 - start_len) + expected_rows = max(0, seqlen - 1 - start_len) if ( req.routed_experts is not None and req.routed_experts.shape[0] != expected_rows ): logger.warning( - "routed_experts row-count mismatch for req %s: got %d, " - "expected %d (seqlen=%d, cached_tokens=%d, start_len=%s). " + "routed_experts row-count mismatch for req %s: got %d, expected %d " + "(seqlen=%d, raw_seqlen=%d, cached_tokens=%d, start_len=%s). " "This indicates a silent bug.", req.rid, req.routed_experts.shape[0], expected_rows, + seqlen, req.seqlen, req.cached_tokens, req.routed_experts_start_len, @@ -154,9 +156,10 @@ def maybe_collect_indexer_topk(self: Scheduler, req: Req): capturer = get_global_indexer_capturer() if capturer is None: return + seqlen = len(req.origin_input_ids) + len(req.output_ids_through_stop) req.indexer_topk = capturer.get_topk( req_pool_idx=req.req_pool_idx, - seqlen=req.seqlen, + seqlen=seqlen, req_to_token_pool=self.req_to_token_pool, ) From c0cac9ebaab85623626eb95cd3ea4d3fdc522981 Mon Sep 17 00:00:00 2001 From: Nan Jiang <59716405+nanjiangwill@users.noreply.github.com> Date: Mon, 25 May 2026 11:25:14 -0400 Subject: [PATCH 35/50] support kimi 2.5/6 lora (logprob diff exist) (#25141) --- python/sglang/srt/entrypoints/engine.py | 11 +- .../sglang/srt/lora/backend/base_backend.py | 11 +- python/sglang/srt/lora/layers.py | 81 ++++++++++--- python/sglang/srt/lora/lora_moe_runners.py | 22 +++- python/sglang/srt/lora/mem_pool.py | 85 +++++++++++-- .../srt/lora/triton_ops/virtual_experts.py | 113 ++++++++++++------ python/sglang/srt/managers/io_struct.py | 2 +- .../srt/managers/tokenizer_control_mixin.py | 16 +-- python/sglang/srt/managers/tp_worker.py | 16 ++- .../sglang/srt/model_executor/model_runner.py | 9 ++ .../deepseek_common/deepseek_weight_loader.py | 17 ++- 11 files changed, 286 insertions(+), 97 deletions(-) diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index 5b7ba8fc7506..4d57d1cf3267 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -1133,15 +1133,16 @@ def load_lora_adapter_from_tensors( load_format: Optional[str] = None, ): if load_format == "flattened_bucket": - serialized_tensors = tensors + serialized_named_tensors = list(tensors) else: - serialized_tensors = MultiprocessingSerializer.serialize( - tensors, output_str=True - ) + serialized_named_tensors = [ + MultiprocessingSerializer.serialize(tensors, output_str=True) + for _ in range(self.server_args.tp_size) + ] lora_req = LoadLoRAAdapterFromTensorsReqInput( lora_name=lora_name, config_dict=config_dict, - serialized_tensors=serialized_tensors, + serialized_named_tensors=serialized_named_tensors, load_format=load_format, ) return self.loop.run_until_complete( diff --git a/python/sglang/srt/lora/backend/base_backend.py b/python/sglang/srt/lora/backend/base_backend.py index 17b7bef1bf7e..2bb59f8eaf8d 100644 --- a/python/sglang/srt/lora/backend/base_backend.py +++ b/python/sglang/srt/lora/backend/base_backend.py @@ -19,6 +19,7 @@ class BaseLoRABackend(LoRABackendLmHeadMixing): def __init__(self, max_loras_per_batch: int, device: torch.device): self.max_loras_per_batch = max_loras_per_batch self.device = device + self.batch_info = None self.init_lm_head_config() def run_lora_a_embedding( @@ -176,10 +177,12 @@ def init_cuda_graph_moe_buffers( """ base = moe_layer.base_layer top_k = base.top_k - qinfo = moe_layer._quant_info - E, N, _ = qinfo.w13_weight.shape - hidden_dim = qinfo.w2_weight.shape[1] - device = qinfo.w13_weight.device + # Derive dims from the base FusedMoE rather than quant-specific tensors, + # so this works for any scheme (FP, WNA16, Marlin-packed, etc.). + E = base.num_local_experts + hidden_dim = base.hidden_size + N = 2 * base.intermediate_size_per_partition + device = next(base.parameters()).device dtype = compute_dtype num_experts = base.num_experts diff --git a/python/sglang/srt/lora/layers.py b/python/sglang/srt/lora/layers.py index 475df00677f7..dacbe52a5a25 100644 --- a/python/sglang/srt/lora/layers.py +++ b/python/sglang/srt/lora/layers.py @@ -41,6 +41,15 @@ def __init__( self.weight = self.base_layer.weight if hasattr(self.base_layer, "bias") and self.base_layer.bias is not None: self.bias = self.base_layer.bias + if hasattr(self.base_layer, "reduce_results"): + self.reduce_results = self.base_layer.reduce_results + # Alias remaining base-layer parameters onto the wrapper so + # `named_parameters(remove_duplicate=True)` yields them at the outer + # path — weight loaders (e.g. FusedMoE's `w13_weight_packed`) lookup + # names without the `.base_layer.` segment. + for _name, _param in base_layer.named_parameters(recurse=False): + if not hasattr(self, _name): + setattr(self, _name, _param) def forward(self, x: torch.Tensor): return self.base_layer.forward(x) @@ -207,8 +216,9 @@ def forward(self, input_: torch.Tensor): ): base_output = self.extra_token_embedding(input_, base_output) - # Apply LoRA if configured - if self.set_lora: + # Apply LoRA if configured. Skip if no batch_info (DP-attention idle + # forward): the base path is correct because no real tokens need LoRA. + if self.set_lora and self.lora_backend.batch_info is not None: # The backend's run_lora_a_embedding now handles both regular # and extra tokens efficiently with CUDA graph support base_output = self.apply_lora(base_output, input_, batch_info) @@ -373,8 +383,8 @@ def forward(self, hidden_states: torch.Tensor): hidden_states, self.weight, bias=getattr(self.base_layer, "bias", None) ) - # Apply LoRA if set - if self.set_lora: + # Apply LoRA if set. Skip in DP-attention idle forward (batch_info unset). + if self.set_lora and self.lora_backend.batch_info is not None: base_output = self.apply_lora(base_output, hidden_states) return base_output @@ -463,7 +473,7 @@ def forward(self, input_: torch.Tensor): self.base_layer, input_, bias ) - if self.set_lora: + if self.set_lora and self.lora_backend.batch_info is not None: output_parallel = self.apply_lora(output_parallel, input_) if self.base_layer.gather_output: @@ -477,9 +487,13 @@ def slice_lora_a_weights(self, A: torch.Tensor, tp_rank: int): return A def slice_lora_b_weights(self, B: torch.Tensor, tp_rank: int): + # See RowParallelLinearWithLoRA.slice_lora_a_weights for why base_layer.tp_rank + # is authoritative: DP-attention makes output_partition_sizes attn_tp-local while + # the caller passes global tp_rank. + local_tp_rank = getattr(self.base_layer, "tp_rank", tp_rank) shard_size = self.base_layer.output_partition_sizes[0] - start_idx = tp_rank * shard_size - end_idx = (tp_rank + 1) * shard_size + start_idx = local_tp_rank * shard_size + end_idx = (local_tp_rank + 1) * shard_size B = B[start_idx:end_idx, :] return B @@ -573,12 +587,15 @@ def slice_lora_a_weights(self, A: torch.Tensor, tp_rank: int): return A def slice_lora_b_weights(self, B: torch.Tensor, tp_rank: int): + # base_layer.tp_rank is authoritative under DP-attention: the caller passes + # the global tp_rank but output_partition_sizes is attn_tp-local. + local_tp_rank = getattr(self.base_layer, "tp_rank", tp_rank) partition_sizes = self.base_layer.output_partition_sizes output_sizes = self.base_layer.output_sizes slices = [] offset = 0 for full_size, part_size in zip(output_sizes, partition_sizes): - start_idx = tp_rank * part_size + start_idx = local_tp_rank * part_size end_idx = start_idx + part_size slices.append(B[offset + start_idx : offset + end_idx, :]) offset += full_size @@ -645,11 +662,14 @@ def slice_lora_b_weights(self, B: torch.Tensor, tp_rank: int) -> torch.Tensor: q_proj_shard_size = base_layer.q_proj_shard_size kv_proj_shard_size = base_layer.kv_proj_shard_size num_kv_head_replicas = base_layer.num_kv_head_replicas + # See RowParallelLinearWithLoRA.slice_lora_a_weights for why base_layer.tp_rank + # is authoritative under DP-attention. + local_tp_rank = getattr(base_layer, "tp_rank", tp_rank) - q_start_idx = q_proj_shard_size * tp_rank + q_start_idx = q_proj_shard_size * local_tp_rank q_end_idx = q_start_idx + q_proj_shard_size - kv_shard_id = tp_rank // num_kv_head_replicas + kv_shard_id = local_tp_rank // num_kv_head_replicas kv_start_idx = kv_proj_shard_size * kv_shard_id kv_end_idx = kv_start_idx + kv_proj_shard_size @@ -731,7 +751,9 @@ def forward(self, input_: torch.Tensor, skip_all_reduce=False, forward_batch=Non and not skip_all_reduce ) - if self.set_lora and should_reduce: + # LoRA skipped when batch_info is None (DP-attention idle forward). + have_batch_info = self.lora_backend.batch_info is not None + if self.set_lora and have_batch_info and should_reduce: lora_a_output = self.lora_backend.run_lora_a_sgemm( input_parallel, self.A_buffer ) @@ -745,7 +767,7 @@ def forward(self, input_: torch.Tensor, skip_all_reduce=False, forward_batch=Non base_output=output_, ) else: - if self.set_lora: + if self.set_lora and have_batch_info: output_parallel = self.apply_lora(output_parallel, input_parallel) if should_reduce: output_ = tensor_model_parallel_all_reduce(output_parallel) @@ -756,9 +778,15 @@ def forward(self, input_: torch.Tensor, skip_all_reduce=False, forward_batch=Non return output_, output_bias def slice_lora_a_weights(self, A: torch.Tensor, tp_rank: int): + # Use base_layer.tp_rank (not the argument) so the slicing rank matches + # the partition group the base layer was built on. For MLA o_proj under + # DP-attention, base_layer.tp_rank is attn_tp_rank while the caller + # passes the global tp_rank; input_size_per_partition is already + # attn_tp-sized, so using global tp_rank overshoots to empty. + local_tp_rank = getattr(self.base_layer, "tp_rank", tp_rank) shard_size = self.base_layer.input_size_per_partition - start_idx = tp_rank * shard_size - end_idx = (tp_rank + 1) * shard_size + start_idx = local_tp_rank * shard_size + end_idx = (local_tp_rank + 1) * shard_size A = A[:, start_idx:end_idx].contiguous() return A @@ -844,7 +872,7 @@ def apply_lora(self, base_output: torch.Tensor, x: torch.Tensor) -> torch.Tensor def forward(self, x: torch.Tensor): bias = self.base_layer.bias if not self.base_layer.skip_bias_add else None output = self.base_layer.quant_method.apply(self.base_layer, x, bias) - if self.set_lora: + if self.set_lora and self.lora_backend.batch_info is not None: output = self.apply_lora(output, x) output_bias = self.base_layer.bias if self.base_layer.skip_bias_add else None return output, output_bias @@ -878,7 +906,19 @@ def __init__( self.experts_shared_outer_loras: bool = False self.lora_use_virtual_experts: bool = False + # Forward for the model's own forward-path dispatch — the outer model + # reads several FusedMoE attributes (e.g. `self.experts.moe_runner_config`, + # `self.experts.dispatcher`, `self.experts.num_local_experts`, + # `self.experts.quant_method`) directly on the wrapper. Quant + # post-processing iterators skip LoRA wrappers via + # `isinstance(module, BaseLayerWithLoRA)` so the packed params on the + # inner FusedMoE get processed there, not here. self.quant_method = base_layer.quant_method + self.moe_runner_config = base_layer.moe_runner_config + self.dispatcher = base_layer.dispatcher + self.num_local_experts = base_layer.num_local_experts + if hasattr(base_layer, "scheme"): + self.scheme = base_layer.scheme self.tp_size = getattr(base_layer, "moe_tp_size", 1) self.tp_rank = getattr(base_layer, "moe_tp_rank", 0) @@ -903,6 +943,12 @@ def __init__( and base_layer.quant_method.runner is not None ): runner_backend = base_layer.quant_method.runner.runner_backend + elif ( + hasattr(base_layer, "scheme") + and hasattr(base_layer.scheme, "runner") + and base_layer.scheme.runner is not None + ): + runner_backend = base_layer.scheme.runner.runner_backend else: runner_backend = MoeRunnerBackend.TRITON @@ -1007,6 +1053,8 @@ def forward(self, hidden_states: torch.Tensor, topk_output: TopKOutput, **kwargs 1. After gate_up projection, before activation 2. After down projection, before final reduction """ + if self.lora_backend.batch_info is None: + return self.base_layer.forward(hidden_states, topk_output, **kwargs) # Build LoRA info for this batch lora_info = self._get_lora_info() @@ -1034,6 +1082,9 @@ def _forward_with_lora( # Use pre-computed quant info (doesn't change so not sure why we need to pass it in every time) quant_info = self._quant_info + quant_info.expert_map = getattr( + base_layer.dispatcher, "local_expert_mapping", None + ) # Run the only lora moe runner (Triton) combine_input = self._lora_runner.run( diff --git a/python/sglang/srt/lora/lora_moe_runners.py b/python/sglang/srt/lora/lora_moe_runners.py index b3f1389b5c01..81105a63f428 100644 --- a/python/sglang/srt/lora/lora_moe_runners.py +++ b/python/sglang/srt/lora/lora_moe_runners.py @@ -205,7 +205,14 @@ def _compute_token_lora_mapping( hidden_states: torch.Tensor, lora_info: LoRAInfo, ) -> torch.Tensor: - """Map each token to its LoRA adapter index (-1 for no LoRA).""" + """Map each token to its LoRA adapter index (-1 for no LoRA). + + Under DP-attention, `hidden_states` is the gathered batch (local + foreign + tokens) but `seg_indptr` / `req_to_lora` cover only local requests, so + `searchsorted` on foreign positions would index one past the end. Pad + `req_to_lora` with a -1 sentinel; foreign outputs are discarded by the + DP-attention scatter anyway. + """ token_positions = torch.arange( hidden_states.shape[0], device=hidden_states.device, dtype=torch.int32 ) @@ -214,7 +221,11 @@ def _compute_token_lora_mapping( token_positions, right=True, ) - return lora_info.req_to_lora.to(torch.int32)[req_indices] + req_to_lora = lora_info.req_to_lora.to(torch.int32) + req_to_lora_padded = torch.cat( + [req_to_lora, req_to_lora.new_full((1,), -1)], dim=0 + ) + return req_to_lora_padded[req_indices] def _compute_lora_alignment( @@ -349,7 +360,6 @@ def _add_lora_gate_up_delta( r = lora_info.max_lora_rank gate_up_a = lora_info.gate_up_lora_a_weights gate_up_b = lora_info.gate_up_lora_b_weights - if lora_info.experts_shared_outer_loras and not lora_info.lora_use_virtual_experts: gate_up_a = gate_up_a.expand(-1, lora_info.num_experts, -1, -1) @@ -359,6 +369,8 @@ def _add_lora_gate_up_delta( if is_gated: inter_size = gate_up_b.shape[2] // 2 lora_a_stacked = [gate_up_a[:, :, :r, :], gate_up_a[:, :, r : 2 * r, :]] + # B halves are also the tuple form the virtual-experts kernel wants + # (one shrink at K=2*r, two expands at K=r each). lora_b_stacked = [ gate_up_b[:, :, :inter_size, :], gate_up_b[:, :, inter_size:, :], @@ -372,7 +384,7 @@ def _add_lora_gate_up_delta( output=intermediate_cache, hidden_states=hidden_states, lora_a=gate_up_a, - lora_b=gate_up_b, + lora_b=tuple(lora_b_stacked) if is_gated else gate_up_b, topk_ids=topk_ids, topk_weights=topk_weights, token_lora_mapping=token_lora_mapping, @@ -447,6 +459,8 @@ def _add_lora_down_delta( down_lora_a = lora_info.down_lora_a_weights down_lora_b = lora_info.down_lora_b_weights if lora_info.experts_shared_outer_loras and not lora_info.lora_use_virtual_experts: + # fused_moe_lora requires B's expert_dim to match A's; expand the + # shared B view. down_lora_b = down_lora_b.expand(-1, lora_info.num_experts, -1, -1) if lora_info.fully_sharded and lora_info.tp_size > 1: diff --git a/python/sglang/srt/lora/mem_pool.py b/python/sglang/srt/lora/mem_pool.py index 2a9bb8b7c34b..dbb5a745231c 100644 --- a/python/sglang/srt/lora/mem_pool.py +++ b/python/sglang/srt/lora/mem_pool.py @@ -284,6 +284,55 @@ def _iter_local_expert_weights( f"Expected dict or 3D torch.Tensor, got {type(weights).__name__}." ) + def _row_parallel_shard_tp( + self, module_name: str, base_model: torch.nn.Module, layer_idx: int + ) -> int: + """Shard count for a non-MoE row-parallel module's activation axis. + + Probes the base module's ``input_size // input_size_per_partition`` so + the LoRA buffer matches the actual shard regardless of which TP group + owns it — covers DP-attention (``o_proj`` uses ``attn_tp_size``) and + shared-expert dense-vs-MoE per-layer-TP differences. Falls back to + ``self.tp_size``. Cached per ``(module_name, layer_idx)``. + + MoE-internal names go through ``self.moe_tp_size`` upstream. + """ + cache = getattr(self, "_row_parallel_tp_cache", None) + if cache is None: + cache = {} + setattr(self, "_row_parallel_tp_cache", cache) + key = (module_name, layer_idx) + if key in cache: + return cache[key] + + layer_markers = (f".layers.{layer_idx}.", f"layers.{layer_idx}.") + + def _probe(m): + in_size = getattr(m, "input_size", None) + per_part = getattr(m, "input_size_per_partition", None) + if in_size is not None and per_part is not None and per_part > 0: + return max(1, in_size // per_part) + inner = getattr(m, "base_layer", None) + if inner is not None and inner is not m: + return _probe(inner) + return None + + suffix = f".{module_name}" + found = None + for _name, module in base_model.named_modules(): + if not _name.endswith(suffix): + continue + if not any(marker in _name for marker in layer_markers): + continue + r = _probe(module) + if r is not None: + found = r + break + + out = found if found is not None else self.tp_size + cache[key] = out + return out + def _get_standard_shape( self, module_name: str, @@ -296,8 +345,12 @@ def _get_standard_shape( module_name, self.base_hf_config, base_model, layer_idx ) c = get_stacked_multiply(module_name, base_model) - if self.tp_size > 1 and module_name in ROW_PARALLELISM_LINEAR_LORA_NAMES: - input_dim = divide(input_dim, self.tp_size) + # Non-MoE row-parallel modules: probe the actual shard size so o_proj / + # down_proj match attn_tp under DP-attention and the shared-experts + # dense-vs-MoE per-layer-TP differences. + row_tp = self._row_parallel_shard_tp(module_name, base_model, layer_idx) + if row_tp > 1 and module_name in ROW_PARALLELISM_LINEAR_LORA_NAMES: + input_dim = divide(input_dim, row_tp) return (self.max_loras_per_batch, max_lora_dim * c, input_dim) def get_lora_A_shape( @@ -318,9 +371,12 @@ def get_lora_A_shape( module_name, self.base_hf_config, base_model, layer_idx ) c = get_stacked_multiply(module_name, base_model) - # MoE modules shard along `moe_tp_size`, not the outer `tp_size`. + # MoE modules shard along `moe_tp_size`; non-MoE row-parallel modules + # use a probed shard that may be attn_tp under DP-attention. effective_tp_size = ( - self.moe_tp_size if self.is_moe_module(module_name) else self.tp_size + self.moe_tp_size + if self.is_moe_module(module_name) + else self._row_parallel_shard_tp(module_name, base_model, layer_idx) ) if ( effective_tp_size > 1 @@ -412,9 +468,12 @@ def get_lora_B_shape( _, output_dim = get_hidden_dim( module_name, self.base_hf_config, base_model, layer_idx ) - # MoE modules shard along `moe_tp_size`, not the outer `tp_size`. + # MoE modules shard along `moe_tp_size`; non-MoE column-parallel modules + # use a probed shard that may be attn_tp under DP-attention. effective_tp_size = ( - self.moe_tp_size if self.is_moe_module(module_name) else self.tp_size + self.moe_tp_size + if self.is_moe_module(module_name) + else self._row_parallel_shard_tp(module_name, base_model, layer_idx) ) if ( effective_tp_size > 1 @@ -765,16 +824,22 @@ def load_lora_weight_tensor( expert_match = re.search(r"experts\.(\d+)\.", name) if expert_match: - # Per-expert MoE weight — 2D tensors, one per expert + # Per-expert MoE weight — 2D tensors, one per expert. + # Init A and B independently: under ``experts_shared_outer_loras``, + # fc1 has shared A (Tensor in temp_A_buffer) + per-expert B + # (dict in temp_B_buffer), and fc2 has the opposite. A shared + # init on both would either clobber the shared Tensor or leave + # the per-expert side as None. target_module = target_module + "_moe" - if temp_A_buffer[target_module] is None: - temp_A_buffer[target_module] = {} - temp_B_buffer[target_module] = {} expert_id = int(expert_match.group(1)) if "lora_A" in name: + if temp_A_buffer[target_module] is None: + temp_A_buffer[target_module] = {} temp_A_buffer[target_module][expert_id] = weights else: + if temp_B_buffer[target_module] is None: + temp_B_buffer[target_module] = {} temp_B_buffer[target_module][expert_id] = weights elif "experts" in name and weights.dim() == 3: # Shared outer MoE weight — 3D tensor [expert_dim, rank, hidden] diff --git a/python/sglang/srt/lora/triton_ops/virtual_experts.py b/python/sglang/srt/lora/triton_ops/virtual_experts.py index 4781dfe504eb..69339ccfe419 100644 --- a/python/sglang/srt/lora/triton_ops/virtual_experts.py +++ b/python/sglang/srt/lora/triton_ops/virtual_experts.py @@ -515,7 +515,7 @@ def _merged_experts_fused_moe_lora_add_impl( output: torch.Tensor, hidden_states: torch.Tensor, lora_a: torch.Tensor, - lora_b: torch.Tensor, + lora_b: torch.Tensor | list[torch.Tensor] | tuple[torch.Tensor, ...], topk_ids: torch.Tensor, topk_weights: torch.Tensor, token_lora_mapping: torch.Tensor, @@ -524,13 +524,30 @@ def _merged_experts_fused_moe_lora_add_impl( experts_shared_outer_loras_b: bool, routing_cache: dict | None = None, ) -> None: + """Fused virtual-experts LoRA delta add. + + ``lora_b`` accepts either a single tensor or a sequence of tensors stacked + along the output dim. Length-2 is the gate_up case where A has rank ``2*r`` + (gate's A and up's A concatenated along rank) and each B has rank ``r``. + The shrink runs once over the full ``2*r`` rank; the expand runs once per + B, each reading its half of the intermediate and writing to its slice of + ``output``. """ - 1. Prepare virtual expert routing metadata from topk_ids + token_lora_mapping * num_experts. - 2. Flatten LoRA weights from [max_loras, num_experts, ...] to [max_loras * num_experts, ...]. - 3. Run regular SGLang fused-MoE kernels for LoRA A and LoRA B. - 4. Mask out tokens with token_lora_mapping == -1 on the add path. - """ + lora_b_list: list[torch.Tensor] = ( + list(lora_b) if isinstance(lora_b, (list, tuple)) else [lora_b] + ) + n_b = len(lora_b_list) + assert n_b in (1, 2), f"lora_b must be length 1 or 2, got {n_b}" + b_rank = lora_b_list[0].shape[3] + for b in lora_b_list[1:]: + assert b.shape == lora_b_list[0].shape, ( + f"all lora_b tensors must share shape; got {[tuple(t.shape) for t in lora_b_list]}" + ) + max_loras, _, max_lora_rank, _ = lora_a.shape + assert max_lora_rank == n_b * b_rank, ( + f"lora_a rank {max_lora_rank} != n_b ({n_b}) * lora_b rank {b_rank}" + ) input_top_k = 1 if hidden_states.shape[0] == topk_ids.numel() else topk_ids.shape[1] def _merge_lora_expert_weight(t: torch.Tensor) -> torch.Tensor: @@ -642,9 +659,10 @@ def _get_routing( ) lora_a_virtual = _merge_lora_expert_weight(lora_a) - lora_b_virtual = _merge_lora_expert_weight(lora_b) + lora_b_virtuals = [_merge_lora_expert_weight(b) for b in lora_b_list] num_experts_a = lora_a.shape[1] - num_experts_b = lora_b.shape[1] + num_experts_b = lora_b_list[0].shape[1] + half_out = lora_b_list[0].shape[2] intermediate = torch.zeros( [token_lora_mapping.shape[0], topk_ids.shape[1], max_lora_rank], @@ -678,7 +696,7 @@ def _get_routing( a_stage_config, ) - b_stage_config = _get_stage_config(lora_b_virtual, 1) + b_stage_config = _get_stage_config(lora_b_virtuals[0], 1) ( sorted_token_ids, expert_ids, @@ -692,33 +710,47 @@ def _get_routing( b_stage_config["BLOCK_SIZE_M"], ) - invoke_fused_moe_kernel( - intermediate.view(-1, max_lora_rank), - lora_b_virtual, - None, - output, - None, - None, - None, - topk_weights, - topk_ids, - sorted_token_ids, - expert_ids, - num_tokens_post_padded, - mul_routed_weight, - 1, - b_stage_config, - tl.bfloat16 if hidden_states.dtype == torch.bfloat16 else tl.float16, - False, - False, - False, - False, - False, - None, - fuse_add_to_output=True, - add_output_mask=token_lora_mask, - router_topk=topk_ids.shape[1], - ) + # n_b expands. For len 1: K=b_rank covers full intermediate, write full output. + # For len 2 (gate_up): split intermediate along rank into [gate, up] halves + # (each contiguous, K=b_rank=r) and output along last dim into [gate, up] + # halves (each of width half_out). Each B in lora_b_virtuals is its own + # half's weight tensor, naturally K=b_rank. + for b_idx, b_virtual in enumerate(lora_b_virtuals): + if n_b == 1: + inter_arg = intermediate.view(-1, b_rank) + out_arg = output + else: + inter_arg = intermediate[..., b_idx * b_rank : (b_idx + 1) * b_rank].contiguous().view(-1, b_rank) + out_arg = output[..., b_idx * half_out : (b_idx + 1) * half_out].contiguous() + invoke_fused_moe_kernel( + inter_arg, + b_virtual, + None, + out_arg, + None, + None, + None, + topk_weights, + topk_ids, + sorted_token_ids, + expert_ids, + num_tokens_post_padded, + mul_routed_weight, + 1, + b_stage_config, + tl.bfloat16 if hidden_states.dtype == torch.bfloat16 else tl.float16, + False, + False, + False, + False, + False, + None, + fuse_add_to_output=True, + add_output_mask=token_lora_mask, + router_topk=topk_ids.shape[1], + ) + if n_b != 1: + output[..., b_idx * half_out : (b_idx + 1) * half_out].copy_(out_arg) def _merged_experts_fused_moe_lora_add_op( @@ -761,7 +793,7 @@ def merged_experts_fused_moe_lora_add( output: torch.Tensor, hidden_states: torch.Tensor, lora_a: torch.Tensor, - lora_b: torch.Tensor, + lora_b: torch.Tensor | list[torch.Tensor] | tuple[torch.Tensor, ...], topk_ids: torch.Tensor, topk_weights: torch.Tensor, token_lora_mapping: torch.Tensor, @@ -770,7 +802,12 @@ def merged_experts_fused_moe_lora_add( experts_shared_outer_loras_b: bool, routing_cache: dict | None = None, ) -> None: - """Public API: wraps the registered op with routing_cache support.""" + """Public API: wraps the registered op with routing_cache support. + + ``lora_b`` accepts a sequence of length 2 for the gate_up case (each B + holds one half of the stacked output, rank ``r``, with A's rank ``2*r``); + a single tensor is used for the down case. + """ _merged_experts_fused_moe_lora_add_impl( output, hidden_states, diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index de740d898751..ebd1e83ce8db 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -1915,7 +1915,7 @@ def to_ref(self) -> LoRARef: class LoadLoRAAdapterFromTensorsReqInput(BaseReq): lora_name: str config_dict: Dict[str, Any] - serialized_tensors: str + serialized_named_tensors: List[Union[str, bytes]] pinned: bool = False added_tokens_config: Optional[Dict[str, Any]] = None lora_id: Optional[str] = None diff --git a/python/sglang/srt/managers/tokenizer_control_mixin.py b/python/sglang/srt/managers/tokenizer_control_mixin.py index a5ae7d4829e0..0fac8dbc736f 100644 --- a/python/sglang/srt/managers/tokenizer_control_mixin.py +++ b/python/sglang/srt/managers/tokenizer_control_mixin.py @@ -586,11 +586,9 @@ async def load_lora_adapter( "LoRA is not enabled. Please set `--enable-lora` to enable LoRA." ) - # TODO (lifuhuang): Remove this after we verify that dynamic lora loading works - # with dp_size > 1. assert ( - self.server_args.dp_size == 1 - ), "dp_size must be 1 for dynamic lora loading" + self.server_args.dp_size == 1 or self.server_args.enable_dp_attention + ), "dp_size must be 1 or dp attention must be enabled for dynamic lora loading" logger.info( "Start load Lora adapter. Lora name=%s, path=%s", obj.lora_name, @@ -665,8 +663,8 @@ async def load_lora_adapter_from_tensors( ) assert ( - self.server_args.dp_size == 1 - ), "dp_size must be 1 for dynamic lora loading" + self.server_args.dp_size == 1 or self.server_args.enable_dp_attention + ), "dp_size must be 1 or dp attention must be enabled for dynamic lora loading" logger.info( "Start load Lora adapter from tensors. Lora name=%s", obj.lora_name, @@ -738,11 +736,9 @@ async def unload_lora_adapter( obj.lora_name is not None ), "lora_name must be provided to unload LoRA adapter" - # TODO (lifuhuang): Remove this after we verify that dynamic lora loading works - # with dp_size > 1. assert ( - self.server_args.dp_size == 1 - ), "dp_size must be 1 for dynamic lora loading" + self.server_args.dp_size == 1 or self.server_args.enable_dp_attention + ), "dp_size must be 1 or dp attention must be enabled for dynamic lora loading" logger.info( "Start unload Lora adapter. Lora name=%s", obj.lora_name, diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py index 1687de74be49..12f1a0784587 100644 --- a/python/sglang/srt/managers/tp_worker.py +++ b/python/sglang/srt/managers/tp_worker.py @@ -194,19 +194,23 @@ def unload_lora_adapter(self, recv_req: UnloadLoRAAdapterReqInput): def load_lora_adapter_from_tensors( self, recv_req: LoadLoRAAdapterFromTensorsReqInput ): - # The LoRA code handles TP sharding internally using slice_lora_a_weights - # and slice_lora_b_weights methods (see lora/layers.py:46-49, mem_pool.py:437-440). + # TP sharding for LoRA happens inside the lora module (see + # lora/layers.py:46-49 and mem_pool.py:437-440). Each TP rank + # deserializes its own producer's bytes — same convention as + # ``update_weights_from_tensor`` above. One producer per one + # consumer means the CUDA-IPC ref counter on the producer's + # bucket drops cleanly each cycle. + monkey_patch_torch_reductions() + serialized = recv_req.serialized_named_tensors[self.tp_rank] if recv_req.load_format == "flattened_bucket": - flattened_data = MultiprocessingSerializer.deserialize( - recv_req.serialized_tensors - ) + flattened_data = MultiprocessingSerializer.deserialize(serialized) bucket = FlattenedTensorBucket( flattened_tensor=flattened_data["flattened_tensor"], metadata=flattened_data["metadata"], ) tensors = dict(bucket.reconstruct_tensors()) else: - tensors = MultiprocessingSerializer.deserialize(recv_req.serialized_tensors) + tensors = MultiprocessingSerializer.deserialize(serialized) result = self.model_runner.load_lora_adapter_from_tensors( recv_req.to_ref(), tensors, diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 48c4f701e1e0..ca4ca4e934c9 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -3723,8 +3723,15 @@ def post_process_weights(self, recv_req): if hasattr(self.model, "post_load_weights"): self.model.post_load_weights() + # LoRA wrappers forward `quant_method` for forward-path dispatch but + # don't own the packed params; skip them here so the inner base layer + # (yielded separately by `named_modules`) handles post-processing. + from sglang.srt.lora.layers import BaseLayerWithLoRA + if recv_req.restore_weights_before_load: for _, module in self.model.named_modules(): + if isinstance(module, BaseLayerWithLoRA): + continue quant_method = getattr(module, "quant_method", None) if quant_method is not None and hasattr( quant_method, "restore_weights_before_loading" @@ -3734,6 +3741,8 @@ def post_process_weights(self, recv_req): if recv_req.post_process_quantization: for _, module in self.model.named_modules(): + if isinstance(module, BaseLayerWithLoRA): + continue quant_method = getattr(module, "quant_method", None) if quant_method is not None and hasattr( quant_method, "process_weights_after_loading" diff --git a/python/sglang/srt/models/deepseek_common/deepseek_weight_loader.py b/python/sglang/srt/models/deepseek_common/deepseek_weight_loader.py index 0f9fbc2590c8..bdcbd369937a 100644 --- a/python/sglang/srt/models/deepseek_common/deepseek_weight_loader.py +++ b/python/sglang/srt/models/deepseek_common/deepseek_weight_loader.py @@ -130,7 +130,15 @@ def do_load_weights( weights, NVFP4_CKPT_FP8_ATTN_QUANT_MODULES, nextn_conf ) - cached_a_proj = {} if self.fuse_qkv_a_proj else None + # Persist across calls: chunked weight updates can split q_a_proj and + # kv_a_proj_with_mqa for the same layer into different chunks, and + # the fusion only fires once both halves have been seen. + if self.fuse_qkv_a_proj: + if not hasattr(self, "_persistent_cached_a_proj"): + self._persistent_cached_a_proj = {} + cached_a_proj = self._persistent_cached_a_proj + else: + cached_a_proj = None if self.num_fused_shared_experts > 0: assert self.num_fused_shared_experts == 1 @@ -261,9 +269,10 @@ def do_load_weights( if self.fuse_qkv_a_proj and ( "q_a_proj" in name or "kv_a_proj_with_mqa" in name ): - cached_a_proj[name] = _clone_if_runai_streamed_tensor( - loaded_weight - ) + # Clone: `loaded_weight` may be a view into an IPC + # bucket that gets reused by the next chunk, and + # the RunAI streamer also relies on cloning. + cached_a_proj[name] = loaded_weight.detach().clone() q_a_proj_name = ( name if "q_a_proj" in name From d9a01589a2a4698d0963156ae80ec2292bf23935 Mon Sep 17 00:00:00 2001 From: Ziang Li Date: Tue, 26 May 2026 13:30:08 -0700 Subject: [PATCH 36/50] [RL] Fix FP8 skip matching for trailing-dot prefixes (#26287) --- .../sglang/srt/layers/quantization/utils.py | 2 + .../test_flashinfer_trtllm_gen_moe_backend.py | 50 +++++++++++++++++++ .../registered/quant/test_is_layer_skipped.py | 11 ++++ 3 files changed, 63 insertions(+) diff --git a/python/sglang/srt/layers/quantization/utils.py b/python/sglang/srt/layers/quantization/utils.py index 2c54d901ebd3..99e3218cfd3f 100644 --- a/python/sglang/srt/layers/quantization/utils.py +++ b/python/sglang/srt/layers/quantization/utils.py @@ -50,6 +50,8 @@ def _module_path_match(ignored: str, prefix: str) -> bool: # match `mlp.gate_up_proj`. Needed for quant configs (e.g. Qwen3.6-FP8) # whose `modules_to_not_convert` lists MoE-template names like `mlp.gate` # that collide with fused dense MLP names by plain substring. + ignored = ignored.rstrip(".") + prefix = prefix.rstrip(".") if ignored == prefix: return True if prefix.startswith(ignored + "."): diff --git a/test/registered/backends/test_flashinfer_trtllm_gen_moe_backend.py b/test/registered/backends/test_flashinfer_trtllm_gen_moe_backend.py index aff581054838..f0528197ba49 100644 --- a/test/registered/backends/test_flashinfer_trtllm_gen_moe_backend.py +++ b/test/registered/backends/test_flashinfer_trtllm_gen_moe_backend.py @@ -155,6 +155,50 @@ def test_gsm8k(self): self.assertGreater(metrics["score"], 0.93) +class FlashinferTrtllmGenMoeBackendMXFP8MixedBF16Base: + backend = None + + @classmethod + def setUpClass(cls): + cls.model = "zianglih/JoyAI-LLM-Flash-MXFP8-last-6-BF16" + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + env={**os.environ, "SGLANG_ENABLE_JIT_DEEPGEMM": "False"}, + other_args=[ + "--kv-cache-dtype", + "bf16", + "--fp8-gemm-backend", + "flashinfer_cutlass", + "--moe-runner-backend", + cls.backend, + "--tp-size", + "4", + "--trust-remote-code", + ], + ) + + @classmethod + def tearDownClass(cls): + kill_process_tree(cls.process.pid) + + def test_gsm8k(self): + args = SimpleNamespace( + base_url=self.base_url, + model=self.model, + eval_name="gsm8k", + api="completion", + max_tokens=512, + num_examples=200, + num_threads=128, + ) + metrics = run_eval(args) + print(f"{metrics=}") + self.assertGreater(metrics["score"], 0.92) + + class FlashinferTrtllmGenMoeBackendNVFP4Base: backend = None @@ -234,6 +278,12 @@ class TestFlashinferTrtllmGenMoeBackendMXFP8Routed( backend = "flashinfer_trtllm_routed" +class TestFlashinferTrtllmRoutedMxfp8MixedBF16( + FlashinferTrtllmGenMoeBackendMXFP8MixedBF16Base, CustomTestCase +): + backend = "flashinfer_trtllm_routed" + + class TestFlashinferTrtllmGenMoeBackendBF16Routed( FlashinferTrtllmGenMoeBackendBF16Base, CustomTestCase ): diff --git a/test/registered/quant/test_is_layer_skipped.py b/test/registered/quant/test_is_layer_skipped.py index 89e2c15ed4a1..5f80846710fe 100644 --- a/test/registered/quant/test_is_layer_skipped.py +++ b/test/registered/quant/test_is_layer_skipped.py @@ -47,6 +47,17 @@ def test_mlp_gate_does_not_match_gate_up_proj(self): ) self.assertTrue(is_layer_skipped("model.layers.0.mlp.gate", ignored, {})) + def test_trailing_dot_prefix_matches_child_modules(self): + # Mixed-precision checkpoints may use a trailing-dot layer prefix to keep + # every module under the layer in higher precision. + ignored = ["model.layers.34."] + self.assertTrue( + is_layer_skipped("model.layers.34.mlp.experts.0.down_proj", ignored, {}) + ) + self.assertFalse( + is_layer_skipped("model.layers.340.mlp.experts.0.down_proj", ignored, {}) + ) + if __name__ == "__main__": unittest.main() From 53e04836902e953eafdbd4d726d595b7b1d0edd1 Mon Sep 17 00:00:00 2001 From: Jiajun Li <48857426+guapisolo@users.noreply.github.com> Date: Thu, 28 May 2026 14:44:43 -0700 Subject: [PATCH 37/50] [sglang-miles] Cherry-pick #26430: Fix GemmaRMSNorm gemma_weight buffer storage for Qwen3.5 (#26429) --- python/sglang/srt/layers/layernorm.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index 7158aabedff5..230a1ac12e4d 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -593,7 +593,9 @@ def __init__( super().__init__() self.weight = nn.Parameter(torch.zeros(hidden_size)) self.variance_epsilon = eps - self.register_buffer("gemma_weight", self.weight.data + 1.0, persistent=False) + self.register_buffer( + "gemma_weight", torch.ones_like(self.weight), persistent=False + ) # (Chen-0210) Gemma weight = standard_weight + 1. Precompute once. # If TRTLLM allreduce fusion ever provides gemma-style norm # natively, this can be removed. @@ -602,7 +604,8 @@ def __init__( def _weight_loader(self, param: torch.Tensor, loaded_weight: torch.Tensor) -> None: assert param.size() == loaded_weight.size() param.data.copy_(loaded_weight) - self.gemma_weight = param.data + 1.0 + # Keep storage stable for CUDA graphs or fused paths that capture this buffer. + torch.add(param.data, 1.0, out=self.gemma_weight) def _forward_impl( self, From e94a84e011c8bd18b1e69415bcad0a44e168e355 Mon Sep 17 00:00:00 2001 From: Shangming Cai Date: Tue, 19 May 2026 02:08:10 +0800 Subject: [PATCH 38/50] [Bugfix] Fix missing group arg in get dp buffer (#25585) --- python/sglang/srt/layers/communicator.py | 6 +++++- python/sglang/srt/models/deepseek_v4.py | 16 +++++++++++++--- 2 files changed, 18 insertions(+), 4 deletions(-) diff --git a/python/sglang/srt/layers/communicator.py b/python/sglang/srt/layers/communicator.py index 48de54b1a09b..5b0196b13ed5 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -1311,8 +1311,12 @@ def _scatter_hidden_states_moe( # DP scatter (if DP attention is enabled) if context.attn_dp_size > 1: + if get_tensor_model_parallel_world_size() == get_attention_dp_size(): + group = get_tp_group() + else: + group = get_attention_tp_group() hidden_states_output, global_hidden_states = ( - get_local_dp_buffer(), + get_local_dp_buffer(group), hidden_states, ) dp_scatter(hidden_states_output, global_hidden_states, forward_batch) diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index b8df0c3135e2..7f0ac2aa744f 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -26,7 +26,11 @@ fused_rope_inplace, ) from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config -from sglang.srt.distributed import get_pp_group, get_tensor_model_parallel_world_size +from sglang.srt.distributed import ( + get_pp_group, + get_tensor_model_parallel_world_size, + get_tp_group, +) from sglang.srt.environ import envs from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation from sglang.srt.layers.attention.dsv4.compressor import Compressor @@ -837,7 +841,10 @@ def forward( input_ids = input_ids[cp_rank::cp_size].contiguous() input_ids_global = input_ids elif _use_tp_moe_gather: - hidden_states, local_hidden_states = get_global_dp_buffer(), hidden_states + hidden_states, local_hidden_states = ( + get_global_dp_buffer(get_tp_group()), + hidden_states, + ) dp_gather_partial(hidden_states, local_hidden_states, forward_batch) _a2a_scatter_chunks: Optional[List[torch.Tensor]] = None if _use_tp_attn_a2a_scatter: @@ -853,7 +860,10 @@ def forward( input_ids_global=input_ids_global, ) if _use_tp_moe_gather: - hidden_states, global_hidden_states = get_local_dp_buffer(), hidden_states + hidden_states, global_hidden_states = ( + get_local_dp_buffer(get_tp_group()), + hidden_states, + ) dp_scatter(hidden_states, global_hidden_states, forward_batch) if _use_tp_attn_a2a_scatter: assert _a2a_scatter_chunks is not None From 3102015cad599f2b87b98d039f9cde052ef47a72 Mon Sep 17 00:00:00 2001 From: Allen Zhu Date: Mon, 1 Jun 2026 17:15:00 -0700 Subject: [PATCH 39/50] fix: is_true_on_policy_enabled() does not take arguments (#26736) --- python/sglang/srt/model_executor/forward_batch_info.py | 8 ++------ 1 file changed, 2 insertions(+), 6 deletions(-) diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 2557c8073c9e..070ad4cf6de2 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -772,9 +772,7 @@ def _compute_mrope_positions( mm_input = batch.multimodal_inputs[batch_idx] if self.forward_mode.is_decode(): # 3 * N - if mm_input is None or is_true_on_policy_enabled( - get_global_server_args() - ): + if mm_input is None or is_true_on_policy_enabled(): mrope_positions_list[batch_idx] = torch.full( (3, 1), self.seq_lens_cpu[batch_idx] - 1, @@ -790,9 +788,7 @@ def _compute_mrope_positions( batch.extend_seq_lens[batch_idx], batch.extend_prefix_lens[batch_idx], ) - if mm_input is None or is_true_on_policy_enabled( - get_global_server_args() - ): + if mm_input is None or is_true_on_policy_enabled(): # text only mrope_positions = torch.tensor( [ From 2467ff360ef03abef667809c4a653b727c85f193 Mon Sep 17 00:00:00 2001 From: Yueming Yuan Date: Tue, 2 Jun 2026 12:19:57 -0700 Subject: [PATCH 40/50] Capture indexer top-k for rollout replay (DP-attention / cuda-graph safe) (#26684) --- .../srt/layers/attention/dsv4/indexer.py | 2 +- .../srt/layers/attention/nsa/nsa_indexer.py | 189 +++++++++++++----- .../srt/layers/attention/nsa_backend.py | 21 +- .../srt/managers/detokenizer_manager.py | 1 + python/sglang/srt/managers/io_struct.py | 6 + .../srt/managers/multi_tokenizer_mixin.py | 2 + .../scheduler_output_processor_mixin.py | 8 + .../sglang/srt/managers/tokenizer_manager.py | 4 + .../attention_forward_methods/forward_mha.py | 8 +- .../sglang/srt/state_capturer/indexer_topk.py | 65 +++++- 10 files changed, 241 insertions(+), 65 deletions(-) diff --git a/python/sglang/srt/layers/attention/dsv4/indexer.py b/python/sglang/srt/layers/attention/dsv4/indexer.py index 5c538631ea21..1a21be61ddfb 100644 --- a/python/sglang/srt/layers/attention/dsv4/indexer.py +++ b/python/sglang/srt/layers/attention/dsv4/indexer.py @@ -467,7 +467,7 @@ def forward_c4_indexer( compress_layer_id = token_to_kv_pool.layer_mapping[ c4_indexer.layer_id ].compress_layer_id - indexer_capturer.capture(compress_layer_id, raw_indices) + indexer_capturer.capture(compress_layer_id, raw_indices, forward_batch) class C4Indexer(nn.Module): diff --git a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py index f3c2c295d22f..193e736bbc56 100644 --- a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py +++ b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py @@ -23,6 +23,7 @@ from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype, is_fp8_fnuz from sglang.srt.layers.utils import MultiPlatformOp from sglang.srt.state_capturer.indexer_topk import ( + get_global_indexer_capturer, maybe_capture_indexer_topk, ) from sglang.srt.utils import ( @@ -431,6 +432,16 @@ def _update_rope_guarded(dst: torch.Tensor, src: torch.Tensor) -> None: return dst.copy_(src) + def _capture_and_return(self, layer_id, topk_result, raw_result, forward_batch): + # capture sequence-relative indices for rollout replay; return the + # transformed (paged/ragged) indices for the attention kernel + maybe_capture_indexer_topk( + layer_id, + raw_result if raw_result is not None else topk_result, + forward_batch, + ) + return topk_result + def _get_topk_paged( self, forward_batch: ForwardBatch, @@ -538,8 +549,15 @@ def _get_topk_paged( clean_logits=False, ) + capture = get_global_indexer_capturer() is not None # NOTE(dark): logits should be cleaned in topk_transform - topk_result = metadata.topk_transform(logits, self.index_topk) + if capture: + topk_result, raw_result = metadata.topk_transform( + logits, self.index_topk, return_raw_indices=True + ) + else: + topk_result = metadata.topk_transform(logits, self.index_topk) + raw_result = None # Restore possible padding exist in the hidden states. if not _is_hip and q_offset < q_fp8.shape[0]: pad_len = q_fp8.shape[0] - q_offset @@ -550,7 +568,9 @@ def _get_topk_paged( device=topk_result.device, ) topk_result = torch.cat([topk_result, padding], dim=0) - return topk_result + if raw_result is not None: + raw_result = torch.cat([raw_result, padding], dim=0) + return topk_result, raw_result def _should_chunk_mqa_logits( self, num_q: int, num_k: int, device: torch.device @@ -622,8 +642,10 @@ def _get_topk_ragged( topk_result = torch.full( (token_nums, self.index_topk), -1, device=device, dtype=torch.int32 ) + capture = get_global_indexer_capturer() is not None + raw_result = torch.full_like(topk_result, -1) if capture else None if batch_size == 0: - return topk_result + return topk_result, raw_result ks, ke = metadata.get_indexer_kvcache_range() @@ -674,9 +696,17 @@ def _get_topk_ragged( assert logits.shape[0] == len(seq_lens_expanded) assert logits.shape[1] == k_offset - raw_topk_result = metadata.topk_transform(logits, self.index_topk, ks=ks) - topk_result[:q_offset] = raw_topk_result - return topk_result + if capture: + transformed, raw = metadata.topk_transform( + logits, self.index_topk, ks=ks, return_raw_indices=True + ) + topk_result[:q_offset] = transformed + raw_result[:q_offset] = raw + else: + topk_result[:q_offset] = metadata.topk_transform( + logits, self.index_topk, ks=ks + ) + return topk_result, raw_result # Chunk path bytes_per_elem = 4 # float32 @@ -739,19 +769,32 @@ def _get_topk_ragged( ) batch_idx_chunk = token_to_batch_idx[start:end] - raw_topk_chunk = metadata.topk_transform( - logits_chunk, - self.index_topk, - ks=ks[start:end], - cu_seqlens_q=cu_seqlens_q_chunk, - ke_offset=lengths_chunk, - batch_idx_list=batch_idx_chunk, - topk_indices_offset_override=topk_offset_chunk, - ) - topk_result[start:end] = raw_topk_chunk + if capture: + transformed, raw = metadata.topk_transform( + logits_chunk, + self.index_topk, + ks=ks[start:end], + cu_seqlens_q=cu_seqlens_q_chunk, + ke_offset=lengths_chunk, + batch_idx_list=batch_idx_chunk, + topk_indices_offset_override=topk_offset_chunk, + return_raw_indices=True, + ) + topk_result[start:end] = transformed + raw_result[start:end] = raw + else: + topk_result[start:end] = metadata.topk_transform( + logits_chunk, + self.index_topk, + ks=ks[start:end], + cu_seqlens_q=cu_seqlens_q_chunk, + ke_offset=lengths_chunk, + batch_idx_list=batch_idx_chunk, + topk_indices_offset_override=topk_offset_chunk, + ) start = end - return topk_result + return topk_result, raw_result def _forward_cuda_k_only( self, @@ -782,7 +825,7 @@ def _forward_cuda_k_only( # MHA doesn't need topk_indices if not return_indices: - return None + return None, None # MLA: use dummy logits with topk kernel's fast path to generate indices # When length <= 2048, naive_topk_cuda directly generates [0,1,...,length-1,-1,...] @@ -793,7 +836,19 @@ def _forward_cuda_k_only( dtype=torch.float32, device=x_meta.device, ) - return metadata.topk_transform(dummy_logits, self.index_topk) + topk_result = metadata.topk_transform(dummy_logits, self.index_topk) + raw_result = None + if get_global_indexer_capturer() is not None: + # fast path selects all keys (seq_len <= index_topk); sequence-relative + # selection is [0..len-1] per token. fast_topk_v2 cannot reproduce this + # from the dummy logits (no naive_topk_cuda special case), so build it. + ar = torch.arange(self.index_topk, device=x_meta.device, dtype=torch.int32) + raw_result = torch.where( + ar.unsqueeze(0) < seq_lens_expanded.to(ar.device).unsqueeze(1), + ar.unsqueeze(0), + torch.full_like(ar, -1).unsqueeze(0), + ) + return topk_result, raw_result def _get_topk_ragged_with_cp( self, @@ -813,6 +868,7 @@ def _get_topk_ragged_with_cp( assert page_size == 64, "only support page size 64" assert len(weights.shape) == 3 weights = weights.squeeze(-1) + capture = get_global_indexer_capturer() is not None k_fp8_list = [] k_scale_list = [] ks_list = [] @@ -885,14 +941,26 @@ def _get_topk_ragged_with_cp( ke, clean_logits=False, ) - topk_result = metadata.topk_transform( - logits, - self.index_topk, - ks=ks, - cu_seqlens_q=actual_seq_q, - ke_offset=ke_offset, - batch_idx_list=batch_idx_list, - ) + if capture: + topk_result, raw_result = metadata.topk_transform( + logits, + self.index_topk, + ks=ks, + cu_seqlens_q=actual_seq_q, + ke_offset=ke_offset, + batch_idx_list=batch_idx_list, + return_raw_indices=True, + ) + else: + topk_result = metadata.topk_transform( + logits, + self.index_topk, + ks=ks, + cu_seqlens_q=actual_seq_q, + ke_offset=ke_offset, + batch_idx_list=batch_idx_list, + ) + raw_result = None else: kv_len = ( forward_batch.seq_lens_cpu[0].item() @@ -934,15 +1002,26 @@ def _get_topk_ragged_with_cp( actual_seq_q = torch.tensor([actual_seq_q], dtype=torch.int32).to( device="cuda", non_blocking=True ) - topk_result = metadata.topk_transform( - logits, - self.index_topk, - ks=ks, - cu_seqlens_q=actual_seq_q, - ke_offset=ke_offset, - ) + if capture: + topk_result, raw_result = metadata.topk_transform( + logits, + self.index_topk, + ks=ks, + cu_seqlens_q=actual_seq_q, + ke_offset=ke_offset, + return_raw_indices=True, + ) + else: + topk_result = metadata.topk_transform( + logits, + self.index_topk, + ks=ks, + cu_seqlens_q=actual_seq_q, + ke_offset=ke_offset, + ) + raw_result = None - return topk_result + return topk_result, raw_result def forward_indexer( self, @@ -1162,18 +1241,18 @@ def forward_cuda( # Optimization: fast path when skipping topk computation if skip_logits_computation and (not self.nsa_enable_prefill_cp): - return maybe_capture_indexer_topk( + topk_result, raw_result = self._forward_cuda_k_only( + x, + positions, + forward_batch, layer_id, - self._forward_cuda_k_only( - x, - positions, - forward_batch, - layer_id, - act_quant, - enable_dual_stream, - metadata, - return_indices, - ), + act_quant, + enable_dual_stream, + metadata, + return_indices, + ) + return self._capture_and_return( + layer_id, topk_result, raw_result, forward_batch ) if enable_dual_stream and forward_batch.forward_mode.is_decode_or_idle(): @@ -1263,6 +1342,7 @@ def forward_cuda( weights = self._get_logits_head_gate(x_for_gate, q_scale) + raw_result = None if _is_cuda or _is_hip: assert forward_batch.seq_lens_cpu is not None if len(forward_batch.seq_lens_cpu) == 0: @@ -1279,6 +1359,7 @@ def forward_cuda( dtype=torch.int, device=x_meta.device, ), + forward_batch, ) if ( @@ -1286,7 +1367,7 @@ def forward_cuda( or forward_batch.forward_mode.is_target_verify() or forward_batch.forward_mode.is_draft_extend(include_v2=True) ): - topk_result = self._get_topk_paged( + topk_result, raw_result = self._get_topk_paged( forward_batch, layer_id, q_fp8, weights, metadata ) else: @@ -1309,7 +1390,7 @@ def forward_cuda( weights_prev, weights_next = torch.split( weights, (weights.shape[0] + 1) // 2, dim=0 ) - topk_result_prev = self._get_topk_ragged_with_cp( + topk_result_prev, raw_prev = self._get_topk_ragged_with_cp( forward_batch, layer_id, q_fp8_prev, @@ -1319,7 +1400,7 @@ def forward_cuda( actual_seq_q_prev, ) - topk_result_next = self._get_topk_ragged_with_cp( + topk_result_next, raw_next = self._get_topk_ragged_with_cp( forward_batch, layer_id, q_fp8_next, @@ -1328,12 +1409,14 @@ def forward_cuda( kv_len_next, actual_seq_q_next, ) - return maybe_capture_indexer_topk( + return self._capture_and_return( layer_id, torch.cat([topk_result_prev, topk_result_next], dim=0), + torch.cat([raw_prev, raw_next], dim=0), + forward_batch, ) else: - topk_result = self._get_topk_ragged( + topk_result, raw_result = self._get_topk_ragged( enable_dual_stream, forward_batch, layer_id, @@ -1349,7 +1432,9 @@ def forward_cuda( topk=self.index_topk, layer_id=layer_id, ) - return maybe_capture_indexer_topk(layer_id, topk_result) + return self._capture_and_return( + layer_id, topk_result, raw_result, forward_batch + ) def forward_npu( self, diff --git a/python/sglang/srt/layers/attention/nsa_backend.py b/python/sglang/srt/layers/attention/nsa_backend.py index 55a52e9912c4..10dda1261f60 100644 --- a/python/sglang/srt/layers/attention/nsa_backend.py +++ b/python/sglang/srt/layers/attention/nsa_backend.py @@ -232,6 +232,7 @@ def topk_transform( ke_offset: torch.Tensor = None, batch_idx_list: List[int] = None, topk_indices_offset_override: Optional[torch.Tensor] = None, + return_raw_indices: bool = False, ) -> torch.Tensor: from sgl_kernel import ( fast_topk_transform_fused, @@ -262,10 +263,18 @@ def topk_transform( page_table_size_1 = self.attn_metadata.page_table_1 if not envs.SGLANG_NSA_FUSE_TOPK.get() or self.force_unfused_topk: - return fast_topk_v2(logits, seq_lens_topk, topk, row_starts=ks) - elif self.topk_transform_method == TopkTransformMethod.PAGED: - # NOTE(dark): if fused, we return a transformed page table directly - return fast_topk_transform_fused( + result = fast_topk_v2(logits, seq_lens_topk, topk, row_starts=ks) + return (result, result) if return_raw_indices else result + + # sequence-relative selection, before the kv-cache coordinate remap + raw_indices = ( + fast_topk_v2(logits, seq_lens_topk, topk, row_starts=ks) + if return_raw_indices + else None + ) + + if self.topk_transform_method == TopkTransformMethod.PAGED: + result = fast_topk_transform_fused( score=logits, lengths=seq_lens_topk, page_table_size_1=page_table_size_1, @@ -279,7 +288,7 @@ def topk_transform( "RAGGED topk_transform requires topk_indices_offset; " "expected extend-without-speculative metadata." ) - return fast_topk_transform_ragged_fused( + result = fast_topk_transform_ragged_fused( score=logits, lengths=seq_lens_topk, topk_indices_offset=cu_topk_indices_offset, @@ -289,6 +298,8 @@ def topk_transform( else: assert False, f"Unsupported {self.topk_transform_method = }" + return (result, raw_indices) if return_raw_indices else result + _NSA_IMPL_T: TypeAlias = Literal[ "flashmla_sparse", "flashmla_kv", "fa3", "tilelang", "trtllm" diff --git a/python/sglang/srt/managers/detokenizer_manager.py b/python/sglang/srt/managers/detokenizer_manager.py index a35e98167b90..8eb542406e48 100644 --- a/python/sglang/srt/managers/detokenizer_manager.py +++ b/python/sglang/srt/managers/detokenizer_manager.py @@ -384,6 +384,7 @@ def handle_batch_token_id_out(self, recv_obj: BatchTokenIDOutput): output_hidden_states=recv_obj.output_hidden_states, routed_experts=routed_experts, indexer_topk=indexer_topk, + indexer_topk_num_layers=recv_obj.indexer_topk_num_layers, customized_info=recv_obj.customized_info, placeholder_tokens_idx=None, placeholder_tokens_val=None, diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index ebd1e83ce8db..ed48ff2a8e78 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -1143,6 +1143,9 @@ class BatchTokenIDOutput(BaseBatchReq, SpeculativeDecodingMetricsMixin): # For observability time_stats: Optional[List[SchedulerReqTimeStats]] = None + # Number of indexer layers, set when indexer_topk is non-empty + indexer_topk_num_layers: Optional[int] = None + @dataclass class BatchStrOutput(BaseBatchReq, SpeculativeDecodingMetricsMixin): @@ -1209,6 +1212,9 @@ class BatchStrOutput(BaseBatchReq, SpeculativeDecodingMetricsMixin): # For observability time_stats: Optional[List[SchedulerReqTimeStats]] = None + # Number of indexer layers, set when indexer_topk is non-empty + indexer_topk_num_layers: Optional[int] = None + @dataclass class BatchEmbeddingOutput(BaseBatchReq): diff --git a/python/sglang/srt/managers/multi_tokenizer_mixin.py b/python/sglang/srt/managers/multi_tokenizer_mixin.py index baf25d332e2e..b72cd6ace087 100644 --- a/python/sglang/srt/managers/multi_tokenizer_mixin.py +++ b/python/sglang/srt/managers/multi_tokenizer_mixin.py @@ -208,6 +208,7 @@ def _handle_output_by_index(output, i): indexer_topk=_extract_field_by_index( output, "indexer_topk", i, check_length=False ), + indexer_topk_num_layers=getattr(output, "indexer_topk_num_layers", None), retraction_counts=_extract_field_by_index(output, "retraction_counts", i), placeholder_tokens_idx=None, placeholder_tokens_val=None, @@ -295,6 +296,7 @@ def _handle_output_by_index(output, i): indexer_topk=_extract_field_by_index( output, "indexer_topk", i, check_length=False ), + indexer_topk_num_layers=getattr(output, "indexer_topk_num_layers", None), customized_info=_extract_field_by_index( output, "customized_info", i, check_length=False ), diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py index d95435dec962..e8d1ebbba076 100644 --- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py +++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py @@ -1263,6 +1263,13 @@ def stream_output_generation( dp_ranks = [self.dp_rank] * len(rids) if rids else None + indexer_topk_num_layers = None + if ( + indexer_topk is not None + and (cap := get_global_indexer_capturer()) is not None + ): + indexer_topk_num_layers = cap.num_layers + # Send to detokenizer if reqs or is_idle_batch: if getattr(self.model_config, "is_multimodal_gen", False): @@ -1304,6 +1311,7 @@ def stream_output_generation( output_hidden_states=output_hidden_states, routed_experts=routed_experts, indexer_topk=indexer_topk, + indexer_topk_num_layers=indexer_topk_num_layers, customized_info=customized_info, placeholder_tokens_idx=None, placeholder_tokens_val=None, diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index e6d337de43c3..d9426a2b80e6 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -1754,6 +1754,10 @@ async def _handle_batch_output( if isinstance(val, torch.Tensor): val = pybase64.b64encode(val.numpy().tobytes()).decode("utf-8") meta_info["indexer_topk"] = val + if ( + n := getattr(recv_obj, "indexer_topk_num_layers", None) + ) is not None: + meta_info["indexer_topk_num_layers"] = n if getattr(recv_obj, "customized_info", None): for k, v in recv_obj.customized_info.items(): meta_info[k] = v[i] diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py index a2846434e846..f0fda546c356 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py @@ -145,13 +145,19 @@ def forward_normal_prepare( q = self.q_b_proj(q_lora)[0].view( -1, self.num_local_heads, self.qk_head_dim ) + from sglang.srt.state_capturer.indexer_topk import ( + get_global_indexer_capturer, + ) + + # dense MHA ignores topk_indices, but capture the prefill + # select-all selection for rollout replay when a capturer is set _ = self.indexer( x=hidden_states, q_lora=q_lora, positions=positions, forward_batch=forward_batch, layer_id=self.layer_id, - return_indices=False, + return_indices=get_global_indexer_capturer() is not None, ) elif _use_aiter_gfx95 and self.q_b_proj.weight.dtype == torch.uint8: # MXFP4: fused RMSNorm + quant diff --git a/python/sglang/srt/state_capturer/indexer_topk.py b/python/sglang/srt/state_capturer/indexer_topk.py index ca2624b99e3b..10190082cb14 100644 --- a/python/sglang/srt/state_capturer/indexer_topk.py +++ b/python/sglang/srt/state_capturer/indexer_topk.py @@ -5,7 +5,13 @@ import pybase64 import torch -from sglang.srt.layers.dp_attention import get_attention_tp_size +from sglang.srt.layers.dp_attention import ( + get_attention_dp_rank, + get_attention_tp_size, + get_dp_local_slice_cpu, + is_dp_attention_enabled, +) +from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.state_capturer.base import BaseTopkCapturer logger = logging.getLogger(__name__) @@ -28,10 +34,13 @@ def __init__( attn_tp_size = get_attention_tp_size() assert attn_tp_size == 1, "IndexerTopkCapturer now only supports DP attention" - # DP-attention capture is per-rank-local: each rank writes [:local_batch, ...] - # to its own device_cache, so the buffer only needs to fit one rank's batch. + # device_cache holds the global DP-rank-padded batch; each rank's slice is + # read back via _get_local_slice (see RoutedExpertsCapturer for the same pattern) server_args = get_global_server_args() - max_batch_size = max(server_args.chunked_prefill_size, max_running_requests) + max_batch_size = max( + server_args.chunked_prefill_size * server_args.dp_size, + max_running_requests, + ) super().__init__( num_tokens=num_tokens, @@ -42,6 +51,48 @@ def __init__( name="indexer_topk", ) + def capture(self, layer_id: int, topk_indices: torch.Tensor, forward_batch=None): + # Each DP rank only computes its own tokens' topk; write them at this + # rank's offset in the shared rank-padded buffer (writing the head would + # clobber across ranks). The offset matches _get_local_slice's read. + if forward_batch is not None and is_dp_attention_enabled(): + gnt = forward_batch.global_num_tokens_cpu + dp_rank = get_attention_dp_rank() + batch = topk_indices.shape[0] + # gnt is None during cuda-graph capture (uniform dummy batch); the + # per-rank stride is then the captured batch size itself. + start = ( + dp_rank * batch if gnt is None else sum(int(n) for n in gnt[:dp_rank]) + ) + self.device_cache.buffer[start : start + batch, layer_id, :] = topk_indices + else: + super().capture(layer_id, topk_indices) + + def _get_local_slice( + self, + forward_batch: ForwardBatch, + can_run_graph: bool, + cuda_graph_batch: Optional[int], + ) -> torch.Tensor: + # Under DP attention the device buffer is rank-padded; read this rank's + # slice (handles eager, cuda-graph-replay, and cuda-graph-capture layouts) + # instead of the head, which would only be correct for DP rank 0. + if not is_dp_attention_enabled(): + num = forward_batch.out_cache_loc.shape[0] + return self.device_cache.buffer[:num, :, : self.topk_size] + gnt = forward_batch.global_num_tokens_cpu + if gnt is None: + # cuda-graph capture: gnt not populated, uniform dummy batch. + num = forward_batch.out_cache_loc.shape[0] + start = get_attention_dp_rank() * ( + cuda_graph_batch if cuda_graph_batch is not None else num + ) + else: + start, num = get_dp_local_slice_cpu( + forward_batch, can_run_graph, cuda_graph_batch + ) + return self.device_cache.buffer[start : start + num, :, : self.topk_size] + _global_indexer_capturer: Optional[IndexerTopkCapturer] = None @@ -56,7 +107,7 @@ def set_global_indexer_capturer(capturer: Optional[IndexerTopkCapturer]): def maybe_capture_indexer_topk( - layer_id: int, topk_indices: Optional[torch.Tensor] + layer_id: int, topk_indices: Optional[torch.Tensor], forward_batch=None ) -> Optional[torch.Tensor]: """Capture topk for layer_id if a capturer is set; pass through unchanged. @@ -66,7 +117,9 @@ def maybe_capture_indexer_topk( if topk_indices is None: return None if (cap := get_global_indexer_capturer()) is not None: - cap.capture(layer_id=layer_id, topk_indices=topk_indices) + cap.capture( + layer_id=layer_id, topk_indices=topk_indices, forward_batch=forward_batch + ) return topk_indices From 505ac1420435af44604cbb3fb1fa137c0d3cea1c Mon Sep 17 00:00:00 2001 From: Jiajun Li <48857426+guapisolo@users.noreply.github.com> Date: Tue, 2 Jun 2026 22:49:51 -0700 Subject: [PATCH 41/50] chore: ignore .humanize/ artifacts (#27098) --- .gitignore | 2 ++ 1 file changed, 2 insertions(+) diff --git a/.gitignore b/.gitignore index 29de5aee37fd..ea241c7effed 100644 --- a/.gitignore +++ b/.gitignore @@ -277,3 +277,5 @@ sgl-kernel/csrc/**/*_musa/ *.npz artifacts/ .claude/scheduled_tasks.lock + +.humanize/ From 8da71333cbdb131a16b52072f6a67295ec764071 Mon Sep 17 00:00:00 2001 From: Yueming Yuan Date: Wed, 3 Jun 2026 00:48:22 -0700 Subject: [PATCH 42/50] fp8 kv cache: treat non-positive k/v_scale as uncalibrated (default 1.0) (#27131) --- python/sglang/srt/layers/quantization/kv_cache.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/python/sglang/srt/layers/quantization/kv_cache.py b/python/sglang/srt/layers/quantization/kv_cache.py index 6415733baa72..e0b9deae0146 100644 --- a/python/sglang/srt/layers/quantization/kv_cache.py +++ b/python/sglang/srt/layers/quantization/kv_cache.py @@ -56,7 +56,7 @@ def process_weights_after_loading(self, layer) -> None: if is_fp8_fnuz(): k_scale *= 2 v_scale *= 2 - elif layer.k_scale < 0.0 and layer.v_scale < 0.0: + elif layer.k_scale <= 0.0 and layer.v_scale <= 0.0: # If no scales were loaded (both scales are invalid negative # values), use the default value of 1.0 k_scale = 1.0 From 73adb23c602e31a4549976aeb1976b7793f20aee Mon Sep 17 00:00:00 2001 From: Yueming Yuan Date: Thu, 4 Jun 2026 14:47:46 -0700 Subject: [PATCH 43/50] Fix DeepSeek V4 DP reduce scatter (#27189) --- python/sglang/srt/models/deepseek_v4.py | 12 ++++++++++-- 1 file changed, 10 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 7f0ac2aa744f..87daea1e2acb 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -52,6 +52,7 @@ get_attention_dp_size, get_attention_tp_rank, get_attention_tp_size, + get_dp_global_num_tokens, get_global_dp_buffer, get_local_dp_buffer, is_dp_attention_enabled, @@ -59,7 +60,7 @@ from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear from sglang.srt.layers.logits_processor import LogitsProcessor -from sglang.srt.layers.moe import get_moe_a2a_backend +from sglang.srt.layers.moe import get_moe_a2a_backend, should_use_dp_reduce_scatterv from sglang.srt.layers.moe.fused_moe_triton import FusedMoE from sglang.srt.layers.quantization.fp8_kernel import sglang_per_token_group_quant_fp8 from sglang.srt.layers.utils import PPMissingLayer, get_layer_id @@ -864,7 +865,14 @@ def forward( get_local_dp_buffer(get_tp_group()), hidden_states, ) - dp_scatter(hidden_states, global_hidden_states, forward_batch) + if should_use_dp_reduce_scatterv(): + get_tp_group().reduce_scatterv( + global_hidden_states, + output=hidden_states, + sizes=get_dp_global_num_tokens(), + ) + else: + dp_scatter(hidden_states, global_hidden_states, forward_batch) if _use_tp_attn_a2a_scatter: assert _a2a_scatter_chunks is not None gathered = [torch.empty_like(t) for t in _a2a_scatter_chunks] From 8d51c433ace559e6fe65fdbc9df5d0c17856387d Mon Sep 17 00:00:00 2001 From: Yueming Yuan Date: Thu, 4 Jun 2026 20:25:57 -0700 Subject: [PATCH 44/50] Fix DeepSeek V4 APE weight update layout (#27306) --- .../srt/layers/attention/dsv4/compressor.py | 20 ++++++++++++------- 1 file changed, 13 insertions(+), 7 deletions(-) diff --git a/python/sglang/srt/layers/attention/dsv4/compressor.py b/python/sglang/srt/layers/attention/dsv4/compressor.py index 332f52977c9a..b3b8659163a7 100644 --- a/python/sglang/srt/layers/attention/dsv4/compressor.py +++ b/python/sglang/srt/layers/attention/dsv4/compressor.py @@ -26,7 +26,7 @@ from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_output from sglang.srt.mem_cache.deepseek_v4_compress_state import CompressStatePool from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool -from sglang.srt.utils import add_prefix +from sglang.srt.utils import add_prefix, set_weight_attrs if TYPE_CHECKING: from sglang.srt.layers.attention.deepseek_v4_backend import DeepseekV4AttnBackend @@ -309,6 +309,7 @@ def __init__( self.ape = nn.Parameter( torch.empty(self.ratio, coff * self.head_dim, dtype=torch.float32) ) + set_weight_attrs(self.ape, {"weight_loader": self.load_ape_weight}) wkv_gate_dtype = torch.bfloat16 self.wkv_gate = ReplicatedLinear( self.dim, @@ -325,14 +326,19 @@ def __init__( self.ape_converted = False - def apply_ape_hotfix(self): - assert not self.ape_converted + def load_ape_weight(self, param: torch.Tensor, loaded_weight: torch.Tensor) -> None: + assert param is self.ape + if self.overlap: + loaded_weight = torch.chunk(loaded_weight, 2, dim=-1) + loaded_weight = torch.cat([loaded_weight[0], loaded_weight[1]], dim=0) + loaded_weight = loaded_weight.view(self.ratio, -1) + assert loaded_weight.shape == param.shape + param.data.copy_(loaded_weight) self.ape_converted = True - if self.overlap: - ape = torch.chunk(self.ape.data, 2, dim=-1) - ape = torch.cat([ape[0], ape[1]], dim=0) - self.ape.data.copy_(ape.view(self.ratio, -1)) + def apply_ape_hotfix(self): + assert not self.ape_converted + self.load_ape_weight(self.ape, self.ape.data) # NOTE: used by v2 compressor backend def get_state_pool(self, forward_batch: ForwardBatch) -> CompressStatePool: From a3d47dacbc849dba727b5cfb30138f5e8609262c Mon Sep 17 00:00:00 2001 From: Yueming Yuan Date: Mon, 8 Jun 2026 15:47:23 -0700 Subject: [PATCH 45/50] Shorten DSV4 unload warning (#27603) --- python/sglang/srt/models/deepseek_v4.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 87daea1e2acb..956074bc1d3e 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -1583,7 +1583,9 @@ def auto_weight_loader(module): } if unloaded_params: logger.warning( - f"Some weights are not initialized from checkpoints: {unloaded_params}" + "Some weights are not initialized from checkpoints: " + f"count={len(unloaded_params)}. " + "Ignore this message for RL weight update." ) self.post_load_weights(is_nextn=is_nextn, weight_names=weight_names) From 914c231984f04ed15c00c7d752081fbeaf29457d Mon Sep 17 00:00:00 2001 From: Byron Hsu Date: Fri, 29 May 2026 15:08:58 -0700 Subject: [PATCH 46/50] [RL] Forward Kimi K2.5 weight hooks to language model (#26744) Co-authored-by: Byron Hsu <24364830+ByronHsu@users.noreply.github.com> (cherry picked from commit 6ea69efb7f65acfbad5dda47464b0cc8e7f01d0e) --- python/sglang/srt/models/kimi_k25.py | 18 ++++++++++++++++++ 1 file changed, 18 insertions(+) diff --git a/python/sglang/srt/models/kimi_k25.py b/python/sglang/srt/models/kimi_k25.py index 832ee74dd00e..07bf31769369 100644 --- a/python/sglang/srt/models/kimi_k25.py +++ b/python/sglang/srt/models/kimi_k25.py @@ -787,6 +787,24 @@ def stream_language_weights(): for _ in stream_language_weights(): pass + def post_load_weights(self): + if self.language_model is not None: + self.language_model.post_load_weights() + + @property + def stacked_params_mapping(self): + return getattr(self.language_model, "stacked_params_mapping", []) + + @property + def expert_params_mapping(self): + return getattr(self.language_model, "expert_params_mapping", []) + + def mutate_weight_preload(self, name): + return self.language_model.mutate_weight_preload(name) + + def custom_scale_remap(self, name): + return self.language_model.custom_scale_remap(name) + @classmethod def get_model_config_for_expert_location(cls, config: KimiK25Config): text_config = config.text_config From 3c4a663aa7096910809bd5c57a8e1873ca6e11a5 Mon Sep 17 00:00:00 2001 From: Mick Date: Sat, 23 May 2026 17:02:44 +0800 Subject: [PATCH 47/50] [VLM] feat: early-return in mm processor if the input is preprocessed (#26117) (cherry picked from commit 774b29dade07d66909094fc6fbae6a5e5fbf0074) --- python/sglang/srt/managers/mm_utils.py | 41 +++--- .../multimodal/processors/base_processor.py | 124 ++++++++++++++---- 2 files changed, 126 insertions(+), 39 deletions(-) diff --git a/python/sglang/srt/managers/mm_utils.py b/python/sglang/srt/managers/mm_utils.py index c71b0e1c4f51..cfc638c10977 100644 --- a/python/sglang/srt/managers/mm_utils.py +++ b/python/sglang/srt/managers/mm_utils.py @@ -1297,6 +1297,24 @@ def _slice_value(value, start, end): return value +def _grid_rows_to_cpu_list(value): + if isinstance(value, torch.Tensor): + value = value.detach() + if value.device.type != "cpu": + value = value.cpu() + return value.tolist() + if isinstance(value, np.ndarray): + return value.tolist() + return value + + +def _prod_grid_values(grid): + result = 1 + for value in grid: + result *= int(value) + return result + + def _slice_model_data( data: dict, index: int, @@ -1382,10 +1400,10 @@ def get_new_expanded_mm_items(original_mm_items): expanded_mm_items.append(item) continue + image_grid_rows = _grid_rows_to_cpu_list(image_grid_thw) patches_per_item = [] - for grid in image_grid_thw: - grid_tensor = torch.as_tensor(grid, dtype=torch.long) - patches_per_item.append(int(torch.prod(grid_tensor).item())) + for grid in image_grid_rows: + patches_per_item.append(_prod_grid_values(grid)) cumulative = torch.cumsum( torch.tensor(patches_per_item, dtype=torch.long), dim=0 @@ -1433,17 +1451,14 @@ def get_new_expanded_mm_items(original_mm_items): # grid_len = num_videos, num_items = sum(T for each video) = total frames grid_len = _get_length(video_grid_thw) num_videos = grid_len + video_grid_rows = _grid_rows_to_cpu_list(video_grid_thw) # Calculate total frames and frames per video frames_per_video = [] total_frames = 0 for i in range(num_videos): - grid = video_grid_thw[i] - if isinstance(grid, torch.Tensor): - T = int(grid[0].item()) # T is the first element [T, H, W] - else: - grid_tensor = torch.as_tensor(grid, dtype=torch.long) - T = int(grid_tensor[0].item()) + grid = video_grid_rows[i] + T = int(grid[0]) # T is the first element [T, H, W] frames_per_video.append(T) total_frames += T @@ -1455,12 +1470,8 @@ def get_new_expanded_mm_items(original_mm_items): # Calculate patches per video: T * H * W for each video patches_per_video = [] for i in range(num_videos): - grid = video_grid_thw[i] - if isinstance(grid, torch.Tensor): - patches_per_video.append(int(torch.prod(grid).item())) - else: - grid_tensor = torch.as_tensor(grid, dtype=torch.long) - patches_per_video.append(int(torch.prod(grid_tensor).item())) + grid = video_grid_rows[i] + patches_per_video.append(_prod_grid_values(grid)) # Calculate cumulative patches to get slice indices for each video cumulative = torch.cumsum( diff --git a/python/sglang/srt/multimodal/processors/base_processor.py b/python/sglang/srt/multimodal/processors/base_processor.py index 5f207f9f2bbf..e8a8d6c44a80 100644 --- a/python/sglang/srt/multimodal/processors/base_processor.py +++ b/python/sglang/srt/multimodal/processors/base_processor.py @@ -49,6 +49,10 @@ class BaseMultiModalProcessorOutput: # input_text with all multimodality placeholder token expanded input_text: str + # original pre-tokenized ids, useful for processor_output/precomputed inputs, + # when they already carry the input ids + input_ids: Optional[Union[List[int], torch.Tensor]] = None + # frames loaded from image, in given order images: Optional[list[Union[Image.Image, dict]]] = dataclasses.field( default_factory=list @@ -516,15 +520,8 @@ def _load_single_item( Class method that can be pickled for multiprocessing """ - if isinstance(data, dict): - data_format = data.get("format") - if data_format in ( - MultimodalInputFormat.PROCESSOR_OUTPUT.name, - MultimodalInputFormat.PRECOMPUTED_EMBEDDING.name, - "processor_output", - "precomputed_embedding", - ): - return data + if cls._is_preprocessed_input(data): + return data try: if modality == Modality.IMAGE: img, _ = load_image(data, cls.gpu_image_decode) @@ -544,6 +541,49 @@ def _load_single_item( except Exception as e: raise RuntimeError(f"Error while loading data {data}: {e}") + @staticmethod + def _get_preprocessed_input_format(data): + """returns the detailed format if the provided data is already preprocessed. + returns none if the provided data is not preprocessed + """ + if not isinstance(data, dict): + return None + data_format = data.get("format") + if isinstance(data_format, MultimodalInputFormat): + return data_format + if data_format in ( + MultimodalInputFormat.PROCESSOR_OUTPUT.name, + "processor_output", + ): + return MultimodalInputFormat.PROCESSOR_OUTPUT + if data_format in ( + MultimodalInputFormat.PRECOMPUTED_EMBEDDING.name, + "precomputed_embedding", + ): + return MultimodalInputFormat.PRECOMPUTED_EMBEDDING + return None + + @classmethod + def _is_preprocessed_input(cls, data): + """returns if the data is already preprocessed (by the vlm processor)""" + return cls._get_preprocessed_input_format(data) is not None + + @classmethod + def _all_mm_data_is_preprocessed(cls, *data_lists): + has_mm_data = False + for data_list in data_lists: + if not data_list: + continue + if not isinstance(data_list, list): + data_list = [data_list] + for item in data_list: + if item is None: + continue + has_mm_data = True + if not cls._is_preprocessed_input(item): + return False + return has_mm_data + def _submit_mm_data_loading_tasks_simple( self, data_list: Optional[list], @@ -668,10 +708,8 @@ def _validate_one_modality(modality: Modality, data_list: Optional[list]): formatted_indices = [] for idx, item in enumerate(data_list): - if isinstance(item, dict): - fmt = item.get("format") - if fmt in {"processor_output", "precomputed_embedding"}: - formatted_indices.append(idx) + if BaseMultimodalProcessor._is_preprocessed_input(item): + formatted_indices.append(idx) if formatted_indices: if len(data_list) != 1: @@ -706,12 +744,7 @@ def validate_mm_data( def _process_loaded_mm_data(self, modality, raw_data, result): images, videos, audios = [], [], [] - is_precomputed = isinstance(raw_data, dict) and raw_data.get("format") in [ - MultimodalInputFormat.PROCESSOR_OUTPUT.name, - MultimodalInputFormat.PRECOMPUTED_EMBEDDING.name, - "processor_output", - "precomputed_embedding", - ] + is_precomputed = self._is_preprocessed_input(raw_data) if modality == Modality.IMAGE: if is_precomputed: @@ -742,6 +775,19 @@ def load_mm_data( BaseMultimodalProcessor.validate_mm_data(image_data, video_data, audio_data) + input_ids = prompt if isinstance(prompt, list) else None + if input_ids is not None and self._all_mm_data_is_preprocessed( + image_data, video_data, audio_data + ): + # fast path for preprocessed data: early return + return BaseMultiModalProcessorOutput( + input_text="", + input_ids=input_ids, + images=list(image_data or []), + videos=list(video_data or []), + audios=list(audio_data or []), + ) + multimodal_tokens_pattern = multimodal_tokens.get_combined_regex() if isinstance(prompt, list) and return_text: assert len(prompt) and isinstance(prompt[0], int) @@ -780,6 +826,7 @@ def load_mm_data( return_text=return_text, discard_alpha_channel=discard_alpha_channel, audio_sample_rate=audio_sample_rate, + input_ids=input_ids, ) # For models other than MiniCPMO and MiniCPMV, # totally align multimodal_tokens, fast path @@ -792,6 +839,7 @@ def load_mm_data( return_text=return_text, discard_alpha_channel=discard_alpha_channel, audio_sample_rate=audio_sample_rate, + input_ids=input_ids, ) def fast_load_mm_data( @@ -804,6 +852,7 @@ def fast_load_mm_data( return_text: Optional[bool] = True, discard_alpha_channel: bool = True, audio_sample_rate: Optional[int] = None, + input_ids: Optional[Union[List[int], torch.Tensor]] = None, ) -> BaseMultiModalProcessorOutput: """ A fast version of `load_mm_data` that loads multimodal data directly. @@ -876,6 +925,7 @@ def fast_load_mm_data( audios=audios, videos=videos, input_text=prompt_str, + input_ids=input_ids, ) def legacy_load_mm_data( @@ -888,6 +938,7 @@ def legacy_load_mm_data( return_text: Optional[bool] = True, discard_alpha_channel: bool = True, audio_sample_rate: Optional[int] = None, + input_ids: Optional[Union[List[int], torch.Tensor]] = None, ) -> BaseMultiModalProcessorOutput: """ Each frame of video/image will be replaced by a single image token @@ -995,6 +1046,7 @@ def legacy_load_mm_data( audios=audios, videos=videos, input_text="".join(new_text_parts), + input_ids=input_ids, ) @staticmethod @@ -1081,6 +1133,15 @@ def _process_and_collect_mm_items( return collected_items, input_ids, ret + @staticmethod + def _ensure_input_ids_is_tensor(input_ids) -> Optional[torch.Tensor]: + """make sure the input_ids is a flattened tensor""" + if input_ids is None: + return None + if isinstance(input_ids, torch.Tensor): + return input_ids.flatten().to(dtype=torch.long) + return torch.tensor(input_ids, dtype=torch.long).flatten() + def process_and_combine_mm_data( self, base_output: BaseMultiModalProcessorOutput, @@ -1134,16 +1195,19 @@ def process_and_combine_mm_data( ret = None # Handle dict items (processed or precomputed) + dict_ret = None for modality, dict_item in dict_items: - input_format = dict_item.get("format", None) - if input_format == "processor_output": + input_format = self._get_preprocessed_input_format(dict_item) + if input_format is not None and dict_ret is None: + dict_ret = dict_item + if input_format == MultimodalInputFormat.PROCESSOR_OUTPUT: items = self.collect_mm_items_from_processor_output(dict_item) for item in items: item.format = MultimodalInputFormat.PROCESSOR_OUTPUT all_collected_items.extend(items) - elif input_format == "precomputed_embedding": - feature = dict_item["feature"] - del dict_item["feature"] + elif input_format == MultimodalInputFormat.PRECOMPUTED_EMBEDDING: + dict_item = dict(dict_item) + feature = dict_item.pop("feature") all_collected_items.append( MultimodalDataItem( modality=modality, @@ -1153,6 +1217,18 @@ def process_and_combine_mm_data( ) ) # Fallback tokenization if no raw items were processed + if ret is None and dict_ret is not None: + ret = dict_ret + + if input_ids is None: + input_ids = self._ensure_input_ids_is_tensor(base_output.input_ids) + + if input_ids is None: + for _, dict_item in dict_items: + input_ids = self._ensure_input_ids_is_tensor(dict_item.get("input_ids")) + if input_ids is not None: + break + if input_ids is None: input_ids = self._tokenizer( base_output.input_text, From 7f52b33dc94d762a547f69626746c0d777501bd0 Mon Sep 17 00:00:00 2001 From: Byron Hsu Date: Mon, 1 Jun 2026 09:58:14 -0700 Subject: [PATCH 48/50] [RL+VLM] Avoid retokenization drift for pre-tokenized (token-id) VLM requests (#26555) Co-authored-by: Byron Hsu Co-authored-by: root Co-authored-by: Cursor Co-authored-by: Mick (cherry picked from commit f6a5a1b59c2875d696b84edb1ee6c42aca065e15) --- python/sglang/srt/environ.py | 3 + .../multimodal/processors/base_processor.py | 120 ++++++++++++++++++ .../srt/multimodal/processors/kimi_common.py | 15 +++ .../vlm/test_token_id_retokenize_e2e.py | 115 +++++++++++++++++ 4 files changed, 253 insertions(+) create mode 100644 test/registered/vlm/test_token_id_retokenize_e2e.py diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index c78ad659826d..cb33e8431a73 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -483,6 +483,9 @@ class Envs: SGLANG_MM_PRECOMPUTE_HASH = EnvBool(False) SGLANG_VIT_ENABLE_CUDA_GRAPH = EnvBool(False) SGLANG_MM_SKIP_COMPUTE_HASH = EnvBool(False) + # For pre-tokenized (list[int]) multimodal prompts, + # preserve the user's original tokens to avoid retokenization drift. + SGLANG_MM_AVOID_RETOKENIZE = EnvBool(True) # VLM Item CUDA IPC Transport diff --git a/python/sglang/srt/multimodal/processors/base_processor.py b/python/sglang/srt/multimodal/processors/base_processor.py index e8a8d6c44a80..a52d7ccb216c 100644 --- a/python/sglang/srt/multimodal/processors/base_processor.py +++ b/python/sglang/srt/multimodal/processors/base_processor.py @@ -1142,6 +1142,84 @@ def _ensure_input_ids_is_tensor(input_ids) -> Optional[torch.Tensor]: return input_ids.flatten().to(dtype=torch.long) return torch.tensor(input_ids, dtype=torch.long).flatten() + def _wrap_tensor_for_cuda_ipc(self, tensor: torch.Tensor): + """helper function to turn a tensor into a cuda-ipc tensor""" + if not tensor.is_cuda: + return tensor + + sync_flag, available_slice, byte_offset = ( + self.cudaipc_mmfeature_pool.return_a_slice_tensor_with_flag(tensor) + ) + if isinstance(available_slice, torch.Tensor): + available_slice.copy_(tensor.view(torch.int8).view(-1), non_blocking=True) + return CudaIpcTensorTransportProxy( + data=available_slice, + info_data=tensor, + sync_buffer_meta=sync_flag, + pool_ipc_handle=( + self.cudaipc_mmfeature_pool._pool_ipc_handle + if _IPC_POOL_HANDLE_CACHE + else None + ), + pool_byte_offset=byte_offset, + pool_device_index=self.cudaipc_mmfeature_pool._pool_device_index, + ) + if self.server_args.keep_mm_feature_on_device: + return tensor + return tensor.cpu() + + def resolve_image_token_counts(self, images: List) -> List[int]: + """Per-image expanded token counts, computed without re-tokenizing. + + Default implementation uses the transformers in-tree convention + ``_get_num_multimodal_tokens(image_sizes=...)`` (present on the in-tree + VLM processors, e.g. Qwen-VL, Gemma3, GLM4V). Models whose processor + does not implement it (e.g. Kimi) override this method. + + """ + assert images is not None + image_sizes = [(image.height, image.width) for image in images] + num_image_tokens = self._processor._get_num_multimodal_tokens( + image_sizes=image_sizes + ).num_image_tokens + return [int(count) for count in num_image_tokens] + + @staticmethod + def _expand_input_ids( + original_ids: List[int], + counts: List[int], + placeholder_token_id: Optional[int], + ) -> List[int]: + """Rebuild final input_ids for a pre-tokenized (list[int]) prompt. + + Keep the user's ORIGINAL tokens verbatim and expand the i-th image + placeholder into ``counts[i]`` copies of ``placeholder_token_id``. The HF + processor's re-tokenization is discarded, so non-media tokens cannot + drift. + + """ + if placeholder_token_id is None: + raise ValueError("placeholder_token_id is not set for this processor") + + num_placeholders = sum( + 1 for token_id in original_ids if token_id == placeholder_token_id + ) + if num_placeholders != len(counts): + raise ValueError( + f"prompt has {num_placeholders} image placeholder token(s) but " + f"{len(counts)} image(s) were provided" + ) + + rebuilt: List[int] = [] + next_image_idx = 0 + for token_id in original_ids: + if token_id == placeholder_token_id: + rebuilt.extend([placeholder_token_id] * counts[next_image_idx]) + next_image_idx += 1 + else: + rebuilt.append(token_id) + return rebuilt + def process_and_combine_mm_data( self, base_output: BaseMultiModalProcessorOutput, @@ -1191,6 +1269,48 @@ def process_and_combine_mm_data( **kwargs, ) all_collected_items = collected_items + + # When SGLANG_MM_AVOID_RETOKENIZE is on, keep the user's exact tokens to avoid retokenize drift. + # Drift happens when Retokenization is not identity: Decode(X) => String => Re-tokenize => Y, X != Y. + if ( + envs.SGLANG_MM_AVOID_RETOKENIZE.get() + and base_output.input_ids is not None + and input_ids is not None + and raw_images + and not raw_audios + and not raw_videos + ): + assert isinstance( + base_output.input_ids, list + ), f"expected list[int] input_ids, got {type(base_output.input_ids)}" + try: + counts = self.resolve_image_token_counts(raw_images) + image_placeholder_token_id = mm_tokens.image_token_id + if image_placeholder_token_id is None: + raise ValueError( + "image placeholder token id is not set for this processor" + ) + processor_placeholder_count = int( + (input_ids == image_placeholder_token_id).sum().item() + ) + if processor_placeholder_count != sum(counts): + raise ValueError( + "processor image placeholder count mismatch: " + f"processor={processor_placeholder_count}, " + f"resolved={sum(counts)}" + ) + input_ids = torch.tensor( + self._expand_input_ids( + base_output.input_ids, + counts, + image_placeholder_token_id, + ), + dtype=input_ids.dtype, + ) + except Exception as e: + logger.warning( + f"Due to {e}, falling back to decode+retokenize, which may change prompt length (token drift)." + ) else: ret = None diff --git a/python/sglang/srt/multimodal/processors/kimi_common.py b/python/sglang/srt/multimodal/processors/kimi_common.py index c2046d32c4ce..5371fee71cb3 100644 --- a/python/sglang/srt/multimodal/processors/kimi_common.py +++ b/python/sglang/srt/multimodal/processors/kimi_common.py @@ -23,6 +23,21 @@ class KimiGridMMDataMixin: - self._tokenizer (with .encode()) """ + def resolve_image_token_counts(self, images): + """Kimi's processor is remote-code and does not implement the + transformers ``_get_num_multimodal_tokens`` convention; use its + ``media_tokens_calculator`` instead. + + """ + assert images is not None + media_tokens_calculator = ( + self._processor.media_processor.media_tokens_calculator + ) + return [ + int(media_tokens_calculator({"type": "image", "image": image})) + for image in images + ] + def _num_image_tokens_from_grid( self, grid_thw: Union[torch.Tensor, np.ndarray, list, tuple] ) -> int: diff --git a/test/registered/vlm/test_token_id_retokenize_e2e.py b/test/registered/vlm/test_token_id_retokenize_e2e.py new file mode 100644 index 000000000000..8a7192bd7ca4 --- /dev/null +++ b/test/registered/vlm/test_token_id_retokenize_e2e.py @@ -0,0 +1,115 @@ +"""E2E test for SGLANG_MM_AVOID_RETOKENIZE on the pre-tokenized VLM path. + +A client may send a multimodal request as input_ids (list[int]) instead of text. +On that path the server decodes the ids back to text and the HF processor +re-tokenizes them. If the original ids were non-canonical (decode -> re-encode is +not identity), that re-tokenization drifts: the reported prompt_tokens changes. + +With SGLANG_MM_AVOID_RETOKENIZE ON (default), the server keeps the user's +original tokens verbatim and only expands the image placeholder, so prompt_tokens +stays faithful to what the client sent. + +For each model we launch a real server twice with the same predefined, +non-canonical prompt ("Describe" split into "D"+"escribe") plus one image: + + * flag OFF -> the prompt re-tokenizes (drift): prompt_tokens shrinks by the + drift delta. + * flag ON -> no drift: prompt_tokens equals the original length (with the + image placeholder expanded). +""" + +import base64 +import io +import unittest + +import requests +from PIL import Image +from transformers import AutoProcessor + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import ( + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + CustomTestCase, + popen_launch_server, +) + +register_cuda_ci(est_time=300, suite="base-b-test-1-gpu-large") + + +def _data_uri(): + img = Image.new("RGB", (64, 64), (128, 128, 128)) + buf = io.BytesIO() + img.save(buf, format="PNG") + return "data:image/png;base64," + base64.b64encode(buf.getvalue()).decode() + + +def _build_drift_prompt(model, image_token): + """Return (input_ids, drift_delta). + + input_ids is a predefined non-canonical prompt: "Describe" is split into + "D"+"escribe" (decodes to the same text but re-encodes to the single merged + token), followed by one image placeholder. drift_delta is how many extra + tokens the non-canonical form carries vs. the canonical re-tokenization. + """ + tok = AutoProcessor.from_pretrained( + model, trust_remote_code=True, use_fast=True + ).tokenizer + + def enc(text): + return tok.encode(text, add_special_tokens=False) + + input_ids = enc("D") + enc("escribe") + enc(" the picture: ") + enc(image_token) + canonical = enc(tok.decode(input_ids)) + drift_delta = len(input_ids) - len(canonical) + return input_ids, drift_delta + + +def _prompt_tokens(base_url, input_ids, image): + resp = requests.post( + base_url + "/generate", + json={ + "input_ids": input_ids, + "image_data": [image], + "sampling_params": {"temperature": 0.0, "max_new_tokens": 1}, + }, + timeout=300, + ) + resp.raise_for_status() + return resp.json()["meta_info"]["prompt_tokens"] + + +class TestQwenVLTokenIdRetokenize(CustomTestCase): + model = "Qwen/Qwen2.5-VL-3B-Instruct" + image_token = "<|vision_start|><|image_pad|><|vision_end|>" + other_args = ["--trust-remote-code", "--mem-fraction-static", "0.7"] + + def test_flag_off_drifts_flag_on_does_not(self): + input_ids, drift_delta = _build_drift_prompt(self.model, self.image_token) + self.assertGreater(drift_delta, 0, "prompt is canonical; no drift to exercise") + image = _data_uri() + + prompt_tokens = {} + for flag in ("0", "1"): + process = popen_launch_server( + self.model, + DEFAULT_URL_FOR_TEST, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=self.other_args, + env={"SGLANG_MM_AVOID_RETOKENIZE": flag}, + ) + try: + prompt_tokens[flag] = _prompt_tokens( + DEFAULT_URL_FOR_TEST, input_ids, image + ) + finally: + kill_process_tree(process.pid) + + # ON keeps the user's original tokens; OFF loses the drift_delta tokens. + pt_off, pt_on = prompt_tokens["0"], prompt_tokens["1"] + self.assertEqual(pt_on - pt_off, drift_delta, f"on={pt_on}, off={pt_off}") + + +if __name__ == "__main__": + unittest.main() From f017fbb62303cd6f3e33bccb3560e4916463fa61 Mon Sep 17 00:00:00 2001 From: Byron Hsu Date: Fri, 29 May 2026 14:31:05 -0700 Subject: [PATCH 49/50] [Model] Add Qwen3-MoE MTP (#26468) Co-authored-by: Byron Hsu Co-authored-by: Cursor Co-authored-by: root (cherry picked from commit cf66693b3530430b8d23aa6eefdfd922c8bfcc1f) --- python/sglang/srt/configs/model_config.py | 4 + python/sglang/srt/models/qwen3_moe.py | 19 ++- python/sglang/srt/models/qwen3_moe_mtp.py | 131 ++++++++++++++++++ python/sglang/srt/speculative/eagle_worker.py | 6 +- .../sglang/srt/speculative/eagle_worker_v2.py | 7 + 5 files changed, 163 insertions(+), 4 deletions(-) create mode 100644 python/sglang/srt/models/qwen3_moe_mtp.py diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index 02afafa0a664..5a7e8d72d007 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -467,6 +467,10 @@ def _config_draft_model(self): self.hf_config.architectures[0] = "Qwen3NextForCausalLMMTP" self.hf_config.num_nextn_predict_layers = 1 + if is_draft_model and self.hf_config.architectures[0] == "Qwen3MoeForCausalLM": + self.hf_config.architectures[0] = "Qwen3MoeForCausalLMMTP" + self.hf_config.num_nextn_predict_layers = 1 + if is_draft_model and self.hf_config.architectures[0] in [ "Qwen3_5ForConditionalGeneration", "Qwen3_5MoeForConditionalGeneration", diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py index fe67985371d0..25e82759ee5f 100644 --- a/python/sglang/srt/models/qwen3_moe.py +++ b/python/sglang/srt/models/qwen3_moe.py @@ -1147,7 +1147,9 @@ def set_dflash_layers_to_capture(self, layer_ids: List[int]): self.capture_aux_hidden_states = True self.model.set_dflash_layers_to_capture([val + 1 for val in layer_ids]) - def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): + def load_weights( + self, weights: Iterable[Tuple[str, torch.Tensor]], is_mtp: bool = False + ): stacked_params_mapping = self.stacked_params_mapping expert_params_mapping = self.expert_params_mapping @@ -1155,6 +1157,21 @@ def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): params_dict = dict(self.named_parameters()) for name, loaded_weight in weights: + if is_mtp: + if "mtp" not in name: + continue + + if name in [ + "mtp.fc.weight", + "mtp.pre_fc_norm_embedding.weight", + "mtp.pre_fc_norm_hidden.weight", + ]: + name = name.replace("mtp.", "") + else: + name = name.replace("mtp", "model") + elif "mtp" in name: + continue + layer_id = get_layer_id(name) if ( layer_id is not None diff --git a/python/sglang/srt/models/qwen3_moe_mtp.py b/python/sglang/srt/models/qwen3_moe_mtp.py new file mode 100644 index 000000000000..e6f825eef630 --- /dev/null +++ b/python/sglang/srt/models/qwen3_moe_mtp.py @@ -0,0 +1,131 @@ +# Copyright 2023-2024 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +"""Inference-only Qwen3-MoE MTP speculative decoding.""" + +import logging +from typing import Iterable, Optional, Tuple + +import torch +from torch import nn +from transformers import PretrainedConfig + +from sglang.srt.distributed import get_pp_group, get_tensor_model_parallel_world_size +from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder +from sglang.srt.layers.layernorm import RMSNorm +from sglang.srt.layers.logits_processor import LogitsProcessor +from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE +from sglang.srt.layers.quantization.base_config import QuantizationConfig +from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead +from sglang.srt.model_executor.forward_batch_info import ForwardBatch +from sglang.srt.models.qwen3_moe import Qwen3MoeForCausalLM, Qwen3MoeModel +from sglang.srt.server_args import get_global_server_args +from sglang.srt.utils import add_prefix + +logger = logging.getLogger(__name__) + + +class Qwen3MoeForCausalLMMTP(Qwen3MoeForCausalLM): + def __init__( + self, + config: PretrainedConfig, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + ) -> None: + nn.Module.__init__(self) + self.config = config + config.num_hidden_layers = 1 + self.tp_size = get_tensor_model_parallel_world_size() + self.quant_config = quant_config + self.pp_group = get_pp_group() + + self.fc = nn.Linear(2 * config.hidden_size, config.hidden_size, bias=False) + self.pre_fc_norm_embedding = RMSNorm( + config.hidden_size, eps=config.rms_norm_eps + ) + self.pre_fc_norm_hidden = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.model = Qwen3MoeModel( + config, quant_config, prefix=add_prefix("model", prefix) + ) + self.lm_head = ParallelLMHead( + config.vocab_size, + config.hidden_size, + quant_config=quant_config, + prefix=add_prefix("lm_head", prefix), + use_attn_tp_group=get_global_server_args().enable_dp_lm_head, + ) + self.logits_processor = LogitsProcessor(config) + + # Required by Qwen3MoeForCausalLM.load_weights(), which we reuse below. + self.stacked_params_mapping = [ + ("qkv_proj", "q_proj", "q"), + ("qkv_proj", "k_proj", "k"), + ("qkv_proj", "v_proj", "v"), + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), + ] + self.expert_params_mapping = FusedMoE.make_expert_params_mapping( + ckpt_gate_proj_name="gate_proj", + ckpt_down_proj_name="down_proj", + ckpt_up_proj_name="up_proj", + num_experts=self.config.num_experts, + ) + self.capture_aux_hidden_states = False + + def set_embed_and_head(self, embed, head): + del self.model.embed_tokens.weight + del self.lm_head.weight + self.model.embed_tokens.weight = embed + self.lm_head.weight = head + torch.cuda.empty_cache() + torch.cuda.synchronize() + + @torch.no_grad() + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + forward_batch: ForwardBatch, + input_embeds: Optional[torch.Tensor] = None, + **kwargs, + ): + if input_embeds is None: + input_embeds = self.model.embed_tokens(input_ids) + + hidden_states = forward_batch.spec_info.hidden_states + + 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) + hidden_states = self.fc(torch.cat((input_embeds, hidden_states), dim=-1)) + + with get_global_expert_distribution_recorder().disable_this_region(): + hidden_states = self.model( + input_ids, + positions, + forward_batch, + hidden_states, + ) + + return self.logits_processor( + input_ids, hidden_states, self.lm_head, forward_batch + ) + + def load_weights( + self, weights: Iterable[Tuple[str, torch.Tensor]], is_mtp: bool = False + ): + return super().load_weights(weights, is_mtp=True) + + +EntryClass = [Qwen3MoeForCausalLMMTP] diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py index a81e20add0ca..a8b9d36e1349 100644 --- a/python/sglang/srt/speculative/eagle_worker.py +++ b/python/sglang/srt/speculative/eagle_worker.py @@ -895,12 +895,12 @@ def draft_forward(self, forward_batch: ForwardBatch): # Set inputs forward_batch.input_ids = input_ids - # This is a temporary fix for the case that the user is using standalone - # speculative decoding and the draft model architecture is gpt-oss. gpt-oss - # rope kernel needs cache_loc to be contiguous. + # Some draft model RoPE kernels need cache_loc to be contiguous. if ( self.server_args.speculative_algorithm == "STANDALONE" and self.model_config.hf_config.architectures[0] == "GptOssForCausalLM" + ) or self.model_config.hf_config.architectures[0] == ( + "Qwen3MoeForCausalLMMTP" ): out_cache_loc = out_cache_loc.contiguous() forward_batch.out_cache_loc = out_cache_loc[i] diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index e845fde3f2ce..1d7ebdffad09 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -458,6 +458,13 @@ def draft_forward(self, forward_batch: ForwardBatch): # Set inputs forward_batch.input_ids = input_ids + # Qwen3-MoE MTP uses a fused RoPE + KV-store path whose cache_loc + # argument must be contiguous. + if ( + self.draft_runner.model_config.hf_config.architectures[0] + == "Qwen3MoeForCausalLMMTP" + ): + out_cache_loc = out_cache_loc.contiguous() forward_batch.out_cache_loc = out_cache_loc[i] forward_batch.attn_backend = self.draft_attn_backend.attn_backends[i] spec_info.hidden_states = hidden_states From 8fef8b14cc41d4afa1ab51fa10a9df469d885206 Mon Sep 17 00:00:00 2001 From: Byron Hsu Date: Fri, 29 May 2026 20:46:10 -0700 Subject: [PATCH 50/50] [RL] Fix crash when the reqs in a batch have a mix of `return_routed_experts` = True and False. (#26423) Co-authored-by: root Co-authored-by: Cursor Co-authored-by: Cheng Wan <54331508+ch-wan@users.noreply.github.com> (cherry picked from commit 6f1c9fc77b0215732cdbcaa9761a6327f8ac1ef7) --- .../scheduler_output_processor_mixin.py | 18 ++- .../sglang/srt/managers/tokenizer_manager.py | 4 +- .../rl/test_return_routed_experts.py | 114 +++++++++++++++--- 3 files changed, 112 insertions(+), 24 deletions(-) diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py index e8d1ebbba076..f82e79dc3376 100644 --- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py +++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py @@ -1233,18 +1233,24 @@ def stream_output_generation( output_token_ids_logprobs_val.append([]) output_token_ids_logprobs_idx.append([]) - if req.return_hidden_states: + if batch.return_hidden_states: if output_hidden_states is None: output_hidden_states = [] - output_hidden_states.append(req.hidden_states) - if req.return_routed_experts: + output_hidden_states.append( + req.hidden_states if req.return_hidden_states else None + ) + if batch.return_routed_experts: if routed_experts is None: routed_experts = [] - routed_experts.append(req.routed_experts) - if req.return_indexer_topk: + routed_experts.append( + req.routed_experts if req.return_routed_experts else None + ) + if batch.return_indexer_topk: if indexer_topk is None: indexer_topk = [] - indexer_topk.append(req.indexer_topk) + indexer_topk.append( + req.indexer_topk if req.return_indexer_topk else None + ) if req.customized_info is not None: for k, v in req.customized_info.items(): diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index d9426a2b80e6..2f2a85d56cce 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -1739,7 +1739,9 @@ async def _handle_batch_output( ] if getattr(recv_obj, "output_hidden_states", None): - meta_info["hidden_states"] = recv_obj.output_hidden_states[i] + hidden_states = recv_obj.output_hidden_states[i] + if hidden_states is not None: + meta_info["hidden_states"] = hidden_states if getattr(recv_obj, "routed_experts", None): val = recv_obj.routed_experts[i] if val is not None: diff --git a/test/registered/rl/test_return_routed_experts.py b/test/registered/rl/test_return_routed_experts.py index 5157616efa27..90007074d30e 100644 --- a/test/registered/rl/test_return_routed_experts.py +++ b/test/registered/rl/test_return_routed_experts.py @@ -77,6 +77,32 @@ def setUpClass(cls): ] cls.reference_args = common cls.sampling_args = {"temperature": 0} + cls.texts = None + cls.baseline_results = None + cls.reference_results = None + cls._endpoints = [ + ( + "/generate", + cls._build_generate_payload, + extract_routed_experts_from_meta_info, + ), + ( + "/v1/chat/completions", + cls._build_chat_payload, + extract_routed_experts_from_openai_response, + ), + ( + "/v1/completions", + cls._build_completion_payload, + extract_routed_experts_from_openai_response, + ), + ] + + @classmethod + def _ensure_comparison_results(cls): + if cls.baseline_results is not None and cls.reference_results is not None: + return + # prepare ShareGPT dataset dataset_path = download_and_cache_hf_file(SHAREGPT_REPO_ID, SHAREGPT_FILENAME) with open(dataset_path) as f: @@ -96,23 +122,6 @@ def setUpClass(cls): if not cls.texts: raise ValueError("No valid texts found in the dataset") cls.texts = cls.texts[:100] - cls._endpoints = [ - ( - "/generate", - cls._build_generate_payload, - extract_routed_experts_from_meta_info, - ), - ( - "/v1/chat/completions", - cls._build_chat_payload, - extract_routed_experts_from_openai_response, - ), - ( - "/v1/completions", - cls._build_completion_payload, - extract_routed_experts_from_openai_response, - ), - ] cls.baseline_results = cls._collect_results(cls.baseline_args) cls.reference_results = cls._collect_results(cls.reference_args) @@ -128,8 +137,13 @@ def test_return_routed_experts_chat_completions(cls): def test_return_routed_experts_completions(cls): cls._run_endpoint_test("/v1/completions") + def test_mixed_return_routed_experts_batch_alignment(self): + self._run_mixed_batch_alignment_case([]) + self._run_mixed_batch_alignment_case(["--tokenizer-worker-num", 2]) + @classmethod def _run_endpoint_test(cls, endpoint): + cls._ensure_comparison_results() captured_baseline_experts = cls.baseline_results[endpoint] captured_reference_experts = cls.reference_results[endpoint] @@ -171,6 +185,72 @@ def _collect_results( finally: kill_process_tree(process.pid) + @classmethod + def _run_mixed_batch_alignment_case(cls, other_args): + process = popen_launch_server( + DEFAULT_ENABLE_ROUTED_EXPERTS_MODEL_NAME_FOR_TEST, + DEFAULT_URL_FOR_TEST, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=[ + "--tp", + 2, + "--enable-return-routed-experts", + "--disable-cuda-graph", + "--disable-piecewise-cuda-graph", + *other_args, + ], + ) + try: + responses = asyncio.run(cls._send_mixed_batch()) + cls._assert_mixed_batch_result(responses) + finally: + kill_process_tree(process.pid) + + @classmethod + async def _send_mixed_batch(cls): + payload_no_rr = { + "text": "The quick brown fox jumps over the lazy dog.", + "sampling_params": { + "temperature": 0, + "max_new_tokens": 16, + "ignore_eos": True, + }, + "return_routed_experts": False, + } + payload_with_rr = { + "text": "The quick brown fox jumps over the lazy dog.", + "sampling_params": { + "temperature": 0, + "max_new_tokens": 16, + "ignore_eos": True, + }, + "return_routed_experts": True, + } + + async with aiohttp.ClientSession() as session: + return await asyncio.gather( + cls._post_generate(session, payload_no_rr), + cls._post_generate(session, payload_with_rr), + ) + + @staticmethod + async def _post_generate(session, payload): + async with session.post( + f"{DEFAULT_URL_FOR_TEST}/generate", json=payload + ) as response: + body = await response.json() + if response.status != 200: + raise AssertionError(f"HTTP {response.status}: {body}") + if "error" in body: + raise AssertionError(f"generate returned error: {body['error']}") + return body + + @classmethod + def _assert_mixed_batch_result(cls, responses): + no_rr, with_rr = responses + assert "routed_experts" not in no_rr.get("meta_info", {}) + assert "routed_experts" in with_rr.get("meta_info", {}) + @classmethod async def _collect_results_async(cls): results = {}