Running vLLM on TPU with Kinetic#
This guide explains how to run vLLM (versatile Large Language Model) serving and inference on Cloud TPUs using the Kinetic framework.
Overview#
Kinetic allows you to easily offload heavy vLLM workloads to Cloud TPUs. This is particularly useful for serving large models like Llama 3.1 that require significant compute and memory.
Prerequisites#
Kinetic Cluster: You need a provisioned Kinetic cluster with TPU nodes (e.g.,
v5litepod).Hugging Face Token: If you are using gated models like Llama 3.1, you need a Hugging Face token (
HF_TOKEN) with access to the model.
Configuration#
To run vLLM successfully on TPU via Kinetic, you need to handle dependencies and environment variables properly.
1. Dependencies#
You need to ensure vllm-tpu is installed in the remote container. You can do this by creating a requirements.txt file in the directory of your script containing:
vllm-tpu
Kinetic will detect this file and build a container with vLLM installed.
2. Environment Variables#
You must pass the following environment variables to ensure correct execution:
VLLM_TARGET_DEVICE="tpu": Tells vLLM to target TPU.VLLM_USE_V1="0": Forces vLLM to use the stable v0 engine (recommended for TPU currently).JAX_PLATFORMS="tpu,cpu": Allows JAX to see both TPU and CPU backends, avoiding initialization crashes.
You can use capture_env_vars in the @kinetic.run decorator to pass these from your local environment.
Example#
# Installation
# before you begin please install VLLM in your remote environment
# pip install vllm-tpu
import os
import kinetic
# We use a TPU accelerator. Adjust as needed (e.g., 'tpu-v5e-1', 'tpu-v5litepod-4')
@kinetic.run(
accelerator="tpu-v5litepod",
# Capture Hugging Face token if using a gated model like Gemma
capture_env_vars=["HF_TOKEN"],
)
def run_vllm_inference():
# Imports must happen inside the decorated function because they need to run
# in the remote container where vllm is installed.
os.environ["VLLM_TARGET_DEVICE"] = "tpu"
os.environ["JAX_PLATFORMS"] = "tpu,cpu"
os.environ["VLLM_USE_V1"] = "0"
from vllm import LLM, SamplingParams
model_id = "meta-llama/Llama-3.1-8B"
print(f"Initializing vLLM with model: {model_id}")
# We use arguments matching the quickstart
llm = LLM(model=model_id, tensor_parallel_size=4, max_model_len=2048)
sampling_params = SamplingParams(temperature=0.8, top_p=0.95)
prompts = [
"Hello, my name is",
"The capital of France is",
"The president of the United States is",
]
print("Generating completions...")
outputs = llm.generate(prompts, sampling_params)
# Print the results
for output in outputs:
prompt = output.prompt
generated_text = output.outputs[0].text
print(f"\nPrompt: {prompt!r}")
print(f"Generated text: {generated_text!r}")
if __name__ == "__main__":
run_vllm_inference()
Running the Example#
To run the script, set the required environment variables locally so they get captured:
HF_TOKEN=your-hf-token \
python3 your_script.py