From 06af0a221a16df09935155a2cdabcc824fa9893d Mon Sep 17 00:00:00 2001 From: Mingbo Cai Date: Fri, 18 Oct 2019 18:02:39 -0400 Subject: [PATCH 1/3] allow sl_rad to be 0 --- brainiak/searchlight/searchlight.py | 10 +++++++--- tests/searchlight/test_searchlight.py | 8 ++++++-- 2 files changed, 13 insertions(+), 5 deletions(-) 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..7b7749795 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 @@ -262,7 +265,7 @@ def block_test(data, mask, max_blk_edge, rad): sl.distribute(data, mask) sl.broadcast(mask) global_outputs = sl.run_block_function(block_test_sfn) - + print(global_outputs) if rank == 0: for d0 in range(rad, global_outputs.shape[0]-rad): for d1 in range(rad, global_outputs.shape[1]-rad): @@ -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) From c0797ef547fb739a13a0c83128b686686a3d6e9f Mon Sep 17 00:00:00 2001 From: Mingbo Cai Date: Fri, 18 Oct 2019 18:04:22 -0400 Subject: [PATCH 2/3] delete unnecessary print function --- tests/searchlight/test_searchlight.py | 1 - 1 file changed, 1 deletion(-) diff --git a/tests/searchlight/test_searchlight.py b/tests/searchlight/test_searchlight.py index 7b7749795..dd9e9cd47 100644 --- a/tests/searchlight/test_searchlight.py +++ b/tests/searchlight/test_searchlight.py @@ -265,7 +265,6 @@ def block_test(data, mask, max_blk_edge, rad): sl.distribute(data, mask) sl.broadcast(mask) global_outputs = sl.run_block_function(block_test_sfn) - print(global_outputs) if rank == 0: for d0 in range(rad, global_outputs.shape[0]-rad): for d1 in range(rad, global_outputs.shape[1]-rad): From d2747fd87b193c4debc0028ce8a5ae7c1d7d42d5 Mon Sep 17 00:00:00 2001 From: Mingbo Cai Date: Fri, 18 Oct 2019 18:05:23 -0400 Subject: [PATCH 3/3] just formatting --- tests/searchlight/test_searchlight.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/searchlight/test_searchlight.py b/tests/searchlight/test_searchlight.py index dd9e9cd47..92fa2f022 100644 --- a/tests/searchlight/test_searchlight.py +++ b/tests/searchlight/test_searchlight.py @@ -265,6 +265,7 @@ def block_test(data, mask, max_blk_edge, rad): sl.distribute(data, mask) sl.broadcast(mask) global_outputs = sl.run_block_function(block_test_sfn) + if rank == 0: for d0 in range(rad, global_outputs.shape[0]-rad): for d1 in range(rad, global_outputs.shape[1]-rad):