Learning random-walk label propagation for weakly-supervised semantic segmentation

Paul Vernaza, Manmohan Chandraker

Introduction

We consider the task of semantic segmentation, which is to learn a predictor capable of accurately assigning a semantic label to each pixel in an image. As with many other popular vision problems, convolutional neural nets (CNNs) have emerged as the leading tool to solve semantic segmentation problems, in part due to their ability to leverage large datasets effectively. However, datasets for semantic segmentation remain orders of magnitude smaller than for tasks such as classification and detection, due chiefly to the much higher annotation expense of this task.

To ease the annotation burden, and in line with previous work , we propose a method for training CNN-based semantic segmentation networks given sparse annotations, such as the scribbles depicted in Fig. 1, which also depicts our proposed training strategy. The idea of our method is to learn mutually-consistent networks for propagating the sparse labels to unlabeled points, and predicting the true labeling given the image alone. Optimizing a mutual-consistency objective obviates the need for dense (or, fully-labeled) supervision. A key innovation of our approach is proposing to use a specific, probabilistic model of sparse label propagation that is a differentiable function of semantic boundary predictions. As it is differentiable, minimizing our loss via gradient-based methods results in simultaneous learning of an image-to-semantic-boundary predictor and an image-to-semantic-segmentation predictor, despite having no direct observations of semantic edges.

Our method is comparable to recent work by Lin et al. , which proposes alternating between propagating sparse labels using a CRF defined over superpixels, and training a CNN to predict the labels thus inferred. A disadvantage of this approach relative to ours is that employs a notion of label smoothness that is non-adaptive: specifically, it is assumed that labels are constant within superpixels, and the CRF binary potentials are not learned. This ultimately places an artificial upper bound on the accuracy of the training data observed by the CNN—an upper bound that never improves as more data is collected. By contrast, by expressing label propagation in terms of a learned semantic boundary predictor, we are able to learn a concept of label propagation that is entirely data-driven, enabling our method to scale fully with the data. Furthermore, the probabilistic nature of our label propagation method allows us to obtain uncertainty estimates that are directly incorporated into our learning process, mitigating the possibilty of training on propagated labels that are incorrect.

A crucial technical component of our approach is defining the label propagation process in terms of random-walk hitting probabilities , which enables efficient inference and gradient-based learning. For this reason, we refer to our approach as RAWKS, a contraction of RAndom-walk WeaKly-supervised Segmentation.

Method

Given densely labeled training images, typical approaches for deep-learning-based semantic segmentation minimize the following cross-entropy loss :

where δy(x)∈Δ∣L∣\delta_{y(x)}\in\Delta^{|\mathcal{L}|} is the indicator vector of the ground truth label y(x)y(x) and Qθ,I(x)Q_{\theta,I}(x) is a label distribution predicted by a CNN with parameters θ\theta evaluated on image II at location xx. See Table 1 for notation. In our case, dense labels y(x)y(x) are not provided: we only have labels y^\hat{y} provided at a subset of points ^X\hat{}X. Our solution is to simultaneously infer the dense labeling yy given the sparse labeling y^\hat{y} and use the inferred labeling to train QQ. To achive this, we propose to train our predictor to minimize the following loss:

Here, QQ is a predictor of the same form as before, while Py(x)∣y^,Bϕ,IP_{y(x)\mid\hat{y},B_{\phi,I}} is a predicted label distribution at xx, conditioned on predicted semantic boundaries Bϕ,IB_{\phi,I} and the sparse labels. Bϕ,IB_{\phi,I} is assumed to be a nonnegative-valued CNN-based boundary predictor, in a similar vein as . See Fig. 1 for a graphical overview of our model.

The key to our method is the definition of the propagated label distributions Py∣y^,Bϕ,IP_{y\mid\hat{y},B_{\phi,I}}. As in prior work on interactive segmentation , we define these distributions in terms of random-walk hitting probabilities, which can be computed analytically, and which allows us to compute the derivatives of the propagated label probabilities with respect to the predicted boundaries. This enables us to minimize (2) by pure backpropagation, without tricks such as the alternating optimization methods employed in prior weakly-supervised segmentation methods .

To elaborate, we first define the set of all 4-connected paths on the image that end at the first labeled point encountered:

We now assign each path ξ∈Ξ\xi\in\Xi a probability that decays exponentially as it crosses boundaries:

