NVIDIA-NeMo / NVIDIA-NeMo/RL

Investigate cuda base version skew with stable torch versions

Open
#486 2 comments 0 reactions 1 assignee Claimed by @chtruong814 View on GitHub
Dominant language
Python
Stars
2k
Forks
561
Avg merge
4d 5h
Merged PRs (30d)
145

Description

https://wandb.ai/nvidia/reinforcer-nemo-ci?nw=i6iqej8ouh

![Image](https://github.com/user-attachments/assets/1a3e29fb-6445-49dd-8381-135749879ecb)

There are several base images we can choose from:
* nvcr cuda (status quo)
* NGC Pytorch
* nvcr cuda-dl-base

The issue uncovered recently is the nvcr cuda image on `main` results is very slow gathers and overall poor e2e perf in our container. Switching to nvcr cuda-dl-base recovers performance, but some combinations of stable torch (e.g., 2.6.0/2.7.0) don't work with some cuda bases despite being the same base cuda version (12.8 or 12.9).

An example is shown above where some base containers crash at step 10 because of a checkpointing error
```
File "/tmp/ray/session_2025-06-04_12-24-09_346850_421975/runtime_resources/working_dir_files/_ray_pkg_93b84cc5752eb70b/nemo_rl/models/policy/dtensor_policy_worker.py", line 837, in save_checkpoint
save_checkpoint(
File "/tmp/ray/session_2025-06-04_12-24-09_346850_421975/runtime_resources/working_dir_files/_ray_pkg_93b84cc5752eb70b/nemo_rl/utils/native_checkpoint.py", line 157, in save_checkpoint
dcp.save(model_state, checkpoint_id=weights_path)
File "/cicd-workspace/tk-2025-06-03-cuda-dl-base-12.9-25.05-grpo-llama-8b-convergence/venvs/nemo_rl.models.policy.dtensor_policy_worker.DTensorPolicyWorker/lib/python3.12/site-packages/torch/distributed/checkpoint/logger.py", line 83, in wrapper
result = func(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^
File "/cicd-workspace/tk-2025-06-03-cuda-dl-base-12.9-25.05-grpo-llama-8b-convergence/venvs/nemo_rl.models.policy.dtensor_policy_worker.DTensorPolicyWorker/lib/python3.12/site-packages/torch/distributed/checkpoint/utils.py", line 429, in inner_func
return func(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^
File "/cicd-workspace/tk-2025-06-03-cuda-dl-base-12.9-25.05-grpo-llama-8b-convergence/venvs/nemo_rl.models.policy.dtensor_policy_worker.DTensorPolicyWorker/lib/python3.12/site-packages/torch/distributed/checkpoint/state_dict_saver.py", line 152, in save
return _save_state_dict(
^^^^^^^^^^^^^^^^^
File "/cicd-workspace/tk-2025-06-03-cuda-dl-base-12.9-25.05-grpo-llama-8b-convergence/venvs/nemo_rl.models.policy.dtensor_policy_worker.DTensorPolicyWorker/lib/python3.12/site-packages/torch/distributed/checkpoint/state_dict_saver.py", line 317, in _save_state_dict
central_plan: SavePlan = distW.reduce_scatter("plan", local_step, global_step)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/cicd-workspace/tk-2025-06-03-cuda-dl-base-12.9-25.05-grpo-llama-8b-convergence/venvs/nemo_rl.models.policy.dtensor_policy_worker.DTensorPolicyWorker/lib/python3.12/site-packages/torch/distributed/checkpoint/utils.py", line 168, in reduce_scatter
all_data = self.gather_object(local_data)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/cicd-workspace/tk-2025-06-03-cuda-dl-base-12.9-25.05-grpo-llama-8b-convergence/venvs/nemo_rl.models.policy.dtensor_policy_worker.DTensorPolicyWorker/lib/python3.12/site-packages/torch/distributed/checkpoint/utils.py", line 107, in gather_object
dist.gather_object(
File "/cicd-workspace/tk-2025-06-03-cuda-dl-base-12.9-25.05-grpo-llama-8b-convergence/venvs/nemo_rl.models.policy.dtensor_policy_worker.DTensorPolicyWorker/lib/python3.12/site-packages/torch/distributed/c10d_logger.py", line 81, in wrapper
return func(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^
File "/cicd-workspace/tk-2025-06-03-cuda-dl-base-12.9-25.05-grpo-llama-8b-convergence/venvs/nemo_rl.models.policy.dtensor_policy_worker.DTensorPolicyWorker/lib/python3.12/site-packages/torch/distributed/distributed_c10d.py", line 3177, in gather_object
object_gather_list[i] = _tensor_to_object(tensor, tensor_size, group)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/cicd-workspace/tk-2025-06-03-cuda-dl-base-12.9-25.05-grpo-llama-8b-convergence/venvs/nemo_rl.models.policy.dtensor_policy_worker.DTensorPolicyWorker/lib/python3.12/site-packages/torch/distributed/distributed_c10d.py", line 2961, in _tensor_to_object
return _unpickler(io.BytesIO(buf)).load()
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
_pickle.UnpicklingError: invalid load key, '\x00'
```

and some will randomly cause the dtensor workers to crash inexplicably during training:
```
/....nemo_rl/algorithms/grpo.py", line 516, in grpo_train
2025-06-04 21:17:18 train_results = policy.train(train_data, loss_fn)
```

We need to track these version skews and document this clearly.

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.