Skip to main content

Reverb 是一种高效且易于使用的数据存储和传输系统,专为机器学习研究而设计。

项目描述

混响

PyPI - Python 版本 PyPI 版本

Reverb 是一种高效且易于使用的数据存储和传输系统,专为机器学习研究而设计。Reverb 主要用作分布式强化学习算法的体验回放系统,但该系统还支持多种数据结构表示,例如 FIFO、LIFO 和优先级队列。

目录

安装

请记住,Reverb 并未针对生产使用进行强化,虽然我们尽最大努力保持工作正常,但可能会出现故障或段错误。

:warning: Reverb 目前只支持基于 Linux 的操作系统。

安装 Reverb 的推荐方法是使用pip. 我们还提供了使用与发布相同的 docker 镜像从源代码构建的说明。

TensorFlow 可以单独安装,也可以作为安装的一部分pip。在安装过程中安装 TensorFlow 可确保兼容性。

$ pip install dm-reverb[tensorflow]

# Without Tensorflow install and version dependency check.
$ pip install dm-reverb

每晚构建

PyPI 版本

$ pip install dm-reverb-nightly[tensorflow]

# Without Tensorflow install and version dependency check.
$ pip install dm-reverb-nightly

调试构建

从 0.6.0 版开始,Reverb 的调试版本会上传到 Google Cloud Storage。pip可以按照以下模式直接下载或安装构建。gsutils可用于导航目录结构以确保文件存在,例如 gsutil ls gs://rl-infra-builds/dm_reverb/builds/dbg. 要构建您自己的调试二进制文件,请参阅 构建说明

对于 Python 3.7:

$ export reverb_version=0.8.0
# Python 3.7
$ export python_version=37
$ pip install https://storage.googleapis.com/rl-infra-builds/dm_reverb/builds/dbg/$reverb_version/dm_reverb-$reverb_version-cp$python_version-cp${python_version}m-manylinux2010_x86_64.whl

对于 python 3.8 和 3.9,请遵循以下模式:

$ export reverb_version=0.8.0
# Python 3.9
$ export python_version=39
$ pip install https://storage.googleapis.com/rl-infra-builds/dm_reverb/builds/dbg/$reverb_version/dm_reverb-$reverb_version-cp$python_version-cp$python_version-manylinux2010_x86_64.whl

从源代码构建

本指南 详细介绍了如何从源代码构建混响。

快速开始

启动混响服务器很简单:

import reverb

server = reverb.Server(tables=[
    reverb.Table(
        name='my_table',
        sampler=reverb.selectors.Uniform(),
        remover=reverb.selectors.Fifo(),
        max_size=100,
        rate_limiter=reverb.rate_limiters.MinSize(1)),
    ],
)

创建一个客户端与服务器通信:

client = reverb.Client(f'localhost:{server.port}')
print(client.server_info())

向表中写入一些数据:

# Creates a single item and data element [0, 1].
client.insert([0, 1], priorities={'my_table': 1.0})

一个项目还可以引用多个数据元素:

# Appends three data elements and inserts a single item which references all
# of them as {'a': [2, 3, 4], 'b': [12, 13, 14]}.
with client.trajectory_writer(num_keep_alive_refs=3) as writer:
  writer.append({'a': 2, 'b': 12})
  writer.append({'a': 3, 'b': 13})
  writer.append({'a': 4, 'b': 14})

  # Create an item referencing all the data.
  writer.create_item(
      table='my_table',
      priority=1.0,
      trajectory={
          'a': writer.history['a'][:],
          'b': writer.history['b'][:],
      })

  # Block until the item has been inserted and confirmed by the server.
  writer.flush()

我们添加到 Reverb 的项目可以通过采样来读取:

# client.sample() returns a generator.
print(list(client.sample('my_table', num_samples=2)))

继续 学习混响教程 以获得交互式教程。

详细概述