where Bϕ,IB_{\phi,I} is assumed to be a nonnegative boundary score prediction. High- and low-probability paths under this model are illustrated in Fig. 2(b). Given sparse labels y^\hat{y}, the probability that a pixel xx has label y′y^{\prime} is then defined as the probability that a path starting at xx eventually hits a point labeled y′y^{\prime}, given the distribution over paths (4):

This quantity can be computed efficiently by solving a sparse linear system, as described in Sec. 2.2. The random-walk model is illustrated in Fig. 2(a). Any pixel in the region labeled (ii), for example, is very likely to be labeled dog instead of chair or background, because any path of significant probability starting in (ii) will hit a pixel labeled dog before it hits a pixel with any other label.

To recap, the overall architecture of our method is summarized in Fig. 1. At training time, an input image is passed to an arbitrary segmentation predictor Qθ,IQ_{\theta,I} and an arbirtary semantic edge predictor Bϕ,IB_{\phi,I}. The semantic edge predictions and sparse training labels y^\hat{y} are passed to a module that computes propagated label probabilities Py∣y^,BP_{y\mid\hat{y},B} using the random-walk model described above. The propagated label probabilities PP and the output of the predictor QQ are then passed to a cross-entropy loss H(P,Q)H(P,Q), which is minimized in the parameters θ\theta and ϕ\phi via backpropagation. At test time, QQ is evaluated and used as the prediction. PP is not evaluated at test time, since labels are not available.

The proposed loss function (2) arises from a natural probabilistic extension of (1) to the case where dense labels are unobserved. Specifically, we consider marginalizing (1) over the unobserved dense labels, given the observed sparse labels. This requires us to define P(y∣y^,I)P(y\mid\hat{y},I). We submit that a natural way to do so is to introduce a new variable BB representing the image’s semantic boundaries. This results in the proposed graphical model in Fig. 3 to represent the independence structure of y,y^,B,Iy,\hat{y},B,I.

It is then straightforward to show that (2) is equivalent to marginalizing (1) with respect to a certain distribution P(y∣y^,I)P(y\mid\hat{y},I), after making a few assumptions. First, the conditional independence structure depicted in Fig. 3 is assumed. y(x)y(x) and y(x′)y(x^{\prime}) are assumed conditionally independent given y^,B,\hat{y},B, ∀x≠x′∈X\forall x\neq x^{\prime}\in X, which allows us to specify P(y∣y^,B)P(y\mid\hat{y},B) in terms of marginal distributions and simplifies inference. Finally, BB is assumed to be a deterministic function of II, defined via parameters ϕ\phi.

We note that the conditional independence assumptions made in Fig. 3 are significant. In particular, yy is assumed independent of II given y^\hat{y} and BB. This essentially implies that there is at least one label for each connected component of the true label image, since knowing the underlying image usually does give us information as to the label of unlabeled connected components. In practice, strictly labeling every connected component is not necessary. However, training data that egregiously violates this assumption will likely yield poor results.

2 Random-walk inference

Key to our method is the efficient computation of the random-walk hitting probabilities (5). It is well-known that such probabilities can be computed efficiently via solving linear systems . We briefly review this result here.

The basic strategy is to compute the partition function ZxlZ_{xl}, which sums the right-hand-side of (4) over all paths starting at xx and ending in a point labeled ll. We can then derive a dynamic programming recursion expressing ZxlZ_{xl} in terms of the same quantity at neighboring points x′x^{\prime}. This recursion defines a set of sparse linear constraints on ZZ, which we can then solve using standard sparse solvers.

We first define Ξxl\mathchar58={ξ∈Ξ∣ξ0=x, y^(ξτ(ξ))=l}\Xi_{xl}\mathrel{\mathop{\mathchar 58\relax}}=\{\xi\in\Xi\mid\xi_{0}=x,\,\hat{y}(\xi_{\tau(\xi)})=l\}. ZxlZ_{xl} is then defined as

The first term in the inner sum can be factored out by introducing a new summation over the four nearest neighbors of xx, denoted x′∼xx^{\prime}\sim x, easily yielding the recursion

Boundary conditions must also be considered in order to fully constrain the solution. Paths exiting the image are assumed to have zero probability: hence, Zxl\mathchar58=0, ∀x∉XZ_{xl}\mathrel{\mathop{\mathchar 58\relax}}=0,\,\forall x\notin X. Paths starting at a labeled point x∈^Xx\in\hat{}X immediately terminate with probability 1; hence, Zxl\mathchar58=1, ∀x∈^XZ_{xl}\mathrel{\mathop{\mathchar 58\relax}}=1,\,\forall x\in\hat{}X. Solving this system yields a unique solution for ZZ, from which the desired probabilities are computed as

