Re-basin via implicit Sinkhorn differentiation
Fidel A. Guerrero-PeñaHeitor Rapela MedeirosThomas DubailMasih AminbeidokhtiEric GrangerMarco Pedersoli
Introduces a fully differentiable neural network re-basin framework using implicit Sinkhorn differentiation, enabling gradient-based optimization of weight permutations for model merging and continual learning via linear mode connectivity.
Deep neural networks are difficult to analyze and merge because their optimization landscapes contain many symmetric solutions. While models trained from different starting points often achieve functionally equivalent performance, combining them directly causes severe performance degradation due to high loss barriers between them. Aligning internal network components via permutation symmetries—known as re-basin—can eliminate these barriers and enable direct model fusion. However, existing alignment methods rely on discrete, non-differentiable algorithms that optimize layers greedily, frequently getting stuck in suboptimal configurations and remaining difficult to adapt to custom training objectives.
The article evaluates a differentiable re-basin framework using implicit Sinkhorn differentiation to find optimal component permutations across neural network layers. The objective is to demonstrate that this differentiable alignment directly enables gradient-based optimization for arbitrary loss functions, lowers interpolation barriers between independently trained models, and mitigates catastrophic forgetting in continual learning.
To evaluate this framework, the authors conducted empirical comparisons across regression and image classification benchmarks, including polynomial approximation tasks, MNIST, CIFAR-10, and Split CIFAR-100. The approach relaxes discrete permutation matrices into continuous doubly stochastic matrices, optimized using standard gradient descent paired with efficient implicit differentiation. The authors tested this method across three settings: reconstructing exact permutations on synthetic benchmarks, minimizing loss barriers along linear paths connecting two trained models, and continually learning successive tasks by fusing past and incoming model parameters with a small replay buffer.
The experiments produced several key findings. First, the proposed differentiable formulation achieved zero alignment error across all tested network depths and initializations, whereas existing greedy weight-matching algorithms frequently settled in local minima with noticeable alignment discrepancies. Second, when connecting independently trained models, the method reduced loss barriers and path areas significantly better than prior techniques; for instance, on CIFAR-10, the data-driven random-point loss reduced the connection barrier to 0.06 compared to 0.15 for the prior Straight-Through Estimator. Third, applying this re-basin approach to continual learning on Rotated MNIST yielded an accuracy of 78.14% with a low forgetting rate of 0.12, outperforming established baselines such as Average Gradient Episodic Memory at 68.47% accuracy and 0.28 forgetting.
These results demonstrate that neural networks can be merged and adapted flexibly across diverse tasks using standard gradient-based pipelines without requiring complex, hand-crafted assignment solvers. This capability provides a practical foundation for lowering compute and communication costs in distributed systems, ensembling, and lifelong machine learning systems by allowing models to be unified in weight space while preserving prior knowledge.
Organizations aiming to implement continual model updates or model merging should adopt differentiable Sinkhorn alignment as a modular alternative to discrete assignment algorithms. Practitioners should tune the stability-plasticity hyperparameter to strike the desired balance between retaining historical performance and rapidly incorporating new data domains. Where memory resources are constrained, developers must account for the higher memory consumption of Sinkhorn iterations compared to traditional Hungarian-based solvers. Additional testing on larger production architectures and specialized real-world domains remains recommended before wide enterprise deployment.
- Paper: Sinkhorn Distances: Lightspeed Computation of Optimal Transportation Distances, Marco Cuturi (2013). Introduces entropic regularization and the Sinkhorn-Knopp matrix scaling algorithm, which form the foundational mathematical basis for the differentiable optimal transport used in the Sinkhorn re-basin network.
- Paper: Computational Optimal Transport, Gabriel Peyré et al. (2018). Provides a comprehensive theoretical and algorithmic treatment of computational optimal transport, including regularized assignment problems critical to model permutation and re-basin methods.
- Paper: Averaging Weights Leads to Wider Optima and Better Generalization, Pavel Izmailov et al. (2018). Establishes foundational insights into loss landscapes, weight averaging, and mode connectivity in deep neural networks that motivate re-basin alignment across solution basins.
- Paper: Riemannian Walk for Incremental Learning: Understanding Forgetting and Intransigence, Arslan Chaudhry et al. (2018). Presents key concepts in incremental and continual learning benchmarks that the source paper builds on to demonstrate the utility of linear mode connectivity via re-basin.
- Paper: Overcoming catastrophic forgetting in neural networks, James Kirkpatrick et al. (2017). Introduces the core challenge of catastrophic forgetting in continual learning, providing the baseline context for parameter preservation evaluated in the source paper.
- Paper: Equivariant Architectures for Learning in Deep Weight Spaces, Aviv Navon et al. (2023). Extends the study of permutation symmetries and functional equivalence across neural network weight spaces by designing neural architectures equivariant to weight permutations.
