Segmentation with a Microscopy-Pretrained Model (NASA MicroNet)

A segmentation model learns faster from few labelled images if it has already seen similar images before. Here we start from an encoder pretrained on NASA’s MicroNet microscopy images and adapt it to our own data (transfer learning).

The task is to find the oxide layer in cross-sectional SEM images of an environmental barrier coating (the EBC1 dataset), using spice.segmentation.run_micronet().

⚠️ Separate environment required

This notebook needs its own virtual environment with pretrained_microscopy_models and its dependencies:

source venvs/.venv_micronet/bin/activate

See the installation guide for how to create it.

JuSPICE Modules and Classes

  • juspice.io: load_data, save_data, SPICEData

  • juspice.segmentation_module: UNetSegmentationAccessor (via spice.segmentation)

  • juspice.synth_data_module: ConfigLoader

  • juspice.tracking: Tracker

[1]:
from juspice.tracking import Tracker

tracker = Tracker(include_metadata=True, notes='Segmentation with NASA MicroNet')

tracker.recording_start()
/Users/amir/GIT_repositories/juspice_pre_release/venvs/.venv_micronet/lib/python3.12/site-packages/tqdm/auto.py:21: TqdmWarning: IProgress not found. Please update jupyter and ipywidgets. See https://ipywidgets.readthedocs.io/en/stable/user_install.html
  from .autonotebook import tqdm as notebook_tqdm
[2]:
# --------- Block 0: Setup ----------

# Standard library
from pathlib import Path
import sys
import os
import json
import logging
import warnings
import shutil

# Third-party
import numpy as np
import matplotlib.pyplot as plt
try:
    import torch
    _torch_available = True
except ImportError:
    torch = None
    _torch_available = False
    print("WARNING: torch is not installed. Device selection will default to 'cpu'.")

# JuSPICE
from juspice.io import load_data, save_data, SPICEData
from juspice.synth_data_module import ConfigLoader

# Paths
try:
    notebook_dir = Path(__file__).resolve().parent
except Exception:
    notebook_dir = Path.cwd()

cur = notebook_dir
repo_root = None
for _ in range(6):
    if (cur / 'juspice').exists() or (cur / 'pyproject.toml').exists():
        repo_root = cur
        break
    if cur.parent == cur:
        break
    cur = cur.parent
if repo_root is None:
    repo_root = notebook_dir
if str(repo_root) not in sys.path:
    sys.path.insert(0, str(repo_root))

[3]:
tracker.recording_stop()

Why start from MicroNet?

An encoder pretrained on everyday photos (ImageNet [40]) knows general edges and textures. MicroNet is a collection of more than 100,000 labelled microscopy images from 54 material classes (Stuckner et al., 2022) [38]; the pretrained models are available from NASA [23]. An encoder pretrained on it has already learned typical microscopy textures, which can help when only a few labelled images are available.

Model. We use UNet++ (Zhou et al., 2018) [37], a U-Net variant with additional skip connections between scales, which helps with objects of different sizes. Its encoder, SE-ResNeXt50, is a larger ResNet variant (ResNeXt [41]) with squeeze-and-excitation blocks [42] that learn which feature channels are most important for each image.

Training. The model is trained on the training images and scored on the validation images after every epoch. The checkpoint with the best validation IoU (overlap between predicted and true mask) is kept. Training stops early if the score does not improve for patience epochs.

For a quick demo, training is limited to 3 epochs, so early stopping (patience=30) never triggers and the model is far from fully trained. The validation set is also small, so its IoU fluctuates strongly between epochs.

Next we set the dataset paths and run the full pipeline with a single call. The figures show the prediction for the first test image and the training and validation loss per epoch.

[4]:
tracker.recording_start()
[5]:
# --------- Block 1: Config, device, and dataset paths ----------

_mps_available = (
    torch is not None
    and hasattr(torch.backends, 'mps')
    and torch.backends.mps.is_available()
)
if torch is not None and torch.cuda.is_available():
    device = 'cuda'
elif _mps_available:
    device = 'mps'
else:
    device = 'cpu'
logging.info(f'Using device: {device}')

config_path = os.path.join(repo_root, 'notebooks', 'notebooks_parameters.yaml')
config = ConfigLoader(config_path)

ebc1_params = config.get_dataset_params('EBC1')
x_train_dir = os.path.join(repo_root, ebc1_params['training_images_dir'])
y_train_dir = os.path.join(repo_root, ebc1_params['training_masks_dir'])
x_valid_dir = os.path.join(repo_root, ebc1_params['val_images'])
y_valid_dir = os.path.join(repo_root, ebc1_params['val_masks'])
x_test_dir  = os.path.join(repo_root, ebc1_params['test_images'])
y_test_dir  = os.path.join(repo_root, ebc1_params['test_masks'])

_img_exts = ('.png', '.jpg', '.jpeg', '.tif', '.tiff')

def _list_images(directory):
    """Return sorted absolute image file paths, skipping hidden/non-image files."""
    return sorted(
        os.path.join(directory, f)
        for f in os.listdir(directory)
        if f.lower().endswith(_img_exts) and not f.startswith('.')
    )

train_images = _list_images(x_train_dir)
train_masks  = _list_images(y_train_dir)
valid_images = _list_images(x_valid_dir)
valid_masks  = _list_images(y_valid_dir)
test_images  = _list_images(x_test_dir)
test_masks   = _list_images(y_test_dir)

pretrained_model_dir = os.path.join(
    repo_root, config['pretrained_models_dir'], 'NASA_Micronet'
)

# Load one training image as the SPICEData entry point / history anchor.
# `spice` is just a variable name for the SPICEData instance load_data()
# returns here — any name would work; we use `spice` throughout these
# notebooks as an intuitive nod to JuSPICE / SPICEData.
spice = load_data(train_images[0])

# run_micronet() combines model creation, dataset preparation, training
# with early stopping, and first-test-image inference into a single call —
# analogous to spice.synth_generation.generate().
# epochs=3 keeps this demo run well under a few minutes on CPU (roughly
# 90s/epoch on the full EBC1 set with se_resnext50_32x4d — the default,
# unbounded except by patience=30 early stopping, can run far longer).
# For real training, raise epochs (or drop it back to None for pure
# patience-based early stopping) and ideally use a GPU.
pred_spice = spice.segmentation.run_micronet(
    train_images=train_images,
    train_masks=train_masks,
    val_images=valid_images,
    val_masks=valid_masks,
    test_images=test_images,
    test_masks=test_masks,
    architecture='UnetPlusPlus',
    encoder='se_resnext50_32x4d',
    pretrained_weights='micronet',
    pretrained_model_dir=pretrained_model_dir,
    device=device,
    epochs=3,
    patience=30,
    lr=2e-4,
    batch_size=6,
    val_batch_size=6,
)

print(f'pred_spice.data.shape: {pred_spice.data.shape}')
print(f'architecture: {pred_spice.metadata["architecture"]}')
print(f'encoder: {pred_spice.metadata["encoder"]}')

plt.figure(figsize=(6, 5))
plt.imshow(pred_spice.data, cmap='gray')
plt.title('Predicted Segmentation (first test image)')
plt.axis('off')
plt.tight_layout()
plt.show()

train_loss_hist = pred_spice.metadata.get('train_loss_history', [])
val_loss_hist   = pred_spice.metadata.get('val_loss_history', [])
if train_loss_hist:
    plt.figure()
    plt.plot(train_loss_hist, label='train_loss')
    plt.plot(val_loss_hist, label='valid_loss')
    plt.legend()
    plt.xlabel('epoch')
    plt.ylabel('loss')
    plt.title('MicroNet Training Loss')
    plt.tight_layout()
    plt.show()
20:08:2026 18:16:47 - Using device: mps

Epoch: 0, lr: 0.00020000, time: 0.00 seconds, patience step: 0, best iou: 0.0000
train:   0%|          | 0/3 [00:00<?, ?it/s]
/Users/amir/GIT_repositories/juspice_pre_release/venvs/.venv_micronet/lib/python3.12/site-packages/torch/utils/data/dataloader.py:759: UserWarning: 'pin_memory' argument is set as true but not supported on MPS now, device pinned memory won't be used.
  super().__init__(loader)
train: 100%|██████████| 3/3 [00:12<00:00,  4.10s/it, DiceBCELoss - 0.7969, iou_score - 0.2504]
valid: 100%|██████████| 1/1 [00:00<00:00,  1.33it/s, DiceBCELoss - 0.6415, iou_score - 0.01376]
Best model saved!

Epoch: 1, lr: 0.00020000, time: 14.46 seconds, patience step: 0, best iou: 0.0138
train: 100%|██████████| 3/3 [00:06<00:00,  2.19s/it, DiceBCELoss - 0.7419, iou_score - 0.4188]
valid: 100%|██████████| 1/1 [00:00<00:00,  2.46it/s, DiceBCELoss - 0.6038, iou_score - 0.6239]
Best model saved!

Epoch: 2, lr: 0.00020000, time: 8.39 seconds, patience step: 0, best iou: 0.6239
train: 100%|██████████| 3/3 [00:06<00:00,  2.24s/it, DiceBCELoss - 0.7048, iou_score - 0.5077]
valid: 100%|██████████| 1/1 [00:00<00:00,  2.81it/s, DiceBCELoss - 0.5589, iou_score - 0.8485]
Best model saved!


Training done! Saving final model
pred_spice.data.shape: (512, 512)
architecture: UnetPlusPlus
encoder: se_resnext50_32x4d
../../_images/notebooks_segmentation_example_nasa_micronet_6_4.png
../../_images/notebooks_segmentation_example_nasa_micronet_6_5.png

Check the fit on training images

We compare a few training images, their true masks, and the model’s predictions. Good agreement here only shows that the model fits the data it has seen; the test prediction above is the fairer check.

[6]:
# --------- Block 1b: Training images vs. predicted segmentation ----------

# pred_spice.extra['model'] is the trained, best-checkpoint model
# run_micronet() already produced — inference below reuses it directly,
# no retraining or reloading the checkpoint from disk needed.
import pretrained_microscopy_models as pmm
from juspice.segmentation_module import get_validation_augmentation, get_preprocessing

trained_model = pred_spice.extra['model']
preprocessing_fn = pred_spice.extra['preprocessing_fn']
trained_model.eval()

n_examples = min(3, len(train_images))
# get_validation_augmentation() (no random augmentation) keeps the
# preview images faithful to what's actually on disk.
preview_dataset = pmm.io.Dataset(
    train_images[:n_examples],
    train_masks[:n_examples],
    augmentation=get_validation_augmentation(),
    preprocessing=get_preprocessing(preprocessing_fn),
    class_values={'oxide': [1]},
)

fig, axes = plt.subplots(n_examples, 3, figsize=(10, 3.2 * n_examples))
if n_examples == 1:
    axes = axes[np.newaxis, :]

col_titles = ['Training image', 'Ground-truth mask', 'Predicted segmentation']

with torch.no_grad():
    for i in range(n_examples):
        image_tensor, mask_tensor = preview_dataset[i]
        x = torch.from_numpy(image_tensor).to(device).unsqueeze(0)
        pred = trained_model.predict(x).squeeze().cpu().numpy().round()

        # image_tensor is (3, H, W), ImageNet-normalised (mean-subtracted,
        # so values go negative) — rescale to [0, 1] just for display.
        display_image = np.moveaxis(image_tensor, 0, -1)
        display_image = (display_image - display_image.min()) / (
            display_image.max() - display_image.min() + 1e-8
        )
        display_mask = np.asarray(mask_tensor).squeeze()

        for j, arr in enumerate([display_image, display_mask, pred]):
            axes[i, j].imshow(arr, cmap=None if j == 0 else 'gray')
            axes[i, j].axis('off')
            if i == 0:
                axes[i, j].set_title(col_titles[j], fontsize=10)

plt.tight_layout()
plt.show()
../../_images/notebooks_segmentation_example_nasa_micronet_8_0.png
[7]:
tracker.recording_stop()

Save and review

We save the test prediction and print the recorded pipeline.

[8]:
tracker.recording_start()
[9]:
# --------- Block 2: Save ----------

# pred_spice IS already the SPICEData returned by spice.segmentation.run_micronet()
# — no manual wrapping needed. It carries the prediction and shares spice's history.
# This folder holds several notebooks, so save_data()'s automatic
# notebook-name detection would be ambiguous when run outside a live
# Jupyter session (e.g. via nbconvert) — pass output_stem explicitly
# to always land on this notebook's own files. See
# docs/repository_structure.rst.
stem = 'example_nasa_micronet'
save_data(
    pred_spice,
    output_stem=str(notebook_dir / stem),
)
print(f'Data:           {os.path.join(notebook_dir, stem + ".npy")}')
print(f'JSON sidecar:   {os.path.join(notebook_dir, stem + ".json")}')
print(f'History script: {os.path.join(notebook_dir, stem + "_history.py")}')
Data:           /Users/amir/GIT_repositories/juspice_pre_release/notebooks/segmentation/example_nasa_micronet.npy
JSON sidecar:   /Users/amir/GIT_repositories/juspice_pre_release/notebooks/segmentation/example_nasa_micronet.json
History script: /Users/amir/GIT_repositories/juspice_pre_release/notebooks/segmentation/example_nasa_micronet_history.py
[10]:
# --------- Block 3: Inspect human-readable history ----------

# pred_spice.history.to_lines() reconstructs the full segmentation pipeline
# applied to this SPICEData object into readable, runnable code: the
# initial load_data(...) call followed by spice.segmentation.run_micronet(...)
# with all runtime argument values (architecture, encoder, dataset directories,
# device, patience, lr, etc.), and a trailing save_data(pred_spice).
readable_lines = pred_spice.history.to_lines()

print('Reconstructed segmentation pipeline for this SPICEData object:')
print('\n'.join(readable_lines))
Reconstructed segmentation pipeline for this SPICEData object:
import juspice
spice = juspice.io.load_data('/Users/amir/GIT_repositories/juspice_pre_release/Sample_data/Train_test_images/EBC1/train/010417#4_S2480009.tif')
pred_spice = spice.segmentation.run_micronet(train_images_dir='/Users/amir/GIT_repositories/juspice_pre_release/Sample_data/Train_test_images/EBC1/train', train_masks_dir='/Users/amir/GIT_repositories/juspice_pre_release/Sample_data/Train_test_images/EBC1/train_annot', val_images_dir='/Users/amir/GIT_repositories/juspice_pre_release/Sample_data/Train_test_images/EBC1/val', val_masks_dir='/Users/amir/GIT_repositories/juspice_pre_release/Sample_data/Train_test_images/EBC1/val_annot', test_images_dir='/Users/amir/GIT_repositories/juspice_pre_release/Sample_data/Train_test_images/EBC1/test', test_masks_dir='/Users/amir/GIT_repositories/juspice_pre_release/Sample_data/Train_test_images/EBC1/test_annot', architecture='UnetPlusPlus', encoder='se_resnext50_32x4d', pretrained_weights='micronet', pretrained_model_dir='/Users/amir/GIT_repositories/juspice_pre_release/Pretrained_models/NASA_Micronet', device='mps', patience=30, lr=0.0002)
juspice.io.save_data(spice)
[11]:
tracker.recording_stop()