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.
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"
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.
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.
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}
}
Python
96.4%
TypeScript
2.5%
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.
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"
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.
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.
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}
}
Python
96.4%
TypeScript
2.5%