Model¶
dt.Module is nn.Module with hyperparameter capture and plotting hooks on
top. How you write PyTorch models does not change.
The three things you fill in¶
class MyNet(dt.Module):
def forward(self, X): ...
def loss(self, y_hat, y): ...
def configure_optimizers(self): ...
loss and configure_optimizers raise NotImplementedError by default. You
have to supply them.
forward is the exception: assign self.net and it is delegated for you.
class MyNet(dt.Module):
def __init__(self):
super().__init__()
self.net = nn.Sequential(nn.Flatten(), nn.LazyLinear(256),
nn.ReLU(), nn.LazyLinear(10))
With neither self.net nor forward, the call site stops you.
save_hyperparameters() and call order¶
It turns every __init__ argument into an instance attribute and collects them
in self.hparams.
class MyNet(dt.Module):
def __init__(self, lr=0.01, num_hiddens=256):
super().__init__()
self.save_hyperparameters()
self.net = nn.Sequential(nn.Flatten(), nn.LazyLinear(num_hiddens),
nn.ReLU(), nn.LazyLinear(10))
model = MyNet(lr=0.1)
model.lr, model.hparams
hparams is stored in checkpoints, so a file alone tells you what settings the
model was trained with.
The order matters
super().__init__() first, save_hyperparameters() second.
The parent __init__ also overwrites hparams with its own arguments.
Reverse the order and your values are erased by the parent's defaults.
Local variables are not picked up. Only declared arguments are read.
To leave an argument out, use ignore.
@add_to_class — attaching methods across cells¶
It removes the notebook problem where changing a class means re-running its definition cell, which then forces you to re-run everything below it.
# cell 3
class MyNet(dt.Module):
def __init__(self, lr=0.01):
super().__init__()
self.save_hyperparameters()
self.net = nn.LazyLinear(10)
# cell 7 — much later
@dt.add_to_class(MyNet)
def loss(self, y_hat, y):
return F.cross_entropy(y_hat, y)
Because it lands on the class, instances you already built pick it up immediately.
The decorator returns the original function, so the name stays usable in the cell that defined it.
@dt.add_to_class(MyNet)
def loss(self, y_hat, y):
return F.cross_entropy(y_hat, y)
loss # <function loss at 0x...> — not None
Batch convention¶
The default training_step and validation_step read a batch this way:
Forward is called as self(*batch[:-1]), so a model with two inputs takes an
(X1, X2, y) batch as is.
Adding an accuracy curve¶
The default only plots loss. To see accuracy as well, override
validation_step.
@dt.add_to_class(MyNet)
def validation_step(self, batch):
y_hat = self(*batch[:-1])
loss = self.loss(y_hat, batch[-1])
self.plot('loss', loss, train=False)
self.log('acc', (y_hat.argmax(-1) == batch[-1]).float().mean())
return loss
Scalars sent to log are averaged over the current record interval into
trainer.history['acc']. That is an epoch average in epoch mode and the span
between record or validation boundaries in step mode. When a board exists, the
current train/eval state also selects a live curve. A
Trainer with log_dir writes the exact key acc to history.jsonl as well.
Using the same key in both training and validation combines all observations
into one interval average. Name them train_acc and val_acc yourself when they
must stay separate. deeptool does not interpret domain metric names.
The return value must still be the loss — Trainer uses it to fill history
and to decide the best epoch.
plot does not reach history
Use plot(key, value, train) when a curve is enough and log(key, value)
when the number must be retained. epoch, step, train_loss, val_loss, lr,
and sec are reserved for Trainer-generated fields.
The live board follows the Trainer's unit: fractional epochs in epoch mode and optimizer steps in step mode.