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