3 Random-walk backpropagation

In order to apply backpropagation, we must ultimately compute the derivative of the loss with respect to a change in the boundary score prediction Bϕ,IB_{\phi,I}. Here, we focus on computing the derivative of the partition function ZZ with respect to the boundary score BB, the other steps being trivial.

Since computing ZZ amounts to solving a linear system, this turns out to be fairly simple. Let us write the constraints (7) in matrix form Az=bAz=b, such that AA is square, zi\mathchar58=Zxiz_{i}\mathrel{\mathop{\mathchar 58\relax}}=Z_{x_{i}} (assigning a unique linear index ii to each xi∈Xx_{i}\in X, and temporarily omitting the dependence on ll), and the iith rows of A,bA,b correspond to the constraints CiZxi=∑x′∼xiZx′C_{i}Z_{x_{i}}=\sum_{x^{\prime}\sim x_{i}}Z_{x^{\prime}}, Zxi=0Z_{x_{i}}=0, or Zxi=1Z_{x_{i}}=1 (as appropriate), where Ci\mathchar58=exp⁡(Bϕ,I(xi))C_{i}\mathrel{\mathop{\mathchar 58\relax}}=\exp(B_{\phi,I}(x_{i})). Let us consider the effect of adding a small variation ϵV\epsilon V to AA, and then re-solving the system. It can be shown that

Substituting the first-order dependence on VV into a Taylor expansion of the loss LL yields:

A first-order variation of CiC_{i} corresponds to Vi=−δiiV_{i}=-\delta_{ii}, which implies that

In summary, this implies that computing the loss derivatives with respect to the boundary score can be implemented efficiently by solving the sparse adjoint system A⊺d ⁣⁡Ld ⁣⁡C=d ⁣⁡Ld ⁣⁡zA^{\intercal}\dfrac{\operatorname{d\!}{}L}{\operatorname{d\!}{C}}=\dfrac{\operatorname{d\!}{}L}{\operatorname{d\!}{z}}, and multiplying the result pointwise by the partition function zz, which in turn allows us to efficiently incorporate sparse label propagation as a function of boundary prediction into an arbitrary deep-learning framework.

4 Uncertainty-weighting the loss

An advantage of our method over prior work is that the random-walk method produces a distribution over dense labelings Py∣y^,Bϕ,IP_{y\mid\hat{y},B_{\phi,I}} given sparse labels, as opposed to a MAP estimate. These uncertainty estimtates can be used to down-weight the loss in areas where the inferred labels may be incorrect, as illustrated in Fig. 4. In this example, the boundary predictor failed to correctly predict parts of object boundaries. In the vicinity of these gaps, the label distribution is uncertain, and the MAP estimate is incorrect. However, we can mitigate the problem by down-weighting the loss proportional to the uncertainty estimate.

More concretely, we actually minimize the following modification of the loss (2):

where we define w(x)\mathchar58=exp⁡(−αH(Py(x)∣y^,Bϕ,I))w(x)\mathrel{\mathop{\mathchar 58\relax}}=\exp(-\alpha H(P_{y(x)\mid\hat{y},B_{\phi,I}})), for some fixed parameter α\alpha. This loss reduces to (2) for the case w(x)=1w(x)=1. Although the KL component of the loss can be avoided by increasing the prediction entropy, the explicit entropy regularization term prevents trivial solutions of very large entropy everywhere.

Related work

The method most comparable to ours is the work of Lin et al. . In contrast to , our method features fully-differentiable, gradient-based training (as opposed to alternating optimization); we learn an inductive rule for predicting boundaries and propagating labels, as opposed to using non-adaptive superpixels and a CRF with non-adaptive binary potentials, which enables us to adapt to large datasets in a data-driven way; and we employ a probabilistic notion of label propagation that enables us to define an uncertainty-weighted loss that mitigates the possibilty of training on propagated labels that are incorrect. Another notable method in the same vein as is the BoxSup method , which also employs alternating optimization, but uses bounding-box annotations as weak supervision.

