training.py 10 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332
  1. import pathlib as pl
  2. from typing import Callable, Tuple, cast
  3. import numpy as np
  4. import pandas as pd
  5. import torch
  6. import torch.nn as nn
  7. from torch.utils.data import DataLoader
  8. from model.dataset import ADNIDataset
  9. type TrainMetrics = Tuple[
  10. float, float, float, float
  11. ] # (train_loss, val_loss, train_acc, val_acc)
  12. type TestMetrics = Tuple[float, float] # (test_loss, test_acc)
  13. type KLLossFn = Callable[[nn.Module], torch.Tensor | None]
  14. def test_model(
  15. model: nn.Module,
  16. test_loader: DataLoader[ADNIDataset],
  17. criterion: nn.Module,
  18. progress: Callable[[int | None, int], None],
  19. write_log: Callable[[str, bool], None],
  20. ) -> TestMetrics:
  21. """
  22. Tests the model on the test dataset.
  23. Args:
  24. model (nn.Module): The model to test.
  25. test_loader (DataLoader[ADNIDataset]): DataLoader for the test dataset.
  26. criterion (nn.Module): Loss function to compute the loss.
  27. Returns:
  28. TrainMetrics: A tuple containing the test loss and test accuracy.
  29. """
  30. model.eval()
  31. test_loss = 0.0
  32. correct = 0
  33. total = 0
  34. for _, (mri, xls, targets, _) in enumerate(test_loader):
  35. outputs = model((mri, xls))
  36. loss = criterion(outputs, targets)
  37. test_loss += loss.item() * (mri.size(0) + xls.size(0))
  38. # Calculate accuracy
  39. predicted = (outputs > 0.5).float()
  40. correct += (predicted == targets).sum().item()
  41. total += targets.numel()
  42. test_loss /= len(test_loader)
  43. test_acc = correct / total if total > 0 else 0.0
  44. return test_loss, test_acc
  45. def train_epoch(
  46. model: nn.Module,
  47. train_loader: DataLoader[ADNIDataset],
  48. val_loader: DataLoader[ADNIDataset],
  49. optimizer: torch.optim.Optimizer,
  50. criterion: nn.Module,
  51. ) -> Tuple[float, float, float, float]:
  52. """
  53. Trains the model for one epoch and evaluates it on the validation set.
  54. Args:
  55. model (nn.Module): The model to train.
  56. train_loader (DataLoader[ADNIDataset]): DataLoader for the training dataset.
  57. val_loader (DataLoader[ADNIDataset]): DataLoader for the validation dataset.
  58. optimizer (torch.optim.Optimizer): Optimizer for updating model parameters.
  59. criterion (nn.Module): Loss function to compute the loss.
  60. Returns:
  61. Tuple[float, float, float, float]: A tuple containing the training loss, validation loss, training accuracy, and validation accuracy.
  62. """
  63. model.train()
  64. train_loss = 0.0
  65. # Training loop
  66. for _, (mri, xls, targets, _) in enumerate(train_loader):
  67. optimizer.zero_grad()
  68. outputs = model((mri, xls))
  69. loss = criterion(outputs, targets)
  70. loss.backward() # type: ignore[reportUnknownMemberType]
  71. optimizer.step()
  72. train_loss += loss.item() * (mri.size(0) + xls.size(0))
  73. train_loss /= len(train_loader)
  74. model.eval()
  75. val_loss = 0.0
  76. correct = 0
  77. total = 0
  78. with torch.no_grad():
  79. for _, (mri, xls, targets, _) in enumerate(val_loader):
  80. outputs = model((mri, xls))
  81. loss = criterion(outputs, targets)
  82. val_loss += loss.item() * (mri.size(0) + xls.size(0))
  83. # Calculate accuracy
  84. predicted = (outputs > 0.5).float()
  85. correct += (predicted == targets).sum().item()
  86. total += targets.numel()
  87. val_loss /= len(val_loader)
  88. val_acc = correct / total if total > 0 else 0.0
  89. train_acc = correct / total if total > 0 else 0.0
  90. return train_loss, val_loss, train_acc, val_acc
  91. def train_model(
  92. model: nn.Module,
  93. train_loader: DataLoader[ADNIDataset],
  94. val_loader: DataLoader[ADNIDataset],
  95. optimizer: torch.optim.Optimizer,
  96. criterion: nn.Module,
  97. num_epochs: int,
  98. output_path: pl.Path,
  99. ) -> Tuple[nn.Module, pd.DataFrame]:
  100. """
  101. Trains the model using the provided training and validation data loaders.
  102. Args:
  103. model (nn.Module): The model to train.
  104. train_loader (DataLoader[ADNIDataset]): DataLoader for the training dataset.
  105. val_loader (DataLoader[ADNIDataset]): DataLoader for the validation dataset.
  106. num_epochs (int): Number of epochs to train the model.
  107. learning_rate (float): Learning rate for the optimizer.
  108. Returns:
  109. Result[nn.Module, str]: A Result object containing the trained model or an error message.
  110. """
  111. # Record the training history
  112. # We record the Epoch, Training Loss, Validation Loss, Training Accuracy, and Validation Accuracy
  113. # use a (num_epochs, 4) shape ndarray to store the history before creating the DataArray
  114. nhist = np.zeros((num_epochs, 4), dtype=np.float32)
  115. for epoch in range(num_epochs):
  116. train_loss, val_loss, train_acc, val_acc = train_epoch(
  117. model,
  118. train_loader,
  119. val_loader,
  120. optimizer,
  121. criterion,
  122. )
  123. # Update the history
  124. nhist[epoch, 0] = train_loss
  125. nhist[epoch, 1] = val_loss
  126. nhist[epoch, 2] = train_acc
  127. nhist[epoch, 3] = val_acc
  128. print(
  129. f"Epoch [{epoch + 1}/{num_epochs}], "
  130. f"Train Loss: {train_loss:.4f}, Val Loss: {val_loss:.4f}, "
  131. f"Train Acc: {train_acc:.4f}, Val Acc: {val_acc:.4f}"
  132. )
  133. # If we are at 25, 50, or 75% of the epochs, save the model
  134. if num_epochs > 4:
  135. if (epoch + 1) % (num_epochs // 4) == 0:
  136. model_save_path = (
  137. output_path / "intermediate_models" / f"model_epoch_{epoch + 1}.pt"
  138. )
  139. torch.save(model.state_dict(), model_save_path)
  140. print(f"Model saved at epoch {epoch + 1}")
  141. # return the trained model and the traning history
  142. history = pd.DataFrame(
  143. data=nhist.astype(np.float32),
  144. columns=["train_loss", "val_loss", "train_acc", "val_acc"],
  145. index=np.arange(1, num_epochs + 1),
  146. )
  147. return model, history
  148. def test_model_bayesian(
  149. model: nn.Module,
  150. test_loader: DataLoader[ADNIDataset],
  151. criterion: nn.Module,
  152. get_kl_loss: KLLossFn,
  153. ) -> TestMetrics:
  154. """
  155. Tests a Bayesian model on the test dataset with KL-augmented loss.
  156. """
  157. model.eval()
  158. test_loss = 0.0
  159. correct = 0
  160. total = 0
  161. with torch.no_grad():
  162. for _, (mri, xls, targets, _) in enumerate(test_loader):
  163. outputs = model((mri, xls))
  164. data_loss = cast(torch.Tensor, criterion(outputs, targets))
  165. batch_size = mri.size(0)
  166. kl_term = get_kl_loss(model)
  167. kl_loss = (
  168. kl_term / batch_size
  169. if kl_term is not None
  170. else torch.tensor(0.0, device=outputs.device)
  171. )
  172. loss: torch.Tensor = data_loss + kl_loss
  173. test_loss += loss.item() * (mri.size(0) + xls.size(0))
  174. predicted = (outputs > 0.5).float()
  175. correct += (predicted == targets).sum().item()
  176. total += targets.numel()
  177. test_loss /= len(test_loader)
  178. test_acc = correct / total if total > 0 else 0.0
  179. return test_loss, test_acc
  180. def train_epoch_bayesian(
  181. model: nn.Module,
  182. train_loader: DataLoader[ADNIDataset],
  183. val_loader: DataLoader[ADNIDataset],
  184. optimizer: torch.optim.Optimizer,
  185. criterion: nn.Module,
  186. get_kl_loss: KLLossFn,
  187. ) -> TrainMetrics:
  188. """
  189. Trains a Bayesian model for one epoch and evaluates on validation data.
  190. """
  191. model.train()
  192. train_loss = 0.0
  193. for _, (mri, xls, targets, _) in enumerate(train_loader):
  194. optimizer.zero_grad()
  195. outputs = model((mri, xls))
  196. data_loss = cast(torch.Tensor, criterion(outputs, targets))
  197. batch_size = mri.size(0)
  198. kl_term = get_kl_loss(model)
  199. kl_loss = (
  200. kl_term / batch_size
  201. if kl_term is not None
  202. else torch.tensor(0.0, device=outputs.device)
  203. )
  204. loss: torch.Tensor = data_loss + kl_loss
  205. loss.backward() # type: ignore[reportUnknownMemberType]
  206. optimizer.step()
  207. train_loss += loss.item() * (mri.size(0) + xls.size(0))
  208. train_loss /= len(train_loader)
  209. model.eval()
  210. val_loss = 0.0
  211. correct = 0
  212. total = 0
  213. with torch.no_grad():
  214. for _, (mri, xls, targets, _) in enumerate(val_loader):
  215. outputs = model((mri, xls))
  216. data_loss = cast(torch.Tensor, criterion(outputs, targets))
  217. batch_size = mri.size(0)
  218. kl_term = get_kl_loss(model)
  219. kl_loss = (
  220. kl_term / batch_size
  221. if kl_term is not None
  222. else torch.tensor(0.0, device=outputs.device)
  223. )
  224. vloss: torch.Tensor = data_loss + kl_loss
  225. val_loss += vloss.item() * (mri.size(0) + xls.size(0))
  226. predicted = (outputs > 0.5).float()
  227. correct += (predicted == targets).sum().item()
  228. total += targets.numel()
  229. val_loss /= len(val_loader)
  230. val_acc = correct / total if total > 0 else 0.0
  231. train_acc = correct / total if total > 0 else 0.0
  232. return train_loss, val_loss, train_acc, val_acc
  233. def train_model_bayesian(
  234. model: nn.Module,
  235. train_loader: DataLoader[ADNIDataset],
  236. val_loader: DataLoader[ADNIDataset],
  237. optimizer: torch.optim.Optimizer,
  238. criterion: nn.Module,
  239. num_epochs: int,
  240. output_path: pl.Path,
  241. get_kl_loss: KLLossFn,
  242. ) -> Tuple[nn.Module, pd.DataFrame]:
  243. """
  244. Trains a Bayesian model with KL-augmented objective.
  245. """
  246. nhist = np.zeros((num_epochs, 4), dtype=np.float32)
  247. for epoch in range(num_epochs):
  248. train_loss, val_loss, train_acc, val_acc = train_epoch_bayesian(
  249. model,
  250. train_loader,
  251. val_loader,
  252. optimizer,
  253. criterion,
  254. get_kl_loss,
  255. )
  256. nhist[epoch, 0] = train_loss
  257. nhist[epoch, 1] = val_loss
  258. nhist[epoch, 2] = train_acc
  259. nhist[epoch, 3] = val_acc
  260. print(
  261. f"Epoch [{epoch + 1}/{num_epochs}], "
  262. f"Train Loss: {train_loss:.4f}, Val Loss: {val_loss:.4f}, "
  263. f"Train Acc: {train_acc:.4f}, Val Acc: {val_acc:.4f}"
  264. )
  265. if num_epochs > 4:
  266. if (epoch + 1) % (num_epochs // 4) == 0:
  267. model_save_path = (
  268. output_path / "intermediate_models" / f"model_epoch_{epoch + 1}.pt"
  269. )
  270. torch.save(model.state_dict(), model_save_path)
  271. print(f"Model saved at epoch {epoch + 1}")
  272. history = pd.DataFrame(
  273. data=nhist.astype(np.float32),
  274. columns=["train_loss", "val_loss", "train_acc", "val_acc"],
  275. index=np.arange(1, num_epochs + 1),
  276. )
  277. return model, history