mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
[docker] fix routing replay with pp and other bugfixes (#395)
This commit is contained in:
@@ -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):
|
||||
"""
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user