Qwen3-32B RL with DTensor backend generate the error:"flash_attn._flash_attn_varlen_forward.default: got mixed torch.Tensor and DTensor, need to convert all torch.Tensor to DTensor before calling distributed operators"
- Dominant language
- Python
- Stars
- 2k
- Forks
- 561
- Avg merge
- 4d 5h
- Merged PRs (30d)
- 145
Description
**Describe the bug**
With the DTensor backend to RL Qwen3-32B. It will have the errors:
```bash
flash_attn._flash_attn_varlen_forward.default: got mixed torch.Tensor and DTensor, need to convert all torch.Tensor to DTensor before calling distributed operators
```
However, the **Megatron backend** does not have this issues.
**Steps/Code to reproduce bug**
Build the container with the main branch. And submit the slurm job with 8 nodes of H100*8.
The job-submit script is as follows:
```bash
#!/bin/bash
NUM_ACTOR_NODES=8
MODEL_PATH="/lustre/fs1/portfolios/coreai/users/jzhai/RL_NeMo_RL/Qwen3-32B"
VLLM_TP=4
TRAIN_TP=8
read -r -d '' COMMAND <
main()
File "/lustre/fs1/portfolios/coreai/users/jzhai/RL_NeMo_RL/nemo-rl/./examples/run_grpo_math.py", line 255, in main
grpo_train(
File "/lustre/fs1/portfolios/coreai/users/jzhai/RL_NeMo_RL/nemo-rl/nemo_rl/algorithms/grpo.py", line 698, in grpo_train
fprop_logprobs = policy.get_logprobs(train_data)["logprobs"]
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/lustre/fs1/portfolios/coreai/users/jzhai/RL_NeMo_RL/nemo-rl/nemo_rl/models/policy/lm_policy.py", line 259, in get_logprobs
self.worker_group.get_all_worker_results(futures)
File "/lustre/fs1/portfolios/coreai/users/jzhai/RL_NeMo_RL/nemo-rl/nemo_rl/distributed/worker_groups.py", line 903, in get_all_worker_results
return future_bundle.get_results(
^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/lustre/fs1/portfolios/coreai/users/jzhai/RL_NeMo_RL/nemo-rl/nemo_rl/distributed/worker_groups.py", line 99, in get_results
all_results = ray.get(object_refs)
^^^^^^^^^^^^^^^^^^^^
File "/opt/nemo_rl_venv/lib/python3.12/site-packages/ray/_private/auto_init_hook.py", line 21, in auto_init_wrapper
return fn(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^
File "/opt/nemo_rl_venv/lib/python3.12/site-packages/ray/_private/client_mode_hook.py", line 103, in wrapper
return func(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^
File "/opt/nemo_rl_venv/lib/python3.12/site-packages/ray/_private/worker.py", line 2822, in get
values, debugger_breakpoint = worker.get_objects(object_refs, timeout=timeout)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/nemo_rl_venv/lib/python3.12/site-packages/ray/_private/worker.py", line 930, in get_objects
raise value.as_instanceof_cause()
ray.exceptions.RayTaskError(RuntimeError): ray::DTensorPolicyWorkerV2.get_logprobs() (pid=458879, ip=10.65.12.205, actor_id=af5e333786de88cedf3a674901000000, repr=DTensorPolicyWorkerV2[rank=43])
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/lustre/fs1/portfolios/coreai/users/jzhai/RL_NeMo_RL/nemo-rl/nemo_rl/utils/nsys.py", line 88, in wrapper
ret = func(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^
File "/lustre/fs1/portfolios/coreai/users/jzhai/RL_NeMo_RL/nemo-rl/nemo_rl/models/policy/dtensor_policy_worker_v2.py", line 921, in get_logprobs
outputs = self.model(
^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1751, in _wrapped_call_impl
return self._call_impl(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1857, in _call_impl
return inner()
^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1805, in inner
result = forward_call(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/nemo-rl/3rdparty/Automodel-workspace/Automodel/nemo_automodel/components/_transformers/auto_model.py", line 80, in wrapper
return func(self, *args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/transformers/utils/generic.py", line 943, in wrapper
output = func(self, *args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/transformers/models/qwen3/modeling_qwen3.py", line 570, in forward
outputs: BaseModelOutputWithPast = self.model(
^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1751, in _wrapped_call_impl
return self._call_impl(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1762, in _call_impl
return forward_call(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/transformers/utils/generic.py", line 943, in wrapper
output = func(self, *args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/transformers/models/qwen3/modeling_qwen3.py", line 458, in forward
layer_outputs = decoder_layer(
^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/transformers/modeling_layers.py", line 83, in __call__
return super().__call__(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1751, in _wrapped_call_impl
return self._call_impl(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1857, in _call_impl
return inner()
^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1805, in inner
result = forward_call(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/transformers/models/qwen3/modeling_qwen3.py", line 262, in forward
hidden_states, self_attn_weights = self.self_attn(
^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1751, in _wrapped_call_impl
return self._call_impl(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1762, in _call_impl
return forward_call(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/transformers/models/qwen3/modeling_qwen3.py", line 217, in forward
attn_output, attn_weights = attention_interface(
^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/transformers/integrations/flash_attention.py", line 65, in flash_attention_forward
attn_output = _flash_attention_forward(
^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/transformers/modeling_flash_attention_utils.py", line 560, in _flash_attention_forward
attn_output = _flash_attn_varlen_func(
^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/flash_attn/flash_attn_interface.py", line 1448, in flash_attn_varlen_func
return FlashAttnVarlenFunc.apply(
^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/autograd/function.py", line 575, in apply
return super().apply(*args, **kwargs) # type: ignore[misc]
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/flash_attn/flash_attn_interface.py", line 930, in forward
out_padded, softmax_lse, S_dmask, rng_state = _wrapped_flash_attn_varlen_forward(
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/_ops.py", line 1158, in __call__
return self._op(*args, **(kwargs or {}))
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/_library/autograd.py", line 113, in autograd_impl
result = forward_no_grad(*args, Metadata(keyset, keyword_only_args))
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/_library/autograd.py", line 40, in forward_no_grad
result = op.redispatch(keyset & _C._after_autograd_keyset, *args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/_ops.py", line 761, in redispatch
return self._handle.redispatch_boxed(keyset, *args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/_compile.py", line 51, in inner
return disable_fn(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/_dynamo/eval_frame.py", line 838, in _fn
return fn(*args, **kwargs)
^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/distributed/tensor/_api.py", line 344, in __torch_dispatch__
return DTensor._op_dispatcher.dispatch(
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/distributed/tensor/_dispatch.py", line 167, in dispatch
op_info = self.unwrap_to_op_info(op_call, args, kwargs)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/distributed/tensor/_dispatch.py", line 366, in unwrap_to_op_info
self._try_replicate_spec_for_scalar_tensor(
File "/opt/ray_venvs/nemo_rl.models.policy.dtensor_policy_worker_v2.DTensorPolicyWorkerV2/lib/python3.12/site-packages/torch/distributed/tensor/_dispatch.py", line 468, in _try_replicate_spec_for_scalar_tensor
raise RuntimeError(
RuntimeError: flash_attn._flash_attn_varlen_forward.default: got mixed torch.Tensor and DTensor, need to convert all torch.Tensor to DTensor before calling distributed operators!
(VllmGenerationWorker pid=602061) INFO 09-03 01:41:48 [block_pool.py:321] Successfully reset prefix cache [repeated 30x across cluster]
(RayWorkerWrapper pid=1864505, ip=10.65.12.207) INFO 09-03 01:41:50 [gpu_worker.py:104] Sleep mode freed 40.20 GiB memory, 6.26 GiB memory is still in use. [repeated 63x across cluster]
(VllmGenerationWorker pid=1863449, ip=10.65.12.207) INFO 09-03 01:41:50 [executor_base.py:187] It took 1.978376 seconds to fall asleep. [repeated 15x across cluster]
2025-09-03 01:42:25,916 INFO worker.py:1694 -- Connecting to existing Ray cluster at address: 10.65.12.145:54514...
2025-09-03 01:42:25,920 INFO worker.py:1879 -- Connected to Ray cluster. View the dashboard at http://127.0.0.1:8265
2025-09-03 01:42:25,981 INFO worker.py:1694 -- Connecting to existing Ray cluster at address: 10.65.12.145:54514...
2025-09-03 01:42:25,984 INFO worker.py:1879 -- Connected to Ray cluster. View the dashboard at http://127.0.0.1:8265
[2025-09-03 01:42:26,011 C 576323 576323] core_worker_process.cc:69: Check failed: !core_worker_process The process is already initialized for core worker.
*** StackTrace Information ***
/opt/nemo_rl_venv/lib/python3.12/site-packages/ray/_raylet.so(+0x14392da) [0x1553c96652da] ray::operator<<()
/opt/nemo_rl_venv/lib/python3.12/site-packages/ray/_raylet.so(_ZN3ray6RayLogD1Ev+0x479) [0x1553c9667d59] ray::RayLog::~RayLog()
/opt/nemo_rl_venv/lib/python3.12/site-packages/ray/_raylet.so(_ZN3ray4core17CoreWorkerProcess10InitializeERKNS0_17CoreWorkerOptionsE+0xe8) [0x1553c8bd7e48] ray::core::CoreWorkerProcess::Initialize()
/opt/nemo_rl_venv/lib/python3.12/site-packages/ray/_raylet.so(+0x81af6e) [0x1553c8a46f6e] __pyx_pf_3ray_7_raylet_10CoreWorker___cinit__()
/opt/nemo_rl_venv/lib/python3.12/site-packages/ray/_raylet.so(+0x81c828) [0x1553c8a48828] __pyx_pw_3ray_7_raylet_10CoreWorker_1__cinit__()
/opt/nemo_rl_venv/lib/python3.12/site-packages/ray/_raylet.so(+0x81d26c) [0x1553c8a4926c] __pyx_tp_new_3ray_7_raylet_CoreWorker()
/root/.local/share/uv/python/cpython-3.12.10-linux-x86_64-gnu/bin/../lib/libpython3.12.so.1.0(+0x3b3d79) [0x1555540d7d79] type_call
/root/.local/share/uv/python/cpython-3.12.10-linux-x86_64-gnu/bin/../lib/libpython3.12.so.1.0(_PyEval_EvalFrameDefault+0x31f0f) [0x1555541377df] _PyEval_EvalFrameDefault
/root/.local/share/uv/python/cpython-3.12.10-linux-x86_64-gnu/bin/../lib/libpython3.12.so.1.0(PyObject_CallOneArg+0x61) [0x155554098dd1] PyObject_CallOneArg
/root/.local/share/uv/python/cpython-3.12.10-linux-x86_64-gnu/bin/../lib/libpython3.12.so.1.0(+0x4c05b8) [0x1555541e45b8] slot_tp_finalize
/root/.local/share/uv/python/cpython-3.12.10-linux-x86_64-gnu/bin/../lib/libpython3.12.so.1.0(+0x3baf75) [0x1555540def75] subtype_dealloc
/root/.local/share/uv/python/cpython-3.12.10-linux-x86_64-gnu/bin/../lib/libpython3.12.so.1.0(+0x37ca29) [0x1555540a0a29] frame_dealloc
/root/.local/share/uv/python/cpython-3.12.10-linux-x86_64-gnu/bin/../lib/libpython3.12.so.1.0(+0x43a16d) [0x15555415e16d] tb_dealloc
/root/.local/share/uv/python/cpython-3.12.10-linux-x86_64-gnu/bin/../lib/libpython3.12.so.1.0(+0x43a161) [0x15555415e161] tb_dealloc
/root/.local/share/uv/python/cpython-3.12.10-linux-x86_64-gnu/bin/../lib/libpython3.12.so.1.0(PyDict_SetItem+0x64f) [0x1555540ba6ef] PyDict_SetItem
/root/.local/share/uv/python/cpython-3.12.10-linux-x86_64-gnu/bin/../lib/libpython3.12.so.1.0(+0x558252) [0x15555427c252] _PySys_ClearAttrString
/root/.local/share/uv/python/cpython-3.12.10-linux-x86_64-gnu/bin/../lib/libpython3.12.so.1.0(+0x54645f) [0x15555426a45f] finalize_modules.llvm.6956522194820775225
/root/.local/share/uv/python/cpython-3.12.10-linux-x86_64-gnu/bin/../lib/libpython3.12.so.1.0(Py_FinalizeEx+0xf1) [0x155554269a21] Py_FinalizeEx
/root/.local/share/uv/python/cpython-3.12.10-linux-x86_64-gnu/bin/../lib/libpython3.12.so.1.0(Py_RunMain+0x183) [0x15555428c2e3] Py_RunMain
/root/.local/share/uv/python/cpython-3.12.10-linux-x86_64-gnu/bin/../lib/libpython3.12.so.1.0(+0x56898f) [0x15555428c98f] pymain_main
/root/.local/share/uv/python/cpython-3.12.10-linux-x86_64-gnu/bin/../lib/libpython3.12.so.1.0(Py_BytesMain+0x2d) [0x15555428ca4d] Py_BytesMain
/usr/lib/x86_64-linux-gnu/libc.so.6(+0x2a1ca) [0x155553a361ca]
/usr/lib/x86_64-linux-gnu/libc.so.6(__libc_start_main+0x8b) [0x155553a3628b] __libc_start_main
/opt/nemo_rl_venv/bin/python3(_start+0x29) [0x6000a9] _start
```
Contributor guide
Assessment
This issue has not been assessed yet.