Our method was initially inspired by , which introduced the idea of training on what we refer to here as sparse labels as a source of weak supervision. Instead of attempting to directly propagate labels, as we do, that method leverages a notion of objectness to mitigate overfitting. A few other works have proposed different modes of weak supervision for segmentation. Notably, and both model weak supervision as imposing linear constraints to be satisfied by the predictor, resulting in models trained by alternating optimization. Our method can be viewed from a similar perspective, since we also impose our weak supervision via linear constraints (7); however, our constraints explicitly model the process of spatial label propagation, whereas the constraints proposed in model only aggregate statistics over regions, and hence make no provision for learning boundaries as we do in this work. Furthermore, our model is differentiable and can be optimized via gradient-based methods.

To the extent that it learns semantic edges with a CNN, our method is similar to previous works such as , which also learns semantic edges using a CNN. However, achieves this using direct supervision of edges—which we do not require—and does not jointly train a semantic labeler, as we do. Learning edges with a weaker form of edge supervision is proposed in ; however, this method relies on a combination of heuristic boundary detectors and bounding boxes for supervision. To reiterate, our method works without any heuristic source of boundaries as input.

Another vein of prior work relates to some form of joint reasoning about boundaries and segments. Most prominently, random walks were previously applied to interactive segmentation in . However, that work did not consider learning these random walks by learning boundary scores, as we do, nor did it consider the task of semantic image segmentation in general. More recently, proposed a method to jointly learn semantic edges and a CNN-based semantic labeler in an gradient-based learning framework. However, their method applies only to the strong-supervision case, and cannot leverage sparse annotations of the kind we employ here.

Experiments

We implemented RAWKS in the Caffe framework. The semantic-boundary-prediction network Bϕ,IB_{\phi,I} and the semantic-segmentation network Qθ,IQ_{\theta,I} were both implemented as fully-convolutional CNNs, based on the same ResNet-101 architecture. The final average-pooling and fully-connected layers were removed, and features from the last resulting layer were upsampled and combined with intermediate-layer features to produce a 4x-downsampled output for both the semantic boundary and label predictions. Both networks were initialized from a model trained for classification on the ImageNet 2012 dataset. We applied no data augmentation techniques in training any of the methods.

RAWKS was trained on the publicly available scribble annotations provided by . We used the same training and validation splits as , for both the PASCAL VOC 2012 and PASCAL CONTEXT datasets: for VOC, the validation set consisted of the VOC 2012 validation set, while the training set consisted of all other images labeled in either the VOC 2012 dataset or the PASCAL Semantic Boundary dataset (10582 training images, and 1449 validation images). We trained models for each of these datasets independently.

We evaluated the performance of both the semantic-segmentation network QQ and the label-propagation network PP (propagating the sparse labels given the learned boundaries). To summarize, evaluating the predicted labels QQ on the validation set, RAWKS slightly underperformed the published results of on VOC 2012, while slightly outperforming on CONTEXT. Our other major observation was that the propagated labels PP on the training set were approximately as accurate as the best possible labeling of a superpixel segmentation of the images.

To elaborate, in Table 2, MIOU refers to the mean-intersection-over-union metric, while w/ CRF refers to the same metric evaluated after post-processing the results with a fully-connected CRF, as in . RAWKS Qθ,IQ_{\theta,I} refers to the evaluation of the predicted labels QQ on the validation set given the image alone, after jointly training PP and QQ via SGD. RAWKS train P,Q then Q also refers to evaluation of QQ, but with a slightly different training protocol: in this case after jointly training PP and QQ with SGD, QQ was fine-tuned with PP fixed. RAWKS training Py∣y^,BP_{y\mid\hat{y},B}, 0% abstain refers to evaluation of PP on the training set, given the sparse training labels and learned boundaries Bϕ,IB_{\phi,I}. RAWKS training Py∣y^,BP_{y\mid\hat{y},B}, 6% abstain consists of the same evaluation, but allowing PP to abstain from prediction on 6% of the pixels, which (for this particular model) corresponds to abstaining on all pixels with a confidence score w(x)w(x) below 0.5 (c.f. Sec. 2.4). The next section of Table 2 reports baselines: sparse-loss baseline consists of training the same base network QQ, but using loss (1) (i.e., without PP), evaluating it only at the sparse locations ^X\hat{}X. train on dense ground truth is the result we obtain training our base network QQ on the dense ground-truth training data. ScribbleSup refers to the result reported by , which we report here verbatim. We note that used a different base segmentation network (DeepLab) than we used in our experiments.

