Split Learning
Instead of putting an entire massive neural network on a small edge device, chop the network in half: the device runs the first few layers and lets a powerful server run the rest.
Why Does This Exist?
In Federated Learning, each client must download and store the entire global neural network. For large modern architectures—such as 7-billion parameter language models, large vision transformers, or multi-modal backbones—a single model requires 14 to 28 gigabytes of memory just to store the weights. Edge clients, such as mobile smartphones, smart sensors, and hospital IoT monitors, simply lack the RAM, thermal headroom, and GPU compute to train full architectures locally.
Furthermore, commercial providers are often unwilling to distribute their proprietary model weights directly to end-user hardware, where weights can be decompiled, stolen, or reverse-engineered.
Split Learning (SL) solves this dual constraint by partitioning the neural network across an execution boundary called the cut layer. The resource-constrained client trains only the first few layers (the client subnetwork), while a cloud server hosts the remaining heavy layers (the server subnetwork). Raw data never leaves the client, and the full model architecture is never exposed to the edge.
Think of It Like This
An assembly line passing semi-finished metal through a security wall
Imagine an artisan making handcrafted security keys. The shape of the key blank depends on confidential internal company measurements that must never leave the facility. However, milling the hardened titanium casing requires an industrial smelting press that only a massive commercial forge possesses.
Instead of shipping the confidential measurements to the commercial forge or buying a multi-million-dollar press for the facility, the artisan takes raw titanium, cuts the basic confidential groove into the blank, and slides the semi-finished key through a locked wall slot to the forge.
The commercial forge applies the heavy heat treatment and stamping, tests the final lock fit, and passes back inspection measurements through the slot. Neither party exposes their private assets: the company keeps its confidential dimensions on-site, and the forge keeps its proprietary industrial press behind locked doors.
How It Actually Works
Cut Layer Partitioning and Distributed Backpropagation
A neural network with parameters is split into two sequential subnetworks at layer :
- Client Subnetwork : Contains layers parameterized by .
- Server Subnetwork : Contains layers parameterized by .
The forward and backward passes execute cooperatively across the network boundary:
-
Client Forward Pass: The client loads a local batch of private raw data and computes intermediate hidden representations called smashed data :
-
Transmission to Server: The client transmits tensor (along with ground-truth labels , unless using U-shaped split learning) over the network to the server.
-
Server Forward & Loss Computation: The server feeds into its subnetwork to produce terminal predictions :
The server evaluates task loss and initiates the backward pass.
-
Server Backward Pass: The server computes gradients with respect to its own parameters and updates via optimizer step:
Simultaneously, it computes the activation gradient at the cut layer using the vector-Jacobian product:
-
Transmission to Client: The server transmits gradient tensor back across the network to the client.
-
Client Backward Pass: The client completes the backward pass through using the chain rule:
and updates local parameters .
In U-Shaped Split Learning, the model is cut twice: the client holds both the first layers and the final output head. This allows the client to keep both the raw features and sensitive labels strictly local, transmitting only intermediate activations into and out of the server's central hidden trunk.
Worked Example
Trace a split linear network where client holds and server holds .
- Private input: .
- Target: .
- Client weights: .
- Server weights: .
-
Client Forward Pass:
Client sends smashed vector to server.
-
Server Forward Pass & Loss:
Squared error loss .
-
Server Backward Pass:
Gradient w.r.t server weights:
Cut-layer gradient to send to client:
Server transmits back to client.
-
Client Backward Pass: Client receives and completes gradient update:
Both client and server have computed exact gradients without exchanging raw inputs or full network parameters .
Code
import torchimport torch.nn as nn
class ClientSubnet(nn.Module): """Client holds input layer and early feature extractor.""" def __init__(self) -> None: super().__init__() self.fc_client = nn.Linear(2, 2, bias=False) self.fc_client.weight.data = torch.tensor([[1.0, 0.5], [0.0, 1.0]])
def forward(self, x: torch.Tensor) -> torch.Tensor: # Returns smashed data return self.fc_client(x)
class ServerSubnet(nn.Module): """Server holds deep layers and loss evaluation.""" def __init__(self) -> None: super().__init__() self.fc_server = nn.Linear(2, 1, bias=False) self.fc_server.weight.data = torch.tensor([[2.0, 1.0]])
def forward(self, smashed_activations: torch.Tensor) -> torch.Tensor: return self.fc_server(smashed_activations)
# Initialize partitioned networksclient_net = ClientSubnet()server_net = ServerSubnet()client_opt = torch.optim.SGD(client_net.parameters(), lr=0.1)server_opt = torch.optim.SGD(server_net.parameters(), lr=0.1)
# Input data on client devicex_private = torch.tensor([[1.0, 2.0]], requires_grad=False)y_target = torch.tensor([[5.0]])
# 1. Client Forward Passsmashed_data = client_net(x_private)# Detach with grad tracking to simulate network transmissionsmashed_transmitted = smashed_data.detach().clone().requires_grad_(True)
# 2. Server Forward Pass & Losspred = server_net(smashed_transmitted)loss = 0.5 * torch.sum((pred - y_target) ** 2)
# 3. Server Backward Passserver_opt.zero_grad()loss.backward()server_opt.step()
# Extract cut-layer gradient to send backcut_gradient = smashed_transmitted.grad.clone()
# 4. Client Backward Passclient_opt.zero_grad()smashed_data.backward(cut_gradient)client_opt.step()
print(f"Prediction: {pred.item():.2f}, Loss: {loss.item():.4f}")print("Updated Server Weight:", server_net.fc_server.weight.data)print("Updated Client Weight:", client_net.fc_client.weight.data)# -> Prediction: 6.00, Loss: 0.5000# -> Updated Server Weight: tensor([[1.8000, 0.8000]])# -> Updated Client Weight: tensor([[0.8000, 0.4000], [-0.1000, 0.8000]])Watch Out For
Feature inversion attacks reconstructing raw data from smashed activations
If the cut layer is positioned too early in the network (e.g., after only 1 convolutional layer), the smashed activations retain high spatial correlation with the raw input. An untrusted server can train an inverse decoder to reconstruct the original input images with near-photographic fidelity.
To prevent activation leakage, place the cut layer deep enough into the network where features are semantic rather than spatial. Add noise via local Differential Privacy, or insert an adversarial training objective (such as NoPeek) that penalizes distance correlation between and during client training.
The Quick Version
- Divides deep neural architectures at a cut layer between an edge client and a central server.
- Clients run early layers on private data, transmitting only intermediate activations (smashed data) across the network.
- Servers compute deep layers and loss, returning cut-layer gradients so the client can complete backpropagation locally.