More Training Data by Augmentation

Segmentation models learn from example images with matching masks, and labelled microscopy images are scarce. Data augmentation [12] creates new training pairs by rotating, scaling, and skewing images we already have.

This notebook generates augmented image–mask pairs from a small set of labelled SEM images of TiO₂ (TiO2_EM) using the Aug method.

JuSPICE Modules and Classes

  • juspice.io: SPICEData, load_data, save_data

  • juspice.synth_data_module: ConfigLoader, SynthGenerationAccessor (via spice.synth_generation)

  • juspice.tracking: Tracker

[1]:
from juspice.tracking import Tracker

tracker = Tracker(
    include_metadata=True, notes='Augmentation-based synthetic data generation'
)

tracker.recording_start()
/Users/amir/GIT_repositories/juspice_pre_release/venvs/.venv/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 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 SPICEData, load_data, save_data
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))

# Logging configuration
logging.basicConfig(
    level=logging.INFO,
    format='%(asctime)s - %(levelname)s - %(message)s',
    datefmt='%Y-%m-%d %H:%M:%S',
)
[3]:
tracker.recording_stop()

Settings

We choose the fastest available device (GPU or CPU) and read the augmentation parameters from notebooks_parameters.yaml.

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

_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)
logging.info(f'Configuration loaded from: {config_path}')

dataset = 'TiO2_EM'
method_name = 'Aug'

aug_params = config.get_gen_method_params(method_name)
logging.info('Augmentation Parameters:')
logging.info(str(aug_params))
20:08:2026 16:39:14 - Using device: mps.
20:08:2026 16:39:14 - Configuration loaded from: /Users/amir/GIT_repositories/juspice_pre_release/notebooks/notebooks_parameters.yaml
20:08:2026 16:39:14 - Augmentation Parameters:
20:08:2026 16:39:14 - {'width': 320, 'height': 320, 'min_scale': 0.8, 'max_scale': 1.5, 'min_rotation': -40, 'max_rotation': 40, 'min_shear': -20, 'max_shear': 20, 'max_exposed_fraction': 0.12, 'inpaintRadius': 3, 'num_repeats_per_image': 2}
[6]:
tracker.recording_stop()

How augmentation works

A particle remains the same particle when it is rotated, slightly resized, or skewed. Showing a model such variations helps it recognise objects regardless of their orientation and size.

The Aug method applies three affine transformations with random values from the configured ranges:

Transformation

Effect

Parameters

Rotation

turns the image

min_rotation, max_rotation (degrees)

Scaling

zooms in or out

min_scale, max_scale (1.0 = no change)

Shear

skews along one axis

min_shear, max_shear (degrees)

The same transformation is applied to the image and its mask, so they stay aligned pixel by pixel. Each image is augmented num_repeats_per_image times.

Rotating or shrinking an image leaves empty corners. These are filled with the LaMa inpainting model [50] via IOPaint [51] so the model does not learn black borders as object edges. If too much of the image would be empty (max_exposed_fraction), the transformation is weakened and retried.

Compared with the other methods: augmentation is the simplest and fastest, but it needs labelled images and cannot create shapes that are not already present. Physics-based generation (PB) needs no labels, and deep-learning methods (DCGAN, SDiff) produce more varied textures at a higher computing cost.

We now load one training image as the starting point and generate the new pairs.

[7]:
tracker.recording_start()
[8]:
# --------- Block 2: Initialize and generate ----------

# Load one training image as the SPICEData entry point / history anchor —
# spice.synth_generation.generate() doesn't transform spice.data itself
# (it produces a brand-new dataset), but the returned dataset shares this
# object's history lineage, so the provenance reads "generated from this
# reference training set".
dataset_params = config.get_dataset_params(dataset)
training_images_dir = os.path.join(repo_root, dataset_params['training_images_dir'])
_img_exts = ('.png', '.jpg', '.jpeg', '.tif', '.tiff')
_seed_candidates = sorted(
    f for f in os.listdir(training_images_dir) if f.lower().endswith(_img_exts)
)
seed_image_name = _seed_candidates[0]
# `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(os.path.join(training_images_dir, seed_image_name))

N_images = 5
print(f'Generating synthetic images using {method_name}...')
synth_spice_aug = spice.synth_generation.generate(
    config=config,
    input_dataset=dataset,
    method_name=method_name,
    repo_root=repo_root,
    N_images=N_images,
    device=device,
)
print(f'Generation complete! dataset_type={synth_spice_aug.dataset_type!r}, '
      f'n_frames={synth_spice_aug.n_frames}')
