burrows99's picture
Update source code and report for 5-experiment structure
744f7db verified
Raw
History Blame Contribute Delete
6.51 kB
import logging
import torch
from pathlib import Path
from sklearn.metrics import accuracy_score, recall_score, f1_score
from tqdm import tqdm
from config import Config
from dataset import MelanomaDataLoaders
from model import MetadataMelanomaModel
from evaluate import Evaluator
from file_io_manager import FileIOManager
logger = logging.getLogger(__name__)
class Trainer:
def __init__(self):
self._device = Config.get_training_config()['device']
loaders = MelanomaDataLoaders()
self._train_loader = loaders.get_train_loader()
self._val_loader = loaders.get_val_loader()
num_metadata_features = loaders.num_metadata_features
self._model = MetadataMelanomaModel.build(num_metadata_features=num_metadata_features)
self._criterion = MetadataMelanomaModel.get_criterion()
self._optimizer = MetadataMelanomaModel.get_optimizer(self._model)
self._scheduler = MetadataMelanomaModel.get_scheduler(self._optimizer)
self._run_name = self._build_run_name()
self._io = FileIOManager.for_run(self._output_name())
self._io.save_preprocessor(loaders.preprocessor)
self._preprocessor = loaders.preprocessor
@staticmethod
def _output_name() -> str:
exp = Config.get_experiment()
if exp is not None:
return f"experiment{exp}"
return Config.MODEL_NAME
def _build_run_name(self) -> str:
cfg = Config.get_training_config()
base = (
f"{Config.MODEL_NAME}_Meta"
f"_LR{cfg['learning_rate']}_BS{cfg['batch_size']}_Ep{cfg['num_epochs']}"
)
aug = Config.get_augmentation_config()
if 'affine' in aug:
suffix = "_AffAug"
elif aug.get('random_erasing_prob', 0) > 0:
suffix = "_EraseAug"
else:
suffix = "_BasicAug"
return base + suffix
def _train_epoch(self) -> tuple:
self._model.train()
total_loss = 0.0
all_preds: list[int] = []
all_labels: list[int] = []
for images, metadata, labels in tqdm(self._train_loader, desc="Training"):
images = images.to(self._device)
metadata = metadata.to(self._device)
labels = labels.to(self._device)
self._optimizer.zero_grad()
outputs = self._model(images, metadata)
loss = self._criterion(outputs, labels.unsqueeze(1).float())
loss.backward()
self._optimizer.step()
total_loss += loss.detach().item()
preds = (torch.sigmoid(outputs.detach()) > 0.5).int()
all_preds.extend(preds.cpu().numpy().flatten())
all_labels.extend(labels.cpu().numpy().flatten())
return total_loss / len(self._train_loader), all_preds, all_labels
def _save_checkpoint(self, epoch: int) -> str:
path = self._io.save_checkpoint(self._model, self._run_name, epoch)
return str(path)
def _final_evaluation(self, best_model_path: str) -> None:
logger.info("Loading best model for final evaluation plots: %s", best_model_path)
final_model = MetadataMelanomaModel.build(num_metadata_features=self._model.num_metadata_features)
self._io.load_checkpoint(final_model, best_model_path, map_location=self._device)
final_model = final_model.to(self._device)
evaluator = Evaluator(final_model, MetadataMelanomaModel.get_criterion(), io=self._io)
_, preds, labels, probs = evaluator.evaluate(
self._val_loader, use_tta=Config.get_evaluation_config()['tta_enabled']
)
evaluator.plot_roc_curve(labels, probs)
evaluator.plot_confusion_matrix(labels, preds)
evaluator.compute_ood_stats(self._val_loader)
if not final_model._image_only:
evaluator.plot_shap(self._val_loader, self._preprocessor._feature_cols)
def train(self) -> None:
num_epochs = Config.get_training_config()['num_epochs']
best_val_f1 = 0.0
best_path = None
best_epoch = 0
for epoch in range(num_epochs):
current_lr = self._optimizer.param_groups[0]['lr']
logger.info("Epoch %d/%d | LR: %s", epoch + 1, num_epochs, current_lr)
train_loss, train_preds, train_labels = self._train_epoch()
val_loss, val_preds, val_labels, _ = Evaluator(self._model, self._criterion).evaluate(
self._val_loader, use_tta=Config.get_evaluation_config()['tta_enabled']
)
train_acc = accuracy_score(train_labels, train_preds)
train_recall = recall_score(train_labels, train_preds)
train_f1 = f1_score(train_labels, train_preds)
val_acc = accuracy_score(val_labels, val_preds)
val_recall = recall_score(val_labels, val_preds)
val_f1 = f1_score(val_labels, val_preds)
logger.info("Train: Loss=%.4f | Acc=%.4f | Recall=%.4f | F1=%.4f", train_loss, train_acc, train_recall, train_f1)
logger.info("Val: Loss=%.4f | Acc=%.4f | Recall=%.4f | F1=%.4f", val_loss, val_acc, val_recall, val_f1)
self._io.append_epoch_metrics({
"epoch": epoch + 1,
"experiment": self._output_name(),
"learning_rate": current_lr,
"train_loss": train_loss, "train_acc": train_acc,
"train_recall": train_recall, "train_f1": train_f1,
"val_loss": val_loss, "val_acc": val_acc,
"val_recall": val_recall, "val_f1": val_f1,
})
if val_f1 > best_val_f1:
best_val_f1 = val_f1
prev_path = best_path
best_epoch = epoch + 1
best_path = self._save_checkpoint(best_epoch)
if prev_path:
Path(prev_path).unlink(missing_ok=True)
logger.info("Removed previous checkpoint: %s", prev_path)
logger.info("New best model saved: %s (Val F1: %.4f)", best_path, best_val_f1)
if self._scheduler:
self._scheduler.step()
if best_path:
self._io.save_gradcam_checkpoint(self._model)
logger.info("Inference checkpoint saved: %s", self._io.gradcam_checkpoint_path())
self._final_evaluation(best_path)
else:
logger.warning("No best model saved; skipping final plots.")