Preemptible fine-tuning with checkpoints#
You run a fine-tune at low priority so that it uses GPUs nobody else needs,
and gives them back the moment someone does. When a higher-priority run
needs its GPU, Astraeus stops the fine-tune gracefully: the container gets
SIGTERM, the script writes a checkpoint and exits, and the run goes back
to the queue without spending its restart budget. When a GPU is free again,
it starts from the checkpoint.
The example fine-tunes an ImageNet-trained ResNet-50 on CIFAR-100 on one GPU, then shows an 8-GPU run preempting it.
What you need:
- A workspace with a GPU machine. To watch the preemption, one machine with 8 GPUs whose other GPUs are free.
- A drive the fine-tune reaches from any machine, for its checkpoints. This
recipe uses a drive on one machine, which needs a host path granted to
the workspace (an organisation admin sets it under
Settings → Clusters → Change → Host paths, for example
/data/astraeus). A drive on a shared filesystem works the same way. - The editor role in the workspace, an API token,
curlandjq(How the recipes are written).
How preemption works#
sequenceDiagram
participant Q as Queue
participant F as finetune-r50-0 (priority -100)
participant U as urgent-eval (priority 0)
U->>Q: waits: no room
Q->>F: Stopping (preempted by urgent-eval)
Note over F: SIGTERM → save ckpt.pt → exit<br/>SIGKILL after the stop grace (30 s)
Q->>U: placed on the GPUs freed
Note over F: Pending: waiting to run again
U-->>Q: completes
Q->>F: placed again → resumes from ckpt.pt
- The queue orders waiting runs by priority (higher first, from −1000 to 1000), then by fair share between workspaces, then by submission time.
- When the run at the head of the queue cannot start and lower-priority work is running, Astraeus stops the fewest units of it that make room: lowest priority first, and among equals the most recently started, so the least work is lost. A gang is stopped whole; an independent worker alone.
- A preempted worker goes back to the queue with its restart count unchanged. Until it is gone, the room it frees is held for the run that preempted it.
- A workspace may ask for priorities up to its max priority (0 unless an organisation admin raised it). Negative priorities are always allowed, which is what makes a scavenger run possible without a grant.
1. Create the checkpoint drive#
{
"metadata": {"name": "ft-ckpt"},
"spec": {
"sources": [
{"node": "gpu-01", "path": "/data/astraeus/ft-ckpt", "mount_path": "/ckpt", "mode": "ReadWrite"}
],
"transport": "Auto",
"permissions": {"create": true}
}
}
The bytes live on gpu-01 under /data/astraeus/ft-ckpt. A worker on
gpu-01 binds the folder; a worker on another machine mounts it over NFS
(over RDMA when both machines have it, else TCP). permissions.create
makes the folder if it does not exist. Astraeus never deletes the files of
a drive on one machine.
- Drives → New drive. Name
ft-ckpt. - Kind: On one machine. Machine:
gpu-01. Path on the machine:/data/astraeus/ft-ckpt. - Mounted in workers at
/ckpt, Access Read-write. Across machines: RDMA when both have it, else TCP. - Create drive.
2. Write a script that checkpoints on SIGTERM#
Three habits make a run safe to preempt:
- Handle
SIGTERMby finishing the current step, saving, and exiting. The container's main process gets the signal, so start Python directly (or withexec), not under a shell that swallows it. - Save atomically: write a temporary file, then rename it over the old checkpoint. A worker killed mid-write leaves the previous checkpoint intact.
- Save periodically as well. The machine kills the container when its stop grace (30 seconds unless the machine's administrator changed it) runs out, and a machine can fail without warning.
import json
import os
import signal
import sys
import time
import torch
import torch.nn as nn
import torchvision
import torchvision.transforms as T
from torch.utils.data import DataLoader, Subset
CKPT_DIR = os.environ.get("CKPT_DIR", "/ckpt")
EPOCHS = int(os.environ.get("EPOCHS", "20"))
BATCH = int(os.environ.get("BATCH_SIZE", "128"))
SAVE_EVERY = int(os.environ.get("SAVE_EVERY", "200")) # steps
RUN = os.environ.get("ASTRAEUS_JOB_NAME", "local").rsplit(".", 1)[-1]
run_dir = os.path.join(CKPT_DIR, RUN)
os.makedirs(run_dir, exist_ok=True)
ckpt_path = os.path.join(run_dir, "ckpt.pt")
stop_requested = False
def on_sigterm(signum, frame):
global stop_requested
stop_requested = True
print("SIGTERM: saving a checkpoint after this step", flush=True)
signal.signal(signal.SIGTERM, on_sigterm)
device = "cuda"
tf = T.Compose([T.Resize(224), T.RandomHorizontalFlip(), T.ToTensor(),
T.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))])
train = torchvision.datasets.CIFAR100(os.path.join(CKPT_DIR, "data"), train=True, download=True, transform=tf)
steps_per_epoch = len(train) // BATCH
model = torchvision.models.resnet50(weights="IMAGENET1K_V2")
model.fc = nn.Linear(model.fc.in_features, 100)
model = model.to(device, memory_format=torch.channels_last)
opt = torch.optim.SGD(model.parameters(), lr=0.01, momentum=0.9, weight_decay=1e-4)
sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=EPOCHS * steps_per_epoch)
loss_fn = nn.CrossEntropyLoss()
start_epoch, start_step, global_step = 0, 0, 0
if os.path.exists(ckpt_path):
ck = torch.load(ckpt_path, map_location=device)
model.load_state_dict(ck["model"])
opt.load_state_dict(ck["opt"])
sched.load_state_dict(ck["sched"])
start_epoch, start_step, global_step = ck["epoch"], ck["step"], ck["global_step"]
print(f"resuming at epoch {start_epoch + 1}, step {start_step} (global step {global_step})", flush=True)
def save(epoch, step):
tmp = ckpt_path + ".tmp"
torch.save({"model": model.state_dict(), "opt": opt.state_dict(), "sched": sched.state_dict(),
"epoch": epoch, "step": step, "global_step": global_step}, tmp)
os.replace(tmp, ckpt_path) # atomic: the old checkpoint stays whole until here
for epoch in range(start_epoch, EPOCHS):
# The same order every time this epoch is run, so a resumed epoch can
# skip exactly the batches already seen.
order = torch.randperm(len(train), generator=torch.Generator().manual_seed(epoch)).tolist()
skip = start_step if epoch == start_epoch else 0
dl = DataLoader(Subset(train, order[skip * BATCH:]), BATCH, num_workers=6, pin_memory=True, drop_last=True)
model.train()
for i, (x, y) in enumerate(dl, start=skip):
x = x.to(device, non_blocking=True, memory_format=torch.channels_last)
y = y.to(device, non_blocking=True)
with torch.autocast("cuda", dtype=torch.bfloat16):
loss = loss_fn(model(x), y)
opt.zero_grad(set_to_none=True)
loss.backward()
opt.step()
sched.step()
global_step += 1
if stop_requested:
t0 = time.time()
save(epoch, i + 1)
print(f"checkpoint at global step {global_step} written in {time.time() - t0:.1f}s; exiting", flush=True)
sys.exit(0)
if global_step % SAVE_EVERY == 0:
save(epoch, i + 1)
print(json.dumps({"epoch": epoch + 1, "global_step": global_step, "loss": round(loss.item(), 4)}), flush=True)
save(epoch + 1, 0)
torch.save(model.state_dict(), os.path.join(run_dir, "model.pt"))
print("done", flush=True)
3. Submit the fine-tune at low priority#
{
"metadata": {"name": "finetune-r50", "labels": {"tier": "scavenger"}},
"spec": {
"priority": -100,
"task_template": {
"image": "pytorch/pytorch:2.4.1-cuda12.4-cudnn9-runtime",
"command": "python",
"args": ["/app/finetune.py"],
"env": {"CKPT_DIR": "/ckpt", "EPOCHS": "20", "SAVE_EVERY": "200", "TORCH_HOME": "/ckpt/torch"},
"restart_policy": "OnFailure",
"requested_resources": {
"cpu_cores": 8,
"memory_bytes": 34359738368,
"gpu_requests": {"count": 1},
"node_selection": {"mode": "Any"}
},
"datavolume_refs": [{"name": "ft-ckpt", "mount_path": "/ckpt", "mode": "ReadWrite"}],
"configs": []
}
}
}
TORCH_HOME on the drive keeps the downloaded ImageNet weights between
attempts.
$ jq --rawfile src finetune.py \
'.spec.task_template.configs = [{"mounts": ["/app/finetune.py"], "value": $src}]' \
finetune.json > finetune.full.json
Runs → New run → Edit as JSON, paste finetune.full.json,
Start run. In the form, priority is under More options →
Priority; its hint shows the highest priority your workspace may ask
for.
The run starts, and the log shows a line every 200 steps:
$ astra astraeus logs finetune-r50-0 | tail -n 2
{"epoch": 1, "global_step": 1000, "loss": 1.8423}
{"epoch": 1, "global_step": 1200, "loss": 1.7011}
4. Preempt it#
Submit a run at the default priority (0) that needs every GPU of the machine. It cannot start while the fine-tune holds one GPU, and the fine-tune's priority is lower, so the fine-tune is stopped for it.
Runs → New run → Edit as JSON, paste urgent.json from the API
tab, Start run. (The form's Arguments field splits on spaces
and does not understand quotes, so a Python one-liner with spaces needs
the JSON.)
$ astra astraeus run --name urgent-eval --image pytorch/pytorch:2.4.1-cuda12.4-cudnn9-runtime \
--gpus 8 --machines 1 --cpus 4 --mem 16G --time 20m \
-- python -c "import torch,time;print(torch.cuda.device_count(),'GPUs');time.sleep(300)"
urgent-eval submitted to main: https://console.astralyx.cloud/o/acme/w/vision/jobs/main/urgent-eval
With GPUs, --cpus and --mem are per GPU: 32 cores and 128 GiB in
all. The run gets the default priority, 0.
{
"metadata": {"name": "urgent-eval"},
"spec": {
"start": "Independent",
"task_template": {
"image": "pytorch/pytorch:2.4.1-cuda12.4-cudnn9-runtime",
"command": "python",
"args": ["-c", "import torch,time;print(torch.cuda.device_count(),'GPUs');time.sleep(300)"],
"restart_policy": "Never",
"time_limit_seconds": 1200,
"requested_resources": {
"cpu_cores": 32, "memory_bytes": 137438953472,
"gpu_requests": {"count": 8}, "machines": 1,
"node_selection": {"mode": "Any"}
}
}
}
}
Within a few seconds the fine-tune is stopped, saves, and waits:
$ astra astraeus logs finetune-r50-0 | tail -n 3
{"epoch": 2, "global_step": 1800, "loss": 1.4120}
SIGTERM: saving a checkpoint after this step
checkpoint at global step 1843 written in 2.7s; exiting
$ astra astraeus runs
NAME CLUSTER STATE REASON
urgent-eval main Running 1 of 1 workers running
finetune-r50 main Pending Preempted by higher-priority work (urgent-eval); waiting to run again
The fine-tune's page shows the same reason, and its History has the stop. The workspace's Events record it too.
5. Watch it resume#
When urgent-eval ends (here after 5 minutes), the GPU is free again and
the fine-tune is placed. On gpu-01 or on another machine, it finds the
checkpoint on the drive:
$ astra astraeus logs finetune-r50-0 | tail -n 2
resuming at epoch 2, step 280 (global step 1843)
{"epoch": 2, "global_step": 2000, "loss": 1.3988}
To verify, check that global step continues from the number the
preempted attempt printed, and that the run's worker shows no restarts:
a preemption is not a failure.
6. Clean up#
$ astra astraeus delete finetune-r50
finetune-r50 deleted
$ astra astraeus delete urgent-eval
urgent-eval deleted
$ curl -fsS -X DELETE "$API/drives/ft-ckpt" -H "Authorization: Bearer $ASTRAEUS_TOKEN"
Deleting the drive leaves /data/astraeus/ft-ckpt on gpu-01 as it is;
remove the folder on the machine if you no longer need it.
Variations#
Higher priority with a grant. If an organisation admin sets your
workspace's Max priority to 50, submit urgent work at "priority": 50
and ordinary work at 0. Asking for more than the maximum is refused with
403 Forbidden and priority 60 is above namespace ws-3f9c2a1b7d4e's
maximum of 50. astra astraeus run has no priority flag: set
"priority" in the specification (console or API).
A run that must not be preempted. Nothing exempts a run: preemption follows priority alone. Give the run a priority no other work in the cluster can exceed, or reserve machines for it (Reservations).
Backfill. A run with a time limit (time_limit_seconds) can start in
room the queue holds for a larger run, if it will end before that run is
due. A fine-tune of unknown length should not set one: a worker stopped at
its time limit fails and is not restarted.
Several GPUs. With torchrun --standalone --nproc-per-node=$ASTRAEUS_GPU_COUNT,
torchrun receives SIGTERM and passes it to its processes; save from
rank 0 only, after a barrier, and keep the save well inside the stop grace.
A multi-machine gang is preempted whole.
Troubleshooting#
| Symptom | Cause | Fix |
|---|---|---|
| The fine-tune restarts from the beginning | The checkpoint is on a drive the next machine does not see (a drive kept on each machine), or the script does not load it. | Use a drive on one machine or a shared filesystem; check the resuming at line. |
SIGTERM never shows in the log |
The main process is a shell that does not pass the signal on. | Make the command python …, or exec python … in sh -c. |
The log stops at SIGTERM: saving… with no written line |
The save took longer than the stop grace and the container was killed. | Save less (weights and optimiser only), save faster storage, or ask the machine's administrator for a longer stop grace. The periodic checkpoint is used. |
urgent-eval waits and nothing is preempted |
Nothing lower is running where it could fit, or the cluster runs with preemption off. | Check its worker's Why it is not running yet. |