You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
{{ message }}
Repository navigation
[BUG] torchrl uses tensordict.utils.Buffer as torch.nn.Buffer, but it is an empty class on torch>=2.5 #4535
torchrl imports Buffer from tensordict.utils in three modules:
torchrl/collectors/utils.py, in _cast and _map_weight
torchrl/weight_update/weight_sync_schemes.py, in _to_module_compatible_tensor
torchrl/objectives/common.py, in _make_target_param
This code expects Buffer to be torch.nn.Buffer. It is not. On torch 2.5 and later, tensordict.utils.Buffer is an empty placeholder class. On torch older than 2.5, it is a copy of torch.nn.Buffer.
On torch 2.5 and later, this has these results:
isinstance(x, Buffer) is always False. The branches that keep a torch.nn.Buffer as a buffer never run, so the three helpers return a plain tensor.
Buffer(x) raises TypeError: Buffer() takes no arguments. No code path reaches these calls now, because of result 1.
_cast raises RuntimeError: Cannot cast tensor ... with gradients for a torch.nn.Buffer that requires grad. When Buffer is torch.nn.Buffer, _cast returns a buffer instead.
I did not find a user-facing failure. In WeightStrategy.apply_weights and in the policy device mapping in collectors/_base.py, the buffers stay registered. The reason is that TensorDict.to_module keeps an existing buffer as a buffer. But the buffer handling in these helpers does not run.
To Reproduce
Run this script with torch 2.5 or later:
importtorchfromtorchimportnnfromtensordict.utilsimportBufferfromtorchrl.collectors.utilsimport_cast, _map_weightfromtorchrl.weight_update.weight_sync_schemesimport_to_module_compatible_tensor# 1. tensordict.utils.Buffer is not torch.nn.Buffer.print(Bufferisnn.Buffer) # Falseprint(isinstance(nn.Buffer(torch.zeros(3)), Buffer)) # False# 2. The torchrl helpers return a plain tensor for a torch.nn.Buffer.buf=nn.Buffer(torch.zeros(3))
print(isinstance(_cast(buf.data, buf), nn.Buffer)) # Falseprint(isinstance(_map_weight(buf, torch.device("cpu")), nn.Buffer)) # Falseprint(isinstance(_to_module_compatible_tensor(torch.zeros(3), buf), nn.Buffer)) # False# 3. Calling tensordict.utils.Buffer fails.Buffer(torch.zeros(3))
On torch 2.5 and later, torchrl uses torch.nn.Buffer, and the calls in step 2 print True. I tested this by setting Buffer = torch.nn.Buffer in the three torchrl modules.
System info
torch 2.14.1+cpu, from the PyTorch CPU index with pip
tensordict 0.14.2 and torchrl 0.14.0, from PyPI with pip
[Versioning] Require torch 2.13 or later tensordict#1830 (open) makes tensordict require torch 2.13 or later. That PR keeps the placeholder, so that tensordict does not change the behavior of torchrl. After torchrl stops importing tensordict.utils.Buffer, tensordict can remove the placeholder.
_make_target_param has a second problem. It calls x.data before isinstance(x, nn.Parameter), and x.data is never a Parameter. So Buffer(x) does not run there, even with torch.nn.Buffer.
Reason and Possible fixes
Import Buffer from torch when it exists. torchrl supports torch 2.1 and later, so keep the tensordict copy for torch older than 2.5:
try:
fromtorch.nnimportBuffer# torch>=2.5exceptImportError: # torch<2.5, where tensordict provides a copyfromtensordict.utilsimportBuffer
This change makes the helpers keep torch.nn.Buffer inputs as buffers, so their output changes. The script above can become a regression test.
Thanks @peterdsharpe , I'm thinking that we first pin torch 2.13 in torchRL and then merge this fix so that we don't need to import the tensordict Buffer at all anymore. Does that work for you?
Describe the bug
torchrl imports
Bufferfromtensordict.utilsin three modules:torchrl/collectors/utils.py, in_castand_map_weighttorchrl/weight_update/weight_sync_schemes.py, in_to_module_compatible_tensortorchrl/objectives/common.py, in_make_target_paramThis code expects
Bufferto betorch.nn.Buffer. It is not. On torch 2.5 and later,tensordict.utils.Bufferis an empty placeholder class. On torch older than 2.5, it is a copy oftorch.nn.Buffer.On torch 2.5 and later, this has these results:
isinstance(x, Buffer)is alwaysFalse. The branches that keep atorch.nn.Bufferas a buffer never run, so the three helpers return a plain tensor.Buffer(x)raisesTypeError: Buffer() takes no arguments. No code path reaches these calls now, because of result 1._castraisesRuntimeError: Cannot cast tensor ... with gradientsfor atorch.nn.Bufferthat requires grad. WhenBufferistorch.nn.Buffer,_castreturns a buffer instead.I did not find a user-facing failure. In
WeightStrategy.apply_weightsand in the policy device mapping incollectors/_base.py, the buffers stay registered. The reason is thatTensorDict.to_modulekeeps an existing buffer as a buffer. But the buffer handling in these helpers does not run.To Reproduce
Run this script with torch 2.5 or later:
Output:
Expected behavior
On torch 2.5 and later, torchrl uses
torch.nn.Buffer, and the calls in step 2 printTrue. I tested this by settingBuffer = torch.nn.Bufferin the three torchrl modules.System info
mainat 648f50c.Additional context
tensordict.utils.Buffer, tensordict can remove the placeholder._make_target_paramhas a second problem. It callsx.databeforeisinstance(x, nn.Parameter), andx.datais never aParameter. SoBuffer(x)does not run there, even withtorch.nn.Buffer.Reason and Possible fixes
Import
Bufferfrom torch when it exists. torchrl supports torch 2.1 and later, so keep the tensordict copy for torch older than 2.5:This change makes the helpers keep
torch.nn.Bufferinputs as buffers, so their output changes. The script above can become a regression test.The code that uses
Buffer:rl/torchrl/collectors/utils.py
Lines 328 to 343 in 648f50c
rl/torchrl/collectors/utils.py
Lines 643 to 659 in 648f50c
rl/torchrl/weight_update/weight_sync_schemes.py
Lines 131 to 142 in 648f50c
rl/torchrl/objectives/common.py
Lines 1030 to 1038 in 648f50c
Checklist