Segmentation Without Training: Segment Anything (SAM)¶
Training a segmentation model requires many hand-labelled images. Meta’s Segment Anything Model (SAM; Kirillov et al., 2023) [19] avoids this: it was trained once on a very large, general image collection and can outline objects in new kinds of images without further training (zero-shot).
This notebook applies SAM to an SEM image in two ways: fully automatically, and guided by a single point that we choose.
JuSPICE Modules and Classes¶
juspice.io:load_data,save_data,SPICEDatajuspice.segmentation_module:UNetSegmentationAccessor(viaspice.segmentation)juspice.synth_data_module:ConfigLoaderjuspice.tracking:Tracker
[1]:
from juspice.tracking import Tracker
tracker = Tracker(include_metadata=True, notes='Zero-shot segmentation with SAM1')
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 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()
Settings¶
We choose a device and set the path to the SAM model file (checkpoint).
[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'
print(f'Using device: {device}')
config_path = os.path.join(repo_root, 'notebooks', 'notebooks_parameters.yaml')
config = ConfigLoader(config_path)
input_image_path = os.path.join(
repo_root, config['sample_images_dir'], '1b935635dd.png'
)
model_type = 'vit_h'
pretrained_model_dir = os.path.join(
repo_root, config['pretrained_models_dir'], 'segment_anything1_META'
)
pretrained_model_path = os.path.join(
pretrained_model_dir, f'sam_{model_type}_4b8939.pth'
)
pretrained_model_url = 'https://dl.fbaipublicfiles.com/segment_anything/sam_vit_h_4b8939.pth'
os.makedirs(pretrained_model_dir, exist_ok=True)
if not os.path.exists(pretrained_model_path):
logging.info(f'Downloading SAM1 checkpoint to: {pretrained_model_path}')
try:
import urllib.request
urllib.request.urlretrieve(pretrained_model_url, pretrained_model_path)
logging.info('Pretrained model downloaded.')
except Exception as _exc:
logging.warning(f'Model download failed: {_exc}')
else:
logging.info('Pretrained model already exists.')
if not os.path.exists(pretrained_model_path):
raise FileNotFoundError(
f'SAM1 checkpoint not found at {pretrained_model_path!r} after the '
f'download attempt. Download it manually from '
f'{pretrained_model_url!r} and place it at that path, then re-run '
'this cell.'
)
20:08:2026 16:21:41 - Downloading SAM1 checkpoint to: /Users/amir/GIT_repositories/juspice_pre_release/Pretrained_models/segment_anything1_META/sam_vit_h_4b8939.pth
Using device: mps
20:08:2026 16:25:25 - Pretrained model downloaded.
[6]:
tracker.recording_stop()
How SAM works¶
SAM was trained on about 11 million images with 1 billion object masks (the SA-1B dataset), so it has learned what object boundaries generally look like. It takes a prompt, such as a point or a box, and returns a mask for the object at that location.
Automatic mode (
automatic=True) places a grid of points over the image and creates a mask for each. JuSPICE keeps the largest one.Point mode (
input_point=[[x, y]]) creates one mask for the object at the given pixel.
We use the largest model, vit_h (about 2.6 GB to download). The smaller vit_l and vit_b need less memory.
Because SAM was not trained on microscopy images, its masks should always be checked visually.
First, the automatic mode:
[7]:
tracker.recording_start()
[8]:
# --------- Block 2: Automatic mask generation ----------
# `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(str(input_image_path))
image = spice.data
# run_sam1() derives a NEW SPICEData with the mask — spice.data is never mutated.
# automatic=True: runs SamAutomaticMaskGenerator and returns the largest mask.
auto_spice = spice.segmentation.run_sam1(
model_path=pretrained_model_path,
model_type=model_type,
automatic=True,
device=device,
)
_score = auto_spice.metadata.get("predicted_score")
logging.info(f'Generated automatic mask, score={_score:.3f}')
NOTE: automatic mask generation is running on 'cpu' instead of 'mps' — segment_anything's SamAutomaticMaskGenerator uses float64 tensors internally, which MPS does not support.
20:08:2026 16:25:42 - Generated automatic mask, score=1.018
[9]:
tracker.recording_stop()
The largest automatic mask (right) usually covers the dominant region of the image, which is not necessarily the object we are interested in.
[10]:
tracker.recording_start()
[11]:
# --------- Block 3: Visualize ----------
fig, ax = plt.subplots(1, 2, figsize=(12, 5))
ax[0].imshow(image)
ax[0].set_title('Original Image')
ax[0].axis('off')
ax[1].imshow(auto_spice.data, cmap='gray')
ax[1].set_title('Automatic Mask (largest region)')
ax[1].axis('off')
plt.show()
[12]:
tracker.recording_stop()
Point-prompt segmentation¶
To segment a specific object, we give SAM one pixel inside it (red star). The mask then refers to that object only.
The printed score is SAM’s own estimate of mask quality (its predicted overlap with the true object). It is a confidence value, not a measured accuracy, since we have no ground-truth mask here.
Both calls are recorded in the same spice.history.
[13]:
tracker.recording_start()
[14]:
# --------- Block 4: Point-based prediction ----------
# run_sam1() with input_point runs SamPredictor for prompt-based segmentation.
# Returns a NEW SPICEData — the same spice.history is shared between both calls,
# so all segmentation operations appear in one provenance chain.
input_point = [[400, 300]]
mask_spice = spice.segmentation.run_sam1(
model_path=pretrained_model_path,
model_type=model_type,
input_point=input_point,
device=device,
)
print(f'Point-based mask score: {mask_spice.metadata.get("predicted_score"):.3f}')
fig, ax = plt.subplots(1, 2, figsize=(12, 5))
ax[0].imshow(image)
ax[0].plot(input_point[0][0], input_point[0][1], 'r*', markersize=20, label='Prompt')
ax[0].legend()
ax[0].set_title('Input Image with Point Prompt')
ax[0].axis('off')
ax[1].imshow(mask_spice.data, cmap='gray')
ax[1].set_title('Predicted Mask')
ax[1].axis('off')
plt.show()
Point-based mask score: 0.995
[15]:
tracker.recording_stop()
Save and review¶
We save the point-prompt mask and print the recorded pipeline.
[16]:
tracker.recording_start()
[17]:
# --------- Block 5: Save ----------
# mask_spice IS already the SPICEData returned by spice.segmentation.run_sam1()
# — no manual wrapping needed. It carries the mask 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_sam1'
save_data(
mask_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_sam1.npy
JSON sidecar: /Users/amir/GIT_repositories/juspice_pre_release/notebooks/segmentation/example_sam1.json
History script: /Users/amir/GIT_repositories/juspice_pre_release/notebooks/segmentation/example_sam1_history.py
[18]:
# --------- Block 6: Inspect human-readable history ----------
# mask_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_sam1(...)
# with all runtime argument values (model_path, model_type, mode, etc.),
# and a trailing save_data(mask_spice).
readable_lines = mask_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/em/1b935635dd.png')
mask_spice = spice.segmentation.run_sam1(model_type='vit_h', mode='automatic', model_path='/Users/amir/GIT_repositories/juspice_pre_release/Pretrained_models/segment_anything1_META/sam_vit_h_4b8939.pth')
mask_spice = spice.segmentation.run_sam1(model_type='vit_h', mode='point_prompt', model_path='/Users/amir/GIT_repositories/juspice_pre_release/Pretrained_models/segment_anything1_META/sam_vit_h_4b8939.pth')
juspice.io.save_data(spice)
[19]:
tracker.recording_stop()