Skip to article frontmatterSkip to article content
Site not loading correctly?

This may be due to an incorrect BASE_URL configuration. See the MyST Documentation for reference.

Training a U-Net with BiaPy

Pretrained models such as StarDist and Cellpose find nuclei in a DNA stain. Here we ask for something they were not made for: find the nuclei in the actin (phalloidin) channel. If that works, we no longer need the DNA stain and can use that channel for another marker.

We train first and then look at the results and different metrics. to point at.

The images come from 00_data_collection_stardist.ipynb. You get them as a zip, so you do not need to run that notebook.

Kernel: biapy on Research Cloud, nlbi26-day3-biapy on your own laptop.

1. The data

Unzip training_data next to this notebook. Two folders with matching file names: images/ is the phalloidin channel, labels/ are the nuclei StarDist found in the DAPI channel. Open both in the file browser on the left and check that the names match.

100 images, 100 label images

Let plot one of the images, you could try look at another image as well.

<Figure size 1200x500 with 2 Axes>

Nobody checked all 100 label images that were automatically created. Let’s count the nuclei in each of the image before we train on them.

count    100.0
mean     107.0
std       38.4
min        0.0
25%       85.8
50%      111.0
75%      135.5
max      173.0
dtype: float64

empty: ['1897201.tif', '1898947.tif']

Two images have no nuclei. They teach the network nothing and they break the scores later (you cannot ask how many of the nuclei we found when there are none), so leave them out of the dataset you build below.

Correct a label image in napari

The labels come from StarDist, so the model can at best learn StarDist’s mistakes too. Open one image in napari and fix its labels, as you did this morning. Then save it, before you build the dataset in section 2.

If import napari fails, run %pip install "napari[all]" once in a new cell and restart the kernel.

2. Splitting the data

Three piles:

  • training — what the network learns from

  • validation — checked every epoch, to see whether it is still improving or has started memorising the training images

  • test — touched once, at the end. It is the only honest number you get, and the moment you use it to pick a setting it stops being one.

Images that are nearly the same must not end up in different piles: frames of a time lapse, slices through one organoid, fields from one patient. Otherwise the test score tells you how well the model remembers, not how well it works. metadata.csv tells you where each image comes from.

Can you split these images at random? What would you do if there were four fields per well?

Exercise: build the dataset

Make this folder structure next to this notebook, from the files in training_data/:

dataset/
├── train/
│   ├── images/   # phalloidin
│   └── labels/   # nuclei
├── val/
│   ├── images/
│   └── labels/
└── test/
    ├── images/
    └── labels/
  • Put about 20% of the images in val/, 20% in test/, the rest in train/.

  • An image and its label keep the same file name and go to the same split.

  • Leave out the images without nuclei.

Do it by hand in the file browser, or write a short script (pathlib and shutil.copy are enough). Then run the cell below to check your result.

3. The settings

BiaPy reads a YAML file. We take the template BiaPy publishes for 2D instance segmentation, change a few settings from Python, and write it back out.

# BiaPy version: 3.7.1

SYSTEM:
    NUM_CPUS: -1

PROBLEM:
    TYPE: INSTANCE_SEG
    NDIM: 2D
    INSTANCE_SEG:
        DATA_CHANNELS: BC
  
DATA: 
    PATCH_SIZE: (256, 256, 1)  
    TRAIN:                                                                                                              
        PATH: /path/to/data
        GT_PATH: /path/to/data
        IN_MEMORY: True
    VAL:
        SPLIT_TRAIN: 0.1
    TEST:                                                                                                               
        PATH: /path/to/data
        GT_PATH: /path/to/data
        IN_MEMORY: True
        LOAD_GT: True
        PADDING: (32,32)

AUGMENTOR:
    ENABLE: True
    RANDOM_ROT: True
    VFLIP: True
    HFLIP: True

MODEL:
    ARCHITECTURE: unet
    LOAD_CHECKPOINT: False

TRAIN:
    ENABLE: True
    OPTIMIZER: ADAMW
    LR: 1.E-4
    BATCH_SIZE: 8
    EPOCHS: 100
    PATIENCE: 20

# Loss function
LOSS:
    CLASS_REBALANCE: True # give the same weight to all problem representation channels

TEST:
    ENABLE: True
    AUGMENTATION: False
    FULL_IMG: False

Read that file into a dictionary, change what we need, save it under a new name. The template stays untouched so you can always start over.

SYSTEM:
  NUM_CPUS: -1
PROBLEM:
  TYPE: INSTANCE_SEG
  NDIM: 2D
  INSTANCE_SEG:
    DATA_CHANNELS:
    - F
DATA:
  PATCH_SIZE: (256, 256, 1)
  TRAIN:
    PATH: dataset/train/images
    GT_PATH: dataset/train/labels
    IN_MEMORY: true
  VAL:
    SPLIT_TRAIN: 0.2
    FROM_TRAIN: true
  TEST:
    PATH: dataset/test/images
    GT_PATH: dataset/test/labels
    IN_MEMORY: true
    LOAD_GT: true
    PADDING: (32,32)
