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}")
def fine_tune_pytorch_nn( model_class: Type[torch.nn.modules.module.Module], model_args: dict, dataset_class: Callable, dataset_args: dict, param: dict):
 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