Semantic Segmentation with Generative Models: Semi-Supervised Learning and Strong Out-of-Domain Generalization
Daiqing Li, Junlin Yang, Karsten Kreis, Antonio Torralba, Sanja Fidler
Introduction
Deep learning is now powering the majority of computer vision applications ranging from autonomous driving and medical imaging to image editing . However, deep networks are extremely data hungry, typically requiring training on large-scale datasets to achieve high accuracy. Even when large datasets are available, generalizing the network’s performance to out-of-distribution data, for example, on images captured by a different sensor, presents challenges, since deep networks tend to overfit to artificial statistics in the training data. Labeling large datasets, particularly for dense pixel-level tasks such as semantic segmentation, is already very time consuming. Re-doing the annotation effort each time the sensor changes is especially undesirable. This is particularly true in the medical domain, where pixel-level annotations are expensive to obtain (require highly-skilled experts), and where imaging sensors vary across sites. In this paper, we aim to significantly reduce the number of training data required for attaining successful performance, while achieving strong out-of-domain generalization.
Semi-supervised learning (SSL) facilitates learning with small labeled data sets by augmenting the training set with large amounts of unlabeled data. The literature on SSL is vast and some classical SSL techniques include pseudo-labeling , consistency regularization , and various data augmentation techniques (also see Sec. 2). State-of-the-art SSL performance is currently achieved by contrastive learning, which aims to train powerful image feature extractors using unsupervised contrastive losses on image transformations . Once the feature extractors are trained, a smaller amount of labels is needed, since the features already implicitly encode semantic information. While SSL approaches have been more widely explored for classification, recent methods also tackle pixel-wise tasks .
Although SSL techniques allow to train models with little labeled data, they usually do not explicitly model the distribution of the input data itself and therefore can still easily overfit to the training data, hampering their generalization capabilities. This is especially critical in semantic segmentation, where annotations are expensive and hence the available amount of labeled data can be particularly small.
To address this, we propose a fully generative approach based on a generative adversarial network (GAN) that models the joint image-label distribution and synthesizes both images and their semantic segmentation masks. We build on top of the StyleGAN2 architecture and augment it with a label generation branch. Our model is trained on a large unlabeled image collection and a small labeled subset using only adversarial objectives. Test-time prediction is framed as first optimizing for the latent code that reconstructs the input image, and then synthesizing the label by applying the generator on the inferred embedding.
We showcase our method in the medical domain and on human faces. It achieves competitive or better in-domain performance even when compared to heavily engineered state-of-the-art approaches, and shows significantly higher generalization ability on out-of-domain tests. We also demonstrate the ability to generalize to domains that are drastically different from the training domain, such as going from CT to MRI volumes, and natural photographs of faces to sculptures, paintings and cartoons, and even animal faces (see Figure 1).
In summary, we make the following contributions: (i) We propose a novel generative model for semantic segmentation that builds on the state-of-the-art StyleGAN2 and naturally allows semi-supervised training. To the best of our knowledge, we are the first work that tackles semantic segmentation with a purely generative method that directly models the joint image-label distribution. (ii) We extensively validate our model in the medical domain and on face images. In the semi-supervised setting, we demonstrate results equal to or better than available competitive baselines. (iii) We show strong generalization capabilities and outperform our baselines on out-of-domain segmentation tasks by a large margin. (iv) We qualitatively demonstrate reasonable performance even on extreme out-of-domain examples.
Related Work
Our paper touches upon various topics, including medical image analysis, semantic segmentation, semi-supervised learning, generative modeling and neural network inversion.
Semi-Supervised Learning and Semantic Segmentation: in the medical domain, semi-supervised semantic segmentation has been tackled via pseudo-labeling , adversarial training , and transformation-consistency in a mean-teacher framework . In computer vision, is the first work using an adversarial objective to train a segmentation network. Later this idea was extended to semi-supervised setups via self-taught losses and discriminator feature matching . Recently, proposed an approach using a flaw detector to approximate pixel-wise prediction confidence. Further relevant approaches to semi-supervised segmentation have been developed in weakly-supervised setups .
For simpler classification tasks, a plethora of SSL methods have been developed, based on pseudo-labeling , self-supervision , entropy-minimization , consistency-regularization , adversarial training , data augmentation , and combinations thereof . However, current state-of-the-art semi-supervised methods are based on self-supervised learning with contrastive objectives . These approaches use unlabeled data in an often task-agnostic manner to learn general feature representations that can be “finetuned” using a smaller amount of labeled data. Related ideas have been applied to semi-supervised semantic segmentation and tailored data augmentation strategies have been explored . Furthermore, many works employ carefully designed pretext tasks to learn useful representations from unlabeled images . Our method is related to these works in the sense that our task for learning strong features is image generation itself, instead of an auxiliary pretext task.
The above works train discriminative models of the form , in contrast to our fully generative approach. However, generative approaches to SSL have been proposed before. leverages variational autoencoders and use GANs in which the discriminator distinguishes between different classes. A related approach to semi-supervised semantic segmentation uses generative models to augment the training data with additional synthesized data . Conceptually, trains a generator together with a pixel-wise discriminator network to perform segmentation, while learn to generate synthetic 3D scenes by matching distributions of real and rendered imagery. In parallel work, exploit GANs to synthesize large labeled data datasets using very few labeled examples. In contrast, to the best of our knowledge, our method is the first fully generative approach to semantic segmentation that uses only adversarial objectives and no cross entropy terms and in which the generator models the joint and directly synthesizes images together with pixel-wise labels. We further use the generative model as a decoder of semantic outputs at test time, which we show leads to better generalization than prior and parallel work.
Generator Inversion: A critical part of our method is the effective inversion of the GAN generator at test-time to infer the latent embedding of a new image to be labeled. We are building on previous works that have studied this task before. Optmization-based methods iteratively optimize a reconstruction objective or perform Markov chain Monte Carlo , while encoder-based techniques directly map target images into the embedding space . Hybrid methods combine these ideas and initialize iterative optimization from an encoder prediction . These works primarily focus on image reconstruction and editing, while we use inferred embeddings for pixel-wise image labeling.
Generative Models for Image Understanding: Our approach is in line with various works that explore the use of generative modeling in different forms for discriminative and image recognition tasks, an idea that dates back until at least and has also been studied using early energy-based models . train deep belief networks to model shapes and learn representations of images in the model’s latent variables. These representations can then be used for image recognition. In , a VAE was used for amodal object instance segmentation. These ideas are closely related to our method, which learns features in a GAN generator that can be used for semantic segmentation.
Recently, demonstrated impressive inpainting as well as colorization and super-resolution results using GANs. In fact, also our method can be interpreted as “inpainting” of missing labels using a generative model of the joint image-label distribution in a similar manner. Along a different line of research, found that generative training of deep classification networks results in better calibrated and more robust models, which is consistent with the strong generalization capabilities we observe in our method.
These related works motivate to also treat semantic segmentation as a generative modeling problem.
Method
We first provide a conceptual overview over our method and discuss its motivation and advantages. Then, we explain the model architecture, training, and inference in detail.
Traditional neural network-based semantic segmentation methods learn a function , mapping images to pixel-wise target labels . The goal of learning is to maximize the conditional probability . This requires large labeled data sets and is prone to overfitting when training with limited amounts of annotated images.
We propose to instead model the joint distribution of images and labels with a GAN-based generative model. In the GAN framework, is implicitly defined as the distribution obtained when mapping latent variables drawn from a noise distribution through a deterministic generator that outputs both images and labels . In this setup, a latent vector explains both the image and its labels, and, given , image and labels are conditionally independent. Hence, we can label a new image by first inferring its embedding via an auxiliary encoder and test-time optimization, and then synthesize the corresponding pixel-wise labels (in practice, we are working directly in StyleGAN2’s -space instead of the “Normal” -space). See Figure 3 for an overview.
2 Motivation
Our fully generative approach to semantic segmentation has several advantages over traditional methods that directly model the conditional .
Semi-supervised Training: Intuitively, a model that can generate realistic images should know how to generate the corresponding pixel-wise labels as well, as they are just capturing semantic information already present in the image itself. This is analogous to rendering, where, if we know how to render a given scene, generating labels of interest, such as segmentation or depth, is simple. A GAN can be viewed as a neural renderer, where the embeddings completely encode and describe the images to be synthesized via a neural network . This connection suggests a similar strategy: If we know how to generate images, the GAN should be able to easily generate associated labels as well. This implies that the feature representations learnt by the GAN can be expected to be useful also for pixel-wise labeling tasks. Hence, we can simply augment the generator with a small additional branch that synthesizes labels from the same features used for image generation. A major benefit of this approach is that training the GAN itself only requires images without labels. A small amount of labels is only necessary for training the small labeling function on top of the main GAN architecture. Therefore, this setup naturally allows for efficient semi-supervised training. Furthermore, by jointly training the GAN for image and label synthesis, its features can be further “finetuned” for semantic label generation.
Note that we can also view this setup as parameter sharing: Given an embedding , the label-generating function shares nearly all its parameters with the image-generating function and only adds few additional parameters that are solely trained with labeled data.
Generalization: After training, we expect the model to synthesize plausible image-label pairs for all embeddings within the noise distribution , from which we drew samples during training. Therefore, we will likely be able to successfully label any new images, whose embeddings are in or sufficiently close. Furthermore, the GAN never sees the same input repeatedly during training, as its input is the resampled noise . Hence, it learns a smooth generator function over the complete latent distribution . In contrast, a purely conditional model is much more likely to overfit to the limited labeled training data and does not take into account the distribution of the data itself. For these reasons, our generative approach can be expected to show significantly better generalization capabilities beyond the training data and even beyond the training domain, which we validate in our experiments.
3 Model
We build our model on top of StyleGAN2 , the current state-of-the-art GAN for image synthesis. It is based on its successor StyleGAN and proposes several modifications, such as latent space path-length regularization to encourage generator smoothness and a redesign of instance normalization to remove generation artifacts. Furthermore, the previous progressive growing strategy is abandoned in favor of a residual skip-connection design. The model achieves remarkable image synthesis quality and has found important applications for example in image editing . We now explain our model design in detail.
Generator: Our generator is based on StyleGAN2’s generator with residual skip-connection design . We add an additional branch at each style layer to output a segmentation mask along with the image output (Figure 3). Like standard StyleGAN2, our generator takes random noise vectors following a simple Normal distribution as input and first transforms them via a fully-connected network to a more complex distribution in a space usually denoted as . After an affine transformation, these complex noise variables are then fed to the generator’s main style layers, which output images and pixel-wise labels . We can formally define this as .
Encoder and -space: During inference, we first need to infer a new image’s embedding. Instead of performing inference in -space, it has been shown that it is beneficial to instead directly work in -space and to model all noise vectors independently, unlike in training, where the same is provided to all style layers . When modeling the ’s independently for each style layer, we can interpret this as an extended space, which is usually denoted as with elements . We are following this previous work and perform embedding inference in . Below, when writing , we indicate generation directly based on , instead of samples .
As explained below, we infer an image’s embedding via test-time optimization. To speed up this optimization process and provide a strong initialization, we are using an additional encoder , mapping images directly to -space. Its architecture is based on , which uses a feature pyramid network as backbone to extract multi-level features. A small fully convolutional network is used to map those features to -space (see Figure 3).
4 Training
We utilize a large unlabeled data set and a small labeled data set , with . We are training in two stages and train generator and discriminators first and encoder second.
Loss Function: The generator and the discriminators are trained with the following standard GAN objectives:
The objective of the discriminators and is to maximize and respectively, while the objective of the generator is to minimize . The second term in Eq. (3) leads to gradients in both the image and label branch of the generator. These gradients are produced by and encourage adjustment of synthesized labels and images. However, we want the synthesized label to be adjusted to match the synthesized image, instead of the other way, \ieperturbing image generation to match the labels. Therefore, we are stopping gradient backpropagation into the generator via the image synthesis branch from the second term in Eq. (3). In this way, the image generation branch is trained purely via the image generation task with feedback from (first term in Eq. (3)), using the complete data set including the unlabeled images. At the same time, the GAN’s main features in the style layers, to which both the image and label synthesis branches are connected, are still experiencing feedback from both the image and the label synthesis branch. Due to this joint training strategy, the generator learns feature representations useful both for realistic image synthesis and corresponding label generation. Note that we use only adversarial losses. There are no pair-wise losses for segmentation, such as cross-entropy between pairs of real and generated label masks, at all.
Encoder: When training the encoder , we freeze the generator . The encoder training objective is
is the supervised loss on labeled images, defined as:
with denoting a pixel-wise cross-entropy loss summed over all pixels and the dice loss as in . The unsupervised loss is
with a hyperparameter trading off different loss contributions, denoting the generator’s image backbone and the label generation branch. is the Learned Perceptual Image Patch Similarity (LPIPS) distance , which measures L2 distance in the feature space of an ImageNet-pretrained VGG19 network. With the above objective, we are training the encoder to map images to embeddings , which re-generate the input images and, for labeled data, also the pixel-wise label masks.
5 Inference
At inference time, we are given a target image and our goal is to find the optimal pixel-wise labels . As explained above, we first embed the target image into the generator’s embedding space, for which we choose instead of . To this end, we are mapping the image to using the encoder and then solve the inversion objective
iteratively via gradient descent-based methods. The first term in Eq. (7) optimizes for reconstruction quality of the given image and the second term regularizes the optimization trajectory to stay in the training domain, where the encoder was training to approximately invert the generator. This strategy was recently proposed in . This regularization, controlled by the hyperparameter , can be particularly beneficial when performing labeling of images outside the training domain. In this case, purely optimizing for reconstruction quality can result in values that lie far outside the distribution of embeddings encountered during training. Since the labeling branch is not trained for such , the predicted label may be incorrect. One may suggest to instead directly regularize with , however, this is not easily possible, since is not actually a tractable distribution. It is only implicitly defined by mapping samples from though the fully-connected noise transformation layers of StyleGAN2.
For the reconstruction term we follow and use LPIPS together with a per-pixel L2 term:
where is another hyperparameter. After obtaining , we pass it back to the generator to get . Since was optimized to minimize the reconstruction error between and , we have . Furthermore, as the generator was trained to align synthesized segmentation labels and images, we can expect to be a correct label of the reconstructed image . Hence, the generated segmentation mask is the almost optimal segmentation of the target image .
Note that we can also look at our inference protocol from a fully probabilistic perspective, where we find the maximum of the log posterior distribution over embeddings given an image. In the supplemental material, we discuss this in more detail and how it relates to other works.
Experiments
Our approach is limited by the expressivity of the generative model. Although GANs have achieved outstanding synthesis quality for “unimodal” data such as images of faces, current generative models cannot model highly complex data, such as images of vivid outdoor scenes. Hence, our method is not applicable to such data. Therefore, in our experiments we focus on human faces as well as the medical domain, where most images can be successfully modeled by StyleGAN2 (see Figure 4), and where annotation is particularly expensive, as it relies on highly skilled experts.
We test our method on three medical tasks, chest X-ray segmentation, skin lesion segmentation, and cross-domain computer tomography to magnetic resonance image (CT-MRI) liver segmentation, as well as face part segmentation. For each task, we assume we have access to a small labeled and a relatively large unlabeled data set. We test our model on in-domain and several out-of-domain data sets. In the following, we explain our experimental setup, and report results by qualitatively and quantitatively comparing to strong baselines. We also analyze the value of labeled, unlabeled and synthetic data. Implementation details are in Appendix.
Datasets. For chest x-ray segmentation, we use two in-domain datasets (for labeled and unlabeled data), on which we train the model. We evaluate on three additional out-of-domain datasets. The datasets vary in terms of sensor quality and patient poses. We follow a similar approach for skin lesion segmentation and combine two datasets for training and evaluate additionally on three out-of-domain datasets. For the cross-domain CT-MRI liver segmentation task we use a CT dataset as our in-domain training data and also evaluate on two MRI datasets. For face part segmentation, we use the CelebA dataset . Furthermore, for out-of-domain evaluation we randomly selected 40 images from the MetFaces dataset , a collection of human face paintings and sculptures, and manually annotated them following the labeling protocol for CelebA.
Metrics. For chest X-ray and CT-MRI liver segmentation, we report per-patient DICE scores, the default metric used in the literature for this task. For skin lesion segmentation, we report the per-patient JC index, following the ISIC challenge . For face part segmentation, we use mean Intersection over Union (mIoU) over all classes, excluding the background class. mIoU is the most widely used metric in computer vision for segmentation tasks.
Baselines. As baselines, we use both fully-supervised approaches, which use only the annotated subset of the data, as well as semi-supervised semantic segmentation methods, which also utilize the additional unlabeled data. With regards to fully supervised methods, the most widely used segmentation network in the medical field is U-Net . Furthermore, following we compare to DeepLabV2 (denoted as DeepLab below), a stable and commonly-used architecture in the computer vision community for segmentation tasks. We also benchmark our approach against several state-of-the-art SLL methods for segmentation that have code available: the mean teacher model with transformation-consistency (MT) , the adversarial training-based method (AdvSSL), and also the recently proposed Guided Collaborative Training (GCT) . All baselines share the same ResNet-50 backbone network architecture. For SSL baselines, we use the default settings as reported in the original paper. The implementations are based on the PixelSSL repository https://github.com/ZHKKKe/PixelSSL.
We consider two versions of our own model. In one, we infer an image’s embedding using the encoder only (denoted as Ours-NO). In the other, we further perform optimization as described in Sec. 3.5 (denoted as Ours).
Further details about the datasets, evaluation metrics, and baselines can be found in the supplemental material.
2 Semi-Supervised Segmentation Results
Chest X-ray Segmentation. Table 1 shows our results for chest x-ray segmentation. We see that when evaluating on in-domain data, our model is on-par or better than all baselines. When evaluating on other, out-of-domain chest x-rays, our model outperforms all baselines, both the fully supervised and semi-supervised ones, often by a large margin. Examples of different segmentations are in Fig. 6.
Skin Lesion Segmentation. Table 2 presents the results for skin lesion segmentation (also see Figure 7 for visualizations). The gap between our method and the baselines is even more pronounced. We consistently outperform all baseslines, both supervised and semi-supervised ones as well as both in-domain and during evaluation on out-of-domain data.
Face Part Segmentation. We observe similar results for face part segmentation, where we outperform all baselines (see Table 4 and Figure 6). In particular for out-of-domain segmentation on the MetFaces data set, we find that we beat the other methods by a large margin. Since our method is designed with semi-supervised training in mind, we trained the models with a limited number of annotations. For reference, we additionally trained a DeepLab model with all 28k mask annotations of the CelebA dataset. This model achieves mIoU when evaluated on CelebA test data and mIoU when evaluated on MetFaces test data. Comparing to Table 4, this means that our method, using only 1.5k labels, even outperforms a modern DeepLab model that was trained with all available 28k labels when evaluated on out-of-domain MetFaces data. This is a testament to our model’s strong generalization and efficient semi-supervised training capabilities.
Encouraged by these results we experiment with evaluating our CelebA model also on more extreme out-of-domain images. We test our model on cartoons, faces of animals and even non-face images that exhibit face-like features (see Figure 8). Qualitatively, we observe that we can generate reasonable segmentations even for these extreme out-of-domain examples, a feat that hasn’t been demonstrated before, to the best of our knowledge.
CT-MRI Transfer. Having observed that our model demonstrates very strong generalization properties in the visual domain, we explore an additional far-out-of-domain problem in medical image analysis: We train our segmentation method on CT images and evaluate on MR images for liver segmentation. Our results in Table 3 demonstrate that our model outperforms the chosen supervised baselines on this very challenging out-of-domain segmentation task by a large margin. Details about this additional experiment are in the supplemental material.
We attribute our model’s strong generalization performance in the semi-supervised setting to its design as a fully generative model. Our experimental results validate our assumptions and motivations discussed in Sec. 3.1. We also find that we generally obtain better results when refining an image’s inferred embedding via optimization, as described in Sec. 3.5, instead of directly using the encoder prediction.
3 Value of Data & Training with Generated Data
We conduct an ablation study on the amount of unlabeled and labeled data used in our method. Traditionally, labeled data is considered more valuable than unlabeled data but there is no clear understanding of how many unlabeled data points boost performance as much as a labeled data sample. We measure the value of data in terms of segmentation performance (mIoU). In Table 5, we report performance for different amounts of labeled and unlabeled data used during training. Interestingly, we observe that the performance with 1500 labeled and 3K unlabeled data is almost equivalent to 150 labeled and 28K unlabeled data samples.
Simulation is often used to directly generate annotated synthetic data, reducing the need for expensive manual labelling. However, it is unclear to which degree synthetic data is useful for downstream tasks, due to the domain gap between simulated and real data. We conduct another experiment to evaluate the value of synthetic labeled data. Since our method models the joint image-label distribution, we can also use our model to generate a large amount of synthetic but annotated images. These can then be used to train a regular segmentation network in a fully-supervised, discriminative manner. Specifically, we sample 20k synthetic face images and their pixel-wise labels, using two different sampling strategies, and then train DeepLabV2 segmentation models with this data. Our results, presented in Table 6, show that high quality synthetic data is useful for the downstream task. We explored different strategies on how to sample and use the data and find that they all beat the baseline that was trained with real data only. However, this approach is sensitive to the sampling strategy used to generate the data. Importantly, we also observe that directly doing segmentation with the generative model, as proposed in this paper, performs best by a large margin. However, doing segmentation with the generative model requires test-time optimization and is thus not suitable for real-time applications. Speed-ups are future work.
Conclusion
In this paper, we propose a fully generative approach to semantic segmentation, based on StyleGAN2, that naturally allows for semi-supervised training and shows very strong generalization capabilities. We validate our method in the medical domain, where annotation can be particularly expensive and where models need to transfer, for example, between different imaging sensors. Quantitatively, we significantly outperform available strong baselines in- as well as out-of-domain. To showcase our method’s versatility, we perform additional experiments on face part segmentation. We find that our model generalizes to paintings, sculptures and cartoons. Interestingly, it produces plausible segmentations even on extreme-out-of-domain examples, such as animal faces. We attribute the model’s remarkable generalization capabilities to its design as a fully generative model.