Inferência em lote com Ray Data e vLLM
Visualização
Este recurso está em Pré-visualização Pública.
Este exemplo executa inferência de LLM em lote offline com
Ray Data e
vLLM em 4 nós A10. Um script de bootstrap inicia um
cluster Ray nos nós, então o Driver usa a API de LLM do Ray Data (ray.data.llm) para
iniciar uma réplica vLLM por nó e Stream um dataset de prompts através deles, gravando o
texto gerado em um volume do Unity Catalog como Parquet.
Ele usa um modelo público (Qwen2.5-7B-Instruct), assim, ele é executado como está sem um token do Hugging Face.
A carga de trabalho executa as seguintes ações:
- Faz upload do projeto local com
code_source: snapshot. - Inicia um head do Ray no nó 0, faz join de 3 nós worker e, em seguida, executa o Driver de inferência em lotes.
- Usa
ray.data.llmpara executar uma réplica 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
airCLI instalada e autenticada. Consulte Instalar a CLI do Runtime de AI.- Um volume do Unity Catalog gravável. Você define o caminho no YAML de carga de trabalho abaixo.
Disposição do projeto
Criar um diretório com os seguintes arquivos.
ray_batch_inference/
├── train.yaml # air workload config (inline dependencies + Ray bootstrap)
└── batch_inference.py # Ray Data + vLLM batch inference driver
O passo 1: Escreva a carga de trabalho YAML
train.yaml solicita 4 GPU_1xA10 nós. As dependências são declaradas em linha sob
environment (com a imagem do cliente version), e o command inicia um cluster Ray
nos nós e, em seguida, executa o Driver, portanto, a carga de trabalho não precisa de um arquivo
de dependência ou script de inicialização separado.
O vLLM não está na imagem base, então é instalado em linha junto com três pins que os nós da GPU
precisam: hf_transfer (a imagem base permite downloads rápidos do Hugging Face e espera este
pacote), um fsspec mais recente (a imagem base vem com um antigo que impede downloads) e um
opencv-python-headless fixado (o vLLM incorpora o OpenCV, cujo wheel default causa a falha do
auto-teste OpenSSL FIPS nos nós da GPU).
Defina OUTPUT_PATH como um volume do Unity Catalog no qual você possa gravar. Defina NUM_GPUS com o mesmo valor de 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
O command em linha inicia um head do Ray com a GPU do nó no nó 0 e, em seguida, executa o driver
com python batch_inference.py. Os nós worker fazem join do head usando MASTER_ADDR e
NODE_RANK, que a plataforma define automaticamente. Cada worker monitora o head e interrompe
seus processos locais do Ray após três falhas consecutivas na verificação de integridade.
Passo 2: Definir o driver de inferência em lote
batch_inference.py cria um Ray Dataset de prompts, configura um processador vLLM com
ray.data.llm e grava os resultados. O driver aguarda que todos os nós façam o join antes de
ler a contagem de GPU. O AIR provisiona um pool de aceleradores fixo, portanto, o driver define
concurrency como uma tupla (minimum, maximum) fixa que solicita uma réplica por GPU.
Como este exemplo usa uma carga de trabalho curta e fixa, o driver aguarda 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 do Ray Data enviadas, incluindo tarefas
em execução e na fila. Isso evita 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 Startup 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 a serem persistidas. Ray Data adiciona uma coluna generated_text com a saída do modelo. O script completo está em Script completo do driver no final desta página.
tensor_parallel_size=1 mantém cada réplica do vLLM em uma GPU A10.
Passo 3: Enviar a execução
air run -f train.yaml --dry-run
air run -f train.yaml --watch
o passo 4: inspeção da execução
air get run <run-id>
air logs <run-id>
Os logs mostram o prompt e a taxa de transferência de geração do motor vLLM enquanto o lote é executado, então uma linha Wrote <n> rows quando a saída é gravada.
Onde os resultados são armazenados
O driver grava um dataset Parquet no volume OUTPUT_PATH, com uma coluna instruction
e uma coluna output. Leia de volta com Spark ou Pandas, por exemplo
spark.read.parquet(OUTPUT_PATH).
Script de driver completo
O batch_inference.py completo 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()