Skip to content

Commit 680824d

Browse files
authored
Support routed_experts_start_len (#2185)
1 parent 474861a commit 680824d

2 files changed

Lines changed: 30 additions & 4 deletions

File tree

slime/backends/megatron_utils/actor.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -445,7 +445,7 @@ def train_actor(self, rollout_id: int, rollout_data: RolloutBatch, external_data
445445
and not self.args.use_critic
446446
and not self.args.keep_old_actor
447447
and not self.args.use_opd
448-
and not self.args.use_routing_replay
448+
and (not self.args.use_routing_replay or self.args.use_rollout_routing_replay)
449449
and self.args.advantage_estimator != "gspo"
450450
)
451451
if (

slime/utils/types.py

Lines changed: 29 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -353,20 +353,46 @@ def _apply_meta_info(
353353
if routed_experts is not None:
354354
if args is None:
355355
raise ValueError("args is required to decode routed experts metadata.")
356-
expected_rows = len(self.tokens) - 1
356+
routed_experts_start_len = int(meta_info.get("routed_experts_start_len", 0) or 0)
357+
if routed_experts_start_len < 0:
358+
raise ValueError(
359+
f"SGLang routed_experts_start_len must be non-negative, got {routed_experts_start_len}."
360+
)
361+
expected_rows = max(0, len(self.tokens) - 1 - routed_experts_start_len)
357362
expected_numel = expected_rows * args.num_layers * args.moe_router_topk
358363
if routed_experts.numel() != expected_numel:
359364
raise ValueError(
360365
"SGLang routed_experts element count does not match sample tokens: "
361366
f"got={routed_experts.numel()}, expected={expected_numel} "
362-
f"(tokens={len(self.tokens)}, num_layers={args.num_layers}, "
367+
f"(tokens={len(self.tokens)}, routed_experts_start_len={routed_experts_start_len}, "
368+
f"num_layers={args.num_layers}, "
363369
f"moe_router_topk={args.moe_router_topk})."
364370
)
365-
self.rollout_routed_experts = routed_experts.reshape(
371+
routed_experts = routed_experts.reshape(
366372
expected_rows,
367373
args.num_layers,
368374
args.moe_router_topk,
369375
)
376+
if routed_experts_start_len == 0:
377+
self.rollout_routed_experts = routed_experts
378+
else:
379+
existing = self.rollout_routed_experts
380+
if existing is None:
381+
raise ValueError(
382+
"Cannot append partial routed experts without existing routed experts "
383+
f"(routed_experts_start_len={routed_experts_start_len})."
384+
)
385+
if not torch.is_tensor(existing):
386+
existing = torch.as_tensor(existing, dtype=routed_experts.dtype)
387+
if existing.shape[0] < routed_experts_start_len:
388+
raise ValueError(
389+
"Existing routed experts shorter than routed_experts_start_len: "
390+
f"existing_rows={existing.shape[0]}, routed_experts_start_len={routed_experts_start_len}."
391+
)
392+
self.rollout_routed_experts = torch.cat(
393+
[existing[:routed_experts_start_len], routed_experts],
394+
dim=0,
395+
)
370396

371397
if not update_terminal_info or "finish_reason" not in meta_info:
372398
return

0 commit comments

Comments
 (0)