Source code for medical_image.algorithms.pfcm

from typing import Optional, List

import torch

from medical_image.algorithms.algorithm import Algorithm
from medical_image.algorithms.fcm import FCMAlgorithm
from medical_image.data.image import Image
from medical_image.utils.image_utils import MathematicalOperations


[docs] class PFCMAlgorithm(Algorithm): """ Possibilistic Fuzzy C-Means (PFCM) algorithm for microcalcification detection. References:: @article{quintanilla2011image, title={Image segmentation by fuzzy and possibilistic clustering algorithms for the identification of microcalcifications}, author={Quintanilla-Dominguez, Joel and others}, journal={Scientia Iranica}, volume={18}, number={3}, pages={580--589}, year={2011}, publisher={Elsevier} } Math and Logic: PFCM extends FCM by adding typicality values that measure how "typical" a sample is for each cluster. Microcalcifications are detected as **atypical** pixels — those with a low maximum typicality. Pipeline: 1. Run standard FCM to warm-start cluster centroids and memberships. 2. Compute initial gamma values and typicality matrix T. 3. Iteratively update prototypes, memberships, gammas, and typicalities. 4. Detect MCs by thresholding the maximum typicality map (atypical pixels). 5. Exclude the darkest background cluster. Attributes (populated after ``apply()``): typicality: ``(c, N)`` typicality matrix T. T_max_map: ``(H, W)`` max typicality per pixel. centroids: ``(c, d)`` cluster centroids. membership: ``(c, N)`` fuzzy membership matrix. labels: ``(H, W)`` int hard cluster assignments. quantized: ``(H, W)`` float quantized image. gamma: ``(c,)`` gamma values per cluster. """
[docs] def __init__( self, c: int = 2, m: float = 2.0, eta: float = 2.0, a: float = 1.0, b: float = 4.0, tau: float = 0.04, max_iter: int = 100, tol: float = 1e-3, fcm_max_iter: int = 100, random_state: int = 42, device: str = "cpu", ): super().__init__(device=device) self.c = c self.m = m self.eta = eta self.a = a self.b = b self.tau = tau self.max_iter = max_iter self.tol = tol self.fcm_max_iter = fcm_max_iter self.random_state = random_state self.compute_distances = ( lambda Z, V: MathematicalOperations.euclidean_distance_sq(Z=Z, V=V) ) self.update_membership = lambda D2: FCMAlgorithm._update_membership( D2=D2, m=self.m ) self.compute_gamma = lambda U, D2: PFCMAlgorithm._compute_gamma( U=U, D2=D2, m=self.m ) self.update_typicality = lambda D2, gamma: PFCMAlgorithm._update_typicality( D2=D2, gamma=gamma, b=self.b, eta=self.eta ) self.update_prototypes = lambda Z, U, T: PFCMAlgorithm._update_prototypes( Z=Z, U=U, T=T, m=self.m, eta=self.eta, a=self.a, b=self.b ) # Results (populated by apply) self.typicality: Optional[torch.Tensor] = None self.T_max_map: Optional[torch.Tensor] = None self.centroids: Optional[torch.Tensor] = None self.membership: Optional[torch.Tensor] = None self.labels: Optional[torch.Tensor] = None self.quantized: Optional[torch.Tensor] = None self.gamma: Optional[torch.Tensor] = None self.n_iter: int = 0 self.converged: bool = False
@staticmethod def _compute_gamma(U: torch.Tensor, D2: torch.Tensor, m: float) -> torch.Tensor: Um = U**m return (Um * D2).sum(dim=1) / (Um.sum(dim=1) + 1e-10) @staticmethod def _update_typicality( D2: torch.Tensor, gamma: torch.Tensor, b: float, eta: float, ) -> torch.Tensor: eps = 1e-10 exp = 1.0 / (eta - 1.0) ratio = (b / (gamma.unsqueeze(1) + eps)) * D2 return 1.0 / (1.0 + (ratio + eps) ** exp) @staticmethod def _update_prototypes( Z: torch.Tensor, U: torch.Tensor, T: torch.Tensor, m: float, eta: float, a: float, b: float, ) -> torch.Tensor: W = a * (U**m) + b * (T**eta) denom = W.sum(dim=1, keepdim=True) + 1e-10 return (W @ Z) / denom
[docs] def apply(self, image: Image, output: Image) -> Image: """ Apply PFCM: warm-start from FCM, iterate PFCM, detect MCs by atypicality. Args: image: Input Image (2D float tensor). output: Output Image — pixel_data = binary MC mask. Returns: The output Image. """ device = self.device img = image.pixel_data.to(device).float() while img.ndim > 2: img = img.squeeze(0) H, W = img.shape image_shape = (H, W) N = H * W Z = img.reshape(N, 1) # Step 1: Warm-start from FCM fcm = FCMAlgorithm( c=self.c, m=self.m, max_iter=self.fcm_max_iter, tol=self.tol, random_state=self.random_state, device=device, ) fcm_output = image.clone() fcm.apply(image, fcm_output) U = fcm.membership.clone() V = fcm.centroids.clone() # Step 2: Initial gamma + typicality D2 = self.compute_distances(Z, V) gamma = self.compute_gamma(U, D2) T = self.update_typicality(D2, gamma) # Step 3: PFCM iterations converged = False n_iter = 0 for iteration in range(self.max_iter): V = self.update_prototypes(Z, U, T) D2 = self.compute_distances(Z, V) U = self.update_membership(D2) gamma = self.compute_gamma(U, D2) T_new = self.update_typicality(D2, gamma) diff = float(torch.norm(T_new - T)) T = T_new n_iter = iteration + 1 if diff < self.tol: converged = True break # Step 4: MC detection via atypicality T_max = T.max(dim=0).values T_max_map = T_max.reshape(image_shape) atypical_mask = T_max_map < self.tau labels = torch.argmax(U, dim=0).to(torch.int64) labels_2d = labels.reshape(image_shape) centroid_vals = V[:, 0] darkest_cluster = int(torch.argmin(centroid_vals)) darkest_mask = labels_2d == darkest_cluster binary_mask = atypical_mask & (~darkest_mask) mx = centroid_vals.max() quant_lut = centroid_vals / (mx + 1e-10) quantized = quant_lut[labels_2d] # Store results self.typicality = T self.T_max_map = T_max_map self.centroids = V self.membership = U self.labels = labels_2d self.quantized = quantized self.gamma = gamma self.n_iter = n_iter self.converged = converged output.pixel_data = binary_mask.float() return output