Aller au contenu principal

Recherche d'hyperparamètres avec Ray Tune

info

Aperçu

Cette fonctionnalité est en Aperçu public.

Cet exemple utilise Ray Tune pour rechercher les hyperparamètres d'affinement LoRA pour Qwen2.5 sur 4 nœuds 1xA10. Une commande d'amorçage démarre un cluster Ray qui s'étend sur les nœuds, et le Driver demande à Ray Tune un GPU par essai. Le cluster exécute 4 essais à la fois, et les autres start à mesure que les GPU se libèrent.

La recherche utilise le planificateur ASHA (Asynchronous Successive Halving). Chaque essai rapporte le eval_loss mis de côté à un intervalle d’étapes fixe, et ASHA arrête les essais qui prennent du retard au lieu d’entraîner chaque candidat jusqu’à son terme.

L’exemple utilise un modèle public (Qwen2.5-0.5B), afin qu’il s’exécute tel quel sans jeton Hugging Face.

La charge de travail réalise les étapes suivantes :

  • Importe le projet local avec code_source: snapshot.
  • Tokenise le dataset une fois sur le driver et le transmet aux essais sous forme de tenseurs.
  • Échantillonne 8 configurations LoRA et en exécute 4 à la fois.
  • Logs les paramètres de balayage, la meilleure configuration et les pertes par essai dans MLflow.

Prérequis

Layout du projet

Créez un répertoire avec les fichiers suivants.

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

Étape 1 : rédiger le YAML de la charge de travail

tune.yaml demande 4 nœuds GPU_1xA10 et déclare ses dépendances en ligne sous environment (avec le runtime version). Le command de la charge de travail démarre un cluster Ray sur les nœuds, puis exécute le Driver, de sorte que l'exemple ne nécessite aucun fichier de dépendance ou script de lancement distinct :

YAML
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

Étape 2 : définissez l’espace de recherche et le planificateur

La fonction main du driver tokenise les données une fois, définit l'espace de recherche, puis configure ASHA :

Python
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}) associe la recherche au cluster. Ray Tune maintient 4 essais en cours car le cluster dispose de 4 GPU ; pour élargir la recherche, augmentez num_accelerators dans le fichier YAML plutôt que de modifier le code.

Chaque essai rapporte tous les EVAL_STEPS pas de l'optimiseur. grace_period définit le nombre de rapports qu'un essai obtient avant de pouvoir être arrêté, max_t limite le nombre de rapports qu'un essai survivant obtient, et reduction_factor=2 arrête environ la moitié inférieure à chaque échelon.

Étape 3 : Signaler la métrique d'élagage de chaque essai

train_fn est un essai. L'appel tune.report est l'endroit où ASHA arrête ou poursuit l'essai :

Python
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 compare les essais sur eval_loss à partir d'un fractionnement mis de côté plutôt que sur la perte d'entraînement, ce qui favoriserait les configurations qui surapprennent le plus rapidement. build_datasets tokenise les données une fois sur le driver et renvoie TensorDataset objets. tune.with_parameters les envoie aux essais sur d'autres nœuds. Les tenseurs sont sérialisés par valeur, tandis qu'un dataset Hugging Face arriverait sous la forme d'un chemin d'accès vers un fichier mappé en mémoire que les autres nœuds ne peuvent pas ouvrir.

Le script complet est répertorié dans Script d'affinement complet à la fin de cette page.

Étape 4 : Soumettre l'exécution

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

Étape 5 : Inspecter l'exécution

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

Le Driver s'exécute sur le nœud 0, le tableau d'état de Ray Tune est donc Stream à partir des Logs de ce nœud, avec une ligne par essai affichant sa configuration échantillonnée, le nombre d'itérations et le dernier eval_loss. Les essais arrêtés par ASHA apparaissent comme TERMINATED avec moins d'itérations que max_t.

Où les résultats sont stockés

À la fin de l’exécution, le driver affiche la meilleure configuration et son eval_loss, et Logs les deux dans l’expérimentation MLflow nommée dans experiment_name, ainsi que les paramètres de sweep et le eval_loss final de chaque essai.

Le Driver génère une erreur si un essai a échoué.

L'exemple ne conserve pas les poids de l'adaptateur. Pour conserver le meilleur adaptateur, donnez à tune.Tuner un RunConfig(storage_path=...) sur un volume Unity Catalog accessible par chaque nœud.

Ajuster la taille du balayage

Les constantes en haut de tune_lora.py contrôlent la taille du balayage. Définissez-les sur une valeur plus petite pour effectuer un test de fumée d'une modification en quelques minutes, bien que les chiffres eval_loss soient alors trop bruités pour classer les configurations.

La durée prévue suit NUM_SAMPLES / num_accelerators, augmentez donc num_accelerators plutôt que de réduire la recherche lorsqu'un balayage prend trop de temps. Pour un modèle plus grand, augmentez accelerator_type vers un GPU plus puissant. Pour choisir des configurations au lieu de les échantillonner de manière aléatoire, passez à TuneConfig un search_alg tel que Optuna.

Script d'affinement complet

Le tune_lora.py complet pour le copier-coller :

Python
#!/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={&quot;gpu&quot;: 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()

Ressources supplémentaires