20:08:2026 16:39:14 - Parameters for the model were set.
20:08:2026 16:39:14 - Preprocessing training images...
Generating synthetic images using Aug...
Preprocessing images: 100%|██████████| 5/5 [00:00<00:00, 60.33it/s]
20:08:2026 16:39:14 - Training images and masks were preprocessed.
20:08:2026 16:39:14 - Test images and masks were preprocessed.
20:08:2026 16:39:14 - *******************
Preprocessing complete.
Images saved to: /Users/amir/GIT_repositories/juspice_pre_release/preprocessed_data/Aug_TiO2_EM/input_images
Masks saved to: /Users/amir/GIT_repositories/juspice_pre_release/preprocessed_data/Aug_TiO2_EM/input_masks
Generating synthetic images: 100%|██████████| 5/5 [00:32<00:00,  6.59s/it]
20:08:2026 16:39:47 - Synthetic image and masks were generated and saved.
Generation complete! dataset_type='multiple_frames', n_frames=10
[9]:
tracker.recording_stop()

Inspect the results

Each generated image (top) has a matching mask (bottom).

[10]:
tracker.recording_start()
[11]:
# --------- Block 3: Visualize generated images ----------

output_masks_dir = synth_spice_aug.metadata['masks_dir']
n_show = min(6, synth_spice_aug.n_frames)

fig, axes = plt.subplots(2, n_show, figsize=(18, 6))
fig.suptitle('Augmented Synthetic Images and Masks', fontsize=16, fontweight='bold')

for idx in range(n_show):
    img_data = synth_spice_aug.frames[idx]
    img_file = os.path.basename(synth_spice_aug.frame_paths[idx])
    mask_path = os.path.join(output_masks_dir, img_file)
    axes[0, idx].imshow(img_data, cmap='gray')
    axes[0, idx].set_title(f'Image {idx+1}')
    axes[0, idx].axis('off')
    if os.path.exists(mask_path):
        mask_data = load_data(mask_path).data
        axes[1, idx].imshow(mask_data, cmap='gray')
    axes[1, idx].axis('off')
plt.tight_layout()
plt.show()
../../_images/notebooks_synthetic_data_example_Aug_14_0.png

Comparing an original with its augmented versions shows the rotation, scaling, and shear, as well as the filled corners.

[12]:
# --------- Block 4: Compare original vs augmented ----------

input_images_dir = synth_spice_aug.metadata['input_images_preprocessed_dir']
input_files = sorted(
    f for f in os.listdir(input_images_dir)
    if f.lower().endswith(('.jpeg', '.jpg', '.png', '.tif', '.tiff'))
)

if input_files:
    original_img_path = os.path.join(input_images_dir, input_files[0])
    original_img = load_data(original_img_path).data
    base_name = os.path.splitext(input_files[0])[0]
    image_files = [os.path.basename(p) for p in synth_spice_aug.frame_paths]
    augmented_files = [f for f in image_files if base_name in f]

    fig, axes = plt.subplots(1, min(3, 1 + len(augmented_files)), figsize=(12, 4))
    axes[0].imshow(original_img, cmap='gray')
    axes[0].set_title('Original')
    axes[0].axis('off')
    for i, af in enumerate(augmented_files[:2], start=1):
        aug_idx = image_files.index(af)
        aug_data = synth_spice_aug.frames[aug_idx]
        axes[i].imshow(aug_data, cmap='gray')
        axes[i].set_title(f'Augmented {i}')
        axes[i].axis('off')
    plt.tight_layout()
    plt.show()
../../_images/notebooks_synthetic_data_example_Aug_16_0.png
[13]:
# --------- Block 5: Save ----------

if synth_spice_aug.n_frames == 0:
    raise RuntimeError('No generated images found to export.')

# Writes <stem>.npy, <stem>.json, and <stem>_history.py next to the notebook
# 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_Aug'
save_data(
    synth_spice_aug,
    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/synthetic_data/example_Aug.npy
JSON sidecar:   /Users/amir/GIT_repositories/juspice_pre_release/notebooks/synthetic_data/example_Aug.json
History script: /Users/amir/GIT_repositories/juspice_pre_release/notebooks/synthetic_data/example_Aug_history.py
[14]:
# --------- Block 6: Validate + Inspect human-readable history ----------

# synth_spice_aug.history.to_lines() reconstructs the full synthetic data
# generation pipeline for this SPICEData object into readable, runnable code:
# the initial load_data(...) call followed by synth_generation.generate() and
# a trailing save_data(synth_spice_aug).
readable_lines = synth_spice_aug.history.to_lines()

print('Reconstructed synthesis pipeline for this SPICEData object:')
print('\n'.join(readable_lines), '...')
Reconstructed synthesis pipeline for this SPICEData object:
import juspice
import juspice.synth_data_module
spice = juspice.io.load_data('/Users/amir/GIT_repositories/juspice_pre_release/Sample_data/Train_test_images/TiO2_EM/images/train/1908248.tif')
synth_spice = spice.synth_generation.generate(config=juspice.synth_data_module.ConfigLoader('/Users/amir/GIT_repositories/juspice_pre_release/notebooks/notebooks_parameters.yaml'), input_dataset='TiO2_EM', method_name='Aug', repo_root='/Users/amir/GIT_repositories/juspice_pre_release', N_images=5, device='mps')
juspice.io.save_data(spice) ...
[15]:
tracker.recording_stop()