Skip to content
AITroveRead. Build. Understand.
Make this comfortable

Structured pruning and latency parity

Last updated: 7 Oct 20265 min read
tutorial
AdvancedBy AITrove Editorial

Removing complete hidden channels can create a smaller dense model; zeroing weights without rebuilding the graph often leaves the original compute shape.

Select a removable structure

An unstructured mask sets individual weights to zero while preserving tensor dimensions. A dense runtime may still execute the same matrix multiplication. Structured pruning removes an entire hidden channel, filter or attention head and rewrites dependent dimensions, which can reduce dense compute. The choice must respect residual branches, normalization and downstream consumers. A channel used in two branches cannot disappear from one without repairing both. For a small classifier, start with one pair of linear layers and prove exact masked-graph parity before fine-tuning.

Copy both sides of the boundary

If the first layer emits eight hidden values and only four are kept, copy those four rows of its weight and matching bias entries. Then copy the same four columns of the next layer. Preserve the next-layer bias. The two-layer example below compares this compact graph with the original graph after explicitly masking removed hidden values; exact equality is expected up to floating-point tolerance before fine-tuning. It does not claim parity with the unpruned model, because dropping useful channels changes its outputs. Logit changes must then be measured on real validation cases.

Choose channels using held-out evidence

Magnitude is one possible ranking signal, but a small weight can still contribute through a correlated or rare activation. Compare candidate prune sets on a selection set that includes weak receipt edges and low-contrast text. Refit weights after structural removal, then use a separate final test to check class recall and confidence. Track the exact channel indices and model revision so the compact graph can be reconstructed. For convolutional backbones, spatial shape and skip-connection constraints make this more than a row-slicing operation.

Benchmark the compiled path

Parameter count and multiply-add estimates are useful screening measures; wall-clock latency depends on kernel choice, memory traffic, batch size and device. An awkward channel width can be slower than a slightly larger aligned width on some hardware. Warm the model, synchronize accelerators where required, and report median and tail latency for the actual serving batch. Include preprocessing and any transfer to the device. The inference baseline should use the same input and runtime.

Combine compression carefully

Prune first or quantize first according to an explicit experiment; changing graph dimensions after calibration can invalidate activation scales. A useful sequence is prune, fine-tune, freeze, calibrate the compact model, convert and retest. Keep each candidate artifact immutable and compare it with the same holdout and device workload. If a speed gain fails a rare-defect recall gate, reduce pruning or restore the prior model. Calibration cannot repair missing capacity by itself.

Implementation

python
import torch
from torch import nn
from torch.nn import functional as functional

torch.manual_seed(47)
feature_layer = nn.Linear(6, 8)
decision_layer = nn.Linear(8, 3)
retained_channels = torch.tensor([0, 2, 5, 7])
compact_feature = nn.Linear(6, len(retained_channels))
compact_decision = nn.Linear(len(retained_channels), 3)
with torch.no_grad():
    compact_feature.weight.copy_(feature_layer.weight[retained_channels])
    compact_feature.bias.copy_(feature_layer.bias[retained_channels])
    compact_decision.weight.copy_(decision_layer.weight[:, retained_channels])
    compact_decision.bias.copy_(decision_layer.bias)

receipt_features = torch.randn(5, 6)
hidden_values = functional.relu(feature_layer(receipt_features))
channel_mask = torch.zeros(8)
channel_mask[retained_channels] = 1
masked_logits = decision_layer(hidden_values * channel_mask)
compact_logits = compact_decision(functional.relu(compact_feature(receipt_features)))
torch.testing.assert_close(masked_logits, compact_logits, atol=1e-6, rtol=1e-6)

Performance and operating cost

The dense two-layer multiply work falls from approximately 6×8 + 8×3 to 6×4 + 4×3 multiplies per example, excluding biases and activation overhead. Parameter storage falls with the same changed dimensions. Selection and copy cost are minor for one layer, but iterative pruning and fine-tuning can dominate project time. Neither multiply count nor a zero mask proves latency gain; use repeated timing on the device and include the cost of any new data movement or operator fallback.

Common Mistakes

  • Do not mistake a zero-valued weight mask for a smaller executable graph.
  • Do not compare the compact graph with the full unpruned graph when testing copy parity.
  • Do not prune one branch of a residual path without reconciling its consumers.

Read next

ai-data
deep-learning
Storage details