给一台装了两张不同显卡的 Windows 机器搭 JAX 计算环境,撞到的坑几乎都是同一类:你查到的东西,不等于实际的东西。最典型的一个是 jax.devices() 对两张不同的卡报了同一个型号名——按名字路由,16 GB 的 batch 会被扔到 8 GB 的卡上。
一、先省掉你两天:WSL2 不是偏好,是硬性要求
在 Windows 上跑 JAX + CUDA,第一个要决定的是走原生 Windows Python 还是 WSL2。这个问题有一个确定的答案,而且不需要试:
jax-cuda12-pjrt 和 jax-cuda13-pjrt 在 PyPI 上只发 manylinux 轮子,没有 win_amd64。
后果不是”装不上”——是装得上但没用。pip 会安静地给你一个 CPU 版的 jaxlib,代码跑得起来,jax.devices() 返回一个 CPU 设备,两张卡一张都用不上。如果你不特意去看设备列表,很可能跑了一整天才发现速度不对。
判断”能不能不用 WSL”时直接查这一条:去 PyPI 看目标包有没有对应平台的 wheel。有没有 CUDA、驱动版本对不对、环境变量设没设,都是这个问题之后的事。
二、两张不同的卡,报同一个名字
这台机器上一张 16 GB、一张 8 GB,型号不同。而在 WSL2 下:
import jax
for d in jax.devices():
print(d.id, d.device_kind)
# 0 NVIDIA RTX 4080
# 1 NVIDIA RTX 4080 ← 这张实际是 8 GB 的另一款
device_kind 对两张卡都报了同一个型号。
这个坑的危险之处在于它的失败方式:如果你按卡名做任务路由——”型号是 4080 的给大 batch”——调度器会把一个需要 16 GB 的任务分给那张 8 GB 的卡,然后你收到一个 OOM。而 OOM 的信息里不会告诉你”你以为你选的是另一张卡”。
实测的真实映射是按 PCI 地址排序的:
device[0] → PCI 01:00.0 → 16 GB 那张
device[1] → PCI 06:00.0 → 8 GB 那张
要可靠地区分,只有两条路:用 device id(顺序稳定,但你得先确认过一次映射),或者绕过 JAX 直接问驱动——WSL 下 nvidia-smi 在 /usr/lib/wsl/lib/nvidia-smi,它报的显存和型号是对的。
更一般的规则:任何”设备自述”的字段都不要用来做资源决策。型号名、驱动上报的描述、容器里看到的 hostname——这些在虚拟化和转译层下面经常是错的或被复用的。要做决策就去问那个真正掌握资源的东西,并且把映射关系记下来。
三、后台进程会被静默杀掉,日志停在半路
在 WSL 发行版里 nohup 起一个长任务,然后关掉那个 wsl.exe 窗口。任务会死。
这本身不算意外,意外的是死法:进程消失,日志文件停在半路,没有任何终止记录。更糟的是 systemd-run 起的瞬态 unit 也一起消失,连 journal 里都不留痕迹。你回来看到的是一个跑了 40% 就没有下文的日志文件,看不出是崩了、被杀了、还是机器重启了。
原因是 WSL 在最后一个会话结束时会把整个发行版关掉,那不是一次进程终止,是整个虚拟机停机。
两个能用的做法:
- 在一个保持打开的 SSH 会话里前台跑——SSH 连接本身就是那个”会话”;
- 装成开机自启的正式 systemd service,不是
systemd-run的瞬态 unit。
如果队列本身也在发行版里,它同样会随之消失。这台机器上的做法是让一个计划任务在开机时跑 wsl --exec sleep infinity,把发行版常驻住——用一个永远不退出的会话,堵住”最后一个会话结束”这个条件。
四、第一版调度策略把大卡饿死了
写作业调度器时,第一版用的是”最佳适配”:按卡的剩余容量排序,把任务放进能装下它的最小的那张卡。这是装箱问题的标准启发式,目的是减少碎片。
结果是两个小任务都被塞进了 8 GB 那张卡,而 16 GB 的那张整个空转。
问题在于最佳适配优化的是”空间碎片”,而这里真正稀缺的资源是并行度——两张卡本来可以同时算两个任务。把两个任务挤在一张卡上,即使装得下,也把吞吐砍了一半。
改成两级判据:
- 优先给完全空闲的卡;
- 没有空闲卡时,才在有余量的卡里挑剩余容量最大的。
第二条也从”最小可容纳”翻转成了”最大剩余”——在只有两张卡的场景下,减少碎片带来的收益远小于保住下一个任务还塞得进去的概率。
套用教科书启发式之前,先确认它优化的目标就是你稀缺的那个资源。最佳适配为内存碎片而生,而这里的瓶颈从来不是碎片。
五、WSL 与 Windows 之间用 stdin 管道会死锁
调度器需要从 Linux 侧调 Windows 侧的命令。直觉写法是把数据从 stdin 管进去:
echo "$payload" | powershell.exe -Command '$in = [Console]::In.ReadToEnd(); ...'
这会挂住。ReadToEnd() 要等一个 EOF,而跨越 WSL 与 Windows 的那条管道在互操作层上不可靠地传递关闭事件——写入端结束了,读取端不一定收到。
而且它的表现是”卡住”而不是报错,所以第一反应通常是去查 PowerShell 的执行策略或者引号转义,查错方向。
可靠的做法是不用管道,用文件:写到一个两边都能看到的路径(WSL 侧的 /mnt/c/... 对应 Windows 侧的 C:\...),把路径当参数传过去。多一次磁盘往返,换掉一个不确定的同步原语。
六、那条只在有人登录时才存在的桥
还有一个更隐蔽的:WSL2 用的是 NAT,它的 IP 每次启动都可能变。要从局域网 SSH 进发行版,就得在 Windows 侧维护一条端口转发,并且在 WSL 重启后重新指向新 IP。
做法是一个脚本 + 一个计划任务。而我最初把那个计划任务的触发条件设成了”用户登录时”。
后果是:机器开着、WSL 跑着、任务在排队,但只要没人登录过桌面,这条桥就不存在,从外面连不进去。而机器本身一切正常,唯一的症状是”连不上”。
解法是准备一条不依赖端口转发也不依赖登录的备用通道——经由 Windows 主机的 SSH 转发进去,它还能顺带把已经关掉的发行版拉起来。
顺带一个 Windows sshd 的坑:如果登录用的账号属于管理员组,Match Group administrators 会让 sshd 改去读 C:\ProgramData\ssh\administrators_authorized_keys,而不是用户目录下的 .ssh/authorized_keys。写进后者完全不生效,而且日志里只是普通的公钥认证失败。那个文件的 ACL 还必须只有 Administrators 和 SYSTEM,多一个条目 sshd 就拒绝使用它。
七、这些坑的共同形状
| 你观测到的 | 实际的 |
|---|---|
| pip 装成功了 | 装的是 CPU 版,两张卡都没用上 |
device_kind 说是同型号 |
一张 16 GB 一张 8 GB |
| 日志停了,没有错误 | 整个发行版被关停 |
| 调度器说”装得下” | 装得下,但吞吐砍半 |
| 命令卡住了 | 管道的 EOF 没传过去 |
| 机器在跑,服务在跑 | 桥没建,因为没人登录 |
| 公钥认证失败 | sshd 在读另一个文件 |
虚拟化和转译层的代价不只是性能,还有可观测性:每多一层,”系统告诉你的”和”实际发生的”之间就多一次翻译,而翻译是会丢信息的。WSL2 之下,GPU 走的是一条转译过的驱动路径,进程生命周期挂在一个 Windows 侧的会话上,网络是 NAT 的,文件系统跨了两个世界——上面每一个坑都恰好落在其中某一层的接缝上。
实用的推论只有一条:在这种环境里,任何要用来做决策的事实,都值得用第二种独立的方法再确认一次。卡的容量问驱动而不是问框架,进程活着不活着看它有没有在写东西而不是看它启动成功了没有,通道通不通从外面实测一次而不是看配置写没写。
References
相关阅读:显存估算的三次修正(”参数量≠显存”是同一类观测偏差)、本地大模型部署:模型与推理引擎、Gemma 4 全变体解析。
