Optimizing ML Models (Pruning) for Mobile Devices
We often encounter situations where a trained model doesn't fit into a smartphone's memory or runs too slowly. Pruning is one of the key methods in our arsenal to solve this problem. It's not just about removing redundant weights — it's a delicate process that requires understanding the architecture and target device. Below, we break down how we perform turnkey pruning, what results we guarantee, and why structured pruning is the #1 choice for mobile applications. If your model doesn't fit device constraints, contact us and we'll help.
Pruning removes part of the weights or neurons from a model. The logic: in a neural network trained on real data, a significant portion of weights are close to zero and barely affect the output. They can be zeroed or removed without substantial accuracy loss, gaining speed and size benefits.
Sounds attractive. In practice, pruning is more complex than quantization, requires fine-tuning after trimming, and doesn't always yield expected speedups on mobile devices due to implementation specifics. Our years of experience show there's no universal recipe. So we approach the task systematically: first analyze the model, then choose the optimal strategy.
Which Pruning for Mobile Apps?
Unstructured pruning — zeroing individual weights (sparse matrices). A matrix with 90% zeros seems like a 10× saving. But GPUs/NPUs work with dense matrices — sparse computations don't accelerate there. Practical benefit: reduced model size after compression (zeros compress well). But not inference speed on ordinary devices.
Structured pruning — removing entire filters (channels) in convolutional layers or heads in attention. The result is a physically smaller graph that actually runs faster on any hardware. This is what genuinely matters for mobile.
| Criterion | Unstructured pruning | Structured pruning |
|---|---|---|
| Size reduction | Significant (compression) | Moderate (channel removal) |
| Speedup on CPU/GPU | Minimal | Proportional to removed channels |
| Implementation complexity | Low | Medium (requires layer synchronization) |
| Requires fine-tuning | Yes | Yes |
| Mobile device support | Limited (few sparse libraries) | Good (any framework) |
Why Structured Pruning Is More Effective
Structured pruning physically reduces the computation graph. On mobile devices, this yields real inference speedup because it doesn't require specialized sparse processors. We use L1-norm to rank filters and remove the least significant ones. Example implementation in PyTorch:
import torch
import torch.nn.utils.prune as prune
# L1-based structured pruning: remove 30% filters from Conv2d layers
# by minimum L1-norm criterion (least important filters)
for name, module in model.named_modules():
if isinstance(module, torch.nn.Conv2d):
prune.ln_structured(
module,
name='weight',
amount=0.3, # 30% channels
n=1, # L1 norm
dim=0 # dim=0 — output filters
)
# After pruning — make weights permanent (remove mask)
for name, module in model.named_modules():
if isinstance(module, torch.nn.Conv2d):
prune.remove(module, 'weight')
After this, the model contains zero filters, but they are still in the graph. The next step is actual removal of zero channels:
# Custom function to remove zero filters
def remove_zero_filters(conv_layer, next_layer=None):
"""Remove filters with zero weights and synchronize next layer"""
weight = conv_layer.weight.data
# Mask: filters with non-zero weights
nonzero_mask = weight.abs().sum(dim=(1,2,3)) > 1e-6
conv_layer.weight = nn.Parameter(weight[nonzero_mask])
if conv_layer.bias is not None:
conv_layer.bias = nn.Parameter(conv_layer.bias.data[nonzero_mask])
conv_layer.out_channels = nonzero_mask.sum().item()
# Synchronize next layer (input channels)
if next_layer is not None and isinstance(next_layer, nn.Conv2d):
next_layer.weight = nn.Parameter(next_layer.weight.data[:, nonzero_mask])
next_layer.in_channels = nonzero_mask.sum().item()
This must be done carefully — BatchNorm layers after Conv also have per-channel parameters and require synchronization.
Fine-tuning After Pruning
After removing 20–40% of filters, the model loses accuracy. Fine-tuning on training data is mandatory. Rule: the more aggressive the pruning, the longer the fine-tuning.
# Fine-tuning after pruning — typically 10-20% of original epochs
optimizer = torch.optim.Adam(pruned_model.parameters(), lr=1e-4) # lower LR
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=20)
for epoch in range(20):
train_one_epoch(pruned_model, train_loader, optimizer)
val_acc = evaluate(pruned_model, val_loader)
scheduler.step()
print(f"Epoch {epoch}: val_acc={val_acc:.4f}")
Iterative pruning — cycle prune → fine-tune → prune — yields better results than a single large removal.
Lottery Ticket Hypothesis: Deeper
For tasks where results are critical, we use the Lottery Ticket approach: train the full network, find "winning tickets" — sparse subnetworks that can be trained from scratch to comparable accuracy. Implementation using the torch_pruning library:
import torch_pruning as tp
# Analyze dependencies between layers
example_inputs = torch.zeros(1, 3, 224, 224)
DG = tp.DependencyGraph()
DG.build_dependency(model, example_inputs=example_inputs)
# Get groups of connected layers (pruning one requires pruning connected)
pruner = tp.pruner.MagnitudePruner(
model,
example_inputs,
importance=tp.importance.MagnitudeImportance(p=1),
pruning_ratio=0.5, # remove 50% channels
global_pruning=False,
iterative_steps=5 # iteratively over 5 steps
)
Why Pruning Doesn't Always Give Speedup
MobileNetV3 is already optimized: depthwise separable convolutions with few channels. Removing 30% filters from a 16-channel layer leaves 11 channels — speed difference is minimal; tensor operation overhead remains.
Pruning works well on large models: ResNet-50, EfficientNet-B4, BERT. On compact models like MobileNet/EfficientNet-lite, the effect is lower. In such cases, it's better to start with a lighter base architecture rather than prune a heavy one.
Combination with Quantization
Pruning + quantization is a standard two-step optimization:
- Structured pruning 30–40% → fine-tuning → reduce graph
- INT8 quantization of the compressed graph → final model
Example result: EfficientNet-B0 (20 MB FP32, 80 ms Android) → pruning 35% + INT8 → 4 MB, 18 ms. Top-1 accuracy dropped from 77.1% to 75.8%.
| Model | Size | Inference time | Top-1 accuracy |
|---|---|---|---|
| Original (FP32) | 20 MB | 80 ms | 77.1% |
| After pruning 35% | 13 MB | 52 ms | 76.5% |
| After pruning + INT8 | 4 MB | 18 ms | 75.8% |
If your model requires such improvements, we are ready to perform the full optimization cycle. Contact us to discuss your project.
How We Conduct Turnkey Pruning
- Model analysis — determine architecture, profile latency and size.
- Pruning strategy selection — structured or lottery ticket, removal percentage.
- Iterative pruning + fine-tuning — 3–5 iterations with accuracy monitoring.
- Testing on target devices — measurements on real smartphones.
- Optional: quantization — INT8 or FP16 for additional compression.
- Documentation and deployment — we provide a report and the final model.
Example libraries used
- PyTorch (torch.nn.utils.prune, torch_pruning)
- TensorFlow Lite (for quantization)
- ONNX Runtime (for cross-platform inference)
- Core ML Tools (for iOS)
What's Included
- Full optimization cycle from analysis to deployment.
- Structured pruning with fine-tuning.
- Testing on customer devices (iOS/Android).
- Documentation on architecture changes and integration instructions.
- 30-day support after delivery.
Our Experience and Guarantees
Our specialists have years of experience optimizing neural networks for mobile devices. We have successfully pruned 50+ projects, including apps with millions of users. We guarantee accuracy retention within 2% of the original, provided fine-tuning recommendations are followed.
We will evaluate your project for free — just contact us. Get a consultation on the optimal pruning method for your model. Leave a request, and we will analyze your model for free.
Pruning (artificial neural network) — Wikipedia torch.nn.utils.prune — PyTorch documentation







