"""Same data and starting parameters as NumPy; only the training machinery changes.""" import torch from torch.utils.data import TensorDataset, DataLoader from common import arguments, settings, load_data, print_scores from network_numpy import initialize from network_torch import build_model, copy_parameters, objective, scores # region one_step def train_step(model, optimizer, X_batch, y_batch, kind): optimizer.zero_grad(set_to_none=True) logits = model(X_batch) loss = objective(logits, y_batch, kind) loss.backward() optimizer.step() return loss.item() # endregion one_step def main(): args = arguments() cfg = settings(args.config) torch.set_num_threads(cfg['threads']) X_np, y_np = load_data(args.config, cfg, 'train') X, y = torch.from_numpy(X_np), torch.from_numpy(y_np) B = cfg['batch_size'] if B > len(X): raise ValueError('Batch is larger than the training set') model = build_model(cfg) copy_parameters(model, initialize(cfg)) optimizer = torch.optim.SGD(model.parameters(), lr=cfg['learning_rate']) if args.mode == 'overfit': xb, yb = X[:B], y[:B] print_scores('Fixed batch BEFORE', *scores(model, xb, yb, cfg['loss']), len(xb)) model.train() for _ in range(cfg['overfit_steps']): train_step(model, optimizer, xb, yb, cfg['loss']) print_scores('Fixed batch AFTER', *scores(model, xb, yb, cfg['loss']), len(xb)) print('This checks memorization of one batch, not generalization.') return print_scores('Training BEFORE', *scores(model, X, y, cfg['loss']), len(X)) # region batches dataset = TensorDataset(X, y) generator = torch.Generator().manual_seed(cfg['shuffle_seed']) loader = DataLoader(dataset, batch_size=B, shuffle=cfg['shuffle'], generator=generator) model.train() updates = 0 for epoch in range(cfg['epochs']): for X_batch, y_batch in loader: train_step(model, optimizer, X_batch, y_batch, cfg['loss']) updates += 1 # endregion batches print(f'Completed {cfg["epochs"]} epochs and {updates} updates.') print_scores('Training AFTER', *scores(model, X, y, cfg['loss']), len(X)) X_test, y_test = load_data(args.config, cfg, 'test') print_scores('Held-out, no updates', *scores(model, torch.from_numpy(X_test), torch.from_numpy(y_test), cfg['loss']), len(X_test)) print('Small teaching subset; this is not a full-MNIST benchmark. Weights are not saved.') if __name__ == '__main__': main()