The last section of Table 2 reports statistics for baselines meant to represent best-case performance bounds for a superpixel-based method such as . These were obtained by segmenting the input images using the method (with author-suggested parameters), and labeling the resulting superpixels in different ways. SPOPT corresponds to labeling each superpixel with the majority label from the ground-truth dense segmentation. SPCON differs from SPOPT only on superpixels containing scribble annotations: to these, SPCON assigns the majority label from the scribbles contained within. Train on SPOPT/SPCON refer to training predictors QQ using the labelings of SPOPT and SPCON, respectively. These results are interesting for a number of reasons. First, the propagated labelings PP that we deduce in the course of training, are nearly as good (VOC) or better (CONTEXT) than the best possible results obtainable using superpixels. Second, we see that training with our propagated labelings PP is competitive with training on optimal superpixel labelings. Finally, we emphasize that while the superpixel baselines cannot improve with training data (as they are not trained), our label propagation model is naturally refined as we train on larger datasets.

In relative terms, RAWKS performed better on the CONTEXT dataset, as evidenced in Table 3. One potential reason for this is the greater number of classes for this dataset (60 vs. 21 for VOC), which naturally calls for finer boundaries. Since our method is able to adaptively learn boundaries suited to the task, while uses non-adaptive heuristics to generate superpixels, this may account for the better relative performance of RAWKS in this context. We also hypothesize that it is easier to learn semanic boundaries when there are a greater number of classes, because low-level edges and features become a more informative cue in this case. Surprisingly, our dense ground-truth baseline performed significantly worse than RAWKS; we hypothesize this is due to overfitting, a consequence of the smaller amount of training data and increased number of classes in CONTEXT, exacerbated by our use of the very-deep ResNet model. Joint training of the propagator network PP in RAWKS seems to have a regularizing effect that may have prevented overfitting to some extent.

Qualitative validation-set results are shown in Fig. 5 for the VOC 2012 dataset, while CONTEXT training-set results are shown in Fig. 6. The training-set results of Fig. 6 demonstrate that RAWKS is able to deduce high-quality semantic boundaries, thereby producing propagated labelings PP that are a close approximation to the ground truth dense labelings (which are not used at training time, to be clear). In the validation-set results of Fig. 5, the loss weights w(x)w(x) and propagated labels PP are shown in addition to the predictions QQ—to be clear, these depend on the sparse labels for these specific examples, which were not used to train this model. Here we remark that our semantic boundary predictions also generalize well to the validation set. Although we did not train on these images, we also observe that had we done so, the loss weights would have behaved appropriately, down-weighting the loss in regions where the propagated labels are incorrect. This seems to happen most often in regions with very fine boundaries (such as the mast of the boat and the airplane’s wing), where our limited resolution sometimes causes missed boundaries.

In general, we note that subjectively, resolution seemed to be a limiting factor in the accuracy of our boundary prediction and label propagation steps. We used quarter-resolution outputs (typicaly around 128x96 pixels) for these steps in order to minimize the computational cost of computing random-walk hitting probabilities. An average forward-backwards pass of the entire network took about 1.1 s per image, with about 800 ms of that spent solving for the random-walk hitting probabilities. This layer was implemented using a CPU-based sparse linear system solver, whereas the rest of the network was run on the GPU (an NVIDIA GTX 1080). We anticipate that implementing this layer using GPU operations will allow us to increase the resolution of these critical steps, which will in turn lead to increased prediction accuracy.

Conclusions

We have presented a novel approach to mitigating the expense of procuring labeled data in semantic segmentation, through a framework that utilizes only sparse clicks or scribbles for training. This has a significant impact on the possibilities for semantic segmentation—for a given dataset, one may obtain competitive labels at a fraction of the cost and conversely, for a given budget, one may obtain labeled data at a much larger scale. Our main technical contribution is a random-walk based label propagation mechanism, which is shown to be differentiable and usable in powerful deep neural network architectures for semantic segmentation. We achieve this through a novel predictor-propagator paradigm, which produces uncertainty estimates for inferred dense labels given sparse labels. We demonstrate encouraging results on challenging benchmarks. More importantly, we argue that our framework has inherent advantages over prior works, since our label propagation is not artificially upper-bounded by superpixel baselines, rather, can keep improving with larger-scale training data. Also, we note that our contribution is equally valid for any state-of-the-art CNN-based semantic segmentation engines. In future work, we will explore other state-of-the-art segmentation architectures and incorporate other forms of weak supervision such as bounding boxes.

References