Source code for arcade.classifieroptimization.pruning

"""Iterative magnitude pruning optimizer."""

from __future__ import annotations

import copy
import logging
from typing import Any

import numpy as np

from .base import BaseOptimizer, OptimizationResult, register_optimizer

log = logging.getLogger(__name__)


def _count_params(model: Any) -> int:
    if hasattr(model, "parameters"):
        return int(sum(p.numel() for p in model.parameters()))
    return 0


[docs] @register_optimizer("pruning") class PruningOptimizer(BaseOptimizer): """Magnitude-based iterative pruning with fine-tuning.""" def __init__( self, sparsity: float = 0.5, method: str = "magnitude", iterative: bool = True, n_steps: int = 10, ) -> None: self.sparsity = sparsity self.method = method self.iterative = iterative self.n_steps = n_steps
[docs] def optimize( self, model: Any, train_data: tuple[Any, Any], val_data: tuple[Any, Any] | None = None, **kwargs: Any, ) -> OptimizationResult: import torch import torch.nn as nn import torch.nn.functional as F from torch.utils.data import DataLoader, TensorDataset train_features, train_labels = train_data if not isinstance(model, nn.Module): log.warning("Model is not a PyTorch module; skipping pruning") return OptimizationResult(model=model) device = "cuda" if torch.cuda.is_available() else "cpu" current = copy.deepcopy(model).to(device) original_params = _count_params(current) x_train = torch.tensor(np.asarray(train_features), dtype=torch.float32, device=device) y_train = torch.tensor(np.asarray(train_labels), dtype=torch.long, device=device) dl = DataLoader(TensorDataset(x_train, y_train), batch_size=128, shuffle=True) if val_data is not None: eval_x = torch.tensor(np.asarray(val_data[0]), dtype=torch.float32, device=device) eval_y = np.asarray(val_data[1]) else: eval_x = x_train eval_y = np.asarray(train_labels) current.eval() with torch.no_grad(): baseline_acc = float((current(eval_x).argmax(1).cpu().numpy() == eval_y).mean()) finetune_epochs = int(kwargs.get("finetune_epochs", 10)) lr = float(kwargs.get("lr", 1e-4)) prune_per_step = self.sparsity / self.n_steps if self.iterative else self.sparsity for step in range(self.n_steps if self.iterative else 1): frac = prune_per_step if self.iterative else self.sparsity for module in current.modules(): if isinstance(module, nn.Linear) and module.weight is not None: w = module.weight.data.abs() k = int(frac * w.numel()) if k > 0: thresh = torch.topk(w.flatten(), w.numel() - k, largest=True).values.min() mask = w >= thresh module.weight.data *= mask.float() opt = torch.optim.Adam( (p for p in current.parameters() if p.requires_grad), lr=lr, ) for _ in range(finetune_epochs): current.train() for xb, yb in dl: opt.zero_grad() loss = F.cross_entropy(current(xb), yb) loss.backward() for module in current.modules(): if isinstance(module, nn.Linear) and module.weight is not None: module.weight.grad.data[module.weight.data == 0] = 0.0 opt.step() current.eval() with torch.no_grad(): final_acc = float((current(eval_x).argmax(1).cpu().numpy() == eval_y).mean()) n_nonzero = sum( int((m.weight.data != 0).sum()) for m in current.modules() if isinstance(m, nn.Linear) and m.weight is not None ) return OptimizationResult( model=current, original_params=original_params, optimized_params=n_nonzero, accuracy_before=baseline_acc, accuracy_after=final_acc, metadata={"sparsity": self.sparsity, "method": self.method, "n_steps": self.n_steps}, )