From 913f966bbd69ff0a37524eca7f0433423199cd1d Mon Sep 17 00:00:00 2001 From: Humaira Firdowse Mohammed Date: Mon, 14 Sep 2026 09:31:47 -0700 Subject: [PATCH 1/7] dyncp changes --- .gitignore | 1 + .../megatron-dynamic-context-parallel.md | 157 +++++++++ docs/index.md | 1 + ...3-32b-4n4g-megatron-dynamiccp-profile.yaml | 10 + ...en3-32b-4n4g-megatron-dynamiccp-quick.yaml | 40 +++ .../distributed/dynamic_context_parallel.py | 186 +++++++++++ nemo_rl/distributed/named_sharding.py | 11 + nemo_rl/distributed/tensor_serialization.py | 46 +++ nemo_rl/models/megatron/data.py | 81 +++-- nemo_rl/models/megatron/dynamic_cp.py | 188 +++++++++++ nemo_rl/models/megatron/setup.py | 18 ++ nemo_rl/models/megatron/train.py | 132 ++++++-- nemo_rl/models/policy/__init__.py | 2 + nemo_rl/models/policy/dynamic_cp.py | 253 +++++++++++++++ nemo_rl/models/policy/lm_policy.py | 92 +++++- nemo_rl/models/policy/tq_policy.py | 2 + .../policy/workers/megatron_policy_worker.py | 83 +++-- pyrefly.toml | 4 + .../functional/dynamic_cp_attention_parity.py | 299 ++++++++++++++++++ tests/functional/dynamic_cp_loss_parity.py | 145 +++++++++ .../test_dynamic_context_parallel.py | 136 ++++++++ .../distributed/test_dynamic_cp_dispatch.py | 175 ++++++++++ .../distributed/test_tensor_serialization.py | 40 +++ .../megatron/test_dynamic_cp_scaling.py | 53 ++++ 24 files changed, 2067 insertions(+), 88 deletions(-) create mode 100644 docs/design-docs/megatron-dynamic-context-parallel.md create mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-profile.yaml create mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-quick.yaml create mode 100644 nemo_rl/distributed/dynamic_context_parallel.py create mode 100644 nemo_rl/distributed/tensor_serialization.py create mode 100644 nemo_rl/models/megatron/dynamic_cp.py create mode 100644 nemo_rl/models/policy/dynamic_cp.py create mode 100644 tests/functional/dynamic_cp_attention_parity.py create mode 100644 tests/functional/dynamic_cp_loss_parity.py create mode 100644 tests/unit/distributed/test_dynamic_context_parallel.py create mode 100644 tests/unit/distributed/test_dynamic_cp_dispatch.py create mode 100644 tests/unit/distributed/test_tensor_serialization.py create mode 100644 tests/unit/models/megatron/test_dynamic_cp_scaling.py diff --git a/.gitignore b/.gitignore index 11242f63adf..0a7ef3abb9e 100644 --- a/.gitignore +++ b/.gitignore @@ -30,6 +30,7 @@ actual_test_nemo_gym_sanity.json tests/functional/*/ tests/unit/unit_results.json tests/unit/unit_results/ +perf_runs/ # Cache uv_cache/ diff --git a/docs/design-docs/megatron-dynamic-context-parallel.md b/docs/design-docs/megatron-dynamic-context-parallel.md new file mode 100644 index 00000000000..254a9cb29a5 --- /dev/null +++ b/docs/design-docs/megatron-dynamic-context-parallel.md @@ -0,0 +1,157 @@ +# Megatron dynamic context parallelism + +Dynamic CP lets the Ray driver assign short sequences to independent model +replicas and longer sequences to larger context-parallel groups within the same +optimizer step. It uses the Megatron-Core version pinned through Megatron Bridge. +No dependency source edits or submodule changes are required. + +## Configuration + +```yaml +policy: + megatron_cfg: + context_parallel_size: 1 + pipeline_model_parallel_size: 1 + dynamic_context_parallel: + enabled: true + min_size: 1 + max_size: 8 + train_tokens_per_rank: 1024 + logprob_tokens_per_rank: 1024 + sequence_packing: + enabled: true + dynamic_batching: + enabled: false +``` + +The static CP size defines the initialized topology. Active CP sizes are powers +of two within the combined DP × CP domain, including sizes larger or smaller +than static CP. `max_size: null` uses the complete domain. The domain and CP +bounds must be powers of two. Budgets count padded tokens per CP rank, before +tensor sequence parallelism. A sequence that cannot fit at the maximum size is +rejected before dispatch. + +This path currently supports PP=1 and the standard Ray policy data path. +TransferQueue, split execution, model-owned multimodal packing, atomic preference +pairs, MTP, draft training, fused linear logprobs, and training CUDA graphs are +not supported. Omit or disable the configuration for existing static behavior. + +For MoE, the minimum active size is raised until `active_CP * TP >= EP`. Workers +also check that each actual expert communication group is contained within its +task's ranks. This constraint prevents expert collectives from crossing task +boundaries; it is not an end-to-end validation of MoE auxiliary losses or router +replay. The smoke recipe uses a dense model. + +## Dispatch and execution + +One scheduling lane contains a complete TP replica. The driver packs sequences +of the same required CP size and assigns each packed task to an aligned, +contiguous group of lanes. The worker verifies those lane IDs against MCore's +initialized DP × CP group and resolves the active attention group using MCore's +hybrid-group API. TP ranks receive the same task payload. EP is a constraint on +group placement, not another data-sharding axis. + +Each phase covers the entire DP × CP domain. Unoccupied lanes execute a small +zero-mask placeholder, so every rank performs the same number of +forward/backward calls and reaches gradient synchronization together. A barrier +between phases keeps changes of active group coordinated, including MCore +reruns. This conservative scheduler introduces more synchronization and padding +than Megatron's balanced scheduler; it is not a throughput-equivalent port of +that scheduler. + +MCore creates hybrid groups during Bridge initialization. NeMo-RL then selects +the standard no-pipeline executor for these driver-planned phases. Optimizer +and DDP groups remain fixed. `PackedSeqParams.local_cp_size` and `cp_group` +carry the active attention topology; size one explicitly uses `cp_group=None`. +The NeMo-RL runtime also initializes the TE CP stream when a model built with +static CP=1 first needs larger groups. + +The pinned attention implementation still reads its constructor process-group +collection for RoPE. NeMo-RL isolates that collection and binds its CP entry to +the active group through forward and backward, then restores the original +collection. The model receives a shallow copy of packed metadata with a real +singleton group for CP=1, because RoPE interprets `None` as a static-group +fallback. Loss and logprob code retain the explicit size-one/None convention. + +Dynamic-CP workers register a Ray serializer for CPU tensor results. It uses +NumPy byte arrays and preserves tensor dtype and shape, including BF16. This +keeps the driver independent of MCore even when MCore replaces PyTorch's tensor +storage loader with a backend-specific function. MCore's checkpoint loader is +left intact. + +## Packing, outputs, and normalization + +Padding is computed before scheduling and checked again after packing on each +worker. Each sequence is aligned to the least common multiple of the user +padding factor and the active CP, sequence-parallel, and precision requirements. +CP>1 uses the two-chunk balanced layout. Padding tokens and placeholder samples +are excluded from the objective and returned outputs. + +Static `REPLICATED_AXES` remains valid for static CP. Dynamic dispatch removes CP +from both input replication and output replication filters. Outputs are gathered +from every DP × CP lane; the first lane of each active task owns its reconstructed +rows. The driver keeps those rows once and restores original sample order, +checking for missing or duplicate sample IDs. A static CP-rank-zero filter would +discard valid results. + +The driver computes valid sequence and token denominators from the unique +global batch before replication, separately for every optimizer step. The +differentiable CP logprob gather replicates the loss over the active CP group. +For this gathered loss, NeMo-RL multiplies by + +```text +number_of_phases / (static_CP * active_CP) +``` + +This cancels the pinned no-pipeline executor's `static_CP / number_of_phases` +scaling and compensates the active-CP gather's backward SUM. DDP sums gradients +over the fixed DP × CP domain. Metrics are retained only on task owners before +global aggregation. `tests/unit/models/megatron/test_dynamic_cp_scaling.py` +checks this multiplier against the installed MCore loss callback. A future +MCore pin must pass that contract test and distributed parity before adoption. + +## GB200 five-step smoke test + +From the RL checkout on the Slurm login node: + +```bash +DRY_RUN=0 bash perf_runs/run_gb200_dynamic_cp.sh +``` + +The launcher requests four nodes with four GPUs each, partition `batch`, QoS +`short`, and a two-hour limit. It uses the container and HF cache defaults in the +launcher, which can be overridden through its environment variables. Omitting +`DRY_RUN=0` prints the submission command without submitting. + +The container preflight checks the Bridge pin against the RL checkout and runs +CPU dispatch tests. A four-GPU test then compares the actual packing, CP +logprob collectives, loss, gradients, and an SGD update against an unsharded +reference for active sizes 4, 2, and 1 with base CP=1 and base CP=2. Another +four-GPU test exercises fused RoPE and TE attention in a small transformer, +including TP=1/2 and base CP=1/2. It also checks actual worker score serialization +and reassembled sample logprobs against a CP=1 reference. Then +`grpo-qwen3-32b-4n4g-megatron-dynamiccp-quick.yaml` runs five GRPO steps with +TP=2, PP=1, base CP=1, and active CP up to 8. Both policy and reference logprob +passes are enabled. Checkpoints and validation are disabled for the smoke test. + +Driver output is in `-logs/ray-driver.log`; TensorBoard metrics are in +`logs/dynamic-cp-`. The final metrics check requires steps 1–5, +finite loss and gradient norm, positive valid-token counts, training/scoring +importance ratios within 0.01 of one, and generation KL below 0.1, and writes +`smoke_result.json` in that log directory. A completed smoke run verifies +execution and finite training metrics, not convergence or a speedup over static CP. + +### Ten-step Nsight profile + +`perf_runs/run_gb200_dynamic_cp_profile.sh` runs the same dense Qwen3-32B setup +for ten steps with TensorBoard and W&B enabled. It profiles only Megatron policy +workers because that is where dynamic CP executes. By default, Nsight captures +all ten steps with `PROFILE_STEP_RANGE=1:11`. Use +`PROFILE_STEP_RANGE=3:6` for a smaller steady-state-only report. The launcher +is a dry run unless `DRY_RUN=0` is explicitly supplied. + +Each runtime phase has an NVTX label such as +`dynamic_cp/phase_7/cp_4/lane_2/data`. The post-run check requires ten dynamic +training plans, transitions between CP=1 and CP>1, valid training metrics, and +at least one completed policy `.nsys-rep` file on the head node. `ray.sub` also +copies reports from all nodes into `-logs/ray/**/nsight/`. diff --git a/docs/index.md b/docs/index.md index b93e8e942fe..b6c4b8f96ff 100644 --- a/docs/index.md +++ b/docs/index.md @@ -387,6 +387,7 @@ design-docs/modelopt-real-quant-architecture.md design-docs/nccl-reshard-refit.md design-docs/media-token-validity-mask.md design-docs/automodel-context-parallel.md +design-docs/megatron-dynamic-context-parallel.md ``` ```{toctree} diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-profile.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-profile.yaml new file mode 100644 index 00000000000..f7f6f0d8387 --- /dev/null +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-profile.yaml @@ -0,0 +1,10 @@ +defaults: ./grpo-qwen3-32b-4n4g-megatron-dynamiccp-quick.yaml +grpo: + max_num_steps: 10 +logger: + log_dir: logs/grpo-qwen3-32b-4n4g-megatron-dynamiccp-profile + wandb_enabled: true + tensorboard_enabled: true + wandb: + project: nemo-rl-dynamic-cp + name: qwen3-32b-4n4g-dynamic-cp-profile diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-quick.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-quick.yaml new file mode 100644 index 00000000000..a5a7c2c660b --- /dev/null +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-quick.yaml @@ -0,0 +1,40 @@ +defaults: ./grpo-qwen3-32b-4n4g.yaml +grpo: + num_prompts_per_step: 8 + num_generations_per_prompt: 8 + max_num_steps: 5 + val_period: 1000 + val_at_start: false + val_at_end: false +loss_fn: + force_on_policy_ratio: false +checkpointing: + enabled: false +policy: + train_global_batch_size: 64 + train_micro_batch_size: 1 + logprob_batch_size: 1 + max_total_sequence_length: 4096 + megatron_cfg: + tensor_model_parallel_size: 2 + pipeline_model_parallel_size: 1 + context_parallel_size: 1 + expert_model_parallel_size: 1 + dynamic_context_parallel: + enabled: true + min_size: 1 + max_size: 8 + train_tokens_per_rank: 1024 + logprob_tokens_per_rank: 1024 + fp8_cfg: + enabled: false + sequence_packing: + enabled: true + dynamic_batching: + enabled: false +logger: + log_dir: logs/grpo-qwen3-32b-4n4g-megatron-dynamiccp-quick + wandb_enabled: false + tensorboard_enabled: true +data_plane: + enabled: false diff --git a/nemo_rl/distributed/dynamic_context_parallel.py b/nemo_rl/distributed/dynamic_context_parallel.py new file mode 100644 index 00000000000..30075180904 --- /dev/null +++ b/nemo_rl/distributed/dynamic_context_parallel.py @@ -0,0 +1,186 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# http://www.apache.org/licenses/LICENSE-2.0 +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""CPU-only plans for Ray-dispatched hybrid DP/CP execution. + +A lane is one complete TP replica. Phases partition the DP*CP lanes into +aligned, power-of-two attention groups. All lanes execute one packed forward +per phase; an empty assignment executes masked padding. This deliberately +uses more phase boundaries than Megatron's load-balancing scheduler, while +preserving its group-transition and gradient-accumulation semantics. +""" + +from dataclasses import dataclass +from math import lcm + +from pydantic import BaseModel, PositiveInt + + +class DynamicContextParallelConfig(BaseModel, extra="forbid"): + """Runtime CP scheduling; token budgets are per CP rank, before SP.""" + + enabled: bool = False + min_size: PositiveInt = 1 + max_size: PositiveInt | None = None + train_tokens_per_rank: PositiveInt = 2048 + logprob_tokens_per_rank: PositiveInt = 2048 + + +@dataclass(frozen=True) +class CPAssignment: + """One packed task, shared by all its participating lanes.""" + + sample_indices: tuple[int, ...] + lane_start: int + cp_size: int + pad_multiple: int + padded_tokens: int + + +@dataclass(frozen=True) +class CPRankStep: + """Assignments and unique-data normalizers for one optimizer step.""" + + assignments: tuple[CPAssignment, ...] + valid_sequences: float + valid_tokens: float + + +@dataclass(frozen=True) +class CPRankPlan: + """Serializable plan; process groups are resolved only on the workers.""" + + lane: int + lane_ranks: tuple[int, ...] + steps: tuple[CPRankStep, ...] + + +def plan_cp_phases( + lengths: list[int], + *, + lanes: int, + min_size: int, + max_size: int, + tokens_per_rank: int, + sequence_parallel_size: int, + user_pad_multiple: int, + token_alignment: int = 1, +) -> tuple[tuple[CPAssignment, ...], ...]: + """Pack equal-CP sequences and place tasks into disjoint aligned groups. + + Sample indices refer to the original batch, never to a sorted copy. + Padding is included when choosing CP size and admitting samples to bins. + """ + if lanes < 1 or lanes & (lanes - 1): + raise ValueError("Dynamic CP requires a power-of-two DP*CP domain") + sizes = (min_size, max_size) + if any(s < 1 or s & (s - 1) for s in sizes) or not min_size <= max_size <= lanes: + raise ValueError("CP bounds must be powers of two within the DP*CP domain") + if ( + min(tokens_per_rank, sequence_parallel_size, user_pad_multiple, token_alignment) + < 1 + ): + raise ValueError("Token budgets and padding factors must be positive") + bins: dict[int, list[tuple[list[int], int, int]]] = {} + for index in sorted(range(len(lengths)), key=lambda i: (-lengths[i], i)): + length = lengths[index] + if length < 2: + raise ValueError("Dynamic CP requires at least two input tokens per sample") + size = min_size + while True: + multiple = lcm( + user_pad_multiple, + (2 * size if size > 1 else 1) + * sequence_parallel_size + * token_alignment, + ) + padded = (length + multiple - 1) // multiple * multiple + if padded <= tokens_per_rank * size: + break + size *= 2 + if size > max_size: + raise ValueError( + f"Sample {index} of length {length} exceeds the dynamic CP token budget" + ) + size_bins = bins.setdefault(size, []) + for bin_index, (members, used, factor) in enumerate(size_bins): + if used + padded <= tokens_per_rank * size: + members.append(index) + size_bins[bin_index] = (members, used + padded, factor) + break + else: + size_bins.append(([index], padded, multiple)) + + pending = [ + CPAssignment(tuple(members), 0, size, factor, used) + for size in sorted(bins, reverse=True) + for members, used, factor in bins[size] + ] + phases = [] + while pending: + phase = [] + remaining = [] + cursor = 0 + for task in pending: + if cursor + task.cp_size <= lanes: + phase.append( + CPAssignment( + task.sample_indices, + cursor, + task.cp_size, + task.pad_multiple, + task.padded_tokens, + ) + ) + cursor += task.cp_size + else: + remaining.append(task) + while cursor < lanes: + factor = lcm( + user_pad_multiple, + (2 * min_size if min_size > 1 else 1) + * sequence_parallel_size + * token_alignment, + ) + phase.append( + CPAssignment( + (), cursor, min_size, factor, ((2 + factor - 1) // factor) * factor + ) + ) + cursor += min_size + phases.append(tuple(phase)) + pending = remaining + return tuple(phases) + + +def assignment_for_lane(phase: tuple[CPAssignment, ...], lane: int) -> CPAssignment: + """Resolve exactly one task for a lane, including padding tasks.""" + matches = [a for a in phase if a.lane_start <= lane < a.lane_start + a.cp_size] + if len(matches) != 1: + raise ValueError(f"Lane {lane} has {len(matches)} assignments in a CP phase") + return matches[0] + + +def cp_loss_multiplier( + *, + active_cp_size: int, + schedule_cp_size: int, + num_microbatches: int, + replicated_cp_loss: bool, +) -> float: + """Cancel MCore's legacy averaging and the gathered loss's duplication. + + The no-pipeline executor scales by static CP / microbatches. Attention + and the differentiable CP gather use active CP, which may differ. + """ + if min(active_cp_size, schedule_cp_size, num_microbatches) < 1: + raise ValueError("CP sizes and microbatch count must be positive") + replicas = active_cp_size if replicated_cp_loss else 1 + return num_microbatches / (schedule_cp_size * replicas) diff --git a/nemo_rl/distributed/named_sharding.py b/nemo_rl/distributed/named_sharding.py index 234a8094e30..467e80f23ce 100644 --- a/nemo_rl/distributed/named_sharding.py +++ b/nemo_rl/distributed/named_sharding.py @@ -28,6 +28,17 @@ ) +def replicated_axes(*, dynamic_cp: bool = False) -> tuple[str, ...]: + """Axes sharing a dispatch result; runtime CP tasks have explicit owners. + + REPLICATED_AXES remains the static-layout default. Dynamic CP callers + collect all CP lanes and deduplicate using their task plan instead. + """ + return tuple( + axis for axis in REPLICATED_AXES if not dynamic_cp or axis != "context_parallel" + ) + + class NamedSharding: """Represents an N-dimensional arrangement of ranks with named axes, facilitating data sharding, replication, and collection based on these axes. diff --git a/nemo_rl/distributed/tensor_serialization.py b/nemo_rl/distributed/tensor_serialization.py new file mode 100644 index 00000000000..b6a9e5a0c93 --- /dev/null +++ b/nemo_rl/distributed/tensor_serialization.py @@ -0,0 +1,46 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# http://www.apache.org/licenses/LICENSE-2.0 +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Backend-independent tensor payloads for policy worker RPC results.""" + +import numpy as np +import torch + +TensorPayload = tuple[np.ndarray, torch.dtype, tuple[int, ...]] + + +def tensor_to_payload(tensor: torch.Tensor) -> TensorPayload: + """Encode a CPU RPC tensor without invoking Torch's storage pickler. + + MCore can replace that pickler's loader with a Megatron function, which + cannot be imported in the Ray driver. Byte arrays also preserve BF16. + """ + if tensor.device.type != "cpu": + raise ValueError("Policy RPC tensors must be moved to CPU before serialization") + data = tensor.detach().contiguous().reshape(-1).view(torch.uint8).numpy() + return data, tensor.dtype, tuple(tensor.shape) + + +def tensor_from_payload(payload: TensorPayload) -> torch.Tensor: + """Restore a writable CPU tensor without any backend dependencies.""" + data, dtype, shape = payload + if data.size == 0: + return torch.empty(0, dtype=dtype).reshape(shape) + return torch.from_numpy(data.copy()).view(dtype).reshape(shape) + + +def register_policy_tensor_serializer() -> None: + """Register the policy worker's Ray serializer, retaining MCore's loader.""" + # Ray is optional in tensor packing/round-trip unit tests. + import ray.util + + ray.util.register_serializer( + torch.Tensor, serializer=tensor_to_payload, deserializer=tensor_from_payload + ) diff --git a/nemo_rl/models/megatron/data.py b/nemo_rl/models/megatron/data.py index 6ef2eaa7f30..dddd076708b 100644 --- a/nemo_rl/models/megatron/data.py +++ b/nemo_rl/models/megatron/data.py @@ -26,18 +26,22 @@ get_context_parallel_rank, get_context_parallel_world_size, ) +from megatron.core.rerun_state_machine import RerunDataIterator from megatron.core.utils import StragglerDetector from nemo_rl.algorithms.loss.interfaces import LossFunction, LossType from nemo_rl.data.multimodal_utils import PACKED_MULTIMODAL_FIELDS from nemo_rl.distributed.batched_data_dict import BatchedDataDict +from nemo_rl.distributed.dynamic_context_parallel import CPRankPlan, CPRankStep from nemo_rl.distributed.model_utils import _get_tokens_on_this_cp_rank from nemo_rl.models.megatron.common import _round_up_to_multiple +from nemo_rl.models.megatron.dynamic_cp import RuntimeCPContext, planned_microbatches from nemo_rl.models.megatron.hybridep import ( get_packed_seq_padding_mask, pad_packed_seq_for_hybridep, uses_hybridep_flex_dispatcher, ) +from nemo_rl.models.policy.dynamic_cp import dynamic_cp_config from nemo_rl.utils.r3_trace import ( r3_trace_verify_forward_enabled, trace_cp_routed_experts, @@ -241,6 +245,8 @@ def get_microbatch_iterator( delegate_pack_to_model: bool = False, delegate_mtp_loss_mask_to_model: bool = False, model_slices_context_parallel_inputs: bool = False, + cp_plan: Optional[CPRankPlan] = None, + cp_step: Optional[CPRankStep] = None, ) -> Tuple[Iterator[ProcessedMicrobatch], int, int, int, int]: """Create a processed microbatch iterator from a batch of data. @@ -262,6 +268,27 @@ def get_microbatch_iterator( - seq_dim_size: Sequence length dimension size - padded_seq_length: Padded sequence length for pipeline parallelism (may differ from seq_length) """ + # This check also protects scoring entrypoints that have not adopted the + # driver plan: falling back to static CP would silently misroute samples. + if cp_plan is None and dynamic_cp_config(cfg) is not None: + raise ValueError("Dynamic CP execution requires a driver-provided rank plan") + if cp_plan is not None: + if cp_step is None: + raise ValueError( + "Dynamic CP iterator requires an explicit optimizer-step plan" + ) + width = data["input_ids"].shape[1] + return ( + RerunDataIterator( + planned_microbatches(data, cp_plan, cp_step, straggler_timer) + ), + len(cp_step.assignments), + 1, + width, + width, + ) + if cp_step is not None: + raise ValueError("CP step supplied without a rank plan") micro_batch_size = mbs pad_factor = 1 pad_full_seq_to = None @@ -363,6 +390,7 @@ def process_microbatch( straggler_timer: Optional[StragglerDetector] = None, create_packed_seq_padding_mask: bool = False, prepad_packed_seq_for_hybridep: bool = False, + cp_context: Optional[RuntimeCPContext] = None, ) -> ProcessedInputs: """Process a microbatch for Megatron model forward pass.""" if create_packed_seq_padding_mask and model_slices_context_parallel_inputs: @@ -374,6 +402,11 @@ def process_microbatch( raise NotImplementedError( "HybridEP input prepadding requires NeMo-owned sequence packing." ) + cp_rank = cp_context.rank if cp_context is not None else get_context_parallel_rank() + cp_size = ( + cp_context.size if cp_context is not None else get_context_parallel_world_size() + ) + cp_group = cp_context.group if cp_context is not None else None ctx = straggler_timer(bdata=True) if straggler_timer is not None else nullcontext() with ctx: input_ids = data_dict["input_ids"] @@ -500,8 +533,8 @@ def process_microbatch( pad_individual_seqs_to_multiple_of, pad_packed_seq_to_multiple_of, pad_full_seq_to, - cp_rank=get_context_parallel_rank(), - cp_size=get_context_parallel_world_size(), + cp_rank=cp_rank, + cp_size=cp_size, ) if model_slices_context_parallel_inputs: packed_seq_params = PackedSeqParams( @@ -545,26 +578,24 @@ def process_microbatch( packed_seq_params=packed_seq_params, cu_seqlens_padded=cu_seqlens_padded, pad_packed_seq_to_multiple_of=pad_packed_seq_to_multiple_of, - cp_rank=get_context_parallel_rank(), - cp_size=get_context_parallel_world_size(), + cp_rank=cp_rank, + cp_size=cp_size, ) full_padding_mask = get_packed_seq_padding_mask( cu_seqlens=cu_seqlens, cu_seqlens_padded=cu_seqlens_padded, total_tokens=input_ids.shape[1], ) - if ( - model_slices_context_parallel_inputs - or get_context_parallel_world_size() == 1 - ): + if model_slices_context_parallel_inputs or cp_size == 1: padding_mask = full_padding_mask else: cp_partition_indices = get_packed_seq_cp_partition_indices( packed_seq_params, total_tokens=input_ids.shape[1], - cp_size=get_context_parallel_world_size(), - cp_rank=get_context_parallel_rank(), + cp_size=cp_size, + cp_rank=cp_rank, device=input_ids.device, + cp_group=cp_group, ) padding_mask = full_padding_mask.index_select( 1, cp_partition_indices @@ -588,16 +619,17 @@ def process_microbatch( seq_lengths, cu_seqlens, cu_seqlens_padded, - get_context_parallel_rank(), - get_context_parallel_world_size(), + cp_rank, + cp_size, ) if model_slices_context_parallel_inputs: cp_partition_indices = get_packed_seq_cp_partition_indices( packed_seq_params, total_tokens=input_ids.shape[1], - cp_size=get_context_parallel_world_size(), - cp_rank=get_context_parallel_rank(), + cp_size=cp_size, + cp_rank=cp_rank, device=input_ids.device, + cp_group=cp_group, ) routed_experts_cp_sharded = routed_experts.index_select( 1, cp_partition_indices @@ -637,8 +669,8 @@ def process_microbatch( else input_ids_cp_sharded ), cp_token_identity_verified_count=verified_token_count, - cp_rank=get_context_parallel_rank(), - cp_size=get_context_parallel_world_size(), + cp_rank=cp_rank, + cp_size=cp_size, ) # Pack pre-computed mtp_loss_mask the same way as input_ids @@ -655,8 +687,8 @@ def process_microbatch( pad_individual_seqs_to_multiple_of, pad_packed_seq_to_multiple_of, pad_full_seq_to, - cp_rank=get_context_parallel_rank(), - cp_size=get_context_parallel_world_size(), + cp_rank=cp_rank, + cp_size=cp_size, ) # Mirror the input_ids layout choice above. A model that # slices CP itself receives the full THD row so it can insert @@ -694,8 +726,8 @@ def process_microbatch( pad_individual_seqs_to_multiple_of, pad_packed_seq_to_multiple_of, pad_full_seq_to, - cp_rank=get_context_parallel_rank(), - cp_size=get_context_parallel_world_size(), + cp_rank=cp_rank, + cp_size=cp_size, ) # Mirror the input_ids layout choice above, for the same # reason the MTP mask does: a model that slices CP itself @@ -741,8 +773,8 @@ def process_microbatch( token_identity_cp_sharded=token_identity_cp_sharded, input_ids_cp_sharded=input_ids_cp_sharded, cp_token_identity_verified_count=verified_token_count, - cp_rank=get_context_parallel_rank(), - cp_size=get_context_parallel_world_size(), + cp_rank=cp_rank, + cp_size=cp_size, ) attention_mask, _, position_ids = get_ltor_masks_and_position_ids( data=input_ids, @@ -761,6 +793,11 @@ def process_microbatch( media_token_validity_mask = data_dict[ "media_token_validity_mask" ].bool() + if cp_context is not None: + if packed_seq_params is None: + raise ValueError("Dynamic CP requires packed THD inputs") + packed_seq_params.local_cp_size = cp_size + packed_seq_params.cp_group = cp_group return ProcessedInputs( input_ids=input_ids, input_ids_cp_sharded=input_ids_cp_sharded, diff --git a/nemo_rl/models/megatron/dynamic_cp.py b/nemo_rl/models/megatron/dynamic_cp.py new file mode 100644 index 00000000000..7d3a54ed5f7 --- /dev/null +++ b/nemo_rl/models/megatron/dynamic_cp.py @@ -0,0 +1,188 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# http://www.apache.org/licenses/LICENSE-2.0 +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""MCore-specific runtime binding for driver-planned context parallelism.""" + +from contextlib import contextmanager +from copy import copy +from dataclasses import dataclass +from typing import Any, Iterator + +import torch +from megatron.core import parallel_state +from megatron.core.transformer.attention import Attention + +from nemo_rl.distributed.dynamic_context_parallel import CPRankPlan, CPRankStep + + +@dataclass(frozen=True) +class RuntimeCPContext: + """Attention group for a single forward/backward, separate from DDP groups.""" + + size: int + rank: int + group: Any + + +def initialize_dynamic_cp_runtime() -> None: + """Initialize CP resources missing on CP=1 builds of the pinned MCore. + + This creates the same stream as TEDotProductAttention's CP constructor; + it does not replace methods, edit dependency files, or change static groups. + Newer MCore creates this lazily, so the initialization is idempotent. + """ + # TE is an optional dependency outside the Megatron worker environment. + from megatron.core.extensions.transformer_engine import TEDotProductAttention + + if TEDotProductAttention.cp_stream is None: + TEDotProductAttention.cp_stream = torch.cuda.Stream() + + +@contextmanager +def preserve_attention_cp_groups(model: torch.nn.Module) -> Iterator[None]: + """Isolate attention's runtime groups from shared model/DDP collections. + + Keep the active group through backward recomputation; restore it after the + complete no-pipeline schedule. TP, DP and optimizer groups are unchanged. + """ + saved = [ + (module, module.pg_collection) + for module in model.modules() + if isinstance(module, Attention) + ] + for module, collection in saved: + module.pg_collection = copy(collection) + try: + yield + finally: + for module, collection in saved: + module.pg_collection = collection + + +def bind_attention_cp_group(model: torch.nn.Module, packed_seq_params: Any) -> Any: + """Bind RoPE as well as TE attention to the microbatch's active group. + + The pinned MCore forwards packed.cp_group to TE but its RoPE still reads + Attention.pg_collection.cp. CP=1 needs a real singleton here because None + means fall back to static CP in MCore's RoPE API. PP=1 supplies that singleton. + Pass a copy to the model so loss/gather metadata retains CP=1's None group. + """ + context = runtime_cp_from_packed(packed_seq_params) + group = context.group + if context.size == 1: + group = parallel_state.get_pipeline_model_parallel_group() + if group.size() != 1: + raise ValueError("Dynamic CP attention requires PP=1") + for module in model.modules(): + if isinstance(module, Attention): + module.pg_collection.cp = group + model_packed = copy(packed_seq_params) + model_packed.cp_group = group + return model_packed + + +def planned_microbatches( + data: Any, plan: CPRankPlan, step: CPRankStep, straggler_timer: Any +) -> Iterator[Any]: + """Yield exactly the driver's phases, binding and validating active groups.""" + # Avoid a cycle: data.py dispatches to this iterator. + from nemo_rl.models.megatron.data import ProcessedMicrobatch, process_microbatch + + domain = parallel_state.get_data_parallel_group(with_context_parallel=True) + tp_rank = parallel_state.get_tensor_model_parallel_rank() + expected = [rank + tp_rank for rank in plan.lane_ranks] + if ( + torch.distributed.get_process_group_ranks(domain) != expected + or domain.rank() != plan.lane + ): + raise ValueError("Ray's DP*CP lane map disagrees with initialized MCore groups") + for phase_index, assignment in enumerate(step.assignments): + size = assignment.cp_size + group = ( + parallel_state.get_hybrid_data_context_parallel_groups(group_size=size) + if size > 1 + else None + ) + rank = plan.lane - assignment.lane_start + if group is not None: + members = expected[assignment.lane_start : assignment.lane_start + size] + if ( + group.size() != size + or group.rank() != rank + or torch.distributed.get_process_group_ranks(group) != members + ): + raise ValueError( + "Active CP group disagrees with the driver's assignment" + ) + expert_group = parallel_state.get_expert_model_parallel_group() + tp_size = parallel_state.get_tensor_model_parallel_world_size() + task_ranks = { + base + offset + for base in plan.lane_ranks[ + assignment.lane_start : assignment.lane_start + size + ] + for offset in range(tp_size) + } + if not set(torch.distributed.get_process_group_ranks(expert_group)).issubset( + task_ranks + ): + raise ValueError( + "Expert communication group crosses dynamic CP task boundaries" + ) + context = RuntimeCPContext(size=size, rank=rank, group=group) + if assignment.sample_indices: + batch = data.select_indices(list(assignment.sample_indices)).to("cuda") + else: + batch = data.select_indices([0]).to("cuda") + # A real attention invocation keeps collective counts aligned, but + # none of this placeholder's targets or metrics belong to the batch. + for key, value in list(batch.items()): + if isinstance(value, torch.Tensor): + batch[key] = torch.zeros_like(value) + batch["input_lengths"].fill_(2) + inputs = process_microbatch( + batch, + seq_length_key="input_lengths", + pack_sequences=True, + pad_individual_seqs_to_multiple_of=assignment.pad_multiple, + straggler_timer=straggler_timer, + cp_context=context, + ) + if inputs.input_ids_cp_sharded.shape[1] * size != assignment.padded_tokens: + raise ValueError( + "Packed worker token count disagrees with the driver's plan" + ) + payload_kind = "data" if assignment.sample_indices else "padding" + # The generator stays paused inside this range while MCore consumes the + # microbatch, so an Nsight trace shows the active CP size for the whole + # forward/backward phase. + with torch.cuda.nvtx.range( + f"dynamic_cp/phase_{phase_index}/cp_{size}/lane_{plan.lane}/{payload_kind}" + ): + yield ProcessedMicrobatch(data_dict=batch, **vars(inputs)) + + +def runtime_cp_from_packed(packed_seq_params: Any) -> RuntimeCPContext: + """Resolve explicit runtime metadata, falling back only for static CP.""" + if packed_seq_params is not None and packed_seq_params.local_cp_size is not None: + size = packed_seq_params.local_cp_size + group = packed_seq_params.cp_group + if size == 1: + if group is not None: + raise ValueError("CP=1 must not carry an attention communication group") + return RuntimeCPContext(size=1, rank=0, group=None) + if group is None or group.size() != size: + raise ValueError("Packed CP size and process group disagree") + return RuntimeCPContext(size=size, rank=group.rank(), group=group) + return RuntimeCPContext( + size=parallel_state.get_context_parallel_world_size(), + rank=parallel_state.get_context_parallel_rank(), + group=parallel_state.get_context_parallel_group(), + ) diff --git a/nemo_rl/models/megatron/setup.py b/nemo_rl/models/megatron/setup.py index ff5793a753c..85cf18e7181 100644 --- a/nemo_rl/models/megatron/setup.py +++ b/nemo_rl/models/megatron/setup.py @@ -75,6 +75,8 @@ from transformers import PreTrainedTokenizerBase from nemo_rl.distributed.model_utils import patch_gpt_model_forward_for_linear_ce_fusion +from nemo_rl.models.megatron.dynamic_cp import initialize_dynamic_cp_runtime +from nemo_rl.models.policy.dynamic_cp import dynamic_cp_config _HF_CONFIG_PATCHED = False @@ -1194,6 +1196,12 @@ def _apply_parallelism_config(model_cfg: Any, config: PolicyConfig) -> None: ] model_cfg.sequence_parallel = config["megatron_cfg"]["sequence_parallel"] model_cfg.context_parallel_size = config["megatron_cfg"]["context_parallel_size"] + if dynamic_cp_config(config) is not None: + # Bridge forwards this to initialize_model_parallel to create hybrid groups. + model_cfg.hybrid_context_parallel = ( + torch.distributed.get_world_size() // model_cfg.tensor_model_parallel_size + > 1 + ) if model_cfg.context_parallel_size > 1: # Either NeMo-RL does the packing+CP-sharding itself (classic mcore @@ -2108,6 +2116,16 @@ def setup_model_and_optimizer( get_position_embedding_ranks=get_position_embedding_ranks, ) + if ( + dynamic_cp_config(policy_cfg) is not None + and megatron_cfg.model.hybrid_context_parallel + ): + initialize_dynamic_cp_runtime() + # The Ray driver supplies already scheduled, uniform phases. Use MCore's + # standard no-pipeline executor, not its TP0 dataloader/redistribution loop. + # Hybrid process groups remain initialized; PackedSeqParams selects them. + megatron_cfg.model.hybrid_context_parallel = False + if megatron_cfg.ft and megatron_cfg.ft.enable_ft_package: fault_tolerance.setup(megatron_cfg, state) fault_tolerance.maybe_setup_simulated_fault(megatron_cfg.ft) diff --git a/nemo_rl/models/megatron/train.py b/nemo_rl/models/megatron/train.py index 5136a696ec5..f78863bb3a5 100644 --- a/nemo_rl/models/megatron/train.py +++ b/nemo_rl/models/megatron/train.py @@ -18,7 +18,7 @@ from typing import Any, Callable, Dict, Iterator, List, Optional, Tuple, Union import torch -from megatron.core import tensor_parallel +from megatron.core import parallel_state, tensor_parallel from megatron.core.models.gpt import GPTModel from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.parallel_state import ( @@ -53,6 +53,7 @@ from nemo_rl.algorithms.loss.utils import _pack_input_ids from nemo_rl.algorithms.utils import mask_out_neg_inf_logprobs from nemo_rl.distributed.batched_data_dict import BatchedDataDict +from nemo_rl.distributed.dynamic_context_parallel import cp_loss_multiplier from nemo_rl.distributed.model_utils import ( allgather_cp_sharded_tensor, distributed_vocab_topk, @@ -64,6 +65,12 @@ from nemo_rl.models.megatron.draft.hidden_capture import ( get_capture_context, ) +from nemo_rl.models.megatron.dynamic_cp import ( + RuntimeCPContext, + bind_attention_cp_group, + preserve_attention_cp_groups, + runtime_cp_from_packed, +) from nemo_rl.models.megatron.opd_full_capture import get_opd_full_capture_context from nemo_rl.models.megatron.router_replay import ( clear_router_replay, @@ -71,6 +78,7 @@ set_router_replay_forward, ) from nemo_rl.models.policy import PolicyConfig +from nemo_rl.models.policy.dynamic_cp import dynamic_cp_config # Union type for any post-processing function (defined after classes below) PostProcessingFunction = Union[ @@ -81,6 +89,21 @@ ] +def _postprocessing_cp_context( + packed_seq_params: Optional[PackedSeqParams], +) -> RuntimeCPContext: + """Use task metadata for dynamic CP and existing accessors for static CP.""" + if packed_seq_params is not None and isinstance( + getattr(packed_seq_params, "local_cp_size", None), int + ): + return runtime_cp_from_packed(packed_seq_params) + return RuntimeCPContext( + size=get_context_parallel_world_size(), + rank=0, + group=get_context_parallel_group(), + ) + + def _prepare_padding_mask_for_model( model: GPTModel, padding_mask: Optional[torch.Tensor], @@ -192,7 +215,11 @@ def model_forward( additional_kwargs = {} # Mamba models currently do not support packed_seq_params if packed_seq_params is not None: - additional_kwargs["packed_seq_params"] = packed_seq_params + additional_kwargs["packed_seq_params"] = ( + bind_attention_cp_group(model, packed_seq_params) + if isinstance(getattr(packed_seq_params, "local_cp_size", None), int) + else packed_seq_params + ) # Pass MTP loss mask to exclude prompt tokens from MTP loss if mtp_loss_mask is not None: @@ -303,6 +330,20 @@ def forward_with_post_processing_fn( attention_mask = processed_mb.attention_mask position_ids = processed_mb.position_ids packed_seq_params = processed_mb.packed_seq_params + cp_context = ( + _postprocessing_cp_context(packed_seq_params) + if packed_seq_params is not None + and isinstance(getattr(packed_seq_params, "local_cp_size", None), int) + else None + ) + if packed_seq_params is not None and isinstance( + getattr(packed_seq_params, "local_cp_size", None), int + ): + # Every phase has one forward/backward per lane. Keep the barrier here, + # rather than in the iterator, so MCore reruns replay it as well. + torch.distributed.barrier( + group=parallel_state.get_data_parallel_group(with_context_parallel=True) + ) cu_seqlens_padded = processed_mb.cu_seqlens_padded mtp_loss_mask = processed_mb.mtp_loss_mask padding_mask = processed_mb.padding_mask @@ -425,6 +466,7 @@ def forward_with_post_processing_fn( input_ids=input_ids, cu_seqlens_padded=cu_seqlens_padded, original_seq_length=original_seq_length, + cp_context=cp_context, ) elif isinstance(post_processing_fn, TeacherFullPayloadPostProcessor): assert original_seq_length is not None @@ -433,6 +475,7 @@ def forward_with_post_processing_fn( input_ids=input_ids, cu_seqlens_padded=cu_seqlens_padded, original_seq_length=original_seq_length, + cp_context=cp_context, hidden_states=( None if opd_full_capture is None @@ -445,6 +488,7 @@ def forward_with_post_processing_fn( data_dict=data_dict, cu_seqlens_padded=cu_seqlens_padded, original_seq_length=original_seq_length, + cp_context=cp_context, ) else: raise TypeError( @@ -519,7 +563,12 @@ def megatron_forward_backward( forward_backward_func = get_forward_backward_func() if use_router_replay: clear_router_replay(model) - with suspend_activation_offload_for_forward_only(model, forward_only): + with ( + suspend_activation_offload_for_forward_only(model, forward_only), + preserve_attention_cp_groups(model) + if dynamic_cp_config(post_processing_fn.cfg) is not None + else nullcontext(), + ): try: return forward_backward_func( forward_step_func=forward_step, @@ -603,6 +652,7 @@ def __call__( Returns: Callable: Function that takes output tensor and returns (loss, metrics) tuple """ + cp_context = _postprocessing_cp_context(packed_seq_params) # A custom prepare_fn (e.g. value models) overrides the default logit prep. logprob_chunk_size = self.cfg.get("logprob_chunk_size", None) if self.prepare_fn is not None: @@ -646,7 +696,7 @@ def __call__( cu_seqlens_q_padded=packed_seq_params.cu_seqlens_q_padded, vocab_parallel_rank=get_tensor_model_parallel_rank(), vocab_parallel_group=get_tensor_model_parallel_group(), - context_parallel_group=get_context_parallel_group(), + context_parallel_group=cp_context.group, ) if "student_logits" in data_dict: # draft + use_fused_linear_logprobs is rejected at setup in @@ -662,7 +712,7 @@ def __call__( loss_weight=float(self.cfg["draft"]["loss_weight"]), vocab_parallel_rank=get_tensor_model_parallel_rank(), vocab_parallel_group=get_tensor_model_parallel_group(), - context_parallel_group=get_context_parallel_group(), + context_parallel_group=cp_context.group, cu_seqlens_q=packed_seq_params.cu_seqlens_q, cu_seqlens_q_padded=packed_seq_params.cu_seqlens_q_padded, d2t=self.d2t, @@ -675,7 +725,7 @@ def __call__( prepare_fn=prepare_loss_input_wrapped, vocab_parallel_rank=get_tensor_model_parallel_rank(), vocab_parallel_group=get_tensor_model_parallel_group(), - context_parallel_group=get_context_parallel_group(), + context_parallel_group=cp_context.group, ) if "student_logits" in data_dict: loss_fn_wrapped = DraftLossWrapper( @@ -685,7 +735,7 @@ def __call__( loss_weight=float(self.cfg["draft"]["loss_weight"]), vocab_parallel_rank=get_tensor_model_parallel_rank(), vocab_parallel_group=get_tensor_model_parallel_group(), - context_parallel_group=get_context_parallel_group(), + context_parallel_group=cp_context.group, ) loss_fn_wrapped = partial( @@ -695,27 +745,19 @@ def __call__( global_valid_toks=global_valid_toks, ) - if self.cp_normalize: - cp_size = get_context_parallel_world_size() - prev_loss_fn = loss_fn_wrapped - - def _div_by_cp_size(*args, **kwargs): - loss, metrics = prev_loss_fn(*args, **kwargs) - return loss / cp_size, metrics - - loss_fn_wrapped = _div_by_cp_size - - # Counteract Megatron's default loss averaging in schedules.py, - # which applies (* cp_size / num_microbatches) to the loss. - cp_size = get_context_parallel_world_size() - num_microbatches = self.num_microbatches - loss_fn_before_mcore_scaling = loss_fn_wrapped - - def _counteract_mcore_loss_averaging(*args, **kwargs): - loss, metrics = loss_fn_before_mcore_scaling(*args, **kwargs) - return loss * num_microbatches / cp_size, metrics + # MCore's unchanged no-pipeline executor scales by STATIC CP / MBs; + # differentiable logprob reconstruction replicates loss over ACTIVE CP. + multiplier = cp_loss_multiplier( + active_cp_size=cp_context.size, + schedule_cp_size=get_context_parallel_world_size(), + num_microbatches=self.num_microbatches, + replicated_cp_loss=self.cp_normalize, + ) + loss_before_scaling = loss_fn_wrapped - loss_fn_wrapped = _counteract_mcore_loss_averaging + def loss_fn_wrapped(*args, **kwargs): + loss, metrics = loss_before_scaling(*args, **kwargs) + return loss * multiplier, metrics return loss_fn_wrapped @@ -737,6 +779,7 @@ def __call__( input_ids: torch.Tensor, cu_seqlens_padded: torch.Tensor, original_seq_length: int, + cp_context: Optional[RuntimeCPContext] = None, ) -> Callable[[torch.Tensor], Tuple[torch.Tensor, Dict[str, torch.Tensor]]]: """Create a post-processing function that computes token log probabilities. @@ -771,7 +814,11 @@ def processor_fn_inner(output_tensor): vocab_end_index=(tp_rank + 1) * output_tensor.shape[-1], group=tp_grp, inference_only=True, - cp_group=get_context_parallel_group(), + cp_group=( + cp_context.group + if cp_context is not None + else get_context_parallel_group() + ), chunk_size=logprob_chunk_size, sampling_params=self.sampling_params, ) @@ -849,6 +896,7 @@ def __call__( input_ids: torch.Tensor, cu_seqlens_padded: torch.Tensor, original_seq_length: int, + cp_context: Optional[RuntimeCPContext] = None, hidden_states: Optional[torch.Tensor] = None, ) -> Callable[[torch.Tensor], Tuple[torch.Tensor, Dict[str, torch.Tensor]]]: """Create the post-processing function for a teacher full-payload forward. @@ -870,9 +918,14 @@ def __call__( input_ids=input_ids, cu_seqlens_padded=cu_seqlens_padded, original_seq_length=original_seq_length, + cp_context=cp_context, ) pack = self.cfg["sequence_packing"]["enabled"] - cp_size = self.cfg["megatron_cfg"]["context_parallel_size"] + cp_size = ( + cp_context.size + if cp_context is not None + else self.cfg["megatron_cfg"]["context_parallel_size"] + ) batch_size = data_dict["input_ids"].shape[0] unpacked_seqlen = data_dict["input_ids"].shape[1] seq_lengths = data_dict["input_lengths"] @@ -908,7 +961,11 @@ def processor_fn_inner(output_tensor): ) if cp_size > 1: - cp_grp = get_context_parallel_group() + cp_grp = ( + cp_context.group + if cp_context is not None + else get_context_parallel_group() + ) if pack: # Per-sequence CP allgather. CP uses a load-balanced # (2 x CP interleaved) layout per sequence, so gathering the @@ -989,6 +1046,7 @@ def __call__( data_dict: BatchedDataDict[Any], cu_seqlens_padded: torch.Tensor, original_seq_length: int, + cp_context: Optional[RuntimeCPContext] = None, ) -> Callable[[torch.Tensor], Tuple[torch.Tensor, Dict[str, torch.Tensor]]]: """Create a post-processing function that computes top-k logits and indices. @@ -1006,7 +1064,11 @@ def __call__( (dummy_loss, {"topk_logits": values, "topk_indices": indices}) """ pack = self.cfg["sequence_packing"]["enabled"] - cp_size = self.cfg["megatron_cfg"]["context_parallel_size"] + cp_size = ( + cp_context.size + if cp_context is not None + else self.cfg["megatron_cfg"]["context_parallel_size"] + ) unpacked_seqlen = data_dict["input_ids"].shape[1] seq_lengths = data_dict["input_lengths"] @@ -1029,8 +1091,12 @@ def processor_fn_inner(output_tensor): chunk_size=chunk_size, ) - if self.cfg["megatron_cfg"]["context_parallel_size"] > 1: - cp_grp = get_context_parallel_group() + if cp_size > 1: + cp_grp = ( + cp_context.group + if cp_context is not None + else get_context_parallel_group() + ) if pack: # Per-sequence CP allgather following packed-sequence logic batch_size = data_dict["input_ids"].shape[0] diff --git a/nemo_rl/models/policy/__init__.py b/nemo_rl/models/policy/__init__.py index 3325edfbfb4..d7ce4c37c01 100644 --- a/nemo_rl/models/policy/__init__.py +++ b/nemo_rl/models/policy/__init__.py @@ -14,6 +14,7 @@ from typing import Any, Literal, NotRequired, TypedDict, Union +from nemo_rl.distributed.dynamic_context_parallel import DynamicContextParallelConfig from nemo_rl.models.generation.interfaces import GenerationConfig from nemo_rl.utils.checkpoint import PretrainedCheckpointConfig @@ -373,6 +374,7 @@ class MegatronConfig(TypedDict): pipeline_model_parallel_size: int num_layers_in_first_pipeline_stage: int | None num_layers_in_last_pipeline_stage: int | None + dynamic_context_parallel: NotRequired[DynamicContextParallelConfig] context_parallel_size: int # Nemotron Omni RADIO/provider booleans. Omit any field to retain the model # provider's checkpoint/default value. diff --git a/nemo_rl/models/policy/dynamic_cp.py b/nemo_rl/models/policy/dynamic_cp.py new file mode 100644 index 00000000000..1734e4ce288 --- /dev/null +++ b/nemo_rl/models/policy/dynamic_cp.py @@ -0,0 +1,253 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# http://www.apache.org/licenses/LICENSE-2.0 +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Ray payload construction and output ownership for dynamic context parallelism.""" + +import logging +from collections import Counter +from dataclasses import dataclass, replace +from typing import Any + +from nemo_rl.distributed.batched_data_dict import BatchedDataDict +from nemo_rl.distributed.dynamic_context_parallel import ( + CPRankPlan, + CPRankStep, + DynamicContextParallelConfig, + assignment_for_lane, + plan_cp_phases, +) +from nemo_rl.distributed.named_sharding import NamedSharding + +logger = logging.getLogger(__name__) + + +def dynamic_cp_config(cfg: dict[str, Any]) -> DynamicContextParallelConfig | None: + """Read optional config without introducing defaults at worker call sites.""" + megatron = cfg.get("megatron_cfg") + raw = megatron.get("dynamic_context_parallel") if megatron is not None else None + if raw is None: + return None + parsed = DynamicContextParallelConfig.model_validate(raw) + return parsed if parsed.enabled else None + + +def validate_dynamic_cp(cfg: dict[str, Any], *, lanes: int) -> None: + """Reject incompatible execution paths before launching model collectives.""" + dynamic = dynamic_cp_config(cfg) + if dynamic is None: + return + mc = cfg["megatron_cfg"] + if not mc["enabled"] or mc["pipeline_model_parallel_size"] != 1: + raise ValueError("Dynamic CP requires Megatron with PP=1") + if not cfg["sequence_packing"]["enabled"] or cfg["dynamic_batching"]["enabled"]: + raise ValueError( + "Dynamic CP requires sequence_packing and disables dynamic_batching" + ) + if mc.get("use_fused_linear_logprobs") or cfg.get("draft", {}).get("enabled"): + raise ValueError( + "Dynamic CP does not support fused linear logprobs or draft training" + ) + if mc.get("cuda_graph_impl") not in (None, "none"): + raise ValueError("Dynamic CP does not support CUDA graph capture") + if mc.get("mtp_num_layers") or mc.get("moe_hybridep_prepad_packed_inputs"): + raise ValueError("Dynamic CP does not support MTP or HybridEP input prepadding") + if cfg["sequence_packing"].get("pair_grouping_key"): + raise ValueError("Dynamic CP does not yet schedule atomic preference pairs") + # Probe even empty plans, so invalid domains/bounds fail at initialization. + tp = mc["tensor_model_parallel_size"] + minimum = dynamic.min_size + while minimum * tp < mc["expert_model_parallel_size"]: + minimum *= 2 + plan_cp_phases( + [], + lanes=lanes, + min_size=minimum, + max_size=dynamic.max_size or lanes, + tokens_per_rank=dynamic.train_tokens_per_rank, + sequence_parallel_size=tp if mc["sequence_parallel"] else 1, + user_pad_multiple=cfg["make_sequence_length_divisible_by"], + ) + + +@dataclass +class CPDispatch: + """Nested DP/CP payloads plus unique result row selection per lane.""" + + data: list[list[BatchedDataDict]] + plans: list[list[CPRankPlan]] + output_rows: list[tuple[list[int], list[int]]] + + +def build_cp_dispatch( + data: BatchedDataDict, + cfg: dict[str, Any], + sharding: NamedSharding, + *, + batch_size: int | None, + training: bool, +) -> CPDispatch: + """Plan before replication; count each original training target once.""" + dynamic = dynamic_cp_config(cfg) + if dynamic is None: + raise ValueError("Dynamic CP dispatch requires enabled configuration") + if data.size == 0: + raise ValueError("Cannot schedule an empty dynamic CP batch") + cp = sharding.shape["context_parallel"] + dp = sharding.shape["data_parallel"] + lanes = cp * dp + mc = cfg["megatron_cfg"] + tp = mc["tensor_model_parallel_size"] + minimum = dynamic.min_size + while minimum * tp < mc["expert_model_parallel_size"]: + minimum *= 2 + group_ranks = tuple( + sharding.get_ranks_by_coord( + pipeline_parallel=0, + data_parallel=i // cp, + context_parallel=i % cp, + tensor_parallel=0, + )[0] + for i in range(lanes) + ) + budget = ( + dynamic.train_tokens_per_rank if training else dynamic.logprob_tokens_per_rank + ) + fp8 = mc.get("fp8_cfg") or {} + alignment = 1 + if fp8.get("enabled"): + alignment = {"blockwise": 128, "mxfp8": 32}.get(fp8["fp8_recipe"], 16) + if ( + mc.get("moe_token_dispatcher_type") == "flex" + and mc.get("moe_flex_dispatcher_backend") == "hybridep" + ): + alignment = max(alignment, 128) + gbs = batch_size if batch_size is not None else data.size + if gbs < 1 or data.size % gbs: + raise ValueError( + "Dynamic CP training data must contain complete global batches" + ) + rank_steps: list[list[CPRankStep]] = [[] for _ in range(lanes)] + for start in range(0, data.size, gbs): + batch = data.select_indices(list(range(start, start + gbs))) + phases = plan_cp_phases( + batch["input_lengths"].tolist(), + lanes=lanes, + min_size=minimum, + max_size=dynamic.max_size or lanes, + tokens_per_rank=budget, + sequence_parallel_size=tp if mc["sequence_parallel"] else 1, + user_pad_multiple=cfg["make_sequence_length_divisible_by"], + token_alignment=alignment, + ) + if training: + if "sample_mask" not in batch or "token_mask" not in batch: + raise ValueError( + "Dynamic CP training requires sample_mask and token_mask" + ) + valid_sequences = float(batch["sample_mask"].sum().item()) + valid_tokens = float( + (batch["token_mask"][:, 1:] * batch["sample_mask"].unsqueeze(-1)) + .sum() + .item() + ) + else: + valid_sequences = valid_tokens = 0.0 + samples_by_cp = Counter() + for phase in phases: + for task in phase: + samples_by_cp[task.cp_size] += len(task.sample_indices) + logger.info( + "Dynamic CP %s: samples=%d phases=%d samples_by_cp=%s valid_sequences=%s valid_tokens=%s", + "train" if training else "score", + gbs, + len(phases), + dict(sorted(samples_by_cp.items())), + valid_sequences, + valid_tokens, + ) + for lane in range(lanes): + assignments = tuple( + replace( + assignment_for_lane(phase, lane), + sample_indices=tuple( + i + start + for i in assignment_for_lane(phase, lane).sample_indices + ), + ) + for phase in phases + ) + rank_steps[lane].append( + CPRankStep(assignments, valid_sequences, valid_tokens) + ) + + payloads, plans, output_rows = [], [], [] + for lane, steps in enumerate(rank_steps): + indices = sorted( + { + i + for step in steps + for task in step.assignments + for i in task.sample_indices + } + ) + # Padding-only lanes still need one prototype row to execute attention. + if not indices: + indices = [0] + local_indices = {index: local for local, index in enumerate(indices)} + remapped_steps = tuple( + replace( + step, + assignments=tuple( + replace( + task, + sample_indices=tuple( + local_indices[i] for i in task.sample_indices + ), + ) + for task in step.assignments + ), + ) + for step in steps + ) + rows, originals, offset = [], [], 0 + for step in steps: + for task in step.assignments: + if lane == task.lane_start: + rows.extend(range(offset, offset + len(task.sample_indices))) + originals.extend(task.sample_indices) + offset += max(1, len(task.sample_indices)) + output_rows.append((rows, originals)) + payloads.append(data.select_indices(indices)) + plans.append(CPRankPlan(lane, group_ranks, remapped_steps)) + return CPDispatch( + data=[payloads[i : i + cp] for i in range(0, lanes, cp)], + plans=[plans[i : i + cp] for i in range(0, lanes, cp)], + output_rows=output_rows, + ) + + +def collect_cp_outputs( + results: list[BatchedDataDict], dispatch: CPDispatch, size: int +) -> BatchedDataDict: + """Keep only task owners and restore original rollout sample order.""" + if len(results) != len(dispatch.output_rows): + raise ValueError("Dynamic CP must return results from every DP*CP lane") + selected, original_indices = [], [] + for result, (rows, indices) in zip(results, dispatch.output_rows): + if rows: + selected.append(result.select_indices(rows)) + original_indices.extend(indices) + if sorted(original_indices) != list(range(size)): + raise ValueError("Dynamic CP output has missing or duplicated sample IDs") + merged = BatchedDataDict.from_batches(selected) + # reorder_data takes each current row's original position, and sorts it + # internally; passing argsort here would apply the inverse permutation. + merged.reorder_data(original_indices) + return merged diff --git a/nemo_rl/models/policy/lm_policy.py b/nemo_rl/models/policy/lm_policy.py index a9e02f94f58..911c07c3c07 100644 --- a/nemo_rl/models/policy/lm_policy.py +++ b/nemo_rl/models/policy/lm_policy.py @@ -30,7 +30,7 @@ SequencePackingArgs, SlicedDataDict, ) -from nemo_rl.distributed.named_sharding import NamedSharding +from nemo_rl.distributed.named_sharding import NamedSharding, replicated_axes from nemo_rl.distributed.ray_actor_environment_registry import get_actor_python_env from nemo_rl.distributed.virtual_cluster import RayVirtualCluster from nemo_rl.distributed.worker_groups import RayWorkerBuilder, RayWorkerGroup @@ -41,6 +41,12 @@ RefitPayloadMode, ) from nemo_rl.models.policy import PolicyConfig +from nemo_rl.models.policy.dynamic_cp import ( + build_cp_dispatch, + collect_cp_outputs, + dynamic_cp_config, + validate_dynamic_cp, +) from nemo_rl.models.policy.interfaces import ( ColocatablePolicyInterface, LogprobOutputSpec, @@ -90,6 +96,8 @@ def _aggregate_megatron_flops_metrics( class Policy(ColocatablePolicyInterface, GenerationInterface): + supports_dynamic_cp_dispatch = True + def __init__( self, cluster: RayVirtualCluster, @@ -330,6 +338,14 @@ def __init__( ], ) + self.dynamic_cp = dynamic_cp_config(config) is not None + if self.dynamic_cp: + if not self.supports_dynamic_cp_dispatch: + raise ValueError( + "Dynamic CP currently requires the Ray payload policy; TQ/split dispatch is not supported" + ) + validate_dynamic_cp(config, lanes=cluster.world_size() // tp_size) + pre_init_queue = RayQueue() worker_kwargs = dict( @@ -651,6 +667,25 @@ def _report_sharded_payload( ) ) + def _get_dynamic_cp_outputs( + self, method: str, data: BatchedDataDict, **kwargs: Any + ) -> BatchedDataDict: + dispatch = build_cp_dispatch( + data, self.cfg, self.sharding_annotations, batch_size=None, training=False + ) + futures = self.worker_group.run_all_workers_sharded_data( + method, + data=dispatch.data, + cp_plan=dispatch.plans, + in_sharded_axes=["data_parallel", "context_parallel"], + replicate_on_axes=list(replicated_axes(dynamic_cp=True)), + output_is_replicated=list(replicated_axes(dynamic_cp=True)), + common_kwargs=kwargs, + ) + return collect_cp_outputs( + self.worker_group.get_all_worker_results(futures), dispatch, data.size + ) + def get_logprobs( self, data: BatchedDataDict[GenerationDatumSpec], @@ -663,6 +698,9 @@ def get_logprobs( We use the convention that the logprob of the first token is 0 so that the sequence length is maintained. The logprob of input token i is specified at position i in the output logprobs tensor. """ + if self.dynamic_cp: + return self._get_dynamic_cp_outputs("get_logprobs", data) + with timer.time("get_logprobs/shard_data") if timer else nullcontext(): sharded_data, unsorted_data_indices = self._shard_for_logprob(data) self._report_sharded_payload(sharded_data, "policy_get_logprobs") @@ -708,6 +746,11 @@ def get_reference_policy_logprobs( Returns: Identical to get_logprobs. """ + if self.dynamic_cp: + return self._get_dynamic_cp_outputs( + "get_reference_policy_logprobs", data, micro_batch_size=micro_batch_size + ) + with ( timer.time("get_reference_policy_logprobs/shard_data") if timer @@ -760,6 +803,10 @@ def get_topk_logits( timer: Optional[Timer] = None, ) -> BatchedDataDict[TopkLogitsOutputSpec]: """Dispatch get_topk_logits to workers (no CP/packed support initially).""" + if self.dynamic_cp: + return self._get_dynamic_cp_outputs( + "get_topk_logits", data, k=k, micro_batch_size=micro_batch_size + ) with timer.time("get_topk_logits/shard_data") if timer else nullcontext(): sharded_data, unsorted_data_indices = self._shard_for_logprob(data) @@ -880,12 +927,28 @@ def train( micro_batch_size = mbs or self.cfg["train_micro_batch_size"] # Shard and replicate the batch with timer.time("policy_training/sharding_data") if timer else nullcontext(): - sharded_data = self._shard_for_train(data, batch_size) - self._report_sharded_payload(sharded_data, "policy_train") + dispatch = ( + build_cp_dispatch( + data, + self.cfg, + self.sharding_annotations, + batch_size=batch_size, + training=True, + ) + if self.dynamic_cp + else None + ) + sharded_data = ( + dispatch.data + if dispatch is not None + else self._shard_for_train(data, batch_size) + ) + if dispatch is None: + self._report_sharded_payload(sharded_data, "policy_train") if self.flops_tracker is not None: self.flops_tracker.reset() - for shard in sharded_data: + for shard in [data] if dispatch is not None else sharded_data: input_lengths = shard["input_lengths"] self.flops_tracker.track_batch(input_lengths.tolist()) @@ -898,17 +961,16 @@ def train( futures = self.worker_group.run_all_workers_sharded_data( "train", data=sharded_data, - in_sharded_axes=["data_parallel"], - replicate_on_axes=[ - "context_parallel", - "tensor_parallel", - "pipeline_parallel", - ], - output_is_replicated=[ - "context_parallel", - "tensor_parallel", - "pipeline_parallel", - ], + **({"cp_plan": dispatch.plans} if dispatch is not None else {}), + in_sharded_axes=["data_parallel", "context_parallel"] + if dispatch is not None + else ["data_parallel"], + replicate_on_axes=list( + replicated_axes(dynamic_cp=dispatch is not None) + ), + output_is_replicated=list( + replicated_axes(dynamic_cp=dispatch is not None) + ), common_kwargs={ "loss_fn": loss_fn, "eval_mode": eval_mode, diff --git a/nemo_rl/models/policy/tq_policy.py b/nemo_rl/models/policy/tq_policy.py index ae674a17c53..2557d3c08bb 100644 --- a/nemo_rl/models/policy/tq_policy.py +++ b/nemo_rl/models/policy/tq_policy.py @@ -112,6 +112,8 @@ class TQPolicy(TQDriverMixin, Policy): rollout actor at first put + driver-/worker-written deltas). """ + supports_dynamic_cp_dispatch = False + def __init__( self, *args: Any, diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index f436b7424ae..8d47a7b6c35 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -58,7 +58,9 @@ ) from nemo_rl.data_plane.worker_mixin import TQWorkerMixin from nemo_rl.distributed.batched_data_dict import BatchedDataDict +from nemo_rl.distributed.dynamic_context_parallel import CPRankPlan from nemo_rl.distributed.named_sharding import NamedSharding +from nemo_rl.distributed.tensor_serialization import register_policy_tensor_serializer from nemo_rl.models.generation.interfaces import GenerationDatumSpec, RefitPayloadMode from nemo_rl.models.generation.megatron.megatron_worker import ( MegatronGenerationMixin, @@ -104,6 +106,7 @@ megatron_forward_backward, ) from nemo_rl.models.policy import PolicyConfig +from nemo_rl.models.policy.dynamic_cp import dynamic_cp_config from nemo_rl.models.policy.interfaces import ( ColocatablePolicyInterface, LogprobOutputSpec, @@ -786,6 +789,13 @@ def __init__( if _model_accepts_media_token_validity_mask(self.model) else None ) + if dynamic_cp_config(self.cfg) is not None: + register_policy_tensor_serializer() + if self.delegate_pack_to_model or self.model_slices_context_parallel_inputs: + raise ValueError("Dynamic CP requires NeMo-owned text-model packing") + if getattr(self._get_model_config(), "mtp_num_layers", None): + raise ValueError("Dynamic CP does not yet support MTP") + if self.model_slices_context_parallel_inputs: if self.delegate_pack_to_model: raise RuntimeError( @@ -956,6 +966,7 @@ def train( gbs: Optional[int] = None, mbs: Optional[int] = None, check_dim_skip_keys: Optional[Iterable[str]] = None, + cp_plan: Optional[CPRankPlan] = None, ) -> dict[str, Any]: """Train the policy on a batch of data with a given loss function. @@ -986,13 +997,16 @@ def train( if mbs is None: mbs = self.cfg["train_micro_batch_size"] local_gbs = gbs // self.dp_size - total_dataset_size = torch.tensor(data.size, device="cuda") - torch.distributed.all_reduce( - total_dataset_size, - op=torch.distributed.ReduceOp.SUM, - group=parallel_state.get_data_parallel_group(), - ) - num_global_batches = int(total_dataset_size.item()) // gbs + if cp_plan is not None: + num_global_batches = len(cp_plan.steps) + else: + total_dataset_size = torch.tensor(data.size, device="cuda") + torch.distributed.all_reduce( + total_dataset_size, + op=torch.distributed.ReduceOp.SUM, + group=parallel_state.get_data_parallel_group(), + ) + num_global_batches = int(total_dataset_size.item()) // gbs if eval_mode: ctx: AbstractContextManager[Any] = torch.no_grad() @@ -1020,16 +1034,26 @@ def train( losses = [] total_num_microbatches = 0 for gb_idx in range(num_global_batches): - gb_result = process_global_batch( - data, - loss_fn=loss_fn, - dp_group=parallel_state.get_data_parallel_group(), - batch_idx=gb_idx, - batch_size=local_gbs, - ) - batch = gb_result["batch"] - global_valid_seqs = gb_result["global_valid_seqs"] - global_valid_toks = gb_result["global_valid_toks"] + cp_step = cp_plan.steps[gb_idx] if cp_plan is not None else None + if cp_step is not None: + batch = data + global_valid_seqs = torch.tensor( + cp_step.valid_sequences, device="cuda" + ) + global_valid_toks = torch.tensor( + cp_step.valid_tokens, device="cuda" + ) + else: + gb_result = process_global_batch( + data, + loss_fn=loss_fn, + dp_group=parallel_state.get_data_parallel_group(), + batch_idx=gb_idx, + batch_size=local_gbs, + ) + batch = gb_result["batch"] + global_valid_seqs = gb_result["global_valid_seqs"] + global_valid_toks = gb_result["global_valid_toks"] # Pre-compute the MTP loss mask, only when MTP is enabled, so # process_microbatch can pack it. @@ -1058,6 +1082,8 @@ def train( delegate_pack_to_model=self.delegate_pack_to_model, delegate_mtp_loss_mask_to_model=self.delegate_mtp_loss_mask_to_model, model_slices_context_parallel_inputs=self.model_slices_context_parallel_inputs, + cp_plan=cp_plan, + cp_step=cp_step, ) # Track total microbatches for MoE aux-loss averaging total_num_microbatches += int(num_microbatches) @@ -1199,6 +1225,17 @@ def train( # keep all microbatch metrics to be normalized later gb_loss_metrics = [] mb_losses = [] + if cp_step is not None: + assert cp_plan is not None + if len(losses_reduced) != len(cp_step.assignments): + raise ValueError( + "MCore executed a different number of phases than planned" + ) + losses_reduced = [ + metric + for metric, task in zip(losses_reduced, cp_step.assignments) + if task.sample_indices and cp_plan.lane == task.lane_start + ] for x in losses_reduced: loss_metrics = {} for k in x.keys(): @@ -1252,7 +1289,9 @@ def train( mb_metrics, global_loss = aggregate_training_statistics( all_mb_metrics=all_mb_metrics, losses=losses, - data_parallel_group=parallel_state.get_data_parallel_group(), + data_parallel_group=parallel_state.get_data_parallel_group( + with_context_parallel=cp_plan is not None + ), ) metrics = { @@ -1341,12 +1380,14 @@ def get_reference_policy_logprobs( *, data: BatchedDataDict[Any], micro_batch_size: Optional[int] = None, + cp_plan: Optional[CPRankPlan] = None, ) -> BatchedDataDict[ReferenceLogprobOutputSpec]: with self.use_reference_model(): reference_logprobs = self.get_logprobs( data=data, micro_batch_size=micro_batch_size, require_router_replay=False, + cp_plan=cp_plan, ) return_data = BatchedDataDict[ReferenceLogprobOutputSpec]() @@ -2110,6 +2151,7 @@ def get_logprobs( data: BatchedDataDict[Any], micro_batch_size: Optional[int] = None, require_router_replay: bool = True, + cp_plan: Optional[CPRankPlan] = None, ) -> BatchedDataDict[LogprobOutputSpec]: """Get the logprobs of the model for a batch of data. @@ -2154,6 +2196,8 @@ def get_logprobs( delegate_pack_to_model=self.delegate_pack_to_model, delegate_mtp_loss_mask_to_model=self.delegate_mtp_loss_mask_to_model, model_slices_context_parallel_inputs=self.model_slices_context_parallel_inputs, + cp_plan=cp_plan, + cp_step=cp_plan.steps[0] if cp_plan is not None else None, ) use_fused_linear_logprobs = self.cfg["megatron_cfg"].get( @@ -2584,6 +2628,7 @@ def get_topk_logits( data: BatchedDataDict[GenerationDatumSpec], k: int, micro_batch_size: Optional[int] = None, + cp_plan: Optional[CPRankPlan] = None, ): """Get the top-k logits and indices for a batch of data. @@ -2621,6 +2666,8 @@ def get_topk_logits( delegate_pack_to_model=self.delegate_pack_to_model, delegate_mtp_loss_mask_to_model=self.delegate_mtp_loss_mask_to_model, model_slices_context_parallel_inputs=self.model_slices_context_parallel_inputs, + cp_plan=cp_plan, + cp_step=cp_plan.steps[0] if cp_plan is not None else None, ) list_of_outputs = megatron_forward_backward( diff --git a/pyrefly.toml b/pyrefly.toml index e4ffa4a3673..4d4ada9921b 100644 --- a/pyrefly.toml +++ b/pyrefly.toml @@ -24,6 +24,10 @@ replace-imports-with-any = [ "zstandard.*", ] project-includes = [ + "nemo_rl/distributed/dynamic_context_parallel.py", + "nemo_rl/distributed/tensor_serialization.py", + "nemo_rl/models/policy/dynamic_cp.py", + "nemo_rl/models/megatron/dynamic_cp.py", # TODO: enable these once we have 100 correctness #"nemo_rl/**/*.py", #"examples/**/*.py", diff --git a/tests/functional/dynamic_cp_attention_parity.py b/tests/functional/dynamic_cp_attention_parity.py new file mode 100644 index 00000000000..eadcc70d473 --- /dev/null +++ b/tests/functional/dynamic_cp_attention_parity.py @@ -0,0 +1,299 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# http://www.apache.org/licenses/LICENSE-2.0 +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Four-GPU fused-RoPE/TE transformer parity across TP and runtime CP groups.""" + +import io +import os +import pickle +from types import SimpleNamespace + +import numpy as np +import ray.cloudpickle +import torch +import torch.distributed as dist +from megatron.core import parallel_state +from megatron.core.models.gpt import GPTModel +from megatron.core.models.gpt.gpt_layer_specs import ( + get_gpt_layer_with_transformer_engine_spec, +) +from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed +from megatron.core.transformer.transformer_config import TransformerConfig +from ray.util.serialization import StandaloneSerializationContext + +from nemo_rl.algorithms.loss.loss_functions import NLLLossFn +from nemo_rl.distributed.batched_data_dict import BatchedDataDict +from nemo_rl.distributed.named_sharding import NamedSharding +from nemo_rl.distributed.tensor_serialization import ( + tensor_from_payload, + tensor_to_payload, +) +from nemo_rl.models.megatron.data import process_microbatch +from nemo_rl.models.megatron.dynamic_cp import ( + RuntimeCPContext, + initialize_dynamic_cp_runtime, + planned_microbatches, + preserve_attention_cp_groups, +) +from nemo_rl.models.megatron.train import ( + LogprobsPostProcessor, + LossPostProcessor, + model_forward, +) +from nemo_rl.models.policy.dynamic_cp import build_cp_dispatch, collect_cp_outputs +from nemo_rl.models.policy.workers.megatron_policy_worker import ( + MegatronPolicyWorkerImpl, +) +from nemo_rl.utils.timer import Timer + + +class DriverPayloadUnpickler(pickle.Unpickler): + """Check that a worker result can load in a driver without MCore.""" + + def find_class(self, module, name): + if module == "megatron" or module.startswith( + ( + "megatron.", + "nemo_rl.models.megatron.", + "nemo_rl.models.policy.workers.megatron", + ) + ): + raise AssertionError( + f"Worker result requires backend class {module}.{name}" + ) + return super().find_class(module, name) + + +def main() -> None: + # Same reducer as the real worker, without starting another Ray cluster. + StandaloneSerializationContext()._register_cloudpickle_serializer( + torch.Tensor, tensor_to_payload, tensor_from_payload + ) + torch.cuda.set_device(int(os.environ["LOCAL_RANK"])) + dist.init_process_group("nccl") + rank, world = dist.get_rank(), dist.get_world_size() + assert world == 4 + lengths = torch.tensor([7, 45, 101, 11, 55, 9]) + ids = (torch.arange(6 * 112).reshape(6, 112) % 31 + 1).long() + data = BatchedDataDict( + input_ids=ids, + input_lengths=lengths, + token_mask=(torch.arange(112)[None, :] < lengths[:, None]).long(), + sample_mask=torch.tensor([1, 1, 1, 0, 1, 1]), + ) + for tp, base_cp in ((1, 1), (1, 2), (2, 1), (2, 2)): + parallel_state.initialize_model_parallel( + tensor_model_parallel_size=tp, + context_parallel_size=base_cp, + hybrid_context_parallel=True, + ) + initialize_dynamic_cp_runtime() + torch.manual_seed(123) + model_parallel_cuda_manual_seed(123) + config = TransformerConfig( + num_layers=2, + hidden_size=256, + num_attention_heads=4, + num_query_groups=2, + ffn_hidden_size=512, + tensor_model_parallel_size=tp, + context_parallel_size=base_cp, + sequence_parallel=tp > 1, + bf16=True, + params_dtype=torch.bfloat16, + attention_dropout=0.0, + hidden_dropout=0.0, + apply_rope_fusion=True, + gradient_accumulation_fusion=False, + ) + # Exercise the same fused RoPE and TE attention path as the Qwen recipe. + model = ( + GPTModel( + config, + get_gpt_layer_with_transformer_engine_spec(), + vocab_size=128, + max_sequence_length=128, + position_embedding_type="rope", + parallel_output=True, + ) + .cuda() + .bfloat16() + ) + mesh = NamedSharding( + np.arange(world).reshape(1, world // tp // base_cp, base_cp, tp), + [ + "pipeline_parallel", + "data_parallel", + "context_parallel", + "tensor_parallel", + ], + ) + cfg = { + "logprob_batch_size": 1, + "make_sequence_length_divisible_by": 1, + "sequence_packing": {"enabled": True}, + "megatron_cfg": { + "tensor_model_parallel_size": tp, + "expert_model_parallel_size": 1, + "sequence_parallel": tp > 1, + "dynamic_context_parallel": { + "enabled": True, + "train_tokens_per_rank": 32 * tp, + "max_size": world // tp, + }, + }, + } + dispatch = build_cp_dispatch(data, cfg, mesh, batch_size=6, training=True) + lane = rank // tp + plan = dispatch.plans[lane // base_cp][lane % base_cp] + payload = dispatch.data[lane // base_cp][lane % base_cp] + step = plan.steps[0] + worker = object.__new__(MegatronPolicyWorkerImpl) + worker.model = model + worker.cfg = cfg + worker.timer = Timer() + worker.mcore_state = SimpleNamespace(straggler_timer=None) + worker.sampling_params = None + worker.defer_fp32_logits = False + worker._router_replay_enabled = False + worker.media_placeholder_token_id = None + worker.delegate_pack_to_model = False + worker.delegate_mtp_loss_mask_to_model = False + worker.model_slices_context_parallel_inputs = False + scores = worker.get_logprobs(data=payload, cp_plan=plan) + DriverPayloadUnpickler(io.BytesIO(ray.cloudpickle.dumps(scores))).load() + gathered_scores = [None] * world + dist.all_gather_object(gathered_scores, scores["logprobs"].numpy()) + restored_scores = collect_cp_outputs( + [ + BatchedDataDict(logprobs=torch.from_numpy(value)) + for value in gathered_scores[::tp] + ], + dispatch, + data.size, + )["logprobs"].cuda() + normalizers = ( + torch.tensor(step.valid_sequences, device="cuda"), + torch.tensor(step.valid_tokens, device="cuda"), + ) + reference_data = data.to("cuda") + reference_batch = process_microbatch( + reference_data, + seq_length_key="input_lengths", + pack_sequences=True, + pad_individual_seqs_to_multiple_of=tp, + cp_context=RuntimeCPContext(1, 0, None), + ) + + def forward(batch, batch_data): + return model_forward( + model=model, + data_dict=batch_data, + input_ids_cp_sharded=batch.input_ids_cp_sharded, + position_ids=batch.position_ids, + attention_mask=batch.attention_mask, + packed_seq_params=batch.packed_seq_params, + ) + + with preserve_attention_cp_groups(model): + with torch.no_grad(): + score_callback = LogprobsPostProcessor(cfg)( + reference_data, + reference_batch.input_ids, + reference_batch.packed_seq_params.cu_seqlens_q_padded, + ids.shape[1], + cp_context=RuntimeCPContext(1, 0, None), + ) + _, reference_scores = score_callback( + forward(reference_batch, reference_data) + ) + valid = reference_data["token_mask"].bool() + torch.testing.assert_close( + restored_scores[valid], + reference_scores["logprobs"][valid], + rtol=0.01, + atol=0.03, + msg=lambda msg: f"Reassembled worker scoring: {msg}", + ) + model.train() + reference_callback = LossPostProcessor(NLLLossFn(), cfg)( + reference_data, reference_batch.packed_seq_params, *normalizers + ) + reference_loss, reference_metrics = reference_callback( + forward(reference_batch, reference_data) + ) + (reference_loss * base_cp).backward() + for parameter in model.parameters(): + if parameter.grad is not None and getattr( + parameter, "sequence_parallel", False + ): + dist.all_reduce( + parameter.grad, + group=parallel_state.get_tensor_model_parallel_group(), + ) + reference_grads = { + name: p.grad.float().clone() + for name, p in model.named_parameters() + if p.grad is not None + } + model.zero_grad(set_to_none=True) + reported = torch.zeros((), device="cuda") + processor = LossPostProcessor(NLLLossFn(), cfg, len(step.assignments)) + sizes = [] + for task, batch in zip( + step.assignments, planned_microbatches(payload, plan, step, None) + ): + dist.barrier() + sizes.append(task.cp_size) + callback = processor( + batch.data_dict, batch.packed_seq_params, *normalizers + ) + loss, metrics = callback(forward(batch, batch.data_dict)) + (loss * base_cp / len(step.assignments)).backward() + if task.sample_indices and lane == task.lane_start: + reported += metrics["loss"] + domain = parallel_state.get_data_parallel_group(with_context_parallel=True) + dist.all_reduce(reported, group=domain) + torch.testing.assert_close( + reported.item(), reference_metrics["loss"], rtol=0.01, atol=0.01 + ) + squared_error = torch.zeros((), device="cuda") + squared_reference = torch.zeros((), device="cuda") + for name, parameter in model.named_parameters(): + if name in reference_grads: + gradient = parameter.grad.float() + if getattr(parameter, "sequence_parallel", False): + dist.all_reduce( + gradient, + group=parallel_state.get_tensor_model_parallel_group(), + ) + dist.all_reduce(gradient, group=domain) + squared_error += (gradient - reference_grads[name]).square().sum() + squared_reference += reference_grads[name].square().sum() + torch.testing.assert_close( + gradient, + reference_grads[name], + rtol=0.08, + atol=0.003, + msg=lambda msg: f"tp={tp}, base_cp={base_cp}, {name}: {msg}", + ) + relative_error = (squared_error / squared_reference).sqrt().item() + assert relative_error < 0.03, (tp, base_cp, relative_error) + print( + f"attention CP parity passed: rank={rank} tp={tp} base_cp={base_cp} " + f"active_sizes={sizes} gradient_relative_error={relative_error:.6f}", + flush=True, + ) + parallel_state.destroy_model_parallel() + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/tests/functional/dynamic_cp_loss_parity.py b/tests/functional/dynamic_cp_loss_parity.py new file mode 100644 index 00000000000..b590585c591 --- /dev/null +++ b/tests/functional/dynamic_cp_loss_parity.py @@ -0,0 +1,145 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# http://www.apache.org/licenses/LICENSE-2.0 +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Four-GPU parity for real NRL packing, CP logprob collectives and gradients. + +Run with the pinned Megatron environment: + torchrun --standalone --nproc-per-node=4 tests/functional/dynamic_cp_loss_parity.py +This uses a tiny trainable token-to-logit table so the assertion isolates the +data/loss integration from transformer numerical differences. Qwen GRPO tests +the actual attention model separately. +""" + +import os + +import numpy as np +import torch +import torch.distributed as dist +import torch.nn.functional as F +from megatron.core import parallel_state + +from nemo_rl.algorithms.loss.loss_functions import NLLLossFn +from nemo_rl.distributed.batched_data_dict import BatchedDataDict +from nemo_rl.distributed.named_sharding import NamedSharding +from nemo_rl.models.megatron.dynamic_cp import planned_microbatches +from nemo_rl.models.megatron.train import LossPostProcessor +from nemo_rl.models.policy.dynamic_cp import build_cp_dispatch + + +def main() -> None: + torch.cuda.set_device(int(os.environ["LOCAL_RANK"])) + dist.init_process_group("nccl") + rank, world = dist.get_rank(), dist.get_world_size() + if world != 4: + raise ValueError("Run this parity test with four ranks") + parallel_state.initialize_model_parallel( + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + context_parallel_size=1, + hybrid_context_parallel=True, + ) + lengths = torch.tensor([7, 45, 101, 11, 55, 9]) + ids = (torch.arange(6 * 112).reshape(6, 112) % 31 + 1).long() + data = BatchedDataDict( + input_ids=ids, + input_lengths=lengths, + token_mask=(torch.arange(112)[None, :] < lengths[:, None]).long(), + sample_mask=torch.tensor([1, 1, 1, 0, 1, 1]), + ) + mesh = NamedSharding( + np.arange(world).reshape(1, world, 1, 1), + ["pipeline_parallel", "data_parallel", "context_parallel", "tensor_parallel"], + ) + cfg = { + "make_sequence_length_divisible_by": 1, + "sequence_packing": {"enabled": True}, + "megatron_cfg": { + "tensor_model_parallel_size": 1, + "expert_model_parallel_size": 1, + "sequence_parallel": False, + "dynamic_context_parallel": { + "enabled": True, + "train_tokens_per_rank": 32, + "logprob_tokens_per_rank": 32, + "max_size": 4, + }, + }, + } + initial = torch.sin(torch.arange(32 * 32, device="cuda").float()).reshape(32, 32) + for base_cp in (1, 2): + # Reinitialize static groups to test active CP both below and above base CP. + if base_cp != 1: + parallel_state.destroy_model_parallel() + parallel_state.initialize_model_parallel( + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + context_parallel_size=base_cp, + hybrid_context_parallel=True, + ) + mesh = NamedSharding( + np.arange(world).reshape(1, world // base_cp, base_cp, 1), mesh.names + ) + dispatch = build_cp_dispatch(data, cfg, mesh, batch_size=6, training=True) + plan = dispatch.plans[rank // base_cp][rank % base_cp] + payload = dispatch.data[rank // base_cp][rank % base_cp] + step = plan.steps[0] + reference = initial.clone().requires_grad_() + labels = ids[:, 1:].cuda() + mask = (data["token_mask"][:, 1:] * data["sample_mask"][:, None]).cuda() + reference_loss = ( + F.cross_entropy( + F.embedding(ids[:, :-1].cuda(), reference).flatten(0, 1), + labels.flatten(), + reduction="none", + ).reshape_as(mask) + * mask + ).sum() / mask.sum() + reference_loss.backward() + weight = initial.clone().requires_grad_() + processor = LossPostProcessor( + NLLLossFn(), cfg, num_microbatches=len(step.assignments) + ) + reported = torch.zeros((), device="cuda") + sizes = [] + for task, batch in zip( + step.assignments, planned_microbatches(payload, plan, step, None) + ): + dist.barrier() + sizes.append(task.cp_size) + callback = processor( + batch.data_dict, + batch.packed_seq_params, + torch.tensor(step.valid_sequences, device="cuda"), + torch.tensor(step.valid_tokens, device="cuda"), + ) + loss, metrics = callback(F.embedding(batch.input_ids_cp_sharded, weight)) + (loss * base_cp / len(step.assignments)).backward() + if task.sample_indices and rank == task.lane_start: + reported += metrics["loss"] + dist.all_reduce(weight.grad) + dist.all_reduce(reported) + torch.testing.assert_close(reported, reference_loss, rtol=3e-5, atol=3e-6) + torch.testing.assert_close(weight.grad, reference.grad, rtol=3e-5, atol=3e-6) + torch.testing.assert_close( + initial - 0.1 * weight.grad, + initial - 0.1 * reference.grad, + rtol=3e-5, + atol=3e-6, + ) + print( + f"dynamic CP parity passed: rank={rank} base_cp={base_cp} active_sizes={sizes}", + flush=True, + ) + parallel_state.destroy_model_parallel() + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/tests/unit/distributed/test_dynamic_context_parallel.py b/tests/unit/distributed/test_dynamic_context_parallel.py new file mode 100644 index 00000000000..756d06602f3 --- /dev/null +++ b/tests/unit/distributed/test_dynamic_context_parallel.py @@ -0,0 +1,136 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# http://www.apache.org/licenses/LICENSE-2.0 +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import random + +import pytest + +from nemo_rl.distributed.dynamic_context_parallel import ( + assignment_for_lane, + cp_loss_multiplier, + plan_cp_phases, +) +from nemo_rl.distributed.named_sharding import REPLICATED_AXES, replicated_axes + + +def make_plan(lengths, *, lanes=8, minimum=1, maximum=None, budget=128, sp=2): + return plan_cp_phases( + lengths, + lanes=lanes, + min_size=minimum, + max_size=maximum or lanes, + tokens_per_rank=budget, + sequence_parallel_size=sp, + user_pad_multiple=sp, + ) + + +def check_plan(phases, lengths, lanes, budget, sp): + owners = [] + for phase in phases: + for lane in range(lanes): + task = assignment_for_lane(phase, lane) + assert task.lane_start % task.cp_size == 0 + assert task.padded_tokens % (task.cp_size * sp) == 0 + assert task.padded_tokens // task.cp_size <= budget + for task in phase: + if not task.sample_indices: + continue + owners.extend(task.sample_indices) + spans = [ + (lengths[i] + task.pad_multiple - 1) + // task.pad_multiple + * task.pad_multiple + for i in task.sample_indices + ] + assert sum(spans) == task.padded_tokens + # Reconstruct the zigzag partition and verify each real token's + # ownership, independently of padding and concatenation order. + for length, span in zip((lengths[i] for i in task.sample_indices), spans): + owned = [] + for rank in range(task.cp_size): + if task.cp_size == 1: + positions = range(span) + else: + chunk = span // (2 * task.cp_size) + positions = [ + p + for c in (rank, 2 * task.cp_size - rank - 1) + for p in range(c * chunk, (c + 1) * chunk) + ] + owned.extend(p for p in positions if p < length) + assert sorted(owned) == list(range(length)) + assert sorted(owners) == list(range(len(lengths))) + + +def test_mixed_cp_and_cross_dp_groups(): + lengths = [400, 200, 90, 70, 9, 3, 1000] + phases = make_plan(lengths) + check_plan(phases, lengths, 8, 128, 2) + assert {a.cp_size for p in phases for a in p if a.sample_indices} == {1, 2, 4, 8} + assert any(a.cp_size == 8 for p in phases for a in p) + + +@pytest.mark.parametrize("lanes", [1, 2, 4, 8, 16]) +def test_randomized_coverage_and_alignment(lanes): + rng = random.Random(42) + for _ in range(30): + lengths = [rng.randint(2, 64 * lanes) for _ in range(rng.randint(1, 40))] + phases = make_plan(lengths, lanes=lanes) + check_plan(phases, lengths, lanes, 128, 2) + assert phases == make_plan(lengths, lanes=lanes) + + +def test_padding_drives_cp_selection(): + phases = make_plan([127], budget=127) + task = next(a for p in phases for a in p if a.sample_indices) + assert task.cp_size == 2 + + +def test_moe_minimum_and_padding_only_lanes(): + phases = make_plan([3, 7], minimum=4) + assert all(a.cp_size >= 4 for p in phases for a in p) + assert any(not a.sample_indices for p in phases for a in p) + check_plan(phases, [3, 7], 8, 128, 2) + + +@pytest.mark.parametrize( + "lanes,minimum,maximum", [(3, 1, 2), (8, 3, 8), (8, 4, 2), (8, 1, 16)] +) +def test_invalid_topologies_rejected(lanes, minimum, maximum): + with pytest.raises(ValueError): + make_plan([], lanes=lanes, minimum=minimum, maximum=maximum) + + +def test_oversized_sample_rejected(): + with pytest.raises(ValueError, match="token budget"): + make_plan([1025]) + + +@pytest.mark.parametrize("base_cp", [1, 2, 4]) +@pytest.mark.parametrize("active_cp", [1, 2, 4, 8]) +@pytest.mark.parametrize("microbatches", [1, 3, 7]) +def test_gather_backward_and_schedule_scaling_cancel(base_cp, active_cp, microbatches): + factor = cp_loss_multiplier( + active_cp_size=active_cp, + schedule_cp_size=base_cp, + num_microbatches=microbatches, + replicated_cp_loss=True, + ) + # MCore scales the returned scalar; gather backward SUMs identical + # contributions. Each owned token must reach DDP SUM with unit weight. + assert factor * base_cp / microbatches * active_cp == pytest.approx(1.0) + + +def test_dynamic_outputs_are_not_static_cp_replicas(): + assert "context_parallel" in REPLICATED_AXES + assert replicated_axes() == REPLICATED_AXES + assert replicated_axes(dynamic_cp=True) == ("tensor_parallel", "pipeline_parallel") diff --git a/tests/unit/distributed/test_dynamic_cp_dispatch.py b/tests/unit/distributed/test_dynamic_cp_dispatch.py new file mode 100644 index 00000000000..5b1d0947660 --- /dev/null +++ b/tests/unit/distributed/test_dynamic_cp_dispatch.py @@ -0,0 +1,175 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# http://www.apache.org/licenses/LICENSE-2.0 +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import unittest +from random import Random + +import numpy as np +import torch + +from nemo_rl.distributed.batched_data_dict import BatchedDataDict +from nemo_rl.distributed.dynamic_context_parallel import cp_loss_multiplier +from nemo_rl.distributed.named_sharding import NamedSharding +from nemo_rl.models.policy.dynamic_cp import build_cp_dispatch, collect_cp_outputs + + +class TestDynamicCPDispatch(unittest.TestCase): + def setUp(self): + self.data = BatchedDataDict( + input_ids=torch.arange(5 * 32).reshape(5, 32), + input_lengths=torch.tensor([5, 11, 17, 25, 29]), + sample_mask=torch.tensor([1, 1, 0, 1, 1]), + sample_ids=torch.arange(5), + ) + self.data["token_mask"] = ( + torch.arange(32)[None, :] < self.data["input_lengths"][:, None] + ).long() + self.cfg = { + "make_sequence_length_divisible_by": 1, + "megatron_cfg": { + "tensor_model_parallel_size": 2, + "expert_model_parallel_size": 1, + "sequence_parallel": False, + "dynamic_context_parallel": { + "enabled": True, + "train_tokens_per_rank": 8, + "logprob_tokens_per_rank": 8, + "max_size": 4, + }, + }, + } + self.mesh = NamedSharding( + np.arange(8).reshape(1, 2, 2, 2), + [ + "pipeline_parallel", + "data_parallel", + "context_parallel", + "tensor_parallel", + ], + ) + + def test_outputs_from_nonzero_static_cp_are_preserved(self): + dispatch = build_cp_dispatch( + self.data, self.cfg, self.mesh, batch_size=None, training=False + ) + results = [] + for lane, plan in enumerate(p for dp in dispatch.plans for p in dp): + payload = dispatch.data[lane // 2][lane % 2] + values = [] + for task in plan.steps[0].assignments: + values.extend( + payload["sample_ids"][list(task.sample_indices)].tolist() + if task.sample_indices + else [-1] + ) + results.append(BatchedDataDict(logprobs=torch.tensor(values)[:, None])) + restored = collect_cp_outputs(results, dispatch, self.data.size) + torch.testing.assert_close(restored["logprobs"].flatten(), torch.arange(5)) + with self.assertRaises(ValueError): + collect_cp_outputs(results[::2], dispatch, self.data.size) + + def test_global_normalizer_and_autograd_do_not_count_cp_copies(self): + dispatch = build_cp_dispatch( + self.data, self.cfg, self.mesh, batch_size=5, training=True + ) + mask = self.data["token_mask"][:, 1:] * self.data["sample_mask"][:, None] + normalizer = mask.sum() + baseline_weight = torch.tensor(0.7, dtype=torch.float64, requires_grad=True) + inputs = self.data["input_ids"][:, 1:].double() / 100 + baseline = ( + torch.nn.functional.softplus(inputs * baseline_weight) * mask + ).sum() / normalizer + baseline.backward() + weight = baseline_weight.detach().clone().requires_grad_() + total = weight * 0 + for lane, plan in enumerate(p for dp in dispatch.plans for p in dp): + payload = dispatch.data[lane // 2][lane % 2] + step = plan.steps[0] + self.assertEqual(step.valid_tokens, normalizer.item()) + self.assertEqual(step.valid_sequences, 4) + nmb = len(step.assignments) + for task in step.assignments: + if not task.sample_indices: + continue + local = payload.select_indices(list(task.sample_indices)) + local_mask = local["token_mask"][:, 1:] * local["sample_mask"][:, None] + loss = ( + torch.nn.functional.softplus( + local["input_ids"][:, 1:].double() / 100 * weight + ) + * local_mask + ).sum() / step.valid_tokens + multiplier = cp_loss_multiplier( + active_cp_size=task.cp_size, + schedule_cp_size=2, + num_microbatches=nmb, + replicated_cp_loss=True, + ) + total = total + loss * multiplier * 2 / nmb + total.backward() + torch.testing.assert_close(total, baseline) + torch.testing.assert_close(weight.grad, baseline_weight.grad) + + def test_arbitrary_sample_permutations_restore_every_output(self): + # A five-row reverse permutation is self-inverse and cannot detect an + # accidental second argsort in BatchedDataDict.reorder_data. + rng = Random(2026) + for count in (9, 17, 64): + for static_cp in (1, 2, 4): + with self.subTest(count=count, static_cp=static_cp): + data = BatchedDataDict( + input_ids=torch.zeros(count, 32, dtype=torch.long), + input_lengths=torch.tensor( + [rng.randint(2, 32) for _ in range(count)] + ), + sample_ids=torch.arange(count), + ) + mesh = NamedSharding( + np.arange(8).reshape(1, 4 // static_cp, static_cp, 2), + self.mesh.names, + ) + dispatch = build_cp_dispatch( + data, self.cfg, mesh, batch_size=None, training=False + ) + results = [] + for payloads, plans in zip(dispatch.data, dispatch.plans): + for payload, plan in zip(payloads, plans): + values = [] + for task in plan.steps[0].assignments: + values.extend( + payload["sample_ids"][ + list(task.sample_indices) + ].tolist() + if task.sample_indices + else [-1] + ) + results.append( + BatchedDataDict(logprobs=torch.tensor(values)[:, None]) + ) + restored = collect_cp_outputs(results, dispatch, count) + torch.testing.assert_close( + restored["logprobs"].flatten(), torch.arange(count) + ) + + def test_steps_keep_separate_denominators(self): + data = BatchedDataDict.from_batches([self.data, self.data]) + data["sample_mask"][5:] = 0 + dispatch = build_cp_dispatch( + data, self.cfg, self.mesh, batch_size=5, training=True + ) + for dp in dispatch.plans: + for plan in dp: + self.assertGreater(plan.steps[0].valid_tokens, 0) + self.assertEqual(plan.steps[1].valid_tokens, 0) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/unit/distributed/test_tensor_serialization.py b/tests/unit/distributed/test_tensor_serialization.py new file mode 100644 index 00000000000..b2f3bccf925 --- /dev/null +++ b/tests/unit/distributed/test_tensor_serialization.py @@ -0,0 +1,40 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# http://www.apache.org/licenses/LICENSE-2.0 +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Portable RPC tensor payloads preserve dtypes, shapes, and writable storage.""" + +import pytest +import torch + +from nemo_rl.distributed.tensor_serialization import ( + tensor_from_payload, + tensor_to_payload, +) + + +@pytest.mark.parametrize( + "dtype", [torch.bfloat16, torch.float32, torch.int64, torch.bool] +) +@pytest.mark.parametrize("shape", [(), (0, 3), (2, 3)]) +def test_tensor_payload_round_trip(dtype, shape): + tensor = torch.ones(shape, dtype=dtype) + if tensor.ndim == 2: + tensor = tensor.T + restored = tensor_from_payload(tensor_to_payload(tensor)) + torch.testing.assert_close(restored, tensor) + restored.zero_() + assert tensor.eq(1).all() + + +def test_tensor_payload_detaches_autograd(): + tensor = torch.tensor(2.0, requires_grad=True) + restored = tensor_from_payload(tensor_to_payload(tensor)) + assert restored.item() == 2.0 + assert not restored.requires_grad diff --git a/tests/unit/models/megatron/test_dynamic_cp_scaling.py b/tests/unit/models/megatron/test_dynamic_cp_scaling.py new file mode 100644 index 00000000000..0361f939202 --- /dev/null +++ b/tests/unit/models/megatron/test_dynamic_cp_scaling.py @@ -0,0 +1,53 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# http://www.apache.org/licenses/LICENSE-2.0 +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Contract test against MCore's actual legacy loss callback scaling.""" + +from types import SimpleNamespace + +import pytest +import torch + +from nemo_rl.distributed.dynamic_context_parallel import cp_loss_multiplier + +pytestmark = pytest.mark.mcore + + +@pytest.mark.parametrize("base_cp", [1, 2]) +@pytest.mark.parametrize("active_cp", [1, 2, 4]) +@pytest.mark.parametrize("num_microbatches", [1, 3]) +@pytest.mark.parametrize("per_token", [False, True]) +def test_mcore_legacy_loss_scaling(base_cp, active_cp, num_microbatches, per_token): + # MCore is optional in the standard CPU test environment. + from megatron.core.pipeline_parallel.schedules import forward_step_calc_loss + + original_loss = torch.tensor(2.0, requires_grad=True) + multiplier = cp_loss_multiplier( + active_cp_size=active_cp, + schedule_cp_size=base_cp, + num_microbatches=num_microbatches, + replicated_cp_loss=True, + ) + metrics = [] + loss, _ = forward_step_calc_loss( + model=None, + output_tensor=original_loss, + loss_func=lambda value: (value * multiplier, {"loss": value.detach().item()}), + config=SimpleNamespace(timers=None, calculate_per_token_loss=per_token), + vp_stage=None, + collect_non_loss_data=False, + num_microbatches=num_microbatches, + forward_data_store=metrics, + cp_group_size=base_cp, + is_last_stage=True, + ) + loss.backward() + assert original_loss.grad.item() * active_cp == pytest.approx(1.0) + assert metrics == [{"loss": 2.0}] From d504ea64b9a7e3fbe51787c4fa2180bbff08a5bb Mon Sep 17 00:00:00 2001 From: Humaira Firdowse Mohammed Date: Wed, 16 Sep 2026 03:01:43 -0700 Subject: [PATCH 2/7] scheduler changes-uneven microbatches --- .../megatron-dynamic-context-parallel.md | 114 +++++-- ...n3-32b-4n4g-megatron-dynamiccp-10step.yaml | 10 + ...en3-32b-4n4g-megatron-dynamiccp-quick.yaml | 6 +- ...en3-32b-4n4g-megatron-staticcp-10step.yaml | 9 + .../distributed/dynamic_context_parallel.py | 300 ++++++++++++++++-- nemo_rl/distributed/tensor_serialization.py | 7 +- nemo_rl/models/megatron/data.py | 3 + nemo_rl/models/megatron/dynamic_cp.py | 133 ++++---- nemo_rl/models/megatron/setup.py | 8 +- nemo_rl/models/megatron/train.py | 9 +- nemo_rl/models/policy/dynamic_cp.py | 258 +++++++++++---- nemo_rl/models/policy/lm_policy.py | 40 ++- .../policy/workers/megatron_policy_worker.py | 6 +- .../functional/dynamic_cp_attention_parity.py | 8 +- tests/functional/dynamic_cp_loss_parity.py | 9 +- .../test_dynamic_context_parallel.py | 89 +++++- .../distributed/test_dynamic_cp_dispatch.py | 153 ++++++++- 17 files changed, 953 insertions(+), 209 deletions(-) create mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-10step.yaml create mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-staticcp-10step.yaml diff --git a/docs/design-docs/megatron-dynamic-context-parallel.md b/docs/design-docs/megatron-dynamic-context-parallel.md index 254a9cb29a5..c6c5d4bd88b 100644 --- a/docs/design-docs/megatron-dynamic-context-parallel.md +++ b/docs/design-docs/megatron-dynamic-context-parallel.md @@ -16,8 +16,7 @@ policy: enabled: true min_size: 1 max_size: 8 - train_tokens_per_rank: 1024 - logprob_tokens_per_rank: 1024 + tokens_per_rank: 4096 sequence_packing: enabled: true dynamic_batching: @@ -31,6 +30,20 @@ bounds must be powers of two. Budgets count padded tokens per CP rank, before tensor sequence parallelism. A sequence that cannot fit at the maximum size is rejected before dispatch. +There is one `tokens_per_rank` budget for scoring and training. Training usually +has the tighter memory limit because it retains activations, so this is the safe +schedule to share. The first policy or reference-logprob call builds the +immutable schedule. Later score calls with the same ordered lengths reuse it, +and training consumes it while recomputing its own valid-token and +valid-sequence denominators. The payload is still rebuilt at each stage because +score and train carry different fields, but `plan_cp_phases` runs only once. + +Set `tokens_per_rank` to the largest packed token count that one rank can safely +train. Setting it lower than the ordinary sequence-packing budget forces extra +CP groups and smaller model calls even when memory does not require them. The +scheduler can still increase CP for work in a partially occupied phase so idle +lanes contribute, matching the balanced hybrid-CP behavior. + This path currently supports PP=1 and the standard Ray policy data path. TransferQueue, split execution, model-owned multimodal packing, atomic preference pairs, MTP, draft training, fused linear logprobs, and training CUDA graphs are @@ -51,16 +64,33 @@ initialized DP × CP group and resolves the active attention group using MCore's hybrid-group API. TP ranks receive the same task payload. EP is a constraint on group placement, not another data-sharding axis. -Each phase covers the entire DP × CP domain. Unoccupied lanes execute a small -zero-mask placeholder, so every rank performs the same number of -forward/backward calls and reaches gradient synchronization together. A barrier -between phases keeps changes of active group coordinated, including MCore -reruns. This conservative scheduler introduces more synchronization and padding -than Megatron's balanced scheduler; it is not a throughput-equivalent port of -that scheduler. +Before finalizing an initial placement, the driver repeatedly expands the smallest real task +to the next CP power of two while unused lanes remain. Its padding factor and +padded-token count are recalculated at the larger CP size. This mirrors MCore's +`fill_empty_gpus` policy and turns idle lanes into useful attention work. If an +explicit `max_size` prevents another expansion, the remaining lanes execute a +small zero-mask placeholder. + +Adjacent placements with the same lane partition are merged into one +synchronization group. Packed tasks are redistributed between equal-size CP +subgroups using longest-processing-time placement with +`sum(sequence_length²) / active_CP` as the estimated attention cost. A subgroup +may consequently execute more packed tasks than another subgroup. Every packed +task still has its own checked token budget; merging never combines their +activations or padding allocations. + +Every lane executes at least one task per synchronization group. A zero-mask +placeholder is used only when a lane has no real task, which lets all DDP ranks +participate in the final gradient synchronization. The standard MCore PP=1 +executor runs every lane's local tasks except its last under `no_sync`; the last +local backward starts gradient synchronization. A domain-wide barrier occurs +only before the first task of each synchronization group. Faster subgroups wait +there after completing their shorter task lists, before any rank changes its CP +topology. The boundary marker is part of `ProcessedMicrobatch`, so an MCore +rerun repeats the barrier. MCore creates hybrid groups during Bridge initialization. NeMo-RL then selects -the standard no-pipeline executor for these driver-planned phases. Optimizer +the standard no-pipeline executor for these driver-planned groups. Optimizer and DDP groups remain fixed. `PackedSeqParams.local_cp_size` and `cp_group` carry the active attention topology; size one explicitly uses `cp_group=None`. The NeMo-RL runtime also initializes the TE CP stream when a model built with @@ -73,11 +103,22 @@ collection. The model receives a shallow copy of packed metadata with a real singleton group for CP=1, because RoPE interprets `None` as a static-group fallback. Loss and logprob code retain the explicit size-one/None convention. -Dynamic-CP workers register a Ray serializer for CPU tensor results. It uses -NumPy byte arrays and preserves tensor dtype and shape, including BF16. This -keeps the driver independent of MCore even when MCore replaces PyTorch's tensor -storage loader with a backend-specific function. MCore's checkpoint loader is -left intact. +During worker setup, dynamic-CP Megatron actors register a Ray serializer for +CPU tensor results. Importing MCore in the worker replaces +`torch.storage._load_from_bytes` with MCore's safe loader. A tensor does not +store a Megatron object, but PyTorch's normal storage pickle records that loader +function by module name; deserializing such a result would therefore make the +lightweight Ray driver import `megatron`. + +The serializer runs when Ray materializes an actor method's return value, after +the worker has copied result tensors to CPU and before the driver's +`get_all_worker_results` completes. It applies once to every tensor nested in +that return value. In the GRPO recipe this includes the full policy-logprob and +reference-logprob result rounds and the small loss/gradient metric tensors from +each training step. It encodes contiguous bytes, dtype, and shape in a NumPy +payload, including BF16 and empty tensors, then reconstructs an independent, +writable CPU tensor in the driver. It does not modify MCore's checkpoint loader +or serialize model parameters and GPU activations. ## Packing, outputs, and normalization @@ -97,13 +138,14 @@ discard valid results. The driver computes valid sequence and token denominators from the unique global batch before replication, separately for every optimizer step. The differentiable CP logprob gather replicates the loss over the active CP group. -For this gathered loss, NeMo-RL multiplies by +For this gathered loss, each lane's task is multiplied by ```text -number_of_phases / (static_CP * active_CP) +number_of_local_tasks / (static_CP * active_CP) ``` -This cancels the pinned no-pipeline executor's `static_CP / number_of_phases` +This cancels the pinned no-pipeline executor's +`static_CP / number_of_local_tasks` scaling and compensates the active-CP gather's backward SUM. DDP sums gradients over the fixed DP × CP domain. Metrics are retained only on task owners before global aggregation. `tests/unit/models/megatron/test_dynamic_cp_scaling.py` @@ -150,8 +192,34 @@ all ten steps with `PROFILE_STEP_RANGE=1:11`. Use `PROFILE_STEP_RANGE=3:6` for a smaller steady-state-only report. The launcher is a dry run unless `DRY_RUN=0` is explicitly supplied. -Each runtime phase has an NVTX label such as -`dynamic_cp/phase_7/cp_4/lane_2/data`. The post-run check requires ten dynamic -training plans, transitions between CP=1 and CP>1, valid training metrics, and -at least one completed policy `.nsys-rep` file on the head node. `ray.sub` also -copies reports from all nodes into `-logs/ray/**/nsight/`. +Each runtime packed task has an NVTX label such as +`dynamic_cp/group_2/task_1/cp_4/lane_2/data`. Within a group, different lanes +may have different maximum task indices. The post-run check requires ten +dynamic training plans, at least two active CP sizes, at least one group with +multiple sequential packed tasks, valid training metrics, and at least one +completed policy `.nsys-rep` file on the head node. It reports how many groups +had uneven per-lane task counts and also requires every training +dispatch to report `schedule=reused`. `ray.sub` copies reports from all nodes +into `-logs/ray/**/nsight/`. + +### Ten-step dynamic/static comparison + +`perf_runs/run_gb200_cp_comparison.sh` runs matched ten-step jobs with the same +model, batch, TP=2, PP=1, base CP=1, generation setup, container, and W&B +project. Set `CP_MODE=dynamic` to allow active CP sizes 1–8, or +`CP_MODE=static` to keep CP=1. Use different `CP_RUN_NAME` values so the W&B +runs and local logs remain distinct. + +The launcher accepts `CP_MAX_TOTAL_SEQUENCE_LENGTH`, `CP_TOKENS_PER_RANK`, +`CP_MAX_SIZE`, and `STATIC_CP_SIZE`. A fair capacity-matched comparison uses the +smallest fixed CP that can accommodate the configured maximum at the same +per-rank token budget. For example, compare dynamic CP1–2 against static CP2 +with an 8192-token maximum and a 4096-token per-rank budget. Static CP1 remains +useful as an unconstrained throughput reference when it fits in memory; static +CP8 is a capacity-matched baseline only when the workload actually requires +CP8. + +After each run, `perf_runs/analyze_cp_sequence_lengths.py` writes +`sequence_length_distribution.json` beside the driver log. It reports length +percentiles, how many samples stopped exactly at the configured ceiling, and +the CP size each sample required before optional idle-lane expansion. diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-10step.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-10step.yaml new file mode 100644 index 00000000000..641bab57116 --- /dev/null +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-10step.yaml @@ -0,0 +1,10 @@ +defaults: ./grpo-qwen3-32b-4n4g-megatron-dynamiccp-quick.yaml +grpo: + max_num_steps: 10 +logger: + log_dir: logs/grpo-qwen3-32b-4n4g-megatron-dynamiccp-10step + wandb_enabled: true + tensorboard_enabled: true + wandb: + project: nemo-rl-cp-comparison + name: qwen3-32b-4n4g-dynamic-cp-10step diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-quick.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-quick.yaml index a5a7c2c660b..009023fdd7e 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-quick.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-quick.yaml @@ -24,8 +24,10 @@ policy: enabled: true min_size: 1 max_size: 8 - train_tokens_per_rank: 1024 - logprob_tokens_per_rank: 1024 + # Match the normal packed-microbatch budget. Smaller values force long + # sequences into CP even when they fit on one GB200 rank, fragmenting one + # step into many poorly utilized model calls. + tokens_per_rank: 4096 fp8_cfg: enabled: false sequence_packing: diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-staticcp-10step.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-staticcp-10step.yaml new file mode 100644 index 00000000000..cc1494b191b --- /dev/null +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-staticcp-10step.yaml @@ -0,0 +1,9 @@ +defaults: ./grpo-qwen3-32b-4n4g-megatron-dynamiccp-10step.yaml +policy: + megatron_cfg: + dynamic_context_parallel: + enabled: false +logger: + log_dir: logs/grpo-qwen3-32b-4n4g-megatron-staticcp-10step + wandb: + name: qwen3-32b-4n4g-static-cp-10step diff --git a/nemo_rl/distributed/dynamic_context_parallel.py b/nemo_rl/distributed/dynamic_context_parallel.py index 30075180904..4897fec984c 100644 --- a/nemo_rl/distributed/dynamic_context_parallel.py +++ b/nemo_rl/distributed/dynamic_context_parallel.py @@ -10,11 +10,11 @@ # limitations under the License. """CPU-only plans for Ray-dispatched hybrid DP/CP execution. -A lane is one complete TP replica. Phases partition the DP*CP lanes into -aligned, power-of-two attention groups. All lanes execute one packed forward -per phase; an empty assignment executes masked padding. This deliberately -uses more phase boundaries than Megatron's load-balancing scheduler, while -preserving its group-transition and gradient-accumulation semantics. +A lane is one complete TP replica. Synchronization groups partition the DP*CP +lanes into aligned, power-of-two attention groups. Each attention group owns +an ordered list of independently bounded packed tasks. Different attention +groups may execute different numbers of tasks before every lane meets at the +next group boundary, matching MCore's uneven hybrid-CP execution model. """ from dataclasses import dataclass @@ -29,8 +29,7 @@ class DynamicContextParallelConfig(BaseModel, extra="forbid"): enabled: bool = False min_size: PositiveInt = 1 max_size: PositiveInt | None = None - train_tokens_per_rank: PositiveInt = 2048 - logprob_tokens_per_rank: PositiveInt = 2048 + tokens_per_rank: PositiveInt = 2048 @dataclass(frozen=True) @@ -44,14 +43,33 @@ class CPAssignment: padded_tokens: int +@dataclass(frozen=True) +class CPSyncGroup: + """A fixed CP partition with one or more sequential tasks per lane.""" + + assignments: tuple[CPAssignment, ...] + + +@dataclass(frozen=True) +class CPRankGroup: + """The ordered packed tasks one lane executes in a synchronization group.""" + + assignments: tuple[CPAssignment, ...] + + @dataclass(frozen=True) class CPRankStep: """Assignments and unique-data normalizers for one optimizer step.""" - assignments: tuple[CPAssignment, ...] + groups: tuple[CPRankGroup, ...] valid_sequences: float valid_tokens: float + @property + def assignments(self) -> tuple[CPAssignment, ...]: + """Flattened execution order, retained for metrics and compatibility.""" + return tuple(task for group in self.groups for task in group.assignments) + @dataclass(frozen=True) class CPRankPlan: @@ -62,6 +80,193 @@ class CPRankPlan: steps: tuple[CPRankStep, ...] +def _padding_for_cp( + cp_size: int, + *, + sequence_parallel_size: int, + user_pad_multiple: int, + token_alignment: int, +) -> int: + return lcm( + user_pad_multiple, + (2 * cp_size if cp_size > 1 else 1) * sequence_parallel_size * token_alignment, + ) + + +def _resize_assignment( + task: CPAssignment, + cp_size: int, + *, + lengths: list[int], + tokens_per_rank: int, + sequence_parallel_size: int, + user_pad_multiple: int, + token_alignment: int, +) -> CPAssignment: + factor = _padding_for_cp( + cp_size, + sequence_parallel_size=sequence_parallel_size, + user_pad_multiple=user_pad_multiple, + token_alignment=token_alignment, + ) + padded_tokens = sum( + (lengths[index] + factor - 1) // factor * factor + for index in task.sample_indices + ) + if padded_tokens > tokens_per_rank * cp_size: + raise ValueError("Expanded dynamic CP assignment exceeds its token budget") + return CPAssignment( + task.sample_indices, + 0, + cp_size, + factor, + padded_tokens, + ) + + +def _fill_idle_lanes( + tasks: list[CPAssignment], + *, + lanes: int, + max_size: int, + lengths: list[int], + tokens_per_rank: int, + sequence_parallel_size: int, + user_pad_multiple: int, + token_alignment: int, +) -> list[CPAssignment]: + """Increase the smallest real CP groups until no legal expansion fits.""" + idle_lanes = lanes - sum(task.cp_size for task in tasks) + while idle_lanes: + candidates = [ + (task.cp_size, index) + for index, task in enumerate(tasks) + if task.cp_size < max_size and task.cp_size <= idle_lanes + ] + if not candidates: + break + _, index = min(candidates) + task = tasks[index] + tasks[index] = _resize_assignment( + task, + task.cp_size * 2, + lengths=lengths, + tokens_per_rank=tokens_per_rank, + sequence_parallel_size=sequence_parallel_size, + user_pad_multiple=user_pad_multiple, + token_alignment=token_alignment, + ) + idle_lanes -= task.cp_size + + # Descending powers of two guarantee that each consecutive lane start is + # aligned for its group. Keep sample indices as a deterministic tie-breaker. + tasks.sort(key=lambda task: (-task.cp_size, task.sample_indices)) + placed = [] + cursor = 0 + for task in tasks: + placed.append( + CPAssignment( + task.sample_indices, + cursor, + task.cp_size, + task.pad_multiple, + task.padded_tokens, + ) + ) + cursor += task.cp_size + return placed + + +def _phase_topology(phase: tuple[CPAssignment, ...]) -> tuple[tuple[int, int], ...]: + return tuple((task.lane_start, task.cp_size) for task in phase) + + +def _assignment_work(task: CPAssignment, lengths: list[int]) -> float: + """Estimate packed attention work per participating CP rank.""" + return ( + sum(float(lengths[index] ** 2) for index in task.sample_indices) / task.cp_size + ) + + +def _merge_compatible_phases( + phases: list[tuple[CPAssignment, ...]], lengths: list[int] +) -> tuple[CPSyncGroup, ...]: + """Merge adjacent equal partitions and balance their tasks across subgroups. + + A topology change still requires a domain-wide barrier. Equal topologies do + not: their independently bounded packed tasks can run sequentially inside + the same fixed CP subgroups. LPT placement minimizes the longest estimated + attention workload and naturally produces uneven task counts. + """ + merged: list[CPSyncGroup] = [] + cursor = 0 + while cursor < len(phases): + topology = _phase_topology(phases[cursor]) + end = cursor + 1 + while end < len(phases) and _phase_topology(phases[end]) == topology: + end += 1 + + slots_by_size: dict[int, list[int]] = {} + for lane_start, cp_size in topology: + slots_by_size.setdefault(cp_size, []).append(lane_start) + assigned: dict[tuple[int, int], list[CPAssignment]] = { + slot: [] for slot in topology + } + loads = {slot: 0.0 for slot in topology} + + real_tasks = [ + task + for phase in phases[cursor:end] + for task in phase + if task.sample_indices + ] + for task in sorted( + real_tasks, + key=lambda item: (-_assignment_work(item, lengths), item.sample_indices), + ): + candidates = [ + (loads[(lane_start, task.cp_size)], lane_start) + for lane_start in slots_by_size[task.cp_size] + ] + _, lane_start = min(candidates) + slot = (lane_start, task.cp_size) + placed = CPAssignment( + task.sample_indices, + lane_start, + task.cp_size, + task.pad_multiple, + task.padded_tokens, + ) + assigned[slot].append(placed) + loads[slot] += _assignment_work(placed, lengths) + + group_assignments: list[CPAssignment] = [] + for phase_task in phases[cursor]: + slot = (phase_task.lane_start, phase_task.cp_size) + tasks = assigned[slot] + if tasks: + group_assignments.extend(tasks) + else: + # Every lane needs at least one local call so the final backward + # on every DDP rank can participate in the same gradient sync. + group_assignments.append( + CPAssignment( + (), + phase_task.lane_start, + phase_task.cp_size, + phase_task.pad_multiple, + ( + (2 + phase_task.pad_multiple - 1) + // phase_task.pad_multiple + * phase_task.pad_multiple + ), + ) + ) + merged.append(CPSyncGroup(tuple(group_assignments))) + cursor = end + return tuple(merged) + + def plan_cp_phases( lengths: list[int], *, @@ -72,11 +277,14 @@ def plan_cp_phases( sequence_parallel_size: int, user_pad_multiple: int, token_alignment: int = 1, -) -> tuple[tuple[CPAssignment, ...], ...]: - """Pack equal-CP sequences and place tasks into disjoint aligned groups. +) -> tuple[CPSyncGroup, ...]: + """Pack sequences and form MCore-style synchronization groups. Sample indices refer to the original batch, never to a sorted copy. Padding is included when choosing CP size and admitting samples to bins. + Every packed task remains within its per-rank token budget. Adjacent phases + with the same CP partition are merged and balanced across their subgroups, + allowing different lanes to execute different sequential task counts. """ if lanes < 1 or lanes & (lanes - 1): raise ValueError("Dynamic CP requires a power-of-two DP*CP domain") @@ -95,11 +303,11 @@ def plan_cp_phases( raise ValueError("Dynamic CP requires at least two input tokens per sample") size = min_size while True: - multiple = lcm( - user_pad_multiple, - (2 * size if size > 1 else 1) - * sequence_parallel_size - * token_alignment, + multiple = _padding_for_cp( + size, + sequence_parallel_size=sequence_parallel_size, + user_pad_multiple=user_pad_multiple, + token_alignment=token_alignment, ) padded = (length + multiple - 1) // multiple * multiple if padded <= tokens_per_rank * size: @@ -125,29 +333,32 @@ def plan_cp_phases( ] phases = [] while pending: - phase = [] + selected = [] remaining = [] cursor = 0 for task in pending: if cursor + task.cp_size <= lanes: - phase.append( - CPAssignment( - task.sample_indices, - cursor, - task.cp_size, - task.pad_multiple, - task.padded_tokens, - ) - ) + selected.append(task) cursor += task.cp_size else: remaining.append(task) + phase = _fill_idle_lanes( + selected, + lanes=lanes, + max_size=max_size, + lengths=lengths, + tokens_per_rank=tokens_per_rank, + sequence_parallel_size=sequence_parallel_size, + user_pad_multiple=user_pad_multiple, + token_alignment=token_alignment, + ) + cursor = sum(task.cp_size for task in phase) while cursor < lanes: - factor = lcm( - user_pad_multiple, - (2 * min_size if min_size > 1 else 1) - * sequence_parallel_size - * token_alignment, + factor = _padding_for_cp( + min_size, + sequence_parallel_size=sequence_parallel_size, + user_pad_multiple=user_pad_multiple, + token_alignment=token_alignment, ) phase.append( CPAssignment( @@ -157,14 +368,33 @@ def plan_cp_phases( cursor += min_size phases.append(tuple(phase)) pending = remaining - return tuple(phases) + return _merge_compatible_phases(phases, lengths) + + +def assignments_for_lane(group: CPSyncGroup, lane: int) -> tuple[CPAssignment, ...]: + """Resolve a lane's ordered tasks in one synchronization group.""" + matches = tuple( + task + for task in group.assignments + if task.lane_start <= lane < task.lane_start + task.cp_size + ) + if not matches: + raise ValueError(f"Lane {lane} has no assignment in a CP synchronization group") + topology = {(task.lane_start, task.cp_size) for task in matches} + if len(topology) != 1: + raise ValueError( + f"Lane {lane} crosses {len(topology)} CP subgroups before synchronization" + ) + return matches -def assignment_for_lane(phase: tuple[CPAssignment, ...], lane: int) -> CPAssignment: - """Resolve exactly one task for a lane, including padding tasks.""" - matches = [a for a in phase if a.lane_start <= lane < a.lane_start + a.cp_size] +def assignment_for_lane(group: CPSyncGroup, lane: int) -> CPAssignment: + """Resolve one-task groups; callers needing uneven work use assignments_for_lane.""" + matches = assignments_for_lane(group, lane) if len(matches) != 1: - raise ValueError(f"Lane {lane} has {len(matches)} assignments in a CP phase") + raise ValueError( + f"Lane {lane} has {len(matches)} sequential assignments in a CP group" + ) return matches[0] diff --git a/nemo_rl/distributed/tensor_serialization.py b/nemo_rl/distributed/tensor_serialization.py index b6a9e5a0c93..e6d2fd1a26c 100644 --- a/nemo_rl/distributed/tensor_serialization.py +++ b/nemo_rl/distributed/tensor_serialization.py @@ -37,7 +37,12 @@ def tensor_from_payload(payload: TensorPayload) -> torch.Tensor: def register_policy_tensor_serializer() -> None: - """Register the policy worker's Ray serializer, retaining MCore's loader.""" + """Register serialization for actor return tensors in this worker process. + + Ray invokes this after an actor method returns and before the result enters + the object store. Registration is process-local and does not change the + driver's serializer or MCore's checkpoint loading behavior. + """ # Ray is optional in tensor packing/round-trip unit tests. import ray.util diff --git a/nemo_rl/models/megatron/data.py b/nemo_rl/models/megatron/data.py index dddd076708b..5083463f025 100644 --- a/nemo_rl/models/megatron/data.py +++ b/nemo_rl/models/megatron/data.py @@ -105,6 +105,9 @@ class ProcessedMicrobatch: routed_experts_cp_sharded: Optional[torch.Tensor] = None original_seq_length: Optional[int] = None media_token_validity_mask: Optional[torch.Tensor] = None + dynamic_cp_group_start: bool = False + dynamic_cp_group_index: Optional[int] = None + dynamic_cp_task_index: Optional[int] = None def make_processed_microbatch_iterator( diff --git a/nemo_rl/models/megatron/dynamic_cp.py b/nemo_rl/models/megatron/dynamic_cp.py index 7d3a54ed5f7..da2f13f02a5 100644 --- a/nemo_rl/models/megatron/dynamic_cp.py +++ b/nemo_rl/models/megatron/dynamic_cp.py @@ -91,7 +91,7 @@ def bind_attention_cp_group(model: torch.nn.Module, packed_seq_params: Any) -> A def planned_microbatches( data: Any, plan: CPRankPlan, step: CPRankStep, straggler_timer: Any ) -> Iterator[Any]: - """Yield exactly the driver's phases, binding and validating active groups.""" + """Yield the lane's uneven task list with explicit group boundaries.""" # Avoid a cycle: data.py dispatches to this iterator. from nemo_rl.models.megatron.data import ProcessedMicrobatch, process_microbatch @@ -103,70 +103,79 @@ def planned_microbatches( or domain.rank() != plan.lane ): raise ValueError("Ray's DP*CP lane map disagrees with initialized MCore groups") - for phase_index, assignment in enumerate(step.assignments): - size = assignment.cp_size - group = ( - parallel_state.get_hybrid_data_context_parallel_groups(group_size=size) - if size > 1 - else None - ) - rank = plan.lane - assignment.lane_start - if group is not None: - members = expected[assignment.lane_start : assignment.lane_start + size] - if ( - group.size() != size - or group.rank() != rank - or torch.distributed.get_process_group_ranks(group) != members - ): + for group_index, rank_group in enumerate(step.groups): + if not rank_group.assignments: + raise ValueError("Every CP synchronization group needs one local task") + for task_index, assignment in enumerate(rank_group.assignments): + size = assignment.cp_size + group = ( + parallel_state.get_hybrid_data_context_parallel_groups(group_size=size) + if size > 1 + else None + ) + rank = plan.lane - assignment.lane_start + if group is not None: + members = expected[assignment.lane_start : assignment.lane_start + size] + if ( + group.size() != size + or group.rank() != rank + or torch.distributed.get_process_group_ranks(group) != members + ): + raise ValueError( + "Active CP group disagrees with the driver's assignment" + ) + expert_group = parallel_state.get_expert_model_parallel_group() + tp_size = parallel_state.get_tensor_model_parallel_world_size() + task_ranks = { + base + offset + for base in plan.lane_ranks[ + assignment.lane_start : assignment.lane_start + size + ] + for offset in range(tp_size) + } + if not set( + torch.distributed.get_process_group_ranks(expert_group) + ).issubset(task_ranks): raise ValueError( - "Active CP group disagrees with the driver's assignment" + "Expert communication group crosses dynamic CP task boundaries" ) - expert_group = parallel_state.get_expert_model_parallel_group() - tp_size = parallel_state.get_tensor_model_parallel_world_size() - task_ranks = { - base + offset - for base in plan.lane_ranks[ - assignment.lane_start : assignment.lane_start + size - ] - for offset in range(tp_size) - } - if not set(torch.distributed.get_process_group_ranks(expert_group)).issubset( - task_ranks - ): - raise ValueError( - "Expert communication group crosses dynamic CP task boundaries" - ) - context = RuntimeCPContext(size=size, rank=rank, group=group) - if assignment.sample_indices: - batch = data.select_indices(list(assignment.sample_indices)).to("cuda") - else: - batch = data.select_indices([0]).to("cuda") - # A real attention invocation keeps collective counts aligned, but - # none of this placeholder's targets or metrics belong to the batch. - for key, value in list(batch.items()): - if isinstance(value, torch.Tensor): - batch[key] = torch.zeros_like(value) - batch["input_lengths"].fill_(2) - inputs = process_microbatch( - batch, - seq_length_key="input_lengths", - pack_sequences=True, - pad_individual_seqs_to_multiple_of=assignment.pad_multiple, - straggler_timer=straggler_timer, - cp_context=context, - ) - if inputs.input_ids_cp_sharded.shape[1] * size != assignment.padded_tokens: - raise ValueError( - "Packed worker token count disagrees with the driver's plan" + context = RuntimeCPContext(size=size, rank=rank, group=group) + if assignment.sample_indices: + batch = data.select_indices(list(assignment.sample_indices)).to("cuda") + else: + batch = data.select_indices([0]).to("cuda") + # A real attention invocation keeps collective counts aligned, but + # none of this placeholder's targets or metrics belong to the batch. + for key, value in list(batch.items()): + if isinstance(value, torch.Tensor): + batch[key] = torch.zeros_like(value) + batch["input_lengths"].fill_(2) + inputs = process_microbatch( + batch, + seq_length_key="input_lengths", + pack_sequences=True, + pad_individual_seqs_to_multiple_of=assignment.pad_multiple, + straggler_timer=straggler_timer, + cp_context=context, ) - payload_kind = "data" if assignment.sample_indices else "padding" - # The generator stays paused inside this range while MCore consumes the - # microbatch, so an Nsight trace shows the active CP size for the whole - # forward/backward phase. - with torch.cuda.nvtx.range( - f"dynamic_cp/phase_{phase_index}/cp_{size}/lane_{plan.lane}/{payload_kind}" - ): - yield ProcessedMicrobatch(data_dict=batch, **vars(inputs)) + if inputs.input_ids_cp_sharded.shape[1] * size != assignment.padded_tokens: + raise ValueError( + "Packed worker token count disagrees with the driver's plan" + ) + payload_kind = "data" if assignment.sample_indices else "padding" + # The generator stays paused inside this range while MCore consumes + # the microbatch, making uneven task counts visible in Nsight. + with torch.cuda.nvtx.range( + f"dynamic_cp/group_{group_index}/task_{task_index}/cp_{size}/" + f"lane_{plan.lane}/{payload_kind}" + ): + yield ProcessedMicrobatch( + data_dict=batch, + dynamic_cp_group_start=task_index == 0, + dynamic_cp_group_index=group_index, + dynamic_cp_task_index=task_index, + **vars(inputs), + ) def runtime_cp_from_packed(packed_seq_params: Any) -> RuntimeCPContext: diff --git a/nemo_rl/models/megatron/setup.py b/nemo_rl/models/megatron/setup.py index 85cf18e7181..1952aff853c 100644 --- a/nemo_rl/models/megatron/setup.py +++ b/nemo_rl/models/megatron/setup.py @@ -2121,9 +2121,11 @@ def setup_model_and_optimizer( and megatron_cfg.model.hybrid_context_parallel ): initialize_dynamic_cp_runtime() - # The Ray driver supplies already scheduled, uniform phases. Use MCore's - # standard no-pipeline executor, not its TP0 dataloader/redistribution loop. - # Hybrid process groups remain initialized; PackedSeqParams selects them. + # The Ray driver supplies packed-task synchronization groups and each + # lane's possibly uneven task list. Use MCore's standard PP=1 executor, + # whose no_sync handling covers every local task except the last, rather + # than its TP0 dataloader/redistribution path. Hybrid process groups stay + # initialized; PackedSeqParams selects the active attention group. megatron_cfg.model.hybrid_context_parallel = False if megatron_cfg.ft and megatron_cfg.ft.enable_ft_package: diff --git a/nemo_rl/models/megatron/train.py b/nemo_rl/models/megatron/train.py index f78863bb3a5..83b4914ea0c 100644 --- a/nemo_rl/models/megatron/train.py +++ b/nemo_rl/models/megatron/train.py @@ -336,11 +336,10 @@ def forward_with_post_processing_fn( and isinstance(getattr(packed_seq_params, "local_cp_size", None), int) else None ) - if packed_seq_params is not None and isinstance( - getattr(packed_seq_params, "local_cp_size", None), int - ): - # Every phase has one forward/backward per lane. Keep the barrier here, - # rather than in the iterator, so MCore reruns replay it as well. + if processed_mb.dynamic_cp_group_start: + # Lanes may run different numbers of packed tasks in a fixed CP + # partition. They meet only before changing to the next partition. + # Keep this in the forward path so MCore reruns replay the barrier. torch.distributed.barrier( group=parallel_state.get_data_parallel_group(with_context_parallel=True) ) diff --git a/nemo_rl/models/policy/dynamic_cp.py b/nemo_rl/models/policy/dynamic_cp.py index 1734e4ce288..e0c694eb2df 100644 --- a/nemo_rl/models/policy/dynamic_cp.py +++ b/nemo_rl/models/policy/dynamic_cp.py @@ -17,10 +17,12 @@ from nemo_rl.distributed.batched_data_dict import BatchedDataDict from nemo_rl.distributed.dynamic_context_parallel import ( + CPRankGroup, CPRankPlan, CPRankStep, + CPSyncGroup, DynamicContextParallelConfig, - assignment_for_lane, + assignments_for_lane, plan_cp_phases, ) from nemo_rl.distributed.named_sharding import NamedSharding @@ -70,7 +72,7 @@ def validate_dynamic_cp(cfg: dict[str, Any], *, lanes: int) -> None: lanes=lanes, min_size=minimum, max_size=dynamic.max_size or lanes, - tokens_per_rank=dynamic.train_tokens_per_rank, + tokens_per_rank=dynamic.tokens_per_rank, sequence_parallel_size=tp if mc["sequence_parallel"] else 1, user_pad_multiple=cfg["make_sequence_length_divisible_by"], ) @@ -83,6 +85,146 @@ class CPDispatch: data: list[list[BatchedDataDict]] plans: list[list[CPRankPlan]] output_rows: list[tuple[list[int], list[int]]] + schedule: "CPBatchSchedule" + + +@dataclass(frozen=True) +class CPBatchSchedule: + """Immutable sample grouping shared by score and train dispatches.""" + + input_lengths: tuple[int, ...] + batch_size: int + lanes: int + min_size: int + max_size: int + tokens_per_rank: int + sequence_parallel_size: int + user_pad_multiple: int + token_alignment: int + groups_by_batch: tuple[tuple[CPSyncGroup, ...], ...] + + +def _schedule_parameters( + cfg: dict[str, Any], sharding: NamedSharding +) -> tuple[int, int, int, int, int, int, int]: + dynamic = dynamic_cp_config(cfg) + if dynamic is None: + raise ValueError("Dynamic CP scheduling requires enabled configuration") + mc = cfg["megatron_cfg"] + cp = sharding.shape["context_parallel"] + lanes = sharding.shape["data_parallel"] * cp + tp = mc["tensor_model_parallel_size"] + minimum = dynamic.min_size + while minimum * tp < mc["expert_model_parallel_size"]: + minimum *= 2 + fp8 = mc.get("fp8_cfg") or {} + alignment = 1 + if fp8.get("enabled"): + alignment = {"blockwise": 128, "mxfp8": 32}.get(fp8["fp8_recipe"], 16) + if ( + mc.get("moe_token_dispatcher_type") == "flex" + and mc.get("moe_flex_dispatcher_backend") == "hybridep" + ): + alignment = max(alignment, 128) + return ( + lanes, + minimum, + dynamic.max_size or lanes, + dynamic.tokens_per_rank, + tp if mc["sequence_parallel"] else 1, + cfg["make_sequence_length_divisible_by"], + alignment, + ) + + +def _schedule_batch_size(data: BatchedDataDict, batch_size: int | None) -> int: + gbs = batch_size if batch_size is not None else data.size + if gbs < 1 or data.size % gbs: + raise ValueError("Dynamic CP data must contain complete schedule batches") + return gbs + + +def _input_lengths(data: BatchedDataDict) -> tuple[int, ...]: + values = data["input_lengths"] + if hasattr(values, "tolist"): + values = values.tolist() + return tuple(int(value) for value in values) + + +def build_cp_schedule( + data: BatchedDataDict, + cfg: dict[str, Any], + sharding: NamedSharding, + *, + batch_size: int | None, +) -> CPBatchSchedule: + """Plan sample groups once, using the training-safe token budget.""" + gbs = _schedule_batch_size(data, batch_size) + lengths = _input_lengths(data) + ( + lanes, + minimum, + maximum, + tokens_per_rank, + sequence_parallel_size, + user_pad_multiple, + alignment, + ) = _schedule_parameters(cfg, sharding) + groups_by_batch = tuple( + plan_cp_phases( + list(lengths[start : start + gbs]), + lanes=lanes, + min_size=minimum, + max_size=maximum, + tokens_per_rank=tokens_per_rank, + sequence_parallel_size=sequence_parallel_size, + user_pad_multiple=user_pad_multiple, + token_alignment=alignment, + ) + for start in range(0, len(lengths), gbs) + ) + return CPBatchSchedule( + lengths, + gbs, + lanes, + minimum, + maximum, + tokens_per_rank, + sequence_parallel_size, + user_pad_multiple, + alignment, + groups_by_batch, + ) + + +def cp_schedule_matches( + schedule: CPBatchSchedule, + data: BatchedDataDict, + cfg: dict[str, Any], + sharding: NamedSharding, + *, + batch_size: int | None, +) -> bool: + """Return whether a cached schedule is valid for this ordered batch.""" + try: + gbs = _schedule_batch_size(data, batch_size) + parameters = _schedule_parameters(cfg, sharding) + except (KeyError, TypeError, ValueError): + return False + return ( + schedule.input_lengths == _input_lengths(data) + and schedule.batch_size == gbs + and ( + schedule.lanes, + schedule.min_size, + schedule.max_size, + schedule.tokens_per_rank, + schedule.sequence_parallel_size, + schedule.user_pad_multiple, + schedule.token_alignment, + ) + == parameters + ) def build_cp_dispatch( @@ -92,6 +234,7 @@ def build_cp_dispatch( *, batch_size: int | None, training: bool, + schedule: CPBatchSchedule | None = None, ) -> CPDispatch: """Plan before replication; count each original training target once.""" dynamic = dynamic_cp_config(cfg) @@ -102,11 +245,6 @@ def build_cp_dispatch( cp = sharding.shape["context_parallel"] dp = sharding.shape["data_parallel"] lanes = cp * dp - mc = cfg["megatron_cfg"] - tp = mc["tensor_model_parallel_size"] - minimum = dynamic.min_size - while minimum * tp < mc["expert_model_parallel_size"]: - minimum *= 2 group_ranks = tuple( sharding.get_ranks_by_coord( pipeline_parallel=0, @@ -116,36 +254,16 @@ def build_cp_dispatch( )[0] for i in range(lanes) ) - budget = ( - dynamic.train_tokens_per_rank if training else dynamic.logprob_tokens_per_rank - ) - fp8 = mc.get("fp8_cfg") or {} - alignment = 1 - if fp8.get("enabled"): - alignment = {"blockwise": 128, "mxfp8": 32}.get(fp8["fp8_recipe"], 16) - if ( - mc.get("moe_token_dispatcher_type") == "flex" - and mc.get("moe_flex_dispatcher_backend") == "hybridep" - ): - alignment = max(alignment, 128) - gbs = batch_size if batch_size is not None else data.size - if gbs < 1 or data.size % gbs: - raise ValueError( - "Dynamic CP training data must contain complete global batches" - ) + gbs = _schedule_batch_size(data, batch_size) + reused_schedule = schedule is not None + if schedule is None: + schedule = build_cp_schedule(data, cfg, sharding, batch_size=batch_size) + elif not cp_schedule_matches(schedule, data, cfg, sharding, batch_size=batch_size): + raise ValueError("Cached dynamic CP schedule does not match this ordered batch") rank_steps: list[list[CPRankStep]] = [[] for _ in range(lanes)] - for start in range(0, data.size, gbs): + for batch_index, start in enumerate(range(0, data.size, gbs)): batch = data.select_indices(list(range(start, start + gbs))) - phases = plan_cp_phases( - batch["input_lengths"].tolist(), - lanes=lanes, - min_size=minimum, - max_size=dynamic.max_size or lanes, - tokens_per_rank=budget, - sequence_parallel_size=tp if mc["sequence_parallel"] else 1, - user_pad_multiple=cfg["make_sequence_length_divisible_by"], - token_alignment=alignment, - ) + groups = schedule.groups_by_batch[batch_index] if training: if "sample_mask" not in batch or "token_mask" not in batch: raise ValueError( @@ -160,31 +278,55 @@ def build_cp_dispatch( else: valid_sequences = valid_tokens = 0.0 samples_by_cp = Counter() - for phase in phases: - for task in phase: + for group in groups: + for task in group.assignments: samples_by_cp[task.cp_size] += len(task.sample_indices) + group_task_ranges = [ + ( + min(len(assignments_for_lane(group, lane)) for lane in range(lanes)), + max(len(assignments_for_lane(group, lane)) for lane in range(lanes)), + ) + for group in groups + ] logger.info( - "Dynamic CP %s: samples=%d phases=%d samples_by_cp=%s valid_sequences=%s valid_tokens=%s", + "Dynamic CP %s: samples=%d groups=%d uneven_groups=%d " + "local_tasks=[%d,%d] " + "samples_by_cp=%s " + "valid_sequences=%s valid_tokens=%s schedule=%s", "train" if training else "score", gbs, - len(phases), + len(groups), + sum(low < high for low, high in group_task_ranges), + min( + sum(len(assignments_for_lane(group, lane)) for group in groups) + for lane in range(lanes) + ), + max( + sum(len(assignments_for_lane(group, lane)) for group in groups) + for lane in range(lanes) + ), dict(sorted(samples_by_cp.items())), valid_sequences, valid_tokens, + "reused" if reused_schedule else "new", ) for lane in range(lanes): - assignments = tuple( - replace( - assignment_for_lane(phase, lane), - sample_indices=tuple( - i + start - for i in assignment_for_lane(phase, lane).sample_indices - ), + rank_groups = tuple( + CPRankGroup( + tuple( + replace( + assignment, + sample_indices=tuple( + i + start for i in assignment.sample_indices + ), + ) + for assignment in assignments_for_lane(group, lane) + ) ) - for phase in phases + for group in groups ) rank_steps[lane].append( - CPRankStep(assignments, valid_sequences, valid_tokens) + CPRankStep(rank_groups, valid_sequences, valid_tokens) ) payloads, plans, output_rows = [], [], [] @@ -204,14 +346,19 @@ def build_cp_dispatch( remapped_steps = tuple( replace( step, - assignments=tuple( - replace( - task, - sample_indices=tuple( - local_indices[i] for i in task.sample_indices - ), + groups=tuple( + CPRankGroup( + tuple( + replace( + task, + sample_indices=tuple( + local_indices[i] for i in task.sample_indices + ), + ) + for task in group.assignments + ) ) - for task in step.assignments + for group in step.groups ), ) for step in steps @@ -230,6 +377,7 @@ def build_cp_dispatch( data=[payloads[i : i + cp] for i in range(0, lanes, cp)], plans=[plans[i : i + cp] for i in range(0, lanes, cp)], output_rows=output_rows, + schedule=schedule, ) diff --git a/nemo_rl/models/policy/lm_policy.py b/nemo_rl/models/policy/lm_policy.py index 911c07c3c07..48bfc467bf5 100644 --- a/nemo_rl/models/policy/lm_policy.py +++ b/nemo_rl/models/policy/lm_policy.py @@ -42,8 +42,10 @@ ) from nemo_rl.models.policy import PolicyConfig from nemo_rl.models.policy.dynamic_cp import ( + CPBatchSchedule, build_cp_dispatch, collect_cp_outputs, + cp_schedule_matches, dynamic_cp_config, validate_dynamic_cp, ) @@ -339,6 +341,7 @@ def __init__( ) self.dynamic_cp = dynamic_cp_config(config) is not None + self._dynamic_cp_schedule: Optional[CPBatchSchedule] = None if self.dynamic_cp: if not self.supports_dynamic_cp_dispatch: raise ValueError( @@ -670,9 +673,22 @@ def _report_sharded_payload( def _get_dynamic_cp_outputs( self, method: str, data: BatchedDataDict, **kwargs: Any ) -> BatchedDataDict: + schedule_batch_size = self.cfg["train_global_batch_size"] + if data.size != schedule_batch_size: + # The score worker currently consumes one CPRankStep. Standalone + # scoring can therefore use dynamic CP as one schedule batch, but + # it will not reuse a multi-global-batch training schedule. + schedule_batch_size = data.size + schedule = self._matching_dynamic_cp_schedule(data, schedule_batch_size) dispatch = build_cp_dispatch( - data, self.cfg, self.sharding_annotations, batch_size=None, training=False + data, + self.cfg, + self.sharding_annotations, + batch_size=schedule_batch_size, + training=False, + schedule=schedule, ) + self._dynamic_cp_schedule = dispatch.schedule futures = self.worker_group.run_all_workers_sharded_data( method, data=dispatch.data, @@ -686,6 +702,22 @@ def _get_dynamic_cp_outputs( self.worker_group.get_all_worker_results(futures), dispatch, data.size ) + def _matching_dynamic_cp_schedule( + self, data: BatchedDataDict, batch_size: int + ) -> Optional[CPBatchSchedule]: + schedule = self._dynamic_cp_schedule + if schedule is None: + return None + if cp_schedule_matches( + schedule, + data, + self.cfg, + self.sharding_annotations, + batch_size=batch_size, + ): + return schedule + return None + def get_logprobs( self, data: BatchedDataDict[GenerationDatumSpec], @@ -934,10 +966,16 @@ def train( self.sharding_annotations, batch_size=batch_size, training=True, + schedule=self._matching_dynamic_cp_schedule(data, batch_size), ) if self.dynamic_cp else None ) + if dispatch is not None: + # A schedule is step-local. Training is the final consumer, so + # release it even though the immutable topology is also stored + # in dispatch for the duration of this call. + self._dynamic_cp_schedule = None sharded_data = ( dispatch.data if dispatch is not None diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index 8d47a7b6c35..601f4b68c7f 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -789,8 +789,12 @@ def __init__( if _model_accepts_media_token_validity_mask(self.model) else None ) + # MCore installs a Megatron-specific torch storage loader process-wide. + # Every Megatron policy result therefore needs the portable CPU-tensor + # serializer, even when dynamic context parallelism is disabled; the + # Ray driver intentionally does not have Megatron on its import path. + register_policy_tensor_serializer() if dynamic_cp_config(self.cfg) is not None: - register_policy_tensor_serializer() if self.delegate_pack_to_model or self.model_slices_context_parallel_inputs: raise ValueError("Dynamic CP requires NeMo-owned text-model packing") if getattr(self._get_model_config(), "mtp_num_layers", None): diff --git a/tests/functional/dynamic_cp_attention_parity.py b/tests/functional/dynamic_cp_attention_parity.py index eadcc70d473..17597140811 100644 --- a/tests/functional/dynamic_cp_attention_parity.py +++ b/tests/functional/dynamic_cp_attention_parity.py @@ -145,7 +145,7 @@ def main() -> None: "sequence_parallel": tp > 1, "dynamic_context_parallel": { "enabled": True, - "train_tokens_per_rank": 32 * tp, + "tokens_per_rank": 32 * tp, "max_size": world // tp, }, }, @@ -247,10 +247,13 @@ def forward(batch, batch_data): reported = torch.zeros((), device="cuda") processor = LossPostProcessor(NLLLossFn(), cfg, len(step.assignments)) sizes = [] + group_boundaries = 0 for task, batch in zip( step.assignments, planned_microbatches(payload, plan, step, None) ): - dist.barrier() + if batch.dynamic_cp_group_start: + group_boundaries += 1 + dist.barrier() sizes.append(task.cp_size) callback = processor( batch.data_dict, batch.packed_seq_params, *normalizers @@ -259,6 +262,7 @@ def forward(batch, batch_data): (loss * base_cp / len(step.assignments)).backward() if task.sample_indices and lane == task.lane_start: reported += metrics["loss"] + assert group_boundaries == len(step.groups) domain = parallel_state.get_data_parallel_group(with_context_parallel=True) dist.all_reduce(reported, group=domain) torch.testing.assert_close( diff --git a/tests/functional/dynamic_cp_loss_parity.py b/tests/functional/dynamic_cp_loss_parity.py index b590585c591..983a50b9c8c 100644 --- a/tests/functional/dynamic_cp_loss_parity.py +++ b/tests/functional/dynamic_cp_loss_parity.py @@ -66,8 +66,7 @@ def main() -> None: "sequence_parallel": False, "dynamic_context_parallel": { "enabled": True, - "train_tokens_per_rank": 32, - "logprob_tokens_per_rank": 32, + "tokens_per_rank": 32, "max_size": 4, }, }, @@ -108,10 +107,13 @@ def main() -> None: ) reported = torch.zeros((), device="cuda") sizes = [] + group_boundaries = 0 for task, batch in zip( step.assignments, planned_microbatches(payload, plan, step, None) ): - dist.barrier() + if batch.dynamic_cp_group_start: + group_boundaries += 1 + dist.barrier() sizes.append(task.cp_size) callback = processor( batch.data_dict, @@ -123,6 +125,7 @@ def main() -> None: (loss * base_cp / len(step.assignments)).backward() if task.sample_indices and rank == task.lane_start: reported += metrics["loss"] + assert group_boundaries == len(step.groups) dist.all_reduce(weight.grad) dist.all_reduce(reported) torch.testing.assert_close(reported, reference_loss, rtol=3e-5, atol=3e-6) diff --git a/tests/unit/distributed/test_dynamic_context_parallel.py b/tests/unit/distributed/test_dynamic_context_parallel.py index 756d06602f3..949bfbdfdd3 100644 --- a/tests/unit/distributed/test_dynamic_context_parallel.py +++ b/tests/unit/distributed/test_dynamic_context_parallel.py @@ -14,7 +14,9 @@ import pytest from nemo_rl.distributed.dynamic_context_parallel import ( + DynamicContextParallelConfig, assignment_for_lane, + assignments_for_lane, cp_loss_multiplier, plan_cp_phases, ) @@ -33,15 +35,28 @@ def make_plan(lengths, *, lanes=8, minimum=1, maximum=None, budget=128, sp=2): ) +def test_one_token_budget_is_shared_by_score_and_train(): + config = DynamicContextParallelConfig(enabled=True, tokens_per_rank=64) + assert config.tokens_per_rank == 64 + with pytest.raises(ValueError): + DynamicContextParallelConfig( + enabled=True, + train_tokens_per_rank=64, + logprob_tokens_per_rank=128, + ) + + def check_plan(phases, lengths, lanes, budget, sp): owners = [] for phase in phases: for lane in range(lanes): - task = assignment_for_lane(phase, lane) - assert task.lane_start % task.cp_size == 0 - assert task.padded_tokens % (task.cp_size * sp) == 0 - assert task.padded_tokens // task.cp_size <= budget - for task in phase: + tasks = assignments_for_lane(phase, lane) + assert len({(task.lane_start, task.cp_size) for task in tasks}) == 1 + for task in tasks: + assert task.lane_start % task.cp_size == 0 + assert task.padded_tokens % (task.cp_size * sp) == 0 + assert task.padded_tokens // task.cp_size <= budget + for task in phase.assignments: if not task.sample_indices: continue owners.extend(task.sample_indices) @@ -75,8 +90,10 @@ def test_mixed_cp_and_cross_dp_groups(): lengths = [400, 200, 90, 70, 9, 3, 1000] phases = make_plan(lengths) check_plan(phases, lengths, 8, 128, 2) - assert {a.cp_size for p in phases for a in p if a.sample_indices} == {1, 2, 4, 8} - assert any(a.cp_size == 8 for p in phases for a in p) + assert ( + len({a.cp_size for p in phases for a in p.assignments if a.sample_indices}) > 1 + ) + assert any(a.cp_size == 8 for p in phases for a in p.assignments) @pytest.mark.parametrize("lanes", [1, 2, 4, 8, 16]) @@ -90,18 +107,66 @@ def test_randomized_coverage_and_alignment(lanes): def test_padding_drives_cp_selection(): - phases = make_plan([127], budget=127) - task = next(a for p in phases for a in p if a.sample_indices) + phases = make_plan([127], lanes=2, budget=127) + task = next(a for p in phases for a in p.assignments if a.sample_indices) assert task.cp_size == 2 -def test_moe_minimum_and_padding_only_lanes(): +def test_moe_minimum_expands_real_work_to_fill_lanes(): phases = make_plan([3, 7], minimum=4) - assert all(a.cp_size >= 4 for p in phases for a in p) - assert any(not a.sample_indices for p in phases for a in p) + assert all(a.cp_size >= 4 for p in phases for a in p.assignments) + assert not any(not a.sample_indices for p in phases for a in p.assignments) + assert [a.cp_size for p in phases for a in p.assignments] == [8] check_plan(phases, [3, 7], 8, 128, 2) +def test_idle_lanes_expand_real_assignment_and_recompute_padding(): + phases = make_plan([3]) + task = phases[0].assignments[0] + assert task.sample_indices == (0,) + assert task.cp_size == 8 + assert task.pad_multiple == 32 + assert task.padded_tokens == 32 + assert not any( + not assignment.sample_indices for assignment in phases[0].assignments + ) + check_plan(phases, [3], 8, 128, 2) + + +def test_maximum_cp_keeps_placeholders_when_real_work_cannot_fill_lanes(): + phases = make_plan([3], maximum=4) + real = [ + assignment for assignment in phases[0].assignments if assignment.sample_indices + ] + placeholders = [ + assignment + for assignment in phases[0].assignments + if not assignment.sample_indices + ] + assert [assignment.cp_size for assignment in real] == [4] + assert sum(assignment.cp_size for assignment in placeholders) == 4 + check_plan(phases, [3], 8, 128, 2) + + +def test_equal_topologies_merge_into_uneven_sequential_task_lists(): + phases = make_plan([40, 40, 40, 40, 40], lanes=2, maximum=1, budget=64, sp=1) + assert len(phases) == 1 + lane_zero = assignments_for_lane(phases[0], 0) + lane_one = assignments_for_lane(phases[0], 1) + assert sorted((len(lane_zero), len(lane_one))) == [2, 3] + assert { + index for task in lane_zero + lane_one for index in task.sample_indices + } == { + 0, + 1, + 2, + 3, + 4, + } + with pytest.raises(ValueError, match="sequential assignments"): + assignment_for_lane(phases[0], 0) + + @pytest.mark.parametrize( "lanes,minimum,maximum", [(3, 1, 2), (8, 3, 8), (8, 4, 2), (8, 1, 16)] ) diff --git a/tests/unit/distributed/test_dynamic_cp_dispatch.py b/tests/unit/distributed/test_dynamic_cp_dispatch.py index 5b1d0947660..978560863d9 100644 --- a/tests/unit/distributed/test_dynamic_cp_dispatch.py +++ b/tests/unit/distributed/test_dynamic_cp_dispatch.py @@ -11,14 +11,22 @@ import unittest from random import Random +from unittest.mock import patch import numpy as np import torch from nemo_rl.distributed.batched_data_dict import BatchedDataDict -from nemo_rl.distributed.dynamic_context_parallel import cp_loss_multiplier +from nemo_rl.distributed.dynamic_context_parallel import ( + cp_loss_multiplier, + plan_cp_phases, +) from nemo_rl.distributed.named_sharding import NamedSharding -from nemo_rl.models.policy.dynamic_cp import build_cp_dispatch, collect_cp_outputs +from nemo_rl.models.policy.dynamic_cp import ( + build_cp_dispatch, + collect_cp_outputs, + cp_schedule_matches, +) class TestDynamicCPDispatch(unittest.TestCase): @@ -40,8 +48,7 @@ def setUp(self): "sequence_parallel": False, "dynamic_context_parallel": { "enabled": True, - "train_tokens_per_rank": 8, - "logprob_tokens_per_rank": 8, + "tokens_per_rank": 8, "max_size": 4, }, }, @@ -170,6 +177,144 @@ def test_steps_keep_separate_denominators(self): self.assertGreater(plan.steps[0].valid_tokens, 0) self.assertEqual(plan.steps[1].valid_tokens, 0) + def test_dispatch_preserves_uneven_packed_task_lists(self): + data = BatchedDataDict( + input_ids=torch.arange(5 * 12).reshape(5, 12), + input_lengths=torch.tensor([10, 10, 10, 10, 10]), + sample_ids=torch.arange(5), + ) + data["sample_mask"] = torch.ones(5, dtype=torch.long) + data["token_mask"] = torch.ones(5, 12, dtype=torch.long) + cfg = {**self.cfg, "megatron_cfg": dict(self.cfg["megatron_cfg"])} + cfg["megatron_cfg"]["dynamic_context_parallel"] = { + "enabled": True, + "tokens_per_rank": 10, + "max_size": 1, + } + dispatch = build_cp_dispatch(data, cfg, self.mesh, batch_size=5, training=True) + local_counts = [ + len(plan.steps[0].assignments) + for dp_plans in dispatch.plans + for plan in dp_plans + ] + self.assertEqual(sorted(local_counts), [1, 1, 1, 2]) + self.assertTrue( + all( + len(plan.steps[0].groups) == 1 + for dp_plans in dispatch.plans + for plan in dp_plans + ) + ) + worker_outputs = [] + for payloads, plans in zip(dispatch.data, dispatch.plans): + for payload, plan in zip(payloads, plans): + values = [] + for task in plan.steps[0].assignments: + values.extend( + payload["sample_ids"][list(task.sample_indices)].tolist() + if task.sample_indices + else [-1] + ) + worker_outputs.append( + BatchedDataDict(logprobs=torch.tensor(values)[:, None]) + ) + restored = collect_cp_outputs(worker_outputs, dispatch, data.size) + torch.testing.assert_close(restored["logprobs"].flatten(), torch.arange(5)) + + baseline_weight = torch.tensor(0.3, dtype=torch.float64, requires_grad=True) + inputs = data["input_ids"][:, 1:].double() / 100 + baseline = torch.nn.functional.softplus(inputs * baseline_weight).sum() / 55 + baseline.backward() + + uneven_weight = baseline_weight.detach().clone().requires_grad_() + uneven_loss = uneven_weight * 0 + for payloads, plans in zip(dispatch.data, dispatch.plans): + for payload, plan in zip(payloads, plans): + step = plan.steps[0] + num_local_tasks = len(step.assignments) + for task in step.assignments: + if not task.sample_indices: + continue + local = payload.select_indices(list(task.sample_indices)) + local_loss = ( + torch.nn.functional.softplus( + local["input_ids"][:, 1:].double() / 100 * uneven_weight + ).sum() + / step.valid_tokens + ) + multiplier = cp_loss_multiplier( + active_cp_size=task.cp_size, + schedule_cp_size=2, + num_microbatches=num_local_tasks, + replicated_cp_loss=True, + ) + # MCore PP=1 applies static_CP / local_microbatch_count. + uneven_loss = ( + uneven_loss + local_loss * multiplier * 2 / num_local_tasks + ) + uneven_loss.backward() + torch.testing.assert_close(uneven_loss, baseline) + torch.testing.assert_close(uneven_weight.grad, baseline_weight.grad) + + def test_score_and_train_reuse_one_schedule(self): + with patch( + "nemo_rl.models.policy.dynamic_cp.plan_cp_phases", + wraps=plan_cp_phases, + ) as planner: + score = build_cp_dispatch( + self.data, + self.cfg, + self.mesh, + batch_size=5, + training=False, + ) + train = build_cp_dispatch( + self.data, + self.cfg, + self.mesh, + batch_size=5, + training=True, + schedule=score.schedule, + ) + + self.assertEqual(planner.call_count, 1) + self.assertIs(train.schedule, score.schedule) + self.assertTrue( + cp_schedule_matches( + score.schedule, + self.data, + self.cfg, + self.mesh, + batch_size=5, + ) + ) + for score_dp, train_dp in zip(score.plans, train.plans): + for score_plan, train_plan in zip(score_dp, train_dp): + self.assertEqual( + score_plan.steps[0].assignments, + train_plan.steps[0].assignments, + ) + self.assertEqual(score_plan.steps[0].valid_tokens, 0) + self.assertGreater(train_plan.steps[0].valid_tokens, 0) + + def test_schedule_reuse_rejects_different_batch_boundaries(self): + score = build_cp_dispatch( + self.data, + self.cfg, + self.mesh, + batch_size=5, + training=False, + ) + self.assertFalse( + cp_schedule_matches( + score.schedule, + self.data, + self.cfg, + self.mesh, + batch_size=1, + ) + ) + if __name__ == "__main__": unittest.main() From 9786faea3a3e2466870062b85ad5c9aa92a7005f Mon Sep 17 00:00:00 2001 From: Humaira Firdowse Mohammed Date: Wed, 16 Sep 2026 17:01:43 -0700 Subject: [PATCH 3/7] moe fixes - loss normalization, padding+mamba dyncp binding --- .../megatron-dynamic-context-parallel.md | 74 +++++-- ...-30ba3b-8n4g-megatron-dynamiccp-quick.yaml | 49 +++++ ...g-async-1off-megatron-dynamiccp-quick.yaml | 54 +++++ ...-async-1off-megatron-dynamiccp-10step.yaml | 30 +++ ...g-async-1off-megatron-dynamiccp-quick.yaml | 52 +++++ ...g-async-1off-megatron-staticcp-10step.yaml | 13 ++ nemo_rl/models/megatron/common.py | 24 ++- nemo_rl/models/megatron/dynamic_cp.py | 188 +++++++++++++++++- nemo_rl/models/megatron/setup.py | 11 +- nemo_rl/models/megatron/train.py | 2 +- nemo_rl/models/policy/dynamic_cp.py | 77 ++++++- .../policy/workers/megatron_policy_worker.py | 26 ++- .../functional/dynamic_cp_attention_parity.py | 2 +- .../distributed/test_dynamic_cp_dispatch.py | 113 +++++++++++ .../models/megatron/test_dynamic_cp_moe.py | 126 ++++++++++++ .../unit/models/megatron/test_moe_metrics.py | 42 ++++ .../test_dynamic_cp_comparison_recipes.py | 69 +++++++ tests/unit/test_dynamic_cp_moe_recipes.py | 88 ++++++++ 18 files changed, 997 insertions(+), 43 deletions(-) create mode 100644 examples/configs/recipes/llm/performance/grpo-nemotron3-nano-30ba3b-8n4g-megatron-dynamiccp-quick.yaml create mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-megatron-dynamiccp-quick.yaml create mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-10step.yaml create mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-quick.yaml create mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-staticcp-10step.yaml create mode 100644 tests/unit/models/megatron/test_dynamic_cp_moe.py create mode 100644 tests/unit/test_dynamic_cp_comparison_recipes.py create mode 100644 tests/unit/test_dynamic_cp_moe_recipes.py diff --git a/docs/design-docs/megatron-dynamic-context-parallel.md b/docs/design-docs/megatron-dynamic-context-parallel.md index c6c5d4bd88b..28cb7d611f2 100644 --- a/docs/design-docs/megatron-dynamic-context-parallel.md +++ b/docs/design-docs/megatron-dynamic-context-parallel.md @@ -51,9 +51,11 @@ not supported. Omit or disable the configuration for existing static behavior. For MoE, the minimum active size is raised until `active_CP * TP >= EP`. Workers also check that each actual expert communication group is contained within its -task's ranks. This constraint prevents expert collectives from crossing task -boundaries; it is not an end-to-end validation of MoE auxiliary losses or router -replay. The smoke recipe uses a dense model. +task's ranks. This keeps every EP collective inside one active-CP task and prevents +experts from communicating across independently scheduled CP blocks. Router +auxiliary losses and per-layer MoE metrics stay attached to the real microbatch +that produced them; placeholder tasks carry zero valid tokens and do not affect +loss normalization. ## Dispatch and execution @@ -183,6 +185,34 @@ importance ratios within 0.01 of one, and generation KL below 0.1, and writes `smoke_result.json` in that log directory. A completed smoke run verifies execution and finite training metrics, not convergence or a speedup over static CP. +### Dynamic-CP MoE quick smoke tests + +`perf_runs/run_gb200_dynamic_cp_moe.sh` runs the distributed loss and attention +preflights followed by a two-to-five-step real-model GRPO smoke test. It defaults +to Qwen3-30B-A3B, five steps, four nodes total (two generation nodes and two +policy nodes inherited from the async 1-off recipe), and a one-hour QoS limit: + +```bash +DYNAMIC_CP_MOE_STEPS=5 DRY_RUN=0 \ + bash perf_runs/run_gb200_dynamic_cp_moe.sh +``` + +Select the other model cases with `DYNAMIC_CP_MOE_MODEL=qwen235b` or +`DYNAMIC_CP_MOE_MODEL=nemotron3-nano`. Qwen3-30B-A3B uses TP1/EP8 and therefore +runs its policy task at CP8. Qwen3-235B-A22B uses TP8/EP16, so its minimum active +CP is two. Nemotron-3-Nano-30B-A3B uses TP2/EP8, so its minimum active CP is four. +Larger active sizes remain available for longer generated sequences. + +The Qwen3-30B-A3B recipe exercises `aux_loss`; Qwen3-235B-A22B exercises +`seq_aux_loss`; and the Nemotron recipe exercises its inherited router setup. +The post-run check requires the expected active CP size, finite training metrics, +the requested number of steps, and the configured MoE metric when applicable. + +The recipes log to both TensorBoard and W&B. The launcher defaults +`WANDB_MODE=online`, uses `/home/humairafirdo/hf_home`, and prints both settings +before submission. Set `WANDB_MODE=offline` explicitly when online logging is not +wanted. + ### Ten-step Nsight profile `perf_runs/run_gb200_dynamic_cp_profile.sh` runs the same dense Qwen3-32B setup @@ -205,19 +235,31 @@ into `-logs/ray/**/nsight/`. ### Ten-step dynamic/static comparison `perf_runs/run_gb200_cp_comparison.sh` runs matched ten-step jobs with the same -model, batch, TP=2, PP=1, base CP=1, generation setup, container, and W&B -project. Set `CP_MODE=dynamic` to allow active CP sizes 1–8, or -`CP_MODE=static` to keep CP=1. Use different `CP_RUN_NAME` values so the W&B -runs and local logs remain distinct. - -The launcher accepts `CP_MAX_TOTAL_SEQUENCE_LENGTH`, `CP_TOKENS_PER_RANK`, -`CP_MAX_SIZE`, and `STATIC_CP_SIZE`. A fair capacity-matched comparison uses the -smallest fixed CP that can accommodate the configured maximum at the same -per-rank token budget. For example, compare dynamic CP1–2 against static CP2 -with an 8192-token maximum and a 4096-token per-rank budget. Static CP1 remains -useful as an unconstrained throughput reference when it fits in memory; static -CP8 is a capacity-matched baseline only when the workload actually requires -CP8. +Qwen3-30B-A3B model, batch, TP4/EP4/PP1 policy topology, generation setup, +container, and W&B project. EP4 is only a sharding change; it does not remove +experts or change model weights. On the two policy nodes, TP4 creates two lanes +and `CP1 * TP4 = EP4`, so a complete expert group fits inside CP1. The dynamic +run can execute two CP1 tasks or one CP2 task, while the capacity-matched static +run stays at CP2. + +The default workload has an 8192-token ceiling, 4096 tokens per rank, and a +global batch of 512 formed from 16 prompts times 32 generations. The launcher +uses the batch QoS and a six-hour limit. Use different `CP_RUN_NAME` values so +the W&B runs and local logs remain distinct. + +The launcher also accepts `CP_NUM_STEPS`, `CP_TRAIN_GLOBAL_BATCH_SIZE`, +`CP_NUM_PROMPTS_PER_STEP`, `CP_NUM_GENERATIONS_PER_PROMPT`, +`CP_MAX_TOTAL_SEQUENCE_LENGTH`, `CP_TOKENS_PER_RANK`, `CP_MAX_SIZE`, and +`STATIC_CP_SIZE`. Prompt count times generations must equal the global batch. +A static CP1 run remains useful as an unconstrained throughput and CP1 sanity +reference when it fits in memory, but it is not the capacity-matched baseline +for 8192 tokens at the 4096-token budget. + +The correctness smoke keeps the original TP1/EP8 topology and is therefore +forced to CP8. The performance pair uses TP4/EP4 specifically to expose an +adaptive CP1/CP2 choice on the same eight policy GPUs. Both sides of the pair +use TP4/EP4, so the measured difference is dynamic versus fixed CP rather than +a model or expert-layout difference between the two runs. After each run, `perf_runs/analyze_cp_sequence_lengths.py` writes `sequence_length_distribution.json` beside the driver log. It reports length diff --git a/examples/configs/recipes/llm/performance/grpo-nemotron3-nano-30ba3b-8n4g-megatron-dynamiccp-quick.yaml b/examples/configs/recipes/llm/performance/grpo-nemotron3-nano-30ba3b-8n4g-megatron-dynamiccp-quick.yaml new file mode 100644 index 00000000000..8d13db7e101 --- /dev/null +++ b/examples/configs/recipes/llm/performance/grpo-nemotron3-nano-30ba3b-8n4g-megatron-dynamiccp-quick.yaml @@ -0,0 +1,49 @@ +defaults: ../grpo-dapomath17k-nanov3-30BA3B-8n4g-megatron-trtllm.yaml +grpo: + num_prompts_per_step: 8 + num_generations_per_prompt: 8 + max_num_steps: 5 + val_period: 1000 + val_at_start: false + val_at_end: false +checkpointing: + enabled: false + checkpoint_dir: results/grpo-nemotron3-nano-30ba3b-8n4g-dynamiccp-quick +policy: + train_global_batch_size: 64 + train_micro_batch_size: 1 + logprob_batch_size: 1 + max_total_sequence_length: 8192 + make_sequence_length_divisible_by: 4 + megatron_cfg: + tensor_model_parallel_size: 2 + pipeline_model_parallel_size: 1 + num_layers_in_first_pipeline_stage: null + num_layers_in_last_pipeline_stage: null + context_parallel_size: 1 + expert_tensor_parallel_size: 1 + expert_model_parallel_size: 8 + sequence_parallel: true + moe_hybridep_prepad_packed_inputs: false + dynamic_context_parallel: + enabled: true + min_size: 1 + max_size: 16 + tokens_per_rank: 4096 + fp8_cfg: + enabled: false + sequence_packing: + enabled: true + dynamic_batching: + enabled: false + generation: + max_new_tokens: 4096 +logger: + log_dir: logs/grpo-nemotron3-nano-30ba3b-8n4g-dynamiccp-quick + wandb_enabled: true + tensorboard_enabled: true + wandb: + project: nemo-rl + name: grpo-nemotron3-nano-30ba3b-8n4g-dynamiccp-quick +data_plane: + enabled: false diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-megatron-dynamiccp-quick.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-megatron-dynamiccp-quick.yaml new file mode 100644 index 00000000000..c3e7e7b8d69 --- /dev/null +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-megatron-dynamiccp-quick.yaml @@ -0,0 +1,54 @@ +defaults: ./grpo-qwen3-235b-32n4g-async-1off.yaml +grpo: + num_prompts_per_step: 8 + num_generations_per_prompt: 8 + max_num_steps: 5 + val_period: 1000 + val_at_start: false + val_at_end: false +checkpointing: + enabled: false + checkpoint_dir: results/grpo-qwen3-235b-32n4g-async-dynamiccp-quick +policy: + hf_config_overrides: + # Cover MCore's sequence-level router loss on a real MoE model. + router_aux_loss_coef: 0.001 + train_global_batch_size: 64 + train_micro_batch_size: 1 + logprob_batch_size: 1 + max_total_sequence_length: 8192 + megatron_cfg: + # PP must be one because every dynamic lane executes its own task list. + # TP8 keeps the dense portion sharded while CP*TP >= EP from CP2 onward. + tensor_model_parallel_size: 8 + pipeline_model_parallel_size: 1 + num_layers_in_first_pipeline_stage: null + num_layers_in_last_pipeline_stage: null + context_parallel_size: 1 + expert_tensor_parallel_size: 1 + expert_model_parallel_size: 16 + sequence_parallel: true + freeze_moe_router: false + moe_router_load_balancing_type: seq_aux_loss + moe_per_layer_logging: true + moe_hybridep_prepad_packed_inputs: false + dynamic_context_parallel: + enabled: true + min_size: 1 + max_size: 8 + tokens_per_rank: 4096 + fp8_cfg: + enabled: false + sequence_packing: + enabled: true + dynamic_batching: + enabled: false +logger: + log_dir: logs/grpo-qwen3-235b-32n4g-async-dynamiccp-quick + wandb_enabled: true + tensorboard_enabled: true + wandb: + project: nemo-rl + name: grpo-qwen3-235b-32n4g-async-dynamiccp-quick +data_plane: + enabled: false diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-10step.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-10step.yaml new file mode 100644 index 00000000000..7489fd0b98e --- /dev/null +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-10step.yaml @@ -0,0 +1,30 @@ +defaults: ./grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-quick.yaml +grpo: + # Keep the original Qwen30B performance recipe's 32 generations per prompt. + num_prompts_per_step: 16 + num_generations_per_prompt: 32 + max_num_steps: 10 +policy: + train_global_batch_size: 512 + max_total_sequence_length: 8192 + make_sequence_length_divisible_by: 4 + megatron_cfg: + # The two policy nodes provide eight ranks. TP4 creates two scheduling + # lanes; EP4 fits completely inside one CP1*TP4 task. Dynamic CP can + # therefore execute two CP1 tasks or one CP2 task without crossing an EP + # collective between independently scheduled tasks. + tensor_model_parallel_size: 4 + expert_model_parallel_size: 4 + sequence_parallel: true + dynamic_context_parallel: + min_size: 1 + max_size: 2 + tokens_per_rank: 4096 +logger: + log_dir: logs/grpo-qwen3-30ba3b-4n4g-async-dynamiccp-10step + wandb_enabled: true + tensorboard_enabled: true + wandb: + project: nemo-rl-cp-comparison + name: qwen3-30ba3b-4n4g-async-dynamiccp-10step + diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-quick.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-quick.yaml new file mode 100644 index 00000000000..fcfd57015b0 --- /dev/null +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-quick.yaml @@ -0,0 +1,52 @@ +defaults: ./grpo-qwen3-30ba3b-4n4g-async-1off.yaml +grpo: + num_prompts_per_step: 8 + num_generations_per_prompt: 8 + max_num_steps: 5 + val_period: 1000 + val_at_start: false + val_at_end: false +checkpointing: + enabled: false + checkpoint_dir: results/grpo-qwen3-30ba3b-4n4g-async-dynamiccp-quick +policy: + hf_config_overrides: + # Exercise the normal per-microbatch router auxiliary loss in this smoke run. + router_aux_loss_coef: 0.001 + train_global_batch_size: 64 + train_micro_batch_size: 1 + logprob_batch_size: 1 + max_total_sequence_length: 4096 + megatron_cfg: + tensor_model_parallel_size: 1 + pipeline_model_parallel_size: 1 + context_parallel_size: 1 + expert_tensor_parallel_size: 1 + expert_model_parallel_size: 8 + sequence_parallel: false + freeze_moe_router: false + moe_router_load_balancing_type: aux_loss + moe_per_layer_logging: true + moe_hybridep_prepad_packed_inputs: false + dynamic_context_parallel: + enabled: true + min_size: 1 + # The inherited async 1-off layout leaves two 4-GPU nodes for policy. + # TP1/EP8 therefore requires and exactly fills one active CP8 task. + max_size: 8 + tokens_per_rank: 4096 + fp8_cfg: + enabled: false + sequence_packing: + enabled: true + dynamic_batching: + enabled: false +logger: + log_dir: logs/grpo-qwen3-30ba3b-4n4g-async-dynamiccp-quick + wandb_enabled: true + tensorboard_enabled: true + wandb: + project: nemo-rl + name: grpo-qwen3-30ba3b-4n4g-async-dynamiccp-quick +data_plane: + enabled: false diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-staticcp-10step.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-staticcp-10step.yaml new file mode 100644 index 00000000000..903b6871c71 --- /dev/null +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-staticcp-10step.yaml @@ -0,0 +1,13 @@ +defaults: ./grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-10step.yaml +policy: + # CP2 is capacity-matched to an 8192-token maximum at 4096 tokens/rank. + # With TP4 and sequence parallelism, packed sequences align to CP*2*TP=16. + make_sequence_length_divisible_by: 16 + megatron_cfg: + context_parallel_size: 2 + dynamic_context_parallel: + enabled: false +logger: + log_dir: logs/grpo-qwen3-30ba3b-4n4g-async-staticcp-10step + wandb: + name: qwen3-30ba3b-4n4g-async-staticcp-10step diff --git a/nemo_rl/models/megatron/common.py b/nemo_rl/models/megatron/common.py index 26991b2aeb2..885ebd1d90b 100644 --- a/nemo_rl/models/megatron/common.py +++ b/nemo_rl/models/megatron/common.py @@ -201,6 +201,7 @@ def get_moe_metrics( num_layers: Optional[int] = None, mtp_num_layers: Optional[int] = None, track_names: Optional[list[str]] = None, + dynamic_parallel_group: Optional[dist.ProcessGroup] = None, ) -> dict[str, Any]: """Returns Mixture of Experts (MoE) auxiliary-loss metrics. @@ -222,6 +223,10 @@ def get_moe_metrics( records for the configured ``moe_router_load_balancing_type``, so callers should derive it via ``get_aux_loss_track_names(model_config)``. Defaults to None, which disables pre-initialization. + dynamic_parallel_group: Fixed TP*DP*CP group used only by dynamic CP. + Dynamic lanes can execute different numbers of microbatches and the + router's per-forward TP*CP group changes with each task, so its last + recorded reduction group is not safe for end-of-step metric sync. Returns: dict[str, Any]: A flat dict of aggregated metrics. For each aux loss name, @@ -258,8 +263,25 @@ def get_moe_metrics( for name in track_names: mcore_tracker.ensure_initialized(name, tracker_num_layers) - reduce_aux_losses_tracker_across_ranks() + dynamic_names: Optional[list[str]] = None + if dynamic_parallel_group is None: + reduce_aux_losses_tracker_across_ranks() + else: + # Each task contributes local router statistics on exactly its active + # TP*CP ranks. A single SUM over the fixed full policy group therefore + # counts every real task once, irrespective of uneven lane task counts. + # The caller's loss_scale is 1 / number_of_unique_real_tasks. + mcore_tracker = get_moe_metrics_tracker() + dynamic_names = ( + track_names if track_names is not None else list(mcore_tracker.metrics) + ) + for name in dynamic_names: + entry = mcore_tracker.metrics.get(name) + if entry is not None: + dist.all_reduce(entry.values, group=dynamic_parallel_group) tracker = get_moe_layer_wise_logging_tracker() + if dynamic_names is not None: + tracker = {name: tracker[name] for name in dynamic_names if name in tracker} metrics: dict[str, Any] = {} if len(tracker) > 0: diff --git a/nemo_rl/models/megatron/dynamic_cp.py b/nemo_rl/models/megatron/dynamic_cp.py index da2f13f02a5..425cf0a0165 100644 --- a/nemo_rl/models/megatron/dynamic_cp.py +++ b/nemo_rl/models/megatron/dynamic_cp.py @@ -18,10 +18,15 @@ import torch from megatron.core import parallel_state from megatron.core.transformer.attention import Attention +from megatron.core.transformer.moe.router import Router from nemo_rl.distributed.dynamic_context_parallel import CPRankPlan, CPRankStep +_DYNAMIC_TP_CP_GROUPS: dict[int, Any] = {} +_ROUTER_CONFIG_BASELINES: dict[int, tuple[Any, Any]] = {} + + @dataclass(frozen=True) class RuntimeCPContext: """Attention group for a single forward/backward, separate from DDP groups.""" @@ -31,12 +36,14 @@ class RuntimeCPContext: group: Any -def initialize_dynamic_cp_runtime() -> None: +def initialize_dynamic_cp_runtime(*, max_cp_size: int) -> None: """Initialize CP resources missing on CP=1 builds of the pinned MCore. This creates the same stream as TEDotProductAttention's CP constructor; - it does not replace methods, edit dependency files, or change static groups. - Newer MCore creates this lazily, so the initialization is idempotent. + it also creates the active TP*CP groups used by MoE routers. Groups must + be created eagerly and in the same order on every rank; creating one from + a microbatch forward would deadlock as different lanes select different CP + sizes. Static DP, EP and optimizer groups are left unchanged. """ # TE is an optional dependency outside the Megatron worker environment. from megatron.core.extensions.transformer_engine import TEDotProductAttention @@ -44,34 +51,165 @@ def initialize_dynamic_cp_runtime() -> None: if TEDotProductAttention.cp_stream is None: TEDotProductAttention.cp_stream = torch.cuda.Stream() + tp_size = parallel_state.get_tensor_model_parallel_world_size() + tp_rank = parallel_state.get_tensor_model_parallel_rank() + lane_group = parallel_state.get_data_parallel_group(with_context_parallel=True) + lane_ranks = torch.distributed.get_process_group_ranks(lane_group) + base_lane_ranks = tuple(rank - tp_rank for rank in lane_ranks) + if any(base % tp_size for base in base_lane_ranks): + raise ValueError("Dynamic CP requires contiguous TP ranks") + if max_cp_size > len(base_lane_ranks): + raise ValueError("Dynamic CP max_size exceeds the available DP*CP lanes") + + _DYNAMIC_TP_CP_GROUPS.clear() + _DYNAMIC_TP_CP_GROUPS[1] = parallel_state.get_tensor_model_parallel_group() + cp_size = 2 + while cp_size <= max_cp_size: + if len(base_lane_ranks) % cp_size: + raise ValueError("Every dynamic CP size must divide the DP*CP lane domain") + local_group = None + for start in range(0, len(base_lane_ranks), cp_size): + ranks = [ + base + offset + for base in base_lane_ranks[start : start + cp_size] + for offset in range(tp_size) + ] + if tp_size == 1: + group = parallel_state.get_hybrid_data_context_parallel_groups( + group_size=cp_size + ) + else: + group = torch.distributed.new_group(ranks=ranks) + if torch.distributed.get_rank() in ranks: + local_group = group + if local_group is None: + raise ValueError("Rank was not assigned to a dynamic TP*CP group") + _DYNAMIC_TP_CP_GROUPS[cp_size] = local_group + cp_size *= 2 + + +def _is_mamba_mixer(module: torch.nn.Module) -> bool: + cp = getattr(module, "cp", None) + return cp is not None and all( + hasattr(cp, name) + for name in ( + "d_inner_local_tp", + "nheads_local_tp", + "ngroups_local_tp", + "conv1d_weight_cp1", + "conv1d_bias_cp1", + "dt_bias_cp1", + "A_log_cp1", + "D_cp1", + ) + ) + + +def _is_gated_delta_net(module: torch.nn.Module) -> bool: + return ( + ".ssm.gated_delta_net" in type(module).__module__ + and hasattr(module, "cp_size") + and hasattr(module, "pg_collection") + ) + + +def _rebuild_mamba_cp(module: torch.nn.Module, group: Any) -> None: + """Rebuild Mamba's cached CP helper for the active microbatch size.""" + cp = module.cp + module.cp = type(cp)( + cp_group=group, + d_inner_local_tp=cp.d_inner_local_tp, + nheads_local_tp=cp.nheads_local_tp, + ngroups_local_tp=cp.ngroups_local_tp, + d_state=cp.d_state, + conv1d_weight_cp1=cp.conv1d_weight_cp1, + conv1d_bias_cp1=cp.conv1d_bias_cp1, + conv1d_padding=cp.conv1d_padding, + dt_bias_cp1=cp.dt_bias_cp1, + A_log_cp1=cp.A_log_cp1, + D_cp1=cp.D_cp1, + D_has_hdim=cp.D_has_hdim, + ) + @contextmanager def preserve_attention_cp_groups(model: torch.nn.Module) -> Iterator[None]: - """Isolate attention's runtime groups from shared model/DDP collections. + """Isolate runtime attention, SSM and router state from fixed groups. Keep the active group through backward recomputation; restore it after the complete no-pipeline schedule. TP, DP and optimizer groups are unchanged. """ - saved = [ + saved_collections = [ (module, module.pg_collection) for module in model.modules() if isinstance(module, Attention) + or _is_mamba_mixer(module) + or _is_gated_delta_net(module) + ] + saved_mamba = [ + (module, module.cp) for module in model.modules() if _is_mamba_mixer(module) + ] + saved_gdn = [ + (module, module.cp_size) + for module in model.modules() + if _is_gated_delta_net(module) + ] + saved_routers = [ + (module, module.tp_cp_group) + for module in model.modules() + if isinstance(module, Router) ] - for module, collection in saved: + router_configs: dict[int, Any] = {} + for module, collection in saved_collections: module.pg_collection = copy(collection) + for module, _ in saved_routers: + config = module.config + router_configs[id(config)] = config + _ROUTER_CONFIG_BASELINES[id(config)] = ( + config.moe_aux_loss_coeff, + config.moe_z_loss_coeff, + ) try: yield finally: - for module, collection in saved: + for module, collection in saved_collections: module.pg_collection = collection + for module, cp in saved_mamba: + module.cp = cp + for module, cp_size in saved_gdn: + module.cp_size = cp_size + for module, group in saved_routers: + module.tp_cp_group = group + for config_id, config in router_configs.items(): + aux_coeff, z_coeff = _ROUTER_CONFIG_BASELINES.pop(config_id) + config.moe_aux_loss_coeff = aux_coeff + config.moe_z_loss_coeff = z_coeff + + +def _bind_router_config(router: Router, *, padding_only: bool) -> None: + baseline = _ROUTER_CONFIG_BASELINES.get(id(router.config)) + if baseline is None: + raise RuntimeError("Dynamic CP router binding escaped its preservation context") + aux_coeff, z_coeff = baseline + if padding_only: + if isinstance(aux_coeff, tuple): + aux_coeff = tuple(0.0 for _ in aux_coeff) + elif isinstance(aux_coeff, list): + aux_coeff = [0.0 for _ in aux_coeff] + else: + aux_coeff = 0.0 + z_coeff = None + router.config.moe_aux_loss_coeff = aux_coeff + router.config.moe_z_loss_coeff = z_coeff def bind_attention_cp_group(model: torch.nn.Module, packed_seq_params: Any) -> Any: - """Bind RoPE as well as TE attention to the microbatch's active group. + """Bind attention, MoE router and SSM modules to the active CP task. The pinned MCore forwards packed.cp_group to TE but its RoPE still reads - Attention.pg_collection.cp. CP=1 needs a real singleton here because None - means fall back to static CP in MCore's RoPE API. PP=1 supplies that singleton. + ``Attention.pg_collection.cp``. Mamba and GatedDeltaNet also cache their CP + helper/group, while MoE routers cache TP*CP for aux-loss token reductions. + CP=1 needs a real singleton because None means static fallback in MCore. Pass a copy to the model so loss/gather metadata retains CP=1's None group. """ context = runtime_cp_from_packed(packed_seq_params) @@ -80,9 +218,29 @@ def bind_attention_cp_group(model: torch.nn.Module, packed_seq_params: Any) -> A group = parallel_state.get_pipeline_model_parallel_group() if group.size() != 1: raise ValueError("Dynamic CP attention requires PP=1") + tp_cp_group = _DYNAMIC_TP_CP_GROUPS.get(context.size) + if tp_cp_group is None: + raise ValueError(f"No dynamic TP*CP group was initialized for CP={context.size}") + expected_tp_cp_size = ( + context.size * parallel_state.get_tensor_model_parallel_world_size() + ) + if tp_cp_group.size() != expected_tp_cp_size: + raise ValueError("Dynamic MoE TP*CP group has the wrong size") + padding_only = bool( + getattr(packed_seq_params, "dynamic_cp_padding_only", False) + ) for module in model.modules(): if isinstance(module, Attention): module.pg_collection.cp = group + if isinstance(module, Router): + module.tp_cp_group = tp_cp_group + _bind_router_config(module, padding_only=padding_only) + if _is_mamba_mixer(module): + module.pg_collection.cp = group + _rebuild_mamba_cp(module, group) + elif _is_gated_delta_net(module): + module.pg_collection.cp = group + module.cp_size = context.size model_packed = copy(packed_seq_params) model_packed.cp_group = group return model_packed @@ -157,7 +315,17 @@ def planned_microbatches( pad_individual_seqs_to_multiple_of=assignment.pad_multiple, straggler_timer=straggler_timer, cp_context=context, + create_packed_seq_padding_mask=True, ) + padding_only = not assignment.sample_indices + inputs.packed_seq_params.dynamic_cp_padding_only = padding_only + if padding_only: + # MCore uses True to mean "padding". Excluding every physical + # placeholder token keeps expert-bias counters clean; router aux + # coefficients are disabled for this forward to avoid a 0/0 loss. + inputs.padding_mask = torch.ones_like( + inputs.input_ids_cp_sharded, dtype=torch.bool + ) if inputs.input_ids_cp_sharded.shape[1] * size != assignment.padded_tokens: raise ValueError( "Packed worker token count disagrees with the driver's plan" diff --git a/nemo_rl/models/megatron/setup.py b/nemo_rl/models/megatron/setup.py index 1952aff853c..b8dd784d2e1 100644 --- a/nemo_rl/models/megatron/setup.py +++ b/nemo_rl/models/megatron/setup.py @@ -2116,11 +2116,12 @@ def setup_model_and_optimizer( get_position_embedding_ranks=get_position_embedding_ranks, ) - if ( - dynamic_cp_config(policy_cfg) is not None - and megatron_cfg.model.hybrid_context_parallel - ): - initialize_dynamic_cp_runtime() + dynamic = dynamic_cp_config(policy_cfg) + if dynamic is not None: + initialize_dynamic_cp_runtime( + max_cp_size=dynamic.max_size + or parallel_state.get_data_parallel_world_size(with_context_parallel=True) + ) # The Ray driver supplies packed-task synchronization groups and each # lane's possibly uneven task list. Use MCore's standard PP=1 executor, # whose no_sync handling covers every local task except the last, rather diff --git a/nemo_rl/models/megatron/train.py b/nemo_rl/models/megatron/train.py index 83b4914ea0c..e84f2889f42 100644 --- a/nemo_rl/models/megatron/train.py +++ b/nemo_rl/models/megatron/train.py @@ -213,7 +213,7 @@ def model_forward( position_ids = None additional_kwargs = {} - # Mamba models currently do not support packed_seq_params + # Packed metadata drives TE attention and the Mamba/GatedDelta CP layouts. if packed_seq_params is not None: additional_kwargs["packed_seq_params"] = ( bind_attention_cp_group(model, packed_seq_params) diff --git a/nemo_rl/models/policy/dynamic_cp.py b/nemo_rl/models/policy/dynamic_cp.py index e0c694eb2df..dc5a3f6b09e 100644 --- a/nemo_rl/models/policy/dynamic_cp.py +++ b/nemo_rl/models/policy/dynamic_cp.py @@ -30,6 +30,51 @@ logger = logging.getLogger(__name__) +def _enabled_global_aux_loss(megatron_cfg: dict[str, Any]) -> bool: + """Return whether a full-DP global aux collective is configured. + + The coefficient can originate in the HF model provider rather than this + dictionary, so the routing type itself must be rejected. + """ + routing_type = megatron_cfg.get("moe_router_load_balancing_type") + routing_types = ( + list(routing_type) + if isinstance(routing_type, (list, tuple)) + else [routing_type] + ) + return "global_aux_loss" in routing_types + + +def _minimum_cp_size_for_experts( + megatron_cfg: dict[str, Any], configured_minimum: int +) -> int: + """Keep every EP collective inside one dynamically scheduled task. + + Dynamic CP tasks contain ``CP * TP`` contiguous model ranks. The pinned + MCore rank order puts a complete EP group inside such a block once it is at + least EP ranks wide. ETP changes that layout, so support it only after it + has a dedicated topology implementation. + """ + expert_parallel = megatron_cfg["expert_model_parallel_size"] + tensor_parallel = megatron_cfg["tensor_model_parallel_size"] + expert_tensor_parallel = megatron_cfg.get("expert_tensor_parallel_size", 1) + if expert_tensor_parallel != 1: + raise ValueError( + "Dynamic CP MoE requires expert_tensor_parallel_size=1" + ) + if expert_parallel <= 1: + return configured_minimum + minimum = configured_minimum + while minimum * tensor_parallel < expert_parallel: + minimum *= 2 + if (minimum * tensor_parallel) % expert_parallel: + raise ValueError( + "Dynamic CP requires expert_model_parallel_size to divide " + "min_dynamic_cp_size * tensor_model_parallel_size" + ) + return minimum + + def dynamic_cp_config(cfg: dict[str, Any]) -> DynamicContextParallelConfig | None: """Read optional config without introducing defaults at worker call sites.""" megatron = cfg.get("megatron_cfg") @@ -62,16 +107,26 @@ def validate_dynamic_cp(cfg: dict[str, Any], *, lanes: int) -> None: raise ValueError("Dynamic CP does not support MTP or HybridEP input prepadding") if cfg["sequence_packing"].get("pair_grouping_key"): raise ValueError("Dynamic CP does not yet schedule atomic preference pairs") + if _enabled_global_aux_loss(mc): + raise ValueError( + "Dynamic CP does not support global_aux_loss: its per-forward " + "TP*DP*CP collective cannot be called by uneven CP task lists. " + "Use aux_loss or seq_aux_loss instead." + ) # Probe even empty plans, so invalid domains/bounds fail at initialization. tp = mc["tensor_model_parallel_size"] - minimum = dynamic.min_size - while minimum * tp < mc["expert_model_parallel_size"]: - minimum *= 2 + minimum = _minimum_cp_size_for_experts(mc, dynamic.min_size) + maximum = dynamic.max_size or lanes + if minimum > maximum: + raise ValueError( + "Dynamic CP cannot contain an EP group: effective min_size " + f"{minimum} exceeds max_size {maximum} (CP*TP must be >= EP)" + ) plan_cp_phases( [], lanes=lanes, min_size=minimum, - max_size=dynamic.max_size or lanes, + max_size=maximum, tokens_per_rank=dynamic.tokens_per_rank, sequence_parallel_size=tp if mc["sequence_parallel"] else 1, user_pad_multiple=cfg["make_sequence_length_divisible_by"], @@ -104,6 +159,16 @@ class CPBatchSchedule: groups_by_batch: tuple[tuple[CPSyncGroup, ...], ...] +def owned_real_task_count(plan: CPRankPlan) -> int: + """Count unique real packed tasks owned by this lane across all steps.""" + return sum( + 1 + for step in plan.steps + for task in step.assignments + if task.sample_indices and plan.lane == task.lane_start + ) + + def _schedule_parameters( cfg: dict[str, Any], sharding: NamedSharding ) -> tuple[int, int, int, int, int, int, int]: @@ -114,9 +179,7 @@ def _schedule_parameters( cp = sharding.shape["context_parallel"] lanes = sharding.shape["data_parallel"] * cp tp = mc["tensor_model_parallel_size"] - minimum = dynamic.min_size - while minimum * tp < mc["expert_model_parallel_size"]: - minimum *= 2 + minimum = _minimum_cp_size_for_experts(mc, dynamic.min_size) fp8 = mc.get("fp8_cfg") or {} alignment = 1 if fp8.get("enabled"): diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index 601f4b68c7f..8dc7e34ae28 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -106,7 +106,7 @@ megatron_forward_backward, ) from nemo_rl.models.policy import PolicyConfig -from nemo_rl.models.policy.dynamic_cp import dynamic_cp_config +from nemo_rl.models.policy.dynamic_cp import dynamic_cp_config, owned_real_task_count from nemo_rl.models.policy.interfaces import ( ColocatablePolicyInterface, LogprobOutputSpec, @@ -1313,7 +1313,28 @@ def train( model_config = getattr(self.model, "config", None) num_moe_experts = getattr(model_config, "num_moe_experts", None) if num_moe_experts is not None and num_moe_experts > 1: - moe_loss_scale = 1.0 / max(1, total_num_microbatches) + dynamic_moe_group = None + if cp_plan is not None: + # Count task owners, not rank-local microbatches: a CP=C task + # appears on C lanes but represents one packed model call. + local_real_tasks = owned_real_task_count(cp_plan) + global_real_tasks = torch.tensor( + local_real_tasks, dtype=torch.int64, device="cuda" + ) + torch.distributed.all_reduce( + global_real_tasks, + group=parallel_state.get_data_parallel_group( + with_context_parallel=True + ), + ) + moe_loss_scale = 1.0 / max(1, int(global_real_tasks.item())) + dynamic_moe_group = ( + parallel_state.get_tensor_and_data_parallel_group( + with_context_parallel=True + ) + ) + else: + moe_loss_scale = 1.0 / max(1, total_num_microbatches) moe_metrics = get_moe_metrics( loss_scale=moe_loss_scale, per_layer_logging=self.cfg["megatron_cfg"]["moe_per_layer_logging"], @@ -1324,6 +1345,7 @@ def train( num_layers=getattr(model_config, "num_layers", None), mtp_num_layers=getattr(model_config, "mtp_num_layers", None), track_names=get_aux_loss_track_names(model_config), + dynamic_parallel_group=dynamic_moe_group, ) if moe_metrics: metrics["moe_metrics"] = moe_metrics diff --git a/tests/functional/dynamic_cp_attention_parity.py b/tests/functional/dynamic_cp_attention_parity.py index 17597140811..d1e333167a7 100644 --- a/tests/functional/dynamic_cp_attention_parity.py +++ b/tests/functional/dynamic_cp_attention_parity.py @@ -94,7 +94,7 @@ def main() -> None: context_parallel_size=base_cp, hybrid_context_parallel=True, ) - initialize_dynamic_cp_runtime() + initialize_dynamic_cp_runtime(max_cp_size=world // tp) torch.manual_seed(123) model_parallel_cuda_manual_seed(123) config = TransformerConfig( diff --git a/tests/unit/distributed/test_dynamic_cp_dispatch.py b/tests/unit/distributed/test_dynamic_cp_dispatch.py index 978560863d9..12e66b2ff65 100644 --- a/tests/unit/distributed/test_dynamic_cp_dispatch.py +++ b/tests/unit/distributed/test_dynamic_cp_dispatch.py @@ -23,9 +23,12 @@ ) from nemo_rl.distributed.named_sharding import NamedSharding from nemo_rl.models.policy.dynamic_cp import ( + _enabled_global_aux_loss, + _minimum_cp_size_for_experts, build_cp_dispatch, collect_cp_outputs, cp_schedule_matches, + owned_real_task_count, ) @@ -63,6 +66,79 @@ def setUp(self): ], ) + def test_moe_minimum_contains_complete_expert_group(self): + self.assertEqual( + _minimum_cp_size_for_experts( + { + "tensor_model_parallel_size": 1, + "expert_tensor_parallel_size": 1, + "expert_model_parallel_size": 8, + }, + 1, + ), + 8, + ) + self.assertEqual( + _minimum_cp_size_for_experts( + { + "tensor_model_parallel_size": 2, + "expert_tensor_parallel_size": 1, + "expert_model_parallel_size": 16, + }, + 1, + ), + 8, + ) + with self.assertRaisesRegex(ValueError, "expert_tensor_parallel_size=1"): + _minimum_cp_size_for_experts( + { + "tensor_model_parallel_size": 2, + "expert_tensor_parallel_size": 2, + "expert_model_parallel_size": 8, + }, + 1, + ) + with self.assertRaisesRegex(ValueError, "expert_tensor_parallel_size=1"): + _minimum_cp_size_for_experts( + { + "tensor_model_parallel_size": 2, + "expert_tensor_parallel_size": 2, + "expert_model_parallel_size": 1, + }, + 1, + ) + with self.assertRaisesRegex(ValueError, "to divide"): + _minimum_cp_size_for_experts( + { + "tensor_model_parallel_size": 2, + "expert_tensor_parallel_size": 1, + "expert_model_parallel_size": 12, + }, + 1, + ) + + def test_global_aux_loss_is_rejected_even_with_provider_coefficient(self): + self.assertTrue( + _enabled_global_aux_loss( + {"moe_router_load_balancing_type": "global_aux_loss"} + ) + ) + self.assertTrue( + _enabled_global_aux_loss( + { + "moe_router_load_balancing_type": [ + "aux_loss", + "global_aux_loss", + ] + } + ) + ) + self.assertFalse( + _enabled_global_aux_loss( + {"moe_router_load_balancing_type": "seq_aux_loss"} + ) + ) + def test_outputs_from_nonzero_static_cp_are_preserved(self): dispatch = build_cp_dispatch( self.data, self.cfg, self.mesh, batch_size=None, training=False @@ -83,6 +159,29 @@ def test_outputs_from_nonzero_static_cp_are_preserved(self): with self.assertRaises(ValueError): collect_cp_outputs(results[::2], dispatch, self.data.size) + def test_moe_dispatch_never_schedules_less_than_ep_over_tp(self): + cfg = {**self.cfg, "megatron_cfg": dict(self.cfg["megatron_cfg"])} + cfg["megatron_cfg"].update( + { + "expert_tensor_parallel_size": 1, + "expert_model_parallel_size": 8, + "dynamic_context_parallel": { + "enabled": True, + "tokens_per_rank": 8, + "max_size": 4, + }, + } + ) + dispatch = build_cp_dispatch( + self.data, cfg, self.mesh, batch_size=None, training=False + ) + assert all( + task.cp_size == 4 + for dp_plans in dispatch.plans + for plan in dp_plans + for task in plan.steps[0].assignments + ) + def test_global_normalizer_and_autograd_do_not_count_cp_copies(self): dispatch = build_cp_dispatch( self.data, self.cfg, self.mesh, batch_size=5, training=True @@ -198,6 +297,20 @@ def test_dispatch_preserves_uneven_packed_task_lists(self): for plan in dp_plans ] self.assertEqual(sorted(local_counts), [1, 1, 1, 2]) + unique_real_tasks = sum( + 1 + for group in dispatch.schedule.groups_by_batch[0] + for task in group.assignments + if task.sample_indices + ) + self.assertEqual( + sum( + owned_real_task_count(plan) + for dp_plans in dispatch.plans + for plan in dp_plans + ), + unique_real_tasks, + ) self.assertTrue( all( len(plan.steps[0].groups) == 1 diff --git a/tests/unit/models/megatron/test_dynamic_cp_moe.py b/tests/unit/models/megatron/test_dynamic_cp_moe.py new file mode 100644 index 00000000000..bde07e3ce31 --- /dev/null +++ b/tests/unit/models/megatron/test_dynamic_cp_moe.py @@ -0,0 +1,126 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# Licensed under the Apache License, Version 2.0 (the "License"); +"""Runtime binding contracts shared by Qwen MoE and Nemotron hybrid models.""" + +from types import SimpleNamespace + +import pytest +import torch + +pytestmark = pytest.mark.mcore + + +class _Group: + def __init__(self, size, rank=0): + self._size = size + self._rank = rank + + def size(self): + return self._size + + def rank(self): + return self._rank + + +class _CPHelper: + def __init__( + self, + *, + cp_group, + d_inner_local_tp=32, + nheads_local_tp=16, + ngroups_local_tp=4, + d_state=8, + conv1d_weight_cp1=None, + conv1d_bias_cp1=None, + conv1d_padding=3, + dt_bias_cp1=None, + A_log_cp1=None, + D_cp1=None, + D_has_hdim=False, + ): + self.cp_group = cp_group + self.d_inner_local_tp = d_inner_local_tp + self.nheads_local_tp = nheads_local_tp + self.ngroups_local_tp = ngroups_local_tp + self.d_state = d_state + self.conv1d_weight_cp1 = conv1d_weight_cp1 + self.conv1d_bias_cp1 = conv1d_bias_cp1 + self.conv1d_padding = conv1d_padding + self.dt_bias_cp1 = dt_bias_cp1 + self.A_log_cp1 = A_log_cp1 + self.D_cp1 = D_cp1 + self.D_has_hdim = D_has_hdim + + +class _Mamba(torch.nn.Module): + def __init__(self, group): + super().__init__() + self.pg_collection = SimpleNamespace(cp=group) + self.cp = _CPHelper(cp_group=group) + + +class _GatedDelta(torch.nn.Module): + def __init__(self, group): + super().__init__() + self.pg_collection = SimpleNamespace(cp=group) + self.cp_size = group.size() + + +_GatedDelta.__module__ = "megatron.core.ssm.gated_delta_net" + + +class _Router(torch.nn.Module): + def __init__(self, group, config): + super().__init__() + self.tp_cp_group = group + self.config = config + + +def test_dynamic_binding_updates_router_and_ssm_then_restores(monkeypatch): + from nemo_rl.models.megatron import dynamic_cp + + original_cp = _Group(1) + active_cp = _Group(2, rank=1) + original_tp_cp = _Group(2) + active_tp_cp = _Group(4) + config = SimpleNamespace(moe_aux_loss_coeff=[0.1, 0.2], moe_z_loss_coeff=0.01) + + model = torch.nn.Module() + model.add_module("mamba", _Mamba(original_cp)) + model.add_module("gdn", _GatedDelta(original_cp)) + model.add_module("router", _Router(original_tp_cp, config)) + packed = SimpleNamespace( + local_cp_size=2, + cp_group=active_cp, + dynamic_cp_padding_only=True, + ) + + monkeypatch.setattr(dynamic_cp, "Router", _Router) + monkeypatch.setattr( + dynamic_cp.parallel_state, + "get_tensor_model_parallel_world_size", + lambda: 2, + ) + monkeypatch.setitem(dynamic_cp._DYNAMIC_TP_CP_GROUPS, 2, active_tp_cp) + + original_mamba_helper = model.mamba.cp + with dynamic_cp.preserve_attention_cp_groups(model): + model_packed = dynamic_cp.bind_attention_cp_group(model, packed) + assert model_packed.cp_group is active_cp + assert model.router.tp_cp_group is active_tp_cp + assert model.mamba.pg_collection.cp is active_cp + assert model.mamba.cp is not original_mamba_helper + assert model.mamba.cp.cp_group is active_cp + assert model.gdn.pg_collection.cp is active_cp + assert model.gdn.cp_size == 2 + assert config.moe_aux_loss_coeff == [0.0, 0.0] + assert config.moe_z_loss_coeff is None + + assert model.router.tp_cp_group is original_tp_cp + assert model.mamba.pg_collection.cp is original_cp + assert model.mamba.cp is original_mamba_helper + assert model.gdn.pg_collection.cp is original_cp + assert model.gdn.cp_size == 1 + assert config.moe_aux_loss_coeff == [0.1, 0.2] + assert config.moe_z_loss_coeff == 0.01 diff --git a/tests/unit/models/megatron/test_moe_metrics.py b/tests/unit/models/megatron/test_moe_metrics.py index 0d76cde0f47..18ee6e2a351 100644 --- a/tests/unit/models/megatron/test_moe_metrics.py +++ b/tests/unit/models/megatron/test_moe_metrics.py @@ -120,6 +120,48 @@ def _clear(): assert cleared["called"], "clear_aux_losses_tracker should be called" +@pytest.mark.mcore +def test_dynamic_cp_metrics_use_one_fixed_sum_group(monkeypatch): + """Uneven lanes must not reduce through the last router task's group.""" + from nemo_rl.models import megatron as megatron_module + from nemo_rl.models.megatron.common import get_moe_metrics + + entry = SimpleNamespace(values=torch.tensor([1.0, 3.0])) + live_tracker = SimpleNamespace(metrics={"load_balancing_loss": entry}) + reductions = [] + fixed_group = object() + + def _all_reduce(values, *, group): + reductions.append(group) + values.mul_(2.0) + + monkeypatch.setattr( + megatron_module.common, "get_moe_metrics_tracker", lambda: live_tracker + ) + monkeypatch.setattr( + megatron_module.common, + "get_moe_layer_wise_logging_tracker", + lambda: {"load_balancing_loss": {"values": entry.values}}, + ) + monkeypatch.setattr( + megatron_module.common, + "reduce_aux_losses_tracker_across_ranks", + lambda: pytest.fail("dynamic CP used the mutable per-task router group"), + ) + monkeypatch.setattr(megatron_module.common.dist, "all_reduce", _all_reduce) + monkeypatch.setattr( + megatron_module.common, "clear_aux_losses_tracker", lambda: None + ) + + metrics = get_moe_metrics( + loss_scale=0.25, + dynamic_parallel_group=fixed_group, + ) + + assert reductions == [fixed_group] + assert metrics["load_balancing_loss"] == pytest.approx(1.0) + + @pytest.mark.mcore @pytest.mark.parametrize( "routing_type,aux_loss_coeff,z_loss_coeff,expected", diff --git a/tests/unit/test_dynamic_cp_comparison_recipes.py b/tests/unit/test_dynamic_cp_comparison_recipes.py new file mode 100644 index 00000000000..72b7374c214 --- /dev/null +++ b/tests/unit/test_dynamic_cp_comparison_recipes.py @@ -0,0 +1,69 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# Licensed under the Apache License, Version 2.0 (the "License"); +"""Validate that the Qwen30B MoE dynamic/static CP pair is matched.""" + +from pathlib import Path + +import pytest +from omegaconf import OmegaConf + +from nemo_rl.models.policy.dynamic_cp import _minimum_cp_size_for_experts +from nemo_rl.utils.config import load_config, register_omegaconf_resolvers + + +RECIPE_DIR = ( + Path(__file__).resolve().parents[2] + / "examples/configs/recipes/llm/performance" +) + + +@pytest.mark.parametrize( + "recipe_name,dynamic_enabled,context_parallel_size,pad_factor", + ( + ( + "grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-10step.yaml", + True, + 1, + 4, + ), + ( + "grpo-qwen3-30ba3b-4n4g-async-1off-megatron-staticcp-10step.yaml", + False, + 2, + 16, + ), + ), +) +def test_cp_comparison_recipe_pair( + recipe_name, dynamic_enabled, context_parallel_size, pad_factor +): + register_omegaconf_resolvers() + resolved = OmegaConf.to_container( + load_config(RECIPE_DIR / recipe_name), resolve=True + ) + assert isinstance(resolved, dict) + + assert resolved["grpo"]["max_num_steps"] == 10 + assert resolved["grpo"]["num_prompts_per_step"] == 16 + assert resolved["grpo"]["num_generations_per_prompt"] == 32 + assert resolved["policy"]["train_global_batch_size"] == 512 + assert resolved["policy"]["max_total_sequence_length"] == 8192 + assert resolved["policy"]["make_sequence_length_divisible_by"] == pad_factor + assert resolved["policy"]["model_name"] == "Qwen/Qwen3-30B-A3B" + assert resolved["cluster"]["num_nodes"] == 4 + assert resolved["policy"]["generation"]["colocated"]["enabled"] is False + assert resolved["policy"]["generation"]["colocated"]["resources"]["num_nodes"] == 2 + + megatron = resolved["policy"]["megatron_cfg"] + assert megatron["tensor_model_parallel_size"] == 4 + assert megatron["pipeline_model_parallel_size"] == 1 + assert megatron["expert_model_parallel_size"] == 4 + assert megatron["expert_tensor_parallel_size"] == 1 + assert _minimum_cp_size_for_experts(megatron, 1) == 1 + assert megatron["context_parallel_size"] == context_parallel_size + assert megatron["dynamic_context_parallel"]["enabled"] is dynamic_enabled + assert megatron["dynamic_context_parallel"]["max_size"] == 2 + assert megatron["dynamic_context_parallel"]["tokens_per_rank"] == 4096 + + assert resolved["logger"]["wandb_enabled"] is True + assert resolved["logger"]["wandb"]["project"] == "nemo-rl-cp-comparison" diff --git a/tests/unit/test_dynamic_cp_moe_recipes.py b/tests/unit/test_dynamic_cp_moe_recipes.py new file mode 100644 index 00000000000..16d892ab234 --- /dev/null +++ b/tests/unit/test_dynamic_cp_moe_recipes.py @@ -0,0 +1,88 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# Licensed under the Apache License, Version 2.0 (the "License"); +"""Resolve the real-model dynamic-CP MoE smoke recipes and validate topology.""" + +from pathlib import Path + +import pytest +from omegaconf import OmegaConf + +from nemo_rl.models.policy.dynamic_cp import ( + _minimum_cp_size_for_experts, + validate_dynamic_cp, +) +from nemo_rl.utils.config import load_config, register_omegaconf_resolvers + + +RECIPE_DIR = ( + Path(__file__).resolve().parents[2] + / "examples/configs/recipes/llm/performance" +) +RECIPES = ( + ( + "grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-quick.yaml", + "aux_loss", + 8, + 4, + 2, + ), + ( + "grpo-qwen3-235b-32n4g-async-1off-megatron-dynamiccp-quick.yaml", + "seq_aux_loss", + 2, + 32, + 16, + ), + ( + "grpo-nemotron3-nano-30ba3b-8n4g-megatron-dynamiccp-quick.yaml", + "none", + 4, + 8, + 8, + ), +) + + +@pytest.mark.parametrize( + "recipe_name,routing_type,effective_minimum,total_nodes,expected_policy_nodes", + RECIPES, +) +def test_dynamic_cp_moe_recipe_topology( + recipe_name, + routing_type, + effective_minimum, + total_nodes, + expected_policy_nodes, +): + register_omegaconf_resolvers() + resolved = OmegaConf.to_container( + load_config(RECIPE_DIR / recipe_name), resolve=True + ) + assert isinstance(resolved, dict) + policy = resolved["policy"] + megatron = policy["megatron_cfg"] + generation = policy["generation"]["colocated"] + cluster = resolved["cluster"] + generation_nodes = generation["resources"]["num_nodes"] or 0 + policy_nodes = ( + cluster["num_nodes"] + if generation["enabled"] + else cluster["num_nodes"] - generation_nodes + ) + policy_world_size = policy_nodes * cluster["gpus_per_node"] + lanes = policy_world_size // ( + megatron["tensor_model_parallel_size"] + * megatron["pipeline_model_parallel_size"] + ) + + assert megatron["pipeline_model_parallel_size"] == 1 + assert megatron["context_parallel_size"] == 1 + assert megatron["expert_tensor_parallel_size"] == 1 + assert megatron["moe_router_load_balancing_type"] == routing_type + assert _minimum_cp_size_for_experts(megatron, 1) == effective_minimum + assert cluster["num_nodes"] == total_nodes + assert policy_nodes == expected_policy_nodes + assert resolved["logger"]["wandb_enabled"] is True + assert resolved["logger"]["wandb"]["project"] == "nemo-rl" + assert resolved["logger"]["wandb"]["name"] + validate_dynamic_cp(policy, lanes=lanes) From 9dcc56b5443a4a503c2e83443dd767fd49110b40 Mon Sep 17 00:00:00 2001 From: humairafirdowse18 Date: Mon, 21 Sep 2026 11:08:50 -0700 Subject: [PATCH 4/7] loss-correction tp-cp groups --- .../megatron-dynamic-context-parallel.md | 117 ++++++++- ...-async-1off-megatron-dynamiccp-10step.yaml | 55 +++++ ...g-async-1off-megatron-staticcp-10step.yaml | 11 + .../distributed/dynamic_context_parallel.py | 27 +- nemo_rl/distributed/tensor_serialization.py | 9 +- nemo_rl/models/megatron/common.py | 32 ++- nemo_rl/models/megatron/dynamic_cp.py | 81 +++++- nemo_rl/models/megatron/train.py | 33 +++ nemo_rl/models/policy/dynamic_cp.py | 61 ++++- nemo_rl/models/policy/lm_policy.py | 8 +- .../policy/workers/megatron_policy_worker.py | 233 +++++++++++------- .../functional/dynamic_cp_attention_parity.py | 20 +- tests/functional/dynamic_cp_loss_parity.py | 20 +- .../test_dynamic_context_parallel.py | 54 ++++ .../distributed/test_dynamic_cp_dispatch.py | 68 +++++ .../distributed/test_tensor_serialization.py | 8 + .../models/megatron/test_dynamic_cp_moe.py | 51 ++++ .../unit/models/megatron/test_moe_metrics.py | 47 ++++ tests/unit/models/megatron/test_train.py | 34 +++ .../models/policy/test_megatron_worker.py | 85 +++++++ .../models/policy/test_policy_validation.py | 38 +++ .../test_dynamic_cp_comparison_recipes.py | 54 +++- 22 files changed, 1003 insertions(+), 143 deletions(-) create mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-async-1off-megatron-dynamiccp-10step.yaml create mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-async-1off-megatron-staticcp-10step.yaml diff --git a/docs/design-docs/megatron-dynamic-context-parallel.md b/docs/design-docs/megatron-dynamic-context-parallel.md index 28cb7d611f2..d3f7d802a75 100644 --- a/docs/design-docs/megatron-dynamic-context-parallel.md +++ b/docs/design-docs/megatron-dynamic-context-parallel.md @@ -47,7 +47,9 @@ lanes contribute, matching the balanced hybrid-CP behavior. This path currently supports PP=1 and the standard Ray policy data path. TransferQueue, split execution, model-owned multimodal packing, atomic preference pairs, MTP, draft training, fused linear logprobs, and training CUDA graphs are -not supported. Omit or disable the configuration for existing static behavior. +not supported. Dynamic batching and HybridEP input prepadding are also disabled +because the driver plan must remain the sole owner of packing and microbatch +boundaries. Omit or disable the configuration for existing static behavior. For MoE, the minimum active size is raised until `active_CP * TP >= EP`. Workers also check that each actual expert communication group is contained within its @@ -55,12 +57,18 @@ task's ranks. This keeps every EP collective inside one active-CP task and preve experts from communicating across independently scheduled CP blocks. Router auxiliary losses and per-layer MoE metrics stay attached to the real microbatch that produced them; placeholder tasks carry zero valid tokens and do not affect -loss normalization. +loss normalization. `global_aux_loss`, expert tensor parallelism, quantile +balancing, and overlapped MoE microbatch execution are rejected because their +collective or scheduling domains cannot follow the active task safely. ## Dispatch and execution -One scheduling lane contains a complete TP replica. The driver packs sequences -of the same required CP size and assigns each packed task to an aligned, +One scheduling lane contains a complete TP replica. The driver visits sequences +in descending length order. Before opening a task at a sample's minimum CP size, +it tries spare capacity in already-required larger-CP tasks, then existing tasks +of the same size. Each admission uses the destination task's padding and token +budget. This avoids leaving holes beside long samples while making extra model +calls for short samples. It assigns each packed task to an aligned, contiguous group of lanes. The worker verifies those lane IDs against MCore's initialized DP × CP group and resolves the active attention group using MCore's hybrid-group API. TP ranks receive the same task payload. EP is a constraint on @@ -105,8 +113,8 @@ collection. The model receives a shallow copy of packed metadata with a real singleton group for CP=1, because RoPE interprets `None` as a static-group fallback. Loss and logprob code retain the explicit size-one/None convention. -During worker setup, dynamic-CP Megatron actors register a Ray serializer for -CPU tensor results. Importing MCore in the worker replaces +During worker setup, only dynamic-CP Megatron actors register a Ray serializer +for tensor results. Importing MCore in the worker replaces `torch.storage._load_from_bytes` with MCore's safe loader. A tensor does not store a Megatron object, but PyTorch's normal storage pickle records that loader function by module name; deserializing such a result would therefore make the @@ -118,8 +126,9 @@ the worker has copied result tensors to CPU and before the driver's that return value. In the GRPO recipe this includes the full policy-logprob and reference-logprob result rounds and the small loss/gradient metric tensors from each training step. It encodes contiguous bytes, dtype, and shape in a NumPy -payload, including BF16 and empty tensors, then reconstructs an independent, -writable CPU tensor in the driver. It does not modify MCore's checkpoint loader +payload, including BF16 and empty tensors. Any unexpected CUDA result is staged +to host memory instead of failing serialization, and the driver reconstructs an +independent, writable CPU tensor. It does not modify MCore's checkpoint loader or serialize model parameters and GPU activations. ## Packing, outputs, and normalization @@ -137,6 +146,11 @@ rows. The driver keeps those rows once and restores original sample order, checking for missing or duplicate sample IDs. A static CP-rank-zero filter would discard valid results. +Score workers execute every step in a multi-global-batch plan and concatenate +their task outputs before owner-row selection. They do not assume +`plan.steps[0]`; this keeps policy, reference, and top-k scoring aligned with the +training-sized schedule batches cached by the driver. + The driver computes valid sequence and token denominators from the unique global batch before replication, separately for every optimizer step. The differentiable CP logprob gather replicates the loss over the active CP group. @@ -154,6 +168,17 @@ global aggregation. `tests/unit/models/megatron/test_dynamic_cp_scaling.py` checks this multiplier against the installed MCore loss callback. A future MCore pin must pass that contract test and distributed parity before adoption. +For MoE load balancing, MCore's per-token path assumes +`local_valid_tokens * TP_CP_size` equals the active task's token count. Dynamic +packing can leave different valid-token counts on participating shards, so the +worker sums the exact valid count over the active TP×CP group and applies a +per-shard correction through `moe_grad_scale_func`. MCore's z-loss coefficient +and attachment factors already cancel to form each rank's valid-token sum; its +temporary coefficient is inversely adjusted so the shared autograd scaler does +not change that gradient. For reporting, ordinary aux +metrics are divided by unique real tasks, while z-loss reproduces MCore's +average over every real TP×CP rank participation. + ## GB200 five-step smoke test From the RL checkout on the Slurm login node: @@ -244,7 +269,8 @@ run stays at CP2. The default workload has an 8192-token ceiling, 4096 tokens per rank, and a global batch of 512 formed from 16 prompts times 32 generations. The launcher -uses the batch QoS and a six-hour limit. Use different `CP_RUN_NAME` values so +uses partition `batch`, inherits the account's default QoS, and requests four +hours by default (`TIME_LIMIT` overrides it). Use different `CP_RUN_NAME` values so the W&B runs and local logs remain distinct. The launcher also accepts `CP_NUM_STEPS`, `CP_TRAIN_GLOBAL_BATCH_SIZE`, @@ -255,6 +281,12 @@ A static CP1 run remains useful as an unconstrained throughput and CP1 sanity reference when it fits in memory, but it is not the capacity-matched baseline for 8192 tokens at the 4096-token budget. +This is a match to the configured memory budget, not proof that CP is necessary +on GB200. Measure peak memory and test static CP1 before concluding that 8192 +tokens require CP2. If CP1 fits and performs better, use that as the practical +baseline; increase `CP_TOKENS_PER_RANK` to the measured training-safe budget. +Do not reduce the budget just to make the scheduler report more CP sizes. + The correctness smoke keeps the original TP1/EP8 topology and is therefore forced to CP8. The performance pair uses TP4/EP4 specifically to expose an adaptive CP1/CP2 choice on the same eight policy GPUs. Both sides of the pair @@ -265,3 +297,70 @@ After each run, `perf_runs/analyze_cp_sequence_lengths.py` writes `sequence_length_distribution.json` beside the driver log. It reports length percentiles, how many samples stopped exactly at the configured ceiling, and the CP size each sample required before optional idle-lane expansion. + +Dynamic schedule logs also report `tasks_by_cp` and `packing_utilization`. +Required CP for an individual sample can differ from its scheduled CP when it +fills an existing larger task or when spare lanes help process a task. Performance +runs do not require a fixed mixture of sizes. For an explicit coverage test, set +`CP_REQUIRE_SIZES="1 2"` (or `"1 2 4"`). The correctness checks still require all +steps, finite metrics, score/train agreement, and schedule reuse. + +Both comparison modes run four-GPU loss and attention parity checks before +training. These tests explicitly retain CP1 coverage after cross-size packing, +including a model initialized at static CP2. Set `CP_GPU_PREFLIGHT=0` only when +reusing validation of the same code/container. Preflight time is outside the +reported training-step timings. + +### Qwen30B packing regression: jobs 7202770 and 7203284 + +Both ten-step jobs passed their smoke checks and processed approximately 26.6M +tokens. Excluding step one, the dynamic job averaged 525.65 seconds per step; +static CP2 averaged 483.72 seconds, so dynamic took 8.67% longer. Training took +390.62 versus 357.51 seconds and policy/reference scoring took 130.03 versus +121.26 seconds. These are independent async rollouts, not identical token batches. + +The old scheduler packed each required CP size separately. At 4096 tokens/rank, +`[6000, 2000]` became two calls even though both samples fit in one 8192-token CP2 +pack. The cross-size packing fix admits the 2000-token sequence into that +already-required call, checks CP2 alignment, and still leaves independent CP1 +tasks when there is remaining short work. EP containment and loss scaling are +unchanged; the worker derives its microbatch count from the resulting plan. + +CPU replay of the ten saved dynamic batches changes the sum of local calls per +lane from 3630 to 3330; static MFFD packing of those same lengths needs 3322. +That is an 8.26% reduction in model calls, not a measured GPU speedup. Rerun the +pair to measure actual performance. A near-zero gain remains possible because +most work in this workload still executes at CP2 after efficient packing. + +Reproduce the recorded timings and replay the current planner: + +```bash +uv run --no-sync python perf_runs/analyze_cp_comparison.py \ + logs/qwen30-gbs512-seq8192-dyncp-20260916T220850Z \ + logs/qwen30-gbs512-seq8192-staticcp2-20260916T220850Z \ + --lanes 2 --tp 4 --tokens-per-rank 4096 +``` + +### Dense Qwen3-32B comparison + +Set `CP_MODEL=qwen32b` to select the new dense recipes. They use four nodes total: +two policy nodes (TP2/PP1/EP1, four CP scheduling lanes) and two generation nodes +(vLLM TP2). Defaults are 16384 total tokens, 4096 tokens/rank, GBS512, ten steps, +activation checkpointing, W&B online, dynamic CP1/2/4 versus static CP4. Both +modes inherit the same workload, optimizer, precision, and generation settings. + +```bash +PAIR_TAG=$(date -u +%Y%m%dT%H%M%SZ) +CP_MODEL=qwen32b CP_MODE=dynamic TIME_LIMIT=06:00:00 \ + CP_RUN_NAME="qwen32-gbs512-seq16384-dynamic-${PAIR_TAG}" DRY_RUN=0 \ + bash perf_runs/run_gb200_cp_comparison.sh +CP_MODEL=qwen32b CP_MODE=static STATIC_CP_SIZE=4 TIME_LIMIT=06:00:00 \ + CP_RUN_NAME="qwen32-gbs512-seq16384-static4-${PAIR_TAG}" DRY_RUN=0 \ + bash perf_runs/run_gb200_cp_comparison.sh +``` + +For a smaller first validation set `CP_NUM_STEPS=2 CP_NUM_PROMPTS_PER_STEP=2 +CP_TRAIN_GLOBAL_BATCH_SIZE=64 QOS=short TIME_LIMIT=01:00:00` and use a distinct +run name. Ten steps are an initial performance sample, not +convergence validation. Report training/scoring separately from total step time +and repeat close results before claiming a speedup. diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-async-1off-megatron-dynamiccp-10step.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-async-1off-megatron-dynamiccp-10step.yaml new file mode 100644 index 00000000000..d3114632266 --- /dev/null +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-async-1off-megatron-dynamiccp-10step.yaml @@ -0,0 +1,55 @@ +defaults: ./grpo-qwen3-32b-8n4g-async-1off.yaml +grpo: + num_prompts_per_step: 16 + num_generations_per_prompt: 32 + max_num_steps: 10 + val_period: 1000 + val_at_start: false + val_at_end: false +checkpointing: + enabled: false +policy: + train_global_batch_size: 512 + train_micro_batch_size: 1 + logprob_batch_size: 1 + max_total_sequence_length: 16384 + make_sequence_length_divisible_by: 2 + megatron_cfg: + tensor_model_parallel_size: 2 + pipeline_model_parallel_size: 1 + context_parallel_size: 1 + expert_model_parallel_size: 1 + sequence_parallel: true + activation_checkpointing: true + dynamic_context_parallel: + enabled: true + min_size: 1 + max_size: 4 + tokens_per_rank: 4096 + fp8_cfg: + enabled: false + sequence_packing: + enabled: true + dynamic_batching: + enabled: false + generation: + colocated: + enabled: false + resources: + num_nodes: 2 + gpus_per_node: 4 + vllm_cfg: + tensor_parallel_size: 2 +logger: + log_dir: logs/grpo-qwen3-32b-4n4g-async-dynamiccp-10step + wandb_enabled: true + tensorboard_enabled: true + wandb: + project: nemo-rl-cp-comparison + name: qwen3-32b-4n4g-async-dynamiccp-10step +cluster: + num_nodes: 4 + gpus_per_node: 4 + segment_size: 2 +data_plane: + enabled: false diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-async-1off-megatron-staticcp-10step.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-async-1off-megatron-staticcp-10step.yaml new file mode 100644 index 00000000000..5b73fc1cac6 --- /dev/null +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-async-1off-megatron-staticcp-10step.yaml @@ -0,0 +1,11 @@ +defaults: ./grpo-qwen3-32b-4n4g-async-1off-megatron-dynamiccp-10step.yaml +policy: + make_sequence_length_divisible_by: 16 + megatron_cfg: + context_parallel_size: 4 + dynamic_context_parallel: + enabled: false +logger: + log_dir: logs/grpo-qwen3-32b-4n4g-async-staticcp-10step + wandb: + name: qwen3-32b-4n4g-async-staticcp-10step diff --git a/nemo_rl/distributed/dynamic_context_parallel.py b/nemo_rl/distributed/dynamic_context_parallel.py index 4897fec984c..2341693ce78 100644 --- a/nemo_rl/distributed/dynamic_context_parallel.py +++ b/nemo_rl/distributed/dynamic_context_parallel.py @@ -317,14 +317,27 @@ def plan_cp_phases( raise ValueError( f"Sample {index} of length {length} exceeds the dynamic CP token budget" ) - size_bins = bins.setdefault(size, []) - for bin_index, (members, used, factor) in enumerate(size_bins): - if used + padded <= tokens_per_rank * size: - members.append(index) - size_bins[bin_index] = (members, used + padded, factor) + # Fill already-required larger-CP calls before opening another call. + # Keeping separate size buckets strands space: e.g. [6000, 2000] at + # 4096 tokens/rank used to require two calls instead of one CP2 pack. + # The destination determines alignment; CP1 padding is not sufficient + # when the short sequence joins a CP2+ task. + placed = False + for target_size in sorted(bins, reverse=True): + if target_size < size: + continue + size_bins = bins[target_size] + for bin_index, (members, used, factor) in enumerate(size_bins): + target_padded = (length + factor - 1) // factor * factor + if used + target_padded <= tokens_per_rank * target_size: + members.append(index) + size_bins[bin_index] = (members, used + target_padded, factor) + placed = True + break + if placed: break - else: - size_bins.append(([index], padded, multiple)) + if not placed: + bins.setdefault(size, []).append(([index], padded, multiple)) pending = [ CPAssignment(tuple(members), 0, size, factor, used) diff --git a/nemo_rl/distributed/tensor_serialization.py b/nemo_rl/distributed/tensor_serialization.py index e6d2fd1a26c..54363acf312 100644 --- a/nemo_rl/distributed/tensor_serialization.py +++ b/nemo_rl/distributed/tensor_serialization.py @@ -17,14 +17,15 @@ def tensor_to_payload(tensor: torch.Tensor) -> TensorPayload: - """Encode a CPU RPC tensor without invoking Torch's storage pickler. + """Encode an RPC tensor without invoking Torch's storage pickler. MCore can replace that pickler's loader with a Megatron function, which cannot be imported in the Ray driver. Byte arrays also preserve BF16. + CUDA tensors are staged to host memory and, like all actor return tensors + using this transport, are restored as CPU tensors in the driver. """ - if tensor.device.type != "cpu": - raise ValueError("Policy RPC tensors must be moved to CPU before serialization") - data = tensor.detach().contiguous().reshape(-1).view(torch.uint8).numpy() + host_tensor = tensor.detach().to(device="cpu").contiguous() + data = host_tensor.reshape(-1).view(torch.uint8).numpy() return data, tensor.dtype, tuple(tensor.shape) diff --git a/nemo_rl/models/megatron/common.py b/nemo_rl/models/megatron/common.py index 885ebd1d90b..e8ac29335f5 100644 --- a/nemo_rl/models/megatron/common.py +++ b/nemo_rl/models/megatron/common.py @@ -202,6 +202,7 @@ def get_moe_metrics( mtp_num_layers: Optional[int] = None, track_names: Optional[list[str]] = None, dynamic_parallel_group: Optional[dist.ProcessGroup] = None, + dynamic_avg_loss_scale: Optional[float] = None, ) -> dict[str, Any]: """Returns Mixture of Experts (MoE) auxiliary-loss metrics. @@ -227,6 +228,12 @@ def get_moe_metrics( Dynamic lanes can execute different numbers of microbatches and the router's per-forward TP*CP group changes with each task, so its last recorded reduction group is not safe for end-of-step metric sync. + dynamic_avg_loss_scale: Scale for dynamic metrics recorded with + ``avg_group`` (currently z-loss). These values have one contribution + per participating TP*CP rank rather than one partial sum per task. + The caller supplies the reciprocal of the global real-task rank + participation count. Required with ``dynamic_parallel_group`` when + an averaged metric is present. Returns: dict[str, Any]: A flat dict of aggregated metrics. For each aux loss name, @@ -264,6 +271,7 @@ def get_moe_metrics( mcore_tracker.ensure_initialized(name, tracker_num_layers) dynamic_names: Optional[list[str]] = None + dynamic_avg_names: set[str] = set() if dynamic_parallel_group is None: reduce_aux_losses_tracker_across_ranks() else: @@ -278,6 +286,11 @@ def get_moe_metrics( for name in dynamic_names: entry = mcore_tracker.metrics.get(name) if entry is not None: + # Padding-only lanes receive a pre-initialized z-loss entry but + # never call record(), so its avg_group remains None locally. + # z_loss is nevertheless an averaged metric on every lane. + if name == "z_loss" or getattr(entry, "avg_group", None) is not None: + dynamic_avg_names.add(name) dist.all_reduce(entry.values, group=dynamic_parallel_group) tracker = get_moe_layer_wise_logging_tracker() if dynamic_names is not None: @@ -285,7 +298,24 @@ def get_moe_metrics( metrics: dict[str, Any] = {} if len(tracker) > 0: - aux_losses = {k: v["values"].float() * loss_scale for k, v in tracker.items()} + if dynamic_avg_names and dynamic_avg_loss_scale is None: + raise ValueError( + "Dynamic averaged MoE metrics require a rank-participation scale" + ) + resolved_dynamic_avg_loss_scale = ( + dynamic_avg_loss_scale + if dynamic_avg_loss_scale is not None + else loss_scale + ) + aux_losses = { + name: value["values"].float() + * ( + resolved_dynamic_avg_loss_scale + if name in dynamic_avg_names + else loss_scale + ) + for name, value in tracker.items() + } for name, loss_list in aux_losses.items(): # Megatron-LM aggregates aux losses across layers and normalizes by number of MoE layers num_tracked_layers = int(loss_list.numel()) if loss_list.numel() > 0 else 1 diff --git a/nemo_rl/models/megatron/dynamic_cp.py b/nemo_rl/models/megatron/dynamic_cp.py index 425cf0a0165..57572f48565 100644 --- a/nemo_rl/models/megatron/dynamic_cp.py +++ b/nemo_rl/models/megatron/dynamic_cp.py @@ -25,6 +25,7 @@ _DYNAMIC_TP_CP_GROUPS: dict[int, Any] = {} _ROUTER_CONFIG_BASELINES: dict[int, tuple[Any, Any]] = {} +_DYNAMIC_MOE_GRAD_SCALE_CORRECTIONS: dict[int, float] = {} @dataclass(frozen=True) @@ -155,20 +156,21 @@ def preserve_attention_cp_groups(model: torch.nn.Module) -> Iterator[None]: if _is_gated_delta_net(module) ] saved_routers = [ - (module, module.tp_cp_group) + (module, module.cp_group, module.tp_cp_group) for module in model.modules() if isinstance(module, Router) ] router_configs: dict[int, Any] = {} for module, collection in saved_collections: module.pg_collection = copy(collection) - for module, _ in saved_routers: + for module, _, _ in saved_routers: config = module.config router_configs[id(config)] = config _ROUTER_CONFIG_BASELINES[id(config)] = ( config.moe_aux_loss_coeff, config.moe_z_loss_coeff, ) + _DYNAMIC_MOE_GRAD_SCALE_CORRECTIONS[id(config)] = 1.0 try: yield finally: @@ -178,10 +180,12 @@ def preserve_attention_cp_groups(model: torch.nn.Module) -> Iterator[None]: module.cp = cp for module, cp_size in saved_gdn: module.cp_size = cp_size - for module, group in saved_routers: - module.tp_cp_group = group + for module, cp_group, tp_cp_group in saved_routers: + module.cp_group = cp_group + module.tp_cp_group = tp_cp_group for config_id, config in router_configs.items(): aux_coeff, z_coeff = _ROUTER_CONFIG_BASELINES.pop(config_id) + _DYNAMIC_MOE_GRAD_SCALE_CORRECTIONS.pop(config_id, None) config.moe_aux_loss_coeff = aux_coeff config.moe_z_loss_coeff = z_coeff @@ -201,6 +205,74 @@ def _bind_router_config(router: Router, *, padding_only: bool) -> None: z_coeff = None router.config.moe_aux_loss_coeff = aux_coeff router.config.moe_z_loss_coeff = z_coeff + _DYNAMIC_MOE_GRAD_SCALE_CORRECTIONS[id(router.config)] = 1.0 + + +def _has_positive_coefficient(value: Any) -> bool: + """Return whether a scalar or coefficient list enables an aux loss.""" + values = value if isinstance(value, (list, tuple)) else (value,) + return any(isinstance(item, (int, float)) and item > 0 for item in values) + + +def configure_dynamic_moe_loss_scaling( + model: torch.nn.Module, padding_mask: torch.Tensor | None +) -> None: + """Correct MCore's fixed-shard aux-loss scaling for the active task. + + MCore multiplies each rank's load-balancing loss by + ``local_valid_tokens * tp_cp_group.size()``. That equals the task-wide token + count only when every TP*CP shard contains the same number of valid tokens. + Packed dynamic tasks do not have that invariant. Compute the exact active + group token count and expose a per-rank correction through MCore's existing + ``moe_grad_scale_func`` hook. + + The same autograd scaler carries z-loss. MCore's z-loss coefficient and + attachment factors already cancel to produce each rank's valid-token sum, + so inversely adjust the temporary coefficient to keep that gradient + unchanged when the shared scaler applies the aux correction. + """ + routers = [module for module in model.modules() if isinstance(module, Router)] + if not routers: + return + + configs = {id(router.config): router.config for router in routers} + for config_id in configs: + _DYNAMIC_MOE_GRAD_SCALE_CORRECTIONS[config_id] = 1.0 + + if not model.training or not torch.is_grad_enabled(): + return + if not any( + _has_positive_coefficient(router.config.moe_aux_loss_coeff) + for router in routers + ): + return + if padding_mask is None: + raise ValueError("Dynamic CP MoE aux loss requires a packed padding mask") + + tp_cp_group = routers[0].tp_cp_group + if any(router.tp_cp_group is not tp_cp_group for router in routers[1:]): + raise ValueError("Dynamic CP routers disagree on the active TP*CP group") + active_size = tp_cp_group.size() + local_valid_tokens = (~padding_mask).sum() + group_valid_tokens = local_valid_tokens.clone() + torch.distributed.all_reduce(group_valid_tokens, group=tp_cp_group) + local_count = int(local_valid_tokens.item()) + correction = ( + float(group_valid_tokens.item()) / (local_count * active_size) + if local_count > 0 + else 1.0 + ) + + for config_id, config in configs.items(): + _DYNAMIC_MOE_GRAD_SCALE_CORRECTIONS[config_id] = correction + _, baseline_z_coeff = _ROUTER_CONFIG_BASELINES[config_id] + if isinstance(baseline_z_coeff, (int, float)): + config.moe_z_loss_coeff = baseline_z_coeff / correction + + +def dynamic_moe_grad_scale_correction(model_config: Any) -> float: + """Return the active task's aux-loss correction, or the static default.""" + return _DYNAMIC_MOE_GRAD_SCALE_CORRECTIONS.get(id(model_config), 1.0) def bind_attention_cp_group(model: torch.nn.Module, packed_seq_params: Any) -> Any: @@ -233,6 +305,7 @@ def bind_attention_cp_group(model: torch.nn.Module, packed_seq_params: Any) -> A if isinstance(module, Attention): module.pg_collection.cp = group if isinstance(module, Router): + module.cp_group = group module.tp_cp_group = tp_cp_group _bind_router_config(module, padding_only=padding_only) if _is_mamba_mixer(module): diff --git a/nemo_rl/models/megatron/train.py b/nemo_rl/models/megatron/train.py index e84f2889f42..1aadfc506e8 100644 --- a/nemo_rl/models/megatron/train.py +++ b/nemo_rl/models/megatron/train.py @@ -68,6 +68,7 @@ from nemo_rl.models.megatron.dynamic_cp import ( RuntimeCPContext, bind_attention_cp_group, + configure_dynamic_moe_loss_scaling, preserve_attention_cp_groups, runtime_cp_from_packed, ) @@ -116,6 +117,13 @@ def _prepare_padding_mask_for_model( if isinstance(core_model, GPTModel) and core_model.pre_process: return padding_mask + return _scatter_padding_mask_to_sequence_parallel_region(padding_mask) + + +def _scatter_padding_mask_to_sequence_parallel_region( + padding_mask: torch.Tensor, +) -> torch.Tensor: + """Scatter a batch-first mask exactly like MCore's GPT preprocessing.""" return ( tensor_parallel.scatter_to_sequence_parallel_region( padding_mask.transpose(0, 1).contiguous(), @@ -126,6 +134,25 @@ def _prepare_padding_mask_for_model( ) +def _prepare_padding_mask_for_router_scaling( + model: GPTModel, + padding_mask: Optional[torch.Tensor], +) -> Optional[torch.Tensor]: + """Return the TP-local mask seen by the router for token normalization.""" + if padding_mask is None or not get_model_config(model).sequence_parallel: + return padding_mask + + core_model = unwrap_model(model) + if isinstance(core_model, GPTModel) and core_model.pre_process: + # GPTModel scatters its forward mask internally. Reproduce that scatter + # here so the dynamic MoE scale counts this TP rank's actual tokens. + return _scatter_padding_mask_to_sequence_parallel_region(padding_mask) + + # Other models and non-embedding GPT stages received a mask already prepared + # by _prepare_padding_mask_for_model. + return padding_mask + + @contextmanager def suspend_activation_offload_for_forward_only( model: Union[GPTModel, List[GPTModel]], forward_only: bool @@ -227,6 +254,12 @@ def model_forward( padding_mask = _prepare_padding_mask_for_model(model, padding_mask) if padding_mask is not None: additional_kwargs["padding_mask"] = padding_mask + if packed_seq_params is not None and isinstance( + getattr(packed_seq_params, "local_cp_size", None), int + ): + configure_dynamic_moe_loss_scaling( + model, _prepare_padding_mask_for_router_scaling(model, padding_mask) + ) # Only sent when the model advertises the parameter, so it never reaches a # forward that would swallow it into **kwargs and quietly ignore it. diff --git a/nemo_rl/models/policy/dynamic_cp.py b/nemo_rl/models/policy/dynamic_cp.py index dc5a3f6b09e..0b3003f85d7 100644 --- a/nemo_rl/models/policy/dynamic_cp.py +++ b/nemo_rl/models/policy/dynamic_cp.py @@ -30,19 +30,28 @@ logger = logging.getLogger(__name__) +def _model_setting(megatron_cfg: dict[str, Any], name: str) -> Any: + """Resolve a model field after applying the Bridge override layer.""" + overrides = megatron_cfg.get("model_overrides") or {} + return overrides[name] if name in overrides else megatron_cfg.get(name) + + +def _routing_types(megatron_cfg: dict[str, Any]) -> list[Any]: + routing_type = _model_setting(megatron_cfg, "moe_router_load_balancing_type") + return ( + list(routing_type) + if isinstance(routing_type, (list, tuple)) + else [routing_type] + ) + + def _enabled_global_aux_loss(megatron_cfg: dict[str, Any]) -> bool: """Return whether a full-DP global aux collective is configured. The coefficient can originate in the HF model provider rather than this dictionary, so the routing type itself must be rejected. """ - routing_type = megatron_cfg.get("moe_router_load_balancing_type") - routing_types = ( - list(routing_type) - if isinstance(routing_type, (list, tuple)) - else [routing_type] - ) - return "global_aux_loss" in routing_types + return "global_aux_loss" in _routing_types(megatron_cfg) def _minimum_cp_size_for_experts( @@ -59,9 +68,7 @@ def _minimum_cp_size_for_experts( tensor_parallel = megatron_cfg["tensor_model_parallel_size"] expert_tensor_parallel = megatron_cfg.get("expert_tensor_parallel_size", 1) if expert_tensor_parallel != 1: - raise ValueError( - "Dynamic CP MoE requires expert_tensor_parallel_size=1" - ) + raise ValueError("Dynamic CP MoE requires expert_tensor_parallel_size=1") if expert_parallel <= 1: return configured_minimum minimum = configured_minimum @@ -103,8 +110,14 @@ def validate_dynamic_cp(cfg: dict[str, Any], *, lanes: int) -> None: ) if mc.get("cuda_graph_impl") not in (None, "none"): raise ValueError("Dynamic CP does not support CUDA graph capture") - if mc.get("mtp_num_layers") or mc.get("moe_hybridep_prepad_packed_inputs"): + if _model_setting(mc, "mtp_num_layers") or _model_setting( + mc, "moe_hybridep_prepad_packed_inputs" + ): raise ValueError("Dynamic CP does not support MTP or HybridEP input prepadding") + if _model_setting(mc, "overlap_moe_expert_parallel_comm"): + raise ValueError( + "Dynamic CP does not support overlap_moe_expert_parallel_comm" + ) if cfg["sequence_packing"].get("pair_grouping_key"): raise ValueError("Dynamic CP does not yet schedule atomic preference pairs") if _enabled_global_aux_loss(mc): @@ -113,6 +126,11 @@ def validate_dynamic_cp(cfg: dict[str, Any], *, lanes: int) -> None: "TP*DP*CP collective cannot be called by uneven CP task lists. " "Use aux_loss or seq_aux_loss instead." ) + if "quantile_balancing" in _routing_types(mc): + raise ValueError( + "Dynamic CP does not support quantile_balancing because its router " + "rejects the packed padding mask required for correct token counts" + ) # Probe even empty plans, so invalid domains/bounds fail at initialization. tp = mc["tensor_model_parallel_size"] minimum = _minimum_cp_size_for_experts(mc, dynamic.min_size) @@ -169,6 +187,16 @@ def owned_real_task_count(plan: CPRankPlan) -> int: ) +def real_task_participation_count(plan: CPRankPlan) -> int: + """Count real task calls made by this lane across all optimizer steps.""" + return sum( + 1 + for step in plan.steps + for task in step.assignments + if task.sample_indices + ) + + def _schedule_parameters( cfg: dict[str, Any], sharding: NamedSharding ) -> tuple[int, int, int, int, int, int, int]: @@ -341,9 +369,16 @@ def build_cp_dispatch( else: valid_sequences = valid_tokens = 0.0 samples_by_cp = Counter() + tasks_by_cp = Counter() + packed_tokens = 0 + packed_capacity = 0 for group in groups: for task in group.assignments: samples_by_cp[task.cp_size] += len(task.sample_indices) + if task.sample_indices: + tasks_by_cp[task.cp_size] += 1 + packed_tokens += task.padded_tokens + packed_capacity += task.cp_size * schedule.tokens_per_rank group_task_ranges = [ ( min(len(assignments_for_lane(group, lane)) for lane in range(lanes)), @@ -354,7 +389,7 @@ def build_cp_dispatch( logger.info( "Dynamic CP %s: samples=%d groups=%d uneven_groups=%d " "local_tasks=[%d,%d] " - "samples_by_cp=%s " + "samples_by_cp=%s tasks_by_cp=%s packing_utilization=%.4f " "valid_sequences=%s valid_tokens=%s schedule=%s", "train" if training else "score", gbs, @@ -369,6 +404,8 @@ def build_cp_dispatch( for lane in range(lanes) ), dict(sorted(samples_by_cp.items())), + dict(sorted(tasks_by_cp.items())), + packed_tokens / packed_capacity if packed_capacity else 0.0, valid_sequences, valid_tokens, "reused" if reused_schedule else "new", diff --git a/nemo_rl/models/policy/lm_policy.py b/nemo_rl/models/policy/lm_policy.py index 48bfc467bf5..e20cad7a041 100644 --- a/nemo_rl/models/policy/lm_policy.py +++ b/nemo_rl/models/policy/lm_policy.py @@ -674,10 +674,10 @@ def _get_dynamic_cp_outputs( self, method: str, data: BatchedDataDict, **kwargs: Any ) -> BatchedDataDict: schedule_batch_size = self.cfg["train_global_batch_size"] - if data.size != schedule_batch_size: - # The score worker currently consumes one CPRankStep. Standalone - # scoring can therefore use dynamic CP as one schedule batch, but - # it will not reuse a multi-global-batch training schedule. + if data.size % schedule_batch_size: + # A standalone score batch need not be a multiple of training GBS. + # In that case plan it as one batch; otherwise preserve every + # training-sized step so score and train can share the schedule. schedule_batch_size = data.size schedule = self._matching_dynamic_cp_schedule(data, schedule_batch_size) dispatch = build_cp_dispatch( diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index 8dc7e34ae28..28063ec8abe 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -76,6 +76,7 @@ get_microbatch_iterator, process_global_batch, ) +from nemo_rl.models.megatron.dynamic_cp import dynamic_moe_grad_scale_correction from nemo_rl.models.megatron.pipeline_parallel import ( broadcast_loss_metrics_from_last_stage, broadcast_obj_from_pp_rank, @@ -106,7 +107,11 @@ megatron_forward_backward, ) from nemo_rl.models.policy import PolicyConfig -from nemo_rl.models.policy.dynamic_cp import dynamic_cp_config, owned_real_task_count +from nemo_rl.models.policy.dynamic_cp import ( + dynamic_cp_config, + owned_real_task_count, + real_task_participation_count, +) from nemo_rl.models.policy.interfaces import ( ColocatablePolicyInterface, LogprobOutputSpec, @@ -789,16 +794,38 @@ def __init__( if _model_accepts_media_token_validity_mask(self.model) else None ) - # MCore installs a Megatron-specific torch storage loader process-wide. - # Every Megatron policy result therefore needs the portable CPU-tensor - # serializer, even when dynamic context parallelism is disabled; the - # Ray driver intentionally does not have Megatron on its import path. - register_policy_tensor_serializer() if dynamic_cp_config(self.cfg) is not None: + # Dynamic dispatch collects every DP*CP result rather than Ray's + # usual replicated-rank subset. Keep the serializer scoped to these + # workers so unrelated Megatron RPCs retain Ray's CUDA behavior. + register_policy_tensor_serializer() if self.delegate_pack_to_model or self.model_slices_context_parallel_inputs: raise ValueError("Dynamic CP requires NeMo-owned text-model packing") - if getattr(self._get_model_config(), "mtp_num_layers", None): + model_config = self._get_model_config() + if getattr(model_config, "mtp_num_layers", None): raise ValueError("Dynamic CP does not yet support MTP") + if getattr( + model_config, + "overlap_moe_expert_parallel_comm", + False, + ): + raise ValueError( + "Dynamic CP does not support overlap_moe_expert_parallel_comm" + ) + if getattr(model_config, "moe_hybridep_prepad_packed_inputs", False): + raise ValueError("Dynamic CP does not support HybridEP input prepadding") + routing_type = getattr( + model_config, "moe_router_load_balancing_type", None + ) + routing_types = ( + list(routing_type) + if isinstance(routing_type, (list, tuple)) + else [routing_type] + ) + if "global_aux_loss" in routing_types: + raise ValueError("Dynamic CP does not support global_aux_loss") + if "quantile_balancing" in routing_types: + raise ValueError("Dynamic CP does not support quantile_balancing") if self.model_slices_context_parallel_inputs: if self.delegate_pack_to_model: @@ -1114,13 +1141,10 @@ def train( self._copy_main_params_to_param_buffer() # Set moe_grad_scale_func for MoE aux-loss gradient scaling. - # With calculate_per_token_loss=True, the router pre-multiplies - # the aux loss by (num_local_tokens * tp_cp_group.size()), and - # MoEAuxLossAutoScaler applies loss_scale to the gradient. Setting - # loss_scale = 1/global_valid_toks (G = global valid token count) - # normalizes the aux gradient consistently with the main per-token - # SFT loss: - # (1/G) * N_local * tp_cp_size * aux_grad -> DDP SUM -> aux_grad / G + # Dynamic CP corrects MCore's local_tokens * TP*CP factor to + # the exact active-task token count. Static CP has correction 1. + # Dividing by global_valid_toks then gives the same global + # per-token normalization as the main loss. self._set_moe_grad_scale_func( # pragma: no cover self._compute_moe_grad_scale(global_valid_toks) ) @@ -1307,10 +1331,10 @@ def train( "grad_norm": torch.tensor([grad_norm]), "train_elapsed_seconds": metrics_train_elapsed, # pragma: no cover } - # Read "config" via getattr-by-string so the token stays out of - # train.__code__.co_names; with torch 2.11 cloudpickle otherwise - # matches torch.distributed.config (a non-pickleable ConfigModuleInstance). - model_config = getattr(self.model, "config", None) + # Keep config lookup in a helper so train.__code__.co_names does not + # include "config" (which torch 2.11 cloudpickle can mistake for + # torch.distributed.config), while still unwrapping Float16Module. + model_config = self._get_model_config() num_moe_experts = getattr(model_config, "num_moe_experts", None) if num_moe_experts is not None and num_moe_experts > 1: dynamic_moe_group = None @@ -1328,13 +1352,23 @@ def train( ), ) moe_loss_scale = 1.0 / max(1, int(global_real_tasks.item())) - dynamic_moe_group = ( - parallel_state.get_tensor_and_data_parallel_group( - with_context_parallel=True - ) + dynamic_moe_group = parallel_state.get_tensor_and_data_parallel_group( + with_context_parallel=True + ) + local_participations = real_task_participation_count(cp_plan) + global_participations = torch.tensor( + local_participations, dtype=torch.int64, device="cuda" + ) + torch.distributed.all_reduce( + global_participations, + group=dynamic_moe_group, + ) + dynamic_avg_loss_scale = 1.0 / max( + 1, int(global_participations.item()) ) else: moe_loss_scale = 1.0 / max(1, total_num_microbatches) + dynamic_avg_loss_scale = None moe_metrics = get_moe_metrics( loss_scale=moe_loss_scale, per_layer_logging=self.cfg["megatron_cfg"]["moe_per_layer_logging"], @@ -1346,6 +1380,7 @@ def train( mtp_num_layers=getattr(model_config, "mtp_num_layers", None), track_names=get_aux_loss_track_names(model_config), dynamic_parallel_group=dynamic_moe_group, + dynamic_avg_loss_scale=dynamic_avg_loss_scale, ) if moe_metrics: metrics["moe_metrics"] = moe_metrics @@ -1386,13 +1421,17 @@ def train( def _compute_moe_grad_scale(self, global_valid_toks): """Build a moe_grad_scale_func that normalizes the aux-loss gradient. - Returns a callable yielding loss_scale = 1/global_valid_toks (clamped to - avoid division by zero) so the MoE aux gradient is normalized consistently - with the main per-token SFT loss. See the call site in train() for the - full derivation. + The base scale is 1/global_valid_toks (clamped to avoid division by + zero). Dynamic CP additionally supplies the current task's exact-token + correction; static CP and dense models retain a correction of one. """ moe_scale = 1.0 / global_valid_toks.clamp(min=1).float() - return lambda: moe_scale + + def _scale() -> torch.Tensor: + model_config = self._get_model_config() if hasattr(self, "model") else None + return moe_scale * dynamic_moe_grad_scale_correction(model_config) + + return _scale def _set_moe_grad_scale_func(self, func): """Set moe_grad_scale_func on the model config for MOE aux loss scaling.""" @@ -2208,24 +2247,6 @@ def get_logprobs( # against a different media alignment than the one trained on. attach_media_token_validity_mask(data, self.media_placeholder_token_id) - ( - mb_iterator, - num_microbatches, - micro_batch_size, - seq_length, - padded_seq_length, - ) = get_microbatch_iterator( - data, - self.cfg, - logprob_batch_size, - straggler_timer=self.mcore_state.straggler_timer, - delegate_pack_to_model=self.delegate_pack_to_model, - delegate_mtp_loss_mask_to_model=self.delegate_mtp_loss_mask_to_model, - model_slices_context_parallel_inputs=self.model_slices_context_parallel_inputs, - cp_plan=cp_plan, - cp_step=cp_plan.steps[0] if cp_plan is not None else None, - ) - use_fused_linear_logprobs = self.cfg["megatron_cfg"].get( "use_fused_linear_logprobs", False ) @@ -2241,22 +2262,47 @@ def get_logprobs( require=require_router_replay, ) + cp_steps = cp_plan.steps if cp_plan is not None else (None,) + if not cp_steps: + raise ValueError("Dynamic CP score plan must contain at least one step") + list_of_logprobs = [] + seq_length = data["input_ids"].shape[1] with maybe_r3_trace_stage("prev-logprob", enabled=use_router_replay): - list_of_logprobs = megatron_forward_backward( - model=self.model, - data_iterator=mb_iterator, - seq_length=padded_seq_length, - mbs=micro_batch_size, - num_microbatches=num_microbatches, - post_processing_fn=logprobs_post_processor, - forward_only=True, - defer_fp32_logits=self.defer_fp32_logits, - sampling_params=self.sampling_params, - straggler_timer=self.mcore_state.straggler_timer, - use_fused_linear_logprobs=use_fused_linear_logprobs, - use_router_replay=use_router_replay, - router_replay_train=False, - ) + for cp_step in cp_steps: + ( + mb_iterator, + num_microbatches, + step_micro_batch_size, + seq_length, + padded_seq_length, + ) = get_microbatch_iterator( + data, + self.cfg, + logprob_batch_size, + straggler_timer=self.mcore_state.straggler_timer, + delegate_pack_to_model=self.delegate_pack_to_model, + delegate_mtp_loss_mask_to_model=self.delegate_mtp_loss_mask_to_model, + model_slices_context_parallel_inputs=self.model_slices_context_parallel_inputs, + cp_plan=cp_plan, + cp_step=cp_step, + ) + list_of_logprobs.extend( + megatron_forward_backward( + model=self.model, + data_iterator=mb_iterator, + seq_length=padded_seq_length, + mbs=step_micro_batch_size, + num_microbatches=num_microbatches, + post_processing_fn=logprobs_post_processor, + forward_only=True, + defer_fp32_logits=self.defer_fp32_logits, + sampling_params=self.sampling_params, + straggler_timer=self.mcore_state.straggler_timer, + use_fused_linear_logprobs=use_fused_linear_logprobs, + use_router_replay=use_router_replay, + router_replay_train=False, + ) + ) if parallel_state.is_pipeline_last_stage(ignore_virtual=True): all_log_probs_padded = [] @@ -2678,36 +2724,43 @@ def get_topk_logits( attach_media_token_validity_mask(data, self.media_placeholder_token_id) - ( - mb_iterator, - num_microbatches, - micro_batch_size, - seq_length, - padded_seq_length, - ) = get_microbatch_iterator( - data, - self.cfg, - logprob_batch_size, - straggler_timer=self.mcore_state.straggler_timer, - delegate_pack_to_model=self.delegate_pack_to_model, - delegate_mtp_loss_mask_to_model=self.delegate_mtp_loss_mask_to_model, - model_slices_context_parallel_inputs=self.model_slices_context_parallel_inputs, - cp_plan=cp_plan, - cp_step=cp_plan.steps[0] if cp_plan is not None else None, - ) - - list_of_outputs = megatron_forward_backward( - model=self.model, - data_iterator=mb_iterator, - seq_length=padded_seq_length, - mbs=micro_batch_size, - num_microbatches=num_microbatches, - post_processing_fn=TopkLogitsPostProcessor(cfg=self.cfg, k=k), - forward_only=True, - defer_fp32_logits=self.defer_fp32_logits, - sampling_params=self.sampling_params, - straggler_timer=self.mcore_state.straggler_timer, - ) + cp_steps = cp_plan.steps if cp_plan is not None else (None,) + if not cp_steps: + raise ValueError("Dynamic CP score plan must contain at least one step") + list_of_outputs = [] + seq_length = data["input_ids"].shape[1] + for cp_step in cp_steps: + ( + mb_iterator, + num_microbatches, + step_micro_batch_size, + seq_length, + padded_seq_length, + ) = get_microbatch_iterator( + data, + self.cfg, + logprob_batch_size, + straggler_timer=self.mcore_state.straggler_timer, + delegate_pack_to_model=self.delegate_pack_to_model, + delegate_mtp_loss_mask_to_model=self.delegate_mtp_loss_mask_to_model, + model_slices_context_parallel_inputs=self.model_slices_context_parallel_inputs, + cp_plan=cp_plan, + cp_step=cp_step, + ) + list_of_outputs.extend( + megatron_forward_backward( + model=self.model, + data_iterator=mb_iterator, + seq_length=padded_seq_length, + mbs=step_micro_batch_size, + num_microbatches=num_microbatches, + post_processing_fn=TopkLogitsPostProcessor(cfg=self.cfg, k=k), + forward_only=True, + defer_fp32_logits=self.defer_fp32_logits, + sampling_params=self.sampling_params, + straggler_timer=self.mcore_state.straggler_timer, + ) + ) if parallel_state.is_pipeline_last_stage(ignore_virtual=True): logits_chunks = [] diff --git a/tests/functional/dynamic_cp_attention_parity.py b/tests/functional/dynamic_cp_attention_parity.py index d1e333167a7..05061bd9de0 100644 --- a/tests/functional/dynamic_cp_attention_parity.py +++ b/tests/functional/dynamic_cp_attention_parity.py @@ -80,13 +80,14 @@ def main() -> None: dist.init_process_group("nccl") rank, world = dist.get_rank(), dist.get_world_size() assert world == 4 - lengths = torch.tensor([7, 45, 101, 11, 55, 9]) - ids = (torch.arange(6 * 112).reshape(6, 112) % 31 + 1).long() + # Preserve CP1 coverage even when short samples fill larger-CP packs. + lengths = torch.tensor([7, 45, 101, 11, 55, 9, 29, 29, 29, 29]) + ids = (torch.arange(len(lengths) * 112).reshape(len(lengths), 112) % 31 + 1).long() data = BatchedDataDict( input_ids=ids, input_lengths=lengths, token_mask=(torch.arange(112)[None, :] < lengths[:, None]).long(), - sample_mask=torch.tensor([1, 1, 1, 0, 1, 1]), + sample_mask=torch.tensor([1, 1, 1, 0, 1, 1, 1, 1, 1, 1]), ) for tp, base_cp in ((1, 1), (1, 2), (2, 1), (2, 2)): parallel_state.initialize_model_parallel( @@ -150,7 +151,18 @@ def main() -> None: }, }, } - dispatch = build_cp_dispatch(data, cfg, mesh, batch_size=6, training=True) + dispatch = build_cp_dispatch( + data, cfg, mesh, batch_size=data.size, training=True + ) + active_sizes = { + task.cp_size + for dp_plans in dispatch.plans + for rank_plan in dp_plans + for rank_step in rank_plan.steps + for task in rank_step.assignments + if task.sample_indices + } + assert active_sizes == ({1, 2, 4} if tp == 1 else {1, 2}), active_sizes lane = rank // tp plan = dispatch.plans[lane // base_cp][lane % base_cp] payload = dispatch.data[lane // base_cp][lane % base_cp] diff --git a/tests/functional/dynamic_cp_loss_parity.py b/tests/functional/dynamic_cp_loss_parity.py index 983a50b9c8c..9b3450abf04 100644 --- a/tests/functional/dynamic_cp_loss_parity.py +++ b/tests/functional/dynamic_cp_loss_parity.py @@ -45,13 +45,14 @@ def main() -> None: context_parallel_size=1, hybrid_context_parallel=True, ) - lengths = torch.tensor([7, 45, 101, 11, 55, 9]) - ids = (torch.arange(6 * 112).reshape(6, 112) % 31 + 1).long() + # Leave enough short work after filling long packs to exercise CP1 too. + lengths = torch.tensor([7, 45, 101, 11, 55, 9, 29, 29, 29, 29]) + ids = (torch.arange(len(lengths) * 112).reshape(len(lengths), 112) % 31 + 1).long() data = BatchedDataDict( input_ids=ids, input_lengths=lengths, token_mask=(torch.arange(112)[None, :] < lengths[:, None]).long(), - sample_mask=torch.tensor([1, 1, 1, 0, 1, 1]), + sample_mask=torch.tensor([1, 1, 1, 0, 1, 1, 1, 1, 1, 1]), ) mesh = NamedSharding( np.arange(world).reshape(1, world, 1, 1), @@ -85,7 +86,18 @@ def main() -> None: mesh = NamedSharding( np.arange(world).reshape(1, world // base_cp, base_cp, 1), mesh.names ) - dispatch = build_cp_dispatch(data, cfg, mesh, batch_size=6, training=True) + dispatch = build_cp_dispatch( + data, cfg, mesh, batch_size=data.size, training=True + ) + active_sizes = { + task.cp_size + for dp_plans in dispatch.plans + for rank_plan in dp_plans + for rank_step in rank_plan.steps + for task in rank_step.assignments + if task.sample_indices + } + assert active_sizes == {1, 2, 4}, active_sizes plan = dispatch.plans[rank // base_cp][rank % base_cp] payload = dispatch.data[rank // base_cp][rank % base_cp] step = plan.steps[0] diff --git a/tests/unit/distributed/test_dynamic_context_parallel.py b/tests/unit/distributed/test_dynamic_context_parallel.py index 949bfbdfdd3..8506df16a34 100644 --- a/tests/unit/distributed/test_dynamic_context_parallel.py +++ b/tests/unit/distributed/test_dynamic_context_parallel.py @@ -112,6 +112,60 @@ def test_padding_drives_cp_selection(): assert task.cp_size == 2 +def test_short_sequences_fill_existing_large_cp_packs(): + lengths = [6000, 2000] + phases = make_plan(lengths, lanes=2, budget=4096, sp=4) + tasks = [task for phase in phases for task in phase.assignments] + assert len(tasks) == 1 + assert tasks[0].sample_indices == (0, 1) + assert tasks[0].cp_size == 2 + assert tasks[0].padded_tokens == 8000 + check_plan(phases, lengths, 2, 4096, 4) + + +def test_cross_size_packing_uses_destination_padding_and_budget(): + # 4089 fits beside 4100 with CP1 padding (4100+4092), but not with + # CP2's required multiple of 16 (4112+4096). Do not overfill that pack. + lengths = [4100, 4089, 4000, 4000] + phases = make_plan(lengths, lanes=2, budget=4096, sp=4) + tasks = [task for phase in phases for task in phase.assignments] + long_task = next(task for task in tasks if 0 in task.sample_indices) + assert 1 not in long_task.sample_indices + assert len(long_task.sample_indices) == 2 + assert long_task.padded_tokens == 8112 + check_plan(phases, lengths, 2, 4096, 4) + + +def test_cross_size_packing_preserves_small_cp_for_remaining_work(): + lengths = [8192, 4000, 4000] + phases = make_plan(lengths, lanes=2, budget=4096, sp=4) + tasks = [task for phase in phases for task in phase.assignments] + assert sorted(task.cp_size for task in tasks) == [1, 1, 2] + check_plan(phases, lengths, 2, 4096, 4) + + +def test_cross_size_packing_respects_expert_minimum(): + lengths = [24000, 8000, 16000, 16000] + phases = make_plan(lengths, lanes=8, minimum=4, budget=4096, sp=2) + assert all(task.cp_size >= 4 for phase in phases for task in phase.assignments) + long_task = next( + task + for phase in phases + for task in phase.assignments + if 0 in task.sample_indices + ) + assert long_task.sample_indices == (0, 1) + check_plan(phases, lengths, 8, 4096, 2) + + +@pytest.mark.parametrize("tp", [1, 2]) +def test_gpu_parity_fixture_still_exercises_cp1_after_cross_size_packing(tp): + lengths = [7, 45, 101, 11, 55, 9, 29, 29, 29, 29] + phases = make_plan(lengths, lanes=4 // tp, budget=32 * tp, sp=tp) + sizes = {task.cp_size for phase in phases for task in phase.assignments} + assert sizes == ({1, 2, 4} if tp == 1 else {1, 2}) + + def test_moe_minimum_expands_real_work_to_fill_lanes(): phases = make_plan([3, 7], minimum=4) assert all(a.cp_size >= 4 for p in phases for a in p.assignments) diff --git a/tests/unit/distributed/test_dynamic_cp_dispatch.py b/tests/unit/distributed/test_dynamic_cp_dispatch.py index 12e66b2ff65..7a5ed828c8b 100644 --- a/tests/unit/distributed/test_dynamic_cp_dispatch.py +++ b/tests/unit/distributed/test_dynamic_cp_dispatch.py @@ -29,6 +29,8 @@ collect_cp_outputs, cp_schedule_matches, owned_real_task_count, + real_task_participation_count, + validate_dynamic_cp, ) @@ -133,12 +135,64 @@ def test_global_aux_loss_is_rejected_even_with_provider_coefficient(self): } ) ) + self.assertTrue( + _enabled_global_aux_loss( + { + "model_overrides": { + "moe_router_load_balancing_type": "global_aux_loss" + } + } + ) + ) self.assertFalse( _enabled_global_aux_loss( {"moe_router_load_balancing_type": "seq_aux_loss"} ) ) + def _validation_cfg(self) -> dict: + return { + "make_sequence_length_divisible_by": 1, + "sequence_packing": {"enabled": True, "pair_grouping_key": None}, + "dynamic_batching": {"enabled": False}, + "draft": {"enabled": False}, + "megatron_cfg": { + "enabled": True, + "pipeline_model_parallel_size": 1, + "tensor_model_parallel_size": 1, + "expert_model_parallel_size": 1, + "expert_tensor_parallel_size": 1, + "sequence_parallel": False, + "dynamic_context_parallel": { + "enabled": True, + "tokens_per_rank": 8, + "max_size": 4, + }, + }, + } + + def test_dynamic_cp_validation_requires_pp_one(self): + cfg = self._validation_cfg() + cfg["megatron_cfg"]["pipeline_model_parallel_size"] = 2 + with self.assertRaisesRegex(ValueError, "PP=1"): + validate_dynamic_cp(cfg, lanes=4) + + def test_dynamic_cp_validation_rejects_quantile_router(self): + cfg = self._validation_cfg() + cfg["megatron_cfg"]["model_overrides"] = { + "moe_router_load_balancing_type": "quantile_balancing" + } + with self.assertRaisesRegex(ValueError, "quantile_balancing"): + validate_dynamic_cp(cfg, lanes=4) + + def test_dynamic_cp_validation_rejects_moe_microbatch_overlap(self): + cfg = self._validation_cfg() + cfg["megatron_cfg"]["model_overrides"] = { + "overlap_moe_expert_parallel_comm": True + } + with self.assertRaisesRegex(ValueError, "overlap_moe_expert_parallel_comm"): + validate_dynamic_cp(cfg, lanes=4) + def test_outputs_from_nonzero_static_cp_are_preserved(self): dispatch = build_cp_dispatch( self.data, self.cfg, self.mesh, batch_size=None, training=False @@ -311,6 +365,20 @@ def test_dispatch_preserves_uneven_packed_task_lists(self): ), unique_real_tasks, ) + real_task_participations = sum( + task.cp_size + for group in dispatch.schedule.groups_by_batch[0] + for task in group.assignments + if task.sample_indices + ) + self.assertEqual( + sum( + real_task_participation_count(plan) + for dp_plans in dispatch.plans + for plan in dp_plans + ), + real_task_participations, + ) self.assertTrue( all( len(plan.steps[0].groups) == 1 diff --git a/tests/unit/distributed/test_tensor_serialization.py b/tests/unit/distributed/test_tensor_serialization.py index b2f3bccf925..f7f252d848d 100644 --- a/tests/unit/distributed/test_tensor_serialization.py +++ b/tests/unit/distributed/test_tensor_serialization.py @@ -38,3 +38,11 @@ def test_tensor_payload_detaches_autograd(): restored = tensor_from_payload(tensor_to_payload(tensor)) assert restored.item() == 2.0 assert not restored.requires_grad + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is unavailable") +def test_tensor_payload_stages_cuda_tensor_to_cpu(): + tensor = torch.arange(6, device="cuda", dtype=torch.float32).reshape(2, 3) + restored = tensor_from_payload(tensor_to_payload(tensor)) + assert restored.device.type == "cpu" + torch.testing.assert_close(restored, tensor.cpu()) diff --git a/tests/unit/models/megatron/test_dynamic_cp_moe.py b/tests/unit/models/megatron/test_dynamic_cp_moe.py index bde07e3ce31..ddc8580dd5d 100644 --- a/tests/unit/models/megatron/test_dynamic_cp_moe.py +++ b/tests/unit/models/megatron/test_dynamic_cp_moe.py @@ -73,6 +73,7 @@ def __init__(self, group): class _Router(torch.nn.Module): def __init__(self, group, config): super().__init__() + self.cp_group = group self.tp_cp_group = group self.config = config @@ -108,6 +109,7 @@ def test_dynamic_binding_updates_router_and_ssm_then_restores(monkeypatch): with dynamic_cp.preserve_attention_cp_groups(model): model_packed = dynamic_cp.bind_attention_cp_group(model, packed) assert model_packed.cp_group is active_cp + assert model.router.cp_group is active_cp assert model.router.tp_cp_group is active_tp_cp assert model.mamba.pg_collection.cp is active_cp assert model.mamba.cp is not original_mamba_helper @@ -117,6 +119,7 @@ def test_dynamic_binding_updates_router_and_ssm_then_restores(monkeypatch): assert config.moe_aux_loss_coeff == [0.0, 0.0] assert config.moe_z_loss_coeff is None + assert model.router.cp_group is original_tp_cp assert model.router.tp_cp_group is original_tp_cp assert model.mamba.pg_collection.cp is original_cp assert model.mamba.cp is original_mamba_helper @@ -124,3 +127,51 @@ def test_dynamic_binding_updates_router_and_ssm_then_restores(monkeypatch): assert model.gdn.cp_size == 1 assert config.moe_aux_loss_coeff == [0.1, 0.2] assert config.moe_z_loss_coeff == 0.01 + + +@pytest.mark.parametrize( + ("active_size", "group_tokens"), + [(1, 3), (4, 10)], +) +def test_dynamic_moe_scaling_uses_exact_active_group_token_count( + monkeypatch, active_size, group_tokens +): + from nemo_rl.models.megatron import dynamic_cp + + active_tp_cp = _Group(active_size) + config = SimpleNamespace(moe_aux_loss_coeff=0.1, moe_z_loss_coeff=0.02) + model = torch.nn.Module() + model.add_module("router", _Router(active_tp_cp, config)) + padding_mask = torch.tensor([[False, False, False, True]]) + + monkeypatch.setattr(dynamic_cp, "Router", _Router) + + def _all_reduce(value, *, group): + assert group is active_tp_cp + value.fill_(group_tokens) + + monkeypatch.setattr(dynamic_cp.torch.distributed, "all_reduce", _all_reduce) + + with dynamic_cp.preserve_attention_cp_groups(model): + dynamic_cp.configure_dynamic_moe_loss_scaling(model, padding_mask) + correction = group_tokens / (3 * active_size) + assert dynamic_cp.dynamic_moe_grad_scale_correction(config) == pytest.approx( + correction + ) + assert config.moe_z_loss_coeff == pytest.approx(0.02 / correction) + + assert dynamic_cp.dynamic_moe_grad_scale_correction(config) == 1.0 + assert config.moe_aux_loss_coeff == 0.1 + assert config.moe_z_loss_coeff == 0.02 + + +def test_dynamic_moe_scaling_is_noop_for_dense_model(monkeypatch): + from nemo_rl.models.megatron import dynamic_cp + + monkeypatch.setattr( + dynamic_cp.torch.distributed, + "all_reduce", + lambda *_args, **_kwargs: pytest.fail("dense model performed MoE reduction"), + ) + + dynamic_cp.configure_dynamic_moe_loss_scaling(torch.nn.Linear(2, 2), None) diff --git a/tests/unit/models/megatron/test_moe_metrics.py b/tests/unit/models/megatron/test_moe_metrics.py index 18ee6e2a351..6b7d2654cac 100644 --- a/tests/unit/models/megatron/test_moe_metrics.py +++ b/tests/unit/models/megatron/test_moe_metrics.py @@ -162,6 +162,53 @@ def _all_reduce(values, *, group): assert metrics["load_balancing_loss"] == pytest.approx(1.0) +@pytest.mark.mcore +def test_dynamic_cp_avg_group_metrics_use_rank_participation_scale(monkeypatch): + """z-loss must reproduce MCore's AVG over every participating rank.""" + from nemo_rl.models import megatron as megatron_module + from nemo_rl.models.megatron.common import get_moe_metrics + + aux_entry = SimpleNamespace(values=torch.tensor([1.0, 3.0]), avg_group=None) + # Padding-only lanes have a pre-initialized z entry whose avg_group was + # never populated by record(); the name must still select AVG semantics. + z_entry = SimpleNamespace(values=torch.tensor([2.0, 4.0]), avg_group=None) + live_tracker = SimpleNamespace( + metrics={"load_balancing_loss": aux_entry, "z_loss": z_entry} + ) + fixed_group = object() + reductions = [] + + def _all_reduce(values, *, group): + reductions.append(group) + values.mul_(2.0) + + monkeypatch.setattr( + megatron_module.common, "get_moe_metrics_tracker", lambda: live_tracker + ) + monkeypatch.setattr( + megatron_module.common, + "get_moe_layer_wise_logging_tracker", + lambda: { + "load_balancing_loss": {"values": aux_entry.values}, + "z_loss": {"values": z_entry.values}, + }, + ) + monkeypatch.setattr(megatron_module.common.dist, "all_reduce", _all_reduce) + monkeypatch.setattr( + megatron_module.common, "clear_aux_losses_tracker", lambda: None + ) + + metrics = get_moe_metrics( + loss_scale=0.25, + dynamic_parallel_group=fixed_group, + dynamic_avg_loss_scale=0.125, + ) + + assert reductions == [fixed_group, fixed_group] + assert metrics["load_balancing_loss"] == pytest.approx(1.0) + assert metrics["z_loss"] == pytest.approx(0.75) + + @pytest.mark.mcore @pytest.mark.parametrize( "routing_type,aux_loss_coeff,z_loss_coeff,expected", diff --git a/tests/unit/models/megatron/test_train.py b/tests/unit/models/megatron/test_train.py index d929b4abd74..0d2165f2ff0 100644 --- a/tests/unit/models/megatron/test_train.py +++ b/tests/unit/models/megatron/test_train.py @@ -226,6 +226,40 @@ def __init__(self): mock_scatter.assert_not_called() assert result is padding_mask + def test_first_gpt_stage_router_scaling_uses_sequence_parallel_mask(self): + """Count the same TP-local tokens that the first-stage router receives.""" + from nemo_rl.models.megatron import train + + class FakeGPTModel: + def __init__(self): + self.config = SimpleNamespace(sequence_parallel=True) + self.pre_process = True + + model = FakeGPTModel() + padding_mask = torch.tensor([[False, True, False, True]]) + scattered = torch.tensor([[False], [False]]) + tp_group = MagicMock() + + with ( + patch.object(train, "GPTModel", FakeGPTModel), + patch.object( + train.tensor_parallel, + "scatter_to_sequence_parallel_region", + return_value=scattered, + ) as mock_scatter, + patch.object( + train, "get_tensor_model_parallel_group", return_value=tp_group + ), + ): + result = train._prepare_padding_mask_for_router_scaling( + model, padding_mask + ) + + mock_scatter.assert_called_once() + assert torch.equal(mock_scatter.call_args.args[0], padding_mask.transpose(0, 1)) + assert mock_scatter.call_args.kwargs["group"] is tp_group + assert torch.equal(result, scattered.transpose(0, 1)) + def test_model_forward_with_defer_fp32_logits(self): """Test model_forward passes fp32_output when defer_fp32_logits is True.""" from nemo_rl.models.megatron.train import model_forward diff --git a/tests/unit/models/policy/test_megatron_worker.py b/tests/unit/models/policy/test_megatron_worker.py index aace0f4cdb3..5c60ce1a348 100644 --- a/tests/unit/models/policy/test_megatron_worker.py +++ b/tests/unit/models/policy/test_megatron_worker.py @@ -1766,6 +1766,26 @@ def test_compute_moe_grad_scale_normalizes_by_valid_tokens(): assert torch.allclose(scale_fn(), torch.tensor(0.25)) +def test_compute_moe_grad_scale_applies_dynamic_task_correction(monkeypatch): + from nemo_rl.models.policy.workers import megatron_policy_worker as worker_module + + worker = object.__new__(worker_module.MegatronPolicyWorkerImpl) + _disable_opd_full(worker) + model_config = SimpleNamespace() + worker.model = SimpleNamespace(config=model_config) + monkeypatch.setattr( + worker_module, + "dynamic_moe_grad_scale_correction", + lambda config: 1.5 if config is model_config else 1.0, + ) + + scale_fn = worker_module.MegatronPolicyWorkerImpl._compute_moe_grad_scale( + worker, torch.tensor(4.0) + ) + + assert torch.allclose(scale_fn(), torch.tensor(0.375)) + + def test_compute_moe_grad_scale_clamps_zero_valid_tokens(): """clamp(min=1) must guard against division by zero when no valid tokens.""" from nemo_rl.models.policy.workers.megatron_policy_worker import ( @@ -1781,6 +1801,71 @@ def test_compute_moe_grad_scale_clamps_zero_valid_tokens(): assert torch.allclose(scale_fn(), torch.tensor(1.0)) +def test_dynamic_cp_get_logprobs_runs_every_score_step(monkeypatch): + from nemo_rl.models.policy.workers import megatron_policy_worker as worker_module + + worker = object.__new__(worker_module.MegatronPolicyWorkerImpl) + _disable_opd_full(worker) + worker.timer = MagicMock() + worker.cfg = { + "logprob_batch_size": 1, + "megatron_cfg": {"use_fused_linear_logprobs": False}, + } + worker.model = MagicMock() + worker.sampling_params = None + worker.defer_fp32_logits = False + worker.mcore_state = SimpleNamespace(straggler_timer=None) + worker.delegate_pack_to_model = False + worker.delegate_mtp_loss_mask_to_model = False + worker.model_slices_context_parallel_inputs = False + worker.media_placeholder_token_id = None + worker._router_replay_enabled = False + data = BatchedDataDict(input_ids=torch.zeros(2, 4, dtype=torch.long)) + first_step = object() + second_step = object() + cp_plan = SimpleNamespace(steps=(first_step, second_step)) + iterator_calls = [] + forward_calls = [] + + def _get_iterator(*args, cp_step, **kwargs): + iterator_calls.append(cp_step) + return iter(()), 1, 1, 4, 4 + + def _forward_backward(**kwargs): + forward_calls.append(kwargs["data_iterator"]) + value = float(len(forward_calls)) + return [{"logprobs": torch.full((1, 4), value)}] + + monkeypatch.setattr( + worker_module, "attach_media_token_validity_mask", lambda *_: None + ) + monkeypatch.setattr(worker_module, "get_microbatch_iterator", _get_iterator) + monkeypatch.setattr(worker_module, "megatron_forward_backward", _forward_backward) + monkeypatch.setattr(worker_module, "_should_use_router_replay", lambda **_: False) + monkeypatch.setattr(worker_module, "LogprobsPostProcessor", MagicMock()) + monkeypatch.setattr( + worker_module.parallel_state, + "is_pipeline_last_stage", + lambda **_: True, + ) + monkeypatch.setattr( + worker_module, + "broadcast_tensors_from_last_stage", + lambda tensors: tensors, + ) + + result = worker_module.MegatronPolicyWorkerImpl.get_logprobs( + worker, data=data, cp_plan=cp_plan + ) + + assert iterator_calls == [first_step, second_step] + assert len(forward_calls) == 2 + torch.testing.assert_close( + result["logprobs"], + torch.tensor([[1.0, 1.0, 1.0, 1.0], [2.0, 2.0, 2.0, 2.0]]), + ) + + @pytest.mark.parametrize( ("kwargs", "expected_param_sync"), [({}, False), ({"param_sync": True}, True)], diff --git a/tests/unit/models/policy/test_policy_validation.py b/tests/unit/models/policy/test_policy_validation.py index d8fc112367e..edbabf8df94 100644 --- a/tests/unit/models/policy/test_policy_validation.py +++ b/tests/unit/models/policy/test_policy_validation.py @@ -20,10 +20,13 @@ when the cluster size is insufficient for the specified parallelism configuration. """ +from types import SimpleNamespace from unittest.mock import MagicMock, patch import pytest +import torch +from nemo_rl.distributed.batched_data_dict import BatchedDataDict from nemo_rl.models.policy import PolicyConfig from nemo_rl.models.policy.lm_policy import Policy @@ -61,6 +64,41 @@ def create_mock_tokenizer(): return tokenizer +def test_dynamic_cp_score_keeps_complete_training_sized_steps() -> None: + policy = Policy.__new__(Policy) + policy.cfg = {"train_global_batch_size": 2} + policy.sharding_annotations = object() + policy._dynamic_cp_schedule = None + policy.worker_group = MagicMock() + policy.worker_group.run_all_workers_sharded_data.return_value = "futures" + policy.worker_group.get_all_worker_results.return_value = ["worker-results"] + data = BatchedDataDict(input_ids=torch.zeros(4, 8, dtype=torch.long)) + schedule = object() + dispatch = SimpleNamespace( + schedule=schedule, + data="sharded-data", + plans="rank-plans", + output_rows=[], + ) + expected = BatchedDataDict(logprobs=torch.zeros(4, 8)) + + with ( + patch( + "nemo_rl.models.policy.lm_policy.build_cp_dispatch", + return_value=dispatch, + ) as build_dispatch, + patch( + "nemo_rl.models.policy.lm_policy.collect_cp_outputs", + return_value=expected, + ), + ): + result = Policy._get_dynamic_cp_outputs(policy, "get_logprobs", data) + + assert result is expected + assert policy._dynamic_cp_schedule is schedule + assert build_dispatch.call_args.kwargs["batch_size"] == 2 + + def create_dtensor_config( model_name: str, tp: int, pp: int = 1, cp: int = 1 ) -> PolicyConfig: diff --git a/tests/unit/test_dynamic_cp_comparison_recipes.py b/tests/unit/test_dynamic_cp_comparison_recipes.py index 72b7374c214..d7faa43e9f2 100644 --- a/tests/unit/test_dynamic_cp_comparison_recipes.py +++ b/tests/unit/test_dynamic_cp_comparison_recipes.py @@ -1,19 +1,17 @@ # Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. # Licensed under the Apache License, Version 2.0 (the "License"); -"""Validate that the Qwen30B MoE dynamic/static CP pair is matched.""" +"""Validate matched Qwen30B MoE and Qwen32B dense CP pairs.""" from pathlib import Path import pytest from omegaconf import OmegaConf -from nemo_rl.models.policy.dynamic_cp import _minimum_cp_size_for_experts from nemo_rl.utils.config import load_config, register_omegaconf_resolvers RECIPE_DIR = ( - Path(__file__).resolve().parents[2] - / "examples/configs/recipes/llm/performance" + Path(__file__).resolve().parents[2] / "examples/configs/recipes/llm/performance" ) @@ -37,6 +35,10 @@ def test_cp_comparison_recipe_pair( recipe_name, dynamic_enabled, context_parallel_size, pad_factor ): + # Policy imports require the optional worker/transformers dependencies; + # keep pure YAML comparisons runnable in the lightweight CPU environment. + from nemo_rl.models.policy.dynamic_cp import _minimum_cp_size_for_experts + register_omegaconf_resolvers() resolved = OmegaConf.to_container( load_config(RECIPE_DIR / recipe_name), resolve=True @@ -67,3 +69,47 @@ def test_cp_comparison_recipe_pair( assert resolved["logger"]["wandb_enabled"] is True assert resolved["logger"]["wandb"]["project"] == "nemo-rl-cp-comparison" + + +@pytest.mark.parametrize("model", ["qwen3-30ba3b", "qwen3-32b"]) +def test_comparison_changes_only_cp_and_logging(model): + register_omegaconf_resolvers() + stem = f"grpo-{model}-4n4g-async-1off-megatron-" + dynamic, static = [ + OmegaConf.to_container( + load_config(RECIPE_DIR / f"{stem}{mode}cp-10step.yaml"), resolve=True + ) + for mode in ("dynamic", "static") + ] + for config in (dynamic, static): + config.pop("logger") + config["policy"].pop("make_sequence_length_divisible_by") + config["policy"]["megatron_cfg"].pop("context_parallel_size") + config["policy"]["megatron_cfg"]["dynamic_context_parallel"].pop("enabled") + assert dynamic == static + + +def test_dense_comparison_capacity_and_async_resources(): + register_omegaconf_resolvers() + resolved = OmegaConf.to_container( + load_config( + RECIPE_DIR / "grpo-qwen3-32b-4n4g-async-1off-megatron-dynamiccp-10step.yaml" + ), + resolve=True, + ) + policy = resolved["policy"] + megatron = policy["megatron_cfg"] + assert policy["model_name"] == "Qwen/Qwen3-32B" + assert policy["train_global_batch_size"] == 512 + assert policy["max_total_sequence_length"] == 16384 + assert megatron["tensor_model_parallel_size"] == 2 + assert megatron["pipeline_model_parallel_size"] == 1 + assert megatron["expert_model_parallel_size"] == 1 + assert megatron["activation_checkpointing"] is True + assert megatron["dynamic_context_parallel"]["max_size"] == 4 + assert megatron["dynamic_context_parallel"]["tokens_per_rank"] == 4096 + assert resolved["cluster"]["num_nodes"] == 4 + assert resolved["grpo"]["async_grpo"]["enabled"] is True + assert policy["generation"]["colocated"]["resources"]["num_nodes"] == 2 + assert policy["generation"]["vllm_cfg"]["tensor_parallel_size"] == 2 + assert resolved["logger"]["wandb_enabled"] is True From 47cd040de60baf1e367da5cca22e5f8b5b53defa Mon Sep 17 00:00:00 2001 From: humairafirdowse18 Date: Mon, 21 Sep 2026 15:04:13 -0700 Subject: [PATCH 5/7] hybridep+mtp+aux loss+ETP fixes --- .../distributed/dynamic_context_parallel.py | 76 +- nemo_rl/models/megatron/common.py | 15 +- nemo_rl/models/megatron/data.py | 38 +- nemo_rl/models/megatron/dynamic_cp.py | 802 ++++++++++++++++-- nemo_rl/models/megatron/hybridep.py | 10 +- nemo_rl/models/policy/dynamic_cp.py | 106 ++- .../policy/workers/megatron_policy_worker.py | 53 +- .../distributed/test_dynamic_cp_dispatch.py | 114 ++- .../models/megatron/test_dynamic_cp_moe.py | 339 +++++++- .../models/megatron/test_hybridep_data.py | 75 +- .../models/megatron/test_megatron_data.py | 31 + .../unit/models/megatron/test_moe_metrics.py | 33 + .../unit/models/megatron/test_mtp_metrics.py | 48 +- 13 files changed, 1581 insertions(+), 159 deletions(-) diff --git a/nemo_rl/distributed/dynamic_context_parallel.py b/nemo_rl/distributed/dynamic_context_parallel.py index 2341693ce78..04257848177 100644 --- a/nemo_rl/distributed/dynamic_context_parallel.py +++ b/nemo_rl/distributed/dynamic_context_parallel.py @@ -267,6 +267,40 @@ def _merge_compatible_phases( return tuple(merged) +def _align_sync_group_tasks(group: CPSyncGroup) -> CPSyncGroup: + """Pad each CP subgroup to the same call count for domain-wide collectives.""" + slots: list[tuple[int, int]] = [] + tasks_by_slot: dict[tuple[int, int], list[CPAssignment]] = {} + for task in group.assignments: + slot = (task.lane_start, task.cp_size) + if slot not in tasks_by_slot: + slots.append(slot) + tasks_by_slot[slot] = [] + tasks_by_slot[slot].append(task) + + round_count = max(len(tasks) for tasks in tasks_by_slot.values()) + aligned: list[CPAssignment] = [] + for slot in slots: + tasks = tasks_by_slot[slot] + aligned.extend(tasks) + template = tasks[0] + for _ in range(round_count - len(tasks)): + aligned.append( + CPAssignment( + sample_indices=(), + lane_start=template.lane_start, + cp_size=template.cp_size, + pad_multiple=template.pad_multiple, + padded_tokens=( + (2 + template.pad_multiple - 1) + // template.pad_multiple + * template.pad_multiple + ), + ) + ) + return CPSyncGroup(tuple(aligned)) + + def plan_cp_phases( lengths: list[int], *, @@ -277,6 +311,8 @@ def plan_cp_phases( sequence_parallel_size: int, user_pad_multiple: int, token_alignment: int = 1, + atomic_groups: list[tuple[int, ...]] | None = None, + align_full_domain_collectives: bool = False, ) -> tuple[CPSyncGroup, ...]: """Pack sequences and form MCore-style synchronization groups. @@ -296,11 +332,23 @@ def plan_cp_phases( < 1 ): raise ValueError("Token budgets and padding factors must be positive") + if atomic_groups is None: + atomic_groups = [(index,) for index in range(len(lengths))] + flattened_indices = [index for group in atomic_groups for index in group] + if sorted(flattened_indices) != list(range(len(lengths))): + raise ValueError( + "Dynamic CP atomic groups must contain every sample exactly once" + ) + if any(not group for group in atomic_groups): + raise ValueError("Dynamic CP atomic groups cannot be empty") + if any(length < 2 for length in lengths): + raise ValueError("Dynamic CP requires at least two input tokens per sample") + bins: dict[int, list[tuple[list[int], int, int]]] = {} - for index in sorted(range(len(lengths)), key=lambda i: (-lengths[i], i)): - length = lengths[index] - if length < 2: - raise ValueError("Dynamic CP requires at least two input tokens per sample") + for group in sorted( + atomic_groups, + key=lambda members: (-sum(lengths[index] for index in members), members), + ): size = min_size while True: multiple = _padding_for_cp( @@ -309,13 +357,16 @@ def plan_cp_phases( user_pad_multiple=user_pad_multiple, token_alignment=token_alignment, ) - padded = (length + multiple - 1) // multiple * multiple + padded = sum( + (lengths[index] + multiple - 1) // multiple * multiple + for index in group + ) if padded <= tokens_per_rank * size: break size *= 2 if size > max_size: raise ValueError( - f"Sample {index} of length {length} exceeds the dynamic CP token budget" + f"Atomic sample group {group} exceeds the dynamic CP token budget" ) # Fill already-required larger-CP calls before opening another call. # Keeping separate size buckets strands space: e.g. [6000, 2000] at @@ -328,16 +379,18 @@ def plan_cp_phases( continue size_bins = bins[target_size] for bin_index, (members, used, factor) in enumerate(size_bins): - target_padded = (length + factor - 1) // factor * factor + target_padded = sum( + (lengths[index] + factor - 1) // factor * factor for index in group + ) if used + target_padded <= tokens_per_rank * target_size: - members.append(index) + members.extend(group) size_bins[bin_index] = (members, used + target_padded, factor) placed = True break if placed: break if not placed: - bins.setdefault(size, []).append(([index], padded, multiple)) + bins.setdefault(size, []).append((list(group), padded, multiple)) pending = [ CPAssignment(tuple(members), 0, size, factor, used) @@ -381,7 +434,10 @@ def plan_cp_phases( cursor += min_size phases.append(tuple(phase)) pending = remaining - return _merge_compatible_phases(phases, lengths) + merged = _merge_compatible_phases(phases, lengths) + if align_full_domain_collectives: + merged = tuple(_align_sync_group_tasks(group) for group in merged) + return merged def assignments_for_lane(group: CPSyncGroup, lane: int) -> tuple[CPAssignment, ...]: diff --git a/nemo_rl/models/megatron/common.py b/nemo_rl/models/megatron/common.py index e8ac29335f5..08d6dc45a8c 100644 --- a/nemo_rl/models/megatron/common.py +++ b/nemo_rl/models/megatron/common.py @@ -203,6 +203,7 @@ def get_moe_metrics( track_names: Optional[list[str]] = None, dynamic_parallel_group: Optional[dist.ProcessGroup] = None, dynamic_avg_loss_scale: Optional[float] = None, + dynamic_global_loss_scale: Optional[float] = None, ) -> dict[str, Any]: """Returns Mixture of Experts (MoE) auxiliary-loss metrics. @@ -234,6 +235,9 @@ def get_moe_metrics( The caller supplies the reciprocal of the global real-task rank participation count. Required with ``dynamic_parallel_group`` when an averaged metric is present. + dynamic_global_loss_scale: Scale for ``global_load_balancing_loss``. + Its full-domain collective produces one complete value per aligned + call round, rather than one value per independently packed task. Returns: dict[str, Any]: A flat dict of aggregated metrics. For each aux loss name, @@ -303,16 +307,19 @@ def get_moe_metrics( "Dynamic averaged MoE metrics require a rank-participation scale" ) resolved_dynamic_avg_loss_scale = ( - dynamic_avg_loss_scale - if dynamic_avg_loss_scale is not None - else loss_scale + dynamic_avg_loss_scale if dynamic_avg_loss_scale is not None else loss_scale ) aux_losses = { name: value["values"].float() * ( resolved_dynamic_avg_loss_scale if name in dynamic_avg_names - else loss_scale + else ( + dynamic_global_loss_scale + if name == "global_load_balancing_loss" + and dynamic_global_loss_scale is not None + else loss_scale + ) ) for name, value in tracker.items() } diff --git a/nemo_rl/models/megatron/data.py b/nemo_rl/models/megatron/data.py index 5083463f025..dcf515f5e33 100644 --- a/nemo_rl/models/megatron/data.py +++ b/nemo_rl/models/megatron/data.py @@ -281,9 +281,19 @@ def get_microbatch_iterator( "Dynamic CP iterator requires an explicit optimizer-step plan" ) width = data["input_ids"].shape[1] + prepad_packed_seq_for_hybridep = bool( + uses_hybridep_flex_dispatcher(cfg["megatron_cfg"]) + and cfg["megatron_cfg"].get("moe_hybridep_prepad_packed_inputs") + ) return ( RerunDataIterator( - planned_microbatches(data, cp_plan, cp_step, straggler_timer) + planned_microbatches( + data, + cp_plan, + cp_step, + straggler_timer, + prepad_packed_seq_for_hybridep=prepad_packed_seq_for_hybridep, + ) ), len(cp_step.assignments), 1, @@ -703,6 +713,19 @@ def process_microbatch( if model_slices_context_parallel_inputs else local_mtp_loss_mask ) + if prepad_packed_seq_for_hybridep: + target_length = input_ids_cp_sharded.shape[1] + current_length = mtp_loss_mask.shape[1] + if current_length > target_length: + raise ValueError( + "HybridEP-prepadded MTP mask exceeds the model input" + ) + if current_length < target_length: + mtp_loss_mask = torch.nn.functional.pad( + mtp_loss_mask, + (0, target_length - current_length), + value=0, + ) # Pack the media-token validity mask the same way as input_ids. # The mask answers a per-token question, so it only means @@ -740,6 +763,19 @@ def process_microbatch( if model_slices_context_parallel_inputs else local_media_mask ).bool() + if prepad_packed_seq_for_hybridep: + target_length = input_ids_cp_sharded.shape[1] + current_length = media_token_validity_mask.shape[1] + if current_length > target_length: + raise ValueError( + "HybridEP-prepadded media mask exceeds the model input" + ) + if current_length < target_length: + media_token_validity_mask = torch.nn.functional.pad( + media_token_validity_mask, + (0, target_length - current_length), + value=False, + ) # For packed sequences, position_ids and attention_mask are typically None # The PackedSeqParams handles all necessary sequence information diff --git a/nemo_rl/models/megatron/dynamic_cp.py b/nemo_rl/models/megatron/dynamic_cp.py index 57572f48565..b0d2f035c53 100644 --- a/nemo_rl/models/megatron/dynamic_cp.py +++ b/nemo_rl/models/megatron/dynamic_cp.py @@ -13,11 +13,10 @@ from contextlib import contextmanager from copy import copy from dataclasses import dataclass -from typing import Any, Iterator +from typing import Any, Callable, Iterator, Optional import torch from megatron.core import parallel_state -from megatron.core.transformer.attention import Attention from megatron.core.transformer.moe.router import Router from nemo_rl.distributed.dynamic_context_parallel import CPRankPlan, CPRankStep @@ -26,6 +25,7 @@ _DYNAMIC_TP_CP_GROUPS: dict[int, Any] = {} _ROUTER_CONFIG_BASELINES: dict[int, tuple[Any, Any]] = {} _DYNAMIC_MOE_GRAD_SCALE_CORRECTIONS: dict[int, float] = {} +_DYNAMIC_MTP_METRICS: dict[str, torch.Tensor] = {} @dataclass(frozen=True) @@ -114,6 +114,41 @@ def _is_gated_delta_net(module: torch.nn.Module) -> bool: ) +def _is_gated_delta_product(module: torch.nn.Module) -> bool: + """Return whether ``module`` owns a headwise GDP CP helper.""" + cp = getattr(module, "cp", None) + return cp is not None and all( + hasattr(cp, name) + for name in ( + "d_inner_local_tp", + "nheads_local_tp", + "ngroups_local_tp", + "num_householder", + "headdim", + "conv1d_cp1", + "dt_bias_cp1", + "A_log_cp1", + "D_cp1", + ) + ) + + +def _uses_direct_cp_group(module: torch.nn.Module) -> bool: + """Return whether pinned MCore caches CP outside ``pg_collection``.""" + module_name = type(module).__module__ + return module_name.startswith( + "megatron.core.transformer.multi_token_prediction" + ) or module_name.startswith("megatron.core.models.hybrid.hybrid_block") + + +def _is_hybrid_stack(module: torch.nn.Module) -> bool: + return ( + type(module).__module__.startswith("megatron.core.models.hybrid.hybrid_block") + and hasattr(module, "layer_config_list") + and hasattr(module, "_cp_layout_manager") + ) + + def _rebuild_mamba_cp(module: torch.nn.Module, group: Any) -> None: """Rebuild Mamba's cached CP helper for the active microbatch size.""" cp = module.cp @@ -133,6 +168,577 @@ def _rebuild_mamba_cp(module: torch.nn.Module, group: Any) -> None: ) +def _rebuild_gdp_cp(module: torch.nn.Module, group: Any) -> None: + """Rebuild GDP's cached headwise CP helper for one runtime task.""" + cp = module.cp + module.cp = type(cp)( + cp_group=group, + d_inner_local_tp=cp.d_inner_local_tp, + nheads_local_tp=cp.nheads_local_tp, + ngroups_local_tp=cp.ngroups_local_tp, + d_state=cp.d_state, + num_householder=cp.num_householder, + headdim=cp.headdim, + conv1d_cp1=cp.conv1d_cp1, + dt_bias_cp1=cp.dt_bias_cp1, + A_log_cp1=cp.A_log_cp1, + D_cp1=cp.D_cp1, + D_has_hdim=cp.D_has_hdim, + sequence_is_contiguous=cp.sequence_is_contiguous, + ) + module.d_inner_local_cp = module.cp.d_inner_local_tpcp + module.nheads_local_cp = module.cp.nheads_local_tpcp + module.ngroups_local_cp = module.cp.ngroups_local_tpcp + + +def _bind_hybrid_stack_layout( + module: torch.nn.Module, *, group: Any, tp_cp_group: Any +) -> None: + """Build the CP layout converter omitted by a static CP=1 construction.""" + from megatron.core.context_parallel import ContextParallelLayoutManager + from megatron.core.models.hybrid.layers import utils as layer_utils + + module.cp_group = group + module.tp_cp_group = tp_cp_group + module._has_linear_layer_with_chunkwise_cp = any( + type(layer_config) is layer_utils.MambaLayerConfig + and layer_config.linear_cp_mode == "chunkwise" + for layer_config in module.layer_config_list + ) + if group.size() == 1: + module._cp_layout_manager = None + return + + layer_layouts = tuple( + ( + layer_config.attention_cp_layout + if type(layer_config) in layer_utils.Symbols.ATTENTION_LAYER_CONFIGS + else layer_config.linear_cp_layout + ) + for layer_config in module.layer_config_list + ) + boundary_layout = ( + module.config.attention_cp_layout + if getattr(module, "is_mtp_layer", False) + else module.config.linear_cp_layout + ) + module._cp_layout_manager = ContextParallelLayoutManager( + layer_layouts=layer_layouts, + boundary_layout=boundary_layout, + sequence_parallel=module.config.sequence_parallel, + cp_group=group, + tp_group=module.tp_group, + tp_cp_group=tp_cp_group, + ) + + +def validate_dynamic_cp_model(model: torch.nn.Module) -> None: + """Reject MCore model internals that cannot be safely rebound at runtime.""" + modules = list(model.modules()) + for module in modules: + module_name = type(module).__module__ + class_name = type(module).__name__ + if ( + module_name.startswith("megatron.core.models.hybrid.hybrid_model") + and getattr(module, "mtp", None) is not None + and getattr(module.config, "linear_cp_layout", None) + != getattr(module.config, "attention_cp_layout", None) + ): + raise ValueError( + "Dynamic CP Hybrid MTP requires matching linear_cp_layout and " + "attention_cp_layout because NeMo-RL supplies one packed layout" + ) + if class_name in {"MLASelfAttention", "AbsorbedMLASelfAttention"} or ( + "multi_latent_attention" in module_name + or "experimental_attention_variant.absorbed_mla" in module_name + ): + raise ValueError( + "Dynamic CP with MLA requires the unmerged MCore runtime-group " + "support; this pinned MCore still rejects dynamic packed metadata" + ) + if ( + hasattr(module, "config") + and getattr(module.config, "linear_cp_mode", None) == "chunkwise" + and ( + _is_mamba_mixer(module) + or _is_gated_delta_product(module) + or ".ssm." in module_name + ) + ): + raise ValueError( + "Dynamic CP does not support chunkwise linear CP with the pinned " + "MCore; use headwise linear_cp_mode" + ) + + +def _save_dynamic_mtp_metrics( + *, + loss_sum: torch.Tensor, + num_tokens: torch.Tensor, + correct: torch.Tensor, + total: torch.Tensor, + layer_number: int, + num_layers: int, +) -> None: + """Accumulate MTP numerators/counts without a stale fixed CP average.""" + values = { + "loss_sums": loss_sum, + "loss_token_counts": num_tokens, + "correct_values": correct, + "total_values": total, + } + for name, value in values.items(): + if name not in _DYNAMIC_MTP_METRICS: + _DYNAMIC_MTP_METRICS[name] = torch.zeros( + num_layers, dtype=torch.float32, device=value.device + ) + _DYNAMIC_MTP_METRICS[name][layer_number] += value.detach().float() + + +def get_dynamic_mtp_metrics( + *, parallel_group: torch.distributed.ProcessGroup +) -> dict[str, float]: + """Reduce token-weighted MTP metrics over the fixed DP*CP lane group.""" + if "loss_sums" not in _DYNAMIC_MTP_METRICS: + return {} + try: + for value in _DYNAMIC_MTP_METRICS.values(): + torch.distributed.all_reduce( + value, op=torch.distributed.ReduceOp.SUM, group=parallel_group + ) + losses = _DYNAMIC_MTP_METRICS["loss_sums"] / _DYNAMIC_MTP_METRICS[ + "loss_token_counts" + ].clamp(min=1) + acceptance = ( + _DYNAMIC_MTP_METRICS["correct_values"] + / _DYNAMIC_MTP_METRICS["total_values"].clamp(min=1) + * 100.0 + ) + metrics: dict[str, float] = {} + for index in range(losses.numel()): + metrics[f"mtp_{index + 1}_loss"] = float(losses[index].item()) + metrics[f"mtp_{index + 1}_acceptance_rate"] = float( + acceptance[index].item() + ) + return metrics + finally: + _DYNAMIC_MTP_METRICS.clear() + + +def _runtime_mtp_token_counts( + original_num_tokens: torch.Tensor, + mtp_num_tokens: torch.Tensor, + cp_group: Optional[torch.distributed.ProcessGroup], +) -> tuple[torch.Tensor, torch.Tensor]: + """Return task-wide main/MTP counts, including empty CP shards.""" + counts = torch.stack((original_num_tokens, mtp_num_tokens)) + if cp_group is not None and cp_group.size() > 1: + torch.distributed.all_reduce( + counts, op=torch.distributed.ReduceOp.SUM, group=cp_group + ) + return counts.unbind() + + +def _dynamic_process_mtp_loss( + hidden_states: torch.Tensor, + labels: Optional[torch.Tensor], + loss_mask: Optional[torch.Tensor], + output_layer: Callable[..., Any], + output_weight: Optional[torch.Tensor], + runtime_gather_output: Optional[bool], + is_training: bool, + compute_language_model_loss: Callable[..., torch.Tensor], + config: Any, + cp_group: Optional[torch.distributed.ProcessGroup] = None, + tp_group: Optional[torch.distributed.ProcessGroup] = None, + packed_seq_params: Optional[Any] = None, + scale_logits_fn: Optional[Callable[[torch.Tensor], torch.Tensor]] = None, + input_ids: Optional[torch.Tensor] = None, + mtp_input_mask: Optional[torch.Tensor] = None, + metric_avg_group: Optional[torch.distributed.ProcessGroup] = None, + main_hidden_states: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Pinned MCore MTP loss with runtime-CP normalization and raw metrics.""" + del metric_avg_group # Dynamic metrics reduce once over the fixed lane group. + from megatron.core.transformer.multi_token_prediction import ( + MTPLossAutoScaler, + _compute_mtp_acceptance_counts, + roll_tensor, + ) + + hidden_states_list = torch.chunk(hidden_states, 1 + config.mtp_num_layers, dim=0) + hidden_states = ( + hidden_states_list[0] if main_hidden_states is None else main_hidden_states + ) + + derived_labels_from_input_ids = False + if labels is None: + if input_ids is None: + return hidden_states + labels, _ = roll_tensor( + input_ids, + shifts=-1, + dims=-1, + cp_group=cp_group, + packed_seq_params=packed_seq_params, + return_sum=False, + ) + derived_labels_from_input_ids = True + + if config.mtp_detach_heads: + output_weight = ( + output_layer.weight.detach() + if output_weight is None + else output_weight.detach() + ) + + mtp_labels = labels.clone() + if loss_mask is None: + loss_mask = torch.ones_like(mtp_labels) + if derived_labels_from_input_ids: + loss_mask, _ = roll_tensor( + loss_mask, + shifts=-1, + dims=-1, + cp_group=cp_group, + packed_seq_params=packed_seq_params, + return_sum=False, + ) + + original_num_tokens = loss_mask.sum() + cumulative_mtp_input_mask = None + rolled_num_tokens = original_num_tokens + if mtp_input_mask is not None: + if mtp_input_mask.shape != loss_mask.shape: + raise ValueError( + f"mtp_input_mask shape {mtp_input_mask.shape} must match " + f"loss_mask shape {loss_mask.shape}" + ) + mtp_input_mask = mtp_input_mask.to(dtype=torch.bool) + + for mtp_layer_number in range(config.mtp_num_layers): + mtp_logits, _ = output_layer( + hidden_states_list[mtp_layer_number + 1], + weight=output_weight, + runtime_gather_output=runtime_gather_output, + ) + if scale_logits_fn is not None: + mtp_logits = scale_logits_fn(mtp_logits) + mtp_labels, _ = roll_tensor( + mtp_labels, + shifts=-1, + dims=-1, + cp_group=cp_group, + packed_seq_params=packed_seq_params, + return_sum=False, + ) + + if mtp_input_mask is not None: + mask_metadata = torch.cat( + (loss_mask, mtp_input_mask.to(dtype=loss_mask.dtype)), dim=0 + ) + mask_metadata, _ = roll_tensor( + mask_metadata, + shifts=-1, + dims=-1, + cp_group=cp_group, + packed_seq_params=packed_seq_params, + return_sum=False, + ) + loss_mask, mtp_input_mask = mask_metadata.chunk(2, dim=0) + mtp_input_mask = mtp_input_mask.to(dtype=torch.bool) + cumulative_mtp_input_mask = ( + mtp_input_mask + if cumulative_mtp_input_mask is None + else cumulative_mtp_input_mask & mtp_input_mask + ) + layer_loss_mask = loss_mask * cumulative_mtp_input_mask + num_tokens = layer_loss_mask.sum() + else: + loss_mask, rolled_num_tokens = roll_tensor( + loss_mask, + shifts=-1, + dims=-1, + cp_group=cp_group, + packed_seq_params=packed_seq_params, + ) + layer_loss_mask = loss_mask + num_tokens = rolled_num_tokens + + mtp_loss = layer_loss_mask * compute_language_model_loss(mtp_labels, mtp_logits) + if is_training: + correct, total = _compute_mtp_acceptance_counts( + mtp_logits, + mtp_labels, + layer_loss_mask, + output_layer, + runtime_gather_output, + tp_group, + ) + _save_dynamic_mtp_metrics( + loss_sum=mtp_loss.sum(), + num_tokens=num_tokens, + correct=correct, + total=total, + layer_number=mtp_layer_number, + num_layers=config.mtp_num_layers, + ) + + mtp_loss_scale = config.mtp_loss_scaling_factor / config.mtp_num_layers + if config.calculate_per_token_loss: + main_num_tokens, task_mtp_num_tokens = _runtime_mtp_token_counts( + original_num_tokens, num_tokens, cp_group + ) + mtp_loss = ( + mtp_loss_scale + * mtp_loss + * (main_num_tokens / task_mtp_num_tokens.clamp(min=1)) + ) + else: + mtp_loss = mtp_loss_scale * mtp_loss / num_tokens.clamp(min=1) + hidden_states = MTPLossAutoScaler.apply(hidden_states, mtp_loss) + + return hidden_states + + +@contextmanager +def _patch_mtp_loss_for_dynamic_cp(enabled: bool) -> Iterator[None]: + """Temporarily route GPT/HybridModel MTP through the NeMo-side fix.""" + if not enabled: + yield + return + + from megatron.core.models.gpt import gpt_model + + targets = [gpt_model] + try: + from megatron.core.models.hybrid import hybrid_model + + targets.append(hybrid_model) + except ImportError: + pass + originals = [(target, target.process_mtp_loss) for target in targets] + for target, _ in originals: + target.process_mtp_loss = _dynamic_process_mtp_loss + try: + yield + finally: + for target, original in originals: + target.process_mtp_loss = original + + +def _dynamic_attach_and_log_load_balancing_loss( + self: Router, + activation: torch.Tensor, + aux_loss_coeff: float, + aux_loss: torch.Tensor, + aux_loss_name: str, + reduce_group: torch.distributed.ProcessGroup, + needs_dp_avg: bool = True, + valid_token_count: int | torch.Tensor | None = None, +) -> torch.Tensor: + """Pinned router attachment with exact runtime-group token scaling. + + MCore's pinned implementation multiplies by ``local_tokens * + tp_cp_group.size()``. That is only equal to the task token count for + equally populated fixed shards. The router pre-hook has already reduced + the current mask over the runtime TP*CP group, so attach that exact count. + Logging intentionally retains the unscaled aux value. + """ + from megatron.core.transformer.moe.moe_logging import get_moe_metrics_tracker + from megatron.core.transformer.moe.moe_utils import MoEAuxLossAutoScaler + + if ( + self.is_mtp_layer + and self.config.mtp_use_repeated_layer + and self.config.mtp_num_layers is not None + ): + aux_loss = aux_loss / self.config.mtp_num_layers + + num_layers = self.config.num_layers + if self.config.mtp_num_layers is not None: + num_layers += self.config.mtp_num_layers + layer_number = ( + self.layer_number + self.config.num_layers + if self.is_mtp_layer + else self.layer_number + ) + get_moe_metrics_tracker().record( + aux_loss_name, + aux_loss / aux_loss_coeff, + layer_number, + num_layers, + reduce_group=reduce_group, + needs_dp_avg=needs_dp_avg, + ) + + if self.calculate_per_token_loss: + task_tokens = getattr(self, "_nemo_dynamic_aux_scale_tokens", None) + if task_tokens is None: + local_tokens = ( + valid_token_count + if valid_token_count is not None + else activation.shape[0] + ) + task_tokens = ( + torch.as_tensor(local_tokens, device=activation.device).detach().clone() + ) + torch.distributed.all_reduce(task_tokens, group=self.tp_cp_group) + return MoEAuxLossAutoScaler.apply(activation, aux_loss * task_tokens) + return MoEAuxLossAutoScaler.apply(activation, aux_loss) + + +@contextmanager +def _patch_hybrid_mtp_padding_masks(model: torch.nn.Module) -> Iterator[None]: + """Carry correct router padding semantics through every dynamic MTP block. + + The pinned HybridModel accepts ``padding_mask`` and sends it through the + backbone, while its nested MTP call accidentally omits the same keyword. + A HybridModel pre-hook captures that missing value and an MTP pre-hook + supplies it without changing static execution or the MCore submodule. + + MCore's MTP roll helper fills newly exposed positions with false. That is + correct for a validity mask, but ``padding_mask`` uses the opposite meaning + (true means padding). Transport validity through MTP and invert it only at + descendant MoE routers. This also guarantees that a padding-only Dynamic + CP call stays empty at every MTP depth. + + Dynamic CP supports headwise linear CP only, where the hybrid backbone and + MTP boundary both use the attention (zigzag) token layout. Consequently + the mask received by HybridModel already has the ordering needed by MTP. + """ + padding_masks: dict[int, torch.Tensor | None] = {} + handles: list[Any] = [] + modules = list(model.modules()) + mtp_blocks: list[torch.nn.Module] = [] + + def capture_padding_mask( + module: torch.nn.Module, args: tuple[Any, ...], kwargs: dict[str, Any] + ) -> None: + del args + padding_masks[id(module.mtp)] = kwargs.get("padding_mask") + + def prepare_mtp_validity_mask( + module: torch.nn.Module, args: tuple[Any, ...], kwargs: dict[str, Any] + ) -> tuple[tuple[Any, ...], dict[str, Any]]: + padding_mask = kwargs.get("padding_mask") + if padding_mask is None and len(args) > 4: + padding_mask = args[4] + if padding_mask is None: + padding_mask = padding_masks.get(id(module)) + if padding_mask is not None: + validity_mask = ~padding_mask.to(dtype=torch.bool) + if len(args) > 4: + args = (*args[:4], validity_mask, *args[5:]) + else: + kwargs["padding_mask"] = validity_mask + return args, kwargs + + for module in modules: + if not type(module).__module__.startswith( + "megatron.core.models.hybrid.hybrid_model" + ): + continue + mtp = getattr(module, "mtp", None) + if mtp is None: + continue + handles.append( + module.register_forward_pre_hook(capture_padding_mask, with_kwargs=True) + ) + mtp_blocks.append(mtp) + + for module in modules: + if ( + type(module).__module__.startswith( + "megatron.core.transformer.multi_token_prediction" + ) + and type(module).__name__ == "MultiTokenPredictionBlock" + and module not in mtp_blocks + ): + mtp_blocks.append(module) + + mtp_router_ids = { + id(nested) + for mtp in mtp_blocks + for nested in mtp.modules() + if isinstance(nested, Router) + } + + def prepare_router_padding_mask( + module: Router, args: tuple[Any, ...], kwargs: dict[str, Any] + ) -> tuple[tuple[Any, ...], dict[str, Any]]: + padding_mask = kwargs.get("padding_mask") + if padding_mask is None and len(args) > 1: + padding_mask = args[1] + if padding_mask is not None and id(module) in mtp_router_ids: + padding_mask = ~padding_mask.to(dtype=torch.bool) + if len(args) > 1: + args = (args[0], padding_mask, *args[2:]) + else: + kwargs["padding_mask"] = padding_mask + + if ( + module.training + and torch.is_grad_enabled() + and getattr(module.config, "calculate_per_token_loss", False) + and _has_positive_coefficient(module.config.moe_aux_loss_coeff) + ): + if padding_mask is None: + raise ValueError( + "Dynamic CP MoE aux loss requires a packed padding mask" + ) + group_tokens = (~padding_mask).sum().detach().clone() + torch.distributed.all_reduce(group_tokens, group=module.tp_cp_group) + module._nemo_dynamic_aux_scale_tokens = group_tokens + else: + module._nemo_dynamic_aux_scale_tokens = None + return args, kwargs + + patched_classes: list[tuple[type[Any], bool, Any]] = [] + for router_class in { + type(module) for module in modules if isinstance(module, Router) + }: + if not hasattr(router_class, "attach_and_log_load_balancing_loss"): + continue + had_direct_method = ( + "attach_and_log_load_balancing_loss" in router_class.__dict__ + ) + original_method = router_class.__dict__.get( + "attach_and_log_load_balancing_loss" + ) + patched_classes.append((router_class, had_direct_method, original_method)) + router_class.attach_and_log_load_balancing_loss = ( + _dynamic_attach_and_log_load_balancing_loss + ) + + for mtp in mtp_blocks: + handles.append( + mtp.register_forward_pre_hook(prepare_mtp_validity_mask, with_kwargs=True) + ) + for module in modules: + if isinstance(module, Router): + handles.append( + module.register_forward_pre_hook( + prepare_router_padding_mask, with_kwargs=True + ) + ) + try: + yield + finally: + for handle in handles: + handle.remove() + for module in modules: + if isinstance(module, Router) and hasattr( + module, "_nemo_dynamic_aux_scale_tokens" + ): + del module._nemo_dynamic_aux_scale_tokens + for router_class, had_direct_method, original_method in patched_classes: + if had_direct_method: + router_class.attach_and_log_load_balancing_loss = original_method + else: + delattr(router_class, "attach_and_log_load_balancing_loss") + + @contextmanager def preserve_attention_cp_groups(model: torch.nn.Module) -> Iterator[None]: """Isolate runtime attention, SSM and router state from fixed groups. @@ -140,24 +746,52 @@ def preserve_attention_cp_groups(model: torch.nn.Module) -> Iterator[None]: Keep the active group through backward recomputation; restore it after the complete no-pipeline schedule. TP, DP and optimizer groups are unchanged. """ + modules = list(model.modules()) saved_collections = [ (module, module.pg_collection) - for module in model.modules() - if isinstance(module, Attention) - or _is_mamba_mixer(module) - or _is_gated_delta_net(module) + for module in modules + if getattr(module, "pg_collection", None) is not None + and hasattr(module.pg_collection, "cp") ] - saved_mamba = [ - (module, module.cp) for module in model.modules() if _is_mamba_mixer(module) + saved_mamba = [(module, module.cp) for module in modules if _is_mamba_mixer(module)] + saved_gdp = [ + ( + module, + module.cp, + module.d_inner_local_cp, + module.nheads_local_cp, + module.ngroups_local_cp, + ) + for module in modules + if _is_gated_delta_product(module) ] saved_gdn = [ - (module, module.cp_size) - for module in model.modules() + (module, module.cp_size, getattr(module, "feat_dim_split", None)) + for module in modules if _is_gated_delta_net(module) ] + saved_direct_groups = [ + ( + module, + module.cp_group, + getattr(module, "tp_cp_group", None), + hasattr(module, "tp_cp_group"), + ) + for module in modules + if _uses_direct_cp_group(module) and hasattr(module, "cp_group") + ] + saved_hybrid_stacks = [ + ( + module, + module._cp_layout_manager, + module._has_linear_layer_with_chunkwise_cp, + ) + for module in modules + if _is_hybrid_stack(module) + ] saved_routers = [ (module, module.cp_group, module.tp_cp_group) - for module in model.modules() + for module in modules if isinstance(module, Router) ] router_configs: dict[int, Any] = {} @@ -171,15 +805,47 @@ def preserve_attention_cp_groups(model: torch.nn.Module) -> Iterator[None]: config.moe_z_loss_coeff, ) _DYNAMIC_MOE_GRAD_SCALE_CORRECTIONS[id(config)] = 1.0 + mtp_enabled = any( + bool(getattr(getattr(module, "config", None), "mtp_num_layers", 0)) + for module in modules + ) try: - yield + with ( + _patch_mtp_loss_for_dynamic_cp(mtp_enabled), + _patch_hybrid_mtp_padding_masks(model), + ): + try: + yield + except Exception: + _DYNAMIC_MTP_METRICS.clear() + raise finally: for module, collection in saved_collections: module.pg_collection = collection for module, cp in saved_mamba: module.cp = cp - for module, cp_size in saved_gdn: + for module, cp_size, feat_dim_split in saved_gdn: module.cp_size = cp_size + if feat_dim_split is not None: + module.feat_dim_split = feat_dim_split + for ( + module, + cp, + d_inner_local_cp, + nheads_local_cp, + ngroups_local_cp, + ) in saved_gdp: + module.cp = cp + module.d_inner_local_cp = d_inner_local_cp + module.nheads_local_cp = nheads_local_cp + module.ngroups_local_cp = ngroups_local_cp + for module, cp_group, tp_cp_group, had_tp_cp_group in saved_direct_groups: + module.cp_group = cp_group + if had_tp_cp_group: + module.tp_cp_group = tp_cp_group + for module, manager, has_chunkwise in saved_hybrid_stacks: + module._cp_layout_manager = manager + module._has_linear_layer_with_chunkwise_cp = has_chunkwise for module, cp_group, tp_cp_group in saved_routers: module.cp_group = cp_group module.tp_cp_group = tp_cp_group @@ -196,11 +862,16 @@ def _bind_router_config(router: Router, *, padding_only: bool) -> None: raise RuntimeError("Dynamic CP router binding escaped its preservation context") aux_coeff, z_coeff = baseline if padding_only: - if isinstance(aux_coeff, tuple): - aux_coeff = tuple(0.0 for _ in aux_coeff) - elif isinstance(aux_coeff, list): - aux_coeff = [0.0 for _ in aux_coeff] - else: + routing_type = getattr(router.config, "moe_router_load_balancing_type", None) + if isinstance(routing_type, (list, tuple)) and isinstance( + aux_coeff, (list, tuple) + ): + filtered = [ + coeff if kind == "global_aux_loss" else 0.0 + for kind, coeff in zip(routing_type, aux_coeff) + ] + aux_coeff = tuple(filtered) if isinstance(aux_coeff, tuple) else filtered + elif routing_type != "global_aux_loss": aux_coeff = 0.0 z_coeff = None router.config.moe_aux_loss_coeff = aux_coeff @@ -217,19 +888,13 @@ def _has_positive_coefficient(value: Any) -> bool: def configure_dynamic_moe_loss_scaling( model: torch.nn.Module, padding_mask: torch.Tensor | None ) -> None: - """Correct MCore's fixed-shard aux-loss scaling for the active task. - - MCore multiplies each rank's load-balancing loss by - ``local_valid_tokens * tp_cp_group.size()``. That equals the task-wide token - count only when every TP*CP shard contains the same number of valid tokens. - Packed dynamic tasks do not have that invariant. Compute the exact active - group token count and expose a per-rank correction through MCore's existing - ``moe_grad_scale_func`` hook. - - The same autograd scaler carries z-loss. MCore's z-loss coefficient and - attachment factors already cancel to produce each rank's valid-token sum, - so inversely adjust the temporary coefficient to keep that gradient - unchanged when the shared scaler applies the aux correction. + """Validate inputs for the temporary exact-token router attachment. + + ``preserve_attention_cp_groups`` installs a router pre-hook which reduces + the current mask and an attachment shim which uses that exact TP*CP token + count. Keeping the check here fails before entering a router collective if + a caller forgot the packed padding mask. The worker's ordinary + ``1/global_valid_tokens`` MoE scale is therefore sufficient. """ routers = [module for module in model.modules() if isinstance(module, Router)] if not routers: @@ -249,25 +914,8 @@ def configure_dynamic_moe_loss_scaling( if padding_mask is None: raise ValueError("Dynamic CP MoE aux loss requires a packed padding mask") - tp_cp_group = routers[0].tp_cp_group - if any(router.tp_cp_group is not tp_cp_group for router in routers[1:]): + if any(router.tp_cp_group is not routers[0].tp_cp_group for router in routers[1:]): raise ValueError("Dynamic CP routers disagree on the active TP*CP group") - active_size = tp_cp_group.size() - local_valid_tokens = (~padding_mask).sum() - group_valid_tokens = local_valid_tokens.clone() - torch.distributed.all_reduce(group_valid_tokens, group=tp_cp_group) - local_count = int(local_valid_tokens.item()) - correction = ( - float(group_valid_tokens.item()) / (local_count * active_size) - if local_count > 0 - else 1.0 - ) - - for config_id, config in configs.items(): - _DYNAMIC_MOE_GRAD_SCALE_CORRECTIONS[config_id] = correction - _, baseline_z_coeff = _ROUTER_CONFIG_BASELINES[config_id] - if isinstance(baseline_z_coeff, (int, float)): - config.moe_z_loss_coeff = baseline_z_coeff / correction def dynamic_moe_grad_scale_correction(model_config: Any) -> float: @@ -292,35 +940,62 @@ def bind_attention_cp_group(model: torch.nn.Module, packed_seq_params: Any) -> A raise ValueError("Dynamic CP attention requires PP=1") tp_cp_group = _DYNAMIC_TP_CP_GROUPS.get(context.size) if tp_cp_group is None: - raise ValueError(f"No dynamic TP*CP group was initialized for CP={context.size}") + raise ValueError( + f"No dynamic TP*CP group was initialized for CP={context.size}" + ) expected_tp_cp_size = ( context.size * parallel_state.get_tensor_model_parallel_world_size() ) if tp_cp_group.size() != expected_tp_cp_size: raise ValueError("Dynamic MoE TP*CP group has the wrong size") - padding_only = bool( - getattr(packed_seq_params, "dynamic_cp_padding_only", False) - ) + padding_only = bool(getattr(packed_seq_params, "dynamic_cp_padding_only", False)) for module in model.modules(): - if isinstance(module, Attention): - module.pg_collection.cp = group + collection = getattr(module, "pg_collection", None) + if collection is not None and hasattr(collection, "cp"): + collection.cp = group + if hasattr(collection, "tp_cp"): + collection.tp_cp = tp_cp_group + if _uses_direct_cp_group(module) and hasattr(module, "cp_group"): + module.cp_group = group + if hasattr(module, "tp_cp_group"): + module.tp_cp_group = tp_cp_group + if _is_hybrid_stack(module): + _bind_hybrid_stack_layout(module, group=group, tp_cp_group=tp_cp_group) if isinstance(module, Router): module.cp_group = group module.tp_cp_group = tp_cp_group _bind_router_config(module, padding_only=padding_only) if _is_mamba_mixer(module): - module.pg_collection.cp = group _rebuild_mamba_cp(module, group) + elif _is_gated_delta_product(module): + _rebuild_gdp_cp(module, group) elif _is_gated_delta_net(module): - module.pg_collection.cp = group + baseline_size = module.cp_size + baseline_split = getattr(module, "feat_dim_split", None) module.cp_size = context.size + if baseline_split is not None: + scaled_split = [] + for value in baseline_split: + numerator = value * baseline_size + if numerator % context.size: + raise ValueError( + "GatedDeltaNet projection dimensions are not divisible " + f"by runtime CP={context.size}" + ) + scaled_split.append(numerator // context.size) + module.feat_dim_split = tuple(scaled_split) model_packed = copy(packed_seq_params) model_packed.cp_group = group return model_packed def planned_microbatches( - data: Any, plan: CPRankPlan, step: CPRankStep, straggler_timer: Any + data: Any, + plan: CPRankPlan, + step: CPRankStep, + straggler_timer: Any, + *, + prepad_packed_seq_for_hybridep: bool = False, ) -> Iterator[Any]: """Yield the lane's uneven task list with explicit group boundaries.""" # Avoid a cycle: data.py dispatches to this iterator. @@ -355,7 +1030,7 @@ def planned_microbatches( raise ValueError( "Active CP group disagrees with the driver's assignment" ) - expert_group = parallel_state.get_expert_model_parallel_group() + expert_group = parallel_state.get_expert_tensor_and_model_parallel_group() tp_size = parallel_state.get_tensor_model_parallel_world_size() task_ranks = { base + offset @@ -368,7 +1043,7 @@ def planned_microbatches( torch.distributed.get_process_group_ranks(expert_group) ).issubset(task_ranks): raise ValueError( - "Expert communication group crosses dynamic CP task boundaries" + "Joint expert TP*EP group crosses dynamic CP task boundaries" ) context = RuntimeCPContext(size=size, rank=rank, group=group) if assignment.sample_indices: @@ -386,16 +1061,19 @@ def planned_microbatches( seq_length_key="input_lengths", pack_sequences=True, pad_individual_seqs_to_multiple_of=assignment.pad_multiple, + pad_packed_seq_to_multiple_of=assignment.pad_multiple, straggler_timer=straggler_timer, cp_context=context, create_packed_seq_padding_mask=True, + prepad_packed_seq_for_hybridep=prepad_packed_seq_for_hybridep, ) padding_only = not assignment.sample_indices inputs.packed_seq_params.dynamic_cp_padding_only = padding_only if padding_only: # MCore uses True to mean "padding". Excluding every physical - # placeholder token keeps expert-bias counters clean; router aux - # coefficients are disabled for this forward to avoid a 0/0 loss. + # placeholder token keeps expert-bias counters clean. Local aux + # coefficients are disabled; an enabled global aux loss still + # enters its aligned full-domain collective with zero tokens. inputs.padding_mask = torch.ones_like( inputs.input_ids_cp_sharded, dtype=torch.bool ) diff --git a/nemo_rl/models/megatron/hybridep.py b/nemo_rl/models/megatron/hybridep.py index bf55abb9452..a86c8f35f5a 100644 --- a/nemo_rl/models/megatron/hybridep.py +++ b/nemo_rl/models/megatron/hybridep.py @@ -57,9 +57,10 @@ def configure_hybridep_packed_input_padding( raise ValueError( "HybridEP input prepadding currently requires pipeline parallel size 1." ) - if megatron_cfg.get("mtp_num_layers"): + dynamic_cp = megatron_cfg.get("dynamic_context_parallel") or {} + if megatron_cfg.get("mtp_num_layers") and not dynamic_cp.get("enabled", False): raise ValueError( - "HybridEP input prepadding currently requires MTP disabled." + "HybridEP input prepadding with MTP currently requires Dynamic CP." ) @@ -180,10 +181,11 @@ def pad_packed_seq_for_hybridep( max_last_sequence_len = int(cu_seqlens_padded[-1] - cu_seqlens_padded[-2]) max_seqlen = max(int(packed_seq_params.max_seqlen_q), max_last_sequence_len) + # Preserve the pre-padding sequence boundaries. MTP must not reinterpret + # the new HybridEP-only tail as tokens; the *_padded boundaries describe + # the enlarged physical buffer consumed by attention and HybridEP. packed_seq_params = replace( packed_seq_params, - cu_seqlens_q=cu_seqlens_padded, - cu_seqlens_kv=cu_seqlens_padded, cu_seqlens_q_padded=cu_seqlens_padded, cu_seqlens_kv_padded=cu_seqlens_padded, max_seqlen_q=max_seqlen, diff --git a/nemo_rl/models/policy/dynamic_cp.py b/nemo_rl/models/policy/dynamic_cp.py index 0b3003f85d7..c359488b28c 100644 --- a/nemo_rl/models/policy/dynamic_cp.py +++ b/nemo_rl/models/policy/dynamic_cp.py @@ -49,7 +49,7 @@ def _enabled_global_aux_loss(megatron_cfg: dict[str, Any]) -> bool: """Return whether a full-DP global aux collective is configured. The coefficient can originate in the HF model provider rather than this - dictionary, so the routing type itself must be rejected. + dictionary, so the routing type is the stable scheduling signal. """ return "global_aux_loss" in _routing_types(megatron_cfg) @@ -57,27 +57,27 @@ def _enabled_global_aux_loss(megatron_cfg: dict[str, Any]) -> bool: def _minimum_cp_size_for_experts( megatron_cfg: dict[str, Any], configured_minimum: int ) -> int: - """Keep every EP collective inside one dynamically scheduled task. + """Keep every joint ETP*EP block inside one dynamically scheduled task. Dynamic CP tasks contain ``CP * TP`` contiguous model ranks. The pinned - MCore rank order puts a complete EP group inside such a block once it is at - least EP ranks wide. ETP changes that layout, so support it only after it - has a dedicated topology implementation. + MCore rank order makes expert TP the fastest expert-grid axis followed by + EP. A task is therefore closed over all expert collectives once its rank + block is a multiple of ``ETP * EP``. """ expert_parallel = megatron_cfg["expert_model_parallel_size"] tensor_parallel = megatron_cfg["tensor_model_parallel_size"] expert_tensor_parallel = megatron_cfg.get("expert_tensor_parallel_size", 1) - if expert_tensor_parallel != 1: - raise ValueError("Dynamic CP MoE requires expert_tensor_parallel_size=1") - if expert_parallel <= 1: + expert_block = expert_parallel * expert_tensor_parallel + if expert_block <= 1: return configured_minimum minimum = configured_minimum - while minimum * tensor_parallel < expert_parallel: + while minimum * tensor_parallel < expert_block: minimum *= 2 - if (minimum * tensor_parallel) % expert_parallel: + if (minimum * tensor_parallel) % expert_block: raise ValueError( - "Dynamic CP requires expert_model_parallel_size to divide " - "min_dynamic_cp_size * tensor_model_parallel_size" + "Dynamic CP requires expert_tensor_parallel_size * " + "expert_model_parallel_size to divide min_dynamic_cp_size * " + "tensor_model_parallel_size" ) return minimum @@ -110,22 +110,8 @@ def validate_dynamic_cp(cfg: dict[str, Any], *, lanes: int) -> None: ) if mc.get("cuda_graph_impl") not in (None, "none"): raise ValueError("Dynamic CP does not support CUDA graph capture") - if _model_setting(mc, "mtp_num_layers") or _model_setting( - mc, "moe_hybridep_prepad_packed_inputs" - ): - raise ValueError("Dynamic CP does not support MTP or HybridEP input prepadding") if _model_setting(mc, "overlap_moe_expert_parallel_comm"): - raise ValueError( - "Dynamic CP does not support overlap_moe_expert_parallel_comm" - ) - if cfg["sequence_packing"].get("pair_grouping_key"): - raise ValueError("Dynamic CP does not yet schedule atomic preference pairs") - if _enabled_global_aux_loss(mc): - raise ValueError( - "Dynamic CP does not support global_aux_loss: its per-forward " - "TP*DP*CP collective cannot be called by uneven CP task lists. " - "Use aux_loss or seq_aux_loss instead." - ) + raise ValueError("Dynamic CP does not support overlap_moe_expert_parallel_comm") if "quantile_balancing" in _routing_types(mc): raise ValueError( "Dynamic CP does not support quantile_balancing because its router " @@ -137,8 +123,8 @@ def validate_dynamic_cp(cfg: dict[str, Any], *, lanes: int) -> None: maximum = dynamic.max_size or lanes if minimum > maximum: raise ValueError( - "Dynamic CP cannot contain an EP group: effective min_size " - f"{minimum} exceeds max_size {maximum} (CP*TP must be >= EP)" + "Dynamic CP cannot contain a joint ETP*EP group: effective min_size " + f"{minimum} exceeds max_size {maximum} (CP*TP must be >= ETP*EP)" ) plan_cp_phases( [], @@ -174,6 +160,8 @@ class CPBatchSchedule: sequence_parallel_size: int user_pad_multiple: int token_alignment: int + pair_grouping: tuple[int, ...] | None + align_full_domain_collectives: bool groups_by_batch: tuple[tuple[CPSyncGroup, ...], ...] @@ -190,10 +178,7 @@ def owned_real_task_count(plan: CPRankPlan) -> int: def real_task_participation_count(plan: CPRankPlan) -> int: """Count real task calls made by this lane across all optimizer steps.""" return sum( - 1 - for step in plan.steps - for task in step.assignments - if task.sample_indices + 1 for step in plan.steps for task in step.assignments if task.sample_indices ) @@ -242,6 +227,36 @@ def _input_lengths(data: BatchedDataDict) -> tuple[int, ...]: return tuple(int(value) for value in values) +def _pair_grouping( + data: BatchedDataDict, cfg: dict[str, Any] +) -> tuple[int, ...] | None: + """Return stable atomic-group ids requested by sequence packing.""" + grouping_key = cfg["sequence_packing"].get("pair_grouping_key") + if grouping_key is None: + return None + if grouping_key not in data: + raise KeyError( + f"sequence_packing pair_grouping_key={grouping_key!r} is not present in the batch" + ) + values = data[grouping_key] + if hasattr(values, "tolist"): + values = values.tolist() + if len(values) != data.size: + raise ValueError("Dynamic CP pair-group ids must have one value per sample") + return tuple(int(value) for value in values) + + +def _atomic_groups_for_batch( + grouping: tuple[int, ...] | None, *, start: int, batch_size: int +) -> list[tuple[int, ...]] | None: + if grouping is None: + return None + members_by_group: dict[int, list[int]] = {} + for index, group_id in enumerate(grouping[start : start + batch_size]): + members_by_group.setdefault(group_id, []).append(index) + return [tuple(members) for _, members in sorted(members_by_group.items())] + + def build_cp_schedule( data: BatchedDataDict, cfg: dict[str, Any], @@ -261,6 +276,22 @@ def build_cp_schedule( user_pad_multiple, alignment, ) = _schedule_parameters(cfg, sharding) + pair_grouping = _pair_grouping(data, cfg) + if pair_grouping is not None: + batches_by_group: dict[int, set[int]] = {} + for index, group_id in enumerate(pair_grouping): + batches_by_group.setdefault(group_id, set()).add(index // gbs) + split_groups = [ + group_id + for group_id, batch_ids in batches_by_group.items() + if len(batch_ids) > 1 + ] + if split_groups: + raise ValueError( + "Dynamic CP atomic groups cannot cross optimizer-step batch " + f"boundaries; split group ids: {split_groups[:8]}" + ) + align_full_domain_collectives = _enabled_global_aux_loss(cfg["megatron_cfg"]) groups_by_batch = tuple( plan_cp_phases( list(lengths[start : start + gbs]), @@ -271,6 +302,10 @@ def build_cp_schedule( sequence_parallel_size=sequence_parallel_size, user_pad_multiple=user_pad_multiple, token_alignment=alignment, + atomic_groups=_atomic_groups_for_batch( + pair_grouping, start=start, batch_size=gbs + ), + align_full_domain_collectives=align_full_domain_collectives, ) for start in range(0, len(lengths), gbs) ) @@ -284,6 +319,8 @@ def build_cp_schedule( sequence_parallel_size, user_pad_multiple, alignment, + pair_grouping, + align_full_domain_collectives, groups_by_batch, ) @@ -304,6 +341,9 @@ def cp_schedule_matches( return False return ( schedule.input_lengths == _input_lengths(data) + and schedule.pair_grouping == _pair_grouping(data, cfg) + and schedule.align_full_domain_collectives + == _enabled_global_aux_loss(cfg["megatron_cfg"]) and schedule.batch_size == gbs and ( schedule.lanes, diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index 28063ec8abe..2a61d0e3e37 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -108,6 +108,7 @@ ) from nemo_rl.models.policy import PolicyConfig from nemo_rl.models.policy.dynamic_cp import ( + _enabled_global_aux_loss, dynamic_cp_config, owned_real_task_count, real_task_participation_count, @@ -802,8 +803,6 @@ def __init__( if self.delegate_pack_to_model or self.model_slices_context_parallel_inputs: raise ValueError("Dynamic CP requires NeMo-owned text-model packing") model_config = self._get_model_config() - if getattr(model_config, "mtp_num_layers", None): - raise ValueError("Dynamic CP does not yet support MTP") if getattr( model_config, "overlap_moe_expert_parallel_comm", @@ -812,20 +811,25 @@ def __init__( raise ValueError( "Dynamic CP does not support overlap_moe_expert_parallel_comm" ) - if getattr(model_config, "moe_hybridep_prepad_packed_inputs", False): - raise ValueError("Dynamic CP does not support HybridEP input prepadding") - routing_type = getattr( - model_config, "moe_router_load_balancing_type", None - ) + routing_type = getattr(model_config, "moe_router_load_balancing_type", None) routing_types = ( list(routing_type) if isinstance(routing_type, (list, tuple)) else [routing_type] ) - if "global_aux_loss" in routing_types: - raise ValueError("Dynamic CP does not support global_aux_loss") + if "global_aux_loss" in routing_types and not _enabled_global_aux_loss( + self.cfg["megatron_cfg"] + ): + raise ValueError( + "The model provider enabled global_aux_loss without exposing it " + "in megatron_cfg; Dynamic CP needs the routing type before " + "dispatch so it can align collective rounds" + ) if "quantile_balancing" in routing_types: raise ValueError("Dynamic CP does not support quantile_balancing") + from nemo_rl.models.megatron.dynamic_cp import validate_dynamic_cp_model + + validate_dynamic_cp_model(self.model) if self.model_slices_context_parallel_inputs: if self.delegate_pack_to_model: @@ -1363,12 +1367,16 @@ def train( global_participations, group=dynamic_moe_group, ) - dynamic_avg_loss_scale = 1.0 / max( - 1, int(global_participations.item()) + dynamic_avg_loss_scale = 1.0 / max(1, int(global_participations.item())) + dynamic_global_loss_scale = ( + 1.0 / max(1, total_num_microbatches) + if _enabled_global_aux_loss(self.cfg["megatron_cfg"]) + else None ) else: moe_loss_scale = 1.0 / max(1, total_num_microbatches) dynamic_avg_loss_scale = None + dynamic_global_loss_scale = None moe_metrics = get_moe_metrics( loss_scale=moe_loss_scale, per_layer_logging=self.cfg["megatron_cfg"]["moe_per_layer_logging"], @@ -1381,6 +1389,7 @@ def train( track_names=get_aux_loss_track_names(model_config), dynamic_parallel_group=dynamic_moe_group, dynamic_avg_loss_scale=dynamic_avg_loss_scale, + dynamic_global_loss_scale=dynamic_global_loss_scale, ) if moe_metrics: metrics["moe_metrics"] = moe_metrics @@ -2889,20 +2898,32 @@ def _collect_mtp_metrics( metrics: Metrics dict to populate with MTP metrics (under "mtp_metrics"). total_num_microbatches: Microbatches accumulated this step. The MTP loss logging helper sums the per-microbatch loss without dividing, so we pass - 1/total_num_microbatches to recover the mean (mirroring the MoE path). + 1/total_num_microbatches to recover the static mean. Dynamic CP instead + reduces raw loss sums and token counts over its fixed lane group. mtp_grad_norm: The MTP parameter group's gradient norm, already reduced across the model-parallel group, or None when unavailable (e.g. clip_grad == 0 or mtp_detach_heads=False). Logged under "mtp_metrics" as "grad_norm". """ mtp_num_layers = getattr(self.model.config, "mtp_num_layers", None) if mtp_num_layers is not None and mtp_num_layers > 0: - from nemo_rl.models.megatron.common import get_mtp_metrics - # MTP layers live only on the last pipeline stage, so the tracker is # populated there alone. Broadcast to all stages so downstream metric # aggregation (which reads rank 0's results) sees them when PP > 1. - mtp_loss_scale = 1.0 / max(1, total_num_microbatches) - mtp_metrics = get_mtp_metrics(loss_scale=mtp_loss_scale) + if dynamic_cp_config(getattr(self, "cfg", {})) is not None: + from nemo_rl.models.megatron.dynamic_cp import ( + get_dynamic_mtp_metrics, + ) + + mtp_metrics = get_dynamic_mtp_metrics( + parallel_group=parallel_state.get_data_parallel_group( + with_context_parallel=True + ) + ) + else: + from nemo_rl.models.megatron.common import get_mtp_metrics + + mtp_loss_scale = 1.0 / max(1, total_num_microbatches) + mtp_metrics = get_mtp_metrics(loss_scale=mtp_loss_scale) mtp_metrics = broadcast_loss_metrics_from_last_stage(mtp_metrics) # mtp_grad_norm is already MP-reduced (same value on every rank); expose it # under the "mtp/" namespace so it logs as train/mtp/grad_norm. diff --git a/tests/unit/distributed/test_dynamic_cp_dispatch.py b/tests/unit/distributed/test_dynamic_cp_dispatch.py index 7a5ed828c8b..a80a5c93e5b 100644 --- a/tests/unit/distributed/test_dynamic_cp_dispatch.py +++ b/tests/unit/distributed/test_dynamic_cp_dispatch.py @@ -91,7 +91,7 @@ def test_moe_minimum_contains_complete_expert_group(self): ), 8, ) - with self.assertRaisesRegex(ValueError, "expert_tensor_parallel_size=1"): + self.assertEqual( _minimum_cp_size_for_experts( { "tensor_model_parallel_size": 2, @@ -99,8 +99,10 @@ def test_moe_minimum_contains_complete_expert_group(self): "expert_model_parallel_size": 8, }, 1, - ) - with self.assertRaisesRegex(ValueError, "expert_tensor_parallel_size=1"): + ), + 8, + ) + self.assertEqual( _minimum_cp_size_for_experts( { "tensor_model_parallel_size": 2, @@ -108,7 +110,9 @@ def test_moe_minimum_contains_complete_expert_group(self): "expert_model_parallel_size": 1, }, 1, - ) + ), + 1, + ) with self.assertRaisesRegex(ValueError, "to divide"): _minimum_cp_size_for_experts( { @@ -119,7 +123,7 @@ def test_moe_minimum_contains_complete_expert_group(self): 1, ) - def test_global_aux_loss_is_rejected_even_with_provider_coefficient(self): + def test_global_aux_loss_is_detected_even_with_provider_coefficient(self): self.assertTrue( _enabled_global_aux_loss( {"moe_router_load_balancing_type": "global_aux_loss"} @@ -145,9 +149,39 @@ def test_global_aux_loss_is_rejected_even_with_provider_coefficient(self): ) ) self.assertFalse( - _enabled_global_aux_loss( - {"moe_router_load_balancing_type": "seq_aux_loss"} - ) + _enabled_global_aux_loss({"moe_router_load_balancing_type": "seq_aux_loss"}) + ) + + def test_global_aux_loss_aligns_full_domain_collective_rounds(self): + data = BatchedDataDict( + input_ids=torch.zeros(5, 12, dtype=torch.long), + input_lengths=torch.full((5,), 10), + sample_mask=torch.ones(5, dtype=torch.long), + token_mask=torch.ones(5, 12, dtype=torch.long), + ) + cfg = self._validation_cfg() + cfg["megatron_cfg"].update( + { + "moe_router_load_balancing_type": "global_aux_loss", + "dynamic_context_parallel": { + "enabled": True, + "tokens_per_rank": 10, + "max_size": 1, + }, + } + ) + + validate_dynamic_cp(cfg, lanes=4) + dispatch = build_cp_dispatch(data, cfg, self.mesh, batch_size=5, training=True) + local_counts = [ + len(plan.steps[0].assignments) for plans in dispatch.plans for plan in plans + ] + assert local_counts == [2, 2, 2, 2] + assert any( + not task.sample_indices + for plans in dispatch.plans + for plan in plans + for task in plan.steps[0].assignments ) def _validation_cfg(self) -> dict: @@ -193,6 +227,70 @@ def test_dynamic_cp_validation_rejects_moe_microbatch_overlap(self): with self.assertRaisesRegex(ValueError, "overlap_moe_expert_parallel_comm"): validate_dynamic_cp(cfg, lanes=4) + def test_dynamic_cp_validation_allows_mtp_and_hybridep_prepad_separately(self): + mtp_cfg = self._validation_cfg() + mtp_cfg["megatron_cfg"]["mtp_num_layers"] = 1 + validate_dynamic_cp(mtp_cfg, lanes=4) + + hybridep_cfg = self._validation_cfg() + hybridep_cfg["megatron_cfg"].update( + { + "moe_token_dispatcher_type": "flex", + "moe_flex_dispatcher_backend": "hybridep", + "moe_hybridep_prepad_packed_inputs": True, + } + ) + validate_dynamic_cp(hybridep_cfg, lanes=4) + + def test_dynamic_cp_validation_allows_mtp_with_hybridep_prepad(self): + cfg = self._validation_cfg() + cfg["megatron_cfg"].update( + { + "mtp_num_layers": 1, + "moe_hybridep_prepad_packed_inputs": True, + } + ) + validate_dynamic_cp(cfg, lanes=4) + + def test_dynamic_cp_keeps_preference_pairs_atomic(self): + data = BatchedDataDict( + input_ids=torch.arange(4 * 16).reshape(4, 16), + input_lengths=torch.tensor([9, 3, 8, 4]), + pair_index=torch.tensor([10, 10, 20, 20]), + sample_mask=torch.ones(4, dtype=torch.long), + token_mask=torch.ones(4, 16, dtype=torch.long), + ) + cfg = self._validation_cfg() + cfg["sequence_packing"]["pair_grouping_key"] = "pair_index" + + validate_dynamic_cp(cfg, lanes=4) + dispatch = build_cp_dispatch(data, cfg, self.mesh, batch_size=4, training=True) + + real_assignments = [ + task.sample_indices + for group in dispatch.schedule.groups_by_batch[0] + for task in group.assignments + if task.sample_indices + ] + for pair in ((0, 1), (2, 3)): + assert any( + set(pair).issubset(assignment) for assignment in real_assignments + ) + + def test_dynamic_cp_rejects_atomic_pair_split_across_steps(self): + data = BatchedDataDict( + input_ids=torch.arange(4 * 16).reshape(4, 16), + input_lengths=torch.tensor([9, 3, 8, 4]), + pair_index=torch.tensor([10, 20, 10, 20]), + sample_mask=torch.ones(4, dtype=torch.long), + token_mask=torch.ones(4, 16, dtype=torch.long), + ) + cfg = self._validation_cfg() + cfg["sequence_packing"]["pair_grouping_key"] = "pair_index" + + with self.assertRaisesRegex(ValueError, "cannot cross"): + build_cp_dispatch(data, cfg, self.mesh, batch_size=2, training=True) + def test_outputs_from_nonzero_static_cp_are_preserved(self): dispatch = build_cp_dispatch( self.data, self.cfg, self.mesh, batch_size=None, training=False diff --git a/tests/unit/models/megatron/test_dynamic_cp_moe.py b/tests/unit/models/megatron/test_dynamic_cp_moe.py index ddc8580dd5d..43c33251ca8 100644 --- a/tests/unit/models/megatron/test_dynamic_cp_moe.py +++ b/tests/unit/models/megatron/test_dynamic_cp_moe.py @@ -53,6 +53,42 @@ def __init__( self.D_has_hdim = D_has_hdim +class _GDPHelper: + def __init__( + self, + *, + cp_group, + d_inner_local_tp=32, + nheads_local_tp=16, + ngroups_local_tp=4, + d_state=8, + num_householder=2, + headdim=2, + conv1d_cp1=None, + dt_bias_cp1=None, + A_log_cp1=None, + D_cp1=None, + D_has_hdim=False, + sequence_is_contiguous=False, + ): + self.cp_group = cp_group + self.d_inner_local_tp = d_inner_local_tp + self.nheads_local_tp = nheads_local_tp + self.ngroups_local_tp = ngroups_local_tp + self.d_state = d_state + self.num_householder = num_householder + self.headdim = headdim + self.conv1d_cp1 = conv1d_cp1 + self.dt_bias_cp1 = dt_bias_cp1 + self.A_log_cp1 = A_log_cp1 + self.D_cp1 = D_cp1 + self.D_has_hdim = D_has_hdim + self.sequence_is_contiguous = sequence_is_contiguous + self.d_inner_local_tpcp = d_inner_local_tp // cp_group.size() + self.nheads_local_tpcp = nheads_local_tp // cp_group.size() + self.ngroups_local_tpcp = max(1, ngroups_local_tp // cp_group.size()) + + class _Mamba(torch.nn.Module): def __init__(self, group): super().__init__() @@ -65,11 +101,22 @@ def __init__(self, group): super().__init__() self.pg_collection = SimpleNamespace(cp=group) self.cp_size = group.size() + self.feat_dim_split = (32, 16, 8, 8) _GatedDelta.__module__ = "megatron.core.ssm.gated_delta_net" +class _GatedDeltaProduct(torch.nn.Module): + def __init__(self, group): + super().__init__() + self.pg_collection = SimpleNamespace(cp=group) + self.cp = _GDPHelper(cp_group=group) + self.d_inner_local_cp = self.cp.d_inner_local_tpcp + self.nheads_local_cp = self.cp.nheads_local_tpcp + self.ngroups_local_cp = self.cp.ngroups_local_tpcp + + class _Router(torch.nn.Module): def __init__(self, group, config): super().__init__() @@ -90,6 +137,7 @@ def test_dynamic_binding_updates_router_and_ssm_then_restores(monkeypatch): model = torch.nn.Module() model.add_module("mamba", _Mamba(original_cp)) model.add_module("gdn", _GatedDelta(original_cp)) + model.add_module("gdp", _GatedDeltaProduct(original_cp)) model.add_module("router", _Router(original_tp_cp, config)) packed = SimpleNamespace( local_cp_size=2, @@ -106,6 +154,7 @@ def test_dynamic_binding_updates_router_and_ssm_then_restores(monkeypatch): monkeypatch.setitem(dynamic_cp._DYNAMIC_TP_CP_GROUPS, 2, active_tp_cp) original_mamba_helper = model.mamba.cp + original_gdp_helper = model.gdp.cp with dynamic_cp.preserve_attention_cp_groups(model): model_packed = dynamic_cp.bind_attention_cp_group(model, packed) assert model_packed.cp_group is active_cp @@ -116,6 +165,10 @@ def test_dynamic_binding_updates_router_and_ssm_then_restores(monkeypatch): assert model.mamba.cp.cp_group is active_cp assert model.gdn.pg_collection.cp is active_cp assert model.gdn.cp_size == 2 + assert model.gdn.feat_dim_split == (16, 8, 4, 4) + assert model.gdp.cp is not original_gdp_helper + assert model.gdp.cp.cp_group is active_cp + assert model.gdp.d_inner_local_cp == 16 assert config.moe_aux_loss_coeff == [0.0, 0.0] assert config.moe_z_loss_coeff is None @@ -125,17 +178,51 @@ def test_dynamic_binding_updates_router_and_ssm_then_restores(monkeypatch): assert model.mamba.cp is original_mamba_helper assert model.gdn.pg_collection.cp is original_cp assert model.gdn.cp_size == 1 + assert model.gdn.feat_dim_split == (32, 16, 8, 8) + assert model.gdp.cp is original_gdp_helper + assert model.gdp.d_inner_local_cp == 32 + assert config.moe_aux_loss_coeff == [0.1, 0.2] + assert config.moe_z_loss_coeff == 0.01 + + +def test_dynamic_padding_task_keeps_only_global_aux_loss(monkeypatch): + from nemo_rl.models.megatron import dynamic_cp + + group = _Group(1) + config = SimpleNamespace( + moe_aux_loss_coeff=[0.1, 0.2], + moe_z_loss_coeff=0.01, + moe_router_load_balancing_type=["aux_loss", "global_aux_loss"], + ) + router = _Router(group, config) + model = torch.nn.Module() + model.add_module("router", router) + packed = SimpleNamespace( + local_cp_size=1, + cp_group=None, + dynamic_cp_padding_only=True, + ) + + monkeypatch.setattr(dynamic_cp, "Router", _Router) + monkeypatch.setattr( + dynamic_cp.parallel_state, "get_pipeline_model_parallel_group", lambda: group + ) + monkeypatch.setattr( + dynamic_cp.parallel_state, "get_tensor_model_parallel_world_size", lambda: 1 + ) + monkeypatch.setitem(dynamic_cp._DYNAMIC_TP_CP_GROUPS, 1, group) + + with dynamic_cp.preserve_attention_cp_groups(model): + dynamic_cp.bind_attention_cp_group(model, packed) + assert config.moe_aux_loss_coeff == [0.0, 0.2] + assert config.moe_z_loss_coeff is None + assert config.moe_aux_loss_coeff == [0.1, 0.2] assert config.moe_z_loss_coeff == 0.01 -@pytest.mark.parametrize( - ("active_size", "group_tokens"), - [(1, 3), (4, 10)], -) -def test_dynamic_moe_scaling_uses_exact_active_group_token_count( - monkeypatch, active_size, group_tokens -): +@pytest.mark.parametrize("active_size", [1, 4]) +def test_dynamic_moe_scaling_is_applied_by_router_attachment(monkeypatch, active_size): from nemo_rl.models.megatron import dynamic_cp active_tp_cp = _Group(active_size) @@ -146,25 +233,64 @@ def test_dynamic_moe_scaling_uses_exact_active_group_token_count( monkeypatch.setattr(dynamic_cp, "Router", _Router) - def _all_reduce(value, *, group): - assert group is active_tp_cp - value.fill_(group_tokens) - - monkeypatch.setattr(dynamic_cp.torch.distributed, "all_reduce", _all_reduce) + monkeypatch.setattr( + dynamic_cp.torch.distributed, + "all_reduce", + lambda *_args, **_kwargs: pytest.fail( + "pre-forward validation performed a token reduction" + ), + ) with dynamic_cp.preserve_attention_cp_groups(model): dynamic_cp.configure_dynamic_moe_loss_scaling(model, padding_mask) - correction = group_tokens / (3 * active_size) - assert dynamic_cp.dynamic_moe_grad_scale_correction(config) == pytest.approx( - correction - ) - assert config.moe_z_loss_coeff == pytest.approx(0.02 / correction) + assert dynamic_cp.dynamic_moe_grad_scale_correction(config) == 1.0 + assert config.moe_z_loss_coeff == 0.02 assert dynamic_cp.dynamic_moe_grad_scale_correction(config) == 1.0 assert config.moe_aux_loss_coeff == 0.1 assert config.moe_z_loss_coeff == 0.02 +def test_dynamic_router_attachment_uses_exact_task_token_count(monkeypatch): + from megatron.core.transformer.moe import moe_logging, moe_utils + from nemo_rl.models.megatron import dynamic_cp + + tracker = SimpleNamespace(record=lambda *args, **kwargs: None) + attached = {} + + def _apply(_activation, aux_loss): + attached["aux_loss"] = aux_loss + return _activation + + monkeypatch.setattr(moe_logging, "get_moe_metrics_tracker", lambda: tracker) + monkeypatch.setattr(moe_utils.MoEAuxLossAutoScaler, "apply", _apply) + + router = SimpleNamespace( + is_mtp_layer=False, + layer_number=1, + calculate_per_token_loss=True, + config=SimpleNamespace( + mtp_use_repeated_layer=False, + mtp_num_layers=None, + num_layers=2, + ), + _nemo_dynamic_aux_scale_tokens=torch.tensor(10.0), + ) + activation = torch.ones(1) + result = dynamic_cp._dynamic_attach_and_log_load_balancing_loss( + router, + activation, + aux_loss_coeff=0.1, + aux_loss=torch.tensor(2.0), + aux_loss_name="load_balancing_loss", + reduce_group=_Group(1), + valid_token_count=torch.tensor(3.0), + ) + + assert result is activation + assert attached["aux_loss"].item() == pytest.approx(20.0) + + def test_dynamic_moe_scaling_is_noop_for_dense_model(monkeypatch): from nemo_rl.models.megatron import dynamic_cp @@ -175,3 +301,182 @@ def test_dynamic_moe_scaling_is_noop_for_dense_model(monkeypatch): ) dynamic_cp.configure_dynamic_moe_loss_scaling(torch.nn.Linear(2, 2), None) + + +def test_runtime_mtp_counts_reduce_before_division(monkeypatch): + from nemo_rl.models.megatron import dynamic_cp + + group = _Group(4) + + def _all_reduce(counts, *, op, group): + del op + assert group.size() == 4 + counts.copy_(torch.tensor([11.0, 7.0])) + + monkeypatch.setattr(dynamic_cp.torch.distributed, "all_reduce", _all_reduce) + main_tokens, mtp_tokens = dynamic_cp._runtime_mtp_token_counts( + torch.tensor(3.0), torch.tensor(1.0), group + ) + + assert main_tokens.item() == 11 + assert mtp_tokens.item() == 7 + + +def test_dynamic_mtp_backward_uses_task_wide_token_ratio(monkeypatch): + from megatron.core.transformer import multi_token_prediction as mcore_mtp + from nemo_rl.models.megatron import dynamic_cp + + group = _Group(2) + + def _roll_tensor(tensor, **kwargs): + return tensor, tensor.sum() if kwargs.get("return_sum", True) else None + + def _all_reduce(counts, *, op, group): + del op + assert group.size() == 2 + counts.copy_(torch.tensor([8.0, 4.0])) + + monkeypatch.setattr(mcore_mtp, "roll_tensor", _roll_tensor) + monkeypatch.setattr(dynamic_cp.torch.distributed, "all_reduce", _all_reduce) + mcore_mtp.MTPLossAutoScaler.set_loss_scale(torch.tensor(1.0)) + + hidden_states = torch.ones(4, 1, requires_grad=True) + labels = torch.zeros(2, 1, dtype=torch.long) + loss_mask = torch.ones(2, 1) + config = SimpleNamespace( + mtp_num_layers=1, + mtp_detach_heads=False, + mtp_loss_scaling_factor=1.0, + calculate_per_token_loss=True, + ) + + output = dynamic_cp._dynamic_process_mtp_loss( + hidden_states=hidden_states, + labels=labels, + loss_mask=loss_mask, + output_layer=lambda value, **_kwargs: (value, None), + output_weight=None, + runtime_gather_output=False, + is_training=False, + compute_language_model_loss=lambda _labels, logits: logits, + config=config, + cp_group=group, + ) + output.sum().backward() + + torch.testing.assert_close(hidden_states.grad[:2], torch.ones(2, 1)) + torch.testing.assert_close(hidden_states.grad[2:], torch.full((2, 1), 2.0)) + + +def test_dynamic_mtp_metrics_are_token_weighted(monkeypatch): + from nemo_rl.models.megatron import dynamic_cp + + dynamic_cp._DYNAMIC_MTP_METRICS.clear() + dynamic_cp._save_dynamic_mtp_metrics( + loss_sum=torch.tensor(6.0), + num_tokens=torch.tensor(2.0), + correct=torch.tensor(1.0), + total=torch.tensor(2.0), + layer_number=0, + num_layers=1, + ) + dynamic_cp._save_dynamic_mtp_metrics( + loss_sum=torch.tensor(9.0), + num_tokens=torch.tensor(3.0), + correct=torch.tensor(2.0), + total=torch.tensor(3.0), + layer_number=0, + num_layers=1, + ) + monkeypatch.setattr( + dynamic_cp.torch.distributed, + "all_reduce", + lambda _value, *, op, group: None, + ) + + metrics = dynamic_cp.get_dynamic_mtp_metrics(parallel_group=_Group(4)) + + assert metrics["mtp_1_loss"] == pytest.approx(3.0) + assert metrics["mtp_1_acceptance_rate"] == pytest.approx(60.0) + assert dynamic_cp._DYNAMIC_MTP_METRICS == {} + + +def test_dynamic_hybrid_mtp_receives_router_padding_mask(monkeypatch): + from nemo_rl.models.megatron import dynamic_cp + + class _MTPRouter(torch.nn.Module): + def __init__(self): + super().__init__() + self.config = SimpleNamespace( + calculate_per_token_loss=False, + moe_aux_loss_coeff=0.0, + ) + self.tp_cp_group = _Group(1) + self.seen_padding_mask = None + + def forward(self, hidden_states, padding_mask=None): + self.seen_padding_mask = padding_mask + return hidden_states + + class _MTP(torch.nn.Module): + def __init__(self): + super().__init__() + self.router = _MTPRouter() + self.seen_padding_mask = None + + def forward(self, *, padding_mask=None): + self.seen_padding_mask = padding_mask + self.router(torch.ones(1), padding_mask=padding_mask) + return self.router.seen_padding_mask + + class _HybridModel(torch.nn.Module): + def __init__(self): + super().__init__() + self.mtp = _MTP() + + def forward(self, *, padding_mask=None): + # This intentionally mirrors the pinned MCore bug: HybridModel has + # the mask, but its MTP call omits it. + return self.mtp() + + _HybridModel.__module__ = "megatron.core.models.hybrid.hybrid_model" + model = _HybridModel() + padding_mask = torch.tensor([[False, True]]) + monkeypatch.setattr(dynamic_cp, "Router", _MTPRouter) + + with dynamic_cp._patch_hybrid_mtp_padding_masks(model): + result = model(padding_mask=padding_mask) + + assert torch.equal(result, padding_mask) + assert torch.equal(model.mtp.seen_padding_mask, ~padding_mask) + model.mtp.seen_padding_mask = None + assert model(padding_mask=padding_mask) is None + + +def test_dynamic_model_validation_rejects_unmerged_mla_support(): + from nemo_rl.models.megatron import dynamic_cp + + class MLASelfAttention(torch.nn.Module): + pass + + model = torch.nn.Module() + model.add_module("mla", MLASelfAttention()) + + with pytest.raises(ValueError, match="unmerged MCore"): + dynamic_cp.validate_dynamic_cp_model(model) + + +def test_dynamic_model_validation_rejects_chunkwise_ssm(): + from nemo_rl.models.megatron import dynamic_cp + + class ChunkwiseMamba(_Mamba): + pass + + ChunkwiseMamba.__module__ = "megatron.core.ssm.mamba_mixer" + mamba = ChunkwiseMamba(_Group(1)) + mamba.config = SimpleNamespace(linear_cp_mode="chunkwise") + model = torch.nn.Module() + model.add_module("mamba", mamba) + + with pytest.raises(ValueError, match="headwise"): + dynamic_cp.validate_dynamic_cp_model(model) diff --git a/tests/unit/models/megatron/test_hybridep_data.py b/tests/unit/models/megatron/test_hybridep_data.py index 070173b64a0..fd9c050fd0c 100644 --- a/tests/unit/models/megatron/test_hybridep_data.py +++ b/tests/unit/models/megatron/test_hybridep_data.py @@ -20,6 +20,31 @@ import torch +@pytest.mark.mcore +def test_hybridep_mtp_prepad_requires_dynamic_cp(): + from nemo_rl.models.megatron.hybridep import ( + configure_hybridep_packed_input_padding, + ) + + megatron_cfg = { + "moe_flex_dispatcher_backend": "hybridep", + "moe_token_dispatcher_type": "flex", + "moe_hybridep_prepad_packed_inputs": True, + "pipeline_model_parallel_size": 1, + "mtp_num_layers": 1, + } + config = { + "megatron_cfg": megatron_cfg, + "sequence_packing": {"enabled": True}, + } + + with pytest.raises(ValueError, match="requires Dynamic CP"): + configure_hybridep_packed_input_padding(None, config) + + megatron_cfg["dynamic_context_parallel"] = {"enabled": True} + configure_hybridep_packed_input_padding(None, config) + + @pytest.mark.mcore def test_hybridep_prepads_packed_inputs_before_model_forward(): from megatron.core.packed_seq_params import PackedSeqParams @@ -30,10 +55,11 @@ def set_group_max(target, **_kwargs): target.fill_(14) input_ids = torch.tensor([[11, 12, 13, 0, 21, 22, 23, 24, 25, 0, 0, 0]]) + cu_seqlens = torch.tensor([0, 3, 8], dtype=torch.int32) cu_seqlens_padded = torch.tensor([0, 4, 12], dtype=torch.int32) packed_seq_params = PackedSeqParams( - cu_seqlens_q=cu_seqlens_padded, - cu_seqlens_kv=cu_seqlens_padded, + cu_seqlens_q=cu_seqlens, + cu_seqlens_kv=cu_seqlens, cu_seqlens_q_padded=cu_seqlens_padded, cu_seqlens_kv_padded=cu_seqlens_padded, max_seqlen_q=8, @@ -82,11 +108,56 @@ def set_group_max(target, **_kwargs): assert torch.equal(padded_input_ids[:, :12], input_ids) assert torch.count_nonzero(padded_input_ids[:, 12:]) == 0 assert torch.equal(padded_cu_seqlens, torch.tensor([0, 4, 16])) + assert torch.equal(padded_params.cu_seqlens_q, cu_seqlens) + assert torch.equal(padded_params.cu_seqlens_kv, cu_seqlens) + assert torch.equal(padded_params.cu_seqlens_q_padded, torch.tensor([0, 4, 16])) assert padded_params.total_tokens == 16 mock_get_group.assert_called_once_with(check_initialized=False) mock_all_reduce.assert_called_once() +@pytest.mark.mcore +def test_hybridep_prepadding_extends_mtp_mask_and_keeps_logical_boundary(): + from nemo_rl.models.megatron.data import process_microbatch + + data = { + "input_ids": torch.tensor([[11, 12, 13, 0, 0, 0], [21, 22, 23, 24, 25, 0]]), + "input_lengths": torch.tensor([3, 5]), + "mtp_loss_mask": torch.tensor([[1, 1, 0, 0, 0, 0], [1, 1, 1, 1, 0, 0]]), + } + + with ( + patch( + "nemo_rl.models.megatron.data.get_context_parallel_rank", + return_value=0, + ), + patch( + "nemo_rl.models.megatron.data.get_context_parallel_world_size", + return_value=1, + ), + patch( + "nemo_rl.models.megatron.hybridep._get_hybridep_aligned_seq_len", + return_value=24, + ), + ): + result = process_microbatch( + data, + seq_length_key="input_lengths", + pad_individual_seqs_to_multiple_of=4, + pad_packed_seq_to_multiple_of=8, + pack_sequences=True, + create_packed_seq_padding_mask=True, + prepad_packed_seq_for_hybridep=True, + ) + + assert result.input_ids_cp_sharded.shape == (1, 24) + assert result.mtp_loss_mask.shape == (1, 24) + assert torch.count_nonzero(result.mtp_loss_mask[:, 16:]) == 0 + assert torch.all(result.padding_mask[:, 16:]) + assert result.packed_seq_params.cu_seqlens_q[-1].item() == 16 + assert result.packed_seq_params.cu_seqlens_q_padded[-1].item() == 24 + + @pytest.mark.mcore def test_hybridep_prepadding_rejects_missing_alignment_group(): from nemo_rl.models.megatron import hybridep diff --git a/tests/unit/models/megatron/test_megatron_data.py b/tests/unit/models/megatron/test_megatron_data.py index 6d99a3d2254..e6ea93bd301 100644 --- a/tests/unit/models/megatron/test_megatron_data.py +++ b/tests/unit/models/megatron/test_megatron_data.py @@ -1348,6 +1348,37 @@ def test_get_microbatch_iterator_sequence_packing( is True ) + @patch("nemo_rl.models.megatron.data.planned_microbatches") + def test_dynamic_cp_forwards_hybridep_prepadding(self, mock_planned): + """Dynamic tasks opt into the same pre-forward HybridEP alignment.""" + from nemo_rl.models.megatron.data import get_microbatch_iterator + + mock_planned.return_value = iter([]) + data = {"input_ids": torch.zeros(1, 128, dtype=torch.long)} + cp_plan = MagicMock() + cp_step = MagicMock() + cp_step.assignments = (MagicMock(),) + cfg = { + "sequence_packing": {"enabled": True}, + "dynamic_batching": {"enabled": False}, + "megatron_cfg": { + "moe_token_dispatcher_type": "flex", + "moe_flex_dispatcher_backend": "hybridep", + "moe_hybridep_prepad_packed_inputs": True, + }, + } + + get_microbatch_iterator( + data=data, + cfg=cfg, + mbs=1, + straggler_timer=MagicMock(), + cp_plan=cp_plan, + cp_step=cp_step, + ) + + assert mock_planned.call_args.kwargs["prepad_packed_seq_for_hybridep"] is True + @patch("nemo_rl.models.megatron.data.get_and_validate_seqlen") @patch("nemo_rl.models.megatron.data.make_processed_microbatch_iterator") def test_get_microbatch_iterator_regular( diff --git a/tests/unit/models/megatron/test_moe_metrics.py b/tests/unit/models/megatron/test_moe_metrics.py index 6b7d2654cac..87576f1cd2c 100644 --- a/tests/unit/models/megatron/test_moe_metrics.py +++ b/tests/unit/models/megatron/test_moe_metrics.py @@ -209,6 +209,39 @@ def _all_reduce(values, *, group): assert metrics["z_loss"] == pytest.approx(0.75) +@pytest.mark.mcore +def test_dynamic_cp_global_aux_uses_aligned_round_scale(monkeypatch): + from nemo_rl.models import megatron as megatron_module + from nemo_rl.models.megatron.common import get_moe_metrics + + entry = SimpleNamespace(values=torch.tensor([2.0]), avg_group=None) + live_tracker = SimpleNamespace(metrics={"global_load_balancing_loss": entry}) + monkeypatch.setattr( + megatron_module.common, "get_moe_metrics_tracker", lambda: live_tracker + ) + monkeypatch.setattr( + megatron_module.common, + "get_moe_layer_wise_logging_tracker", + lambda: {"global_load_balancing_loss": {"values": entry.values}}, + ) + monkeypatch.setattr( + megatron_module.common.dist, + "all_reduce", + lambda values, *, group: values.mul_(2.0), + ) + monkeypatch.setattr( + megatron_module.common, "clear_aux_losses_tracker", lambda: None + ) + + metrics = get_moe_metrics( + loss_scale=0.1, + dynamic_parallel_group=object(), + dynamic_global_loss_scale=0.25, + ) + + assert metrics["global_load_balancing_loss"] == pytest.approx(1.0) + + @pytest.mark.mcore @pytest.mark.parametrize( "routing_type,aux_loss_coeff,z_loss_coeff,expected", diff --git a/tests/unit/models/megatron/test_mtp_metrics.py b/tests/unit/models/megatron/test_mtp_metrics.py index fb1a2c111df..e0d9e8dcad2 100644 --- a/tests/unit/models/megatron/test_mtp_metrics.py +++ b/tests/unit/models/megatron/test_mtp_metrics.py @@ -130,10 +130,11 @@ def test_get_mtp_metrics_default_loss_scale_is_identity(monkeypatch): assert get_mtp_metrics()["mtp_1_loss"] == pytest.approx(3.0) -def _fake_worker(mtp_num_layers): +def _fake_worker(mtp_num_layers, cfg=None): """A minimal stand-in for MegatronPolicyWorkerImpl for calling _collect_mtp_metrics.""" return SimpleNamespace( - model=SimpleNamespace(config=SimpleNamespace(mtp_num_layers=mtp_num_layers)) + model=SimpleNamespace(config=SimpleNamespace(mtp_num_layers=mtp_num_layers)), + cfg=cfg or {}, ) @@ -193,6 +194,49 @@ def test_collect_mtp_metrics_omits_grad_norm_when_none(monkeypatch): assert "grad_norm" not in metrics["mtp_metrics"] +@pytest.mark.mcore +def test_collect_mtp_metrics_uses_dynamic_token_weighted_tracker(monkeypatch): + """Dynamic CP bypasses MCore's stale fixed-group mean-of-means tracker.""" + from nemo_rl.models.policy.workers import megatron_policy_worker as mpw + + fixed_group = object() + monkeypatch.setattr( + mpw.parallel_state, + "get_data_parallel_group", + lambda *, with_context_parallel: fixed_group, + ) + captured = {} + + def fake_dynamic_metrics(*, parallel_group): + captured["parallel_group"] = parallel_group + return {"mtp_1_loss": 0.25} + + monkeypatch.setattr( + "nemo_rl.models.megatron.dynamic_cp.get_dynamic_mtp_metrics", + fake_dynamic_metrics, + ) + monkeypatch.setattr(mpw, "broadcast_loss_metrics_from_last_stage", lambda d: d) + cfg = { + "megatron_cfg": { + "dynamic_context_parallel": { + "enabled": True, + "tokens_per_rank": 1024, + } + } + } + + metrics: dict = {} + mpw.MegatronPolicyWorkerImpl._collect_mtp_metrics( + _fake_worker(mtp_num_layers=1, cfg=cfg), + metrics, + total_num_microbatches=99, + mtp_grad_norm=None, + ) + + assert captured["parallel_group"] is fixed_group + assert metrics["mtp_metrics"]["mtp_1_loss"] == pytest.approx(0.25) + + @pytest.mark.mcore def test_collect_mtp_metrics_noop_when_mtp_disabled(monkeypatch): """MTP disabled (mtp_num_layers=0) -> nothing added and get_mtp_metrics is not called.""" From b1c3d34da98c7ce2088ba03b03092ab9408a21a2 Mon Sep 17 00:00:00 2001 From: humairafirdowse18 Date: Tue, 22 Sep 2026 13:17:37 -0700 Subject: [PATCH 6/7] prerun validations --- .gitignore | 15 + .../megatron-dynamic-context-parallel.md | 241 ++----- ...-30ba3b-8n4g-megatron-dynamiccp-quick.yaml | 49 -- ...g-async-1off-megatron-dynamiccp-quick.yaml | 54 -- ...-async-1off-megatron-dynamiccp-10step.yaml | 30 - ...g-async-1off-megatron-dynamiccp-quick.yaml | 52 -- ...g-async-1off-megatron-staticcp-10step.yaml | 13 - ...-async-1off-megatron-dynamiccp-10step.yaml | 55 -- ...g-async-1off-megatron-staticcp-10step.yaml | 11 - ...n3-32b-4n4g-megatron-dynamiccp-10step.yaml | 10 - ...3-32b-4n4g-megatron-dynamiccp-profile.yaml | 10 - ...en3-32b-4n4g-megatron-dynamiccp-quick.yaml | 42 -- ...en3-32b-4n4g-megatron-staticcp-10step.yaml | 9 - .../distributed/dynamic_context_parallel.py | 39 +- nemo_rl/models/megatron/common.py | 38 +- nemo_rl/models/megatron/dynamic_cp.py | 206 +++--- nemo_rl/models/megatron/hybridep.py | 4 + nemo_rl/models/policy/dynamic_cp.py | 149 +++-- nemo_rl/models/policy/lm_policy.py | 107 ++- nemo_rl/models/policy/teacher_worker_group.py | 4 + .../policy/workers/megatron_policy_worker.py | 15 +- .../L1_Functional_Tests_Megatron_1.sh | 1 + tests/functional/dynamic_cp.sh | 21 + .../test_dynamic_context_parallel.py | 20 + .../distributed/test_dynamic_cp_dispatch.py | 40 ++ .../models/megatron/test_dynamic_cp_moe.py | 32 +- .../megatron/test_dynamic_cp_scaling.py | 10 +- .../models/megatron/test_hybridep_data.py | 1 + .../unit/models/megatron/test_moe_metrics.py | 27 +- .../models/policy/test_megatron_worker.py | 20 - .../models/policy/test_policy_validation.py | 20 +- .../policy/test_teacher_worker_group.py | 15 +- .../test_dynamic_cp_comparison_recipes.py | 13 + tests/unit/test_dynamic_cp_moe_recipes.py | 10 +- tools/analyze_dynamic_cp_comparison.py | 629 ++++++++++++++++++ tools/launch_dynamic_cp_comparison.sh | 219 ++++++ 36 files changed, 1459 insertions(+), 772 deletions(-) delete mode 100644 examples/configs/recipes/llm/performance/grpo-nemotron3-nano-30ba3b-8n4g-megatron-dynamiccp-quick.yaml delete mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-megatron-dynamiccp-quick.yaml delete mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-10step.yaml delete mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-quick.yaml delete mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-staticcp-10step.yaml delete mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-async-1off-megatron-dynamiccp-10step.yaml delete mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-async-1off-megatron-staticcp-10step.yaml delete mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-10step.yaml delete mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-profile.yaml delete mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-quick.yaml delete mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-staticcp-10step.yaml create mode 100644 tests/functional/dynamic_cp.sh create mode 100755 tools/analyze_dynamic_cp_comparison.py create mode 100755 tools/launch_dynamic_cp_comparison.sh diff --git a/.gitignore b/.gitignore index 0a7ef3abb9e..1ddc37ae6ac 100644 --- a/.gitignore +++ b/.gitignore @@ -62,3 +62,18 @@ code_snapshots*/ # Named rather than ignoring the directory shape: tests/test_suites/llm/ # performance/ sits at that depth and is tracked. /tests/test_suites/**/metrics.json + +# Local dynamic/static context-parallel benchmark recipes (not in upstream main) +/examples/configs/recipes/llm/performance/grpo-nemotron3-nano-30ba3b-8n4g-megatron-dynamiccp-quick.yaml +/examples/configs/recipes/llm/performance/grpo-nemotron3-nano-30ba3b-8n4g-megatron-staticcp-quick.yaml +/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-megatron-dynamiccp-quick.yaml +/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-megatron-staticcp-quick.yaml +/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-10step.yaml +/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-quick.yaml +/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-staticcp-10step.yaml +/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-async-1off-megatron-dynamiccp-10step.yaml +/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-async-1off-megatron-staticcp-10step.yaml +/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-10step.yaml +/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-profile.yaml +/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-quick.yaml +/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-staticcp-10step.yaml diff --git a/docs/design-docs/megatron-dynamic-context-parallel.md b/docs/design-docs/megatron-dynamic-context-parallel.md index d3f7d802a75..57240c0cf83 100644 --- a/docs/design-docs/megatron-dynamic-context-parallel.md +++ b/docs/design-docs/megatron-dynamic-context-parallel.md @@ -44,22 +44,30 @@ CP groups and smaller model calls even when memory does not require them. The scheduler can still increase CP for work in a partially occupied phase so idle lanes contribute, matching the balanced hybrid-CP behavior. -This path currently supports PP=1 and the standard Ray policy data path. -TransferQueue, split execution, model-owned multimodal packing, atomic preference -pairs, MTP, draft training, fused linear logprobs, and training CUDA graphs are -not supported. Dynamic batching and HybridEP input prepadding are also disabled -because the driver plan must remain the sole owner of packing and microbatch -boundaries. Omit or disable the configuration for existing static behavior. - -For MoE, the minimum active size is raised until `active_CP * TP >= EP`. Workers +This path currently supports PP=1 and the standard Ray policy data path. Atomic +preference groups, MTP, HybridEP input prepadding, `global_aux_loss`, and expert +tensor parallelism are supported. Hybrid MTP requires matching attention and +linear CP layouts so the same packed layout can be used safely. HybridEP keeps +logical sequence boundaries for MTP while marking its enlarged physical buffer +as padded for Transformer Engine attention. Non-colocated distillation teachers +keep their own static CP topology; Dynamic CP remains enabled only on the student. + +TransferQueue/split execution, model-owned multimodal packing, draft training, +fused linear logprobs, dynamic batching, training CUDA graphs, MLA, and chunkwise +linear CP are not supported. Omit or disable the configuration for existing +static behavior. + +For MoE, the minimum active size is raised until +`active_CP * TP >= ETP * EP`. Workers also check that each actual expert communication group is contained within its task's ranks. This keeps every EP collective inside one active-CP task and prevents experts from communicating across independently scheduled CP blocks. Router auxiliary losses and per-layer MoE metrics stay attached to the real microbatch that produced them; placeholder tasks carry zero valid tokens and do not affect -loss normalization. `global_aux_loss`, expert tensor parallelism, quantile -balancing, and overlapped MoE microbatch execution are rejected because their -collective or scheduling domains cannot follow the active task safely. +loss normalization. `global_aux_loss` uses aligned placeholder rounds so every +rank enters its full-domain collective in the same order. Quantile balancing and +overlapped MoE microbatch execution remain rejected because their mask or +scheduling requirements cannot follow the active task safely. ## Dispatch and execution @@ -77,8 +85,9 @@ group placement, not another data-sharding axis. Before finalizing an initial placement, the driver repeatedly expands the smallest real task to the next CP power of two while unused lanes remain. Its padding factor and padded-token count are recalculated at the larger CP size. This mirrors MCore's -`fill_empty_gpus` policy and turns idle lanes into useful attention work. If an -explicit `max_size` prevents another expansion, the remaining lanes execute a +`fill_empty_gpus` policy and turns idle lanes into useful attention work. If the +larger alignment would exceed the task's token budget, or an explicit `max_size` +prevents expansion, the task keeps its valid size and remaining lanes execute a small zero-mask placeholder. Adjacent placements with the same lane partition are merged into one @@ -170,197 +179,31 @@ MCore pin must pass that contract test and distributed parity before adoption. For MoE load balancing, MCore's per-token path assumes `local_valid_tokens * TP_CP_size` equals the active task's token count. Dynamic -packing can leave different valid-token counts on participating shards, so the -worker sums the exact valid count over the active TP×CP group and applies a -per-shard correction through `moe_grad_scale_func`. MCore's z-loss coefficient -and attachment factors already cancel to form each rank's valid-token sum; its -temporary coefficient is inversely adjusted so the shared autograd scaler does -not change that gradient. For reporting, ordinary aux +packing can leave different valid-token counts on participating shards. A router +pre-hook sums the exact valid count over the active TP×CP group, and the temporary +router attachment uses that task count directly. The worker's ordinary +`1/global_valid_tokens` gradient scale therefore remains sufficient and does not +need a second per-shard correction. For reporting, ordinary aux metrics are divided by unique real tasks, while z-loss reproduces MCore's average over every real TP×CP rank participation. -## GB200 five-step smoke test - -From the RL checkout on the Slurm login node: - -```bash -DRY_RUN=0 bash perf_runs/run_gb200_dynamic_cp.sh -``` - -The launcher requests four nodes with four GPUs each, partition `batch`, QoS -`short`, and a two-hour limit. It uses the container and HF cache defaults in the -launcher, which can be overridden through its environment variables. Omitting -`DRY_RUN=0` prints the submission command without submitting. - -The container preflight checks the Bridge pin against the RL checkout and runs -CPU dispatch tests. A four-GPU test then compares the actual packing, CP -logprob collectives, loss, gradients, and an SGD update against an unsharded -reference for active sizes 4, 2, and 1 with base CP=1 and base CP=2. Another -four-GPU test exercises fused RoPE and TE attention in a small transformer, -including TP=1/2 and base CP=1/2. It also checks actual worker score serialization -and reassembled sample logprobs against a CP=1 reference. Then -`grpo-qwen3-32b-4n4g-megatron-dynamiccp-quick.yaml` runs five GRPO steps with -TP=2, PP=1, base CP=1, and active CP up to 8. Both policy and reference logprob -passes are enabled. Checkpoints and validation are disabled for the smoke test. - -Driver output is in `-logs/ray-driver.log`; TensorBoard metrics are in -`logs/dynamic-cp-`. The final metrics check requires steps 1–5, -finite loss and gradient norm, positive valid-token counts, training/scoring -importance ratios within 0.01 of one, and generation KL below 0.1, and writes -`smoke_result.json` in that log directory. A completed smoke run verifies -execution and finite training metrics, not convergence or a speedup over static CP. - -### Dynamic-CP MoE quick smoke tests - -`perf_runs/run_gb200_dynamic_cp_moe.sh` runs the distributed loss and attention -preflights followed by a two-to-five-step real-model GRPO smoke test. It defaults -to Qwen3-30B-A3B, five steps, four nodes total (two generation nodes and two -policy nodes inherited from the async 1-off recipe), and a one-hour QoS limit: - -```bash -DYNAMIC_CP_MOE_STEPS=5 DRY_RUN=0 \ - bash perf_runs/run_gb200_dynamic_cp_moe.sh -``` - -Select the other model cases with `DYNAMIC_CP_MOE_MODEL=qwen235b` or -`DYNAMIC_CP_MOE_MODEL=nemotron3-nano`. Qwen3-30B-A3B uses TP1/EP8 and therefore -runs its policy task at CP8. Qwen3-235B-A22B uses TP8/EP16, so its minimum active -CP is two. Nemotron-3-Nano-30B-A3B uses TP2/EP8, so its minimum active CP is four. -Larger active sizes remain available for longer generated sequences. - -The Qwen3-30B-A3B recipe exercises `aux_loss`; Qwen3-235B-A22B exercises -`seq_aux_loss`; and the Nemotron recipe exercises its inherited router setup. -The post-run check requires the expected active CP size, finite training metrics, -the requested number of steps, and the configured MoE metric when applicable. - -The recipes log to both TensorBoard and W&B. The launcher defaults -`WANDB_MODE=online`, uses `/home/humairafirdo/hf_home`, and prints both settings -before submission. Set `WANDB_MODE=offline` explicitly when online logging is not -wanted. - -### Ten-step Nsight profile - -`perf_runs/run_gb200_dynamic_cp_profile.sh` runs the same dense Qwen3-32B setup -for ten steps with TensorBoard and W&B enabled. It profiles only Megatron policy -workers because that is where dynamic CP executes. By default, Nsight captures -all ten steps with `PROFILE_STEP_RANGE=1:11`. Use -`PROFILE_STEP_RANGE=3:6` for a smaller steady-state-only report. The launcher -is a dry run unless `DRY_RUN=0` is explicitly supplied. - -Each runtime packed task has an NVTX label such as -`dynamic_cp/group_2/task_1/cp_4/lane_2/data`. Within a group, different lanes -may have different maximum task indices. The post-run check requires ten -dynamic training plans, at least two active CP sizes, at least one group with -multiple sequential packed tasks, valid training metrics, and at least one -completed policy `.nsys-rep` file on the head node. It reports how many groups -had uneven per-lane task counts and also requires every training -dispatch to report `schedule=reused`. `ray.sub` copies reports from all nodes -into `-logs/ray/**/nsight/`. - -### Ten-step dynamic/static comparison - -`perf_runs/run_gb200_cp_comparison.sh` runs matched ten-step jobs with the same -Qwen3-30B-A3B model, batch, TP4/EP4/PP1 policy topology, generation setup, -container, and W&B project. EP4 is only a sharding change; it does not remove -experts or change model weights. On the two policy nodes, TP4 creates two lanes -and `CP1 * TP4 = EP4`, so a complete expert group fits inside CP1. The dynamic -run can execute two CP1 tasks or one CP2 task, while the capacity-matched static -run stays at CP2. - -The default workload has an 8192-token ceiling, 4096 tokens per rank, and a -global batch of 512 formed from 16 prompts times 32 generations. The launcher -uses partition `batch`, inherits the account's default QoS, and requests four -hours by default (`TIME_LIMIT` overrides it). Use different `CP_RUN_NAME` values so -the W&B runs and local logs remain distinct. - -The launcher also accepts `CP_NUM_STEPS`, `CP_TRAIN_GLOBAL_BATCH_SIZE`, -`CP_NUM_PROMPTS_PER_STEP`, `CP_NUM_GENERATIONS_PER_PROMPT`, -`CP_MAX_TOTAL_SEQUENCE_LENGTH`, `CP_TOKENS_PER_RANK`, `CP_MAX_SIZE`, and -`STATIC_CP_SIZE`. Prompt count times generations must equal the global batch. -A static CP1 run remains useful as an unconstrained throughput and CP1 sanity -reference when it fits in memory, but it is not the capacity-matched baseline -for 8192 tokens at the 4096-token budget. - -This is a match to the configured memory budget, not proof that CP is necessary -on GB200. Measure peak memory and test static CP1 before concluding that 8192 -tokens require CP2. If CP1 fits and performs better, use that as the practical -baseline; increase `CP_TOKENS_PER_RANK` to the measured training-safe budget. -Do not reduce the budget just to make the scheduler report more CP sizes. - -The correctness smoke keeps the original TP1/EP8 topology and is therefore -forced to CP8. The performance pair uses TP4/EP4 specifically to expose an -adaptive CP1/CP2 choice on the same eight policy GPUs. Both sides of the pair -use TP4/EP4, so the measured difference is dynamic versus fixed CP rather than -a model or expert-layout difference between the two runs. - -After each run, `perf_runs/analyze_cp_sequence_lengths.py` writes -`sequence_length_distribution.json` beside the driver log. It reports length -percentiles, how many samples stopped exactly at the configured ceiling, and -the CP size each sample required before optional idle-lane expansion. - -Dynamic schedule logs also report `tasks_by_cp` and `packing_utilization`. -Required CP for an individual sample can differ from its scheduled CP when it -fills an existing larger task or when spare lanes help process a task. Performance -runs do not require a fixed mixture of sizes. For an explicit coverage test, set -`CP_REQUIRE_SIZES="1 2"` (or `"1 2 4"`). The correctness checks still require all -steps, finite metrics, score/train agreement, and schedule reuse. - -Both comparison modes run four-GPU loss and attention parity checks before -training. These tests explicitly retain CP1 coverage after cross-size packing, -including a model initialized at static CP2. Set `CP_GPU_PREFLIGHT=0` only when -reusing validation of the same code/container. Preflight time is outside the -reported training-step timings. - -### Qwen30B packing regression: jobs 7202770 and 7203284 - -Both ten-step jobs passed their smoke checks and processed approximately 26.6M -tokens. Excluding step one, the dynamic job averaged 525.65 seconds per step; -static CP2 averaged 483.72 seconds, so dynamic took 8.67% longer. Training took -390.62 versus 357.51 seconds and policy/reference scoring took 130.03 versus -121.26 seconds. These are independent async rollouts, not identical token batches. - -The old scheduler packed each required CP size separately. At 4096 tokens/rank, -`[6000, 2000]` became two calls even though both samples fit in one 8192-token CP2 -pack. The cross-size packing fix admits the 2000-token sequence into that -already-required call, checks CP2 alignment, and still leaves independent CP1 -tasks when there is remaining short work. EP containment and loss scaling are -unchanged; the worker derives its microbatch count from the resulting plan. - -CPU replay of the ten saved dynamic batches changes the sum of local calls per -lane from 3630 to 3330; static MFFD packing of those same lengths needs 3322. -That is an 8.26% reduction in model calls, not a measured GPU speedup. Rerun the -pair to measure actual performance. A near-zero gain remains possible because -most work in this workload still executes at CP2 after efficient packing. - -Reproduce the recorded timings and replay the current planner: - -```bash -uv run --no-sync python perf_runs/analyze_cp_comparison.py \ - logs/qwen30-gbs512-seq8192-dyncp-20260916T220850Z \ - logs/qwen30-gbs512-seq8192-staticcp2-20260916T220850Z \ - --lanes 2 --tp 4 --tokens-per-rank 4096 -``` - -### Dense Qwen3-32B comparison +## Validation -Set `CP_MODEL=qwen32b` to select the new dense recipes. They use four nodes total: -two policy nodes (TP2/PP1/EP1, four CP scheduling lanes) and two generation nodes -(vLLM TP2). Defaults are 16384 total tokens, 4096 tokens/rank, GBS512, ten steps, -activation checkpointing, W&B online, dynamic CP1/2/4 versus static CP4. Both -modes inherit the same workload, optimizer, precision, and generation settings. +CPU planner and dispatch behavior is covered by the unit tests under +`tests/unit/distributed/`. Four-GPU correctness coverage is available through: ```bash -PAIR_TAG=$(date -u +%Y%m%dT%H%M%SZ) -CP_MODEL=qwen32b CP_MODE=dynamic TIME_LIMIT=06:00:00 \ - CP_RUN_NAME="qwen32-gbs512-seq16384-dynamic-${PAIR_TAG}" DRY_RUN=0 \ - bash perf_runs/run_gb200_cp_comparison.sh -CP_MODEL=qwen32b CP_MODE=static STATIC_CP_SIZE=4 TIME_LIMIT=06:00:00 \ - CP_RUN_NAME="qwen32-gbs512-seq16384-static4-${PAIR_TAG}" DRY_RUN=0 \ - bash perf_runs/run_gb200_cp_comparison.sh +bash tests/functional/dynamic_cp.sh ``` -For a smaller first validation set `CP_NUM_STEPS=2 CP_NUM_PROMPTS_PER_STEP=2 -CP_TRAIN_GLOBAL_BATCH_SIZE=64 QOS=short TIME_LIMIT=01:00:00` and use a distinct -run name. Ten steps are an initial performance sample, not -convergence validation. Report training/scoring separately from total step time -and repeat close results before claiming a speedup. +The functional test compares loss, gradients, an optimizer update, output +reassembly, fused RoPE, and Transformer Engine attention across active CP sizes +1, 2, and 4, including TP and initialized-static-CP variants. + +End-to-end performance should be measured with matched local recipes. Keep the +model, generated token batch, parallel topology, optimizer, and precision fixed; +compare dynamic CP with the smallest static CP that fits the same workload. +Record score, reference-score, refit, training, and total step time separately, +plus peak memory, MFU, task counts by CP size, packing utilization, schedule +reuse, and the sequence-length distribution. Ten or twenty steps are useful for +a smoke/performance sample, but are not enough to claim training convergence. diff --git a/examples/configs/recipes/llm/performance/grpo-nemotron3-nano-30ba3b-8n4g-megatron-dynamiccp-quick.yaml b/examples/configs/recipes/llm/performance/grpo-nemotron3-nano-30ba3b-8n4g-megatron-dynamiccp-quick.yaml deleted file mode 100644 index 8d13db7e101..00000000000 --- a/examples/configs/recipes/llm/performance/grpo-nemotron3-nano-30ba3b-8n4g-megatron-dynamiccp-quick.yaml +++ /dev/null @@ -1,49 +0,0 @@ -defaults: ../grpo-dapomath17k-nanov3-30BA3B-8n4g-megatron-trtllm.yaml -grpo: - num_prompts_per_step: 8 - num_generations_per_prompt: 8 - max_num_steps: 5 - val_period: 1000 - val_at_start: false - val_at_end: false -checkpointing: - enabled: false - checkpoint_dir: results/grpo-nemotron3-nano-30ba3b-8n4g-dynamiccp-quick -policy: - train_global_batch_size: 64 - train_micro_batch_size: 1 - logprob_batch_size: 1 - max_total_sequence_length: 8192 - make_sequence_length_divisible_by: 4 - megatron_cfg: - tensor_model_parallel_size: 2 - pipeline_model_parallel_size: 1 - num_layers_in_first_pipeline_stage: null - num_layers_in_last_pipeline_stage: null - context_parallel_size: 1 - expert_tensor_parallel_size: 1 - expert_model_parallel_size: 8 - sequence_parallel: true - moe_hybridep_prepad_packed_inputs: false - dynamic_context_parallel: - enabled: true - min_size: 1 - max_size: 16 - tokens_per_rank: 4096 - fp8_cfg: - enabled: false - sequence_packing: - enabled: true - dynamic_batching: - enabled: false - generation: - max_new_tokens: 4096 -logger: - log_dir: logs/grpo-nemotron3-nano-30ba3b-8n4g-dynamiccp-quick - wandb_enabled: true - tensorboard_enabled: true - wandb: - project: nemo-rl - name: grpo-nemotron3-nano-30ba3b-8n4g-dynamiccp-quick -data_plane: - enabled: false diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-megatron-dynamiccp-quick.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-megatron-dynamiccp-quick.yaml deleted file mode 100644 index c3e7e7b8d69..00000000000 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-megatron-dynamiccp-quick.yaml +++ /dev/null @@ -1,54 +0,0 @@ -defaults: ./grpo-qwen3-235b-32n4g-async-1off.yaml -grpo: - num_prompts_per_step: 8 - num_generations_per_prompt: 8 - max_num_steps: 5 - val_period: 1000 - val_at_start: false - val_at_end: false -checkpointing: - enabled: false - checkpoint_dir: results/grpo-qwen3-235b-32n4g-async-dynamiccp-quick -policy: - hf_config_overrides: - # Cover MCore's sequence-level router loss on a real MoE model. - router_aux_loss_coef: 0.001 - train_global_batch_size: 64 - train_micro_batch_size: 1 - logprob_batch_size: 1 - max_total_sequence_length: 8192 - megatron_cfg: - # PP must be one because every dynamic lane executes its own task list. - # TP8 keeps the dense portion sharded while CP*TP >= EP from CP2 onward. - tensor_model_parallel_size: 8 - pipeline_model_parallel_size: 1 - num_layers_in_first_pipeline_stage: null - num_layers_in_last_pipeline_stage: null - context_parallel_size: 1 - expert_tensor_parallel_size: 1 - expert_model_parallel_size: 16 - sequence_parallel: true - freeze_moe_router: false - moe_router_load_balancing_type: seq_aux_loss - moe_per_layer_logging: true - moe_hybridep_prepad_packed_inputs: false - dynamic_context_parallel: - enabled: true - min_size: 1 - max_size: 8 - tokens_per_rank: 4096 - fp8_cfg: - enabled: false - sequence_packing: - enabled: true - dynamic_batching: - enabled: false -logger: - log_dir: logs/grpo-qwen3-235b-32n4g-async-dynamiccp-quick - wandb_enabled: true - tensorboard_enabled: true - wandb: - project: nemo-rl - name: grpo-qwen3-235b-32n4g-async-dynamiccp-quick -data_plane: - enabled: false diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-10step.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-10step.yaml deleted file mode 100644 index 7489fd0b98e..00000000000 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-10step.yaml +++ /dev/null @@ -1,30 +0,0 @@ -defaults: ./grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-quick.yaml -grpo: - # Keep the original Qwen30B performance recipe's 32 generations per prompt. - num_prompts_per_step: 16 - num_generations_per_prompt: 32 - max_num_steps: 10 -policy: - train_global_batch_size: 512 - max_total_sequence_length: 8192 - make_sequence_length_divisible_by: 4 - megatron_cfg: - # The two policy nodes provide eight ranks. TP4 creates two scheduling - # lanes; EP4 fits completely inside one CP1*TP4 task. Dynamic CP can - # therefore execute two CP1 tasks or one CP2 task without crossing an EP - # collective between independently scheduled tasks. - tensor_model_parallel_size: 4 - expert_model_parallel_size: 4 - sequence_parallel: true - dynamic_context_parallel: - min_size: 1 - max_size: 2 - tokens_per_rank: 4096 -logger: - log_dir: logs/grpo-qwen3-30ba3b-4n4g-async-dynamiccp-10step - wandb_enabled: true - tensorboard_enabled: true - wandb: - project: nemo-rl-cp-comparison - name: qwen3-30ba3b-4n4g-async-dynamiccp-10step - diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-quick.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-quick.yaml deleted file mode 100644 index fcfd57015b0..00000000000 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-quick.yaml +++ /dev/null @@ -1,52 +0,0 @@ -defaults: ./grpo-qwen3-30ba3b-4n4g-async-1off.yaml -grpo: - num_prompts_per_step: 8 - num_generations_per_prompt: 8 - max_num_steps: 5 - val_period: 1000 - val_at_start: false - val_at_end: false -checkpointing: - enabled: false - checkpoint_dir: results/grpo-qwen3-30ba3b-4n4g-async-dynamiccp-quick -policy: - hf_config_overrides: - # Exercise the normal per-microbatch router auxiliary loss in this smoke run. - router_aux_loss_coef: 0.001 - train_global_batch_size: 64 - train_micro_batch_size: 1 - logprob_batch_size: 1 - max_total_sequence_length: 4096 - megatron_cfg: - tensor_model_parallel_size: 1 - pipeline_model_parallel_size: 1 - context_parallel_size: 1 - expert_tensor_parallel_size: 1 - expert_model_parallel_size: 8 - sequence_parallel: false - freeze_moe_router: false - moe_router_load_balancing_type: aux_loss - moe_per_layer_logging: true - moe_hybridep_prepad_packed_inputs: false - dynamic_context_parallel: - enabled: true - min_size: 1 - # The inherited async 1-off layout leaves two 4-GPU nodes for policy. - # TP1/EP8 therefore requires and exactly fills one active CP8 task. - max_size: 8 - tokens_per_rank: 4096 - fp8_cfg: - enabled: false - sequence_packing: - enabled: true - dynamic_batching: - enabled: false -logger: - log_dir: logs/grpo-qwen3-30ba3b-4n4g-async-dynamiccp-quick - wandb_enabled: true - tensorboard_enabled: true - wandb: - project: nemo-rl - name: grpo-qwen3-30ba3b-4n4g-async-dynamiccp-quick -data_plane: - enabled: false diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-staticcp-10step.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-staticcp-10step.yaml deleted file mode 100644 index 903b6871c71..00000000000 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-staticcp-10step.yaml +++ /dev/null @@ -1,13 +0,0 @@ -defaults: ./grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-10step.yaml -policy: - # CP2 is capacity-matched to an 8192-token maximum at 4096 tokens/rank. - # With TP4 and sequence parallelism, packed sequences align to CP*2*TP=16. - make_sequence_length_divisible_by: 16 - megatron_cfg: - context_parallel_size: 2 - dynamic_context_parallel: - enabled: false -logger: - log_dir: logs/grpo-qwen3-30ba3b-4n4g-async-staticcp-10step - wandb: - name: qwen3-30ba3b-4n4g-async-staticcp-10step diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-async-1off-megatron-dynamiccp-10step.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-async-1off-megatron-dynamiccp-10step.yaml deleted file mode 100644 index d3114632266..00000000000 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-async-1off-megatron-dynamiccp-10step.yaml +++ /dev/null @@ -1,55 +0,0 @@ -defaults: ./grpo-qwen3-32b-8n4g-async-1off.yaml -grpo: - num_prompts_per_step: 16 - num_generations_per_prompt: 32 - max_num_steps: 10 - val_period: 1000 - val_at_start: false - val_at_end: false -checkpointing: - enabled: false -policy: - train_global_batch_size: 512 - train_micro_batch_size: 1 - logprob_batch_size: 1 - max_total_sequence_length: 16384 - make_sequence_length_divisible_by: 2 - megatron_cfg: - tensor_model_parallel_size: 2 - pipeline_model_parallel_size: 1 - context_parallel_size: 1 - expert_model_parallel_size: 1 - sequence_parallel: true - activation_checkpointing: true - dynamic_context_parallel: - enabled: true - min_size: 1 - max_size: 4 - tokens_per_rank: 4096 - fp8_cfg: - enabled: false - sequence_packing: - enabled: true - dynamic_batching: - enabled: false - generation: - colocated: - enabled: false - resources: - num_nodes: 2 - gpus_per_node: 4 - vllm_cfg: - tensor_parallel_size: 2 -logger: - log_dir: logs/grpo-qwen3-32b-4n4g-async-dynamiccp-10step - wandb_enabled: true - tensorboard_enabled: true - wandb: - project: nemo-rl-cp-comparison - name: qwen3-32b-4n4g-async-dynamiccp-10step -cluster: - num_nodes: 4 - gpus_per_node: 4 - segment_size: 2 -data_plane: - enabled: false diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-async-1off-megatron-staticcp-10step.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-async-1off-megatron-staticcp-10step.yaml deleted file mode 100644 index 5b73fc1cac6..00000000000 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-async-1off-megatron-staticcp-10step.yaml +++ /dev/null @@ -1,11 +0,0 @@ -defaults: ./grpo-qwen3-32b-4n4g-async-1off-megatron-dynamiccp-10step.yaml -policy: - make_sequence_length_divisible_by: 16 - megatron_cfg: - context_parallel_size: 4 - dynamic_context_parallel: - enabled: false -logger: - log_dir: logs/grpo-qwen3-32b-4n4g-async-staticcp-10step - wandb: - name: qwen3-32b-4n4g-async-staticcp-10step diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-10step.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-10step.yaml deleted file mode 100644 index 641bab57116..00000000000 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-10step.yaml +++ /dev/null @@ -1,10 +0,0 @@ -defaults: ./grpo-qwen3-32b-4n4g-megatron-dynamiccp-quick.yaml -grpo: - max_num_steps: 10 -logger: - log_dir: logs/grpo-qwen3-32b-4n4g-megatron-dynamiccp-10step - wandb_enabled: true - tensorboard_enabled: true - wandb: - project: nemo-rl-cp-comparison - name: qwen3-32b-4n4g-dynamic-cp-10step diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-profile.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-profile.yaml deleted file mode 100644 index f7f6f0d8387..00000000000 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-profile.yaml +++ /dev/null @@ -1,10 +0,0 @@ -defaults: ./grpo-qwen3-32b-4n4g-megatron-dynamiccp-quick.yaml -grpo: - max_num_steps: 10 -logger: - log_dir: logs/grpo-qwen3-32b-4n4g-megatron-dynamiccp-profile - wandb_enabled: true - tensorboard_enabled: true - wandb: - project: nemo-rl-dynamic-cp - name: qwen3-32b-4n4g-dynamic-cp-profile diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-quick.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-quick.yaml deleted file mode 100644 index 009023fdd7e..00000000000 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-dynamiccp-quick.yaml +++ /dev/null @@ -1,42 +0,0 @@ -defaults: ./grpo-qwen3-32b-4n4g.yaml -grpo: - num_prompts_per_step: 8 - num_generations_per_prompt: 8 - max_num_steps: 5 - val_period: 1000 - val_at_start: false - val_at_end: false -loss_fn: - force_on_policy_ratio: false -checkpointing: - enabled: false -policy: - train_global_batch_size: 64 - train_micro_batch_size: 1 - logprob_batch_size: 1 - max_total_sequence_length: 4096 - megatron_cfg: - tensor_model_parallel_size: 2 - pipeline_model_parallel_size: 1 - context_parallel_size: 1 - expert_model_parallel_size: 1 - dynamic_context_parallel: - enabled: true - min_size: 1 - max_size: 8 - # Match the normal packed-microbatch budget. Smaller values force long - # sequences into CP even when they fit on one GB200 rank, fragmenting one - # step into many poorly utilized model calls. - tokens_per_rank: 4096 - fp8_cfg: - enabled: false - sequence_packing: - enabled: true - dynamic_batching: - enabled: false -logger: - log_dir: logs/grpo-qwen3-32b-4n4g-megatron-dynamiccp-quick - wandb_enabled: false - tensorboard_enabled: true -data_plane: - enabled: false diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-staticcp-10step.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-staticcp-10step.yaml deleted file mode 100644 index cc1494b191b..00000000000 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-megatron-staticcp-10step.yaml +++ /dev/null @@ -1,9 +0,0 @@ -defaults: ./grpo-qwen3-32b-4n4g-megatron-dynamiccp-10step.yaml -policy: - megatron_cfg: - dynamic_context_parallel: - enabled: false -logger: - log_dir: logs/grpo-qwen3-32b-4n4g-megatron-staticcp-10step - wandb: - name: qwen3-32b-4n4g-static-cp-10step diff --git a/nemo_rl/distributed/dynamic_context_parallel.py b/nemo_rl/distributed/dynamic_context_parallel.py index 04257848177..26564063c67 100644 --- a/nemo_rl/distributed/dynamic_context_parallel.py +++ b/nemo_rl/distributed/dynamic_context_parallel.py @@ -80,7 +80,7 @@ class CPRankPlan: steps: tuple[CPRankStep, ...] -def _padding_for_cp( +def padding_for_cp( cp_size: int, *, sequence_parallel_size: int, @@ -103,7 +103,7 @@ def _resize_assignment( user_pad_multiple: int, token_alignment: int, ) -> CPAssignment: - factor = _padding_for_cp( + factor = padding_for_cp( cp_size, sequence_parallel_size=sequence_parallel_size, user_pad_multiple=user_pad_multiple, @@ -137,25 +137,36 @@ def _fill_idle_lanes( ) -> list[CPAssignment]: """Increase the smallest real CP groups until no legal expansion fits.""" idle_lanes = lanes - sum(task.cp_size for task in tasks) + blocked: set[int] = set() while idle_lanes: candidates = [ (task.cp_size, index) for index, task in enumerate(tasks) - if task.cp_size < max_size and task.cp_size <= idle_lanes + if index not in blocked + and task.cp_size < max_size + and task.cp_size <= idle_lanes ] if not candidates: break _, index = min(candidates) task = tasks[index] - tasks[index] = _resize_assignment( - task, - task.cp_size * 2, - lengths=lengths, - tokens_per_rank=tokens_per_rank, - sequence_parallel_size=sequence_parallel_size, - user_pad_multiple=user_pad_multiple, - token_alignment=token_alignment, - ) + try: + tasks[index] = _resize_assignment( + task, + task.cp_size * 2, + lengths=lengths, + tokens_per_rank=tokens_per_rank, + sequence_parallel_size=sequence_parallel_size, + user_pad_multiple=user_pad_multiple, + token_alignment=token_alignment, + ) + except ValueError: + # CP growth also increases the per-sequence padding multiple. A task + # can therefore overflow after expansion even though the aggregate + # budget grew. Idle filling is optional, so retain the valid task and + # try another candidate instead of failing the whole rollout. + blocked.add(index) + continue idle_lanes -= task.cp_size # Descending powers of two guarantee that each consecutive lane start is @@ -351,7 +362,7 @@ def plan_cp_phases( ): size = min_size while True: - multiple = _padding_for_cp( + multiple = padding_for_cp( size, sequence_parallel_size=sequence_parallel_size, user_pad_multiple=user_pad_multiple, @@ -420,7 +431,7 @@ def plan_cp_phases( ) cursor = sum(task.cp_size for task in phase) while cursor < lanes: - factor = _padding_for_cp( + factor = padding_for_cp( min_size, sequence_parallel_size=sequence_parallel_size, user_pad_multiple=user_pad_multiple, diff --git a/nemo_rl/models/megatron/common.py b/nemo_rl/models/megatron/common.py index 08d6dc45a8c..97056ec9059 100644 --- a/nemo_rl/models/megatron/common.py +++ b/nemo_rl/models/megatron/common.py @@ -216,15 +216,15 @@ def get_moe_metrics( per_layer_logging: If True, include per-layer values in the returned dict. num_layers: Total number of transformer layers. When provided together with a non-empty ``track_names``, the aux-loss tracker is pre-initialized on every - rank before the reduction (see Note). Defaults to None, which disables - pre-initialization. + rank before the reduction (see Note). Required for dynamic CP. mtp_num_layers: Extra layers contributed by Multi-Token Prediction, added to ``num_layers`` to size the pre-initialized tensor, matching the size the router uses when recording. Defaults to None (treated as 0). track_names: Aux-loss names to pre-initialize; must mirror what the router records for the configured ``moe_router_load_balancing_type``, so callers should derive it via ``get_aux_loss_track_names(model_config)``. Defaults to - None, which disables pre-initialization. + None, which disables pre-initialization. Dynamic CP requires an explicit + list (which may be empty) so every rank uses the same collective order. dynamic_parallel_group: Fixed TP*DP*CP group used only by dynamic CP. Dynamic lanes can execute different numbers of microbatches and the router's per-forward TP*CP group changes with each task, so its last @@ -252,6 +252,14 @@ def get_moe_metrics( record an aux loss this step (e.g. a stage with no MoE layer, or an MTP MoE layer that lives only on the last stage). """ + if dynamic_parallel_group is not None and ( + track_names is None or num_layers is None + ): + raise ValueError( + "Dynamic CP MoE metrics require explicit track_names and num_layers " + "so every rank reduces the same tensors in the same order" + ) + # Pre-initialize the aux-loss tracker so every PP rank has the same set of # named, equally-sized tensors BEFORE the collective all_reduce below. # @@ -284,18 +292,20 @@ def get_moe_metrics( # counts every real task once, irrespective of uneven lane task counts. # The caller's loss_scale is 1 / number_of_unique_real_tasks. mcore_tracker = get_moe_metrics_tracker() - dynamic_names = ( - track_names if track_names is not None else list(mcore_tracker.metrics) - ) + # The validation above makes this list rank-identical. Never derive + # collective order from lazily populated, rank-local tracker state. + assert track_names is not None + dynamic_names = track_names for name in dynamic_names: - entry = mcore_tracker.metrics.get(name) - if entry is not None: - # Padding-only lanes receive a pre-initialized z-loss entry but - # never call record(), so its avg_group remains None locally. - # z_loss is nevertheless an averaged metric on every lane. - if name == "z_loss" or getattr(entry, "avg_group", None) is not None: - dynamic_avg_names.add(name) - dist.all_reduce(entry.values, group=dynamic_parallel_group) + # ensure_initialized above guarantees this lookup on every rank; + # fail locally rather than conditionally skipping a collective. + entry = mcore_tracker.metrics[name] + # Padding-only lanes receive a pre-initialized z-loss entry but + # never call record(), so its avg_group remains None locally. + # z_loss is nevertheless an averaged metric on every lane. + if name == "z_loss": + dynamic_avg_names.add(name) + dist.all_reduce(entry.values, group=dynamic_parallel_group) tracker = get_moe_layer_wise_logging_tracker() if dynamic_names is not None: tracker = {name: tracker[name] for name in dynamic_names if name in tracker} diff --git a/nemo_rl/models/megatron/dynamic_cp.py b/nemo_rl/models/megatron/dynamic_cp.py index b0d2f035c53..6e69a9de6e1 100644 --- a/nemo_rl/models/megatron/dynamic_cp.py +++ b/nemo_rl/models/megatron/dynamic_cp.py @@ -24,8 +24,8 @@ _DYNAMIC_TP_CP_GROUPS: dict[int, Any] = {} _ROUTER_CONFIG_BASELINES: dict[int, tuple[Any, Any]] = {} -_DYNAMIC_MOE_GRAD_SCALE_CORRECTIONS: dict[int, float] = {} _DYNAMIC_MTP_METRICS: dict[str, torch.Tensor] = {} +_ACTIVE_BIND_TARGETS: dict[int, "_BindTargets"] = {} @dataclass(frozen=True) @@ -149,6 +149,52 @@ def _is_hybrid_stack(module: torch.nn.Module) -> bool: ) +@dataclass(frozen=True) +class _BindTargets: + """Stable module classifications reused by every task in one model call.""" + + modules: tuple[torch.nn.Module, ...] + pg_collections: tuple[torch.nn.Module, ...] + direct_groups: tuple[torch.nn.Module, ...] + hybrid_stacks: tuple[torch.nn.Module, ...] + routers: tuple[Router, ...] + mamba_mixers: tuple[torch.nn.Module, ...] + gated_delta_products: tuple[torch.nn.Module, ...] + gated_delta_nets: tuple[torch.nn.Module, ...] + + +def _classify_bind_targets(model: torch.nn.Module) -> _BindTargets: + modules = tuple(model.modules()) + return _BindTargets( + modules=modules, + pg_collections=tuple( + module + for module in modules + if getattr(module, "pg_collection", None) is not None + and hasattr(module.pg_collection, "cp") + ), + direct_groups=tuple( + module + for module in modules + if _uses_direct_cp_group(module) and hasattr(module, "cp_group") + ), + hybrid_stacks=tuple(module for module in modules if _is_hybrid_stack(module)), + routers=tuple(module for module in modules if isinstance(module, Router)), + mamba_mixers=tuple(module for module in modules if _is_mamba_mixer(module)), + gated_delta_products=tuple( + module for module in modules if _is_gated_delta_product(module) + ), + gated_delta_nets=tuple( + module for module in modules if _is_gated_delta_net(module) + ), + ) + + +def _bind_targets(model: torch.nn.Module) -> _BindTargets: + """Reuse the outer schedule's classification, with a safe direct-call fallback.""" + return _ACTIVE_BIND_TARGETS.get(id(model)) or _classify_bind_targets(model) + + def _rebuild_mamba_cp(module: torch.nn.Module, group: Any) -> None: """Rebuild Mamba's cached CP helper for the active microbatch size.""" cp = module.cp @@ -299,6 +345,9 @@ def get_dynamic_mtp_metrics( *, parallel_group: torch.distributed.ProcessGroup ) -> dict[str, float]: """Reduce token-weighted MTP metrics over the fixed DP*CP lane group.""" + # This branch is collective-safe: training gives every lane at least one + # real or placeholder task, and even a zero-token placeholder records all + # four tensors. Evaluation records MTP metrics on no lane, so all lanes exit. if "loss_sums" not in _DYNAMIC_MTP_METRICS: return {} try: @@ -503,7 +552,11 @@ def _dynamic_process_mtp_loss( @contextmanager def _patch_mtp_loss_for_dynamic_cp(enabled: bool) -> Iterator[None]: - """Temporarily route GPT/HybridModel MTP through the NeMo-side fix.""" + """Temporarily route GPT/HybridModel MTP through the NeMo-side fix. + + This module-level patch relies on the current worker contract: one training + model executes at a time and generation does not share the worker process. + """ if not enabled: yield return @@ -589,7 +642,9 @@ def _dynamic_attach_and_log_load_balancing_loss( @contextmanager -def _patch_hybrid_mtp_padding_masks(model: torch.nn.Module) -> Iterator[None]: +def _patch_hybrid_mtp_padding_masks( + model: torch.nn.Module, modules: tuple[torch.nn.Module, ...] | None = None +) -> Iterator[None]: """Carry correct router padding semantics through every dynamic MTP block. The pinned HybridModel accepts ``padding_mask`` and sends it through the @@ -609,7 +664,7 @@ def _patch_hybrid_mtp_padding_masks(model: torch.nn.Module) -> Iterator[None]: """ padding_masks: dict[int, torch.Tensor | None] = {} handles: list[Any] = [] - modules = list(model.modules()) + modules = modules or tuple(model.modules()) mtp_blocks: list[torch.nn.Module] = [] def capture_padding_mask( @@ -695,6 +750,8 @@ def prepare_router_padding_mask( return args, kwargs patched_classes: list[tuple[type[Any], bool, Any]] = [] + # This class-level patch has the same single-model worker assumption as the + # MTP patch above; the preservation context always restores it in finally. for router_class in { type(module) for module in modules if isinstance(module, Router) }: @@ -746,14 +803,12 @@ def preserve_attention_cp_groups(model: torch.nn.Module) -> Iterator[None]: Keep the active group through backward recomputation; restore it after the complete no-pipeline schedule. TP, DP and optimizer groups are unchanged. """ - modules = list(model.modules()) + targets = _classify_bind_targets(model) + modules = targets.modules saved_collections = [ - (module, module.pg_collection) - for module in modules - if getattr(module, "pg_collection", None) is not None - and hasattr(module.pg_collection, "cp") + (module, module.pg_collection) for module in targets.pg_collections ] - saved_mamba = [(module, module.cp) for module in modules if _is_mamba_mixer(module)] + saved_mamba = [(module, module.cp) for module in targets.mamba_mixers] saved_gdp = [ ( module, @@ -762,13 +817,11 @@ def preserve_attention_cp_groups(model: torch.nn.Module) -> Iterator[None]: module.nheads_local_cp, module.ngroups_local_cp, ) - for module in modules - if _is_gated_delta_product(module) + for module in targets.gated_delta_products ] saved_gdn = [ (module, module.cp_size, getattr(module, "feat_dim_split", None)) - for module in modules - if _is_gated_delta_net(module) + for module in targets.gated_delta_nets ] saved_direct_groups = [ ( @@ -777,8 +830,7 @@ def preserve_attention_cp_groups(model: torch.nn.Module) -> Iterator[None]: getattr(module, "tp_cp_group", None), hasattr(module, "tp_cp_group"), ) - for module in modules - if _uses_direct_cp_group(module) and hasattr(module, "cp_group") + for module in targets.direct_groups ] saved_hybrid_stacks = [ ( @@ -786,33 +838,42 @@ def preserve_attention_cp_groups(model: torch.nn.Module) -> Iterator[None]: module._cp_layout_manager, module._has_linear_layer_with_chunkwise_cp, ) - for module in modules - if _is_hybrid_stack(module) + for module in targets.hybrid_stacks ] saved_routers = [ - (module, module.cp_group, module.tp_cp_group) - for module in modules - if isinstance(module, Router) + (module, module.cp_group, module.tp_cp_group) for module in targets.routers ] router_configs: dict[int, Any] = {} - for module, collection in saved_collections: - module.pg_collection = copy(collection) - for module, _, _ in saved_routers: - config = module.config - router_configs[id(config)] = config - _ROUTER_CONFIG_BASELINES[id(config)] = ( - config.moe_aux_loss_coeff, - config.moe_z_loss_coeff, - ) - _DYNAMIC_MOE_GRAD_SCALE_CORRECTIONS[id(config)] = 1.0 mtp_enabled = any( bool(getattr(getattr(module, "config", None), "mtp_num_layers", 0)) for module in modules ) + model_id = id(model) + if model_id in _ACTIVE_BIND_TARGETS: + raise RuntimeError("Dynamic CP model binding is not re-entrant") try: + _ACTIVE_BIND_TARGETS[model_id] = targets + for module, collection in saved_collections: + module.pg_collection = copy(collection) + for module, _, _ in saved_routers: + config = module.config + config_id = id(config) + # Transformer layers commonly share one TransformerConfig. Record + # that shared object once, while still rejecting overlap with a + # different active model binding. + if config_id in router_configs: + continue + if config_id in _ROUTER_CONFIG_BASELINES: + raise RuntimeError("Dynamic CP router binding is not re-entrant") + baseline = ( + config.moe_aux_loss_coeff, + config.moe_z_loss_coeff, + ) + _ROUTER_CONFIG_BASELINES[config_id] = baseline + router_configs[config_id] = config with ( _patch_mtp_loss_for_dynamic_cp(mtp_enabled), - _patch_hybrid_mtp_padding_masks(model), + _patch_hybrid_mtp_padding_masks(model, modules), ): try: yield @@ -820,6 +881,7 @@ def preserve_attention_cp_groups(model: torch.nn.Module) -> Iterator[None]: _DYNAMIC_MTP_METRICS.clear() raise finally: + _ACTIVE_BIND_TARGETS.pop(model_id, None) for module, collection in saved_collections: module.pg_collection = collection for module, cp in saved_mamba: @@ -851,7 +913,6 @@ def preserve_attention_cp_groups(model: torch.nn.Module) -> Iterator[None]: module.tp_cp_group = tp_cp_group for config_id, config in router_configs.items(): aux_coeff, z_coeff = _ROUTER_CONFIG_BASELINES.pop(config_id) - _DYNAMIC_MOE_GRAD_SCALE_CORRECTIONS.pop(config_id, None) config.moe_aux_loss_coeff = aux_coeff config.moe_z_loss_coeff = z_coeff @@ -876,7 +937,6 @@ def _bind_router_config(router: Router, *, padding_only: bool) -> None: z_coeff = None router.config.moe_aux_loss_coeff = aux_coeff router.config.moe_z_loss_coeff = z_coeff - _DYNAMIC_MOE_GRAD_SCALE_CORRECTIONS[id(router.config)] = 1.0 def _has_positive_coefficient(value: Any) -> bool: @@ -896,14 +956,10 @@ def configure_dynamic_moe_loss_scaling( a caller forgot the packed padding mask. The worker's ordinary ``1/global_valid_tokens`` MoE scale is therefore sufficient. """ - routers = [module for module in model.modules() if isinstance(module, Router)] + routers = _bind_targets(model).routers if not routers: return - configs = {id(router.config): router.config for router in routers} - for config_id in configs: - _DYNAMIC_MOE_GRAD_SCALE_CORRECTIONS[config_id] = 1.0 - if not model.training or not torch.is_grad_enabled(): return if not any( @@ -918,11 +974,6 @@ def configure_dynamic_moe_loss_scaling( raise ValueError("Dynamic CP routers disagree on the active TP*CP group") -def dynamic_moe_grad_scale_correction(model_config: Any) -> float: - """Return the active task's aux-loss correction, or the static default.""" - return _DYNAMIC_MOE_GRAD_SCALE_CORRECTIONS.get(id(model_config), 1.0) - - def bind_attention_cp_group(model: torch.nn.Module, packed_seq_params: Any) -> Any: """Bind attention, MoE router and SSM modules to the active CP task. @@ -949,41 +1000,40 @@ def bind_attention_cp_group(model: torch.nn.Module, packed_seq_params: Any) -> A if tp_cp_group.size() != expected_tp_cp_size: raise ValueError("Dynamic MoE TP*CP group has the wrong size") padding_only = bool(getattr(packed_seq_params, "dynamic_cp_padding_only", False)) - for module in model.modules(): - collection = getattr(module, "pg_collection", None) - if collection is not None and hasattr(collection, "cp"): - collection.cp = group - if hasattr(collection, "tp_cp"): - collection.tp_cp = tp_cp_group - if _uses_direct_cp_group(module) and hasattr(module, "cp_group"): - module.cp_group = group - if hasattr(module, "tp_cp_group"): - module.tp_cp_group = tp_cp_group - if _is_hybrid_stack(module): - _bind_hybrid_stack_layout(module, group=group, tp_cp_group=tp_cp_group) - if isinstance(module, Router): - module.cp_group = group + targets = _bind_targets(model) + for module in targets.pg_collections: + module.pg_collection.cp = group + if hasattr(module.pg_collection, "tp_cp"): + module.pg_collection.tp_cp = tp_cp_group + for module in targets.direct_groups: + module.cp_group = group + if hasattr(module, "tp_cp_group"): module.tp_cp_group = tp_cp_group - _bind_router_config(module, padding_only=padding_only) - if _is_mamba_mixer(module): - _rebuild_mamba_cp(module, group) - elif _is_gated_delta_product(module): - _rebuild_gdp_cp(module, group) - elif _is_gated_delta_net(module): - baseline_size = module.cp_size - baseline_split = getattr(module, "feat_dim_split", None) - module.cp_size = context.size - if baseline_split is not None: - scaled_split = [] - for value in baseline_split: - numerator = value * baseline_size - if numerator % context.size: - raise ValueError( - "GatedDeltaNet projection dimensions are not divisible " - f"by runtime CP={context.size}" - ) - scaled_split.append(numerator // context.size) - module.feat_dim_split = tuple(scaled_split) + for module in targets.hybrid_stacks: + _bind_hybrid_stack_layout(module, group=group, tp_cp_group=tp_cp_group) + for module in targets.routers: + module.cp_group = group + module.tp_cp_group = tp_cp_group + _bind_router_config(module, padding_only=padding_only) + for module in targets.mamba_mixers: + _rebuild_mamba_cp(module, group) + for module in targets.gated_delta_products: + _rebuild_gdp_cp(module, group) + for module in targets.gated_delta_nets: + baseline_size = module.cp_size + baseline_split = getattr(module, "feat_dim_split", None) + module.cp_size = context.size + if baseline_split is not None: + scaled_split = [] + for value in baseline_split: + numerator = value * baseline_size + if numerator % context.size: + raise ValueError( + "GatedDeltaNet projection dimensions are not divisible " + f"by runtime CP={context.size}" + ) + scaled_split.append(numerator // context.size) + module.feat_dim_split = tuple(scaled_split) model_packed = copy(packed_seq_params) model_packed.cp_group = group return model_packed diff --git a/nemo_rl/models/megatron/hybridep.py b/nemo_rl/models/megatron/hybridep.py index a86c8f35f5a..245584ed1c7 100644 --- a/nemo_rl/models/megatron/hybridep.py +++ b/nemo_rl/models/megatron/hybridep.py @@ -190,6 +190,10 @@ def pad_packed_seq_for_hybridep( cu_seqlens_kv_padded=cu_seqlens_padded, max_seqlen_q=max_seqlen, max_seqlen_kv=max_seqlen, + # HybridEP just added physical padding beyond the logical boundaries. + # Set this explicitly because model-owned/static packing may have + # already resolved TE's optional inference to False. + pad_between_seqs=True, total_tokens=target_seq_len, ) return input_ids, input_ids_cp_sharded, packed_seq_params, cu_seqlens_padded diff --git a/nemo_rl/models/policy/dynamic_cp.py b/nemo_rl/models/policy/dynamic_cp.py index c359488b28c..9c4a2ee7a25 100644 --- a/nemo_rl/models/policy/dynamic_cp.py +++ b/nemo_rl/models/policy/dynamic_cp.py @@ -23,6 +23,7 @@ CPSyncGroup, DynamicContextParallelConfig, assignments_for_lane, + padding_for_cp, plan_cp_phases, ) from nemo_rl.distributed.named_sharding import NamedSharding @@ -82,6 +83,20 @@ def _minimum_cp_size_for_experts( return minimum +def _dynamic_cp_token_alignment(megatron_cfg: dict[str, Any]) -> int: + """Return the precision/dispatcher alignment used by the planner.""" + fp8 = megatron_cfg.get("fp8_cfg") or {} + alignment = 1 + if fp8.get("enabled"): + alignment = {"blockwise": 128, "mxfp8": 32}.get(fp8["fp8_recipe"], 16) + if ( + megatron_cfg.get("moe_token_dispatcher_type") == "flex" + and megatron_cfg.get("moe_flex_dispatcher_backend") == "hybridep" + ): + alignment = max(alignment, 128) + return alignment + + def dynamic_cp_config(cfg: dict[str, Any]) -> DynamicContextParallelConfig | None: """Read optional config without introducing defaults at worker call sites.""" megatron = cfg.get("megatron_cfg") @@ -110,6 +125,22 @@ def validate_dynamic_cp(cfg: dict[str, Any], *, lanes: int) -> None: ) if mc.get("cuda_graph_impl") not in (None, "none"): raise ValueError("Dynamic CP does not support CUDA graph capture") + cp_comm_type = _model_setting(mc, "cp_comm_type") + cp_comm_types = ( + cp_comm_type + if isinstance(cp_comm_type, (list, tuple)) + else [cp_comm_type] + ) + if ( + "a2a+p2p" in cp_comm_types + or _model_setting(mc, "hierarchical_context_parallel_sizes") is not None + ): + raise ValueError( + "Dynamic CP does not support hierarchical context parallelism " + "(cp_comm_type='a2a+p2p' or " + "hierarchical_context_parallel_sizes). Set cp_comm_type to 'p2p' or " + "'a2a' and remove hierarchical_context_parallel_sizes." + ) if _model_setting(mc, "overlap_moe_expert_parallel_comm"): raise ValueError("Dynamic CP does not support overlap_moe_expert_parallel_comm") if "quantile_balancing" in _routing_types(mc): @@ -135,6 +166,24 @@ def validate_dynamic_cp(cfg: dict[str, Any], *, lanes: int) -> None: sequence_parallel_size=tp if mc["sequence_parallel"] else 1, user_pad_multiple=cfg["make_sequence_length_divisible_by"], ) + maximum_padding = padding_for_cp( + maximum, + sequence_parallel_size=tp if mc["sequence_parallel"] else 1, + user_pad_multiple=cfg["make_sequence_length_divisible_by"], + token_alignment=_dynamic_cp_token_alignment(mc), + ) + sequence_ceiling = cfg["max_total_sequence_length"] + padded_ceiling = ( + (sequence_ceiling + maximum_padding - 1) // maximum_padding * maximum_padding + ) + maximum_capacity = dynamic.tokens_per_rank * maximum + if padded_ceiling > maximum_capacity: + raise ValueError( + "Dynamic CP cannot fit policy.max_total_sequence_length=" + f"{sequence_ceiling} at max_size={maximum}: padding requires " + f"{padded_ceiling} tokens but the configured capacity is " + f"tokens_per_rank * max_size = {maximum_capacity}" + ) @dataclass @@ -193,15 +242,7 @@ def _schedule_parameters( lanes = sharding.shape["data_parallel"] * cp tp = mc["tensor_model_parallel_size"] minimum = _minimum_cp_size_for_experts(mc, dynamic.min_size) - fp8 = mc.get("fp8_cfg") or {} - alignment = 1 - if fp8.get("enabled"): - alignment = {"blockwise": 128, "mxfp8": 32}.get(fp8["fp8_recipe"], 16) - if ( - mc.get("moe_token_dispatcher_type") == "flex" - and mc.get("moe_flex_dispatcher_backend") == "hybridep" - ): - alignment = max(alignment, 128) + alignment = _dynamic_cp_token_alignment(mc) return ( lanes, minimum, @@ -408,48 +449,54 @@ def build_cp_dispatch( ) else: valid_sequences = valid_tokens = 0.0 - samples_by_cp = Counter() - tasks_by_cp = Counter() - packed_tokens = 0 - packed_capacity = 0 - for group in groups: - for task in group.assignments: - samples_by_cp[task.cp_size] += len(task.sample_indices) - if task.sample_indices: - tasks_by_cp[task.cp_size] += 1 - packed_tokens += task.padded_tokens - packed_capacity += task.cp_size * schedule.tokens_per_rank - group_task_ranges = [ - ( - min(len(assignments_for_lane(group, lane)) for lane in range(lanes)), - max(len(assignments_for_lane(group, lane)) for lane in range(lanes)), - ) + assignments_by_group_lane = tuple( + tuple(assignments_for_lane(group, lane) for lane in range(lanes)) for group in groups - ] - logger.info( - "Dynamic CP %s: samples=%d groups=%d uneven_groups=%d " - "local_tasks=[%d,%d] " - "samples_by_cp=%s tasks_by_cp=%s packing_utilization=%.4f " - "valid_sequences=%s valid_tokens=%s schedule=%s", - "train" if training else "score", - gbs, - len(groups), - sum(low < high for low, high in group_task_ranges), - min( - sum(len(assignments_for_lane(group, lane)) for group in groups) - for lane in range(lanes) - ), - max( - sum(len(assignments_for_lane(group, lane)) for group in groups) - for lane in range(lanes) - ), - dict(sorted(samples_by_cp.items())), - dict(sorted(tasks_by_cp.items())), - packed_tokens / packed_capacity if packed_capacity else 0.0, - valid_sequences, - valid_tokens, - "reused" if reused_schedule else "new", ) + if logger.isEnabledFor(logging.INFO): + samples_by_cp = Counter() + tasks_by_cp = Counter() + packed_tokens = 0 + packed_capacity = 0 + for group in groups: + for task in group.assignments: + samples_by_cp[task.cp_size] += len(task.sample_indices) + if task.sample_indices: + tasks_by_cp[task.cp_size] += 1 + packed_tokens += task.padded_tokens + packed_capacity += task.cp_size * schedule.tokens_per_rank + group_task_ranges = [ + ( + min(len(assignments) for assignments in group_assignments), + max(len(assignments) for assignments in group_assignments), + ) + for group_assignments in assignments_by_group_lane + ] + tasks_per_lane = [ + sum( + len(assignments_by_group_lane[group_index][lane]) + for group_index in range(len(groups)) + ) + for lane in range(lanes) + ] + logger.info( + "Dynamic CP %s: samples=%d groups=%d uneven_groups=%d " + "local_tasks=[%d,%d] " + "samples_by_cp=%s tasks_by_cp=%s packing_utilization=%.4f " + "valid_sequences=%s valid_tokens=%s schedule=%s", + "train" if training else "score", + gbs, + len(groups), + sum(low < high for low, high in group_task_ranges), + min(tasks_per_lane), + max(tasks_per_lane), + dict(sorted(samples_by_cp.items())), + dict(sorted(tasks_by_cp.items())), + packed_tokens / packed_capacity if packed_capacity else 0.0, + valid_sequences, + valid_tokens, + "reused" if reused_schedule else "new", + ) for lane in range(lanes): rank_groups = tuple( CPRankGroup( @@ -460,10 +507,10 @@ def build_cp_dispatch( i + start for i in assignment.sample_indices ), ) - for assignment in assignments_for_lane(group, lane) + for assignment in assignments_by_group_lane[group_index][lane] ) ) - for group in groups + for group_index in range(len(groups)) ) rank_steps[lane].append( CPRankStep(rank_groups, valid_sequences, valid_tokens) diff --git a/nemo_rl/models/policy/lm_policy.py b/nemo_rl/models/policy/lm_policy.py index e20cad7a041..98f6270610a 100644 --- a/nemo_rl/models/policy/lm_policy.py +++ b/nemo_rl/models/policy/lm_policy.py @@ -15,7 +15,7 @@ import warnings from collections import defaultdict from contextlib import nullcontext -from typing import Any, Iterable, Optional, Union +from typing import Any, Iterable, Mapping, Optional, Sequence, Union import numpy as np import ray @@ -656,7 +656,7 @@ def _shard_for_train( def _report_sharded_payload( self, - sharded_data: list["SlicedDataDict"], + sharded_data: Sequence[Mapping[str, Any]], boundary: str, ) -> None: """Measure the exact unique per-DP-shard Ray arguments.""" @@ -671,33 +671,63 @@ def _report_sharded_payload( ) def _get_dynamic_cp_outputs( - self, method: str, data: BatchedDataDict, **kwargs: Any + self, + method: str, + data: BatchedDataDict, + *, + timer: Optional[Timer] = None, + **kwargs: Any, ) -> BatchedDataDict: - schedule_batch_size = self.cfg["train_global_batch_size"] - if data.size % schedule_batch_size: - # A standalone score batch need not be a multiple of training GBS. - # In that case plan it as one batch; otherwise preserve every - # training-sized step so score and train can share the schedule. - schedule_batch_size = data.size - schedule = self._matching_dynamic_cp_schedule(data, schedule_batch_size) - dispatch = build_cp_dispatch( - data, - self.cfg, - self.sharding_annotations, - batch_size=schedule_batch_size, - training=False, - schedule=schedule, - ) + labels = { + "get_logprobs": ( + "get_logprobs/shard_data", + "get_logprobs/submit_logprob_futures", + "policy_get_logprobs", + ), + "get_reference_policy_logprobs": ( + "get_reference_policy_logprobs/shard_data", + "get_reference_policy_logprobs/submit_reference_policy_logprob_futures", + "policy_get_reference_logprobs", + ), + "get_topk_logits": ( + "get_topk_logits/shard_data", + "get_topk_logits/submit_topk_logits_futures", + None, + ), + } + shard_label, submit_label, payload_boundary = labels[method] + with timer.time(shard_label) if timer else nullcontext(): + schedule_batch_size = self.cfg["train_global_batch_size"] + if data.size % schedule_batch_size: + # A standalone score batch need not be a multiple of training GBS. + # In that case plan it as one batch; otherwise preserve every + # training-sized step so score and train can share the schedule. + schedule_batch_size = data.size + schedule = self._matching_dynamic_cp_schedule(data, schedule_batch_size) + dispatch = build_cp_dispatch( + data, + self.cfg, + self.sharding_annotations, + batch_size=schedule_batch_size, + training=False, + schedule=schedule, + ) self._dynamic_cp_schedule = dispatch.schedule - futures = self.worker_group.run_all_workers_sharded_data( - method, - data=dispatch.data, - cp_plan=dispatch.plans, - in_sharded_axes=["data_parallel", "context_parallel"], - replicate_on_axes=list(replicated_axes(dynamic_cp=True)), - output_is_replicated=list(replicated_axes(dynamic_cp=True)), - common_kwargs=kwargs, - ) + if payload_boundary is not None: + self._report_sharded_payload( + [shard for dp_shards in dispatch.data for shard in dp_shards], + payload_boundary, + ) + with timer.time(submit_label) if timer else nullcontext(): + futures = self.worker_group.run_all_workers_sharded_data( + method, + data=dispatch.data, + cp_plan=dispatch.plans, + in_sharded_axes=["data_parallel", "context_parallel"], + replicate_on_axes=list(replicated_axes(dynamic_cp=True)), + output_is_replicated=list(replicated_axes(dynamic_cp=True)), + common_kwargs=kwargs, + ) return collect_cp_outputs( self.worker_group.get_all_worker_results(futures), dispatch, data.size ) @@ -731,7 +761,7 @@ def get_logprobs( The logprob of input token i is specified at position i in the output logprobs tensor. """ if self.dynamic_cp: - return self._get_dynamic_cp_outputs("get_logprobs", data) + return self._get_dynamic_cp_outputs("get_logprobs", data, timer=timer) with timer.time("get_logprobs/shard_data") if timer else nullcontext(): sharded_data, unsorted_data_indices = self._shard_for_logprob(data) @@ -780,7 +810,10 @@ def get_reference_policy_logprobs( """ if self.dynamic_cp: return self._get_dynamic_cp_outputs( - "get_reference_policy_logprobs", data, micro_batch_size=micro_batch_size + "get_reference_policy_logprobs", + data, + timer=timer, + micro_batch_size=micro_batch_size, ) with ( @@ -837,7 +870,11 @@ def get_topk_logits( """Dispatch get_topk_logits to workers (no CP/packed support initially).""" if self.dynamic_cp: return self._get_dynamic_cp_outputs( - "get_topk_logits", data, k=k, micro_batch_size=micro_batch_size + "get_topk_logits", + data, + timer=timer, + k=k, + micro_batch_size=micro_batch_size, ) with timer.time("get_topk_logits/shard_data") if timer else nullcontext(): sharded_data, unsorted_data_indices = self._shard_for_logprob(data) @@ -981,8 +1018,14 @@ def train( if dispatch is not None else self._shard_for_train(data, batch_size) ) - if dispatch is None: - self._report_sharded_payload(sharded_data, "policy_train") + self._report_sharded_payload( + ( + sharded_data + if dispatch is None + else [shard for dp_shards in sharded_data for shard in dp_shards] + ), + "policy_train", + ) if self.flops_tracker is not None: self.flops_tracker.reset() diff --git a/nemo_rl/models/policy/teacher_worker_group.py b/nemo_rl/models/policy/teacher_worker_group.py index cd2a2a8bd4d..dd4ca496a85 100644 --- a/nemo_rl/models/policy/teacher_worker_group.py +++ b/nemo_rl/models/policy/teacher_worker_group.py @@ -169,6 +169,10 @@ def __init__( cfg["megatron_cfg"]["peft"]["enabled"] = False if "draft" in cfg: cfg["draft"]["enabled"] = False + # Teacher inference is dispatched through TQ metadata and never receives + # the driver's dynamic-CP rank plan. It may use its configured static CP, + # but must not inherit the student's dynamic execution block. + cfg["megatron_cfg"].pop("dynamic_context_parallel", None) # Router replay keeps the student's rollout and training logprobs # consistent. A frozen teacher has no training pass, and its text-only # TQ fetch does not carry routed_experts, so replay must stay off. diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index 2a61d0e3e37..80eddecda57 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -76,7 +76,6 @@ get_microbatch_iterator, process_global_batch, ) -from nemo_rl.models.megatron.dynamic_cp import dynamic_moe_grad_scale_correction from nemo_rl.models.megatron.pipeline_parallel import ( broadcast_loss_metrics_from_last_stage, broadcast_obj_from_pp_rank, @@ -1428,19 +1427,9 @@ def train( return metrics def _compute_moe_grad_scale(self, global_valid_toks): - """Build a moe_grad_scale_func that normalizes the aux-loss gradient. - - The base scale is 1/global_valid_toks (clamped to avoid division by - zero). Dynamic CP additionally supplies the current task's exact-token - correction; static CP and dense models retain a correction of one. - """ + """Build a moe_grad_scale_func normalized by valid tokens.""" moe_scale = 1.0 / global_valid_toks.clamp(min=1).float() - - def _scale() -> torch.Tensor: - model_config = self._get_model_config() if hasattr(self, "model") else None - return moe_scale * dynamic_moe_grad_scale_correction(model_config) - - return _scale + return lambda: moe_scale def _set_moe_grad_scale_func(self, func): """Set moe_grad_scale_func on the model config for MOE aux loss scaling.""" diff --git a/tests/functional/L1_Functional_Tests_Megatron_1.sh b/tests/functional/L1_Functional_Tests_Megatron_1.sh index 000867eda15..3c165492bfa 100644 --- a/tests/functional/L1_Functional_Tests_Megatron_1.sh +++ b/tests/functional/L1_Functional_Tests_Megatron_1.sh @@ -35,6 +35,7 @@ run_test() { } run_test fast uv run --no-sync bash ./tests/functional/audio_grpo_megatron.sh +run_test uv run --no-sync bash ./tests/functional/dynamic_cp.sh run_test uv run --no-sync bash ./tests/functional/grpo_megatron.sh run_test uv run --no-sync bash ./tests/functional/grpo_megatron_mbridge_restore.sh run_test fast uv run --no-sync bash ./tests/functional/grpo_megatron_eagle3_online.sh diff --git a/tests/functional/dynamic_cp.sh b/tests/functional/dynamic_cp.sh new file mode 100644 index 00000000000..8d3a4b0552f --- /dev/null +++ b/tests/functional/dynamic_cp.sh @@ -0,0 +1,21 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# http://www.apache.org/licenses/LICENSE-2.0 +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +#!/bin/bash +set -euo pipefail + +SCRIPT_DIR=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" &>/dev/null && pwd) +PROJECT_ROOT=$(realpath "${SCRIPT_DIR}/../..") + +cd "${PROJECT_ROOT}" + +torchrun --standalone --nproc-per-node=4 tests/functional/dynamic_cp_loss_parity.py +torchrun --standalone --nproc-per-node=4 tests/functional/dynamic_cp_attention_parity.py diff --git a/tests/unit/distributed/test_dynamic_context_parallel.py b/tests/unit/distributed/test_dynamic_context_parallel.py index 8506df16a34..9c16119b542 100644 --- a/tests/unit/distributed/test_dynamic_context_parallel.py +++ b/tests/unit/distributed/test_dynamic_context_parallel.py @@ -187,6 +187,26 @@ def test_idle_lanes_expand_real_assignment_and_recompute_padding(): check_plan(phases, [3], 8, 128, 2) +def test_idle_lane_expansion_skips_task_when_larger_cp_padding_overflows(): + lengths = [100] * 32 + phases = plan_cp_phases( + lengths, + lanes=2, + min_size=1, + max_size=2, + tokens_per_rank=4096, + sequence_parallel_size=1, + user_pad_multiple=1, + token_alignment=128, + ) + + real = [task for task in phases[0].assignments if task.sample_indices] + placeholders = [task for task in phases[0].assignments if not task.sample_indices] + assert [(task.cp_size, task.padded_tokens) for task in real] == [(1, 4096)] + assert [task.cp_size for task in placeholders] == [1] + check_plan(phases, lengths, 2, 4096, 1) + + def test_maximum_cp_keeps_placeholders_when_real_work_cannot_fill_lanes(): phases = make_plan([3], maximum=4) real = [ diff --git a/tests/unit/distributed/test_dynamic_cp_dispatch.py b/tests/unit/distributed/test_dynamic_cp_dispatch.py index a80a5c93e5b..2c4d14f9cbc 100644 --- a/tests/unit/distributed/test_dynamic_cp_dispatch.py +++ b/tests/unit/distributed/test_dynamic_cp_dispatch.py @@ -160,6 +160,7 @@ def test_global_aux_loss_aligns_full_domain_collective_rounds(self): token_mask=torch.ones(5, 12, dtype=torch.long), ) cfg = self._validation_cfg() + cfg["max_total_sequence_length"] = 10 cfg["megatron_cfg"].update( { "moe_router_load_balancing_type": "global_aux_loss", @@ -187,6 +188,7 @@ def test_global_aux_loss_aligns_full_domain_collective_rounds(self): def _validation_cfg(self) -> dict: return { "make_sequence_length_divisible_by": 1, + "max_total_sequence_length": 32, "sequence_packing": {"enabled": True, "pair_grouping_key": None}, "dynamic_batching": {"enabled": False}, "draft": {"enabled": False}, @@ -227,6 +229,30 @@ def test_dynamic_cp_validation_rejects_moe_microbatch_overlap(self): with self.assertRaisesRegex(ValueError, "overlap_moe_expert_parallel_comm"): validate_dynamic_cp(cfg, lanes=4) + def test_dynamic_cp_validation_rejects_hierarchical_cp(self): + cases = { + "direct cp_comm_type": {"cp_comm_type": "a2a+p2p"}, + "override cp_comm_type list": { + "model_overrides": {"cp_comm_type": ["p2p", "a2a+p2p"]} + }, + "direct hierarchy sizes": { + "hierarchical_context_parallel_sizes": [2, 2] + }, + "override hierarchy sizes": { + "model_overrides": { + "hierarchical_context_parallel_sizes": [2, 2] + } + }, + } + for name, settings in cases.items(): + with self.subTest(name=name): + cfg = self._validation_cfg() + cfg["megatron_cfg"].update(settings) + with self.assertRaisesRegex( + ValueError, "does not support hierarchical context parallelism" + ): + validate_dynamic_cp(cfg, lanes=4) + def test_dynamic_cp_validation_allows_mtp_and_hybridep_prepad_separately(self): mtp_cfg = self._validation_cfg() mtp_cfg["megatron_cfg"]["mtp_num_layers"] = 1 @@ -238,6 +264,11 @@ def test_dynamic_cp_validation_allows_mtp_and_hybridep_prepad_separately(self): "moe_token_dispatcher_type": "flex", "moe_flex_dispatcher_backend": "hybridep", "moe_hybridep_prepad_packed_inputs": True, + "dynamic_context_parallel": { + "enabled": True, + "tokens_per_rank": 256, + "max_size": 4, + }, } ) validate_dynamic_cp(hybridep_cfg, lanes=4) @@ -252,6 +283,15 @@ def test_dynamic_cp_validation_allows_mtp_with_hybridep_prepad(self): ) validate_dynamic_cp(cfg, lanes=4) + def test_dynamic_cp_validation_rejects_unreachable_sequence_ceiling(self): + cfg = self._validation_cfg() + cfg["max_total_sequence_length"] = 33 + + with self.assertRaisesRegex( + ValueError, "cannot fit.*max_total_sequence_length" + ): + validate_dynamic_cp(cfg, lanes=4) + def test_dynamic_cp_keeps_preference_pairs_atomic(self): data = BatchedDataDict( input_ids=torch.arange(4 * 16).reshape(4, 16), diff --git a/tests/unit/models/megatron/test_dynamic_cp_moe.py b/tests/unit/models/megatron/test_dynamic_cp_moe.py index 43c33251ca8..86780cb719c 100644 --- a/tests/unit/models/megatron/test_dynamic_cp_moe.py +++ b/tests/unit/models/megatron/test_dynamic_cp_moe.py @@ -139,6 +139,8 @@ def test_dynamic_binding_updates_router_and_ssm_then_restores(monkeypatch): model.add_module("gdn", _GatedDelta(original_cp)) model.add_module("gdp", _GatedDeltaProduct(original_cp)) model.add_module("router", _Router(original_tp_cp, config)) + # Real transformer layers share one config across many routers. + model.add_module("router2", _Router(original_tp_cp, config)) packed = SimpleNamespace( local_cp_size=2, cp_group=active_cp, @@ -160,6 +162,8 @@ def test_dynamic_binding_updates_router_and_ssm_then_restores(monkeypatch): assert model_packed.cp_group is active_cp assert model.router.cp_group is active_cp assert model.router.tp_cp_group is active_tp_cp + assert model.router2.cp_group is active_cp + assert model.router2.tp_cp_group is active_tp_cp assert model.mamba.pg_collection.cp is active_cp assert model.mamba.cp is not original_mamba_helper assert model.mamba.cp.cp_group is active_cp @@ -174,6 +178,8 @@ def test_dynamic_binding_updates_router_and_ssm_then_restores(monkeypatch): assert model.router.cp_group is original_tp_cp assert model.router.tp_cp_group is original_tp_cp + assert model.router2.cp_group is original_tp_cp + assert model.router2.tp_cp_group is original_tp_cp assert model.mamba.pg_collection.cp is original_cp assert model.mamba.cp is original_mamba_helper assert model.gdn.pg_collection.cp is original_cp @@ -221,6 +227,30 @@ def test_dynamic_padding_task_keeps_only_global_aux_loss(monkeypatch): assert config.moe_z_loss_coeff == 0.01 +def test_dynamic_binding_setup_failure_cleans_global_state(monkeypatch): + from nemo_rl.models.megatron import dynamic_cp + + group = _Group(1) + config = SimpleNamespace(moe_aux_loss_coeff=0.1, moe_z_loss_coeff=0.01) + # The second config fails after the first baseline has been installed. + invalid_config = SimpleNamespace(moe_aux_loss_coeff=0.2) + model = torch.nn.Module() + model.add_module("router", _Router(group, config)) + model.add_module("invalid_router", _Router(group, invalid_config)) + + monkeypatch.setattr(dynamic_cp, "Router", _Router) + + with pytest.raises(AttributeError): + with dynamic_cp.preserve_attention_cp_groups(model): + pytest.fail("setup unexpectedly completed") + + assert id(model) not in dynamic_cp._ACTIVE_BIND_TARGETS + assert id(config) not in dynamic_cp._ROUTER_CONFIG_BASELINES + assert id(invalid_config) not in dynamic_cp._ROUTER_CONFIG_BASELINES + assert model.router.cp_group is group + assert model.router.tp_cp_group is group + + @pytest.mark.parametrize("active_size", [1, 4]) def test_dynamic_moe_scaling_is_applied_by_router_attachment(monkeypatch, active_size): from nemo_rl.models.megatron import dynamic_cp @@ -243,10 +273,8 @@ def test_dynamic_moe_scaling_is_applied_by_router_attachment(monkeypatch, active with dynamic_cp.preserve_attention_cp_groups(model): dynamic_cp.configure_dynamic_moe_loss_scaling(model, padding_mask) - assert dynamic_cp.dynamic_moe_grad_scale_correction(config) == 1.0 assert config.moe_z_loss_coeff == 0.02 - assert dynamic_cp.dynamic_moe_grad_scale_correction(config) == 1.0 assert config.moe_aux_loss_coeff == 0.1 assert config.moe_z_loss_coeff == 0.02 diff --git a/tests/unit/models/megatron/test_dynamic_cp_scaling.py b/tests/unit/models/megatron/test_dynamic_cp_scaling.py index 0361f939202..71c81c0df64 100644 --- a/tests/unit/models/megatron/test_dynamic_cp_scaling.py +++ b/tests/unit/models/megatron/test_dynamic_cp_scaling.py @@ -24,7 +24,10 @@ @pytest.mark.parametrize("active_cp", [1, 2, 4]) @pytest.mark.parametrize("num_microbatches", [1, 3]) @pytest.mark.parametrize("per_token", [False, True]) -def test_mcore_legacy_loss_scaling(base_cp, active_cp, num_microbatches, per_token): +@pytest.mark.parametrize("replicated_cp_loss", [False, True]) +def test_mcore_legacy_loss_scaling( + base_cp, active_cp, num_microbatches, per_token, replicated_cp_loss +): # MCore is optional in the standard CPU test environment. from megatron.core.pipeline_parallel.schedules import forward_step_calc_loss @@ -33,7 +36,7 @@ def test_mcore_legacy_loss_scaling(base_cp, active_cp, num_microbatches, per_tok active_cp_size=active_cp, schedule_cp_size=base_cp, num_microbatches=num_microbatches, - replicated_cp_loss=True, + replicated_cp_loss=replicated_cp_loss, ) metrics = [] loss, _ = forward_step_calc_loss( @@ -49,5 +52,6 @@ def test_mcore_legacy_loss_scaling(base_cp, active_cp, num_microbatches, per_tok is_last_stage=True, ) loss.backward() - assert original_loss.grad.item() * active_cp == pytest.approx(1.0) + expected_grad = 1.0 / active_cp if replicated_cp_loss else 1.0 + assert original_loss.grad.item() == pytest.approx(expected_grad) assert metrics == [{"loss": 2.0}] diff --git a/tests/unit/models/megatron/test_hybridep_data.py b/tests/unit/models/megatron/test_hybridep_data.py index fd9c050fd0c..92b4abd6ecb 100644 --- a/tests/unit/models/megatron/test_hybridep_data.py +++ b/tests/unit/models/megatron/test_hybridep_data.py @@ -111,6 +111,7 @@ def set_group_max(target, **_kwargs): assert torch.equal(padded_params.cu_seqlens_q, cu_seqlens) assert torch.equal(padded_params.cu_seqlens_kv, cu_seqlens) assert torch.equal(padded_params.cu_seqlens_q_padded, torch.tensor([0, 4, 16])) + assert padded_params.pad_between_seqs is True assert padded_params.total_tokens == 16 mock_get_group.assert_called_once_with(check_initialized=False) mock_all_reduce.assert_called_once() diff --git a/tests/unit/models/megatron/test_moe_metrics.py b/tests/unit/models/megatron/test_moe_metrics.py index 87576f1cd2c..d1af13f8d96 100644 --- a/tests/unit/models/megatron/test_moe_metrics.py +++ b/tests/unit/models/megatron/test_moe_metrics.py @@ -127,7 +127,10 @@ def test_dynamic_cp_metrics_use_one_fixed_sum_group(monkeypatch): from nemo_rl.models.megatron.common import get_moe_metrics entry = SimpleNamespace(values=torch.tensor([1.0, 3.0])) - live_tracker = SimpleNamespace(metrics={"load_balancing_loss": entry}) + live_tracker = SimpleNamespace( + metrics={"load_balancing_loss": entry}, + ensure_initialized=lambda *_args: None, + ) reductions = [] fixed_group = object() @@ -155,6 +158,8 @@ def _all_reduce(values, *, group): metrics = get_moe_metrics( loss_scale=0.25, + num_layers=2, + track_names=["load_balancing_loss"], dynamic_parallel_group=fixed_group, ) @@ -162,6 +167,14 @@ def _all_reduce(values, *, group): assert metrics["load_balancing_loss"] == pytest.approx(1.0) +@pytest.mark.mcore +def test_dynamic_cp_metrics_require_deterministic_collective_inputs(): + from nemo_rl.models.megatron.common import get_moe_metrics + + with pytest.raises(ValueError, match="explicit track_names and num_layers"): + get_moe_metrics(loss_scale=1.0, dynamic_parallel_group=object()) + + @pytest.mark.mcore def test_dynamic_cp_avg_group_metrics_use_rank_participation_scale(monkeypatch): """z-loss must reproduce MCore's AVG over every participating rank.""" @@ -173,7 +186,8 @@ def test_dynamic_cp_avg_group_metrics_use_rank_participation_scale(monkeypatch): # never populated by record(); the name must still select AVG semantics. z_entry = SimpleNamespace(values=torch.tensor([2.0, 4.0]), avg_group=None) live_tracker = SimpleNamespace( - metrics={"load_balancing_loss": aux_entry, "z_loss": z_entry} + metrics={"load_balancing_loss": aux_entry, "z_loss": z_entry}, + ensure_initialized=lambda *_args: None, ) fixed_group = object() reductions = [] @@ -200,6 +214,8 @@ def _all_reduce(values, *, group): metrics = get_moe_metrics( loss_scale=0.25, + num_layers=2, + track_names=["load_balancing_loss", "z_loss"], dynamic_parallel_group=fixed_group, dynamic_avg_loss_scale=0.125, ) @@ -215,7 +231,10 @@ def test_dynamic_cp_global_aux_uses_aligned_round_scale(monkeypatch): from nemo_rl.models.megatron.common import get_moe_metrics entry = SimpleNamespace(values=torch.tensor([2.0]), avg_group=None) - live_tracker = SimpleNamespace(metrics={"global_load_balancing_loss": entry}) + live_tracker = SimpleNamespace( + metrics={"global_load_balancing_loss": entry}, + ensure_initialized=lambda *_args: None, + ) monkeypatch.setattr( megatron_module.common, "get_moe_metrics_tracker", lambda: live_tracker ) @@ -235,6 +254,8 @@ def test_dynamic_cp_global_aux_uses_aligned_round_scale(monkeypatch): metrics = get_moe_metrics( loss_scale=0.1, + num_layers=1, + track_names=["global_load_balancing_loss"], dynamic_parallel_group=object(), dynamic_global_loss_scale=0.25, ) diff --git a/tests/unit/models/policy/test_megatron_worker.py b/tests/unit/models/policy/test_megatron_worker.py index 5c60ce1a348..80bcad4592e 100644 --- a/tests/unit/models/policy/test_megatron_worker.py +++ b/tests/unit/models/policy/test_megatron_worker.py @@ -1766,26 +1766,6 @@ def test_compute_moe_grad_scale_normalizes_by_valid_tokens(): assert torch.allclose(scale_fn(), torch.tensor(0.25)) -def test_compute_moe_grad_scale_applies_dynamic_task_correction(monkeypatch): - from nemo_rl.models.policy.workers import megatron_policy_worker as worker_module - - worker = object.__new__(worker_module.MegatronPolicyWorkerImpl) - _disable_opd_full(worker) - model_config = SimpleNamespace() - worker.model = SimpleNamespace(config=model_config) - monkeypatch.setattr( - worker_module, - "dynamic_moe_grad_scale_correction", - lambda config: 1.5 if config is model_config else 1.0, - ) - - scale_fn = worker_module.MegatronPolicyWorkerImpl._compute_moe_grad_scale( - worker, torch.tensor(4.0) - ) - - assert torch.allclose(scale_fn(), torch.tensor(0.375)) - - def test_compute_moe_grad_scale_clamps_zero_valid_tokens(): """clamp(min=1) must guard against division by zero when no valid tokens.""" from nemo_rl.models.policy.workers.megatron_policy_worker import ( diff --git a/tests/unit/models/policy/test_policy_validation.py b/tests/unit/models/policy/test_policy_validation.py index edbabf8df94..4afe95ce053 100644 --- a/tests/unit/models/policy/test_policy_validation.py +++ b/tests/unit/models/policy/test_policy_validation.py @@ -20,8 +20,9 @@ when the cluster size is insufficient for the specified parallelism configuration. """ +from contextlib import nullcontext from types import SimpleNamespace -from unittest.mock import MagicMock, patch +from unittest.mock import MagicMock, call, patch import pytest import torch @@ -67,8 +68,10 @@ def create_mock_tokenizer(): def test_dynamic_cp_score_keeps_complete_training_sized_steps() -> None: policy = Policy.__new__(Policy) policy.cfg = {"train_global_batch_size": 2} + policy.debug_payload_metrics = False policy.sharding_annotations = object() policy._dynamic_cp_schedule = None + policy._report_sharded_payload = MagicMock() policy.worker_group = MagicMock() policy.worker_group.run_all_workers_sharded_data.return_value = "futures" policy.worker_group.get_all_worker_results.return_value = ["worker-results"] @@ -76,11 +79,13 @@ def test_dynamic_cp_score_keeps_complete_training_sized_steps() -> None: schedule = object() dispatch = SimpleNamespace( schedule=schedule, - data="sharded-data", + data=[[BatchedDataDict(input_ids=torch.zeros(2, 8, dtype=torch.long))]], plans="rank-plans", output_rows=[], ) expected = BatchedDataDict(logprobs=torch.zeros(4, 8)) + timer = MagicMock() + timer.time.side_effect = lambda _label: nullcontext() with ( patch( @@ -92,11 +97,20 @@ def test_dynamic_cp_score_keeps_complete_training_sized_steps() -> None: return_value=expected, ), ): - result = Policy._get_dynamic_cp_outputs(policy, "get_logprobs", data) + result = Policy._get_dynamic_cp_outputs( + policy, "get_logprobs", data, timer=timer + ) assert result is expected assert policy._dynamic_cp_schedule is schedule assert build_dispatch.call_args.kwargs["batch_size"] == 2 + assert timer.time.call_args_list == [ + call("get_logprobs/shard_data"), + call("get_logprobs/submit_logprob_futures"), + ] + policy._report_sharded_payload.assert_called_once_with( + [dispatch.data[0][0]], "policy_get_logprobs" + ) def create_dtensor_config( diff --git a/tests/unit/models/policy/test_teacher_worker_group.py b/tests/unit/models/policy/test_teacher_worker_group.py index 0591ca0b0db..97ea44efd41 100644 --- a/tests/unit/models/policy/test_teacher_worker_group.py +++ b/tests/unit/models/policy/test_teacher_worker_group.py @@ -94,8 +94,8 @@ def test_create_teacher_configs_deduplicates(): assert len(configs) == 2 -def test_teacher_worker_group_disables_student_router_replay(monkeypatch): - """Frozen teachers do not require rollout-to-training route consistency.""" +def test_teacher_worker_group_disables_student_only_runtime_features(monkeypatch): + """Frozen teachers do not use student router replay or dynamic dispatch.""" import nemo_rl.distributed.worker_groups as worker_groups from nemo_rl.models.policy.teacher_worker_group import ( TeacherConfig, @@ -119,7 +119,13 @@ def __init__(self, cluster, worker_builder, **kwargs): cluster.world_size.return_value = 1 policy_config = { "model_name": "/ckpt/student", - "megatron_cfg": {"enabled": True}, + "megatron_cfg": { + "enabled": True, + "dynamic_context_parallel": { + "enabled": True, + "tokens_per_rank": 4096, + }, + }, "dtensor_cfg": {"enabled": False}, "sequence_packing": {"enabled": False}, "dynamic_batching": {"enabled": False}, @@ -148,7 +154,10 @@ def __init__(self, cluster, worker_builder, **kwargs): assert captured["cfg"]["router_replay"]["enabled"] is False assert teacher.cfg["router_replay"]["enabled"] is False + assert "dynamic_context_parallel" not in captured["cfg"]["megatron_cfg"] + assert "dynamic_context_parallel" not in teacher.cfg["megatron_cfg"] assert policy_config["router_replay"]["enabled"] is True + assert policy_config["megatron_cfg"]["dynamic_context_parallel"]["enabled"] is True def test_teacher_worker_group_drops_the_student_pretrained_checkpoint(monkeypatch): diff --git a/tests/unit/test_dynamic_cp_comparison_recipes.py b/tests/unit/test_dynamic_cp_comparison_recipes.py index d7faa43e9f2..46e3248e51b 100644 --- a/tests/unit/test_dynamic_cp_comparison_recipes.py +++ b/tests/unit/test_dynamic_cp_comparison_recipes.py @@ -13,6 +13,19 @@ RECIPE_DIR = ( Path(__file__).resolve().parents[2] / "examples/configs/recipes/llm/performance" ) +_LOCAL_RECIPES = ( + "grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-10step.yaml", + "grpo-qwen3-30ba3b-4n4g-async-1off-megatron-staticcp-10step.yaml", + "grpo-qwen3-32b-4n4g-async-1off-megatron-dynamiccp-10step.yaml", + "grpo-qwen3-32b-4n4g-async-1off-megatron-staticcp-10step.yaml", +) + +# The comparison recipes are personal benchmark inputs, not repository +# examples. Run these assertions when that local set exists and skip otherwise. +pytestmark = pytest.mark.skipif( + any(not (RECIPE_DIR / recipe_name).exists() for recipe_name in _LOCAL_RECIPES), + reason="local dynamic/static-CP comparison recipes are not present", +) @pytest.mark.parametrize( diff --git a/tests/unit/test_dynamic_cp_moe_recipes.py b/tests/unit/test_dynamic_cp_moe_recipes.py index 16d892ab234..59b82d1afe2 100644 --- a/tests/unit/test_dynamic_cp_moe_recipes.py +++ b/tests/unit/test_dynamic_cp_moe_recipes.py @@ -15,8 +15,7 @@ RECIPE_DIR = ( - Path(__file__).resolve().parents[2] - / "examples/configs/recipes/llm/performance" + Path(__file__).resolve().parents[2] / "examples/configs/recipes/llm/performance" ) RECIPES = ( ( @@ -42,6 +41,13 @@ ), ) +# These benchmark recipes are intentionally local/ignored. Keep the topology +# checks useful for local recipe development without breaking a clean checkout. +pytestmark = pytest.mark.skipif( + any(not (RECIPE_DIR / recipe_name).exists() for recipe_name, *_ in RECIPES), + reason="local dynamic-CP benchmark recipes are not present", +) + @pytest.mark.parametrize( "recipe_name,routing_type,effective_minimum,total_nodes,expected_policy_nodes", diff --git a/tools/analyze_dynamic_cp_comparison.py b/tools/analyze_dynamic_cp_comparison.py new file mode 100755 index 00000000000..76fe136c518 --- /dev/null +++ b/tools/analyze_dynamic_cp_comparison.py @@ -0,0 +1,629 @@ +#!/usr/bin/env python3 +"""Build CSV/Markdown speedup sheets for matched Dynamic-CP/static-CP runs. + +The launcher in ``tools/launch_dynamic_cp_comparison.sh`` writes runs as: + + //// + +This tool accepts the ```` directory. It reads TensorBoard +events when TensorBoard is installed and falls back to the human-readable +timing blocks in Slurm/Ray logs. +""" + +from __future__ import annotations + +import argparse +import ast +import csv +import json +import math +import re +import statistics +import sys +from dataclasses import dataclass +from pathlib import Path +from typing import Iterable + + +MODEL_METADATA = { + "qwen30": { + "gbs": 512, + "max_sequence_length": 8192, + "static_cp": 2, + "dynamic_cp": "1-2", + }, + "qwen32": { + "gbs": 512, + "max_sequence_length": 16384, + "static_cp": 4, + "dynamic_cp": "1-4", + }, + "nano": { + "gbs": 64, + "max_sequence_length": 8192, + "static_cp": 4, + "dynamic_cp": "4-16 effective (1-16 configured)", + }, +} + +MODEL_ALIASES = { + "qwen30": "qwen30", + "qwen3-30b": "qwen30", + "qwen3-30ba3b": "qwen30", + "qwen32": "qwen32", + "qwen3-32b": "qwen32", + "nano": "nano", + "nt3-nano": "nano", + "nemotron3-nano": "nano", +} + +PREFERRED_METRICS = [ + "timing/train/total_step_time", + "timing/train/policy_training", + "timing/train/policy_and_reference_logprobs", + "timing/train/generation", + "timing/train/prepare_for_generation/total", + "timing/train/prepare_for_generation/transfer_and_update_weights", + "timing/train/training_prep", + "timing/train/logprob_inference_prep", + "timing/train/reward_calculation", + "timing/train/data_processing", + "timing/train/valid_tokens_per_sec_per_gpu", +] + +ANSI_RE = re.compile(r"\x1b\[[0-?]*[ -/]*[@-~]") +STEP_RE = re.compile(r"\bStep\s+(\d+)\s*/\s*\d+", re.IGNORECASE) +TRAIN_STEP_RE = re.compile(r"\btrain step\s+(\d+)\s*/\s*\d+", re.IGNORECASE) +FLOAT = r"[-+]?(?:\d+(?:\.\d*)?|\.\d+)(?:[eE][-+]?\d+)?" +TOTAL_TIME_RE = re.compile(rf"Total step time\s*:\s*({FLOAT})s\b", re.IGNORECASE) +TIMING_VALUE_RE = re.compile(rf"[•*-]\s*([A-Za-z0-9_./-]+)\s*:\s*({FLOAT})s(?:\s|$)") +FLAT_METRIC_RE = re.compile( + rf"(timing/train/[A-Za-z0-9_./-]+)\s*[=:]\s*({FLOAT})" +) +DYNAMIC_CP_RE = re.compile( + rf"Dynamic CP .*?samples_by_cp=(\{{.*?\}}).*?tasks_by_cp=(\{{.*?\}})" + rf".*?packing_utilization=({FLOAT})" +) + + +@dataclass(frozen=True) +class RunSpec: + model: str + mode: str + path: Path + + +@dataclass +class RunData: + spec: RunSpec + # metric -> step -> value. Later event/log files replace an earlier value. + metrics: dict[str, dict[int, float]] + packing_utilizations: list[float] + samples_by_cp: dict[int, int] + tasks_by_cp: dict[int, int] + source: str + + +def canonical_model(value: str) -> str: + key = value.strip().lower() + if key not in MODEL_ALIASES: + raise ValueError(f"unknown model '{value}'") + return MODEL_ALIASES[key] + + +def canonical_mode(value: str) -> str: + key = value.strip().lower() + if key in {"dynamic", "dyncp", "dyn"}: + return "dynamic" + if key in {"static", "staticcp", "no-dyncp", "nodyncp"}: + return "static" + raise ValueError(f"unknown mode '{value}'") + + +def parse_run_arg(value: str) -> RunSpec: + try: + model, mode, path = value.split(":", 2) + except ValueError as exc: + raise argparse.ArgumentTypeError( + "--run must be MODEL:MODE:PATH (for example qwen30:dynamic:/logs/run)" + ) from exc + try: + return RunSpec(canonical_model(model), canonical_mode(mode), Path(path)) + except ValueError as exc: + raise argparse.ArgumentTypeError(str(exc)) from exc + + +def discover_runs(root: Path) -> list[RunSpec]: + runs: list[RunSpec] = [] + for model in MODEL_METADATA: + for mode in ("dynamic", "static"): + path = root / model / mode + if path.is_dir(): + runs.append(RunSpec(model, mode, path)) + return runs + + +def _record( + metrics: dict[str, dict[int, float]], metric: str, step: int, value: float +) -> None: + if math.isfinite(value): + metrics.setdefault(metric, {})[step] = value + + +def read_tensorboard(path: Path) -> tuple[dict[str, dict[int, float]], bool]: + event_files = sorted( + path.rglob("events*tfevents*"), key=lambda item: item.stat().st_mtime + ) + if not event_files: + return {}, False + + try: + from tensorboard.backend.event_processing import event_accumulator + except ImportError: + print( + "WARNING: TensorBoard event files exist, but the 'tensorboard' package " + "is unavailable; falling back to text logs.", + file=sys.stderr, + ) + return {}, False + + size_guidance = { + event_accumulator.SCALARS: 0, + event_accumulator.TENSORS: 0, + } + metrics: dict[str, dict[int, float]] = {} + for event_file in event_files: + accumulator = event_accumulator.EventAccumulator( + str(event_file), size_guidance=size_guidance + ) + accumulator.Reload() + for metric in accumulator.scalars.Keys(): + if not metric.startswith("timing/train/"): + continue + for scalar in accumulator.Scalars(metric): + _record(metrics, metric, int(scalar.step), float(scalar.value)) + return metrics, bool(metrics) + + +def _parse_cp_dict(value: str) -> dict[int, int]: + parsed = ast.literal_eval(value) + if not isinstance(parsed, dict): + return {} + return {int(key): int(count) for key, count in parsed.items()} + + +def _merge_counts(target: dict[int, int], values: dict[int, int]) -> None: + for key, value in values.items(): + target[key] = target.get(key, 0) + value + + +def _candidate_text_files(path: Path) -> list[Path]: + candidates: set[Path] = set() + for pattern in ("*.out", "*.log", "*.txt"): + candidates.update(path.rglob(pattern)) + return sorted(candidates, key=lambda item: item.stat().st_mtime) + + +def _walk_json(value: object, prefix: str = "") -> Iterable[tuple[str, object]]: + if not isinstance(value, dict): + return + for key, child in value.items(): + full_key = f"{prefix}/{key}" if prefix else str(key) + if isinstance(child, dict): + yield from _walk_json(child, full_key) + else: + yield full_key, child + + +def read_text_logs( + path: Path, *, read_timing: bool +) -> tuple[dict[str, dict[int, float]], list[float], dict[int, int], dict[int, int]]: + metrics: dict[str, dict[int, float]] = {} + packing_utilizations: list[float] = [] + samples_by_cp: dict[int, int] = {} + tasks_by_cp: dict[int, int] = {} + + for log_path in _candidate_text_files(path): + current_step: int | None = None + inferred_step = 0 + in_timing = False + try: + handle = log_path.open("r", encoding="utf-8", errors="replace") + except OSError as exc: + print(f"WARNING: cannot read {log_path}: {exc}", file=sys.stderr) + continue + + with handle: + for raw_line in handle: + line = ANSI_RE.sub("", raw_line).strip() + step_match = STEP_RE.search(line) or TRAIN_STEP_RE.search(line) + if step_match: + current_step = int(step_match.group(1)) + in_timing = False + + cp_match = DYNAMIC_CP_RE.search(line) + if cp_match: + try: + _merge_counts(samples_by_cp, _parse_cp_dict(cp_match.group(1))) + _merge_counts(tasks_by_cp, _parse_cp_dict(cp_match.group(2))) + packing_utilizations.append(float(cp_match.group(3))) + except (SyntaxError, TypeError, ValueError): + pass + + if not read_timing: + continue + + if "Timing:" in line: + in_timing = True + continue + if "Performance Metrics:" in line or "Training Results:" in line: + in_timing = False + + total_match = TOTAL_TIME_RE.search(line) + if total_match: + if current_step is None: + inferred_step += 1 + current_step = inferred_step + _record( + metrics, + "timing/train/total_step_time", + current_step, + float(total_match.group(1)), + ) + continue + + if in_timing and current_step is not None: + timing_match = TIMING_VALUE_RE.search(line) + if timing_match: + _record( + metrics, + f"timing/train/{timing_match.group(1)}", + current_step, + float(timing_match.group(2)), + ) + + for flat_match in FLAT_METRIC_RE.finditer(line): + if current_step is not None: + _record( + metrics, + flat_match.group(1), + current_step, + float(flat_match.group(2)), + ) + + if line.startswith("{") and len(line) < 1_000_000: + try: + payload = json.loads(line) + except json.JSONDecodeError: + continue + if not isinstance(payload, dict): + continue + json_step = payload.get("step", payload.get("_step", current_step)) + if not isinstance(json_step, (int, float)): + continue + for key, value in _walk_json(payload): + if key.startswith("timing/train/") and isinstance( + value, (int, float) + ): + _record(metrics, key, int(json_step), float(value)) + + return metrics, packing_utilizations, samples_by_cp, tasks_by_cp + + +def load_run(spec: RunSpec) -> RunData: + if not spec.path.is_dir(): + raise FileNotFoundError(f"run directory does not exist: {spec.path}") + + metrics, used_tensorboard = read_tensorboard(spec.path) + text_metrics, utilization, samples_by_cp, tasks_by_cp = read_text_logs( + spec.path, read_timing=not used_tensorboard + ) + if not used_tensorboard: + metrics = text_metrics + return RunData( + spec=spec, + metrics=metrics, + packing_utilizations=utilization, + samples_by_cp=samples_by_cp, + tasks_by_cp=tasks_by_cp, + source="tensorboard" if used_tensorboard else "text", + ) + + +def kept_values(steps: dict[int, float], warmup_steps: int) -> list[float]: + ordered = [value for _, value in sorted(steps.items())] + return ordered[warmup_steps:] + + +def metric_order(metric: str) -> tuple[int, str]: + try: + return PREFERRED_METRICS.index(metric), metric + except ValueError: + return len(PREFERRED_METRICS), metric + + +def higher_is_better(metric: str) -> bool: + name = metric.lower() + return "per_sec" in name or "throughput" in name or name.endswith("_tps") + + +def fmt(value: float | None, digits: int = 3) -> str: + return "" if value is None else f"{value:.{digits}f}" + + +def build_rows( + runs: dict[tuple[str, str], RunData], warmup_steps: int +) -> tuple[list[dict[str, object]], list[str]]: + rows: list[dict[str, object]] = [] + warnings: list[str] = [] + for model, metadata in MODEL_METADATA.items(): + dynamic = runs.get((model, "dynamic")) + static = runs.get((model, "static")) + if dynamic is None or static is None: + missing = "dynamic" if dynamic is None else "static" + warnings.append(f"{model}: missing {missing} run") + continue + + common_metrics = set(dynamic.metrics) & set(static.metrics) + if not common_metrics: + warnings.append(f"{model}: dynamic/static runs have no shared timing metrics") + continue + + for metric in sorted(common_metrics, key=metric_order): + dynamic_values = kept_values(dynamic.metrics[metric], warmup_steps) + static_values = kept_values(static.metrics[metric], warmup_steps) + if not dynamic_values or not static_values: + warnings.append( + f"{model}/{metric}: no samples remain after dropping " + f"{warmup_steps} warmup step(s)" + ) + continue + dynamic_mean = statistics.fmean(dynamic_values) + static_mean = statistics.fmean(static_values) + is_higher_better = higher_is_better(metric) + if dynamic_mean == 0 or static_mean == 0: + speedup = math.nan + improvement = math.nan + elif is_higher_better: + speedup = dynamic_mean / static_mean + improvement = (dynamic_mean - static_mean) / static_mean * 100 + else: + speedup = static_mean / dynamic_mean + improvement = (static_mean - dynamic_mean) / static_mean * 100 + rows.append( + { + "model": model, + **metadata, + "metric": metric, + "direction": "higher is better" if is_higher_better else "lower is better", + "dynamic_samples": len(dynamic_values), + "static_samples": len(static_values), + "dynamic_mean": dynamic_mean, + "dynamic_median": statistics.median(dynamic_values), + "static_mean": static_mean, + "static_median": statistics.median(static_values), + "speedup_x": speedup, + "improvement_percent": improvement, + } + ) + return rows, warnings + + +def write_speedup_csv(path: Path, rows: list[dict[str, object]]) -> None: + fieldnames = [ + "model", + "gbs", + "max_sequence_length", + "static_cp", + "dynamic_cp", + "metric", + "direction", + "dynamic_samples", + "static_samples", + "dynamic_mean", + "dynamic_median", + "static_mean", + "static_median", + "speedup_x", + "improvement_percent", + ] + with path.open("w", newline="", encoding="utf-8") as handle: + writer = csv.DictWriter(handle, fieldnames=fieldnames) + writer.writeheader() + for row in rows: + output = dict(row) + for key in ( + "dynamic_mean", + "dynamic_median", + "static_mean", + "static_median", + "speedup_x", + "improvement_percent", + ): + output[key] = fmt(float(output[key]), 6) + writer.writerow(output) + + +def write_runs_csv(path: Path, runs: dict[tuple[str, str], RunData]) -> None: + fields = [ + "model", + "mode", + "path", + "metric_source", + "timing_metrics", + "timing_steps", + "mean_packing_utilization", + "samples_by_cp", + "tasks_by_cp", + ] + with path.open("w", newline="", encoding="utf-8") as handle: + writer = csv.DictWriter(handle, fieldnames=fields) + writer.writeheader() + for key in sorted(runs): + run = runs[key] + all_steps = {step for values in run.metrics.values() for step in values} + writer.writerow( + { + "model": run.spec.model, + "mode": run.spec.mode, + "path": run.spec.path, + "metric_source": run.source, + "timing_metrics": len(run.metrics), + "timing_steps": len(all_steps), + "mean_packing_utilization": fmt( + statistics.fmean(run.packing_utilizations) + if run.packing_utilizations + else None, + 6, + ), + "samples_by_cp": json.dumps(run.samples_by_cp, sort_keys=True), + "tasks_by_cp": json.dumps(run.tasks_by_cp, sort_keys=True), + } + ) + + +def write_markdown( + path: Path, + rows: list[dict[str, object]], + runs: dict[tuple[str, str], RunData], + warnings: list[str], + warmup_steps: int, +) -> None: + lines = [ + "# Dynamic CP vs static CP speedup", + "", + f"Warmup steps excluded per metric: **{warmup_steps}**.", + "Speedup for time metrics is `static mean / dynamic mean`; values above 1.0 favor Dynamic CP.", + "The static arms use CP2/CP4/CP4—not CP1/no-CP.", + "", + "| Model | GBS | Max seq | Static CP | Dynamic CP | Metric | Dynamic | Static | Speedup | Improvement |", + "|---|---:|---:|---:|---|---|---:|---:|---:|---:|", + ] + for row in rows: + metric = str(row["metric"]).removeprefix("timing/train/") + unit = "" if higher_is_better(str(row["metric"])) else " s" + lines.append( + f"| {row['model']} | {row['gbs']} | {row['max_sequence_length']} " + f"| {row['static_cp']} | {row['dynamic_cp']} | `{metric}` " + f"| {fmt(float(row['dynamic_mean']))}{unit} " + f"| {fmt(float(row['static_mean']))}{unit} " + f"| {fmt(float(row['speedup_x']))}x " + f"| {fmt(float(row['improvement_percent']), 2)}% |" + ) + + lines.extend(["", "## Run diagnostics", ""]) + for key in sorted(runs): + run = runs[key] + all_steps = sorted({step for values in run.metrics.values() for step in values}) + utilization = ( + fmt(statistics.fmean(run.packing_utilizations), 4) + if run.packing_utilizations + else "n/a" + ) + lines.append( + f"- `{run.spec.model}/{run.spec.mode}`: source={run.source}, " + f"steps={all_steps or 'none'}, mean Dynamic-CP packing utilization={utilization}, " + f"tasks_by_cp={run.tasks_by_cp or 'n/a'}" + ) + + if warnings: + lines.extend(["", "## Warnings", ""]) + lines.extend(f"- {warning}" for warning in warnings) + path.write_text("\n".join(lines) + "\n", encoding="utf-8") + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--root", + type=Path, + help="Comparison directory containing qwen30/, qwen32/, and nano/", + ) + parser.add_argument( + "--run", + action="append", + default=[], + type=parse_run_arg, + metavar="MODEL:MODE:PATH", + help="Explicit run directory; repeat for each arm instead of --root", + ) + parser.add_argument( + "--output-dir", + type=Path, + help="Output directory (default: ROOT/analysis or ./dynamic-cp-analysis)", + ) + parser.add_argument( + "--warmup-steps", + type=int, + default=1, + help="Discard this many earliest samples from each metric (default: 1)", + ) + parser.add_argument( + "--strict", + action="store_true", + help="Exit nonzero when a model/mode pair or usable metric is missing", + ) + args = parser.parse_args() + if args.root is None and not args.run: + parser.error("provide --root or at least one --run") + if args.warmup_steps < 0: + parser.error("--warmup-steps must be zero or greater") + return args + + +def main() -> int: + args = parse_args() + specs = list(args.run) + if args.root is not None: + if not args.root.is_dir(): + print(f"ERROR: root directory does not exist: {args.root}", file=sys.stderr) + return 2 + specs.extend(discover_runs(args.root)) + if not specs: + print( + "ERROR: no runs found; expected ROOT/{qwen30,qwen32,nano}/{dynamic,static}", + file=sys.stderr, + ) + return 2 + + runs: dict[tuple[str, str], RunData] = {} + for spec in specs: + key = (spec.model, spec.mode) + if key in runs: + print(f"ERROR: duplicate run for {spec.model}/{spec.mode}", file=sys.stderr) + return 2 + try: + runs[key] = load_run(spec) + except (FileNotFoundError, OSError, ValueError) as exc: + print(f"ERROR: {exc}", file=sys.stderr) + return 2 + + rows, warnings = build_rows(runs, args.warmup_steps) + if not rows: + warnings.append("no comparison rows were generated") + + output_dir = args.output_dir + if output_dir is None: + output_dir = ( + args.root / "analysis" if args.root else Path("dynamic-cp-analysis") + ) + output_dir.mkdir(parents=True, exist_ok=True) + speedup_csv = output_dir / "dynamic_cp_speedup.csv" + runs_csv = output_dir / "dynamic_cp_runs.csv" + markdown = output_dir / "dynamic_cp_speedup.md" + write_speedup_csv(speedup_csv, rows) + write_runs_csv(runs_csv, runs) + write_markdown(markdown, rows, runs, warnings, args.warmup_steps) + + for warning in warnings: + print(f"WARNING: {warning}", file=sys.stderr) + print(f"Wrote {speedup_csv}") + print(f"Wrote {runs_csv}") + print(f"Wrote {markdown}") + if args.strict and warnings: + return 1 + return 0 if rows else 1 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tools/launch_dynamic_cp_comparison.sh b/tools/launch_dynamic_cp_comparison.sh new file mode 100755 index 00000000000..f82230dfe35 --- /dev/null +++ b/tools/launch_dynamic_cp_comparison.sh @@ -0,0 +1,219 @@ +#!/bin/bash +set -euo pipefail + +# Submit one arm of the matched Dynamic-CP versus static-CP comparison. +# +# Usage: +# tools/launch_dynamic_cp_comparison.sh qwen30 dynamic +# tools/launch_dynamic_cp_comparison.sh qwen30 static +# +# "static" means Dynamic CP is disabled, but CP is still greater than one. It +# is deliberately not a no-CP (CP1) baseline. + +usage() { + cat <<'EOF' +Usage: launch_dynamic_cp_comparison.sh MODEL MODE + +MODEL: qwen30 | qwen32 | nano +MODE: dynamic | static + +Useful environment variables: + COMPARISON_NAME Common name for all six runs (required for pairing) + NRL_MAX_STEPS Training steps; default: 3 + DYNAMIC_CP_RESULTS_ROOT Shared output root on /lustre + CONTAINER NeMo-RL squashfs/image + SLURM_ACCOUNT Default: coreai_dlalgo_nemorl + SLURM_PARTITION Default: batch + WALLTIME Default: 1:59:00 + ENABLE_WANDB 0 (default) or 1 + DRY_RUN 1 prints the submission without calling sbatch +EOF +} + +if [[ $# -ne 2 ]]; then + usage >&2 + exit 2 +fi + +MODEL="$1" +MODE="$2" +case "${MODE}" in + dynamic|static) ;; + *) + echo "ERROR: MODE must be 'dynamic' or 'static' (got '${MODE}')." >&2 + exit 2 + ;; +esac + +SCRIPT_DIR="$(cd -L -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd -L)" +NEMO_RL_ROOT="${NEMO_RL_ROOT:-$(cd -L -- "${SCRIPT_DIR}/.." && pwd -L)}" + +case "${MODEL}" in + qwen30) + NODES=4 + SEGMENT_SIZE=2 + GBS=512 + MAX_SEQUENCE_LENGTH=8192 + STATIC_CP=2 + DYNAMIC_CP_RANGE="1-2" + DYNAMIC_RECIPE="examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-dynamiccp-10step.yaml" + STATIC_RECIPE="examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-megatron-staticcp-10step.yaml" + ;; + qwen32) + NODES=4 + SEGMENT_SIZE=2 + GBS=512 + MAX_SEQUENCE_LENGTH=16384 + STATIC_CP=4 + DYNAMIC_CP_RANGE="1-4" + DYNAMIC_RECIPE="examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-async-1off-megatron-dynamiccp-10step.yaml" + STATIC_RECIPE="examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-async-1off-megatron-staticcp-10step.yaml" + ;; + nano) + NODES=8 + # The Nano recipe does not set cluster.segment_size, so do not claim an + # allocation-side segment that the runtime config does not also express. + SEGMENT_SIZE="" + GBS=64 + MAX_SEQUENCE_LENGTH=8192 + STATIC_CP=4 + # Configured min=1; TP2*CP must contain EP8, so the effective minimum is 4. + DYNAMIC_CP_RANGE="4-16 (effective; configured 1-16)" + DYNAMIC_RECIPE="examples/configs/recipes/llm/performance/grpo-nemotron3-nano-30ba3b-8n4g-megatron-dynamiccp-quick.yaml" + STATIC_RECIPE="examples/configs/recipes/llm/performance/grpo-nemotron3-nano-30ba3b-8n4g-megatron-staticcp-quick.yaml" + ;; + *) + echo "ERROR: MODEL must be qwen30, qwen32, or nano (got '${MODEL}')." >&2 + exit 2 + ;; +esac + +if [[ "${MODE}" == "dynamic" ]]; then + RECIPE="${DYNAMIC_RECIPE}" + CP_DESCRIPTION="Dynamic CP ${DYNAMIC_CP_RANGE}" +else + RECIPE="${STATIC_RECIPE}" + CP_DESCRIPTION="static CP${STATIC_CP} (Dynamic CP disabled)" +fi + +if [[ ! -f "${NEMO_RL_ROOT}/${RECIPE}" ]]; then + echo "ERROR: recipe does not exist: ${NEMO_RL_ROOT}/${RECIPE}" >&2 + echo "The comparison recipes are local to the Dynamic-CP checkout." >&2 + exit 1 +fi + +CONTAINER="${CONTAINER:-/lustre/fs1/portfolios/coreai/projects/coreai_dlalgo_ci/nemo_rl_ci/sqsh_files/rl.nightly.sqsh}" +if [[ "${CONTAINER}" == /* && ! -e "${CONTAINER}" ]]; then + echo "ERROR: container does not exist: ${CONTAINER}" >&2 + echo "Set CONTAINER to a NeMo-RL image available on the compute nodes." >&2 + exit 1 +fi + +if [[ -z "${COMPARISON_NAME:-}" ]]; then + echo "ERROR: set COMPARISON_NAME once and reuse it for all six runs." >&2 + echo 'Example: export COMPARISON_NAME="cp-smoke-$(date +%Y%m%d-%H%M%S)"' >&2 + exit 2 +fi +NRL_MAX_STEPS="${NRL_MAX_STEPS:-3}" +if ! [[ "${NRL_MAX_STEPS}" =~ ^[1-9][0-9]*$ ]]; then + echo "ERROR: NRL_MAX_STEPS must be a positive integer." >&2 + exit 2 +fi + +DYNAMIC_CP_RESULTS_ROOT="${DYNAMIC_CP_RESULTS_ROOT:-/lustre/fsw/portfolios/coreai/users/${USER}/projects/nemo-rl-workspace/dynamic-cp-comparisons}" +RUN_DIR="${DYNAMIC_CP_RESULTS_ROOT}/${COMPARISON_NAME}/${MODEL}/${MODE}" +METRICS_DIR="${RUN_DIR}/metrics" +BASE_LOG_DIR="${RUN_DIR}/slurm" + +ENABLE_WANDB="${ENABLE_WANDB:-0}" +case "${ENABLE_WANDB}" in + 0) WANDB_ENABLED=false ;; + 1) + if [[ -z "${WANDB_API_KEY:-}" ]]; then + echo "ERROR: ENABLE_WANDB=1 requires WANDB_API_KEY." >&2 + exit 2 + fi + WANDB_ENABLED=true + ;; + *) + echo "ERROR: ENABLE_WANDB must be 0 or 1." >&2 + exit 2 + ;; +esac + +if [[ "${DRY_RUN:-0}" != "1" ]] && [[ -d "${RUN_DIR}" ]] && \ + [[ -n "$(find "${RUN_DIR}" -mindepth 1 -print -quit)" ]]; then + echo "ERROR: run directory is not empty: ${RUN_DIR}" >&2 + echo "Use a new COMPARISON_NAME so old and new measurements are not mixed." >&2 + exit 1 +fi +if [[ "${DRY_RUN:-0}" != "1" ]]; then + mkdir -p "${RUN_DIR}" "${BASE_LOG_DIR}" +fi + +COMMAND_PARTS=( + uv run python examples/run_grpo.py + --config "${RECIPE}" + "grpo.max_num_steps=${NRL_MAX_STEPS}" + checkpointing.enabled=false + "logger.log_dir=${METRICS_DIR}" + "logger.wandb_enabled=${WANDB_ENABLED}" + logger.tensorboard_enabled=true + logger.wandb.project=nemo-rl-cp-comparison + "logger.wandb.name=${COMPARISON_NAME}-${MODEL}-${MODE}" +) +printf -v COMMAND '%q ' "${COMMAND_PARTS[@]}" + +REQUIRED_MOUNTS="/lustre:/lustre,${NEMO_RL_ROOT}:${NEMO_RL_ROOT}" +MOUNTS="${MOUNTS:-${REQUIRED_MOUNTS}}" +if [[ -n "${EXTRA_MOUNTS:-}" ]]; then + MOUNTS="${MOUNTS},${EXTRA_MOUNTS}" +fi + +export CONTAINER MOUNTS COMMAND BASE_LOG_DIR +export GPUS_PER_NODE="${GPUS_PER_NODE:-4}" +export RAY_LOG_SYNC_FREQUENCY="${RAY_LOG_SYNC_FREQUENCY:-30}" +export HF_HOME="${HF_HOME:-/lustre/fsw/portfolios/coreai/users/${USER}/hf_home}" + +SLURM_ACCOUNT="${SLURM_ACCOUNT:-coreai_dlalgo_nemorl}" +SLURM_PARTITION="${SLURM_PARTITION:-batch}" +WALLTIME="${WALLTIME:-1:59:00}" +JOB_NAME="dcp-${MODEL}-${MODE}" +SBATCH_ARGS=( + --nodes="${NODES}" + --account="${SLURM_ACCOUNT}" + --partition="${SLURM_PARTITION}" + --time="${WALLTIME}" + --gres="gpu:${GPUS_PER_NODE}" + --exclusive + --mem=0 + --job-name="${JOB_NAME}" + --output="${RUN_DIR}/slurm-%j.out" +) +if [[ -n "${SEGMENT_SIZE}" ]]; then + SBATCH_ARGS+=(--segment="${SEGMENT_SIZE}") +fi +if [[ -n "${SLURM_QOS:-}" ]]; then + SBATCH_ARGS+=(--qos="${SLURM_QOS}") +fi + +cat < Date: Tue, 22 Sep 2026 21:27:10 -0700 Subject: [PATCH 7/7] fixed freq all reduce --- nemo_rl/models/megatron/data.py | 3 + nemo_rl/models/megatron/dynamic_cp.py | 270 +++++++++++------- nemo_rl/models/megatron/hybridep.py | 15 +- .../models/megatron/test_dynamic_cp_moe.py | 115 +++++++- .../models/megatron/test_hybridep_data.py | 41 +++ 5 files changed, 320 insertions(+), 124 deletions(-) diff --git a/nemo_rl/models/megatron/data.py b/nemo_rl/models/megatron/data.py index dcf515f5e33..48cc2d665f5 100644 --- a/nemo_rl/models/megatron/data.py +++ b/nemo_rl/models/megatron/data.py @@ -593,6 +593,9 @@ def process_microbatch( pad_packed_seq_to_multiple_of=pad_packed_seq_to_multiple_of, cp_rank=cp_rank, cp_size=cp_size, + # Dynamic CP's planner pads every task to the same + # HybridEP-aligned length on all participating ranks. + group_aligned=cp_context is not None, ) full_padding_mask = get_packed_seq_padding_mask( cu_seqlens=cu_seqlens, diff --git a/nemo_rl/models/megatron/dynamic_cp.py b/nemo_rl/models/megatron/dynamic_cp.py index 6e69a9de6e1..952e80449fb 100644 --- a/nemo_rl/models/megatron/dynamic_cp.py +++ b/nemo_rl/models/megatron/dynamic_cp.py @@ -26,6 +26,7 @@ _ROUTER_CONFIG_BASELINES: dict[int, tuple[Any, Any]] = {} _DYNAMIC_MTP_METRICS: dict[str, torch.Tensor] = {} _ACTIVE_BIND_TARGETS: dict[int, "_BindTargets"] = {} +_ACTIVE_BIND_SIGNATURES: dict[int, tuple[int, int, int, bool]] = {} @dataclass(frozen=True) @@ -351,24 +352,30 @@ def get_dynamic_mtp_metrics( if "loss_sums" not in _DYNAMIC_MTP_METRICS: return {} try: - for value in _DYNAMIC_MTP_METRICS.values(): - torch.distributed.all_reduce( - value, op=torch.distributed.ReduceOp.SUM, group=parallel_group + totals = torch.stack( + tuple( + _DYNAMIC_MTP_METRICS[name] + for name in ( + "loss_sums", + "loss_token_counts", + "correct_values", + "total_values", + ) ) - losses = _DYNAMIC_MTP_METRICS["loss_sums"] / _DYNAMIC_MTP_METRICS[ - "loss_token_counts" - ].clamp(min=1) - acceptance = ( - _DYNAMIC_MTP_METRICS["correct_values"] - / _DYNAMIC_MTP_METRICS["total_values"].clamp(min=1) - * 100.0 ) - metrics: dict[str, float] = {} - for index in range(losses.numel()): - metrics[f"mtp_{index + 1}_loss"] = float(losses[index].item()) - metrics[f"mtp_{index + 1}_acceptance_rate"] = float( - acceptance[index].item() + if parallel_group.size() > 1: + torch.distributed.all_reduce( + totals, op=torch.distributed.ReduceOp.SUM, group=parallel_group ) + losses = totals[0] / totals[1].clamp(min=1) + acceptance = totals[2] / totals[3].clamp(min=1) * 100.0 + loss_values, acceptance_values = torch.stack((losses, acceptance)).tolist() + metrics: dict[str, float] = {} + for index, (loss, rate) in enumerate( + zip(loss_values, acceptance_values, strict=True) + ): + metrics[f"mtp_{index + 1}_loss"] = float(loss) + metrics[f"mtp_{index + 1}_acceptance_rate"] = float(rate) return metrics finally: _DYNAMIC_MTP_METRICS.clear() @@ -594,9 +601,9 @@ def _dynamic_attach_and_log_load_balancing_loss( MCore's pinned implementation multiplies by ``local_tokens * tp_cp_group.size()``. That is only equal to the task token count for - equally populated fixed shards. The router pre-hook has already reduced - the current mask over the runtime TP*CP group, so attach that exact count. - Logging intentionally retains the unscaled aux value. + equally populated fixed shards. The task setup (or the MTP mask hook) has + already reduced the current mask over the runtime TP*CP group, so attach + that exact count. Logging intentionally retains the unscaled aux value. """ from megatron.core.transformer.moe.moe_logging import get_moe_metrics_tracker from megatron.core.transformer.moe.moe_utils import MoEAuxLossAutoScaler @@ -628,15 +635,9 @@ def _dynamic_attach_and_log_load_balancing_loss( if self.calculate_per_token_loss: task_tokens = getattr(self, "_nemo_dynamic_aux_scale_tokens", None) if task_tokens is None: - local_tokens = ( - valid_token_count - if valid_token_count is not None - else activation.shape[0] - ) - task_tokens = ( - torch.as_tensor(local_tokens, device=activation.device).detach().clone() + raise RuntimeError( + "Dynamic CP MoE token scaling was not prepared before router forward" ) - torch.distributed.all_reduce(task_tokens, group=self.tp_cp_group) return MoEAuxLossAutoScaler.apply(activation, aux_loss * task_tokens) return MoEAuxLossAutoScaler.apply(activation, aux_loss) @@ -663,6 +664,8 @@ def _patch_hybrid_mtp_padding_masks( the mask received by HybridModel already has the ordering needed by MTP. """ padding_masks: dict[int, torch.Tensor | None] = {} + mtp_token_counts: dict[int, tuple[torch.Tensor, Any, torch.Tensor]] = {} + active_task_marker: object | None = None handles: list[Any] = [] modules = modules or tuple(model.modules()) mtp_blocks: list[torch.nn.Module] = [] @@ -722,10 +725,12 @@ def prepare_mtp_validity_mask( def prepare_router_padding_mask( module: Router, args: tuple[Any, ...], kwargs: dict[str, Any] ) -> tuple[tuple[Any, ...], dict[str, Any]]: + nonlocal active_task_marker padding_mask = kwargs.get("padding_mask") if padding_mask is None and len(args) > 1: padding_mask = args[1] - if padding_mask is not None and id(module) in mtp_router_ids: + validity_mask = padding_mask + if padding_mask is not None: padding_mask = ~padding_mask.to(dtype=torch.bool) if len(args) > 1: args = (args[0], padding_mask, *args[2:]) @@ -735,18 +740,44 @@ def prepare_router_padding_mask( if ( module.training and torch.is_grad_enabled() - and getattr(module.config, "calculate_per_token_loss", False) + and getattr(module, "calculate_per_token_loss", False) and _has_positive_coefficient(module.config.moe_aux_loss_coeff) ): - if padding_mask is None: + if validity_mask is None: raise ValueError( "Dynamic CP MoE aux loss requires a packed padding mask" ) - group_tokens = (~padding_mask).sum().detach().clone() - torch.distributed.all_reduce(group_tokens, group=module.tp_cp_group) + task_marker = getattr(module, "_nemo_dynamic_moe_task_marker", None) + if task_marker is None: + raise RuntimeError( + "Dynamic CP MoE token scaling was not configured for this task" + ) + if task_marker is not active_task_marker: + mtp_token_counts.clear() + active_task_marker = task_marker + + cache_key = id(validity_mask) + cached = mtp_token_counts.get(cache_key) + if ( + cached is not None + and cached[0] is validity_mask + and cached[1] is module.tp_cp_group + ): + group_tokens = cached[2] + else: + group_tokens = validity_mask.sum().detach() + if module.tp_cp_group.size() > 1: + torch.distributed.all_reduce( + group_tokens, + op=torch.distributed.ReduceOp.SUM, + group=module.tp_cp_group, + ) + mtp_token_counts[cache_key] = ( + validity_mask, + module.tp_cp_group, + group_tokens, + ) module._nemo_dynamic_aux_scale_tokens = group_tokens - else: - module._nemo_dynamic_aux_scale_tokens = None return args, kwargs patched_classes: list[tuple[type[Any], bool, Any]] = [] @@ -773,7 +804,7 @@ def prepare_router_padding_mask( mtp.register_forward_pre_hook(prepare_mtp_validity_mask, with_kwargs=True) ) for module in modules: - if isinstance(module, Router): + if isinstance(module, Router) and id(module) in mtp_router_ids: handles.append( module.register_forward_pre_hook( prepare_router_padding_mask, with_kwargs=True @@ -789,6 +820,10 @@ def prepare_router_padding_mask( module, "_nemo_dynamic_aux_scale_tokens" ): del module._nemo_dynamic_aux_scale_tokens + if isinstance(module, Router) and hasattr( + module, "_nemo_dynamic_moe_task_marker" + ): + del module._nemo_dynamic_moe_task_marker for router_class, had_direct_method, original_method in patched_classes: if had_direct_method: router_class.attach_and_log_load_balancing_loss = original_method @@ -882,6 +917,7 @@ def preserve_attention_cp_groups(model: torch.nn.Module) -> Iterator[None]: raise finally: _ACTIVE_BIND_TARGETS.pop(model_id, None) + _ACTIVE_BIND_SIGNATURES.pop(model_id, None) for module, collection in saved_collections: module.pg_collection = collection for module, cp in saved_mamba: @@ -948,14 +984,7 @@ def _has_positive_coefficient(value: Any) -> bool: def configure_dynamic_moe_loss_scaling( model: torch.nn.Module, padding_mask: torch.Tensor | None ) -> None: - """Validate inputs for the temporary exact-token router attachment. - - ``preserve_attention_cp_groups`` installs a router pre-hook which reduces - the current mask and an attachment shim which uses that exact TP*CP token - count. Keeping the check here fails before entering a router collective if - a caller forgot the packed padding mask. The worker's ordinary - ``1/global_valid_tokens`` MoE scale is therefore sufficient. - """ + """Reduce a task token count once and share it across all MoE routers.""" routers = _bind_targets(model).routers if not routers: return @@ -963,7 +992,8 @@ def configure_dynamic_moe_loss_scaling( if not model.training or not torch.is_grad_enabled(): return if not any( - _has_positive_coefficient(router.config.moe_aux_loss_coeff) + getattr(router, "calculate_per_token_loss", False) + and _has_positive_coefficient(router.config.moe_aux_loss_coeff) for router in routers ): return @@ -973,6 +1003,19 @@ def configure_dynamic_moe_loss_scaling( if any(router.tp_cp_group is not routers[0].tp_cp_group for router in routers[1:]): raise ValueError("Dynamic CP routers disagree on the active TP*CP group") + task_marker = object() + task_tokens = (~padding_mask.to(dtype=torch.bool)).sum().detach() + tp_cp_group = routers[0].tp_cp_group + if tp_cp_group.size() > 1: + torch.distributed.all_reduce( + task_tokens, + op=torch.distributed.ReduceOp.SUM, + group=tp_cp_group, + ) + for router in routers: + router._nemo_dynamic_moe_task_marker = task_marker + router._nemo_dynamic_aux_scale_tokens = task_tokens + def bind_attention_cp_group(model: torch.nn.Module, packed_seq_params: Any) -> Any: """Bind attention, MoE router and SSM modules to the active CP task. @@ -1001,39 +1044,48 @@ def bind_attention_cp_group(model: torch.nn.Module, packed_seq_params: Any) -> A raise ValueError("Dynamic MoE TP*CP group has the wrong size") padding_only = bool(getattr(packed_seq_params, "dynamic_cp_padding_only", False)) targets = _bind_targets(model) - for module in targets.pg_collections: - module.pg_collection.cp = group - if hasattr(module.pg_collection, "tp_cp"): - module.pg_collection.tp_cp = tp_cp_group - for module in targets.direct_groups: - module.cp_group = group - if hasattr(module, "tp_cp_group"): + model_id = id(model) + binding_signature = (context.size, id(group), id(tp_cp_group), padding_only) + binding_cache_active = model_id in _ACTIVE_BIND_TARGETS + if ( + not binding_cache_active + or _ACTIVE_BIND_SIGNATURES.get(model_id) != binding_signature + ): + for module in targets.pg_collections: + module.pg_collection.cp = group + if hasattr(module.pg_collection, "tp_cp"): + module.pg_collection.tp_cp = tp_cp_group + for module in targets.direct_groups: + module.cp_group = group + if hasattr(module, "tp_cp_group"): + module.tp_cp_group = tp_cp_group + for module in targets.hybrid_stacks: + _bind_hybrid_stack_layout(module, group=group, tp_cp_group=tp_cp_group) + for module in targets.routers: + module.cp_group = group module.tp_cp_group = tp_cp_group - for module in targets.hybrid_stacks: - _bind_hybrid_stack_layout(module, group=group, tp_cp_group=tp_cp_group) - for module in targets.routers: - module.cp_group = group - module.tp_cp_group = tp_cp_group - _bind_router_config(module, padding_only=padding_only) - for module in targets.mamba_mixers: - _rebuild_mamba_cp(module, group) - for module in targets.gated_delta_products: - _rebuild_gdp_cp(module, group) - for module in targets.gated_delta_nets: - baseline_size = module.cp_size - baseline_split = getattr(module, "feat_dim_split", None) - module.cp_size = context.size - if baseline_split is not None: - scaled_split = [] - for value in baseline_split: - numerator = value * baseline_size - if numerator % context.size: - raise ValueError( - "GatedDeltaNet projection dimensions are not divisible " - f"by runtime CP={context.size}" - ) - scaled_split.append(numerator // context.size) - module.feat_dim_split = tuple(scaled_split) + _bind_router_config(module, padding_only=padding_only) + for module in targets.mamba_mixers: + _rebuild_mamba_cp(module, group) + for module in targets.gated_delta_products: + _rebuild_gdp_cp(module, group) + for module in targets.gated_delta_nets: + baseline_size = module.cp_size + baseline_split = getattr(module, "feat_dim_split", None) + module.cp_size = context.size + if baseline_split is not None: + scaled_split = [] + for value in baseline_split: + numerator = value * baseline_size + if numerator % context.size: + raise ValueError( + "GatedDeltaNet projection dimensions are not divisible " + f"by runtime CP={context.size}" + ) + scaled_split.append(numerator // context.size) + module.feat_dim_split = tuple(scaled_split) + if binding_cache_active: + _ACTIVE_BIND_SIGNATURES[model_id] = binding_signature model_packed = copy(packed_seq_params) model_packed.cp_group = group return model_packed @@ -1059,43 +1111,51 @@ def planned_microbatches( or domain.rank() != plan.lane ): raise ValueError("Ray's DP*CP lane map disagrees with initialized MCore groups") + expert_group = parallel_state.get_expert_tensor_and_model_parallel_group() + expert_ranks = set(torch.distributed.get_process_group_ranks(expert_group)) + tp_size = parallel_state.get_tensor_model_parallel_world_size() + runtime_contexts: dict[tuple[int, int], RuntimeCPContext] = {} for group_index, rank_group in enumerate(step.groups): if not rank_group.assignments: raise ValueError("Every CP synchronization group needs one local task") for task_index, assignment in enumerate(rank_group.assignments): size = assignment.cp_size - group = ( - parallel_state.get_hybrid_data_context_parallel_groups(group_size=size) - if size > 1 - else None - ) - rank = plan.lane - assignment.lane_start - if group is not None: - members = expected[assignment.lane_start : assignment.lane_start + size] - if ( - group.size() != size - or group.rank() != rank - or torch.distributed.get_process_group_ranks(group) != members - ): - raise ValueError( - "Active CP group disagrees with the driver's assignment" + context_key = (assignment.lane_start, size) + context = runtime_contexts.get(context_key) + if context is None: + group = ( + parallel_state.get_hybrid_data_context_parallel_groups( + group_size=size ) - expert_group = parallel_state.get_expert_tensor_and_model_parallel_group() - tp_size = parallel_state.get_tensor_model_parallel_world_size() - task_ranks = { - base + offset - for base in plan.lane_ranks[ - assignment.lane_start : assignment.lane_start + size - ] - for offset in range(tp_size) - } - if not set( - torch.distributed.get_process_group_ranks(expert_group) - ).issubset(task_ranks): - raise ValueError( - "Joint expert TP*EP group crosses dynamic CP task boundaries" + if size > 1 + else None ) - context = RuntimeCPContext(size=size, rank=rank, group=group) + rank = plan.lane - assignment.lane_start + if group is not None: + members = expected[ + assignment.lane_start : assignment.lane_start + size + ] + if ( + group.size() != size + or group.rank() != rank + or torch.distributed.get_process_group_ranks(group) != members + ): + raise ValueError( + "Active CP group disagrees with the driver's assignment" + ) + task_ranks = { + base + offset + for base in plan.lane_ranks[ + assignment.lane_start : assignment.lane_start + size + ] + for offset in range(tp_size) + } + if not expert_ranks.issubset(task_ranks): + raise ValueError( + "Joint expert TP*EP group crosses dynamic CP task boundaries" + ) + context = RuntimeCPContext(size=size, rank=rank, group=group) + runtime_contexts[context_key] = context if assignment.sample_indices: batch = data.select_indices(list(assignment.sample_indices)).to("cuda") else: diff --git a/nemo_rl/models/megatron/hybridep.py b/nemo_rl/models/megatron/hybridep.py index 245584ed1c7..ca40381ec46 100644 --- a/nemo_rl/models/megatron/hybridep.py +++ b/nemo_rl/models/megatron/hybridep.py @@ -146,6 +146,8 @@ def pad_packed_seq_for_hybridep( pad_packed_seq_to_multiple_of: int, cp_rank: int, cp_size: int, + *, + group_aligned: bool = False, ) -> tuple[torch.Tensor, torch.Tensor, PackedSeqParams, torch.Tensor]: """Align packed inputs once, before model collectives can overlap.""" local_seq_len = input_ids_cp_sharded.shape[1] @@ -153,11 +155,14 @@ def pad_packed_seq_for_hybridep( pad_packed_seq_to_multiple_of, cp_size, ) - target_seq_len = _get_hybridep_aligned_seq_len( - local_seq_len, - local_pad_multiple, - input_ids_cp_sharded.device, - ) + if group_aligned: + target_seq_len = _round_up_to_multiple(local_seq_len, local_pad_multiple) + else: + target_seq_len = _get_hybridep_aligned_seq_len( + local_seq_len, + local_pad_multiple, + input_ids_cp_sharded.device, + ) if target_seq_len == local_seq_len: return input_ids, input_ids_cp_sharded, packed_seq_params, cu_seqlens_padded diff --git a/tests/unit/models/megatron/test_dynamic_cp_moe.py b/tests/unit/models/megatron/test_dynamic_cp_moe.py index 86780cb719c..77990149d5e 100644 --- a/tests/unit/models/megatron/test_dynamic_cp_moe.py +++ b/tests/unit/models/megatron/test_dynamic_cp_moe.py @@ -123,6 +123,9 @@ def __init__(self, group, config): self.cp_group = group self.tp_cp_group = group self.config = config + self.calculate_per_token_loss = getattr( + config, "calculate_per_token_loss", False + ) def test_dynamic_binding_updates_router_and_ssm_then_restores(monkeypatch): @@ -176,6 +179,13 @@ def test_dynamic_binding_updates_router_and_ssm_then_restores(monkeypatch): assert config.moe_aux_loss_coeff == [0.0, 0.0] assert config.moe_z_loss_coeff is None + # Consecutive tasks in one partition reuse all runtime bindings. + bound_mamba_helper = model.mamba.cp + bound_gdp_helper = model.gdp.cp + dynamic_cp.bind_attention_cp_group(model, packed) + assert model.mamba.cp is bound_mamba_helper + assert model.gdp.cp is bound_gdp_helper + assert model.router.cp_group is original_tp_cp assert model.router.tp_cp_group is original_tp_cp assert model.router2.cp_group is original_tp_cp @@ -252,29 +262,41 @@ def test_dynamic_binding_setup_failure_cleans_global_state(monkeypatch): @pytest.mark.parametrize("active_size", [1, 4]) -def test_dynamic_moe_scaling_is_applied_by_router_attachment(monkeypatch, active_size): +def test_dynamic_moe_scaling_reduces_once_per_task(monkeypatch, active_size): from nemo_rl.models.megatron import dynamic_cp active_tp_cp = _Group(active_size) - config = SimpleNamespace(moe_aux_loss_coeff=0.1, moe_z_loss_coeff=0.02) + config = SimpleNamespace( + calculate_per_token_loss=True, + moe_aux_loss_coeff=0.1, + moe_z_loss_coeff=0.02, + ) model = torch.nn.Module() model.add_module("router", _Router(active_tp_cp, config)) + model.add_module("router2", _Router(active_tp_cp, config)) padding_mask = torch.tensor([[False, False, False, True]]) + reductions = [] monkeypatch.setattr(dynamic_cp, "Router", _Router) - monkeypatch.setattr( - dynamic_cp.torch.distributed, - "all_reduce", - lambda *_args, **_kwargs: pytest.fail( - "pre-forward validation performed a token reduction" - ), - ) + def _all_reduce(count, *, op, group): + reductions.append((op, group)) + count.fill_(10) + + monkeypatch.setattr(dynamic_cp.torch.distributed, "all_reduce", _all_reduce) with dynamic_cp.preserve_attention_cp_groups(model): dynamic_cp.configure_dynamic_moe_loss_scaling(model, padding_mask) + expected_tokens = 3 if active_size == 1 else 10 + assert model.router._nemo_dynamic_aux_scale_tokens.item() == expected_tokens + assert ( + model.router._nemo_dynamic_aux_scale_tokens + is model.router2._nemo_dynamic_aux_scale_tokens + ) + assert len(reductions) == (1 if active_size > 1 else 0) assert config.moe_z_loss_coeff == 0.02 + assert not hasattr(model.router, "_nemo_dynamic_aux_scale_tokens") assert config.moe_aux_loss_coeff == 0.1 assert config.moe_z_loss_coeff == 0.02 @@ -292,6 +314,13 @@ def _apply(_activation, aux_loss): monkeypatch.setattr(moe_logging, "get_moe_metrics_tracker", lambda: tracker) monkeypatch.setattr(moe_utils.MoEAuxLossAutoScaler, "apply", _apply) + monkeypatch.setattr( + dynamic_cp.torch.distributed, + "all_reduce", + lambda *_args, **_kwargs: pytest.fail( + "router attachment performed a fallback reduction" + ), + ) router = SimpleNamespace( is_mtp_layer=False, @@ -416,16 +445,18 @@ def test_dynamic_mtp_metrics_are_token_weighted(monkeypatch): layer_number=0, num_layers=1, ) - monkeypatch.setattr( - dynamic_cp.torch.distributed, - "all_reduce", - lambda _value, *, op, group: None, - ) + reductions = [] + + def _all_reduce(_value, *, op, group): + reductions.append((op, group)) + + monkeypatch.setattr(dynamic_cp.torch.distributed, "all_reduce", _all_reduce) metrics = dynamic_cp.get_dynamic_mtp_metrics(parallel_group=_Group(4)) assert metrics["mtp_1_loss"] == pytest.approx(3.0) assert metrics["mtp_1_acceptance_rate"] == pytest.approx(60.0) + assert len(reductions) == 1 assert dynamic_cp._DYNAMIC_MTP_METRICS == {} @@ -481,6 +512,62 @@ def forward(self, *, padding_mask=None): assert model(padding_mask=padding_mask) is None +def test_dynamic_mtp_routers_share_one_count_reduction(monkeypatch): + from nemo_rl.models.megatron import dynamic_cp + + group = _Group(2) + config = SimpleNamespace( + calculate_per_token_loss=True, + moe_aux_loss_coeff=0.1, + moe_z_loss_coeff=None, + ) + + class _MTPRouter(_Router): + def forward(self, hidden_states, padding_mask=None): + return hidden_states + + class _MTP(torch.nn.Module): + def __init__(self): + super().__init__() + self.router1 = _MTPRouter(group, config) + self.router2 = _MTPRouter(group, config) + + def forward(self, *, padding_mask=None): + value = self.router1(torch.ones(1), padding_mask=padding_mask) + return self.router2(value, padding_mask=padding_mask) + + class _HybridModel(torch.nn.Module): + def __init__(self): + super().__init__() + self.mtp = _MTP() + + def forward(self, *, padding_mask=None): + return self.mtp() + + _HybridModel.__module__ = "megatron.core.models.hybrid.hybrid_model" + model = _HybridModel() + padding_mask = torch.tensor([[False, True]]) + reductions = [] + + def _all_reduce(count, *, op, group): + reductions.append((op, group)) + count.mul_(2) + + monkeypatch.setattr(dynamic_cp, "Router", _MTPRouter) + monkeypatch.setattr(dynamic_cp.torch.distributed, "all_reduce", _all_reduce) + + with dynamic_cp._patch_hybrid_mtp_padding_masks(model): + dynamic_cp.configure_dynamic_moe_loss_scaling(model, padding_mask) + model(padding_mask=padding_mask) + assert ( + model.mtp.router1._nemo_dynamic_aux_scale_tokens + is model.mtp.router2._nemo_dynamic_aux_scale_tokens + ) + + # One main-task reduction plus one shared reduction for the shifted MTP mask. + assert len(reductions) == 2 + + def test_dynamic_model_validation_rejects_unmerged_mla_support(): from nemo_rl.models.megatron import dynamic_cp diff --git a/tests/unit/models/megatron/test_hybridep_data.py b/tests/unit/models/megatron/test_hybridep_data.py index 92b4abd6ecb..54c9722615c 100644 --- a/tests/unit/models/megatron/test_hybridep_data.py +++ b/tests/unit/models/megatron/test_hybridep_data.py @@ -312,6 +312,47 @@ def test_hybridep_prepadding_returns_original_objects_when_already_aligned() -> assert result[3] is cu_seqlens_padded +@pytest.mark.mcore +def test_dynamic_cp_hybridep_alignment_skips_group_reduction() -> None: + from megatron.core.packed_seq_params import PackedSeqParams + + from nemo_rl.models.megatron import hybridep + + input_ids = torch.arange(1, 13).view(1, 12) + cu_seqlens_padded = torch.tensor([0, 12], dtype=torch.int32) + packed_seq_params = PackedSeqParams( + cu_seqlens_q=cu_seqlens_padded, + cu_seqlens_kv=cu_seqlens_padded, + cu_seqlens_q_padded=cu_seqlens_padded, + cu_seqlens_kv_padded=cu_seqlens_padded, + max_seqlen_q=12, + max_seqlen_kv=12, + qkv_format="thd", + total_tokens=12, + ) + + with patch.object( + hybridep.torch.distributed, + "all_reduce", + side_effect=AssertionError("planner-aligned input performed a reduction"), + ): + result = hybridep.pad_packed_seq_for_hybridep( + input_ids=input_ids, + input_ids_cp_sharded=input_ids, + packed_seq_params=packed_seq_params, + cu_seqlens_padded=cu_seqlens_padded, + pad_packed_seq_to_multiple_of=8, + cp_rank=0, + cp_size=1, + group_aligned=True, + ) + + assert result[0].shape == (1, 16) + assert result[1].shape == (1, 16) + assert result[2].total_tokens == 16 + assert torch.equal(result[3], torch.tensor([0, 16])) + + @pytest.mark.mcore @patch("nemo_rl.models.megatron.data.get_context_parallel_rank", return_value=0) @patch("nemo_rl.models.megatron.data.get_context_parallel_world_size", return_value=2)