Skip to content
94 changes: 41 additions & 53 deletions deeptrack/tests/test_noises.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,100 +9,88 @@
from deeptrack.image import Image
from deeptrack import noises

from deeptrack.backend import TORCH_AVAILABLE
from deeptrack.backend import TORCH_AVAILABLE, xp
from deeptrack.tests import BackendTestBase

if TORCH_AVAILABLE:
import torch

class TestNoises_Numpy(BackendTestBase):
class TestNoises_NumPy(BackendTestBase):
BACKEND = "numpy"


@property
def array_type(self):
if self.BACKEND == "numpy":
return np.ndarray
elif self.BACKEND == "torch":
return torch.Tensor
else:
raise ValueError(f"Unsupported backend: {self.BACKEND}")

def test_Offset(self):
noise = noises.Offset(offset=0.5)
input_image = Image(np.zeros((256, 256)))
input_image = Image(xp.zeros((256, 256)))
output_image = noise.resolve(input_image)

self.assertIsInstance(output_image, np.ndarray)
self.assertIsInstance(output_image, self.array_type)
self.assertEqual(output_image.shape, (256, 256))
self.assertTrue(np.all(np.array(output_image) == 0.5))
self.assertTrue(xp.all(xp.asarray(output_image) == 0.5))

def test_Background(self):
# Test with DeepTrack Image
noise = noises.Background(offset=0.5)
input_image = Image(np.zeros((256, 256)))
input_image = Image(xp.zeros((256, 256)))
output_image = noise.resolve(input_image)

self.assertIsInstance(output_image, np.ndarray)
self.assertIsInstance(output_image, self.array_type)
self.assertEqual(output_image.shape, (256, 256))
self.assertTrue(np.all(np.array(output_image) == 0.5))
self.assertTrue(xp.all(xp.asarray(output_image) == 0.5))

# Test with NumPy array
# Test with arrays
noise = noises.Background(offset=0.5)
input_image = np.ones((10, 10))
input_image = xp.ones((10, 10))
output_image = noise.resolve(input_image)
self.assertIsInstance(output_image, np.ndarray)

self.assertIsInstance(output_image, self.array_type)
self.assertEqual(output_image.shape, (10, 10))
self.assertTrue(np.all(np.array(output_image) == 1.5))
self.assertTrue(xp.all(xp.asarray(output_image) == 1.5))

def test_Gaussian(self):
noise = noises.Gaussian(mu=0.1, sigma=0.05)
input_image = Image(np.zeros((256, 256)))
input_image = Image(xp.zeros((256, 256)))
output_image = noise.resolve(input_image)
self.assertIsInstance(output_image, np.ndarray)

self.assertIsInstance(output_image, self.array_type)
self.assertEqual(output_image.shape, (256, 256))

def test_ComplexGaussian(self):
noise = noises.ComplexGaussian(mu=0.1, sigma=0.05)
input_image = Image(np.zeros((256, 256)))
input_image = Image(xp.zeros((256, 256)))
output_image = noise.resolve(input_image)
self.assertIsInstance(output_image, np.ndarray)

self.assertIsInstance(output_image, self.array_type)
self.assertEqual(output_image.shape, (256, 256))
self.assertTrue(np.iscomplexobj(output_image))
self.assertTrue(xp.any(output_image.imag != 0))

if self.BACKEND == "numpy":
self.assertTrue(np.iscomplexobj(output_image))
elif self.BACKEND == "torch":
self.assertTrue(torch.is_complex(output_image))

def test_Poisson(self):
noise = noises.Poisson(snr=20)
input_image = Image(np.ones((256, 256)) * 0.1)
input_image = xp.ones((256, 256)) * 0.1
output_image = noise.resolve(input_image)
self.assertIsInstance(output_image, np.ndarray)

self.assertIsInstance(output_image, self.array_type)
self.assertEqual(output_image.shape, (256, 256))


# Extending the test and setting the backend to torch
@unittest.skipUnless(TORCH_AVAILABLE, "PyTorch is not installed.")
class TestNoises_Torch(TestNoises_Numpy):
class TestNoises_PyTorch(TestNoises_NumPy):
BACKEND = "torch"

def test_Backgroud(self):
noise = noises.Background(offset=0.25)
input_image = torch.zeros(5,5)
output_image = noise.resolve(input_image)

self.assertIsInstance(output_image, torch.Tensor)
self.assertEqual(output_image.shape, (5,5))
self.assertTrue(torch.all(output_image == 0.25).item())
pass

def test_Gaussian(self):
noise = noises.Gaussian(mu=0.1, sigma=0.05)
input_image = torch.zeros((256, 256))
output_image = noise.resolve(input_image)
self.assertIsInstance(output_image, torch.Tensor)
self.assertEqual(output_image.shape, (256, 256))

def test_ComplexGaussian(self):
noise = noises.ComplexGaussian(mu=0.1, sigma=0.05)
input_image = torch.zeros((256, 256))
output_image = noise.resolve(input_image)
self.assertIsInstance(output_image, torch.Tensor)
self.assertEqual(output_image.shape, (256, 256))
self.assertTrue(torch.is_complex(output_image))

def test_Poisson(self):
noise = noises.Poisson(snr=20)
input_image = torch.ones((256, 256)) * 0.1
output_image = noise.resolve(input_image)

self.assertIsInstance(output_image, torch.Tensor)
self.assertEqual(output_image.shape, (256, 256))

if __name__ == "__main__":
unittest.main()