Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
155 changes: 74 additions & 81 deletions extract_spk_with_axisfile_matlab.m
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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)));
Expand All @@ -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
12 changes: 11 additions & 1 deletion rasterplot.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = (
Expand Down