From 8990b258756b018c5230b707e9cc875e77cfb5f9 Mon Sep 17 00:00:00 2001 From: Hemil Desai Date: Wed, 24 Dec 2025 02:32:37 +0000 Subject: [PATCH 01/12] feat: RL support for custom moe models in dtensor v2 Signed-off-by: Hemil Desai --- ...grpo-moonlight-16b-automodel-1n8g-ep8.yaml | 64 ++++ .../workers/dtensor_policy_worker_v2.py | 305 +++++++++++------- .../grpo-moonlight-16b-automodel-1n8g-ep8.sh | 42 +++ 3 files changed, 290 insertions(+), 121 deletions(-) create mode 100644 examples/configs/recipes/llm/grpo-moonlight-16b-automodel-1n8g-ep8.yaml create mode 100755 tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh diff --git a/examples/configs/recipes/llm/grpo-moonlight-16b-automodel-1n8g-ep8.yaml b/examples/configs/recipes/llm/grpo-moonlight-16b-automodel-1n8g-ep8.yaml new file mode 100644 index 00000000000..77bc2db8b92 --- /dev/null +++ b/examples/configs/recipes/llm/grpo-moonlight-16b-automodel-1n8g-ep8.yaml @@ -0,0 +1,64 @@ +defaults: ../../grpo_math_1B.yaml +grpo: + num_prompts_per_step: 64 + num_generations_per_prompt: 32 + max_num_steps: 30 +checkpointing: + checkpoint_dir: results/grpo-moonlight-16b-a3b-instruct-automodel-1n8g-ep8 +policy: + model_name: moonshotai/Moonlight-16B-A3B-Instruct + train_micro_batch_size: 1 + max_total_sequence_length: 4096 + activation_checkpointing_enabled: false + dtensor_cfg: + expert_parallel_size: 8 + clear_cache_every_n_steps: 50 + automodel_kwargs: + use_liger_kernel: false + backend: + _target_: nemo_automodel.components.moe.utils.BackendConfig + attn: te + linear: te + rms_norm: te + enable_deepep: true + fake_balanced_gate: false + enable_hf_state_dict_adapter: true + enable_fsdp_optimizations: false + gate_precision: float64 + dynamic_batching: + enabled: true + sequence_packing: + enabled: false + make_sequence_length_divisible_by: 4 + optimizer: + kwargs: + lr: 3.0e-07 + scheduler: + - name: torch.optim.lr_scheduler.LinearLR + kwargs: + start_factor: 0.1 + end_factor: 1 + total_iters: 13 + - name: torch.optim.lr_scheduler.ConstantLR + kwargs: + factor: 1 + total_iters: 10000000000 + - milestones: + - 13 + generation: + _pad_token_id: 163838 + max_new_tokens: 4096 + pad_token_id: 163838 + stop_token_ids: + - 163585 +data: + max_input_seq_length: 4096 +cluster: + gpus_per_node: 8 +logger: + log_dir: logs/grpo-moonlight-16b-a3b-instruct-automodel-1n8g-ep8 + wandb_enabled: true + tensorboard_enabled: true + wandb: + project: nemo-rl + name: grpo-moonlight-16b-a3b-instruct-automodel-1n8g-ep8 diff --git a/nemo_rl/models/policy/workers/dtensor_policy_worker_v2.py b/nemo_rl/models/policy/workers/dtensor_policy_worker_v2.py index 785568cc761..afd66e4e2b6 100644 --- a/nemo_rl/models/policy/workers/dtensor_policy_worker_v2.py +++ b/nemo_rl/models/policy/workers/dtensor_policy_worker_v2.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import contextlib import gc import itertools import os @@ -35,7 +36,9 @@ from nemo_automodel.components.config.loader import _resolve_target from nemo_automodel.components.distributed.cp_utils import ( create_context_parallel_ctx, - get_train_context, +) +from nemo_automodel.components.distributed.cp_utils import ( + get_train_context as get_train_context_automodel, ) from nemo_automodel.components.distributed.fsdp2 import ( FSDP2Manager, @@ -99,6 +102,49 @@ } +def _maybe_adapt_tensor_to_hf( + model_part: nn.Module, fqn: str, tensor: torch.Tensor, quantization: bool = False +) -> list[tuple[str, torch.Tensor]]: + adapter = getattr(model_part, "state_dict_adapter", None) + if adapter: + return adapter.convert_single_tensor_to_hf( + fqn, + tensor, + exclude_key_regex=r".*_extra_state.*", + quantization=quantization, + ) + return [(fqn, tensor)] + + +@contextlib.contextmanager +def get_train_context( + cp_size: int, + cp_mesh: Any, + cp_buffers: list, + sequence_dim: int, + dtype: torch.dtype, + autocast_enabled: bool = True, +) -> Generator[None, None, None]: + """Create combined context manager for training with context parallel and autocast.""" + with contextlib.ExitStack() as stack: + context_parallel_ctx = None + if cp_size > 1: + # Create context parallel context + context_parallel_ctx = create_context_parallel_ctx( + cp_mesh=cp_mesh, + cp_buffers=cp_buffers, + cp_seq_dims=[sequence_dim] * len(cp_buffers), + cp_no_restore_buffers=set(cp_buffers), + ) + + stack.enter_context( + get_train_context_automodel(False, False, context_parallel_ctx)() + ) + if autocast_enabled: + stack.enter_context(torch.autocast(device_type="cuda", dtype=dtype)) + yield + + @ray.remote( runtime_env=get_runtime_env_for_policy_worker("dtensor_policy_worker_v2") ) # pragma: no cover @@ -404,14 +450,16 @@ def __init__( self.cp_size = manager.cp_size # Parallelize model - is_moe_model = any(["expert" in key for key in self.model_state_dict_keys]) - is_hf_model = ( + self.is_moe_model = any(["expert" in key for key in self.model_state_dict_keys]) + self.is_hf_model = ( model_config.architectures[0] not in ModelRegistry.model_arch_name_to_cls ) + # Autocast is disabled for custom MoE models (non-HF) to avoid numerical issues + self.autocast_enabled = not (self.is_moe_model and not self.is_hf_model) if ( not isinstance(self.model, PreTrainedModel) - and is_moe_model - and not is_hf_model + and self.is_moe_model + and not self.is_hf_model ): assert self.tp_size == 1, ( "Using custom implementation {self.model.__class__.__name__} for MoE model {model_name} which doesn't support tp_size > 1. Please use expert_parallel_size > 1 for custom implementation or set force_hf=True in your config at policy->dtensor_cfg->automodel_kwargs to use the HuggingFace implementation." @@ -726,7 +774,6 @@ def train( "Sequence parallel is not supported with multimodal since there's an issue when you do not pass position_ids. See https://github.com/NVIDIA-NeMo/Automodel/issues/652" ) - context_parallel_ctx = None if self.cp_size > 1: assert len(vlm_kwargs) == 0, ( f"multimodal kwargs={vlm_kwargs} are not supported for context parallel" @@ -734,48 +781,45 @@ def train( seq_index = torch.arange( seq_len, device=input_ids.device ).repeat(1, 1) - cp_buffers = ( - [input_ids, position_ids, seq_index] - if self.cp_size > 1 - else [] - ) + cp_buffers = [input_ids, position_ids, seq_index] + else: + cp_buffers = [] + seq_index = None - # Create context parallel context - context_parallel_ctx = create_context_parallel_ctx( - cp_mesh=self.cp_mesh, - cp_buffers=cp_buffers, - cp_seq_dims=[sequence_dim] * len(cp_buffers), - cp_no_restore_buffers=set(cp_buffers), + with get_train_context( + cp_size=self.cp_size, + cp_mesh=self.cp_mesh, + cp_buffers=cp_buffers, + sequence_dim=sequence_dim, + dtype=self.dtype, + autocast_enabled=self.autocast_enabled, + ): + model_args = dict( + input_ids=input_ids, + attention_mask=attention_mask, + position_ids=position_ids, + use_cache=False, + flash_attn_kwargs=flash_attn_kwargs, + **vlm_kwargs, ) - with get_train_context(False, False, context_parallel_ctx)(): - with torch.autocast(device_type="cuda", dtype=self.dtype): - model_args = dict( - input_ids=input_ids, - attention_mask=attention_mask, - position_ids=position_ids, - use_cache=False, - flash_attn_kwargs=flash_attn_kwargs, - **vlm_kwargs, - ) - - if self._is_reward_model: - # `flash_attn_kwarg` is not supported for `LlamaForSequenceClassification`. - # Note that it should be empty anyway since sequence packing - # is not supported for reward models. - assert not flash_attn_kwargs - del model_args["flash_attn_kwargs"] - # remove flash_attn_kwargs if there are multimodal kwargs - if len(vlm_kwargs) > 0: - del model_args["flash_attn_kwargs"] + if self._is_reward_model: + # `flash_attn_kwarg` is not supported for `LlamaForSequenceClassification`. + # Note that it should be empty anyway since sequence packing + # is not supported for reward models. + assert not flash_attn_kwargs + del model_args["flash_attn_kwargs"] + # remove flash_attn_kwargs if there are multimodal kwargs + if len(vlm_kwargs) > 0: + del model_args["flash_attn_kwargs"] - if ( - not self.allow_flash_attn_args - and "flash_attn_kwargs" in model_args - ): - del model_args["flash_attn_kwargs"] + if ( + not self.allow_flash_attn_args + and "flash_attn_kwargs" in model_args + ): + del model_args["flash_attn_kwargs"] - outputs = self.model(**model_args) + outputs = self.model(**model_args) # Get logprobs if isinstance(outputs, (torch.Tensor, DTensor)): @@ -1064,7 +1108,6 @@ def get_logprobs( if len(vlm_kwargs) > 0: position_ids = None - context_parallel_ctx = None if self.cp_size > 1: assert len(vlm_kwargs) == 0, ( "multimodal kwargs are not supported for context parallel" @@ -1073,35 +1116,36 @@ def get_logprobs( 1, 1 ) cp_buffers = [input_ids, position_ids, seq_index] - - # Create context parallel context - context_parallel_ctx = create_context_parallel_ctx( - cp_mesh=self.cp_mesh, - cp_buffers=cp_buffers, - cp_seq_dims=[sequence_dim] * len(cp_buffers), - cp_no_restore_buffers=set(cp_buffers), + else: + cp_buffers = [] + seq_index = None + + with get_train_context( + cp_size=self.cp_size, + cp_mesh=self.cp_mesh, + cp_buffers=cp_buffers, + sequence_dim=sequence_dim, + dtype=self.dtype, + autocast_enabled=self.autocast_enabled, + ): + model_args = dict( + input_ids=input_ids, + attention_mask=attention_mask, + position_ids=position_ids, + use_cache=False, + flash_attn_kwargs=flash_attn_kwargs, + **vlm_kwargs, ) + if len(vlm_kwargs) > 0: + del model_args["flash_attn_kwargs"] - with get_train_context(False, False, context_parallel_ctx)(): - with torch.autocast(device_type="cuda", dtype=self.dtype): - model_args = dict( - input_ids=input_ids, - attention_mask=attention_mask, - position_ids=position_ids, - use_cache=False, - flash_attn_kwargs=flash_attn_kwargs, - **vlm_kwargs, - ) - if len(vlm_kwargs) > 0: - del model_args["flash_attn_kwargs"] - - if ( - not self.allow_flash_attn_args - and "flash_attn_kwargs" in model_args - ): - del model_args["flash_attn_kwargs"] + if ( + not self.allow_flash_attn_args + and "flash_attn_kwargs" in model_args + ): + del model_args["flash_attn_kwargs"] - outputs = self.model(**model_args) + outputs = self.model(**model_args) logits = outputs.logits if hasattr(outputs, "logits") else outputs @@ -1326,29 +1370,30 @@ def score(self, data: BatchedDataDict) -> BatchedDataDict[ScoreOutputSpec]: dtype=torch.bool, device=input_ids.device, ) - context_parallel_ctx = None if self.cp_size > 1: seq_index = torch.arange(seq_len, device=input_ids.device).repeat( 1, 1 ) cp_buffers = [input_ids, position_ids, seq_index] - - # Create context parallel context - context_parallel_ctx = create_context_parallel_ctx( - cp_mesh=self.cp_mesh, - cp_buffers=cp_buffers, - cp_seq_dims=[sequence_dim] * len(cp_buffers), - cp_no_restore_buffers=set(cp_buffers), + else: + cp_buffers = [] + seq_index = None + + with get_train_context( + cp_size=self.cp_size, + cp_mesh=self.cp_mesh, + cp_buffers=cp_buffers, + sequence_dim=sequence_dim, + dtype=self.dtype, + autocast_enabled=self.autocast_enabled, + ): + model_args = dict( + input_ids=input_ids, + attention_mask=attention_mask, + position_ids=position_ids, + use_cache=False, ) - with get_train_context(False, False, context_parallel_ctx)(): - with torch.autocast(device_type="cuda", dtype=self.dtype): - model_args = dict( - input_ids=input_ids, - attention_mask=attention_mask, - position_ids=position_ids, - use_cache=False, - ) - outputs = self.model(**model_args) + outputs = self.model(**model_args) if not hasattr(outputs, "logits"): logits = self.model.lm_head(outputs.last_hidden_state) @@ -1477,31 +1522,31 @@ def get_topk_logits( (batch_size, seq_len), dtype=torch.long, device=input_ids.device ) - context_parallel_ctx = None if self.cp_size > 1: seq_index = torch.arange(seq_len, device=input_ids.device).repeat( 1, 1 ) cp_buffers = [input_ids, position_ids, seq_index] - - # Create context parallel context - context_parallel_ctx = create_context_parallel_ctx( - cp_mesh=self.cp_mesh, - cp_buffers=cp_buffers, - cp_seq_dims=[sequence_dim] * len(cp_buffers), - cp_no_restore_buffers=set(cp_buffers), + else: + cp_buffers = [] + seq_index = None + + with get_train_context( + cp_size=self.cp_size, + cp_mesh=self.cp_mesh, + cp_buffers=cp_buffers, + sequence_dim=sequence_dim, + dtype=self.dtype, + autocast_enabled=self.autocast_enabled, + ): + outputs = self.model( + input_ids=input_ids, + attention_mask=attention_mask_input_all_ones, + position_ids=position_ids, + use_cache=False, + flash_attn_kwargs=flash_attn_kwargs, ) - with get_train_context(False, False, context_parallel_ctx)(): - with torch.autocast(device_type="cuda", dtype=self.dtype): - outputs = self.model( - input_ids=input_ids, - attention_mask=attention_mask_input_all_ones, - position_ids=position_ids, - use_cache=False, - flash_attn_kwargs=flash_attn_kwargs, - ) - if not hasattr(outputs, "logits"): logits = self.model.lm_head(outputs.last_hidden_state) else: @@ -1708,8 +1753,15 @@ def prepare_refit_info(self) -> Optional[dict[str, Any]]: """Prepare state dict metadata for weight refitting and IPC streaming.""" state_dict_info = {} for name, tensor in self.model.state_dict().items(): + full_tensor = ( + tensor.full_tensor() if isinstance(tensor, DTensor) else tensor + ) # all tensor will be casted to self.dtype in stream_weights_via_ipc_zmq/broadcast_weights_for_collective - state_dict_info[name] = (tensor.shape, self.dtype) + adapted_fqn_tensors = _maybe_adapt_tensor_to_hf( + self.model, name, full_tensor + ) + for adapted_fqn, adapted_tensor in adapted_fqn_tensors: + state_dict_info[adapted_fqn] = (adapted_tensor.shape, self.dtype) return state_dict_info @@ -1750,17 +1802,18 @@ def stream_weights_via_ipc_zmq( def dtensor_params_generator(): """Generator that yields (name, tensor) pairs, converting DTensors to local tensors.""" for name, tensor in self.model.state_dict().items(): - if isinstance(tensor, DTensor): - # Convert DTensor to full tensor for streaming - full_tensor = tensor.full_tensor() + full_tensor = ( + tensor.full_tensor() if isinstance(tensor, DTensor) else tensor + ) + adapted_fqn_tensors = _maybe_adapt_tensor_to_hf( + self.model, name, full_tensor + ) + for adapted_fqn, adapted_tensor in adapted_fqn_tensors: # Convert to target dtype yield ( - name, - full_tensor.to(self.dtype, non_blocking=True).contiguous(), + adapted_fqn, + adapted_tensor.to(self.dtype, non_blocking=True).contiguous(), ) - else: - # Convert to target dtype - yield name, tensor.to(self.dtype, non_blocking=True).contiguous() # Use the shared implementation stream_weights_via_ipc_zmq_impl( @@ -1789,17 +1842,27 @@ def broadcast_weights_for_collective( ) self.model = self.move_to_cuda(self.model) - def _dtensor_post_iter_func(tensor, dtype): - if isinstance(tensor, DTensor): - tensor = tensor.full_tensor() - tensor = tensor.to(dtype, non_blocking=True) - return tensor + def dtensor_params_generator(): + """Generator that yields (name, tensor) pairs, converting DTensors to local tensors and adapting to HF format.""" + for name, tensor in self.model.state_dict().items(): + full_tensor = ( + tensor.full_tensor() if isinstance(tensor, DTensor) else tensor + ) + adapted_fqn_tensors = _maybe_adapt_tensor_to_hf( + self.model, name, full_tensor + ) + for adapted_fqn, adapted_tensor in adapted_fqn_tensors: + # Convert to target dtype + yield ( + adapted_fqn, + adapted_tensor.to(self.dtype, non_blocking=True).contiguous(), + ) # param_iterator will return (name, tensor), we only need tensor - dtensor_post_iter_func = lambda x: _dtensor_post_iter_func(x[1], self.dtype) + dtensor_post_iter_func = lambda x: x[1] packed_broadcast_producer( - iterator=iter(self.model.state_dict().items()), + iterator=dtensor_params_generator(), group=self.model_update_group, src=0, post_iter_func=dtensor_post_iter_func, diff --git a/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh b/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh new file mode 100755 index 00000000000..0d18854c1f6 --- /dev/null +++ b/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh @@ -0,0 +1,42 @@ +#!/bin/bash +SCRIPT_DIR=$( cd -- "$( dirname -- "${BASH_SOURCE[0]}" )" &> /dev/null && pwd) +source $SCRIPT_DIR/common.env + +# ===== BEGIN CONFIG ===== +NUM_NODES=1 +STEPS_PER_RUN=30 +MAX_STEPS=30 +NUM_RUNS=$(( (MAX_STEPS + STEPS_PER_RUN - 1) / STEPS_PER_RUN )) # Round up +NUM_MINUTES=180 +# ===== END CONFIG ===== + +exit_if_max_steps_reached + +# Run the experiment +cd $PROJECT_ROOT +uv run examples/run_grpo_math.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 \ + 'mean(data["train/gen_kl_error"]) < 0.001' \ + 'data["train/gen_kl_error"]["30"] < 0.001 ' \ + 'data["train/reward"]["30"] > 0.4' \ + 'data["train/grad_norm"] < 0.5' \ + 'data["train/grad_norm"] > 0.05' \ +fi From d34d6a2520a4f8cee195401b5dca077952a61215 Mon Sep 17 00:00:00 2001 From: Hemil Desai Date: Wed, 24 Dec 2025 18:02:23 +0000 Subject: [PATCH 02/12] fix Signed-off-by: Hemil Desai --- .../llm/grpo-moonlight-16b-automodel-1n8g-ep8.yaml | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/examples/configs/recipes/llm/grpo-moonlight-16b-automodel-1n8g-ep8.yaml b/examples/configs/recipes/llm/grpo-moonlight-16b-automodel-1n8g-ep8.yaml index 77bc2db8b92..16d123d711d 100644 --- a/examples/configs/recipes/llm/grpo-moonlight-16b-automodel-1n8g-ep8.yaml +++ b/examples/configs/recipes/llm/grpo-moonlight-16b-automodel-1n8g-ep8.yaml @@ -1,10 +1,12 @@ defaults: ../../grpo_math_1B.yaml grpo: - num_prompts_per_step: 64 - num_generations_per_prompt: 32 - max_num_steps: 30 + val_period: -1 +loss_fn: + reference_policy_kl_penalty: 0.04 checkpointing: checkpoint_dir: results/grpo-moonlight-16b-a3b-instruct-automodel-1n8g-ep8 + enabled: false + save_period: 10000 policy: model_name: moonshotai/Moonlight-16B-A3B-Instruct train_micro_batch_size: 1 From 7fa7dcc9ddaaa2e067de1f2fdf1c73d6bdbf977b Mon Sep 17 00:00:00 2001 From: Hemil Desai Date: Wed, 24 Dec 2025 18:53:26 +0000 Subject: [PATCH 03/12] fix Signed-off-by: Hemil Desai --- tests/test_suites/nightly.txt | 3 +++ 1 file changed, 3 insertions(+) diff --git a/tests/test_suites/nightly.txt b/tests/test_suites/nightly.txt index ee1fda01b11..4c93e4fcb9c 100644 --- a/tests/test_suites/nightly.txt +++ b/tests/test_suites/nightly.txt @@ -10,6 +10,9 @@ tests/test_suites/llm/grpo-gemma3-1b-it-1n8g-fsdp2tp1.sh # Dtensor (Qwen/Qwen2.5-7B-Instruct) tests/test_suites/llm/grpo-qwen2.5-7b-instruct-4n8g-fsdp2tp4.v3.sh +# Dtensor (moonshotai/Moonlight-16B-A3B-Instruct) +tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh + # Megatron tests/test_suites/llm/grpo-llama3.2-1b-instruct-1n8g-megatron.sh tests/test_suites/llm/grpo-llama3.2-1b-instruct-1n8g-megatron_generation.sh From 7a3b940823d0d0f9d8d9bd0ac2ec3a0782ee744e Mon Sep 17 00:00:00 2001 From: Hemil Desai Date: Sun, 4 Jan 2026 23:40:55 +0000 Subject: [PATCH 04/12] fix Signed-off-by: Hemil Desai --- .../llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh | 3 +++ tests/unit/test_recipes_and_test_suites.py | 6 +++--- 2 files changed, 6 insertions(+), 3 deletions(-) diff --git a/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh b/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh index 0d18854c1f6..48e441b4268 100755 --- a/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh +++ b/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh @@ -12,6 +12,9 @@ NUM_MINUTES=180 exit_if_max_steps_reached +# This is needed for correctly importing transformers_modules for moonlight +export HF_HOME=$HF_HOME + # Run the experiment cd $PROJECT_ROOT uv run examples/run_grpo_math.py \ diff --git a/tests/unit/test_recipes_and_test_suites.py b/tests/unit/test_recipes_and_test_suites.py index f41fe31eaee..88553567c73 100644 --- a/tests/unit/test_recipes_and_test_suites.py +++ b/tests/unit/test_recipes_and_test_suites.py @@ -180,7 +180,7 @@ def test_all_recipe_yamls_accounted_for_in_test_suites( ) -def test_nightly_compute_stays_below_1180_hours(nightly_test_suite, tracker): +def test_nightly_compute_stays_below_1210_hours(nightly_test_suite, tracker): command = f"DRYRUN=1 HF_HOME=... HF_DATASETS_CACHE=... CONTAINER= ACCOUNT= PARTITION= ./tools/launch {' '.join(nightly_test_suite)}" print(f"Running command: {command}") @@ -212,8 +212,8 @@ def test_nightly_compute_stays_below_1180_hours(nightly_test_suite, tracker): f"Last line of output was not as expected: '{last_line}'" ) total_gpu_hours = float(last_line.split(":")[-1].strip()) - assert total_gpu_hours <= 1180, ( - f"Total GPU hours exceeded 1180: {last_line}. We should revisit the test suites to reduce the total GPU hours." + assert total_gpu_hours <= 1210, ( + f"Total GPU hours exceeded 1210: {last_line}. We should revisit the test suites to reduce the total GPU hours." ) tracker.track("total_nightly_gpu_hours", total_gpu_hours) From a96d528726f6545906a9fa50dd92c79b55a7a95c Mon Sep 17 00:00:00 2001 From: Hemil Desai Date: Sun, 4 Jan 2026 17:22:47 -0800 Subject: [PATCH 05/12] fix Signed-off-by: Hemil Desai --- tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh | 3 --- 1 file changed, 3 deletions(-) diff --git a/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh b/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh index 48e441b4268..0d18854c1f6 100755 --- a/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh +++ b/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh @@ -12,9 +12,6 @@ NUM_MINUTES=180 exit_if_max_steps_reached -# This is needed for correctly importing transformers_modules for moonlight -export HF_HOME=$HF_HOME - # Run the experiment cd $PROJECT_ROOT uv run examples/run_grpo_math.py \ From c86bf9ce5e06afed9d8499126a9e26fdac72b690 Mon Sep 17 00:00:00 2001 From: Hemil Desai Date: Mon, 5 Jan 2026 09:40:31 -0800 Subject: [PATCH 06/12] fix Signed-off-by: Hemil Desai --- tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh b/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh index 0d18854c1f6..4db32695088 100755 --- a/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh +++ b/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh @@ -38,5 +38,5 @@ if [[ $(jq 'to_entries | .[] | select(.key == "train/loss") | .value | keys | ma 'data["train/gen_kl_error"]["30"] < 0.001 ' \ 'data["train/reward"]["30"] > 0.4' \ 'data["train/grad_norm"] < 0.5' \ - 'data["train/grad_norm"] > 0.05' \ + 'data["train/grad_norm"] > 0.05' fi From 33eccd0f420edf8a45898f5f62cc81677058d6ed Mon Sep 17 00:00:00 2001 From: Hemil Desai Date: Mon, 5 Jan 2026 17:59:03 +0000 Subject: [PATCH 07/12] fix Signed-off-by: Hemil Desai --- .../models/policy/test_dtensor_worker_v2.py | 292 ++++++++++++++++++ 1 file changed, 292 insertions(+) diff --git a/tests/unit/models/policy/test_dtensor_worker_v2.py b/tests/unit/models/policy/test_dtensor_worker_v2.py index 4e9f33f99e2..66409bd2056 100644 --- a/tests/unit/models/policy/test_dtensor_worker_v2.py +++ b/tests/unit/models/policy/test_dtensor_worker_v2.py @@ -12,12 +12,15 @@ # See the License for the specific language governing permissions and # limitations under the License. +import contextlib import os import tempfile +from unittest.mock import MagicMock, Mock, patch import pytest import ray import torch +import torch.nn as nn from nemo_rl.algorithms.utils import get_tokenizer from nemo_rl.distributed.batched_data_dict import BatchedDataDict @@ -26,6 +29,16 @@ from nemo_rl.models.policy.lm_policy import Policy from tests.unit.test_utils import SimpleLoss +try: + from nemo_rl.models.policy.workers.dtensor_policy_worker_v2 import ( + _maybe_adapt_tensor_to_hf, + get_train_context, + ) + + NEMO_AUTOMODEL_AVAILABLE = True +except ImportError: + NEMO_AUTOMODEL_AVAILABLE = False + def create_test_config( model_name: str, @@ -456,3 +469,282 @@ def test_dtensor_v2_mixed_precision_training_and_logprobs( assert worker_info is not None, "Should get worker info" finally: policy.shutdown() + + +@pytest.mark.automodel +@pytest.mark.skipif(not NEMO_AUTOMODEL_AVAILABLE, reason="nemo_automodel not available") +class TestMaybeAdaptTensorToHF: + """Tests for the _maybe_adapt_tensor_to_hf helper function.""" + + def test_no_adapter_returns_single_tuple(self): + """Test that when model has no adapter, returns single FQN-tensor tuple.""" + # Arrange + model = nn.Linear(10, 10) + fqn = "layer.weight" + tensor = torch.randn(10, 10) + + # Act + result = _maybe_adapt_tensor_to_hf(model, fqn, tensor) + + # Assert + assert len(result) == 1, "Should return single tuple when no adapter" + assert result[0][0] == fqn, "FQN should be unchanged" + assert torch.equal(result[0][1], tensor), "Tensor should be unchanged" + + def test_adapter_converts_single_tensor(self): + """Test that adapter is called when present on model.""" + # Arrange + model = nn.Linear(10, 10) + adapter_mock = Mock() + adapter_mock.convert_single_tensor_to_hf.return_value = [ + ("adapted.weight", torch.randn(10, 10)), + ("adapted.bias", torch.randn(10)), + ] + model.state_dict_adapter = adapter_mock + + fqn = "layer.weight" + tensor = torch.randn(10, 10) + + # Act + result = _maybe_adapt_tensor_to_hf(model, fqn, tensor) + + # Assert + adapter_mock.convert_single_tensor_to_hf.assert_called_once_with( + fqn, + tensor, + exclude_key_regex=r".*_extra_state.*", + quantization=False, + ) + assert len(result) == 2, "Should return multiple adapted tensors" + assert result[0][0] == "adapted.weight" + assert result[1][0] == "adapted.bias" + + def test_adapter_with_quantization_flag(self): + """Test that quantization flag is passed to adapter correctly.""" + # Arrange + model = nn.Linear(10, 10) + adapter_mock = Mock() + adapter_mock.convert_single_tensor_to_hf.return_value = [ + ("quantized.weight", torch.randn(10, 10)) + ] + model.state_dict_adapter = adapter_mock + + fqn = "layer.weight" + tensor = torch.randn(10, 10) + + # Act + result = _maybe_adapt_tensor_to_hf(model, fqn, tensor, quantization=True) + + # Assert + adapter_mock.convert_single_tensor_to_hf.assert_called_once_with( + fqn, + tensor, + exclude_key_regex=r".*_extra_state.*", + quantization=True, + ) + assert len(result) == 1 + + def test_adapter_excludes_extra_state_regex(self): + """Test that _extra_state regex is always passed to exclude such tensors.""" + # Arrange + model = nn.Linear(10, 10) + adapter_mock = Mock() + adapter_mock.convert_single_tensor_to_hf.return_value = [] + model.state_dict_adapter = adapter_mock + + fqn = "layer._extra_state" + tensor = torch.randn(10) + + # Act + _maybe_adapt_tensor_to_hf(model, fqn, tensor) + + # Assert + adapter_mock.convert_single_tensor_to_hf.assert_called_once() + call_kwargs = adapter_mock.convert_single_tensor_to_hf.call_args[1] + assert call_kwargs["exclude_key_regex"] == r".*_extra_state.*", ( + "Should exclude extra_state tensors" + ) + + +@pytest.mark.automodel +@pytest.mark.skipif(not NEMO_AUTOMODEL_AVAILABLE, reason="nemo_automodel not available") +class TestGetTrainContext: + """Tests for the get_train_context context manager function.""" + + @patch( + "nemo_rl.models.policy.workers.dtensor_policy_worker_v2.get_train_context_automodel" + ) + @patch( + "nemo_rl.models.policy.workers.dtensor_policy_worker_v2.create_context_parallel_ctx" + ) + def test_no_cp_with_autocast(self, mock_create_cp_ctx, mock_get_train_ctx_am): + """Test context creation without context parallel but with autocast.""" + # Arrange + mock_get_train_ctx_am.return_value = lambda: contextlib.nullcontext() + + cp_size = 1 + cp_mesh = None + cp_buffers = [] + sequence_dim = 1 + dtype = torch.bfloat16 + + # Act + with get_train_context( + cp_size=cp_size, + cp_mesh=cp_mesh, + cp_buffers=cp_buffers, + sequence_dim=sequence_dim, + dtype=dtype, + autocast_enabled=True, + ): + pass + + # Assert - CP context should not be created when cp_size=1 + mock_create_cp_ctx.assert_not_called() + mock_get_train_ctx_am.assert_called_once_with(False, False, None) + + @patch( + "nemo_rl.models.policy.workers.dtensor_policy_worker_v2.get_train_context_automodel" + ) + @patch( + "nemo_rl.models.policy.workers.dtensor_policy_worker_v2.create_context_parallel_ctx" + ) + def test_with_cp_and_autocast(self, mock_create_cp_ctx, mock_get_train_ctx_am): + """Test context creation with context parallel and autocast.""" + # Arrange + mock_cp_ctx = MagicMock() + mock_create_cp_ctx.return_value = mock_cp_ctx + mock_get_train_ctx_am.return_value = lambda: contextlib.nullcontext() + + cp_size = 2 + cp_mesh = MagicMock() + cp_buffers = [torch.randn(2, 10), torch.randn(2, 10)] + sequence_dim = 1 + dtype = torch.bfloat16 + + # Act + with get_train_context( + cp_size=cp_size, + cp_mesh=cp_mesh, + cp_buffers=cp_buffers, + sequence_dim=sequence_dim, + dtype=dtype, + autocast_enabled=True, + ): + pass + + # Assert - CP context should be created when cp_size > 1 + mock_create_cp_ctx.assert_called_once_with( + cp_mesh=cp_mesh, + cp_buffers=cp_buffers, + cp_seq_dims=[sequence_dim] * len(cp_buffers), + cp_no_restore_buffers=set(cp_buffers), + ) + mock_get_train_ctx_am.assert_called_once_with(False, False, mock_cp_ctx) + + @patch( + "nemo_rl.models.policy.workers.dtensor_policy_worker_v2.get_train_context_automodel" + ) + def test_autocast_disabled(self, mock_get_train_ctx_am): + """Test context creation with autocast disabled.""" + # Arrange + mock_get_train_ctx_am.return_value = lambda: contextlib.nullcontext() + + cp_size = 1 + cp_mesh = None + cp_buffers = [] + sequence_dim = 1 + dtype = torch.bfloat16 + + # Act + with get_train_context( + cp_size=cp_size, + cp_mesh=cp_mesh, + cp_buffers=cp_buffers, + sequence_dim=sequence_dim, + dtype=dtype, + autocast_enabled=False, + ): + # Verify we're NOT in autocast mode + assert not torch.is_autocast_enabled("cuda"), ( + "Autocast should be disabled when autocast_enabled=False" + ) + + # Assert + mock_get_train_ctx_am.assert_called_once_with(False, False, None) + + @patch( + "nemo_rl.models.policy.workers.dtensor_policy_worker_v2.get_train_context_automodel" + ) + @patch( + "nemo_rl.models.policy.workers.dtensor_policy_worker_v2.create_context_parallel_ctx" + ) + def test_cp_buffers_empty_when_cp_size_one( + self, mock_create_cp_ctx, mock_get_train_ctx_am + ): + """Test that CP context is not created when cp_size is 1.""" + # Arrange + mock_get_train_ctx_am.return_value = lambda: contextlib.nullcontext() + + cp_size = 1 + cp_mesh = MagicMock() + cp_buffers = [] # Empty buffers for cp_size=1 + sequence_dim = 1 + dtype = torch.float32 + + # Act + with get_train_context( + cp_size=cp_size, + cp_mesh=cp_mesh, + cp_buffers=cp_buffers, + sequence_dim=sequence_dim, + dtype=dtype, + autocast_enabled=True, + ): + pass + + # Assert - CP context should not be created when cp_size=1 + mock_create_cp_ctx.assert_not_called() + + @patch( + "nemo_rl.models.policy.workers.dtensor_policy_worker_v2.get_train_context_automodel" + ) + @patch( + "nemo_rl.models.policy.workers.dtensor_policy_worker_v2.create_context_parallel_ctx" + ) + def test_multiple_cp_buffers_sequence_dim_replication( + self, mock_create_cp_ctx, mock_get_train_ctx_am + ): + """Test that sequence_dim is properly replicated for each CP buffer.""" + # Arrange + mock_cp_ctx = MagicMock() + mock_create_cp_ctx.return_value = mock_cp_ctx + mock_get_train_ctx_am.return_value = lambda: contextlib.nullcontext() + + cp_size = 2 + cp_mesh = MagicMock() + # Three buffers + cp_buffers = [torch.randn(2, 10), torch.randn(2, 10), torch.randn(2, 10)] + sequence_dim = 1 + dtype = torch.float16 + + # Act + with get_train_context( + cp_size=cp_size, + cp_mesh=cp_mesh, + cp_buffers=cp_buffers, + sequence_dim=sequence_dim, + dtype=dtype, + autocast_enabled=True, + ): + pass + + # Assert - sequence_dim should be replicated for each buffer + mock_create_cp_ctx.assert_called_once() + call_kwargs = mock_create_cp_ctx.call_args[1] + assert call_kwargs["cp_seq_dims"] == [ + sequence_dim, + sequence_dim, + sequence_dim, + ], "sequence_dim should be replicated for each buffer" + assert len(call_kwargs["cp_seq_dims"]) == 3 From b256333f645a6e5529c1704af1020f440f798018 Mon Sep 17 00:00:00 2001 From: Hemil Desai Date: Tue, 6 Jan 2026 12:57:12 +0000 Subject: [PATCH 08/12] fix Signed-off-by: Hemil Desai --- nemo_rl/models/policy/__init__.py | 2 +- .../workers/dtensor_policy_worker_v2.py | 59 ++++---- .../grpo-moonlight-16b-automodel-1n8g-ep8.sh | 4 +- .../models/policy/test_dtensor_worker_v2.py | 137 ++++++++++++++++++ 4 files changed, 165 insertions(+), 37 deletions(-) diff --git a/nemo_rl/models/policy/__init__.py b/nemo_rl/models/policy/__init__.py index 781c8cdf175..5cba9457a6c 100644 --- a/nemo_rl/models/policy/__init__.py +++ b/nemo_rl/models/policy/__init__.py @@ -47,7 +47,7 @@ class AutomodelBackendConfig(TypedDict): enable_deepep: NotRequired[bool] # Use fake balanced gate for testing/debugging MoE fake_balanced_gate: NotRequired[bool] - # Enable HuggingFace state dict adapter for checkpoint loading + # Enable HuggingFace state dict adapter for checkpoint saving/loading plus refit support for RL enable_hf_state_dict_adapter: NotRequired[bool] # Enable FSDP-specific optimizations enable_fsdp_optimizations: NotRequired[bool] diff --git a/nemo_rl/models/policy/workers/dtensor_policy_worker_v2.py b/nemo_rl/models/policy/workers/dtensor_policy_worker_v2.py index afd66e4e2b6..99fa3218fbb 100644 --- a/nemo_rl/models/policy/workers/dtensor_policy_worker_v2.py +++ b/nemo_rl/models/policy/workers/dtensor_policy_worker_v2.py @@ -102,6 +102,29 @@ } +def dtensor_params_generator( + model: nn.Module, target_dtype: torch.dtype +) -> Generator[tuple[str, torch.Tensor], None, None]: + """Generator that yields (name, tensor) pairs, converting DTensors to local tensors and adapting to HF format. + + Args: + model: The model whose parameters to generate. + target_dtype: The dtype to convert tensors to. + + Yields: + Tuples of (fully_qualified_name, tensor) where tensors are converted to target dtype and made contiguous. + """ + for name, tensor in model.state_dict().items(): + full_tensor = tensor.full_tensor() if isinstance(tensor, DTensor) else tensor + adapted_fqn_tensors = _maybe_adapt_tensor_to_hf(model, name, full_tensor) + for adapted_fqn, adapted_tensor in adapted_fqn_tensors: + # Convert to target dtype + yield ( + adapted_fqn, + adapted_tensor.to(target_dtype, non_blocking=True).contiguous(), + ) + + def _maybe_adapt_tensor_to_hf( model_part: nn.Module, fqn: str, tensor: torch.Tensor, quantization: bool = False ) -> list[tuple[str, torch.Tensor]]: @@ -1799,25 +1822,9 @@ def stream_weights_via_ipc_zmq( from nemo_rl.models.policy.utils import stream_weights_via_ipc_zmq_impl - def dtensor_params_generator(): - """Generator that yields (name, tensor) pairs, converting DTensors to local tensors.""" - for name, tensor in self.model.state_dict().items(): - full_tensor = ( - tensor.full_tensor() if isinstance(tensor, DTensor) else tensor - ) - adapted_fqn_tensors = _maybe_adapt_tensor_to_hf( - self.model, name, full_tensor - ) - for adapted_fqn, adapted_tensor in adapted_fqn_tensors: - # Convert to target dtype - yield ( - adapted_fqn, - adapted_tensor.to(self.dtype, non_blocking=True).contiguous(), - ) - # Use the shared implementation stream_weights_via_ipc_zmq_impl( - params_generator=dtensor_params_generator(), + params_generator=dtensor_params_generator(self.model, self.dtype), buffer_size_bytes=buffer_size_bytes, zmq_socket=self.zmq_socket, rank=self.rank, @@ -1842,27 +1849,11 @@ def broadcast_weights_for_collective( ) self.model = self.move_to_cuda(self.model) - def dtensor_params_generator(): - """Generator that yields (name, tensor) pairs, converting DTensors to local tensors and adapting to HF format.""" - for name, tensor in self.model.state_dict().items(): - full_tensor = ( - tensor.full_tensor() if isinstance(tensor, DTensor) else tensor - ) - adapted_fqn_tensors = _maybe_adapt_tensor_to_hf( - self.model, name, full_tensor - ) - for adapted_fqn, adapted_tensor in adapted_fqn_tensors: - # Convert to target dtype - yield ( - adapted_fqn, - adapted_tensor.to(self.dtype, non_blocking=True).contiguous(), - ) - # param_iterator will return (name, tensor), we only need tensor dtensor_post_iter_func = lambda x: x[1] packed_broadcast_producer( - iterator=dtensor_params_generator(), + iterator=dtensor_params_generator(self.model, self.dtype), group=self.model_update_group, src=0, post_iter_func=dtensor_post_iter_func, diff --git a/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh b/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh index 4db32695088..22ab50424a1 100755 --- a/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh +++ b/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh @@ -37,6 +37,6 @@ if [[ $(jq 'to_entries | .[] | select(.key == "train/loss") | .value | keys | ma 'mean(data["train/gen_kl_error"]) < 0.001' \ 'data["train/gen_kl_error"]["30"] < 0.001 ' \ 'data["train/reward"]["30"] > 0.4' \ - 'data["train/grad_norm"] < 0.5' \ - 'data["train/grad_norm"] > 0.05' + 'data["train/grad_norm"]["30"] < 0.2' \ + 'data["train/grad_norm"]["30"] > 0.1' fi diff --git a/tests/unit/models/policy/test_dtensor_worker_v2.py b/tests/unit/models/policy/test_dtensor_worker_v2.py index 66409bd2056..0a257baa86c 100644 --- a/tests/unit/models/policy/test_dtensor_worker_v2.py +++ b/tests/unit/models/policy/test_dtensor_worker_v2.py @@ -32,6 +32,7 @@ try: from nemo_rl.models.policy.workers.dtensor_policy_worker_v2 import ( _maybe_adapt_tensor_to_hf, + dtensor_params_generator, get_train_context, ) @@ -566,6 +567,142 @@ def test_adapter_excludes_extra_state_regex(self): ) +@pytest.mark.automodel +@pytest.mark.skipif(not NEMO_AUTOMODEL_AVAILABLE, reason="nemo_automodel not available") +class TestDTensorParamsGenerator: + """Tests for the dtensor_params_generator helper function.""" + + def test_simple_model_yields_adapted_tensors(self): + """Test that generator yields correct (name, tensor) pairs for a simple model.""" + # Arrange + model = nn.Linear(10, 5) + target_dtype = torch.float32 + + # Act + results = list(dtensor_params_generator(model, target_dtype)) + + # Assert + assert len(results) == 2, "Linear layer should have weight and bias" + names = [name for name, _ in results] + assert "weight" in names + assert "bias" in names + + # Check that tensors are in the correct dtype and contiguous + for name, tensor in results: + assert tensor.dtype == target_dtype, ( + f"Tensor {name} should be {target_dtype}" + ) + assert tensor.is_contiguous(), f"Tensor {name} should be contiguous" + + def test_dtype_conversion(self): + """Test that tensors are converted to target dtype.""" + # Arrange + model = nn.Linear(10, 5) + # Initialize with float32 + model = model.to(torch.float32) + target_dtype = torch.bfloat16 + + # Act + results = list(dtensor_params_generator(model, target_dtype)) + + # Assert + for name, tensor in results: + assert tensor.dtype == target_dtype, ( + f"Tensor {name} should be converted to {target_dtype}" + ) + + def test_contiguous_output(self): + """Test that output tensors are contiguous.""" + # Arrange + model = nn.Linear(10, 5) + target_dtype = torch.float32 + + # Act + results = list(dtensor_params_generator(model, target_dtype)) + + # Assert + for name, tensor in results: + assert tensor.is_contiguous(), f"Tensor {name} should be contiguous" + + def test_with_adapter_model(self): + """Test that adapter is used when present on model.""" + # Arrange + model = nn.Linear(10, 5) + adapter_mock = Mock() + # Mock adapter to return multiple tensors for a single input + adapter_mock.convert_single_tensor_to_hf.return_value = [ + ("adapted.weight.1", torch.randn(5, 10)), + ("adapted.weight.2", torch.randn(5, 10)), + ] + model.state_dict_adapter = adapter_mock + target_dtype = torch.float32 + + # Act + results = list(dtensor_params_generator(model, target_dtype)) + + # Assert + # Each state_dict entry (weight, bias) goes through adapter + # Adapter returns 2 tensors for each, so we expect 4 total + assert len(results) >= 4, "Should have adapted tensors from adapter" + + # Verify adapter was called + assert adapter_mock.convert_single_tensor_to_hf.call_count >= 2 + + def test_empty_model(self): + """Test handling of model with no parameters.""" + # Arrange + model = nn.Module() # Empty module with no parameters + target_dtype = torch.float32 + + # Act + results = list(dtensor_params_generator(model, target_dtype)) + + # Assert + assert len(results) == 0, "Empty model should yield no parameters" + + def test_generator_is_iterable(self): + """Test that dtensor_params_generator returns an iterable generator.""" + # Arrange + model = nn.Linear(10, 5) + target_dtype = torch.float32 + + # Act + gen = dtensor_params_generator(model, target_dtype) + + # Assert + from collections.abc import Generator as ABCGenerator + + assert isinstance(gen, ABCGenerator), "Should return a generator" + + # Verify we can iterate over it + results = list(gen) + assert len(results) > 0, "Should yield at least one item" + + def test_multiple_layers(self): + """Test generator with a more complex model with multiple layers.""" + # Arrange + model = nn.Sequential( + nn.Linear(10, 5), + nn.ReLU(), + nn.Linear(5, 2), + ) + target_dtype = torch.float32 + + # Act + results = list(dtensor_params_generator(model, target_dtype)) + + # Assert + # Should have 4 parameters: 2 weights + 2 biases from the Linear layers + assert len(results) == 4, ( + "Sequential with 2 Linear layers should have 4 parameters" + ) + + # Check all tensors + for name, tensor in results: + assert tensor.dtype == target_dtype + assert tensor.is_contiguous() + + @pytest.mark.automodel @pytest.mark.skipif(not NEMO_AUTOMODEL_AVAILABLE, reason="nemo_automodel not available") class TestGetTrainContext: From 0a7d562b5c931474f09ddeac323c6cb383f38e42 Mon Sep 17 00:00:00 2001 From: Hemil Desai Date: Tue, 6 Jan 2026 18:26:21 +0000 Subject: [PATCH 09/12] fix Signed-off-by: Hemil Desai --- nemo_rl/models/policy/__init__.py | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/nemo_rl/models/policy/__init__.py b/nemo_rl/models/policy/__init__.py index 5cba9457a6c..f10da21f08b 100644 --- a/nemo_rl/models/policy/__init__.py +++ b/nemo_rl/models/policy/__init__.py @@ -35,6 +35,13 @@ class LoRAConfig(TypedDict): class AutomodelBackendConfig(TypedDict): + """Configuration for custom MoE implementation backend in Automodel. + + Used when setting the backend in automodel_kwargs in your config. + Alternatively, pass `force_hf: true` in automodel_kwargs to fall back + to the HuggingFace implementation. + """ + # Hydra target class path (e.g., "nemo_automodel.components.moe.utils.BackendConfig") _target_: str # Attention implementation: "te" (Transformer Engine), "flex" (FlexAttention), etc. @@ -60,6 +67,8 @@ class AutomodelKwargs(TypedDict): use_liger_kernel: NotRequired[bool] # Backend configuration for MoE models backend: NotRequired[AutomodelBackendConfig] + # Whether to force use of the HuggingFace implementation for MoE models + force_hf: NotRequired[bool] class DTensorConfigDisabled(TypedDict): From 7dd58086022d332d2a8b877fdeb57831fd09c53b Mon Sep 17 00:00:00 2001 From: Hemil Desai Date: Thu, 8 Jan 2026 02:35:04 +0000 Subject: [PATCH 10/12] fix Signed-off-by: Hemil Desai --- nemo_rl/models/policy/__init__.py | 1 + 1 file changed, 1 insertion(+) diff --git a/nemo_rl/models/policy/__init__.py b/nemo_rl/models/policy/__init__.py index f10da21f08b..1a934a26d4e 100644 --- a/nemo_rl/models/policy/__init__.py +++ b/nemo_rl/models/policy/__init__.py @@ -55,6 +55,7 @@ class AutomodelBackendConfig(TypedDict): # Use fake balanced gate for testing/debugging MoE fake_balanced_gate: NotRequired[bool] # Enable HuggingFace state dict adapter for checkpoint saving/loading plus refit support for RL + # This should almost always be set to True when using a custom MoE implementation. Set to False only for specific use cases like debugging or performance testing. enable_hf_state_dict_adapter: NotRequired[bool] # Enable FSDP-specific optimizations enable_fsdp_optimizations: NotRequired[bool] From 23b5525517ddf4cfe6c2cfeaeebf35c57b459e42 Mon Sep 17 00:00:00 2001 From: Hemil Desai Date: Thu, 8 Jan 2026 15:10:28 +0000 Subject: [PATCH 11/12] fix Signed-off-by: Hemil Desai --- tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh | 3 +++ 1 file changed, 3 insertions(+) diff --git a/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh b/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh index 22ab50424a1..fc98b560a51 100755 --- a/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh +++ b/tests/test_suites/llm/grpo-moonlight-16b-automodel-1n8g-ep8.sh @@ -39,4 +39,7 @@ if [[ $(jq 'to_entries | .[] | select(.key == "train/loss") | .value | keys | ma 'data["train/reward"]["30"] > 0.4' \ 'data["train/grad_norm"]["30"] < 0.2' \ 'data["train/grad_norm"]["30"] > 0.1' + + # Clean up checkpoint directory after successful run to save space. + rm -rf "$CKPT_DIR" fi From 3b64a9ceb06c331a491be2f73464a98ec167d356 Mon Sep 17 00:00:00 2001 From: Hemil Desai Date: Thu, 8 Jan 2026 16:43:42 -0800 Subject: [PATCH 12/12] fix Signed-off-by: Hemil Desai --- .../recipes/llm/grpo-moonlight-16b-automodel-1n8g-ep8.yaml | 6 ------ 1 file changed, 6 deletions(-) diff --git a/examples/configs/recipes/llm/grpo-moonlight-16b-automodel-1n8g-ep8.yaml b/examples/configs/recipes/llm/grpo-moonlight-16b-automodel-1n8g-ep8.yaml index 16d123d711d..d0fc72fbe99 100644 --- a/examples/configs/recipes/llm/grpo-moonlight-16b-automodel-1n8g-ep8.yaml +++ b/examples/configs/recipes/llm/grpo-moonlight-16b-automodel-1n8g-ep8.yaml @@ -47,12 +47,6 @@ policy: total_iters: 10000000000 - milestones: - 13 - generation: - _pad_token_id: 163838 - max_new_tokens: 4096 - pad_token_id: 163838 - stop_token_ids: - - 163585 data: max_input_seq_length: 4096 cluster: