|
11 | 11 | from nwbinspector import Importance, InspectorMessage |
12 | 12 | from nwbinspector.checks import ( |
13 | 13 | check_ascending_spike_times, |
| 14 | + check_data_orientation, |
14 | 15 | check_electrical_series_dims, |
15 | 16 | check_electrical_series_reference_electrodes_table, |
16 | 17 | check_negative_spike_times, |
@@ -156,44 +157,72 @@ def test_trigger_check_electrical_series_reference_electrodes_table(self): |
156 | 157 | ) |
157 | 158 |
|
158 | 159 |
|
159 | | -def test_spikeeventseries_dims_check(): |
160 | | - """ |
161 | | - Test that 2D SpikeEventSeries does not trigger a warning, |
162 | | - but 3D SpikeEventSeries with mismatched electrodes does. |
163 | | - """ |
164 | | - |
165 | | - nwbfile = NWBFile(session_description="", identifier=str(uuid4()), session_start_time=datetime.now().astimezone()) |
166 | | - device = nwbfile.create_device(name="dev") |
167 | | - group = nwbfile.create_electrode_group(name="electrode_group", description="desc", location="loc", device=device) |
168 | | - for _ in range(3): |
169 | | - nwbfile.add_electrode(x=3.0, y=3.0, z=3.0, imp=-1.0, location="unknown", filtering="unknown", group=group) |
170 | | - electrodes = nwbfile.create_electrode_table_region(region=[0, 1, 2], description="three elecs") |
171 | | - |
172 | | - # 2D data: [num_events, num_samples] (should NOT trigger warning) |
173 | | - ses_2d = SpikeEventSeries( |
174 | | - name="spike_events_2d", |
175 | | - data=np.zeros((10, 5)), |
176 | | - electrodes=electrodes, |
177 | | - timestamps=[0.1 * i for i in range(10)], |
178 | | - ) |
179 | | - assert check_electrical_series_dims(ses_2d) is None |
180 | | - |
181 | | - # 3D data: [num_events, num_channels, num_samples] with mismatched num_channels (should trigger warning) |
182 | | - ses_3d = SpikeEventSeries( |
183 | | - name="spike_events_3d", |
184 | | - data=np.zeros((10, 4, 5)), # 4 != 3 electrodes |
185 | | - electrodes=electrodes, |
186 | | - timestamps=[0.1 * i for i in range(10)], |
187 | | - ) |
188 | | - result = check_electrical_series_dims(ses_3d) |
189 | | - assert result == InspectorMessage( |
190 | | - message=("The second dimension of data does not match the length of electrodes. Your data may be transposed."), |
191 | | - importance=Importance.CRITICAL, |
192 | | - check_function_name="check_electrical_series_dims", |
193 | | - object_type="SpikeEventSeries", |
194 | | - object_name="spike_events_3d", |
195 | | - location="/", |
196 | | - ) |
| 160 | +class TestCheckSpikeEventSeries(TestCase): |
| 161 | + def setUp(self): |
| 162 | + nwbfile = NWBFile( |
| 163 | + session_description="", identifier=str(uuid4()), session_start_time=datetime.now().astimezone() |
| 164 | + ) |
| 165 | + device = nwbfile.create_device(name="dev") |
| 166 | + group = nwbfile.create_electrode_group( |
| 167 | + name="electrode_group", description="desc", location="loc", device=device |
| 168 | + ) |
| 169 | + for _ in range(3): |
| 170 | + nwbfile.add_electrode(location="unknown", group=group) |
| 171 | + self.nwbfile = nwbfile |
| 172 | + |
| 173 | + def test_check_data_orientation_spike_event_series(self): |
| 174 | + """Test that SpikeEventSeries with more waveform samples than events doesn't trigger the data orientation check.""" |
| 175 | + |
| 176 | + # create data with shape (events, channels, waveform_samples) where waveform_samples > events |
| 177 | + data = np.zeros((5, 3, 10)) |
| 178 | + timestamps = np.arange(5) |
| 179 | + electrodes = self.nwbfile.create_electrode_table_region(region=[0, 1, 2], description="three elecs") |
| 180 | + |
| 181 | + spike_event_series = SpikeEventSeries( |
| 182 | + name="spike_events", |
| 183 | + description="test spike events", |
| 184 | + data=data, |
| 185 | + timestamps=timestamps, |
| 186 | + electrodes=electrodes, |
| 187 | + ) |
| 188 | + |
| 189 | + assert check_data_orientation(spike_event_series) is None |
| 190 | + |
| 191 | + def test_spikeeventseries_dims_check(self): |
| 192 | + """ |
| 193 | + Test that 2D SpikeEventSeries does not trigger a warning, |
| 194 | + but 3D SpikeEventSeries with mismatched electrodes does. |
| 195 | + """ |
| 196 | + |
| 197 | + electrodes = self.nwbfile.create_electrode_table_region(region=[0, 1, 2], description="three elecs") |
| 198 | + |
| 199 | + # 2D data: [num_events, num_samples] (should NOT trigger warning) |
| 200 | + ses_2d = SpikeEventSeries( |
| 201 | + name="spike_events_2d", |
| 202 | + data=np.zeros((10, 5)), |
| 203 | + electrodes=electrodes, |
| 204 | + timestamps=[0.1 * i for i in range(10)], |
| 205 | + ) |
| 206 | + assert check_electrical_series_dims(ses_2d) is None |
| 207 | + |
| 208 | + # 3D data: [num_events, num_channels, num_samples] with mismatched num_channels (should trigger warning) |
| 209 | + ses_3d = SpikeEventSeries( |
| 210 | + name="spike_events_3d", |
| 211 | + data=np.zeros((10, 4, 5)), # 4 != 3 electrodes |
| 212 | + electrodes=electrodes, |
| 213 | + timestamps=[0.1 * i for i in range(10)], |
| 214 | + ) |
| 215 | + result = check_electrical_series_dims(ses_3d) |
| 216 | + assert result == InspectorMessage( |
| 217 | + message=( |
| 218 | + "The second dimension of data does not match the length of electrodes. Your data may be transposed." |
| 219 | + ), |
| 220 | + importance=Importance.CRITICAL, |
| 221 | + check_function_name="check_electrical_series_dims", |
| 222 | + object_type="SpikeEventSeries", |
| 223 | + object_name="spike_events_3d", |
| 224 | + location="/", |
| 225 | + ) |
197 | 226 |
|
198 | 227 |
|
199 | 228 | def test_check_spike_times_not_in_unobserved_interval_pass(): |
|
0 commit comments