SparseMAP: Differentiable Sparse Structured Inference

Vlad Niculae, André F. T. Martins, Mathieu Blondel, Claire Cardie

Introduction

Structured prediction involves the manipulation of discrete, combinatorial structures, e.g., trees and alignments (Bakır et al., 2007; Smith, 2011; Nowozin et al., 2014). Such structures arise naturally as machine learning outputs, and as intermediate representations in deep pipelines. However, the set of possible structures is typically prohibitively large. As such, inference is a core challenge, often sidestepped by greedy search, factorization assumptions, or continuous relaxations (Belanger & McCallum, 2016).

In this paper, we propose an appealing alternative: a new inference strategy, dubbed SparseMAP⁡\operatorname{\mathsf{SparseMAP}}, which encourages sparsity in the structured representations. Namely, we seek solutions explicitly expressed as a combination of a small, enumerable set of global structures.

Our framework departs from the two most common inference strategies in structured prediction: maximum a posteriori (MAP) inference, which returns the highest-scoring structure, and marginal inference, which yields a dense probability distribution over structures. Neither of these strategies is fully satisfactory: for latent structure models, marginal inference is appealing, since it can represent uncertainty and, unlike MAP inference, it is continuous and differentiable, hence amenable for use in structured hidden layers in neural networks (Kim et al., 2017). It has, however, several limitations. First, there are useful problems for which MAP is tractable, but marginal inference is not, e.g., linear assignment (Valiant, 1979; Taskar, 2004). Even when marginal inference is available, case-by-case derivation of the backward pass is needed, sometimes producing fairly complicated algorithms, e.g., second-order expectation semirings (Li & Eisner, 2009). Finally, marginal inference is dense: it assigns nonzero probabilities to all structures and cannot completely rule out irrelevant ones. This can be statistically and computationally wasteful, as well as qualitatively harder to interpret.

In this work, we make the following contributions:

We propose SparseMAP⁡\operatorname{\mathsf{SparseMAP}}: a new framework for sparse structured inference (§3.1). The main idea is illustrated in Figure 1. SparseMAP⁡\operatorname{\mathsf{SparseMAP}} is a twofold generalization: first, as a structured extension of the sparsemax⁡\operatorname{\mathsf{sparsemax}} transformation (Martins & Astudillo, 2016); second, as a continuous yet sparse relaxation of MAP inference. MAP yields a single structure and marginal inference yields a dense distribution over all structures. In contrast, the SparseMAP⁡\operatorname{\mathsf{SparseMAP}} solutions are sparse combinations of a small number of often-overlapping structures.

We show how to compute SparseMAP⁡\operatorname{\mathsf{SparseMAP}} effectively, requiring only a MAP solver as a subroutine (§3.2), by exploiting the problem’s sparsity and quadratic curvature. Noticeably, the MAP oracle can be any arbitrary solver, e.g., the Hungarian algorithm for linear assignment, which permits tackling problems for which marginal inference is intractable.

We derive expressions for gradient backpropagation through SparseMAP⁡\operatorname{\mathsf{SparseMAP}} inference, which, unlike MAP, is differentiable almost everywhere (§3.3). The backward pass is fully general (applicable to any type of structure), and it is efficient, thanks to the sparsity of the solutions and to reusing quantities computed in the forward pass.

We introduce a novel SparseMAP⁡\operatorname{\mathsf{SparseMAP}} loss for structured prediction, placing it into a family of loss functions which generalizes the CRF and structured SVM losses (§4). Inheriting the desirable properties of SparseMAP⁡\operatorname{\mathsf{SparseMAP}} inference, the SparseMAP⁡\operatorname{\mathsf{SparseMAP}} loss and its gradients can be computed efficiently, provided access to MAP inference.

Our experiments demonstrate that SparseMAP⁡\operatorname{\mathsf{SparseMAP}} is useful both for predicting structured outputs, as well as for learning latent structured representations. On dependency parsing (§5.1), structured output networks trained with the SparseMAP⁡\operatorname{\mathsf{SparseMAP}} loss yield more accurate models with sparse, interpretable predictions, adapting to the ambiguity (or lack thereof) of test examples. On natural language inference (§5.2), we learn latent structured alignments, obtaining good predictive performance, as well as useful natural visualizations concentrated on a small number of structures. General-purpose dynet and pytorch implementations available at https://github.com/vene/sparsemap.

Preliminaries

