From 61d501427d07b80080140e8ca3537f0510cceac8 Mon Sep 17 00:00:00 2001 From: Shuwen Wang <47200617+alphabetc1@users.noreply.github.com> Date: Tue, 8 Sep 2026 13:10:22 +0800 Subject: [PATCH] fix: make dfs weight ordering iterative (#38313) --- .../sglang/srt/mem_cache/base_prefix_cache.py | 35 +++++++++++-------- 1 file changed, 21 insertions(+), 14 deletions(-) diff --git a/python/sglang/srt/mem_cache/base_prefix_cache.py b/python/sglang/srt/mem_cache/base_prefix_cache.py index 530d630e8..66df08eda 100644 --- a/python/sglang/srt/mem_cache/base_prefix_cache.py +++ b/python/sglang/srt/mem_cache/base_prefix_cache.py @@ -266,25 +266,32 @@ def _dfs_weight_order( node: len(indices) for node, indices in last_node_to_indices.items() } - def calc_weight(node: Any) -> None: - for child in node.children.values(): - calc_weight(child) - node_to_weight[node] = node_to_weight.get(node, 0) + node_to_weight.get( - child, 0 - ) - - calc_weight(root_node) + stack: list[tuple[Any, bool]] = [(root_node, False)] + while stack: + node, visited = stack.pop() + if visited: + weight = node_to_weight.get(node, 0) + for child in node.children.values(): + weight += node_to_weight.get(child, 0) + node_to_weight[node] = weight + continue + stack.append((node, True)) + for child in reversed(list(node.children.values())): + stack.append((child, False)) order: list[int] = [] - def append_dfs(node: Any) -> None: + stack = [(root_node, False)] + while stack: + node, visited = stack.pop() + if visited: + order.extend(last_node_to_indices.get(node, ())) + continue children = list(node.children.values()) children.sort(key=lambda child: -node_to_weight.get(child, 0)) - for child in children: - append_dfs(child) - order.extend(last_node_to_indices.get(node, ())) - - append_dfs(root_node) + stack.append((node, True)) + for child in reversed(children): + stack.append((child, False)) return order