Dávkové odvozování s Ray Data a vLLM

Important

Tato funkce je ve verzi Public Preview.

Tento příklad spouští offline dávkovou inferenci LLM s Ray Data a vLLM na 4 uzlech A10. Bootstrap skript spustí cluster Ray napříč uzly, poté driver pomocí LLM API Ray Data (ray.data.llm) spustí jednu repliku vLLM na každý uzel a průběžně přes ně zpracovává datovou sadu promptů, přičemž vygenerovaný text zapisuje do svazku Unity Catalog ve formátu Parquet.

Používá veřejný model (Qwen2.5-7B-Instruct), takže běží as-is bez tokenu Hugging Face.

Úloha provádí následující akce:

  • Nahraje místní projekt pomocí code_source: snapshot.
  • Spustí Ray hlavu na uzlu 0, připojí 3 pracovní uzly a pak spustí dávkový inferenční ovladač.
  • Používá ray.data.llm k provozování jedné repliky vLLM na uzel a k paralelnímu zpracování promptů.
  • Zapíše prompty a vygenerované výstupy do svazku v Unity Catalogu ve formátu Parquet.

Předpoklady

Rozložení projektu

Vytvořte adresář s následujícími soubory.

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

Krok 1: Zápis YAML úlohy

train.yaml požaduje 4 GPU_1xA10 uzly. Závislosti jsou deklarovány inline pod environment (s klientským obrazem version) a command spustí cluster Ray napříč uzly a poté driver, takže úloha nepotřebuje samostatný soubor se závislostmi ani spouštěcí skript.

vLLM není v základní imagi, takže se instaluje přímo spolu se třemi fixovanými verzemi balíčků, které uzly GPU vyžadují: hf_transfer (základní image umožňuje rychlé stahování z Hugging Face a očekává tento balíček), novější fsspec (základní image obsahuje starou verzi, která způsobuje selhávání stahování) a fixovaná verze opencv-python-headless (vLLM si jako závislost přitahuje OpenCV, jehož výchozí wheel balíček způsobuje pád samokontroly OpenSSL FIPS na uzlech GPU).

Nastavte OUTPUT_PATH na svazek v katalogu Unity Catalog, do kterého máte oprávnění zapisovat. Nastavte NUM_GPUS stejnou hodnotu jako 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

Vložený příkaz command spustí hlavní uzel Ray s GPU daného uzlu na uzlu 0 a poté spustí driver pomocí python batch_inference.py. Pracovní uzly se připojují k hlavnímu uzlu pomocí MASTER_ADDR a NODE_RANK, které platforma automaticky nastavuje. Každý pracovník monitoruje hlavu a po třech po sobě jdoucích neúspěšných zdravotních kontrolách zastaví své lokální Ray procesy.

Krok 2: Definujte ovladač pro dávkovou inferenci

batch_inference.py sestaví datovou sadu Ray z promptů, nakonfiguruje procesor vLLM pomocí ray.data.llm a zapíše výsledky. Ovladač čeká, až se všechny uzly připojí, než načte počet GPU. AIR poskytuje pevný fond akcelerátorů, takže ovladač nastaví concurrency na pevnou dvojici (minimum, maximum), která vyžaduje jednu repliku pro každé GPU. Protože tento příklad používá krátkou pevně danou pracovní zátěž, ovladač čeká až 300 sekund, než se všechny repliky inicializují, a teprve poté přidělí práci. Každý aktér zpracovává až dvě dávky současně a má maximálně dvě odevzdané úlohy Ray Data, včetně běžících a frontovaných úloh. To zabraňuje tomu, aby první aktér, který inicializuje, rezervoval většinu pracovní zátěže. 2 000 promptů je rozděleno do 32 vstupních bloků, přičemž na repliku je k dispozici osm bloků. Pro delší pracovní zátěže ladíte tato nastavení podle času startu a požadavků na propustnost:

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 změní každý vstupní řádek na žádost chatu a postprocess zachová sloupce tak, aby zůstaly zachovány. Ray Data přidá generated_text sloupec s výstupem modelu. Úplný skript je v úplném skriptu ovladače na konci této stránky.

tensor_parallel_size=1 každá replika vLLM zůstává na jednom A10 GPU.

Krok 3: Odeslání spuštění

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

Krok 4: Kontrola spuštění

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

V protokolech se během běhu dávky zobrazuje propustnost zpracování promptů a generování enginu vLLM a poté se při zápisu výstupu objeví řádek Wrote <n> rows.

Kam se výsledky uloží

Ovladač zapíše jednu datovou sadu Parquet do svazku OUTPUT_PATH se sloupcem instruction a sloupcem output . Znovu je načtěte pomocí Apache Spark nebo knihovny pandas, například spark.read.parquet(OUTPUT_PATH).

Úplný skript ovladače

Kompletní batch_inference.py pro kopírování a vložení:

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

Další kroky