diff --git a/brainiak/searchlight/searchlight.py b/brainiak/searchlight/searchlight.py index db017782b..b640bb883 100644 --- a/brainiak/searchlight/searchlight.py +++ b/brainiak/searchlight/searchlight.py @@ -155,6 +155,8 @@ class Searchlight: """ def __init__(self, sl_rad=1, max_blk_edge=10, shape=Cube, min_active_voxels_proportion=0): + assert sl_rad >= 0, 'sl_rad should not be negative' + assert max_blk_edge > 0, 'max_blk_edge should be positive' self.sl_rad = sl_rad self.max_blk_edge = max_blk_edge self.min_active_voxels_proportion = min_active_voxels_proportion @@ -534,9 +536,11 @@ def _singlenode_searchlight(l, msk, mysl_rad, bcast_var, extra_params): voxel_fn = extra_params[0] shape_mask = extra_params[1] min_active_voxels_proportion = extra_params[2] - outmat = np.empty(msk.shape, dtype=np.object)[mysl_rad:-mysl_rad, - mysl_rad:-mysl_rad, - mysl_rad:-mysl_rad] + outmat = np.empty(msk.shape, dtype=np.object) + if mysl_rad > 0: + outmat = outmat[mysl_rad:-mysl_rad, + mysl_rad:-mysl_rad, + mysl_rad:-mysl_rad] for i in range(0, outmat.shape[0]): for j in range(0, outmat.shape[1]): for k in range(0, outmat.shape[2]): diff --git a/tests/searchlight/test_searchlight.py b/tests/searchlight/test_searchlight.py index 7d472e6f5..92fa2f022 100644 --- a/tests/searchlight/test_searchlight.py +++ b/tests/searchlight/test_searchlight.py @@ -206,7 +206,10 @@ def voxel_test_sfn(l, msk, myrad, bcast): def block_test_sfn(l, msk, myrad, bcast_var, extra_params): outmat = l[0][:, :, :, 0] outmat[~msk] = None - return outmat[myrad:-myrad, myrad:-myrad, myrad:-myrad] + if myrad == 0: + return outmat + else: + return outmat[myrad:-myrad, myrad:-myrad, myrad:-myrad] def test_correctness(): # noqa: C901 @@ -286,6 +289,7 @@ def do_test(dim0, dim1, dim2, ntr, nsubj, max_blk_edge, rad): block_test(data, mask, max_blk_edge, rad) do_test(dim0=7, dim1=5, dim2=9, ntr=5, nsubj=1, max_blk_edge=4, rad=1) + do_test(dim0=7, dim1=5, dim2=9, ntr=5, nsubj=1, max_blk_edge=4, rad=0) do_test(dim0=7, dim1=5, dim2=9, ntr=5, nsubj=5, max_blk_edge=4, rad=1) do_test(dim0=1, dim1=5, dim2=9, ntr=5, nsubj=5, max_blk_edge=4, rad=1) do_test(dim0=0, dim1=10, dim2=8, ntr=5, nsubj=5, max_blk_edge=4, rad=1)