OSCR

Developmental molecular signatures define de novo cortico-brainstem circuit for skilled forelimb movement.

Code ↔ Paper

5 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

The 5 matches
  1. [1] § Method › Single pellet reaching task ↔ ReachingClassificationGUI_reach_hm2.m, lines 1–60 · score 0.61 · single pellet reaching, 1–5 days, classified, offset, weights, front
  2. [2] § Method › 3D reconstruction of SCPN injection volumes ↔ MATLAB/VOL3D_Step2_GUI.m, lines 2283–2424 · score 0.57 · selected ABA structures, voxel, VOL3D, overlap, volumes, atlas
  3. [3] § Method › Quantification of retrogradely labeled SCPN ↔ MATLAB/CELL3D_Step1_Animal_GUI.m, lines 1841–1940 · score 0.54 · CELL3D, Cell coordinates, transformed, cingulate, FIJI, spaced
  4. [4] § Method › Kinematic analysis of skilled reaching ↔ ReachingClassificationGUI_reach_hm2.m, lines 3406–3489 · score 0.54 · pellet centered, accuracy, acceleration, metrics, contacted, retrieval
  5. [5] § Method › Kinematic analysis of skilled reaching ↔ ReachingClassificationGUI_reach_hm2.m, lines 3975–4011 · score 0.54 · pellet contact, pause, digits, tracked, likelihood, frames

Paper

Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC

The paper is loaded when this pane is shown.

The authors' code

