synapse_net.inference

This submodule implements SynapseNet's segmentation functionality.

1"""This submodule implements SynapseNet's segmentation functionality.
2"""
3from .inference import compute_scale_from_voxel_size, get_model, get_segmentation_function, run_segmentation
4
5
6__all__ = ["compute_scale_from_voxel_size", "get_model", "get_segmentation_function", "run_segmentation"]
def compute_scale_from_voxel_size(voxel_size: Dict[str, float], model_type: str) -> List[float]:
152def compute_scale_from_voxel_size(
153    voxel_size: Dict[str, float],
154    model_type: str
155) -> List[float]:
156    """Compute the appropriate scale factor for inference with a given pretrained model.
157
158    Args:
159        voxel_size: The voxel size of the data for inference.
160        model_type: The name of the pretrained model.
161
162    Returns:
163        The scale factor, as a list in zyx order.
164    """
165    training_voxel_size = get_model_training_resolution(model_type)
166    scale = [
167        voxel_size["x"] / training_voxel_size["x"],
168        voxel_size["y"] / training_voxel_size["y"],
169    ]
170    if len(voxel_size) == 3 and len(training_voxel_size) == 3:
171        scale.append(
172            voxel_size["z"] / training_voxel_size["z"]
173        )
174    return scale

Compute the appropriate scale factor for inference with a given pretrained model.

Arguments:
  • voxel_size: The voxel size of the data for inference.
  • model_type: The name of the pretrained model.
Returns:

The scale factor, as a list in zyx order.

def get_model( model_type: str, device: Union[str, torch.device, NoneType] = None) -> torch.nn.modules.module.Module:
 97def get_model(model_type: str, device: Optional[Union[str, torch.device]] = None) -> torch.nn.Module:
 98    """Get the model for a specific segmentation type.
 99
100    Args:
101        model_type: The model for one of the following segmentation tasks:
102            'vesicles_3d', 'active_zone', 'compartments', 'mitochondria', 'ribbon', 'vesicles_2d', 'vesicles_cryo'.
103        device: The device to use.
104
105    Returns:
106        The model.
107    """
108    if device is None:
109        device = get_device(device)
110    model_path = get_model_path(model_type)
111    model = torch.load(model_path, weights_only=False)
112    model.to(device)
113    return model

Get the model for a specific segmentation type.

Arguments:
  • model_type: The model for one of the following segmentation tasks: 'vesicles_3d', 'active_zone', 'compartments', 'mitochondria', 'ribbon', 'vesicles_2d', 'vesicles_cryo'.
  • device: The device to use.
Returns:

The model.

def get_segmentation_function(model_type: str) -> Callable:
244def get_segmentation_function(model_type: str) -> Callable:
245    """Get the segmentation function associated with a model type.
246
247    Args:
248        model_type: The name of the pretrained model.
249
250    Returns:
251        The segmentation function used for the model type.
252
253    Raises:
254        ValueError: If the model type is unknown.
255    """
256    if model_type.startswith("vesicles"):
257        return segment_vesicles
258    if model_type in ("mitochondria", "mitochondria2"):
259        return segment_mitochondria
260    if model_type == "active_zone":
261        return segment_active_zone
262    if model_type == "compartments":
263        return segment_compartments
264    if model_type == "ribbon":
265        return _segment_ribbon_AZ
266    if "cristae" in model_type:
267        return segment_cristae
268    raise ValueError(f"Unknown model type: {model_type}")

Get the segmentation function associated with a model type.

Arguments:
  • model_type: The name of the pretrained model.
Returns:

The segmentation function used for the model type.

Raises:
  • ValueError: If the model type is unknown.
def run_segmentation( image: numpy.ndarray, model: torch.nn.modules.module.Module, model_type: str, tiling: Optional[Dict[str, Dict[str, int]]] = None, scale: Optional[List[float]] = None, verbose: bool = False, **kwargs) -> Union[numpy.ndarray, Dict[str, numpy.ndarray]]:
271def run_segmentation(
272    image: np.ndarray,
273    model: torch.nn.Module,
274    model_type: str,
275    tiling: Optional[Dict[str, Dict[str, int]]] = None,
276    scale: Optional[List[float]] = None,
277    verbose: bool = False,
278    **kwargs,
279) -> np.ndarray | Dict[str, np.ndarray]:
280    """Run synaptic structure segmentation.
281
282    Args:
283        image: The input image or image volume.
284        model: The segmentation model.
285        model_type: The model type. This will determine which segmentation post-processing is used.
286        tiling: The tiling settings for inference.
287        scale: A scale factor for resizing the input before applying the model.
288            The output will be scaled back to the initial size.
289        verbose: Whether to print detailed information about the prediction and segmentation.
290        kwargs: Optional parameters for the segmentation function.
291
292    Returns:
293        The segmentation. For models that return multiple segmentations, this function returns a dictionary.
294    """
295    segmentation_function = get_segmentation_function(model_type)
296    if segmentation_function is segment_cristae:
297        training_resolution = get_model_training_resolution(model_type)
298        voxel_size = np.mean(list(training_resolution.values()))
299        segmentation = segmentation_function(
300            image, model=model, tiling=tiling, scale=scale, verbose=verbose, voxel_size=voxel_size, **kwargs
301        )
302    else:
303        segmentation = segmentation_function(
304            image, model=model, tiling=tiling, scale=scale, verbose=verbose, **kwargs
305        )
306    return segmentation

Run synaptic structure segmentation.

Arguments:
  • image: The input image or image volume.
  • model: The segmentation model.
  • model_type: The model type. This will determine which segmentation post-processing is used.
  • tiling: The tiling settings for inference.
  • scale: A scale factor for resizing the input before applying the model. The output will be scaled back to the initial size.
  • verbose: Whether to print detailed information about the prediction and segmentation.
  • kwargs: Optional parameters for the segmentation function.
Returns:

The segmentation. For models that return multiple segmentations, this function returns a dictionary.