Skip to content

Commit f1af3af

Browse files
Fix non-consecutive ROIs bug.
If ROIs were missing in the ROI mask, the ROI-IDs in the Traces table were no longer matched with the ROI-IDs in the ROI mask and the ROI table.
1 parent 7bdd7bb commit f1af3af

10 files changed

Lines changed: 156 additions & 123 deletions

File tree

djimaging/autorois/roi_canvas.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1405,7 +1405,8 @@ def split_name(stim_condition):
14051405
return stim, condition
14061406

14071407
def insert_database(self, roi_mask_tab, field_key):
1408-
from djimaging.utils.mask_utils import to_igor_format, compare_roi_masks
1408+
from djimaging.utils.mask_utils import compare_roi_masks
1409+
from djimaging.utils.mask_format_utils import to_igor_format
14091410

14101411
pres_and_roi_mask = []
14111412
for pres_key in self.pres_names:

djimaging/tables/core/roi.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55

66
from djimaging.utils.dj_utils import get_primary_key
77
from djimaging.utils.plot_utils import plot_field
8-
from djimaging.utils.scanm.roi_utils import extract_roi_idxs
8+
from djimaging.utils.scanm.roi_utils import extract_roi_ids
99

1010

1111
class RoiTemplate(dj.Computed):
@@ -59,7 +59,7 @@ def make(self, key):
5959
if not np.any(roi_mask):
6060
return
6161

62-
roi_idxs = extract_roi_idxs(roi_mask)
62+
roi_idxs = extract_roi_ids(roi_mask)
6363

6464
# add every roi to list and the bulk add to roi table
6565
for roi_idx in roi_idxs:

djimaging/tables/core/roi_mask.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -12,8 +12,9 @@
1212

1313
from djimaging.utils.filesystem_utils import as_pre_filepath
1414
from djimaging.utils.dj_utils import get_primary_key, check_unique_one
15-
from djimaging.utils.mask_utils import to_igor_format, to_python_format, to_roi_mask_file, sort_roi_mask_files, \
15+
from djimaging.utils.mask_utils import to_roi_mask_file, sort_roi_mask_files, \
1616
load_preferred_roi_mask_igor, load_preferred_roi_mask_pickle, compare_roi_masks
17+
from djimaging.utils.mask_format_utils import to_igor_format, to_python_format
1718
from djimaging.utils.plot_utils import plot_field
1819

1920

djimaging/tables/misc/sr_index.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -24,8 +24,8 @@ class SrIndex(misc.SrIndexTemplate):
2424
import numpy as np
2525
from matplotlib import pyplot as plt
2626

27+
from djimaging.utils import mask_format_utils
2728
from djimaging.utils.scanm import read_utils
28-
from djimaging.utils import mask_utils
2929
from djimaging.utils.dj_utils import get_primary_key
3030

3131

@@ -133,7 +133,7 @@ def plot1(self, key=None, sr_threshold=0.5):
133133

134134
def compute_sr_idxs(ch1_stack, roi_ids, roi_mask, npixartifact):
135135
"""Compute SR index for each ROI in roi_ids. SR index is defined as """
136-
roi_mask = mask_utils.as_python_format(roi_mask)
136+
roi_mask = mask_format_utils.as_python_format(roi_mask)
137137

138138
ch1_avg = np.mean(ch1_stack, axis=2)
139139

@@ -154,7 +154,7 @@ def compute_sr_idxs(ch1_stack, roi_ids, roi_mask, npixartifact):
154154

155155
def plot_stack_sr_idxs(ch0_avg, ch1_avg, roi_mask, roi_ids, sr_idxs, sr_threshold=0.5):
156156
"""Plot SR idxs on top of stack averages."""
157-
roi_mask = mask_utils.as_python_format(roi_mask)
157+
roi_mask = mask_format_utils.as_python_format(roi_mask)
158158

159159
binary_sr_mask = np.isin(roi_mask, roi_ids[sr_idxs >= sr_threshold])
160160

