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,
curlandjq(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):
{
"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}
}
}
}
}
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.
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#
{
"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).
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. |