Skip to content

Hyperparameter sweep with an array#

You run a grid of 12 trials (4 learning rates × 3 weight decays) as one array run: 12 independent workers built from one template, each with its own index in ASTRAEUS_ARRAY_TASK_ID, at most 4 at a time. Each trial trains a ResNet-18 on CIFAR-10 for 5 epochs on one GPU and prints its score; you then rank the trials. The dataset is downloaded once per machine onto a drive, not once per trial.

What you need:

  • A workspace with GPU machines that have a data location chosen.
  • The editor role in the workspace, an API token, curl and jq (How the recipes are written).

How an array runs#

Workers <run>-0 … <run>-11, one per index. Each has the template's whole ask (here 1 GPU).
Index ASTRAEUS_ARRAY_TASK_ID (0 to size − 1) and ASTRAEUS_ARRAY_SIZE, unless you set them in env.
Start Each worker starts when it fits, independently; several may share a machine.
max_parallel At most this many workers run at once; the others wait in Pending. 0 or absent: no cap.
Limits size from 1 to 10,000; max_parallel at most size. An array cannot be a service, cannot have worker groups, and starts Independent.
End Completed when every worker completed; Failed if any failed for good.

1. Put the dataset on a drive#

Twelve trials downloading CIFAR-10 at once on the same machine would fight over the same files. A drive kept on each machine with a fill downloads it once per machine, before any trial there starts (A dataset on a drive, shared by runs):

cifar10-drive.json
{
  "metadata": {"name": "cifar10"},
  "spec": {
    "sources": [{"node_scope": "placed", "path": "", "mount_path": "/data", "mode": "ReadOnly"}],
    "size_hint_bytes": 400000000,
    "evictable": true,
    "fill": {
      "template": {
        "image": "pytorch/pytorch:2.4.1-cuda12.4-cudnn9-runtime",
        "command": "sh",
        "args": ["-c", "set -e; python -c \"import torchvision as tv; tv.datasets.CIFAR10('/drive', download=True); tv.datasets.CIFAR10('/drive', train=False, download=True)\"; touch /drive/.astralyx-filled"],
        "requested_resources": {"cpu_cores": 2, "memory_bytes": 4294967296}
      }
    }
  }
}

Drives → New drive → Edit as JSON, paste cifar10-drive.json, Create drive.

$ curl -fsS -X POST "$API/drives" -H "Authorization: Bearer $ASTRAEUS_TOKEN" \
    -H 'content-type: application/json' -d @cifar10-drive.json | jq -r .metadata.name
cifar10

2. Write the trial script#

The script turns its index into one point of the grid and prints a single RESULT line at the end.

trial.py
import itertools
import json
import os
import time

import torch
import torch.nn as nn
import torchvision
import torchvision.transforms as T

LRS = [0.05, 0.1, 0.2, 0.4]
WDS = [1e-4, 5e-4, 1e-3]
GRID = list(itertools.product(LRS, WDS))  # 12 points

trial = int(os.environ["ASTRAEUS_ARRAY_TASK_ID"])
size = int(os.environ["ASTRAEUS_ARRAY_SIZE"])
assert size == len(GRID), f"array size {size} does not match the grid ({len(GRID)})"
lr, wd = GRID[trial]
epochs = int(os.environ.get("EPOCHS", "5"))
torch.manual_seed(trial)
print(f"trial {trial}/{size}: lr={lr} wd={wd} on {os.environ['ASTRAEUS_NODE_NAME']}", flush=True)

norm = T.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616))
train = torchvision.datasets.CIFAR10("/data", train=True, download=False,
                                     transform=T.Compose([T.RandomCrop(32, padding=4), T.RandomHorizontalFlip(), T.ToTensor(), norm]))
test = torchvision.datasets.CIFAR10("/data", train=False, download=False, transform=T.Compose([T.ToTensor(), norm]))
train_dl = torch.utils.data.DataLoader(train, 256, shuffle=True, num_workers=4, pin_memory=True, drop_last=True)
test_dl = torch.utils.data.DataLoader(test, 1024, num_workers=4)

