pytorch / pytorch/rl

[BUG] A2C fails with functional=True and shifted=True for ValueEstimator

Open
#2,265 5 comments 0 reactions 1 assignee View on GitHub

@vmoens is already working on this.

Since Jul 2, 2024.

bug
Dominant language
Python
Stars
3.6k
Forks
484
Avg merge
1d 1h
Merged PRs (30d)
207

Description

Describe the bug

Not quite sure if this is supported behavior, but if I set functional=True for the A2C loss and shifted=True for TD0Estimator I get an internal error.

To Reproduce

import gymnasium as gym
import torchrl.envs
import torch
import torchrl
from torchrl.objectives import ValueEstimators
from torchrl.objectives.value import TD0Estimator
from torchrl.modules import MLP, ValueOperator, ProbabilisticActor, Actor

time_dim = 4

    gym_env = torchrl.envs.GymEnv("MountainCar-v0", device="cpu")
    observation_shape = gym_env.observation_spec["observation"].shape[0]

    actor_net_mock = torch.nn.Linear(
        in_features=observation_shape,
        out_features=gym_env.action_spec.shape[-1],
    )

    value_net_mock = torch.nn.Linear(
        in_features=observation_shape,
        out_features=1,
    )
    probabilistic_actor = ProbabilisticActor(
        module=Actor(
            actor_net_mock,out_keys=["logits"]
        ),
        in_keys=["logits"],
        distribution_class=torch.distributions.OneHotCategorical,
    )
    value_operator = ValueOperator(module=value_net_mock, in_keys=["observation"])

    rollout = gym_env.rollout(max_steps=time_dim, policy=probabilistic_actor)
    loss = torchrl.objectives.a2c.A2CLoss(
        probabilistic_actor,
        value_operator,
        functional=True,
    )
    loss.make_value_estimator(ValueEstimators.TD0, gamma=0.9, shifted=True)

    rollout_loss = loss(rollout)
../venv-nightly/lib/python3.11/site-packages/torch/nn/modules/module.py:1657: in _wrapped_call_impl
    return self._call_impl(*args, **kwargs)
../venv-nightly/lib/python3.11/site-packages/torch/nn/modules/module.py:1709: in _call_impl
    result = forward_call(*args, **kwargs)
../venv-nightly/lib/python3.11/site-packages/tensordict/_contextlib.py:126: in decorate_context
    return func(*args, **kwargs)
../venv-nightly/lib/python3.11/site-packages/tensordict/nn/common.py:289: in wrapper
    return func(_self, tensordict, *args, **kwargs)
../venv-nightly/lib/python3.11/site-packages/torchrl/objectives/a2c.py:470: in forward
    self.value_estimator(
../venv-nightly/lib/python3.11/site-packages/torch/nn/modules/module.py:1657: in _wrapped_call_impl
    return self._call_impl(*args, **kwargs)
../venv-nightly/lib/python3.11/site-packages/torch/nn/modules/module.py:1668: in _call_impl
    return forward_call(*args, **kwargs)
../venv-nightly/lib/python3.11/site-packages/torchrl/objectives/value/advantages.py:68: in new_func
    return fun(self, *args, **kwargs)
../venv-nightly/lib/python3.11/site-packages/torchrl/objectives/value/advantages.py:57: in new_fun
    return fun(self, *args, **kwargs)
../venv-nightly/lib/python3.11/site-packages/tensordict/nn/common.py:289: in wrapper
    return func(_self, tensordict, *args, **kwargs)
../venv-nightly/lib/python3.11/site-packages/torchrl/objectives/value/advantages.py:632: in forward
    value, next_value = _call_value_nets(
_ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ 

value_net = ValueOperator(
    module=Linear(in_features=2, out_features=1, bias=True),
    device=cpu,
    in_keys=['observation'],
    out_keys=['state_value'])
data = TensorDict(
    fields={
        action: Tensor(shape=torch.Size([4, 3]), device=cpu, dtype=torch.float32, is_shared=F..., device=cpu, dtype=torch.bool, is_shared=False)},
    batch_size=torch.Size([4]),
    device=cpu,
    is_shared=False)
params = TensorDict(
    fields={
        module: TensorDict(
            fields={
                bias: Tensor(shape=torch.Siz...       device=None,
            is_shared=False)},
    batch_size=torch.Size([]),
    device=None,
    is_shared=False)
next_params = TensorDict(
    fields={
        module: TensorDict(
            fields={
                bias: Tensor(shape=torch.Siz...       device=None,
            is_shared=False)},
    batch_size=torch.Size([]),
    device=None,
    is_shared=False)
single_call = True, value_key = 'state_value', detach_next = True
vmap_randomness = 'error'

    def _call_value_nets(
        value_net: TensorDictModuleBase,
        data: TensorDictBase,
        params: TensorDictBase,
        next_params: TensorDictBase,
        single_call: bool,
        value_key: NestedKey,
        detach_next: bool,
        vmap_randomness: str = "error",
    ):
        in_keys = value_net.in_keys
        if single_call:
            for i, name in enumerate(data.names):
                if name == "time":
                    ndim = i + 1
                    break
            else:
                ndim = None
            if ndim is not None:
                # get data at t and last of t+1
                idx0 = (slice(None),) * (ndim - 1) + (slice(-1, None),)
                idx = (slice(None),) * (ndim - 1) + (slice(None, -1),)
                idx_ = (slice(None),) * (ndim - 1) + (slice(1, None),)
                data_in = torch.cat(
                    [
                        data.select(*in_keys, value_key, strict=False),
                        data.get("next").select(*in_keys, value_key, strict=False)[idx0],
                    ],
                    ndim - 1,
                )
            else:
                if RL_WARNINGS:
                    warnings.warn(
                        "Got a tensordict without a time-marked dimension, assuming time is along the last dimension. "
                        "This warning can be turned off by setting the environment variable RL_WARNINGS to False."
                    )
                ndim = data.ndim
                idx = (slice(None),) * (ndim - 1) + (slice(None, data.shape[ndim - 1]),)
                idx_ = (slice(None),) * (ndim - 1) + (slice(data.shape[ndim - 1], None),)
                data_in = torch.cat(
                    [
                        data.select(*in_keys, value_key, strict=False),
                        data.get("next").select(*in_keys, value_key, strict=False),
                    ],
                    ndim - 1,
                )
    
            # next_params should be None or be identical to params
            if next_params is not None and next_params is not params:
>               raise ValueError(
                    "the value at t and t+1 cannot be retrieved in a single call without recurring to vmap when both params and next params are passed."
                )
E               ValueError: the value at t and t+1 cannot be retrieved in a single call without recurring to vmap when both params and next params are passed.

../venv-nightly/lib/python3.11/site-packages/torchrl/objectives/value/advantages.py:122: ValueError

Process finished with exit code 1

Expected behavior

The losses are calculated correctly and the value_network is only called once in the computation of the advantage.

System info

import torchrl, numpy, sys
print(torchrl.__version__, numpy.__version__, sys.version, sys.platform)
2024.6.23 2.0.0 3.11.5 (main, Sep 11 2023, 13:54:46) [GCC 11.2.0] linux

Reason and Possible fixes

The problem seems to be in this snippet, where detached parameter are used for params which makes them unequal.

self.value_estimator(
                tensordict,
                params=self._cached_detach_critic_network_params,
                target_params=self.target_critic_network_params,
            )
´´´´
## Checklist

- [x] I have checked that there is no similar issue in the repo (**required**)
- [x] I have read the [documentation](https://github.com/pytorch/rl/tree/main/docs/) (**required**)
- [x] I have provided a minimal working example to reproduce the bug (**required**)

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.