Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
63 commits
Select commit Hold shift + click to select a range
dbce3b3
docs(issue26): record correctness branch base and Step 0 status
fwyc0573 Sep 21, 2026
a1f8e45
docs(issue26): land the candidate and vLLM 0.10.2 source audit
fwyc0573 Sep 21, 2026
0adf04d
Merge the completed oversized-module split into the correctness branch
fwyc0573 Sep 21, 2026
8d46c1e
docs(issue26): record that the split prerequisite is met and Step 2 i…
fwyc0573 Sep 21, 2026
6ab521d
fix(scheduler): keep round-robin DP rotation across scheduling calls
fwyc0573 Sep 21, 2026
5d4d169
Merge the refactor tip so this branch carries its own DP-placement co…
fwyc0573 Sep 21, 2026
c18eb2c
docs(issue26): record the Step 2 measurement and scope Step 3
fwyc0573 Sep 21, 2026
33f0d5a
docs(issue26): correct what the sync state actually partitions
fwyc0573 Sep 21, 2026
d244bde
docs(issue26): record the maintainer review decisions D1 and D2
fwyc0573 Sep 21, 2026
7f079e1
docs(issue26): self-review of the decision record; correct D1 mechani…
fwyc0573 Sep 21, 2026
18edbbe
docs(issue26): transcribe the contaminated label's manifest before de…
fwyc0573 Sep 21, 2026
3d8c8a0
Merge the corrected fidelity gate and the split's final evidence
fwyc0573 Sep 22, 2026
ceac2b4
test(scheduler): make the W2 placement tests state where requests land
fwyc0573 Sep 22, 2026
3d47417
docs(issue26): re-measure W2 with harness and source at one revision
fwyc0573 Sep 22, 2026
65ed8a7
fix(scheduler): give a monolithic Replica one shared forward across p…
fwyc0573 Sep 22, 2026
bdff4aa
docs(review): record the W3 fidelity measurement and rebuilt controls
fwyc0573 Sep 22, 2026
cdfcdf5
docs(review): mark Step 3 published and record the PR 35 W3 section
fwyc0573 Sep 22, 2026
10dd474
feat(scheduler): add an opt-in vLLM-style DP request placement policy
fwyc0573 Sep 22, 2026
0fd12c4
docs(review): record the W4 measurement, controls and fidelity result
fwyc0573 Sep 22, 2026
cb8ba14
docs(review): mark Step 4 published and record the PR 35 W4 section
fwyc0573 Sep 22, 2026
de2bee8
docs(review): close W5 as not ported and record why
fwyc0573 Sep 22, 2026
bbbfcaa
docs(review): mark the W5 closure published and record the PR 35 section
fwyc0573 Sep 22, 2026
7269bac
fix(profiling): complete the legacy fused-MoE expert computation
fwyc0573 Sep 22, 2026
697f219
test(profiling): add native fused-MoE expert parity against vLLM
fwyc0573 Sep 22, 2026
79f599a
docs(profiling): record the grouped-GEMM scope and its identity limits
fwyc0573 Sep 22, 2026
d695679
docs(review): record the W6 native parity submission and close artifa…
fwyc0573 Sep 22, 2026
cad3afd
docs(review): record the verified W7 facts and its companion-reposito…
fwyc0573 Sep 22, 2026
e0e265c
docs(review): record the PR 35 W6 and W7 sections and the re-run base…
fwyc0573 Sep 22, 2026
1b95187
fix(cc_backend): accept an empty collective through the collective-si…
fwyc0573 Sep 22, 2026
beded3c
test(governance): scan only the sources Frontier owns under frontier/
fwyc0573 Sep 22, 2026
ff2ed0d
docs(review): record the W6 native parity result and the W7 delivery
fwyc0573 Sep 22, 2026
d881357
docs: complete the cluster-scheduler list and record the opt-in DP po…
fwyc0573 Sep 22, 2026
8730509
docs(records): record the Step 8 combined regression and the deferred…
fwyc0573 Sep 22, 2026
75f355e
docs(records): close Step 8 and write the completion archive
fwyc0573 Sep 22, 2026
0137269
docs(records): plan Step 9, PP>1 support for the opt-in vLLM DP place…
fwyc0573 Sep 22, 2026
0d025f8
Merge branch 'refactor/oversized-module-split' into fix/issue26-corre…
fwyc0573 Sep 22, 2026
f7c31e4
fix(scheduler): credit decoding requests at a dense layer inside a mi…
fwyc0573 Sep 22, 2026
ca1b9b6
test(profiling): pass the FP8 block shape and skip the CPU boundary t…
fwyc0573 Sep 22, 2026
57ffa5b
docs(records): apply the 2026-09-22 external review to the records an…
fwyc0573 Sep 22, 2026
c231322
docs(records): record the review-correction publication SHAs and the …
fwyc0573 Sep 22, 2026
9df18d1
docs(records): FP8 native rerun passes; second Step 9 plan review aga…
fwyc0573 Sep 22, 2026
e955406
docs(plan): R9-01 keeps the pipeline-room test in the policy, compute…
fwyc0573 Sep 22, 2026
cf1cb74
test(dp-placement): add the vLLM engine-iteration reference loop for …
fwyc0573 Sep 22, 2026
9e3afe8
docs(step9): record the P1 evidence, the design checkpoint, and W9-01
fwyc0573 Sep 22, 2026
4a510d4
docs(step9): record G1 instrumentation, the case binding, and Step 9 …
fwyc0573 Sep 22, 2026
4c2d573
docs(step9): record the W9-01 fix on PR 36 and the resume order
fwyc0573 Sep 23, 2026
f4f12a0
docs(step9): drop the moving PR 36 head from the W9-01 note
fwyc0573 Sep 23, 2026
887d34b
docs(step9): record the PR 36 round-2 remediation for W9-01
fwyc0573 Sep 23, 2026
8315d9b
docs(step9): record the PR 36 pre-merge untrack for W9-01
fwyc0573 Sep 23, 2026
4ab1964
fix(scheduler): stage admission deadlock with attention-DP lanes unde…
fwyc0573 Sep 23, 2026
dd9b8d9
Merge remote-tracking branch 'origin/main' into fix/issue26-correctne…
fwyc0573 Sep 23, 2026
03d5f24
test(stage-admission): skip dispatched sync rooms in the drain report
fwyc0573 Sep 23, 2026
d3e6e78
docs(step9): record the W9-01 merge-forward and composition check
fwyc0573 Sep 23, 2026
9b5d7be
docs(step9): complete P1(b) and propose the D9-2 report key
fwyc0573 Sep 23, 2026
d1a2a06
docs(step9): record the D9-2 decision and the C1 PP3 amendment
fwyc0573 Sep 23, 2026
2ffe78d
feat(dp-placement): publish vLLM DP loads per engine iteration under PP
fwyc0573 Sep 23, 2026
bacdbb4
test(dp-placement): cover schedule-time reports in the real loop unde…
fwyc0573 Sep 23, 2026
c1a570d
docs(dp-placement): record Step 9 validation and extend the policy's …
fwyc0573 Sep 23, 2026
104b6ff
docs(review): state the Step 9 source line counts and W9-04 reachabil…
fwyc0573 Sep 23, 2026
339e6bd
docs(progress): record the Step 9 publication and PR 35 body update
fwyc0573 Sep 23, 2026
2ffb062
fix(scheduler): withdraw a placeholder when its lane joins the forward
fwyc0573 Sep 23, 2026
a9ff5d5
docs(w9-04): record the placeholder-withdrawal fix and its validation
fwyc0573 Sep 23, 2026
33a1c8e
docs(progress): record the W9-04 publication and PR 35 body update
fwyc0573 Sep 23, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 6 additions & 1 deletion AGENTS.md
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,8 @@
| 2026-09-05 | Clarified Replica-local collective backend materialization. |
| 2026-09-05 | Added Python module size and plain ML-system naming guidance for refactors. |
| 2026-09-06 | Added cleanup-first and split-analysis requirements for critical modules above 2,000 lines. |
| 2026-09-22 | Completed the cluster-scheduler implementation list and recorded the opt-in vLLM DP placement policy and its supported scope. |
| 2026-09-23 | Extended the vLLM DP placement policy's supported scope to pipeline parallelism. |

