Disentanglement from Nonlinear Task Sparsity
Abstract
In this work, we study representation learning through task sparsity: learning a representation in which each task depends on only small subset of latent factors. Prior work has theoretically shown that such representations can be recovered through sparsity-aware multi-task learning. However, practical realizations have been limited to linear task heads, limiting their scalability. Extending task sparsity to nonlinear predictors is challenging because sparsity is inherently discrete and therefore difficult to optimize with gradient-based methods. To address this challenge, we propose Disentanglement from nonlinear Task Sparsity (DTS), a general sparsity-aware representation learning framework for recovering disentangled representations from nonlinear multi-task datasets. \methodname formulates task sparsity as a task-conditioned binary masking problem, imposing sparsity only through the masking operation and allowing nonlinear task heads to learn from the masked features. Empirically, we evaluate \methodname on a multi-task dataset generated following the CLEVR protocol johnson2017clevr. We show that \methodname recovers the underlying disentangled factors from nonlinear tasks and the learned representation enables strong OOD generalization through sparse downstream task learning. Finally, our analysis shows that both the masking architecture and the sparsity objective are crucial for learning disentangled representations from nonlinear tasks.
est. 32% chance this paper gets accepted at ICLR 2027.
What do you think this paper will get?
All positions stay anonymous.