catglossop/cast-vla

Python

4

1,332 commits

updated Apr 28, 2026

See the code

README

Quickstart for CAST finetuning

Note that the script to run on tpus will automatically push your code. If you want to isolate your development please make a new branch. Read the launch script run_tpu.sh carefully before launching.

Installation

Clone the repo:

git clone https://github.com/catglossop/cast-vla.git --recursive

Setup your environment

cd cast-vla
uv venv
uv sync
source .venv/bin/activate

To launch on a single tpu vm (v4-8)

bash run_tpu.sh <name of tpu> <initialize (true for first job on new tpu)> <update (true if code is changed)> <wandb api key> <config-name>

To ssh into a single tpu vm (v4-8)

gcloud alpha compute tpus tpu-vm ssh <name of tpu> --zone=us-central2-b

To launch on a pod

bash run_tpu_pod.sh <name of pod> <initialize (true for first job on new tpu)> <update (true if code is changed)> <wandb api key> <config-name>

To ssh into a pod

bash ssh_pod.sh <name of pod>

Note that on initialization of a new pod or tpu vm, you will need to login to hugging face (to be fixed). To do this, ssh into the single vm or pod,

gcloud compute tpus tpu-vm ssh <name-of-tpu-vm> --zone=<your-region>
source ~/cast-vla/.venv/bin/activate
huggingface-cli login

and input your token.

To run inference,

export CUDA_VISIBLE_DEVICES=0
cd cast-vla
python scripts/inference_server.py --platform <gpu or tpu> --checkpoint_dir <your/path/to/checkpoint> --checkpoint_step <0> --prompt <the prompt to the model>

NOTE: Please make sure you have installed jax[cuda]==0.4.34 if you'd like to use a GPU for inference. Note that the model will use about 18 GB of memory and experiments were performed with a single GPU.

For example,

python scripts/inference_server.py --platform tpu --checkpoint_dir ~/cast_checkpoint --checkpoint_step 0 --prompt "Move along the wall"

PaliVLA

This is a framework for training multimodal vision-language-action (VLA) model for robotics in JAX. It primarily supports PaliGemma for now, though more base models will be added in the future.

Training

To train a model, run:

python scripts/train.py --config <your config name>

For example,

python scripts/train.py --config configs/cast_config.py

This repository is (for now) a fork of big_vision.

Citation

If you use PaliVLA in your own project, please cite this repository:

@misc{palivla,
  author       = {Kyle Stachowicz},
  title        = {PaliVLA},
  year         = {2024},
  url          = {https://github.com/kylestach/bigvision-palivla},
  note         = {GitHub repository}
}

catglossop/cast-vla

Python

4

1,332 commits

updated Apr 28, 2026

See the code

README

Quickstart for CAST finetuning

Note that the script to run on tpus will automatically push your code. If you want to isolate your development please make a new branch. Read the launch script run_tpu.sh carefully before launching.

Installation

Clone the repo:

git clone https://github.com/catglossop/cast-vla.git --recursive

Setup your environment

cd cast-vla
uv venv
uv sync
source .venv/bin/activate

To launch on a single tpu vm (v4-8)

bash run_tpu.sh <name of tpu> <initialize (true for first job on new tpu)> <update (true if code is changed)> <wandb api key> <config-name>

To ssh into a single tpu vm (v4-8)

gcloud alpha compute tpus tpu-vm ssh <name of tpu> --zone=us-central2-b

To launch on a pod

bash run_tpu_pod.sh <name of pod> <initialize (true for first job on new tpu)> <update (true if code is changed)> <wandb api key> <config-name>

To ssh into a pod

bash ssh_pod.sh <name of pod>

Note that on initialization of a new pod or tpu vm, you will need to login to hugging face (to be fixed). To do this, ssh into the single vm or pod,

gcloud compute tpus tpu-vm ssh <name-of-tpu-vm> --zone=<your-region>
source ~/cast-vla/.venv/bin/activate
huggingface-cli login

and input your token.

To run inference,

export CUDA_VISIBLE_DEVICES=0
cd cast-vla
python scripts/inference_server.py --platform <gpu or tpu> --checkpoint_dir <your/path/to/checkpoint> --checkpoint_step <0> --prompt <the prompt to the model>

NOTE: Please make sure you have installed jax[cuda]==0.4.34 if you'd like to use a GPU for inference. Note that the model will use about 18 GB of memory and experiments were performed with a single GPU.

For example,

python scripts/inference_server.py --platform tpu --checkpoint_dir ~/cast_checkpoint --checkpoint_step 0 --prompt "Move along the wall"

PaliVLA

This is a framework for training multimodal vision-language-action (VLA) model for robotics in JAX. It primarily supports PaliGemma for now, though more base models will be added in the future.

Training

To train a model, run:

python scripts/train.py --config <your config name>

For example,

python scripts/train.py --config configs/cast_config.py

This repository is (for now) a fork of big_vision.

Citation

If you use PaliVLA in your own project, please cite this repository:

@misc{palivla,
  author       = {Kyle Stachowicz},
  title        = {PaliVLA},
  year         = {2024},
  url          = {https://github.com/kylestach/bigvision-palivla},
  note         = {GitHub repository}
}

Languages

Python

96.4%

TypeScript

2.5%