From ae53fbf11f586d0579a7f0e82650c03a1b4d7c34 Mon Sep 17 00:00:00 2001 From: Stefan van der Walt Date: Sat, 22 Aug 2009 17:03:11 -0700 Subject: [PATCH] Fix and test filter API. --- scikits/image/filter/lpi_filter.py | 15 +++++++++------ scikits/image/filter/tests/test_lpi_filter.py | 5 +++++ 2 files changed, 14 insertions(+), 6 deletions(-) diff --git a/scikits/image/filter/lpi_filter.py b/scikits/image/filter/lpi_filter.py index 094c642f..16a2c34c 100644 --- a/scikits/image/filter/lpi_filter.py +++ b/scikits/image/filter/lpi_filter.py @@ -71,6 +71,9 @@ class LPIFilter2D(object): >>> filter = LPIFilter2D(filt_func) """ + if impulse_response is None: + raise ValueError("Impulse response must be a callable.") + self.impulse_response = impulse_response self.filter_params = filter_params self._cache = None @@ -152,7 +155,7 @@ def forward(data, impulse_response=None, filter_params={}, predefined_filter = LPIFilter2D(impulse_response, **filter_params) return predefined_filter(data) -def inverse(data, max_gain=2, impulse_response=None, filter_params={}, +def inverse(data, impulse_response=None, filter_params={}, max_gain=2, predefined_filter=None): """Apply the filter in reverse to the given data. @@ -160,15 +163,15 @@ def inverse(data, max_gain=2, impulse_response=None, filter_params={}, ---------- data : (M,N) ndarray Input data. + impulse_response : callable f(r, c, **filter_params) + Impulse response of the filter. See LPIFilter2D.__init__. + filter_params : dict + Additional keyword parameters to the impulse_response function. max_gain : float Limit the filter gain. Often, the filter contains zeros, which would cause the inverse filter to have infinite gain. High gain causes amplification of artefacts, so a conservative limit is recommended. - impulse_response : callable f(r, c, **filter_params) - Impulse response of the filter. See LPIFilter2D.__init__. - filter_params : dict - Additional keyword parameters to the impulse_response function. Additional Parameters --------------------- @@ -192,7 +195,7 @@ def inverse(data, max_gain=2, impulse_response=None, filter_params={}, return _centre(np.abs(ifftshift(np.dual.ifftn(G * F))), data.shape) -def wiener(data, K=0.25, impulse_response=None, filter_params={}, +def wiener(data, impulse_response=None, filter_params={}, K=0.25, predefined_filter=None): """Minimum Mean Square Error (Wiener) inverse filter. diff --git a/scikits/image/filter/tests/test_lpi_filter.py b/scikits/image/filter/tests/test_lpi_filter.py index 8b209e21..2262f1b9 100644 --- a/scikits/image/filter/tests/test_lpi_filter.py +++ b/scikits/image/filter/tests/test_lpi_filter.py @@ -41,6 +41,9 @@ class TestLPIFilter2D(): g1 = inverse(F[::-1,::-1], predefined_filter=self.f) assert ((g - g1[::-1,::-1]).sum() < 55) + g1 = inverse(F[::-1, ::-1], self.filt_func) + assert ((g - g1[::-1,::-1]).sum() < 55) + def test_wiener(self): F = self.f(self.img) g = wiener(F, predefined_filter=self.f) @@ -49,6 +52,8 @@ class TestLPIFilter2D(): g1 = wiener(F[::-1,::-1], predefined_filter=self.f) assert ((g - g1[::-1,::-1]).sum() < 1) + g1 = wiener(F[::-1,::-1], self.filt_func) + assert ((g - g1[::-1,::-1]).sum() < 1) if __name__ == "__main__": run_module_suite()