On-the-Fly Test-time Adaptation for Medical Image Segmentation

Jeya Maria Jose Valanarasu, Pengfei Guo, Vibashan VS, Vishal M. Patel

Introduction

Image segmentation is a major task in medical imaging as it is essential for computer-aided diagnosis and image-guided surgery systems. In the past few years, deep learning-based solutions have been widely popular for medical image segmentation. Many convolutional methods and transformer-based methods have been proposed for various medical image segmentation tasks showing very good performance. However, a major problem with deep neural networks (DNN) is that they are highly dependent on the dataset that they are trained on. If a DNN is trained on a specific dataset and tested on a different dataset, the performance usually drops even if they are of the same modality. This happens due to occurrence of many shifts like camera/scanner parameters, resolution, intensity, and contrast variations. This drop in performance makes DNN-based solutions for medical imaging tasks impractical to be adopted for real-time clinical use. For clinical use, the model needs to be robust as there can be small changes in the test data distribution even if they are of the same modality.

Many domain adaptation techniques for medical image segmentation have looked into solving this problem . However, this setting assumes that we have access to the source model, source data as well as the target data. Another setting very close to real-time is fully test-time adaptation where we assume that we do not have access to the source data and adapt the model to the target data by performing one back propagation per sample. However, the model is adapted to the complete test distribution as the model weights are updated for at least one complete epoch. This setting can be considered one-shot adaptation as the model sees all the data in the distribution at least once. This can also be extended for few-shot adaptation to further adapt the model during test time. However, this setting is also not clinically deployable as we need a complete distribution to perform the adaptation and get the new model weights. Also, performing back-propagation during test-time means that we still have to do some training during the deployment-time although it is unsupervised.

In this work, we propose a clinically motivated setting called On-the-Fly test-time adaptation where the model adapts to a single image/volume at a time. Here, we do not perform any back-propagation during test time and just attempt to adapt our model to the new data instance. Also, as the model is reset for every data instance there is no need to assume the availability of the complete target distribution to perform adaptation since patient data may come with privacy concerns. On-the-Fly Adaptation is a more useful scenario in the current trend as there has been a shift of laboratory to bed-side settings for medical imaging . To solve this problem, we propose a new framework called Adaptive-UNet where the model is equipped with adaptive batch-norm layers in both the encoder and decoder to adapt to a select domain code.

In summary, the following are the contributions of this work: 1) We introduce On-the-Fly Adaptation which is more closer to real-world clinical scenarios where the adaptation is zero-shot and episodic removing the assumption of the availability of complete target distribution and back-propagation during test-phase. 2) We propose Adaptive-UNet, a new framework that learns to adapt to a new test data instance making use of a domain code and adaptive batch normalization. 3) We validate our method for 9 domain shifts in medical image segmentation for 2D fundus images and 3D MRI volumes where we get better performance than recent test-time adaptation methods.

Related Works

Unsupervised Domain Adaptation for medical image segmentation is a widely explored topic. Methods like feature alignment using adversarial training , disentangling the representation , ensembling and using soft labels have been proposed. These methods, however, use the training distribution of both source and target data for adaptation which is not always feasible for medical imaging due to privacy concerns.

Source-free Unsupervised Domain Adaptation works for medical image segmentation assume no availability of source data. In , a label-free entropy loss is defined over target distribution with a domain-invariant prior. In , an uncertainty aware denoised pseudo label method is proposed.

Test-time Adaptation methods such as TENT uses entropy minimization of batch norm statistics to adapt to a new target distribution. Recently, Hu et al. proposed using new losses like regional nuclear norm and contour regularization to improve test-time performance for medical image segmentation. Self domain adapted networks use auto-encoder based adaptors to rapidly adapt to a new task at test-time. Karani et al. proposed a per-test-image adaptation method where they adapt the image so as to obtain plausible segmentation. They still update the weights during test time by assessing the similarity of a given segmentation to those in the source data.

On-the-Fly Adaptation

Let XX and YY represent the set of input data and labels, respectively. Let us assume that the source distribution is represented as XsX_{s}, YsY_{s} and the target distribution as XtX_{t}, YtY_{t}. In normal source training, we train the model using XstrainX_{s}^{train}, YstrainY_{s}^{train} and test the model on XstestX_{s}^{test}. For direct testing (no adaptation), we use the model trained on XstrainX_{s}^{train}, YstrainY_{s}^{train} and test on XttestX_{t}^{test}. The best performance would be obtained when the model is trained using XttrainX_{t}^{train}, YttrainY_{t}^{train} and tested on XttestX_{t}^{test} which corresponds to the oracle performance.

In general test-time adaptation, we assume the availability of the entire target test distribution XttestX_{t}^{test} during test-time. The weights are optimized on the gradients calculated using an unsupervised loss function loss Ltta(Xttest)\mathcal{L}_{tta}(X^{test}_{t}) obtained using test data distribution XttestX_{t}^{test}. Some works like perform adaptation for each test-image. However, they do optimize the weights according to the test-image at hand during inference. Optimizing the network weights across all the data in test distribution and later using the new model to validate on XttestX_{t}^{test} again is not suitable in a clinical setting. It makes more sense to adapt the model for each test-image as using the entire test distribution in medical setting involves using a variety of patient data at test-time. It is also difficult to update model weights during deployment time as it requires massive computational power.

