Skip to content

Commit 16cbb44

Browse files
lcnaturemihaic
authored andcommitted
reprsimil: Add new SNR prior to GBRSA (#387)
Add an option in BRSA to set the prior of SNR to be "equal" across all voxels.
1 parent 224d2ee commit 16cbb44

2 files changed

Lines changed: 27 additions & 7 deletions

File tree

‎brainiak/reprsimil/brsa.py‎

Lines changed: 20 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -2818,11 +2818,13 @@ class GBRSA(BRSA):
28182818
In this case, the standard deviation of log(SNR) is set
28192819
by the parameter logS_range.
28202820
If set to 'unif', a uniform prior in [0,1] is imposed.
2821-
In all these cases, SNR is numerically
2821+
In all above cases, SNR is numerically
28222822
marginalized on a grid of parameters. So the parameter SNR_bins
28232823
determines how accurate the numerical integration is. The more
28242824
number of bins are used, the more accurate the numerical
28252825
integration becomes.
2826+
If set to 'equal', all voxels are assumed to have the same fixed
2827+
SNR. Pseudo-SNR is 1.0 for all voxels.
28262828
In all the cases, the grids used for pseudo-SNR do not really
28272829
set an upper bound for SNR, because the real SNR is determined
28282830
by both pseudo-SNR and U, the shared covariance structure.
@@ -3003,11 +3005,14 @@ def __init__(
30033005
if type(logS_range) is int:
30043006
logS_range = float(logS_range)
30053007
self.logS_range = logS_range
3006-
assert SNR_prior in ['unif', 'lognorm', 'exp'], \
3008+
assert SNR_prior in ['unif', 'lognorm', 'exp', 'equal'], \
30073009
'SNR_prior can only be chosen from ''unif'', ''lognorm''' \
3008-
' and ''exp'''
3010+
' ''exp'' and ''equal'''
30093011
self.SNR_prior = SNR_prior
3010-
self.SNR_bins = SNR_bins
3012+
if self.SNR_prior == 'equal':
3013+
self.SNR_bins = 1
3014+
else:
3015+
self.SNR_bins = SNR_bins
30113016
self.rho_bins = rho_bins
30123017
self.tol = tol
30133018
self.optimizer = optimizer
@@ -3094,9 +3099,14 @@ def fit(self, X, design, nuisance=None, scan_onsets=None):
30943099
# However, in fit(), we keep the scikit-learn API that
30953100
# X is the input data to fit and y, a reserved name not used, is
30963101
# the label to map to from X.
3097-
assert self.SNR_bins >= 10 and self.rho_bins >= 10, \
3102+
assert self.SNR_bins >= 10 and self.SNR_prior != 'equal' or \
3103+
self.SNR_bins == 1 and self.SNR_prior == 'equal', \
3104+
'At least 10 bins are required to perform the numerical'\
3105+
' integration over SNR, unless choosing SNR_prior=''equal'','\
3106+
' in which case SNR_bins should be 1.'
3107+
assert self.rho_bins >= 10, \
30983108
'At least 10 bins are required to perform the numerical'\
3099-
' integration over SNR and rho'
3109+
' integration over rho'
31003110
assert self.logS_range * 6 / self.SNR_bins < 0.5 \
31013111
or self.SNR_prior != 'lognorm', \
31023112
'The minimum grid of log(SNR) should not be larger than 0.5 '\
@@ -4128,9 +4138,12 @@ def _set_SNR_grids(self):
41284138
# Center of mass of each segment between consecutive
41294139
# bounds are set as the grids for SNR.
41304140
SNR_weights = np.ones(self.SNR_bins) / self.SNR_bins
4131-
else: # SNR_prior == 'exp'
4141+
elif self.SNR_prior == 'exp':
41324142
SNR_grids = self._bin_exp(self.SNR_bins)
41334143
SNR_weights = np.ones(self.SNR_bins) / self.SNR_bins
4144+
else:
4145+
SNR_grids = np.ones(1)
4146+
SNR_weights = np.ones(1)
41344147
SNR_weights = SNR_weights / np.sum(SNR_weights)
41354148
return SNR_grids, SNR_weights
41364149

‎tests/reprsimil/test_gbrsa.py‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -361,6 +361,13 @@ def test_SNR_grids():
361361
and np.all(np.diff(SNR_grids) > 0)
362362
), 'SNR_grids or SNR_weights not correct for exponential prior'
363363

364+
s = brainiak.reprsimil.brsa.GBRSA(SNR_prior='equal')
365+
SNR_grids, SNR_weights = s._set_SNR_grids()
366+
assert (np.all(SNR_grids == 1)
367+
and np.all(SNR_weights == 1)
368+
and np.size(SNR_grids) == 1
369+
), 'SNR_grids or SNR_weights not correct for equal prior'
370+
364371

365372
def test_n_nureg():
366373
import brainiak.reprsimil.brsa

0 commit comments

Comments
 (0)