Búsqueda de hiperparámetros con Ray Tune

Importante

Esta característica está en versión preliminar pública.

Este ejemplo utiliza Ray Tune para buscar hiperparámetros de ajuste fino LoRA para Qwen2.5 en 4 nodos 1xA10. Un comando de arranque inicia un clúster Ray que abarca los nodos, y el controlador pide a Ray Tune una GPU por prueba. El clúster ejecuta 4 pruebas a la vez, y el resto empiezan a medida que las GPUs se liberan.

La búsqueda utiliza el planificador ASHA (Reducción sucesiva asíncrona). Cada ensayo informa de eval_loss en datos reservados a intervalos de pasos fijos y ASHA detiene las pruebas que se quedan atrás en lugar de entrenar hasta el final cada candidato.

El ejemplo utiliza un modelo público (Qwen2.5-0.5B), por lo que se ejecuta as-is sin un token de Hugging Face.

La carga de trabajo hace lo siguiente:

  • Carga el proyecto local con code_source: snapshot.
  • Tokeniza el conjunto de datos una sola vez en el controlador y lo pasa a las pruebas como tensores.
  • Muestra 8 configuraciones LoRA y ejecuta 4 a la vez.
  • Registra los parámetros de la barrido, la mejor configuración y las pérdidas de cada prueba en MLflow.

Prerequisites

Diseño del proyecto

Cree un directorio con los siguientes archivos.

ray_tune_lora/
├── tune.yaml           # air workload config (inline dependencies + Ray bootstrap)
└── tune_lora.py        # Ray Tune driver + per-trial LoRA fine-tuning

Paso 1: Escribir la carga de trabajo YAML

tune.yaml solicita 4 nodos GPU_1xA10 y declara sus dependencias en línea en environment (con version en tiempo de ejecución). La carga command de trabajo inicia un clúster Ray a través de los nodos, luego ejecuta el controlador, por lo que el ejemplo no necesita ningún archivo de dependencia ni script de lanzador separado:

experiment_name: air-ray-tune-lora

environment:
  version: 'databricks_ai_v5'
  dependencies:
    # databricks_ai_v5 ships ray, transformers, and datasets. It does not ship peft
    # and needs a newer fsspec for huggingface_hub.
    - peft>=0.13
    - fsspec>=2024.6.1

# 4 1xA10 nodes. Ray Tune runs one trial per GPU.
compute:
  num_accelerators: 4
  accelerator_type: GPU_1xA10

code_source:
  type: snapshot
  snapshot:
    root_path: .

command: |
  cd $CODE_SOURCE_PATH
  set -e
  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
    # Stop the cluster on exit, even if the driver fails, so workers don't wait out the timeout.
    trap 'ray stop --grace-period 5' EXIT
    python tune_lora.py
  else
    echo "NODE_RANK=$NODE_RANK: connecting to Ray head at $MASTER_ADDR:$RAY_HEAD_PORT..."
    # `ray start` returns as soon as this node joins, so the worker controls its own exit.
    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; aborting." >&2
      exit 1
    fi

    # health-check exits non-zero once the head runs `ray stop`, which is this worker's cue
    # to exit. The timeout keeps each probe short so the job finishes promptly; the counter
    # caps the total wait.
    echo "Worker joined; waiting for the head to finish the sweep..."
    for _ in $(seq 1 360); do
      if ! timeout 5 ray health-check --address "$MASTER_ADDR:$RAY_HEAD_PORT" 2>/dev/null; then
        break
      fi
      sleep 5
    done
    echo "Head is no longer healthy; stopping local Ray and exiting."
    ray stop --grace-period 5
  fi

max_retries: 0
timeout_minutes: 45

env_variables:
  NCCL_SOCKET_IFNAME: eth0
  HF_HOME: /tmp/hf

Paso 2: Define el espacio de búsqueda y el planificador

La función del main controlador tokeniza los datos una vez, define el espacio de búsqueda y luego configura ASHA:

