Setting up a JAX compute environment on a Windows box with two different GPUs turned up a series of traps that all share one shape: what you can observe is not what is actually there. The clearest of them is jax.devices() reporting the same model name for two different cards — route work by that name and a 16 GB batch lands on the 8 GB card.
1. Saving you two days: WSL2 is a requirement, not a preference
Running JAX with CUDA on Windows starts with one decision: native Windows Python or WSL2. That question has a definite answer, and you do not need to experiment to find it:
jax-cuda12-pjrt and jax-cuda13-pjrt publish only manylinux wheels on PyPI. There is no win_amd64 build.
The consequence is not “it won’t install” — it is that it installs and does nothing. pip quietly gives you a CPU-only jaxlib, the code runs, jax.devices() returns a CPU device, and neither GPU is used at all. Unless you deliberately inspect the device list, you can lose a full day before noticing the speed is wrong.
When deciding whether WSL can be avoided, check exactly this: go to PyPI and see whether the package ships a wheel for your platform. Whether CUDA is present, whether driver versions line up, whether environment variables are set — all of that comes after.
2. Two different cards reporting the same name
This machine has one 16 GB card and one 8 GB card, different models. Under WSL2:
import jax
for d in jax.devices():
print(d.id, d.device_kind)
# 0 NVIDIA RTX 4080
# 1 NVIDIA RTX 4080 <- actually the 8 GB card, a different model
device_kind reports the same model for both cards.
What makes this dangerous is its failure mode. Route tasks by card name — “give the big batches to the 4080” — and the scheduler hands a job needing 16 GB to the 8 GB card, and you get an OOM. Nothing in that OOM tells you that you thought you had selected the other card.
The real mapping, measured, follows PCI address order:
device[0] -> PCI 01:00.0 -> the 16 GB card
device[1] -> PCI 06:00.0 -> the 8 GB card
There are exactly two reliable ways to tell them apart: use the device id (the order is stable, but you have to confirm the mapping once), or bypass JAX and ask the driver — under WSL, nvidia-smi lives at /usr/lib/wsl/lib/nvidia-smi, and its memory and model figures are correct.
The general rule: never make resource decisions from a field the device reports about itself. Model names, driver-supplied descriptions, the hostname visible inside a container — under virtualisation and translation layers these are routinely wrong or reused. Ask whatever actually owns the resource, and write the mapping down.
3. Background processes are killed silently, logs stop mid-file
nohup a long job inside the WSL distribution, then close that wsl.exe window. The job dies.
That much is unsurprising. What surprises is how it dies: the process vanishes, the log file stops mid-way, and there is no termination record anywhere. Worse, transient units started with systemd-run vanish too, leaving nothing in the journal either. What you come back to is a log that stops at 40% with no way to tell whether it crashed, was killed, or the machine rebooted.
The cause is that WSL shuts the whole distribution down when the last session ends. That is not a process termination; it is the VM stopping.
Two things that work:
- Run in the foreground inside an SSH session you keep open — the SSH connection is the session;
- Install a real systemd service that starts at boot, not a
systemd-runtransient unit.
If the queue itself also lives inside the distribution, it disappears along with everything else. The approach here is a scheduled task that runs wsl --exec sleep infinity at startup, pinning the distribution up — one session that never exits, blocking the “last session ended” condition from ever becoming true.
4. The first scheduling policy starved the big card
The first version of the job scheduler used best fit: sort cards by remaining capacity and place each job on the smallest card that can hold it. That is the textbook bin-packing heuristic, aimed at reducing fragmentation.
The result was that two small jobs both landed on the 8 GB card while the 16 GB card sat completely idle.
Best fit optimises for space fragmentation, and the scarce resource here is parallelism — two cards could have been running two jobs simultaneously. Packing both onto one card, even though they fit, halves the throughput.
The replacement is a two-level rule:
- Prefer a completely idle card;
- Only when no card is idle, pick the one with the most remaining capacity.
Note that the second rule also flipped from “smallest that fits” to “largest remaining.” With only two cards, the benefit of reducing fragmentation is far smaller than the benefit of keeping room for whatever arrives next.
Before reaching for a textbook heuristic, check that what it optimises is what you are actually short of. Best fit exists for memory fragmentation, and fragmentation was never the bottleneck here.
5. stdin pipes deadlock across the WSL/Windows boundary
The scheduler needs to invoke Windows-side commands from Linux. The instinctive form pipes data in through stdin:
echo "$payload" | powershell.exe -Command '$in = [Console]::In.ReadToEnd(); ...'
This hangs. ReadToEnd() waits for EOF, and a pipe crossing the WSL/Windows interop layer does not reliably propagate the close — the writer finishes, and the reader may never learn about it.
And because it presents as a hang rather than an error, the instinctive next step is to go looking at PowerShell execution policy or quoting. Wrong direction entirely.
The reliable form is a file instead of a pipe: write to a path both sides can see (/mnt/c/... on the WSL side is C:\... on the Windows side) and pass the path as an argument. One extra disk round trip in exchange for removing an unreliable synchronisation primitive.
6. A bridge that only exists while somebody is logged in
Subtler still: WSL2 is NAT’d, and its IP can change on every start. Reaching the distribution over SSH from the LAN therefore needs a port forward on the Windows side, re-pointed at the new IP after each WSL restart.
That is a script plus a scheduled task. I originally set the task’s trigger to “at user logon.”
The consequence: the machine is on, WSL is running, jobs are queued — but unless somebody has logged into the desktop, the bridge does not exist and nothing can connect. Everything about the machine is healthy; the only symptom is that you cannot reach it.
The fix is having a fallback path that depends on neither port forwarding nor a logon — forwarding in through the Windows host’s own SSH, which can also start the distribution if it has already shut down.
One related Windows sshd trap: if the login account belongs to the administrators group, Match Group administrators makes sshd read C:\ProgramData\ssh\administrators_authorized_keys instead of the user’s .ssh/authorized_keys. Keys written to the latter have no effect at all, and the log shows only an ordinary public-key authentication failure. That file’s ACL must also contain nothing but Administrators and SYSTEM — one extra entry and sshd refuses to use it.
7. The shape they all share
| What you observe | What is true |
|---|---|
| pip installed successfully | it installed the CPU build; neither GPU is in use |
device_kind says the same model |
one is 16 GB, one is 8 GB |
| the log stopped, no error | the whole distribution was shut down |
| the scheduler says it fits | it fits, and throughput is halved |
| the command is hanging | EOF never crossed the pipe |
| machine up, services up | no bridge, because nobody logged in |
| public-key auth failed | sshd is reading a different file |
The price of virtualisation and translation layers is not only performance — it is observability. Each layer adds another translation between “what the system tells you” and “what is happening,” and translations lose information. Under WSL2 the GPU arrives over a translated driver path, process lifetime hangs off a Windows-side session, the network is NAT’d, and the filesystem spans two worlds. Every trap above sits precisely on one of those seams.
The practical takeaway is a single sentence: in an environment like this, any fact you intend to make a decision from is worth confirming a second time by an independent method. Ask the driver for card capacity rather than the framework; judge whether a process is alive by whether it is still writing rather than by whether it started; test whether a channel works from outside rather than reading the config.
References
- JAX installation docs: platform support matrix
- systemd support and lifecycle in WSL
- Windows OpenSSH: administrators_authorized_keys and its ACL requirements
Related: three corrections to VRAM estimation (“parameter count is not memory” is the same class of observation error), local LLM deployment: models and inference engines, and the Gemma 4 variants.
