Inferenza batch con Ray Data e vLLM

Importante

Questa funzionalità è in Anteprima Pubblica.

Questo esempio esegue inferenza batch offline di LLM con Ray Data e vLLM su 4 nodi A10. Uno script di bootstrap avvia un cluster Ray su tutti i nodi, quindi il driver utilizza l'API LLM di Ray Data (ray.data.llm) per avviare una replica di vLLM per nodo e far passare attraverso di esse un set di dati di prompt, scrivendo il testo generato in formato Parquet in un volume di Unity Catalog.

Usa un modello pubblico (Qwen2.5-7B-Instruct), quindi può essere eseguito così com'è senza bisogno di un token di Hugging Face.

Il carico di lavoro esegue le operazioni seguenti:

  • Carica il progetto locale con code_source: snapshot.
  • Avvia una testa Ray sul nodo 0, si unisce a 3 nodi worker e poi esegue il driver di inferenza batch.
  • Usa ray.data.llm per eseguire una replica di vLLM per nodo ed elaborare i prompt in parallelo.
  • Scrive i prompt e gli output generati in un volume di Unity Catalog in formato Parquet.

Prerequisiti

Layout del progetto

Creare una directory con i file seguenti.

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

Passaggio 1: Scrivere il carico di lavoro YAML

train.yaml richiede 4 GPU_1xA10 nodi. Le dipendenze vengono dichiarate direttamente in environment (con l’immagine client version), e command avvia un cluster Ray tra i nodi ed esegue quindi il driver, perciò il carico di lavoro non richiede un file delle dipendenze separato né uno script di avvio.

VLLM non si trova nell'immagine di base, quindi viene installato inline insieme a tre pin necessari per i nodi GPU: hf_transfer (l'immagine di base consente download rapidi di Hugging Face e prevede questo pacchetto), una versione più recente fsspec (l'immagine di base fornisce un vecchio che interrompe i download) e un vLLM pull in opencv-python-headless OpenCV, il cui volante predefinito arresta il self-test FIPS OpenSSL nei nodi GPU.

Imposta OUTPUT_PATH su un volume di Unity Catalog su cui è possibile scrivere. Impostato NUM_GPUS allo stesso valore di 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

L'inline command avvia una testata Ray con la GPU del nodo sul nodo 0, poi esegue il driver con python batch_inference.py. I nodi worker si collegano al nodo head usando MASTER_ADDR e NODE_RANK, che la piattaforma imposta automaticamente. Ogni lavoratore monitora la testa e interrompe i processi locali Ray dopo tre fallimenti consecutivi dei controlli sanitari.

Passaggio 2: Definire il driver di inferenza batch

batch_inference.py compila un set di dati Ray di richieste, configura un processore vLLM con ray.data.llme scrive i risultati. Il driver aspetta che tutti i nodi si uniscano prima di leggere il conteggio della GPU. AIR fornisce un pool di acceleratori fissi, così il driver imposta concurrency su una tupla fissa (minimum, maximum) che richiede una replica per GPU. Poiché questo esempio utilizza un carico di lavoro breve e fisso, il driver attende fino a 300 secondi affinché tutte le repliche vengano inizializzate prima di inviare il lavoro. Ogni attore elabora fino a due lotti contemporaneamente e ha al massimo due compiti Ray Data inviati, inclusi compiti in esecuzione e in coda. Questo impedisce al primo attore che inizializza di riservare la maggior parte del carico di lavoro. I 2.000 prompt sono suddivisi in 32 blocchi di input, con otto blocchi disponibili per replica. Per carichi di lavoro più lunghi, regola queste impostazioni in base ai tempi di avvio e ai requisiti di 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 trasforma ogni riga di input in una richiesta di chat e postprocess mantiene le colonne in modo permanente. Ray Data aggiunge una generated_text colonna con l'output del modello. Lo script completo è incluso nello script completo del driver alla fine di questa pagina.

tensor_parallel_size=1 conserva ogni replica vLLM su una GPU A10.

Passaggio 3: Invia il run

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

Passaggio 4: Controllare l'esecuzione

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

I log mostrano il prompt e la velocità effettiva di generazione del motore vLLM durante l'esecuzione del batch, quindi una Wrote <n> rows riga quando viene scritto l'output.

Dove atterrare i risultati

Il driver scrive un dataset Parquet nel volume OUTPUT_PATH, con una colonna instruction e una colonna output. Rileggilo con Spark o pandas, ad esempio spark.read.parquet(OUTPUT_PATH).

Script completo del driver

Il file completo batch_inference.py per la copia-incolla:

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

Passaggi successivi