In the proposed On-the-Fly adaptation, we focus on adapting to a single test image/volume at a time as illustrated in Fig. 1. Instead of assuming we have the entire target test distribution XttestX_{t}^{test}, we assume we only have a single data-instance xtix_{t}^{i} which is actually the case during clinical deployment. This makes On-the-Fly adaptation episodic as it resets to original weight for adapting to each data-instance. Also, we constrain the setting to not perform any back-propagation during the test-phase as it involves the availability of compute power or some cloud resource during testing. This makes On-the-Fly adaptation zero-shot as it does not really involve any gradient back-propagation during test-time. The setting of On-the-Fly adaptation is summarized in Table 1 and compared against other frameworks.

Method - Adaptive UNet

To solve On-the-Fly Adaptation, we propose an Adaptive UNet framework where we make use of adaptive batch normalization and a domain prior.

Network Details: We follow the skeleton of a generic UNet architecture . We use 5 conv blocks in both the encoder and decoder, respectively. Each conv block in the encoder consists of a conv layer, adaptive batch normalization, ReLU activation and a max-pooling layer. Each conv block in the decoder consists of a conv layer, adaptive batch normalization, ReLU activation and an upsampling layer. For upsampling, we use bilinear interpolation. For our experiments on 3D volumes, we use a 3D UNet architecture with the same setup replacing 2D conv layer with 3D conv layers, 2D max-pooling with 3D max-pooling and bilinear upsampling with trilinear upsampling.

Adaptive Batch Normalization: Batch Normalization (BN) layers are used in DNNs to mitigate the issue of internal co-variate shifts. It normalizes the features in the network helping in training and faster convergence. BN can be defined as:

where xx is the input batch, zz represents the output, μ\mu represents the mean E[X]E[X], σ\sigma represents standard deviation Var(X)\sqrt{Var(X)}. Here, γ\gamma and β\beta are learnable parameters which control the scaling and shifting while normalizing.

Adaptive Instance normalization (AdaIN) is used to align the mean and standard deviation of two feature codes (usually one being context and another being style). AdaIN can be defined as:

where xx and yy are the two feature codes and zz is the normalized output.

In Adaptive UNet, we make use of Adaptive Batch Normalization (AdaBN) which basically learns scaling and shifting operation while adaptively normalizing the batch statistics between two codes. AdaBN can be defined as:

Note that our formulation is a bit different from AdaBN as explained in as it tries to shift the model to test data’s mean and standard deviation instead of aligning them. Here, we align the codes while also learning how to align them by learning the shift and scale parameters. In Adaptive UNet, the input to AdaBN layers are the feature codes from the UNet represented as XX and domain codes YY generated from the domain prior generator.

Domain Prior Generator: The Domain Prior Generator (DPG) is an encoder of a pre-trained auto-encoder. We first pre-train a UNet as an auto-encoder for medical images. This task is self-supervised as we just try to predict the original image while feeding an augmented version of the data as input. Doing this helps the model learn an abstract code in the latent space. We train the model on a variety of medical data consisting of different modalities. More details can be found in the supplementary file. We make sure that the distribution of data that we conduct experiments to validate Adaptive UNet do not overlap with the data that the auto-encoder is trained on. However, it does have the data of similar modality. This helps the encoder generate different domain codes for different modalities. For example, two images of T1 MRI would have their corresponding domain codes closer in the latent space when compared to T1 MRI and T2 MRI.

Training source model: During the training phase of source model, we feed in the input image to both the encoder of UNet and the pre-trained domain prior generator. The domain code obtained from the domain prior generator is passed to the AdaBN layers in the encoder and decoder. The feature maps are normalized according to the domain code using AdaBN. So, in the training itself the model has learned to adapt to the domain code of the current modality/distribution. The learnable parameters γ\gamma and β\beta of AdaBN layers learn the scale and shift necessary to adapt to the style code at each level to provide the optimal segmentation prediction. Note that the weights are updated only for the UNet segmentation network. The pre-trained domain prior generator is frozen during training the source model.

Inference on target data instance: When a model trained on XstrainX_{s}^{train} is validated on a target data instance xtx_{t}, we pass the image xtx_{t} to both the domain prior generator and the source model. First, we generate the new domain code using the domain prior generator for the new image xtx_{t}. Next we pass this domain code to all the AdaBN layers in Adaptive UNet. During feed forward, the features extracted at each layer of Adaptive UNet are adapted to the new domain code. So, the model is thus adapted according to the code of the new modality/target domain. There is no back-propagation involved as the features are adapted in feed-forward itself. Also, as the model weights are not changed, this framework is episodic and does not depend on the entire test data distribution for validation. An overview of the framework is illustrated in Fig. 2.

Experiments and Results

