做强化学习、神经进化方向研究的读者,一定绕不开一个名字:JAX。作为Google在深度学习领域的一张王牌,它用NumPy风格API直接给出了极致的自动微分、JIT编译和GPU/TPU并行能力,DeepMind大量论文的核心源码都跑在JAX上面。而EvoRL,则是近几年基于JAX长出来的进化强化学习库,把进化策略和强化学习算法统一在同一套框架里,特别适合想做大规模并行实验的小团队和实验室。
问题在于,JAX的安装比PyTorch麻烦一个量级。它跟Python版本、CUDA、cuDNN、甚至显卡驱动版本绑定得很死,装错一个环节就得推倒重来;EvoRL又是比较新的项目,依赖链长,文档未必跟得上版本变化。我这两套东西来回装过好几次,把完整的安装流程、版本匹配规则、验证方法和常见报错一次性整理出来。无论你是想先拿纯CPU体验一下语法,还是准备直接用GPU跑实验,照着下面的顺序走,基本不会出大问题。
1. 先搞清楚这几件事:JAX到底牛在哪,EvoRL又为什么必须用它
1.1 JAX的核心定位,和它跟PyTorch的差异
JAX常被简单理解成“带自动微分的NumPy”,这个说法方向没错,但低估了它。它真正的杀手锏是一套组合变换:grad只负责求梯度,但配合jit可以把整个训练循环编译成XLA指令,配合vmap可以把处理单条数据的函数自动向量化到批数据,配合pmap可以把计算均匀摊到多块GPU上。很多人在JAX里写代码,从单样本逻辑开始,改很少几行就变成了大规模并行版本,这在PyTorch里通常要手写分布式逻辑。
差异最直观的地方体现在函数式风格上。PyTorch默认是面向对象、带可变状态的模型类,JAX则强调纯函数、不可变数组,模型参数一般作为参数传进函数而不是挂在对象上。这个设计让JIT优化变得极其激进,也让代码推理更简单——只要函数输入输出类型确定,编译后的执行计划往往非常高效。我用一个生活化类比:PyTorch像自动挡汽车,踩油门就能跑,方便省心;JAX更像手动挡的性能车,你需要先了解换挡逻辑,比如想清楚jit之后哪些操作不能做,但一旦掌握了,你能把硬件的每一分性能都压榨出来,尤其在GPU这种并行设备上,收益是数量级的。
1.2 EvoRL的设计思路:进化与强化学习为什么要在JAX上融合
EvoRL的全称是Evolutionary Reinforcement Learning,它不是某一个单一算法,而是一套进化计算与深度强化学习融合的框架。传统进化策略比如OpenES、CMA-ES,不依赖反向传播,直接对参数做扰动再评估适应度,优点是稳定、全局搜索能力强,但深度网络参数量巨大,一点扰动就要跑完一整轮完整评估,计算开销非常夸张。EvoRL做的事情,就是把这类进化算法和PPO、SAC等主流强化学习算法放到同一个框架里,让两者各取所长。
为什么它必须基于JAX?核心原因是进化策略天然需要大规模并行:每一代的每个种群成员都要独立评估,这种模式用JAX的vmap和pmap几乎可以零成本展开。如果换成PyTorch,你得自己写进程池、自己处理采样和梯度同步,代码量翻好几倍。再加上JAX的JIT编译,几百个actor同时跑在GPU上时,单步吞吐量的优势是跨量级的。所以EvoRL从底层就绑定JAX,而不是把它当成一个可以随时替换的后端。你甚至可以说,没有JAX的并行原语,EvoRL这类项目的工程成本会高到劝退大多数研究组。
1.3 安装前必须先确定的版本组合,顺序反了基本返工
很多人装JAX失败,不是命令写错,而是没先确定版本组合。JAX官方提供CPU、CUDA 11、CUDA 12、TPU几种不同的构建产物,它们对底层CUDA版本和驱动版本都有要求。我在动手前强烈建议把下面这个版本组合表看清楚:
| 使用场景 | 推荐Python版本 | 安装命令 | 说明 |
|---|---|---|---|
| 纯CPU开发调试 | 3.9 - 3.12 | pip install -U jax jaxlib | 最稳,适合语法学习和无GPU服务器 |
| Linux + NVIDIA GPU | 3.9 - 3.12 | pip install -U "jax[cuda12]" | 驱动需支持CUDA 12,优先推荐 |
| 旧驱动老环境 | 3.9 - 3.12 | pip install -U "jax[cuda11]" | 驱动只支持CUDA 11时使用 |
| Apple Silicon | 3.9 - 3.12 | pip install -U jax | 默认CPU,Metal支持仍在完善 |
| TPU | 3.9 - 3.12 | pip install -U "jax[tpu]" | 需要云TPU环境,一般用不到 |
JAX的Python支持覆盖3.9到3.12,EvoRL这类项目通常要求JAX不低于0.4.x。如果你的驱动比较新,就优先选CUDA 12;如果老旧环境不打算动驱动,再看CUDA 11。驱动版本太低会导致装上wheel之后运行时报“CUDA driver too old”,这个错不是重装JAX能解决的,只能升级驱动或者换低版本CUDA的wheel。所以我的建议顺序是:先查显卡驱动支持的CUDA版本,再决定JAX装哪个分支,最后按EvoRL仓库要求核对JAX版本。顺序反了的,几乎都会返工。
2. JAX安装全流程:从纯CPU到CUDA GPU,一步步来
2.1 创建虚拟环境,别再把依赖堆进系统Python
我强烈不建议图省事把JAX直接装到系统Python里。JAX的依赖(numpy、scipy、opt-einsum)跟深度学习生态有大量交叉,装在系统环境里早晚出依赖冲突,而且系统Python往往被系统包管理器锁定了部分包版本,升级权限也不足。用conda创建一个独立环境,是最保险的开局:
conda create -n evorl python=3.11 -y conda activate evorl这里Python版本选3.11,是目前JAX和EvoRL支持度最均衡的版本。如果实验室服务器上只有3.9或3.12,也能跑,但后面遇到“某个依赖版本不兼容”的概率会高一点。创建完环境后,建议先把pip升级一下:pip install -U pip setuptools wheel。有些老环境里pip版本太低,在解析JAX这种混合依赖时会出现莫名其妙的冲突提示,升级后很多问题会自动消失。
2.2 纯CPU版本安装与快速验证
如果只是想试一下JAX的语法、跑跑教学代码,或者在没有GPU的服务器上先开发调试,装CPU版本就够了。安装命令非常简单:
pip install -U jax jaxlib新版也可以用pip install "jax[cpu]",效果等价。装完之后打开Python验证:
import jax print(jax.__version__) print(jax.devices())正常会输出类似0.4.36的版本号,以及[TFRT_CPU_0]或[CpuDevice]之类的设备列表。看到CPU设备,就说明基础安装已经成功。
我实测过,CPU版本的JAX在普通笔记本上跑小规模线性回归、小型MLP是没有任何问题的。但它有一个特点:所有操作默认走XLA编译,第一次执行某个函数时会有一个编译预热过程,看起来比PyTorch慢,第二次开始明显变快。刚上手的人看到第一次运行慢,千万别急着怀疑装坏了,多跑几次对比一下就明白了。预热时间在CPU上通常是几百毫秒级别,在GPU上第一次跑大模型可能达到几十秒,这些都是正常的。
2.3 GPU版本安装:CUDA版本怎么选、命令怎么写
GPU版本安装前,先确认两个东西:NVIDIA驱动版本和CUDA Toolkit版本。其实对JAX来说,系统里装没装CUDA Toolkit不是最关键的,因为官方wheel自带运行时依赖,关键点是你的显卡驱动要支持目标CUDA版本。执行nvidia-smi,看右上角的“CUDA Version”,那就是驱动支持的最高CUDA版本。
如果驱动支持CUDA 12,推荐直接安装:
pip install -U "jax[cuda12]"这会在安装jax的同时,拉取兼容的jaxlib和CUDA运行时依赖。如果驱动只支持到CUDA 11,换成:
pip install -U "jax[cuda11]"安装过程中如果网速不快,可能在看jaxlib下载那一步卡很久,因为jaxlib的wheel包含预编译的XLA运行时和CUDA库,体积通常在500MB到1GB之间。遇到下载超时或速度极慢,建议临时换成国内pip镜像源,比如清华源或阿里源,-i参数指定即可。
验证GPU安装是否成功,执行:
python -c "import jax; print(jax.default_backend()); print(jax.devices())"如果输出中看到gpu和[CudaDevice(id=0)],说明GPU已经接管了计算。如果仍然显示cpu或CpuDevice,多半是环境变量指向了旧的jaxlib,或者conda环境里残留了CPU版本包。这时候先执行pip uninstall -y jax jaxlib,再重新执行上面的GPU安装命令,基本能解决。
2.4 顺手做一个自动微分和JIT测试
设备验证通过后,我还习惯再做一轮功能验证,确认自动微分和JIT都正常,因为这两个是后来EvoRL依赖的重中之重:
import jax import jax.numpy as jnp from jax import grad, jit def simple_loss(w, x, y): return jnp.mean((x @ w - y) ** 2) w = jnp.ones((3, 1)) x = jnp.ones((5, 3)) y = jnp.ones((5, 1)) g = grad(simple_loss)(w, x, y) loss = jit(simple_loss)(w, x, y) print("grad shape:", g.shape) print("loss:", loss)如果这个脚本跑得通,说明JAX核心功能完好,后面装EvoRL就等于成功了一大半。我还会顺手跑一次jax.device_count(),确认机器上到底有几块GPU,这个数字在EvoRL做pmap并行时会用到,提前知道有助于后续配置并行度。
3. EvoRL安装与测试:装完还要能真正跑起来
3.1 EvoRL的两种安装方式
EvoRL的安装一般有两种方式。第一种是从PyPI直接安装:
pip install evorl这种方式适合只想把EvoRL当工具包调用、不打算改源码的场景。第二种是从GitHub把源码clone到本地,再用可编辑模式安装:
git clone https://github.com/evolutionrl/evorl.git cd evorl pip install -e .如果你的实验需要在算法层做较多改动,或者需要逐行调试EvoRL内部实现,我建议用第二种方式,这样改了代码后不用重新安装,修改即时生效。第二种方式对git和网络下载的依赖较多,但等待是值得的,因为你能直接看到每个模块的源码,遇到问题时可以快速定位到底层逻辑。
我个人的习惯是装之前先看一眼EvoRL仓库里的requirements文件,确认它对jax、numpy、gymnasium等依赖的版本范围。再把当前环境里已有的jax版本跟它对照一下。很多报错并不是EvoRL本身有bug,而是它要求的jax版本和环境里的版本差了一两个小版本,导致某个API签名对不上。
3.2 依赖冲突检查,别让flax和jax各自为政
EvoRL的依赖列表除了JAX,通常还包括flax(神经网络库)、optax(优化器)、gymnasium(强化学习环境接口)、ml_collections(配置管理)等。这些包之间偶尔会打架,尤其是flax和jax的版本要保持同代,否则flax内部调用不存在的jax API时会抛出AttributeError。这一点在你只装JAX时感受不到,但一跑EvoRL就会立刻暴露。
检查依赖冲突有一个笨但有效的办法:安装完后执行pip check。这个命令会扫描当前环境里所有包的依赖关系,把版本不合、缺少依赖的问题直接列出来。如果输出“No broken requirements found”,说明依赖层面已经没问题,后面再遇到报错就可以安心地排查算法逻辑和代码路径。
如果环境里之前装过其他深度学习框架,强烈建议做一次这个检查。我自己曾经因为一个老版本TensorFlow残留的依赖,导致gymnasium一直给EvoRL返回异常格式的动作空间,折腾了大半天,最后新建环境才解决。这个教训让我后来对所有涉及强化学习的项目都保持环境洁癖。
3.3 跑第一个EvoRL示例,体验一次完整的JIT预热
装完后验证是否可用,可以打开Python终端执行:
import evorl print(evorl)能成功导入,再接着尝试导入内部的子模块,比如from evorl.algorithms import ppo这类路径,确认没有缺失依赖。更完整的上手方式,是直接看EvoRL仓库examples目录下的训练脚本,然后运行一个简单的经典控制任务,比如CartPole或倒立摆:
python examples/train_ppo.py --config=examples/configs/ppo_cartpole.yaml具体的配置文件名要以你clone下来的仓库为准,不同版本可能略有差异。第一次跑会有一段较长的时间花在JIT编译上,之后终端会开始打印训练轮次、奖励均值等指标。如果你看到奖励曲线逐步上升,说明JAX、EvoRL和环境库这一整条链路已经彻底打通。
这一步我建议耐住性子把它跑完,不要只确认能导入包就关掉。后续实验改动都是在这个基础上进行的,把第一次的编译耗时、日志格式、配置加载方式都弄明白,后面改参数、加算法会快很多。如果你打算在EvoRL里用Brax或Gymnax这些JAX生态环境库,也建议在同一个虚拟环境里提前装好,它们是EvoRL官方示例中经常出现的依赖项。装这几个扩展库要注意版本号尽量选中性版本,避免过旧或过新导致环境API不匹配。
4. 我在安装和测试中踩过的坑:常见报错与对应解法
4.1 最经典的“No module named 'jaxlib'”报错
“No module named 'jaxlib'”大概是JAX安装问题里出现频率最高的一条。出现这个错误,基本可以断定jax和jaxlib两个包没有对齐:要么只装了jax,要么两个包版本不一致。解决方法很直接,把两个包同时卸载再一起安装:
pip uninstall -y jax jaxlib pip install -U "jax[cuda12]"不要指望单独更新jaxlib能解决,它的版本必须和jax严格匹配。官方在PyPI上会把这两个包绑定发布,所以最稳妥的方式就是通过方括号里的extra参数一次装齐。另外一个容易踩的点是,pip在解析时可能把jaxlib当成了系统已有的旧包而不去更新。如果pip show jaxlib显示的版本和安装时不一致,也要先卸载再重装。
4.2 CUDA版本太旧或驱动不匹配,运行时直接崩溃
在GPU版本上最常见的运行时报错是:
RuntimeError: CUDA driver too old / CUDA driver version is insufficient for CUDA runtime version这个错误的原因很直接:显卡驱动支持的最高CUDA版本,低于JAX wheel里编译时使用的CUDA版本。解决办法有几个方向,优先级从高到低:
- 更新显卡驱动,让驱动支持目标CUDA版本,注意去NVIDIA官网找对型号的驱动包;
- 如果不想动驱动,就换对应CUDA 11的JAX wheel,不要反向硬顶CUDA 12;
- 检查多环境中是否配置过CUDA相关的环境变量,某些shell配置里写死的旧版CUDA路径会干扰运行时加载。
这类问题比安装时直接报错更难排查,因为它发生在运行时,而且报错信息直指CUDA,不会主动提示JAX版本问题。我总结的排查顺序是:先用nvidia-smi看驱动支持的CUDA版本,再确认jaxlib对应了哪个CUDA构建,两边对不上就直接处理驱动或换wheel。
4.3 JIT编译带来的Python控制流陷阱
装好之后,实际跑EvoRL或自己写JAX代码时,另一个高频问题来自JIT的静态编译特性:Python的原生类型(int、bool、list)在JIT编译时必须是确定性的,如果函数里依赖运行时的Python if分支或可变全局变量,编译会报错或给出奇怪结果。常见错误像ConcretizationTypeError、TracerBoolConversionError,都是同一个原因。
解决办法是:把非数值逻辑写成jnp.where这种向量化操作,或者把需要动态判断的值作为静态参数传给jit。这个坑在EvoRL中尤其常见,因为强化学习的动作采样、PPO的clip判断,新手往往会下意识去写Python if。如果你在EvoRL示例代码基础上改自己的环境时遇到这类错误,先检查是不是把Python控制流放进了被jit装饰的函数里。
4.4 WSL2用户需要注意的Windows特有问题
在Windows下跑JAX,官方不支持原生的Windows wheel,推荐做法是安装WSL2。这里有个我踩过的细节:WSL2里的显卡驱动其实是在Windows侧安装的,WSL内部并不需要单独装驱动,但前提是Windows侧驱动版本足够新,否则WSL内部同样会报CUDA版本过旧。
如果WSL2里执行nvidia-smi显示不出显卡,多半是Windows没装WSL对应的GPU驱动,或者路径配置有问题。建议先专心解决驱动可见性,再考虑装JAX,因为驱动不可见时,即使你装的是GPU版本JAX,运行时也会静默回退到CPU,而你不会第一时间发现。很多WSL2用户跑实验跑完一个通宵,第二天一看log才发现全程用的CPU,这种浪费完全是可以通过先检查jax.devices()来避免的。
5. 装完之后,怎么确认整条链子是通的
5.1 一个综合体检脚本,把四大核心能力一次验证完
经过上面的步骤,环境照理说已经可用了。但我还会花两分钟跑一个综合体检,把JAX设备、自动微分、jit、vmap四个核心能力一次验证到位:
import jax import jax.numpy as jnp from jax import grad, jit, vmap print("JAX version:", jax.__version__) print("Backend:", jax.default_backend()) print("Devices:", jax.devices()) def f(w, x): return jnp.sum(jnp.tanh(x @ w)) batch = vmap(lambda x: f(jnp.ones((4, 4)), x)) print("vmap out:", batch(jnp.ones((8, 4))))输出正常,说明JAX的计算核心、编译核心、批处理能力全部在线。对EvoRL来说,这三者是后续训练能否快速跑起来的前提。如果jax.default_backend()返回的是gpu,而jax.devices()列出的设备数量符合预期,那就可以放心继续配置EvoRL的并行规模了。
5.2 EvoRL的应用层验证,三个清单逐项过
JAX验证完,再回到EvoRL。我在实际评估时不会只做导入测试,而是确认三件事:能导入evorl包;能正常读取配置文件;能在CPU或GPU上启动一个小的训练脚本并完成至少几个完整迭代。只有这三项都通过,才能认定安装成功。只做导入验证很多环境问题会被掩盖,因为有些依赖只有实际运行时才会触发导入,比如环境包装器、环境注册表、特定的渲染后端。
这些验证做完后,我还有一个习惯:记录当前环境的关键版本到一个文本文件里,包括Python版本、JAX版本、jaxlib版本、EvoRL版本、显卡驱动版本。这个动作看起来简单,但能救大命。过几个月再回来看实验,如果复现不出结果,先对照这份版本记录,往往能迅速定位是环境漂移还是代码改动引起的差异。
5.3 顺带聊聊whisper jax这类JAX生态的热门应用
这里提一个能让你的安装体验变得有价值的热门应用:whisper jax。它是OpenAI Whisper语音识别模型的JAX实现,社区里很多人就是冲着“把语音识别推理速度翻倍”来的。如果你装好了JAX,顺手跑一次whisper jax的demo,大概率能直观感受到JAX在GPU上的加速红利。EvoRL是训练和演化方向的代表,whisper jax则是推理方向的代表,两者放在一起看,能帮你更清楚把握JAX在不同任务中的定位。
不过whisper jax自身依赖比较多,包括transformers、librosa这些音频处理库,我建议在已经通过EvoRL验证的环境里新建一个独立虚拟环境再装,不要和EvoRL混在一起。语音识别库更新非常频繁,经常会把依赖升到新版,和强化学习库放一起很容易互相拖垮。我在本地就是两个环境分开维护,遇到实验需求切换时用conda activate切换,省心很多。
最后再分享一个我的习惯
装JAX和EvoRL这件事,难不在“敲命令”,难在版本匹配和环境隔离。我自己每个项目都开新的conda环境,装完之后写一个环境说明文档,把python版本、jax版本、cuda分支、显卡驱动版本都记录在案。这样过一两个月回去再看实验,不用对着报错猜环境。前面这些坑基本都踩过一遍,如果你照着这个流程走,多数问题都能在第一次安装时绕开。把基础环境这块地基打稳,后面无论是跑EvoRL的进化实验,还是去尝试whisper jax这类生态应用,都会顺利得多。