Attention head partitioning and inter-layer speculation for distributed LLM inference at the edge.
Reference implementation of the method presented in:
Dimitrios Kafetzis and Iordanis Koutsopoulos, "HeadSplit: Attention Head Partitioning and Inter-Layer Speculation for Distributed Inference at the Edge", IEEE International Workshop on Signal Processing and Artificial Intelligence for Wireless Communications (SPAWC), 2026.
Even a small LLM exceeds the memory of a single edge device, so the model must be split across several devices connected over a wireless network. The standard approach assigns each layer of the model to a different device; because layers execute one after the other, at any moment a single device computes while all the others sit idle.
HeadSplit exploits this idle compute capacity through two coupled ideas:
-
Head-level partitioning. Each layer's attention heads are independent and each has its own key-value cache, so they can be separated into independently schedulable units and distributed across devices. All devices then execute every layer's attention in parallel. This moves the bottleneck to the feed-forward network (FFN) block, which still runs on a single device while the others wait.
-
Inter-layer speculation. After attention and projection complete but before the FFN executes, an intermediate output Z of the layer is available that closely approximates the final layer output X (the FFN's contribution enters through a residual connection). HeadSplit forwards Z early, so the next layer's attention runs in parallel with the ongoing FFN. An acceptance test compares Z with X once the FFN finishes: if the relative Frobenius-norm error is below a tolerance (5% by default) the speculative attention is kept; otherwise it is re-executed with the exact input, so no unbounded approximation error is ever introduced. The mechanism pipelines consecutive layers within a single token and, unlike token-level speculative decoding, needs no auxiliary draft model.
The two techniques are coupled: head partitioning creates the idle interval that speculation fills, and speculation changes the placement target, since a speculating layer's attention only needs to finish within the previous layer's FFN time.
┌─────────────┐
│ controller │ embeddings, LM head, placement,
│ │ speculation decisions, token loop
└──────┬───────┘
broadcast ┌──────┼──────────┐ PyTorch RPC over WiFi
┌───────────▼──┐ ┌─▼──────────┐ ┌▼─────────────┐
│ worker1 │ │ worker2 │ │ worker3 … │
│ heads, FFN, │ │ heads, │ │ heads, │
│ projection │ │ projection │ │ FFN blocks │
└──────────────┘ └────────────┘ └──────────────┘
└── gather: head outputs ──┘
at the projection device
- Offline profiling (
scripts/profile_offline.py): exact forward passes over WikiText-103 calibration data record, at every layer transition, whether the acceptance test would pass; the per-layer acceptance probability α is the fraction of passing sequences. - Step 1 — speculation decision: before generation, a layer is enabled for speculation when its profiled α reaches α_min (0.7 by default) and the expected time saving,
min(T_FFN_prev, T_attn) - (1 - α) · T_attn, is positive, evaluated on a preliminary placement (heads uniform, FFN on the median-capacity device). - Step 2 — greedy placement: FFN and projection blocks (the largest components) go first, each to the least-loaded device with enough free memory; then each attention head goes to the device minimizing the layer's attention time, and for a speculating layer the head assignment must keep the attention within the previous layer's FFN time. The controller re-runs the placement whenever a device's memory usage crosses 85% of capacity as the KV caches grow; heads that move take their caches with them.
- Runtime: workers exchange activations via PyTorch RPC with retries for wireless jitter. Head outputs are gathered at the layer's projection device and the result is broadcast to the next layer's devices. Round-trip times are measured at startup and increased by a 15% margin before the placement uses them. For a speculating layer the controller forwards Z as soon as projection completes, the next layer's workers start attention on it (staging their KV entries), and when the exact X arrives the controller evaluates the acceptance test and commits or discards the staged state.
headsplit/
├── config.py configuration objects (YAML-backed)
├── model.py GQA-aware per-head decomposition of Llama-family
│ models; schedulable head / projection / FFN blocks;
│ partial safetensors loading so a worker never holds
│ the full model
├── speculation.py the acceptance test
├── pipeline.py single-process assembly (reference + profiling)
├── profiling.py offline alpha profiling and block-time benchmarks
├── placement.py the two-step placement and speculation algorithm
└── runtime/
├── comm.py RPC setup, retries, RTT measurement with margin
├── worker.py block hosting, gather, stage/commit/discard
└── controller.py orchestration, speculation protocol, re-placement
scripts/
├── profile_offline.py offline profiling step
├── run_worker.py start one worker device
├── run_controller.py start the controller and run generation
└── run_local_sim.py full protocol on one machine (no testbed needed)
tests/ numerical equivalence and algorithm properties
The decomposed model is verified against the HuggingFace reference forward pass: prefill, incremental decoding with KV caches, and greedy generation are numerically identical, and with a tolerance of zero the speculative path is bit-for-bit equal to exact execution (tests/).
git clone https://github.com/Dimitrios-Kafetzis/headsplit.git
cd headsplit
python3 -m venv .venv && source .venv/bin/activate
pip install -r requirements.txt # CPU builds of torch are enoughOn Raspberry Pi OS, install the CPU wheel of PyTorch first (for example from the piwheels index) and then the remaining requirements.
The whole protocol — head partitioning, gather and broadcast, speculation with commit/discard, placement re-runs — can be exercised on one machine with local worker processes:
python scripts/run_local_sim.py --model tiny --workers 4 --tokens 16--model tiny uses a small random Llama with grouped-query attention (no download). The simulation checks that every mode reproduces the exact reference output and that forced speculation behaves correctly (all rejections must leave the output bit-for-bit identical to exact execution). Latency numbers in this mode are not meaningful, since all workers share one CPU.
Hardware used in the paper: five Raspberry Pi 4 devices (4 GB RAM), 802.11ac WiFi, one controller and up to four workers. Any set of Linux devices reachable over TCP works.
-
Profile the model offline (once, on any machine with enough memory for the full model; copy the JSON to the controller):
python scripts/profile_offline.py \ --model TinyLlama/TinyLlama-1.1B-Chat-v1.0 --out artifacts/profile.json -
Configure the devices in
config.yaml(names, addresses, memory budgets, speculation parameters). -
Start the workers, one per device. Each worker loads only the blocks the placement assigns to it, through partial safetensors reads, so the full model never needs to fit on one device:
# on worker device k (k = 1..4), with the controller's address: python scripts/run_worker.py --name worker{k} --rank {k} \ --world-size 5 --master-addr <controller-ip>
-
Run the controller:
python scripts/run_controller.py --config config.yaml \ --profile artifacts/profile.json --tokens 256 --mode all
--mode all runs the three configurations compared in the paper — layer-based partitioning (each layer on one device), head-level partitioning without speculation, and full HeadSplit — on the same prompt, and prints per-token latency, payload volume, speculation acceptance counts, and the speedups over the layer-based baseline.
| Parameter | Default | Meaning |
|---|---|---|
speculation.epsilon |
0.05 | acceptance-test tolerance (relative Frobenius error) |
speculation.alpha_min |
0.70 | minimum profiled acceptance probability for a layer to speculate |
memory_replace_threshold |
0.85 | device memory fraction that triggers a placement re-run |
rtt_margin |
1.15 | safety margin applied to measured round-trip times |
The speculation decisions follow from the profiled α of the chosen model and the measured block times of the devices; layers whose α falls below alpha_min are disabled by the speculation-decision step.
On the five-device Raspberry Pi 4 testbed with TinyLlama-1.1B (22 layers, 32 heads, WikiText-103 prompts, 256 generated tokens), the paper reports that HeadSplit reduces end-to-end inference latency by 1.54 to 1.67 times compared to layer-based partitioning, with 6% lower total energy and a 47% better energy-delay product; head-level partitioning alone contributes a 1.24 to 1.31 times speedup, and speculation contributes the rest. See the paper for the measurement methodology and the full discussion.
pip install pytest
pytest tests/ -v@inproceedings{kafetzis2026headsplit,
author = {Kafetzis, Dimitrios and Koutsopoulos, Iordanis},
title = {{HeadSplit}: Attention Head Partitioning and Inter-Layer
Speculation for Distributed Inference at the Edge},
booktitle = {IEEE International Workshop on Signal Processing and
Artificial Intelligence for Wireless Communications (SPAWC)},
year = {2026}
}MIT — see LICENSE.
Dimitrios Kafetzis — Department of Informatics, Athens University of Economics and Business (kafetzis@aueb.gr)
This research was carried out in the framework of the H.F.R.I. project "Towards advancing the Mathematical and Computational Foundations for Digital Twins of Wireless Ad Hoc Networks" (Project Number 23767).