Skip to content

[BUG] torchrl uses tensordict.utils.Buffer as torch.nn.Buffer, but it is an empty class on torch>=2.5 #4535

Description

@peterdsharpe

Describe the bug

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:

  1. 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.
  2. Buffer(x) raises TypeError: Buffer() takes no arguments. No code path reaches these calls now, because of result 1.
  3. _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:

import torch
from torch import nn
from tensordict.utils import Buffer
from torchrl.collectors.utils import _cast, _map_weight
from torchrl.weight_update.weight_sync_schemes import _to_module_compatible_tensor

# 1. tensordict.utils.Buffer is not torch.nn.Buffer.
print(Buffer is nn.Buffer)  # False
print(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))  # False
print(isinstance(_map_weight(buf, torch.device("cpu")), nn.Buffer))  # False
print(isinstance(_to_module_compatible_tensor(torch.zeros(3), buf), nn.Buffer))  # False

# 3. Calling tensordict.utils.Buffer fails.
Buffer(torch.zeros(3))

Output:

False
False
False
False
False
Traceback (most recent call last):
  File "repro.py", line 18, in <module>
    Buffer(torch.zeros(3))
TypeError: Buffer() takes no arguments

Expected behavior

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
  • Python 3.12.12 on Linux
  • The same code is on torchrl main at 648f50c.
>>> import torchrl, numpy, sys
>>> print(torchrl.__version__, numpy.__version__, sys.version, sys.platform)
0.14.0 2.5.3 3.12.12 (main, Oct 28 2025, 12:10:49) [Clang 20.1.4 ] linux

Additional context

  • [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:
    from torch.nn import Buffer  # torch>=2.5
except ImportError:  # torch<2.5, where tensordict provides a copy
    from tensordict.utils import Buffer

This change makes the helpers keep torch.nn.Buffer inputs as buffers, so their output changes. The script above can become a regression test.

The code that uses Buffer:

@implement_for("torch", "2.5.0")
def _cast(
p: nn.Parameter | torch.Tensor,
param_maybe_buffer: nn.Parameter | torch.Tensor | None = None,
) -> nn.Parameter | torch.Tensor:
if param_maybe_buffer is None:
param_maybe_buffer = p
p = p.data
if isinstance(param_maybe_buffer, Parameter):
# Create parameter without gradients to avoid serialization issues
return Parameter(p, requires_grad=False)
if isinstance(param_maybe_buffer, Buffer):
return Buffer(p)
if p.requires_grad:
raise RuntimeError(f"Cannot cast tensor {p} with gradients")
return p

def _map_weight(
weight,
policy_device,
):
is_param = isinstance(weight, Parameter)
is_buffer = isinstance(weight, Buffer)
weight = weight.data
if weight.device != policy_device:
weight = weight.to(policy_device)
elif weight.device.type in ("cpu",):
weight = weight.share_memory_()
if is_param:
weight = Parameter(weight, requires_grad=False)
elif is_buffer:
weight = Buffer(weight)
return weight

def _to_module_compatible_tensor(
tensor: torch.Tensor,
destination_tensor: torch.Tensor,
) -> torch.Tensor:
"""Cast a tensor so ``TensorDict.to_module`` preserves module registration."""
if isinstance(tensor, (nn.Parameter, Buffer)):
return tensor
if isinstance(destination_tensor, nn.Parameter):
return nn.Parameter(tensor, requires_grad=destination_tensor.requires_grad)
if isinstance(destination_tensor, Buffer):
return Buffer(tensor)
return tensor

class _make_target_param:
def __init__(self, clone):
self.clone = clone
def __call__(self, x):
x = x.data.clone() if self.clone else x.data
if isinstance(x, nn.Parameter):
return Buffer(x)
return x

Checklist

  • I have checked that there is no similar issue in the repo (required)
  • I have read the documentation (required)
  • I have provided a minimal working example to reproduce the bug (required)

Activity

  1. ahlag commented on Oct 8, 2026

    @ahlag

    /assign

  2. added a commit that references this issue on Oct 9, 2026
    2dd48bb
  3. peterdsharpe commented on Oct 9, 2026

    @peterdsharpe
    Author

    After a fix for this merges, we can mark Buffer for deprecation in the tensordict side - see #4535

  4. bsprenger commented on Oct 10, 2026

    @bsprenger
    Collaborator

    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?

  5. peterdsharpe commented on Oct 10, 2026

    @peterdsharpe
    Author

    @bsprenger yep! Sounds good to me :)

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

Labels

No labels
No labels

Type

No type

Projects

No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions