TPUs for Reinforcement Learning and Behavioral Modeling
How to squeeze a year of compute time into a single day.

Given that you have created a TPU deployment on GCP, you are now probably trying to figure out how to run your carefully crafted training on this compute. This overview will highlight my learning from working with TPUs and orchestrating training.
Once you have a TPU such as a v4-32 deployed in whatever region and zone, I would recommend deploying with the tpu-ubuntu2204-base software version, as it provides the most complete setup for whatever training needs. The next step is to SSH into your machine and check it's still there ;). The process I developed for working with spot instances which might be preempted at any point by the provider starts with a bootstrap run which turns the bare VM into a Ray cluster. This starts with per-heating all the workers in the topology:
gcloud compute tpus tpu-vm ssh v4-32-us-spot \
--zone ZONE --project PROJECT --worker=all \
--command='python3 -m pip install --user "jax[tpu]" \
-f https://storage.googleapis.com/jax-releases/libtpu_releases.html \
stable-baselines3 gymnasium wandb tensorboard "ray[default]"'
If you want actual isolation, create and activate a venv on each node first. At this point it would be good to run a test to ask each worker to report on its world:
gcloud compute tpus tpu-vm ssh v4-32-us-spot \
--zone ZONE --project PROJECT --worker=all \
--command='python3 -c "import jax; jax.distributed.initialize(); print(jax.process_index(), jax.device_count())"'
This will print out a report from each node on its position and chip count.
Creating the Head
Now we have setup ray on all the workers, next we establish the head which we can then connect to with:
gcloud compute tpus tpu-vm ssh v4-32-us-spot \
--zone ZONE --project PROJECT --worker=0 \
--command='export PATH=\(HOME/.local/bin:\)PATH && \
ray start --head --port=6379 --dashboard-host=0.0.0.0 \
--resources="{\"TPU\":8}" --disable-usage-stats'
The head will let us connect to the cluster and make the deployment of our code much easier (we will later create an SSH tunnel for the head portal). Now we need to configure each of the workers to connect to the head and listen for jobs:
for w in 1 2 3; do
gcloud compute tpus tpu-vm ssh v4-32-us-spot \
--zone ZONE --project PROJECT --worker=$w \
--command='export PATH=\(HOME/.local/bin:\)PATH && \
ray start --address=INTERNAL_IP:6379 \
--resources="{\"TPU\":8}" --disable-usage-stats'
done
This should be tailored to your worker count but it will use the internal IP of the main node which you can find on the GCP console or with a quick command. And with that we have aligned all the TPU workers to enable the full parallelization across the entire cluster. Great, but we still have no access from our local machine, to do that we need to create an ssh tunnel for the port 8265 or whichever you want to configure.
gcloud compute tpus tpu-vm ssh v4-32-us-spot \
--zone ZONE --project PROJECT --worker=0 \
-- -L 8265:localhost:8265 -L 10001:localhost:10001 -N
Now if you open http://localhost:8265 you will see the dashboard with your ray cluster with all the nodes we configured.
Submitting Jobs
Now we are finally at the fun part: actually sending training to the cluster.
The key idea is to submit jobs from your local machine to the Ray head through the tunnel on localhost:8265. Ray will package your working directory, install dependencies from a runtime env, and execute your entrypoint remotely.
I usually start with a runtime payload so every run is reproducible and self-contained:
{
"pip": [
"stable-baselines3>=2.2.0",
"gymnasium>=0.29.0",
"wandb",
"tensorboard",
"python-dotenv",
"pandas"
],
"env_vars": {
"WANDB_PROJECT": "PROJECT",
"WANDB_ENTITY": "YOUR_ENTITY",
"WANDB_API_KEY": os.getenv("WANDB_API_KEY", "")
}
}
With this in place, do a smoke test first:
ray job submit \
--address http://localhost:8265 \
--working-dir . \
--runtime-env-json "$RUNTIME_ENV_JSON" \
-- \
python -m engine.train --algo ppo --total-timesteps 5000 --alpha 0.3 --device cpu
If that succeeds, scale to your real training command.
For distributed execution (one trainer per host or per TPU allocation), use your distributed launcher which will make the distribution of parameters or allocation of configuration easier for us to handle:
ray job submit \
--address http://localhost:8265 \
--working-dir . \
--runtime-env-json "$RUNTIME_ENV_JSON" \
-- \
python scripts/ray_distributed_train.py \
--train-args "--algo ppo --total-timesteps 100000 --alpha 0.4 --device cpu" \
--num-nodes 4 \
--tpu-per-task 8 \
--base-seed 42
If you wrapped this in submit_ray_job.sh (which I recommend), the same flow is shorter and harder to mess up:
RAY_MODE=distributed \
NUM_NODES=4 \
TPU_PER_TASK=8 \
TRAIN_ARGS="--algo ppo --total-timesteps 100000 --alpha 0.4 --device cpu" \
bash ./submit_ray_job.sh
Once single runs look healthy, launch sweep agents so Ray fans out trials massively in parallel.
RAY_MODE=sweep \
SWEEP_KIND=ppo_rl_study \
SWEEP_METHOD=random \
SWEEP_RUN_CAP=900 \
NUM_NODES=4 \
AGENTS_PER_NODE=16 \
TPU_PER_TASK=0 \
WANDB_PROJECT=capstone_tpu \
WANDB_ENTITY=YOUR_ENTITY \
bash ./submit_ray_job.sh
A practical note: for SB3-based algorithms, most compute is CPU-side, so high agent parallelism is usually more useful than trying to reserve TPU for each actor. TPU matters most for JAX-heavy paths. The dashboard gives the visual view, but I still use CLI for reliability:
ray status
ray job list
ray job status <submission_id>
ray job logs -f <submission_id>
If something is off, these are the first things I check:
Are all expected nodes alive in
ray status? (Maybe got preempted)Are two Ray nodes accidentally on the same TPU host IP?
Did the job start with the expected env vars (
PHANTOM_JAX_PLATFORM, W&B settings)?Are metrics arriving in W&B at steady cadence? (Or did the run crash?)
Checkpoints and Artifact Uploads
For spot instances, assume interruption will happen. The minimum reliable setup is:
Save final model weights locally at end of training.
Upload that final checkpoint to W&B as an artifact.
Use deterministic aliases (
latest,step-<N>) so resuming logic stays simple.
A minimal upload pattern looks like this:
artifact = wandb.Artifact(name="phantom-ppo-final", type="checkpoint")
artifact.add_file("engine/models/phantom_ppo.zip", name="model.zip")
wandb.log_artifact(artifact, aliases=["latest", f"step-{num_timesteps}"])
I intentionally upload only final weights for sweep runs to avoid artifact noise and storage bloat (you can maybe also use the new HF Buckets). When a spot VM gets reclaimed, the recovery loop is straightforward:
Re-run the bootstrap install on all workers. You can put this into a watchdog script that runs either in a foreground shell on your computer or in a docker container.
Start Ray head + workers again.
Re-open the SSH tunnel to
8265.Re-launch sweep agents with the same
SWEEP_ID.
Because sweep state is on W&B, agents can continue pulling remaining runs after cluster recovery. When you are done, stop Ray everywhere so the next startup is clean:
gcloud compute tpus tpu-vm ssh v4-32-us-spot \
--zone ZONE --project PROJECT --worker=all \
--command='export PATH=\(HOME/.local/bin:\)PATH && ray stop -f'
That is basically the full lifecycle: bootstrap, form cluster, tunnel, submit, monitor, recover, and tear down. Once this loop is scripted, TPU experimentation becomes much less painful and much more reproducible.
Behavior Modeling Parallelization

