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
234 changes: 115 additions & 119 deletions extract_spk_with_axisfile_octave.m
Original file line number Diff line number Diff line change
Expand Up @@ -3,9 +3,12 @@
% Checkout the feature/octave branch from git@github.com:ar-jan/AxionFileLoader.git
% in the vendored AxionFileLoader, and select Octave in the process_spk notebook.

function extract_spk_with_axisfile_octave(spk_path, output_csv, loader_dir)
%EXTRACT_SPK_WITH_AXISFILE_OCTAVE Convert one Axion .spk file to CSV.
% extract_spk_with_axisfile_octave(spk_path, output_csv, loader_dir)
function output_csv_path = extract_spk_with_axisfile_octave(spk_path, output_csv, loader_dir)
%EXTRACT_SPK_WITH_AXISFILE_OCTAVE Extract spike timings from one Axion .spk file.
% output_csv_path = extract_spk_with_axisfile_octave(spk_path, output_csv, loader_dir)
% loads spike data with the Octave-compatible AxionFileLoader API,
% 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 All @@ -18,153 +21,146 @@ function extract_spk_with_axisfile_octave(spk_path, output_csv, loader_dir)
error('extract_spk_with_axisfile_octave:MissingInput', 'spk_path is required.');
end

if ~(ischar(spk_path) && ~isempty(spk_path))
error('extract_spk_with_axisfile_octave:InvalidPath', ...
'spk_path must be a non-empty character vector.');
end

if nargin < 2 || isempty(output_csv)
[spk_parent, spk_name, ~] = fileparts(spk_path);
output_csv = fullfile(spk_parent, [spk_name '.csv']);
output_csv_path = fullfile(spk_parent, [spk_name '.csv']);
else
output_csv_path = output_csv;
end

if nargin < 3 || isempty(loader_dir)
this_dir = fileparts(mfilename('fullpath'));
loader_dir = fullfile(this_dir, 'vendor', 'AxionFileLoader', 'AxionFileLoader');
script_dir = fileparts(mfilename('fullpath'));
loader_dir = fullfile(script_dir, 'vendor', 'AxionFileLoader', 'AxionFileLoader');
end

if exist(loader_dir, 'dir') ~= 7
error('extract_spk_with_axisfile_octave:MissingLoader', ...
'AxionFileLoader directory not found: %s', loader_dir);
end

run_timer = tic;
addpath(loader_dir);
log_progress('Starting Octave SPK extraction.');
log_progress(sprintf('Input SPK: %s', spk_path));
log_progress(sprintf('Output CSV: %s', output_csv_path));
log_progress(sprintf('AxionFileLoader directory: %s', loader_dir));

if exist(spk_path, 'file') ~= 2
error('extract_spk_with_axisfile_octave:MissingFile', 'SPK file not found: %s', spk_path);
error('extract_spk_with_axisfile_octave:MissingFile', ...
'SPK file not found: %s', spk_path);
end

if exist(loader_dir, 'dir') ~= 7
error('extract_spk_with_axisfile_octave:MissingLoader', 'AxionFileLoader directory not found: %s', loader_dir);
[output_parent, ~, ~] = fileparts(output_csv_path);
if ~isempty(output_parent) && exist(output_parent, 'dir') ~= 7
mkdir(output_parent);
end

% This addpath is expected to point at the vendored feature/octave
% AxionFileLoader tree. That branch adds Octave compatibility shims and
% the LoadAllSpikesDetailed() helper used below.
addpath(loader_dir);
[~, ~, spk_ext] = fileparts(spk_path);
if ~strcmpi(spk_ext, '.spk')
error('extract_spk_with_axisfile_octave:InvalidExtension', ...
'Input file must have a .spk extension: %s', spk_path);
end

out_parent = fileparts(output_csv);
if ~isempty(out_parent) && exist(out_parent, 'dir') ~= 7
mkdir(out_parent);
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_octave:UnexpectedSpikeDataSets', ...
'Expected one spike dataset, found %d.', numel(spike_data));
end

