Skip to content

From a notebook to an Astraeus training run#

A notebook is the quickest place to try an idea; a run is the right place to train it: it queues, runs to the end with nobody watching, keeps its log, and does not stop when you close the tab. In this recipe you write a training script from a notebook onto the notebook's drive, try it for a few seconds there, then start a run that mounts the same drive on the same machine, trains for longer, and leaves a checkpoint the notebook reads back.

The script trains a small classifier on data it synthesises, so it needs no download.

Before you begin#

  • A machine with a GPU in your workspace, with a Data location (Drives): the notebooks drive keeps its files there.
  • The editor or admin role.
  • For the CLI tabs, astra signed in (Install the CLI). For the API tabs, ASTRAEUS_TOKEN, CONSOLE and API as in API.

Below, <machine> stands for the GPU machine's name.

Why the same machine

The notebooks drive is kept on each machine's data location: each machine has its own copy, and copies are not synchronised. A run on another machine would see that machine's copy, without your script. So the run is pinned to the machine the notebook ran on.

1. Make the notebook#

  1. Open Hesperus → Notebooks and press New notebook.
  2. Title: Prototype. Name: prototype. Leave Drive (notebooks) and File on the drive (notebooks/prototype.ipynb).
  3. Under Machine, choose <machine>. Under Image, choose PyTorch; set GPUs to 1.
  4. Press Create and open. The editor opens and connects once the runtime is ready (the first start pulls the PyTorch image, about 4 GB).

astra has no notebook commands: use the console or the API.

$ curl -fsS -X POST "$API/notebooks" -H "Authorization: Bearer $ASTRAEUS_TOKEN" \
    -H 'content-type: application/json' \
    -d '{"metadata": {"name": "prototype"}, "spec": {"title": "Prototype",
         "runtime_defaults": {"image": "pytorch", "resources": {"gpus": {"count": 1}}, "placement": {"node": "<machine>"}}}}' > /dev/null
$ curl -fsS -X POST "$API/notebooks/prototype/open" -H "Authorization: Bearer $ASTRAEUS_TOKEN" \
    | jq -r '.runtime.metadata.name + " " + .runtime.status.state'
prototype-3fa91c Pending

Then open the notebook in the console to run cells.

2. Write the script from a cell#

The drive is at /content in the runtime. Run this cell:

%%writefile /content/train.py
import argparse, os
import torch
from torch import nn

p = argparse.ArgumentParser()
p.add_argument("--epochs", type=int, default=5)
args = p.parse_args()

dev = "cuda" if torch.cuda.is_available() else "cpu"
torch.manual_seed(0)
X = torch.randn(20000, 32)
y = ((X @ torch.randn(32, 1)).squeeze() > 0).long()
X, y = X.to(dev), y.to(dev)

model = nn.Sequential(nn.Linear(32, 64), nn.ReLU(), nn.Linear(64, 2)).to(dev)
opt = torch.optim.Adam(model.parameters(), lr=1e-2)
for epoch in range(args.epochs):
    for i in range(0, len(X), 512):
        loss = nn.functional.cross_entropy(model(X[i:i + 512]), y[i:i + 512])
        opt.zero_grad()
        loss.backward()
        opt.step()
    with torch.no_grad():
        acc = (model(X).argmax(1) == y).float().mean().item()
    print(f"epoch {epoch + 1} loss {loss.item():.4f} accuracy {acc:.3f} on {dev}", flush=True)

os.makedirs("/content/checkpoints", exist_ok=True)
torch.save(model.state_dict(), "/content/checkpoints/model.pt")
print("saved /content/checkpoints/model.pt")
Writing /content/train.py

Try it for two epochs, in the next cell:

!python /content/train.py --epochs 2
epoch 1 loss 0.0961 accuracy 0.978 on cuda
epoch 2 loss 0.0601 accuracy 0.986 on cuda
saved /content/checkpoints/model.pt

The exact numbers vary. The runtime's chip at the top of the editor shows the machine it is on: that is <machine>.

3. Train it as a run#

The run uses the same image as the PyTorch preset (on NVIDIA), mounts the notebooks drive at /content, and is pinned to <machine>:

train-from-notebook.json
{
  "metadata": {"name": "train-from-notebook"},
  "spec": {
    "task_template": {
      "image": "pytorch/pytorch:2.14.1-cuda12.6-cudnn9-runtime",
      "command": "python",
      "args": ["/content/train.py", "--epochs", "50"],
      "restart_policy": "Never",
      "requested_resources": {
        "cpu_cores": 4,
        "memory_bytes": 17179869184,
        "gpu_requests": {"count": 1},
        "node_selection": {"names": ["<machine>"]}
      },
      "datavolume_refs": [
        {"name": "notebooks", "mount_path": "/content", "mode": "ReadWrite"}
      ]
    }
  }
}

The notebook's runtime holds one GPU: a machine with a single GPU must free it first. Stop the runtime (its chip → Stop the runtime (frees the GPU)) unless the machine has a second GPU.

  1. Open Runs → New run and choose the cluster.
  2. Select Edit as JSON (worker groups, privileges, anything else) and replace the text with train-from-notebook.json, with your machine's name.
  3. Select Start run.

astra astraeus run has no flag for drives or for a machine's name: start the run in the console or with the API, then follow it here:

$ astra astraeus logs train-from-notebook -f
epoch 1 loss 0.0961 accuracy 0.978 on cuda
…
epoch 50 loss 0.0083 accuracy 0.997 on cuda
saved /content/checkpoints/model.pt

$ curl -fsS -X POST "$API/jobs" -H "Authorization: Bearer $ASTRAEUS_TOKEN" \
    -H 'content-type: application/json' --data-binary @train-from-notebook.json | jq -r .metadata.name
train-from-notebook

The run goes Pending → Running → Completed. If it waits, its page's Why? says what for (the GPU still held, for example).

4. Read the checkpoint in the notebook#

Open the notebook again (Connect). Its files are on the drive, with the run's checkpoint beside them:

import torch
from torch import nn
model = nn.Sequential(nn.Linear(32, 64), nn.ReLU(), nn.Linear(64, 2))
model.load_state_dict(torch.load("/content/checkpoints/model.pt", map_location="cpu"))
print(sum(p.numel() for p in model.parameters()), "parameters loaded")
2242 parameters loaded

Going further#

  • On several machines: see Multi-machine runs. Put data and checkpoints that several machines share on a shared drive and mount it in the notebook too (Change runtime… → Drives).
  • Again and again: keep the JSON beside the notebook, or save it as a run template.