Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -0,0 +1,38 @@
defaults: ../../grpo_math_1B.yaml
grpo:
num_prompts_per_step: 2
num_generations_per_prompt: 8
checkpointing:
checkpoint_dir: results/grpo-nanov3-30BA3B-2n8g-megatron-pack-cp
policy:
model_name: nvidia/NVIDIA-Nemotron-3-Nano-30B-A3B-Base-BF16
tokenizer:
name: nvidia/NVIDIA-Nemotron-3-Nano-30B-A3B-BF16
train_global_batch_size: 16
train_micro_batch_size: 1
logprob_batch_size: 1
max_total_sequence_length: 8192
make_sequence_length_divisible_by: 8
dtensor_cfg:
enabled: false
megatron_cfg:
enabled: true
bias_activation_fusion: false
tensor_model_parallel_size: 2
pipeline_model_parallel_size: 2
context_parallel_size: 2
expert_model_parallel_size: 8
sequence_parallel: true
generation:
vllm_cfg:
tensor_parallel_size: 4
gpu_memory_utilization: 0.7
logger:
wandb_enabled: true
tensorboard_enabled: true
wandb:
project: nemo-rl
name: grpo-nanov3-30BA3B-2n8g-megatron-pack-cp
cluster:
gpus_per_node: 8
num_nodes: 2
3 changes: 3 additions & 0 deletions nemo_rl/models/megatron/data.py
Original file line number Diff line number Diff line change
Expand Up @@ -509,6 +509,8 @@ def _pack_sequences_for_megatron(
if cu_seqlens_padded is None:
cu_seqlens_padded = cu_seqlens.clone()

# total_tokens is required for PackedSeqParams.__post_init__ to build
# seq_idx, which Mamba uses to reset SSM state at sample boundaries.
packed_seq_params = PackedSeqParams(
cu_seqlens_q=cu_seqlens_padded,
cu_seqlens_kv=cu_seqlens_padded,
Expand All @@ -517,6 +519,7 @@ def _pack_sequences_for_megatron(
max_seqlen_q=int(max_seqlen),
max_seqlen_kv=int(max_seqlen),
qkv_format="thd",
total_tokens=packed_input_ids.shape[1],
)

return (
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,44 @@
#!/bin/bash
SCRIPT_DIR=$( cd -- "$( dirname -- "${BASH_SOURCE[0]}" )" &> /dev/null && pwd)
source $SCRIPT_DIR/common.env

# ===== BEGIN CONFIG =====
NUM_NODES=2
GPUS_PER_NODE=8
STEPS_PER_RUN=3
MAX_STEPS=3
NUM_RUNS=$(( (MAX_STEPS + STEPS_PER_RUN - 1) / STEPS_PER_RUN )) # Round up
NUM_MINUTES=30
# ===== END CONFIG =====

exit_if_max_steps_reached

# Run the experiment
cd $PROJECT_ROOT
uv run examples/run_grpo.py \
--config $CONFIG_PATH \
grpo.max_num_steps=$MAX_STEPS \
logger.log_dir=$LOG_DIR \
logger.wandb_enabled=True \
logger.wandb.project=nemo-rl \
logger.wandb.name=$EXP_NAME \
logger.monitor_gpus=True \
logger.tensorboard_enabled=True \
checkpointing.enabled=True \
checkpointing.checkpoint_dir=$CKPT_DIR \
$@ \
2>&1 | tee $RUN_LOG

# Convert tensorboard logs to json
uv run tests/json_dump_tb_logs.py $LOG_DIR --output_path $JSON_METRICS

# Only run metrics if the target step is reached
if [[ $(jq 'to_entries | .[] | select(.key == "train/loss") | .value | keys | map(tonumber) | max' $JSON_METRICS) -ge $MAX_STEPS ]]; then
uv run tests/check_metrics.py $JSON_METRICS \
'median(data["train/token_mult_prob_error"]) < 1.1' \
'mean(data["train/gen_kl_error"]) < 0.02' \
'max(data["train/reward"]) > 0.05'

# Clean up checkpoint directory after successful run to save space.
rm -rf "$CKPT_DIR"
fi
1 change: 1 addition & 0 deletions tests/test_suites/nightly.txt
Original file line number Diff line number Diff line change
Expand Up @@ -84,6 +84,7 @@ tests/test_suites/llm/grpo-nano-v2-12b-2n8g-fsdp2tp1.sh
tests/test_suites/llm/grpo-nanov3-30BA3B-2n8g-fsdp2.sh
tests/test_suites/llm/grpo-nanov3-30BA3B-2n8g-fsdp2-lora.sh
tests/test_suites/llm/grpo-nanov3-30BA3B-2n8g-megatron-lora.sh
tests/test_suites/llm/grpo-nanov3-30BA3B-2n8g-megatron-pack-cp.sh

# Lora
tests/test_suites/llm/grpo-qwen3-8B-base-1n8g-fsdp2-lora.sh
Expand Down
35 changes: 35 additions & 0 deletions tests/unit/models/megatron/megatron_data_actors.py
Original file line number Diff line number Diff line change
Expand Up @@ -497,6 +497,25 @@ def _test_context_parallel(
"error": f"CP wrong qkv_format: expected 'thd', got {packed_seq_params.qkv_format}",
}

# Mamba SSM state reset relies on packed_seq_params.seq_idx, which
# __post_init__ only builds when total_tokens is set.
if packed_seq_params.total_tokens != expected_total_tokens:
return {
"success": False,
"error": f"CP packed_seq_params.total_tokens mismatch: expected {expected_total_tokens}, got {packed_seq_params.total_tokens}",
}
if packed_seq_params.seq_idx is None:
return {
"success": False,
"error": "CP packed_seq_params.seq_idx is None",
}
expected_seq_idx_len = int(cu_seqlens_padded[-1].item())
if packed_seq_params.seq_idx.shape != (1, expected_seq_idx_len):
return {
"success": False,
"error": f"CP packed_seq_params.seq_idx shape mismatch: expected (1, {expected_seq_idx_len}), got {tuple(packed_seq_params.seq_idx.shape)}",
}

# Test 2: CP packing with full sequence padding
pad_full_seq_to = (batch_size * seq_len) + 8 # Add some padding
(
Expand Down Expand Up @@ -532,6 +551,22 @@ def _test_context_parallel(
"error": f"CP full padding cu_seqlens_padded mismatch: expected {expected_cu_seqlens_padded_full}, got {cu_seqlens_padded_full}",
}

if packed_seq_params_full.total_tokens != expected_tokens_per_rank_full:
return {
"success": False,
"error": f"CP (full pad) packed_seq_params.total_tokens mismatch: expected {expected_tokens_per_rank_full}, got {packed_seq_params_full.total_tokens}",
}
if packed_seq_params_full.seq_idx is None:
return {
"success": False,
"error": "CP (full pad) packed_seq_params.seq_idx is None",
}
if packed_seq_params_full.seq_idx.shape != (1, pad_full_seq_to):
return {
"success": False,
"error": f"CP (full pad) packed_seq_params.seq_idx shape mismatch: expected (1, {pad_full_seq_to}), got {tuple(packed_seq_params_full.seq_idx.shape)}",
}

correct_ids_0 = torch.tensor(
[0, 1, 2, 3, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 0, 0, 0, 0, 0, 0],
device="cuda",
Expand Down
Loading