Skip to content

2. The model

models/my_net.py → cfg.model_name, cfg.model_params

class Model(nn.Module):
    def forward(self, x, mask=None) -> dict     # parameter names are prep_batch's keys

cfg.model_params is splatted into Model(**params), and prep_batch's model_input is splatted into forward, so name the parameters after the keys you send, and a mistyped key raises TypeError instead of going quiet. Forward returns a dict: that is what lets a loss or a metric select one output by name.

examples/synthetic/models/mlp.py
import torch.nn as nn


class Model(nn.Module):
    def __init__(self, in_dim=64, hidden=128, n_classes=10):
        super().__init__()
        self.net = nn.Sequential(nn.Linear(in_dim, hidden), nn.ReLU(), nn.Linear(hidden, n_classes))

    def forward(self, x):
        return {"logits": self.net(x)}

The dict keys are an interface

Whatever you return here is what the loss and metrics select by name. Rename one and get_value raises ConfigError: could not resolve 'logits'; available: ['logit']. It names the keys it found, so read that list before suspecting the model.

Templates: model