When there are no ties, arg max⁡\operatorname*{\mathsf{arg\,max}} has a unique solution ei\bm{e}_{i} peaking at the index ii of the highest value of θ\bm{\theta}. When there are ties, arg max⁡\operatorname*{\mathsf{arg\,max}} is set-valued. Even assuming no ties, arg max⁡\operatorname*{\mathsf{arg\,max}} is piecewise constant, and thus is ill-suited for direct use within neural networks, e.g., in an attention mechanism. Instead, it is common to use softmax⁡\operatorname{\mathsf{softmax}}, a continuous and differentiable approximation to arg max⁡\operatorname*{\mathsf{arg\,max}}, which can be seen as an entropy-regularized arg max⁡\operatorname*{\mathsf{arg\,max}}

where H(y)=−∑iyiln\displaylimitsyiH(\bm{y})=-\sum_{i}y_{i}\mathop{\mathsf{ln}}\displaylimits y_{i}, i.e. the negative Shannon entropy. Since exp\displaylimits⋅>0\mathop{\mathsf{exp}}\displaylimits\cdot>0 strictly, softmax⁡\operatorname{\mathsf{softmax}} outputs are dense.

Both softmax⁡\operatorname{\mathsf{softmax}} and sparsemax⁡\operatorname{\mathsf{sparsemax}} are continuous and differentiable almost everywhere; however, sparsemax⁡\operatorname{\mathsf{sparsemax}} encourages sparsity in its outputs. This is because it corresponds to an Euclidean projection onto the simplex, which is likely to hit its boundary as the magnitude of θ\bm{\theta} increases. Both mechanisms, as well as variants with different penalties (Niculae & Blondel, 2017), have been successfully used in attention mechanisms, for mapping a score vector θ\bm{\theta} to a dd-dimensional normalized discrete probability distribution over a small set of choices. The relationship between arg max⁡\operatorname*{\mathsf{arg\,max}}, softmax⁡\operatorname{\mathsf{softmax}}, and sparsemax⁡\operatorname{\mathsf{sparsemax}}, illustrated in Figure 1, sits at the foundation of SparseMAP⁡\operatorname{\mathsf{SparseMAP}}.

2 Structured Inference

In structured prediction, the space of possible outputs is typically very large: for instance, all possible labelings of a length-nn sequence, spanning trees over nn nodes, or one-to-one alignments between two sets. We may still write optimization problems such as max\displaylimitss=1Dθs\mathop{\mathsf{max}}\displaylimits_{s=1}^{D}\theta_{s}, but it is impractical to enumerate all of the DD possible structures and, in turn, to specify the scores for each structure in θ\bm{\theta}.

where ηU\bm{\eta}_{U} and ηF\bm{\eta}_{F} are unary and higher-order log-potentials, and sis_{i} and sfs_{f} are local configurations at variable and factor nodes. This can be written in matrix notation as θ=M⊤ηU+N⊤ηF\bm{\theta}=\bm{M}^{\top}\bm{\eta}_{U}+\bm{N}^{\top}\bm{\eta}_{F} for suitable matrices {M,N}\{\bm{M},\bm{N}\}, fitting the assumption above with A=[M;N]\bm{A}=[\bm{M};\bm{N}] and η=[ηU;ηF]\bm{\bm{\eta}}=[\bm{\eta}_{U};\bm{\eta}_{F}].

where MA≔{[u;v]:u=My, v=Ny, y∈△D}\mathcal{M}_{\bm{A}}\coloneqq\{[\bm{u};\bm{v}]:\bm{u}=\bm{My},~{}\bm{v}=\bm{Ny},~{}\bm{y}\in\triangle^{D}\} is the marginal polytope (Wainwright & Jordan, 2008), with one vertex for each possible structure (Figure 1). However, as previously said, since it is equivalent to a DD-dimensional arg max⁡\operatorname*{\mathsf{arg\,max}}, MAP is piecewise constant and discontinuous.

Negative entropy regularization over y\bm{y}, on the other hand, yields marginal inference,

Marginal inference is differentiable, but may be more difficult to compute; the entropy HA(u,v)=H(y)H_{\bm{A}}(\bm{u},\bm{v})=H(\bm{y}) itself lacks a closed form (Wainwright & Jordan, 2008, §4.1.2). Gradient backpropagation is available only to specialized problem instances, e.g. those solvable by dynamic programming (Li & Eisner, 2009). The entropic term regularizes y\bm{y} toward more uniform distributions, resulting in strictly dense solutions, just like in the case of softmax⁡\operatorname{\mathsf{softmax}} (Equation 1).

Interesting types of structures, which we use in the experiments described in Section 5, include the following.