spike_dataset = AxisFile(spk_path).SpikeData;
spike_rows = struct( ...
'WellRow', [], ...
'WellColumn', [], ...
'ElectrodeColumn', [], ...
'ElectrodeRow', [], ...
'WaveformStartTime', [], ...
'SpikeTime', [], ...
'MaximumAmplitude', [], ...
'MinimumAmplitude', [], ...
'PeakToPeakAmplitude', []);

if ~isempty(spike_dataset)
for ds_idx = 1:numel(spike_dataset)
% Octave-specific dependency:
% the original Axion MATLAB example uses
% SpikeData.LoadData + waveform.GetTimeVoltageVector().
% This wrapper instead depends on the vendored
% LoadAllSpikesDetailed() method, which was added in the
% feature/octave loader branch to avoid MATLAB-only waveform
% object construction paths while preserving equivalent fields.
current_rows = spike_dataset(ds_idx).LoadAllSpikesDetailed();
if isempty(spike_rows.WaveformStartTime)
spike_rows = current_rows;
elseif ~isempty(current_rows.WaveformStartTime)
spike_rows.WellRow = [spike_rows.WellRow, current_rows.WellRow];
spike_rows.WellColumn = [spike_rows.WellColumn, current_rows.WellColumn];
spike_rows.ElectrodeColumn = [spike_rows.ElectrodeColumn, current_rows.ElectrodeColumn];
spike_rows.ElectrodeRow = [spike_rows.ElectrodeRow, current_rows.ElectrodeRow];
spike_rows.WaveformStartTime = [spike_rows.WaveformStartTime, current_rows.WaveformStartTime];
spike_rows.SpikeTime = [spike_rows.SpikeTime, current_rows.SpikeTime];
spike_rows.MaximumAmplitude = [spike_rows.MaximumAmplitude, current_rows.MaximumAmplitude];
spike_rows.MinimumAmplitude = [spike_rows.MinimumAmplitude, current_rows.MinimumAmplitude];
spike_rows.PeakToPeakAmplitude = [spike_rows.PeakToPeakAmplitude, current_rows.PeakToPeakAmplitude];
end
[hardware_channels, timestamp] = spike_data.LoadAllSpikes();
timestamp = timestamp(:);
total_rows = numel(timestamp);
log_progress(sprintf('Loaded %d trigger-aligned spikes.', total_rows));

channel_label = zeros(0, 1);
well_label = cell(0, 1);
if total_rows > 0
if numel(hardware_channels.Achk) ~= total_rows || ...
numel(hardware_channels.Channel) ~= total_rows
error('extract_spk_with_axisfile_octave:SpikeDataLengthMismatch', ...
'LoadAllSpikes returned inconsistent channel and timestamp lengths.');
end
end

final_results = struct();
if ~isempty(spike_rows.WaveformStartTime)
spike_rows = sort_spike_rows_for_csv(spike_rows);

electrode_digits = floor(log10(spike_rows.ElectrodeRow)) + 1;
channel_label = spike_rows.ElectrodeColumn .* (10 .^ electrode_digits) + spike_rows.ElectrodeRow;
well_label = arrayfun(@well_label_from_indices, ...
spike_rows.WellRow(:), ...
spike_rows.WellColumn(:), ...
'UniformOutput', false);

% Keep the CSV headers aligned with the example: use "timestamp"
final_results.Channel_Label = channel_label(:);
final_results.Well_Label = well_label(:);
final_results.Timestamp = spike_rows.WaveformStartTime(:);
final_results.Maximum_Amplitude = spike_rows.MaximumAmplitude(:);
final_results.Minimum_Amplitude = spike_rows.MinimumAmplitude(:);
final_results.Peak_to_peak_Amplitude = spike_rows.PeakToPeakAmplitude(:);
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 ...
);
log_progress(sprintf( ...
'Mapped %d spikes across %d hardware channels.', ...
total_rows, unique_channel_count));
end

fid = fopen(output_csv, 'w');
log_progress(sprintf('Writing CSV table with %d rows.', total_rows));
fid = fopen(output_csv_path, 'w');
if fid < 0
error('extract_spk_with_axisfile_octave:OpenOutput', 'Unable to open output CSV: %s', output_csv);
error('extract_spk_with_axisfile_octave:OpenOutput', ...
'Unable to open output CSV: %s', output_csv_path);
end

