Batch-Inferenz mit Ray Data und vLLM

Important

Dieses Feature befindet sich in der Public Preview.

Dieses Beispiel führt eine Offline-LLM-Batch-Inferenz mit Ray Data und vLLM über 8 A10-Knoten aus. Ein Bootstrap-Skript startet einen Ray-Cluster über die Knoten, dann verwendet der Treiber die LLM-API (ray.data.llm) von Ray Data, um pro Knoten eine vLLM-Replik zu starten und einen Datensatz von Prompts durch diese zu streamen, wobei der generierte Text als Parquet auf ein Unity-Katalog-Volume geschrieben wird.

Es verwendet ein öffentliches Modell (Qwen2.5-7B-Instruct), sodass es ohne weitere Anpassungen und ohne Hugging-Face-Token läuft.

Die Workload führt Folgendes aus:

  • Lädt das lokale Projekt mit code_source: snapshot hoch.
  • Startet einen Ray-Head auf Knoten 0, fügt 7 Worker-Knoten hinzu und führt anschließend den Batch-Inferenz-Treiber aus.
  • Verwendet ray.data.llm, um pro Knoten ein vLLM-Replikat auszuführen und Prompts parallel zu verarbeiten.
  • Schreibt die Prompts und generierten Ausgaben als Parquet in ein Unity Catalog-Volume.

Prerequisites

  • Die air CLI wurde installiert und authentifiziert. Siehe Installieren der AI-Runtime CLI.
  • Ein Unity-Katalogvolume, in das Sie schreiben können. Sie legen Ihren Pfad in der unten stehenden Workload-YAML-Datei fest.

Projektlayout

Erstellen Sie ein Verzeichnis mit den folgenden Dateien.

ray_batch_inference/
├── train.yaml            # air workload config (inline dependencies + Ray bootstrap)
└── batch_inference.py    # Ray Data + vLLM batch inference driver

Schritt 1: Schreiben Sie das Workload-YAML

train.yaml fordert 8 GPU_1xA10 Knoten an. Abhängigkeiten werden inline unter environment (zusammen mit dem Client-Image version) deklariert, und command startet einen Ray-Cluster über die Knoten hinweg und führt anschließend den Treiber aus, sodass die Workload keine separate Abhängigkeitsdatei oder kein separates Startskript benötigt.

vLLM ist nicht im Basisimage enthalten, daher wird es inline zusammen mit drei angehefteten Paketen installiert, die die GPU-Knoten benötigen: hf_transfer (das Basisimage ermöglicht schnelle Hugging Face-Downloads und erwartet dieses Paket), ein neueres fsspec (das Basisimage enthält ein altes Paket, das Downloads unterbricht) und ein angeheftetes opencv-python-headless (vLLM zieht OpenCV ein, dessen Standard-Wheel den OpenSSL FIPS-Selbsttest auf den GPU-Knoten zum Absturz bringt).

Legen Sie den Satz OUTPUT_PATH auf ein Unity-Katalogvolume fest, in das Sie schreiben können. Setze NUM_GPUS auf denselben Wert wie 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

# 8 A10 nodes, one GPU each. Ray Data runs one vLLM replica per node.
compute:
  num_accelerators: 8
  accelerator_type: GPU_1xA10

code_source:
  type: snapshot
  snapshot:
    root_path: .

command: |
  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
    python batch_inference.py
    ray stop
  else
    echo "NODE_RANK=$NODE_RANK: connecting to Ray head at $MASTER_ADDR:$RAY_HEAD_PORT..."
    for i in $(seq 1 12); do
      if ray start --address="$MASTER_ADDR:$RAY_HEAD_PORT" --num-gpus="$GPUS_PER_NODE" --block 2>/dev/null; then
        break
      fi
      echo "Attempt $i failed, retrying in 5s..."
      sleep 5
    done
  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: '8' # must match num_accelerators

Das inline-Element command startet einen Ray-Head mit der GPU des Knotens auf Knoten 0, führt den Treiber mit python batch_inference.py aus und hält anschließend den Cluster an. Worker-Knoten treten dem Head über MASTER_ADDR und NODE_RANK bei, die von der Plattform automatisch festgelegt werden.

Schritt 2: Definieren des Batch-Ableitungstreibers

batch_inference.py erstellt ein Ray-Dataset von Eingabeaufforderungen, konfiguriert einen vLLM-Prozessor mit ray.data.llmund schreibt die Ergebnisse. concurrency ist die Anzahl der vLLM-Replikate, die Ray Data parallel ausführt. Der Treiber wartet, bis alle Knoten beitreten, bevor er die GPU-Zählung ausliest, daher wird jeder Knoten verwendet:

import os
import time

import ray
from ray.data.llm import build_processor, vLLMEngineProcessorConfig

ray.init(address="auto")
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}.")

config = vLLMEngineProcessorConfig(
    model_source="Qwen/Qwen2.5-7B-Instruct",
    engine_kwargs={"max_model_len": 4096, "tensor_parallel_size": 1},
    concurrency=total_gpus,   # one vLLM replica per GPU in the cluster
    batch_size=64,
)

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 Wandelt jede Eingabezeile in eine Chatanfrage um, und postprocess die Spalten bleiben erhalten. Ray Data fügt eine generated_text Spalte mit der Ausgabe des Modells hinzu. Das vollständige Skript befindet sich am Ende dieser Seite im vollständigen Treiberskript .

Setzen Sie für größere Modelle tensor_parallel_size so, dass ein Replikat auf mehrere GPUs verteilt wird, und teilen Sie total_gpus durch diesen Wert, damit die Replikate den Cluster weiterhin vollständig ausfüllen, zum Beispiel concurrency=total_gpus // 2 mit tensor_parallel_size=2.

Schritt 3: Ausführung senden

air run -f train.yaml --dry-run
air run -f train.yaml --watch

Schritt 4: Ausführung überprüfen

air get run <run-id>
air logs <run-id>

Die Protokolle zeigen den Prompt- und Generierungsdurchsatz der vLLM-Engine, während der Batch läuft, und dann eine Zeile mit Wrote <n> rows, sobald die Ausgabe geschrieben wird.

Wo die Ergebnisse landen

Der Treiber schreibt ein Parquet-Dataset in das OUTPUT_PATH-Volume, mit einer instruction-Spalte und einer output-Spalte. Lesen Sie es beispielsweise mit Spark oder pandas wieder ein spark.read.parquet(OUTPUT_PATH).

Vollständiges Treiberskript

Der komplette batch_inference.py zum Kopieren und Einfügen:

#!/usr/bin/env python3
"""Offline batch inference with Ray Data + vLLM across 8 A10 nodes.

The workload `command` starts a Ray head on node 0 and joins 7 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.llm import build_processor, vLLMEngineProcessorConfig

MODEL_SOURCE = "Qwen/Qwen2.5-7B-Instruct"
NUM_PROMPTS = 1000
# 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")
    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()

    # vLLM engine config. concurrency = number of replicas Ray Data runs in parallel;
    # one per GPU in the cluster here. engine_kwargs are passed through to the vLLM engine.
    config = vLLMEngineProcessorConfig(
        model_source=MODEL_SOURCE,
        engine_kwargs={
            "max_model_len": 4096,
            "tensor_parallel_size": 1,
            "enable_chunked_prefill": True,
        },
        concurrency=total_gpus,
        batch_size=64,
    )

    # 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()

Nächste Schritte