mirror of
https://github.com/cdtools-developers/cdtools.git
synced 2026-09-09 13:02:41 +02:00
Attempted to update the normalizations for the loss functions in a way which allows for a nice normalization of the Poisson NLL. The slow tests will now likely fail
This commit is contained in:
@@ -220,9 +220,11 @@ class FancyPtycho(CDIModel):
|
||||
if (loss.lower().strip() == 'amplitude mse'
|
||||
or loss.lower().strip() == 'amplitude_mse'):
|
||||
self.loss = tools.losses.amplitude_mse
|
||||
self.loss_normalizer = tools.losses.AmplitudeMSENormalizer()
|
||||
elif (loss.lower().strip() == 'poisson nll'
|
||||
or loss.lower().strip() == 'poisson_nll'):
|
||||
self.loss = tools.losses.poisson_nll
|
||||
self.loss_normalizer = tools.losses.SimplePoissonNLLNormalizer()
|
||||
else:
|
||||
raise KeyError('Specified loss function not supported')
|
||||
|
||||
|
||||
@@ -162,7 +162,9 @@ class Reconstructor:
|
||||
# The data loader is responsible for setting the minibatch
|
||||
# size, so each set is a minibatch
|
||||
for inputs, patterns in self.data_loader:
|
||||
normalization += t.sum(patterns).cpu().numpy()
|
||||
if hasattr(self.model, 'loss_normalizer') and \
|
||||
self.model.loss_normalizer is not None:
|
||||
self.model.loss_normalizer.accumulate(patterns)
|
||||
N += 1
|
||||
|
||||
def closure():
|
||||
@@ -218,7 +220,9 @@ class Reconstructor:
|
||||
# This takes the step for this minibatch
|
||||
loss += self.optimizer.step(closure).detach().cpu().numpy()
|
||||
|
||||
loss /= normalization
|
||||
if hasattr(self.model, 'loss_normalizer') and \
|
||||
self.model.loss_normalizer is not None:
|
||||
loss = self.model.loss_normalizer.normalize_loss(loss)
|
||||
|
||||
# We step the scheduler after the full epoch
|
||||
if self.scheduler is not None:
|
||||
|
||||
@@ -8,7 +8,14 @@ maximum likelihood metric for a system with Poisson statistics.
|
||||
|
||||
import torch as t
|
||||
|
||||
__all__ = ['amplitude_mse', 'intensity_mse', 'poisson_nll']
|
||||
__all__ = [
|
||||
'amplitude_mse',
|
||||
'AmplitudeMSENormalizer',
|
||||
'intensity_mse',
|
||||
'IntensityMSENormalizer',
|
||||
'poisson_nll',
|
||||
'SimplePoissonNLLNormalizer',
|
||||
]
|
||||
|
||||
|
||||
def amplitude_mse(intensities, sim_intensities, mask=None):
|
||||
@@ -17,14 +24,20 @@ def amplitude_mse(intensities, sim_intensities, mask=None):
|
||||
Calculates the mean squared error between a given set of
|
||||
measured diffraction intensities and a simulated set.
|
||||
|
||||
|
||||
This function calculates the mean squared error between their
|
||||
associated amplitudes. Because this is not well defined for negative
|
||||
numbers, make sure that all the intensities are >0 before using this
|
||||
loss. Note that this is actually a sum-squared error, because this
|
||||
formulation makes it vastly simpler to compare error calculations
|
||||
between reconstructions with different minibatch size. I hope to
|
||||
find a better way to do this that is more honest with this
|
||||
cost function, though.
|
||||
loss.
|
||||
|
||||
Note that this is actually, by defauly, a sum-squared error. In this
|
||||
case, it is intended to be used with the loss normalization
|
||||
|
||||
|
||||
|
||||
in a ptychography model. This formulation makes it easier to compare
|
||||
error calculations between reconstructions with different minibatch
|
||||
size while keeping the loss function formally equivalent to the MSE.
|
||||
|
||||
It can accept intensity and simulated intensity tensors of any shape
|
||||
as long as their shapes match, and the provided mask array can be
|
||||
@@ -40,6 +53,8 @@ def amplitude_mse(intensities, sim_intensities, mask=None):
|
||||
A tensor of simulated detector intensities
|
||||
mask : torch.Tensor
|
||||
A mask with ones for pixels to include and zeros for pixels to exclude
|
||||
use_sum : bool
|
||||
Default is True. If set to True, actually performs the sum squared error
|
||||
|
||||
Returns
|
||||
-------
|
||||
@@ -60,6 +75,48 @@ def amplitude_mse(intensities, sim_intensities, mask=None):
|
||||
t.sqrt(masked_intensities))**2)
|
||||
|
||||
|
||||
class AmplitudeMSENormalizer(object):
|
||||
""" Normalizer for the amplitude MSE loss, used with recon.optimize
|
||||
|
||||
This is a normalizer designed for use with the recon.optimize function. The
|
||||
normalization is done separately from the loss, in order to make it simple to
|
||||
use different normalization strategies for different loss metrics and to make it
|
||||
easier to work with different minibatch sizes.
|
||||
|
||||
This normalizer accumulates the total number of pixels across all patterns
|
||||
during the first epoch, then divides the summed loss by this count to
|
||||
convert from sum-squared error to mean-squared error.
|
||||
|
||||
The normalizer is stateful: it completes its accumulation phase on the
|
||||
first epoch and then applies the same normalization factor for all
|
||||
subsequent epochs.
|
||||
|
||||
Methods
|
||||
-------
|
||||
accumulate(patterns, mask=None)
|
||||
Accumulate the normalization factor (called once per minibatch).
|
||||
normalize_loss(loss)
|
||||
Apply the accumulated normalization (called once per epoch).
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.first_pass_complete = False
|
||||
self.num_pix = 0
|
||||
|
||||
def accumulate(self, patterns, mask=None):
|
||||
if not self.first_pass_complete:
|
||||
if mask is None:
|
||||
self.num_pix += patterns.numel()
|
||||
else:
|
||||
self.num_pix += patterns.masked_select(mask).numel()
|
||||
|
||||
def normalize_loss(self, loss):
|
||||
if not self.first_pass_complete:
|
||||
self.first_pass_complete = True
|
||||
|
||||
return loss / self.num_pix
|
||||
|
||||
|
||||
def intensity_mse(intensities, sim_intensities, mask=None):
|
||||
""" Returns the mean squared error of a simulated dataset's intensities
|
||||
@@ -98,6 +155,72 @@ def intensity_mse(intensities, sim_intensities, mask=None):
|
||||
/ masked_intensities.shape[0]
|
||||
|
||||
|
||||
class IntensityMSENormalizer(object):
|
||||
""" Normalizer for the intensity MSE loss, used with recon.optimize
|
||||
|
||||
This is a normalizer designed for use with the recon.optimize function. The
|
||||
normalization is done separately from the loss, in order to make it simple to
|
||||
use different normalization strategies for different loss metrics and to make it
|
||||
easier to work with different minibatch sizes.
|
||||
|
||||
This normalizer accumulates the total number of pixels across all patterns
|
||||
during the first epoch, then divides the summed loss by this count to
|
||||
convert from sum-squared error to mean-squared error.
|
||||
|
||||
The normalizer is stateful: it completes its accumulation phase on the
|
||||
first epoch and then applies the same normalization factor for all
|
||||
subsequent epochs.
|
||||
|
||||
Methods
|
||||
-------
|
||||
accumulate(patterns, mask=None)
|
||||
Accumulate the normalization factor (called once per minibatch).
|
||||
normalize_loss(loss)
|
||||
Apply the accumulated normalization (called once per epoch).
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.first_pass_complete = False
|
||||
self.num_pix = 0
|
||||
|
||||
def accumulate(self, patterns, mask=None):
|
||||
"""Accumulate pixel counts from a batch of patterns.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
patterns : torch.Tensor
|
||||
A tensor of measured detector patterns
|
||||
mask : torch.Tensor, optional
|
||||
A mask with ones for pixels to include and zeros for pixels to
|
||||
exclude. If provided, only masked pixels are counted.
|
||||
|
||||
"""
|
||||
if not self.first_pass_complete:
|
||||
if mask is None:
|
||||
self.num_pix += patterns.numel()
|
||||
else:
|
||||
self.num_pix += patterns.masked_select(mask).numel()
|
||||
|
||||
def normalize_loss(self, loss):
|
||||
"""Convert summed loss to mean loss by dividing by pixel count.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
loss : torch.Tensor
|
||||
The accumulated summed loss across minibatches in an epoch
|
||||
|
||||
Returns
|
||||
-------
|
||||
normalized_loss : torch.Tensor
|
||||
The loss divided by the total number of pixels
|
||||
|
||||
"""
|
||||
if not self.first_pass_complete:
|
||||
self.first_pass_complete = True
|
||||
|
||||
return loss / self.num_pix
|
||||
|
||||
|
||||
def poisson_nll(
|
||||
intensities,
|
||||
@@ -135,6 +258,8 @@ def poisson_nll(
|
||||
A mask with ones for pixels to include and zeros for pixels to exclude
|
||||
eps : float
|
||||
Optional, a small number to add to the simulated intensities
|
||||
subtract_min : bool
|
||||
Default is False, whether to subtract a min to produce a nonnegative output
|
||||
|
||||
Returns
|
||||
-------
|
||||
@@ -144,8 +269,7 @@ def poisson_nll(
|
||||
"""
|
||||
if mask is None:
|
||||
nll = t.sum(sim_intensities+eps -
|
||||
t.xlogy(intensities,sim_intensities+eps)) \
|
||||
/ intensities.view(-1).shape[0]
|
||||
t.xlogy(intensities,sim_intensities+eps))
|
||||
if subtract_min:
|
||||
nll -= t.sum(intensities - t.xlogy(intensities,intensities))
|
||||
|
||||
@@ -155,17 +279,108 @@ def poisson_nll(
|
||||
masked_sims = sim_intensities.masked_select(mask)
|
||||
|
||||
nll = t.sum(masked_sims + eps - \
|
||||
t.xlogy(masked_intensities, masked_sims+eps)) \
|
||||
/ masked_intensities.shape[0]
|
||||
t.xlogy(masked_intensities, masked_sims+eps))
|
||||
|
||||
if subtract_min:
|
||||
nll -= t.nansum(masked_intensities - \
|
||||
t.xlogy(masked_intensities, masked_intensities)) \
|
||||
/ masked_intensities.shape[0]
|
||||
t.xlogy(masked_intensities, masked_intensities))
|
||||
|
||||
return nll
|
||||
|
||||
|
||||
class SimplePoissonNLLNormalizer(object):
|
||||
""" Normalizer for the intensity MSE loss, used with recon.optimize
|
||||
|
||||
This is a normalizer designed for use with the recon.optimize function. The
|
||||
normalization is done separately from the loss, in order to make it simple to
|
||||
use different normalization strategies for different loss metrics and to make it
|
||||
easier to work with different minibatch sizes.
|
||||
|
||||
This normalizer converts raw Poisson negative log likelihood values into
|
||||
a statistic that is more interpretable for comparing reconstructions. It
|
||||
performs two operations:
|
||||
|
||||
1. **Offset subtraction**: Subtracts the NLL calculated when comparing
|
||||
measured patterns to themselves (i.e., poisson_nll(data, data)). This
|
||||
represents the best-case scenario and makes the loss non-negative.
|
||||
|
||||
2. **Normalization scaling**: Divides by 0.5 times the count of non-zero
|
||||
pixels in the measured patterns. This is because, roughly, each non-zero
|
||||
pixel is expected to contribute to the Poisson NLL, if Poisson noise were
|
||||
the only relevant source of noise in the data.
|
||||
|
||||
The normalizer is stateful: it completes its accumulation phase on the
|
||||
first epoch by processing all patterns in the data, then applies the
|
||||
same normalization factors for all subsequent epochs.
|
||||
|
||||
Methods
|
||||
-------
|
||||
accumulate(patterns, mask=None)
|
||||
Accumulate the normalization factor (called once per minibatch).
|
||||
normalize_loss(loss)
|
||||
Apply the accumulated normalization (called once per epoch).
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.first_pass_complete = False
|
||||
self.sum_nonzero = 0
|
||||
self.offset = 0
|
||||
|
||||
def accumulate(self, patterns, mask=None):
|
||||
"""Accumulate statistics needed for normalization from a batch.
|
||||
|
||||
During the first epoch, this method counts non-zero pixels and
|
||||
computes the Poisson NLL comparing patterns to themselves, which
|
||||
defines the offset baseline for the loss.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
patterns : torch.Tensor
|
||||
A tensor of measured detector patterns
|
||||
mask : torch.Tensor, optional
|
||||
A mask with ones for pixels to include and zeros for pixels to
|
||||
exclude. If provided, only masked pixels are counted.
|
||||
|
||||
"""
|
||||
if not self.first_pass_complete:
|
||||
if mask is None:
|
||||
self.sum_nonzero += t.sum(patterns >= 1)
|
||||
self.offset += poisson_nll(patterns, patterns)
|
||||
else:
|
||||
masked_pats = patterns.masked_select(mask)
|
||||
self.sum_nonzero += t.sum(masked_pats >= 1)
|
||||
self.offset += poisson_nll(masked_pats, masked_pats)
|
||||
|
||||
|
||||
def normalize_loss(self, loss):
|
||||
"""Normalize the Poisson NLL for interpretability across datasets.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
loss : torch.Tensor
|
||||
The accumulated Poisson NLL across minibatches in an epoch
|
||||
|
||||
Returns
|
||||
-------
|
||||
normalized_loss : torch.Tensor
|
||||
The offset-corrected and scaled loss value
|
||||
|
||||
"""
|
||||
if not self.first_pass_complete:
|
||||
self.normalization = 0.5 * self.sum_nonzero
|
||||
self.first_pass_complete = True
|
||||
|
||||
return (loss - self.offset) / self.normalization
|
||||
|
||||
|
||||
#
|
||||
# Note: I have two other ideas for how to normalize the Poisson NLL
|
||||
#
|
||||
# Idea 2: Use the mean pattern to estimate the expected error
|
||||
# Idea 3: Use the simulated intensities to estimate it, but use detach
|
||||
# so it doesn't hit the backward pass
|
||||
#
|
||||
|
||||
def poisson_plus_fixed_nll(
|
||||
intensities,
|
||||
|
||||
Reference in New Issue
Block a user