CoolFace
Apppublic

JimmyChin1998/Pytorch-Learning-File

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
engine.py193 linesDownload Raw Back to root
1"""
2Contains functions for training and testing a PyTorch model.
3"""
4import torch
5
6from tqdm.auto import tqdm
7from typing import Dict, List, Tuple
8
9def train_step(model: torch.nn.Module, 
10               dataloader: torch.utils.data.DataLoader, 
11               loss_fn: torch.nn.Module, 
12               optimizer: torch.optim.Optimizer,
13               device: torch.device) -> Tuple[float, float]:
14  """Trains a PyTorch model for a single epoch.
15
16  Turns a target PyTorch model to training mode and then
17  runs through all of the required training steps (forward
18  pass, loss calculation, optimizer step).
19
20  Args:
21    model: A PyTorch model to be trained.
22    dataloader: A DataLoader instance for the model to be trained on.
23    loss_fn: A PyTorch loss function to minimize.
24    optimizer: A PyTorch optimizer to help minimize the loss function.
25    device: A target device to compute on (e.g. "cuda" or "cpu").
26
27  Returns:
28    A tuple of training loss and training accuracy metrics.
29    In the form (train_loss, train_accuracy). For example:
30
31    (0.1112, 0.8743)
32  """
33  # Put model in train mode
34  model.train()
35
36  # Setup train loss and train accuracy values
37  train_loss, train_acc = 0, 0
38
39  # Loop through data loader data batches
40  for batch, (X, y) in enumerate(dataloader):
41      # Send data to target device
42      X, y = X.to(device), y.to(device)
43
44      # 1. Forward pass
45      y_pred = model(X)
46
47      # 2. Calculate  and accumulate loss
48      loss = loss_fn(y_pred, y)
49      train_loss += loss.item() 
50
51      # 3. Optimizer zero grad
52      optimizer.zero_grad()
53
54      # 4. Loss backward
55      loss.backward()
56
57      # 5. Optimizer step
58      optimizer.step()
59
60      # Calculate and accumulate accuracy metric across all batches
61      y_pred_class = torch.argmax(torch.softmax(y_pred, dim=1), dim=1)
62      train_acc += (y_pred_class == y).sum().item()/len(y_pred)
63
64  # Adjust metrics to get average loss and accuracy per batch 
65  train_loss = train_loss / len(dataloader)
66  train_acc = train_acc / len(dataloader)
67  return train_loss, train_acc
68
69def test_step(model: torch.nn.Module, 
70              dataloader: torch.utils.data.DataLoader, 
71              loss_fn: torch.nn.Module,
72              device: torch.device) -> Tuple[float, float]:
73  """Tests a PyTorch model for a single epoch.
74
75  Turns a target PyTorch model to "eval" mode and then performs
76  a forward pass on a testing dataset.
77
78  Args:
79    model: A PyTorch model to be tested.
80    dataloader: A DataLoader instance for the model to be tested on.
81    loss_fn: A PyTorch loss function to calculate loss on the test data.
82    device: A target device to compute on (e.g. "cuda" or "cpu").
83
84  Returns:
85    A tuple of testing loss and testing accuracy metrics.
86    In the form (test_loss, test_accuracy). For example:
87
88    (0.0223, 0.8985)
89  """
90  # Put model in eval mode
91  model.eval() 
92
93  # Setup test loss and test accuracy values
94  test_loss, test_acc = 0, 0
95
96  # Turn on inference context manager
97  with torch.inference_mode():
98      # Loop through DataLoader batches
99      for batch, (X, y) in enumerate(dataloader):
100          # Send data to target device
101          X, y = X.to(device), y.to(device)
102
103          # 1. Forward pass
104          test_pred_logits = model(X)
105
106          # 2. Calculate and accumulate loss
107          loss = loss_fn(test_pred_logits, y)
108          test_loss += loss.item()
109
110          # Calculate and accumulate accuracy
111          test_pred_labels = test_pred_logits.argmax(dim=1)
112          test_acc += ((test_pred_labels == y).sum().item()/len(test_pred_labels))
113
114  # Adjust metrics to get average loss and accuracy per batch 
115  test_loss = test_loss / len(dataloader)
116  test_acc = test_acc / len(dataloader)
117  return test_loss, test_acc
118
119def train(model: torch.nn.Module, 
120          train_dataloader: torch.utils.data.DataLoader, 
121          test_dataloader: torch.utils.data.DataLoader, 
122          optimizer: torch.optim.Optimizer,
123          loss_fn: torch.nn.Module,
124          epochs: int,
125          device: torch.device) -> Dict[str, List]:
126  """Trains and tests a PyTorch model.
127
128  Passes a target PyTorch models through train_step() and test_step()
129  functions for a number of epochs, training and testing the model
130  in the same epoch loop.
131
132  Calculates, prints and stores evaluation metrics throughout.
133
134  Args:
135    model: A PyTorch model to be trained and tested.
136    train_dataloader: A DataLoader instance for the model to be trained on.
137    test_dataloader: A DataLoader instance for the model to be tested on.
138    optimizer: A PyTorch optimizer to help minimize the loss function.
139    loss_fn: A PyTorch loss function to calculate loss on both datasets.
140    epochs: An integer indicating how many epochs to train for.
141    device: A target device to compute on (e.g. "cuda" or "cpu").
142
143  Returns:
144    A dictionary of training and testing loss as well as training and
145    testing accuracy metrics. Each metric has a value in a list for 
146    each epoch.
147    In the form: {train_loss: [...],
148                  train_acc: [...],
149                  test_loss: [...],
150                  test_acc: [...]} 
151    For example if training for epochs=2: 
152                 {train_loss: [2.0616, 1.0537],
153                  train_acc: [0.3945, 0.3945],
154                  test_loss: [1.2641, 1.5706],
155                  test_acc: [0.3400, 0.2973]} 
156  """
157  # Create empty results dictionary
158  results = {"train_loss": [],
159      "train_acc": [],
160      "test_loss": [],
161      "test_acc": []
162  }
163
164  # Loop through training and testing steps for a number of epochs
165  for epoch in tqdm(range(epochs)):
166      train_loss, train_acc = train_step(model=model,
167                                          dataloader=train_dataloader,
168                                          loss_fn=loss_fn,
169                                          optimizer=optimizer,
170                                          device=device)
171      test_loss, test_acc = test_step(model=model,
172          dataloader=test_dataloader,
173          loss_fn=loss_fn,
174          device=device)
175
176      # Print out what's happening
177      print(
178          f"Epoch: {epoch+1} | "
179          f"train_loss: {train_loss:.4f} | "
180          f"train_acc: {train_acc:.4f} | "
181          f"test_loss: {test_loss:.4f} | "
182          f"test_acc: {test_acc:.4f}"
183      )
184
185      # Update results dictionary
186      results["train_loss"].append(train_loss)
187      results["train_acc"].append(train_acc)
188      results["test_loss"].append(test_loss)
189      results["test_acc"].append(test_acc)
190
191  # Return the filled results at the end of the epochs
192  return results
193