tuner = tune.Tuner(
    # with_resources gives each trial a whole GPU so trials never share a device.
    tune.with_resources(
        tune.with_parameters(train_fn, train_data=train_data, eval_data=eval_data),
        resources={"gpu": 1},
    ),
    param_space={
        "lr": tune.loguniform(1e-5, 1e-3),
        "lora_r": tune.choice([8, 16, 32]),
        "lora_alpha_ratio": tune.choice([1, 2]),
        "lora_dropout": tune.uniform(0.0, 0.1),
        "weight_decay": tune.choice([0.0, 0.01]),
        "batch_size": tune.choice([4, 8]),
    },
    tune_config=tune.TuneConfig(
        metric="eval_loss",
        mode="min",
        scheduler=ASHAScheduler(
            max_t=MAX_ITERATIONS, grace_period=GRACE_PERIOD, reduction_factor=2
        ),
        num_samples=NUM_SAMPLES,
    ),
)
results = tuner.fit()

tune.with_resources(..., resources={"gpu": 1}) mapea la búsqueda en el clúster. Ray Tune mantiene 4 pruebas en ejecución porque el clúster tiene 4 GPU, así que para ampliar el barrido, aumenta num_accelerators en el YAML en lugar de cambiar el código.

Cada prueba informa de todos los EVAL_STEPS pasos del optimizador. grace_period establece cuántos informes recibe una prueba antes de que pueda detenerse, max_t limita cuántos puede recibir una prueba que continúa y reduction_factor=2 detiene aproximadamente a la mitad inferior en cada nivel.

Paso 3: Informa de la métrica de poda de cada ensayo

train_fn es una prueba. La llamada de tune.report es donde ASHA detiene o continúa la prueba:

def train_fn(config, train_data=None, eval_data=None):
    # Ray Tune pins one GPU per trial via CUDA_VISIBLE_DEVICES, so cuda:0 is this trial's.
    device = torch.device("cuda")

    model = AutoModelForCausalLM.from_pretrained(MODEL_NAME, dtype=torch.bfloat16)
    model.config.use_cache = False
    lora = LoraConfig(
        r=config["lora_r"],
        lora_alpha=config["lora_r"] * config["lora_alpha_ratio"],
        lora_dropout=config["lora_dropout"],
        target_modules=["q_proj", "k_proj", "v_proj", "o_proj"],
        task_type="CAUSAL_LM",
    )
    model = get_peft_model(model, lora).to(device)
    ...
    if step % EVAL_STEPS == 0:
        tune.report({
            "eval_loss": evaluate(model, eval_loader, device),
            "train_loss": out.loss.item(),
            "step": step,
        })

ASHA compara las pruebas en eval_loss en una partición de datos reservada, en lugar de hacerlo en función de la pérdida de entrenamiento, lo que favorecería las configuraciones que se sobreajustan más rápidamente. build_datasets tokeniza los datos una vez en el driver y devuelve objetos TensorDataset. tune.with_parameters Los envía a pruebas en otros nodos. Los tensores se serializan por valor, mientras que un conjunto de datos de Hugging Face llegaría como una ruta a un archivo asignado en memoria que los demás nodos no pueden abrir.

El guion completo está listado en el script de ajuste completo al final de esta página.

Paso 4: Enviar la ejecución

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

Paso 5: Inspección de la ejecución

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

El controlador se ejecuta en el nodo 0, por lo que la tabla de estado de Ray Tune se transmite desde los registros de ese nodo, con una fila por prueba que muestra su configuración muestreada, el número de iteraciones y el eval_loss más reciente. Los ensayos que ASHA detuvo aparecen como TERMINATED con menos iteraciones que max_t.

Dónde llegan los resultados

Al final de la ejecución, el controlador imprime la mejor configuración y su eval_loss, y registra ambos en el experimento de MLflow indicado en experiment_name, junto con la configuración del barrido y el eval_loss final de cada prueba.

El controlador genera un error si alguna prueba ha fallado.

El ejemplo no conserva los pesos del adaptador. Para conservar el mejor adaptador, dé a tune.Tuner un RunConfig(storage_path=...) en un volumen de Unity Catalog al que todos los nodos puedan acceder.

Ajusta el tamaño del barrido

Las constantes en la parte superior de tune_lora.py controlan el tamaño del barrido. Ajústelos a un tamaño más pequeño para realizar una prueba rápida de los cambios en un par de minutos, aunque los valores de eval_loss tienen demasiado ruido como para clasificar configuraciones.

El tempo de reloj sigue NUM_SAMPLES / num_accelerators, así que aumente num_accelerators en lugar de reducir la búsqueda cuando un barrido se prolonga demasiado. Para un modelo más grande, sube accelerator_type a una GPU más grande. Para elegir configuraciones en lugar de muestrearlas aleatoriamente, pasa TuneConfig un search_alg como Optuna.