One thing that gave me a real throughput jump was moving the behavioral-profile logic into JAX parallel batches instead of evaluating profiles in Python loops. In practice, I keep a profile bank for human and agent transition kernels (and any synthetic contamination profiles) as fixed-shape tensors, then run divergence scoring and robust alpha selection with jax.jit + jax.vmap so each step evaluates all profile candidates in one compiled pass. If you want to scale beyond a single host, the same computation can be sharded across workers with pmap/multi-process JAX, where each process gets a slice of the profile bank and we aggregate worst-case signals at the end; this is especially useful when you are screening many levels at once. The important engineering detail is to keep shapes static (pad trajectories, use masks, avoid Python-side branching) so recompilation does not kill performance, and to pre-place profile tensors on device once to avoid host-to-device transfer overhead on every step. I also strongly recommend deterministic PRNG handling with jax.random.fold_in(process_index, global_step) so profile comparisons remain reproducible across nodes, which makes sweep analysis and equilibrium claims much easier to trust in your results.
Conclusion
At the end of the day, running RL at TPU scale is less about a single launch command and more about building a resilient workflow: bootstrap all workers, form a clean Ray cluster, submit jobs through a reproducible runtime env, and treat spot preemption as a normal event instead of a failure. Once that loop is scripted, the infrastructure gets out of your way and you can focus on the actual science and using JAX-parallel logic to evaluate policies fast and consistently. That is when the cluster stops feeling like ops overhead or MLOps project and starts feeling like a research accelerator.
PS: Let your AI agents help you with configuration and testing of the clusters itself.