Non-projective dependency parsing. Consider a sentence of length nn. Here, a structure ss is a dependency tree: a rooted spanning tree over the n2n^{2} possible arcs (for example, the arcs above the sentences in Figure 3). Each column ms∈{0,1}n2\bm{{{m}}}_{s}\in\{0,1\}^{n^{2}} encodes a tree by assigning a 11 to its arcs. N\bm{N} is empty, MA\mathcal{M}_{\bm{A}} is known as the arborescence polytope (Martins et al., 2009). MAP inference may be performed by maximal arborescence algorithms (Chu & Liu, 1965; Edmonds, 1967; McDonald et al., 2005), and the Matrix-Tree theorem (Kirchhoff, 1847) provides a way to perform marginal inference (Koo et al., 2007; Smith & Smith, 2007).

Linear assignment. Consider a one-to-one matching (linear assignment) between two sets of nn nodes. A global structure ss is a nn-permutation, and a column ms∈{0,1}n2\bm{{{m}}}_{s}\in\{0,1\}^{n^{2}} can be seen as a flattening of the corresponding permutation matrix. Again, N\bm{N} is empty. MA\mathcal{M}_{\bm{A}} is the Birkhoff polytope (Birkhoff, 1946), and MAP inference can be performed by, e.g., the Hungarian algorithm (Kuhn, 1955) or the Jonker-Volgenant algorithm (Jonker & Volgenant, 1987). Noticeably, marginal inference is known to be #P-complete (Valiant, 1979; Taskar, 2004, Section 3.5). This makes it an open problem how to use matchings as latent variables.

𝗦𝗽𝗮𝗿𝘀𝗲𝗠𝗔𝗣𝗦𝗽𝗮𝗿𝘀𝗲𝗠𝗔𝗣\operatorname{\mathsf{SparseMAP}}

Armed with the parallel between structured inference and regularized max\displaylimits\mathop{\mathsf{max}}\displaylimits operators described in §2, we are now ready to introduce SparseMAP⁡\operatorname{\mathsf{SparseMAP}}, a novel inference optimization problem which returns sparse solutions.

The quadratic penalty replaces the entropic penalty from marginal inference (Equation 4), which pushes the solutions to the strict interior of the marginal polytope. In consequence, SparseMAP⁡\operatorname{\mathsf{SparseMAP}} favors sparse solutions from the faces of the marginal polytope MA\mathcal{M}_{\bm{A}}, as illustrated in Figure 1. For the structured prediction problems mentioned in Section 2.2, SparseMAP⁡\operatorname{\mathsf{SparseMAP}} would be able to return, for example, a sparse combination of sequence labelings, parse trees, or matchings. Moreover, the strongly convex regularization on u\bm{u} ensures that SparseMAP⁡\operatorname{\mathsf{SparseMAP}} has a unique solution and is differentiable almost everywhere, as we will see.

2 Solving 𝗦𝗽𝗮𝗿𝘀𝗲𝗠𝗔𝗣𝗦𝗽𝗮𝗿𝘀𝗲𝗠𝗔𝗣\operatorname{\mathsf{SparseMAP}}

We now tackle the optimization problem in Equation 5. Although SparseMAP⁡\operatorname{\mathsf{SparseMAP}} is a QP over a polytope, even describing it in standard form is infeasible, since enumerating the exponentially-large set of vertices is infeasible. This prevents direct application of, e.g., the generic differentiable QP solver of Amos & Kolter (2017). We instead focus on SparseMAP⁡\operatorname{\mathsf{SparseMAP}} solvers that involve a sequence of MAP problems as a subroutine—this makes SparseMAP⁡\operatorname{\mathsf{SparseMAP}} widely applicable, given the availability of MAP implementations for various structures. We discuss two such methods, one based on the conditional gradient algorithm and another based on the active set method for quadratic programming. We provide a full description of both methods in Appendix A.

Conditional gradient. One family of such solvers is based on the conditional gradient (CG) algorithm (Frank & Wolfe, 1956; Lacoste-Julien & Jaggi, 2015), considered in prior work for solving approximations of the marginal inference problem (Belanger et al., 2013; Krishnan et al., 2015). Each step must solve a linearized subproblem. Denote by ff the SparseMAP⁡\operatorname{\mathsf{SparseMAP}} objective from Equation 5,

The gradients of ff with respect to the two variables are

A linear approximation to ff around a point [u′;v′][\bm{u}^{\prime};\bm{v}^{\prime}] is