Script de afinación completo

El tune_lora.py completo para copiar y pegar:

#!/usr/bin/env python3
"""LoRA hyperparameter search for Qwen2.5-0.5B with Ray Tune + ASHA on 4 1xA10 nodes.

The workload's `command` starts a Ray head on node 0 and joins the other nodes as workers,
then runs this script on the head. Ray Tune requests one GPU per trial, so every node runs
one trial at a time. ASHA concentrates GPU time on the promising configurations by stopping
trials that fall behind at each rung.

Uses a public model (no Hugging Face token required) so the example runs as-is.
"""

import os

import mlflow
import ray
import torch
from datasets import load_dataset
from peft import LoraConfig, get_peft_model
from ray import tune
from ray.tune.schedulers import ASHAScheduler
from torch.utils.data import DataLoader, TensorDataset
from transformers import AutoModelForCausalLM, AutoTokenizer

MODEL_NAME = "Qwen/Qwen2.5-0.5B"
DATASET_NAME = "tatsu-lab/alpaca"
MAX_SEQ_LEN = 512

# Trials report every EVAL_STEPS optimizer steps, so ASHA sees at most MAX_ITERATIONS
# reports per trial and can start pruning once a trial has sent GRACE_PERIOD of them.
EVAL_STEPS = 25
MAX_ITERATIONS = 12
GRACE_PERIOD = 3

NUM_SAMPLES = 8
TRAIN_EXAMPLES = 2000
EVAL_EXAMPLES = 200


def build_datasets():
    """Tokenizes the SFT data once on the driver.

    Returns TensorDatasets so the tokenized splits serialize by value, which is what lets
    tune.with_parameters hand them to trials on any node in the cluster.
    """
    tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
    if tokenizer.pad_token is None:
        tokenizer.pad_token = tokenizer.eos_token

    raw = load_dataset(DATASET_NAME, split=f"train[:{TRAIN_EXAMPLES + EVAL_EXAMPLES}]")

    def format_example(row):
        prompt = f"### Instruction:\n{row['instruction']}\n\n"
        if row.get("input"):
            prompt += f"### Input:\n{row['input']}\n\n"
        text = f"{prompt}### Response:\n{row['output']}{tokenizer.eos_token}"
        out = tokenizer(text, truncation=True, max_length=MAX_SEQ_LEN, padding="max_length")
        # -100 is cross-entropy's ignore_index, so the loss covers only real tokens and
        # eval_loss stays a meaningful signal for ASHA to rank trials by.
        out["labels"] = [token if mask == 1 else -100 for token, mask in zip(out["input_ids"], out["attention_mask"])]
        return out

    tokenized = raw.map(format_example, remove_columns=raw.column_names)
    split = tokenized.train_test_split(test_size=EVAL_EXAMPLES, shuffle=True, seed=0)

    def to_tensors(ds):
        return TensorDataset(
            torch.tensor(ds["input_ids"], dtype=torch.long),
            torch.tensor(ds["attention_mask"], dtype=torch.long),
            torch.tensor(ds["labels"], dtype=torch.long),
        )

    return to_tensors(split["train"]), to_tensors(split["test"])


def evaluate(model, loader, device):
    """Mean cross-entropy over the held-out split. This is the metric ASHA prunes on."""
    model.eval()
    total, batches = 0.0, 0
    with torch.no_grad():
        for input_ids, attention_mask, labels in loader:
            out = model(
                input_ids=input_ids.to(device),
                attention_mask=attention_mask.to(device),
                labels=labels.to(device),
            )
            total += out.loss.item()
            batches += 1
    model.train()
    return total / max(batches, 1)


