diff --git a/extract_spk_with_axisfile_matlab.m b/extract_spk_with_axisfile_matlab.m index 6f0c640..4d54329 100644 --- a/extract_spk_with_axisfile_matlab.m +++ b/extract_spk_with_axisfile_matlab.m @@ -2,7 +2,8 @@ %EXTRACT_SPK_WITH_AXISFILE_MATLAB Extract spike timings from one Axion .spk file. % output_csv_path = extract_spk_with_axisfile_matlab(spk_path, output_csv, loader_dir) % loads spike data with the AxionFileLoader MATLAB API, -% extracts timing and amplitude columns and writes a .csv file. +% maps hardware channels to well/electrode coordinates, and writes +% trigger-aligned spike timings to a .csv file. % % - spk_path: path to input .spk file % - output_csv: path to output .csv file (optional; defaults to spk basename) @@ -68,83 +69,77 @@ 'Input file must have a .spk extension: %s', spk_path); end - log_progress('Loading spike data via AxisFile(...).SpikeData.LoadData ...'); - all_data = AxisFile(spk_path).SpikeData.LoadData; - [nwr, nwc, nec, ner] = size(all_data); - result_blocks = cell(0, 1); - total_groups = nwr * nwc * nec * ner; - scanned_groups = 0; - nonempty_groups = 0; - total_rows = 0; - - log_progress(sprintf( ... - 'Loaded spike grid: WellRows=%d, WellColumns=%d, ElectrodeColumns=%d, ElectrodeRows=%d (%d groups).', ... - nwr, nwc, nec, ner, total_groups)); - - for wr = 1:nwr - for wc = 1:nwc - log_progress(sprintf( ... - 'Scanning well %s (%d of %d).', ... - well_label_from_indices(wr, wc), ... - ((wr - 1) * nwc) + wc, ... - nwr * nwc)); - for ec = 1:nec - for er = 1:ner - scanned_groups = scanned_groups + 1; - data = all_data{wr, wc, ec, er}; - if isempty(data) - if mod(scanned_groups, 250) == 0 || scanned_groups == total_groups - log_progress(sprintf( ... - 'Scanned %d/%d groups; non-empty=%d; rows=%d; elapsed=%.1fs.', ... - scanned_groups, total_groups, nonempty_groups, total_rows, toc(run_timer))); - end - continue - end - - [t, v] = data.GetTimeVoltageVector; - timestamp = t(1, :); - timestamp_length = length(timestamp); - nonempty_groups = nonempty_groups + 1; - total_rows = total_rows + timestamp_length; - - channel_label = repmat(str2double(strcat(num2str(ec), num2str(er))), timestamp_length, 1); - well_label = repmat({well_label_from_indices(wr, wc)}, timestamp_length, 1); - timestamp = timestamp(:); - min_amplitude = min(v, [], 1)'; - max_amplitude = max(v, [], 1)'; - peak_to_peak_amplitude = max_amplitude - min_amplitude; - - result_blocks{end + 1, 1} = table( ... - channel_label, ... - well_label, ... - timestamp, ... - max_amplitude, ... - min_amplitude, ... - peak_to_peak_amplitude, ... - 'VariableNames', { ... - 'Channel_Label', ... - 'Well_Label', ... - 'Timestamp', ... - 'Maximum_Amplitude', ... - 'Minimum_Amplitude', ... - 'Peak_to_peak_Amplitude' ... - } ... - ); - - if nonempty_groups <= 5 || mod(nonempty_groups, 25) == 0 || scanned_groups == total_groups - log_progress(sprintf( ... - 'Processed group wr=%d wc=%d ec=%d er=%d; spikes=%d; non-empty=%d; rows=%d; elapsed=%.1fs.', ... - wr, wc, ec, er, timestamp_length, nonempty_groups, total_rows, toc(run_timer))); - end - end - end - end + log_progress('Loading trigger-aligned spikes via AxisFile(...).SpikeData.LoadAllSpikes ...'); + axis_file = AxisFile(spk_path); + spike_data = axis_file.SpikeData; + if numel(spike_data) ~= 1 + error('extract_spk_with_axisfile_matlab:UnexpectedSpikeDataSets', ... + 'Expected one spike dataset, found %d.', numel(spike_data)); end - if isempty(result_blocks) + [hardware_channels, timestamp] = spike_data.LoadAllSpikes; + timestamp = timestamp(:); + total_rows = numel(timestamp); + log_progress(sprintf('Loaded %d trigger-aligned spikes.', total_rows)); + + if total_rows == 0 final_results = create_empty_result_table(); else - final_results = vertcat(result_blocks{:}); + if numel(hardware_channels.Achk) ~= total_rows || ... + numel(hardware_channels.Channel) ~= total_rows + error('extract_spk_with_axisfile_matlab:SpikeDataLengthMismatch', ... + 'LoadAllSpikes returned inconsistent channel and timestamp lengths.'); + end + + hardware_keys = bitor( ... + bitshift(uint16(hardware_channels.Achk(:)), 8), ... + uint16(hardware_channels.Channel(:)) ... + ); + [~, first_channel_indices, channel_group_indices] = unique(hardware_keys); + unique_channel_count = numel(first_channel_indices); + log_progress(sprintf( ... + 'Mapping %d unique hardware channels through ChannelArray.', ... + unique_channel_count)); + + channel_mappings = spike_data.ChannelArray.LookupChannelMapping( ... + hardware_channels.Achk(first_channel_indices), ... + hardware_channels.Channel(first_channel_indices) ... + ); + mapped_well_rows = double([channel_mappings.WellRow]).'; + mapped_well_columns = double([channel_mappings.WellColumn]).'; + mapped_electrode_columns = double([channel_mappings.ElectrodeColumn]).'; + mapped_electrode_rows = double([channel_mappings.ElectrodeRow]).'; + + well_rows = mapped_well_rows(channel_group_indices); + well_columns = mapped_well_columns(channel_group_indices); + electrode_columns = mapped_electrode_columns(channel_group_indices); + electrode_rows = mapped_electrode_rows(channel_group_indices); + + well_label = arrayfun( ... + @well_label_from_indices, ... + well_rows, ... + well_columns, ... + 'UniformOutput', false ... + ); + channel_label = arrayfun( ... + @channel_label_from_indices, ... + electrode_columns, ... + electrode_rows ... + ); + + final_results = table( ... + channel_label(:), ... + well_label(:), ... + timestamp, ... + 'VariableNames', { ... + 'Channel_Label', ... + 'Well_Label', ... + 'Timestamp' ... + } ... + ); + log_progress(sprintf( ... + 'Mapped %d spikes across %d hardware channels.', ... + total_rows, unique_channel_count)); end log_progress(sprintf('Writing CSV table with %d rows.', height(final_results))); @@ -162,21 +157,19 @@ function log_progress(message) label = sprintf('%s%d', char(double('A') + well_row - 1), well_column); end +function label = channel_label_from_indices(electrode_column, electrode_row) + label = str2double(sprintf('%d%d', electrode_column, electrode_row)); +end + function empty_table = create_empty_result_table() empty_table = table( ... zeros(0, 1), ... cell(0, 1), ... zeros(0, 1), ... - zeros(0, 1), ... - zeros(0, 1), ... - zeros(0, 1), ... 'VariableNames', { ... 'Channel_Label', ... 'Well_Label', ... - 'Timestamp', ... - 'Maximum_Amplitude', ... - 'Minimum_Amplitude', ... - 'Peak_to_peak_Amplitude' ... + 'Timestamp' ... } ... ); end diff --git a/rasterplot.py b/rasterplot.py index 350712e..0588a07 100644 --- a/rasterplot.py +++ b/rasterplot.py @@ -85,7 +85,17 @@ def combine_available_wells(*well_lists: list[str]) -> list[str]: if well_label not in seen: well_labels.append(well_label) seen.add(well_label) - return well_labels + + def _well_sort_key(well_label: str) -> tuple[str, float, str]: + normalized_label = well_label.strip() + row_label = normalized_label.rstrip("0123456789") + column_label = normalized_label[len(row_label) :] + column_number = ( + float(column_label) if column_label.isdigit() else float("inf") + ) + return row_label.casefold(), column_number, normalized_label.casefold() + + return sorted(well_labels, key=_well_sort_key) def available_channels(df: pl.DataFrame) -> list[str]: channels = (