Minimizing f^\hat{f} over M\mathcal{M} is exactly MAP inference with adjusted variable scores ηU−u′\bm{\eta}_{U}-\bm{u}^{\prime}. Intuitively, at each step we seek a high-scoring structure while penalizing sharing variables with already-selected structures Vanilla CG simply adds the new structure to the active set at every iteration. The pairwise and away-step variants trade off between the direction toward the new structure, and away from one of the already-selected structures. More sophisticated variants have been proposed (Garber & Meshi, 2016) which can provide sparse solutions when optimizing over a polytope.

Active set method. Importantly, the SparseMAP⁡\operatorname{\mathsf{SparseMAP}} problem in Equation 5 has quadratic curvature, which the general CG algorithms may not optimally leverage. For this reason, we consider the active set method for constrained QPs: a generalization of Wolfe’s min-norm point algorithm (Wolfe, 1976), also used in structured prediction for the quadratic subproblems by Martins et al. (2015). The active set algorithm, at each iteration, updates an estimate of the solution support by adding or removing one constraint to/from the active set; then it solves the Karush–Kuhn–Tucker (KKT) system of a relaxed QP restricted to the current support.

Comparison. Both algorithms enjoy global linear convergence with similar rates (Lacoste-Julien & Jaggi, 2015), but the active set algorithm also exhibits exact finite convergence—this allows it, for instance, to capture the optimal sparsity pattern (Nocedal & Wright, 1999, Ch. 16.4 & 16.5). Vinyes & Obozinski (2017) provide a more in-depth discussion of the connections between the two algorithms. We perform an empirical comparison on a dependency parsing instance with random potentials. Figure 2 shows that active set substantially outperforms all CG variants, both in terms of objective value as well as in the solution sparsity, suggesting that the quadratic curvature makes SparseMAP⁡\operatorname{\mathsf{SparseMAP}} solvable in very few iterations to high accuracy. We therefore use the active set solver in the remainder of the paper.

3 Backpropagating Gradients through 𝗦𝗽𝗮𝗿𝘀𝗲𝗠𝗔𝗣𝗦𝗽𝗮𝗿𝘀𝗲𝗠𝗔𝗣\operatorname{\mathsf{SparseMAP}}

In order to use SparseMAP⁡\operatorname{\mathsf{SparseMAP}} as a neural network layer trained with backpropagation, one must compute products of the SparseMAP⁡\operatorname{\mathsf{SparseMAP}} Jacobian with a vector p\bm{p}. Computing the Jacobian of an optimization problem is an active research topic known as argmin differentiation, and is generally difficult. Fortunately, as we show next, argmin differentiation is always easy and efficient in the case of SparseMAP⁡\operatorname{\mathsf{SparseMAP}}.

Denote a SparseMAP⁡\operatorname{\mathsf{SparseMAP}} solution by y⋆\bm{y}^{\star} and its support by I≔{s : ys>0}\mathcal{I}\coloneqq\{s~{}:~{}y_{s}>0\}. Then, SparseMAP⁡\operatorname{\mathsf{SparseMAP}} is differentiable almost everywhere with Jacobian

The proof, given in Appendix B, relies on the KKT conditions of the SparseMAP⁡\operatorname{\mathsf{SparseMAP}} QP. Importantly, because D(I)\bm{D}(\mathcal{I}) is zero outside of the support of the solution, computing the Jacobian only requires the columns of M\bm{M} and A\bm{A} corresponding to the structures in the active set. Moreover, when using the active set algorithm discussed in §3.2, the matrix Z\bm{Z} is readily available as a byproduct of the forward pass. The backward pass can, therefore, be computed in O(k∣I∣)\mathcal{O}(k|\mathcal{I}|).

Our approach for gradient computation draws its efficiency from the solution sparsity and does not depend on the type of structure considered. This is contrasted with two related lines of research. The first is “unrolling” iterative inference algorithms, for instance belief propagation (Stoyanov et al., 2011; Domke, 2013) and gradient descent (Belanger et al., 2017), where the backward pass complexity scales with the number of iterations. In the second, employed by Kim et al. (2017), when inference can be performed via dynamic programming, backpropagation can be performed using second-order expectation semirings (Li & Eisner, 2009) or more general smoothing (Mensch & Blondel, 2018), in the same time complexity as the forward pass. Moreover, in our approach, neither the forward nor the backward passes involve logarithms, exponentiations or log-domain classes, avoiding the slowdown and stability issues normally incurred.

