Importance-Aware Feature Sparsification for Wireless Split Learning
Authors: Bumjun Kim, Yoon Huh, Wan Choi
Organizations: Department of Electrical and Computer Engineering, Seoul National University (SNU), and the Institute of New Media and Communications, SNU, Seoul 08826, Korea
Wireless split learning (SL) reduces on-device computation by offloading upper layers to a server, yet transmitting high-dimensional intermediate features at each iteration remains a major communication bottleneck. Existing methods select features at the client side using task-agnostic criteria such as magnitude, statistics, or clustering, which increases client-side processing and often degrades accuracy under non-independent and identically distributed (non-i.i.d.) client data. We propose importance-aware class-balanced sparsification (ICS), a lightweight approach in which the server ranks feature channels using Grad-CAM-based scores obtained from the true-class logit during backpropagation. The per-class scores are aggregated into a class-balanced, label-agnostic importance vector that mitigates head-class bias under label skew, and each client reuses this vector in the next round to retain the top-N feature channels, incurring no additional client-side forward or backward passes. We further derive a non-asymptotic convergence bound that isolates the sparsification-induced error and characterizes how the sparsification ratio and mini-batch size jointly affect convergence under a fixed communication budget, and we analyze the communication and computational overhead of ICS against representative baselines. Beyond sequential CNN-based SL, we extend ICS to parallel split learning and to transformer-based models. Experiments show that ICS consistently outperforms the baselines, with larger gains under severe non-i.i.d. partitions.
Figures & tables
Paradigm
Main idea
Client-side computation
Uplink communication overhead
FL
Each client trains the full model locally, and the server aggregates the uploaded model updates.
Full forward/backward pass over the entire model at each client, i.e., O(Cd+Cs) .
Model updates are uploaded from the clients, i.e., O(K(Pd+Ps)) in total.
Sequential SL
The model is split into client-side and server-side parts, and clients communicate with the server sequentially.
Forward/backward pass only over the client-side model at each client, i.e., O(Cd) .
Intermediate features are uploaded sequentially, i.e., O(KBCUV) in total.
PSL
Multiple clients perform client-side computation and transmit intermediate features to the server in parallel. Client-side models are synchronized after the local updates.
Forward/backward pass over the client-side model at each client, i.e., O(Cd) .
Intermediate features and client-side model updates are uploaded, i.e., O(K(BCUV+Pd)) in total.
TABLE I: Comparison of FL, sequential SL and PSL.
Method
Selection criterion
Task-aware importance
No extra client-side computation
Non-i.i.d. consideration
RS [ 20 ]
Random sampling
✗
✓
✗
TS [ 17 ]
Feature magnitude
✗
✗
✗
RTS [ 23 ]
Feature magnitude + random sampling
✗
✗
✗
FedLite [ 19 ]
Feature clustering
✗
✗
✗
SplitFC [ 18 ]
Feature statistics
✗
✗
✗
Proposed ICS
Grad-CAM-based channel importance
✓
✓
✓
TABLE II: Comparison of communication-efficient split learning methods.
Symbol
Description
K
Number of clients.
T
Number of iterations.
E
Number of local steps.
B
Mini-batch size.
L
Number of output classes.
L
Per-sample loss function.
TABLE III: Notations and Descriptions.
Fig. 1: System model of proposed framework with the numbered stages .
Method
Client-side sparsification (msec/batch)
Server-side additional runtime (msec/batch)
RS
0.4133
0.0
TS
0.5099
0.0
RTS
0.5357
0.0
FedLite
31.7882
0.0
SplitFC
0.8106
0.0
ICS
0.0928
6.1702 (effectively 0.0)
TABLE IV: Runtime comparison of communication-efficient SL methods.
Fig. 2: Test accuracy of CIFAR-100 classification under i.i.d. datasets.
Fig. 3: Test accuracy of CIFAR-100 classification with Dirichlet α=0.5 .
Fig. 4: Test accuracy of CIFAR-100 classification with Dirichlet α=0.1 .
Fig. 5: Test accuracy of CIFAR-100 classification with Dirichlet α=0.05 .
Fig. 6: Test accuracy of ICS and ICS w/o CB on CIFAR-100 with Dirichlet α=0.1 and α=0.05 .
Fig. 7: Test accuracy of CIFAR-100 classification under different sparsification ratio and mini-batch size pair.
Fig. 8: Test accuracy on CIFAR-100 classification under PSL with Dirichlet α=0.5 .
Fig. 9: Test accuracy of CIFAR-100 classification with transformer-based SL.
Fig. 10: Test accuracy of CIFAR-10 classification under i.i.d. datasets.
Fig. 11: mIoU of segmentation under Oxford-IIIT Pet datasets.
Fig. 12: Pixel accuracy of segmentation under Oxford-IIIT Pet datasets.
Fig. 13: Comparison of feature saliency maps for sampled images from five different labels. For each sample, the first, second, and third columns show the original image, the Grad-CAM saliency map, and their overlay, respectively. Each row corresponds to a different sparsification method: ICS (row 1), TS (row 2), RS (row 3), RTS (row 4), SplitFC (row 5), and FedLite (row 6).