Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 2 additions & 3 deletions pywt/_thresholding.py
Original file line number Diff line number Diff line change
Expand Up @@ -250,7 +250,6 @@ def threshold_firm(data, value_low, value_high):
thresholded[magnitude == 0] = 0

# restore hard-thresholding behavior for values > value_high
large_vals = np.where(magnitude > value_high)
if np.any(large_vals[0]):
thresholded[large_vals] = data[large_vals]
large_vals = magnitude > value_high
thresholded[large_vals] = data[large_vals]
return thresholded
12 changes: 12 additions & 0 deletions pywt/tests/test_thresholding.py
Original file line number Diff line number Diff line change
Expand Up @@ -200,3 +200,15 @@ def test_threshold_zero_value_with_zeros():
assert_(not np.isnan(out_soft).any())
assert_(not np.isnan(out_garrote).any())
assert_(not np.isnan(out_firm).any())


def test_threshold_firm_large_value_first():
# Values above value_high must be returned unchanged even when the only
# such value sits at index 0 (of the first axis).
assert_allclose(pywt.threshold_firm(np.array([5.0, 0.5]), 1.0, 2.0),
[5.0, 0.0], rtol=1e-12)
assert_allclose(pywt.threshold_firm(np.array([[-5.0, 0.5], [0.5, 0.5]]),
1.0, 2.0),
[[-5.0, 0.0], [0.0, 0.0]], rtol=1e-12)
assert_allclose(pywt.threshold_firm(np.array([3.0 + 4.0j, 0.5]), 1.0, 2.0),
[3.0 + 4.0j, 0.0], rtol=1e-12)