经验回放已经成为训练离线强化学习策略的重要工具。它被诸如 Deep Q-Networks (DQN)Soft Actor-Critic (SAC)Deep Deterministic Policy Gradients (DDPG)Hindsight Experience Replay等算法所使用……然而,构建一个高效、易于使用和可扩展的回放系统可能具有挑战性。为了获得良好的性能,Reverb 是用 C++ 实现的,并且为了实现分布式使用,它提供了一个 gRPC 服务,用于添加、采样和更新表的内容。Python 客户端以易于使用的方式公开服务的全部功能。此外,原生 TensorFlow 操作可用于与 TensorFlow 和tf.data.

尽管最初是为离策略强化学习而设计的,但 Reverb 的灵活性使其对策略强化甚至(非)监督学习同样有用。有创意的用户甚至使用 Reverb 来存储和分发经常更新的数据(例如模型权重),作为分布式文件系统的内存中轻量级替代品,其中每个表代表一个文件。

混响Server由一个或多个表组成。一个表包含项目,每个项目引用一个或多个数据元素。表格还定义了样本和移除选择策略、最大项目容量和速率限制器

多个项目可以引用相同的数据元素,即使这些项目存在于不同的表中。这是因为项目只包含对数据元素的引用(而不是数据本身的副本)。这也意味着只有在不存在包含对它的引用的项目时才删除数据元素。

例如,可以将一个表设置为转换(长度为 2 的序列)的优先体验重放 (PER),并将另一个表设置为长度为 3 的序列的 (FIFO) 队列。在这种情况下,PER 数据可以用于训练 DQN,FIFO 数据用于训练环境的转换模型。

使用多个表

当满足以下两个条件之一时,项目会自动从表中删除:

  1. 插入新项目会导致表中的项目数超过其最大容量。表的删除策略用于确定要删除的项目。

  2. 一个项目的采样次数超过了表的速率限制器允许的最大次数。这样的项目被删除。

任何项目不再引用的数据元素也将被删除。

用户可以完全控制如何从混响表中采样和删除数据。该行为主要由 提供给 as和的项目选择策略控制。结合 和,可以实现范围广泛的行为。一些常用的配置包括:Tablesamplerremoverrate_limitermax_times_sampled

统一体验回放

维护一组N=1000最近插入的项目。通过设置 sampler=reverb.selectors.Uniform(),选择一个项目的概率对于所有项目都是相同的。由于reverb.rate_limiters.MinSize(100),采样请求将被阻止,直到插入 100 个项目。通过设置 remover=reverb.selectors.Fifo()何时需要删除项目,首先删除最旧的项目。

reverb.Table(
     name='my_uniform_experience_replay_buffer',
     sampler=reverb.selectors.Uniform(),
     remover=reverb.selectors.Fifo(),
     max_size=1000,
     rate_limiter=reverb.rate_limiters.MinSize(100),
)

利用统一经验回放的算法示例包括SACDDPG

优先体验重播

一组N=1000最近插入的项目。通过设置 sampler=reverb.selectors.Prioritized(priority_exponent=0.8),选择项目的概率与项目的优先级成正比。

注:参见Schaul、Tom 等人。用于优先体验重放的此实现中使用的算法。

reverb.Table(
     name='my_prioritized_experience_replay_buffer',
     sampler=reverb.selectors.Prioritized(0.8),
     remover=reverb.selectors.Fifo(),
     max_size=1000,
     rate_limiter=reverb.rate_limiters.MinSize(100),
)

利用优先体验重放的算法示例是 DQN(及其变体)和 分布式分布确定性策略梯度

队列

最多可收集N=1000在同一操作中选择并删除最旧项目的项目。如果集合包含 1000 个项目,则插入调用被阻塞,直到它不再满,如果集合为空,则样本调用被阻塞,直到至少有一个项目。

reverb.Table(
    name='my_queue',
    sampler=reverb.selectors.Fifo(),
    remover=reverb.selectors.Fifo(),
    max_size=1000,
    max_times_sampled=1,
    rate_limiter=reverb.rate_limiters.Queue(size=1000),
)

