Batch-slutsatsdragning med Ray Data och vLLM

Important

Den här funktionen finns som allmänt tillgänglig förhandsversion.

Detta exempel kör offline LLM-batchinferens med Ray Data och vLLM över 4 A10-noder. Ett bootstrap-skript startar ett Ray-kluster över noderna, sedan använder drivrutinen Ray Datas LLM API (ray.data.llm) för att starta en vLLM-replika per nod och strömma en datamängd av prompts genom dem, och skriver den genererade texten till en Unity Catalog-volym som Parquet.

Den använder en publik modell (Qwen2.5-7B-Instruct), så den körs direkt utan någon Hugging Face-token.

Arbetsbelastningen gör följande:

  • Laddar upp det lokala projektet med code_source: snapshot.
  • Startar en Ray-huvudnod på nod 0, ansluter 3 arbetsnoder och kör sedan drivrutinen för batchinferens.
  • Används ray.data.llm för att köra en vLLM-replika per nod och processa promptar parallellt.
  • Skriver uppmaningarna och genererade utdata till en Unity Catalog-volym i Parquet-format.

Förutsättningar

  • air CLI har installerats och autentiserats. Se Installera AI Runtime CLI.
  • En Unity Catalog-volym som du kan skriva till. Du anger dess sökväg i arbetsbelastningen YAML nedan.

Projektlayout

Skapa en katalog med följande filer.

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

Steg 1: Skriv arbetsbelastningens YAML-fil

train.yaml begär 4 GPU_1xA10 noder. Beroenden deklareras inline under environment (med klientbilden version), och command startar ett Ray-kluster på alla noder och kör sedan drivrutinen, vilket innebär att arbetsbelastningen inte behöver en separat beroendefil eller ett startskript.

vLLM finns inte i basimagen, så det installeras inline tillsammans med tre versionslåsningar som GPU-noderna behöver: hf_transfer (basimagen möjliggör snabba nedladdningar från Hugging Face och förutsätter det här paketet), en nyare fsspec (basimagen levereras med en gammal version som gör att nedladdningar slutar fungera) och en versionslåst opencv-python-headless (vLLM drar in OpenCV, vars standard-wheel kraschar OpenSSL FIPS-självtestet på GPU-noderna).

Ange OUTPUT_PATH till en Unity Catalog-volym som du kan skriva till. Sätt NUM_GPUS till samma värde som 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

Inline-enheten command startar ett Ray-huvud med nodens GPU på nod 0, och kör sedan drivrutinen med python batch_inference.py. Arbetarnoder ansluter till huvudnoden med hjälp av MASTER_ADDR och NODE_RANK, som plattformen ställer in automatiskt. Varje arbetare övervakar huvudet och stoppar sina lokala Ray-processer efter tre på varandra följande hälsokontrollfel.

Steg 2: Definiera drivrutinen för batchinferens

batch_inference.py skapar en Ray Dataset med prompter, konfigurerar en vLLM-processor med ray.data.llmoch skriver resultatet. Drivrutinen väntar tills alla noder ansluter innan den läser GPU-räkningen. AIR tilldelar en fast acceleratorpool, så drivrutinen anger concurrency till en fast (minimum, maximum)-tupl som begär en replika per GPU. Eftersom detta exempel använder en kort, fast arbetsbelastning väntar drivrutinen upp till 300 sekunder för att alla repliker ska initiera innan arbetet kan skickas igång. Varje aktör bearbetar upp till två batchar samtidigt och har högst två inskickade Ray Data-uppgifter, inklusive körande och köade uppgifter. Detta förhindrar att den första aktören som initierar reserverar större delen av arbetsbelastningen. De 2 000 promptarna är uppdelade i 32 inmatningsblock, med åtta block tillgängliga per replika. För längre arbetsbelastningar, justera dessa inställningar baserat på starttid och genomströmningskrav:

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 omvandlar varje indatarad till en chattbegäran och postprocess håller kolumnerna kvar. Ray Data lägger till en generated_text kolumn med modellens utdata. Det fullständiga skriptet finns i skriptet Fullständig drivrutin i slutet av den här sidan.

tensor_parallel_size=1 Varje vLLM-replika finns på ett A10-grafikkort.

Steg 3: Skicka in körningen

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

Steg 4: Granska körningen

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

Loggarna visar vLLM-motorns prompt och genereringsgenomströmning under batchkörningen, följt av en rad med Wrote <n> rows när utdata skrivs.

Där resultat landar

Drivrutinen skriver en Parquet-datamängd till OUTPUT_PATH volymen, med en instruction kolumn och en output kolumn. Läs in det igen med Spark eller pandas, till exempel spark.read.parquet(OUTPUT_PATH).

Fullständigt drivrutinsskript

Den kompletta batch_inference.py för att kopiera och klistra in:

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

Nästa steg