Inferência em lote com Dados do Ray e vLLM

Importante

Esse recurso está em Visualização Pública.

Este exemplo executa inferência em lote offline de LLM com Ray Data e vLLM em 4 nós A10. Um script de inicialização inicia um cluster do Ray nos nós e, em seguida, o driver usa a API LLM do Ray Data (ray.data.llm) para iniciar uma réplica de vLLM por nó e transmitir um conjunto de dados de prompts por meio deles, gravando o texto gerado em um volume do Unity Catalog como Parquet.

Ele usa um modelo público (Qwen2.5-7B-Instruct), então funciona como está, sem um token do Hugging Face.

A carga de trabalho faz o seguinte:

  • Faz upload de um snapshot do projeto local com code_source.
  • Inicia um nó principal do Ray no nó 0, adiciona 3 nós de trabalho e depois executa o driver de inferência em lote.
  • Usa ray.data.llm para executar uma réplica do vLLM por nó e processar prompts em paralelo.
  • Grava os prompts e as saídas geradas em um volume do Unity Catalog como Parquet.

Pré-requisitos

Layout do projeto

Crie um diretório com os arquivos a seguir.

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

Etapa 1: Escreva o YAML da carga de trabalho

train.yaml solicita 4 GPU_1xA10 nós. As dependências são declaradas diretamente em environment (com a imagem do cliente version), e o command inicia um cluster Ray em todos os nós e, em seguida, executa o programa driver, de modo que a carga de trabalho não precisa de um arquivo de dependências separado nem de um script de inicialização.

O vLLM não está na imagem base, logo, ele é instalado em linha com três marcadores de que os nós de GPU precisam: hf_transfer (a imagem base habilita downloads do Hugging Face rápidos e espera esse pacote), um fsspec mais recente (a imagem base acompanha uma anterior que quebra os downloads) e um opencv-python-headless marcado (vLLM abre em OpenCV, cujo wheel padrão falha no autoteste FIPS OpenSSL nos nós de GPU).

Defina OUTPUT_PATH como um volume do Catálogo do Unity no qual você pode gravar. Defina NUM_GPUS para o mesmo valor que 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:
  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

O command embutido inicia um cabeçalho do Ray usando a GPU do nó no nó 0, e então executa o driver com python batch_inference.py. Nós de trabalho se conectam ao nó principal usando MASTER_ADDR e NODE_RANK, que a plataforma configura automaticamente. Cada trabalhador monitora o cabeçalho e interrompe seus processos locais do Ray após três falhas consecutivas na verificação de integridade.

Etapa 2: Definir o driver de inferência em lote

batch_inference.py compila um Ray Dataset de prompts, um processador vLLM com ray.data.llm e grava os resultados. O driver espera que todos os nós entrem antes de ler o número de GPUs. O AIR fornece um pool de aceleradores fixos, então o driver configura concurrency para uma tupla fixa (minimum, maximum) que solicita uma réplica por GPU. Como este exemplo usa uma carga de trabalho curta e fixa, o motorista espera até 300 segundos para que todas as réplicas sejam inicializadas antes de despachar o trabalho. Cada ator processa até dois lotes simultaneamente e tem no máximo duas tarefas de Ray Data submetidas, incluindo tarefas em execução e em fila. Isso impede que o primeiro ator a inicializar reserve a maior parte da carga de trabalho. Os 2.000 prompts são divididos em 32 blocos de entrada, com oito blocos disponíveis por réplica. Para cargas de trabalho mais longas, ajuste essas configurações com base no tempo de inicialização e nos requisitos de 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 transforma cada linha de entrada em uma solicitação de chat e postprocess mantém as colunas para persistir. O Ray Data adiciona uma generated_text coluna com a saída do modelo. O script completo está no script de driver completo no final desta página.

tensor_parallel_size=1 mantém cada réplica de vLLM em uma GPU A10.

Etapa 3: Enviar a execução

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

Etapa 4: inspecionar a execução

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

Os logs mostram o prompt do mecanismo vLLM e a taxa de transferência à medida que o lote é executado e, em seguida, uma linha Wrote <n> rows quando a saída é gravada.

Onde os resultados chegam

O driver grava um conjunto de dados Parquet no volume OUTPUT_PATH, com uma coluna instruction e uma coluna output. Releia com Spark ou pandas, por exemplo spark.read.parquet(OUTPUT_PATH).

Script de driver completo

O conteúdo completo batch_inference.py para copiar e colar:

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

Próximas Etapas