Skip to content
Closed
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
59 changes: 44 additions & 15 deletions monai/losses/image_dissimilarity.py
Original file line number Diff line number Diff line change
Expand Up @@ -316,40 +316,69 @@ def parzen_windowing_gaussian(self, img: torch.Tensor) -> tuple[torch.Tensor, to
Note: the input is expected to range between 0 and 1
Args:
img: the shape should be B[NDHW].
Returns:
A tuple containing per-sample Gaussian bin weights and the
corresponding marginal probability.
Raises:
ValueError: if the Gaussian kernel buffers are unavailable.
"""
img = torch.clamp(img, 0, 1)
output_dtype = img.dtype
if output_dtype == torch.float16:
img = img.float()
compute_dtype = torch.float32 if img.dtype in (torch.float16, torch.bfloat16) else img.dtype
img = torch.clamp(img, 0, 1).to(dtype=compute_dtype)
img = img.reshape(img.shape[0], -1, 1) # (batch, num_sample, 1)
if self.bin_centers is None or self.preterm is None:
raise ValueError("bin_centers and preterm must be defined for gaussian parzen windowing.")
weight = torch.exp(
-self.preterm.to(img) * (img - self.bin_centers.to(img)) ** 2
) # (batch, num_sample, num_bin)
preterm = self.preterm.to(device=img.device, dtype=compute_dtype)
bin_centers = self.bin_centers.to(device=img.device, dtype=compute_dtype)
weight = torch.exp(-preterm * (img - bin_centers) ** 2) # (batch, num_sample, num_bin)
weight = weight / torch.sum(weight, dim=-1, keepdim=True) # (batch, num_sample, num_bin)
probability = torch.mean(weight, dim=-2, keepdim=True) # (batch, 1, num_bin)
if output_dtype in (torch.float16, torch.bfloat16):
weight = weight.to(dtype=output_dtype)
probability = probability.to(dtype=output_dtype)
return weight, probability

def forward(self, pred: torch.Tensor, target: torch.Tensor) -> torch.Tensor:
"""
Args:
pred: the shape should be B[NDHW].
target: the shape should be same as the pred shape.
Returns:
Reduced negative mutual information loss.
Raises:
ValueError: When ``self.reduction`` is not one of ["mean", "sum", "none"].
"""
if target.shape != pred.shape:
raise ValueError(f"ground truth has differing shape ({target.shape}) from pred ({pred.shape})")
wa, pa, wb, pb = self.parzen_windowing(pred, target) # (batch, num_sample, num_bin), (batch, 1, num_bin)

pab = torch.bmm(wa.permute(0, 2, 1), wb.to(wa)).div(wa.shape[1]) # (batch, num_bins, num_bins)
papb = torch.bmm(pa.permute(0, 2, 1), pb.to(pa)) # (batch, num_bins, num_bins)
mi = torch.sum(
pab * torch.log((pab + self.smooth_nr) / (papb + self.smooth_dr) + self.smooth_dr), dim=(1, 2)
) # (batch)
# Half-precision matrix multiplication can overflow before the joint
# histogram is divided by the number of samples. Accumulate histogram
# products in float32 while preserving float64 inputs.
output_dtype = wa.dtype
compute_dtype = torch.float32 if wa.dtype in (torch.float16, torch.bfloat16) else wa.dtype
wa = wa.to(dtype=compute_dtype)
wb = wb.to(wa)
pa = pa.to(dtype=compute_dtype)
pb = pb.to(pa)
with torch.autocast(device_type=wa.device.type, enabled=False):
pab = torch.bmm(wa.permute(0, 2, 1), wb).div(wa.shape[1]) # (batch, num_bins, num_bins)
papb = torch.bmm(pa.permute(0, 2, 1), pb) # (batch, num_bins, num_bins)
mi = torch.sum(
pab * torch.log((pab + self.smooth_nr) / (papb + self.smooth_dr) + self.smooth_dr), dim=(1, 2)
) # (batch)

if self.reduction == LossReduction.SUM.value:
return torch.sum(mi).neg() # sum over the batch and channel ndims
if self.reduction == LossReduction.NONE.value:
return mi.neg()
if self.reduction == LossReduction.MEAN.value:
return torch.mean(mi).neg() # average over the batch and channel ndims
raise ValueError(f'Unsupported reduction: {self.reduction}, available options are ["mean", "sum", "none"].')
loss = torch.sum(mi).neg() # sum over the batch and channel ndims
elif self.reduction == LossReduction.NONE.value:
loss = mi.neg()
elif self.reduction == LossReduction.MEAN.value:
loss = torch.mean(mi).neg() # average over the batch and channel ndims
else:
raise ValueError(f'Unsupported reduction: {self.reduction}, available options are ["mean", "sum", "none"].')
return loss.to(dtype=output_dtype)
Original file line number Diff line number Diff line change
Expand Up @@ -164,6 +164,59 @@ def test_ill_opts(self, num_bins, reduction, expected_exception, expected_messag
GlobalMutualInformationLoss(num_bins=num_bins, reduction=reduction)(pred, target)


class TestGlobalMutualInformationLossHalfPrecision(unittest.TestCase):
"""Test stable Gaussian mutual information in reduced-precision modes."""

@parameterized.expand([(torch.float16,), (torch.bfloat16,)])
def test_half_precision_gaussian_weights_with_many_bins_are_finite(self, dtype):
"""Verify many-bin Parzen outputs remain finite and preserve metadata."""
image = torch.zeros((1, 1, 2), dtype=dtype)
loss = GlobalMutualInformationLoss(kernel_type="gaussian", num_bins=256)

weight, probability = loss.parzen_windowing_gaussian(image)

self.assertTrue(torch.isfinite(weight).all())
self.assertTrue(torch.isfinite(probability).all())
self.assertEqual(weight.dtype, image.dtype)
self.assertEqual(probability.dtype, image.dtype)
self.assertEqual(weight.device, image.device)
self.assertEqual(probability.device, image.device)

@parameterized.expand([(torch.float16,), (torch.bfloat16,)])
def test_half_precision_large_constant_volume_is_finite(self, dtype):
"""Verify reduced-precision loss and gradients remain finite."""
pred = torch.zeros((1, 1, 48, 48, 48), dtype=dtype, requires_grad=True)
target = torch.zeros_like(pred)
loss = GlobalMutualInformationLoss(kernel_type="gaussian")

result = loss(pred, target)

self.assertTrue(torch.isfinite(result))
self.assertEqual(result.dtype, pred.dtype)
self.assertEqual(result.device, pred.device)
result.backward()
self.assertIsNotNone(pred.grad)
self.assertTrue(torch.isfinite(pred.grad).all())
self.assertEqual(pred.grad.dtype, pred.dtype)
self.assertEqual(pred.grad.device, pred.device)

def test_cpu_float16_autocast_large_volume_is_finite(self):
"""Verify CPU float16 autocast avoids histogram accumulation overflow."""
pred = torch.zeros((1, 1, 48, 48, 48), requires_grad=True)
target = torch.zeros_like(pred)
loss = GlobalMutualInformationLoss(kernel_type="gaussian")

with torch.autocast(device_type="cpu", dtype=torch.float16):
result = loss(pred, target)

self.assertTrue(torch.isfinite(result))
result.backward()
self.assertIsNotNone(pred.grad)
self.assertTrue(torch.isfinite(pred.grad).all())
self.assertEqual(pred.grad.dtype, pred.dtype)
self.assertEqual(pred.grad.device, pred.device)


class TestGlobalMutualInformationLossBuffers(unittest.TestCase):
def test_gaussian_kernel_registers_buffers(self):
"""Verify gaussian kernel registers preterm and bin_centers as non-trainable, non-persistent buffers."""
Expand Down
Loading