garmentiq.classification.fine_tune_pytorch_nn
Fine-tuning a pretrained classification model on new data.
Freezes the backbone and retrains only the layers matching unfreeze_patterns, so a
model can be adapted to a new catalogue with far less data than full training.
1"""Fine-tuning a pretrained classification model on new data. 2 3Freezes the backbone and retrains only the layers matching `unfreeze_patterns`, so a 4model can be adapted to a new catalogue with far less data than full training. 5""" 6import torch 7import torch.nn as nn 8from torch.utils.data import DataLoader 9from typing import Callable, Type 10from tqdm.auto import tqdm 11import os 12from sklearn.model_selection import StratifiedKFold 13from garmentiq.utils.device import empty_cache 14from garmentiq.classification.utils import ( 15 CachedDataset, 16 seed_worker, 17 train_epoch, 18 validate_epoch, 19 save_best_model, 20 validate_train_param, 21 validate_test_param, 22) 23 24def fine_tune_pytorch_nn( 25 model_class: Type[torch.nn.Module], 26 model_args: dict, 27 dataset_class: Callable, 28 dataset_args: dict, 29 param: dict, 30): 31 """ 32 Fine-tunes a pretrained PyTorch model using k-fold cross-validation, early stopping, and checkpointing. 33 34 This function loads pretrained weights, optionally freezes specified layers, and trains the model on a new dataset 35 while preserving original learned features. It performs stratified k-fold CV, monitors validation loss, and saves 36 the best performing model. 37 38 Args: 39 model_class (Type[torch.nn.Module]): Class of the PyTorch model (inherits from `torch.nn.Module`). 40 model_args (dict): Arguments for model initialization. 41 dataset_class (Callable): Callable that returns a Dataset given indices and cached tensors. 42 dataset_args (dict): Dict containing: 43 - 'metadata_df': DataFrame for stratification 44 - 'raw_labels': Labels array for KFold 45 - 'cached_images': Tensor of images 46 - 'cached_labels': Tensor of labels 47 param (dict): Training configuration dict. Must include: 48 - 'pretrained_path' (str): Path to pretrained weights (.pt) 49 - 'freeze_layers' (bool): Whether to freeze base layers 50 - 'optimizer_class', 'optimizer_args' 51 - optional: 'device', 'n_fold', 'n_epoch', 'patience', 52 'batch_size', 'model_save_dir', 'seed', 53 'seed_worker', 'max_workers', 'pin_memory', 54 'persistent_workers', 'best_model_name' 55 56 The optional 'device' entry accepts either a string or a `torch.device`, e.g. `"cpu"`, 57 `"cuda"`, `"cuda:0"`, or `"mps"`. Hardware acceleration is opt-in; pass it explicitly 58 to use a GPU or Apple Silicon. Default is `"cpu"`. 59 60 Raises: 61 ValueError: If required keys are missing, or if the requested 'device' is invalid or 62 unavailable on this machine. 63 Returns: None 64 """ 65 # Validate parameters 66 validate_train_param(param) 67 os.makedirs(param.get("model_save_dir", "./models"), exist_ok=True) 68 overall_best_loss = float("inf") 69 best_model_path = os.path.join(param["model_save_dir"], param["best_model_name"]) 70 71 # Stratified KFold 72 kfold = StratifiedKFold( 73 n_splits=param.get("n_fold", 5), shuffle=True, random_state=param.get("seed", 88) 74 ) 75 76 for fold, (train_idx, val_idx) in enumerate( 77 kfold.split(dataset_args["metadata_df"], dataset_args["raw_labels"]) 78 ): 79 print(f"\nFold {fold+1}/{param.get('n_fold',5)}") 80 81 # Prepare data loaders 82 train_dataset = dataset_class( 83 train_idx, dataset_args["cached_images"], dataset_args["cached_labels"] 84 ) 85 val_dataset = dataset_class( 86 val_idx, dataset_args["cached_images"], dataset_args["cached_labels"] 87 ) 88 89 g = torch.Generator() 90 g.manual_seed(param.get("seed", 88)) 91 92 train_loader = DataLoader( 93 train_dataset, 94 batch_size=param.get("batch_size", 64), 95 shuffle=True, 96 num_workers=param.get("max_workers", 1), 97 worker_init_fn=param.get("seed_worker", seed_worker), 98 generator=g, 99 pin_memory=param.get("pin_memory", True), 100 persistent_workers=param.get("persistent_workers", False), 101 ) 102 val_loader = DataLoader( 103 val_dataset, 104 batch_size=param.get("batch_size", 64), 105 shuffle=False, 106 num_workers=param.get("max_workers", 1), 107 worker_init_fn=param.get("seed_worker", seed_worker), 108 generator=g, 109 pin_memory=param.get("pin_memory", True), 110 persistent_workers=param.get("persistent_workers", False), 111 ) 112 113 # Initialize model and load pretrained weights 114 device = param["device"] 115 model = model_class(**model_args).to(device) 116 117 # Load pretrained weights 118 state_dict = torch.load(param["pretrained_path"], map_location=device) 119 cleaned = {k.replace("module.", ""): v for k, v in state_dict.items()} 120 model.load_state_dict(cleaned, strict=False) 121 122 # Freeze base layers if requested 123 if param.get("freeze_layers", False): 124 for name, p in model.named_parameters(): 125 if not any(x in name for x in param.get("unfreeze_patterns", [])): 126 p.requires_grad = False 127 128 # DataParallel if multiple GPUs 129 if device.type == "cuda" and torch.cuda.device_count() > 1: 130 model = nn.DataParallel(model) 131 132 optimizer = param["optimizer_class"]( 133 filter(lambda p: p.requires_grad, model.parameters()), 134 **param["optimizer_args"] 135 ) 136 empty_cache(device) 137 138 best_fold_loss = float("inf") 139 patience_counter = 0 140 epoch_pbar = tqdm(range(param.get("n_epoch", 100)), desc="Epoch", leave=False) 141 142 # Training loop 143 for epoch in epoch_pbar: 144 train_loss = train_epoch(model, train_loader, optimizer, param) 145 val_loss, val_f1, val_acc = validate_epoch(model, val_loader, param) 146 147 best_fold_loss, patience_counter, overall_best_loss = save_best_model( 148 model, val_loss, best_fold_loss, patience_counter, 149 overall_best_loss, param, fold, best_model_path 150 ) 151 152 epoch_pbar.set_postfix({ 153 'train_loss': f"{train_loss:.4f}", 154 'val_loss': f"{val_loss:.4f}", 155 'val_acc': f"{val_acc:.4f}", 156 'val_f1': f"{val_f1:.4f}", 157 'patience': patience_counter, 158 }) 159 160 print(f"Fold {fold+1} | Epoch {epoch+1} | Val Loss: {val_loss:.4f} | Acc: {val_acc:.4f} | F1: {val_f1:.4f}") 161 if patience_counter >= param.get("patience", 5): 162 print(f"Early stopping at epoch {epoch+1}") 163 break 164 165 empty_cache(param["device"]) 166 print(f"\nFine-tuning completed. Best model saved at: {best_model_path}")
25def fine_tune_pytorch_nn( 26 model_class: Type[torch.nn.Module], 27 model_args: dict, 28 dataset_class: Callable, 29 dataset_args: dict, 30 param: dict, 31): 32 """ 33 Fine-tunes a pretrained PyTorch model using k-fold cross-validation, early stopping, and checkpointing. 34 35 This function loads pretrained weights, optionally freezes specified layers, and trains the model on a new dataset 36 while preserving original learned features. It performs stratified k-fold CV, monitors validation loss, and saves 37 the best performing model. 38 39 Args: 40 model_class (Type[torch.nn.Module]): Class of the PyTorch model (inherits from `torch.nn.Module`). 41 model_args (dict): Arguments for model initialization. 42 dataset_class (Callable): Callable that returns a Dataset given indices and cached tensors. 43 dataset_args (dict): Dict containing: 44 - 'metadata_df': DataFrame for stratification 45 - 'raw_labels': Labels array for KFold 46 - 'cached_images': Tensor of images 47 - 'cached_labels': Tensor of labels 48 param (dict): Training configuration dict. Must include: 49 - 'pretrained_path' (str): Path to pretrained weights (.pt) 50 - 'freeze_layers' (bool): Whether to freeze base layers 51 - 'optimizer_class', 'optimizer_args' 52 - optional: 'device', 'n_fold', 'n_epoch', 'patience', 53 'batch_size', 'model_save_dir', 'seed', 54 'seed_worker', 'max_workers', 'pin_memory', 55 'persistent_workers', 'best_model_name' 56 57 The optional 'device' entry accepts either a string or a `torch.device`, e.g. `"cpu"`, 58 `"cuda"`, `"cuda:0"`, or `"mps"`. Hardware acceleration is opt-in; pass it explicitly 59 to use a GPU or Apple Silicon. Default is `"cpu"`. 60 61 Raises: 62 ValueError: If required keys are missing, or if the requested 'device' is invalid or 63 unavailable on this machine. 64 Returns: None 65 """ 66 # Validate parameters 67 validate_train_param(param) 68 os.makedirs(param.get("model_save_dir", "./models"), exist_ok=True) 69 overall_best_loss = float("inf") 70 best_model_path = os.path.join(param["model_save_dir"], param["best_model_name"]) 71 72 # Stratified KFold 73 kfold = StratifiedKFold( 74 n_splits=param.get("n_fold", 5), shuffle=True, random_state=param.get("seed", 88) 75 ) 76 77 for fold, (train_idx, val_idx) in enumerate( 78 kfold.split(dataset_args["metadata_df"], dataset_args["raw_labels"]) 79 ): 80 print(f"\nFold {fold+1}/{param.get('n_fold',5)}") 81 82 # Prepare data loaders 83 train_dataset = dataset_class( 84 train_idx, dataset_args["cached_images"], dataset_args["cached_labels"] 85 ) 86 val_dataset = dataset_class( 87 val_idx, dataset_args["cached_images"], dataset_args["cached_labels"] 88 ) 89 90 g = torch.Generator() 91 g.manual_seed(param.get("seed", 88)) 92 93 train_loader = DataLoader( 94 train_dataset, 95 batch_size=param.get("batch_size", 64), 96 shuffle=True, 97 num_workers=param.get("max_workers", 1), 98 worker_init_fn=param.get("seed_worker", seed_worker), 99 generator=g, 100 pin_memory=param.get("pin_memory", True), 101 persistent_workers=param.get("persistent_workers", False), 102 ) 103 val_loader = DataLoader( 104 val_dataset, 105 batch_size=param.get("batch_size", 64), 106 shuffle=False, 107 num_workers=param.get("max_workers", 1), 108 worker_init_fn=param.get("seed_worker", seed_worker), 109 generator=g, 110 pin_memory=param.get("pin_memory", True), 111 persistent_workers=param.get("persistent_workers", False), 112 ) 113 114 # Initialize model and load pretrained weights 115 device = param["device"] 116 model = model_class(**model_args).to(device) 117 118 # Load pretrained weights 119 state_dict = torch.load(param["pretrained_path"], map_location=device) 120 cleaned = {k.replace("module.", ""): v for k, v in state_dict.items()} 121 model.load_state_dict(cleaned, strict=False) 122 123 # Freeze base layers if requested 124 if param.get("freeze_layers", False): 125 for name, p in model.named_parameters(): 126 if not any(x in name for x in param.get("unfreeze_patterns", [])): 127 p.requires_grad = False 128 129 # DataParallel if multiple GPUs 130 if device.type == "cuda" and torch.cuda.device_count() > 1: 131 model = nn.DataParallel(model) 132 133 optimizer = param["optimizer_class"]( 134 filter(lambda p: p.requires_grad, model.parameters()), 135 **param["optimizer_args"] 136 ) 137 empty_cache(device) 138 139 best_fold_loss = float("inf") 140 patience_counter = 0 141 epoch_pbar = tqdm(range(param.get("n_epoch", 100)), desc="Epoch", leave=False) 142 143 # Training loop 144 for epoch in epoch_pbar: 145 train_loss = train_epoch(model, train_loader, optimizer, param) 146 val_loss, val_f1, val_acc = validate_epoch(model, val_loader, param) 147 148 best_fold_loss, patience_counter, overall_best_loss = save_best_model( 149 model, val_loss, best_fold_loss, patience_counter, 150 overall_best_loss, param, fold, best_model_path 151 ) 152 153 epoch_pbar.set_postfix({ 154 'train_loss': f"{train_loss:.4f}", 155 'val_loss': f"{val_loss:.4f}", 156 'val_acc': f"{val_acc:.4f}", 157 'val_f1': f"{val_f1:.4f}", 158 'patience': patience_counter, 159 }) 160 161 print(f"Fold {fold+1} | Epoch {epoch+1} | Val Loss: {val_loss:.4f} | Acc: {val_acc:.4f} | F1: {val_f1:.4f}") 162 if patience_counter >= param.get("patience", 5): 163 print(f"Early stopping at epoch {epoch+1}") 164 break 165 166 empty_cache(param["device"]) 167 print(f"\nFine-tuning completed. Best model saved at: {best_model_path}")
Fine-tunes a pretrained PyTorch model using k-fold cross-validation, early stopping, and checkpointing.
This function loads pretrained weights, optionally freezes specified layers, and trains the model on a new dataset while preserving original learned features. It performs stratified k-fold CV, monitors validation loss, and saves the best performing model.
Arguments:
- model_class (Type[torch.nn.Module]): Class of the PyTorch model (inherits from
torch.nn.Module). - model_args (dict): Arguments for model initialization.
- dataset_class (Callable): Callable that returns a Dataset given indices and cached tensors.
- dataset_args (dict): Dict containing:
- 'metadata_df': DataFrame for stratification
- 'raw_labels': Labels array for KFold
- 'cached_images': Tensor of images
- 'cached_labels': Tensor of labels
param (dict): Training configuration dict. Must include:
- 'pretrained_path' (str): Path to pretrained weights (.pt)
- 'freeze_layers' (bool): Whether to freeze base layers
- 'optimizer_class', 'optimizer_args'
- optional: 'device', 'n_fold', 'n_epoch', 'patience', 'batch_size', 'model_save_dir', 'seed', 'seed_worker', 'max_workers', 'pin_memory', 'persistent_workers', 'best_model_name'
The optional 'device' entry accepts either a string or a
torch.device, e.g."cpu","cuda","cuda:0", or"mps". Hardware acceleration is opt-in; pass it explicitly to use a GPU or Apple Silicon. Default is"cpu".
Raises:
- ValueError: If required keys are missing, or if the requested 'device' is invalid or unavailable on this machine.
- Returns: None