- 机器学习
- 深度学习
【免费下载链接】dopamine
Dopamine is a research framework for fast prototyping of reinforcement learning algorithms.
ClassicControlRainbowNetwork 是 Dopamine 强化学习框架中为 CartPole、Acrobot、LunarLander、MountainCar 等经典控制(classic control)环境量身打造的 JAX Rainbow 网络,它以全连接多层感知机(MLP)替代 Atari 场景下的卷积网络,并内置特征归一化逻辑,直接适配 Gym 低维连续观测。本文将以官方 API 文档为核心,结合仓库源码、Gin 配置文件与 Colab 示例,完整讲解该网络的字段含义、前向计算流程、归一化原理以及如何在 Rainbow/C51 智能体中将其配置落地。
一、网络定位:从 Atari 像素到经典控制向量
Dopamine 的 JAX 实现(dopamine/jax/networks.py)中同时存在两类网络:面向 Atari 2600 图像输入的卷积网络(如RainbowNetwork、ImplicitQuantileNetwork),以及面向 Gym 经典控制任务的稠密网络。ClassicControlRainbowNetwork属于后者,官方文档对其的定位描述为 “Jax Rainbow network for classic control environments”,即经典控制环境专用的 JAX Rainbow 网络。
经典控制环境的观测不是 84×84 的像素帧,而是低维连续向量,例如:
| 环境 | 观测维度(形状) | 栈大小(stack_size) |
|---|---|---|
| CartPole | (4, 1) | 1 |
| Acrobot | (6, 1) | 1 |
| LunarLander | (8, 1) | 1 |
| MountainCar | (2, 1) | 1 |
上述形状与栈大小常量定义于 dopamine/discrete_domains/gym_lib.py(如gym_lib.CARTPOLE_OBSERVATION_SHAPE、gym_lib.CARTPOLE_STACK_SIZE等)。正因为观测是 4~8 维的小向量,网络完全不需要卷积层,一个数层 MLP 即可高效拟合价值分布。
二、字段(Attributes)完整解析
官方 API 文档将该类标记为 dataclass 风格模块,公开字段如下,全部为 Flax Linennn.Module的 dataclass 字段,在实例化网络时以关键字参数传入:
| 字段 | 含义 | 默认值 |
|---|---|---|
num_actions | 智能体可选动作数,由环境动作空间决定,必填 | 无 |
num_atoms | 分布强化学习中的原子(atom)数量,即 C51/Rainbow 的回报分布支撑点个数 | 无 |
num_layers | 隐藏层数量 | 2 |
hidden_units | 每个隐藏层的神经元数量 | 512 |
min_vals | 观测归一化使用的最小值元组,对应环境各观测维度的下界 | None |
max_vals | 观测归一化使用的最大值元组,对应环境各观测维度的上界 | None |
inputs_preprocessed | 输入是否已预处理(为 True 时跳过网络内部的归一化与展平) | False |
parent | Flax Module 的父模块(框架通用字段) | 无 |
name | Flax Module 名称(框架通用字段) | 无 |
其中parent与name是 Flaxnn.Module的通用元数据字段,num_actions、num_atoms、num_layers、hidden_units、min_vals、max_vals、inputs_preprocessed才是该网络的核心配置。
三、源码级实现:setup 与前向计算
ClassicControlRainbowNetwork定义于 dopamine/jax/networks.py#L337-L377,其核心实现分为setup与__call__两部分。
3.1 setup:构建权重层
def setup(self): if self.min_vals is not None: self._min_vals = jnp.array(self.min_vals) self._max_vals = jnp.array(self.max_vals) initializer = nn.initializers.xavier_uniform() self.layers = [ nn.Dense(features=self.hidden_units, kernel_init=initializer) for _ in range(self.num_layers) ] self.final_layer = nn.Dense( features=self.num_actions * self.num_atoms, kernel_init=initializer )要点:
- 使用
xavier_uniform初始化器初始化所有 Dense 层,这一初始化策略适合 ReLU 之前的线性变换,有助于训练稳定; - 网络由
num_layers个隐藏 Dense 层(每层hidden_units个神经元)加一个最终 Dense 层构成; - 最终层输出维度为
num_actions * num_atoms,即每个动作对应一条长度为num_atoms的回报分布支撑; - 若传入了
min_vals,则将其缓存为 JAX 数组供前向归一化使用。
3.2call:归一化、前向传播与分布输出
def __call__(self, x, support): if not self.inputs_preprocessed: x = x.astype(jnp.float32) x = x.reshape((-1)) # flatten if self.min_vals is not None: x -= self._min_vals x /= self._max_vals - self._min_vals x = 2.0 * x - 1.0 # Rescale in range [-1, 1]. for layer in self.layers: x = layer(x) x = nn.relu(x) x = self.final_layer(x) logits = x.reshape((self.num_actions, self.num_atoms)) probabilities = nn.softmax(logits) q_values = jnp.sum(support * probabilities, axis=1) return atari_lib.RainbowNetworkType(q_values, logits, probabilities)前向流程可拆解为四步:
- 类型与形状归一:输入先转为
jnp.float32并展平为一维向量(reshape((-1)))。注意inputs_preprocessed=True时跳过这一步,适用于输入已经完成预处理与展平的场景。 - 特征缩放(核心亮点):当配置了
min_vals/max_vals时,对每个观测维度执行 min-max 归一化后再映射到 [-1, 1] 区间:x = 2.0 * (x - min_vals) / (max_vals - min_vals) - 1.0。这一设计对经典控制环境至关重要——不同环境观测的数值量纲差异极大(例如 CartPole 的角度范围约 ±0.26 rad,而 MountainCar 的位置范围是 [-1.2, 0.6]),统一到 [-1, 1] 可显著改善 MLP 的训练稳定性。 - MLP 特征提取:依次经过
num_layers个 Dense + ReLU 隐藏层,再通过最终 Dense 层输出num_actions * num_atoms的原始 logits。 - 分布输出:logits 重塑为
(num_actions, num_atoms),经 softmax 得到每个动作上的回报概率分布probabilities,再与支撑向量support做加权求和得到期望 Q 值q_values,最终以RainbowNetworkType(q_values, logits, probabilities)三元组返回。
3.3 返回类型:RainbowNetworkType
返回的命名元组RainbowNetworkType定义于 dopamine/discrete_domains/atari_lib.py#L53-L55:
RainbowNetworkType = collections.namedtuple( 'c51_network', ['q_values', 'logits', 'probabilities'] )q_values:每个动作的期望 Q 值,用于动作选择(argmax);logits:每个动作的原子 logits,用于 C51 交叉熵损失计算;probabilities:每个动作的回报分布概率,用于构建目标分布。
该命名元组同时被 Atari 版RainbowNetwork复用,保证了网络接口的一致性,智能体无需关心底层是 CNN 还是 MLP。
四、与 Rainbow 智能体的协作机制
ClassicControlRainbowNetwork由JaxRainbowAgent(dopamine/jax/agents/rainbow/rainbow_agent.py)驱动。在训练时,智能体调用network_def.apply(params, state, support)获取 logits 并计算 C51 交叉熵损失;在目标分布构建时,target_distribution使用next_state_target_outputs.q_values选最优动作、用probabilities做支撑投影(project_distribution);动作选择阶段则直接取network_def.apply(params, state, support).q_values的 argmax(见 rainbow_agent.py#L200-L204)。support由智能体根据vmin/vmax与num_atoms生成,默认参数为num_atoms=51、vmax=10.0。
五、Gin 配置实战:四环境完整示例
在 Dopamine 中,网络通过 Gin 配置文件绑定到智能体。以 CartPole 的 C51 配置 dopamine/jax/agents/rainbow/configs/c51_cartpole.gin 为例:
import dopamine.jax.agents.rainbow.rainbow_agent import dopamine.jax.networks import dopamine.discrete_domains.gym_lib import dopamine.discrete_domains.run_experiment JaxRainbowAgent.observation_shape = %gym_lib.CARTPOLE_OBSERVATION_SHAPE JaxRainbowAgent.observation_dtype = %jax_networks.CARTPOLE_OBSERVATION_DTYPE JaxRainbowAgent.stack_size = %gym_lib.CARTPOLE_STACK_SIZE JaxRainbowAgent.network = @networks.ClassicControlRainbowNetwork JaxRainbowAgent.num_atoms = 201 JaxRainbowAgent.vmax = 100. JaxRainbowAgent.gamma = 0.99 JaxRainbowAgent.epsilon_eval = 0. JaxRainbowAgent.epsilon_train = 0.01 JaxRainbowAgent.update_horizon = 1 JaxRainbowAgent.min_replay_history = 500 JaxRainbowAgent.update_period = 1 JaxRainbowAgent.target_update_period = 1 JaxRainbowAgent.epsilon_fn = @dqn_agent.identity_epsilon JaxRainbowAgent.replay_scheme = 'uniform' create_optimizer.learning_rate = 0.00001 create_optimizer.eps = 0.00000390625 ClassicControlRainbowNetwork.min_vals = %jax_networks.CARTPOLE_MIN_VALS ClassicControlRainbowNetwork.max_vals = %jax_networks.CARTPOLE_MAX_VALS create_gym_environment.environment_name = 'CartPole' create_gym_environment.version = 'v0' create_runner.schedule = 'continuous_train' create_agent.agent_name = 'jax_rainbow' create_agent.debug_mode = True TrainRunner.create_environment_fn = @gym_lib.create_gym_environment Runner.num_iterations = 400 Runner.training_steps = 1_000 Runner.evaluation_steps = 1_000 Runner.max_steps_per_episode = 200 # Default max episode length. ReplayBuffer.max_capacity = 50_000 ReplayBuffer.batch_size = 128 PrioritizedSamplingDistribution.max_capacity = 50_0005.1 归一化边界常量:min_vals / max_vals 的来源
ClassicControlRainbowNetwork.min_vals与max_vals直接引用 dopamine/jax/networks.py#L35-L51 中注册的 Gin 常量,各环境取值如下:
| 环境 | MIN_VALS | MAX_VALS |
|---|---|---|
| CartPole | (-2.4, -5.0, -π/12, -2π) | (2.4, 5.0, π/12, 2π) |
| Acrobot | (-1, -1, -1, -1, -5, -5) | (1, 1, 1, 1, 5, 5) |
| MountainCar | (-1.2, -0.07) | (0.6, 0.07) |
| LunarLander | —(不配置 min_vals/max_vals) | — |
这些常量对应各环境的物理边界(如 CartPole 的位置 ±2.4、速度 ±5、角度 ±π/12 rad、角速度 ±2π rad/s),因此归一化无需从数据中统计,直接使用环境定义即可。值得注意:CARTPOLE_OBSERVATION_DTYPE、ACROBOT_OBSERVATION_DTYPE等常量将观测 dtype 注册为jnp.float64,以保证数值精度;而 Acrobot/MountainCar/CartPole 均有配套的 rainbow 与 c51 两套配置(如 rainbow_acrobot.gin、c51_mountaincar.gin、rainbow_lunarlander.gin),区别主要在于num_atoms(C51 配置用 51,Rainbow 配置用 201)与vmax取值。
5.2 运行方式
配置完成后,可通过dopamine.discrete_domains.train启动训练:
python -m dopamine.discrete_domains.train \ --base_dir=/tmp/dopamine/cartpole \ --gin_files=dopamine/jax/agents/rainbow/configs/c51_cartpole.gin另外,官方 Colab 示例 dopamine/colab/cartpole.ipynb 中也以相同的 Gin 绑定方式(JaxRainbowAgent.network = @networks.ClassicControlRainbowNetwork、ClassicControlRainbowNetwork.min_vals/max_vals)演示了 CartPole 上的 Rainbow 训练,可作为交互式入门参考。
六、测试与验证
仓库中 tests/dopamine/jax/networks_test.py 对 JAX 网络族(含经典控制网络)进行单元验证,tests/dopamine/jax/agents/rainbow/下的智能体测试则覆盖了JaxRainbowAgent与网络绑定的端到端行为。如需修改或扩展该网络,建议同步补充对应测试以验证输出形状、归一化效果与分布计算逻辑。
七、小结
ClassicControlRainbowNetwork是 Dopamine JAX 体系中将分布强化学习(C51/Rainbow)从 Atari 视觉任务迁移到经典控制任务的关键桥梁:它以可配置深度的 MLP 处理低维连续观测,用环境物理边界驱动的 min-max 归一化保证训练稳定,并通过RainbowNetworkType(q_values, logits, probabilities)这一统一接口无缝接入 Rainbow 智能体的损失计算与动作选择流程。理解它的字段语义与前向流程,即可在 CartPole、Acrobot、MountainCar 等环境中快速复现和定制分布强化学习算法。
- 机器学习
- 深度学习
【免费下载链接】dopamine
Dopamine is a research framework for fast prototyping of reinforcement learning algorithms.
相关推荐
Dopamine legacy_networks 模块解析:TensorFlow 离散域网络架构与实战配置指南
Dopamine legacy_networks 模块解析:TensorFlow 离散域网络架构与实战配置指南 导读 dopamine.discrete_dom
机器学习深度学习maths-cs-ai-compendium 计算机视觉:卷积神经网络全解——卷积机制、经典架构演进与 JAX 实战
maths cs ai compendium 计算机视觉:卷积神经网络全解——卷积机制、经典架构演进与 JAX 实战 卷积神经网络(CNN)不依赖人工设计的滤波
文档教程知识库使用Matcha-gtk-theme打造专业开发环境:程序员桌面美化的10个技巧
使用Matcha gtk theme打造专业开发环境:程序员桌面美化的10个技巧 想要为你的Linux桌面打造一个既美观又高效的开发环境吗?Matcha gtk
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考