Investigate cuda base version skew with stable torch versions
- 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

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
Assessment
This issue has not been assessed yet.