Datasets: For 2D experiments, we focus on the task of retinal vessel segmentation from fundus images. We make use of the following datasets: CHASE , RITE and HRF . CHASE contains 28 retina images with a resolution of 999×\times960 collected from 14 school children with a hand-held Nidek NM-200-D fundus camera. RITE consists of 40 images of resolution 768×\times584 collected from people aging from 25 to 90 using a Canon CR5 non-mydriatic 3CCD camera. HRF contains 18 images collected from 18 human subjects using a Canon CR-1 fundus camera of around resolution 3504×\times2336. There exists a domain shift among these datasets as they vary with respect to camera properties, age of patients and resolution etc. The datasets are separated into a randomized 80-20 split wherever test split was not given. For 2D experiments, DPG is pre-trained on fundus images from

For 3D experiments, we focus on brain tumor segmentation from MRI volumes. We make use of the BraTS 2019 dataset which consists of four modalities- FLAIR, T1, T1ce and T2. We study the domain shift problems between these four modalities for volumetric segmentation of brain tumor. This is a multi-class segmentation problem with 4 labels. We randomly split the dataset into 266 for training and 69 for validation. We do this as the ground truth is not provided publicly for the original validation dataset. For the MRI experiments, we pre-train DPG on MRI images from Kaggle MRI dataset and IXI dataset .

Implementation Details: We use Pytorch framework for implementing Adaptive UNet. For 2D experiments, we use a combination of binary cross entropy (BCE) and dice loss to train Adaptive UNet. The loss L\mathcal{L} between the prediction y^\hat{y} and the target yy is formulated as:

We use an Adam optimizer with a learning rate of 0.0001 and momentum of 0.9. We also use a cosine annealing learning rate scheduler with a minimum learning rate upto 0.00001. The batch size is set equal to 8. For 3D experiments, we use a similar loss but use a learning rate of 0.001 while also reducing the batch size to 2. More details can be found in code and the supplementary file.

Performance Comparison: We compare our proposed method with recent test-time adaptation methods like TENT , Hu et al. (RN+CR loss) , and self domain adapted network (SDA) . In Table 2, we present the results of 2D experiments for 6 different domain shifts in fundus image. In Table 3, we present the results of 3D experiments for 3 different domain shifts in MRI modality. In both the tables, the first row corresponds to the direct adaptation results where we train the model on source domain and report the results of those models while testing on the target domain without any adaptation. The last row corresponds to the oracle which is the maximum possible performance when the model is trained on the target train distribution and tested on the target test distribution. Note that the 3D experiments have two target-training configurations- Uni-modal and Multi-modal. Uni-modal oracle corresponds to the configuration where we only use one modality and multi-modal oracle corresponds to the case where we use all four modalities to train the model. The compared test-time adaptation methods are presented in both one-shot and ten-shot settings. In one-shot setting, the model weights are adapted by back-propagation for one epoch using the test distribution. In ten-shot setting, the model weights are adapted for ten epochs by back-propagation using the test distribution. For our proposed method, we adapt the model once per image during feed-forward.

From Tables 2 and 3, it can be inferred that there is a considerable drop in performance while directly testing the source model on the target domain. This drop is expected as there exits a domain shift between the source and the target distribution. SDA does not perform well, especially in relatively small datasets (Table 2). Since SDA requires training a set of auto-encoders to provide supervision during test-time adaption, only training on a small amount of data may result in overfitting and consequently lower adaptation performance. TENT and RN+CR methods improve the performance in most cases as they try to reduce the entropy and regularize the batch-norm statistics for the target distribution. Our proposed method shows a considerable improvement over the direct testing as well as test-time baselines on almost all domain shifts achieving state-of-the-art adaptation performance. Note that brain tumor segmentation from 3D volumes is a multi-class segmentation problem which shows that our method can be successfully adopted for multi-class problems as well.

We also provide sample qualitative results in Fig. 3. It can be observed that without any adaptation, the segmentation predictions are noisy and contain over-segmentation. While recent test-time methods improve the prediction, they still suffer from mis-classification of pixels (see MRI predictions in Fig. 3 for TENT/RN+CR) and also over-segmentation (see fundus predictions in Fig. 3 for TENT/RN+CR). Our method achieves good segmentation prediction that is very close to the oracle prediction and the ground-truth.

Discussion: From our experiments, we find that our method works pretty well for 2D experiments when compared to 3D experiments. This observation can be understood as the 2D experiments consider domain shifts within the same modality with differences in camera/sensor properties and type of patient. The 3D experiments consider cross-modality domain shifts where the MRI sequences are themselves different. This is a more difficult task as each sequence extracts different types of features. However, we get a considerable boost over other test-time methods even though our setting is episodic and zero-shot.

Conclusion

In this work, we propose a new adaptation setting called On-the-Fly adaptation. In this setting, we constrain the adaptation to be episodic and zero-shot thus assuming no availability of the entire target distribution and model update during test-time. This makes On-the-Fly adaptation very close to the scenario while deploying deep-learning models in real-world clinical settings. We propose a new framework- Adaptive UNet to solve this adaptation problem by making using of adaptive batch normalization and domain priors. We validate our model on both 2D and 3D domain shifts of fundus images and MRI volumes and show that the proposed method achieves a competitive performance over the recent test-time adaptation methods even with the tighter constraint of On-the-Fly adaptation.

References