- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
Google Research
jrl/envs/README.md是 JRL(Jax 离线强化学习研究代码库)中关于环境(Environment)接入规范的核心文档。本文以该文档为主体骨架,结合 jrl/envs/init.py、jrl/envs/d4rl.py、jrl/envs/dm_control.py 以及 jrl/localized/runner.py 等源码,系统讲解 JRL 环境的统一抽象、内置环境实现、注册方式与实战调用流程。读完本文,你将掌握如何在 JRL 中新增自定义环境,并理解 dm_env 规范、Acme wrappers、单精度封装等底层机制。
环境模块定位:JRL 的离线 RL 研究基石
JRL 是基于 Jax 与 Acme RL 库构建的离线强化学习研究代码库,其核心论文为 "Why so pessimistic? Estimating uncertainties for offline RL through ensembles, and why their independence matters."(jrl/README.md)。整个代码库按职责划分为四个模块:
- Agents:训练算法实现(详见 jrl/agents/README.md)
- Datasets:离线数据集加载(详见 jrl/data/README.md)
- Environments:环境接入与统一封装(即本文主题,详见 jrl/envs/README.md)
- evaluation / localized / utils:评估、训练入口 runner 与工具函数
jrl/envs/目录下仅有三个文件:README.md、init.py、d4rl.py、dm_control.py,是一个轻量但职责清晰的环境注册与工厂模块。它解决了离线 RL 研究中一个关键痛点:不同来源的环境(如 Gym/D4RL、DM Control)拥有完全不同的 API,而 JRL 的 Agent 与 Runner 只消费统一的dm_env接口,因此需要一个集中式的环境工厂来完成适配与注册。
核心规范:所有环境必须封装为dm_env
jrl/envs/README.md全文只有两条核心指令,它们是整个环境模块的设计契约:
- 环境必须被封装为
dm_env(DeepMind Environment)接口; - 封装后的环境必须注册到 jrl/envs/init.py 的工厂函数中。
这两条规范直接决定了 JRL 全链路代码的写法。从 jrl/localized/runner.py 可以看到,训练入口通过统一的工厂函数创建环境:
create_env_fn = lambda: envs.create_environment( FLAGS.task_class, FLAGS.task_name, FLAGS.single_precision_env) environment = create_env_fn() spec = specs.make_environment_spec(environment)也就是说,Agent 的训练、评估乃至环境规格推导(make_environment_spec)全部只依赖dm_env.Environment接口。任何不满足该接口的环境,都无法直接进入 JRL 的训练管线。这正是 README 强调封装规范的根本原因。
工厂函数源码剖析
jrl/envs/init.py 实现了两个层次的工厂函数:
def _create_environment(task_class, task_name, **kwargs): if task_class == 'd4rl': from jrl.envs import d4rl return d4rl.create_d4rl_env(task_name, **kwargs) elif task_class == 'dm_control': from jrl.envs import dm_control return dm_control.create_dm_control_env(task_name) else: raise NotImplementedError('task class not handled!') def create_environment(task_class, task_name, single_precision=False, **kwargs): env = _create_environment(task_class, task_name, **kwargs) if single_precision: env = wrappers.SinglePrecisionWrapper(env) return env关键设计点:
- 以
task_class分派,以task_name定位具体任务:task_class决定环境来源(当前支持'd4rl'与'dm_control'),task_name是具体任务标识(如'antmaze-large-diverse-v0')。 single_precision参数:默认False。开启后会用 Acme 的wrappers.SinglePrecisionWrapper将环境内部所有 spec 与 timestep 强制转为 float32,避免 float64 导致的 Jax XLA 编译问题——这是 Jax 训练管线中的常见坑,值得所有 RL 研究者注意。- 延迟导入(lazy import):工厂函数内部才
from jrl.envs import d4rl,避免导入 JRL 时强制加载 gym/d4rl 等重依赖,加快模块加载速度。 - 未支持类目直接抛
NotImplementedError:新增环境来源必须同时修改该工厂函数,否则运行时报错。
内置环境实现一:D4RL(Gym 生态)
jrl/envs/d4rl.py 提供了 D4RL 任务的创建逻辑:
def create_d4rl_env(task_name): env = gym.make(task_name) env = wrappers.GymWrapper(env) return env实现要点:
- 使用
gym.make(task_name)创建 Gym 环境,task_name即 Gym/D4RL 的标准任务名(例如'antmaze-large-diverse-v0'、'halfcheetah-medium-expert-v2'); - 通过 Acme 的
wrappers.GymWrapper将 Gym 环境包装为dm_env.Environment接口,从而满足 README 的第一条规范; - D4RL 是离线 RL 的标准基准数据集,其环境直接支持通过 gym 注册表创建,因此这一路径是 JRL 训练 D4RL 基准任务(如 BC、CQL、MSG 等算法)的默认选择。
从 jrl/agents/bc/README.md 可以看到实际运行 D4RL 任务的命令行用法:
--task_class 'd4rl' \ --task_name 'antmaze-large-diverse-v0' \这与 jrl/data/d4rl.py 中的数据集加载逻辑一一对应(task_class == 'd4rl'时加载 D4RL 数据集),保证环境与数据来源一致。
内置环境实现二:DM Control
jrl/envs/dm_control.py 提供 DeepMind Control Suite 支持,其实现比 D4RL 多了一层自定义封装:
class FlatObservationWrapper(wrappers.EnvironmentWrapper): def _convert_obs(self, obs): flat_obs = [v.flatten() for v in obs.values()] return np.concatenate(flat_obs, axis=-1) def reset(self): ts = self._environment.reset() return ts._replace(observation=self._convert_obs(ts.observation)) def step(self, action): ts = self._environment.step(action) return ts._replace(observation=self._convert_obs(ts.observation)) def observation_spec(self): original_obs_spec = self._environment.observation_spec() types = [] sizes = [] for k, v in original_obs_spec.items(): types.append(v.dtype) sizes.append(np.prod(v.shape)) assert all(x == types[0] for x in types), 'All types not the same!' total_size = sum(sizes) return specs.Array(shape=(total_size,), dtype=types[0], name='flat_obs_spec') def create_dm_control_env(task_name): split_name = task_name.split('__') domain_name, task_name = split_name[0], split_name[1] env = suite.load(domain_name=domain_name, task_name=task_name) env = FlatObservationWrapper(env) return env实现要点:
- 任务名编码约定:DM Control 任务的
task_name使用domain__task格式,例如'cartpole__swingup',工厂函数内部用split('__')拆分为domain_name与task_name,再传给suite.load; FlatObservationWrapper:DM Control 的观测是 dict 结构(如关节位置、速度等),而 JRL 的 Agent 期望扁平化向量观测。该 wrapper 继承 Acme 的EnvironmentWrapper,将 dict 观测展平拼接为单一向量,并同步重写了observation_spec(校验所有子观测 dtype 一致后返回拼接后的specs.Array);- 这是 README 规范的最佳实践示范:在不改动底层环境的前提下,通过 wrapper 层完成接口与格式适配。
如何在 JRL 中新增自定义环境(实战步骤)
结合 README 的规范与源码结构,接入一个新环境(例如自定义 Gym 环境)需要四步:
- 封装为 dm_env:确保环境实现
dm_env.Environment的reset/step/observation_spec/action_spec/reward_spec接口。若你的环境是 Gym 格式,可直接复用 Acme 的wrappers.GymWrapper;若是其他格式,可参考 FlatObservationWrapper 继承wrappers.EnvironmentWrapper自行包装。 - 新建创建函数:参考
create_d4rl_env/create_dm_control_env,在jrl/envs/下新增模块(如my_env.py),提供接收task_name并返回 dm_env 的工厂函数。 - 注册到工厂:修改 jrl/envs/init.py 的
_create_environment,新增task_class分支并调用你的创建函数。 - 联调验证:通过 jrl/localized/runner.py 运行
--task_class 'my_env_class' --task_name '<your_task>'启动训练,检查环境规格推导与 rollout 是否正常;若遇到 float64 精度问题,可加--single_precision_env开启单精度封装。
完整运行示例
以 JRL 自带的 BC 算法 + D4RL 任务为例(完整命令见 jrl/agents/bc/README.md):
python3 -m jrl.localized.runner \ --pdb_post_mortem \ --debug_nans=False \ --create_saved_model_actor=False \ --num_steps 11000 \ --eval_every_steps 500 \ --episodes_per_eval 100 \ --batch_size 51200 \ --root_dir '/tmp/test_bc' \ --seed 42 \ --algorithm 'bc' \ --task_class 'd4rl' \ --task_name 'antmaze-large-diverse-v0' \ --gin_bindings='bc.config.BCConfig.num_sgd_steps_per_step=200' \ --gin_bindings='bc.config.BCConfig.policy_lr=1e-4' \ --gin_bindings='bc.config.BCConfig.loss_type="MLE"' \ --gin_bindings='bc.config.BCConfig.entropy_regularization_weight=0'其中--task_class 'd4rl'与--task_name 'antmaze-large-diverse-v0'正是通过本文介绍的环境工厂被解析并创建;若将--task_class改为'dm_control'、--task_name改为'cartpole__swingup'这类格式,即可切换为 DM Control 环境。--gin_bindings则通过 gin 配置库注入算法超参(各算法参数定义见对应config.py)。
与其他模块的协作关系
环境模块并非孤立存在,它与 JRL 的其余模块构成完整闭环:
- 与数据模块对应:jrl/data/init.py 采用与
envs完全相同的task_class分派模式('d4rl'→d4rl.create_d4rl_data_iter),保证离线数据与环境一一对应; - 与 Runner 衔接:jrl/localized/runner.py 先创建环境、推导
environment_spec,再将其交给agents.create_agent构建智能体,环境规范(观测/动作空间)是 Agent 网络结构定义的输入; - 与算法模块联动:
jrl/agents/下各算法(bc、cql、msg、snr、batch_ensemble_msg)的 README 均以--task_class 'd4rl'作为标准示例,说明该工厂模式是全部算法的通用环境入口。
小结
jrl/envs/README.md虽然简短,却定义了 JRL 环境接入的两条铁律——统一 dm_env 抽象 + 集中式工厂注册。从源码可见,这一设计带来了三方面收益:环境来源可扩展(新增task_class即可)、接口对 Agent 完全透明(一律消费 dm_env)、精度/观测格式等共性问题可在 wrapper 层统一解决。如果你打算在 JRL 中复现或扩展离线 RL 实验,掌握本模块的封装与注册机制是接入新基准的第一步。
- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
Google Research
相关推荐
MarkItDown 快速上手:把办公文档转成 Markdown 的保姆级指南
MarkItDown 快速上手:把办公文档转成 Markdown 的保姆级指南 你有没有被这种场景折腾过:手里一堆 PDF、Word、Excel 和 PPT,想
人工智能深度学习NLP计算机视觉强化学习JRL 中的 CQL 离线强化学习实现:配置参数、BC 预热与 D4RL 训练实战指南
JRL 中的 CQL 离线强化学习实现:配置参数、BC 预热与 D4RL 训练实战指南 本指南以 Google Research 的 JRL(Jax Reinf
人工智能深度学习NLP计算机视觉强化学习signal-back高级技巧:解决备份文件损坏、密码错误与附件提取失败的实用方法
signal back高级技巧:解决备份文件损坏、密码错误与附件提取失败的实用方法 signal back 是一款用于在应用外解密Signal加密备份的工具,采
人工智能深度学习NLP计算机视觉强化学习
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考