In the unstructured case, since M=I\bm{M}=\bm{I}, Z\bm{Z} is also an identity matrix, uncovering the sparsemax⁡\operatorname{\mathsf{sparsemax}} Jacobian (Martins & Astudillo, 2016). In general, structures are not necessarily orthogonal, but may have degrees of overlap.

Structured Fenchel-Young Losses and the 𝗦𝗽𝗮𝗿𝘀𝗲𝗠𝗔𝗣𝗦𝗽𝗮𝗿𝘀𝗲𝗠𝗔𝗣\operatorname{\mathsf{SparseMAP}} Loss

The Fenchel convex conjugate of Ω△\Omega_{\triangle} is

We next introduce a family of structured prediction losses, named after the corresponding Fenchel-Young duality gap.

This family, studied in more detail in (Blondel et al., 2018), includes the commonly-used structured losses:

Structured perceptron (Collins, 2002): Ω≡0\Omega\equiv 0;

Structured SVM (Taskar et al., 2003; Tsochantaridis et al., 2004): Ω≡ρ(⋅,yˉ)\Omega\equiv\rho(\cdot,\bar{\bm{y}}) for a cost function ρ\rho, where yˉ\bar{\bm{y}} is the true output;

CRF (Lafferty et al., 2001): Ω≡−H\Omega\equiv-H;

Margin CRF (Gimpel & Smith, 2010): Ω≡−H+ρ(⋅,yˉ)\Omega\equiv-H+\rho(\cdot,\bar{\bm{y}}).

This leads to a natural way of defining SparseMAP⁡\operatorname{\mathsf{SparseMAP}} losses, by plugging the following into Equation 6:

SparseMAP⁡\operatorname{\mathsf{SparseMAP}} loss: Ω(y)=12∥My∥22\Omega(\bm{y})=\frac{1}{2}\left\lVert\bm{My}\right\rVert^{2}_{2},

Margin SparseMAP⁡\operatorname{\mathsf{SparseMAP}}: Ω(y)=12∥My∥22+ρ(y,yˉ)\Omega(\bm{y})=\frac{1}{2}\left\lVert\bm{My}\right\rVert^{2}_{2}+\rho(\bm{y},\bar{\bm{y}}).

It is well-known that the subgradients of structured perceptron and SVM losses consist of MAP inference, while the CRF loss gradient requires marginal inference. Similarly, the subgradients of the SparseMAP⁡\operatorname{\mathsf{SparseMAP}} loss can be computed via SparseMAP⁡\operatorname{\mathsf{SparseMAP}} inference, which in turn only requires MAP. The next proposition states properties of structured Fenchel-Young losses, including a general connection between a loss and its corresponding inference method.

Experimental Results

In this section, we experimentally validate SparseMAP⁡\operatorname{\mathsf{SparseMAP}} on two natural language processing applications, illustrating the two main use cases presented: structured output prediction with the SparseMAP⁡\operatorname{\mathsf{SparseMAP}} loss (§5.1) and structured hidden layers (§5.2). All models are implemented using the dynet library v2.0.2 (Neubig et al., 2017).

We evaluate the SparseMAP⁡\operatorname{\mathsf{SparseMAP}} losses against the commonly used CRF and structured SVM losses. The task we focus on is non-projective dependency parsing: a structured output task consisting of predicting the directed tree of grammatical dependencies between words in a sentence (Jurafsky & Martin, 2018, Ch. 14). We use annotated Universal Dependency data (Nivre et al., 2016), as used in the CoNLL 2017 shared task (Zeman et al., 2017). To isolate the effect of the loss, we use the provided gold tokenization and part-of-speech tags. We follow closely the bidirectional LSTM arc-factored parser of Kiperwasser & Goldberg (2016), using the same model configuration; the only exception is not using externally pretrained embeddings. Parameters are trained using Adam (Kingma & Ba, 2015), tuning the learning rate on the grid {.5,1,2,4,8}×10−3\{.5,1,2,4,8\}\times 10^{-3}, expanded by a factor of 2 if the best model is at either end.

We experiment with 5 languages, diverse both in terms of family and in terms of the amount of training data (ranging from 1,400 sentences for Vietnamese to 12,525 for English). Test set results (Table 1) indicate that the SparseMAP⁡\operatorname{\mathsf{SparseMAP}} losses outperform the SVM and CRF losses on 4 out of the 5 languages considered. This suggests that SparseMAP⁡\operatorname{\mathsf{SparseMAP}} is a good middle ground between MAP-based and marginal-based losses in terms of smoothness and gradient sparsity.