Lines changed: 93 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,93 @@
1+
import numpy as np
2+
3+
4+
def assert_igor_format(roi_mask):
5+
vmin = np.min(roi_mask)
6+
vmax = np.max(roi_mask)
7+
8+
if roi_mask.ndim != 2:
9+
raise ValueError(f'ROI mask must be 2D, but has shape {roi_mask.shape}')
10+
11+
if np.any(roi_mask != roi_mask.astype(int)):
12+
raise ValueError(f'ROI mask must be integer, but has non-integer values')
13+
14+
if not ((vmax in [0, 1]) and vmin <= 1):
15+
raise ValueError(f'ROI mask has unexpected values: vmin={vmin}, vmax={vmax}, value={np.unique(roi_mask)}')
16+
17+
18+
def is_igor_format(roi_mask):
19+
"""This method can fail if the igor format is not used consistently."""
20+
vmin = np.min(roi_mask)
21+
vmax = np.max(roi_mask)
22+
23+
if roi_mask.ndim != 2:
24+
return False
25+
26+
if np.any(roi_mask != roi_mask.astype(int)):
27+
return False
28+
29+
if vmax in [0, 1] and vmin <= 1:
30+
return True
31+
else:
32+
return False
33+
34+
35+
def as_igor_format(roi_mask):
36+
"""Convert from python format to igor format"""
37+
if is_igor_format(roi_mask):
38+
return roi_mask
39+
else:
40+
return to_igor_format(roi_mask)
41+
42+
43+
def to_igor_format(roi_mask):
44+
if not is_python_format(roi_mask):
45+
raise ValueError(f'ROI mask is not in python format; unique values in mask: {np.unique(roi_mask)}')
46+
47+
roi_mask = roi_mask.copy()
48+
roi_mask[roi_mask == 0] = -1
49+
roi_mask = -roi_mask
50+
51+
return roi_mask
52+
53+
54+
def is_python_format(roi_mask):
55+
vmin = np.min(roi_mask)
56+
vmax = np.max(roi_mask)
57+
58+
if roi_mask.ndim != 2:
59+
return False
60+
61+
if np.any(roi_mask != roi_mask.astype(int)):
62+
return False
63+
64+
if vmin == 0 and vmax >= 0:
65+
return True
66+
else:
67+
return False
68+
69+
70+
def as_python_format(roi_mask):
71+
"""Convert from igor format to python format if necessary"""
72+
if is_python_format(roi_mask):
73+
return roi_mask
74+
else:
75+
return to_python_format(roi_mask)
76+
77+
78+
def to_python_format(roi_mask):
79+
# Some ROI masks have 11 instead of 1's for some reason
80+
rm_vals = np.unique(roi_mask)
81+
if set(rm_vals[rm_vals > 0]) == {1, 11}:
82+
roi_mask = roi_mask.copy()
83+
roi_mask[roi_mask == 11] = 1
84+
85+
assert_igor_format(roi_mask)
86+
87+
vmax = np.max(roi_mask)
88+
89+
roi_mask = roi_mask.copy()
90+
roi_mask[roi_mask == vmax] = 0
91+
roi_mask = np.abs(roi_mask)
92+
93+
return roi_mask

djimaging/utils/mask_utils.py

Lines changed: 1 addition & 86 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
from matplotlib import pyplot as plt
88

99
from djimaging.utils.alias_utils import check_shared_alias_str
10+
from djimaging.utils.mask_format_utils import assert_igor_format, to_igor_format
1011
from djimaging.utils.scanm import read_h5_utils
1112
from djimaging.utils.cellpose_utils import intersection_over_union
1213