- Current public branch supports `co-location`, sequential PDD / `pd-disaggregation`, and sequential PD-AF / `pd-af-disaggregation`.
- The public co-location, PDD, and PD-AF examples explicitly select `--cc_backend_config_type analytical` for one-click smoke runs using the built-in analytical model.
Expand Down Expand Up @@ -611,9 +613,12 @@ The scheduling logic is split across four distinct layers to mirror real-world s
2. **Cluster Scheduler** (`ClusterSchedulerRegistry`):
- **Role**: Manages workload distribution within a specific `ClusterType` (e.g., selecting which Replica gets a request).
- **Implementations**:
- `RoundRobinClusterScheduler`: Distributes requests cyclically.
- `RoundRobinClusterScheduler`: Distributes requests cyclically over replicas and, inside each replica, over attention-DP lanes. The ordinal persists across scheduling calls, so an identical ordered request stream lands identically however it is divided between calls.
- `LORClusterScheduler`: Least Outstanding Requests (load balancing).
- `RandomClusterScheduler`: Random assignment.
- `StickyRoundRobinClusterScheduler`: Round-robin over targets, pinned per session so a session's later requests return to the same target.
- `StickyLORClusterScheduler`: Least Outstanding Requests with the same per-session pinning.
- `VllmLoadBalancingClusterScheduler`: Models vLLM V1's internal DP selection, choosing the lane with the lowest `waiting * 4 + running` score from a load snapshot the frontend observes with a delay. Opt-in and deliberately narrow: one `co-location` replica, the `vllm_v1` replica scheduler, and either a MoE model or `attn_dp=1`, at any pipeline depth. As in vLLM, a lane publishes its load when it admits a batch while its pipeline still has room, and otherwise with its next completion. The constructor rejects everything else. No placement or timing equivalence with a real vLLM deployment is claimed.

