Hyperparameterzoektocht met Ray Tune

Belangrijk

Deze functie bevindt zich in openbare preview-versie.

Dit voorbeeld gebruikt Ray Tune om te zoeken naar LoRA-fine-tuning hyperparameters voor Qwen2.5 over 4 1xA10-knooppunten. Een bootstrap-commando start een Ray-cluster die de nodes overspant, en de driver vraagt Ray Tune om één GPU per proef. De cluster draait 4 trials tegelijk, en de rest begint zodra GPU's vrij worden.

De zoekactie gebruikt de ASHA-scheduler (Asynchronous Successive Halving). Elke proef rapporteert dat ze op een vast interval worden gevoerd eval_loss , en ASHA stopt de proeven die achterblijven in plaats van elke kandidaat tot voltooiing te trainen.

Het voorbeeld gebruikt een publiek model (Qwen2.5-0.5B), dus het draait as-is zonder een Hugging Face-token.

De werkbelasting doet het volgende:

  • Uploadt het lokale project met code_source: snapshot.
  • Tokeniseert de dataset één keer op de driver en geeft deze als tensoren door aan de proefuitvoeringen.
  • Bemonstert 8 LoRA-configuraties en voert er 4 tegelijk uit.
  • Registreert de sweep-instellingen, de beste configuratie en de verliezen per proef voor MLflow.

Prerequisites

Projectindeling

Maak een map met de volgende bestanden.

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

Stap 1: de YAML van de workload schrijven

tune.yaml vraagt 4 GPU_1xA10 nodes aan en verklaart zijn afhankelijkheden inline onder environment (met de runtime version). De workload command start een Ray-cluster over de knooppunten en voert vervolgens de driver uit, dus het voorbeeld heeft geen apart afhankelijkheidsbestand of launcherscript nodig:

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

Stap 2: Definieer de zoekruimte en de scheduler

De main driverfunctie tokeniseert de data één keer, definieert de zoekruimte en configureert vervolgens 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}) wijst de zoekopdracht toe aan het cluster. Ray Tune houdt 4 trials gelijktijdig actief omdat het cluster 4 GPU's heeft, dus vergroot je de zoekruimte door num_accelerators in de YAML te verhogen in plaats van de code aan te passen.

Elke run rapporteert na elke EVAL_STEPS optimalisatiestappen. grace_period bepaalt hoeveel evaluaties een run krijgt voordat die kan worden gestopt, max_t begrenst hoeveel evaluaties een overlevende run krijgt, en reduction_factor=2 stopt op elk niveau grofweg de slechtst presterende helft.

Stap 3: Rapporteer de snoeimaatstaf van elke proef

train_fn is een proefversie. Het tune.report gesprek is het moment waarop ASHA het proces stopt of voortzet:

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 vergelijkt proeven op basis eval_loss van een vastgehouden split in plaats van op trainingsverlies, wat de configuraties bevoordeelt die het snelst overfitten. build_datasetstokeniseert de data eenmaal op de driver en geeft objecten terug.TensorDataset tune.with_parameters ze naar proeven op andere knooppunten stuurt. Tensoren serialiseren op waarde, terwijl een Hugging Face-dataset als pad naar een geheugen-gemapt bestand aankomt dat de andere nodes niet kunnen openen.

Het volledige script staat in Full tuning script aan het einde van deze pagina.

Stap 4: Dien de uitvoering in

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

Stap 5: Controleer de uitvoering

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

De driver draait op node 0, dus de Ray Tune-statustabel streamt vanuit de logs van die node, met één rij per proef die de gesamplede configuratie, iteratietelling en laatste eval_losstoont. Runs die door ASHA zijn gestopt, worden weergegeven als TERMINATED met minder iteraties dan max_t.

Waar resultaten terechtkomen

Aan het einde van de run drukt de driver de beste configuratie en de bijbehorende eval_loss af, en logt beide naar het MLflow-experiment dat in experiment_name is genoemd, samen met de sweep-instellingen en de uiteindelijke eval_loss van elke poging.

Het stuurprogramma genereert een fout als een van de pogingen is mislukt.

In het voorbeeld worden adaptergewichten niet opgeslagen. Om de beste adapter te behouden, geef tune.Tuner een RunConfig(storage_path=...) op een Unity Catalog-volume waartoe elke node toegang heeft.

Pas de grootte van de sweep aan

De constanten bovenaan bepalen tune_lora.py de grootte van de sweep. Stel ze kleiner in om in een paar minuten snel een wijziging te testen, al bevatten de eval_loss cijfers dan te veel ruis om configuraties te rangschikken.

De klok volgt NUM_SAMPLES / num_accelerators, dus verhoog num_accelerators in plaats van de zoektocht te verkleinen wanneer een sweep te lang duurt. Voor een groter model, verhoog accelerator_type naar een grotere GPU. Om configuraties te kiezen in plaats van ze willekeurig te bemonsteren, geef een search_alg door TuneConfig zoals Optuna.

Volledig afstemmingsscript

Het volledige tune_lora.py voor kopiëren en plakken:

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

Aanvullende informatiebronnen