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))