Wnioskowanie wsadowe z użyciem Ray Data i vLLM

Ważna

Ta funkcja jest dostępna w publicznej wersji testowej.

Ten przykład uruchamia wsadowe wnioskowanie offline modeli LLM z użyciem Ray Data i vLLM na 4 węzłach A10. Skrypt rozruchowy uruchamia klaster Ray na wszystkich węzłach, a następnie program sterujący używa interfejsu API LLM Ray Data (ray.data.llm), aby uruchomić po jednej replice vLLM na każdym węźle i strumieniowo przetwarzać przez nie zbiór promptów, zapisując wygenerowany tekst do woluminu Unity Catalog w formacie Parquet.

Korzysta z publicznego modelu (Qwen2.5-7B-Instruct), więc działa od razu bez tokenu Hugging Face.

Obciążenie robocze realizuje następujące działania:

  • Przesyła lokalny projekt za pomocą polecenia code_source: snapshot.
  • Uruchamia głowę Ray na węźle 0, łączy 3 węzły robocze, a następnie uruchamia sterownik wnioskowania wsadowego.
  • Używa się ray.data.llm do uruchamiania jednej repliki vLLM na węzeł i przetwarzania promptów równolegle.
  • Zapisuje prompty i wygenerowane wyniki w woluminie Unity Catalog w formacie Parquet.

Wymagania wstępne

Układ projektu

Utwórz katalog z następującymi plikami.

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

Krok 1. Zapisywanie obciążenia YAML

train.yaml żąda 4 GPU_1xA10 węzłów. Zależności są deklarowane bezpośrednio pod environment (wraz z obrazem klienta version), a command uruchamia klaster Ray na wszystkich węzłach, a następnie uruchamia sterownik, więc zadanie nie wymaga osobnego pliku zależności ani skryptu uruchamiającego.

vLLM nie znajduje się w obrazie bazowym, więc jest instalowany bezpośrednio w poleceniu wraz z trzema przypiętymi wersjami pakietów wymaganymi przez węzły GPU: hf_transfer (obraz bazowy umożliwia szybkie pobieranie z Hugging Face i oczekuje tego pakietu), nowszą wersją fsspec (obraz bazowy zawiera starszą, która psuje pobieranie) oraz przypiętą wersją opencv-python-headless (vLLM instaluje OpenCV, którego domyślny wheel powoduje awarię autotestu FIPS biblioteki OpenSSL na węzłach GPU).

Ustaw OUTPUT_PATH na wolumin Unity Catalog, do którego można zapisywać. Ustaw NUM_GPUS tę samą wartość co 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

Polecenie command uruchamia węzeł główny Ray z użyciem procesora GPU węzła 0, a następnie uruchamia proces sterownika za pomocą python batch_inference.py. Węzły robocze dołączają do węzła głównego przy użyciu MASTER_ADDR i NODE_RANK, które platforma ustawia automatycznie. Każdy pracownik monitoruje głowicę i zatrzymuje lokalne procesy promieniowe po trzech kolejnych niepowodzeniach kontroli stanu zdrowia.

Krok 2: Definiowanie sterownika do wnioskowania wsadowego

batch_inference.py tworzy zbiór danych Ray z promptami, konfiguruje procesor vLLM przy użyciu ray.data.llm i zapisuje wyniki. Sterownik czeka, aż wszystkie węzły się dołączą, zanim odczyta liczbę GPU. AIR udostępnia stałą pulę akceleratorów, więc program sterujący ustawia concurrency na stałą krotkę (minimum, maximum), która określa jedną replikę na każdy procesor GPU. Ponieważ ten przykład wykorzystuje krótkie, stałe obciążenie, kierowca czeka do 300 sekund na inicjalizację wszystkich replik przed wysłaniem pracy. Każdy aktor przetwarza jednocześnie maksymalnie dwie partie i ma co najwyżej dwa przesłane zadania Ray Data, w tym zadania uruchomione i znajdujące się w kolejce. Zapobiega to temu, by pierwszy zainicjalizowany aktor zarezerwował sobie większość obciążenia. 2 000 promptów podzielono na 32 bloki wejściowe, przy czym dla każdej repliki dostępnych jest osiem bloków. Dla dłuższych obciążeń dostosuj te ustawienia do czasu uruchomienia i wymagań dotyczących przepustowości:

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 przekształca każdy wiersz danych wejściowych w żądanie czatu, a postprocess zachowuje kolumny, które mają zostać zapisane. Ray Data dodaje kolumnę generated_text z danymi wyjściowymi modelu. Pełny skrypt znajduje się w skrygcie pełnego sterownika na końcu tej strony.

tensor_parallel_size=1 każda replika vLLM jest przechowywana na jednej kartie graficznej A10.

Krok 3: Prześlij uruchomienie

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

Krok 4. Sprawdzanie przebiegu

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

Dzienniki pokazują przepustowość przetwarzania monitu i generowania silnika vLLM podczas wykonywania wsadu, a następnie wiersz Wrote <n> rows, gdy dane wyjściowe zostaną zapisane.

Gdzie wyniki lądują

Sterownik zapisuje jeden zestaw danych Parquet w woluminie OUTPUT_PATH z kolumną instruction i kolumną output . Odczytaj go z powrotem za pomocą platformy Spark lub biblioteki pandas, na przykład spark.read.parquet(OUTPUT_PATH).

Pełny skrypt sterownika

Kompletny batch_inference.py do kopiowania i wklejania:

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

Następne kroki