3. **Replica Scheduler** (`ReplicaSchedulerRegistry`):
- **Role**: Operates at the level of a single `Replica` (GPU node/instance).
Expand Down
40 changes: 40 additions & 0 deletions docs/profiling/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@

| Date | Summary of Changes |
| ---------- | ------------------ |
| 2026-09-22 | Recorded the `moe_grouped_gemm` measurement scope and the limits of the backend identity columns. |
| 2026-09-17 | Corrected TP8 GDN launch to use eight distributed processes. |
| 2026-09-14 | Documented standard ROCm/GDN output contracts and the experimental SGLang boundary. |

Expand Down Expand Up @@ -207,6 +208,45 @@ TP domain independently of these MoE EP values. At runtime, Step3 shared-expert
work uses `attn_tp` in co-location/PREFILL/unified DECODE and the role-local
`moe_tp` in DECODE_FFN.

#### What `moe_grouped_gemm` measures

`moe_grouped_gemm` is the complete local expert computation: the first expert
GEMM, the gated activation, the optional activation quantization, the second
expert GEMM with the routing weights applied, and the local reduction over one
token's top-k expert outputs. That reduction is a per-token sum, not a
collective, so it carries no communication cost. `MOE_FAMILY` has no separate
operator for it, and the vLLM functional backend has always included it, so
counting it here counts the reduction once on each backend.

The two backends still time different envelopes. The low-level path runs block
alignment (`moe_align_block_size`) once before the timed step on caller-prepared
buffers; vLLM's `fused_experts` aligns inside the call, per chunk, with its own
workspace. A `moe_grouped_gemm` row from the functional backend therefore
includes preparation work that a low-level row does not, while Frontier adds
`moe_shuffling` as a separate term for both. Treat rows from the two backends as
different measurements under one operator name, and do not mix them in one
dataset.

Before 2026-09-22 the vLLM 0.10.x low-level profiling path omitted the gated
activation and the reduction. Rows produced by that path under-measure
`moe_grouped_gemm`, and the gap grows with token count: on the checked-in
`a800/qwen3-a3b-30b-moe` dataset the two missing kernels are an estimated 16.5%
of the corrected value at 4096 tokens, against 6.8% at the median row. Rows
produced by the functional backend, or by any path after that date, are
complete.

#### Backend identity columns cannot date a row

`moe_grouped_gemm_backend` records `vllm_fused` for both the vLLM 0.10.x
low-level path and the current functional path, so it does not distinguish a
row measured before that fix from one measured after it.
`profiling_patch_tag` appears in one historical CSV header, but nothing in the
source writes it, so it is not a live mechanism either.

Re-profile rather than infer. If you need corrected `moe_grouped_gemm` timings
from a `vllm>=0.10,<0.11` environment, re-run the producer; an existing row
cannot be checked for completeness from its own metadata.

