3.9c — Pruning from scratch : magnitude, structured, Lottery Ticket
Sous-grain Bloc A.3 #16060 : trois familles de pruning écrites en numpy pur, puis comparées a torch.nn.utils.prune (Bloc B.2 #16060) sur un MLP applique a MNIST.
Unstructured magnitude pruning : top-k% des poids par layer, one-shot (Han et al. 2015). Le geste le plus simple : pour chaque tensor de poids, garder les k% de poids de plus grande magnitude absolue, masquer les autres. La sparsite s = 1 - k/100 est appliquee a chaque layer independamment.
Structured filter pruning : pour un kernel de convolution K, calculer la L1-norm par filtre de sortie, supprimer les s% de plus petite norme (Li et al. 2017). Plus dur a atteindre en accuracy mais acceleration reelle (le filtre disparait du calcul).
Lottery Ticket Hypothesis (Frankle & Carlin 2019) : un réseau dense contient un sous-réseau sparse (mask m *) qui, remis aux poids d’initialisation (pas aux poids entraines) et re-entraine, atteint la même accuracy que le réseau dense original. Procedure : (i) entraîner dense, (ii) extraire le mask top-k%, (iii) reset les poids non-masques a l’init, (iv) re-entraîner. Le reset est la clef.
Stack : numpy pour les opérations bas niveau, torch pour le modèle et l’entraînement (sans torch.nn.utils.prune dans la boucle from-scratch). Modèle = MLP (2 couches cachees) sur MNIST subset 20k, CPU-compatible, <10 min d’exécution.
import numpy as npimport torchimport osos.environ.setdefault("CUBLAS_WORKSPACE_CONFIG", ":4096:8") # determinisme cuBLAS (#16795)# Determinisme (#16795) : la graine seule ne garantit PAS la reproductibilite# (heuristiques cuDNN, kernels non deterministes). warn_only=True au premier# passage pour inventorier les ops fautives sans faire echouer le run.torch.use_deterministic_algorithms(True, warn_only=True)torch.backends.cudnn.deterministic =Truetorch.backends.cudnn.benchmark =Falseimport torch.nn as nnimport torch.nn.functional as Ffrom torch.utils.data import DataLoader, Subsetfrom torchvision import datasets, transformsimport timeimport matplotlib.pyplot as plttorch.manual_seed(42)np.random.seed(42)device = torch.device('cpu')print(f'Device : {device}')print(f'torch : {torch.__version__}, numpy : {np.__version__}')
Device : cpu
torch : 2.8.0+cu126, numpy : 2.3.5
1. Unstructured magnitude pruning
magnitude_prune(W, sparsity) retourne un mask binaire M de même shape que W : 1.0 = poids conserve, 0.0 = poids masque. On applique le mask en multipliant le poids par M (les poids masques deviennent exactement 0, ce qui reduit la memoire du modèle — mais n’accelere pas le calcul sur CPU/GPU classique sauf avec des noyaux sparse dedies).
def magnitude_prune(weights, sparsity):"""Pruning unstructured par magnitude : masque les poids de |w| les plus petits.""" w = np.asarray(weights, dtype=float).flatten() n =len(w) k =max(1, int(round((1.0- sparsity) * n)))if k >= n:return np.ones_like(w).reshape(np.asarray(weights).shape) threshold = np.partition(np.abs(w), n - k)[n - k] mask = (np.abs(w) >= threshold).astype(float)return mask.reshape(np.asarray(weights).shape)def count_nonzero(weights):returnint(np.sum(np.abs(weights) >0))# Demonstration sur un tensor exemple (matmul 100x100)W_demo = np.random.standard_normal((100, 100))print(f'Tensor demo : shape {W_demo.shape}, |W|_0 = {count_nonzero(W_demo)}')for s in [0.0, 0.5, 0.9, 0.99]: mask = magnitude_prune(W_demo, sparsity=s)print(f' sparsite {s:.2f} : conserve {int(mask.sum() / mask.size *100)}% des poids')
Tensor demo : shape (100, 100), |W|_0 = 10000
sparsite 0.00 : conserve 100% des poids
sparsite 0.50 : conserve 50% des poids
sparsite 0.90 : conserve 10% des poids
sparsite 0.99 : conserve 1% des poids
2. Structured filter pruning (L1-norm)
Pour un kernel Conv2d K, on calcule la L1-norm par filtre de sortie et on supprime les s% de plus petite norme. Pour un MLP (qui n’a pas de filtres conv), on applique le même principe sur les colonnes du layer lineaire (les « unites cachees ») : L1-norm par colonne de W, suppression des s% plus petites colonnes.
def l1_filter_norms(kernel):"""L1-norm par filtre de sortie. Pour un Linear weight (out, in), c'est l'axis=1."""return np.abs(kernel).sum(axis=tuple(range(1, kernel.ndim)))def structured_prune_filters(kernel, sparsity):"""Pruning structured : supprime les filtres/colonnes de plus petite L1-norm. Retourne un mask booleen de shape (kernel.shape[0],).""" norms = l1_filter_norms(kernel) n =len(norms) n_keep =max(1, int(round((1.0- sparsity) * n))) threshold = np.partition(norms, n - n_keep)[n - n_keep]return norms >= threshold# Demo sur un kernel Conv2dkernel_demo = np.random.standard_normal((16, 8, 3, 3))print(f'Kernel Conv2d : {kernel_demo.shape}, L1 norms = {l1_filter_norms(kernel_demo).round(2)}')for s in [0.0, 0.25, 0.5]: keep = structured_prune_filters(kernel_demo, sparsity=s)print(f' sparsite {s:.2f} : garde {keep.sum()}/{len(keep)} filtres')# Demo sur un poids Linear (MLP)linear_demo = np.random.standard_normal((64, 128))print(f'\nLinear weight : {linear_demo.shape}')for s in [0.0, 0.25, 0.5]: keep = structured_prune_filters(linear_demo, sparsity=s)print(f' sparsite {s:.2f} : garde {keep.sum()}/{len(keep)} neurones de sortie')
Architecture : 784 → 256 → 128 → 10 (3 couches lineaires). ~235K paramètres. Entraînement sur un sous-ensemble de 20k exemples MNIST (au lieu de 60k) pour rester en <2 min par entraînement. Suffisant pour demontrer les trois mécanismes (LTH, structured, unstructured) avec une accuracy mesurable autour de 0.95.
class MLP(nn.Module):def__init__(self, input_dim=784, hidden_dims=(256, 128), num_classes=10):super().__init__()self.fc1 = nn.Linear(input_dim, hidden_dims[0])self.fc2 = nn.Linear(hidden_dims[0], hidden_dims[1])self.fc3 = nn.Linear(hidden_dims[1], num_classes)self.hidden_dims = hidden_dimsdef forward(self, x): x = x.flatten(1) x = F.relu(self.fc1(x)) x = F.relu(self.fc2(x))returnself.fc3(x)model = MLP().to(device)n_params =sum(p.numel() for p in model.parameters())print(f'MLP : {n_params:,} parametres (~{n_params/1e3:.0f}K)')
def train_epoch(model, loader, opt, masks=None): model.train() total, correct, loss_sum =0, 0, 0.0for x, y in loader: x, y = x.to(device), y.to(device) opt.zero_grad() out = model(x) loss = F.cross_entropy(out, y) loss.backward() opt.step()if masks isnotNone:with torch.no_grad():for name, param in model.named_parameters():if name in masks: param.mul_(masks[name]) total += y.size(0) correct += (out.argmax(1) == y).sum().item() loss_sum += loss.item() * y.size(0)return loss_sum / total, correct / totaldef eval_model(model, loader): model.eval() total, correct =0, 0with torch.no_grad():for x, y in loader: x, y = x.to(device), y.to(device) out = model(x) total += y.size(0) correct += (out.argmax(1) == y).sum().item()return correct / total# Sanity check : 1 epoch doit donner accuracy ~92-95% sur MNIST subsetmodel_sanity = MLP().to(device)opt = torch.optim.Adam(model_sanity.parameters(), lr=1e-3)t0 = time.time()loss, acc = train_epoch(model_sanity, train_loader, opt)print(f'Epoch 1 sanity : loss = {loss:.4f}, train acc = {acc:.4f}, temps = {time.time()-t0:.1f}s')print(f'Test acc apres 1 epoch : {eval_model(model_sanity, test_loader):.4f}')
Epoch 1 sanity : loss = 0.4583, train acc = 0.8688, temps = 4.7s
Test acc apres 1 epoch : 0.9245
4. Lottery Ticket Hypothesis (LTH)
Frankle & Carlin 2019 : un réseau dense initialise contient un sous-réseau sparse (mask) qui, re-initialise a ses poids d’initiaux et re-entraine, atteint la même accuracy que le réseau dense original.
Procedure canonique : 1. Initialiser le réseau (sauver state_dict_init). 2. Entraîner le réseau dense, obtenir state_dict_trained. 3. Extraire le maskm = top-k%(|state_dict_trained|) par layer. 4. Reset : state_dict_test = state_dict_init (les poids NON masques reprennent leur valeur d’init). 5. Appliquer le mask : state_dict_test = state_dict_test * m. 6. Re-entraîner.
Le reset a l’init (et non pas au trained state) est la clef de la decouverte.
def lth_prune_and_reset(model, sparsity, init_state):"""Applique le mask LTH : top-k% des poids du modele courant, reset des poids NON masques a init_state. Retourne {name: mask_tensor} pour reappliquer pendant le re-entrainement (Finding 1 Hermes c.1166).""" new_state = {} masks = {}for name, param in model.state_dict().items():# On ne prune que les poids 2D (lineaires)if'weight'notin name or param.dim() !=2: new_state[name] = init_state[name].clone()continue w_trained = param.detach().cpu().numpy() mask = magnitude_prune(w_trained, sparsity=sparsity) masks[name] = torch.from_numpy(mask).to(param.device)# Reset : on repart de init_state, puis on applique le mask w_init = init_state[name].cpu().numpy() w_new = w_init * mask new_state[name] = torch.from_numpy(w_new).to(param.device) model.load_state_dict(new_state)return masks
5. Entraînement comparatif : dense vs LTH vs pruning simple vs torch.nn.utils.prune
Quatre protocoles, 5 epochs chacun sur MNIST subset 20k :
Dense baseline : MLP entrainé tel quel, lr=1e-3, 5 epochs.
import torch.nn.utils.prune as prune# On repart du modele dense entrainé et on applique prune.global_unstructuredmodel_torchprune = MLP().to(device)model_torchprune.load_state_dict(model_dense.state_dict())parameters_to_prune = [ (module, 'weight') for module in model_torchprune.modules()ifisinstance(module, nn.Linear)]print(f'{len(parameters_to_prune)} modules Linear eligible au prune.')prune.global_unstructured(parameters_to_prune, pruning_method=prune.L1Unstructured, amount=SPARSITY)opt = torch.optim.Adam(model_torchprune.parameters(), lr=1e-3)t0 = time.time()for epoch inrange(N_EPOCHS): loss, train_acc = train_epoch(model_torchprune, train_loader, opt) test_acc = eval_model(model_torchprune, test_loader)print(f'torch.prune Epoch {epoch+1}/{N_EPOCHS} : loss={loss:.4f}, train={train_acc:.4f}, test={test_acc:.4f}')print(f'Temps total torch.prune : {time.time()-t0:.1f}s')torchprune_test_acc = eval_model(model_torchprune, test_loader)print(f'torch.prune final test acc = {torchprune_test_acc:.4f}')
3 modules Linear eligible au prune.
torch.prune Epoch 1/5 : loss=0.0518, train=0.9866, test=0.9704
torch.prune Epoch 2/5 : loss=0.0304, train=0.9923, test=0.9717
torch.prune Epoch 3/5 : loss=0.0201, train=0.9953, test=0.9710
torch.prune Epoch 4/5 : loss=0.0135, train=0.9972, test=0.9717
torch.prune Epoch 5/5 : loss=0.0094, train=0.9984, test=0.9716
Temps total torch.prune : 19.1s
torch.prune final test acc = 0.9716
# Tableau comparatif finalprint('='*60)print(f'Comparaison a {SPARSITY:.0%} de sparsite (5 epochs, MLP, MNIST subset 20k)')print('='*60)print(f'Dense baseline : test acc = {dense_test_acc:.4f}')print(f'LTH (one-shot, reset) : test acc = {lth_test_acc:.4f}')print(f'Pruning simple (no res.): test acc = {pruned_test_acc:.4f}')print(f'torch.nn.utils.prune : test acc = {torchprune_test_acc:.4f}')print()n_total =sum(p.numel() for p in model_dense.parameters() if p.dim() ==2)n_kept_lth =sum(int((p.abs() >0).sum().item()) for p in model_lth.parameters() if p.dim() ==2)print(f'Poids conserves LTH : {n_kept_lth}/{n_total} = {n_kept_lth/n_total:.1%}')print()print('Observations mesurees sur ce terrain leger (MLP/MNIST 5 epochs) :')print(' - LTH ('+f'{lth_test_acc:.4f}'+') < Pruning simple ('+f'{pruned_test_acc:.4f}'+') a meme sparsite effective 20 %')print(' - Inversion vs attendu Frankle & Carlin 2019 : LTH ne gagne PAS sur ce terrain leger')print(' - torch.nn.utils.prune comparable (global unstructured, pas de reset)')print(' - Dense baseline ('+f'{dense_test_acc:.4f}'+') sous le trio prune (effet MLP/MNIST leger)')
============================================================
Comparaison a 80% de sparsite (5 epochs, MLP, MNIST subset 20k)
============================================================
Dense baseline : test acc = 0.9669
LTH (one-shot, reset) : test acc = 0.9683
Pruning simple (no res.): test acc = 0.9709
torch.nn.utils.prune : test acc = 0.9716
Poids conserves LTH : 46951/234752 = 20.0%
Observations mesurees sur ce terrain leger (MLP/MNIST 5 epochs) :
- LTH (0.9683) < Pruning simple (0.9709) a meme sparsite effective 20 %
- Inversion vs attendu Frankle & Carlin 2019 : LTH ne gagne PAS sur ce terrain leger
- torch.nn.utils.prune comparable (global unstructured, pas de reset)
- Dense baseline (0.9669) sous le trio prune (effet MLP/MNIST leger)
6. Synthese
Méthode
Sparsite
Test acc
Mécanisme
Dense baseline
0%
dense_test_acc (cell 17)
Adam 5 epochs, lr=1e-3, MLP 784-256-128-10
LTH (one-shot)
80% (mesuree)
lth_test_acc (cell 17)
Entraînement → mask → reset init → re-entraînement avec mask reapplique après opt.step()
Pruning simple
80% (mesuree)
pruned_test_acc (cell 17)
Entraînement → mask (pas de reset) → re-entraînement avec mask reapplique après opt.step()
torch.nn.utils.prune
80% (mesuree)
torchprune_test_acc (cell 17)
global_unstructured L1Unstructured, fine-tune
Observations : - LTH (0.9678) < Pruning simple (0.9704) a même sparsite effective 20 % sur ce terrain leger MLP/MNIST 5 epochs. L’inversion est explicable : la garde mask reapplied apres opt.step() (Hermes c.1166-L1 ★★) montre que la mesure initiale etait 98 % (pas 80 %), donc le mécanisme LTH vs Pruning simple etait teste sur du quasi-dense vs du quasi-dense — la comparaison ne discriminait rien. Après correction, la sparsite reelle est 20 %, et sur ce terrain leger, Pruning simple gagne de 0.26 pt. L’effet LTH documente (Frankle & Carlin 2019) sur modèles plus grands (ResNet-20/CIFAR-10, #16190) ou l’ecart est attendu plus marque. - Acceleration hardware : unstructured magnitude = sparse storage uniquement (matmul dense). Structured filter pruning = acceleration reelle (le filtre/neurone disparait du calcul). - torch.nn.utils.prune est l’API standard PyTorch mais ne reproduit pas exactement le LTH (pas de reset a init). Note : torch.prune applique nativement un hook forward_pre_hook qui reapplique le mask après chaque forward — d’ou le fait qu’il ne souffrait pas du bug dans la version initiale.
Limites : - Le mask unstructured n’accelere pas le calcul sur GPU/CPU classique (sauf avec cuSPARSE). - L’effet LTH est plus marque sur des modèles plus grands et des problemes plus durs (notre MLP/MNIST est volontairement leger).
7. Exercices
Trois exercices C.1. Le notebook s’execute de bout en bout même exercices non completes.
def exercice_1_iterative_pruning(model, sparsity_target, n_iter, n_finetune_epochs, init_state):"""Exercice 1 : pruning iteratif avec fine-tuning entre chaque etape (Han et al. 2015). Algorithme : s_step = 1 - (1 - sparsity_target) ** (1/n_iter) repeter n_iter fois : appliquer mask top-(1 - s_step) sur le modele fine-tune pendant n_finetune_epochs mesurer accuracy Retourne un historique [(sparsity, test_acc), ...]. """# TODO etudiant : implementez le pruning iteratif.# Indication : s_step = 1 - (1 - sparsity_target)**(1/n_iter)# pour chaque iter : mask = magnitude_prune(current_weights, s_step),# multiplier current_weights par mask (garder une copie avant),# fine-tune pendant n_finetune_epochs.returnNone# TODO etudiantprint('Stub exercice 1 defini.')
Stub exercice 1 defini.
def exercice_2_sparsity_vs_accuracy(model_dense, sparsities, n_epochs_each):"""Exercice 2 : courbe sparsite vs accuracy en LTH (one-shot). Pour chaque sparsite, faire LTH et mesurer l'accuracy apres n_epochs_each. Retourne un dict {sparsity: test_acc}. Indication : reprendre la procedure LTH (lth_prune_and_reset + re-entrainement). """# TODO etudiant : trace la courbe sparsite vs accuracy pour LTH.returnNone# TODO etudiantprint('Stub exercice 2 defini.')
Stub exercice 2 defini.
def exercice_3_structured_pruning_mlp(model, sparsity):"""Exercice 3 : structured pruning d'un MLP via L1-norm sur les colonnes de fc1. Pour chaque couche lineaire, supprimer les `sparsity%` de colonnes de plus petite L1-norm (entree = nombre de neurones du layer precedent). Mettre les poids a 0 sur les colonnes supprimees et les lignes correspondantes dans le layer suivant. """# TODO etudiant : implementez le structured pruning du MLP.returnNone# TODO etudiantprint('Stub exercice 3 defini.')
Stub exercice 3 defini.
8. References
Han, S., Pool, J., Tran, J., & Dally, W. J. (2015). Learning both Weights and Connections for Efficient Neural Networks. NeurIPS.
Li, H., Kadav, A., Durdanovic, I., Samet, H., & Graf, H. P. (2017). Pruning Filters for Efficient ConvNets. ICLR.
Frankle, J., & Carlin, M. (2019). The Lottery Ticket Hypothesis: Finding Sparse, Trainable Neural Networks. ICLR.
Implementation : numpy pour les opérations de pruning bas niveau (magnitude mask, L1-norm), torch pour le modèle MLP et l’entraînement. Pas d’utilisation de torch.nn.utils.prune dans la boucle from-scratch (Bloc A.3) — c’est l’objet du Bloc B.2.