Moreover, as illustrated in Figure 4, the SparseMAP⁡\operatorname{\mathsf{SparseMAP}} loss encourages sparse predictions: models converge towards sparser solutions as they train, yielding very few ambiguous arcs. When confident, SparseMAP⁡\operatorname{\mathsf{SparseMAP}} can predict a single tree. Otherwise, the small set of candidate parses returned can be easily visualized, often indicating genuine linguistic ambiguities (Figure 3). Returning a small set of parses, also sought concomittantly by Keith et al. (2018), is valuable in pipeline systems, e.g., when the parse is an input to a downstream application: error propagation is diminished in cases where the highest-scoring tree is incorrect (which is the case for the sentences in Figure 3). Unlike KK-best heuristics, SparseMAP⁡\operatorname{\mathsf{SparseMAP}} dynamically adjusts its output sparsity, which is desirable on realistic data where most instances are easy.

2 Latent Structured Alignment for Natural Language Inference

In this section, we demonstrate SparseMAP⁡\operatorname{\mathsf{SparseMAP}} for inferring latent structure in large-scale deep neural networks. We focus on the task of natural language inference, defined as the classification problem of deciding, given two sentences (a premise and a hypothesis), whether the premise entails the hypothesis, contradicts it, or is neutral with respect to it.

We consider novel structured variants of the state-of-the-art ESIM model (Chen et al., 2017). Given a premise P\mathsf{P} of length mm and a hypothesis H\mathsf{H} of length nn, ESIM:

Encodes P\mathsf{P} and H\mathsf{H} with an LSTM.

Computes P\mathsf{P}-to-H\mathsf{H} and H\mathsf{H}-to-P\mathsf{P} alignments using row-wise, respectively column-wise softmax⁡\operatorname{\mathsf{softmax}} on G\bm{G}.

Augments P\mathsf{P} words with the weighted average of its aligned H\mathsf{H} words, and vice-versa.

Passes the result through another LSTM, then predicts.

We consider the following structured replacements for the independent row-wise and column-wise softmax⁡\operatorname{\mathsf{softmax}}es (step 3):

Sequential alignment. We model the alignment of p\bm{p} to h\bm{h} as a sequence tagging instance of length mm, with nn possible tags corresponding to the nn words of the hypothesis. Through transition scores, we enable the model to capture continuity and monotonicity of alignments: we parametrize transitioning from word t1t_{1} to t2t_{2} by binning the distance t2−t1t_{2}-t_{1} into 5 groups, {−2 or less,−1,0,1,2 or more}\{-2\text{ or less},-1,0,1,2\text{ or more}\}. We similarly parametrize the initial alignment using bins {1,2 or more}\{1,2\text{ or more}\} and the final alignment as {−2 or less,−1}\{-2\text{ or less},-1\}, allowing the model to express whether an alignment starts at the beginning or ends on the final word of h\bm{h}; formally

We align p\bm{p} to h\bm{h} applying the same method in the other direction, with different transition scores w\bm{w}. Overall, sequential alignment requires learning 18 additional scalar parameters.

