JAX与EvoRL安装实战:进化强化学习环境搭建避坑指南
上周给实验室的新机器配了一套进化强化学习环境目标是装好JAX和EvoRL跑群智能体实验。原以为驱动装好、再 pip 两个包就算完事结果在EvoRL 安装这一步踩了整整一下午的坑jaxlib 版本对不上、依赖被 pip 偷偷替换、CUDA 识别不出来……回头看问题大多出在没搞清 JAX 和 EvoRL 之间的依赖层级就急着敲命令。这篇文章按我最后的成功路径重新梳理了一遍从机器预检、JAX 选型、EvoRL 依赖逻辑到安装完成后的验证和排错全部记录下来。准备用 EvoRL 做进化策略实验、或者想在 JAX 生态里跑强化学习算法的读者可以直接照着流程操作能少走不少弯路。1. 为什么要把 JAX 和 EvoRL 放在一起装先弄清这层依赖关系很多人对 JAX 的理解停留在能用 GPU 加速的 NumPy装上就跑对 EvoRL 就更模糊以为它是个独立的强化学习工具箱装完就能调算法。实际上EvoRL 是构建在 JAX 之上的算法套件它的梯度计算、种群并行评估、环境批量交互底层全都依赖 JAX 的自动微分和 XLA 编译能力。JAX 装不对后面啥都跑不动。1.1 JAX 到底是干什么的JAX 是一个函数式数值计算框架核心卖点是四个转化grad自动微分、jit即时编译、vmap自动向量化和pmap多设备并行。跟 NumPy 比它最大的优势是同一份 Python 代码既能逐条求梯度又能自动套上批量维度还能被编译成高效的低级指令。举个例子在传统强化学习里你要并行采样 256 条轨迹通常得写一个for循环依次跑环境或者在 C 侧开线程池。但在 JAX 里用vmap把策略函数映射到批量纬上一次性推 256 份参数和环境状态整个 rollout 的并行度直接交给 XLA 编译器和 GPU 去处理。这就是 EvoRL 选择它当后端的核心原因。1.2 EvoRL 到底是干什么的EvoRL 把进化计算和强化学习放进了同一个框架。你可以用 CMA-ES、OpenES、PGPE、遗传算法这类进化优化器去训练策略也能切换成 PPO、DDPG、TD3、SAC 这些常规 RL 算法做对照实验。进化策略的训练方式跟梯度下降不太一样它不反向传播算梯度而是维护一个参数种群每个个体是候选的策略权重然后把整批个体丢到环境里评估适应度再用适应度排序去更新下一轮分布。这个过程天然适合并行——每个个体跑环境互不干扰所以种群评估部分几乎能无损地拆到多设备上。JAX 的vmap加上jit恰好能把这种大规模并行 rollout 分布更新的计算吞得干干净净。1.3 安装顺序的底层逻辑EvoRL 的依赖列表里有 JAX所以你直接pip install evorl时pip 会顺手装一个 JAX。但问题也出在这如果不显式指定版本pip 默认拉下来的很可能是 CPU 版 JAX。你在有 N 卡的机器上费劲装好了驱动结果 EvoRL 还是在用 CPU 硬算性能差距可能有好几十倍。所以最稳的顺序是先根据机器情况手动装好指定版本的 JAX确认 GPU 能被识别再去装 EvoRL装完再一次确认 JAX 有没有被覆盖。下面就从机器预检开始讲。2. 装之前先看机器这套预检清单能筛掉八成的装不成功安装类任务最忌讳的就是闭眼敲命令报错再查。大多数安装失败其实在敲第一条命令之前就已经注定了——要么是 Python 版本太新要么是驱动根本不支持目标 CUDA 版本。我现在的习惯是先花五分钟把下面这张清单过一遍。2.1 操作系统与 Python 版本JAX 和 EvoRL 的官方支持主线是 Linux在 Windows 上直接装 EvoRL 有概率遇到编译报错所以如果你是 Windows 机器我建议直接开 WSL2在里面操作。macOS 的 Intel 和 Apple Silicon 都能装 CPU 版 JAX但 EvoRL 某些依赖依然更适合在 Linux 下跑。# 看系统和架构 uname -a # 看 Python 版本 python -VPython 版本建议用3.10 或 3.11。3.9 能用但有些新版本的依赖已经放弃对它做完整测试3.12 和更高的版本目前还时不时有编译链不兼容的问题尤其是涉及 C 扩展的那些库没必要赌。2.2 GPU 驱动与 CUDA 检查如果机器上有 N 卡先执行nvidia-sminvidia-smi重点看右上角的CUDA Version和 Driver 版本。新版 JAX0.4.x 之后已经内置了 CUDA 运行时和 cuDNN 库不需要你再手动装整套 CUDA Toolkit这条容易误解我特别说明一下你只需要显卡驱动支持的目标 CUDA 版本 12.0JAX 的cuda12变体就能跑了。只有用很老版本的 JAX 时才需要手动去匹配 cuDNN。如果是 CPU 机器或者只有核显直接跳过 GPU 这条通篇看 CPU 版安装即可。2.3 内存与磁盘余量XLA 在第一次 JIT 编译代码时会生成本地指令缓存占用 CPU 资源比较大如果编译大型网络内存 8G 会有点紧16G 以上比较舒服。磁盘方面JAX、EvoRL 和它的依赖安装完加上缓存预留 5G 比较踏实。2.4 虚拟环境隔离强烈建议单独给这个项目建一个虚拟环境不要直接装到 base 环境里。我用的是一个干净环境conda create -n evorl python3.10 -y conda activate evorl后面所有 pip 命令都在这个环境里执行。这样一旦搞坏删掉环境重来就行不会污染其他项目的依赖。3. JAX 本体安装的完整链路CPU、GPU 与 Apple Silicon 的选型JAX 安装的核心就一句话选择和你机器匹配的 wheel别让 pip 猜。很多人装上之后发现jax.devices()不显示 GPU十有八九是 pip 猜了个 CPU 版给你。3.1 CPU 版安装纯 CPU 场景下最简单python -m pip install --upgrade pip pip install -U jax老教程会让你额外装jaxlib新版 JAX 已经把它作为依赖自动拉取了所以不用手动指定。装完可以先快速验证import jax print(jax.__version__) print(jax.devices())输出里出现CpuDevice就说明基础环境 OK。如果机器没有 GPU这步通过后可以直接跳到 EvoRL 的安装。3.2 GPU 版安装Linux CUDA 12.x在确认nvidia-smi显示的 CUDA 12.0 之后用带 extras 的安装方式pip install -U jax[cuda12]这里引号不能省否则部分 shell 会解析错误。cuda12这个标签会把 JAX、jaxlib 和对应的 CUDA 插件一起装好。import jax print(jax.devices())如果输出里有CudaDevice(id0)之类的条目说明 GPU 版生效。我见过不少人在这步发现只有CpuDevice原因通常是之前装过 CPU 版 jax缓存里优先用了旧包这时候pip install --upgrade jax[cuda12]再强制重装一次就好。3.3 Apple Silicon 与其它场景M 系列芯片的 Mac 上CPU 版直接pip install -U jax就行。如果要利用 GPU 做加速需要另外装一个专门的后端pip install -U jax pip install jax-metal不过 EvoRL 的很多依赖在 macOS 上的兼容性不如 Linux我自己的建议是Mac 上跑简单 demo 没问题做完整实验还是弄台 Linux 机器或者 WSL2 更省心。3.4 两个实用的环境变量JAX 默认是 32 位浮点计算对某些需要高精度的实验来说不够用。可以在环境变量里开启 x64export JAX_ENABLE_X641另外 GPU 版 JAX 默认会在一开始预分配显存可能把整个显存都占住导致其他程序没得用。如果机器上要跑多个任务建议加export XLA_PYTHON_CLIENT_PREALLOCATEfalse装完 JAX 并验证通过后下面再碰 EvoRL 就不容易踩雷了。4. EvoRL 的依赖逻辑与你需要懂的那些库在真正执行安装命令之前先花点时间理解 EvoRL 装了以后会带来哪些依赖以及它们各自在框架里干什么。这能帮你快速判断报错到底出在哪一层。4.1 EvoRL 的运行链路EvoRL 的一个典型实验循环是这样的首先是进化优化器根据当前分布生成一批候选参数向量然后把这批参数分发给并行环境评估器。评估器用 JAX 编译过的函数批量跑环境得到每条轨迹的累计奖励最后进化优化器根据这些适应度更新参数分布。这个过程里神经网络参数定义、概率分布采样、奖励归一化、优化器更新分别由不同库承担。所以 EvoRL 的依赖不是随便堆上去的每个都有明确分工。4.2 核心依赖逐个说chexJAX 生态里的类型检查和调试工具。EvoRL 用它来保证参数数组的形状和 dtype 在多层函数间传递时不跑偏。distrax处理概率分布。进化策略里的高斯扰动、PPO 用到的策略分布都靠它实现。flax神经网络模块。EvoRL 里的 Actor 和 Critic 网络默认用 flax 来定义和存储参数。optax优化器集合。标准 RL 算法用到的 Adam 这类更新规则走的就是 optax。hydra-core配置管理。EvoRL 支持通过 yaml 文件配置算法超参hydra 负责把这些配置跟命令行参数串起来。gymnax / jumanjiJAX 生态的环境库。提供用 JAX 重写过的强化学习环境支持批量并行 rollout。看到这些依赖你就能理解一个常见现象安装 EvoRL 的时候pip 会重新解析一整套 JAX 生态包的版本。如果其中某个版本强制依赖了老版 JAX就可能把你刚装好的 GPU 版覆盖掉。4.3 版本共振问题JAX 迭代速度很快它的 API 每隔几个版本就会调整。比如jax.tree_util在旧版是jax.utilchex 和 distrax 会跟着 JAX 一起换版本。所以安装 EvoRL 时尽量不要手动固定老版本 JAX 去迎合某个旧依赖除非你有特殊原因。我的做法是让 EvoRL 的依赖在合理范围内取最新然后单独锁定 JAX 为 GPU 版两者冲突的概率会小很多。5. EvoRL 安装实操pip、源码与依赖锁定的取舍EvoRL 的安装有两条常规路线一条是直接用 pip 装发布版适合大多数人另一条是从源码安装适合要改算法源码或者需要跑最新特性的场景。5.1 pip 快速安装如果只是想跑实验最直接的方式是pip install evorl这一步会自动安装前面说的 chex、distrax、flax、optax 等依赖。如果你之前已经手动装好了 GPU 版 JAX装完 evorl 之后务必要再执行一次pip install --upgrade jax[cuda12]原因我在前面说过evorl 的依赖解析过程不保证不会动 jax 的版本。多敲这一条能把 JAX 覆盖回 GPU 版成本极低。装完再跑一次jax.devices()确认 GPU 还在。5.2 源码安装需要改 EvoRL 内部实现或者官方发布版还没包含你想要的新算法时走源码安装。基本流程是git clone evorl 官方仓库地址 cd evorl pip install -e . --no-deps这里我用--no-deps是因为不想让 pip 重新解析整套依赖链以免意外改动已安装好的 JAX GPU 版。然后根据项目的 requirements 手动补齐缺失部分pip install hydra-core chex distrax flax optax gymnax jumanji源码安装的优点是pip install -e会建立软链接你直接改源码下次运行就生效非常适合调试算法。5.3 用 pip check 验证依赖一致性安装完成后我建议立刻跑一条压箱底命令pip check这条命令会扫描当前环境里所有已安装包的依赖关系把版本冲突和不满足项列出来。如果输出里出现No broken requirements found说明依赖层面是干净的如果列出了冲突就按它提示的包名和版本去调整。6. 验证安装是否真的成功了从 JAX 到 EvoRL 的四步检查很多教程到安装成功就结束了但实际上import成功和能跑实验之间还有不小距离。我一般按下面四步做验证每一步都看得见输出出问题也能定位到具体环节。6.1 JAX 的基本功能测试先验证自动微分和 JIT 是否正常import jax import jax.numpy as jnp from jax import grad, jit def f(x): return x ** 3 2 * x # 自动微分d/dx (x^3 2x) 在 x2 处应为 14 print(grad(f)(2.0)) # JIT 编译后重跑 print(jit(f)(2.0)) print(jax version:, jax.__version__)如果grad输出不是14.0那说明 JAX 的自动微分链路有问题基本上要回退到版本检查。6.2 设备识别测试print(jax.devices())GPU 机器上应该能看到CudaDeviceCPU 机器则是CpuDevice。这一步出错通常发生在装了 CPU 版的 jax、或者显存被其他进程占满导致设备不可用。前者回看第 3 章后者检查正在跑的进程。6.3 EvoRL 的导入测试接下来测试 EvoRL 本身能不能正常导入import evorl from evorl.agent import Agent from evorl.envs import make_evorl_env不同版本 API 略有差异以你装的那个版本实际支持的导入路径为准。只要这一步不报ModuleNotFoundError就说明依赖层的包都被正确引入了。6.4 跑一个最小训练 demo最关键的验证是让 EvoRL 真跑一段训练。通常项目仓库里会带 demo 目录进入后按 README 给的命令启动一个最简单的进化策略训练脚本观察有没有正常的日志输出。cd demo python train.py --config-name配置名这里配置名以仓库 README 和 configs 目录下的实际文件名为准。训练脚本跑起来之后只要日志在更新说明 JAX 的计算和 EvoRL 的训练循环是通的。跑个二三十轮确认不崩安装就算真正完成。7. 安装过程中最常踩的坑以及我的排查链路最后分享几个我实际遇到过的坑。每个坑我都尽量还原排查过程而不只是给答案因为同样的现象可能由不同原因引起思路比命令重要。7.1 jaxlib 版本错位现象import jax 时直接报ValueError: jax and jaxlib must have the same version或者类似Jax runtime was built with ...的信息。排查思路先看pip list | grep jax确认 jax 和 jaxlib 的版本号是否一致。不一致的原因多半是 pip 在装 EvoRL 时单独升级了其中一个而另一个还留在旧版。解决办法是让它们重新对齐pip install -U jax jaxlibGPU 场景再补一条pip install --upgrade jax[cuda12]把插件也一起更新。7.2 GPU 明明存在jax.devices() 却只有 CPU现象nvidia-smi正常但jax.devices()输出只有CpuDevice。排查思路第一步看pip list | grep jax确认装的到底是不是jax-cuda12系列插件。如果只是普通 jax说明当初装的时候 pip 选了 CPU 版用 GPU 版命令强制重装。第二步检查环境变量CUDA_VISIBLE_DEVICES如果被设成了空字符串设备自然不可见。第三步看显存其他进程把显存占满后JAX 初始化时可能直接放弃 GPU。7.3 protobuf 版本冲突现象导入某些环境模块时报TypeError: Descriptors cannot not be created directly或者Unknown name for this Entity之类的错误。排查思路这是 protobuf 版本和 gym 系列包不匹配的经典问题。EvoRL 的依赖链里有不同包对 protobuf 的上下限要求不一致尤其容易发生在混装了老 gym 和新 protobuf 的时候。我的处理方式是用pip list | grep protobuf看当前版本然后调整到依赖允许的区间pip install protobuf3.20.3不过 3.20.3 不是万能解具体版本以pip check给出的提示为准。7.4 装完 EvoRLJAX 又被覆盖回 CPU 版现象安装前jax.devices()还能看到 GPU装完 EvoRL 之后再看只剩 CPU 了。排查思路这是 EvoRL 依赖解析时动到了 JAX。我当时处理的办法是装完立刻重装一次 GPU 版顺序就是第 5 章写的两行命令。如果以后频繁遇到可以把它写进项目目录下的 requirements 文件每次重建环境时统一装。7.5 Windows 下源码安装失败现象在 Windows PowerShell 里执行源码安装编译某些依赖时报 C 编译器错误。排查思路JAX 生态里的不少包对 Windows 的原生编译链支持不完善。不要硬刚直接开 WSL2 或者换一台 Linux 机器。装环境的时间也是时间这点早点放弃反而高效。最后再分享一个小技巧整套环境装好之后把nvidia-smi、pip check、第 6 章的四步验证脚本存成一个check_env.sh或verify_env.py每次在新机器上部署时跑一遍五分钟就能确认整条链路是否健康。我自己吃过太多当时能跑、过两周就跑不了的亏这套最小验证脚本帮我省了很多排查时间。