### Standard GDN on ROCm

The standard GDN producer uses the vLLM Qwen3.5 module path and writes a
Expand Down
9 changes: 9 additions & 0 deletions frontier/config/cluster_scheduler_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,3 +46,12 @@ class StickyLORClusterSchedulerConfig(BaseClusterSchedulerConfig):
@staticmethod
def get_type():
return ClusterSchedulerType.STICKY_LOR


@dataclass
class VllmLoadBalancingClusterSchedulerConfig(BaseClusterSchedulerConfig):
"""Select vLLM V1's internal DP placement. Selecting it is the only knob."""

@staticmethod
def get_type():
return ClusterSchedulerType.VLLM_LOAD_BALANCING
1 change: 1 addition & 0 deletions frontier/config/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -120,6 +120,7 @@
RoundRobinClusterSchedulerConfig,
StickyLORClusterSchedulerConfig,
StickyRoundRobinClusterSchedulerConfig,
VllmLoadBalancingClusterSchedulerConfig,
)
from frontier.config.execution_time_predictor_config import (
BaseExecutionTimePredictorConfig,
Expand Down
2 changes: 1 addition & 1 deletion frontier/events/cluster_schedule_event.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,7 +39,7 @@ def handle_event(
logger.info(f"Cluster scheduling started at {self.time:.3f}s: "
f"{self._cluster_type.name} cluster with {queue_size} requests in queue")

self._request_mapping = cluster_scheduler.schedule()
self._request_mapping = cluster_scheduler.schedule_at(self.time)

# DEBUG: Log request mapping
mapping_summary = {}
Expand Down
5 changes: 5 additions & 0 deletions frontier/events/global_batch_end_event.py
Original file line number Diff line number Diff line change
Expand Up @@ -178,6 +178,11 @@ def _current_request_entries() -> list[tuple[int, object]]:
thinking_round_start_times=self._thinking_round_start_times,
)
replica_scheduler.on_batch_end(self._batch) # decrement running batches
# After the lane's request-state transition, so a routing policy that
# reads lane populations here observes the post-step load.
cluster_scheduler.on_replica_batch_end(
self.time, self._replica_id, self._replica_local_id, self._batch
)

thinking_requeue_events: List[BaseEvent] = []
for index, request in pre_batch_request_entries:
Expand Down
59 changes: 54 additions & 5 deletions frontier/profiling/moe/moe_vllm_kernel.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,8 @@
VLLM_API_VERSION = None
FP8_QUANT_AVAILABLE = False
_functional_fused_experts = None
# Bound only by the low-level import below, like `_functional_fused_experts`.
_vllm_custom_ops = None
_FUNCTIONAL_MXFP4_STATES: Dict[Tuple[Any, ...], Dict[str, Any]] = {}


Expand Down Expand Up @@ -126,6 +128,11 @@ def plan_mxfp4_weight_layout(
try_get_optimal_moe_config,
get_config_dtype_str,
)
# The gated activation and the local top-k reduction belong to the same
# low-level API as the two kernel invocations. Importing them here means
# a build without them selects the functional path instead of running an
# incomplete expert computation.
from vllm import _custom_ops as _vllm_custom_ops

VLLM_API_VERSION = "0.10.x"
except ImportError:
Expand Down Expand Up @@ -368,13 +375,14 @@ def _run_fused_moe_iteration(
w2: torch.Tensor,
intermediate_cache1: torch.Tensor,
intermediate_cache2: torch.Tensor,
intermediate_cache3: torch.Tensor,
out_hidden_states: torch.Tensor,
topk_weights: torch.Tensor,
sorted_token_ids: torch.Tensor,
expert_ids: torch.Tensor,
num_tokens_post_padded: torch.Tensor,
top_k: int,
config: Dict,
expert_hidden_dim_per_partition: int,
block_dims: Optional[Tuple[int, int]],
A_scale: Optional[torch.Tensor] = None,
w1_scale: Optional[torch.Tensor] = None,
Expand All @@ -383,6 +391,24 @@ def _run_fused_moe_iteration(
per_channel_quant: bool = False,
block_shape: Optional[List[int]] = None,
) -> None:
"""Run one complete local expert computation, as vLLM's own path does.

Reference: `fused_experts_impl` in vLLM 0.10.x `fused_moe.py`, which runs
the first expert GEMM, a gated activation, the optional activation
quantization, the second expert GEMM with the routing weights, and a local
reduction of the top-k expert outputs. The buffers follow the same naming:

```text
intermediate_cache1 (M, top_k, 2 * E) first GEMM output, gate | up
intermediate_cache2 (M * top_k, E) gated activation output
intermediate_cache3 (M, top_k, H) second GEMM output, per expert
out_hidden_states (M, H) local top-k reduction
```

The reduction is a local sum over the `top_k` expert outputs of one token.
It is not a collective and adds no communication cost.
"""

_invoke_kernel(
A=A.contiguous(),
B=w1.contiguous(),
Expand All @@ -401,8 +427,15 @@ def _run_fused_moe_iteration(
block_shape=block_shape,
)

intermediate_cache1_flat = intermediate_cache1.view(-1, intermediate_cache1.shape[-1])
intermediate_cache2_input = intermediate_cache1_flat[:, :expert_hidden_dim_per_partition].contiguous()
# Gated SiLU over the two halves of the first projection. Taking the gate
# half alone would skip this kernel and feed the wrong operand to the second
# GEMM. The caller always materializes `w1` with `2 * E` rows, so the gated
# layout is the only one this path can produce.
torch.ops._C.silu_and_mul(
intermediate_cache2,
intermediate_cache1.view(-1, intermediate_cache1.shape[-1]),
)
intermediate_cache2_input = intermediate_cache2

intermediate_A_scale = None
if use_fp8:
Expand All @@ -415,7 +448,7 @@ def _run_fused_moe_iteration(
_invoke_kernel(
A=intermediate_cache2_input,
B=w2.contiguous(),
C=intermediate_cache2.contiguous(),
C=intermediate_cache3.contiguous(),
topk_weights=topk_weights.contiguous(),
sorted_token_ids=sorted_token_ids.contiguous(),
expert_ids=expert_ids.contiguous(),
Expand All @@ -430,6 +463,9 @@ def _run_fused_moe_iteration(
block_shape=block_shape,
)

# Local reduction of one token's top-k expert outputs.
_vllm_custom_ops.moe_sum(intermediate_cache3, out_hidden_states)


def _run_functional_fused_experts_iteration(
A: torch.Tensor,
Expand Down Expand Up @@ -910,12 +946,24 @@ def _step() -> None:
dtype=output_dtype,
)
intermediate_cache2 = torch.empty(
num_tokens * top_k,
expert_hidden_dim_per_partition,
device=device,
dtype=output_dtype,
)
intermediate_cache3 = torch.empty(
num_tokens,
top_k,
hidden_dim,
device=device,
dtype=output_dtype,
)
out_hidden_states = torch.empty(
num_tokens,
hidden_dim,
device=device,
dtype=output_dtype,
)

def _step() -> None:
_run_fused_moe_iteration(
Expand All @@ -924,13 +972,14 @@ def _step() -> None:
w2=w2,
intermediate_cache1=intermediate_cache1,
intermediate_cache2=intermediate_cache2,
intermediate_cache3=intermediate_cache3,
out_hidden_states=out_hidden_states,
topk_weights=topk_weights,
sorted_token_ids=sorted_token_ids,
expert_ids=expert_ids,
num_tokens_post_padded=num_tokens_post_padded,
top_k=top_k,
config=config,
expert_hidden_dim_per_partition=expert_hidden_dim_per_partition,
block_dims=block_dims,
A_scale=A_scale,
w1_scale=w1_scale,
Expand Down
88 changes: 87 additions & 1 deletion frontier/scheduler/cluster_scheduler/base_cluster_scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -106,7 +106,11 @@
handle_decode_attn_arrival as handle_m2n_decode_attn_arrival,
handle_decode_ffn_arrival,
)
from frontier.scheduler.utils.sync_entry import enter_decode_sync, enter_prefill_sync
from frontier.scheduler.utils.sync_entry import (
enter_decode_sync,
enter_prefill_sync,
uses_shared_forward_room,
)
from frontier.scheduler.utils.pdaf_phase import (
prepare_decode_attn_batch_phase,
apply_decode_attn_batch_phase,
Expand All @@ -128,6 +132,7 @@
)
from frontier.scheduler.utils.prefill_collective import handle_prefill_sync_collective
from frontier.scheduler.utils.decode_collective import handle_decode_sync_collective
from frontier.scheduler.utils.forward_collective import handle_forward_sync_collective
from frontier.scheduler.utils.afd_metadata import aggregate_afd_metadata
from frontier.scheduler.utils.request_selection import collect_active_requests
from frontier.scheduler.utils.replica_schedulers import build_replica_scheduler_maps
Expand Down Expand Up @@ -422,6 +427,47 @@ def _schedule_batch_mode(self) -> List[Tuple[int, int, Request]]:
def add_request(self, request: Request) -> None:
self._request_queue.append(request)

def schedule_at(self, time: float) -> List[Tuple[int, int, Request]]:
"""Route the queue at a known simulation time.

The default ignores the time and routes exactly as before, so every
policy that does not need it is unaffected. A policy whose placement
depends on wall-clock progress -- a delayed load snapshot, for instance
-- overrides this instead of reaching for a mutable time bridge.
"""

return self.schedule()

def on_replica_batch_scheduled(
self,
time: float,
replica_id: int,
replica_local_id: int | None,
batch: Batch,
) -> None:
"""Observe one batch admitted by a MONOLITHIC or PREFILL lane. Inert by default.

Called after the lane counts the batch as running, so a policy that
reads lane populations here sees the post-admission state.
"""

return None

def on_replica_batch_end(
self,
time: float,
replica_id: int,
replica_local_id: int | None,
batch: Batch,
) -> None:
"""Observe one Replica-local batch completion. Inert by default.

Called after the batch's request-state transition, so a policy that
reads lane populations here sees the post-step state.
"""

return None

def get_replica(self, replica_id: int) -> Replica:
return self._cluster.replicas[replica_id]

Expand Down Expand Up @@ -1073,6 +1119,36 @@ def _uses_shared_decode_ep_wave(self, batch: Batch, layer_id: int) -> bool:
require_moe_layer=True,
)

def _on_forward_ep_wave_ready(self, *, time: float, replica_id: int, stage_id: int, batch: Batch, layer_id: int, replica_local_id: int | None = None, cohort_batches: dict[int, Batch] | None = None, metrics_store=None) -> List:
"""Schedule one shared monolithic forward wave, across mixed lanes."""

return schedule_layer_wave(
self,
mode="forward",
time=time,
replica_id=replica_id,
stage_id=stage_id,
batch=batch,
layer_id=layer_id,
replica_local_id=replica_local_id,
cohort_batches=cohort_batches,
metrics_store=metrics_store,
)

def on_forward_sync_collective(self, time: float, replica_id: int, stage_id: int, batch_global_id: int, sync_stage: str, layer_id: int, metrics_store):
"""Complete one shared monolithic forward through the utility handler."""

return handle_forward_sync_collective(
self,
time,
replica_id,
stage_id,
batch_global_id,
sync_stage,
layer_id,
metrics_store,
)

def _uses_shared_decode_layer_path(self, batch: Batch, layer_id: int) -> bool:
"""Return whether a shared-domain DECODE model needs layer stepping."""
model_config = getattr(getattr(self._config, "replica_config", None), "model_config", None)
Expand Down Expand Up @@ -1122,6 +1198,11 @@ def on_dense_layer_complete(

def on_prefill_sync_collective(self, time: float, replica_id: int, stage_id: int, batch_global_id: int, sync_stage: str, layer_id: int, metrics_store, *, direct_batch: Optional[Batch] = None):
"""Delegate PREFILL collective completion to the utility handler."""
if direct_batch is None and uses_shared_forward_room(self):
return self.on_forward_sync_collective(
time, replica_id, stage_id, batch_global_id, sync_stage,
layer_id, metrics_store,
)
return handle_prefill_sync_collective(
self,
time,
Expand Down Expand Up @@ -1208,6 +1289,11 @@ def on_decode_sync(
def on_decode_sync_collective(self, time: float, replica_id: int, stage_id: int, batch_global_id: int, sync_stage: str, layer_id: int, metrics_store, *, direct_batch: Optional[Batch] = None):
"""Delegate DECODE collective completion to the utility handler."""

if direct_batch is None and uses_shared_forward_room(self):
return self.on_forward_sync_collective(
time, replica_id, stage_id, batch_global_id, sync_stage,
layer_id, metrics_store,
)
return handle_decode_sync_collective(
self,
time,
Expand Down
Loading