# Or use the helper classmethod `.queue`.
reverb.Table.queue(name='my_queue', max_size=1000)

使用队列的算法示例是 IMPALA和近端策略优化的异步实现 。

项目选择策略

Reverb 定义了几个可用于项目采样或移除的选择器:

  • 均匀:在所有项目中均匀采样。
  • 优先级:与存储的优先级成比例的样本。
  • FIFO:选择最旧的数据。
  • LIFO:选择最新数据。
  • MinHeap:选择优先级最低的数据。
  • MaxHeap:选择优先级最高的数据。

这些策略中的任何一个都可用于从表中采样或删除项目。这使用户可以灵活地创建最适合其需求的自定义表格。

速率限制

速率限制器允许用户对何时可以插入和/或从表中采样项目实施条件。以下是 Reverb 中当前可用的速率限制器列表:

  • MinSize:设置在可以对任何内容进行采样之前必须在表中的最小项目数。
  • SampleToInsertRatio:通过阻止插入和/或样本请求来设置插入与样本的平均比率。这对于控制每​​个项目在被删除之前的采样次数很有用。
  • 队列:项目在被删除之前只被采样一次。
  • 堆栈:项目在被移除之前只被采样一次。

分片

混响服务器彼此不知道,当将系统扩展到多服务器时,设置数据不会跨多个节点复制。这使得 Reverb 不适合作为传统数据库,但它的好处是在可以接受一定程度的数据丢失的情况下扩展系统变得微不足道。

分布式系统可以通过简单地增加 Reverb 服务器的数量来进行水平扩展。当与 gRPC 兼容的负载均衡器结合使用时,负载均衡目标的地址可以简单地提供给 Reverb 客户端,操作将自动分布在不同的节点上。您将在相关方法和类的文档中找到有关特定行为的详细信息。

如果您的设置中没有负载平衡器,或者需要更多控制,那么系统仍然可以以几乎相同的方式进行扩展。只需增加 Reverb 服务器的数量并为每个服务器创建单独的客户端。

检查点

Reverb 支持检查点;Reverb 服务器的状态和内容可以存储到永久存储中。指向时,Server序列化重建它所需的所有数据和元数据。在此过程中,将Server 阻止所有传入的插入、采样、更新和删除请求。

检查点是通过来自 Reverb 的调用完成的Client

# client.checkpoint() returns the path the checkpoint was written to.
checkpoint_path = client.checkpoint()

要从检查点恢复reverb.Server

checkpointer = reverb.checkpointers.DefaultCheckpointer(path=checkpoint_path)
# The arguments passed to `tables=` must be the same as those used by the
# `Server` that wrote the checkpoint.
server = reverb.Server(tables=[...], checkpointer=checkpointer)

有关在 Reverb 中执行检查点的详细信息,请参阅 tfrecord_checkpointer.h 。

reverb_server使用(测试版)启动混响

安装dm-reverbusingpip将安装一个reverb_server脚本,该脚本接受其配置作为 textproto。例如:

$ reverb_server --config="
port: 8000
tables: {
  table_name: \"my_table\"
  sampler: {
    fifo: true
  }
  remover: {
    fifo: true
  }
  max_size: 200 max_times_sampled: 5
  rate_limiter: {
    min_size_to_sample: 1
    samples_per_insert: 1
    min_diff: $(python3 -c "import sys; print(-sys.float_info.max)")
    max_diff: $(python3 -c "import sys; print(sys.float_info.max)")
  }
}"

rate_limiter配置等效于 Python 表达式MinSize(1),请参阅rate_limiters.py

引文

如果您使用此代码,请将 Reverb 论文引用为

@misc{cassirer2021reverb,
      title={Reverb: A Framework For Experience Replay},
      author={Albin Cassirer and Gabriel Barth-Maron and Eugene Brevdo and Sabela Ramos and Toby Boyd and Thibault Sottiaux and Manuel Kroiss},
      year={2021},
      eprint={2102.04736},
      archivePrefix={arXiv},
      primaryClass={cs.LG}
}