"""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},
)