model = torchvision.models.resnet18(num_classes=10)
model.conv1 = nn.Conv2d(3, 64, 3, 1, 1, bias=False)
model.maxpool = nn.Identity()
model = model.cuda().to(memory_format=torch.channels_last)
opt = torch.optim.SGD(model.parameters(), lr=lr, momentum=0.9, weight_decay=wd, nesterov=True)
sched = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=lr, epochs=epochs, steps_per_epoch=len(train_dl))
loss_fn = nn.CrossEntropyLoss()

t0 = time.time()
for epoch in range(epochs):
    model.train()
    for x, y in train_dl:
        with torch.autocast("cuda", dtype=torch.bfloat16):
            loss = loss_fn(model(x.cuda(non_blocking=True).to(memory_format=torch.channels_last)), y.cuda(non_blocking=True))
        opt.zero_grad(set_to_none=True)
        loss.backward()
        opt.step()
        sched.step()
    print(f"epoch {epoch + 1}: loss {loss.item():.4f}", flush=True)

model.eval()
correct = 0
with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16):
    for x, y in test_dl:
        correct += (model(x.cuda().to(memory_format=torch.channels_last)).argmax(1).cpu() == y).sum().item()
acc = correct / len(test)
if acc != acc:  # NaN: a diverged trial still reports
    acc = 0.0
print("RESULT " + json.dumps({"trial": trial, "lr": lr, "wd": wd, "test_acc": round(acc, 4), "seconds": round(time.time() - t0)}), flush=True)

3. Write the array run#

sweep.json
{
  "metadata": {"name": "sweep-r18", "labels": {"sweep": "r18-lr-wd"}},
  "spec": {
    "start": "Independent",
    "array": {"size": 12, "max_parallel": 4},
    "task_template": {
      "image": "pytorch/pytorch:2.4.1-cuda12.4-cudnn9-runtime",
      "command": "python",
      "args": ["/app/trial.py"],
      "env": {"EPOCHS": "5"},
      "restart_policy": "Never",
      "time_limit_seconds": 900,
      "requested_resources": {
        "cpu_cores": 6,
        "memory_bytes": 17179869184,
        "gpu_requests": {"count": 1},
        "node_selection": {"mode": "Any"}
      },
      "datavolume_refs": [{"name": "cifar10", "mount_path": "/data", "mode": "ReadOnly"}],
      "configs": []
    }
  }
}
$ jq --rawfile src trial.py \
    '.spec.task_template.configs = [{"mounts": ["/app/trial.py"], "value": $src}]' \
    sweep.json > sweep.full.json
Field Why
array.size: 12 One worker per grid point. The script checks the size matches.
array.max_parallel: 4 At most 4 GPUs for the sweep at any time, so it does not crowd out the rest of the workspace.
restart_policy: Never A trial that crashes is a result, not something to retry ten times.
time_limit_seconds: 900 Each trial is stopped after 15 minutes. Stating it also lets Astraeus backfill trials into room it holds for a larger run that starts later, if they end before it is due.

4. Submit it#

Runs → New run → Edit as JSON, paste sweep.full.json, Start run. In the form, the same settings are under More options: Array size 12, At most in parallel 4, Time limit 15m.

astra astraeus run has no array flags; submit arrays from the console or the API. Follow them with the CLI (step 5).

$ curl -fsS -X POST "$API/runs" -H "Authorization: Bearer $ASTRAEUS_TOKEN" \
    -H 'content-type: application/json' -d @sweep.full.json | jq -r .metadata.name
sweep-r18

5. Follow it#

$ astra astraeus runs
NAME                          CLUSTER       STATE       REASON
sweep-r18                     main          Running     4 of 12 workers running
$ curl -fsS "$API/workers?job=sweep-r18" -H "Authorization: Bearer $ASTRAEUS_TOKEN" \
    | jq -r '.items[] | [.metadata.name, .status.state, (.spec.assigned_node.name // "-")] | @tsv'
