(tree: ProtoTree,
train_loader: DataLoader,
optimizer: torch.optim.Optimizer,
epoch: int,
disable_derivative_free_leaf_optim: bool,
device,
log: Log = None,
log_prefix: str = 'log_train_epochs',
progress_prefix: str = 'Train Epoch'
)
| 12 | from util.log import Log |
| 13 | |
| 14 | def train_epoch(tree: ProtoTree, |
| 15 | train_loader: DataLoader, |
| 16 | optimizer: torch.optim.Optimizer, |
| 17 | epoch: int, |
| 18 | disable_derivative_free_leaf_optim: bool, |
| 19 | device, |
| 20 | log: Log = None, |
| 21 | log_prefix: str = 'log_train_epochs', |
| 22 | progress_prefix: str = 'Train Epoch' |
| 23 | ) -> dict: |
| 24 | |
| 25 | tree = tree.to(device) |
| 26 | # Make sure the model is in eval mode |
| 27 | tree.eval() |
| 28 | # Store info about the procedure |
| 29 | train_info = dict() |
| 30 | total_loss = 0. |
| 31 | total_acc = 0. |
| 32 | # Create a log if required |
| 33 | log_loss = f'{log_prefix}_losses' |
| 34 | |
| 35 | nr_batches = float(len(train_loader)) |
| 36 | with torch.no_grad(): |
| 37 | _old_dist_params = dict() |
| 38 | for leaf in tree.leaves: |
| 39 | _old_dist_params[leaf] = leaf._dist_params.detach().clone() |
| 40 | # Optimize class distributions in leafs |
| 41 | eye = torch.eye(tree._num_classes).to(device) |
| 42 | |
| 43 | # Show progress on progress bar |
| 44 | train_iter = tqdm(enumerate(train_loader), |
| 45 | total=len(train_loader), |
| 46 | desc=progress_prefix+' %s'%epoch, |
| 47 | ncols=0) |
| 48 | # Iterate through the data set to update leaves, prototypes and network |
| 49 | for i, (xs, ys) in train_iter: |
| 50 | # Make sure the model is in train mode |
| 51 | tree.train() |
| 52 | # Reset the gradients |
| 53 | optimizer.zero_grad() |
| 54 | |
| 55 | xs, ys = xs.to(device), ys.to(device) |
| 56 | |
| 57 | # Perform a forward pass through the network |
| 58 | ys_pred, info = tree.forward(xs) |
| 59 | |
| 60 | # Learn prototypes and network with gradient descent. |
| 61 | # If disable_derivative_free_leaf_optim, leaves are optimized with gradient descent as well. |
| 62 | # Compute the loss |
| 63 | if tree._log_probabilities: |
| 64 | loss = F.nll_loss(ys_pred, ys) |
| 65 | else: |
| 66 | loss = F.nll_loss(torch.log(ys_pred), ys) |
| 67 | |
| 68 | # Compute the gradient |
| 69 | loss.backward() |
| 70 | # Update model parameters |
| 71 | optimizer.step() |
no test coverage detected