@@ -417,92 +418,6 @@ def generate_roi_suggestions(mask_pred, mask_true, n_artifact, threshold=0.1, ve
417418
return rois_to_add
418419

419420

420-
def assert_igor_format(roi_mask):
421-
vmin = np.min(roi_mask)
422-
vmax = np.max(roi_mask)
423-
424-
if roi_mask.ndim != 2:
425-
raise ValueError(f'ROI mask must be 2D, but has shape {roi_mask.shape}')
426-
427-
if np.any(roi_mask != roi_mask.astype(int)):
428-
raise ValueError(f'ROI mask must be integer, but has non-integer values')
429-
430-
if not ((vmax in [0, 1]) and vmin <= 1):
431-
raise ValueError(f'ROI mask has unexpected values: vmin={vmin}, vmax={vmax}, value={np.unique(roi_mask)}')
432-
433-
434-
def is_igor_format(roi_mask):
435-
"""This method can fail if the igor format is not used consistently."""
436-
vmin = np.min(roi_mask)
437-
vmax = np.max(roi_mask)
438-
439-
if roi_mask.ndim != 2:
440-
return False
441-
442-
if np.any(roi_mask != roi_mask.astype(int)):
443-
return False
444-
445-
if vmax in [0, 1] and vmin <= 1:
446-
return True
447-
else:
448-
return False
449-
450-
451-
def as_igor_format(roi_mask):
452-
"""Convert from python format to igor format"""
453-
if is_igor_format(roi_mask):
454-
return roi_mask
455-
else:
456-
return to_igor_format(roi_mask)
457-
458-
459-
def to_igor_format(roi_mask):
460-
if not is_python_format(roi_mask):
461-
raise ValueError(f'ROI mask is not in python format; unique values in mask: {np.unique(roi_mask)}')
462-
463-
roi_mask = roi_mask.copy()
464-
roi_mask[roi_mask == 0] = -1
465-
roi_mask = -roi_mask
466-
467-
return roi_mask
468-
469-
470-
def is_python_format(roi_mask):
471-
vmin = np.min(roi_mask)
472-
vmax = np.max(roi_mask)
473-
474-
if roi_mask.ndim != 2:
475-
return False
476-
477-
if np.any(roi_mask != roi_mask.astype(int)):
478-
return False
479-
480-
if vmin == 0 and vmax >= 0:
481-
return True
482-
else:
483-
return False
484-
485-
486-
def as_python_format(roi_mask):
487-
"""Convert from igor format to python format if necessary"""
488-
if is_python_format(roi_mask):
489-
return roi_mask
490-
else:
491-
return to_python_format(roi_mask)
492-
493-
494-
def to_python_format(roi_mask):
495-
assert_igor_format(roi_mask)
496-
497-
vmax = np.max(roi_mask)
498-
499-
roi_mask = roi_mask.copy()
500-
roi_mask[roi_mask == vmax] = 0
501-
roi_mask = np.abs(roi_mask)
502-
503-
return roi_mask
504-
505-
506421
def to_roi_mask_file(data_file, old_suffix=None, new_suffix='_ROIs.pkl',
507422
roi_mask_dir=None, old_prefix=None, new_prefix=None):
508423
"""Get ROI mask file path from data file path"""

djimaging/utils/scanm/read_h5_utils.py

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44
import numpy as np
55

66
from djimaging.utils.scanm.wparams_utils import check_dims_ch_stack_wparams
7-
from djimaging.utils.scanm.roi_utils import get_roi2trace
7+
from djimaging.utils.scanm.roi_utils import get_roi2trace, extract_roi_ids
88

99

1010
def load_stacks_and_wparams(filepath, ch_names=('wDataCh0', 'wDataCh1')) -> (dict, dict):
@@ -54,9 +54,17 @@ def load_roi2trace(filepath: str, roi_ids: np.ndarray):
5454
try:
5555
with h5py.File(filepath, "r", driver="stdio") as h5_file:
5656
traces, traces_times = extract_traces(h5_file)
57+
roi_ids_traces = extract_roi_ids(extract_roi_mask(h5_file, ignore_not_found=False), npixartifact=0)
5758
except OSError as e:
5859
raise OSError(f"Error loading file {filepath}: {e}")
59-
roi2trace = get_roi2trace(traces=traces, traces_times=traces_times, roi_ids=roi_ids)
60+
61+
if not set(roi_ids).issubset(set(roi_ids_traces)):
62+
raise ValueError(f"roi_ids {roi_ids} do not match roi_ids in traces {roi_ids_traces}")
63+
if traces.shape[-1] != len(roi_ids_traces):
64+
raise ValueError(f"Number of roi_ids {len(roi_ids_traces)} does not match traces shape {traces.shape[-1]}.")
65+
66+
roi2trace = get_roi2trace(traces=traces, traces_times=traces_times,
67+
roi_ids_traces=roi_ids_traces, roi_ids_subset=roi_ids)
6068
frame_dt = np.mean(np.diff(traces_times, axis=0))
6169
return roi2trace, frame_dt
6270

djimaging/utils/scanm/roi_utils.py

Lines changed: 26 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -2,16 +2,16 @@
22

33
import numpy as np
44

5-
from djimaging.utils.mask_utils import assert_igor_format
5+
from djimaging.utils.mask_format_utils import assert_igor_format
66

77

8-
def extract_roi_idxs(roi_mask, npixartifact=0):
8+
def extract_roi_ids(roi_mask, npixartifact=0):
99
"""Return roi idxs as in ROI mask (i.e. negative values)"""
1010
assert roi_mask.ndim == 2
11-
roi_idxs = np.unique(roi_mask[npixartifact:, :])
12-
roi_idxs = roi_idxs[roi_idxs < 0] # remove background indexes (0 or 1)
13-
roi_idxs = roi_idxs[np.argsort(np.abs(roi_idxs))] # Sort by value
14-
return roi_idxs.astype(int)
11+
roi_ids = np.unique(roi_mask[npixartifact:, :])
12+
roi_ids = roi_ids[roi_ids < 0] # remove background indexes (0 or 1)
13+
roi_ids = roi_ids[np.argsort(np.abs(roi_ids))] # Sort by value
14+
return roi_ids.astype(int)
1515

1616

1717
def fix_first_or_last_n_nan(trace, n):
@@ -25,22 +25,32 @@ def fix_first_or_last_n_nan(trace, n):
2525
return trace
2626

2727

28-
def get_roi2trace(traces, traces_times, roi_ids):
28+
def get_roi2trace(traces, traces_times, roi_ids_traces, roi_ids_subset=None):
2929
"""Get dict that holds traces and times accessible by roi_id"""
30-
assert np.all(roi_ids >= 1)
30+
roi_ids_subset = roi_ids_traces if roi_ids_subset is not None else roi_ids_subset
31+
32+
assert np.all(roi_ids_traces >= 1)
33+
assert np.all(roi_ids_subset >= 1)
34+
35+
if traces.shape != traces_times.shape:
36+
raise ValueError(f"traces shape {traces.shape} does not match traces_times shape {traces_times.shape}.")
37+
38+
if traces.shape[-1] != len(roi_ids_traces):
39+
warnings.warn(f"Number of roi_ids {len(roi_ids_traces)} does not match traces shape {traces.shape[-1]}.")
3140

3241
roi2trace = dict()
3342

34-
for roi_id in roi_ids:
35-
idx = roi_id - 1
43+
for i, roi_id in enumerate(roi_ids_traces):
44+
if roi_id not in roi_ids_subset:
45+
continue
3646

37-
if traces.ndim == 3 and idx < traces.shape[-1]:
38-
trace = traces[:, :, idx]
39-
trace_times = traces_times[:, :, idx]
47+
if traces.ndim == 3 and i < traces.shape[-1]:
48+
trace = traces[:, :, i]
49+
trace_times = traces_times[:, :, i]
4050
trace_valid = 1
41-
elif traces.ndim == 2 and idx < traces.shape[-1]:
42-
trace = traces[:, idx]
43-
trace_times = traces_times[:, idx]
51+
elif traces.ndim == 2 and i < traces.shape[-1]:
52+
trace = traces[:, i]
53+
trace_times = traces_times[:, i]
4454
trace_valid = 1
4555
else:
4656
trace_valid = 0

0 commit comments

Comments
 (0)