fprintf(fid, '%s,%s,%s,%s,%s,%s\n', ...
fprintf(fid, '%s,%s,%s\n', ...
'Channel_Label', ...
'Well_Label', ...
'Timestamp', ...
'Maximum_Amplitude', ...
'Minimum_Amplitude', ...
'Peak_to_peak_Amplitude');

if isfield(final_results, 'Timestamp')
for row_idx = 1:numel(final_results.Timestamp)
fprintf(fid, '%d,%s,%.17g,%.17g,%.17g,%.17g\n', ...
final_results.Channel_Label(row_idx), ...
final_results.Well_Label{row_idx}, ...
final_results.Timestamp(row_idx), ...
final_results.Maximum_Amplitude(row_idx), ...
final_results.Minimum_Amplitude(row_idx), ...
final_results.Peak_to_peak_Amplitude(row_idx));
end
'Timestamp');

for row_idx = 1:total_rows
fprintf(fid, '%d,%s,%.15g\n', ...
channel_label(row_idx), ...
well_label{row_idx}, ...
timestamp(row_idx));
end
fclose(fid);

fprintf('Wrote CSV: %s\n', output_csv);
log_progress(sprintf('Completed extraction in %.1fs.', toc(run_timer)));
fprintf('Wrote CSV: %s\n', output_csv_path);
end

function sorted_rows = sort_spike_rows_for_csv(spike_rows)
%SORT_SPIKE_ROWS_FOR_CSV Recreate the MATLAB script row ordering.
% The reference scripts iterate in nested WellRow/WellColumn/
% ElectrodeColumn/ElectrodeRow order and then emit spikes within each
% waveform group by waveform start time. SpikeTime is kept as a final
% tie-breaker so the sort stays deterministic if two spikes share the
% same start time.

if isempty(spike_rows.WaveformStartTime)
sorted_rows = spike_rows;
return;
end

sort_keys = [ ...
spike_rows.WellRow(:), ...
spike_rows.WellColumn(:), ...
spike_rows.ElectrodeColumn(:), ...
spike_rows.ElectrodeRow(:), ...
spike_rows.WaveformStartTime(:), ...
spike_rows.SpikeTime(:) ...
];
[~, sort_idx] = sortrows(sort_keys, [1 2 3 4 5 6]);
sort_idx = reshape(sort_idx, 1, []);

sorted_rows = spike_rows;
field_names = fieldnames(spike_rows);
for field_idx = 1:numel(field_names)
field_name = field_names{field_idx};
sorted_rows.(field_name) = spike_rows.(field_name)(sort_idx);
end
function log_progress(message)
fprintf('[%s] %s\n', datestr(now, 'yyyy-mm-dd HH:MM:SS'), message);
fflush(stdout);
end

function label = well_label_from_indices(well_row, well_column)
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
5 changes: 4 additions & 1 deletion process_spk.py
Original file line number Diff line number Diff line change
Expand Up @@ -65,6 +65,7 @@ def run_logged_subprocess(

def octave_loader_is_compatible(loader_dir: Path) -> bool:
spike_dataset_file = loader_dir / "SpikeDataSet.m"
axion_empty_helper = loader_dir / "axion_empty.m"
heterogeneous_shim = loader_dir / "+matlab" / "+mixin" / "Heterogeneous.m"
custom_display_shim = loader_dir / "+matlab" / "+mixin" / "CustomDisplay.m"

Expand All @@ -79,7 +80,9 @@ def octave_loader_is_compatible(loader_dir: Path) -> bool:
return False

return (
"LoadAllSpikesDetailed" in spike_dataset_text
"function [aElectrodes, aTimes] = LoadAllSpikes" in spike_dataset_text
and "BuildMappedDataWithoutMemmap" in spike_dataset_text
and axion_empty_helper.is_file()
and heterogeneous_shim.is_file()
and custom_display_shim.is_file()
)
Expand Down