MCPcopy Create free account
hub / github.com/M-Nauta/ProtoTree / train_epoch

Function train_epoch

prototree/train.py:14–107  ·  view source on GitHub ↗
(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'
                )

Source from the content-addressed store, hash-verified

12from util.log import Log
13
14def 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()

Callers 1

run_treeFunction · 0.90

Calls 3

distributionMethod · 0.80
log_valuesMethod · 0.80
forwardMethod · 0.45

Tested by

no test coverage detected