From 88aedcc44843db594fac10e5bfe54b7ede975183 Mon Sep 17 00:00:00 2001 From: Abe Levitan Date: Tue, 15 Aug 2023 10:47:50 -0700 Subject: [PATCH] Update the poisson loss function to use xlogy and start work on the poisson + detector loss function --- pyproject.toml | 3 + .../tools/interactions/interactions.py | 2 +- src/cdtools/tools/losses/losses.py | 121 ++++++++++++++---- 3 files changed, 103 insertions(+), 23 deletions(-) create mode 100644 pyproject.toml diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..b7fd46a --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,3 @@ +[tool.ruff] +# Decrease the maximum line length to 79 characters. +line-length = 79 \ No newline at end of file diff --git a/src/cdtools/tools/interactions/interactions.py b/src/cdtools/tools/interactions/interactions.py index f290e9c..247b515 100644 --- a/src/cdtools/tools/interactions/interactions.py +++ b/src/cdtools/tools/interactions/interactions.py @@ -41,7 +41,7 @@ def translations_to_pixel(basis, translations, surface_normal=t.Tensor([0.,0.,1. Returns ------- pixel_translations : torch.Tensor - A Jx2 stack of translations in internal (i,j) pixel-space, or a single translation + A Jx2 stack of translations, or a single translation, in (i,j) pixel-space """ projection_1 = t.as_tensor(np.array([[1,0,0], [0,1,0], diff --git a/src/cdtools/tools/losses/losses.py b/src/cdtools/tools/losses/losses.py index aeb5345..7f50a6c 100644 --- a/src/cdtools/tools/losses/losses.py +++ b/src/cdtools/tools/losses/losses.py @@ -99,8 +99,14 @@ def intensity_mse(intensities, sim_intensities, mask=None): -def poisson_nll(intensities, sim_intensities, mask=None, eps=1e-6, subtract_min=False): - """ Returns the Poisson negative log likelihood for a simulated dataset's intensities +def poisson_nll( + intensities, + sim_intensities, + mask=None, + eps=1e-6, + subtract_min=False + ): + """ Returns the Poisson negative log likelihood for simulated intensities Calculates the overall Poisson maximum likelihood metric using diffraction intensities - the measured set of detector intensities - @@ -136,21 +142,12 @@ def poisson_nll(intensities, sim_intensities, mask=None, eps=1e-6, subtract_min= A single value for the poisson negative log likelihood """ - #When x.logy gets into the regular build, add it by uncommenting! if mask is None: - nll = t.sum(sim_intensities+eps - - intensities * t.log(sim_intensities+eps)) \ + t.xlogy(intensities,sim_intensities+eps)) \ / intensities.view(-1).shape[0] - #nll = t.sum(sim_intensities+epsa - - # t.xlogy(intensities,sim_intensities+eps)) \ - # / intensities.view(-1).shape[0] if subtract_min: - nll -= t.nansum(intensities - intensities*t.log(intensities))\ - / intensities.view(-1).shape[0] - #nll -= t.sum(intensities - t.xlogy(intensities,intensities)) - # We don't need to include the log factorial part here, because - # it will get subtracted off in the min anyway. + nll -= t.sum(intensities - t.xlogy(intensities,intensities)) return nll else: @@ -158,17 +155,97 @@ def poisson_nll(intensities, sim_intensities, mask=None, eps=1e-6, subtract_min= masked_sims = sim_intensities.masked_select(mask) nll = t.sum(masked_sims + eps - \ - masked_intensities * t.log(masked_sims+eps)) \ + t.xlogy(masked_intensities, masked_sims+eps)) \ / masked_intensities.shape[0] - - #nll = t.sum(masked_sims + eps - \ - # t.xlogy(masked_intensities, masked_sims+eps)) \ - # / masked_intensities.shape[0] - if subtract_min: - nll -= t.nansum(masked_intensities - - masked_intensities*t.log(masked_intensities)) \ - / masked_intensities.shape[0] + nll -= t.nansum(masked_intensities - \ + t.xlogy(masked_intensities, masked_intensities)) \ + / masked_intensities.shape[0] return nll + + + +def poisson_plus_fixed_nll( + intensities, + sim_intensities, + fixed_nll, + range, + mask=None, + eps=1e-6, + subtract_min=False): + """ Return a combined negative log likelihood for Poisson and fixed noise. + + This is unpublished as far as I know. First, it calculates the log + likelihoods for all the possible true photon counts within (range) + of the measured value, using the detector noise model. + + Next, it calculates the poisson log likelihood for each of those + possible true photon counts within the minimum and maximum offset + defined by range. It then compares them to the simulated poisson + reciprocal arrival rate (simulated intensity). Finally, it combines + them using a logsumexp to calculate the negative log likelihood for + the combined noise model. + + This loss is appropriate for detectors with significant fixed readout + noise which is independent of signal strength, after conversion from + detector native units to units of photons. In the simple case, + a gaussian nll can be passed to the fixed_nll function, but if the detector + noise has been well characterized, + + Note that this calculation ignores the log(intensities!) term in the + full expression for Poisson negative log likelihood. This term doesn't + change the calculated gradients so isn't worth taking the time to compute + + It can accept intensity and simulated intensity tensors of any shape + as long as their shapes match, and the provided mask array can be + broadcast correctly along them. + + The default value of eps is 1e-6 - a nonzero value here helps avoid + divergence of the log function near zero. + + Parameters + ---------- + intensities : torch.Tensor + A tensor with measured detector intensities. + sim_intensities : torch.Tensor + A tensor of simulated detector intensities + fixed_nll : function + A function which calculates the fixed negative log likelihood part + rang : tuple + A pair (min, max) defining the relative search range for photon counts + mask : torch.Tensor + 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 + + Returns + ------- + loss : torch.Tensor + A single value for the poisson negative log likelihood + + """ + raise NotImplementedError() + if mask is None: + nll = t.sum(sim_intensities+eps - + t.xlogy(intensities,sim_intensities+eps)) \ + / intensities.view(-1).shape[0] + if subtract_min: + nll -= t.sum(intensities - t.xlogy(intensities,intensities)) + + return nll + else: + masked_intensities = intensities.masked_select(mask) + masked_sims = sim_intensities.masked_select(mask) + + nll = t.sum(masked_sims + eps - \ + t.xlogy(masked_intensities, masked_sims+eps)) \ + / masked_intensities.shape[0] + + if subtract_min: + nll -= t.nansum(masked_intensities - \ + t.xlogy(masked_intensities, masked_intensities)) \ + / masked_intensities.shape[0] + + return nll \ No newline at end of file