synapse_net.inference
This submodule implements SynapseNet's segmentation functionality.
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.