AUGMENTOR:
  ENABLE: true
  RANDOM_ROT: true
  VFLIP: true
  HFLIP: true
MODEL:
  ARCHITECTURE: unet
  LOAD_CHECKPOINT: false
TRAIN:
  ENABLE: true
  OPTIMIZER: ADAMW
  LR: 0.0001
  BATCH_SIZE: 8
  EPOCHS: 20
  PATIENCE: 20
LOSS:
  CLASS_REBALANCE: true
TEST:
  ENABLE: true
  AUGMENTATION: false
  FULL_IMG: false

4. Train

The output is long, so we use %%capture keeps it out of the notebook. Run biapy_out.show() in a new cell if you want to read it.

[13:52:18.931082] aug
[13:52:18.931168] charts
[13:52:18.931184] instance_associations
[13:52:18.931197] per_image
[13:52:18.931209] per_image_instances
[13:52:18.931222] tensorboard
[13:52:18.931235] test_F_instance_channels
[13:52:18.931247] test_results_metrics.csv
[13:52:18.931260] train_F_instance_channels

5. Check the training graphs

BiaPy automatically plots the graphs and saves them in charts/.

<IPython.core.display.Image object>
<IPython.core.display.Image object>

The loss is how wrong the network is, so lower is better. Compare the training and validation curves:

  • both still going down at the last epoch: it can train longer

  • both flat for a while: it has learned what it can

  • training loss going down while validation loss goes up: it is memorising the training images. Stop earlier, augment more, or get more data.

Look at your own curves and decide which of these it is before changing anything.

The IoU plot shows the overlap between prediction and target, per channel.

6. What does it predict?

TEST.ENABLE is on in the template, so BiaPy already ran the model on the test images.

<Figure size 1600x500 with 3 Axes>
<Figure size 1600x500 with 3 Axes>
<Figure size 1600x500 with 3 Axes>

Look at a dense patch: do you see two nuclei with one colour? Then they were segmented as one object.

BiaPy also scored the test images. The scores are in test_results_metrics.csv in the results folder.

The simplest score: how many objects. If your question is how many cells are in a well, this may be all you need.

[13:52:50.204847]          image  nuclei in the labels  objects found
0  1895788.tif                   137            158
1  1895923.tif                   135            154
2  1895950.tif                   107            120
3  1896103.tif                     2            114
4  1896670.tif                    35             76
5  1896787.tif                   131            149
6  1896841.tif                   112            138
7  1897129.tif                    39             62
8  1897615.tif                    68             87
9  1897687.tif                    83            114
[13:52:50.208211] 
[13:52:50.208591] nuclei in the labels    1950
objects found           2450
dtype: int64

But a count says nothing about where the objects are. To decide whether a predicted object is the same as a labelled one we can compare their overlapping pixels: intersection over union (IoU), the pixels they share divided by the pixels they cover together. 1.0 is identical, 0.5 is a reasonable overlap, 0.0 is no overlap at all.

<Figure size 1100x400 with 3 Axes>

7. How to separate merged nuclei?

A U-Net takes an image and returns an image. It cannot output “nucleus 17”: the numbers in the label image are arbitrary, and nothing tells the network that this blob should be 17 and the next one 18. So BiaPy turns the labels into target images the network can learn, and rebuilds the objects from the prediction afterwards. Those targets are the data channels we set in section 3.

With ['F'] (foreground) the network only learns nucleus or not. Two touching nuclei become one blob. With ['F', 'C'] it also learns the contour of each nucleus, and BiaPy uses that to cut touching nuclei apart.

8. Run 2: add the contour to the target

[13:54:26.030152] ['F', 'C'] -> run2_foreground_contour
Loading...
<Figure size 1600x500 with 3 Axes>

Did the contour change the number of objects, their outlines, or both? Which column tells you that? Did FN go down? That is where the merged nuclei were hiding.

Both runs got 20 epochs. Two channels is a harder job than one, so open run 2’s charts and see whether it had finished learning.

9. If you have time

Each of these is one line, then re-run the config cell, the training cell and the scoring cells. Give every run its own job_name. They take as long as the runs above, so pick one or two.

config['PROBLEM']['INSTANCE_SEG']['DATA_CHANNELS'] = ['F', 'D']   # distance instead of contour
config['TRAIN']['EPOCHS'] = 60                                    # if the loss was still falling
config['AUGMENTOR']['ENABLE'] = False                             # it is on in the template
config['DATA']['PATCH_SIZE'] = '(128, 128, 1)'                    # smaller crops, less context
config['MODEL']['ARCHITECTURE'] = 'resunet'                       # a different network

The last one is what people try first and it is rarely where the gain is. Compare it with what the contour channel did.

And the question worth asking before any of them: train on 20, 40 and all training images (move images out of dataset/train/, then re-run the check cell) and plot F1 against the number of training images. If that curve is still rising, annotating more images beats every setting on this page.