diff --git a/doc/source/rawio.rst b/doc/source/rawio.rst index 54cc8dbfa..dcde5c46b 100644 --- a/doc/source/rawio.rst +++ b/doc/source/rawio.rst @@ -77,11 +77,11 @@ Then browse the internal header and display information:: nb_block: 1 nb_segment: [1] signal_channels: [V1] - unit_channels: [Wspk1u, Wspk2u, Wspk4u, Wspk5u ... Wspk29u Wspk30u Wspk31u Wspk32u] + spike_channels: [Wspk1u, Wspk2u, Wspk4u, Wspk5u ... Wspk29u Wspk30u Wspk31u Wspk32u] event_channels: [] You get the number of blocks and segments per block. You have information -about channels: **signal_channels**, **unit_channels**, **event_channels**. +about channels: **signal_channels**, **spike_channels**, **event_channels**. All this information is internally available in the *header* dict:: @@ -91,7 +91,7 @@ All this information is internally available in the *header* dict:: event_channels [] nb_segment [1] nb_block 1 - unit_channels [('Wspk1u', 'ch1#0', '', 0.00146484, 0., 0, 30000.) + spike_channels [('Wspk1u', 'ch1#0', '', 0.00146484, 0., 0, 30000.) ('Wspk2u', 'ch2#0', '', 0.00146484, 0., 0, 30000.) ... @@ -141,7 +141,7 @@ Inspect units channel. Each channel gives a SpikeTrain for each Segment. Note that for many formats a physical channel can have several units after spike sorting. So the nb_unit could be more than physical channel or signal channels. - >>> nb_unit = reader.unit_channels_count() + >>> nb_unit = reader.spike_channels_count() >>> print('nb_unit', nb_unit) nb_unit 30 >>> for unit_index in range(nb_unit): diff --git a/examples/read_files_neo_rawio.py b/examples/read_files_neo_rawio.py index 7d3a8a33c..86ef20820 100644 --- a/examples/read_files_neo_rawio.py +++ b/examples/read_files_neo_rawio.py @@ -31,8 +31,8 @@ print(float_sigs.shape, float_sigs.dtype) print(sampling_rate, t_start, units) -# Count unit and spike per units -nb_unit = reader.unit_channels_count() +# Count units and spikes per unit +nb_unit = reader.spike_channels_count() print('nb_unit', nb_unit) for unit_index in range(nb_unit): nb_spike = reader.spike_count(block_index=0, seg_index=0, unit_index=unit_index) diff --git a/neo/io/basefromrawio.py b/neo/io/basefromrawio.py index ba2549dab..e4025f753 100644 --- a/neo/io/basefromrawio.py +++ b/neo/io/basefromrawio.py @@ -134,31 +134,24 @@ def read_block(self, block_index=0, lazy=False, bl = Block(**bl_annotations) - # Group for AnalogSignals + # Group for AnalogSignals coming from signal_streams if create_group_across_segment['AnalogSignal']: - all_channels = self.header['signal_channels'] - channel_indexes_list = self.get_group_signal_channel_indexes() - sig_groups = [] - for channel_index in channel_indexes_list: - for i, (ind_within, ind_abs) in self._make_signal_channel_subgroups( - channel_index, signal_group_mode=signal_group_mode).items(): - group = Group(name='AnalogSignal group {}'.format(i)) - # @andrew @ julia @michael : do we annotate group across segment with this arrays ? - group.annotate(ch_names=all_channels[ind_abs]['name'].astype('U')) # ?? - group.annotate(channel_ids=all_channels[ind_abs]['id']) # ?? - bl.groups.append(group) - sig_groups.append(group) + signal_streams = self.header['signal_streams'] + sub_streams = self.get_sub_signal_streams(signal_group_mode) + sub_stream_groups = [] + for sub_stream in sub_streams: + stream_index, inner_stream_channels, name = sub_stream + group = Group(name=name, stream_id=signal_streams[stream_index]['id']) + bl.groups.append(group) + sub_stream_groups.append(group) if create_group_across_segment['SpikeTrain']: - unit_channels = self.header['unit_channels'] + spike_channels = self.header['spike_channels'] st_groups = [] - for c in range(unit_channels.size): + for c in range(spike_channels.size): group = Group(name='SpikeTrain group {}'.format(c)) - group.annotate(unit_name=unit_channels[c]['name']) - group.annotate(unit_id=unit_channels[c]['id']) - unit_annotations = self.raw_annotations['unit_channels'][c] - unit_annotations = check_annotations(unit_annotations) - group.annotate(**unit_annotations) + group.annotate(unit_name=spike_channels[c]['name']) + group.annotate(unit_id=spike_channels[c]['id']) bl.groups.append(group) st_groups.append(group) @@ -183,7 +176,7 @@ def read_block(self, block_index=0, lazy=False, for seg in bl.segments: if create_group_across_segment['AnalogSignal']: for c, anasig in enumerate(seg.analogsignals): - sig_groups[c].add(anasig) + sub_stream_groups[c].add(anasig) if create_group_across_segment['SpikeTrain']: for c, sptr in enumerate(seg.spiketrains): @@ -231,38 +224,35 @@ def read_segment(self, block_index=0, seg_index=0, lazy=False, signal_group_mode = self._prefered_signal_group_mode # annotations - seg_annotations = dict(self.raw_annotations['blocks'][block_index]['segments'][seg_index]) - for k in ('signals', 'units', 'events'): + seg_annotations = self.raw_annotations['blocks'][block_index]['segments'][seg_index].copy() + for k in ('signals', 'spikes', 'events'): seg_annotations.pop(k) seg_annotations = check_annotations(seg_annotations) seg = Segment(index=seg_index, **seg_annotations) # AnalogSignal - signal_channels = self.header['signal_channels'] - if signal_channels.size > 0: - channel_indexes_list = self.get_group_signal_channel_indexes() - for channel_indexes in channel_indexes_list: - for i, (ind_within, ind_abs) in self._make_signal_channel_subgroups( - channel_indexes, - signal_group_mode=signal_group_mode).items(): - # make a proxy... - anasig = AnalogSignalProxy(rawio=self, global_channel_indexes=ind_abs, - block_index=block_index, seg_index=seg_index) - - if not lazy: - # ... and get the real AnalogSIgnal if not lazy - anasig = anasig.load(time_slice=time_slice, strict_slicing=strict_slicing) - # TODO magnitude_mode='rescaled'/'raw' - - anasig.segment = seg - seg.analogsignals.append(anasig) + signal_streams = self.header['signal_streams'] + sub_streams = self.get_sub_signal_streams(signal_group_mode) + for sub_stream in sub_streams: + stream_index, inner_stream_channels, name = sub_stream + anasig = AnalogSignalProxy(rawio=self, stream_index=stream_index, + inner_stream_channels=inner_stream_channels, + block_index=block_index, seg_index=seg_index) + anasig.name = name + + if not lazy: + # ... and get the real AnalogSignal if not lazy + anasig = anasig.load(time_slice=time_slice, strict_slicing=strict_slicing) + + anasig.segment = seg + seg.analogsignals.append(anasig) # SpikeTrain and waveforms (optional) - unit_channels = self.header['unit_channels'] - for unit_index in range(len(unit_channels)): + spike_channels = self.header['spike_channels'] + for spike_channel_index in range(len(spike_channels)): # make a proxy... - sptr = SpikeTrainProxy(rawio=self, unit_index=unit_index, + sptr = SpikeTrainProxy(rawio=self, spike_channel_index=spike_channel_index, block_index=block_index, seg_index=seg_index) if not lazy: @@ -286,7 +276,7 @@ def read_segment(self, block_index=0, seg_index=0, lazy=False, seg.events.append(e) elif event_channels['type'][chan_ind] == b'epoch': e = EpochProxy(rawio=self, event_channel_index=chan_ind, - block_index=block_index, seg_index=seg_index) + block_index=block_index, seg_index=seg_index) if not lazy: e = e.load(time_slice=time_slice, strict_slicing=strict_slicing) e.segment = seg @@ -295,37 +285,50 @@ def read_segment(self, block_index=0, seg_index=0, lazy=False, seg.create_many_to_one_relationship() return seg - def _make_signal_channel_subgroups(self, channel_indexes, - signal_group_mode='group-by-same-units'): + def get_sub_signal_streams(self, signal_group_mode='group-by-same-units'): """ - For some RawIO channel are already splitted in groups. - But in any cases, channel need to be splitted again in sub groups - because they do not have the same units. - - They can also be splitted one by one to match previous behavior for - some IOs in older version of neo (<=0.5). + When signal streams don't have homogeneous SI units across channels, + they have to be split in sub streams to construct AnalogSignal objects with unique units. - This method aggregate signal channels with same units or split them all. + For backward compatibility (neo version <= 0.5) sub-streams can also be + used to generate one AnalogSignal per channel. """ - all_channels = self.header['signal_channels'] - if channel_indexes is None: - channel_indexes = np.arange(all_channels.size, dtype=int) - channels = all_channels[channel_indexes] - - groups = collections.OrderedDict() - if signal_group_mode == 'group-by-same-units': - all_units = np.unique(channels['units']) - - for i, unit in enumerate(all_units): - ind_within, = np.nonzero(channels['units'] == unit) - ind_abs = channel_indexes[ind_within] - groups[i] = (ind_within, ind_abs) - - elif signal_group_mode == 'split-all': - for i, chan_index in enumerate(channel_indexes): - ind_within = [i] - ind_abs = channel_indexes[ind_within] - groups[i] = (ind_within, ind_abs) - else: - raise (NotImplementedError) - return groups + signal_streams = self.header['signal_streams'] + signal_channels = self.header['signal_channels'] + + sub_streams = [] + for stream_index in range(len(signal_streams)): + stream_id = signal_streams[stream_index]['id'] + stream_name = signal_streams[stream_index]['name'] + mask = signal_channels['stream_id'] == stream_id + channels = signal_channels[mask] + if signal_group_mode == 'group-by-same-units': + # this does not keep the original order + _, idx = np.unique(channels['units'], return_index=True) + all_units = channels['units'][np.sort(idx)] + + if len(all_units) == 1: + # no substream + #  None iwill be transform as slice later + inner_stream_channels = None + name = stream_name + sub_stream = (stream_index, inner_stream_channels, name) + sub_streams.append(sub_stream) + else: + for units in all_units: + inner_stream_channels, = np.nonzero(channels['units'] == units) + chan_names = channels[inner_stream_channels]['name'] + name = 'Channels: (' + ' '.join(chan_names) + ')' + sub_stream = (stream_index, inner_stream_channels, name) + sub_streams.append(sub_stream) + elif signal_group_mode == 'split-all': + # mimic all neo <= 0.5 behavior + for i, channel in enumerate(channels): + inner_stream_channels = [i] + name = channels[i]['name'] + sub_stream = (stream_index, inner_stream_channels, name) + sub_streams.append(sub_stream) + else: + raise (NotImplementedError) + + return sub_streams diff --git a/neo/io/proxyobjects.py b/neo/io/proxyobjects.py index 1059b2918..b17a7c5d4 100644 --- a/neo/io/proxyobjects.py +++ b/neo/io/proxyobjects.py @@ -86,17 +86,32 @@ class AnalogSignalProxy(BaseProxy): _recommended_attrs = BaseNeo._recommended_attrs proxy_for = AnalogSignal - def __init__(self, rawio=None, global_channel_indexes=None, block_index=0, seg_index=0): + def __init__(self, rawio=None, stream_index=None, inner_stream_channels=None, + block_index=0, seg_index=0): + # stream_index: indicate the stream stream_id can be retreive easily + # inner_stream_channels: are channel index inside the stream None means all channels + # if inner_stream_channels is not None: + #  * then this is a "substream" + # * handle the case where channels have different units inside a stream + # * is related to BaseFromRaw.get_sub_signal_streams() + self._rawio = rawio self._block_index = block_index self._seg_index = seg_index - if global_channel_indexes is None: - global_channel_indexes = slice(None) - total_nb_chan = self._rawio.header['signal_channels'].size - self._global_channel_indexes = np.arange(total_nb_chan)[global_channel_indexes] + self._stream_index = stream_index + if inner_stream_channels is None: + inner_stream_channels = slice(inner_stream_channels) + self._inner_stream_channels = inner_stream_channels + + signal_streams = self._rawio.header['signal_streams'] + stream_id = signal_streams[stream_index]['id'] + signal_channels = self._rawio.header['signal_channels'] + global_inds, = np.nonzero(signal_channels['stream_id'] == stream_id) + self._nb_total_chann_in_stream = global_inds.size + self._global_channel_indexes = global_inds[inner_stream_channels] self._nb_chan = self._global_channel_indexes.size - sig_chans = self._rawio.header['signal_channels'][self._global_channel_indexes] + sig_chans = signal_channels[self._global_channel_indexes] assert np.unique(sig_chans['units']).size == 1, 'Channel do not have same units' assert np.unique(sig_chans['dtype']).size == 1, 'Channel do not have same dtype' @@ -108,10 +123,9 @@ def __init__(self, rawio=None, global_channel_indexes=None, block_index=0, seg_i self.sampling_rate = sig_chans['sampling_rate'][0] * pq.Hz self.sampling_period = 1. / self.sampling_rate sigs_size = self._rawio.get_signal_size(block_index=block_index, seg_index=seg_index, - channel_indexes=self._global_channel_indexes) + stream_index=stream_index) self.shape = (sigs_size, self._nb_chan) - self.t_start = self._rawio.get_signal_t_start(block_index, seg_index, - self._global_channel_indexes) * pq.s + self.t_start = self._rawio.get_signal_t_start(block_index, seg_index, stream_index) * pq.s # magnitude_mode='raw' is supported only if all offset=0 # and all gain are the same @@ -120,45 +134,19 @@ def __init__(self, rawio=None, global_channel_indexes=None, block_index=0, seg_i if support_raw_magnitude: str_units = ensure_signal_units(sig_chans['units'][0]).units.dimensionality.string - self._raw_units = pq.CompoundUnit('{}*{}'.format(sig_chans['gain'][0], str_units)) + gain0 = sig_chans['gain'][0] + self._raw_units = pq.CompoundUnit(f'{gain0}*{str_units}') else: self._raw_units = None - # both necessary attr and annotations - annotations = {} - annotations['name'] = self._make_name(None) - if len(sig_chans) == 1: - # when only one channel raw_annotations are set to standart annotations - d = self._rawio.raw_annotations['blocks'][block_index]['segments'][seg_index][ - 'signals'][self._global_channel_indexes[0]] - annotations.update(d) - - array_annotations = { - 'channel_names': np.array(sig_chans['name'], copy=True), - 'channel_ids': np.array(sig_chans['id'], copy=True), - } - # array annotations for signal can be at 2 places - # global at signal channel level - d = self._rawio.raw_annotations['signal_channels'] - array_annotations.update(create_analogsignal_array_annotations( - d, self._global_channel_indexes)) - # or specific to block/segment/signals - d = self._rawio.raw_annotations['blocks'][block_index]['segments'][seg_index]['signals'] - array_annotations.update(create_analogsignal_array_annotations( - d, self._global_channel_indexes)) + # retrieve annotations and array annotations + seg_ann = self._rawio.raw_annotations['blocks'][block_index]['segments'][seg_index] + annotations = seg_ann['signals'][stream_index].copy() + array_annotations = annotations.pop('__array_annotations__') + array_annotations = {k: v[inner_stream_channels] for k, v in array_annotations.items()} BaseProxy.__init__(self, array_annotations=array_annotations, **annotations) - def _make_name(self, channel_indexes): - sig_chans = self._rawio.header['signal_channels'][self._global_channel_indexes] - if channel_indexes is not None: - sig_chans = sig_chans[channel_indexes] - if len(sig_chans) == 1: - name = sig_chans['name'][0] - else: - name = 'Channel bundle ({}) '.format(','.join(sig_chans['name'])) - return name - @property def duration(self): '''Signal duration''' @@ -190,8 +178,28 @@ def load(self, time_slice=None, strict_slicing=True, (t_start or t_stop) is outside the real time range of the segment. ''' - if channel_indexes is None: - channel_indexes = slice(None) + # fixed_chan_indexes is channel index (or slice) in the stream + # channel_indexes is channel index (or slice) in the substream + if isinstance(self._inner_stream_channels, slice): + if self._inner_stream_channels == slice(None): + # sub stream is the entire stream + if channel_indexes is None: + fixed_chan_indexes = None + else: + fixed_chan_indexes = channel_indexes + else: + # sub stream is part of stream with slice + if channel_indexes is None: + fixed_chan_indexes = self._inner_stream_channels + else: + global_inds = np.arange(self._nb_total_chann_in_stream) + fixed_chan_indexes = global_inds[self._inner_stream_channels][channel_indexes] + else: + # sub stream is part of stream with indexes + if channel_indexes is None: + fixed_chan_indexes = self._inner_stream_channels + else: + fixed_chan_indexes = self._inner_stream_channels[channel_indexes] sr = self.sampling_rate @@ -227,12 +235,15 @@ def load(self, time_slice=None, strict_slicing=True, raw_signal = self._rawio.get_analogsignal_chunk(block_index=self._block_index, seg_index=self._seg_index, i_start=i_start, i_stop=i_stop, - channel_indexes=self._global_channel_indexes[channel_indexes]) + stream_index=self._stream_index, channel_indexes=fixed_chan_indexes) # if slice in channel : change name and array_annotations if raw_signal.shape[1] != self._nb_chan: - name = self._make_name(channel_indexes) - array_annotations = {k: v[channel_indexes] for k, v in self.array_annotations.items()} + name = 'slice of ' + self.name + channel_indexes2 = channel_indexes + if channel_indexes2 is None: + channel_indexes2 = slice(None) + array_annotations = {k: v[channel_indexes2] for k, v in self.array_annotations.items()} else: name = self.name array_annotations = self.array_annotations @@ -249,7 +260,8 @@ def load(self, time_slice=None, strict_slicing=True, else: dtype = 'float32' sig = self._rawio.rescale_signal_raw_to_float(raw_signal, dtype=dtype, - channel_indexes=self._global_channel_indexes[channel_indexes]) + stream_index=self._stream_index, + channel_indexes=fixed_chan_indexes) units = self.units anasig = AnalogSignal(sig, units=units, copy=False, t_start=sig_t_start, @@ -290,32 +302,29 @@ class SpikeTrainProxy(BaseProxy): _parent_objects = ('Segment', 'Unit') _quantity_attr = 'times' _necessary_attrs = (('t_start', pq.Quantity, 0), - ('t_stop', pq.Quantity, 0)) + ('t_stop', pq.Quantity, 0)) _recommended_attrs = () proxy_for = SpikeTrain - def __init__(self, rawio=None, unit_index=None, block_index=0, seg_index=0): + def __init__(self, rawio=None, spike_channel_index=None, block_index=0, seg_index=0): self._rawio = rawio self._block_index = block_index self._seg_index = seg_index - self._unit_index = unit_index + self._spike_channel_index = spike_channel_index nb_spike = self._rawio.spike_count(block_index=block_index, seg_index=seg_index, - unit_index=unit_index) + spike_channel_index=spike_channel_index) self.shape = (nb_spike, ) self.t_start = self._rawio.segment_t_start(block_index, seg_index) * pq.s self.t_stop = self._rawio.segment_t_stop(block_index, seg_index) * pq.s - # both necessary attr and annotations - annotations = {} - for k in ('name', 'id'): - annotations[k] = self._rawio.header['unit_channels'][unit_index][k] - ann = self._rawio.raw_annotations['blocks'][block_index]['segments'][seg_index]['units'][unit_index] - annotations.update(ann) + seg_ann = self._rawio.raw_annotations['blocks'][block_index]['segments'][seg_index] + annotations = seg_ann['spikes'][spike_channel_index].copy() + array_annotations = annotations.pop('__array_annotations__') - h = self._rawio.header['unit_channels'][unit_index] + h = self._rawio.header['spike_channels'][spike_channel_index] wf_sampling_rate = h['wf_sampling_rate'] if not np.isnan(wf_sampling_rate) and wf_sampling_rate > 0: self.sampling_rate = wf_sampling_rate * pq.Hz @@ -325,7 +334,7 @@ def __init__(self, rawio=None, unit_index=None, block_index=0, seg_index=0): self.sampling_rate = None self.left_sweep = None - BaseProxy.__init__(self, **annotations) + BaseProxy.__init__(self, array_annotations=array_annotations, **annotations) def load(self, time_slice=None, strict_slicing=True, magnitude_mode='rescaled', load_waveforms=False): @@ -345,8 +354,8 @@ def load(self, time_slice=None, strict_slicing=True, _t_start, _t_stop = prepare_time_slice(time_slice) spike_timestamps = self._rawio.get_spike_timestamps(block_index=self._block_index, - seg_index=self._seg_index, unit_index=self._unit_index, t_start=_t_start, - t_stop=_t_stop) + seg_index=self._seg_index, spike_channel_index=self._spike_channel_index, + t_start=_t_start, t_stop=_t_stop) if magnitude_mode == 'raw': # we must modify a bit the neo.rawio interface to also read the spike_timestamps @@ -361,11 +370,11 @@ def load(self, time_slice=None, strict_slicing=True, assert self.sampling_rate is not None, 'Do not have waveforms' raw_wfs = self._rawio.get_spike_raw_waveforms(block_index=self._block_index, - seg_index=self._seg_index, unit_index=self._unit_index, + seg_index=self._seg_index, spike_channel_index=self._spike_channel_index, t_start=_t_start, t_stop=_t_stop) if magnitude_mode == 'rescaled': float_wfs = self._rawio.rescale_waveforms_to_float(raw_wfs, - dtype='float32', unit_index=self._unit_index) + dtype='float32', spike_channel_index=self._spike_channel_index) waveforms = pq.Quantity(float_wfs, units=self._wf_units, dtype='float32', copy=False) elif magnitude_mode == 'raw': @@ -381,6 +390,12 @@ def load(self, time_slice=None, strict_slicing=True, waveforms=waveforms, left_sweep=self.left_sweep, name=self.name, file_origin=self.file_origin, description=self.description, **self.annotations) + if time_slice is None: + sptr.array_annotate(**self.array_annotations) + else: + # TODO handle array_annotations with time_slice + pass + return sptr @@ -406,10 +421,12 @@ def __init__(self, rawio=None, event_channel_index=None, block_index=0, seg_inde annotations = {} for k in ('name', 'id'): annotations[k] = self._rawio.header['event_channels'][event_channel_index][k] - ann = self._rawio.raw_annotations['blocks'][block_index]['segments'][seg_index]['events'][event_channel_index] - annotations.update(ann) - BaseProxy.__init__(self, **annotations) + seg_ann = self._rawio.raw_annotations['blocks'][block_index]['segments'][seg_index] + ann = seg_ann['events'][event_channel_index] + annotations = ann.copy() + array_annotations = annotations.pop('__array_annotations__') + BaseProxy.__init__(self, array_annotations=array_annotations, **annotations) def load(self, time_slice=None, strict_slicing=True): ''' @@ -447,6 +464,12 @@ def load(self, time_slice=None, strict_slicing=True): name=self.name, file_origin=self.file_origin, description=self.description, **self.annotations) + if time_slice is None: + ret.array_annotate(**self.array_annotations) + else: + # TODO handle array_annotations with time_slice + pass + return ret @@ -605,33 +628,3 @@ def consolidate_time_slice(time_slice, seg_t_start, seg_t_stop, strict_slicing): t_stop = ensure_second(t_stop) return (t_start, t_stop) - - -def create_analogsignal_array_annotations(sig_annotations, global_channel_indexes): - """ - Create array_annotations from raw_annoations. - Since raw_annotation are not np.array but nested dict, this func - try to find keys in raw_annotation that are shared by all channel - and make array_annotation with it. - """ - # intersection of keys across channels - common_keys = None - for ind in global_channel_indexes: - keys = [k for k, v in sig_annotations[ind].items() if not \ - isinstance(v, (list, tuple, np.ndarray))] - if common_keys is None: - common_keys = keys - else: - common_keys = [k for k in common_keys if k in keys] - - # this is redundant and done with other name - for k in ['name', 'channel_id']: - if k in common_keys: - common_keys.remove(k) - - array_annotations = {} - for k in common_keys: - values = [sig_annotations[ind][k] for ind in global_channel_indexes] - array_annotations[k] = np.array(values) - - return array_annotations diff --git a/neo/rawio/axographrawio.py b/neo/rawio/axographrawio.py index a1c31a9c6..f3396a443 100644 --- a/neo/rawio/axographrawio.py +++ b/neo/rawio/axographrawio.py @@ -154,8 +154,8 @@ Intervals". """ -from .baserawio import (BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, - _event_channel_dtype) +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype) import os from datetime import datetime @@ -267,20 +267,20 @@ def _segment_t_stop(self, block_index, seg_index): ### # signal and channel zone - def _get_signal_size(self, block_index, seg_index, channel_indexes): + def _get_signal_size(self, block_index, seg_index, stream_index): # same for all signals in all segments return len(self._raw_signals[seg_index][0]) - def _get_signal_t_start(self, block_index, seg_index, channel_indexes): + def _get_signal_t_start(self, block_index, seg_index, stream_index): # same for all signals in all segments return self._t_start def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, - channel_indexes): + stream_index, channel_indexes): if channel_indexes is None or \ np.all(channel_indexes == slice(None, None, None)): - channel_indexes = range(self.signal_channels_count()) + channel_indexes = range(self.signal_channels_count(stream_index)) raw_signals = [self._raw_signals [seg_index] @@ -381,13 +381,13 @@ def _get_event_timestamps(self, block_index, seg_index, return timestamps, durations, labels - def _rescale_event_timestamp(self, event_timestamps, dtype): + def _rescale_event_timestamp(self, event_timestamps, dtype, event_channel_index): # Scale either event or epoch start times from sample index to seconds # (t_start shouldn't be added) event_times = event_timestamps.astype(dtype) * self._sampling_period return event_times - def _rescale_epoch_duration(self, raw_duration, dtype): + def _rescale_epoch_duration(self, raw_duration, dtype, event_channel_index): # Scale epoch durations from samples to seconds epoch_durations = raw_duration.astype(dtype) * self._sampling_period return epoch_durations @@ -865,8 +865,8 @@ def _scan_axograph_file(self): # channel_info will be cast to _signal_channel_dtype channel_info = ( - name, i, 1 / sampling_period, f.byte_order + dtype, - units, gain, offset, 0) + name, str(i), 1 / sampling_period, f.byte_order + dtype, + units, gain, offset, '0') self.logger.debug('channel_info: {}'.format(channel_info)) self.logger.debug('') @@ -1309,15 +1309,18 @@ def _scan_axograph_file(self): event_channels.append(('AxoGraph Tags', '', 'event')) event_channels.append(('AxoGraph Intervals', '', 'epoch')) + if len(sig_channels) > 0: + signal_streams = [('Signals', '0')] + else: + signal_streams = [] + # organize header self.header['nb_block'] = 1 self.header['nb_segment'] = [1] - self.header['signal_channels'] = \ - np.array(sig_channels, dtype=_signal_channel_dtype) - self.header['event_channels'] = \ - np.array(event_channels, dtype=_event_channel_dtype) - self.header['unit_channels'] = \ - np.array([], dtype=_unit_channel_dtype) + self.header['signal_streams'] = np.array(signal_streams, dtype=_signal_stream_dtype) + self.header['signal_channels'] = np.array(sig_channels, dtype=_signal_channel_dtype) + self.header['event_channels'] = np.array(event_channels, dtype=_event_channel_dtype) + self.header['spike_channels'] = np.array([], dtype=_spike_channel_dtype) ############################################## # DATA OBJECTS diff --git a/neo/rawio/axonrawio.py b/neo/rawio/axonrawio.py index d094e9f4e..cef7626de 100644 --- a/neo/rawio/axonrawio.py +++ b/neo/rawio/axonrawio.py @@ -33,8 +33,8 @@ reads abf files - would be good to cross-check """ -from .baserawio import (BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, - _event_channel_dtype) +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype) import numpy as np @@ -149,7 +149,7 @@ def _parse_header(self): else: channel_ids = list(range(nbchannel)) - sig_channels = [] + signal_channels = [] adc_nums = [] for chan_index, chan_id in enumerate(channel_ids): if version < 2.: @@ -191,12 +191,14 @@ def _parse_header(self): offset -= info['listADCInfo'][chan_id]['fSignalOffset'] else: gain, offset = 1., 0. - group_id = 0 - sig_channels.append((name, chan_id, self._sampling_rate, - sig_dtype, units, gain, offset, group_id)) + stream_id = '0' + signal_channels.append((name, str(chan_id), self._sampling_rate, + sig_dtype, units, gain, offset, stream_id)) - sig_channels = np.array(sig_channels, dtype=_signal_channel_dtype) + signal_channels = np.array(signal_channels, dtype=_signal_channel_dtype) + # one unique signal stream + signal_streams = np.array([('Signals', '0')], dtype=_signal_stream_dtype) # only one events channel : tag # In ABF timstamps are not attached too any particular segment @@ -215,15 +217,16 @@ def _parse_header(self): event_channels = np.array(event_channels, dtype=_event_channel_dtype) # No spikes - unit_channels = [] - unit_channels = np.array(unit_channels, dtype=_unit_channel_dtype) + spike_channels = [] + spike_channels = np.array(spike_channels, dtype=_spike_channel_dtype) # fille into header dict self.header = {} self.header['nb_block'] = 1 self.header['nb_segment'] = [nb_segment] - self.header['signal_channels'] = sig_channels - self.header['unit_channels'] = unit_channels + self.header['signal_streams'] = signal_streams + self.header['signal_channels'] = signal_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels # insert some annotation at some place @@ -237,9 +240,9 @@ def _parse_header(self): seg_annotations = bl_annotations['segments'][seg_index] seg_annotations['abf_version'] = version - for c in range(sig_channels.size): - anasig_an = seg_annotations['signals'][c] - anasig_an['nADCNum'] = adc_nums[c] + signal_an = self.raw_annotations['blocks'][0]['segments'][seg_index]['signals'][0] + nADCNum = np.array([adc_nums[c] for c in range(signal_channels.size)]) + signal_an['__array_annotations__']['nADCNum'] = nADCNum for c in range(event_channels.size): ev_ann = seg_annotations['events'][c] @@ -256,14 +259,15 @@ def _segment_t_stop(self, block_index, seg_index): self._raw_signals[seg_index].shape[0] / self._sampling_rate return t_stop - def _get_signal_size(self, block_index, seg_index, channel_indexes): + def _get_signal_size(self, block_index, seg_index, stream_index): shape = self._raw_signals[seg_index].shape return shape[0] - def _get_signal_t_start(self, block_index, seg_index, channel_indexes): + def _get_signal_t_start(self, block_index, seg_index, stream_index): return self._t_starts[seg_index] - def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, stream_index, + channel_indexes): if channel_indexes is None: channel_indexes = slice(None) raw_signals = self._raw_signals[seg_index][slice(i_start, i_stop), channel_indexes] @@ -291,7 +295,7 @@ def _get_event_timestamps(self, block_index, seg_index, event_channel_index, t_s return timestamp, durations, labels - def _rescale_event_timestamp(self, event_timestamps, dtype): + def _rescale_event_timestamp(self, event_timestamps, dtype, event_channel_index): event_times = event_timestamps.astype(dtype) / self._sampling_rate return event_times @@ -613,6 +617,7 @@ def safe_decode_units(s): s = s.decode('utf-8') return s + BLOCKSIZE = 512 headerDescriptionV1 = [ diff --git a/neo/rawio/baserawio.py b/neo/rawio/baserawio.py index 602a4d758..f63e5d5c6 100644 --- a/neo/rawio/baserawio.py +++ b/neo/rawio/baserawio.py @@ -8,10 +8,10 @@ BaseRawIO abstract class which should be overridden to write a RawIO. -RawIO is a new API in neo that is supposed to acces as fast as possible -raw data. All IO with theses characteristics should/could be rewritten: - * internally use of memmap (or hdf5) - * reading header is quite cheap (not read all the file) +RawIO is a low level API in neo that provides fast access to the raw data. +When possible, all IOs should/implement this level following these guidelines: + * internal use of memmap (or hdf5) + * fast reading of the header (do not read the complete file) * neo tree object is symetric and logical: same channel/units/event along all block and segments. @@ -21,26 +21,28 @@ * Only one channel set for SpikeTrain (aka Unit) stable along Segment * AnalogSignal have all the same sampling_rate acroos all Segment * t_start/t_stop are the same for many object (SpikeTrain, Event) inside a Segment - * AnalogSignal should all have the same sampling_rate otherwise the won't be read - a the same time. So signal_group_mode=='split-all' in BaseFromRaw +signal channels are handled by group of "stream". +one stream will at neo.io level one AnalogSignal with multi-channel. + + +A helper class `neo.io.basefromrawio.BaseFromRaw` transform a RawIO to +neo legacy IO. In short all "neo.rawio" classes are also "neo.io" +with lazy reading capability. -A helper class `neo.io.basefromrawio.BaseFromRaw` should transform a RawIO to -neo legacy IO from free. With this API the IO have an attributes `header` with necessary keys. +This `header` attribute is done in `_parse_header(...)` method. See ExampleRawIO as example. BaseRawIO implement a possible presistent cache system that can be used by some IOs to avoid very long parse_header(). The idea is that some variable -or vector can be store somewhere (near the fiel, /tmp, any path) +or vector can be store somewhere (near the file, /tmp, any path) """ -# from __future__ import unicode_literals, print_function, division, absolute_import - import logging import numpy as np import os @@ -59,20 +61,27 @@ error_header = 'Header is not read yet, do parse_header() first' +_signal_stream_dtype = [ + ('name', 'U64'), # not necessary unique + ('id', 'U64'), # must be unique +] + _signal_channel_dtype = [ - ('name', 'U64'), - ('id', 'int64'), + ('name', 'U64'), # not necessarily unique + ('id', 'U64'), # must be unique ('sampling_rate', 'float64'), ('dtype', 'U16'), ('units', 'U64'), ('gain', 'float64'), ('offset', 'float64'), - ('group_id', 'int64'), + ('stream_id', 'U64'), ] +# TODO for later: add t_start and length in _signal_channel_dtype +# this would simplify all t_start/t_stop stuff for each RawIO class -_common_sig_characteristics = ['sampling_rate', 'dtype', 'group_id'] +_common_sig_characteristics = ['sampling_rate', 'dtype', 'stream_id'] -_unit_channel_dtype = [ +_spike_channel_dtype = [ ('name', 'U64'), ('id', 'U64'), # for waveform @@ -83,10 +92,12 @@ ('wf_sampling_rate', 'float64'), ] +# in rawio event and epoch are handled the same way +# except, that duration is `None` for events _event_channel_dtype = [ ('name', 'U64'), ('id', 'U64'), - ('type', 'S5'), # epoch ot event + ('type', 'S5'), # epoch or event ] @@ -139,15 +150,16 @@ def parse_header(self): This must create self.header['nb_block'] self.header['nb_segment'] + self.header['signal_streams'] self.header['signal_channels'] - self.header['units_channels'] + self.header['spike_channels'] self.header['event_channels'] """ self._parse_header() - self._group_signal_channel_characteristics() + self._check_stream_signal_channel_characteristics() def source_name(self): """Return fancy name of file source""" @@ -161,117 +173,136 @@ def __repr__(self): nb_seg = [self.segment_count(i) for i in range(nb_block)] txt += 'nb_segment: {}\n'.format(nb_seg) - for k in ('signal_channels', 'unit_channels', 'event_channels'): + # signal streams + v = [s['name'] + f' (chans: {self.signal_channels_count(i)})' + for i, s in enumerate(self.header['signal_streams'])] + v = pprint_vector(v) + txt += f'signal_streams: {v}\n' + + for k in ('signal_channels', 'spike_channels', 'event_channels'): ch = self.header[k] - if len(ch) > 8: - chantxt = "[{} ... {}]".format(', '.join(e for e in ch['name'][:4]), - ' '.join(e for e in ch['name'][-4:])) - else: - chantxt = "[{}]".format(', '.join(e for e in ch['name'])) - txt += '{}: {}\n'.format(k, chantxt) + v = pprint_vector(self.header[k]['name']) + txt += f'{k}: {v}\n' return txt def _generate_minimal_annotations(self): """ - Helper function that generate a nested dict - of all annotations. - must be called when these are Ok: + Helper function that generate a nested dict for annotations. + + Must be called when these are Ok after self.header is done + And so when theses function are ready: * block_count() * segment_count() + * signal_streams_count() * signal_channels_count() - * unit_channels_count() + * spike_channels_count() * event_channels_count() - Usage: - raw_annotations['blocks'][block_index] = { 'nickname' : 'super block', 'segments' : ...} - raw_annotations['blocks'][block_index] = { 'nickname' : 'super block', 'segments' : ...} - raw_annotations['blocks'][block_index]['segments'][seg_index]['signals'][channel_index] = {'nickname': 'super channel'} - raw_annotations['blocks'][block_index]['segments'][seg_index]['units'][unit_index] = {'nickname': 'super neuron'} - raw_annotations['blocks'][block_index]['segments'][seg_index]['events'][ev_chan] = {'nickname': 'super trigger'} + There are several sources and kinds of annotations that will + be forwarded to the neo.io level and used to enrich neo objects: + * annotations of objects common across segments + * signal_streams > neo.AnalogSignal annotations + * signal_channels > neo.AnalogSignal array_annotations split by stream + * spike_channels > neo.SpikeTrain + * event_channels > neo.Event and neo.Epoch + * annotations that depend of the block_id/segment_id of the object: + * nested in raw_annotations['blocks'][block_index]['segments'][seg_index]['signals'] + + Usage after a call to this function we can do this to populate more annotations: + + raw_annotations['blocks'][block_index][ 'nickname'] = 'super block' + raw_annotations['blocks'][block_index] + ['segments']['important_key'] = 'important value' + raw_annotations['blocks'][block_index] + ['segments'][seg_index] + ['signals']['nickname'] = 'super signals stream' + raw_annotations['blocks'][block_index] + ['segments'][seg_index] + ['signals']['__array_annotations__'] + ['channels_quality'] = ['bad', 'good', 'medium', 'good'] + raw_annotations['blocks'][block_index] + ['segments'][seg_index] + ['spikes'][spike_chan]['nickname'] = 'super neuron' + raw_annotations['blocks'][block_index] + ['segments'][seg_index] + ['spikes'][spike_chan] + ['__array_annotations__']['spike_amplitudes'] = [-1.2, -10., ...] + raw_annotations['blocks'][block_index] + ['segments'][seg_index] + ['events'][ev_chan]['nickname'] = 'super trigger' + raw_annotations['blocks'][block_index] + ['segments'][seg_index] + ['events'][ev_chan] + Z['__array_annotations__']['additional_label'] = ['A', 'B', 'A', 'C', ...] + Theses annotations will be used at the neo.io API directly in objects. Standard annotation like name/id/file_origin are already generated here. """ + signal_streams = self.header['signal_streams'] signal_channels = self.header['signal_channels'] - unit_channels = self.header['unit_channels'] + spike_channels = self.header['spike_channels'] event_channels = self.header['event_channels'] - a = {'blocks': [], 'signal_channels': [], 'unit_channels': [], 'event_channels': []} - for block_index in range(self.block_count()): - d = {'segments': []} - d['file_origin'] = self.source_name() - a['blocks'].append(d) - for seg_index in range(self.segment_count(block_index)): - d = {'signals': [], 'units': [], 'events': []} - d['file_origin'] = self.source_name() - a['blocks'][block_index]['segments'].append(d) - - for c in range(signal_channels.size): - # use for AnalogSignal.annotations - d = {} - d['name'] = signal_channels['name'][c] - d['channel_id'] = signal_channels['id'][c] - a['blocks'][block_index]['segments'][seg_index]['signals'].append(d) - - for c in range(unit_channels.size): - # use for SpikeTrain.annotations - d = {} - d['name'] = unit_channels['name'][c] - d['id'] = unit_channels['id'][c] - a['blocks'][block_index]['segments'][seg_index]['units'].append(d) - - for c in range(event_channels.size): - # use for Event.annotations - d = {} - d['name'] = event_channels['name'][c] - d['id'] = event_channels['id'][c] - d['file_origin'] = self._source_name() - a['blocks'][block_index]['segments'][seg_index]['events'].append(d) - - for c in range(signal_channels.size): - # use for ChannelIndex.annotations + # use for AnalogSignal.annotations and AnalogSignal.array_annotations + signal_stream_annotations = [] + for c in range(signal_streams.size): + stream_id = signal_streams[c]['id'] + channels = signal_channels[signal_channels['stream_id'] == stream_id] d = {} - d['name'] = signal_channels['name'][c] - d['channel_id'] = signal_channels['id'][c] + d['name'] = signal_streams['name'][c] + d['stream_id'] = stream_id d['file_origin'] = self._source_name() - a['signal_channels'].append(d) - - for c in range(unit_channels.size): + d['__array_annotations__'] = {} + for key in ('name', 'id'): + values = np.array([channels[key][chan] for chan in range(channels.size)]) + d['__array_annotations__']['channel_' + key + 's'] = values + signal_stream_annotations.append(d) + + # used for SpikeTrain.annotations and SpikeTrain.array_annotations + spike_annotations = [] + for c in range(spike_channels.size): # use for Unit.annotations d = {} - d['name'] = unit_channels['name'][c] - d['id'] = unit_channels['id'][c] + d['name'] = spike_channels['name'][c] + d['id'] = spike_channels['id'][c] d['file_origin'] = self._source_name() - a['unit_channels'].append(d) + d['__array_annotations__'] = {} + spike_annotations.append(d) + # used for Event/Epoch.annotations and Event/Epoch.array_annotations + event_annotations = [] for c in range(event_channels.size): # not used in neo.io at the moment could usefull one day d = {} d['name'] = event_channels['name'][c] d['id'] = event_channels['id'][c] d['file_origin'] = self._source_name() - a['event_channels'].append(d) + d['__array_annotations__'] = {} + event_annotations.append(d) - self.raw_annotations = a + # duplicate this signal_stream_annotations/spike_annotations/event_annotations + # accros blocks and segments and create annotations + ann = {} + ann['blocks'] = [] + for block_index in range(self.block_count()): + d = {} + d['file_origin'] = self.source_name() + d['segments'] = [] + ann['blocks'].append(d) - def _raw_annotate(self, obj_name, chan_index=0, block_index=0, seg_index=0, **kargs): - """ - Annotate an object in the list/dict tree annotations. - """ - bl_annotations = self.raw_annotations['blocks'][block_index] - seg_annotations = bl_annotations['segments'][seg_index] - if obj_name == 'blocks': - bl_annotations.update(kargs) - elif obj_name == 'segments': - seg_annotations.update(kargs) - elif obj_name in ['signals', 'events', 'units']: - obj_annotations = seg_annotations[obj_name][chan_index] - obj_annotations.update(kargs) - elif obj_name in ['signal_channels', 'unit_channels', 'event_channel']: - obj_annotations = self.raw_annotations[obj_name][chan_index] - obj_annotations.update(kargs) + for seg_index in range(self.segment_count(block_index)): + d = {} + d['file_origin'] = self.source_name() + # copy nested + d['signals'] = signal_stream_annotations.copy() + d['spikes'] = spike_annotations.copy() + d['events'] = event_annotations.copy() + ann['blocks'][block_index]['segments'].append(d) + + self.raw_annotations = ann def _repr_annotations(self): txt = 'Raw annotations\n' @@ -286,19 +317,30 @@ def _repr_annotations(self): seg_a = bl_a['segments'][seg_index] txt += ' *Segment {}\n'.format(seg_index) for k, v in seg_a.items(): - if k in ('signals', 'units', 'events',): + if k in ('signals', 'spikes', 'events',): continue txt += ' -{}: {}\n'.format(k, v) - for child in ('signals', 'units', 'events'): - n = self.header[child[:-1] + '_channels'].shape[0] + # annotations by channels for spikes/events/epochs + for child in ('signals', 'events', 'spikes', ): + if child == 'signals': + n = self.header['signal_streams'].shape[0] + else: + n = self.header[child[:-1] + '_channels'].shape[0] for c in range(n): neo_name = {'signals': 'AnalogSignal', - 'units': 'SpikeTrain', 'events': 'Event/Epoch'}[child] - txt += ' *{} {}\n'.format(neo_name, c) + 'spikes': 'SpikeTrain', + 'events': 'Event/Epoch'}[child] + txt += f' *{neo_name} {c}\n' child_a = seg_a[child][c] for k, v in child_a.items(): - txt += ' -{}: {}\n'.format(k, v) + if k == '__array_annotations__': + continue + txt += f' -{k}: {v}\n' + for k, values in child_a['__array_annotations__'].items(): + values = ', '.join([str(v) for v in values[:4]]) + values = '[ ' + values + ' ...' + txt += f' -{k}: {values}\n' return txt @@ -314,17 +356,26 @@ def segment_count(self, block_index): """return number of segment for a given block""" return self.header['nb_segment'][block_index] - def signal_channels_count(self): - """Return the number of signal channels. + def signal_streams_count(self): + """Return the number of signal stream channels. Same along all Blocks and Segments. """ - return len(self.header['signal_channels']) + return len(self.header['signal_streams']) - def unit_channels_count(self): + def signal_channels_count(self, stream_index): + """Return the number of signal channels for a given stream + This number is constant across Blocks and Segments. + """ + stream_id = self.header['signal_streams'][stream_index]['id'] + channels = self.header['signal_channels'] + channels = channels[channels['stream_id'] == stream_id] + return len(channels) + + def spike_channels_count(self): """Return the number of unit (aka spike) channels. Same along all Blocks and Segment. """ - return len(self.header['unit_channels']) + return len(self.header['spike_channels']) def event_channels_count(self): """Return the number of event/epoch channels. @@ -347,152 +398,144 @@ def segment_t_stop(self, block_index, seg_index): ### # signal and channel zone - def _group_signal_channel_characteristics(self): + def _check_stream_signal_channel_characteristics(self): """ - Useful for few IOs (TdtrawIO, NeuroExplorerRawIO, ...). - - Group signals channels by same characteristics: + This check that all channel that belong to the same stream_id + have common characteristics: + * stream_id (explicite channel group) * sampling_rate (global along block and segment) - * group_id (explicite channel group) - - If all channels have the same characteristics then - `get_analogsignal_chunk` can be call wihtout restriction. - If not, then **channel_indexes** must be specified - in `get_analogsignal_chunk` and only channels with same - characteristics can be read at the same time. - - This is useful for some IO than - have internally several signals channels family. - - For many RawIO all channels have the same - sampling_rate/size/t_start. In that cases, internal flag - **self._several_channel_groups will be set to False, so - `get_analogsignal_chunk(..)` won't suffer in performance. - - Note that at neo.io level this have an impact on - `signal_group_mode`. 'split-all' will work in any situation - But grouping channel in the same AnalogSignal - with 'group-by-XXX' will depend on common characteristics - of course. - - """ - - characteristics = self.header['signal_channels'][_common_sig_characteristics] - unique_characteristics = np.unique(characteristics) - if len(unique_characteristics) == 1: - self._several_channel_groups = False - else: - self._several_channel_groups = True - - def _check_common_characteristics(self, channel_indexes): - """ - Useful for few IOs (TdtrawIO, NeuroExplorerRawIO, ...). - - Check that a set a signal channel_indexes share common - characteristics (**sampling_rate/t_start/size**). - Useful only when RawIO propose differents channels groups - with different sampling_rate for instance. + * units + * dtype """ - # ~ print('_check_common_characteristics', channel_indexes) + signal_streams = self.header['signal_streams'] + signal_channels = self.header['signal_channels'] + if signal_streams.size > 0: + assert signal_channels.size > 0, 'Signal stream but no signal_channels!!!' - assert channel_indexes is not None, \ - 'You must specify channel_indexes' - characteristics = self.header['signal_channels'][_common_sig_characteristics] - # ~ print(characteristics[channel_indexes]) - assert np.unique(characteristics[channel_indexes]).size == 1, \ - 'This channel set has varied characteristics' + for stream_index in range(signal_streams.size): + stream_id = signal_streams[stream_index]['id'] + mask = signal_channels['stream_id'] == stream_id + characteristics = signal_channels[mask][_common_sig_characteristics] + unique_characteristics = np.unique(characteristics) + assert unique_characteristics.size == 1, \ + f'Some channel in stream_id {stream_id} ' \ + f'do not have same {_common_sig_characteristics} {unique_characteristics}' - def get_group_signal_channel_indexes(self): - """ - Useful for few IOs (TdtrawIO, NeuroExplorerRawIO, ...). + # also check that id is unique inside a stream + channel_ids = signal_channels[mask]['id'] + assert np.unique(channel_ids).size == channel_ids.size, \ + f'signal_channels dont have unique ids for stream {stream_index}' - Return a list of channel_indexes than have same characteristics - """ - if self._several_channel_groups: - characteristics = self.header['signal_channels'][_common_sig_characteristics] - unique_characteristics = np.unique(characteristics) - channel_indexes_list = [] - for e in unique_characteristics: - channel_indexes, = np.nonzero(characteristics == e) - channel_indexes_list.append(channel_indexes) - return channel_indexes_list - else: - return [None] + self._several_channel_groups = signal_streams.size > 1 - def channel_name_to_index(self, channel_names): + def channel_name_to_index(self, stream_index, channel_names): """ - Transform channel_names to channel_indexes. + Inside a stream, transform channel_names to channel_indexes. Based on self.header['signal_channels'] - """ - ch = self.header['signal_channels'] - channel_indexes, = np.nonzero(np.in1d(ch['name'], channel_names)) - assert len(channel_indexes) == len(channel_names), 'not match' + channel_indexes is local to stream + """ + stream_id = self.header['signal_streams'][stream_index]['id'] + mask = self.header['signal_channels']['stream_id'] == stream_id + signal_channels = self.header['signal_channels'][mask] + chan_names = list(signal_channels['name']) + assert signal_channels.size == np.unique(chan_names).size, 'Channel names not unique' + channel_indexes = np.array([chan_names.index(name) for name in channel_names]) return channel_indexes - def channel_id_to_index(self, channel_ids): + def channel_id_to_index(self, stream_index, channel_ids): """ - Transform channel_ids to channel_indexes. + Inside a stream, transform channel_ids to channel_indexes. Based on self.header['signal_channels'] - """ - ch = self.header['signal_channels'] - channel_indexes, = np.nonzero(np.in1d(ch['id'], channel_ids)) - assert len(channel_indexes) == len(channel_ids), 'not match' + channel_indexes is local to stream + """ + # unique ids is already check in _check_stream_signal_channel_characteristics + stream_id = self.header['signal_streams'][stream_index]['id'] + mask = self.header['signal_channels']['stream_id'] == stream_id + signal_channels = self.header['signal_channels'][mask] + chan_ids = list(signal_channels['id']) + channel_indexes = np.array([chan_ids.index(chan_id) for chan_id in channel_ids]) return channel_indexes - def _get_channel_indexes(self, channel_indexes, channel_names, channel_ids): + def _get_channel_indexes(self, stream_index, channel_indexes, channel_names, channel_ids): """ Select channel_indexes from channel_indexes/channel_names/channel_ids depending which is not None. """ if channel_indexes is None and channel_names is not None: - channel_indexes = self.channel_name_to_index(channel_names) - - if channel_indexes is None and channel_ids is not None: - channel_indexes = self.channel_id_to_index(channel_ids) - + channel_indexes = self.channel_name_to_index(stream_index, channel_names) + elif channel_indexes is None and channel_ids is not None: + channel_indexes = self.channel_id_to_index(stream_index, channel_ids) return channel_indexes - def get_signal_size(self, block_index, seg_index, channel_indexes=None): - if self._several_channel_groups: - self._check_common_characteristics(channel_indexes) - return self._get_signal_size(block_index, seg_index, channel_indexes) - - def get_signal_t_start(self, block_index, seg_index, channel_indexes=None): - if self._several_channel_groups: - self._check_common_characteristics(channel_indexes) - return self._get_signal_t_start(block_index, seg_index, channel_indexes) - - def get_signal_sampling_rate(self, channel_indexes=None): - if self._several_channel_groups: - self._check_common_characteristics(channel_indexes) - chan_index0 = channel_indexes[0] + def _get_stream_index(self, stream_index): + if stream_index is None: + assert self.header['signal_streams'].size == 1 + stream_index = 0 else: - chan_index0 = 0 - sr = self.header['signal_channels'][chan_index0]['sampling_rate'] + assert 0 <= stream_index < self.header['signal_streams'].size + return stream_index + + def get_signal_size(self, block_index, seg_index, stream_index=None): + stream_index = self._get_stream_index(stream_index) + return self._get_signal_size(block_index, seg_index, stream_index) + + def get_signal_t_start(self, block_index, seg_index, stream_index=None): + stream_index = self._get_stream_index(stream_index) + return self._get_signal_t_start(block_index, seg_index, stream_index) + + def get_signal_sampling_rate(self, stream_index=None): + stream_index = self._get_stream_index(stream_index) + stream_id = self.header['signal_streams'][stream_index]['id'] + mask = self.header['signal_channels']['stream_id'] == stream_id + signal_channels = self.header['signal_channels'][mask] + sr = signal_channels[0]['sampling_rate'] return float(sr) def get_analogsignal_chunk(self, block_index=0, seg_index=0, i_start=None, i_stop=None, - channel_indexes=None, channel_names=None, channel_ids=None): + stream_index=None, channel_indexes=None, channel_names=None, + channel_ids=None, prefer_slice=False): """ Return a chunk of raw signal. """ - channel_indexes = self._get_channel_indexes(channel_indexes, channel_names, channel_ids) - if self._several_channel_groups: - self._check_common_characteristics(channel_indexes) + stream_index = self._get_stream_index(stream_index) + channel_indexes = self._get_channel_indexes(stream_index, channel_indexes, + channel_names, channel_ids) + + # some check on channel_indexes + if isinstance(channel_indexes, list): + channel_indexes = np.asarray(channel_indexes) + + if isinstance(channel_indexes, np.ndarray): + if channel_indexes.dtype == 'bool': + assert self.signal_channels_count(stream_index) == channel_indexes.size + channel_indexes, = np.nonzero(channel_indexes) + + if prefer_slice and isinstance(channel_indexes, np.ndarray): + # check if channel_indexes are coninuous and transform to slice + # this is usefull for memmap or hdf5 where slice make read lazy + # contrary to indexes that make a copy (like numpy.take()) + if np.all(np.diff(channel_indexes) == 1): + channel_indexes = slice(channel_indexes[0], channel_indexes[-1] + 1) raw_chunk = self._get_analogsignal_chunk( - block_index, seg_index, i_start, i_stop, channel_indexes) + block_index, seg_index, i_start, i_stop, stream_index, channel_indexes) return raw_chunk - def rescale_signal_raw_to_float(self, raw_signal, dtype='float32', + def rescale_signal_raw_to_float(self, raw_signal, dtype='float32', stream_index=None, channel_indexes=None, channel_names=None, channel_ids=None): - - channel_indexes = self._get_channel_indexes(channel_indexes, channel_names, channel_ids) + stream_index = self._get_stream_index(stream_index) + channel_indexes = self._get_channel_indexes(stream_index, channel_indexes, + channel_names, channel_ids) if channel_indexes is None: channel_indexes = slice(None) - channels = self.header['signal_channels'][channel_indexes] + stream_id = self.header['signal_streams'][stream_index]['id'] + mask = self.header['signal_channels']['stream_id'] == stream_id + channels = self.header['signal_channels'][mask] + if channel_indexes is None: + channel_indexes = slice(None) + channels = channels[channel_indexes] float_signal = raw_signal.astype(dtype) @@ -505,10 +548,10 @@ def rescale_signal_raw_to_float(self, raw_signal, dtype='float32', return float_signal # spiketrain and unit zone - def spike_count(self, block_index=0, seg_index=0, unit_index=0): - return self._spike_count(block_index, seg_index, unit_index) + def spike_count(self, block_index=0, seg_index=0, spike_channel_index=0): + return self._spike_count(block_index, seg_index, spike_channel_index) - def get_spike_timestamps(self, block_index=0, seg_index=0, unit_index=0, + def get_spike_timestamps(self, block_index=0, seg_index=0, spike_channel_index=0, t_start=None, t_stop=None): """ The timestamp is as close to the format itself. Sometimes float/int32/int64. @@ -516,9 +559,9 @@ def get_spike_timestamps(self, block_index=0, seg_index=0, unit_index=0, The conversion to second or index_on_signal is done outside here. t_start/t_sop are limits in seconds. - """ - timestamp = self._get_spike_timestamps(block_index, seg_index, unit_index, t_start, t_stop) + timestamp = self._get_spike_timestamps(block_index, seg_index, + spike_channel_index, t_start, t_stop) return timestamp def rescale_spike_timestamp(self, spike_timestamps, dtype='float64'): @@ -528,14 +571,16 @@ def rescale_spike_timestamp(self, spike_timestamps, dtype='float64'): return self._rescale_spike_timestamp(spike_timestamps, dtype) # spiketrain waveform zone - def get_spike_raw_waveforms(self, block_index=0, seg_index=0, unit_index=0, + def get_spike_raw_waveforms(self, block_index=0, seg_index=0, spike_channel_index=0, t_start=None, t_stop=None): - wf = self._get_spike_raw_waveforms(block_index, seg_index, unit_index, t_start, t_stop) + wf = self._get_spike_raw_waveforms(block_index, seg_index, + spike_channel_index, t_start, t_stop) return wf - def rescale_waveforms_to_float(self, raw_waveforms, dtype='float32', unit_index=0): - wf_gain = self.header['unit_channels']['wf_gain'][unit_index] - wf_offset = self.header['unit_channels']['wf_offset'][unit_index] + def rescale_waveforms_to_float(self, raw_waveforms, dtype='float32', + spike_channel_index=0): + wf_gain = self.header['spike_channels']['wf_gain'][spike_channel_index] + wf_offset = self.header['spike_channels']['wf_offset'][spike_channel_index] float_waveforms = raw_waveforms.astype(dtype) @@ -569,17 +614,19 @@ def get_event_timestamps(self, block_index=0, seg_index=0, event_channel_index=0 block_index, seg_index, event_channel_index, t_start, t_stop) return timestamp, durations, labels - def rescale_event_timestamp(self, event_timestamps, dtype='float64'): + def rescale_event_timestamp(self, event_timestamps, dtype='float64', + event_channel_index=0): """ Rescale event timestamps to s """ - return self._rescale_event_timestamp(event_timestamps, dtype) + return self._rescale_event_timestamp(event_timestamps, dtype, event_channel_index) - def rescale_epoch_duration(self, raw_duration, dtype='float64'): + def rescale_epoch_duration(self, raw_duration, dtype='float64', + event_channel_index=0): """ Rescale epoch raw duration to s """ - return self._rescale_epoch_duration(raw_duration, dtype) + return self._rescale_epoch_duration(raw_duration, dtype, event_channel_index) def setup_cache(self, cache_path, **init_kargs): if self.rawmode in ('one-file', 'multi-file'): @@ -653,7 +700,7 @@ def _segment_t_stop(self, block_index, seg_index): ### # signal and channel zone - def _get_signal_size(self, block_index, seg_index, channel_indexes): + def _get_signal_size(self, block_index, seg_index, stream_index): """ Return the size of a set of AnalogSignals indexed by channel_indexes. @@ -661,7 +708,7 @@ def _get_signal_size(self, block_index, seg_index, channel_indexes): """ raise (NotImplementedError) - def _get_signal_t_start(self, block_index, seg_index, channel_indexes): + def _get_signal_t_start(self, block_index, seg_index, stream_index): """ Return the t_start of a set of AnalogSignals indexed by channel_indexes. @@ -669,9 +716,11 @@ def _get_signal_t_start(self, block_index, seg_index, channel_indexes): """ raise (NotImplementedError) - def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, + stream_index, channel_indexes): """ - Return the samples from a set of AnalogSignals indexed by channel_indexes. + Return the samples from a set of AnalogSignals indexed + by channel_indexes (local index inner stream). All channels indexed must have the same size and t_start. @@ -684,10 +733,11 @@ def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, chann ### # spiketrain and unit zone - def _spike_count(self, block_index, seg_index, unit_index): + def _spike_count(self, block_index, seg_index, spike_channel_index): raise (NotImplementedError) - def _get_spike_timestamps(self, block_index, seg_index, unit_index, t_start, t_stop): + def _get_spike_timestamps(self, block_index, seg_index, + spike_channel_index, t_start, t_stop): raise (NotImplementedError) def _rescale_spike_timestamp(self, spike_timestamps, dtype): @@ -695,7 +745,8 @@ def _rescale_spike_timestamp(self, spike_timestamps, dtype): ### # spike waveforms zone - def _get_spike_raw_waveforms(self, block_index, seg_index, unit_index, t_start, t_stop): + def _get_spike_raw_waveforms(self, block_index, seg_index, + spike_channel_index, t_start, t_stop): raise (NotImplementedError) ### @@ -711,3 +762,16 @@ def _rescale_event_timestamp(self, event_timestamps, dtype): def _rescale_epoch_duration(self, raw_duration, dtype): raise (NotImplementedError) + + +def pprint_vector(vector, lim=8): + vector = np.asarray(vector) + assert vector.ndim == 1 + if len(vector) > lim: + part1 = ', '.join(e for e in vector[:lim // 2]) + part2 = ' , '.join(e for e in vector[-lim // 2:]) + txt = f"[{part1} ... {part2}]" + else: + part1 = ', '.join(e for e in vector[:lim // 2]) + txt = f"[{part1}]" + return txt diff --git a/neo/rawio/bci2000rawio.py b/neo/rawio/bci2000rawio.py index 2b3b3041d..8299cf5a0 100644 --- a/neo/rawio/bci2000rawio.py +++ b/neo/rawio/bci2000rawio.py @@ -3,7 +3,9 @@ https://www.bci2000.org/mediawiki/index.php/Technical_Reference:BCI2000_File_Format """ -from .baserawio import BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, _event_channel_dtype +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype) + import numpy as np import re @@ -36,11 +38,17 @@ def _parse_header(self): self.header['nb_block'] = 1 self.header['nb_segment'] = [1] + # one unique stream + signal_streams = np.array([('Signals', '0')], dtype=_signal_stream_dtype) + self.header['signal_streams'] = signal_streams + sig_channels = [] for chan_ix in range(file_info['SourceCh']): - ch_name = param_defs['ChannelNames']['value'][chan_ix] \ - if 'ChannelNames' in param_defs and param_defs['ChannelNames']['value'] is not np.nan else 'ch' + str(chan_ix) - chan_id = chan_ix + 1 + if 'ChannelNames' in param_defs and not np.isnan(param_defs['ChannelNames']['value']): + ch_name = param_defs['ChannelNames']['value'][chan_ix] + else: + ch_name = 'ch' + str(chan_ix) + chan_id = str(chan_ix + 1) sr = param_defs['SamplingRate']['value'] # Hz dtype = file_info['DataFormat'] units = 'uV' @@ -58,11 +66,11 @@ def _parse_header(self): if isinstance(offset, str): offset = float(offset) - group_id = 0 - sig_channels.append((ch_name, chan_id, sr, dtype, units, gain, offset, group_id)) + stream_id = '0' + sig_channels.append((ch_name, chan_id, sr, dtype, units, gain, offset, stream_id)) self.header['signal_channels'] = np.array(sig_channels, dtype=_signal_channel_dtype) - self.header['unit_channels'] = np.array([], dtype=_unit_channel_dtype) + self.header['spike_channels'] = np.array([], dtype=_spike_channel_dtype) # creating event channel for each state variable event_channels = [] @@ -79,7 +87,8 @@ def _parse_header(self): 'file_info': file_info, 'param_defs': param_defs }) - for ev_ix, ev_dict in enumerate(self.raw_annotations['event_channels']): + event_annotations = self.raw_annotations['blocks'][0]['segments'][0]['events'] + for ev_ix, ev_dict in enumerate(event_annotations): ev_dict.update({ 'length': state_defs[ev_ix][1], 'startVal': state_defs[ev_ix][2], @@ -125,13 +134,16 @@ def _segment_t_start(self, block_index, seg_index): def _segment_t_stop(self, block_index, seg_index): return self._read_info['n_samps'] / self._read_info['sampling_rate'] - def _get_signal_size(self, block_index, seg_index, channel_indexes=None): + def _get_signal_size(self, block_index, seg_index, stream_index): + assert stream_index == 0 return self._read_info['n_samps'] def _get_signal_t_start(self, block_index, seg_index, channel_indexes): return 0. - def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, + stream_index, channel_indexes): + assert stream_index == 0 if i_start is None: i_start = 0 if i_stop is None: @@ -170,19 +182,22 @@ def _get_event_timestamps(self, block_index, seg_index, event_channel_index, t_s keep = np.logical_and(keep, ts <= t_stop) return ts[keep], dur[keep], labels[keep] - def _rescale_event_timestamp(self, event_timestamps, dtype): + def _rescale_event_timestamp(self, event_timestamps, dtype, event_channel_index): event_times = (event_timestamps / float(self._read_info['sampling_rate'])).astype(dtype) return event_times - def _rescale_epoch_duration(self, raw_duration, dtype): + def _rescale_epoch_duration(self, raw_duration, dtype, event_channel_index): durations = (raw_duration / float(self._read_info['sampling_rate'])).astype(dtype) return durations @property def _event_arrays_list(self): if self._my_events is None: + event_annotations = self.raw_annotations['blocks'][0]['segments'][0]['events'] + self._my_events = [] - for s_ix, sd in enumerate(self.raw_annotations['event_channels']): + for event_channel_index in range(self.event_channels_count()): + sd = event_annotations[event_channel_index] ev_times = durs = vals = np.array([]) # Skip these big but mostly useless (?) states. if sd['name'] not in ['SourceTime', 'StimulusTime']: diff --git a/neo/rawio/blackrockrawio.py b/neo/rawio/blackrockrawio.py index 54c23276a..dce9bf639 100644 --- a/neo/rawio/blackrockrawio.py +++ b/neo/rawio/blackrockrawio.py @@ -65,8 +65,8 @@ import numpy as np import quantities as pq -from .baserawio import (BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, - _event_channel_dtype) +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype) class BlackrockRawIO(BaseRawIO): @@ -223,8 +223,9 @@ def _parse_header(self): main_sampling_rate = 30000. event_channels = [] - unit_channels = [] - sig_channels = [] + spike_channels = [] + signal_streams = [] + signal_channels = [] # Step1 NEV file if self._avail_files['nev']: @@ -241,7 +242,7 @@ def _parse_header(self): spikes, spike_segment_ids = self.nev_data['Spikes'] # scan all channel to get number of Unit - unit_channels = [] + spike_channels = [] self.internal_unit_ids = [] # pair of chan['packet_id'], spikes['unit_class_nb'] for i in range(len(self.__nev_ext_header[b'NEUEVWAV'])): @@ -261,7 +262,7 @@ def _parse_header(self): # default value: threshold crossing after 10 samples of waveform wf_left_sweep = 10 wf_sampling_rate = main_sampling_rate - unit_channels.append((name, _id, wf_units, wf_gain, + spike_channels.append((name, _id, wf_units, wf_gain, wf_offset, wf_left_sweep, wf_sampling_rate)) # scan events @@ -311,7 +312,7 @@ def _parse_header(self): raise(ValueError('nsx_to_load is wrong')) assert all(nsx_nb in self._avail_nsx for nsx_nb in self.nsx_to_load),\ - 'nsx_to_load do not match available nsx list' + 'nsx_to_load do not match available nsx list' # check that all files come from the same specification all_spec = [self.__nsx_spec[nsx_nb] for nsx in self.nsx_to_load] @@ -334,9 +335,6 @@ def _parse_header(self): for nsx_nb in self.nsx_to_load: self.__match_nsx_and_nev_segment_ids(nsx_nb) - # usefull to get local channel index in nsX from the global channel index - local_sig_indexes = [] - self.nsx_datas = {} self.sig_sampling_rates = {} if len(self.nsx_to_load) > 0: @@ -360,14 +358,16 @@ def _parse_header(self): d[key] = params[key][i] ext_header.append(d) + if len(ext_header) > 0: + signal_streams.append((f'nsx{nsx_nb}', str(nsx_nb))) for i, chan in enumerate(ext_header): if spec in ['2.2', '2.3']: ch_name = chan['electrode_label'].decode() - ch_id = chan['electrode_id'] + ch_id = str(chan['electrode_id']) units = chan['units'].decode() elif spec == '2.1': ch_name = chan['labels'] - ch_id = self.__nsx_ext_header[nsx_nb][i]['electrode_id'] + ch_id = str(self.__nsx_ext_header[nsx_nb][i]['electrode_id']) units = chan['units'] sig_dtype = 'int16' # max_analog_val/min_analog_val/max_digital_val/min_analog_val are int16!!!!! @@ -380,17 +380,14 @@ def _parse_header(self): (float(chan['max_digital_val']) - float(chan['min_digital_val'])) offset = -float(chan['min_digital_val']) \ * gain + float(chan['min_analog_val']) - group_id = nsx_nb - sig_channels.append((ch_name, ch_id, sr, sig_dtype, - units, gain, offset, group_id,)) - local_sig_indexes.extend(range(len(ext_header))) - - self._local_sig_indexes = np.array(local_sig_indexes) + stream_id = str(nsx_nb) + signal_channels.append((ch_name, ch_id, sr, sig_dtype, + units, gain, offset, stream_id)) # check nb segment per nsx nb_segments_for_nsx = [len(self.nsx_datas[nsx_nb]) for nsx_nb in self.nsx_to_load] assert all(nb == nb_segments_for_nsx[0] for nb in nb_segments_for_nsx),\ - 'Segment nb not consistanent across nsX files' + 'Segment nb not consistanent across nsX files' self._nb_segment = nb_segments_for_nsx[0] self.__delete_empty_segments() @@ -458,15 +455,17 @@ def _parse_header(self): self._sigs_t_starts = [None] * self._nb_segment # finalize header - unit_channels = np.array(unit_channels, dtype=_unit_channel_dtype) + spike_channels = np.array(spike_channels, dtype=_spike_channel_dtype) event_channels = np.array(event_channels, dtype=_event_channel_dtype) - sig_channels = np.array(sig_channels, dtype=_signal_channel_dtype) + signal_channels = np.array(signal_channels, dtype=_signal_channel_dtype) + signal_streams = np.array(signal_streams, dtype=_signal_stream_dtype) self.header = {} self.header['nb_block'] = 1 self.header['nb_segment'] = [self._nb_segment] - self.header['signal_channels'] = sig_channels - self.header['unit_channels'] = unit_channels + self.header['signal_streams'] = signal_streams + self.header['signal_channels'] = signal_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels rec_datetime = self.__nev_params('rec_datetime') if self._avail_files['nev'] else None @@ -474,6 +473,7 @@ def _parse_header(self): # Put annotations at some places for compatibility # with previous BlackrockIO version self._generate_minimal_annotations() + block_ann = self.raw_annotations['blocks'][0] block_ann['description'] = 'Block of data from Blackrock file set.' block_ann['file_origin'] = self.filename @@ -487,22 +487,15 @@ def _parse_header(self): # block_ann['avail_ccf'] = self._avail_files['ccf'] block_ann['rec_pauses'] = False - for c in range(unit_channels.size): - unit_ann = self.raw_annotations['unit_channels'][c] - channel_id, unit_id = self.internal_unit_ids[c] - unit_ann['channel_id'] = self.internal_unit_ids[c][0] - unit_ann['unit_id'] = self.internal_unit_ids[c][1] - unit_ann['unit_tag'] = {0: 'unclassified', 255: 'noise'}.get(unit_id, str(unit_id)) - unit_ann['description'] = 'Unit channel_id: {}, unit_id: {}, unit_tag: {}'.format( - channel_id, unit_id, unit_ann['unit_tag']) - + # this is not used anymore because not more ChannelIndex + """ flt_type = {0: 'None', 1: 'Butterworth'} - for c in range(sig_channels.size): + for c in range(signal_channels.size): chidx_ann = self.raw_annotations['signal_channels'][c] if self._avail_files['nev']: neuevwav = self.__nev_ext_header[b'NEUEVWAV'] - if sig_channels[c]['id'] in neuevwav['electrode_id']: - get_idx = list(neuevwav['electrode_id']).index(sig_channels[c]['id']) + if signal_channels[c]['id'] in neuevwav['electrode_id']: + get_idx = list(neuevwav['electrode_id']).index(signal_channels[c]['id']) chidx_ann['connector_ID'] = neuevwav['physical_connector'][get_idx] chidx_ann['connector_pinID'] = neuevwav['connector_pin'][get_idx] chidx_ann['nev_dig_factor'] = neuevwav['digitization_factor'][get_idx] @@ -514,12 +507,12 @@ def _parse_header(self): 'nev_dig_factor'] / 1000 * pq.uV chidx_ann['nb_sorted_units'] = neuevwav['nb_sorted_units'][get_idx] chidx_ann['waveform_size'] = self.__waveform_size[self.__nev_spec]( - )[sig_channels[c]['id']] * self.__nev_params('waveform_time_unit') + )[signal_channels[c]['id']] * self.__nev_params('waveform_time_unit') if self.__nev_spec in ['2.2', '2.3']: neuevflt = self.__nev_ext_header[b'NEUEVFLT'] get_idx = list( neuevflt['electrode_id']).index( - sig_channels[c]['id']) + signal_channels[c]['id']) # filter type codes (extracted from blackrock manual) chidx_ann['nev_hi_freq_corner'] = neuevflt['hi_freq_corner'][ get_idx] / 1000. * pq.Hz @@ -531,29 +524,7 @@ def _parse_header(self): chidx_ann['nev_lo_freq_order'] = neuevflt['lo_freq_order'][get_idx] chidx_ann['nev_lo_freq_type'] = flt_type[neuevflt['lo_freq_type'][ get_idx]] - if self.__nsx_spec[self.nsx_to_load[0]] in ['2.2', '2.3'] and self.__nsx_ext_header: - # It does not matter which nsX file to ask for this info - k = list(self.__nsx_ext_header.keys())[0] - if sig_channels[c]['id'] in self.__nsx_ext_header[k]['electrode_id']: - get_idx = list( - self.__nsx_ext_header[k]['electrode_id']).index( - sig_channels[c]['id']) - chidx_ann['connector_ID'] = self.__nsx_ext_header[k]['physical_connector'][ - get_idx] - chidx_ann['connector_pinID'] = self.__nsx_ext_header[k]['connector_pin'][ - get_idx] - chidx_ann['nsx_hi_freq_corner'] = \ - self.__nsx_ext_header[k]['hi_freq_corner'][get_idx] / 1000. * pq.Hz - chidx_ann['nsx_lo_freq_corner'] = \ - self.__nsx_ext_header[k]['lo_freq_corner'][get_idx] / 1000. * pq.Hz - chidx_ann['nsx_hi_freq_order'] = self.__nsx_ext_header[k][ - 'hi_freq_order'][get_idx] - chidx_ann['nsx_lo_freq_order'] = self.__nsx_ext_header[k][ - 'lo_freq_order'][get_idx] - chidx_ann['nsx_hi_freq_type'] = flt_type[ - self.__nsx_ext_header[k]['hi_freq_type'][get_idx]] - chidx_ann['nsx_lo_freq_type'] = flt_type[ - self.__nsx_ext_header[k]['hi_freq_type'][get_idx]] + """ for seg_index in range(self._nb_segment): seg_ann = block_ann['segments'][seg_index] @@ -565,25 +536,37 @@ def _parse_header(self): # so datetime is valide only for seg_index=0 seg_ann['rec_datetime'] = rec_datetime - for c in range(sig_channels.size): - nsx_nb = sig_channels['group_id'][c] - anasig_an = seg_ann['signals'][c] - desc = "AnalogSignal {} from channel_id: {}, label: {}, nsx: {}".format( - c, sig_channels['id'][c], sig_channels['name'][c], nsx_nb) - anasig_an['description'] = desc - anasig_an['file_origin'] = self._filenames['nsx'] + '.ns' + str(nsx_nb) - anasig_an['nsx'] = nsx_nb - chidx_ann = self.raw_annotations['signal_channels'][c] - chidx_ann['description'] = 'Container for Units and AnalogSignals of ' \ - 'one recording channel across segments.' - - for c in range(unit_channels.size): + for c in range(signal_streams.size): + sig_ann = seg_ann['signals'][c] + stream_id = signal_streams['id'][c] + nsx_nb = int(stream_id) + sig_ann['description'] = f'AnalogSignal from nsx{nsx_nb}' + sig_ann['file_origin'] = self._filenames['nsx'] + '.ns' + str(nsx_nb) + sig_ann['nsx'] = nsx_nb + # handle signal array annotations from nsx header + if self.__nsx_spec[nsx_nb] in ['2.2', '2.3'] and nsx_nb in self.__nsx_ext_header: + mask = signal_channels['stream_id'] == stream_id + channels = signal_channels[mask] + nsx_header = self.__nsx_ext_header[nsx_nb] + for key in ('physical_connector', 'connector_pin', 'hi_freq_corner', + 'lo_freq_corner', 'hi_freq_order', 'lo_freq_order', + 'hi_freq_type', 'lo_freq_type'): + values = [] + for chan_id in channels['id']: + chan_id = int(chan_id) + idx = list(nsx_header['electrode_id']).index(chan_id) + values.append(nsx_header[key][idx]) + values = np.array(values) + sig_ann['__array_annotations__'][key] = values + + for c in range(spike_channels.size): + st_ann = seg_ann['spikes'][c] channel_id, unit_id = self.internal_unit_ids[c] - st_ann = seg_ann['units'][c] - unit_ann = self.raw_annotations['unit_channels'][c] - st_ann.update(unit_ann) - st_ann['description'] = 'SpikeTrain channel_id: {}, unit_id: {}'.format( - channel_id, unit_id) + unit_tag = {0: 'unclassified', 255: 'noise'}.get(unit_id, str(unit_id)) + st_ann['channel_id'] = channel_id + st_ann['unit_id'] = unit_id + st_ann['unit_tag'] = unit_tag + st_ann['description'] = f'SpikeTrain channel_id: {channel_id}, unit_id: {unit_id}' st_ann['file_origin'] = self._filenames['nev'] + '.nev' if self._avail_files['nev']: @@ -610,33 +593,25 @@ def _segment_t_start(self, block_index, seg_index): def _segment_t_stop(self, block_index, seg_index): return self._seg_t_stops[seg_index] - def _get_nsx_and_local_indexes(self, channel_indexes): - # internal helper to get nsx number and local channel index - # from global channel indexes - # when this is called channell_indexes are alwas in the same group_id - # this is checked at BaseRaw level - if channel_indexes is None: - channel_indexes = slice(None) - nsx_nb = self.header['signal_channels'][channel_indexes]['group_id'][0] - if channel_indexes is None: - local_indexes = slice(None) - else: - local_indexes = self._local_sig_indexes[channel_indexes] - return nsx_nb, local_indexes - - def _get_signal_size(self, block_index, seg_index, channel_indexes): - nsx_nb, local_indexes = self._get_nsx_and_local_indexes(channel_indexes) + def _get_signal_size(self, block_index, seg_index, stream_index): + stream_id = self.header['signal_streams'][stream_index]['id'] + nsx_nb = int(stream_id) memmap_data = self.nsx_datas[nsx_nb][seg_index] return memmap_data.shape[0] - def _get_signal_t_start(self, block_index, seg_index, channel_indexes): - nsx_nb, local_indexes = self._get_nsx_and_local_indexes(channel_indexes) + def _get_signal_t_start(self, block_index, seg_index, stream_index): + stream_id = self.header['signal_streams'][stream_index]['id'] + nsx_nb = int(stream_id) return self._sigs_t_starts[nsx_nb][seg_index] - def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): - nsx_nb, local_indexes = self._get_nsx_and_local_indexes(channel_indexes) + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, + stream_index, channel_indexes): + stream_id = self.header['signal_streams'][stream_index]['id'] + nsx_nb = int(stream_id) memmap_data = self.nsx_datas[nsx_nb][seg_index] - sig_chunk = memmap_data[i_start:i_stop, local_indexes] + if channel_indexes is None: + channel_indexes = slice(None) + sig_chunk = memmap_data[i_start:i_stop, channel_indexes] return sig_chunk def _spike_count(self, block_index, seg_index, unit_index): @@ -766,7 +741,7 @@ def _get_event_timestamps(self, block_index, seg_index, event_channel_index, t_s return timestamp, durations, labels - def _rescale_event_timestamp(self, event_timestamps, dtype): + def _rescale_event_timestamp(self, event_timestamps, dtype, event_channel_index): ev_times = event_timestamps.astype(dtype) ev_times /= self.__nev_basic_header['timestamp_resolution'] return ev_times diff --git a/neo/rawio/brainvisionrawio.py b/neo/rawio/brainvisionrawio.py index 253a611d1..7cc95043d 100644 --- a/neo/rawio/brainvisionrawio.py +++ b/neo/rawio/brainvisionrawio.py @@ -7,8 +7,8 @@ Author: Samuel Garcia """ -from .baserawio import (BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, - _event_channel_dtype) +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype) import numpy as np @@ -57,6 +57,8 @@ def _parse_header(self): sigs = sigs[:-sigs.size % nb_channel] self._raw_signals = sigs.reshape(-1, nb_channel) + signal_streams = np.array([('Signals', '0')], dtype=_signal_stream_dtype) + sig_channels = [] channel_infos = vhdr_header['Channel Infos'] for c in range(nb_channel): @@ -66,20 +68,20 @@ def _parse_header(self): channel_desc = channel_infos['ch%d' % (c + 1,)] name, ref, res, units = channel_desc.split(',') units = units.replace('µ', 'u') - chan_id = c + 1 + chan_id = str(c + 1) if sig_dtype == np.int16 or sig_dtype == np.int32: gain = float(res) else: gain = 1 offset = 0 - group_id = 0 + stream_id = '0' sig_channels.append((name, chan_id, self._sampling_rate, sig_dtype, - units, gain, offset, group_id)) + units, gain, offset, stream_id)) sig_channels = np.array(sig_channels, dtype=_signal_channel_dtype) # No spikes - unit_channels = [] - unit_channels = np.array(unit_channels, dtype=_unit_channel_dtype) + spike_channels = [] + spike_channels = np.array(spike_channels, dtype=_spike_channel_dtype) # read all markers in memory @@ -112,18 +114,21 @@ def _parse_header(self): self.header = {} self.header['nb_block'] = 1 self.header['nb_segment'] = [1] + self.header['signal_streams'] = signal_streams self.header['signal_channels'] = sig_channels - self.header['unit_channels'] = unit_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels self._generate_minimal_annotations() if 'Coordinates' in vhdr_header: + sig_annotations = self.raw_annotations['blocks'][0]['segments'][0]['signals'][0] + all_coords = [] for c in range(sig_channels.size): coords = vhdr_header['Coordinates']['Ch{}'.format(c + 1)] - coords = [float(v) for v in coords.split(',')] - if coords[0] > 0.: - # if radius is 0 we do not have coordinates. - self.raw_annotations['signal_channels'][c]['coordinates'] = coords + all_coords.append([float(v) for v in coords.split(',')]) + all_coords = np.array(all_coords) + for dim in range(all_coords.shape[1]): + sig_annotations['__array_annotations__'][f'coordinates_{dim}'] = all_coords[:, dim] def _source_name(self): return self.filename @@ -136,13 +141,15 @@ def _segment_t_stop(self, block_index, seg_index): return t_stop ### - def _get_signal_size(self, block_index, seg_index, channel_indexes): + def _get_signal_size(self, block_index, seg_index, stream_index): + assert stream_index == 0 return self._raw_signals.shape[0] - def _get_signal_t_start(self, block_index, seg_index, channel_indexes): + def _get_signal_t_start(self, block_index, seg_index, stream_index): return 0. - def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, + stream_index, channel_indexes): if channel_indexes is None: channel_indexes = slice(None) raw_signals = self._raw_signals[slice(i_start, i_stop), channel_indexes] @@ -177,7 +184,7 @@ def _get_event_timestamps(self, block_index, seg_index, event_channel_index, t_s raise (NotImplementedError) - def _rescale_event_timestamp(self, event_timestamps, dtype): + def _rescale_event_timestamp(self, event_timestamps, dtype, event_channel_index): event_times = event_timestamps.astype(dtype) / self._sampling_rate return event_times diff --git a/neo/rawio/elanrawio.py b/neo/rawio/elanrawio.py index 451898672..447032168 100644 --- a/neo/rawio/elanrawio.py +++ b/neo/rawio/elanrawio.py @@ -16,8 +16,8 @@ """ -from .baserawio import (BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, - _event_channel_dtype) +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype) import numpy as np @@ -57,10 +57,10 @@ def _parse_header(self): # strange 2 line for datetime # line1 - l = f.readline() - r1 = re.findall(r'(\d+)-(\d+)-(\d+) (\d+):(\d+):(\d+)', l) - r2 = re.findall(r'(\d+):(\d+):(\d+)', l) - r3 = re.findall(r'(\d+)-(\d+)-(\d+)', l) + line = f.readline() + r1 = re.findall(r'(\d+)-(\d+)-(\d+) (\d+):(\d+):(\d+)', line) + r2 = re.findall(r'(\d+):(\d+):(\d+)', line) + r3 = re.findall(r'(\d+)-(\d+)-(\d+)', line) YY, MM, DD, hh, mm, ss = (None,) * 6 if len(r1) != 0: DD, MM, YY, hh, mm, ss = r1[0] @@ -70,10 +70,10 @@ def _parse_header(self): DD, MM, YY = r3[0] # line2 - l = f.readline() - r1 = re.findall(r'(\d+)-(\d+)-(\d+) (\d+):(\d+):(\d+)', l) - r2 = re.findall(r'(\d+):(\d+):(\d+)', l) - r3 = re.findall(r'(\d+)-(\d+)-(\d+)', l) + line = f.readline() + r1 = re.findall(r'(\d+)-(\d+)-(\d+) (\d+):(\d+):(\d+)', line) + r2 = re.findall(r'(\d+):(\d+):(\d+)', line) + r3 = re.findall(r'(\d+)-(\d+)-(\d+)', line) if len(r1) != 0: DD, MM, YY, hh, mm, ss = r1[0] elif len(r2) != 0: @@ -86,17 +86,17 @@ def _parse_header(self): except: fulldatetime = None - l = f.readline() - l = f.readline() - l = f.readline() + line = f.readline() + line = f.readline() + line = f.readline() # sampling rate sample - l = f.readline() - self._sampling_rate = 1. / float(l) + line = f.readline() + self._sampling_rate = 1. / float(line) # nb channel - l = f.readline() - nb_channel = int(l) - 2 + line = f.readline() + nb_channel = int(line) - 2 channel_infos = [{} for c in range(nb_channel + 2)] # channel label @@ -123,21 +123,23 @@ def _parse_header(self): for c in range(nb_channel + 2): channel_infos[c]['info_filter'] = f.readline()[:-1] - n = int(round(np.log(channel_infos[0]['max_logic'] - - channel_infos[0]['min_logic']) / np.log(2)) / 8) + n = int(round(np.log(channel_infos[0]['max_logic'] + - channel_infos[0]['min_logic']) / np.log(2)) / 8) sig_dtype = np.dtype('>i' + str(n)) + signal_streams = np.array([('Signals', '0')], dtype=_signal_stream_dtype) + sig_channels = [] for c, chan_info in enumerate(channel_infos[:-2]): chan_name = chan_info['label'] - chan_id = c + chan_id = str(c) gain = (chan_info['max_physic'] - chan_info['min_physic']) / \ (chan_info['max_logic'] - chan_info['min_logic']) offset = - chan_info['min_logic'] * gain + chan_info['min_physic'] - gourp_id = 0 + stream_id = '0' sig_channels.append((chan_name, chan_id, self._sampling_rate, sig_dtype, - chan_info['units'], gain, offset, gourp_id)) + chan_info['units'], gain, offset, stream_id)) sig_channels = np.array(sig_channels, dtype=_signal_channel_dtype) @@ -151,8 +153,8 @@ def _parse_header(self): self._raw_event_timestamps = [] self._event_labels = [] self._reject_codes = [] - for l in f.readlines(): - r = re.findall(r' *(\d+)\s* *(\d+)\s* *(\d+) *', l) + for line in f.readlines(): + r = re.findall(r' *(\d+)\s* *(\d+)\s* *(\d+) *', line) self._raw_event_timestamps.append(int(r[0][0])) self._event_labels.append(str(r[0][1])) self._reject_codes.append(str(r[0][2])) @@ -166,28 +168,34 @@ def _parse_header(self): event_channels = np.array(event_channels, dtype=_event_channel_dtype) # No spikes - unit_channels = [] - unit_channels = np.array(unit_channels, dtype=_unit_channel_dtype) + spike_channels = [] + spike_channels = np.array(spike_channels, dtype=_spike_channel_dtype) # fille into header dict self.header = {} self.header['nb_block'] = 1 self.header['nb_segment'] = [1] + self.header['signal_streams'] = signal_streams self.header['signal_channels'] = sig_channels - self.header['unit_channels'] = unit_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels # insert some annotation at some place self._generate_minimal_annotations() extra_info = dict(rec_datetime=fulldatetime, elan_version=version, info1=info1, info2=info2) - for obj_name in ('blocks', 'segments'): - self._raw_annotate(obj_name, **extra_info) - for c in range(nb_channel): - d = channel_infos[c] - self._raw_annotate('signals', chan_index=c, info_filter=d['info_filter']) - self._raw_annotate('signals', chan_index=c, kind=d['kind']) - self._raw_annotate('events', chan_index=0, reject_codes=self._reject_codes) + block_annotations = self.raw_annotations['blocks'][0] + block_annotations.update(extra_info) + seg_annotations = self.raw_annotations['blocks'][0]['segments'][0] + seg_annotations.update(extra_info) + + sig_annotations = self.raw_annotations['blocks'][0]['segments'][0]['signals'][0] + for key in ('info_filter', 'kind'): + values = [channel_infos[c][key] for c in range(nb_channel)] + sig_annotations['__array_annotations__'][key] = np.array(values) + + event_annotations = self.raw_annotations['blocks'][0]['segments'][0]['events'][0] + event_annotations['__array_annotations__']['reject_codes'] = self._reject_codes def _source_name(self): return self.filename @@ -199,13 +207,16 @@ def _segment_t_stop(self, block_index, seg_index): t_stop = self._raw_signals.shape[0] / self._sampling_rate return t_stop - def _get_signal_size(self, block_index, seg_index, channel_indexes=None): + def _get_signal_size(self, block_index, seg_index, stream_index): + assert stream_index == 0 return self._raw_signals.shape[0] - def _get_signal_t_start(self, block_index, seg_index, channel_indexes=None): + def _get_signal_t_start(self, block_index, seg_index, stream_index): + assert stream_index == 0 return 0. - def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, + stream_index, channel_indexes): if channel_indexes is None: channel_indexes = slice(None) raw_signals = self._raw_signals[slice(i_start, i_stop), channel_indexes] @@ -217,7 +228,8 @@ def _spike_count(self, block_index, seg_index, unit_index): def _event_count(self, block_index, seg_index, event_channel_index): return self._raw_event_timestamps.size - def _get_event_timestamps(self, block_index, seg_index, event_channel_index, t_start, t_stop): + def _get_event_timestamps(self, block_index, seg_index, + event_channel_index, t_start, t_stop): timestamp = self._raw_event_timestamps labels = self._event_labels durations = None @@ -234,6 +246,6 @@ def _get_event_timestamps(self, block_index, seg_index, event_channel_index, t_s return timestamp, durations, labels - def _rescale_event_timestamp(self, event_timestamps, dtype): + def _rescale_event_timestamp(self, event_timestamps, dtype, event_channel_index): event_times = event_timestamps.astype(dtype) / self._sampling_rate return event_times diff --git a/neo/rawio/examplerawio.py b/neo/rawio/examplerawio.py index b0516b5ca..d5dc0e071 100644 --- a/neo/rawio/examplerawio.py +++ b/neo/rawio/examplerawio.py @@ -6,18 +6,18 @@ Rules for creating a new class: 1. Step 1: Create the main class * Create a file in **neo/rawio/** that endith with "rawio.py" - * Create the class that inherits BaseRawIO + * Create the class that inherits from BaseRawIO * copy/paste all methods that need to be implemented. - See the end a neo.rawio.baserawio.BaseRawIO - * code hard! The main difficulty **is _parse_header()**. + * code hard! The main difficulty is `_parse_header()`. In short you have a create a mandatory dict than contains channel informations:: self.header = {} self.header['nb_block'] = 2 self.header['nb_segment'] = [2, 3] - self.header['signal_channels'] = sig_channels - self.header['unit_channels'] = unit_channels + self.header['signal_streams'] = signal_streams + self.header['signal_channels'] = signal_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels 2. Step 2: RawIO test: @@ -37,8 +37,8 @@ """ -from .baserawio import (BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, - _event_channel_dtype) +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype) import numpy as np @@ -60,8 +60,8 @@ class ExampleRawIO(BaseRawIO): This fake IO: * has 2 blocks * blocks have 2 and 3 segments - * has 16 signal_channels sample_rate = 10000 - * has 3 unit_channels + * has 2 signals streams of 8 channel each (sample_rate = 10000) so 16 channels in total + * has 3 spike_channels * has 2 event channels: one has *type=event*, the other has *type=epoch* @@ -75,7 +75,8 @@ class ExampleRawIO(BaseRawIO): i_start=0, i_stop=1024, channel_names=channel_names) >>> float_chunk = reader.rescale_signal_raw_to_float(raw_chunk, dtype='float64', channel_indexes=[0, 3, 6]) - >>> spike_timestamp = reader.spike_timestamps(unit_index=0, t_start=None, t_stop=None) + >>> spike_timestamp = reader.spike_timestamps(spike_channel_index=0, + t_start=None, t_stop=None) >>> spike_times = reader.rescale_spike_timestamp(spike_timestamp, 'float64') >>> ev_timestamps, _, ev_labels = reader.event_timestamps(event_channel_index=0) @@ -96,19 +97,27 @@ def _source_name(self): return self.filename def _parse_header(self): - # This is the central of a RawIO - # we need to collect in the original format all - # informations needed for further fast access + # This is the central part of a RawIO + # we need to collect from the original format all + # information required for fast access # at any place in the file - # In short _parse_header can be slow but - # _get_analogsignal_chunk need to be as fast as possible - - # create signals channels information + # In short `_parse_header()` can be slow but + # `_get_analogsignal_chunk()` need to be as fast as possible + + # create fake signals stream information + signal_streams = [] + for c in range(2): + name = f'stream {c}' + stream_id = c + signal_streams.append((name, stream_id)) + signal_streams = np.array(signal_streams, dtype=_signal_stream_dtype) + + # create fake signals channels information # This is mandatory!!!! # gain/offset/units are really important because # the scaling to real value will be done with that - # at the end real_signal = (raw_signal * gain + offset) * pq.Quantity(units) - sig_channels = [] + # The real signal will be evaluated as `(raw_signal * gain + offset) * pq.Quantity(units)` + signal_channels = [] for c in range(16): ch_name = 'ch{}'.format(c) # our channel id is c+1 just for fun @@ -121,20 +130,27 @@ def _parse_header(self): units = 'uV' gain = 1000. / 2 ** 16 offset = 0. - # group_id is only for special cases when channels have different - # sampling rate for instance. See TdtIO for that. - # Here this is the general case: all channel have the same characteritics - group_id = 0 - sig_channels.append((ch_name, chan_id, sr, dtype, units, gain, offset, group_id)) - sig_channels = np.array(sig_channels, dtype=_signal_channel_dtype) - - # creating units channels + # stream_id indicates how to group channels + # channels inside a "stream" share same characteristics + # (sampling rate/dtype/t_start/units/...) + stream_id = str(c // 8) + signal_channels.append((ch_name, chan_id, sr, dtype, units, gain, offset, stream_id)) + signal_channels = np.array(signal_channels, dtype=_signal_channel_dtype) + + # A stream can contain signals with different physical units. + # Here, the two last channels will have different units (pA) + # Since AnalogSignals must have consistent units across channels, + # this stream will be split in 2 parts on the neo.io level and finally 3 AnalogSignals + # will be generated per Segment. + signal_channels[-2:]['units'] = 'pA' + + # create fake units channels # This is mandatory!!!! # Note that if there is no waveform at all in the file # then wf_units/wf_gain/wf_offset/wf_left_sweep/wf_sampling_rate # can be set to any value because _spike_raw_waveforms # will return None - unit_channels = [] + spike_channels = [] for c in range(3): unit_name = 'unit{}'.format(c) unit_id = '#{}'.format(c) @@ -143,9 +159,9 @@ def _parse_header(self): wf_offset = 0. wf_left_sweep = 20 wf_sampling_rate = 10000. - unit_channels.append((unit_name, unit_id, wf_units, wf_gain, + spike_channels.append((unit_name, unit_id, wf_units, wf_gain, wf_offset, wf_left_sweep, wf_sampling_rate)) - unit_channels = np.array(unit_channels, dtype=_unit_channel_dtype) + spike_channels = np.array(spike_channels, dtype=_spike_channel_dtype) # creating event/epoch channel # This is mandatory!!!! @@ -160,16 +176,26 @@ def _parse_header(self): self.header = {} self.header['nb_block'] = 2 self.header['nb_segment'] = [2, 3] - self.header['signal_channels'] = sig_channels - self.header['unit_channels'] = unit_channels + self.header['signal_streams'] = signal_streams + self.header['signal_channels'] = signal_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels - # insert some annotation at some place - # at neo.io level IO are free to add some annoations + # insert some annotations/array_annotations at some place + # at neo.io level. IOs can add annotations # to any object. To keep this functionality with the wrapper - # BaseFromRaw you can add annoations in a nested dict. + # BaseFromRaw you can add annotations in a nested dict. + + # `_generate_minimal_annotations()` must be called to generate the nested + # dict of annotations/array_annotations self._generate_minimal_annotations() - # If you are a lazy dev you can stop here. + # this pprint lines really help for understand the nested (and complicated sometimes) dict + # from pprint import pprint + # pprint(self.raw_annotations) + + # Until here all mandatory operations for setting up a rawio are implemented. + # The following lines provide additional, recommended annotations for the + # final neo objects. for block_index in range(2): bl_ann = self.raw_annotations['blocks'][block_index] bl_ann['name'] = 'Block #{}'.format(block_index) @@ -180,16 +206,27 @@ def _parse_header(self): seg_index, block_index) seg_ann['seg_extra_info'] = 'This is the seg {} of block {}'.format( seg_index, block_index) - for c in range(16): - anasig_an = seg_ann['signals'][c] - anasig_an['info'] = 'This is a good signals' + for c in range(2): + sig_an = seg_ann['signals'][c]['nickname'] = \ + f'This stream {c} is from a subdevice' + # add some array annotations (8 channels) + sig_an = seg_ann['signals'][c]['__array_annotations__']['impedance'] = \ + np.random.rand(8) * 10000 for c in range(3): - spiketrain_an = seg_ann['units'][c] + spiketrain_an = seg_ann['spikes'][c] spiketrain_an['quality'] = 'Good!!' + # add some array annotations + num_spikes = self.spike_count(block_index, seg_index, c) + spiketrain_an['__array_annotations__']['amplitudes'] = \ + np.random.randn(num_spikes) + for c in range(2): event_an = seg_ann['events'][c] if c == 0: event_an['nickname'] = 'Miss Event 0' + # add some array annotations + num_ev = self.event_count(block_index, seg_index, c) + event_an['__array_annotations__']['button'] = ['A'] * num_ev elif c == 1: event_an['nickname'] = 'MrEpoch 1' @@ -205,16 +242,18 @@ def _segment_t_stop(self, block_index, seg_index): all_stops = [[10., 25.], [10., 30., 70.]] return all_stops[block_index][seg_index] - def _get_signal_size(self, block_index, seg_index, channel_indexes=None): - # we are lucky: signals in all segment have the same shape!! (10.0 seconds) - # it is not always the case + def _get_signal_size(self, block_index, seg_index, stream_index): + # We generate fake data in which the two stream signals have the same shape + # across all segments (10.0 seconds) + # This is not the case for real data, instead you should return the signal + # size depending on the block_index and segment_index # this must return an int = the number of sample # Note that channel_indexes can be ignored for most cases # except for several sampling rate. return 100000 - def _get_signal_t_start(self, block_index, seg_index, channel_indexes): + def _get_signal_t_start(self, block_index, seg_index, stream_index): # This give the t_start of signals. # Very often this equal to _segment_t_start but not # always. @@ -227,17 +266,19 @@ def _get_signal_t_start(self, block_index, seg_index, channel_indexes): # this is not always the case return self._segment_t_start(block_index, seg_index) - def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): - # this must return a signal chunk limited with - # i_start/i_stop (can be None) - # channel_indexes can be None (=all channel) or a list or numpy.array + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, + stream_index, channel_indexes): + # this must return a signal chunk in a signal stream + # limited with i_start/i_stop (can be None) + # channel_indexes can be None (=all channel in the stream) or a list or numpy.array # This must return a numpy array 2D (even with one channel). # This must return the orignal dtype. No conversion here. # This must as fast as possible. - # Everything that can be done in _parse_header() must not be here. + # To speed up this call all preparatory calculations should be implemented + # in _parse_header(). # Here we are lucky: our signals is always zeros!! - # it is not always the case + # it is not always the case :) # internally signals are int16 # convertion to real units is done with self.header['signal_channels'] @@ -246,25 +287,35 @@ def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, chann if i_stop is None: i_stop = 100000 - assert i_start >= 0, "I don't like your jokes" - assert i_stop <= 100000, "I don't like your jokes" + if i_start < 0 or i_stop > 100000: + # some check + raise IndexError("I don't like your jokes") if channel_indexes is None: - nb_chan = 16 + nb_chan = 8 + elif isinstance(channel_indexes, slice): + channel_indexes = np.arange(8, dtype='int')[channel_indexes] + nb_chan = len(channel_indexes) else: + channel_indexes = np.asarray(channel_indexes) + if any(channel_indexes < 0): + raise IndexError('bad boy') + if any(channel_indexes >= 8): + raise IndexError('big bad wolf') nb_chan = len(channel_indexes) + raw_signals = np.zeros((i_stop - i_start, nb_chan), dtype='int16') return raw_signals - def _spike_count(self, block_index, seg_index, unit_index): - # Must return the nb of spike for given (block_index, seg_index, unit_index) + def _spike_count(self, block_index, seg_index, spike_channel_index): + # Must return the nb of spikes for given (block_index, seg_index, spike_channel_index) # we are lucky: our units have all the same nb of spikes!! # it is not always the case nb_spikes = 20 return nb_spikes - def _get_spike_timestamps(self, block_index, seg_index, unit_index, t_start, t_stop): - # In our IO, timstamp are internally coded 'int64' and they + def _get_spike_timestamps(self, block_index, seg_index, spike_channel_index, t_start, t_stop): + # In our IO, timestamp are internally coded 'int64' and they # represent the index of the signals 10kHz # we are lucky: spikes have the same discharge in all segments!! # incredible neuron!! This is not always the case @@ -291,7 +342,8 @@ def _rescale_spike_timestamp(self, spike_timestamps, dtype): spike_times /= 10000. # because 10kHz return spike_times - def _get_spike_raw_waveforms(self, block_index, seg_index, unit_index, t_start, t_stop): + def _get_spike_raw_waveforms(self, block_index, seg_index, spike_channel_index, + t_start, t_stop): # this must return a 3D numpy array (nb_spike, nb_channel, nb_sample) # in the original dtype # this must be as fast as possible. @@ -302,13 +354,14 @@ def _get_spike_raw_waveforms(self, block_index, seg_index, unit_index, t_start, # In our IO waveforms come from all channels # they are int16 - # convertion to real units is done with self.header['unit_channels'] + # convertion to real units is done with self.header['spike_channels'] # Here, we have a realistic case: all waveforms are only noise. # it is not always the case # we 20 spikes with a sweep of 50 (5ms) # trick to get how many spike in the slice - ts = self._get_spike_timestamps(block_index, seg_index, unit_index, t_start, t_stop) + ts = self._get_spike_timestamps(block_index, seg_index, + spike_channel_index, t_start, t_stop) nb_spike = ts.size np.random.seed(2205) # a magic number (my birthday) @@ -357,7 +410,7 @@ def _get_event_timestamps(self, block_index, seg_index, event_channel_index, t_s return timestamp, durations, labels - def _rescale_event_timestamp(self, event_timestamps, dtype): + def _rescale_event_timestamp(self, event_timestamps, dtype, event_channel_index): # must rescale to second a particular event_timestamps # with a fixed dtype so the user can choose the precisino he want. @@ -365,7 +418,7 @@ def _rescale_event_timestamp(self, event_timestamps, dtype): event_times = event_timestamps.astype(dtype) return event_times - def _rescale_epoch_duration(self, raw_duration, dtype): + def _rescale_epoch_duration(self, raw_duration, dtype, event_channel_index): # really easy here because in our case it is already seconds durations = raw_duration.astype(dtype) return durations diff --git a/neo/rawio/intanrawio.py b/neo/rawio/intanrawio.py index eef771e07..e9a27f7b4 100644 --- a/neo/rawio/intanrawio.py +++ b/neo/rawio/intanrawio.py @@ -16,10 +16,9 @@ Author: Samuel Garcia """ -# from __future__ import unicode_literals is not compatible with numpy.dtype both py2 py3 -from .baserawio import (BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, - _event_channel_dtype) +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype, _common_sig_characteristics) import numpy as np from collections import OrderedDict @@ -58,22 +57,28 @@ def _parse_header(self): assert np.all(np.diff(timestamp) == 1), 'timestamp have gaps' # signals - sig_channels = [] + signal_channels = [] for c, chan_info in enumerate(self._ordered_channels): name = chan_info['native_channel_name'] - chan_id = c # the chan_id have no meaning in intan + chan_id = str(c) # the chan_id have no meaning in intan if chan_info['signal_type'] == 20: # exception for temperature sig_dtype = 'int16' else: sig_dtype = 'uint16' - group_id = 0 - sig_channels.append((name, chan_id, chan_info['sampling_rate'], + stream_id = str(chan_info['signal_type']) + signal_channels.append((name, chan_id, chan_info['sampling_rate'], sig_dtype, chan_info['units'], chan_info['gain'], - chan_info['offset'], chan_info['signal_type'])) - sig_channels = np.array(sig_channels, dtype=_signal_channel_dtype) + chan_info['offset'], stream_id)) + signal_channels = np.array(signal_channels, dtype=_signal_channel_dtype) - self._max_sampling_rate = np.max(sig_channels['sampling_rate']) + stream_ids = np.unique(signal_channels['stream_id']) + signal_streams = np.zeros(stream_ids.size, dtype=_signal_stream_dtype) + signal_streams['id'] = stream_ids + for stream_index, stream_id in enumerate(stream_ids): + signal_streams['name'][stream_index] = stream_type_to_name.get(int(stream_id), '') + + self._max_sampling_rate = np.max(signal_channels['sampling_rate']) self._max_sigs_length = self._raw_data.size * self._block_size # No events @@ -81,15 +86,16 @@ def _parse_header(self): event_channels = np.array(event_channels, dtype=_event_channel_dtype) # No spikes - unit_channels = [] - unit_channels = np.array(unit_channels, dtype=_unit_channel_dtype) + spike_channels = [] + spike_channels = np.array(spike_channels, dtype=_spike_channel_dtype) # fille into header dict self.header = {} self.header['nb_block'] = 1 self.header['nb_segment'] = [1] - self.header['signal_channels'] = sig_channels - self.header['unit_channels'] = unit_channels + self.header['signal_streams'] = signal_streams + self.header['signal_channels'] = signal_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels self._generate_minimal_annotations() @@ -101,27 +107,32 @@ def _segment_t_stop(self, block_index, seg_index): t_stop = self._max_sigs_length / self._max_sampling_rate return t_stop - def _get_signal_size(self, block_index, seg_index, channel_indexes): - assert channel_indexes is not None, 'channel_indexes cannot be None, several signal size' - assert np.unique(self.header['signal_channels'][channel_indexes]['group_id']).size == 1 - channel_names = self.header['signal_channels'][channel_indexes]['name'] - chan_name = channel_names[0] - size = self._raw_data[chan_name].size + def _get_signal_size(self, block_index, seg_index, stream_index): + stream_id = self.header['signal_streams'][stream_index]['id'] + mask = self.header['signal_channels']['stream_id'] == stream_id + signal_channels = self.header['signal_channels'][mask] + channel_names = signal_channels['name'] + chan_name0 = channel_names[0] + size = self._raw_data[chan_name0].size return size - def _get_signal_t_start(self, block_index, seg_index, channel_indexes): + def _get_signal_t_start(self, block_index, seg_index, stream_index): return 0. - def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, + stream_index, channel_indexes): if i_start is None: i_start = 0 if i_stop is None: - i_stop = self._get_signal_size(block_index, seg_index, channel_indexes) + i_stop = self._get_signal_size(block_index, seg_index, stream_index) + stream_id = self.header['signal_streams'][stream_index]['id'] + mask = self.header['signal_channels']['stream_id'] == stream_id + signal_channels = self.header['signal_channels'][mask] if channel_indexes is None: channel_indexes = slice(None) - channel_names = self.header['signal_channels'][channel_indexes]['name'] + channel_names = signal_channels['name'][channel_indexes] shape = self._raw_data[channel_names[0]].shape @@ -400,6 +411,15 @@ def read_rhs(filename): ('electrode_impedance_phase', 'float32'), ] +stream_type_to_name = { + 0: 'RHD2000 amplifier channel', + 1: 'RHD2000 auxiliary input channel', + 2: 'RHD2000 supply voltage channel', + 3: 'USB board ADC input channel', + 4: 'USB board digital input channel', + 5: 'USB board digital output channel', +} + def read_rhd(filename): with open(filename, mode='rb') as f: diff --git a/neo/rawio/mearecrawio.py b/neo/rawio/mearecrawio.py index 79bff2fd0..9ea2e10a4 100644 --- a/neo/rawio/mearecrawio.py +++ b/neo/rawio/mearecrawio.py @@ -9,8 +9,8 @@ Author : Alessio Buccino """ -from .baserawio import (BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, - _event_channel_dtype) +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype) import numpy as np from copy import deepcopy @@ -57,21 +57,24 @@ def _parse_header(self): self._sampling_rate = self._recgen.info['recordings']['fs'] self._recordings = self._recgen.recordings self._num_frames, self._num_channels = self._recordings.shape + + signal_streams = np.array([('Signals', '0')], dtype=_signal_stream_dtype) + sig_channels = [] for c in range(self._num_channels): ch_name = 'ch{}'.format(c) - chan_id = c + 1 + chan_id = str(c + 1) sr = self._sampling_rate # Hz dtype = self._recordings.dtype units = 'uV' gain = 1. offset = 0. - group_id = 0 - sig_channels.append((ch_name, chan_id, sr, dtype, units, gain, offset, group_id)) + stream_id = '0' + sig_channels.append((ch_name, chan_id, sr, dtype, units, gain, offset, stream_id)) sig_channels = np.array(sig_channels, dtype=_signal_channel_dtype) # creating units channels - unit_channels = [] + spike_channels = [] self._spiketrains = self._recgen.spiketrains for c in range(len(self._spiketrains)): unit_name = 'unit{}'.format(c) @@ -82,9 +85,9 @@ def _parse_header(self): wf_offset = 0. wf_left_sweep = 0 wf_sampling_rate = self._sampling_rate - unit_channels.append((unit_name, unit_id, wf_units, wf_gain, + spike_channels.append((unit_name, unit_id, wf_units, wf_gain, wf_offset, wf_left_sweep, wf_sampling_rate)) - unit_channels = np.array(unit_channels, dtype=_unit_channel_dtype) + spike_channels = np.array(spike_channels, dtype=_spike_channel_dtype) event_channels = [] event_channels = np.array(event_channels, dtype=_event_channel_dtype) @@ -92,8 +95,9 @@ def _parse_header(self): self.header = {} self.header['nb_block'] = 1 self.header['nb_segment'] = [1] + self.header['signal_streams'] = signal_streams self.header['signal_channels'] = sig_channels - self.header['unit_channels'] = unit_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels self._generate_minimal_annotations() @@ -110,13 +114,16 @@ def _segment_t_stop(self, block_index, seg_index): all_stops = [[t_stop]] return all_stops[block_index][seg_index] - def _get_signal_size(self, block_index, seg_index, channel_indexes=None): + def _get_signal_size(self, block_index, seg_index, stream_index): + assert stream_index == 0 return self._num_frames - def _get_signal_t_start(self, block_index, seg_index, channel_indexes): + def _get_signal_t_start(self, block_index, seg_index, stream_index): + assert stream_index == 0 return self._segment_t_start(block_index, seg_index) - def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, + stream_index, channel_indexes): if i_start is None: i_start = 0 if i_stop is None: @@ -144,17 +151,6 @@ def _get_spike_timestamps(self, block_index, seg_index, unit_index, t_start, t_s def _rescale_spike_timestamp(self, spike_timestamps, dtype): return spike_timestamps.astype(dtype) - def _get_spike_raw_waveforms(self, block_index, seg_index, unit_index, t_start, t_stop): - return None - - def _event_count(self, block_index, seg_index, event_channel_index): - return None - - def _get_event_timestamps(self, block_index, seg_index, event_channel_index, t_start, t_stop): - return None - - def _rescale_event_timestamp(self, event_timestamps, dtype): - return None - - def _rescale_epoch_duration(self, raw_duration, dtype): + def _get_spike_raw_waveforms(self, block_index, seg_index, + spike_channel_index, t_start, t_stop): return None diff --git a/neo/rawio/micromedrawio.py b/neo/rawio/micromedrawio.py index 4b2e6ea7b..6b7b4824f 100644 --- a/neo/rawio/micromedrawio.py +++ b/neo/rawio/micromedrawio.py @@ -9,8 +9,8 @@ # from __future__ import unicode_literals is not compatible with numpy.dtype both py2 py3 -from .baserawio import (BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, - _event_channel_dtype) +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype) import numpy as np @@ -83,7 +83,7 @@ def _parse_header(self): units_code = {-1: 'nV', 0: 'uV', 1: 'mV', 2: 1, 100: 'percent', 101: 'dimensionless', 102: 'dimensionless'} - sig_channels = [] + signal_channels = [] sig_grounds = [] for c in range(Num_Chan): zname2, pos, length = zones['LABCOD'] @@ -105,14 +105,17 @@ def _parse_header(self): f.seek(8, 1) sampling_rate, = f.read_f('H') sampling_rate *= Rate_Min - chan_id = c - group_id = 0 - sig_channels.append((chan_name, chan_id, sampling_rate, sig_dtype, - units, gain, offset, group_id)) + chan_id = str(c) + stream_id = '0' + signal_channels.append((chan_name, chan_id, sampling_rate, sig_dtype, + units, gain, offset, stream_id)) - sig_channels = np.array(sig_channels, dtype=_signal_channel_dtype) - assert np.unique(sig_channels['sampling_rate']).size == 1 - self._sampling_rate = float(np.unique(sig_channels['sampling_rate'])[0]) + signal_channels = np.array(signal_channels, dtype=_signal_channel_dtype) + + signal_streams = np.array([('Signals', '0')], dtype=_signal_stream_dtype) + + assert np.unique(signal_channels['sampling_rate']).size == 1 + self._sampling_rate = float(np.unique(signal_channels['sampling_rate'])[0]) # Event channels event_channels = [] @@ -142,15 +145,16 @@ def _parse_header(self): self._raw_events.append(rawevent) # No spikes - unit_channels = [] - unit_channels = np.array(unit_channels, dtype=_unit_channel_dtype) + spike_channels = [] + spike_channels = np.array(spike_channels, dtype=_spike_channel_dtype) # fille into header dict self.header = {} self.header['nb_block'] = 1 self.header['nb_segment'] = [1] - self.header['signal_channels'] = sig_channels - self.header['unit_channels'] = unit_channels + self.header['signal_streams'] = signal_streams + self.header['signal_channels'] = signal_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels # insert some annotation at some place @@ -164,11 +168,8 @@ def _parse_header(self): d['surname'] = surname d['header_version'] = header_version - for c in range(sig_channels.size): - anasig_an = seg_annotations['signals'][c] - anasig_an['ground'] = sig_grounds[c] - channel_an = self.raw_annotations['signal_channels'][c] - channel_an['ground'] = sig_grounds[c] + sig_annotations = self.raw_annotations['blocks'][0]['segments'][0]['signals'][0] + sig_annotations['__array_annotations__']['ground'] = np.array(sig_grounds) def _source_name(self): return self.filename @@ -180,13 +181,16 @@ def _segment_t_stop(self, block_index, seg_index): t_stop = self._raw_signals.shape[0] / self._sampling_rate return t_stop - def _get_signal_size(self, block_index, seg_index, channel_indexes): + def _get_signal_size(self, block_index, seg_index, stream_index): + assert stream_index == 0 return self._raw_signals.shape[0] - def _get_signal_t_start(self, block_index, seg_index, channel_indexes): + def _get_signal_t_start(self, block_index, seg_index, stream_index): + assert stream_index == 0 return 0. - def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, + stream_index, channel_indexes): if channel_indexes is None: channel_indexes = slice(channel_indexes) raw_signals = self._raw_signals[slice(i_start, i_stop), channel_indexes] @@ -225,10 +229,10 @@ def _get_event_timestamps(self, block_index, seg_index, event_channel_index, t_s return timestamp, durations, labels - def _rescale_event_timestamp(self, event_timestamps, dtype): + def _rescale_event_timestamp(self, event_timestamps, dtype, event_channel_index): event_times = event_timestamps.astype(dtype) / self._sampling_rate return event_times - def _rescale_epoch_duration(self, raw_duration, dtype): - durations = raw_duration.astype(dtype) // self._sampling_rate + def _rescale_epoch_duration(self, raw_duration, dtype, event_channel_index): + durations = raw_duration.astype(dtype) / self._sampling_rate return durations diff --git a/neo/rawio/neuralynxrawio/neuralynxrawio.py b/neo/rawio/neuralynxrawio/neuralynxrawio.py index cd4a047da..2d5eae19e 100644 --- a/neo/rawio/neuralynxrawio/neuralynxrawio.py +++ b/neo/rawio/neuralynxrawio/neuralynxrawio.py @@ -14,11 +14,10 @@ Author: Julia Sprenger, Carlos Canova, Samuel Garcia, Peter N. Steinmetz. """ -# from __future__ import unicode_literals is not compatible with numpy.dtype both py2 py3 -from neo.rawio.baserawio import (BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, - _event_channel_dtype) +from ..baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype) import numpy as np import os @@ -69,8 +68,9 @@ def _source_name(self): def _parse_header(self): - sig_channels = [] - unit_channels = [] + stream_channels = [] + signal_channels = [] + spike_channels = [] event_channels = [] self.ncs_filenames = OrderedDict() # (chan_name, chan_id): filename @@ -122,9 +122,9 @@ def _parse_header(self): if info.get('input_inverted', False): gain *= -1 offset = 0. - group_id = 0 - sig_channels.append((chan_name, chan_id, info['sampling_rate'], - 'int16', units, gain, offset, group_id)) + stream_id = 0 + signal_channels.append((chan_name, str(chan_id), info['sampling_rate'], + 'int16', units, gain, offset, stream_id)) self.ncs_filenames[chan_uid] = filename keys = [ 'DspFilterDelay_µs', @@ -183,7 +183,7 @@ def _parse_header(self): wf_offset = 0. wf_left_sweep = -1 # NOT KNOWN wf_sampling_rate = info['sampling_rate'] - unit_channels.append( + spike_channels.append( (unit_name, '{}'.format(unit_id), wf_units, wf_gain, wf_offset, wf_left_sweep, wf_sampling_rate)) unit_annotations.append(dict(file_origin=filename)) @@ -211,15 +211,19 @@ def _parse_header(self): self._nev_memmap[chan_id] = data - sig_channels = np.array(sig_channels, dtype=_signal_channel_dtype) - unit_channels = np.array(unit_channels, dtype=_unit_channel_dtype) + signal_channels = np.array(signal_channels, dtype=_signal_channel_dtype) + spike_channels = np.array(spike_channels, dtype=_spike_channel_dtype) event_channels = np.array(event_channels, dtype=_event_channel_dtype) # require all sampled signals, ncs files, to have same sampling rate - if sig_channels.size > 0: - sampling_rate = np.unique(sig_channels['sampling_rate']) + if signal_channels.size > 0: + sampling_rate = np.unique(signal_channels['sampling_rate']) assert sampling_rate.size == 1 self._sigs_sampling_rate = sampling_rate[0] + signal_streams = [('signals', '0')] + else: + signal_streams = [] + signal_streams = np.array(signal_streams, dtype=_signal_stream_dtype) # set 2 attributes needed later for header in case there are no ncs files in dataset, # e.g. Pegasus @@ -280,8 +284,9 @@ def _parse_header(self): self.header = {} self.header['nb_block'] = 1 self.header['nb_segment'] = [self._nb_segment] - self.header['signal_channels'] = sig_channels - self.header['unit_channels'] = unit_channels + self.header['signal_streams'] = signal_streams + self.header['signal_channels'] = signal_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels # Annotations @@ -291,12 +296,22 @@ def _parse_header(self): for seg_index in range(self._nb_segment): seg_annotations = bl_annotations['segments'][seg_index] - for c in range(sig_channels.size): + for c in range(signal_streams.size): + # one or no signal stream sig_ann = seg_annotations['signals'][c] - sig_ann.update(signal_annotations[c]) - - for c in range(unit_channels.size): - unit_ann = seg_annotations['units'][c] + # handle array annotations + for key in signal_annotations[0].keys(): + values = [] + for c in range(signal_channels.size): + value = signal_annotations[0][key] + values.append(value) + values = np.array(values) + if values.ndim == 1: + # 'InputRange': is 2D and make bugs + sig_ann['__array_annotations__'][key] = values + + for c in range(spike_channels.size): + unit_ann = seg_annotations['spikes'][c] unit_ann.update(unit_annotations[c]) for c in range(event_channels.size): @@ -319,13 +334,14 @@ def _segment_t_start(self, block_index, seg_index): def _segment_t_stop(self, block_index, seg_index): return self._seg_t_stops[seg_index] - self.global_t_start - def _get_signal_size(self, block_index, seg_index, channel_indexes): + def _get_signal_size(self, block_index, seg_index, stream_index): return self._sigs_length[seg_index] - def _get_signal_t_start(self, block_index, seg_index, channel_indexes): + def _get_signal_t_start(self, block_index, seg_index, stream_index): return self._sigs_t_start[seg_index] - self.global_t_start - def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, + stream_index, channel_indexes): """ Retrieve chunk of analog signal, a chunk being a set of contiguous samples. @@ -359,7 +375,7 @@ def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, chann if channel_indexes is None: channel_indexes = slice(None) - channel_ids = self.header['signal_channels'][channel_indexes]['id'] + channel_ids = self.header['signal_channels'][channel_indexes]['id'].astype(int) channel_names = self.header['signal_channels'][channel_indexes]['name'] # create buffer for samples @@ -461,7 +477,7 @@ def _get_event_timestamps(self, block_index, seg_index, event_channel_index, t_s durations = None return timestamps, durations, labels - def _rescale_event_timestamp(self, event_timestamps, dtype): + def _rescale_event_timestamp(self, event_timestamps, dtype, event_channel_index): event_times = event_timestamps.astype(dtype) event_times /= 1e6 event_times -= self.global_t_start @@ -503,8 +519,8 @@ def scan_ncs_files(self, ncs_filenames): nlxHeader = NlxHeader(ncs_filename) if not chanSectMap or (chanSectMap and - not NcsSectionsFactory._verifySectionsStructure(data, - lastNcsSections)): + not NcsSectionsFactory._verifySectionsStructure(data, + lastNcsSections)): lastNcsSections = NcsSectionsFactory.build_for_ncs_file(data, nlxHeader) chanSectMap[chan_uid] = [lastNcsSections, nlxHeader, data] diff --git a/neo/rawio/neuroexplorerrawio.py b/neo/rawio/neuroexplorerrawio.py index d7897bff2..11e21285f 100644 --- a/neo/rawio/neuroexplorerrawio.py +++ b/neo/rawio/neuroexplorerrawio.py @@ -22,10 +22,9 @@ Author: Samuel Garcia, luc estebanez, mark hollenbeck """ -# from __future__ import unicode_literals is not compatible with numpy.dtype both py2 py3 -from .baserawio import (BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, - _event_channel_dtype) +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype) import numpy as np from collections import OrderedDict @@ -57,14 +56,14 @@ def _parse_header(self): self._sig_lengths = [] self._sig_t_starts = [] sig_channels = [] - unit_channels = [] + spike_channels = [] event_channels = [] for i in range(self.global_header['nvar']): entity_header = self._entity_headers[i] name = entity_header['name'] - _id = i + _id = str(i) if entity_header['type'] == 0: # Unit - unit_channels.append((name, _id, '', 0, 0, 0, 0)) + spike_channels.append((name, _id, '', 0, 0, 0, 0)) elif entity_header['type'] == 1: # Event event_channels.append((name, _id, 'event')) @@ -78,7 +77,7 @@ def _parse_header(self): wf_offset = entity_header['MVOffset'] wf_left_sweep = 0 wf_sampling_rate = entity_header['WFrequency'] - unit_channels.append((name, _id, wf_units, wf_gain, wf_offset, + spike_channels.append((name, _id, wf_units, wf_gain, wf_offset, wf_left_sweep, wf_sampling_rate)) elif entity_header['type'] == 4: @@ -91,9 +90,9 @@ def _parse_header(self): dtype = 'int16' gain = entity_header['ADtoMV'] offset = entity_header['MVOffset'] - group_id = 0 + stream_id = str(_id) sig_channels.append((name, _id, sampling_rate, dtype, units, - gain, offset, group_id)) + gain, offset, stream_id)) self._sig_lengths.append(entity_header['NPointsWave']) # sig t_start is the first timestamp if datablock offset = entity_header['offset'] @@ -105,19 +104,23 @@ def _parse_header(self): event_channels.append((name, _id, 'event')) sig_channels = np.array(sig_channels, dtype=_signal_channel_dtype) - unit_channels = np.array(unit_channels, dtype=_unit_channel_dtype) + spike_channels = np.array(spike_channels, dtype=_spike_channel_dtype) event_channels = np.array(event_channels, dtype=_event_channel_dtype) # each signal channel have a dierent groups that force reading # them one by one - sig_channels['group_id'] = np.arange(sig_channels.size) + sig_channels['stream_id'] = np.arange(sig_channels.size).astype('U') + signal_streams = np.zeros(sig_channels.size, dtype=_signal_stream_dtype) + signal_streams['name'] = sig_channels['name'] + signal_streams['id'] = sig_channels['stream_id'] # fill into header dict self.header = {} self.header['nb_block'] = 1 self.header['nb_segment'] = [1] + self.header['signal_streams'] = signal_streams self.header['signal_channels'] = sig_channels - self.header['unit_channels'] = unit_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels # Annotations @@ -136,17 +139,15 @@ def _segment_t_stop(self, block_index, seg_index): t_stop = self.global_header['tend'] / self.global_header['freq'] return t_stop - def _get_signal_size(self, block_index, seg_index, channel_indexes): - assert len(channel_indexes) == 1, 'only one channel by one channel' - return self._sig_lengths[channel_indexes[0]] + def _get_signal_size(self, block_index, seg_index, stream_index): + return self._sig_lengths[stream_index] - def _get_signal_t_start(self, block_index, seg_index, channel_indexes): - assert len(channel_indexes) == 1, 'only one channel by one channel' - return self._sig_t_starts[channel_indexes[0]] + def _get_signal_t_start(self, block_index, seg_index, stream_index): + return self._sig_t_starts[stream_index] - def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): - assert len(channel_indexes) == 1, 'only one channel by one channel' - channel_index = channel_indexes[0] + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, + stream_index, channel_indexes): + channel_index = stream_index entity_index = int(self.header['signal_channels'][channel_index]['id']) entity_header = self._entity_headers[entity_index] n = entity_header['n'] @@ -161,13 +162,13 @@ def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, chann return raw_signal def _spike_count(self, block_index, seg_index, unit_index): - entity_index = int(self.header['unit_channels'][unit_index]['id']) + entity_index = int(self.header['spike_channels'][unit_index]['id']) entity_header = self._entity_headers[entity_index] nb_spike = entity_header['n'] return nb_spike def _get_spike_timestamps(self, block_index, seg_index, unit_index, t_start, t_stop): - entity_index = int(self.header['unit_channels'][unit_index]['id']) + entity_index = int(self.header['spike_channels'][unit_index]['id']) entity_header = self._entity_headers[entity_index] n = entity_header['n'] offset = entity_header['offset'] @@ -188,7 +189,7 @@ def _rescale_spike_timestamp(self, spike_timestamps, dtype): return spike_times def _get_spike_raw_waveforms(self, block_index, seg_index, unit_index, t_start, t_stop): - entity_index = int(self.header['unit_channels'][unit_index]['id']) + entity_index = int(self.header['spike_channels'][unit_index]['id']) entity_header = self._entity_headers[entity_index] if entity_header['type'] == 0: return None @@ -245,12 +246,12 @@ def _get_event_timestamps(self, block_index, seg_index, event_channel_index, t_s return timestamps, durations, labels - def _rescale_event_timestamp(self, event_timestamps, dtype): + def _rescale_event_timestamp(self, event_timestamps, dtype, event_channel_index): event_times = event_timestamps.astype(dtype) event_times /= self.global_header['freq'] return event_times - def _rescale_epoch_duration(self, raw_duration, dtype): + def _rescale_epoch_duration(self, raw_duration, dtype, event_channel_index): durations = raw_duration.astype(dtype) durations /= self.global_header['freq'] return durations diff --git a/neo/rawio/neuroscoperawio.py b/neo/rawio/neuroscoperawio.py index 38f243434..809ac19b2 100644 --- a/neo/rawio/neuroscoperawio.py +++ b/neo/rawio/neuroscoperawio.py @@ -16,8 +16,8 @@ """ -from .baserawio import (BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, - _event_channel_dtype) +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype) import numpy as np from xml.etree import ElementTree @@ -67,16 +67,19 @@ def _parse_header(self): self._raw_signals = np.memmap(filename + '.dat', dtype=sig_dtype, mode='r', offset=0).reshape(-1, nb_channel) + # one unique stream + signal_streams = np.array([('Signals', '0')], dtype=_signal_stream_dtype) + # signals sig_channels = [] for c in range(nb_channel): name = 'ch{}grp{}'.format(c, channel_group[c]) - chan_id = c + chan_id = str(c) units = 'mV' offset = 0. - group_id = 0 + stream_id = '0' sig_channels.append((name, chan_id, self._sampling_rate, - sig_dtype, units, gain, offset, group_id)) + sig_dtype, units, gain, offset, stream_id)) sig_channels = np.array(sig_channels, dtype=_signal_channel_dtype) # No events @@ -84,15 +87,16 @@ def _parse_header(self): event_channels = np.array(event_channels, dtype=_event_channel_dtype) # No spikes - unit_channels = [] - unit_channels = np.array(unit_channels, dtype=_unit_channel_dtype) + spike_channels = [] + spike_channels = np.array(spike_channels, dtype=_spike_channel_dtype) # fille into header dict self.header = {} self.header['nb_block'] = 1 self.header['nb_segment'] = [1] + self.header['signal_streams'] = signal_streams self.header['signal_channels'] = sig_channels - self.header['unit_channels'] = unit_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels self._generate_minimal_annotations() @@ -104,13 +108,16 @@ def _segment_t_stop(self, block_index, seg_index): t_stop = self._raw_signals.shape[0] / self._sampling_rate return t_stop - def _get_signal_size(self, block_index, seg_index, channel_indexes): + def _get_signal_size(self, block_index, seg_index, stream_index): + assert stream_index == 0 return self._raw_signals.shape[0] - def _get_signal_t_start(self, block_index, seg_index, channel_indexes): + def _get_signal_t_start(self, block_index, seg_index, stream_index): + assert stream_index == 0 return 0. - def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, + stream_index, channel_indexes): if channel_indexes is None: channel_indexes = slice(None) raw_signals = self._raw_signals[slice(i_start, i_stop), channel_indexes] diff --git a/neo/rawio/nixrawio.py b/neo/rawio/nixrawio.py index e120f9e94..f67d08af1 100644 --- a/neo/rawio/nixrawio.py +++ b/neo/rawio/nixrawio.py @@ -7,8 +7,9 @@ Author: Chek Yin Choi """ -from .baserawio import (BaseRawIO, _signal_channel_dtype, - _unit_channel_dtype, _event_channel_dtype) +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype) + from ..io.nixio import NixIO from ..io.nixio import check_nix_version import numpy as np @@ -48,8 +49,9 @@ def _source_name(self): def _parse_header(self): self.file = nix.File.open(self.filename, nix.FileMode.ReadOnly) - sig_channels = [] + signal_channels = [] size_list = [] + stream_ids = [] for bl in self.file.blocks: for seg in bl.groups: for da_idx, da in enumerate(seg.data_arrays): @@ -62,22 +64,24 @@ def _parse_header(self): da_leng = da.size if da_leng not in size_list: size_list.append(da_leng) - group_id = 0 - for sid, li_leng in enumerate(size_list): - if li_leng == da_leng: - group_id = sid - # very important! group_id use to store - # channel groups!!! - # use only for different signal length + stream_ids.append(str(len(size_list))) + # very important! group_id use to store + # channel groups!!! + # use only for different signal length + stream_index = size_list.index(da_leng) + stream_id = stream_ids[stream_index] gain = 1 offset = 0. - sig_channels.append((ch_name, chan_id, sr, dtype, - units, gain, offset, group_id)) + signal_channels.append((ch_name, chan_id, sr, dtype, + units, gain, offset, stream_id)) break break - sig_channels = np.array(sig_channels, dtype=_signal_channel_dtype) + signal_channels = np.array(signal_channels, dtype=_signal_channel_dtype) + signal_streams = np.zeros(len(stream_ids), dtype=_signal_stream_dtype) + signal_streams['id'] = stream_ids + signal_streams['name'] = '' - unit_channels = [] + spike_channels = [] unit_name = "" unit_id = "" for bl in self.file.blocks: @@ -99,13 +103,13 @@ def _parse_header(self): wf_left_sweep = wf.metadata["left_sweep"] wf_gain = 1 wf_offset = 0. - unit_channels.append( + spike_channels.append( (unit_name, unit_id, wf_units, wf_gain, wf_offset, wf_left_sweep, wf_sampling_rate) ) break break - unit_channels = np.array(unit_channels, dtype=_unit_channel_dtype) + spike_channels = np.array(spike_channels, dtype=_spike_channel_dtype) event_channels = [] event_count = 0 @@ -169,8 +173,7 @@ def _parse_header(self): segment['spiketrains'].append(st.positions) segment['spiketrains_id'].append(st.id) wftypestr = "neo.waveforms" - if (st.features - and st.features[0].data.type == wftypestr): + if (st.features and st.features[0].data.type == wftypestr): waveforms = st.features[0].data stdict = segment['spiketrains_unit'][st_idx] if waveforms: @@ -183,8 +186,9 @@ def _parse_header(self): self.header = {} self.header['nb_block'] = len(self.file.blocks) self.header['nb_segment'] = [len(bl.groups) for bl in self.file.blocks] - self.header['signal_channels'] = sig_channels - self.header['unit_channels'] = unit_channels + self.header['signal_streams'] = signal_streams + self.header['signal_channels'] = signal_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels self._generate_minimal_annotations() @@ -192,27 +196,31 @@ def _parse_header(self): bl_ann = self.raw_annotations['blocks'][blk_idx] props = blk.metadata.inherited_properties() bl_ann.update(self._filter_properties(props, "block")) - for grp_idx, grp in enumerate(blk.groups): + for grp_idx, group in enumerate(blk.groups): seg_ann = bl_ann['segments'][grp_idx] - props = grp.metadata.inherited_properties() + props = group.metadata.inherited_properties() seg_ann.update(self._filter_properties(props, "segment")) - sig_idx = 0 - groupdas = NixIO._group_signals(grp.data_arrays) - for nix_name, signals in groupdas.items(): - da = signals[0] - if da.type == 'neo.analogsignal' and seg_ann['signals']: - # collect and group DataArrays - sig_ann = seg_ann['signals'][sig_idx] - sig_chan_ann = self.raw_annotations['signal_channels'][sig_idx] - props = da.metadata.inherited_properties() - sig_ann.update(self._filter_properties(props, 'analogsignal')) - sig_chan_ann.update(self._filter_properties(props, 'analogsignal')) - sig_idx += 1 + + # TODO handle annotation at stream level + ''' + sig_idx = 0 + groupdas = NixIO._group_signals(grp.data_arrays) + for nix_name, signals in groupdas.items(): +   da = signals[0] +   if da.type == 'neo.analogsignal' and seg_ann['signals']: +   # collect and group DataArrays +   sig_ann = seg_ann['signals'][sig_idx] +   sig_chan_ann = self.raw_annotations['signal_channels'][sig_idx] +   props = da.metadata.inherited_properties() +   sig_ann.update(self._filter_properties(props, 'analogsignal')) +   sig_chan_ann.update(self._filter_properties(props, 'analogsignal')) +   sig_idx += 1 + ''' sp_idx = 0 ev_idx = 0 - for mt in grp.multi_tags: - if mt.type == 'neo.spiketrain' and seg_ann['units']: - st_ann = seg_ann['units'][sp_idx] + for mt in group.multi_tags: + if mt.type == 'neo.spiketrain' and seg_ann['spikes']: + st_ann = seg_ann['spikes'][sp_idx] props = mt.metadata.inherited_properties() st_ann.update(self._filter_properties(props, 'spiketrain')) sp_idx += 1 @@ -225,12 +233,6 @@ def _parse_header(self): event_ann.update(self._filter_properties(props, 'event')) ev_idx += 1 - # populate ChannelIndex annotations - for srcidx, source in enumerate(blk.sources): - chx_ann = self.raw_annotations["signal_channels"][srcidx] - props = source.metadata.inherited_properties() - chx_ann.update(self._filter_properties(props, "channelindex")) - def _segment_t_start(self, block_index, seg_index): t_start = 0 for mt in self.file.blocks[block_index].groups[seg_index].multi_tags: @@ -245,18 +247,20 @@ def _segment_t_stop(self, block_index, seg_index): t_stop = mt.metadata['t_stop'] return t_stop - def _get_signal_size(self, block_index, seg_index, channel_indexes): - if channel_indexes is None: - channel_indexes = list(range(self.header['signal_channels'].size)) + def _get_signal_size(self, block_index, seg_index, stream_index): + stream_id = self.header['signal_streams'][stream_index]['id'] + keep = self.header['signal_channels']['stream_id'] == stream_id + channel_indexes, = np.nonzero(keep) ch_idx = channel_indexes[0] block = self.da_list['blocks'][block_index] segment = block['segments'][seg_index] size = segment['data_size'][ch_idx] return size # size is per signal, not the sum of all channel_indexes - def _get_signal_t_start(self, block_index, seg_index, channel_indexes): - if channel_indexes is None: - channel_indexes = list(range(self.header['signal_channels'].size)) + def _get_signal_t_start(self, block_index, seg_index, stream_index): + stream_id = self.header['signal_streams'][stream_index]['id'] + keep = self.header['signal_channels']['stream_id'] == stream_id + channel_indexes, = np.nonzero(keep) ch_idx = channel_indexes[0] block = self.file.blocks[block_index] das = [da for da in block.groups[seg_index].data_arrays] @@ -264,22 +268,22 @@ def _get_signal_t_start(self, block_index, seg_index, channel_indexes): sig_t_start = float(da.metadata['t_start']) return sig_t_start # assume same group_id always same t_start - def _get_analogsignal_chunk(self, block_index, seg_index, - i_start, i_stop, channel_indexes): - if channel_indexes is None: - channel_indexes = list(range(self.header['signal_channels'].size)) + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, + stream_index, channel_indexes): + stream_id = self.header['signal_streams'][stream_index]['id'] + keep = self.header['signal_channels']['stream_id'] == stream_id + global_channel_indexes, = np.nonzero(keep) + if channel_indexes is not None: + global_channel_indexes = global_channel_indexes[channel_indexes] + if i_start is None: i_start = 0 if i_stop is None: - block = self.da_list['blocks'][block_index] - segment = block['segments'][seg_index] - for c in channel_indexes: - i_stop = segment['data_size'][c] - break + i_stop = self.get_signal_size(block_index, seg_index, stream_index) raw_signals_list = [] da_list = self.da_list['blocks'][block_index]['segments'][seg_index] - for idx in channel_indexes: + for idx in global_channel_indexes: da = da_list['data'][idx] raw_signals_list.append(da[i_start:i_stop]) @@ -289,7 +293,7 @@ def _get_analogsignal_chunk(self, block_index, seg_index, def _spike_count(self, block_index, seg_index, unit_index): count = 0 - head_id = self.header['unit_channels'][unit_index][1] + head_id = self.header['spike_channels'][unit_index][1] for mt in self.file.blocks[block_index].groups[seg_index].multi_tags: for src in mt.sources: if mt.type == 'neo.spiketrain' and [src.type == "neo.unit"]: @@ -358,8 +362,7 @@ def _get_event_timestamps(self, block_index, seg_index, if mt.type == "neo.event" or mt.type == "neo.epoch": labels.append(mt.positions.dimensions[0].labels) po = mt.positions - if (po.type == "neo.event.times" - or po.type == "neo.epoch.times"): + if (po.type == "neo.event.times" or po.type == "neo.epoch.times"): timestamp.append(po) channel = self.header['event_channels'][event_channel_index] if channel['type'] == b'epoch' and mt.extents: @@ -379,7 +382,7 @@ def _get_event_timestamps(self, block_index, seg_index, timestamp, labels = timestamp[keep], labels[keep] return timestamp, durations, labels # only the first fits in rescale - def _rescale_event_timestamp(self, event_timestamps, dtype='float64'): + def _rescale_event_timestamp(self, event_timestamps, dtype, event_channel_index): ev_unit = '' for mt in self.file.blocks[0].groups[0].multi_tags: if mt.type == "neo.event": @@ -391,7 +394,7 @@ def _rescale_event_timestamp(self, event_timestamps, dtype='float64'): # supposing unit is second, other possibilities maybe mS microS... return event_times # return in seconds - def _rescale_epoch_duration(self, raw_duration, dtype='float64'): + def _rescale_epoch_duration(self, raw_duration, dtype, event_channel_index): ep_unit = '' for mt in self.file.blocks[0].groups[0].multi_tags: if mt.type == "neo.epoch": diff --git a/neo/rawio/openephysrawio.py b/neo/rawio/openephysrawio.py index 1981bc0ef..000409a94 100644 --- a/neo/rawio/openephysrawio.py +++ b/neo/rawio/openephysrawio.py @@ -9,8 +9,8 @@ import numpy as np -from .baserawio import (BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, - _event_channel_dtype) +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype) RECORD_SIZE = 1024 @@ -78,7 +78,7 @@ def _parse_header(self): self._sigs_memmap = {} self._sig_length = {} self._sig_timestamp0 = {} - sig_channels = [] + signal_channels = [] oe_indices = sorted(list(info['continuous'].keys())) for seg_index, oe_index in enumerate(oe_indices): self._sigs_memmap[seg_index] = {} @@ -116,8 +116,8 @@ def _parse_header(self): if seg_index == 0: # add in channel list - sig_channels.append((ch_name, chan_id, chan_info['sampleRate'], - 'int16', 'V', chan_info['bitVolts'], 0., int(processor_id))) + signal_channels.append((ch_name, chan_id, chan_info['sampleRate'], + 'int16', 'V', chan_info['bitVolts'], 0., processor_id)) # In some cases, continuous do not have the same lentgh because # one record block is missing when the "OE GUI is freezing" @@ -148,22 +148,33 @@ def _parse_header(self): all_first_timestamps.append(data_chan[0]['timestamp']) all_last_timestamps.append(data_chan[-1]['timestamp']) - # chech that all signals have the same lentgh and timestamp0 for this segment + # check that all signals have the same lentgh and timestamp0 for this segment assert all(all_sigs_length[0] == e for e in all_sigs_length),\ - 'All signals do not have the same lentgh' + 'Not all signals have the same length' assert all(all_first_timestamps[0] == e for e in all_first_timestamps),\ - 'All signals do not have the same first timestamp' + 'Not all signals have the same first timestamp' assert all(all_samplerate[0] == e for e in all_samplerate),\ - 'All signals do not have the same sample rate' + 'Not all signals have the same sample rate' self._sig_length[seg_index] = all_sigs_length[0] self._sig_timestamp0[seg_index] = all_first_timestamps[0] - sig_channels = np.array(sig_channels, dtype=_signal_channel_dtype) - self._sig_sampling_rate = sig_channels['sampling_rate'][0] # unique for channel + signal_channels = np.array(signal_channels, dtype=_signal_channel_dtype) + self._sig_sampling_rate = signal_channels['sampling_rate'][0] # unique for channel + + # split channels in stream depending the name CHxxx ADCxxx + chan_stream_ids = [name[:2] if name.startswith('CH') else name[:3] + for name in signal_channels['name']] + signal_channels['stream_id'] = chan_stream_ids + + # and create streams channels (keep natural order 'CH' first) + stream_ids, order = np.unique(chan_stream_ids, return_index=True) + stream_ids = stream_ids[order] + signal_streams = [(f'Signals {stream_id}', f'{stream_id}') for stream_id in stream_ids] + signal_streams = np.array(signal_streams, dtype=_signal_stream_dtype) # scan for spikes files - unit_channels = [] + spike_channels = [] if len(info['spikes']) > 0: @@ -216,10 +227,10 @@ def _parse_header(self): for sorted_id in all_sorted_ids: unit_name = "{}#{}".format(name, sorted_id) unit_id = "{}#{}".format(name, sorted_id) - unit_channels.append((unit_name, unit_id, wf_units, + spike_channels.append((unit_name, unit_id, wf_units, wf_gain, wf_offset, wf_left_sweep, wf_sampling_rate)) - unit_channels = np.array(unit_channels, dtype=_unit_channel_dtype) + spike_channels = np.array(spike_channels, dtype=_spike_channel_dtype) # event file are: # * all_channel.events (header + binray) --> event 0 @@ -236,7 +247,7 @@ def _parse_header(self): event_info = read_file_header(fullname) self._event_sampling_rate = event_info['sampleRate'] data_event = np.memmap(fullname, mode='r', offset=HEADER_SIZE, - dtype=events_dtype) + dtype=events_dtype) self._events_memmap[seg_index] = data_event event_channels.append(('all_channels', '', 'event')) @@ -247,8 +258,9 @@ def _parse_header(self): self.header = {} self.header['nb_block'] = 1 self.header['nb_segment'] = [nb_segment] - self.header['signal_channels'] = sig_channels - self.header['unit_channels'] = unit_channels + self.header['signal_streams'] = signal_streams + self.header['signal_channels'] = signal_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels # Annotate some objects from coninuous files @@ -272,13 +284,14 @@ def _segment_t_stop(self, block_index, seg_index): return (self._sig_timestamp0[seg_index] + self._sig_length[seg_index])\ / self._sig_sampling_rate - def _get_signal_size(self, block_index, seg_index, channel_indexes=None): + def _get_signal_size(self, block_index, seg_index, stream_index): return self._sig_length[seg_index] - def _get_signal_t_start(self, block_index, seg_index, channel_indexes): + def _get_signal_t_start(self, block_index, seg_index, stream_index): return self._sig_timestamp0[seg_index] / self._sig_sampling_rate - def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, + stream_index, channel_indexes): if i_start is None: i_start = 0 if i_stop is None: @@ -289,20 +302,23 @@ def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, chann sl0 = i_start % RECORD_SIZE sl1 = sl0 + (i_stop - i_start) + stream_id = self.header['signal_streams'][stream_index]['id'] + mask = self.header['signal_channels']['stream_id'] + global_channel_indexes, = np.nonzero(mask == stream_id) if channel_indexes is None: channel_indexes = slice(None) - channel_indexes = np.arange(self.header['signal_channels'].size)[channel_indexes] + global_channel_indexes = global_channel_indexes[channel_indexes] - sigs_chunk = np.zeros((i_stop - i_start, len(channel_indexes)), dtype='int16') - for i, chan_index in enumerate(channel_indexes): - data = self._sigs_memmap[seg_index][chan_index] + sigs_chunk = np.zeros((i_stop - i_start, len(global_channel_indexes)), dtype='int16') + for i, global_chan_index in enumerate(global_channel_indexes): + data = self._sigs_memmap[seg_index][global_chan_index] sub = data[block_start:block_stop] sigs_chunk[:, i] = sub['samples'].flatten()[sl0:sl1] return sigs_chunk def _get_spike_slice(self, seg_index, unit_index, t_start, t_stop): - name, sorted_id = self.header['unit_channels'][unit_index]['name'].split('#') + name, sorted_id = self.header['spike_channels'][unit_index]['name'].split('#') sorted_id = int(sorted_id) data_spike = self._spikes_memmap[seg_index][name] @@ -366,11 +382,11 @@ def _get_event_timestamps(self, block_index, seg_index, event_channel_index, t_s return timestamps, durations, labels - def _rescale_event_timestamp(self, event_timestamps, dtype): + def _rescale_event_timestamp(self, event_timestamps, dtype, event_channel_index): event_times = event_timestamps.astype(dtype) / self._event_sampling_rate return event_times - def _rescale_epoch_duration(self, raw_duration, dtype): + def _rescale_epoch_duration(self, raw_duration, dtype, event_channel_index): return None diff --git a/neo/rawio/phyrawio.py b/neo/rawio/phyrawio.py index 23919d6d7..1bfa8cf2f 100644 --- a/neo/rawio/phyrawio.py +++ b/neo/rawio/phyrawio.py @@ -8,8 +8,8 @@ Author: Regimantas Jurkus """ -from .baserawio import (BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, - _event_channel_dtype) +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype) import numpy as np from pathlib import Path @@ -28,7 +28,7 @@ class PhyRawIO(BaseRawIO): >>> r.parse_header() >>> print(r) >>> spike_timestamp = r.get_spike_timestamps(block_index=0, - ... seg_index=0, unit_index=0, t_start=None, t_stop=None) + ... seg_index=0, spike_channel_index=0, t_start=None, t_stop=None) >>> spike_times = r.rescale_spike_timestamp(spike_timestamp, 'float64') """ @@ -84,10 +84,13 @@ def _parse_header(self): self._t_start = 0. self._t_stop = max(self._spike_times).item() / self._sampling_frequency - sig_channels = [] - sig_channels = np.array(sig_channels, dtype=_signal_channel_dtype) + signal_streams = [] + signal_streams = np.array(signal_streams, dtype=_signal_stream_dtype) - unit_channels = [] + signal_channels = [] + signal_channels = np.array(signal_channels, dtype=_signal_channel_dtype) + + spike_channels = [] for i, clust_id in enumerate(clust_ids): unit_name = f'unit {clust_id}' unit_id = f'{clust_id}' @@ -96,9 +99,9 @@ def _parse_header(self): wf_offset = 0. wf_left_sweep = 0 wf_sampling_rate = 0 - unit_channels.append((unit_name, unit_id, wf_units, wf_gain, + spike_channels.append((unit_name, unit_id, wf_units, wf_gain, wf_offset, wf_left_sweep, wf_sampling_rate)) - unit_channels = np.array(unit_channels, dtype=_unit_channel_dtype) + spike_channels = np.array(spike_channels, dtype=_spike_channel_dtype) event_channels = [] event_channels = np.array(event_channels, dtype=_event_channel_dtype) @@ -106,8 +109,9 @@ def _parse_header(self): self.header = {} self.header['nb_block'] = 1 self.header['nb_segment'] = [1] - self.header['signal_channels'] = sig_channels - self.header['unit_channels'] = unit_channels + self.header['signal_streams'] = signal_streams + self.header['signal_channels'] = signal_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels self._generate_minimal_annotations() @@ -126,7 +130,7 @@ def _parse_header(self): seg_ann = bl_ann['segments'][0] seg_ann['name'] = 'Seg #0 Block #0' for index, clust_id in enumerate(clust_ids): - spiketrain_an = seg_ann['units'][index] + spiketrain_an = seg_ann['spikes'][index] # Loop over list of list of dict and annotate each st for annotation_list in annotation_lists: @@ -160,27 +164,26 @@ def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): return None - def _spike_count(self, block_index, seg_index, unit_index): + def _spike_count(self, block_index, seg_index, spike_channel_index): assert block_index == 0 spikes = self._spike_clusters - unit_label = self.unit_labels[unit_index] + unit_label = self.unit_labels[spike_channel_index] mask = spikes == unit_label nb_spikes = np.sum(mask) return nb_spikes - def _get_spike_timestamps(self, block_index, seg_index, unit_index, + def _get_spike_timestamps(self, block_index, seg_index, spike_channel_index, t_start, t_stop): assert block_index == 0 assert seg_index == 0 - unit_label = self.unit_labels[unit_index] + unit_label = self.unit_labels[spike_channel_index] mask = self._spike_clusters == unit_label spike_timestamps = self._spike_times[mask] if t_start is not None: start_frame = int(t_start * self._sampling_frequency) - spike_timestamps = spike_timestamps[spike_timestamps >= - start_frame] + spike_timestamps = spike_timestamps[spike_timestamps >= start_frame] if t_stop is not None: end_frame = int(t_stop * self._sampling_frequency) spike_timestamps = spike_timestamps[spike_timestamps < end_frame] @@ -192,7 +195,7 @@ def _rescale_spike_timestamp(self, spike_timestamps, dtype): spike_times /= self._sampling_frequency return spike_times - def _get_spike_raw_waveforms(self, block_index, seg_index, unit_index, + def _get_spike_raw_waveforms(self, block_index, seg_index, spike_channel_index, t_start, t_stop): return None diff --git a/neo/rawio/plexonrawio.py b/neo/rawio/plexonrawio.py index eff9c7ff1..233b63ab9 100644 --- a/neo/rawio/plexonrawio.py +++ b/neo/rawio/plexonrawio.py @@ -21,10 +21,9 @@ Author: Samuel Garcia """ -# from __future__ import unicode_literals is not compatible with numpy.dtype both py2 py3 -from .baserawio import (BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, - _event_channel_dtype) +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype) import numpy as np from collections import OrderedDict @@ -165,13 +164,18 @@ def _parse_header(self): .5 * (2 ** global_header['BitsPerSpikeSample']) * h['Gain'] * h['PreampGain']) offset = 0. - group_id = 0 - sig_channels.append((name, chan_id, sampling_rate, sig_dtype, - units, gain, offset, group_id)) + stream_id = '0' + sig_channels.append((name, str(chan_id), sampling_rate, sig_dtype, + units, gain, offset, stream_id)) if len(all_sig_length) > 0: self._signal_length = min(all_sig_length) sig_channels = np.array(sig_channels, dtype=_signal_channel_dtype) + if sig_channels.size > 0: + signal_streams = np.array([('Signals', '0')], dtype=_signal_stream_dtype) + else: + signal_streams = np.array([], dtype=_signal_stream_dtype) + self._global_ssampling_rate = global_header['ADFrequency'] if slowChannelHeaders.size > 0: assert np.unique(slowChannelHeaders['ADFreq'] @@ -186,7 +190,7 @@ def _parse_header(self): self.internal_unit_ids.append((chan_id, unit_id)) # Spikes channels - unit_channels = [] + spike_channels = [] for unit_index, (chan_id, unit_id) in enumerate(self.internal_unit_ids): c = np.nonzero(dspChannelHeaders['Channel'] == chan_id)[0][0] h = dspChannelHeaders[c] @@ -207,9 +211,9 @@ def _parse_header(self): wf_offset = 0. wf_left_sweep = -1 # DONT KNOWN wf_sampling_rate = global_header['WaveformFreq'] - unit_channels.append((name, _id, wf_units, wf_gain, wf_offset, + spike_channels.append((name, _id, wf_units, wf_gain, wf_offset, wf_left_sweep, wf_sampling_rate)) - unit_channels = np.array(unit_channels, dtype=_unit_channel_dtype) + spike_channels = np.array(spike_channels, dtype=_spike_channel_dtype) # Event channels event_channels = [] @@ -225,8 +229,9 @@ def _parse_header(self): self.header = {} self.header['nb_block'] = 1 self.header['nb_segment'] = [1] + self.header['signal_streams'] = signal_streams self.header['signal_channels'] = sig_channels - self.header['unit_channels'] = unit_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels # Annotations @@ -248,13 +253,16 @@ def _segment_t_stop(self, block_index, seg_index): else: return t_stop1 - def _get_signal_size(self, block_index, seg_index, channel_indexes): + def _get_signal_size(self, block_index, seg_index, stream_index): + assert stream_index == 0 return self._signal_length - def _get_signal_t_start(self, block_index, seg_index, channel_indexes): + def _get_signal_t_start(self, block_index, seg_index, stream_index): + assert stream_index == 0 return 0. - def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, + stream_index, channel_indexes): if i_start is None: i_start = 0 if i_stop is None: @@ -262,11 +270,13 @@ def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, chann if channel_indexes is None: channel_indexes = np.arange(self.header['signal_channels'].size) + elif isinstance(channel_indexes, slice): + channel_indexes = np.arange(self.header['signal_channels'].size)[channel_indexes] raw_signals = np.zeros((i_stop - i_start, len(channel_indexes)), dtype='int16') for c, channel_index in enumerate(channel_indexes): - chan_header = self.header['signal_channels'][channel_index] - chan_id = chan_header['id'] + chan_id = self.header['signal_channels'][channel_index]['id'] + chan_id = np.int32(chan_id) data_blocks = self._data_blocks[5][chan_id] @@ -369,7 +379,7 @@ def _get_event_timestamps(self, block_index, seg_index, event_channel_index, t_s return timestamps, durations, labels - def _rescale_event_timestamp(self, event_timestamps, dtype): + def _rescale_event_timestamp(self, event_timestamps, dtype, event_channel_index): event_times = event_timestamps.astype(dtype) event_times /= self._global_ssampling_rate return event_times diff --git a/neo/rawio/rawbinarysignalrawio.py b/neo/rawio/rawbinarysignalrawio.py index e8f618263..6f4029471 100644 --- a/neo/rawio/rawbinarysignalrawio.py +++ b/neo/rawio/rawbinarysignalrawio.py @@ -17,8 +17,8 @@ class RawBinarySignalIO Author: Samuel Garcia """ -from .baserawio import (BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, - _event_channel_dtype) +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype) import numpy as np @@ -53,32 +53,39 @@ def _parse_header(self): # The the neo.io.RawBinarySignalIO is used for write_segment self._raw_signals = None - sig_channels = [] + signal_channels = [] if self._raw_signals is not None: for c in range(self.nb_channel): - name = 'ch{}'.format(c) - chan_id = c + name = f'ch{c}' + chan_id = f'{c}' units = '' - group_id = 0 - sig_channels.append((name, chan_id, self.sampling_rate, self.dtype, - units, self.signal_gain, self.signal_offset, group_id)) + stream_id = '0' + signal_channels.append((name, chan_id, self.sampling_rate, self.dtype, + units, self.signal_gain, self.signal_offset, stream_id)) - sig_channels = np.array(sig_channels, dtype=_signal_channel_dtype) + signal_channels = np.array(signal_channels, dtype=_signal_channel_dtype) + + # one unique stream + if signal_channels.size > 0: + signal_streams = np.array([('Signals', '0')], dtype=_signal_stream_dtype) + else: + signal_streams = np.array([], dtype=_signal_stream_dtype) # No events event_channels = [] event_channels = np.array(event_channels, dtype=_event_channel_dtype) # No spikes - unit_channels = [] - unit_channels = np.array(unit_channels, dtype=_unit_channel_dtype) + spike_channels = [] + spike_channels = np.array(spike_channels, dtype=_spike_channel_dtype) # fille into header dict self.header = {} self.header['nb_block'] = 1 self.header['nb_segment'] = [1] - self.header['signal_channels'] = sig_channels - self.header['unit_channels'] = unit_channels + self.header['signal_streams'] = signal_streams + self.header['signal_channels'] = signal_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels # insert some annotation at some place @@ -91,15 +98,17 @@ def _segment_t_stop(self, block_index, seg_index): t_stop = self._raw_signals.shape[0] / self.sampling_rate return t_stop - def _get_signal_size(self, block_index, seg_index, channel_indexes): + def _get_signal_size(self, block_index, seg_index, stream_index): + assert stream_index == 0 return self._raw_signals.shape[0] - def _get_signal_t_start(self, block_index, seg_index, channel_indexes): + def _get_signal_t_start(self, block_index, seg_index, stream_index): + assert stream_index == 0 return 0. - def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, + stream_index, channel_indexes): if channel_indexes is None: channel_indexes = slice(None) raw_signals = self._raw_signals[slice(i_start, i_stop), channel_indexes] - return raw_signals diff --git a/neo/rawio/rawmcsrawio.py b/neo/rawio/rawmcsrawio.py index 0f8394c45..4b2d97471 100644 --- a/neo/rawio/rawmcsrawio.py +++ b/neo/rawio/rawmcsrawio.py @@ -14,8 +14,8 @@ Author: Samuel Garcia """ -from .baserawio import (BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, - _event_channel_dtype) +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype) import numpy as np @@ -41,16 +41,19 @@ def _parse_header(self): self.sampling_rate = info['sampling_rate'] self.nb_channel = len(info['channel_names']) + # one unique stream + signal_streams = np.array([('Signals', '0')], dtype=_signal_stream_dtype) + self._raw_signals = np.memmap(self.filename, dtype=self.dtype, mode='r', offset=info['header_size']).reshape(-1, self.nb_channel) sig_channels = [] for c in range(self.nb_channel): - chan_id = c - group_id = 0 + chan_id = str(c) + stream_id = '0' sig_channels.append((info['channel_names'][c], chan_id, self.sampling_rate, self.dtype, info['signal_units'], info['signal_gain'], - info['signal_offset'], group_id)) + info['signal_offset'], stream_id)) sig_channels = np.array(sig_channels, dtype=_signal_channel_dtype) # No events @@ -58,15 +61,16 @@ def _parse_header(self): event_channels = np.array(event_channels, dtype=_event_channel_dtype) # No spikes - unit_channels = [] - unit_channels = np.array(unit_channels, dtype=_unit_channel_dtype) + spike_channels = [] + spike_channels = np.array(spike_channels, dtype=_spike_channel_dtype) # fille into header dict self.header = {} self.header['nb_block'] = 1 self.header['nb_segment'] = [1] + self.header['signal_streams'] = signal_streams self.header['signal_channels'] = sig_channels - self.header['unit_channels'] = unit_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels # insert some annotation at some place @@ -79,17 +83,17 @@ def _segment_t_stop(self, block_index, seg_index): t_stop = self._raw_signals.shape[0] / self.sampling_rate return t_stop - def _get_signal_size(self, block_index, seg_index, channel_indexes): + def _get_signal_size(self, block_index, seg_index, stream_index): return self._raw_signals.shape[0] - def _get_signal_t_start(self, block_index, seg_index, channel_indexes): + def _get_signal_t_start(self, block_index, seg_index, stream_index): return 0. - def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, + stream_index, channel_indexes): if channel_indexes is None: channel_indexes = slice(None) raw_signals = self._raw_signals[slice(i_start, i_stop), channel_indexes] - return raw_signals diff --git a/neo/rawio/spike2rawio.py b/neo/rawio/spike2rawio.py index bad47a8fa..220553335 100644 --- a/neo/rawio/spike2rawio.py +++ b/neo/rawio/spike2rawio.py @@ -17,10 +17,9 @@ Author: Samuel Garcia """ -# from __future__ import unicode_literals is not compatible with numpy.dtype both py2 py3 -from .baserawio import (BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, - _event_channel_dtype) +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype) import numpy as np from collections import OrderedDict @@ -34,7 +33,7 @@ class Spike2RawIO(BaseRawIO): rawmode = 'one-file' def __init__(self, filename='', take_ideal_sampling_rate=False, ced_units=True, - try_signal_grouping=True): + try_signal_grouping=True): BaseRawIO.__init__(self) self.filename = filename @@ -191,8 +190,8 @@ def _parse_header(self): self._seg_t_stops.append(t_stop) # create typed channels - sig_channels = [] - unit_channels = [] + signal_channels = [] + spike_channels = [] event_channels = [] self.internal_unit_ids = {} @@ -219,9 +218,9 @@ def _parse_header(self): gain = 1. offset = 0. sig_dtype = 'float32' - group_id = 0 - sig_channels.append((name, chan_id, sampling_rate, sig_dtype, - units, gain, offset, group_id)) + stream_id = '0' # set it after the loop + signal_channels.append((name, str(chan_id), sampling_rate, sig_dtype, + units, gain, offset, stream_id)) elif chan_info['kind'] in [2, 3, 4, 5, 8]: # Event @@ -255,78 +254,94 @@ def _parse_header(self): # All spike from one channel are group in one SpikeTrain unit_ids = ['all'] for unit_id in unit_ids: - unit_index = len(unit_channels) + unit_index = len(spike_channels) self.internal_unit_ids[unit_index] = (chan_id, unit_id) _id = "ch{}#{}".format(chan_id, unit_id) - unit_channels.append((name, _id, wf_units, wf_gain, wf_offset, + spike_channels.append((name, _id, wf_units, wf_gain, wf_offset, wf_left_sweep, wf_sampling_rate)) - sig_channels = np.array(sig_channels, dtype=_signal_channel_dtype) - unit_channels = np.array(unit_channels, dtype=_unit_channel_dtype) + signal_channels = np.array(signal_channels, dtype=_signal_channel_dtype) + spike_channels = np.array(spike_channels, dtype=_spike_channel_dtype) event_channels = np.array(event_channels, dtype=_event_channel_dtype) - if len(sig_channels) > 0: + if len(signal_channels) > 0: if self.try_signal_grouping: # try to group signals channel if same sampling_rate/dtype/... # it can raise error for some files (when they do not have signal length) common_keys = ['sampling_rate', 'dtype', 'units', 'gain', 'offset'] - characteristics = sig_channels[common_keys] + characteristics = signal_channels[common_keys] unique_characteristics = np.unique(characteristics) self._sig_dtypes = {} - for group_id, charact in enumerate(unique_characteristics): + signal_streams = [] + for stream_index, charact in enumerate(unique_characteristics): chan_grp_indexes, = np.nonzero(characteristics == charact) - sig_channels['group_id'][chan_grp_indexes] = group_id + stream_id = str(stream_index) + signal_channels['stream_id'][chan_grp_indexes] = stream_id # check same size for channel in groups for seg_index in range(nb_segment): sig_sizes = [] for ind in chan_grp_indexes: - chan_id = sig_channels[ind]['id'] + chan_id = int(signal_channels[ind]['id']) sig_size = np.sum(self._by_seg_data_blocks[chan_id][seg_index]['size']) sig_sizes.append(sig_size) sig_sizes = np.array(sig_sizes) assert np.all(sig_sizes == sig_sizes[0]),\ - 'Signal channel in groups do not have same size'\ - ', use try_signal_grouping=False' - self._sig_dtypes[group_id] = np.dtype(charact['dtype']) + 'Signal channel in groups do not have same size,'\ + 'use try_signal_grouping=False' + self._sig_dtypes[stream_id] = np.dtype(charact['dtype']) + signal_streams.append((f'Signal stream {stream_id}', stream_id)) + signal_streams = np.array(signal_streams, dtype=_signal_stream_dtype) else: # if try_signal_grouping fail the user can try to split each channel in # separate group - sig_channels['group_id'] = np.arange(sig_channels.size) - self._sig_dtypes = {s['group_id']: np.dtype(s['dtype']) for s in sig_channels} + signal_channels['stream_id'] = np.arange(signal_channels.size) + signal_streams = np.zeros(signal_channels.size, dtype=_signal_stream_dtype) + signal_streams['id'] = signal_channels['stream_id'] + signal_streams['name'] = signal_channels['name'] + self._sig_dtypes = {s['stream_id']: np.dtype(s['dtype']) for s in signal_channels} + else: + signal_streams = np.array([], dtype=_signal_stream_dtype) # fille into header dict self.header = {} self.header['nb_block'] = 1 self.header['nb_segment'] = [nb_segment] - self.header['signal_channels'] = sig_channels - self.header['unit_channels'] = unit_channels + self.header['signal_streams'] = signal_streams + self.header['signal_channels'] = signal_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels # Annotations self._generate_minimal_annotations() bl_ann = self.raw_annotations['blocks'][0] bl_ann['system_id'] = info['system_id'] - seg_ann = bl_ann['segments'][0] - seg_ann['system_id'] = info['system_id'] - - for c, sig_channel in enumerate(sig_channels): - chan_id = sig_channel['id'] - anasig_an = seg_ann['signals'][c] - anasig_an['physical_channel_index'] = self._channel_infos[chan_id]['phy_chan'] - anasig_an['comment'] = self._channel_infos[chan_id]['comment'] - - for c, unit_channel in enumerate(unit_channels): - chan_id, unit_id = self.internal_unit_ids[c] - unit_an = seg_ann['units'][c] - unit_an['physical_channel_index'] = self._channel_infos[chan_id]['phy_chan'] - unit_an['comment'] = self._channel_infos[chan_id]['comment'] - - for c, event_channel in enumerate(event_channels): - chan_id = int(event_channel['id']) - ev_an = seg_ann['events'][c] - ev_an['physical_channel_index'] = self._channel_infos[chan_id]['phy_chan'] - ev_an['comment'] = self._channel_infos[chan_id]['comment'] + for seg_index in range(nb_segment): + seg_ann = self.raw_annotations['blocks'][0]['segments'][seg_index] + seg_ann['system_id'] = info['system_id'] + + for c, stream_channel in enumerate(signal_streams): + stream_id = stream_channel['id'] + signal_an = self.raw_annotations['blocks'][0]['segments'][seg_index]['signals'][c] + mask = (signal_channels['stream_id'] == stream_id) + + for key in ('phy_chan', 'comment'): + values = [] + for chan_id in signal_channels[mask]['id']: + values.append(self._channel_infos[int(chan_id)][key]) + signal_an['__array_annotations__'][key] = np.array(values) + + for c, unit_channel in enumerate(spike_channels): + chan_id, unit_id = self.internal_unit_ids[c] + unit_an = self.raw_annotations['blocks'][0]['segments'][seg_index]['spikes'][c] + unit_an['physical_channel_index'] = self._channel_infos[chan_id]['phy_chan'] + unit_an['comment'] = self._channel_infos[chan_id]['comment'] + + for c, event_channel in enumerate(event_channels): + chan_id = int(event_channel['id']) + ev_an = self.raw_annotations['blocks'][0]['segments'][seg_index]['events'][c] + ev_an['physical_channel_index'] = self._channel_infos[chan_id]['phy_chan'] + ev_an['comment'] = self._channel_infos[chan_id]['comment'] def _source_name(self): return self.filename @@ -337,43 +352,48 @@ def _segment_t_start(self, block_index, seg_index): def _segment_t_stop(self, block_index, seg_index): return self._seg_t_stops[seg_index] * self._time_factor - def _check_channel_indexes(self, channel_indexes): - if channel_indexes is None: - channel_indexes = slice(None) - channel_indexes = np.arange(self.header['signal_channels'].size)[channel_indexes] - return channel_indexes - - def _get_signal_size(self, block_index, seg_index, channel_indexes): - channel_indexes = self._check_channel_indexes(channel_indexes) - chan_id = self.header['signal_channels'][channel_indexes[0]]['id'] - sig_size = np.sum(self._by_seg_data_blocks[chan_id][seg_index]['size']) + def _get_signal_size(self, block_index, seg_index, stream_index): + stream_id = self.header['signal_streams'][stream_index]['id'] + mask = self.header['signal_channels']['stream_id'] == stream_id + signal_channels = self.header['signal_channels'][mask] + chan_ids = signal_channels['id'] + chan_id0 = int(chan_ids[0]) + sig_size = np.sum(self._by_seg_data_blocks[chan_id0][seg_index]['size']) return sig_size - def _get_signal_t_start(self, block_index, seg_index, channel_indexes): - channel_indexes = self._check_channel_indexes(channel_indexes) - chan_id = self.header['signal_channels'][channel_indexes[0]]['id'] - return self._sig_t_starts[chan_id][seg_index] * self._time_factor + def _get_signal_t_start(self, block_index, seg_index, stream_index): + stream_id = self.header['signal_streams'][stream_index]['id'] + mask = self.header['signal_channels']['stream_id'] == stream_id + signal_channels = self.header['signal_channels'][mask] + chan_ids = signal_channels['id'] + chan_id0 = int(chan_ids[0]) + return self._sig_t_starts[chan_id0][seg_index] * self._time_factor - def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, + stream_index, channel_indexes): if i_start is None: i_start = 0 if i_stop is None: - i_stop = self._get_signal_size(block_index, seg_index, channel_indexes) + i_stop = self._get_signal_size(block_index, seg_index, stream_index) + + stream_id = self.header['signal_streams'][stream_index]['id'] + mask = self.header['signal_channels']['stream_id'] == stream_id + signal_channels = self.header['signal_channels'][mask] + chan_ids = signal_channels['id'] + self._sig_dtypes[stream_id] + + if channel_indexes is not None: + chan_ids = chan_ids[channel_indexes] - channel_indexes = self._check_channel_indexes(channel_indexes) - chan_index = channel_indexes[0] - chan_id = self.header['signal_channels'][chan_index]['id'] - group_id = self.header['signal_channels'][channel_indexes[0]]['group_id'] - dt = self._sig_dtypes[group_id] + dt = self._sig_dtypes[stream_id] - raw_signals = np.zeros((i_stop - i_start, len(channel_indexes)), dtype=dt) - for c, channel_index in enumerate(channel_indexes): + raw_signals = np.zeros((i_stop - i_start, len(chan_ids)), dtype=dt) + for c, chan_id in enumerate(chan_ids): + chan_id = int(chan_id) # NOTE: this actual way is slow because we run throught # the file for each channel. The loop should be reversed. # But there is no garanty that channels shared the same data block # indexes. So this make the job too difficult. - chan_header = self.header['signal_channels'][channel_index] - chan_id = chan_header['id'] data_blocks = self._by_seg_data_blocks[chan_id][seg_index] # loop over data blocks and get chunks @@ -481,7 +501,7 @@ def _spike_count(self, block_index, seg_index, unit_index): lim0, lim1, marker_filter=marker_filter) def _get_spike_timestamps(self, block_index, seg_index, unit_index, t_start, t_stop): - unit_header = self.header['unit_channels'][unit_index] + unit_header = self.header['spike_channels'][unit_index] chan_id, unit_id = self.internal_unit_ids[unit_index] if self.ced_units: @@ -501,7 +521,7 @@ def _rescale_spike_timestamp(self, spike_timestamps, dtype): return spike_times def _get_spike_raw_waveforms(self, block_index, seg_index, unit_index, t_start, t_stop): - unit_header = self.header['unit_channels'][unit_index] + unit_header = self.header['spike_channels'][unit_index] chan_id, unit_id = self.internal_unit_ids[unit_index] if self.ced_units: @@ -548,7 +568,7 @@ def _get_event_timestamps(self, block_index, seg_index, event_channel_index, t_s return timestamps, durations, labels - def _rescale_event_timestamp(self, event_timestamps, dtype): + def _rescale_event_timestamp(self, event_timestamps, dtype, event_channel_index): event_times = event_timestamps.astype(dtype) event_times *= self._time_factor return event_times @@ -569,8 +589,8 @@ def read_as_dict(fid, dtype): if dt[k].kind == 'S': v = v.decode('iso-8859-1') if len(v) > 0: - l = ord(v[0]) - v = v[1:l + 1] + length = ord(v[0]) + v = v[1:length + 1] info[k] = v return info @@ -608,8 +628,8 @@ def get_sample_interval(info, chan_info): Get sample interval for one channel """ if info['system_id'] in [1, 2, 3, 4, 5]: # Before version 5 - sample_interval = (int(chan_info['divide']) * info['us_per_time'] * - info['time_per_adc']) * 1e-6 + sample_interval = (int(chan_info['divide']) * info['us_per_time'] + * info['time_per_adc']) * 1e-6 else: sample_interval = (int(chan_info['l_chan_dvd']) * info['us_per_time'] * info['dtime_base']) diff --git a/neo/rawio/spikeglxrawio.py b/neo/rawio/spikeglxrawio.py index f43917d41..21cc4058a 100644 --- a/neo/rawio/spikeglxrawio.py +++ b/neo/rawio/spikeglxrawio.py @@ -5,7 +5,7 @@ Here an adaptation of the spikeglx tools into the neo rawio API. -Note that each pair of ".bin"/."meta" files is represented as a group of channels +Note that each pair of ".bin"/."meta" files is represented as a stream of channels that share the same sampling rate. It will be one AnalogSignal multi channel at neo.io level. @@ -39,7 +39,8 @@ Author : Samuel Garcia """ -from .baserawio import BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, _event_channel_dtype +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype) from pathlib import Path import os @@ -84,52 +85,34 @@ def _parse_header(self): self._memmaps[key] = data # create channel header - self._global_channel_to_stream = {} - self._global_channel_to_local_channel = [] - self._channel_location = {} - sig_channels = [] - global_chan = 0 - signal_annotations = [] + signal_streams = [] + signal_channels = [] for stream_name in stream_names: # take first segment info = self.signals_info_dict[0, stream_name] - group_id = stream_names.index(info['stream_name']) + stream_id = stream_name + stream_index = stream_names.index(info['stream_name']) + signal_streams.append((stream_name, stream_id)) # add channels to global list for local_chan in range(info['num_chan']): - self._global_channel_to_stream[global_chan] = info['stream_name'] - self._global_channel_to_local_channel.append(local_chan) chan_name = info['channel_names'][local_chan] - sig_channels.append((chan_name, global_chan, info['sampling_rate'], 'int16', + chan_id = f'{stream_name}#{chan_name}' + signal_channels.append((chan_name, chan_id, info['sampling_rate'], 'int16', info['units'], info['channel_gains'][local_chan], - info['channel_offsets'][local_chan], group_id)) + info['channel_offsets'][local_chan], stream_id)) - # annotation - ann = {} - ann['stream'] = info['stream_name'] - signal_annotations.append(ann) - - # channel location - if 'channel_location' in info: - self._channel_location[info['seg_index'], info['device']] = \ - info['channel_location'] - - # the channel id is a global counter and so equivalent to channel_index - # this is bad and should rather be changed to a string based id - global_chan += 1 - - sig_channels = np.array(sig_channels, dtype=_signal_channel_dtype) - self._global_channel_to_local_channel = np.array(self._global_channel_to_local_channel, - dtype='int64') + signal_streams = np.array(signal_streams, dtype=_signal_stream_dtype) + signal_channels = np.array(signal_channels, dtype=_signal_channel_dtype) # No events event_channels = [] event_channels = np.array(event_channels, dtype=_event_channel_dtype) # No spikes - unit_channels = [] - unit_channels = np.array(unit_channels, dtype=_unit_channel_dtype) + spike_channels = [] + spike_channels = np.array(spike_channels, dtype=_spike_channel_dtype) # deal with nb_segment and t_start/t_stop per segment self._t_starts = {seg_index: 0. for seg_index in range(nb_segment)} @@ -144,24 +127,32 @@ def _parse_header(self): self.header = {} self.header['nb_block'] = 1 self.header['nb_segment'] = [nb_segment] - self.header['signal_channels'] = sig_channels - self.header['unit_channels'] = unit_channels + self.header['signal_streams'] = signal_streams + self.header['signal_channels'] = signal_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels # insert some annotation at some place self._generate_minimal_annotations() self._generate_minimal_annotations() block_ann = self.raw_annotations['blocks'][0] - block_ann['file_origin'] = self.dirname for seg_index in range(nb_segment): - seg_ann = block_ann['segments'][seg_index] - seg_ann['file_origin'] = self.dirname + seg_ann = self.raw_annotations['blocks'][0]['segments'][seg_index] seg_ann['name'] = "Segment {}".format(seg_index) - for c in range(sig_channels.size): - anasig_an = seg_ann['signals'][c] - anasig_an.update(signal_annotations[c]) + for c, signal_stream in enumerate(signal_streams): + stream_name = signal_stream['name'] + sig_ann = self.raw_annotations['blocks'][0]['segments'][seg_index]['signals'][c] + + # channel location + info = self.signals_info_dict[seg_index, stream_name] + if 'channel_location' in info: + loc = info['channel_location'] + # one fake channel for "sys0" + loc = np.concatenate((loc, [[0., 0.]]), axis=0) + for ndim in range(loc.shape[1]): + sig_ann['__array_annotations__'][f'channel_location_{ndim}'] = loc[:, ndim] def _segment_t_start(self, block_index, seg_index): return 0. @@ -169,44 +160,31 @@ def _segment_t_start(self, block_index, seg_index): def _segment_t_stop(self, block_index, seg_index): return self._t_stops[seg_index] - def _get_signal_size(self, block_index, seg_index, channel_indexes=None): - assert channel_indexes is not None - stream_name = self._global_channel_to_stream[channel_indexes[0]] - memmap = self._memmaps[seg_index, stream_name] + def _get_signal_size(self, block_index, seg_index, stream_index): + stream_id = self.header['signal_streams'][stream_index]['id'] + memmap = self._memmaps[seg_index, stream_id] return int(memmap.shape[0]) - def _get_signal_t_start(self, block_index, seg_index, channel_indexes): + def _get_signal_t_start(self, block_index, seg_index, stream_index): return 0. - def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): - assert channel_indexes is not None - stream_name = self._global_channel_to_stream[channel_indexes[0]] - memmap = self._memmaps[seg_index, stream_name] + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, + stream_index, channel_indexes): + stream_id = self.header['signal_streams'][stream_index]['id'] + memmap = self._memmaps[seg_index, stream_id] - local_chans = self._global_channel_to_local_channel[channel_indexes] - if np.all(np.diff(local_chans) == 1): - # consecutive channel then slice this avoid a copy (because of ndarray.take(...) - # and so keep the underlying memmap - local_chans = slice(local_chans[0], local_chans[0] + len(local_chans)) + if channel_indexes is None: + channel_indexes = slice(channel_indexes) - raw_signals = memmap[slice(i_start, i_stop), local_chans] + if not isinstance(channel_indexes, slice): + if np.all(np.diff(channel_indexes) == 1): + # consecutive channel then slice this avoid a copy (because of ndarray.take(...) + # and so keep the underlying memmap + local_chans = slice(channel_indexes[0], channel_indexes[0] + len(channel_indexes)) - return raw_signals + raw_signals = memmap[slice(i_start, i_stop), channel_indexes] - def get_channel_location(self, seg_index=0, device=None, x_pitch=21, y_pitch=20): - # x_pitch=21, y_pitch=2 are taken from spikeinterface implementation. - # See also `_parse_spikeglx_metafile` in spikeglxrecordingextractor.py - # in https://github.com/SpikeInterface/spikeextractors - # This need to be check. - if device is None: - if len(self._channel_location) == 1: - locations = list(self._channel_location.values())[0] - else: - raise ValueError('device must specified') - else: - locations = self._channel_location[seg_index, device] - locations = locations * [[x_pitch, y_pitch]] - return locations + return raw_signals def scan_files(dirname): diff --git a/neo/rawio/tdtrawio.py b/neo/rawio/tdtrawio.py index d1938c55a..473bba0d7 100644 --- a/neo/rawio/tdtrawio.py +++ b/neo/rawio/tdtrawio.py @@ -22,7 +22,8 @@ Author: Samuel Garcia, SummitKwan, Chadwick Boulay """ -from .baserawio import BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, _event_channel_dtype +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype) import numpy as np import os @@ -96,7 +97,8 @@ def _parse_header(self): tsq_filename = os.path.join(path, tankname + '_' + segment_name + '.tsq') tsq = np.fromfile(tsq_filename, dtype=tsq_dtype) self._tsq.append(tsq) - # Start and stop times are only found in the second and last header row, respectively. + # Start and stop times are only found in the second + # and last header row, respectively. if tsq[1]['evname'] == chr(EVMARK_STARTBLOCK).encode(): self._seg_t_starts.append(tsq[1]['timestamp']) else: @@ -136,6 +138,7 @@ def _parse_header(self): self._global_t_start = self._seg_t_starts[0] # signal channels EVTYPE_STREAM + signal_streams = [] signal_channels = [] self._sigs_data_buf = {seg_index: {} for seg_index in range(nb_segment)} self._sigs_index = {seg_index: {} for seg_index in range(nb_segment)} @@ -147,12 +150,16 @@ def _parse_header(self): for seg_index in range(nb_segment)} # key = seg_index then group_id keep = info_channel_groups['TankEvType'] == EVTYPE_STREAM - for group_id, info in enumerate(info_channel_groups[keep]): - self._sig_sample_per_chunk[group_id] = info['NumPoints'] + for stream_index, info in enumerate(info_channel_groups[keep]): + self._sig_sample_per_chunk[stream_index] = info['NumPoints'] + + stream_name = str(info['StoreName']) + stream_id = f'{stream_index}' + signal_streams.append((stream_name, stream_id)) for c in range(info['NumChan']): - chan_index = len(signal_channels) - chan_id = c + 1 # If several StoreName then chan_id is not unique in TDT!!!!! + global_chan_index = len(signal_channels) + chan_id = c + 1 # several StoreName can have same chan_id: this is ok # loop over segment to get sampling_rate/data_index/data_buffer sampling_rate = None @@ -164,20 +171,20 @@ def _parse_header(self): (tsq['evname'] == info['StoreName']) & \ (tsq['channel'] == chan_id) data_index = tsq[mask].copy() - self._sigs_index[seg_index][chan_index] = data_index + self._sigs_index[seg_index][global_chan_index] = data_index size = info['NumPoints'] * data_index.size - if group_id not in self._sigs_lengths[seg_index]: - self._sigs_lengths[seg_index][group_id] = size + if stream_index not in self._sigs_lengths[seg_index]: + self._sigs_lengths[seg_index][stream_index] = size else: - assert self._sigs_lengths[seg_index][group_id] == size + assert self._sigs_lengths[seg_index][stream_index] == size # signal start time, relative to start of segment t_start = data_index['timestamp'][0] - if group_id not in self._sigs_t_start[seg_index]: - self._sigs_t_start[seg_index][group_id] = t_start + if stream_index not in self._sigs_t_start[seg_index]: + self._sigs_t_start[seg_index][stream_index] = t_start else: - assert self._sigs_t_start[seg_index][group_id] == t_start + assert self._sigs_t_start[seg_index][stream_index] == t_start # sampling_rate and dtype _sampling_rate = float(data_index['frequency'][0]) @@ -185,10 +192,10 @@ def _parse_header(self): if sampling_rate is None: sampling_rate = _sampling_rate dtype = _dtype - if group_id not in self._sig_dtype_by_group: - self._sig_dtype_by_group[group_id] = np.dtype(dtype) + if stream_index not in self._sig_dtype_by_group: + self._sig_dtype_by_group[stream_index] = np.dtype(dtype) else: - assert self._sig_dtype_by_group[group_id] == dtype + assert self._sig_dtype_by_group[stream_index] == dtype else: assert sampling_rate == _sampling_rate, 'sampling is changing!!!' assert dtype == _dtype, 'sampling is changing!!!' @@ -203,22 +210,23 @@ def _parse_header(self): else: data = self._tev_datas[seg_index] assert data is not None, 'no TEV nor SEV' - self._sigs_data_buf[seg_index][chan_index] = data + self._sigs_data_buf[seg_index][global_chan_index] = data chan_name = '{} {}'.format(info['StoreName'], c + 1) sampling_rate = sampling_rate units = 'V' # WARNING this is not sur at all gain = 1. offset = 0. - signal_channels.append((chan_name, chan_id, sampling_rate, dtype, - units, gain, offset, group_id)) + signal_channels.append((chan_name, str(chan_id), sampling_rate, dtype, + units, gain, offset, stream_id)) + signal_streams = np.array(signal_streams, dtype=_signal_stream_dtype) signal_channels = np.array(signal_channels, dtype=_signal_channel_dtype) # unit channels EVTYPE_SNIP self.internal_unit_ids = {} self._waveforms_size = [] self._waveforms_dtype = [] - unit_channels = [] + spike_channels = [] keep = info_channel_groups['TankEvType'] == EVTYPE_SNIP tsq = np.hstack(self._tsq) # If there is no chance the differet TSQ files will have different units, @@ -231,7 +239,7 @@ def _parse_header(self): (tsq['channel'] == chan_id) unit_ids = np.unique(tsq[mask]['sortcode']) for unit_id in unit_ids: - unit_index = len(unit_channels) + unit_index = len(spike_channels) self.internal_unit_ids[unit_index] = (info['StoreName'], chan_id, unit_id) unit_name = "ch{}#{}".format(chan_id, unit_id) @@ -240,14 +248,14 @@ def _parse_header(self): wf_offset = 0. wf_left_sweep = info['NumPoints'] // 2 wf_sampling_rate = info['SampleFreq'] - unit_channels.append((unit_name, '{}'.format(unit_id), + spike_channels.append((unit_name, '{}'.format(unit_id), wf_units, wf_gain, wf_offset, wf_left_sweep, wf_sampling_rate)) self._waveforms_size.append(info['NumPoints']) self._waveforms_dtype.append(np.dtype(data_formats[info['DataFormat']])) - unit_channels = np.array(unit_channels, dtype=_unit_channel_dtype) + spike_channels = np.array(spike_channels, dtype=_spike_channel_dtype) # signal channels EVTYPE_STRON event_channels = [] @@ -263,8 +271,9 @@ def _parse_header(self): self.header = {} self.header['nb_block'] = 1 self.header['nb_segment'] = [nb_segment] + self.header['signal_streams'] = signal_streams self.header['signal_channels'] = signal_channels - self.header['unit_channels'] = unit_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels # Annotations only standard ones: @@ -282,36 +291,42 @@ def _segment_t_start(self, block_index, seg_index): def _segment_t_stop(self, block_index, seg_index): return self._seg_t_stops[seg_index] - self._global_t_start - def _get_signal_size(self, block_index, seg_index, channel_indexes): - group_id = self.header['signal_channels'][channel_indexes[0]]['group_id'] - size = self._sigs_lengths[seg_index][group_id] + def _get_signal_size(self, block_index, seg_index, stream_index): + size = self._sigs_lengths[seg_index][stream_index] return size - def _get_signal_t_start(self, block_index, seg_index, channel_indexes): - group_id = self.header['signal_channels'][channel_indexes[0]]['group_id'] - return self._sigs_t_start[seg_index][group_id] - self._global_t_start - - def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): - # check of channel_indexes is same group_id is done outside (BaseRawIO) - # so first is identique to others - group_id = self.header['signal_channels'][channel_indexes[0]]['group_id'] + def _get_signal_t_start(self, block_index, seg_index, stream_index): + return self._sigs_t_start[seg_index][stream_index] - self._global_t_start + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, + stream_index, channel_indexes): if i_start is None: i_start = 0 if i_stop is None: - i_stop = self._sigs_lengths[seg_index][group_id] + i_stop = self._sigs_lengths[seg_index][stream_index] + + stream_id = self.header['signal_streams'][stream_index]['id'] + signal_channels = self.header['signal_channels'] + mask = signal_channels['stream_id'] == stream_id + global_chan_indexes = np.arange(signal_channels.size)[mask] + signal_channels = signal_channels[mask] + + if channel_indexes is None: + channel_indexes = slice(None) + global_chan_indexes = global_chan_indexes[channel_indexes] + signal_channels = signal_channels[channel_indexes] - dt = self._sig_dtype_by_group[group_id] - raw_signals = np.zeros((i_stop - i_start, len(channel_indexes)), dtype=dt) + dt = self._sig_dtype_by_group[stream_index] + raw_signals = np.zeros((i_stop - i_start, signal_channels.size), dtype=dt) - sample_per_chunk = self._sig_sample_per_chunk[group_id] + sample_per_chunk = self._sig_sample_per_chunk[stream_index] bl0 = i_start // sample_per_chunk bl1 = int(np.ceil(i_stop / sample_per_chunk)) chunk_nb_bytes = sample_per_chunk * dt.itemsize - for c, channel_index in enumerate(channel_indexes): - data_index = self._sigs_index[seg_index][channel_index] - data_buf = self._sigs_data_buf[seg_index][channel_index] + for c, global_index in enumerate(global_chan_indexes): + data_index = self._sigs_index[seg_index][global_index] + data_buf = self._sigs_data_buf[seg_index][global_index] # loop over data blocks and get chunks ind = 0 @@ -420,7 +435,7 @@ def _get_event_timestamps(self, block_index, seg_index, event_channel_index, t_s # it was not implemented in previous IO. return timestamps, durations, labels - def _rescale_event_timestamp(self, event_timestamps, dtype): + def _rescale_event_timestamp(self, event_timestamps, dtype, event_channel_index): # already in s ev_times = event_timestamps.astype(dtype) return ev_times diff --git a/neo/rawio/winedrrawio.py b/neo/rawio/winedrrawio.py index 148c683ed..fde928aa4 100644 --- a/neo/rawio/winedrrawio.py +++ b/neo/rawio/winedrrawio.py @@ -9,8 +9,8 @@ """ -from .baserawio import (BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, - _event_channel_dtype) +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype, _common_sig_characteristics) import numpy as np @@ -56,7 +56,7 @@ def _parse_header(self): DT *= .001 self._sampling_rate = 1. / DT - sig_channels = [] + signal_channels = [] for c in range(header['NC']): YCF = float(header['YCF%d' % c].replace(',', '.')) YAG = float(header['YAG%d' % c].replace(',', '.')) @@ -69,26 +69,36 @@ def _parse_header(self): units = header['YU%d' % c] gain = AD / (YCF * YAG * (ADCMAX + 1)) offset = -YZ * gain - group_id = 0 - sig_channels.append((name, chan_id, self._sampling_rate, 'int16', - units, gain, offset, group_id)) + stream_id = '0' + signal_channels.append((name, str(chan_id), self._sampling_rate, 'int16', + units, gain, offset, stream_id)) - sig_channels = np.array(sig_channels, dtype=_signal_channel_dtype) + signal_channels = np.array(signal_channels, dtype=_signal_channel_dtype) + + characteristics = signal_channels[_common_sig_characteristics] + unique_characteristics = np.unique(characteristics) + signal_streams = [] + for i in range(unique_characteristics.size): + mask = unique_characteristics[i] == characteristics + signal_channels['stream_id'][mask] = str(i) + signal_streams.append((f'stream {i}', str(i))) + signal_streams = np.array(signal_streams, dtype=_signal_stream_dtype) # No events event_channels = [] event_channels = np.array(event_channels, dtype=_event_channel_dtype) # No spikes - unit_channels = [] - unit_channels = np.array(unit_channels, dtype=_unit_channel_dtype) + spike_channels = [] + spike_channels = np.array(spike_channels, dtype=_spike_channel_dtype) # fille into header dict self.header = {} self.header['nb_block'] = 1 self.header['nb_segment'] = [1] - self.header['signal_channels'] = sig_channels - self.header['unit_channels'] = unit_channels + self.header['signal_streams'] = signal_streams + self.header['signal_channels'] = signal_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels # insert some annotation at some place @@ -101,20 +111,19 @@ def _segment_t_stop(self, block_index, seg_index): t_stop = self._raw_signals.shape[0] / self._sampling_rate return t_stop - def _get_signal_size(self, block_index, seg_index, channel_indexes): + def _get_signal_size(self, block_index, seg_index, stream_index): return self._raw_signals.shape[0] - def _get_signal_t_start(self, block_index, seg_index, channel_indexes): + def _get_signal_t_start(self, block_index, seg_index, stream_index): return 0. - def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): - # WARNING check if id or index for signals (in the old IO it was ids - # ~ raw_signals = self._raw_signals[slice(i_start, i_stop), channel_indexes] + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, + stream_index, channel_indexes): + stream_id = self.header['signal_streams'][stream_index]['id'] + global_channel_indexes, = np.nonzero(self.header['signal_channels'] + ['stream_id'] == stream_id) if channel_indexes is None: - channel_indexes = np.arange(self.header['signal_channels'].size) - - l = self.header['signal_channels']['id'].tolist() - channel_ids = [l.index(c) for c in channel_indexes] - raw_signals = self._raw_signals[slice(i_start, i_stop), channel_ids] - + channel_indexes = slice(None) + global_channel_indexes = global_channel_indexes[channel_indexes] + raw_signals = self._raw_signals[slice(i_start, i_stop), global_channel_indexes] return raw_signals diff --git a/neo/rawio/winwcprawio.py b/neo/rawio/winwcprawio.py index 7547c5f12..9bbb1c5cb 100644 --- a/neo/rawio/winwcprawio.py +++ b/neo/rawio/winwcprawio.py @@ -5,11 +5,10 @@ WinWCP is free: http://spider.science.strath.ac.uk/sipbs/software.htm -Author : sgarcia Author: Samuel Garcia """ -from .baserawio import (BaseRawIO, _signal_channel_dtype, _unit_channel_dtype, - _event_channel_dtype) +from .baserawio import (BaseRawIO, _signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype, _common_sig_characteristics) import numpy as np @@ -77,7 +76,7 @@ def _parse_header(self): self._sampling_rate = 1. / all_sampling_interval[0] - sig_channels = [] + signal_channels = [] for c in range(header['NC']): YG = float(header['YG%d' % c].replace(',', '.')) ADCMAX = header['ADCMAX'] @@ -88,26 +87,36 @@ def _parse_header(self): units = header['YU%d' % c] gain = VMax / ADCMAX / YG offset = 0. - group_id = 0 - sig_channels.append((name, chan_id, self._sampling_rate, 'int16', - units, gain, offset, group_id)) + stream_id = '0' + signal_channels.append((name, chan_id, self._sampling_rate, 'int16', + units, gain, offset, stream_id)) - sig_channels = np.array(sig_channels, dtype=_signal_channel_dtype) + signal_channels = np.array(signal_channels, dtype=_signal_channel_dtype) + + characteristics = signal_channels[_common_sig_characteristics] + unique_characteristics = np.unique(characteristics) + signal_streams = [] + for i in range(unique_characteristics.size): + mask = unique_characteristics[i] == characteristics + signal_channels['stream_id'][mask] = str(i) + signal_streams.append((f'stream {i}', str(i))) + signal_streams = np.array(signal_streams, dtype=_signal_stream_dtype) # No events event_channels = [] event_channels = np.array(event_channels, dtype=_event_channel_dtype) # No spikes - unit_channels = [] - unit_channels = np.array(unit_channels, dtype=_unit_channel_dtype) + spike_channels = [] + spike_channels = np.array(spike_channels, dtype=_spike_channel_dtype) # fille into header dict self.header = {} self.header['nb_block'] = 1 self.header['nb_segment'] = [nb_segment] - self.header['signal_channels'] = sig_channels - self.header['unit_channels'] = unit_channels + self.header['signal_streams'] = signal_streams + self.header['signal_channels'] = signal_channels + self.header['spike_channels'] = spike_channels self.header['event_channels'] = event_channels # insert some annotation at some place @@ -120,21 +129,21 @@ def _segment_t_stop(self, block_index, seg_index): t_stop = self._raw_signals[seg_index].shape[0] / self._sampling_rate return t_stop - def _get_signal_size(self, block_index, seg_index, channel_indexes): + def _get_signal_size(self, block_index, seg_index, stream_index): return self._raw_signals[seg_index].shape[0] - def _get_signal_t_start(self, block_index, seg_index, channel_indexes): + def _get_signal_t_start(self, block_index, seg_index, stream_index): return 0. - def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, channel_indexes): - # WARNING check if id or index for signals (in the old IO it was ids - # ~ raw_signals = self._raw_signals[seg_index][slice(i_start, i_stop), channel_indexes] + def _get_analogsignal_chunk(self, block_index, seg_index, i_start, i_stop, + stream_index, channel_indexes): + stream_id = self.header['signal_streams'][stream_index]['id'] + global_channel_indexes, = np.nonzero(self.header['signal_channels'] + ['stream_id'] == stream_id) if channel_indexes is None: - channel_indexes = np.arange(self.header['signal_channels'].size) - - ids = self.header['signal_channels']['id'].tolist() - channel_ids = [ids.index(c) for c in channel_indexes] - raw_signals = self._raw_signals[seg_index][slice(i_start, i_stop), channel_ids] + channel_indexes = slice(None) + inds = global_channel_indexes[channel_indexes] + raw_signals = self._raw_signals[seg_index][slice(i_start, i_stop), inds] return raw_signals diff --git a/neo/test/coretest/test_segment.py b/neo/test/coretest/test_segment.py index a0784a7fa..bb6222e0d 100644 --- a/neo/test/coretest/test_segment.py +++ b/neo/test/coretest/test_segment.py @@ -239,11 +239,12 @@ def test_times(self): reader = ExampleRawIO(filename='my_filename.fake') reader.parse_header() - proxy_anasig = AnalogSignalProxy(rawio=reader, global_channel_indexes=None, block_index=0, - seg_index=0) + proxy_anasig = AnalogSignalProxy(rawio=reader, + stream_index=0, inner_stream_channels=None, + block_index=0, seg_index=0) seg.analogsignals.append(proxy_anasig) - proxy_st = SpikeTrainProxy(rawio=reader, unit_index=0, block_index=0, seg_index=0) + proxy_st = SpikeTrainProxy(rawio=reader, spike_channel_index=0, block_index=0, seg_index=0) seg.spiketrains.append(proxy_st) proxy_event = EventProxy(rawio=reader, event_channel_index=0, block_index=0, seg_index=0) @@ -939,11 +940,11 @@ def test__time_slice(self): reader.parse_header() proxy_anasig = AnalogSignalProxy(rawio=reader, - global_channel_indexes=None, - block_index=0, seg_index=0) + stream_index=0, inner_stream_channels=None, + block_index=0, seg_index=0) seg.analogsignals.append(proxy_anasig) - proxy_st = SpikeTrainProxy(rawio=reader, unit_index=0, + proxy_st = SpikeTrainProxy(rawio=reader, spike_channel_index=0, block_index=0, seg_index=0) seg.spiketrains.append(proxy_st) diff --git a/neo/test/iotest/test_axographio.py b/neo/test/iotest/test_axographio.py index 11fc73e36..1caedb038 100644 --- a/neo/test/iotest/test_axographio.py +++ b/neo/test/iotest/test_axographio.py @@ -259,8 +259,8 @@ def test_group_by_same_units(self): assert_equal(len(blk.groups), 1) assert_equal(len(blk.segments[0].analogsignals), 1) - names = [sig.name for sig in blk.segments[0].analogsignals] - assert_equal(names, ['Channel bundle (CAP,STIM) ']) + chan_names = blk.segments[0].analogsignals[0].array_annotations['channel_names'] + assert_equal(chan_names, ['CAP', 'STIM']) sig = blk.segments[0].analogsignals[0][:5] arr = sig.as_array('V') diff --git a/neo/test/iotest/test_exampleio.py b/neo/test/iotest/test_exampleio.py index 7f548635f..cb0674a13 100644 --- a/neo/test/iotest/test_exampleio.py +++ b/neo/test/iotest/test_exampleio.py @@ -15,8 +15,6 @@ # This run standart tests, this is mandatory for all IO - - class TestExampleIO(BaseTestIO, unittest.TestCase, ): ioclass = ExampleIO files_to_test = ['fake1', @@ -24,7 +22,10 @@ class TestExampleIO(BaseTestIO, unittest.TestCase, ): ] files_to_download = [] - +# This is the minimal variables that are required +# to run the common IO tests. IO specific tests +# can be added here and will be run automatically +# in addition to the common tests. class Specific_TestExampleIO(unittest.TestCase): def test_read_segment_lazy(self): r = ExampleIO(filename=None) diff --git a/neo/test/iotest/test_proxyobjects.py b/neo/test/iotest/test_proxyobjects.py index ae97463b6..bdae9f4a7 100644 --- a/neo/test/iotest/test_proxyobjects.py +++ b/neo/test/iotest/test_proxyobjects.py @@ -29,8 +29,9 @@ def setUp(self): class TestAnalogSignalProxy(BaseProxyTest): def test_AnalogSignalProxy(self): - proxy_anasig = AnalogSignalProxy(rawio=self.reader, global_channel_indexes=None, - block_index=0, seg_index=0,) + proxy_anasig = AnalogSignalProxy(rawio=self.reader, + stream_index=0, inner_stream_channels=None, + block_index=0, seg_index=0,) assert proxy_anasig.sampling_rate == 10 * pq.kHz assert proxy_anasig.t_start == 0 * pq.s @@ -47,14 +48,14 @@ def test_AnalogSignalProxy(self): anasig = proxy_anasig.load(time_slice=(2. * pq.s, 5 * pq.s)) assert anasig.t_start == 2. * pq.s assert anasig.duration == 3. * pq.s - assert anasig.shape == (30000, 16) + assert anasig.shape == (30000, 8) assert_same_attributes(proxy_anasig.time_slice(2. * pq.s, 5 * pq.s), anasig) # ceil next sample when slicing anasig = proxy_anasig.load(time_slice=(1.99999 * pq.s, 5.000001 * pq.s)) assert anasig.t_start == 2. * pq.s assert anasig.duration == 3. * pq.s - assert anasig.shape == (30000, 16) + assert anasig.shape == (30000, 8) # buggy time slice with self.assertRaises(AssertionError): @@ -63,11 +64,11 @@ def test_AnalogSignalProxy(self): assert proxy_anasig.t_stop == 10 * pq.s # select channels - anasig = proxy_anasig.load(channel_indexes=[3, 4, 9]) + anasig = proxy_anasig.load(channel_indexes=[3, 4, 5]) assert anasig.shape[1] == 3 # select channels and slice times - anasig = proxy_anasig.load(time_slice=(2. * pq.s, 5 * pq.s), channel_indexes=[3, 4, 9]) + anasig = proxy_anasig.load(time_slice=(2. * pq.s, 5 * pq.s), channel_indexes=[3, 4, 5]) assert anasig.shape == (30000, 3) # magnitude mode rescaled @@ -84,27 +85,31 @@ def test_AnalogSignalProxy(self): assert_arrays_almost_equal(anasig_float, anasig_int.rescale('uV'), 1e-9) # test array_annotations - assert 'info' in proxy_anasig.array_annotations - assert proxy_anasig.array_annotations['info'].size == 16 - assert 'info' in anasig_float.array_annotations - assert anasig_float.array_annotations['info'].size == 16 + assert '__array_annotations__' not in proxy_anasig.annotations + assert 'impedance' in proxy_anasig.array_annotations + assert proxy_anasig.array_annotations['impedance'].size == 8 + assert 'impedance' in anasig_float.array_annotations + assert anasig_float.array_annotations['impedance'].size == 8 def test_global_local_channel_indexes(self): proxy_anasig = AnalogSignalProxy(rawio=self.reader, - global_channel_indexes=slice(0, 10, 2), block_index=0, seg_index=0) + stream_index=0, inner_stream_channels=slice(0, 8, 2), + block_index=0, seg_index=0) - assert proxy_anasig.shape == (100000, 5) - assert '(ch0,ch2,ch4,ch6,ch8)' in proxy_anasig.name + assert proxy_anasig.shape == (100000, 4) + assert np.array_equal(proxy_anasig.array_annotations['channel_names'], + ['ch0', 'ch2', 'ch4', 'ch6']) # should be channel ch0 and ch6 anasig = proxy_anasig.load(channel_indexes=[0, 3]) assert anasig.shape == (100000, 2) - assert '(ch0,ch6)' in anasig.name + assert np.array_equal(anasig.array_annotations['channel_names'], + ['ch0', 'ch6']) class TestSpikeTrainProxy(BaseProxyTest): def test_SpikeTrainProxy(self): - proxy_sptr = SpikeTrainProxy(rawio=self.reader, unit_index=0, + proxy_sptr = SpikeTrainProxy(rawio=self.reader, spike_channel_index=0, block_index=0, seg_index=0) assert proxy_sptr.name == 'unit0' @@ -160,6 +165,10 @@ def test_SpikeTrainProxy(self): sptr = proxy_sptr.load(load_waveforms=True, time_slice=(250 * pq.ms, 500 * pq.ms)) assert sptr.waveforms.shape == (6, 1, 50) + # test array_annotations + assert '__array_annotations__' not in proxy_sptr.annotations + assert 'amplitudes' in proxy_sptr.array_annotations + class TestEventProxy(BaseProxyTest): def test_EventProxy(self): @@ -186,6 +195,11 @@ def test_EventProxy(self): event = proxy_event.load(time_slice=(2 * pq.s, 15 * pq.s)) event = proxy_event.load(time_slice=(2 * pq.s, 15 * pq.s), strict_slicing=False) + # test annotations/array_annotations + assert '__array_annotations__' not in proxy_event.annotations + assert 'nickname' in proxy_event.annotations + assert 'button' in proxy_event.array_annotations + class TestEpochProxy(BaseProxyTest): def test_EpochProxy(self): @@ -213,17 +227,21 @@ def test_EpochProxy(self): epoch = proxy_epoch.load(time_slice=(2 * pq.s, 15 * pq.s)) epoch = proxy_epoch.load(time_slice=(2 * pq.s, 15 * pq.s), strict_slicing=False) + # test annotations/array_annotations + assert '__array_annotations__' not in proxy_epoch.annotations + assert 'nickname' in proxy_epoch.annotations + class TestSegmentWithProxy(BaseProxyTest): def test_segment_with_proxy(self): seg = Segment() proxy_anasig = AnalogSignalProxy(rawio=self.reader, - global_channel_indexes=None, + stream_index=0, inner_stream_channels=None, block_index=0, seg_index=0,) seg.analogsignals.append(proxy_anasig) - proxy_sptr = SpikeTrainProxy(rawio=self.reader, unit_index=0, + proxy_sptr = SpikeTrainProxy(rawio=self.reader, spike_channel_index=0, block_index=0, seg_index=0) seg.spiketrains.append(proxy_sptr) diff --git a/neo/test/iotest/test_tdtio.py b/neo/test/iotest/test_tdtio.py index 530622a42..49d251a88 100644 --- a/neo/test/iotest/test_tdtio.py +++ b/neo/test/iotest/test_tdtio.py @@ -33,30 +33,19 @@ def test_signal_group_mode(self): filename='aep_05', directory=self.local_test_dir, clean=False) - # TdtIO is a hard case they are 3 groups at rawio level - # there are 3 groups of signals - nb_sigs_by_group = [1, 16, 16] + # In this TDT dataset there are 3 signal streams + nb_sigs_by_stream = [16, 1, 16] - signal_group_mode = 'group-by-same-units' reader = TdtIO(dirname=dirname) - bl = reader.read_block(signal_group_mode=signal_group_mode) + bl = reader.read_block() for seg in bl.segments: assert len(seg.analogsignals) == 3 i = 0 for anasig in seg.analogsignals: - # print(anasig.shape, anasig.sampling_rate) - assert anasig.shape[1] == nb_sigs_by_group[i] + # print(anasig.shape, anasig.sampling_rate, nb_sigs_by_stream[i]) + assert anasig.shape[1] == nb_sigs_by_stream[i] i += 1 - signal_group_mode = 'split-all' - reader = TdtIO(dirname=dirname) - bl = reader.read_block(signal_group_mode=signal_group_mode) - for seg in bl.segments: - assert len(seg.analogsignals) == np.sum(nb_sigs_by_group) - for anasig in seg.analogsignals: - # print(anasig.shape, anasig.sampling_rate) - assert anasig.shape[1] == 1 - if __name__ == "__main__": unittest.main() diff --git a/neo/test/rawiotest/common_rawio_test.py b/neo/test/rawiotest/common_rawio_test.py index 88c9e6aa5..d40734506 100644 --- a/neo/test/rawiotest/common_rawio_test.py +++ b/neo/test/rawiotest/common_rawio_test.py @@ -132,7 +132,7 @@ def test_read_all(self): # ~ ax.plot(sigs[:, 0]) # ~ plt.show() - # ~ nb_unit = reader.unit_channels_count() + # ~ nb_unit = reader.spike_channels_count() # ~ for unit_index in range(nb_unit): # ~ wfs = reader.spike_raw_waveforms(block_index=0, seg_index=0, # ~ unit_index=unit_index) diff --git a/neo/test/rawiotest/rawio_compliance.py b/neo/test/rawiotest/rawio_compliance.py index b8f6cd2ae..a40cb5d27 100644 --- a/neo/test/rawiotest/rawio_compliance.py +++ b/neo/test/rawiotest/rawio_compliance.py @@ -16,8 +16,8 @@ import numpy as np -from neo.rawio.baserawio import (_signal_channel_dtype, _unit_channel_dtype, - _event_channel_dtype, _common_sig_characteristics) +from neo.rawio.baserawio import (_signal_channel_dtype, _signal_stream_dtype, + _spike_channel_dtype, _event_channel_dtype, _common_sig_characteristics) def print_class(reader): @@ -27,30 +27,43 @@ def print_class(reader): def header_is_total(reader): """ Test if hedaer contains: + * 'nb_block' + * 'nb_segment' + * 'signal_streams' * 'signal_channels' - * 'unit_channels' + * 'spike_channels' * 'event_channels' """ h = reader.header + assert 'nb_block' in h, "`nb_block`missing in header" + assert 'nb_segment' in h, "`nb_segment`missing in header" + assert len(h['nb_segment']) == h['nb_block'] + + assert 'signal_streams' in h, 'signal_streams missing in header' + if h['signal_streams'] is not None: + dt = h['signal_streams'].dtype + for k, _ in _signal_stream_dtype: + assert k in dt.fields, f'{k} not in signal_streams.dtype' + assert 'signal_channels' in h, 'signal_channels missing in header' if h['signal_channels'] is not None: dt = h['signal_channels'].dtype for k, _ in _signal_channel_dtype: - assert k in dt.fields, '%s not in signal_channels.dtype' % k + assert k in dt.fields, f'{k} not in signal_channels.dtype' - assert 'unit_channels' in h, 'unit_channels missing in header' - if h['unit_channels'] is not None: - dt = h['unit_channels'].dtype - for k, _ in _unit_channel_dtype: - assert k in dt.fields, '%s not in unit_channels.dtype' % k + assert 'spike_channels' in h, 'spike_channels missing in header' + if h['spike_channels'] is not None: + dt = h['spike_channels'].dtype + for k, _ in _spike_channel_dtype: + assert k in dt.fields, f'{k} not in spike_channels.dtype' assert 'event_channels' in h, 'event_channels missing in header' if h['event_channels'] is not None: dt = h['event_channels'].dtype for k, _ in _event_channel_dtype: - assert k in dt.fields, '%s not in event_channels.dtype' % k + assert k in dt.fields, f'{k} not in event_channels.dtype' def count_element(reader): @@ -59,8 +72,10 @@ def count_element(reader): """ - nb_sig = reader.signal_channels_count() - nb_unit = reader.unit_channels_count() + nb_stream = reader.signal_streams_count() + for stream_index in range(nb_stream): + nb_chan = reader.signal_channels_count(stream_index) + nb_unit = reader.spike_channels_count() nb_event_channel = reader.event_channels_count() nb_block = reader.block_count() @@ -74,50 +89,37 @@ def count_element(reader): t_stop = reader.segment_t_stop(block_index=block_index, seg_index=seg_index) assert t_stop > t_start - if nb_sig > 0: - if reader._several_channel_groups: - channel_indexes_list = reader.get_group_signal_channel_indexes() - for channel_indexes in channel_indexes_list: - sig_size = reader.get_signal_size(block_index, seg_index, - channel_indexes=channel_indexes) - else: - sig_size = reader.get_signal_size(block_index, seg_index, - channel_indexes=None) - - for unit_index in range(nb_unit): + for stream_index in range(nb_stream): + sig_size = reader.get_signal_size(block_index, seg_index, + stream_index=stream_index) + + for spike_channel_index in range(nb_unit): nb_spike = reader.spike_count(block_index=block_index, seg_index=seg_index, - unit_index=unit_index) + spike_channel_index=spike_channel_index) for event_channel_index in range(nb_event_channel): nb_event = reader.event_count(block_index=block_index, seg_index=seg_index, event_channel_index=event_channel_index) -def iter_over_sig_chunks(reader, channel_indexes, chunksize=1024): - if channel_indexes is None: - nb_sig = reader.signal_channels_count() - else: - nb_sig = len(channel_indexes) - if nb_sig == 0: - return - +def iter_over_sig_chunks(reader, stream_index, channel_indexes, chunksize=1024): nb_block = reader.block_count() # read all chunk in RAW data - chunksize = 1024 for block_index in range(nb_block): nb_seg = reader.segment_count(block_index) for seg_index in range(nb_seg): - sig_size = reader.get_signal_size(block_index, seg_index, channel_indexes) + sig_size = reader.get_signal_size(block_index, seg_index, stream_index) nb = sig_size // chunksize + 1 for i in range(nb): i_start = i * chunksize i_stop = min((i + 1) * chunksize, sig_size) raw_chunk = reader.get_analogsignal_chunk(block_index=block_index, - seg_index=seg_index, - i_start=i_start, i_stop=i_stop, - channel_indexes=channel_indexes) + seg_index=seg_index, + i_start=i_start, i_stop=i_stop, + stream_index=stream_index, + channel_indexes=channel_indexes) yield raw_chunk @@ -128,94 +130,105 @@ def read_analogsignals(reader): Test special case when signal_channels do not have same sampling_rate. AKA _need_chan_index_check """ - nb_sig = reader.signal_channels_count() - if nb_sig == 0: + nb_stream = reader.signal_streams_count() + if nb_stream == 0: return - if reader._several_channel_groups: - channel_indexes_list = reader.get_group_signal_channel_indexes() - else: - channel_indexes_list = [None] + for stream_index in range(nb_stream): + sr = reader.get_signal_sampling_rate(stream_index=stream_index) + assert type(sr) == float, 'Type of sampling is {} should float'.format(type(sr)) - # read all chunk for all channel all block all segment - for channel_indexes in channel_indexes_list: - for raw_chunk in iter_over_sig_chunks(reader, channel_indexes, chunksize=1024): - assert raw_chunk.ndim == 2 - # ~ pass + # make other test on the first chunk of first block first block + block_index = 0 + seg_index = 0 - for channel_indexes in channel_indexes_list: - sr = reader.get_signal_sampling_rate(channel_indexes=channel_indexes) - assert type(sr) == float, 'Type of sampling is {} should float'.format(type(sr)) + sig_size = reader.get_signal_size(block_index, seg_index, stream_index) + + # read all chunk for all channel all block all segment + channel_indexes = None + for raw_chunk in iter_over_sig_chunks(reader, stream_index, + channel_indexes, chunksize=1024): + assert raw_chunk.ndim == 2 - # make other test on the first chunk of first block first block - block_index = 0 - seg_index = 0 - for channel_indexes in channel_indexes_list: i_start = 0 - sig_size = reader.get_signal_size(block_index, seg_index, - channel_indexes=channel_indexes) + sig_size = reader.get_signal_size(block_index, seg_index, stream_index) i_stop = min(1024, sig_size) - if channel_indexes is None: - nb_sig = reader.header['signal_channels'].size - channel_indexes = np.arange(nb_sig, dtype=int) - - all_signal_channels = reader.header['signal_channels'] - - signal_names = all_signal_channels['name'][channel_indexes] - signal_ids = all_signal_channels['id'][channel_indexes] + nb_chan = reader.signal_channels_count(stream_index) + channel_indexes = np.arange(nb_chan, dtype=int) - unique_chan_name = (np.unique(signal_names).size == all_signal_channels.size) - unique_chan_id = (np.unique(signal_ids).size == all_signal_channels.size) + signal_channels = reader.header['signal_channels'] + stream_id = reader.header['signal_streams'][stream_index]['id'] + mask = signal_channels['stream_id'] == stream_id + channel_names = signal_channels['name'][mask] + channel_ids = signal_channels['id'][mask] # acces by channel inde/ids/names should give the same chunk channel_indexes2 = channel_indexes[::2] - channel_names2 = signal_names[::2] - channel_ids2 = signal_ids[::2] + channel_names2 = channel_names[::2] + channel_ids2 = channel_ids[::2] + # slice by index raw_chunk0 = reader.get_analogsignal_chunk(block_index=block_index, seg_index=seg_index, - i_start=i_start, i_stop=i_stop, - channel_indexes=channel_indexes2) + i_start=i_start, i_stop=i_stop, + stream_index=stream_index, + channel_indexes=channel_indexes2) assert raw_chunk0.ndim == 2 assert raw_chunk0.shape[0] == i_stop assert raw_chunk0.shape[1] == len(channel_indexes2) + # slice by ids + raw_chunk2 = reader.get_analogsignal_chunk(block_index=block_index, seg_index=seg_index, + i_start=i_start, i_stop=i_stop, + stream_index=stream_index, + channel_ids=channel_ids2) + np.testing.assert_array_equal(raw_chunk0, raw_chunk2) + + # channel names are not always unique inside a stream + unique_chan_name = (np.unique(channel_names).size == channel_names.size) if unique_chan_name: raw_chunk1 = reader.get_analogsignal_chunk(block_index=block_index, seg_index=seg_index, - i_start=i_start, i_stop=i_stop, - channel_names=channel_names2) + i_start=i_start, i_stop=i_stop, + stream_index=stream_index, + channel_names=channel_names2) np.testing.assert_array_equal(raw_chunk0, raw_chunk1) - if unique_chan_id: - raw_chunk2 = reader.get_analogsignal_chunk(block_index=block_index, seg_index=seg_index, - i_start=i_start, i_stop=i_stop, - channel_ids=channel_ids2) - np.testing.assert_array_equal(raw_chunk0, raw_chunk2) + # test prefer_slice=True/False + if nb_chan >= 3: + for prefer_slice in (True, False): + raw_chunk3 = reader.get_analogsignal_chunk(block_index=block_index, + seg_index=seg_index, + i_start=i_start, i_stop=i_stop, + stream_index=stream_index, + channel_indexes=[1, 2]) # convert to float32/float64 for dt in ('float32', 'float64'): float_chunk0 = reader.rescale_signal_raw_to_float(raw_chunk0, dtype=dt, - channel_indexes=channel_indexes2) + stream_index=stream_index, + channel_indexes=channel_indexes2) + float_chunk2 = reader.rescale_signal_raw_to_float(raw_chunk2, dtype=dt, + stream_index=stream_index, + channel_ids=channel_ids2) if unique_chan_name: - float_chunk1 = reader.rescale_signal_raw_to_float(raw_chunk1, dtype=dt, - channel_names=channel_names2) - if unique_chan_id: - float_chunk2 = reader.rescale_signal_raw_to_float(raw_chunk2, dtype=dt, - channel_ids=channel_ids2) + float_chunk1 = reader.rescale_signal_raw_to_float(raw_chunk1, + dtype=dt, + stream_index=stream_index, + channel_names=channel_names2) assert float_chunk0.dtype == dt + np.testing.assert_array_equal(float_chunk0, float_chunk2) if unique_chan_name: np.testing.assert_array_equal(float_chunk0, float_chunk1) - if unique_chan_id: - np.testing.assert_array_equal(float_chunk0, float_chunk2) # read 500ms with several chunksize - sr = reader.get_signal_sampling_rate(channel_indexes=channel_indexes) + sr = reader.get_signal_sampling_rate(stream_index=stream_index) lenght_to_read = int(.5 * sr) if lenght_to_read < sig_size: ref_raw_sigs = reader.get_analogsignal_chunk(block_index=block_index, seg_index=seg_index, i_start=0, i_stop=lenght_to_read, + stream_index=stream_index, channel_indexes=channel_indexes) for chunksize in (511, 512, 513, 1023, 1024, 1025): i_start = 0 @@ -225,6 +238,7 @@ def read_analogsignals(reader): raw_chunk = reader.get_analogsignal_chunk(block_index=block_index, seg_index=seg_index, i_start=i_start, i_stop=i_stop, + stream_index=stream_index, channel_indexes=channel_indexes) chunks.append(raw_chunk) i_start += chunksize @@ -238,32 +252,28 @@ def benchmark_speed_read_signals(reader): in a file. """ - if reader._several_channel_groups: - channel_indexes_list = reader.get_group_signal_channel_indexes() - else: - channel_indexes_list = [None] + nb_stream = reader.signal_streams_count() + if nb_stream == 0: + return - for channel_indexes in channel_indexes_list: - if channel_indexes is None: - nb_sig = reader.signal_channels_count() - else: - nb_sig = len(channel_indexes) - if nb_sig == 0: - continue + for stream_index in range(nb_stream): nb_samples = 0 + channel_indexes = None + nb_chan = reader.signal_channels_count(stream_index) + t0 = time.perf_counter() - for raw_chunk in iter_over_sig_chunks(reader, channel_indexes, chunksize=1024): + for raw_chunk in iter_over_sig_chunks(reader, stream_index, + channel_indexes, chunksize=1024): nb_samples += raw_chunk.shape[0] t1 = time.perf_counter() if t0 != t1: - speed = (nb_samples * nb_sig) / (t1 - t0) / 1e6 + speed = (nb_samples * nb_chan) / (t1 - t0) / 1e6 else: speed = np.inf logging.info( - '{} read ({}signals x {}samples) in {:0.3f} s so speed {:0.3f} MSPS from {}'.format( - print_class(reader), - nb_sig, nb_samples, t1 - t0, speed, reader.source_name())) + f'{print_class(reader)} read ({nb_chan}channels x {nb_samples}samples)' + f'in {t1 - t0:0.3f} s so speed {speed:0.3f} MSPS from {reader.source_name()}') def read_spike_times(reader): @@ -272,21 +282,22 @@ def read_spike_times(reader): """ nb_block = reader.block_count() - nb_unit = reader.unit_channels_count() + nb_unit = reader.spike_channels_count() for block_index in range(nb_block): nb_seg = reader.segment_count(block_index) for seg_index in range(nb_seg): - for unit_index in range(nb_unit): + for spike_channel_index in range(nb_unit): nb_spike = reader.spike_count(block_index=block_index, - seg_index=seg_index, unit_index=unit_index) + seg_index=seg_index, + spike_channel_index=spike_channel_index) if nb_spike == 0: continue spike_timestamp = reader.get_spike_timestamps(block_index=block_index, - seg_index=seg_index, - unit_index=unit_index, t_start=None, - t_stop=None) + seg_index=seg_index, + spike_channel_index=spike_channel_index, + t_start=None, t_stop=None) assert spike_timestamp.shape[0] == nb_spike, 'nb_spike {} != {}'.format( spike_timestamp.shape[0], nb_spike) @@ -299,9 +310,9 @@ def read_spike_times(reader): t_stop = spike_times[1] + 0.001 spike_timestamp2 = reader.get_spike_timestamps(block_index=block_index, - seg_index=seg_index, - unit_index=unit_index, - t_start=t_start, t_stop=t_stop) + seg_index=seg_index, + spike_channel_index=spike_channel_index, + t_start=t_start, t_stop=t_stop) assert spike_timestamp2.shape[0] == 1 spike_times2 = reader.rescale_spike_timestamp(spike_timestamp2, 'float64') @@ -313,21 +324,22 @@ def read_spike_waveforms(reader): Read and convert some all waveforms. """ nb_block = reader.block_count() - nb_unit = reader.unit_channels_count() + nb_unit = reader.spike_channels_count() for block_index in range(nb_block): nb_seg = reader.segment_count(block_index) for seg_index in range(nb_seg): - for unit_index in range(nb_unit): + for spike_channel_index in range(nb_unit): nb_spike = reader.spike_count(block_index=block_index, - seg_index=seg_index, unit_index=unit_index) + seg_index=seg_index, + spike_channel_index=spike_channel_index) if nb_spike == 0: continue raw_waveforms = reader.get_spike_raw_waveforms(block_index=block_index, - seg_index=seg_index, - unit_index=unit_index, - t_start=None, t_stop=None) + seg_index=seg_index, + spike_channel_index=spike_channel_index, + t_start=None, t_stop=None) if raw_waveforms is None: continue assert raw_waveforms.shape[0] == nb_spike @@ -335,7 +347,7 @@ def read_spike_waveforms(reader): for dt in ('float32', 'float64'): float_waveforms = reader.rescale_waveforms_to_float( - raw_waveforms, dtype=dt, unit_index=unit_index) + raw_waveforms, dtype=dt, spike_channel_index=spike_channel_index) assert float_waveforms.dtype == dt assert float_waveforms.shape == raw_waveforms.shape diff --git a/neo/test/rawiotest/test_blackrockrawio.py b/neo/test/rawiotest/test_blackrockrawio.py index c061ea2b3..9ccbbbf90 100644 --- a/neo/test/rawiotest/test_blackrockrawio.py +++ b/neo/test/rawiotest/test_blackrockrawio.py @@ -66,16 +66,18 @@ def test_compare_blackrockio_with_matlabloader(self): reader.parse_header() # Check if analog data on channels 1-8 are equal - self.assertGreater(reader.signal_channels_count(), 0) + stream_index = 0 + self.assertGreater(reader.signal_channels_count(stream_index), 0) for c in range(0, 8): - raw_sigs = reader.get_analogsignal_chunk(channel_indexes=[c]) + raw_sigs = reader.get_analogsignal_chunk(channel_indexes=[c], + stream_index=stream_index) raw_sigs = raw_sigs.flatten() assert_equal(raw_sigs[:-1], lfp_ml[c, :]) # Check if spikes in channels are equal - nb_unit = reader.unit_channels_count() - for unit_index in range(nb_unit): - unit_name = reader.header['unit_channels'][unit_index]['name'] + nb_unit = reader.spike_channels_count() + for spike_channel_index in range(nb_unit): + unit_name = reader.header['spike_channels'][spike_channel_index]['name'] # name is chXX#YY where XX is channel_id and YY is unit_id channel_id, unit_id = unit_name.split('#') channel_id = int(channel_id.replace('ch', '')) @@ -83,12 +85,13 @@ def test_compare_blackrockio_with_matlabloader(self): matlab_spikes = ts_ml[(elec_ml == channel_id) & (unit_ml == unit_id)] - io_spikes = reader.get_spike_timestamps(unit_index=unit_index) + io_spikes = reader.get_spike_timestamps(spike_channel_index=spike_channel_index) assert_equal(io_spikes, matlab_spikes) # Check waveforms of channel 1, unit 0 if channel_id == 1 and unit_id == 0: - io_waveforms = reader.get_spike_raw_waveforms(unit_index=unit_index) + io_waveforms = reader.get_spike_raw_waveforms( + spike_channel_index=spike_channel_index) io_waveforms = io_waveforms[:, 0, :] # remove dim 1 assert_equal(io_waveforms, wf_ml) @@ -142,7 +145,8 @@ def test_compare_blackrockio_with_matlabloader_v21(self): reader.parse_header() # Check if analog data are equal - self.assertGreater(reader.signal_channels_count(), 0) + stream_index = 0 + self.assertGreater(reader.signal_channels_count(stream_index), 0) for c in range(0, param[2]): raw_sigs = reader.get_analogsignal_chunk(channel_indexes=[c]) @@ -150,9 +154,9 @@ def test_compare_blackrockio_with_matlabloader_v21(self): assert_equal(raw_sigs[:], lfp_ml[c, :]) # Check if spikes in channels are equal - nb_unit = reader.unit_channels_count() - for unit_index in range(nb_unit): - unit_name = reader.header['unit_channels'][unit_index]['name'] + nb_unit = reader.spike_channels_count() + for spike_channel_index in range(nb_unit): + unit_name = reader.header['spike_channels'][spike_channel_index]['name'] # name is chXX#YY where XX is channel_id and YY is unit_id channel_id, unit_id = unit_name.split('#') channel_id = int(channel_id.replace('ch', '')) @@ -160,11 +164,12 @@ def test_compare_blackrockio_with_matlabloader_v21(self): matlab_spikes = ts_ml[(elec_ml == channel_id) & (unit_ml == unit_id)] - io_spikes = reader.get_spike_timestamps(unit_index=unit_index) + io_spikes = reader.get_spike_timestamps(spike_channel_index=spike_channel_index) assert_equal(io_spikes, matlab_spikes) # Check all waveforms - io_waveforms = reader.get_spike_raw_waveforms(unit_index=unit_index) + io_waveforms = reader.get_spike_raw_waveforms( + spike_channel_index=spike_channel_index) io_waveforms = io_waveforms[:, 0, :] # remove dim 1 matlab_wf = wf_ml[np.nonzero( np.logical_and(elec_ml == channel_id, unit_ml == unit_id)), :][0] diff --git a/neo/test/rawiotest/test_examplerawio.py b/neo/test/rawiotest/test_examplerawio.py index 84c4e9544..546ff8a5e 100644 --- a/neo/test/rawiotest/test_examplerawio.py +++ b/neo/test/rawiotest/test_examplerawio.py @@ -33,11 +33,10 @@ class TestExampleRawIO(BaseTestRawIO, unittest.TestCase, ): rawioclass = ExampleRawIO # here obsvisously there is nothing to download: files_to_download = [] - # here we will test 2 fake files - # not that IO base on dirname you can put the dirname here. - entities_to_test = ['fake1', - 'fake2', - ] + # here we will test 1 fake files + # note that for IOs based on directory names you can put the directory + # name here instead of the file name. + entities_to_test = ['fake1'] if __name__ == "__main__": diff --git a/neo/test/test_utils.py b/neo/test/test_utils.py index 7ad0bedd8..37eaa833e 100644 --- a/neo/test/test_utils.py +++ b/neo/test/test_utils.py @@ -561,11 +561,11 @@ def test__cut_block_by_epochs(self): seg = Segment() proxy_anasig = AnalogSignalProxy(rawio=self.reader, - global_channel_indexes=None, - block_index=0, seg_index=0) + stream_index=0, inner_stream_channels=None, + block_index=0, seg_index=0) seg.analogsignals.append(proxy_anasig) - proxy_st = SpikeTrainProxy(rawio=self.reader, unit_index=0, + proxy_st = SpikeTrainProxy(rawio=self.reader, spike_channel_index=0, block_index=0, seg_index=0) seg.spiketrains.append(proxy_st)