sweep-r18-0 Running gpu-01
sweep-r18-1 Running gpu-01
sweep-r18-2 Running gpu-02
sweep-r18-3 Running gpu-02
sweep-r18-4 Pending -
sweep-r18-5 Pending -
…

The first trial on each machine waits for the cifar10 fill there (Waiting for drive <namespace>.cifar10 to be filled on gpu-01); the others on that machine then start at once. The run's Workers tab shows every trial, its machine and its exit code.

6. Collect and rank the results#

When the run is Completed, read each trial's RESULT line and sort:

$ for i in $(seq 0 11); do astra astraeus logs sweep-r18-$i | tail -n 5 | sed -n 's/^RESULT //p'; done \
    | jq -s 'sort_by(-.test_acc) | .[:3]'
[
  {"trial": 4, "lr": 0.1, "wd": 0.0005, "test_acc": 0.9012, "seconds": 171},
  {"trial": 7, "lr": 0.2, "wd": 0.0005, "test_acc": 0.8987, "seconds": 169},
  {"trial": 3, "lr": 0.1, "wd": 0.0001, "test_acc": 0.8874, "seconds": 172}
]

The same with the API only:

$ for i in $(seq 0 11); do
    curl -fsS "$API/workers/sweep-r18-$i/logs?tail=5" -H "Authorization: Bearer $ASTRAEUS_TOKEN" \
      | jq -r .logs | sed -n 's/^RESULT //p'
  done | jq -s 'sort_by(-.test_acc) | .[0]'

To check the sweep worked: 12 RESULT lines, one per trial, and the run Completed with All workers completed. If a trial failed, the run is Failed and its reason names the first worker that failed and why.

7. Clean up#

$ astra astraeus delete sweep-r18
sweep-r18 deleted
$ curl -fsS -X DELETE "$API/drives/cifar10" -H "Authorization: Bearer $ASTRAEUS_TOKEN"

Keep the cifar10 drive if you will sweep again: its copies make the next sweep start without downloading.

Variations#

Results on a drive, and a report run. For more than a line per trial (curves, checkpoints), write to a drive every trial reaches, such as a drive on one machine or on a shared filesystem, mounted read-write at /results: /results/sweep-r18/trial-<i>.json. Then submit a small CPU run that mounts it read-only, reads every file and writes summary.json. Submit it once the sweep is Completed; there is no dependency between runs.

Random search. Draw the point from a seeded generator instead of a grid: random.Random(trial).choice(...), 10 ** random.Random(trial).uniform(-3, -1). The same index always gives the same point, so a trial that is run again repeats its configuration.

Bigger trials. Each worker of an array gets the template's whole ask: with "gpu_requests": {"count": 8} every trial gets 8 GPUs on one machine (use torchrun --standalone --nproc-per-node=$ASTRAEUS_GPU_COUNT).

Scripts written for Slurm arrays. Astraeus does not set SLURM_ARRAY_TASK_ID. Set it from the index at the start of the command: sh -c 'export SLURM_ARRAY_TASK_ID=$ASTRAEUS_ARRAY_TASK_ID; exec python train.py'. See Slurm compatibility.

Lower priority. Add "priority": -10 so the sweep gives way to other work in the workspace: queued trials go after higher-priority runs, and running trials may be preempted (each one alone) and go back to the queue.

Troubleshooting#

Symptom Cause Fix
Must not exceed array.size (12), got 16 max_parallel larger than size. Lower it, or leave it out for no cap.
An array's workers are independent: use policy Independent start: Gang (or policy: Gang) with an array. Use "start": "Independent" or leave it out.
FileNotFoundError / Dataset not found in a trial The trial ran before the copy was whole, or /data is the wrong folder. Check the fill's log and that the trial mounts cifar10 at /data.
Only some trials run at once, fewer than max_parallel No more free GPUs in your pools, or your workspace's quota is reached. The pending workers' pages say which.