Matching alignment. We now seek a symmetrical alignment in both directions simultaneously. To this end, we cast the alignment problem as finding a maximal weight bipartite matching. We recall from §2.2 that a solution can be found via the Hungarian algorithm (in contrast to marginal inference, which is #P-complete). When n=mn=m, maximal matchings can be represented as permutation matrices, and when n≠mn\neq m some words remain unaligned. SparseMAP⁡\operatorname{\mathsf{SparseMAP}} returns a weighted average of a few maximal matchings. This method requires no additional learned parameters.

We evaluate the two models alongside the softmax⁡\operatorname{\mathsf{softmax}} baseline on the SNLI (Bowman et al., 2015) and MultiNLI (Williams et al., 2018) datasets.We split the MultiNLI matched validation set into equal validation and test sets; for SNLI we use the provided split. All models are trained by SGD, with 0.9×0.9\times learning rate decay at epochs when the validation accuracy is not the best seen. We tune the learning rate on the grid \big{\{}2^{k}:k\in\{-6,-5,-4,-3\}\big{\}}, extending the range if the best model is at either end. The results in Table 2 show that structured alignments are competitive with softmax⁡\operatorname{\mathsf{softmax}} in terms of accuracy, but are orders of magnitude sparser. This sparsity allows them to produce global alignment structures that are interpretable, as illustrated in Figure 5.

Interestingly, we observe computational advantages of sparsity. Despite the overhead of GPU memory copying, both training and validation in our latent structure models take roughly the same time as with softmax⁡\operatorname{\mathsf{softmax}} and become faster as the models grow more certain. For the sake of comparison, Kim et al. (2017) report a 5×5\times slow-down in their structured attention networks, where they use marginal inference.

Related Work

Structured attention networks. Kim et al. (2017) and Liu & Lapata (2018) take advantage of the tractability of marginal inference in certain structured models and derive specialized backward passes for structured attention. In contrast, our approach is modular and general: with SparseMAP⁡\operatorname{\mathsf{SparseMAP}}, the forward pass only requires MAP inference, and the backward pass is efficiently computed based on the forward pass results. Moreover, unlike marginal inference, SparseMAP⁡\operatorname{\mathsf{SparseMAP}} yields sparse solutions, which is an appealing property statistically, computationally, and visually.

KK-best inference. As it returns a small set of structures, SparseMAP⁡\operatorname{\mathsf{SparseMAP}} brings to mind KK-best inference, often used in pipeline NLP systems for increasing recall and handling uncertainty (Yang & Cardie, 2013). KK-best inference can be approximated (or, in some cases, solved), roughly KK times slower than MAP inference (Yanover & Weiss, 2004; Camerini et al., 1980; Chegireddy & Hamacher, 1987; Fromer & Globerson, 2009). The main advantages of SparseMAP⁡\operatorname{\mathsf{SparseMAP}} are convexity, differentiablity, and modularity, as SparseMAP⁡\operatorname{\mathsf{SparseMAP}} can be computed in terms of MAP subproblems. Moreover, it yields a distribution, unlike KK-best, which does not reveal the gap between selected structures,

Learning permutations. A popular approach for differentiable permutation learning involves mean-entropic optimal transport relaxations (Adams & Zemel, 2011; Mena et al., 2018). Unlike SparseMAP⁡\operatorname{\mathsf{SparseMAP}}, this does not apply to general structures, and solutions are not directly expressible as combinations of a few permutations.

Conclusion

We introduced a new framework for sparse structured inference, SparseMAP⁡\operatorname{\mathsf{SparseMAP}}, along with a corresponding loss function. We proposed efficient ways to compute the forward and backward passes of SparseMAP⁡\operatorname{\mathsf{SparseMAP}}. Experimental results illustrate two use cases where sparse inference is well-suited. For structured prediction, the SparseMAP⁡\operatorname{\mathsf{SparseMAP}} loss leads to strong models that make sparse, interpretable predictions, a good fit for tasks where local ambiguities are common, like many natural language processing tasks. For structured hidden layers, we demonstrated that SparseMAP⁡\operatorname{\mathsf{SparseMAP}} leads to strong, interpretable networks trained end-to-end. Modular by design, SparseMAP⁡\operatorname{\mathsf{SparseMAP}} can be applied readily to any structured problem for which MAP inference is available, including combinatorial problems such as linear assignment.

Acknowledgements

We thank Tim Vieira, David Belanger, Jack Hessel, Justine Zhang, Sydney Zink, the Unbabel AI Research team, and the three anonymous reviewers for their insightful comments. This work was supported by the European Research Council (ERC StG DeepSPIN 758969) and by the Fundação para a Ciência e Tecnologia through contracts UID/EEA/50008/2013, PTDC/EEI-SII/7092/2014 (LearnBig), and CMUPERI/TIC/0046/2014 (GoLocal).

References

Appendix A Implementation Details for 𝗦𝗽𝗮𝗿𝘀𝗲𝗠𝗔𝗣𝗦𝗽𝗮𝗿𝘀𝗲𝗠𝗔𝗣\operatorname{\mathsf{SparseMAP}} Solvers

We adapt the presentation of vanilla, away-step and pairwise conditional gradient of Lacoste-Julien & Jaggi (2015).

Recall the SparseMAP⁡\operatorname{\mathsf{SparseMAP}} optimization problem (Equation 5), which we rewrite below as a minimization, to align with the formulation in (Lacoste-Julien & Jaggi, 2015)

The gradients of the objective function ff w.r.t. the two variables are

The ingredients required to apply conditional gradient algorithms are solving linear minimization problem, selecting the away step, computing the Wolfe gap, and performing line search.

For SparseMAP⁡\operatorname{\mathsf{SparseMAP}}, this amounts to a MAP inference call, since

where we assume MAP⁡A\operatorname{\mathsf{MAP}}_{\bm{A}} yields the set of maximally-scoring structures.

Away step selection.

This step involves searching the currently selected structures in the active set I\mathcal{I} with the opposite goal: finding the structure maximizing the linearization

Wolfe gap.

The gap at a point d=[du;dv]\bm{d}=[\bm{d}_{\bm{u}};\bm{d}_{\bm{v}}] is given by

Line search.

Once we have picked a direction d=[du;dv]\bm{d}=[\bm{d}_{\bm{u}};\bm{d}_{\bm{v}}], we can pick the optimal step size by solving a simple optimization problem. Let uγ≔u′+γdu\bm{u}_{\gamma}\coloneqq\bm{u}^{\prime}+\gamma\bm{d}_{\bm{u}}, and vγ≔v′+γdv\bm{v}_{\gamma}\coloneqq\bm{v}^{\prime}+\gamma\bm{d}_{\bm{v}}. We seek γ\gamma so as to optimize

Setting the gradient w.r.t. γ\gamma to yields

We may therefore compute the optimal step size γ\gamma as

A.2 The Active Set Algorithm

We use a variant of the active set algorithm (Nocedal & Wright, 1999, Ch. 16.4 & 16.5) as proposed for the quadratic subproblems of the AD3 algorithm; our presentation follows (Martins et al., 2015, Algorithm 3). At each step, the active set algorithm solves a relaxed variant of the SparseMAP⁡\operatorname{\mathsf{SparseMAP}} QP, relaxing the non-negativity constraint on y\bm{y}, and restricting the solution to the current active set I\mathcal{I}

whose solution can be found by solving the KKT system

At each iteration, the (symmetric) design matrix in Equation 9 is updated by adding or removing a row and a column; therefore its inverse (or a decomposition) may be efficiently maintained and updated.

The optimal step size for moving a feasible current estimate y′\bm{y}^{\prime} toward a solution y^\hat{\bm{y}} of Equation 9, while keeping feasibility, is given by (Martins et al., 2015, Equation 31)

When γ≤1\gamma\leq 1 this update zeros out a coordinate of y′\bm{y}^{\prime}; otherwise, I\mathcal{I} remains the same.

Recall that SparseMAP⁡\operatorname{\mathsf{SparseMAP}} is defined as the u⋆\bm{u}^{\star} that maximizes the value of the quadratic program (Equation 5),

We now rewrite the QP in Equation 11 in terms of the convex combination of vertices of the marginal polytope

We use the optimality conditions of problem 12 to derive an explicit relationship between u⋆\bm{u}^{\star} and x\bm{x}. At an optimum, the following KKT conditions hold

Let I\mathcal{I} denote the support of y⋆\bm{y}^{\star}, i.e., I={s : ys⋆>0}\mathcal{I}=\{s~{}:~{}y^{\star}_{s}>0\}. From Equation 17 we have λI=0\bm{\lambda}_{\mathcal{I}}=\bm{0} and therefore

Solving for yI⋆\bm{y}_{\mathcal{I}}^{\star} in Equation 18 we get a direct expression

where we introduced Z=(M⊤M)−1\bm{Z}=(\bm{M}^{\top}\bm{M})^{-1}. Solving for τ⋆\tau^{\star} yields

Plugging this back and left-multiplying by MI\bm{M}_{\mathcal{I}} we get

Note that, in a neighborhood of η\bm{\eta}, the support of the solution I\mathcal{I} is constant. (On the measure-zero set of points where the support changes, SparseMAP⁡\operatorname{\mathsf{SparseMAP}} is subdifferentiable and our assumption yields a generalized Jacobian (Clarke, 1990).) Differentiating w.r.t. the score of a configuration θs\theta_{s}, we get the expression

Since θs=as⊤η\theta_{s}=\bm{{{a}}}_{s}^{\top}\bm{\eta}, by the chain rule, we get the desired result

Appendix C Fenchel-Young Losses: Proof of Proposition 2

Since Ω△\Omega_{\triangle} is the restriction of a convex function to a convex set, it is convex (Boyd & Vandenberghe, 2004, Section 3.1.2).

From the Fenchel-Young inequality (Fenchel, 1949; Boyd & Vandenberghe, 2004, Section 3.3.2), we have

In particular, when θ=A⊤η\bm{\theta}=\bm{A}^{\top}\bm{\eta},

where we used the fact that y∈△d\bm{y}\in\triangle^{d}. The second part of the claim follows.

Property 2.

To prove convexity in η\bm{\eta}, we rewrite the loss, for fixed y\bm{y}, as

Property 3.

This follows from the scaling property of the convex conjugate (Boyd & Vandenberghe, 2004, Section 3.3.2)

Denoting η′=t−1η\bm{\eta}^{\prime}=t^{-1}\bm{\eta}, we have that