Hinweis
Für den Zugriff auf diese Seite ist eine Autorisierung erforderlich. Sie können versuchen, sich anzumelden oder das Verzeichnis zu wechseln.
Für den Zugriff auf diese Seite ist eine Autorisierung erforderlich. Sie können versuchen, das Verzeichnis zu wechseln.
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: snapshothoch. - 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
airCLI 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()