mirror of
https://github.com/cdtools-developers/cdtools.git
synced 2026-09-10 05:22:41 +02:00
111 lines
4.6 KiB
Python
111 lines
4.6 KiB
Python
import numpy as np
|
|
import torch as t
|
|
from scipy.special import xlogy
|
|
|
|
from cdtools.tools import losses
|
|
|
|
|
|
# The idea here is to use a simple numpy calculation of the various
|
|
# objective functions to check the torch implementations and make sure
|
|
# that any optimizations in the future don't change the results
|
|
|
|
|
|
def test_amplitude_mse():
|
|
# Make some fake data
|
|
data = np.random.rand(10, 100, 100)
|
|
# And add some noise to it
|
|
sim = data + 0.1 * np.random.rand(10, 100, 100)
|
|
# and define a simple mask that needs to be broadcast
|
|
mask = (np.random.rand(100, 100) > 0.1).astype(bool)
|
|
|
|
# First, test without a mask
|
|
np_result = np.sum((np.sqrt(data) - np.sqrt(sim))**2)
|
|
# np_result /= data.size
|
|
torch_result = losses.amplitude_mse(t.from_numpy(data), t.from_numpy(sim),
|
|
use_sum=True)
|
|
assert np.isclose(np_result, np.take(torch_result.numpy(), 0))
|
|
|
|
# Then, test with a mask
|
|
np_result = np.sum(mask * (np.sqrt(data) - np.sqrt(sim))**2)
|
|
# np_result /= np.count_nonzero(mask * np.ones_like(data))
|
|
torch_result = losses.amplitude_mse(t.from_numpy(data), t.from_numpy(sim),
|
|
mask=t.from_numpy(mask), use_sum=True)
|
|
assert np.isclose(np_result, np.take(torch_result.numpy(), 0))
|
|
|
|
# Now, test the version with use_sum=False, the default
|
|
|
|
# First, test without a mask
|
|
np_result = np.mean((np.sqrt(data) - np.sqrt(sim))**2)
|
|
# np_result /= data.size
|
|
torch_result = losses.amplitude_mse(t.from_numpy(data), t.from_numpy(sim))
|
|
assert np.isclose(np_result, np.take(torch_result.numpy(), 0))
|
|
|
|
# Then, test with a mask. Note that with a mask, the masked pixels
|
|
# should not contribute to the denominator for the mean.
|
|
np_result = np.sum(mask * (np.sqrt(data) - np.sqrt(sim))**2)
|
|
np_result /= np.count_nonzero(mask * np.ones_like(data))
|
|
torch_result = losses.amplitude_mse(t.from_numpy(data), t.from_numpy(sim),
|
|
mask=t.from_numpy(mask), use_sum=False)
|
|
assert np.isclose(np_result, np.take(torch_result.numpy(), 0))
|
|
|
|
|
|
|
|
def test_intensity_mse():
|
|
# Make some fake data
|
|
data = np.random.rand(10, 100, 100)
|
|
# And add some noise to it
|
|
sim = data + 0.1 * np.random.rand(10, 100, 100)
|
|
# and define a simple mask that needs to be broadcast
|
|
mask = (np.random.rand(100, 100) > 0.1).astype(bool)
|
|
|
|
# First, test without a mask
|
|
np_result = np.sum((data - sim)**2)
|
|
torch_result = losses.intensity_mse(t.from_numpy(data), t.from_numpy(sim),
|
|
use_sum=True)
|
|
assert np.isclose(np_result, np.take(torch_result.numpy(), 0))
|
|
|
|
# Then, test with a mask
|
|
np_result = np.sum(mask * (data - sim)**2)
|
|
torch_result = losses.intensity_mse(t.from_numpy(data), t.from_numpy(sim),
|
|
mask=t.from_numpy(mask), use_sum=True)
|
|
assert np.isclose(np_result, np.take(torch_result.numpy(), 0))
|
|
|
|
# Now, test the version with use_sum=False, the default
|
|
|
|
# First, test without a mask
|
|
np_result = np.sum((data - sim)**2)
|
|
np_result /= data.size
|
|
torch_result = losses.intensity_mse(t.from_numpy(data), t.from_numpy(sim))
|
|
assert np.isclose(np_result, np.take(torch_result.numpy(), 0))
|
|
|
|
# Then, test with a mask
|
|
np_result = np.sum(mask * (data - sim)**2)
|
|
np_result /= np.count_nonzero(mask * np.ones_like(data))
|
|
torch_result = losses.intensity_mse(t.from_numpy(data), t.from_numpy(sim),
|
|
mask=t.from_numpy(mask), use_sum=False)
|
|
assert np.isclose(np_result, np.take(torch_result.numpy(), 0))
|
|
|
|
|
|
def test_poisson_nll():
|
|
# Make some fake data spread over a realistic photon-count range,
|
|
# with ~5% of pixels set to zero
|
|
data = 10 * np.random.rand(10, 100, 100)
|
|
data[np.random.rand(10, 100, 100) < 0.05] = 0
|
|
# Add some noise, but set ~5% of sim pixels to exactly match data
|
|
sim = data + 0.1 * np.random.rand(10, 100, 100)
|
|
exact_match = np.random.rand(10, 100, 100) < 0.05
|
|
sim[exact_match] = data[exact_match]
|
|
# and define a simple mask that needs to be broadcast
|
|
mask = (np.random.rand(100, 100) > 0.1).astype(bool)
|
|
|
|
# First, test without a mask
|
|
np_result = np.sum(sim - xlogy(data, sim))
|
|
torch_result = losses.poisson_nll(t.from_numpy(data), t.from_numpy(sim), eps=0)
|
|
assert np.isclose(np_result, np.take(torch_result.numpy(), 0))
|
|
|
|
# Then, test with a mask
|
|
np_result = np.sum(mask * (sim - xlogy(data, sim)))
|
|
torch_result = losses.poisson_nll(t.from_numpy(data), t.from_numpy(sim),
|
|
mask=t.from_numpy(mask), eps=0)
|
|
assert np.isclose(np_result, np.take(torch_result.numpy(), 0))
|