Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -116,10 +116,10 @@ You can edit the config file afterward to add additional metadata like `emdb_id`
Segmentations are stored in the copick project at:

```
{overlay_root}/ExperimentRuns/{run_name}/VoxelSpacing{voxel_size:.3f}/Segmentations/
{overlay_root}/ExperimentRuns/{run_name}/Segmentations/
```

Each segmentation is stored as a zarr array with OME-Zarr metadata.
Each single-label segmentation is stored as a `uint8` OME-Zarr 0.5 / Zarr v3 array.

## Requirements

Expand Down
2 changes: 2 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -50,8 +50,10 @@ Issues = "https://github.com/copick/copick-easymode/issues"

[project.optional-dependencies]
test = [
"ome-zarr-models>=1.6",
"pytest>=8.0.0",
"pytest-cov>=6.2.1",
"zarr>=3.1.6,<4",
]
dev = [
"black>=25.1.0",
Expand Down
242 changes: 135 additions & 107 deletions src/copick_easymode/core/inference.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
from typing import TYPE_CHECKING, Optional

import numpy as np
from copick.util.escape import sanitize_name

from copick_easymode.core import easymode_adapter

Expand Down Expand Up @@ -198,139 +199,166 @@ def run_easymode_inference(
logger.warning("No runs found to process.")
return stats

sanitized_user_id = sanitize_name(user_id, suppress_warnings=True)
sanitized_session_id = sanitize_name(session_id, suppress_warnings=True)
if logger and sanitized_user_id != user_id:
logger.info(f"Using copick user ID '{sanitized_user_id}' for requested user ID '{user_id}'")
if logger and sanitized_session_id != session_id:
logger.info(f"Using copick session ID '{sanitized_session_id}' for requested session ID '{session_id}'")

# Track if config was modified
config_modified = False

# Process each model
for model_name in models:
if logger:
logger.info(f"Loading model: {model_name}")

# Get and load model
model_path, metadata = runtime.get_model(model_name)
if model_path is None:
error_msg = f"Model '{model_name}' not found. Skipping."
if logger:
logger.error(error_msg)
stats["errors"].append(error_msg)
continue

model_apix = (metadata or {}).get("apix", 10.0)

if logger:
logger.info(f"Model loaded from {model_path}, inference at {model_apix} A/px")

model = runtime.load_model(model_path)

# Add object definition if needed
if add_objects:
existing_obj = root.get_object(model_name)
if existing_obj is None:
model = None
try:
try:
entity_name = sanitize_name(model_name, suppress_warnings=True)
except Exception as error:
error_msg = f"Error preparing copick entity for model '{model_name}': {error}. Skipping."
if logger:
logger.info(f"Adding object definition for '{model_name}'")
root.new_object(
name=model_name,
is_particle=False, # Segmentation, not particle picks
# label and color will be auto-assigned
)
config_modified = True
logger.exception(error_msg)
stats["errors"].append(error_msg)
continue

# Process each run
for run in runs:
if logger:
logger.info(f"Processing run: {run.name}")
logger.info(f"Loading model: {model_name}")
if entity_name != model_name:
logger.info(f"Using copick entity '{entity_name}' for easymode model '{model_name}'")

# Get tomogram
# Get and load model. Model failures do not prevent later models from running.
try:
vs = run.get_voxel_spacing(voxel_size)
if vs is None:
if logger:
logger.warning(f"Voxel spacing {voxel_size} not found in {run.name}")
stats["skipped"] += 1
continue

tomo = vs.get_tomogram(tomo_type)
if tomo is None:
if logger:
logger.warning(f"Tomogram {tomo_type}@{voxel_size} not found in {run.name}")
stats["skipped"] += 1
continue
except Exception as e:
error_msg = f"Error getting tomogram in {run.name}: {e}"
model_path, metadata = runtime.get_model(model_name)
if model_path is None:
raise FileNotFoundError(f"Model '{model_name}' not found")
model = runtime.load_model(model_path)
except Exception as error:
error_msg = f"Error loading model '{model_name}': {error}. Skipping."
if logger:
logger.warning(error_msg)
logger.exception(error_msg)
stats["errors"].append(error_msg)
continue

# Check if segmentation already exists
existing_segs = run.get_segmentations(
name=model_name,
user_id=user_id,
session_id=session_id,
voxel_size=voxel_size,
is_multilabel=False,
)
model_apix = (metadata or {}).get("apix", 10.0)

if existing_segs and not overwrite:
if logger:
logger.info(f"Segmentation already exists for {model_name} in {run.name}, skipping")
stats["skipped"] += 1
continue

try:
# Read tomogram as numpy
if logger:
logger.info(f"Reading tomogram from {run.name}")
tomo_data = tomo.numpy()
if logger:
logger.info(f"Model loaded from {model_path}, inference at {model_apix} A/px")

# Add object definition if needed
if add_objects:
try:
existing_obj = root.get_object(entity_name)
if existing_obj is None:
if logger:
logger.info(f"Adding object definition for '{entity_name}'")
root.new_object(
name=entity_name,
is_particle=False, # Segmentation, not particle picks
# label and color will be auto-assigned
)
config_modified = True
except Exception as error:
error_msg = f"Error creating object for model '{model_name}': {error}. Skipping."
if logger:
logger.exception(error_msg)
stats["errors"].append(error_msg)
continue

# Run inference
# Process each run
for run in runs:
if logger:
logger.info(f"Running inference for {model_name} on {run.name}")

seg_data = segmenter(
model=model,
volume=tomo_data,
input_apix=voxel_size,
model_apix=model_apix,
tta=tta,
batch_size=batch_size,
)

# Binarize using threshold and convert to uint8 (0 or 1)
seg_data = (seg_data >= threshold).astype(np.uint8)

# Create or get segmentation
if existing_segs and overwrite:
# Delete existing segmentation first
for _existing_seg in existing_segs:
# Note: copick doesn't have a delete method, so we overwrite via exist_ok
pass
logger.info(f"Processing run: {run.name}")

# Get tomogram
try:
vs = run.get_voxel_spacing(voxel_size)
if vs is None:
if logger:
logger.warning(f"Voxel spacing {voxel_size} not found in {run.name}")
stats["skipped"] += 1
continue

tomograms = vs.get_tomograms(tomo_type=tomo_type)
if not tomograms:
if logger:
logger.warning(f"Tomogram {tomo_type}@{voxel_size} not found in {run.name}")
stats["skipped"] += 1
continue
tomo = tomograms[0]
except Exception as error:
error_msg = f"Error getting tomogram in {run.name}: {error}"
if logger:
logger.warning(error_msg)
stats["errors"].append(error_msg)
continue

seg = run.new_segmentation(
name=model_name,
# Check if segmentation already exists
existing_segs = run.get_segmentations(
name=entity_name,
user_id=sanitized_user_id,
session_id=sanitized_session_id,
voxel_size=voxel_size,
user_id=user_id,
session_id=session_id,
is_multilabel=False,
)

# Write segmentation
seg.from_numpy(seg_data)
if existing_segs and not overwrite:
if logger:
logger.info(f"Segmentation already exists for {model_name} in {run.name}, skipping")
stats["skipped"] += 1
continue

if logger:
logger.info(f"Saved segmentation for {model_name} in {run.name}")
try:
# Read tomogram as numpy
if logger:
logger.info(f"Reading tomogram from {run.name}")
tomo_data = tomo.numpy()

stats["processed"] += 1
# Run inference
if logger:
logger.info(f"Running inference for {model_name} on {run.name}")

seg_data = segmenter(
model=model,
volume=tomo_data,
input_apix=voxel_size,
model_apix=model_apix,
tta=tta,
batch_size=batch_size,
)

# Binarize using threshold and convert to uint8 (0 or 1)
seg_data = (seg_data >= threshold).astype(np.uint8)

# Reuse a matching entity only after inference has completed. This lets a
# failed inference leave an existing segmentation untouched.
seg = run.new_segmentation(
name=entity_name,
voxel_size=voxel_size,
user_id=sanitized_user_id,
session_id=sanitized_session_id,
is_multilabel=False,
exist_ok=overwrite,
)

# Write segmentation
seg.from_numpy(seg_data)

except Exception as e:
error_msg = f"Error processing {model_name} in {run.name}: {e}"
if logger:
logger.exception(error_msg)
stats["errors"].append(error_msg)
if logger:
logger.info(f"Saved segmentation for {model_name} in {run.name}")

# Clean up model to free GPU memory
tf.keras.backend.clear_session()
gc.collect()
stats["processed"] += 1

except Exception as error:
error_msg = f"Error processing {model_name} in {run.name}: {error}"
if logger:
logger.exception(error_msg)
stats["errors"].append(error_msg)
finally:
# Release the loaded model after every model attempt, including failures.
model = None
tf.keras.backend.clear_session()
gc.collect()

# Save config if modified
if config_modified and config_path:
Expand Down
Loading
Loading