【问题标题】:tf_agents and reverb produce incompatible tensortf_agents 和 reverb 产生不兼容的张量
【发布时间】:2022-10-15 20:26:05
【问题描述】:

我正在尝试使用tf_agents 和reverb 实现DDPG,但我无法弄清楚这两个库如何协同工作。为此,我尝试将来自 tf_agents 的DQL-Tutorial 中的代码与我自己的代理和健身房环境一起使用。当我尝试从混响中检索数据并且张量形状不匹配时,会发生错误。我创建了我能想到的最小的示例来显示问题:

进口

import gym
from gym import spaces
from gym.utils.env_checker import check_env
from gym.envs.registration import register

import tensorflow as tf
import numpy as np
import reverb

from tf_agents.agents import DdpgAgent
from tf_agents.drivers.py_driver import PyDriver
from tf_agents.environments import TFPyEnvironment, suite_gym, validate_py_environment
from tf_agents.networks import Sequential
from tf_agents.policies import PyTFEagerPolicy
from tf_agents.replay_buffers import ReverbReplayBuffer, ReverbAddTrajectoryObserver
from tf_agents.specs import tensor_spec, BoundedArraySpec

健身房环境示例

class TestGym(gym.Env):
    metadata = {"render_modes": ["human"]}
    def __init__(self):
        self.observation_space = spaces.Box(low=-1, high=1, shape=(30,), dtype=np.float32)
        self.action_space = spaces.Box(low=-1, high=1, shape=(2,), dtype=np.float32)
        self.__count = 0
    def step(self, action):
        self.__count += 1
        return np.zeros(30, dtype=np.float32), 0, self.__count >= 100, {}
    def render(self, mode="human"):
        return None
    def reset(self, seed=None, return_info=False, options=None):
        super().reset(seed=seed, options=options)
        self.__count = 0
        if return_info:
            return np.zeros(30, dtype=np.float32), {}
        else:
            return np.zeros(30, dtype=np.float32)

register(
    id="TestGym-v0",
    entry_point="reverb_test:TestGym",
    nondeterministic=False
)

创建一个 TFAgent 并使用混响来存储和检索

def main():
    # make sure the gym environment is ok
    check_env(gym.make("TestGym-v0"))

    # create tf-py-environment
    env = TFPyEnvironment(suite_gym.load("TestGym-v0"))

    # make sure the py environment is ok
    validate_py_environment(env.pyenv, episodes=5)

    # example actor network
    actor_network = Sequential([
        tf.keras.layers.Dense(40),
        tf.keras.layers.Dense(2, activation=None)
    ], input_spec=env.observation_spec())

    # example critic network
    n_actions = env.action_spec().shape[0]
    n_observ = env.observation_spec().shape[0]
    critic_input_spec: BoundedArraySpec = BoundedArraySpec((n_actions + n_observ,), "float32", minimum=-1, maximum=1)
    critic_network = Sequential([
        tf.keras.layers.Dense(40),
        tf.keras.layers.Dense(1, activation=None)
    ], input_spec=critic_input_spec)

    # example rl agent
    agent = DdpgAgent(
        time_step_spec=env.time_step_spec(),
        action_spec=env.action_spec(),
        actor_network=actor_network,
        critic_network=critic_network,
    )

    # create reverb table
    table_name = "uniform_table"
    replay_buffer_signature = tensor_spec.from_spec(agent.collect_data_spec)
    replay_buffer_signature = tensor_spec.add_outer_dim(replay_buffer_signature)
    table = reverb.Table(
        table_name,
        max_size=100_000,
        sampler=reverb.selectors.Uniform(),
        remover=reverb.selectors.Fifo(),
        rate_limiter=reverb.rate_limiters.MinSize(1),
        signature=replay_buffer_signature
    )

    # create reverb server
    reverb_server = reverb.Server([table])

    # create replay buffer for this table and server
    replay_buffer = ReverbReplayBuffer(
        agent.collect_data_spec,
        table_name=table_name,
        sequence_length=2,
        local_server=reverb_server
    )

    # create observer to store experiences
    observer = ReverbAddTrajectoryObserver(
        replay_buffer.py_client,
        table_name,
        sequence_length=2
    )

    # run a view steps to ill the replay buffer
    driver = PyDriver(env.pyenv, PyTFEagerPolicy(agent.collect_policy, use_tf_function=True), [observer], max_steps=100)
    driver.run(env.reset())

    # create a dataset to access the replay buffer
    dataset = replay_buffer.as_dataset(num_parallel_calls=3, sample_batch_size=20, num_steps=2).prefetch(3)
    iterator = iter(dataset)

    # retrieve a sample
    print(next(iterator)) # <===== ERROR


if __name__ == '__main__':
    main()

当我运行此代码时,我收到以下错误消息:

tensorflow.python.framework.errors_impl.InvalidArgumentError:
{{function_node __wrapped__IteratorGetNext_output_types_11_device_/job:localhost/replica:0/task:0/device:CPU:0}}
Received incompatible tensor at flattened index 0 from table 'uniform_table'.
    Specification has (dtype, shape): (int32, [?]).
    Tensor has        (dtype, shape): (int32, [2,1]).
Table signature:
    0: Tensor<name: 'step_type/step_type', dtype: int32, shape: [?]>,
    1: Tensor<name: 'observation/observation', dtype: float, shape: [?,30]>,
    2: Tensor<name: 'action/action', dtype: float, shape: [?,2]>,
    3: Tensor<name: 'next_step_type/step_type', dtype: int32, shape: [?]>,
    4: Tensor<name: 'reward/reward', dtype: float, shape: [?]>,
    5: Tensor<name: 'discount/discount', dtype: float, shape: [?]>
[Op:IteratorGetNext]

在我的健身房环境中,我将动作空间定义为一个 2 元素向量,我猜这个动作向量在某种程度上是个问题。我尝试对每个输入和输出使用张量规范,但我想我在某个地方犯了错误。有没有人知道我在这里做错了什么?

【问题讨论】:

  • 这个[?] 可能建议一维数据,但您有[2,1] 建议二维数据。有时它只需要flatten() 数据。
  • 这里的想法是成对检索数据点。出于这个原因,重播缓冲区、观察者和数据集的序列长度为 2。我假设因此,张量在索引 0 处有 2 个元素。因为我正在使用所有这些框架(TFPyEnvironment、DdpgAgent、混响,PyDriver等...),我无法真正手动将其展平,并且我正在努力寻找可以设置来修复它的参数。

标签: python tensorflow reinforcement-learning openai-gym tf-agent


【解决方案1】:

我终于想通了:

PyDriver 需要 PyEnvironment 才能正常工作。在我的代码中,我使用了TFPyEnvironment 的pyenv 属性,尽管它的名称如此,但它不会返回常规的PyEnvironment,而是返回一个批处理的。

通过以下方式更改代码可解决此问题:

...

def main():
    # make sure the gym environment is ok
    check_env(gym.make("TestGym-v0"))

    # create py-environment
    pyenv = suite_gym.load("TestGym-v0")  # <=============

    # create tf-py-environment
    env = TFPyEnvironment(pyenv)

    ...

    driver = PyDriver(py_env, PyTFEagerPolicy(agent.collect_policy, use_tf_function=True), [observer], max_steps=100)
    driver.run(py_env.reset())

    ...

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 2022-12-12
    • 2016-04-12
    • 1970-01-01
    • 2021-08-16
    • 2020-05-03
    • 1970-01-01
    • 1970-01-01
    • 2019-03-18
    相关资源
    最近更新 更多