NVIDIA-NeMo / NVIDIA-NeMo/RL

Gemma3 27B crashes on 16 nodes w/ FSDP

Open
#1,086 1 comment 0 reactions 1 assignee Claimed by @joyang-nv View on GitHub
bug t-pytdensor
Dominant language
Python
Stars
2k
Forks
561
Avg merge
4d 5h
Merged PRs (30d)
145

Description

## 16 nodes full w/ FSDP (v1 and v2 dtensor backends)

Gemma3 27B doesn't work on 16 nodes anymore

https://github.com/NVIDIA-NeMo/RL/blob/ae89e126bedd57020dcb3b5173e393bac287d634/examples/configs/recipes/llm/grpo-gemma3-27b-it-16n8g-fsdp2tp8sp-actckpt-long.yaml

I believe it's due to this PR https://github.com/NVIDIA-NeMo/RL/commit/9f7825ec805863a4651d6337a186b2cb92ee8167 which added the input/output embedding to the parallelize plan whereas before it was replicated.

repro: `uv run tests/test_suites/llm/grpo-gemma3-27b-it-16n8g-fsdp2tp8sp-actckpt-long.sh`

```
(DTensorPolicyWorkerV2 pid=1276099, ip=10.65.6.79) ^^^^^^^^^^^^^^^^^^^^^^^^^^^ [repeated 24x across cluster]
(DTensorPolicyWorkerV2 pid=1276099, ip=10.65.6.79) File "/code_snapshots/grpo-gemma3-27b-it-16n8g-fsdp2tp8sp-actckpt-long/3rdparty/Automodel-workspace/Automodel/nemo_automodel/components/distributed/parallelizer.py", line 589, in fsdp2_strategy_parallelize [repeated 24x across cluster]
(DTensorPolicyWorkerV2 pid=1276099, ip=10.65.6.79) model = fully_shard( [repeated 24x across cluster]
(DTensorPolicyWorkerV2 pid=1276099, ip=10.65.6.79) ^^^^^^^^^^^^ [repeated 24x across cluster]
(DTensorPolicyWorkerV2 pid=1276099, ip=10.65.6.79) File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/distributed/_composable/contract.py", line 150, in wrapper [repeated 24x across cluster]
(DTensorPolicyWorkerV2 pid=1276099, ip=10.65.6.79) updated = func(inp_module, *args, **kwargs) [repeated 24x across cluster]
(DTensorPolicyWorkerV2 pid=1276099, ip=10.65.6.79) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ [repeated 24x across cluster]
(DTensorPolicyWorkerV2 pid=1276099, ip=10.65.6.79) File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/distributed/fsdp/_fully_shard/_fully_shard.py", line 217, in fully_shard [repeated 24x across cluster]
(DTensorPolicyWorkerV2 pid=1276099, ip=10.65.6.79) state._fsdp_param_group = FSDPParamGroup( [repeated 24x across cluster]
(DTensorPolicyWorkerV2 pid=1276099, ip=10.65.6.79) ^^^^^^^^^^^^^^^ [repeated 24x across cluster]
(DTensorPolicyWorkerV2 pid=1276099, ip=10.65.6.79) FSDPParam( [repeated 24x across cluster]
(DTensorPolicyWorkerV2 pid=1276099, ip=10.65.6.79) self._init_sharded_param(param, device, shard_placement_fn) [repeated 24x across cluster]
(DTensorPolicyWorkerV2 pid=1276099, ip=10.65.6.79) File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/utils/_contextlib.py", line 116, in decorate_context [repeated 24x across cluster]
(DTensorPolicyWorkerV2 pid=1276099, ip=10.65.6.79) return func(*args, **kwargs) [repeated 24x across cluster]
(DTensorPolicyWorkerV2 pid=1276099, ip=10.65.6.79) ^^^^^^^^^^^^^^^^^^^^^ [repeated 24x across cluster]
(DTensorPolicyWorkerV2 pid=1276099, ip=10.65.6.79) File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/distributed/fsdp/_fully_shard/_fsdp_param.py", line 335, in _init_sharded_param [repeated 24x across cluster]
(DTensorPolicyWorkerV2 pid=1276099, ip=10.65.6.79) raise NotImplementedError( [repeated 24x across cluster]
(DTensorPolicyWorkerV2 pid=1276099, ip=10.65.6.79) NotImplementedError: FSDP+TP sharding does not support uneven sharding for now: tensor dim 0 has size 262208 which cannot be evenly sharded into 128 shards. [repeated 24x across cluster]
```

https://wandb.ai/nvidia/nemo-rl/runs/z55enf3w

Also filed against Automodel: https://github.com/NVIDIA-NeMo/Automodel/issues/425

Would HSDP help?

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.