- MantraMindAI
- Blog
- AI Engineering & MLOps
Edge AI in practice: memory bandwidth, quantisation, and updates
Jai Rao
August 22, 202620 min read
Why you would run a model on the device, why memory bandwidth usually sets the speed limit, how int8 and 4-bit quantisation really work, and how to ship updates.
A camera on a bottling line has to decide whether the label in front of it is crooked, and it has to decide before the bottle leaves the frame — a window of tens of milliseconds. A phone keyboard has to predict the next word between one keystroke and the next. Neither has a cloud option, and not because a cloud model would be worse. The round trip is the problem. Even on a good mobile network, a request to a nearby region and back costs somewhere in the tens of milliseconds before your model has done any work at all, and a cold connection costs several round trips on top of that. No amount of GPU budget gets you under a physical latency floor.
So the model moves onto the device. That trade usually gets labelled "edge AI" and left there, which hides the interesting part: you are swapping tens of gigabytes of high-bandwidth memory and a mains supply for a few gigabytes of shared LPDDR, a couple of watts of sustained power, and a fleet you cannot log into. What follows is what actually constrains on-device inference (rarely FLOPs), what quantisation does to your weights and your accuracy, and what it means to ship a model to a million devices you will never touch again.
Four reasons to move inference onto the device
The first is the latency floor. If a response has to land inside roughly 20 milliseconds — a wake word, a gesture, an autofocus decision, a frame-level safety gate — the round trip alone eats the budget. Local inference is not faster because the chip is better; it is faster because there is no wire.
The second is data that must not leave. Keyboard text, camera frames from inside a home, an ECG trace, a photo of a document: if the raw signal never crosses the network boundary, you are no longer storing, transmitting or deleting that data on someone's behalf, and the compliance surface shrinks accordingly.
Sending a 200-byte event ("hard hat missing, camera 4, 14:02") rather than a video stream is a smaller privacy problem and a smaller bandwidth bill at once.
The third is operation without a network: a mine, a ship, a tunnel, an aircraft cabin, a rural clinic, a warehouse with dead spots between the racks. A feature that degrades to "no connection" there does not exist there.
The fourth is per-inference cost. Cloud inference is a marginal cost on every call, forever; on-device inference is a fixed engineering cost plus a slice of the user's battery. When call volume is high and value per call is low — every video frame, every keystroke, every accelerometer window — the economics fail long before the latency does.
The honest counterweight: if call volume is modest, the latency budget is a second or more, and the data is allowed to leave, call a cloud API and move on. On-device deployment is an engineering programme with a long maintenance tail, and it should be chosen for one of those four reasons, not for aesthetics.
Where the device runs out first
Memory bandwidth, not FLOPs
Spec sheets advertise compute — a phone NPU might claim several trillion operations per second. Compute is usually not what limits you. Consider a matrix multiply at batch size 1, which is what single-user inference looks like: every weight is read from memory exactly once and used for one multiply and one add. That is two arithmetic operations per weight, and at fp32 you moved four bytes to get them. No accelerator is fed fast enough for that ratio.
Put the numbers side by side. A phone-class memory system delivers on the order of tens of gigabytes per second, shared with the display compositor, the OS and the rest of your app; a server GPU with HBM delivers hundreds of gigabytes to a couple of terabytes per second. So a 1 GB model on a phone needs roughly 20 milliseconds simply to stream its weights past the compute units once. That is a hard floor for producing one token before any arithmetic happens, capping you somewhere near 50 tokens per second in the best case — and you will not get the best case. These are order-of-magnitude figures; the ratio is the point, not the digits.
This reframes the optimisation problem. Quantisation is first and foremost a bandwidth optimisation: fp32 to int8 quarters the bytes you move, and on a bandwidth-bound layer that buys a roughly proportional speedup whether or not the chip has fast integer math. Not every workload is in that regime: a convolution over a large feature map reuses each kernel weight across many positions and can genuinely be compute-bound, so check first. Divide bytes moved by measured time and compare against rated bandwidth; do the same for operations against rated throughput. Close to bandwidth and nowhere near peak compute means you should stop tuning kernels and start shrinking weights.
A thermal budget, not a power rating
A phone has no fan, and the limit is skin temperature. The whole SoC can dissipate a few watts sustained, and burst several times that for a few seconds; frequency scaling closes the gap by dropping clocks after seconds to minutes of load. So the benchmark you ran for thirty seconds, plugged in, on a cool device on a desk is the best number your feature will ever produce. Run a ten-minute loop, unplugged, in the kind of case a user actually has, and report the last minute. Continuous accelerator load also drains a battery fast, which is why the standard pattern is duty cycling: a tiny always-on model gates a much larger one that runs only when the gate fires.
RAM, including the parts you forgot to count
The file on disk is parameters times bytes per parameter, plus quantisation scales. Resident memory is bigger: add the activation workspace, the runtime and its accelerator delegate, and — for anything transformer-shaped — a key/value cache that grows linearly with context length and can rival the weights once the context is long. iOS terminates an app that exceeds a per-device memory limit; on Android the low-memory killer takes your backgrounded process. Map the weights read-only from a file instead of allocating them on the heap, so the OS can evict pages under pressure rather than killing you. Then test on the cheapest device you support, because your dev phone has twice the RAM of the median device in the field.
What quantisation actually does to a weight
Training leaves you with fp32 weights, and their distribution is usually a narrow bell around zero. Quantisation throws away that continuous representation and keeps an integer grid plus a scale factor. For weights, the standard choice is symmetric int8: pick s = max|w| / 127, store q = round(w / s) clamped to the range -127 to 127, and recover an approximation as s * q. Eight bits per weight instead of thirty-two, and the matrix multiply becomes integer accumulation with a single scale multiply at the end.
Why does that not wreck the model? Each weight picks up a rounding error of at most half a grid step, a dot product sums thousands of such terms whose errors are roughly independent and partly cancel, and networks are over-parameterised enough to absorb small perturbations. What breaks the argument is a single outlier: one enormous weight stretches max|w| so far that every other weight in the tensor collapses onto a handful of levels. The fix is per-channel scales — one per output row rather than one for the whole tensor. It costs almost nothing in storage and is often the difference between "int8 was fine" and "int8 destroyed it".
Activations are the harder half, because their range depends on the input. You either measure it offline from calibration data and bake in a fixed scale (static quantisation), or compute it per call at runtime (dynamic quantisation), which adapts but costs a pass over the tensor. Transformer activations are known for a few outlier channels whose magnitudes sit far above the rest, and one per-tensor activation scale handles them badly. Hence the default first move for language models on device: quantise weights only, since that is where the bandwidth goes, and leave activations in floating point.
Below eight bits
Four bits gives sixteen levels: hopeless for a whole tensor, workable when the scale is local. So 4-bit formats are groupwise — split each row into contiguous groups of 32, 64 or 128 weights and store a scale, sometimes a zero point too, per group. That overhead is real. A 16-bit scale per 64 weights adds a quarter of a bit per weight, so "4-bit" on disk is nearer 4.25 to 4.5 bits, and a group size of 32 doubles the overhead while usually buying back accuracy. Methods like GPTQ and AWQ go further, using calibration activations to decide which weights matter most and adjusting the rounding of later weights to compensate for error already introduced.
Accuracy loss is best described as a shape rather than a digit. For a well-behaved vision model, int8 with per-channel weights typically costs a fraction of a percentage point of top-1 accuracy, small enough to sit inside run-to-run noise. 4-bit weight-only on a large language model moves perplexity a little, and the change is often hard to spot in casual use; the same recipe on a model of a few hundred million parameters is visible immediately, because there is far less redundancy to spend. Below four bits, quality falls off sharply. The rule of thumb worth carrying is that the bigger the model, the fewer bits it tolerates — and the thing that actually bites at 4 bits is long-form coherence and rare-token behaviour, which averaged metrics hide. So evaluate the quantised artifact on your own task, and look at per-slice numbers, because aggregate accuracy can stay flat while one rare but important class quietly collapses.
Post-training quantisation first, QAT when it earns its keep
Post-training quantisation (PTQ) converts a model that is already trained. Weight scales are deterministic — read them off the tensors. Activation scales come from pushing a few hundred representative inputs through the network and recording ranges. The range estimator matters more than most people expect: plain min/max is hostage to one outlier batch, while percentile clipping or a threshold chosen to minimise mean squared error usually wins. Calibration data must look like production data, same preprocessing and same distribution; calibrating an outdoor camera model on clean studio images is a common own goal. PTQ costs minutes, needs no labels and no training loop, and should always be the first attempt.
When PTQ loses too much, the next step is not quantisation-aware training — it is mixed precision. Find the layers that hurt and leave them alone. The usual suspects are the first and last layers, embeddings, anything with a very wide activation range, and depthwise convolutions, which have so few weights per channel that per-channel statistics are thin. A per-layer sensitivity sweep — quantise one layer at a time, measure, rank — costs an afternoon and frequently recovers most of the gap.
Quantisation-aware training (QAT) inserts fake-quantise operations into the graph so the forward pass rounds exactly the way the deployed kernel will, then fine-tunes through it. The gradient of round() is zero almost everywhere, so QAT uses a straight-through estimator: pass the gradient through as if the rounding were not there, clipped outside the representable range. The network then learns weights that survive the grid. The cost is a training pipeline, labelled data, hyperparameters and a fine-tune over a fraction of the original schedule — days of engineering, not minutes.
QAT earns that cost in four situations: you are targeting 4 bits or fewer; the model is small and has no redundancy to spare; the hardware is integer-only so activations must be quantised too, which is normal for microcontroller accelerators and some NPU paths; or PTQ plus mixed precision still leaves you outside budget. If PTQ already lands inside noise, QAT is a week spent for nothing.
Making it smaller, then making it portable
The cheapest possible experiment is PyTorch dynamic quantisation: one call, no calibration data, weights stored as int8 and activations quantised per call. It rewrites only Linear and recurrent modules, leaves everything else in fp32, and targets CPU. This snippet builds a small Linear-heavy stack, then reports state-dict size and per-call latency for both versions.
import io, time, torch, torch.nn as nnmodel = nn.Sequential(nn.Linear(1024, 4096), nn.GELU(), nn.Linear(4096, 1024)).eval()def size_mb(m): buf = io.BytesIO() torch.save(m.state_dict(), buf) return buf.getbuffer().nbytes / 1e6def bench(m, x, n=200): with torch.no_grad(): for _ in range(20): m(x) t0 = time.perf_counter() for _ in range(n): m(x) return (time.perf_counter() - t0) / n * 1e3qmodel = torch.ao.quantization.quantize_dynamic(model, {nn.Linear}, dtype=torch.qint8)x = torch.randn(1, 1024)print(f"fp32 {size_mb(model):6.1f} MB {bench(model, x):5.2f} ms")print(f"int8 {size_mb(qmodel):6.1f} MB {bench(qmodel, x):5.2f} ms")The size change is arithmetic you can predict: 8.4 million parameters at four bytes each is about 34 MB, and at one byte each about 8.5 MB plus scales. The latency change is what people get wrong — do not expect four times faster. Dynamic quantisation pays to quantise activations on every call, and at batch size 1 with small tensors you may be dominated by per-op overhead rather than bandwidth. On the feed-forward blocks of a real transformer the win is substantial; on a small stack it can be a wash or slower.
Getting off PyTorch is a separate step, because the runtime on the device is usually not PyTorch. ONNX plus ONNX Runtime is the common path for desktop, mobile and embedded Linux. Export with a fixed opset and explicit dynamic axes, then quantise the graph.
import numpy as np, torch, onnxruntime as ortfrom onnxruntime.quantization import quantize_dynamic, QuantTypetorch.onnx.export( model, torch.randn(1, 1024), "model.onnx", input_names=["input"], output_names=["logits"], dynamic_axes={"input": {0: "batch"}}, opset_version=17,)quantize_dynamic("model.onnx", "model.int8.onnx", weight_type=QuantType.QInt8)sess = ort.InferenceSession("model.int8.onnx", providers=["CPUExecutionProvider"])x = np.random.randn(1, 1024).astype(np.float32)ref = model(torch.from_numpy(x)).detach().numpy()print("max abs diff:", np.abs(sess.run(None, {"input": x})[0] - ref).max())The last two lines are the ones that save you: always compare the exported graph against the original on the same input, because a silently different operator mapping is far more common than a crash and surfaces later as an unexplained quality regression. ONNX Runtime also has a static path, quantize_static, which takes a calibration data reader and quantises activations too — that is the one you want for integer-only hardware.
For Android and integer-only accelerators, TFLite's full-integer conversion is the equivalent, and it makes calibration concrete: you hand the converter a generator that yields real inputs.
import tensorflow as tfdef representative_data(): for sample in calibration_samples[:200]: # a few hundred real production-like inputs yield [sample.astype("float32")[None, ...]]conv = tf.lite.TFLiteConverter.from_saved_model("saved_model/")conv.optimizations = [tf.lite.Optimize.DEFAULT]conv.representative_dataset = representative_dataconv.target_spec.supported_ops = [tf.lite.OpsSet.TFLITE_BUILTINS_INT8]conv.inference_input_type = tf.int8conv.inference_output_type = tf.int8open("model_int8.tflite", "wb").write(conv.convert())Setting the input and output types to int8 matters: quantisation and dequantisation now happen outside the graph, so the whole network can sit on an accelerator with no floating-point path at all. It also means your app code owns the input scaling — one more thing that must be versioned alongside the weights.
Pruning and distillation, honestly
Unstructured magnitude pruning zeroes individual small weights. It reliably shrinks the file, because sparse formats compress well, and it reliably does nothing for latency: commodity CPUs, GPUs and NPUs have no fast path for irregular sparsity, so you traverse the same tiles and do the same work, multiplying by zero. Treat it as a storage optimisation and be sceptical of anyone selling it as a speed one.
Some hardware does support a semi-structured pattern such as two non-zeros in every four, which is a genuine middle ground.
Structured pruning is the version that gets faster: remove whole channels, attention heads or layers so the remaining tensors are simply smaller and still dense. The catch is keeping shapes consistent through the graph, needing a fine-tune afterwards to recover, and paying a real accuracy cost. Dropping a fraction of the heads and trimming feed-forward width, then fine-tuning, is a well-trodden way to buy a meaningful speedup for a modest loss.
Distillation trains a small model to reproduce a large model's full output distribution — not just its argmax — over a large pool of unlabelled inputs. It tends to be the most effective of the three, because you are not compressing a fixed model but training a better small one, and the teacher's soft targets carry far more information per example than a one-hot label. It is also the most expensive: you need the teacher, the input pool, and a real training run.
Ranked by benefit per unit of effort: quantisation first, every time; distillation for the largest gains, budgeted as a project; unstructured pruning mostly for file size. In practice the shipped artifact is often distilled and then quantised, since the two compose cleanly.
Choosing a target: cloud GPU, phone NPU, microcontroller
The figures below are order-of-magnitude and vary enormously by part. Read them as ratios between rows, not as specifications.
| Target | Memory for the model | Memory bandwidth | Sustained power | Typical model size | Numeric formats |
|---|---|---|---|---|---|
| Server GPU | 16-80 GB dedicated | hundreds of GB/s to a few TB/s | 300-700 W, mains | 1 GB to hundreds of GB | fp16, bf16, fp8, int8, int4 |
| Laptop or desktop CPU / iGPU | 8-64 GB shared | tens to ~200 GB/s | 15-45 W, mains or large battery | tens of MB to a few GB | fp32, fp16, int8, int4 |
| Phone SoC (CPU / GPU / NPU) | 4-16 GB shared, app budget a fraction of it | tens of GB/s | 2-5 W for the whole SoC | a few MB to ~3 GB | fp16, int8, int4 weight-only |
| Microcontroller (Cortex-M class) | 256 KB - 2 MB SRAM, a few MB flash | well under 1 GB/s | tens of mW to ~0.5 W, coin cell to small battery | 10 KB to ~1 MB | int8 only, sometimes int4 |
The jumps are what matter. Between a phone and a microcontroller you lose roughly three orders of magnitude of memory and about the same in power. That is why microcontroller ML is a different discipline rather than a smaller version of the same one: with 256 KB of SRAM the binding constraint is often the activation working set rather than the weights, because a single intermediate feature map can exceed your entire memory. You design around the memory budget from the first layer, and int8 stops being an optimisation applied at the end and becomes the only format available.
Shipping a fix to hardware you do not control
Once a build is in the field, every version you ever released is running somewhere. App review takes hours to days, a staged rollout takes days more, a long tail of users never updates, and plenty of devices are offline for weeks. Meanwhile the model is the component most likely to need changing, because new failure modes arrive as inputs your test set never contained.
The structural answer is to decouple the model from the binary. Ship weights as a signed, versioned asset the app downloads and caches, with a server-side pointer to the current version and the ability to move that pointer backwards; a bad model then becomes a config change rather than a release. Two conditions come attached. The bundle must carry everything that has to match those weights — tokeniser or label map, normalisation constants, input shape, preprocessing parameters — because a mismatch raises no error, it just makes the model quietly worse. And the download needs signature verification, since a model file is close enough to executable input to be worth attacking.
You are also flying with less telemetry than a cloud service has, by construction: raw inputs are not allowed to leave, which was the point. What you can collect is aggregate and derived — latency percentiles by device model, confidence histograms, escalation rates, results from a small held-out evaluation set shipped inside the app and run on-device, and counts of crashes and accelerator fallbacks by chip and driver version. That last one matters more than it sounds: accelerator delegates differ by vendor and driver version, and a graph that runs entirely on the NPU on one chip may silently fall back to CPU on another, or return slightly different numbers. Keep a working CPU path, detect at runtime which path you got, and report it. Then roll to a small slice first and watch escalation rate and p99 latency, because those move before user complaints do.
The pattern that resolves most of this tension is a small local model with cloud escalation. The local model handles the common case and, importantly, emits a signal about its own uncertainty — a maximum-probability threshold, an entropy measure, or an abstain class trained in deliberately. Above the bar, answer on device: fast, private, free. Below it, escalate to a larger cloud model. You get on-device latency and privacy for the bulk of traffic, cloud-grade quality where it matters, and a knob that trades spend against quality. Set the threshold from the actual curve — escalation rate times cost per cloud call, against measured accuracy — not from a round number that looked sensible.
Two parts of the hybrid design are easy to get wrong. Escalation has to degrade sensibly with no network: queue the request, or answer locally and mark the answer provisional, but never hang. And escalation changes the privacy claim. The moment uncertain inputs travel to a server, "your data stays on your device" stops being true for some fraction of requests — and that fraction is precisely the unusual, often most sensitive inputs. Say so in the interface and the policy, or escalate a redacted or derived representation rather than the raw signal.
What to verify before the build goes out
Most on-device regressions are caught by measuring the right thing on the right hardware, and missed by measuring a convenient thing on a workstation. Before a release, confirm each of these on real devices from your support matrix, including the cheapest:
- Steady-state latency, not peak. Ten-minute loop, unplugged, device in a case. Report the final minute, at p95 and p99, not the mean.
- Cold start, measured separately. Loading and mapping weights plus warming the accelerator delegate often costs more than dozens of inferences, and it is the first thing the user experiences.
- Peak resident memory on your lowest-spec supported device, including a background-and-return cycle to check you survive memory pressure.
- Energy per hundred inferences. A feature that is fast and costs several percent of battery per hour is not shippable.
- Accuracy on the exact artifact you are shipping — post-quantisation, post-export, executing on the device — sliced by the categories you care about, not the fp32 checkpoint on a server.
- Numeric parity as a CI test. Fixed inputs, expected outputs, run through the on-device path. This is what catches preprocessing drift, the most common cause of "the model got worse after export".
- The fallback path. Run with the accelerator disabled and confirm the model still works and still passes the accuracy bar.
- The rollback. Actually revert a model version in staging and confirm devices pick up the older one.
None of that is exotic. It is the discipline of treating a model as a shipped binary artifact rather than a service you can redeploy at will. The devices in the field are indifferent to your architecture diagram; they care about bytes moved per second, watts turned into heat, and whether the version they happen to be running is the one you meant to send.