Inférence par lots avec Ray Data et vLLM

Important

Cette fonctionnalité est disponible en préversion publique.

Cet exemple exécute une inférence par lots hors ligne de LLM avec Ray Data et vLLM sur 4 nœuds A10. Un script d’amorçage démarre un cluster Ray sur l’ensemble des nœuds, puis le processus pilote utilise l’API LLM de Ray Data (ray.data.llm) pour lancer une réplique vLLM par nœud et faire transiter en flux un jeu de données de prompts via celles-ci, puis écrit 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 as-is sans jeton Hugging Face.

La charge de travail effectue les opérations suivantes :

  • Charge le projet local avec code_source: snapshot.
  • Démarre un nœud principal Ray sur le nœud 0, y connecte 3 nœuds de travail, puis exécute le programme pilote d’inférence par lots.
  • Utilise ray.data.llm pour exécuter une réplique vLLM par nœud et traiter les invites en parallèle.
  • Écrit les prompts et les résultats générés dans un volume Unity Catalog au format Parquet.

Prerequisites

  • L’interface air CLI a été installée et authentifiée. Consultez Installer l’interface CLI d’AI Runtime.
  • Un volume de catalogue Unity dans lequel vous pouvez écrire. Vous définissez son chemin dans la charge de travail YAML ci-dessous.

Disposition du projet

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 : Écrire la charge de travail YAML

train.yaml demande 4 GPU_1xA10 nœuds. Les dépendances sont déclarées en ligne sous environment (avec l’image versionclient ), puis command lance un cluster Ray à travers les nœuds puis exécute le pilote, donc la charge de travail n’a pas besoin d’un fichier de dépendance séparé ou d’un script de lanceur.

vLLM n’est pas inclus dans l’image de base ; il est donc installé en ligne, avec trois dépendances épinglées nécessaires aux nœuds GPU : hf_transfer (l’image de base permet des téléchargements rapides depuis Hugging Face et attend ce paquet), une version plus récente de fsspec (l’image de base inclut une version ancienne qui empêche les téléchargements) et une version épinglée de opencv-python-headless (vLLM dépend d’OpenCV, dont le wheel par défaut provoque l’échec du test FIPS OpenSSL sur les nœuds GPU).

Définissez OUTPUT_PATH sur un volume Unity Catalog sur lequel vous disposez d’un accès en écriture. Fixer NUM_GPUS à 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 une infrastructure Ray avec le GPU du nœud sur le nœud 0, puis exécute le pilote avec python batch_inference.py. Les nœuds de calcul rejoignent le nœud principal à l’aide de MASTER_ADDR et de NODE_RANK, que la plateforme définit automatiquement. Chaque travailleur surveille la tête et arrête ses processus Ray locaux après trois échecs consécutifs de contrôle de santé.

Étape 2 : Définir le pilote d’inférence par lots

batch_inference.py crée un jeu de données Ray de prompts, configure un processeur vLLM avec ray.data.llm et écrit les résultats. Le pilote attend que tous les nœuds se joignent avant de lire le nombre de GPU. AIR met à disposition un pool fixe d’accélérateurs, de sorte que le pilote définit concurrency sur le tuple fixe (minimum, maximum) qui demande une réplique par GPU. Comme cet exemple utilise une charge de travail courte et fixe, le pilote attend jusqu’à 300 secondes que toutes les répliques s’initialisent avant de répartir le travail. Chaque acteur traite jusqu’à deux lots simultanément et dispose au maximum de deux tâches Ray Data soumises, incluant des tâches en cours d’exécution et des tâches en file d’attente. Cela empêche le premier acteur à initialiser de réserver la majeure partie de la charge de travail. Les 2 000 requêtes sont réparties en 32 blocs d’entrée, à raison de huit blocs disponibles par réplica. Pour des charges de travail plus longues, ajustez ces réglages en fonction du temps de démarrage et des exigences de débit :

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 requête de chat et postprocess conserve les colonnes. Ray Data ajoute une generated_text colonne avec la sortie du modèle. Le script complet se trouve dans le script de pilote complet à la fin de cette page.

tensor_parallel_size=1 il conserve chaque réplique vLLM sur un GPU A10.

Étape 3 : Envoyer 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 journaux affichent la requête et le débit de génération du moteur vLLM pendant l’exécution du lot, puis une ligne Wrote <n> rows quand la sortie est enregistrée.

Où s’affichent les résultats

Le pilote écrit un jeu de données Parquet sur le volume OUTPUT_PATH, avec une colonne instruction et une colonne output. Vous pouvez les relire avec Spark ou pandas, par exemple spark.read.parquet(OUTPUT_PATH).

Script de pilote complet

L’intégralité batch_inference.py à 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()

Étapes suivantes