def train_fn(config, train_data=None, eval_data=None):
    """One trial: LoRA fine-tunes Qwen on a single GPU and reports eval_loss to ASHA."""
    # Ray Tune pins one GPU per trial via CUDA_VISIBLE_DEVICES, so cuda:0 is this trial's.
    device = torch.device("cuda")

    model = AutoModelForCausalLM.from_pretrained(MODEL_NAME, dtype=torch.bfloat16)
    model.config.use_cache = False
    lora = LoraConfig(
        r=config["lora_r"],
        lora_alpha=config["lora_r"] * config["lora_alpha_ratio"],
        lora_dropout=config["lora_dropout"],
        target_modules=["q_proj", "k_proj", "v_proj", "o_proj"],
        task_type="CAUSAL_LM",
    )
    model = get_peft_model(model, lora).to(device)

    train_loader = DataLoader(train_data, batch_size=config["batch_size"], shuffle=True, drop_last=True)
    eval_loader = DataLoader(eval_data, batch_size=config["batch_size"])

    optimizer = torch.optim.AdamW(
        (p for p in model.parameters() if p.requires_grad),
        lr=config["lr"],
        weight_decay=config["weight_decay"],
    )

    model.train()
    step = 0
    max_steps = EVAL_STEPS * MAX_ITERATIONS
    # Cycle the loader over multiple epochs until the step budget is spent.
    while step < max_steps:
        for input_ids, attention_mask, labels in train_loader:
            out = model(
                input_ids=input_ids.to(device),
                attention_mask=attention_mask.to(device),
                labels=labels.to(device),
            )
            out.loss.backward()
            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
            optimizer.step()
            optimizer.zero_grad()
            step += 1

            if step % EVAL_STEPS == 0:
                # ASHA stops or continues the trial based on this report.
                tune.report(
                    {
                        "eval_loss": evaluate(model, eval_loader, device),
                        "train_loss": out.loss.item(),
                        "step": step,
                    }
                )
            if step >= max_steps:
                break


def main():
    ray.init(address="auto")

    num_nodes = int(os.environ.get("NUM_NODES", 1))
    total_gpus = int(ray.cluster_resources().get("GPU", 0))
    if total_gpus < 1:
        raise SystemExit("No GPUs registered with Ray; check GPU discovery on the cluster.")
    print(f"Cluster ready: {num_nodes} node(s), {total_gpus} GPU(s)", flush=True)
    print(f"Running {NUM_SAMPLES} trials, up to {total_gpus} concurrently\n", flush=True)

    train_data, eval_data = build_datasets()

    param_space = {
        "lr": tune.loguniform(1e-5, 1e-3),
        "lora_r": tune.choice([8, 16, 32]),
        "lora_alpha_ratio": tune.choice([1, 2]),
        "lora_dropout": tune.uniform(0.0, 0.1),
        "weight_decay": tune.choice([0.0, 0.01]),
        "batch_size": tune.choice([4, 8]),
    }

    tuner = tune.Tuner(
        # with_resources gives each trial a whole GPU so trials never share a device.
        tune.with_resources(
            tune.with_parameters(train_fn, train_data=train_data, eval_data=eval_data),
            resources={"gpu": 1},
        ),
        param_space=param_space,
        tune_config=tune.TuneConfig(
            metric="eval_loss",
            mode="min",
            scheduler=ASHAScheduler(
                max_t=MAX_ITERATIONS,
                grace_period=GRACE_PERIOD,
                reduction_factor=2,
            ),
            num_samples=NUM_SAMPLES,
        ),
    )

    results = tuner.fit()

    # Surface trial failures: a best result is only meaningful when the whole sweep ran.
    if results.num_errors:
        raise RuntimeError(
            f"{results.num_errors} of {len(results)} trials errored; see the per-trial error files above."
        )

    best = results.get_best_result("eval_loss", "min")
    print(f"\nBest config:    {best.config}", flush=True)
    print(f"Best eval_loss: {best.metrics['eval_loss']:.4f}", flush=True)

    # AI Runtime injects MLFLOW_RUN_ID and configures the databricks tracking URI on the
    # node, so logging needs no credentials here. Gating on the variable keeps the script
    # runnable off-platform, where it is unset.
    if os.environ.get("MLFLOW_RUN_ID"):
        with mlflow.start_run(run_id=os.environ["MLFLOW_RUN_ID"]):
            mlflow.log_params(
                {
                    "model": MODEL_NAME,
                    "dataset": DATASET_NAME,
                    "num_samples": NUM_SAMPLES,
                    "scheduler": "ASHA",
                    "asha_max_t": MAX_ITERATIONS,
                    "asha_grace_period": GRACE_PERIOD,
                    **{f"best_{k}": v for k, v in best.config.items()},
                }
            )
            mlflow.log_metric("best_eval_loss", best.metrics["eval_loss"])
            for i, result in enumerate(results):
                if result.metrics and "eval_loss" in result.metrics:
                    mlflow.log_metric("trial_eval_loss", result.metrics["eval_loss"], step=i)

    ray.shutdown()


if __name__ == "__main__":
    main()

Recursos adicionales