Inférence par batch avec Ray Data et vLLM
Aperçu
Cette fonctionnalité est en aperçu public.
Cet exemple exécute une inférence LLM par batch hors ligne avec
Ray Data et
vLLM sur 4 nœuds A10. Un script de bootstrap starts un
cluster Ray sur les nœuds, puis le Driver utilise l’API LLM de Ray Data (ray.data.llm) pour
lancer une réplique vLLM par nœud et Stream un dataset de prompts à travers eux, en écrivant le
texte généré dans un volume Unity Catalog au format Parquet.
Il utilise un modèle public (Qwen2.5-7B-Instruct), de sorte qu’il s’exécute tel quel sans jeton Hugging Face.
La charge de travail réalise les étapes suivantes :
- Importe le projet local avec
code_source: snapshot. - Démarre un head Ray sur le nœud 0, joint 3 nœuds worker, puis exécute le driver d’inférence de batch.
- Utilise
ray.data.llmpour exécuter une réplique vLLM par nœud et traiter les invites en parallèle. - Écrit les invites et les sorties générées vers un volume Unity Catalog au format Parquet.
Prérequis
- Le CLI
airest installé et authentifié. Voir Installer le CLI AI Runtime. - Un volume Unity Catalog dans lequel vous pouvez écrire. Vous définissez son chemin dans le fichier YAML de la charge de travail ci-dessous.
Project Layout
Créez un répertoire avec les fichiers suivants.
ray_batch_inference/
├── train.yaml # air workload config (inline dependencies + Ray bootstrap)
└── batch_inference.py # Ray Data + vLLM batch inference driver
Étape 1 : Rédiger le YAML de la charge de travail
train.yaml demande 4 nœuds GPU_1xA10. Les dépendances sont déclarées en ligne sous
environment (avec l’image client version), et le command start un cluster Ray
sur les nœuds puis exécute le Driver, de sorte que la charge de travail n’a pas besoin d’un fichier de dépendances ou d’un script de lancement distinct.
vLLM n'est pas dans l'image de base, il est donc installé en ligne avec trois pins dont les nœuds GPU
ont besoin : hf_transfer (l'image de base permet des downloads rapides depuis Hugging Face et s'attend à ce que ce
package soit présent), un fsspec plus récent (l'image de base contient une ancienne version qui interrompt les downloads), et un opencv-python-headless
pins (vLLM intègre OpenCV, dont la roue par default fait planter l'autotest FIPS d'OpenSSL
sur les nœuds GPU).
Définissez OUTPUT_PATH sur un volume Unity Catalog sur lequel vous pouvez écrire. Définissez NUM_GPUS sur la même valeur que num_accelerators.
experiment_name: air-ray-batch-inference
environment:
version: '5'
dependencies:
- ray[data]==2.56.1
- vllm
- datasets>=3.0
- huggingface_hub>=0.34
# The base image sets HF_HUB_ENABLE_HF_TRANSFER=1; install the package it expects
# so model and dataset downloads don't error out.
- hf_transfer
# The base image ships fsspec 2023.5.0, which is too old for modern
# huggingface_hub and breaks dataset/model downloads. Pin a newer fsspec.
- fsspec>=2024.6.1
# vLLM pulls in opencv; its default wheel crashes the OpenSSL FIPS self-test
# on the GPU nodes. This pinned headless build avoids the crash.
- opencv-python-headless==4.12.0.88
# 4 A10 nodes, one GPU each. Ray Data runs one vLLM replica per node.
compute:
num_accelerators: 4
accelerator_type: GPU_1xA10
code_source:
type: snapshot
snapshot:
root_path: .
command: |
set -e
cd $CODE_SOURCE_PATH
RAY_HEAD_PORT=6379
GPUS_PER_NODE=${LOCAL_WORLD_SIZE:-1}
if [ "${NODE_RANK:-0}" = "0" ]; then
echo "NODE_RANK=0: starting Ray head with $GPUS_PER_NODE GPU(s)..."
ray start --head --port=$RAY_HEAD_PORT --num-gpus="$GPUS_PER_NODE" --dashboard-host=0.0.0.0
trap 'ray stop || true' EXIT
python batch_inference.py
else
echo "NODE_RANK=$NODE_RANK: connecting to Ray head at $MASTER_ADDR:$RAY_HEAD_PORT..."
joined=""
for i in $(seq 1 12); do
if ray start --address="$MASTER_ADDR:$RAY_HEAD_PORT" --num-gpus="$GPUS_PER_NODE" 2>/dev/null; then
joined=1
break
fi
echo "Attempt $i failed, retrying in 5s..."
sleep 5
done
if [ -z "$joined" ]; then
echo "Worker failed to join the Ray head after all retries." >&2
exit 1
fi
echo "Worker joined. Waiting for the head to finish..."
consecutive_failures=0
for _ in $(seq 1 720); do
if timeout 5 ray health-check --address "$MASTER_ADDR:$RAY_HEAD_PORT" 2>/dev/null; then
consecutive_failures=0
else
consecutive_failures=$((consecutive_failures + 1))
if [ "$consecutive_failures" -ge 3 ]; then
echo "Head is no longer healthy. Stopping local Ray processes..."
ray stop || true
exit 0
fi
echo "Head health check failed ($consecutive_failures/3). Retrying..."
fi
sleep 5
done
echo "Timed out waiting for the Ray head to finish." >&2
ray stop || true
exit 1
fi
max_retries: 0
timeout_minutes: 60
env_variables:
NCCL_SOCKET_IFNAME: eth0
# Unity Catalog volume where results land as Parquet. Replace with your volume.
OUTPUT_PATH: /Volumes/main/default/air_examples/ray_batch_inference
NUM_GPUS: '4' # must match num_accelerators
Le command en ligne démarre un head Ray avec le GPU du nœud sur le nœud 0, puis exécute le driver
avec python batch_inference.py. Les nœuds worker rejoignent le head à l’aide de MASTER_ADDR et
NODE_RANK, que la plateforme définit automatiquement. Chaque worker surveille le head et arrête
ses processus Ray locaux après trois échecs consécutifs de vérification de l’état.
Étape 2 : Définir le driver d'inférence par batch
batch_inference.py crée un Ray Dataset de prompts, configure un processeur vLLM avec
ray.data.llm et écrit les résultats. Le driver attend que tous les nœuds rejoignent le cluster avant
de lire le nombre de GPU. AIR provisionne un pool d’accélérateurs fixe, de sorte que le driver définit
concurrency sur un tuple (minimum, maximum) fixe qui demande une réplique par GPU.
Comme cet exemple utilise une charge de travail courte et fixe, le driver attend jusqu’à 300 secondes que
toutes les répliques s’initialisent avant de répartir le travail. Chaque acteur traite jusqu’à deux batchs
simultanément et dispose d’au plus deux tâches Ray Data soumises, incluant les tâches
en cours d’exécution et en file d’attente. Cela empêche le premier acteur à s’initialiser de réserver la majeure partie de la charge de travail. Les
2 000 prompts sont divisés en 32 blocs d’entrée, avec huit blocs disponibles par réplique. Pour les
charges de travail plus longues, ajustez ces paramètres en fonction du temps de startup et des exigences de throughput :
import os
import time
import ray
from ray.data import DataContext
from ray.data.llm import build_processor, vLLMEngineProcessorConfig
ray.init(address="auto")
data_context = DataContext.get_current()
data_context.wait_for_min_actors_s = 300
num_gpus = int(os.environ["NUM_GPUS"])
for _ in range(60):
if int(ray.cluster_resources().get("GPU", 0)) >= num_gpus:
break
time.sleep(5)
total_gpus = int(ray.cluster_resources().get("GPU", 0))
if total_gpus < num_gpus:
raise SystemExit(f"Expected {num_gpus} GPU(s) but Ray only sees {total_gpus}.")
ds = build_prompts().repartition(total_gpus * 8)
config = vLLMEngineProcessorConfig(
model_source="Qwen/Qwen2.5-7B-Instruct",
engine_kwargs={"max_model_len": 4096, "tensor_parallel_size": 1},
concurrency=(total_gpus, total_gpus),
batch_size=64,
max_concurrent_batches=2,
max_tasks_in_flight_per_actor=2,
)
processor = build_processor(
config,
preprocess=lambda row: dict(
messages=[{"role": "user", "content": row["instruction"]}],
sampling_params=dict(max_tokens=256, temperature=0.7),
),
postprocess=lambda row: dict(instruction=row["instruction"], output=row["generated_text"]),
)
out = processor(ds) # ds is a Ray Dataset with an "instruction" column
out.write_parquet(OUTPUT_PATH)
preprocess transforme chaque ligne d'entrée en une requête de chat, et postprocess conserve les colonnes
à persister. Ray Data ajoute une colonne generated_text avec la sortie du modèle. Le script
complet se trouve dans Script complet du driver à la fin de cette page.
tensor_parallel_size=1 maintient chaque réplique vLLM sur un GPU A10.
Étape 3 : Soumettre l'exécution
air run -f train.yaml --dry-run
air run -f train.yaml --watch
Étape 4 : Inspecter l'exécution
air get run <run-id>
air logs <run-id>
Les Logs affichent le prompt et le throughput de génération du moteur vLLM pendant l'exécution du batch, puis une ligne
Wrote <n> rows lorsque la sortie est écrite.
Où les résultats aboutissent
Le Driver écrit un dataset Parquet dans le volume OUTPUT_PATH, avec une colonne instruction
et une colonne output. Relisez-le avec Spark ou pandas, par exemple
spark.read.parquet(OUTPUT_PATH).
Script de driver complet
Le batch_inference.py complet pour le copier-coller :
#!/usr/bin/env python3
"""Offline batch inference with Ray Data + vLLM across 4 A10 nodes.
The workload `command` starts a Ray head on node 0 and joins 3 worker nodes, each
contributing 1 GPU. Ray Data's LLM API (`ray.data.llm`) launches one vLLM replica
per GPU and streams a dataset of prompts through them, then writes the generated text
to a Unity Catalog volume as Parquet.
Uses a public model (no Hugging Face token required) so the example runs as-is.
"""
import os
import time
import ray
from datasets import load_dataset
from ray.data import DataContext
from ray.data.llm import build_processor, vLLMEngineProcessorConfig
MODEL_SOURCE = "Qwen/Qwen2.5-7B-Instruct"
NUM_PROMPTS = 2000
BATCH_SIZE = 64
BLOCKS_PER_REPLICA = 8
# Unity Catalog volume path where results land as Parquet. Set this in train.yaml.
OUTPUT_PATH = os.environ.get("OUTPUT_PATH", "/Volumes/main/default/air_examples/ray_batch_inference")
def build_prompts():
"""Build a Ray Dataset of prompts from a public instruction dataset."""
raw = load_dataset("tatsu-lab/alpaca", split=f"train[:{NUM_PROMPTS}]")
items = []
for row in raw:
instruction = row["instruction"]
if row.get("input"):
instruction = f"{instruction}\n\n{row['input']}"
items.append({"instruction": instruction})
return ray.data.from_items(items)
def main():
ray.init(address="auto")
data_context = DataContext.get_current()
data_context.wait_for_min_actors_s = 300
num_gpus = int(os.environ["NUM_GPUS"])
for _ in range(60):
if int(ray.cluster_resources().get("GPU", 0)) >= num_gpus:
break
time.sleep(5)
total_gpus = int(ray.cluster_resources().get("GPU", 0))
if total_gpus < num_gpus:
raise SystemExit(
f"Expected {num_gpus} GPU(s) but Ray only sees {total_gpus}; "
"check GPU discovery / node join on all nodes."
)
print(f"Ray cluster ready: {total_gpus} GPU(s)", flush=True)
ds = build_prompts().repartition(total_gpus * BLOCKS_PER_REPLICA)
# AIR provisions a fixed accelerator pool. Bound prefetching so the first ready
# actor cannot reserve the small workload before the other actors initialize.
config = vLLMEngineProcessorConfig(
model_source=MODEL_SOURCE,
engine_kwargs={
"max_model_len": 4096,
"tensor_parallel_size": 1,
"enable_chunked_prefill": True,
},
concurrency=(total_gpus, total_gpus),
batch_size=BATCH_SIZE,
max_concurrent_batches=2,
max_tasks_in_flight_per_actor=2,
)
# preprocess maps each input row to a chat request; postprocess keeps the columns
# we want to persist. ray.data.llm adds a `generated_text` column.
processor = build_processor(
config,
preprocess=lambda row: dict(
messages=[
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": row["instruction"]},
],
sampling_params=dict(max_tokens=256, temperature=0.7),
),
postprocess=lambda row: dict(
instruction=row["instruction"],
output=row["generated_text"],
),
)
# materialize once so the write and the sample print don't re-run inference.
out = processor(ds).materialize()
out.write_parquet(OUTPUT_PATH)
print(f"Wrote {out.count()} rows to {OUTPUT_PATH}", flush=True)
for row in out.take(2):
print("INSTRUCTION:", row["instruction"][:120], flush=True)
print("OUTPUT:", row["output"][:200], flush=True)
ray.shutdown()
if __name__ == "__main__":
main()