[docker] fix routing replay with pp and other bugfixes (#395)

This commit is contained in:
Zilin Zhu
2025-09-28 14:55:24 +08:00
committed by GitHub
parent d862870c6a
commit eb5604e20f
2 changed files with 53 additions and 23 deletions
+16 -21
View File
@@ -364,7 +364,7 @@ index 63ee9d1f..b90b744c 100644
ops.append(recv_next_op)
if len(ops) > 0:
diff --git a/megatron/core/transformer/moe/moe_utils.py b/megatron/core/transformer/moe/moe_utils.py
index 235b6f6a..640273ba 100644
index 235b6f6a..39693a03 100644
--- a/megatron/core/transformer/moe/moe_utils.py
+++ b/megatron/core/transformer/moe/moe_utils.py
@@ -1,9 +1,11 @@
@@ -379,7 +379,7 @@ index 235b6f6a..640273ba 100644
from megatron.core import parallel_state
from megatron.core.process_groups_config import ModelCommProcessGroups
@@ -506,6 +508,48 @@ def pad_routing_map(routing_map: torch.Tensor, pad_multiple: int) -> torch.Tenso
@@ -506,6 +508,47 @@ def pad_routing_map(routing_map: torch.Tensor, pad_multiple: int) -> torch.Tenso
return routing_map
@@ -395,8 +395,8 @@ index 235b6f6a..640273ba 100644
+ all_routing_replays = []
+
+ def __init__(self):
+ self.forward_indices = 0
+ self.backward_indices = []
+ self.forward_index = 0
+ self.backward_index = 0
+ self.top_indices_list = []
+ RoutingReplay.all_routing_replays.append(self)
+
@@ -404,19 +404,18 @@ index 235b6f6a..640273ba 100644
+ self.top_indices_list.append(top_indices)
+
+ def pop_forward(self):
+ top_indices = self.top_indices_list[self.forward_indices]
+ self.backward_indices.append(self.forward_indices)
+ self.forward_indices += 1
+ top_indices = self.top_indices_list[self.forward_index]
+ self.forward_index += 1
+ return top_indices
+
+ def pop_backward(self):
+ backward_indices = self.backward_indices.pop()
+ top_indices = self.top_indices_list[backward_indices]
+ top_indices = self.top_indices_list[self.backward_index]
+ self.backward_index += 1
+ return top_indices
+
+ def clear(self):
+ self.forward_indices = 0
+ self.backward_indices = []
+ self.forward_index = 0
+ self.backward_index = 0
+ self.top_indices_list = []
+
+ @staticmethod
@@ -428,7 +427,7 @@ index 235b6f6a..640273ba 100644
def topk_routing_with_score_function(
logits: torch.Tensor,
topk: int,
@@ -553,7 +597,7 @@ def topk_routing_with_score_function(
@@ -553,7 +596,7 @@ def topk_routing_with_score_function(
expert_bias=expert_bias,
)
@@ -437,7 +436,7 @@ index 235b6f6a..640273ba 100644
if group_topk:
return group_limited_topk(
scores=scores,
@@ -566,6 +610,30 @@ def topk_routing_with_score_function(
@@ -566,6 +609,30 @@ def topk_routing_with_score_function(
else:
return torch.topk(scores, k=topk, dim=1)
@@ -469,7 +468,7 @@ index 235b6f6a..640273ba 100644
if use_pre_softmax:
scores = torch.softmax(logits, dim=-1, dtype=torch.float32).type_as(logits)
diff --git a/megatron/core/transformer/moe/router.py b/megatron/core/transformer/moe/router.py
index 6b20b862..80786f84 100644
index 6b20b862..405d95f8 100644
--- a/megatron/core/transformer/moe/router.py
+++ b/megatron/core/transformer/moe/router.py
@@ -1,5 +1,6 @@
@@ -479,7 +478,7 @@ index 6b20b862..80786f84 100644
from abc import ABC, abstractmethod
from typing import Optional
@@ -156,6 +157,19 @@ class TopKRouter(Router):
@@ -156,6 +157,15 @@ class TopKRouter(Router):
self.local_tokens_per_expert = None
self.expert_bias = None
@@ -487,14 +486,10 @@ index 6b20b862..80786f84 100644
+ from .moe_utils import RoutingReplay, set_routing_replay
+ self.routing_replay = RoutingReplay()
+
+ def forward_hook(*args, **kwargs):
+ def pre_forward_hook(*args, **kwargs):
+ set_routing_replay(self.routing_replay)
+
+ def backward_hook(*args, **kwargs):
+ set_routing_replay(self.routing_replay)
+
+ self.register_forward_pre_hook(forward_hook)
+ self.register_full_backward_pre_hook(backward_hook)
+ self.register_forward_pre_hook(pre_forward_hook)
+
def _maintain_float32_expert_bias(self):
"""
+37 -2
View File
@@ -248,8 +248,25 @@ index 867ffe91b..708d69075 100644
-
EntryClass = [Glm4MoeForCausalLM]
diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py
index 845baff56..f04c50787 100644
--- a/python/sglang/srt/server_args.py
+++ b/python/sglang/srt/server_args.py
@@ -1002,8 +1002,10 @@ class ServerArgs:
# Check TP size
if self.tp_size > 1:
- raise ValueError(
- "Currently only TP size 1 is supported for deterministic inference."
+ os.environ["NCCL_ALGO"] = "allreduce:tree"
+ self.disable_custom_all_reduce = True
+ logger.warning(
+ "NCCL_ALGO is set to 'allreduce:tree' and custom all reduce is disabled for deterministic inference when TP size > 1."
)
# Warnings on MoE models
diff --git a/python/sglang/srt/speculative/eagle_utils.py b/python/sglang/srt/speculative/eagle_utils.py
index e6c55df18..263845736 100644
index e6c55df18..2b0df0f37 100644
--- a/python/sglang/srt/speculative/eagle_utils.py
+++ b/python/sglang/srt/speculative/eagle_utils.py
@@ -189,6 +189,10 @@ class EagleDraftInput:
@@ -263,7 +280,25 @@ index e6c55df18..263845736 100644
else:
# in some cases(e.g draft_extend), we have not filtered the batch by `unfinished_index`
self.topk_p = self.topk_p[new_indices]
@@ -1077,7 +1081,7 @@ def create_accept_length_filter(
@@ -211,6 +215,17 @@ class EagleDraftInput:
self.verified_id = torch.cat([self.verified_id, spec_info.verified_id], axis=0)
self.topk_p = torch.cat([self.topk_p, spec_info.topk_p])
self.topk_index = torch.cat([self.topk_index, spec_info.topk_index])
+ if self.accept_length is not None and spec_info.accept_length is not None:
+ self.accept_length = torch.cat([self.accept_length, spec_info.accept_length])
+ self.accept_length_cpu = self.accept_length.tolist()
+ elif self.accept_length is not None:
+ zeros = torch.zeros([spec_info.verified_id.shape[0]],dtype=self.accept_length.dtype,device=self.accept_length.device)
+ self.accept_length = torch.cat([self.accept_length, zeros])
+ self.accept_length_cpu = self.accept_length.tolist()
+ elif spec_info.accept_length is not None:
+ zeros = torch.zeros([self.verified_id.shape[0]],dtype=self.accept_length.dtype,device=self.accept_length.device)
+ self.accept_length = torch.cat([zeros, spec_info.accept_length])
+ self.accept_length_cpu = self.accept_length.tolist()
@dataclass
@@ -1077,7 +1092,7 @@ def create_accept_length_filter(
return accept_length_filter