MATLAB · 5,438 lines · 182 KB · no license · 3 matches

  1. function ReachingClassificationGUI_reach()
  2. % Create main UI
  3. handles = struct();
  4. handles.baseDir = pwd;
  5. handles.colors.errorRed = [0.8 0.2 0.2]; % For error messages
  6. handles.colors.statusPending = [0.62 0.64 0.70];
  7. handles.colors.successColor = [0.35 0.75 0.45]; % success green #59BF73
  8. outDir = fullfile(handles.baseDir,'OUT');
  9. if ~exist(outDir) mkdir(outDir); end
  10. fig = uifigure('Name', 'Single Pellet Reaching Processing', ...
  11. 'Position', [100, 100, 1000, 600]);
  12. % Define simple grid layout
  13. gl = uigridlayout(fig, [4,5]);
  14. gl.RowHeight = {30, '1x',50, 50, 100};
  15. gl.ColumnWidth = {'1x', '1x', '1x', '1x'};
  16. % --- Row 1: Folder Selection + Animal Dropdown + Alignment Button ---
  17. headerLabel = uilabel(gl, 'Text', 'Single Pellet Reaching Assessment', ...
  18. 'FontSize', 18, ...
  19. 'FontWeight', 'bold', ...
  20. 'HorizontalAlignment', 'left');
  21. headerLabel.Layout.Row = 1;
  22. headerLabel.Layout.Column = [1 5];
  23. [~,lblFolder] = fileparts(handles.baseDir); % stores foldername
  24. % Add table to show video pair status
  25. tblStatus = uitable(gl, ...
  26. 'Data', {}, ...
  27. 'FontSize', 12, ...
  28. 'ColumnName', { ...
  29. 'coreID', ...
  30. 'SideVideo', ...
  31. 'FrontVideo', ...
  32. 'SideDLC', ...
  33. 'FrontDLC', ...
  34. 'OffsetStart', ...
  35. 'OffsetEnd', ...
  36. 'TotalReaches', ...
  37. 'Exclude'}, ...
  38. 'ColumnFormat', { ...
  39. 'char', ...
  40. 'logical', ...
  41. 'logical', ...
  42. 'logical', ...
  43. 'logical', ...
  44. 'numeric', ...
  45. 'numeric', ...
  46. 'numeric', ...
  47. 'logical'}, ...
  48. 'ColumnEditable', [false false false false false false false false true]);
  49. tblStatus.Layout.Row = 2;
  50. tblStatus.Layout.Column = [1 5];
  51. % New row for likelihood input above buttons
  52. likelihoodLabel = uilabel(gl, ...
  53. 'Text', 'Paw Likelihood Threshold:', ...
  54. 'HorizontalAlignment', 'right', ...
  55. 'FontWeight', 'bold', ...
  56. 'Tooltip', 'Set minimum likelihood threshold for paw detection (0 to 1)', ...
  57. 'FontSize', 12);
  58. likelihoodLabel.Layout.Row = 3;
  59. likelihoodLabel.Layout.Column = 1;
  60. defaultLikelihood = getLatestPawLikelihoodFromLog(handles.baseDir);
  61. likelihoodEdit = uieditfield(gl, 'numeric', ...
  62. 'Limits', [0 1], ...
  63. 'Value', defaultLikelihood, ... % default value
  64. 'RoundFractionalValues', false, ...
  65. 'Tooltip', 'Enter a value between 0 (low) and 1 (high likelihood)', ...
  66. 'FontSize', 12);
  67. likelihoodEdit.Layout.Row = 3;
  68. likelihoodEdit.Layout.Column = 2;
  69. % Store in handles for access in callbacks
  70. handles.likelihoodEdit = likelihoodEdit;
  71. guidata(fig, handles);
  72. btnAlignVideos = uibutton(gl, 'Text', '🛠 Calculate Offset (front/side)');
  73. btnAlignVideos.Layout.Row = 4;
  74. btnAlignVideos.Layout.Column = 1;
  75. btnCalibratePole = uibutton(gl, 'Text', '🛠 Calculate Pole Width (Calibrate)');
  76. btnCalibratePole.Layout.Row = 4;
  77. btnCalibratePole.Layout.Column = 2;
  78. btnDetectReaches = uibutton(gl, 'Text', '🐾 Detect Reaches');
  79. btnDetectReaches.Layout.Row = 4;
  80. btnDetectReaches.Layout.Column = 3;
  81. btnClassifyReaches = uibutton(gl, 'Text', '🎯 Classify Reaches');
  82. btnClassifyReaches.Layout.Row = 4;
  83. btnClassifyReaches.Layout.Column = 4;
  84. btnAnalyzeReaches = uibutton(gl, 'Text', '📊 Analyze Paw Kinematics');
  85. btnAnalyzeReaches.Layout.Row = 4;
  86. btnAnalyzeReaches.Layout.Column = 5;
  87. msgLayout = uigridlayout(gl, [1, 1]); % One cell grid layout
  88. msgLayout.RowHeight = {'1x'};
  89. msgLayout.ColumnWidth = {'1x'};
  90. msgLayout.Padding = [0 0 0 0]; % No padding
  91. msgLayout.Layout.Row = 5;
  92. msgLayout.Layout.Column = [1 5];
  93. msgLabel = uilabel(msgLayout, ...
  94. 'Text', 'No errors in the setup', ...
  95. 'FontWeight', 'bold', ...
  96. 'FontSize', 13, ...
  97. 'HorizontalAlignment', 'center', ... % Center horizontally
  98. 'VerticalAlignment', 'center', ... % Center vertically
  99. 'FontColor', '#000000', ...
  100. 'BackgroundColor', '#ffffff', ...
  101. 'WordWrap', 'on');
  102. % --- Store Handles for Later Use ---
  103. handles.fig = fig;
  104. handles.outDir = outDir;
  105. % handles.btnFolder = btnFolder;
  106. handles.btnAlignVideos = btnAlignVideos;
  107. handles.btnCalibratePole = btnCalibratePole;
  108. handles.btnDetectReaches = btnDetectReaches;
  109. handles.btnClassifyReaches = btnClassifyReaches;
  110. handles.btnAnalyzeReaches = btnAnalyzeReaches;
  111. handles.lblFolder = lblFolder;
  112. handles.tblStatus = tblStatus;
  113. handles.msgLabel = msgLabel;
  114. guidata(fig, handles);
  115. % --- Callbacks ---
  116. btnFolder.ButtonPushedFcn = @(src, evt) selectFolderCallback(fig);
  117. % btnConvertVideos.ButtonPushedFcn = @(src, evt) precomputeStacksCallback(fig);
  118. btnAlignVideos.ButtonPushedFcn = @(src, evt) alignVideosCallback(fig);
  119. btnCalibratePole.ButtonPushedFcn = @(src,evt) calibratePoleByClick(fig);
  120. btnDetectReaches.ButtonPushedFcn = @(src, evt) detectReachesCallback(fig);
  121. btnClassifyReaches.ButtonPushedFcn = @(src, evt) classifyReachesCallback(fig);
  122. btnAnalyzeReaches.ButtonPushedFcn = @(src, evt) analyzeReachesCallback(fig);
  123. handles.tblStatus.CellEditCallback = @(src, evt) excludeCellEditCallback(fig, src, evt);
  124. % Immediately try loading current folder
  125. LoadFolder(fig);
  126. end
  127. %% GUI table and interaction functions
  128. function LoadFolder(fig)
  129. handles = guidata(fig);
  130. % Cache directories
  131. baseDir = handles.baseDir;
  132. outDir = handles.outDir;
  133. sideDir = fullfile(baseDir, 'Side');
  134. frontDir = fullfile(baseDir, 'Front');
  135. % Show status
  136. handles.msgLabel.Text = 'Scanning video files...';
  137. handles.msgLabel.FontColor = handles.colors.statusPending;
  138. drawnow;
  139. % Scan for videos
  140. sideFiles = dir(fullfile(sideDir, '*_Side_*.mp4'));
  141. frontFiles = dir(fullfile(frontDir, '*_Front_*.mp4'));
  142. % Load alignment if it exists
  143. alignmentFile = fullfile(outDir, 'alignment_table.mat');
  144. alignmentData = table();
  145. if exist(alignmentFile, 'file')
  146. S = load(alignmentFile);
  147. if isfield(S, 'results')
  148. alignmentData = S.results;
  149. elseif isfield(S, 'allResults')
  150. alignmentData = S.allResults;
  151. end
  152. % Normalize CoreID case
  153. if ~ismember('CoreID', alignmentData.Properties.VariableNames)
  154. if ismember('coreID', alignmentData.Properties.VariableNames)
  155. alignmentData.Properties.VariableNames{'coreID'} = 'CoreID';
  156. end
  157. end
  158. end
  159. % Find video pairs
  160. filePairs = findVideoPairs(sideFiles, sideDir, frontFiles, frontDir, handles);
  161. % --- Map GUI column names into valid struct field names ---
  162. colNamesGUI = { ...
  163. 'coreID', ...
  164. 'SideVideo', ...
  165. 'FrontVideo', ...
  166. 'SideDLC', ...
  167. 'FrontDLC', ...
  168. 'OffsetStart', ...
  169. 'OffsetEnd', ...
  170. 'TotalReaches', ...
  171. 'Exclude'};
  172. % Preallocate struct array
  173. emptyRow = cell2struct(cell(1,numel(colNamesGUI)), colNamesGUI, 2);
  174. emptyRow.OffsetStart = NaN;
  175. emptyRow.OffsetEnd = NaN;
  176. emptyRow.TotalReaches = 0;
  177. emptyRow.Exclude = false;
  178. pairStruct = repmat(emptyRow, 1, numel(filePairs));
  179. % Fill values
  180. parfor i = 1:numel(filePairs)
  181. pair = filePairs(i);
  182. coreID = pair.coreID;
  183. row = pairStruct(i); % template row with all fields
  184. % Fill known columns (must match colNamesGUI)
  185. row.coreID = coreID;
  186. row.SideVideo = isfile(pair.sideVideo);
  187. row.FrontVideo = isfile(pair.frontVideo);
  188. row.SideDLC = ~isempty(fastReadDLC(pair.sideVideo, 'Side', 'existsonly'));
  189. row.FrontDLC = ~isempty(fastReadDLC(pair.frontVideo, 'Front', 'existsonly'));
  190. % Alignment lookup from MAT
  191. if ~isempty(alignmentData) && ismember('CoreID', alignmentData.Properties.VariableNames)
  192. matchIdx = strcmp(alignmentData.CoreID, coreID);
  193. if any(matchIdx)
  194. if ismember('OffsetStart', alignmentData.Properties.VariableNames)
  195. row.OffsetStart = alignmentData.OffsetStart(matchIdx);
  196. elseif ismember('Offset', alignmentData.Properties.VariableNames)
  197. row.OffsetStart = alignmentData.Offset(matchIdx); % fallback
  198. end
  199. if ismember('OffsetEnd', alignmentData.Properties.VariableNames)
  200. row.OffsetEnd = alignmentData.OffsetEnd(matchIdx);
  201. end
  202. end
  203. end
  204. % Reach stats
  205. reachFile = fullfile(outDir, sprintf('%s_reaches.mat', coreID));
  206. if exist(reachFile, 'file')
  207. R = load(reachFile, 'reaches');
  208. if isfield(R, 'reaches')
  209. row.TotalReaches = numel(R.reaches);
  210. end
  211. end
  212. pairStruct(i) = row;
  213. end
  214. % Convert to table
  215. pairSummary = struct2table(pairStruct, 'AsArray', true);
  216. excludeFile = fullfile(outDir, 'exclude_table.mat');
  217. if exist(excludeFile, 'file')
  218. S = load(excludeFile);
  219. if isfield(S, 'excludeTable')
  220. [tf, loc] = ismember(pairSummary.coreID, S.excludeTable.CoreID);
  221. pairSummary.Exclude(tf) = S.excludeTable.Exclude(loc(tf));
  222. end
  223. end
  224. % Push into GUI
  225. handles.tblStatus.Data = pairSummary;
  226. handles.msgLabel.Text = 'Pipeline ready for processing';
  227. handles.msgLabel.FontColor = handles.colors.statusPending;
  228. drawnow;
  229. % Store pairs
  230. handles.pairs = filePairs;
  231. guidata(fig, handles);
  232. end
  233. function filePairs = findVideoPairs(sideFiles, sideDir, frontFiles, frontDir, handles)
  234. filePairs = struct('coreID', {}, 'sideVideo', {}, 'frontVideo', {}); % initialize
  235. for i = 1:length(sideFiles)
  236. sideName = sideFiles(i).name;
  237. % Extract animal + condition from the side filename
  238. tokens = regexp(sideName, '(.*)_Side_(.*)\.mp4', 'tokens', 'once');
  239. if isempty(tokens)
  240. fprintf('⚠️ Could not parse: %s\n', sideName);
  241. continue;
  242. end
  243. animal = tokens{1};
  244. condition = tokens{2};
  245. coreID = sprintf('%s_%s', animal, condition);
  246. % Rebuild the expected front filename pattern
  247. frontPattern = sprintf('%s_Front_%s.mp4', animal, condition);
  248. matchIdx = find(strcmp({frontFiles.name}, frontPattern), 1);
  249. if isempty(matchIdx)
  250. fprintf('❌ Skipping %s: No exact match for %s\n', sideName, frontPattern);
  251. handles.msgLabel.Text = sprintf('❌ Skipping %s: No exact match for %s\n', sideName, frontPattern);
  252. handles.msgLabel.FontColor = handles.colors.errorRed;
  253. drawnow;
  254. continue;
  255. end
  256. newPair = struct();
  257. newPair.coreID = coreID;
  258. newPair.sideVideo = fullfile(sideDir, sideName);
  259. newPair.frontVideo = fullfile(frontDir, frontFiles(matchIdx).name);
  260. filePairs(end+1) = newPair; %#ok<AGROW>
  261. end
  262. end
  263. function logMsg(msg, showInCommand, fid)
  264. try
  265. if showInCommand
  266. handles.msgLabel.Text = msg;
  267. handles.msgLabel.FontColor = handles.colors.statusPending;
  268. drawnow;
  269. fprintf('%s\n', msg);
  270. end
  271. catch
  272. % If handles doesn't exist, just print to command window
  273. if showInCommand
  274. fprintf('%s\n', msg);
  275. end
  276. end
  277. fprintf(fid, '%s\n', msg);
  278. end
  279. function defaultLikelihood = getLatestPawLikelihoodFromLog(outDir)
  280. % Default fallback value if no file or value found
  281. defaultLikelihood = 0.6;
  282. % Find log files matching pattern
  283. logFiles = dir(fullfile(outDir, 'reach_detection_log_*.txt'));
  284. if isempty(logFiles)
  285. fprintf('No reach detection log files found, using default likelihood %.2f\n', defaultLikelihood);
  286. return;
  287. end
  288. % Sort files by date, latest first
  289. [~, idx] = sort([logFiles.datenum], 'descend');
  290. newestLogFile = fullfile(logFiles(idx(1)).folder, logFiles(idx(1)).name);
  291. % Open the newest log file
  292. fid = fopen(newestLogFile, 'r');
  293. if fid == -1
  294. fprintf('Failed to open log file, using default likelihood %.2f\n', defaultLikelihood);
  295. return;
  296. end
  297. % Read line by line to find "paw likelihood:"
  298. while ~feof(fid)
  299. tline = fgetl(fid);
  300. if contains(tline, 'paw likelihood:', 'IgnoreCase', true)
  301. tokens = regexp(tline, 'paw likelihood:\s*([0-9.]+)', 'tokens', 'once');
  302. if ~isempty(tokens)
  303. val = str2double(tokens{1});
  304. if ~isnan(val) && val >= 0 && val <= 1
  305. defaultLikelihood = val;
  306. fprintf('Loaded paw likelihood threshold %.2f from %s\n', val, newestLogFile);
  307. fclose(fid);
  308. return;
  309. end
  310. end
  311. end
  312. end
  313. fclose(fid);
  314. fprintf('paw likelihood line not found, using default value %.2f\n', defaultLikelihood);
  315. end
  316. function excludeCellEditCallback(fig, src, evt)
  317. handles = guidata(fig);
  318. % Get the row and new value
  319. row = evt.Indices(1);
  320. newVal = evt.NewData;
  321. % Update table data in handles
  322. handles.tblStatus.Data.Exclude(row) = newVal;
  323. % Save to file immediately
  324. excludeTable = table(handles.tblStatus.Data.coreID, ...
  325. handles.tblStatus.Data.Exclude, ...
  326. 'VariableNames', {'CoreID','Exclude'});
  327. outDir = handles.outDir;
  328. save(fullfile(outDir, 'exclude_table.mat'), 'excludeTable');
  329. % Push updated handles back
  330. guidata(fig, handles);
  331. excludeTable = table( ...
  332. handles.tblStatus.Data.coreID, ...
  333. handles.tblStatus.Data.Exclude, ...
  334. 'VariableNames', {'CoreID','Exclude'} );
  335. outDir = handles.outDir;
  336. save(fullfile(outDir, 'exclude_table.mat'), 'excludeTable');
  337. % Also save CSV for Fiji
  338. writetable(excludeTable, fullfile(outDir, 'exclude_table.csv'));
  339. % Optional: show status message
  340. handles.msgLabel.Text = sprintf('Updated Exclude for %s', ...
  341. handles.tblStatus.Data.coreID{row});
  342. end
  343. %% ---- Read DLC data
  344. function dlcData = fastReadDLC(videoPath, viewType, varargin)
  345. % Cache DLC data in memory to avoid repeated file reads
  346. % Optional third parameter: 'existsonly' to just check if file exists
  347. % Parse optional parameter
  348. checkExistsOnly = false;
  349. if nargin > 2 && strcmpi(varargin{1}, 'existsonly')
  350. checkExistsOnly = true;
  351. end
  352. % If just checking existence, use same logic as readDLCcsv
  353. if checkExistsOnly
  354. % Extract base path and coreID (filename without extension)
  355. [videoFolder, videoNameNoExt, ~] = fileparts(videoPath);
  356. % Construct path to expected subfolder (Front or Side)
  357. expectedFolder = fullfile(videoFolder, '..', viewType);
  358. expectedFolder = fullfile(expectedFolder); % resolve any relative paths
  359. % Look for DLC CSVs in the expected folder
  360. allCSV = dir(fullfile(expectedFolder, '*.csv'));
  361. % Match by video core name
  362. matches = contains({allCSV.name}, videoNameNoExt) & ...
  363. endsWith({allCSV.name}, '.csv') & ...
  364. ~contains({allCSV.name}, 'meta');
  365. dlcData = any(matches); % Return true if any matches found
  366. return;
  367. end
  368. persistent dlcCache
  369. if isempty(dlcCache)
  370. dlcCache = containers.Map();
  371. end
  372. [~, fname] = fileparts(videoPath);
  373. cacheKey = [fname '_' viewType];
  374. if dlcCache.isKey(cacheKey)
  375. dlcData = dlcCache(cacheKey);
  376. return;
  377. end
  378. % Original DLC reading logic here
  379. dlcData = readDLCcsv(videoPath, viewType);
  380. % Cache the result
  381. dlcCache(cacheKey) = dlcData;
  382. end
  383. function data_table = readDLCcsv(videoFile, expectedView)
  384. % Validate expectedView input
  385. if ~ischar(expectedView) || ~ismember(expectedView, {'Front', 'Side'})
  386. error('expectedView must be ''Front'' or ''Side''.');
  387. end
  388. % Extract base path and coreID (filename without extension)
  389. [videoFolder, videoNameNoExt, ~] = fileparts(videoFile);
  390. % Construct path to expected subfolder (Front or Side)
  391. expectedFolder = fullfile(videoFolder, '..', expectedView);
  392. expectedFolder = fullfile(expectedFolder); % resolve any relative paths
  393. % Look for DLC CSVs in the expected folder
  394. allCSV = dir(fullfile(expectedFolder, '*.csv'));
  395. % Match by video core name
  396. matches = contains({allCSV.name}, videoNameNoExt) & ...
  397. endsWith({allCSV.name}, '.csv') & ...
  398. ~contains({allCSV.name}, 'meta');
  399. if ~any(matches)
  400. warning('No DLC file found for %s in folder %s', videoNameNoExt, expectedView);
  401. data_table = [];
  402. return;
  403. end
  404. % Read the first match
  405. csv_file = fullfile(allCSV(find(matches, 1)).folder, allCSV(find(matches, 1)).name);
  406. %fprintf('Loaded DLC file: %s\n', csv_file);
  407. % --- DLC-specific parsing ---
  408. opts = detectImportOptions(csv_file); %#ok
  409. opts.DataLine = 4; % DLC data starts at line 4
  410. % Read header lines 2 and 3
  411. header_lines = readcell(csv_file, 'Range', '2:3');
  412. header_names = strcat(header_lines(1,:), '.', header_lines(2,:));
  413. header_names{1} = 'frames'; % Rename first column
  414. opts.VariableNames = header_names;
  415. % Read final table
  416. warnState = warning('off', 'MATLAB:table:ModifiedAndSavedVarnames');
  417. data_table = readtable(csv_file, opts);
  418. warning(warnState); % restore original warning state
  419. % Normalize column names to lowercase
  420. data_table.Properties.VariableNames = lower(data_table.Properties.VariableNames);
  421. end
  422. %% alignment of videos (FIJI based)
  423. function alignVideosCallback(fig)
  424. handles = guidata(fig);
  425. fijiPath = 'C:\Fiji.app\fiji-win64.exe';
  426. dataDir = handles.baseDir;
  427. % Double-escape backslashes for the dir argument
  428. dirArg = sprintf('dir=%s', strrep(dataDir, '\', '\\'));
  429. % Build Fiji system call
  430. cmd = sprintf('"%s" --ij2 --run "SPG_VideoSyncTool " "%s"', fijiPath, dirArg);
  431. system(cmd);
  432. results = loadAlignmentResults(dataDir);
  433. % Match CoreIDs
  434. if istable(handles.tblStatus.Data)
  435. guiCoreIDs = handles.tblStatus.Data.coreID; % if stored as a table with variable coreID
  436. else
  437. guiCoreIDs = handles.tblStatus.Data(:,1); % if it's still a cell array
  438. end
  439. % Force everything to cell array of char
  440. if isstring(guiCoreIDs)
  441. guiCoreIDs = cellstr(guiCoreIDs);
  442. elseif iscell(guiCoreIDs)
  443. guiCoreIDs = cellfun(@char, guiCoreIDs, 'UniformOutput', false);
  444. elseif isnumeric(guiCoreIDs)
  445. guiCoreIDs = cellstr(string(guiCoreIDs));
  446. end
  447. fijiCoreIDs = cellstr(results.CoreID);
  448. [tf, loc] = ismember(guiCoreIDs, fijiCoreIDs);
  449. % Column indices
  450. colStart = find(strcmp(handles.tblStatus.ColumnName, 'OffsetStart'));
  451. colEnd = find(strcmp(handles.tblStatus.ColumnName, 'OffsetEnd'));
  452. % Update GUI table
  453. for i = 1:numel(guiCoreIDs)
  454. if tf(i)
  455. handles.tblStatus.Data{i,colStart} = results.OffsetStart(loc(i));
  456. handles.tblStatus.Data{i,colEnd} = results.OffsetEnd(loc(i));
  457. end
  458. end
  459. guidata(fig, handles);
  460. end
  461. function results = loadAlignmentResults(baseDir)
  462. outDir = fullfile(baseDir, 'OUT');
  463. csvFile = fullfile(outDir, 'alignment_table.csv');
  464. if ~isfile(csvFile)
  465. warning('⚠️ No alignment_table.csv found in %s', outDir);
  466. results = table(); % return empty table
  467. return;
  468. end
  469. results = readtable(csvFile);
  470. % --- Normalize column names ---
  471. if ~ismember('CoreID', results.Properties.VariableNames)
  472. error('alignment_table.csv missing CoreID column');
  473. end
  474. % Handle old format (single Offset column)
  475. if ismember('Offset', results.Properties.VariableNames)
  476. results.OffsetStart = results.Offset;
  477. results.OffsetEnd = results.Offset;
  478. results.Offset = []; % drop old col
  479. end
  480. % Handle new format (Offset1/Offset2)
  481. if ismember('Offset1', results.Properties.VariableNames)
  482. results.Properties.VariableNames{'Offset1'} = 'OffsetStart';
  483. end
  484. if ismember('Offset2', results.Properties.VariableNames)
  485. results.Properties.VariableNames{'Offset2'} = 'OffsetEnd';
  486. end
  487. % Ensure columns exist even if missing
  488. if ~ismember('OffsetStart', results.Properties.VariableNames)
  489. results.OffsetStart = NaN(height(results),1);
  490. end
  491. if ~ismember('OffsetEnd', results.Properties.VariableNames)
  492. results.OffsetEnd = NaN(height(results),1);
  493. end
  494. % Save MAT version for speed
  495. save(fullfile(outDir, 'alignment_table.mat'), 'results');
  496. fprintf('✅ Loaded %d entries from alignment_table.csv\n', height(results));
  497. end
  498. % ----------- Offset calculation on these alignments
  499. function offset = getDynamicOffset(results, coreID, sideFrame)
  500. %GETDYNAMICOFFSET Interpolates offset for a given frame
  501. % results = table from alignment_table.csv
  502. % coreID = string identifying the trial
  503. % sideFrame = the frame number in Side video
  504. rowIdx = find(strcmpi(results.CoreID, coreID), 1);
  505. if isempty(rowIdx)
  506. error('CoreID %s not found in results.', coreID);
  507. end
  508. r = results(rowIdx,:);
  509. % If only one offset exists, return it
  510. if isnan(r.OffsetEnd) || r.OffsetEnd == r.OffsetStart
  511. offset = r.OffsetStart;
  512. return;
  513. end
  514. % Linear interpolation between the two anchor points
  515. offset = r.OffsetStart + ...
  516. (r.OffsetEnd - r.OffsetStart) * ...
  517. ( (sideFrame - r.SideFrame1) / (r.SideFrame2 - r.SideFrame1) );
  518. offset = round(offset); % return integer frame offset
  519. end
  520. function offset = getRangeOffset(results, coreID, fStart, fEnd)
  521. %GETRANGEOFFSET Average offset across a frame range
  522. ofs1 = getDynamicOffset(results, coreID, fStart);
  523. ofs2 = getDynamicOffset(results, coreID, fEnd);
  524. % Use average offset across the reach
  525. offset = round(mean([ofs1 ofs2]));
  526. end
  527. %% ---- Define Reaching Frames and store
  528. function detectReachesCallback(fig)
  529. rerunFrontOnly = true; % <--- set to true if you want to recompute only front-mapping
  530. handles = guidata(fig);
  531. baseDir = handles.baseDir;
  532. outDir = handles.outDir;
  533. pairs = handles.pairs;
  534. % Load alignment info
  535. alignmentFile = fullfile(outDir, 'alignment_table.mat');
  536. if ~isfile(alignmentFile)
  537. uialert(fig, 'No alignment_table.mat found. Please run alignment first.', 'Missing File');
  538. return;
  539. end
  540. results = load(alignmentFile, 'results').results;
  541. % Setup log
  542. logFile = fullfile(baseDir, sprintf('reach_detection_log_%s.txt', datestr(now, 'yyyymmdd_HHMMSS')));
  543. fid = fopen(logFile, 'w');
  544. cleanupObj = onCleanup(@() fclose(fid));
  545. % --- Reach Detection Parameters ---
  546. params = struct( ...
  547. 'paw_likelihood', handles.likelihoodEdit.Value,... %get user input
  548. 'seq_likelihood', 0.3, ...
  549. 'pellet_likelihood_threshold', 0.4, ...
  550. 'pellet_check_frames', 10, ...
  551. 'min_frames', 10, ...
  552. 'frame_buffer', 15, ...
  553. 'gap_tolerance', 5, ...
  554. 'gauss_smooth', 50 ...
  555. );
  556. logParams(params, fid);
  557. % Prepare table
  558. tblData = handles.tblStatus.Data;
  559. tblData = ensureReachColumns(tblData);
  560. % --- Loop through animals ---
  561. for pairIdx = 1:numel(pairs)
  562. coreID = pairs(pairIdx).coreID;
  563. % Skip if reach output already exists
  564. reachMatFile = fullfile(outDir, sprintf('%s_reaches.mat', coreID));
  565. if isfile(reachMatFile) && ~rerunFrontOnly
  566. logMsg(sprintf('⏩ Skipping %s (reach files already exist)', coreID), true, fid);
  567. continue;
  568. end
  569. %%% NEW: load BOTH DLC tables for alignment model
  570. sideCSV = readDLCcsv(pairs(pairIdx).sideVideo, 'Side');
  571. frontCSV = readDLCcsv(pairs(pairIdx).frontVideo, 'Front');
  572. rowIdx = find(strcmpi(results.CoreID, coreID), 1);
  573. if ~isempty(rowIdx)
  574. OffsetStart = results.OffsetStart(rowIdx);
  575. OffsetEnd = results.OffsetEnd(rowIdx);
  576. SideFrame1 = results.SideFrame1(rowIdx);
  577. FrontFrame1 = results.FrontFrame1(rowIdx);
  578. SideFrame2 = results.SideFrame2(rowIdx);
  579. FrontFrame2 = results.FrontFrame2(rowIdx);
  580. % initOffset = coarse hint for candidate search
  581. initOffset = mean([OffsetStart, OffsetEnd], 'omitnan');
  582. if isnan(initOffset), initOffset = 0; end
  583. % If anything is missing, set to NaN
  584. if isnan(OffsetStart), OffsetStart = []; end
  585. if isnan(OffsetEnd), OffsetEnd = []; end
  586. else
  587. OffsetStart = [];
  588. OffsetEnd = [];
  589. SideFrame1 = [];
  590. FrontFrame1 = [];
  591. SideFrame2 = [];
  592. FrontFrame2 = [];
  593. end
  594. alignOpts = struct( ...
  595. 'hi', 0.90, ...
  596. 'prom', 0.20, ...
  597. 'minSep', 350, ...
  598. 'minWidth', 12, ...
  599. 'medianFront', 11, ...
  600. 'nnWindow', 1500, ...
  601. 'initOffset', round(initOffset), ...
  602. 'OffsetStart', OffsetStart, ...
  603. 'OffsetEnd', OffsetEnd, ...
  604. 'SideFrame1', SideFrame1, ...
  605. 'FrontFrame1', FrontFrame1, ...
  606. 'SideFrame2', SideFrame2, ...
  607. 'FrontFrame2', FrontFrame2, ...
  608. 'alignMode', 'segmented', ... % 'segmented', 'affine_weighted', 'piecewise', or 'pchip'
  609. 'anchorWeight', 15, ...
  610. 'minPairs', 3);
  611. alignOpts.fp.MinPeakHeight = 'auto';
  612. alignOpts.fp.MinPeakProminence = 'auto';
  613. model = fitAlignmentByPeaksAUC(sideCSV, frontCSV, coreID, outDir, alignOpts, fid);
  614. % --- rerun front only ---
  615. if rerunFrontOnly && isfile(reachMatFile)
  616. S = load(reachMatFile, 'reaches');
  617. reaches = S.reaches;
  618. % Remap using the NEW model
  619. for j = 1:numel(reaches)
  620. oldFront = reaches(j).frontFrames; % store original
  621. newFront = model.mapSide2Front(reaches(j).sideFrames); % recompute
  622. % Print first/last few values to avoid flooding log
  623. logMsg(sprintf('Reach %d: oldFront(1:3)=%s ... %s | newFront(1:3)=%s ... %s', ...
  624. j, mat2str(oldFront(1:min(3,end))), mat2str(oldFront(max(end-2,1):end)), ...
  625. mat2str(newFront(1:min(3,end))), mat2str(newFront(max(end-2,1):end))), ...
  626. true, fid);
  627. reaches(j).frontFrames = newFront; % overwrite
  628. end
  629. save(reachMatFile, 'reaches');
  630. logMsg(sprintf('🔄 Updated front frames for %s with new alignment model', coreID), true, fid);
  631. continue;
  632. end
  633. logMsg(sprintf('🐾 Detecting reaches for %s...', coreID), true, fid);
  634. if isempty(sideCSV) || isempty(frontCSV)
  635. logMsg('Missing DLC tables for this pair. Skipping.', true, fid);
  636. continue;
  637. end
  638. % Define slit threshold
  639. params.slit_threshold = getSlitThreshold(sideCSV, fid);
  640. params.pellet_merge_gap = 40;
  641. % Find reaches using sideView only
  642. [starts, ends, isReach, reachID] = detectReaches(sideCSV, params, fid, coreID, outDir);
  643. % Pack reach structs
  644. reaches = struct('startFrame', {}, 'endFrame', {}, 'sideFrames', {}, 'frontFrames', {}, 'label', {});
  645. for j = 1:length(starts)
  646. reaches(j).startFrame = starts(j);
  647. reaches(j).endFrame = ends(j);
  648. reaches(j).sideFrames = starts(j):ends(j);
  649. reaches(j).frontFrames = model.mapSide2Front(reaches(j).sideFrames);
  650. % --- Pellet presence check ---
  651. f1 = starts(j);
  652. f2 = ends(j);
  653. len = f2 - f1 + 1;
  654. if len > 0
  655. checkFrames = f1 : f1 + floor(len/4); % first quarter of reach
  656. checkFrames = checkFrames(checkFrames <= height(sideCSV));
  657. if ismember('pellet_likelihood', sideCSV.Properties.VariableNames)
  658. pelLh = sideCSV.pellet_likelihood(checkFrames);
  659. pelLh = pelLh(~isnan(pelLh));
  660. if isempty(pelLh) || mean(pelLh) < 0.2 % threshold adjustable
  661. reaches(j).label = 'Attempt - No Pellet';
  662. else
  663. reaches(j).label = '';
  664. end
  665. else
  666. reaches(j).label = ''; % pellet not tracked
  667. end
  668. else
  669. reaches(j).label = '';
  670. end
  671. end
  672. save(fullfile(outDir, sprintf('%s_reaches.mat', coreID)), 'reaches');
  673. tblData.TotalReaches(pairIdx) = numel(reaches);
  674. end
  675. % Final update
  676. handles.tblStatus.Data = tblData;
  677. handles.msgLabel.Text = 'Reach detection completed!';
  678. handles.msgLabel.FontColor = handles.colors.successColor;
  679. guidata(fig, handles);
  680. % ===== Helper Functions =====
  681. function logParams(p, fid)
  682. logMsg('--- Reach Detection Parameters ---', false, fid);
  683. fns = fieldnames(p);
  684. for i = 1:numel(fns)
  685. logMsg(sprintf('%s: %s', strrep(fns{i}, '_', ' '), num2str(p.(fns{i}))), false, fid);
  686. end
  687. end
  688. function slit_x = getSlitThreshold(data, fid)
  689. if all(ismember({'slit_bottom__x', 'slit_top__x'}, data.Properties.VariableNames))
  690. slit_x_all = [data.slit_bottom__x; data.slit_top__x];
  691. slit_lh_all = [data.slit_bottom__likelihood; data.slit_top__likelihood];
  692. [~, idx] = sort(slit_lh_all, 'descend');
  693. topX = slit_x_all(idx(1:max(10, round(0.05 * numel(idx)))));
  694. slit_x = mean(topX, 'omitnan');
  695. logMsg(sprintf('Slit threshold: %.2f', slit_x), true, fid);
  696. else
  697. error('Missing slit coordinates');
  698. end
  699. end
  700. function [starts, ends, isReach, reachID] = detectReaches(data, p, fid, coreID, outDir)
  701. % ---------------------------
  702. % Step 0: parameter defaults
  703. % ---------------------------
  704. if ~isfield(p,'gauss_smooth') || isempty(p.gauss_smooth), p.gauss_smooth = 5; end
  705. if ~isfield(p,'slit_threshold') || isempty(p.slit_threshold), p.slit_threshold = 240;end
  706. if ~isfield(p,'min_frames') || isempty(p.min_frames), p.min_frames = 6; end
  707. if ~isfield(p,'gap_tolerance') || isempty(p.gap_tolerance), p.gap_tolerance = 5; end
  708. if ~isfield(p,'pellet_contact_dist')|| isempty(p.pellet_contact_dist),p.pellet_contact_dist= 15; end
  709. if ~isfield(p,'pellet_merge_gap') || isempty(p.pellet_merge_gap), p.pellet_merge_gap = 20; end
  710. if ~isfield(p,'overshoot_margin') || isempty(p.overshoot_margin), p.overshoot_margin = 10; end
  711. if ~isfield(p,'frame_buffer') || isempty(p.frame_buffer), p.frame_buffer = 0; end
  712. % peak-finding (for overshoot inside reaches)
  713. if ~isfield(p,'pk_min_prom') || isempty(p.pk_min_prom), p.pk_min_prom = 5; end
  714. if ~isfield(p,'pk_min_dist') || isempty(p.pk_min_dist), p.pk_min_dist = 15; end
  715. % ---------------------------
  716. % Step 1: Extract variables
  717. % ---------------------------
  718. pt_x = data.paw_tip__x;
  719. pc_x = data.paw_center__x;
  720. pt_lh = data.paw_tip__likelihood;
  721. pc_lh = data.paw_center__likelihood;
  722. % ---------------------------
  723. % Step 2: Smooth trajectories
  724. % ---------------------------
  725. pt_x_smooth = smoothdata(pt_x, 'gaussian', p.gauss_smooth);
  726. pc_x_smooth = smoothdata(pc_x, 'gaussian', p.gauss_smooth);
  727. % ---------------------------
  728. % Step 3: Candidate (above slit)
  729. % ---------------------------
  730. above_tip = pt_x_smooth > p.slit_threshold;
  731. above_center = pc_x_smooth > p.slit_threshold;
  732. isCandidate = above_tip | above_center;
  733. % Likelihood threshold (dynamic)
  734. lh_all = max(pt_lh, pc_lh);
  735. lh_cand = lh_all(isCandidate);
  736. lh_cand = lh_cand(~isnan(lh_cand));
  737. if ~isempty(lh_cand)
  738. pL = prctile(lh_cand, 5);
  739. pH = prctile(lh_cand, 99);
  740. lh_win = lh_cand;
  741. lh_win(lh_win < pL) = pL;
  742. lh_win(lh_win > pH) = pH;
  743. else
  744. lh_win = lh_all(~isnan(lh_all));
  745. end
  746. t_otsu = graythresh(lh_win);
  747. t_q = prctile(lh_win, 80);
  748. dyn_lh_thresh = max([(t_otsu + t_q)/2, 0.25]);
  749. p.paw_likelihood = dyn_lh_thresh;
  750. logMsg(sprintf('Dynamic paw likelihood: Otsu=%.3f, Q80=%.3f, chosen=%.3f', ...
  751. t_otsu, t_q, dyn_lh_thresh), true, fid);
  752. % ---------------------------
  753. % Step 3.5: Pellet contact heuristic
  754. % ---------------------------
  755. if ismember('pellet_x', data.Properties.VariableNames)
  756. dx = pt_x_smooth - data.pellet_x;
  757. if ismember('pellet_y', data.Properties.VariableNames) && ...
  758. ismember('paw_tip__y', data.Properties.VariableNames)
  759. dy = data.paw_tip__y - data.pellet_y;
  760. dist_tip = hypot(dx, dy);
  761. else
  762. dist_tip = abs(dx);
  763. end
  764. if ismember('pellet_likelihood', data.Properties.VariableNames)
  765. valid_lh = (pt_lh >= p.paw_likelihood) & (data.pellet_likelihood >= 0.3);
  766. else
  767. valid_lh = (pt_lh >= p.paw_likelihood);
  768. end
  769. pellet_contact = (dist_tip < p.pellet_contact_dist) & valid_lh;
  770. else
  771. pellet_contact = false(height(data),1);
  772. end
  773. % ---------------------------
  774. % Step 4: Likelihood & length filtering
  775. % ---------------------------
  776. [starts_raw, ends_raw] = getSegments(isCandidate);
  777. isReach = false(height(data), 1);
  778. too_short = []; too_short_e = [];
  779. low_lh = []; low_lh_e = [];
  780. for i = 1:numel(starts_raw)
  781. s = starts_raw(i);
  782. e = ends_raw(i);
  783. len = e - s + 1;
  784. if len < p.min_frames
  785. too_short(end+1) = s; %#ok<AGROW>
  786. too_short_e(end+1) = e;
  787. continue;
  788. end
  789. frames_above_lh = (pt_lh(s:e) >= p.paw_likelihood) | ...
  790. (pc_lh(s:e) >= p.paw_likelihood);
  791. if sum(frames_above_lh) >= p.min_frames
  792. isReach(s:e) = true;
  793. else
  794. low_lh(end+1) = s; %#ok<AGROW>
  795. low_lh_e(end+1) = e;
  796. end
  797. end
  798. % ---------------------------
  799. % Step 5: Morphological cleanup
  800. % ---------------------------
  801. se = ones(max(1, round(p.gap_tolerance)), 1);
  802. isReach = imclose(isReach, se);
  803. % ---------------------------
  804. % Step 6: Split reaches if several peaks with overshoot/pellet contact
  805. % ---------------------------
  806. [starts, ends] = getSegments(isReach);
  807. nFrames = height(data); % total number of frames, needed for buffer clamping
  808. % Precompute pellet anchors (merge contacts into areas)
  809. [pc_s_all, pc_e_all] = getSegments(pellet_contact);
  810. if ~isempty(pc_s_all)
  811. merged_s = pc_s_all(1);
  812. merged_e = pc_e_all(1);
  813. new_pc_s = []; new_pc_e = [];
  814. for k = 2:numel(pc_s_all)
  815. if pc_s_all(k) - merged_e <= p.pellet_merge_gap
  816. merged_e = pc_e_all(k);
  817. else
  818. new_pc_s(end+1) = merged_s; %#ok<AGROW>
  819. new_pc_e(end+1) = merged_e;
  820. merged_s = pc_s_all(k);
  821. merged_e = pc_e_all(k);
  822. end
  823. end
  824. new_pc_s(end+1) = merged_s; %#ok<AGROW>
  825. new_pc_e(end+1) = merged_e; %#ok<AGROW>
  826. pc_s_all = new_pc_s;
  827. pc_e_all = new_pc_e;
  828. end
  829. % Center of each pellet area = pellet anchors
  830. pellet_anchors_all = round((pc_s_all + pc_e_all)/2);
  831. % ---------------------------
  832. % Step 6b: Refine reaches by anchors (pellet + overshoot)
  833. % ---------------------------
  834. new_starts = [];
  835. new_ends = [];
  836. overshoot_anchor_list = [];
  837. for r = 1:numel(starts)
  838. s = starts(r);
  839. e = ends(r);
  840. % ---- pellet anchors (merged areas, already computed globally) ----
  841. pel_here = pellet_anchors_all(pellet_anchors_all >= s & pellet_anchors_all <= e);
  842. % ---- overshoot anchors (true peaks beyond pellet) ----
  843. seg_x = pt_x_smooth(s:e);
  844. if ismember('pellet_x', data.Properties.VariableNames)
  845. pellet_here = nanmedian(data.pellet_x(s:e));
  846. else
  847. pellet_here = prctile(seg_x,95); % fallback if pellet not tracked
  848. end
  849. segLen = numel(seg_x);
  850. if segLen < 2
  851. % Too short for peak detection
  852. continue; % skip this segment
  853. end
  854. baseDist = 60;
  855. minDist = min(baseDist, segLen - 1);
  856. if minDist >= (segLen - 1)
  857. minDist = segLen - 2;
  858. end
  859. if minDist < 1
  860. minDist = 1;
  861. end
  862. fprintf('segLen: %.2f, minDist: %.2f\n',segLen, minDist);
  863. [pks, locs] = findpeaks(seg_x, 'MinPeakProminence', 0.15, 'MinPeakHeight', 0.6, 'MinPeakDistance', minDist);
  864. keep = pks > (pellet_here + p.overshoot_margin);
  865. over_here = s + locs(keep) - 1;
  866. overshoot_anchor_list = [overshoot_anchor_list, over_here(:)']; %#ok<AGROW>
  867. % ---- combine anchors ----
  868. anchors = unique([pel_here(:); over_here(:)]);
  869. % ---- refine: find local minima around each anchor ----
  870. if isempty(anchors)
  871. % no anchors → keep full reach
  872. new_starts(end+1) = s;
  873. new_ends(end+1) = e;
  874. else
  875. for a = 1:numel(anchors)
  876. this_anchor = anchors(a);
  877. % find nearest local minimum to left
  878. left_idx = this_anchor-1;
  879. while left_idx > s+1 && seg_x(left_idx-s+1) > seg_x(left_idx-s)
  880. left_idx = left_idx-1;
  881. end
  882. % find nearest local minimum to right
  883. right_idx = this_anchor+1;
  884. while right_idx < e-1 && seg_x(right_idx-s+1) > seg_x(right_idx-s+2)
  885. right_idx = right_idx+1;
  886. end
  887. % add refined segment
  888. new_starts(end+1) = max(s,left_idx);
  889. new_ends(end+1) = min(e,right_idx);
  890. end
  891. end
  892. end
  893. % replace
  894. starts = new_starts(:)';
  895. ends = new_ends(:)';
  896. % ============================================================
  897. % Step 6c: Merge overlapping/adjacent segments
  898. % ============================================================
  899. if ~isempty(starts)
  900. % sort just in case
  901. [starts, sortIdx] = sort(starts);
  902. ends = ends(sortIdx);
  903. merged_s = starts(1);
  904. merged_e = ends(1);
  905. clean_starts = [];
  906. clean_ends = [];
  907. for k = 2:numel(starts)
  908. if starts(k) <= merged_e % overlap or touching
  909. merged_e = max(merged_e, ends(k));
  910. else
  911. clean_starts(end+1) = merged_s; %#ok<AGROW>
  912. clean_ends(end+1) = merged_e; %#ok<AGROW>
  913. merged_s = starts(k);
  914. merged_e = ends(k);
  915. end
  916. end
  917. % add last segment
  918. clean_starts(end+1) = merged_s;
  919. clean_ends(end+1) = merged_e;
  920. starts = clean_starts;
  921. ends = clean_ends;
  922. end
  923. % ---------------------------
  924. % Step 7: Apply buffer (playback only)
  925. % ---------------------------
  926. starts = max(starts - p.frame_buffer, 1);
  927. ends = min(ends + p.frame_buffer, nFrames);
  928. % ---------------------------
  929. % Step 8: Assign reachID
  930. % ---------------------------
  931. reachID = zeros(nFrames, 1);
  932. for i = 1:numel(starts)
  933. reachID(starts(i):ends(i)) = i;
  934. end
  935. % ==========================
  936. % === FINAL QC FIGURE ===
  937. % ==========================
  938. qcDir = fullfile(outDir, 'QC'); if ~exist(qcDir,'dir'), mkdir(qcDir); end
  939. qcFile = fullfile(qcDir, sprintf('REACH_pawlikelihood_cutoff_%s.png', coreID));
  940. % ---- Summary plot ----
  941. f = figure('Visible','on','Name',sprintf('QC Summary — %s',coreID));
  942. tiledlayout(f,1,1,'Padding','compact','TileSpacing','compact');
  943. ax = nexttile; hold(ax,'on');
  944. % Smoothed tip (colored by likelihood)
  945. scatter(ax, 1:height(data), pt_x_smooth, 12, pt_lh, 'filled');
  946. colormap(ax, parula);
  947. cb = colorbar(ax); cb.Label.String = 'Tip likelihood';
  948. % Highlights (kept reaches with border)
  949. highlightSegments(ax, too_short, too_short_e, [1 0 0]); % too short (red)
  950. highlightSegments(ax, low_lh, low_lh_e, [1 0.5 0]); % low likelihood (orange)
  951. highlightSegments(ax, starts, ends, 'c'); % final kept (cyan)
  952. uistack(findobj(ax,'Type','patch'),'bottom');
  953. % --- Pellet contact dots (magenta circles) ---
  954. if exist('pellet_contact','var') && any(pellet_contact)
  955. scatter(ax, find(pellet_contact), pt_x_smooth(pellet_contact), ...
  956. 25, 'mo', 'filled', 'MarkerFaceAlpha', 0.7, 'DisplayName','Pellet contact');
  957. end
  958. % --- Pellet anchors (black triangles) ---
  959. if exist('pellet_anchors_all','var') && ~isempty(pellet_anchors_all)
  960. scatter(ax, pellet_anchors_all, pt_x_smooth(pellet_anchors_all), ...
  961. 40, 'k^', 'filled', 'MarkerFaceAlpha', 0.9, 'DisplayName','Pellet anchor');
  962. end
  963. % --- Overshoot anchors (green diamonds) ---
  964. if exist('overshoot_anchor_list','var') && ~isempty(overshoot_anchor_list)
  965. scatter(ax, overshoot_anchor_list, pt_x_smooth(overshoot_anchor_list), ...
  966. 40, 'gd', 'filled', 'MarkerFaceAlpha', 0.9, 'DisplayName','Overshoot anchor');
  967. end
  968. % Legend stubs
  969. hShort = plot(ax,nan,nan,'s','MarkerFaceColor',[1 0 0],'MarkerEdgeColor','none');
  970. hLow = plot(ax,nan,nan,'s','MarkerFaceColor',[1 0.5 0],'MarkerEdgeColor','none');
  971. hFinal = plot(ax,nan,nan,'s','MarkerFaceColor','c','MarkerEdgeColor','none');
  972. hPel = plot(ax,nan,nan,'o','MarkerFaceColor','m','MarkerEdgeColor','none');
  973. hPelA = plot(ax,nan,nan,'^','MarkerFaceColor','k','MarkerEdgeColor','none');
  974. hOver = plot(ax,nan,nan,'d','MarkerFaceColor','g','MarkerEdgeColor','none');
  975. legend(ax,[hShort hLow hFinal hPel hPelA hOver], ...
  976. {'Too short','Low likelihood','Final','Pellet contact','Pellet anchor','Overshoot anchor'}, ...
  977. 'Location','bestoutside');
  978. % Save
  979. qcFileSummary_png = fullfile(qcDir, sprintf('REACH_%s_QC_summary.png', coreID));
  980. saveas(f, qcFileSummary_png);
  981. qcFileSummary_fig = fullfile(qcDir, sprintf('REACH_%s_QC_summary.fig', coreID));
  982. saveas(f, qcFileSummary_fig);
  983. close(f);
  984. logMsg(sprintf('Num Reaches detected (side): %d', numel(starts)), true, fid);
  985. % ---- helper for colored patches with border ----
  986. function highlightSegments(ax, s, e, color)
  987. yl = ylim(ax);
  988. for k = 1:numel(s)
  989. patch(ax, [s(k) e(k) e(k) s(k)], ...
  990. [yl(1) yl(1) yl(2) yl(2)], ...
  991. color, 'FaceAlpha', 0.20, ...
  992. 'EdgeColor', 'k', 'LineWidth', 0.5, ...
  993. 'HandleVisibility','off');
  994. end
  995. end
  996. end
  997. end
  998. function [s, e] = getSegments(logicalVec)
  999. s = []; e = []; in = false;
  1000. for i = 1:length(logicalVec)
  1001. if ~in && logicalVec(i), s(end+1) = i; in = true;
  1002. elseif in && ~logicalVec(i), e(end+1) = i-1; in = false; end
  1003. end
  1004. if in, e(end+1) = length(logicalVec); end
  1005. end
  1006. function tbl = ensureReachColumns(tbl)
  1007. if ~ismember('TotalReaches', tbl.Properties.VariableNames)
  1008. tbl.TotalReaches = zeros(height(tbl),1);
  1009. elseif iscell(tbl.TotalReaches)
  1010. tbl.TotalReaches = cellfun(@(x) ifempty(x,0), tbl.TotalReaches);
  1011. end
  1012. end
  1013. %
  1014. % function val = ifempty(x, def)
  1015. % if isempty(x), val = def; else, val = x; end
  1016. % end
  1017. function v = getField(T, preferName, altName)
  1018. % grab T.(preferName) if present, else T.(altName). Returns zeros if missing.
  1019. v = [];
  1020. if istable(T)
  1021. if ismember(preferName, T.Properties.VariableNames)
  1022. v = T.(preferName);
  1023. elseif ismember(altName, T.Properties.VariableNames)
  1024. v = T.(altName);
  1025. end
  1026. end
  1027. if isempty(v)
  1028. v = zeros(height(T),1);
  1029. end
  1030. end
  1031. function model = fitAlignmentByPeaksAUC(Tside, Tfront, coreID, outDir, alignOpts, fid)
  1032. % fitAlignmentByPeaksAUC (anchors = plateau starts; peaks only for QC)
  1033. % Robust to missing fields in alignOpts (fills sane defaults).
  1034. if nargin < 6 || isempty(fid), fid = 1; end
  1035. % ---------- defaults (filled if missing) ----------
  1036. DEF.initOffset = 0; % hint only (coarse offset overrides)
  1037. DEF.sideSmooth = 101; % Gaussian win (numeric); struct ok (see getSmoothWin)
  1038. DEF.frontSmooth = 401;
  1039. DEF.timeWindow = 2000; % initial candidate window (frames)
  1040. DEF.alpha = 1.0; % timing weight
  1041. DEF.beta = 0.7; % area weight (z-diff)
  1042. DEF.maxAllow = 1.2; % initial max matching cost
  1043. DEF.minMatchesWant = 6; % target pairs before fitting
  1044. DEF.maxRelaxIters = 3; % relax rounds
  1045. DEF.coarseLagMax = 10000; % coarse xcorr max lag (frames)
  1046. DEF.downsample = 20; % downsample for coarse xcorr
  1047. DEF.plateauThresh = 0.4; % passed into buildPelletEvents (also adapts inside)
  1048. DEF.plateauMinDur = 100;
  1049. if nargin < 5 || isempty(alignOpts), alignOpts = struct(); end
  1050. fn = fieldnames(DEF);
  1051. for k=1:numel(fn)
  1052. if ~isfield(alignOpts, fn{k}) || isempty(alignOpts.(fn{k}))
  1053. alignOpts.(fn{k}) = DEF.(fn{k});
  1054. end
  1055. end
  1056. % Peak detection thresholds (fill missing subfields)
  1057. FPDEF = struct('MinPeakHeight',0.6,'MinPeakProminence',0.15, ...
  1058. 'MinPeakDistance',1500,'MinPeakWidth',20);
  1059. if ~isfield(alignOpts,'fp') || isempty(alignOpts.fp), alignOpts.fp = FPDEF; end
  1060. sub = fieldnames(FPDEF);
  1061. for k=1:numel(sub)
  1062. if ~isfield(alignOpts.fp, sub{k}) || isempty(alignOpts.fp.(sub{k}))
  1063. alignOpts.fp.(sub{k}) = FPDEF.(sub{k});
  1064. end
  1065. end
  1066. % Allow smoothing window as numeric or struct with .gauss
  1067. sideWin = getSmoothWin(alignOpts.sideSmooth, 201);
  1068. frontWin = getSmoothWin(alignOpts.frontSmooth, 401);
  1069. % ---------- extract pellet likelihoods ----------
  1070. pelSide = getField(Tside,'pellet__likelihood','pellet_likelihood');
  1071. pelFront = getField(Tfront,'pellet__likelihood','pellet_likelihood');
  1072. % ---------- build events (your buildPelletEvents does dynamic tuning) ----------
  1073. [eSide, fpSideEff, plSideEff] = buildPelletEvents(pelSide, sideWin, alignOpts.fp, alignOpts.plateauThresh, alignOpts.plateauMinDur, 'Side', fid);
  1074. [eFront, fpFrontEff, plFrontEff] = buildPelletEvents(pelFront, frontWin, alignOpts.fp, alignOpts.plateauThresh, alignOpts.plateauMinDur, 'Front', fid);
  1075. logMsg(sprintf('[%s] Side dynThr=%.3f (base=%.2f) | kept plateaus=%d (minDur=%d)', ...
  1076. coreID, plSideEff.plateauThresh, alignOpts.plateauThresh, size(eSide.plat.se,1), alignOpts.plateauMinDur), true, fid);
  1077. logMsg(sprintf('[%s] Side: peaks=%d, plateaus=%d', ...
  1078. coreID, numel(eSide.peak.idx), size(eSide.plat.se,1)), true, fid);
  1079. logMsg(sprintf('[%s] Front: peaks=%d, plateaus=%d', ...
  1080. coreID, numel(eFront.peak.idx), size(eFront.plat.se,1)), true, fid);
  1081. % ---------- choose anchors = plateau STARTS (fallback to peaks) ----------
  1082. if ~isempty(eSide.plat.se), anchorS = eSide.plat.se(:,1); else, anchorS = eSide.peak.idx; end
  1083. if ~isempty(eFront.plat.se), anchorF = eFront.plat.se(:,1); else, anchorF = eFront.peak.idx; end
  1084. % ---------- coarse offset (front->side) from xcorr of smoothed signals ----------
  1085. rough = estimate_coarse_offset(eSide.sig, eFront.sig, alignOpts.downsample, alignOpts.coarseLagMax);
  1086. logMsg(sprintf('[%s] Coarse offset (front->side): %+d frames', coreID, rough), true, fid);
  1087. % ---------- restrict anchors to manual boundaries ----------
  1088. if isfield(alignOpts,'FrontFrame1') && isfield(alignOpts,'FrontFrame2') && ...
  1089. isfield(alignOpts,'SideFrame1') && isfield(alignOpts,'SideFrame2')
  1090. manualFront = [alignOpts.FrontFrame1, alignOpts.FrontFrame2];
  1091. manualSide = [alignOpts.SideFrame1, alignOpts.SideFrame2];
  1092. % Sort in case user provided out of order
  1093. [manualFront, order] = sort(manualFront);
  1094. manualSide = manualSide(order);
  1095. fMin = manualFront(1); fMax = manualFront(end);
  1096. sMin = manualSide(1); sMax = manualSide(end);
  1097. keepF = anchorF >= fMin & anchorF <= fMax;
  1098. keepS = anchorS >= sMin & anchorS <= sMax;
  1099. anchorF = anchorF(keepF);
  1100. anchorS = anchorS(keepS);
  1101. logMsg(sprintf('Restricted auto anchors to manual window: Front[%d..%d], Side[%d..%d]', ...
  1102. fMin,fMax,sMin,sMax), true, fid);
  1103. end
  1104. % ---------- matching with relaxation cascade ----------
  1105. [a,b, matchPairs, usedCost, matchparam, model] = match_with_relax( ...
  1106. eSide, eFront, anchorS, anchorF, rough, alignOpts, fid);
  1107. % Ensure matchPairs indices are valid for the restricted anchor arrays
  1108. valid = matchPairs(:,1) <= numel(anchorS) & matchPairs(:,2) <= numel(anchorF);
  1109. matchPairs = matchPairs(valid,:);
  1110. model.frontN = height(Tfront);
  1111. model.sideN = height(Tside);
  1112. % ---------- QC ----------
  1113. figDir = fullfile(outDir,'QC'); if ~exist(figDir,'dir'), mkdir(figDir); end
  1114. cmap = lines(max(1,size(matchPairs,1)));
  1115. % (1) events + matched anchors (colored)
  1116. f = figure('Visible','off','Name',sprintf('Matched anchors — %s',coreID));
  1117. tiledlayout(2,1,'Padding','compact','TileSpacing','compact');
  1118. % --- SIDE ---
  1119. ax1 = nexttile; hold(ax1,'on');
  1120. plot(ax1, eSide.sig, 'm'); % raw signal in Side time
  1121. for k=1:size(eSide.plat.se,1)
  1122. S=eSide.plat.se(k,1); E=eSide.plat.se(k,2);
  1123. patch(ax1,[S E E S],[0 0 1 1],'m','FaceAlpha',0.1,'EdgeColor','none');
  1124. end
  1125. scatter(ax1, eSide.peak.idx, eSide.sig(eSide.peak.idx), 16,'k','filled');
  1126. if ~isempty(matchPairs)
  1127. for p = 1:size(matchPairs,1)
  1128. i = matchPairs(p,1);
  1129. if i <= numel(anchorS)
  1130. scatter(ax1, anchorS(i), eSide.sig(anchorS(i)), 36, cmap(p,:), 'filled');
  1131. end
  1132. text(anchorS(i), min(0.95, eSide.sig(anchorS(i))+0.06), sprintf('%d',p), ...
  1133. 'Color', cmap(p,:), 'FontWeight','bold','HorizontalAlignment','center');
  1134. end
  1135. end
  1136. title(ax1,'Side: matched anchors (original time)'); ylabel(ax1,'Lh');
  1137. ylim(ax1,[0 1]); xlim(ax1,[1 numel(eSide.sig)]);
  1138. % --- FRONT ---
  1139. ax2 = nexttile; hold(ax2,'on');
  1140. plot(ax2, eFront.sig, 'g');
  1141. for k=1:size(eFront.plat.se,1)
  1142. S=eFront.plat.se(k,1); E=eFront.plat.se(k,2);
  1143. patch(ax2,[S E E S],[0 0 1 1],'g','FaceAlpha',0.1,'EdgeColor','none');
  1144. end
  1145. scatter(ax2, eFront.peak.idx, eFront.sig(eFront.peak.idx), 16,'k','filled');
  1146. if ~isempty(matchPairs)
  1147. for p = 1:size(matchPairs,1)
  1148. j = matchPairs(p,2);
  1149. scatter(ax2, anchorF(j), eFront.sig(anchorF(j)), 36, cmap(p,:), 'filled');
  1150. text(anchorF(j), min(0.95, eFront.sig(anchorF(j))+0.06), sprintf('%d',p), ...
  1151. 'Color', cmap(p,:), 'FontWeight','bold','HorizontalAlignment','center');
  1152. end
  1153. end
  1154. % --- Boundary anchors (blue stars) ---
  1155. if isfield(alignOpts,'OffsetStart') && ~isnan(alignOpts.OffsetStart) ...
  1156. && isfield(alignOpts,'SideFrame1') && isfield(alignOpts,'FrontFrame1')
  1157. scatter(ax1, alignOpts.SideFrame1, eSide.sig(min(end, alignOpts.SideFrame1)), ...
  1158. 60, 'b*', 'LineWidth',1.5);
  1159. scatter(ax2, alignOpts.FrontFrame1, eFront.sig(min(end, alignOpts.FrontFrame1)), ...
  1160. 60, 'b*', 'LineWidth',1.5);
  1161. end
  1162. if isfield(alignOpts,'OffsetEnd') && ~isnan(alignOpts.OffsetEnd) ...
  1163. && isfield(alignOpts,'SideFrame2') && isfield(alignOpts,'FrontFrame2')
  1164. scatter(ax1, alignOpts.SideFrame2, eSide.sig(min(end, alignOpts.SideFrame2)), ...
  1165. 60, 'b*', 'LineWidth',1.5);
  1166. scatter(ax2, alignOpts.FrontFrame2, eFront.sig(min(end, alignOpts.FrontFrame2)), ...
  1167. 60, 'b*', 'LineWidth',1.5);
  1168. end
  1169. title(ax2,'Front: matched anchors (original time)');
  1170. xlabel(ax2,'Frame'); ylabel(ax2,'Lh');
  1171. ylim(ax2,[0 1]); xlim(ax2,[1 numel(eFront.sig)]);
  1172. saveas(f, fullfile(figDir, sprintf('ALIGN_EVENT-MATCH_%s.png', coreID)));
  1173. close(f);
  1174. % (2) Frame mapping QC
  1175. if ~isempty(matchPairs)
  1176. fQC = figure('Visible','off','Name',sprintf('Frame Mapping QC — %s', coreID));
  1177. hold on; grid on;
  1178. % --- Anchors driving the mapping ---
  1179. usedF = model.tF_sorted(:);
  1180. usedS = model.tS_sorted(:);
  1181. % --- Mapping curve ---
  1182. modeStr = lower(alignOpts.alignMode);
  1183. if any(strcmpi(modeStr, {'piecewise','segmented'}))
  1184. % broken-stick polyline through anchors
  1185. plot(usedF, usedS, 'r-', 'LineWidth',1.5, 'MarkerSize',4, ...
  1186. 'DisplayName',[upper(modeStr(1)) modeStr(2:end) ' mapping']);
  1187. else
  1188. % dense sample for affine/pchip
  1189. xx = linspace(0, model.frontN, 400);
  1190. yy = model.mapFront2Side(xx);
  1191. plot(xx, yy, 'r-', 'LineWidth',1.5, 'DisplayName','Mapping curve');
  1192. % overlay polyline through used anchors for context
  1193. plot(usedF, usedS, 'r--o', 'LineWidth',1.0, 'MarkerSize',4, ...
  1194. 'DisplayName','Anchor polyline');
  1195. end
  1196. % Dropped anchors
  1197. if isfield(model,'tF_all')
  1198. droppedMask = ~ismember(model.tF_all, model.tF_sorted);
  1199. droppedManual = droppedMask & model.w_all > 1;
  1200. droppedAuto = droppedMask & model.w_all == 1;
  1201. scatter(model.tF_all(droppedAuto), model.tS_all(droppedAuto), ...
  1202. 60, 'rx', 'LineWidth',1.5, 'DisplayName','Dropped auto anchors');
  1203. scatter(model.tF_all(droppedManual), model.tS_all(droppedManual), ...
  1204. 70, 'ms','filled','LineWidth',1.5,'DisplayName','Dropped manual anchors');
  1205. end
  1206. % Kept anchors
  1207. if isfield(model,'isManual')
  1208. scatter(model.tF_sorted(~model.isManual), model.tS_sorted(~model.isManual), ...
  1209. 50, 'b*','LineWidth',1.5,'DisplayName','Kept auto anchors');
  1210. scatter(model.tF_sorted(model.isManual), model.tS_sorted(model.isManual), ...
  1211. 70, 'gd','filled','LineWidth',1.5,'DisplayName','Kept manual anchors');
  1212. end
  1213. % --- Axis formatting ---
  1214. xlim([0 model.frontN*1.05]);
  1215. ylim([0 model.sideN*1.05]);
  1216. xlabel('Front anchor frame (video)');
  1217. ylabel('Side anchor frame (video)');
  1218. title(sprintf('Frame mapping with %s fit', alignOpts.alignMode));
  1219. xline(model.frontN,'g--','End Front');
  1220. yline(model.sideN,'m--','End Side');
  1221. legend('show','Location','best');
  1222. saveas(fQC, fullfile(figDir, sprintf('ALIGN_FRAMEMAP_%s.png', coreID)));
  1223. close(fQC);
  1224. end
  1225. % (3) Offset drift QC
  1226. if ~isempty(matchPairs)
  1227. f5 = figure('Visible','off','Name',sprintf('Offset Drift — %s', coreID));
  1228. hold on; grid on;
  1229. offsets = model.tS_sorted(:) - model.tF_sorted(:);
  1230. plot(model.tF_sorted, offsets, 'o-','LineWidth',1.5,'MarkerSize',8);
  1231. xx = linspace(min(model.tF_sorted), max(model.tF_sorted), 300);
  1232. yy = model.mapFront2Side(xx) - xx;
  1233. plot(xx, yy, 'r-','LineWidth',1.5,'DisplayName','Spline-predicted');
  1234. % Manual anchors
  1235. if isfield(alignOpts,'FrontFrame1') && isfield(alignOpts,'SideFrame1')
  1236. off1 = alignOpts.SideFrame1 - alignOpts.FrontFrame1;
  1237. scatter(alignOpts.FrontFrame1, off1, 60,'b*','LineWidth',1.5);
  1238. text(alignOpts.FrontFrame1, off1, 'Start anchor', ...
  1239. 'VerticalAlignment','top','HorizontalAlignment','left','Color','b');
  1240. end
  1241. if isfield(alignOpts,'FrontFrame2') && isfield(alignOpts,'SideFrame2')
  1242. off2 = alignOpts.SideFrame2 - alignOpts.FrontFrame2;
  1243. scatter(alignOpts.FrontFrame2, off2, 60,'b*','LineWidth',1.5);
  1244. text(alignOpts.FrontFrame2, off2, 'End anchor', ...
  1245. 'VerticalAlignment','bottom','HorizontalAlignment','right','Color','b');
  1246. end
  1247. xlabel('Front anchor frame'); ylabel('Offset (Side - Front)');
  1248. title('Offset drift across video');
  1249. saveas(f5, fullfile(figDir, sprintf('ALIGN_OFFSET_%s.png', coreID)));
  1250. close(f5);
  1251. end
  1252. % (4) Local slope QC
  1253. if ~isempty(matchPairs)
  1254. fSlope = figure('Visible','off','Name',sprintf('Local Slope — %s', coreID));
  1255. xx = linspace(model.tF_sorted(1), model.tF_sorted(end), 300);
  1256. yy = model.mapFront2Side(xx);
  1257. slope = gradient(yy) ./ gradient(xx);
  1258. plot(xx, slope, 'o-','LineWidth',1.5); grid on
  1259. xlabel('Front frame'); ylabel('dSide/dFront');
  1260. title('Local slope (effective rate drift)');
  1261. saveas(fSlope, fullfile(figDir, sprintf('ALIGN_SLOPE_%s.png', coreID)));
  1262. close(fSlope);
  1263. end
  1264. end
  1265. % ============================ helpers ==============================
  1266. function w = getSmoothWin(val, def)
  1267. if isstruct(val)
  1268. if isfield(val,'gauss') && ~isempty(val.gauss), w = val.gauss; else, w = def; end
  1269. elseif isnumeric(val) && ~isempty(val)
  1270. w = val;
  1271. else
  1272. w = def;
  1273. end
  1274. if w < 1, w = 1; end
  1275. end
  1276. function off = estimate_coarse_offset(sigSide, sigFront, ds, maxLag)
  1277. if ds < 1, ds = 1; end
  1278. s1 = downsample(double(sigSide), max(1,ds));
  1279. s2 = downsample(double(sigFront), max(1,ds));
  1280. % focus on dips (pellet disappearing)
  1281. s1 = 1 - (s1 - min(s1))/max(eps, (max(s1)-min(s1)));
  1282. s2 = 1 - (s2 - min(s2))/max(eps, (max(s2)-min(s2)));
  1283. L = min(numel(s1), numel(s2));
  1284. s1 = s1(1:L); s2 = s2(1:L);
  1285. ml = min(maxLag, L-5);
  1286. [c,lags] = xcorr(s1 - mean(s1), s2 - mean(s2), ml, 'coeff');
  1287. [~,idx] = max(c);
  1288. off = lags(idx) * ds; % frames front->side
  1289. end
  1290. function [a,b,matchPairs,usedCost,matchparam,model] = match_with_relax(eSide,eFront,anchorS,anchorF,initOff,opts,fid)
  1291. % Defaults
  1292. a = 1; b = initOff; matchPairs = []; usedCost = opts.maxAllow;
  1293. alpha0 = opts.alpha; beta = opts.beta; win = opts.timeWindow; maxAllow = opts.maxAllow;
  1294. % --- Candidate matching loop ---
  1295. for attempt = 0:opts.maxRelaxIters
  1296. [cI,cJ,cost,m,n] = candidates(eSide,eFront,anchorS,anchorF,initOff,win,alpha0,beta);
  1297. logMsg(sprintf(' attempt %d: candidates=%d (m=%d,n=%d) | win=%d | beta=%.2f | maxAllow=%.2f', ...
  1298. attempt,numel(cost),m,n,win,beta,maxAllow), true,fid);
  1299. if isempty(cost)
  1300. matchPairs = [];
  1301. else
  1302. C = inf(m,n); C(sub2ind([m n],cI,cJ)) = cost;
  1303. try
  1304. [ai,aj] = matchpairs(C,maxAllow);
  1305. matchPairs = [ai,aj];
  1306. catch
  1307. [~,order] = sort(cost,'ascend');
  1308. usedI = false(m,1); usedJ = false(n,1); tmp = [];
  1309. for k = 1:numel(order)
  1310. ii = cI(order(k)); jj = cJ(order(k));
  1311. if ~usedI(ii) && ~usedJ(jj) && cost(order(k)) <= maxAllow
  1312. tmp(end+1,:) = [ii jj]; %#ok<AGROW>
  1313. usedI(ii) = true; usedJ(jj) = true;
  1314. end
  1315. end
  1316. matchPairs = tmp;
  1317. end
  1318. end
  1319. logMsg(sprintf(' attempt %d: matched=%d',attempt,size(matchPairs,1)), true,fid);
  1320. if size(matchPairs,1) >= opts.minMatchesWant, usedCost = maxAllow; break; end
  1321. % relax constraints
  1322. beta = max(0.2,beta*0.6);
  1323. win = round(win*1.5);
  1324. maxAllow = maxAllow + 0.3;
  1325. if attempt==1 && size(matchPairs,1) <= 1 && ...
  1326. ~isempty(eSide.peak.idx) && ~isempty(eFront.peak.idx)
  1327. logMsg(' switching to PEAK MAXIMA anchors for matching', true,fid);
  1328. anchorS = eSide.peak.idx;
  1329. anchorF = eFront.peak.idx;
  1330. matchPairs = [];
  1331. end
  1332. end
  1333. % --- Gather matched anchors ---
  1334. if ~isempty(matchPairs)
  1335. tS = anchorS(matchPairs(:,1));
  1336. tF = anchorF(matchPairs(:,2));
  1337. else
  1338. tS = []; tF = [];
  1339. end
  1340. weights = ones(numel(tF),1);
  1341. % Inject manual anchors directly
  1342. if isfield(opts,'FrontFrame1') && isfield(opts,'SideFrame1')
  1343. tF = [opts.FrontFrame1; tF(:)];
  1344. tS = [opts.SideFrame1; tS(:)];
  1345. weights = [opts.anchorWeight; weights];
  1346. end
  1347. if isfield(opts,'FrontFrame2') && isfield(opts,'SideFrame2')
  1348. tF = [opts.FrontFrame2; tF(:)];
  1349. tS = [opts.SideFrame2; tS(:)];
  1350. weights = [opts.anchorWeight; weights];
  1351. end
  1352. % ---------- SAVE RAW ANCHORS HERE ----------
  1353. model.tF_all = tF(:);
  1354. model.tS_all = tS(:);
  1355. model.w_all = weights(:);
  1356. % before the gate
  1357. pinManuals = isfield(opts,'alignMode') && strcmpi(opts.alignMode,'segmented');
  1358. % ================== CONSISTENCY GATE (diagnostic-only if segmented) ==================
  1359. autoMask = (weights == 1);
  1360. tF_auto = tF(autoMask); tS_auto = tS(autoMask);
  1361. if numel(tF_auto) >= 3
  1362. ab0 = lscov([tF_auto(:) ones(numel(tF_auto),1)], tS_auto(:), ones(numel(tF_auto),1));
  1363. a0 = ab0(1); b0 = ab0(2);
  1364. r = (a0 .* tF(:) + b0) - tS(:); % residuals vs. auto-only line
  1365. madR = mad(r,1);
  1366. tau = max(6*max(madR,1), 400);
  1367. isManual = weights > 1;
  1368. % Always log residuals for manuals
  1369. if any(isManual)
  1370. for k = find(isManual)'
  1371. logMsg(sprintf('Manual anchor: F=%d | S_meas=%d | S_calc=%.0f | resid=%.0f', ...
  1372. tF(k), tS(k), a0*tF(k)+b0, r(k)), true, fid);
  1373. end
  1374. end
  1375. if ~pinManuals
  1376. % original behavior (allowed to change manual weights)
  1377. badMan = isManual & abs(r) > tau;
  1378. if any(badMan)
  1379. % Softer penalty but still >1 keeps them as manuals (optional):
  1380. weights(badMan) = max(weights(badMan) * 0.5, 1.1);
  1381. % Or: weights(badMan) = 0.1; % hard drop (NOT used for segmented)
  1382. logMsg(sprintf('Manual anchor(s) softened by gate: idx=%s', ...
  1383. mat2str(find(badMan)')), true, fid);
  1384. end
  1385. else
  1386. % segmented mode: diagnostics only, don't touch weights
  1387. model.flaggedManual = isManual & abs(r) > tau; % for QC highlighting
  1388. if any(model.flaggedManual)
  1389. logMsg('Segmented mode: manual anchors flagged by residuals (kept).', true, fid);
  1390. end
  1391. end
  1392. end
  1393. % =====================================================================
  1394. %%% 1) Prune automatic anchors too close to manual ones
  1395. distThresh = 1000; % distance in frames, adjust to taste
  1396. isManual = weights > 1;
  1397. keepMask = true(size(tF));
  1398. for k = find(isManual)' % loop over manual anchors
  1399. closeIdx = abs(tF - tF(k)) < distThresh & ~isManual;
  1400. if any(closeIdx)
  1401. logMsg(sprintf('Dropping %d auto anchors near manual anchor F=%d', ...
  1402. sum(closeIdx), tF(k)), true, fid);
  1403. end
  1404. keepMask(closeIdx) = false;
  1405. end
  1406. %%% 2) Drop manual–manual conflicts (if r available)
  1407. if exist('r','var')
  1408. manualIdx = find(isManual);
  1409. for k = 1:numel(manualIdx)
  1410. i = manualIdx(k);
  1411. if ~keepMask(i), continue; end
  1412. tooClose = abs(tF(manualIdx) - tF(i)) < distThresh;
  1413. tooClose(manualIdx==i) = false;
  1414. if any(tooClose)
  1415. closeMans = manualIdx(tooClose);
  1416. group = [i; closeMans(:)];
  1417. [~,bestIdx] = min(abs(r(group)));
  1418. keepMask(group) = false;
  1419. keepMask(group(bestIdx)) = true;
  1420. logMsg(sprintf('Dropping %d manual anchors near F=%d (kept best residual)', ...
  1421. numel(group)-1, tF(group(bestIdx))), true, fid);
  1422. end
  1423. end
  1424. else
  1425. logMsg('Warning: residuals r not found, skipping manual–manual pruning', true, fid);
  1426. end
  1427. %%% 3) Apply pruning
  1428. tF = tF(keepMask);
  1429. tS = tS(keepMask);
  1430. weights = weights(keepMask);
  1431. isManual = isManual(keepMask);
  1432. %%% 4) NOW add Manual-line gating + slope guard
  1433. % ---------------------------------------------------------------
  1434. % MANUAL-LINE GATING + SLOPE GUARD
  1435. % ---------------------------------------------------------------
  1436. haveTwoManuals = isfield(opts,'FrontFrame1') && isfield(opts,'SideFrame1') && ...
  1437. isfield(opts,'FrontFrame2') && isfield(opts,'SideFrame2');
  1438. if haveTwoManuals
  1439. aM = (opts.SideFrame2 - opts.SideFrame1) / max(eps, (opts.FrontFrame2 - opts.FrontFrame1));
  1440. bM = opts.SideFrame1 - aM*opts.FrontFrame1;
  1441. res = tS - (aM*tF + bM);
  1442. tauRes = max(4*mad(res(~isManual),1), 600);
  1443. keepR = isManual | abs(res) <= tauRes;
  1444. if any(~keepR & ~isManual)
  1445. logMsg(sprintf('Pruned %d auto anchors by manual-line residual (|res|>%d)', ...
  1446. sum(~keepR & ~isManual), round(tauRes)), true, fid);
  1447. end
  1448. tF = tF(keepR);
  1449. tS = tS(keepR);
  1450. weights = weights(keepR);
  1451. isManual = isManual(keepR);
  1452. end
  1453. % If too few anchors left, we'll fall back later anyway
  1454. if numel(tF) < 2
  1455. a = NaN; b = NaN; usedCost = [];
  1456. params = struct('beta',beta,'win',win,'maxAllow',maxAllow);
  1457. model = struct();
  1458. return
  1459. end
  1460. % --- Final fit depending on opts.alignMode ---
  1461. if ~isfield(opts,'alignMode') || isempty(opts.alignMode)
  1462. opts.alignMode = 'affine_weighted';
  1463. end
  1464. % sort and dedup
  1465. [tF,ord] = sort(tF(:)); tS = tS(ord); weights = weights(ord);
  1466. [tF,ia] = unique(tF,'stable'); tS = tS(ia); weights = weights(ia);
  1467. %%% DEBUG: how many anchors survived
  1468. nAnchors = numel(tF);
  1469. logMsg(sprintf('Final anchors used in model: %d', nAnchors), true, fid);
  1470. logMsg(sprintf('Manual anchors kept: %d / %d', sum(isManual), sum(weights>1)), true, fid);
  1471. % Save final pruned set
  1472. [tF_sorted, ord] = sort(tF(:));
  1473. tS_sorted = tS(ord);
  1474. w_sorted = weights(ord);
  1475. [tF_sorted, ia] = unique(tF_sorted,'stable');
  1476. tS_sorted = tS_sorted(ia);
  1477. w_sorted = w_sorted(ia);
  1478. model.tF_sorted = tF_sorted;
  1479. model.tS_sorted = tS_sorted;
  1480. model.w_sorted = w_sorted;
  1481. model.isManual = w_sorted > 1; % manual anchors in final set
  1482. switch lower(opts.alignMode)
  1483. case 'affine_weighted'
  1484. X = [tF(:), ones(size(tF(:)))];
  1485. ab = lscov(X,tS(:),weights(:));
  1486. a = ab(1); b = ab(2);
  1487. model.mapFront2Side = @(tf) a*tf + b;
  1488. model.mapSide2Front = @(ts) (ts - b)./max(a,eps);
  1489. model.a=a; model.b=b;
  1490. case 'piecewise'
  1491. %%% NEW: multi-segment piecewise linear fit
  1492. % Sort anchors by front-frame time
  1493. [tF_sorted, ord] = sort(tF(:));
  1494. tS_sorted = tS(ord);
  1495. w_sorted = weights(ord);
  1496. % Deduplicate
  1497. [tF_sorted, ia] = unique(tF_sorted,'stable');
  1498. tS_sorted = tS_sorted(ia);
  1499. w_sorted = w_sorted(ia);
  1500. % Build a piecewise-linear map by interpolation
  1501. model.mapFront2Side = @(tf) interp1(tF_sorted, tS_sorted, tf, 'linear','extrap');
  1502. model.mapSide2Front = @(ts) interp1(tS_sorted, tF_sorted, ts, 'linear','extrap');
  1503. % For consistency with affine, also return an average slope
  1504. if numel(tF_sorted) >= 2
  1505. a = (tS_sorted(end)-tS_sorted(1)) / max(eps,(tF_sorted(end)-tF_sorted(1)));
  1506. b = tS_sorted(1) - a*tF_sorted(1);
  1507. else
  1508. a = NaN; b = NaN;
  1509. end
  1510. model.a = a;
  1511. model.b = b;
  1512. model.tF_sorted = tF_sorted;
  1513. model.tS_sorted = tS_sorted;
  1514. model.w_sorted = w_sorted;
  1515. model.isManual = w_sorted > 1; % flag manual anchors
  1516. case 'pchip'
  1517. keep = [true; diff(tF)>0] & [true; diff(tS)>0];
  1518. tf2=tF(keep); ts2=tS(keep);
  1519. ppF2S = pchip(tf2,ts2);
  1520. ppS2F = pchip(ts2,tf2);
  1521. model.mapFront2Side=@(tf) ppval(ppF2S,tf);
  1522. model.mapSide2Front=@(ts) ppval(ppS2F,ts);
  1523. a=(ts2(end)-ts2(1))/max(eps,(tf2(end)-tf2(1)));
  1524. b=ts2(1)-a*tf2(1);
  1525. model.a=a; model.b=b;
  1526. case 'segmented'
  1527. % Sort anchors by front time
  1528. [tF_sorted, ord] = sort(tF(:));
  1529. tS_sorted = tS(ord);
  1530. w_sorted = weights(ord);
  1531. % Need >=2 anchors
  1532. if numel(tF_sorted) < 2
  1533. warning('Segmented mode requires >=2 anchors; reverting to affine.');
  1534. a = (tS_sorted(end)-tS_sorted(1)) / max(eps,(tF_sorted(end)-tF_sorted(1)));
  1535. b = tS_sorted(1) - a*tF_sorted(1);
  1536. model.mapFront2Side = @(tf) a.*tf + b;
  1537. model.mapSide2Front = @(ts) (ts-b)./max(a,eps);
  1538. segSlopes = a; segIntercepts = b;
  1539. else
  1540. % Local helper functions with safety rails
  1541. mapFront2Side_local = @(x) localMapFront2Side(x, tF_sorted, tS_sorted);
  1542. mapSide2Front_local = @(y) localMapSide2Front(y, tS_sorted, tF_sorted);
  1543. % Vectorized
  1544. model.mapFront2Side = @(tf) arrayfun(mapFront2Side_local, tf);
  1545. model.mapSide2Front = @(ts) arrayfun(mapSide2Front_local, ts);
  1546. % Segment info for debugging/QC
  1547. segSlopes = diff(tS_sorted) ./ max(diff(tF_sorted), eps);
  1548. segIntercepts = tS_sorted(1:end-1) - segSlopes .* tF_sorted(1:end-1);
  1549. end
  1550. % Global slope for summary
  1551. if numel(tF_sorted) >= 2
  1552. a = (tS_sorted(end)-tS_sorted(1)) / max(eps,(tF_sorted(end)-tF_sorted(1)));
  1553. b = tS_sorted(1) - a*tF_sorted(1);
  1554. else
  1555. a = NaN; b = NaN;
  1556. end
  1557. % Save for QC
  1558. model.a = a; model.b = b;
  1559. model.tF_sorted = tF_sorted;
  1560. model.tS_sorted = tS_sorted;
  1561. model.w_sorted = w_sorted;
  1562. model.isManual = w_sorted > 1;
  1563. model.segSlopes = segSlopes;
  1564. model.segIntercepts = segIntercepts;
  1565. otherwise
  1566. error('Unknown opts.alignMode: %s',opts.alignMode);
  1567. end
  1568. % --- normalize model fields for QC (ALL MODES) ---
  1569. if ~isfield(model, 'tF_sorted') || ~isfield(model, 'tS_sorted')
  1570. [tF_sorted_plot, ord_plot] = sort(tF(:));
  1571. tS_sorted_plot = tS(ord_plot);
  1572. w_sorted_plot = weights(ord_plot);
  1573. model.tF_sorted = tF_sorted_plot;
  1574. model.tS_sorted = tS_sorted_plot;
  1575. model.w_sorted = w_sorted_plot;
  1576. model.isManual = w_sorted_plot > 1;
  1577. end
  1578. model.mode = lower(opts.alignMode); % for QC titles
  1579. % Verify mapping passes through the anchors the model uses
  1580. if isfield(model,'tF_sorted') && isfield(model,'tS_sorted')
  1581. err_used = max(abs(model.mapFront2Side(model.tF_sorted) - model.tS_sorted));
  1582. logMsg(sprintf('Model/anchor consistency: max |F2S(used anchors)-S| = %.3f frames', err_used), true, fid);
  1583. end
  1584. % return everything needed for QC
  1585. matchparam = struct( ...
  1586. 'beta',beta, ...
  1587. 'win',win, ...
  1588. 'maxAllow',maxAllow, ...
  1589. 'tF_sorted',tF, ...
  1590. 'tS_sorted',tS, ...
  1591. 'weights',weights);
  1592. end
  1593. % ---------- helper ----------
  1594. function ys = localMapFront2Side(x, tF_sorted, tS_sorted)
  1595. if x <= tF_sorted(1)
  1596. % Extrapolate using first two anchors
  1597. slope = (tS_sorted(2)-tS_sorted(1)) / max(eps,(tF_sorted(2)-tF_sorted(1)));
  1598. intercept = tS_sorted(1) - slope*tF_sorted(1);
  1599. ys = slope*x + intercept;
  1600. elseif x >= tF_sorted(end)
  1601. % Extrapolate using last two anchors
  1602. slope = (tS_sorted(end)-tS_sorted(end-1)) / max(eps,(tF_sorted(end)-tF_sorted(end-1)));
  1603. intercept = tS_sorted(end) - slope*tF_sorted(end);
  1604. ys = slope*x + intercept;
  1605. else
  1606. % Inside range → find segment
  1607. idx = find(tF_sorted(1:end-1) <= x & x <= tF_sorted(2:end), 1, 'last');
  1608. slope = (tS_sorted(idx+1)-tS_sorted(idx)) / max(eps,(tF_sorted(idx+1)-tF_sorted(idx)));
  1609. intercept = tS_sorted(idx) - slope*tF_sorted(idx);
  1610. ys = slope*x + intercept;
  1611. end
  1612. end
  1613. function xf = localMapSide2Front(y, tS_sorted, tF_sorted)
  1614. if y <= tS_sorted(1)
  1615. slope = (tF_sorted(2)-tF_sorted(1)) / max(eps,(tS_sorted(2)-tS_sorted(1)));
  1616. intercept = tF_sorted(1) - slope*tS_sorted(1);
  1617. xf = slope*y + intercept;
  1618. elseif y >= tS_sorted(end)
  1619. slope = (tF_sorted(end)-tF_sorted(end-1)) / max(eps,(tS_sorted(end)-tS_sorted(end-1)));
  1620. intercept = tF_sorted(end) - slope*tS_sorted(end);
  1621. xf = slope*y + intercept;
  1622. else
  1623. idx = find(tS_sorted(1:end-1) <= y & y <= tS_sorted(2:end), 1, 'last');
  1624. slope = (tF_sorted(idx+1)-tF_sorted(idx)) / max(eps,(tS_sorted(idx+1)-tS_sorted(idx)));
  1625. intercept = tF_sorted(idx) - slope*tS_sorted(idx);
  1626. xf = slope*y + intercept;
  1627. end
  1628. end
  1629. function tf = piecewise_inv(ts,a,b,c,k)
  1630. tf=zeros(size(ts));
  1631. yk = a*k + b;
  1632. idx1 = ts <= yk;
  1633. tf(idx1) = (ts(idx1)-b)./max(a,eps);
  1634. idx2 = ~idx1;
  1635. tf(idx2) = (ts(idx2)-(b-c*k))./max(a+c,eps);
  1636. end
  1637. function [cI,cJ,cost,m,n] = candidates(eSide,eFront,anchorS,anchorF,initOff,win,alpha,beta)
  1638. m = numel(anchorS); n = numel(anchorF);
  1639. cI=[]; cJ=[]; cost=[];
  1640. if ~m || ~n, return; end
  1641. for i = 1:m
  1642. s_t = anchorS(i);
  1643. expF = s_t - initOff;
  1644. JJ = find(abs(anchorF - expF) <= win);
  1645. if isempty(JJ), continue; end
  1646. dt = abs((anchorF(JJ) + initOff) - s_t); % timing cost
  1647. % area z for this side anchor
  1648. kS = map_idx(eSide, s_t);
  1649. aZs = 0; if kS>0 && kS<=numel(eSide.areaZ) && ~isnan(eSide.areaZ(kS)), aZs = eSide.areaZ(kS); end
  1650. % area z for each candidate front anchor
  1651. kF = map_idx(eFront, anchorF(JJ));
  1652. aZf = zeros(size(kF));
  1653. good = kF>0 & kF<=numel(eFront.areaZ);
  1654. aZf(good) = eFront.areaZ(kF(good));
  1655. aZf(~good) = 0;
  1656. dA = abs(aZf - aZs);
  1657. dA = min(dA, 2.5); % cap extremes
  1658. c = alpha*(dt./win) + beta*dA;
  1659. cI = [cI; i*ones(numel(JJ),1)];
  1660. cJ = [cJ; JJ(:)];
  1661. cost = [cost; c(:)];
  1662. end
  1663. end
  1664. function k = map_idx(E, anchorIdx)
  1665. % map anchor time(s) to indices in E.idx (unified event list)
  1666. [~,k] = ismember(anchorIdx, E.idx);
  1667. end
  1668. % =====================================================================
  1669. % ========================== HELPER FUNCTIONS ==========================
  1670. % =====================================================================
  1671. function [events, fpEff, plEff] = buildPelletEvents(sig, smoothWin, fp, plateauThresh, plateauMinDur, viewName, fid)
  1672. % Build peaks + plateaus from a pellet-likelihood trace.
  1673. % - Peaks via findpeaks (prominence*width area proxy)
  1674. % - Plateaus via dynamic high-threshold; anchor at START of plateau
  1675. % Returns per-class areas and z-scores, plus unified arrays.
  1676. if nargin < 6 || isempty(viewName), viewName = 'View'; end
  1677. sig = sig(:);
  1678. sig(isnan(sig)) = 0;
  1679. if smoothWin > 1
  1680. sig = smoothdata(sig,'gaussian',smoothWin);
  1681. end
  1682. % ---------------- Normalization ----------------
  1683. % Robustly stretch so tallest peak is ~1
  1684. if any(sig > 0)
  1685. lo = prctile(sig,1);
  1686. hi = prctile(sig,99);
  1687. sig = (sig - lo) ./ max(eps, hi-lo);
  1688. sig = min(max(sig,0),1); % clamp to [0,1]
  1689. end
  1690. % Effective fp struct we may overwrite adaptively
  1691. fpEff = fp;
  1692. % Collect candidate local maxima for adaptive thresholds
  1693. isLM = islocalmax(sig);
  1694. LMamp = sig(isLM);
  1695. if isempty(LMamp), LMamp = sig; end
  1696. base = median(sig);
  1697. spread = mad(sig,1);
  1698. % Adaptive MinPeakHeight
  1699. if ~isfield(fpEff,'MinPeakHeight') || isempty(fpEff.MinPeakHeight) || ...
  1700. (ischar(fpEff.MinPeakHeight) && strcmpi(fpEff.MinPeakHeight,'auto'))
  1701. k = 2.0; % MAD multiplier
  1702. q = 0.70; % percentile of local maxima
  1703. thrMAD = base + k*spread;
  1704. thrQ = quantile(LMamp,q);
  1705. mpHeight = max(thrMAD, thrQ);
  1706. fpEff.MinPeakHeight = min(mpHeight,0.95);
  1707. end
  1708. % Adaptive MinPeakProminence
  1709. if ~isfield(fpEff,'MinPeakProminence') || isempty(fpEff.MinPeakProminence) || ...
  1710. (ischar(fpEff.MinPeakProminence) && strcmpi(fpEff.MinPeakProminence,'auto'))
  1711. promMAD = 1.5*spread;
  1712. tailGap = max(0, quantile(LMamp,0.85)-base);
  1713. mpProm = max(promMAD, 0.5*tailGap);
  1714. fpEff.MinPeakProminence = min(mpProm,0.5);
  1715. end
  1716. % -------- peaks --------
  1717. [pks, locs, widths, prom] = findpeaks(sig, ...
  1718. 'MinPeakHeight', fpEff.MinPeakHeight, ...
  1719. 'MinPeakProminence', fpEff.MinPeakProminence, ...
  1720. 'MinPeakDistance', fpEff.MinPeakDistance, ...
  1721. 'MinPeakWidth', fpEff.MinPeakWidth);
  1722. area_pk = prom .* widths;
  1723. areaZ_pk = zscore_robust(area_pk);
  1724. % -------- adaptive plateau threshold (baseline-aware) --------
  1725. nz = sig(sig > 0); % nonzero samples
  1726. if isempty(nz)
  1727. thrDyn = plateauThresh; % degenerate fallback
  1728. hi=0;base=0;sLow=0;minFracHi=0;kMAD=0;
  1729. else
  1730. hi = prctile(nz,95); % robust "high"
  1731. base = prctile(nz,30); % baseline-ish level
  1732. sLow = mad(nz(nz <= hi), 1); % robust spread below hi
  1733. % Candidates:
  1734. minFracHi = 0.70; % try 0.65–0.75
  1735. kMAD = 2.0; % try 1.5–2.5
  1736. cand1 = minFracHi * hi;
  1737. cand2 = base + kMAD * sLow;
  1738. thrDyn = max(cand1, cand2); % pick the more conservative
  1739. thrDyn = min(thrDyn, 0.98*hi); % don't exceed the very top
  1740. end
  1741. above = sig >= thrDyn;
  1742. % bridge small dips
  1743. gap = 40;
  1744. above = imclose(above, ones(gap,1));
  1745. logMsg(sprintf('%s: plateau thr=%.3f | hi=%.3f base=%.3f MAD=%.3f (take=max(%.2f*hi, base+%.1f*MAD))', ...
  1746. viewName, thrDyn, hi, base, sLow, minFracHi, kMAD), true, fid);
  1747. d = diff([0; above; 0]);
  1748. S = find(d==1);
  1749. E = find(d==-1) - 1;
  1750. % keep only long-enough plateaus
  1751. keep = (E - S + 1) >= plateauMinDur;
  1752. S = S(keep); E = E(keep);
  1753. % plateau area proxy and anchor = START
  1754. nP = numel(S);
  1755. area_pl = zeros(nP,1);
  1756. idx_pl = zeros(nP,1);
  1757. for k=1:nP
  1758. seg = sig(S(k):E(k));
  1759. area_pl(k) = sum(seg);
  1760. idx_pl(k) = S(k);
  1761. end
  1762. areaZ_pl = zscore_robust(area_pl);
  1763. % -------- unify for "both" mode --------
  1764. idx_all_raw = [locs(:); idx_pl(:)];
  1765. area_all_raw = [area_pk(:); area_pl(:)];
  1766. type_all = [ones(numel(locs),1); 2*ones(numel(idx_pl),1)];
  1767. areaZ_all = zscore_robust(area_all_raw);
  1768. % -------- package --------
  1769. events.sig = sig;
  1770. events.peak.idx = locs(:);
  1771. events.peak.area = area_pk(:);
  1772. events.peak.areaZ = areaZ_pk(:);
  1773. events.plat.se = [S(:) E(:)];
  1774. events.plat.idxStart = idx_pl(:);
  1775. events.plat.area = area_pl(:);
  1776. events.plat.areaZ = areaZ_pl(:);
  1777. events.plat.thresh = thrDyn;
  1778. [events.idx, sortOrder] = sort(idx_all_raw(:));
  1779. events.area = area_all_raw(sortOrder);
  1780. events.areaZ = areaZ_all(sortOrder);
  1781. events.type = type_all(sortOrder);
  1782. % effective thresholds for logging
  1783. plEff.plateauThresh = thrDyn;
  1784. plEff.plateauMinDur = plateauMinDur;
  1785. end
  1786. % ---- utils ----
  1787. function z = zscore_robust(x)
  1788. x = x(:);
  1789. if isempty(x), z = x; return; end
  1790. medx = median(x,'omitnan');
  1791. madx = mad(x,1);
  1792. z = (x - medx) ./ max(madx, eps);
  1793. end
  1794. %% Classify reaches manually
  1795. function classifyReachesCallback(fig)
  1796. handles = guidata(fig);
  1797. handles.msgLabel.Text = 'Starting Classification Pipeline in FIJI.';
  1798. handles.msgLabel.FontColor = handles.colors.statusPending;
  1799. drawnow;
  1800. fijiPath = 'C:\Fiji.app\fiji-win64.exe';
  1801. dataDir = handles.baseDir;
  1802. % Double-escape backslashes for the dir argument
  1803. dirArg = sprintf('dir=%s', strrep(dataDir, '\', '\\'));
  1804. % Build Fiji system call
  1805. cmd = sprintf('"%s" --ij2 --run "SPG_ClassifyReaches " "%s"', fijiPath, dirArg);
  1806. system(cmd);
  1807. handles.msgLabel.Text = 'Classification pipeline shutdown. After classifying all, continue with analysis';
  1808. handles.msgLabel.FontColor = handles.colors.successColor;
  1809. drawnow;
  1810. return
  1811. end
  1812. %% Calibration
  1813. function calibratePoleByClick(fig)
  1814. % CALIBRATEPOLEBYCLICK Interactively calibrates pole widths in side/front videos.
  1815. % Assumes known pole width = 12.7 mm and computes px, mm/px, and mm.
  1816. % Saves results to pole_calibration.mat
  1817. % --- Setup ---
  1818. handles = guidata(fig);
  1819. baseDir = handles.baseDir;
  1820. outDir = handles.outDir;
  1821. pairs = handles.pairs;
  1822. known_mm = 12.7; % Known physical pole width in mm
  1823. calibrationFile = fullfile(outDir, 'pole_calibration.mat');
  1824. if isfile(calibrationFile)
  1825. load(calibrationFile, 'poleCal');
  1826. else
  1827. poleCal = struct();
  1828. end
  1829. % --- Loop through animals ---
  1830. for idx = 1:numel(pairs)
  1831. coreID = pairs(idx).coreID;
  1832. fprintf('\n📏 Calibrating pole width for %s...\n', coreID);
  1833. sideVid = pairs(idx).sideVideo;
  1834. frontVid = pairs(idx).frontVideo;
  1835. data = struct();
  1836. % === SIDE VIDEO ===
  1837. if isfile(sideVid)
  1838. v = VideoReader(sideVid);
  1839. frameNum = 50; %load frame nr 50 in case video starts blacl
  1840. for i = 1:frameNum
  1841. f = readFrame(v);
  1842. end
  1843. figure('Name', sprintf('%s - SIDE View', coreID));
  1844. imshow(f); axis on; hold on;
  1845. title('SIDE view: Click LEFT and RIGHT edges of pole');
  1846. [x, y] = getTwoClicks();
  1847. px_width = abs(x(2) - x(1));
  1848. mm_per_pixel = known_mm / px_width;
  1849. data.side.px_width = px_width;
  1850. data.side.mm_per_pixel = mm_per_pixel;
  1851. data.side.mm_width = px_width * mm_per_pixel;
  1852. close;
  1853. else
  1854. warning('⚠️ Side video for %s not found.', coreID);
  1855. end
  1856. % === FRONT VIDEO ===
  1857. if isfile(frontVid)
  1858. v = VideoReader(frontVid);
  1859. f = readFrame(v);
  1860. figure('Name', sprintf('%s - FRONT View', coreID));
  1861. imshow(f); axis on; hold on;
  1862. title('FRONT view: Click LEFT and RIGHT edges of pole');
  1863. [x, y] = getTwoClicks();
  1864. px_width = abs(x(2) - x(1));
  1865. mm_per_pixel = known_mm / px_width;
  1866. data.front.px_width = px_width;
  1867. data.front.mm_per_pixel = mm_per_pixel;
  1868. data.front.mm_width = px_width * mm_per_pixel;
  1869. close;
  1870. else
  1871. warning('⚠️ Front video for %s not found.', coreID);
  1872. end
  1873. field = matlab.lang.makeValidName(coreID);
  1874. poleCal.(field) = data;
  1875. save(calibrationFile, 'poleCal');
  1876. side_px = NaN;
  1877. front_px = NaN;
  1878. if isfield(data, 'side') && isfield(data.side, 'px_width')
  1879. side_px = data.side.px_width;
  1880. end
  1881. if isfield(data, 'front') && isfield(data.front, 'px_width')
  1882. front_px = data.front.px_width;
  1883. end
  1884. fprintf('✅ Side: %.1f px | Front: %.1f px\n', side_px, front_px);
  1885. end
  1886. % --- Finalize ---
  1887. handles.msgLabel.Text = 'Pole calibration completed!';
  1888. handles.msgLabel.FontColor = handles.colors.successColor;
  1889. guidata(fig, handles);
  1890. end
  1891. % === Helper: Clicks with feedback ===
  1892. function [x, y] = getTwoClicks()
  1893. x = zeros(1,2); y = zeros(1,2);
  1894. for i = 1:2
  1895. [x(i), y(i)] = ginput(1);
  1896. plot(x(i), y(i), 'ro', 'MarkerSize', 10, 'LineWidth', 2);
  1897. drawnow;
  1898. end
  1899. plot(x, y, 'r--', 'LineWidth', 2); % Connect clicks
  1900. pause(0.2);
  1901. end
  1902. %% -------- Analysis
  1903. %interpolate outliers from trajectory (not many but some) - if in
  1904. %completely different space! - average trace needs some more thinking
  1905. %about...
  1906. function analyzeReachesCallback(fig)
  1907. handles = guidata(fig);
  1908. outDir = handles.outDir;
  1909. pairs = handles.pairs;
  1910. handles.msgLabel.Text = 'Starting analysis...';
  1911. handles.msgLabel.FontColor = handles.colors.statusPending;
  1912. drawnow;
  1913. % --- Exclusions ---
  1914. excludeFile = fullfile(outDir,'exclude_table.csv');
  1915. if isfile(excludeFile)
  1916. T_excl = readtable(excludeFile);
  1917. % Only keep rows marked with Exclude == 1
  1918. if any(strcmpi(T_excl.Properties.VariableNames,'Exclude'))
  1919. excludeCoreIDs = T_excl.CoreID(T_excl.Exclude == 1);
  1920. else
  1921. warning('Exclude column not found, excluding none.');
  1922. excludeCoreIDs = {};
  1923. end
  1924. elseif isfile(fullfile(outDir,'exclude_table.mat'))
  1925. S_excl = load(fullfile(outDir,'exclude_table.mat'));
  1926. T_excl = S_excl.exclude_table;
  1927. if any(strcmpi(T_excl.Properties.VariableNames,'Exclude'))
  1928. excludeCoreIDs = T_excl.CoreID(T_excl.Exclude == 1);
  1929. else
  1930. warning('Exclude column not found in MAT, excluding none.');
  1931. excludeCoreIDs = {};
  1932. end
  1933. else
  1934. warning('No exclusion file found, not excluding any animals.');
  1935. excludeCoreIDs = {};
  1936. end
  1937. animalData = struct([]);
  1938. %% ----------- Pass 1: Build animalData & collect global info ----------
  1939. fprintf('[%s] Starting Pass 1 (building animalData)\n', datestr(now,'HH:MM:SS.FFF'));
  1940. % Temporary storage (parfor-safe)
  1941. animalDataCell = cell(numel(pairs),1);
  1942. allGroupsScan = cell(numel(pairs),1);
  1943. allTestDaysScan = cell(numel(pairs),1);
  1944. allLabelsCollect= cell(numel(pairs),1);
  1945. NewLabelMap = containers.Map( ...
  1946. string({ "Attempt - No Touch", "Miss - Targeting", "Miss - Knock", ...
  1947. "Error - During Grasp", "Error - Retrieve Failure", "Success After Many" }), ...
  1948. string({ "Error - Approach", "Error - Approach", "Error - Approach", ...
  1949. "Error - Grasp", "Error - Retrieval", "Success" }));
  1950. parfor i = 1:numel(pairs)
  1951. coreID = pairs(i).coreID;
  1952. % Skip excluded IDs
  1953. if ismember(coreID, excludeCoreIDs)
  1954. fprintf('Excluding %s\n', coreID);
  1955. continue;
  1956. end
  1957. % --- Load reachLabels.csv ---
  1958. labelFile = fullfile(outDir, sprintf('%s_reachLabels.csv', coreID));
  1959. if ~isfile(labelFile)
  1960. fprintf('Missing reachLabels for %s\n', coreID);
  1961. continue;
  1962. end
  1963. T_labels = readtable(labelFile);
  1964. T_labels.Label = strrep(T_labels.Label,"–","-");
  1965. if ~ismember('Label', T_labels.Properties.VariableNames)
  1966. warning('File %s missing Label column, skipping.', coreID);
  1967. continue;
  1968. end
  1969. % rename labels to new error approach / grasp / retrieve
  1970. oldLabels = string(T_labels.Label);
  1971. newLabels = oldLabels;
  1972. for k = 1:numel(oldLabels)
  1973. if NewLabelMap.isKey(oldLabels(k))
  1974. newLabels(k) = NewLabelMap(oldLabels(k));
  1975. end
  1976. end
  1977. T_labels.Label = newLabels;
  1978. % --- Exclude unwanted labels ---
  1979. excludeLabels = ["Skip (Not a Reach)", "Attempt - No Pellet", "Unknown / Hard to Say"];
  1980. T_labels(ismember(string(T_labels.Label), excludeLabels), :) = [];
  1981. % --- DLC load ---
  1982. dlc_side = readDLCcsv(pairs(i).sideVideo,'Side');
  1983. dlc_front = readDLCcsv(pairs(i).frontVideo,'Front');
  1984. if isempty(dlc_side) || isempty(dlc_front)
  1985. warning('Missing DLC data for %s', coreID);
  1986. continue;
  1987. end
  1988. % --- Calibration ---
  1989. poleCalibrationStruct = load(fullfile(outDir,'pole_calibration.mat'));
  1990. poleCalibrationStruct = poleCalibrationStruct.poleCal;
  1991. coreID_field = strrep(coreID,'-','_');
  1992. if ~isfield(poleCalibrationStruct, coreID_field)
  1993. warning('No pole calibration for %s', coreID);
  1994. continue;
  1995. end
  1996. poleCalibration = poleCalibrationStruct.(coreID_field);
  1997. % Side calibration
  1998. mmPerPixel_side = poleCalibration.side.mm_per_pixel;
  1999. varNames = dlc_side.Properties.VariableNames;
  2000. xyCols = contains(varNames,'_x') | contains(varNames,'_y');
  2001. for c = find(xyCols)
  2002. dlc_side.(varNames{c}) = dlc_side.(varNames{c}) * mmPerPixel_side;
  2003. end
  2004. % Front calibration (fallback to side)
  2005. if isfield(poleCalibration,'front') && isfield(poleCalibration.front,'mm_per_pixel')
  2006. mmPerPixel_front = poleCalibration.front.mm_per_pixel;
  2007. else
  2008. mmPerPixel_front = mmPerPixel_side;
  2009. end
  2010. varNamesF = dlc_front.Properties.VariableNames;
  2011. xyColsF = contains(varNamesF,'_x') | contains(varNamesF,'_y');
  2012. for c = find(xyColsF)
  2013. dlc_front.(varNamesF{c}) = dlc_front.(varNamesF{c}) * mmPerPixel_front;
  2014. end
  2015. % --- Parse metadata ---
  2016. [group, animal, test_day] = parseCoreID(coreID);
  2017. % --- Build struct ---
  2018. s = struct( ...
  2019. 'coreID', coreID, ...
  2020. 'group', group, ...
  2021. 'animal', animal, ...
  2022. 'test_day', test_day, ...
  2023. 'reaches', T_labels, ...
  2024. 'offsetVals', T_labels.Offset, ...
  2025. 'dlc_side', dlc_side, ...
  2026. 'dlc_front', dlc_front ...
  2027. );
  2028. % Save into cell
  2029. animalDataCell{i} = s;
  2030. allGroupsScan{i} = group;
  2031. allTestDaysScan{i} = test_day;
  2032. allLabelsCollect{i} = unique(string(T_labels.Label));
  2033. end
  2034. % Collapse cells → struct array
  2035. animalData = [animalDataCell{~cellfun('isempty',animalDataCell)}];
  2036. allGroupsScan = allGroupsScan(~cellfun('isempty',allGroupsScan));
  2037. allTestDaysScan = allTestDaysScan(~cellfun('isempty',allTestDaysScan));
  2038. allLabelsCollect= vertcat(allLabelsCollect{~cellfun('isempty',allLabelsCollect)});
  2039. % Unique lists
  2040. uniqueGroups = unique(allGroupsScan,'stable');
  2041. uniqueTestDays = unique(allTestDaysScan,'stable');
  2042. uniqueLabelsGlobal = unique(allLabelsCollect,'stable');
  2043. % Precompute fast maps for lookups
  2044. groupMap = containers.Map(uniqueGroups, 1:numel(uniqueGroups));
  2045. dayMap = containers.Map(uniqueTestDays, 1:numel(uniqueTestDays));
  2046. fprintf('[%s] Pass 1 complete: %d animals retained\n', ...
  2047. datestr(now,'HH:MM:SS.FFF'), numel(animalData));
  2048. nGroups = numel(uniqueGroups);
  2049. nTestdays = numel(uniqueTestDays);
  2050. nLabels = numel(uniqueLabelsGlobal);
  2051. % ---------- Pass 2: Analyze and accumulate heatmaps ----------
  2052. for i = 1:numel(animalData)
  2053. coreID = animalData(i).coreID;
  2054. group = animalData(i).group;
  2055. animal = animalData(i).animal;
  2056. test_day = animalData(i).test_day;
  2057. coreID = animalData(i).coreID;
  2058. reaches = animalData(i).reaches;
  2059. dlc_side = animalData(i).dlc_side;
  2060. dlc_front = animalData(i).dlc_front;
  2061. offsets = animalData(i).offsetVals;
  2062. % Thresholds
  2063. likelihoodThresh_side = 0.3;
  2064. likelihoodThresh_front = 0.2;
  2065. % --- progress bar update ---
  2066. nBlocks = 20;
  2067. progress = i / numel(animalData);
  2068. filledBlocks = round(progress * nBlocks);
  2069. barStr = [repmat('█', 1, filledBlocks), repmat('░', 1, nBlocks - filledBlocks)];
  2070. handles.msgLabel.Text = sprintf('Computing trajectories for animal %d of %d [%s]...', ...
  2071. i, numel(animalData), barStr);
  2072. handles.msgLabel.FontColor = handles.colors.statusPending;
  2073. drawnow;
  2074. % Storage for this animal
  2075. perReachRows = table(); % merged per-reach metrics
  2076. allFields = {}; % running superset of metrics
  2077. perReachTrajSide = struct('coreID', {}, 'reachID', {}, 'label', {}, 'broadLabel', {}, 'traj', {});
  2078. perReachTrajFront = struct('coreID', {}, 'reachID', {}, 'label', {}, 'broadLabel', {}, 'traj', {});
  2079. failedReaches = table('Size',[0 5], ...
  2080. 'VariableTypes', {'string','double','string','string','string'}, ...
  2081. 'VariableNames', {'coreID','reachID','label','view','reason'});
  2082. % Skip animals with too few reaches
  2083. if height(reaches) < 2
  2084. warning('Skipping %s: not enough reaches (%d).', coreID, height(reaches));
  2085. continue;
  2086. end
  2087. for r = 1:height(reaches)
  2088. reachID = reaches.ReachIndex(r);
  2089. label = string(reaches.Label(r));
  2090. % --- Broad label grouping (success vs error only) ---
  2091. if startsWith(lower(label), "success")
  2092. broadLabel = "success";
  2093. else
  2094. broadLabel = "error";
  2095. end
  2096. % --- SIDE ---
  2097. [sideMetrics, sideTraj, failReason_side, sStartRef, sEndRef] = ...
  2098. analyzeReach_Side(coreID, reaches(r,:), dlc_side, likelihoodThresh_side, 'tip', outDir);
  2099. if strlength(failReason_side) > 0
  2100. failedReaches = [failedReaches; {coreID, reachID, label, "side", failReason_side}];
  2101. end
  2102. % --- FRONT ---
  2103. [frontMetrics, frontTraj, failReason_front] = ...
  2104. analyzeReach_Front(coreID, reaches(r,:), dlc_front, offsets(r), likelihoodThresh_front, sStartRef, sEndRef);
  2105. if strlength(failReason_front) > 0
  2106. failedReaches = [failedReaches; {coreID, reachID, label, "front", failReason_front}];
  2107. end
  2108. %save CSV for traj in R
  2109. % Paths to source CSVs (handy in R)
  2110. sideCSVPath = pairs(i).sideVideo; % absolute or relative, as stored
  2111. frontCSVPath = pairs(i).frontVideo;
  2112. labelsCSVPath = fullfile(outDir, sprintf('%s_reachLabels.csv', coreID));
  2113. % Save per-reach smoothed trajectories and remember the file paths
  2114. sideTrajPath = "";
  2115. frontTrajPath = "";
  2116. if ~isempty(sideTraj) && ~isempty(fieldnames(sideTraj))
  2117. sideTrajPath = saveTrajectoryCSV(sideTraj, coreID, 'Side', outDir);
  2118. end
  2119. if ~isempty(frontTraj) && ~isempty(fieldnames(frontTraj))
  2120. frontTrajPath = saveTrajectoryCSV(frontTraj, coreID, 'Front', outDir);
  2121. end
  2122. % derive start/end frames (cropped) for convenience columns
  2123. sideStartFrame = NaN; sideEndFrame = NaN;
  2124. frontStartFrame = NaN; frontEndFrame = NaN;
  2125. if ~isempty(sideTraj) && isfield(sideTraj,'frames') && ~isempty(sideTraj.frames)
  2126. sideStartFrame = sideTraj.frames(1);
  2127. sideEndFrame = sideTraj.frames(end);
  2128. end
  2129. if ~isempty(frontTraj) && isfield(frontTraj,'frames') && ~isempty(frontTraj.frames)
  2130. frontStartFrame = frontTraj.frames(1);
  2131. frontEndFrame = frontTraj.frames(end);
  2132. end
  2133. % --- Merge metrics into one table row ---
  2134. metaStruct = struct( ...
  2135. 'coreID', coreID, ...
  2136. 'group', group, ...
  2137. 'animal', animal, ...
  2138. 'test_day', test_day, ...
  2139. 'label', label, ...
  2140. 'broadLabel', broadLabel, ...
  2141. 'reachID', reachID, ...
  2142. 'sideStartFrame', sideStartFrame, ...
  2143. 'sideEndFrame', sideEndFrame, ...
  2144. 'frontStartFrame', frontStartFrame, ...
  2145. 'frontEndFrame', frontEndFrame, ...
  2146. 'sideTrajCSV', string(sideTrajPath), ...
  2147. 'frontTrajCSV', string(frontTrajPath), ...
  2148. 'sideCSVPath', string(sideCSVPath), ...
  2149. 'frontCSVPath', string(frontCSVPath), ...
  2150. 'reachLabelsCSV', string(labelsCSVPath));
  2151. [perReachRows, allFields] = joinStructsFlexible( ...
  2152. perReachRows, sideMetrics, frontMetrics, metaStruct, allFields);
  2153. if ~isempty(sideTraj) && ~isempty(fieldnames(sideTraj))
  2154. perReachTrajSide(end+1) = struct( ...
  2155. 'coreID', coreID, ...
  2156. 'reachID', reachID, ...
  2157. 'label', label, ...
  2158. 'broadLabel', broadLabel, ...
  2159. 'traj', sideTraj ...
  2160. );
  2161. end
  2162. if ~isempty(frontTraj) && ~isempty(fieldnames(frontTraj))
  2163. perReachTrajFront(end+1) = struct( ...
  2164. 'coreID', coreID, ...
  2165. 'reachID', reachID, ...
  2166. 'label', label, ...
  2167. 'broadLabel', broadLabel, ...
  2168. 'traj', frontTraj ...
  2169. );
  2170. end
  2171. % --- debug ---
  2172. if strlength(failReason_side) > 0
  2173. fprintf('[DEBUG] %s Reach %d (%s): side fail (%s)\n', ...
  2174. coreID, reachID, label, failReason_side);
  2175. end
  2176. if strlength(failReason_front) > 0
  2177. fprintf('[DEBUG] %s Reach %d (%s): front fail (%s)\n', ...
  2178. coreID, reachID, label, failReason_front);
  2179. end
  2180. end
  2181. % ---- Save per animal ----
  2182. csvDir = fullfile(outDir,'CSV');
  2183. if ~exist(csvDir,'dir'), mkdir(csvDir); end
  2184. writetable(perReachRows, fullfile(csvDir, sprintf('Params_%s.csv', coreID)));
  2185. matDir = fullfile(outDir,'MAT');
  2186. if ~exist(matDir,'dir'), mkdir(matDir); end
  2187. results = struct( ...
  2188. 'coreID', coreID, ...
  2189. 'group', group, ...
  2190. 'animal', animal, ...
  2191. 'test_day', test_day, ...
  2192. 'perReachRows', perReachRows, ...
  2193. 'failedReaches', failedReaches, ...
  2194. 'side', struct( ...
  2195. 'metrics', perReachRows(:, contains(perReachRows.Properties.VariableNames,'side','IgnoreCase',true)), ...
  2196. 'trajectories', perReachTrajSide ...
  2197. ), ...
  2198. 'front', struct( ...
  2199. 'metrics', perReachRows(:, contains(perReachRows.Properties.VariableNames,'front','IgnoreCase',true)), ...
  2200. 'trajectories', perReachTrajFront ...
  2201. ) ...
  2202. );
  2203. save(fullfile(matDir, sprintf('%s_results.mat', coreID)), 'results','-v7.3');
  2204. fprintf('[%s] Finished %s (%d reaches, %d fails)\n', ...
  2205. datestr(now,'HH:MM:SS.FFF'), coreID, height(reaches), height(failedReaches));
  2206. end
  2207. handles.msgLabel.Text = sprintf('Computation done and saved, moving on to visualization');
  2208. handles.msgLabel.FontColor = handles.colors.statusPending;
  2209. drawnow;
  2210. %% Pass 3 - Trajectory and Heatmap
  2211. % -------- Load MAT files --------
  2212. matDir = fullfile(outDir,'MAT');
  2213. figDir = fullfile(outDir,'FIG');
  2214. if ~exist(figDir,'dir'), mkdir(figDir); end
  2215. matFiles = dir(fullfile(matDir,'*_results.mat'));
  2216. if isempty(matFiles)
  2217. error('No *_results.mat files found in %s',matDir);
  2218. end
  2219. allResults = cell(numel(matFiles),1);
  2220. for i = 1:numel(matFiles)
  2221. tmp = load(fullfile(matDir,matFiles(i).name));
  2222. if isfield(tmp,'results')
  2223. allResults{i} = tmp.results;
  2224. else
  2225. warning('File %s missing results struct, skipping',matFiles(i).name);
  2226. end
  2227. end
  2228. allResults = [allResults{~cellfun('isempty',allResults)}];
  2229. unifiedParams = vertcat(allResults.perReachRows); % now consistent
  2230. % -------- Merge parameter tables into one Masterfile--------
  2231. csvDir = fullfile(outDir,'CSV');
  2232. writetable(unifiedParams, fullfile(csvDir,'All_Params.csv'));
  2233. handles.msgLabel.Text = sprintf('Master Parameter table (Front/Side) saved in /CSV.%sProceeding with visualization', newline);
  2234. handles.msgLabel.FontColor = handles.colors.statusPending;
  2235. drawnow;
  2236. % Collect metadata
  2237. allGroups = string({allResults.group});
  2238. allDays = string({allResults.test_day});
  2239. allCoreIDs = string({allResults.coreID});
  2240. uniqueGroups = unique(allGroups,'stable');
  2241. uniqueDays = unique(allDays,'stable');
  2242. uniqueLabelsGlobal = unique(unifiedParams.label,'stable');
  2243. % --- enforce custom ordering (optional) --- (I dont think this is being
  2244. % used in the functions but thats ok)
  2245. desiredLabelOrder = ["Success","Error - Approach","Error - Grasp", "Error - Retrieval"];
  2246. desiredDayOrder = ["Baseline","Drug","Washout"];
  2247. % Reorder groups
  2248. [~, idxG] = ismember(desiredLabelOrder, uniqueLabelsGlobal);
  2249. idxG = idxG(idxG>0); % keep only those that exist
  2250. uniqueLabelsGlobal = uniqueLabelsGlobal(idxG);
  2251. % Reorder days
  2252. [~, idxD] = ismember(desiredDayOrder, uniqueDays);
  2253. idxD = idxD(idxD>0); % keep only those that exist
  2254. uniqueDays = uniqueTestDays(idxD);
  2255. %
  2256. % % -------- Per animal plots --------
  2257. % for i = 1:numel(allResults)
  2258. % R = allResults(i);
  2259. % coreID = R.coreID;
  2260. %
  2261. % % SIDE trajectories per label
  2262. % if isfield(R,'side') && ~isempty(R.side.trajectories)
  2263. % plotPerAnimalTraj(R.side.trajectories,coreID,'Side',figDir, 'label');
  2264. % plotPerAnimalHeatmap(R.side.trajectories,coreID,'Side',figDir, 'label');
  2265. % plotPerAnimalTraj(R.side.trajectories,coreID,'Side',figDir, 'broadLabel');
  2266. % plotPerAnimalHeatmap(R.side.trajectories,coreID,'Side',figDir, 'broadLabel');
  2267. % end
  2268. %
  2269. % % FRONT trajectories per label
  2270. % if isfield(R,'front') && ~isempty(R.front.trajectories)
  2271. % plotPerAnimalTraj(R.front.trajectories,coreID,'Front',figDir, 'label');
  2272. % plotPerAnimalHeatmap(R.front.trajectories,coreID,'Front',figDir, 'label');
  2273. % plotPerAnimalTraj(R.front.trajectories,coreID,'Front',figDir, 'broadLabel');
  2274. % plotPerAnimalHeatmap(R.front.trajectories,coreID,'Front',figDir, 'broadLabel');
  2275. % end
  2276. % end
  2277. %
  2278. % handles.msgLabel.Text = sprintf([ ...
  2279. % 'Per-animal trajectories and heatmaps complete.' newline ...
  2280. % 'Proceeding with group-level visualizations...' ]);
  2281. % handles.msgLabel.FontColor = handles.colors.statusPending;
  2282. % drawnow;
  2283. % -------- Group-level heatmaps (side + front) --------
  2284. for g = 1:numel(uniqueGroups)
  2285. grpName = uniqueGroups(g);
  2286. grpMask = strcmp(allGroups,grpName);
  2287. Rgrp = allResults(grpMask);
  2288. if ~isempty(Rgrp)
  2289. % plotGroupHeatmaps(Rgrp, grpName, 'Side', figDir, 'label');
  2290. % plotGroupHeatmaps(Rgrp, grpName, 'Side', figDir, 'broadLabel');
  2291. % plotGroupHeatmaps(Rgrp, grpName, 'Front', figDir,'label');
  2292. % plotGroupHeatmaps(Rgrp, grpName, 'Front', figDir,'broadLabel');
  2293. % plotGroupTraj(Rgrp, grpName, 'Side', figDir, 'label');
  2294. % plotGroupTraj(Rgrp, grpName, 'Side', figDir, 'broadLabel');
  2295. % plotGroupTraj(Rgrp, grpName, 'Front', figDir, 'label');
  2296. % plotGroupTraj(Rgrp, grpName, 'Front', figDir, 'broadLabel');
  2297. end
  2298. end
  2299. handles.msgLabel.Text = sprintf([ ...
  2300. 'Group-level heatmaps complete.' newline ...
  2301. 'Proceeding with global label × group × day visualizations...' ]);
  2302. handles.msgLabel.FontColor = handles.colors.statusPending;
  2303. drawnow;
  2304. % -------- Global heatmaps grid (rows=test_day, cols=group) --------
  2305. allParams = vertcat(allResults.perReachRows);
  2306. uniqueLabelsGlobal = unique(allParams.label, 'stable');
  2307. uniqueBroadGlobal = unique(allParams.broadLabel, 'stable');
  2308. % % label-level
  2309. % for l = uniqueLabelsGlobal'
  2310. % plotGlobalLabelHeatmaps(allResults, uniqueGroups, uniqueDays, l, 'Side', figDir, 'label');
  2311. % plotGlobalLabelHeatmaps(allResults, uniqueGroups, uniqueDays, l, 'Front', figDir, 'label');
  2312. % end
  2313. %
  2314. %
  2315. % % broadLabel-level
  2316. % for bl = uniqueBroadGlobal'
  2317. % plotGlobalLabelHeatmaps(allResults, uniqueGroups, uniqueDays, bl, 'Side', figDir, 'broadLabel');
  2318. % plotGlobalLabelHeatmaps(allResults, uniqueGroups, uniqueDays, bl, 'Front', figDir, 'broadLabel');
  2319. % end
  2320. % label-level
  2321. plotDifferenceHeatmaps(allResults, uniqueGroups, uniqueDays, uniqueLabelsGlobal, [], 'Side', figDir, 'label');
  2322. plotDifferenceHeatmaps(allResults, uniqueGroups, uniqueDays, uniqueLabelsGlobal, [], 'Front', figDir, 'label');
  2323. % % broadLabel-level
  2324. % plotDifferenceHeatmaps(allResults, uniqueGroups, uniqueDays, [], uniqueBroadGlobal, 'Side', figDir, 'broadLabel');
  2325. % plotDifferenceHeatmaps(allResults, uniqueGroups, uniqueDays, [], uniqueBroadGlobal, 'Front', figDir, 'broadLabel');
  2326. %
  2327. %
  2328. % %over all reaches
  2329. % plotGlobalAndDifferenceHeatmaps(allResults, uniqueGroups, 'Side', figDir);
  2330. % plotGlobalAndDifferenceHeatmaps(allResults, uniqueGroups, 'Front', figDir);
  2331. handles.msgLabel.Text = '✅ Kinematic analysis complete, check output files in CSV and FIG folder.';
  2332. handles.msgLabel.FontColor = handles.colors.successColor;
  2333. guidata(fig,handles);
  2334. end
  2335. function [perReachRows, allFields] = joinStructsFlexible(perReachRows, sideMetrics, frontMetrics, metaStruct, allFields)
  2336. % ---- Step 1: merge metrics ----
  2337. unified = sideMetrics;
  2338. if ~isempty(frontMetrics)
  2339. f2 = fieldnames(frontMetrics);
  2340. for k = 1:numel(f2)
  2341. unified.(f2{k}) = frontMetrics.(f2{k});
  2342. end
  2343. end
  2344. % ---- Step 2: add metadata ----
  2345. unified.coreID = string(metaStruct.coreID);
  2346. unified.group = string(metaStruct.group);
  2347. unified.animal = string(metaStruct.animal);
  2348. unified.test_day = string(metaStruct.test_day);
  2349. unified.label = string(metaStruct.label);
  2350. unified.broadLabel = string(metaStruct.broadLabel);
  2351. unified.reachID = double(metaStruct.reachID);
  2352. % ---- Step 3: update running field list ----
  2353. fn = fieldnames(unified);
  2354. allFields = union(allFields, fn, 'stable');
  2355. % ---- Step 4: patch unified struct with NaN where needed ----
  2356. for k = 1:numel(allFields)
  2357. fld = allFields{k};
  2358. if ~isfield(unified,fld)
  2359. unified.(fld) = NaN;
  2360. end
  2361. end
  2362. % ---- Step 5: convert to table ----
  2363. rowT = struct2table(unified);
  2364. % ---- Step 6: harmonize existing table with new row ----
  2365. if ~isempty(perReachRows)
  2366. % Add missing vars to existing table
  2367. missingInExisting = setdiff(rowT.Properties.VariableNames, perReachRows.Properties.VariableNames);
  2368. for m = missingInExisting
  2369. perReachRows.(m{1}) = NaN(height(perReachRows),1);
  2370. end
  2371. % Add missing vars to new row
  2372. missingInNew = setdiff(perReachRows.Properties.VariableNames, rowT.Properties.VariableNames);
  2373. for m = missingInNew
  2374. rowT.(m{1}) = NaN(height(rowT),1);
  2375. end
  2376. % Reorder rowT to match perReachRows
  2377. rowT = rowT(:, perReachRows.Properties.VariableNames);
  2378. end
  2379. % ---- Step 7: append ----
  2380. perReachRows = [perReachRows; rowT];
  2381. end
  2382. function [metrics, traj, failReason, sideStartRefined, sideEndRefined] = ...
  2383. analyzeReach_Side(coreID, reachRow, dlc_side, likelihoodThresh, pawPart, outDir)
  2384. if nargin < 5 || isempty(pawPart), pawPart = 'tip'; end
  2385. if nargin < 6, outDir = ''; end
  2386. metrics = struct();
  2387. traj = struct();
  2388. failReason = "";
  2389. sideStartRefined = NaN;
  2390. sideEndRefined = NaN;
  2391. % ----------- basic guards -----------
  2392. if isempty(reachRow) || ~istable(reachRow) || height(reachRow)~=1
  2393. failReason = "reachRow must be a single table row";
  2394. return;
  2395. end
  2396. needVars = {'SideStart','SideEnd','Label'};
  2397. if ~all(ismember(needVars, reachRow.Properties.VariableNames))
  2398. failReason = "reachRow missing SideStart/SideEnd/Label";
  2399. return;
  2400. end
  2401. startF = double(reachRow.SideStart);
  2402. endF = double(reachRow.SideEnd);
  2403. if isnan(startF) || isnan(endF) || endF <= startF
  2404. failReason = "invalid SideStart/SideEnd window";
  2405. return;
  2406. end
  2407. % normalize label (fix en-dash artifact & make string)
  2408. label = string(reachRow.Label);
  2409. % figure out reachID field
  2410. if ismember('ReachID', reachRow.Properties.VariableNames)
  2411. reachID = reachRow.ReachID;
  2412. elseif ismember('ReachIndex', reachRow.Properties.VariableNames)
  2413. reachID = reachRow.ReachIndex;
  2414. else
  2415. reachID = startF; % fallback (not ideal, but stable)
  2416. end
  2417. % --- Paw selection ---
  2418. switch lower(pawPart)
  2419. case 'center'
  2420. xAll = dlc_side.paw_center__x;
  2421. yAll = dlc_side.paw_center__y;
  2422. likelihoods = dlc_side.paw_center__likelihood;
  2423. case 'tip'
  2424. xAll = dlc_side.paw_tip__x;
  2425. yAll = dlc_side.paw_tip__y;
  2426. likelihoods = dlc_side.paw_tip__likelihood;
  2427. otherwise
  2428. error('Unknown pawPart: choose either "center" or "tip".');
  2429. end
  2430. % likelihood filtering → NaN
  2431. xAll(likelihoods < likelihoodThresh) = NaN;
  2432. yAll(likelihoods < likelihoodThresh) = NaN;
  2433. %jumping outliers)
  2434. vel = hypot(diff(xAll), diff(yAll));
  2435. zscoreVel = (vel - mean(vel,'omitnan'))/std(vel,'omitnan');
  2436. outlierIdx = [false; abs(zscoreVel) > 5]; % mark crazy jumps
  2437. xAll(outlierIdx) = NaN;
  2438. yAll(outlierIdx) = NaN;
  2439. %filling in those gaps
  2440. maxGap = 5; % frames
  2441. xAll = fillmissing(xAll, 'linear', 'MaxGap',maxGap,'EndValues','nearest');
  2442. yAll = fillmissing(yAll, 'linear', 'MaxGap',maxGap,'EndValues','nearest');
  2443. nFrames = height(dlc_side);
  2444. frames = max(1,startF):min(endF,nFrames);
  2445. frames = frames(frames>0 & frames<=nFrames);
  2446. if isempty(frames)
  2447. failReason = "empty clipped frame window";
  2448. return;
  2449. end
  2450. % slit mean position (mm)
  2451. slitX = mean([mean(dlc_side.slit_bottom__x,'omitnan'), mean(dlc_side.slit_top__x,'omitnan')]);
  2452. % Parameters for segmentation
  2453. par.frame_buffer = 100; %to extend trajectory if necessary
  2454. par.rightwardVelocityThreshold = 2;
  2455. par.minSustainFrames = 3;
  2456. par.velocityStopThresh = 1;
  2457. par.windowLen = 3;
  2458. par.minPeakDistance = 40; %Only peaks at least 10 frames apart will be considered separate; closer peaks are merged.
  2459. par.minPeakProminence = 1; %A detected peak must rise at least 0.5 mm above surrounding troughs to be considered valid.
  2460. par.peakTolerance = 1.0; % mm tolerance for equivalent peaks
  2461. par.showQCplot = false;
  2462. par.gauss_smooth = 20;
  2463. par.pelletContactTolerance = 2;
  2464. if ~isfield(par,'plateauTol'), par.plateauTol = 1.0; end % mm range above min-dist to count as plateau
  2465. if ~isfield(par,'boundarySlack'), par.boundarySlack = 2; end % +/- frames to pad the crop
  2466. % ----------- extract/crop, compute metrics -----------
  2467. try
  2468. [xPlot, yPlot, pelletXreach, pelletYreach, slitX_norm, m, sStart, sEnd, croppedFrames] = ...
  2469. extractReachTrajectory(frames, xAll, yAll, dlc_side, slitX, par, reachID, coreID, outDir, label);
  2470. sideStartRefined = sStart;
  2471. sideEndRefined = sEnd;
  2472. if isempty(xPlot)
  2473. failReason = "Empty trajectory after extraction";
  2474. return;
  2475. end
  2476. % Ensure required scalar fields are present (mirror your original guards)
  2477. scalarFields = {'peakOutwardVelocity','meanOutwardVelocity','outwardMovementDuration', ...
  2478. 'initialReachAngle','y_at_1_3','y_at_2_3','y_at_pellet','corrections','pauseCount', ...
  2479. 'totalPauseDuration','retrievalDuration','peakRetrievalVelocity','meanRetrievalVelocity', ...
  2480. 'retrievalStraightness','endpoint_x','endpoint_y','timeSlitToContact_frames', ...
  2481. 'timeSlitToContact_sec','slitToPelletDistance_mm','normReachTime_s_per_mm','meanOutwardSpeed_mm_per_s', ...
  2482. 'peakAcceleration','peakDeceleration','meanAbsJerk','peakJerk','accelSignChanges', ...
  2483. 'trajectoryLength_outward','retrievalArc','retrievalReversals','reachDuration','maxHeight', ...
  2484. 'trajectoryLength','trajectoryStraightness','pathTortuosity','nPelletContactPeaks','contactDuration'};
  2485. for f = 1:numel(scalarFields)
  2486. fld = scalarFields{f};
  2487. if ~isfield(m,fld) || isempty(m.(fld))
  2488. m.(fld) = NaN;
  2489. end
  2490. end
  2491. if ~isfield(m,'pelletContactPeakFrames') || ~iscell(m.pelletContactPeakFrames)
  2492. m.pelletContactPeakFrames = {[]};
  2493. end
  2494. if ~isfield(m,'peakOutwardVelocityFrame'), m.peakOutwardVelocityFrame = NaN; end
  2495. if ~isfield(m,'timeToPelletContact'), m.timeToPelletContact = m.timeSlitToContact_frames; end
  2496. % ----------- build outputs -----------
  2497. metrics = m; % full struct from extractReachTrajectory
  2498. traj = struct( ...
  2499. 'reachID', reachID, ...
  2500. 'label', label, ...
  2501. 'frames', croppedFrames, ...
  2502. 'x', xPlot, ...
  2503. 'y', yPlot, ...
  2504. 'pelletX', pelletXreach, ...
  2505. 'pelletY', pelletYreach, ...
  2506. 'slitX_norm', slitX_norm, ...
  2507. 'sideStartRefined', sideStartRefined, ...
  2508. 'sideEndRefined', sideEndRefined ...
  2509. );
  2510. catch ME
  2511. failReason = sprintf('error: %s', ME.message);
  2512. end
  2513. end
  2514. function outPath = saveTrajectoryCSV(traj, coreID, viewStr, outDir)
  2515. % Save a per-reach trajectory CSV with global frames + smoothed x/y
  2516. % Returns full file path.
  2517. trajDir = fullfile(outDir,'CSV','Trajectories',viewStr);
  2518. if ~exist(trajDir,'dir'), mkdir(trajDir); end
  2519. fn = sprintf('%s_reach%04d_%s_traj.csv', coreID, round(double(traj.reachID)), lower(viewStr));
  2520. outPath = fullfile(trajDir, fn);
  2521. % build table: one row per sample in the cropped, smoothed trajectory
  2522. T = table( ...
  2523. repmat(string(coreID), numel(traj.frames), 1), ...
  2524. repmat(double(traj.reachID), numel(traj.frames), 1), ...
  2525. repmat(string(viewStr), numel(traj.frames), 1), ...
  2526. traj.frames(:), ...
  2527. traj.x(:), ...
  2528. traj.y(:), ...
  2529. repmat(traj.pelletX, numel(traj.frames), 1), ...
  2530. repmat(traj.pelletY, numel(traj.frames), 1), ...
  2531. repmat(traj.slitX_norm, numel(traj.frames), 1), ...
  2532. 'VariableNames', {'coreID','reachID','view','frame','x_mm','y_mm','pelletX_mm','pelletY_mm','slitX_norm_mm'} ...
  2533. );
  2534. writetable(T, outPath);
  2535. end
  2536. function [xPlot, yPlot, pelletXreach, pelletYreach, slitX_norm, metrics, sideStartFrame, sideEndFrame, croppedFrames] = ...
  2537. extractReachTrajectory(frames, xAll, yAll, dlc_side, slitX, par, reachID, coreID, outDir, label)
  2538. nFrames = length(xAll);
  2539. sideStartFrame = NaN;
  2540. sideEndFrame = NaN;
  2541. %% -----------------------
  2542. % PREPROCESSING
  2543. % -----------------------
  2544. extendedFrames = max(frames(1)-par.frame_buffer,1):min(frames(end)+par.frame_buffer,nFrames);
  2545. pawX = xAll(extendedFrames);
  2546. pawY = yAll(extendedFrames);
  2547. validMask = ~isnan(pawX) & ~isnan(pawY);
  2548. pawX = pawX(validMask);
  2549. pawY = pawY(validMask);
  2550. framesFiltered = extendedFrames(validMask);
  2551. if length(framesFiltered) < 10
  2552. return;
  2553. end
  2554. % Pellet position
  2555. pelletFrames = framesFiltered(dlc_side.pellet_likelihood(framesFiltered) > 0.8);
  2556. if ~isempty(pelletFrames)
  2557. pelletXreach = median(dlc_side.pellet_x(pelletFrames),'omitnan');
  2558. pelletYreach = median(dlc_side.pellet_y(pelletFrames),'omitnan');
  2559. else
  2560. pelletXreach = median(dlc_side.pellet_x(dlc_side.pellet_likelihood > 0.8),'omitnan');
  2561. pelletYreach = median(dlc_side.pellet_y(dlc_side.pellet_likelihood > 0.8),'omitnan');
  2562. end
  2563. slitX_norm = slitX - pelletXreach;
  2564. % Normalize paw coords
  2565. pawX_norm = pawX - pelletXreach;
  2566. pawY_norm = pawY - pelletYreach;
  2567. % Smooth
  2568. if length(pawX_norm) > 5
  2569. pawX_smooth = smoothdata(pawX_norm,'gaussian',par.gauss_smooth);
  2570. pawY_smooth = smoothdata(pawY_norm,'gaussian',par.gauss_smooth);
  2571. else
  2572. pawX_smooth = pawX_norm;
  2573. pawY_smooth = pawY_norm;
  2574. end
  2575. %% -----------------------
  2576. % GLOBAL CONTACT PHASE
  2577. % -----------------------
  2578. distToPellet_global = sqrt(pawX_smooth.^2 + pawY_smooth.^2);
  2579. % -----------------------
  2580. % CONTACT PHASE (based on X beyond pellet)
  2581. % -----------------------
  2582. contactMask_global = pawX_smooth >= 0; % paw passes pellet's x-position
  2583. contactRuns_global = bwconncomp(contactMask_global);
  2584. if contactRuns_global.NumObjects > 0
  2585. % keep each bout separately
  2586. pelletContactPhaseIdx_global = contactRuns_global.PixelIdxList;
  2587. % count individual bouts
  2588. nPeaks = contactRuns_global.NumObjects;
  2589. else
  2590. pelletContactPhaseIdx_global = {};
  2591. nPeaks = 0;
  2592. end
  2593. %% -----------------------
  2594. % REACH BOUNDARIES (global)
  2595. % -----------------------
  2596. if ~isempty(pelletContactPhaseIdx_global)
  2597. % --- pick bout depending on label ---
  2598. boutLengths = cellfun(@numel, pelletContactPhaseIdx_global);
  2599. if strcmpi(label,'Success') || strcmpi(label,'SuccessAfterMany')
  2600. chosenBoutIdx = numel(pelletContactPhaseIdx_global); % last bout
  2601. else
  2602. chosenBoutIdx = 1; % first bout
  2603. end
  2604. chosenBout_global = pelletContactPhaseIdx_global{chosenBoutIdx};
  2605. else
  2606. % --- fallback: no contact bouts found ---
  2607. [~, peakIdx] = max(pawX_smooth);
  2608. chosenBout_global = peakIdx; % treat peak as 1-frame bout
  2609. end
  2610. % Convenience handles
  2611. firstContactIdx = chosenBout_global(1); % start of chosen bout
  2612. searchStart = chosenBout_global(end); % end of chosen bout
  2613. % --- slit hysteresis ---
  2614. if ~isfield(par,'slitHyst'), par.slitHyst = 0.5; end
  2615. slitLower = slitX_norm - par.slitHyst;
  2616. slitUpper = slitX_norm + par.slitHyst;
  2617. %% -----------------------
  2618. % REACH BOUNDARIES (bout-local)
  2619. % -----------------------
  2620. if ~isempty(pelletContactPhaseIdx_global)
  2621. % --- choose bout depending on label ---
  2622. if strcmpi(label,'Success') || strcmpi(label,'SuccessAfterMany')
  2623. chosenIdx = numel(pelletContactPhaseIdx_global); % last bout
  2624. else
  2625. chosenIdx = 1; % first bout
  2626. end
  2627. chosenBout = pelletContactPhaseIdx_global{chosenIdx};
  2628. firstContactIdx = chosenBout(1);
  2629. searchStart = chosenBout(end);
  2630. % --- define left/right windows around chosen bout ---
  2631. if chosenIdx > 1
  2632. leftBound = pelletContactPhaseIdx_global{chosenIdx-1}(end);
  2633. else
  2634. leftBound = 1;
  2635. end
  2636. if chosenIdx < numel(pelletContactPhaseIdx_global)
  2637. rightBound = pelletContactPhaseIdx_global{chosenIdx+1}(1);
  2638. else
  2639. rightBound = numel(pawX_smooth);
  2640. end
  2641. else
  2642. % --- fallback: no contact bouts at all ---
  2643. [~, peakIdx] = max(pawX_smooth);
  2644. chosenBout = peakIdx;
  2645. firstContactIdx = peakIdx;
  2646. searchStart = peakIdx;
  2647. leftBound = 1;
  2648. rightBound = numel(pawX_smooth);
  2649. end
  2650. % --- slit hysteresis thresholds ---
  2651. if ~isfield(par,'slitHyst'), par.slitHyst = 0.5; end
  2652. slitLower = slitX_norm - par.slitHyst;
  2653. slitUpper = slitX_norm + par.slitHyst;
  2654. %% Start boundary: only within [leftBound … firstContactIdx]
  2655. lastInside = find(pawX_smooth(leftBound:firstContactIdx) <= slitLower, 1, 'last');
  2656. if ~isempty(lastInside)
  2657. firstOutside = find(pawX_smooth(leftBound-1+lastInside:firstContactIdx) >= slitUpper, 1, 'first');
  2658. if ~isempty(firstOutside)
  2659. slitCrossStartIdx = (leftBound-1) + lastInside + firstOutside - 1;
  2660. else
  2661. slitCrossStartIdx = (leftBound-1) + lastInside;
  2662. end
  2663. else
  2664. % fallback = local minimum in that window
  2665. [~, relMin] = min(pawX_smooth(leftBound:firstContactIdx));
  2666. slitCrossStartIdx = (leftBound-1) + relMin;
  2667. end
  2668. %% End boundary: only within [searchStart … rightBound]
  2669. crossBackSlit = find(pawX_smooth(searchStart:rightBound) <= slitX_norm, 1, 'first');
  2670. if ~isempty(crossBackSlit)
  2671. reachEndIdx = searchStart + crossBackSlit - 1;
  2672. else
  2673. % fallback = local minimum in that window
  2674. [~, relMin] = min(pawX_smooth(searchStart:rightBound));
  2675. reachEndIdx = searchStart + relMin - 1;
  2676. end
  2677. %% Apply slack
  2678. reachStartIdx = max(1, slitCrossStartIdx - par.boundarySlack);
  2679. reachEndIdx = min(numel(pawX_smooth), reachEndIdx + par.boundarySlack);
  2680. %% -----------------------
  2681. % METRICS based on chosen bout
  2682. % -----------------------
  2683. metrics.nPelletContactPeaks = numel(pelletContactPhaseIdx_global);
  2684. if ~isempty(chosenBout_global)
  2685. metrics.yAtContact = pawY(chosenBout_global(1));
  2686. metrics.xErrorAtContact = pawX(chosenBout_global(1));
  2687. else
  2688. metrics.yAtContact = NaN;
  2689. metrics.xErrorAtContact = NaN;
  2690. end
  2691. %% -----------------------
  2692. % CROP SEGMENT
  2693. % -----------------------
  2694. reachSegment = reachStartIdx:reachEndIdx;
  2695. xPlot = pawX_norm(reachSegment);
  2696. yPlot = pawY_norm(reachSegment);
  2697. xPlot_raw = pawX(reachSegment);
  2698. yPlot_raw = pawY(reachSegment);
  2699. croppedFrames = framesFiltered(reachSegment);
  2700. if numel(xPlot) < 3
  2701. [xPlot,yPlot,pelletXreach,pelletYreach,slitX_norm,metrics] = deal([]);
  2702. return;
  2703. end
  2704. sideStartFrame = croppedFrames(1);
  2705. sideEndFrame = croppedFrames(end);
  2706. %% -----------------------
  2707. % CROPPED CONTACT
  2708. % -----------------------
  2709. distToPellet_crop = sqrt(xPlot.^2 + yPlot.^2);
  2710. contactMask_crop = distToPellet_crop <= par.pelletContactTolerance;
  2711. contactRuns_crop = bwconncomp(contactMask_crop);
  2712. if contactRuns_crop.NumObjects > 0
  2713. firstContactCrop = contactRuns_crop.PixelIdxList{1}(1);
  2714. lastContactCrop = contactRuns_crop.PixelIdxList{end}(end);
  2715. pelletContactPhaseIdx = firstContactCrop:lastContactCrop;
  2716. % store the peak frames (e.g. first frame of each bout)
  2717. peakFrames = cellfun(@(idxs) idxs(1), contactRuns_crop.PixelIdxList);
  2718. metrics.pelletContactPeakFrames = {peakFrames}; % always a cell
  2719. else
  2720. pelletContactPhaseIdx = [];
  2721. metrics.pelletContactPeakFrames = {[]}; % still a cell
  2722. end
  2723. %% -----------------------
  2724. % ROBUST ENDPOINT (max X within chosen bout; de-spiked)
  2725. % -----------------------
  2726. % Work on pellet-centered cropped X
  2727. xForArgmax = xPlot;
  2728. if ~isfield(par,'endpointSmoothWin'), par.endpointSmoothWin = 5; end
  2729. if numel(xForArgmax) >= par.endpointSmoothWin && par.endpointSmoothWin > 1
  2730. xForArgmax = movmedian(xForArgmax, par.endpointSmoothWin);
  2731. end
  2732. % Simply take the global max X in the cropped range
  2733. [~, relIdx] = max(xForArgmax);
  2734. endpointIdx = relIdx; % index into cropped arrays
  2735. metrics.endpoint_x = xPlot(endpointIdx);
  2736. metrics.endpoint_y = yPlot(endpointIdx);
  2737. %% -----------------------
  2738. % SLIT → CONTACT TIMING
  2739. % -----------------------
  2740. slitStartInCropped = max(1, min(slitCrossStartIdx - reachStartIdx + 1, numel(xPlot)));
  2741. if ~isempty(pelletContactPhaseIdx)
  2742. contactIdx = pelletContactPhaseIdx(1); % first frame of chosen band
  2743. else
  2744. contactIdx = slitStartInCropped; % fallback
  2745. end
  2746. metrics.timeSlitToContact_frames = max(0, contactIdx - slitStartInCropped);
  2747. metrics.slitToPelletDistance_mm = abs(slitX_norm);
  2748. if isfield(par,'fps') && par.fps > 0
  2749. metrics.timeSlitToContact_sec = metrics.timeSlitToContact_frames / par.fps;
  2750. if metrics.slitToPelletDistance_mm > 0
  2751. metrics.normReachTime_s_per_mm = metrics.timeSlitToContact_sec / metrics.slitToPelletDistance_mm;
  2752. metrics.meanOutwardSpeed_mm_per_s = metrics.slitToPelletDistance_mm / metrics.timeSlitToContact_sec;
  2753. else
  2754. metrics.normReachTime_s_per_mm = NaN;
  2755. metrics.meanOutwardSpeed_mm_per_s = NaN;
  2756. end
  2757. else
  2758. metrics.timeSlitToContact_sec = NaN;
  2759. metrics.normReachTime_s_per_mm = NaN;
  2760. metrics.meanOutwardSpeed_mm_per_s = NaN;
  2761. end
  2762. %% -----------------------
  2763. % QC PLOT
  2764. % -----------------------
  2765. if isfield(par,'showQCplot') && par.showQCplot
  2766. figQC = figure('Visible','off', ...
  2767. 'Name', sprintf('QC Reach %s - %d [%s]', coreID, reachID, label), ...
  2768. 'Color','w','Position',[100 100 1200 500]);
  2769. hold on;
  2770. % Trajectory colored by distance
  2771. cmap = flipud(jet(256));
  2772. normDist = (distToPellet_global - min(distToPellet_global)) / ...
  2773. (max(distToPellet_global) - min(distToPellet_global) + eps);
  2774. for i = 1:(numel(framesFiltered)-1)
  2775. cIdx = max(1,min(256,round(normDist(i)*255)+1));
  2776. plot(framesFiltered(i:i+1), pawX_smooth(i:i+1), ...
  2777. 'Color', cmap(cIdx,:), 'LineWidth',2,'HandleVisibility','off');
  2778. end
  2779. % Overlay raw pawX (gray, thin)
  2780. plot(framesFiltered, pawX_norm, 'Color',[0.5 0.5 0.5 0.6], ...
  2781. 'LineWidth',1, 'DisplayName','Raw pawX (norm)');
  2782. % Colorbar
  2783. cb = colorbar;
  2784. caxis([0 max(distToPellet_global)]);
  2785. ylabel(cb,'Distance to pellet (mm)');
  2786. legend('show');
  2787. % Y-limits
  2788. ylo = min(pawX_smooth)-5;
  2789. yhi = max(pawX_smooth)+5;
  2790. % --- Shade ALL contact bouts (global detection) ---
  2791. for r = 1:contactRuns_global.NumObjects
  2792. boutIdx = contactRuns_global.PixelIdxList{r};
  2793. x0 = framesFiltered(boutIdx(1));
  2794. x1 = framesFiltered(boutIdx(end));
  2795. patch([x0 x1 x1 x0], [ylo ylo yhi yhi], [0.2 0.8 0.2], ...
  2796. 'FaceAlpha',0.2,'EdgeColor','none','HandleVisibility','off');
  2797. end
  2798. % --- Shade the CHOSEN contact bout (global -> full green band) ---
  2799. if ~isempty(chosenBout_global)
  2800. x0 = framesFiltered(chosenBout_global(1));
  2801. x1 = framesFiltered(chosenBout_global(end));
  2802. patch([x0 x1 x1 x0], [ylo ylo yhi yhi], [0.2 0.2 0.9], ...
  2803. 'FaceAlpha',0.25,'EdgeColor','none','DisplayName','Chosen contact');
  2804. end
  2805. % --- Endpoint marker (in cropped indices) ---
  2806. plot(croppedFrames(endpointIdx), xPlot(endpointIdx), 'ro', ...
  2807. 'MarkerFaceColor','r','DisplayName','Endpoint');
  2808. % Boundaries
  2809. xline(framesFiltered(reachStartIdx),'--m','LineWidth',1.75,'DisplayName','Slit start');
  2810. xline(framesFiltered(reachEndIdx),'--r','LineWidth',1.75,'DisplayName','Retrieval end');
  2811. yline(0,'-k','LineWidth',2,'DisplayName','Pellet (x=0)');
  2812. yline(slitX_norm,':k','LineWidth',1.5,'DisplayName','Slit');
  2813. xlabel('Frame'); ylabel('Paw X [mm] (pellet-centered)');
  2814. title(sprintf('%s - Reach %d [%s]', coreID, reachID, label));
  2815. xlim([framesFiltered(1) framesFiltered(end)]);
  2816. ylim([ylo yhi]);
  2817. % --- Save QC plot ---
  2818. if exist('outDir','var') && ~isempty(outDir)
  2819. qcDir = fullfile(outDir,'QC','Reach');
  2820. if ~exist(qcDir,'dir'), mkdir(qcDir); end
  2821. saveas(figQC, fullfile(qcDir, sprintf('%s_%d.png', coreID, reachID)));
  2822. end
  2823. close(figQC); % prevent too many open figs
  2824. end
  2825. % -----------------------
  2826. % METRICS (all expected fields)
  2827. % -----------------------
  2828. % Outward movement
  2829. xOut = xPlot(1:contactIdx);
  2830. yOut = yPlot(1:contactIdx);
  2831. velOut = diff(xOut);
  2832. metrics.outwardMovementDuration = numel(xOut)-1;
  2833. metrics.peakOutwardVelocity = max(velOut, [], 'omitnan');
  2834. metrics.meanOutwardVelocity = mean(velOut, 'omitnan');
  2835. if ~isempty(velOut)
  2836. [~, pOut] = max(velOut);
  2837. metrics.timeToPeakVelocity = pOut;
  2838. metrics.timeToPeakVelocity_norm = pOut / max(1,numel(xPlot));
  2839. metrics.peakOutwardVelocityFrame = pOut;
  2840. else
  2841. metrics.timeToPeakVelocity = NaN;
  2842. metrics.timeToPeakVelocity_norm = NaN;
  2843. metrics.peakOutwardVelocityFrame = NaN;
  2844. end
  2845. accOut = diff(velOut);
  2846. if ~isempty(accOut)
  2847. metrics.peakAcceleration = max(accOut, [], 'omitnan');
  2848. metrics.peakDeceleration = min(accOut, [], 'omitnan');
  2849. jerkOut = diff(accOut);
  2850. metrics.meanAbsJerk = mean(abs(jerkOut), 'omitnan');
  2851. if ~isempty(jerkOut)
  2852. metrics.peakJerk = max(abs(jerkOut), [], 'omitnan');
  2853. else
  2854. metrics.peakJerk = NaN;
  2855. end
  2856. metrics.accelSignChanges = sum(diff(sign(accOut))~=0);
  2857. else
  2858. metrics.peakAcceleration = NaN;
  2859. metrics.peakDeceleration = NaN;
  2860. metrics.meanAbsJerk = NaN;
  2861. metrics.peakJerk = NaN;
  2862. metrics.accelSignChanges = 0;
  2863. end
  2864. metrics.trajectoryLength_outward = sum(hypot(diff(xOut), diff(yOut)), 'omitnan');
  2865. if numel(xOut) >= 2
  2866. metrics.initialReachAngle = atan2d(yOut(2)-yOut(1), xOut(2)-xOut(1));
  2867. else
  2868. metrics.initialReachAngle = NaN;
  2869. end
  2870. % Y positions at fractions of slit→pellet
  2871. dist3 = slitX_norm/3;
  2872. x_points = [slitX_norm + dist3, slitX_norm + 2*dist3, 0];
  2873. [xUnique, uniqueIdx] = unique(xOut, 'stable');
  2874. yUnique = yOut(uniqueIdx);
  2875. if numel(xUnique) >= 2
  2876. y_interp = interp1(xUnique, yUnique, x_points, 'linear','extrap');
  2877. else
  2878. y_interp = [NaN NaN NaN];
  2879. end
  2880. metrics.y_at_1_3 = y_interp(1);
  2881. metrics.y_at_2_3 = y_interp(2);
  2882. metrics.y_at_pellet = y_interp(3);
  2883. % Contact-related
  2884. metrics.nPelletContactPeaks = nPeaks;
  2885. if exist('pelletContactPhaseIdx','var') && ~isempty(pelletContactPhaseIdx)
  2886. metrics.contactDuration = numel(pelletContactPhaseIdx);
  2887. else
  2888. metrics.contactDuration = 0;
  2889. end
  2890. % Pauses & corrections
  2891. velX = diff(xPlot);
  2892. metrics.corrections = sum(velX < 0);
  2893. velThreshold = 0.5;
  2894. lowVel = abs(velX) < velThreshold;
  2895. dLow = diff([0; lowVel; 0]);
  2896. pauseS = find(dLow == 1);
  2897. pauseE = find(dLow == -1) - 1;
  2898. metrics.pauseCount = numel(pauseS);
  2899. metrics.totalPauseDuration= sum(max(0, pauseE - pauseS + 1));
  2900. % Retrieval
  2901. xRetr = xPlot(contactIdx:end);
  2902. yRetr = yPlot(contactIdx:end);
  2903. velRetr = diff(xRetr);
  2904. metrics.retrievalDuration = numel(xRetr)-1;
  2905. metrics.peakRetrievalVelocity = min(velRetr, [], 'omitnan');
  2906. metrics.meanRetrievalVelocity = mean(velRetr, 'omitnan');
  2907. if numel(xRetr) >= 2
  2908. retrDiffs = diff([xRetr(:), yRetr(:)]);
  2909. retrLen = sum(hypot(retrDiffs(:,1), retrDiffs(:,2)),'omitnan');
  2910. retrEuc = norm([xRetr(end)-xRetr(1), yRetr(end)-yRetr(1)]);
  2911. metrics.retrievalStraightness = retrEuc / max(retrLen, eps);
  2912. metrics.retrievalArc = max(yRetr) - min(yRetr);
  2913. metrics.retrievalReversals = sum(diff(xRetr) > 0);
  2914. accRetr = diff(velRetr);
  2915. metrics.retrievalMeanAbsJerk = mean(abs(diff(accRetr)), 'omitnan');
  2916. else
  2917. metrics.retrievalStraightness = NaN;
  2918. metrics.retrievalArc = NaN;
  2919. metrics.retrievalReversals = 0;
  2920. metrics.retrievalMeanAbsJerk = NaN;
  2921. end
  2922. % Global
  2923. metrics.reachDuration = sideEndFrame - sideStartFrame + 1;
  2924. metrics.maxHeight = max(yPlot);
  2925. metrics.trajectoryLength = sum(hypot(diff(xPlot), diff(yPlot)), 'omitnan');
  2926. metrics.trajectoryStraightness = norm([xPlot(end)-xPlot(1), yPlot(end)-yPlot(1)]) / ...
  2927. max(metrics.trajectoryLength, eps);
  2928. diffVecs = diff([xPlot(:), yPlot(:)]);
  2929. angles = atan2(diffVecs(:,2), diffVecs(:,1));
  2930. metrics.pathTortuosity = sum(abs(diff(angles)), 'omitnan');
  2931. % Attempts (helper function required)
  2932. metrics.reachAttempts = computeReachAttempts(pawX_smooth, slitX, 0);
  2933. % Placeholder fields if missing
  2934. if ~isfield(metrics,'timeToPelletContact')
  2935. metrics.timeToPelletContact = metrics.timeSlitToContact_frames; % keep in frames
  2936. end
  2937. if ~isfield(metrics,'pelletContactPeakFrames')
  2938. metrics.pelletContactPeakFrames = [];
  2939. end
  2940. % Ensure scalar values for metrics
  2941. scalarFields = {'peakOutwardVelocity','meanOutwardVelocity','outwardMovementDuration', ...
  2942. 'initialReachAngle','y_at_1_3','y_at_2_3','y_at_pellet', ...
  2943. 'overshootDistance','corrections','pauseCount','totalPauseDuration', ...
  2944. 'retrievalDuration','peakRetrievalVelocity','meanRetrievalVelocity', ...
  2945. 'retrievalStraightness','endpoint_x','endpoint_y', ...
  2946. 'timeSlitToContact_frames','timeSlitToContact_sec', ...
  2947. 'slitToPelletDistance_mm','normReachTime_s_per_mm','meanOutwardSpeed_mm_per_s'};
  2948. for f = 1:numel(scalarFields)
  2949. fld = scalarFields{f};
  2950. if ~isfield(metrics,fld) || isempty(metrics.(fld))
  2951. metrics.(fld) = NaN;
  2952. end
  2953. end
  2954. % Cell fields (must always be a cell, even if empty)
  2955. cellFields = {'pelletContactPeakFrames'};
  2956. for f = 1:numel(cellFields)
  2957. fld = cellFields{f};
  2958. if ~isfield(metrics,fld)
  2959. metrics.(fld) = {[]};
  2960. elseif ~iscell(metrics.(fld))
  2961. metrics.(fld) = {metrics.(fld)};
  2962. end
  2963. end
  2964. end
  2965. function attempts = computeReachAttempts(pawX_smooth, slitX, pelletXreach)
  2966. isOutsideSlit = pawX_smooth > slitX;
  2967. dSlit = diff([0; isOutsideSlit; 0]);
  2968. slitStarts = find(dSlit == 1);
  2969. slitEnds = find(dSlit == -1) - 1;
  2970. attemptsPerBout = zeros(length(slitStarts),1);
  2971. for b = 1:length(slitStarts)
  2972. seg = pawX_smooth(slitStarts(b):slitEnds(b));
  2973. contactMaskBout = seg > pelletXreach;
  2974. dBout = diff([0; contactMaskBout; 0]);
  2975. attemptsPerBout(b) = sum(dBout == 1);
  2976. end
  2977. if isempty(attemptsPerBout)
  2978. attempts = 0;
  2979. else
  2980. attempts = max(attemptsPerBout);
  2981. end
  2982. end
  2983. function idx = lastLocalMinBefore(sig, idxPeak)
  2984. % Return index of the last local minimum BEFORE idxPeak (>=1).
  2985. % Robust to short segments; falls back to global min if needed.
  2986. if idxPeak <= 1
  2987. idx = 1;
  2988. return;
  2989. end
  2990. seg = sig(1:idxPeak-1);
  2991. n = numel(seg);
  2992. if n >= 3
  2993. [~, locs] = findpeaks(-seg); % minima = peaks on inverted signal
  2994. if ~isempty(locs)
  2995. idx = locs(end); % last minimum before peak
  2996. return;
  2997. end
  2998. % fallback: global min in segment
  2999. [~, idx] = min(seg);
  3000. elseif n >= 1
  3001. [~, idx] = min(seg); % too short for findpeaks → pick min
  3002. else
  3003. idx = 1; % no samples
  3004. end
  3005. end
  3006. function idx = firstLocalMinAfter(sig, idxPeak)
  3007. % Return index of the first local minimum AFTER idxPeak (<= length(sig)).
  3008. % Robust to short segments; falls back to global min if needed.
  3009. nSig = numel(sig);
  3010. if idxPeak >= nSig
  3011. idx = nSig;
  3012. return;
  3013. end
  3014. seg = sig(idxPeak+1:end);
  3015. n = numel(seg);
  3016. if n >= 3
  3017. [~, locs] = findpeaks(-seg);
  3018. if ~isempty(locs)
  3019. idx = idxPeak + locs(1); % map back to full-signal index
  3020. return;
  3021. end
  3022. % fallback: global min after peak
  3023. [~, rel] = min(seg);
  3024. idx = idxPeak + rel;
  3025. elseif n >= 1
  3026. [~, rel] = min(seg); % too short for findpeaks → pick min
  3027. idx = idxPeak + rel;
  3028. else
  3029. idx = nSig; % no samples
  3030. end
  3031. end
  3032. %%
  3033. function [metrics, traj, failReason] = analyzeReach_Front( ...
  3034. coreID, reachRow, dlc_front, offsetVal, likelihoodThresh, sideStartRef, sideEndRef)
  3035. metrics = struct();
  3036. traj = struct();
  3037. failReason = "";
  3038. % Parameters
  3039. par.showQCplot = false;
  3040. par.frame_buffer = 200; % extend ± buffer
  3041. par.boundarySlack= 10; % slack around peaks
  3042. par.minLik = 0.6; % high likelihood threshold
  3043. par.minRun = 5; % min consecutive high-likelihood frames
  3044. par.minRunKeep = 4; % minimum run length to keep after crop
  3045. % ----- pellet reference -----
  3046. pelletMask = dlc_front.pellet_likelihood > 0.9;
  3047. pelletX = median(dlc_front.pellet_x(pelletMask), 'omitnan');
  3048. pelletY = median(dlc_front.pellet_y(pelletMask), 'omitnan');
  3049. % Paw = mean of digit2 + digit5
  3050. pawX = mean([dlc_front.digit2_x, dlc_front.digit5_x], 2, 'omitnan') - pelletX;
  3051. pawY = mean([dlc_front.digit2_y, dlc_front.digit5_y], 2, 'omitnan') - pelletY;
  3052. pawLik= mean([dlc_front.digit2_likelihood, dlc_front.digit5_likelihood], 2, 'omitnan');
  3053. % Digit coords relative to pellet
  3054. digit2X = dlc_front.digit2_x - pelletX;
  3055. digit2Y = dlc_front.digit2_y - pelletY;
  3056. digit5X = dlc_front.digit5_x - pelletX;
  3057. digit5Y = dlc_front.digit5_y - pelletY;
  3058. % ----- reach window mapping -----
  3059. if isnan(sideStartRef) || isnan(sideEndRef) || sideEndRef <= sideStartRef
  3060. failReason = "invalid side refined window";
  3061. return;
  3062. end
  3063. frontStart = sideStartRef + offsetVal;
  3064. frontEnd = sideEndRef + offsetVal;
  3065. frontStart = max(1, min(frontStart, height(dlc_front)));
  3066. frontEnd = max(1, min(frontEnd, height(dlc_front)));
  3067. frames = max(1, frontStart - par.frame_buffer) : ...
  3068. min(frontEnd + par.frame_buffer, height(dlc_front));
  3069. % ----- digit spread for peak detection -----
  3070. d2x = dlc_front.digit2_x(frames) - pelletX;
  3071. d2y = dlc_front.digit2_y(frames) - pelletY;
  3072. d5x = dlc_front.digit5_x(frames) - pelletX;
  3073. d5y = dlc_front.digit5_y(frames) - pelletY;
  3074. spreadRaw = sqrt((d2x - d5x).^2 + (d2y - d5y).^2);
  3075. spreadLik = min(dlc_front.digit2_likelihood(frames), dlc_front.digit5_likelihood(frames));
  3076. spread = spreadRaw;
  3077. spread(spreadLik < likelihoodThresh) = NaN;
  3078. % --- peak detection ---
  3079. [pkVals, locs] = findpeaks(spread, ...
  3080. 'MinPeakProminence', 0.5, ...
  3081. 'MinPeakDistance', 8);
  3082. peakLik = arrayfun(@(i) mean(pawLik(frames(max(1,i-2):min(end,i+2))), 'omitnan'), locs);
  3083. if isempty(locs)
  3084. failReason = "No spread peaks found";
  3085. return;
  3086. end
  3087. score = pkVals .* peakLik;
  3088. [~, bestIdx] = max(score);
  3089. idxMaxSpread = locs(bestIdx);
  3090. % --- smoothing and run segmentation ---
  3091. spreadSm = smoothdata(spread, 'sgolay', 17, 'omitnan');
  3092. likSm = smoothdata(spreadLik, 'movmean', 7, 'omitnan');
  3093. validMask = ~isnan(spreadSm);
  3094. edges = diff([false; validMask; false]);
  3095. runStarts = find(edges == 1);
  3096. runEnds = find(edges == -1) - 1;
  3097. r = find(idxMaxSpread >= runStarts & idxMaxSpread <= runEnds, 1, 'first');
  3098. thisRun = runStarts(r):runEnds(r);
  3099. % Restrict valley search
  3100. valleyMask = false(size(spreadSm));
  3101. valleyMask(thisRun) = islocalmin(spreadSm(thisRun));
  3102. valleyLocs = find(valleyMask);
  3103. % Score valleys
  3104. peakVal = spreadSm(idxMaxSpread);
  3105. scoreValley = @(i, signFlip) ...
  3106. (peakVal - spreadSm(i)) / max(peakVal, eps) + ...
  3107. 0.5 * max(0, signFlip * (likSm(min(i+4,end)) - likSm(max(i-4,1))));
  3108. leftCands = valleyLocs(valleyLocs < idxMaxSpread - 5);
  3109. if ~isempty(leftCands)
  3110. scores = arrayfun(@(i) scoreValley(i,+1), leftCands);
  3111. [~,k] = max(scores);
  3112. leftValley = leftCands(k);
  3113. else
  3114. leftValley = thisRun(1);
  3115. end
  3116. rightCands = valleyLocs(valleyLocs > idxMaxSpread + 5);
  3117. if ~isempty(rightCands)
  3118. scores = arrayfun(@(i) scoreValley(i,-1), rightCands);
  3119. [~,k] = max(scores);
  3120. rightValley = rightCands(k);
  3121. else
  3122. rightValley = thisRun(end);
  3123. end
  3124. % redefine reach window
  3125. reachStartFrame = frames(leftValley);
  3126. reachEndFrame = frames(rightValley);
  3127. frames = reachStartFrame:reachEndFrame;
  3128. % ----- build cropped signals -----
  3129. xTraj = pawX(frames);
  3130. yTraj = pawY(frames);
  3131. d2Traj = [digit2X(frames), digit2Y(frames)];
  3132. d5Traj = [digit5X(frames), digit5Y(frames)];
  3133. % ----- post-crop cleanup -----
  3134. % 1) Remove jumps
  3135. dx = diff(xTraj); dy = diff(yTraj);
  3136. stepDist = hypot(dx,dy);
  3137. medStep = median(stepDist,'omitnan');
  3138. madStep = mad(stepDist,1);
  3139. jumpThresh = medStep + 10*madStep;
  3140. jumps = [false; stepDist > jumpThresh];
  3141. xTraj(jumps) = NaN; yTraj(jumps) = NaN;
  3142. % 2) Remove spikes
  3143. vel = hypot(diff(xTraj), diff(yTraj));
  3144. zv = (vel - mean(vel,'omitnan'))/std(vel,'omitnan');
  3145. spike = [false; abs(zv) > 5];
  3146. xTraj(spike) = NaN; yTraj(spike) = NaN;
  3147. % 3) Interpolate short gaps
  3148. maxGap = 5;
  3149. xTraj = fillmissing(xTraj,'linear','MaxGap',maxGap,'EndValues','nearest');
  3150. yTraj = fillmissing(yTraj,'linear','MaxGap',maxGap,'EndValues','nearest');
  3151. % 4) Remove short runs (≤3 frames)
  3152. validMask = ~isnan(xTraj);
  3153. edges = diff([false; validMask; false]);
  3154. runStarts = find(edges==1);
  3155. runEnds = find(edges==-1)-1;
  3156. for rr = 1:numel(runStarts)
  3157. if runEnds(rr) - runStarts(rr) + 1 <= par.minRunKeep
  3158. xTraj(runStarts(rr):runEnds(rr)) = NaN;
  3159. yTraj(runStarts(rr):runEnds(rr)) = NaN;
  3160. end
  3161. end
  3162. xTraj = fillmissing(xTraj,'linear','MaxGap',maxGap,'EndValues','nearest');
  3163. yTraj = fillmissing(yTraj,'linear','MaxGap',maxGap,'EndValues','nearest');
  3164. % recompute spread
  3165. spread = sqrt((d2Traj(:,1)-d5Traj(:,1)).^2 + (d2Traj(:,2)-d5Traj(:,2)).^2);
  3166. if sum(~isnan(xTraj)) < 5
  3167. failReason = "Not enough valid paw points after cleanup";
  3168. return;
  3169. end
  3170. % ------------------ INDEX CONVERSIONS ------------------
  3171. idxMax_c = idxMaxSpread - leftValley + 1;
  3172. idxMax_c = max(1, min(idxMax_c, numel(frames)));
  3173. pelletContactIdx_c = max(1, min(idxMax_c + 1, numel(frames)));
  3174. absMaxFrame = frames(idxMax_c);
  3175. absClosureFrame = frames(pelletContactIdx_c);
  3176. absLeftFrame = frames(1);
  3177. absRightFrame = frames(end);
  3178. % ----- metrics -----
  3179. trajLen = sum(sqrt(diff(xTraj).^2 + diff(yTraj).^2), 'omitnan');
  3180. trajWidth = range(xTraj);
  3181. trajHeight = range(yTraj);
  3182. lineVec = [0 0] - [xTraj(1), yTraj(1)];
  3183. normLine = norm(lineVec);
  3184. if normLine > 0
  3185. proj = (xTraj - xTraj(1))*lineVec(1) + (yTraj - yTraj(1))*lineVec(2);
  3186. proj = proj / normLine^2 * lineVec;
  3187. dev = sqrt((xTraj - (xTraj(1)+proj(:,1))).^2 + (yTraj - (yTraj(1)+proj(:,2))).^2);
  3188. frontDeviation = mean(dev,'omitnan');
  3189. else
  3190. frontDeviation = NaN;
  3191. end
  3192. vx = diff(xTraj);
  3193. frontZigZags = sum(diff(sign(vx))~=0);
  3194. distToPellet = sqrt(xTraj.^2 + yTraj.^2);
  3195. minDistToPellet = min(distToPellet,[],'omitnan');
  3196. [~, minDistIdx] = min(distToPellet);
  3197. digitSpreadMean = mean(spread,'omitnan');
  3198. digitSpreadPeak = max(spread,[],'omitnan');
  3199. digitSpreadAtContact = spread(1);
  3200. digitSpreadAtRetrieval = spread(end);
  3201. digitSpreadTiming = idxMax_c / numel(spread);
  3202. metrics.frontTrajLen = trajLen;
  3203. metrics.frontTrajWidth = trajWidth;
  3204. metrics.frontTrajHeight = trajHeight;
  3205. metrics.frontDeviation = frontDeviation;
  3206. metrics.frontZigZags = frontZigZags;
  3207. metrics.minDistToPellet = minDistToPellet;
  3208. metrics.digitSpreadMean = digitSpreadMean;
  3209. metrics.digitSpreadPeak = digitSpreadPeak;
  3210. metrics.digitSpreadAtContact = digitSpreadAtContact;
  3211. metrics.digitSpreadAtRetrieval= digitSpreadAtRetrieval;
  3212. metrics.digitSpreadTiming = digitSpreadTiming;
  3213. metrics.digitSpreadAtClosure = spread(pelletContactIdx_c);
  3214. metrics.distToPelletAtClosure = hypot(xTraj(pelletContactIdx_c), yTraj(pelletContactIdx_c));
  3215. metrics.closureFrameNorm = pelletContactIdx_c / numel(spread);
  3216. traj = struct( ...
  3217. 'reachID', reachRow.ReachIndex, ...
  3218. 'label', string(reachRow.Label), ...
  3219. 'frames', frames, ...
  3220. 'x', xTraj, ...
  3221. 'y', yTraj, ...
  3222. 'digit2', d2Traj, ...
  3223. 'pelletX', 0, ... % you computed these above
  3224. 'pelletY', 0, ...
  3225. 'slitX_norm', 0, ...
  3226. 'digit5', d5Traj, ...
  3227. 'spread', spread);
  3228. % ---------- QC PLOT ----------
  3229. if par.showQCplot
  3230. figQC = figure('Visible','on', ...
  3231. 'Name', sprintf('Front QC Reach %s - %s [%s]', ...
  3232. coreID, string(reachRow.ReachIndex), string(reachRow.Label)), ...
  3233. 'Color','w','Position',[100 100 1200 500]);
  3234. fullFrames = max(1, frontStart - par.frame_buffer) : ...
  3235. min(frontEnd + par.frame_buffer, height(dlc_front));
  3236. d2x_full = dlc_front.digit2_x(fullFrames) - pelletX;
  3237. d2y_full = dlc_front.digit2_y(fullFrames) - pelletY;
  3238. d5x_full = dlc_front.digit5_x(fullFrames) - pelletX;
  3239. d5y_full = dlc_front.digit5_y(fullFrames) - pelletY;
  3240. spreadFull_raw = sqrt((d2x_full - d5x_full).^2 + (d2y_full - d5y_full).^2);
  3241. likFull = min(dlc_front.digit2_likelihood(fullFrames), ...
  3242. dlc_front.digit5_likelihood(fullFrames));
  3243. spreadFull = spreadFull_raw;
  3244. spreadFull(likFull < likelihoodThresh) = NaN;
  3245. dS = diff(spreadFull);
  3246. medD = median(dS,'omitnan');
  3247. madD = mad(dS,1);
  3248. jump = [false; abs(dS - medD) > 8*madD];
  3249. spreadFull(jump) = NaN;
  3250. spreadFull = fillmissing(spreadFull,'linear','MaxGap',5,'EndValues','nearest');
  3251. spreadFullSm = smoothdata(spreadFull,'sgolay',11);
  3252. toFullIdx = @(absF) max(1, min(numel(fullFrames), absF - fullFrames(1) + 1));
  3253. iMax = toFullIdx(absMaxFrame);
  3254. iClosure = toFullIdx(absClosureFrame);
  3255. iLeft = toFullIdx(absLeftFrame);
  3256. iRight = toFullIdx(absRightFrame);
  3257. yMax = spreadFullSm(iMax);
  3258. yClosure = spreadFullSm(iClosure);
  3259. subplot(1,2,1); hold on;
  3260. cmap = jet(256);
  3261. likNorm = round(1 + 255*max(0,min(1,likFull)));
  3262. likNorm(isnan(likNorm)) = 1;
  3263. for i = 1:(numel(fullFrames)-1)
  3264. if any(isnan(spreadFullSm(i:i+1))), continue; end
  3265. cIdx = min(max(likNorm(i),1),256);
  3266. plot(fullFrames(i:i+1), spreadFullSm(i:i+1), '-', ...
  3267. 'Color', cmap(cIdx,:), 'LineWidth', 2, 'HandleVisibility','off');
  3268. end
  3269. plot(absMaxFrame, yMax, 'bo','MarkerFaceColor','b','MarkerSize',8,'DisplayName','Chosen max spread');
  3270. plot(absClosureFrame, yClosure, 'ro','MarkerFaceColor','r','MarkerSize',8,'DisplayName','Closure');
  3271. xline(absLeftFrame, '--k','Left valley');
  3272. xline(absRightFrame, '--k','Right valley');
  3273. text(absMaxFrame, yMax, sprintf(' Max @ %d', absMaxFrame), ...
  3274. 'VerticalAlignment','bottom','Color','b','FontWeight','bold');
  3275. text(absClosureFrame, yClosure, sprintf(' Closure @ %d', absClosureFrame), ...
  3276. 'VerticalAlignment','top','Color','r','FontWeight','bold');
  3277. xlabel('Frame'); ylabel('Digit 2–5 spread (mm)');
  3278. title('Digit spread QC (colored by likelihood)');
  3279. legend('show');
  3280. cb1 = colorbar; caxis([0 1]); ylabel(cb1,'Tracking likelihood');
  3281. subplot(1,2,2); hold on;
  3282. likCrop = pawLik(frames);
  3283. likCrop(isnan(likCrop)) = 0;
  3284. likCropNorm = round(1 + 255*max(0,min(1,likCrop)));
  3285. for i = 1:(numel(xTraj)-1)
  3286. if any(isnan([xTraj(i:i+1); yTraj(i:i+1)])), continue; end
  3287. cIdx = min(max(likCropNorm(i),1),256);
  3288. plot(xTraj(i:i+1), yTraj(i:i+1), '-', ...
  3289. 'Color', cmap(cIdx,:), 'LineWidth', 2, 'HandleVisibility','off');
  3290. end
  3291. scatter(xTraj(1), yTraj(1), 60, 'g', 'filled', 'DisplayName','Start');
  3292. scatter(xTraj(end), yTraj(end), 60, 'm', 'filled', 'DisplayName','End');
  3293. scatter(xTraj(pelletContactIdx_c), yTraj(pelletContactIdx_c), ...
  3294. 60, 'r', 'filled','DisplayName','Closure');
  3295. text(xTraj(pelletContactIdx_c), yTraj(pelletContactIdx_c), ...
  3296. sprintf(' %d', absClosureFrame), ...
  3297. 'VerticalAlignment','bottom','Color','r');
  3298. xlabel('X (pellet-centered)'); ylabel('Y (pellet-centered)');
  3299. title('Paw trajectory QC (colored by likelihood)');
  3300. legend('show');
  3301. cb2 = colorbar; caxis([0 1]); ylabel(cb2,'Tracking likelihood');
  3302. qcDir = fullfile(pwd,'OUT','QC','FrontReaches');
  3303. if ~exist(qcDir,'dir'), mkdir(qcDir); end
  3304. saveas(figQC, fullfile(qcDir, sprintf('%s_front_%s.png', ...
  3305. coreID, string(reachRow.ReachIndex))));
  3306. pause;
  3307. close(figQC);
  3308. end
  3309. end
  3310. %%
  3311. function [group, animal, test_day] = parseCoreID(coreID)
  3312. % Split by both '-' and '_'
  3313. parts = regexp(coreID, '[-_]', 'split');
  3314. if numel(parts) < 4
  3315. error('coreID format error: expected four parts split by "-" and "_".');
  3316. end
  3317. group = parts{1};
  3318. % Combine parts 2 and 3 for animal (with hyphen)
  3319. animal = strcat(parts{2}, '-', parts{3});
  3320. test_day = parts{4};
  3321. end
  3322. %% Graph functions
  3323. function params = getViewParams(viewType)
  3324. % Returns x/y limits, binning and edges/centers for the two views
  3325. switch lower(viewType)
  3326. case 'side'
  3327. params.xlim = [-15 10];
  3328. params.ylim = [-10 15];
  3329. params.bin = 1; % mm
  3330. params.xEdges = params.xlim(1):params.bin:params.xlim(2);
  3331. params.yEdges = params.ylim(1):params.bin:params.ylim(2);
  3332. % centers
  3333. params.xCtr = (params.xEdges(1:end-1)+params.xEdges(2:end))/2;
  3334. params.yCtr = (params.yEdges(1:end-1)+params.yEdges(2:end))/2;
  3335. params.showSlit = true;
  3336. case 'front'
  3337. params.xlim = [-8 8];
  3338. params.ylim = [-7 1];
  3339. params.bin = 0.5; % mm
  3340. params.xEdges = params.xlim(1):params.bin:params.xlim(2);
  3341. params.yEdges = params.ylim(1):params.bin:params.ylim(2);
  3342. % centers from edges
  3343. params.xCtr = (params.xEdges(1:end-1)+params.xEdges(2:end))/2;
  3344. params.yCtr = (params.yEdges(1:end-1)+params.yEdges(2:end))/2;
  3345. params.showSlit = false; % no slit for front view
  3346. otherwise
  3347. error('Unknown viewType: %s', viewType);
  3348. end
  3349. end
  3350. function plotPerAnimalTraj(trajArray, coreID, viewType, figDir, groupBy)
  3351. if nargin < 5
  3352. groupBy = "label"; % default if not specified
  3353. end
  3354. if isempty(trajArray), return; end
  3355. P = getViewParams(viewType);
  3356. % --- Choose which field to group on ---
  3357. switch lower(groupBy)
  3358. case "label"
  3359. labelsAll = string({trajArray.label});
  3360. case "broadlabel"
  3361. labelsAll = string({trajArray.broadLabel});
  3362. otherwise
  3363. error("Unknown groupBy option: %s (must be 'label' or 'broadLabel')", groupBy);
  3364. end
  3365. uLabels = unique(labelsAll,'stable');
  3366. % --- reorder: Success first, then Errors alphabetically ---
  3367. isSuccess = strcmpi(uLabels, "Success");
  3368. isError = startsWith(uLabels, "Error", 'IgnoreCase', true);
  3369. if ~any(isSuccess | isError)
  3370. warning('No Success or Error labels found, keeping original order');
  3371. % keep uLabels as-is
  3372. else
  3373. errorLabels = sort(uLabels(isError));
  3374. uLabels = [uLabels(isSuccess), errorLabels, uLabels(~(isSuccess | isError))];
  3375. end
  3376. nLabels = numel(uLabels);
  3377. nRows = 1;
  3378. nCols = nLabels;
  3379. cols = lines(nLabels);
  3380. panelSize = 300; % px per subplot, adjust as needed
  3381. f = figure('Visible','off','Color','w','Name',sprintf('%s %s Traj',coreID,viewType), ...
  3382. 'Position', [100 100 panelSize*nCols panelSize*nRows]);
  3383. for k = 1:nLabels
  3384. L = uLabels(k);
  3385. sel = strcmp(labelsAll, L);
  3386. T = [trajArray(sel).traj];
  3387. subplot(nRows,nCols,k); hold on; grid on;
  3388. title(sprintf('%s (%d)', L, numel(T)), 'Interpreter','none');
  3389. xlabel('X (mm)'); ylabel('Y (mm)');
  3390. xlim(P.xlim); ylim(P.ylim);
  3391. set(gca,'YDir','reverse'); % <-- restore flipped Y
  3392. base = cols(k,:);
  3393. % individual trajectories
  3394. for t = 1:numel(T)
  3395. c = base * (t/numel(T));
  3396. plot(T(t).x, T(t).y, 'LineWidth',1.5, 'Color', [c 0.20]);
  3397. end
  3398. % --- average trace (median) ---
  3399. if ~isempty(T)
  3400. nPts = 100; % number of resample points
  3401. Xrs = nan(numel(T), nPts);
  3402. Yrs = nan(numel(T), nPts);
  3403. nIncluded = 0;
  3404. nExcluded = 0;
  3405. for t = 1:numel(T)
  3406. x = T(t).x(:); % force column
  3407. y = T(t).y(:); % force column
  3408. if numel(x) < 2
  3409. nExcluded = nExcluded + 1;
  3410. continue; % skip too-short trajectories
  3411. end
  3412. % find NaN runs
  3413. mask = isnan(x) | isnan(y);
  3414. d = diff([0; mask; 0]);
  3415. starts = find(d == 1);
  3416. ends = find(d == -1);
  3417. if isempty(starts)
  3418. maxRun = 0;
  3419. else
  3420. runLengths = ends - starts;
  3421. maxRun = max(runLengths);
  3422. end
  3423. if maxRun > 2 % threshold: reject if longest NaN run > 3 samples
  3424. nExcluded = nExcluded + 1;
  3425. continue;
  3426. end
  3427. % patch small NaNs by interpolation
  3428. x = fillmissing(x,'linear','EndValues','extrap');
  3429. y = fillmissing(y,'linear','EndValues','extrap');
  3430. % safety: if still non-finite, skip
  3431. if any(~isfinite(x)) || any(~isfinite(y))
  3432. nExcluded = nExcluded + 1;
  3433. continue;
  3434. end
  3435. dx = diff(x);
  3436. dy = diff(y);
  3437. arc = [0; cumsum(sqrt(dx.^2 + dy.^2))];
  3438. if arc(end) == 0
  3439. nExcluded = nExcluded + 1;
  3440. continue; % skip degenerate
  3441. end
  3442. arcNorm = arc ./ arc(end);
  3443. % remove duplicates in arcNorm
  3444. [arcNormUnique, ia] = unique(arcNorm, 'stable');
  3445. xUnique = x(ia);
  3446. yUnique = y(ia);
  3447. if numel(arcNormUnique) < 2
  3448. nExcluded = nExcluded + 1;
  3449. continue;
  3450. end
  3451. arcGrid = linspace(0,1,nPts);
  3452. Xrs(t,:) = interp1(arcNormUnique, xUnique, arcGrid, 'linear', 'extrap');
  3453. Yrs(t,:) = interp1(arcNormUnique, yUnique, arcGrid, 'linear', 'extrap');
  3454. nIncluded = nIncluded + 1;
  3455. end
  3456. xMed = median(Xrs,1,'omitnan');
  3457. yMed = median(Yrs,1,'omitnan');
  3458. plot(xMed, yMed, 'k-', 'LineWidth',1.0);
  3459. end
  3460. fprintf('Label %s: included %d trajectories, excluded %d\n', string(L), nIncluded, nExcluded);
  3461. % pellet at origin
  3462. scatter(0,0,60,'o','MarkerEdgeColor','k','MarkerFaceColor',[0.5 0.5 0.5],'LineWidth',1.25);
  3463. % mean slit (side only)
  3464. if P.showSlit
  3465. slitVals = [T.slitX_norm];
  3466. if ~isempty(slitVals)
  3467. mSlit = mean(slitVals,'omitnan');
  3468. plot([mSlit mSlit], get(gca,'YLim'), ':', 'LineWidth',1.5, 'Color',[0.7 0.7 0.7]);
  3469. end
  3470. end
  3471. end
  3472. outDir = fullfile(figDir, char(groupBy), viewType);
  3473. if ~exist(outDir, 'dir'), mkdir(outDir); end
  3474. pos = get(f,'Position');
  3475. fprintf('Figure position: width=%.1f, height=%.1f\n', pos(3), pos(4));
  3476. % build stub without extension
  3477. savename = fullfile(outDir, sprintf('RAW_%s_Traj_%s_%s', ...
  3478. upper(viewType(1)), coreID, char(groupBy)));
  3479. set(f, 'Color', 'w'); % figure background
  3480. export_fig(savename, '-pdf', '-png', '-r300', f);
  3481. close(f);
  3482. end
  3483. function plotGroupTraj(Rgrp, grpName, viewType, figDir, groupBy)
  3484. if nargin < 5
  3485. groupBy = "label"; % default
  3486. end
  3487. % --- smoothing options (tunable) ---
  3488. S.smoothIndividuals = true; % smooth rainbow traces?
  3489. S.smoothAverage = true; % smooth black median path?
  3490. S.method = 'movmean'; % 'sgolay' preserves shape; 'movmean' also fine
  3491. S.windowPtsRaw = 10; % window for raw (per-frame) smoothing
  3492. S.sgOrder = 2; % sgolay poly order
  3493. S.windowPtsAvg = 10; % window for the resampled average (nPts-grid)
  3494. % local helper
  3495. smoothXY = @(x,y,w,method,ord) deal( ...
  3496. smoothdata(x, method, w, 'SamplePoints', 1:numel(x), 'sgolaydegree', ord), ...
  3497. smoothdata(y, method, w, 'SamplePoints', 1:numel(y), 'sgolaydegree', ord) );
  3498. if isempty(Rgrp), return; end
  3499. P = getViewParams(viewType);
  3500. % --- Collect wrappers with test_day + animal ID ---
  3501. trajArray = [];
  3502. for r = 1:numel(Rgrp)
  3503. if isfield(Rgrp(r), lower(viewType)) && ...
  3504. isfield(Rgrp(r).(lower(viewType)), 'trajectories')
  3505. W = Rgrp(r).(lower(viewType)).trajectories;
  3506. [W.test_day] = deal(Rgrp(r).test_day);
  3507. [W.animal] = deal(Rgrp(r).animal);
  3508. trajArray = [trajArray W]; %#ok<AGROW>
  3509. end
  3510. end
  3511. if isempty(trajArray), return; end
  3512. % --- Days and labels ---
  3513. switch lower(groupBy)
  3514. case "label"
  3515. labelsAll = string({trajArray.label});
  3516. case "broadlabel"
  3517. labelsAll = string({trajArray.broadLabel});
  3518. otherwise
  3519. error("Unknown groupBy option: %s", groupBy);
  3520. end
  3521. uDays = string(unique({trajArray.test_day}, 'stable'));
  3522. nDays = numel(uDays);
  3523. labelsAll = string(labelsAll); % enforce string array
  3524. uLabels = string(unique(labelsAll, 'stable'));
  3525. % --- reorder: Success first, then Errors alphabetically ---
  3526. isSuccess = strcmpi(uLabels, "Success");
  3527. isError = startsWith(uLabels, "Error", 'IgnoreCase', true);
  3528. if any(isSuccess | isError)
  3529. errorLabels = sort(uLabels(isError));
  3530. others = uLabels(~(isSuccess | isError));
  3531. uLabels = [uLabels(isSuccess), errorLabels, others];
  3532. end
  3533. nLabels = numel(uLabels);
  3534. % --- Animals for coloring ---
  3535. uAnimals = unique({trajArray.animal}, 'stable');
  3536. nAnimals = numel(uAnimals);
  3537. animalColors = lines(nAnimals);
  3538. % --- Subplot grid: rows = labels, cols = days ---
  3539. [LL, DD] = ndgrid(string(uLabels), string(uDays));
  3540. labelsGrid = strcat(LL," / ",DD);
  3541. nRows = nLabels;
  3542. nCols = nDays;
  3543. f = figure('Visible','off','Color','w', ...
  3544. 'Name', sprintf('%s %s GroupTraj', grpName, viewType), ...
  3545. 'Position', [100 100 300*nCols 300*nRows]);
  3546. for d = 1:nDays
  3547. daySel = strcmpi({trajArray.test_day}, uDays(d));
  3548. for k = 1:nLabels
  3549. labSel = strcmpi(labelsAll, uLabels(k));
  3550. sel = daySel & labSel;
  3551. T = [trajArray(sel).traj];
  3552. A = {trajArray(sel).animal};
  3553. idx = (k-1)*nCols + d; % manual linear index, column-major to row-major fix
  3554. subplot(nRows,nCols,idx); hold on; grid on;
  3555. title(sprintf('%s / %s (%d)', string(uLabels(k)), uDays(d), numel(T)), 'Interpreter','none');
  3556. xlabel('X (mm)'); ylabel('Y (mm)');
  3557. xlim(P.xlim); ylim(P.ylim);
  3558. set(gca,'YDir','reverse');
  3559. % --- plot individual trajectories ---
  3560. for t = 1:numel(T)
  3561. % pick animal color
  3562. aIdx = find(strcmp(uAnimals, A{t}));
  3563. col = animalColors(aIdx,:);
  3564. % grab raw coords
  3565. x = T(t).x(:);
  3566. y = T(t).y(:);
  3567. % optional smoothing
  3568. if exist('S','var') && isfield(S,'smoothIndividuals') && S.smoothIndividuals
  3569. switch lower(S.method)
  3570. case 'sgolay'
  3571. % direct Savitzky–Golay filter
  3572. x = sgolayfilt(x, 2, S.windowPtsRaw);
  3573. y = sgolayfilt(y, 2, S.windowPtsRaw);
  3574. otherwise
  3575. % smoothdata with movmean or other methods
  3576. x = smoothdata(x, S.method, S.windowPtsRaw);
  3577. y = smoothdata(y, S.method, S.windowPtsRaw);
  3578. end
  3579. end
  3580. % plot (with transparency)
  3581. plot(x, y, 'Color', [col 0.2], 'LineWidth', 1.0);
  3582. end
  3583. % --- compute average trajectory across this subset ---
  3584. if ~isempty(T)
  3585. nPts = 100; % number of normalized points
  3586. Xrs = nan(numel(T), nPts);
  3587. Yrs = nan(numel(T), nPts);
  3588. nIncluded = 0;
  3589. nExcluded = 0;
  3590. for t = 1:numel(T)
  3591. x = T(t).x(:); % force column
  3592. y = T(t).y(:); % force column
  3593. if numel(x) < 2
  3594. continue; % skip too-short trajectories
  3595. end
  3596. % find NaN runs
  3597. mask = isnan(x) | isnan(y);
  3598. dMask = diff([0; mask; 0]);
  3599. starts = find(dMask == 1);
  3600. ends = find(dMask == -1);
  3601. if isempty(starts)
  3602. maxRun = 0;
  3603. else
  3604. runLengths = ends - starts;
  3605. maxRun = max(runLengths);
  3606. end
  3607. if maxRun > 2 % threshold: reject if longest NaN run > 3 samples
  3608. nExcluded = nExcluded + 1;
  3609. continue;
  3610. end
  3611. % patch small NaNs by interpolation
  3612. x = fillmissing(x,'linear','EndValues','extrap');
  3613. y = fillmissing(y,'linear','EndValues','extrap');
  3614. % safety: if still non-finite, skip
  3615. if any(~isfinite(x)) || any(~isfinite(y))
  3616. nExcluded = nExcluded + 1;
  3617. continue;
  3618. end
  3619. dx = diff(x);
  3620. dy = diff(y);
  3621. arc = [0; cumsum(sqrt(dx.^2 + dy.^2))];
  3622. if arc(end) == 0
  3623. nExcluded = nExcluded + 1;
  3624. continue; % skip degenerate
  3625. end
  3626. arcNorm = arc ./ arc(end);
  3627. % remove duplicates in arcNorm
  3628. [arcNormUnique, ia] = unique(arcNorm, 'stable');
  3629. xUnique = x(ia);
  3630. yUnique = y(ia);
  3631. if numel(arcNormUnique) < 2
  3632. nExcluded = nExcluded + 1;
  3633. continue;
  3634. end
  3635. arcGrid = linspace(0,1,nPts);
  3636. Xrs(t,:) = interp1(arcNormUnique, xUnique, arcGrid, 'linear', 'extrap');
  3637. Yrs(t,:) = interp1(arcNormUnique, yUnique, arcGrid, 'linear', 'extrap');
  3638. nIncluded = nIncluded + 1;
  3639. end
  3640. xAvg = median(Xrs,1,'omitnan');
  3641. yAvg = median(Yrs,1,'omitnan');
  3642. plot(xAvg, yAvg, 'c-', 'LineWidth', 1); %plot unsmoothed
  3643. % optional smoothing of the median path
  3644. if exist('S','var') && isfield(S,'smoothAverage') && S.smoothAverage
  3645. switch lower(S.method)
  3646. case 'sgolay'
  3647. % direct Savitzky–Golay filter
  3648. xAvg = sgolayfilt(xAvg, S.sgOrder, S.windowPtsAvg);
  3649. yAvg = sgolayfilt(yAvg, S.sgOrder, S.windowPtsAvg);
  3650. otherwise
  3651. % smoothdata for movmean, gaussian, etc.
  3652. xAvg = smoothdata(xAvg, S.method, S.windowPtsAvg);
  3653. yAvg = smoothdata(yAvg, S.method, S.windowPtsAvg);
  3654. end
  3655. end
  3656. plot(xAvg, yAvg, 'k-', 'LineWidth', 1); %plot smoothed
  3657. end
  3658. % pellet marker
  3659. scatter(0,0,60,'o','MarkerEdgeColor','k','MarkerFaceColor',[0.5 0.5 0.5],'LineWidth',1.25);
  3660. % slit line (side view only)
  3661. if P.showSlit && isfield(T, 'slitX_norm')
  3662. slitVals = [T.slitX_norm];
  3663. if ~isempty(slitVals)
  3664. mSlit = mean(slitVals,'omitnan');
  3665. plot([mSlit mSlit], get(gca,'YLim'), ':','LineWidth',1.5,'Color',[0.7 0.7 0.7]);
  3666. end
  3667. end
  3668. end
  3669. end
  3670. % --- Legend (one for all animals, top right outside grid) ---
  3671. ax = axes(f,'Visible','off'); %#ok<LAXES>
  3672. hold(ax,'on');
  3673. h = gobjects(nAnimals,1);
  3674. for a = 1:nAnimals
  3675. h(a) = plot(ax, nan, nan, 'Color', animalColors(a,:), 'LineWidth',2); %#ok<AGROW>
  3676. end
  3677. legend(ax, h, uAnimals, 'Location','northeastoutside', 'Box','off');
  3678. title(ax, 'Animals', 'FontWeight','bold');
  3679. % --- Save ---
  3680. outDir = fullfile(figDir, char(groupBy), viewType);
  3681. if ~exist(outDir, 'dir'), mkdir(outDir); end
  3682. outName = fullfile(outDir, sprintf('%s_%s_GroupTraj_%s', grpName, upper(viewType(1)), lower(groupBy)));
  3683. % saveas(f, fullfile(outDir, outName));
  3684. set(f, 'Color', 'w'); % figure background
  3685. export_fig(outName, '-pdf', '-png', '-r300', f);
  3686. close(f);
  3687. end
  3688. function plotPerAnimalHeatmap(trajArray, coreID, viewType, figDir, groupBy)
  3689. if nargin < 5
  3690. groupBy = "label"; % default if not specified
  3691. end
  3692. if isempty(trajArray), return; end
  3693. P = getViewParams(viewType);
  3694. % --- Choose which field to group on ---
  3695. switch lower(groupBy)
  3696. case "label"
  3697. labelsAll = string({trajArray.label});
  3698. case "broadlabel"
  3699. labelsAll = string({trajArray.broadLabel});
  3700. otherwise
  3701. error("Unknown groupBy option: %s (must be 'label' or 'broadLabel')", groupBy);
  3702. end
  3703. if isempty(trajArray), return; end
  3704. P = getViewParams(viewType);
  3705. uLabels = unique(labelsAll,'stable');
  3706. nLabels = numel(uLabels);
  3707. nRows = 1; % force one row
  3708. nCols = nLabels; % one column per label
  3709. % -------- accumulate heatmaps per label --------
  3710. Hraw = cell(nLabels,1);
  3711. nReaches = zeros(nLabels,1);
  3712. for k = 1:nLabels
  3713. L = uLabels(k);
  3714. sel = strcmp(labelsAll, L);
  3715. T = [trajArray(sel).traj];
  3716. nReaches(k) = numel(T);
  3717. Hk = zeros(numel(P.yEdges)-1, numel(P.xEdges)-1);
  3718. slitVals = []; % collect slit positions
  3719. for t = 1:numel(T)
  3720. if isempty(T(t).x) || isempty(T(t).y), continue; end
  3721. Hk = Hk + histcounts2(T(t).y, T(t).x, P.yEdges, P.xEdges);
  3722. if isfield(T(t),'slitX_norm')
  3723. slitVals(end+1) = T(t).slitX_norm; %#ok<AGROW>
  3724. end
  3725. end
  3726. Hraw{k} = struct('data', Hk, 'slitVals', slitVals);
  3727. end
  3728. % -------- derived heatmaps --------
  3729. Hnorm = cell(nLabels,1); % normalized per reach
  3730. Hprob = cell(nLabels,1); % normalized to probability
  3731. for k = 1:nLabels
  3732. Hk = Hraw{k}; % struct with fields .data and .slitVals
  3733. if nReaches(k) > 0
  3734. Hnorm{k} = struct('data', Hk.data ./ nReaches(k), ...
  3735. 'slitVals', Hk.slitVals);
  3736. else
  3737. Hnorm{k} = Hk; % just copy the struct
  3738. end
  3739. s = sum(Hk.data(:));
  3740. if s > 0
  3741. Hprob{k} = struct('data', Hk.data ./ s, ...
  3742. 'slitVals', Hk.slitVals);
  3743. else
  3744. Hprob{k} = Hk;
  3745. end
  3746. end
  3747. % -------- generate outputs --------
  3748. baseName = sprintf('RAW_%s_Heatmap_%s_%s', upper(viewType(1)), coreID, lower(groupBy));
  3749. plotHeatmapSet(Hnorm, uLabels, viewType, figDir, baseName, 'HM', P, groupBy, [1 nLabels]);
  3750. % plotHeatmapSet(Hprob, uLabels, viewType, figDir, baseName, 'PROB', P, groupBy);
  3751. end
  3752. function plotGroupHeatmaps(GrArray, grpName, viewType, figDir, groupBy)
  3753. if nargin < 5
  3754. groupBy = "label"; % default if not specified
  3755. end
  3756. if isempty(GrArray), return; end
  3757. P = getViewParams(viewType);
  3758. % --- Collect wrappers, attach test_day ---
  3759. trajArray = [];
  3760. for r = 1:numel(GrArray)
  3761. if isfield(GrArray(r), lower(viewType)) && ...
  3762. isfield(GrArray(r).(lower(viewType)), 'trajectories')
  3763. W = GrArray(r).(lower(viewType)).trajectories;
  3764. [W.test_day] = deal(GrArray(r).test_day); % attach test_day
  3765. trajArray = [trajArray W]; %#ok<AGROW>
  3766. end
  3767. end
  3768. if isempty(trajArray), return; end
  3769. % --- Unique days and labels ---
  3770. uDays = string(unique(string({trajArray.test_day}), 'stable'));
  3771. nDays = numel(uDays);
  3772. switch lower(groupBy)
  3773. case "label"
  3774. labelsAll = string({trajArray.label});
  3775. case "broadlabel"
  3776. labelsAll = string({trajArray.broadLabel});
  3777. otherwise
  3778. error("Unknown groupBy option: %s (must be 'label' or 'broadLabel')", groupBy);
  3779. end
  3780. uLabels = string(unique(string(labelsAll), 'stable'));
  3781. nLabels = numel(uLabels);
  3782. % allocate
  3783. Hraw = cell(nDays, nLabels);
  3784. nReaches = zeros(nDays, nLabels);
  3785. % --- accumulate heatmaps per (day,label) ---
  3786. for d = 1:nDays
  3787. daySel = strcmp({trajArray.test_day}, uDays(d));
  3788. for k = 1:nLabels
  3789. labSel = strcmp(labelsAll, uLabels(k));
  3790. sel = daySel & labSel;
  3791. T = [trajArray(sel).traj];
  3792. nReaches(d,k) = numel(T);
  3793. Hk = zeros(numel(P.yEdges)-1, numel(P.xEdges)-1);
  3794. slitVals = []; % collect slit positions
  3795. for t = 1:numel(T)
  3796. if isempty(T(t).x) || isempty(T(t).y), continue; end
  3797. Hk = Hk + histcounts2(T(t).y, T(t).x, P.yEdges, P.xEdges);
  3798. % collect slit positions if available
  3799. if isfield(T(t), 'slitX_norm')
  3800. slitVals(end+1) = T(t).slitX_norm; %#ok<AGROW>
  3801. end
  3802. end
  3803. % wrap into struct so plotHeatmapSet can handle slitVals
  3804. Hraw{d,k} = struct('data', Hk, 'slitVals', slitVals);
  3805. end
  3806. end
  3807. % --- derived heatmaps ---
  3808. Hnorm = cell(size(Hraw));
  3809. Hprob = cell(size(Hraw));
  3810. for d = 1:nDays
  3811. for k = 1:nLabels
  3812. Hk = Hraw{d,k}; % struct with fields .data and .slitVals
  3813. if isempty(Hk), continue; end
  3814. % Normalize per reach
  3815. if nReaches(d,k) > 0
  3816. Hnorm{d,k} = struct('data', Hk.data ./ nReaches(d,k), ...
  3817. 'slitVals', Hk.slitVals);
  3818. else
  3819. Hnorm{d,k} = Hk; % just carry forward struct
  3820. end
  3821. % Normalize to probability
  3822. s = sum(Hk.data(:));
  3823. if s > 0
  3824. Hprob{d,k} = struct('data', Hk.data ./ s, ...
  3825. 'slitVals', Hk.slitVals);
  3826. else
  3827. Hprob{d,k} = Hk;
  3828. end
  3829. end
  3830. end
  3831. % --- label grid ---
  3832. [DD, LL] = ndgrid(string(uDays), string(uLabels)); % days-major
  3833. labelsGrid = strcat(LL, " / ", DD);
  3834. labelsGrid = labelsGrid(:)';
  3835. % --- plot ---
  3836. baseName = sprintf('%s_%s_GroupHeatmap_%s', grpName, upper(viewType(1)), lower(groupBy));
  3837. % Now reshape heatmap array in day-major order
  3838. Hflat = reshape(Hnorm,1,[]); % transpose to flip orientation
  3839. plotHeatmapSet(Hflat, labelsGrid, viewType, figDir, baseName, 'HM', P, groupBy, [nDays nLabels]);
  3840. % Hflat = reshape(Hprob,1,[]);
  3841. % plotHeatmapSet(Hflat, labelsGrid, viewType, figDir, baseName, 'PROB', P, groupBy, [nDays nLabels]);
  3842. end
  3843. function plotGlobalLabelHeatmaps(allResults, uniqueGroups, uniqueDays, labelName, viewType, figDir, groupBy)
  3844. if nargin < 7
  3845. groupBy = "label"; % default
  3846. end
  3847. P = getViewParams(viewType);
  3848. nGroups = numel(uniqueGroups);
  3849. nDays = numel(uniqueDays);
  3850. % -------- accumulate per (day, group) --------
  3851. Hraw = cell(nDays, nGroups);
  3852. nReaches = zeros(nDays, nGroups);
  3853. for d = 1:nDays
  3854. for g = 1:nGroups
  3855. sel = strcmp({allResults.test_day}, uniqueDays(d)) & ...
  3856. strcmp({allResults.group}, uniqueGroups(g));
  3857. Rsub = allResults(sel);
  3858. if isempty(Rsub), continue; end
  3859. % collect trajectories with the given label/broadLabel
  3860. allTraj = []; % <-- will store the *inner* traj structs with x/y
  3861. for r = 1:numel(Rsub)
  3862. if isfield(Rsub(r), lower(viewType)) && ...
  3863. isfield(Rsub(r).(lower(viewType)), 'trajectories')
  3864. Wrappers = Rsub(r).(lower(viewType)).trajectories; % wrapper structs: coreID/reachID/label/broadLabel/traj
  3865. % filter the wrappers by groupBy
  3866. switch lower(groupBy)
  3867. case "label"
  3868. keep = strcmp(string({Wrappers.label}), string(labelName));
  3869. case "broadlabel"
  3870. if isfield(Wrappers, 'broadLabel')
  3871. keep = strcmp(string({Wrappers.broadLabel}), string(labelName));
  3872. else
  3873. keep = false(size(Wrappers));
  3874. end
  3875. otherwise
  3876. error("Unknown groupBy option: %s", groupBy);
  3877. end
  3878. Wrappers = Wrappers(keep);
  3879. if ~isempty(Wrappers)
  3880. % unwrap to inner traj structs (with x/y)
  3881. inner = [Wrappers.traj]; % this is now an array of structs with fields x,y,(maybe slitX_norm)
  3882. allTraj = [allTraj inner];
  3883. end
  3884. end
  3885. end
  3886. if isempty(allTraj), continue; end
  3887. % build heatmap for this (day, group)
  3888. Hk = zeros(numel(P.yEdges)-1, numel(P.xEdges)-1);
  3889. slitVals = []; % collect slit positions
  3890. for t = 1:numel(allTraj)
  3891. if isempty(allTraj(t).x) || isempty(allTraj(t).y), continue; end
  3892. Hk = Hk + histcounts2(allTraj(t).y, allTraj(t).x, P.yEdges, P.xEdges);
  3893. % collect slit positions if available
  3894. if isfield(allTraj(t), 'slitX_norm')
  3895. slitVals(end+1) = allTraj(t).slitX_norm; %#ok<AGROW>
  3896. end
  3897. end
  3898. Hraw{d,g} = struct('data', Hk, 'slitVals', slitVals);
  3899. nReaches(d,g) = numel(allTraj);
  3900. end
  3901. end
  3902. % -------- derived sets --------
  3903. Hnorm = cell(size(Hraw));
  3904. Hprob = cell(size(Hraw));
  3905. groupMax = zeros(1, nGroups);
  3906. for g = 1:nGroups
  3907. vals = [];
  3908. for d = 1:nDays
  3909. if ~isempty(Hnorm{d,g})
  3910. % unwrap struct if needed
  3911. if isstruct(Hnorm{d,g})
  3912. vals = [vals; Hnorm{d,g}.data(:)];
  3913. else
  3914. vals = [vals; Hnorm{d,g}(:)];
  3915. end
  3916. end
  3917. end
  3918. if ~isempty(vals)
  3919. groupMax(g) = max(vals);
  3920. else
  3921. groupMax(g) = 0;
  3922. end
  3923. end
  3924. for d = 1:nDays
  3925. for g = 1:nGroups
  3926. Hk = Hraw{d,g}; % struct with fields .data and .slitVals
  3927. if isempty(Hk), continue; end
  3928. % norm per reach
  3929. if nReaches(d,g) > 0
  3930. Hnorm{d,g} = struct('data', Hk.data ./ nReaches(d,g), ...
  3931. 'slitVals', Hk.slitVals);
  3932. else
  3933. Hnorm{d,g} = Hk;
  3934. end
  3935. % probability density
  3936. s = sum(Hk.data(:));
  3937. if s > 0
  3938. Hprob{d,g} = struct('data', Hk.data ./ s, ...
  3939. 'slitVals', Hk.slitVals);
  3940. else
  3941. Hprob{d,g} = Hk;
  3942. end
  3943. end
  3944. end
  3945. [GG, DD] = ndgrid(string(uniqueGroups), string(uniqueDays)); % group-major
  3946. labelsGrid = strcat(GG, " / ", DD);
  3947. labelsGrid = labelsGrid(:)';
  3948. % -------- Global heatmaps (grid: rows=test_day, cols=group) --------
  3949. baseName = sprintf('RAW_%s_GlobalHeatmap_%s_%s', upper(viewType(1)), labelName, lower(groupBy));
  3950. % NORM
  3951. Hflat = reshape(Hnorm,1,[]);
  3952. plotHeatmapSet(Hflat, labelsGrid, viewType, figDir, baseName, 'HM', P, groupBy, [nDays nGroups], 'pergroup');
  3953. % % PROB
  3954. % Hflat = reshape(Hprob,1,[]);
  3955. % plotHeatmapSet(Hflat, labelsGrid, viewType, figDir, baseName, 'PROB', P, groupBy, [nDays nGroups], 'pergroup');
  3956. end
  3957. function plotHeatmapSet(Hcell, labels, viewType, figDir, outName, tag, P, groupBy, varargin)
  3958. % --- default groupBy ---
  3959. if nargin < 8 || isempty(groupBy)
  3960. groupBy = "label"; % fallback
  3961. end
  3962. cmap = load('C:\Users\juk4004\Documents\MATLAB\myColormaps.mat', 'jet2');
  3963. jet2 = cmap.jet2;
  3964. % --- reorder labels (Success first, then Errors alphabetically) ---
  3965. isSuccess = strcmpi(labels, "Success"); % case-insensitive
  3966. isError = startsWith(labels, "Error", 'IgnoreCase', true);
  3967. others = ~(isSuccess | isError);
  3968. if ~any(isSuccess | isError)
  3969. % Fallback: keep original order
  3970. warning('plotHeatmapSet: no Success/Error labels found, keeping original order');
  3971. % do nothing, labels and Hcell stay as they are
  3972. else
  3973. successLabels = labels(isSuccess);
  3974. errorLabels = sort(labels(isError));
  3975. otherLabels = labels(others);
  3976. labels = [successLabels, errorLabels, otherLabels];
  3977. % Reorder Hcell to match new label order
  3978. newIdx = [find(isSuccess), find(isError), find(others)];
  3979. Hcell = Hcell(newIdx);
  3980. end
  3981. % grid size
  3982. if nargin >= 9 && ~isempty(varargin) && ~isempty(varargin{1}) && isnumeric(varargin{1})
  3983. nRows = varargin{1}(1);
  3984. nCols = varargin{1}(2);
  3985. varargin(1) = []; % pop it off so later args shift down
  3986. else
  3987. nLabels = numel(labels);
  3988. nCols = min(3, nLabels);
  3989. nRows = ceil(nLabels/nCols);
  3990. end
  3991. % scaling mode
  3992. if ~isempty(varargin) && ischar(varargin{1}) && strcmpi(varargin{1}, 'pergroup')
  3993. scaleMode = 'pergroup';
  3994. else
  3995. scaleMode = 'global';
  3996. end
  3997. % compute global/group maxima for caxis scaling
  3998. if strcmp(scaleMode, 'pergroup')
  3999. groupMax = zeros(1, nRows);
  4000. for r = 1:nRows
  4001. vals = [];
  4002. for c = 1:nCols
  4003. idx = (r-1)*nCols + c;
  4004. if idx <= numel(Hcell) && ~isempty(Hcell{idx})
  4005. if isstruct(Hcell{idx})
  4006. vals = [vals; Hcell{idx}.data(:)];
  4007. else
  4008. vals = [vals; Hcell{idx}(:)];
  4009. end
  4010. end
  4011. end
  4012. if ~isempty(vals)
  4013. groupMax(r) = max(vals);
  4014. end
  4015. end
  4016. else
  4017. nonempty = Hcell(~cellfun('isempty',Hcell));
  4018. if all(cellfun(@isstruct, nonempty))
  4019. groupMax = max(cellfun(@(m) max(m.data(:)), nonempty));
  4020. else
  4021. groupMax = max(cellfun(@(m) max(m(:)), nonempty));
  4022. end
  4023. end
  4024. % Each subplot ~250 px wide, ~250 px tall
  4025. figW = 250 * nCols;
  4026. figH = 250 * nRows;
  4027. f = figure('Visible','off','Color','w', ...
  4028. 'Name', sprintf('%s Heatmaps (%s)', outName, tag), ...
  4029. 'Position', [100 100 figW figH]);
  4030. for k = 1:numel(labels)
  4031. % compute row, col in column-major order
  4032. r = mod(k-1, nRows) + 1; % row index
  4033. c = floor((k-1)/nRows) + 1; % col index
  4034. subplot(nRows, nCols, (r-1)*nCols + c);
  4035. if isempty(Hcell{k})
  4036. axis off; continue;
  4037. end
  4038. % unwrap struct vs numeric
  4039. if isstruct(Hcell{k})
  4040. Hdata = Hcell{k}.data;
  4041. else
  4042. Hdata = Hcell{k};
  4043. end
  4044. % Pad Hdata by repeating last row and col
  4045. Hpad = [Hdata, Hdata(:,end)]; % add last column again
  4046. Hpad = [Hpad; Hpad(end,:)]; % add last row again
  4047. [X,Y] = meshgrid(P.xEdges, P.yEdges); % 51x51 if Hdata is 50x50
  4048. pcolor(X, Y, Hpad);
  4049. shading flat; % avoid grid lines
  4050. axis xy; set(gca,'YDir','reverse');
  4051. xlim(P.xlim); ylim(P.ylim);
  4052. colormap(jet2); colorbar;
  4053. % scaling
  4054. [r,c] = ind2sub([nRows nCols], k);
  4055. if strcmp(scaleMode, 'pergroup')
  4056. if groupMax(r) > 0, caxis([0 groupMax(r)]); end
  4057. else
  4058. if groupMax > 0, caxis([0 groupMax]); end
  4059. end
  4060. title(sprintf('%s - %s', tag, labels(k)), 'Interpreter','none'); hold on;
  4061. % pellet marker
  4062. scatter(0,0,60,'o','MarkerEdgeColor','w','MarkerFaceColor','w','LineWidth',1.25);
  4063. % slit line (Side view only)
  4064. if P.showSlit && isstruct(Hcell{k}) && isfield(Hcell{k},'slitVals')
  4065. mSlit = mean(Hcell{k}.slitVals,'omitnan');
  4066. if ~isnan(mSlit)
  4067. plot([mSlit mSlit], get(gca,'YLim'), ':','LineWidth',1.5,'Color',[0.7 0.7 0.7]);
  4068. end
  4069. end
  4070. end
  4071. outDir= fullfile(figDir, char(groupBy), viewType, tag);
  4072. if ~exist(outDir, 'dir'), mkdir(outDir); end
  4073. savename = fullfile(outDir, outName);
  4074. set(f, 'PaperUnits', 'inches', 'PaperPosition', [0 0 nCols*3 nRows*3]);
  4075. set(f, 'Color', 'w'); % figure background
  4076. export_fig(savename, '-png','-pdf', '-r300', f);
  4077. close(f);
  4078. end
  4079. function plotDifferenceHeatmaps(allResults, uniqueGroups, uniqueDays, uniqueLabels, uniqueBroadLabels, viewType, figDir, groupBy)
  4080. % Plot difference heatmaps (Drug - Baseline) for broadLabel, label, and global levels
  4081. %
  4082. % Debugging version: prints info about groups, labels, counts, and skips.
  4083. P = getViewParams(viewType);
  4084. dayBaseline = "Baseline";
  4085. dayDrug = "Drug";
  4086. cmap = load('C:\Users\juk4004\Documents\MATLAB\myColormaps.mat', 'jet2');
  4087. jet2 = cmap.jet2;
  4088. % normalize label inputs
  4089. if ischar(uniqueLabels) || isstring(uniqueLabels)
  4090. uniqueLabels = string(uniqueLabels);
  4091. end
  4092. if ischar(uniqueBroadLabels) || isstring(uniqueBroadLabels)
  4093. uniqueBroadLabels = cellstr(uniqueBroadLabels);
  4094. end
  4095. % decide which set of categories to use
  4096. switch lower(groupBy)
  4097. case 'label'
  4098. labelSet = uniqueLabels;
  4099. fieldName = 'label';
  4100. % Custom reordering: "Success" first, then Errors alphabetically
  4101. isSuccess = strcmp(labelSet, "Success");
  4102. isError = startsWith(labelSet, "Error");
  4103. successLabels = labelSet(isSuccess);
  4104. errorLabels = sort(labelSet(isError));
  4105. % force row orientation to avoid horzcat dimension mismatch
  4106. successLabels = successLabels(:).';
  4107. errorLabels = errorLabels(:).';
  4108. % Preserve "Success" first, then sorted Errors
  4109. labelSet = [successLabels, errorLabels];
  4110. case 'broadlabel'
  4111. labelSet = uniqueBroadLabels;
  4112. fieldName = 'broadLabel';
  4113. % For broadLabel: "Success" first if present, then the rest sorted
  4114. isSuccess = strcmp(labelSet, "Success");
  4115. successLabels = labelSet(isSuccess);
  4116. rest = labelSet(~isSuccess);
  4117. successLabels = successLabels(:).';
  4118. rest = rest(:).';
  4119. labelSet = [successLabels, sort(rest)];
  4120. otherwise
  4121. error('groupBy must be ''label'' or ''broadLabel''');
  4122. end
  4123. % Prepare storage
  4124. heatmapsBase = cell(numel(uniqueGroups), numel(labelSet));
  4125. heatmapsDrug = cell(numel(uniqueGroups), numel(labelSet));
  4126. % -------- Build per-group, per-label histograms --------
  4127. for g = 1:numel(uniqueGroups)
  4128. grpName = uniqueGroups(g);
  4129. for l = 1:numel(labelSet)
  4130. labelName = labelSet{l};
  4131. % Initialize
  4132. Hbase = zeros(numel(P.yEdges)-1, numel(P.xEdges)-1);
  4133. Hdrug = zeros(size(Hbase));
  4134. slitBase = [];
  4135. slitDrug = [];
  4136. countBase = 0;
  4137. countDrug = 0;
  4138. % Walk through results
  4139. for i = 1:numel(allResults)
  4140. R = allResults(i);
  4141. if string(R.group) ~= string(grpName), continue; end
  4142. viewField = lower(viewType);
  4143. if ~isfield(R, viewField) || ~isfield(R.(viewField), 'trajectories')
  4144. continue;
  4145. end
  4146. trajArray = R.(viewField).trajectories;
  4147. if isempty(trajArray), continue; end
  4148. selIdx = strcmp(string({trajArray.(fieldName)}), string(labelName));
  4149. trajArray = trajArray(selIdx);
  4150. if isempty(trajArray), continue; end
  4151. if string(R.test_day) == dayBaseline
  4152. for t = 1:numel(trajArray)
  4153. xt = trajArray(t).traj.x;
  4154. yt = trajArray(t).traj.y;
  4155. if isempty(xt) || isempty(yt), continue; end
  4156. Hk = histcounts2(yt, xt, P.yEdges, P.xEdges);
  4157. Hbase = Hbase + Hk;
  4158. countBase = countBase + 1;
  4159. if isfield(trajArray(t).traj, 'slitX_norm')
  4160. slitBase(end+1) = trajArray(t).traj.slitX_norm; %#ok<AGROW>
  4161. end
  4162. end
  4163. elseif string(R.test_day) == dayDrug
  4164. for t = 1:numel(trajArray)
  4165. xt = trajArray(t).traj.x;
  4166. yt = trajArray(t).traj.y;
  4167. if isempty(xt) || isempty(yt), continue; end
  4168. Hk = histcounts2(yt, xt, P.yEdges, P.xEdges);
  4169. Hdrug = Hdrug + Hk;
  4170. countDrug = countDrug + 1;
  4171. if isfield(trajArray(t).traj, 'slitX_norm')
  4172. slitDrug(end+1) = trajArray(t).traj.slitX_norm; %#ok<AGROW>
  4173. end
  4174. end
  4175. end
  4176. end
  4177. % Normalize
  4178. if countBase > 0, Hbase = Hbase / countBase; end
  4179. if countDrug > 0, Hdrug = Hdrug / countDrug; end
  4180. % Pack into structs for later use
  4181. heatmapsBase{g,l} = struct('data', Hbase, 'slitVals', slitBase);
  4182. heatmapsDrug{g,l} = struct('data', Hdrug, 'slitVals', slitDrug);
  4183. end
  4184. end
  4185. % -------- Compute and plot differences --------
  4186. for g = 1:numel(uniqueGroups)
  4187. grpName = string(uniqueGroups(g));
  4188. % Precompute shared scale across all labels for Baseline/Drug
  4189. allMaxVals = [];
  4190. for l = 1:numel(labelSet)
  4191. tmpBase = heatmapsBase{g, l}.data;
  4192. tmpDrug = heatmapsDrug{g, l}.data;
  4193. allMaxVals = [allMaxVals; tmpBase(:); tmpDrug(:)];
  4194. end
  4195. sharedMaxVal = max(allMaxVals);
  4196. if sharedMaxVal == 0, sharedMaxVal = 1; end
  4197. % ===== 1) Individual per-label figs =====
  4198. for l = 1:numel(labelSet)
  4199. labelName = string(labelSet{l});
  4200. % Peel struct into numeric + slit arrays
  4201. tmp = heatmapsBase{g, l};
  4202. Hbase = tmp.data;
  4203. slitBase = tmp.slitVals;
  4204. tmp = heatmapsDrug{g, l};
  4205. Hdrug = tmp.data;
  4206. slitDrug = tmp.slitVals;
  4207. diffHeatmap = Hdrug - Hbase;
  4208. % Skip if trivial
  4209. if all(diffHeatmap(:) == 0)
  4210. fprintf(' WARNING: diffHeatmap all zeros for group=%s, label=%s\n', char(grpName), char(labelName));
  4211. continue;
  4212. end
  4213. % shared scale for baseline & drug
  4214. maxVal = max([Hbase(:); Hdrug(:)]);
  4215. if maxVal == 0, maxVal = 1; end
  4216. % Custom diverging colormap: blue (#2596be) → white → red (#be2525)
  4217. nColors = 256;
  4218. mid = round(nColors/2);
  4219. blueRGB = [37 150 190] / 255; % #2596be
  4220. redRGB = [190 37 37] / 255; % #be2525
  4221. blue2white = [linspace(blueRGB(1),1,mid)', ...
  4222. linspace(blueRGB(2),1,mid)', ...
  4223. linspace(blueRGB(3),1,mid)'];
  4224. white2red = [linspace(1,redRGB(1),mid)', ...
  4225. linspace(1,redRGB(2),mid)', ...
  4226. linspace(1,redRGB(3),mid)'];
  4227. cmapDiff = [blue2white; white2red];
  4228. % figure with 3 panels
  4229. figure('Visible','off');
  4230. tiledlayout(1,3);
  4231. [X, Y] = meshgrid(P.xEdges, P.yEdges);
  4232. % Baseline
  4233. nexttile;
  4234. Hpad = [Hbase, Hbase(:,end)];
  4235. Hpad = [Hpad; Hpad(end,:)];
  4236. pcolor(X, Y, Hpad);
  4237. shading flat;
  4238. axis xy; set(gca,'YDir','reverse');
  4239. colormap(gca, jet2); colorbar; caxis([0 maxVal]);
  4240. if P.showSlit
  4241. mSlitBase = mean(slitBase, 'omitnan'); %#ok<NASGU>
  4242. end
  4243. title('Baseline');
  4244. % Drug
  4245. nexttile;
  4246. Hpad = [Hdrug, Hdrug(:,end)];
  4247. Hpad = [Hpad; Hpad(end,:)];
  4248. pcolor(X, Y, Hpad);
  4249. shading flat;
  4250. axis xy; set(gca,'YDir','reverse');
  4251. colormap(gca, jet2); colorbar; caxis([0 maxVal]);
  4252. if P.showSlit
  4253. mSlitDrug = mean(slitDrug, 'omitnan'); %#ok<NASGU>
  4254. end
  4255. title('Drug');
  4256. % Difference
  4257. nexttile;
  4258. Hpad = [diffHeatmap, diffHeatmap(:,end)];
  4259. Hpad = [Hpad; Hpad(end,:)];
  4260. pcolor(X, Y, Hpad);
  4261. shading flat;
  4262. axis xy; set(gca,'YDir','reverse');
  4263. colormap(gca, cmapDiff); colorbar;
  4264. clim = max(abs(diffHeatmap(:))); if clim==0, clim=1; end
  4265. caxis([-clim clim]);
  4266. title('Drug - Baseline');
  4267. % overlay average slit (if side view)
  4268. if P.showSlit
  4269. allSlits = [slitBase, slitDrug];
  4270. if ~isempty(allSlits)
  4271. mSlit = mean(allSlits, 'omitnan');
  4272. hold on;
  4273. plot([mSlit mSlit], get(gca,'YLim'), ':', 'LineWidth', 1.5, 'Color', [0.7 0.7 0.7]);
  4274. hold off;
  4275. end
  4276. end
  4277. sgtitle(sprintf('Group: %s | Label: %s | View: %s', grpName, labelName, viewType));
  4278. outDir = fullfile(figDir, 'DifferenceHeatmaps', groupBy, char(viewType));
  4279. if ~exist(outDir, 'dir'), mkdir(outDir); end
  4280. savename = fullfile(outDir, sprintf('DiffHeatmap_%s_%s_%s', char(grpName), char(labelName), viewType));
  4281. set(gcf, 'Color', 'w');
  4282. export_fig(savename, '-pdf', '-png', '-r300', gcf);
  4283. close;
  4284. end
  4285. % ===== 2) Group-level combined fig =====
  4286. figure('Visible','off'); tiledlayout(3, numel(labelSet), 'TileSpacing','compact');
  4287. % Reuse the same diverging cmap for consistency
  4288. nColors = 256;
  4289. mid = round(nColors/2);
  4290. blueRGB = [37 150 190] / 255; % #2596be
  4291. redRGB = [190 37 37] / 255; % #be2525
  4292. blue2white = [linspace(blueRGB(1),1,mid)', ...
  4293. linspace(blueRGB(2),1,mid)', ...
  4294. linspace(blueRGB(3),1,mid)'];
  4295. white2red = [linspace(1,redRGB(1),mid)', ...
  4296. linspace(1,redRGB(2),mid)', ...
  4297. linspace(1,redRGB(3),mid)'];
  4298. cmapDiff = [blue2white; white2red];
  4299. for l = 1:numel(labelSet)
  4300. labelName = string(labelSet{l});
  4301. % Peel struct into numeric + slit arrays
  4302. tmp = heatmapsBase{g, l};
  4303. Hbase = tmp.data;
  4304. tmp = heatmapsDrug{g, l};
  4305. Hdrug = tmp.data;
  4306. diffHeatmap = Hdrug - Hbase;
  4307. maxVal = sharedMaxVal;
  4308. clim = max(abs(diffHeatmap(:))); if clim==0, clim=1; end
  4309. [X, Y] = meshgrid(P.xEdges, P.yEdges);
  4310. % row 1 = Baseline
  4311. nexttile(l);
  4312. Hpad = [Hbase, Hbase(:,end)];
  4313. Hpad = [Hpad; Hpad(end,:)];
  4314. pcolor(X, Y, Hpad);
  4315. shading flat;
  4316. axis xy; set(gca,'YDir','reverse');
  4317. colormap(gca, jet2); colorbar; caxis([0 maxVal]);
  4318. if l==1, ylabel('Baseline'); end
  4319. title(labelName);
  4320. % row 2 = Drug
  4321. nexttile(l+numel(labelSet));
  4322. Hpad = [Hdrug, Hdrug(:,end)];
  4323. Hpad = [Hpad; Hpad(end,:)];
  4324. pcolor(X, Y, Hpad);
  4325. shading flat;
  4326. axis xy; set(gca,'YDir','reverse');
  4327. colormap(gca, jet2); colorbar; caxis([0 maxVal]);
  4328. if l==1, ylabel('Drug'); end
  4329. % row 3 = Difference
  4330. nexttile(l+2*numel(labelSet));
  4331. Hpad = [diffHeatmap, diffHeatmap(:,end)];
  4332. Hpad = [Hpad; Hpad(end,:)];
  4333. pcolor(X, Y, Hpad);
  4334. shading flat;
  4335. axis xy; set(gca,'YDir','reverse');
  4336. colormap(gca, cmapDiff); colorbar; caxis([-clim clim]);
  4337. if l==1, ylabel('Drug - Baseline'); end
  4338. end
  4339. sgtitle(sprintf('Group: %s | View: %s (%s)', grpName, viewType, groupBy));
  4340. outDir = fullfile(figDir, 'DifferenceHeatmaps', groupBy, char(viewType));
  4341. if ~exist(outDir, 'dir'), mkdir(outDir); end
  4342. savename = fullfile(outDir, sprintf('DiffHeatmapSet_ALLLABELS_%s_%s', grpName, viewType));
  4343. set(gcf, 'Color', 'w');
  4344. export_fig(savename, '-pdf', '-png', '-r300', gcf);
  4345. close; % combined fig
  4346. end
  4347. end
  4348. function plotGlobalAndDifferenceHeatmaps(allResults, uniqueGroups, viewType, figDir)
  4349. % Plot per-group global heatmaps (Baseline, Drug, Difference) in one figure
  4350. % with consistent scaling, colormaps, and slit overlays.
  4351. P = getViewParams(viewType);
  4352. dayBaseline = "Baseline";
  4353. dayDrug = "Drug";
  4354. cmap = load('C:\Users\juk4004\Documents\MATLAB\myColormaps.mat', 'jet2');
  4355. jet2 = cmap.jet2;
  4356. for g = 1:numel(uniqueGroups)
  4357. grpName = string(uniqueGroups(g));
  4358. % Initialize
  4359. Hbase = zeros(numel(P.yEdges)-1, numel(P.xEdges)-1);
  4360. Hdrug = zeros(size(Hbase));
  4361. slitBase = [];
  4362. slitDrug = [];
  4363. countBase = 0;
  4364. countDrug = 0;
  4365. % Collect all reaches for this group
  4366. for i = 1:numel(allResults)
  4367. R = allResults(i);
  4368. if string(R.group) ~= grpName
  4369. continue;
  4370. end
  4371. viewField = lower(viewType);
  4372. if ~isfield(R, viewField) || ~isfield(R.(viewField), 'trajectories')
  4373. continue;
  4374. end
  4375. trajArray = R.(viewField).trajectories;
  4376. if isempty(trajArray), continue; end
  4377. thisDay = string(R.test_day);
  4378. if thisDay == dayBaseline
  4379. for t = 1:numel(trajArray)
  4380. xt = trajArray(t).traj.x;
  4381. yt = trajArray(t).traj.y;
  4382. if isempty(xt) || isempty(yt), continue; end
  4383. Hk = histcounts2(yt, xt, P.yEdges, P.xEdges);
  4384. Hbase = Hbase + Hk;
  4385. countBase = countBase + 1;
  4386. if isfield(trajArray(t).traj,'slitX_norm')
  4387. slitBase(end+1) = trajArray(t).traj.slitX_norm; %#ok<AGROW>
  4388. end
  4389. end
  4390. elseif thisDay == dayDrug
  4391. for t = 1:numel(trajArray)
  4392. xt = trajArray(t).traj.x;
  4393. yt = trajArray(t).traj.y;
  4394. if isempty(xt) || isempty(yt), continue; end
  4395. Hk = histcounts2(yt, xt, P.yEdges, P.xEdges);
  4396. Hdrug = Hdrug + Hk;
  4397. countDrug = countDrug + 1;
  4398. if isfield(trajArray(t).traj,'slitX_norm')
  4399. slitDrug(end+1) = trajArray(t).traj.slitX_norm; %#ok<AGROW>
  4400. end
  4401. end
  4402. end
  4403. end
  4404. % Normalize
  4405. if countBase > 0, Hbase = Hbase / countBase; end
  4406. if countDrug > 0, Hdrug = Hdrug / countDrug; end
  4407. diffMap = Hdrug - Hbase;
  4408. % Shared scale for Baseline/Drug
  4409. maxVal = max([Hbase(:); Hdrug(:)]);
  4410. if maxVal == 0, maxVal = 1; end
  4411. % Custom red-white-blue colormap for differences
  4412. n = 128;
  4413. blue = [37 150 190] / 255; % #2596be
  4414. red = [190 37 37] / 255; % #be2525
  4415. white = [1 1 1];
  4416. cmapDiff = [linspace(blue(1),white(1),n)' linspace(blue(2),white(2),n)' linspace(blue(3),white(3),n)';
  4417. linspace(white(1),red(1),n)' linspace(white(2),red(2),n)' linspace(white(3),red(3),n)'];
  4418. % Create one figure with 3 panels
  4419. figure('Visible','off');
  4420. tiledlayout(1,3);
  4421. [X, Y] = meshgrid(P.xEdges, P.yEdges);
  4422. % Panel 1: Baseline
  4423. nexttile;
  4424. Hpad = [Hbase, Hbase(:,end)];
  4425. Hpad = [Hpad; Hpad(end,:)];
  4426. pcolor(X, Y, Hpad);
  4427. shading flat;
  4428. axis xy; set(gca,'YDir','reverse');
  4429. colorbar; colormap(gca, jet2);
  4430. caxis([0 maxVal]);
  4431. title('Baseline');
  4432. if P.showSlit && ~isempty(slitBase)
  4433. mSlitBase = mean(slitBase,'omitnan');
  4434. plot([mSlitBase mSlitBase], get(gca,'YLim'), ':','LineWidth',1.5,'Color',[0.7 0.7 0.7]);
  4435. end
  4436. % Panel 2: Drug
  4437. nexttile;
  4438. Hpad = [Hdrug, Hdrug(:,end)];
  4439. Hpad = [Hpad; Hpad(end,:)];
  4440. pcolor(X, Y, Hpad);
  4441. shading flat;
  4442. axis xy; set(gca,'YDir','reverse');
  4443. colorbar; colormap(gca, jet2);
  4444. caxis([0 maxVal]);
  4445. title('Drug');
  4446. if P.showSlit && ~isempty(slitDrug)
  4447. mSlitDrug = mean(slitDrug,'omitnan');
  4448. plot([mSlitDrug mSlitDrug], get(gca,'YLim'), ':','LineWidth',1.5,'Color',[0.7 0.7 0.7]);
  4449. end
  4450. % Panel 3: Difference
  4451. nexttile;
  4452. Hpad = [diffMap, diffMap(:,end)];
  4453. Hpad = [Hpad; Hpad(end,:)];
  4454. pcolor(X, Y, Hpad);
  4455. shading flat;
  4456. axis xy; set(gca,'YDir','reverse');
  4457. colorbar; colormap(gca, cmapDiff);
  4458. clim = max(abs(diffMap(:))); if clim==0, clim=1; end
  4459. caxis([-clim clim]);
  4460. title('Drug - Baseline');
  4461. if P.showSlit
  4462. allSlits = [slitBase slitDrug];
  4463. if ~isempty(allSlits)
  4464. mSlit = mean(allSlits,'omitnan');
  4465. plot([mSlit mSlit], get(gca,'YLim'), ':','LineWidth',1.5,'Color',[0.7 0.7 0.7]);
  4466. end
  4467. end
  4468. sgtitle(sprintf('Group: %s | View: %s', grpName, viewType));
  4469. % Save per-group figure
  4470. outDir = fullfile(figDir, 'DifferenceHeatmaps', char(viewType));
  4471. if ~exist(outDir,'dir'), mkdir(outDir); end
  4472. savename = fullfile(outDir, sprintf('AllReaches_%s_%s', grpName, viewType));
  4473. set(gcf, 'Color', 'w'); % figure background
  4474. export_fig(savename, '-pdf', '-png', '-r300', gcf);
  4475. close;
  4476. end
  4477. end

ReachingClassificationGUI_reach_hm2.m at commit 8eecee1, no license · at the source

Overview

Authors: Julia Kaiser1,2, Payal Patel1, Sam Fedde1, Alexander Lammers1, Matthew Kenwood3, Asim Iqbal1,4, Mark P Goldberg3, Vibhu Sahni1,2,5
  1. Burke Neurological Institute, White Plains, NY USA
  2. Feil Family Brain and Mind Research Institute, Weill Cornell Medicine, New York, NY USA
  3. Department of Neurology, UT San Antonio, San Antonio, TX USA
  4. Tibbling Technologies, Seattle, WA USA
  5. Weill Cornell Graduate School of Medical Sciences, New York, NY USA
Journal: Nature communications, volume 17, issue 1, article 6716
Dates: received 18 March 2025; accepted 7 May 2026; published online 22 May 2026
Type: Research article · Language: English
License: CC BY-NC-ND
Identifiers: DOI 10.1038/s41467-026-73476-4 · PMID 42173922 · PMCID PMC13385823 · OpenAlex W7162142989
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: mouse (organism)
Methods: Statistics, Smoothing, state filtering, decompositions, Machine learning, Preprocessing, Evoked potentials
Keywords: Development of the nervous system, Motor control, Neural circuits
MeSH: Brain Stem*, Cerebral Cortex*, Forelimb*, Motor Cortex*, Animals, Female, Male, Mice, Movement, Neural Pathways, Neurons, Neuropeptide Y, Spinal Cord (* major topic)
Topic: Zebrafish Biomedical Research Applications (Cell Biology, Biochemistry, Genetics and Molecular Biology), according to OpenAlex
Funding: New York State Department of Health - Wadsworth Center (Department of Health, Wadsworth Center) (C39069GG); U.S. Department of Health & Human Services | NIH | NIH Office of the Director (OD) (S10OD036432, S10OD030383); U.S. Department of Health & Human Services | NIH | National Institute of Neurological Disorders and Stroke (NINDS) (R01NS131662); NICHD NIH HHS (K12 HD093427); U.S. Department of Health & Human Services | NIH | Eunice Kennedy Shriver National Institute of Child Health and Human Development (NICHD) (K12HD093427); Swiss National Science Foundation (P2EZP3_191858, 191858); NINDS NIH HHS (R01 NS131662); NIH HHS (S10 OD030383, S10 OD036432); Craig H. Neilsen Foundation (Neilsen Foundation) (727694)
Citations: not cited yet (Europe PMC); 109 references in the paper

Abstract

The abstract is not reproduced here: the paper's license (CC BY-NC-ND) does not allow it. Read it in the paper, at the publisher or on Europe PMC.

Repositories

Its files are read in the Code ↔ Paper reader above, with 5 matches between paragraphs and lines of code.

jkaiser87

License: none: the authors keep all their rights
State: the link answers, verified on 28 September 2026
Evidence: the link answers
Software Heritage: not checked
Found in: “Code availability”
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Availability: 1 check, the latest on 28 September 2026: the link answers (HTTP 200)
  • 28 September 2026: the link answers (HTTP 200)
At the source: github.com/jkaiser87/

itsasimiqbal/StARQ

License: MIT
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Commit: afc803c98ce8af21cc8bfb4d4c2a7108c99a71fc, 3 October 2024
Languages: Jupyter (4)
Size: 10 files, 4 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README, license file, 4 notebooks
Not found: CITATION.cff, environment file, tests, continuous integration, documentation
Tools: Matplotlib (3 files), NumPy (3 files), OpenCV (3 files), PyTorch (3 files), pandas (1 file), scikit-image (1 file), tifffile (1 file)
Availability: 1 check, the latest on 28 September 2026: the link answers
  • 28 September 2026: the link answers
5 files

Zenodo 19409819

License: CC-BY-4.0
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Size: 1 file
Software Heritage: not checked
Found in: the references
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Availability: 1 check, the latest on 28 September 2026: the link answers (HTTP 200)
  • 28 September 2026: the link answers (HTTP 200)
3 files
At the source:

Zenodo 19409814

License: CC-BY-4.0
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Size: 1 file
Software Heritage: not checked
Found in: the references
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Availability: 1 check, the latest on 28 September 2026: the link answers (HTTP 200)
  • 28 September 2026: the link answers (HTTP 200)
3 files
At the source:

Zenodo 19409800

License: CC-BY-4.0
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Size: 1 file
Software Heritage: not checked
Found in: the references
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Availability: 1 check, the latest on 28 September 2026: the link answers (HTTP 200)
  • 28 September 2026: the link answers (HTTP 200)
2 files
At the source:

jkaiser87/sahni_spg_analysis

License: none: the authors keep all their rights
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Commit: 8eecee124631747513d6dacb060c0b124995d79b, 6 April 2026
Languages: MATLAB (1)
Size: 6 files, 1 script
Software Heritage: not archived
Found in: the Zenodo archive record
Holds: README, CITATION.cff
Not found: license file, environment file, tests, continuous integration, documentation
Availability: 1 check, the latest on 28 September 2026: the link answers
  • 28 September 2026: the link answers
2 files

jkaiser87/cell3d

License: none: the authors keep all their rights
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Commit: 8c9c7e2e4129d8eb21929d7086563276756192cf, 6 April 2026
Languages: MATLAB (2)
Size: 49 files, 2 scripts
Software Heritage: not archived
Found in: the Zenodo archive record
Holds: README, CITATION.cff
Not found: license file, environment file, tests, continuous integration, documentation
Availability: 1 check, the latest on 28 September 2026: the link answers
  • 28 September 2026: the link answers
3 files

jkaiser87/vol3d

License: none: the authors keep all their rights
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Commit: 2a4bc77547ed1a6cf4ec462ce67307e5fcc2c89b, 6 April 2026
Languages: MATLAB (2)
Size: 41 files, 2 scripts
Software Heritage: not archived
Found in: the Zenodo archive record
Holds: README, CITATION.cff
Not found: license file, environment file, tests, continuous integration, documentation
Availability: 1 check, the latest on 28 September 2026: the link answers
  • 28 September 2026: the link answers
3 files

Code availability statement

The paper has a code availability statement. Its license (CC BY-NC-ND) does not allow reproducing it here; in short, from what the harvester recognized in it:

Read it in the paper: doi.org/10.1038/s41467-026-73476-4.

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 8 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 13 scripts, each with its path and the digest of its content;
  • 5 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.

Data

Datasets cited

Code and data availability statement

The paper has a code and data availability statement. Its license (CC BY-NC-ND) does not allow reproducing it here; in short, from what the harvester recognized in it:

Read it in the paper: doi.org/10.1038/s41467-026-73476-4.

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 1, 28 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 8 authors, 3 keywords, 13 MeSH terms, 9 funders, 103 references.

Cite

This paper

Kaiser, J., Patel, P., Fedde, S., Lammers, A., Kenwood, M., Iqbal, A., Goldberg, M. P., & Sahni, V. (2026). Developmental molecular signatures define de novo cortico-brainstem circuit for skilled forelimb movement. Nature communications, 17(1), 6716. https://doi.org/10.1038/s41467-026-73476-4

BibTeX

@article{kaiser2026developmental,
author = {Kaiser, Julia and Patel, Payal and Fedde, Sam and Lammers, Alexander and Kenwood, Matthew and Iqbal, Asim and Goldberg, Mark P and Sahni, Vibhu},
title = {{Developmental molecular signatures define de novo cortico-brainstem circuit for skilled forelimb movement}},
journal = {Nature communications},
year = {2026},
month = may,
volume = {17},
number = {1},
pages = {6716},
publisher = {Nature Publishing Group},
issn = {2041-1723},
doi = {10.1038/s41467-026-73476-4},
url = {https://doi.org/10.1038/s41467-026-73476-4},
pmid = {42173922},
pmcid = {PMC13385823}
}

RIS

TY - JOUR
AU - Kaiser, Julia
AU - Patel, Payal
AU - Fedde, Sam
AU - Lammers, Alexander
AU - Kenwood, Matthew
AU - Iqbal, Asim
AU - Goldberg, Mark P
AU - Sahni, Vibhu
TI - Developmental molecular signatures define de novo cortico-brainstem circuit for skilled forelimb movement
T2 - Nature communications
J2 - Nat Commun
PY - 2026
DA - 2026/05/22
VL - 17
IS - 1
SP - 6716
SN - 2041-1723
PB - Nature Publishing Group
DO - 10.1038/s41467-026-73476-4
UR - https://doi.org/10.1038/s41467-026-73476-4
LA - en
ER -

CSL-JSON

{
"id": "10.1038/s41467-026-73476-4",
"type": "article-journal",
"title": "Developmental molecular signatures define de novo cortico-brainstem circuit for skilled forelimb movement",
"container-title": "Nature communications",
"author": [
{
"family": "Kaiser",
"given": "Julia"
},
{
"family": "Patel",
"given": "Payal"
},
{
"family": "Fedde",
"given": "Sam"
},
{
"family": "Lammers",
"given": "Alexander"
},
{
"family": "Kenwood",
"given": "Matthew"
},
{
"family": "Iqbal",
"given": "Asim"
},
{
"family": "Goldberg",
"given": "Mark P"
},
{
"family": "Sahni",
"given": "Vibhu"
}
],
"container-title-short": "Nat Commun",
"volume": "17",
"issue": "1",
"page": "6716",
"DOI": "10.1038/s41467-026-73476-4",
"PMID": "42173922",
"PMCID": "PMC13385823",
"ISSN": "2041-1723",
"publisher": "Nature Publishing Group",
"URL": "https://doi.org/10.1038/s41467-026-73476-4",
"language": "en",
"issued": {
"date-parts": [
[
2026,
5,
22
]
]
}
}

The tracing map gets a citation of its own once an author has validated it and it has a DOI.

Similar papers

The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.

[1] doi:10.7554/elife.109240 [code]
Neural activity profiles reveal overlapping, intermingled subpopulations spanning area borders in mouse sensorimotor cortex.
Journal: eLife
In common: Image Processing Toolbox, mouse, 11 references
[2] doi:10.1038/s41586-026-10679-1 [code]
Cortical development dynamics across autism spectrum disorder mouse models.
Journal: Nature
In common: tifffile, OpenCV, scikit-image, 3 other tools, mouse, 6 references
[3] doi:10.1016/j.celrep.2026.117419 [code]
Conserved role of primary motor cortex in the control of prehension in mice and macaques.
Journal: Cell reports
In common: OpenCV, pandas, Matplotlib, 1 other tool, mouse, 6 references
[4] doi:10.1038/s41593-026-02232-0 [code]
Entorhinal cortex represents task-relevant remote locations independently of CA1.
Journal: Nature neuroscience
In common: OpenCV, scikit-image, Image Processing Toolbox, 5 other tools, mouse, 3 references
[5] doi:10.1186/s12974-026-03885-1 [code]
Shared transcriptomic signatures in perilesional and contralesional cortex after ischemic stroke.
Journal: Journal of neuroinflammation
In common: Image Processing Toolbox, mouse, 2 authors
[6] doi:10.1038/s41467-026-74823-1 [code]
Cerebellar activity is triggered by reach endpoint during learning of a complex locomotor task.
Journal: Nature communications
In common: tifffile, OpenCV, scikit-image, 3 other tools, mouse, 3 references
[7] doi:10.1038/s41467-026-74569-w [code]
Motor cortex directly excites the substantia nigra pars reticulata, the basal ganglia output nucleus.
Journal: Nature communications
In common: pandas, NumPy, mouse, 7 references
[8] doi: [code]
Real-time closed-loop feedback system for mouse mesoscale cortical signal and movement control
Journal: eLife
In common: tifffile, OpenCV, scikit-image, 3 other tools, mouse, 3 references
[9] doi:10.1126/sciadv.adw5487 [code]
Early fate diversification of radial glial progenitors during corticogenesis.
Journal: Science advances
In common: mouse, 7 references
[10] doi:10.1016/j.celrep.2026.117646 [code]
Medial entorhinal-hippocampal desynchronization parallels the emergence of memory impairment in a mouse model of Alzheimer's disease pathology.
Journal: Cell reports
In common: export_fig, OpenCV, scikit-image, 5 other tools, mouse

Contribute

The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.

Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.

Request its removal

To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).

Discussion, reproductions, activity

Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.

Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.

Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.