Skip to content
Original file line number Diff line number Diff line change
Expand Up @@ -14,10 +14,6 @@ policy:
megatron_cfg:
context_parallel_size: 1
moe_router_dtype: fp32
freeze_config:
freeze_vision_model: true
freeze_vision_projection: true
freeze_language_model: false
fp8_cfg:
enabled: true
optimizer:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,10 @@ policy:
context_parallel_size: 2
sequence_parallel: true
expert_model_parallel_size: 16
freeze_config:
freeze_vision_model: true
freeze_vision_projection: true
freeze_language_model: false
apply_rope_fusion: false
activation_checkpointing: true
defer_fp32_logits: true
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,10 @@ policy:
moe_token_dispatcher_type: allgather
apply_rope_fusion: false
defer_fp32_logits: true
freeze_config:
freeze_vision_model: true
freeze_vision_projection: true
freeze_language_model: false
optimizer:
lr: 1.0e-06
min_lr: 1.0e-06
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -18,10 +18,6 @@ policy:
make_sequence_length_divisible_by: 32
megatron_cfg:
moe_router_dtype: fp32
freeze_config:
freeze_vision_model: true
freeze_vision_projection: true
freeze_language_model: false
fp8_cfg:
enabled: true
optimizer:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,10 @@ policy:
apply_rope_fusion: false
activation_checkpointing: true
defer_fp32_logits: true
freeze_config:
freeze_vision_model: true
freeze_vision_projection: true
freeze_language_model: false
generation:
vllm_cfg:
tensor_parallel_size: 4
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
from pathlib import Path
from typing import Any

import pytest
from omegaconf import OmegaConf

from nemo_rl.utils.config import (
Expand All @@ -24,11 +25,18 @@

PROJECT_ROOT = Path(__file__).resolve().parents[4]
RECIPE_NAME = "grpo-qwen3.5-35ba3b-6n4g-async-1off-bf16-trtllm.yaml"
TEXT_RECIPE_NAMES = (
RECIPE_NAME,
"grpo-qwen3.5-9b-1n8g-megatron.yaml",
"grpo-qwen3.5-9b-1n8g-megatron-fp8.yaml",
"grpo-qwen3.5-35ba3b-2n8g-megatron-ep16tp2-fp8.yaml",
"grpo-qwen3.5-397ba17b-32n8g-megatron.v2.yaml",
)


def _load_recipe() -> dict[str, Any]:
def _load_recipe(recipe_name: str = RECIPE_NAME) -> dict[str, Any]:
register_omegaconf_resolvers()
recipe_path = PROJECT_ROOT / "examples/configs/recipes/llm" / RECIPE_NAME
recipe_path = PROJECT_ROOT / "examples/configs/recipes/llm" / recipe_name
recipe = OmegaConf.to_container(
load_config_with_inheritance(recipe_path), resolve=True
)
Expand Down Expand Up @@ -69,3 +77,16 @@ def test_qwen35_bf16_trtllm_recipe_uses_supported_expert_layout() -> None:
"moe_backend": "flashinfer_trtllm",
"expert_placement_strategy": "linear",
}


@pytest.mark.parametrize("recipe_name", TEXT_RECIPE_NAMES)
def test_qwen35_text_recipe_freezes_unused_vision_modules(
recipe_name: str,
) -> None:
recipe = _load_recipe(recipe_name)

assert recipe["policy"]["megatron_cfg"]["freeze_config"] == {
"freeze_vision_model": True,
"freeze_vision_projection": True,
"freeze_language_model": False,
}
Loading